From ff7f0faf3daab460d7e38f59b17e7cac86e6734d Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 18 May 2026 15:10:43 +0200 Subject: [PATCH 001/250] alignment basically working --- tests/integration/alignment/__init__.py | 0 tests/integration/alignment/profile_fit.py | 138 ++ .../alignment/run_random_pdb_fit.py | 242 ++ tests/integration/alignment/submit_sweep.sh | 35 + tests/integration/alignment/sweep_slurm.sh | 27 + .../integration/alignment/test_fit_to_data.py | 112 + .../alignment/test_fit_to_data_translation.py | 111 + .../alignment/test_pipeline_recovery.py | 133 ++ .../alignment/test_real_data_mr.py | 170 ++ .../alignment/test_rotation_recovery.py | 209 ++ tests/unit/alignment/__init__.py | 0 .../alignment/test_patterson_translation.py | 123 + tests/unit/alignment/test_sh.py | 172 ++ .../unit/alignment/test_variance_weighting.py | 137 ++ tests/unit/alignment/test_wigner.py | 159 ++ tests/unit/model/test_rotate.py | 96 + tests/unit/model/test_spacegroup_setter.py | 127 ++ tests/unit/model/test_translate.py | 74 + torchref/alignment/__init__.py | 191 +- torchref/alignment/align.py | 727 ++++++ torchref/alignment/ball_search.py | 591 +++++ torchref/alignment/ball_transform.py | 2027 ----------------- torchref/alignment/jax_subpixel_peaks.py | 67 - torchref/alignment/lattman_love.py | 219 ++ torchref/alignment/ml_rotation.py | 434 ++++ torchref/alignment/patterson_filter.py | 163 ++ torchref/alignment/pipeline.py | 44 +- torchref/alignment/rigid_body.py | 36 +- torchref/alignment/sh.py | 520 +++++ torchref/alignment/translation.py | 694 ++++++ torchref/alignment/wigner.py | 371 +++ torchref/model/model.py | 184 +- torchref/model/model_ft.py | 28 +- torchref/model/sf_fft.py | 78 +- torchref/scaling/solvent.py | 8 +- 35 files changed, 6150 insertions(+), 2297 deletions(-) create mode 100644 tests/integration/alignment/__init__.py create mode 100644 tests/integration/alignment/profile_fit.py create mode 100644 tests/integration/alignment/run_random_pdb_fit.py create mode 100644 tests/integration/alignment/submit_sweep.sh create mode 100644 tests/integration/alignment/sweep_slurm.sh create mode 100644 tests/integration/alignment/test_fit_to_data.py create mode 100644 tests/integration/alignment/test_fit_to_data_translation.py create mode 100644 tests/integration/alignment/test_pipeline_recovery.py create mode 100644 tests/integration/alignment/test_real_data_mr.py create mode 100644 tests/integration/alignment/test_rotation_recovery.py create mode 100644 tests/unit/alignment/__init__.py create mode 100644 tests/unit/alignment/test_patterson_translation.py create mode 100644 tests/unit/alignment/test_sh.py create mode 100644 tests/unit/alignment/test_variance_weighting.py create mode 100644 tests/unit/alignment/test_wigner.py create mode 100644 tests/unit/model/test_rotate.py create mode 100644 tests/unit/model/test_spacegroup_setter.py create mode 100644 tests/unit/model/test_translate.py create mode 100644 torchref/alignment/align.py create mode 100644 torchref/alignment/ball_search.py delete mode 100644 torchref/alignment/ball_transform.py delete mode 100644 torchref/alignment/jax_subpixel_peaks.py create mode 100644 torchref/alignment/lattman_love.py create mode 100644 torchref/alignment/ml_rotation.py create mode 100644 torchref/alignment/patterson_filter.py create mode 100644 torchref/alignment/sh.py create mode 100644 torchref/alignment/wigner.py diff --git a/tests/integration/alignment/__init__.py b/tests/integration/alignment/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/integration/alignment/profile_fit.py b/tests/integration/alignment/profile_fit.py new file mode 100644 index 00000000..0d29521e --- /dev/null +++ b/tests/integration/alignment/profile_fit.py @@ -0,0 +1,138 @@ +#!/usr/bin/env python +""" +Profile `ModelFT.fit_to_data` to find where the time goes. + +Run from the repo root: + .venv/bin/python tests/integration/alignment/profile_fit.py [--pdb 1DAW] \ + [--n-rotation-candidates 3] [--n-translation-candidates 3] \ + [--translation-grid-steps 16] + +Output: top-50 cumulative-time entries from cProfile + a custom per-stage timer +breakdown (rotation search, ML rescore, TF, local refine, joint refine, +final Scaler refit). +""" +from __future__ import annotations + +import argparse +import cProfile +import pstats +import time +from contextlib import contextmanager +from pathlib import Path + +import torch + +from torchref.alignment.ball_search import rotation_matrix_from_edmonds_euler +from torchref.io.datasets.reflection_data import ReflectionData +from torchref.model import ModelFT +from torchref.symmetry import SpaceGroup + + +TEST_FILES = Path("/das/work/p17/p17490/Peter/Library/work_trees_torchref/fix_alignment/tests/files") + +PAIRS = { + "1DAW": (TEST_FILES / "pdb" / "1DAW.pdb", TEST_FILES / "mtz" / "1DAW.mtz"), + "1AK5": (TEST_FILES / "pdb" / "1AK5_with_H.pdb", TEST_FILES / "mtz" / "1AK5.mtz"), + "3A5V": (TEST_FILES / "pdb" / "3A5V.pdb", TEST_FILES / "mtz" / "3A5V.mtz"), +} + + +_TIMINGS: dict[str, float] = {} + + +@contextmanager +def _stage(name: str): + t0 = time.time() + try: + yield + finally: + _TIMINGS[name] = _TIMINGS.get(name, 0.0) + (time.time() - t0) + print(f" [{name}] {time.time()-t0:.2f}s", flush=True) + + +def _patch_for_timing(): + """Wrap key fit_to_data stages so we get an inline breakdown.""" + from torchref.alignment import ball_search, ml_rotation, translation + from torchref.alignment import lattman_love, rigid_body + from torchref import scaling + + originals = {} + + def wrap(module, attr, label): + original = getattr(module, attr) + originals[(module, attr)] = original + + def wrapper(*args, **kwargs): + with _stage(label): + return original(*args, **kwargs) + + setattr(module, attr, wrapper) + + wrap(ball_search, "ball_rotation_search", "ball_rotation_search") + wrap(ml_rotation, "sim_mlrf_rescore", "sim_mlrf_rescore") + wrap(translation, "amplitude_translation_search", "amplitude_translation_search") + wrap(translation, "local_translation_refine", "local_translation_refine") + wrap(translation, "precompute_G_for_rotation", "precompute_G_for_rotation") + wrap(lattman_love, "LattmanLoveInterpolator", "LL_interp_build") + wrap(rigid_body, "RigidBodyRefinement", "RigidBodyRefinement_init") + # Scaler is heavy; track its calls + return originals + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--pdb", default="1DAW", choices=sorted(PAIRS.keys())) + ap.add_argument("--n-rotation-candidates", type=int, default=3) + ap.add_argument("--n-translation-candidates", type=int, default=3) + ap.add_argument("--translation-grid-steps", type=int, default=16) + ap.add_argument("--top", type=int, default=40, help="top N cProfile entries") + args = ap.parse_args() + + pdb_path, mtz_path = PAIRS[args.pdb] + print(f"=== Profiling fit_to_data on {args.pdb} ===", flush=True) + + data = ReflectionData().load_mtz(str(mtz_path)) + canonical = ModelFT().load_pdb(str(pdb_path)) + canonical.spacegroup = SpaceGroup("P 1") + R_true = rotation_matrix_from_edmonds_euler(0.6, 0.4, 1.2) + rotated_p = canonical.rotate( + R_true.to(canonical.dtype_float), center=canonical.xyz().mean(dim=0), + ) + t_true = torch.tensor([0.18, -0.07, 0.23], dtype=canonical.dtype_float) + perturbed = rotated_p.translate(t_true, fractional=True) + print(f" spacegroup={data.spacegroup}, n_atoms={canonical.xyz().shape[0]}, " + f"n_hkl={data.hkl.shape[0]}", flush=True) + + _patch_for_timing() + + profiler = cProfile.Profile() + t0 = time.time() + profiler.enable() + aligned = perturbed.fit_to_data( + data, + n_rotation_candidates=args.n_rotation_candidates, + n_translation_candidates=args.n_translation_candidates, + translation_grid_steps=args.translation_grid_steps, + verbose=0, + ) + profiler.disable() + total = time.time() - t0 + + print(f"\n=== Stage breakdown (total {total:.2f}s) ===", flush=True) + other = total - sum(_TIMINGS.values()) + for name, t in sorted(_TIMINGS.items(), key=lambda kv: -kv[1]): + print(f" {name:40s} {t:8.2f}s ({100*t/total:5.1f}%)", flush=True) + print(f" {'(unattributed)':40s} {other:8.2f}s ({100*other/total:5.1f}%)", + flush=True) + + print(f"\n=== Top-{args.top} cProfile (cumulative time) ===", flush=True) + stats = pstats.Stats(profiler).sort_stats("cumulative") + stats.print_stats(args.top) + + print(f"\n=== Top-{args.top} cProfile (own time) ===", flush=True) + stats = pstats.Stats(profiler).sort_stats("tottime") + stats.print_stats(args.top) + + +if __name__ == "__main__": + main() diff --git a/tests/integration/alignment/run_random_pdb_fit.py b/tests/integration/alignment/run_random_pdb_fit.py new file mode 100644 index 00000000..56af43f8 --- /dev/null +++ b/tests/integration/alignment/run_random_pdb_fit.py @@ -0,0 +1,242 @@ +#!/usr/bin/env python +""" +End-to-end demo / sanity check for the alignment pipeline. + +Flow: + 1. Pick a random PDB / MTZ pair from `tests/files`. + 2. Load model + data (real cell, real F_obs). + 3. Apply a random rotation to the model atoms. + 4. Run `ModelFT.fit_to_data` to recover an aligned orientation. + 5. Fit an anisotropic Scaler against the data. + 6. Report R-work / R-free before and after. + +Run as a script (NOT a pytest test) — it's a one-off integration probe: + + cd /das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/fix_alignment + python tests/integration/alignment/run_random_pdb_fit.py [--seed N] [--pdb 1DAW] +""" + + +from __future__ import annotations + +import argparse +import math +import random +import time +from pathlib import Path + +import torch + +from torchref.alignment.ball_search import rotation_angular_distance_deg +from torchref.io.datasets.reflection_data import ReflectionData +from torchref.model import ModelFT +from torchref.scaling import Scaler + + +TEST_FILES =Path('/das/work/p17/p17490/Peter/Library/work_trees_torchref/fix_alignment/tests/files') + +PAIRS = { + # PDB stem → (pdb_path, mtz_path). Some PDBs use a non-standard filename. + "1AK5": (TEST_FILES / "pdb" / "1AK5_with_H.pdb", TEST_FILES / "mtz" / "1AK5.mtz"), + "1DAW": (TEST_FILES / "pdb" / "1DAW.pdb", TEST_FILES / "mtz" / "1DAW.mtz"), + "2DQ6": (TEST_FILES / "pdb" / "2DQ6.pdb", TEST_FILES / "mtz" / "2DQ6.mtz"), + "3A5V": (TEST_FILES / "pdb" / "3A5V.pdb", TEST_FILES / "mtz" / "3A5V.mtz"), + "3E98": (TEST_FILES / "pdb" / "3E98.pdb", TEST_FILES / "mtz" / "3E98.mtz"), + "3GR5": (TEST_FILES / "pdb" / "3GR5.pdb", TEST_FILES / "mtz" / "3GR5.mtz"), + "3K7M": (TEST_FILES / "pdb" / "3K7M.pdb", TEST_FILES / "mtz" / "3K7M.mtz"), + "3VRJ": (TEST_FILES / "pdb" / "3VRJ.pdb", TEST_FILES / "mtz" / "3VRJ.mtz"), + "4BX9": (TEST_FILES / "pdb" / "4BX9.pdb", TEST_FILES / "mtz" / "4BX9.mtz"), + "5BOV": (TEST_FILES / "pdb" / "5BOV.pdb", TEST_FILES / "mtz" / "5BOV.mtz"), + "6G9X": (TEST_FILES / "pdb" / "6G9X.pdb", TEST_FILES / "mtz" / "6G9X.mtz"), +} + + +def _random_rotation(seed: int) -> torch.Tensor: + """Uniform random rotation on SO(3) via QR of a Gaussian matrix.""" + g = torch.Generator().manual_seed(int(seed)) + A = torch.randn(3, 3, generator=g, dtype=torch.float64) + Q, R = torch.linalg.qr(A) + Q = Q @ torch.diag(torch.sign(torch.diag(R))) + if torch.det(Q) < 0: + Q[:, 0] = -Q[:, 0] + return Q + + +def _min_err_over_sym(R_test: torch.Tensor, R_ref: torch.Tensor, + sym_mats: torch.Tensor) -> float: + """Minimum angular distance of R_test to any S·R_ref over sym_mats.""" + best = float("inf") + R_test = R_test.to(torch.float64) + R_ref = R_ref.to(torch.float64) + for k in range(sym_mats.shape[0]): + e = rotation_angular_distance_deg(R_test, sym_mats[k] @ R_ref) + if e < best: + best = e + return best + + +def _kabsch_rotation(xyz_a: torch.Tensor, xyz_b: torch.Tensor) -> torch.Tensor: + """Return R minimising ||xyz_a - xyz_b @ R^T|| (both centred).""" + a = (xyz_a.detach() - xyz_a.detach().mean(0)).to(torch.float64) + b = (xyz_b.detach() - xyz_b.detach().mean(0)).to(torch.float64) + H = b.T @ a + U, _, Vt = torch.linalg.svd(H) + d = float(torch.sign(torch.det(Vt.T @ U.T))) + D = torch.diag(torch.tensor([1.0, 1.0, d], dtype=H.dtype)) + return Vt.T @ D @ U.T + + +def run(pdb_key: str, seed: int, verbose: int = 1, + device: torch.device = torch.device("cpu")) -> dict: + pdb_path, mtz_path = PAIRS[pdb_key] + print(f"\n=== {pdb_key}: {pdb_path.name} + {mtz_path.name} ===", flush=True) + + # 1. Construct model + data directly on `device` so the alignment + # pipeline runs end-to-end on GPU without re-creating the SfFFT. + # Post-hoc `.to(device)` reassigns Model.cell via the setter, which + # triggers `_maybe_initialize_fft` and drops the already-built grid / + # map_symmetry state. + t0 = time.time() + model = ModelFT(device=device).load_pdb(str(pdb_path)) + data = ReflectionData(device=str(device)).load_mtz(str(mtz_path)) + sym_mats = data.spacegroup.matrices.to(torch.float64) + print(f" spacegroup: {data.spacegroup} cell: " + f"a={data.cell.a:.1f}, b={data.cell.b:.1f}, c={data.cell.c:.1f} Å " + f"atoms: {model.xyz().shape[0]} ({time.time()-t0:.1f}s)", flush=True) + + def _scale_and_r(m: ModelFT) -> tuple[float, float]: + s = Scaler(model=m, data=data, nbins=20, verbose=0, + device=m.xyz().device) + fcalc = m(data.hkl) + s.initialize(fcalc) + s.refine_lbfgs(fcalc=fcalc) + rw, rf = s.rfactor(fcalc) + # rfactor may return Python floats or 0-d tensors depending on backend + rw = rw.item() if hasattr(rw, "item") else float(rw) + rf = rf.item() if hasattr(rf, "item") else float(rf) + return rw, rf + + # Reference R-factor of the un-rotated model (the optimal we could hope for). + rwork_ref, rfree_ref = _scale_and_r(model) + print(f" reference R-work (un-rotated model): {rwork_ref:.4f} " + f"R-free: {rfree_ref:.4f}", flush=True) + + # 2. Apply random rotation to atom coords. + R_true = _random_rotation(seed) + xyz_canonical = model.xyz().clone() + centroid = xyz_canonical.mean(0) + rotated_search = model.rotate( + R_true.to(model.dtype_float).to(device), center=centroid, + ) + + # R-factor of the rotated search model (should be ~50% — random). + rwork_pre, rfree_pre = _scale_and_r(rotated_search) + print(f" rotated search R-work (no alignment): {rwork_pre:.4f} " + f"R-free: {rfree_pre:.4f} (should be ~0.5)", flush=True) + + # 3. Run fit_to_data: recover the alignment. + t1 = time.time() + aligned = rotated_search.fit_to_data( + data, + d_min=4.0, d_max=15.0, + L=32, n_shells=20, + n_rotation_peaks=200, n_ml_refine=200, + verbose=verbose, + ) + fit_time = time.time() - t1 + print(f" fit_to_data took {fit_time:.1f}s", flush=True) + + # Effective rotation between aligned coords and the canonical reference. + R_residual = _kabsch_rotation(aligned.xyz(), xyz_canonical) + err_to_canonical = min( + rotation_angular_distance_deg(R_residual.to(torch.float64), sym_mats[k]) + for k in range(sym_mats.shape[0]) + ) + print(f" aligned-vs-canonical angular distance " + f"(mod {data.spacegroup}-symmetry): {err_to_canonical:.2f}°", flush=True) + + # 4. Scale the aligned model and report R-factor. + rwork_post, rfree_post = _scale_and_r(aligned) + print(f" aligned R-work: {rwork_post:.4f} R-free: {rfree_post:.4f}", flush=True) + + return { + "pdb": pdb_key, + "spacegroup": str(data.spacegroup), + "ref_rwork": rwork_ref, + "ref_rfree": rfree_ref, + "pre_rwork": rwork_pre, + "pre_rfree": rfree_pre, + "post_rwork": rwork_post, + "post_rfree": rfree_post, + "err_canonical_deg": err_to_canonical, + "fit_time_s": fit_time, + } + + +def main(): + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--seed", type=int, default=None, + help="Random seed (rotation + PDB pick). Default: time-based.") + ap.add_argument("--pdb", default=None, choices=sorted(PAIRS.keys()), + help="PDB key to use. Default: random.") + ap.add_argument("--n-trials", type=int, default=1, + help="Number of random trials with different rotations / PDBs.") + ap.add_argument("--sweep", action="store_true", + help="Iterate over every PDB in PAIRS, n-trials per PDB.") + ap.add_argument("--verbose", type=int, default=0, + help="Verbosity passed to fit_to_data.") + ap.add_argument("--device", default="cpu", choices=["cpu", "cuda"], + help="Run the alignment on this device.") + args = ap.parse_args() + device = torch.device(args.device) + if device.type == "cuda" and not torch.cuda.is_available(): + raise SystemExit("--device cuda requested but torch.cuda not available") + + if args.seed is None: + args.seed = int(time.time()) + rng = random.Random(args.seed) + print(f"seed = {args.seed}", flush=True) + + # Build the (pdb, seed) work list. + if args.sweep: + worklist = [(pdb, rng.randint(0, 10 ** 9)) + for pdb in sorted(PAIRS.keys()) + for _ in range(args.n_trials)] + else: + worklist = [] + for _ in range(args.n_trials): + pdb = args.pdb if args.pdb is not None else rng.choice(list(PAIRS.keys())) + worklist.append((pdb, rng.randint(0, 10 ** 9))) + + results = [] + for pdb_key, trial_seed in worklist: + try: + r = run(pdb_key, trial_seed, verbose=args.verbose, device=device) + results.append(r) + except Exception as exc: + print(f" TRIAL FAILED on {pdb_key}: {exc!r}", flush=True) + results.append({"pdb": pdb_key, "error": repr(exc)}) + finally: + # Release CUDA allocator caches between trials. Without this a + # failed trial leaves its ~tens-of-GB residue in the allocator + # pool and starves every subsequent trial of memory. + if device.type == "cuda": + import gc + gc.collect() + torch.cuda.empty_cache() + + print("\n=== summary ===", flush=True) + print(f"{'pdb':>6} {'sg':>8} {'rwork_ref':>10} {'rwork_pre':>10} " + f"{'rwork_post':>10} {'err_deg':>8} {'time_s':>7}", flush=True) + for r in results: + if "error" in r: + print(f"{r['pdb']:>6} FAILED: {r['error']}", flush=True) + continue + print(f"{r['pdb']:>6} {r['spacegroup'][12:20]:>8} " + f"{r['ref_rwork']:>10.4f} {r['pre_rwork']:>10.4f} " + f"{r['post_rwork']:>10.4f} {r['err_canonical_deg']:>8.2f} " + f"{r['fit_time_s']:>7.1f}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/tests/integration/alignment/submit_sweep.sh b/tests/integration/alignment/submit_sweep.sh new file mode 100644 index 00000000..8a2c2f27 --- /dev/null +++ b/tests/integration/alignment/submit_sweep.sh @@ -0,0 +1,35 @@ +#!/bin/bash +# Submit one SLURM job per (PDB × trial) — 11 PDBs × 3 trials = 33 jobs. +# +# Usage: ./submit_sweep.sh [N_TRIALS_PER_PDB] (default 3) + +set -euo pipefail +N=${1:-3} + +PDBS=(1AK5 1DAW 2DQ6 3A5V 3E98 3GR5 3K7M 3VRJ 4BX9 5BOV 6G9X) +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +SUBMIT_DIR="$SCRIPT_DIR" +LOG_DIR="$(pwd)/sweep_logs_$(date +%Y%m%d_%H%M%S)" +mkdir -p "$LOG_DIR" +cd "$LOG_DIR" + +# Pick a base seed (fixed for reproducibility across the sweep; per-trial +# seed is base + trial_idx so each (pdb, trial) gets a distinct seed). +BASE_SEED=${BASE_SEED:-42} + +echo "Submitting $((${#PDBS[@]} * N)) jobs into $LOG_DIR" >&2 +JOBIDS=() +for pdb in "${PDBS[@]}"; do + for trial in $(seq 0 $((N - 1))); do + seed=$((BASE_SEED + trial * 1000003)) + jobid=$(sbatch --parsable --job-name="fit_${pdb}_t${trial}" \ + "$SUBMIT_DIR/sweep_slurm.sh" "$pdb" "$seed") + echo " $pdb trial $trial → seed=$seed jobid=$jobid" + JOBIDS+=("$jobid") + done +done + +echo +echo "Submitted ${#JOBIDS[@]} jobs. Track with:" +echo " squeue --user \$USER --jobs=$(IFS=,; echo "${JOBIDS[*]}")" +echo "Logs in: $LOG_DIR" diff --git a/tests/integration/alignment/sweep_slurm.sh b/tests/integration/alignment/sweep_slurm.sh new file mode 100644 index 00000000..4e375460 --- /dev/null +++ b/tests/integration/alignment/sweep_slurm.sh @@ -0,0 +1,27 @@ +#!/bin/bash +#SBATCH --job-name=fit_to_data +#SBATCH --output=slurm-%j.out +#SBATCH --ntasks=1 +#SBATCH --cpus-per-task=16 +#SBATCH --mem=64G +#SBATCH --time=03:00:00 +#SBATCH --partition=day + +# Usage: sbatch sweep_slurm.sh +# E.g. sbatch sweep_slurm.sh 1AK5 12345 + +set -euo pipefail + +PDB_KEY="${1:?need PDB key}" +SEED="${2:?need seed}" + +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/fix_alignment +PYTHON=/das/work/units/LBR-FEL/p17490/CONDA/torchref/bin/python + +cd "$REPO" +export PYTHONPATH="$REPO" +export TORCHREF_NUM_THREADS="${SLURM_CPUS_PER_TASK:-8}" + +echo "Running fit on PDB=$PDB_KEY seed=$SEED on $(hostname)" +$PYTHON tests/integration/alignment/run_random_pdb_fit.py \ + --pdb "$PDB_KEY" --seed "$SEED" --n-trials 1 --verbose 1 diff --git a/tests/integration/alignment/test_fit_to_data.py b/tests/integration/alignment/test_fit_to_data.py new file mode 100644 index 00000000..b775140c --- /dev/null +++ b/tests/integration/alignment/test_fit_to_data.py @@ -0,0 +1,112 @@ +""" +Integration test for `ModelFT.fit_to_data`: end-to-end Patterson rotation +search + Sim MLRF rescoring, returning a re-oriented ModelFT. + +Setup: +- F_obs: real 1DAW.mtz at its native C2 spacegroup. +- Search model: P1 copy of 1DAW.pdb whose atomic coordinates have been + rotated by a random R_true. + +Acceptance: after `fit_to_data`, the returned model's atom coordinates are +within 8° rotation distance (modulo C2 symmetry of F_obs) of the un-rotated +canonical orientation, for 5/5 random trials. +""" +import math +from pathlib import Path + +import pytest +import torch + +from torchref.alignment.ball_search import rotation_angular_distance_deg +from torchref.io.datasets.reflection_data import ReflectionData +from torchref.model import ModelFT +from torchref.symmetry import SpaceGroup + + +TEST_FILES = Path(__file__).resolve().parents[2] / "files" +PDB_1DAW = TEST_FILES / "pdb" / "1DAW.pdb" +MTZ_1DAW = TEST_FILES / "mtz" / "1DAW.mtz" + + +def _load_p1_search_model() -> ModelFT: + """Load 1DAW and force spacegroup to P1 via the proper setter.""" + m = ModelFT().load_pdb(str(PDB_1DAW)) + m.spacegroup = SpaceGroup("P 1") + return m + + +@pytest.fixture(scope="module") +def real_setup(): + """Real C2 F_obs + P1 search model factory.""" + data = ReflectionData().load_mtz(str(MTZ_1DAW)) + return data, _load_p1_search_model + + +def _random_rotation(seed: int) -> torch.Tensor: + g = torch.Generator().manual_seed(seed) + A = torch.randn(3, 3, generator=g, dtype=torch.float64) + Q, R = torch.linalg.qr(A) + Q = Q @ torch.diag(torch.sign(torch.diag(R))) + if torch.det(Q) < 0: + Q[:, 0] = -Q[:, 0] + return Q + + +def _best_alignment_rotation(xyz_a: torch.Tensor, xyz_b: torch.Tensor) -> torch.Tensor: + """ + Kabsch: return R that minimises ||xyz_a - xyz_b @ R.T|| (with both centred). + Used here to recover the effective rotation between two atom sets. + """ + a = xyz_a - xyz_a.mean(0) + b = xyz_b - xyz_b.mean(0) + H = b.T @ a + U, _, Vt = torch.linalg.svd(H) + d = torch.sign(torch.det(Vt.T @ U.T)) + D = torch.diag(torch.tensor([1.0, 1.0, d], dtype=H.dtype)) + R = Vt.T @ D @ U.T + return R + + +@pytest.mark.integration +@pytest.mark.slow +@pytest.mark.parametrize("trial", range(5)) +def test_fit_to_data_real_1daw(real_setup, trial): + """ + Apply a random R_true to a P1 search model, call `fit_to_data(real_F_obs)`, + and verify the returned model is within 8° rotation distance of the + canonical orientation (modulo C2 symmetry of F_obs). + """ + data, make_model = real_setup + sym_mats = data.spacegroup.matrices.to(torch.float64) + + canonical = make_model() + xyz_canonical = canonical.xyz().clone() + centroid = xyz_canonical.mean(0) + + R_true = _random_rotation(seed=5000 + trial) + search = canonical.rotate(R_true.to(canonical.dtype_float), center=centroid) + + aligned = search.fit_to_data( + data, + d_min=4.0, d_max=15.0, + L=32, n_shells=20, + n_rotation_peaks=200, n_ml_refine=200, + do_translation=False, # this test only checks rotation accuracy + verbose=0, + ) + + # The effective rotation between aligned.xyz() and xyz_canonical should + # be a C2 symmetry operator (i.e. nearly identity or 2-fold along b). + R_residual = _best_alignment_rotation( + aligned.xyz().to(torch.float64), xyz_canonical.to(torch.float64), + ) + # Compare R_residual to identity / each C2 op. + best_err = float("inf") + for k in range(sym_mats.shape[0]): + err = rotation_angular_distance_deg(R_residual, sym_mats[k]) + if err < best_err: + best_err = err + assert best_err < 8.0, ( + f"trial {trial}: aligned model is {best_err:.2f}° from a C2-equivalent " + f"canonical orientation" + ) diff --git a/tests/integration/alignment/test_fit_to_data_translation.py b/tests/integration/alignment/test_fit_to_data_translation.py new file mode 100644 index 00000000..ee14f87a --- /dev/null +++ b/tests/integration/alignment/test_fit_to_data_translation.py @@ -0,0 +1,111 @@ +""" +Integration test for the new translation + joint R+t refinement in +`ModelFT.fit_to_data` (Phase 3 component B). + +Setup: 1DAW.mtz (C2) F_obs + P1 search model. Apply a small known rotation and +fractional translation, then ask `fit_to_data` to recover both. Acceptance: +recovered `(R_residual, t_residual)` brings the model close to canonical +(within 8° rotation modulo C2 symmetry and within 0.1 fractional shift along +any axis modulo unit cell), and the post-refinement R-work drops well below +the pre-fit value. +""" +import math +from pathlib import Path + +import pytest +import torch + +from torchref.alignment.ball_search import ( + rotation_angular_distance_deg, + rotation_matrix_from_edmonds_euler, +) +from torchref.io.datasets.reflection_data import ReflectionData +from torchref.model import ModelFT +from torchref.scaling import Scaler +from torchref.symmetry import SpaceGroup + + +TEST_FILES = Path(__file__).resolve().parents[2] / "files" +PDB_1DAW = TEST_FILES / "pdb" / "1DAW.pdb" +MTZ_1DAW = TEST_FILES / "mtz" / "1DAW.mtz" + + +def _scale_and_rwork(model: ModelFT, data) -> float: + scaler = Scaler(model=model, data=data, nbins=20, verbose=0) + fcalc = model(data.hkl) + scaler.initialize(fcalc) + scaler.refine_lbfgs(fcalc=fcalc) + rw, _ = scaler.rfactor(fcalc) + return rw.item() if hasattr(rw, "item") else float(rw) + + +def _wrap_frac(t: torch.Tensor) -> torch.Tensor: + """Wrap fractional coords into [-0.5, 0.5).""" + return (t + 0.5) % 1.0 - 0.5 + + +@pytest.mark.integration +@pytest.mark.slow +def test_fit_to_data_recovers_rotation_and_translation(): + data = ReflectionData().load_mtz(str(MTZ_1DAW)) + canonical = ModelFT().load_pdb(str(PDB_1DAW)) + canonical.spacegroup = SpaceGroup("P 1") + + # Apply a known random rotation + fractional translation. + R_true = rotation_matrix_from_edmonds_euler(0.6, 0.4, 1.2) + R_apply = R_true.to(canonical.dtype_float) + rotated = canonical.rotate(R_apply, center=canonical.xyz().mean(dim=0)) + t_frac_true = torch.tensor([0.18, -0.07, 0.23], dtype=canonical.dtype_float) + perturbed = rotated.translate(t_frac_true, fractional=True) + + rwork_pre = _scale_and_rwork(perturbed, data) + + aligned = perturbed.fit_to_data( + data, + d_min=4.0, d_max=15.0, + L=32, n_shells=20, + n_rotation_peaks=200, n_ml_refine=200, + do_translation=True, + do_joint_refine=True, + verbose=0, + ) + + # Recovered rotation (modulo C2). Compare canonical vs aligned via centroid. + xyz_canon = canonical.xyz().to(torch.float64) + xyz_aligned = aligned.xyz().to(torch.float64) + c_canon = xyz_canon.mean(dim=0) + c_aligned = xyz_aligned.mean(dim=0) + a = xyz_canon - c_canon + b = xyz_aligned - c_aligned + H = b.T @ a + U, _, Vt = torch.linalg.svd(H) + d = float(torch.sign(torch.det(Vt.T @ U.T))) + D = torch.diag(torch.tensor([1.0, 1.0, d], dtype=H.dtype)) + R_residual = Vt.T @ D @ U.T + sym_mats = data.spacegroup.matrices.to(torch.float64) + best_rot_err = min( + rotation_angular_distance_deg(R_residual, sym_mats[k]) + for k in range(sym_mats.shape[0]) + ) + assert best_rot_err < 8.0, ( + f"residual rotation {best_rot_err:.2f}° > 8° gate" + ) + + # We do not pin down the recovered translation directly — for spacegroups + # with polar / non-unique origins (e.g. C2's free origin along y) the + # recovered translation may differ from the applied one by an allowed + # origin shift. The crystallographic test that this is a valid solution is + # the R-factor of the scaled model. + # The translation function brings R-work close to the canonical-native + # reference (0.21 for 1DAW). The residual gap (~0.12) is from the + # rotation function's ~2° angular error — a separate refinement that's + # not part of the translation function. A 1.87° rotation residual on a + # 100Å molecule moves atoms by ~3Å, which costs ~0.12 in R-work even + # with the exactly correct translation. + rwork_post = _scale_and_rwork(aligned, data) + canonical_native = ModelFT().load_pdb(str(PDB_1DAW)) # native C2 + rwork_ref = _scale_and_rwork(canonical_native, data) + assert rwork_post < rwork_ref + 0.18, ( + f"R-work {rwork_post:.4f} > reference {rwork_ref:.4f} + 0.18 " + f"(pre-fit was {rwork_pre:.4f})" + ) diff --git a/tests/integration/alignment/test_pipeline_recovery.py b/tests/integration/alignment/test_pipeline_recovery.py new file mode 100644 index 00000000..82e8da8a --- /dev/null +++ b/tests/integration/alignment/test_pipeline_recovery.py @@ -0,0 +1,133 @@ +""" +End-to-end pipeline recovery test. + +Setup mirrors the rotation-recovery test: instead of rotating the *atoms* of a +symmetric crystal (which couples to the C2 symmetry expansion in ModelFT and +breaks the simple "rotated model" assumption), we exercise the pipeline by +providing a known true rotation and verifying that the rotation-search + +clustering produces a candidate within tolerance. + +This test is intentionally narrower than `test_rotation_recovery.py` — it +covers the ball-search → cluster pipeline integration and exercises the +public `cluster_rotation_peaks` / `rotation_angular_distance` helpers, but +does NOT exercise the full `MolecularReplacementPipeline.run()`, which would +require a P1 search-model setup with synthetic F_obs and is out of scope for +this gate test. + +The full pipeline gate (rotation + translation + rigid body in P1) is covered +by a follow-up test once the underlying rotation function is fully proven on +real data (this test plus `test_rotation_recovery.py`). +""" +import math +from pathlib import Path + +import numpy as np +import pytest +import torch + +from torchref.alignment.ball_search import ( + ball_rotation_search, + rotation_matrix_from_edmonds_euler, + rotation_angular_distance_deg, +) +from torchref.alignment.pipeline import cluster_rotation_peaks +from torchref.io.datasets.reflection_data import ReflectionData +from torchref.model import ModelFT + + +TEST_FILES = Path(__file__).resolve().parents[2] / "files" +PDB_1DAW = TEST_FILES / "pdb" / "1DAW.pdb" +MTZ_1DAW = TEST_FILES / "mtz" / "1DAW.mtz" + + +@pytest.fixture(scope="module") +def model_1daw(): + return ModelFT().load_pdb(str(PDB_1DAW)) + + +@pytest.fixture(scope="module") +def data_p1(): + return ReflectionData().load_mtz(str(MTZ_1DAW)).expand_to_p1(include_friedel=False) + + +def _random_rotation(seed: int) -> torch.Tensor: + g = torch.Generator().manual_seed(seed) + A = torch.randn(3, 3, generator=g, dtype=torch.float64) + Q, R = torch.linalg.qr(A) + Q = Q @ torch.diag(torch.sign(torch.diag(R))) + if torch.det(Q) < 0: + Q[:, 0] = -Q[:, 0] + return Q + + +def _hkl_to_s(hkl, cell): + rec_basis = cell.reciprocal_basis_matrix + if callable(rec_basis): + rec_basis = rec_basis() + return hkl.to(torch.float64) @ rec_basis.to(torch.float64) + + +def _normalize_by_shell(F, s_mag, P): + sorted_idx = torch.argsort(s_mag) + shell_idx = torch.zeros(s_mag.shape[0], dtype=torch.int64) + chunk = s_mag.shape[0] // P + for k in range(P): + a = k * chunk + b = (k + 1) * chunk if k < P - 1 else s_mag.shape[0] + shell_idx[sorted_idx[a:b]] = k + norm = torch.zeros_like(s_mag, dtype=torch.float64) + for k in range(P): + mask = shell_idx == k + norm[mask] = (F[mask] ** 2).mean().clamp(min=1e-30).sqrt() + return F / norm + + +@pytest.mark.integration +@pytest.mark.parametrize("trial", range(5)) +def test_pipeline_clustering_preserves_truth(model_1daw, data_p1, trial): + """ + After clustering, the true rotation must still be represented in the + top-3 clusters. This guards the `cluster_rotation_peaks` step from + rejecting the correct peak as a duplicate of a higher-scoring artifact. + """ + R_true = _random_rotation(seed=3000 + trial) + + with torch.no_grad(): + F_orig = model_1daw(data_p1.hkl).abs().to(torch.float64) + s_obs = _hkl_to_s(data_p1.hkl, model_1daw.cell) + s_mag = s_obs.norm(dim=-1) + keep = (s_mag >= 1.0 / 15.0) & (s_mag <= 1.0 / 4.0) + s_obs = s_obs[keep] + F_orig = F_orig[keep] + s_mag = s_mag[keep] + + e_obs = _normalize_by_shell(F_orig, s_mag, P=20) + s_calc = s_obs @ R_true.to(torch.float64).T + e_calc = e_obs + + _C, _a, _b, _g, peaks = ball_rotation_search( + s_obs, e_obs, s_calc, e_calc, + L=32, P=20, n_peaks=60, refine_subvoxel=True, n_refine=20, + sigma_threshold=0.0, + ) + peak_tuples = [(p.alpha, p.beta, p.gamma, p.score, p.sigma) for p in peaks] + clustered = cluster_rotation_peaks(peak_tuples, threshold_deg=8.0) + assert len(clustered) > 0 + + # The true R (mod point-group symmetry of the obs field) must be within + # 8° of one of the top-5 clustered peaks. + sym = data_p1.spacegroup.matrices.to(torch.float64) + best_err = float("inf") + best_rank = None + for rank, peak in enumerate(clustered[:5]): + R_p = rotation_matrix_from_edmonds_euler(peak[0], peak[1], peak[2]).to(torch.float64) + for k in range(sym.shape[0]): + R_eq = sym[k] @ R_true.to(torch.float64) + err = rotation_angular_distance_deg(R_p, R_eq) + if err < best_err: + best_err = err + best_rank = rank + assert best_err < 8.0, ( + f"trial {trial}: best clustered peak {best_err:.2f}° from R_true " + f"(rank {best_rank} of {len(clustered)})" + ) diff --git a/tests/integration/alignment/test_real_data_mr.py b/tests/integration/alignment/test_real_data_mr.py new file mode 100644 index 00000000..1ffc6c8e --- /dev/null +++ b/tests/integration/alignment/test_real_data_mr.py @@ -0,0 +1,170 @@ +""" +Phase-2 real-data rotation recovery test. + +Setup: +- F_obs: real measured data from 1DAW.mtz (C2 spacegroup, intermolecular + Patterson contributions present). +- Search model: P1 copy of 1DAW.pdb whose atomic coordinates have been + rotated by a random R_true. This simulates "user has a search model in + some arbitrary orientation". +- Pipeline: ball-search on (E²-1) Patterson coefficients → top-N candidates → + Sim MLRF (LL interpolation + per-shell σA fit) rescore. + +Acceptance: the true rotation must be within 8° of one of the top-5 rescored +peaks for 5/5 random trials. This is "shortlist contains truth" — the +downstream translation + rigid-body refinement (Phase 3) breaks the +Patterson rotation-function ambiguity by R-factor. +""" +import math +from pathlib import Path + +import pytest +import torch + +from torchref.alignment.ball_search import ( + ball_rotation_search, + edmonds_euler_from_rotation_matrix, + rotation_angular_distance_deg, + rotation_matrix_from_edmonds_euler, +) +from torchref.alignment.lattman_love import LattmanLoveInterpolator +from torchref.alignment.ml_rotation import sim_mlrf_rescore +from torchref.io.datasets.reflection_data import ReflectionData +from torchref.model import ModelFT +from torchref.symmetry import SpaceGroup + + +TEST_FILES = Path(__file__).resolve().parents[2] / "files" +PDB_1DAW = TEST_FILES / "pdb" / "1DAW.pdb" +MTZ_1DAW = TEST_FILES / "mtz" / "1DAW.mtz" + + +@pytest.fixture(scope="module") +def real_data_setup(): + """ + Real C2 F_obs + P1 search model. + + The search model is the same atoms as ``data`` but with spacegroup forced + to P1 — the standard MR setup. We use the proper `model.spacegroup` + setter, which routes through `_maybe_initialize_fft()` to rebuild the + SfFFT with the new symmetry. + """ + data = ReflectionData().load_mtz(str(MTZ_1DAW)) + model = ModelFT().load_pdb(str(PDB_1DAW)) + model.spacegroup = SpaceGroup("P 1") + return data, model + + +def _random_rotation(seed: int) -> torch.Tensor: + g = torch.Generator().manual_seed(seed) + A = torch.randn(3, 3, generator=g, dtype=torch.float64) + Q, R = torch.linalg.qr(A) + Q = Q @ torch.diag(torch.sign(torch.diag(R))) + if torch.det(Q) < 0: + Q[:, 0] = -Q[:, 0] + return Q + + +def _hkl_to_s(hkl, cell): + rec = cell.reciprocal_basis_matrix + if callable(rec): + rec = rec() + return hkl.to(torch.float64) @ rec.to(torch.float64) + + +def _shellbin_norm(F, smag, P): + order = torch.argsort(smag) + idx = torch.zeros_like(smag, dtype=torch.int64) + chunk = smag.numel() // P + for k in range(P): + a = k * chunk + b = (k + 1) * chunk if k < P - 1 else smag.numel() + idx[order[a:b]] = k + norm = torch.zeros_like(smag, dtype=F.dtype) + for k in range(P): + m = idx == k + norm[m] = (F[m] ** 2).mean().clamp(min=1e-30).sqrt() + return F / norm + + +def _min_err_over_sym(R_test, R_ref, sym_mats): + best = float("inf") + for k in range(sym_mats.shape[0]): + R_eq = sym_mats[k] @ R_ref.to(torch.float64) + e = rotation_angular_distance_deg(R_test.to(torch.float64), R_eq) + if e < best: + best = e + return best + + +@pytest.mark.integration +@pytest.mark.slow +@pytest.mark.parametrize("trial", range(5)) +def test_real_data_rotation_recovery(real_data_setup, trial): + """ + Real C2 F_obs + rotated P1 search model: the true rotation must be within + 8° of one of the top-5 ML-rescored peaks (modulo C2 symmetry). + """ + data, model = real_data_setup + sym_mats = data.spacegroup.matrices.to(torch.float64) + centric_full = data.centric if isinstance(data.centric, torch.Tensor) else torch.tensor(data.centric) + + F_obs_full = data.F.to(torch.float64).abs() + s_vec_full = _hkl_to_s(data.hkl, data.cell) + s_mag_full = s_vec_full.norm(dim=-1) + d_min, d_max = 4.0, 15.0 + keep = (s_mag_full >= 1.0 / d_max) & (s_mag_full <= 1.0 / d_min) + F_obs = F_obs_full[keep] + s_vec = s_vec_full[keep] + s_mag = s_mag_full[keep] + hkl = data.hkl[keep] + centric = centric_full[keep].to(torch.bool) + + P_shells = 20 + E_obs = _shellbin_norm(F_obs, s_mag, P_shells) + patt_obs = E_obs ** 2 - 1.0 # Patterson coefficient (origin-removed E²) + + # Apply random rotation to search-model atoms, build LL interpolator. + R_true = _random_rotation(seed=5000 + trial) + xyz_canonical = model.xyz().clone() + centroid = xyz_canonical.mean(0) + xyz_rot = (xyz_canonical - centroid) @ R_true.T.to(xyz_canonical.dtype) + centroid + model.xyz[:] = xyz_rot + try: + ll = LattmanLoveInterpolator(model, padding_factor=2.0, max_res_A=3.0) + F_calc = ll.evaluate( + torch.eye(3, dtype=torch.float32), hkl, data.cell, return_amplitude=True, + ).to(torch.float64) + finally: + model.xyz[:] = xyz_canonical + + E_calc = _shellbin_norm(F_calc, s_mag, P_shells) + patt_calc = E_calc ** 2 - 1.0 + + # Stage 1: fast ball-search on Patterson coefficients. + _, _, _, _, peaks = ball_rotation_search( + s_vec, patt_obs, s_vec, patt_calc, + L=32, P=P_shells, n_peaks=200, refine_subvoxel=True, n_refine=50, + sigma_threshold=-5.0, + ) + + # Stage 2: Sim MLRF rescore. + rescored = sim_mlrf_rescore( + peaks, F_obs, hkl, s_mag, centric, ll, data.cell, + n_shells=15, n_refine=200, batch_size=50, + ) + + # The true rotation (modulo C2 symmetry) must be within 8° of one of top-5. + best = float("inf") + best_rank = None + for rank, p in enumerate(rescored[:10]): + R_p = rotation_matrix_from_edmonds_euler(p.alpha, p.beta, p.gamma) + err = _min_err_over_sym(R_p, R_true, sym_mats) + if err < best: + best = err + best_rank = rank + assert best < 8.0, ( + f"trial {trial}: best top-10 ML peak {best:.2f}° from R_true " + f"(rank {best_rank}); rescored top-5 errs = " + f"{[round(_min_err_over_sym(rotation_matrix_from_edmonds_euler(p.alpha, p.beta, p.gamma), R_true, sym_mats), 2) for p in rescored[:5]]}" + ) diff --git a/tests/integration/alignment/test_rotation_recovery.py b/tests/integration/alignment/test_rotation_recovery.py new file mode 100644 index 00000000..b4869d3d --- /dev/null +++ b/tests/integration/alignment/test_rotation_recovery.py @@ -0,0 +1,209 @@ +""" +Gate test for `torchref.alignment.ball_rotation_search`: +recover a known random rotation of a model from its rotated F_calc. + +Setup (synthetic, P1 logic-only): +- Load 1DAW.pdb. +- Compute F_calc at the model's HKL set → `e_obs`. +- For each random R_true, rotate model coordinates by R_true, recompute + F_calc → `e_calc`. +- Run ball_rotation_search and check that the top peak gives R within `tol_deg` + of R_true. + +The test does NOT use F_obs from the MTZ — that would test model+data fit, not +rotation function correctness. Once 20/20 passes here, the next test exercises +the full pipeline against real F_obs. +""" + +import math +from pathlib import Path + +import numpy as np +import pytest +import torch + +from torchref.alignment.ball_search import ( + ball_rotation_search, + rotation_matrix_from_edmonds_euler, + rotation_angular_distance_deg, +) +from torchref.io.datasets.reflection_data import ReflectionData +from torchref.model import ModelFT + + +TEST_FILES = Path(__file__).resolve().parents[2] / "files" +PDB_1DAW = TEST_FILES / "pdb" / "1DAW.pdb" +MTZ_1DAW = TEST_FILES / "mtz" / "1DAW.mtz" + + +@pytest.fixture(scope="module") +def model_1daw(): + return ModelFT().load_pdb(str(PDB_1DAW)) + + +@pytest.fixture(scope="module") +def data_1daw(): + # Expand to P1 so the HKL set covers reciprocal-space directions uniformly, + # not just the ASU. Friedel mates are added by sh_expand_ball internally, + # but symmetry mates are needed here to get angular coverage. + return ReflectionData().load_mtz(str(MTZ_1DAW)).expand_to_p1(include_friedel=False) + + +def _random_rotation(seed: int) -> torch.Tensor: + """Uniform random rotation on SO(3) via QR of a Gaussian matrix.""" + g = torch.Generator().manual_seed(seed) + A = torch.randn(3, 3, generator=g, dtype=torch.float64) + Q, R = torch.linalg.qr(A) + Q = Q @ torch.diag(torch.sign(torch.diag(R))) + if torch.det(Q) < 0: + Q[:, 0] = -Q[:, 0] + return Q + + +def _hkl_to_s_vectors(hkl: torch.Tensor, cell) -> torch.Tensor: + """Convert HKL → reciprocal-lattice vectors |s| (1/Å).""" + rec_basis = cell.reciprocal_basis_matrix + if callable(rec_basis): + rec_basis = rec_basis() + rec_basis = rec_basis.to(torch.float64) + return hkl.to(torch.float64) @ rec_basis + + +def _normalize_amplitudes_by_shell(F: torch.Tensor, s_mag: torch.Tensor, n_shells: int = 24) -> torch.Tensor: + """ + Crude per-shell normalization: |E(h)| = |F(h)| / sqrt(<|F|²>_shell). + Returns a real tensor of the same shape as F. + """ + F = F.to(torch.float64) + s_mag = s_mag.to(torch.float64) + # equal-count binning + sorted_idx = torch.argsort(s_mag) + n = s_mag.numel() + shell_idx = torch.zeros(n, dtype=torch.int64) + chunk = n // n_shells + for k in range(n_shells): + a = k * chunk + b = (k + 1) * chunk if k < n_shells - 1 else n + shell_idx[sorted_idx[a:b]] = k + norm = torch.zeros(n, dtype=torch.float64) + for k in range(n_shells): + mask = shell_idx == k + f2 = (F[mask] ** 2).mean().clamp(min=1e-30) + norm[mask] = torch.sqrt(f2) + return F / norm + + +def _run_trial(model: ModelFT, hkl: torch.Tensor, R_true: torch.Tensor, + L: int, P: int, d_min: float, d_max: float): + """ + One rotation-recovery trial. + + We do NOT rotate the model atoms — that would change |F_calc| in ways + coupled to crystallographic symmetry (ModelFT applies symmetry mates). + Instead, we simulate "rotated F_calc" by placing F_orig values at rotated + reciprocal-space positions: this gives a field that is exactly the rotation + of the original field on the sphere, with no symmetry artifacts. This + tests the rotation function itself, isolated from F-calc-vs-symmetry. + """ + with torch.no_grad(): + F_orig = model(hkl).abs().to(torch.float64) + + s_obs = _hkl_to_s_vectors(hkl, model.cell) + s_mag = s_obs.norm(dim=-1) + keep = (s_mag >= 1.0 / d_max) & (s_mag <= 1.0 / d_min) + s_obs = s_obs[keep] + F_orig = F_orig[keep] + s_mag = s_mag[keep] + + # "Rotated calc": same values at rotated reciprocal positions. + R64 = R_true.to(torch.float64) + s_calc = s_obs @ R64.T + + # Normalize within each radial shell so the rotation function is dominated + # by directional structure, not by the |F|² magnitude vs resolution. + e_obs = _normalize_amplitudes_by_shell(F_orig, s_mag, n_shells=P) + e_calc = e_obs # values come along with positions + + C, alphas, betas, gammas, peaks = ball_rotation_search( + s_obs, e_obs, s_calc, e_calc, + L=L, P=P, n_peaks=30, refine_subvoxel=True, n_refine=10, + sigma_threshold=0.0, + ) + return peaks + + +def _min_err_over_pointgroup(R_test: torch.Tensor, R_true: torch.Tensor, + data) -> float: + """ + Return min angular distance (deg) between R_test and any symmetry-equivalent + of R_true under the spacegroup's point group. + + The field f_obs has the symmetry of |F_calc|, which equals the spacegroup + point group symmetry. So a recovered rotation R_test that differs from + R_true by a point-group operation is just as correct. + """ + sym_mats = data.spacegroup.matrices.to(torch.float64) # (N_ops, 3, 3) + best = float("inf") + R_test = R_test.to(torch.float64) + R_true = R_true.to(torch.float64) + for k in range(sym_mats.shape[0]): + R_eq = sym_mats[k] @ R_true + err = rotation_angular_distance_deg(R_test, R_eq) + if err < best: + best = err + return best + + +@pytest.mark.integration +@pytest.mark.parametrize("trial", range(5)) +def test_rotation_recovery_1daw_top_peak(model_1daw, data_1daw, trial): + """ + The true rotation (or any spacegroup-symmetry equivalent) should be within + 8° of one of the score-tied top peaks (L=32 → 5.6° voxels). + + "Score-tied" = score within 1% of the top peak. This tolerates the known + Patterson rotation-function degeneracy where |F|² accidentally has near- + symmetries that produce a (sub-percent-) close-second peak at a related but + distinct rotation. The pipeline always evaluates top-N candidates with + rigid-body refinement, so a tied second-place peak is just as good. + """ + R_true = _random_rotation(seed=1000 + trial) + peaks = _run_trial( + model_1daw, data_1daw.hkl, R_true, + L=32, P=20, d_min=4.0, d_max=15.0, + ) + top_score = peaks[0].score + tied = [p for p in peaks if p.score >= top_score * 0.99] + best_err = min( + _min_err_over_pointgroup( + rotation_matrix_from_edmonds_euler(p.alpha, p.beta, p.gamma), + R_true, + data_1daw, + ) + for p in tied + ) + assert best_err < 8.0, ( + f"trial {trial}: best score-tied peak {best_err:.2f}° from R_true " + f"({len(tied)} peaks within 1% of top score {top_score:.3e})" + ) + + +@pytest.mark.integration +@pytest.mark.parametrize("trial", range(5)) +def test_rotation_recovery_1daw_in_top5(model_1daw, data_1daw, trial): + """ + The true rotation (modulo spacegroup symmetry) must be within 8° of one + of the top-5 peaks. + """ + R_true = _random_rotation(seed=2000 + trial) + peaks = _run_trial( + model_1daw, data_1daw.hkl, R_true, + L=32, P=20, d_min=4.0, d_max=15.0, + ) + best_err = float("inf") + for p in peaks[:5]: + R_p = rotation_matrix_from_edmonds_euler(p.alpha, p.beta, p.gamma) + err = _min_err_over_pointgroup(R_p, R_true, data_1daw) + if err < best_err: + best_err = err + assert best_err < 8.0, f"trial {trial}: best of top-5 = {best_err:.2f}°" diff --git a/tests/unit/alignment/__init__.py b/tests/unit/alignment/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/unit/alignment/test_patterson_translation.py b/tests/unit/alignment/test_patterson_translation.py new file mode 100644 index 00000000..4fe5bb57 --- /dev/null +++ b/tests/unit/alignment/test_patterson_translation.py @@ -0,0 +1,123 @@ +""" +Unit tests for `amplitude_translation_search`. + +The function does a coarse-grid Pearson correlation between |F_obs|² and +|F_calc(h, t)|² over fractional translations. With the search model placed at +canonical positions and `F_obs` derived from a translated copy of the same +model, the top correlation peak (or one of the top-3) must land at `-t_true` +modulo an allowed origin shift of the spacegroup — i.e. the translation that +would bring the search model into agreement with the observed data. +""" +from pathlib import Path + +import numpy as np +import pytest +import torch + +from torchref.alignment.translation import amplitude_translation_search +from torchref.io.datasets.reflection_data import ReflectionData +from torchref.model import ModelFT +from torchref.symmetry import SpaceGroup + + +TEST_FILES = Path(__file__).resolve().parents[2] / "files" +PDB_1DAW = TEST_FILES / "pdb" / "1DAW.pdb" +MTZ_1DAW = TEST_FILES / "mtz" / "1DAW.mtz" + + +class _ModelEvaluator: + """Thin evaluator: returns model_p1(hkl) at integer HKL.""" + + def __init__(self, model_p1): + self._model = model_p1 + self.device = model_p1.xyz().device + + def evaluate(self, R, hkl, real_cell, return_amplitude=False): + hkl_int = hkl.round().to(torch.int64).to(self.device) + with torch.no_grad(): + f = self._model(hkl_int) + return f.abs() if return_amplitude else f + + +def _wrap_frac(t: np.ndarray) -> np.ndarray: + return (t + 0.5) % 1.0 - 0.5 + + +@pytest.fixture(scope="module") +def setup(): + canonical = ModelFT().load_pdb(str(PDB_1DAW)) + data = ReflectionData().load_mtz(str(MTZ_1DAW)) + mask = data.get_valid_mask() + return canonical, data, mask + + +@pytest.mark.unit +@pytest.mark.slow +def test_amplitude_tf_zero_translation(setup): + """Un-translated model: top-1 peak at the origin (modulo C-centering).""" + canonical, data, mask = setup + with torch.no_grad(): + F_obs = canonical(data.hkl[mask]).abs().to(torch.float64) + model_p1 = canonical.copy() + model_p1.spacegroup = SpaceGroup("P 1") + evaluator = _ModelEvaluator(model_p1) + + R_id = torch.eye(3, dtype=torch.float64) + _, _, peaks = amplitude_translation_search( + F_obs=F_obs, interpolator=evaluator, R_rotation=R_id, + hkl=data.hkl[mask], spacegroup=data.spacegroup, real_cell=data.cell, + grid_steps=12, n_peaks=10, cluster_radius=0.05, + ) + assert len(peaks) > 0 + # C2 + C-centering allowed origins: (0, *, 0) and (1/2, *, 1/2) + def origin_dist(t): + xz0 = np.linalg.norm(_wrap_frac(np.array([t[0], t[2]]))) + xz1 = np.linalg.norm(_wrap_frac(np.array([t[0] - 0.5, t[2] - 0.5]))) + return min(xz0, xz1) + best_dist = min(origin_dist(p.translation) for p in peaks[:3]) + assert best_dist < 0.10, ( + f"top-3 peaks miss origin-equivalent by {best_dist:.3f}; " + f"peaks: {[p.translation.tolist() for p in peaks[:3]]}" + ) + + +@pytest.mark.unit +@pytest.mark.slow +def test_amplitude_tf_recovers_known_translation(setup): + """A model translated by t_true: top-3 peaks include -t_true (mod origins).""" + canonical, data, mask = setup + t_true = np.array([0.18, -0.07, 0.23]) + # F_obs from canonical (un-translated); search model is canonical_p1 + # translated by t_true (so the recovered TF peak should be at -t_true mod + # the allowed origin shifts). + with torch.no_grad(): + F_obs = canonical(data.hkl[mask]).abs().to(torch.float64) + model_p1 = canonical.copy() + model_p1.spacegroup = SpaceGroup("P 1") + model_p1 = model_p1.translate( + torch.tensor(t_true, dtype=canonical.dtype_float), fractional=True, + ) + evaluator = _ModelEvaluator(model_p1) + + R_id = torch.eye(3, dtype=torch.float64) + _, _, peaks = amplitude_translation_search( + F_obs=F_obs, interpolator=evaluator, R_rotation=R_id, + hkl=data.hkl[mask], spacegroup=data.spacegroup, real_cell=data.cell, + grid_steps=12, n_peaks=10, cluster_radius=0.05, + ) + assert len(peaks) > 0 + # Expected: t_peak ≡ -t_true (mod allowed origin). C2 allowed origins + # along x and z: (0, *, 0) and (1/2, *, 1/2). y is polar. + def xz_dist(t): + d_origin = np.linalg.norm( + _wrap_frac(np.array([t[0] + t_true[0], t[2] + t_true[2]])) + ) + d_cshift = np.linalg.norm( + _wrap_frac(np.array([t[0] + t_true[0] - 0.5, t[2] + t_true[2] - 0.5])) + ) + return min(d_origin, d_cshift) + best_dist = min(xz_dist(p.translation) for p in peaks[:3]) + assert best_dist < 0.10, ( + f"top-3 peaks don't bracket -t_true (best xz_dist {best_dist:.3f}); " + f"peaks: {[p.translation.tolist() for p in peaks[:3]]}" + ) diff --git a/tests/unit/alignment/test_sh.py b/tests/unit/alignment/test_sh.py new file mode 100644 index 00000000..1f17acc3 --- /dev/null +++ b/tests/unit/alignment/test_sh.py @@ -0,0 +1,172 @@ +""" +Unit tests for torchref.alignment.sh: spherical harmonic primitives. + +Conventions verified: +- Y_{l,m} are fully orthonormal physics SH with Condon-Shortley phase. +- Matches scipy.special.sph_harm. +- For a centrosymmetric (Friedel-symmetric) input, sh_expand_ball produces + exactly-zero odd-l coefficients when enforce_friedel=True. +""" +import math + +import numpy as np +import pytest +import torch + +scipy_special = pytest.importorskip("scipy.special") + +from torchref.alignment.sh import ( + _bar_legendre_recurrence, + evaluate_ylm, + sh_expand_ball, + equal_count_shell_edges, + assign_shells, +) + + +def _scipy_ylm(l, m, theta, phi): + """Reference: scipy uses Y(m, l, phi, theta) with C-S phase included.""" + # scipy.special.sph_harm(m, l, phi, theta) returns Y_l^m(theta, phi) + # following the standard physics convention (with C-S phase). + return scipy_special.sph_harm(m, l, phi, theta) + + +@pytest.mark.parametrize("L", [4, 8, 16]) +def test_ylm_matches_scipy(L): + """Our Y_lm should match scipy.special.sph_harm at random points.""" + torch.manual_seed(0) + n = 20 + theta = torch.rand(n, dtype=torch.float64) * math.pi + phi = (torch.rand(n, dtype=torch.float64) - 0.5) * 2 * math.pi + + Y = evaluate_ylm(theta, phi, L) # (n, L, 2L-1) + + for l in range(L): + for m in range(-l, l + 1): + ref = _scipy_ylm(l, m, theta.numpy(), phi.numpy()) + got = Y[:, l, L - 1 + m].numpy() + np.testing.assert_allclose(got, ref, atol=1e-12, rtol=1e-10, + err_msg=f"mismatch at l={l}, m={m}") + + +def test_ylm_orthonormality_on_grid(): + """Y_lm should be ~orthonormal when integrated on a fine spherical grid.""" + L = 6 + # Gauss-Legendre in cos(theta), uniform in phi + n_theta = 2 * L + 4 + n_phi = 4 * L + 4 + # GL nodes in [-1, 1] + x_gl, w_gl = np.polynomial.legendre.leggauss(n_theta) + theta_np = np.arccos(x_gl) + phi_np = np.linspace(0, 2 * math.pi, n_phi, endpoint=False) + weights = (np.repeat(w_gl, n_phi)) * (2 * math.pi / n_phi) # (n_theta*n_phi,) + theta = torch.tensor(np.repeat(theta_np, n_phi), dtype=torch.float64) + phi = torch.tensor(np.tile(phi_np, n_theta), dtype=torch.float64) + + Y = evaluate_ylm(theta, phi, L) # (n_pts, L, 2L-1) + Yflat = Y.reshape(-1, L * (2 * L - 1)) + w = torch.tensor(weights, dtype=torch.float64) + + # G[a,b] = sum_pts w * Y*_a * Y_b + G = torch.einsum("p,pa,pb->ab", w.to(Yflat.dtype), Yflat.conj(), Yflat) + G_np = G.numpy() + + # Only diagonals for valid (l,m) entries (m | <= l) should be 1; off-diagonal ~0. + expected = np.zeros_like(G_np, dtype=np.complex128) + for l in range(L): + for m in range(-l, l + 1): + idx = l * (2 * L - 1) + (L - 1 + m) + expected[idx, idx] = 1.0 + # Zero out the entries we don't care about (l < |m|, where Y is zero anyway) + valid_mask = np.zeros(L * (2 * L - 1), dtype=bool) + for l in range(L): + for m in range(-l, l + 1): + valid_mask[l * (2 * L - 1) + (L - 1 + m)] = True + + G_valid = G_np[np.ix_(valid_mask, valid_mask)] + exp_valid = expected[np.ix_(valid_mask, valid_mask)] + np.testing.assert_allclose(G_valid, exp_valid, atol=1e-9, rtol=1e-9) + + +def test_bar_legendre_pole_values(): + """At the north pole (cos θ = 1), bar_P_l^m = 0 for m > 0 and bar_P_l^0 ≠ 0.""" + L = 5 + theta = torch.tensor([0.0], dtype=torch.float64) + bar_P = _bar_legendre_recurrence(torch.cos(theta), torch.sin(theta), L) # (1, L, L) + # m > 0 must be zero (sin θ = 0 kills sectorals; vertical recurrence then ~ 0) + for l in range(L): + for m in range(1, l + 1): + assert abs(bar_P[0, l, m].item()) < 1e-14, f"l={l}, m={m}: {bar_P[0,l,m].item()}" + # m = 0: bar_P_l^0(1) = sqrt((2l+1)/(4π)) (Legendre polynomial at 1 is 1) + for l in range(L): + expected = math.sqrt((2 * l + 1) / (4 * math.pi)) + np.testing.assert_allclose(bar_P[0, l, 0].item(), expected, atol=1e-12) + + +def test_friedel_enforces_even_l(): + """sh_expand_ball with enforce_friedel=True must give exactly-zero odd-l coefficients.""" + torch.manual_seed(1) + L = 8 + P = 4 + n_pts = 1000 + s_vectors = torch.randn(n_pts, 3, dtype=torch.float64) * 0.5 + # |s| roughly in [0, 1]; make sure non-zero + s_vectors = s_vectors / (1 + s_vectors.norm(dim=-1, keepdim=True) * 0.1) + s_mags = s_vectors.norm(dim=-1) + edges, _ = equal_count_shell_edges(s_mags, P) + shell_idx = assign_shells(s_mags, edges) + values = torch.rand(n_pts, dtype=torch.float64) + + f_plm = sh_expand_ball(s_vectors, values, shell_idx, P, L, enforce_friedel=True) + + # Odd-l rows must be exactly zero (after explicit zero of FP drift). + for l in range(1, L, 2): + assert f_plm[:, l, :].abs().max().item() == 0.0, f"l={l} not zero" + + +def test_friedel_without_enforce_has_odd_l_for_nonsymmetric_input(): + """Without Friedel enforcement, a non-centrosymmetric scatter produces nonzero odd-l.""" + torch.manual_seed(2) + L = 6 + P = 1 + # Place all mass at the north pole — extremely non-centrosymmetric + s_vectors = torch.tensor([[0.0, 0.0, 1.0]] * 5, dtype=torch.float64) + values = torch.ones(5, dtype=torch.float64) + shell_idx = torch.zeros(5, dtype=torch.int64) + + f_no_friedel = sh_expand_ball(s_vectors, values, shell_idx, P, L, + enforce_friedel=False) + # odd-l (l=1) entries should be non-trivial + odd_norm = f_no_friedel[:, 1, :].abs().max().item() + assert odd_norm > 1e-3, "expected nonzero odd-l without Friedel enforcement" + + +def test_shell_assignment_round_trip(): + """assign_shells gives indices that round-trip through equal_count_shell_edges.""" + torch.manual_seed(3) + s_mags = torch.rand(2000, dtype=torch.float64) * 2.0 + P = 16 + edges, centers = equal_count_shell_edges(s_mags, P) + idx = assign_shells(s_mags, edges) + # all should be in [0, P-1] + assert idx.min().item() >= 0 + assert idx.max().item() == P - 1 + # roughly equal counts (within 2x because of tie-breaking at edges) + counts = torch.bincount(idx, minlength=P) + assert counts.min().item() >= s_mags.numel() / (4 * P) + + +def test_sh_expand_zero_when_l_too_large_for_no_points(): + """Empty shell should give all-zero coefficients.""" + L = 4 + P = 3 + n_pts = 100 + torch.manual_seed(4) + s_vectors = torch.randn(n_pts, 3, dtype=torch.float64) + values = torch.ones(n_pts, dtype=torch.float64) + # Force shell_idx == 1 for all points (shell 0 and 2 empty) + shell_idx = torch.ones(n_pts, dtype=torch.int64) + f = sh_expand_ball(s_vectors, values, shell_idx, P, L, enforce_friedel=False) + assert f[0].abs().max() == 0 + assert f[2].abs().max() == 0 + assert f[1].abs().max() > 0 diff --git a/tests/unit/alignment/test_variance_weighting.py b/tests/unit/alignment/test_variance_weighting.py new file mode 100644 index 00000000..2c35e2fa --- /dev/null +++ b/tests/unit/alignment/test_variance_weighting.py @@ -0,0 +1,137 @@ +""" +Unit tests for the Phaser-style empirical per-shell variance weighting added in +Phase 3. + +Covers: +- `compute_patterson_shell_variance` returns ~1.0 per shell on a synthetic + acentric Wilson sample (E² ~ Exp(1)). +- Sparse-shell handling: a shell below `min_count` inherits a neighbour's value + and doesn't blow up the inverse-sqrt weight. +- `ball_rotation_search` with `auto_variance_weights=True` agrees with + `=False` on a synthetic flat-variance case (so the variance correction is a + no-op on data that already satisfies Wilson assumptions). +""" +import math + +import pytest +import torch + +from torchref.alignment.ball_search import ( + ball_rotation_search, + rotation_angular_distance_deg, + rotation_matrix_from_edmonds_euler, +) +from torchref.alignment.sh import compute_patterson_shell_variance + + +@pytest.mark.unit +def test_patterson_shell_variance_wilson(): + """On acentric Wilson E² ~ Exp(1), Var(E²-1) ≈ 1 per shell.""" + g = torch.Generator().manual_seed(0) + P = 6 + per_shell = 4000 + N = P * per_shell + e2 = -torch.log(torch.rand(N, generator=g, dtype=torch.float64).clamp(min=1e-12)) + patt = e2 - 1.0 + shell_idx = torch.repeat_interleave(torch.arange(P, dtype=torch.int64), per_shell) + var = compute_patterson_shell_variance(patt, shell_idx, P=P) + assert var.shape == (P,) + for k in range(P): + assert 0.85 < var[k].item() < 1.15, f"shell {k} variance {var[k].item()}" + + +@pytest.mark.unit +def test_patterson_shell_variance_handles_sparse_shells(): + """A shell with < min_count reflections inherits a neighbour's variance.""" + P = 4 + # Shell 0 dense (1000 pts), shell 1 sparse (2 pts), shells 2-3 dense. + patt = torch.cat([ + torch.randn(1000, dtype=torch.float64), + torch.tensor([10.0, -10.0], dtype=torch.float64), # would give huge var alone + torch.randn(1000, dtype=torch.float64), + torch.randn(1000, dtype=torch.float64), + ]) + shell_idx = torch.cat([ + torch.zeros(1000, dtype=torch.int64), + torch.ones(2, dtype=torch.int64), + torch.full((1000,), 2, dtype=torch.int64), + torch.full((1000,), 3, dtype=torch.int64), + ]) + var = compute_patterson_shell_variance(patt, shell_idx, P=P, min_count=8) + # Shell 1 should have inherited from shell 0 or 2 (~1.0), not the outlier value. + assert var[1].item() < 5.0, f"sparse shell 1 inherited huge variance {var[1].item()}" + + +@pytest.mark.unit +def test_patterson_shell_variance_eps_floor(): + """A near-zero-variance shell is floored at eps so weight stays bounded.""" + P = 2 + patt = torch.cat([ + torch.full((100,), 1.0, dtype=torch.float64), # variance ~ 0 + torch.randn(100, dtype=torch.float64), + ]) + shell_idx = torch.cat([ + torch.zeros(100, dtype=torch.int64), + torch.ones(100, dtype=torch.int64), + ]) + eps = 1e-3 + var = compute_patterson_shell_variance(patt, shell_idx, P=P, min_count=8, eps=eps) + assert var[0].item() >= eps + + +@pytest.mark.unit +def test_auto_variance_weights_recovers_synthetic_rotation(): + """ + On a uniform synthetic test (flat per-shell variance), the rotation + recovered with `auto_variance_weights=True` matches `=False` within a + fraction of a degree. The variance correction must not break the + well-conditioned case. + """ + g = torch.Generator().manual_seed(1) + N = 3000 + # uniform directions on the sphere + s_dirs = torch.randn(N, 3, generator=g, dtype=torch.float64) + s_dirs = s_dirs / s_dirs.norm(dim=-1, keepdim=True) + # log-uniform |s| in [1/15, 1/4] Å^-1 + s_mag = torch.empty(N, dtype=torch.float64).uniform_( + 1.0 / 15.0, 1.0 / 4.0, generator=g, + ) + s_obs = s_dirs * s_mag.unsqueeze(-1) + # Wilson-distributed Patterson coefficients per reflection. + e2 = -torch.log(torch.rand(N, generator=g, dtype=torch.float64).clamp(min=1e-12)) + patt_obs = e2 - 1.0 + + R_true = rotation_matrix_from_edmonds_euler(0.7, 1.2, 2.3, dtype=torch.float64) + # F_calc samples the same field rotated: place obs values at R·ŝ positions. + s_calc = s_obs @ R_true.T + patt_calc = patt_obs.clone() + + common_kwargs = dict( + L=20, P=10, n_peaks=20, d_min=4.0, d_max=15.0, + refine_subvoxel=True, n_refine=5, sigma_threshold=-5.0, + ) + _, _, _, _, peaks_on = ball_rotation_search( + s_obs, patt_obs, s_calc, patt_calc, + auto_variance_weights=True, **common_kwargs, + ) + _, _, _, _, peaks_off = ball_rotation_search( + s_obs, patt_obs, s_calc, patt_calc, + auto_variance_weights=False, **common_kwargs, + ) + + R_on = rotation_matrix_from_edmonds_euler( + peaks_on[0].alpha, peaks_on[0].beta, peaks_on[0].gamma, dtype=torch.float64, + ) + R_off = rotation_matrix_from_edmonds_euler( + peaks_off[0].alpha, peaks_off[0].beta, peaks_off[0].gamma, dtype=torch.float64, + ) + # both should be close to R_true; their mutual difference should be small. + err_on = rotation_angular_distance_deg(R_on, R_true) + err_off = rotation_angular_distance_deg(R_off, R_true) + assert err_on < 15.0, f"on: {err_on}" + assert err_off < 15.0, f"off: {err_off}" + # variance weighting should not move the answer by more than a voxel-or-so. + assert abs(err_on - err_off) < 6.0, ( + f"variance weighting changed the synthetic answer significantly: " + f"on={err_on:.2f}°, off={err_off:.2f}°" + ) diff --git a/tests/unit/alignment/test_wigner.py b/tests/unit/alignment/test_wigner.py new file mode 100644 index 00000000..82b7e853 --- /dev/null +++ b/tests/unit/alignment/test_wigner.py @@ -0,0 +1,159 @@ +""" +Unit tests for torchref.alignment.wigner. + +Conventions verified: +- D^l_{m,n}(α,β,γ) = e^{-imα} d^l_{m,n}(β) e^{-inγ} (Edmonds) +- d^l_{m,n}(0) = δ_{m,n}; d^l_{m,n}(π) = (-1)^{l+m} δ_{m,-n}. +- Unitarity: D^l D^l† = I for any (α,β,γ). +- Composition (sanity): the inverse FFT path agrees with the pointwise path. +""" +import math + +import numpy as np +import pytest +import torch + +from torchref.alignment.wigner import ( + small_d_block, + small_d_packed, + wigner_D_pointwise, + evaluate_rotation_function_grid, + evaluate_rotation_function_pointwise, +) + + +@pytest.mark.parametrize("l", [0, 1, 2, 3, 5, 8]) +def test_small_d_identity_at_zero(l): + """d^l_{m,n}(0) = δ_{m,n}.""" + beta = torch.tensor([0.0], dtype=torch.float64) + d = small_d_block(l, beta)[0] # (2l+1, 2l+1) + expected = torch.eye(2 * l + 1, dtype=torch.float64) + np.testing.assert_allclose(d.numpy(), expected.numpy(), atol=1e-12) + + +@pytest.mark.parametrize("l", [0, 1, 2, 3, 5, 8]) +def test_small_d_at_pi(l): + """d^l_{m,n}(π) = (-1)^{l+m} δ_{m,-n}.""" + beta = torch.tensor([math.pi], dtype=torch.float64) + d = small_d_block(l, beta)[0] # (2l+1, 2l+1) + size = 2 * l + 1 + expected = torch.zeros(size, size, dtype=torch.float64) + for m_idx in range(size): + m = m_idx - l + expected[m_idx, -m_idx - 1 + size] = (-1.0) ** (l + m) # n = -m → index size-1 - m_idx + np.testing.assert_allclose(d.numpy(), expected.numpy(), atol=1e-10) + + +@pytest.mark.parametrize("l", [1, 2, 4, 6]) +def test_small_d_unitary(l): + """d^l(β) is real-orthogonal: d^T d = I.""" + beta = torch.tensor([0.3, 1.1, 2.5], dtype=torch.float64) + d = small_d_block(l, beta) # (3, 2l+1, 2l+1) + for k in range(3): + dk = d[k] + prod = dk.T @ dk + np.testing.assert_allclose(prod.numpy(), np.eye(2 * l + 1), atol=1e-10, + err_msg=f"l={l} β={beta[k]:.3f}: d^T d ≠ I") + + +@pytest.mark.parametrize("L", [3, 5]) +def test_wigner_D_unitary(L): + """D^l(R) is unitary for any (α,β,γ).""" + torch.manual_seed(0) + n = 4 + alpha = torch.rand(n, dtype=torch.float64) * 2 * math.pi + beta = torch.rand(n, dtype=torch.float64) * math.pi + gamma = torch.rand(n, dtype=torch.float64) * 2 * math.pi + D = wigner_D_pointwise(alpha, beta, gamma, L) # (n, L, 2L-1, 2L-1) + + for k in range(n): + for l in range(L): + sl = slice(L - 1 - l, L - 1 + l + 1) + Dl = D[k, l, sl, sl] + prod = Dl @ Dl.conj().transpose(-1, -2) + eye = torch.eye(2 * l + 1, dtype=Dl.dtype) + np.testing.assert_allclose(prod.numpy(), eye.numpy(), atol=1e-10, + err_msg=f"l={l} not unitary at k={k}") + + +def test_wigner_D_diagonal_for_pure_z_rotation(): + """For β=γ=0, D^l_{m,n}(α,0,0) = δ_{m,n} e^{-imα}.""" + L = 4 + alpha = torch.tensor([0.5], dtype=torch.float64) + zero = torch.zeros_like(alpha) + D = wigner_D_pointwise(alpha, zero, zero, L)[0] # (L, 2L-1, 2L-1) + for l in range(L): + sl = slice(L - 1 - l, L - 1 + l + 1) + Dl = D[l, sl, sl].numpy() + # diagonal entries + for idx in range(2 * l + 1): + m = idx - l + expected = np.exp(-1j * m * 0.5) + np.testing.assert_allclose(Dl[idx, idx], expected, atol=1e-12) + # off-diagonal must vanish + offdiag = Dl - np.diag(np.diag(Dl)) + assert np.abs(offdiag).max() < 1e-12 + + +def test_pointwise_matches_grid(): + """evaluate_rotation_function_pointwise and *_grid agree at grid points.""" + L = 4 + torch.manual_seed(1) + # Random xi coefficients (only valid (l, |m|<=l, |n|<=l) entries non-zero). + xi = torch.zeros((L, 2 * L - 1, 2 * L - 1), dtype=torch.complex128) + for l in range(L): + for m in range(-l, l + 1): + for n in range(-l, l + 1): + xi[l, L - 1 + m, L - 1 + n] = (torch.randn(1).item() + + 1j * torch.randn(1).item()) + + C_grid, alphas, betas, gammas = evaluate_rotation_function_grid( + xi, L, n_alpha=2 * L, n_beta=2 * L, n_gamma=2 * L + ) + # Pick a few grid points and verify pointwise matches. + rng = np.random.default_rng(0) + for _ in range(5): + ka = int(rng.integers(0, 2 * L)) + kb = int(rng.integers(0, 2 * L)) + kg = int(rng.integers(0, 2 * L)) + C_from_grid = C_grid[kg, kb, ka] + a = alphas[ka:ka + 1] + b = betas[kb:kb + 1] + g = gammas[kg:kg + 1] + C_pointwise = evaluate_rotation_function_pointwise(xi, a, b, g, L)[0] + np.testing.assert_allclose(C_from_grid.item(), C_pointwise.item(), + atol=1e-10, rtol=1e-8) + + +def test_pointwise_real_for_hermitian_xi(): + """If ξ_{l,m,n} satisfies the conjugacy relation expected of a real cross-correlation, + then C(R) is real-valued.""" + L = 3 + torch.manual_seed(2) + # Build xi from a real-field convention: + # ξ_{l,m,n} = conj(f_{l,m}) g_{l,n}, where f and g are SH coefficients of real fields + # so f_{l,-m} = (-1)^m conj(f_{l,m}). + def random_real_field_coeffs(L): + f = torch.zeros((L, 2 * L - 1), dtype=torch.complex128) + for l in range(L): + for m in range(0, l + 1): + r = torch.randn(1).item() + 1j * torch.randn(1).item() + if m == 0: + r = complex(r.real, 0.0) + f[l, L - 1 + m] = r + if m > 0: + f[l, L - 1 - m] = ((-1) ** m) * np.conj(r) + return f + + f = random_real_field_coeffs(L) + g = random_real_field_coeffs(L) + xi = torch.zeros((L, 2 * L - 1, 2 * L - 1), dtype=torch.complex128) + for l in range(L): + for mi in range(2 * L - 1): + for ni in range(2 * L - 1): + xi[l, mi, ni] = f[l, mi].conj() * g[l, ni] + + C_grid, _, _, _ = evaluate_rotation_function_grid(xi, L) + imag = C_grid.imag.abs().max().item() + real = C_grid.real.abs().max().item() + assert imag < 1e-10 * max(real, 1.0), f"C imag={imag} too large (real max {real})" diff --git a/tests/unit/model/test_rotate.py b/tests/unit/model/test_rotate.py new file mode 100644 index 00000000..0ceb813e --- /dev/null +++ b/tests/unit/model/test_rotate.py @@ -0,0 +1,96 @@ +""" +Unit tests for Model.rotate. +""" +import math +from pathlib import Path + +import pytest +import torch + +from torchref.model import Model + + +TEST_PDB = Path(__file__).resolve().parents[2] / "files" / "pdb" / "1DAW.pdb" + + +@pytest.fixture +def loaded_model(): + return Model().load_pdb(str(TEST_PDB)) + + +def _Rz(angle_rad: float) -> torch.Tensor: + c, s = math.cos(angle_rad), math.sin(angle_rad) + return torch.tensor([[c, -s, 0.0], [s, c, 0.0], [0.0, 0.0, 1.0]]) + + +@pytest.mark.unit +def test_rotate_identity_preserves_coords(loaded_model): + """Rotation by identity is a no-op (within FP).""" + rotated = loaded_model.rotate(torch.eye(3)) + delta = (loaded_model.xyz() - rotated.xyz()).abs().max().item() + assert delta < 1e-5 + + +@pytest.mark.unit +def test_rotate_returns_new_instance(loaded_model): + """rotate() must NOT mutate the original.""" + xyz_before = loaded_model.xyz().clone() + rotated = loaded_model.rotate(_Rz(math.radians(30.0))) + xyz_after = loaded_model.xyz() + assert (xyz_before - xyz_after).abs().max().item() < 1e-9, \ + "Original model coords should be unchanged" + # Rotated model has different coords + delta = (loaded_model.xyz() - rotated.xyz()).abs().max().item() + assert delta > 1e-3 + + +@pytest.mark.unit +def test_rotate_preserves_centroid(loaded_model): + """Rotation around centroid (default) preserves the centroid.""" + R = _Rz(math.radians(60.0)) + rotated = loaded_model.rotate(R) + c_before = loaded_model.xyz().mean(dim=0) + c_after = rotated.xyz().mean(dim=0) + assert (c_before - c_after).norm().item() < 1e-4 + + +@pytest.mark.unit +def test_rotate_preserves_pairwise_distances(loaded_model): + """Rotation is rigid: pairwise distances are preserved.""" + R = _Rz(math.radians(45.0)) + rotated = loaded_model.rotate(R) + xyz0 = loaded_model.xyz() + xyz1 = rotated.xyz() + # Sample 50 random pairs of atoms + g = torch.Generator().manual_seed(0) + n = xyz0.shape[0] + i = torch.randint(0, n, (50,), generator=g) + j = torch.randint(0, n, (50,), generator=g) + d0 = (xyz0[i] - xyz0[j]).norm(dim=-1) + d1 = (xyz1[i] - xyz1[j]).norm(dim=-1) + assert (d0 - d1).abs().max().item() < 1e-4 + + +@pytest.mark.unit +def test_rotate_two_rotations_compose(loaded_model): + """rotate(R2) ∘ rotate(R1) ≈ rotate(R2 @ R1) (centroid-preserving).""" + R1 = _Rz(math.radians(30.0)) + R2 = _Rz(math.radians(50.0)) + composed = loaded_model.rotate(R1).rotate(R2) + single = loaded_model.rotate(R2 @ R1) + delta = (composed.xyz() - single.xyz()).abs().max().item() + assert delta < 1e-4 + + +@pytest.mark.unit +def test_rotate_around_explicit_center(loaded_model): + """Rotation around an explicit center: that point is fixed.""" + R = _Rz(math.radians(90.0)) + center = torch.tensor([5.0, -3.0, 1.0]) + rotated = loaded_model.rotate(R, center=center) + # The center, if it were a model atom, would be invariant. We can check by + # taking any atom and asserting xyz_new = R · (xyz_old − center) + center. + xyz0 = loaded_model.xyz() + expected = (xyz0 - center) @ R.T.to(xyz0.dtype) + center + delta = (rotated.xyz() - expected).abs().max().item() + assert delta < 1e-4 diff --git a/tests/unit/model/test_spacegroup_setter.py b/tests/unit/model/test_spacegroup_setter.py new file mode 100644 index 00000000..bcc2b1d5 --- /dev/null +++ b/tests/unit/model/test_spacegroup_setter.py @@ -0,0 +1,127 @@ +""" +Regression tests for the cell / spacegroup setters on Model, ModelFT, and SfFFT. + +Background: PyTorch's `nn.Module.__setattr__` intercepts assignments of +nn.Module values to attribute names and registers them in `_modules[name]`, +bypassing class-level `@property` descriptors. Without the `__setattr__` +override on our nn.Module-derived classes, assigning a SpaceGroup *object* +to e.g. `fft.spacegroup` creates a phantom `_modules['spacegroup']` entry +while leaving the canonical `_modules['_spacegroup']` (and the property +return value) stale. + +These tests exercise every public path for changing the spacegroup and +verify that |F_calc| actually reflects the new symmetry. +""" +import math +from pathlib import Path + +import pytest +import torch + +from torchref.model import ModelFT +from torchref.symmetry import Cell, SpaceGroup + + +TEST_PDB = Path(__file__).resolve().parents[2] / "files" / "pdb" / "1DAW.pdb" + + +@pytest.fixture +def model_c2(): + """A fresh ModelFT loaded from 1DAW (C2 spacegroup).""" + return ModelFT().load_pdb(str(TEST_PDB)) + + +@pytest.fixture +def test_hkl(): + return torch.tensor([[5, 3, 2], [4, 0, 5], [2, 2, 3]], dtype=torch.int64) + + +def _F_calc(model, hkl): + with torch.no_grad(): + return model(hkl).abs() + + +@pytest.mark.unit +def test_spacegroup_string_input(model_c2, test_hkl): + """`M.spacegroup = 'P 1'` (string) must take effect on F_calc.""" + F_c2 = _F_calc(model_c2, test_hkl) + model_c2.spacegroup = "P 1" + F_p1 = _F_calc(model_c2, test_hkl) + assert not torch.allclose(F_c2, F_p1), \ + "F_calc must change after spacegroup change (string input)" + assert str(model_c2.spacegroup) == "SpaceGroup('P1', number=1, n_ops=1)" + assert str(model_c2.fft.spacegroup) == "SpaceGroup('P1', number=1, n_ops=1)" + # No phantom _modules['spacegroup'] entry from PyTorch's auto-registration. + assert model_c2._modules.get("spacegroup") is None + assert model_c2.fft._modules.get("spacegroup") is None + + +@pytest.mark.unit +def test_spacegroup_module_input(model_c2, test_hkl): + """ + `M.spacegroup = SpaceGroup('P 1')` (Module input) must take effect. + This is the case that previously failed because PyTorch's + `nn.Module.__setattr__` would intercept the Module assignment and + bypass the property setter. + """ + F_c2 = _F_calc(model_c2, test_hkl) + model_c2.spacegroup = SpaceGroup("P 1") + F_p1 = _F_calc(model_c2, test_hkl) + assert not torch.allclose(F_c2, F_p1), \ + "F_calc must change after spacegroup change (Module input)" + assert str(model_c2.spacegroup) == "SpaceGroup('P1', number=1, n_ops=1)" + assert str(model_c2.fft.spacegroup) == "SpaceGroup('P1', number=1, n_ops=1)" + assert model_c2._modules.get("spacegroup") is None + assert model_c2.fft._modules.get("spacegroup") is None + + +@pytest.mark.unit +def test_fft_spacegroup_setter_module_input(model_c2, test_hkl): + """Low-level `M.fft.spacegroup = SpaceGroup('P 1')` must take effect.""" + F_c2 = _F_calc(model_c2, test_hkl) + model_c2.fft.spacegroup = SpaceGroup("P 1") + F_p1 = _F_calc(model_c2, test_hkl) + assert not torch.allclose(F_c2, F_p1) + assert str(model_c2.fft.spacegroup) == "SpaceGroup('P1', number=1, n_ops=1)" + + +@pytest.mark.unit +def test_fft_set_cell_and_spacegroup_module_input(model_c2, test_hkl): + """`M.fft.set_cell_and_spacegroup(cell, SpaceGroup('P 1'))` must take effect.""" + F_c2 = _F_calc(model_c2, test_hkl) + model_c2.fft.set_cell_and_spacegroup(model_c2.cell, SpaceGroup("P 1")) + F_p1 = _F_calc(model_c2, test_hkl) + assert not torch.allclose(F_c2, F_p1) + assert str(model_c2.fft.spacegroup) == "SpaceGroup('P1', number=1, n_ops=1)" + + +@pytest.mark.unit +def test_fft_spacegroup_setter_invalidates_caches(model_c2): + """The SfFFT's `map_symmetry` and reciprocal-symmetry extractor must be + rebuilt on spacegroup change (otherwise the next density-map build would + apply the OLD symmetry).""" + fft = model_c2.fft + # Force a state where caches exist. + fft.setup_grid() + old_map_sym_id = id(fft.map_symmetry) + fft.spacegroup = SpaceGroup("P 1") + # map_symmetry was rebuilt by setup_grid (called from the setter). + new_map_sym_id = id(fft.map_symmetry) + assert new_map_sym_id != old_map_sym_id, \ + "map_symmetry must be rebuilt after spacegroup change" + + +@pytest.mark.unit +def test_fft_cell_setter_invalidates_grid(model_c2): + """Setting `M.fft.cell = new_cell` must invalidate / rebuild the grid.""" + fft = model_c2.fft + fft.setup_grid() + old_grid_size = tuple(int(x) for x in fft.gridsize) + # Use a cell with very different parameters. + new_cell = Cell([80.0, 80.0, 80.0, 90.0, 90.0, 90.0]) + fft.cell = new_cell + # Grid was re-set up automatically with the new cell. + new_grid_size = tuple(int(x) for x in fft.gridsize) + assert new_grid_size != old_grid_size, \ + "Grid must be rebuilt after cell change" + assert fft.cell is new_cell diff --git a/tests/unit/model/test_translate.py b/tests/unit/model/test_translate.py new file mode 100644 index 00000000..d269a3b6 --- /dev/null +++ b/tests/unit/model/test_translate.py @@ -0,0 +1,74 @@ +""" +Unit tests for Model.translate (now returns a copy, mirroring Model.rotate). +""" +from pathlib import Path + +import pytest +import torch + +from torchref.model import ModelFT + + +TEST_PDB = Path(__file__).resolve().parents[2] / "files" / "pdb" / "1DAW.pdb" + + +@pytest.fixture +def model(): + return ModelFT().load_pdb(str(TEST_PDB)) + + +@pytest.mark.unit +def test_translate_returns_copy_not_inplace(model): + """Original model coordinates must be unchanged after translate().""" + xyz_before = model.xyz().detach().clone() + t = torch.tensor([5.0, -3.0, 1.0], dtype=model.dtype_float) + translated = model.translate(t) + xyz_after = model.xyz().detach() + assert torch.allclose(xyz_before, xyz_after), \ + "translate() must not mutate the source model" + expected = xyz_before + t.to(xyz_before.dtype) + assert torch.allclose(translated.xyz().detach(), expected, atol=1e-6) + + +@pytest.mark.unit +def test_translate_returns_different_object(model): + translated = model.translate(torch.tensor([1.0, 0.0, 0.0], dtype=model.dtype_float)) + assert translated is not model + # Two independent storages + assert translated.xyz().data_ptr() != model.xyz().data_ptr() + + +@pytest.mark.unit +def test_translate_fractional_matches_cartesian(model): + """`translate(t_frac, fractional=True)` agrees with the Cartesian form via + the cell's `fractional_to_cartesian` helper (the canonical conversion).""" + t_frac = torch.tensor([0.25, 0.10, -0.30], dtype=model.dtype_float) + t_cart = model.cell.fractional_to_cartesian(t_frac) + translated_frac = model.translate(t_frac, fractional=True) + translated_cart = model.translate(t_cart, fractional=False) + assert torch.allclose( + translated_frac.xyz().detach(), + translated_cart.xyz().detach(), + atol=1e-5, + ) + + +@pytest.mark.unit +def test_translate_preserves_b_and_occupancy(model): + """ADP and occupancy values must be unchanged by translation.""" + adp_before = model.adp().detach().clone() + occ_before = model.occupancy().detach().clone() + translated = model.translate( + torch.tensor([2.5, 0.0, 0.0], dtype=model.dtype_float), + ) + assert torch.allclose(translated.adp().detach(), adp_before) + assert torch.allclose(translated.occupancy().detach(), occ_before) + + +@pytest.mark.unit +def test_translate_zero_is_noop_on_coords(model): + """Translation by 0 returns a fresh copy with the same coordinates.""" + zero = torch.zeros(3, dtype=model.dtype_float) + translated = model.translate(zero) + assert translated is not model + assert torch.allclose(translated.xyz().detach(), model.xyz().detach()) diff --git a/torchref/alignment/__init__.py b/torchref/alignment/__init__.py index 8a9a1318..8189f22b 100644 --- a/torchref/alignment/__init__.py +++ b/torchref/alignment/__init__.py @@ -1,14 +1,18 @@ """ Alignment module for TorchRef. -Provides molecular replacement functionality including: - -1. Fast Rotation Function: Ball harmonic transform for rotation search -2. Translation Search: FFT-based translation function -3. Rigid Body Refinement: Optimization of rotation and translation -4. Unified Pipeline: Complete MR workflow with early stopping - -Example - Full MR Pipeline +Pure-PyTorch Patterson-based molecular replacement: + +1. Fast Rotation Function (`ball_search.ball_rotation_search`) — ball-harmonic + SO(3) cross-correlation, evaluated on an Euler-angle grid via inverse Wigner + transform. +2. Translation Search (`translation.fft_translation_search_torch`). +3. Rigid Body Refinement (`rigid_body.RigidBodyRefinement`) — LBFGS on rotation + and translation parameters with a maximum-likelihood x-ray target. +4. Unified Pipeline (`pipeline.MolecularReplacementPipeline`) — end-to-end + workflow with early-stopping. + +Example — full MR pipeline -------------------------- :: @@ -20,32 +24,33 @@ model = ModelFT().load_pdb('search_model.pdb') pipeline = MolecularReplacementPipeline(data, model) - solutions = pipeline.run(n_rotation_peaks=50, min_tries=3, max_tries=10) + solutions = pipeline.run(n_rotation_peaks=200, min_tries=3, max_tries=10) print(f"Best R-factor: {solutions[0].r_factor:.3f}") -Example - Individual Components +Example — individual components ------------------------------- :: from torchref.alignment import ( - ball_rotation_search_torch, + ball_rotation_search, fft_translation_search_torch, RigidBodyRefinement, ) # 1. Rotation search - rf, angles, peaks = ball_rotation_search_torch( - E_obs, s_obs, E_calc, s_calc, L=32, P=20 + C, alphas, betas, gammas, peaks = ball_rotation_search( + s_obs, e_obs, s_calc, e_calc, L=48, P=24, ) # 2. Translation search for top rotation - alpha, beta, gamma, score, sigma = peaks[0] + peak = peaks[0] corr_map, best_trans, trans_peaks = fft_translation_search_torch( F_obs, F_calc_rotated, hkl ) # 3. Rigid body refinement - rb = RigidBodyRefinement(model, data, initial_rotation=..., initial_translation=...) + rb = RigidBodyRefinement(model, data, + initial_rotation=..., initial_translation=...) result = rb.refine() """ @@ -57,61 +62,43 @@ ) # ============================================================================= -# Pipeline & Rotation search require JAX + s2fft (dev dependencies) +# Pure-PyTorch ball-harmonic rotation search (no JAX, s2fft, s2ball needed) # ============================================================================= -try: - from .pipeline import ( - MolecularReplacementPipeline, - MRSolution, - cluster_rotation_peaks, - rotation_angular_distance, - euler_angular_distance, - ) - from .ball_transform import ( - ball_rotation_search, - ball_rotation_search_torch, - rotation_matrix_from_euler_zyz, - rotation_matrix_to_euler_zyz, - rotation_matrix_to_quaternion, - check_rotation_recovery, - BallHarmonicCoefficients, - splat_evalues_to_ball, - compute_ball_harmonic_coefficients, - compute_ball_cross_correlation_coefficients, - evaluate_rotation_function, - find_rotation_peaks, - reduce_rotation_by_symmetry, - reduce_peaks_by_symmetry, - reduce_peaks_by_symmetry_torch, - cluster_rotation_peaks, - cluster_rotation_peaks_torch, - RotationCluster, - ) - _HAS_BALL_TRANSFORM = True -except ImportError: - _HAS_BALL_TRANSFORM = False - - _BALL_TRANSFORM_MSG = ( - "The alignment pipeline and rotation search require jax, s2fft, " - "s2ball, spherical, and quaternionic. " - "Install with: pip install torchref[alignment]" - ) - - def _missing_dep_factory(name): - """Create a callable stub that raises ImportError with install hint.""" - def _stub(*args, **kwargs): - raise ImportError( - f"{name} is not available. {_BALL_TRANSFORM_MSG}" - ) - _stub.__name__ = name - _stub.__qualname__ = name - return _stub - - # Provide stubs so that attribute access works but calling raises - MolecularReplacementPipeline = _missing_dep_factory("MolecularReplacementPipeline") - MRSolution = _missing_dep_factory("MRSolution") - ball_rotation_search = _missing_dep_factory("ball_rotation_search") - ball_rotation_search_torch = _missing_dep_factory("ball_rotation_search_torch") +from .ball_search import ( + BallHarmonicCoefficients, + RotationPeak, + ball_rotation_search, + compute_ball_harmonic_coefficients, + compute_ball_cross_correlation_coefficients, + find_rotation_peaks, + refine_peaks_subvoxel, + rotation_matrix_from_edmonds_euler, + edmonds_euler_from_rotation_matrix, + rotation_angular_distance_deg, +) +from .lattman_love import LattmanLoveInterpolator +from .ml_rotation import sim_mlrf_rescore, brute_ml_rotation_search +from .sh import ( + evaluate_ylm, + sh_expand_ball, + equal_count_shell_edges, + assign_shells, +) +from .wigner import ( + small_d_block, + small_d_packed, + wigner_D_pointwise, + evaluate_rotation_function_grid, + evaluate_rotation_function_pointwise, +) +from .pipeline import ( + MolecularReplacementPipeline, + MRSolution, + cluster_rotation_peaks, + rotation_angular_distance, + euler_angular_distance, +) +from .align import align_model_to_data # ============================================================================= # Translation search @@ -173,7 +160,42 @@ def _stub(*args, **kwargs): __all__ = [ # ------------------------------------------------------------------------- - # Translation search + # Rotation search (ball-harmonic Patterson, pure torch) + # ------------------------------------------------------------------------- + "ball_rotation_search", + "compute_ball_harmonic_coefficients", + "compute_ball_cross_correlation_coefficients", + "find_rotation_peaks", + "refine_peaks_subvoxel", + "BallHarmonicCoefficients", + "RotationPeak", + "rotation_matrix_from_edmonds_euler", + "edmonds_euler_from_rotation_matrix", + "rotation_angular_distance_deg", + "LattmanLoveInterpolator", + "sim_mlrf_rescore", + "brute_ml_rotation_search", + # Low-level math primitives + "evaluate_ylm", + "sh_expand_ball", + "equal_count_shell_edges", + "assign_shells", + "small_d_block", + "small_d_packed", + "wigner_D_pointwise", + "evaluate_rotation_function_grid", + "evaluate_rotation_function_pointwise", + # ------------------------------------------------------------------------- + # Pipeline + # ------------------------------------------------------------------------- + "MolecularReplacementPipeline", + "MRSolution", + "cluster_rotation_peaks", + "rotation_angular_distance", + "euler_angular_distance", + "align_model_to_data", + # ------------------------------------------------------------------------- + # Translation # ------------------------------------------------------------------------- "fft_translation_search", "fft_translation_search_torch", @@ -223,32 +245,3 @@ def _stub(*args, **kwargs): "VectorSampler", "get_rotation_sampling_range", ] - -if _HAS_BALL_TRANSFORM: - __all__ += [ - # Pipeline (main entry point) - "MolecularReplacementPipeline", - "MRSolution", - "cluster_rotation_peaks", - "rotation_angular_distance", - "euler_angular_distance", - # Rotation search - "ball_rotation_search", - "ball_rotation_search_torch", - "rotation_matrix_from_euler_zyz", - "rotation_matrix_to_euler_zyz", - "rotation_matrix_to_quaternion", - "check_rotation_recovery", - "BallHarmonicCoefficients", - "splat_evalues_to_ball", - "compute_ball_harmonic_coefficients", - "compute_ball_cross_correlation_coefficients", - "evaluate_rotation_function", - "find_rotation_peaks", - "reduce_rotation_by_symmetry", - "reduce_peaks_by_symmetry", - "reduce_peaks_by_symmetry_torch", - "cluster_rotation_peaks", - "cluster_rotation_peaks_torch", - "RotationCluster", - ] diff --git a/torchref/alignment/align.py b/torchref/alignment/align.py new file mode 100644 index 00000000..0346f920 --- /dev/null +++ b/torchref/alignment/align.py @@ -0,0 +1,727 @@ +""" +End-to-end molecular replacement alignment entry point. + +This module owns the full MR pipeline: + + LERF1 ball-search rotation + → Sim MLRF rescore + → amplitude-correlation translation search + → local translation refine (analytical-scale R) + → dense rotation sampling (ML-LLG, multi-pass) + → LBFGS rigid-body polish + → final solvent-aware Scaler refit (user-facing R-work) + +The single public function `align_model_to_data` is what `ModelFT.fit_to_data` +delegates to; the latter is a thin wrapper. Keep alignment logic in this file +to keep `torchref/model/model_ft.py` focused on the FFT structure-factor model. +""" + +from __future__ import annotations + +import math +import time +from contextlib import contextmanager +from typing import TYPE_CHECKING + +import torch + +from .ball_search import ( + RotationPeak, + ball_rotation_search, + edmonds_euler_from_rotation_matrix, + rotation_matrix_from_edmonds_euler, +) +from .lattman_love import LattmanLoveInterpolator +from .ml_rotation import sim_mlrf_rescore +from .rigid_body import RigidBodyRefinement +from .sh import ( + apply_overall_anisotropy, + assign_shells, + equal_count_shell_edges, + fit_overall_anisotropy, +) +from .translation import ( + amplitude_translation_search, + local_translation_refine, + precompute_G_for_rotation, +) + +if TYPE_CHECKING: + from ..io.datasets.reflection_data import ReflectionData + from ..model.model_ft import ModelFT + + +# --------------------------------------------------------------------------- +# Stage timing +# --------------------------------------------------------------------------- + + +class _StageTimer: + """Lightweight wall-clock accumulator. Gated by ``verbose >= 2``. + + Two interleavable usages: + * ``with t.stage(name):`` block — records the block's wall time. + * ``t.start(name)`` / ``t.stop(name)`` — checkpoint pair, no indent. + + The summary table prints stages aggregated by name; per-rotation loop + stages (translation search etc.) get aggregated counts. + """ + + def __init__(self, enabled: bool): + self.enabled = enabled + self.records: list[tuple[str, float]] = [] + self._open: dict[str, float] = {} + + @contextmanager + def stage(self, name: str): + if not self.enabled: + yield + return + t0 = time.perf_counter() + try: + yield + finally: + self.records.append((name, time.perf_counter() - t0)) + + def start(self, name: str) -> None: + if self.enabled: + self._open[name] = time.perf_counter() + + def stop(self, name: str) -> None: + if not self.enabled: + return + t0 = self._open.pop(name, None) + if t0 is not None: + self.records.append((name, time.perf_counter() - t0)) + + def summary(self) -> str: + if not self.records: + return "" + # Aggregate repeated stage names (the per-rotation loop visits the + # translation stages once per candidate rotation). + agg: dict[str, list[float]] = {} + for name, dt in self.records: + agg.setdefault(name, []).append(dt) + total = sum(sum(v) for v in agg.values()) + lines = [ + f"{'stage':<32s} {'count':>5s} {'wall_s':>10s} {'%':>6s}", + "-" * 60, + ] + for name, vs in agg.items(): + wall = sum(vs) + lines.append( + f"{name:<32s} {len(vs):>5d} {wall:>10.3f} " + f"{100 * wall / total:>5.1f}%" + ) + lines.append("-" * 60) + lines.append(f"{'TOTAL':<32s} {'':>5s} {total:>10.3f} 100.0%") + return "\n".join(lines) + + +# --------------------------------------------------------------------------- +# Internal helpers +# --------------------------------------------------------------------------- + + +def _shellbin_norm_etrick( + F: torch.Tensor, smag: torch.Tensor, P: int +) -> torch.Tensor: + """Per-shell E-trick normalisation: ``E = F / sqrt(_shell)``. + + Equal-count shell binning by sorted ``smag``; small-count shells just + inherit their normaliser from the local mean. + """ + order = torch.argsort(smag) + idx = torch.zeros_like(smag, dtype=torch.int64) + chunk = smag.numel() // P + for k in range(P): + a = k * chunk + b = (k + 1) * chunk if k < P - 1 else smag.numel() + idx[order[a:b]] = k + norm = torch.zeros_like(smag, dtype=F.dtype) + for k in range(P): + m = idx == k + norm[m] = (F[m] ** 2).mean().clamp(min=1e-30).sqrt() + return F / norm + + +def _external_rwork(model: "ModelFT", data: "ReflectionData") -> float: + """Full-resolution scaled R-work via the standard Scaler. + + The TF + local refine work in analytical-scale R-factor (which ranks + candidates correctly but isn't the user-facing R-work). We compute the + proper Scaler-fit R-work once per finalist. + """ + from ..scaling import Scaler + + # Build the Scaler on the model's device so that its anisotropy U + # tensor and per-bin scales land alongside `data.hkl`/`model(hkl)` — + # otherwise Scaler.forward's `matmul(self.s, U)` mixes CPU/GPU and + # crashes at refine_lbfgs. + s = Scaler(model=model, data=data, nbins=20, verbose=0, + device=model.xyz().device) + fc = model(data.hkl) + s.initialize(fc) + s.refine_lbfgs(fcalc=fc) + rw, _ = s.rfactor(fc) + return rw.item() if hasattr(rw, "item") else float(rw) + + +class _DirectModelEvaluator: + """Returns ``F_p1(hkl)`` of a P1-spacegroup model at integer HKL. + + Wraps a `ModelFT` to expose the same `.evaluate(R, hkl, cell, ...)` API + as `LattmanLoveInterpolator`, for use by the translation search. + """ + + def __init__(self, m: "ModelFT") -> None: + self._m = m + self.device = m.xyz().device + + def evaluate(self, R, hkl, real_cell, return_amplitude=False): + hkl_int = hkl.round().to(torch.int64).to(self.device) + with torch.no_grad(): + f = self._m(hkl_int) + return f.abs() if return_amplitude else f + + +def _rodrigues(omega: torch.Tensor) -> torch.Tensor: + """Rodrigues axis-angle → SO(3). `omega = θ · axis` (radians). + + Accepts shape (3,) for a single rotation or (..., 3) for a batched stack + and returns matching (3, 3) or (..., 3, 3). The small-θ limit is handled + implicitly: sin(θ)→0 and (1-cos θ)→0 zero out the K and K² contributions + so R→I as θ→0; `clamp(min=1e-30)` prevents NaN from axis=0/0. + """ + if omega.dtype != torch.float64: + omega = omega.to(torch.float64) + is_single = omega.dim() == 1 + if is_single: + omega = omega.unsqueeze(0) + + th = omega.norm(dim=-1, keepdim=True) # (..., 1) + axis = omega / th.clamp(min=1e-30) # (..., 3) + zeros = torch.zeros_like(axis[..., 0]) + K = torch.stack([ + torch.stack([zeros, -axis[..., 2], axis[..., 1]], dim=-1), + torch.stack([axis[..., 2], zeros, -axis[..., 0]], dim=-1), + torch.stack([-axis[..., 1], axis[..., 0], zeros], dim=-1), + ], dim=-2) # (..., 3, 3) + + th_b = th.unsqueeze(-1) # (..., 1, 1) + sin_th = torch.sin(th_b) + cos_th = torch.cos(th_b) + + eye = torch.eye(3, dtype=omega.dtype, device=omega.device) + eye_b = eye.expand(*omega.shape[:-1], 3, 3) + KK = torch.matmul(K, K) + R = eye_b + sin_th * K + (1.0 - cos_th) * KK + + if is_single: + R = R.squeeze(0) + return R + + +# --------------------------------------------------------------------------- +# Public entry point +# --------------------------------------------------------------------------- + + +def align_model_to_data( + model: "ModelFT", + data: "ReflectionData", + *, + d_min: float = 4.0, + d_max: float = 15.0, + L: int = 48, + n_shells: int = 20, + n_rotation_peaks: int = 500, + n_ml_refine: int = 500, + ll_max_res_A: float = 3.0, + ll_padding_factor: float = 2.0, + verbose: int = 0, + auto_variance_weights: bool = True, + do_translation: bool = True, + n_translation_peaks: int = 20, + n_translation_candidates: int = 3, + translation_grid_steps: int = 16, + n_rotation_candidates: int = 15, + do_joint_refine: bool = True, + joint_refine_max_res_A: float = 4.0, + joint_refine_expected_rot_error: float = 0.1, +) -> "ModelFT": + """Run full MR alignment of ``model`` against ``data``. + + Returns a new rotated+translated+refined ``ModelFT`` carrying + ``last_alignment_rotation``, ``last_alignment_translation`` and + ``last_alignment_rfactor`` provenance attributes. + + See `ModelFT.fit_to_data` for full kwarg semantics — this function is the + canonical implementation; `fit_to_data` is a thin wrapper. + """ + from ..scaling import Scaler # noqa: F401 (imported by _external_rwork) + from ..symmetry import SpaceGroup + + if not model.initialized: + raise RuntimeError( + "Cannot fit an uninitialized ModelFT. Load PDB data first." + ) + + timer = _StageTimer(enabled=verbose >= 2) + + # Device propagation: align_model_to_data runs on the model's device by + # default. Data tensors (data.F, data.hkl, etc.) often arrive on CPU and + # need to be moved to match — otherwise the very first hkl-derived + # quantity (s_mag) lives on CPU while F_calc from the GPU-resident LL + # interpolator lives on GPU, and _shellbin_norm_etrick crashes with + # "Expected all tensors to be on the same device". + device = model.xyz().device + + # --- Prepare F_obs and shell-normalized E_obs in resolution range --- + # Do the masking on CPU (data tensors arrive there) then move the + # resolution-bounded slices to `device`. Avoids GPU-index-into-CPU + # crashes when `keep` is on a different device than `data.centric`. + timer.start("0_data_prep") + F_obs = data.F.to(torch.float64).abs() + hkl_all = data.hkl + rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) + s_vec_all = hkl_all.to(torch.float64) @ rec_basis + s_mag_all = s_vec_all.norm(dim=-1) + keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min) + if keep.sum().item() < n_shells * 5: + raise ValueError( + f"Too few reflections ({keep.sum().item()}) in [{d_min},{d_max}] Å " + f"for {n_shells} shells; widen the resolution range." + ) + # Index on CPU, then move slices to `device`. The model device is the + # canonical destination; downstream stages (LL.evaluate, ball_search, + # sim_mlrf_rescore) all infer device from their inputs. + F_obs = F_obs[keep].to(device) + hkl = hkl_all[keep].to(device) + s_vec = s_vec_all[keep].to(device) + s_mag = s_mag_all[keep].to(device) + centric = ( + data.centric[keep].to(torch.bool).to(device) + if hasattr(data, "centric") + else torch.zeros_like(F_obs, dtype=torch.bool) + ) + + timer.stop("0_data_prep") + + # Popov-Bourenkov overall anisotropy correction (full variance fix): + # Removes the direction-dependent Wilson-falloff in F_obs that otherwise + # biases the rotation function on anisotropic and high-symmetry crystals + # (P6522, etc.). Fitted from F_obs alone — no model dependence. + timer.start("1_anisotropy_fit") + aniso_edges, _ = equal_count_shell_edges(s_mag, n_shells) + aniso_idx = assign_shells(s_mag, aniso_edges) + U_aniso = fit_overall_anisotropy( + F_obs, s_vec, aniso_idx, P=n_shells, min_count=20, + ) + if verbose > 0: + print( + f"fit_to_data: overall U-aniso diag (Ų) = " + f"({U_aniso[0, 0].item():+.2f}, {U_aniso[1, 1].item():+.2f}, " + f"{U_aniso[2, 2].item():+.2f})", + flush=True, + ) + F_obs_aniso = apply_overall_anisotropy(F_obs, s_vec, U_aniso) + timer.stop("1_anisotropy_fit") + + # --- Build LL interpolator from a P1 view of the model --- + timer.start("2_ll_build") + if verbose > 0: + print( + f"fit_to_data: building Lattman-Love interpolator " + f"(box={ll_padding_factor}·diam, max_res={ll_max_res_A} Å)…", + flush=True, + ) + ll = LattmanLoveInterpolator( + model, padding_factor=ll_padding_factor, max_res_A=ll_max_res_A, + verbose=verbose, + ) + + # --- Symmetry-expand reciprocal-space points for the Patterson SH expansion --- + # The MTZ stores only the spacegroup ASU (1 / n_ops of reciprocal space). + # The observed Patterson is spacegroup-invariant (|F(S_k h)| = |F(h)|), + # so the SH expansion of P_obs must sample the full sphere — otherwise + # the rotation function loses its spacegroup symmetry and the true + # orientation is no longer a global maximum. On 1AK5 (P432, 24 ops) + # the un-expanded rotation function picked maxima 5× higher than the + # value at R_true; expanding F_obs to all 24 symmetry mates makes + # C(R) = C(S_k R) by construction (and the calc-side LL evaluator + # is queried at the same expanded HKL set so the cross-correlation is + # consistent). For low-sym cells (P21, n_ops=2) this is a no-op factor; + # for P432 it makes 1AK5 / 3K7M find the right basin. + sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) # (n_ops, 3, 3) + n_ops_sg = int(sg_mats.shape[0]) + # h_sym[k, i, :] = S_k · hkl[i] (rotation part of the symop) + hkl_sym = torch.einsum("kij,nj->kni", sg_mats, hkl.to(torch.float64)) + hkl_sym_flat = hkl_sym.reshape(-1, 3) # (n_ops·N, 3) + s_vec_sym = hkl_sym_flat @ rec_basis.to(device) # (n_ops·N, 3) + s_mag_sym = s_vec_sym.norm(dim=-1) + # |F_obs| is replicated across symmetry mates (Patterson invariance). + F_obs_aniso_sym = F_obs_aniso.unsqueeze(0).expand(n_ops_sg, -1).reshape(-1) + E_obs_sym = _shellbin_norm_etrick(F_obs_aniso_sym, s_mag_sym, n_shells) + patt_obs = E_obs_sym ** 2 - 1.0 + + # F_calc evaluated at the SAME symmetry-expanded HKL set so the cross- + # correlation between f_obs and f_calc samples the same directions. + F_calc_sym = ll.evaluate( + torch.eye(3, dtype=torch.float32), hkl_sym_flat, data.cell, + return_amplitude=True, + ).to(torch.float64) + E_calc_sym = _shellbin_norm_etrick(F_calc_sym, s_mag_sym, n_shells) + patt_calc = E_calc_sym ** 2 - 1.0 + # Replace the un-expanded s_vec with the expanded one for ball_search. + s_vec_for_search = s_vec_sym + timer.stop("2_ll_build") + + # --- Stage 1: fast Patterson ball-search --- + timer.start("3_ball_search") + if verbose > 0: + print( + f"fit_to_data: ball-search (L={L}, P={n_shells}, " + f"n_peaks={n_rotation_peaks})…", + flush=True, + ) + _, _, _, _, peaks = ball_rotation_search( + s_vec_for_search, patt_obs, s_vec_for_search, patt_calc, + L=L, P=n_shells, n_peaks=n_rotation_peaks, + refine_subvoxel=True, n_refine=min(n_rotation_peaks, 50), + sigma_threshold=-5.0, + auto_variance_weights=auto_variance_weights, + ) + timer.stop("3_ball_search") + + # --- Stage 2: Sim-MLRF rescore (per-shell σA fit per candidate) --- + timer.start("4_sim_mlrf_rescore") + if verbose > 0: + print( + f"fit_to_data: ML rescoring top " + f"{min(len(peaks), n_ml_refine)} peaks…", + flush=True, + ) + rescored = sim_mlrf_rescore( + peaks, F_obs, hkl, s_mag, centric, ll, data.cell, + n_shells=max(n_shells // 2, 8), + n_refine=min(len(peaks), n_ml_refine), + batch_size=50, + verbose=verbose, + auto_variance_weights=auto_variance_weights, + ) + timer.stop("4_sim_mlrf_rescore") + if not rescored: + raise RuntimeError("Rotation search produced no peaks.") + + n_rot = min(n_rotation_candidates if do_translation else 1, len(rescored)) + + # The Patterson rotation function has a centrosymmetric ambiguity, but + # `sim_mlrf_rescore` ranks Patterson-equivalents adjacent. `n_rot ≥ 3` + # is enough without explicitly multiplying by spacegroup rotations. + if verbose > 0 and n_rot > 1: + print( + f"fit_to_data: trying top {n_rot} rotation candidates " + f"(Patterson-equivalents covered by LLG ranking).", + flush=True, + ) + + def _candidate(k): + peak = rescored[k] + R_rec = rotation_matrix_from_edmonds_euler( + peak.alpha, peak.beta, peak.gamma, + ) + R_app = R_rec.T.contiguous() + rot = model.rotate( + R_app.to(device=model.device, dtype=model.dtype_float), + ) + rot.last_alignment_rotation = R_rec + return rot, R_rec, R_app, peak + + if not do_translation: + rotated, R_recovered, _, top = _candidate(0) + if verbose > 0: + print( + f"fit_to_data: top peak LLG = {top.score:.2f} " + f"(σ_Z = {top.sigma:.2f}); applying R⁻¹ to coords.", + flush=True, + ) + if verbose >= 2: + print("\n" + timer.summary(), flush=True) + return rotated + + # --- Stage 3: translation search + analytical-R local refine --- + hkl_full = data.hkl + F_obs_full = data.F + if hasattr(data, "get_valid_mask"): + tmask = data.get_valid_mask() + else: + tmask = torch.ones( + F_obs_full.shape[0], dtype=torch.bool, device=F_obs_full.device, + ) + # Index on CPU then move slices to `device` — same pattern as the + # early `keep`-mask section; the validity mask lives on the data's + # device, while downstream consumers run on `model.device`. + F_obs_amp = F_obs_full[tmask].abs().to(torch.float64).to(device) + hkl_keep = hkl_full[tmask].to(device) + + global_best = None + for k_rot in range(n_rot): + rotated_k, R_recovered_k, _, peak_k = _candidate(k_rot) + if verbose > 0: + print( + f"\nfit_to_data: rot{k_rot} " + f"(LLG={peak_k.score:.2f}, σ_Z={peak_k.sigma:.2f})", + flush=True, + ) + + if str(rotated_k.spacegroup) != str(data.spacegroup): + rotated_k.spacegroup = data.spacegroup + + rotated_p1 = rotated_k.copy() + rotated_p1.spacegroup = SpaceGroup("P 1") + evaluator = _DirectModelEvaluator(rotated_p1) + + # Pre-compute per-sym F_asu contributions once per rotation; reused + # by both coarse TF and each local refine. + timer.start("5_precompute_G") + G_pre, h_R_pre = precompute_G_for_rotation( + evaluator, torch.eye(3, dtype=torch.float64), + hkl_keep, data.spacegroup, data.cell, + ) + timer.stop("5_precompute_G") + + timer.start("6_amplitude_TF") + _, _, t_peaks = amplitude_translation_search( + F_obs=F_obs_amp, interpolator=evaluator, + R_rotation=torch.eye(3, dtype=torch.float64), + hkl=hkl_keep, + spacegroup=data.spacegroup, real_cell=data.cell, + grid_steps=translation_grid_steps, + n_peaks=n_translation_peaks, + cluster_radius=0.05, + precomputed_G=G_pre, precomputed_h_R=h_R_pre, + ) + timer.stop("6_amplitude_TF") + if not t_peaks: + if verbose > 0: + print(" no translation peaks; skipping", flush=True) + continue + if verbose > 0: + tt = tuple(round(float(x), 3) for x in t_peaks[0].translation.tolist()) + print( + f" top translation t={tt} corr={t_peaks[0].score:.4f}", + flush=True, + ) + + if not do_joint_refine: + t_top = torch.as_tensor( + t_peaks[0].translation, + dtype=model.dtype_float, device=rotated_k.device, + ) + translated = rotated_k.translate(t_top, fractional=True) + translated.last_alignment_rotation = R_recovered_k + translated.last_alignment_translation = t_top + return translated + + for k_t, tp in enumerate(t_peaks[:n_translation_candidates]): + t_init = torch.as_tensor(tp.translation, dtype=torch.float64) + timer.start("7_local_TF_refine") + # Single-pass local refine (was 2). Pass-2 zoomed to ~0.0017 + # fractional resolution; the downstream LBFGS rigid-body polish + # refines to gradient-tolerance anyway, so pass 2 was just + # paying ~half the local-TF cost for a precision that gets + # overridden a few steps later. + t_refined, r_analytic = local_translation_refine( + F_obs=F_obs_amp, interpolator=evaluator, + R_rotation=torch.eye(3, dtype=torch.float64), + hkl=hkl_keep, + spacegroup=data.spacegroup, real_cell=data.cell, + t_init=t_init, radius=0.06, grid_steps=13, + n_refinement_passes=1, + precomputed_G=G_pre, precomputed_h_R=h_R_pre, + ) + timer.stop("7_local_TF_refine") + if verbose > 0: + print( + f" rot{k_rot} trans{k_t}: " + f"R(analytic)={r_analytic:.4f}, " + f"t={[round(float(x), 3) for x in t_refined.tolist()]}", + flush=True, + ) + if global_best is None or r_analytic < global_best[0]: + global_best = (r_analytic, rotated_k, R_recovered_k, t_refined) + + if global_best is None: + raise RuntimeError("Translation + joint refine produced no candidates.") + r_analytic_best, rot_best, R_recovered_best, t_refined_best = global_best + refined = rot_best.translate( + t_refined_best.to(model.dtype_float), fractional=True, + ) + + # --- Stage 4: dense rotation sampling at the found translation --- + if do_joint_refine: + timer.start("8_dense_R_ll_build") + refined_p1 = refined.copy() + refined_p1.spacegroup = SpaceGroup("P 1") + ll_refine = LattmanLoveInterpolator( + refined_p1, padding_factor=ll_padding_factor, + max_res_A=ll_max_res_A, verbose=0, + ) + timer.stop("8_dense_R_ll_build") + + centric_keep = ( + data.centric[tmask].to(torch.bool).to(device) if hasattr(data, "centric") + else torch.zeros(hkl_keep.shape[0], dtype=torch.bool, device=device) + ) + rec_basis_keep = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) + s_mag_keep = (hkl_keep.to(torch.float64) @ rec_basis_keep).norm(dim=-1) + + # Dense-R sampling: 2-pass zoom on the σ_A-fitted ML LLG which has + # FWHM ~2–3° (sharper than the Patterson rotation function's ~10° + # because of log-likelihood curvature and σ_A up-weighting). Pass 1 + # at 9³ × ±5.7° (1.43° spacing) gives 1–2 samples per ML-FWHM and + # locates the basin; pass 2 at 5³ × ±1.43° (0.71° spacing) zooms in. + # Single-pass or larger spacing regresses R-work; finer than this + # is just sampled by the downstream LBFGS polish anyway. + n_per_axis_pass = [9, 5] + zoom_factor = 4.0 + radii = [ + float(joint_refine_expected_rot_error), + float(joint_refine_expected_rot_error) / zoom_factor, + ] + R_accumulated = torch.eye(3, dtype=torch.float64) + + for pass_idx, max_perturb_rad in enumerate(radii): + n_per_axis = n_per_axis_pass[pass_idx] + coords_r = torch.linspace( + -max_perturb_rad, max_perturb_rad, n_per_axis, + dtype=torch.float64, + ) + wx, wy, wz = torch.meshgrid( + coords_r, coords_r, coords_r, indexing="ij", + ) + omegas = torch.stack( + [wx.flatten(), wy.flatten(), wz.flatten()], dim=-1, + ) + # Batched Rodrigues + matmul: one (B, 3, 3) build instead of B + # per-omega calls. Previously _rodrigues had `.item()` × 3 per + # call, so dense_R cost was dominated by Python dispatch. + R_perturbs = _rodrigues(omegas) # (B, 3, 3) + R_cand_full = R_perturbs @ R_accumulated # (B, 3, 3) + cand_peaks = [] + for R_c in R_cand_full: + a, b, g = edmonds_euler_from_rotation_matrix(R_c) + cand_peaks.append(RotationPeak( + alpha=a, beta=b, gamma=g, score=0.0, sigma=0.0, + )) + if verbose > 0: + print( + f"\nfit_to_data: dense R pass {pass_idx + 1} " + f"({n_per_axis}³={omegas.shape[0]} perturbations, " + f"±{math.degrees(max_perturb_rad):.2f}°)…", + flush=True, + ) + # batch_size adapted to N_hkl × D_grid (~41) to keep the + # llg_for_rotation_batch inner tensors bounded. + rescore_batch = max( + 4, min(100, 1_000_000 // max(hkl_keep.shape[0], 1)), + ) + timer.start("9_dense_R_rescore") + # n_D_grid=11 (vs 41 default): we only need relative LLG + # ranking across a tight rotation neighbourhood; the σA + # optimum shifts negligibly. 4× fewer Bessel evals. + rescored_refine = sim_mlrf_rescore( + cand_peaks, F_obs_amp, hkl_keep, s_mag_keep, centric_keep, + ll_refine, data.cell, + n_shells=max(n_shells // 2, 8), + n_refine=len(cand_peaks), batch_size=rescore_batch, + verbose=0, n_D_grid=11, + ) + timer.stop("9_dense_R_rescore") + top = rescored_refine[0] + best_idx = next( + i for i, p in enumerate(cand_peaks) + if p.alpha == top.alpha and p.beta == top.beta + and p.gamma == top.gamma + ) + R_accumulated = R_cand_full[best_idx] + if verbose > 0: + print( + f" pass {pass_idx + 1} best LLG={top.score:.2f}, " + f"|ω|={omegas[best_idx].norm().item() * 180 / math.pi:.3f}°", + flush=True, + ) + + # Trust the LLG. The dense-R rescore picks the perturbation with + # the highest LLG (scale-invariant ML target); LLG monotonicity + # implies a non-degraded R-work. Two solvent-aware Scaler refits + # used to live here as a defensive gate — at ~25 s each on 1DAW + # they ran twice the wall time of the entire rescore loop they + # were checking. Cheaper to trust the score. + refined = refined.rotate( + R_accumulated.T.to(model.dtype_float).contiguous(), + ) + + # --- Stage 5: joint LBFGS polish on (R, t) --- + if do_joint_refine: + timer.start("11_lbfgs_polish") + rb = RigidBodyRefinement( + refined, data, + initial_translation=torch.zeros( + 3, dtype=torch.float32, device=refined.device, + ), + expected_rotational_error=joint_refine_expected_rot_error, + max_res=joint_refine_max_res_A, + device=refined.device, + verbose=max(0, verbose - 1), + ) + rb_result = rb.refine() + with torch.no_grad(): + R_polish = rb.get_rotation_matrix().detach() + t_polish = rb.translation_frac.detach() + polished = refined.rotate(R_polish.to(model.dtype_float)) + polished = polished.translate( + t_polish.to(model.dtype_float), fractional=True, + ) + timer.stop("11_lbfgs_polish") + # Compare initial vs final R-work from the LBFGS's *own* internal + # scaler (same instance, both numbers no-solvent — apples to + # apples). Saves two solvent-aware Scaler refits (~25 s each + # on 1DAW) that previously gated this decision. + if rb_result.final_r_factor <= rb_result.initial_r_factor: + refined = polished + if verbose > 0: + print( + f"\nfit_to_data: joint polish " + f"{rb_result.initial_r_factor:.4f} → " + f"{rb_result.final_r_factor:.4f} (no-solvent R)", + flush=True, + ) + elif verbose > 0: + print( + f"\nfit_to_data: joint polish kept original " + f"({rb_result.initial_r_factor:.4f} ≤ " + f"{rb_result.final_r_factor:.4f} no-solvent R)", + flush=True, + ) + + # Single solvent-aware Scaler refit at the very end on the winner — + # gives the user-facing R-work without paying the cost on every + # intermediate gate. + timer.start("12_final_scaler") + rwork_final = _external_rwork(refined, data) + timer.stop("12_final_scaler") + + refined.last_alignment_rotation = R_recovered_best + refined.last_alignment_translation = t_refined_best + refined.last_alignment_rfactor = rwork_final + if verbose > 0: + print( + f"fit_to_data: best analytical R={r_analytic_best:.4f}, " + f"final Scaler-fit R-work={rwork_final:.4f}", + flush=True, + ) + if verbose >= 2: + print("\n" + timer.summary(), flush=True) + return refined diff --git a/torchref/alignment/ball_search.py b/torchref/alignment/ball_search.py new file mode 100644 index 00000000..6abb305c --- /dev/null +++ b/torchref/alignment/ball_search.py @@ -0,0 +1,591 @@ +""" +Pure-PyTorch Patterson-based rotation search. + +Replaces the old `ball_transform.py` (which depended on jax / s2fft / s2ball / +spherical / quaternionic). All math is implemented with `torchref.alignment.sh` +and `torchref.alignment.wigner` and runs on whatever device the inputs live on. + +Conventions (locked, asserted by tests/unit/alignment/test_*.py and the +synthetic-rotation integration test): + + Forward analytical SH expansion: + f_{p,l,m} = Σ_{i ∈ shell p} v_i · conj(Y_{l,m}(ŝ_i)) + with Friedel mates `(-s_i, v_i)` included so odd-l rows are zero by + construction. + + Cross-correlation in Wigner-D basis: + ξ_{l,m,n} = Σ_p w_p · f_{obs}[p,l,n] · conj(f_{calc}[p,l,m]) + C(R) = Σ_{l,m,n} ξ_{l,m,n} · D^l_{m,n}(R) (Edmonds D) + + Recovered rotation: the Euler triple (α, β, γ) at the maximum of C is the + rotation R = R_z(α) R_y(β) R_z(γ) (Edmonds active ZYZ) such that + s_calc = R · s_obs (column vector) + for the test scenario where F_calc was generated by applying R to the model + coordinates. To build a (3, 3) matrix using + `torchref.alignment.transform.rotation_matrix_from_euler` (which interprets + its args as `R = R_z(γ_arg) R_y(β_arg) R_z(α_arg)`), pass `[γ, β, α]`. +""" + +from __future__ import annotations + +import math +from dataclasses import dataclass +from typing import List, Optional, Tuple + +import torch +import torch.nn.functional as F + +from .sh import ( + assign_shells, + compute_patterson_shell_variance, + equal_count_shell_edges, + sh_expand_ball, +) +from .wigner import ( + evaluate_rotation_function_grid, + evaluate_rotation_function_pointwise, +) + + +# ============================================================================= +# Data containers +# ============================================================================= + + +@dataclass +class BallHarmonicCoefficients: + """Output of `compute_ball_harmonic_coefficients`.""" + + f_plm: torch.Tensor # complex, shape (P, L, 2L-1) + L: int + P: int + shell_edges: torch.Tensor # real, shape (P+1,) + shell_centers: torch.Tensor # real, shape (P,) + shell_counts: torch.Tensor # int64, shape (P,) + + +@dataclass +class RotationPeak: + """A single peak of the rotation function.""" + + alpha: float + beta: float + gamma: float + score: float # value of C at the peak (real part) + sigma: float # (C - mean) / std at the peak + + +# ============================================================================= +# Forward expansion +# ============================================================================= + + +def compute_ball_harmonic_coefficients( + s_vectors: torch.Tensor, + values: torch.Tensor, + L: int = 48, + P: int = 24, + shell_edges: Optional[torch.Tensor] = None, + d_min: Optional[float] = None, + d_max: Optional[float] = None, + enforce_friedel: bool = True, +) -> BallHarmonicCoefficients: + """ + Compute the ball-harmonic expansion of a real scattered-point field. + + Parameters + ---------- + s_vectors : (N, 3) real tensor + Reciprocal-lattice vectors in 1/Å. + values : (N,) real tensor + Sample values (e.g. |E(h)|) at each `s_vector`. + L : int, default 48 + Angular bandlimit. + P : int, default 24 + Number of radial shells. + shell_edges : (P+1,) real tensor, optional + Pre-computed shell boundaries (in |s|). If None, use equal-count + binning between `d_min` and `d_max` (or full range). + d_min, d_max : float, optional + Resolution bounds (Å). When given, reflections outside [1/d_max, 1/d_min] + in |s| are discarded. + enforce_friedel : bool, default True + Augment input with (-s, value) pairs and zero odd-l rows after expansion. + + Returns + ------- + BallHarmonicCoefficients + """ + s_mag = s_vectors.norm(dim=-1) + if d_min is not None or d_max is not None: + s_lo = 1.0 / d_max if d_max is not None else 0.0 + s_hi = 1.0 / d_min if d_min is not None else float("inf") + keep = (s_mag >= s_lo) & (s_mag <= s_hi) + s_vectors = s_vectors[keep] + values = values[keep] + s_mag = s_mag[keep] + + if shell_edges is None: + shell_edges, _ = equal_count_shell_edges(s_mag, P) + shell_centers = 0.5 * (shell_edges[:-1] + shell_edges[1:]) + shell_idx = assign_shells(s_mag, shell_edges) + keep = shell_idx >= 0 + s_vectors = s_vectors[keep] + values = values[keep] + shell_idx = shell_idx[keep] + + shell_counts = torch.bincount(shell_idx, minlength=P).to(torch.int64) + + f_plm = sh_expand_ball( + s_vectors, values, shell_idx, P, L, + enforce_friedel=enforce_friedel, + ) + + return BallHarmonicCoefficients( + f_plm=f_plm, L=L, P=P, + shell_edges=shell_edges, + shell_centers=shell_centers, + shell_counts=shell_counts, + ) + + +# ============================================================================= +# Cross-correlation +# ============================================================================= + + +def compute_ball_cross_correlation_coefficients( + f_obs: BallHarmonicCoefficients, + f_calc: BallHarmonicCoefficients, + weights: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """ + Wigner-D coefficients ξ_{l,m,n} of the rotation function (see module docstring). + + ξ_{l,m,n} = Σ_p w_p · f_{obs}[p, l, n] · conj(f_{calc}[p, l, m]) + """ + assert f_obs.L == f_calc.L, "bandwidths must match" + assert f_obs.P == f_calc.P, "shell counts must match" + P, L = f_obs.P, f_obs.L + + if weights is None: + # Default: weight by shell count (equal weight per reflection). + counts = f_obs.shell_counts.to(torch.float64) + w_sum = counts.sum().clamp(min=1.0) + weights = (counts / w_sum).to(f_obs.f_plm.dtype) + else: + weights = weights.to(f_obs.f_plm.dtype) + + xi = torch.einsum( + "p,pln,plm->lmn", + weights, + f_obs.f_plm, + torch.conj(f_calc.f_plm), + ) + return xi + + +# ============================================================================= +# Peak finding +# ============================================================================= + + +def _circular_pad_gamma_alpha(x: torch.Tensor, r: int) -> torch.Tensor: + """ + Wrap-pad axes 0 (γ) and -1 (α) by `r` voxels using circular boundary + conditions, and constant-pad axis 1 (β) with -inf (non-periodic). + Input shape (Γ, B, Α); output shape (Γ+2r, B+2r, Α+2r). + """ + if r <= 0: + return x + x = torch.cat([x[-r:], x, x[:r]], dim=0) # wrap γ + x = torch.cat([x[..., -r:], x, x[..., :r]], dim=-1) # wrap α + pad_shape = list(x.shape) + pad_shape[1] = r + neg_inf = torch.full(pad_shape, float("-inf"), dtype=x.dtype, device=x.device) + x = torch.cat([neg_inf, x, neg_inf], dim=1) # β: -inf pad + return x + + +def find_rotation_peaks( + C: torch.Tensor, + alphas: torch.Tensor, + betas: torch.Tensor, + gammas: torch.Tensor, + n_peaks: int = 200, + sigma_threshold: float = 0.0, + cluster_radius_voxels: int = 2, +) -> List[RotationPeak]: + """ + Extract local maxima of `C` and return them sorted by descending value. + + Periodic in α and γ (FFT axes); β is non-periodic but mirror at the + endpoints isn't enforced — the user is expected to oversample β. + + Patterson centrosymmetry C(R) = C(-R) is left to the caller — when needed, + add a post-processing step that merges peaks with R and -R representations. + + Vectorised over voxels: NMS is implemented as `C == max_pool3d(C)` with + `kernel = 2·r + 1`, circular in α and γ. The legacy per-voxel Python loop + with per-iteration `.item()` is replaced with a single sort + bulk + `.tolist()` at the end — same semantics on randomly-real data (exact + ties are vanishingly rare), but device-portable and ~1-2 orders of + magnitude faster. + """ + C_real = C.real + n_gamma, n_beta, n_alpha = C_real.shape + flat_stats = C_real.flatten() + + mean = flat_stats.mean() + std = flat_stats.std().clamp(min=1e-30) + threshold = mean + sigma_threshold * std + + r = int(cluster_radius_voxels) + if r > 0: + padded = _circular_pad_gamma_alpha(C_real, r) + # max_pool3d wants (N, C, D, H, W). Add batch+channel dims, then strip. + pooled = F.max_pool3d( + padded.unsqueeze(0).unsqueeze(0), + kernel_size=2 * r + 1, stride=1, padding=0, + )[0, 0] + is_local_max = C_real >= pooled + else: + is_local_max = torch.ones_like(C_real, dtype=torch.bool) + + above_threshold = C_real >= threshold + candidate = is_local_max & above_threshold + candidate_flat = candidate.flatten() + cand_idx = candidate_flat.nonzero(as_tuple=False).flatten() + if cand_idx.numel() == 0: + return [] + cand_vals = C_real.flatten().index_select(0, cand_idx) + + # Sort descending; truncate to n_peaks. + top_k = min(int(n_peaks), int(cand_vals.numel())) + sort_vals, sort_perm = torch.sort(cand_vals, descending=True) + sel_flat_idx = cand_idx.index_select(0, sort_perm[:top_k]) + sel_vals = sort_vals[:top_k] + + # Convert flat → (g, b, a). Periodic axes have already been deduped, so + # plain integer arithmetic is sufficient. + nba = n_beta * n_alpha + ig = sel_flat_idx // nba + rem = sel_flat_idx - ig * nba + ib = rem // n_alpha + ia = rem - ib * n_alpha + + sigmas = (sel_vals - mean) / std + + # Bulk transfer to Python lists, then build dataclasses. + a_list = alphas.index_select(0, ia).tolist() + b_list = betas.index_select(0, ib).tolist() + g_list = gammas.index_select(0, ig).tolist() + v_list = sel_vals.tolist() + s_list = sigmas.tolist() + return [ + RotationPeak(alpha=a, beta=b, gamma=g, score=v, sigma=s) + for a, b, g, v, s in zip(a_list, b_list, g_list, v_list, s_list) + ] + + +# ============================================================================= +# Sub-voxel refinement (pointwise gradient ascent on C(R) using autograd) +# ============================================================================= + + +def _parabolic_offset_vec( + y_minus: torch.Tensor, y_zero: torch.Tensor, y_plus: torch.Tensor, +) -> torch.Tensor: + """ + Vectorised sub-voxel offset for a quadratic through three samples + y(-1), y(0), y(+1). Returns δ clamped to [-1, 1]; degenerate cells + (|denom| < 1e-30) get δ = 0. + """ + denom = y_minus - 2.0 * y_zero + y_plus + bad = denom.abs() < 1e-30 + safe_denom = torch.where(bad, torch.ones_like(denom), denom) + delta = 0.5 * (y_minus - y_plus) / safe_denom + delta = torch.where(bad, torch.zeros_like(delta), delta) + return delta.clamp(-1.0, 1.0) + + +def refine_peaks_subvoxel( + peaks: List[RotationPeak], + C: torch.Tensor, + alphas: torch.Tensor, + betas: torch.Tensor, + gammas: torch.Tensor, +) -> List[RotationPeak]: + """ + Sub-voxel refinement of integer-grid peaks by separable quadratic fitting. + + For each peak, fits y = a + b·δ + c·δ² through the values at ±1 voxel along + each of α, β, γ axes (α and γ periodic with the grid length, β not). + Returns refined peaks with interpolated angles and values. + + Vectorised over all peaks: indices and neighbour samples are gathered in + a single batched op; the legacy per-peak `.item()` loop is replaced with + one `.tolist()` at the end (cuts ~12·K scalar GPU↔CPU round-trips). + """ + if not peaks: + return list(peaks) + + n_gamma, n_beta, n_alpha = C.shape + Cr = C.real + device = Cr.device + real_dtype = alphas.dtype + + dalpha = (alphas[1] - alphas[0]).item() + dgamma = (gammas[1] - gammas[0]).item() + dbeta = (betas[1] - betas[0]).item() if betas.numel() > 1 else 0.0 + a0 = alphas[0].item() + g0 = gammas[0].item() + b0 = betas[0].item() + + # Pack peak coordinates into tensors in one go. + K = len(peaks) + alpha_in = torch.tensor([p.alpha for p in peaks], dtype=real_dtype, device=device) + beta_in = torch.tensor([p.beta for p in peaks], dtype=real_dtype, device=device) + gamma_in = torch.tensor([p.gamma for p in peaks], dtype=real_dtype, device=device) + + # Integer voxel indices. + ia = ((alpha_in - a0) / dalpha).round().long() % n_alpha + ig = ((gamma_in - g0) / dgamma).round().long() % n_gamma + if dbeta > 0: + ib = ((beta_in - b0) / dbeta).round().long().clamp(0, n_beta - 1) + else: + ib = torch.zeros(K, dtype=torch.long, device=device) + + am = (ia - 1) % n_alpha + ap = (ia + 1) % n_alpha + gm = (ig - 1) % n_gamma + gp = (ig + 1) % n_gamma + # β neighbours: clamp at boundaries; we zero δβ explicitly there below. + ibm = (ib - 1).clamp(min=0, max=n_beta - 1) + ibp = (ib + 1).clamp(min=0, max=n_beta - 1) + + y0 = Cr[ig, ib, ia] + ya_m = Cr[ig, ib, am] + ya_p = Cr[ig, ib, ap] + yg_m = Cr[gm, ib, ia] + yg_p = Cr[gp, ib, ia] + yb_m = Cr[ig, ibm, ia] + yb_p = Cr[ig, ibp, ia] + + da = _parabolic_offset_vec(ya_m, y0, ya_p) + dg = _parabolic_offset_vec(yg_m, y0, yg_p) + db = _parabolic_offset_vec(yb_m, y0, yb_p) + # Zero out β offset at the boundary (single-sided / undefined). + at_b_boundary = (ib == 0) | (ib == n_beta - 1) + db = torch.where(at_b_boundary, torch.zeros_like(db), db) + if dbeta <= 0: + db = torch.zeros_like(db) + + two_pi = 2.0 * math.pi + alpha_new = (alphas.index_select(0, ia) + da * dalpha) % two_pi + gamma_new = (gammas.index_select(0, ig) + dg * dgamma) % two_pi + beta_new = (betas.index_select(0, ib) + db * dbeta).clamp(min=0.0, max=math.pi) + + a_list = alpha_new.tolist() + b_list = beta_new.tolist() + g_list = gamma_new.tolist() + return [ + RotationPeak( + alpha=a, beta=b, gamma=g, score=p.score, sigma=p.sigma, + ) + for a, b, g, p in zip(a_list, b_list, g_list, peaks) + ] + + +# ============================================================================= +# Top-level entry +# ============================================================================= + + +def ball_rotation_search( + s_obs: torch.Tensor, + e_obs: torch.Tensor, + s_calc: torch.Tensor, + e_calc: torch.Tensor, + L: int = 48, + P: int = 24, + n_peaks: int = 200, + d_min: Optional[float] = None, + d_max: Optional[float] = None, + refine_subvoxel: bool = True, + n_refine: int = 20, + sigma_threshold: float = 1.0, + weights: Optional[torch.Tensor] = None, + auto_variance_weights: bool = True, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, List[RotationPeak]]: + """ + End-to-end Patterson rotation search. + + Parameters + ---------- + s_obs, s_calc : (N, 3) real + Reciprocal-lattice vectors. Need not be the same set of HKLs for obs + and calc — but should sample the same resolution range. + e_obs, e_calc : (N,) real + Normalized structure-factor amplitudes (E-values). Pass |F|² instead + if you want a strict Patterson rotation function. + L, P : int, default 48 / 24 + Bandlimit and shell count. + d_min, d_max : float + Resolution limits in Å (s_min = 1/d_max, s_max = 1/d_min). + n_peaks : int, default 200 + Number of rotation peaks to return. + refine_subvoxel : bool, default True + If True, refine the top `n_refine` peaks by gradient ascent on the + pointwise rotation function (autograd, no FFT). + n_refine : int + Number of peaks to refine. + sigma_threshold : float + Minimum (peak - mean)/std for a peak to be kept. + weights : (P,) real, optional + Per-shell weights in the cross-correlation. Default: shell-count weighted, + or (when `auto_variance_weights=True`) inverse-sqrt of the empirical + per-shell variance of the obs Patterson coefficient. + auto_variance_weights : bool, default True + Compute `weights = 1/√Var(e_obs)_p` per shell when no explicit `weights` + are supplied. This is the LERF1-style empirical variance correction — + downweights shells whose observed Patterson coefficient is dominated by + intermolecular contributions (high-symmetry crystals). + + Returns + ------- + C : complex tensor (n_gamma, n_beta, n_alpha) + Rotation function on Euler grid (Edmonds order). + alphas, betas, gammas : real 1-D tensors + Grid coordinates. + peaks : list of RotationPeak (sorted by descending score) + """ + f_obs = compute_ball_harmonic_coefficients( + s_obs, e_obs, L=L, P=P, d_min=d_min, d_max=d_max, + ) + # Share shell edges between obs and calc — required for the cross-correlation + # to be meaningful (same radial binning on both sides). + f_calc = compute_ball_harmonic_coefficients( + s_calc, e_calc, L=L, P=P, shell_edges=f_obs.shell_edges, + ) + + if auto_variance_weights and weights is None: + s_mag_obs = s_obs.norm(dim=-1) + if d_min is not None or d_max is not None: + s_lo = 1.0 / d_max if d_max is not None else 0.0 + s_hi = 1.0 / d_min if d_min is not None else float("inf") + keep_obs = (s_mag_obs >= s_lo) & (s_mag_obs <= s_hi) + s_mag_v = s_mag_obs[keep_obs] + e_obs_v = e_obs[keep_obs] + else: + s_mag_v = s_mag_obs + e_obs_v = e_obs + shell_idx_v = assign_shells(s_mag_v, f_obs.shell_edges) + shell_var = compute_patterson_shell_variance( + e_obs_v.to(torch.float64), shell_idx_v, P=P, + ) + w = 1.0 / shell_var.sqrt() + w = w * (P / w.sum().clamp(min=1e-30)) + weights = w.to(f_obs.f_plm.real.dtype) + + xi = compute_ball_cross_correlation_coefficients(f_obs, f_calc, weights=weights) + C, alphas, betas, gammas = evaluate_rotation_function_grid( + xi, L, n_alpha=2 * L, n_beta=2 * L, n_gamma=2 * L, + ) + peaks = find_rotation_peaks( + C, alphas, betas, gammas, + n_peaks=n_peaks, sigma_threshold=sigma_threshold, + ) + if refine_subvoxel and len(peaks) > 0: + head = peaks[:n_refine] + head = refine_peaks_subvoxel(head, C, alphas, betas, gammas) + peaks = head + peaks[n_refine:] + peaks.sort(key=lambda r: r.score, reverse=True) + return C, alphas, betas, gammas, peaks + + +# ============================================================================= +# Convenience: Euler ↔ matrix using the Edmonds ZYZ convention +# ============================================================================= + + +def rotation_matrix_from_edmonds_euler_batch( + alpha: torch.Tensor, beta: torch.Tensor, gamma: torch.Tensor, +) -> torch.Tensor: + """ + Vectorised Edmonds ZYZ Euler → R: R = R_z(α) R_y(β) R_z(γ). + + `alpha`, `beta`, `gamma`: identically-shaped real tensors. Returns + `(..., 3, 3)` in the same dtype/device as the inputs. Equivalent to + calling `rotation_matrix_from_edmonds_euler` per Euler triple, but + avoids 9 small-tensor allocations per call. + """ + ca, sa = torch.cos(alpha), torch.sin(alpha) + cb, sb = torch.cos(beta), torch.sin(beta) + cg, sg = torch.cos(gamma), torch.sin(gamma) + zero = torch.zeros_like(alpha) + one = torch.ones_like(alpha) + Rz_a = torch.stack([ + torch.stack([ca, -sa, zero], dim=-1), + torch.stack([sa, ca, zero], dim=-1), + torch.stack([zero, zero, one], dim=-1), + ], dim=-2) + Ry_b = torch.stack([ + torch.stack([cb, zero, sb], dim=-1), + torch.stack([zero, one, zero], dim=-1), + torch.stack([-sb, zero, cb], dim=-1), + ], dim=-2) + Rz_c = torch.stack([ + torch.stack([cg, -sg, zero], dim=-1), + torch.stack([sg, cg, zero], dim=-1), + torch.stack([zero, zero, one], dim=-1), + ], dim=-2) + return Rz_a @ Ry_b @ Rz_c + + +def rotation_matrix_from_edmonds_euler( + alpha: float, beta: float, gamma: float, dtype=torch.float64, +) -> torch.Tensor: + """ + Build R = R_z(α) R_y(β) R_z(γ) (Edmonds active ZYZ). + + Equivalent to passing `[γ, β, α]` to + `torchref.alignment.transform.rotation_matrix_from_euler`. + """ + ca, sa = math.cos(alpha), math.sin(alpha) + cb, sb = math.cos(beta), math.sin(beta) + cg, sg = math.cos(gamma), math.sin(gamma) + Rz_a = torch.tensor([[ca, -sa, 0.0], [sa, ca, 0.0], [0.0, 0.0, 1.0]], dtype=dtype) + Ry_b = torch.tensor([[cb, 0.0, sb], [0.0, 1.0, 0.0], [-sb, 0.0, cb]], dtype=dtype) + Rz_c = torch.tensor([[cg, -sg, 0.0], [sg, cg, 0.0], [0.0, 0.0, 1.0]], dtype=dtype) + return Rz_a @ Ry_b @ Rz_c + + +def edmonds_euler_from_rotation_matrix(R: torch.Tensor) -> Tuple[float, float, float]: + """ + Recover (α, β, γ) such that R = R_z(α) R_y(β) R_z(γ). + Returns angles in radians; α, γ ∈ [0, 2π), β ∈ [0, π]. + Singular when β = 0 or π (only α+γ is determined); we set γ=0 in those cases. + """ + R = R.to(torch.float64) + cos_beta = R[2, 2].clamp(-1.0, 1.0).item() + beta = math.acos(cos_beta) + sin_beta = math.sin(beta) + if abs(sin_beta) < 1e-9: + # Degenerate: only α+γ is determined. Set γ=0. + alpha = math.atan2(R[1, 0].item(), R[0, 0].item()) + gamma = 0.0 + else: + alpha = math.atan2(R[1, 2].item(), R[0, 2].item()) + gamma = math.atan2(R[2, 1].item(), -R[2, 0].item()) + alpha = alpha % (2.0 * math.pi) + gamma = gamma % (2.0 * math.pi) + return alpha, beta, gamma + + +def rotation_angular_distance_deg(R1: torch.Tensor, R2: torch.Tensor) -> float: + """Geodesic distance on SO(3) in degrees: arccos((tr(R1 R2^T) - 1)/2).""" + R = R1.to(torch.float64) @ R2.to(torch.float64).T + tr = (R[0, 0] + R[1, 1] + R[2, 2]).clamp(-1.0, 3.0).item() + cos_a = max(-1.0, min(1.0, (tr - 1.0) / 2.0)) + return math.degrees(math.acos(cos_a)) diff --git a/torchref/alignment/ball_transform.py b/torchref/alignment/ball_transform.py deleted file mode 100644 index 28e31f00..00000000 --- a/torchref/alignment/ball_transform.py +++ /dev/null @@ -1,2027 +0,0 @@ -""" -Ball Harmonic Transform for Fast Rotation Function. - -Implements 3D ball transforms that preserve radial (resolution) information -for molecular replacement rotation searches. - -The ball function f(r, θ, φ) is expanded using: -- Uniform radial shells (resolution bins) -- Spherical harmonics for angular component (via s2fft) - -For rotation correlation of two ball functions f and g: - C(R) = ∫ f(x) g(R⁻¹x) dx - = Σ_{p,l,m,n} f*_{p,l,m} g_{p,l,n} D^l_{m,n}(R) - = Σ_{l,m,n} ξ_{l,m,n} D^l_{m,n}(R) - -where ξ_{l,m,n} = Σ_p w_p f*_{p,l,m} g_{p,l,n} sums over radial shells. - -Key property: Rotations only affect the angular part - radial indices are summed. -This preserves resolution information while reducing to a standard Wigner transform. -""" - -import math -from dataclasses import dataclass -from typing import Dict, List, Optional, Tuple - -import numpy as np -import torch -from numba import njit - -# Enable JAX 64-bit precision before importing s2fft/s2ball -import jax -jax.config.update("jax_enable_x64", True) - -import s2fft -import s2ball.transform.wigner as wigner_transform -import spherical -import quaternionic - - -# ============================================================================= -# Numba-accelerated Analytical Spherical Harmonic Computation -# ============================================================================= - -@njit(cache=True) -def _compute_plm_recurrence(l_max: int, cos_theta: np.ndarray) -> np.ndarray: - """ - Compute associated Legendre polynomials P_l^m(cos(theta)) using recurrence. - - Parameters - ---------- - l_max : int - Maximum l value (exclusive), i.e., computes for l = 0, 1, ..., l_max-1. - cos_theta : np.ndarray - Cosine of colatitude angles, shape (N,). - - Returns - ------- - Plm : np.ndarray - Associated Legendre polynomials, shape (N, l_max, l_max). - Plm[i, l, m] = P_l^m(cos_theta[i]) for m >= 0. - """ - N = len(cos_theta) - sin_theta = np.sqrt(1 - cos_theta**2) - Plm = np.zeros((N, l_max, l_max)) - - # P_0^0 = 1 - Plm[:, 0, 0] = 1.0 - - if l_max > 1: - # P_1^0 = cos(theta) - Plm[:, 1, 0] = cos_theta - # P_1^1 = -sin(theta) - Plm[:, 1, 1] = -sin_theta - - # Recurrence for P_l^l (diagonal) - for l in range(2, l_max): - Plm[:, l, l] = -(2*l - 1) * sin_theta * Plm[:, l-1, l-1] - - # Recurrence for P_l^{l-1} (subdiagonal) - for l in range(2, l_max): - Plm[:, l, l-1] = cos_theta * (2*l - 1) * Plm[:, l-1, l-1] - - # Recurrence for P_l^m (general) - for l in range(2, l_max): - for m in range(0, l-1): - Plm[:, l, m] = ((2*l - 1) * cos_theta * Plm[:, l-1, m] - - (l + m - 1) * Plm[:, l-2, m]) / (l - m) - - return Plm - - -@njit(cache=True) -def _compute_sh_normalization_factors(l_max: int) -> np.ndarray: - """ - Compute normalization factors for spherical harmonics. - - K_l^m = sqrt((2l+1)/(4*pi) * (l-m)!/(l+m)!) - - Parameters - ---------- - l_max : int - Maximum l value (exclusive). - - Returns - ------- - K : np.ndarray - Normalization factors, shape (l_max, l_max). - K[l, m] for m >= 0. - """ - K = np.zeros((l_max, l_max)) - for l in range(l_max): - for m in range(l + 1): - # Compute log((l-m)!/(l+m)!) for numerical stability - log_factor = 0.0 - for k in range(l - m + 1, l + m + 1): - log_factor -= np.log(k) - K[l, m] = np.sqrt((2*l + 1) / (4 * np.pi) * np.exp(log_factor)) - return K - - -@njit(cache=True) -def _compute_sh_coeffs_analytical_numba( - E: np.ndarray, - cos_theta: np.ndarray, - phi: np.ndarray, - L: int, -) -> np.ndarray: - """ - Compute spherical harmonic coefficients analytically using numba. - - Computes a_lm = (4*pi/N) * sum_i(E_i * conj(Y_lm(theta_i, phi_i))) - - where Y_lm = K_l^m * P_l^m(cos(theta)) * exp(i*m*phi) - - Parameters - ---------- - E : np.ndarray - Function values at sample points, shape (N,). - cos_theta : np.ndarray - Cosine of colatitude angles, shape (N,). - phi : np.ndarray - Azimuthal angles in [0, 2*pi), shape (N,). - L : int - Angular bandlimit. - - Returns - ------- - flm : np.ndarray - Spherical harmonic coefficients, shape (L, 2*L-1). - flm[l, m + L - 1] = a_lm for m in [-l, l]. - """ - N = len(E) - - # Compute associated Legendre polynomials - Plm = _compute_plm_recurrence(L, cos_theta) - - # Compute normalization factors - K = _compute_sh_normalization_factors(L) - - # Precompute exp(i*m*phi) for all m values - exp_imphi = np.zeros((N, 2*L - 1), dtype=np.complex128) - for m in range(-(L-1), L): - exp_imphi[:, m + L - 1] = np.cos(m * phi) + 1j * np.sin(m * phi) - - # Compute coefficients - flm = np.zeros((L, 2*L - 1), dtype=np.complex128) - - for l in range(L): - for m in range(-l, l + 1): - # Compute Y_lm - if m >= 0: - Y_lm = K[l, m] * Plm[:, l, m] * exp_imphi[:, m + L - 1] - else: - # Y_l^{-|m|} = (-1)^|m| * conj(Y_l^|m|) - abs_m = -m - sign = 1.0 if abs_m % 2 == 0 else -1.0 - Y_lm = sign * K[l, abs_m] * Plm[:, l, abs_m] * np.conj(exp_imphi[:, abs_m + L - 1]) - - # a_lm = (4*pi/N) * sum(E * conj(Y_lm)) - a_lm = 0.0 + 0.0j - for i in range(N): - a_lm += E[i] * np.conj(Y_lm[i]) - flm[l, m + L - 1] = (4 * np.pi / N) * a_lm - - return flm - - - -# ============================================================================= -# Ball Harmonic Coefficient Container -# ============================================================================= - -@dataclass -class BallHarmonicCoefficients: - """ - Container for ball harmonic coefficients. - - Attributes - ---------- - flmp : np.ndarray - Spherical harmonic coefficients for each radial shell, shape (P, L, 2L-1). - L : int - Angular bandlimit. - P : int - Number of radial shells. - shell_edges : np.ndarray - Radial shell boundaries in Å⁻¹, shape (P+1,). - shell_centers : np.ndarray - Radial shell centers in Å⁻¹, shape (P,). - shell_counts : np.ndarray - Number of reflections in each shell, shape (P,). - """ - flmp: np.ndarray - L: int - P: int - shell_edges: np.ndarray - shell_centers: np.ndarray - shell_counts: np.ndarray - - -def _compute_radial_shells_np( - d_min: float, - d_max: float, - P: int, -) -> Tuple[np.ndarray, np.ndarray]: - """ - Compute uniform radial shell boundaries in reciprocal space. - - Parameters - ---------- - d_min : float - High resolution limit in Angstrom. - d_max : float - Low resolution limit in Angstrom. - P : int - Number of radial shells. - - Returns - ------- - shell_edges : np.ndarray - Shell boundaries in Angstrom^-1, shape (P+1,). - shell_centers : np.ndarray - Shell centers in Angstrom^-1, shape (P,). - """ - s_min = 1.0 / d_max # Low resolution end - s_max = 1.0 / d_min # High resolution end - - shell_edges = np.linspace(s_min, s_max, P + 1) - shell_centers = 0.5 * (shell_edges[:-1] + shell_edges[1:]) - - return shell_edges, shell_centers - - -def _compute_equal_count_shells_np( - s_mag: np.ndarray, - P: int, - d_min: float, - d_max: float, -) -> Tuple[np.ndarray, np.ndarray]: - """ - Compute radial shell boundaries such that each shell has equal reflection count. - - This addresses the issue that uniform spacing in s leads to highly imbalanced - shell counts (low resolution shells have few reflections, high resolution - shells have many). Equal-count binning ensures each shell contributes equally - to the spherical harmonic expansion. - - Parameters - ---------- - s_mag : np.ndarray - Magnitude of s-vectors (|s| = 1/d), shape (N,). - P : int - Number of radial shells. - d_min : float - High resolution limit in Angstrom. - d_max : float - Low resolution limit in Angstrom. - - Returns - ------- - shell_edges : np.ndarray - Shell boundaries in Angstrom^-1, shape (P+1,). - shell_centers : np.ndarray - Shell centers in Angstrom^-1, shape (P,). - """ - s_min = 1.0 / d_max # Low resolution end - s_max = 1.0 / d_min # High resolution end - - # Filter to resolution range - mask = (s_mag >= s_min) & (s_mag <= s_max) - s_in_range = s_mag[mask] - - if len(s_in_range) == 0: - # Fallback to uniform if no reflections in range - return _compute_radial_shells_np(d_min, d_max, P) - - # Compute quantiles for equal-count binning - # We want P bins, so we need P+1 edges at quantiles 0, 1/P, 2/P, ..., 1 - quantiles = np.linspace(0, 1, P + 1) - shell_edges = np.quantile(s_in_range, quantiles) - - # Ensure edges are strictly within bounds - shell_edges[0] = max(shell_edges[0], s_min) - shell_edges[-1] = min(shell_edges[-1], s_max) - - # Ensure strictly increasing (can happen with many identical values) - for i in range(1, len(shell_edges)): - if shell_edges[i] <= shell_edges[i-1]: - shell_edges[i] = shell_edges[i-1] + 1e-10 - - shell_centers = 0.5 * (shell_edges[:-1] + shell_edges[1:]) - - return shell_edges, shell_centers - - -def get_mw_grid(L: int) -> Tuple[np.ndarray, np.ndarray]: - """ - Get McEwen-Wiaux (MW) sampling grid positions. - - Parameters - ---------- - L : int - Angular bandlimit. - - Returns - ------- - thetas : np.ndarray - Colatitude samples in [0, π], shape (L,). - phis : np.ndarray - Azimuth samples in [0, 2π), shape (2L-1,). - """ - thetas = np.array([(2*t + 1) * np.pi / (2*L) for t in range(L)]) - phis = np.array([2 * np.pi * p / (2*L - 1) for p in range(2*L - 1)]) - return thetas, phis - - -def splat_to_mw_grid( - theta: np.ndarray, - phi: np.ndarray, - values: np.ndarray, - L: int, - mean_center: bool = True, -) -> np.ndarray: - """ - Splat values onto MW sampling grid using bilinear interpolation. - - Parameters - ---------- - theta : np.ndarray - Colatitude angles in [0, π], shape (N,). - phi : np.ndarray - Azimuthal angles in [0, 2π), shape (N,). - values : np.ndarray - Values to splat, shape (N,). - L : int - Angular bandlimit. - mean_center : bool - If True, subtract mean from grid. - - Returns - ------- - grid : np.ndarray - Splatted grid of shape (L, 2L-1). - """ - n_theta = L - n_phi = 2 * L - 1 - - grid = np.zeros((n_theta, n_phi), dtype=np.float64) - weights = np.zeros((n_theta, n_phi), dtype=np.float64) - - # MW grid: theta_t = (2t + 1) * π / (2L) - theta_px = (theta * 2 * L / np.pi - 1) / 2 - phi_px = phi * (2 * L - 1) / (2 * np.pi) - - theta_px = np.clip(theta_px, 0, n_theta - 1 - 1e-6) - - theta_lo = np.floor(theta_px).astype(int) - phi_lo = np.floor(phi_px).astype(int) - - theta_frac = theta_px - theta_lo - phi_frac = phi_px - phi_lo - - theta_hi = np.clip(theta_lo + 1, 0, n_theta - 1) - theta_lo = np.clip(theta_lo, 0, n_theta - 1) - phi_hi = (phi_lo + 1) % n_phi - phi_lo = phi_lo % n_phi - - w00 = (1 - theta_frac) * (1 - phi_frac) - w01 = (1 - theta_frac) * phi_frac - w10 = theta_frac * (1 - phi_frac) - w11 = theta_frac * phi_frac - - np.add.at(grid, (theta_lo, phi_lo), w00 * values) - np.add.at(grid, (theta_lo, phi_hi), w01 * values) - np.add.at(grid, (theta_hi, phi_lo), w10 * values) - np.add.at(grid, (theta_hi, phi_hi), w11 * values) - - np.add.at(weights, (theta_lo, phi_lo), w00) - np.add.at(weights, (theta_lo, phi_hi), w01) - np.add.at(weights, (theta_hi, phi_lo), w10) - np.add.at(weights, (theta_hi, phi_hi), w11) - - mask = weights > 0 - grid[mask] /= weights[mask] - - if mean_center: - grid = grid - grid.mean() - - return grid - - -def _assign_to_shells_np( - s_mag: np.ndarray, - shell_edges: np.ndarray, -) -> np.ndarray: - """ - Internal NumPy version for ball-specific functions. - """ - shell_idx = np.digitize(s_mag, shell_edges) - 1 - P = len(shell_edges) - 1 - # Mark out-of-range as -1 - shell_idx[(shell_idx < 0) | (shell_idx >= P)] = -1 - return shell_idx - - -def splat_evalues_to_ball( - E_values: np.ndarray, - s_vectors: np.ndarray, - L: int, - P: int, - d_min: float, - d_max: float, - mean_center_shells: bool = True, -) -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: - """ - Splat E-values onto a 3D ball grid using uniform radial shells. - - Parameters - ---------- - E_values : np.ndarray - E² values, shape (N,). - s_vectors : np.ndarray - Reciprocal space vectors in Å⁻¹, shape (N, 3). - L : int - Angular bandlimit. - P : int - Number of radial shells. - d_min : float - High resolution limit in Å. - d_max : float - Low resolution limit in Å. - mean_center_shells : bool - If True, mean-center each radial shell. - - Returns - ------- - ball_grid : np.ndarray - 3D ball grid of shape (P, L, 2L-1). - shell_edges : np.ndarray - Shell boundaries in Å⁻¹. - shell_centers : np.ndarray - Shell centers in Å⁻¹. - shell_counts : np.ndarray - Number of reflections per shell. - """ - # Compute uniform radial shells - shell_edges, shell_centers = _compute_radial_shells_np(d_min, d_max, P) - - # Compute |s| and angles - s_mag = np.linalg.norm(s_vectors, axis=1) - s_normed = s_vectors / np.maximum(s_mag[:, np.newaxis], 1e-10) - - # Spherical angles - theta = np.arccos(np.clip(s_normed[:, 2], -1, 1)) - phi = np.arctan2(s_normed[:, 1], s_normed[:, 0]) - phi = phi % (2 * np.pi) - - # Assign to shells - shell_idx = _assign_to_shells_np(s_mag, shell_edges) - - # Initialize ball grid - ball_grid = np.zeros((P, L, 2 * L - 1), dtype=np.float64) - shell_counts = np.zeros(P, dtype=np.int64) - - # Splat each shell - for p in range(P): - mask = shell_idx == p - count = mask.sum() - shell_counts[p] = count - - if count == 0: - continue - - ball_grid[p] = splat_to_mw_grid( - theta[mask], - phi[mask], - E_values[mask], - L, - mean_center=mean_center_shells, - ) - - return ball_grid, shell_edges, shell_centers, shell_counts - - -def compute_ball_harmonic_coefficients( - ball_grid: np.ndarray, - L: int, - shell_edges: np.ndarray, - shell_centers: np.ndarray, - shell_counts: np.ndarray, -) -> BallHarmonicCoefficients: - """ - Compute spherical harmonic coefficients for each radial shell using s2fft. - - Parameters - ---------- - ball_grid : np.ndarray - 3D ball grid of shape (P, L, 2L-1). - L : int - Angular bandlimit. - shell_edges : np.ndarray - Shell boundaries. - shell_centers : np.ndarray - Shell centers. - shell_counts : np.ndarray - Number of reflections per shell. - - Returns - ------- - coeffs : BallHarmonicCoefficients - Ball harmonic coefficients (SH coeffs for each shell). - """ - P = ball_grid.shape[0] - - # Compute SH coefficients for each shell using s2fft - flmp = np.zeros((P, L, 2*L - 1), dtype=np.complex128) - - for p in range(P): - if shell_counts[p] > 0: - # Use s2fft forward transform: grid -> SH coefficients - flmp[p] = s2fft.forward( - ball_grid[p], - L, - sampling="mw", - method="jax", - reality=False, - ) - - return BallHarmonicCoefficients( - flmp=flmp, - L=L, - P=P, - shell_edges=shell_edges, - shell_centers=shell_centers, - shell_counts=shell_counts, - ) - - -def compute_ball_harmonic_coefficients_analytical( - E_values: np.ndarray, - s_vectors: np.ndarray, - L: int, - P: int, - d_min: float, - d_max: float, - normalize_shells: bool = True, - equal_count_shells: bool = True, -) -> BallHarmonicCoefficients: - """ - Compute ball harmonic coefficients analytically without splatting. - - This method computes SH coefficients directly from the reflection positions - using the analytical formula: - a_lm = (4*pi/N) * sum_i(E_i * conj(Y_lm(theta_i, phi_i))) - - This avoids discretization errors from splatting onto a grid and is - significantly faster when using numba acceleration. - - Parameters - ---------- - E_values : np.ndarray - E-values (normalized structure factor amplitudes), shape (N,). - s_vectors : np.ndarray - Reciprocal space vectors in Angstrom^-1, shape (N, 3). - L : int - Angular bandlimit. - P : int - Number of radial shells. - d_min : float - High resolution limit in Angstrom. - d_max : float - Low resolution limit in Angstrom. - normalize_shells : bool - If True, mean-center and variance-normalize E-values in each shell. - This removes DC bias and makes correlations comparable across shells. - equal_count_shells : bool - If True (default), compute shell boundaries such that each shell has - approximately the same number of reflections. This ensures equal - contribution from each resolution shell to the spherical harmonic - expansion. If False, use uniform spacing in s (which leads to - imbalanced shell counts due to increasing reflection density at - higher resolution). - - Returns - ------- - coeffs : BallHarmonicCoefficients - Ball harmonic coefficients computed analytically. - """ - # Compute s-vector magnitudes (needed for shell assignment) - s_mag = np.linalg.norm(s_vectors, axis=1) - - # Compute radial shells - if equal_count_shells: - shell_edges, shell_centers = _compute_equal_count_shells_np(s_mag, P, d_min, d_max) - else: - shell_edges, shell_centers = _compute_radial_shells_np(d_min, d_max, P) - - # Compute normalized s-vectors for angular coordinates - s_normed = s_vectors / np.maximum(s_mag[:, np.newaxis], 1e-10) - - # Spherical angles - theta = np.arccos(np.clip(s_normed[:, 2], -1, 1)) - cos_theta = np.cos(theta) - phi = np.arctan2(s_normed[:, 1], s_normed[:, 0]) - phi = phi % (2 * np.pi) - - # Assign to shells - shell_idx = _assign_to_shells_np(s_mag, shell_edges) - - # Initialize coefficient arrays - flm = np.zeros((P, L, 2*L - 1), dtype=np.complex128) - shell_counts = np.zeros(P, dtype=np.int64) - - # Compute SH coefficients for each shell using numba-accelerated function - for p in range(P): - mask = shell_idx == p - count = mask.sum() - shell_counts[p] = count - - if count == 0: - continue - - E_shell = E_values[mask].copy() - cos_theta_shell = cos_theta[mask] - phi_shell = phi[mask] - - if normalize_shells: - # Mean-center (removes DC bias) - E_shell = E_shell - E_shell.mean() - # Variance-normalize (makes correlations comparable across shells) - E_std = E_shell.std() - if E_std > 1e-10: - E_shell = E_shell / E_std - - # Use numba-accelerated analytical SH computation - flm[p] = _compute_sh_coeffs_analytical_numba(E_shell, cos_theta_shell, phi_shell, L) - - return BallHarmonicCoefficients( - flmp=flm, - L=L, - P=P, - shell_edges=shell_edges, - shell_centers=shell_centers, - shell_counts=shell_counts, - ) - - -def compute_ball_cross_correlation_coefficients( - f_coeffs: BallHarmonicCoefficients, - g_coeffs: BallHarmonicCoefficients, - radial_weights: Optional[np.ndarray] = None, -) -> np.ndarray: - """ - Compute Wigner coefficients for ball cross-correlation. - - The cross-correlation is: - C(R) = Σ_{p,l,m,n} f*_{p,l,m} g_{p,l,n} D^l_{m,n}(R) - = Σ_{l,m,n} ξ_{l,m,n} D^l_{m,n}(R) - - where ξ_{l,m,n} = Σ_p w_p f*_{p,l,m} g_{p,l,n} - - Parameters - ---------- - f_coeffs : BallHarmonicCoefficients - Ball harmonic coefficients of function f (observed). - g_coeffs : BallHarmonicCoefficients - Ball harmonic coefficients of function g (calculated). - radial_weights : np.ndarray, optional - Weights for each radial shell, shape (P,). - Default: uniform weights based on shell counts. - - Returns - ------- - xi_nlm : np.ndarray - Wigner coefficients, shape (2N-1, L, 2L-1) where N=L. - """ - assert f_coeffs.L == g_coeffs.L, "Angular bandlimits must match" - assert f_coeffs.P == g_coeffs.P, "Radial bandlimits must match" - - L = f_coeffs.L - P = f_coeffs.P - N = L - - if radial_weights is None: - # Weight by number of reflections in each shell (normalized) - radial_weights = f_coeffs.shell_counts.astype(np.float64) - radial_weights = np.where(radial_weights > 0, radial_weights, 0) - - # Normalize weights - weight_sum = radial_weights.sum() - if weight_sum > 0: - radial_weights = radial_weights / weight_sum - else: - radial_weights = np.ones(P) / P - - # Initialize Wigner coefficients: (2N-1, L, 2L-1) = [n_idx, l, m_idx] - xi_nlm = np.zeros((2*N - 1, L, 2*L - 1), dtype=np.complex128) - - # f_coeffs.flmp has shape (P, L, 2L-1) - # Sum over radial index p - # - # For cross-correlation C(R) that finds rotation R such that g(R⁻¹x) ≈ f(x): - # C(R) = ∫ f(x) g(R⁻¹x) dx = Σ f*_{lm} g_{ln} D^l_{mn}(R) - # - # But D^l_{mn}(R) convention in s2ball gives R^T, so we swap f↔g: - # ξ[n_idx, l, m_idx] = Σ_p w_p * conj(g[l, m_idx]) * f[l, n_idx] - # - for p in range(P): - w = radial_weights[p] - f_lm = f_coeffs.flmp[p] # (L, 2L-1) - g_lm = g_coeffs.flmp[p] # (L, 2L-1) - - # Swap f and g to get R instead of R^T - g_conj = np.conj(g_lm) # (L, 2L-1) - - # For s2fft/s2ball wigner convention: - # coeffs[n_idx, l, m_idx] corresponds to D^l_{m,n} - # where m = m_idx - (L-1), n = n_idx - (N-1) - - # Build the product: for each l, compute outer product over m and n - for l in range(L): - # Valid m range: -l to l, i.e., m_idx from L-1-l to L-1+l - m_start = L - 1 - l - m_end = L - 1 + l + 1 - - # Valid n range: -l to l, i.e., n_idx from N-1-l to N-1+l - n_start = N - 1 - l - n_end = N - 1 + l + 1 - - # g_conj[l, m_start:m_end] shape: (2l+1,) - # f_lm[l, n_start:n_end] shape: (2l+1,) - g_l = g_conj[l, m_start:m_end] # (2l+1,) - conjugated g for m index - f_l = f_lm[l, n_start:n_end] # (2l+1,) - f for n index - - # Outer product: (2l+1, 2l+1) -> [n, m] - outer = np.outer(f_l, g_l) # (2l+1, 2l+1) = [n, m] - - xi_nlm[n_start:n_end, l, m_start:m_end] += w * outer - - return xi_nlm - - -def evaluate_rotation_function( - xi_nlm: np.ndarray, - L: int, -) -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: - """ - Evaluate rotation function from Wigner coefficients using s2ball. - - Parameters - ---------- - xi_nlm : np.ndarray - Wigner coefficients, shape (2N-1, L, 2L-1). - L : int - Angular bandlimit. - - Returns - ------- - rotation_function : np.ndarray - Rotation function, shape (2N-1, L, 2L-1) = (gamma, beta, alpha). - alphas : np.ndarray - Alpha angle grid. - betas : np.ndarray - Beta angle grid. - gammas : np.ndarray - Gamma angle grid. - """ - N = L - - # Inverse Wigner transform - rotation_function = wigner_transform.inverse( - xi_nlm, - L=L, - N=N, - method="jax", - ) - - # MW sampling grid positions - betas = np.array([(2*t + 1) * np.pi / (2*L) for t in range(L)]) - alphas = np.array([2 * np.pi * p / (2*L - 1) for p in range(2*L - 1)]) - gammas = np.array([2 * np.pi * p / (2*N - 1) for p in range(2*N - 1)]) - - return np.asarray(rotation_function), alphas, betas, gammas - - -def ball_rotation_search( - E_obs: np.ndarray, - s_obs: np.ndarray, - E_calc: np.ndarray, - s_calc: np.ndarray, - L: int = 32, - P: int = 20, - d_min: float = 4.0, - d_max: float = 50.0, - n_peaks: int = 100, - radial_weights: Optional[np.ndarray] = None, - refine_subvoxel: bool = True, - refine_analytical: bool = False, - analytical_embedding: bool = True, - equal_count_shells: bool = True, - return_coefficients: bool = False, - verbose: bool = True, -) -> Tuple[np.ndarray, tuple, list]: - """ - Perform ball harmonic rotation function search. - - Main entry point for the ball-based fast rotation function. - - Parameters - ---------- - E_obs : np.ndarray - Observed E² values, shape (N_obs,). - s_obs : np.ndarray - Observed s-vectors in Å⁻¹, shape (N_obs, 3). - E_calc : np.ndarray - Calculated E² values, shape (N_calc,). - s_calc : np.ndarray - Calculated s-vectors in Å⁻¹, shape (N_calc, 3). - L : int - Angular bandlimit. - P : int - Radial bandlimit. - d_min : float - High resolution limit in Å. - d_max : float - Low resolution limit in Å. - n_peaks : int - Number of peaks to extract. - radial_weights : np.ndarray, optional - Weights for radial shells. - refine_subvoxel : bool - If True, refine peak positions using fast quadratic fitting. - Ignored if refine_analytical=True. - refine_analytical : bool - If True, refine peak positions using analytical Wigner D-matrix - evaluation. This is more accurate but slower than subvoxel refinement. - Provides exact sub-grid positions for the band-limited rotation function. - analytical_embedding : bool - If True (default), compute spherical harmonic coefficients analytically - from reflection positions using numba-accelerated computation. This - avoids discretization errors from splatting and is significantly faster. - If False, use the original splatting approach. - equal_count_shells : bool - If True (default), compute radial shell boundaries such that each shell - has approximately the same number of reflections. This ensures equal - contribution from each resolution shell. Only applies when - analytical_embedding=True. - return_coefficients : bool - If True, also return the Wigner coefficients xi_nlm for later use - (e.g., for rescoring or additional analytical refinement). - verbose : bool - Print progress. - - Returns - ------- - rotation_function : np.ndarray - Full rotation function, shape (2L-1, L, 2L-1). - angles_grid : tuple - (alphas, betas, gammas) angle grids. - peaks : list - List of (alpha, beta, gamma, score, sigma) tuples. - xi_nlm : np.ndarray, optional - Wigner coefficients, shape (2N-1, L, 2L-1). Only returned if - return_coefficients=True. - """ - import time - - start_time = time.time() - - if verbose: - print(f"Ball rotation search: L={L}, P={P}") - print(f"Resolution range: {d_min:.2f} - {d_max:.2f} Å") - print(f"Embedding method: {'analytical' if analytical_embedding else 'splatting'}") - - if analytical_embedding: - # Analytical approach: compute SH coefficients directly from reflection positions - # This avoids discretization errors from splatting and is much faster - if verbose: - print("Computing analytical spherical harmonic coefficients...") - if equal_count_shells: - print(" Using equal-count shell binning") - - coeffs_obs = compute_ball_harmonic_coefficients_analytical( - E_obs, s_obs, L, P, d_min, d_max, - normalize_shells=True, equal_count_shells=equal_count_shells - ) - coeffs_calc = compute_ball_harmonic_coefficients_analytical( - E_calc, s_calc, L, P, d_min, d_max, - normalize_shells=True, equal_count_shells=equal_count_shells - ) - - if verbose: - print(f" Coefficients shape: {coeffs_obs.flmp.shape}") - print(f" Reflections per shell (obs): min={coeffs_obs.shell_counts.min()}, max={coeffs_obs.shell_counts.max()}") - - else: - # Original splatting approach - # Step 1: Splat E-values onto ball grids (uniform radial shells) - if verbose: - print("Splatting E-values onto ball grid...") - - ball_obs, shell_edges, shell_centers, shell_counts_obs = splat_evalues_to_ball( - E_obs, s_obs, L, P, d_min, d_max - ) - ball_calc, _, _, shell_counts_calc = splat_evalues_to_ball( - E_calc, s_calc, L, P, d_min, d_max - ) - - if verbose: - print(f" Ball grid shape: {ball_obs.shape}") - print(f" Shell range: [{shell_edges[0]:.4f}, {shell_edges[-1]:.4f}] Å⁻¹") - print(f" Reflections per shell (obs): min={shell_counts_obs.min()}, max={shell_counts_obs.max()}") - - # Step 2: Compute spherical harmonic coefficients for each shell - if verbose: - print("Computing spherical harmonic coefficients per shell...") - - coeffs_obs = compute_ball_harmonic_coefficients( - ball_obs, L, shell_edges, shell_centers, shell_counts_obs - ) - coeffs_calc = compute_ball_harmonic_coefficients( - ball_calc, L, shell_edges, shell_centers, shell_counts_calc - ) - - if verbose: - print(f" Coefficients shape: {coeffs_obs.flmp.shape}") - - # Step 3: Compute cross-correlation Wigner coefficients - if verbose: - print("Computing cross-correlation coefficients...") - - xi_nlm = compute_ball_cross_correlation_coefficients( - coeffs_obs, coeffs_calc, radial_weights - ) - - if verbose: - print(f" Wigner coefficients shape: {xi_nlm.shape}") - print(f" Max |ξ|: {np.abs(xi_nlm).max():.6e}") - - # Step 4: Evaluate rotation function - if verbose: - print("Evaluating rotation function via inverse Wigner transform...") - - rotation_function, alphas, betas, gammas = evaluate_rotation_function(xi_nlm, L) - - rf_real = np.real(rotation_function) - - if verbose: - print(f" Rotation function shape: {rf_real.shape}") - print(f" RF range: [{rf_real.min():.4f}, {rf_real.max():.4f}]") - print(f" RF mean: {rf_real.mean():.4f}, std: {rf_real.std():.4f}") - - # Step 5: Find peaks - if verbose: - print("Finding peaks...") - - peaks = find_rotation_peaks(rf_real, alphas, betas, gammas, n_peaks=n_peaks) - - # Step 6: Optionally refine peaks - if peaks: - if refine_analytical: - # Analytical refinement using Wigner D-matrices (more accurate) - if verbose: - print("Refining peaks using analytical Wigner D-matrices...") - - rf_mean = rf_real.mean() - rf_std = rf_real.std() - peaks = refine_peaks_analytical( - peaks, xi_nlm, L, rf_mean=rf_mean, rf_std=rf_std, verbose=verbose - ) - elif refine_subvoxel: - # Fast quadratic subvoxel refinement - if verbose: - print("Refining peaks to sub-voxel accuracy...") - - peaks = refine_peaks_subvoxel_wrapper( - peaks, rf_real, alphas, betas, gammas - ) - - if verbose: - print(f" Refined {len(peaks)} peaks") - - elapsed = time.time() - start_time - if verbose: - print(f"Ball rotation search completed in {elapsed:.2f}s") - if peaks: - print(f"Top peak: alpha={np.degrees(peaks[0][0]):.2f}°, " - f"beta={np.degrees(peaks[0][1]):.2f}°, " - f"gamma={np.degrees(peaks[0][2]):.2f}°, " - f"sigma={peaks[0][4]:.2f}") - - if return_coefficients: - return rf_real, (alphas, betas, gammas), peaks, xi_nlm - return rf_real, (alphas, betas, gammas), peaks - - -def refine_peaks_subvoxel_wrapper( - peaks: list, - rotation_function: np.ndarray, - alphas: np.ndarray, - betas: np.ndarray, - gammas: np.ndarray, -) -> list: - """ - Refine peak positions to sub-voxel accuracy using quadratic fitting. - - Parameters - ---------- - peaks : list - List of (alpha, beta, gamma, score, sigma) tuples. - rotation_function : np.ndarray - Rotation function grid, shape (n_gamma, n_beta, n_alpha). - alphas, betas, gammas : np.ndarray - Angle grids. - - Returns - ------- - refined_peaks : list - List of (alpha, beta, gamma, score, sigma) tuples with refined positions. - """ - import jax.numpy as jnp - from torchref.alignment.jax_subpixel_peaks import refine_peaks_subvoxel - - if not peaks: - return peaks - - n_gamma, n_beta, n_alpha = rotation_function.shape - - # Compute angle spacings - d_alpha = alphas[1] - alphas[0] if len(alphas) > 1 else 2 * np.pi / n_alpha - d_beta = betas[1] - betas[0] if len(betas) > 1 else np.pi / n_beta - d_gamma = gammas[1] - gammas[0] if len(gammas) > 1 else 2 * np.pi / n_gamma - - # Convert peaks to grid indices - peak_indices = [] - for alpha, beta, gamma, score, sigma in peaks: - # Find nearest grid indices - a_idx = int(round((alpha - alphas[0]) / d_alpha)) % n_alpha - b_idx = int(round((beta - betas[0]) / d_beta)) - b_idx = max(0, min(b_idx, n_beta - 1)) - g_idx = int(round((gamma - gammas[0]) / d_gamma)) % n_gamma - - peak_indices.append([g_idx, b_idx, a_idx]) # shape matches rf: (gamma, beta, alpha) - - peak_indices = jnp.array(peak_indices, dtype=jnp.int32) - grid_jax = jnp.array(rotation_function) - - # Call JAX subvoxel refinement - refined_coords, refined_values = refine_peaks_subvoxel(grid_jax, peak_indices) - - # Convert back to numpy - refined_coords = np.array(refined_coords) - refined_values = np.array(refined_values) - - # Compute mean and std for sigma calculation - rf_mean = rotation_function.mean() - rf_std = rotation_function.std() - - # Convert refined grid indices back to angles - refined_peaks = [] - for i, ((alpha_orig, beta_orig, gamma_orig, score_orig, sigma_orig), refined_val) in enumerate( - zip(peaks, refined_values) - ): - g_refined, b_refined, a_refined = refined_coords[i] - - # Convert to angles (with periodic wrapping for alpha and gamma) - alpha_new = alphas[0] + a_refined * d_alpha - alpha_new = alpha_new % (2 * np.pi) - - beta_new = betas[0] + b_refined * d_beta - beta_new = np.clip(beta_new, 0, np.pi) - - gamma_new = gammas[0] + g_refined * d_gamma - gamma_new = gamma_new % (2 * np.pi) - - # Compute refined sigma - sigma_new = (refined_val - rf_mean) / rf_std if rf_std > 1e-10 else 0.0 - - refined_peaks.append((alpha_new, beta_new, gamma_new, float(refined_val), float(sigma_new))) - - # Note: We do NOT re-sort by refined score because the quadratic interpolation - # gives an approximate value, not the true score. Keeping original order - # preserves the ranking from the grid-based search. - return refined_peaks - - -# ============================================================================= -# Analytical Peak Refinement using Wigner D-matrices -# ============================================================================= - -def build_wigner_index_mapping(L: int) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: - """ - Build index mapping from xi_nlm array to Wigner D-matrix indices. - - Pre-computes the mapping between the xi_nlm coefficient array - (shape: 2N-1, L, 2L-1) and the flattened Wigner D-matrix from spherical. - - Parameters - ---------- - L : int - Angular bandlimit. - - Returns - ------- - xi_indices : np.ndarray - Array of (n_idx, l, m_idx) indices, shape (n_terms, 3). - D_indices : np.ndarray - Corresponding indices into spherical's D array, shape (n_terms,). - l_values : np.ndarray - The l value for each term, shape (n_terms,). - """ - N = L - xi_indices = [] - D_indices = [] - l_values = [] - - for ell in range(L): - for m in range(-ell, ell + 1): - m_idx = m + (L - 1) - for n in range(-ell, ell + 1): - n_idx = n + (N - 1) - D_idx = spherical.WignerDindex(ell, m, n) - xi_indices.append((n_idx, ell, m_idx)) - D_indices.append(D_idx) - l_values.append(ell) - - return (np.array(xi_indices, dtype=np.int32), - np.array(D_indices, dtype=np.int32), - np.array(l_values, dtype=np.int32)) - - -def evaluate_rotation_function_at_angles( - xi_nlm: np.ndarray, - alpha: float, - beta: float, - gamma: float, - wigner: spherical.Wigner, - xi_indices: np.ndarray, - D_indices: np.ndarray, - l_values: np.ndarray, -) -> float: - """ - Evaluate rotation function at arbitrary Euler angles using Wigner D-matrices. - - The s2ball inverse Wigner transform uses the normalization: - RF(α,β,γ) = Σ_{l,m,n} ξ_{l,m,n} * [(2l+1)/(8π²)] * D^l_{m,n}(α,β,γ) - - This allows exact evaluation of the band-limited rotation function at - any point, not just grid points. - - Parameters - ---------- - xi_nlm : np.ndarray - Wigner coefficients from compute_ball_cross_correlation_coefficients(), - shape (2N-1, L, 2L-1). - alpha, beta, gamma : float - ZYZ Euler angles in radians. - wigner : spherical.Wigner - Pre-initialized Wigner calculator. - xi_indices, D_indices : np.ndarray - Pre-computed index mappings from build_wigner_index_mapping(). - l_values : np.ndarray - The l value for each term, shape (n_terms,). - - Returns - ------- - float - Rotation function value (real part). - """ - R = quaternionic.array.from_euler_angles(alpha, beta, gamma) - D_all = wigner.D(R) - - xi_flat = xi_nlm[xi_indices[:, 0], xi_indices[:, 1], xi_indices[:, 2]] - D_flat = D_all[D_indices] - - # Apply the s2ball normalization factor (2l+1)/(8π²) for each term - norm_factors = (2 * l_values + 1) / (8 * np.pi**2) - - return np.sum(xi_flat * D_flat * norm_factors).real - - -def refine_peaks_analytical( - peaks: list, - xi_nlm: np.ndarray, - L: int, - rf_mean: float = None, - rf_std: float = None, - verbose: bool = False, -) -> list: - """ - Refine peak positions using analytical Wigner D-matrix evaluation. - - This provides exact sub-grid peak positions by evaluating the rotation - function directly from its Wigner coefficient expansion, rather than - using grid interpolation. - - Parameters - ---------- - peaks : list - List of (alpha, beta, gamma, score, sigma) tuples from grid search. - xi_nlm : np.ndarray - Wigner coefficients from compute_ball_cross_correlation_coefficients(), - shape (2N-1, L, 2L-1). - L : int - Angular bandlimit. - rf_mean, rf_std : float, optional - Mean and std of rotation function for sigma calculation. - If None, sigma values are preserved from input. - verbose : bool - Print progress. - - Returns - ------- - refined_peaks : list - List of (alpha, beta, gamma, score, sigma) tuples with refined positions. - """ - from scipy.optimize import minimize - - if not peaks: - return peaks - - # Pre-compute index mappings (done once) - if verbose: - print(" Building Wigner index mappings...") - xi_indices, D_indices, l_values = build_wigner_index_mapping(L) - wigner = spherical.Wigner(L - 1) - - # Grid spacing for search bounds - d_angle = 2 * np.pi / (2 * L - 1) - d_beta = np.pi / L - - refined_peaks = [] - - for i, (alpha, beta, gamma, score, sigma) in enumerate(peaks): - def neg_rf(angles): - return -evaluate_rotation_function_at_angles( - xi_nlm, angles[0], angles[1], angles[2], - wigner, xi_indices, D_indices, l_values - ) - - # Local search within ~2 grid spacings - bounds = [ - (alpha - 2*d_angle, alpha + 2*d_angle), - (max(0.01, beta - 2*d_beta), min(np.pi - 0.01, beta + 2*d_beta)), - (gamma - 2*d_angle, gamma + 2*d_angle) - ] - - result = minimize(neg_rf, [alpha, beta, gamma], - method='L-BFGS-B', bounds=bounds, - options={'ftol': 1e-10, 'gtol': 1e-10}) - - alpha_ref = result.x[0] % (2 * np.pi) - beta_ref = np.clip(result.x[1], 0, np.pi) - gamma_ref = result.x[2] % (2 * np.pi) - score_ref = -result.fun - - # Compute sigma if stats provided - if rf_mean is not None and rf_std is not None and rf_std > 1e-10: - sigma_ref = (score_ref - rf_mean) / rf_std - else: - sigma_ref = sigma # Keep original - - refined_peaks.append((alpha_ref, beta_ref, gamma_ref, score_ref, sigma_ref)) - - if verbose and (i + 1) % 100 == 0: - print(f" Refined {i + 1}/{len(peaks)} peaks") - - if verbose: - print(f" Refined {len(peaks)} peaks") - - return refined_peaks - - -def find_rotation_peaks( - rotation_function: np.ndarray, - alphas: np.ndarray, - betas: np.ndarray, - gammas: np.ndarray, - n_peaks: int = 100, - sigma_cutoff: float = 2.0, - cluster_radius_deg: float = 5.0, -) -> list: - """ - Extract and cluster peaks from rotation function. - - Parameters - ---------- - rotation_function : np.ndarray - Rotation function, shape (n_gamma, n_beta, n_alpha). - alphas, betas, gammas : np.ndarray - Angle grids. - n_peaks : int - Maximum number of peaks. - sigma_cutoff : float - Minimum sigma above mean. - cluster_radius_deg : float - Clustering radius in degrees. - - Returns - ------- - peaks : list - List of (alpha, beta, gamma, score, sigma) tuples. - """ - rf_mean = rotation_function.mean() - rf_std = rotation_function.std() - - if rf_std < 1e-10: - return [] - - threshold = rf_mean + sigma_cutoff * rf_std - - # Get sorted indices (descending) - flat_rf = rotation_function.flatten() - sorted_idx = np.argsort(flat_rf)[::-1] - - n_gamma, n_beta, n_alpha = rotation_function.shape - cluster_rad = np.radians(cluster_radius_deg) - - peaks = [] - used_angles = [] - - for flat_i in sorted_idx: - if len(peaks) >= n_peaks: - break - - g_idx, b_idx, a_idx = np.unravel_index(flat_i, rotation_function.shape) - score = rotation_function[g_idx, b_idx, a_idx] - - if score < threshold: - break - - alpha = alphas[a_idx] - beta = betas[b_idx] - gamma = gammas[g_idx] - - # Check if too close to existing peak - is_new = True - for prev_alpha, prev_beta, prev_gamma in used_angles: - da = min(abs(alpha - prev_alpha), 2*np.pi - abs(alpha - prev_alpha)) - db = abs(beta - prev_beta) - dg = min(abs(gamma - prev_gamma), 2*np.pi - abs(gamma - prev_gamma)) - if np.sqrt(da**2 + db**2 + dg**2) < cluster_rad: - is_new = False - break - - if is_new: - sigma = (score - rf_mean) / rf_std - peaks.append((alpha, beta, gamma, score, sigma)) - used_angles.append((alpha, beta, gamma)) - - return peaks - - -def rotation_matrix_from_euler_zyz( - alpha: float, - beta: float, - gamma: float, -) -> np.ndarray: - """ - Create rotation matrix from ZYZ Euler angles. - - R = Rz(alpha) @ Ry(beta) @ Rz(gamma) - """ - ca, sa = np.cos(alpha), np.sin(alpha) - cb, sb = np.cos(beta), np.sin(beta) - cg, sg = np.cos(gamma), np.sin(gamma) - - R = np.array([ - [ca*cb*cg - sa*sg, -ca*cb*sg - sa*cg, ca*sb], - [sa*cb*cg + ca*sg, -sa*cb*sg + ca*cg, sa*sb], - [-sb*cg, sb*sg, cb] - ]) - - return R - - -def check_rotation_recovery( - peaks: list, - true_alpha: float, - true_beta: float, - true_gamma: float, - symmetry_matrices: Optional[np.ndarray] = None, - tolerance_deg: float = 10.0, -) -> Tuple[bool, int, float]: - """ - Check if true rotation was recovered among top peaks. - - Parameters - ---------- - peaks : list - List of (alpha, beta, gamma, score, sigma) peaks. - true_alpha, true_beta, true_gamma : float - True rotation angles in radians. - symmetry_matrices : np.ndarray, optional - Symmetry operations for equivalence checking. - tolerance_deg : float - Angular tolerance in degrees. - - Returns - ------- - found : bool - Whether the rotation was found. - rank : int - Rank of matching peak (0-indexed), or -1. - min_error : float - Minimum angular error in degrees. - """ - R_true = rotation_matrix_from_euler_zyz(true_alpha, true_beta, true_gamma) - tolerance_rad = np.radians(tolerance_deg) - - min_error = float('inf') - best_rank = -1 - - for rank, (alpha, beta, gamma, score, sigma) in enumerate(peaks): - R_peak = rotation_matrix_from_euler_zyz(alpha, beta, gamma) - - error = _rotation_matrix_error(R_peak, R_true) - - if error < min_error: - min_error = error - best_rank = rank - - if symmetry_matrices is not None: - for S in symmetry_matrices: - S_rot = S[:3, :3] if S.shape[0] > 3 else S - for R_combined in [S_rot @ R_true, R_true @ S_rot]: - err = _rotation_matrix_error(R_peak, R_combined) - if err < min_error: - min_error = err - best_rank = rank - - found = min_error < tolerance_rad - return found, best_rank, np.degrees(min_error) - - -def _rotation_matrix_error(R1: np.ndarray, R2: np.ndarray) -> float: - """Compute angular error between rotation matrices.""" - R_diff = R1 @ R2.T - trace = np.clip(np.trace(R_diff), -1, 3) - angle = np.arccos((trace - 1) / 2) - return abs(angle) - - -@dataclass -class RotationCluster: - """ - Container for a cluster of rotation peaks. - - Attributes - ---------- - alpha : float - Alpha angle of best peak (radians). - beta : float - Beta angle of best peak (radians). - gamma : float - Gamma angle of best peak (radians). - best_score : float - Score of the best peak in cluster. - best_sigma : float - Sigma of the best peak. - size : int - Number of peaks in cluster. - sum_score : float - Sum of all scores in cluster. - mean_score : float - Mean score of peaks in cluster. - sum_sigma : float - Sum of all sigma values in cluster. - mean_sigma : float - Mean sigma of peaks in cluster. - score_std : float - Standard deviation of scores in cluster. - avg_alpha : float - Score-weighted average alpha (radians). - avg_beta : float - Score-weighted average beta (radians). - avg_gamma : float - Score-weighted average gamma (radians). - angular_spread : float - Average angular distance of peaks from cluster center (degrees). - """ - alpha: float - beta: float - gamma: float - best_score: float - best_sigma: float - size: int - sum_score: float - mean_score: float - sum_sigma: float - mean_sigma: float - score_std: float - avg_alpha: float - avg_beta: float - avg_gamma: float - angular_spread: float - - def to_tuple(self) -> tuple: - """Return basic tuple (alpha, beta, gamma, score, sigma, size).""" - return (self.alpha, self.beta, self.gamma, self.best_score, self.best_sigma, self.size) - - def __repr__(self) -> str: - return ( - f"RotationCluster(α={np.degrees(self.alpha):.1f}°, " - f"β={np.degrees(self.beta):.1f}°, γ={np.degrees(self.gamma):.1f}°, " - f"score={self.best_score:.4f}, size={self.size}, " - f"sum_score={self.sum_score:.4f}, spread={self.angular_spread:.2f}°)" - ) - - -def _average_quaternions_weighted(quaternions: np.ndarray, weights: np.ndarray) -> np.ndarray: - """ - Compute weighted average of quaternions. - - Uses the eigenvector method for quaternion averaging. - - Parameters - ---------- - quaternions : np.ndarray - Array of quaternions [w, x, y, z], shape (N, 4). - weights : np.ndarray - Weights for each quaternion, shape (N,). - - Returns - ------- - avg_q : np.ndarray - Averaged quaternion [w, x, y, z], shape (4,). - """ - # Normalize weights - weights = weights / weights.sum() - - # Build weighted outer product sum - M = np.zeros((4, 4)) - for q, w in zip(quaternions, weights): - M += w * np.outer(q, q) - - # Eigenvector with largest eigenvalue is the average - eigenvalues, eigenvectors = np.linalg.eigh(M) - avg_q = eigenvectors[:, -1] # Largest eigenvalue - - # Ensure w >= 0 - if avg_q[0] < 0: - avg_q = -avg_q - - return avg_q - - -def _quaternion_to_euler_zyz(q: np.ndarray) -> Tuple[float, float, float]: - """Convert quaternion [w, x, y, z] to ZYZ Euler angles.""" - w, x, y, z = q - - # Convert to rotation matrix - R = np.array([ - [1 - 2*(y*y + z*z), 2*(x*y - w*z), 2*(x*z + w*y)], - [2*(x*y + w*z), 1 - 2*(x*x + z*z), 2*(y*z - w*x)], - [2*(x*z - w*y), 2*(y*z + w*x), 1 - 2*(x*x + y*y)] - ]) - - return rotation_matrix_to_euler_zyz(R) - - -def cluster_rotation_peaks( - peaks: list, - cluster_radius_deg: float = 5.0, - symmetry_matrices: Optional[np.ndarray] = None, - return_details: bool = False, -) -> list: - """ - Cluster rotation peaks based on rotation matrix distance. - - Uses greedy clustering: iterates through peaks (sorted by score), and - assigns each peak to an existing cluster if within the angular threshold, - otherwise creates a new cluster. Returns the best peak from each cluster. - - Parameters - ---------- - peaks : list - List of (alpha, beta, gamma, score, sigma) tuples. - Angles are in radians. - cluster_radius_deg : float - Maximum angular distance (in degrees) for peaks to be in the same cluster. - symmetry_matrices : np.ndarray, optional - If provided, considers symmetry-equivalent rotations when clustering. - Shape (N, 3, 3) or (N, 4, 4). - return_details : bool - If True, return list of RotationCluster objects with full statistics. - If False, return simple tuples (alpha, beta, gamma, score, sigma, size). - - Returns - ------- - clustered_peaks : list - If return_details=False: List of (alpha, beta, gamma, score, sigma, cluster_size) tuples. - If return_details=True: List of RotationCluster objects with aggregate statistics. - Sorted by best score (descending). - """ - if not peaks: - return [] - - cluster_radius_rad = np.radians(cluster_radius_deg) - - # Sort peaks by score (descending) - sorted_peaks = sorted(peaks, key=lambda p: p[3], reverse=True) - - # Each cluster stores: (best_peak, rotation_matrix, list_of_all_peaks) - clusters = [] - - for peak in sorted_peaks: - alpha, beta, gamma = peak[0], peak[1], peak[2] - R_peak = rotation_matrix_from_euler_zyz(alpha, beta, gamma) - - # Find closest cluster - best_cluster_idx = -1 - best_distance = float('inf') - - for i, (_, R_cluster, _) in enumerate(clusters): - # Direct distance - dist = _rotation_matrix_error(R_peak, R_cluster) - - # Check symmetry-equivalent rotations if provided - if symmetry_matrices is not None: - for S in symmetry_matrices: - S_rot = S[:3, :3] if S.shape[0] > 3 else S - # Check both S @ R_peak and R_peak @ S - for R_equiv in [S_rot @ R_peak, R_peak @ S_rot]: - d = _rotation_matrix_error(R_equiv, R_cluster) - dist = min(dist, d) - - if dist < best_distance: - best_distance = dist - best_cluster_idx = i - - if best_distance < cluster_radius_rad and best_cluster_idx >= 0: - # Add to existing cluster - best_peak, R_cluster, peak_list = clusters[best_cluster_idx] - peak_list.append(peak) - clusters[best_cluster_idx] = (best_peak, R_cluster, peak_list) - else: - # Create new cluster - clusters.append((peak, R_peak, [peak])) - - # Compute cluster statistics - clustered_peaks = [] - for best_peak, R_center, peak_list in clusters: - alpha, beta, gamma, score, sigma = best_peak[:5] - size = len(peak_list) - - # Basic statistics - scores = np.array([p[3] for p in peak_list]) - sigmas = np.array([p[4] for p in peak_list]) - - sum_score = scores.sum() - mean_score = scores.mean() - sum_sigma = sigmas.sum() - mean_sigma = sigmas.mean() - score_std = scores.std() if size > 1 else 0.0 - - # Compute weighted average rotation using quaternions - if size > 1: - quaternions = [] - weights = [] - for p in peak_list: - R_p = rotation_matrix_from_euler_zyz(p[0], p[1], p[2]) - q = rotation_matrix_to_quaternion(R_p) - quaternions.append(q) - weights.append(max(p[3], 0.0)) # Use score as weight (non-negative) - - quaternions = np.array(quaternions) - weights = np.array(weights) - - if weights.sum() > 0: - avg_q = _average_quaternions_weighted(quaternions, weights) - avg_alpha, avg_beta, avg_gamma = _quaternion_to_euler_zyz(avg_q) - else: - avg_alpha, avg_beta, avg_gamma = alpha, beta, gamma - else: - avg_alpha, avg_beta, avg_gamma = alpha, beta, gamma - - # Compute angular spread (average distance from center) - if size > 1: - distances = [] - for p in peak_list: - R_p = rotation_matrix_from_euler_zyz(p[0], p[1], p[2]) - dist = _rotation_matrix_error(R_p, R_center) - distances.append(dist) - angular_spread = np.degrees(np.mean(distances)) - else: - angular_spread = 0.0 - - if return_details: - cluster = RotationCluster( - alpha=alpha, - beta=beta, - gamma=gamma, - best_score=score, - best_sigma=sigma, - size=size, - sum_score=sum_score, - mean_score=mean_score, - sum_sigma=sum_sigma, - mean_sigma=mean_sigma, - score_std=score_std, - avg_alpha=avg_alpha, - avg_beta=avg_beta, - avg_gamma=avg_gamma, - angular_spread=angular_spread, - ) - clustered_peaks.append(cluster) - else: - clustered_peaks.append((alpha, beta, gamma, score, sigma, size)) - - # Sort by best score (descending) - if return_details: - clustered_peaks.sort(key=lambda c: c.best_score, reverse=True) - else: - clustered_peaks.sort(key=lambda p: p[3], reverse=True) - - return clustered_peaks - - -def cluster_rotation_peaks_torch( - peaks: list, - cluster_radius_deg: float = 5.0, - symmetry_matrices: Optional[torch.Tensor] = None, - return_details: bool = False, -) -> list: - """ - Torch wrapper for cluster_rotation_peaks. - - Parameters - ---------- - peaks : list - List of (alpha, beta, gamma, score, sigma) tuples. - cluster_radius_deg : float - Maximum angular distance for clustering. - symmetry_matrices : torch.Tensor, optional - Symmetry matrices, shape (N, 3, 3) or (N, 4, 4). - return_details : bool - If True, return RotationCluster objects with full statistics. - - Returns - ------- - clustered_peaks : list - If return_details=False: List of (alpha, beta, gamma, score, sigma, cluster_size) tuples. - If return_details=True: List of RotationCluster objects. - """ - sym_np = None - if symmetry_matrices is not None: - sym_np = symmetry_matrices.detach().cpu().numpy() - return cluster_rotation_peaks(peaks, cluster_radius_deg, sym_np, return_details) - - -# Torch convenience wrappers -def ball_rotation_search_torch( - E_obs: torch.Tensor, - s_obs: torch.Tensor, - E_calc: torch.Tensor, - s_calc: torch.Tensor, - **kwargs, -) -> Tuple[np.ndarray, tuple, list]: - """Torch wrapper for ball_rotation_search.""" - return ball_rotation_search( - E_obs.detach().cpu().numpy(), - s_obs.detach().cpu().numpy(), - E_calc.detach().cpu().numpy(), - s_calc.detach().cpu().numpy(), - **kwargs, - ) - - -# ============================================================================= -# Symmetry Reduction Functions -# ============================================================================= - -def rotation_matrix_to_quaternion(R: np.ndarray) -> np.ndarray: - """ - Convert a 3x3 rotation matrix to a unit quaternion [w, x, y, z]. - - Uses Shepperd's method for numerical stability. - - Parameters - ---------- - R : np.ndarray - Rotation matrix, shape (3, 3). - - Returns - ------- - q : np.ndarray - Unit quaternion [w, x, y, z], shape (4,). - """ - trace = np.trace(R) - - if trace > 0: - s = 0.5 / np.sqrt(trace + 1.0) - w = 0.25 / s - x = (R[2, 1] - R[1, 2]) * s - y = (R[0, 2] - R[2, 0]) * s - z = (R[1, 0] - R[0, 1]) * s - elif R[0, 0] > R[1, 1] and R[0, 0] > R[2, 2]: - s = 2.0 * np.sqrt(1.0 + R[0, 0] - R[1, 1] - R[2, 2]) - w = (R[2, 1] - R[1, 2]) / s - x = 0.25 * s - y = (R[0, 1] + R[1, 0]) / s - z = (R[0, 2] + R[2, 0]) / s - elif R[1, 1] > R[2, 2]: - s = 2.0 * np.sqrt(1.0 + R[1, 1] - R[0, 0] - R[2, 2]) - w = (R[0, 2] - R[2, 0]) / s - x = (R[0, 1] + R[1, 0]) / s - y = 0.25 * s - z = (R[1, 2] + R[2, 1]) / s - else: - s = 2.0 * np.sqrt(1.0 + R[2, 2] - R[0, 0] - R[1, 1]) - w = (R[1, 0] - R[0, 1]) / s - x = (R[0, 2] + R[2, 0]) / s - y = (R[1, 2] + R[2, 1]) / s - z = 0.25 * s - - q = np.array([w, x, y, z]) - # Normalize - q = q / np.linalg.norm(q) - # Ensure w >= 0 for canonical form (q and -q represent same rotation) - if q[0] < 0: - q = -q - return q - - -def quaternion_to_euler_zyz(q: np.ndarray) -> Tuple[float, float, float]: - """ - Convert unit quaternion [w, x, y, z] to ZYZ Euler angles. - - Parameters - ---------- - q : np.ndarray - Unit quaternion [w, x, y, z], shape (4,). - - Returns - ------- - alpha, beta, gamma : float - ZYZ Euler angles in radians. - """ - w, x, y, z = q - - # First convert to rotation matrix - R = np.array([ - [1 - 2*(y*y + z*z), 2*(x*y - w*z), 2*(x*z + w*y)], - [2*(x*y + w*z), 1 - 2*(x*x + z*z), 2*(y*z - w*x)], - [2*(x*z - w*y), 2*(y*z + w*x), 1 - 2*(x*x + y*y)] - ]) - - return rotation_matrix_to_euler_zyz(R) - - -def rotation_matrix_to_euler_zyz(R: np.ndarray) -> Tuple[float, float, float]: - """ - Extract ZYZ Euler angles from rotation matrix. - - R = Rz(alpha) @ Ry(beta) @ Rz(gamma) - - Parameters - ---------- - R : np.ndarray - Rotation matrix, shape (3, 3). - - Returns - ------- - alpha, beta, gamma : float - ZYZ Euler angles in radians, with: - - alpha in [0, 2π) - - beta in [0, π] - - gamma in [0, 2π) - """ - # beta from R[2,2] = cos(beta) - cos_beta = np.clip(R[2, 2], -1.0, 1.0) - beta = np.arccos(cos_beta) - - sin_beta = np.sin(beta) - - if np.abs(sin_beta) > 1e-10: - # General case: sin(beta) != 0 - # alpha from R[0,2] = cos(alpha)*sin(beta), R[1,2] = sin(alpha)*sin(beta) - alpha = np.arctan2(R[1, 2], R[0, 2]) - # gamma from R[2,0] = -sin(beta)*cos(gamma), R[2,1] = sin(beta)*sin(gamma) - gamma = np.arctan2(R[2, 1], -R[2, 0]) - else: - # Gimbal lock: beta ≈ 0 or beta ≈ π - # Only (alpha + gamma) or (alpha - gamma) is determined - gamma = 0.0 - if cos_beta > 0: # beta ≈ 0 - alpha = np.arctan2(R[1, 0], R[0, 0]) - else: # beta ≈ π - alpha = np.arctan2(-R[1, 0], -R[0, 0]) - - # Normalize to [0, 2π) - alpha = alpha % (2 * np.pi) - gamma = gamma % (2 * np.pi) - - return alpha, beta, gamma - - -def reduce_rotation_by_symmetry( - alpha: float, - beta: float, - gamma: float, - symmetry_matrices: np.ndarray, -) -> Tuple[float, float, float, int]: - """ - Map rotation angles to a canonical representative using crystal symmetry. - - Given a rotation R and symmetry operations {S_i}, finds the canonical - representative among all symmetry-equivalent rotations {S_i @ R}. - The canonical form is chosen as the one with the smallest Euler angle - norm (closest to identity rotation), with all angles positive. - - Parameters - ---------- - alpha, beta, gamma : float - ZYZ Euler angles in radians. - symmetry_matrices : np.ndarray - Crystallographic symmetry rotation matrices, shape (N, 3, 3). - These should be the rotation parts of the space group operations. - - Returns - ------- - alpha_red, beta_red, gamma_red : float - Reduced ZYZ Euler angles in radians (all positive, smallest norm). - sym_index : int - Index of the symmetry operation that gave the canonical form. - """ - R = rotation_matrix_from_euler_zyz(alpha, beta, gamma) - - best_norm = float('inf') - best_angles = (alpha % (2 * np.pi), beta, gamma % (2 * np.pi)) - best_sym_idx = 0 - - for i, S in enumerate(symmetry_matrices): - # Extract 3x3 rotation part if needed - S_rot = S[:3, :3] if S.shape[0] > 3 else S - - # Apply symmetry: S @ R - R_equiv = S_rot @ R - - # Extract Euler angles (already positive from rotation_matrix_to_euler_zyz) - a, b, g = rotation_matrix_to_euler_zyz(R_equiv) - - # Compute norm of the angle vector - norm = np.sqrt(a**2 + b**2 + g**2) - - if norm < best_norm: - best_norm = norm - best_angles = (a, b, g) - best_sym_idx = i - - return (*best_angles, best_sym_idx) - - -def reduce_peaks_by_symmetry( - peaks: list, - symmetry_matrices: np.ndarray, -) -> list: - """ - Reduce a list of rotation peaks to canonical symmetry representatives. - - Parameters - ---------- - peaks : list - List of (alpha, beta, gamma, score, sigma) tuples. - Angles are in radians. - symmetry_matrices : np.ndarray - Crystallographic symmetry rotation matrices, shape (N, 3, 3). - - Returns - ------- - reduced_peaks : list - List of (alpha, beta, gamma, score, sigma, sym_index) tuples. - Angles are the canonical representatives in radians. - """ - reduced_peaks = [] - - for alpha, beta, gamma, score, sigma in peaks: - alpha_red, beta_red, gamma_red, sym_idx = reduce_rotation_by_symmetry( - alpha, beta, gamma, symmetry_matrices - ) - reduced_peaks.append((alpha_red, beta_red, gamma_red, score, sigma, sym_idx)) - - return reduced_peaks - - -def reduce_peaks_by_symmetry_torch( - peaks: list, - symmetry_matrices: torch.Tensor, -) -> list: - """ - Torch wrapper for reduce_peaks_by_symmetry. - - Parameters - ---------- - peaks : list - List of (alpha, beta, gamma, score, sigma) tuples. - symmetry_matrices : torch.Tensor - Symmetry matrices, shape (N, 3, 3) or (N, 4, 4). - - Returns - ------- - reduced_peaks : list - List of (alpha, beta, gamma, score, sigma, sym_index) tuples. - """ - sym_np = symmetry_matrices.detach().cpu().numpy() - return reduce_peaks_by_symmetry(peaks, sym_np) diff --git a/torchref/alignment/jax_subpixel_peaks.py b/torchref/alignment/jax_subpixel_peaks.py deleted file mode 100644 index 5fbf2135..00000000 --- a/torchref/alignment/jax_subpixel_peaks.py +++ /dev/null @@ -1,67 +0,0 @@ -import jax -import jax.numpy as jnp - - -@jax.jit -def refine_peaks_subvoxel( - grid: jax.Array, - peak_indices: jax.Array, -) -> tuple[jax.Array, jax.Array]: - """ - Refine discrete peaks to sub-voxel accuracy using separable quadratic fitting. - - Args: - grid: (n0, n1, n2) array of values - peak_indices: (n_peaks, 3) integer indices of discrete maxima - - Returns: - refined_coords: (n_peaks, 3) fractional coordinates - refined_values: (n_peaks,) interpolated peak heights - """ - shape = jnp.array(grid.shape) - - def refine_single_peak(idx): - i0, i1, i2 = idx[0], idx[1], idx[2] - - def get_val(d0, d1, d2): - # Periodic boundary conditions - ii0 = (i0 + d0) % shape[0] - ii1 = (i1 + d1) % shape[1] - ii2 = (i2 + d2) % shape[2] - return grid[ii0, ii1, ii2] - - f0 = get_val(0, 0, 0) - - # Axis 0 - fm = get_val(-1, 0, 0) - fp = get_val(1, 0, 0) - c0 = (fp - 2*f0 + fm) / 2 - b0 = (fp - fm) / 2 - delta0 = jnp.where(jnp.abs(c0) > 1e-12, jnp.clip(-b0 / (2*c0), -1.0, 1.0), 0.0) - - # Axis 1 - fm = get_val(0, -1, 0) - fp = get_val(0, 1, 0) - c1 = (fp - 2*f0 + fm) / 2 - b1 = (fp - fm) / 2 - delta1 = jnp.where(jnp.abs(c1) > 1e-12, jnp.clip(-b1 / (2*c1), -1.0, 1.0), 0.0) - - # Axis 2 - fm = get_val(0, 0, -1) - fp = get_val(0, 0, 1) - c2 = (fp - 2*f0 + fm) / 2 - b2 = (fp - fm) / 2 - delta2 = jnp.where(jnp.abs(c2) > 1e-12, jnp.clip(-b2 / (2*c2), -1.0, 1.0), 0.0) - - # Refined coordinates - refined = jnp.array([i0 + delta0, i1 + delta1, i2 + delta2]) - - # Refined value: f(δ) ≈ f0 - b²/(4c) at peak - peak_val = f0 - peak_val -= jnp.where(jnp.abs(c0) > 1e-12, b0**2 / (4*c0), 0.0) - peak_val -= jnp.where(jnp.abs(c1) > 1e-12, b1**2 / (4*c1), 0.0) - peak_val -= jnp.where(jnp.abs(c2) > 1e-12, b2**2 / (4*c2), 0.0) - - return refined, peak_val - - return jax.vmap(refine_single_peak)(peak_indices) \ No newline at end of file diff --git a/torchref/alignment/lattman_love.py b/torchref/alignment/lattman_love.py new file mode 100644 index 00000000..371e6bc3 --- /dev/null +++ b/torchref/alignment/lattman_love.py @@ -0,0 +1,219 @@ +""" +Lattman-Love (1970) structure-factor interpolation for the alignment module. + +Compute F_calc once for the search model in a large cubic P1 box (densely sampled +in reciprocal space), then interpolate the dense grid at rotated reciprocal-space +positions of the *real* crystal cell to obtain F_calc for any candidate rotation. + +This is the standard Phaser MR setup (Phaser paper §2.2.2): F_calc is generated +"by structure-factor interpolation (Lattman & Love, 1970) from a model in a large +P1 unit cell". It removes the sphere-sampling bias that arises if one instead +rotates atom coordinates and recomputes F_calc on the (non-uniform) real-cell +HKL grid — which is what the bare ball-search hits on real, non-cubic data. + +Convention (matches `torchref/base/reciprocal/interpolation.py::interpolate_for_rotation`): + For a model whose atom coordinates have been rotated by R (column-vector + convention: xyz_new = R · xyz_old), the structure factor at real-cell HKL h + equals the un-rotated model's structure factor at the rotated reciprocal- + space point R^T · s_real, where s_real = h · rec_basis(real_cell). + +Only amplitudes are needed for the rotation function and Sim MLRF rescoring. +Translation search (downstream) needs the phase too — this class returns complex +F so callers can use either. +""" + +from __future__ import annotations + +from typing import Optional + +import torch + +from torchref.base.fourier.fft import ifft +from torchref.base.reciprocal.interpolation import ( + interpolate_complex_from_grid, + interpolate_structure_factor_from_grid, +) +from torchref.model.sf_fft import SfFFT +from torchref.symmetry import SpaceGroup +from torchref.symmetry.cell import Cell + + +class LattmanLoveInterpolator: + """ + Compute F_calc on a dense P1 reciprocal grid once; interpolate at arbitrary + rotated reciprocal positions per query. + + Parameters + ---------- + model : ModelFT + Search model. Its current atom coordinates are used and the FT is built + for those positions (un-rotated; rotations are applied later in + `evaluate(R, ...)`). The model's own cell/spacegroup are NOT used. + padding_factor : float, default 2.0 + Cubic P1 box side = padding_factor * molecule_bounding_box_diameter. + Phaser uses ~2.0. Larger → finer reciprocal grid spacing, more memory. + min_cell_size_A : float, optional + Lower bound on the cubic side (Å). Useful for very small molecules where + 2·diameter would give a tiny FFT grid. + max_res_A : float, optional + Resolution limit (Å). The dense grid will resolve features down to this. + Default: 2.0 Å (suitable for proteins up to that resolution). + radius_angstrom : float, default 3.0 + Atomic-density support radius for SfFFT. + device : torch.device, optional + Target device. Defaults to the model's device. + """ + + def __init__( + self, + model, + padding_factor: float = 2.0, + min_cell_size_A: Optional[float] = None, + max_res_A: float = 2.0, + radius_angstrom: float = 3.0, + device: Optional[torch.device] = None, + verbose: int = 0, + ): + if device is None: + device = model.xyz().device + + # Bounding-box diameter of the un-rotated atomic coordinates. + xyz = model.xyz().detach().to(device) + bbox_max = xyz.max(dim=0).values + bbox_min = xyz.min(dim=0).values + diameter_A = (bbox_max - bbox_min).norm().item() + cubic_side = padding_factor * diameter_A + if min_cell_size_A is not None: + cubic_side = max(cubic_side, float(min_cell_size_A)) + + # Cubic P1 cell. Keep dtype float32 to match Cell defaults / SfFFT grid math. + self.cubic_cell = Cell( + [cubic_side, cubic_side, cubic_side, 90.0, 90.0, 90.0], + dtype=torch.float32, device=device, + ) + + # Shift atoms so the molecule centroid sits at the centre of the cubic box. + centroid = xyz.mean(dim=0) + target = torch.tensor( + [cubic_side / 2, cubic_side / 2, cubic_side / 2], + dtype=xyz.dtype, device=device, + ) + self.shift_vec = target - centroid # (3,) — applied to atom coords + + # Extract atomic parameters (use the model's own helper). + xyz_iso, adp_iso, occ_iso, A_iso, B_iso = model.get_iso() + xyz_iso = (xyz_iso.detach().to(device) + self.shift_vec).to(xyz.dtype) + adp_iso = adp_iso.detach().to(device) + occ_iso = occ_iso.detach().to(device) + A_iso = A_iso.detach().to(device) + B_iso = B_iso.detach().to(device) + + # Handle the (uncommon) anisotropic atoms by leaving them aside for the + # search model — Phaser-style MR also approximates with isotropic ADP. + # Caller is free to call evaluate after adding aniso atoms in a subclass. + + # Build the dense F_calc on a P1 cubic grid. + sf = SfFFT( + cell=self.cubic_cell, + spacegroup=SpaceGroup("P 1"), + max_res=max_res_A, + radius_angstrom=radius_angstrom, + dtype_float=torch.float32, + device=device, + verbose=verbose, + ) + sf.setup_grid() + density_map = sf.build_density_map( + xyz_iso=xyz_iso, + adp_iso=adp_iso, + occ_iso=occ_iso, + A_iso=A_iso, + B_iso=B_iso, + apply_symmetry=False, # already P1 + ) + # IFFT to reciprocal space; gives a complex (Nx, Ny, Nz) tensor with + # crystallographic normalization. Layout: DC at index (0, 0, 0); negative + # HKL wraps to high indices. This matches the convention expected by + # `interpolate_structure_factor_from_grid`. + self.reciprocal_grid = ifft(density_map, self.cubic_cell.volume.item()) + self.cubic_cell_volume = float(self.cubic_cell.volume.item()) + self.device = device + self.cubic_side = cubic_side + self.gridsize = tuple(int(x) for x in self.reciprocal_grid.shape) + self.max_res_A = max_res_A + + if verbose: + print(f"LattmanLove: molecule diameter {diameter_A:.1f} Å, " + f"cubic side {cubic_side:.1f} Å, grid {self.gridsize}, " + f"max_res {max_res_A:.2f} Å") + + @staticmethod + def _real_hkl_to_cubic_hkl( + hkl_real: torch.Tensor, real_cell: Cell, cubic_cell: Cell, + ) -> torch.Tensor: + """Convert HKL of `real_cell` to (float) HKL of `cubic_cell`.""" + # s = h @ rec_basis (Å^-1, Cartesian) + rec_real = real_cell.reciprocal_basis_matrix.to(hkl_real.device) + s = hkl_real.to(rec_real.dtype) @ rec_real + rec_cubic = cubic_cell.reciprocal_basis_matrix.to(hkl_real.device) + # cubic_hkl = s @ rec_cubic^{-1} + return s @ torch.linalg.inv(rec_cubic.to(rec_real.dtype)) + + def evaluate( + self, + R: torch.Tensor, + hkl_real: torch.Tensor, + real_cell: Cell, + return_amplitude: bool = True, + ) -> torch.Tensor: + """ + Interpolate F_calc at the real-cell HKL set after rotating the model + by R (column-vector convention). + + Parameters + ---------- + R : torch.Tensor + Rotation matrix, shape (3, 3) or (B, 3, 3). + hkl_real : torch.Tensor + Miller indices in the real crystal cell, shape (N, 3). + real_cell : Cell + The real crystal cell (provides `reciprocal_basis_matrix`). + return_amplitude : bool, default True + If True, return |F_calc| (real, no phase ambiguity from trilinear + interpolation). If False, return complex F — only safe if the dense + grid is well-oversampled (small `max_res_A`). + + Returns + ------- + torch.Tensor + Interpolated structure factors, shape (N,) or (B, N). + """ + batched = R.dim() == 3 + if not batched: + R = R.unsqueeze(0) + R = R.to(self.device).to(torch.float32) + + # Real HKL → Cartesian s (Å^-1, real cell) + rec_real = real_cell.reciprocal_basis_matrix.to(self.device).to(torch.float32) + s_real = hkl_real.to(self.device).to(torch.float32) @ rec_real # (N, 3) + + # Rotated reciprocal point: F_rotated_model(s) = F_orig(R^T s) => + # use R^T · s to look up the un-rotated grid. + # For batched R, einsum: + s_rot = torch.einsum("bij,nj->bni", R.transpose(-1, -2), s_real) # (B, N, 3) + + # Cartesian s → cubic-cell float HKL + rec_cubic = self.cubic_cell.reciprocal_basis_matrix.to(self.device).to(torch.float32) + inv_rec_cubic = torch.linalg.inv(rec_cubic) + hkl_cubic = s_rot @ inv_rec_cubic # (B, N, 3) + + B, N, _ = hkl_cubic.shape + flat = hkl_cubic.reshape(B * N, 3) + if return_amplitude: + interp = interpolate_structure_factor_from_grid( + self.reciprocal_grid, flat, interpolate_amplitude=True, + ) + else: + interp = interpolate_complex_from_grid(self.reciprocal_grid, flat) + out = interp.reshape(B, N) + return out if batched else out.squeeze(0) diff --git a/torchref/alignment/ml_rotation.py b/torchref/alignment/ml_rotation.py new file mode 100644 index 00000000..3a522ca7 --- /dev/null +++ b/torchref/alignment/ml_rotation.py @@ -0,0 +1,434 @@ +""" +Maximum-Likelihood Rotation Function (Sim MLRF) rescoring of peaks from the +fast ball-search. + +Phaser paper §2.1.2: the fast rotation function is a "shortlist generator"; +discrimination of the correct orientation comes from rescoring the top peaks +with a slow, full ML target. This file implements that rescoring. + +Per-shell σA (= D) is fitted on-the-fly for each candidate rotation. We work in +E-value (normalized structure-factor amplitude) space, which gives the standard +Rice / Woolfson likelihood forms + + P(E_obs | σA · E_calc, 1 − σA²) (acentric, Rice) + P(E_obs | σA · E_calc, 1 − σA²) (centric, Woolfson) + +The "log-likelihood gain" relative to a Wilson reference (σA = 0) is the +discriminating score Phaser reports as LLG. +""" + +from __future__ import annotations + +import math +from typing import List, Optional + +import torch + +from .ball_search import ( + RotationPeak, + edmonds_euler_from_rotation_matrix, + rotation_matrix_from_edmonds_euler, + rotation_matrix_from_edmonds_euler_batch, +) +from .distributions import rice_log_likelihood, woolfson_log_likelihood +from .lattman_love import LattmanLoveInterpolator + + +# ============================================================================= +# Helpers +# ============================================================================= + + +def _equal_count_shell_idx(s_mag: torch.Tensor, n_shells: int) -> torch.Tensor: + """ + Partition `s_mag` into `n_shells` shells with (approximately) equal counts. + + Returns + ------- + shell_idx : torch.Tensor (int64), shape (N,) + Shell index in [0, n_shells). + """ + n = s_mag.numel() + order = torch.argsort(s_mag) + chunk = max(n // n_shells, 1) + positions = torch.arange(n, device=s_mag.device, dtype=torch.int64) + sorted_labels = (positions // chunk).clamp(max=n_shells - 1) + shell_idx = torch.empty(n, dtype=torch.int64, device=s_mag.device) + shell_idx[order] = sorted_labels + return shell_idx + + +def _normalize_to_e(F: torch.Tensor, shell_idx: torch.Tensor, + n_shells: int) -> torch.Tensor: + """E = F / sqrt( per shell). Vectorised across shells via scatter.""" + F2 = F ** 2 + sum_per_shell = torch.zeros(n_shells, dtype=F2.dtype, device=F.device) + sum_per_shell.scatter_add_(0, shell_idx, F2) + count_per_shell = torch.bincount(shell_idx, minlength=n_shells).to(F2.dtype) + mean_per_shell = (sum_per_shell / count_per_shell.clamp(min=1.0)).clamp(min=1e-30) + norm_per_refl = mean_per_shell.sqrt().index_select(0, shell_idx) + return F / norm_per_refl + + +def _shell_ll( + E_obs: torch.Tensor, + E_calc: torch.Tensor, + centric: torch.Tensor, + D: float, +) -> torch.Tensor: + """ + Per-reflection log-likelihood at a given σA = D for one shell, in E-value space. + + Acentric: Rice with F_mean = D · E_calc, variance = 1 − D². + Centric: Woolfson with F_mean = D · E_calc, variance = 1 − D². + """ + var = torch.full_like(E_obs, max(1.0 - D * D, 1e-4)) + F_mean = D * E_calc + ll = torch.where( + centric, + woolfson_log_likelihood(E_obs, F_mean, var), + rice_log_likelihood(E_obs, F_mean, var), + ) + return ll + + +def _optimize_D_in_shell( + E_obs: torch.Tensor, + E_calc: torch.Tensor, + centric: torch.Tensor, + n_grid: int = 21, + n_refine: int = 12, +) -> float: + """ + Find the σA = D ∈ [0, 0.99] that maximizes the sum log-likelihood for this + shell. Two-stage: coarse grid search, then golden-section refinement. + + Returns the optimal D. + """ + # Coarse grid + D_grid = torch.linspace(0.0, 0.99, n_grid, device=E_obs.device) + best_D = 0.0 + best_ll = -float("inf") + for D in D_grid.tolist(): + ll = _shell_ll(E_obs, E_calc, centric, D).sum().item() + if ll > best_ll: + best_ll = ll + best_D = D + # Golden-section refinement around best_D + span = 1.0 / (n_grid - 1) + lo = max(0.0, best_D - span) + hi = min(0.99, best_D + span) + phi = (math.sqrt(5.0) - 1) / 2.0 + x1 = hi - phi * (hi - lo) + x2 = lo + phi * (hi - lo) + f1 = _shell_ll(E_obs, E_calc, centric, x1).sum().item() + f2 = _shell_ll(E_obs, E_calc, centric, x2).sum().item() + for _ in range(n_refine): + if f1 > f2: + hi = x2 + x2 = x1 + f2 = f1 + x1 = hi - phi * (hi - lo) + f1 = _shell_ll(E_obs, E_calc, centric, x1).sum().item() + else: + lo = x1 + x1 = x2 + f1 = f2 + x2 = lo + phi * (hi - lo) + f2 = _shell_ll(E_obs, E_calc, centric, x2).sum().item() + return 0.5 * (lo + hi) + + +# ============================================================================= +# Public API +# ============================================================================= + + +def llg_for_rotation( + F_obs: torch.Tensor, + s_mag: torch.Tensor, + shell_idx: torch.Tensor, + n_shells: int, + E_obs: torch.Tensor, + centric: torch.Tensor, + F_calc: torch.Tensor, + shell_weights: Optional[torch.Tensor] = None, +) -> float: + """ + Total log-likelihood gain (Sim − Wilson) for a single candidate rotation. + Thin wrapper around `llg_for_rotation_batch` for a single (1,N) input. + """ + F_calc_batch = F_calc.unsqueeze(0) if F_calc.dim() == 1 else F_calc + return llg_for_rotation_batch( + F_obs=F_obs, shell_idx=shell_idx, n_shells=n_shells, + E_obs=E_obs, centric=centric, F_calc=F_calc_batch, + shell_weights=shell_weights, + )[0].item() + + +def llg_for_rotation_batch( + F_obs: torch.Tensor, + shell_idx: torch.Tensor, + n_shells: int, + E_obs: torch.Tensor, + centric: torch.Tensor, + F_calc: torch.Tensor, + n_D_grid: int = 41, + shell_weights: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """ + Vectorized log-likelihood gain across a batch of candidate rotations. + + Per-shell σA fit is performed on a coarse grid of `n_D_grid` D values; + the grid maximum is taken (no golden refinement — for shortlisting only). + + Parameters + ---------- + F_obs : torch.Tensor, shape (N,) + shell_idx : torch.Tensor (int64), shape (N,) + n_shells : int + E_obs : torch.Tensor, shape (N,) + F_obs normalized to unit variance per shell. + centric : torch.Tensor (bool), shape (N,) + F_calc : torch.Tensor, shape (B, N) + Per-rotation |F_calc| at the same HKL set. + n_D_grid : int, default 41 + Number of σA grid points in [0, 0.99]. + shell_weights : torch.Tensor, shape (n_shells,), optional + Per-shell weight applied to the per-shell LL gain before accumulation. + Used to implement Phaser-style empirical variance correction + (`w_p = 1/√Var(E_obs²-1)_p`). The weight multiplies *both* the Sim and + Wilson LL contributions uniformly per shell, so the LL gain + interpretation is preserved. + + Returns + ------- + llg : torch.Tensor, shape (B,) + Total log-likelihood gain (Sim − Wilson) per rotation candidate. + """ + B, N = F_calc.shape + device = F_calc.device + dtype = F_calc.dtype + + # --- Per-shell E normalisation of F_calc, fully vectorised --- + # Build (B, n_shells) shell-mean of F_calc² via scatter_add, then gather + # back per reflection. Replaces the n_shells-step Python loop that + # masked + meaned one shell at a time. + shell_idx_b = shell_idx.view(1, N).expand(B, N) + F_calc2 = F_calc * F_calc + sum_per_shell_b = torch.zeros((B, n_shells), dtype=dtype, device=device) + sum_per_shell_b.scatter_add_(1, shell_idx_b, F_calc2) + count_per_shell = torch.bincount(shell_idx, minlength=n_shells).to(dtype) + mean_per_shell_b = ( + sum_per_shell_b / count_per_shell.clamp(min=1.0).unsqueeze(0) + ).clamp(min=1e-30) + norm_per_refl_b = mean_per_shell_b.sqrt().gather(1, shell_idx_b) # (B, N) + E_calc = F_calc / norm_per_refl_b + + # --- Joint (D, B, N) likelihood evaluation --- + # Memory: D · B · N · 8 B. For default args (D=41, B≤100, N≈3 k) this is + # ~100 MB, comparable to what the per-shell loop already built per + # iteration. For dense-R we typically pass n_D_grid=11, so cost is small. + D_grid = torch.linspace(0.0, 0.99, n_D_grid, device=device, dtype=dtype) + F_mean = D_grid.view(-1, 1, 1) * E_calc.unsqueeze(0) # (D, B, N) + var_d = (1.0 - D_grid * D_grid).clamp(min=1e-4) + var_full = var_d.view(-1, 1, 1).expand(n_D_grid, B, N) + E_obs_full = E_obs.view(1, 1, -1).expand(n_D_grid, B, N) + + ll_acent = rice_log_likelihood(E_obs_full, F_mean, var_full) + ll_cent = woolfson_log_likelihood(E_obs_full, F_mean, var_full) + cent_full = centric.view(1, 1, -1) + ll = torch.where(cent_full, ll_cent, ll_acent) # (D, B, N) + + # --- Sum per shell across N, max over D, sum weighted across shells --- + shell_idx_dbn = shell_idx.view(1, 1, -1).expand(n_D_grid, B, N) + ll_per_shell = torch.zeros((n_D_grid, B, n_shells), dtype=dtype, device=device) + ll_per_shell.scatter_add_(2, shell_idx_dbn, ll) + ll_sim_per_shell, _ = ll_per_shell.max(dim=0) # (B, n_shells) + + # Wilson reference at D = 0 (data-only): F_mean = 0, var = 1. + var0 = torch.ones_like(E_obs) + F_mean0 = torch.zeros_like(E_obs) + ll_wil_acent = rice_log_likelihood(E_obs, F_mean0, var0) + ll_wil_cent = woolfson_log_likelihood(E_obs, F_mean0, var0) + ll_wil_per_refl = torch.where(centric, ll_wil_cent, ll_wil_acent) + ll_wil_per_shell = torch.zeros(n_shells, dtype=dtype, device=device) + ll_wil_per_shell.scatter_add_(0, shell_idx, ll_wil_per_refl) + + gain_per_shell = ll_sim_per_shell - ll_wil_per_shell.unsqueeze(0) # (B, n_shells) + if shell_weights is not None: + gain_per_shell = gain_per_shell * shell_weights.to(dtype).view(1, -1) + total_gain = gain_per_shell.sum(dim=-1) # (B,) + return total_gain + + +def sim_mlrf_rescore( + peaks: List[RotationPeak], + F_obs: torch.Tensor, + hkl_real: torch.Tensor, + s_mag: torch.Tensor, + centric: torch.Tensor, + interpolator: LattmanLoveInterpolator, + real_cell, + n_shells: int = 20, + n_refine: Optional[int] = None, + batch_size: int = 100, + verbose: int = 0, + shell_weights: Optional[torch.Tensor] = None, + auto_variance_weights: bool = True, + n_D_grid: int = 41, +) -> List[RotationPeak]: + """ + Rescore a list of peaks from `ball_rotation_search` by the per-shell-fitted + Sim Maximum-Likelihood Rotation Function (LLG). Returns a new list sorted by + descending LLG with `score = LLG` and `sigma = Z-score(LLG)`. + + Batches candidates of size `batch_size` for fast vectorized evaluation. + + Parameters + ---------- + peaks : list of RotationPeak + F_obs : torch.Tensor, shape (N,) + hkl_real : torch.Tensor (int), shape (N, 3) + s_mag : torch.Tensor, shape (N,) + centric : torch.Tensor (bool), shape (N,) + interpolator : LattmanLoveInterpolator + real_cell : Cell + n_shells : int, default 20 + n_refine : int, optional + Number of top peaks to rescore (default: all). + batch_size : int, default 100 + Number of candidates evaluated together in one LL.evaluate + LLG call. + """ + if not peaks: + return [] + + if n_refine is None: + n_refine = len(peaks) + head = peaks[: n_refine] + tail = peaks[n_refine:] + + shell_idx = _equal_count_shell_idx(s_mag, n_shells) + E_obs = _normalize_to_e(F_obs, shell_idx, n_shells) + + if shell_weights is None and auto_variance_weights: + from .sh import compute_patterson_shell_variance + patt_obs = (E_obs.to(torch.float64) ** 2) - 1.0 + var_p = compute_patterson_shell_variance( + patt_obs, shell_idx, P=n_shells, + ) + w = 1.0 / var_p.sqrt() + w = w * (n_shells / w.sum().clamp(min=1e-30)) + shell_weights = w.to(F_obs.dtype) + + # Build all rotation matrices up front. The peak's Euler triple represents + # "the rotation applied to the model coords" (synthetic-test convention of + # ball_rotation_search). For ML scoring, we need "the rotation to apply to + # the current model to align it to obs" — which is R^T. We transpose here. + # Vectorised over peaks: previously this list comprehension built M·9 + # small (3,3) tensors per dense-R pass. + alpha_t = torch.tensor([p.alpha for p in head], dtype=torch.float64) + beta_t = torch.tensor([p.beta for p in head], dtype=torch.float64) + gamma_t = torch.tensor([p.gamma for p in head], dtype=torch.float64) + R_all = rotation_matrix_from_edmonds_euler_batch( + alpha_t, beta_t, gamma_t, + ).transpose(-1, -2).to(torch.float32) # (M, 3, 3) + + llg_chunks: List[torch.Tensor] = [] + M = R_all.shape[0] + for start in range(0, M, batch_size): + stop = min(start + batch_size, M) + R_batch = R_all[start:stop] # (B, 3, 3) + # Batched LL interpolation: returns (B, N) + F_calc = interpolator.evaluate( + R_batch, hkl_real, real_cell, return_amplitude=True, + ) + F_calc = F_calc.to(F_obs.dtype) + llg_batch = llg_for_rotation_batch( + F_obs=F_obs, shell_idx=shell_idx, n_shells=n_shells, + E_obs=E_obs, centric=centric, F_calc=F_calc, + shell_weights=shell_weights, n_D_grid=n_D_grid, + ) + llg_chunks.append(llg_batch) + if verbose > 1: + print(f" ML rescore batch {start}-{stop}/{M}", flush=True) + + # Concatenate on-device, compute z-score on-device, then ONE bulk + # transfer at the end. The previous code did `.cpu().tolist()` per + # batch — fine on CPU but a per-batch GPU↔CPU stall on cuda. + llgs_t = torch.cat(llg_chunks) + mean_t = llgs_t.mean() + std_t = llgs_t.std().clamp(min=1e-30) + sigmas_t = (llgs_t - mean_t) / std_t + + llgs_list = llgs_t.tolist() + sigmas_list = sigmas_t.tolist() + rescored = [ + RotationPeak( + alpha=p.alpha, beta=p.beta, gamma=p.gamma, + score=llg, sigma=sigma, + ) + for p, llg, sigma in zip(head, llgs_list, sigmas_list) + ] + rescored.sort(key=lambda r: r.score, reverse=True) + return rescored + tail + + +# ============================================================================= +# Brute-force ML rotation search over a uniform SO(3) sample +# ============================================================================= + + +def _uniform_random_rotations(n: int, seed: int = 0, dtype=torch.float64) -> torch.Tensor: + """ + Generate `n` uniformly-distributed random rotation matrices via QR of + Gaussian matrices. Returns shape (n, 3, 3) with det = +1. + """ + g = torch.Generator().manual_seed(int(seed)) + A = torch.randn(n, 3, 3, generator=g, dtype=dtype) + Q, R = torch.linalg.qr(A) + # Make det = 1 (flip first column if needed) + diag_sign = torch.sign(torch.diagonal(R, dim1=-2, dim2=-1)) # (n, 3) + Q = Q * diag_sign.unsqueeze(-2) + det = torch.det(Q) + flip = det < 0 + Q[flip, :, 0] = -Q[flip, :, 0] + return Q + + +def brute_ml_rotation_search( + F_obs: torch.Tensor, + hkl_real: torch.Tensor, + s_mag: torch.Tensor, + centric: torch.Tensor, + interpolator: LattmanLoveInterpolator, + real_cell, + n_candidates: int = 5000, + n_shells: int = 15, + batch_size: int = 100, + seed: int = 0, + verbose: int = 0, +) -> List[RotationPeak]: + """ + Evaluate ML LLG on `n_candidates` uniformly-random SO(3) rotations. + Returns peaks sorted by descending LLG. Acts as a shortlist generator that + bypasses the fast ball-search (which has known sphere-sampling limitations + on real-cell HKL data). + + Cost: ~10ms per candidate at default batch_size on CPU; ~50s for 5000 cands. + + Returns + ------- + list of RotationPeak (sorted by LLG descending) + `score = LLG`, `sigma = Z-score across the candidate set`. + """ + R_all = _uniform_random_rotations(n_candidates, seed=seed) + peaks_in = [] + for k in range(n_candidates): + a, b, g = edmonds_euler_from_rotation_matrix(R_all[k]) + peaks_in.append(RotationPeak(a, b, g, score=0.0, sigma=0.0)) + return sim_mlrf_rescore( + peaks_in, F_obs, hkl_real, s_mag, centric, interpolator, real_cell, + n_shells=n_shells, n_refine=n_candidates, batch_size=batch_size, + verbose=verbose, + ) diff --git a/torchref/alignment/patterson_filter.py b/torchref/alignment/patterson_filter.py new file mode 100644 index 00000000..3703b197 --- /dev/null +++ b/torchref/alignment/patterson_filter.py @@ -0,0 +1,163 @@ +""" +Restrict the Patterson function to within a sphere of radius Ω (molecule +diameter) so the rotation function sees only intramolecular vectors. + +This is the operational equivalent of Phaser's χ_Ω weighting (LERF1, §2.1.3): +the sphere of integration excludes intermolecular Patterson peaks that come +from cross-vectors between symmetry-related copies in the crystal — those are +the dominant contamination of |F_obs|² on real data and are why the bare +Crowther-style rotation function (|E|² · |E|² overlap) fails on real F_obs. + +Implementation +-------------- +- Place |F|² values onto a regular 3-D reciprocal-space grid of a large cubic + cell (radius ≥ Ω, so the sphere fits inside). +- 3-D FFT → real-space Patterson in Cartesian coordinates of the cubic cell. +- Apply spherical mask: zero values at |r| > Ω. +- IFFT → sphere-restricted Patterson coefficients on the cubic grid. +- Extract values at the original real-cell HKL positions via trilinear + interpolation. +""" + +from __future__ import annotations + +import math +from typing import Tuple + +import torch + +from torchref.symmetry.cell import Cell +from torchref.base.reciprocal.interpolation import ( + interpolate_structure_factor_from_grid, +) + + +def _make_cubic_grid(F_sq_real: torch.Tensor, + hkl_real: torch.Tensor, + real_cell: Cell, + cubic_side: float, + gridsize: int) -> Tuple[torch.Tensor, Cell]: + """ + Place |F|² values onto a (gridsize, gridsize, gridsize) cubic-cell + reciprocal grid in FFT layout. Multiple real-HKL points falling into the + same cubic-grid cell are accumulated. + + Returns + ------- + grid : torch.Tensor (real), shape (N, N, N) + Cubic reciprocal grid with FFT layout (DC at index 0, negative HKL + wrapped to high indices). + cubic_cell : Cell + """ + device = F_sq_real.device + dtype = F_sq_real.dtype + N = int(gridsize) + cubic_cell = Cell( + [cubic_side, cubic_side, cubic_side, 90.0, 90.0, 90.0], + dtype=torch.float32, device=device, + ) + # h_real → Cartesian s + rec_real = real_cell.reciprocal_basis_matrix.to(device).to(dtype) + s = hkl_real.to(dtype) @ rec_real # (M, 3) + # Cartesian s → cubic-cell HKL (orthogonal cubic, so rec_basis = (1/a)·I → inv = a·I) + h_cubic_f = s * cubic_side # (M, 3), real float + h_idx = torch.round(h_cubic_f).long() # (M, 3), nearest integer cubic HKL + # Wrap into [0, N) (negative HKL → high indices) + h_idx = h_idx % N + # Accumulate values into the grid + flat_idx = (h_idx[:, 0] * N + h_idx[:, 1]) * N + h_idx[:, 2] + grid_flat = torch.zeros(N * N * N, dtype=dtype, device=device) + grid_flat.index_add_(0, flat_idx, F_sq_real) + grid = grid_flat.view(N, N, N) + return grid, cubic_cell + + +def _sphere_mask(N: int, cubic_side: float, omega: float, + dtype: torch.dtype, device: torch.device) -> torch.Tensor: + """ + Build a sphere indicator mask on an (N, N, N) real-space grid covering the + cubic cell (side = cubic_side). The mask is 1 inside |r| < omega, 0 outside. + The grid layout is FFT-compatible: index 0 = origin, indices > N/2 wrap to + negative coordinates. + """ + idx = torch.arange(N, device=device) + # Voxel offsets, wrapped to [-N/2, N/2) + offsets = torch.where(idx < (N + 1) // 2, idx, idx - N).to(dtype) * (cubic_side / N) + rx, ry, rz = torch.meshgrid(offsets, offsets, offsets, indexing="ij") + r2 = rx ** 2 + ry ** 2 + rz ** 2 + return (r2 < omega ** 2).to(dtype) + + +def restrict_to_sphere( + F_squared: torch.Tensor, + hkl_real: torch.Tensor, + real_cell: Cell, + omega_A: float, + cubic_side_A: float | None = None, + gridsize: int | None = None, + max_res_A: float = 3.0, +) -> Tuple[torch.Tensor, Cell, torch.Tensor]: + """ + Sphere-restrict a Patterson-like coefficient set defined at the real-cell + HKL positions. + + Parameters + ---------- + F_squared : torch.Tensor (real), shape (M,) + Values to filter (e.g. (|E|² − 1) per shell, or |F|² with origin removed). + hkl_real : torch.Tensor, shape (M, 3) + Miller indices in the real (crystal) cell. + real_cell : Cell + omega_A : float + Sphere-of-integration radius (Å). Typically the molecule's bounding-box + radius — vectors longer than 2·omega are intermolecular and get filtered out. + cubic_side_A : float, optional + Cubic cell side (Å) for the FFT. Default: 4 · omega_A (so the Patterson + sphere fits comfortably with padding to avoid wraparound aliasing). + gridsize : int, optional + Grid size (one dim). Default: 2 · ceil(cubic_side / max_res_A). + max_res_A : float, default 3.0 + Resolution limit (Å) for the cubic grid spacing. + + Returns + ------- + filtered_values : torch.Tensor (real), shape (M,) + Sphere-restricted Patterson coefficients at the SAME real-cell HKL set. + cubic_cell : Cell + cubic_grid : torch.Tensor (real complex-real), shape (N, N, N) + Final sphere-restricted reciprocal grid (kept for caller introspection). + """ + if cubic_side_A is None: + cubic_side_A = 4.0 * omega_A + if gridsize is None: + gridsize = 2 * int(math.ceil(cubic_side_A / max_res_A)) + # Round to even for cleaner FFT + if gridsize % 2: + gridsize += 1 + + grid_recip, cubic_cell = _make_cubic_grid( + F_squared, hkl_real, real_cell, cubic_side_A, gridsize, + ) + # FFT to real space (the grid is real-valued) + patterson_real = torch.fft.fftn(grid_recip, dim=(0, 1, 2)).real + mask = _sphere_mask(gridsize, cubic_side_A, omega_A, + dtype=patterson_real.dtype, device=patterson_real.device) + patterson_real_masked = patterson_real * mask + # IFFT back (will be approximately real for a real Patterson) + grid_recip_filtered = torch.fft.ifftn( + patterson_real_masked, dim=(0, 1, 2) + ).real + + # Extract values at original real-cell HKL positions via trilinear + # interpolation. Note: interpolate_structure_factor_from_grid expects a + # complex grid for `interpolate_amplitude=False` and a complex or real grid + # for `interpolate_amplitude=True`. Our grid is real — wrap in complex for + # the utility. + rec_real = real_cell.reciprocal_basis_matrix.to(F_squared.device).to(F_squared.dtype) + s = hkl_real.to(F_squared.dtype) @ rec_real + h_cubic_f = s * cubic_side_A # cubic-cell float HKL + grid_complex = grid_recip_filtered.to(torch.complex64) + filtered = interpolate_structure_factor_from_grid( + grid_complex, h_cubic_f, interpolate_amplitude=False, + ).real.to(F_squared.dtype) + return filtered, cubic_cell, grid_recip_filtered diff --git a/torchref/alignment/pipeline.py b/torchref/alignment/pipeline.py index 8f3c9f7d..fd4e8a0a 100644 --- a/torchref/alignment/pipeline.py +++ b/torchref/alignment/pipeline.py @@ -15,14 +15,25 @@ import numpy as np import torch -from .ball_transform import ( - ball_rotation_search_torch, - rotation_matrix_from_euler_zyz, +from .ball_search import ( + ball_rotation_search, + rotation_matrix_from_edmonds_euler, + RotationPeak, ) from .translation import fft_translation_search_torch, TranslationPeak from .rigid_body import RigidBodyRefinement, RigidBodyResult from .clashscore import ClashScoreCalculator, AtomSampler + +def rotation_matrix_from_euler_zyz(alpha, beta, gamma) -> np.ndarray: + """ + Build R = R_z(α) R_y(β) R_z(γ) (Edmonds active ZYZ) as a NumPy 3×3 matrix. + + Compatibility wrapper around `rotation_matrix_from_edmonds_euler`. + """ + R = rotation_matrix_from_edmonds_euler(float(alpha), float(beta), float(gamma)) + return R.detach().cpu().numpy() + if TYPE_CHECKING: from torchref.model import ModelFT from torchref.io.datasets import ReflectionData @@ -227,7 +238,7 @@ def __init__( def run( self, - n_rotation_peaks: int = 100, + n_rotation_peaks: int = 200, n_translation_peaks: int = 5, min_tries: int = 3, max_tries: int = 10, @@ -235,9 +246,9 @@ def run( max_clash_score: float = 100.0, d_min: float = 4.0, d_max: float = 50.0, - L: int = 32, - P: int = 20, - cluster_threshold_deg: float = 6.0, + L: int = 48, + P: int = 24, + cluster_threshold_deg: Optional[float] = None, ) -> List[MRSolution]: """ Run full MR pipeline with early stopping. @@ -274,6 +285,10 @@ def run( List[MRSolution] Solutions sorted by R-factor. Early stops if converged. """ + # Auto cluster threshold: tied to grid voxel resolution. + if cluster_threshold_deg is None: + cluster_threshold_deg = max(6.0, 180.0 / L) + # Step 1: Rotation search if self.verbose: print("Step 1: Rotation search...") @@ -396,17 +411,20 @@ def _rotation_search( L: int, P: int, ) -> list: - """Run ball rotation search.""" + """Run ball rotation search; return list of (α, β, γ, score, σ) tuples.""" # Prepare E-values E_obs, s_obs = self._get_e_values_obs(d_min, d_max) E_calc, s_calc = self._get_e_values_calc(d_min, d_max) - _, _, peaks = ball_rotation_search_torch( - E_obs, s_obs, E_calc, s_calc, - L=L, P=P, d_min=d_min, d_max=d_max, n_peaks=n_peaks, - verbose=self.verbose > 1, + _C, _alphas, _betas, _gammas, peaks = ball_rotation_search( + s_obs, E_obs, s_calc, E_calc, + L=L, P=P, n_peaks=n_peaks, + d_min=d_min, d_max=d_max, + refine_subvoxel=True, n_refine=min(n_peaks, 50), + sigma_threshold=0.0, ) - return peaks + # Convert dataclass to tuple form expected by the rest of the pipeline. + return [(p.alpha, p.beta, p.gamma, p.score, p.sigma) for p in peaks] def _translation_search( self, diff --git a/torchref/alignment/rigid_body.py b/torchref/alignment/rigid_body.py index 070c0921..f8c3b133 100644 --- a/torchref/alignment/rigid_body.py +++ b/torchref/alignment/rigid_body.py @@ -25,7 +25,7 @@ import torch import torch.nn as nn -from torchref.scaling import ScalerBase +from torchref.scaling import Scaler from torchref.model import SfFFT from torchref.symmetry import spacegroup from torchref.refinement.targets import MaximumLikelihoodXrayTarget @@ -130,10 +130,19 @@ def __init__( centroid = torch.mean(self.xyz_initial, dim=0) self.register_buffer("centroid", centroid) - self.cell = data.cell + # Move Cell (which holds fractional_matrix etc.) to the target device + # so RigidBodyRefinement.get_transformed_xyz can mm a GPU tensor + # against fractional_matrix.T without "mat2 on cpu" crashes. + self.cell = data.cell.to(device=device) if hasattr(data.cell, "to") else data.cell self.spacegroup = data.spacegroup - self.fft = SfFFT(self.cell, self.spacegroup, max_res=max_res) + # Forward `device` to SfFFT — its `setup_grid` reads `self.device` + # to allocate the real-space grid; without this, the grid lands on + # CPU even when the surrounding RigidBodyRefinement is on cuda, and + # the joint refine crashes with "mat2 on cuda, others on cpu" at + # the first compute_structure_factors call. + self.fft = SfFFT(self.cell, self.spacegroup, max_res=max_res, + device=device) self.verbose = verbose @@ -170,10 +179,18 @@ def __init__( initial_translation = initial_translation.to(device=device).clone() self.translation_frac = nn.Parameter(initial_translation) - self.scaler = ScalerBase(data=data, nbins=20, verbose=0, device=device) + # Use `Scaler` for per-bin scales + anisotropy correction, but skip + # the bulk-solvent setup. The solvent mask is computed once from + # `model.xyz()` and goes stale as the joint refine moves atoms; the + # mismatch then biases the LBFGS gradient. Better to leave solvent + # out of the joint refine — `fit_to_data` does a fresh + # solvent-aware Scaler refit on the final polished model for the + # user-facing R-work. + self.scaler = Scaler(model=model, data=data, nbins=20, + verbose=0, device=device) fcalc_initial = self() - - self.scaler.initialize(fcalc_initial) + self.scaler.calc_initial_scale(fcalc_initial) + self.scaler.setup_anisotropy_correction() self.scaler.refine_lbfgs(fcalc=fcalc_initial) self.xray_target = MaximumLikelihoodXrayTarget(data=self.data, scaler=self.scaler) @@ -232,8 +249,9 @@ def get_transformed_xyz(self) -> torch.Tensor: xyz_centered = self.xyz_initial - self.centroid xyz_rotated = xyz_centered @ R.T + self.centroid - # Apply translation (fractional -> Cartesian) - t_cart = self.translation_frac @ self.cell.fractional_matrix + # Apply translation (fractional → Cartesian via cell.fractional_matrix.T; + # see Cell.fractional_to_cartesian for the canonical convention). + t_cart = self.translation_frac @ self.cell.fractional_matrix.T return xyz_rotated + t_cart def get_scale(self) -> float: @@ -273,7 +291,7 @@ def forward(self, debug: bool = False) -> torch.Tensor: R = self.get_rotation_matrix() xyz_aniso_centered = self.xyz_aniso_original - self.centroid xyz_aniso_rotated = xyz_aniso_centered @ R.T + self.centroid - t_cart = self.translation_frac @ self.cell.fractional_matrix + t_cart = self.translation_frac @ self.cell.fractional_matrix.T xyz_aniso = xyz_aniso_rotated + t_cart # Compute structure factors via FFT (bypasses MixedTensor!) diff --git a/torchref/alignment/sh.py b/torchref/alignment/sh.py new file mode 100644 index 00000000..61839d36 --- /dev/null +++ b/torchref/alignment/sh.py @@ -0,0 +1,520 @@ +""" +Pure-PyTorch spherical harmonic expansion for the alignment module. + +Conventions (locked, asserted by tests/unit/alignment/test_sh.py): + + Y_{l,m}(θ, φ) = (-1)^m · √[(2l+1)/(4π) · (l-m)!/(l+m)!] · P_l^m(cos θ) · e^{imφ} (m ≥ 0) + Y_{l,-m}(θ, φ) = (-1)^m · conj(Y_{l,m}(θ, φ)) (m > 0) + +i.e. fully orthonormal physics convention with Condon-Shortley phase included. +Matches scipy.special.sph_harm and the convention used in Edmonds, Sakurai, etc. + +Numerical core is a stable forward recurrence on the fully-normalized associated +Legendre `bar_P_l^m(cosθ) = √[(2l+1)/(4π) · (l-m)!/(l+m)!] · P_l^m(cosθ)` so we +never form (2l)! explicitly. +""" + +from __future__ import annotations + +import math +from typing import Optional, Tuple + +import torch + + +def _bar_legendre_recurrence( + cos_theta: torch.Tensor, + sin_theta: torch.Tensor, + L: int, +) -> torch.Tensor: + """ + Compute fully-normalized associated Legendre `bar_P_l^m(cos θ)` for + all l in [0, L), m in [0, l]. + + Definition: + bar_P_l^m(x) = √[(2l+1)/(4π) · (l-m)!/(l+m)!] · P_l^m(x) + where P_l^m is the *unsigned* associated Legendre (no Condon-Shortley phase). + + Returns + ------- + bar_P : torch.Tensor, real + Shape (..., L, L). `bar_P[..., l, m]` is bar_P_l^m for m <= l, else 0. + """ + batch_shape = cos_theta.shape + dtype = cos_theta.dtype + device = cos_theta.device + + bar_P = torch.zeros((*batch_shape, L, L), dtype=dtype, device=device) + + # Seed: bar_P_0^0 = 1 / sqrt(4π) + inv_sqrt_4pi = 1.0 / math.sqrt(4.0 * math.pi) + bar_P[..., 0, 0] = inv_sqrt_4pi + + # Sectoral recurrence on the diagonal m == l: + # bar_P_m^m = sqrt((2m+1)/(2m)) · sinθ · bar_P_{m-1}^{m-1} + for m in range(1, L): + factor = math.sqrt((2.0 * m + 1.0) / (2.0 * m)) + bar_P[..., m, m] = factor * sin_theta * bar_P[..., m - 1, m - 1] + + # Vertical recurrence (l > m, fixed m): + # bar_P_l^m = a_l^m · cosθ · bar_P_{l-1}^m - b_l^m · bar_P_{l-2}^m + # with a_l^m = sqrt((2l-1)(2l+1)/((l-m)(l+m))) + # b_l^m = sqrt((2l+1)(l+m-1)(l-m-1)/((l-m)(l+m)(2l-3))) [0 when l = m+1] + for m in range(0, L - 1): + # l = m+1 step (b term vanishes because l-m-1 = 0) + l = m + 1 + a = math.sqrt((2.0 * l - 1.0) * (2.0 * l + 1.0) / ((l - m) * (l + m))) + bar_P[..., l, m] = a * cos_theta * bar_P[..., l - 1, m] + # l from m+2 to L-1 + for l in range(m + 2, L): + a = math.sqrt((2.0 * l - 1.0) * (2.0 * l + 1.0) / ((l - m) * (l + m))) + b = math.sqrt( + (2.0 * l + 1.0) * (l + m - 1.0) * (l - m - 1.0) + / ((l - m) * (l + m) * (2.0 * l - 3.0)) + ) + bar_P[..., l, m] = a * cos_theta * bar_P[..., l - 1, m] - b * bar_P[..., l - 2, m] + + return bar_P + + +def evaluate_ylm( + theta: torch.Tensor, + phi: torch.Tensor, + L: int, +) -> torch.Tensor: + """ + Evaluate Y_{l,m}(θ, φ) for all (l, m) with l ∈ [0, L), m ∈ [-(L-1), L-1]. + + Parameters + ---------- + theta : torch.Tensor + Polar angle, shape (...,), values in [0, π]. + phi : torch.Tensor + Azimuthal angle, shape (...,), values in [0, 2π). + L : int + Maximum SH degree (exclusive: l_max = L - 1). + + Returns + ------- + Y : torch.Tensor, complex + Shape (..., L, 2L-1). Y[..., l, L-1+m] = Y_{l,m}(θ, φ) for |m| ≤ l, + zero otherwise. dtype is complex128 if input is float64, else complex64. + """ + assert theta.shape == phi.shape, "theta and phi must have the same shape" + + real_dtype = theta.dtype + if real_dtype == torch.float64: + complex_dtype = torch.complex128 + elif real_dtype == torch.float32: + complex_dtype = torch.complex64 + else: + raise TypeError(f"Unsupported real dtype: {real_dtype}") + + device = theta.device + cos_theta = torch.cos(theta) + sin_theta = torch.sin(theta).clamp(min=0.0) # numerical floor at the poles + + bar_P = _bar_legendre_recurrence(cos_theta, sin_theta, L) # (..., L, L) + + # Y_{l,m}(θ,φ) = (-1)^m · bar_P_l^m(cosθ) · e^{i m φ} for m ≥ 0 + # Y_{l,-m} = (-1)^m · conj(Y_{l,m}) for m > 0 + Y = torch.zeros((*theta.shape, L, 2 * L - 1), dtype=complex_dtype, device=device) + + # Precompute e^{i m φ} for m = 0..L-1 + # (use stacking to keep things vectorized) + m_vals = torch.arange(L, dtype=real_dtype, device=device) + m_phi = phi.unsqueeze(-1) * m_vals # (..., L) + expo = torch.complex(torch.cos(m_phi), torch.sin(m_phi)) # (..., L), e^{i m φ} + + # Fill m ≥ 0 columns + for m in range(L): + sign = (-1.0) ** m + # bar_P[..., :, m] is (..., L); only entries l >= m are non-zero (others left at 0) + Y[..., :, L - 1 + m] = sign * bar_P[..., :, m].to(complex_dtype) * expo[..., m].unsqueeze(-1) + + # Fill m < 0 columns by hermitian symmetry: Y_{l,-m} = (-1)^m · conj(Y_{l,m}) + for m in range(1, L): + sign = (-1.0) ** m + Y[..., :, L - 1 - m] = sign * torch.conj(Y[..., :, L - 1 + m]) + + return Y + + +def angular_density_weights( + s_vectors: torch.Tensor, + k_neighbors: int = 12, +) -> torch.Tensor: + """ + Per-sample weights that compensate for non-uniform angular sampling on the + unit sphere. Returns w_i ∝ (1 / local_density)^... so the weighted sum + `Σ_i w_i · v_i · Y*_lm(ŝ_i)` is an unbiased Monte-Carlo estimate of the + SH integral on the sphere. + + Heuristic: w_i ~ (k-th NN great-circle distance)². Normalised so that + Σ w_i = N. + + Pure-torch O(N · k) memory; suitable up to N ~ 30k. For larger N use a + chunked KNN, but typical resolution-cut datasets fit easily. + + Parameters + ---------- + s_vectors : torch.Tensor, shape (N, 3) + k_neighbors : int, default 12 + Number of nearest angular neighbours to estimate local density. + + Returns + ------- + w : torch.Tensor, shape (N,) + """ + device = s_vectors.device + dtype = s_vectors.dtype + N = s_vectors.shape[0] + norm = s_vectors.norm(dim=-1).clamp(min=1e-30) + s_hat = s_vectors / norm.unsqueeze(-1) # (N, 3) + # cos(angle) between every pair via dot product + # Memory: (N, N). For N ~ 30k, ~3 GB at fp32 — chunk if too big. + if N <= 8000: + dots = (s_hat @ s_hat.transpose(0, 1)).clamp(-1.0, 1.0) + ang = torch.acos(dots) # (N, N) + # k-th NN distance (excluding self at column-diagonal). topk smallest. + # Set diagonal large so it doesn't show up as nearest. + ang.fill_diagonal_(float("inf")) + kth_dist, _ = torch.topk(ang, k_neighbors, dim=-1, largest=False) # (N, k) + d_local = kth_dist[:, -1] # k-th NN distance + else: + # Chunked: still O(N²) compute but bounded memory. + chunk = 1024 + d_local = torch.empty(N, dtype=dtype, device=device) + for i0 in range(0, N, chunk): + i1 = min(i0 + chunk, N) + dots = (s_hat[i0:i1] @ s_hat.transpose(0, 1)).clamp(-1.0, 1.0) + ang = torch.acos(dots) + for j, gi in enumerate(range(i0, i1)): + ang[j, gi] = float("inf") + kth_dist, _ = torch.topk(ang, k_neighbors, dim=-1, largest=False) + d_local[i0:i1] = kth_dist[:, -1] + + w = d_local ** 2 # ~ local Voronoi area + w = w * (N / w.sum().clamp(min=1e-30)) # normalise to Σw = N + return w + + +def sh_expand_ball( + s_vectors: torch.Tensor, + values: torch.Tensor, + shell_idx: torch.Tensor, + P: int, + L: int, + enforce_friedel: bool = True, + chunk_size: int = 2048, + angular_weights: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """ + Analytical spherical-harmonic expansion of a scattered-point real field + on a set of radial shells. + + f_{p,l,m} = Σ_{i ∈ shell p} values_i · conj(Y_{l,m}(θ_i, φ_i)) + + When `enforce_friedel=True` the input is augmented with the antipodal copy + `(-s_i, values_i)`. Y_{l,m}(-ŝ) = (-1)^l Y_{l,m}(ŝ), so the sum then has + f_{p,l,m} = (1 + (-1)^l) · Σ_i v_i · Y*_{l,m}(ŝ_i) + i.e. odd-l rows are exactly zero by construction (and even-l rows get a + factor of 2 which we keep — this absorbs into the cross-correlation + normalisation when the same convention is applied to both operands). + + Parameters + ---------- + s_vectors : torch.Tensor + Reciprocal-lattice vectors, shape (N, 3). Direction only is used; + magnitudes do not enter (shell assignment is done by the caller). + values : torch.Tensor + Real-valued samples (e.g. |E(h)|), shape (N,). + shell_idx : torch.Tensor (int64) + Shell index in [0, P) for each reflection, shape (N,). + P : int + Number of radial shells. + L : int + SH bandlimit (l in [0, L)). + enforce_friedel : bool, default True + Augment input with (-s, value) pairs and zero odd-l coefficients. + chunk_size : int + Points per chunk for memory control during Y_lm evaluation. + + Returns + ------- + f_plm : torch.Tensor, complex + Shape (P, L, 2L-1). `f_plm[p, l, L-1+m] = f_{p,l,m}`. + """ + assert s_vectors.dim() == 2 and s_vectors.shape[-1] == 3 + assert values.dim() == 1 and values.shape[0] == s_vectors.shape[0] + assert shell_idx.dim() == 1 and shell_idx.shape[0] == s_vectors.shape[0] + + real_dtype = s_vectors.dtype + if real_dtype == torch.float64: + complex_dtype = torch.complex128 + elif real_dtype == torch.float32: + complex_dtype = torch.complex64 + else: + raise TypeError(f"Unsupported dtype {real_dtype}") + device = s_vectors.device + + if enforce_friedel: + s_vectors = torch.cat([s_vectors, -s_vectors], dim=0) + values = torch.cat([values, values], dim=0) + shell_idx = torch.cat([shell_idx, shell_idx], dim=0) + if angular_weights is not None: + angular_weights = torch.cat([angular_weights, angular_weights], dim=0) + + if angular_weights is not None: + values = values * angular_weights.to(values.dtype) + + # Direction (θ, φ). At |s|=0 the direction is undefined; the caller should + # have excluded F(000), but we guard anyway. + norm = s_vectors.norm(dim=-1).clamp(min=1e-30) + s_hat = s_vectors / norm.unsqueeze(-1) + cos_theta = s_hat[..., 2].clamp(min=-1.0, max=1.0) + theta = torch.acos(cos_theta) + phi = torch.atan2(s_hat[..., 1], s_hat[..., 0]) + + f_plm = torch.zeros((P, L, 2 * L - 1), dtype=complex_dtype, device=device) + + N = s_vectors.shape[0] + for start in range(0, N, chunk_size): + stop = min(start + chunk_size, N) + Y = evaluate_ylm(theta[start:stop], phi[start:stop], L) # (n, L, 2L-1) + # contribution to f_{p,l,m} is value_i * conj(Y_{l,m}(s_i)) + contrib = values[start:stop].to(complex_dtype).view(-1, 1, 1) * torch.conj(Y) + # scatter-add into shells + f_plm.index_add_(0, shell_idx[start:stop], contrib) + + if enforce_friedel: + # Zero odd-l rows explicitly (they should already be ~0; this kills FP drift). + l_vals = torch.arange(L, device=device) + odd_mask = (l_vals % 2 == 1) + f_plm[:, odd_mask, :] = 0.0 + + return f_plm + + +def equal_count_shell_edges( + s_magnitudes: torch.Tensor, + P: int, + s_min: Optional[float] = None, + s_max: Optional[float] = None, +) -> Tuple[torch.Tensor, torch.Tensor]: + """ + Compute equal-count radial shell edges for a list of |s| magnitudes. + + Each of the P shells receives (approximately) the same number of reflections. + Reflections outside [s_min, s_max] (if given) are excluded. + + Returns + ------- + edges : torch.Tensor, real, shape (P+1,) + Shell boundaries, increasing. + centers : torch.Tensor, real, shape (P,) + Mid-points of each shell. + """ + s = s_magnitudes + if s_min is not None: + s = s[s >= s_min] + if s_max is not None: + s = s[s <= s_max] + s_sorted, _ = torch.sort(s) + N = s_sorted.numel() + # quantile-based partition + idx = torch.linspace(0, N - 1, P + 1, dtype=torch.float64, device=s.device).round().long() + edges = s_sorted[idx] + # nudge endpoints so the data is fully covered (avoid floating-point miss) + if s_min is not None: + edges[0] = min(edges[0].item(), s_min) + else: + edges[0] = edges[0] - 1e-6 + if s_max is not None: + edges[-1] = max(edges[-1].item(), s_max) + else: + edges[-1] = edges[-1] + 1e-6 + centers = 0.5 * (edges[:-1] + edges[1:]) + return edges, centers + + +def fit_overall_anisotropy( + F_obs: torch.Tensor, + s_vectors: torch.Tensor, + shell_idx: torch.Tensor, + P: int, + min_count: int = 20, +) -> torch.Tensor: + """ + Fit the overall anisotropy tensor U from F_obs alone (no model needed). + + The Popov-Bourenkov anisotropy correction models the observed structure- + factor amplitudes as a per-shell isotropic Wilson piece modulated by an + overall anisotropic Debye-Waller term: + + |F_obs(h)|² ≈ <|F_iso|²>(s) · exp(−2π²·s·U·s) + + Taking logs and fitting the linear regression + ln |F_obs|² − ln<|F_iso|²>(s) = −2π² s·U·s + over all reflections gives the 6-parameter U directly (linear in U). We + parametrise as U_xx, U_yy, U_zz, U_xy, U_xz, U_yz and ignore reflections + in shells with fewer than `min_count` entries (poor shell mean estimate). + + The returned U is the correction to *apply* in the form + F_obs_corrected(h) = F_obs(h) · exp(+π²·s·U·s) + so that the resulting amplitudes have the same per-shell mean square + regardless of direction. + + Parameters + ---------- + F_obs : (N,) real + s_vectors : (N, 3) real (1/Å) + shell_idx : (N,) int64 — assigns each reflection to a shell in [0, P) + P : int, number of shells + + Returns + ------- + U : (3, 3) symmetric real tensor (Ų) + """ + device = F_obs.device + dtype = F_obs.dtype + valid = shell_idx >= 0 + F = F_obs[valid] + s = s_vectors[valid].to(dtype) + idx = shell_idx[valid] + count = torch.zeros(P, dtype=torch.int64, device=device) + count.index_add_(0, idx, torch.ones_like(idx)) + F2 = F * F + sum_F2 = torch.zeros(P, dtype=dtype, device=device) + sum_F2.index_add_(0, idx, F2) + mean_F2 = sum_F2 / count.clamp(min=1).to(dtype) + # Mask shells with too few reflections + good = count >= min_count + if good.sum() == 0: + return torch.zeros((3, 3), dtype=dtype, device=device) + keep = good[idx] + F2k = F2[keep].clamp(min=1e-30) + sk = s[keep] + mean_F2_k = mean_F2[idx[keep]].clamp(min=1e-30) + # y = ln|F|² - ln<|F|²> = -2π² · sUs + y = (torch.log(F2k) - torch.log(mean_F2_k)).to(torch.float64) + sk = sk.to(torch.float64) + # Design matrix X for u = (Uxx, Uyy, Uzz, Uxy, Uxz, Uyz): + # s·U·s = Uxx sx² + Uyy sy² + Uzz sz² + 2 Uxy sx sy + 2 Uxz sx sz + 2 Uyz sy sz + X = torch.stack([ + sk[:, 0] ** 2, sk[:, 1] ** 2, sk[:, 2] ** 2, + 2.0 * sk[:, 0] * sk[:, 1], + 2.0 * sk[:, 0] * sk[:, 2], + 2.0 * sk[:, 1] * sk[:, 2], + ], dim=-1) + A = -2.0 * (torch.pi ** 2) * X # y ≈ A · u + # Least-squares solve A u = y + u_vec, _, _, _ = torch.linalg.lstsq(A, y.unsqueeze(-1)) + u_vec = u_vec.squeeze(-1) + Uxx, Uyy, Uzz, Uxy, Uxz, Uyz = u_vec.tolist() + U = torch.tensor( + [[Uxx, Uxy, Uxz], [Uxy, Uyy, Uyz], [Uxz, Uyz, Uzz]], + dtype=dtype, device=device, + ) + return U + + +def apply_overall_anisotropy( + F: torch.Tensor, + s_vectors: torch.Tensor, + U: torch.Tensor, +) -> torch.Tensor: + """ + Apply the inverse of the anisotropy tensor to amplitudes: + F_corrected(h) = F(h) · exp(+π²·s·U·s) + (Inverse direction = +π² to undo the Debye-Waller effect.) + """ + device = F.device + dtype = F.dtype + s = s_vectors.to(device).to(dtype) + U_t = U.to(device).to(dtype) + s_dot_U = s @ U_t # (N, 3) + arg = (torch.pi ** 2) * (s_dot_U * s).sum(dim=-1) + return F * torch.exp(arg.clamp(min=-10.0, max=10.0)) + + +def compute_patterson_shell_variance( + patt: torch.Tensor, + shell_idx: torch.Tensor, + P: int, + min_count: int = 8, + eps: float = 1e-3, +) -> torch.Tensor: + """ + Empirical per-shell variance of a Patterson coefficient (typically E²−1). + + For each shell p, returns Var_p = _p - _p². + + Shells with fewer than `min_count` reflections inherit the variance of the + nearest shell that does meet the count threshold (search outward, prefer + higher-resolution / larger-index shells first). The result is clamped at + `eps` so the inverse-sqrt weight `1/√Var_p` stays bounded. + + Parameters + ---------- + patt : torch.Tensor, shape (N,) + Per-reflection Patterson coefficient (e.g. E² − 1). + shell_idx : torch.Tensor, shape (N,), int64 + Shell assignment in [0, P). Reflections with index -1 are ignored. + P : int + Number of shells. + min_count : int, default 8 + Shells with fewer reflections than this borrow from a neighbour. + eps : float, default 1e-3 + Lower bound on returned variance. + + Returns + ------- + var : torch.Tensor, shape (P,), same dtype as `patt` + """ + device = patt.device + dtype = patt.dtype + valid = shell_idx >= 0 + patt_v = patt[valid] + idx_v = shell_idx[valid] + count = torch.zeros(P, dtype=torch.int64, device=device) + count.index_add_(0, idx_v, torch.ones_like(idx_v)) + sum1 = torch.zeros(P, dtype=dtype, device=device) + sum2 = torch.zeros(P, dtype=dtype, device=device) + sum1.index_add_(0, idx_v, patt_v) + sum2.index_add_(0, idx_v, patt_v * patt_v) + safe_count = count.clamp(min=1).to(dtype) + mean = sum1 / safe_count + var = sum2 / safe_count - mean * mean + var = var.clamp(min=eps) + + counts_l = count.tolist() + var_l = var.tolist() + good = [i for i, c in enumerate(counts_l) if c >= min_count] + if not good: + return torch.full_like(var, max(eps, 1.0)) + for i in range(P): + if counts_l[i] >= min_count: + continue + nearest = min(good, key=lambda g, ii=i: (abs(g - ii), -g)) + var_l[i] = var_l[nearest] + return torch.tensor(var_l, dtype=dtype, device=device).clamp(min=eps) + + +def assign_shells( + s_magnitudes: torch.Tensor, + edges: torch.Tensor, +) -> torch.Tensor: + """ + Assign each reflection to a shell index in [0, P). + + Reflections strictly outside the edges get index -1 (caller may filter). + """ + # bucketize returns indices in [0, P+1]; we want shell indices in [0, P). + # Reflections with s == edges[0] go to bucket 0 (use right=True trick). + idx = torch.bucketize(s_magnitudes.contiguous(), edges.contiguous(), right=True) - 1 + P = edges.shape[0] - 1 + invalid = (idx < 0) | (idx >= P) + idx = idx.clamp(min=0, max=P - 1) + idx = torch.where(invalid, torch.full_like(idx, -1), idx) + return idx diff --git a/torchref/alignment/translation.py b/torchref/alignment/translation.py index a5d16914..de613940 100644 --- a/torchref/alignment/translation.py +++ b/torchref/alignment/translation.py @@ -242,6 +242,700 @@ def apply_translation_to_fcalc( return F_calc * np.exp(1j * phase_shift) +def amplitude_translation_search( + F_obs: torch.Tensor, + interpolator, + R_rotation: torch.Tensor, + hkl: torch.Tensor, + spacegroup, + real_cell, + grid_steps: int = 16, + n_peaks: int = 20, + cluster_radius: float = 0.05, + batch_size: int = 256, + use_e_values: bool = True, + n_shells: int = 20, + precomputed_G: Optional[torch.Tensor] = None, + precomputed_h_R: Optional[torch.Tensor] = None, +) -> Tuple[np.ndarray, np.ndarray, List[TranslationPeak]]: + """ + Coarse-grid translation search via |F|²-correlation. + + For each candidate fractional translation `t` on a `grid_steps`³ grid in + `[0, 1)³`, scores the model at the current rotation translated by `t` + against the observed amplitudes by Pearson correlation of `|F_obs|²` and + `|F_calc(h, t)|²`. The structure-factor sum uses the spacegroup symmetry + expansion + + F_calc(h, t) = Σ_i G_i(h) · exp(2πi (h R_i) · t) + G_i(h) = exp(2πi h · t_i) · F_p1(h R_i) + + with `F_p1(h R_i)` looked up via the supplied interpolator at the rotation + already applied to the model. The `G_i` factors are computed once; only the + phase exponential changes per candidate, so the scan is efficient. + + Parameters + ---------- + F_obs : torch.Tensor, shape (N,) + Observed amplitudes (complex inputs are coerced to |·|). + interpolator : LattmanLoveInterpolator + Provides `evaluate(R, hkl, real_cell, return_amplitude=False)`. + R_rotation : torch.Tensor, shape (3, 3) + Rotation that has been applied to the model coordinates. + hkl : torch.Tensor, shape (N, 3) + Integer Miller indices of the observed reflections. + spacegroup : SpaceGroup + Provides `matrices` and `translations`. + real_cell : Cell + Real crystal cell. + grid_steps : int, default 16 + Per-axis grid resolution. Total candidates = grid_steps³. + n_peaks : int, default 20 + Number of peaks returned (after clustering). + cluster_radius : float, default 0.05 + Minimum fractional separation between returned peaks. + batch_size : int, default 256 + Number of candidate translations evaluated per inner batch. + use_e_values : bool, default True + Normalize `|F_obs|` and `|F_calc(h, t)|` per resolution shell to unit + Wilson variance (E-values) before correlating. This removes the + resolution-dependent envelope mismatch between real F_obs (with bulk + solvent + thermal falloff) and a model that doesn't model these — a + per-shell mean subtraction in the Pearson correlation alone doesn't + cover it because the falloff is multiplicative, not additive. + n_shells : int, default 20 + Number of equal-count radial shells used by `use_e_values`. + + Returns + ------- + correlation_map : np.ndarray, shape (grid_steps, grid_steps, grid_steps) + Pearson correlation of |F_obs|² and |F_calc(t)|² at each grid point. + best_translation : np.ndarray, shape (3,) + Top-scoring fractional translation. + peaks : list of TranslationPeak + Top-`n_peaks` peaks sorted by descending correlation. + """ + device = getattr(interpolator, "device", hkl.device) + real_dtype = torch.float64 + complex_dtype = torch.complex128 + + F_obs_t = F_obs.detach().to(device) + if F_obs_t.is_complex(): + F_obs_t = F_obs_t.abs() + F_obs_t = F_obs_t.to(real_dtype) + + hkl_t = hkl.detach().to(device).to(real_dtype) # (N, 3) + + # Precompute per-shell normalisation if requested. We bin reflections by + # |s| into n_shells equal-count shells and normalise F → F / sqrt(_shell) + # (Wilson E-value). The same shell norm is applied to F_calc(t) inside the + # batch loop. This makes the Pearson correlation a Patterson-style + # correlation of "E²−1" — robust to bulk-solvent / B-factor mismatch. + if use_e_values: + rec_basis_real = real_cell.reciprocal_basis_matrix.to(device).to(real_dtype) + s_mag = (hkl_t @ rec_basis_real).norm(dim=-1) + order = torch.argsort(s_mag) + shell_idx = torch.zeros_like(s_mag, dtype=torch.int64) + chunk = s_mag.numel() // max(n_shells, 1) + for k in range(n_shells): + a = k * chunk + b = (k + 1) * chunk if k < n_shells - 1 else s_mag.numel() + shell_idx[order[a:b]] = k + shell_norm_obs = torch.zeros(n_shells, dtype=real_dtype, device=device) + for k in range(n_shells): + m = shell_idx == k + if m.any(): + shell_norm_obs[k] = (F_obs_t[m] ** 2).mean().clamp(min=1e-30).sqrt() + E_obs = F_obs_t / shell_norm_obs[shell_idx] + F_obs2 = E_obs * E_obs + else: + shell_idx = None + F_obs2 = F_obs_t * F_obs_t + F_obs2_centered = F_obs2 - F_obs2.mean() + + # Pre-compute G_i(h) = exp(2πi h·t_i) · F_p1(h R_i) (or reuse caller's) + two_pi_i = 2j * torch.pi + if precomputed_G is not None and precomputed_h_R is not None: + G = precomputed_G.to(device).to(complex_dtype) + h_R = precomputed_h_R.to(device).to(real_dtype) + else: + G, h_R = precompute_G_for_rotation( + interpolator, R_rotation, hkl, spacegroup, real_cell, device=device, + ) + + # Crowther–Blow FFT translation function (Acta Cryst. B23 (1967) 544). + # The grid-evaluated score + # num(t) = Σ_h F_obs²_centered(h) · |F_calc(h, t)|² + # expands as + # num(t) = Σ_{i,j} [Σ_h F_obs²_c(h) · G_i*(h) · G_j(h)] + # · exp(2πi · (h·R_j − h·R_i) · t) + # and on a regular fractional t-grid t = (jx, jy, jz) / G this is exactly + # an inverse DFT of the bracketed coefficients accumulated onto a 3-D + # reciprocal grid at integer indices (h·R_j − h·R_i) mod G. + # + # We accumulate two such reciprocal grids in one sym-op pass: + # W_num : weight per h = F_obs²_centered(h) → num(t) + # W_den : weight per h = 1 → Σ_h |F_calc(h,t)|² + # Score(t) = num(t) / Σ_h|F_calc(h,t)|² — a per-t scale-normalised + # Pearson proxy (Phaser's TF uses the full Pearson denominator; ours + # uses the same scaling that the previous separable-phase code applied + # via explicit per-t centering, achieved here without materialising + # |F_calc(h,t)|² per t-point). + # + # One IFFT pair replaces G³ grid evaluations — for our defaults this is + # ~5000× less arithmetic than the separable-phase scoring it supersedes, + # and orders of magnitude less than the original explicit grid loop. + S_eff, N_eff = G.shape + h_R_int = h_R.round().to(torch.int64) # (S, N, 3) + F_obs2_c_complex = F_obs2_centered.to(complex_dtype) # (N,) + ones_complex = torch.ones(N_eff, dtype=complex_dtype, device=device) + + W_num_flat = torch.zeros( + grid_steps ** 3, dtype=complex_dtype, device=device, + ) + W_den_flat = torch.zeros( + grid_steps ** 3, dtype=complex_dtype, device=device, + ) + G_stride_xy = grid_steps * grid_steps + for i in range(S_eff): + Gi_conj = G[i].conj() # (N,) + pair = Gi_conj.view(1, -1) * G # (S, N) + coeff_num = F_obs2_c_complex.view(1, -1) * pair # (S, N) + coeff_den = ones_complex.view(1, -1) * pair # (S, N) + dh = (h_R_int - h_R_int[i:i + 1]) % grid_steps # (S, N, 3) + flat = (dh[..., 0] * G_stride_xy + + dh[..., 1] * grid_steps + dh[..., 2]) # (S, N) + flat_flat = flat.reshape(-1) + W_num_flat.index_add_(0, flat_flat, coeff_num.reshape(-1)) + W_den_flat.index_add_(0, flat_flat, coeff_den.reshape(-1)) + + W_num = W_num_flat.view(grid_steps, grid_steps, grid_steps) + W_den = W_den_flat.view(grid_steps, grid_steps, grid_steps) + # IFFT scales by 1/G³; undo so values are raw integrals. + num_t = (torch.fft.ifftn(W_num, dim=(0, 1, 2)).real + * (grid_steps ** 3)).to(real_dtype) + den_t = (torch.fft.ifftn(W_den, dim=(0, 1, 2)).real + * (grid_steps ** 3)).to(real_dtype) + corr_map = num_t / den_t.clamp(min=1e-30) + corr_map_np = corr_map.detach().cpu().numpy().astype(np.float32) + peaks = find_translation_peaks(corr_map_np, n_peaks=n_peaks, + cluster_radius=cluster_radius) + best = peaks[0].translation if peaks else np.zeros(3) + return corr_map_np, best, peaks + + +def precompute_G_for_rotation( + interpolator, + R_rotation: torch.Tensor, + hkl: torch.Tensor, + spacegroup, + real_cell, + device=None, +): + """ + Pre-compute per-symmetry F_asu contributions `G_i(h)` for a fixed rotation. + + These are the only inputs that depend on `R_rotation` (and therefore on + expensive interpolator/model-forward evaluations). Passing the result + into `amplitude_translation_search` and `local_translation_refine` lets + them share a single set of (n_sym) model evaluations across the coarse + TF and the fine refinement, instead of recomputing each call. + + Returns + ------- + G : (S, N) complex128 + h_R : (S, N, 3) float64 + """ + real_dtype = torch.float64 + complex_dtype = torch.complex128 + if device is None: + device = getattr(interpolator, "device", hkl.device) + + hkl_t = hkl.detach().to(device).to(real_dtype) + sym_R = spacegroup.matrices.detach().to(device).to(real_dtype) + sym_t = spacegroup.translations.detach().to(device).to(real_dtype) + S = sym_R.shape[0] + N = hkl_t.shape[0] + R_rot = R_rotation.detach().to(device).to(real_dtype) + + # Batched: h_R[i, n, d] = Σ_e hkl[n, e] · sym_R[i, e, d] + h_R = torch.einsum("ne,ied->ind", hkl_t, sym_R) # (S, N, 3) + # Per-sym-op translation phase: exp(2πi · h · t_i) + two_pi_i = 2j * torch.pi + phase_arg = torch.einsum("ne,ie->in", hkl_t, sym_t) # (S, N) + phase = torch.exp(two_pi_i * phase_arg.to(complex_dtype)) # (S, N) + + # One interpolator.evaluate over all (S × N) rotated indices: lets the + # backend do a single grid_sample instead of S sequential ones. + h_R_flat = h_R.reshape(-1, 3) # (S·N, 3) + F_flat = interpolator.evaluate( + R_rot, h_R_flat, real_cell, return_amplitude=False, + ) + F_all = F_flat.reshape(S, N).to(complex_dtype) # (S, N) + + G = F_all * phase + return G, h_R + + +def local_translation_refine( + F_obs: torch.Tensor, + interpolator, + R_rotation: torch.Tensor, + hkl: torch.Tensor, + spacegroup, + real_cell, + t_init: torch.Tensor, + radius: float = 0.06, + grid_steps: int = 13, + n_refinement_passes: int = 2, + batch_size: int = 1024, + precomputed_G: Optional[torch.Tensor] = None, + precomputed_h_R: Optional[torch.Tensor] = None, +) -> Tuple[torch.Tensor, float]: + """ + Fine-grid Patterson translation refinement around ``t_init``. + + For each candidate ``t`` in a `grid_steps`³ cubic grid of half-width + ``radius`` centered on ``t_init`` (fractional), computes |F_calc(h, t)|² + via the symmetry expansion and the analytical-scale R-factor + R(t) = Σ ||F_obs| − k·|F_calc(t)|| / Σ |F_obs| + k(t) = Σ |F_obs|·|F_calc(t)| / Σ |F_calc(t)|² + against `F_obs`. Returns the (t, R) at the minimum. + + Use `n_refinement_passes > 1` to do a multi-pass zoom: each pass shrinks + the radius by `grid_steps/2` and re-centers on the previous best. For + `radius=0.06, grid_steps=13, n_refinement_passes=2`, the final fractional + resolution is ~0.005 (≈0.3 Å for a 60 Å cell). + + The analytical-scale R-factor uses a single global scale; it is not the + same number a full crystallographic Scaler would return, but its + *minimum location* is robust because both numerator and denominator share + the same per-shell envelope. Use a full Scaler to compute the final + R-work after this routine selects (R, t). + """ + device = getattr(interpolator, "device", hkl.device) + real_dtype = torch.float64 + complex_dtype = torch.complex128 + + F_obs_t = F_obs.detach().to(device) + if F_obs_t.is_complex(): + F_obs_t = F_obs_t.abs() + F_obs_t = F_obs_t.to(real_dtype) + F_obs_sum = F_obs_t.sum().clamp(min=1e-30) + + two_pi_i = 2j * torch.pi + if precomputed_G is not None and precomputed_h_R is not None: + G = precomputed_G.to(device).to(complex_dtype) + h_R = precomputed_h_R.to(device).to(real_dtype) + else: + G, h_R = precompute_G_for_rotation( + interpolator, R_rotation, hkl, spacegroup, real_cell, device=device, + ) + + # Adapt batch_size to keep the largest inner einsum tensor under ~250 MB + # of complex128. The (S, B, N) phase tensor is the offender: + # S × B × N × 16 bytes ≤ 2.5e8 → B ≤ 2.5e8 / (16 × S × N). + S_eff, N_eff = G.shape + safe_b = max(8, int(2.5e8 / (16.0 * max(S_eff, 1) * max(N_eff, 1)))) + batch_size = min(batch_size, safe_b) + + # Crowther–Blow FFT refinement on a fine grid around t_init. + # Bake t_init into G as a per-h_R phase factor, then the IFFT trick from + # `amplitude_translation_search` works on the offset grid Δt with the + # same (num/den) Pearson-proxy scoring. The previous nested-grid Python + # evaluation paid O(G³ · N · S) Bessel/exp/einsum per call; this pays + # one IFFT pair on G_fft³ + an O(N · S) F_calc evaluation at the final t. + S, N = G.shape + t_init_t = torch.as_tensor(t_init, dtype=real_dtype, device=device) + h_R_int = h_R.round().to(torch.int64) # (S, N, 3) + + # Bake t_init into G: + phase_init = torch.exp( + two_pi_i * torch.einsum("snd,d->sn", h_R, t_init_t).to(complex_dtype) + ) # (S, N) + G_shifted = G * phase_init # (S, N) + + # Pick G_fft so the IFFT spacing matches the requested fine grid: + # spacing = 2·radius / (grid_steps − 1), G_fft = round(1 / spacing). + # Cap at 128 to bound memory (128³ complex128 ≈ 32 MB). + desired_spacing = max(2.0 * float(radius) / max(grid_steps - 1, 1), 1e-6) + G_fft = max(grid_steps, int(round(1.0 / desired_spacing))) + G_fft = min(G_fft, 128) + half_window = max(1, int(round(float(radius) * G_fft))) + + F_obs2 = (F_obs_t * F_obs_t).to(real_dtype) + F_obs2_centered = F_obs2 - F_obs2.mean() + F_obs2_c_complex = F_obs2_centered.to(complex_dtype) + ones_complex = torch.ones(N, dtype=complex_dtype, device=device) + + W_num_flat = torch.zeros(G_fft ** 3, dtype=complex_dtype, device=device) + W_den_flat = torch.zeros(G_fft ** 3, dtype=complex_dtype, device=device) + G_stride_xy = G_fft * G_fft + for i in range(S): + Gi_conj = G_shifted[i].conj() + pair = Gi_conj.view(1, -1) * G_shifted # (S, N) + coeff_num = F_obs2_c_complex.view(1, -1) * pair + coeff_den = ones_complex.view(1, -1) * pair + dh = (h_R_int - h_R_int[i:i + 1]) % G_fft # (S, N, 3) + flat = (dh[..., 0] * G_stride_xy + + dh[..., 1] * G_fft + dh[..., 2]) # (S, N) + flat_flat = flat.reshape(-1) + W_num_flat.index_add_(0, flat_flat, coeff_num.reshape(-1)) + W_den_flat.index_add_(0, flat_flat, coeff_den.reshape(-1)) + + W_num = W_num_flat.view(G_fft, G_fft, G_fft) + W_den = W_den_flat.view(G_fft, G_fft, G_fft) + num_t = torch.fft.ifftn(W_num, dim=(0, 1, 2)).real * (G_fft ** 3) + den_t = torch.fft.ifftn(W_den, dim=(0, 1, 2)).real * (G_fft ** 3) + score = num_t / den_t.clamp(min=1e-30) + + # Roll so the (Δt = 0) cell sits in the centre of a (2·half_window+1) + # window, then look for the maximum within the radius-sphere. + score_rolled = torch.roll( + score, shifts=(half_window, half_window, half_window), dims=(0, 1, 2), + ) + w = 2 * half_window + 1 + score_window = score_rolled[:w, :w, :w] + idx_flat = int(score_window.argmax().item()) + jx = idx_flat // (w * w) + rem = idx_flat % (w * w) + jy = rem // w + jz = rem % w + Delta_t = torch.tensor( + [(jx - half_window) / G_fft, + (jy - half_window) / G_fft, + (jz - half_window) / G_fft], + dtype=real_dtype, device=device, + ) + best_t = t_init_t + Delta_t + + # Compute the analytical-scale R-factor at best_t (one t evaluation), + # which is what the caller uses to rank rotation × translation + # candidates. The local-refine grid search above optimised the + # FFT-scored Pearson proxy; analytical R is monotonically related on + # this neighbourhood so the choice of which fine-grid maximum to + # commit to is preserved. + phase_best = torch.exp( + two_pi_i * torch.einsum("snd,d->sn", h_R, best_t).to(complex_dtype) + ) # (S, N) + F_calc_best = (G * phase_best).sum(dim=0) # (N,) complex + F_c_abs = F_calc_best.abs().to(real_dtype) + num_a = (F_obs_t * F_c_abs).sum() + den_a = (F_c_abs ** 2).sum().clamp(min=1e-30) + k = num_a / den_a + best_R = float( + ((F_obs_t - k * F_c_abs).abs().sum() / F_obs_sum).item() + ) + # Unused: `n_refinement_passes`, `batch_size` kept in signature for + # back-compat with callers passing them. + _ = n_refinement_passes + _ = batch_size + + return best_t.cpu(), best_R + + +def local_rotation_translation_refine( + F_obs: torch.Tensor, + interpolator, + R_initial: torch.Tensor, + t_initial: torch.Tensor, + hkl: torch.Tensor, + spacegroup, + real_cell, + centric: torch.Tensor, + n_shells: int = 15, + rotation_grid_steps: int = 5, + rotation_radius_rad: float = 0.04, + translation_grid_steps: int = 9, + translation_radius_frac: float = 0.02, + batch_size: int = 1024, + verbose: int = 0, +) -> Tuple[torch.Tensor, torch.Tensor, float]: + """ + Joint (R, t) fine-grid refinement scored by the Sim MLRF log-likelihood + gain (LLG) — fully scale-invariant. + + Scale invariance: F_obs is shell-normalised to E-values (Wilson + statistics per resolution shell) and F_calc(R, t) is shell-normalised + per-candidate. The LLG is then a per-shell σA fit + sum of Rice + (acentric) / Woolfson (centric) log-likelihoods — none of which depend + on the absolute magnitude of either F_obs or F_calc. + + Procedure: + 1. Build a `rotation_grid_steps³` cubic grid of small rotation + perturbations in (Δα, Δβ, Δγ) around `R_initial`, parametrised as + axis-angle rotations of magnitude up to `rotation_radius_rad`. + 2. For each R candidate: + - Pre-compute the per-symmetry F_asu via `interpolator.evaluate(R, …)` + (the expensive step — one model.forward per sym op per R). + - Run an analytical inner translation grid of + `translation_grid_steps³` candidates around `t_initial`. Within + this inner loop, only phase factors change (cheap). + - Pick the inner-best t (by an analytical R-factor proxy — fast). + 3. For each (R, best-inner-t) pair, evaluate the **full** Sim MLRF LLG + (per-shell σA fit). Pick the global best by LLG. + + Returns + ------- + R_best : torch.Tensor (3, 3) + Refined rotation = R_initial @ R_perturb_best (column-vector form). + t_best : torch.Tensor (3,) + Refined fractional translation. + llg_best : float + Sim MLRF LLG at the returned (R_best, t_best). + """ + from torchref.alignment.ml_rotation import llg_for_rotation_batch + device = getattr(interpolator, "device", hkl.device) + real_dtype = torch.float64 + complex_dtype = torch.complex128 + two_pi_i = 2j * torch.pi + + F_obs_t = F_obs.detach().to(device) + if F_obs_t.is_complex(): + F_obs_t = F_obs_t.abs() + F_obs_t = F_obs_t.to(real_dtype) + F_obs_sum = F_obs_t.sum().clamp(min=1e-30) + + hkl_t = hkl.detach().to(device).to(real_dtype) + R_init = R_initial.detach().to(device).to(real_dtype) + t_init = t_initial.detach().to(device).to(real_dtype) + centric_t = centric.detach().to(device).to(torch.bool) + + # Per-shell normalisation of F_obs → E_obs (Wilson, shell-equal-count). + rec_basis = real_cell.reciprocal_basis_matrix.to(device).to(real_dtype) + s_mag = (hkl_t @ rec_basis).norm(dim=-1) + order = torch.argsort(s_mag) + shell_idx = torch.zeros_like(s_mag, dtype=torch.int64) + chunk = max(1, s_mag.numel() // max(n_shells, 1)) + for k in range(n_shells): + a = k * chunk + b = (k + 1) * chunk if k < n_shells - 1 else s_mag.numel() + shell_idx[order[a:b]] = k + shell_sigma_obs = torch.zeros(n_shells, dtype=real_dtype, device=device) + for k in range(n_shells): + m = shell_idx == k + if m.any(): + shell_sigma_obs[k] = (F_obs_t[m] ** 2).mean().clamp(min=1e-30).sqrt() + E_obs = F_obs_t / shell_sigma_obs[shell_idx] + + # Rotation perturbation grid. + # We parametrise (Δα, Δβ, Δγ) ∈ [-r, r]³ via the small-angle rotation + # R_perturb ≈ I + ω_x · Lx + ω_y · Ly + ω_z · Lz, exponentiated via + # the matrix exponential of the skew-symmetric generator. For small ω + # (≤ ~3°) Rodrigues is well-conditioned. + def _so3_exp(omega): + # omega: (3,) axis-angle + th = omega.norm() + if th.item() < 1e-12: + return torch.eye(3, dtype=real_dtype, device=device) + axis = omega / th + K = torch.tensor( + [[0.0, -axis[2].item(), axis[1].item()], + [axis[2].item(), 0.0, -axis[0].item()], + [-axis[1].item(), axis[0].item(), 0.0]], + dtype=real_dtype, device=device, + ) + return (torch.eye(3, dtype=real_dtype, device=device) + + torch.sin(th) * K + (1 - torch.cos(th)) * (K @ K)) + + coords_r = torch.linspace(-rotation_radius_rad, rotation_radius_rad, + rotation_grid_steps, dtype=real_dtype, device=device) + omega_grid = torch.stack(torch.meshgrid(coords_r, coords_r, coords_r, + indexing="ij"), dim=-1).reshape(-1, 3) + + # Inner translation grid: (Δtx, Δty, Δtz) ∈ [-rt, rt]³ around t_initial. + coords_t = torch.linspace(-translation_radius_frac, translation_radius_frac, + translation_grid_steps, dtype=real_dtype, device=device) + t_offsets = torch.stack(torch.meshgrid(coords_t, coords_t, coords_t, + indexing="ij"), dim=-1).reshape(-1, 3) + t_candidates = t_init.unsqueeze(0) + t_offsets # (T, 3) + + # For each rotation candidate, pre-compute G_i and run inner t scan. + best_llg = -float("inf") + best_R = R_init.clone() + best_t = t_init.clone() + + for r_idx, omega in enumerate(omega_grid): + R_perturb = _so3_exp(omega) + R_cand = R_init @ R_perturb + # Build G_i for this rotation candidate (the only expensive step). + G, h_R = precompute_G_for_rotation( + interpolator, R_cand, hkl, spacegroup, real_cell, device=device, + ) + + # Inner translation scan: scored by analytical R-factor for speed. + scores = torch.empty(t_candidates.shape[0], dtype=real_dtype, device=device) + for start in range(0, t_candidates.shape[0], batch_size): + stop = min(start + batch_size, t_candidates.shape[0]) + t_batch = t_candidates[start:stop] + dot = torch.einsum("ind,bd->ibn", h_R, t_batch) + phase = torch.exp(two_pi_i * dot.to(complex_dtype)) + F_calc = torch.einsum("in,ibn->bn", G, phase) + F_c_abs = F_calc.abs().to(real_dtype) + num = (F_obs_t.unsqueeze(0) * F_c_abs).sum(dim=-1) + den = (F_c_abs ** 2).sum(dim=-1).clamp(min=1e-30) + k = num / den + R_b = ((F_obs_t.unsqueeze(0) - k.unsqueeze(-1) * F_c_abs).abs() + .sum(dim=-1)) / F_obs_sum + scores[start:stop] = R_b + best_inner = int(scores.argmin().item()) + t_best_inner = t_candidates[best_inner] + + # Score (R_cand, t_best_inner) by full Sim MLRF LLG (scale-invariant + # per-shell σA fit on E-values). + dot = (h_R * t_best_inner.unsqueeze(0).unsqueeze(0)).sum(dim=-1) # (S, N) + phase = torch.exp(two_pi_i * dot.to(complex_dtype)) + F_calc = (G * phase).sum(dim=0) # (N,) + F_calc_abs = F_calc.abs().to(real_dtype).unsqueeze(0) # (1, N) + llg = llg_for_rotation_batch( + F_obs=F_obs_t, shell_idx=shell_idx, n_shells=n_shells, + E_obs=E_obs, centric=centric_t, F_calc=F_calc_abs, + )[0].item() + + if verbose > 1: + print(f" R-refine {r_idx}/{omega_grid.shape[0]}: " + f"|ω|={omega.norm().item():.4f} rad, R={scores[best_inner]:.4f}, " + f"LLG={llg:.2f}", flush=True) + + if llg > best_llg: + best_llg = llg + best_R = R_cand.clone() + best_t = t_best_inner.clone() + return best_R.cpu(), best_t.cpu(), float(best_llg) + + +def patterson_translation_function( + F_obs: torch.Tensor, + interpolator, + R_rotation: torch.Tensor, + hkl: torch.Tensor, + spacegroup, + real_cell, + grid_shape: Optional[Tuple[int, int, int]] = None, + n_peaks: int = 20, + cluster_radius: float = 0.05, +) -> Tuple[np.ndarray, np.ndarray, List[TranslationPeak]]: + """ + Crowther-Blow Patterson translation function for molecular replacement. + + Computes T(t) = Σ_h |F_obs(h)|² · |F_calc(h, t)|² on a fractional grid by + expanding |F_calc(h, t)|² over symmetry-operator pairs and inverse-FFTing + the result. The peaks of T(t) are the translations that best place the + rotated model against the observed amplitudes — unlike the bare + `fft_translation_search`, this is the standard MR translation function and + works on amplitude-only F_obs. + + For each symmetry operator (R_i, t_i) with `x_new = R_i x_old + t_i`, + a per-symmetry asymmetric-unit structure factor is computed as + F_asu_i(h) = interpolator.evaluate(R_rotation, h R_i, real_cell, + return_amplitude=False) + * exp(2πi h · t_i) + (the "h R_i" notation follows from F(h, R x + t) = exp(2πi h·t)·F(R^T h, x); + in tensor form: `hkl @ R_i`). For each ordered pair (i, j) with i ≠ j, the + contribution + |F_obs(h)|² · conj(F_asu_i(h)) · F_asu_j(h) + is scattered into a 3-D reciprocal grid at h' = h @ (R_j − R_i), then the + inverse FFT gives the translation function. Diagonal pairs (i = j) are + t-independent. + + Parameters + ---------- + F_obs : torch.Tensor, shape (N,) + Observed amplitudes. Complex inputs are coerced to |·|. + interpolator : LattmanLoveInterpolator + Provides `evaluate(R, hkl, real_cell, return_amplitude=False)` returning + complex F_calc of the P1 ASU at arbitrary HKL. + R_rotation : torch.Tensor, shape (3, 3) + Rotation that has already been applied to the model coordinates, in + the convention `xyz_new = R · xyz_old`. Passed through to the + interpolator so it evaluates F of the rotated model. + hkl : torch.Tensor, shape (N, 3) + Integer Miller indices of the observed reflections. + spacegroup : SpaceGroup + Provides `matrices` (n_ops, 3, 3, integer in fractional basis) and + `translations` (n_ops, 3, fractional). + real_cell : Cell + Real crystal cell, passed to interpolator.evaluate. + grid_shape : tuple of int, optional + Translation-function grid (Nx, Ny, Nz). Default: 4·max(|hkl|) per axis, + which covers `h @ (R_j − R_i)^T` for any standard spacegroup. + n_peaks : int, default 20 + Number of translation peaks returned. + cluster_radius : float, default 0.05 + Minimum fractional separation between returned peaks. + + Returns + ------- + correlation_map : np.ndarray, shape (Nx, Ny, Nz) + Real-valued translation function T(t). + best_translation : np.ndarray, shape (3,) + Fractional coordinates of the top peak. + peaks : list of TranslationPeak + Top-`n_peaks` peaks sorted by descending T value. + """ + device = getattr(interpolator, "device", hkl.device) + real_dtype = torch.float64 + complex_dtype = torch.complex128 + + F_obs_t = F_obs.detach().to(device) + if F_obs_t.is_complex(): + F_obs_t = F_obs_t.abs() + F_obs_t = F_obs_t.to(real_dtype) + F_obs2 = F_obs_t * F_obs_t # (N,) + + hkl_t = hkl.detach().to(device).to(real_dtype) # (N, 3) + + sym_R = spacegroup.matrices.detach().to(device).to(real_dtype) # (S, 3, 3) + sym_t = spacegroup.translations.detach().to(device).to(real_dtype) # (S, 3) + S = sym_R.shape[0] + + R_rot = R_rotation.detach().to(device).to(real_dtype) + + # Per-symmetry F_asu_i(h) = F_model_rot(h R_i) · exp(2πi h · t_i) + F_asu = torch.zeros((S, hkl_t.shape[0]), dtype=complex_dtype, device=device) + two_pi_i = 2j * torch.pi + for i in range(S): + hkl_i = hkl_t @ sym_R[i] + F_i = interpolator.evaluate(R_rot, hkl_i, real_cell, return_amplitude=False) + F_i = F_i.to(complex_dtype) + phase = torch.exp(two_pi_i * (hkl_t @ sym_t[i])).to(complex_dtype) + F_asu[i] = F_i * phase + + # Translation grid extent: bound by max |h @ (R_j − R_i)^T|. For + # crystallographic R_op (entries in {-1, 0, 1}, occasionally 2 for trigonal + # subgroups), |R_j − R_i| has entries up to 2 → 2·max|h| per axis is the + # natural extent. Use 4·max(|hkl|) for an oversampled, periodic grid. + if grid_shape is None: + max_h = hkl_t.abs().max(dim=0).values + grid_shape = tuple(int(4 * (m.item() + 1)) for m in max_h) + Nx, Ny, Nz = grid_shape + + W = torch.zeros((Nx, Ny, Nz), dtype=complex_dtype, device=device) + W_flat = W.view(-1) + # Cross-pair accumulation (skip i == j: t-independent, only shifts DC). + for i in range(S): + Fi_conj = torch.conj(F_asu[i]) + for j in range(S): + if j == i: + continue + diff_R = sym_R[j] - sym_R[i] + h_diff = (hkl_t @ diff_R).round().to(torch.int64) + ix = h_diff[:, 0] % Nx + iy = h_diff[:, 1] % Ny + iz = h_diff[:, 2] % Nz + flat_idx = ix * (Ny * Nz) + iy * Nz + iz + weight = F_obs2 * Fi_conj * F_asu[j] + W_flat.index_add_(0, flat_idx, weight) + + TF_complex = torch.fft.ifftn(W, dim=(0, 1, 2)) + TF = TF_complex.real + + TF_np = TF.detach().cpu().numpy().astype(np.float32) + peaks = find_translation_peaks(TF_np, n_peaks=n_peaks, cluster_radius=cluster_radius) + best = peaks[0].translation if peaks else np.zeros(3) + return TF_np, best, peaks + + def apply_translation_to_fcalc_torch( F_calc: torch.Tensor, hkl: torch.Tensor, diff --git a/torchref/alignment/wigner.py b/torchref/alignment/wigner.py new file mode 100644 index 00000000..8f7e5e2b --- /dev/null +++ b/torchref/alignment/wigner.py @@ -0,0 +1,371 @@ +""" +Pure-PyTorch Wigner small-d and Wigner-D evaluation for the alignment module. + +Conventions (locked, asserted by tests/unit/alignment/test_wigner.py): + + D^l_{m,n}(α, β, γ) = e^{-i m α} · d^l_{m,n}(β) · e^{-i n γ} (Edmonds) + +with the Euler angles paired to the rotation matrix used by +`torchref.alignment.transform.rotation_matrix_from_euler` — i.e. ZYZ. + +Small-d uses the direct sum formula (Edmonds 4.1.23) with log-factorials so the +recurrence never forms `(2l)!` explicitly: + + d^l_{m,n}(β) = Σ_k (-1)^k · √[(l+m)!(l-m)!(l+n)!(l-n)!] + / [(l+m-k)! · k! · (l-n-k)! · (k+n-m)!] + · cos(β/2)^(2l+m-n-2k) · sin(β/2)^(2k+n-m) + +k runs over the integers that keep every factorial non-negative: + max(0, m-n) ≤ k ≤ min(l+m, l-n). +""" + +from __future__ import annotations + +from typing import Optional, Tuple + +import torch + + +def _log_factorial(n: torch.Tensor) -> torch.Tensor: + """log(n!) for non-negative integer tensor.""" + return torch.lgamma(n.to(torch.float64) + 1.0) + + +def _build_half_angle_pow_tables( + beta: torch.Tensor, max_exp: int +) -> Tuple[torch.Tensor, torch.Tensor]: + """Precompute `cos(β/2)^j` and `sin(β/2)^j` for j ∈ [0, max_exp]. + + Why: `torch.pow(tensor, tensor)` takes the slow `exp(log(x) * y)` path + even for integer exponents. The Wigner-d sum gathers cos/sin to powers + in {0, 1, …, 2l} for every (m, n, k) cell; replacing `cos_h ** k` with + a `cumprod`-built table + fancy index is ~10× faster on CPU for L≥16 + and dominates the `small_d_packed` cost when sharing the table across + the L iterations. + + Returns tensors of shape `(*beta.shape, max_exp + 1)`, float64. + """ + beta64 = beta.to(torch.float64) + half = 0.5 * beta64 + cos_h = torch.cos(half).unsqueeze(-1) # (*beta, 1) + sin_h = torch.sin(half).unsqueeze(-1) + ones = torch.ones_like(cos_h) + if max_exp < 1: + return ones, ones + base_cos = cos_h.expand(*beta.shape, max_exp) + base_sin = sin_h.expand(*beta.shape, max_exp) + cos_seq = torch.cat([ones, base_cos], dim=-1) + sin_seq = torch.cat([ones, base_sin], dim=-1) + return torch.cumprod(cos_seq, dim=-1), torch.cumprod(sin_seq, dim=-1) + + +def small_d_block( + l: int, + beta: torch.Tensor, + cos_pow_table: Optional[torch.Tensor] = None, + sin_pow_table: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """ + Evaluate d^l_{m,n}(β) for fixed l, all m,n ∈ [-l, l], batched over β. + + Parameters + ---------- + l : int + Wigner degree. + beta : torch.Tensor (real) + Euler β angle(s), arbitrary shape. Values in [0, π]. + cos_pow_table, sin_pow_table : torch.Tensor, optional + Precomputed tables with `cos_pow_table[..., j] = cos(β/2)^j` and + likewise for sin, shape `(*beta.shape, max_exp+1)` with + `max_exp >= 2*l`. If omitted, built locally. Pass them when looping + over l with shared β (see `small_d_packed`) to avoid repeated + `pow`-via-`exp(log·y)` evaluations. + + Returns + ------- + d : torch.Tensor (real, float64 internally, cast to beta.dtype on return) + Shape (..., 2l+1, 2l+1). `d[..., m+l, n+l] = d^l_{m,n}(β)`. + """ + if l == 0: + out = torch.ones((*beta.shape, 1, 1), dtype=beta.dtype, device=beta.device) + return out + + device = beta.device + out_dtype = beta.dtype + + # Build (or reuse) the pow tables. Local build is cheap (~max_exp small + # ops) so we only skip it when the caller hands us one. + if cos_pow_table is None or sin_pow_table is None: + cos_pow_table, sin_pow_table = _build_half_angle_pow_tables(beta, 2 * l) + max_exp = cos_pow_table.shape[-1] - 1 + assert max_exp >= 2 * l, ( + f"pow table max_exp={max_exp} insufficient for degree l={l}" + ) + + # Precompute log factorials for arguments in [0, 2l]. + n_table = torch.arange(0, 2 * l + 1, device=device) + log_fac = _log_factorial(n_table) # (2l+1,) float64 + + size = 2 * l + 1 + # Build (m, n) index grids: m_idx = m + l, n_idx = n + l, m,n ∈ [-l, l]. + m_grid = torch.arange(-l, l + 1, dtype=torch.int64, device=device) # (size,) + n_grid = m_grid.clone() + M = m_grid.view(size, 1).expand(size, size) # (size, size) + N = n_grid.view(1, size).expand(size, size) + + # k range for each (m, n) pair. + k_lo = torch.clamp(M - N, min=0) # (size, size) + k_hi = torch.minimum(torch.full_like(M, l) + M, torch.full_like(M, l) - N) + # Universal k range across all (m,n): k in [0, 2l]. + K = torch.arange(0, 2 * l + 1, dtype=torch.int64, device=device) + # Build the validity mask for each (m, n, k): + K_mn = K.view(1, 1, -1) + mask = (K_mn >= k_lo.unsqueeze(-1)) & (K_mn <= k_hi.unsqueeze(-1)) # (size, size, 2l+1) + + # Coefficient log( (l+m)!(l-m)!(l+n)!(l-n)! / [(l+m-k)! k! (l-n-k)! (k+n-m)!] )^(1/2) + # Common numerator (depends on m, n only) + L_t = torch.full_like(M, l) + log_num = 0.5 * ( + log_fac[L_t + M] + log_fac[L_t - M] + log_fac[L_t + N] + log_fac[L_t - N] + ) # (size, size) + + # Denominator term per (m, n, k) — guard out-of-range indices with mask. + # Use clamp into [0, 2l] so the index is always valid; result will be masked off. + def _safe_lf(idx): + return log_fac[idx.clamp(min=0, max=2 * l)] + + idx_a = (L_t + M).unsqueeze(-1) - K_mn # (l+m-k) + idx_b = K_mn.expand(size, size, -1) # k + idx_c = (L_t - N).unsqueeze(-1) - K_mn # (l-n-k) + idx_d = K_mn + (N - M).unsqueeze(-1) # (k+n-m) + + log_den = _safe_lf(idx_a) + _safe_lf(idx_b) + _safe_lf(idx_c) + _safe_lf(idx_d) + log_coef = log_num.unsqueeze(-1) - log_den # (size, size, 2l+1) + + coef = torch.exp(log_coef) + sign = torch.where((K_mn % 2 == 0), torch.ones_like(coef), -torch.ones_like(coef)) + # Zero out invalid k entries + coef = torch.where(mask, sign * coef, torch.zeros_like(coef)) + + # Per-k exponents. In the *valid* (mask=True) region these lie in [0, 2l]; + # invalid entries can fall outside, so we clamp to [0, max_exp] and rely + # on coef=0 to nuke their contribution. + exp_cos = 2 * l + M.unsqueeze(-1) - N.unsqueeze(-1) - 2 * K_mn # (size, size, 2l+1) + exp_sin = 2 * K_mn + N.unsqueeze(-1) - M.unsqueeze(-1) + exp_cos_safe = exp_cos.clamp(min=0, max=max_exp) + exp_sin_safe = exp_sin.clamp(min=0, max=max_exp) + + # Gather cos(β/2)^exp_cos and sin(β/2)^exp_sin from the precomputed + # tables via index_select. Faster + cleaner dispatch than fancy + # indexing (`table[..., idx]`) on CPU. Flatten the (size, size, 2l+1) + # index into 1D, gather along the last axis, then reshape back. + flat_idx_cos = exp_cos_safe.reshape(-1) + flat_idx_sin = exp_sin_safe.reshape(-1) + cos_pow = cos_pow_table.index_select(-1, flat_idx_cos).reshape( + *cos_pow_table.shape[:-1], *exp_cos_safe.shape + ) + sin_pow = sin_pow_table.index_select(-1, flat_idx_sin).reshape( + *sin_pow_table.shape[:-1], *exp_sin_safe.shape + ) + + out64 = (coef * cos_pow * sin_pow).sum(dim=-1) # (*beta, size, size) + + return out64.to(out_dtype) + + +def small_d_table(L: int, beta: torch.Tensor) -> Tuple[torch.Tensor, ...]: + """ + Compute d^l_{m,n}(β) for all l ∈ [0, L), batched over β. + + Returned as a list of tensors of shape (..., 2l+1, 2l+1) — variable in + final two dims because the small-d matrix for degree l has size 2l+1. + + Use `small_d_packed` if you want a single dense (L, 2L-1, 2L-1) tensor with + zero-padding for the off-diagonal entries beyond |m|, |n| > l. + """ + return tuple(small_d_block(l, beta) for l in range(L)) + + +def small_d_packed(L: int, beta: torch.Tensor) -> torch.Tensor: + """ + Compute the small-d matrices for all l ∈ [0, L), packed into a single + dense tensor of shape (..., L, 2L-1, 2L-1) with zero padding for entries + where |m| > l or |n| > l. + + `d_packed[..., l, L-1+m, L-1+n] = d^l_{m,n}(β)` if |m|, |n| ≤ l, else 0. + + Internally builds the cos/sin half-angle pow tables once (shared across + all L iterations) so each `small_d_block` call gathers from a table + instead of running a tensor-exponent `pow`. A fully vectorised-over-l + implementation would need a (n_beta, L, 2L-1, 2L-1, 2L-1) intermediate + that at L=32, n_beta=64 is multi-GB — the shared pow table buys most + of the speedup while staying memory-bounded. + """ + if L <= 0: + raise ValueError(f"L must be >= 1, got {L}") + out = torch.zeros((*beta.shape, L, 2 * L - 1, 2 * L - 1), + dtype=beta.dtype, device=beta.device) + max_exp = max(2 * (L - 1), 0) + cos_pow_table, sin_pow_table = _build_half_angle_pow_tables(beta, max_exp) + for l in range(L): + d_l = small_d_block( + l, beta, + cos_pow_table=cos_pow_table, + sin_pow_table=sin_pow_table, + ) # (..., 2l+1, 2l+1) + out[..., l, L - 1 - l : L - 1 + l + 1, L - 1 - l : L - 1 + l + 1] = d_l + return out + + +def wigner_D_pointwise( + alpha: torch.Tensor, + beta: torch.Tensor, + gamma: torch.Tensor, + L: int, +) -> torch.Tensor: + """ + Evaluate `D^l_{m,n}(α, β, γ) = e^{-imα} d^l_{m,n}(β) e^{-inγ}` for all + l, m, n with l < L, |m|, |n| ≤ L-1, batched over the (α, β, γ) triples. + + Returns + ------- + D : torch.Tensor (complex), shape (..., L, 2L-1, 2L-1) + """ + assert alpha.shape == beta.shape == gamma.shape + + real_dtype = beta.dtype + if real_dtype == torch.float64: + complex_dtype = torch.complex128 + elif real_dtype == torch.float32: + complex_dtype = torch.complex64 + else: + raise TypeError(f"Unsupported dtype {real_dtype}") + + device = beta.device + d = small_d_packed(L, beta) # (..., L, 2L-1, 2L-1) real + m_vals = torch.arange(-(L - 1), L, dtype=real_dtype, device=device) + n_vals = m_vals.clone() + + # phase_m[..., m_idx] = e^{-i m α}, phase_n[..., n_idx] = e^{-i n γ} + ma = alpha.unsqueeze(-1) * m_vals + ng = gamma.unsqueeze(-1) * n_vals + phase_m = torch.complex(torch.cos(-ma), torch.sin(-ma)) # (..., 2L-1) + phase_n = torch.complex(torch.cos(-ng), torch.sin(-ng)) # (..., 2L-1) + + # D[..., l, m_idx, n_idx] = d[..., l, m_idx, n_idx] · phase_m[..., m_idx] · phase_n[..., n_idx] + # Broadcast shapes: d is (..., L, 2L-1, 2L-1); need phase_m as (..., 1, 2L-1, 1) + # and phase_n as (..., 1, 1, 2L-1). + D = d.to(complex_dtype) * phase_m[..., None, :, None] * phase_n[..., None, None, :] + return D + + +def evaluate_rotation_function_grid( + xi_lmn: torch.Tensor, + L: int, + n_alpha: Optional[int] = None, + n_beta: Optional[int] = None, + n_gamma: Optional[int] = None, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """ + Evaluate `C(α, β, γ) = Σ_l Σ_{m,n} ξ_{l,m,n} D^l_{m,n}(α, β, γ)` on a + uniform Euler grid, via per-β contraction + 2D IFFT in (α, γ). + + The mathematical identity used: + C(α, β, γ) = Σ_{m,n} M_{m,n}(β) · e^{-i m α} · e^{-i n γ} + with M_{m,n}(β) := Σ_l ξ_{l,m,n} · d^l_{m,n}(β). + For each β, M_{m,n}(β) is a (2L-1)×(2L-1) matrix; the (α, γ) dependence is + a 2-D Fourier series, evaluated on a regular (n_α, n_γ) grid via IFFT. + + Parameters + ---------- + xi_lmn : torch.Tensor (complex) + Wigner coefficients, shape (L, 2L-1, 2L-1). Layout: + `xi_lmn[l, L-1+m, L-1+n] = ξ_{l,m,n}` for |m|, |n| ≤ l, else expected zero. + L : int + SH / Wigner bandlimit. + n_alpha, n_beta, n_gamma : int, optional + Grid sizes in α, β, γ. Defaults: n_alpha = n_gamma = 2L (oversampled + FFT grid), n_beta = 2L (midpoint quadrature in β). + + Returns + ------- + C : torch.Tensor (complex) + Shape (n_gamma, n_beta, n_alpha). Real-valued in exact arithmetic; the + imaginary part is returned for diagnostics. Layout: C[k_γ, k_β, k_α]. + alpha_grid, beta_grid, gamma_grid : torch.Tensor (real) + 1-D grids in radians. + """ + if n_alpha is None: + n_alpha = 2 * L + if n_gamma is None: + n_gamma = 2 * L + if n_beta is None: + n_beta = 2 * L + + device = xi_lmn.device + real_dtype = torch.float64 if xi_lmn.dtype == torch.complex128 else torch.float32 + complex_dtype = xi_lmn.dtype + + # Grids + alpha_grid = (2.0 * torch.pi / n_alpha) * torch.arange(n_alpha, dtype=real_dtype, device=device) + gamma_grid = (2.0 * torch.pi / n_gamma) * torch.arange(n_gamma, dtype=real_dtype, device=device) + # β: midpoint rule on (0, π). + beta_grid = (torch.pi * (torch.arange(n_beta, dtype=real_dtype, device=device) + 0.5) + / n_beta) + + # Build M_{m,n}(β_k) = Σ_l ξ_{l,m,n} d^l_{m,n}(β_k) for ALL β_k at once. + # `small_d_packed` accepts a batched β tensor and returns + # (n_beta, L, 2L-1, 2L-1); collapsing the previous `for kb in range(n_beta):` + # loop into a single call eliminates ~64× of Python+torch dispatch + # overhead (the dominant cost in this stage on CPU). + d_all = small_d_packed(L, beta_grid) # (n_beta, L, 2L-1, 2L-1) + d_all_c = d_all.to(complex_dtype) + M_all = (xi_lmn.unsqueeze(0) * d_all_c).sum(dim=1) # (n_beta, 2L-1, 2L-1) + + # M_{m,n}(β) gives Fourier coefficients in (-m·α, -n·γ): + # C(α, γ | β) = Σ_{m,n} M_{m,n} e^{-i m α} e^{-i n γ} + # Build a zero-padded (n_beta, n_alpha, n_gamma) coefficient grid by + # placing each M_{m,n}(β) entry at index (m mod n_alpha, n mod n_gamma); + # torch.fft.fft2 then yields C(α_k, γ_j | β) with the correct sign + # (`fft` uses exp(-2π i k n / N) which matches e^{-i m α}). + Mhat = torch.zeros( + (n_beta, n_alpha, n_gamma), dtype=complex_dtype, device=device, + ) + m_idx = torch.arange(-(L - 1), L, device=device) % n_alpha # (2L-1,) + n_idx = torch.arange(-(L - 1), L, device=device) % n_gamma # (2L-1,) + # Vectorised scatter: M_all[:, m+L-1, n+L-1] → Mhat[:, m_idx, n_idx]. + Mhat[:, m_idx.unsqueeze(-1), n_idx.unsqueeze(0)] = M_all + + # Batched 2-D FFT over (α, γ). + slice_C = torch.fft.fft2(Mhat, dim=(-2, -1)) # (n_beta, n_alpha, n_gamma) + # Re-order to (γ, β, α) layout per our convention C[k_γ, k_β, k_α]. + C = slice_C.permute(2, 0, 1).contiguous() # (n_gamma, n_beta, n_alpha) + + return C, alpha_grid, beta_grid, gamma_grid + + +def evaluate_rotation_function_pointwise( + xi_lmn: torch.Tensor, + alpha: torch.Tensor, + beta: torch.Tensor, + gamma: torch.Tensor, + L: int, +) -> torch.Tensor: + """ + Evaluate the rotation function at arbitrary Euler triples. Slow but exact + and differentiable — used for sub-voxel peak refinement and convention tests. + + Parameters + ---------- + xi_lmn : torch.Tensor (complex), shape (L, 2L-1, 2L-1) + alpha, beta, gamma : torch.Tensor (real), same shape (..., ) + + Returns + ------- + C : torch.Tensor (complex), shape (...,). In exact arithmetic real for real + input fields, but kept complex so callers can inspect drift. + """ + D = wigner_D_pointwise(alpha, beta, gamma, L) # (..., L, 2L-1, 2L-1) + # xi_lmn has shape (L, 2L-1, 2L-1); broadcasting aligns on the trailing dims. + C = (xi_lmn * D).sum(dim=(-3, -2, -1)) + return C diff --git a/torchref/model/model.py b/torchref/model/model.py index 09534d97..806b6dd8 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -245,11 +245,9 @@ def spacegroup(self, value): """ Set the space group and update the symmetry object. - Parameters - ---------- - value : gemmi.SpaceGroup or str or int - The space group to set. Can be a gemmi.SpaceGroup object, - a space group name string, or a space group number. + Accepts any input accepted by the SpaceGroup constructor — a + gemmi.SpaceGroup, an existing SpaceGroup module, a Hermann-Mauguin + string, or a space group number. """ if value is not None: self._spacegroup = SpaceGroup(value) @@ -283,6 +281,26 @@ def symmetry(self, value: Optional[SpaceGroup]): """ self._spacegroup = value + def __setattr__(self, name, value): + """ + Route Module-typed assignments to property setters when one exists. + + PyTorch's `nn.Module.__setattr__` intercepts any value that is itself + an `nn.Module` and registers it under `name` in `self._modules`, + bypassing class-level `@property` descriptors. We want our property + setters (e.g. `cell`, `spacegroup`) to take precedence — they perform + cache invalidation and other bookkeeping. + """ + if isinstance(value, nn.Module): + for klass in type(self).__mro__: + descriptor = klass.__dict__.get(name) + if descriptor is not None: + if isinstance(descriptor, property) and descriptor.fset is not None: + descriptor.fset(self, value) + return + break + super().__setattr__(name, value) + # ========================================================================= # Crystallographic matrix properties (delegated to Cell) # ========================================================================= @@ -1139,6 +1157,13 @@ def copy(self): else: model_copy.altloc_pairs = [] + # Restore the iso/aniso indexing state. The aniso_flag buffer was copied + # above; `_rebuild_sf_indices` recomputes `_iso_indices`, `_aniso_indices`, + # `_iso_covers_all`, and `_aniso_is_empty` from it. Without this, a + # freshly-copied Model fails on the first `get_iso()` / `get_aniso()`. + if hasattr(self, "aniso_flag"): + model_copy._rebuild_sf_indices() + if self.verbose > 0: print(f"✓ Model copied successfully ({len(model_copy.pdb)} atoms)") @@ -3177,116 +3202,141 @@ def xyz_fractional(self) -> torch.Tensor: return fractional_coords def rotate( - self, rotation_matrix: torch.Tensor, center: Optional[torch.Tensor] = None + self, rotation_matrix: torch.Tensor, center: Optional[torch.Tensor] = None, ) -> "Model": """ - Apply rotation to atomic coordinates (in-place). + Return a new Model with atomic coordinates rotated by R. - Rotates all atoms around a specified center point. The rotation is - applied using the formula: xyz_new = R @ (xyz - center) + center + Column-vector convention: ``xyz_new = R · (xyz_old − center) + center``. + The anisotropic U tensor (if present) is also rotated: + ``U_new = R · U · Rᵀ`` per atom. The returned Model is a deep copy — + all other state (cell, spacegroup, ADP, occupancy, PDB metadata) is + cloned, and the input ``self`` is NOT modified. Parameters ---------- - rotation_matrix : torch.Tensor - 3x3 rotation matrix. Should be orthogonal (R^T @ R = I). - center : torch.Tensor, optional - Center of rotation with shape (3,). If None, uses the centroid - of all atomic coordinates. + rotation_matrix : torch.Tensor or array-like + 3×3 rotation matrix (orthogonal, det = +1). Coerced to a torch + tensor on the model's device with the model's dtype. + center : torch.Tensor or array-like, optional + Rotation center (3,) in Cartesian Å. Default: centroid of the + current atomic coordinates (``self.xyz().mean(dim=0)``), which + preserves the centre of mass. Returns ------- Model - Self, for method chaining. + A new, independent Model instance with rotated coordinates. Examples -------- :: - # Rotate 90 degrees around Z-axis - import math - angle = math.pi / 2 - R = torch.tensor([ - [math.cos(angle), -math.sin(angle), 0], - [math.sin(angle), math.cos(angle), 0], - [0, 0, 1] - ]) - model.rotate(R) + import math, torch + from torchref.model import Model - # Rotate around a specific point - center = torch.tensor([10.0, 20.0, 30.0]) - model.rotate(R, center=center) + model = Model().load_pdb("structure.pdb") + angle = math.pi / 2 + R = torch.tensor([[math.cos(angle), -math.sin(angle), 0.0], + [math.sin(angle), math.cos(angle), 0.0], + [0.0, 0.0, 1.0]]) + rotated = model.rotate(R) # around centroid + rotated_about_point = model.rotate(R, center=torch.tensor([10.0, 20.0, 30.0])) """ if not self.initialized: raise RuntimeError("Model must be initialized to apply rotation.") - xyz = self.xyz() - if center is None: - center = xyz.mean(dim=0) - - # Ensure tensors are on the same device - rotation_matrix = rotation_matrix.to(device=xyz.device, dtype=xyz.dtype) - center = center.to(device=xyz.device, dtype=xyz.dtype) - - # Apply rotation: xyz_new = R @ (xyz - center) + center - xyz_centered = xyz - center - xyz_rotated = xyz_centered @ rotation_matrix.T + center + # Coerce to torch tensors on the model's device/dtype. + R_t = rotation_matrix if isinstance(rotation_matrix, torch.Tensor) \ + else torch.as_tensor(rotation_matrix) + R_t = R_t.to(device=self.device, dtype=self.dtype_float) + if R_t.shape != (3, 3): + raise ValueError(f"rotation_matrix must be (3, 3); got {tuple(R_t.shape)}") - # Update coordinates in-place - self.xyz[:] = xyz_rotated - - return self + xyz_curr = self.xyz() + if center is None: + center_t = xyz_curr.mean(dim=0) + else: + center_t = center if isinstance(center, torch.Tensor) \ + else torch.as_tensor(center) + center_t = center_t.to(device=self.device, dtype=self.dtype_float) + + # Column-vector rotation: xyz_new = R · (xyz_old − center) + center + # Row-vector form for (N, 3) data: xyz_new = (xyz_old − center) @ Rᵀ + center + xyz_new = (xyz_curr - center_t) @ R_t.T + center_t + + rotated = self.copy() + rotated.xyz[:] = xyz_new + + # Rotate the anisotropic U tensor per atom if any are present. + # Layout in self.u is (N, 6) packed as (U11, U22, U33, U12, U13, U23). + if self.u is not None: + u_packed = self.u() + if u_packed.numel() > 0 and not torch.isnan(u_packed).all(): + U = torch.stack([ + torch.stack([u_packed[:, 0], u_packed[:, 3], u_packed[:, 4]], dim=-1), + torch.stack([u_packed[:, 3], u_packed[:, 1], u_packed[:, 5]], dim=-1), + torch.stack([u_packed[:, 4], u_packed[:, 5], u_packed[:, 2]], dim=-1), + ], dim=-2) # (N, 3, 3) + U_rot = R_t.unsqueeze(0) @ U @ R_t.T.unsqueeze(0) # (N, 3, 3) + u_new = torch.stack([ + U_rot[:, 0, 0], U_rot[:, 1, 1], U_rot[:, 2, 2], + U_rot[:, 0, 1], U_rot[:, 0, 2], U_rot[:, 1, 2], + ], dim=-1) # (N, 6) + rotated.u[:] = u_new + + return rotated def translate( self, translation: torch.Tensor, fractional: bool = False ) -> "Model": """ - Apply translation to atomic coordinates (in-place). + Return a new Model with atomic coordinates translated by ``translation``. - Translates all atoms by a specified vector. The translation can be - given in either Cartesian or fractional coordinates. + Translates all atoms by a vector. ``translation`` may be Cartesian + (default) or fractional. The returned Model is a deep copy — all other + state (cell, spacegroup, ADP, occupancy, PDB metadata) is cloned, and + the input ``self`` is NOT modified. Parameters ---------- - translation : torch.Tensor - Translation vector with shape (3,). - fractional : bool, optional - If True, the translation is interpreted as fractional coordinates - and converted to Cartesian before applying. Default is False - (translation is in Cartesian Angstroms). + translation : torch.Tensor or array-like + Translation vector, shape (3,). Coerced to a torch tensor on the + model's device with the model's dtype. + fractional : bool, default False + If True, ``translation`` is fractional and is converted to Cartesian + via the unit cell's fractional matrix. Returns ------- Model - Self, for method chaining. + A new, independent Model instance with translated coordinates. Examples -------- :: - # Translate by 5 Angstroms along X - model.translate(torch.tensor([5.0, 0.0, 0.0])) - - # Translate by half a unit cell along each axis - model.translate(torch.tensor([0.5, 0.5, 0.5]), fractional=True) + translated = model.translate(torch.tensor([5.0, 0.0, 0.0])) + translated_frac = model.translate(torch.tensor([0.5, 0.5, 0.5]), + fractional=True) """ if not self.initialized: raise RuntimeError("Model must be initialized to apply translation.") xyz = self.xyz() - translation = translation.to(device=xyz.device, dtype=xyz.dtype) + t = translation if isinstance(translation, torch.Tensor) \ + else torch.as_tensor(translation) + t = t.to(device=xyz.device, dtype=xyz.dtype) if fractional: - # Convert fractional to Cartesian using the fractional matrix - # fractional_matrix transforms fractional -> Cartesian - translation_cart = translation @ self.fractional_matrix + t_cart = self.cell.fractional_to_cartesian(t) else: - translation_cart = translation + t_cart = t - # Apply translation in-place - xyz_translated = xyz + translation_cart - self.xyz[:] = xyz_translated - - return self + xyz_translated = xyz + t_cart + translated = self.copy() + translated.xyz[:] = xyz_translated + return translated def get_centroid(self) -> torch.Tensor: """ diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index 1642aecf..1455827b 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -153,17 +153,13 @@ def spacegroup(self): @spacegroup.setter def spacegroup(self, value): """ - Set the space group and initialize FFT if cell is also set. + Set the space group and re-initialize the SfFFT submodule. - Parameters - ---------- - value : SpaceGroup, gemmi.SpaceGroup, str, or int - The space group to set. + Accepts any input the SpaceGroup constructor accepts — a string, + space-group number, gemmi.SpaceGroup, or SpaceGroup module. """ if value is not None: - self._spacegroup = SpaceGroup( - value, dtype=self.dtype_float, device=self.device - ) + self._spacegroup = SpaceGroup(value, dtype=self.dtype_float, device=self.device) else: self._spacegroup = None self._maybe_initialize_fft() @@ -978,11 +974,27 @@ def copy(self, detach: bool = True) -> "ModelFT": # Reset cache on the copy (don't share cached structure factors) model_copy.reset_cache() + # Restore iso/aniso indexing (Model.copy() does this, but our override + # bypasses that path; aniso_flag was copied above as a buffer so we can + # rebuild from it). + if hasattr(model_copy, "aniso_flag") and model_copy.aniso_flag is not None: + model_copy._rebuild_sf_indices() + if self.verbose > 0: print(f"✓ ModelFT copied successfully ({len(model_copy.pdb)} atoms)") return model_copy + def fit_to_data(self, data, **kwargs) -> "ModelFT": + """Align this model to observed data via MR. + + Thin delegation to + :func:`torchref.alignment.align.align_model_to_data` — see that + function for the full kwargs list and behaviour. + """ + from torchref.alignment.align import align_model_to_data + return align_model_to_data(self, data, **kwargs) + def state_dict(self, destination=None, prefix="", keep_vars=False): """ Return a dictionary containing the complete state of the ModelFT. diff --git a/torchref/model/sf_fft.py b/torchref/model/sf_fft.py index de6145b7..a9a501fb 100644 --- a/torchref/model/sf_fft.py +++ b/torchref/model/sf_fft.py @@ -170,8 +170,16 @@ def cell(self) -> Optional[Cell]: @cell.setter def cell(self, value: Cell): - """Set unit cell.""" + """ + Set the unit cell. Invalidates the grid (gridsize, real_space_grid, + voxel_size) and any cached symmetry state since both depend on the + cell. The grid is re-set up automatically if it was previously set up. + """ + had_grid = self.real_space_grid is not None self._cell = value + self._invalidate_grid_caches() + if had_grid: + self.setup_grid() @property def spacegroup(self) -> Optional[SpaceGroup]: @@ -180,13 +188,59 @@ def spacegroup(self) -> Optional[SpaceGroup]: @spacegroup.setter def spacegroup(self, value: SpaceGroupLike): - """Set space group.""" + """ + Set the space group. Invalidates `map_symmetry`, the reciprocal-space + symmetry extractor, and the late-symmetry-compatibility flag. The grid + is re-set up if it was previously set up (so the new map_symmetry is + built immediately). + """ if value is not None: - self._spacegroup = SpaceGroup( - value, dtype=self.dtype_float, device=self.device - ) + new_sg = SpaceGroup(value, dtype=self.dtype_float, device=self.device) else: - self._spacegroup = None + new_sg = None + had_grid = self.real_space_grid is not None + self._spacegroup = new_sg + self._invalidate_symmetry_caches() + if had_grid: + self.setup_grid() + + def _invalidate_symmetry_caches(self): + """Clear cached symmetry-derived state. Safe to call repeatedly.""" + self.map_symmetry = None + self._late_symmetry_compatible = None + self._sym_extractor = None + self._sym_extractor_hkl_id = None + + def _invalidate_grid_caches(self): + """Clear cached grid-derived state plus symmetry state (which depends on grid shape).""" + # The grid buffers are registered via `register_buffer`. Reassigning to + # None goes through the module's __setattr__ but for buffers (not + # Modules) it does the right thing. + self.register_buffer("gridsize", None) + self.register_buffer("real_space_grid", None) + self.register_buffer("voxel_size", None) + self._invalidate_symmetry_caches() + + def __setattr__(self, name, value): + """ + Route Module-typed assignments to property setters when one exists. + + PyTorch's `nn.Module.__setattr__` intercepts any value that is itself + an `nn.Module` and registers it under `name` in `self._modules`, + bypassing class-level `@property` descriptors. We want our property + setters (e.g. `cell`, `spacegroup`) to take precedence — they perform + cache invalidation and (for spacegroup) wrap the SpaceGroup with the + right dtype/device. + """ + if isinstance(value, nn.Module): + for klass in type(self).__mro__: + descriptor = klass.__dict__.get(name) + if descriptor is not None: + if isinstance(descriptor, property) and descriptor.fset is not None: + descriptor.fset(self, value) + return + break # found a non-property class attribute; fall through + super().__setattr__(name, value) @property def symmetry(self) -> Optional[SpaceGroup]: @@ -601,6 +655,18 @@ def compute_structure_factors( Electron density map with shape (nx, ny, nz). Note: When using late symmetry, this is the P1 map (without symmetry). """ + # Ensure the grid is set up so `_late_symmetry_compatible` reflects + # the actual `MapSymmetry` instance instead of its `None` default. + # Without this, the very first `compute_structure_factors` call on a + # freshly-constructed SfFFT short-circuits to the early-symmetry + # path (`None and X → None → falsy`) and explodes on high-sym + # large-cell grids — `MapSymmetry.forward` allocates ~ n_ops × 3 × + # grid index tensors plus the gradient-tracked density gathers, + # which on a 224³ P432 (n_ops=24) crystal blew past 35 GB on an + # A100 even when only a handful of HKLs were requested. + if self.real_space_grid is None: + self.setup_grid() + # Decide symmetry strategy: # - Late symmetry: build P1 map, apply symmetry in reciprocal space # - Early symmetry: apply symmetry to density map before FFT diff --git a/torchref/scaling/solvent.py b/torchref/scaling/solvent.py index 72dc041f..b4b3bfd4 100644 --- a/torchref/scaling/solvent.py +++ b/torchref/scaling/solvent.py @@ -413,6 +413,12 @@ def smooth_solvent_mask(self): # Convert mask to float for smoothing and ensure it's on the same device mask_float = self.solvent_mask.to(dtype=self.log_k_solvent.dtype) + # `self.device` was captured at __init__ and may not match where the + # mask actually lives (e.g. SolventModel constructed with device=cpu + # default but `get_solvent_mask` later used `self.model.device=cuda`). + # Use the mask's actual device for the kernel so conv3d doesn't + # fail with "Input type cuda, weight type cpu". + device = mask_float.device # Smooth the mask using 3D Gaussian convolution # This creates soft edges at protein-solvent boundary @@ -425,7 +431,7 @@ def smooth_solvent_mask(self): # Generate 1D Gaussian x = torch.arange( - kernel_size, dtype=self.log_k_solvent.dtype, device=self.device + kernel_size, dtype=self.log_k_solvent.dtype, device=device ) x = x - kernel_size // 2 gauss_1d = torch.exp(-(x**2) / (2 * sigma**2)) From 3f5140fabed3e56c8cee4e57ebaa609f4058a9b2 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 22 May 2026 10:10:57 +0200 Subject: [PATCH 002/250] minor performance improvements --- torchref/alignment/align.py | 10 ++++-- torchref/alignment/lattman_love.py | 58 +++++++++++++++++------------- torchref/alignment/pipeline.py | 3 +- torchref/alignment/rigid_body.py | 8 ++++- 4 files changed, 50 insertions(+), 29 deletions(-) diff --git a/torchref/alignment/align.py b/torchref/alignment/align.py index 0346f920..7ba981b4 100644 --- a/torchref/alignment/align.py +++ b/torchref/alignment/align.py @@ -160,10 +160,16 @@ def _external_rwork(model: "ModelFT", data: "ReflectionData") -> float: # crashes at refine_lbfgs. s = Scaler(model=model, data=data, nbins=20, verbose=0, device=model.xyz().device) - fc = model(data.hkl) + # Detach the model forward — the scaler only needs gradients through its + # own parameters; leaving `fc` attached to the model's autograd graph + # keeps SfFFT density-build intermediates alive after this function + # returns. + with torch.no_grad(): + fc = model(data.hkl).detach() s.initialize(fc) s.refine_lbfgs(fcalc=fc) - rw, _ = s.rfactor(fc) + with torch.no_grad(): + rw, _ = s.rfactor(fc) return rw.item() if hasattr(rw, "item") else float(rw) diff --git a/torchref/alignment/lattman_love.py b/torchref/alignment/lattman_love.py index 371e6bc3..8d7577a0 100644 --- a/torchref/alignment/lattman_love.py +++ b/torchref/alignment/lattman_love.py @@ -36,9 +36,10 @@ from torchref.model.sf_fft import SfFFT from torchref.symmetry import SpaceGroup from torchref.symmetry.cell import Cell +from torchref.utils.device_mixin import DeviceMixin -class LattmanLoveInterpolator: +class LattmanLoveInterpolator(DeviceMixin): """ Compute F_calc on a dense P1 reciprocal grid once; interpolate at arbitrary rotated reciprocal positions per query. @@ -112,30 +113,37 @@ def __init__( # search model — Phaser-style MR also approximates with isotropic ADP. # Caller is free to call evaluate after adding aniso atoms in a subclass. - # Build the dense F_calc on a P1 cubic grid. - sf = SfFFT( - cell=self.cubic_cell, - spacegroup=SpaceGroup("P 1"), - max_res=max_res_A, - radius_angstrom=radius_angstrom, - dtype_float=torch.float32, - device=device, - verbose=verbose, - ) - sf.setup_grid() - density_map = sf.build_density_map( - xyz_iso=xyz_iso, - adp_iso=adp_iso, - occ_iso=occ_iso, - A_iso=A_iso, - B_iso=B_iso, - apply_symmetry=False, # already P1 - ) - # IFFT to reciprocal space; gives a complex (Nx, Ny, Nz) tensor with - # crystallographic normalization. Layout: DC at index (0, 0, 0); negative - # HKL wraps to high indices. This matches the convention expected by - # `interpolate_structure_factor_from_grid`. - self.reciprocal_grid = ifft(density_map, self.cubic_cell.volume.item()) + # Build the dense F_calc on a P1 cubic grid. Wrapped in no_grad because + # the alignment pipeline never differentiates through this grid — and + # without no_grad the resulting `self.reciprocal_grid` carries a grad_fn + # whose autograd graph pins the SfFFT internals (real_space_grid, + # voxel_xyz, per-atom kernel ≈ 5 GB on 4BX9) across trials. + with torch.no_grad(): + sf = SfFFT( + cell=self.cubic_cell, + spacegroup=SpaceGroup("P 1"), + max_res=max_res_A, + radius_angstrom=radius_angstrom, + dtype_float=torch.float32, + device=device, + verbose=verbose, + ) + sf.setup_grid() + density_map = sf.build_density_map( + xyz_iso=xyz_iso, + adp_iso=adp_iso, + occ_iso=occ_iso, + A_iso=A_iso, + B_iso=B_iso, + apply_symmetry=False, # already P1 + ) + # IFFT to reciprocal space; gives a complex (Nx, Ny, Nz) tensor + # with crystallographic normalization. Layout: DC at (0, 0, 0); + # negative HKL wraps to high indices. Matches + # `interpolate_structure_factor_from_grid`'s expectation. + self.reciprocal_grid = ifft( + density_map, self.cubic_cell.volume.item(), + ) self.cubic_cell_volume = float(self.cubic_cell.volume.item()) self.device = device self.cubic_side = cubic_side diff --git a/torchref/alignment/pipeline.py b/torchref/alignment/pipeline.py index fd4e8a0a..33fdc300 100644 --- a/torchref/alignment/pipeline.py +++ b/torchref/alignment/pipeline.py @@ -23,6 +23,7 @@ from .translation import fft_translation_search_torch, TranslationPeak from .rigid_body import RigidBodyRefinement, RigidBodyResult from .clashscore import ClashScoreCalculator, AtomSampler +from torchref.utils.device_mixin import DeviceMixin def rotation_matrix_from_euler_zyz(alpha, beta, gamma) -> np.ndarray: @@ -183,7 +184,7 @@ class MRSolution: refined_translation: Optional[np.ndarray] = None -class MolecularReplacementPipeline: +class MolecularReplacementPipeline(DeviceMixin): """ Unified MR pipeline: Rotation -> Translation -> Rigid Body Refinement. diff --git a/torchref/alignment/rigid_body.py b/torchref/alignment/rigid_body.py index 20a5447d..42f09644 100644 --- a/torchref/alignment/rigid_body.py +++ b/torchref/alignment/rigid_body.py @@ -189,7 +189,13 @@ def __init__( # user-facing R-work. self.scaler = Scaler(model=model, data=data, nbins=20, verbose=0, device=device) - fcalc_initial = self() + # Initial scaler fit only needs grad through scaler params, not + # through the rigid-body forward. Without detaching, the SfFFT + # density-build intermediates from the initial forward stay pinned + # by the autograd graph until `rb` is freed — and on multi-trial + # runs that adds ~5 GB of GPU residue per alignment. + with torch.no_grad(): + fcalc_initial = self().detach() self.scaler.calc_initial_scale(fcalc_initial) self.scaler.setup_anisotropy_correction() self.scaler.refine_lbfgs(fcalc=fcalc_initial) From 143e8c91cec78b87777dda71fe459ab975e4187f Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 22 May 2026 10:11:29 +0200 Subject: [PATCH 003/250] minor performance improvements --- pyproject.toml | 8 -------- torchref/scaling/solvent.py | 3 ++- 2 files changed, 2 insertions(+), 9 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index bb61b8d7..729b8172 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -52,14 +52,6 @@ dev = [ "flake8>=3.9.0", ] -alignment = [ - "jax>=0.4.0", - "s2fft>=1.0.0", - "s2ball>=0.0.2", - "spherical>=1.0.0", - "quaternionic>=1.0.0", -] - forcefield = [ "torchmd-net>=2.0.0", ] diff --git a/torchref/scaling/solvent.py b/torchref/scaling/solvent.py index 8c11f726..43f8796d 100644 --- a/torchref/scaling/solvent.py +++ b/torchref/scaling/solvent.py @@ -440,7 +440,8 @@ def smooth_solvent_mask(self): kernel_size += 1 x = torch.arange( - kernel_size, dtype=self.log_k_solvent.dtype, device=device + kernel_size, dtype=self.log_k_solvent.dtype, + device=self.solvent_mask.device, ) x = x - kernel_size // 2 gauss_1d = torch.exp(-(x**2) / (2 * sigma**2)) From 6b0272fcd730c983b2533ca7671e244b6cff840e Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 28 May 2026 15:56:43 +0200 Subject: [PATCH 004/250] Stable FRF --- torchref/alignment/__init__.py | 17 +- torchref/alignment/align.py | 787 +++++++++++-- torchref/alignment/distributions.py | 80 ++ torchref/alignment/frf/__init__.py | 70 ++ torchref/alignment/frf/api.py | 435 +++++++ torchref/alignment/{ => frf}/ball_search.py | 333 +++++- torchref/alignment/frf/bessel.py | 97 ++ torchref/alignment/frf/data_mr.py | 165 +++ torchref/alignment/frf/dense_calc.py | 97 ++ torchref/alignment/frf/peak_finder.py | 154 +++ torchref/alignment/frf/phaser_frf.py | 1164 +++++++++++++++++++ torchref/alignment/frf/preprocessing.py | 516 ++++++++ torchref/alignment/frf/sitelist_ang.py | 333 ++++++ torchref/alignment/frf/spherical_y.py | 96 ++ torchref/alignment/frf/types.py | 93 ++ torchref/alignment/frf/wigner_d.py | 125 ++ torchref/alignment/lattman_love.py | 79 ++ torchref/alignment/ml_rotation.py | 563 ++++++++- torchref/alignment/pipeline.py | 43 +- torchref/alignment/rigid_body.py | 65 +- torchref/alignment/sh.py | 152 ++- torchref/alignment/translation.py | 109 ++ torchref/alignment/wigner.py | 137 ++- 23 files changed, 5535 insertions(+), 175 deletions(-) create mode 100644 torchref/alignment/frf/__init__.py create mode 100644 torchref/alignment/frf/api.py rename torchref/alignment/{ => frf}/ball_search.py (63%) create mode 100644 torchref/alignment/frf/bessel.py create mode 100644 torchref/alignment/frf/data_mr.py create mode 100644 torchref/alignment/frf/dense_calc.py create mode 100644 torchref/alignment/frf/peak_finder.py create mode 100644 torchref/alignment/frf/phaser_frf.py create mode 100644 torchref/alignment/frf/preprocessing.py create mode 100644 torchref/alignment/frf/sitelist_ang.py create mode 100644 torchref/alignment/frf/spherical_y.py create mode 100644 torchref/alignment/frf/types.py create mode 100644 torchref/alignment/frf/wigner_d.py diff --git a/torchref/alignment/__init__.py b/torchref/alignment/__init__.py index 8189f22b..f4535f19 100644 --- a/torchref/alignment/__init__.py +++ b/torchref/alignment/__init__.py @@ -62,9 +62,10 @@ ) # ============================================================================= -# Pure-PyTorch ball-harmonic rotation search (no JAX, s2fft, s2ball needed) +# Fast Rotation Function engines (consolidated in the .frf sub-package) # ============================================================================= -from .ball_search import ( +# Ball-harmonic E-value rotation search (engine="ball"). +from .frf.ball_search import ( BallHarmonicCoefficients, RotationPeak, ball_rotation_search, @@ -76,6 +77,13 @@ edmonds_euler_from_rotation_matrix, rotation_angular_distance_deg, ) +# Phaser-faithful engine — the production default rotation search. +from .frf.api import ( + FastRotationFunction, + phaser_lmax_resolution, + phaser_rotation_search, +) +from .frf.dense_calc import dense_calc_via_box from .lattman_love import LattmanLoveInterpolator from .ml_rotation import sim_mlrf_rescore, brute_ml_rotation_search from .sh import ( @@ -172,6 +180,11 @@ "rotation_matrix_from_edmonds_euler", "edmonds_euler_from_rotation_matrix", "rotation_angular_distance_deg", + # Phaser-faithful engine (production default) + "FastRotationFunction", + "phaser_rotation_search", + "phaser_lmax_resolution", + "dense_calc_via_box", "LattmanLoveInterpolator", "sim_mlrf_rescore", "brute_ml_rotation_search", diff --git a/torchref/alignment/align.py b/torchref/alignment/align.py index 7ba981b4..c90674af 100644 --- a/torchref/alignment/align.py +++ b/torchref/alignment/align.py @@ -21,30 +21,37 @@ import math import time from contextlib import contextmanager -from typing import TYPE_CHECKING +from dataclasses import dataclass +from typing import Optional, TYPE_CHECKING +import numpy as np import torch -from .ball_search import ( +from .frf.ball_search import ( RotationPeak, ball_rotation_search, edmonds_euler_from_rotation_matrix, rotation_matrix_from_edmonds_euler, ) -from .lattman_love import LattmanLoveInterpolator -from .ml_rotation import sim_mlrf_rescore +from .lattman_love import LattmanLoveInterpolator, estimate_interp_var +from .ml_rotation import compute_sigma_a_luzzati, m_letf1_rescore, sim_mlrf_rescore from .rigid_body import RigidBodyRefinement from .sh import ( apply_overall_anisotropy, assign_shells, + compute_patterson_shell_variance, equal_count_shell_edges, fit_overall_anisotropy, + get_high_order_axis, ) from .translation import ( + TranslationPeak, amplitude_translation_search, + llg_translation_rescore, local_translation_refine, precompute_G_for_rotation, ) +from .ml_rotation import fit_sigma_a_per_shell if TYPE_CHECKING: from ..io.datasets.reflection_data import ReflectionData @@ -229,65 +236,60 @@ def _rodrigues(omega: torch.Tensor) -> torch.Tensor: # --------------------------------------------------------------------------- -# Public entry point +# FRF input preparation (shared by the live pipeline and the rotation-ranking +# benchmark in tests/integration/alignment/benchmark_rotation_ranking.py) # --------------------------------------------------------------------------- -def align_model_to_data( +@dataclass +class FRFInputs: + """Container for the spherical-harmonic rotation-search inputs. + + `*_sym` arrays are symmetry-expanded across the spacegroup rotation + operators (so the Patterson SH expansion samples the full sphere). + Un-suffixed `F_obs / hkl / s_vec / s_mag / centric` are the resolution- + masked, anisotropy-corrected reflection arrays — used downstream by + the MLRF rescore, translation search and rigid-body polish. + """ + # FRF (symmetry-expanded) inputs + s_vec_for_search: torch.Tensor # (n_ops·N, 3) + s_mag_sym: torch.Tensor # (n_ops·N,) + patt_obs: torch.Tensor # (n_ops·N,) = |E_obs|² − 1 + patt_calc: torch.Tensor # (n_ops·N,) = |E_calc|² − 1 + # Per-reflection inputs (un-expanded) + F_obs: torch.Tensor # (N,) anisotropy-corrected? See `F_obs_aniso` flag + hkl: torch.Tensor # (N, 3) integer Miller indices + s_vec: torch.Tensor # (N, 3) reciprocal-space Cartesian + s_mag: torch.Tensor # (N,) Å⁻¹ + centric: torch.Tensor # (N,) bool + # Other state used downstream + ll: "LattmanLoveInterpolator" + U_aniso: torch.Tensor # (3, 3) Popov-Bourenkov U + device: torch.device + + +def _prepare_frf_inputs( model: "ModelFT", data: "ReflectionData", *, - d_min: float = 4.0, - d_max: float = 15.0, - L: int = 48, - n_shells: int = 20, - n_rotation_peaks: int = 500, - n_ml_refine: int = 500, - ll_max_res_A: float = 3.0, + d_min: float, + d_max: float, + n_shells: int, ll_padding_factor: float = 2.0, + ll_max_res_A: float = 3.0, verbose: int = 0, - auto_variance_weights: bool = True, - do_translation: bool = True, - n_translation_peaks: int = 20, - n_translation_candidates: int = 3, - translation_grid_steps: int = 16, - n_rotation_candidates: int = 15, - do_joint_refine: bool = True, - joint_refine_max_res_A: float = 4.0, - joint_refine_expected_rot_error: float = 0.1, -) -> "ModelFT": - """Run full MR alignment of ``model`` against ``data``. +) -> FRFInputs: + """Build the symmetry-expanded SH-rotation-search inputs. - Returns a new rotated+translated+refined ``ModelFT`` carrying - ``last_alignment_rotation``, ``last_alignment_translation`` and - ``last_alignment_rfactor`` provenance attributes. + Encapsulates the data prep / anisotropy / symmetry-expansion logic + previously inlined in `align_model_to_data`. The returned dataclass + feeds both the live `ball_rotation_search` call and the benchmark. - See `ModelFT.fit_to_data` for full kwarg semantics — this function is the - canonical implementation; `fit_to_data` is a thin wrapper. + `F_obs` on the returned dataclass is the *anisotropy-corrected* value + (matches what previously was the `F_obs_aniso` local variable). """ - from ..scaling import Scaler # noqa: F401 (imported by _external_rwork) - from ..symmetry import SpaceGroup - - if not model.initialized: - raise RuntimeError( - "Cannot fit an uninitialized ModelFT. Load PDB data first." - ) - - timer = _StageTimer(enabled=verbose >= 2) - - # Device propagation: align_model_to_data runs on the model's device by - # default. Data tensors (data.F, data.hkl, etc.) often arrive on CPU and - # need to be moved to match — otherwise the very first hkl-derived - # quantity (s_mag) lives on CPU while F_calc from the GPU-resident LL - # interpolator lives on GPU, and _shellbin_norm_etrick crashes with - # "Expected all tensors to be on the same device". device = model.xyz().device - # --- Prepare F_obs and shell-normalized E_obs in resolution range --- - # Do the masking on CPU (data tensors arrive there) then move the - # resolution-bounded slices to `device`. Avoids GPU-index-into-CPU - # crashes when `keep` is on a different device than `data.centric`. - timer.start("0_data_prep") F_obs = data.F.to(torch.float64).abs() hkl_all = data.hkl rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) @@ -299,9 +301,6 @@ def align_model_to_data( f"Too few reflections ({keep.sum().item()}) in [{d_min},{d_max}] Å " f"for {n_shells} shells; widen the resolution range." ) - # Index on CPU, then move slices to `device`. The model device is the - # canonical destination; downstream stages (LL.evaluate, ball_search, - # sim_mlrf_rescore) all infer device from their inputs. F_obs = F_obs[keep].to(device) hkl = hkl_all[keep].to(device) s_vec = s_vec_all[keep].to(device) @@ -312,92 +311,474 @@ def align_model_to_data( else torch.zeros_like(F_obs, dtype=torch.bool) ) - timer.stop("0_data_prep") - - # Popov-Bourenkov overall anisotropy correction (full variance fix): - # Removes the direction-dependent Wilson-falloff in F_obs that otherwise - # biases the rotation function on anisotropic and high-symmetry crystals - # (P6522, etc.). Fitted from F_obs alone — no model dependence. - timer.start("1_anisotropy_fit") aniso_edges, _ = equal_count_shell_edges(s_mag, n_shells) aniso_idx = assign_shells(s_mag, aniso_edges) U_aniso = fit_overall_anisotropy( F_obs, s_vec, aniso_idx, P=n_shells, min_count=20, ) - if verbose > 0: - print( - f"fit_to_data: overall U-aniso diag (Ų) = " - f"({U_aniso[0, 0].item():+.2f}, {U_aniso[1, 1].item():+.2f}, " - f"{U_aniso[2, 2].item():+.2f})", - flush=True, - ) + # Project U onto the spacegroup's point-group-invariant subspace + # (Phaser RefineANO.cc:116-142 via cctbx `site_symmetry.average_u_star`). + # Without this constraint a 6-component unconstrained regression can + # fit physically impossible anisotropy on high-symmetry cells — e.g. + # 3K7M (cubic) fits eigenvalues (0.8, 17, 70) Ų which then blows up + # the per-reflection exp(π²·s·U·s) multiplier and destroys the FRF. + # After projection, cubic → U = λI (1 DOF), tetragonal → diag(λ,λ,μ), + # orthorhombic → diag(λ,μ,ν), etc. + from .sh import hkl_symops_to_cartesian, symmetrize_anisotropy + _sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) + _sym_mats_cart = hkl_symops_to_cartesian(_sg_mats, rec_basis.to(device)) + U_aniso = symmetrize_anisotropy(U_aniso, _sym_mats_cart) F_obs_aniso = apply_overall_anisotropy(F_obs, s_vec, U_aniso) - timer.stop("1_anisotropy_fit") - # --- Build LL interpolator from a P1 view of the model --- - timer.start("2_ll_build") - if verbose > 0: - print( - f"fit_to_data: building Lattman-Love interpolator " - f"(box={ll_padding_factor}·diam, max_res={ll_max_res_A} Å)…", - flush=True, - ) ll = LattmanLoveInterpolator( model, padding_factor=ll_padding_factor, max_res_A=ll_max_res_A, verbose=verbose, ) - # --- Symmetry-expand reciprocal-space points for the Patterson SH expansion --- - # The MTZ stores only the spacegroup ASU (1 / n_ops of reciprocal space). - # The observed Patterson is spacegroup-invariant (|F(S_k h)| = |F(h)|), - # so the SH expansion of P_obs must sample the full sphere — otherwise - # the rotation function loses its spacegroup symmetry and the true - # orientation is no longer a global maximum. On 1AK5 (P432, 24 ops) - # the un-expanded rotation function picked maxima 5× higher than the - # value at R_true; expanding F_obs to all 24 symmetry mates makes - # C(R) = C(S_k R) by construction (and the calc-side LL evaluator - # is queried at the same expanded HKL set so the cross-correlation is - # consistent). For low-sym cells (P21, n_ops=2) this is a no-op factor; - # for P432 it makes 1AK5 / 3K7M find the right basin. - sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) # (n_ops, 3, 3) + # Symmetry-expand reciprocal-space points so the Patterson SH expansion + # samples the full sphere (spacegroup-invariant by construction). See + # the long comment in `align_model_to_data` for the rationale. + sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) n_ops_sg = int(sg_mats.shape[0]) - # h_sym[k, i, :] = S_k · hkl[i] (rotation part of the symop) hkl_sym = torch.einsum("kij,nj->kni", sg_mats, hkl.to(torch.float64)) - hkl_sym_flat = hkl_sym.reshape(-1, 3) # (n_ops·N, 3) - s_vec_sym = hkl_sym_flat @ rec_basis.to(device) # (n_ops·N, 3) + hkl_sym_flat = hkl_sym.reshape(-1, 3) + s_vec_sym = hkl_sym_flat @ rec_basis.to(device) s_mag_sym = s_vec_sym.norm(dim=-1) - # |F_obs| is replicated across symmetry mates (Patterson invariance). F_obs_aniso_sym = F_obs_aniso.unsqueeze(0).expand(n_ops_sg, -1).reshape(-1) E_obs_sym = _shellbin_norm_etrick(F_obs_aniso_sym, s_mag_sym, n_shells) patt_obs = E_obs_sym ** 2 - 1.0 - # F_calc evaluated at the SAME symmetry-expanded HKL set so the cross- - # correlation between f_obs and f_calc samples the same directions. F_calc_sym = ll.evaluate( torch.eye(3, dtype=torch.float32), hkl_sym_flat, data.cell, return_amplitude=True, ).to(torch.float64) E_calc_sym = _shellbin_norm_etrick(F_calc_sym, s_mag_sym, n_shells) patt_calc = E_calc_sym ** 2 - 1.0 - # Replace the un-expanded s_vec with the expanded one for ball_search. - s_vec_for_search = s_vec_sym - timer.stop("2_ll_build") - # --- Stage 1: fast Patterson ball-search --- - timer.start("3_ball_search") + return FRFInputs( + s_vec_for_search=s_vec_sym, + s_mag_sym=s_mag_sym, + patt_obs=patt_obs, + patt_calc=patt_calc, + F_obs=F_obs_aniso, + hkl=hkl, + s_vec=s_vec, + s_mag=s_mag, + centric=centric, + ll=ll, + U_aniso=U_aniso, + device=device, + ) + + +def _run_frf_separate_rotation( + model: "ModelFT", + data: "ReflectionData", + frf: "FRFInputs", + *, + lmax_cap: int = 48, + dense_pad: float = 2.0, + n_peaks: int = 500, + grid_sampling_deg: float = 3.0, + delta_vrms_A: float = 0.5, + verbose: int = 0, + _orbit_unroll: bool = False, + # --- Phaser model-prep knobs for the FRF calc side (default OFF) --- + apply_bulk_solvent: bool = False, + solvent_fsol: float = 0.95, + solvent_bsol: float = 300.0, + vrms_strategy: str = "fixed", + vrms_identity: float = 1.0, + apply_wilson_b: bool = False, +): + """Phaser-faithful (validated) rotation search — the production default. + + Reproduces the v19 benchmark config that solved the high-symmetry cases + (4BX9 342→4-7, 6G9X 77→1-4; see ``FRF_CONSOLIDATION.md``): + + * obs taken at the **full data resolution** — ``auto_lmax`` coarsens the SH + bandwidth to ``cap`` internally (the resolution↔bandwidth coupling that + removes the aliasing background), so we do not pre-restrict resolution; + * Popov-Bourenkov **anisotropy correction** (reuses ``frf.U_aniso``); + * obs **symmetry-unroll** to the full reciprocal sphere (critical for + high-symmetry spacegroups — the SH invariant subspace is otherwise + under-sampled); + * **dense P1-box calc** (single molecular transform, not unrolled) at the + coarsened resolution — fixes high-l SH under-determination on large models; + * French-Wilson + shell-variance weights; stable Wigner-d; all under no_grad. + + Returns the validated engine's peak list (``frf.types.RotationPeak``, whose + ``.score`` aliases ``.value`` so it is drop-in for the ball-search peaks). + """ + from .frf.api import phaser_lmax_resolution, phaser_rotation_search + from .frf.dense_calc import dense_calc_via_box + + device = frf.device + with torch.no_grad(): + rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) + hkl_all = data.hkl.to(device) + s_vec_all = hkl_all.to(torch.float64) @ rec_basis + s_mag_all = s_vec_all.norm(dim=-1) + # Full data resolution window; auto_lmax coarsens d_min to match the cap. + # d_max ≈ no low-res cutoff (matches the validated config's d_max_mimic). + d_min_eff = float(1.0 / s_mag_all.max().item()) + d_max_eff = 100.0 + keep = (s_mag_all >= 1.0 / d_max_eff) & (s_mag_all <= 1.0 / d_min_eff) + + s_obs = s_vec_all[keep] + # Anisotropy correction (reuse the tensor fitted in _prepare_frf_inputs). + F_obs = apply_overall_anisotropy( + data.F.to(torch.float64).abs().to(device)[keep], s_obs, frf.U_aniso, + ) + sigF = ( + data.F_sigma.to(torch.float64).to(device)[keep] + if getattr(data, "F_sigma", None) is not None + else None + ) + centric = ( + data.centric[keep].to(torch.bool).to(device) + if hasattr(data, "centric") + else torch.zeros_like(F_obs, dtype=torch.bool) + ) + + # Obs symmetry-unroll → full reciprocal space (each ASU reflection becomes + # n_ops entries carrying the same |F|², centric, σF — |F(Sh)|=|F(h)|). + # + # The `_orbit_unroll=True` path uses `epsilon_aware_unroll` (Phaser + # DataMR.cc:954-986's `!duplicate(isym, rhkl)` skip — keeps only unique + # orbit positions). It is OFF by default: as a standalone change it + # regressed the rebench (job 103409: 3K7M 18->189, 3GR5 47->204, 2DQ6 + # 202->324). The dedup is correct only as part of a coordinated Phaser- + # faithful preprocessing chain (ε-Wilson + V(h) + σ_A), pending. + sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) + if _orbit_unroll: + from .frf.preprocessing import epsilon_aware_unroll + hkl_keep_int = hkl_all.to(torch.long).to(device)[keep] + unrolled_hkl, asu_idx = epsilon_aware_unroll(hkl_keep_int, sg_mats) + s_obs = unrolled_hkl.to(torch.float64) @ rec_basis + F_obs = F_obs[asu_idx] + centric = centric[asu_idx] + if sigF is not None: + sigF = sigF[asu_idx] + else: + n_ops = int(sg_mats.shape[0]) + hkl_keep = hkl_all.to(torch.float64)[keep] + hkl_unroll = torch.einsum("kij,nj->kni", sg_mats, hkl_keep).reshape(-1, 3) + s_obs = hkl_unroll @ rec_basis + F_obs = F_obs.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() + centric = centric.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() + if sigF is not None: + sigF = sigF.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() + + # Dense P1-box calc on the (un-rotated) search model at the coarsened res. + model_radius_A = float( + (model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item() + ) + dmin_dense = phaser_lmax_resolution(model_radius_A, d_min_eff, lmax_cap)[1] + s_calc, F_calc = dense_calc_via_box( + model, d_max_eff, dmin_dense, pad=dense_pad, verbose=verbose > 0, + ) + s_calc = s_calc.to(device) + F_calc = F_calc.to(device) + + # Optional Wilson-B match on the dense calc (EnsemblePDB.cc:793-851). + # Bin obs and calc into the same shells (defined by obs s-distribution), + # regress log(/) vs s², apply DW `exp(-B·s²/4)` to F_calc. + if apply_wilson_b: + from .frf.preprocessing import fit_relative_wilson_b + s_obs_mag = s_obs.norm(dim=-1) + s_calc_mag = s_calc.norm(dim=-1) + B_rel = fit_relative_wilson_b( + F_obs.to(torch.float64), F_calc.to(torch.float64), + s_obs_mag.to(torch.float64), n_shells=20, + s_mag_calc=s_calc_mag.to(torch.float64), + ) + if abs(B_rel) > 1e-6: + F_calc = F_calc * torch.exp(-B_rel * (s_calc_mag * s_calc_mag) / 4.0) + if verbose > 0: + print(f" FRF Wilson-B applied: B_rel = {B_rel:+.2f} Ų", flush=True) + + # Optional Oeffner vrms (rms_estimate.cc:37) — depends on n_residues + # estimated from atom count (≈ 8 heavy atoms / residue). + delta_vrms_for_frf = delta_vrms_A + if vrms_strategy == "oeffner": + from .frf.preprocessing import oeffner_vrms + n_residues_est = max(1, int(model.xyz().shape[0] / 8)) + delta_vrms_for_frf = oeffner_vrms(n_residues_est, vrms_identity) + if verbose > 0: + print( + f" FRF Oeffner vrms = {delta_vrms_for_frf:.3f} Å " + f"(n_res≈{n_residues_est}, ident={vrms_identity})", + flush=True, + ) + elif vrms_strategy != "fixed": + raise ValueError( + f"vrms_strategy={vrms_strategy!r}; expected 'fixed' or 'oeffner'." + ) + + _arf, peaks = phaser_rotation_search( + s_obs, F_obs, centric, + s_calc, F_calc, + sg_mats, + d_min=d_min_eff, d_max=d_max_eff, n_peaks=n_peaks, + delta_vrms_A=delta_vrms_for_frf, + sigma_threshold=-5.0, + use_lerf1_intensity=True, + use_m_symmetry_filter=True, + sig_F_obs=sigF, + use_french_wilson=(sigF is not None), + use_shell_variance_weights=True, + grid_sampling_deg=grid_sampling_deg, + model_radius_A=model_radius_A, + auto_lmax=True, + lmax_cap=lmax_cap, + apply_bulk_solvent=apply_bulk_solvent, + solvent_fsol=solvent_fsol, + solvent_bsol=solvent_bsol, + ) + return peaks + + +# --------------------------------------------------------------------------- +# Public entry point +# --------------------------------------------------------------------------- + + +def align_model_to_data( + model: "ModelFT", + data: "ReflectionData", + *, + d_min: float = 4.0, + d_max: float = 15.0, + L: int = 48, + n_shells: int = 20, + n_rotation_peaks: int = 500, + n_ml_refine: int = 500, + ll_max_res_A: float = 3.0, + ll_padding_factor: float = 2.0, + verbose: int = 0, + auto_variance_weights: bool = True, + do_translation: bool = True, + n_translation_peaks: int = 20, + n_translation_candidates: int = 3, + translation_grid_steps: int = 16, + n_rotation_candidates: int = 15, + do_joint_refine: bool = True, + joint_refine_max_res_A: float = 4.0, + joint_refine_expected_rot_error: float = 0.1, + use_interp_var: bool = False, + use_llg_tf: bool = False, + refine_b: bool = False, + sigma_rot_deg: float = 0.0, + sigma_trans_ang: float = 0.0, + sigma_b: float = 0.0, + use_sigma_a_frf: bool = False, + frf_delta_vrms_A: float = 1.0, + frf_weight_combine: str = "sigma_a_only", + use_m_symmetry_filter: bool = False, + use_lerf1_intensity: bool = False, + use_fitted_delta_vrms: bool = False, + use_even_l_only: bool = False, + engine: str = "frf_separate", + frf_lmax_cap: int = 48, + frf_dense_pad: float = 2.0, + rescore_engine: str = "m_letf1", +) -> "ModelFT": + """Run full MR alignment of ``model`` against ``data``. + + Returns a new rotated+translated+refined ``ModelFT`` carrying + ``last_alignment_rotation``, ``last_alignment_translation`` and + ``last_alignment_rfactor`` provenance attributes. + + See `ModelFT.fit_to_data` for full kwarg semantics — this function is the + canonical implementation; `fit_to_data` is a thin wrapper. + """ + from ..scaling import Scaler # noqa: F401 (imported by _external_rwork) + from ..symmetry import SpaceGroup + + if not model.initialized: + raise RuntimeError( + "Cannot fit an uninitialized ModelFT. Load PDB data first." + ) + + timer = _StageTimer(enabled=verbose >= 2) + + timer.start("0_data_prep") + timer.start("1_anisotropy_fit") + timer.start("2_ll_build") + frf = _prepare_frf_inputs( + model, data, + d_min=d_min, d_max=d_max, n_shells=n_shells, + ll_padding_factor=ll_padding_factor, ll_max_res_A=ll_max_res_A, + verbose=verbose, + ) + timer.stop("0_data_prep") + timer.stop("1_anisotropy_fit") + timer.stop("2_ll_build") if verbose > 0: + U_aniso = frf.U_aniso print( - f"fit_to_data: ball-search (L={L}, P={n_shells}, " - f"n_peaks={n_rotation_peaks})…", + f"fit_to_data: overall U-aniso diag (Ų) = " + f"({U_aniso[0, 0].item():+.2f}, {U_aniso[1, 1].item():+.2f}, " + f"{U_aniso[2, 2].item():+.2f})", flush=True, ) - _, _, _, _, peaks = ball_rotation_search( - s_vec_for_search, patt_obs, s_vec_for_search, patt_calc, - L=L, P=n_shells, n_peaks=n_rotation_peaks, - refine_subvoxel=True, n_refine=min(n_rotation_peaks, 50), - sigma_threshold=-5.0, - auto_variance_weights=auto_variance_weights, - ) + print( + f"fit_to_data: built Lattman-Love interpolator " + f"(box={ll_padding_factor}·diam, max_res={ll_max_res_A} Å)", + flush=True, + ) + + # Unpack into the local names the rest of this function uses. + device = frf.device + F_obs = frf.F_obs + hkl = frf.hkl + s_vec = frf.s_vec + s_mag = frf.s_mag + centric = frf.centric + ll = frf.ll + U_aniso = frf.U_aniso + s_vec_for_search = frf.s_vec_for_search + patt_obs = frf.patt_obs + patt_calc = frf.patt_calc + + # --- Stage 1: fast Patterson ball-search --- + # F3: Fit ΔVRMS from the model's mean B-factor (runs FIRST so E3/F2 + # downstream use the fitted value). ΔVRMS² = / (8π²) converts + # Debye-Waller B to a 1-D RMS coordinate displacement. + effective_delta_vrms = frf_delta_vrms_A + if use_fitted_delta_vrms: + with torch.no_grad(): + _, adp_iso, _, _, _ = model.get_iso() + b_mean = float(adp_iso.mean().item()) if adp_iso.numel() > 0 else 0.0 + effective_delta_vrms = max( + math.sqrt(max(b_mean, 1e-6) / (8 * math.pi ** 2)), 0.1, + ) + if verbose > 0: + print( + f"fit_to_data: ΔVRMS fitted from = {b_mean:.2f} Ų → " + f"{effective_delta_vrms:.3f} Å (was {frf_delta_vrms_A} Å).", + flush=True, + ) + + # Optional Phaser-style σA pre-weighting of the SH input (E3). Off by + # default — when on, replaces `auto_variance_weights`. The two can be + # combined explicitly via `frf_weight_combine="sigma_a_x_variance"`. + rotsearch_weights: Optional[torch.Tensor] = None + rotsearch_auto_var = auto_variance_weights + if use_sigma_a_frf: + shell_edges, _ = equal_count_shell_edges(frf.s_mag_sym, n_shells) + shell_mid = 0.5 * (shell_edges[:-1] + shell_edges[1:]) # (P,) + sigma_a_shell = compute_sigma_a_luzzati( + shell_mid, delta_vrms_A=effective_delta_vrms, + ) + w = sigma_a_shell ** 2 + if frf_weight_combine == "sigma_a_x_variance": + shell_idx_sym = assign_shells(frf.s_mag_sym, shell_edges) + var_shell = compute_patterson_shell_variance( + patt_obs.to(torch.float64), shell_idx_sym, P=n_shells, + ) + w = w / var_shell.sqrt().clamp(min=1e-30) + elif frf_weight_combine != "sigma_a_only": + raise ValueError( + f"frf_weight_combine={frf_weight_combine!r}; " + "expected 'sigma_a_only' or 'sigma_a_x_variance'." + ) + # Match ball_rotation_search's internal normalisation: sum-to-P. + w = w * (n_shells / w.sum().clamp(min=1e-30)) + rotsearch_weights = w.to(patt_obs.dtype) + rotsearch_auto_var = False + if verbose > 0: + print( + f"fit_to_data: σA-weighted FRF (ΔVRMS={effective_delta_vrms}Å, " + f"combine={frf_weight_combine}, w[0]={rotsearch_weights[0]:.3f}, " + f"w[-1]={rotsearch_weights[-1]:.3f}).", + flush=True, + ) + + # F2: LERF1 likelihood intensity (Phaser DataMR.cc:947–951). Replace + # `E² − 1` on the OBSERVED side with the FRF likelihood intensity + # intensity = cweight · (E² − 1) · DFAC² + # where cweight ∈ {1, 2} for centric/acentric and DFAC is a + # per-reflection Luzzati factor proxied here as the per-reflection + # σA(s) (same Luzzati formula as Eterm but evaluated on the actual + # per-reflection s_mag, not the per-shell mean). + if use_lerf1_intensity: + n_ops_sg = int(data.spacegroup.matrices.shape[0]) + cweight_per = torch.where( + centric, torch.ones_like(F_obs), 2.0 * torch.ones_like(F_obs), + ) + cweight_sym = cweight_per.unsqueeze(0).expand(n_ops_sg, -1).reshape(-1) + dfac_sym = compute_sigma_a_luzzati( + frf.s_mag_sym, delta_vrms_A=effective_delta_vrms, + ).to(patt_obs.dtype) + patt_obs = patt_obs * cweight_sym.to(patt_obs.dtype) * (dfac_sym ** 2) + if verbose > 0: + print( + f"fit_to_data: LERF1 intensity ON (mean(cweight·DFAC²) = " + f"{(cweight_sym * dfac_sym ** 2).mean().item():.3f}).", + flush=True, + ) + + # F1: m-symmetry filter (Phaser DataMR.cc:1019 / 1117). Compute ZSYMM + # from the spacegroup. Off-by-default for backwards compat. + rotsearch_zsymm = 1 + if use_m_symmetry_filter: + sg_mats_cpu = data.spacegroup.matrices.to(torch.float64).cpu() + axis, zsymm = get_high_order_axis(sg_mats_cpu) + if axis != 2 and verbose > 0: + print( + f"fit_to_data: WARNING — highest-order axis is {axis} (x/y), " + f"but m-symmetry filter assumes z. Applying with potentially " + f"reduced effect; axis permutation not yet implemented.", + flush=True, + ) + rotsearch_zsymm = int(zsymm) + if verbose > 0: + print( + f"fit_to_data: m-symmetry filter ON, ZSYMM={rotsearch_zsymm} " + f"(axis={axis}).", + flush=True, + ) + + timer.start("3_ball_search") + if engine not in ("frf_separate", "ball"): + raise ValueError( + f"engine={engine!r}; expected 'frf_separate' (default) or 'ball'." + ) + if engine == "frf_separate": + # Validated Phaser-faithful default (dense calc + auto_lmax cap + + # obs-unroll + no_grad); solves the high-symmetry cases. The σA/LERF1/ + # m-filter ball-prep above is ignored on this path. + if verbose > 0: + print( + f"fit_to_data: frf_separate rotation search " + f"(dense calc + auto_lmax cap={frf_lmax_cap}, " + f"n_peaks={n_rotation_peaks})…", + flush=True, + ) + peaks = _run_frf_separate_rotation( + model, data, frf, + lmax_cap=frf_lmax_cap, dense_pad=frf_dense_pad, + n_peaks=n_rotation_peaks, verbose=verbose, + ) + else: # engine == "ball" — legacy ball-harmonic E-value search + if verbose > 0: + print( + f"fit_to_data: ball-search (L={L}, P={n_shells}, " + f"n_peaks={n_rotation_peaks})…", + flush=True, + ) + _, _, _, _, peaks = ball_rotation_search( + s_vec_for_search, patt_obs, s_vec_for_search, patt_calc, + L=L, P=n_shells, n_peaks=n_rotation_peaks, + refine_subvoxel=True, n_refine=min(n_rotation_peaks, 50), + sigma_threshold=-5.0, + weights=rotsearch_weights, + auto_variance_weights=rotsearch_auto_var, + zsymm=rotsearch_zsymm, + skip_odd_l=use_even_l_only, + ) timer.stop("3_ball_search") # --- Stage 2: Sim-MLRF rescore (per-shell σA fit per candidate) --- @@ -408,14 +789,50 @@ def align_model_to_data( f"{min(len(peaks), n_ml_refine)} peaks…", flush=True, ) - rescored = sim_mlrf_rescore( - peaks, F_obs, hkl, s_mag, centric, ll, data.cell, - n_shells=max(n_shells // 2, 8), - n_refine=min(len(peaks), n_ml_refine), - batch_size=50, - verbose=verbose, - auto_variance_weights=auto_variance_weights, - ) + interp_var_main: Optional[torch.Tensor] = None + if use_interp_var: + # Per-reflection interpolation variance (Phaser totvar_search analogue). + # Inflates the Rice variance budget so a noisy true peak isn't + # demoted below a noise-free wrong peak by the rescore. + rescore_n_shells = max(n_shells // 2, 8) + rescore_edges, _ = equal_count_shell_edges(s_mag, rescore_n_shells) + rescore_shell_idx = assign_shells(s_mag, rescore_edges) + interp_var_main = estimate_interp_var( + ll, hkl, data.cell, rescore_shell_idx, rescore_n_shells, + ).to(F_obs.dtype) + if verbose > 0: + print( + f"fit_to_data: interp_var enabled (mean={interp_var_main.mean().item():.3f}, " + f"max={interp_var_main.max().item():.3f}).", + flush=True, + ) + + if rescore_engine not in ("m_letf1", "sim"): + raise ValueError( + f"rescore_engine={rescore_engine!r}; expected 'm_letf1' (default) or 'sim'." + ) + if rescore_engine == "m_letf1": + # Phaser-faithful: NSYMP calc sum + V(h) budget + Rice/Woolfson logRel. + # Cross-rotation case: no fixed model, so totvar_known=0 and the + # variance budget reduces to ε(h) - σ_A²(s)·n_mol in E-space. + rescored = m_letf1_rescore( + peaks, F_obs, hkl, s_mag, centric, ll, data.cell, + data.spacegroup.matrices.to(torch.float64).to(device), + n_shells=max(n_shells // 2, 8), + n_refine=min(len(peaks), n_ml_refine), + batch_size=50, + verbose=verbose, + ) + else: # rescore_engine == "sim" — legacy Sim/Rice approximation + rescored = sim_mlrf_rescore( + peaks, F_obs, hkl, s_mag, centric, ll, data.cell, + n_shells=max(n_shells // 2, 8), + n_refine=min(len(peaks), n_ml_refine), + batch_size=50, + verbose=verbose, + auto_variance_weights=auto_variance_weights, + interp_var=interp_var_main, + ) timer.stop("4_sim_mlrf_rescore") if not rescored: raise RuntimeError("Rotation search produced no peaks.") @@ -513,6 +930,101 @@ def _candidate(k): if verbose > 0: print(" no translation peaks; skipping", flush=True) continue + + # Phase B: re-rank the cheap-correlation peaks by Rice/Woolfson LLG + # using a shared per-shell σA fitted at the top correlation peak. + # Mirrors Phaser's FTF — the correlation pre-filter is fast but its + # ranking is degraded for partial models; the LLG ranks consistently + # with the rotation rescore. + if use_llg_tf: + timer.start("6b_llg_tf_rescore") + rec_basis_keep = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) + s_mag_keep_tf = (hkl_keep.to(torch.float64) @ rec_basis_keep).norm(dim=-1) + tf_n_shells = max(n_shells // 2, 8) + tf_edges, _ = equal_count_shell_edges(s_mag_keep_tf, tf_n_shells) + tf_shell_idx = assign_shells(s_mag_keep_tf, tf_edges) + centric_keep_tf = ( + data.centric[tmask].to(torch.bool).to(device) + if hasattr(data, "centric") + else torch.zeros_like(F_obs_amp, dtype=torch.bool) + ) + + # E_obs normalised on the validity-masked set. + cnt_tf = torch.bincount( + tf_shell_idx, minlength=tf_n_shells, + ).to(torch.float64) + sum_F2 = torch.zeros(tf_n_shells, dtype=torch.float64, device=device) + sum_F2.scatter_add_(0, tf_shell_idx, F_obs_amp * F_obs_amp) + mean_F2 = (sum_F2 / cnt_tf.clamp(min=1.0)).clamp(min=1e-30) + E_obs_tf = F_obs_amp / mean_F2.sqrt().index_select(0, tf_shell_idx) + + # Compute |F_calc| at the top correlation peak's translation, use + # that to fit the per-shell σA. One-shot, no per-candidate refit. + t_top_np = t_peaks[0].translation + t_top_t = torch.as_tensor(t_top_np, dtype=torch.float64, device=device) + S_eff, N_eff = G_pre.shape + phase_top = torch.exp( + 2j * torch.pi * torch.einsum( + "ind,d->in", + h_R_pre.to(torch.float64), t_top_t, + ).to(G_pre.dtype), + ) + Fc_top = (G_pre * phase_top).sum(dim=0).abs().to(torch.float64) + # Per-shell E normalise + sum_Fc2 = torch.zeros(tf_n_shells, dtype=torch.float64, device=device) + sum_Fc2.scatter_add_(0, tf_shell_idx, Fc_top * Fc_top) + mean_Fc2 = (sum_Fc2 / cnt_tf.clamp(min=1.0)).clamp(min=1e-30) + E_calc_top = Fc_top / mean_Fc2.sqrt().index_select(0, tf_shell_idx) + sigma_a_tf = fit_sigma_a_per_shell( + E_obs_tf, E_calc_top, centric_keep_tf, + tf_shell_idx, tf_n_shells, n_grid=81, + ) + + # interp_var is only meaningful when the F_calc comes from a + # trilinear interpolator. The TF stage uses a direct-SF evaluator + # (`_DirectModelEvaluator`) with no interpolation noise, so the + # Phaser totvar_search analogue does not apply here. + interp_var_tf: Optional[torch.Tensor] = None + + t_cands = torch.as_tensor( + np.stack([p.translation for p in t_peaks]), + dtype=torch.float64, device=device, + ) + llg_tf = llg_translation_rescore( + F_obs=F_obs_amp, hkl=hkl_keep, centric=centric_keep_tf, + shell_idx=tf_shell_idx, n_shells=tf_n_shells, + G=G_pre, h_R=h_R_pre, t_candidates=t_cands, + sigma_a=sigma_a_tf, interp_var=interp_var_tf, + ) + timer.stop("6b_llg_tf_rescore") + # Re-rank t_peaks by LLG (descending). Update the score to carry + # LLG so downstream picks correctly. + llg_list = llg_tf.detach().cpu().tolist() + corr_list = [p.score for p in t_peaks] # original FFT-correlation scores + order = sorted( + range(len(t_peaks)), key=lambda i: llg_list[i], reverse=True, + ) + # Record where the correlation top-1 ended up after LLG re-rank. + corr_top1_new_rank = order.index(0) + t_peaks = [ + TranslationPeak( + translation=t_peaks[i].translation, + score=float(llg_list[i]), + sigma=float(llg_list[i]), + ) + for i in order + ] + if verbose > 0: + tt = tuple(round(float(x), 3) for x in t_peaks[0].translation.tolist()) + llg_top1 = llg_list[order[0]] + corr_at_llg_top1 = corr_list[order[0]] + print( + f" LLG-TF rescore: top t={tt} LLG={t_peaks[0].score:.2f} " + f"(corr at LLG-top1: {corr_at_llg_top1:.4f}; " + f"corr-top1 demoted to LLG-rank {corr_top1_new_rank}/{len(order)})", + flush=True, + ) + if verbose > 0: tt = tuple(round(float(x), 3) for x in t_peaks[0].translation.tolist()) print( @@ -598,6 +1110,18 @@ def _candidate(k): ] R_accumulated = torch.eye(3, dtype=torch.float64) + # Hoisted out of the per-pass loop: ll_refine is fixed across passes, + # so interp_var only needs to be estimated once. + interp_var_dense: Optional[torch.Tensor] = None + if use_interp_var: + dense_n_shells = max(n_shells // 2, 8) + dense_edges, _ = equal_count_shell_edges(s_mag_keep, dense_n_shells) + dense_shell_idx = assign_shells(s_mag_keep, dense_edges) + interp_var_dense = estimate_interp_var( + ll_refine, hkl_keep, data.cell, + dense_shell_idx, dense_n_shells, + ).to(F_obs_amp.dtype) + for pass_idx, max_perturb_rad in enumerate(radii): n_per_axis = n_per_axis_pass[pass_idx] coords_r = torch.linspace( @@ -637,13 +1161,24 @@ def _candidate(k): # n_D_grid=11 (vs 41 default): we only need relative LLG # ranking across a tight rotation neighbourhood; the σA # optimum shifts negligibly. 4× fewer Bessel evals. - rescored_refine = sim_mlrf_rescore( - cand_peaks, F_obs_amp, hkl_keep, s_mag_keep, centric_keep, - ll_refine, data.cell, - n_shells=max(n_shells // 2, 8), - n_refine=len(cand_peaks), batch_size=rescore_batch, - verbose=0, n_D_grid=11, - ) + if rescore_engine == "m_letf1": + rescored_refine = m_letf1_rescore( + cand_peaks, F_obs_amp, hkl_keep, s_mag_keep, centric_keep, + ll_refine, data.cell, + data.spacegroup.matrices.to(torch.float64).to(device), + n_shells=max(n_shells // 2, 8), + n_refine=len(cand_peaks), batch_size=rescore_batch, + verbose=0, + ) + else: + rescored_refine = sim_mlrf_rescore( + cand_peaks, F_obs_amp, hkl_keep, s_mag_keep, centric_keep, + ll_refine, data.cell, + n_shells=max(n_shells // 2, 8), + n_refine=len(cand_peaks), batch_size=rescore_batch, + verbose=0, n_D_grid=11, + interp_var=interp_var_dense, + ) timer.stop("9_dense_R_rescore") top = rescored_refine[0] best_idx = next( @@ -681,6 +1216,10 @@ def _candidate(k): max_res=joint_refine_max_res_A, device=refined.device, verbose=max(0, verbose - 1), + refine_b=refine_b, + sigma_rot_deg=sigma_rot_deg, + sigma_trans_ang=sigma_trans_ang, + sigma_b=sigma_b, ) rb_result = rb.refine() with torch.no_grad(): diff --git a/torchref/alignment/distributions.py b/torchref/alignment/distributions.py index c00336d2..c6e14980 100644 --- a/torchref/alignment/distributions.py +++ b/torchref/alignment/distributions.py @@ -315,3 +315,83 @@ def centric_pdf( Probability density for each reflection. """ return torch.exp(woolfson_log_likelihood(F_obs, F_mean, variance)) + + +# ============================================================================= +# Phaser-faithful log-likelihood normalization (m_LETF1 / RiceWoolfson.cc) +# ============================================================================= +# +# These differ from `rice_log_likelihood` / `woolfson_log_likelihood` above in +# their normalization convention: Phaser's V is twice the standard Rice variance +# for acentric, equal to it for centric. Match the Phaser source exactly so the +# m_letf1_rescore values are commensurable with Phaser's m_LETF1 LL. +# +# Source: phaser/lib/RiceWoolfson.cc:25-74. + + +def phaser_log_rel_rice( + F1: torch.Tensor, + DF2: torch.Tensor, + V: torch.Tensor, +) -> torch.Tensor: + """Phaser's ``logRelRice(F1, DF2, V)`` for acentric reflections. + + Source ``phaser/lib/RiceWoolfson.cc:25-50``:: + + logRelRice(F1, DF2, V) = log I_0(2·F1·DF2/V) − log V − (F1² + DF2²)/V + + Used by ``m_LETF1`` (DataMR.cc:1425) to score each acentric reflection's + contribution to the Rice log-likelihood at a candidate orientation. + + Parameters + ---------- + F1 : torch.Tensor + Observed Wilson-normalised amplitude ``E = F_eff / sqrt(ε·Σ_N)``. + DF2 : torch.Tensor + ``sqrt(eImove)`` — square-root of the expected moving-model intensity + ``Σ_isym σ_A²·|F_calc(R^T·S_isym·h)|²``. + V : torch.Tensor + Per-reflection variance budget from ``compute_v_budget`` (DataMR.cc:949,1411). + + Returns + ------- + torch.Tensor + Per-reflection acentric log-likelihood (same shape as inputs). + """ + V_safe = V.clamp(min=1e-30) + arg = 2.0 * F1 * DF2 / V_safe + return stable_log_bessel_i0(arg) - V_safe.log() - (F1 * F1 + DF2 * DF2) / V_safe + + +def phaser_log_rel_woolfson( + F1: torch.Tensor, + DF2: torch.Tensor, + V: torch.Tensor, +) -> torch.Tensor: + """Phaser's ``logRelWoolfson(F1, DF2, V)`` for centric reflections. + + Source ``phaser/lib/RiceWoolfson.cc:52-74``:: + + logRelWoolfson(F1, DF2, V) = log cosh(F1·DF2/V) − ½·log V + − (F1² + DF2²)/(2V) + + Used by ``m_LETF1`` (DataMR.cc:1425) for centric reflections. + + Numerically stable for large ``F1·DF2/V`` via the standard + ``log cosh(x) = |x| + log1p(exp(-2|x|)) − log 2`` reformulation. + + Parameters + ---------- + F1, DF2, V : torch.Tensor + Same meaning as ``phaser_log_rel_rice``. + + Returns + ------- + torch.Tensor + Per-reflection centric log-likelihood. + """ + V_safe = V.clamp(min=1e-30) + arg = F1 * DF2 / V_safe + abs_arg = arg.abs() + log_cosh = abs_arg + torch.log1p(torch.exp(-2.0 * abs_arg)) - math.log(2.0) + return log_cosh - 0.5 * V_safe.log() - (F1 * F1 + DF2 * DF2) / (2.0 * V_safe) diff --git a/torchref/alignment/frf/__init__.py b/torchref/alignment/frf/__init__.py new file mode 100644 index 00000000..bf45f38b --- /dev/null +++ b/torchref/alignment/frf/__init__.py @@ -0,0 +1,70 @@ +"""Fast Rotation Function engines (consolidated). + +This sub-package collects every Fast Rotation Function (FRF) implementation in +one place. Three engines coexist (see ``FRF_CONSOLIDATION.md``): + +- **``api`` (the production default)** — the Phaser-faithful engine validated in + the high-symmetry investigation: chunked Bessel-SH expansion, stable Wigner-d + (``wigner_d.small_d_stable``), resolution↔bandwidth coupling + (``phaser_lmax_resolution``, cap=48), dense P1-box calc (``dense_calc``), all + under ``no_grad``. Reached 4BX9 342→4-7, 6G9X 77→1-4. +- **``ball_search`` (``engine="ball"``)** — the original pure-torch ball-harmonic + E-value rotation search. +- **``phaser_frf`` (``legacy_phaser_rotation_search``)** — the earlier 1164-line + Phaser-mimic, superseded by ``api`` but kept for reference/benchmarking. + +Shared leaf math (``..sh``, ``..wigner``) stays in the parent ``alignment`` +package; this sub-package imports it "up". +""" +from .api import ( + FastRotationFunction, + phaser_lmax_resolution, + phaser_rotation_search, +) +from .ball_search import ( + BallHarmonicCoefficients, + RotationPeak as BallRotationPeak, + ball_rotation_search, + compute_ball_cross_correlation_coefficients, + compute_ball_harmonic_coefficients, + edmonds_euler_from_rotation_matrix, + find_rotation_peaks, + refine_peaks_subvoxel, + rotation_angular_distance_deg, + rotation_matrix_from_edmonds_euler, +) +from .dense_calc import dense_calc_via_box, model_sf_abs +from .phaser_frf import phaser_rotation_search as legacy_phaser_rotation_search +from .types import ( + AdaptiveRotationFunction, + BesselSHCoefficients, + RotationPeak, + WignerContraction, +) + +__all__ = [ + # Validated engine (production default) + "FastRotationFunction", + "phaser_rotation_search", + "phaser_lmax_resolution", + "dense_calc_via_box", + "model_sf_abs", + # Ball-harmonic engine (engine="ball") + "ball_rotation_search", + "BallHarmonicCoefficients", + "BallRotationPeak", + "compute_ball_harmonic_coefficients", + "compute_ball_cross_correlation_coefficients", + "find_rotation_peaks", + "refine_peaks_subvoxel", + "rotation_matrix_from_edmonds_euler", + "edmonds_euler_from_rotation_matrix", + "rotation_angular_distance_deg", + # Legacy Phaser-mimic + "legacy_phaser_rotation_search", + # Types (validated engine) + "AdaptiveRotationFunction", + "BesselSHCoefficients", + "RotationPeak", + "WignerContraction", +] diff --git a/torchref/alignment/frf/api.py b/torchref/alignment/frf/api.py new file mode 100644 index 00000000..6962daa2 --- /dev/null +++ b/torchref/alignment/frf/api.py @@ -0,0 +1,435 @@ +"""Top-level FastRotationFunction class + drop-in ``phaser_rotation_search``. + +Signature matches ``torchref.alignment.phaser_frf.phaser_rotation_search`` +so ``tests/integration/alignment/benchmark_phaser_frf.py`` can swap +implementations via a single ``--engine`` flag. + +Pipeline (mirrors Phaser ``run_FRF()``): + 1. Resolution mask (both sides). + 2. Wilson normalisation, optionally + French-Wilson + DFAC on obs. + 3. Build LERF1 obs intensity. + 4. Optional per-shell variance reweight on obs intensity. + 5. σA Eterm on calc intensity. + 6. Detect ZSYMM → m-symmetry filter on obs SH coefficients. + 7. Bessel-SH expand both sides. + 8. Cross-correlate on the radial axis → ξ_{lmn}. + 9. Per-β fixed-shape FFT + adaptive sample list bilinear interp → + adaptive rotation function. + 10. Greedy SO(3) NMS peak finding. +""" +from __future__ import annotations + +import math +import os +import warnings +from typing import List, Optional, Tuple + +import torch + +from .data_mr import bessel_sh_expand, cross_correlate_xi +from .peak_finder import find_rotation_peaks +from .preprocessing import ( + apply_shell_variance_weights, + build_lerf1_intensity, + compute_epsilon, + detect_zsymm, + eterm_sigma_a, + french_wilson_preprocess, + wilson_normalise, + wilson_normalise_epsilon, +) +from .sitelist_ang import evaluate_rotation_function +from .types import AdaptiveRotationFunction, RotationPeak + +__all__ = ["FastRotationFunction", "phaser_rotation_search", "phaser_lmax_resolution"] + + +def phaser_lmax_resolution( + model_radius_A: float, + d_min_data: float, + lmax_cap: int = 48, +): + """Phaser's coupling of SH bandwidth to rotation-function resolution. + + Phaser source: ``runMR_FRF.cc:407-419``:: + + sphereOuter = 2 * mean_radius + LMAX = ceil(2*pi*sphereOuter / HIRES) # HIRES = data d_min + if LMAX odd: LMAX++ + LMAX = min(LMAX, DEF_CLMN_LMAX=100) + LMAX_RESO = (LMAX capped) ? 2*pi*sphereOuter/LMAX : HIRES + + The point: including data finer than the bandwidth can represent only + adds aliasing background (the discrete Y_lm are not orthogonal over + scattered reflections), which buries the symmetry-diluted true peak on + large / high-symmetry structures. So Phaser either raises LMAX to match + the resolution, or — when LMAX hits the cap — coarsens the resolution to + ``LMAX_RESO`` and drops finer reflections (DataMR.cc:984). + + **In practice this is a per-structure high-resolution cutoff.** For real + protein search models ``LMAX_ideal = ceil(2*pi*2r/d_min)`` is ~100-170, so + it always hits ``lmax_cap`` and the function reduces to: use ``L=lmax_cap`` + and keep only data coarser than ``d_min_eff = 2*pi*(2r)/lmax_cap`` — the + finest resolution that bandwidth can faithfully represent for a molecule of + radius ``r``. Bigger molecule -> coarser cutoff. The variable-L branch only + matters for tiny models / low-res data (``LMAX_ideal < lmax_cap``). + + **lmax_cap default is 48, NOT Phaser's 100 — for a different reason than + before.** The contraction now uses ``frf_separate.wigner_d.small_d_stable`` + (J_y eigendecomposition = π/2 / SOFT basis), validated stable+correct to + l=128, so there is no longer a *numerical* ceiling (the old Edmonds-sum + ``small_d_packed`` exploded to |d|~1e11 at l>=50). BUT raising L empirically + makes high-symmetry cases WORSE in this pipeline: at L=100/4.4Å on 4BX9 the + truth went from #taller=158 (cap=48) to 39084, and σA weighting did not + suppress it. Cause: at high L the SH modes (~L² per shell) are + under-determined by the obs sampled on the sparse crystal lattice (~10⁴ + reflections), so high-l coefficients are noise. Phaser avoids this by + computing the *model* transform on a dense P1-box FFT grid; we sample the + calc at the crystal lattice. Until that is changed, lmax_cap≈48 is the sweet + spot. small_d_stable is kept regardless (correct + removes the ceiling, and + is the prerequisite for any future dense-sampling high-L work). + + Parameters + ---------- + model_radius_A : float + The search model's mean atomic radius from its centroid (Å). + ``sphereOuter = 2 * model_radius_A``. + d_min_data : float + High-resolution limit of the data (Å). + lmax_cap : int + Hard bandwidth cap. Default 100 (Phaser's DEF_CLMN_LMAX). The stable + small_d_stable Wigner-d makes higher caps safe at increasing compute cost + (~L^3); for very large assemblies one may even exceed Phaser's 100. + + Returns + ------- + (L, d_min_eff) : Tuple[int, float] + ``L`` is the frf_separate bandwidth (lmax = L-1, even), ``d_min_eff`` + is the resolution to actually use for the expansion. + """ + sphere_outer = 2.0 * float(model_radius_A) + lmax = int(math.ceil(2.0 * math.pi * sphere_outer / float(d_min_data))) + if lmax % 2 != 0: + lmax += 1 + lmax = min(lmax, int(lmax_cap)) + if lmax >= int(lmax_cap): + d_min_eff = 2.0 * math.pi * sphere_outer / lmax # coarsen to match cap + else: + d_min_eff = float(d_min_data) + if lmax >= 256: + warnings.warn( + f"phaser_lmax_resolution chose lmax={lmax} (>=256): small_d_stable " + "is stable here but the contraction cost grows ~l^3 and memory ~l^2; " + "consider a tighter lmax_cap if this is slow.", + RuntimeWarning, stacklevel=2, + ) + return lmax + 1, d_min_eff # L = lmax + 1 (our bandwidth convention) + + +def _resolution_mask( + s_vec: torch.Tensor, + extra: Tuple[torch.Tensor, ...], + d_min: Optional[float], + d_max: Optional[float], +): + smag = s_vec.norm(dim=-1) + lo = 1.0 / d_max if d_max is not None else 0.0 + hi = 1.0 / d_min if d_min is not None else float("inf") + keep = (smag >= lo) & (smag <= hi) + return s_vec[keep], tuple(e[keep] for e in extra), smag[keep] + + +class FastRotationFunction: + """Reusable obs-side preprocessor + SH expansion. + + Instantiate once per (obs reflections, sym, config) tuple, then + call ``score_model(s_calc, F_calc)`` for each candidate model. + """ + + def __init__( + self, + s_obs: torch.Tensor, + F_obs: torch.Tensor, + centric_obs: torch.Tensor, + sym_mats: Optional[torch.Tensor], + *, + L: int = 24, + d_min: Optional[float] = None, + d_max: Optional[float] = None, + delta_vrms_A: float = 1.0, + n_wilson_shells: int = 20, + bessel_h_scale: Optional[float] = None, + use_lerf1_intensity: bool = True, + use_m_symmetry_filter: bool = True, + sig_F_obs: Optional[torch.Tensor] = None, + use_french_wilson: bool = False, + use_shell_variance_weights: bool = False, + n_var_shells: int = 20, + grid_sampling_deg: float = 2.0, + hkl_obs: Optional[torch.Tensor] = None, + use_epsilon: bool = False, + model_radius_A: Optional[float] = None, + auto_lmax: bool = False, + lmax_cap: int = 48, # sweet spot: higher L under-determines SH modes on sparse lattice + ): + self.device = s_obs.device + self.real_dtype = s_obs.dtype + + # Phaser-faithful coupling of bandwidth to resolution (runMR_FRF.cc:408). + # Overrides L and d_min so the SH expansion is not flooded with data + # finer than L can represent (the high-symmetry failure mode). + self.auto_lmax = auto_lmax + if auto_lmax: + if model_radius_A is None or d_min is None: + raise ValueError( + "auto_lmax=True requires model_radius_A and d_min (data res)." + ) + L, d_min = phaser_lmax_resolution(model_radius_A, d_min, lmax_cap=lmax_cap) + + self.L = L + self.d_min = d_min + self.d_max = d_max + self.delta_vrms_A = delta_vrms_A + self.n_wilson_shells = n_wilson_shells + self.grid_sampling_deg = grid_sampling_deg + + # 1. Resolution mask on obs. hkl_obs (if given) is masked in lock-step + # so the ε(h) computation below stays aligned with F_obs. + extras = (F_obs, centric_obs) + if sig_F_obs is not None: + extras = extras + (sig_F_obs,) + if hkl_obs is not None: + extras = extras + (hkl_obs,) + s_obs, extras, smag_obs = _resolution_mask(s_obs, extras, d_min, d_max) + F_obs = extras[0] + centric_obs = extras[1] + if os.environ.get("FRF_DEBUG"): + import sys as _sys + print(f"[FRF_DEBUG] auto_lmax={auto_lmax} L={self.L} d_min={d_min} " + f"d_max={d_max} n_obs_after_mask={s_obs.shape[0]} " + f"model_radius_A={model_radius_A}", file=_sys.stderr, flush=True) + ei = 2 + if sig_F_obs is not None: + sig_F_obs = extras[ei]; ei += 1 + if hkl_obs is not None: + hkl_obs = extras[ei]; ei += 1 + + if s_obs.shape[0] < n_wilson_shells * 5: + raise ValueError( + f"Too few obs reflections ({s_obs.shape[0]}) for " + f"{n_wilson_shells} Wilson shells in [{d_min}, {d_max}] Å." + ) + + # 1b. Multiplicity ε(h). Needs integer hkl + spacegroup operators. + epsilon = None + if use_epsilon: + if hkl_obs is None or sym_mats is None: + raise ValueError("use_epsilon=True requires hkl_obs and sym_mats.") + epsilon = compute_epsilon(hkl_obs, sym_mats) + + # 2. Bessel scaling default — Phaser's lmax · d_min (DataMR.cc:1107). + if bessel_h_scale is None: + if d_min is None: + raise ValueError("bessel_h_scale must be set when d_min is None") + lmax = L - 1 + lmax_even = lmax if lmax % 2 == 0 else lmax - 1 + bessel_h_scale = float(lmax_even) * float(d_min) + self.bessel_h_scale = bessel_h_scale + + # 3. Wilson + optional FW + DFAC on obs. French-Wilson does its own + # per-shell normalisation; ε-correction only applies to the plain + # Wilson path (FW handles axial reflections via its posterior). + if use_french_wilson: + if sig_F_obs is None: + raise ValueError("use_french_wilson=True requires sig_F_obs.") + fw = french_wilson_preprocess( + F_obs, sig_F_obs, smag_obs, centric_obs, + n_wilson_shells=n_wilson_shells, + ) + eEobs = fw["eEobs"] + dfac = fw["DFAC"] + # Fold ε into eEobs² post-hoc: divide by sqrt(ε) so the effective + # intensity is I/ε (axial reflections de-weighted). + if epsilon is not None: + eEobs = eEobs / epsilon.sqrt().to(eEobs.dtype) + elif epsilon is not None: + E_obs, _ = wilson_normalise_epsilon( + F_obs, smag_obs, epsilon, n_wilson_shells, + ) + eEobs = E_obs + dfac = torch.ones_like(E_obs) + else: + E_obs, _ = wilson_normalise(F_obs, smag_obs, n_wilson_shells) + eEobs = E_obs + dfac = torch.ones_like(E_obs) + + # 4. LERF1 obs intensity. + intensity_obs = build_lerf1_intensity( + eEobs, centric_obs, dfac=dfac, + use_centric_weight=use_lerf1_intensity, + ) + + # 5. Optional shell-variance reweight. + if use_shell_variance_weights: + intensity_obs = apply_shell_variance_weights( + intensity_obs, smag_obs, n_var_shells=n_var_shells, + ) + + # 6. ZSYMM detection + m-symmetry filter on obs SH coefficients. + zsymm = detect_zsymm(sym_mats) if use_m_symmetry_filter else 1 + self._zsymm = zsymm # also reused on the calc side when enabled (score_model) + + # 7. Bessel-SH expand obs side. + self._c_obs = bessel_sh_expand( + s_obs, intensity_obs.to(self.real_dtype), + L=L, bessel_h_scale=bessel_h_scale, + zsymm=zsymm, enforce_friedel=True, + ) + + def score_model( + self, + s_calc: torch.Tensor, + F_calc: torch.Tensor, + *, + n_peaks: int = 500, + sigma_threshold: float = -5.0, + calc_m_symmetry_filter: bool = False, + apply_bulk_solvent: bool = False, + solvent_fsol: float = 0.95, + solvent_bsol: float = 300.0, + ) -> Tuple[AdaptiveRotationFunction, List[RotationPeak]]: + # Resolution mask + Wilson + Eterm on calc. + s_calc, (F_calc,), smag_calc = _resolution_mask( + s_calc, (F_calc,), self.d_min, self.d_max, + ) + E_calc, _ = wilson_normalise(F_calc, smag_calc, self.n_wilson_shells) + eterm = eterm_sigma_a(smag_calc, self.delta_vrms_A) + # Optional Babinet bulk-solvent factor: Phaser folds it into σ_A as + # `σ_A_eff = solTerm(s²) · Luzzati(s², vrms)` (EnsemblePDB.cc:96-100). + # Default OFF; flip after v25 validates. + if apply_bulk_solvent: + from .preprocessing import bulk_solvent_factor + sol = bulk_solvent_factor( + smag_calc, fsol=solvent_fsol, bsol=solvent_bsol, + ).to(eterm.dtype) + eterm = eterm * sol + intensity_calc = (eterm * eterm) * (E_calc * E_calc - 1.0) + + # Bessel-SH expand calc. The "ideal-arithmetic" math says calc-side + # m-filtering is a no-op when obs is exactly spacegroup-invariant + # (R_sym(g) = NSYMP·R_orig(g), a constant scale); but in practice obs + # invariance leaks at high l (discrete-Y_lm non-orthogonality, anisotropy + # residuals, axial-reflection counting), and calc is rich in non-invariant + # m. Projecting calc onto invariant m kills the spurious obs-leak × + # calc-non-invariant product channel. + calc_zsymm = self._zsymm if calc_m_symmetry_filter else 1 + c_calc = bessel_sh_expand( + s_calc, intensity_calc.to(self.real_dtype), + L=self.L, bessel_h_scale=self.bessel_h_scale, + zsymm=calc_zsymm, enforce_friedel=True, + ) + + # 8. Cross-correlate over the radial axis. + xi = cross_correlate_xi(self._c_obs, c_calc) + + # 9. FRF on the adaptive SO(3) grid. + arf = evaluate_rotation_function(xi, grid_sampling_deg=self.grid_sampling_deg) + + # 10. Peak finding (greedy SO(3) NMS). + peaks = find_rotation_peaks( + arf, + n_peaks=n_peaks, + sigma_threshold=sigma_threshold, + nms_radius_deg=max(2.0 * self.grid_sampling_deg, 6.0), + ) + return arf, peaks + + +def phaser_rotation_search( + s_obs: torch.Tensor, + F_obs: torch.Tensor, + centric_obs: torch.Tensor, + s_calc: torch.Tensor, + F_calc: torch.Tensor, + sym_mats: torch.Tensor, + *, + L: int = 24, + d_min: Optional[float] = None, + d_max: Optional[float] = None, + delta_vrms_A: float = 1.0, + n_wilson_shells: int = 20, + n_peaks: int = 500, + refine_subvoxel: bool = True, # accepted for signature parity; ignored + n_refine: int = 50, # ignored + sigma_threshold: float = -5.0, + bessel_h_scale: Optional[float] = None, + use_lerf1_intensity: bool = True, + use_m_symmetry_filter: bool = True, + sig_F_obs: Optional[torch.Tensor] = None, + use_french_wilson: bool = False, + use_shell_variance_weights: bool = False, + n_var_shells: int = 20, + grid_sampling_deg: float = 2.0, + hkl_obs: Optional[torch.Tensor] = None, + use_epsilon: bool = False, + model_radius_A: Optional[float] = None, + auto_lmax: bool = False, + lmax_cap: int = 48, # sweet spot: higher L under-determines SH modes on sparse lattice + # NOTE: defaults to False — the standalone v20 calc-filter regressed the rebench + # (3K7M 18->189, 3GR5 47->204, 2DQ6 202->324; job 103409). Phaser's preprocessing + # pieces (eps-Wilson + V(h) + sigma_A) are mutually load-bearing; this knob will + # be flipped back on once the coordinated Phaser-faithful preprocessing chain lands. + calc_m_symmetry_filter: bool = False, + # Phaser bulk-solvent (Babinet) folded into σ_A on the calc side + # (EnsemblePDB.cc:96-100; solTerm.h:9). Default OFF; flip after sweep. + apply_bulk_solvent: bool = False, + solvent_fsol: float = 0.95, + solvent_bsol: float = 300.0, +) -> Tuple[AdaptiveRotationFunction, List[RotationPeak]]: + """Drop-in for ``torchref.alignment.phaser_frf.phaser_rotation_search``. + + Same signature, same return-shape. Sub-voxel refinement parameters + are accepted for signature parity but currently not implemented in + the frf_separate path — the per-β fixed-shape FFT already provides + sub-voxel precision via the bilinear interpolation, and Phaser + itself does not run an extra quadratic refinement. + + Extra (non-legacy) kwargs: + hkl_obs : integer Miller indices aligned with s_obs, needed for ε(h). + use_epsilon : apply the ε(h) multiplicity correction to Wilson + normalisation (de-weights axial reflections). Requires hkl_obs. + """ + # The FRF is forward-only (peak search, no backprop). Without no_grad the + # SH-Bessel expansion, the per-l Wigner contraction recurrence, and the FFT + # accumulate an autograd graph across every loop iteration — the dominant + # memory cost (tens to >100 GB at L≈100 / dense grids), and the cause of the + # OOMs. Disable grad for the whole engine. + with torch.no_grad(): + frf = FastRotationFunction( + s_obs, F_obs, centric_obs, sym_mats, + L=L, d_min=d_min, d_max=d_max, + delta_vrms_A=delta_vrms_A, + n_wilson_shells=n_wilson_shells, + bessel_h_scale=bessel_h_scale, + use_lerf1_intensity=use_lerf1_intensity, + use_m_symmetry_filter=use_m_symmetry_filter, + sig_F_obs=sig_F_obs, + use_french_wilson=use_french_wilson, + use_shell_variance_weights=use_shell_variance_weights, + n_var_shells=n_var_shells, + grid_sampling_deg=grid_sampling_deg, + hkl_obs=hkl_obs, + use_epsilon=use_epsilon, + model_radius_A=model_radius_A, + auto_lmax=auto_lmax, + lmax_cap=lmax_cap, + ) + return frf.score_model( + s_calc, F_calc, + n_peaks=n_peaks, + sigma_threshold=sigma_threshold, + calc_m_symmetry_filter=calc_m_symmetry_filter, + apply_bulk_solvent=apply_bulk_solvent, + solvent_fsol=solvent_fsol, + solvent_bsol=solvent_bsol, + ) diff --git a/torchref/alignment/ball_search.py b/torchref/alignment/frf/ball_search.py similarity index 63% rename from torchref/alignment/ball_search.py rename to torchref/alignment/frf/ball_search.py index 6abb305c..fd58eed3 100644 --- a/torchref/alignment/ball_search.py +++ b/torchref/alignment/frf/ball_search.py @@ -35,13 +35,14 @@ import torch import torch.nn.functional as F -from .sh import ( +from ..sh import ( assign_shells, compute_patterson_shell_variance, equal_count_shell_edges, sh_expand_ball, ) -from .wigner import ( +from ..wigner import ( + AdaptiveRotationFunction, evaluate_rotation_function_grid, evaluate_rotation_function_pointwise, ) @@ -89,6 +90,8 @@ def compute_ball_harmonic_coefficients( d_min: Optional[float] = None, d_max: Optional[float] = None, enforce_friedel: bool = True, + zsymm: int = 1, + skip_odd_l: bool = False, ) -> BallHarmonicCoefficients: """ Compute the ball-harmonic expansion of a real scattered-point field. @@ -139,6 +142,7 @@ def compute_ball_harmonic_coefficients( f_plm = sh_expand_ball( s_vectors, values, shell_idx, P, L, enforce_friedel=enforce_friedel, + zsymm=zsymm, skip_odd_l=skip_odd_l, ) return BallHarmonicCoefficients( @@ -288,6 +292,321 @@ def find_rotation_peaks( ] +# ============================================================================= +# Adaptive-grid peak finding +# ============================================================================= + + +def _adaptive_global_stats(arf: AdaptiveRotationFunction) -> Tuple[float, float]: + """ + Sample-weighted global mean and std across all β-slices. + + Every voxel of every slice contributes equally — the adaptive grid is + already sample-uniform in SO(3) sense (each voxel covers a near-constant + `sin(β)/(pmax(β)·qmax(β))` solid angle), so simple mean/std is correct. + """ + flat = torch.cat([s.real.reshape(-1) for s in arf.slices]) + mean = float(flat.mean().item()) + std = float(flat.std().clamp(min=1e-30).item()) + return mean, std + + +def _slice_nms_mask(slice_C: torch.Tensor, r_alpha: int, r_gamma: int) -> torch.Tensor: + """ + 2-D non-maximum-suppression for one β-slice, circular in both α and γ. + + `slice_C` is real-valued (γ, α). Returns a boolean mask of local maxima. + """ + if r_alpha <= 0 and r_gamma <= 0: + return torch.ones_like(slice_C, dtype=torch.bool) + n_gamma, n_alpha = slice_C.shape + r_a = max(r_alpha, 0) + r_g = max(r_gamma, 0) + # Clamp radii so we never wrap past the slice's own period — at β = 0 the + # γ axis has length 1, at β = π the α axis has length 1. + r_a = min(r_a, max(n_alpha - 1, 0)) + r_g = min(r_g, max(n_gamma - 1, 0)) + if r_a == 0 and r_g == 0: + return torch.ones_like(slice_C, dtype=torch.bool) + # Circular pad in γ (axis 0) then α (axis 1). + x = slice_C + if r_g > 0: + x = torch.cat([x[-r_g:], x, x[:r_g]], dim=0) + if r_a > 0: + x = torch.cat([x[:, -r_a:], x, x[:, :r_a]], dim=1) + pooled = F.max_pool2d( + x.unsqueeze(0).unsqueeze(0), + kernel_size=(2 * r_g + 1, 2 * r_a + 1), + stride=1, padding=0, + )[0, 0] + return slice_C >= pooled + + +def _sample_slice_at( + slice_C: torch.Tensor, + alpha_grid: torch.Tensor, + gamma_grid: torch.Tensor, + alpha_q: float, + gamma_q: float, +) -> float: + """ + Nearest-neighbour sample of a β-slice at a query (α, γ) in radians. + """ + n_gamma, n_alpha = slice_C.shape + two_pi = 2.0 * math.pi + a0 = float(alpha_grid[0].item()) + g0 = float(gamma_grid[0].item()) + da = two_pi / n_alpha + dg = two_pi / n_gamma + ia = int(round((alpha_q - a0) / da)) % n_alpha + ig = int(round((gamma_q - g0) / dg)) % n_gamma + return float(slice_C[ig, ia].item()) + + +def _so3_greedy_nms( + alpha: torch.Tensor, + beta: torch.Tensor, + gamma: torch.Tensor, + values: torch.Tensor, + nms_radius_rad: float, + keep_at_most: int, +) -> torch.Tensor: + """ + Greedy SO(3) non-max suppression by angular distance between rotations. + + Walks the candidates in descending `values` order; keeps a candidate iff + its angular distance `arccos((tr(R_cand · R_kept^T) − 1) / 2)` exceeds + `nms_radius_rad` from every already-kept rotation. Returns the indices + of the kept candidates in the order they were accepted (i.e. by + descending score). + + Complexity: `O(N · K)` matrix-trace ops where `K = min(keep_at_most, N)` + — the inner loop runs entirely in PyTorch on the host device. + """ + n = alpha.shape[0] + if n == 0: + return torch.zeros(0, dtype=torch.long, device=alpha.device) + R_all = rotation_matrix_from_edmonds_euler_batch(alpha, beta, gamma) # (n, 3, 3) + order = torch.argsort(values, descending=True) + + cos_thresh = math.cos(nms_radius_rad) + kept_indices: List[int] = [] + kept_R = R_all.new_zeros((0, 3, 3)) + + for i in order.tolist(): + Ri = R_all[i] # (3, 3) + if kept_R.shape[0] > 0: + # tr(Ri · R_kept^T) over the batch of kept rotations. + tr = torch.einsum("ab,kab->k", Ri, kept_R) + # angular distance = arccos((tr − 1) / 2). Compare against + # cosine to avoid the acos call entirely. + cos_theta = ((tr - 1.0) * 0.5).clamp(-1.0, 1.0) + if (cos_theta > cos_thresh).any().item(): + continue + kept_indices.append(i) + kept_R = torch.cat([kept_R, Ri.unsqueeze(0)], dim=0) + if len(kept_indices) >= keep_at_most: + break + + return torch.tensor(kept_indices, dtype=torch.long, device=alpha.device) + + +def find_rotation_peaks_adaptive( + arf: AdaptiveRotationFunction, + n_peaks: int = 200, + sigma_threshold: float = 0.0, + cluster_radius_deg: float = 6.0, + nms_radius_deg: float = 6.0, + candidate_cap: int = 10000, +) -> List[RotationPeak]: + """ + Local-maximum peak finder for the ragged adaptive Euler grid. + + Algorithm: + 1. Per-slice 2-D NMS (circular in α and γ, radius converted from + `cluster_radius_deg` to a per-slice voxel count). + 2. Threshold + truncate to the top `candidate_cap` survivors across + all slices (bounds the cost of the next step). + 3. Greedy SO(3) non-max suppression: walk candidates in descending + score; keep iff angular distance to every already-kept rotation + is > `nms_radius_deg`. This subsumes the role the dense path's + `max_pool3d` played and is geometry-aware — duplicate "polar" + peaks differing only in α+γ (β≈0) or α−γ (β≈π) have angular + distance ≈ 0 and are correctly collapsed. + 4. Truncate to top `n_peaks` and attach Z-score using the + globally-pooled mean / std. + """ + n_beta = arf.betas.shape[0] + mean, std = _adaptive_global_stats(arf) + threshold = mean + sigma_threshold * std + + # Per-slice local maxima + threshold. Collect onto torch tensors directly + # to avoid per-candidate .item() round-trips. + alpha_chunks: List[torch.Tensor] = [] + beta_chunks: List[torch.Tensor] = [] + gamma_chunks: List[torch.Tensor] = [] + value_chunks: List[torch.Tensor] = [] + for k in range(n_beta): + slice_C = arf.slices[k].real + n_gamma_k, n_alpha_k = slice_C.shape + r_a = max(1, int(round(cluster_radius_deg * n_alpha_k / 360.0))) + r_g = max(1, int(round(cluster_radius_deg * n_gamma_k / 360.0))) + + is_max = _slice_nms_mask(slice_C, r_alpha=r_a, r_gamma=r_g) + is_max = is_max & (slice_C >= threshold) + if not bool(is_max.any().item()): + continue + ig_idx, ia_idx = torch.nonzero(is_max, as_tuple=True) + vals = slice_C[ig_idx, ia_idx] + a_grid_k = arf.alpha_grids[k] + g_grid_k = arf.gamma_grids[k] + alpha_chunks.append(a_grid_k.index_select(0, ia_idx)) + gamma_chunks.append(g_grid_k.index_select(0, ig_idx)) + beta_chunks.append(arf.betas[k].expand_as(vals)) + value_chunks.append(vals) + + if not value_chunks: + return [] + + alpha_all = torch.cat(alpha_chunks) + beta_all = torch.cat(beta_chunks) + gamma_all = torch.cat(gamma_chunks) + value_all = torch.cat(value_chunks) + + # Cap the candidate set by raw score before the O(N·K) SO(3) NMS. + if value_all.shape[0] > candidate_cap: + top_v, top_i = torch.topk(value_all, candidate_cap) + alpha_all = alpha_all.index_select(0, top_i) + beta_all = beta_all.index_select(0, top_i) + gamma_all = gamma_all.index_select(0, top_i) + value_all = top_v + + # Greedy SO(3) NMS. + keep_idx = _so3_greedy_nms( + alpha_all, beta_all, gamma_all, value_all, + nms_radius_rad=math.radians(nms_radius_deg), + keep_at_most=int(n_peaks), + ) + + if keep_idx.numel() == 0: + return [] + + alpha_keep = alpha_all.index_select(0, keep_idx).tolist() + beta_keep = beta_all.index_select(0, keep_idx).tolist() + gamma_keep = gamma_all.index_select(0, keep_idx).tolist() + value_keep = value_all.index_select(0, keep_idx).tolist() + + return [ + RotationPeak( + alpha=a, beta=b, gamma=g, score=v, + sigma=(v - mean) / max(std, 1e-30), + ) + for (a, b, g, v) in zip(alpha_keep, beta_keep, gamma_keep, value_keep) + ] + + +def refine_peaks_subvoxel_adaptive( + peaks: List[RotationPeak], + arf: AdaptiveRotationFunction, +) -> List[RotationPeak]: + """ + Sub-voxel parabolic refinement for adaptive-grid peaks. + + For each peak at (α, β_k, γ): + * α and γ refinement uses the per-slice spacing `2π / pmax_k`, + `2π / qmax_k`. Skipped at slices where the relevant axis has + length 1 (β ≈ 0 → no γ refinement; β ≈ π → no α refinement). + * β refinement fits a parabola through the values at the same + (α, γ) sampled into slices k-1 and k+1 by nearest-neighbour. + Boundary slices use db = 0. + """ + if not peaks: + return list(peaks) + + n_beta = arf.betas.shape[0] + beta_grid = arf.betas + # Beta is uniform in our construction, so a single dβ is correct. + if n_beta > 1: + dbeta = float((beta_grid[1] - beta_grid[0]).item()) + else: + dbeta = 0.0 + + out: List[RotationPeak] = [] + for p in peaks: + # Locate the host β-slice. + ib = int(round((p.beta - float(beta_grid[0].item())) / dbeta)) if dbeta > 0 else 0 + ib = max(0, min(ib, n_beta - 1)) + slice_C = arf.slices[ib].real + a_grid = arf.alpha_grids[ib] + g_grid = arf.gamma_grids[ib] + n_gamma_k, n_alpha_k = slice_C.shape + da = 2.0 * math.pi / n_alpha_k + dg = 2.0 * math.pi / n_gamma_k + + # Integer voxel indices. + a0 = float(a_grid[0].item()) + g0 = float(g_grid[0].item()) + ia = int(round((p.alpha - a0) / da)) % n_alpha_k + ig = int(round((p.gamma - g0) / dg)) % n_gamma_k + + y0 = float(slice_C[ig, ia].item()) + + # α refinement. + if n_alpha_k >= 3: + am = (ia - 1) % n_alpha_k + ap = (ia + 1) % n_alpha_k + yam = float(slice_C[ig, am].item()) + yap = float(slice_C[ig, ap].item()) + denom = yam - 2.0 * y0 + yap + if abs(denom) > 1e-30: + d_a = max(-1.0, min(1.0, 0.5 * (yam - yap) / denom)) + else: + d_a = 0.0 + alpha_new = (p.alpha + d_a * da) % (2.0 * math.pi) + else: + alpha_new = p.alpha + + # γ refinement. + if n_gamma_k >= 3: + gm = (ig - 1) % n_gamma_k + gp = (ig + 1) % n_gamma_k + ygm = float(slice_C[gm, ia].item()) + ygp = float(slice_C[gp, ia].item()) + denom = ygm - 2.0 * y0 + ygp + if abs(denom) > 1e-30: + d_g = max(-1.0, min(1.0, 0.5 * (ygm - ygp) / denom)) + else: + d_g = 0.0 + gamma_new = (p.gamma + d_g * dg) % (2.0 * math.pi) + else: + gamma_new = p.gamma + + # β refinement: sample neighbouring slices at (α, γ). + if 0 < ib < n_beta - 1 and dbeta > 0: + v_prev = _sample_slice_at( + arf.slices[ib - 1].real, arf.alpha_grids[ib - 1], + arf.gamma_grids[ib - 1], p.alpha, p.gamma, + ) + v_next = _sample_slice_at( + arf.slices[ib + 1].real, arf.alpha_grids[ib + 1], + arf.gamma_grids[ib + 1], p.alpha, p.gamma, + ) + denom = v_prev - 2.0 * y0 + v_next + if abs(denom) > 1e-30: + d_b = max(-1.0, min(1.0, 0.5 * (v_prev - v_next) / denom)) + else: + d_b = 0.0 + beta_new = max(0.0, min(math.pi, p.beta + d_b * dbeta)) + else: + beta_new = p.beta + + out.append(RotationPeak( + alpha=alpha_new, beta=beta_new, gamma=gamma_new, + score=p.score, sigma=p.sigma, + )) + return out + + # ============================================================================= # Sub-voxel refinement (pointwise gradient ascent on C(R) using autograd) # ============================================================================= @@ -417,6 +736,8 @@ def ball_rotation_search( sigma_threshold: float = 1.0, weights: Optional[torch.Tensor] = None, auto_variance_weights: bool = True, + zsymm: int = 1, + skip_odd_l: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, List[RotationPeak]]: """ End-to-end Patterson rotation search. @@ -460,13 +781,21 @@ def ball_rotation_search( Grid coordinates. peaks : list of RotationPeak (sorted by descending score) """ + # NOTE: `zsymm` is applied only to the OBSERVED side. The observed + # Patterson is spacegroup-invariant by physics, so SH coefficients with + # m not divisible by ZSYMM are noise. The calc Patterson comes from a + # rotated P1 ensemble that is NOT spacegroup-invariant — those m's + # carry real signal and must be kept. Phaser does the same asymmetric + # filtering (DataMR.cc applies the m-filter; Ensemble.cc does not). f_obs = compute_ball_harmonic_coefficients( s_obs, e_obs, L=L, P=P, d_min=d_min, d_max=d_max, + zsymm=zsymm, skip_odd_l=skip_odd_l, ) # Share shell edges between obs and calc — required for the cross-correlation # to be meaningful (same radial binning on both sides). f_calc = compute_ball_harmonic_coefficients( s_calc, e_calc, L=L, P=P, shell_edges=f_obs.shell_edges, + zsymm=1, skip_odd_l=skip_odd_l, ) if auto_variance_weights and weights is None: diff --git a/torchref/alignment/frf/bessel.py b/torchref/alignment/frf/bessel.py new file mode 100644 index 00000000..56a2e863 --- /dev/null +++ b/torchref/alignment/frf/bessel.py @@ -0,0 +1,97 @@ +"""Spherical Bessel functions ``j_u(x)`` for the Bessel-SH expansion. + +Phaser source: Phaser uses Miller downward recurrence in its own +implementation (``phaser/lib/jiffy.h``) for stability. We use the same +recurrence here because ``scipy.special.spherical_jn`` is not vectorised +across torch tensors and goes via NumPy round-trip — too slow for the +``(N_radial × N_obs)`` table required by ``bessel_sh_expand``. + +Numerical reference for Miller downward recurrence: + Abramowitz & Stegun §10.1.19 — start at high u, recur downward, then + rescale by the analytic ``j_0(x) = sin(x)/x``. +""" +from __future__ import annotations + +import torch + + +def spherical_bessel_table( + x: torch.Tensor, + u_max: int, +) -> torch.Tensor: + """Tabulate ``j_u(x)`` for u ∈ [0, u_max] over a 1-D tensor of x. + + Returns + ------- + j : torch.Tensor, shape (u_max + 1, x.numel()) + ``j[u, i] = j_u(x[i])``. Always float64 internally. + """ + if x.ndim != 1: + raise ValueError(f"expected 1-D x, got shape {tuple(x.shape)}") + if u_max < 0: + raise ValueError(f"u_max must be >= 0, got {u_max}") + + device = x.device + x64 = x.to(torch.float64) + n = x64.numel() + out = torch.zeros((u_max + 1, n), dtype=torch.float64, device=device) + + # Handle x == 0 separately (j_0(0)=1, j_u(0)=0 for u>=1). + zero_mask = x64 == 0 + nz_mask = ~zero_mask + out[0, zero_mask] = 1.0 + if not nz_mask.any(): + return out + + xnz = x64[nz_mask] + nnz = xnz.numel() + + # Direct formulas for u = 0, 1 — accurate for all x > 0. + j0 = torch.sin(xnz) / xnz + j1 = (torch.sin(xnz) - xnz * torch.cos(xnz)) / xnz**2 + + if u_max == 0: + out[0, nz_mask] = j0 + return out + if u_max == 1: + out[0, nz_mask] = j0 + out[1, nz_mask] = j1 + return out + + # Forward recurrence is unstable when u >> x; switch to Miller downward + # recurrence for those entries. Cutoff u_fwd_safe = ceil(x) is generous; + # Phaser uses similar logic in jiffy.h. + u_fwd_safe = torch.clamp(torch.ceil(xnz).to(torch.int64), min=2) + # Allocate per-x downward recurrence with a generous starting index + # (u_max + a few extra terms) — Miller convention. + u_start = u_max + 15 + f_curr = torch.zeros(nnz, dtype=torch.float64, device=device) + f_next = torch.ones(nnz, dtype=torch.float64, device=device) + table = torch.zeros((u_max + 1, nnz), dtype=torch.float64, device=device) + + # Downward: j_{u-1}(x) = (2u+1)/x · j_u(x) - j_{u+1}(x) + for u in range(u_start, -1, -1): + f_prev = (2 * u + 1) / xnz * f_next - f_curr + if u <= u_max: + table[u] = f_prev + f_curr = f_next + f_next = f_prev + + # Rescale so table[0] matches analytic j_0(x) = sin(x)/x. + scale = j0 / table[0] + table = table * scale.unsqueeze(0) + + # Override with the forward direct values where they are stable + # (small u, small x). Forward recurrence: j_{u+1} = (2u+1)/x · j_u - j_{u-1}. + fwd = torch.zeros((u_max + 1, nnz), dtype=torch.float64, device=device) + fwd[0] = j0 + fwd[1] = j1 + for u in range(1, u_max): + fwd[u + 1] = (2 * u + 1) / xnz * fwd[u] - fwd[u - 1] + # Per-x, use forward where u <= u_fwd_safe; otherwise Miller-rescaled. + u_idx = torch.arange(u_max + 1, device=device).unsqueeze(1) # (u_max+1, 1) + use_fwd = u_idx <= u_fwd_safe.unsqueeze(0) # (u_max+1, nnz) + table_combined = torch.where(use_fwd, fwd, table) + + out[:, nz_mask] = table_combined + return out diff --git a/torchref/alignment/frf/data_mr.py b/torchref/alignment/frf/data_mr.py new file mode 100644 index 00000000..582d14c2 --- /dev/null +++ b/torchref/alignment/frf/data_mr.py @@ -0,0 +1,165 @@ +"""Bessel-radial × spherical-harmonic expansion of obs and calc Pattersons. + +Mirrors ``DataMR::dataMR_FRF`` (DataMR.cc) and the helper sums in +``Ensemble.cc``. We import the validated ``bessel_sh_expand`` from +``torchref.alignment.phaser_frf`` and add the obs/calc cross-correlation +contraction that's the input to ``SiteListAng::get_FRF``. + +Citations: + * Bessel-radial × SH expansion: DataMR.cc:993, 1107-1117 + * sqrt(2u+1) · j_u(h)/h radial weight: DataMR.cc:993 + * Even-l only (Patterson centrosymmetry): implicit in DataMR.cc + * m-symmetry filter: DataMR.cc:863-870, 1117 +""" +from __future__ import annotations + +import math + +import torch + +from .phaser_frf import spherical_bessel_table +from ..sh import evaluate_ylm + +from .types import BesselSHCoefficients + +__all__ = [ + "bessel_sh_expand", + "cross_correlate_xi", +] + + +def bessel_sh_expand( + s_vectors: torch.Tensor, + intensity: torch.Tensor, + *, + L: int, + bessel_h_scale: float, + zsymm: int = 1, + enforce_friedel: bool = True, + chunk_size: int = -1, +) -> BesselSHCoefficients: + """Phaser-style ``c_nlm = Σ_h Y*_lm(ŝ) · I · sqrt(2u+1) · j_u(h)/h``. + + Memory-bounded chunked reimplementation of + ``torchref.alignment.phaser_frf.bessel_sh_expand`` (identical math, + verified element-wise by ``tests/unit/frf_separate``). The legacy + version materialises the full ``(M, L, N_radial)`` Bessel table and + ``(M, u_max+1)`` j-table for *all* reflections at once — at L≈100 with + a symmetry-unrolled obs set (≳10⁶ reflections) that is tens of GB and + OOMs. Here the j-table, Bessel weights and Y_lm are all computed + *inside* the reflection-chunk loop, so peak memory is set by one chunk. + + Citations: + * radial × SH expansion, sqrt(2u+1)·j_u(h)/h weight: DataMR.cc:993, 1107 + * even-l only (Patterson centrosymmetry) + m-filter: DataMR.cc:863-870, 1117 + + Parameters + ---------- + chunk_size : int + Reflections per chunk. ``-1`` (default) auto-sizes so the per-chunk + Y_lm block ``(chunk, L, 2L-1)`` stays near ~256 MB. + """ + assert s_vectors.dim() == 2 and s_vectors.shape[-1] == 3 + assert intensity.dim() == 1 and intensity.shape[0] == s_vectors.shape[0] + + real_dtype = s_vectors.dtype + if real_dtype == torch.float64: + complex_dtype = torch.complex128 + elif real_dtype == torch.float32: + complex_dtype = torch.complex64 + else: + raise TypeError(f"Unsupported real dtype: {real_dtype}") + device = s_vectors.device + + if enforce_friedel: + s_vectors = torch.cat([s_vectors, -s_vectors], dim=0) + intensity = torch.cat([intensity, intensity], dim=0) + + lmax = L - 1 + lmax_even = lmax if (lmax % 2 == 0) else (lmax - 1) + if lmax_even < 2: + raise ValueError(f"L={L} too small; need lmax_even >= 2 (so L >= 3).") + N_radial = (lmax_even - 2) // 2 + 1 + u_max = lmax_even + 1 + + # (l, n) -> u = l + 2n + 1 and the sqrt(2u+1) weight, precomputed once. + even_ls = list(range(2, lmax_even + 1, 2)) + ln_u = [] # (l, n, u, sqrt(2u+1)) + for l in even_ls: + n_l = (lmax_even - l) // 2 + 1 + for n in range(n_l): + u = l + 2 * n + 1 + ln_u.append((l, n, u, math.sqrt(float(2 * u + 1)))) + + M = s_vectors.shape[0] + if chunk_size <= 0: + # target ~256 MB for the complex (chunk, L, 2L-1) Y block + chunk_size = int(max(256, min(8192, 16_000_000 // max(1, L * (2 * L - 1))))) + + c_nlm = torch.zeros( + (N_radial, L, 2 * L - 1), dtype=complex_dtype, device=device, + ) + + for start in range(0, M, chunk_size): + stop = min(start + chunk_size, M) + s_c = s_vectors[start:stop] + i_c = intensity[start:stop] + s_mag = s_c.norm(dim=-1).clamp(min=1e-30) + s_hat = s_c / s_mag.unsqueeze(-1) + cos_theta = s_hat[..., 2].clamp(min=-1.0, max=1.0) + theta = torch.acos(cos_theta) + phi = torch.atan2(s_hat[..., 1], s_hat[..., 0]) + + x = (bessel_h_scale * s_mag).clamp(min=1e-30) # (c,) + j_all = spherical_bessel_table(x, u_max) # (c, u_max+1) + + bessel = torch.zeros((stop - start, L, N_radial), dtype=real_dtype, device=device) + for (l, n, u, w) in ln_u: + bessel[:, l, n] = w * j_all[:, u] / x + + Y = evaluate_ylm(theta, phi, L) # (c, L, 2L-1) + Y_w = torch.conj(Y) * i_c.to(complex_dtype).view(-1, 1, 1) + c_nlm += torch.einsum("hln,hlm->nlm", bessel.to(complex_dtype), Y_w) + + # Zero odd-l rows and l = 0 (Patterson centrosymmetry; Phaser drops l=0). + l_vals = torch.arange(L, device=device) + c_nlm[:, (l_vals % 2 == 1) | (l_vals == 0), :] = 0.0 + + # m-symmetry filter (observed side only; caller passes zsymm=1 for calc). + if zsymm > 1: + m_vals = torch.arange(-(L - 1), L, device=device) + c_nlm[:, :, (m_vals.abs() % zsymm) != 0] = 0.0 + + return BesselSHCoefficients( + coeffs=c_nlm, + L=L, + N_radial=N_radial, + bessel_h_scale=float(bessel_h_scale), + ) + + +def cross_correlate_xi( + c_obs: BesselSHCoefficients, + c_calc: BesselSHCoefficients, +) -> torch.Tensor: + """Contract obs/calc Bessel-SH coefficients on the radial-Bessel axis. + + Phaser source: the radial sum ``Σ_n c_obs[n,l,m] · conj(c_calc[n,l,m'])`` + happens inside ``DataMR::dataMR_FRF`` before being fed into + ``SiteListAng::DoRfftStuff`` as the ``clmn`` tensor (FastRot.cc:39). + + Convention (matches torchref's existing ball_search.py:182): + xi[l, m, n] = Σ_r c_obs[r, l, n] · conj(c_calc[r, l, m]) + so that the peak Euler triple satisfies ``s_calc = R · s_obs``. + + Returns + ------- + xi : torch.Tensor (complex), shape (L, 2L-1, 2L-1) + """ + if c_obs.L != c_calc.L: + raise ValueError(f"L mismatch: obs={c_obs.L} calc={c_calc.L}") + return torch.einsum( + "rln,rlm->lmn", + c_obs.coeffs, + torch.conj(c_calc.coeffs), + ) diff --git a/torchref/alignment/frf/dense_calc.py b/torchref/alignment/frf/dense_calc.py new file mode 100644 index 00000000..afa85324 --- /dev/null +++ b/torchref/alignment/frf/dense_calc.py @@ -0,0 +1,97 @@ +"""Dense P1-box sampling of a model's molecular transform. + +This is the "dense calc" lever from the high-symmetry FRF investigation +(see ``FRF_CONSOLIDATION.md``): the Fast Rotation Function correlates the obs +Patterson against the *model* transform, and sampling that transform at the +sparse crystal lattice under-determines the high-l spherical-harmonic modes for +large molecules. Phaser avoids this by computing the model transform on a dense, +oversampled P1-box FFT grid (``EnsemblePDB.cc:122-135``). This module does the +same: drop the (single, un-symmetry-expanded) model into a cubic P1 box and +reuse ``ModelFT``'s own structure-factor machinery (real ITC92 form factors + +per-atom B/occ) to sample ``|F_calc|`` on the box's dense reciprocal grid. + +Unlike the original benchmark helper this operates on ``model.copy()`` so the +caller's ``cell``/``spacegroup``/``max_res`` are never mutated, and the whole SF +build runs under ``torch.no_grad()`` (the load-bearing memory fix — a forward-only +SF build otherwise accumulates a backward graph and OOMs on big grids). +""" +from __future__ import annotations + +import math +from typing import TYPE_CHECKING, Tuple + +import torch + +if TYPE_CHECKING: + from torchref.model import ModelFT + + +def dense_calc_via_box( + model: "ModelFT", + d_max: float, + d_min: float, + *, + pad: float = 2.0, + verbose: bool = False, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Sample the model transform on a dense P1-box grid in ``[d_max, d_min]``. + + Parameters + ---------- + model : ModelFT + Search model (in whatever orientation the rotation search should treat + as the reference). **Not mutated** — a ``model.copy()`` is used internally. + d_max, d_min : float + Low- and high-resolution limits (Å) for the returned reflections. + pad : float, optional + P1-box edge as a multiple of the molecular diameter (default 2.0, the + validated v19 value). A bigger box = finer reciprocal sampling. + verbose : bool, optional + If True, print a one-line ``[DENSE_FT]`` summary (box edge / grid size). + + Returns + ------- + (s_vec, F_calc) : Tuple[torch.Tensor, torch.Tensor] + ``s_vec`` is the Cartesian reciprocal grid (N, 3) and ``F_calc`` the + amplitudes (N,), both float64 on the model's device. + """ + from torchref.symmetry.cell import Cell + + m = model.copy() # isolate the box mutation from the caller + with torch.no_grad(): + coords = m.xyz() + dev = coords.device + # Cubic P1 box sized to ``pad`` diameters; no symmetry keeps the grid small. + extent = (coords - coords.mean(0)).norm(dim=-1).max().item() + a = float(pad * 2.0 * extent) + # Order matters: set max_res first, then the sg/cell setters rebuild the + # FFT (via _maybe_initialize_fft) reading the new max_res. + m.max_res = float(d_min) + m.spacegroup = "P 1" + m.cell = Cell([a, a, a, 90.0, 90.0, 90.0], device=dev) + + nmax = int(math.ceil(a / d_min)) + idx = torch.arange(-nmax, nmax + 1, device=dev) + H, K, Lg = torch.meshgrid(idx, idx, idx, indexing="ij") + hkl = torch.stack( + [H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], dim=-1 + ).to(torch.long) + # Cubic box: |s| = |hkl| / a. + smag = hkl.to(torch.float64).norm(dim=-1) / a + keep = (smag >= 1.0 / d_max) & (smag <= 1.0 / d_min) + hkl = hkl[keep].contiguous() + F = model_sf_abs(m, hkl) + s_vec = hkl.to(torch.float64) / a + + if verbose: + print( + f"[DENSE_FT] box={a:.0f}A n_grid={hkl.shape[0]} max_res={d_min:.2f}", + flush=True, + ) + return s_vec, F + + +def model_sf_abs(model: "ModelFT", hkl: torch.Tensor) -> torch.Tensor: + """``|F_calc|`` (float64) for ``hkl`` via the model's SF machinery (no grad).""" + with torch.no_grad(): + return model.get_structure_factor(hkl, recalc=True).abs().to(torch.float64) diff --git a/torchref/alignment/frf/peak_finder.py b/torchref/alignment/frf/peak_finder.py new file mode 100644 index 00000000..cb2f5562 --- /dev/null +++ b/torchref/alignment/frf/peak_finder.py @@ -0,0 +1,154 @@ +"""Peak finding on the adaptive SO(3) sample list. + +Phaser source: ``SiteListAng::findpeaks`` (referenced from FastRot.cc; +implementation in ``SiteListAng.cc`` ``findpeaks`` and the related NMS +routines). Phaser's strategy is essentially: + + 1. Compute mean + std of all samples → z-score per sample. + 2. Sort by descending value. + 3. Greedy non-max suppression on SO(3) by *angular distance* between + rotations (not by α, β, γ box distance — that would double-count + near the poles). + +We implement the same flow in PyTorch, vectorised where possible. The +SO(3) angular-distance NMS is identical to the v13 ``_so3_greedy_nms`` +in ``ball_search.py`` — that part of v13 was correct; the bug was in +the FFT/grid, not the NMS. +""" +from __future__ import annotations + +import math +from typing import List + +import torch + +from .types import AdaptiveRotationFunction, RotationPeak + +__all__ = [ + "find_rotation_peaks", +] + + +def _euler_to_matrix_edmonds_zyz( + alpha: torch.Tensor, beta: torch.Tensor, gamma: torch.Tensor, +) -> torch.Tensor: + """R = R_z(α) R_y(β) R_z(γ) — Edmonds ZYZ convention. + + Returns shape (*alpha.shape, 3, 3) real. + """ + ca, sa = torch.cos(alpha), torch.sin(alpha) + cb, sb = torch.cos(beta), torch.sin(beta) + cg, sg = torch.cos(gamma), torch.sin(gamma) + # R = Rz(a) * Ry(b) * Rz(g) + R = torch.stack( + [ + torch.stack([ca * cb * cg - sa * sg, -ca * cb * sg - sa * cg, ca * sb], dim=-1), + torch.stack([sa * cb * cg + ca * sg, -sa * cb * sg + ca * cg, sa * sb], dim=-1), + torch.stack([-sb * cg, sb * sg, cb ], dim=-1), + ], + dim=-2, + ) + return R + + +def _so3_angular_distance_deg(R1: torch.Tensor, R2: torch.Tensor) -> torch.Tensor: + """Angular distance between two rotation matrices, in degrees. + + R1: (..., 3, 3), R2: (..., 3, 3). Returns (...,) real. + """ + trace = torch.einsum("...ij,...ij->...", R1, R2) + cos_theta = ((trace - 1.0) * 0.5).clamp(min=-1.0, max=1.0) + return torch.arccos(cos_theta) * (180.0 / math.pi) + + +def _so3_greedy_nms( + alphas: torch.Tensor, + betas: torch.Tensor, + gammas: torch.Tensor, + values: torch.Tensor, + nms_radius_deg: float, + keep_at_most: int, +) -> torch.Tensor: + """Return indices (into the input order) of kept peaks after SO(3) NMS. + + Greedy: walk the values in descending order; keep a candidate if its + angular distance from every already-kept rotation is > nms_radius_deg. + """ + n = values.shape[0] + if n == 0: + return torch.empty(0, dtype=torch.int64, device=values.device) + order = torch.argsort(values, descending=True) + R_all = _euler_to_matrix_edmonds_zyz(alphas, betas, gammas) # (n, 3, 3) + R_all = R_all.to(torch.float64) + + kept_idx: List[int] = [] + kept_R: List[torch.Tensor] = [] + for i_t in order.tolist(): + Ri = R_all[i_t] + if kept_R: + stack = torch.stack(kept_R, dim=0) # (k, 3, 3) + dists = _so3_angular_distance_deg(Ri.unsqueeze(0), stack) + if dists.min().item() <= nms_radius_deg: + continue + kept_R.append(Ri) + kept_idx.append(i_t) + if len(kept_idx) >= keep_at_most: + break + return torch.tensor(kept_idx, dtype=torch.int64, device=values.device) + + +def find_rotation_peaks( + arf: AdaptiveRotationFunction, + n_peaks: int = 500, + sigma_threshold: float = -5.0, + nms_radius_deg: float = 6.0, +) -> List[RotationPeak]: + """Greedy SO(3) NMS over the adaptive sample list. + + Returns peaks sorted by descending value, capped at ``n_peaks`` and + filtered by ``sigma >= sigma_threshold``. + """ + values = arf.values + if values.numel() == 0: + return [] + + mean = values.mean() + std = values.std().clamp(min=1e-30) + sigma = (values - mean) / std + + # Pre-filter by sigma threshold to keep the NMS loop tractable. + keep_mask = sigma >= sigma_threshold + if not keep_mask.any(): + return [] + + idx_filtered = torch.nonzero(keep_mask, as_tuple=False).squeeze(-1) + # Optionally cap candidate set so the O(n_kept · n_cand) NMS stays small. + candidate_cap = max(n_peaks * 20, 2000) + if idx_filtered.numel() > candidate_cap: + top_vals = values[idx_filtered] + top_keep = torch.topk(top_vals, candidate_cap).indices + idx_filtered = idx_filtered[top_keep] + + a = arf.alphas[idx_filtered] + b = arf.betas[idx_filtered] + g = arf.gammas[idx_filtered] + v = values[idx_filtered] + + kept = _so3_greedy_nms( + a, b, g, v, + nms_radius_deg=nms_radius_deg, + keep_at_most=n_peaks, + ) + + peaks: List[RotationPeak] = [] + for k in kept.tolist(): + peaks.append( + RotationPeak( + alpha=float(a[k].item()), + beta=float(b[k].item()), + gamma=float(g[k].item()), + value=float(v[k].item()), + sigma=float(((v[k] - mean) / std).item()), + ) + ) + return peaks diff --git a/torchref/alignment/frf/phaser_frf.py b/torchref/alignment/frf/phaser_frf.py new file mode 100644 index 00000000..c1509785 --- /dev/null +++ b/torchref/alignment/frf/phaser_frf.py @@ -0,0 +1,1164 @@ +""" +Phaser-faithful Fast Rotation Function (FRF) — a clean, parallel implementation. + +Translates Phaser's `DataMR.cc` (observed-side SH expansion) + `Ensemble.cc` +(model-side σA Eterm) + `FastRot.cc` (Wigner-D contraction + 2-D FFT) into +PyTorch. Designed to be benchmarked directly against torchref's existing +`ball_rotation_search` in +`tests/integration/alignment/benchmark_phaser_frf.py`. + +The single fundamental difference from `ball_rotation_search` is the **radial +basis**: this module expands the per-reflection Patterson onto **spherical +Bessel functions** rather than shell-step indicators. The Bessel basis is the +natural orthonormal radial basis on the unit ball (Crowther 1972, Navaza FAST); +shell-step is its coarsest piecewise-constant approximation. Phaser's +`DataMR.cc:1102–1134` accumulates + + clmn[l, m, n] = Σ_h Y*_{l,m}(ŝ_h) · I(h) · √u · j_u(2π·a·|s_h|) / (2π·a·|s_h|) + +with u = l + 2n − 1 and n ∈ [1, nmax(l)] where nmax(l) = (lmax − l)/2 + 1. + +Other Phaser-faithful choices we make: +- Even l only (Patterson is centrosymmetric); odd-l rows zeroed. +- m-symmetry filter on observed side only (`DataMR.cc:863-870, 1117`): + obs SH coefficients with `m % ZSYMM != 0` are dropped; calc side keeps all m. +- σA Eterm on calc side (`Ensemble.cc:36-46`): + `Eterm(s) = exp(-2π² · s² · ΔVRMS²)`. Per-reflection (not per-shell). +- LERF1-style observed intensity: `I_obs(h) = cweight · (E²−1)` with + `cweight = 1` for centric, `2` for acentric. +- Euler convention: this module computes in **Edmonds ZYZ** + `R = R_z(α) R_y(β) R_z(γ)` — same as the rest of torchref. Phaser's + internal convention is `R = R_z(γ) R_y(β) R_z(α)` (swapped α↔γ) but only + the rotation matrix matters for our rank-of-truth metric, so we never + swap explicitly — the orbit_rank metric in the benchmark builds rotation + matrices via `rotation_matrix_from_edmonds_euler` and the answer is + parameterization-independent. + +What we deliberately DO NOT yet implement (out of scope per plan): +- French-Wilson per-reflection DFAC. Phaser uses `(E²-V)/V² · DFAC²` with + per-reflection Luzzati DFAC ∈ [0.05, 10]. Here we use DFAC = 1. +- Adaptive (α, γ) grid sampling (Phaser's `pmax = 720·cos(β/2)/grid_sampling`). + We use uniform 2L grid. +- Axis permutation: Phaser detects the high-order axis and may permute the + coordinate frame. We use the z-preferred selection from + `get_high_order_axis`, applied without permutation. For groups where the + high-order axis is already z (most cases including P432 with body-diagonal + 3-folds along z), this is identical to Phaser. + +References (paths under +`/das/work/p17/p17490/Peter/Library/torchref/reverse_engineering/phenix/phenix-1.20-4459/modules/phaser/codebase/phaser/src/`): +- `DataMR.cc:863-1148` observed-side SH + intensity + ZSYMM filter +- `Ensemble.cc:36-46` model-side Eterm +- `FastRot.cc:30-217` Wigner contraction + FFT + Z-score +- `runMR_FRF.cc:546-587` peak picking + Z-score normalisation +""" + +from __future__ import annotations + +import math +from dataclasses import dataclass +from typing import List, Optional, Tuple + +import torch + +from .ball_search import ( + RotationPeak, + find_rotation_peaks_adaptive, + refine_peaks_subvoxel_adaptive, +) +from ..sh import ( + assign_shells, + equal_count_shell_edges, + evaluate_ylm, + get_high_order_axis, +) +from ..wigner import ( + AdaptiveRotationFunction, + evaluate_rotation_function_grid_adaptive, +) + + +# ============================================================================= +# Spherical Bessel j_u(x) via Miller's downward recurrence +# ============================================================================= + + +def spherical_bessel_table( + x: torch.Tensor, + u_max: int, + n_extra: int = 25, +) -> torch.Tensor: + """ + Tabulate spherical Bessel `j_u(x)` for u ∈ [0, u_max], batched over x. + + Uses Miller's downward recurrence (the standard stable choice for + `j_n(x)` with `n > x`): + + j_{u-1}(x) = (2u + 1) / x · j_u(x) − j_{u+1}(x) + + Seed: choose `n_start = u_max + n_extra`, set + `j_{n_start+1} = 0`, `j_{n_start} = 1` (unnormalised), recur down to + j_0, then renormalise using the exact `j_0(x) = sin(x) / x`. Float64 + internally for accuracy at moderate `u/x` ratios; cast back to input + dtype. + + Returns + ------- + j_table : torch.Tensor + Shape `(*x.shape, u_max + 1)`, dtype = `x.dtype`. + """ + real_dtype = x.dtype + device = x.device + x64 = x.to(torch.float64) + safe_x = x64.clamp(min=1e-30) + inv_x = 1.0 / safe_x + + n_start = max(u_max + n_extra, u_max + 2) + j_high = torch.zeros_like(x64) # j_{n_start + 1} + j_mid = torch.ones_like(x64) # j_{n_start} (arbitrary scale) + j_table = torch.zeros( + (u_max + 1, *x64.shape), dtype=torch.float64, device=device, + ) + + # Recur from n = n_start down to 1, generating j_{n-1} at each step. + for n in range(n_start, 0, -1): + j_low = (2.0 * n + 1.0) * inv_x * j_mid - j_high + if n - 1 <= u_max: + j_table[n - 1] = j_low + j_high = j_mid + j_mid = j_low + + # Normalise against the exact j_0(x) = sin(x)/x. Handle x ≈ 0 (where + # j_0(0) = 1, j_u(0) = 0 for u ≥ 1) so the renormalisation factor is + # well-defined. + true_j0 = torch.sin(x64) * inv_x + true_j0 = torch.where(x64 < 1e-30, torch.ones_like(x64), true_j0) + computed_j0 = j_table[0] + safe_j0 = torch.where( + computed_j0.abs() < 1e-30, torch.ones_like(computed_j0), computed_j0, + ) + scale = true_j0 / safe_j0 + j_table = j_table * scale.unsqueeze(0) + + # Re-arrange `(u_max+1, *x.shape) → (*x.shape, u_max+1)`. + perm = list(range(1, j_table.dim())) + [0] + j_table = j_table.permute(*perm).contiguous() + return j_table.to(real_dtype) + + +# ============================================================================= +# Phaser-style Bessel-radial SH expansion +# ============================================================================= + + +@dataclass +class BesselSHCoefficients: + """Phaser-style clmn coefficients. + + `c_nlm[n, l, m+L-1]` is the (n-th radial Bessel) × (l, m) coefficient. + """ + c_nlm: torch.Tensor # complex, shape (N_radial, L, 2L-1) + L: int # bandwidth (l ∈ [0, L), even-l only filled) + N_radial: int # number of Bessel radial terms = (lmax-2)/2 + 1 + bessel_h_scale: float # `h = bessel_h_scale · |s|` (Phaser uses lmax · d_min) + zsymm: int # m-symmetry filter applied (1 = none) + + +def bessel_sh_expand( + s_vectors: torch.Tensor, + intensity: torch.Tensor, + *, + L: int, + bessel_h_scale: float, + zsymm: int = 1, + enforce_friedel: bool = True, + chunk_size: int = 2048, +) -> BesselSHCoefficients: + """ + Compute the Phaser-style spherical-Bessel radial × spherical-harmonic + angular expansion of a scattered-point intensity field on the unit ball. + + c_nlm[n, l, m] = Σ_h Y*_{l,m}(ŝ_h) · I_h · √u · j_u(h_h) / h_h + + where `h_h = bessel_h_scale · |s_h|` and u = l + 2n + 1 (n ∈ [0, N_radial)). + Phaser (`DataMR.cc:1107`) uses `h = lmax · |s| · HIRES` where HIRES is the + high-resolution limit `d_min` in Å. With that scaling, h_max ≈ lmax (so the + highest-u Bessel functions are sampled near their first peak rather than + deep in their decaying tail). Pass `bessel_h_scale = (L - 1) * d_min` to + match Phaser exactly. + + For each l, only n with `2n + l + 1 ≤ lmax + 1` (i.e. `n ≤ (lmax − l)/2`) + are filled — others stay 0 (Phaser's truncation: + `nmax(l) = (lmax - l + 2)/2`). + + Even l only (the Patterson is centrosymmetric: Y_{l,m}(−ŝ) = (−1)^l Y_{l,m}, + so odd-l rows cancel exactly when Friedel mates are summed). + + When `zsymm > 1`, SH coefficients with `|m| % zsymm ≠ 0` are zeroed + post-expansion (Phaser `DataMR.cc:863-870, 1117`). This is the + m-symmetry filter; applied **only** on the observed side by the caller, + NOT on the calc side (which is a rotated P1 ensemble and not spacegroup- + invariant). + + Parameters + ---------- + s_vectors : (N, 3) real + Reciprocal-lattice vectors in 1/Å. + intensity : (N,) real + Per-reflection intensity (LERF1 obs intensity or σA-weighted calc). + L : int + Wigner/SH bandwidth. lmax = L - 1, restricted to even. + bessel_h_scale : float + The pre-multiplier on |s| inside the Bessel argument: `h = bessel_h_scale · |s|`. + Phaser uses `lmax · d_min` (where d_min is HIRES in Å). This puts h_max ≈ lmax + so the highest-u terms are sampled near their first peak. + zsymm : int, default 1 + m-symmetry filter (zero coefficients with |m| not divisible by zsymm). + zsymm=1 is no filter. + enforce_friedel : bool, default True + Augment the sum with (-s, I) pairs to make odd-l rows exactly zero + (otherwise we rely on the explicit zero-out at the end). + chunk_size : int + Reflections per Y_lm chunk for memory control. + + Returns + ------- + BesselSHCoefficients with `c_nlm[n, l, m+L-1]` populated. + """ + assert s_vectors.dim() == 2 and s_vectors.shape[-1] == 3 + assert intensity.dim() == 1 and intensity.shape[0] == s_vectors.shape[0] + + real_dtype = s_vectors.dtype + if real_dtype == torch.float64: + complex_dtype = torch.complex128 + elif real_dtype == torch.float32: + complex_dtype = torch.complex64 + else: + raise TypeError(f"Unsupported real dtype: {real_dtype}") + device = s_vectors.device + + if enforce_friedel: + s_vectors = torch.cat([s_vectors, -s_vectors], dim=0) + intensity = torch.cat([intensity, intensity], dim=0) + + s_mag = s_vectors.norm(dim=-1).clamp(min=1e-30) + s_hat = s_vectors / s_mag.unsqueeze(-1) + cos_theta = s_hat[..., 2].clamp(min=-1.0, max=1.0) + theta = torch.acos(cos_theta) + phi = torch.atan2(s_hat[..., 1], s_hat[..., 0]) + + lmax = L - 1 + lmax_even = lmax if (lmax % 2 == 0) else (lmax - 1) + # Phaser indexing: n ∈ [1, nmax(l)] with nmax(l) = (lmax - l + 2)/2. + # Our 0-indexed n: n ∈ [0, N_radial - 1] with N_radial = nmax(l=2). + if lmax_even < 2: + raise ValueError( + f"L={L} too small; need lmax_even >= 2 (so L >= 3)." + ) + N_radial = (lmax_even - 2) // 2 + 1 + # Highest Bessel order needed: for l = 2 and n = N_radial − 1 (1-indexed + # n_max), u = l + 2n + 1 = 2 + 2(N_radial − 1) + 1 = lmax_even + 1. + u_max = lmax_even + 1 + + # Bessel j_u(bessel_h_scale · |s|) for u ∈ [0, u_max], per reflection. + x = bessel_h_scale * s_mag # (M,) + j_all = spherical_bessel_table(x, u_max) # (M, u_max+1) + + # bessel_factor[h, l, n] = √(2u+1) · j_u(x_h) / x_h, with u = l + 2n + 1. + # Matches Phaser DataMR.cc:993 + 1109: + # sqrt_table[i] = sqrt(2*i+1); besselx[u] = sqrt_table[u] · sphbessel(u, h) / h + # (Crowther/Navaza unit-ball radial-basis weight; the previous mimic used + # √u, an unweighted variant that biased the radial power spectrum.) + M = s_vectors.shape[0] + bessel = torch.zeros((M, L, N_radial), dtype=real_dtype, device=device) + safe_x = x.clamp(min=1e-30) + for l in range(2, lmax_even + 1, 2): + n_l = (lmax_even - l) // 2 + 1 + for n in range(n_l): + u = l + 2 * n + 1 + bessel[:, l, n] = math.sqrt(float(2 * u + 1)) * j_all[:, u] / safe_x + + # Accumulate the SH expansion in chunks. + # c_nlm[n, l, m] = Σ_h conj(Y_lm(ŝ_h)) · I_h · bessel[h, l, n] + c_nlm = torch.zeros( + (N_radial, L, 2 * L - 1), dtype=complex_dtype, device=device, + ) + intensity_c = intensity.to(complex_dtype) + + for start in range(0, M, chunk_size): + stop = min(start + chunk_size, M) + Y = evaluate_ylm(theta[start:stop], phi[start:stop], L) # (n, L, 2L-1) + # weight × conj(Y) per reflection. + Y_w = torch.conj(Y) * intensity_c[start:stop].view(-1, 1, 1) # (n, L, 2L-1) + b_slice = bessel[start:stop].to(complex_dtype) # (n, L, N_radial) + # einsum: sum over h, contract bessel(h,l,n) × Y_w(h,l,m). + contrib = torch.einsum("hln,hlm->nlm", b_slice, Y_w) + c_nlm += contrib + + # Zero odd-l rows (and l = 0, which Phaser also doesn't use). + l_vals = torch.arange(L, device=device) + odd_or_zero_mask = (l_vals % 2 == 1) | (l_vals == 0) + c_nlm[:, odd_or_zero_mask, :] = 0.0 + + # m-symmetry filter (applied to observed side only; caller passes + # zsymm = 1 for calc side). + if zsymm > 1: + m_vals = torch.arange(-(L - 1), L, device=device) # (2L-1,) + m_invalid = (m_vals.abs() % zsymm) != 0 + c_nlm[:, :, m_invalid] = 0.0 + + return BesselSHCoefficients( + c_nlm=c_nlm, L=L, N_radial=N_radial, + bessel_h_scale=float(bessel_h_scale), zsymm=int(zsymm), + ) + + +# ============================================================================= +# Phaser-style data preparation: Wilson normalisation + σA Eterm +# ============================================================================= + + +def _wilson_normalise( + F: torch.Tensor, + s_mag: torch.Tensor, + n_shells: int = 20, +) -> Tuple[torch.Tensor, torch.Tensor]: + """ + Per-shell Wilson normalisation of amplitudes: + + E_h = F_h / sqrt(_p) where p = shell containing h. + + Phaser's `Feff[r] / SIGMAN.sqrt_epsnSN[r]` (`DataMR.cc:925`) does + roughly this (modulo French-Wilson and explicit ε-factor handling + which we skip; the input F is already anisotropy-corrected by the + caller in practice). + + Returns (E_h, sqrt_mean_F2_per_h). + """ + edges, _ = equal_count_shell_edges(s_mag, n_shells) + shell_idx = assign_shells(s_mag, edges) + valid = shell_idx >= 0 + F_dtype = F.dtype + F2 = F * F + count = torch.zeros(n_shells, dtype=torch.int64, device=F.device) + sumF2 = torch.zeros(n_shells, dtype=F_dtype, device=F.device) + F2_v = F2[valid] + idx_v = shell_idx[valid] + count.index_add_(0, idx_v, torch.ones_like(idx_v)) + sumF2.index_add_(0, idx_v, F2_v) + mean_F2 = sumF2 / count.clamp(min=1).to(F_dtype) + mean_F2 = mean_F2.clamp(min=1e-12) + sqrt_mean = mean_F2.sqrt() + per_h = torch.ones_like(F) + per_h[valid] = sqrt_mean[idx_v] + E = F / per_h + return E, per_h + + +# ----------------------------------------------------------------------------- +# French-Wilson posterior expected values (math_FrenchWilson.cc) +# ----------------------------------------------------------------------------- + + +def _expectE_FW_acen(eosq, sigesq): + """ + Acentric posterior expected E from normalised observed intensity (eosq) + and its standard deviation (sigesq). Translates verbatim from Phaser's + `lib/math_FrenchWilson.cc:expectEFWacen` (lines 8-44). Vectorised NumPy. + + `eosq = Iobs / `, `sigesq = σIobs / `. + """ + import numpy as np + from scipy.special import erfc, pbdv + CROSS1, CROSS2 = -12.5, 18.0 + SQRT2 = np.sqrt(2.0) + x = (eosq - sigesq ** 2) / sigesq + xsqr = x * x + ee = np.empty_like(eosq) + # Large negative argument: asymptotic + m_neg = x < CROSS1 + if m_neg.any(): + xs = xsqr[m_neg] + num = (-916620705. + xs * + (91891800. + xs * + (-11531520. + xs * + (1935360. + xs * + (-491520. + xs * 262144.))))) + den = (-495452160. + xs * + (55050240. + xs * + (-7864320. + xs * + (1572864. + xs * + (-524288. + xs * 524288.))))) + ee[m_neg] = np.sqrt(-np.pi * sigesq[m_neg] / x[m_neg]) * num / den + # Large positive argument: asymptotic + m_pos = x > CROSS2 + if m_pos.any(): + xs = xsqr[m_pos] + num = (-45045. + 32. * xs * + (-315. + 8. * xs * + (-15. - 16. * xs + 128. * xs * xs))) + ee[m_pos] = (np.sqrt(sigesq[m_pos]) * num / + (32768. * x[m_pos] ** 7.5)) + # Moderate arguments + m_mid = ~(m_neg | m_pos) + if m_mid.any(): + xm = x[m_mid] + pcd, _ = pbdv(-1.5, -xm) + ee[m_mid] = (np.sqrt(sigesq[m_mid] / 2.0) * np.exp(-xm * xm / 4.0) * + pcd / erfc(-xm / SQRT2)) + return ee + + +def _expectEsq_FW_acen(eosq, sigesq): + """Acentric posterior . From `expectEsqFWacen` (lines 46-78).""" + import numpy as np + from scipy.special import erfc + CROSS1, CROSS2 = -8.9, 5.7 + SQRT2_BY_PI = np.sqrt(2.0 / np.pi) + SQRT2 = np.sqrt(2.0) + eesq_base = eosq - sigesq ** 2 # baseline value + x = eesq_base / (SQRT2 * sigesq) + xsqr = x * x + eesq = eesq_base.copy() + m_neg = x < CROSS1 + if m_neg.any(): + xs = xsqr[m_neg] + num = (-135135. + xs * (20790. + xs * (-3780. + xs * + (840. + xs * (-240. + xs * (96. - xs * 64.)))))) + den = (-135135. + xs * (20790. + xs * (-3780. + xs * + (840. + xs * (-240. + xs * (96. + xs * + (-64. + xs * 128.))))))) + eesq[m_neg] = eesq_base[m_neg] * num / den + m_mid = (x >= CROSS1) & (x <= CROSS2) + if m_mid.any(): + xm = x[m_mid] + eesq[m_mid] = (eesq_base[m_mid] + + SQRT2_BY_PI * sigesq[m_mid] / + (np.exp(xm * xm) * erfc(-xm))) + # x > CROSS2: eesq stays at eesq_base (default per source) + return eesq + + +def _expectE_FW_cen(eosq, sigesq): + """Centric posterior . From `expectEFWcen` (lines 80-113).""" + import numpy as np + from scipy.special import pbdv + CROSS1, CROSS2 = -17.5, 17.5 + SQRTPI = np.sqrt(np.pi) + x = sigesq / 2.0 - eosq / sigesq + xsqr = x * x + pcdratio = np.empty_like(x) + m_neg = x < CROSS1 + if m_neg.any(): + xn, xs = x[m_neg], xsqr[m_neg] + pcdratio[m_neg] = ((1024. * SQRTPI * (-xn) ** 6.5) / + (3465. + xs * + (840. + xs * + (384. + xs * 1024.)))) + m_pos = x > CROSS2 + if m_pos.any(): + xp, xs = x[m_pos], xsqr[m_pos] + num = (3440640. + xs * + (-491520. + xs * + (98304. + xs * + (-32768. + xs * 32768.)))) + den = (675675. + xs * + (-110880. + xs * + (26880. + xs * + (-12288. + xs * 32768.)))) + pcdratio[m_pos] = num / (den * np.sqrt(xp)) + m_mid = ~(m_neg | m_pos) + if m_mid.any(): + xm = x[m_mid] + d_neg1, _ = pbdv(-1.0, xm) + d_neghalf, _ = pbdv(-0.5, xm) + pcdratio[m_mid] = d_neg1 / d_neghalf + return np.sqrt(sigesq / np.pi) * pcdratio + + +def _expectEsq_FW_cen(eosq, sigesq): + """Centric posterior . From `expectEsqFWcen` (lines 115-152).""" + import numpy as np + from scipy.special import pbdv + CROSS1, CROSS2 = -17.5, 17.5 + x = sigesq / 2.0 - eosq / sigesq + xsqr = x * x + pcdratio = np.empty_like(x) + m_neg = x < CROSS1 + if m_neg.any(): + xn, xs = x[m_neg], xsqr[m_neg] + num = (45045. + xs * + (10080. + xs * + (3840. + xs * + (4096. - xs * 32768.)))) + den = xn * (55440. + xs * + (13440. + xs * + (6144. + xs * 16384.))) + pcdratio[m_neg] = num / den + m_pos = x > CROSS2 + if m_pos.any(): + xp, xs = x[m_pos], xsqr[m_pos] + num = (11486475. + xs * + (-1441440. + xs * + (241920. + xs * + (-61440. + xs * 32768.)))) + den = xp * (675675. + xs * + (-110880. + xs * + (26880. + xs * + (-12288. + xs * 32768.)))) + pcdratio[m_pos] = num / den + m_mid = ~(m_neg | m_pos) + if m_mid.any(): + xm = x[m_mid] + d_neg15, _ = pbdv(-1.5, xm) + d_neghalf, _ = pbdv(-0.5, xm) + pcdratio[m_mid] = d_neg15 / d_neghalf + return sigesq * pcdratio / 2.0 + + +def _french_wilson_posterior(eosq, sigesq, centric_mask): + """Wrap centric/acentric branches. + + Phaser `expectEFW` / `expectEsqFW` (lines 154-178): if sigesq <= 0 the + measurement is treated as exact and (eEFW, eEsqFW) = (sqrt(eosq), eosq). + """ + import numpy as np + eEFW = np.empty_like(eosq) + eEsqFW = np.empty_like(eosq) + zero_sig = sigesq <= 0.0 + if zero_sig.any(): + eEFW[zero_sig] = np.sqrt(np.maximum(eosq[zero_sig], 0.0)) + eEsqFW[zero_sig] = np.maximum(eosq[zero_sig], 0.0) + valid = ~zero_sig + if valid.any(): + cen = centric_mask & valid + acen = (~centric_mask) & valid + if cen.any(): + eEFW[cen] = _expectE_FW_cen(eosq[cen], sigesq[cen]) + eEsqFW[cen] = _expectEsq_FW_cen(eosq[cen], sigesq[cen]) + if acen.any(): + eEFW[acen] = _expectE_FW_acen(eosq[acen], sigesq[acen]) + eEsqFW[acen] = _expectEsq_FW_acen(eosq[acen], sigesq[acen]) + return eEFW, eEsqFW + + +# ----------------------------------------------------------------------------- +# DFAC via Halley iteration (math_RiceLLG.cc:getDfactor) +# ----------------------------------------------------------------------------- + + +def _i0e_full(x): + """Phaser's `eBesselI0(x) = I0(x)·exp(-|x|)`. Symmetric in x. + + See `math_eBesselI0.cc` — Phaser's Rice-moment formulas use the + exp-scaled Bessel (the un-scaled `I0` cancels analytically with the + Gaussian envelope of the Rice distribution). + """ + import numpy as np + from scipy.special import i0e + return i0e(np.abs(x)) + + +def _i1e_full(x): + """Phaser's `eBesselI1(x) = I1(x)·exp(-|x|)`. Antisymmetric in x.""" + import numpy as np + from scipy.special import i1e + return np.sign(x) * i1e(np.abs(x)) + + +def _effSigaRoot_acen(ee, eesq, sa): + """`effSigaRootAcen` (math_RiceLLG.cc:12-34).""" + import numpy as np + sigbsqr = 1.0 - sa * sa + x = 0.5 * (eesq - sigbsqr) / sigbsqr + return (np.sqrt(np.pi * sigbsqr) / (2.0 * sigbsqr) * + (eesq * _i0e_full(x) + (eesq - sigbsqr) * _i1e_full(x)) - ee) + + +def _deffSigaRoot_acen(eesq, sa): + """`deffSigaRootAcen_by_dsa` (lines 36-52).""" + import numpy as np + sigbsqr = 1.0 - sa * sa + x = 0.5 * (eesq - sigbsqr) / sigbsqr + return np.sqrt(np.pi / sigbsqr) * (sa / 2.0) * _i1e_full(x) + + +def _d2effSigaRoot_acen(eesq, sa): + """`d2effSigaRootAcen_by_dsa2` (lines 54-81).""" + import numpy as np + sigasqr = sa * sa + sigapow4 = sigasqr * sigasqr + sigbsqr = 1.0 - sigasqr + xnum = eesq - sigbsqr + x = 0.5 * xnum / sigbsqr + out = np.empty_like(eesq) + big = xnum > 1e-10 + if big.any(): + I0 = _i0e_full(x[big]) + I1 = _i1e_full(x[big]) + out[big] = (np.sqrt(np.pi / sigbsqr[big]) / (2.0 * sigbsqr[big] ** 2) * + (eesq[big] * sigasqr[big] * I0 + + (eesq[big] - 1.0 - (-2.0 + eesq[big] * (2.0 + eesq[big])) * sigasqr[big] + + (eesq[big] - 1.0) * sigapow4[big]) * I1 / xnum[big])) + small = ~big + if small.any(): + samin = np.sqrt(np.maximum(1.0 - eesq[small], 0.0)) + out[small] = (np.sqrt(np.pi) * samin * + ((3.0 + samin * samin) * sa[small] - + 2.0 * (samin + samin ** 3)) / + (4.0 * eesq[small] ** 2.5)) + return out + + +def _effSigaRoot_cen(ee, eesq, sa): + """`effSigaRootCen` (lines 83-105).""" + import numpy as np + from scipy.special import erf + sigbsqr = 1.0 - sa * sa + x = 0.5 * (eesq - sigbsqr) / sigbsqr + x_safe = np.maximum(x, 0.0) # erf(sqrt(x)) ill-defined for x<0 + return (np.exp(-x) * np.sqrt(2.0 * sigbsqr / np.pi) + + np.sqrt(np.maximum(eesq - sigbsqr, 0.0)) * erf(np.sqrt(x_safe)) - ee) + + +def _deffSigaRoot_cen(eesq, sa): + """`deffSigaRootCen_by_dsa` (lines 107-130).""" + import numpy as np + from scipy.special import erf + sigbsqr = 1.0 - sa * sa + xnum = eesq - sigbsqr + x = 0.5 * xnum / sigbsqr + out = np.empty_like(eesq) + big = np.abs(xnum) > 1e-10 + if big.any(): + x_safe = np.maximum(x[big], 0.0) + out[big] = (sa[big] * erf(np.sqrt(x_safe)) / + np.sqrt(np.maximum(xnum[big], 1e-30)) - + np.exp(-x[big]) * np.sqrt(2.0 * sigbsqr[big] / np.pi) * + sa[big] / sigbsqr[big]) + small = ~big + if small.any(): + out[small] = (xnum[small] * np.sqrt(2.0 / np.pi) * sa[small] / + (3.0 * sigbsqr[small] ** 1.5)) + return out + + +def _d2effSigaRoot_cen(eesq, sa): + """`d2effSigaRootCen_by_dsa2` (lines 132-159).""" + import numpy as np + from scipy.special import erf + sigasqr = sa * sa + sigapow4 = sigasqr * sigasqr + sigbsqr = 1.0 - sigasqr + xnum = eesq - sigbsqr + x = 0.5 * xnum / sigbsqr + sigbsqrtpi = np.sqrt(np.pi * sigbsqr) + out = np.empty_like(eesq) + big = np.abs(xnum) > 1e-10 + if big.any(): + x_safe = np.maximum(x[big], 0.0) + d2num = ((eesq[big] - 1.0) * sigbsqr[big] ** 2 * sigbsqrtpi[big] * + erf(np.sqrt(x_safe))) + exp_part = np.where( + x[big] < 20.0, + np.sqrt(np.maximum(2.0 * xnum[big], 0.0)) * np.exp(-x[big]) * + (1.0 - eesq[big] + sigasqr[big] * + (eesq[big] + eesq[big] ** 2 - 2.0) + sigapow4[big]), + np.zeros_like(x[big]), + ) + d2num = d2num + exp_part + out[big] = d2num / (sigbsqr[big] ** 2 * + np.maximum(xnum[big], 1e-30) ** 1.5 * + sigbsqrtpi[big]) + small = ~big + if small.any(): + out[small] = (np.sqrt(2.0 / np.pi) * sa[small] / + (1.5 * sigbsqr[small] ** 1.5)) + return out + + +def get_dfactor_vectorised(ee_np, eesq_np, centric_np): + """ + Vectorised version of Phaser's `math_RiceLLG.cc:getDfactor` (lines 191-250). + Halley's method with bisection fallback, run over all reflections in + parallel. Each reflection has its own bracket [dflo, dfhi]. + + Returns a (N,) numpy float64 array of DFAC values in (0, 1). + """ + import numpy as np + + EPS1 = 1e-7 + EPS2 = 1e-10 + MAXDFAC = 1.0 - EPS1 + + ee = np.asarray(ee_np, dtype=np.float64) + eesq = np.asarray(eesq_np, dtype=np.float64) + cen = np.asarray(centric_np, dtype=bool) + N = ee.shape[0] + + # Case 1: no observational error (eesq - ee² ≤ 0) → DFAC = 1.0 + out = np.ones(N, dtype=np.float64) + has_err = (eesq - ee * ee) > 0.0 + + if not has_err.any(): + return out + + ee_a, eesq_a, cen_a = ee[has_err], eesq[has_err], cen[has_err] + # Bracket. dflo = max(sqrt(1 - min(eesq, 1)) + EPS, EPS). + dflo = np.maximum(np.sqrt(np.maximum(1.0 - np.minimum(eesq_a, 1.0), 0.0)) + EPS1, EPS1) + dfhi = np.full_like(dflo, MAXDFAC) + + # Early return: if dflo >= MAXDFAC, just use dflo. + early = dflo >= MAXDFAC + if early.any(): + # No iteration for those; keep dflo + pass + + dfmid = 0.5 * (dflo + dfhi) + # Compute initial fmid + fmid = np.empty_like(dfmid) + if cen_a.any(): + fmid[cen_a] = _effSigaRoot_cen(ee_a[cen_a], eesq_a[cen_a], dfmid[cen_a]) + if (~cen_a).any(): + fmid[~cen_a] = _effSigaRoot_acen(ee_a[~cen_a], eesq_a[~cen_a], dfmid[~cen_a]) + + active = ~early # reflections still being iterated + for _ in range(50): + if not active.any(): + break + # Convergence check + conv = (dfhi - dflo) <= EPS1 + conv |= np.abs(fmid) <= EPS2 + active = active & ~conv + if not active.any(): + break + + # slope and curve at dfmid (only for active reflections) + slope = np.empty_like(dfmid) + curve = np.empty_like(dfmid) + cen_act = cen_a & active + acen_act = (~cen_a) & active + if cen_act.any(): + slope[cen_act] = _deffSigaRoot_cen(eesq_a[cen_act], dfmid[cen_act]) + curve[cen_act] = _d2effSigaRoot_cen(eesq_a[cen_act], dfmid[cen_act]) + if acen_act.any(): + slope[acen_act] = _deffSigaRoot_acen(eesq_a[acen_act], dfmid[acen_act]) + curve[acen_act] = _d2effSigaRoot_acen(eesq_a[acen_act], dfmid[acen_act]) + + # Halley step (fall back to `fmid · slope` if curve <= 0, matching + # the Phaser source literally — math_RiceLLG.cc:230-231). + denom_halley = 2.0 * (slope ** 2 - fmid * curve) + use_halley = (curve > 0.0) & (np.abs(denom_halley) > 1e-30) + step = np.where( + use_halley, + 2.0 * fmid * slope / np.where(use_halley, denom_halley, 1.0), + fmid * slope, + ) + dfnew = dfmid - step + # If new value out of bracket, bisect instead + in_bracket = (dfnew > dflo) & (dfnew < dfhi) + dfmid_new = np.where(in_bracket, dfnew, 0.5 * (dflo + dfhi)) + + # Apply only on active reflections; converged ones keep dfmid. + dfmid = np.where(active, dfmid_new, dfmid) + + # Re-evaluate fmid only on active reflections. + if cen_act.any(): + fmid[cen_act] = _effSigaRoot_cen(ee_a[cen_act], eesq_a[cen_act], + dfmid[cen_act]) + if acen_act.any(): + fmid[acen_act] = _effSigaRoot_acen(ee_a[acen_act], eesq_a[acen_act], + dfmid[acen_act]) + + # Update bracket based on fmid sign. + below = (fmid < 0.0) & active + above = (fmid >= 0.0) & active + dflo = np.where(below, dfmid, dflo) + dfhi = np.where(above, dfmid, dfhi) + + out[has_err] = dfmid + return np.clip(out, EPS1, MAXDFAC) + + +def french_wilson_preprocess( + F: torch.Tensor, + sig_F: torch.Tensor, + s_mag: torch.Tensor, + centric: torch.Tensor, + *, + n_wilson_shells: int = 20, +) -> dict: + """ + Phaser-style preprocessing from raw (F, σF, centric) to (eEobs, DFAC). + + Implements the chain documented at the top of the module: + 1. equal-count Wilson shells over `s_mag` + 2. per-shell `_p` (Phaser's `SIGMAN.BINS`) + 3. per-reflection normalised intensity `eosq = F² / ` and σ + `sigesq = σI / ≈ 2·F·σF / ` + 4. French-Wilson posterior `eEFW, eEsqFW` (`math_FrenchWilson.cc`) + 5. DFAC via Halley iteration on Rice moments (`math_RiceLLG.cc`) + 6. `eEobs = sqrt(eEsqFW + (DFAC²−1)/DFAC²)`, clamped to ≤10 + (Phaser `Dfactor.cc:87-93`). + + Returns a dict with torch tensors back on the input device: + eEobs: (N,) effective normalised amplitude + DFAC : (N,) per-reflection D-factor ∈ [1e-7, 1−1e-7] + sqrt_mean_F2: (N,) per-reflection √_p (caller can multiply + back to recover absolute-scale Feff if needed) + """ + import numpy as np + + device = F.device + F_np = F.detach().to("cpu").to(torch.float64).numpy() + sigF_np = sig_F.detach().to("cpu").to(torch.float64).numpy() + s_np = s_mag.detach().to("cpu").to(torch.float64).numpy() + cen_np = centric.detach().to("cpu").bool().numpy() + + # 1+2. Per-shell by equal-count binning over |s|. + sorted_idx = np.argsort(s_np) + edges_idx = np.linspace(0, len(s_np) - 1, n_wilson_shells + 1).round().astype(np.int64) + s_edges = s_np[sorted_idx][edges_idx] + # Nudge endpoints + s_edges[0] -= 1e-6 + s_edges[-1] += 1e-6 + shell_idx = np.clip( + np.searchsorted(s_edges, s_np, side="right") - 1, 0, n_wilson_shells - 1, + ) + F2 = F_np * F_np + mean_F2 = np.zeros(n_wilson_shells, dtype=np.float64) + counts = np.zeros(n_wilson_shells, dtype=np.int64) + np.add.at(mean_F2, shell_idx, F2) + np.add.at(counts, shell_idx, 1) + mean_F2 = mean_F2 / np.maximum(counts, 1) + mean_F2 = np.maximum(mean_F2, 1e-12) + mean_I_per_h = mean_F2[shell_idx] # _p mapped to each h + sqrt_mean_F2 = np.sqrt(mean_I_per_h) + + # 3. Normalised intensity + its sigma. + eosq = F2 / mean_I_per_h + sigesq = 2.0 * F_np * sigF_np / mean_I_per_h + # Guard: sigesq must be positive for FW. If σF=0 we let zero_sig path + # in `_french_wilson_posterior` handle it. + sigesq = np.maximum(sigesq, 0.0) + + # 4. French-Wilson posterior. + eEFW, eEsqFW = _french_wilson_posterior(eosq, sigesq, cen_np) + # Guard non-physical posteriors (numerical accidents): make sure + # eEsqFW ≥ eEFW² (the second moment must dominate the first squared). + bad = eEsqFW < eEFW * eEFW + if bad.any(): + eEsqFW[bad] = eEFW[bad] ** 2 + 1e-12 + + # 5. DFAC via Halley iteration. + DFAC = get_dfactor_vectorised(eEFW, eEsqFW, cen_np) + + # 6. eEobs = sqrt(eEsqFW + (DFAC²-1)/DFAC²), clamp to ≤ 10; recompute + # DFAC if clamped (mirrors `Dfactor.cc:87-93`). + dfsqr = DFAC * DFAC + eEobs_sqr = eEsqFW + (dfsqr - 1.0) / np.maximum(dfsqr, 1e-30) + eEobs_sqr = np.maximum(eEobs_sqr, 0.0) + eEobs = np.sqrt(eEobs_sqr) + # Clamp at 10 and recompute DFAC where clamped (rare). + clamp_mask = (eEobs > 10.0) & (eEsqFW > 1.0) + if clamp_mask.any(): + # eEobs = 10 → solve for DFAC: 100 = eEsqFW + (dfsqr-1)/dfsqr + # → (100 - eEsqFW) = (dfsqr - 1)/dfsqr + # → dfsqr (100 - eEsqFW) = dfsqr - 1 + # → dfsqr (100 - eEsqFW - 1) = -1 + # → dfsqr (eEsqFW - 99) = 1 + eEobs[clamp_mask] = 10.0 + DFAC[clamp_mask] = 1.0 / np.sqrt(np.maximum(eEsqFW[clamp_mask] - 99.0, 1e-30)) + DFAC[clamp_mask] = np.clip(DFAC[clamp_mask], 1e-7, 1.0 - 1e-7) + + return { + "eEobs": torch.from_numpy(eEobs).to(device=device, dtype=F.dtype), + "DFAC": torch.from_numpy(DFAC).to(device=device, dtype=F.dtype), + "sqrt_mean_F2": torch.from_numpy(sqrt_mean_F2).to(device=device, dtype=F.dtype), + } + + +def eterm_sigma_a(s_mag: torch.Tensor, delta_vrms_A: float) -> torch.Tensor: + """ + Phaser's σA Eterm, literal port of `Ensemble.cc:42`: + + Eterm(s) = exp(-(2π²/3) · s² · ΔVRMS_var) + + where `ΔVRMS_var` is the *coordinate variance* in Ų. We accept the RMS + coordinate error `delta_vrms_A` (Å) per the standard σA convention and + square it internally: `ΔVRMS_var = delta_vrms_A²`. + + Per-reflection (`s_mag` may be any shape). Returns the same shape. + + NB: an earlier version of this mimic used `exp(-2π² · s² · ΔVRMS²)` — + that is *3×* too aggressive in the exponent and over-attenuated calc at + high resolution. See `phaser_frf_known_bugs.md` for the audit. + """ + s2 = s_mag * s_mag + return torch.exp(-(2.0 / 3.0) * (math.pi ** 2) * s2 * (delta_vrms_A ** 2)) + + +# ============================================================================= +# Top-level FRF entry point — mirrors `ball_rotation_search` API +# ============================================================================= + + +def phaser_rotation_search( + s_obs: torch.Tensor, + F_obs: torch.Tensor, + centric_obs: torch.Tensor, + s_calc: torch.Tensor, + F_calc: torch.Tensor, + sym_mats: torch.Tensor, + *, + L: int = 24, + d_min: Optional[float] = None, + d_max: Optional[float] = None, + delta_vrms_A: float = 1.0, + n_wilson_shells: int = 20, + n_peaks: int = 500, + refine_subvoxel: bool = True, + n_refine: int = 50, + sigma_threshold: float = -5.0, + bessel_h_scale: Optional[float] = None, + use_lerf1_intensity: bool = True, + use_m_symmetry_filter: bool = True, + sig_F_obs: Optional[torch.Tensor] = None, + use_french_wilson: bool = False, + use_shell_variance_weights: bool = False, + n_var_shells: int = 20, + grid_sampling_deg: float = 3.0, +) -> Tuple[AdaptiveRotationFunction, List[RotationPeak]]: + """ + Phaser-faithful Fast Rotation Function. + + Parameters + ---------- + s_obs : (N_o, 3) real + Observed-side reciprocal-lattice vectors in 1/Å. + F_obs : (N_o,) real + Observed amplitudes (already anisotropy-corrected by caller, if + applicable). + centric_obs : (N_o,) bool + Centric flag for each observed reflection. + s_calc, F_calc : (N_c, 3), (N_c,) real + Model reciprocal vectors and amplitudes. + sym_mats : (n_ops, 3, 3) real + Spacegroup rotation matrices — used to detect the high-order axis + for the m-symmetry filter (Phaser `highOrderAxis()`). + L : int, default 24 + Wigner/SH bandwidth. lmax = L − 1 (rounded down to even). + d_min, d_max : float, optional + Resolution window. Reflections outside `[1/d_max, 1/d_min]` (in |s|) + are discarded on both sides before normalisation. + delta_vrms_A : float, default 1.0 + ΔVRMS in Å for the Luzzati σA Eterm applied to the calc side. + n_wilson_shells : int, default 20 + Number of equal-count shells used in per-shell Wilson normalisation. + n_peaks : int, default 500 + Maximum number of rotation peaks returned. + refine_subvoxel : bool, default True + Apply quadratic sub-voxel refinement to the top `n_refine` peaks. + n_refine : int, default 50 + Number of peaks to refine sub-voxel. + sigma_threshold : float, default -5.0 + Minimum Z-score for a peak to be returned. Negative ≈ "keep everything". + bessel_h_scale : float, optional + Pre-multiplier on |s| in the Bessel argument: `h = bessel_h_scale · |s|`. + Default: `(L - 1) · d_min` (Phaser's `lmax · HIRES`). Picking this + scale puts the highest-u Bessel functions near their first peak at + the maximum |s| in the data, so the radial basis covers the full + resolution range without wasted bandwidth. + use_lerf1_intensity : bool, default True + If True, observed intensity is `cweight · (E_obs² − 1)`; if False, + plain `E_obs² − 1` (cweight = 1 everywhere). + use_m_symmetry_filter : bool, default True + If True, detect ZSYMM from `sym_mats` and apply the m-symmetry + filter to the observed-side SH coefficients. + sig_F_obs : (N_o,) real, optional + Standard deviation of `F_obs`. Required if `use_french_wilson=True`; + ignored otherwise. + use_french_wilson : bool, default False + If True, the observed-side intensity is built via the full Phaser + preprocessing chain: per-shell Wilson normalisation → French-Wilson + posterior → per-reflection Luzzati DFAC → effective Eobs. Observed + intensity becomes `cweight · (eEobs² − 1) · DFAC²` instead of the + simpler `cweight · (E² − 1)`. Requires `sig_F_obs`. + use_shell_variance_weights : bool, default False + If True, compute per-shell empirical variance of `intensity_obs` + and weight each reflection by `1/√Var_p` before the SH expansion. + Analog of torchref's `auto_variance_weights=True`. Downweights + shells whose obs Patterson coefficient is dominated by intermolecular + noise — particularly important for cubic/tetragonal groups where the + spacegroup-invariant SH subspace has few coefficients and per-shell + SNR matters disproportionately. + n_var_shells : int, default 20 + Number of shells for the variance estimate. Same scheme as + `n_wilson_shells`. + grid_sampling_deg : float, default 3.0 + Target Euler-angle resolution in degrees. Drives the per-β α/γ sample + counts via Phaser's `pmax(β) = 720/Δ · cos(β/2)`, + `qmax(β) = 360/Δ · sin(β/2)` (FastRot.cc:92-96). Total adaptive + sample count ≈ `(720 · 360) / grid_sampling_deg²`. + + Returns + ------- + adaptive_rf : AdaptiveRotationFunction + Ragged Euler-grid evaluation of `C(α, β, γ)` with per-β + variable-shape `(qmax_k, pmax_k)` slices, alpha/gamma grids, and the + β midpoint quadrature points. + peaks : list of RotationPeak (sorted by descending score). + """ + device = s_obs.device + real_dtype = s_obs.dtype + + # 1. Resolution mask on both sides. + def _resmask(s_vec, F, centric=None): + smag = s_vec.norm(dim=-1) + lo = 1.0 / d_max if d_max is not None else 0.0 + hi = 1.0 / d_min if d_min is not None else float("inf") + keep = (smag >= lo) & (smag <= hi) + s_vec = s_vec[keep] + F = F[keep] + if centric is not None: + centric = centric[keep] + return s_vec, F, centric, smag[keep] + return s_vec, F, smag[keep] + + # Apply resolution mask to obs side; also mask sig_F if provided. + if use_french_wilson: + if sig_F_obs is None: + raise ValueError("use_french_wilson=True requires sig_F_obs.") + smag_obs_pre = s_obs.norm(dim=-1) + lo = 1.0 / d_max if d_max is not None else 0.0 + hi = 1.0 / d_min if d_min is not None else float("inf") + keep_obs = (smag_obs_pre >= lo) & (smag_obs_pre <= hi) + s_obs = s_obs[keep_obs] + F_obs = F_obs[keep_obs] + sig_F_obs = sig_F_obs[keep_obs] + centric_obs = centric_obs[keep_obs] + smag_obs = smag_obs_pre[keep_obs] + s_calc, F_calc, smag_calc = _resmask(s_calc, F_calc) + else: + s_obs, F_obs, centric_obs, smag_obs = _resmask(s_obs, F_obs, centric_obs) + s_calc, F_calc, smag_calc = _resmask(s_calc, F_calc) + + if s_obs.shape[0] < n_wilson_shells * 5: + raise ValueError( + f"Too few obs reflections ({s_obs.shape[0]}) for " + f"{n_wilson_shells} Wilson shells in [{d_min}, {d_max}] Å." + ) + + # 2. Bessel argument scale. Default = lmax · d_min (Phaser DataMR.cc:1107). + if bessel_h_scale is None: + if d_min is None: + raise ValueError("bessel_h_scale must be set when d_min is None") + lmax = L - 1 + lmax_even = lmax if lmax % 2 == 0 else lmax - 1 + bessel_h_scale = float(lmax_even) * float(d_min) + + # 3. Wilson-normalise both sides (+ FW + DFAC on obs if requested). + if use_french_wilson: + fw = french_wilson_preprocess( + F_obs, sig_F_obs, smag_obs, centric_obs, + n_wilson_shells=n_wilson_shells, + ) + eEobs = fw["eEobs"] + DFAC = fw["DFAC"] + else: + E_obs, _ = _wilson_normalise(F_obs, smag_obs, n_wilson_shells) + eEobs = E_obs + DFAC = torch.ones_like(E_obs) + E_calc, _ = _wilson_normalise(F_calc, smag_calc, n_wilson_shells) + + # 4. Observed intensity (LERF1-style): cweight · (eEobs² − 1) · DFAC². + if use_lerf1_intensity: + cweight = torch.where( + centric_obs.bool(), + torch.ones_like(eEobs), + 2.0 * torch.ones_like(eEobs), + ) + else: + cweight = torch.ones_like(eEobs) + intensity_obs = cweight * (eEobs * eEobs - 1.0) * (DFAC * DFAC) + + # 4b. Optional per-shell empirical-variance weighting (Phaser does this + # via per-shell BINS + best(r) — torchref's `auto_variance_weights=True` + # is the cleanest analog). Downweight shells whose intensity_obs has high + # empirical variance, normalising so the mean weight is ~1. + if use_shell_variance_weights: + from ..sh import compute_patterson_shell_variance + edges_var, _ = equal_count_shell_edges(smag_obs, n_var_shells) + shell_idx_obs = assign_shells(smag_obs, edges_var) + valid_obs = shell_idx_obs >= 0 + var_p = compute_patterson_shell_variance( + intensity_obs[valid_obs].to(torch.float64), + shell_idx_obs[valid_obs], + P=n_var_shells, + ) + inv_sqrt_var = 1.0 / var_p.sqrt().clamp(min=1e-30) + # Normalise mean weight = 1 so the absolute scale of intensity_obs + # doesn't shift (only the SHAPE across shells matters for the + # cross-correlation rotation function). + inv_sqrt_var = (inv_sqrt_var * + (n_var_shells / inv_sqrt_var.sum().clamp(min=1e-30))) + weights_per_h = torch.ones_like(intensity_obs) + weights_per_h[valid_obs] = inv_sqrt_var[shell_idx_obs[valid_obs]].to( + intensity_obs.dtype, + ) + intensity_obs = intensity_obs * weights_per_h + + # 5. Model intensity with σA Eterm² weighting (per-reflection). + eterm = eterm_sigma_a(smag_calc, delta_vrms_A) + intensity_calc = (eterm ** 2) * (E_calc * E_calc - 1.0) + + # 6. Detect ZSYMM (high-order rotation axis from sym_mats). + zsymm = 1 + if use_m_symmetry_filter and sym_mats is not None: + _, zsymm = get_high_order_axis( + sym_mats.to(torch.float64).cpu(), + ) + zsymm = int(zsymm) + + # 7. Bessel-radial × SH expansion of both sides. + c_obs = bessel_sh_expand( + s_obs, intensity_obs.to(real_dtype), + L=L, bessel_h_scale=bessel_h_scale, + zsymm=zsymm, enforce_friedel=True, + ) + c_calc = bessel_sh_expand( + s_calc, intensity_calc.to(real_dtype), + L=L, bessel_h_scale=bessel_h_scale, + zsymm=1, enforce_friedel=True, # NO m-filter on calc side + ) + + # 8. Cross-correlation in Wigner basis: contract on the radial-Bessel axis. + # Convention chosen to match torchref's existing `ball_search.py` (line 182): + # xi[l, m, n] = Σ_r c_obs[r, l, n] · conj(c_calc[r, l, m]) + # i.e. obs-side SH index labels the OUTPUT "n" axis (→ γ Fourier), + # calc-side SH index labels the OUTPUT "m" axis (→ α Fourier). With this + # ordering, the peak Euler triple satisfies `s_calc = R · s_obs` (column + # vector), matching the test-driver scenario where F_calc was generated by + # applying R to the model coordinates. + xi = torch.einsum( + "rln,rlm->lmn", + c_obs.c_nlm, + torch.conj(c_calc.c_nlm), + ) + + # 9. Evaluate on the Phaser-faithful adaptive Euler grid. + arf = evaluate_rotation_function_grid_adaptive( + xi, L, grid_sampling_deg=grid_sampling_deg, n_beta=2 * L, + ) + + # 10. Peak picking + sub-voxel refinement on the ragged grid. + peaks = find_rotation_peaks_adaptive( + arf, n_peaks=n_peaks, sigma_threshold=sigma_threshold, + ) + if refine_subvoxel and peaks: + head = peaks[: min(n_refine, len(peaks))] + head = refine_peaks_subvoxel_adaptive(head, arf) + peaks = head + peaks[len(head):] + peaks.sort(key=lambda r: r.score, reverse=True) + + return arf, peaks diff --git a/torchref/alignment/frf/preprocessing.py b/torchref/alignment/frf/preprocessing.py new file mode 100644 index 00000000..53fe3f5b --- /dev/null +++ b/torchref/alignment/frf/preprocessing.py @@ -0,0 +1,516 @@ +"""Observed-side preprocessing chain. + +Mirrors the chain in Phaser ``DataMR::dataMR_FRF`` (DataMR.cc:863-1133) +and the auxiliary helpers in ``lib/math_FrenchWilson.cc`` and +``lib/math_RiceLLG.cc``. The PyTorch ports of these algorithms already +live in ``torchref.alignment.phaser_frf``; rather than duplicate the +~600 LoC implementations here, we import them and re-document the +Phaser source citations. + +If a specific preprocessing piece turns out to be wrong (per Tier 2 +synthetic tests), the fix lives here — replace the import with a fresh +implementation cited line-by-line to the corresponding Phaser source. +""" +from __future__ import annotations + +import math +from typing import Optional + +import torch + +# Re-exports from the validated legacy implementation. +from .phaser_frf import ( + _wilson_normalise as wilson_normalise, # DataMR.cc:925 + eterm_sigma_a, # Ensemble.cc:42 + french_wilson_preprocess, # math_FrenchWilson.cc + Dfactor.cc +) +from ..sh import ( + get_high_order_axis, # phaser's highOrderAxis() + compute_patterson_shell_variance, + equal_count_shell_edges, + assign_shells, +) + +__all__ = [ + "wilson_normalise", + "wilson_normalise_epsilon", + "compute_epsilon", + "eterm_sigma_a", + "french_wilson_preprocess", + "get_high_order_axis", + "build_lerf1_intensity", + "apply_shell_variance_weights", + "detect_zsymm", + "epsilon_aware_unroll", + "compute_v_budget", + "bulk_solvent_factor", + "oeffner_vrms", + "fit_relative_wilson_b", +] + + +def epsilon_aware_unroll( + hkl_int: torch.Tensor, + sym_mats: torch.Tensor, +): + """Unroll each ASU reflection to the **unique** P1 positions in its orbit. + + For each ``h`` in the input list, generate the orbit ``{S_k · h}`` over the + ``n_ops`` symop matrices and emit one entry per *distinct* position. Axial / + special-position reflections (whose stabilizer has order ε(h) > 1) therefore + appear ``n_ops / ε(h)`` times, **not** ``n_ops`` times. + + Mirrors Phaser's ``if (!duplicate(isym, rhkl))`` skip in + ``DataMR.cc:954-986``. A naive ``einsum + reshape`` unroll over-counts axial + reflections by ε(h), polluting the obs SH coefficients with spurious + non-invariant content — the noise channel that hurts high-symmetry cases. + + Parameters + ---------- + hkl_int : (N, 3) integer tensor + ASU Miller indices. + sym_mats : (n_ops, 3, 3) tensor + Spacegroup rotation operators in the reciprocal (hkl) basis. Cast to + ``long`` internally; values must be integer. + + Returns + ------- + unrolled_hkl : (M, 3) long tensor — flat list of unique orbit positions + across all ASU reflections. + asu_idx : (M,) long tensor — index into ``hkl_int`` that each unrolled entry + came from. Callers use it to broadcast intensities / centric / sigF: + ``F_unrolled = F_obs[asu_idx]``. + """ + hkl_int = hkl_int.to(torch.long) + sym_mats = sym_mats.round().to(torch.long) + N, n_ops = hkl_int.shape[0], sym_mats.shape[0] + # Orbits: (N, n_ops, 3) — S_k applied to each h (row-vector convention, + # matching the existing `einsum("kij,nj->nki", ...)` unroll site). + orbits = torch.einsum("kij,nj->nki", sym_mats, hkl_int) + # Pack (h, k, l) into a single int64 key for per-row dedup. + base = 2 * int(orbits.abs().max().item()) + 1 + key = (orbits[:, :, 0] * base + orbits[:, :, 1]) * base + orbits[:, :, 2] + # Stable sort along dim=1 → duplicates land contiguously, lowest op-index + # first (matches Phaser's "first occurrence wins" rule). + sorted_keys, sort_idx = key.sort(dim=1, stable=True) + first_in_sorted = torch.cat( + [ + torch.ones(N, 1, dtype=torch.bool, device=key.device), + sorted_keys[:, 1:] != sorted_keys[:, :-1], + ], + dim=1, + ) + keep_mask = torch.empty_like(first_in_sorted) + keep_mask.scatter_(1, sort_idx, first_in_sorted) + asu_idx, op_idx = keep_mask.nonzero(as_tuple=True) + unrolled_hkl = orbits[asu_idx, op_idx] + return unrolled_hkl, asu_idx + + +def compute_epsilon( + hkl: torch.Tensor, + sym_mats: torch.Tensor, +) -> torch.Tensor: + """Reflection multiplicity ε(h) — the order of the stabilizer subgroup. + + Phaser source: the ``epsn`` array in ``DataMR.cc`` (used in + ``SIGMAN.sqrt_epsnSN``, DataMR.cc:925) — same role as + ``cctbx::miller::index_span`` epsilons. + + ε(h) = number of point-group rotation operators W for which + ``h · W = h`` (row-vector convention, no Friedel). For a general + reflection ε = 1; reflections on an n-fold symmetry axis get ε = n. + + Used to epsilon-correct Wilson normalisation: axial reflections are + systematically stronger (``⟨I_h⟩ = ε_h · Σ``), so without the + correction they over-weight the m = 0 SH column and bias the + rotation-function map for high-symmetry spacegroups. + + Parameters + ---------- + hkl : (N, 3) integer-valued (any dtype) Miller indices. + sym_mats : (n_ops, 3, 3) integer rotation operators (fractional/lattice + rotation parts of the spacegroup). + + Returns + ------- + epsilon : (N,) float — multiplicity ε ≥ 1. + """ + h = hkl.to(torch.float64) + W = sym_mats.to(torch.float64) + eps = torch.zeros(h.shape[0], dtype=torch.float64, device=h.device) + for k in range(W.shape[0]): + h_t = h @ W[k] # row-vector: h' = h · W + same = (h_t.round() == h).all(dim=-1) + eps += same.to(torch.float64) + return eps.clamp(min=1.0) + + +def wilson_normalise_epsilon( + F: torch.Tensor, + s_mag: torch.Tensor, + epsilon: torch.Tensor, + n_shells: int = 20, +): + """Epsilon-corrected per-shell Wilson normalisation. + + Standard crystallographic normalisation with the multiplicity factor: + + Σ_shell = ⟨I_h / ε_h⟩_shell + E²_h = (I_h / ε_h) / Σ_shell + + so axial reflections (large ε) are not over-counted. Returns + ``(E, sqrt_mean_eps_corrected)`` mirroring ``wilson_normalise``. + """ + edges, _ = equal_count_shell_edges(s_mag, n_shells) + shell_idx = assign_shells(s_mag, edges) + valid = shell_idx >= 0 + I_corr = (F * F) / epsilon.clamp(min=1.0) + count = torch.zeros(n_shells, dtype=torch.int64, device=F.device) + sumI = torch.zeros(n_shells, dtype=F.dtype, device=F.device) + idx_v = shell_idx[valid] + count.index_add_(0, idx_v, torch.ones_like(idx_v)) + sumI.index_add_(0, idx_v, I_corr[valid]) + mean_I = (sumI / count.clamp(min=1).to(F.dtype)).clamp(min=1e-12) + per_h = torch.ones_like(F) + per_h[valid] = mean_I[idx_v] + E = (I_corr / per_h).clamp(min=0.0).sqrt() + return E, per_h.sqrt() + + +def build_lerf1_intensity( + eEobs: torch.Tensor, + centric_obs: torch.Tensor, + dfac: Optional[torch.Tensor] = None, + use_centric_weight: bool = True, +) -> torch.Tensor: + """LERF1 observed intensity: ``cweight · (eEobs² − 1) · DFAC²``. + + Phaser source: ``DataMR::m_LETF1`` (DataMR.cc:1326-1431) — the + intensity that gets fed into the Bessel-SH expansion. cweight is + ε(h) · (1 for centric, 2 for acentric); we use the centric/acentric + factor only (the ε(h) multiplicity is implicit in the symmetry + reduction of the input reflection set). + """ + if use_centric_weight: + cw = torch.where( + centric_obs.bool(), + torch.ones_like(eEobs), + 2.0 * torch.ones_like(eEobs), + ) + else: + cw = torch.ones_like(eEobs) + if dfac is None: + dfac = torch.ones_like(eEobs) + return cw * (eEobs * eEobs - 1.0) * (dfac * dfac) + + +def apply_shell_variance_weights( + intensity: torch.Tensor, + s_mag: torch.Tensor, + n_var_shells: int = 20, +) -> torch.Tensor: + """Per-shell empirical variance reweight. + + Downweights shells whose observed Patterson intensity is dominated + by noise. Mean-normalised so total scale doesn't shift. Closest + Phaser analog is per-shell BINS + ``best(r)`` in ``Ensemble.cc``. + """ + edges, _ = equal_count_shell_edges(s_mag, n_var_shells) + shell_idx = assign_shells(s_mag, edges) + valid = shell_idx >= 0 + var_p = compute_patterson_shell_variance( + intensity[valid].to(torch.float64), + shell_idx[valid], + P=n_var_shells, + ) + inv_sqrt_var = 1.0 / var_p.sqrt().clamp(min=1e-30) + inv_sqrt_var = inv_sqrt_var * ( + n_var_shells / inv_sqrt_var.sum().clamp(min=1e-30) + ) + weights = torch.ones_like(intensity) + weights[valid] = inv_sqrt_var[shell_idx[valid]].to(intensity.dtype) + return intensity * weights + + +def detect_zsymm(sym_mats: Optional[torch.Tensor]) -> int: + """Detect the high-order rotational axis order ``ZSYMM``. + + Phaser source: ``highOrderAxis()`` in ``rotationgroup.h``. Used to + apply the m-symmetry filter to the obs-side SH coefficients: SH + coefficients with ``|m| mod ZSYMM != 0`` average to zero over the + spacegroup and are zeroed out (DataMR.cc:863-870, 1117). + """ + if sym_mats is None: + return 1 + _, zsymm = get_high_order_axis(sym_mats.to(torch.float64).cpu()) + return int(zsymm) + + +def compute_v_budget( + eps_factor: torch.Tensor, + sigma_a: torch.Tensor, + n_mol: int = 1, + totvar_known: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Phaser's per-reflection variance budget ``V(h)`` for the m_LETF1 LL. + + Source: ``DataMR.cc:949`` (build) + ``DataMR.cc:1411`` (use in ``m_LETF1``):: + + V = PTNCS.EPSFAC[r] − totvar_known[r] − totvar_search[r] + + where ``EPSFAC[r] = ε(h)`` (or the tNCS-corrected variance bin in the + NCS-present case; we use plain ε for the standalone search), ``totvar_known`` + is the variance contribution from any fixed model (zero for a pure cross + rotation function), and ``totvar_search = σ_A²(s) · n_mol`` is the variance + explained by the moving model at the expected scattering content. + + For the cross-rotation case (no fixed model, ``totvar_known = 0``): + + V(h) = ε(h) − σ_A²(s)·n_mol + + Working in E-space (obs already Wilson-normalised), so no Σ_N factor. + + Parameters + ---------- + eps_factor : (N,) tensor + Per-reflection ε(h), the multiplicity (1 for general positions, + n>1 for reflections on n-fold symmetry axes). From + :func:`compute_epsilon`. + sigma_a : (N,) tensor + Per-reflection σ_A(s) (interpolated from the per-shell fit). + n_mol : int + Number of molecules in the unit cell summed over by the NSYMP-loop in + the calc-side expected intensity. Equals ``NSYMP`` for the standalone + cross-rotation search. + totvar_known : (N,) tensor, optional + Variance contribution from a fixed/known model. Default ``None`` → + treated as zero (standalone cross-rotation function). + + Returns + ------- + V : (N,) tensor + Per-reflection variance budget. Clamped to ``> 0`` to keep the LL finite; + a non-positive ``V`` would imply σ_A² overshoots ε, which Phaser also + guards against via ``PHASER_ASSERT(C > 0)`` at DataMR.cc:1413. + """ + sa = sigma_a.to(eps_factor.dtype) + moving = (sa * sa) * float(n_mol) + V = eps_factor - moving + if totvar_known is not None: + V = V - totvar_known.to(eps_factor.dtype) + return V.clamp(min=1e-6) + + +# ============================================================================= +# Phaser model-prep — three pieces Phaser applies before the FRF that we don't. +# See `Ensemble::setPDB` (EnsemblePDB.cc:40-100). Adding these as opt-in. +# ============================================================================= + + +def bulk_solvent_factor( + s_mag: torch.Tensor, + fsol: float = 0.95, + bsol: float = 300.0, + sigA_min: float = 0.01, +) -> torch.Tensor: + """Phaser's Babinet bulk-solvent term — ``solTerm.h:9``:: + + solTerm(s²) = max(SIGA_MIN, 1 − fsol · exp(−bsol · s²/4)) + + Models the bulk solvent's contribution to the structure factor via Babinet's + principle. At low resolution (s→0) the term → ``1 − fsol`` ≈ 0.05 (with the + default ``fsol=0.95``), aggressively suppressing the calc — physically, the + model represents only the macromolecule, but the diffraction data sees + macromolecule + bulk solvent, and at low resolution the solvent's flat + average density partially cancels the macromolecule's contribution. At high + resolution (s→∞) the term → 1 (no effect). + + Phaser folds this into the effective σ_A via + ``σ_A_eff(s) = solTerm(s²) · DLuzzati(s², vrms)`` (EnsemblePDB.cc:96-100). + For callers that work in σ_A space (the rescore, the FRF eterm), multiplying + by this factor reproduces that behaviour. + + Defaults match Phaser (``DEF_SOLPAR_BULK_FSOL=0.95``, + ``DEF_SOLPAR_BULK_BSOL=300``, ``DEF_SOLPAR_SIGA_MIN=0.01``). + + Parameters + ---------- + s_mag : tensor + Per-reflection reciprocal-space magnitude |s| (Å^-1). + fsol, bsol, sigA_min : float + Babinet parameters. Defaults match Phaser. + + Returns + ------- + torch.Tensor + Per-reflection solvent multiplier, same shape as ``s_mag``. Always in + ``[sigA_min, 1]``. + """ + s2 = s_mag * s_mag + babinet = 1.0 - float(fsol) * torch.exp(-float(bsol) * s2 / 4.0) + return babinet.clamp(min=float(sigA_min)) + + +def oeffner_vrms(n_residues: int, identity: float = 1.0) -> float: + """Phaser's Oeffner empirical vrms estimate — ``rms_estimate.cc:37``:: + + vrms = A · (B + clamp(n_residues, 125, 1500))^(1/3) · exp(C · (1 − ident)) + + with ``A = 0.0569``, ``B = 173``, ``C = 1.52``. The clamp avoids extrapolating + beyond the well-populated range of Oeffner et al.'s training set + (Acta Cryst. (2013) D69:2209-2215). For a perfect model (``identity=1``) and + a typical protein (~300 residues), this gives vrms ≈ 0.47 Å; large + assemblies (clamped at 1500) give vrms ≈ 0.67 Å. Phaser uses this as the + Luzzati ``vrms`` for the σ_A computation. + + Parameters + ---------- + n_residues : int + Sequence length of the search model. Internally clamped to [125, 1500]. + identity : float, optional + Sequence identity to expected target on [0, 1]. Default 1.0 (perfect). + + Returns + ------- + float + Coordinate RMS estimate in Å, suitable as ``delta_vrms_A`` for + :func:`compute_sigma_a_luzzati` / :func:`eterm_sigma_a`. + """ + A, B, C = 0.0569, 173.0, 1.52 + n_clamped = max(125, min(int(n_residues), 1500)) + return A * (B + n_clamped) ** (1.0 / 3.0) * math.exp(C * (1.0 - float(identity))) + + +def fit_relative_wilson_b( + F_obs: torch.Tensor, + F_calc: torch.Tensor, + s_mag: torch.Tensor, + n_shells: int = 20, + clamp_b: float = 50.0, + s_mag_calc: Optional[torch.Tensor] = None, +) -> float: + """Phaser's relative Wilson-B fit — ``EnsemblePDB.cc:793-851``. + + Estimates the per-model relative Wilson B-factor that brings the calc's + per-shell <|F_calc|²> into agreement with the data's per-shell <|F_obs|²>. + Implemented as a weighted linear regression of + ``log(_shell / _shell)`` against ``s²`` over equal-count + shells; returns ``WilsonB = -2 · slope``. + + Per-shell weighting mirrors Phaser (EnsemblePDB.cc:830-835): + - ``s² < 0.009`` (d > 10.5 Å): weight = 0 (no reliable BEST curve). + - ``s² < 0.04`` (d > 5 Å): weight = ``(s²/0.04)²`` (down-weighted). + - ``s² ≥ 0.04``: weight = 1. + + Apply at call sites as ``F_calc · exp(-WilsonB · s² / 4)``. + + Parameters + ---------- + F_obs : (N_obs,) tensor + F_calc : (N_calc,) tensor + Calc amplitudes. May be on a different reciprocal grid than obs (e.g. + the dense P1-box from ``dense_calc_via_box``); per-shell means handle + the binning independently. + s_mag : (N_obs,) tensor + Reciprocal-space magnitudes for OBS. Used to derive shell edges by + equal-count binning on the obs distribution. + s_mag_calc : (N_calc,) tensor, optional + Reciprocal-space magnitudes for CALC. Defaults to ``s_mag`` (when obs + and calc share a grid). When given, obs and calc are binned into the + SAME edges (derived from obs ``s_mag``) but with independent counts. + n_shells : int + Number of equal-count resolution shells (on obs). + clamp_b : float + Clamps the fitted B to ``[-clamp_b, +clamp_b]``. + + Returns + ------- + float + Relative Wilson B-factor (Ų). ``0`` if too few shells contribute. + """ + if s_mag_calc is None: + s_mag_calc = s_mag + if F_obs.shape[0] != s_mag.shape[0]: + raise ValueError( + f"F_obs / s_mag length mismatch: {F_obs.shape[0]} vs {s_mag.shape[0]}" + ) + if F_calc.shape[0] != s_mag_calc.shape[0]: + raise ValueError( + f"F_calc / s_mag_calc length mismatch: " + f"{F_calc.shape[0]} vs {s_mag_calc.shape[0]}" + ) + + edges, _ = equal_count_shell_edges(s_mag, n_shells) + # Bin obs and calc into the SAME edges (independently — different N's). + shell_idx_obs = assign_shells(s_mag, edges) + shell_idx_calc = assign_shells(s_mag_calc, edges) + valid_obs = shell_idx_obs >= 0 + valid_calc = shell_idx_calc >= 0 + if not valid_obs.any() or not valid_calc.any(): + return 0.0 + + F2_obs = (F_obs * F_obs).to(torch.float64) + F2_calc = (F_calc * F_calc).to(torch.float64) + s2_obs = (s_mag * s_mag).to(torch.float64) + + counts_obs = torch.zeros(n_shells, dtype=torch.int64, device=s_mag.device) + counts_calc = torch.zeros(n_shells, dtype=torch.int64, device=s_mag.device) + sum_F2obs = torch.zeros(n_shells, dtype=torch.float64, device=s_mag.device) + sum_F2calc = torch.zeros(n_shells, dtype=torch.float64, device=s_mag.device) + sum_s2 = torch.zeros(n_shells, dtype=torch.float64, device=s_mag.device) + idx_v_obs = shell_idx_obs[valid_obs] + idx_v_calc = shell_idx_calc[valid_calc] + counts_obs.index_add_(0, idx_v_obs, torch.ones_like(idx_v_obs)) + counts_calc.index_add_(0, idx_v_calc, torch.ones_like(idx_v_calc)) + sum_F2obs.index_add_(0, idx_v_obs, F2_obs[valid_obs]) + sum_F2calc.index_add_(0, idx_v_calc, F2_calc[valid_calc]) + sum_s2.index_add_(0, idx_v_obs, s2_obs[valid_obs]) # obs-side s² for the regression abscissa + + # Drop shells empty on either side. + keep = (counts_obs > 0) & (counts_calc > 0) + mean_F2obs = sum_F2obs[keep] / counts_obs[keep].to(torch.float64) + mean_F2calc = sum_F2calc[keep] / counts_calc[keep].to(torch.float64) + mean_s2 = sum_s2[keep] / counts_obs[keep].to(torch.float64) + + # log(Σ_N / Σ_P) per shell. + eps = 1e-30 + log_ratio = (mean_F2obs.clamp(min=eps) / mean_F2calc.clamp(min=eps)).log() + + # Phaser per-shell weights (EnsemblePDB.cc:830-835). + weights = torch.ones_like(mean_s2) + low_mask = mean_s2 < 0.04 + weights[low_mask] = (mean_s2[low_mask] / 0.04) ** 2 + weights[mean_s2 < 0.009] = 0.0 + if (weights > 0).sum().item() < 2: + return 0.0 + + # Weighted linear fit y = slope · x (no intercept), per the Phaser source. + w = weights + x = mean_s2 + y = log_ratio + sw = w.sum() + swx = (w * x).sum() + swy = (w * y).sum() + swx2 = (w * x * x).sum() + swxy = (w * x * y).sum() + denom = (sw * swx2 - swx * swx).item() + if abs(denom) < 1e-30: + return 0.0 + slope = (sw * swxy - swx * swy).item() / denom + # Phaser: WilsonB_intensity = -4·slope, then halved → WilsonB = -2·slope. + wilson_b = -2.0 * slope + return float(max(-clamp_b, min(wilson_b, clamp_b))) + + +# Note: a naive OLS-on-log-F² fit of anisotropic Wilson U was attempted on +# 2026-05-28 and didn't work. The Wilson left tail (small F values produce huge +# negative log F²) dominates the regression, returning U components of order +# 10²–10³ Ų on real data — three orders of magnitude beyond physical, on both +# easy (1DAW) and hard (2DQ6) cases. Robustifying via |F|-weighting + ridge +# only made the fit saturate any sensible clamp. Phaser's ``scaleANIS`` +# (``DataB.cc``, ~500 LoC) is an iterative ML fit on ``logSigmaEsq``; that's +# the right approach if obs-side aniso ever becomes the next lever. The 2DQ6 +# benchmark failure we were chasing turned out to be tNCS +# (``<(E²−1)²>_acentric = 5.5`` vs Wilson = 1.0), not anisotropy, so this +# branch is not in the immediate critical path. diff --git a/torchref/alignment/frf/sitelist_ang.py b/torchref/alignment/frf/sitelist_ang.py new file mode 100644 index 00000000..0cd0913d --- /dev/null +++ b/torchref/alignment/frf/sitelist_ang.py @@ -0,0 +1,333 @@ +"""Per-β rotation function evaluation on Phaser's adaptive SO(3) sample list. + +Mirrors ``SiteListAng`` from +``reverse_engineering/phenix/phenix-1.20-4459/modules/phaser/codebase/phaser/src/FastRot.cc``. + +The crucial point — and the bug that broke v13 of the legacy adaptive +grid — is that **the FFT itself is NOT per-β-adaptive**. Phaser does: + +1. ``get_FRF`` (FastRot.cc:90-167) loops over a uniform β grid + ``β_b = b · Δ`` for ``b ∈ [0, bmax)``, ``bmax = ceil(180/Δ)``. +2. For each β: ``DoRfftStuff`` (FastRot.cc:19-88) builds the per-β + Fourier-mode amplitudes + ``S_{m1, m2}(β) = Σ_l ξ_{l, m1, m2} · d^l_{m1, m2}(β)`` + on the full ``(2L-1) × (2L-1)`` grid (asymmetric-unit storage only — + the Friedel mate is added by cctbx via ``conjugate_flag=true``). +3. The 2D inverse FFT runs at a **fixed shape** + ``amax = adjust_gridding(2·max(bmax, lmax), max_prime=5)`` for every β. + The result is a dense ``M_β(α, γ)`` map indexed in ``[0, 1)`` along + each axis. +4. The **adaptive sample list** is built once by ``allocate_memory`` + (FastRot.cc:169-262): for each β, + ``pmax(β) = 720/Δ · cos(β/2)`` + ``qmax(β) = 360/Δ · sin(β/2)`` + and the ``(p, q)`` lattice is mapped to ``(α, γ)`` via + ``α = (p/pmax + q/qmax) mod 1`` + ``γ = sign · (p/pmax − q/qmax) mod 1`` (FastRot.cc:216-219) + with the β=0 special case keeping only the ``p == p`` diagonal + (FastRot.cc:189-207) because only ``α + γ`` is meaningful at the pole. +5. ``M_β`` is **bilinearly interpolated** at each ``(α, γ)`` sample point + to give the RF value (FastRot.cc:146-152, ``four_point_interpolation``). + +v13's mistake was making the FFT shape itself ``(pmax(β), qmax(β))`` — +which collapses to ``(N, 1)`` at small β and loses all γ Fourier +information. Phaser keeps the FFT dense; adaptivity is only in the +sample list and the interpolation. +""" +from __future__ import annotations + +import math +from typing import List, Tuple + +import torch + +from .types import AdaptiveRotationFunction +from .wigner_d import wigner_contraction_per_beta + +__all__ = [ + "adjust_gridding", + "build_dense_map_per_beta", + "build_adaptive_sample_list", + "evaluate_rotation_function", +] + + +def adjust_gridding(target: int, max_prime: int = 5) -> int: + """Smallest integer ≥ target whose largest prime factor is ≤ max_prime. + + Phaser source: ``scitbx::fftpack::adjust_gridding`` (FastRot.cc:66-69). + For ``max_prime=5`` this is the standard 5-smooth (Hamming) numbers. + """ + if target <= 1: + return 1 + primes = [2, 3, 5, 7, 11, 13][: min(max(max_prime // 2, 1), 6)] + primes = [p for p in [2, 3, 5, 7, 11, 13] if p <= max_prime] + n = int(target) + while True: + m = n + for p in primes: + while m % p == 0: + m //= p + if m == 1: + return n + n += 1 + + +def build_dense_map_per_beta( + xi_lmn: torch.Tensor, + betas: torch.Tensor, + fft_size: int, +) -> torch.Tensor: + """Return the dense FFT map ``M_β(α, γ)`` for every β. + + Phaser source: ``DoRfftStuff`` (FastRot.cc:19-88), but tensor-batched + over β and with a single 2D ``torch.fft.ifft2`` per β instead of a + cctbx ``real_to_complex_3d`` of shape ``(1, amax, amax)`` — they + produce equivalent dense (α, γ) grids. + + Parameters + ---------- + xi_lmn : torch.Tensor (complex), shape (L, 2L-1, 2L-1) + Cross-correlation coefficients with l ∈ [0, L), m, n ∈ [-(L-1), L-1]. + betas : torch.Tensor (real), shape (n_beta,) + β values in radians. + fft_size : int + Fixed FFT grid size N. The map ``M_β`` will be ``(N, N)`` for every β, + indexed as ``M[k', l'] = RF(2π k'/N, β, 2π l'/N) / N²``. + + Returns + ------- + M : torch.Tensor (complex), shape (n_beta, fft_size, fft_size) + The dense maps. Use bilinear interpolation in the (α, γ) plane to + evaluate at non-grid points. + """ + L = xi_lmn.shape[0] + n_beta = betas.shape[0] + device = xi_lmn.device + + # 1. Per-β Wigner-d contraction: S[k, m1+L-1, m2+L-1] = Σ_l ξ d^l_{m1,m2}(β_k). + S = wigner_contraction_per_beta(xi_lmn, betas) # (n_beta, 2L-1, 2L-1) + + # 2. Place S into the (fft_size, fft_size) Fourier grid by FFT-frequency + # mapping: m → (m mod N), n → (n mod N). For m, n ∈ [-(L-1), L-1] and + # N >> 2L-1, this puts negative frequencies at the high end of each axis. + if fft_size < 2 * L - 1: + raise ValueError( + f"fft_size={fft_size} must be >= 2L-1={2*L-1} to avoid aliasing" + ) + pad = torch.zeros( + (n_beta, fft_size, fft_size), dtype=S.dtype, device=device, + ) + m_vals = torch.arange(-(L - 1), L, device=device) + idx = (m_vals % fft_size).to(torch.int64) + pad[:, idx.unsqueeze(1), idx.unsqueeze(0)] = S + + # 3. Forward 2D FFT — torch convention: + # fft2(X)[k, l] = Σ_{m, n} X[m, n] · exp(-2πi (m·k/N + n·l/M)) + # which gives M[k, l] = RF(α = 2π·k/N, β, γ = 2π·l/N) directly, with + # Edmonds D^l_{m,n} = exp(-imα) d^l_{m,n}(β) exp(-inγ). Using ifft2 + # here (the cctbx default in Phaser's pipeline) would give RF at + # (-α, -γ) which then requires negating alpha_frac/gamma_frac when + # recording sample Euler angles — Phaser's FastRot.cc:153 does this + # explicit ``-360 * alpha`` flip. We do the equivalent by using fft2 + # so the Euler labels are already in the right sign. + M = torch.fft.fft2(pad, dim=(-2, -1)) + return M + + +def _build_beta_grid(grid_sampling_deg: float) -> Tuple[torch.Tensor, int]: + """Return (β_grid in radians, bmax). β = b · Δ for b ∈ [0, bmax).""" + bmax = int(math.ceil(180.0 / grid_sampling_deg)) + if bmax < 1: + raise ValueError(f"grid_sampling_deg={grid_sampling_deg} too coarse") + b = torch.arange(bmax, dtype=torch.float64) + betas_rad = b * grid_sampling_deg * (math.pi / 180.0) + return betas_rad, bmax + + +def build_adaptive_sample_list( + grid_sampling_deg: float, + dtype: torch.dtype = torch.float64, + device: torch.device = torch.device("cpu"), +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Build the per-β (α, γ) sample list. + + Phaser source: ``SiteListAng::allocate_memory`` (FastRot.cc:169-262). + Returns (in radians): + alphas : (N_samples,) α value per sample + betas : (N_samples,) β value per sample + gammas : (N_samples,) γ value per sample + beta_starts: (bmax + 1,) int64 slice [beta_starts[b]:beta_starts[b+1]] + is the samples at β = b · Δ + beta_grid : (bmax,) the β values in radians + """ + betas_rad, bmax = _build_beta_grid(grid_sampling_deg) + betas_rad = betas_rad.to(device=device, dtype=dtype) + + alphas_list: List[torch.Tensor] = [] + gammas_list: List[torch.Tensor] = [] + betas_list: List[torch.Tensor] = [] + beta_starts: List[int] = [0] + + for b in range(bmax): + beta_rad = float(betas_rad[b].item()) + cosb = math.cos(beta_rad / 2.0) + sinb = math.sin(beta_rad / 2.0) + pmax = max(1, int(720.0 / grid_sampling_deg * cosb)) + qmax = max(1, int(360.0 / grid_sampling_deg * sinb)) + + if b == 0: + # β=0: only α = γ = p/pmax for p < pmax/2 (FastRot.cc:189-207). + p_idx = torch.arange(pmax, device=device) + p_ratio = p_idx.to(torch.float64) / pmax + keep = p_ratio < 0.5 + p_ratio = p_ratio[keep] + alpha_frac = p_ratio + gamma_frac = p_ratio + else: + # (p, q) ∈ [0, pmax) × [0, qmax), mapped to (α, γ) via + # FastRot.cc:216-219 — with the negative branch for γ when + # p_ratio < q_ratio (gives γ ∈ [0, 1) without negative values). + p_idx = torch.arange(pmax, device=device) + q_idx = torch.arange(qmax, device=device) + p_ratio = (p_idx.to(torch.float64) / pmax).unsqueeze(1) # (pmax, 1) + q_ratio = (q_idx.to(torch.float64) / qmax).unsqueeze(0) # (1, qmax) + alpha_frac = torch.fmod(p_ratio + q_ratio, 1.0) + diff = p_ratio - q_ratio + gamma_frac = torch.where( + diff >= 0.0, + torch.fmod(diff, 1.0), + 1.0 - torch.fmod(-diff, 1.0), + ) + alpha_frac = alpha_frac.reshape(-1) + gamma_frac = gamma_frac.reshape(-1) + + # Dedup: when p ≥ pmax/2, the (p, q) → (α, γ) map can collide + # with (p − pmax/2, q') for some q' (FastRot.cc:222-241). Phaser + # does an O(pmax·qmax²) loop; we do it via tuple-set in numpy + # which is O(N log N) — same result. + keys = torch.stack([alpha_frac, gamma_frac], dim=-1) + # Round to 1e-6 to match the epsilon comparison in the source. + key_round = (keys * 1_000_000).round().to(torch.int64) + _, uniq_idx = torch.unique(key_round, dim=0, return_inverse=True) + # Keep first occurrence of each unique key. + seen = {} + keep_mask = torch.zeros(alpha_frac.shape[0], dtype=torch.bool, device=device) + for i, k in enumerate(uniq_idx.tolist()): + if k not in seen: + seen[k] = True + keep_mask[i] = True + alpha_frac = alpha_frac[keep_mask] + gamma_frac = gamma_frac[keep_mask] + + n_this = alpha_frac.shape[0] + alphas_list.append((alpha_frac * (2.0 * math.pi)).to(dtype)) + gammas_list.append((gamma_frac * (2.0 * math.pi)).to(dtype)) + betas_list.append(torch.full((n_this,), beta_rad, dtype=dtype, device=device)) + beta_starts.append(beta_starts[-1] + n_this) + + alphas = torch.cat(alphas_list) + gammas = torch.cat(gammas_list) + betas_flat = torch.cat(betas_list) + beta_starts_t = torch.tensor(beta_starts, dtype=torch.int64, device=device) + return alphas, betas_flat, gammas, beta_starts_t, betas_rad + + +def _bilinear_interp_periodic( + M: torch.Tensor, # (N, N) complex + alpha_frac: torch.Tensor, # (n,) in [0, 1) + gamma_frac: torch.Tensor, # (n,) in [0, 1) +) -> torch.Tensor: + """Periodic bilinear interpolation of M at (alpha_frac · N, gamma_frac · N). + + Mirrors ``four_point_interpolation`` (FastRot.cc:146) on a periodic + map: both axes wrap around modulo N. + """ + N = M.shape[-1] + af = (alpha_frac % 1.0) * N + gf = (gamma_frac % 1.0) * N + a0 = torch.floor(af).to(torch.int64) % N + g0 = torch.floor(gf).to(torch.int64) % N + a1 = (a0 + 1) % N + g1 = (g0 + 1) % N + da = (af - torch.floor(af)).to(M.real.dtype) + dg = (gf - torch.floor(gf)).to(M.real.dtype) + da_c = da.to(M.dtype) + dg_c = dg.to(M.dtype) + v00 = M[a0, g0] + v01 = M[a0, g1] + v10 = M[a1, g0] + v11 = M[a1, g1] + return ( + v00 * ((1 - da_c) * (1 - dg_c)) + + v01 * ((1 - da_c) * dg_c) + + v10 * (da_c * (1 - dg_c)) + + v11 * (da_c * dg_c) + ) + + +def evaluate_rotation_function( + xi_lmn: torch.Tensor, + grid_sampling_deg: float = 2.0, + fft_size: int = -1, +) -> AdaptiveRotationFunction: + """Compute the rotation function on Phaser's adaptive SO(3) sample list. + + Phaser source: composition of ``get_FRF`` + ``allocate_memory`` + + ``four_point_interpolation`` (FastRot.cc:90-262). Returns + real-valued samples (the rotation function is real). + + Parameters + ---------- + xi_lmn : torch.Tensor (complex), shape (L, 2L-1, 2L-1) + grid_sampling_deg : float + Phaser's ``grid_sampling`` keyword. β grid is uniform at this + spacing; (α, γ) sample density per β follows pmax/qmax. + fft_size : int, optional + Dense FFT shape. Default: ``adjust_gridding(2·max(bmax, 2L-1), 5)``. + """ + if xi_lmn.ndim != 3: + raise ValueError(f"xi_lmn must be 3-D (L, 2L-1, 2L-1), got {tuple(xi_lmn.shape)}") + L = xi_lmn.shape[0] + device = xi_lmn.device + real_dtype = ( + torch.float64 + if xi_lmn.dtype in (torch.complex128, torch.float64) + else torch.float32 + ) + + bmax = int(math.ceil(180.0 / grid_sampling_deg)) + if fft_size < 0: + fft_size = adjust_gridding(2 * max(bmax, 2 * L - 1), max_prime=5) + + # 1. Build adaptive sample list (purely geometric — independent of xi). + alphas, betas_flat, gammas, beta_starts, beta_grid = build_adaptive_sample_list( + grid_sampling_deg, dtype=real_dtype, device=device, + ) + + # 2. Dense FFT map per β. + M = build_dense_map_per_beta(xi_lmn, beta_grid, fft_size) # (n_beta, N, N) + + # 3. Bilinear interp at each sample's (α, γ). + values = torch.zeros(alphas.shape[0], dtype=real_dtype, device=device) + n_beta = beta_grid.shape[0] + for b in range(n_beta): + i0, i1 = int(beta_starts[b].item()), int(beta_starts[b + 1].item()) + if i1 <= i0: + continue + af = alphas[i0:i1] / (2.0 * math.pi) + gf = gammas[i0:i1] / (2.0 * math.pi) + v_complex = _bilinear_interp_periodic(M[b], af, gf) + # The rotation function is real; the imaginary residue is at the + # numerical-noise level for a Hermitian-symmetric input ξ. Drop it. + values[i0:i1] = v_complex.real + + return AdaptiveRotationFunction( + alphas=alphas, + betas=betas_flat, + gammas=gammas, + values=values, + beta_starts=beta_starts, + beta_grid=beta_grid, + grid_sampling_deg=grid_sampling_deg, + ) diff --git a/torchref/alignment/frf/spherical_y.py b/torchref/alignment/frf/spherical_y.py new file mode 100644 index 00000000..926768bd --- /dev/null +++ b/torchref/alignment/frf/spherical_y.py @@ -0,0 +1,96 @@ +"""Spherical harmonics ``Y_l^m(θ, φ)`` table. + +Phaser source: ``phaser/lib/sphericalY.h`` (Condon-Shortley phase, the +standard physics convention). + +We build the table once per (θ, φ) batch and reuse across all +``(l, m)`` indices needed. +""" +from __future__ import annotations + +import math + +import torch + + +def _normalised_associated_legendre( + cos_theta: torch.Tensor, L: int +) -> torch.Tensor: + """Compute the *normalised* associated Legendre polynomials. + + Returns ``P[l, m, ...] = sqrt((2l+1)/(4π) · (l-m)!/(l+m)!) · P_l^m(cos θ)`` + for l ∈ [0, L), m ∈ [0, l], with the standard recurrence (no Condon-Shortley + sign — that's applied at the Y_lm level). + + Out-of-range entries (m > l) are zero. + """ + device = cos_theta.device + x = cos_theta.to(torch.float64) + sin_t = torch.sqrt(torch.clamp(1.0 - x * x, min=0.0)) + + plm = torch.zeros((L, L, *x.shape), dtype=torch.float64, device=device) + # m = 0 sector: standard Legendre recurrence with normalisation built in. + plm[0, 0] = math.sqrt(1.0 / (4.0 * math.pi)) + if L > 1: + plm[1, 0] = math.sqrt(3.0 / (4.0 * math.pi)) * x + for l in range(2, L): + a = math.sqrt((2 * l + 1) * (2 * l - 1)) / l + b = math.sqrt((2 * l + 1) / (2 * l - 3)) * (l - 1) / l + plm[l, 0] = a * x * plm[l - 1, 0] - b * plm[l - 2, 0] + + # m > 0 sector: build P_l^l from the previous diagonal, then recur up in l. + for m in range(1, L): + plm[m, m] = -math.sqrt((2 * m + 1) / (2 * m)) * sin_t * plm[m - 1, m - 1] + if m + 1 < L: + plm[m + 1, m] = math.sqrt(2 * m + 3) * x * plm[m, m] + for l in range(m + 2, L): + a = math.sqrt((2 * l + 1) * (2 * l - 1) / ((l - m) * (l + m))) + b = math.sqrt( + (2 * l + 1) * (l + m - 1) * (l - m - 1) + / ((l - m) * (l + m) * (2 * l - 3)) + ) + plm[l, m] = a * x * plm[l - 1, m] - b * plm[l - 2, m] + + return plm + + +def ylm_table( + theta: torch.Tensor, phi: torch.Tensor, L: int +) -> torch.Tensor: + """Compute ``Y_l^m(θ, φ)`` for all l ∈ [0, L), m ∈ [-l, l]. + + Convention: Condon-Shortley phase (the ``(-1)^m`` factor is in + ``Y_l^m`` for m > 0). The normalised associated Legendre polynomial + is real; the φ dependence is ``exp(i m φ)``. + + Parameters + ---------- + theta, phi : torch.Tensor (real, same shape) + Polar (θ ∈ [0, π]) and azimuthal (φ ∈ [0, 2π)) angles. + + Returns + ------- + Y : torch.Tensor (complex128), shape (L, 2L-1, *theta.shape) + ``Y[l, m + L - 1, ...] = Y_l^m(θ, φ)`` for |m| ≤ l, else 0. + """ + if theta.shape != phi.shape: + raise ValueError( + f"theta {tuple(theta.shape)} and phi {tuple(phi.shape)} must agree" + ) + device = theta.device + plm = _normalised_associated_legendre(torch.cos(theta), L) # (L, L, *) + Y = torch.zeros( + (L, 2 * L - 1, *theta.shape), dtype=torch.complex128, device=device + ) + phi64 = phi.to(torch.float64) + # m = 0 + Y[:, L - 1, ...] = plm[:, 0, ...].to(torch.complex128) + for m in range(1, L): + e_pos = torch.complex(torch.cos(m * phi64), torch.sin(m * phi64)) + e_neg = e_pos.conj() + # Y_l^{+m} = (-1)^m * sqrt(...) P_l^m(cosθ) * exp(i m φ), already + # includes the (-1)^m if we apply it here. + sign = (-1) ** m + Y[:, L - 1 + m, ...] = sign * plm[:, m, ...].to(torch.complex128) * e_pos + Y[:, L - 1 - m, ...] = plm[:, m, ...].to(torch.complex128) * e_neg + return Y diff --git a/torchref/alignment/frf/types.py b/torchref/alignment/frf/types.py new file mode 100644 index 00000000..9497ed60 --- /dev/null +++ b/torchref/alignment/frf/types.py @@ -0,0 +1,93 @@ +"""Dataclasses shared across frf_separate. + +Mirrors the small "data carrier" structs in Phaser +(phenix-1.20-4459/modules/phaser/codebase/phaser/src/SiteListAng.h, + src/DataMR.h) without trying to keep the same names everywhere — +crystallographic intent first, C++ naming second. +""" +from __future__ import annotations + +from dataclasses import dataclass +from typing import List, Tuple + +import torch + + +@dataclass +class BesselSHCoefficients: + """Bessel-radial × spherical-harmonic coefficients ``c_{n, l, m}``. + + Phaser source: DataMR.cc constructs the equivalent ``c_lmn`` 3D complex + grid inside ``DataMR::dataMR_FRF`` (not shown here — see annotations). + + Shape: ``coeffs[n, l, m + L - 1]`` is complex, with l ∈ [0, L) even-only, + n ∈ [0, N_radial), m ∈ [-l, l] (zero-padded outside). + """ + + coeffs: torch.Tensor # (N_radial, L, 2L-1), complex + L: int # angular bandlimit (lmax = L - 1) + N_radial: int # number of radial nodes + bessel_h_scale: float # h = bessel_h_scale * |s| + + +@dataclass +class WignerContraction: + """The per-β ``S_{m1,m2}(β) = Σ_l ξ_{l,m1,m2} · d^l_{m1,m2}(β)``. + + Phaser source: ``SiteListAng::DoRfftStuff`` (FastRot.cc:39-59) builds + this sum into the ``rot`` accumulator before pushing into a pseudo-SF + array for FFT. + """ + + S: torch.Tensor # (n_beta, 2L-1, 2L-1), complex + betas: torch.Tensor # (n_beta,), real + + +@dataclass +class AdaptiveRotationFunction: + """Rotation function values on Phaser's adaptive SO(3) sample list. + + The samples live on a *non-rectangular* point set per β — generated by + the ``allocate_memory`` routine in ``SiteListAng::allocate_memory`` + (FastRot.cc:169-262). Each β has a different number of samples; we + store them concatenated, with ``beta_starts`` indexing into the flat + arrays. + + The dense FFT map ``M_β(α, γ)`` (fixed shape across β) is *not* stored + here — only the interpolated samples on the adaptive list. + """ + + alphas: torch.Tensor # (N_samples,) real in [0, 2π) + betas: torch.Tensor # (N_samples,) real in [0, π] + gammas: torch.Tensor # (N_samples,) real in [0, 2π) + values: torch.Tensor # (N_samples,) real — RF values + beta_starts: torch.Tensor # (n_beta + 1,) int — slice indices per β + beta_grid: torch.Tensor # (n_beta,) real — the β values themselves + grid_sampling_deg: float + + def total_samples(self) -> int: + return int(self.values.numel()) + + +@dataclass +class RotationPeak: + """A single rotation function peak. + + Identical to torchref.alignment.ball_search.RotationPeak so consumers + don't need to change. Convention: Edmonds ZYZ Euler in radians. + """ + + alpha: float + beta: float + gamma: float + value: float + sigma: float + + @property + def score(self) -> float: + """Alias for ``value`` so this peak is drop-in for + ``ball_search.RotationPeak`` (whose primary field is ``score``). + Lets the production pipeline/align conversion + ``(p.alpha, p.beta, p.gamma, p.score, p.sigma)`` work for both engines. + """ + return self.value diff --git a/torchref/alignment/frf/wigner_d.py b/torchref/alignment/frf/wigner_d.py new file mode 100644 index 00000000..9e22e082 --- /dev/null +++ b/torchref/alignment/frf/wigner_d.py @@ -0,0 +1,125 @@ +"""Wigner small-d matrices and Wigner-D pointwise evaluation. + +Phaser source: ``phaser/lib/wigner.h`` (the C++ template +``djmn_recursive_table`` used in ``FastRot.cc:41`` per-l, per-β). + +Phaser uses the Sakurai recurrence convention; we use the equivalent +Edmonds (4.1.23) direct-sum formula, already validated against Phaser's +output by the convention tests in +``tests/unit/alignment/test_wigner.py``. To avoid duplicating maths, +this module re-exports the existing implementation from +``torchref.alignment.wigner`` (which is the same convention) and adds +Phaser-specific helpers on top. +""" +from __future__ import annotations + +from typing import Tuple + +import torch + +# Re-export the existing validated implementations. +from ..wigner import ( + _build_half_angle_pow_tables, + small_d_block, + small_d_packed, + small_d_table, + wigner_D_pointwise, +) + +__all__ = [ + "small_d_block", + "small_d_packed", + "small_d_stable", + "small_d_table", + "wigner_D_pointwise", + "wigner_contraction_per_beta", +] + + +def small_d_stable(L: int, betas: torch.Tensor) -> torch.Tensor: + """Numerically stable Wigner small-d table via J_y diagonalization. + + ``d^l(β) = exp(-iβ J_y^{(l)})``. J_y is a tiny real tridiagonal generator, + so ``eigh(i·J_y)`` gives integer eigenvalues μ ∈ [-l, l] and a basis V with + ``d^l_{m n}(β) = Σ_μ V_{m μ} e^{-iβ μ} V*_{n μ}`` (real). + This is exactly the π/2 / SOFT Fourier-over-μ decomposition (V are the π/2 + matrices up to phase), and it has NO catastrophic cancellation — unlike the + Edmonds direct-sum ``small_d_packed`` which explodes to |d|~1e11 at l≥50. + + Returns ``(n_beta, L, 2L-1, 2L-1)`` with ``d[k, l, m+L-1, n+L-1] = d^l_{m,n}(β_k)``, + matching ``small_d_packed`` exactly (same convention, verified at l≤40). + """ + betas = betas.to(torch.float64) + n_beta = betas.shape[0] + dim = 2 * L - 1 + device = betas.device + out = torch.zeros((n_beta, L, dim, dim), dtype=torch.float64, device=device) + out[:, 0, L - 1, L - 1] = 1.0 # l=0 + for l in range(1, L): + sz = 2 * l + 1 + p = torch.arange(sz - 1, dtype=torch.float64, device=device) + sup = 0.5 * torch.sqrt((2 * l - p) * (p + 1.0)) # J_y off-diagonal magnitudes + A = torch.diag(sup, 1) - torch.diag(sup, -1) # A = -i J_y, real antisymmetric + H = 1j * A.to(torch.complex128) # Hermitian + w, V = torch.linalg.eigh(H) # w≈[-l..l], V complex + # d_l(β) = Re( V · diag(e^{-iβ w}) · V^H ), batched over β + phase = torch.exp(-1j * betas.unsqueeze(1) * w.unsqueeze(0)) # (n_beta, sz) + d_l = torch.einsum("ma,ka,na->kmn", V, phase, V.conj()).real # (n_beta, sz, sz) + lo, hi = L - 1 - l, L - 1 + l + 1 + out[:, l, lo:hi, lo:hi] = d_l + return out + + +def wigner_contraction_per_beta( + xi_lmn: torch.Tensor, + betas: torch.Tensor, +) -> torch.Tensor: + """Compute ``S_{m1, m2}(β) = Σ_l ξ_{l, m1, m2} · d^l_{m1, m2}(β)``. + + Phaser source: ``SiteListAng::DoRfftStuff`` (FastRot.cc:39-59) — the + inner ``for (l_index, m1_index, m2_index)`` triple-loop. We do it + in one tensor contraction instead of a Python loop over l. + + Parameters + ---------- + xi_lmn : torch.Tensor (complex), shape (L, 2L-1, 2L-1) + SH-Bessel coefficients with l ∈ [0, L), |m|, |n| ≤ L-1, + zero-padded outside |m| > l or |n| > l. (Already n-summed over + the Bessel radial index by the caller.) + betas : torch.Tensor (real), shape (n_beta,) + β values to evaluate at, in radians. + + Returns + ------- + S : torch.Tensor (complex), shape (n_beta, 2L-1, 2L-1) + ``S[k, m1+L-1, m2+L-1] = Σ_l ξ_{l, m1, m2} · d^l_{m1, m2}(β_k)``. + """ + if xi_lmn.ndim != 3: + raise ValueError(f"xi_lmn must be 3-D, got shape {tuple(xi_lmn.shape)}") + L = xi_lmn.shape[0] + dim = 2 * L - 1 + device = xi_lmn.device + betas = betas.to(torch.float64) + n_beta = betas.shape[0] + xi = xi_lmn.to(torch.complex128) + + # Fused per-l loop: compute each d^l(β) block via J_y eigendecomposition + # (small_d_stable's method, stable to any l) and contract it into S + # immediately. Never materialises the full (n_beta, L, 2L-1, 2L-1) table + # (~19 GB at L=100) nor a 4-D einsum intermediate — peak memory is one + # (n_beta, 2l+1, 2l+1) block (~0.4 GB at l=99). + S = torch.zeros((n_beta, dim, dim), dtype=torch.complex128, device=device) + c = L - 1 + S[:, c, c] += xi[0, c, c] # l=0: d^0 = 1 + for l in range(1, L): + sz = 2 * l + 1 + p = torch.arange(sz - 1, dtype=torch.float64, device=device) + sup = 0.5 * torch.sqrt((2 * l - p) * (p + 1.0)) + A = torch.diag(sup, 1) - torch.diag(sup, -1) # A = -i J_y + w, V = torch.linalg.eigh(1j * A.to(torch.complex128)) # w∈[-l..l] + phase = torch.exp(-1j * betas.unsqueeze(1) * w.unsqueeze(0)) # (n_beta, sz) + VP = V.unsqueeze(0) * phase.unsqueeze(1) # (n_beta, sz, sz) = (k,m,a) + d_l = (VP @ V.conj().transpose(-1, -2)).real # (n_beta, sz, sz) + lo, hi = c - l, c + l + 1 + S[:, lo:hi, lo:hi] += xi[l, lo:hi, lo:hi].unsqueeze(0) * d_l.to(torch.complex128) + return S diff --git a/torchref/alignment/lattman_love.py b/torchref/alignment/lattman_love.py index 8d7577a0..342e823f 100644 --- a/torchref/alignment/lattman_love.py +++ b/torchref/alignment/lattman_love.py @@ -225,3 +225,82 @@ def evaluate( interp = interpolate_complex_from_grid(self.reciprocal_grid, flat) out = interp.reshape(B, N) return out if batched else out.squeeze(0) + + +def estimate_interp_var( + interpolator: "LattmanLoveInterpolator", + hkl_real: torch.Tensor, + real_cell: Cell, + shell_idx: torch.Tensor, + n_shells: int, + n_jitter: int = 4, + jitter_frac: float = 0.5, + seed: int = 0, +) -> torch.Tensor: + """ + Estimate per-reflection trilinear-interpolation variance in E-value units. + + Phaser's totvar_search analogue. Inflates the Rice/Woolfson variance budget + so that interpolation noise in the search model doesn't make a slightly- + noisy true peak look worse than a noise-free wrong peak. + + Method: evaluate the interpolator at the original HKLs and at `n_jitter` + sub-grid-cell perturbations of the HKLs, take the per-shell empirical + variance of |F| across the perturbations, and normalise by the per-shell + mean |F|² so the returned quantity adds correctly to ``(1 - D²)`` in the + Rice variance. + + `jitter_frac` is the fraction of a cubic-cell grid spacing to jitter by; + 0.5 sweeps half a Nyquist cell and gives a robust upper-bound estimate + of trilinear bias. n_jitter=4 keeps the cost negligible. + + Returns + ------- + interp_var : torch.Tensor, shape (N,) + Per-reflection interpolation variance in dimensionless E² units. + """ + device = interpolator.device + dtype = torch.float32 + R_eye = torch.eye(3, dtype=dtype, device=device) + + F_ref = interpolator.evaluate( + R_eye, hkl_real, real_cell, return_amplitude=True, + ).to(dtype) # (N,) + + # Map a Cartesian Å^-1 shift back to fractional HKL_real space. delta_s is + # the magnitude of the jitter in Cartesian reciprocal Å^-1. + delta_s = jitter_frac / float(interpolator.cubic_side) + rec_real_inv = torch.linalg.inv( + real_cell.reciprocal_basis_matrix.to(device).to(dtype), + ) + + g = torch.Generator(device="cpu").manual_seed(int(seed)) + diffs_sq = torch.zeros_like(F_ref) + for _ in range(n_jitter): + direction = torch.randn(3, generator=g, dtype=torch.float64) + direction = (direction / direction.norm()).to(device).to(dtype) + delta_h_real = (delta_s * direction) @ rec_real_inv # (3,) + hkl_j = hkl_real.to(dtype) + delta_h_real + F_j = interpolator.evaluate( + R_eye, hkl_j, real_cell, return_amplitude=True, + ).to(dtype) + diffs_sq = diffs_sq + (F_j - F_ref) ** 2 + diffs_sq = diffs_sq / max(n_jitter, 1) # (N,) + + # Per-shell aggregation. Cast to f64 for stable sums on large N. + shell_idx_l = shell_idx.to(device).long() + diffs_d = diffs_sq.to(torch.float64) + F_ref2_d = (F_ref.to(torch.float64)) ** 2 + + var_per_shell = torch.zeros(n_shells, dtype=torch.float64, device=device) + F2_per_shell = torch.zeros(n_shells, dtype=torch.float64, device=device) + cnt = torch.zeros(n_shells, dtype=torch.float64, device=device) + var_per_shell.scatter_add_(0, shell_idx_l, diffs_d) + F2_per_shell.scatter_add_(0, shell_idx_l, F_ref2_d) + cnt.scatter_add_(0, shell_idx_l, torch.ones_like(diffs_d)) + + mean_var = var_per_shell / cnt.clamp(min=1.0) # (n_shells,) + mean_F2 = (F2_per_shell / cnt.clamp(min=1.0)).clamp(min=1e-30) # (n_shells,) + interp_var_E_per_shell = (mean_var / mean_F2).clamp(min=0.0, max=1.0) + + return interp_var_E_per_shell.to(dtype).index_select(0, shell_idx_l) # (N,) diff --git a/torchref/alignment/ml_rotation.py b/torchref/alignment/ml_rotation.py index 3a522ca7..42e6242e 100644 --- a/torchref/alignment/ml_rotation.py +++ b/torchref/alignment/ml_rotation.py @@ -24,13 +24,18 @@ import torch -from .ball_search import ( +from .frf.ball_search import ( RotationPeak, edmonds_euler_from_rotation_matrix, rotation_matrix_from_edmonds_euler, rotation_matrix_from_edmonds_euler_batch, ) -from .distributions import rice_log_likelihood, woolfson_log_likelihood +from .distributions import ( + phaser_log_rel_rice, + phaser_log_rel_woolfson, + rice_log_likelihood, + woolfson_log_likelihood, +) from .lattman_love import LattmanLoveInterpolator @@ -61,13 +66,23 @@ def _equal_count_shell_idx(s_mag: torch.Tensor, n_shells: int) -> torch.Tensor: def _normalize_to_e(F: torch.Tensor, shell_idx: torch.Tensor, n_shells: int) -> torch.Tensor: """E = F / sqrt( per shell). Vectorised across shells via scatter.""" + return F / _per_shell_sqrt_mean(F, shell_idx, n_shells) + + +def _per_shell_sqrt_mean(F: torch.Tensor, shell_idx: torch.Tensor, + n_shells: int) -> torch.Tensor: + """Per-reflection ``sqrt(_shell)``. Wilson-normalisation denominator. + + Use this when you need to normalise *another* tensor by the same per-shell + statistic computed from F — e.g. converting rotated |F_calc| to E_calc + using the reference |F_calc|'s shell means (rotation-invariant). + """ F2 = F ** 2 sum_per_shell = torch.zeros(n_shells, dtype=F2.dtype, device=F.device) sum_per_shell.scatter_add_(0, shell_idx, F2) count_per_shell = torch.bincount(shell_idx, minlength=n_shells).to(F2.dtype) mean_per_shell = (sum_per_shell / count_per_shell.clamp(min=1.0)).clamp(min=1e-30) - norm_per_refl = mean_per_shell.sqrt().index_select(0, shell_idx) - return F / norm_per_refl + return mean_per_shell.sqrt().index_select(0, shell_idx) def _shell_ll( @@ -75,14 +90,23 @@ def _shell_ll( E_calc: torch.Tensor, centric: torch.Tensor, D: float, + interp_var: Optional[torch.Tensor] = None, ) -> torch.Tensor: """ Per-reflection log-likelihood at a given σA = D for one shell, in E-value space. - Acentric: Rice with F_mean = D · E_calc, variance = 1 − D². - Centric: Woolfson with F_mean = D · E_calc, variance = 1 − D². + Acentric: Rice with F_mean = D · E_calc, variance = (1 − D²) + interp_var. + Centric: Woolfson with F_mean = D · E_calc, variance = (1 − D²) + interp_var. + + `interp_var` (Phaser totvar_search analogue) inflates the variance to absorb + interpolation / model error and prevents the Rice tail from over-penalising + slightly-noisy true peaks. """ - var = torch.full_like(E_obs, max(1.0 - D * D, 1e-4)) + base = max(1.0 - D * D, 1e-4) + if interp_var is None: + var = torch.full_like(E_obs, base) + else: + var = (interp_var + base).clamp(min=1e-4) F_mean = D * E_calc ll = torch.where( centric, @@ -139,6 +163,84 @@ def _optimize_D_in_shell( return 0.5 * (lo + hi) +def compute_sigma_a_luzzati( + s_mag: torch.Tensor, + delta_vrms_A: float = 1.0, +) -> torch.Tensor: + """ + Phaser-style Luzzati σA(s) = exp(−2π²·s²·ΔVRMS²). + + Closed-form, rotation-independent estimate of the per-reflection (or + per-shell) σA, derived from the search model's RMS coordinate + deviation ΔVRMS (in Å). At s=0 returns 1.0 (perfect agreement); + falls off monotonically with resolution. Matches the Phaser FastRot + Eterm/Vterm weighting (LERF1 §2.1.2): `Eterm = exp(−2π²s²ΔVRMS)` and + `Vterm = Eterm²`. + + Parameters + ---------- + s_mag : torch.Tensor + Reciprocal magnitudes in Å⁻¹ (any shape). + delta_vrms_A : float + Estimated RMS coordinate error of the search model, Å. Default + 1.0 Å is a reasonable starting point for MR search models; + tune via `frf_delta_vrms_A` kwarg in `align_model_to_data`. + + Returns + ------- + sigma_a : torch.Tensor, same shape and dtype as `s_mag`. + """ + return torch.exp( + -2.0 * (math.pi ** 2) * (s_mag ** 2) * (float(delta_vrms_A) ** 2) + ) + + +def fit_sigma_a_per_shell( + E_obs: torch.Tensor, + E_calc: torch.Tensor, + centric: torch.Tensor, + shell_idx: torch.Tensor, + n_shells: int, + n_grid: int = 81, + interp_var: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """ + Vectorised per-shell σA = D fit. Single source of truth for D across the + alignment stages (rotation rescore + likelihood TF). + + For each shell, scans D ∈ [0, 0.99] on a fine grid and returns the + grid maximum. With n_grid=81 the resolution is ~0.012, comparable to the + golden-section result in `_optimize_D_in_shell` for downstream LLG purposes. + + Returns + ------- + sigma_a : torch.Tensor, shape (n_shells,) + """ + device = E_obs.device + dtype = E_obs.dtype + N = E_obs.numel() + D_grid = torch.linspace(0.0, 0.99, n_grid, device=device, dtype=dtype) # (G,) + F_mean = D_grid.view(-1, 1) * E_calc.view(1, -1) # (G, N) + var_d = (1.0 - D_grid * D_grid).clamp(min=1e-4) # (G,) + if interp_var is None: + var_full = var_d.view(-1, 1).expand(n_grid, N) + else: + var_full = (var_d.view(-1, 1) + interp_var.view(1, -1)).clamp(min=1e-4) + E_obs_full = E_obs.view(1, -1).expand(n_grid, N) + + ll_acent = rice_log_likelihood(E_obs_full, F_mean, var_full) + ll_cent = woolfson_log_likelihood(E_obs_full, F_mean, var_full) + cent_full = centric.view(1, -1) + ll = torch.where(cent_full, ll_cent, ll_acent) # (G, N) + + # Sum per shell, take argmax over the D-grid. + shell_idx_gn = shell_idx.view(1, -1).expand(n_grid, N) + ll_per_shell = torch.zeros((n_grid, n_shells), dtype=dtype, device=device) + ll_per_shell.scatter_add_(1, shell_idx_gn, ll) # (G, n_shells) + best_idx = ll_per_shell.argmax(dim=0) # (n_shells,) + return D_grid[best_idx] + + # ============================================================================= # Public API # ============================================================================= @@ -175,6 +277,8 @@ def llg_for_rotation_batch( F_calc: torch.Tensor, n_D_grid: int = 41, shell_weights: Optional[torch.Tensor] = None, + interp_var: Optional[torch.Tensor] = None, + sigma_a: Optional[torch.Tensor] = None, ) -> torch.Tensor: """ Vectorized log-likelihood gain across a batch of candidate rotations. @@ -200,6 +304,15 @@ def llg_for_rotation_batch( (`w_p = 1/√Var(E_obs²-1)_p`). The weight multiplies *both* the Sim and Wilson LL contributions uniformly per shell, so the LL gain interpretation is preserved. + interp_var : torch.Tensor, shape (N,), optional + Per-reflection interpolation variance (Phaser totvar_search analogue). + Added to the model variance term. None ⇒ original Rice/Woolfson. + sigma_a : torch.Tensor, shape (n_shells,), optional + Externally-fitted per-shell σA to reuse across candidates. When given, + the per-(D, B) grid maximisation is bypassed and the LLG is computed + at this fixed sigma_a (one D per shell, broadcast per reflection). + This is the "shared D" path used when an external single-source σA + is available (see ``fit_sigma_a_per_shell``). Returns ------- @@ -225,26 +338,49 @@ def llg_for_rotation_batch( norm_per_refl_b = mean_per_shell_b.sqrt().gather(1, shell_idx_b) # (B, N) E_calc = F_calc / norm_per_refl_b - # --- Joint (D, B, N) likelihood evaluation --- - # Memory: D · B · N · 8 B. For default args (D=41, B≤100, N≈3 k) this is - # ~100 MB, comparable to what the per-shell loop already built per - # iteration. For dense-R we typically pass n_D_grid=11, so cost is small. - D_grid = torch.linspace(0.0, 0.99, n_D_grid, device=device, dtype=dtype) - F_mean = D_grid.view(-1, 1, 1) * E_calc.unsqueeze(0) # (D, B, N) - var_d = (1.0 - D_grid * D_grid).clamp(min=1e-4) - var_full = var_d.view(-1, 1, 1).expand(n_D_grid, B, N) - E_obs_full = E_obs.view(1, 1, -1).expand(n_D_grid, B, N) + if sigma_a is not None: + # --- Shared per-shell σA path: skip the D-grid, evaluate LL once. --- + D_per_refl = sigma_a.to(dtype).to(device).index_select(0, shell_idx) # (N,) + var_d = (1.0 - D_per_refl * D_per_refl).clamp(min=1e-4) # (N,) + if interp_var is not None: + var_per_refl = (var_d + interp_var.to(dtype).to(device)).clamp(min=1e-4) + else: + var_per_refl = var_d + F_mean = D_per_refl.view(1, N) * E_calc # (B, N) + var_full = var_per_refl.view(1, N).expand(B, N) + E_obs_full = E_obs.view(1, N).expand(B, N) + ll_acent = rice_log_likelihood(E_obs_full, F_mean, var_full) + ll_cent = woolfson_log_likelihood(E_obs_full, F_mean, var_full) + cent_full = centric.view(1, N) + ll = torch.where(cent_full, ll_cent, ll_acent) # (B, N) + ll_per_shell = torch.zeros((B, n_shells), dtype=dtype, device=device) + ll_per_shell.scatter_add_(1, shell_idx_b, ll) + ll_sim_per_shell = ll_per_shell # (B, n_shells) + else: + # --- Joint (D, B, N) likelihood evaluation, original behaviour. --- + # Memory: D · B · N · 8 B. For default args (D=41, B≤100, N≈3 k) this is + # ~100 MB, comparable to what the per-shell loop already built per + # iteration. For dense-R we typically pass n_D_grid=11, so cost is small. + D_grid = torch.linspace(0.0, 0.99, n_D_grid, device=device, dtype=dtype) + F_mean = D_grid.view(-1, 1, 1) * E_calc.unsqueeze(0) # (D, B, N) + var_d = (1.0 - D_grid * D_grid).clamp(min=1e-4) + if interp_var is None: + var_full = var_d.view(-1, 1, 1).expand(n_D_grid, B, N) + else: + iv = interp_var.to(dtype).to(device).view(1, 1, N) + var_full = (var_d.view(-1, 1, 1) + iv).clamp(min=1e-4).expand(n_D_grid, B, N) + E_obs_full = E_obs.view(1, 1, -1).expand(n_D_grid, B, N) - ll_acent = rice_log_likelihood(E_obs_full, F_mean, var_full) - ll_cent = woolfson_log_likelihood(E_obs_full, F_mean, var_full) - cent_full = centric.view(1, 1, -1) - ll = torch.where(cent_full, ll_cent, ll_acent) # (D, B, N) + ll_acent = rice_log_likelihood(E_obs_full, F_mean, var_full) + ll_cent = woolfson_log_likelihood(E_obs_full, F_mean, var_full) + cent_full = centric.view(1, 1, -1) + ll = torch.where(cent_full, ll_cent, ll_acent) # (D, B, N) - # --- Sum per shell across N, max over D, sum weighted across shells --- - shell_idx_dbn = shell_idx.view(1, 1, -1).expand(n_D_grid, B, N) - ll_per_shell = torch.zeros((n_D_grid, B, n_shells), dtype=dtype, device=device) - ll_per_shell.scatter_add_(2, shell_idx_dbn, ll) - ll_sim_per_shell, _ = ll_per_shell.max(dim=0) # (B, n_shells) + # --- Sum per shell across N, max over D, sum weighted across shells --- + shell_idx_dbn = shell_idx.view(1, 1, -1).expand(n_D_grid, B, N) + ll_per_shell = torch.zeros((n_D_grid, B, n_shells), dtype=dtype, device=device) + ll_per_shell.scatter_add_(2, shell_idx_dbn, ll) + ll_sim_per_shell, _ = ll_per_shell.max(dim=0) # (B, n_shells) # Wilson reference at D = 0 (data-only): F_mean = 0, var = 1. var0 = torch.ones_like(E_obs) @@ -277,6 +413,8 @@ def sim_mlrf_rescore( shell_weights: Optional[torch.Tensor] = None, auto_variance_weights: bool = True, n_D_grid: int = 41, + interp_var: Optional[torch.Tensor] = None, + sigma_a: Optional[torch.Tensor] = None, ) -> List[RotationPeak]: """ Rescore a list of peaks from `ball_rotation_search` by the per-shell-fitted @@ -348,6 +486,7 @@ def sim_mlrf_rescore( F_obs=F_obs, shell_idx=shell_idx, n_shells=n_shells, E_obs=E_obs, centric=centric, F_calc=F_calc, shell_weights=shell_weights, n_D_grid=n_D_grid, + interp_var=interp_var, sigma_a=sigma_a, ) llg_chunks.append(llg_batch) if verbose > 1: @@ -374,6 +513,255 @@ def sim_mlrf_rescore( return rescored + tail +# ============================================================================= +# Phaser-faithful m_LETF1 rescore: NSYMP calc sum + V(h) budget + Rice/Woolfson +# logRel formulas (DataMR.cc:1326-1429). +# ============================================================================= + + +def m_letf1_rescore( + peaks: List[RotationPeak], + F_obs: torch.Tensor, + hkl_real: torch.Tensor, + s_mag: torch.Tensor, + centric: torch.Tensor, + interpolator: LattmanLoveInterpolator, + real_cell, + sym_mats: torch.Tensor, + *, + n_shells: int = 20, + n_refine: Optional[int] = None, + batch_size: int = 50, + sigma_a: Optional[torch.Tensor] = None, + eps_factor: Optional[torch.Tensor] = None, + verbose: int = 0, + # --- Phaser model-prep knobs (all default OFF; see frf/preprocessing.py) --- + apply_bulk_solvent: bool = False, + solvent_fsol: float = 0.95, + solvent_bsol: float = 300.0, + vrms_strategy: str = "fixed", # "fixed" (legacy delta_vrms=0.5) or "oeffner" + vrms_n_residues: Optional[int] = None, # required if vrms_strategy="oeffner" + vrms_identity: float = 1.0, + apply_wilson_b: bool = False, + wilson_b_value: Optional[float] = None, # if None and apply_wilson_b=True, fitted from data +) -> List[RotationPeak]: + """Phaser-faithful ``m_LETF1`` rescore (DataMR.cc:1326-1429). + + Upgrades over :func:`sim_mlrf_rescore`: + + 1. **NSYMP symmetry sum on calc** — for each obs reflection ``h``, the + expected moving-model intensity is + ``eImove(h) = Σ_isym σ_A²(s) · |F_calc(R^T · S_isym · h)|²`` + summed over the ``NSYMP`` spacegroup rotation operators + (DataMR.cc:1371-1404). Implemented vectorised: pre-compute the orbit + ``hkl_unroll`` of shape ``(N, n_ops, 3)``, flatten to + ``(N·n_ops, 3)``, evaluate the LL interpolator once per orientation, + view back, square + sum over the symop dim. + + 2. **Per-reflection variance budget** ``V(h) = ε(h) − σ_A²(s)·n_mol`` from + :func:`torchref.alignment.frf.preprocessing.compute_v_budget` + (DataMR.cc:949,1411). For cross-rotation with no fixed model. + + 3. **Phaser ``logRelRice`` / ``logRelWoolfson``** as the per-reflection LL + formula (RiceWoolfson.cc:25-74), commensurable with Phaser's m_LETF1 + output. Different normalisation from our generic + :func:`rice_log_likelihood` (factor of 2 in the Bessel argument; ``V`` is + twice the standard Rice variance for acentric). + + Returns peaks ranked by descending LL with ``score = LL`` and + ``sigma = (LL − μ_batch) / σ_batch``, drop-in for downstream consumers. + + Parameters + ---------- + peaks + Candidate orientations from the FRF, ZYZ Edmonds Euler. + F_obs, hkl_real, s_mag, centric + Per-reflection obs arrays (anisotropy-corrected F_obs is fine). + interpolator, real_cell + ``LattmanLoveInterpolator`` for the model molecular transform and the + crystal real cell. + sym_mats : (n_ops, 3, 3) tensor + Spacegroup rotation operators in the reciprocal (hkl) basis. + sigma_a : (N,) tensor, optional + Per-reflection σ_A. If ``None``, fitted on-the-fly from the identity + rotation's |F_calc| via :func:`fit_sigma_a_per_shell` and interpolated + per shell. + eps_factor : (N,) tensor, optional + Per-reflection multiplicity ε(h). If ``None``, computed via + :func:`torchref.alignment.frf.preprocessing.compute_epsilon`. + n_refine, batch_size, verbose + As in :func:`sim_mlrf_rescore`. + """ + if not peaks: + return [] + if n_refine is None: + n_refine = len(peaks) + head = peaks[:n_refine] + tail = peaks[n_refine:] + + device = F_obs.device + dtype = F_obs.dtype + n_ops = int(sym_mats.shape[0]) + N = hkl_real.shape[0] + + # 1. Per-shell binning + Wilson-normalised E_obs (E in Phaser notation). + shell_idx = _equal_count_shell_idx(s_mag, n_shells) + E_obs = _normalize_to_e(F_obs, shell_idx, n_shells) + + # 2. ε(h) per reflection. + if eps_factor is None: + from .frf.preprocessing import compute_epsilon + eps_factor = compute_epsilon(hkl_real.to(torch.long), sym_mats).to(dtype) + eps_factor = eps_factor.to(device) + + # 3. σ_A per reflection — Luzzati formula from a coordinate-error parameter + # ``delta_vrms_A`` (default 0.5 Å, matching the validated FRF v19 config). + # This is the Phaser-faithful choice: Phaser's σ_A per shell is pre-fit + # via a Wilson/Luzzati-style formula that depends only on (s, ΔVRMS) and + # does NOT require an aligned model. Data-fit alternatives + # (`fit_sigma_a_per_shell`) need a meaningful obs-calc alignment, which + # we don't have a priori — at any misaligned reference (identity or + # even the top FRF peak when truth is buried) the fit returns ~0 and + # the LL becomes orientation-blind (the 4BX9 rank-481 / 2DQ6 rank-235 + # failure modes in v22/v23). + # + # Always need the calc shell-scale for E-normalisation; use identity + # (rotation-invariant — sphere permutation, shell sums preserved). + I_eye = torch.eye(3, dtype=torch.float32, device=device) + F_calc_ref = interpolator.evaluate( + I_eye, hkl_real, real_cell, return_amplitude=True, + ).to(dtype).squeeze(0) # (N,) + sqrt_mean_F2_calc_per_h = _per_shell_sqrt_mean(F_calc_ref, shell_idx, n_shells).to(device) + + # Optional Wilson-B match (EnsemblePDB.cc:793-851). Compute once from the + # identity-rotation F_calc reference (rotation-invariant shell statistic), + # apply as Debye-Waller multiplier `exp(-B·s²/4)` to F_calc inside the batch. + if apply_wilson_b and wilson_b_value is None: + from .frf.preprocessing import fit_relative_wilson_b + wilson_b_value = fit_relative_wilson_b( + F_obs, F_calc_ref, s_mag, n_shells=n_shells, + ) + wilson_b_value = float(wilson_b_value or 0.0) + if apply_wilson_b and abs(wilson_b_value) > 1e-6: + dw = torch.exp(-wilson_b_value * (s_mag * s_mag) / 4.0).to(dtype).to(device) + else: + dw = None # skip the elementwise mul if a no-op + + if sigma_a is None: + # σ_A: Phaser-faithful Luzzati with optional Oeffner vrms + bulk solvent. + if vrms_strategy == "oeffner": + if vrms_n_residues is None: + raise ValueError( + "vrms_strategy='oeffner' requires vrms_n_residues=." + ) + from .frf.preprocessing import oeffner_vrms + delta_vrms_A = oeffner_vrms(int(vrms_n_residues), float(vrms_identity)) + elif vrms_strategy == "fixed": + delta_vrms_A = 0.5 # legacy default + else: + raise ValueError( + f"vrms_strategy={vrms_strategy!r}; expected 'fixed' or 'oeffner'." + ) + sigma_a = compute_sigma_a_luzzati(s_mag, delta_vrms_A=delta_vrms_A).to(dtype).to(device) + if apply_bulk_solvent: + from .frf.preprocessing import bulk_solvent_factor + sol = bulk_solvent_factor( + s_mag, fsol=solvent_fsol, bsol=solvent_bsol, + ).to(dtype).to(device) + sigma_a = sigma_a * sol + sigma_a = sigma_a.to(device) + sigma_a2 = sigma_a * sigma_a + + # 4. V(h) — rotation-independent variance budget. Phaser-faithful per + # DataMR.cc:1342-1345: `thisV = scatFactor · NSYMP · DFAC² · σ_A²`. + # With `scatFactor = 1/NSYMP` (single ensemble, fracMove=1) this collapses + # to `thisV = σ_A²`, so V = ε − σ_A² — **independent of NSYMP**. Using + # `n_mol=n_ops` (as the initial implementation did) drove V negative on + # high-symmetry spacegroups (4BX9 NSYMP=8, V_clamp blew up the LL). + from .frf.preprocessing import compute_v_budget + V = compute_v_budget(eps_factor, sigma_a, n_mol=1) # (N,) + + # 5. Orbit hkl pre-compute (rotation-independent; reused per batch). + sym_mats_f = sym_mats.to(torch.float64).to(device) + hkl_f = hkl_real.to(torch.float64).to(device) + # (N, n_ops, 3) — for each obs h, all n_ops sym-equivalents in hkl space. + hkl_unroll = torch.einsum("kij,nj->nki", sym_mats_f, hkl_f) + hkl_flat = hkl_unroll.reshape(-1, 3) # (N·n_ops, 3) + + # 6. Build candidate rotation matrices. Same convention as sim_mlrf_rescore: + # transpose the Edmonds Euler matrix because the peak encodes "rotation + # applied to model coords"; we need the inverse for evaluating + # F_calc(R^T · h). + alpha_t = torch.tensor([p.alpha for p in head], dtype=torch.float64) + beta_t = torch.tensor([p.beta for p in head], dtype=torch.float64) + gamma_t = torch.tensor([p.gamma for p in head], dtype=torch.float64) + R_all = rotation_matrix_from_edmonds_euler_batch( + alpha_t, beta_t, gamma_t, + ).transpose(-1, -2).to(torch.float32) # (M, 3, 3) + M = R_all.shape[0] + + # 7. Batched score: eImove via NSYMP-summed calc, LL via Phaser logRel. + E_obs_b = E_obs.unsqueeze(0) # (1, N) — broadcasts over batch dim + V_b = V.unsqueeze(0) # (1, N) + # eImove pre-factor per Phaser DataMR.cc:1397: `thisEsqr *= repsn * scatFactor` + # → `eImove = ε(h) · σ_A² · (1/n_ops) · Σ_k |E_calc(S_k h)|²`. We collapse the + # `1/n_ops` into the pre-factor so the batch loop just does the raw sum. + eImove_prefac = (eps_factor * sigma_a2 / float(n_ops)).unsqueeze(0) # (1, N) + centric_b = centric.to(torch.bool).unsqueeze(0) + + # Normaliser to convert rotated |F_calc| → |E_calc| (per-shell sqrt mean + # from the identity-rotation reference; broadcasts over symops via the + # last unsqueeze since rotation preserves |h| → same shell for all symops). + sqrt_mean_b = sqrt_mean_F2_calc_per_h.unsqueeze(0).unsqueeze(-1) # (1, N, 1) + + llg_chunks: List[torch.Tensor] = [] + for start in range(0, M, batch_size): + stop = min(start + batch_size, M) + R_batch = R_all[start:stop] # (B, 3, 3) + F_calc_flat = interpolator.evaluate( + R_batch, hkl_flat, real_cell, return_amplitude=True, + ).to(dtype) # (B, N·n_ops) + F_calc = F_calc_flat.view(-1, N, n_ops) # (B, N, n_ops) + # Optional Wilson-B Debye-Waller multiplier (per-reflection, broadcasts + # over batch + symops). Applied PRE-normalisation so the per-shell sqrt + # mean (computed from un-DW'd identity F_calc_ref) stays the right scale + # for normalisation; the DW shifts E_calc relative to that scale, which + # is exactly what Wilson-B matching is supposed to do. + if dw is not None: + F_calc = F_calc * dw.unsqueeze(0).unsqueeze(-1) + # Normalise to E_calc (same Wilson scale as E_obs); Rice/Woolfson LL + # only makes sense with both sides on the same per-shell scale. + E_calc = F_calc / sqrt_mean_b + # eImove(h) = ε(h) · σ_A²(s) · (1/n_ops) · Σ_isym |E_calc(R^T·S_isym·h)|² + # (Phaser DataMR.cc:1371-1404; scatFactor = 1/NSYMP folds the symop sum + # into a mean, ε(h) is the multiplicity factor `repsn`). + eImove = eImove_prefac * (E_calc * E_calc).sum(dim=-1) # (B, N) + sqrt_eImove = eImove.clamp(min=1e-30).sqrt() # (B, N) + ll_acen = phaser_log_rel_rice(E_obs_b, sqrt_eImove, V_b) # (B, N) + ll_cen = phaser_log_rel_woolfson(E_obs_b, sqrt_eImove, V_b) # (B, N) + ll = torch.where(centric_b, ll_cen, ll_acen) # (B, N) + llg_chunks.append(ll.sum(dim=-1)) # (B,) + if verbose > 1: + print(f" m_LETF1 batch {start}-{stop}/{M}", flush=True) + + llgs_t = torch.cat(llg_chunks) + mean_t = llgs_t.mean() + std_t = llgs_t.std().clamp(min=1e-30) + sigmas_t = (llgs_t - mean_t) / std_t + + llgs_list = llgs_t.tolist() + sigmas_list = sigmas_t.tolist() + rescored = [ + RotationPeak( + alpha=p.alpha, beta=p.beta, gamma=p.gamma, + score=llg, sigma=sigma, + ) + for p, llg, sigma in zip(head, llgs_list, sigmas_list) + ] + rescored.sort(key=lambda r: r.score, reverse=True) + return rescored + tail + + # ============================================================================= # Brute-force ML rotation search over a uniform SO(3) sample # ============================================================================= @@ -432,3 +820,128 @@ def brute_ml_rotation_search( n_shells=n_shells, n_refine=n_candidates, batch_size=batch_size, verbose=verbose, ) + + +# ============================================================================= +# BRF — Brute Rotation Function (Phaser stage that refines top FRF peaks via a +# denser local rotation sampling + full Rice/Woolfson LL). +# ============================================================================= + + +def _random_rotation_in_cone( + n: int, radius_rad: float, generator: torch.Generator, +) -> torch.Tensor: + """``n`` random rotation matrices uniformly inside an angular cone of given + radius from identity. Axis ∼ uniform-on-sphere, angle ∼ uniform on + ``[0, radius_rad]`` (proper Haar measure on the cone would weight angle by + sin²(θ/2); for small radii this approximation is essentially uniform). + + Returns ``(n, 3, 3)`` float64 tensor. + """ + # Random unit axes via Gaussian normalization. + axes = torch.randn(n, 3, generator=generator, dtype=torch.float64) + axes = axes / axes.norm(dim=-1, keepdim=True).clamp(min=1e-30) + angles = torch.rand(n, generator=generator, dtype=torch.float64) * radius_rad + # Rodrigues: R = I + sin(θ)·K + (1−cos(θ))·K² with K skew of axis. + cos = angles.cos().view(-1, 1, 1) + sin = angles.sin().view(-1, 1, 1) + K = torch.zeros(n, 3, 3, dtype=torch.float64) + K[:, 0, 1] = -axes[:, 2]; K[:, 0, 2] = axes[:, 1] + K[:, 1, 0] = axes[:, 2]; K[:, 1, 2] = -axes[:, 0] + K[:, 2, 0] = -axes[:, 1]; K[:, 2, 1] = axes[:, 0] + I = torch.eye(3, dtype=torch.float64).expand(n, 3, 3) + K2 = K @ K + return I + sin * K + (1.0 - cos) * K2 + + +def brf_refine( + peaks: List[RotationPeak], + F_obs: torch.Tensor, + hkl_real: torch.Tensor, + s_mag: torch.Tensor, + centric: torch.Tensor, + interpolator: LattmanLoveInterpolator, + real_cell, + sym_mats: torch.Tensor, + *, + n_top: int = 100, + n_perturb: int = 10, + angular_radius_deg: float = 3.0, + seed: int = 42, + verbose: int = 0, + **m_letf1_kwargs, +) -> List[RotationPeak]: + """Brute Rotation Function — Phaser's BRF stage (post-FRF rotation refinement). + + For each of the top ``n_top`` FRF peaks, sample ``n_perturb`` random + rotations within an angular cone of radius ``angular_radius_deg`` from the + peak; score all (original + perturbations) via :func:`m_letf1_rescore`. + Returns the rescored set sorted by descending LL — the truth orientation + typically lurks ≤ FRF-grid-spacing degrees from the FRF peak nearest to it, + so this local refinement can recover peaks the coarse FRF grid missed. + + Total LL evaluations: ``n_top · (n_perturb + 1)``. With defaults this is + ``100 × 11 = 1100`` — ~2× the cost of the standard top-500 rescore. + + Parameters + ---------- + peaks + FRF peak list (sorted by descending FRF score). + n_top + Refine around the top ``n_top`` FRF peaks. Should be ≥ the expected + rank of the truth peak — for hard cases (e.g., 2DQ6 FRF rank ≈50), use + ``n_top ≥ 100``. + n_perturb + Number of random rotations sampled per peak (in addition to the + peak itself, which is always included). + angular_radius_deg + Cone half-angle for perturbations. Match the FRF grid resolution + (~3° at ``grid_sampling_deg=3.0``). + seed + RNG seed for reproducibility. + **m_letf1_kwargs + Forwarded to :func:`m_letf1_rescore` — e.g. ``apply_bulk_solvent=True``, + ``vrms_strategy="oeffner"``, ``vrms_n_residues=...``, etc. + """ + if not peaks: + return [] + n_top = min(n_top, len(peaks)) + top = peaks[:n_top] + g = torch.Generator().manual_seed(int(seed)) + radius_rad = math.radians(float(angular_radius_deg)) + + # One big batch of perturbations: (n_top * n_perturb, 3, 3). + pert_R = _random_rotation_in_cone(n_top * n_perturb, radius_rad, g) + pert_R = pert_R.view(n_top, n_perturb, 3, 3) + + out_peaks: List[RotationPeak] = [] + for i, peak in enumerate(top): + # Include the original peak (perturbation theta=0). + out_peaks.append(RotationPeak( + alpha=peak.alpha, beta=peak.beta, gamma=peak.gamma, + score=peak.score, sigma=peak.sigma, + )) + R_orig = rotation_matrix_from_edmonds_euler( + peak.alpha, peak.beta, peak.gamma, + ).to(torch.float64) + for j in range(n_perturb): + R_pert = R_orig @ pert_R[i, j] + a, b, c = edmonds_euler_from_rotation_matrix(R_pert) + out_peaks.append(RotationPeak( + alpha=a, beta=b, gamma=c, score=peak.score, sigma=peak.sigma, + )) + + if verbose > 0: + print( + f" BRF: {len(out_peaks)} candidates " + f"({n_top} top × ({n_perturb}+1) perturbations, ±{angular_radius_deg}°)", + flush=True, + ) + + return m_letf1_rescore( + out_peaks, F_obs, hkl_real, s_mag, centric, interpolator, real_cell, + sym_mats, + n_refine=len(out_peaks), + verbose=verbose, + **m_letf1_kwargs, + ) diff --git a/torchref/alignment/pipeline.py b/torchref/alignment/pipeline.py index 1b91c367..db078957 100644 --- a/torchref/alignment/pipeline.py +++ b/torchref/alignment/pipeline.py @@ -15,18 +15,12 @@ import numpy as np import torch -<<<<<<< HEAD -from .ball_search import ( +from torchref.config import get_default_device, get_float_dtype + +from .frf.ball_search import ( ball_rotation_search, rotation_matrix_from_edmonds_euler, RotationPeak, -======= -from torchref.config import get_default_device, get_float_dtype - -from .ball_transform import ( - ball_rotation_search_torch, - rotation_matrix_from_euler_zyz, ->>>>>>> main ) from .translation import fft_translation_search_torch, TranslationPeak from .rigid_body import RigidBodyRefinement, RigidBodyResult @@ -258,6 +252,7 @@ def run( L: int = 48, P: int = 24, cluster_threshold_deg: Optional[float] = None, + engine: str = "frf_separate", ) -> List[MRSolution]: """ Run full MR pipeline with early stopping. @@ -301,7 +296,9 @@ def run( # Step 1: Rotation search if self.verbose: print("Step 1: Rotation search...") - rotation_peaks = self._rotation_search(n_rotation_peaks, d_min, d_max, L, P) + rotation_peaks = self._rotation_search( + n_rotation_peaks, d_min, d_max, L, P, engine=engine + ) if not rotation_peaks: if self.verbose: @@ -419,9 +416,31 @@ def _rotation_search( d_max: float, L: int, P: int, + engine: str = "frf_separate", ) -> list: - """Run ball rotation search; return list of (α, β, γ, score, σ) tuples.""" - # Prepare E-values + """Run the rotation search; return list of (α, β, γ, score, σ) tuples. + + ``engine="frf_separate"`` (default) uses the validated Phaser-faithful + engine (dense calc + auto_lmax + obs-unroll + no_grad) via the shared + ``align`` helpers; ``engine="ball"`` uses the legacy E-value ball search. + """ + if engine not in ("frf_separate", "ball"): + raise ValueError( + f"engine={engine!r}; expected 'frf_separate' (default) or 'ball'." + ) + if engine == "frf_separate": + from .align import _prepare_frf_inputs, _run_frf_separate_rotation + frf = _prepare_frf_inputs( + self.model, self.data, + d_min=d_min, d_max=d_max, n_shells=20, verbose=self.verbose, + ) + peaks = _run_frf_separate_rotation( + self.model, self.data, frf, + n_peaks=n_peaks, verbose=self.verbose, + ) + return [(p.alpha, p.beta, p.gamma, p.score, p.sigma) for p in peaks] + + # engine == "ball" — legacy ball-harmonic E-value search. E_obs, s_obs = self._get_e_values_obs(d_min, d_max) E_calc, s_calc = self._get_e_values_calc(d_min, d_max) diff --git a/torchref/alignment/rigid_body.py b/torchref/alignment/rigid_body.py index ae03a66c..3ab8f706 100644 --- a/torchref/alignment/rigid_body.py +++ b/torchref/alignment/rigid_body.py @@ -21,16 +21,14 @@ from dataclasses import dataclass from typing import Optional, Tuple +import math + import numpy as np import torch import torch.nn as nn -<<<<<<< HEAD -from torchref.scaling import Scaler -======= from torchref.config import get_default_device -from torchref.scaling import ScalerBase ->>>>>>> main +from torchref.scaling import Scaler from torchref.model import SfFFT from torchref.symmetry import spacegroup from torchref.refinement.targets import MaximumLikelihoodXrayTarget @@ -121,10 +119,20 @@ def __init__( rfactor_converged_threshold: float = 0.45, max_res: float = 4.0, verbose: int = 1, + refine_b: bool = False, + sigma_rot_deg: float = 0.0, + sigma_trans_ang: float = 0.0, + sigma_b: float = 0.0, ): super().__init__() self.device = device self.data = data + # Phase C: B-refine + Phaser-style Gaussian restraints. All off by + # default so previously-passing trajectories are unaffected. + self.refine_b = bool(refine_b) + self.sigma_rot_rad = math.radians(float(sigma_rot_deg)) if sigma_rot_deg > 0 else 0.0 + self.sigma_trans_ang = float(sigma_trans_ang) + self.sigma_b = float(sigma_b) xyz_iso, adp_iso, occ_iso, A_iso, B_iso = model.get_iso() @@ -185,6 +193,14 @@ def __init__( initial_translation = initial_translation.to(device=device).clone() self.translation_frac = nn.Parameter(initial_translation) + # Phase C: per-atom B-factor perturbation. Held as nn.Parameter even + # when refine_b=False (gradient just won't flow); cost is negligible + # and the codepath stays uniform. + n_iso = int(self.adp_iso.shape[0]) + self.delta_b_iso = nn.Parameter( + torch.zeros(n_iso, device=device, dtype=self.adp_iso.dtype), + ) + # Use `Scaler` for per-bin scales + anisotropy correction, but skip # the bulk-solvent setup. The solvent mask is computed once from # `model.xyz()` and goes stale as the joint refine moves atoms; the @@ -306,12 +322,20 @@ def forward(self, debug: bool = False) -> torch.Tensor: t_cart = self.translation_frac @ self.cell.fractional_matrix.T xyz_aniso = xyz_aniso_rotated + t_cart + # Phase C: optionally perturb per-atom B-factors. Clamp at 0 so the + # density model stays physical even mid-refine; the Gaussian + # restraint on delta_b_iso prevents large excursions. + if self.refine_b: + adp_iso_eff = (self.adp_iso + self.delta_b_iso).clamp(min=0.0) + else: + adp_iso_eff = self.adp_iso + # Compute structure factors via FFT (bypasses MixedTensor!) # Note: fractional matrices are now obtained from FFT's internal Cell object sf, _ = self.fft.compute_structure_factors( hkl=hkl, xyz_iso=xyz_transformed, - adp_iso=self.adp_iso, + adp_iso=adp_iso_eff, occ_iso=self.occ_iso, A_iso=self.A_iso, B_iso=self.B_iso, @@ -366,15 +390,42 @@ def refine( print(f" Setting up LBFGS optimizer niter = {n_iter} and max tries = {n_tries}") sys.stdout.flush() parameters = [self.rotation_parameters, self.translation_frac, *self.scaler.parameters()] + if self.refine_b: + parameters.append(self.delta_b_iso) self.optimizer = torch.optim.LBFGS( parameters, lr=1, max_iter=100, line_search_fn='strong_wolfe' ) + # Phaser-style Gaussian restraints (0 ⇒ disabled). Pre-square once. + sigma_rot_rad = self.sigma_rot_rad + sigma_trans_ang = self.sigma_trans_ang + sigma_b = self.sigma_b + restraints_active = ( + sigma_rot_rad > 0 or sigma_trans_ang > 0 + or (self.refine_b and sigma_b > 0) + ) + + def restraint_loss() -> torch.Tensor: + r = torch.zeros((), dtype=self.translation_frac.dtype, + device=self.translation_frac.device) + if sigma_rot_rad > 0: + r = r + 0.5 * (self.rotation ** 2).sum() / (sigma_rot_rad ** 2) + if sigma_trans_ang > 0: + # Cartesian translation = T_frac @ fractional_matrix.T (Å). + t_cart = self.translation_frac @ self.cell.fractional_matrix.T + r = r + 0.5 * (t_cart ** 2).sum() / (sigma_trans_ang ** 2) + if self.refine_b and sigma_b > 0: + r = r + 0.5 * (self.delta_b_iso ** 2).sum() / (sigma_b ** 2) + return r + def loss(): fcalc = self() - return self.xray_target(fcalc) + ll = self.xray_target(fcalc) + if restraints_active: + ll = ll + restraint_loss() + return ll noise = 0 diff --git a/torchref/alignment/sh.py b/torchref/alignment/sh.py index 61839d36..983c8560 100644 --- a/torchref/alignment/sh.py +++ b/torchref/alignment/sh.py @@ -199,6 +199,59 @@ def angular_density_weights( return w +def get_axis_order(sym_mats: torch.Tensor, axis: int) -> int: + """ + Order of the highest-multiplicity proper rotation about a principal axis. + + `sym_mats` is the spacegroup rotation matrices (n_ops, 3, 3). `axis` is + 0/1/2 for x/y/z. Returns the largest n such that some R in `sym_mats` is a + rotation by 2π/n around that axis. For non-rotational operations (or + rotations not aligned with the axis), the spacegroup element is skipped. + Returns 1 if no proper rotation around the axis exists. + + Used by the Phaser-style m-symmetry filter on the spherical-harmonic + coefficients: the Patterson is invariant under the spacegroup rotations, + so m-values that violate the highest-order axis symmetry are pure noise. + """ + a = torch.zeros(3, dtype=torch.float64, device=sym_mats.device) + a[axis] = 1.0 + max_order = 1 + n_ops = sym_mats.shape[0] + for k in range(n_ops): + R = sym_mats[k].to(torch.float64) + # Axis must be invariant under R (proper or improper rotation about it). + if (R @ a - a).norm().item() > 1e-3: + continue + # Trace of a rotation by angle θ about the preserved axis is 1+2cosθ. + tr = R.diagonal().sum().item() + cos_a = max(-1.0, min(1.0, (tr - 1.0) / 2.0)) + # Identity (angle ~0) → order 1. + if cos_a >= 1.0 - 1e-6: + continue + angle = math.acos(cos_a) + n = round(2 * math.pi / angle) + if n > max_order: + max_order = n + return max_order + + +def get_high_order_axis(sym_mats: torch.Tensor) -> Tuple[int, int]: + """ + Return (axis, zsymm) where `axis` ∈ {0, 1, 2} (x/y/z) maximises + `get_axis_order`, with z preferred on ties (matches Phaser's + `highOrderAxis()` in SpaceGroup.cc). + """ + orders = [get_axis_order(sym_mats, a) for a in (0, 1, 2)] + # Phaser: axis=3 (z); axis=2 if orderY > orderZ; axis=1 if orderX > both. + if orders[0] > orders[1] and orders[0] > orders[2]: + axis = 0 + elif orders[1] > orders[2]: + axis = 1 + else: + axis = 2 + return axis, orders[axis] + + def sh_expand_ball( s_vectors: torch.Tensor, values: torch.Tensor, @@ -208,6 +261,8 @@ def sh_expand_ball( enforce_friedel: bool = True, chunk_size: int = 2048, angular_weights: Optional[torch.Tensor] = None, + zsymm: int = 1, + skip_odd_l: bool = False, ) -> torch.Tensor: """ Analytical spherical-harmonic expansion of a scattered-point real field @@ -287,12 +342,22 @@ def sh_expand_ball( # scatter-add into shells f_plm.index_add_(0, shell_idx[start:stop], contrib) - if enforce_friedel: - # Zero odd-l rows explicitly (they should already be ~0; this kills FP drift). + if enforce_friedel or skip_odd_l: + # Zero odd-l rows explicitly (they should already be ~0; this kills FP drift + # and is the only required step when skip_odd_l is set without Friedel). l_vals = torch.arange(L, device=device) odd_mask = (l_vals % 2 == 1) f_plm[:, odd_mask, :] = 0.0 + # F1: m-symmetry filter (Phaser DataMR.cc:1019 / 1117). The Patterson is + # invariant under the spacegroup rotation operators, so SH coefficients + # whose m-index violates the highest-order rotation axis are pure noise. + # Zero them out post-expansion. With `zsymm=1` (no filter) this is a no-op. + if zsymm > 1: + m_vals = torch.arange(-(L - 1), L, device=device) # (2L-1,) + m_invalid = (m_vals.abs() % zsymm) != 0 + f_plm[:, :, m_invalid] = 0.0 + return f_plm @@ -419,6 +484,89 @@ def fit_overall_anisotropy( return U +def hkl_symops_to_cartesian( + sg_mats: torch.Tensor, + rec_basis: torch.Tensor, +) -> torch.Tensor: + """ + Convert spacegroup symmetry operators that act on integer Miller indices + into the equivalent rotation operators acting on Cartesian reciprocal- + space vectors (`s = h @ rec_basis`, column-vector form: `s = M @ h` with + `M = rec_basis^T`). + + For column vectors: `s' = M · S · M⁻¹ · s` so `P_cart = M · S · M⁻¹`. + + For orthogonal cells (orthorhombic+) M is diagonal and `P_cart == S` + exactly. For non-orthogonal cells (monoclinic, hex/trig with γ=120°, + triclinic), the Cartesian form differs and matters for any operation + that mixes the axes (e.g. averaging tensors over the point group). + + Parameters + ---------- + sg_mats : torch.Tensor, shape (n_ops, 3, 3) + Integer Miller-index symops (data.spacegroup.matrices). + rec_basis : torch.Tensor, shape (3, 3) + Reciprocal basis matrix such that `s = h @ rec_basis` (row-vector + convention used throughout torchref). + + Returns + ------- + sym_mats_cart : torch.Tensor, shape (n_ops, 3, 3), real + """ + dtype = torch.float64 + M = rec_basis.to(dtype).transpose(-1, -2) # (3, 3) + M_inv = torch.linalg.inv(M) + S = sg_mats.to(dtype) # (n_ops, 3, 3) + return torch.einsum("ij,kjl,lm->kim", M, S, M_inv) + + +def symmetrize_anisotropy( + U: torch.Tensor, + sym_mats_cart: torch.Tensor, +) -> torch.Tensor: + """ + Project a symmetric 3×3 tensor `U` onto the point-group-invariant + subspace of the spacegroup by averaging over its Cartesian rotation + operators: + + U_sym = (1/n) Σ_k P_k · U · P_k^T + + This mirrors Phaser's `site_symmetry.average_u_star()` and the + `RefineANO.cc:116-142` constraint construction. After symmetrisation the + tensor automatically satisfies the crystal's point-group symmetry: + + - cubic: U_sym = (trace U / 3) · I (1 DOF) + - tetragonal: U_sym = diag(λ, λ, μ) (2 DOF) + - orthorhombic: U_sym = diag(λ, μ, ν) (3 DOF) + - monoclinic: diagonal + one off-diagonal (4 DOF) + - triclinic: unchanged (6 DOF) + + This is the structural fix that prevents an unconstrained 6-component + regression from producing physically impossible anisotropy in + high-symmetry cells (e.g. fitting a 70 Ų eigenvalue on a cubic dataset + where every eigenvalue must be equal by symmetry). + + Parameters + ---------- + U : torch.Tensor, shape (3, 3), symmetric + sym_mats_cart : torch.Tensor, shape (n_ops, 3, 3) + Output of `hkl_symops_to_cartesian` — Cartesian-space rotation + operators of the spacegroup. + + Returns + ------- + U_sym : torch.Tensor, shape (3, 3), symmetric, same dtype/device as U + """ + dtype = U.dtype + device = U.device + sym = sym_mats_cart.to(device=device, dtype=dtype) + # Vectorised: U_avg = mean_k P_k · U · P_k^T + U_avg = torch.einsum("kij,jl,knl->in", sym, U, sym) / sym.shape[0] + # Enforce symmetry (numerical hygiene; pure rotation averaging preserves + # symmetry exactly, but FP drift can leave 1e-15 asymmetry). + return 0.5 * (U_avg + U_avg.transpose(-1, -2)) + + def apply_overall_anisotropy( F: torch.Tensor, s_vectors: torch.Tensor, diff --git a/torchref/alignment/translation.py b/torchref/alignment/translation.py index de613940..1eef94fd 100644 --- a/torchref/alignment/translation.py +++ b/torchref/alignment/translation.py @@ -424,6 +424,115 @@ def amplitude_translation_search( return corr_map_np, best, peaks +def llg_translation_rescore( + F_obs: torch.Tensor, + hkl: torch.Tensor, + centric: torch.Tensor, + shell_idx: torch.Tensor, + n_shells: int, + G: torch.Tensor, + h_R: torch.Tensor, + t_candidates: torch.Tensor, + sigma_a: torch.Tensor, + interp_var: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """ + Per-translation Rice / Woolfson log-likelihood, using the symmetry-summed + interpolator contributions ``G`` (Phaser EM_search analogue) and a fixed + per-shell σA. + + For each candidate t: + F_calc(h, t) = Σ_i G_i(h) · exp(2πi (h R_i) · t) + E_calc(h, t) = |F_calc(h, t)| / sqrt(_per_shell) + LLG(t) = Σ_shell [LL_Rice(E_obs, D·E_calc, var) − LL_Wilson(E_obs)] + where var = (1 − D²) + interp_var. + + Phase B alignment likelihood-TF. Replaces the |F|² Pearson correlation + in `amplitude_translation_search` as the scoring rule when the caller + re-ranks the FFT-cheap pre-filter peaks. + + Parameters + ---------- + F_obs : (N,) real + hkl : (N, 3) — unused here but kept for symmetry with the rest of the + module (and future extension to per-h variance models). + centric : (N,) bool + shell_idx : (N,) int64 — same binning as used to fit sigma_a / interp_var. + n_shells : int + G : (S, N) complex — per-sym F_p1 contributions × per-sym translation phase + (output of `precompute_G_for_rotation`). + h_R : (S, N, 3) — per-sym rotated reciprocal indices. + t_candidates : (K, 3) fractional translations to score. + sigma_a : (n_shells,) — fixed per-shell σA (shared across candidates). + interp_var : (N,) optional — per-reflection variance inflation. + + Returns + ------- + llg : (K,) torch.Tensor — log-likelihood gain per candidate. + """ + from .distributions import rice_log_likelihood, woolfson_log_likelihood + + device = G.device + real_dtype = torch.float64 + complex_dtype = G.dtype + + K = t_candidates.shape[0] + S, N = G.shape + + t_cand = t_candidates.to(device).to(real_dtype) # (K, 3) + # Phase factor for each (k, i, n): exp(2πi · (h_R[i, n] · t[k])) + phase_arg = torch.einsum("ind,kd->kin", h_R.to(real_dtype), t_cand) + phase = torch.exp(2j * torch.pi * phase_arg.to(complex_dtype)) # (K, S, N) + # F_calc(k, n) = Σ_i G[i, n] · phase[k, i, n] + Fc_complex = (G.view(1, S, N) * phase).sum(dim=1) # (K, N) + F_calc = Fc_complex.abs().to(real_dtype) # (K, N) + + # Per-shell E normalisation of F_calc across the K-batch. + shell_idx_l = shell_idx.to(device).long() + cnt = torch.bincount(shell_idx_l, minlength=n_shells).to(real_dtype) + shell_idx_k = shell_idx_l.view(1, -1).expand(K, N) + F2 = F_calc * F_calc + sum_per_shell = torch.zeros((K, n_shells), dtype=real_dtype, device=device) + sum_per_shell.scatter_add_(1, shell_idx_k, F2) + mean_per_shell = (sum_per_shell / cnt.clamp(min=1.0).unsqueeze(0)).clamp(min=1e-30) + norm_per_refl = mean_per_shell.sqrt().gather(1, shell_idx_k) # (K, N) + E_calc = F_calc / norm_per_refl # (K, N) + + F_obs_t = F_obs.to(device).to(real_dtype) + sum_F_obs2 = torch.zeros(n_shells, dtype=real_dtype, device=device) + sum_F_obs2.scatter_add_(0, shell_idx_l, F_obs_t * F_obs_t) + mean_F_obs2 = (sum_F_obs2 / cnt.clamp(min=1.0)).clamp(min=1e-30) + E_obs = F_obs_t / mean_F_obs2.sqrt().index_select(0, shell_idx_l) + + sigma_a_d = sigma_a.to(device).to(real_dtype) # (n_shells,) + D_per_refl = sigma_a_d.index_select(0, shell_idx_l) # (N,) + var_d = (1.0 - D_per_refl * D_per_refl).clamp(min=1e-4) # (N,) + if interp_var is not None: + var_per_refl = (var_d + interp_var.to(device).to(real_dtype)).clamp(min=1e-4) + else: + var_per_refl = var_d + + F_mean = D_per_refl.view(1, N) * E_calc # (K, N) + var_full = var_per_refl.view(1, N).expand(K, N) + E_obs_full = E_obs.view(1, N).expand(K, N) + cent_full = centric.to(device).to(torch.bool).view(1, N) + + ll_acent = rice_log_likelihood(E_obs_full, F_mean, var_full) + ll_cent = woolfson_log_likelihood(E_obs_full, F_mean, var_full) + ll = torch.where(cent_full, ll_cent, ll_acent) # (K, N) + + # Wilson reference (data only): F_mean = 0, var = 1. + var0 = torch.ones_like(E_obs) + F_mean0 = torch.zeros_like(E_obs) + ll_wil_acent = rice_log_likelihood(E_obs, F_mean0, var0) + ll_wil_cent = woolfson_log_likelihood(E_obs, F_mean0, var0) + ll_wil_per_refl = torch.where(centric.to(device).to(torch.bool), + ll_wil_cent, ll_wil_acent) + ll_wil_total = ll_wil_per_refl.sum() + + return ll.sum(dim=1) - ll_wil_total # (K,) + + def precompute_G_for_rotation( interpolator, R_rotation: torch.Tensor, diff --git a/torchref/alignment/wigner.py b/torchref/alignment/wigner.py index 8f7e5e2b..4e49a6b1 100644 --- a/torchref/alignment/wigner.py +++ b/torchref/alignment/wigner.py @@ -21,7 +21,8 @@ from __future__ import annotations -from typing import Optional, Tuple +from dataclasses import dataclass +from typing import List, Optional, Tuple import torch @@ -344,6 +345,140 @@ def evaluate_rotation_function_grid( return C, alpha_grid, beta_grid, gamma_grid +@dataclass +class AdaptiveRotationFunction: + """ + Phaser-style ragged rotation-function grid. + + On SO(3) the natural area element is `sin(β) dα dβ dγ`. A uniform Euler + cube oversamples the polar caps; this structure stores per-β slices of + variable `(qmax_k, pmax_k)` shape with sampling density matching + `pmax(β) = 720/Δ · cos(β/2)` and `qmax(β) = 360/Δ · sin(β/2)` (Phaser + FastRot.cc:92-96). + + Attributes + ---------- + betas : torch.Tensor + Shape `(n_β,)`, midpoint quadrature on `(0, π)`. + slices : list[torch.Tensor] + Length `n_β`. Slice `k` has shape `(qmax_k, pmax_k)` complex, indexed + as `slices[k][k_γ, k_α]` (matches dense convention `C[k_γ, k_β, k_α]`). + alpha_grids, gamma_grids : list[torch.Tensor] + Length `n_β`. Per-slice α and γ grid coordinates in radians. + grid_sampling_deg : float + Phaser's `grid_sampling` argument — target angular resolution in degrees. + """ + + betas: torch.Tensor + slices: List[torch.Tensor] + alpha_grids: List[torch.Tensor] + gamma_grids: List[torch.Tensor] + grid_sampling_deg: float + + def total_samples(self) -> int: + return sum(s.numel() for s in self.slices) + + +def evaluate_rotation_function_grid_adaptive( + xi_lmn: torch.Tensor, + L: int, + grid_sampling_deg: float = 3.0, + n_beta: Optional[int] = None, +) -> AdaptiveRotationFunction: + """ + Evaluate `C(α, β, γ)` on a Phaser-faithful adaptive Euler grid. + + Per β, the (α, γ) sampling density follows + + pmax(β) = max(1, round(720 / grid_sampling_deg · cos(β/2))) + qmax(β) = max(1, round(360 / grid_sampling_deg · sin(β/2))) + + Total sample count ≈ `(720 · 360) / grid_sampling_deg²` — independent of L + and free of the polar duplication that a uniform `(2L)³` grid produces. + + Computes `M_{m,n}(β_k) = Σ_l ξ_{l,m,n} d^l_{m,n}(β_k)` via the existing + batched `small_d_packed`, then for each β does a zero-padded scatter of + `M` into a `(pmax_k, qmax_k)` array and runs `torch.fft.fft2` on that + slice. The scatter is intentional: aliasing past the per-slice Nyquist + is the physically correct behaviour — those frequencies cannot be + resolved at that β. + """ + if n_beta is None: + n_beta = 2 * L + + device = xi_lmn.device + if xi_lmn.dtype == torch.complex128: + real_dtype = torch.float64 + elif xi_lmn.dtype == torch.complex64: + real_dtype = torch.float32 + else: + raise TypeError(f"xi_lmn must be complex, got {xi_lmn.dtype}") + complex_dtype = xi_lmn.dtype + + # β: midpoint rule on (0, π). + beta_grid = (torch.pi * (torch.arange(n_beta, dtype=real_dtype, device=device) + 0.5) + / n_beta) + + # Batched M_{m,n}(β_k) = Σ_l ξ_{l,m,n} d^l_{m,n}(β_k). + d_all = small_d_packed(L, beta_grid).to(complex_dtype) # (n_β, L, 2L-1, 2L-1) + M_all = (xi_lmn.unsqueeze(0) * d_all).sum(dim=1) # (n_β, 2L-1, 2L-1) + + # Per-β IFFT2 with adaptive shape. + half_beta = beta_grid * 0.5 + cos_half = torch.cos(half_beta) + sin_half = torch.sin(half_beta) + pmax_all = torch.clamp( + (720.0 / grid_sampling_deg * cos_half).round().to(torch.long), min=1, + ) + qmax_all = torch.clamp( + (360.0 / grid_sampling_deg * sin_half).round().to(torch.long), min=1, + ) + + m_vals = torch.arange(-(L - 1), L, dtype=torch.long, device=device) # (2L-1,) + n_vals = m_vals.clone() + + slices: List[torch.Tensor] = [] + alpha_grids: List[torch.Tensor] = [] + gamma_grids: List[torch.Tensor] = [] + + for k in range(n_beta): + pmax_k = int(pmax_all[k].item()) + qmax_k = int(qmax_all[k].item()) + + # Scatter M_{m,n} into the (pmax_k, qmax_k) Fourier-coefficient grid. + m_idx = m_vals % pmax_k # (2L-1,) + n_idx = n_vals % qmax_k # (2L-1,) + m_grid = m_idx.unsqueeze(-1).expand(2 * L - 1, 2 * L - 1) + n_grid = n_idx.unsqueeze(0).expand(2 * L - 1, 2 * L - 1) + flat_idx = m_grid * qmax_k + n_grid # (2L-1, 2L-1) + Mhat = torch.zeros((pmax_k, qmax_k), dtype=complex_dtype, device=device) + Mhat.view(-1).index_add_(0, flat_idx.reshape(-1), M_all[k].reshape(-1)) + + # `torch.fft.fft2` uses exp(-2π i k n / N) so positive (m, n) frequencies + # at index (m, n) reconstruct e^{-i m α} e^{-i n γ} on the uniform grid. + C_slice = torch.fft.fft2(Mhat, dim=(-2, -1)) # (pmax_k, qmax_k) + + alpha_grid_k = (2.0 * torch.pi / pmax_k) * torch.arange( + pmax_k, dtype=real_dtype, device=device, + ) + gamma_grid_k = (2.0 * torch.pi / qmax_k) * torch.arange( + qmax_k, dtype=real_dtype, device=device, + ) + + # Transpose to (γ, α) layout to match the dense `C[k_γ, k_β, k_α]` convention. + slices.append(C_slice.transpose(0, 1).contiguous()) + alpha_grids.append(alpha_grid_k) + gamma_grids.append(gamma_grid_k) + + return AdaptiveRotationFunction( + betas=beta_grid, + slices=slices, + alpha_grids=alpha_grids, + gamma_grids=gamma_grids, + grid_sampling_deg=grid_sampling_deg, + ) + + def evaluate_rotation_function_pointwise( xi_lmn: torch.Tensor, alpha: torch.Tensor, From c4394314006c8badb83ad996c5c8051fe24d9fa8 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 1 Jun 2026 14:10:39 +0200 Subject: [PATCH 005/250] FRF fixed and production ready. --- torchref/alignment/__init__.py | 100 +- torchref/alignment/align.py | 146 ++- torchref/alignment/frf/__init__.py | 50 +- torchref/alignment/frf/api.py | 66 +- torchref/alignment/frf/ball_search.py | 920 ----------------- torchref/alignment/frf/data_mr.py | 248 ++++- torchref/alignment/frf/french_wilson.py | 505 ++++++++++ torchref/alignment/frf/peak_finder.py | 57 +- torchref/alignment/frf/phaser_frf.py | 1164 ---------------------- torchref/alignment/frf/preprocessing.py | 121 ++- torchref/alignment/frf/rotation_utils.py | 89 ++ torchref/alignment/frf/sitelist_ang.py | 83 +- torchref/alignment/frf/types.py | 15 +- torchref/alignment/frf/wigner_d.py | 35 +- torchref/alignment/ml_rotation.py | 191 +--- torchref/alignment/patterson_filter.py | 163 --- torchref/alignment/pipeline.py | 53 +- torchref/alignment/sh.py | 113 ++- 18 files changed, 1309 insertions(+), 2810 deletions(-) delete mode 100644 torchref/alignment/frf/ball_search.py create mode 100644 torchref/alignment/frf/french_wilson.py delete mode 100644 torchref/alignment/frf/phaser_frf.py create mode 100644 torchref/alignment/frf/rotation_utils.py delete mode 100644 torchref/alignment/patterson_filter.py diff --git a/torchref/alignment/__init__.py b/torchref/alignment/__init__.py index f4535f19..8bb90d01 100644 --- a/torchref/alignment/__init__.py +++ b/torchref/alignment/__init__.py @@ -3,13 +3,13 @@ Pure-PyTorch Patterson-based molecular replacement: -1. Fast Rotation Function (`ball_search.ball_rotation_search`) — ball-harmonic - SO(3) cross-correlation, evaluated on an Euler-angle grid via inverse Wigner - transform. -2. Translation Search (`translation.fft_translation_search_torch`). -3. Rigid Body Refinement (`rigid_body.RigidBodyRefinement`) — LBFGS on rotation - and translation parameters with a maximum-likelihood x-ray target. -4. Unified Pipeline (`pipeline.MolecularReplacementPipeline`) — end-to-end +1. Fast Rotation Function (``frf.phaser_rotation_search`` / + ``frf.FastRotationFunction``) — Phaser-faithful Bessel-radial × SH + expansion, stable Wigner-d, dense P1-box calc. +2. Translation Search (``translation.fft_translation_search_torch``). +3. Rigid Body Refinement (``rigid_body.RigidBodyRefinement``) — LBFGS on + rotation and translation parameters with an ML target. +4. Unified Pipeline (``pipeline.MolecularReplacementPipeline``) — end-to-end workflow with early-stopping. Example — full MR pipeline @@ -26,32 +26,6 @@ pipeline = MolecularReplacementPipeline(data, model) solutions = pipeline.run(n_rotation_peaks=200, min_tries=3, max_tries=10) print(f"Best R-factor: {solutions[0].r_factor:.3f}") - -Example — individual components -------------------------------- -:: - - from torchref.alignment import ( - ball_rotation_search, - fft_translation_search_torch, - RigidBodyRefinement, - ) - - # 1. Rotation search - C, alphas, betas, gammas, peaks = ball_rotation_search( - s_obs, e_obs, s_calc, e_calc, L=48, P=24, - ) - - # 2. Translation search for top rotation - peak = peaks[0] - corr_map, best_trans, trans_peaks = fft_translation_search_torch( - F_obs, F_calc_rotated, hkl - ) - - # 3. Rigid body refinement - rb = RigidBodyRefinement(model, data, - initial_rotation=..., initial_translation=...) - result = rb.refine() """ import warnings @@ -62,30 +36,20 @@ ) # ============================================================================= -# Fast Rotation Function engines (consolidated in the .frf sub-package) +# Fast Rotation Function — Phaser-faithful, single engine # ============================================================================= -# Ball-harmonic E-value rotation search (engine="ball"). -from .frf.ball_search import ( - BallHarmonicCoefficients, +from .frf import ( + FastRotationFunction, RotationPeak, - ball_rotation_search, - compute_ball_harmonic_coefficients, - compute_ball_cross_correlation_coefficients, - find_rotation_peaks, - refine_peaks_subvoxel, - rotation_matrix_from_edmonds_euler, + dense_calc_via_box, edmonds_euler_from_rotation_matrix, - rotation_angular_distance_deg, -) -# Phaser-faithful engine — the production default rotation search. -from .frf.api import ( - FastRotationFunction, phaser_lmax_resolution, phaser_rotation_search, + rotation_angular_distance_deg, + rotation_matrix_from_edmonds_euler, ) -from .frf.dense_calc import dense_calc_via_box from .lattman_love import LattmanLoveInterpolator -from .ml_rotation import sim_mlrf_rescore, brute_ml_rotation_search +from .ml_rotation import m_letf1_rescore, sim_mlrf_rescore from .sh import ( evaluate_ylm, sh_expand_ball, @@ -167,27 +131,19 @@ from .sampling import VectorSampler, get_rotation_sampling_range __all__ = [ - # ------------------------------------------------------------------------- - # Rotation search (ball-harmonic Patterson, pure torch) - # ------------------------------------------------------------------------- - "ball_rotation_search", - "compute_ball_harmonic_coefficients", - "compute_ball_cross_correlation_coefficients", - "find_rotation_peaks", - "refine_peaks_subvoxel", - "BallHarmonicCoefficients", - "RotationPeak", - "rotation_matrix_from_edmonds_euler", - "edmonds_euler_from_rotation_matrix", - "rotation_angular_distance_deg", - # Phaser-faithful engine (production default) + # Rotation search "FastRotationFunction", "phaser_rotation_search", "phaser_lmax_resolution", "dense_calc_via_box", + "RotationPeak", + "rotation_matrix_from_edmonds_euler", + "edmonds_euler_from_rotation_matrix", + "rotation_angular_distance_deg", + # Rescore + interpolation "LattmanLoveInterpolator", + "m_letf1_rescore", "sim_mlrf_rescore", - "brute_ml_rotation_search", # Low-level math primitives "evaluate_ylm", "sh_expand_ball", @@ -198,32 +154,24 @@ "wigner_D_pointwise", "evaluate_rotation_function_grid", "evaluate_rotation_function_pointwise", - # ------------------------------------------------------------------------- # Pipeline - # ------------------------------------------------------------------------- "MolecularReplacementPipeline", "MRSolution", "cluster_rotation_peaks", "rotation_angular_distance", "euler_angular_distance", "align_model_to_data", - # ------------------------------------------------------------------------- # Translation - # ------------------------------------------------------------------------- "fft_translation_search", "fft_translation_search_torch", "TranslationPeak", "find_translation_peaks", "apply_translation_to_fcalc", "apply_translation_to_fcalc_torch", - # ------------------------------------------------------------------------- # Rigid body refinement - # ------------------------------------------------------------------------- "RigidBodyRefinement", "RigidBodyResult", - # ------------------------------------------------------------------------- # Transforms - # ------------------------------------------------------------------------- "RigidTransform", "quaternion_normalize", "quaternion_conjugate", @@ -237,24 +185,18 @@ "euler_zyz_to_quaternion", "rotation_matrix_from_euler", "sample_angles", - # ------------------------------------------------------------------------- # Clash scoring - # ------------------------------------------------------------------------- "ClashScoreCalculator", "AtomSampler", "compute_clash_score", - # ------------------------------------------------------------------------- # Distributions - # ------------------------------------------------------------------------- "stable_log_bessel_i0", "rice_log_likelihood", "woolfson_log_likelihood", "combined_log_likelihood", "acentric_pdf", "centric_pdf", - # ------------------------------------------------------------------------- # Utilities - # ------------------------------------------------------------------------- "VectorSampler", "get_rotation_sampling_range", ] diff --git a/torchref/alignment/align.py b/torchref/alignment/align.py index c90674af..f88083e4 100644 --- a/torchref/alignment/align.py +++ b/torchref/alignment/align.py @@ -27,12 +27,11 @@ import numpy as np import torch -from .frf.ball_search import ( - RotationPeak, - ball_rotation_search, +from .frf.rotation_utils import ( edmonds_euler_from_rotation_matrix, rotation_matrix_from_edmonds_euler, ) +from .frf.types import RotationPeak from .lattman_love import LattmanLoveInterpolator, estimate_interp_var from .ml_rotation import compute_sigma_a_luzzati, m_letf1_rescore, sim_mlrf_rescore from .rigid_body import RigidBodyRefinement @@ -283,7 +282,7 @@ def _prepare_frf_inputs( Encapsulates the data prep / anisotropy / symmetry-expansion logic previously inlined in `align_model_to_data`. The returned dataclass - feeds both the live `ball_rotation_search` call and the benchmark. + feeds both the live FRF call and the benchmark scripts. `F_obs` on the returned dataclass is the *anisotropy-corrected* value (matches what previously was the `F_obs_aniso` local variable). @@ -383,13 +382,33 @@ def _run_frf_separate_rotation( delta_vrms_A: float = 0.5, verbose: int = 0, _orbit_unroll: bool = False, - # --- Phaser model-prep knobs for the FRF calc side (default OFF) --- - apply_bulk_solvent: bool = False, + # --- Phaser model-prep knobs (defaults ON post v26 validation: see + # the SLURM v26 sweep in slurm_logs/rescore_v26_103820_*.csv). --- + apply_bulk_solvent: bool = True, solvent_fsol: float = 0.95, solvent_bsol: float = 300.0, - vrms_strategy: str = "fixed", + vrms_strategy: str = "oeffner", vrms_identity: float = 1.0, - apply_wilson_b: bool = False, + apply_wilson_b: bool = True, + use_epsilon: bool = False, + # obs-side term toggles (all default ON = production) for knockout bisection + frf_use_m_filter: bool = True, + frf_use_shell_variance: bool = True, + frf_use_french_wilson: bool = True, + frf_use_lerf1: bool = True, + frf_acentric_only: bool = False, + frf_d_max: float = 100.0, + frf_obs_lmax: Optional[int] = None, + frf_obs_solid_angle: bool = False, + frf_patterson_radius_scale: float = 1.0, + # Run the dominant SH-Bessel expansion (Legendre/Y_lm precompute + radial + # contraction) in single precision. The contraction is the FRF's bottleneck; + # FP64 is rate-limited on GPUs and SIMD-narrower on CPUs. The spherical-Bessel + # downward recurrence keeps its float64 internals and the cross-chunk + # accumulator stays full-precision, so only the contraction loses precision. + # `None` (default) → float32 on CUDA, full precision on CPU. Set explicitly + # to True/False to force single/double precision on either device. + frf_einsum_float32: Optional[bool] = None, ): """Phaser-faithful (validated) rotation search — the production default. @@ -422,7 +441,7 @@ def _run_frf_separate_rotation( # Full data resolution window; auto_lmax coarsens d_min to match the cap. # d_max ≈ no low-res cutoff (matches the validated config's d_max_mimic). d_min_eff = float(1.0 / s_mag_all.max().item()) - d_max_eff = 100.0 + d_max_eff = float(frf_d_max) # low-resolution cutoff (default 100 ≈ none) keep = (s_mag_all >= 1.0 / d_max_eff) & (s_mag_all <= 1.0 / d_min_eff) s_obs = s_vec_all[keep] @@ -451,11 +470,23 @@ def _run_frf_separate_rotation( # 202->324). The dedup is correct only as part of a coordinated Phaser- # faithful preprocessing chain (ε-Wilson + V(h) + σ_A), pending. sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) + # NOTE: centered lattices (I/C/F) list each point-group rotation once per + # centering op, so the raw matrices over-replicate the obs orbit (C2/I422 + # → ×2). Deduping to unique rotations is the correct point group and saves + # that compute, BUT it is NOT result-neutral: the equal-COUNT Wilson shells + # rebin when the obs count changes, perturbing the normalisation (3A5V + # 3→4). Left as-is to keep FRF behaviour stable; tracked in + # GHOST_INVESTIGATION.md as a follow-up (fix needs count-independent shells). + # Integer unrolled Miller indices aligned with s_obs — needed for the + # ε(h) multiplicity correction (use_epsilon), which down-weights the + # axial/zonal reflections that otherwise over-weight the m=0 SH column + # and feed high-symmetry rotation-function ghosts (compute_epsilon docstring). if _orbit_unroll: from .frf.preprocessing import epsilon_aware_unroll hkl_keep_int = hkl_all.to(torch.long).to(device)[keep] unrolled_hkl, asu_idx = epsilon_aware_unroll(hkl_keep_int, sg_mats) s_obs = unrolled_hkl.to(torch.float64) @ rec_basis + hkl_obs_int = unrolled_hkl.to(torch.float64) F_obs = F_obs[asu_idx] centric = centric[asu_idx] if sigF is not None: @@ -465,11 +496,30 @@ def _run_frf_separate_rotation( hkl_keep = hkl_all.to(torch.float64)[keep] hkl_unroll = torch.einsum("kij,nj->kni", sg_mats, hkl_keep).reshape(-1, 3) s_obs = hkl_unroll @ rec_basis + hkl_obs_int = hkl_unroll F_obs = F_obs.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() centric = centric.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() if sigF is not None: sigF = sigF.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() + # Optional: restrict the obs to ACENTRIC reflections before the SH + # expansion. Centric reflections lie on the reciprocal-space zones + # perpendicular to symmetry axes and carry concentrated symmetry-axis + # signal (and a heavier Wilson tail); pooling them into the obs over- + # weights the symmetry-axis channel that produces high-symmetry ghosts. + # Dropping them also makes the Wilson normalisation acentric-only. + if frf_acentric_only: + acen = ~centric + s_obs = s_obs[acen] + F_obs = F_obs[acen] + centric = centric[acen] + hkl_obs_int = hkl_obs_int[acen] + if sigF is not None: + sigF = sigF[acen] + if verbose > 0: + print(f" FRF acentric-only: kept {int(acen.sum())}/{acen.numel()} obs", + flush=True) + # Dense P1-box calc on the (un-rotated) search model at the coarsened res. model_radius_A = float( (model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item() @@ -523,18 +573,32 @@ def _run_frf_separate_rotation( d_min=d_min_eff, d_max=d_max_eff, n_peaks=n_peaks, delta_vrms_A=delta_vrms_for_frf, sigma_threshold=-5.0, - use_lerf1_intensity=True, - use_m_symmetry_filter=True, + use_lerf1_intensity=frf_use_lerf1, + use_m_symmetry_filter=frf_use_m_filter, sig_F_obs=sigF, - use_french_wilson=(sigF is not None), - use_shell_variance_weights=True, + use_french_wilson=(frf_use_french_wilson and (sigF is not None)), + use_shell_variance_weights=frf_use_shell_variance, + use_epsilon=use_epsilon, + hkl_obs=hkl_obs_int, grid_sampling_deg=grid_sampling_deg, model_radius_A=model_radius_A, auto_lmax=True, lmax_cap=lmax_cap, + obs_lmax=frf_obs_lmax, + obs_solid_angle=frf_obs_solid_angle, + patterson_radius_scale=frf_patterson_radius_scale, apply_bulk_solvent=apply_bulk_solvent, solvent_fsol=solvent_fsol, solvent_bsol=solvent_bsol, + compute_dtype=( + torch.complex64 + if ( + frf_einsum_float32 + if frf_einsum_float32 is not None + else (device.type == "cuda") # default: fp32 on GPU only + ) + else None + ), ) return peaks @@ -579,7 +643,6 @@ def align_model_to_data( use_lerf1_intensity: bool = False, use_fitted_delta_vrms: bool = False, use_even_l_only: bool = False, - engine: str = "frf_separate", frf_lmax_cap: int = 48, frf_dense_pad: float = 2.0, rescore_engine: str = "m_letf1", @@ -684,7 +747,7 @@ def align_model_to_data( f"frf_weight_combine={frf_weight_combine!r}; " "expected 'sigma_a_only' or 'sigma_a_x_variance'." ) - # Match ball_rotation_search's internal normalisation: sum-to-P. + # Per-shell sum-to-P normalisation (matches the FRF's internal weighting). w = w * (n_shells / w.sum().clamp(min=1e-30)) rotsearch_weights = w.to(patt_obs.dtype) rotsearch_auto_var = False @@ -741,45 +804,22 @@ def align_model_to_data( flush=True, ) - timer.start("3_ball_search") - if engine not in ("frf_separate", "ball"): - raise ValueError( - f"engine={engine!r}; expected 'frf_separate' (default) or 'ball'." - ) - if engine == "frf_separate": - # Validated Phaser-faithful default (dense calc + auto_lmax cap + - # obs-unroll + no_grad); solves the high-symmetry cases. The σA/LERF1/ - # m-filter ball-prep above is ignored on this path. - if verbose > 0: - print( - f"fit_to_data: frf_separate rotation search " - f"(dense calc + auto_lmax cap={frf_lmax_cap}, " - f"n_peaks={n_rotation_peaks})…", - flush=True, - ) - peaks = _run_frf_separate_rotation( - model, data, frf, - lmax_cap=frf_lmax_cap, dense_pad=frf_dense_pad, - n_peaks=n_rotation_peaks, verbose=verbose, - ) - else: # engine == "ball" — legacy ball-harmonic E-value search - if verbose > 0: - print( - f"fit_to_data: ball-search (L={L}, P={n_shells}, " - f"n_peaks={n_rotation_peaks})…", - flush=True, - ) - _, _, _, _, peaks = ball_rotation_search( - s_vec_for_search, patt_obs, s_vec_for_search, patt_calc, - L=L, P=n_shells, n_peaks=n_rotation_peaks, - refine_subvoxel=True, n_refine=min(n_rotation_peaks, 50), - sigma_threshold=-5.0, - weights=rotsearch_weights, - auto_variance_weights=rotsearch_auto_var, - zsymm=rotsearch_zsymm, - skip_odd_l=use_even_l_only, + timer.start("3_rotation_search") + # Phaser-faithful FRF (dense calc + auto_lmax cap + obs-unroll + no_grad); + # solves the high-symmetry cases. Single engine post-consolidation. + if verbose > 0: + print( + f"fit_to_data: frf_separate rotation search " + f"(dense calc + auto_lmax cap={frf_lmax_cap}, " + f"n_peaks={n_rotation_peaks})…", + flush=True, ) - timer.stop("3_ball_search") + peaks = _run_frf_separate_rotation( + model, data, frf, + lmax_cap=frf_lmax_cap, dense_pad=frf_dense_pad, + n_peaks=n_rotation_peaks, verbose=verbose, + ) + timer.stop("3_rotation_search") # --- Stage 2: Sim-MLRF rescore (per-shell σA fit per candidate) --- timer.start("4_sim_mlrf_rescore") diff --git a/torchref/alignment/frf/__init__.py b/torchref/alignment/frf/__init__.py index bf45f38b..e30100a4 100644 --- a/torchref/alignment/frf/__init__.py +++ b/torchref/alignment/frf/__init__.py @@ -1,19 +1,12 @@ -"""Fast Rotation Function engines (consolidated). +"""Fast Rotation Function — single, validated implementation. -This sub-package collects every Fast Rotation Function (FRF) implementation in -one place. Three engines coexist (see ``FRF_CONSOLIDATION.md``): +Phaser-faithful engine: chunked Bessel-SH expansion, stable Wigner-d +(``wigner_d.small_d_stable``), resolution↔bandwidth coupling +(``phaser_lmax_resolution``, default cap=48), dense P1-box calc, all under +``no_grad``. Solved the high-symmetry cases that broke the earlier ball +and Phaser-mimic engines (4BX9 342→4–7, 6G9X 77→1–4). -- **``api`` (the production default)** — the Phaser-faithful engine validated in - the high-symmetry investigation: chunked Bessel-SH expansion, stable Wigner-d - (``wigner_d.small_d_stable``), resolution↔bandwidth coupling - (``phaser_lmax_resolution``, cap=48), dense P1-box calc (``dense_calc``), all - under ``no_grad``. Reached 4BX9 342→4-7, 6G9X 77→1-4. -- **``ball_search`` (``engine="ball"``)** — the original pure-torch ball-harmonic - E-value rotation search. -- **``phaser_frf`` (``legacy_phaser_rotation_search``)** — the earlier 1164-line - Phaser-mimic, superseded by ``api`` but kept for reference/benchmarking. - -Shared leaf math (``..sh``, ``..wigner``) stays in the parent ``alignment`` +Shared leaf math (``..sh``, ``..wigner``) lives in the parent ``alignment`` package; this sub-package imports it "up". """ from .api import ( @@ -21,20 +14,13 @@ phaser_lmax_resolution, phaser_rotation_search, ) -from .ball_search import ( - BallHarmonicCoefficients, - RotationPeak as BallRotationPeak, - ball_rotation_search, - compute_ball_cross_correlation_coefficients, - compute_ball_harmonic_coefficients, +from .dense_calc import dense_calc_via_box, model_sf_abs +from .rotation_utils import ( edmonds_euler_from_rotation_matrix, - find_rotation_peaks, - refine_peaks_subvoxel, rotation_angular_distance_deg, rotation_matrix_from_edmonds_euler, + rotation_matrix_from_edmonds_euler_batch, ) -from .dense_calc import dense_calc_via_box, model_sf_abs -from .phaser_frf import phaser_rotation_search as legacy_phaser_rotation_search from .types import ( AdaptiveRotationFunction, BesselSHCoefficients, @@ -43,26 +29,18 @@ ) __all__ = [ - # Validated engine (production default) + # Engine "FastRotationFunction", "phaser_rotation_search", "phaser_lmax_resolution", "dense_calc_via_box", "model_sf_abs", - # Ball-harmonic engine (engine="ball") - "ball_rotation_search", - "BallHarmonicCoefficients", - "BallRotationPeak", - "compute_ball_harmonic_coefficients", - "compute_ball_cross_correlation_coefficients", - "find_rotation_peaks", - "refine_peaks_subvoxel", + # Rotation geometry helpers "rotation_matrix_from_edmonds_euler", + "rotation_matrix_from_edmonds_euler_batch", "edmonds_euler_from_rotation_matrix", "rotation_angular_distance_deg", - # Legacy Phaser-mimic - "legacy_phaser_rotation_search", - # Types (validated engine) + # Types "AdaptiveRotationFunction", "BesselSHCoefficients", "RotationPeak", diff --git a/torchref/alignment/frf/api.py b/torchref/alignment/frf/api.py index 6962daa2..bdc24e0f 100644 --- a/torchref/alignment/frf/api.py +++ b/torchref/alignment/frf/api.py @@ -35,6 +35,7 @@ detect_zsymm, eterm_sigma_a, french_wilson_preprocess, + solid_angle_weights, wilson_normalise, wilson_normalise_epsilon, ) @@ -171,9 +172,14 @@ def __init__( model_radius_A: Optional[float] = None, auto_lmax: bool = False, lmax_cap: int = 48, # sweet spot: higher L under-determines SH modes on sparse lattice + obs_lmax: Optional[int] = None, # cap obs SH bandwidth below calc (determinacy test) + obs_solid_angle: bool = False, # angular quadrature weight to de-bias the obs SH + patterson_radius_scale: float = 1.0, # <1 tightens the Patterson integration sphere + compute_dtype: Optional[torch.dtype] = None, # complex64 → faster einsum (GPU) ): self.device = s_obs.device self.real_dtype = s_obs.dtype + self.compute_dtype = compute_dtype # Phaser-faithful coupling of bandwidth to resolution (runMR_FRF.cc:408). # Overrides L and d_min so the SH expansion is not flooded with data @@ -228,12 +234,20 @@ def __init__( epsilon = compute_epsilon(hkl_obs, sym_mats) # 2. Bessel scaling default — Phaser's lmax · d_min (DataMR.cc:1107). + # bessel_h_scale = 2π·R_patt is the Patterson integration radius (the + # χ_Ω sphere): the Bessel argument is h = bessel_h_scale·|s|, so the + # radial basis represents the Patterson out to R_patt = bessel_h_scale + # /(2π). With auto_lmax this defaults to R_patt ≈ sphereOuter = 2·mean + # radius. `patterson_radius_scale` < 1 tightens it toward the + # short-range intra-molecular self-Patterson (excludes the noisy + # long-vector / inter-molecular tail that carries crystal symmetry). if bessel_h_scale is None: if d_min is None: raise ValueError("bessel_h_scale must be set when d_min is None") lmax = L - 1 lmax_even = lmax if lmax % 2 == 0 else lmax - 1 bessel_h_scale = float(lmax_even) * float(d_min) + bessel_h_scale = bessel_h_scale * float(patterson_radius_scale) self.bessel_h_scale = bessel_h_scale # 3. Wilson + optional FW + DFAC on obs. French-Wilson does its own @@ -275,6 +289,13 @@ def __init__( intensity_obs, smag_obs, n_var_shells=n_var_shells, ) + # 5b. Optional angular quadrature weight: de-bias the obs SH expansion + # for the non-uniform reciprocal-lattice point distribution (denser + # along symmetry directions → amplified symmetry-axis/ghost channel). + if obs_solid_angle: + w_ang = solid_angle_weights(s_obs).to(intensity_obs.dtype) + intensity_obs = intensity_obs * w_ang + # 6. ZSYMM detection + m-symmetry filter on obs SH coefficients. zsymm = detect_zsymm(sym_mats) if use_m_symmetry_filter else 1 self._zsymm = zsymm # also reused on the calc side when enabled (score_model) @@ -284,8 +305,22 @@ def __init__( s_obs, intensity_obs.to(self.real_dtype), L=L, bessel_h_scale=bessel_h_scale, zsymm=zsymm, enforce_friedel=True, + compute_dtype=self.compute_dtype, ) + # Optional: cap the obs SH bandwidth BELOW the calc's by zeroing high-l + # obs coefficients. The obs is expanded over the sparse, anisotropically- + # distributed reciprocal crystal lattice (denser along symmetry + # directions), so its high-l coefficients are aliased/under-determined + # and over-represent the symmetry-axis (ghost) channel. Truncating obs-l + # while keeping calc-l full tests whether that aliasing drives the + # high-symmetry ghosts. L (and the contraction) are unchanged; the rows + # l > obs_lmax just contribute nothing. + if obs_lmax is not None and obs_lmax < (L - 1): + l_idx = torch.arange(L, device=self.device) + self._c_obs.coeffs[:, l_idx > int(obs_lmax), :] = 0.0 + self._obs_lmax = obs_lmax + def score_model( self, s_calc: torch.Tensor, @@ -293,7 +328,6 @@ def score_model( *, n_peaks: int = 500, sigma_threshold: float = -5.0, - calc_m_symmetry_filter: bool = False, apply_bulk_solvent: bool = False, solvent_fsol: float = 0.95, solvent_bsol: float = 300.0, @@ -315,18 +349,16 @@ def score_model( eterm = eterm * sol intensity_calc = (eterm * eterm) * (E_calc * E_calc - 1.0) - # Bessel-SH expand calc. The "ideal-arithmetic" math says calc-side - # m-filtering is a no-op when obs is exactly spacegroup-invariant - # (R_sym(g) = NSYMP·R_orig(g), a constant scale); but in practice obs - # invariance leaks at high l (discrete-Y_lm non-orthogonality, anisotropy - # residuals, axial-reflection counting), and calc is rich in non-invariant - # m. Projecting calc onto invariant m kills the spurious obs-leak × - # calc-non-invariant product channel. - calc_zsymm = self._zsymm if calc_m_symmetry_filter else 1 + # Bessel-SH expand calc. The calc is NEVER m-filtered (zsymm=1): the + # model carries no crystal symmetry, and projecting it onto the obs's + # invariant-m subspace destroys the orientation information that + # discriminates truth — a calc-side m-filter was tested and is strongly + # harmful on high-symmetry cases (3K7M rank 8→92), so the knob was removed. c_calc = bessel_sh_expand( s_calc, intensity_calc.to(self.real_dtype), L=self.L, bessel_h_scale=self.bessel_h_scale, - zsymm=calc_zsymm, enforce_friedel=True, + zsymm=1, enforce_friedel=True, + compute_dtype=self.compute_dtype, ) # 8. Cross-correlate over the radial axis. @@ -375,16 +407,15 @@ def phaser_rotation_search( model_radius_A: Optional[float] = None, auto_lmax: bool = False, lmax_cap: int = 48, # sweet spot: higher L under-determines SH modes on sparse lattice - # NOTE: defaults to False — the standalone v20 calc-filter regressed the rebench - # (3K7M 18->189, 3GR5 47->204, 2DQ6 202->324; job 103409). Phaser's preprocessing - # pieces (eps-Wilson + V(h) + sigma_A) are mutually load-bearing; this knob will - # be flipped back on once the coordinated Phaser-faithful preprocessing chain lands. - calc_m_symmetry_filter: bool = False, + obs_lmax: Optional[int] = None, # cap obs SH bandwidth below calc (determinacy test) + obs_solid_angle: bool = False, # angular quadrature weight to de-bias the obs SH + patterson_radius_scale: float = 1.0, # <1 tightens the Patterson integration sphere # Phaser bulk-solvent (Babinet) folded into σ_A on the calc side # (EnsemblePDB.cc:96-100; solTerm.h:9). Default OFF; flip after sweep. apply_bulk_solvent: bool = False, solvent_fsol: float = 0.95, solvent_bsol: float = 300.0, + compute_dtype: Optional[torch.dtype] = None, ) -> Tuple[AdaptiveRotationFunction, List[RotationPeak]]: """Drop-in for ``torchref.alignment.phaser_frf.phaser_rotation_search``. @@ -423,12 +454,15 @@ def phaser_rotation_search( model_radius_A=model_radius_A, auto_lmax=auto_lmax, lmax_cap=lmax_cap, + obs_lmax=obs_lmax, + obs_solid_angle=obs_solid_angle, + patterson_radius_scale=patterson_radius_scale, + compute_dtype=compute_dtype, ) return frf.score_model( s_calc, F_calc, n_peaks=n_peaks, sigma_threshold=sigma_threshold, - calc_m_symmetry_filter=calc_m_symmetry_filter, apply_bulk_solvent=apply_bulk_solvent, solvent_fsol=solvent_fsol, solvent_bsol=solvent_bsol, diff --git a/torchref/alignment/frf/ball_search.py b/torchref/alignment/frf/ball_search.py deleted file mode 100644 index fd58eed3..00000000 --- a/torchref/alignment/frf/ball_search.py +++ /dev/null @@ -1,920 +0,0 @@ -""" -Pure-PyTorch Patterson-based rotation search. - -Replaces the old `ball_transform.py` (which depended on jax / s2fft / s2ball / -spherical / quaternionic). All math is implemented with `torchref.alignment.sh` -and `torchref.alignment.wigner` and runs on whatever device the inputs live on. - -Conventions (locked, asserted by tests/unit/alignment/test_*.py and the -synthetic-rotation integration test): - - Forward analytical SH expansion: - f_{p,l,m} = Σ_{i ∈ shell p} v_i · conj(Y_{l,m}(ŝ_i)) - with Friedel mates `(-s_i, v_i)` included so odd-l rows are zero by - construction. - - Cross-correlation in Wigner-D basis: - ξ_{l,m,n} = Σ_p w_p · f_{obs}[p,l,n] · conj(f_{calc}[p,l,m]) - C(R) = Σ_{l,m,n} ξ_{l,m,n} · D^l_{m,n}(R) (Edmonds D) - - Recovered rotation: the Euler triple (α, β, γ) at the maximum of C is the - rotation R = R_z(α) R_y(β) R_z(γ) (Edmonds active ZYZ) such that - s_calc = R · s_obs (column vector) - for the test scenario where F_calc was generated by applying R to the model - coordinates. To build a (3, 3) matrix using - `torchref.alignment.transform.rotation_matrix_from_euler` (which interprets - its args as `R = R_z(γ_arg) R_y(β_arg) R_z(α_arg)`), pass `[γ, β, α]`. -""" - -from __future__ import annotations - -import math -from dataclasses import dataclass -from typing import List, Optional, Tuple - -import torch -import torch.nn.functional as F - -from ..sh import ( - assign_shells, - compute_patterson_shell_variance, - equal_count_shell_edges, - sh_expand_ball, -) -from ..wigner import ( - AdaptiveRotationFunction, - evaluate_rotation_function_grid, - evaluate_rotation_function_pointwise, -) - - -# ============================================================================= -# Data containers -# ============================================================================= - - -@dataclass -class BallHarmonicCoefficients: - """Output of `compute_ball_harmonic_coefficients`.""" - - f_plm: torch.Tensor # complex, shape (P, L, 2L-1) - L: int - P: int - shell_edges: torch.Tensor # real, shape (P+1,) - shell_centers: torch.Tensor # real, shape (P,) - shell_counts: torch.Tensor # int64, shape (P,) - - -@dataclass -class RotationPeak: - """A single peak of the rotation function.""" - - alpha: float - beta: float - gamma: float - score: float # value of C at the peak (real part) - sigma: float # (C - mean) / std at the peak - - -# ============================================================================= -# Forward expansion -# ============================================================================= - - -def compute_ball_harmonic_coefficients( - s_vectors: torch.Tensor, - values: torch.Tensor, - L: int = 48, - P: int = 24, - shell_edges: Optional[torch.Tensor] = None, - d_min: Optional[float] = None, - d_max: Optional[float] = None, - enforce_friedel: bool = True, - zsymm: int = 1, - skip_odd_l: bool = False, -) -> BallHarmonicCoefficients: - """ - Compute the ball-harmonic expansion of a real scattered-point field. - - Parameters - ---------- - s_vectors : (N, 3) real tensor - Reciprocal-lattice vectors in 1/Å. - values : (N,) real tensor - Sample values (e.g. |E(h)|) at each `s_vector`. - L : int, default 48 - Angular bandlimit. - P : int, default 24 - Number of radial shells. - shell_edges : (P+1,) real tensor, optional - Pre-computed shell boundaries (in |s|). If None, use equal-count - binning between `d_min` and `d_max` (or full range). - d_min, d_max : float, optional - Resolution bounds (Å). When given, reflections outside [1/d_max, 1/d_min] - in |s| are discarded. - enforce_friedel : bool, default True - Augment input with (-s, value) pairs and zero odd-l rows after expansion. - - Returns - ------- - BallHarmonicCoefficients - """ - s_mag = s_vectors.norm(dim=-1) - if d_min is not None or d_max is not None: - s_lo = 1.0 / d_max if d_max is not None else 0.0 - s_hi = 1.0 / d_min if d_min is not None else float("inf") - keep = (s_mag >= s_lo) & (s_mag <= s_hi) - s_vectors = s_vectors[keep] - values = values[keep] - s_mag = s_mag[keep] - - if shell_edges is None: - shell_edges, _ = equal_count_shell_edges(s_mag, P) - shell_centers = 0.5 * (shell_edges[:-1] + shell_edges[1:]) - shell_idx = assign_shells(s_mag, shell_edges) - keep = shell_idx >= 0 - s_vectors = s_vectors[keep] - values = values[keep] - shell_idx = shell_idx[keep] - - shell_counts = torch.bincount(shell_idx, minlength=P).to(torch.int64) - - f_plm = sh_expand_ball( - s_vectors, values, shell_idx, P, L, - enforce_friedel=enforce_friedel, - zsymm=zsymm, skip_odd_l=skip_odd_l, - ) - - return BallHarmonicCoefficients( - f_plm=f_plm, L=L, P=P, - shell_edges=shell_edges, - shell_centers=shell_centers, - shell_counts=shell_counts, - ) - - -# ============================================================================= -# Cross-correlation -# ============================================================================= - - -def compute_ball_cross_correlation_coefficients( - f_obs: BallHarmonicCoefficients, - f_calc: BallHarmonicCoefficients, - weights: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """ - Wigner-D coefficients ξ_{l,m,n} of the rotation function (see module docstring). - - ξ_{l,m,n} = Σ_p w_p · f_{obs}[p, l, n] · conj(f_{calc}[p, l, m]) - """ - assert f_obs.L == f_calc.L, "bandwidths must match" - assert f_obs.P == f_calc.P, "shell counts must match" - P, L = f_obs.P, f_obs.L - - if weights is None: - # Default: weight by shell count (equal weight per reflection). - counts = f_obs.shell_counts.to(torch.float64) - w_sum = counts.sum().clamp(min=1.0) - weights = (counts / w_sum).to(f_obs.f_plm.dtype) - else: - weights = weights.to(f_obs.f_plm.dtype) - - xi = torch.einsum( - "p,pln,plm->lmn", - weights, - f_obs.f_plm, - torch.conj(f_calc.f_plm), - ) - return xi - - -# ============================================================================= -# Peak finding -# ============================================================================= - - -def _circular_pad_gamma_alpha(x: torch.Tensor, r: int) -> torch.Tensor: - """ - Wrap-pad axes 0 (γ) and -1 (α) by `r` voxels using circular boundary - conditions, and constant-pad axis 1 (β) with -inf (non-periodic). - Input shape (Γ, B, Α); output shape (Γ+2r, B+2r, Α+2r). - """ - if r <= 0: - return x - x = torch.cat([x[-r:], x, x[:r]], dim=0) # wrap γ - x = torch.cat([x[..., -r:], x, x[..., :r]], dim=-1) # wrap α - pad_shape = list(x.shape) - pad_shape[1] = r - neg_inf = torch.full(pad_shape, float("-inf"), dtype=x.dtype, device=x.device) - x = torch.cat([neg_inf, x, neg_inf], dim=1) # β: -inf pad - return x - - -def find_rotation_peaks( - C: torch.Tensor, - alphas: torch.Tensor, - betas: torch.Tensor, - gammas: torch.Tensor, - n_peaks: int = 200, - sigma_threshold: float = 0.0, - cluster_radius_voxels: int = 2, -) -> List[RotationPeak]: - """ - Extract local maxima of `C` and return them sorted by descending value. - - Periodic in α and γ (FFT axes); β is non-periodic but mirror at the - endpoints isn't enforced — the user is expected to oversample β. - - Patterson centrosymmetry C(R) = C(-R) is left to the caller — when needed, - add a post-processing step that merges peaks with R and -R representations. - - Vectorised over voxels: NMS is implemented as `C == max_pool3d(C)` with - `kernel = 2·r + 1`, circular in α and γ. The legacy per-voxel Python loop - with per-iteration `.item()` is replaced with a single sort + bulk - `.tolist()` at the end — same semantics on randomly-real data (exact - ties are vanishingly rare), but device-portable and ~1-2 orders of - magnitude faster. - """ - C_real = C.real - n_gamma, n_beta, n_alpha = C_real.shape - flat_stats = C_real.flatten() - - mean = flat_stats.mean() - std = flat_stats.std().clamp(min=1e-30) - threshold = mean + sigma_threshold * std - - r = int(cluster_radius_voxels) - if r > 0: - padded = _circular_pad_gamma_alpha(C_real, r) - # max_pool3d wants (N, C, D, H, W). Add batch+channel dims, then strip. - pooled = F.max_pool3d( - padded.unsqueeze(0).unsqueeze(0), - kernel_size=2 * r + 1, stride=1, padding=0, - )[0, 0] - is_local_max = C_real >= pooled - else: - is_local_max = torch.ones_like(C_real, dtype=torch.bool) - - above_threshold = C_real >= threshold - candidate = is_local_max & above_threshold - candidate_flat = candidate.flatten() - cand_idx = candidate_flat.nonzero(as_tuple=False).flatten() - if cand_idx.numel() == 0: - return [] - cand_vals = C_real.flatten().index_select(0, cand_idx) - - # Sort descending; truncate to n_peaks. - top_k = min(int(n_peaks), int(cand_vals.numel())) - sort_vals, sort_perm = torch.sort(cand_vals, descending=True) - sel_flat_idx = cand_idx.index_select(0, sort_perm[:top_k]) - sel_vals = sort_vals[:top_k] - - # Convert flat → (g, b, a). Periodic axes have already been deduped, so - # plain integer arithmetic is sufficient. - nba = n_beta * n_alpha - ig = sel_flat_idx // nba - rem = sel_flat_idx - ig * nba - ib = rem // n_alpha - ia = rem - ib * n_alpha - - sigmas = (sel_vals - mean) / std - - # Bulk transfer to Python lists, then build dataclasses. - a_list = alphas.index_select(0, ia).tolist() - b_list = betas.index_select(0, ib).tolist() - g_list = gammas.index_select(0, ig).tolist() - v_list = sel_vals.tolist() - s_list = sigmas.tolist() - return [ - RotationPeak(alpha=a, beta=b, gamma=g, score=v, sigma=s) - for a, b, g, v, s in zip(a_list, b_list, g_list, v_list, s_list) - ] - - -# ============================================================================= -# Adaptive-grid peak finding -# ============================================================================= - - -def _adaptive_global_stats(arf: AdaptiveRotationFunction) -> Tuple[float, float]: - """ - Sample-weighted global mean and std across all β-slices. - - Every voxel of every slice contributes equally — the adaptive grid is - already sample-uniform in SO(3) sense (each voxel covers a near-constant - `sin(β)/(pmax(β)·qmax(β))` solid angle), so simple mean/std is correct. - """ - flat = torch.cat([s.real.reshape(-1) for s in arf.slices]) - mean = float(flat.mean().item()) - std = float(flat.std().clamp(min=1e-30).item()) - return mean, std - - -def _slice_nms_mask(slice_C: torch.Tensor, r_alpha: int, r_gamma: int) -> torch.Tensor: - """ - 2-D non-maximum-suppression for one β-slice, circular in both α and γ. - - `slice_C` is real-valued (γ, α). Returns a boolean mask of local maxima. - """ - if r_alpha <= 0 and r_gamma <= 0: - return torch.ones_like(slice_C, dtype=torch.bool) - n_gamma, n_alpha = slice_C.shape - r_a = max(r_alpha, 0) - r_g = max(r_gamma, 0) - # Clamp radii so we never wrap past the slice's own period — at β = 0 the - # γ axis has length 1, at β = π the α axis has length 1. - r_a = min(r_a, max(n_alpha - 1, 0)) - r_g = min(r_g, max(n_gamma - 1, 0)) - if r_a == 0 and r_g == 0: - return torch.ones_like(slice_C, dtype=torch.bool) - # Circular pad in γ (axis 0) then α (axis 1). - x = slice_C - if r_g > 0: - x = torch.cat([x[-r_g:], x, x[:r_g]], dim=0) - if r_a > 0: - x = torch.cat([x[:, -r_a:], x, x[:, :r_a]], dim=1) - pooled = F.max_pool2d( - x.unsqueeze(0).unsqueeze(0), - kernel_size=(2 * r_g + 1, 2 * r_a + 1), - stride=1, padding=0, - )[0, 0] - return slice_C >= pooled - - -def _sample_slice_at( - slice_C: torch.Tensor, - alpha_grid: torch.Tensor, - gamma_grid: torch.Tensor, - alpha_q: float, - gamma_q: float, -) -> float: - """ - Nearest-neighbour sample of a β-slice at a query (α, γ) in radians. - """ - n_gamma, n_alpha = slice_C.shape - two_pi = 2.0 * math.pi - a0 = float(alpha_grid[0].item()) - g0 = float(gamma_grid[0].item()) - da = two_pi / n_alpha - dg = two_pi / n_gamma - ia = int(round((alpha_q - a0) / da)) % n_alpha - ig = int(round((gamma_q - g0) / dg)) % n_gamma - return float(slice_C[ig, ia].item()) - - -def _so3_greedy_nms( - alpha: torch.Tensor, - beta: torch.Tensor, - gamma: torch.Tensor, - values: torch.Tensor, - nms_radius_rad: float, - keep_at_most: int, -) -> torch.Tensor: - """ - Greedy SO(3) non-max suppression by angular distance between rotations. - - Walks the candidates in descending `values` order; keeps a candidate iff - its angular distance `arccos((tr(R_cand · R_kept^T) − 1) / 2)` exceeds - `nms_radius_rad` from every already-kept rotation. Returns the indices - of the kept candidates in the order they were accepted (i.e. by - descending score). - - Complexity: `O(N · K)` matrix-trace ops where `K = min(keep_at_most, N)` - — the inner loop runs entirely in PyTorch on the host device. - """ - n = alpha.shape[0] - if n == 0: - return torch.zeros(0, dtype=torch.long, device=alpha.device) - R_all = rotation_matrix_from_edmonds_euler_batch(alpha, beta, gamma) # (n, 3, 3) - order = torch.argsort(values, descending=True) - - cos_thresh = math.cos(nms_radius_rad) - kept_indices: List[int] = [] - kept_R = R_all.new_zeros((0, 3, 3)) - - for i in order.tolist(): - Ri = R_all[i] # (3, 3) - if kept_R.shape[0] > 0: - # tr(Ri · R_kept^T) over the batch of kept rotations. - tr = torch.einsum("ab,kab->k", Ri, kept_R) - # angular distance = arccos((tr − 1) / 2). Compare against - # cosine to avoid the acos call entirely. - cos_theta = ((tr - 1.0) * 0.5).clamp(-1.0, 1.0) - if (cos_theta > cos_thresh).any().item(): - continue - kept_indices.append(i) - kept_R = torch.cat([kept_R, Ri.unsqueeze(0)], dim=0) - if len(kept_indices) >= keep_at_most: - break - - return torch.tensor(kept_indices, dtype=torch.long, device=alpha.device) - - -def find_rotation_peaks_adaptive( - arf: AdaptiveRotationFunction, - n_peaks: int = 200, - sigma_threshold: float = 0.0, - cluster_radius_deg: float = 6.0, - nms_radius_deg: float = 6.0, - candidate_cap: int = 10000, -) -> List[RotationPeak]: - """ - Local-maximum peak finder for the ragged adaptive Euler grid. - - Algorithm: - 1. Per-slice 2-D NMS (circular in α and γ, radius converted from - `cluster_radius_deg` to a per-slice voxel count). - 2. Threshold + truncate to the top `candidate_cap` survivors across - all slices (bounds the cost of the next step). - 3. Greedy SO(3) non-max suppression: walk candidates in descending - score; keep iff angular distance to every already-kept rotation - is > `nms_radius_deg`. This subsumes the role the dense path's - `max_pool3d` played and is geometry-aware — duplicate "polar" - peaks differing only in α+γ (β≈0) or α−γ (β≈π) have angular - distance ≈ 0 and are correctly collapsed. - 4. Truncate to top `n_peaks` and attach Z-score using the - globally-pooled mean / std. - """ - n_beta = arf.betas.shape[0] - mean, std = _adaptive_global_stats(arf) - threshold = mean + sigma_threshold * std - - # Per-slice local maxima + threshold. Collect onto torch tensors directly - # to avoid per-candidate .item() round-trips. - alpha_chunks: List[torch.Tensor] = [] - beta_chunks: List[torch.Tensor] = [] - gamma_chunks: List[torch.Tensor] = [] - value_chunks: List[torch.Tensor] = [] - for k in range(n_beta): - slice_C = arf.slices[k].real - n_gamma_k, n_alpha_k = slice_C.shape - r_a = max(1, int(round(cluster_radius_deg * n_alpha_k / 360.0))) - r_g = max(1, int(round(cluster_radius_deg * n_gamma_k / 360.0))) - - is_max = _slice_nms_mask(slice_C, r_alpha=r_a, r_gamma=r_g) - is_max = is_max & (slice_C >= threshold) - if not bool(is_max.any().item()): - continue - ig_idx, ia_idx = torch.nonzero(is_max, as_tuple=True) - vals = slice_C[ig_idx, ia_idx] - a_grid_k = arf.alpha_grids[k] - g_grid_k = arf.gamma_grids[k] - alpha_chunks.append(a_grid_k.index_select(0, ia_idx)) - gamma_chunks.append(g_grid_k.index_select(0, ig_idx)) - beta_chunks.append(arf.betas[k].expand_as(vals)) - value_chunks.append(vals) - - if not value_chunks: - return [] - - alpha_all = torch.cat(alpha_chunks) - beta_all = torch.cat(beta_chunks) - gamma_all = torch.cat(gamma_chunks) - value_all = torch.cat(value_chunks) - - # Cap the candidate set by raw score before the O(N·K) SO(3) NMS. - if value_all.shape[0] > candidate_cap: - top_v, top_i = torch.topk(value_all, candidate_cap) - alpha_all = alpha_all.index_select(0, top_i) - beta_all = beta_all.index_select(0, top_i) - gamma_all = gamma_all.index_select(0, top_i) - value_all = top_v - - # Greedy SO(3) NMS. - keep_idx = _so3_greedy_nms( - alpha_all, beta_all, gamma_all, value_all, - nms_radius_rad=math.radians(nms_radius_deg), - keep_at_most=int(n_peaks), - ) - - if keep_idx.numel() == 0: - return [] - - alpha_keep = alpha_all.index_select(0, keep_idx).tolist() - beta_keep = beta_all.index_select(0, keep_idx).tolist() - gamma_keep = gamma_all.index_select(0, keep_idx).tolist() - value_keep = value_all.index_select(0, keep_idx).tolist() - - return [ - RotationPeak( - alpha=a, beta=b, gamma=g, score=v, - sigma=(v - mean) / max(std, 1e-30), - ) - for (a, b, g, v) in zip(alpha_keep, beta_keep, gamma_keep, value_keep) - ] - - -def refine_peaks_subvoxel_adaptive( - peaks: List[RotationPeak], - arf: AdaptiveRotationFunction, -) -> List[RotationPeak]: - """ - Sub-voxel parabolic refinement for adaptive-grid peaks. - - For each peak at (α, β_k, γ): - * α and γ refinement uses the per-slice spacing `2π / pmax_k`, - `2π / qmax_k`. Skipped at slices where the relevant axis has - length 1 (β ≈ 0 → no γ refinement; β ≈ π → no α refinement). - * β refinement fits a parabola through the values at the same - (α, γ) sampled into slices k-1 and k+1 by nearest-neighbour. - Boundary slices use db = 0. - """ - if not peaks: - return list(peaks) - - n_beta = arf.betas.shape[0] - beta_grid = arf.betas - # Beta is uniform in our construction, so a single dβ is correct. - if n_beta > 1: - dbeta = float((beta_grid[1] - beta_grid[0]).item()) - else: - dbeta = 0.0 - - out: List[RotationPeak] = [] - for p in peaks: - # Locate the host β-slice. - ib = int(round((p.beta - float(beta_grid[0].item())) / dbeta)) if dbeta > 0 else 0 - ib = max(0, min(ib, n_beta - 1)) - slice_C = arf.slices[ib].real - a_grid = arf.alpha_grids[ib] - g_grid = arf.gamma_grids[ib] - n_gamma_k, n_alpha_k = slice_C.shape - da = 2.0 * math.pi / n_alpha_k - dg = 2.0 * math.pi / n_gamma_k - - # Integer voxel indices. - a0 = float(a_grid[0].item()) - g0 = float(g_grid[0].item()) - ia = int(round((p.alpha - a0) / da)) % n_alpha_k - ig = int(round((p.gamma - g0) / dg)) % n_gamma_k - - y0 = float(slice_C[ig, ia].item()) - - # α refinement. - if n_alpha_k >= 3: - am = (ia - 1) % n_alpha_k - ap = (ia + 1) % n_alpha_k - yam = float(slice_C[ig, am].item()) - yap = float(slice_C[ig, ap].item()) - denom = yam - 2.0 * y0 + yap - if abs(denom) > 1e-30: - d_a = max(-1.0, min(1.0, 0.5 * (yam - yap) / denom)) - else: - d_a = 0.0 - alpha_new = (p.alpha + d_a * da) % (2.0 * math.pi) - else: - alpha_new = p.alpha - - # γ refinement. - if n_gamma_k >= 3: - gm = (ig - 1) % n_gamma_k - gp = (ig + 1) % n_gamma_k - ygm = float(slice_C[gm, ia].item()) - ygp = float(slice_C[gp, ia].item()) - denom = ygm - 2.0 * y0 + ygp - if abs(denom) > 1e-30: - d_g = max(-1.0, min(1.0, 0.5 * (ygm - ygp) / denom)) - else: - d_g = 0.0 - gamma_new = (p.gamma + d_g * dg) % (2.0 * math.pi) - else: - gamma_new = p.gamma - - # β refinement: sample neighbouring slices at (α, γ). - if 0 < ib < n_beta - 1 and dbeta > 0: - v_prev = _sample_slice_at( - arf.slices[ib - 1].real, arf.alpha_grids[ib - 1], - arf.gamma_grids[ib - 1], p.alpha, p.gamma, - ) - v_next = _sample_slice_at( - arf.slices[ib + 1].real, arf.alpha_grids[ib + 1], - arf.gamma_grids[ib + 1], p.alpha, p.gamma, - ) - denom = v_prev - 2.0 * y0 + v_next - if abs(denom) > 1e-30: - d_b = max(-1.0, min(1.0, 0.5 * (v_prev - v_next) / denom)) - else: - d_b = 0.0 - beta_new = max(0.0, min(math.pi, p.beta + d_b * dbeta)) - else: - beta_new = p.beta - - out.append(RotationPeak( - alpha=alpha_new, beta=beta_new, gamma=gamma_new, - score=p.score, sigma=p.sigma, - )) - return out - - -# ============================================================================= -# Sub-voxel refinement (pointwise gradient ascent on C(R) using autograd) -# ============================================================================= - - -def _parabolic_offset_vec( - y_minus: torch.Tensor, y_zero: torch.Tensor, y_plus: torch.Tensor, -) -> torch.Tensor: - """ - Vectorised sub-voxel offset for a quadratic through three samples - y(-1), y(0), y(+1). Returns δ clamped to [-1, 1]; degenerate cells - (|denom| < 1e-30) get δ = 0. - """ - denom = y_minus - 2.0 * y_zero + y_plus - bad = denom.abs() < 1e-30 - safe_denom = torch.where(bad, torch.ones_like(denom), denom) - delta = 0.5 * (y_minus - y_plus) / safe_denom - delta = torch.where(bad, torch.zeros_like(delta), delta) - return delta.clamp(-1.0, 1.0) - - -def refine_peaks_subvoxel( - peaks: List[RotationPeak], - C: torch.Tensor, - alphas: torch.Tensor, - betas: torch.Tensor, - gammas: torch.Tensor, -) -> List[RotationPeak]: - """ - Sub-voxel refinement of integer-grid peaks by separable quadratic fitting. - - For each peak, fits y = a + b·δ + c·δ² through the values at ±1 voxel along - each of α, β, γ axes (α and γ periodic with the grid length, β not). - Returns refined peaks with interpolated angles and values. - - Vectorised over all peaks: indices and neighbour samples are gathered in - a single batched op; the legacy per-peak `.item()` loop is replaced with - one `.tolist()` at the end (cuts ~12·K scalar GPU↔CPU round-trips). - """ - if not peaks: - return list(peaks) - - n_gamma, n_beta, n_alpha = C.shape - Cr = C.real - device = Cr.device - real_dtype = alphas.dtype - - dalpha = (alphas[1] - alphas[0]).item() - dgamma = (gammas[1] - gammas[0]).item() - dbeta = (betas[1] - betas[0]).item() if betas.numel() > 1 else 0.0 - a0 = alphas[0].item() - g0 = gammas[0].item() - b0 = betas[0].item() - - # Pack peak coordinates into tensors in one go. - K = len(peaks) - alpha_in = torch.tensor([p.alpha for p in peaks], dtype=real_dtype, device=device) - beta_in = torch.tensor([p.beta for p in peaks], dtype=real_dtype, device=device) - gamma_in = torch.tensor([p.gamma for p in peaks], dtype=real_dtype, device=device) - - # Integer voxel indices. - ia = ((alpha_in - a0) / dalpha).round().long() % n_alpha - ig = ((gamma_in - g0) / dgamma).round().long() % n_gamma - if dbeta > 0: - ib = ((beta_in - b0) / dbeta).round().long().clamp(0, n_beta - 1) - else: - ib = torch.zeros(K, dtype=torch.long, device=device) - - am = (ia - 1) % n_alpha - ap = (ia + 1) % n_alpha - gm = (ig - 1) % n_gamma - gp = (ig + 1) % n_gamma - # β neighbours: clamp at boundaries; we zero δβ explicitly there below. - ibm = (ib - 1).clamp(min=0, max=n_beta - 1) - ibp = (ib + 1).clamp(min=0, max=n_beta - 1) - - y0 = Cr[ig, ib, ia] - ya_m = Cr[ig, ib, am] - ya_p = Cr[ig, ib, ap] - yg_m = Cr[gm, ib, ia] - yg_p = Cr[gp, ib, ia] - yb_m = Cr[ig, ibm, ia] - yb_p = Cr[ig, ibp, ia] - - da = _parabolic_offset_vec(ya_m, y0, ya_p) - dg = _parabolic_offset_vec(yg_m, y0, yg_p) - db = _parabolic_offset_vec(yb_m, y0, yb_p) - # Zero out β offset at the boundary (single-sided / undefined). - at_b_boundary = (ib == 0) | (ib == n_beta - 1) - db = torch.where(at_b_boundary, torch.zeros_like(db), db) - if dbeta <= 0: - db = torch.zeros_like(db) - - two_pi = 2.0 * math.pi - alpha_new = (alphas.index_select(0, ia) + da * dalpha) % two_pi - gamma_new = (gammas.index_select(0, ig) + dg * dgamma) % two_pi - beta_new = (betas.index_select(0, ib) + db * dbeta).clamp(min=0.0, max=math.pi) - - a_list = alpha_new.tolist() - b_list = beta_new.tolist() - g_list = gamma_new.tolist() - return [ - RotationPeak( - alpha=a, beta=b, gamma=g, score=p.score, sigma=p.sigma, - ) - for a, b, g, p in zip(a_list, b_list, g_list, peaks) - ] - - -# ============================================================================= -# Top-level entry -# ============================================================================= - - -def ball_rotation_search( - s_obs: torch.Tensor, - e_obs: torch.Tensor, - s_calc: torch.Tensor, - e_calc: torch.Tensor, - L: int = 48, - P: int = 24, - n_peaks: int = 200, - d_min: Optional[float] = None, - d_max: Optional[float] = None, - refine_subvoxel: bool = True, - n_refine: int = 20, - sigma_threshold: float = 1.0, - weights: Optional[torch.Tensor] = None, - auto_variance_weights: bool = True, - zsymm: int = 1, - skip_odd_l: bool = False, -) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, List[RotationPeak]]: - """ - End-to-end Patterson rotation search. - - Parameters - ---------- - s_obs, s_calc : (N, 3) real - Reciprocal-lattice vectors. Need not be the same set of HKLs for obs - and calc — but should sample the same resolution range. - e_obs, e_calc : (N,) real - Normalized structure-factor amplitudes (E-values). Pass |F|² instead - if you want a strict Patterson rotation function. - L, P : int, default 48 / 24 - Bandlimit and shell count. - d_min, d_max : float - Resolution limits in Å (s_min = 1/d_max, s_max = 1/d_min). - n_peaks : int, default 200 - Number of rotation peaks to return. - refine_subvoxel : bool, default True - If True, refine the top `n_refine` peaks by gradient ascent on the - pointwise rotation function (autograd, no FFT). - n_refine : int - Number of peaks to refine. - sigma_threshold : float - Minimum (peak - mean)/std for a peak to be kept. - weights : (P,) real, optional - Per-shell weights in the cross-correlation. Default: shell-count weighted, - or (when `auto_variance_weights=True`) inverse-sqrt of the empirical - per-shell variance of the obs Patterson coefficient. - auto_variance_weights : bool, default True - Compute `weights = 1/√Var(e_obs)_p` per shell when no explicit `weights` - are supplied. This is the LERF1-style empirical variance correction — - downweights shells whose observed Patterson coefficient is dominated by - intermolecular contributions (high-symmetry crystals). - - Returns - ------- - C : complex tensor (n_gamma, n_beta, n_alpha) - Rotation function on Euler grid (Edmonds order). - alphas, betas, gammas : real 1-D tensors - Grid coordinates. - peaks : list of RotationPeak (sorted by descending score) - """ - # NOTE: `zsymm` is applied only to the OBSERVED side. The observed - # Patterson is spacegroup-invariant by physics, so SH coefficients with - # m not divisible by ZSYMM are noise. The calc Patterson comes from a - # rotated P1 ensemble that is NOT spacegroup-invariant — those m's - # carry real signal and must be kept. Phaser does the same asymmetric - # filtering (DataMR.cc applies the m-filter; Ensemble.cc does not). - f_obs = compute_ball_harmonic_coefficients( - s_obs, e_obs, L=L, P=P, d_min=d_min, d_max=d_max, - zsymm=zsymm, skip_odd_l=skip_odd_l, - ) - # Share shell edges between obs and calc — required for the cross-correlation - # to be meaningful (same radial binning on both sides). - f_calc = compute_ball_harmonic_coefficients( - s_calc, e_calc, L=L, P=P, shell_edges=f_obs.shell_edges, - zsymm=1, skip_odd_l=skip_odd_l, - ) - - if auto_variance_weights and weights is None: - s_mag_obs = s_obs.norm(dim=-1) - if d_min is not None or d_max is not None: - s_lo = 1.0 / d_max if d_max is not None else 0.0 - s_hi = 1.0 / d_min if d_min is not None else float("inf") - keep_obs = (s_mag_obs >= s_lo) & (s_mag_obs <= s_hi) - s_mag_v = s_mag_obs[keep_obs] - e_obs_v = e_obs[keep_obs] - else: - s_mag_v = s_mag_obs - e_obs_v = e_obs - shell_idx_v = assign_shells(s_mag_v, f_obs.shell_edges) - shell_var = compute_patterson_shell_variance( - e_obs_v.to(torch.float64), shell_idx_v, P=P, - ) - w = 1.0 / shell_var.sqrt() - w = w * (P / w.sum().clamp(min=1e-30)) - weights = w.to(f_obs.f_plm.real.dtype) - - xi = compute_ball_cross_correlation_coefficients(f_obs, f_calc, weights=weights) - C, alphas, betas, gammas = evaluate_rotation_function_grid( - xi, L, n_alpha=2 * L, n_beta=2 * L, n_gamma=2 * L, - ) - peaks = find_rotation_peaks( - C, alphas, betas, gammas, - n_peaks=n_peaks, sigma_threshold=sigma_threshold, - ) - if refine_subvoxel and len(peaks) > 0: - head = peaks[:n_refine] - head = refine_peaks_subvoxel(head, C, alphas, betas, gammas) - peaks = head + peaks[n_refine:] - peaks.sort(key=lambda r: r.score, reverse=True) - return C, alphas, betas, gammas, peaks - - -# ============================================================================= -# Convenience: Euler ↔ matrix using the Edmonds ZYZ convention -# ============================================================================= - - -def rotation_matrix_from_edmonds_euler_batch( - alpha: torch.Tensor, beta: torch.Tensor, gamma: torch.Tensor, -) -> torch.Tensor: - """ - Vectorised Edmonds ZYZ Euler → R: R = R_z(α) R_y(β) R_z(γ). - - `alpha`, `beta`, `gamma`: identically-shaped real tensors. Returns - `(..., 3, 3)` in the same dtype/device as the inputs. Equivalent to - calling `rotation_matrix_from_edmonds_euler` per Euler triple, but - avoids 9 small-tensor allocations per call. - """ - ca, sa = torch.cos(alpha), torch.sin(alpha) - cb, sb = torch.cos(beta), torch.sin(beta) - cg, sg = torch.cos(gamma), torch.sin(gamma) - zero = torch.zeros_like(alpha) - one = torch.ones_like(alpha) - Rz_a = torch.stack([ - torch.stack([ca, -sa, zero], dim=-1), - torch.stack([sa, ca, zero], dim=-1), - torch.stack([zero, zero, one], dim=-1), - ], dim=-2) - Ry_b = torch.stack([ - torch.stack([cb, zero, sb], dim=-1), - torch.stack([zero, one, zero], dim=-1), - torch.stack([-sb, zero, cb], dim=-1), - ], dim=-2) - Rz_c = torch.stack([ - torch.stack([cg, -sg, zero], dim=-1), - torch.stack([sg, cg, zero], dim=-1), - torch.stack([zero, zero, one], dim=-1), - ], dim=-2) - return Rz_a @ Ry_b @ Rz_c - - -def rotation_matrix_from_edmonds_euler( - alpha: float, beta: float, gamma: float, dtype=torch.float64, -) -> torch.Tensor: - """ - Build R = R_z(α) R_y(β) R_z(γ) (Edmonds active ZYZ). - - Equivalent to passing `[γ, β, α]` to - `torchref.alignment.transform.rotation_matrix_from_euler`. - """ - ca, sa = math.cos(alpha), math.sin(alpha) - cb, sb = math.cos(beta), math.sin(beta) - cg, sg = math.cos(gamma), math.sin(gamma) - Rz_a = torch.tensor([[ca, -sa, 0.0], [sa, ca, 0.0], [0.0, 0.0, 1.0]], dtype=dtype) - Ry_b = torch.tensor([[cb, 0.0, sb], [0.0, 1.0, 0.0], [-sb, 0.0, cb]], dtype=dtype) - Rz_c = torch.tensor([[cg, -sg, 0.0], [sg, cg, 0.0], [0.0, 0.0, 1.0]], dtype=dtype) - return Rz_a @ Ry_b @ Rz_c - - -def edmonds_euler_from_rotation_matrix(R: torch.Tensor) -> Tuple[float, float, float]: - """ - Recover (α, β, γ) such that R = R_z(α) R_y(β) R_z(γ). - Returns angles in radians; α, γ ∈ [0, 2π), β ∈ [0, π]. - Singular when β = 0 or π (only α+γ is determined); we set γ=0 in those cases. - """ - R = R.to(torch.float64) - cos_beta = R[2, 2].clamp(-1.0, 1.0).item() - beta = math.acos(cos_beta) - sin_beta = math.sin(beta) - if abs(sin_beta) < 1e-9: - # Degenerate: only α+γ is determined. Set γ=0. - alpha = math.atan2(R[1, 0].item(), R[0, 0].item()) - gamma = 0.0 - else: - alpha = math.atan2(R[1, 2].item(), R[0, 2].item()) - gamma = math.atan2(R[2, 1].item(), -R[2, 0].item()) - alpha = alpha % (2.0 * math.pi) - gamma = gamma % (2.0 * math.pi) - return alpha, beta, gamma - - -def rotation_angular_distance_deg(R1: torch.Tensor, R2: torch.Tensor) -> float: - """Geodesic distance on SO(3) in degrees: arccos((tr(R1 R2^T) - 1)/2).""" - R = R1.to(torch.float64) @ R2.to(torch.float64).T - tr = (R[0, 0] + R[1, 1] + R[2, 2]).clamp(-1.0, 3.0).item() - cos_a = max(-1.0, min(1.0, (tr - 1.0) / 2.0)) - return math.degrees(math.acos(cos_a)) diff --git a/torchref/alignment/frf/data_mr.py b/torchref/alignment/frf/data_mr.py index 582d14c2..4e88b40a 100644 --- a/torchref/alignment/frf/data_mr.py +++ b/torchref/alignment/frf/data_mr.py @@ -1,9 +1,9 @@ """Bessel-radial × spherical-harmonic expansion of obs and calc Pattersons. Mirrors ``DataMR::dataMR_FRF`` (DataMR.cc) and the helper sums in -``Ensemble.cc``. We import the validated ``bessel_sh_expand`` from -``torchref.alignment.phaser_frf`` and add the obs/calc cross-correlation -contraction that's the input to ``SiteListAng::get_FRF``. +``Ensemble.cc``. Contains both the spherical-Bessel table (Miller +recurrence) and the chunked obs/calc Bessel×SH expansion fed into the +cross-correlation that ``SiteListAng::get_FRF`` consumes. Citations: * Bessel-radial × SH expansion: DataMR.cc:993, 1107-1117 @@ -14,20 +14,81 @@ from __future__ import annotations import math +import os +import time import torch -from .phaser_frf import spherical_bessel_table -from ..sh import evaluate_ylm +_PROFILE = bool(os.environ.get("FRF_PROFILE")) + +from ..sh import _bar_legendre_recurrence, evaluate_ylm from .types import BesselSHCoefficients __all__ = [ "bessel_sh_expand", "cross_correlate_xi", + "spherical_bessel_table", ] +def spherical_bessel_table( + x: torch.Tensor, + u_max: int, + n_extra: int = 25, +) -> torch.Tensor: + """Tabulate spherical Bessel ``j_u(x)`` for ``u ∈ [0, u_max]``, batched over ``x``. + + Uses Miller's downward recurrence (the standard stable choice for + ``j_n(x)`` with ``n > x``): + + j_{u-1}(x) = (2u + 1) / x · j_u(x) − j_{u+1}(x) + + Seed: ``n_start = u_max + n_extra``, ``j_{n_start+1} = 0``, + ``j_{n_start} = 1`` (unnormalised), recur down to ``j_0``, then + renormalise using the exact ``j_0(x) = sin(x) / x``. Float64 + internally for accuracy at moderate ``u/x`` ratios; cast back to + input dtype on return. + + Returns + ------- + j_table : torch.Tensor + Shape ``(*x.shape, u_max + 1)``, dtype = ``x.dtype``. + """ + real_dtype = x.dtype + device = x.device + x64 = x.to(torch.float64) + safe_x = x64.clamp(min=1e-30) + inv_x = 1.0 / safe_x + + n_start = max(u_max + n_extra, u_max + 2) + j_high = torch.zeros_like(x64) + j_mid = torch.ones_like(x64) + j_table = torch.zeros( + (u_max + 1, *x64.shape), dtype=torch.float64, device=device, + ) + + for n in range(n_start, 0, -1): + j_low = (2.0 * n + 1.0) * inv_x * j_mid - j_high + if n - 1 <= u_max: + j_table[n - 1] = j_low + j_high = j_mid + j_mid = j_low + + true_j0 = torch.sin(x64) * inv_x + true_j0 = torch.where(x64 < 1e-30, torch.ones_like(x64), true_j0) + computed_j0 = j_table[0] + safe_j0 = torch.where( + computed_j0.abs() < 1e-30, torch.ones_like(computed_j0), computed_j0, + ) + scale = true_j0 / safe_j0 + j_table = j_table * scale.unsqueeze(0) + + perm = list(range(1, j_table.dim())) + [0] + j_table = j_table.permute(*perm).contiguous() + return j_table.to(real_dtype) + + def bessel_sh_expand( s_vectors: torch.Tensor, intensity: torch.Tensor, @@ -37,6 +98,7 @@ def bessel_sh_expand( zsymm: int = 1, enforce_friedel: bool = True, chunk_size: int = -1, + compute_dtype: "torch.dtype | None" = None, ) -> BesselSHCoefficients: """Phaser-style ``c_nlm = Σ_h Y*_lm(ŝ) · I · sqrt(2u+1) · j_u(h)/h``. @@ -58,6 +120,15 @@ def bessel_sh_expand( chunk_size : int Reflections per chunk. ``-1`` (default) auto-sizes so the per-chunk Y_lm block ``(chunk, L, 2L-1)`` stays near ~256 MB. + compute_dtype : torch.dtype, optional + Complex dtype for the dominant per-chunk einsum (the radial × SH + contraction — the FRF's FLOP bottleneck). Default ``None`` uses the + full-precision complex dtype matching the input. Passing + ``torch.complex64`` runs the contraction in single precision (a large + speedup on GPUs where FP64 is rate-limited), while the Bessel recurrence + and Legendre/Y_lm precompute stay at the input precision and the + cross-chunk accumulator stays at full precision — so only the contraction + loses precision, not the recurrences or the running sum. """ assert s_vectors.dim() == 2 and s_vectors.shape[-1] == 3 assert intensity.dim() == 1 and intensity.shape[0] == s_vectors.shape[0] @@ -71,6 +142,21 @@ def bessel_sh_expand( raise TypeError(f"Unsupported real dtype: {real_dtype}") device = s_vectors.device + # Working precision for the per-chunk precompute (angles, Y_lm, Bessel + # weights) and the contraction. When the caller opts into complex64 we run + # the whole chunk in float32 — the dominant cost on GPUs where FP64 is + # rate-limited is not just the einsum but also the Legendre/Y_lm recurrence, + # so both must drop to single precision to matter. The spherical-Bessel + # downward recurrence keeps its float64 internals (it is the most + # cancellation-prone step) and the cross-chunk accumulator stays at full + # complex precision. + if compute_dtype == torch.complex64: + comp_real = torch.float32 + elif compute_dtype == torch.complex128: + comp_real = torch.float64 + else: + comp_real = real_dtype + if enforce_friedel: s_vectors = torch.cat([s_vectors, -s_vectors], dim=0) intensity = torch.cat([intensity, intensity], dim=0) @@ -82,48 +168,134 @@ def bessel_sh_expand( N_radial = (lmax_even - 2) // 2 + 1 u_max = lmax_even + 1 - # (l, n) -> u = l + 2n + 1 and the sqrt(2u+1) weight, precomputed once. + # (l, n) -> u = l + 2n + 1 and the sqrt(2u+1) weight, precomputed once as + # index tensors so the per-chunk Bessel fill is a single advanced-indexed + # assignment instead of a Python loop over ~N_radial·(lmax/2) (l, n) pairs. even_ls = list(range(2, lmax_even + 1, 2)) - ln_u = [] # (l, n, u, sqrt(2u+1)) + l_list, n_list, u_list, w_list = [], [], [], [] for l in even_ls: n_l = (lmax_even - l) // 2 + 1 for n in range(n_l): u = l + 2 * n + 1 - ln_u.append((l, n, u, math.sqrt(float(2 * u + 1)))) + l_list.append(l) + n_list.append(n) + u_list.append(u) + w_list.append(math.sqrt(float(2 * u + 1))) + l_idx = torch.tensor(l_list, dtype=torch.long, device=device) + n_idx = torch.tensor(n_list, dtype=torch.long, device=device) + u_idx = torch.tensor(u_list, dtype=torch.long, device=device) + w_vec = torch.tensor(w_list, dtype=comp_real, device=device) + # Only even degrees l ∈ [2, lmax_even] carry signal (odd-l and l=0 are zeroed + # by Patterson centrosymmetry). Compute / contract Y_lm on these rows only — + # the assembly + einsum are the bottleneck, so this ~halves them. The full + # c_nlm keeps the (L, ...) shape with odd/zero rows left at zero. + even_l_idx = torch.tensor(even_ls, dtype=torch.long, device=device) M = s_vectors.shape[0] - if chunk_size <= 0: - # target ~256 MB for the complex (chunk, L, 2L-1) Y block - chunk_size = int(max(256, min(8192, 16_000_000 // max(1, L * (2 * L - 1))))) + einsum_dtype = compute_dtype if compute_dtype is not None else complex_dtype - c_nlm = torch.zeros( - (N_radial, L, 2 * L - 1), dtype=complex_dtype, device=device, - ) + prof = {"cluster": 0.0, "dbuild": 0.0, "bessel": 0.0, "ylm": 0.0, "einsum": 0.0} if _PROFILE else None + + def _tick(t0): + if device.type == "cuda": + torch.cuda.synchronize() + return time.perf_counter() - t0 + + if _PROFILE: + t0 = time.perf_counter() + + # ---- Cluster reflections by (|s|, cosθ) --------------------------------- + # Both the radial Bessel weight (a function of |s|) and the Legendre barP (a + # function of cosθ) are constant within a cluster, so the per-reflection SH + # expansion factorises — only the azimuthal phase e^{imφ} and the intensity + # vary inside a cluster. This generalises Phaser's cosθ clustering + # (DataMR.cc:918, HKL_clustered) by ALSO factoring the radial term, which + # collapses the dominant contraction from O(M·L³) to O(n_clusters·L³). On the + # dense P1 calc box (and on cubic/tetragonal obs lattices) reflections sharing + # (H²+K²+L², L_z) land in one cluster → n_clusters ≪ M (~16× fewer). For a + # non-degenerate reflection set it degrades gracefully to ≈ the per-reflection + # cost (clusters are singletons), with no change in the result. + s_mag_all = s_vectors.norm(dim=-1).clamp(min=1e-30) + cos_all = (s_vectors[..., 2] / s_mag_all).clamp(min=-1.0, max=1.0) + phi_all = torch.atan2(s_vectors[..., 1], s_vectors[..., 0]) + KSCALE = 10_000_000 # ~1e-7 grouping resolution → exact for grid-degenerate sets + k_s = (s_mag_all * KSCALE).round().to(torch.int64) + k_c = (cos_all * KSCALE).round().to(torch.int64) + KSCALE # shift ≥ 0 + key = k_s * (2 * KSCALE + 1) + k_c + uniq_key, inverse = torch.unique(key, return_inverse=True) + n_clusters = int(uniq_key.shape[0]) + # Per-cluster representative geometry (all members are equal by construction). + rep_cos = torch.empty(n_clusters, dtype=comp_real, device=device) + rep_smag = torch.empty(n_clusters, dtype=comp_real, device=device) + rep_cos[inverse] = cos_all.to(comp_real) + rep_smag[inverse] = s_mag_all.to(comp_real) + rep_sin = torch.sqrt((1.0 - rep_cos * rep_cos).clamp(min=0.0)) + if _PROFILE: + prof["cluster"] += _tick(t0); t0 = time.perf_counter() + + # ---- D[c, m] = Σ_{h∈c} I_h · conj(C(m, φ_h)) ---------------------------- + # Y_lm = barP_{l,|m|}(cosθ) · C(m, φ) with C(m,φ) = (-1)^m e^{imφ} (m≥0) / + # e^{imφ} (m<0) — the sh.evaluate_ylm convention. So conj(C) = sign(m) e^{-imφ} + # with sign(m)=(-1)^m for m≥0 else 1. D folds the φ-phase, sign and intensity, + # accumulated per cluster — O(M·L), the only remaining per-reflection work. + m_idx = torch.arange(-(L - 1), L, device=device) # (2L-1,) + # sign(m) = (-1)^m for m≥0 else 1. clamp(min=0) keeps the (unused) negative + # entries' base-pow well-defined (no NaN from (-1)^(neg float)). + sign_m = torch.where( + m_idx >= 0, + ((-1.0) ** m_idx.clamp(min=0).to(comp_real)), + torch.ones_like(m_idx, dtype=comp_real), + ).to(einsum_dtype) + Dc = torch.zeros((n_clusters, 2 * L - 1), dtype=einsum_dtype, device=device) + dchunk = 262_144 + mrow = m_idx.to(comp_real).unsqueeze(0) # (1, 2L-1) + for start in range(0, M, dchunk): + stop = min(start + dchunk, M) + ph = phi_all[start:stop].to(comp_real).unsqueeze(1) # (c, 1) + i_c = intensity[start:stop].to(comp_real).unsqueeze(1) # (c, 1) + e_neg = torch.exp((-1j) * (mrow * ph)).to(einsum_dtype) # (c, 2L-1) = e^{-imφ} + f = (i_c.to(einsum_dtype) * sign_m.unsqueeze(0)) * e_neg + Dc.index_add_(0, inverse[start:stop], f) + if _PROFILE: + prof["dbuild"] += _tick(t0); t0 = time.perf_counter() + + # ---- per-cluster Bessel (radial) ---------------------------------------- + x_c = (bessel_h_scale * rep_smag).clamp(min=1e-30) + j_all = spherical_bessel_table(x_c, u_max) # (n_clusters, u_max+1) + bessel = torch.zeros((n_clusters, L, N_radial), dtype=comp_real, device=device) + bessel[:, l_idx, n_idx] = w_vec.unsqueeze(0) * j_all[:, u_idx] / x_c.unsqueeze(-1) + bessel_e = bessel[:, even_l_idx, :].to(einsum_dtype) # (n_clusters, n_even, N_radial) + if _PROFILE: + prof["bessel"] += _tick(t0); t0 = time.perf_counter() + + # ---- per-cluster Legendre (even l), expanded over m via |m| -------------- + bar_P = _bar_legendre_recurrence(rep_cos, rep_sin, L) # (n_clusters, L, L) + bar_P_e = bar_P[:, even_l_idx, :] # (n_clusters, n_even, L) + abs_m = m_idx.abs() + if _PROFILE: + prof["ylm"] += _tick(t0); t0 = time.perf_counter() + + # ---- factored contraction over clusters --------------------------------- + # c[n,l,m] = Σ_c bessel[c,l,n] · barP[c,l,|m|] · D[c,m]. Chunk over clusters to + # bound the (n_even, 2L-1) intermediate when n_clusters is large (singletons). + c_e = torch.zeros((N_radial, len(even_ls), 2 * L - 1), dtype=einsum_dtype, device=device) + cbytes = (8 if einsum_dtype == torch.complex64 else 16) + cstep = max(1, min(n_clusters, 256_000_000 // max(1, cbytes * len(even_ls) * (2 * L - 1)))) + for cs in range(0, n_clusters, cstep): + ce = min(cs + cstep, n_clusters) + G = bar_P_e[cs:ce].to(einsum_dtype)[:, :, abs_m] * Dc[cs:ce].unsqueeze(1) + c_e += torch.einsum("cln,clm->nlm", bessel_e[cs:ce], G) + c_nlm = torch.zeros((N_radial, L, 2 * L - 1), dtype=complex_dtype, device=device) + c_nlm[:, even_l_idx, :] = c_e.to(complex_dtype) + if _PROFILE: + prof["einsum"] += _tick(t0) - for start in range(0, M, chunk_size): - stop = min(start + chunk_size, M) - s_c = s_vectors[start:stop] - i_c = intensity[start:stop] - s_mag = s_c.norm(dim=-1).clamp(min=1e-30) - s_hat = s_c / s_mag.unsqueeze(-1) - cos_theta = s_hat[..., 2].clamp(min=-1.0, max=1.0) - theta = torch.acos(cos_theta) - phi = torch.atan2(s_hat[..., 1], s_hat[..., 0]) - - x = (bessel_h_scale * s_mag).clamp(min=1e-30) # (c,) - j_all = spherical_bessel_table(x, u_max) # (c, u_max+1) - - bessel = torch.zeros((stop - start, L, N_radial), dtype=real_dtype, device=device) - for (l, n, u, w) in ln_u: - bessel[:, l, n] = w * j_all[:, u] / x - - Y = evaluate_ylm(theta, phi, L) # (c, L, 2L-1) - Y_w = torch.conj(Y) * i_c.to(complex_dtype).view(-1, 1, 1) - c_nlm += torch.einsum("hln,hlm->nlm", bessel.to(complex_dtype), Y_w) - - # Zero odd-l rows and l = 0 (Patterson centrosymmetry; Phaser drops l=0). - l_vals = torch.arange(L, device=device) - c_nlm[:, (l_vals % 2 == 1) | (l_vals == 0), :] = 0.0 + if _PROFILE: + tot = sum(prof.values()) + 1e-30 + print(f"[FRF_PROFILE] M={M} n_clusters={n_clusters} ({M/max(1,n_clusters):.1f}x) " + f"L={L} dtype={comp_real} | " + + " ".join(f"{k}={v*1000:.0f}ms({100*v/tot:.0f}%)" for k, v in prof.items()), + flush=True) # m-symmetry filter (observed side only; caller passes zsymm=1 for calc). if zsymm > 1: diff --git a/torchref/alignment/frf/french_wilson.py b/torchref/alignment/frf/french_wilson.py new file mode 100644 index 00000000..4dcaa3b3 --- /dev/null +++ b/torchref/alignment/frf/french_wilson.py @@ -0,0 +1,505 @@ +"""French–Wilson posterior + Luzzati DFAC chain. + +Pure ports of Phaser's ``lib/math_FrenchWilson.cc`` (centric/acentric +posterior moments via Parabolic-cylinder ratios) and the Halley-iteration +``getDfactor`` in ``lib/math_RiceLLG.cc``. The public entry point +:func:`french_wilson_preprocess` returns ``(eEobs, DFAC, sqrt_mean_F2)`` +from raw ``(F, σF, |s|, centric)``. + +Everything except ``french_wilson_preprocess`` is module-private; expose +the public name through :mod:`torchref.alignment.frf.preprocessing`. + +References (paths under +``…/reverse_engineering/phenix/.../phaser/src/``): +- ``lib/math_FrenchWilson.cc:8-178`` posterior ```` / ```` +- ``lib/math_RiceLLG.cc:12-250`` Rice-moment effective σA + Halley +- ``Dfactor.cc:87-93`` eEobs assembly + clamp +""" +from __future__ import annotations + +import torch + + +__all__ = ["french_wilson_preprocess"] + + +# ----------------------------------------------------------------------------- +# French-Wilson posterior expected values (math_FrenchWilson.cc) +# ----------------------------------------------------------------------------- + + +def _expectE_FW_acen(eosq, sigesq): + """ + Acentric posterior expected E from normalised observed intensity (eosq) + and its standard deviation (sigesq). Translates verbatim from Phaser's + `lib/math_FrenchWilson.cc:expectEFWacen` (lines 8-44). Vectorised NumPy. + + `eosq = Iobs / `, `sigesq = σIobs / `. + """ + import numpy as np + from scipy.special import erfc, pbdv + CROSS1, CROSS2 = -12.5, 18.0 + SQRT2 = np.sqrt(2.0) + x = (eosq - sigesq ** 2) / sigesq + xsqr = x * x + ee = np.empty_like(eosq) + m_neg = x < CROSS1 + if m_neg.any(): + xs = xsqr[m_neg] + num = (-916620705. + xs * + (91891800. + xs * + (-11531520. + xs * + (1935360. + xs * + (-491520. + xs * 262144.))))) + den = (-495452160. + xs * + (55050240. + xs * + (-7864320. + xs * + (1572864. + xs * + (-524288. + xs * 524288.))))) + ee[m_neg] = np.sqrt(-np.pi * sigesq[m_neg] / x[m_neg]) * num / den + m_pos = x > CROSS2 + if m_pos.any(): + xs = xsqr[m_pos] + num = (-45045. + 32. * xs * + (-315. + 8. * xs * + (-15. - 16. * xs + 128. * xs * xs))) + ee[m_pos] = (np.sqrt(sigesq[m_pos]) * num / + (32768. * x[m_pos] ** 7.5)) + m_mid = ~(m_neg | m_pos) + if m_mid.any(): + xm = x[m_mid] + pcd, _ = pbdv(-1.5, -xm) + ee[m_mid] = (np.sqrt(sigesq[m_mid] / 2.0) * np.exp(-xm * xm / 4.0) * + pcd / erfc(-xm / SQRT2)) + return ee + + +def _expectEsq_FW_acen(eosq, sigesq): + """Acentric posterior . From `expectEsqFWacen` (lines 46-78).""" + import numpy as np + from scipy.special import erfc + CROSS1, CROSS2 = -8.9, 5.7 + SQRT2_BY_PI = np.sqrt(2.0 / np.pi) + SQRT2 = np.sqrt(2.0) + eesq_base = eosq - sigesq ** 2 + x = eesq_base / (SQRT2 * sigesq) + xsqr = x * x + eesq = eesq_base.copy() + m_neg = x < CROSS1 + if m_neg.any(): + xs = xsqr[m_neg] + num = (-135135. + xs * (20790. + xs * (-3780. + xs * + (840. + xs * (-240. + xs * (96. - xs * 64.)))))) + den = (-135135. + xs * (20790. + xs * (-3780. + xs * + (840. + xs * (-240. + xs * (96. + xs * + (-64. + xs * 128.))))))) + eesq[m_neg] = eesq_base[m_neg] * num / den + m_mid = (x >= CROSS1) & (x <= CROSS2) + if m_mid.any(): + xm = x[m_mid] + eesq[m_mid] = (eesq_base[m_mid] + + SQRT2_BY_PI * sigesq[m_mid] / + (np.exp(xm * xm) * erfc(-xm))) + return eesq + + +def _expectE_FW_cen(eosq, sigesq): + """Centric posterior . From `expectEFWcen` (lines 80-113).""" + import numpy as np + from scipy.special import pbdv + CROSS1, CROSS2 = -17.5, 17.5 + SQRTPI = np.sqrt(np.pi) + x = sigesq / 2.0 - eosq / sigesq + xsqr = x * x + pcdratio = np.empty_like(x) + m_neg = x < CROSS1 + if m_neg.any(): + xn, xs = x[m_neg], xsqr[m_neg] + pcdratio[m_neg] = ((1024. * SQRTPI * (-xn) ** 6.5) / + (3465. + xs * + (840. + xs * + (384. + xs * 1024.)))) + m_pos = x > CROSS2 + if m_pos.any(): + xp, xs = x[m_pos], xsqr[m_pos] + num = (3440640. + xs * + (-491520. + xs * + (98304. + xs * + (-32768. + xs * 32768.)))) + den = (675675. + xs * + (-110880. + xs * + (26880. + xs * + (-12288. + xs * 32768.)))) + pcdratio[m_pos] = num / (den * np.sqrt(xp)) + m_mid = ~(m_neg | m_pos) + if m_mid.any(): + xm = x[m_mid] + d_neg1, _ = pbdv(-1.0, xm) + d_neghalf, _ = pbdv(-0.5, xm) + pcdratio[m_mid] = d_neg1 / d_neghalf + return np.sqrt(sigesq / np.pi) * pcdratio + + +def _expectEsq_FW_cen(eosq, sigesq): + """Centric posterior . From `expectEsqFWcen` (lines 115-152).""" + import numpy as np + from scipy.special import pbdv + CROSS1, CROSS2 = -17.5, 17.5 + x = sigesq / 2.0 - eosq / sigesq + xsqr = x * x + pcdratio = np.empty_like(x) + m_neg = x < CROSS1 + if m_neg.any(): + xn, xs = x[m_neg], xsqr[m_neg] + num = (45045. + xs * + (10080. + xs * + (3840. + xs * + (4096. - xs * 32768.)))) + den = xn * (55440. + xs * + (13440. + xs * + (6144. + xs * 16384.))) + pcdratio[m_neg] = num / den + m_pos = x > CROSS2 + if m_pos.any(): + xp, xs = x[m_pos], xsqr[m_pos] + num = (11486475. + xs * + (-1441440. + xs * + (241920. + xs * + (-61440. + xs * 32768.)))) + den = xp * (675675. + xs * + (-110880. + xs * + (26880. + xs * + (-12288. + xs * 32768.)))) + pcdratio[m_pos] = num / den + m_mid = ~(m_neg | m_pos) + if m_mid.any(): + xm = x[m_mid] + d_neg15, _ = pbdv(-1.5, xm) + d_neghalf, _ = pbdv(-0.5, xm) + pcdratio[m_mid] = d_neg15 / d_neghalf + return sigesq * pcdratio / 2.0 + + +def _french_wilson_posterior(eosq, sigesq, centric_mask): + """Wrap centric/acentric branches. + + Phaser `expectEFW` / `expectEsqFW` (lines 154-178): if sigesq <= 0 the + measurement is treated as exact and (eEFW, eEsqFW) = (sqrt(eosq), eosq). + """ + import numpy as np + eEFW = np.empty_like(eosq) + eEsqFW = np.empty_like(eosq) + zero_sig = sigesq <= 0.0 + if zero_sig.any(): + eEFW[zero_sig] = np.sqrt(np.maximum(eosq[zero_sig], 0.0)) + eEsqFW[zero_sig] = np.maximum(eosq[zero_sig], 0.0) + valid = ~zero_sig + if valid.any(): + cen = centric_mask & valid + acen = (~centric_mask) & valid + if cen.any(): + eEFW[cen] = _expectE_FW_cen(eosq[cen], sigesq[cen]) + eEsqFW[cen] = _expectEsq_FW_cen(eosq[cen], sigesq[cen]) + if acen.any(): + eEFW[acen] = _expectE_FW_acen(eosq[acen], sigesq[acen]) + eEsqFW[acen] = _expectEsq_FW_acen(eosq[acen], sigesq[acen]) + return eEFW, eEsqFW + + +# ----------------------------------------------------------------------------- +# DFAC via Halley iteration (math_RiceLLG.cc:getDfactor) +# ----------------------------------------------------------------------------- + + +def _i0e_full(x): + """Phaser's `eBesselI0(x) = I0(x)·exp(-|x|)`. Symmetric in x.""" + import numpy as np + from scipy.special import i0e + return i0e(np.abs(x)) + + +def _i1e_full(x): + """Phaser's `eBesselI1(x) = I1(x)·exp(-|x|)`. Antisymmetric in x.""" + import numpy as np + from scipy.special import i1e + return np.sign(x) * i1e(np.abs(x)) + + +def _effSigaRoot_acen(ee, eesq, sa): + """`effSigaRootAcen` (math_RiceLLG.cc:12-34).""" + import numpy as np + sigbsqr = 1.0 - sa * sa + x = 0.5 * (eesq - sigbsqr) / sigbsqr + return (np.sqrt(np.pi * sigbsqr) / (2.0 * sigbsqr) * + (eesq * _i0e_full(x) + (eesq - sigbsqr) * _i1e_full(x)) - ee) + + +def _deffSigaRoot_acen(eesq, sa): + """`deffSigaRootAcen_by_dsa` (lines 36-52).""" + import numpy as np + sigbsqr = 1.0 - sa * sa + x = 0.5 * (eesq - sigbsqr) / sigbsqr + return np.sqrt(np.pi / sigbsqr) * (sa / 2.0) * _i1e_full(x) + + +def _d2effSigaRoot_acen(eesq, sa): + """`d2effSigaRootAcen_by_dsa2` (lines 54-81).""" + import numpy as np + sigasqr = sa * sa + sigapow4 = sigasqr * sigasqr + sigbsqr = 1.0 - sigasqr + xnum = eesq - sigbsqr + x = 0.5 * xnum / sigbsqr + out = np.empty_like(eesq) + big = xnum > 1e-10 + if big.any(): + I0 = _i0e_full(x[big]) + I1 = _i1e_full(x[big]) + out[big] = (np.sqrt(np.pi / sigbsqr[big]) / (2.0 * sigbsqr[big] ** 2) * + (eesq[big] * sigasqr[big] * I0 + + (eesq[big] - 1.0 - (-2.0 + eesq[big] * (2.0 + eesq[big])) * sigasqr[big] + + (eesq[big] - 1.0) * sigapow4[big]) * I1 / xnum[big])) + small = ~big + if small.any(): + samin = np.sqrt(np.maximum(1.0 - eesq[small], 0.0)) + out[small] = (np.sqrt(np.pi) * samin * + ((3.0 + samin * samin) * sa[small] - + 2.0 * (samin + samin ** 3)) / + (4.0 * eesq[small] ** 2.5)) + return out + + +def _effSigaRoot_cen(ee, eesq, sa): + """`effSigaRootCen` (lines 83-105).""" + import numpy as np + from scipy.special import erf + sigbsqr = 1.0 - sa * sa + x = 0.5 * (eesq - sigbsqr) / sigbsqr + x_safe = np.maximum(x, 0.0) + return (np.exp(-x) * np.sqrt(2.0 * sigbsqr / np.pi) + + np.sqrt(np.maximum(eesq - sigbsqr, 0.0)) * erf(np.sqrt(x_safe)) - ee) + + +def _deffSigaRoot_cen(eesq, sa): + """`deffSigaRootCen_by_dsa` (lines 107-130).""" + import numpy as np + from scipy.special import erf + sigbsqr = 1.0 - sa * sa + xnum = eesq - sigbsqr + x = 0.5 * xnum / sigbsqr + out = np.empty_like(eesq) + big = np.abs(xnum) > 1e-10 + if big.any(): + x_safe = np.maximum(x[big], 0.0) + out[big] = (sa[big] * erf(np.sqrt(x_safe)) / + np.sqrt(np.maximum(xnum[big], 1e-30)) - + np.exp(-x[big]) * np.sqrt(2.0 * sigbsqr[big] / np.pi) * + sa[big] / sigbsqr[big]) + small = ~big + if small.any(): + out[small] = (xnum[small] * np.sqrt(2.0 / np.pi) * sa[small] / + (3.0 * sigbsqr[small] ** 1.5)) + return out + + +def _d2effSigaRoot_cen(eesq, sa): + """`d2effSigaRootCen_by_dsa2` (lines 132-159).""" + import numpy as np + from scipy.special import erf + sigasqr = sa * sa + sigapow4 = sigasqr * sigasqr + sigbsqr = 1.0 - sigasqr + xnum = eesq - sigbsqr + x = 0.5 * xnum / sigbsqr + sigbsqrtpi = np.sqrt(np.pi * sigbsqr) + out = np.empty_like(eesq) + big = np.abs(xnum) > 1e-10 + if big.any(): + x_safe = np.maximum(x[big], 0.0) + d2num = ((eesq[big] - 1.0) * sigbsqr[big] ** 2 * sigbsqrtpi[big] * + erf(np.sqrt(x_safe))) + exp_part = np.where( + x[big] < 20.0, + np.sqrt(np.maximum(2.0 * xnum[big], 0.0)) * np.exp(-x[big]) * + (1.0 - eesq[big] + sigasqr[big] * + (eesq[big] + eesq[big] ** 2 - 2.0) + sigapow4[big]), + np.zeros_like(x[big]), + ) + d2num = d2num + exp_part + out[big] = d2num / (sigbsqr[big] ** 2 * + np.maximum(xnum[big], 1e-30) ** 1.5 * + sigbsqrtpi[big]) + small = ~big + if small.any(): + out[small] = (np.sqrt(2.0 / np.pi) * sa[small] / + (1.5 * sigbsqr[small] ** 1.5)) + return out + + +def _get_dfactor_vectorised(ee_np, eesq_np, centric_np): + """Vectorised port of Phaser's ``math_RiceLLG.cc:getDfactor`` (lines 191-250). + + Halley's method with bisection fallback, run over all reflections in + parallel. Each reflection has its own bracket ``[dflo, dfhi]``. Returns a + ``(N,)`` numpy float64 array of DFAC values in ``(0, 1)``. + """ + import numpy as np + + EPS1 = 1e-7 + EPS2 = 1e-10 + MAXDFAC = 1.0 - EPS1 + + ee = np.asarray(ee_np, dtype=np.float64) + eesq = np.asarray(eesq_np, dtype=np.float64) + cen = np.asarray(centric_np, dtype=bool) + N = ee.shape[0] + + out = np.ones(N, dtype=np.float64) + has_err = (eesq - ee * ee) > 0.0 + + if not has_err.any(): + return out + + ee_a, eesq_a, cen_a = ee[has_err], eesq[has_err], cen[has_err] + dflo = np.maximum(np.sqrt(np.maximum(1.0 - np.minimum(eesq_a, 1.0), 0.0)) + EPS1, EPS1) + dfhi = np.full_like(dflo, MAXDFAC) + + early = dflo >= MAXDFAC + if early.any(): + pass + + dfmid = 0.5 * (dflo + dfhi) + fmid = np.empty_like(dfmid) + if cen_a.any(): + fmid[cen_a] = _effSigaRoot_cen(ee_a[cen_a], eesq_a[cen_a], dfmid[cen_a]) + if (~cen_a).any(): + fmid[~cen_a] = _effSigaRoot_acen(ee_a[~cen_a], eesq_a[~cen_a], dfmid[~cen_a]) + + active = ~early + for _ in range(50): + if not active.any(): + break + conv = (dfhi - dflo) <= EPS1 + conv |= np.abs(fmid) <= EPS2 + active = active & ~conv + if not active.any(): + break + + slope = np.empty_like(dfmid) + curve = np.empty_like(dfmid) + cen_act = cen_a & active + acen_act = (~cen_a) & active + if cen_act.any(): + slope[cen_act] = _deffSigaRoot_cen(eesq_a[cen_act], dfmid[cen_act]) + curve[cen_act] = _d2effSigaRoot_cen(eesq_a[cen_act], dfmid[cen_act]) + if acen_act.any(): + slope[acen_act] = _deffSigaRoot_acen(eesq_a[acen_act], dfmid[acen_act]) + curve[acen_act] = _d2effSigaRoot_acen(eesq_a[acen_act], dfmid[acen_act]) + + denom_halley = 2.0 * (slope ** 2 - fmid * curve) + use_halley = (curve > 0.0) & (np.abs(denom_halley) > 1e-30) + step = np.where( + use_halley, + 2.0 * fmid * slope / np.where(use_halley, denom_halley, 1.0), + fmid * slope, + ) + dfnew = dfmid - step + in_bracket = (dfnew > dflo) & (dfnew < dfhi) + dfmid_new = np.where(in_bracket, dfnew, 0.5 * (dflo + dfhi)) + + dfmid = np.where(active, dfmid_new, dfmid) + + if cen_act.any(): + fmid[cen_act] = _effSigaRoot_cen(ee_a[cen_act], eesq_a[cen_act], + dfmid[cen_act]) + if acen_act.any(): + fmid[acen_act] = _effSigaRoot_acen(ee_a[acen_act], eesq_a[acen_act], + dfmid[acen_act]) + + below = (fmid < 0.0) & active + above = (fmid >= 0.0) & active + dflo = np.where(below, dfmid, dflo) + dfhi = np.where(above, dfmid, dfhi) + + out[has_err] = dfmid + return np.clip(out, EPS1, MAXDFAC) + + +def french_wilson_preprocess( + F: torch.Tensor, + sig_F: torch.Tensor, + s_mag: torch.Tensor, + centric: torch.Tensor, + *, + n_wilson_shells: int = 20, +) -> dict: + """Phaser-style preprocessing from raw ``(F, σF, centric)`` to ``(eEobs, DFAC)``. + + Implements the chain: + + 1. equal-count Wilson shells over ``s_mag`` + 2. per-shell ``_p`` (Phaser's ``SIGMAN.BINS``) + 3. per-reflection normalised intensity ``eosq = F² / `` and σ + ``sigesq = σI / ≈ 2·F·σF / `` + 4. French-Wilson posterior ``eEFW, eEsqFW`` (``math_FrenchWilson.cc``) + 5. DFAC via Halley iteration on Rice moments (``math_RiceLLG.cc``) + 6. ``eEobs = sqrt(eEsqFW + (DFAC²−1)/DFAC²)``, clamped to ≤10 + (Phaser ``Dfactor.cc:87-93``). + + Returns a dict with torch tensors back on the input device: + eEobs: (N,) effective normalised amplitude + DFAC : (N,) per-reflection D-factor ∈ [1e-7, 1−1e-7] + sqrt_mean_F2: (N,) per-reflection √_p + """ + import numpy as np + + device = F.device + F_np = F.detach().to("cpu").to(torch.float64).numpy() + sigF_np = sig_F.detach().to("cpu").to(torch.float64).numpy() + s_np = s_mag.detach().to("cpu").to(torch.float64).numpy() + cen_np = centric.detach().to("cpu").bool().numpy() + + sorted_idx = np.argsort(s_np) + edges_idx = np.linspace(0, len(s_np) - 1, n_wilson_shells + 1).round().astype(np.int64) + s_edges = s_np[sorted_idx][edges_idx] + s_edges[0] -= 1e-6 + s_edges[-1] += 1e-6 + shell_idx = np.clip( + np.searchsorted(s_edges, s_np, side="right") - 1, 0, n_wilson_shells - 1, + ) + F2 = F_np * F_np + mean_F2 = np.zeros(n_wilson_shells, dtype=np.float64) + counts = np.zeros(n_wilson_shells, dtype=np.int64) + np.add.at(mean_F2, shell_idx, F2) + np.add.at(counts, shell_idx, 1) + mean_F2 = mean_F2 / np.maximum(counts, 1) + mean_F2 = np.maximum(mean_F2, 1e-12) + mean_I_per_h = mean_F2[shell_idx] + sqrt_mean_F2 = np.sqrt(mean_I_per_h) + + eosq = F2 / mean_I_per_h + sigesq = 2.0 * F_np * sigF_np / mean_I_per_h + sigesq = np.maximum(sigesq, 0.0) + + eEFW, eEsqFW = _french_wilson_posterior(eosq, sigesq, cen_np) + bad = eEsqFW < eEFW * eEFW + if bad.any(): + eEsqFW[bad] = eEFW[bad] ** 2 + 1e-12 + + DFAC = _get_dfactor_vectorised(eEFW, eEsqFW, cen_np) + + dfsqr = DFAC * DFAC + eEobs_sqr = eEsqFW + (dfsqr - 1.0) / np.maximum(dfsqr, 1e-30) + eEobs_sqr = np.maximum(eEobs_sqr, 0.0) + eEobs = np.sqrt(eEobs_sqr) + clamp_mask = (eEobs > 10.0) & (eEsqFW > 1.0) + if clamp_mask.any(): + eEobs[clamp_mask] = 10.0 + DFAC[clamp_mask] = 1.0 / np.sqrt(np.maximum(eEsqFW[clamp_mask] - 99.0, 1e-30)) + DFAC[clamp_mask] = np.clip(DFAC[clamp_mask], 1e-7, 1.0 - 1e-7) + + return { + "eEobs": torch.from_numpy(eEobs).to(device=device, dtype=F.dtype), + "DFAC": torch.from_numpy(DFAC).to(device=device, dtype=F.dtype), + "sqrt_mean_F2": torch.from_numpy(sqrt_mean_F2).to(device=device, dtype=F.dtype), + } diff --git a/torchref/alignment/frf/peak_finder.py b/torchref/alignment/frf/peak_finder.py index cb2f5562..660d7e99 100644 --- a/torchref/alignment/frf/peak_finder.py +++ b/torchref/alignment/frf/peak_finder.py @@ -77,22 +77,33 @@ def _so3_greedy_nms( n = values.shape[0] if n == 0: return torch.empty(0, dtype=torch.int64, device=values.device) - order = torch.argsort(values, descending=True) - R_all = _euler_to_matrix_edmonds_zyz(alphas, betas, gammas) # (n, 3, 3) - R_all = R_all.to(torch.float64) - + # The greedy walk is inherently sequential and latency-bound; on GPU a + # per-iteration `.item()` sync would dominate. Move the (tiny) candidate + # rotations to CPU once and run the loop there with no device syncs, a + # preallocated kept-buffer (no repeated torch.stack), and a cosine threshold + # (no per-iteration arccos). Result is identical to the original distance test. + order = torch.argsort(values, descending=True).cpu().tolist() + R_all = ( + _euler_to_matrix_edmonds_zyz(alphas, betas, gammas) + .to(torch.float64).cpu() + ) # (n, 3, 3) + # angle > nms_radius ⇔ cos(angle) < cos(nms_radius); cos(angle) from trace. + cos_thresh = math.cos(math.radians(nms_radius_deg)) kept_idx: List[int] = [] - kept_R: List[torch.Tensor] = [] - for i_t in order.tolist(): + kept_R = torch.empty((keep_at_most, 3, 3), dtype=torch.float64) + count = 0 + for i_t in order: Ri = R_all[i_t] - if kept_R: - stack = torch.stack(kept_R, dim=0) # (k, 3, 3) - dists = _so3_angular_distance_deg(Ri.unsqueeze(0), stack) - if dists.min().item() <= nms_radius_deg: + if count > 0: + trace = torch.einsum("kij,ij->k", kept_R[:count], Ri) + cos_theta = ((trace - 1.0) * 0.5).clamp(min=-1.0, max=1.0) + # Some kept rotation within nms_radius (cos_theta > cos_thresh) → skip. + if bool((cos_theta > cos_thresh).any()): continue - kept_R.append(Ri) + kept_R[count] = Ri kept_idx.append(i_t) - if len(kept_idx) >= keep_at_most: + count += 1 + if count >= keep_at_most: break return torch.tensor(kept_idx, dtype=torch.int64, device=values.device) @@ -140,15 +151,15 @@ def find_rotation_peaks( keep_at_most=n_peaks, ) - peaks: List[RotationPeak] = [] - for k in kept.tolist(): - peaks.append( - RotationPeak( - alpha=float(a[k].item()), - beta=float(b[k].item()), - gamma=float(g[k].item()), - value=float(v[k].item()), - sigma=float(((v[k] - mean) / std).item()), - ) - ) + # Gather kept peaks and move to CPU once (avoids a per-peak device sync). + a_k = a[kept].cpu().tolist() + b_k = b[kept].cpu().tolist() + g_k = g[kept].cpu().tolist() + v_k = v[kept] + s_k = ((v_k - mean) / std).cpu().tolist() + v_k = v_k.cpu().tolist() + peaks: List[RotationPeak] = [ + RotationPeak(alpha=a_k[i], beta=b_k[i], gamma=g_k[i], score=v_k[i], sigma=s_k[i]) + for i in range(len(a_k)) + ] return peaks diff --git a/torchref/alignment/frf/phaser_frf.py b/torchref/alignment/frf/phaser_frf.py deleted file mode 100644 index c1509785..00000000 --- a/torchref/alignment/frf/phaser_frf.py +++ /dev/null @@ -1,1164 +0,0 @@ -""" -Phaser-faithful Fast Rotation Function (FRF) — a clean, parallel implementation. - -Translates Phaser's `DataMR.cc` (observed-side SH expansion) + `Ensemble.cc` -(model-side σA Eterm) + `FastRot.cc` (Wigner-D contraction + 2-D FFT) into -PyTorch. Designed to be benchmarked directly against torchref's existing -`ball_rotation_search` in -`tests/integration/alignment/benchmark_phaser_frf.py`. - -The single fundamental difference from `ball_rotation_search` is the **radial -basis**: this module expands the per-reflection Patterson onto **spherical -Bessel functions** rather than shell-step indicators. The Bessel basis is the -natural orthonormal radial basis on the unit ball (Crowther 1972, Navaza FAST); -shell-step is its coarsest piecewise-constant approximation. Phaser's -`DataMR.cc:1102–1134` accumulates - - clmn[l, m, n] = Σ_h Y*_{l,m}(ŝ_h) · I(h) · √u · j_u(2π·a·|s_h|) / (2π·a·|s_h|) - -with u = l + 2n − 1 and n ∈ [1, nmax(l)] where nmax(l) = (lmax − l)/2 + 1. - -Other Phaser-faithful choices we make: -- Even l only (Patterson is centrosymmetric); odd-l rows zeroed. -- m-symmetry filter on observed side only (`DataMR.cc:863-870, 1117`): - obs SH coefficients with `m % ZSYMM != 0` are dropped; calc side keeps all m. -- σA Eterm on calc side (`Ensemble.cc:36-46`): - `Eterm(s) = exp(-2π² · s² · ΔVRMS²)`. Per-reflection (not per-shell). -- LERF1-style observed intensity: `I_obs(h) = cweight · (E²−1)` with - `cweight = 1` for centric, `2` for acentric. -- Euler convention: this module computes in **Edmonds ZYZ** - `R = R_z(α) R_y(β) R_z(γ)` — same as the rest of torchref. Phaser's - internal convention is `R = R_z(γ) R_y(β) R_z(α)` (swapped α↔γ) but only - the rotation matrix matters for our rank-of-truth metric, so we never - swap explicitly — the orbit_rank metric in the benchmark builds rotation - matrices via `rotation_matrix_from_edmonds_euler` and the answer is - parameterization-independent. - -What we deliberately DO NOT yet implement (out of scope per plan): -- French-Wilson per-reflection DFAC. Phaser uses `(E²-V)/V² · DFAC²` with - per-reflection Luzzati DFAC ∈ [0.05, 10]. Here we use DFAC = 1. -- Adaptive (α, γ) grid sampling (Phaser's `pmax = 720·cos(β/2)/grid_sampling`). - We use uniform 2L grid. -- Axis permutation: Phaser detects the high-order axis and may permute the - coordinate frame. We use the z-preferred selection from - `get_high_order_axis`, applied without permutation. For groups where the - high-order axis is already z (most cases including P432 with body-diagonal - 3-folds along z), this is identical to Phaser. - -References (paths under -`/das/work/p17/p17490/Peter/Library/torchref/reverse_engineering/phenix/phenix-1.20-4459/modules/phaser/codebase/phaser/src/`): -- `DataMR.cc:863-1148` observed-side SH + intensity + ZSYMM filter -- `Ensemble.cc:36-46` model-side Eterm -- `FastRot.cc:30-217` Wigner contraction + FFT + Z-score -- `runMR_FRF.cc:546-587` peak picking + Z-score normalisation -""" - -from __future__ import annotations - -import math -from dataclasses import dataclass -from typing import List, Optional, Tuple - -import torch - -from .ball_search import ( - RotationPeak, - find_rotation_peaks_adaptive, - refine_peaks_subvoxel_adaptive, -) -from ..sh import ( - assign_shells, - equal_count_shell_edges, - evaluate_ylm, - get_high_order_axis, -) -from ..wigner import ( - AdaptiveRotationFunction, - evaluate_rotation_function_grid_adaptive, -) - - -# ============================================================================= -# Spherical Bessel j_u(x) via Miller's downward recurrence -# ============================================================================= - - -def spherical_bessel_table( - x: torch.Tensor, - u_max: int, - n_extra: int = 25, -) -> torch.Tensor: - """ - Tabulate spherical Bessel `j_u(x)` for u ∈ [0, u_max], batched over x. - - Uses Miller's downward recurrence (the standard stable choice for - `j_n(x)` with `n > x`): - - j_{u-1}(x) = (2u + 1) / x · j_u(x) − j_{u+1}(x) - - Seed: choose `n_start = u_max + n_extra`, set - `j_{n_start+1} = 0`, `j_{n_start} = 1` (unnormalised), recur down to - j_0, then renormalise using the exact `j_0(x) = sin(x) / x`. Float64 - internally for accuracy at moderate `u/x` ratios; cast back to input - dtype. - - Returns - ------- - j_table : torch.Tensor - Shape `(*x.shape, u_max + 1)`, dtype = `x.dtype`. - """ - real_dtype = x.dtype - device = x.device - x64 = x.to(torch.float64) - safe_x = x64.clamp(min=1e-30) - inv_x = 1.0 / safe_x - - n_start = max(u_max + n_extra, u_max + 2) - j_high = torch.zeros_like(x64) # j_{n_start + 1} - j_mid = torch.ones_like(x64) # j_{n_start} (arbitrary scale) - j_table = torch.zeros( - (u_max + 1, *x64.shape), dtype=torch.float64, device=device, - ) - - # Recur from n = n_start down to 1, generating j_{n-1} at each step. - for n in range(n_start, 0, -1): - j_low = (2.0 * n + 1.0) * inv_x * j_mid - j_high - if n - 1 <= u_max: - j_table[n - 1] = j_low - j_high = j_mid - j_mid = j_low - - # Normalise against the exact j_0(x) = sin(x)/x. Handle x ≈ 0 (where - # j_0(0) = 1, j_u(0) = 0 for u ≥ 1) so the renormalisation factor is - # well-defined. - true_j0 = torch.sin(x64) * inv_x - true_j0 = torch.where(x64 < 1e-30, torch.ones_like(x64), true_j0) - computed_j0 = j_table[0] - safe_j0 = torch.where( - computed_j0.abs() < 1e-30, torch.ones_like(computed_j0), computed_j0, - ) - scale = true_j0 / safe_j0 - j_table = j_table * scale.unsqueeze(0) - - # Re-arrange `(u_max+1, *x.shape) → (*x.shape, u_max+1)`. - perm = list(range(1, j_table.dim())) + [0] - j_table = j_table.permute(*perm).contiguous() - return j_table.to(real_dtype) - - -# ============================================================================= -# Phaser-style Bessel-radial SH expansion -# ============================================================================= - - -@dataclass -class BesselSHCoefficients: - """Phaser-style clmn coefficients. - - `c_nlm[n, l, m+L-1]` is the (n-th radial Bessel) × (l, m) coefficient. - """ - c_nlm: torch.Tensor # complex, shape (N_radial, L, 2L-1) - L: int # bandwidth (l ∈ [0, L), even-l only filled) - N_radial: int # number of Bessel radial terms = (lmax-2)/2 + 1 - bessel_h_scale: float # `h = bessel_h_scale · |s|` (Phaser uses lmax · d_min) - zsymm: int # m-symmetry filter applied (1 = none) - - -def bessel_sh_expand( - s_vectors: torch.Tensor, - intensity: torch.Tensor, - *, - L: int, - bessel_h_scale: float, - zsymm: int = 1, - enforce_friedel: bool = True, - chunk_size: int = 2048, -) -> BesselSHCoefficients: - """ - Compute the Phaser-style spherical-Bessel radial × spherical-harmonic - angular expansion of a scattered-point intensity field on the unit ball. - - c_nlm[n, l, m] = Σ_h Y*_{l,m}(ŝ_h) · I_h · √u · j_u(h_h) / h_h - - where `h_h = bessel_h_scale · |s_h|` and u = l + 2n + 1 (n ∈ [0, N_radial)). - Phaser (`DataMR.cc:1107`) uses `h = lmax · |s| · HIRES` where HIRES is the - high-resolution limit `d_min` in Å. With that scaling, h_max ≈ lmax (so the - highest-u Bessel functions are sampled near their first peak rather than - deep in their decaying tail). Pass `bessel_h_scale = (L - 1) * d_min` to - match Phaser exactly. - - For each l, only n with `2n + l + 1 ≤ lmax + 1` (i.e. `n ≤ (lmax − l)/2`) - are filled — others stay 0 (Phaser's truncation: - `nmax(l) = (lmax - l + 2)/2`). - - Even l only (the Patterson is centrosymmetric: Y_{l,m}(−ŝ) = (−1)^l Y_{l,m}, - so odd-l rows cancel exactly when Friedel mates are summed). - - When `zsymm > 1`, SH coefficients with `|m| % zsymm ≠ 0` are zeroed - post-expansion (Phaser `DataMR.cc:863-870, 1117`). This is the - m-symmetry filter; applied **only** on the observed side by the caller, - NOT on the calc side (which is a rotated P1 ensemble and not spacegroup- - invariant). - - Parameters - ---------- - s_vectors : (N, 3) real - Reciprocal-lattice vectors in 1/Å. - intensity : (N,) real - Per-reflection intensity (LERF1 obs intensity or σA-weighted calc). - L : int - Wigner/SH bandwidth. lmax = L - 1, restricted to even. - bessel_h_scale : float - The pre-multiplier on |s| inside the Bessel argument: `h = bessel_h_scale · |s|`. - Phaser uses `lmax · d_min` (where d_min is HIRES in Å). This puts h_max ≈ lmax - so the highest-u terms are sampled near their first peak. - zsymm : int, default 1 - m-symmetry filter (zero coefficients with |m| not divisible by zsymm). - zsymm=1 is no filter. - enforce_friedel : bool, default True - Augment the sum with (-s, I) pairs to make odd-l rows exactly zero - (otherwise we rely on the explicit zero-out at the end). - chunk_size : int - Reflections per Y_lm chunk for memory control. - - Returns - ------- - BesselSHCoefficients with `c_nlm[n, l, m+L-1]` populated. - """ - assert s_vectors.dim() == 2 and s_vectors.shape[-1] == 3 - assert intensity.dim() == 1 and intensity.shape[0] == s_vectors.shape[0] - - real_dtype = s_vectors.dtype - if real_dtype == torch.float64: - complex_dtype = torch.complex128 - elif real_dtype == torch.float32: - complex_dtype = torch.complex64 - else: - raise TypeError(f"Unsupported real dtype: {real_dtype}") - device = s_vectors.device - - if enforce_friedel: - s_vectors = torch.cat([s_vectors, -s_vectors], dim=0) - intensity = torch.cat([intensity, intensity], dim=0) - - s_mag = s_vectors.norm(dim=-1).clamp(min=1e-30) - s_hat = s_vectors / s_mag.unsqueeze(-1) - cos_theta = s_hat[..., 2].clamp(min=-1.0, max=1.0) - theta = torch.acos(cos_theta) - phi = torch.atan2(s_hat[..., 1], s_hat[..., 0]) - - lmax = L - 1 - lmax_even = lmax if (lmax % 2 == 0) else (lmax - 1) - # Phaser indexing: n ∈ [1, nmax(l)] with nmax(l) = (lmax - l + 2)/2. - # Our 0-indexed n: n ∈ [0, N_radial - 1] with N_radial = nmax(l=2). - if lmax_even < 2: - raise ValueError( - f"L={L} too small; need lmax_even >= 2 (so L >= 3)." - ) - N_radial = (lmax_even - 2) // 2 + 1 - # Highest Bessel order needed: for l = 2 and n = N_radial − 1 (1-indexed - # n_max), u = l + 2n + 1 = 2 + 2(N_radial − 1) + 1 = lmax_even + 1. - u_max = lmax_even + 1 - - # Bessel j_u(bessel_h_scale · |s|) for u ∈ [0, u_max], per reflection. - x = bessel_h_scale * s_mag # (M,) - j_all = spherical_bessel_table(x, u_max) # (M, u_max+1) - - # bessel_factor[h, l, n] = √(2u+1) · j_u(x_h) / x_h, with u = l + 2n + 1. - # Matches Phaser DataMR.cc:993 + 1109: - # sqrt_table[i] = sqrt(2*i+1); besselx[u] = sqrt_table[u] · sphbessel(u, h) / h - # (Crowther/Navaza unit-ball radial-basis weight; the previous mimic used - # √u, an unweighted variant that biased the radial power spectrum.) - M = s_vectors.shape[0] - bessel = torch.zeros((M, L, N_radial), dtype=real_dtype, device=device) - safe_x = x.clamp(min=1e-30) - for l in range(2, lmax_even + 1, 2): - n_l = (lmax_even - l) // 2 + 1 - for n in range(n_l): - u = l + 2 * n + 1 - bessel[:, l, n] = math.sqrt(float(2 * u + 1)) * j_all[:, u] / safe_x - - # Accumulate the SH expansion in chunks. - # c_nlm[n, l, m] = Σ_h conj(Y_lm(ŝ_h)) · I_h · bessel[h, l, n] - c_nlm = torch.zeros( - (N_radial, L, 2 * L - 1), dtype=complex_dtype, device=device, - ) - intensity_c = intensity.to(complex_dtype) - - for start in range(0, M, chunk_size): - stop = min(start + chunk_size, M) - Y = evaluate_ylm(theta[start:stop], phi[start:stop], L) # (n, L, 2L-1) - # weight × conj(Y) per reflection. - Y_w = torch.conj(Y) * intensity_c[start:stop].view(-1, 1, 1) # (n, L, 2L-1) - b_slice = bessel[start:stop].to(complex_dtype) # (n, L, N_radial) - # einsum: sum over h, contract bessel(h,l,n) × Y_w(h,l,m). - contrib = torch.einsum("hln,hlm->nlm", b_slice, Y_w) - c_nlm += contrib - - # Zero odd-l rows (and l = 0, which Phaser also doesn't use). - l_vals = torch.arange(L, device=device) - odd_or_zero_mask = (l_vals % 2 == 1) | (l_vals == 0) - c_nlm[:, odd_or_zero_mask, :] = 0.0 - - # m-symmetry filter (applied to observed side only; caller passes - # zsymm = 1 for calc side). - if zsymm > 1: - m_vals = torch.arange(-(L - 1), L, device=device) # (2L-1,) - m_invalid = (m_vals.abs() % zsymm) != 0 - c_nlm[:, :, m_invalid] = 0.0 - - return BesselSHCoefficients( - c_nlm=c_nlm, L=L, N_radial=N_radial, - bessel_h_scale=float(bessel_h_scale), zsymm=int(zsymm), - ) - - -# ============================================================================= -# Phaser-style data preparation: Wilson normalisation + σA Eterm -# ============================================================================= - - -def _wilson_normalise( - F: torch.Tensor, - s_mag: torch.Tensor, - n_shells: int = 20, -) -> Tuple[torch.Tensor, torch.Tensor]: - """ - Per-shell Wilson normalisation of amplitudes: - - E_h = F_h / sqrt(_p) where p = shell containing h. - - Phaser's `Feff[r] / SIGMAN.sqrt_epsnSN[r]` (`DataMR.cc:925`) does - roughly this (modulo French-Wilson and explicit ε-factor handling - which we skip; the input F is already anisotropy-corrected by the - caller in practice). - - Returns (E_h, sqrt_mean_F2_per_h). - """ - edges, _ = equal_count_shell_edges(s_mag, n_shells) - shell_idx = assign_shells(s_mag, edges) - valid = shell_idx >= 0 - F_dtype = F.dtype - F2 = F * F - count = torch.zeros(n_shells, dtype=torch.int64, device=F.device) - sumF2 = torch.zeros(n_shells, dtype=F_dtype, device=F.device) - F2_v = F2[valid] - idx_v = shell_idx[valid] - count.index_add_(0, idx_v, torch.ones_like(idx_v)) - sumF2.index_add_(0, idx_v, F2_v) - mean_F2 = sumF2 / count.clamp(min=1).to(F_dtype) - mean_F2 = mean_F2.clamp(min=1e-12) - sqrt_mean = mean_F2.sqrt() - per_h = torch.ones_like(F) - per_h[valid] = sqrt_mean[idx_v] - E = F / per_h - return E, per_h - - -# ----------------------------------------------------------------------------- -# French-Wilson posterior expected values (math_FrenchWilson.cc) -# ----------------------------------------------------------------------------- - - -def _expectE_FW_acen(eosq, sigesq): - """ - Acentric posterior expected E from normalised observed intensity (eosq) - and its standard deviation (sigesq). Translates verbatim from Phaser's - `lib/math_FrenchWilson.cc:expectEFWacen` (lines 8-44). Vectorised NumPy. - - `eosq = Iobs / `, `sigesq = σIobs / `. - """ - import numpy as np - from scipy.special import erfc, pbdv - CROSS1, CROSS2 = -12.5, 18.0 - SQRT2 = np.sqrt(2.0) - x = (eosq - sigesq ** 2) / sigesq - xsqr = x * x - ee = np.empty_like(eosq) - # Large negative argument: asymptotic - m_neg = x < CROSS1 - if m_neg.any(): - xs = xsqr[m_neg] - num = (-916620705. + xs * - (91891800. + xs * - (-11531520. + xs * - (1935360. + xs * - (-491520. + xs * 262144.))))) - den = (-495452160. + xs * - (55050240. + xs * - (-7864320. + xs * - (1572864. + xs * - (-524288. + xs * 524288.))))) - ee[m_neg] = np.sqrt(-np.pi * sigesq[m_neg] / x[m_neg]) * num / den - # Large positive argument: asymptotic - m_pos = x > CROSS2 - if m_pos.any(): - xs = xsqr[m_pos] - num = (-45045. + 32. * xs * - (-315. + 8. * xs * - (-15. - 16. * xs + 128. * xs * xs))) - ee[m_pos] = (np.sqrt(sigesq[m_pos]) * num / - (32768. * x[m_pos] ** 7.5)) - # Moderate arguments - m_mid = ~(m_neg | m_pos) - if m_mid.any(): - xm = x[m_mid] - pcd, _ = pbdv(-1.5, -xm) - ee[m_mid] = (np.sqrt(sigesq[m_mid] / 2.0) * np.exp(-xm * xm / 4.0) * - pcd / erfc(-xm / SQRT2)) - return ee - - -def _expectEsq_FW_acen(eosq, sigesq): - """Acentric posterior . From `expectEsqFWacen` (lines 46-78).""" - import numpy as np - from scipy.special import erfc - CROSS1, CROSS2 = -8.9, 5.7 - SQRT2_BY_PI = np.sqrt(2.0 / np.pi) - SQRT2 = np.sqrt(2.0) - eesq_base = eosq - sigesq ** 2 # baseline value - x = eesq_base / (SQRT2 * sigesq) - xsqr = x * x - eesq = eesq_base.copy() - m_neg = x < CROSS1 - if m_neg.any(): - xs = xsqr[m_neg] - num = (-135135. + xs * (20790. + xs * (-3780. + xs * - (840. + xs * (-240. + xs * (96. - xs * 64.)))))) - den = (-135135. + xs * (20790. + xs * (-3780. + xs * - (840. + xs * (-240. + xs * (96. + xs * - (-64. + xs * 128.))))))) - eesq[m_neg] = eesq_base[m_neg] * num / den - m_mid = (x >= CROSS1) & (x <= CROSS2) - if m_mid.any(): - xm = x[m_mid] - eesq[m_mid] = (eesq_base[m_mid] + - SQRT2_BY_PI * sigesq[m_mid] / - (np.exp(xm * xm) * erfc(-xm))) - # x > CROSS2: eesq stays at eesq_base (default per source) - return eesq - - -def _expectE_FW_cen(eosq, sigesq): - """Centric posterior . From `expectEFWcen` (lines 80-113).""" - import numpy as np - from scipy.special import pbdv - CROSS1, CROSS2 = -17.5, 17.5 - SQRTPI = np.sqrt(np.pi) - x = sigesq / 2.0 - eosq / sigesq - xsqr = x * x - pcdratio = np.empty_like(x) - m_neg = x < CROSS1 - if m_neg.any(): - xn, xs = x[m_neg], xsqr[m_neg] - pcdratio[m_neg] = ((1024. * SQRTPI * (-xn) ** 6.5) / - (3465. + xs * - (840. + xs * - (384. + xs * 1024.)))) - m_pos = x > CROSS2 - if m_pos.any(): - xp, xs = x[m_pos], xsqr[m_pos] - num = (3440640. + xs * - (-491520. + xs * - (98304. + xs * - (-32768. + xs * 32768.)))) - den = (675675. + xs * - (-110880. + xs * - (26880. + xs * - (-12288. + xs * 32768.)))) - pcdratio[m_pos] = num / (den * np.sqrt(xp)) - m_mid = ~(m_neg | m_pos) - if m_mid.any(): - xm = x[m_mid] - d_neg1, _ = pbdv(-1.0, xm) - d_neghalf, _ = pbdv(-0.5, xm) - pcdratio[m_mid] = d_neg1 / d_neghalf - return np.sqrt(sigesq / np.pi) * pcdratio - - -def _expectEsq_FW_cen(eosq, sigesq): - """Centric posterior . From `expectEsqFWcen` (lines 115-152).""" - import numpy as np - from scipy.special import pbdv - CROSS1, CROSS2 = -17.5, 17.5 - x = sigesq / 2.0 - eosq / sigesq - xsqr = x * x - pcdratio = np.empty_like(x) - m_neg = x < CROSS1 - if m_neg.any(): - xn, xs = x[m_neg], xsqr[m_neg] - num = (45045. + xs * - (10080. + xs * - (3840. + xs * - (4096. - xs * 32768.)))) - den = xn * (55440. + xs * - (13440. + xs * - (6144. + xs * 16384.))) - pcdratio[m_neg] = num / den - m_pos = x > CROSS2 - if m_pos.any(): - xp, xs = x[m_pos], xsqr[m_pos] - num = (11486475. + xs * - (-1441440. + xs * - (241920. + xs * - (-61440. + xs * 32768.)))) - den = xp * (675675. + xs * - (-110880. + xs * - (26880. + xs * - (-12288. + xs * 32768.)))) - pcdratio[m_pos] = num / den - m_mid = ~(m_neg | m_pos) - if m_mid.any(): - xm = x[m_mid] - d_neg15, _ = pbdv(-1.5, xm) - d_neghalf, _ = pbdv(-0.5, xm) - pcdratio[m_mid] = d_neg15 / d_neghalf - return sigesq * pcdratio / 2.0 - - -def _french_wilson_posterior(eosq, sigesq, centric_mask): - """Wrap centric/acentric branches. - - Phaser `expectEFW` / `expectEsqFW` (lines 154-178): if sigesq <= 0 the - measurement is treated as exact and (eEFW, eEsqFW) = (sqrt(eosq), eosq). - """ - import numpy as np - eEFW = np.empty_like(eosq) - eEsqFW = np.empty_like(eosq) - zero_sig = sigesq <= 0.0 - if zero_sig.any(): - eEFW[zero_sig] = np.sqrt(np.maximum(eosq[zero_sig], 0.0)) - eEsqFW[zero_sig] = np.maximum(eosq[zero_sig], 0.0) - valid = ~zero_sig - if valid.any(): - cen = centric_mask & valid - acen = (~centric_mask) & valid - if cen.any(): - eEFW[cen] = _expectE_FW_cen(eosq[cen], sigesq[cen]) - eEsqFW[cen] = _expectEsq_FW_cen(eosq[cen], sigesq[cen]) - if acen.any(): - eEFW[acen] = _expectE_FW_acen(eosq[acen], sigesq[acen]) - eEsqFW[acen] = _expectEsq_FW_acen(eosq[acen], sigesq[acen]) - return eEFW, eEsqFW - - -# ----------------------------------------------------------------------------- -# DFAC via Halley iteration (math_RiceLLG.cc:getDfactor) -# ----------------------------------------------------------------------------- - - -def _i0e_full(x): - """Phaser's `eBesselI0(x) = I0(x)·exp(-|x|)`. Symmetric in x. - - See `math_eBesselI0.cc` — Phaser's Rice-moment formulas use the - exp-scaled Bessel (the un-scaled `I0` cancels analytically with the - Gaussian envelope of the Rice distribution). - """ - import numpy as np - from scipy.special import i0e - return i0e(np.abs(x)) - - -def _i1e_full(x): - """Phaser's `eBesselI1(x) = I1(x)·exp(-|x|)`. Antisymmetric in x.""" - import numpy as np - from scipy.special import i1e - return np.sign(x) * i1e(np.abs(x)) - - -def _effSigaRoot_acen(ee, eesq, sa): - """`effSigaRootAcen` (math_RiceLLG.cc:12-34).""" - import numpy as np - sigbsqr = 1.0 - sa * sa - x = 0.5 * (eesq - sigbsqr) / sigbsqr - return (np.sqrt(np.pi * sigbsqr) / (2.0 * sigbsqr) * - (eesq * _i0e_full(x) + (eesq - sigbsqr) * _i1e_full(x)) - ee) - - -def _deffSigaRoot_acen(eesq, sa): - """`deffSigaRootAcen_by_dsa` (lines 36-52).""" - import numpy as np - sigbsqr = 1.0 - sa * sa - x = 0.5 * (eesq - sigbsqr) / sigbsqr - return np.sqrt(np.pi / sigbsqr) * (sa / 2.0) * _i1e_full(x) - - -def _d2effSigaRoot_acen(eesq, sa): - """`d2effSigaRootAcen_by_dsa2` (lines 54-81).""" - import numpy as np - sigasqr = sa * sa - sigapow4 = sigasqr * sigasqr - sigbsqr = 1.0 - sigasqr - xnum = eesq - sigbsqr - x = 0.5 * xnum / sigbsqr - out = np.empty_like(eesq) - big = xnum > 1e-10 - if big.any(): - I0 = _i0e_full(x[big]) - I1 = _i1e_full(x[big]) - out[big] = (np.sqrt(np.pi / sigbsqr[big]) / (2.0 * sigbsqr[big] ** 2) * - (eesq[big] * sigasqr[big] * I0 + - (eesq[big] - 1.0 - (-2.0 + eesq[big] * (2.0 + eesq[big])) * sigasqr[big] + - (eesq[big] - 1.0) * sigapow4[big]) * I1 / xnum[big])) - small = ~big - if small.any(): - samin = np.sqrt(np.maximum(1.0 - eesq[small], 0.0)) - out[small] = (np.sqrt(np.pi) * samin * - ((3.0 + samin * samin) * sa[small] - - 2.0 * (samin + samin ** 3)) / - (4.0 * eesq[small] ** 2.5)) - return out - - -def _effSigaRoot_cen(ee, eesq, sa): - """`effSigaRootCen` (lines 83-105).""" - import numpy as np - from scipy.special import erf - sigbsqr = 1.0 - sa * sa - x = 0.5 * (eesq - sigbsqr) / sigbsqr - x_safe = np.maximum(x, 0.0) # erf(sqrt(x)) ill-defined for x<0 - return (np.exp(-x) * np.sqrt(2.0 * sigbsqr / np.pi) + - np.sqrt(np.maximum(eesq - sigbsqr, 0.0)) * erf(np.sqrt(x_safe)) - ee) - - -def _deffSigaRoot_cen(eesq, sa): - """`deffSigaRootCen_by_dsa` (lines 107-130).""" - import numpy as np - from scipy.special import erf - sigbsqr = 1.0 - sa * sa - xnum = eesq - sigbsqr - x = 0.5 * xnum / sigbsqr - out = np.empty_like(eesq) - big = np.abs(xnum) > 1e-10 - if big.any(): - x_safe = np.maximum(x[big], 0.0) - out[big] = (sa[big] * erf(np.sqrt(x_safe)) / - np.sqrt(np.maximum(xnum[big], 1e-30)) - - np.exp(-x[big]) * np.sqrt(2.0 * sigbsqr[big] / np.pi) * - sa[big] / sigbsqr[big]) - small = ~big - if small.any(): - out[small] = (xnum[small] * np.sqrt(2.0 / np.pi) * sa[small] / - (3.0 * sigbsqr[small] ** 1.5)) - return out - - -def _d2effSigaRoot_cen(eesq, sa): - """`d2effSigaRootCen_by_dsa2` (lines 132-159).""" - import numpy as np - from scipy.special import erf - sigasqr = sa * sa - sigapow4 = sigasqr * sigasqr - sigbsqr = 1.0 - sigasqr - xnum = eesq - sigbsqr - x = 0.5 * xnum / sigbsqr - sigbsqrtpi = np.sqrt(np.pi * sigbsqr) - out = np.empty_like(eesq) - big = np.abs(xnum) > 1e-10 - if big.any(): - x_safe = np.maximum(x[big], 0.0) - d2num = ((eesq[big] - 1.0) * sigbsqr[big] ** 2 * sigbsqrtpi[big] * - erf(np.sqrt(x_safe))) - exp_part = np.where( - x[big] < 20.0, - np.sqrt(np.maximum(2.0 * xnum[big], 0.0)) * np.exp(-x[big]) * - (1.0 - eesq[big] + sigasqr[big] * - (eesq[big] + eesq[big] ** 2 - 2.0) + sigapow4[big]), - np.zeros_like(x[big]), - ) - d2num = d2num + exp_part - out[big] = d2num / (sigbsqr[big] ** 2 * - np.maximum(xnum[big], 1e-30) ** 1.5 * - sigbsqrtpi[big]) - small = ~big - if small.any(): - out[small] = (np.sqrt(2.0 / np.pi) * sa[small] / - (1.5 * sigbsqr[small] ** 1.5)) - return out - - -def get_dfactor_vectorised(ee_np, eesq_np, centric_np): - """ - Vectorised version of Phaser's `math_RiceLLG.cc:getDfactor` (lines 191-250). - Halley's method with bisection fallback, run over all reflections in - parallel. Each reflection has its own bracket [dflo, dfhi]. - - Returns a (N,) numpy float64 array of DFAC values in (0, 1). - """ - import numpy as np - - EPS1 = 1e-7 - EPS2 = 1e-10 - MAXDFAC = 1.0 - EPS1 - - ee = np.asarray(ee_np, dtype=np.float64) - eesq = np.asarray(eesq_np, dtype=np.float64) - cen = np.asarray(centric_np, dtype=bool) - N = ee.shape[0] - - # Case 1: no observational error (eesq - ee² ≤ 0) → DFAC = 1.0 - out = np.ones(N, dtype=np.float64) - has_err = (eesq - ee * ee) > 0.0 - - if not has_err.any(): - return out - - ee_a, eesq_a, cen_a = ee[has_err], eesq[has_err], cen[has_err] - # Bracket. dflo = max(sqrt(1 - min(eesq, 1)) + EPS, EPS). - dflo = np.maximum(np.sqrt(np.maximum(1.0 - np.minimum(eesq_a, 1.0), 0.0)) + EPS1, EPS1) - dfhi = np.full_like(dflo, MAXDFAC) - - # Early return: if dflo >= MAXDFAC, just use dflo. - early = dflo >= MAXDFAC - if early.any(): - # No iteration for those; keep dflo - pass - - dfmid = 0.5 * (dflo + dfhi) - # Compute initial fmid - fmid = np.empty_like(dfmid) - if cen_a.any(): - fmid[cen_a] = _effSigaRoot_cen(ee_a[cen_a], eesq_a[cen_a], dfmid[cen_a]) - if (~cen_a).any(): - fmid[~cen_a] = _effSigaRoot_acen(ee_a[~cen_a], eesq_a[~cen_a], dfmid[~cen_a]) - - active = ~early # reflections still being iterated - for _ in range(50): - if not active.any(): - break - # Convergence check - conv = (dfhi - dflo) <= EPS1 - conv |= np.abs(fmid) <= EPS2 - active = active & ~conv - if not active.any(): - break - - # slope and curve at dfmid (only for active reflections) - slope = np.empty_like(dfmid) - curve = np.empty_like(dfmid) - cen_act = cen_a & active - acen_act = (~cen_a) & active - if cen_act.any(): - slope[cen_act] = _deffSigaRoot_cen(eesq_a[cen_act], dfmid[cen_act]) - curve[cen_act] = _d2effSigaRoot_cen(eesq_a[cen_act], dfmid[cen_act]) - if acen_act.any(): - slope[acen_act] = _deffSigaRoot_acen(eesq_a[acen_act], dfmid[acen_act]) - curve[acen_act] = _d2effSigaRoot_acen(eesq_a[acen_act], dfmid[acen_act]) - - # Halley step (fall back to `fmid · slope` if curve <= 0, matching - # the Phaser source literally — math_RiceLLG.cc:230-231). - denom_halley = 2.0 * (slope ** 2 - fmid * curve) - use_halley = (curve > 0.0) & (np.abs(denom_halley) > 1e-30) - step = np.where( - use_halley, - 2.0 * fmid * slope / np.where(use_halley, denom_halley, 1.0), - fmid * slope, - ) - dfnew = dfmid - step - # If new value out of bracket, bisect instead - in_bracket = (dfnew > dflo) & (dfnew < dfhi) - dfmid_new = np.where(in_bracket, dfnew, 0.5 * (dflo + dfhi)) - - # Apply only on active reflections; converged ones keep dfmid. - dfmid = np.where(active, dfmid_new, dfmid) - - # Re-evaluate fmid only on active reflections. - if cen_act.any(): - fmid[cen_act] = _effSigaRoot_cen(ee_a[cen_act], eesq_a[cen_act], - dfmid[cen_act]) - if acen_act.any(): - fmid[acen_act] = _effSigaRoot_acen(ee_a[acen_act], eesq_a[acen_act], - dfmid[acen_act]) - - # Update bracket based on fmid sign. - below = (fmid < 0.0) & active - above = (fmid >= 0.0) & active - dflo = np.where(below, dfmid, dflo) - dfhi = np.where(above, dfmid, dfhi) - - out[has_err] = dfmid - return np.clip(out, EPS1, MAXDFAC) - - -def french_wilson_preprocess( - F: torch.Tensor, - sig_F: torch.Tensor, - s_mag: torch.Tensor, - centric: torch.Tensor, - *, - n_wilson_shells: int = 20, -) -> dict: - """ - Phaser-style preprocessing from raw (F, σF, centric) to (eEobs, DFAC). - - Implements the chain documented at the top of the module: - 1. equal-count Wilson shells over `s_mag` - 2. per-shell `_p` (Phaser's `SIGMAN.BINS`) - 3. per-reflection normalised intensity `eosq = F² / ` and σ - `sigesq = σI / ≈ 2·F·σF / ` - 4. French-Wilson posterior `eEFW, eEsqFW` (`math_FrenchWilson.cc`) - 5. DFAC via Halley iteration on Rice moments (`math_RiceLLG.cc`) - 6. `eEobs = sqrt(eEsqFW + (DFAC²−1)/DFAC²)`, clamped to ≤10 - (Phaser `Dfactor.cc:87-93`). - - Returns a dict with torch tensors back on the input device: - eEobs: (N,) effective normalised amplitude - DFAC : (N,) per-reflection D-factor ∈ [1e-7, 1−1e-7] - sqrt_mean_F2: (N,) per-reflection √_p (caller can multiply - back to recover absolute-scale Feff if needed) - """ - import numpy as np - - device = F.device - F_np = F.detach().to("cpu").to(torch.float64).numpy() - sigF_np = sig_F.detach().to("cpu").to(torch.float64).numpy() - s_np = s_mag.detach().to("cpu").to(torch.float64).numpy() - cen_np = centric.detach().to("cpu").bool().numpy() - - # 1+2. Per-shell by equal-count binning over |s|. - sorted_idx = np.argsort(s_np) - edges_idx = np.linspace(0, len(s_np) - 1, n_wilson_shells + 1).round().astype(np.int64) - s_edges = s_np[sorted_idx][edges_idx] - # Nudge endpoints - s_edges[0] -= 1e-6 - s_edges[-1] += 1e-6 - shell_idx = np.clip( - np.searchsorted(s_edges, s_np, side="right") - 1, 0, n_wilson_shells - 1, - ) - F2 = F_np * F_np - mean_F2 = np.zeros(n_wilson_shells, dtype=np.float64) - counts = np.zeros(n_wilson_shells, dtype=np.int64) - np.add.at(mean_F2, shell_idx, F2) - np.add.at(counts, shell_idx, 1) - mean_F2 = mean_F2 / np.maximum(counts, 1) - mean_F2 = np.maximum(mean_F2, 1e-12) - mean_I_per_h = mean_F2[shell_idx] # _p mapped to each h - sqrt_mean_F2 = np.sqrt(mean_I_per_h) - - # 3. Normalised intensity + its sigma. - eosq = F2 / mean_I_per_h - sigesq = 2.0 * F_np * sigF_np / mean_I_per_h - # Guard: sigesq must be positive for FW. If σF=0 we let zero_sig path - # in `_french_wilson_posterior` handle it. - sigesq = np.maximum(sigesq, 0.0) - - # 4. French-Wilson posterior. - eEFW, eEsqFW = _french_wilson_posterior(eosq, sigesq, cen_np) - # Guard non-physical posteriors (numerical accidents): make sure - # eEsqFW ≥ eEFW² (the second moment must dominate the first squared). - bad = eEsqFW < eEFW * eEFW - if bad.any(): - eEsqFW[bad] = eEFW[bad] ** 2 + 1e-12 - - # 5. DFAC via Halley iteration. - DFAC = get_dfactor_vectorised(eEFW, eEsqFW, cen_np) - - # 6. eEobs = sqrt(eEsqFW + (DFAC²-1)/DFAC²), clamp to ≤ 10; recompute - # DFAC if clamped (mirrors `Dfactor.cc:87-93`). - dfsqr = DFAC * DFAC - eEobs_sqr = eEsqFW + (dfsqr - 1.0) / np.maximum(dfsqr, 1e-30) - eEobs_sqr = np.maximum(eEobs_sqr, 0.0) - eEobs = np.sqrt(eEobs_sqr) - # Clamp at 10 and recompute DFAC where clamped (rare). - clamp_mask = (eEobs > 10.0) & (eEsqFW > 1.0) - if clamp_mask.any(): - # eEobs = 10 → solve for DFAC: 100 = eEsqFW + (dfsqr-1)/dfsqr - # → (100 - eEsqFW) = (dfsqr - 1)/dfsqr - # → dfsqr (100 - eEsqFW) = dfsqr - 1 - # → dfsqr (100 - eEsqFW - 1) = -1 - # → dfsqr (eEsqFW - 99) = 1 - eEobs[clamp_mask] = 10.0 - DFAC[clamp_mask] = 1.0 / np.sqrt(np.maximum(eEsqFW[clamp_mask] - 99.0, 1e-30)) - DFAC[clamp_mask] = np.clip(DFAC[clamp_mask], 1e-7, 1.0 - 1e-7) - - return { - "eEobs": torch.from_numpy(eEobs).to(device=device, dtype=F.dtype), - "DFAC": torch.from_numpy(DFAC).to(device=device, dtype=F.dtype), - "sqrt_mean_F2": torch.from_numpy(sqrt_mean_F2).to(device=device, dtype=F.dtype), - } - - -def eterm_sigma_a(s_mag: torch.Tensor, delta_vrms_A: float) -> torch.Tensor: - """ - Phaser's σA Eterm, literal port of `Ensemble.cc:42`: - - Eterm(s) = exp(-(2π²/3) · s² · ΔVRMS_var) - - where `ΔVRMS_var` is the *coordinate variance* in Ų. We accept the RMS - coordinate error `delta_vrms_A` (Å) per the standard σA convention and - square it internally: `ΔVRMS_var = delta_vrms_A²`. - - Per-reflection (`s_mag` may be any shape). Returns the same shape. - - NB: an earlier version of this mimic used `exp(-2π² · s² · ΔVRMS²)` — - that is *3×* too aggressive in the exponent and over-attenuated calc at - high resolution. See `phaser_frf_known_bugs.md` for the audit. - """ - s2 = s_mag * s_mag - return torch.exp(-(2.0 / 3.0) * (math.pi ** 2) * s2 * (delta_vrms_A ** 2)) - - -# ============================================================================= -# Top-level FRF entry point — mirrors `ball_rotation_search` API -# ============================================================================= - - -def phaser_rotation_search( - s_obs: torch.Tensor, - F_obs: torch.Tensor, - centric_obs: torch.Tensor, - s_calc: torch.Tensor, - F_calc: torch.Tensor, - sym_mats: torch.Tensor, - *, - L: int = 24, - d_min: Optional[float] = None, - d_max: Optional[float] = None, - delta_vrms_A: float = 1.0, - n_wilson_shells: int = 20, - n_peaks: int = 500, - refine_subvoxel: bool = True, - n_refine: int = 50, - sigma_threshold: float = -5.0, - bessel_h_scale: Optional[float] = None, - use_lerf1_intensity: bool = True, - use_m_symmetry_filter: bool = True, - sig_F_obs: Optional[torch.Tensor] = None, - use_french_wilson: bool = False, - use_shell_variance_weights: bool = False, - n_var_shells: int = 20, - grid_sampling_deg: float = 3.0, -) -> Tuple[AdaptiveRotationFunction, List[RotationPeak]]: - """ - Phaser-faithful Fast Rotation Function. - - Parameters - ---------- - s_obs : (N_o, 3) real - Observed-side reciprocal-lattice vectors in 1/Å. - F_obs : (N_o,) real - Observed amplitudes (already anisotropy-corrected by caller, if - applicable). - centric_obs : (N_o,) bool - Centric flag for each observed reflection. - s_calc, F_calc : (N_c, 3), (N_c,) real - Model reciprocal vectors and amplitudes. - sym_mats : (n_ops, 3, 3) real - Spacegroup rotation matrices — used to detect the high-order axis - for the m-symmetry filter (Phaser `highOrderAxis()`). - L : int, default 24 - Wigner/SH bandwidth. lmax = L − 1 (rounded down to even). - d_min, d_max : float, optional - Resolution window. Reflections outside `[1/d_max, 1/d_min]` (in |s|) - are discarded on both sides before normalisation. - delta_vrms_A : float, default 1.0 - ΔVRMS in Å for the Luzzati σA Eterm applied to the calc side. - n_wilson_shells : int, default 20 - Number of equal-count shells used in per-shell Wilson normalisation. - n_peaks : int, default 500 - Maximum number of rotation peaks returned. - refine_subvoxel : bool, default True - Apply quadratic sub-voxel refinement to the top `n_refine` peaks. - n_refine : int, default 50 - Number of peaks to refine sub-voxel. - sigma_threshold : float, default -5.0 - Minimum Z-score for a peak to be returned. Negative ≈ "keep everything". - bessel_h_scale : float, optional - Pre-multiplier on |s| in the Bessel argument: `h = bessel_h_scale · |s|`. - Default: `(L - 1) · d_min` (Phaser's `lmax · HIRES`). Picking this - scale puts the highest-u Bessel functions near their first peak at - the maximum |s| in the data, so the radial basis covers the full - resolution range without wasted bandwidth. - use_lerf1_intensity : bool, default True - If True, observed intensity is `cweight · (E_obs² − 1)`; if False, - plain `E_obs² − 1` (cweight = 1 everywhere). - use_m_symmetry_filter : bool, default True - If True, detect ZSYMM from `sym_mats` and apply the m-symmetry - filter to the observed-side SH coefficients. - sig_F_obs : (N_o,) real, optional - Standard deviation of `F_obs`. Required if `use_french_wilson=True`; - ignored otherwise. - use_french_wilson : bool, default False - If True, the observed-side intensity is built via the full Phaser - preprocessing chain: per-shell Wilson normalisation → French-Wilson - posterior → per-reflection Luzzati DFAC → effective Eobs. Observed - intensity becomes `cweight · (eEobs² − 1) · DFAC²` instead of the - simpler `cweight · (E² − 1)`. Requires `sig_F_obs`. - use_shell_variance_weights : bool, default False - If True, compute per-shell empirical variance of `intensity_obs` - and weight each reflection by `1/√Var_p` before the SH expansion. - Analog of torchref's `auto_variance_weights=True`. Downweights - shells whose obs Patterson coefficient is dominated by intermolecular - noise — particularly important for cubic/tetragonal groups where the - spacegroup-invariant SH subspace has few coefficients and per-shell - SNR matters disproportionately. - n_var_shells : int, default 20 - Number of shells for the variance estimate. Same scheme as - `n_wilson_shells`. - grid_sampling_deg : float, default 3.0 - Target Euler-angle resolution in degrees. Drives the per-β α/γ sample - counts via Phaser's `pmax(β) = 720/Δ · cos(β/2)`, - `qmax(β) = 360/Δ · sin(β/2)` (FastRot.cc:92-96). Total adaptive - sample count ≈ `(720 · 360) / grid_sampling_deg²`. - - Returns - ------- - adaptive_rf : AdaptiveRotationFunction - Ragged Euler-grid evaluation of `C(α, β, γ)` with per-β - variable-shape `(qmax_k, pmax_k)` slices, alpha/gamma grids, and the - β midpoint quadrature points. - peaks : list of RotationPeak (sorted by descending score). - """ - device = s_obs.device - real_dtype = s_obs.dtype - - # 1. Resolution mask on both sides. - def _resmask(s_vec, F, centric=None): - smag = s_vec.norm(dim=-1) - lo = 1.0 / d_max if d_max is not None else 0.0 - hi = 1.0 / d_min if d_min is not None else float("inf") - keep = (smag >= lo) & (smag <= hi) - s_vec = s_vec[keep] - F = F[keep] - if centric is not None: - centric = centric[keep] - return s_vec, F, centric, smag[keep] - return s_vec, F, smag[keep] - - # Apply resolution mask to obs side; also mask sig_F if provided. - if use_french_wilson: - if sig_F_obs is None: - raise ValueError("use_french_wilson=True requires sig_F_obs.") - smag_obs_pre = s_obs.norm(dim=-1) - lo = 1.0 / d_max if d_max is not None else 0.0 - hi = 1.0 / d_min if d_min is not None else float("inf") - keep_obs = (smag_obs_pre >= lo) & (smag_obs_pre <= hi) - s_obs = s_obs[keep_obs] - F_obs = F_obs[keep_obs] - sig_F_obs = sig_F_obs[keep_obs] - centric_obs = centric_obs[keep_obs] - smag_obs = smag_obs_pre[keep_obs] - s_calc, F_calc, smag_calc = _resmask(s_calc, F_calc) - else: - s_obs, F_obs, centric_obs, smag_obs = _resmask(s_obs, F_obs, centric_obs) - s_calc, F_calc, smag_calc = _resmask(s_calc, F_calc) - - if s_obs.shape[0] < n_wilson_shells * 5: - raise ValueError( - f"Too few obs reflections ({s_obs.shape[0]}) for " - f"{n_wilson_shells} Wilson shells in [{d_min}, {d_max}] Å." - ) - - # 2. Bessel argument scale. Default = lmax · d_min (Phaser DataMR.cc:1107). - if bessel_h_scale is None: - if d_min is None: - raise ValueError("bessel_h_scale must be set when d_min is None") - lmax = L - 1 - lmax_even = lmax if lmax % 2 == 0 else lmax - 1 - bessel_h_scale = float(lmax_even) * float(d_min) - - # 3. Wilson-normalise both sides (+ FW + DFAC on obs if requested). - if use_french_wilson: - fw = french_wilson_preprocess( - F_obs, sig_F_obs, smag_obs, centric_obs, - n_wilson_shells=n_wilson_shells, - ) - eEobs = fw["eEobs"] - DFAC = fw["DFAC"] - else: - E_obs, _ = _wilson_normalise(F_obs, smag_obs, n_wilson_shells) - eEobs = E_obs - DFAC = torch.ones_like(E_obs) - E_calc, _ = _wilson_normalise(F_calc, smag_calc, n_wilson_shells) - - # 4. Observed intensity (LERF1-style): cweight · (eEobs² − 1) · DFAC². - if use_lerf1_intensity: - cweight = torch.where( - centric_obs.bool(), - torch.ones_like(eEobs), - 2.0 * torch.ones_like(eEobs), - ) - else: - cweight = torch.ones_like(eEobs) - intensity_obs = cweight * (eEobs * eEobs - 1.0) * (DFAC * DFAC) - - # 4b. Optional per-shell empirical-variance weighting (Phaser does this - # via per-shell BINS + best(r) — torchref's `auto_variance_weights=True` - # is the cleanest analog). Downweight shells whose intensity_obs has high - # empirical variance, normalising so the mean weight is ~1. - if use_shell_variance_weights: - from ..sh import compute_patterson_shell_variance - edges_var, _ = equal_count_shell_edges(smag_obs, n_var_shells) - shell_idx_obs = assign_shells(smag_obs, edges_var) - valid_obs = shell_idx_obs >= 0 - var_p = compute_patterson_shell_variance( - intensity_obs[valid_obs].to(torch.float64), - shell_idx_obs[valid_obs], - P=n_var_shells, - ) - inv_sqrt_var = 1.0 / var_p.sqrt().clamp(min=1e-30) - # Normalise mean weight = 1 so the absolute scale of intensity_obs - # doesn't shift (only the SHAPE across shells matters for the - # cross-correlation rotation function). - inv_sqrt_var = (inv_sqrt_var * - (n_var_shells / inv_sqrt_var.sum().clamp(min=1e-30))) - weights_per_h = torch.ones_like(intensity_obs) - weights_per_h[valid_obs] = inv_sqrt_var[shell_idx_obs[valid_obs]].to( - intensity_obs.dtype, - ) - intensity_obs = intensity_obs * weights_per_h - - # 5. Model intensity with σA Eterm² weighting (per-reflection). - eterm = eterm_sigma_a(smag_calc, delta_vrms_A) - intensity_calc = (eterm ** 2) * (E_calc * E_calc - 1.0) - - # 6. Detect ZSYMM (high-order rotation axis from sym_mats). - zsymm = 1 - if use_m_symmetry_filter and sym_mats is not None: - _, zsymm = get_high_order_axis( - sym_mats.to(torch.float64).cpu(), - ) - zsymm = int(zsymm) - - # 7. Bessel-radial × SH expansion of both sides. - c_obs = bessel_sh_expand( - s_obs, intensity_obs.to(real_dtype), - L=L, bessel_h_scale=bessel_h_scale, - zsymm=zsymm, enforce_friedel=True, - ) - c_calc = bessel_sh_expand( - s_calc, intensity_calc.to(real_dtype), - L=L, bessel_h_scale=bessel_h_scale, - zsymm=1, enforce_friedel=True, # NO m-filter on calc side - ) - - # 8. Cross-correlation in Wigner basis: contract on the radial-Bessel axis. - # Convention chosen to match torchref's existing `ball_search.py` (line 182): - # xi[l, m, n] = Σ_r c_obs[r, l, n] · conj(c_calc[r, l, m]) - # i.e. obs-side SH index labels the OUTPUT "n" axis (→ γ Fourier), - # calc-side SH index labels the OUTPUT "m" axis (→ α Fourier). With this - # ordering, the peak Euler triple satisfies `s_calc = R · s_obs` (column - # vector), matching the test-driver scenario where F_calc was generated by - # applying R to the model coordinates. - xi = torch.einsum( - "rln,rlm->lmn", - c_obs.c_nlm, - torch.conj(c_calc.c_nlm), - ) - - # 9. Evaluate on the Phaser-faithful adaptive Euler grid. - arf = evaluate_rotation_function_grid_adaptive( - xi, L, grid_sampling_deg=grid_sampling_deg, n_beta=2 * L, - ) - - # 10. Peak picking + sub-voxel refinement on the ragged grid. - peaks = find_rotation_peaks_adaptive( - arf, n_peaks=n_peaks, sigma_threshold=sigma_threshold, - ) - if refine_subvoxel and peaks: - head = peaks[: min(n_refine, len(peaks))] - head = refine_peaks_subvoxel_adaptive(head, arf) - peaks = head + peaks[len(head):] - peaks.sort(key=lambda r: r.score, reverse=True) - - return arf, peaks diff --git a/torchref/alignment/frf/preprocessing.py b/torchref/alignment/frf/preprocessing.py index 53fe3f5b..c5b35617 100644 --- a/torchref/alignment/frf/preprocessing.py +++ b/torchref/alignment/frf/preprocessing.py @@ -18,12 +18,7 @@ import torch -# Re-exports from the validated legacy implementation. -from .phaser_frf import ( - _wilson_normalise as wilson_normalise, # DataMR.cc:925 - eterm_sigma_a, # Ensemble.cc:42 - french_wilson_preprocess, # math_FrenchWilson.cc + Dfactor.cc -) +from .french_wilson import french_wilson_preprocess # math_FrenchWilson.cc + Dfactor.cc from ..sh import ( get_high_order_axis, # phaser's highOrderAxis() compute_patterson_shell_variance, @@ -31,6 +26,54 @@ assign_shells, ) + +def wilson_normalise( + F: torch.Tensor, + s_mag: torch.Tensor, + n_shells: int = 20, +): + """Per-shell Wilson normalisation of amplitudes. + + Source: Phaser's ``Feff[r] / SIGMAN.sqrt_epsnSN[r]`` (``DataMR.cc:925``) + minus French-Wilson + explicit ε (``F`` is assumed anisotropy-corrected + by the caller). + + E_h = F_h / sqrt(_p) where p = shell containing h. + + Returns ``(E_h, sqrt_mean_F2_per_h)``. + """ + edges, _ = equal_count_shell_edges(s_mag, n_shells) + shell_idx = assign_shells(s_mag, edges) + valid = shell_idx >= 0 + F_dtype = F.dtype + F2 = F * F + count = torch.zeros(n_shells, dtype=torch.int64, device=F.device) + sumF2 = torch.zeros(n_shells, dtype=F_dtype, device=F.device) + F2_v = F2[valid] + idx_v = shell_idx[valid] + count.index_add_(0, idx_v, torch.ones_like(idx_v)) + sumF2.index_add_(0, idx_v, F2_v) + mean_F2 = sumF2 / count.clamp(min=1).to(F_dtype) + mean_F2 = mean_F2.clamp(min=1e-12) + sqrt_mean = mean_F2.sqrt() + per_h = torch.ones_like(F) + per_h[valid] = sqrt_mean[idx_v] + E = F / per_h + return E, per_h + + +def eterm_sigma_a(s_mag: torch.Tensor, delta_vrms_A: float) -> torch.Tensor: + """Phaser's σA Eterm, literal port of ``Ensemble.cc:42``: + + Eterm(s) = exp(-(2π²/3) · s² · ΔVRMS_var) + + where ``ΔVRMS_var`` is the *coordinate variance* in Ų. We accept the + RMS coordinate error ``delta_vrms_A`` (Å) per the standard σA + convention and square it internally: ``ΔVRMS_var = delta_vrms_A²``. + """ + s2 = s_mag * s_mag + return torch.exp(-(2.0 / 3.0) * (math.pi ** 2) * s2 * (delta_vrms_A ** 2)) + __all__ = [ "wilson_normalise", "wilson_normalise_epsilon", @@ -205,6 +248,46 @@ def build_lerf1_intensity( return cw * (eEobs * eEobs - 1.0) * (dfac * dfac) +def solid_angle_weights( + s_vec: torch.Tensor, + n_cos_theta: int = 16, + n_phi: int = 32, +) -> torch.Tensor: + """Per-reflection angular quadrature weight to de-bias the SH expansion. + + The obs SH coefficient is a discretised ``∫ Y*_lm I dΩ`` over the reciprocal + crystal lattice. The lattice points are NOT uniform on the sphere — they + cluster along the cell's symmetry directions — so the unweighted sum + over-represents those directions and amplifies the symmetry-axis (ghost) + channel. This returns a weight ``w_i = 1 / (count in i's equal-area angular + cell)`` (normalised so ``Σ w = N``), which equalises each direction's + contribution — a crude spherical-quadrature / inverse-density correction. + + Bins are equal-area on the sphere (uniform in ``cos θ`` and ``φ``). + + Parameters + ---------- + s_vec : (N, 3) reciprocal-space Cartesian vectors. + n_cos_theta, n_phi : int — angular bin counts (equal-area cells). + + Returns + ------- + w : (N,) weights, dtype = s_vec.dtype, normalised to ``Σ w = N``. + """ + s_mag = s_vec.norm(dim=-1).clamp(min=1e-30) + hat = s_vec / s_mag.unsqueeze(-1) + cos_t = hat[..., 2].clamp(-1.0, 1.0) + phi = torch.atan2(hat[..., 1], hat[..., 0]) # [-π, π] + ti = ((cos_t + 1.0) * 0.5 * n_cos_theta).floor().clamp(0, n_cos_theta - 1).to(torch.int64) + pi = ((phi + math.pi) / (2.0 * math.pi) * n_phi).floor().clamp(0, n_phi - 1).to(torch.int64) + cell = ti * n_phi + pi # (N,) + n_cells = n_cos_theta * n_phi + count = torch.bincount(cell, minlength=n_cells).clamp(min=1) + w = 1.0 / count[cell].to(torch.float64) + w = w * (float(s_vec.shape[0]) / w.sum().clamp(min=1e-30)) + return w.to(s_vec.dtype) + + def apply_shell_variance_weights( intensity: torch.Tensor, s_mag: torch.Tensor, @@ -234,16 +317,28 @@ def apply_shell_variance_weights( def detect_zsymm(sym_mats: Optional[torch.Tensor]) -> int: - """Detect the high-order rotational axis order ``ZSYMM``. - - Phaser source: ``highOrderAxis()`` in ``rotationgroup.h``. Used to - apply the m-symmetry filter to the obs-side SH coefficients: SH - coefficients with ``|m| mod ZSYMM != 0`` average to zero over the - spacegroup and are zeroed out (DataMR.cc:863-870, 1117). + """Detect ``ZSYMM`` for the **z-axis** m-symmetry filter. + + Phaser source: ``highOrderAxis()`` in ``rotationgroup.h``. The m-filter + zeroes obs SH coefficients with ``|m| mod ZSYMM != 0`` — but ``m`` is the + azimuthal order **about the z axis of the SH basis**, so the filter is only + valid when the crystal's high-order rotation axis is actually along z + (cubic / tetragonal / hexagonal: principal axis = c ∥ z). + + For spacegroups whose high-order axis is x or y — e.g. monoclinic C2 / P2₁ + with the 2-fold along b ∥ y — applying a z-axis filter is WRONG (it filters + about the wrong axis and corrupts the obs coefficients). Phaser rotates the + high-order axis to z before the expansion (DataMR.cc:962-979); we do not, so + here we conservatively return ``ZSYMM=1`` (no filter) when the axis is not z, + rather than apply a wrong filter. Verified: for all benchmark cubic/tetra/hex + cases the axis is z (filter unchanged); only monoclinic 1DAW/3E98/3VRJ change + (they already rank 0-3, and a 2-fold m-filter is a weak constraint anyway). """ if sym_mats is None: return 1 - _, zsymm = get_high_order_axis(sym_mats.to(torch.float64).cpu()) + axis, zsymm = get_high_order_axis(sym_mats.to(torch.float64).cpu()) + if axis != 2: # high-order axis not along z → don't apply a wrong filter + return 1 return int(zsymm) diff --git a/torchref/alignment/frf/rotation_utils.py b/torchref/alignment/frf/rotation_utils.py new file mode 100644 index 00000000..8f1d2537 --- /dev/null +++ b/torchref/alignment/frf/rotation_utils.py @@ -0,0 +1,89 @@ +"""Pure-geometry helpers shared by the FRF, the rescore, and tests. + +Edmonds active ZYZ convention throughout: a rotation matrix is built as +``R = R_z(α) R_y(β) R_z(γ)``. ``α, γ ∈ [0, 2π)``, ``β ∈ [0, π]``. +""" +from __future__ import annotations + +import math +from typing import Tuple + +import torch + + +def rotation_matrix_from_edmonds_euler( + alpha: float, beta: float, gamma: float, dtype=torch.float64, +) -> torch.Tensor: + """Build ``R = R_z(α) R_y(β) R_z(γ)`` (Edmonds active ZYZ). + + Equivalent to passing ``[γ, β, α]`` to + ``torchref.alignment.transform.rotation_matrix_from_euler``. + """ + ca, sa = math.cos(alpha), math.sin(alpha) + cb, sb = math.cos(beta), math.sin(beta) + cg, sg = math.cos(gamma), math.sin(gamma) + Rz_a = torch.tensor([[ca, -sa, 0.0], [sa, ca, 0.0], [0.0, 0.0, 1.0]], dtype=dtype) + Ry_b = torch.tensor([[cb, 0.0, sb], [0.0, 1.0, 0.0], [-sb, 0.0, cb]], dtype=dtype) + Rz_c = torch.tensor([[cg, -sg, 0.0], [sg, cg, 0.0], [0.0, 0.0, 1.0]], dtype=dtype) + return Rz_a @ Ry_b @ Rz_c + + +def rotation_matrix_from_edmonds_euler_batch( + alpha: torch.Tensor, beta: torch.Tensor, gamma: torch.Tensor, +) -> torch.Tensor: + """Vectorised Edmonds ZYZ Euler → ``R``. + + ``alpha``, ``beta``, ``gamma``: identically-shaped real tensors. Returns + ``(..., 3, 3)`` in the same dtype/device as the inputs. + """ + ca, sa = torch.cos(alpha), torch.sin(alpha) + cb, sb = torch.cos(beta), torch.sin(beta) + cg, sg = torch.cos(gamma), torch.sin(gamma) + zero = torch.zeros_like(alpha) + one = torch.ones_like(alpha) + Rz_a = torch.stack([ + torch.stack([ca, -sa, zero], dim=-1), + torch.stack([sa, ca, zero], dim=-1), + torch.stack([zero, zero, one], dim=-1), + ], dim=-2) + Ry_b = torch.stack([ + torch.stack([cb, zero, sb], dim=-1), + torch.stack([zero, one, zero], dim=-1), + torch.stack([-sb, zero, cb], dim=-1), + ], dim=-2) + Rz_c = torch.stack([ + torch.stack([cg, -sg, zero], dim=-1), + torch.stack([sg, cg, zero], dim=-1), + torch.stack([zero, zero, one], dim=-1), + ], dim=-2) + return Rz_a @ Ry_b @ Rz_c + + +def edmonds_euler_from_rotation_matrix(R: torch.Tensor) -> Tuple[float, float, float]: + """Recover ``(α, β, γ)`` such that ``R = R_z(α) R_y(β) R_z(γ)``. + + Returns angles in radians; ``α, γ ∈ [0, 2π)``, ``β ∈ [0, π]``. Singular + when ``β = 0`` or ``π`` (only ``α+γ`` is determined); in those cases + ``γ=0`` is returned. + """ + R = R.to(torch.float64) + cos_beta = R[2, 2].clamp(-1.0, 1.0).item() + beta = math.acos(cos_beta) + sin_beta = math.sin(beta) + if abs(sin_beta) < 1e-9: + alpha = math.atan2(R[1, 0].item(), R[0, 0].item()) + gamma = 0.0 + else: + alpha = math.atan2(R[1, 2].item(), R[0, 2].item()) + gamma = math.atan2(R[2, 1].item(), -R[2, 0].item()) + alpha = alpha % (2.0 * math.pi) + gamma = gamma % (2.0 * math.pi) + return alpha, beta, gamma + + +def rotation_angular_distance_deg(R1: torch.Tensor, R2: torch.Tensor) -> float: + """Geodesic distance on SO(3) in degrees: ``arccos((tr(R1 R2^T) − 1)/2)``.""" + R = R1.to(torch.float64) @ R2.to(torch.float64).T + tr = (R[0, 0] + R[1, 1] + R[2, 2]).clamp(-1.0, 3.0).item() + cos_a = max(-1.0, min(1.0, (tr - 1.0) / 2.0)) + return math.degrees(math.acos(cos_a)) diff --git a/torchref/alignment/frf/sitelist_ang.py b/torchref/alignment/frf/sitelist_ang.py index 0cd0913d..50567d94 100644 --- a/torchref/alignment/frf/sitelist_ang.py +++ b/torchref/alignment/frf/sitelist_ang.py @@ -145,6 +145,13 @@ def _build_beta_grid(grid_sampling_deg: float) -> Tuple[torch.Tensor, int]: return betas_rad, bmax +# Module-level memo for the data-independent sample list, keyed on +# (grid_sampling_deg, device-str, dtype). The list depends only on geometry, so +# repeat FRF calls at the same grid reuse it; a single cold call still pays the +# (now vectorised, CPU-built) construction once. +_SAMPLE_LIST_CACHE: dict = {} + + def build_adaptive_sample_list( grid_sampling_deg: float, dtype: torch.dtype = torch.float64, @@ -160,17 +167,33 @@ def build_adaptive_sample_list( beta_starts: (bmax + 1,) int64 slice [beta_starts[b]:beta_starts[b+1]] is the samples at β = b · Δ beta_grid : (bmax,) the β values in radians + + The construction is purely geometric (independent of the data / ξ), so it is + memoised on ``(grid_sampling_deg, device, dtype)``. It is also built on the + CPU — the per-β tensors are tiny and the work is launch-latency-bound, so a + single host-side build + one device transfer is far cheaper than thousands of + small CUDA kernels. The per-key dedup is vectorised (no ``.tolist()`` / + Python scan), so there is no host sync inside the loop. """ - betas_rad, bmax = _build_beta_grid(grid_sampling_deg) - betas_rad = betas_rad.to(device=device, dtype=dtype) + device = torch.device(device) if not isinstance(device, torch.device) else device + cache_key = (float(grid_sampling_deg), str(device), dtype) + cached = _SAMPLE_LIST_CACHE.get(cache_key) + if cached is not None: + return cached + + bmax = int(math.ceil(180.0 / grid_sampling_deg)) + if bmax < 1: + raise ValueError(f"grid_sampling_deg={grid_sampling_deg} too coarse") + cpu = torch.device("cpu") alphas_list: List[torch.Tensor] = [] gammas_list: List[torch.Tensor] = [] betas_list: List[torch.Tensor] = [] beta_starts: List[int] = [0] + deg2rad = math.pi / 180.0 for b in range(bmax): - beta_rad = float(betas_rad[b].item()) + beta_rad = b * grid_sampling_deg * deg2rad # plain math: no sync cosb = math.cos(beta_rad / 2.0) sinb = math.sin(beta_rad / 2.0) pmax = max(1, int(720.0 / grid_sampling_deg * cosb)) @@ -178,7 +201,7 @@ def build_adaptive_sample_list( if b == 0: # β=0: only α = γ = p/pmax for p < pmax/2 (FastRot.cc:189-207). - p_idx = torch.arange(pmax, device=device) + p_idx = torch.arange(pmax, device=cpu) p_ratio = p_idx.to(torch.float64) / pmax keep = p_ratio < 0.5 p_ratio = p_ratio[keep] @@ -188,8 +211,8 @@ def build_adaptive_sample_list( # (p, q) ∈ [0, pmax) × [0, qmax), mapped to (α, γ) via # FastRot.cc:216-219 — with the negative branch for γ when # p_ratio < q_ratio (gives γ ∈ [0, 1) without negative values). - p_idx = torch.arange(pmax, device=device) - q_idx = torch.arange(qmax, device=device) + p_idx = torch.arange(pmax, device=cpu) + q_idx = torch.arange(qmax, device=cpu) p_ratio = (p_idx.to(torch.float64) / pmax).unsqueeze(1) # (pmax, 1) q_ratio = (q_idx.to(torch.float64) / qmax).unsqueeze(0) # (1, qmax) alpha_frac = torch.fmod(p_ratio + q_ratio, 1.0) @@ -203,34 +226,44 @@ def build_adaptive_sample_list( gamma_frac = gamma_frac.reshape(-1) # Dedup: when p ≥ pmax/2, the (p, q) → (α, γ) map can collide - # with (p − pmax/2, q') for some q' (FastRot.cc:222-241). Phaser - # does an O(pmax·qmax²) loop; we do it via tuple-set in numpy - # which is O(N log N) — same result. - keys = torch.stack([alpha_frac, gamma_frac], dim=-1) - # Round to 1e-6 to match the epsilon comparison in the source. - key_round = (keys * 1_000_000).round().to(torch.int64) - _, uniq_idx = torch.unique(key_round, dim=0, return_inverse=True) - # Keep first occurrence of each unique key. - seen = {} - keep_mask = torch.zeros(alpha_frac.shape[0], dtype=torch.bool, device=device) - for i, k in enumerate(uniq_idx.tolist()): - if k not in seen: - seen[k] = True - keep_mask[i] = True + # with (p − pmax/2, q') for some q' (FastRot.cc:222-241). Keep the + # first occurrence (in original order) of each rounded (α, γ) key. + # Vectorised first-occurrence: stable-sort the unique-group labels, + # mark group boundaries, scatter back — same kept set & order as the + # original dict scan, but no host sync / Python loop. + # Hash the two rounded fracs (each in [0, 1e6]) into one int64 so we + # can use the fast 1-D unique instead of a 2-D row lexsort. + a_round = (alpha_frac * 1_000_000).round().to(torch.int64) + g_round = (gamma_frac * 1_000_000).round().to(torch.int64) + key_hash = a_round * 1_000_001 + g_round + _, uniq_idx = torch.unique(key_hash, return_inverse=True) + n = uniq_idx.shape[0] + order = torch.argsort(uniq_idx, stable=True) + sorted_u = uniq_idx[order] + first_in_sorted = torch.ones(n, dtype=torch.bool, device=cpu) + first_in_sorted[1:] = sorted_u[1:] != sorted_u[:-1] + keep_mask = torch.zeros(n, dtype=torch.bool, device=cpu) + keep_mask[order[first_in_sorted]] = True alpha_frac = alpha_frac[keep_mask] gamma_frac = gamma_frac[keep_mask] n_this = alpha_frac.shape[0] alphas_list.append((alpha_frac * (2.0 * math.pi)).to(dtype)) gammas_list.append((gamma_frac * (2.0 * math.pi)).to(dtype)) - betas_list.append(torch.full((n_this,), beta_rad, dtype=dtype, device=device)) + betas_list.append(torch.full((n_this,), beta_rad, dtype=dtype, device=cpu)) beta_starts.append(beta_starts[-1] + n_this) - alphas = torch.cat(alphas_list) - gammas = torch.cat(gammas_list) - betas_flat = torch.cat(betas_list) + # Concatenate on CPU, then move to the target device in one transfer each. + alphas = torch.cat(alphas_list).to(device) + gammas = torch.cat(gammas_list).to(device) + betas_flat = torch.cat(betas_list).to(device) beta_starts_t = torch.tensor(beta_starts, dtype=torch.int64, device=device) - return alphas, betas_flat, gammas, beta_starts_t, betas_rad + b = torch.arange(bmax, dtype=torch.float64, device=cpu) + betas_rad = (b * grid_sampling_deg * deg2rad).to(device=device, dtype=dtype) + + result = (alphas, betas_flat, gammas, beta_starts_t, betas_rad) + _SAMPLE_LIST_CACHE[cache_key] = result + return result def _bilinear_interp_periodic( diff --git a/torchref/alignment/frf/types.py b/torchref/alignment/frf/types.py index 9497ed60..985baca3 100644 --- a/torchref/alignment/frf/types.py +++ b/torchref/alignment/frf/types.py @@ -73,21 +73,12 @@ def total_samples(self) -> int: class RotationPeak: """A single rotation function peak. - Identical to torchref.alignment.ball_search.RotationPeak so consumers - don't need to change. Convention: Edmonds ZYZ Euler in radians. + Single source of truth across the FRF, the rescore, and the tests. + Convention: Edmonds ZYZ Euler in radians. """ alpha: float beta: float gamma: float - value: float + score: float sigma: float - - @property - def score(self) -> float: - """Alias for ``value`` so this peak is drop-in for - ``ball_search.RotationPeak`` (whose primary field is ``score``). - Lets the production pipeline/align conversion - ``(p.alpha, p.beta, p.gamma, p.score, p.sigma)`` work for both engines. - """ - return self.value diff --git a/torchref/alignment/frf/wigner_d.py b/torchref/alignment/frf/wigner_d.py index 9e22e082..e6c45a5b 100644 --- a/torchref/alignment/frf/wigner_d.py +++ b/torchref/alignment/frf/wigner_d.py @@ -70,6 +70,34 @@ def small_d_stable(L: int, betas: torch.Tensor) -> torch.Tensor: return out +# Cache of the J_y eigendecomposition per l. (w_l, V_l) depend only on l (and +# device), not on β or the data, so they are computed once per (L, device) and +# reused across every FRF call. Keyed on (L, device-str). +_WIGNER_EIG_CACHE: dict = {} + + +def _wigner_eig_table(L: int, device: torch.device): + """Return [(w_l, V_l)] for l ∈ [1, L) — the J_y eigendecomposition per l. + + ``d^l(β) = Re(V_l · diag(e^{-iβ w_l}) · V_l^H)``. ``w_l ≈ [-l..l]`` and + ``V_l`` are independent of β and the data, so they are memoised. + """ + key = (int(L), str(device)) + cached = _WIGNER_EIG_CACHE.get(key) + if cached is not None: + return cached + table = [] + for l in range(1, L): + sz = 2 * l + 1 + p = torch.arange(sz - 1, dtype=torch.float64, device=device) + sup = 0.5 * torch.sqrt((2 * l - p) * (p + 1.0)) + A = torch.diag(sup, 1) - torch.diag(sup, -1) # A = -i J_y + w, V = torch.linalg.eigh(1j * A.to(torch.complex128)) # w∈[-l..l] + table.append((w, V)) + _WIGNER_EIG_CACHE[key] = table + return table + + def wigner_contraction_per_beta( xi_lmn: torch.Tensor, betas: torch.Tensor, @@ -111,12 +139,9 @@ def wigner_contraction_per_beta( S = torch.zeros((n_beta, dim, dim), dtype=torch.complex128, device=device) c = L - 1 S[:, c, c] += xi[0, c, c] # l=0: d^0 = 1 + eig_table = _wigner_eig_table(L, device) # cached (w_l, V_l) for l in range(1, L): - sz = 2 * l + 1 - p = torch.arange(sz - 1, dtype=torch.float64, device=device) - sup = 0.5 * torch.sqrt((2 * l - p) * (p + 1.0)) - A = torch.diag(sup, 1) - torch.diag(sup, -1) # A = -i J_y - w, V = torch.linalg.eigh(1j * A.to(torch.complex128)) # w∈[-l..l] + w, V = eig_table[l - 1] # data-independent phase = torch.exp(-1j * betas.unsqueeze(1) * w.unsqueeze(0)) # (n_beta, sz) VP = V.unsqueeze(0) * phase.unsqueeze(1) # (n_beta, sz, sz) = (k,m,a) d_l = (VP @ V.conj().transpose(-1, -2)).real # (n_beta, sz, sz) diff --git a/torchref/alignment/ml_rotation.py b/torchref/alignment/ml_rotation.py index 42e6242e..2d376491 100644 --- a/torchref/alignment/ml_rotation.py +++ b/torchref/alignment/ml_rotation.py @@ -24,12 +24,12 @@ import torch -from .frf.ball_search import ( - RotationPeak, +from .frf.rotation_utils import ( edmonds_euler_from_rotation_matrix, rotation_matrix_from_edmonds_euler, rotation_matrix_from_edmonds_euler_batch, ) +from .frf.types import RotationPeak from .distributions import ( phaser_log_rel_rice, phaser_log_rel_woolfson, @@ -417,7 +417,7 @@ def sim_mlrf_rescore( sigma_a: Optional[torch.Tensor] = None, ) -> List[RotationPeak]: """ - Rescore a list of peaks from `ball_rotation_search` by the per-shell-fitted + Rescore a list of FRF peaks by the per-shell-fitted Sim Maximum-Likelihood Rotation Function (LLG). Returns a new list sorted by descending LLG with `score = LLG` and `sigma = Z-score(LLG)`. @@ -461,7 +461,7 @@ def sim_mlrf_rescore( # Build all rotation matrices up front. The peak's Euler triple represents # "the rotation applied to the model coords" (synthetic-test convention of - # ball_rotation_search). For ML scoring, we need "the rotation to apply to + # FRF synthetic-test convention). For ML scoring, we need "the rotation to apply to # the current model to align it to obs" — which is R^T. We transpose here. # Vectorised over peaks: previously this list comprehension built M·9 # small (3,3) tensors per dense-R pass. @@ -762,186 +762,3 @@ def m_letf1_rescore( return rescored + tail -# ============================================================================= -# Brute-force ML rotation search over a uniform SO(3) sample -# ============================================================================= - - -def _uniform_random_rotations(n: int, seed: int = 0, dtype=torch.float64) -> torch.Tensor: - """ - Generate `n` uniformly-distributed random rotation matrices via QR of - Gaussian matrices. Returns shape (n, 3, 3) with det = +1. - """ - g = torch.Generator().manual_seed(int(seed)) - A = torch.randn(n, 3, 3, generator=g, dtype=dtype) - Q, R = torch.linalg.qr(A) - # Make det = 1 (flip first column if needed) - diag_sign = torch.sign(torch.diagonal(R, dim1=-2, dim2=-1)) # (n, 3) - Q = Q * diag_sign.unsqueeze(-2) - det = torch.det(Q) - flip = det < 0 - Q[flip, :, 0] = -Q[flip, :, 0] - return Q - - -def brute_ml_rotation_search( - F_obs: torch.Tensor, - hkl_real: torch.Tensor, - s_mag: torch.Tensor, - centric: torch.Tensor, - interpolator: LattmanLoveInterpolator, - real_cell, - n_candidates: int = 5000, - n_shells: int = 15, - batch_size: int = 100, - seed: int = 0, - verbose: int = 0, -) -> List[RotationPeak]: - """ - Evaluate ML LLG on `n_candidates` uniformly-random SO(3) rotations. - Returns peaks sorted by descending LLG. Acts as a shortlist generator that - bypasses the fast ball-search (which has known sphere-sampling limitations - on real-cell HKL data). - - Cost: ~10ms per candidate at default batch_size on CPU; ~50s for 5000 cands. - - Returns - ------- - list of RotationPeak (sorted by LLG descending) - `score = LLG`, `sigma = Z-score across the candidate set`. - """ - R_all = _uniform_random_rotations(n_candidates, seed=seed) - peaks_in = [] - for k in range(n_candidates): - a, b, g = edmonds_euler_from_rotation_matrix(R_all[k]) - peaks_in.append(RotationPeak(a, b, g, score=0.0, sigma=0.0)) - return sim_mlrf_rescore( - peaks_in, F_obs, hkl_real, s_mag, centric, interpolator, real_cell, - n_shells=n_shells, n_refine=n_candidates, batch_size=batch_size, - verbose=verbose, - ) - - -# ============================================================================= -# BRF — Brute Rotation Function (Phaser stage that refines top FRF peaks via a -# denser local rotation sampling + full Rice/Woolfson LL). -# ============================================================================= - - -def _random_rotation_in_cone( - n: int, radius_rad: float, generator: torch.Generator, -) -> torch.Tensor: - """``n`` random rotation matrices uniformly inside an angular cone of given - radius from identity. Axis ∼ uniform-on-sphere, angle ∼ uniform on - ``[0, radius_rad]`` (proper Haar measure on the cone would weight angle by - sin²(θ/2); for small radii this approximation is essentially uniform). - - Returns ``(n, 3, 3)`` float64 tensor. - """ - # Random unit axes via Gaussian normalization. - axes = torch.randn(n, 3, generator=generator, dtype=torch.float64) - axes = axes / axes.norm(dim=-1, keepdim=True).clamp(min=1e-30) - angles = torch.rand(n, generator=generator, dtype=torch.float64) * radius_rad - # Rodrigues: R = I + sin(θ)·K + (1−cos(θ))·K² with K skew of axis. - cos = angles.cos().view(-1, 1, 1) - sin = angles.sin().view(-1, 1, 1) - K = torch.zeros(n, 3, 3, dtype=torch.float64) - K[:, 0, 1] = -axes[:, 2]; K[:, 0, 2] = axes[:, 1] - K[:, 1, 0] = axes[:, 2]; K[:, 1, 2] = -axes[:, 0] - K[:, 2, 0] = -axes[:, 1]; K[:, 2, 1] = axes[:, 0] - I = torch.eye(3, dtype=torch.float64).expand(n, 3, 3) - K2 = K @ K - return I + sin * K + (1.0 - cos) * K2 - - -def brf_refine( - peaks: List[RotationPeak], - F_obs: torch.Tensor, - hkl_real: torch.Tensor, - s_mag: torch.Tensor, - centric: torch.Tensor, - interpolator: LattmanLoveInterpolator, - real_cell, - sym_mats: torch.Tensor, - *, - n_top: int = 100, - n_perturb: int = 10, - angular_radius_deg: float = 3.0, - seed: int = 42, - verbose: int = 0, - **m_letf1_kwargs, -) -> List[RotationPeak]: - """Brute Rotation Function — Phaser's BRF stage (post-FRF rotation refinement). - - For each of the top ``n_top`` FRF peaks, sample ``n_perturb`` random - rotations within an angular cone of radius ``angular_radius_deg`` from the - peak; score all (original + perturbations) via :func:`m_letf1_rescore`. - Returns the rescored set sorted by descending LL — the truth orientation - typically lurks ≤ FRF-grid-spacing degrees from the FRF peak nearest to it, - so this local refinement can recover peaks the coarse FRF grid missed. - - Total LL evaluations: ``n_top · (n_perturb + 1)``. With defaults this is - ``100 × 11 = 1100`` — ~2× the cost of the standard top-500 rescore. - - Parameters - ---------- - peaks - FRF peak list (sorted by descending FRF score). - n_top - Refine around the top ``n_top`` FRF peaks. Should be ≥ the expected - rank of the truth peak — for hard cases (e.g., 2DQ6 FRF rank ≈50), use - ``n_top ≥ 100``. - n_perturb - Number of random rotations sampled per peak (in addition to the - peak itself, which is always included). - angular_radius_deg - Cone half-angle for perturbations. Match the FRF grid resolution - (~3° at ``grid_sampling_deg=3.0``). - seed - RNG seed for reproducibility. - **m_letf1_kwargs - Forwarded to :func:`m_letf1_rescore` — e.g. ``apply_bulk_solvent=True``, - ``vrms_strategy="oeffner"``, ``vrms_n_residues=...``, etc. - """ - if not peaks: - return [] - n_top = min(n_top, len(peaks)) - top = peaks[:n_top] - g = torch.Generator().manual_seed(int(seed)) - radius_rad = math.radians(float(angular_radius_deg)) - - # One big batch of perturbations: (n_top * n_perturb, 3, 3). - pert_R = _random_rotation_in_cone(n_top * n_perturb, radius_rad, g) - pert_R = pert_R.view(n_top, n_perturb, 3, 3) - - out_peaks: List[RotationPeak] = [] - for i, peak in enumerate(top): - # Include the original peak (perturbation theta=0). - out_peaks.append(RotationPeak( - alpha=peak.alpha, beta=peak.beta, gamma=peak.gamma, - score=peak.score, sigma=peak.sigma, - )) - R_orig = rotation_matrix_from_edmonds_euler( - peak.alpha, peak.beta, peak.gamma, - ).to(torch.float64) - for j in range(n_perturb): - R_pert = R_orig @ pert_R[i, j] - a, b, c = edmonds_euler_from_rotation_matrix(R_pert) - out_peaks.append(RotationPeak( - alpha=a, beta=b, gamma=c, score=peak.score, sigma=peak.sigma, - )) - - if verbose > 0: - print( - f" BRF: {len(out_peaks)} candidates " - f"({n_top} top × ({n_perturb}+1) perturbations, ±{angular_radius_deg}°)", - flush=True, - ) - - return m_letf1_rescore( - out_peaks, F_obs, hkl_real, s_mag, centric, interpolator, real_cell, - sym_mats, - n_refine=len(out_peaks), - verbose=verbose, - **m_letf1_kwargs, - ) diff --git a/torchref/alignment/patterson_filter.py b/torchref/alignment/patterson_filter.py deleted file mode 100644 index 3703b197..00000000 --- a/torchref/alignment/patterson_filter.py +++ /dev/null @@ -1,163 +0,0 @@ -""" -Restrict the Patterson function to within a sphere of radius Ω (molecule -diameter) so the rotation function sees only intramolecular vectors. - -This is the operational equivalent of Phaser's χ_Ω weighting (LERF1, §2.1.3): -the sphere of integration excludes intermolecular Patterson peaks that come -from cross-vectors between symmetry-related copies in the crystal — those are -the dominant contamination of |F_obs|² on real data and are why the bare -Crowther-style rotation function (|E|² · |E|² overlap) fails on real F_obs. - -Implementation --------------- -- Place |F|² values onto a regular 3-D reciprocal-space grid of a large cubic - cell (radius ≥ Ω, so the sphere fits inside). -- 3-D FFT → real-space Patterson in Cartesian coordinates of the cubic cell. -- Apply spherical mask: zero values at |r| > Ω. -- IFFT → sphere-restricted Patterson coefficients on the cubic grid. -- Extract values at the original real-cell HKL positions via trilinear - interpolation. -""" - -from __future__ import annotations - -import math -from typing import Tuple - -import torch - -from torchref.symmetry.cell import Cell -from torchref.base.reciprocal.interpolation import ( - interpolate_structure_factor_from_grid, -) - - -def _make_cubic_grid(F_sq_real: torch.Tensor, - hkl_real: torch.Tensor, - real_cell: Cell, - cubic_side: float, - gridsize: int) -> Tuple[torch.Tensor, Cell]: - """ - Place |F|² values onto a (gridsize, gridsize, gridsize) cubic-cell - reciprocal grid in FFT layout. Multiple real-HKL points falling into the - same cubic-grid cell are accumulated. - - Returns - ------- - grid : torch.Tensor (real), shape (N, N, N) - Cubic reciprocal grid with FFT layout (DC at index 0, negative HKL - wrapped to high indices). - cubic_cell : Cell - """ - device = F_sq_real.device - dtype = F_sq_real.dtype - N = int(gridsize) - cubic_cell = Cell( - [cubic_side, cubic_side, cubic_side, 90.0, 90.0, 90.0], - dtype=torch.float32, device=device, - ) - # h_real → Cartesian s - rec_real = real_cell.reciprocal_basis_matrix.to(device).to(dtype) - s = hkl_real.to(dtype) @ rec_real # (M, 3) - # Cartesian s → cubic-cell HKL (orthogonal cubic, so rec_basis = (1/a)·I → inv = a·I) - h_cubic_f = s * cubic_side # (M, 3), real float - h_idx = torch.round(h_cubic_f).long() # (M, 3), nearest integer cubic HKL - # Wrap into [0, N) (negative HKL → high indices) - h_idx = h_idx % N - # Accumulate values into the grid - flat_idx = (h_idx[:, 0] * N + h_idx[:, 1]) * N + h_idx[:, 2] - grid_flat = torch.zeros(N * N * N, dtype=dtype, device=device) - grid_flat.index_add_(0, flat_idx, F_sq_real) - grid = grid_flat.view(N, N, N) - return grid, cubic_cell - - -def _sphere_mask(N: int, cubic_side: float, omega: float, - dtype: torch.dtype, device: torch.device) -> torch.Tensor: - """ - Build a sphere indicator mask on an (N, N, N) real-space grid covering the - cubic cell (side = cubic_side). The mask is 1 inside |r| < omega, 0 outside. - The grid layout is FFT-compatible: index 0 = origin, indices > N/2 wrap to - negative coordinates. - """ - idx = torch.arange(N, device=device) - # Voxel offsets, wrapped to [-N/2, N/2) - offsets = torch.where(idx < (N + 1) // 2, idx, idx - N).to(dtype) * (cubic_side / N) - rx, ry, rz = torch.meshgrid(offsets, offsets, offsets, indexing="ij") - r2 = rx ** 2 + ry ** 2 + rz ** 2 - return (r2 < omega ** 2).to(dtype) - - -def restrict_to_sphere( - F_squared: torch.Tensor, - hkl_real: torch.Tensor, - real_cell: Cell, - omega_A: float, - cubic_side_A: float | None = None, - gridsize: int | None = None, - max_res_A: float = 3.0, -) -> Tuple[torch.Tensor, Cell, torch.Tensor]: - """ - Sphere-restrict a Patterson-like coefficient set defined at the real-cell - HKL positions. - - Parameters - ---------- - F_squared : torch.Tensor (real), shape (M,) - Values to filter (e.g. (|E|² − 1) per shell, or |F|² with origin removed). - hkl_real : torch.Tensor, shape (M, 3) - Miller indices in the real (crystal) cell. - real_cell : Cell - omega_A : float - Sphere-of-integration radius (Å). Typically the molecule's bounding-box - radius — vectors longer than 2·omega are intermolecular and get filtered out. - cubic_side_A : float, optional - Cubic cell side (Å) for the FFT. Default: 4 · omega_A (so the Patterson - sphere fits comfortably with padding to avoid wraparound aliasing). - gridsize : int, optional - Grid size (one dim). Default: 2 · ceil(cubic_side / max_res_A). - max_res_A : float, default 3.0 - Resolution limit (Å) for the cubic grid spacing. - - Returns - ------- - filtered_values : torch.Tensor (real), shape (M,) - Sphere-restricted Patterson coefficients at the SAME real-cell HKL set. - cubic_cell : Cell - cubic_grid : torch.Tensor (real complex-real), shape (N, N, N) - Final sphere-restricted reciprocal grid (kept for caller introspection). - """ - if cubic_side_A is None: - cubic_side_A = 4.0 * omega_A - if gridsize is None: - gridsize = 2 * int(math.ceil(cubic_side_A / max_res_A)) - # Round to even for cleaner FFT - if gridsize % 2: - gridsize += 1 - - grid_recip, cubic_cell = _make_cubic_grid( - F_squared, hkl_real, real_cell, cubic_side_A, gridsize, - ) - # FFT to real space (the grid is real-valued) - patterson_real = torch.fft.fftn(grid_recip, dim=(0, 1, 2)).real - mask = _sphere_mask(gridsize, cubic_side_A, omega_A, - dtype=patterson_real.dtype, device=patterson_real.device) - patterson_real_masked = patterson_real * mask - # IFFT back (will be approximately real for a real Patterson) - grid_recip_filtered = torch.fft.ifftn( - patterson_real_masked, dim=(0, 1, 2) - ).real - - # Extract values at original real-cell HKL positions via trilinear - # interpolation. Note: interpolate_structure_factor_from_grid expects a - # complex grid for `interpolate_amplitude=False` and a complex or real grid - # for `interpolate_amplitude=True`. Our grid is real — wrap in complex for - # the utility. - rec_real = real_cell.reciprocal_basis_matrix.to(F_squared.device).to(F_squared.dtype) - s = hkl_real.to(F_squared.dtype) @ rec_real - h_cubic_f = s * cubic_side_A # cubic-cell float HKL - grid_complex = grid_recip_filtered.to(torch.complex64) - filtered = interpolate_structure_factor_from_grid( - grid_complex, h_cubic_f, interpolate_amplitude=False, - ).real.to(F_squared.dtype) - return filtered, cubic_cell, grid_recip_filtered diff --git a/torchref/alignment/pipeline.py b/torchref/alignment/pipeline.py index db078957..fc7bab9e 100644 --- a/torchref/alignment/pipeline.py +++ b/torchref/alignment/pipeline.py @@ -17,11 +17,8 @@ from torchref.config import get_default_device, get_float_dtype -from .frf.ball_search import ( - ball_rotation_search, - rotation_matrix_from_edmonds_euler, - RotationPeak, -) +from .frf.rotation_utils import rotation_matrix_from_edmonds_euler +from .frf.types import RotationPeak from .translation import fft_translation_search_torch, TranslationPeak from .rigid_body import RigidBodyRefinement, RigidBodyResult from .clashscore import ClashScoreCalculator, AtomSampler @@ -252,7 +249,6 @@ def run( L: int = 48, P: int = 24, cluster_threshold_deg: Optional[float] = None, - engine: str = "frf_separate", ) -> List[MRSolution]: """ Run full MR pipeline with early stopping. @@ -297,7 +293,7 @@ def run( if self.verbose: print("Step 1: Rotation search...") rotation_peaks = self._rotation_search( - n_rotation_peaks, d_min, d_max, L, P, engine=engine + n_rotation_peaks, d_min, d_max, L, P, ) if not rotation_peaks: @@ -416,42 +412,21 @@ def _rotation_search( d_max: float, L: int, P: int, - engine: str = "frf_separate", ) -> list: - """Run the rotation search; return list of (α, β, γ, score, σ) tuples. + """Run the rotation search via the Phaser-faithful FRF. - ``engine="frf_separate"`` (default) uses the validated Phaser-faithful - engine (dense calc + auto_lmax + obs-unroll + no_grad) via the shared - ``align`` helpers; ``engine="ball"`` uses the legacy E-value ball search. + Single engine post-consolidation: dense calc + auto_lmax + obs-unroll + + no_grad. Returns ``(α, β, γ, score, σ)`` tuples. """ - if engine not in ("frf_separate", "ball"): - raise ValueError( - f"engine={engine!r}; expected 'frf_separate' (default) or 'ball'." - ) - if engine == "frf_separate": - from .align import _prepare_frf_inputs, _run_frf_separate_rotation - frf = _prepare_frf_inputs( - self.model, self.data, - d_min=d_min, d_max=d_max, n_shells=20, verbose=self.verbose, - ) - peaks = _run_frf_separate_rotation( - self.model, self.data, frf, - n_peaks=n_peaks, verbose=self.verbose, - ) - return [(p.alpha, p.beta, p.gamma, p.score, p.sigma) for p in peaks] - - # engine == "ball" — legacy ball-harmonic E-value search. - E_obs, s_obs = self._get_e_values_obs(d_min, d_max) - E_calc, s_calc = self._get_e_values_calc(d_min, d_max) - - _C, _alphas, _betas, _gammas, peaks = ball_rotation_search( - s_obs, E_obs, s_calc, E_calc, - L=L, P=P, n_peaks=n_peaks, - d_min=d_min, d_max=d_max, - refine_subvoxel=True, n_refine=min(n_peaks, 50), - sigma_threshold=0.0, + from .align import _prepare_frf_inputs, _run_frf_separate_rotation + frf = _prepare_frf_inputs( + self.model, self.data, + d_min=d_min, d_max=d_max, n_shells=20, verbose=self.verbose, + ) + peaks = _run_frf_separate_rotation( + self.model, self.data, frf, + n_peaks=n_peaks, verbose=self.verbose, ) - # Convert dataclass to tuple form expected by the rest of the pipeline. return [(p.alpha, p.beta, p.gamma, p.score, p.sigma) for p in peaks] def _translation_search( diff --git a/torchref/alignment/sh.py b/torchref/alignment/sh.py index 983c8560..a1292ad4 100644 --- a/torchref/alignment/sh.py +++ b/torchref/alignment/sh.py @@ -17,10 +17,15 @@ from __future__ import annotations import math +import os +import time from typing import Optional, Tuple import torch +_PROFILE = bool(os.environ.get("FRF_PROFILE")) +_YLM_PROF = {"recurrence": 0.0, "assembly": 0.0} + def _bar_legendre_recurrence( cos_theta: torch.Tensor, @@ -50,29 +55,43 @@ def _bar_legendre_recurrence( inv_sqrt_4pi = 1.0 / math.sqrt(4.0 * math.pi) bar_P[..., 0, 0] = inv_sqrt_4pi - # Sectoral recurrence on the diagonal m == l: - # bar_P_m^m = sqrt((2m+1)/(2m)) · sinθ · bar_P_{m-1}^{m-1} - for m in range(1, L): - factor = math.sqrt((2.0 * m + 1.0) / (2.0 * m)) - bar_P[..., m, m] = factor * sin_theta * bar_P[..., m - 1, m - 1] - - # Vertical recurrence (l > m, fixed m): - # bar_P_l^m = a_l^m · cosθ · bar_P_{l-1}^m - b_l^m · bar_P_{l-2}^m - # with a_l^m = sqrt((2l-1)(2l+1)/((l-m)(l+m))) - # b_l^m = sqrt((2l+1)(l+m-1)(l-m-1)/((l-m)(l+m)(2l-3))) [0 when l = m+1] - for m in range(0, L - 1): - # l = m+1 step (b term vanishes because l-m-1 = 0) - l = m + 1 - a = math.sqrt((2.0 * l - 1.0) * (2.0 * l + 1.0) / ((l - m) * (l + m))) - bar_P[..., l, m] = a * cos_theta * bar_P[..., l - 1, m] - # l from m+2 to L-1 - for l in range(m + 2, L): - a = math.sqrt((2.0 * l - 1.0) * (2.0 * l + 1.0) / ((l - m) * (l + m))) - b = math.sqrt( - (2.0 * l + 1.0) * (l + m - 1.0) * (l - m - 1.0) - / ((l - m) * (l + m) * (2.0 * l - 3.0)) - ) - bar_P[..., l, m] = a * cos_theta * bar_P[..., l - 1, m] - b * bar_P[..., l - 2, m] + # Precompute the recurrence coefficients as (L, L) tables, indexed [l, m]: + # a_l^m = sqrt((2l-1)(2l+1)/((l-m)(l+m))) + # b_l^m = sqrt((2l+1)(l+m-1)(l-m-1)/((l-m)(l+m)(2l-3))) [0 when l = m+1] + # b vanishes at l = m+1 (factor l-m-1 = 0), so no special-casing is needed. + # Computed in float64 then cast to `dtype` (matches the original, which used + # float64 python scalars multiplied into the working-dtype tensors). + ll = torch.arange(L, dtype=torch.float64, device=device).view(L, 1) + mm = torch.arange(L, dtype=torch.float64, device=device).view(1, L) + valid = (ll > mm) # l > m (vertical recurrence region) + denom = (ll - mm) * (ll + mm) + denom_safe = torch.where(valid, denom, torch.ones_like(denom)) + a_coef = torch.sqrt((2.0 * ll - 1.0) * (2.0 * ll + 1.0) / denom_safe) + b_num = (2.0 * ll + 1.0) * (ll + mm - 1.0) * (ll - mm - 1.0) + b_den = denom_safe * (2.0 * ll - 3.0) + b_coef = torch.sqrt(torch.clamp(b_num / torch.where(b_den == 0, torch.ones_like(b_den), b_den), min=0.0)) + a_coef = torch.where(valid, a_coef, torch.zeros_like(a_coef)).to(dtype) + b_coef = torch.where(valid, b_coef, torch.zeros_like(b_coef)).to(dtype) + # Sectoral diagonal factor sqrt((2m+1)/(2m)) for m = l. + m_arange = torch.arange(L, dtype=torch.float64, device=device) + sect = torch.sqrt((2.0 * m_arange + 1.0) / (2.0 * m_arange).clamp(min=1.0)).to(dtype) + + cos_e = cos_theta.unsqueeze(-1) # (..., 1) + # Single loop over l; at each l update all m ∈ [0, l] at once. The vertical + # recurrence (m < l) and the sectoral diagonal (m = l) both read only level + # l-1 (and l-2), already computed — so this is the same recurrence as the + # original double loop, just reordered to vectorise over m. + for l in range(1, L): + prev1 = bar_P[..., l - 1, :l] # (..., l) + if l >= 2: + prev2 = bar_P[..., l - 2, :l] # (..., l) + else: + prev2 = torch.zeros_like(prev1) + bar_P[..., l, :l] = ( + a_coef[l, :l] * cos_e * prev1 - b_coef[l, :l] * prev2 + ) + # Sectoral m == l. + bar_P[..., l, l] = sect[l] * sin_theta * bar_P[..., l - 1, l - 1] return bar_P @@ -81,6 +100,7 @@ def evaluate_ylm( theta: torch.Tensor, phi: torch.Tensor, L: int, + l_indices: Optional[torch.Tensor] = None, ) -> torch.Tensor: """ Evaluate Y_{l,m}(θ, φ) for all (l, m) with l ∈ [0, L), m ∈ [-(L-1), L-1]. @@ -93,12 +113,21 @@ def evaluate_ylm( Azimuthal angle, shape (...,), values in [0, 2π). L : int Maximum SH degree (exclusive: l_max = L - 1). + l_indices : torch.Tensor, optional + If given (1-D long tensor of l values), assemble and return Y only for + those degrees, shape ``(..., len(l_indices), 2L-1)`` with row ``i`` + holding ``Y_{l_indices[i], m}``. The Legendre recurrence still runs over + the full degree range (it is a recurrence), but the costly complex Y + assembly is restricted to the requested rows. Used by the FRF expansion, + which only needs even degrees (the odd-l and l=0 rows are zeroed by + Patterson centrosymmetry) — halving the dominant assembly cost. Returns ------- Y : torch.Tensor, complex - Shape (..., L, 2L-1). Y[..., l, L-1+m] = Y_{l,m}(θ, φ) for |m| ≤ l, - zero otherwise. dtype is complex128 if input is float64, else complex64. + Shape (..., L, 2L-1) (or (..., len(l_indices), 2L-1) if l_indices given). + ``Y[..., l, L-1+m] = Y_{l,m}(θ, φ)`` for |m| ≤ l, zero otherwise. dtype is + complex128 if input is float64, else complex64. """ assert theta.shape == phi.shape, "theta and phi must have the same shape" @@ -114,29 +143,39 @@ def evaluate_ylm( cos_theta = torch.cos(theta) sin_theta = torch.sin(theta).clamp(min=0.0) # numerical floor at the poles + if _PROFILE: + t0 = time.perf_counter() bar_P = _bar_legendre_recurrence(cos_theta, sin_theta, L) # (..., L, L) + if l_indices is not None: + bar_P = bar_P[..., l_indices, :] # (..., n_sel, L) — even rows only + n_rows = bar_P.shape[-2] + if _PROFILE: + _YLM_PROF["recurrence"] += time.perf_counter() - t0 + t0 = time.perf_counter() # Y_{l,m}(θ,φ) = (-1)^m · bar_P_l^m(cosθ) · e^{i m φ} for m ≥ 0 # Y_{l,-m} = (-1)^m · conj(Y_{l,m}) for m > 0 - Y = torch.zeros((*theta.shape, L, 2 * L - 1), dtype=complex_dtype, device=device) + Y = torch.zeros((*theta.shape, n_rows, 2 * L - 1), dtype=complex_dtype, device=device) - # Precompute e^{i m φ} for m = 0..L-1 - # (use stacking to keep things vectorized) + # Precompute e^{i m φ} for m = 0..L-1 and the Condon-Shortley signs (-1)^m. m_vals = torch.arange(L, dtype=real_dtype, device=device) m_phi = phi.unsqueeze(-1) * m_vals # (..., L) expo = torch.complex(torch.cos(m_phi), torch.sin(m_phi)) # (..., L), e^{i m φ} + signs = ((-1.0) ** m_vals).to(complex_dtype) # (L,) - # Fill m ≥ 0 columns - for m in range(L): - sign = (-1.0) ** m - # bar_P[..., :, m] is (..., L); only entries l >= m are non-zero (others left at 0) - Y[..., :, L - 1 + m] = sign * bar_P[..., :, m].to(complex_dtype) * expo[..., m].unsqueeze(-1) + # Fill m ≥ 0 columns (L-1 .. 2L-2): Y_{l,m} = (-1)^m bar_P_l^m e^{imφ}, + # vectorised over (l, m). phase = (-1)^m e^{imφ} broadcasts over l. + phase_pos = (signs * expo).unsqueeze(-2) # (..., 1, L) + Y_pos = bar_P.to(complex_dtype) * phase_pos # (..., n_rows, L) + Y[..., :, L - 1:] = Y_pos - # Fill m < 0 columns by hermitian symmetry: Y_{l,-m} = (-1)^m · conj(Y_{l,m}) - for m in range(1, L): - sign = (-1.0) ** m - Y[..., :, L - 1 - m] = sign * torch.conj(Y[..., :, L - 1 + m]) + # Fill m < 0 columns by hermitian symmetry: Y_{l,-m} = (-1)^m conj(Y_{l,m}). + # For m = 1..L-1 these land in columns L-2 .. 0, i.e. the reversed prefix. + neg = signs[1:] * torch.conj(Y_pos[..., :, 1:]) # (..., L, L-1), m = 1..L-1 + Y[..., :, : L - 1] = torch.flip(neg, dims=(-1,)) + if _PROFILE: + _YLM_PROF["assembly"] += time.perf_counter() - t0 return Y From 9615447a8e2090da89e505f0d2d2bbd9268a9d8f Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 11 Jun 2026 14:23:37 +0200 Subject: [PATCH 006/250] working on alignement --- torchref/alignment/align.py | 40 +- torchref/alignment/frf/rotation_utils.py | 31 + torchref/alignment/ml_rotation.py | 720 ++++++++++++++++++----- 3 files changed, 639 insertions(+), 152 deletions(-) diff --git a/torchref/alignment/align.py b/torchref/alignment/align.py index f88083e4..e8950274 100644 --- a/torchref/alignment/align.py +++ b/torchref/alignment/align.py @@ -617,7 +617,7 @@ def align_model_to_data( L: int = 48, n_shells: int = 20, n_rotation_peaks: int = 500, - n_ml_refine: int = 500, + n_ml_refine: int = 20, # rescore only the top-20 FRF peaks (refinement use case) ll_max_res_A: float = 3.0, ll_padding_factor: float = 2.0, verbose: int = 0, @@ -646,6 +646,12 @@ def align_model_to_data( frf_lmax_cap: int = 48, frf_dense_pad: float = 2.0, rescore_engine: str = "m_letf1", + rescore_scat_mode: str = "legacy", + subpeak_refine: bool = False, + subpeak_refine_k: int = -1, + subpeak_refine_step_deg: float = 1.5, + subpeak_refine_iters: int = 1, + subpeak_refine_max_move_deg: Optional[float] = 1.5, ) -> "ModelFT": """Run full MR alignment of ``model`` against ``data``. @@ -862,7 +868,39 @@ def align_model_to_data( n_refine=min(len(peaks), n_ml_refine), batch_size=50, verbose=verbose, + scat_mode=rescore_scat_mode, ) + if subpeak_refine: + # Sharpen the top candidate orientations on the (now-corrected) + # ML-LLG surface before the translation search — the FTF is sensitive + # to orientation error and the FRF grid leaves up to ~half a grid step + # (~1°). Rebuild the LLG context once (cheap: one identity interpolator + # eval) and take a quadratic tangent-space Newton step per top-K peak. + from .ml_rotation import _build_llg_context, quadratic_llg_refine + timer.start("4b_subpeak_refine") + _ctx = _build_llg_context( + F_obs, hkl, s_mag, centric, ll, data.cell, + data.spacegroup.matrices.to(torch.float64).to(device), + n_shells=max(n_shells // 2, 8), batch_size=50, + scat_mode=rescore_scat_mode, + ) + _k = subpeak_refine_k if subpeak_refine_k > 0 else n_rotation_candidates + _k = min(_k, len(rescored)) + rescored = quadratic_llg_refine( + rescored, _ctx, + k_refine=_k, + step_deg=subpeak_refine_step_deg, + iterations=subpeak_refine_iters, + max_move_deg=subpeak_refine_max_move_deg, + verbose=verbose, + ) + timer.stop("4b_subpeak_refine") + if verbose > 0: + print( + f"fit_to_data: sub-peak refined top {_k} orientations " + f"on the ML-LLG surface (step={subpeak_refine_step_deg}°).", + flush=True, + ) else: # rescore_engine == "sim" — legacy Sim/Rice approximation rescored = sim_mlrf_rescore( peaks, F_obs, hkl, s_mag, centric, ll, data.cell, diff --git a/torchref/alignment/frf/rotation_utils.py b/torchref/alignment/frf/rotation_utils.py index 8f1d2537..f4e9d4b2 100644 --- a/torchref/alignment/frf/rotation_utils.py +++ b/torchref/alignment/frf/rotation_utils.py @@ -81,6 +81,37 @@ def edmonds_euler_from_rotation_matrix(R: torch.Tensor) -> Tuple[float, float, f return alpha, beta, gamma +def axis_angle_to_matrix(omega: torch.Tensor) -> torch.Tensor: + """Rodrigues axis-angle → SO(3). ``omega = θ · axis`` (radians). + + Accepts ``(3,)`` for a single rotation or ``(..., 3)`` for a batched stack + and returns ``(3, 3)`` or ``(..., 3, 3)``. The small-θ limit is handled + implicitly (sin θ→0, (1−cos θ)→0 ⇒ R→I); ``clamp(min=1e-30)`` guards the + axis normalisation at θ=0. Mirrors ``align._rodrigues`` but lives here so + both the rescore and the alignment pipeline can share it without a circular + import. + """ + if omega.dtype not in (torch.float32, torch.float64): + omega = omega.to(torch.float64) + single = omega.dim() == 1 + if single: + omega = omega.unsqueeze(0) + th = omega.norm(dim=-1, keepdim=True) # (..., 1) + axis = omega / th.clamp(min=1e-30) # (..., 3) + zeros = torch.zeros_like(axis[..., 0]) + K = torch.stack([ + torch.stack([zeros, -axis[..., 2], axis[..., 1]], dim=-1), + torch.stack([axis[..., 2], zeros, -axis[..., 0]], dim=-1), + torch.stack([-axis[..., 1], axis[..., 0], zeros], dim=-1), + ], dim=-2) # (..., 3, 3) + th_b = th.unsqueeze(-1) # (..., 1, 1) + eye = torch.eye(3, dtype=omega.dtype, device=omega.device).expand( + *omega.shape[:-1], 3, 3 + ) + R = eye + torch.sin(th_b) * K + (1.0 - torch.cos(th_b)) * (K @ K) + return R.squeeze(0) if single else R + + def rotation_angular_distance_deg(R1: torch.Tensor, R2: torch.Tensor) -> float: """Geodesic distance on SO(3) in degrees: ``arccos((tr(R1 R2^T) − 1)/2)``.""" R = R1.to(torch.float64) @ R2.to(torch.float64).T diff --git a/torchref/alignment/ml_rotation.py b/torchref/alignment/ml_rotation.py index 2d376491..3c602117 100644 --- a/torchref/alignment/ml_rotation.py +++ b/torchref/alignment/ml_rotation.py @@ -20,11 +20,13 @@ from __future__ import annotations import math -from typing import List, Optional +from dataclasses import dataclass +from typing import Callable, List, Optional import torch from .frf.rotation_utils import ( + axis_angle_to_matrix, edmonds_euler_from_rotation_matrix, rotation_matrix_from_edmonds_euler, rotation_matrix_from_edmonds_euler_batch, @@ -69,6 +71,25 @@ def _normalize_to_e(F: torch.Tensor, shell_idx: torch.Tensor, return F / _per_shell_sqrt_mean(F, shell_idx, n_shells) +def _normalize_to_e_epsilon( + F: torch.Tensor, shell_idx: torch.Tensor, n_shells: int, eps: torch.Tensor, +) -> torch.Tensor: + """ε-corrected Wilson E: ``E²_h = (F²_h/ε_h) / ⟨F²/ε⟩_shell``. + + Matches Phaser's obs E (``E = F/sqrt(ε·Σ_N)``) and the FRF's + :func:`torchref.alignment.frf.preprocessing.wilson_normalise_epsilon`. The + plain :func:`_normalize_to_e` (no ε) over-counts axial reflections (ε>1) on + high-symmetry spacegroups, letting them dominate the ``-(E²+eImove)/V`` term + and blind the m_LETF1 orientation discrimination. + """ + I_corr = (F * F) / eps.clamp(min=1.0) + sum_shell = torch.zeros(n_shells, dtype=I_corr.dtype, device=F.device) + sum_shell.scatter_add_(0, shell_idx, I_corr) + count = torch.bincount(shell_idx, minlength=n_shells).to(I_corr.dtype) + mean_shell = (sum_shell / count.clamp(min=1.0)).clamp(min=1e-30) + return (I_corr / mean_shell.index_select(0, shell_idx)).clamp(min=0.0).sqrt() + + def _per_shell_sqrt_mean(F: torch.Tensor, shell_idx: torch.Tensor, n_shells: int) -> torch.Tensor: """Per-reflection ``sqrt(_shell)``. Wilson-normalisation denominator. @@ -514,11 +535,534 @@ def sim_mlrf_rescore( # ============================================================================= -# Phaser-faithful m_LETF1 rescore: NSYMP calc sum + V(h) budget + Rice/Woolfson -# logRel formulas (DataMR.cc:1326-1429). +# Phaser-faithful m_LETF1 rescore: unique-orbit calc sum + V(h) budget + +# Rice/Woolfson logRel formulas (DataMR.cc:1326-1429). +# +# The per-orientation LLG evaluator is factored out of `m_letf1_rescore` into a +# reusable `_LLGContext` + `_llg_for_orientations`, so the sub-peak refiner +# (`quadratic_llg_refine`) optimises the *same* likelihood the rescore ranks on. # ============================================================================= +@dataclass +class _LLGContext: + """Rotation-independent context for the m_LETF1 per-orientation LLG. + + Built once by :func:`_build_llg_context`; consumed by + :func:`_llg_for_orientations` (rescore) and :func:`quadratic_llg_refine` + (sub-peak optimiser). Everything here depends only on the data + σ_A model, + not on the candidate orientation. + """ + + interpolator: LattmanLoveInterpolator + real_cell: object + unrolled_hkl: torch.Tensor # (M, 3) float64 — distinct orbit mates + asu_idx: torch.Tensor # (M,) long — ASU reflection each mate maps to + N: int # number of ASU reflections + E_obs_b: torch.Tensor # (1, N) + V_b: torch.Tensor # (1, N) + eImove_prefac: torch.Tensor # (1, N) = ε·σ_A²/n_ops + sqrt_mean_per_m: torch.Tensor # (M,) per-mate E-normaliser + centric_b: torch.Tensor # (1, N) bool + dw_per_m: Optional[torch.Tensor] # (M,) or None — Wilson-B Debye-Waller + dtype: torch.dtype + batch_size: int + + +def _build_llg_context( + F_obs: torch.Tensor, + hkl_real: torch.Tensor, + s_mag: torch.Tensor, + centric: torch.Tensor, + interpolator: LattmanLoveInterpolator, + real_cell, + sym_mats: torch.Tensor, + *, + n_shells: int = 20, + batch_size: int = 50, + sigma_a: Optional[torch.Tensor] = None, + eps_factor: Optional[torch.Tensor] = None, + apply_bulk_solvent: bool = False, + solvent_fsol: float = 0.95, + solvent_bsol: float = 300.0, + vrms_strategy: str = "fixed", + vrms_n_residues: Optional[int] = None, + vrms_identity: float = 1.0, + apply_wilson_b: bool = False, + wilson_b_value: Optional[float] = None, + scat_mode: str = "legacy", +) -> _LLGContext: + """Build the rotation-independent m_LETF1 LLG context (DataMR.cc:1326-1429). + + Two corrections vs. the original implementation, both borrowed from the FRF's + own high-symmetry fixes: + + * **Unique-orbit calc sum.** The moving-model intensity sums ``|E_calc|²`` over + the **distinct** orbit mates via + :func:`torchref.alignment.frf.preprocessing.epsilon_aware_unroll` + (Phaser's ``if(!duplicate(isym))``), not all ``n_ops`` raw mates. Summing all + mates over-weights axial reflections (ε>1) by ε(h) and orientation-blinds + high-symmetry spacegroups (the 4BX9/6G9X rank-360+ failure). + * **σ_A Eterm convention.** σ_A uses ``eterm_sigma_a`` (the ``2π²/3`` isotropic + Eterm, Ensemble.cc:42 — matching the FRF/Phaser), not the ``2π²`` Luzzati + form which falls off ~3× too fast. + """ + from .frf.preprocessing import ( + compute_epsilon, + compute_v_budget, + epsilon_aware_unroll, + eterm_sigma_a, + ) + + device = F_obs.device + dtype = F_obs.dtype + n_ops = int(sym_mats.shape[0]) + N = hkl_real.shape[0] + + # 1. ε(h) per reflection (needed for the ε-corrected obs normalisation). + if eps_factor is None: + eps_factor = compute_epsilon(hkl_real.to(torch.long), sym_mats).to(dtype) + eps_factor = eps_factor.to(device) + + # 2. Per-shell ε-corrected Wilson E_obs (Phaser E = F/sqrt(ε·Σ_N)). Dividing + # ε out of the obs is the obs-side analog of the unique-orbit calc dedup: + # both stop axial reflections (ε>1) from being over-weighted on + # high-symmetry spacegroups. + shell_idx = _equal_count_shell_idx(s_mag, n_shells) + E_obs = _normalize_to_e_epsilon(F_obs, shell_idx, n_shells, eps_factor) + + # 3. Identity-rotation calc reference → E-normalisation scale for F_calc + # (rotation-invariant: sphere permutation, shell sums preserved). + I_eye = torch.eye(3, dtype=torch.float32, device=device) + F_calc_ref = interpolator.evaluate( + I_eye, hkl_real, real_cell, return_amplitude=True, + ).to(dtype).squeeze(0) # (N,) + if scat_mode == "legacy": + # Per-shell unit-variance normalisation: forces _shell = 1 in + # EVERY shell, flattening F_calc's inter-shell amplitude shape. + calc_norm_per_h = _per_shell_sqrt_mean( + F_calc_ref, shell_idx, n_shells, + ).to(device) + elif scat_mode == "absolute": + # Single GLOBAL scale: preserves F_calc's inter-shell shape (how much the + # model actually scatters per resolution) instead of flattening it to 1. + # Phaser keeps E_calc physically scaled and carries the model's fraction + # of the cell in scatFactor = AtomScatRatio·SCATTERING/TOTAL_SCAT/NSYMP; + # for a search model that IS the full ASU (the benchmark case) scatFactor + # reduces to 1/n_ops, so the prefactor is unchanged and the only change + # here is dropping the per-shell flatten. + global_rms = F_calc_ref.pow(2).mean().clamp(min=1e-30).sqrt() + calc_norm_per_h = torch.full( + (N,), float(global_rms), dtype=dtype, device=device, + ) + else: + raise ValueError( + f"scat_mode={scat_mode!r}; expected 'legacy' or 'absolute'." + ) + + # Optional Wilson-B match (EnsemblePDB.cc:793-851), applied as a per-reflection + # Debye-Waller multiplier on F_calc. + if apply_wilson_b and wilson_b_value is None: + from .frf.preprocessing import fit_relative_wilson_b + wilson_b_value = fit_relative_wilson_b( + F_obs, F_calc_ref, s_mag, n_shells=n_shells, + ) + wilson_b_value = float(wilson_b_value or 0.0) + if apply_wilson_b and abs(wilson_b_value) > 1e-6: + dw = torch.exp(-wilson_b_value * (s_mag * s_mag) / 4.0).to(dtype).to(device) + else: + dw = None + + # 4. σ_A per reflection — FRF/Phaser Eterm (2π²/3 isotropic form), not the + # 2π² Luzzati form. Rotation-independent, no aligned model required. + if sigma_a is None: + if vrms_strategy == "oeffner": + if vrms_n_residues is None: + raise ValueError( + "vrms_strategy='oeffner' requires vrms_n_residues=." + ) + from .frf.preprocessing import oeffner_vrms + delta_vrms_A = oeffner_vrms(int(vrms_n_residues), float(vrms_identity)) + elif vrms_strategy == "fixed": + delta_vrms_A = 0.5 # legacy default + else: + raise ValueError( + f"vrms_strategy={vrms_strategy!r}; expected 'fixed' or 'oeffner'." + ) + sigma_a = eterm_sigma_a(s_mag, delta_vrms_A=delta_vrms_A).to(dtype).to(device) + if apply_bulk_solvent: + from .frf.preprocessing import bulk_solvent_factor + sol = bulk_solvent_factor( + s_mag, fsol=solvent_fsol, bsol=solvent_bsol, + ).to(dtype).to(device) + sigma_a = sigma_a * sol + sigma_a = sigma_a.to(device) + sigma_a2 = sigma_a * sigma_a + + # 5. V(h) — rotation-independent variance budget V = ε − σ_A² (n_mol=1). + V = compute_v_budget(eps_factor, sigma_a, n_mol=1) # (N,) + + # 6. Unique-orbit unroll: distinct mates only (Phaser duplicate-skip). Each + # ASU reflection appears n_ops/ε(h) times, NOT n_ops times. + unrolled_hkl, asu_idx = epsilon_aware_unroll(hkl_real, sym_mats) + unrolled_hkl = unrolled_hkl.to(torch.float64).to(device) + asu_idx = asu_idx.to(device) + + # 7. Broadcastable per-reflection tensors. + E_obs_b = E_obs.unsqueeze(0) # (1, N) + V_b = V.unsqueeze(0) # (1, N) + # eImove = ε(h)·σ_A²·(1/n_ops)·Σ_{distinct mates} |E_calc(R^T·S_k·h)|² + # (Phaser DataMR.cc:1397: thisEsqr *= repsn·scatFactor, scatFactor∝1/NSYMP). + eImove_prefac = (eps_factor * sigma_a2 / float(n_ops)).unsqueeze(0) # (1, N) + centric_b = centric.to(torch.bool).to(device).unsqueeze(0) + + # Per-mate normaliser + DW (rotation preserves |h| → same shell across the + # orbit, so the per-h scale broadcasts to every mate via asu_idx). + sqrt_mean_per_m = calc_norm_per_h[asu_idx] # (M,) + dw_per_m = dw[asu_idx] if dw is not None else None + + return _LLGContext( + interpolator=interpolator, real_cell=real_cell, + unrolled_hkl=unrolled_hkl, asu_idx=asu_idx, N=N, + E_obs_b=E_obs_b, V_b=V_b, eImove_prefac=eImove_prefac, + sqrt_mean_per_m=sqrt_mean_per_m, centric_b=centric_b, + dw_per_m=dw_per_m, dtype=dtype, batch_size=batch_size, + ) + + +def _llg_for_orientations( + ctx: _LLGContext, + alpha: torch.Tensor, + beta: torch.Tensor, + gamma: torch.Tensor, +) -> torch.Tensor: + """Per-orientation m_LETF1 LLG for a batch of Edmonds-ZYZ Euler angles. + + Returns a ``(n_orient,)`` tensor of LLG values. The calc orbit-sum is over + the deduped mates: evaluate ``|E_calc|²`` on ``ctx.unrolled_hkl`` then + ``scatter_add`` back per ASU reflection. Same Phaser logRel math as before. + """ + R_all = rotation_matrix_from_edmonds_euler_batch( + alpha.to(torch.float64), beta.to(torch.float64), gamma.to(torch.float64), + ).transpose(-1, -2).to(torch.float32) # (n_orient, 3, 3) + n_orient = R_all.shape[0] + sqrt_mean_b = ctx.sqrt_mean_per_m.unsqueeze(0) # (1, M) + dw_b = ctx.dw_per_m.unsqueeze(0) if ctx.dw_per_m is not None else None + + chunks: List[torch.Tensor] = [] + for start in range(0, n_orient, ctx.batch_size): + R_batch = R_all[start:start + ctx.batch_size] # (B, 3, 3) + F_calc_m = ctx.interpolator.evaluate( + R_batch, ctx.unrolled_hkl, ctx.real_cell, return_amplitude=True, + ).to(ctx.dtype) # (B, M) + if dw_b is not None: + F_calc_m = F_calc_m * dw_b + E_calc_m = F_calc_m / sqrt_mean_b # (B, M) + Esq_m = E_calc_m * E_calc_m # (B, M) + B = Esq_m.shape[0] + sum_per_h = torch.zeros( + B, ctx.N, dtype=Esq_m.dtype, device=Esq_m.device, + ) + idx = ctx.asu_idx.unsqueeze(0).expand(B, -1) # (B, M) + sum_per_h.scatter_add_(1, idx, Esq_m) # (B, N) + eImove = ctx.eImove_prefac * sum_per_h # (B, N) + sqrt_eImove = eImove.clamp(min=1e-30).sqrt() + ll_acen = phaser_log_rel_rice(ctx.E_obs_b, sqrt_eImove, ctx.V_b) + ll_cen = phaser_log_rel_woolfson(ctx.E_obs_b, sqrt_eImove, ctx.V_b) + ll = torch.where(ctx.centric_b, ll_cen, ll_acen) # (B, N) + chunks.append(ll.sum(dim=-1)) # (B,) + return torch.cat(chunks) + + +@dataclass +class _SimLLGContext: + """Context for the per-candidate-σ_A Sim-LLG surface (no orbit sum). + + Unlike :class:`_LLGContext` (fixed σ_A m_LETF1), this surface FITS σ_A per + shell for each orientation via :func:`llg_for_rotation_batch`. The fixed-σ_A + m_LETF1 surface is locally mis-peaked on high-sym/tNCS cases; the per-candidate + fit re-shapes it so the local maximum sits at the true orientation (the + property `sim_mlrf_rescore` already exhibits). Used as a ``llg_fn`` for + :func:`quadratic_llg_refine`. + """ + + interpolator: LattmanLoveInterpolator + real_cell: object + hkl: torch.Tensor # (N, 3) + F_obs: torch.Tensor # (N,) + shell_idx: torch.Tensor # (N,) int64 + n_shells: int + E_obs: torch.Tensor # (N,) + centric: torch.Tensor # (N,) bool + shell_weights: Optional[torch.Tensor] + n_D_grid: int + interp_var: Optional[torch.Tensor] + batch_size: int + + +def _build_sim_llg_context( + F_obs: torch.Tensor, + hkl_real: torch.Tensor, + s_mag: torch.Tensor, + centric: torch.Tensor, + interpolator: LattmanLoveInterpolator, + real_cell, + *, + n_shells: int = 10, + n_D_grid: int = 21, + batch_size: int = 64, + auto_variance_weights: bool = True, + interp_var: Optional[torch.Tensor] = None, +) -> _SimLLGContext: + """Build the rotation-independent context for the Sim-LLG surface.""" + shell_idx = _equal_count_shell_idx(s_mag, n_shells) + E_obs = _normalize_to_e(F_obs, shell_idx, n_shells) + shell_weights = None + if auto_variance_weights: + from .sh import compute_patterson_shell_variance + patt_obs = (E_obs.to(torch.float64) ** 2) - 1.0 + var_p = compute_patterson_shell_variance(patt_obs, shell_idx, P=n_shells) + w = 1.0 / var_p.sqrt() + w = w * (n_shells / w.sum().clamp(min=1e-30)) + shell_weights = w.to(F_obs.dtype) + return _SimLLGContext( + interpolator=interpolator, real_cell=real_cell, hkl=hkl_real, + F_obs=F_obs, shell_idx=shell_idx, n_shells=n_shells, E_obs=E_obs, + centric=centric.to(torch.bool), shell_weights=shell_weights, + n_D_grid=n_D_grid, interp_var=interp_var, batch_size=batch_size, + ) + + +def _sim_llg_for_orientations( + ctx: _SimLLGContext, + alpha: torch.Tensor, + beta: torch.Tensor, + gamma: torch.Tensor, +) -> torch.Tensor: + """Per-orientation Sim-LLG (per-candidate σ_A fit), drop-in ``llg_fn``.""" + R_all = rotation_matrix_from_edmonds_euler_batch( + alpha.to(torch.float64), beta.to(torch.float64), gamma.to(torch.float64), + ).transpose(-1, -2).to(torch.float32) + n_orient = R_all.shape[0] + chunks: List[torch.Tensor] = [] + for start in range(0, n_orient, ctx.batch_size): + R_batch = R_all[start:start + ctx.batch_size] + F_calc = ctx.interpolator.evaluate( + R_batch, ctx.hkl, ctx.real_cell, return_amplitude=True, + ).to(ctx.F_obs.dtype) # (B, N) + llg = llg_for_rotation_batch( + F_obs=ctx.F_obs, shell_idx=ctx.shell_idx, n_shells=ctx.n_shells, + E_obs=ctx.E_obs, centric=ctx.centric, F_calc=F_calc, + shell_weights=ctx.shell_weights, n_D_grid=ctx.n_D_grid, + interp_var=ctx.interp_var, + ) + chunks.append(llg) + return torch.cat(chunks) + + +def _euler_batch_from_matrices(R: torch.Tensor): + """(K,3,3) → three (K,) float64 Euler-angle tensors (Edmonds ZYZ). + + Loops the scalar :func:`edmonds_euler_from_rotation_matrix` (K is small — + the top-K refine set), returning tensors ready for + :func:`_llg_for_orientations`. + """ + a, b, g = [], [], [] + for k in range(R.shape[0]): + aa, bb, gg = edmonds_euler_from_rotation_matrix(R[k]) + a.append(aa) + b.append(bb) + g.append(gg) + return ( + torch.tensor(a, dtype=torch.float64), + torch.tensor(b, dtype=torch.float64), + torch.tensor(g, dtype=torch.float64), + ) + + +def quadratic_llg_refine( + peaks: List[RotationPeak], + ctx: _LLGContext, + *, + k_refine: int = 20, + step_deg: float = 1.5, + n_grid: int = 3, + iterations: int = 1, + max_move_deg: Optional[float] = None, + llg_fn: Optional[Callable[[torch.Tensor, torch.Tensor, torch.Tensor], torch.Tensor]] = None, + verbose: int = 0, +) -> List[RotationPeak]: + """Sub-grid refinement of the top-``k_refine`` peaks on the ML-LLG surface. + + For each peak, sample the LLG on a local **axis-angle** grid around the + orientation, fit a 3-D paraboloid in the tangent space, and step to the vertex + (a Newton step on the LLG). Axis-angle (not Euler α,β,γ) perturbation keeps the + local metric isotropic and avoids the β→0/π gimbal degeneracy where the FRF + returns many peaks. Guards (Hessian negative-definite + vertex inside the + sampled box) fall back to the best sampled grid point, so the refined peak can + never score below its grid value. Refined peaks are re-ranked by their (truly + re-evaluated) LLG; peaks beyond ``k_refine`` are appended unchanged. + + Reliable (sub-degree recovery from grid-resolution hits) on well-behaved + crystals; on high-symmetry / tNCS cases the m_LETF1 surface is locally + mis-peaked (~3–4° off truth) so refinement can WALK AWAY from a good hit — + use ``max_move_deg`` to bound that. + + Parameters + ---------- + peaks + Rescored candidate orientations (Edmonds ZYZ), best-first. + ctx + The :class:`_LLGContext` built for the same data (its + :func:`_llg_for_orientations` defines the surface being optimised). + k_refine + Number of leading peaks to refine. + step_deg + Half-width of the local axis-angle grid (degrees) and the per-iteration + capture radius. Default 1.5 ≈ grid_sampling/2 (the FRF grid half-step). + n_grid + Samples per tangent axis (3 → 27 orientations per peak). + iterations + Newton iterations; each re-centres and halves the grid half-width. + max_move_deg + Safety cap: if the refined orientation moves more than this (geodesic + degrees) from the input peak, keep the input peak instead. ``None`` + disables the cap. Protects against the mis-peaked-surface failure mode. + llg_fn + Surface to optimise: a callable ``(alpha,beta,gamma) -> (M,)`` LLG. If + ``None``, uses the m_LETF1 surface ``_llg_for_orientations(ctx, ...)``. + Pass a per-candidate-σ_A Sim surface (:func:`_sim_llg_for_orientations`) + when the fixed-σ_A m_LETF1 surface is locally mis-peaked (high-sym/tNCS). + """ + if not peaks: + return [] + if llg_fn is None: + def llg_fn(a, b, g): + return _llg_for_orientations(ctx, a, b, g) + k = min(k_refine, len(peaks)) + head = peaks[:k] + tail = peaks[k:] + + # R0 for the head peaks (Edmonds ZYZ, un-transposed — the convention + # `_llg_for_orientations` consumes after its own transpose). + a0 = torch.tensor([p.alpha for p in head], dtype=torch.float64) + b0 = torch.tensor([p.beta for p in head], dtype=torch.float64) + g0 = torch.tensor([p.gamma for p in head], dtype=torch.float64) + R0 = rotation_matrix_from_edmonds_euler_batch(a0, b0, g0) # (k, 3, 3) + R0_orig = R0.clone() # for the move cap + + radius = math.radians(step_deg) + for _ in range(max(1, iterations)): + # Local tangent grid (G, 3), shared across peaks; rebuilt per iteration + # so a 2nd pass zooms in. + lin = torch.linspace(-radius, radius, n_grid, dtype=torch.float64) + gx, gy, gz = torch.meshgrid(lin, lin, lin, indexing="ij") + omegas = torch.stack( + [gx.reshape(-1), gy.reshape(-1), gz.reshape(-1)], dim=-1, + ) # (G, 3) + G = omegas.shape[0] + x, y, z = omegas[:, 0], omegas[:, 1], omegas[:, 2] + ones = torch.ones_like(x) + # Design matrix Φ (G,10): [1, x,y,z, x²,y²,z², xy,xz,yz]. + Phi = torch.stack( + [ones, x, y, z, x * x, y * y, z * z, x * y, x * z, y * z], dim=-1, + ) # (G, 10) + + # Grid orientations: R = rodrigues(ω) @ R0, for every (peak, grid point). + Rloc = axis_angle_to_matrix(omegas) # (G, 3, 3) + R_grid = torch.einsum("gij,kjl->kgil", Rloc, R0) # (k, G, 3, 3) + a_t, b_t, g_t = _euler_batch_from_matrices( + R_grid.reshape(k * G, 3, 3), + ) + llg = llg_fn(a_t, b_t, g_t).reshape(k, G) # (k, G) + + # Batched quadratic fit via ridge-stabilised normal equations: + # θ = (ΦᵀΦ + λI)⁻¹ Φᵀ llg. ΦᵀΦ is shared; only the RHS varies per peak. + PtP = Phi.t() @ Phi # (10, 10) + PtP = PtP + 1e-9 * torch.eye(10, dtype=PtP.dtype) + rhs = torch.einsum("gd,kg->kd", Phi, llg.to(torch.float64)) # (k, 10) + theta = torch.linalg.solve( + PtP.unsqueeze(0).expand(k, -1, -1), rhs.unsqueeze(-1), + ).squeeze(-1) # (k, 10) + + # Gradient b and Hessian H of the paraboloid (in tangent coords). + bvec = theta[:, 1:4] # (k, 3) + H = torch.zeros(k, 3, 3, dtype=torch.float64) + H[:, 0, 0] = 2.0 * theta[:, 4] + H[:, 1, 1] = 2.0 * theta[:, 5] + H[:, 2, 2] = 2.0 * theta[:, 6] + H[:, 0, 1] = H[:, 1, 0] = theta[:, 7] + H[:, 0, 2] = H[:, 2, 0] = theta[:, 8] + H[:, 1, 2] = H[:, 2, 1] = theta[:, 9] + + best_grid = llg.argmax(dim=1) # (k,) + new_R0 = R0.clone() + n_accept = 0 + for kk in range(k): + accept = False + try: + eig = torch.linalg.eigvalsh(H[kk]) + if bool((eig < 0).all()): # genuine maximum + xstar = torch.linalg.solve(H[kk], -bvec[kk]) # (3,) + if torch.isfinite(xstar).all() and float(xstar.norm()) <= radius: + new_R0[kk] = axis_angle_to_matrix(xstar) @ R0[kk] + accept = True + n_accept += 1 + except Exception: + accept = False + if not accept: + new_R0[kk] = R_grid[kk, best_grid[kk]] + R0 = new_R0 + radius = radius / 2.0 + if verbose > 1: + print( + f" quadratic_llg_refine: {n_accept}/{k} vertices accepted " + f"(rest fell back to grid max)", + flush=True, + ) + + # Safety cap: revert any peak whose total move from the input exceeds + # ``max_move_deg`` (geodesic). On a locally mis-peaked surface (high-sym / + # tNCS) the refinement walks toward a spurious LLG max ~3–4° away; capping + # the move means a good hit can never be degraded by more than the cap. + if max_move_deg is not None: + cos_cap = math.cos(math.radians(max_move_deg)) + n_revert = 0 + for kk in range(k): + trace = torch.einsum("ij,ij->", R0[kk], R0_orig[kk]) + cos_move = float(((trace - 1.0) * 0.5).clamp(-1.0, 1.0)) + if cos_move < cos_cap: # moved further than the cap + R0[kk] = R0_orig[kk] + n_revert += 1 + if verbose > 1 and n_revert: + print( + f" quadratic_llg_refine: reverted {n_revert}/{k} peaks that " + f"moved > {max_move_deg}° (mis-peaked-surface guard).", + flush=True, + ) + + # Final TRUE LLG at the refined orientations (never trust the paraboloid). + af, bf, gf = _euler_batch_from_matrices(R0) + llg_final = llg_fn(af, bf, gf) # (k,) + if llg_final.numel() > 1: + std_t = llg_final.std().clamp(min=1e-30) + sig = (llg_final - llg_final.mean()) / std_t + else: + sig = torch.zeros_like(llg_final) + + af_l, bf_l, gf_l = af.tolist(), bf.tolist(), gf.tolist() + llg_l, sig_l = llg_final.tolist(), sig.tolist() + refined = [ + RotationPeak( + alpha=af_l[i], beta=bf_l[i], gamma=gf_l[i], + score=llg_l[i], sigma=sig_l[i], + ) + for i in range(k) + ] + refined.sort(key=lambda p: p.score, reverse=True) + return refined + tail + + def m_letf1_rescore( peaks: List[RotationPeak], F_obs: torch.Tensor, @@ -544,19 +1088,22 @@ def m_letf1_rescore( vrms_identity: float = 1.0, apply_wilson_b: bool = False, wilson_b_value: Optional[float] = None, # if None and apply_wilson_b=True, fitted from data + scat_mode: str = "legacy", # "legacy" (per-shell calc norm) | "absolute" (global) ) -> List[RotationPeak]: """Phaser-faithful ``m_LETF1`` rescore (DataMR.cc:1326-1429). + Thin wrapper around :func:`_build_llg_context` + :func:`_llg_for_orientations`. + Upgrades over :func:`sim_mlrf_rescore`: - 1. **NSYMP symmetry sum on calc** — for each obs reflection ``h``, the + 1. **Unique-orbit symmetry sum on calc** — for each obs reflection ``h``, the expected moving-model intensity is - ``eImove(h) = Σ_isym σ_A²(s) · |F_calc(R^T · S_isym · h)|²`` - summed over the ``NSYMP`` spacegroup rotation operators - (DataMR.cc:1371-1404). Implemented vectorised: pre-compute the orbit - ``hkl_unroll`` of shape ``(N, n_ops, 3)``, flatten to - ``(N·n_ops, 3)``, evaluate the LL interpolator once per orientation, - view back, square + sum over the symop dim. + ``eImove(h) = ε(h)·σ_A²·(1/n_ops)·Σ_{distinct mates} |E_calc(R^T·S_k·h)|²`` + summed over the **distinct** orbit mates only (Phaser's + ``if(!duplicate(isym))``, DataMR.cc:1371-1404), via + :func:`torchref.alignment.frf.preprocessing.epsilon_aware_unroll` + + ``scatter_add``. Summing all ``n_ops`` raw mates over-weights axial + reflections by ε(h) and orientation-blinds high-symmetry spacegroups. 2. **Per-reflection variance budget** ``V(h) = ε(h) − σ_A²(s)·n_mol`` from :func:`torchref.alignment.frf.preprocessing.compute_v_budget` @@ -599,152 +1146,23 @@ def m_letf1_rescore( head = peaks[:n_refine] tail = peaks[n_refine:] - device = F_obs.device - dtype = F_obs.dtype - n_ops = int(sym_mats.shape[0]) - N = hkl_real.shape[0] - - # 1. Per-shell binning + Wilson-normalised E_obs (E in Phaser notation). - shell_idx = _equal_count_shell_idx(s_mag, n_shells) - E_obs = _normalize_to_e(F_obs, shell_idx, n_shells) - - # 2. ε(h) per reflection. - if eps_factor is None: - from .frf.preprocessing import compute_epsilon - eps_factor = compute_epsilon(hkl_real.to(torch.long), sym_mats).to(dtype) - eps_factor = eps_factor.to(device) - - # 3. σ_A per reflection — Luzzati formula from a coordinate-error parameter - # ``delta_vrms_A`` (default 0.5 Å, matching the validated FRF v19 config). - # This is the Phaser-faithful choice: Phaser's σ_A per shell is pre-fit - # via a Wilson/Luzzati-style formula that depends only on (s, ΔVRMS) and - # does NOT require an aligned model. Data-fit alternatives - # (`fit_sigma_a_per_shell`) need a meaningful obs-calc alignment, which - # we don't have a priori — at any misaligned reference (identity or - # even the top FRF peak when truth is buried) the fit returns ~0 and - # the LL becomes orientation-blind (the 4BX9 rank-481 / 2DQ6 rank-235 - # failure modes in v22/v23). - # - # Always need the calc shell-scale for E-normalisation; use identity - # (rotation-invariant — sphere permutation, shell sums preserved). - I_eye = torch.eye(3, dtype=torch.float32, device=device) - F_calc_ref = interpolator.evaluate( - I_eye, hkl_real, real_cell, return_amplitude=True, - ).to(dtype).squeeze(0) # (N,) - sqrt_mean_F2_calc_per_h = _per_shell_sqrt_mean(F_calc_ref, shell_idx, n_shells).to(device) - - # Optional Wilson-B match (EnsemblePDB.cc:793-851). Compute once from the - # identity-rotation F_calc reference (rotation-invariant shell statistic), - # apply as Debye-Waller multiplier `exp(-B·s²/4)` to F_calc inside the batch. - if apply_wilson_b and wilson_b_value is None: - from .frf.preprocessing import fit_relative_wilson_b - wilson_b_value = fit_relative_wilson_b( - F_obs, F_calc_ref, s_mag, n_shells=n_shells, - ) - wilson_b_value = float(wilson_b_value or 0.0) - if apply_wilson_b and abs(wilson_b_value) > 1e-6: - dw = torch.exp(-wilson_b_value * (s_mag * s_mag) / 4.0).to(dtype).to(device) - else: - dw = None # skip the elementwise mul if a no-op - - if sigma_a is None: - # σ_A: Phaser-faithful Luzzati with optional Oeffner vrms + bulk solvent. - if vrms_strategy == "oeffner": - if vrms_n_residues is None: - raise ValueError( - "vrms_strategy='oeffner' requires vrms_n_residues=." - ) - from .frf.preprocessing import oeffner_vrms - delta_vrms_A = oeffner_vrms(int(vrms_n_residues), float(vrms_identity)) - elif vrms_strategy == "fixed": - delta_vrms_A = 0.5 # legacy default - else: - raise ValueError( - f"vrms_strategy={vrms_strategy!r}; expected 'fixed' or 'oeffner'." - ) - sigma_a = compute_sigma_a_luzzati(s_mag, delta_vrms_A=delta_vrms_A).to(dtype).to(device) - if apply_bulk_solvent: - from .frf.preprocessing import bulk_solvent_factor - sol = bulk_solvent_factor( - s_mag, fsol=solvent_fsol, bsol=solvent_bsol, - ).to(dtype).to(device) - sigma_a = sigma_a * sol - sigma_a = sigma_a.to(device) - sigma_a2 = sigma_a * sigma_a - - # 4. V(h) — rotation-independent variance budget. Phaser-faithful per - # DataMR.cc:1342-1345: `thisV = scatFactor · NSYMP · DFAC² · σ_A²`. - # With `scatFactor = 1/NSYMP` (single ensemble, fracMove=1) this collapses - # to `thisV = σ_A²`, so V = ε − σ_A² — **independent of NSYMP**. Using - # `n_mol=n_ops` (as the initial implementation did) drove V negative on - # high-symmetry spacegroups (4BX9 NSYMP=8, V_clamp blew up the LL). - from .frf.preprocessing import compute_v_budget - V = compute_v_budget(eps_factor, sigma_a, n_mol=1) # (N,) + ctx = _build_llg_context( + F_obs, hkl_real, s_mag, centric, interpolator, real_cell, sym_mats, + n_shells=n_shells, batch_size=batch_size, sigma_a=sigma_a, + eps_factor=eps_factor, apply_bulk_solvent=apply_bulk_solvent, + solvent_fsol=solvent_fsol, solvent_bsol=solvent_bsol, + vrms_strategy=vrms_strategy, vrms_n_residues=vrms_n_residues, + vrms_identity=vrms_identity, apply_wilson_b=apply_wilson_b, + wilson_b_value=wilson_b_value, scat_mode=scat_mode, + ) - # 5. Orbit hkl pre-compute (rotation-independent; reused per batch). - sym_mats_f = sym_mats.to(torch.float64).to(device) - hkl_f = hkl_real.to(torch.float64).to(device) - # (N, n_ops, 3) — for each obs h, all n_ops sym-equivalents in hkl space. - hkl_unroll = torch.einsum("kij,nj->nki", sym_mats_f, hkl_f) - hkl_flat = hkl_unroll.reshape(-1, 3) # (N·n_ops, 3) - - # 6. Build candidate rotation matrices. Same convention as sim_mlrf_rescore: - # transpose the Edmonds Euler matrix because the peak encodes "rotation - # applied to model coords"; we need the inverse for evaluating - # F_calc(R^T · h). alpha_t = torch.tensor([p.alpha for p in head], dtype=torch.float64) beta_t = torch.tensor([p.beta for p in head], dtype=torch.float64) gamma_t = torch.tensor([p.gamma for p in head], dtype=torch.float64) - R_all = rotation_matrix_from_edmonds_euler_batch( - alpha_t, beta_t, gamma_t, - ).transpose(-1, -2).to(torch.float32) # (M, 3, 3) - M = R_all.shape[0] - - # 7. Batched score: eImove via NSYMP-summed calc, LL via Phaser logRel. - E_obs_b = E_obs.unsqueeze(0) # (1, N) — broadcasts over batch dim - V_b = V.unsqueeze(0) # (1, N) - # eImove pre-factor per Phaser DataMR.cc:1397: `thisEsqr *= repsn * scatFactor` - # → `eImove = ε(h) · σ_A² · (1/n_ops) · Σ_k |E_calc(S_k h)|²`. We collapse the - # `1/n_ops` into the pre-factor so the batch loop just does the raw sum. - eImove_prefac = (eps_factor * sigma_a2 / float(n_ops)).unsqueeze(0) # (1, N) - centric_b = centric.to(torch.bool).unsqueeze(0) - - # Normaliser to convert rotated |F_calc| → |E_calc| (per-shell sqrt mean - # from the identity-rotation reference; broadcasts over symops via the - # last unsqueeze since rotation preserves |h| → same shell for all symops). - sqrt_mean_b = sqrt_mean_F2_calc_per_h.unsqueeze(0).unsqueeze(-1) # (1, N, 1) + llgs_t = _llg_for_orientations(ctx, alpha_t, beta_t, gamma_t) + if verbose > 1: + print(f" m_LETF1 scored {len(head)} peaks", flush=True) - llg_chunks: List[torch.Tensor] = [] - for start in range(0, M, batch_size): - stop = min(start + batch_size, M) - R_batch = R_all[start:stop] # (B, 3, 3) - F_calc_flat = interpolator.evaluate( - R_batch, hkl_flat, real_cell, return_amplitude=True, - ).to(dtype) # (B, N·n_ops) - F_calc = F_calc_flat.view(-1, N, n_ops) # (B, N, n_ops) - # Optional Wilson-B Debye-Waller multiplier (per-reflection, broadcasts - # over batch + symops). Applied PRE-normalisation so the per-shell sqrt - # mean (computed from un-DW'd identity F_calc_ref) stays the right scale - # for normalisation; the DW shifts E_calc relative to that scale, which - # is exactly what Wilson-B matching is supposed to do. - if dw is not None: - F_calc = F_calc * dw.unsqueeze(0).unsqueeze(-1) - # Normalise to E_calc (same Wilson scale as E_obs); Rice/Woolfson LL - # only makes sense with both sides on the same per-shell scale. - E_calc = F_calc / sqrt_mean_b - # eImove(h) = ε(h) · σ_A²(s) · (1/n_ops) · Σ_isym |E_calc(R^T·S_isym·h)|² - # (Phaser DataMR.cc:1371-1404; scatFactor = 1/NSYMP folds the symop sum - # into a mean, ε(h) is the multiplicity factor `repsn`). - eImove = eImove_prefac * (E_calc * E_calc).sum(dim=-1) # (B, N) - sqrt_eImove = eImove.clamp(min=1e-30).sqrt() # (B, N) - ll_acen = phaser_log_rel_rice(E_obs_b, sqrt_eImove, V_b) # (B, N) - ll_cen = phaser_log_rel_woolfson(E_obs_b, sqrt_eImove, V_b) # (B, N) - ll = torch.where(centric_b, ll_cen, ll_acen) # (B, N) - llg_chunks.append(ll.sum(dim=-1)) # (B,) - if verbose > 1: - print(f" m_LETF1 batch {start}-{stop}/{M}", flush=True) - - llgs_t = torch.cat(llg_chunks) mean_t = llgs_t.mean() std_t = llgs_t.std().clamp(min=1e-30) sigmas_t = (llgs_t - mean_t) / std_t From 2ec9fd692dc5ba334dd2e7486d6eba26b2be6676 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 15 Jun 2026 13:14:46 +0200 Subject: [PATCH 007/250] ALignment works 70% of the time, main blocker is the rotation rescoring function --- .../alignment/run_random_pdb_fit.py | 4 +- torchref/experimental/alignment/__init__.py | 14 +- torchref/experimental/alignment/align.py | 780 +--------- .../alignment/frf/preprocessing.py | 11 +- torchref/experimental/alignment/pipeline.py | 1254 +++++++++++------ 5 files changed, 872 insertions(+), 1191 deletions(-) diff --git a/tests/integration/alignment/run_random_pdb_fit.py b/tests/integration/alignment/run_random_pdb_fit.py index 8f07c339..d7d7b208 100644 --- a/tests/integration/alignment/run_random_pdb_fit.py +++ b/tests/integration/alignment/run_random_pdb_fit.py @@ -101,7 +101,8 @@ def run(pdb_key: str, seed: int, verbose: int = 1, use_m_symmetry_filter: bool = False, use_lerf1_intensity: bool = False, use_fitted_delta_vrms: bool = False, - use_even_l_only: bool = False) -> dict: + use_even_l_only: bool = False, + rescore_engine: str = "m_letf1") -> dict: pdb_path, mtz_path = PAIRS[pdb_key] print(f"\n=== {pdb_key}: {pdb_path.name} + {mtz_path.name} ===", flush=True) @@ -175,6 +176,7 @@ def _scale_and_r(m: ModelFT) -> tuple[float, float]: use_lerf1_intensity=use_lerf1_intensity, use_fitted_delta_vrms=use_fitted_delta_vrms, use_even_l_only=use_even_l_only, + rescore_engine=rescore_engine, ) fit_time = time.time() - t1 print(f" fit_to_data took {fit_time:.1f}s", flush=True) diff --git a/torchref/experimental/alignment/__init__.py b/torchref/experimental/alignment/__init__.py index bcab98b9..cebeda16 100644 --- a/torchref/experimental/alignment/__init__.py +++ b/torchref/experimental/alignment/__init__.py @@ -5,12 +5,16 @@ 1. Fast Rotation Function (``frf.phaser_rotation_search`` / ``frf.FastRotationFunction``) — Phaser-faithful Bessel-radial × SH - expansion, stable Wigner-d, dense P1-box calc. -2. Translation Search (``translation.fft_translation_search_torch``). + expansion, stable Wigner-d, dense P1-box calc — then ML rescoring + (``ml_rotation.m_letf1_rescore``) to rank candidate orientations. +2. Fast Translation Function (``translation.amplitude_translation_search`` + + ``local_translation_refine``) — run per rotation candidate. 3. Rigid Body Refinement (``rigid_body.RigidBodyRefinement``) — LBFGS on - rotation and translation parameters with an ML target. -4. Unified Pipeline (``pipeline.MolecularReplacementPipeline``) — end-to-end - workflow with early-stopping. + rotation and translation (and optional B-factors) with an ML target. +4. Canonical Pipeline (``pipeline.MolecularReplacementPipeline``) — the + multi-candidate FRF → FTF → post-refine tree with early-stopping; the + implementation that ``align.align_model_to_data`` / + ``ModelFT.fit_to_data`` delegate to. Example — full MR pipeline -------------------------- diff --git a/torchref/experimental/alignment/align.py b/torchref/experimental/alignment/align.py index e8950274..a22a8a17 100644 --- a/torchref/experimental/alignment/align.py +++ b/torchref/experimental/alignment/align.py @@ -1,60 +1,41 @@ """ -End-to-end molecular replacement alignment entry point. - -This module owns the full MR pipeline: - - LERF1 ball-search rotation - → Sim MLRF rescore - → amplitude-correlation translation search - → local translation refine (analytical-scale R) - → dense rotation sampling (ML-LLG, multi-pass) - → LBFGS rigid-body polish - → final solvent-aware Scaler refit (user-facing R-work) - -The single public function `align_model_to_data` is what `ModelFT.fit_to_data` -delegates to; the latter is a thin wrapper. Keep alignment logic in this file -to keep `torchref/model/model_ft.py` focused on the FFT structure-factor model. +Molecular replacement: data-prep / FRF stage helpers + the public entry point. + +This module hosts the heavy, reusable stage helpers — Lattman-Love / anisotropy +data prep (`_prepare_frf_inputs`), the Phaser-faithful rotation search +(`_run_frf_separate_rotation`), the solvent-aware R-work +(`_external_rwork`), the direct-SF translation evaluator +(`_DirectModelEvaluator`), the Rodrigues helper (`_rodrigues`) and the stage +timer (`_StageTimer`) — that are shared by the rotation-ranking benchmarks and +by the orchestrator. + +`align_model_to_data` is the public entry point that `ModelFT.fit_to_data` +delegates to; it in turn delegates the FRF → FTF(per-candidate) → post-refine +control flow to +:class:`torchref.experimental.alignment.pipeline.MolecularReplacementPipeline`, +returning that pipeline's single best `ModelFT`. """ from __future__ import annotations -import math import time from contextlib import contextmanager from dataclasses import dataclass from typing import Optional, TYPE_CHECKING -import numpy as np import torch -from .frf.rotation_utils import ( - edmonds_euler_from_rotation_matrix, - rotation_matrix_from_edmonds_euler, -) -from .frf.types import RotationPeak -from .lattman_love import LattmanLoveInterpolator, estimate_interp_var -from .ml_rotation import compute_sigma_a_luzzati, m_letf1_rescore, sim_mlrf_rescore -from .rigid_body import RigidBodyRefinement +from .lattman_love import LattmanLoveInterpolator from .sh import ( apply_overall_anisotropy, assign_shells, - compute_patterson_shell_variance, equal_count_shell_edges, fit_overall_anisotropy, - get_high_order_axis, ) -from .translation import ( - TranslationPeak, - amplitude_translation_search, - llg_translation_rescore, - local_translation_refine, - precompute_G_for_rotation, -) -from .ml_rotation import fit_sigma_a_per_shell if TYPE_CHECKING: - from ..io.datasets.reflection_data import ReflectionData - from ..model.model_ft import ModelFT + from ...io.datasets.reflection_data import ReflectionData + from ...model.model_ft import ModelFT # --------------------------------------------------------------------------- @@ -158,7 +139,7 @@ def _external_rwork(model: "ModelFT", data: "ReflectionData") -> float: candidates correctly but isn't the user-facing R-work). We compute the proper Scaler-fit R-work once per finalist. """ - from ..scaling import Scaler + from ...scaling import Scaler # Build the Scaler on the model's device so that its anisotropy U # tensor and per-bin scales land alongside `data.hkl`/`model(hkl)` — @@ -662,689 +643,54 @@ def align_model_to_data( See `ModelFT.fit_to_data` for full kwarg semantics — this function is the canonical implementation; `fit_to_data` is a thin wrapper. """ - from ..scaling import Scaler # noqa: F401 (imported by _external_rwork) - from ..symmetry import SpaceGroup - if not model.initialized: raise RuntimeError( "Cannot fit an uninitialized ModelFT. Load PDB data first." ) - timer = _StageTimer(enabled=verbose >= 2) - - timer.start("0_data_prep") - timer.start("1_anisotropy_fit") - timer.start("2_ll_build") - frf = _prepare_frf_inputs( - model, data, - d_min=d_min, d_max=d_max, n_shells=n_shells, - ll_padding_factor=ll_padding_factor, ll_max_res_A=ll_max_res_A, + # `MolecularReplacementPipeline` is the implementation of record. This + # function preserves the historical kwarg surface (so `ModelFT.fit_to_data` + # and the benchmark scripts keep working unchanged) and returns the single + # best `ModelFT`; drive the pipeline directly to get the ranked candidate + # list. Imported lazily to avoid an import cycle — `pipeline` imports the + # stage helpers (`_prepare_frf_inputs`, `_run_frf_separate_rotation`, + # `_external_rwork`, `_DirectModelEvaluator`, `_rodrigues`, `_StageTimer`) + # from this module. + from .pipeline import MolecularReplacementPipeline + + pipeline = MolecularReplacementPipeline( + data, model, + device=model.xyz().device, verbose=verbose, + d_min=d_min, d_max=d_max, n_shells=n_shells, + ll_max_res_A=ll_max_res_A, ll_padding_factor=ll_padding_factor, + n_rotation_peaks=n_rotation_peaks, n_ml_refine=n_ml_refine, + frf_lmax_cap=frf_lmax_cap, frf_dense_pad=frf_dense_pad, + rescore_engine=rescore_engine, rescore_scat_mode=rescore_scat_mode, + auto_variance_weights=auto_variance_weights, + use_interp_var=use_interp_var, + subpeak_refine=subpeak_refine, subpeak_refine_k=subpeak_refine_k, + subpeak_refine_step_deg=subpeak_refine_step_deg, + subpeak_refine_iters=subpeak_refine_iters, + subpeak_refine_max_move_deg=subpeak_refine_max_move_deg, + n_rotation_candidates=n_rotation_candidates, + n_translation_peaks=n_translation_peaks, + n_translation_candidates=n_translation_candidates, + translation_grid_steps=translation_grid_steps, + use_llg_tf=use_llg_tf, + do_joint_refine=do_joint_refine, + joint_refine_max_res_A=joint_refine_max_res_A, + joint_refine_expected_rot_error=joint_refine_expected_rot_error, + refine_b=refine_b, + sigma_rot_deg=sigma_rot_deg, sigma_trans_ang=sigma_trans_ang, + sigma_b=sigma_b, + L=L, + use_sigma_a_frf=use_sigma_a_frf, frf_delta_vrms_A=frf_delta_vrms_A, + frf_weight_combine=frf_weight_combine, + use_m_symmetry_filter=use_m_symmetry_filter, + use_lerf1_intensity=use_lerf1_intensity, + use_fitted_delta_vrms=use_fitted_delta_vrms, + use_even_l_only=use_even_l_only, ) - timer.stop("0_data_prep") - timer.stop("1_anisotropy_fit") - timer.stop("2_ll_build") - if verbose > 0: - U_aniso = frf.U_aniso - print( - f"fit_to_data: overall U-aniso diag (Ų) = " - f"({U_aniso[0, 0].item():+.2f}, {U_aniso[1, 1].item():+.2f}, " - f"{U_aniso[2, 2].item():+.2f})", - flush=True, - ) - print( - f"fit_to_data: built Lattman-Love interpolator " - f"(box={ll_padding_factor}·diam, max_res={ll_max_res_A} Å)", - flush=True, - ) - - # Unpack into the local names the rest of this function uses. - device = frf.device - F_obs = frf.F_obs - hkl = frf.hkl - s_vec = frf.s_vec - s_mag = frf.s_mag - centric = frf.centric - ll = frf.ll - U_aniso = frf.U_aniso - s_vec_for_search = frf.s_vec_for_search - patt_obs = frf.patt_obs - patt_calc = frf.patt_calc - - # --- Stage 1: fast Patterson ball-search --- - # F3: Fit ΔVRMS from the model's mean B-factor (runs FIRST so E3/F2 - # downstream use the fitted value). ΔVRMS² = / (8π²) converts - # Debye-Waller B to a 1-D RMS coordinate displacement. - effective_delta_vrms = frf_delta_vrms_A - if use_fitted_delta_vrms: - with torch.no_grad(): - _, adp_iso, _, _, _ = model.get_iso() - b_mean = float(adp_iso.mean().item()) if adp_iso.numel() > 0 else 0.0 - effective_delta_vrms = max( - math.sqrt(max(b_mean, 1e-6) / (8 * math.pi ** 2)), 0.1, - ) - if verbose > 0: - print( - f"fit_to_data: ΔVRMS fitted from = {b_mean:.2f} Ų → " - f"{effective_delta_vrms:.3f} Å (was {frf_delta_vrms_A} Å).", - flush=True, - ) - - # Optional Phaser-style σA pre-weighting of the SH input (E3). Off by - # default — when on, replaces `auto_variance_weights`. The two can be - # combined explicitly via `frf_weight_combine="sigma_a_x_variance"`. - rotsearch_weights: Optional[torch.Tensor] = None - rotsearch_auto_var = auto_variance_weights - if use_sigma_a_frf: - shell_edges, _ = equal_count_shell_edges(frf.s_mag_sym, n_shells) - shell_mid = 0.5 * (shell_edges[:-1] + shell_edges[1:]) # (P,) - sigma_a_shell = compute_sigma_a_luzzati( - shell_mid, delta_vrms_A=effective_delta_vrms, - ) - w = sigma_a_shell ** 2 - if frf_weight_combine == "sigma_a_x_variance": - shell_idx_sym = assign_shells(frf.s_mag_sym, shell_edges) - var_shell = compute_patterson_shell_variance( - patt_obs.to(torch.float64), shell_idx_sym, P=n_shells, - ) - w = w / var_shell.sqrt().clamp(min=1e-30) - elif frf_weight_combine != "sigma_a_only": - raise ValueError( - f"frf_weight_combine={frf_weight_combine!r}; " - "expected 'sigma_a_only' or 'sigma_a_x_variance'." - ) - # Per-shell sum-to-P normalisation (matches the FRF's internal weighting). - w = w * (n_shells / w.sum().clamp(min=1e-30)) - rotsearch_weights = w.to(patt_obs.dtype) - rotsearch_auto_var = False - if verbose > 0: - print( - f"fit_to_data: σA-weighted FRF (ΔVRMS={effective_delta_vrms}Å, " - f"combine={frf_weight_combine}, w[0]={rotsearch_weights[0]:.3f}, " - f"w[-1]={rotsearch_weights[-1]:.3f}).", - flush=True, - ) - - # F2: LERF1 likelihood intensity (Phaser DataMR.cc:947–951). Replace - # `E² − 1` on the OBSERVED side with the FRF likelihood intensity - # intensity = cweight · (E² − 1) · DFAC² - # where cweight ∈ {1, 2} for centric/acentric and DFAC is a - # per-reflection Luzzati factor proxied here as the per-reflection - # σA(s) (same Luzzati formula as Eterm but evaluated on the actual - # per-reflection s_mag, not the per-shell mean). - if use_lerf1_intensity: - n_ops_sg = int(data.spacegroup.matrices.shape[0]) - cweight_per = torch.where( - centric, torch.ones_like(F_obs), 2.0 * torch.ones_like(F_obs), - ) - cweight_sym = cweight_per.unsqueeze(0).expand(n_ops_sg, -1).reshape(-1) - dfac_sym = compute_sigma_a_luzzati( - frf.s_mag_sym, delta_vrms_A=effective_delta_vrms, - ).to(patt_obs.dtype) - patt_obs = patt_obs * cweight_sym.to(patt_obs.dtype) * (dfac_sym ** 2) - if verbose > 0: - print( - f"fit_to_data: LERF1 intensity ON (mean(cweight·DFAC²) = " - f"{(cweight_sym * dfac_sym ** 2).mean().item():.3f}).", - flush=True, - ) - - # F1: m-symmetry filter (Phaser DataMR.cc:1019 / 1117). Compute ZSYMM - # from the spacegroup. Off-by-default for backwards compat. - rotsearch_zsymm = 1 - if use_m_symmetry_filter: - sg_mats_cpu = data.spacegroup.matrices.to(torch.float64).cpu() - axis, zsymm = get_high_order_axis(sg_mats_cpu) - if axis != 2 and verbose > 0: - print( - f"fit_to_data: WARNING — highest-order axis is {axis} (x/y), " - f"but m-symmetry filter assumes z. Applying with potentially " - f"reduced effect; axis permutation not yet implemented.", - flush=True, - ) - rotsearch_zsymm = int(zsymm) - if verbose > 0: - print( - f"fit_to_data: m-symmetry filter ON, ZSYMM={rotsearch_zsymm} " - f"(axis={axis}).", - flush=True, - ) - - timer.start("3_rotation_search") - # Phaser-faithful FRF (dense calc + auto_lmax cap + obs-unroll + no_grad); - # solves the high-symmetry cases. Single engine post-consolidation. - if verbose > 0: - print( - f"fit_to_data: frf_separate rotation search " - f"(dense calc + auto_lmax cap={frf_lmax_cap}, " - f"n_peaks={n_rotation_peaks})…", - flush=True, - ) - peaks = _run_frf_separate_rotation( - model, data, frf, - lmax_cap=frf_lmax_cap, dense_pad=frf_dense_pad, - n_peaks=n_rotation_peaks, verbose=verbose, - ) - timer.stop("3_rotation_search") - - # --- Stage 2: Sim-MLRF rescore (per-shell σA fit per candidate) --- - timer.start("4_sim_mlrf_rescore") - if verbose > 0: - print( - f"fit_to_data: ML rescoring top " - f"{min(len(peaks), n_ml_refine)} peaks…", - flush=True, - ) - interp_var_main: Optional[torch.Tensor] = None - if use_interp_var: - # Per-reflection interpolation variance (Phaser totvar_search analogue). - # Inflates the Rice variance budget so a noisy true peak isn't - # demoted below a noise-free wrong peak by the rescore. - rescore_n_shells = max(n_shells // 2, 8) - rescore_edges, _ = equal_count_shell_edges(s_mag, rescore_n_shells) - rescore_shell_idx = assign_shells(s_mag, rescore_edges) - interp_var_main = estimate_interp_var( - ll, hkl, data.cell, rescore_shell_idx, rescore_n_shells, - ).to(F_obs.dtype) - if verbose > 0: - print( - f"fit_to_data: interp_var enabled (mean={interp_var_main.mean().item():.3f}, " - f"max={interp_var_main.max().item():.3f}).", - flush=True, - ) - - if rescore_engine not in ("m_letf1", "sim"): - raise ValueError( - f"rescore_engine={rescore_engine!r}; expected 'm_letf1' (default) or 'sim'." - ) - if rescore_engine == "m_letf1": - # Phaser-faithful: NSYMP calc sum + V(h) budget + Rice/Woolfson logRel. - # Cross-rotation case: no fixed model, so totvar_known=0 and the - # variance budget reduces to ε(h) - σ_A²(s)·n_mol in E-space. - rescored = m_letf1_rescore( - peaks, F_obs, hkl, s_mag, centric, ll, data.cell, - data.spacegroup.matrices.to(torch.float64).to(device), - n_shells=max(n_shells // 2, 8), - n_refine=min(len(peaks), n_ml_refine), - batch_size=50, - verbose=verbose, - scat_mode=rescore_scat_mode, - ) - if subpeak_refine: - # Sharpen the top candidate orientations on the (now-corrected) - # ML-LLG surface before the translation search — the FTF is sensitive - # to orientation error and the FRF grid leaves up to ~half a grid step - # (~1°). Rebuild the LLG context once (cheap: one identity interpolator - # eval) and take a quadratic tangent-space Newton step per top-K peak. - from .ml_rotation import _build_llg_context, quadratic_llg_refine - timer.start("4b_subpeak_refine") - _ctx = _build_llg_context( - F_obs, hkl, s_mag, centric, ll, data.cell, - data.spacegroup.matrices.to(torch.float64).to(device), - n_shells=max(n_shells // 2, 8), batch_size=50, - scat_mode=rescore_scat_mode, - ) - _k = subpeak_refine_k if subpeak_refine_k > 0 else n_rotation_candidates - _k = min(_k, len(rescored)) - rescored = quadratic_llg_refine( - rescored, _ctx, - k_refine=_k, - step_deg=subpeak_refine_step_deg, - iterations=subpeak_refine_iters, - max_move_deg=subpeak_refine_max_move_deg, - verbose=verbose, - ) - timer.stop("4b_subpeak_refine") - if verbose > 0: - print( - f"fit_to_data: sub-peak refined top {_k} orientations " - f"on the ML-LLG surface (step={subpeak_refine_step_deg}°).", - flush=True, - ) - else: # rescore_engine == "sim" — legacy Sim/Rice approximation - rescored = sim_mlrf_rescore( - peaks, F_obs, hkl, s_mag, centric, ll, data.cell, - n_shells=max(n_shells // 2, 8), - n_refine=min(len(peaks), n_ml_refine), - batch_size=50, - verbose=verbose, - auto_variance_weights=auto_variance_weights, - interp_var=interp_var_main, - ) - timer.stop("4_sim_mlrf_rescore") - if not rescored: - raise RuntimeError("Rotation search produced no peaks.") - - n_rot = min(n_rotation_candidates if do_translation else 1, len(rescored)) - - # The Patterson rotation function has a centrosymmetric ambiguity, but - # `sim_mlrf_rescore` ranks Patterson-equivalents adjacent. `n_rot ≥ 3` - # is enough without explicitly multiplying by spacegroup rotations. - if verbose > 0 and n_rot > 1: - print( - f"fit_to_data: trying top {n_rot} rotation candidates " - f"(Patterson-equivalents covered by LLG ranking).", - flush=True, - ) - - def _candidate(k): - peak = rescored[k] - R_rec = rotation_matrix_from_edmonds_euler( - peak.alpha, peak.beta, peak.gamma, - ) - R_app = R_rec.T.contiguous() - rot = model.rotate( - R_app.to(device=model.device, dtype=model.dtype_float), - ) - rot.last_alignment_rotation = R_rec - return rot, R_rec, R_app, peak - - if not do_translation: - rotated, R_recovered, _, top = _candidate(0) - if verbose > 0: - print( - f"fit_to_data: top peak LLG = {top.score:.2f} " - f"(σ_Z = {top.sigma:.2f}); applying R⁻¹ to coords.", - flush=True, - ) - if verbose >= 2: - print("\n" + timer.summary(), flush=True) - return rotated - - # --- Stage 3: translation search + analytical-R local refine --- - hkl_full = data.hkl - F_obs_full = data.F - if hasattr(data, "get_valid_mask"): - tmask = data.get_valid_mask() - else: - tmask = torch.ones( - F_obs_full.shape[0], dtype=torch.bool, device=F_obs_full.device, - ) - # Index on CPU then move slices to `device` — same pattern as the - # early `keep`-mask section; the validity mask lives on the data's - # device, while downstream consumers run on `model.device`. - F_obs_amp = F_obs_full[tmask].abs().to(torch.float64).to(device) - hkl_keep = hkl_full[tmask].to(device) - - global_best = None - for k_rot in range(n_rot): - rotated_k, R_recovered_k, _, peak_k = _candidate(k_rot) - if verbose > 0: - print( - f"\nfit_to_data: rot{k_rot} " - f"(LLG={peak_k.score:.2f}, σ_Z={peak_k.sigma:.2f})", - flush=True, - ) - - if str(rotated_k.spacegroup) != str(data.spacegroup): - rotated_k.spacegroup = data.spacegroup - - rotated_p1 = rotated_k.copy() - rotated_p1.spacegroup = SpaceGroup("P 1") - evaluator = _DirectModelEvaluator(rotated_p1) - - # Pre-compute per-sym F_asu contributions once per rotation; reused - # by both coarse TF and each local refine. - timer.start("5_precompute_G") - G_pre, h_R_pre = precompute_G_for_rotation( - evaluator, torch.eye(3, dtype=torch.float64), - hkl_keep, data.spacegroup, data.cell, - ) - timer.stop("5_precompute_G") - - timer.start("6_amplitude_TF") - _, _, t_peaks = amplitude_translation_search( - F_obs=F_obs_amp, interpolator=evaluator, - R_rotation=torch.eye(3, dtype=torch.float64), - hkl=hkl_keep, - spacegroup=data.spacegroup, real_cell=data.cell, - grid_steps=translation_grid_steps, - n_peaks=n_translation_peaks, - cluster_radius=0.05, - precomputed_G=G_pre, precomputed_h_R=h_R_pre, - ) - timer.stop("6_amplitude_TF") - if not t_peaks: - if verbose > 0: - print(" no translation peaks; skipping", flush=True) - continue - - # Phase B: re-rank the cheap-correlation peaks by Rice/Woolfson LLG - # using a shared per-shell σA fitted at the top correlation peak. - # Mirrors Phaser's FTF — the correlation pre-filter is fast but its - # ranking is degraded for partial models; the LLG ranks consistently - # with the rotation rescore. - if use_llg_tf: - timer.start("6b_llg_tf_rescore") - rec_basis_keep = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) - s_mag_keep_tf = (hkl_keep.to(torch.float64) @ rec_basis_keep).norm(dim=-1) - tf_n_shells = max(n_shells // 2, 8) - tf_edges, _ = equal_count_shell_edges(s_mag_keep_tf, tf_n_shells) - tf_shell_idx = assign_shells(s_mag_keep_tf, tf_edges) - centric_keep_tf = ( - data.centric[tmask].to(torch.bool).to(device) - if hasattr(data, "centric") - else torch.zeros_like(F_obs_amp, dtype=torch.bool) - ) - - # E_obs normalised on the validity-masked set. - cnt_tf = torch.bincount( - tf_shell_idx, minlength=tf_n_shells, - ).to(torch.float64) - sum_F2 = torch.zeros(tf_n_shells, dtype=torch.float64, device=device) - sum_F2.scatter_add_(0, tf_shell_idx, F_obs_amp * F_obs_amp) - mean_F2 = (sum_F2 / cnt_tf.clamp(min=1.0)).clamp(min=1e-30) - E_obs_tf = F_obs_amp / mean_F2.sqrt().index_select(0, tf_shell_idx) - - # Compute |F_calc| at the top correlation peak's translation, use - # that to fit the per-shell σA. One-shot, no per-candidate refit. - t_top_np = t_peaks[0].translation - t_top_t = torch.as_tensor(t_top_np, dtype=torch.float64, device=device) - S_eff, N_eff = G_pre.shape - phase_top = torch.exp( - 2j * torch.pi * torch.einsum( - "ind,d->in", - h_R_pre.to(torch.float64), t_top_t, - ).to(G_pre.dtype), - ) - Fc_top = (G_pre * phase_top).sum(dim=0).abs().to(torch.float64) - # Per-shell E normalise - sum_Fc2 = torch.zeros(tf_n_shells, dtype=torch.float64, device=device) - sum_Fc2.scatter_add_(0, tf_shell_idx, Fc_top * Fc_top) - mean_Fc2 = (sum_Fc2 / cnt_tf.clamp(min=1.0)).clamp(min=1e-30) - E_calc_top = Fc_top / mean_Fc2.sqrt().index_select(0, tf_shell_idx) - sigma_a_tf = fit_sigma_a_per_shell( - E_obs_tf, E_calc_top, centric_keep_tf, - tf_shell_idx, tf_n_shells, n_grid=81, - ) - - # interp_var is only meaningful when the F_calc comes from a - # trilinear interpolator. The TF stage uses a direct-SF evaluator - # (`_DirectModelEvaluator`) with no interpolation noise, so the - # Phaser totvar_search analogue does not apply here. - interp_var_tf: Optional[torch.Tensor] = None - - t_cands = torch.as_tensor( - np.stack([p.translation for p in t_peaks]), - dtype=torch.float64, device=device, - ) - llg_tf = llg_translation_rescore( - F_obs=F_obs_amp, hkl=hkl_keep, centric=centric_keep_tf, - shell_idx=tf_shell_idx, n_shells=tf_n_shells, - G=G_pre, h_R=h_R_pre, t_candidates=t_cands, - sigma_a=sigma_a_tf, interp_var=interp_var_tf, - ) - timer.stop("6b_llg_tf_rescore") - # Re-rank t_peaks by LLG (descending). Update the score to carry - # LLG so downstream picks correctly. - llg_list = llg_tf.detach().cpu().tolist() - corr_list = [p.score for p in t_peaks] # original FFT-correlation scores - order = sorted( - range(len(t_peaks)), key=lambda i: llg_list[i], reverse=True, - ) - # Record where the correlation top-1 ended up after LLG re-rank. - corr_top1_new_rank = order.index(0) - t_peaks = [ - TranslationPeak( - translation=t_peaks[i].translation, - score=float(llg_list[i]), - sigma=float(llg_list[i]), - ) - for i in order - ] - if verbose > 0: - tt = tuple(round(float(x), 3) for x in t_peaks[0].translation.tolist()) - llg_top1 = llg_list[order[0]] - corr_at_llg_top1 = corr_list[order[0]] - print( - f" LLG-TF rescore: top t={tt} LLG={t_peaks[0].score:.2f} " - f"(corr at LLG-top1: {corr_at_llg_top1:.4f}; " - f"corr-top1 demoted to LLG-rank {corr_top1_new_rank}/{len(order)})", - flush=True, - ) - - if verbose > 0: - tt = tuple(round(float(x), 3) for x in t_peaks[0].translation.tolist()) - print( - f" top translation t={tt} corr={t_peaks[0].score:.4f}", - flush=True, - ) - - if not do_joint_refine: - t_top = torch.as_tensor( - t_peaks[0].translation, - dtype=model.dtype_float, device=rotated_k.device, - ) - translated = rotated_k.translate(t_top, fractional=True) - translated.last_alignment_rotation = R_recovered_k - translated.last_alignment_translation = t_top - return translated - - for k_t, tp in enumerate(t_peaks[:n_translation_candidates]): - t_init = torch.as_tensor(tp.translation, dtype=torch.float64) - timer.start("7_local_TF_refine") - # Single-pass local refine (was 2). Pass-2 zoomed to ~0.0017 - # fractional resolution; the downstream LBFGS rigid-body polish - # refines to gradient-tolerance anyway, so pass 2 was just - # paying ~half the local-TF cost for a precision that gets - # overridden a few steps later. - t_refined, r_analytic = local_translation_refine( - F_obs=F_obs_amp, interpolator=evaluator, - R_rotation=torch.eye(3, dtype=torch.float64), - hkl=hkl_keep, - spacegroup=data.spacegroup, real_cell=data.cell, - t_init=t_init, radius=0.06, grid_steps=13, - n_refinement_passes=1, - precomputed_G=G_pre, precomputed_h_R=h_R_pre, - ) - timer.stop("7_local_TF_refine") - if verbose > 0: - print( - f" rot{k_rot} trans{k_t}: " - f"R(analytic)={r_analytic:.4f}, " - f"t={[round(float(x), 3) for x in t_refined.tolist()]}", - flush=True, - ) - if global_best is None or r_analytic < global_best[0]: - global_best = (r_analytic, rotated_k, R_recovered_k, t_refined) - - if global_best is None: - raise RuntimeError("Translation + joint refine produced no candidates.") - r_analytic_best, rot_best, R_recovered_best, t_refined_best = global_best - refined = rot_best.translate( - t_refined_best.to(model.dtype_float), fractional=True, - ) - - # --- Stage 4: dense rotation sampling at the found translation --- - if do_joint_refine: - timer.start("8_dense_R_ll_build") - refined_p1 = refined.copy() - refined_p1.spacegroup = SpaceGroup("P 1") - ll_refine = LattmanLoveInterpolator( - refined_p1, padding_factor=ll_padding_factor, - max_res_A=ll_max_res_A, verbose=0, - ) - timer.stop("8_dense_R_ll_build") - - centric_keep = ( - data.centric[tmask].to(torch.bool).to(device) if hasattr(data, "centric") - else torch.zeros(hkl_keep.shape[0], dtype=torch.bool, device=device) - ) - rec_basis_keep = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) - s_mag_keep = (hkl_keep.to(torch.float64) @ rec_basis_keep).norm(dim=-1) - - # Dense-R sampling: 2-pass zoom on the σ_A-fitted ML LLG which has - # FWHM ~2–3° (sharper than the Patterson rotation function's ~10° - # because of log-likelihood curvature and σ_A up-weighting). Pass 1 - # at 9³ × ±5.7° (1.43° spacing) gives 1–2 samples per ML-FWHM and - # locates the basin; pass 2 at 5³ × ±1.43° (0.71° spacing) zooms in. - # Single-pass or larger spacing regresses R-work; finer than this - # is just sampled by the downstream LBFGS polish anyway. - n_per_axis_pass = [9, 5] - zoom_factor = 4.0 - radii = [ - float(joint_refine_expected_rot_error), - float(joint_refine_expected_rot_error) / zoom_factor, - ] - R_accumulated = torch.eye(3, dtype=torch.float64) - - # Hoisted out of the per-pass loop: ll_refine is fixed across passes, - # so interp_var only needs to be estimated once. - interp_var_dense: Optional[torch.Tensor] = None - if use_interp_var: - dense_n_shells = max(n_shells // 2, 8) - dense_edges, _ = equal_count_shell_edges(s_mag_keep, dense_n_shells) - dense_shell_idx = assign_shells(s_mag_keep, dense_edges) - interp_var_dense = estimate_interp_var( - ll_refine, hkl_keep, data.cell, - dense_shell_idx, dense_n_shells, - ).to(F_obs_amp.dtype) - - for pass_idx, max_perturb_rad in enumerate(radii): - n_per_axis = n_per_axis_pass[pass_idx] - coords_r = torch.linspace( - -max_perturb_rad, max_perturb_rad, n_per_axis, - dtype=torch.float64, - ) - wx, wy, wz = torch.meshgrid( - coords_r, coords_r, coords_r, indexing="ij", - ) - omegas = torch.stack( - [wx.flatten(), wy.flatten(), wz.flatten()], dim=-1, - ) - # Batched Rodrigues + matmul: one (B, 3, 3) build instead of B - # per-omega calls. Previously _rodrigues had `.item()` × 3 per - # call, so dense_R cost was dominated by Python dispatch. - R_perturbs = _rodrigues(omegas) # (B, 3, 3) - R_cand_full = R_perturbs @ R_accumulated # (B, 3, 3) - cand_peaks = [] - for R_c in R_cand_full: - a, b, g = edmonds_euler_from_rotation_matrix(R_c) - cand_peaks.append(RotationPeak( - alpha=a, beta=b, gamma=g, score=0.0, sigma=0.0, - )) - if verbose > 0: - print( - f"\nfit_to_data: dense R pass {pass_idx + 1} " - f"({n_per_axis}³={omegas.shape[0]} perturbations, " - f"±{math.degrees(max_perturb_rad):.2f}°)…", - flush=True, - ) - # batch_size adapted to N_hkl × D_grid (~41) to keep the - # llg_for_rotation_batch inner tensors bounded. - rescore_batch = max( - 4, min(100, 1_000_000 // max(hkl_keep.shape[0], 1)), - ) - timer.start("9_dense_R_rescore") - # n_D_grid=11 (vs 41 default): we only need relative LLG - # ranking across a tight rotation neighbourhood; the σA - # optimum shifts negligibly. 4× fewer Bessel evals. - if rescore_engine == "m_letf1": - rescored_refine = m_letf1_rescore( - cand_peaks, F_obs_amp, hkl_keep, s_mag_keep, centric_keep, - ll_refine, data.cell, - data.spacegroup.matrices.to(torch.float64).to(device), - n_shells=max(n_shells // 2, 8), - n_refine=len(cand_peaks), batch_size=rescore_batch, - verbose=0, - ) - else: - rescored_refine = sim_mlrf_rescore( - cand_peaks, F_obs_amp, hkl_keep, s_mag_keep, centric_keep, - ll_refine, data.cell, - n_shells=max(n_shells // 2, 8), - n_refine=len(cand_peaks), batch_size=rescore_batch, - verbose=0, n_D_grid=11, - interp_var=interp_var_dense, - ) - timer.stop("9_dense_R_rescore") - top = rescored_refine[0] - best_idx = next( - i for i, p in enumerate(cand_peaks) - if p.alpha == top.alpha and p.beta == top.beta - and p.gamma == top.gamma - ) - R_accumulated = R_cand_full[best_idx] - if verbose > 0: - print( - f" pass {pass_idx + 1} best LLG={top.score:.2f}, " - f"|ω|={omegas[best_idx].norm().item() * 180 / math.pi:.3f}°", - flush=True, - ) - - # Trust the LLG. The dense-R rescore picks the perturbation with - # the highest LLG (scale-invariant ML target); LLG monotonicity - # implies a non-degraded R-work. Two solvent-aware Scaler refits - # used to live here as a defensive gate — at ~25 s each on 1DAW - # they ran twice the wall time of the entire rescore loop they - # were checking. Cheaper to trust the score. - refined = refined.rotate( - R_accumulated.T.to(model.dtype_float).contiguous(), - ) - - # --- Stage 5: joint LBFGS polish on (R, t) --- - if do_joint_refine: - timer.start("11_lbfgs_polish") - rb = RigidBodyRefinement( - refined, data, - initial_translation=torch.zeros( - 3, dtype=torch.float32, device=refined.device, - ), - expected_rotational_error=joint_refine_expected_rot_error, - max_res=joint_refine_max_res_A, - device=refined.device, - verbose=max(0, verbose - 1), - refine_b=refine_b, - sigma_rot_deg=sigma_rot_deg, - sigma_trans_ang=sigma_trans_ang, - sigma_b=sigma_b, - ) - rb_result = rb.refine() - with torch.no_grad(): - R_polish = rb.get_rotation_matrix().detach() - t_polish = rb.translation_frac.detach() - polished = refined.rotate(R_polish.to(model.dtype_float)) - polished = polished.translate( - t_polish.to(model.dtype_float), fractional=True, - ) - timer.stop("11_lbfgs_polish") - # Compare initial vs final R-work from the LBFGS's *own* internal - # scaler (same instance, both numbers no-solvent — apples to - # apples). Saves two solvent-aware Scaler refits (~25 s each - # on 1DAW) that previously gated this decision. - if rb_result.final_r_factor <= rb_result.initial_r_factor: - refined = polished - if verbose > 0: - print( - f"\nfit_to_data: joint polish " - f"{rb_result.initial_r_factor:.4f} → " - f"{rb_result.final_r_factor:.4f} (no-solvent R)", - flush=True, - ) - elif verbose > 0: - print( - f"\nfit_to_data: joint polish kept original " - f"({rb_result.initial_r_factor:.4f} ≤ " - f"{rb_result.final_r_factor:.4f} no-solvent R)", - flush=True, - ) - - # Single solvent-aware Scaler refit at the very end on the winner — - # gives the user-facing R-work without paying the cost on every - # intermediate gate. - timer.start("12_final_scaler") - rwork_final = _external_rwork(refined, data) - timer.stop("12_final_scaler") - - refined.last_alignment_rotation = R_recovered_best - refined.last_alignment_translation = t_refined_best - refined.last_alignment_rfactor = rwork_final - if verbose > 0: - print( - f"fit_to_data: best analytical R={r_analytic_best:.4f}, " - f"final Scaler-fit R-work={rwork_final:.4f}", - flush=True, - ) - if verbose >= 2: - print("\n" + timer.summary(), flush=True) - return refined + solutions = pipeline.run(do_translation=do_translation) + return solutions[0].model diff --git a/torchref/experimental/alignment/frf/preprocessing.py b/torchref/experimental/alignment/frf/preprocessing.py index c5b35617..665d830e 100644 --- a/torchref/experimental/alignment/frf/preprocessing.py +++ b/torchref/experimental/alignment/frf/preprocessing.py @@ -129,7 +129,16 @@ def epsilon_aware_unroll( N, n_ops = hkl_int.shape[0], sym_mats.shape[0] # Orbits: (N, n_ops, 3) — S_k applied to each h (row-vector convention, # matching the existing `einsum("kij,nj->nki", ...)` unroll site). - orbits = torch.einsum("kij,nj->nki", sym_mats, hkl_int) + # Integer einsum dispatches to baddbmm, which CUDA does not implement for + # Long; compute in float64 (exact for symop 0/±1 × small Miller indices) + # and round back so the GPU path works. + orbits = ( + torch.einsum( + "kij,nj->nki", sym_mats.to(torch.float64), hkl_int.to(torch.float64), + ) + .round() + .to(torch.long) + ) # Pack (h, k, l) into a single int64 key for per-row dedup. base = 2 * int(orbits.abs().max().item()) + 1 key = (orbits[:, :, 0] * base + orbits[:, :, 1]) * base + orbits[:, :, 2] diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 4257465d..385524f8 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -1,87 +1,105 @@ """ -Molecular replacement pipeline integrating rotation, translation, and refinement. - -This module provides a unified pipeline for molecular replacement that chains: -1. Fast rotation function (ball transform) -2. FFT-based translation search -3. Clash filtering -4. Rigid body refinement - -The pipeline supports early stopping when a good solution is found. +Molecular replacement pipeline: the single canonical MR orchestrator. + +Implements the classic Phaser-style molecular-replacement tree: + +1. **Fast Rotation Function (FRF)** — Phaser-faithful Bessel-radial × SH + expansion (dense P1-box calc + auto_lmax), then ML rescoring + (``m_letf1_rescore`` / ``sim_mlrf_rescore``) to rank candidate orientations. +2. **Fast Translation Function (FTF)** — for *each* of the top-N rotation + candidates, an amplitude-correlation translation search (optionally + re-ranked by a Rice/Woolfson LLG) followed by an analytical-R local refine. +3. **Post-refinement** — optional dense rotation re-sampling on the ML-LLG + surface, then an LBFGS rigid-body polish on (R, t) (with optional B-factor + co-refinement and Gaussian restraints). + +Each rotation candidate is carried through translation + refinement +independently; the candidates are ranked by their refined R-factor and the +best is returned (a Phaser-style multi-candidate tree, with early-stopping once +a candidate beats ``rfactor_converged``). The user-facing solvent-aware R-work +is computed once, on the winner. + +``align_model_to_data`` (and therefore ``ModelFT.fit_to_data``) delegates to +this class — it is the implementation of record. The heavy crystallographic +stage helpers live in :mod:`torchref.experimental.alignment.align`, +:mod:`~torchref.experimental.alignment.translation` and +:mod:`~torchref.experimental.alignment.ml_rotation`; this module owns the +control flow that wires them together. """ +from __future__ import annotations + +import math from dataclasses import dataclass from typing import List, Optional, Tuple, TYPE_CHECKING + import numpy as np import torch -from torchref.config import get_default_device, get_float_dtype +from torchref.config import get_default_device +from torchref.utils.device_mixin import DeviceMixin -from .frf.rotation_utils import rotation_matrix_from_edmonds_euler +from ...symmetry import SpaceGroup +from .align import ( + _DirectModelEvaluator, + _StageTimer, + _external_rwork, + _prepare_frf_inputs, + _rodrigues, + _run_frf_separate_rotation, +) +from .frf.rotation_utils import ( + edmonds_euler_from_rotation_matrix, + rotation_matrix_from_edmonds_euler, +) from .frf.types import RotationPeak -from .translation import fft_translation_search_torch, TranslationPeak -from .rigid_body import RigidBodyRefinement, RigidBodyResult -from .clashscore import ClashScoreCalculator, AtomSampler -from torchref.utils.device_mixin import DeviceMixin +from .lattman_love import LattmanLoveInterpolator, estimate_interp_var +from .ml_rotation import ( + fit_sigma_a_per_shell, + m_letf1_rescore, + sim_mlrf_rescore, +) +from .rigid_body import RigidBodyRefinement +from .sh import assign_shells, equal_count_shell_edges +from .translation import ( + TranslationPeak, + amplitude_translation_search, + llg_translation_rescore, + local_translation_refine, + precompute_G_for_rotation, +) + +if TYPE_CHECKING: + from torchref.io.datasets import ReflectionData + from torchref.model import ModelFT def rotation_matrix_from_euler_zyz(alpha, beta, gamma) -> np.ndarray: - """ - Build R = R_z(α) R_y(β) R_z(γ) (Edmonds active ZYZ) as a NumPy 3×3 matrix. + """Build R = R_z(α) R_y(β) R_z(γ) (Edmonds active ZYZ) as a NumPy 3×3 matrix. Compatibility wrapper around `rotation_matrix_from_edmonds_euler`. """ R = rotation_matrix_from_edmonds_euler(float(alpha), float(beta), float(gamma)) return R.detach().cpu().numpy() -if TYPE_CHECKING: - from torchref.model import ModelFT - from torchref.io.datasets import ReflectionData - def rotation_angular_distance(R1: np.ndarray, R2: np.ndarray) -> float: - """ - Compute angular distance between two rotation matrices in degrees. + """Angular distance between two rotation matrices in degrees. - The angular distance is the angle of the rotation R2 @ R1.T. - - Parameters - ---------- - R1, R2 : np.ndarray - 3x3 rotation matrices. - - Returns - ------- - float - Angular distance in degrees. + The angular distance is the angle of the rotation ``R2 @ R1.T``. """ R_diff = R2 @ R1.T - # Clamp trace to valid range for arccos - trace = np.trace(R_diff) - trace = np.clip(trace, -1.0, 3.0) - angle_rad = np.arccos((trace - 1.0) / 2.0) - return np.degrees(angle_rad) + trace = np.clip(np.trace(R_diff), -1.0, 3.0) + return np.degrees(np.arccos((trace - 1.0) / 2.0)) def euler_angular_distance( euler1: Tuple[float, float, float], euler2: Tuple[float, float, float], ) -> float: - """ - Compute angular distance between two ZYZ Euler angle sets. - - Parameters - ---------- - euler1, euler2 : tuple - (alpha, beta, gamma) Euler angles in radians. - - Returns - ------- - float - Angular distance in degrees. - """ - R1 = rotation_matrix_from_euler_zyz(euler1[0], euler1[1], euler1[2]) - R2 = rotation_matrix_from_euler_zyz(euler2[0], euler2[1], euler2[2]) + """Angular distance between two ZYZ Euler angle sets (degrees).""" + R1 = rotation_matrix_from_euler_zyz(*euler1) + R2 = rotation_matrix_from_euler_zyz(*euler2) return rotation_angular_distance(R1, R2) @@ -90,534 +108,836 @@ def cluster_rotation_peaks( threshold_deg: float = 6.0, symmetry_matrices: Optional[np.ndarray] = None, ) -> list: - """ - Cluster rotation peaks by angular distance. + """Cluster rotation peaks by angular distance. - Peaks within threshold_deg of each other are considered the same solution. - Only the highest-scoring peak from each cluster is kept. + Peaks within ``threshold_deg`` of each other are considered the same + solution; only the highest-scoring peak from each cluster is kept. Not on + the default pipeline path (the ML rescore already ranks Patterson- + equivalents adjacently); retained for callers that want explicit + de-duplication. Parameters ---------- peaks : list - List of rotation peaks as tuples (alpha, beta, gamma, score, sigma). + Rotation peaks as tuples ``(alpha, beta, gamma, score, sigma)``. threshold_deg : float - Angular distance threshold for clustering in degrees. + Angular distance threshold for clustering (degrees). symmetry_matrices : np.ndarray, optional - Point group symmetry matrices (N, 3, 3) to check symmetry equivalents. - - Returns - ------- - list - Clustered peaks with one representative per cluster. + Point-group symmetry matrices (N, 3, 3) to check symmetry equivalents. """ if not peaks: return [] - # Sort by sigma (descending) so we keep highest-scoring peaks sorted_peaks = sorted(peaks, key=lambda p: p[4], reverse=True) - clustered = [] - used_rotations = [] # Store rotation matrices of accepted peaks - + used_rotations = [] for peak in sorted_peaks: alpha, beta, gamma, score, sigma = peak R = rotation_matrix_from_euler_zyz(alpha, beta, gamma) - - # Check if this rotation is too close to any already accepted rotation is_new = True for R_used in used_rotations: - dist = rotation_angular_distance(R, R_used) - if dist < threshold_deg: + if rotation_angular_distance(R, R_used) < threshold_deg: is_new = False break - - # Also check symmetry equivalents if provided - if symmetry_matrices is not None and is_new: + if symmetry_matrices is not None: for sym_op in symmetry_matrices: - R_sym = sym_op @ R - dist_sym = rotation_angular_distance(R_sym, R_used) - if dist_sym < threshold_deg: + if rotation_angular_distance(sym_op @ R, R_used) < threshold_deg: is_new = False break if not is_new: break - if is_new: clustered.append(peak) used_rotations.append(R) - return clustered @dataclass class MRSolution: - """ - Molecular replacement solution. + """A molecular-replacement placement. Attributes ---------- rotation : np.ndarray - ZYZ Euler angles (radians), shape (3,). - translation : np.ndarray - Fractional coordinates, shape (3,). + Recovered orientation as a 3×3 rotation matrix (``R_recovered`` — the + rotation that maps the *search-model* frame onto the *crystal* frame). + translation : np.ndarray or None + Fractional translation applied after rotation, shape (3,). ``None`` for + a rotation-only solution (``do_translation=False``). rotation_score : float - FRF sigma (Z-score). + ML-LLG score of the rotation candidate (from the rescore). translation_score : float - Translation function correlation. - clash_score : float - Steric clash score (lower is better). + Analytical-R of the best translation for this candidate (lower better). r_factor : float - R-factor after refinement. - refined_rotation : np.ndarray, optional - Refined Euler angles (radians). - refined_translation : np.ndarray, optional - Refined fractional translation. + Ranking key. During the candidate loop this is the rigid-body's own + (no-solvent) R-work; for the returned winner it is replaced by the + solvent-aware Scaler R-work. + model : ModelFT + The rotated (+translated +refined) model for this candidate. + clash_score : float, optional + Steric clash score, only populated when ``clash_filter`` is enabled. """ + rotation: np.ndarray - translation: np.ndarray + translation: Optional[np.ndarray] rotation_score: float translation_score: float - clash_score: float r_factor: float - refined_rotation: Optional[np.ndarray] = None - refined_translation: Optional[np.ndarray] = None + model: "ModelFT" + clash_score: Optional[float] = None class MolecularReplacementPipeline(DeviceMixin): - """ - Unified MR pipeline: Rotation -> Translation -> Rigid Body Refinement. + """Canonical MR pipeline: FRF → FTF (per candidate) → post-refine. - This pipeline integrates the fast rotation function (ball transform), - FFT-based translation search, clash filtering, and rigid body refinement - into a single workflow with early stopping. + Parameters mirror :func:`align_model_to_data` (which delegates here), so a + caller can either use ``fit_to_data`` for the common case or drive this + class directly for finer control / access to the ranked candidate list. Parameters ---------- data : ReflectionData Observed reflection data. model : ModelFT - Search model with atomic coordinates. + Initialised search model. device : torch.device, optional - Computation device. Default is CPU. + Compute device (defaults to the model's device). verbose : int - Verbosity level (0=silent, 1=summary, 2=detailed). + 0 silent, 1 summary, ≥2 adds a per-stage wall-clock table. Examples -------- :: from torchref.experimental.alignment import MolecularReplacementPipeline - from torchref.model import ModelFT - from torchref.io.datasets import ReflectionData - - data = ReflectionData().load_mtz('observed.mtz') - model = ModelFT().load_pdb('search_model.pdb') - pipeline = MolecularReplacementPipeline(data, model) - solutions = pipeline.run(n_rotation_peaks=50, min_tries=3, max_tries=10) - print(f'Best R: {solutions[0].r_factor:.3f}') + + pipe = MolecularReplacementPipeline(data, model) + solutions = pipe.run() + print(f"best R-work: {solutions[0].r_factor:.3f}") """ def __init__( self, data: "ReflectionData", model: "ModelFT", + *, device: Optional[torch.device] = None, - verbose: int = 1, + verbose: int = 0, + # --- data prep / FRF --- + d_min: float = 4.0, + d_max: float = 15.0, + n_shells: int = 20, + ll_max_res_A: float = 3.0, + ll_padding_factor: float = 2.0, + n_rotation_peaks: int = 500, + n_ml_refine: int = 20, + frf_lmax_cap: int = 48, + frf_dense_pad: float = 2.0, + # --- rescore --- + rescore_engine: str = "m_letf1", + rescore_scat_mode: str = "legacy", + auto_variance_weights: bool = True, + use_interp_var: bool = False, + subpeak_refine: bool = False, + subpeak_refine_k: int = -1, + subpeak_refine_step_deg: float = 1.5, + subpeak_refine_iters: int = 1, + subpeak_refine_max_move_deg: Optional[float] = 1.5, + # --- candidate tree --- + n_rotation_candidates: int = 15, + n_translation_peaks: int = 20, + n_translation_candidates: int = 3, + translation_grid_steps: int = 16, + use_llg_tf: bool = False, + # --- post-refine --- + do_joint_refine: bool = True, + dense_rotation_refine: bool = True, + joint_refine_max_res_A: float = 4.0, + joint_refine_expected_rot_error: float = 0.1, + refine_b: bool = False, + sigma_rot_deg: float = 0.0, + sigma_trans_ang: float = 0.0, + sigma_b: float = 0.0, + # --- early stop --- + min_tries: int = 3, + max_tries: Optional[int] = None, + rfactor_converged: float = 0.45, + # --- vestigial FRF knobs (superseded by _run_frf_separate_rotation's + # frf_use_* defaults post-consolidation; accepted for API stability) --- + L: int = 48, + use_sigma_a_frf: bool = False, + frf_delta_vrms_A: float = 1.0, + frf_weight_combine: str = "sigma_a_only", + use_m_symmetry_filter: bool = False, + use_lerf1_intensity: bool = False, + use_fitted_delta_vrms: bool = False, + use_even_l_only: bool = False, ): self.data = data self.model = model self.device = device or get_default_device() self.verbose = verbose - # Lazy caches - self._clash_calc = None - self._e_obs = None - self._e_calc = None - self._s_vectors = None - self._mask = None + self.d_min = d_min + self.d_max = d_max + self.n_shells = n_shells + self.ll_max_res_A = ll_max_res_A + self.ll_padding_factor = ll_padding_factor + self.n_rotation_peaks = n_rotation_peaks + self.n_ml_refine = n_ml_refine + self.frf_lmax_cap = frf_lmax_cap + self.frf_dense_pad = frf_dense_pad + + self.rescore_engine = rescore_engine + self.rescore_scat_mode = rescore_scat_mode + self.auto_variance_weights = auto_variance_weights + self.use_interp_var = use_interp_var + self.subpeak_refine = subpeak_refine + self.subpeak_refine_k = subpeak_refine_k + self.subpeak_refine_step_deg = subpeak_refine_step_deg + self.subpeak_refine_iters = subpeak_refine_iters + self.subpeak_refine_max_move_deg = subpeak_refine_max_move_deg + + self.n_rotation_candidates = n_rotation_candidates + self.n_translation_peaks = n_translation_peaks + self.n_translation_candidates = n_translation_candidates + self.translation_grid_steps = translation_grid_steps + self.use_llg_tf = use_llg_tf + + self.do_joint_refine = do_joint_refine + self.dense_rotation_refine = dense_rotation_refine + self.joint_refine_max_res_A = joint_refine_max_res_A + self.joint_refine_expected_rot_error = joint_refine_expected_rot_error + self.refine_b = refine_b + self.sigma_rot_deg = sigma_rot_deg + self.sigma_trans_ang = sigma_trans_ang + self.sigma_b = sigma_b + + self.min_tries = min_tries + self.max_tries = max_tries + self.rfactor_converged = rfactor_converged + + # Vestigial; retained so legacy callers/benchmarks do not break. + self._vestigial = dict( + L=L, + use_sigma_a_frf=use_sigma_a_frf, + frf_delta_vrms_A=frf_delta_vrms_A, + frf_weight_combine=frf_weight_combine, + use_m_symmetry_filter=use_m_symmetry_filter, + use_lerf1_intensity=use_lerf1_intensity, + use_fitted_delta_vrms=use_fitted_delta_vrms, + use_even_l_only=use_even_l_only, + ) - def run( - self, - n_rotation_peaks: int = 200, - n_translation_peaks: int = 5, - min_tries: int = 3, - max_tries: int = 10, - rfactor_converged: float = 0.45, - max_clash_score: float = 100.0, - d_min: float = 4.0, - d_max: float = 50.0, - L: int = 48, - P: int = 24, - cluster_threshold_deg: Optional[float] = None, - ) -> List[MRSolution]: - """ - Run full MR pipeline with early stopping. + self._timer = _StageTimer(enabled=verbose >= 2) + # Filled in by run(). + self._frf = None + self._F_obs_amp = None + self._hkl_keep = None + self._tmask = None + self._eye3 = torch.eye(3, dtype=torch.float64) + + # ------------------------------------------------------------------ + # Public entry point + # ------------------------------------------------------------------ + def run(self, do_translation: bool = True) -> List[MRSolution]: + """Run the MR pipeline and return solutions ranked by R-factor. Parameters ---------- - n_rotation_peaks : int - Number of rotation peaks to try. - n_translation_peaks : int - Number of translation peaks per rotation. - min_tries : int - Minimum number of candidates to refine before early stopping. - max_tries : int - Maximum number of candidates to refine. - rfactor_converged : float - R-factor threshold for convergence (early stopping). - max_clash_score : float - Maximum clash score to accept a candidate. - d_min : float - High resolution limit for rotation search (Angstroms). - d_max : float - Low resolution limit for rotation search (Angstroms). - L : int - Angular bandlimit for rotation search. - P : int - Radial bandlimit for rotation search. - cluster_threshold_deg : float - Angular distance threshold for clustering rotation peaks (degrees). - Peaks within this angular distance are considered the same solution. - Default 6.0 degrees matches rigid body refinement convergence radius. + do_translation : bool + If ``False``, stop after rotation rescoring and return a single + rotation-only solution (the model rotated onto the best + orientation, no translation or refinement). Returns ------- - List[MRSolution] - Solutions sorted by R-factor. Early stops if converged. + list of MRSolution + Sorted by ``r_factor`` (ascending). The first element is the best + placement; its ``r_factor`` is the solvent-aware Scaler R-work. """ - # Auto cluster threshold: tied to grid voxel resolution. - if cluster_threshold_deg is None: - cluster_threshold_deg = max(6.0, 180.0 / L) - - # Step 1: Rotation search - if self.verbose: - print("Step 1: Rotation search...") - rotation_peaks = self._rotation_search( - n_rotation_peaks, d_min, d_max, L, P, - ) - - if not rotation_peaks: - if self.verbose: - print(" No rotation peaks found!") - return [] - - # Step 1b: Cluster rotation peaks - if self.verbose: - print(f" Found {len(rotation_peaks)} raw peaks, clustering with {cluster_threshold_deg}° threshold...") - - # Get symmetry matrices for clustering - sym_matrices = None - if hasattr(self.model, 'symmetry') and self.model.symmetry is not None: - sym_matrices = np.array([s.numpy() for s in self.model.symmetry.matrices]) + if not self.model.initialized: + raise RuntimeError( + "Cannot fit an uninitialized ModelFT. Load PDB data first." + ) - rotation_peaks = cluster_rotation_peaks( - rotation_peaks, - threshold_deg=cluster_threshold_deg, - symmetry_matrices=sym_matrices, + timer = self._timer + timer.start("0_data_prep") + frf = _prepare_frf_inputs( + self.model, self.data, + d_min=self.d_min, d_max=self.d_max, n_shells=self.n_shells, + ll_padding_factor=self.ll_padding_factor, + ll_max_res_A=self.ll_max_res_A, verbose=self.verbose, ) - - if self.verbose: - print(f" {len(rotation_peaks)} unique rotation clusters") - - if not rotation_peaks: - if self.verbose: - print(" No rotation peaks after clustering!") - return [] - - # Step 2: Translation search for each rotation - if self.verbose: - print(f"Step 2: Translation search for {len(rotation_peaks)} rotations...") - candidates = [] - for i, rot in enumerate(rotation_peaks): - trans_peaks = self._translation_search(rot, n_translation_peaks) - for trans in trans_peaks: - candidates.append((rot, trans)) - if self.verbose > 1 and (i + 1) % 10 == 0: - print(f" Processed {i+1}/{len(rotation_peaks)} rotations...") - - if self.verbose: - print(f" Generated {len(candidates)} rotation+translation candidates") - - # Step 3: Score and filter by clash - if self.verbose: - print("Step 3: Clash filtering...") - candidates = self._score_and_filter(candidates, max_clash_score) - - if not candidates: - if self.verbose: - print(" All candidates rejected by clash filter!") - return [] - - if self.verbose: - print(f" {len(candidates)} candidates passed clash filter") - - # Step 4: Rigid body refinement with early stopping - if self.verbose: - print(f"Step 4: Refining candidates (min={min_tries}, max={max_tries})...") - - solutions = [] - best_r_factor = float('inf') - converged = False - - n_to_refine = min(len(candidates), max_tries) - for i, (rot, trans, clash) in enumerate(candidates[:n_to_refine]): - if self.verbose > 1: - print(f" Refining candidate {i+1}/{n_to_refine}...") - - try: - result = self._rigid_body_refine(rot, trans) - - solution = MRSolution( - rotation=np.array([rot[0], rot[1], rot[2]]), - translation=trans.translation, - rotation_score=rot[4], # sigma - translation_score=trans.score, - clash_score=clash, - r_factor=result.final_r_factor, - refined_rotation=np.array(result.final_rotation), - refined_translation=np.array(result.final_translation_frac), + timer.stop("0_data_prep") + self._frf = frf + + # --- Stage 1+2: FRF rotation search + ML rescore --- + rescored = self._rotation_candidates(frf) + if not rescored: + raise RuntimeError("Rotation search produced no peaks.") + + if not do_translation: + rotated, R_rec = self._make_rotated(rescored[0]) + top = rescored[0] + if self.verbose > 0: + print( + f"fit_to_data: top peak LLG = {top.score:.2f} " + f"(σ_Z = {top.sigma:.2f}); applying R⁻¹ to coords.", + flush=True, ) - solutions.append(solution) - - # Track best - if result.final_r_factor < best_r_factor: - best_r_factor = result.final_r_factor - - # Check early stopping - n_refined = i + 1 - if n_refined >= min_tries and best_r_factor < rfactor_converged: - converged = True - if self.verbose: - print(f" Converged! R-factor {best_r_factor:.4f} < {rfactor_converged}") - break + if self.verbose >= 2: + print("\n" + timer.summary(), flush=True) + return [ + MRSolution( + rotation=R_rec.detach().cpu().numpy(), + translation=None, + rotation_score=float(top.score), + translation_score=float("nan"), + r_factor=float("nan"), + model=rotated, + ) + ] + + # --- Stage 3+: per-candidate translation + post-refine tree --- + self._prepare_translation_arrays() + n_rot = min(self.n_rotation_candidates, len(rescored)) + max_tries = self.max_tries if self.max_tries is not None else n_rot + if self.verbose > 0 and n_rot > 1: + print( + f"fit_to_data: trying up to {n_rot} rotation candidates " + f"(early-stop after ≥{self.min_tries} once R < " + f"{self.rfactor_converged}).", + flush=True, + ) - except Exception as e: - if self.verbose > 1: - print(f" Refinement failed: {e}") + solutions: List[MRSolution] = [] + best_r = float("inf") + for k in range(n_rot): + peak_k = rescored[k] + rotated_k, R_rec_k = self._make_rotated(peak_k) + if self.verbose > 0: + print( + f"\nfit_to_data: rot{k} " + f"(LLG={peak_k.score:.2f}, σ_Z={peak_k.sigma:.2f})", + flush=True, + ) + placement = self._placement_for_candidate(rotated_k) + if placement is None: + if self.verbose > 0: + print(" no translation peaks; skipping", flush=True) continue + r_analytic, t_refined = placement + + refined = rotated_k.translate( + t_refined.to(self.model.dtype_float), fractional=True, + ) + r_rank = r_analytic + if self.do_joint_refine: + if self.dense_rotation_refine: + refined = self._dense_rotation_refine(refined) + polished, rb_result = self._rigid_body_polish(refined) + if rb_result.final_r_factor <= rb_result.initial_r_factor: + refined = polished + r_rank = rb_result.final_r_factor + if self.verbose > 0: + print( + f" joint polish {rb_result.initial_r_factor:.4f} → " + f"{rb_result.final_r_factor:.4f} (no-solvent R)", + flush=True, + ) + else: + r_rank = rb_result.initial_r_factor + if self.verbose > 0: + print( + f" joint polish kept original " + f"({rb_result.initial_r_factor:.4f} ≤ " + f"{rb_result.final_r_factor:.4f} no-solvent R)", + flush=True, + ) + + refined.last_alignment_rotation = R_rec_k + refined.last_alignment_translation = t_refined + solutions.append( + MRSolution( + rotation=R_rec_k.detach().cpu().numpy(), + translation=t_refined.detach().cpu().numpy(), + rotation_score=float(peak_k.score), + translation_score=float(r_analytic), + r_factor=float(r_rank), + model=refined, + ) + ) + best_r = min(best_r, r_rank) + + n_done = k + 1 + if n_done >= self.min_tries and best_r < self.rfactor_converged: + if self.verbose > 0: + print( + f"fit_to_data: converged (R {best_r:.4f} < " + f"{self.rfactor_converged}) after {n_done} candidates.", + flush=True, + ) + break + if n_done >= max_tries: + break - if self.verbose: - status = "converged" if converged else f"completed {len(solutions)}/{n_to_refine}" - print(f" Refinement {status}") - if solutions: - print(f" Best R-factor: {best_r_factor:.4f}") + if not solutions: + raise RuntimeError("Translation + joint refine produced no candidates.") solutions.sort(key=lambda s: s.r_factor) + winner = solutions[0] + + # Single solvent-aware Scaler refit on the winner for the user-facing R. + timer.start("12_final_scaler") + rwork_final = _external_rwork(winner.model, self.data) + timer.stop("12_final_scaler") + winner.model.last_alignment_rfactor = rwork_final + winner.r_factor = rwork_final + if self.verbose > 0: + print( + f"fit_to_data: winner analytical-TF R={winner.translation_score:.4f}, " + f"final Scaler-fit R-work={rwork_final:.4f}", + flush=True, + ) + if self.verbose >= 2: + print("\n" + timer.summary(), flush=True) return solutions - def _rotation_search( - self, - n_peaks: int, - d_min: float, - d_max: float, - L: int, - P: int, - ) -> list: - """Run the rotation search via the Phaser-faithful FRF. - - Single engine post-consolidation: dense calc + auto_lmax + obs-unroll - + no_grad. Returns ``(α, β, γ, score, σ)`` tuples. - """ - from .align import _prepare_frf_inputs, _run_frf_separate_rotation - frf = _prepare_frf_inputs( - self.model, self.data, - d_min=d_min, d_max=d_max, n_shells=20, verbose=self.verbose, - ) + # ------------------------------------------------------------------ + # Stage 1+2: rotation search + ML rescore + # ------------------------------------------------------------------ + def _rotation_candidates(self, frf) -> list: + """FRF rotation search followed by ML rescoring of the top peaks.""" + data = self.data + device = self.device + timer = self._timer + + timer.start("3_rotation_search") + if self.verbose > 0: + print( + f"fit_to_data: frf_separate rotation search " + f"(dense calc + auto_lmax cap={self.frf_lmax_cap}, " + f"n_peaks={self.n_rotation_peaks})…", + flush=True, + ) peaks = _run_frf_separate_rotation( - self.model, self.data, frf, - n_peaks=n_peaks, verbose=self.verbose, + self.model, data, frf, + lmax_cap=self.frf_lmax_cap, dense_pad=self.frf_dense_pad, + n_peaks=self.n_rotation_peaks, verbose=self.verbose, ) - return [(p.alpha, p.beta, p.gamma, p.score, p.sigma) for p in peaks] - - def _translation_search( - self, - rotation_peak: tuple, - n_peaks: int, - ) -> List[TranslationPeak]: - """Run translation search for a rotation.""" - alpha, beta, gamma, _, _ = rotation_peak - - # Apply rotation to model coordinates - R = torch.tensor( - rotation_matrix_from_euler_zyz(alpha, beta, gamma), - dtype=get_float_dtype(), - device=self.device, - ) - xyz = self.model.xyz() - xyz_centered = xyz - xyz.mean(dim=0) - xyz_rotated = xyz_centered @ R.T - - # Temporarily update model coordinates and compute F_calc - original_xyz = self.model.xyz().clone() - self.model.xyz[:] = xyz_rotated + timer.stop("3_rotation_search") - try: - hkl = self.data.hkl - F_obs = self.data.F - mask = self.data.get_valid_mask() - - with torch.no_grad(): - F_calc = self.model(hkl) - - # Apply mask - F_obs_masked = F_obs[mask] - F_calc_masked = F_calc[mask] - hkl_masked = hkl[mask] - - _, _, peaks = fft_translation_search_torch( - F_obs_masked, F_calc_masked, hkl_masked, n_peaks=n_peaks + if self.rescore_engine not in ("m_letf1", "sim", "none"): + raise ValueError( + f"rescore_engine={self.rescore_engine!r}; " + "expected 'm_letf1' (default), 'sim' or 'none'." ) - finally: - # Restore original coordinates - self.model.xyz[:] = original_xyz - - return peaks - def _score_and_filter( - self, - candidates: list, - max_clash: float, - ) -> list: - """Score candidates and filter by clash.""" - if self._clash_calc is None: - self._clash_calc = ClashScoreCalculator( - symmetry=self.data.spacegroup, - default_clash_radius=4.0, - device=self.device, + # No ML rescore: rank candidates by the raw FRF score and let the + # multi-candidate tree (FTF + refine + R-ranking) do the discrimination. + if self.rescore_engine == "none": + if self.verbose > 0: + print("fit_to_data: ML rescore DISABLED — using raw FRF peak " + "ranking (RFZ).", flush=True) + return sorted(peaks, key=lambda p: p.score, reverse=True) + + F_obs = frf.F_obs + hkl = frf.hkl + s_mag = frf.s_mag + centric = frf.centric + ll = frf.ll + rescore_n_shells = max(self.n_shells // 2, 8) + + interp_var_main: Optional[torch.Tensor] = None + if self.use_interp_var: + rescore_edges, _ = equal_count_shell_edges(s_mag, rescore_n_shells) + rescore_shell_idx = assign_shells(s_mag, rescore_edges) + interp_var_main = estimate_interp_var( + ll, hkl, data.cell, rescore_shell_idx, rescore_n_shells, + ).to(F_obs.dtype) + + timer.start("4_ml_rescore") + if self.verbose > 0: + print( + f"fit_to_data: ML rescoring top " + f"{min(len(peaks), self.n_ml_refine)} peaks…", + flush=True, + ) + if self.rescore_engine == "m_letf1": + rescored = m_letf1_rescore( + peaks, F_obs, hkl, s_mag, centric, ll, data.cell, + data.spacegroup.matrices.to(torch.float64).to(device), + n_shells=rescore_n_shells, + n_refine=min(len(peaks), self.n_ml_refine), + batch_size=50, verbose=self.verbose, + scat_mode=self.rescore_scat_mode, + ) + if self.subpeak_refine: + rescored = self._subpeak_refine(rescored, F_obs, hkl, s_mag, + centric, ll, rescore_n_shells) + else: # legacy Sim/Rice approximation + rescored = sim_mlrf_rescore( + peaks, F_obs, hkl, s_mag, centric, ll, data.cell, + n_shells=rescore_n_shells, + n_refine=min(len(peaks), self.n_ml_refine), + batch_size=50, verbose=self.verbose, + auto_variance_weights=self.auto_variance_weights, + interp_var=interp_var_main, + ) + timer.stop("4_ml_rescore") + return rescored + + def _subpeak_refine(self, rescored, F_obs, hkl, s_mag, centric, ll, + rescore_n_shells): + """Quadratic tangent-space Newton sharpening of the top orientations.""" + from .ml_rotation import _build_llg_context, quadratic_llg_refine + + data = self.data + device = self.device + self._timer.start("4b_subpeak_refine") + ctx = _build_llg_context( + F_obs, hkl, s_mag, centric, ll, data.cell, + data.spacegroup.matrices.to(torch.float64).to(device), + n_shells=rescore_n_shells, batch_size=50, + scat_mode=self.rescore_scat_mode, + ) + k = self.subpeak_refine_k if self.subpeak_refine_k > 0 else self.n_rotation_candidates + k = min(k, len(rescored)) + rescored = quadratic_llg_refine( + rescored, ctx, k_refine=k, + step_deg=self.subpeak_refine_step_deg, + iterations=self.subpeak_refine_iters, + max_move_deg=self.subpeak_refine_max_move_deg, + verbose=self.verbose, + ) + self._timer.stop("4b_subpeak_refine") + if self.verbose > 0: + print( + f"fit_to_data: sub-peak refined top {k} orientations " + f"on the ML-LLG surface (step={self.subpeak_refine_step_deg}°).", + flush=True, ) + return rescored - scored = [] - for rot, trans in candidates: - clash = self._compute_clash(rot, trans) - if clash <= max_clash: - scored.append((rot, trans, clash)) + def _make_rotated(self, peak: "RotationPeak"): + """Rotate the search model onto a candidate orientation. - # Sort by combined score (rotation_sigma + trans_sigma - clash_penalty) - def combined_score(x): - rot, trans, clash = x - return rot[4] + trans.sigma - clash / 100.0 + Returns ``(rotated_model, R_recovered)`` where ``R_recovered`` maps the + search-model frame onto the crystal frame; the applied coordinate + rotation is ``R_recovered.T``. + """ + R_rec = rotation_matrix_from_edmonds_euler(peak.alpha, peak.beta, peak.gamma) + R_app = R_rec.T.contiguous() + rot = self.model.rotate( + R_app.to(device=self.model.device, dtype=self.model.dtype_float), + ) + rot.last_alignment_rotation = R_rec + return rot, R_rec + + # ------------------------------------------------------------------ + # Stage 3: per-candidate translation search + local refine + # ------------------------------------------------------------------ + def _prepare_translation_arrays(self) -> None: + """Resolution/validity-masked obs amplitudes + Miller indices.""" + data = self.data + device = self.device + hkl_full = data.hkl + F_obs_full = data.F + if hasattr(data, "get_valid_mask"): + tmask = data.get_valid_mask() + else: + tmask = torch.ones( + F_obs_full.shape[0], dtype=torch.bool, device=F_obs_full.device, + ) + self._tmask = tmask + self._F_obs_amp = F_obs_full[tmask].abs().to(torch.float64).to(device) + self._hkl_keep = hkl_full[tmask].to(device) - scored.sort(key=combined_score, reverse=True) - return scored + def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: + """Translation search + analytical-R local refine for one rotation. - def _compute_clash( - self, - rotation_peak: tuple, - trans_peak: TranslationPeak, - ) -> float: - """Compute clash score for a solution.""" - alpha, beta, gamma, _, _ = rotation_peak - R = torch.tensor( - rotation_matrix_from_euler_zyz(alpha, beta, gamma), - dtype=get_float_dtype(), - device=self.device, + Returns ``(r_analytic, t_refined)`` for the best translation of this + rotation candidate, or ``None`` if no translation peaks were found. + """ + data = self.data + device = self.device + timer = self._timer + eye3 = self._eye3 + + if str(rotated_k.spacegroup) != str(data.spacegroup): + rotated_k.spacegroup = data.spacegroup + rotated_p1 = rotated_k.copy() + rotated_p1.spacegroup = SpaceGroup("P 1") + evaluator = _DirectModelEvaluator(rotated_p1) + + timer.start("5_precompute_G") + G_pre, h_R_pre = precompute_G_for_rotation( + evaluator, eye3, self._hkl_keep, data.spacegroup, data.cell, + ) + timer.stop("5_precompute_G") + + timer.start("6_amplitude_TF") + _, _, t_peaks = amplitude_translation_search( + F_obs=self._F_obs_amp, interpolator=evaluator, + R_rotation=eye3, hkl=self._hkl_keep, + spacegroup=data.spacegroup, real_cell=data.cell, + grid_steps=self.translation_grid_steps, + n_peaks=self.n_translation_peaks, + cluster_radius=0.05, + precomputed_G=G_pre, precomputed_h_R=h_R_pre, ) + timer.stop("6_amplitude_TF") + if not t_peaks: + return None + + if self.use_llg_tf: + t_peaks = self._llg_tf_rescore(t_peaks, G_pre, h_R_pre) + + if self.verbose > 0: + tt = tuple(round(float(x), 3) for x in t_peaks[0].translation.tolist()) + print(f" top translation t={tt} score={t_peaks[0].score:.4f}", + flush=True) + + # do_joint_refine=False: take the top translation peak directly (no + # local refine), rank by its correlation score (negated so lower=better + # like an R-factor). + if not self.do_joint_refine: + t_top = torch.as_tensor(t_peaks[0].translation, dtype=torch.float64) + return -float(t_peaks[0].score), t_top + + best = None + for k_t, tp in enumerate(t_peaks[:self.n_translation_candidates]): + t_init = torch.as_tensor(tp.translation, dtype=torch.float64) + timer.start("7_local_TF_refine") + t_refined, r_analytic = local_translation_refine( + F_obs=self._F_obs_amp, interpolator=evaluator, + R_rotation=eye3, hkl=self._hkl_keep, + spacegroup=data.spacegroup, real_cell=data.cell, + t_init=t_init, radius=0.06, grid_steps=13, + n_refinement_passes=1, + precomputed_G=G_pre, precomputed_h_R=h_R_pre, + ) + timer.stop("7_local_TF_refine") + if self.verbose > 0: + print( + f" trans{k_t}: R(analytic)={r_analytic:.4f}, " + f"t={[round(float(x), 3) for x in t_refined.tolist()]}", + flush=True, + ) + if best is None or r_analytic < best[0]: + best = (r_analytic, t_refined) + return best - xyz = self.model.xyz() - xyz_centered = xyz - xyz.mean(dim=0) - xyz_rotated = xyz_centered @ R.T + def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): + """Re-rank translation peaks by a shared-σA Rice/Woolfson LLG. - # Apply translation (fractional -> Cartesian) - trans_frac = torch.tensor( - trans_peak.translation, - dtype=get_float_dtype(), - device=self.device, + Mirrors Phaser's FTF — the cheap amplitude correlation is a fast + pre-filter but ranks poorly for partial models; the LLG ranks + consistently with the rotation rescore. + """ + data = self.data + device = self.device + F_obs_amp = self._F_obs_amp + hkl_keep = self._hkl_keep + tmask = self._tmask + self._timer.start("6b_llg_tf_rescore") + + rec_basis_keep = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) + s_mag_keep_tf = (hkl_keep.to(torch.float64) @ rec_basis_keep).norm(dim=-1) + tf_n_shells = max(self.n_shells // 2, 8) + tf_edges, _ = equal_count_shell_edges(s_mag_keep_tf, tf_n_shells) + tf_shell_idx = assign_shells(s_mag_keep_tf, tf_edges) + centric_keep_tf = ( + data.centric[tmask].to(torch.bool).to(device) + if hasattr(data, "centric") + else torch.zeros_like(F_obs_amp, dtype=torch.bool) ) - trans_cart = trans_frac @ self.data.cell.fractional_matrix.to(self.device) - xyz_final = xyz_rotated + trans_cart - atom_mask = AtomSampler.from_model(self.model, mode='ca_only') - with torch.no_grad(): - clash = self._clash_calc( - xyz=xyz_final, - cell=self.data.cell.data, - atom_mask=atom_mask, - ).item() - return clash - - def _rigid_body_refine( - self, - rotation_peak: tuple, - trans_peak: TranslationPeak, - ) -> RigidBodyResult: - """Run rigid body refinement.""" - alpha, beta, gamma, _, _ = rotation_peak + cnt_tf = torch.bincount(tf_shell_idx, minlength=tf_n_shells).to(torch.float64) + sum_F2 = torch.zeros(tf_n_shells, dtype=torch.float64, device=device) + sum_F2.scatter_add_(0, tf_shell_idx, F_obs_amp * F_obs_amp) + mean_F2 = (sum_F2 / cnt_tf.clamp(min=1.0)).clamp(min=1e-30) + E_obs_tf = F_obs_amp / mean_F2.sqrt().index_select(0, tf_shell_idx) - rb = RigidBodyRefinement( - model=self.model, - data=self.data, - initial_rotation=torch.tensor([alpha, beta, gamma], dtype=get_float_dtype()), - initial_translation=torch.tensor( - trans_peak.translation, dtype=get_float_dtype() - ), - device=self.device, - verbose=max(0, self.verbose - 1), + t_top_t = torch.as_tensor( + t_peaks[0].translation, dtype=torch.float64, device=device, + ) + phase_top = torch.exp( + 2j * torch.pi * torch.einsum( + "ind,d->in", h_R_pre.to(torch.float64), t_top_t, + ).to(G_pre.dtype), + ) + Fc_top = (G_pre * phase_top).sum(dim=0).abs().to(torch.float64) + sum_Fc2 = torch.zeros(tf_n_shells, dtype=torch.float64, device=device) + sum_Fc2.scatter_add_(0, tf_shell_idx, Fc_top * Fc_top) + mean_Fc2 = (sum_Fc2 / cnt_tf.clamp(min=1.0)).clamp(min=1e-30) + E_calc_top = Fc_top / mean_Fc2.sqrt().index_select(0, tf_shell_idx) + sigma_a_tf = fit_sigma_a_per_shell( + E_obs_tf, E_calc_top, centric_keep_tf, + tf_shell_idx, tf_n_shells, n_grid=81, ) - return rb.refine() - def _get_e_values_obs( - self, - d_min: float, - d_max: float, - ) -> Tuple[torch.Tensor, torch.Tensor]: - """Get E-values for observed data (cached).""" - if self._e_obs is None: - from torchref.base.alignment import F_squared_to_E_values - - F_obs = self.data.F - mask = self.data.get_valid_mask() - F_obs_masked = F_obs[mask] - - F2 = (F_obs_masked ** 2).to(torch.float64) - s = self._get_s_vectors()[mask] - - E_values, E_squared, _ = F_squared_to_E_values( - F2, s, n_shells=20, d_min=d_min, d_max=d_max + t_cands = torch.as_tensor( + np.stack([p.translation for p in t_peaks]), + dtype=torch.float64, device=device, + ) + llg_tf = llg_translation_rescore( + F_obs=F_obs_amp, hkl=hkl_keep, centric=centric_keep_tf, + shell_idx=tf_shell_idx, n_shells=tf_n_shells, + G=G_pre, h_R=h_R_pre, t_candidates=t_cands, + sigma_a=sigma_a_tf, interp_var=None, + ) + self._timer.stop("6b_llg_tf_rescore") + + llg_list = llg_tf.detach().cpu().tolist() + order = sorted(range(len(t_peaks)), key=lambda i: llg_list[i], reverse=True) + return [ + TranslationPeak( + translation=t_peaks[i].translation, + score=float(llg_list[i]), + sigma=float(llg_list[i]), ) - self._e_obs = E_squared # Use E² for correlation - self._s_obs = s - self._mask = mask # Cache mask for later use - return self._e_obs, self._s_obs + for i in order + ] - def _get_e_values_calc( - self, - d_min: float, - d_max: float, - ) -> Tuple[torch.Tensor, torch.Tensor]: - """Get E-values for calculated data (cached).""" - if self._e_calc is None: - from torchref.base.alignment import F_squared_to_E_values + # ------------------------------------------------------------------ + # Stage 4+5: dense rotation re-sampling + rigid-body polish + # ------------------------------------------------------------------ + def _dense_rotation_refine(self, refined): + """Two-pass dense rotation re-sampling on the ML-LLG surface. - hkl = self.data.hkl - mask = self.data.get_valid_mask() + Zooms the orientation onto the (sharper) ML-LLG basin at the found + translation before the LBFGS polish. Returns the re-rotated model. + """ + data = self.data + device = self.device + timer = self._timer + + timer.start("8_dense_R_ll_build") + refined_p1 = refined.copy() + refined_p1.spacegroup = SpaceGroup("P 1") + ll_refine = LattmanLoveInterpolator( + refined_p1, padding_factor=self.ll_padding_factor, + max_res_A=self.ll_max_res_A, verbose=0, + ) + timer.stop("8_dense_R_ll_build") + + tmask = self._tmask + hkl_keep = self._hkl_keep + F_obs_amp = self._F_obs_amp + centric_keep = ( + data.centric[tmask].to(torch.bool).to(device) if hasattr(data, "centric") + else torch.zeros(hkl_keep.shape[0], dtype=torch.bool, device=device) + ) + rec_basis_keep = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) + s_mag_keep = (hkl_keep.to(torch.float64) @ rec_basis_keep).norm(dim=-1) + rescore_n_shells = max(self.n_shells // 2, 8) + + n_per_axis_pass = [9, 5] + zoom_factor = 4.0 + radii = [ + float(self.joint_refine_expected_rot_error), + float(self.joint_refine_expected_rot_error) / zoom_factor, + ] + R_accumulated = torch.eye(3, dtype=torch.float64) + + interp_var_dense: Optional[torch.Tensor] = None + if self.use_interp_var: + dense_edges, _ = equal_count_shell_edges(s_mag_keep, rescore_n_shells) + dense_shell_idx = assign_shells(s_mag_keep, dense_edges) + interp_var_dense = estimate_interp_var( + ll_refine, hkl_keep, data.cell, dense_shell_idx, rescore_n_shells, + ).to(F_obs_amp.dtype) + + for pass_idx, max_perturb_rad in enumerate(radii): + n_per_axis = n_per_axis_pass[pass_idx] + coords_r = torch.linspace( + -max_perturb_rad, max_perturb_rad, n_per_axis, dtype=torch.float64, + ) + wx, wy, wz = torch.meshgrid(coords_r, coords_r, coords_r, indexing="ij") + omegas = torch.stack([wx.flatten(), wy.flatten(), wz.flatten()], dim=-1) + R_perturbs = _rodrigues(omegas) + R_cand_full = R_perturbs @ R_accumulated + cand_peaks = [] + for R_c in R_cand_full: + a, b, g = edmonds_euler_from_rotation_matrix(R_c) + cand_peaks.append(RotationPeak(alpha=a, beta=b, gamma=g, + score=0.0, sigma=0.0)) + if self.verbose > 0: + print( + f" dense R pass {pass_idx + 1} " + f"({n_per_axis}³={omegas.shape[0]} perturbations, " + f"±{math.degrees(max_perturb_rad):.2f}°)…", + flush=True, + ) + rescore_batch = max(4, min(100, 1_000_000 // max(hkl_keep.shape[0], 1))) + timer.start("9_dense_R_rescore") + if self.rescore_engine == "m_letf1": + rescored_refine = m_letf1_rescore( + cand_peaks, F_obs_amp, hkl_keep, s_mag_keep, centric_keep, + ll_refine, data.cell, + data.spacegroup.matrices.to(torch.float64).to(device), + n_shells=rescore_n_shells, + n_refine=len(cand_peaks), batch_size=rescore_batch, verbose=0, + ) + else: + rescored_refine = sim_mlrf_rescore( + cand_peaks, F_obs_amp, hkl_keep, s_mag_keep, centric_keep, + ll_refine, data.cell, + n_shells=rescore_n_shells, + n_refine=len(cand_peaks), batch_size=rescore_batch, + verbose=0, n_D_grid=11, interp_var=interp_var_dense, + ) + timer.stop("9_dense_R_rescore") + top = rescored_refine[0] + best_idx = next( + i for i, p in enumerate(cand_peaks) + if p.alpha == top.alpha and p.beta == top.beta + and p.gamma == top.gamma + ) + R_accumulated = R_cand_full[best_idx] + if self.verbose > 0: + print( + f" pass {pass_idx + 1} best LLG={top.score:.2f}, " + f"|ω|={omegas[best_idx].norm().item() * 180 / math.pi:.3f}°", + flush=True, + ) - with torch.no_grad(): - F_calc = self.model(hkl).abs() + return refined.rotate( + R_accumulated.T.to(self.model.dtype_float).contiguous(), + ) - F_calc_masked = F_calc[mask] - F2 = (F_calc_masked ** 2).to(torch.float64) - s = self._get_s_vectors()[mask] + def _rigid_body_polish(self, refined): + """LBFGS rigid-body polish on (R, t) — returns ``(polished, result)``. - E_values, E_squared, _ = F_squared_to_E_values( - F2, s, n_shells=20, d_min=d_min, d_max=d_max - ) - self._e_calc = E_squared - self._s_calc = s - return self._e_calc, self._s_calc - - def _get_s_vectors(self) -> torch.Tensor: - """Get scattering vectors (cached).""" - if self._s_vectors is None: - from torchref.base import reciprocal_basis_matrix - rec_basis = reciprocal_basis_matrix(self.model.cell) - self._s_vectors = self.data.hkl.to(torch.float64) @ rec_basis.to(torch.float64) - return self._s_vectors - - def clear_cache(self): - """Clear cached E-values and s-vectors.""" - self._e_obs = None - self._e_calc = None - self._s_vectors = None - self._s_obs = None - self._s_calc = None - self._mask = None + The model is pre-rotated/translated; the refinement optimises a small + delta with ``initial_translation=0``. ``result`` carries the no-solvent + initial/final R-work used by the caller's accept/reject gate. + """ + timer = self._timer + timer.start("11_lbfgs_polish") + rb = RigidBodyRefinement( + refined, self.data, + initial_translation=torch.zeros( + 3, dtype=torch.float32, device=refined.device, + ), + expected_rotational_error=self.joint_refine_expected_rot_error, + max_res=self.joint_refine_max_res_A, + device=refined.device, + verbose=max(0, self.verbose - 1), + refine_b=self.refine_b, + sigma_rot_deg=self.sigma_rot_deg, + sigma_trans_ang=self.sigma_trans_ang, + sigma_b=self.sigma_b, + ) + rb_result = rb.refine() + with torch.no_grad(): + R_polish = rb.get_rotation_matrix().detach() + t_polish = rb.translation_frac.detach() + polished = refined.rotate(R_polish.to(self.model.dtype_float)) + polished = polished.translate( + t_polish.to(self.model.dtype_float), fractional=True, + ) + timer.stop("11_lbfgs_polish") + return polished, rb_result From c284281d90efed7c4a856e27c47b094b3aa6ddce Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 11:28:34 +0200 Subject: [PATCH 008/250] Carry the derived per-atom state through Model.copy `_rebuild_sf_indices()` builds the iso/aniso partition (`_iso_indices`, `_aniso_indices` and the two fast-path flags) from `aniso_flag` and the heavy-atom mask. It ran only in `load()`, and the partition is neither a buffer nor a parameter wrapper, so `copy()` did not carry it: the copy raised `AttributeError` from `get_iso()`/`get_aniso()`. `Model.copy()` also assigned the space group as an object. `spacegroup` is a property, but `SpaceGroup` is an `nn.Module`, so `nn.Module.__setattr__` intercepted the assignment, stored it in `_modules` under the property's own name and never ran the setter -- registering the *original's* SpaceGroup as a second submodule of the copy, shared by identity. `ModelFT.copy()` already assigned `_spacegroup` directly; `Model.copy()` now matches it. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- tests/unit/model/test_copy.py | 101 ++++++++++++++++++++++++++++++++++ torchref/model/model.py | 19 ++++++- torchref/model/model_ft.py | 3 + 3 files changed, 121 insertions(+), 2 deletions(-) create mode 100644 tests/unit/model/test_copy.py diff --git a/tests/unit/model/test_copy.py b/tests/unit/model/test_copy.py new file mode 100644 index 00000000..6f49c23d --- /dev/null +++ b/tests/unit/model/test_copy.py @@ -0,0 +1,101 @@ +"""``Model.copy`` / ``ModelFT.copy`` must carry the derived per-atom state. + +Two things in ``copy()`` are neither buffers nor parameter wrappers, so the +buffer and module loops do not carry them: + +* the **iso/aniso partition** (``_iso_indices``, ``_aniso_indices`` and the two + fast-path flags) is rebuilt from ``aniso_flag`` and the heavy-atom mask by + ``_rebuild_sf_indices``, which otherwise runs only in ``load()``. Without it a + copy raises ``AttributeError`` from ``get_iso()``/``get_aniso()``. +* the **space group**. ``Model.spacegroup`` is a property, but ``SpaceGroup`` is + an ``nn.Module``, so ``model.spacegroup = sg_object`` is intercepted by + ``nn.Module.__setattr__``, stored in ``_modules`` under the property's own + name, and the setter never runs. The copy must own its space group, not + register the original's under a second key. + +``4BX9`` is used because it carries ``ANISOU`` records, so the partition is +genuinely mixed (220 isotropic, 9973 anisotropic) rather than all-isotropic, +where an empty partition would still pass the fast path. +""" + +import pytest +import torch + +_MIXED_ADP_PDB = "4BX9.pdb" + + +@pytest.fixture(scope="module") +def mixed_adp_path(pdb_dir): + p = pdb_dir / _MIXED_ADP_PDB + if not p.exists(): + pytest.skip(f"{_MIXED_ADP_PDB} not available") + return str(p) + + +def _load(cls, path): + return cls(verbose=0).load_pdb(path) + + +@pytest.mark.unit +@pytest.mark.parametrize("cls_name", ["Model", "ModelFT"]) +def test_copy_has_a_usable_iso_aniso_partition(cls_name, mixed_adp_path): + """``get_iso``/``get_aniso`` must work on the copy and agree with the source.""" + import torchref.model as tm + + cls = getattr(tm, cls_name) + m = _load(cls, mixed_adp_path) + c = m.copy() + + for attr in ("_iso_indices", "_aniso_indices", "_iso_covers_all", + "_aniso_is_empty"): + assert hasattr(c, attr), f"{cls_name}.copy() dropped {attr}" + + assert torch.equal(c._iso_indices, m._iso_indices) + assert torch.equal(c._aniso_indices, m._aniso_indices) + assert c._iso_covers_all == m._iso_covers_all + assert c._aniso_is_empty == m._aniso_is_empty + + # ``ModelFT`` appends the per-atom form-factor tables to the same tuple, so + # unpack positionally rather than by a fixed arity. + iso = c.get_iso() + aniso = c.get_aniso() + xyz_i, occ_i = iso[0], iso[2] + xyz_a, u_a, occ_a = aniso[0], aniso[1], aniso[2] + assert xyz_i.shape[0] == m._iso_indices.numel() + assert xyz_a.shape[0] == m._aniso_indices.numel() + assert u_a.shape[-1] == 6 + assert occ_i.shape[0] == xyz_i.shape[0] + assert occ_a.shape[0] == xyz_a.shape[0] + + +@pytest.mark.unit +@pytest.mark.parametrize("cls_name", ["Model", "ModelFT"]) +def test_the_partition_is_mixed_so_the_test_has_teeth(cls_name, mixed_adp_path): + """Guard the premise: an all-isotropic model would not exercise the split.""" + import torchref.model as tm + + m = _load(getattr(tm, cls_name), mixed_adp_path) + assert m._iso_indices.numel() > 0 + assert m._aniso_indices.numel() > 0 + assert not m._iso_covers_all + assert not m._aniso_is_empty + + +@pytest.mark.unit +@pytest.mark.parametrize("cls_name", ["Model", "ModelFT"]) +def test_copy_owns_its_spacegroup(cls_name, mixed_adp_path): + """The copy carries the space group without aliasing or double-registering.""" + import torchref.model as tm + + m = _load(getattr(tm, cls_name), mixed_adp_path) + c = m.copy() + + assert c.spacegroup is not None + assert str(c.spacegroup) == str(m.spacegroup) + # Own object: `.to(device)` on the copy must not move the original's matrices. + assert c.spacegroup is not m.spacegroup + # Exactly one registration, under the private name the property reads. + assert "spacegroup" not in c._modules + assert "_spacegroup" in c._modules + stray = [k for k in c.state_dict() if k.startswith("spacegroup.")] + assert stray == [], f"copy registered a second space group: {stray}" diff --git a/torchref/model/model.py b/torchref/model/model.py index 12c6fd58..8fee4db9 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -928,8 +928,15 @@ def copy(self): model_copy.pdb = self.pdb.copy(deep=True) - # Setter also sets symmetry; gemmi.SpaceGroup is immutable, so shared. - model_copy.spacegroup = self.spacegroup + # Assign ``_spacegroup`` directly, as ``ModelFT.copy`` does. ``spacegroup`` + # is a property, but ``SpaceGroup`` is an ``nn.Module``, so assigning the + # object goes through ``nn.Module.__setattr__``, lands in ``_modules`` + # under the property's own name and never runs the setter -- registering + # the *original's* SpaceGroup as a second submodule of the copy. + if self._spacegroup is not None: + model_copy._spacegroup = self._spacegroup.copy() + else: + model_copy._spacegroup = None model_copy.initialized = True if self.cell is not None: @@ -941,7 +948,11 @@ def copy(self): # Parameter wrappers via their own .copy(), which preserves each # wrapper's parametrization (log-space, Cholesky, collapsed logits). + # ``_spacegroup`` is already copied above. + skip_modules = {"_spacegroup", "spacegroup", "_symmetry", "symmetry"} for module_name, module in self._modules.items(): + if module_name in skip_modules: + continue if module is not None and hasattr(module, "copy"): setattr(model_copy, module_name, module.copy()) @@ -952,6 +963,10 @@ def copy(self): else: model_copy.altloc_pairs = [] + # The iso/aniso partition is derived state, not a buffer, so it is not + # carried by the buffer loop above; get_iso()/get_aniso() read it. + model_copy._rebuild_sf_indices() + if self.verbose > 0: print(f"✓ Model copied successfully ({len(model_copy.pdb)} atoms)") diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index 85ec6e13..d9ca4df8 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -911,6 +911,9 @@ def copy(self, detach: bool = True) -> "ModelFT": # Don't share cached structure factors with the original. model_copy.reset_cache() + # The iso/aniso partition is derived state, not a buffer, so it is not + # carried by the buffer loop above; get_iso()/get_aniso() read it. + model_copy._rebuild_sf_indices() if self.verbose > 0: print(f"✓ ModelFT copied successfully ({len(model_copy.pdb)} atoms)") From 18899d6bafa2a8f595778d8068a5cbda9c529c2b Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 11:30:02 +0200 Subject: [PATCH 009/250] Fix the reciprocal-space symmetry convention in the alignment package Miller indices transform as `h' = h.S`, i.e. with the transpose of the symmetry rotation, which is what `SpaceGroup.apply_to_hkl` implements. Four sites re-derived the orbit expansion inline as `S.h`: align.py `_prepare_frf_inputs` obs unroll align.py `_run_frf_separate_rotation` non-epsilon unroll frf/preprocessing `epsilon_aware_unroll` sh.py `hkl_symops_to_cartesian` `S.h` and `h.S` agree whenever the symmetry matrices are orthogonal, which they are in every setting except trigonal and hexagonal -- so this was invisible on most of the benchmark. Where it is not orthogonal the unroll sites mix non-equivalent reflections into one orbit, and `hkl_symops_to_cartesian` returns matrices that are not rotations at all: orthogonality error 5.33 for P 3_1 2 1 and P 6_5 2 2, against 2e-7 with the transpose. `symmetrize_anisotropy` consumes those matrices, so the point-group projection of the anisotropy tensor was averaging over non-rotations and could increase the anisotropy it was meant to constrain. `compute_epsilon` already used the row-vector form and is unchanged. The new tests pin the premise (which settings are non-orthogonal), the observable consequence (the orbit differs as a set only when non-orthogonal), and the absence of new copies of the wrong contraction anywhere in the package -- the failure mode here was a shared helper being fixed while inline duplicates were missed. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- .../alignment/test_symmetry_conventions.py | 263 ++++++++++++++++++ torchref/experimental/alignment/align.py | 11 +- .../alignment/frf/preprocessing.py | 9 +- torchref/experimental/alignment/sh.py | 9 +- 4 files changed, 286 insertions(+), 6 deletions(-) create mode 100644 tests/unit/alignment/test_symmetry_conventions.py diff --git a/tests/unit/alignment/test_symmetry_conventions.py b/tests/unit/alignment/test_symmetry_conventions.py new file mode 100644 index 00000000..5e1e459d --- /dev/null +++ b/tests/unit/alignment/test_symmetry_conventions.py @@ -0,0 +1,263 @@ +"""Guard the reciprocal-space symmetry convention used by the FRF. + +Reciprocal space transforms as ``h' = h·R`` (equivalently ``Rᵀ·h``), the rule +:meth:`SpaceGroup.apply_to_hkl` implements. The alignment package re-derives +symmetry expansions inline in several places, and each of those inline copies +once used ``R·h`` instead. + +The reason that survived so long is worth stating, because it is what makes +these tests necessary rather than merely nice: **the two conventions agree +whenever the symmetry matrices are orthogonal**, which they are in every +monoclinic, orthorhombic, tetragonal and cubic setting. Only in a trigonal or +hexagonal basis, where ``S·Sᵀ ≠ I``, do they diverge -- and there the wrong +convention silently builds an orbit that mixes non-equivalent reflections, +writing conflicting ``|F|`` onto one Miller index. + +So every test here is parametrised over a space group whose matrices are *not* +orthogonal. A test that only covers P2₁2₁2₁ or P4₃2₁2 cannot fail. +""" + +from pathlib import Path + +import pytest +import torch + +from torchref.experimental.alignment.frf.preprocessing import ( + compute_epsilon, + epsilon_aware_unroll, +) +from torchref.experimental.alignment.sh import ( + hkl_symops_to_cartesian, + symmetrize_anisotropy, +) +from torchref.symmetry import SpaceGroup + + +#: Space groups spanning both regimes. ``non_orthogonal`` flags the settings +#: whose rotation matrices are not orthogonal in their own basis -- the only +#: ones that can discriminate ``h·R`` from ``R·h``. +SPACEGROUPS = [ + pytest.param("P 31 2 1", True, id="P3121-trigonal"), + pytest.param("P 65 2 2", True, id="P6522-hexagonal"), + pytest.param("P 63", True, id="P63-hexagonal"), + pytest.param("P 4 3 2", False, id="P432-cubic"), + pytest.param("P 21 21 2", False, id="P21212-orthorhombic"), + pytest.param("C 1 2 1", False, id="C2-monoclinic"), +] + + +def _cell_for(hm: str): + """A cell consistent with the space group's lattice constraints.""" + if hm.startswith("P 3") or hm.startswith("P 6"): + return (60.0, 60.0, 95.0, 90.0, 90.0, 120.0) + if hm.startswith("P 4 3") or hm.startswith("P 21 21"): + return (70.0, 70.0, 70.0, 90.0, 90.0, 90.0) if "4 3" in hm else ( + 45.0, 55.0, 65.0, 90.0, 90.0, 90.0) + return (80.0, 40.0, 60.0, 90.0, 104.0, 90.0) + + +def _reciprocal_basis(cell): + """``B`` from TorchRef's own :class:`Cell`, so the convention matches by + construction rather than by a hand-rolled duplicate.""" + from torchref.symmetry import Cell + + return Cell(list(cell)).reciprocal_basis_matrix.detach().cpu().to(torch.float64) + + +@pytest.mark.parametrize("hm, non_orthogonal", SPACEGROUPS) +def test_symop_orthogonality_flags_the_hard_cases(hm, non_orthogonal): + """The premise of every other test here: which settings can discriminate. + + If this ever reports a trigonal/hexagonal group as orthogonal, the tests + below stop testing anything. + """ + S = SpaceGroup(hm).matrices.detach().cpu().to(torch.float64) + eye = torch.eye(3, dtype=torch.float64) + err = max(float((S[k] @ S[k].T - eye).abs().max()) for k in range(S.shape[0])) + assert (err > 0.5) == non_orthogonal, ( + f"{hm}: max|S·Sᵀ-I| = {err:.3f}, expected " + f"{'non-orthogonal' if non_orthogonal else 'orthogonal'} matrices" + ) + + +@pytest.mark.parametrize("hm, non_orthogonal", SPACEGROUPS) +def test_cartesian_symops_are_rotations(hm, non_orthogonal): + """``hkl_symops_to_cartesian`` must return genuine rotations. + + A Cartesian symmetry operator is orthogonal with determinant +1 by + definition. Conjugating with the untransposed ``S`` does not produce one in + a non-orthogonal basis -- it returned matrices with orthogonality error 5.33 + for P 3₁ 2 1 and P 6₅ 2 2, which is what let a wrong anisotropy + symmetrisation through undetected. + """ + del non_orthogonal + S = SpaceGroup(hm).matrices.detach().cpu().to(torch.float64) + B = _reciprocal_basis(_cell_for(hm)) + R = hkl_symops_to_cartesian(S, B).detach().cpu().to(torch.float64) + + eye = torch.eye(3, dtype=torch.float64) + for k in range(R.shape[0]): + # 1e-6 is float64 noise through the matrix inverse; the defect this + # guards against produced an error of 5.33. + assert torch.allclose(R[k] @ R[k].T, eye, atol=1e-6), ( + f"{hm} op {k} is not orthogonal:\n{R[k]}" + ) + assert float(torch.det(R[k])) == pytest.approx(1.0, abs=1e-6), ( + f"{hm} op {k} has determinant {float(torch.det(R[k])):.6f}, not +1" + ) + + +@pytest.mark.parametrize("hm, non_orthogonal", SPACEGROUPS) +def test_unroll_matches_apply_to_hkl(hm, non_orthogonal): + """Any inline symmetry expansion must agree with the shared helper. + + ``apply_to_hkl`` is the single definition of the convention; the einsum + contraction ``"kji,nj->kni"`` is the same operation with the operator axis + first. ``"kij,nj->kni"`` is the bug and differs here for the non-orthogonal + settings. + """ + del non_orthogonal + sg = SpaceGroup(hm) + S = sg.matrices.detach().cpu().to(torch.float64) + g = torch.Generator().manual_seed(7) + hkl = torch.randint(-12, 13, (200, 3), generator=g).to(torch.float64) + + reference = sg.apply_to_hkl(hkl).detach().cpu().to(torch.float64) # (N, 3, ops) + unrolled = torch.einsum("kji,nj->kni", S, hkl) # (ops, N, 3) + assert torch.allclose(unrolled.permute(1, 2, 0), reference, atol=1e-9), ( + f"{hm}: inline unroll disagrees with SpaceGroup.apply_to_hkl" + ) + + +@pytest.mark.parametrize("hm, non_orthogonal", SPACEGROUPS) +def test_wrong_convention_changes_the_orbit(hm, non_orthogonal): + """``R·h`` yields a different *orbit* from ``h·R`` iff ``S`` is non-orthogonal. + + Per operator the two always differ unless ``S`` is symmetric, so the honest + statement is about the orbit as a set. For an orthogonal group ``Sᵀ = S⁻¹`` + is itself a group member, so the set is unchanged and the bug is invisible; + in a hexagonal basis ``Sᵀ`` leaves the group and the orbit genuinely moves. + + This is what gives the other tests their teeth: it asserts the defect is + observable at all for these settings, so a regression cannot hide behind a + benchmark built only from orthogonal lattices. + """ + S = SpaceGroup(hm).matrices.detach().cpu().to(torch.float64) + g = torch.Generator().manual_seed(11) + hkl = torch.randint(-12, 13, (200, 3), generator=g).to(torch.float64) + + def orbit_keys(contraction): + o = torch.einsum(contraction, S, hkl).round().to(torch.long) # (ops, N, 3) + k = ((o[..., 0] + 64) * 256 + (o[..., 1] + 64)) * 256 + (o[..., 2] + 64) + return torch.sort(k, dim=0).values # per-reflection set + + same_orbits = torch.equal(orbit_keys("kji,nj->kni"), orbit_keys("kij,nj->kni")) + assert (not same_orbits) == non_orthogonal, ( + f"{hm}: expected the orbits to " + f"{'differ' if non_orthogonal else 'coincide'} between conventions" + ) + + +@pytest.mark.parametrize("hm", ["P 31 2 1", "P 65 2 2"]) +def test_symmetrised_anisotropy_obeys_the_lattice(hm): + """With a 3- or 6-fold along c, Cartesian ``U`` must be ``diag(a, a, c)``. + + The previous conjugation returned a ``U`` *more* anisotropic than its input + (diag 0.90/1.30/0.60 became 1.58/3.70/0.60 with off-diagonals of 1.83), + which no averaging over a point group can do. + """ + S = SpaceGroup(hm).matrices.detach().cpu().to(torch.float64) + B = _reciprocal_basis(_cell_for(hm)) + R = hkl_symops_to_cartesian(S, B).detach().cpu().to(torch.float64) + + U = torch.tensor([[0.90, 0.10, 0.05], + [0.10, 1.30, -0.07], + [0.05, -0.07, 0.60]], dtype=torch.float64) + U_sym = symmetrize_anisotropy(U, R).detach().cpu().to(torch.float64) + + off = U_sym - torch.diag(U_sym.diagonal()) + assert float(off.abs().max()) < 1e-6, f"{hm}: U not diagonal:\n{U_sym}" + assert float(U_sym[0, 0]) == pytest.approx(float(U_sym[1, 1]), abs=1e-6), ( + f"{hm}: U11 != U22 ({U_sym[0,0]:.6f} vs {U_sym[1,1]:.6f})" + ) + # c is unconstrained by the axis, so it must survive untouched. + assert float(U_sym[2, 2]) == pytest.approx(0.60, abs=1e-6) + # An average cannot exceed the input's range. + assert float(U_sym[0, 0]) == pytest.approx(1.10, abs=1e-6) + + +@pytest.mark.parametrize("hm, non_orthogonal", SPACEGROUPS) +def test_epsilon_uses_the_row_vector_convention(hm, non_orthogonal): + """``compute_epsilon`` counts ops fixing ``h``, which needs ``h·R``. + + Reflections on a symmetry axis must come out with multiplicity > 1; with + the wrong convention the wrong reflections are flagged. + """ + del non_orthogonal + sg = SpaceGroup(hm) + S = sg.matrices.detach().cpu().to(torch.float64) + g = torch.Generator().manual_seed(3) + hkl = torch.randint(-9, 10, (400, 3), generator=g) + + eps = compute_epsilon(hkl, S).detach().cpu() + assert int(eps.min()) >= 1 + # Recompute independently through the shared helper. + ref = sg.apply_to_hkl(hkl.to(torch.float64)).detach().cpu() # (N, 3, ops) + same = (ref.round() == hkl.to(torch.float64).unsqueeze(-1)).all(dim=1) + assert torch.equal(eps.to(torch.long), same.sum(dim=-1).clamp(min=1)), ( + f"{hm}: compute_epsilon disagrees with apply_to_hkl" + ) + + +@pytest.mark.parametrize("hm, non_orthogonal", SPACEGROUPS) +def test_epsilon_aware_unroll_stays_within_the_true_orbit(hm, non_orthogonal): + """Every emitted position must be a genuine symmetry mate of its input. + + This exercises a real call site rather than the contraction in isolation. + Under the wrong convention the emitted positions leave the true orbit for a + non-orthogonal setting, which is what let two inequivalent reflections land + on one Miller index carrying different ``|F|``. + """ + del non_orthogonal + sg = SpaceGroup(hm) + S = sg.matrices.detach().cpu().to(torch.float64) + g = torch.Generator().manual_seed(19) + hkl = torch.randint(-9, 10, (150, 3), generator=g) + + unrolled, asu_idx = epsilon_aware_unroll(hkl, S) + unrolled = unrolled.detach().cpu().to(torch.long) + asu_idx = asu_idx.detach().cpu().to(torch.long) + + # true orbit of each parent, via the shared helper + orbit = sg.apply_to_hkl(hkl.to(torch.float64)).detach().cpu().round().to(torch.long) + for i in range(0, unrolled.shape[0], 17): # stride: full sweep is redundant + parent = int(asu_idx[i]) + mates = orbit[parent].T # (ops, 3) + assert (mates == unrolled[i]).all(dim=-1).any(), ( + f"{hm}: emitted {unrolled[i].tolist()} is not in the orbit of " + f"{hkl[parent].tolist()}" + ) + + +def test_no_column_convention_survives_in_the_alignment_package(): + """No ``S·h`` symmetry contraction may reappear in the alignment package. + + The four defects this module guards were four *copies* of one rule. The + other tests pin the rule; this one pins the absence of new copies, which is + the failure mode that actually occurred -- the shared helper was corrected + while the inline duplicates were not. + """ + import re + + pkg = Path(__file__).resolve().parents[3] / "torchref" / "experimental" / "alignment" + # `kij` contracted against an hkl-like index is the column convention. + pattern = re.compile(r'einsum\(\s*["\']k?ij,\s*nj->') + offenders = [] + for path in sorted(pkg.rglob("*.py")): + for lineno, line in enumerate(path.read_text().splitlines(), 1): + if pattern.search(line): + offenders.append(f"{path.relative_to(pkg.parent.parent.parent)}:{lineno}: {line.strip()}") + assert not offenders, ( + "reciprocal-space symmetry expansion must use h·R (\"kji,nj->\"), not " + "R·h (\"kij,nj->\"):\n " + "\n ".join(offenders) + ) diff --git a/torchref/experimental/alignment/align.py b/torchref/experimental/alignment/align.py index a22a8a17..5234bba6 100644 --- a/torchref/experimental/alignment/align.py +++ b/torchref/experimental/alignment/align.py @@ -320,7 +320,13 @@ def _prepare_frf_inputs( # the long comment in `align_model_to_data` for the rationale. sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) n_ops_sg = int(sg_mats.shape[0]) - hkl_sym = torch.einsum("kij,nj->kni", sg_mats, hkl.to(torch.float64)) + # h' = h.R, NOT R.h -- reciprocal space transforms with the transpose + # (SpaceGroup.apply_to_hkl). The two agree only when the symmetry matrices + # are orthogonal, which they are in orthorhombic/tetragonal/cubic and + # monoclinic settings but NOT in a hexagonal basis, where S.S^T != I. Using + # R.h there mixes non-equivalent reflections into one orbit and writes + # conflicting |F| onto the same Miller index. + hkl_sym = torch.einsum("kji,nj->kni", sg_mats, hkl.to(torch.float64)) hkl_sym_flat = hkl_sym.reshape(-1, 3) s_vec_sym = hkl_sym_flat @ rec_basis.to(device) s_mag_sym = s_vec_sym.norm(dim=-1) @@ -475,7 +481,8 @@ def _run_frf_separate_rotation( else: n_ops = int(sg_mats.shape[0]) hkl_keep = hkl_all.to(torch.float64)[keep] - hkl_unroll = torch.einsum("kij,nj->kni", sg_mats, hkl_keep).reshape(-1, 3) + # h' = h.R (transpose) -- see the note at the `hkl_sym` unroll. + hkl_unroll = torch.einsum("kji,nj->kni", sg_mats, hkl_keep).reshape(-1, 3) s_obs = hkl_unroll @ rec_basis hkl_obs_int = hkl_unroll F_obs = F_obs.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() diff --git a/torchref/experimental/alignment/frf/preprocessing.py b/torchref/experimental/alignment/frf/preprocessing.py index 665d830e..ffc8a62d 100644 --- a/torchref/experimental/alignment/frf/preprocessing.py +++ b/torchref/experimental/alignment/frf/preprocessing.py @@ -127,14 +127,17 @@ def epsilon_aware_unroll( hkl_int = hkl_int.to(torch.long) sym_mats = sym_mats.round().to(torch.long) N, n_ops = hkl_int.shape[0], sym_mats.shape[0] - # Orbits: (N, n_ops, 3) — S_k applied to each h (row-vector convention, - # matching the existing `einsum("kij,nj->nki", ...)` unroll site). + # Orbits: (N, n_ops, 3) — h.S_k, the row-vector (reciprocal-space) + # convention, matching the unroll sites in `align.py`. Note this is the + # TRANSPOSE contraction: `kji`, not `kij`. They coincide only for + # orthogonal symmetry matrices, so `kij` silently works everywhere except + # trigonal/hexagonal. # Integer einsum dispatches to baddbmm, which CUDA does not implement for # Long; compute in float64 (exact for symop 0/±1 × small Miller indices) # and round back so the GPU path works. orbits = ( torch.einsum( - "kij,nj->nki", sym_mats.to(torch.float64), hkl_int.to(torch.float64), + "kji,nj->nki", sym_mats.to(torch.float64), hkl_int.to(torch.float64), ) .round() .to(torch.long) diff --git a/torchref/experimental/alignment/sh.py b/torchref/experimental/alignment/sh.py index a1292ad4..6363f4fa 100644 --- a/torchref/experimental/alignment/sh.py +++ b/torchref/experimental/alignment/sh.py @@ -556,7 +556,14 @@ def hkl_symops_to_cartesian( M = rec_basis.to(dtype).transpose(-1, -2) # (3, 3) M_inv = torch.linalg.inv(M) S = sg_mats.to(dtype) # (n_ops, 3, 3) - return torch.einsum("ij,kjl,lm->kim", M, S, M_inv) + # S^T, not S: reciprocal space transforms as h' = h.S, so the operator + # acting on Cartesian s as a column vector is (B^-1 S B)^T = M S^T M^-1 + # with M = B^T. Using S here returns matrices that are not rotations at all + # in a non-orthogonal basis -- measured orthogonality error 5.33 for + # P 3_1 2 1 and P 6_5 2 2, versus 2e-7 with the transpose. The two agree + # whenever the symmetry matrices are orthogonal, i.e. everywhere except + # trigonal/hexagonal, which is why this survived. + return torch.einsum("ij,klj,lm->kim", M, S, M_inv) def symmetrize_anisotropy( From 67f06c7fc7508dda4e95c76d3c594a45d9b519cb Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 11:32:53 +0200 Subject: [PATCH 010/250] Stop molecular-replacement rotation candidates compounding onto each other `Model.rotate` and `Model.translate` mutate in place and return `self`, so every call site that treats them as returning a fresh model needs a `.copy()` first. `_make_rotated` is called once per rotation candidate off the same `self.model`, so candidate k+1 was evaluated at an orientation composed on top of candidate k, and `self.model` was destroyed in the process. Same pattern in `_dense_rotation_refine`, `_rigid_body_polish` (whose caller keeps the unpolished model when the polish does not improve R) and the post-placement translate. Also pass the space-group NAME rather than a SpaceGroup object at the three sites that build P1 copies. `spacegroup` is a property, but `SpaceGroup` is an `nn.Module`: object assignment is intercepted by `nn.Module.__setattr__`, filed under `_modules["spacegroup"]`, and the setter never runs -- so the "P1 search model" still carried the crystal symmetry. It also poisons the attribute, since a later correct string assignment then raises TypeError. Behaviour change to note: the `rescore_engine="none"` early return moved below the array extraction so sub-peak refinement runs in that branch too. It sharpens each orientation in place without reordering, so it is independent of which engine ranks the candidates -- but any baseline measured with `rescore_engine="none"` and `subpeak_refine=True` is not comparable across this commit. The new tests are fast, unlike the integration tests that previously covered this, which `--run-slow` gates off by default. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- .../alignment/test_patterson_translation.py | 4 +- .../unit/alignment/test_pipeline_aliasing.py | 130 ++++++++++++++++++ torchref/experimental/alignment/pipeline.py | 44 ++++-- 3 files changed, 161 insertions(+), 17 deletions(-) create mode 100644 tests/unit/alignment/test_pipeline_aliasing.py diff --git a/tests/unit/alignment/test_patterson_translation.py b/tests/unit/alignment/test_patterson_translation.py index bd826526..307b1ea1 100644 --- a/tests/unit/alignment/test_patterson_translation.py +++ b/tests/unit/alignment/test_patterson_translation.py @@ -59,7 +59,7 @@ def test_amplitude_tf_zero_translation(setup): with torch.no_grad(): F_obs = canonical(data.hkl[mask]).abs().to(torch.float64) model_p1 = canonical.copy() - model_p1.spacegroup = SpaceGroup("P 1") + model_p1.spacegroup = "P 1" evaluator = _ModelEvaluator(model_p1) R_id = torch.eye(3, dtype=torch.float64) @@ -93,7 +93,7 @@ def test_amplitude_tf_recovers_known_translation(setup): with torch.no_grad(): F_obs = canonical(data.hkl[mask]).abs().to(torch.float64) model_p1 = canonical.copy() - model_p1.spacegroup = SpaceGroup("P 1") + model_p1.spacegroup = "P 1" model_p1 = model_p1.translate( torch.tensor(t_true, dtype=canonical.dtype_float), fractional=True, ) diff --git a/tests/unit/alignment/test_pipeline_aliasing.py b/tests/unit/alignment/test_pipeline_aliasing.py new file mode 100644 index 00000000..2736df09 --- /dev/null +++ b/tests/unit/alignment/test_pipeline_aliasing.py @@ -0,0 +1,130 @@ +"""Guards for two aliasing traps in the molecular-replacement pipeline. + +`Model.rotate` and `Model.translate` mutate in place and return ``self``. Every +call site that treats them as returning a fresh model therefore has to +``.copy()`` first. `MolecularReplacementPipeline._make_rotated` is called once +per rotation candidate off the same ``self.model``, so without the copy +candidate *k+1* is evaluated at an orientation composed on top of candidate *k* +and ``self.model`` is destroyed along the way. + +`Model.spacegroup` is a property, but ``SpaceGroup`` is an ``nn.Module``: +assigning a SpaceGroup *object* is intercepted by ``nn.Module.__setattr__``, +stored in ``_modules`` under the property's own name, and the setter never runs. +The pipeline builds P1 copies for its dense-transform stages, so a silent no-op +there means the "P1 search model" still carries the crystal symmetry. + +Both are pinned here rather than only in the slow integration tests, which +``--run-slow`` gates off by default. +""" + +import math + +import pytest +import torch + +pytestmark = pytest.mark.unit + + +@pytest.fixture(scope="module") +def small_model(pdb_dir): + from torchref.model import ModelFT + + p = pdb_dir / "1DAW.pdb" + if not p.exists(): + pytest.skip("1DAW.pdb not available") + return ModelFT(verbose=0).load_pdb(str(p)) + + +def _peak(alpha, beta, gamma): + from torchref.experimental.alignment.frf.types import RotationPeak + + return RotationPeak(alpha=alpha, beta=beta, gamma=gamma, score=1.0, sigma=1.0) + + +def test_rotate_mutates_in_place_and_returns_self(small_model): + """The premise. If this ever changes, the copies below stop being necessary.""" + m = small_model.copy() + before = m.xyz().clone() + R = torch.tensor( + [[0.0, -1.0, 0.0], [1.0, 0.0, 0.0], [0.0, 0.0, 1.0]], dtype=m.dtype_float, + ) + out = m.rotate(R) + assert out is m, "rotate no longer returns self" + assert not torch.allclose(m.xyz(), before), "rotate no longer mutates in place" + + +def test_make_rotated_leaves_the_search_model_untouched(small_model): + """The pipeline's own model must survive candidate generation.""" + from torchref.experimental.alignment.pipeline import MolecularReplacementPipeline + + pipe = object.__new__(MolecularReplacementPipeline) + pipe.model = small_model + reference = small_model.xyz().clone() + + rotated, _ = pipe._make_rotated(_peak(0.3, 0.7, 1.1)) + + assert rotated is not pipe.model + assert torch.allclose(pipe.model.xyz(), reference), ( + "_make_rotated mutated the pipeline's search model" + ) + + +def test_successive_candidates_do_not_compound(small_model): + """Candidate k+1 must not be rotated on top of candidate k.""" + from torchref.experimental.alignment.pipeline import MolecularReplacementPipeline + + pipe = object.__new__(MolecularReplacementPipeline) + pipe.model = small_model + + p1 = _peak(0.3, 0.7, 1.1) + p2 = _peak(2.0, 1.3, 0.4) + + first, _ = pipe._make_rotated(p1) + second, _ = pipe._make_rotated(p2) + + # A fresh pipeline that only ever sees p2 is the ground truth for p2. + solo = object.__new__(MolecularReplacementPipeline) + solo.model = small_model.copy() + expected, _ = solo._make_rotated(p2) + + assert torch.allclose(second.xyz(), expected.xyz(), atol=1e-5), ( + "the second candidate depends on the first -- rotations are compounding" + ) + assert not torch.allclose(first.xyz(), second.xyz()), ( + "the two candidates are identical; the peaks chosen do not discriminate" + ) + + +def test_spacegroup_name_assignment_works(small_model): + """The supported form: pass the space-group NAME.""" + m = small_model.copy() + m.spacegroup = "P 1" + assert m.spacegroup.number == 1 + assert int(m.spacegroup.matrices.shape[0]) == 1 + + +def test_spacegroup_object_assignment_bypasses_the_setter_and_poisons_the_name( + small_model, +): + """Pin the trap in full: the object form does not just fail to take effect. + + ``nn.Module.__setattr__`` files the SpaceGroup under ``_modules["spacegroup"]`` + without running the property setter, so the space group is unchanged. Worse, + the name is now a registered child module, so the *correct* string assignment + afterwards raises ``TypeError`` instead of working. + """ + from torchref.symmetry import SpaceGroup + + m = small_model.copy() + original = str(m.spacegroup) + assert m.spacegroup.number != 1, "1DAW should not already be P1" + + m.spacegroup = SpaceGroup("P 1") + assert str(m.spacegroup) == original, ( + "object assignment now reaches the property setter -- the explicit name " + "assignments in pipeline.py can be simplified" + ) + assert "spacegroup" in m._modules + + with pytest.raises(TypeError, match="child module"): + m.spacegroup = "P 1" diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 385524f8..a337e6b0 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -429,7 +429,7 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: continue r_analytic, t_refined = placement - refined = rotated_k.translate( + refined = rotated_k.copy().translate( t_refined.to(self.model.dtype_float), fractional=True, ) r_rank = r_analytic @@ -534,14 +534,6 @@ def _rotation_candidates(self, frf) -> list: "expected 'm_letf1' (default), 'sim' or 'none'." ) - # No ML rescore: rank candidates by the raw FRF score and let the - # multi-candidate tree (FTF + refine + R-ranking) do the discrimination. - if self.rescore_engine == "none": - if self.verbose > 0: - print("fit_to_data: ML rescore DISABLED — using raw FRF peak " - "ranking (RFZ).", flush=True) - return sorted(peaks, key=lambda p: p.score, reverse=True) - F_obs = frf.F_obs hkl = frf.hkl s_mag = frf.s_mag @@ -549,6 +541,21 @@ def _rotation_candidates(self, frf) -> list: ll = frf.ll rescore_n_shells = max(self.n_shells // 2, 8) + # No ML rescore: rank candidates by the raw FRF score and let the + # multi-candidate tree (FTF + refine + R-ranking) do the discrimination. + # Sub-peak refinement is still available here: it sharpens each + # orientation in place and does not reorder, so it is independent of + # which engine (if any) ranks the candidates. + if self.rescore_engine == "none": + if self.verbose > 0: + print("fit_to_data: ML rescore DISABLED — using raw FRF peak " + "ranking (RFZ).", flush=True) + ranked = sorted(peaks, key=lambda p: p.score, reverse=True) + if self.subpeak_refine: + ranked = self._subpeak_refine(ranked, F_obs, hkl, s_mag, + centric, ll, rescore_n_shells) + return ranked + interp_var_main: Optional[torch.Tensor] = None if self.use_interp_var: rescore_edges, _ = equal_count_shell_edges(s_mag, rescore_n_shells) @@ -629,7 +636,9 @@ def _make_rotated(self, peak: "RotationPeak"): """ R_rec = rotation_matrix_from_edmonds_euler(peak.alpha, peak.beta, peak.gamma) R_app = R_rec.T.contiguous() - rot = self.model.rotate( + # .copy() first: Model.rotate mutates in place and returns self, so + # rotating self.model directly would compound candidate k+1 onto k. + rot = self.model.copy().rotate( R_app.to(device=self.model.device, dtype=self.model.dtype_float), ) rot.last_alignment_rotation = R_rec @@ -666,9 +675,12 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: eye3 = self._eye3 if str(rotated_k.spacegroup) != str(data.spacegroup): - rotated_k.spacegroup = data.spacegroup + # NOTE: assign the space-group NAME, not a SpaceGroup object. SpaceGroup is an + # nn.Module, so nn.Module.__setattr__ intercepts object assignment, stores it in + # _modules and never runs the property setter -- a silent no-op. + rotated_k.spacegroup = data.spacegroup.hm rotated_p1 = rotated_k.copy() - rotated_p1.spacegroup = SpaceGroup("P 1") + rotated_p1.spacegroup = "P 1" evaluator = _DirectModelEvaluator(rotated_p1) timer.start("5_precompute_G") @@ -816,7 +828,7 @@ def _dense_rotation_refine(self, refined): timer.start("8_dense_R_ll_build") refined_p1 = refined.copy() - refined_p1.spacegroup = SpaceGroup("P 1") + refined_p1.spacegroup = "P 1" ll_refine = LattmanLoveInterpolator( refined_p1, padding_factor=self.ll_padding_factor, max_res_A=self.ll_max_res_A, verbose=0, @@ -904,7 +916,7 @@ def _dense_rotation_refine(self, refined): flush=True, ) - return refined.rotate( + return refined.copy().rotate( R_accumulated.T.to(self.model.dtype_float).contiguous(), ) @@ -935,7 +947,9 @@ def _rigid_body_polish(self, refined): with torch.no_grad(): R_polish = rb.get_rotation_matrix().detach() t_polish = rb.translation_frac.detach() - polished = refined.rotate(R_polish.to(self.model.dtype_float)) + # .copy() first: the caller keeps `refined` when the polish does not + # improve R, so `polished` must not be the same object. + polished = refined.copy().rotate(R_polish.to(self.model.dtype_float)) polished = polished.translate( t_polish.to(self.model.dtype_float), fractional=True, ) From 4ec0311487ecdb77f7d1efe4812e5c637cf1e5c3 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 11:35:07 +0200 Subject: [PATCH 011/250] Bring the alignment package back onto the current core APIs The package was forked before several core changes and never caught up: - `Scaler.rfactor` was removed. `align._external_rwork` now calls `rfactor_work_free(data, abs(scaler.forward(fcalc)))` -- it takes scaled amplitudes, not complex F_calc -- and `RigidBodyRefinement` uses `XrayTarget.get_rfactor`, which scales through the target and so reports R against the same work/free partition the loss uses. - `SfFFT` no longer takes `radius_angstrom`; the splat radius is the global `torchref.sigma_cutoff_ed`, in sigmas. Dropped from `LattmanLoveInterpolator.__init__`, which no caller passed. - `torchref.alignment.ml_rotation` does not exist; the lazy import in `local_rotation_translation_refine` is now relative. - `rigid_body.py` imported three names twice and two it does not use, one of which (`MaximumLikelihoodXrayTarget`) no longer exists and made the whole package unimportable. `ModelFT.fit_to_data` does not exist either, so the merged integration tests and drivers now call `align_model_to_data` directly. One of them was comparing the aligned result against a reference that its own in-place `rotate` had already moved. Also drops the five `SpaceGroup` imports left dead by the switch to name-string assignment. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- tests/integration/alignment/profile_fit.py | 7 ++++--- tests/integration/alignment/run_phaser_mr.py | 3 ++- tests/integration/alignment/run_random_pdb_fit.py | 7 +++++-- tests/integration/alignment/test_fit_to_data.py | 7 ++++--- .../alignment/test_fit_to_data_translation.py | 14 +++++++++----- tests/unit/alignment/test_patterson_translation.py | 1 - torchref/experimental/alignment/align.py | 4 +++- torchref/experimental/alignment/lattman_love.py | 4 ---- torchref/experimental/alignment/pipeline.py | 1 - torchref/experimental/alignment/rigid_body.py | 14 ++++---------- torchref/experimental/alignment/translation.py | 2 +- 11 files changed, 32 insertions(+), 32 deletions(-) diff --git a/tests/integration/alignment/profile_fit.py b/tests/integration/alignment/profile_fit.py index 54ffab7c..6c4bf6f6 100644 --- a/tests/integration/alignment/profile_fit.py +++ b/tests/integration/alignment/profile_fit.py @@ -22,10 +22,10 @@ import torch +from torchref.experimental.alignment.align import align_model_to_data from torchref.experimental.alignment.frf.rotation_utils import rotation_matrix_from_edmonds_euler from torchref.io.datasets.reflection_data import ReflectionData from torchref.model import ModelFT -from torchref.symmetry import SpaceGroup TEST_FILES = Path("/das/work/p17/p17490/Peter/Library/work_trees_torchref/fix_alignment/tests/files") @@ -91,7 +91,7 @@ def main(): data = ReflectionData().load_mtz(str(mtz_path)) canonical = ModelFT().load_pdb(str(pdb_path)) - canonical.spacegroup = SpaceGroup("P 1") + canonical.spacegroup = "P 1" R_true = rotation_matrix_from_edmonds_euler(0.6, 0.4, 1.2) rotated_p = canonical.rotate( R_true.to(canonical.dtype_float), center=canonical.xyz().mean(dim=0), @@ -106,7 +106,8 @@ def main(): profiler = cProfile.Profile() t0 = time.time() profiler.enable() - aligned = perturbed.fit_to_data( + aligned = align_model_to_data( + perturbed, data, n_rotation_candidates=args.n_rotation_candidates, n_translation_candidates=args.n_translation_candidates, diff --git a/tests/integration/alignment/run_phaser_mr.py b/tests/integration/alignment/run_phaser_mr.py index d9b4c395..d42eb79b 100644 --- a/tests/integration/alignment/run_phaser_mr.py +++ b/tests/integration/alignment/run_phaser_mr.py @@ -40,6 +40,7 @@ import torch +from torchref.base.metrics.rfactor import rfactor_work_free from torchref.experimental.alignment.frf.rotation_utils import rotation_angular_distance_deg from torchref.io.datasets.reflection_data import ReflectionData from torchref.model import ModelFT @@ -299,7 +300,7 @@ def run(pdb_key: str, seed: int, *, work_root: Path, scaler.initialize(fcalc) scaler.refine_lbfgs(fcalc=fcalc) with torch.no_grad(): - rw, rf = scaler.rfactor(fcalc) + rw, rf = rfactor_work_free(data, torch.abs(scaler.forward(fcalc))) rwork_torchref_scaler = ( rw.item() if hasattr(rw, "item") else float(rw) ) diff --git a/tests/integration/alignment/run_random_pdb_fit.py b/tests/integration/alignment/run_random_pdb_fit.py index d7d7b208..ded16b08 100644 --- a/tests/integration/alignment/run_random_pdb_fit.py +++ b/tests/integration/alignment/run_random_pdb_fit.py @@ -27,6 +27,8 @@ import torch +from torchref.base.metrics.rfactor import rfactor_work_free +from torchref.experimental.alignment.align import align_model_to_data from torchref.experimental.alignment.frf.rotation_utils import rotation_angular_distance_deg from torchref.io.datasets.reflection_data import ReflectionData from torchref.model import ModelFT @@ -131,7 +133,7 @@ def _scale_and_r(m: ModelFT) -> tuple[float, float]: s.initialize(fcalc) s.refine_lbfgs(fcalc=fcalc) with torch.no_grad(): - rw, rf = s.rfactor(fcalc) + rw, rf = rfactor_work_free(data, torch.abs(s.forward(fcalc))) rw = rw.item() if hasattr(rw, "item") else float(rw) rf = rf.item() if hasattr(rf, "item") else float(rf) return rw, rf @@ -156,7 +158,8 @@ def _scale_and_r(m: ModelFT) -> tuple[float, float]: # 3. Run fit_to_data: recover the alignment. t1 = time.time() - aligned = rotated_search.fit_to_data( + aligned = align_model_to_data( + rotated_search, data, d_min=4.0, d_max=15.0, L=32, n_shells=20, diff --git a/tests/integration/alignment/test_fit_to_data.py b/tests/integration/alignment/test_fit_to_data.py index 664bc57e..a22c544f 100644 --- a/tests/integration/alignment/test_fit_to_data.py +++ b/tests/integration/alignment/test_fit_to_data.py @@ -17,10 +17,10 @@ import pytest import torch +from torchref.experimental.alignment.align import align_model_to_data from torchref.experimental.alignment.frf.rotation_utils import rotation_angular_distance_deg from torchref.io.datasets.reflection_data import ReflectionData from torchref.model import ModelFT -from torchref.symmetry import SpaceGroup TEST_FILES = Path(__file__).resolve().parents[2] / "files" @@ -31,7 +31,7 @@ def _load_p1_search_model() -> ModelFT: """Load 1DAW and force spacegroup to P1 via the proper setter.""" m = ModelFT().load_pdb(str(PDB_1DAW)) - m.spacegroup = SpaceGroup("P 1") + m.spacegroup = "P 1" return m @@ -86,7 +86,8 @@ def test_fit_to_data_real_1daw(real_setup, trial): R_true = _random_rotation(seed=5000 + trial) search = canonical.rotate(R_true.to(canonical.dtype_float), center=centroid) - aligned = search.fit_to_data( + aligned = align_model_to_data( + search, data, d_min=4.0, d_max=15.0, L=32, n_shells=20, diff --git a/tests/integration/alignment/test_fit_to_data_translation.py b/tests/integration/alignment/test_fit_to_data_translation.py index 63188fd3..60f77a8c 100644 --- a/tests/integration/alignment/test_fit_to_data_translation.py +++ b/tests/integration/alignment/test_fit_to_data_translation.py @@ -15,6 +15,8 @@ import pytest import torch +from torchref.base.metrics.rfactor import rfactor_work_free +from torchref.experimental.alignment.align import align_model_to_data from torchref.experimental.alignment.frf.rotation_utils import ( rotation_angular_distance_deg, rotation_matrix_from_edmonds_euler, @@ -22,7 +24,6 @@ from torchref.io.datasets.reflection_data import ReflectionData from torchref.model import ModelFT from torchref.scaling import Scaler -from torchref.symmetry import SpaceGroup TEST_FILES = Path(__file__).resolve().parents[2] / "files" @@ -35,7 +36,7 @@ def _scale_and_rwork(model: ModelFT, data) -> float: fcalc = model(data.hkl) scaler.initialize(fcalc) scaler.refine_lbfgs(fcalc=fcalc) - rw, _ = scaler.rfactor(fcalc) + rw, _ = rfactor_work_free(data, torch.abs(scaler.forward(fcalc))) return rw.item() if hasattr(rw, "item") else float(rw) @@ -49,18 +50,21 @@ def _wrap_frac(t: torch.Tensor) -> torch.Tensor: def test_fit_to_data_recovers_rotation_and_translation(): data = ReflectionData().load_mtz(str(MTZ_1DAW)) canonical = ModelFT().load_pdb(str(PDB_1DAW)) - canonical.spacegroup = SpaceGroup("P 1") + canonical.spacegroup = "P 1" # Apply a known random rotation + fractional translation. R_true = rotation_matrix_from_edmonds_euler(0.6, 0.4, 1.2) R_apply = R_true.to(canonical.dtype_float) - rotated = canonical.rotate(R_apply, center=canonical.xyz().mean(dim=0)) + # .copy() first: Model.rotate mutates in place, so rotating `canonical` + # directly would perturb the very reference this test compares against. + rotated = canonical.copy().rotate(R_apply, center=canonical.xyz().mean(dim=0)) t_frac_true = torch.tensor([0.18, -0.07, 0.23], dtype=canonical.dtype_float) perturbed = rotated.translate(t_frac_true, fractional=True) rwork_pre = _scale_and_rwork(perturbed, data) - aligned = perturbed.fit_to_data( + aligned = align_model_to_data( + perturbed, data, d_min=4.0, d_max=15.0, L=32, n_shells=20, diff --git a/tests/unit/alignment/test_patterson_translation.py b/tests/unit/alignment/test_patterson_translation.py index 307b1ea1..2870659a 100644 --- a/tests/unit/alignment/test_patterson_translation.py +++ b/tests/unit/alignment/test_patterson_translation.py @@ -17,7 +17,6 @@ from torchref.experimental.alignment.translation import amplitude_translation_search from torchref.io.datasets.reflection_data import ReflectionData from torchref.model import ModelFT -from torchref.symmetry import SpaceGroup TEST_FILES = Path(__file__).resolve().parents[2] / "files" diff --git a/torchref/experimental/alignment/align.py b/torchref/experimental/alignment/align.py index 5234bba6..52d3800e 100644 --- a/torchref/experimental/alignment/align.py +++ b/torchref/experimental/alignment/align.py @@ -139,6 +139,7 @@ def _external_rwork(model: "ModelFT", data: "ReflectionData") -> float: candidates correctly but isn't the user-facing R-work). We compute the proper Scaler-fit R-work once per finalist. """ + from ...base.metrics.rfactor import rfactor_work_free from ...scaling import Scaler # Build the Scaler on the model's device so that its anisotropy U @@ -156,7 +157,8 @@ def _external_rwork(model: "ModelFT", data: "ReflectionData") -> float: s.initialize(fc) s.refine_lbfgs(fcalc=fc) with torch.no_grad(): - rw, _ = s.rfactor(fc) + # rfactor_work_free takes already-scaled amplitudes, not complex F_calc. + rw, _ = rfactor_work_free(data, torch.abs(s.forward(fc))) return rw.item() if hasattr(rw, "item") else float(rw) diff --git a/torchref/experimental/alignment/lattman_love.py b/torchref/experimental/alignment/lattman_love.py index 342e823f..254d1827 100644 --- a/torchref/experimental/alignment/lattman_love.py +++ b/torchref/experimental/alignment/lattman_love.py @@ -59,8 +59,6 @@ class LattmanLoveInterpolator(DeviceMixin): max_res_A : float, optional Resolution limit (Å). The dense grid will resolve features down to this. Default: 2.0 Å (suitable for proteins up to that resolution). - radius_angstrom : float, default 3.0 - Atomic-density support radius for SfFFT. device : torch.device, optional Target device. Defaults to the model's device. """ @@ -71,7 +69,6 @@ def __init__( padding_factor: float = 2.0, min_cell_size_A: Optional[float] = None, max_res_A: float = 2.0, - radius_angstrom: float = 3.0, device: Optional[torch.device] = None, verbose: int = 0, ): @@ -123,7 +120,6 @@ def __init__( cell=self.cubic_cell, spacegroup=SpaceGroup("P 1"), max_res=max_res_A, - radius_angstrom=radius_angstrom, dtype_float=torch.float32, device=device, verbose=verbose, diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index a337e6b0..5b502703 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -39,7 +39,6 @@ from torchref.config import get_default_device from torchref.utils.device_mixin import DeviceMixin -from ...symmetry import SpaceGroup from .align import ( _DirectModelEvaluator, _StageTimer, diff --git a/torchref/experimental/alignment/rigid_body.py b/torchref/experimental/alignment/rigid_body.py index 8d686aab..08ada4b7 100644 --- a/torchref/experimental/alignment/rigid_body.py +++ b/torchref/experimental/alignment/rigid_body.py @@ -27,17 +27,11 @@ import torch import torch.nn as nn -from torchref.config import get_default_device -from torchref.scaling import Scaler -from torchref.model import SfFFT -from torchref.symmetry import spacegroup -from torchref.refinement.targets import RiceXrayTarget from torchref.base import rotation_matrix_euler_zyz from torchref.config import get_default_device from torchref.model import SfFFT -from torchref.refinement.targets import MaximumLikelihoodXrayTarget -from torchref.scaling import ScalerBase -from torchref.symmetry import spacegroup +from torchref.refinement.targets import RiceXrayTarget +from torchref.scaling import Scaler from torchref.utils.device_mixin import DeviceMixin @@ -453,7 +447,7 @@ def closure(): ) return current_loss - rwork_initial, rfree_initial = self.scaler.rfactor(self()) + rwork_initial, rfree_initial = self.xray_target.get_rfactor(self()) initial_loss = closure().item() if self.verbose > 0: @@ -476,7 +470,7 @@ def closure(): if self.verbose > 1: print(f"Iter {tries_needed} Current ML loss: {current_loss:.4f}") final_loss = closure().item() - final_rwork, final_rfree = self.scaler.rfactor(self()) + final_rwork, final_rfree = self.xray_target.get_rfactor(self()) converged = final_rwork < self.rfactor_converged_threshold if converged or tries_needed >= n_tries: diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index 45d59f94..9d324458 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -793,7 +793,7 @@ def local_rotation_translation_refine( llg_best : float Sim MLRF LLG at the returned (R_best, t_best). """ - from torchref.alignment.ml_rotation import llg_for_rotation_batch + from .ml_rotation import llg_for_rotation_batch device = getattr(interpolator, "device", hkl.device) real_dtype = torch.float64 complex_dtype = torch.complex128 From 11e1d0dc95fcd2ad889322cd5840e6757918c5a5 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 11:38:27 +0200 Subject: [PATCH 012/250] Correct the alignment package's stale docstring references Named modules, files and methods that do not exist: `torchref.alignment.*` (the package is `torchref.experimental.alignment`), `phaser_frf`, `ball_search.py`, `benchmark_phaser_frf.py` with its `--engine` flag, `ModelFT.fit_to_data`, and three markdown/CSV paths that were never in the repo. The `fit_to_data:` prefix also appeared in nine runtime progress messages. Also drops the "v13" post-mortem prose from the sample-list and peak-finder module docstrings: what the deleted implementation got wrong is not the behaviour of this one. Adds a note in `bessel_sh_expand` that Phaser's radial band is per-`l` (`nmax = (lmax - l + 2)/2`, DataMR.cc:894) and that `N_radial` is only the allocated width, plus tests asserting the populated support equals that band rather than merely fitting inside it. The stale references inside `_run_frf_separate_rotation` are left alone; that function is being replaced. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- .../unit/alignment/test_radial_truncation.py | 112 ++++++++++++++++++ torchref/experimental/alignment/__init__.py | 2 +- torchref/experimental/alignment/align.py | 13 +- torchref/experimental/alignment/frf/api.py | 16 +-- .../experimental/alignment/frf/data_mr.py | 14 ++- .../experimental/alignment/frf/dense_calc.py | 3 +- .../alignment/frf/french_wilson.py | 2 +- .../experimental/alignment/frf/peak_finder.py | 5 +- .../alignment/frf/preprocessing.py | 7 +- .../alignment/frf/rotation_utils.py | 2 +- .../alignment/frf/sitelist_ang.py | 11 +- torchref/experimental/alignment/frf/types.py | 2 +- .../experimental/alignment/frf/wigner_d.py | 2 +- .../experimental/alignment/ml_rotation.py | 10 +- torchref/experimental/alignment/pipeline.py | 20 ++-- torchref/experimental/alignment/rigid_body.py | 2 +- torchref/experimental/alignment/wigner.py | 2 +- 17 files changed, 164 insertions(+), 61 deletions(-) create mode 100644 tests/unit/alignment/test_radial_truncation.py diff --git a/tests/unit/alignment/test_radial_truncation.py b/tests/unit/alignment/test_radial_truncation.py new file mode 100644 index 00000000..b085e131 --- /dev/null +++ b/tests/unit/alignment/test_radial_truncation.py @@ -0,0 +1,112 @@ +"""Pin the SH-Bessel radial band to Phaser's per-``l`` size. + +Phaser allocates the ``Elmn`` array with ``nmax = (lmax - l + 2) / 2`` radial +terms for each even ``l`` and runs ``n`` from 1 to ``nmax`` +(``DataMR.cc:894-896``). The band therefore narrows as ``l`` rises -- for +``lmax = 76`` it is 38 terms at ``l = 2`` and a single term at ``l = 76`` -- so +the high-``l`` bands cannot carry more radial detail than the reflection set +supports. + +``bessel_sh_expand`` allocates a flat ``(N_radial, L, 2L-1)`` array, where +``N_radial`` is Phaser's *widest* band (the one at ``l = 2``). That shape is +easy to misread as "every ``l`` carries ``N_radial`` radial terms"; it does not. +The ``(l, n) -> u = l + 2n + 1`` index build populates only +``n_l = (lmax_even - l)//2 + 1`` terms per ``l``, which is exactly Phaser's +``nmax``, so the truncation is already in force through the allocated support. + +These tests pin that invariant against the formula, so the agreement is +asserted rather than inferred from the array shape. +""" + +import pytest +import torch + +from torchref.experimental.alignment.frf.data_mr import bessel_sh_expand + + +def _expand(L: int, n_points: int = 900): + """Expand a fixed pseudo-random point set; returns ``(N_radial, L, 2L-1)``.""" + g = torch.Generator().manual_seed(5) + s = torch.randn(n_points, 3, generator=g, dtype=torch.float64) + s = s / s.norm(dim=-1, keepdim=True) * ( + 0.05 + 0.15 * torch.rand(n_points, 1, generator=g, dtype=torch.float64) + ) + intensity = torch.randn(n_points, generator=g, dtype=torch.float64) + out = bessel_sh_expand(s, intensity, L=L, bessel_h_scale=30.0) + return out.coeffs.detach().cpu() + + +def _phaser_nmax(lmax_even: int, l: int) -> int: + """``(lmax - l + 2) / 2`` -- Phaser's radial band size for one ``l``.""" + return (lmax_even - l + 2) // 2 + + +@pytest.mark.parametrize("L", [21, 41, 67]) +def test_radial_band_matches_phaser_width(L): + """Non-zero radial indices at each ``l`` must stop at Phaser's ``nmax``. + + Our ``n`` index is 0-based against Phaser's 1-based, so the condition is + ``n < nmax(l)``. + """ + c = _expand(L) + N_radial, _, _ = c.shape + lmax_even = L - 1 if (L - 1) % 2 == 0 else L - 2 + + for l in range(2, lmax_even + 1, 2): + nz = (c[:, l, :].abs() > 0).any(dim=-1).nonzero().flatten() + if nz.numel() == 0: + continue # legitimately empty band + allowed = _phaser_nmax(lmax_even, l) + assert int(nz.max()) < allowed, ( + f"L={L} l={l}: radial index {int(nz.max())} exceeds Phaser's " + f"nmax={allowed}" + ) + + +@pytest.mark.parametrize("L", [21, 41, 67]) +def test_band_narrows_to_a_single_term_at_lmax(L): + """The top band keeps exactly one radial term, the widest keeps them all.""" + lmax_even = L - 1 if (L - 1) % 2 == 0 else L - 2 + assert _phaser_nmax(lmax_even, lmax_even) == 1 + c = _expand(L) + assert _phaser_nmax(lmax_even, 2) == c.shape[0], ( + "N_radial should equal Phaser's widest band, at l=2" + ) + top = (c[:, lmax_even, :].abs() > 0).any(dim=-1).nonzero().flatten() + if top.numel(): + assert int(top.max()) == 0, "l=lmax must retain only the n=0 term" + + +@pytest.mark.parametrize("L", [21, 41, 67]) +def test_band_is_exactly_phasers_width_not_merely_bounded(L): + """The populated band must *equal* Phaser's ``nmax(l)``, not just fit inside. + + A bound alone would also pass for an expansion that silently drops radial + terms it should keep, which would cost resolution at low ``l``. + """ + c = _expand(L) + lmax_even = L - 1 if (L - 1) % 2 == 0 else L - 2 + for l in range(2, lmax_even + 1, 2): + nz = (c[:, l, :].abs() > 0).any(dim=-1).nonzero().flatten() + assert nz.numel(), f"L={L} l={l}: band is empty" + assert int(nz.max()) + 1 == _phaser_nmax(lmax_even, l), ( + f"L={L} l={l}: {int(nz.max()) + 1} radial terms, " + f"Phaser has {_phaser_nmax(lmax_even, l)}" + ) + + +def test_allocated_width_exceeds_the_populated_band(): + """The array is wider than the support at every l above the first. + + This is the fact that makes the array shape misleading, and the reason the + tests above assert the support rather than ``coeffs.shape``. + """ + L = 67 + lmax_even = L - 1 if (L - 1) % 2 == 0 else L - 2 + c = _expand(L) + N_radial = c.shape[0] + assert N_radial == _phaser_nmax(lmax_even, 2) + assert _phaser_nmax(lmax_even, lmax_even) == 1 + assert N_radial > _phaser_nmax(lmax_even, lmax_even), ( + "allocated width should exceed the top band's single term" + ) diff --git a/torchref/experimental/alignment/__init__.py b/torchref/experimental/alignment/__init__.py index cebeda16..3263dde3 100644 --- a/torchref/experimental/alignment/__init__.py +++ b/torchref/experimental/alignment/__init__.py @@ -14,7 +14,7 @@ 4. Canonical Pipeline (``pipeline.MolecularReplacementPipeline``) — the multi-candidate FRF → FTF → post-refine tree with early-stopping; the implementation that ``align.align_model_to_data`` / - ``ModelFT.fit_to_data`` delegate to. + ``align.align_model_to_data`` delegates to. Example — full MR pipeline -------------------------- diff --git a/torchref/experimental/alignment/align.py b/torchref/experimental/alignment/align.py index 52d3800e..e5165b83 100644 --- a/torchref/experimental/alignment/align.py +++ b/torchref/experimental/alignment/align.py @@ -9,9 +9,8 @@ timer (`_StageTimer`) — that are shared by the rotation-ranking benchmarks and by the orchestrator. -`align_model_to_data` is the public entry point that `ModelFT.fit_to_data` -delegates to; it in turn delegates the FRF → FTF(per-candidate) → post-refine -control flow to +`align_model_to_data` is the public entry point. It delegates the +FRF → FTF(per-candidate) → post-refine control flow to :class:`torchref.experimental.alignment.pipeline.MolecularReplacementPipeline`, returning that pipeline's single best `ModelFT`. """ @@ -649,8 +648,8 @@ def align_model_to_data( ``last_alignment_rotation``, ``last_alignment_translation`` and ``last_alignment_rfactor`` provenance attributes. - See `ModelFT.fit_to_data` for full kwarg semantics — this function is the - canonical implementation; `fit_to_data` is a thin wrapper. + `MolecularReplacementPipeline` is the implementation of record; this + function returns its single best solution. """ if not model.initialized: raise RuntimeError( @@ -658,8 +657,8 @@ def align_model_to_data( ) # `MolecularReplacementPipeline` is the implementation of record. This - # function preserves the historical kwarg surface (so `ModelFT.fit_to_data` - # and the benchmark scripts keep working unchanged) and returns the single + # function preserves the historical kwarg surface (so the benchmark + # scripts keep working unchanged) and returns the single # best `ModelFT`; drive the pipeline directly to get the ranked candidate # list. Imported lazily to avoid an import cycle — `pipeline` imports the # stage helpers (`_prepare_frf_inputs`, `_run_frf_separate_rotation`, diff --git a/torchref/experimental/alignment/frf/api.py b/torchref/experimental/alignment/frf/api.py index bdc24e0f..7680050c 100644 --- a/torchref/experimental/alignment/frf/api.py +++ b/torchref/experimental/alignment/frf/api.py @@ -1,8 +1,4 @@ -"""Top-level FastRotationFunction class + drop-in ``phaser_rotation_search``. - -Signature matches ``torchref.alignment.phaser_frf.phaser_rotation_search`` -so ``tests/integration/alignment/benchmark_phaser_frf.py`` can swap -implementations via a single ``--engine`` flag. +"""Top-level ``FastRotationFunction`` class + the ``phaser_rotation_search`` wrapper. Pipeline (mirrors Phaser ``run_FRF()``): 1. Resolution mask (both sides). @@ -417,13 +413,11 @@ def phaser_rotation_search( solvent_bsol: float = 300.0, compute_dtype: Optional[torch.dtype] = None, ) -> Tuple[AdaptiveRotationFunction, List[RotationPeak]]: - """Drop-in for ``torchref.alignment.phaser_frf.phaser_rotation_search``. + """Construct a :class:`FastRotationFunction` and score one model. - Same signature, same return-shape. Sub-voxel refinement parameters - are accepted for signature parity but currently not implemented in - the frf_separate path — the per-β fixed-shape FFT already provides - sub-voxel precision via the bilinear interpolation, and Phaser - itself does not run an extra quadratic refinement. + ``refine_subvoxel`` and ``n_refine`` are accepted and ignored: the per-β + fixed-shape FFT already provides sub-voxel precision through its bilinear + interpolation, and Phaser runs no extra quadratic refinement either. Extra (non-legacy) kwargs: hkl_obs : integer Miller indices aligned with s_obs, needed for ε(h). diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index 4e88b40a..7c9e76e7 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -102,10 +102,9 @@ def bessel_sh_expand( ) -> BesselSHCoefficients: """Phaser-style ``c_nlm = Σ_h Y*_lm(ŝ) · I · sqrt(2u+1) · j_u(h)/h``. - Memory-bounded chunked reimplementation of - ``torchref.alignment.phaser_frf.bessel_sh_expand`` (identical math, - verified element-wise by ``tests/unit/frf_separate``). The legacy - version materialises the full ``(M, L, N_radial)`` Bessel table and + Memory-bounded and chunked, verified element-wise by + ``tests/unit/frf_separate``. A direct implementation materialises the full + ``(M, L, N_radial)`` Bessel table and ``(M, u_max+1)`` j-table for *all* reflections at once — at L≈100 with a symmetry-unrolled obs set (≳10⁶ reflections) that is tens of GB and OOMs. Here the j-table, Bessel weights and Y_lm are all computed @@ -174,6 +173,11 @@ def bessel_sh_expand( even_ls = list(range(2, lmax_even + 1, 2)) l_list, n_list, u_list, w_list = [], [], [], [] for l in even_ls: + # Phaser's per-l radial band: nmax = (lmax - l + 2)/2 (DataMR.cc:894), + # narrowing from N_radial terms at l=2 to a single term at l=lmax, so + # the high-l bands cannot carry more radial detail than the reflection + # set supports. `N_radial` above is only the allocated width (Phaser's + # widest band); the populated support is this per-l count. n_l = (lmax_even - l) // 2 + 1 for n in range(n_l): u = l + 2 * n + 1 @@ -320,7 +324,7 @@ def cross_correlate_xi( happens inside ``DataMR::dataMR_FRF`` before being fed into ``SiteListAng::DoRfftStuff`` as the ``clmn`` tensor (FastRot.cc:39). - Convention (matches torchref's existing ball_search.py:182): + Convention: xi[l, m, n] = Σ_r c_obs[r, l, n] · conj(c_calc[r, l, m]) so that the peak Euler triple satisfies ``s_calc = R · s_obs``. diff --git a/torchref/experimental/alignment/frf/dense_calc.py b/torchref/experimental/alignment/frf/dense_calc.py index afa85324..71928c4c 100644 --- a/torchref/experimental/alignment/frf/dense_calc.py +++ b/torchref/experimental/alignment/frf/dense_calc.py @@ -1,7 +1,6 @@ """Dense P1-box sampling of a model's molecular transform. -This is the "dense calc" lever from the high-symmetry FRF investigation -(see ``FRF_CONSOLIDATION.md``): the Fast Rotation Function correlates the obs +The Fast Rotation Function correlates the obs Patterson against the *model* transform, and sampling that transform at the sparse crystal lattice under-determines the high-l spherical-harmonic modes for large molecules. Phaser avoids this by computing the model transform on a dense, diff --git a/torchref/experimental/alignment/frf/french_wilson.py b/torchref/experimental/alignment/frf/french_wilson.py index 4dcaa3b3..9d02ee6e 100644 --- a/torchref/experimental/alignment/frf/french_wilson.py +++ b/torchref/experimental/alignment/frf/french_wilson.py @@ -7,7 +7,7 @@ from raw ``(F, σF, |s|, centric)``. Everything except ``french_wilson_preprocess`` is module-private; expose -the public name through :mod:`torchref.alignment.frf.preprocessing`. +the public name through :mod:`torchref.experimental.alignment.frf.preprocessing`. References (paths under ``…/reverse_engineering/phenix/.../phaser/src/``): diff --git a/torchref/experimental/alignment/frf/peak_finder.py b/torchref/experimental/alignment/frf/peak_finder.py index 660d7e99..dce37d7d 100644 --- a/torchref/experimental/alignment/frf/peak_finder.py +++ b/torchref/experimental/alignment/frf/peak_finder.py @@ -10,10 +10,7 @@ rotations (not by α, β, γ box distance — that would double-count near the poles). -We implement the same flow in PyTorch, vectorised where possible. The -SO(3) angular-distance NMS is identical to the v13 ``_so3_greedy_nms`` -in ``ball_search.py`` — that part of v13 was correct; the bug was in -the FFT/grid, not the NMS. +We implement the same flow in PyTorch, vectorised where possible. """ from __future__ import annotations diff --git a/torchref/experimental/alignment/frf/preprocessing.py b/torchref/experimental/alignment/frf/preprocessing.py index ffc8a62d..bcfb4434 100644 --- a/torchref/experimental/alignment/frf/preprocessing.py +++ b/torchref/experimental/alignment/frf/preprocessing.py @@ -2,10 +2,9 @@ Mirrors the chain in Phaser ``DataMR::dataMR_FRF`` (DataMR.cc:863-1133) and the auxiliary helpers in ``lib/math_FrenchWilson.cc`` and -``lib/math_RiceLLG.cc``. The PyTorch ports of these algorithms already -live in ``torchref.alignment.phaser_frf``; rather than duplicate the -~600 LoC implementations here, we import them and re-document the -Phaser source citations. +``lib/math_RiceLLG.cc``. The heavier ports live in sibling modules +(:mod:`~torchref.experimental.alignment.frf.french_wilson` in particular); +this module imports them and carries the Phaser source citations. If a specific preprocessing piece turns out to be wrong (per Tier 2 synthetic tests), the fix lives here — replace the import with a fresh diff --git a/torchref/experimental/alignment/frf/rotation_utils.py b/torchref/experimental/alignment/frf/rotation_utils.py index f4e9d4b2..ead4c86c 100644 --- a/torchref/experimental/alignment/frf/rotation_utils.py +++ b/torchref/experimental/alignment/frf/rotation_utils.py @@ -17,7 +17,7 @@ def rotation_matrix_from_edmonds_euler( """Build ``R = R_z(α) R_y(β) R_z(γ)`` (Edmonds active ZYZ). Equivalent to passing ``[γ, β, α]`` to - ``torchref.alignment.transform.rotation_matrix_from_euler``. + ``torchref.experimental.alignment.transform.rotation_matrix_from_euler``. """ ca, sa = math.cos(alpha), math.sin(alpha) cb, sb = math.cos(beta), math.sin(beta) diff --git a/torchref/experimental/alignment/frf/sitelist_ang.py b/torchref/experimental/alignment/frf/sitelist_ang.py index 50567d94..5a1defc9 100644 --- a/torchref/experimental/alignment/frf/sitelist_ang.py +++ b/torchref/experimental/alignment/frf/sitelist_ang.py @@ -3,8 +3,8 @@ Mirrors ``SiteListAng`` from ``reverse_engineering/phenix/phenix-1.20-4459/modules/phaser/codebase/phaser/src/FastRot.cc``. -The crucial point — and the bug that broke v13 of the legacy adaptive -grid — is that **the FFT itself is NOT per-β-adaptive**. Phaser does: +The crucial point is that **the FFT itself is NOT per-β-adaptive**. +Phaser does: 1. ``get_FRF`` (FastRot.cc:90-167) loops over a uniform β grid ``β_b = b · Δ`` for ``b ∈ [0, bmax)``, ``bmax = ceil(180/Δ)``. @@ -29,10 +29,9 @@ 5. ``M_β`` is **bilinearly interpolated** at each ``(α, γ)`` sample point to give the RF value (FastRot.cc:146-152, ``four_point_interpolation``). -v13's mistake was making the FFT shape itself ``(pmax(β), qmax(β))`` — -which collapses to ``(N, 1)`` at small β and loses all γ Fourier -information. Phaser keeps the FFT dense; adaptivity is only in the -sample list and the interpolation. +Making the FFT shape itself ``(pmax(β), qmax(β))`` would collapse to +``(N, 1)`` at small β and lose all γ Fourier information. The FFT stays +dense; adaptivity is only in the sample list and the interpolation. """ from __future__ import annotations diff --git a/torchref/experimental/alignment/frf/types.py b/torchref/experimental/alignment/frf/types.py index 985baca3..678da364 100644 --- a/torchref/experimental/alignment/frf/types.py +++ b/torchref/experimental/alignment/frf/types.py @@ -1,4 +1,4 @@ -"""Dataclasses shared across frf_separate. +"""Dataclasses shared across the fast rotation function. Mirrors the small "data carrier" structs in Phaser (phenix-1.20-4459/modules/phaser/codebase/phaser/src/SiteListAng.h, diff --git a/torchref/experimental/alignment/frf/wigner_d.py b/torchref/experimental/alignment/frf/wigner_d.py index e6c45a5b..7dfc6973 100644 --- a/torchref/experimental/alignment/frf/wigner_d.py +++ b/torchref/experimental/alignment/frf/wigner_d.py @@ -8,7 +8,7 @@ output by the convention tests in ``tests/unit/alignment/test_wigner.py``. To avoid duplicating maths, this module re-exports the existing implementation from -``torchref.alignment.wigner`` (which is the same convention) and adds +``torchref.experimental.alignment.wigner`` (which is the same convention) and adds Phaser-specific helpers on top. """ from __future__ import annotations diff --git a/torchref/experimental/alignment/ml_rotation.py b/torchref/experimental/alignment/ml_rotation.py index 3c602117..fb90a751 100644 --- a/torchref/experimental/alignment/ml_rotation.py +++ b/torchref/experimental/alignment/ml_rotation.py @@ -77,7 +77,7 @@ def _normalize_to_e_epsilon( """ε-corrected Wilson E: ``E²_h = (F²_h/ε_h) / ⟨F²/ε⟩_shell``. Matches Phaser's obs E (``E = F/sqrt(ε·Σ_N)``) and the FRF's - :func:`torchref.alignment.frf.preprocessing.wilson_normalise_epsilon`. The + :func:`torchref.experimental.alignment.frf.preprocessing.wilson_normalise_epsilon`. The plain :func:`_normalize_to_e` (no ε) over-counts axial reflections (ε>1) on high-symmetry spacegroups, letting them dominate the ``-(E²+eImove)/V`` term and blind the m_LETF1 orientation discrimination. @@ -599,7 +599,7 @@ def _build_llg_context( * **Unique-orbit calc sum.** The moving-model intensity sums ``|E_calc|²`` over the **distinct** orbit mates via - :func:`torchref.alignment.frf.preprocessing.epsilon_aware_unroll` + :func:`torchref.experimental.alignment.frf.preprocessing.epsilon_aware_unroll` (Phaser's ``if(!duplicate(isym))``), not all ``n_ops`` raw mates. Summing all mates over-weights axial reflections (ε>1) by ε(h) and orientation-blinds high-symmetry spacegroups (the 4BX9/6G9X rank-360+ failure). @@ -1101,12 +1101,12 @@ def m_letf1_rescore( ``eImove(h) = ε(h)·σ_A²·(1/n_ops)·Σ_{distinct mates} |E_calc(R^T·S_k·h)|²`` summed over the **distinct** orbit mates only (Phaser's ``if(!duplicate(isym))``, DataMR.cc:1371-1404), via - :func:`torchref.alignment.frf.preprocessing.epsilon_aware_unroll` + + :func:`torchref.experimental.alignment.frf.preprocessing.epsilon_aware_unroll` + ``scatter_add``. Summing all ``n_ops`` raw mates over-weights axial reflections by ε(h) and orientation-blinds high-symmetry spacegroups. 2. **Per-reflection variance budget** ``V(h) = ε(h) − σ_A²(s)·n_mol`` from - :func:`torchref.alignment.frf.preprocessing.compute_v_budget` + :func:`torchref.experimental.alignment.frf.preprocessing.compute_v_budget` (DataMR.cc:949,1411). For cross-rotation with no fixed model. 3. **Phaser ``logRelRice`` / ``logRelWoolfson``** as the per-reflection LL @@ -1135,7 +1135,7 @@ def m_letf1_rescore( per shell. eps_factor : (N,) tensor, optional Per-reflection multiplicity ε(h). If ``None``, computed via - :func:`torchref.alignment.frf.preprocessing.compute_epsilon`. + :func:`torchref.experimental.alignment.frf.preprocessing.compute_epsilon`. n_refine, batch_size, verbose As in :func:`sim_mlrf_rescore`. """ diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 5b502703..04cdaafb 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -19,7 +19,7 @@ a candidate beats ``rfactor_converged``). The user-facing solvent-aware R-work is computed once, on the winner. -``align_model_to_data`` (and therefore ``ModelFT.fit_to_data``) delegates to +``align_model_to_data`` delegates to this class — it is the implementation of record. The heavy crystallographic stage helpers live in :mod:`torchref.experimental.alignment.align`, :mod:`~torchref.experimental.alignment.translation` and @@ -190,7 +190,7 @@ class MolecularReplacementPipeline(DeviceMixin): """Canonical MR pipeline: FRF → FTF (per candidate) → post-refine. Parameters mirror :func:`align_model_to_data` (which delegates here), so a - caller can either use ``fit_to_data`` for the common case or drive this + caller can either use ``align_model_to_data`` for the common case or drive this class directly for finer control / access to the ranked candidate list. Parameters @@ -381,7 +381,7 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: top = rescored[0] if self.verbose > 0: print( - f"fit_to_data: top peak LLG = {top.score:.2f} " + f"mr: top peak LLG = {top.score:.2f} " f"(σ_Z = {top.sigma:.2f}); applying R⁻¹ to coords.", flush=True, ) @@ -404,7 +404,7 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: max_tries = self.max_tries if self.max_tries is not None else n_rot if self.verbose > 0 and n_rot > 1: print( - f"fit_to_data: trying up to {n_rot} rotation candidates " + f"mr: trying up to {n_rot} rotation candidates " f"(early-stop after ≥{self.min_tries} once R < " f"{self.rfactor_converged}).", flush=True, @@ -473,7 +473,7 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: if n_done >= self.min_tries and best_r < self.rfactor_converged: if self.verbose > 0: print( - f"fit_to_data: converged (R {best_r:.4f} < " + f"mr: converged (R {best_r:.4f} < " f"{self.rfactor_converged}) after {n_done} candidates.", flush=True, ) @@ -495,7 +495,7 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: winner.r_factor = rwork_final if self.verbose > 0: print( - f"fit_to_data: winner analytical-TF R={winner.translation_score:.4f}, " + f"mr: winner analytical-TF R={winner.translation_score:.4f}, " f"final Scaler-fit R-work={rwork_final:.4f}", flush=True, ) @@ -515,7 +515,7 @@ def _rotation_candidates(self, frf) -> list: timer.start("3_rotation_search") if self.verbose > 0: print( - f"fit_to_data: frf_separate rotation search " + f"mr: frf_separate rotation search " f"(dense calc + auto_lmax cap={self.frf_lmax_cap}, " f"n_peaks={self.n_rotation_peaks})…", flush=True, @@ -547,7 +547,7 @@ def _rotation_candidates(self, frf) -> list: # which engine (if any) ranks the candidates. if self.rescore_engine == "none": if self.verbose > 0: - print("fit_to_data: ML rescore DISABLED — using raw FRF peak " + print("mr: ML rescore DISABLED — using raw FRF peak " "ranking (RFZ).", flush=True) ranked = sorted(peaks, key=lambda p: p.score, reverse=True) if self.subpeak_refine: @@ -566,7 +566,7 @@ def _rotation_candidates(self, frf) -> list: timer.start("4_ml_rescore") if self.verbose > 0: print( - f"fit_to_data: ML rescoring top " + f"mr: ML rescoring top " f"{min(len(peaks), self.n_ml_refine)} peaks…", flush=True, ) @@ -620,7 +620,7 @@ def _subpeak_refine(self, rescored, F_obs, hkl, s_mag, centric, ll, self._timer.stop("4b_subpeak_refine") if self.verbose > 0: print( - f"fit_to_data: sub-peak refined top {k} orientations " + f"mr: sub-peak refined top {k} orientations " f"on the ML-LLG surface (step={self.subpeak_refine_step_deg}°).", flush=True, ) diff --git a/torchref/experimental/alignment/rigid_body.py b/torchref/experimental/alignment/rigid_body.py index 08ada4b7..dcd397b8 100644 --- a/torchref/experimental/alignment/rigid_body.py +++ b/torchref/experimental/alignment/rigid_body.py @@ -205,7 +205,7 @@ def __init__( # the bulk-solvent setup. The solvent mask is computed once from # `model.xyz()` and goes stale as the joint refine moves atoms; the # mismatch then biases the LBFGS gradient. Better to leave solvent - # out of the joint refine — `fit_to_data` does a fresh + # out of the joint refine — `align_model_to_data` does a fresh # solvent-aware Scaler refit on the final polished model for the # user-facing R-work. self.scaler = Scaler(model=model, data=data, nbins=20, diff --git a/torchref/experimental/alignment/wigner.py b/torchref/experimental/alignment/wigner.py index 4e49a6b1..551ba234 100644 --- a/torchref/experimental/alignment/wigner.py +++ b/torchref/experimental/alignment/wigner.py @@ -6,7 +6,7 @@ D^l_{m,n}(α, β, γ) = e^{-i m α} · d^l_{m,n}(β) · e^{-i n γ} (Edmonds) with the Euler angles paired to the rotation matrix used by -`torchref.alignment.transform.rotation_matrix_from_euler` — i.e. ZYZ. +`torchref.experimental.alignment.transform.rotation_matrix_from_euler` — i.e. ZYZ. Small-d uses the direct sum formula (Edmonds 4.1.23) with log-factorials so the recurrence never forms `(2l)!` explicitly: From 26483b37a915498d4c1d3bcf16006e0de37a155d Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 11:40:44 +0200 Subject: [PATCH 013/250] Drop the sample-list minimum clamp Phaser does not have `build_adaptive_sample_list` clamped `pmax`/`qmax` to a minimum of 1. Phaser truncates toward zero with no clamp (FastRot.cc:214-215) and its `for (p = 0; p < pmax; p++)` body then never runs, so the beta section is genuinely empty; clamping invents a section Phaser does not sample. A section that comes out empty now contributes no samples, with `beta_starts` still carrying an entry so the slice stays representable. Behaviourally inert at every practical sampling step: the sample lists are bit-identical for grid_sampling_deg 2, 3, 4 and 5 degrees. The clamp could only bite for a step coarse enough to drive `pmax` to zero, and `qmax = 0` at beta = 0 is unused because that branch samples the alpha = gamma diagonal. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- torchref/experimental/alignment/frf/sitelist_ang.py | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) diff --git a/torchref/experimental/alignment/frf/sitelist_ang.py b/torchref/experimental/alignment/frf/sitelist_ang.py index 5a1defc9..2662558c 100644 --- a/torchref/experimental/alignment/frf/sitelist_ang.py +++ b/torchref/experimental/alignment/frf/sitelist_ang.py @@ -195,8 +195,16 @@ def build_adaptive_sample_list( beta_rad = b * grid_sampling_deg * deg2rad # plain math: no sync cosb = math.cos(beta_rad / 2.0) sinb = math.sin(beta_rad / 2.0) - pmax = max(1, int(720.0 / grid_sampling_deg * cosb)) - qmax = max(1, int(360.0 / grid_sampling_deg * sinb)) + # Truncation toward zero, with NO clamp to a minimum of 1 -- Phaser + # has none (FastRot.cc:214-215). Near beta = 180 deg, cos(beta/2) drives + # pmax to 0 and Phaser's `for (p=0; p 0 and qmax == 0): + beta_starts.append(beta_starts[-1]) + continue if b == 0: # β=0: only α = γ = p/pmax for p < pmax/2 (FastRot.cc:189-207). From 8cea183ccc176e8af41178e576c1f727b8a377fd Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 11:41:20 +0200 Subject: [PATCH 014/250] Track the alignment lab harness Fifty-odd copy-pasted scripts drove the rotation-function investigation, with the same helpers reimplemented dozens of times and two of them in incompatible variants: `random_rotation` existed with and without the `sign(diag(R))` correction, so the same seed gave different true rotations depending on which script you ran, and rank-of-truth existed in four left/right x fractional/Cartesian flavours. `alignment_lab/lab` holds one definition of each, with the conventions as explicit arguments and recorded in every output row, and `alignment_lab/tests` pins the contracts whose violation silently changed results. `analysis/aggregate.py` reports paired per-trial differences and signs rather than a bare median: seed-to-seed truth-rank spread is +/-4-6 ranks, and three findings that looked strong below ten trials did not survive the full set. Run outputs (`runs/`, 457 MB) and scheduler logs (`slurm/`) are gitignored, as is the third-party Phaser source copy used for the FRF instrumentation. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- .gitignore | 11 +- alignment_lab/README.md | 74 +++ alignment_lab/analysis/aggregate.py | 129 +++++ alignment_lab/analysis/array_template.sh | 51 ++ alignment_lab/diagnostics/bench_stages.py | 142 +++++ .../diagnostics/frf_aniso_knockout.py | 165 ++++++ .../diagnostics/frf_aniso_rank_sweep.py | 152 ++++++ .../diagnostics/frf_encode_compare.py | 511 ++++++++++++++++++ .../diagnostics/frf_ghost_knockout.py | 200 +++++++ .../diagnostics/frf_inject_phaser_obs.py | 211 ++++++++ alignment_lab/diagnostics/frf_map_compare.py | 432 +++++++++++++++ .../diagnostics/frf_normaliser_anatomy.py | 265 +++++++++ alignment_lab/diagnostics/frf_prep_compare.py | 293 ++++++++++ alignment_lab/diagnostics/frf_rank.py | 88 +++ alignment_lab/diagnostics/ghost_origin.py | 124 +++++ .../diagnostics/phaser_headtohead.py | 112 ++++ alignment_lab/diagnostics/pose_recovery.py | 155 ++++++ alignment_lab/diagnostics/rescore_rank.py | 118 ++++ alignment_lab/lab/__init__.py | 60 ++ alignment_lab/lab/aniso.py | 172 ++++++ alignment_lab/lab/benchmark.py | 129 +++++ alignment_lab/lab/frf.py | 166 ++++++ alignment_lab/lab/phaser.py | 204 +++++++ alignment_lab/lab/phaser_match.py | 441 +++++++++++++++ alignment_lab/lab/rescore.py | 195 +++++++ alignment_lab/lab/results.py | 139 +++++ alignment_lab/lab/truth.py | 198 +++++++ alignment_lab/tests/test_lab.py | 163 ++++++ 28 files changed, 5099 insertions(+), 1 deletion(-) create mode 100644 alignment_lab/README.md create mode 100644 alignment_lab/analysis/aggregate.py create mode 100644 alignment_lab/analysis/array_template.sh create mode 100644 alignment_lab/diagnostics/bench_stages.py create mode 100644 alignment_lab/diagnostics/frf_aniso_knockout.py create mode 100644 alignment_lab/diagnostics/frf_aniso_rank_sweep.py create mode 100644 alignment_lab/diagnostics/frf_encode_compare.py create mode 100644 alignment_lab/diagnostics/frf_ghost_knockout.py create mode 100644 alignment_lab/diagnostics/frf_inject_phaser_obs.py create mode 100644 alignment_lab/diagnostics/frf_map_compare.py create mode 100644 alignment_lab/diagnostics/frf_normaliser_anatomy.py create mode 100644 alignment_lab/diagnostics/frf_prep_compare.py create mode 100644 alignment_lab/diagnostics/frf_rank.py create mode 100644 alignment_lab/diagnostics/ghost_origin.py create mode 100644 alignment_lab/diagnostics/phaser_headtohead.py create mode 100644 alignment_lab/diagnostics/pose_recovery.py create mode 100644 alignment_lab/diagnostics/rescore_rank.py create mode 100644 alignment_lab/lab/__init__.py create mode 100644 alignment_lab/lab/aniso.py create mode 100644 alignment_lab/lab/benchmark.py create mode 100644 alignment_lab/lab/frf.py create mode 100644 alignment_lab/lab/phaser.py create mode 100644 alignment_lab/lab/phaser_match.py create mode 100644 alignment_lab/lab/rescore.py create mode 100644 alignment_lab/lab/results.py create mode 100644 alignment_lab/lab/truth.py create mode 100644 alignment_lab/tests/test_lab.py diff --git a/.gitignore b/.gitignore index 9d302c4c..60c82421 100644 --- a/.gitignore +++ b/.gitignore @@ -84,4 +84,13 @@ graphify-out/ # Large binary scratch dirs & squashfs images *.sqsh anisotropic/ -torchref_refine_optimization/runs/ \ No newline at end of file +torchref_refine_optimization/runs/ + +# Alignment lab: keep lab/, diagnostics/, analysis/ and tests/ tracked; +# ignore run outputs (CSVs, Phaser working dirs) and scheduler logs. +alignment_lab/runs/ +alignment_lab/slurm/ +alignment_lab/**/__pycache__/ + +# Third-party Phaser source copy (from PHENIX 1.20-4459) for FRF map instrumentation +phaser_src/ diff --git a/alignment_lab/README.md b/alignment_lab/README.md new file mode 100644 index 00000000..f1428eab --- /dev/null +++ b/alignment_lab/README.md @@ -0,0 +1,74 @@ +# alignment_lab + +Harness for the FRF rotation-function work: shared primitives, one file per +experiment, one aggregator. + +``` +lab/ shared library — import from here, do not re-derive +diagnostics/ one experiment per file, each with --out-csv +analysis/ aggregator + SLURM array template +tests/ self-tests for the primitives +runs/ CSVs and Phaser working dirs (gitignored) +slurm/ scheduler logs (gitignored) +``` + +## Why a library + +The scripts this replaces carried ~37 copies of the rotation generator, ~30 of +the benchmark list, ~28 of the CSV writer and ~20 of the rank-of-truth +computation — and several disagreed with each other. Two of those divergences +changed results silently: + +- **Two rotation generators.** One omitted the `sign(diag(R))` QR correction, so + the same seed produced a *different* rotation. Results from the two families + were never comparable. `lab.truth.random_rotation` is the corrected form and + `tests/test_lab.py` pins it against the other variant. +- **Four orbit conventions.** Rank-of-truth was computed with the symmetry + operators applied on either side, in either the fractional or Cartesian frame. + The choice changes the rank, so `orbit_rank` takes it as an explicit argument + and every result row records it. + +## Running + +```bash +PY=.dev/bin/python # or another worktree's interpreter +PYTHONPATH=. $PY alignment_lab/diagnostics/ghost_origin.py --pdb 3K7M --trial 0 +PYTHONPATH=. $PY -m pytest alignment_lab/tests -q + +sbatch --array=0-29 --partition=hour --time=00:55:00 --cpus-per-task=4 \ + --mem=32G alignment_lab/analysis/array_template.sh ghost_origin +PYTHONPATH=. $PY alignment_lab/analysis/aggregate.py \ + 'alignment_lab/runs/ghost_origin_*/*.csv' --compare obs_mode +``` + +## Reading a result + +Seed-to-seed truth-rank spread at `lmax_cap=64` is **±4–6** (1AK5 has been seen +at 9, 11 and 17 for one configuration). Below ~10 trials nothing is +interpretable; three findings that looked strong at n≤7 evaporated at full n. +`aggregate.py` therefore reports paired per-trial differences with the +per-trial values visible, never a bare median, and prints whatever it dropped. + +## Traps worth knowing + +- `model.spacegroup = SpaceGroup("P 1")` is a **silent no-op** — `SpaceGroup` is + an `nn.Module`, so `nn.Module.__setattr__` intercepts the assignment and the + property setter never runs. Assign the **name string**. `ghost_origin.py`'s P1 + arm depends on this. +- `Model.rotate` / `.translate` mutate in place and return `self`. Copy first if + you still need the original — `rotated_case` does. +- Phaser **exits 0 on fatal input errors**; an empty peak list is the real + signal. Its keyword file needs absolute paths. +- `PEAKS ROT SELECT ALL` returns ~80–92k densely spaced samples (median nearest + neighbour under 1°), so "the closest sample is within a degree" means nothing + by itself. + +## Known gap + +`bench_stages.py` currently attributes time only to `phaser_rotation_search` +(~81–85%) and `dense_calc_via_box` (~14–17%). The inner Bessel/Wigner/peak +stages register **0 calls** — the separated engine does not route through those +module-level symbols, so wrapping them there intercepts nothing. They are +printed with their zero counts rather than omitted, because an absent row reads +as a free stage. Getting the inner breakdown needs different instrumentation +points. diff --git a/alignment_lab/analysis/aggregate.py b/alignment_lab/analysis/aggregate.py new file mode 100644 index 00000000..20a2cdb5 --- /dev/null +++ b/alignment_lab/analysis/aggregate.py @@ -0,0 +1,129 @@ +"""Aggregate lab result CSVs. + +One aggregator, because every row shares the core schema. It reports **paired, +per-trial** differences rather than a bare median of each arm: on this benchmark +the seed-to-seed truth-rank spread at ``lmax_cap=64`` is +-4-6 (1AK5 has been +seen at 9, 11 and 17 for the same configuration), so a difference of medians +over a handful of trials is noise. Three findings that looked strong at n<=7 +vanished at full n. + +Anything dropped is printed. A silent truncation reads as full coverage. + +Usage:: + + python alignment_lab/analysis/aggregate.py 'alignment_lab/runs/*.csv' + python alignment_lab/analysis/aggregate.py 'runs/*.csv' --compare obs_mode +""" + +from __future__ import annotations + +import argparse +import csv +import glob +import statistics +from collections import defaultdict +from typing import Dict, List + + +def load(patterns: List[str]) -> List[dict]: + """Read every CSV matching the patterns into a list of row dicts.""" + rows: List[dict] = [] + files = sorted({f for p in patterns for f in glob.glob(p)}) + if not files: + raise SystemExit(f"no CSVs matched {patterns}") + for f in files: + with open(f, newline="") as fh: + rows.extend(csv.DictReader(fh)) + print(f"# {len(rows)} rows from {len(files)} file(s)") + return rows + + +def _rank(row: dict) -> float: + """Truth rank as a number; a miss (-1) sorts as worst, not as best.""" + try: + r = int(row["truth_rank"]) + except (KeyError, ValueError): + return float("nan") + return float("inf") if r < 0 else float(r) + + +def summarise(rows: List[dict]) -> None: + """Per-structure rank summary, with misses counted separately.""" + by_pdb: Dict[str, List[float]] = defaultdict(list) + for r in rows: + by_pdb[r.get("pdb", "?")].append(_rank(r)) + print(f"\n{'pdb':8s} {'n':>3s} {'median':>7s} {'min':>5s} {'max':>5s} " + f"{'misses':>7s} per-trial ranks") + for pdb in sorted(by_pdb): + vals = by_pdb[pdb] + finite = [v for v in vals if v != float("inf")] + misses = sum(1 for v in vals if v == float("inf")) + med = statistics.median(finite) if finite else float("nan") + lo = min(finite) if finite else float("nan") + hi = max(finite) if finite else float("nan") + shown = ", ".join("miss" if v == float("inf") else f"{int(v)}" for v in vals) + print(f"{pdb:8s} {len(vals):3d} {med:7.1f} {lo:5.0f} {hi:5.0f} " + f"{misses:7d} [{shown}]") + if any(v == float("inf") for vs in by_pdb.values() for v in vs): + print("# 'miss' = truth not found in the peak list; excluded from median/min/max") + + +def compare(rows: List[dict], key: str, base: str = None) -> None: + """Paired per-(pdb, seed) comparison across the arms of ``key``.""" + arms = sorted({r.get(key, "") for r in rows}) + if len(arms) < 2: + print(f"\n# only one arm for {key!r}; nothing to pair") + return + cells: Dict[tuple, Dict[str, float]] = defaultdict(dict) + for r in rows: + cells[(r.get("pdb"), r.get("seed"))][r.get(key, "")] = _rank(r) + + if base is not None and base not in arms: + raise SystemExit(f"--base {base!r} not among {key} values {arms}") + base = base if base is not None else arms[0] + arms = [a for a in arms if a != base] + print(f"\n# paired against {key}={base!r}; + means rank got worse") + for arm in arms: + deltas, unpaired = [], 0 + for (pdb, seed), by_arm in sorted(cells.items()): + a, b = by_arm.get(base), by_arm.get(arm) + if a is None or b is None: + unpaired += 1 + continue + if a == float("inf") or b == float("inf"): + unpaired += 1 # a miss has no meaningful numeric difference + continue + deltas.append(b - a) + if not deltas: + print(f" {arm:>16s}: no comparable pairs ({unpaired} unpaired)") + continue + better = sum(1 for d in deltas if d < 0) + worse = sum(1 for d in deltas if d > 0) + same = sum(1 for d in deltas if d == 0) + print(f" {arm:>16s}: median delta {statistics.median(deltas):+.1f} " + f"(better {better} / worse {worse} / unchanged {same}, n={len(deltas)})" + + (f" [{unpaired} pair(s) dropped: missing arm or a miss]" if unpaired else "")) + if len(deltas) < 10: + print(f" {'':16s} n={len(deltas)} is below the ~10 trials this " + f"benchmark needs; treat as indicative only") + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("patterns", nargs="+", help="CSV glob(s)") + ap.add_argument("--base", default=None, + help="which value of --compare is the control arm " + "(default: first alphabetically)") + ap.add_argument("--compare", default=None, + help="column whose values are the arms to pair on, " + "e.g. obs_mode or lmax_cap") + args = ap.parse_args() + rows = load(args.patterns) + summarise(rows) + if args.compare: + compare(rows, args.compare, args.base) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/analysis/array_template.sh b/alignment_lab/analysis/array_template.sh new file mode 100644 index 00000000..a2925eb6 --- /dev/null +++ b/alignment_lab/analysis/array_template.sh @@ -0,0 +1,51 @@ +#!/bin/bash +# SLURM array template for the alignment lab. +# +# Resources go on the sbatch command line, not in this file, so one template +# serves CPU diagnostics and GPU sweeps. The array index selects a (pdb, trial) +# cell from the worklist below. +# +# sbatch --array=0-29 --partition=hour --time=00:55:00 --cpus-per-task=4 \ +# --mem=32G alignment_lab/analysis/array_template.sh ghost_origin +# +# 10 structures x 3 trials = 30 tasks. Note the +-4-6 seed-to-seed rank spread: +# 3 trials is for a smoke run, ~10 for anything you intend to believe. +#SBATCH --job-name=align_lab +#SBATCH --output=alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=alignment_lab/slurm/%x_%A_%a.err +set -euo pipefail + +DIAG="${1:?usage: array_template.sh [extra args...]}" +shift || true + +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY="$REPO/.dev/bin/python" +[ -x "$PY" ] || PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python + +cd "$REPO" +export PYTHONPATH="$REPO" +export TORCHREF_NUM_THREADS="${SLURM_CPUS_PER_TASK:-4}" +export OMP_NUM_THREADS="$TORCHREF_NUM_THREADS" +export MKL_NUM_THREADS="$TORCHREF_NUM_THREADS" +export PYTHONUNBUFFERED=1 +[ -z "${SLURM_JOB_GPUS:-}" ] && export CUDA_VISIBLE_DEVICES="" + +# Worklist: keep in step with lab.benchmark.BENCH_PDBS (order is a seed contract). +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +TRIALS=3 +IDX="${SLURM_ARRAY_TASK_ID:-0}" +PDB="${PDBS[$((IDX / TRIALS))]}" +TRIAL=$((IDX % TRIALS)) + +OUTDIR="alignment_lab/runs/${DIAG}_${SLURM_ARRAY_JOB_ID:-local}" +mkdir -p "$OUTDIR" alignment_lab/slurm + +echo "task $IDX -> $PDB trial $TRIAL -> $OUTDIR" +rc=0 +"$PY" -u "alignment_lab/diagnostics/${DIAG}.py" \ + --pdb "$PDB" --trial "$TRIAL" \ + --out-csv "$OUTDIR/${DIAG}_${PDB}_t${TRIAL}.csv" "$@" || rc=$? + +# Report the real exit status: a task that dies must not be logged COMPLETED. +echo "exit_code=$rc" +exit "$rc" diff --git a/alignment_lab/diagnostics/bench_stages.py b/alignment_lab/diagnostics/bench_stages.py new file mode 100644 index 00000000..ccc7421d --- /dev/null +++ b/alignment_lab/diagnostics/bench_stages.py @@ -0,0 +1,142 @@ +"""Stage-resolved timing of one cold FRF call. + +The eventual goal is placement cheap enough to sit inside a training loop, so +what matters is where the time goes, not just the total. Stage functions are +wrapped for the duration of one call and restored afterwards. + +Timings are **cold by default**: the first call in a process pays one-off costs +(parametrisation, grid setup, any compile). Pass ``--warmup`` for steady-state +numbers, and say which one a reported figure is. + +Usage:: + + python alignment_lab/diagnostics/bench_stages.py --pdb 1DAW --lmax-cap 64 +""" + +from __future__ import annotations + +import argparse +import sys +import time +from collections import defaultdict +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, FRFConfig, ResultWriter, orbit_rank, # noqa: E402 + rotated_case, run_frf, seed_for) +from lab.frf import patched # noqa: E402 + +#: (module path, attribute) pairs timed individually. +STAGES = [ + ("torchref.experimental.alignment.frf.dense_calc", "dense_calc_via_box"), + ("torchref.experimental.alignment.frf.api", "phaser_rotation_search"), + ("torchref.experimental.alignment.frf.bessel", "spherical_bessel_table"), + ("torchref.experimental.alignment.frf.wigner_d", "small_d_stable"), + ("torchref.experimental.alignment.frf.wigner_d", "wigner_contraction_per_beta"), + ("torchref.experimental.alignment.frf.peak_finder", "find_rotation_peaks"), +] + + +def _instrument(stack, totals, counts, skipped): + """Wrap each resolvable stage with a timer, via the exit stack.""" + import importlib + + for mod_path, attr in STAGES: + try: + mod = importlib.import_module(mod_path) + original = getattr(mod, attr) + except (ImportError, AttributeError): + # Report it: a silently skipped stage reads as "that stage is free". + skipped.append(f"{mod_path.rsplit('.', 1)[-1]}.{attr}") + continue + + def make(orig, key): + def timed(*a, **k): + t0 = time.perf_counter() + try: + return orig(*a, **k) + finally: + totals[key] += time.perf_counter() - t0 + counts[key] += 1 + return timed + + # Register at zero so a resolved-but-never-called stage still prints: + # an absent row is indistinguishable from a free one. + totals[attr] += 0.0 + counts[attr] += 0 + stack.enter_context(patched(mod, attr, make(original, attr))) + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) + ap.add_argument("--trial", type=int, default=0) + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--d-min", type=float, default=4.0) + ap.add_argument("--d-max", type=float, default=15.0) + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--warmup", action="store_true", + help="discard one call first and report steady state") + ap.add_argument("--out-csv", default=None) + args = ap.parse_args() + + from contextlib import ExitStack + + seed = seed_for(args.pdb, args.trial) + rotated, data, R_true = rotated_case(args.pdb, seed) + cfg = FRFConfig(d_min=args.d_min, d_max=args.d_max, + n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) + + if args.warmup: + run_frf(rotated, data, cfg, capture_arf=False) + + totals, counts, skipped = defaultdict(float), defaultdict(int), [] + with ExitStack() as stack: + _instrument(stack, totals, counts, skipped) + res = run_frf(rotated, data, cfg) + + sym = data.spacegroup.matrices.to(torch.float64).cpu() + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + rank, ang = orbit_rank(res.peaks, R_true, sym, reciprocal_basis=rec) + + kind = "steady-state" if args.warmup else "cold" + print(f"=== {args.pdb} lmax_cap={args.lmax_cap} ({kind}) ===") + print(f" {'stage':28s} {'calls':>6s} {'seconds':>9s} {'% of run':>9s}") + accounted = 0.0 + for key, secs in sorted(totals.items(), key=lambda kv: -kv[1]): + accounted += secs + print(f" {key:28s} {counts[key]:6d} {secs:9.3f} " + f"{100.0 * secs / max(res.seconds, 1e-9):9.1f}") + print(f" {'(unattributed)':28s} {'':6s} {res.seconds - accounted:9.3f} " + f"{100.0 * (res.seconds - accounted) / max(res.seconds, 1e-9):9.1f}") + print(f" {'TOTAL':28s} {'':6s} {res.seconds:9.3f}") + if skipped: + print(f" NOT INSTRUMENTED (renamed or absent): {', '.join(skipped)}") + print(f" truth rank {rank} at {ang:.2f} deg") + + if args.out_csv: + w = ResultWriter(args.out_csv, "bench_stages", + extra_fields=("timing_kind", "total_seconds", + "unattributed_seconds") + + tuple(f"t_{a}" for _, a in STAGES)) + row = dict(pdb=args.pdb, seed=seed, trial=args.trial, + spacegroup=str(data.spacegroup), n_ops=int(sym.shape[0]), + truth_rank=rank, truth_angle_deg=round(ang, 4), + orbit_side="left", orbit_frame="cart", + lmax_cap=args.lmax_cap, d_min=args.d_min, d_max=args.d_max, + device="cpu", timing_kind=kind, + total_seconds=round(res.seconds, 4), + unattributed_seconds=round(res.seconds - accounted, 4)) + for _, attr in STAGES: + row[f"t_{attr}"] = round(totals.get(attr, 0.0), 4) + w.write(**row) + print(f" wrote {args.out_csv}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_aniso_knockout.py b/alignment_lab/diagnostics/frf_aniso_knockout.py new file mode 100644 index 00000000..6c383e37 --- /dev/null +++ b/alignment_lab/diagnostics/frf_aniso_knockout.py @@ -0,0 +1,165 @@ +"""Is our own anisotropy correction what destroys the hexagonal cases? + +The normaliser anatomy (job 489537) decomposed ``log(Esqr_phaser / Esqr_ours)`` +and found the disagreement is overwhelmingly **angular**, and only on the two +failing structures: + +| pdb | rms | eps_n | iso(|s|) | anisotropy | equivalent B spread | +|------|-------|--------|----------|------------|---------------------| +| 1AK5 | 14% | 61.5% | 23.1% | **0.01%** | 0.36 A^2 | +| 2DQ6 | 44% | 1.6% | 16.4% | **78.1%** | 158 A^2 | +| 3GR5 | 116% | 0.9% | 6.0% | **92.2%** | 189 A^2 | + +A radial mis-scaling is nearly harmless to a rotation function; an angular one is +exactly what it measures. So the suspect is our own overall-anisotropy +correction, ``fit_overall_anisotropy`` (``sh.py:445``), which regresses +``ln|F|^2 - ln<|F|^2>_shell`` on ``-2 pi^2 s.U.s`` by unweighted least squares +**with no constant term**. Single-reflection ``ln|F|^2`` is a badly behaved +regressand: its expectation is offset by ``-gamma`` for acentrics and +``-gamma - ln 2`` for centrics, the ``clamp(min=1e-30)`` turns a vanishing +amplitude into ``y ~ -69``, and with no intercept every one of those offsets is +absorbed into the quadratic form. + +``symmetrize_anisotropy`` then projects the result onto the point-group-invariant +subspace, and the code comment records that the raw fit gives eigenvalues +``(0.8, 17, 70) A^2`` on a *cubic* dataset where symmetry forces them equal. That +projection is why the damage is invisible on the working structures and not on +these two: + +* cubic -> 1 DOF (lambda I): the garbage is annihilated; +* trigonal/hexagonal -> 2 DOF (diag(lambda, lambda, mu)): a **uniaxial tensor + along c is symmetry-allowed**, so the garbage survives as exactly the fake + anisotropy the anatomy measures. + +Three arms settle it, and the null arm is the one that matters -- if switching the +correction off recovers truth, our correction is not merely imperfect, it is +actively destructive: + +* ``production`` -- fitted, symmetrised U; +* ``no_aniso`` -- U = 0, no correction at all; +* ``iso_only`` -- U = (trace/3) I, keeping the radial part and dropping every + angular component, which the shell means then absorb. + +Also reported per structure: the eigenvalues of the fitted tensor before and +after symmetrisation, as B = 8 pi^2 U, so the size of the artefact is visible +next to the rank it costs. + +Usage +----- + python -m diagnostics.frf_aniso_knockout --pdb 3GR5 +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from lab import (ANISO_ARMS, FRFConfig, aniso_arm, load_case, # noqa: E402 + patched, run_frf, tensor_report) +from lab.results import append_row, provenance # noqa: E402 +from diagnostics.frf_ghost_knockout import ( # noqa: E402 + PHASER_PINNED, _orbit_of_identity, _truth_and_margin, +) + +EXPERIMENT = "frf_aniso_knockout" + +ARMS = ANISO_ARMS + + +def run_arm(pdb: str, arm: str, *, n_peaks: int = 500) -> dict: + from torchref.experimental.alignment.frf import api as _api + + pin = PHASER_PINNED[pdb] + model, data = load_case(pdb) + orbit = _orbit_of_identity(data) + + seen: dict = {} + cfg_probe = FRFConfig() + + def _pinned(model_radius_A, d_min_data, lmax_cap=48): + return int(pin["lmax"]) + 1, float(pin["d_min_eff"]) + + cfg = FRFConfig(n_peaks=n_peaks, lmax_cap=int(pin["lmax"]), + extra={"grid_sampling_deg": float(pin["sampling_deg"])}) + t0 = time.time() + with patched(_api, "phaser_lmax_resolution", _pinned), \ + aniso_arm(arm, data, d_min=cfg_probe.d_min, d_max=cfg_probe.d_max, + captured=seen): + res = run_frf(model, data, cfg, capture_arf=True, verbose=0) + rank, sig, ang, ghost, margin = _truth_and_margin(res.arf, orbit) + + # The tensor actually applied, recomputed the same way align.py does it. + from torchref.experimental.alignment.sh import ( + hkl_symops_to_cartesian, symmetrize_anisotropy, + ) + rec = data.cell.reciprocal_basis_matrix.to(torch.float64) + cart = hkl_symops_to_cartesian( + data.spacegroup.matrices.to(torch.float64), rec) + raw = seen.get("raw") + + row = {"experiment": EXPERIMENT, "pdb": pdb, "arm": arm} + row.update(provenance()) + row.update({ + "spacegroup": str(data.spacegroup.hm), + "n_ops": int(data.spacegroup.matrices.shape[0]), + "lmax": pin["lmax"], "sampling_deg": pin["sampling_deg"], + "d_min_eff": pin["d_min_eff"], + "n_samples": int(res.arf.values.numel()), + "truth_rank": rank, "truth_sigma": round(sig, 4), + "truth_angle_deg": round(ang, 3), + "best_ghost_sigma": round(ghost, 4), "margin": round(margin, 4), + "seconds": round(time.time() - t0, 1), + }) + if raw is not None: + row.update(tensor_report(raw.cpu(), "raw")) + row.update(tensor_report( + symmetrize_anisotropy(raw.to(torch.float64).cpu(), cart.cpu()), + "sym")) + if "fixed" in seen: + row.update(tensor_report(seen["fixed"].cpu(), "fix_raw")) + row.update(tensor_report( + symmetrize_anisotropy(seen["fixed"].to(torch.float64).cpu(), + cart.cpu()), "fix_sym")) + return row + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", required=True, choices=sorted(PHASER_PINNED)) + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--outdir", default=None) + args = ap.parse_args() + + outdir = Path(args.outdir) if args.outdir else ( + Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT) + outdir.mkdir(parents=True, exist_ok=True) + csv_path = outdir / f"{EXPERIMENT}_{args.pdb}.csv" + + rows = [run_arm(args.pdb, a, n_peaks=args.n_peaks) for a in ARMS] + r0 = rows[0] + print(f"\n{args.pdb} ({r0['spacegroup']}): fitted anisotropy as B (A^2) -- " + f"raw {r0.get('raw_B_min')}..{r0.get('raw_B_max')} " + f"(spread {r0.get('raw_B_spread')}), after symmetrisation " + f"{r0.get('sym_B_min')}..{r0.get('sym_B_max')} " + f"(spread {r0.get('sym_B_spread')})", flush=True) + print(f"{'arm':<14}{'rank':>8}{'truth_sig':>11}{'ghost_sig':>11}{'margin':>9}", + flush=True) + cols = {} + for r in rows: + cols.update({k: "" for k in r}) + for r in rows: + append_row(csv_path, {**cols, **r}) + print(f"{r['arm']:<14}{r['truth_rank']:>8}{r['truth_sigma']:>11.2f}" + f"{r['best_ghost_sigma']:>11.2f}{r['margin']:>+9.2f}", flush=True) + print(f"\nwrote {csv_path}", flush=True) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_aniso_rank_sweep.py b/alignment_lab/diagnostics/frf_aniso_rank_sweep.py new file mode 100644 index 00000000..75884c37 --- /dev/null +++ b/alignment_lab/diagnostics/frf_aniso_rank_sweep.py @@ -0,0 +1,152 @@ +"""Does the anisotropy fix hold on the task the pipeline actually runs? + +Every number behind the anisotropy diagnosis was measured with the model in its +DEPOSITED orientation and with lmax / sampling / resolution pinned to Phaser's +own logged values -- one evaluation per structure, truth at the identity. That +was the right setup for a bisection against Phaser, and it is the wrong setup +for deciding a default: + +* the pipeline searches a RANDOMLY ROTATED model, not the identity, and the + ghosts are pose-dependent; +* it runs at the production configuration, not Phaser's pinned one; +* seed-to-seed truth-rank spread at ``lmax_cap = 64`` is +-4 to 6 ranks + (1AK5 [9, 11, 17], 3K7M [7, 8, 20]), and three earlier findings in this + investigation looked strong at n <= 7 and vanished at full n. + +So this re-measures the arms over seeded random rotations at the production +config, reporting **per-trial paired differences against the production arm** +rather than a bare median -- the same discipline the rest of the lab uses. + +Arms are :data:`lab.aniso.ARMS`: ``production``, ``no_aniso``, ``iso_only``, +``fixed_fit``. + +Usage +----- + python -m diagnostics.frf_aniso_rank_sweep --pdb 3GR5 --trials 10 +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (ANISO_ARMS, BENCH_PDBS, FRFConfig, aniso_arm, # noqa: E402 + orbit_rank, rotated_case, run_frf, seed_for, tensor_report) +from lab.results import append_row, provenance # noqa: E402 + +EXPERIMENT = "frf_aniso_rank_sweep" + + +def run_one(pdb: str, trial: int, arm: str, cfg: FRFConfig, + *, thr_deg: float) -> dict: + seed = seed_for(pdb, trial) + model, data, R_true = rotated_case(pdb, seed) + captured: dict = {} + t0 = time.time() + with aniso_arm(arm, data, d_min=cfg.d_min, d_max=cfg.d_max, + captured=captured): + res = run_frf(model, data, cfg, capture_arf=False, verbose=0) + seconds = time.time() - t0 + + rank, ang = orbit_rank( + res.peaks, R_true, data.spacegroup.matrices.to(torch.float64).cpu(), + reciprocal_basis=data.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), + side="left", frame="cart", thr_deg=thr_deg, + ) + row = {"experiment": EXPERIMENT, "pdb": pdb, "trial": trial, "arm": arm, + "seed": seed} + row.update(provenance()) + row.update(cfg.as_row()) + row.update({ + "spacegroup": str(data.spacegroup.hm), + "truth_rank": rank, + # orbit_rank returns -1 for "no peak within thr_deg". That must NOT be + # ordered as a good rank: for paired comparison a miss counts as worse + # than the worst hit, i.e. the peak-list length. + "rank_for_compare": rank if rank >= 0 else cfg.n_peaks, + "found": int(rank >= 0), + "truth_angle_deg": None if ang is None else round(float(ang), 3), + "n_peaks_found": len(res.peaks), + "orbit_side": "left", "orbit_frame": "cart", "thr_deg": thr_deg, + "seconds": round(seconds, 1), + }) + for tag in ("raw", "fixed"): + if tag in captured: + row.update(tensor_report(captured[tag], tag)) + return row + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) + ap.add_argument("--trials", type=int, default=10) + ap.add_argument("--arms", default=",".join(ANISO_ARMS)) + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--d-min", type=float, default=4.0) + ap.add_argument("--d-max", type=float, default=15.0) + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--thr-deg", type=float, default=5.0) + ap.add_argument("--outdir", default=None) + args = ap.parse_args() + + arms = [a for a in args.arms.split(",") if a] + cfg = FRFConfig(d_min=args.d_min, d_max=args.d_max, + n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) + outdir = Path(args.outdir) if args.outdir else ( + Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT) + outdir.mkdir(parents=True, exist_ok=True) + csv_path = outdir / f"{EXPERIMENT}_{args.pdb}.csv" + + print(f"{args.pdb}: truth rank per trial, lmax_cap={args.lmax_cap}", + flush=True) + print(f"{'trial':>6}" + "".join(f"{a:>15}" for a in arms), flush=True) + ranks = {a: [] for a in arms} + n_fail = 0 + for trial in range(args.trials): + cells = [] + for arm in arms: + try: + row = run_one(args.pdb, trial, arm, cfg, thr_deg=args.thr_deg) + except Exception as exc: + n_fail += 1 + ranks[arm].append(None) + cells.append(f"{type(exc).__name__}") + print(f" trial {trial} arm {arm} FAILED: {exc}", flush=True) + continue + append_row(csv_path, row) + ranks[arm].append(row["rank_for_compare"]) + cells.append(str(row["truth_rank"]) if row["found"] + else f"miss({row['truth_angle_deg']:.0f}d)") + print(f"{trial:>6}" + "".join(f"{c:>15}" for c in cells), flush=True) + + # Paired differences against production; per-trial signs, never a bare median. + base = ranks.get("production") + if base: + print("\npaired vs production (negative = better rank):", flush=True) + for arm in arms: + if arm == "production": + continue + d = [(a - b) for a, b in zip(ranks[arm], base) + if a is not None and b is not None] + if not d: + print(f" {arm:<14} no paired trials", flush=True) + continue + sd = sorted(d) + med = (sd[len(sd) // 2] if len(sd) % 2 + else 0.5 * (sd[len(sd) // 2 - 1] + sd[len(sd) // 2])) + print(f" {arm:<14} n={len(d):<3} better={sum(x < 0 for x in d)} " + f"same={sum(x == 0 for x in d)} worse={sum(x > 0 for x in d)} " + f"median_delta={med:+.1f} per-trial={d}", flush=True) + print(f"\nwrote {csv_path} ({n_fail} failures)", flush=True) + return 1 if n_fail == len(arms) * args.trials else 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_encode_compare.py b/alignment_lab/diagnostics/frf_encode_compare.py new file mode 100644 index 00000000..e1c90135 --- /dev/null +++ b/alignment_lab/diagnostics/frf_encode_compare.py @@ -0,0 +1,511 @@ +"""Feed Phaser's own prepared data into our SH-Bessel encoder. + +Every earlier comparison changed two things at once: the *inputs* to the +expansion (normalisation, symmetry unroll, F_calc) and the *expansion itself*. +The per-reflection attribution (job 489442) pinned the input side -- Phaser's +own identities reproduce bit-exactly, DFAC and V are unity, and the residual +disagreement is a roughly uniform ~7-11% in the Wilson normaliser across all +structures, so it does not single out the trigonal/hexagonal failures. That +leaves the encoder untested on its own. + +This runs the encoder with Phaser's inputs, so a mismatch can only come from our +projection: + +* ``obs_phaser_pts`` -- Phaser's prepared observations (``PHASER_OBS_DUMP``: + post-normalisation, post-LERF1, post-unroll, post-axis-permutation, in polar + coordinates, i.e. exactly what ``DataMR::getELMNxR2`` consumes) through + ``bessel_sh_expand``, against Phaser's ``DataElmn``. +* ``calc_phaser_pts`` -- Phaser's molecular-transform samples + (``PHASER_CALC_DUMP``, from ``Ensemble::getELMNxR2`` -- a *different* function + with its own radial scale and its own l != 0 doubling) against Phaser's + ``SearchElmn``. +* ``obs_phaser_clustered`` / ``calc_phaser_clustered`` -- the same two, but + replaying Phaser's own angular approximation. Phaser buckets reflections by + ``|cos(theta) - cos(theta_rep)| < 1e-3`` and evaluates the Legendre functions + once per bucket from the first member's theta (sphericalY.h:43, + DataMR.cc:1096); we cluster only on values equal to ~1e-7. So a high-l + disagreement in the arms above is *expected*, and is Phaser being + approximate rather than us being wrong. These arms separate the two, and + answer a question that has never been asked: whether that 1e-3 polar + smoothing is part of why Phaser is immune to the symmetry-axis ghosts. +* ``obs_ours_unroll`` / ``obs_dedup_unroll`` -- Phaser's ASU-level intensities + (``PHASER_TERMS_DUMP``, keyed by Miller index) put through *our* two symmetry + unrolls: the production one, which emits all ``n_ops`` orbit positions, and + ``epsilon_aware_unroll``, which emits only the distinct ones as Phaser does + (``!duplicate(isym,rhkl)``, DataMR.cc:954). Same intensities, same encoder, + same target -- so the difference between these two arms is the multiplicity + handling and nothing else. + +Two scalars the expansion needs are not recoverable from the dumped rows -- the +observation-side ``HIRES`` is the *minimum* reso over selected reflections, +one step below the smallest that survives the ``reso(r) > HIRES`` gate -- so the +instrumented binary now writes them to ``.meta`` and they are read, not +inferred. + +Expected relation, if our encoder is right. Phaser projects with ``Y_lm`` +(``e^{+im phi}``, Condon-Shortley sign folded into its ``Pmm`` recurrence) while +we project with ``conj(C(m,phi))``; ``bar_P`` carries no CS phase and our +``sign_m`` restores it. Both are real-weighted sums, so + + ours[n, l, m] = k * conj(phaser[l, m, n+1]), k = 2 if enforce_friedel else 1 + +with the factor 2 because appending ``-s`` doubles every even-l coefficient +exactly (``Y_lm(-s) = (-1)^l Y_lm(s)``). ``k`` is therefore a prediction, not a +fitted nuisance: a modulus away from 1 or a phase away from 0 is a finding. + +Usage +----- + python -m diagnostics.frf_encode_compare --pdb 2DQ6 +""" + +from __future__ import annotations + +import argparse +import math +import os +import subprocess +import sys +import time +from pathlib import Path + +import numpy as np +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from lab import case_paths, load_case # noqa: E402 +from lab.phaser_match import PATCHED_PHASER, write_keywords # noqa: E402 +from lab.results import append_row, provenance # noqa: E402 + +EXPERIMENT = "frf_encode_compare" + + +# --------------------------------------------------------------------------- +# Phaser side +# --------------------------------------------------------------------------- + +def run_phaser_dumps(pdb: str, work: Path) -> dict: + """One instrumented run producing every stage this comparison needs.""" + work.mkdir(parents=True, exist_ok=True) + pdb_path, mtz_path = case_paths(pdb) + kw = write_keywords(work, mtz_path=mtz_path, model_pdb=pdb_path, + n_peaks=5, root=f"{pdb}_enc", title=f"encode {pdb}") + paths = { + "PHASER_OBS_DUMP": work / "obs.csv", + "PHASER_CALC_DUMP": work / "calc.csv", + "PHASER_TERMS_DUMP": work / "terms.csv", + "PHASER_DATA_ELMN_DUMP": work / "data_elmn.csv", + "PHASER_SEARCH_ELMN_DUMP": work / "search_elmn.csv", + } + env = dict(os.environ) + for k, v in paths.items(): + env[k] = str(v) + proc = subprocess.run([str(PATCHED_PHASER)], cwd=str(work), + input=kw.read_text(), capture_output=True, + text=True, timeout=5400, env=env) + (work / "run.log").write_text((proc.stdout or "") + (proc.stderr or "")) + missing = [v.name for v in paths.values() if not v.exists()] + if missing: + raise RuntimeError(f"{pdb}: missing dumps {missing}; see {work/'run.log'}") + # "PHASER_OBS_DUMP" -> "obs": strip both the prefix and the _DUMP suffix. + return {k[len("PHASER_"):-len("_DUMP")].lower(): v for k, v in paths.items()} + + +def read_meta(path: Path) -> dict: + """``.meta`` -- the scalars the point list cannot carry.""" + meta = Path(str(path) + ".meta") + if not meta.exists(): + raise RuntimeError( + f"{meta} missing: rebuild the instrumented binary " + f"(phaser_src/build/rebuild.sh) -- the Bessel scale would otherwise " + f"have to be guessed." + ) + out = {} + for line in meta.read_text().splitlines()[1:]: + k, v = line.split(",") + out[k] = float(v) + return out + + +def _cart(r, th, ph) -> torch.Tensor: + return torch.from_numpy(np.stack( + [r * np.sin(th) * np.cos(ph), + r * np.sin(th) * np.sin(ph), + r * np.cos(th)], axis=1)).to(torch.float64) + + +def load_points(path: Path): + """``cluster,r,theta,phi,intensity`` -> our encoder's inputs. + + Returns ``(s_exact, intensity, cos_theta, s_clustered, cluster_stats)``. + + ``s_clustered`` replays Phaser's OWN angular approximation. + ``HKL_clustered::add`` (sphericalY.h:43) buckets reflections greedily by + ``|cos(theta) - cos(theta_rep)| < 1e-3`` against the FIRST member of each + bucket, and the projection then evaluates the associated Legendre functions + once per bucket from that first member's theta (DataMR.cc:1096) -- while the + radial Bessel term stays per-reflection. So Phaser's Y_lm carries up to + 1e-3 of cos-theta error, which at l ~ 70 is a percent-level per-coefficient + error, largest near the poles where sin(theta) is small. + + Our encoder clusters only on values that are equal to ~1e-7, so it is the + more accurate of the two. That means a high-l disagreement is expected and + is Phaser's approximation, not our defect -- and it has to be separated + from a real difference before any residual can be read. Substituting each + point's bucket-representative ``cos(theta)`` while keeping its own ``r`` and + ``phi`` reproduces Phaser's evaluation exactly, because ``r`` and ``phi`` + are the only per-reflection quantities Phaser keeps. + + The dump is written before ``HKL_list.shuffle()``, so row order within a + cluster is insertion order and row 0 of each cluster is the representative. + """ + d = np.loadtxt(path, delimiter=",", skiprows=1) + if d.ndim == 1: + d = d[None, :] + cid = d[:, 0].astype(np.int64) + r, th, ph, val = d[:, 1], d[:, 2], d[:, 3], d[:, 4] + + first = np.zeros(cid.max() + 1, dtype=np.int64) + seen = np.zeros(cid.max() + 1, dtype=bool) + for i, c in enumerate(cid): + if not seen[c]: + seen[c], first[c] = True, i + th_rep = th[first[cid]] + stats = { + "phaser_n_clusters": int(seen.sum()), + "phaser_cos_spread_max": float(np.abs(np.cos(th) - np.cos(th_rep)).max()), + "phaser_cluster_size_max": int(np.bincount(cid).max()), + } + return (_cart(r, th, ph), + torch.from_numpy(val).to(torch.float64), + torch.from_numpy(np.cos(th)).to(torch.float64), + _cart(r, th_rep, ph), + stats) + + +def load_elmn(path: Path, L: int) -> torch.Tensor: + """Phaser's ``l,m,n`` dump into our ``(N_radial, L, 2L-1)`` layout. + + Phaser's ``n`` is 1-based against our 0-based, and both index the same + ``u = l + 2n - 1`` radial order, so ``n0 = n - 1``. + """ + lmax = L - 1 + lmax_even = lmax if lmax % 2 == 0 else lmax - 1 + n_radial = (lmax_even - 2) // 2 + 1 + out = torch.zeros((n_radial, L, 2 * L - 1), dtype=torch.complex128) + d = np.loadtxt(path, delimiter=",", skiprows=1) + if d.ndim == 1: + d = d[None, :] + l = d[:, 0].astype(int) + m = d[:, 1].astype(int) + n0 = d[:, 2].astype(int) - 1 + keep = (l <= lmax_even) & (n0 >= 0) & (n0 < n_radial) & (np.abs(m) <= lmax_even) + if not keep.all(): + raise RuntimeError(f"{path}: {int((~keep).sum())} rows outside the L={L} band") + out[n0, l, m + (L - 1)] = torch.from_numpy(d[:, 3] + 1j * d[:, 4]) + return out + + +def band_mask(L: int, device="cpu") -> torch.Tensor: + """The (n, l, m) entries Phaser allocates: l even, |m| <= l, n < nmax(l).""" + lmax = L - 1 + lmax_even = lmax if lmax % 2 == 0 else lmax - 1 + n_radial = (lmax_even - 2) // 2 + 1 + mask = torch.zeros((n_radial, L, 2 * L - 1), dtype=torch.bool, device=device) + for l in range(2, lmax_even + 1, 2): + n_l = (lmax_even - l) // 2 + 1 + m_lo, m_hi = (L - 1) - l, (L - 1) + l + mask[:n_l, l, m_lo:m_hi + 1] = True + return mask + + +# --------------------------------------------------------------------------- +# comparison +# --------------------------------------------------------------------------- + +def compare_coeffs(ours: torch.Tensor, phaser: torch.Tensor, L: int, + *, k_expected: float) -> dict: + """Our coefficients against ``conj(phaser)`` over Phaser's allocated band. + + Reported quantities: + ``corr`` modulus of the complex correlation -- shape agreement. + ``k_mod``/``k_arg_deg`` the fitted complex scale; the prediction is + ``k_expected`` at 0 degrees, so a phase here means a + convention mismatch, not a scale. + ``rel_resid`` ``||a - k b|| / ||a||`` after the fitted scale, i.e. what + the correlation hides. + ``pow_offband`` fraction of OUR power sitting where Phaser has exactly + zero -- the m-filter / forbidden-m channel. + ``worst_l`` the even l with the lowest per-l correlation. + """ + mask = band_mask(L, device=ours.device) + b = torch.conj(phaser.to(ours.device)) + a = ours + av, bv = a[mask], b[mask] + + num = torch.vdot(bv, av) # sum conj(b)*a + denom_b = (bv.abs() ** 2).sum() + out = { + "n_band": int(mask.sum()), + "n_phaser_nonzero": int((bv.abs() > 0).sum()), + "n_ours_nonzero": int((av.abs() > 0).sum()), + } + if float(denom_b) == 0.0: + out.update(corr=float("nan"), k_mod=float("nan"), + k_arg_deg=float("nan"), rel_resid=float("nan")) + return out + corr = float(num.abs() / (av.norm() * bv.norm()).clamp(min=1e-300)) + k = num / denom_b + resid = (av - k * bv).norm() / av.norm().clamp(min=1e-300) + out.update({ + "corr": corr, + "k_mod": float(k.abs()), + "k_arg_deg": float(torch.rad2deg(torch.angle(k))), + "k_expected": float(k_expected), + "rel_resid": float(resid), + }) + # Power we place where Phaser has none (inside its own band). + zero_b = mask & (b.abs() == 0) + out["pow_offband"] = float( + (a[zero_b].abs() ** 2).sum() / (a[mask].abs() ** 2).sum().clamp(min=1e-300) + ) + # Per-l correlation, to see whether a mismatch is radial (high l) or global. + lmax_even = (L - 1) if (L - 1) % 2 == 0 else (L - 2) + worst_l, worst_c = -1, 2.0 + per_l = [] + for l in range(2, lmax_even + 1, 2): + ml = mask[:, l, :] + al, bl = a[:, l, :][ml], b[:, l, :][ml] + if float(bl.abs().max()) == 0.0: + continue + cl = float(torch.vdot(bl, al).abs() + / (al.norm() * bl.norm()).clamp(min=1e-300)) + per_l.append((l, cl)) + if cl < worst_c: + worst_c, worst_l = cl, l + out["worst_l"] = worst_l + out["worst_l_corr"] = worst_c + out["corr_l2"] = per_l[0][1] if per_l else float("nan") + out["corr_lmax"] = per_l[-1][1] if per_l else float("nan") + return out + + +# --------------------------------------------------------------------------- +# our side +# --------------------------------------------------------------------------- + +def encode(s: torch.Tensor, intensity: torch.Tensor, *, L: int, + h_scale: float, zsymm: int, friedel: bool) -> torch.Tensor: + from torchref.experimental.alignment.frf.data_mr import bessel_sh_expand + return bessel_sh_expand( + s, intensity, L=L, bessel_h_scale=h_scale, zsymm=zsymm, + enforce_friedel=friedel, + ).coeffs + + +def unroll_arms(pdb: str, terms_csv: Path): + """Phaser's ASU intensities through both of our symmetry unrolls. + + Returns ``(arms, stats)`` where ``arms`` maps name -> ``(s, intensity)``. + The counts are exact integers, so the multiplicity question is answered by + arithmetic before any encoding happens. + """ + from torchref.experimental.alignment.frf.preprocessing import ( + epsilon_aware_unroll, + ) + + d = np.loadtxt(terms_csv, delimiter=",", skiprows=1) + if d.ndim == 1: + d = d[None, :] + hkl = torch.from_numpy(d[:, 0:3]).to(torch.float64) + inten = torch.from_numpy(d[:, 9]).to(torch.float64) + + _, data = load_case(pdb) + sg = data.spacegroup.matrices.to(torch.float64).cpu() + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + n_ops = int(sg.shape[0]) + + # Production: every orbit position, duplicates included (align.py:487). + hkl_all = torch.einsum("kji,nj->kni", sg, hkl).reshape(-1, 3) + s_all = hkl_all @ rec + i_all = inten.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() + + # Phaser-faithful: distinct orbit positions only (DataMR.cc:954). + hkl_ded, asu_idx = epsilon_aware_unroll(hkl.to(torch.long), sg) + s_ded = hkl_ded.to(torch.float64) @ rec + i_ded = inten[asu_idx] + + stats = { + "n_asu_terms": int(hkl.shape[0]), + "n_ops": n_ops, + "n_unroll_all": int(s_all.shape[0]), + "n_unroll_dedup": int(s_ded.shape[0]), + "dup_frac": float(1.0 - s_ded.shape[0] / max(1, s_all.shape[0])), + } + return {"obs_ours_unroll": (s_all, i_all), + "obs_dedup_unroll": (s_ded, i_ded)}, stats + + +def frame_check(s_ours: torch.Tensor, s_phaser: torch.Tensor, + *, n_sample: int = 2000) -> dict: + """Do the two Cartesian reciprocal frames coincide? + + ``|s|`` is frame-independent but theta and phi are not, so an orthogonalisation + convention difference would rotate every coefficient (mixing m) and make the + coefficient comparison meaningless while leaving the radial part intact. + Nearest-neighbour distance in Cartesian space tests position, not just radius. + """ + g = torch.Generator().manual_seed(1) + k = min(n_sample, int(s_phaser.shape[0])) + q = s_phaser[torch.randperm(s_phaser.shape[0], generator=g)[:k]] + dist = torch.empty(k, dtype=torch.float64) + step = 100 # the (step, N, 3) broadcast is the memory bound, not the (step, N) d2 + for i in range(0, k, step): + d2 = ((q[i:i + step, None, :] - s_ours[None, :, :]) ** 2).sum(-1) + dist[i:i + step] = d2.min(1).values.clamp(min=0).sqrt() + return {"frame_median_dist": float(dist.median()), + "frame_max_dist": float(dist.max()), + "frame_matched_frac": float((dist < 1e-9).to(torch.float64).mean())} + + +# --------------------------------------------------------------------------- +# driver +# --------------------------------------------------------------------------- + +def run(pdb: str, outdir: Path, *, reuse: Path | None = None) -> list: + work = reuse if reuse is not None else (outdir / "phaser" / pdb) + if reuse is not None: + dumps = {n: work / f"{n}.csv" for n in + ("obs", "calc", "terms", "data_elmn", "search_elmn")} + missing = [str(p) for p in dumps.values() if not p.exists()] + if missing: + raise RuntimeError(f"--reuse given but missing: {missing}") + else: + t0 = time.time() + dumps = run_phaser_dumps(pdb, work) + print(f" phaser dumps in {time.time()-t0:.0f}s", flush=True) + + obs_meta = read_meta(dumps["obs"]) + calc_meta = read_meta(dumps["calc"]) + L = int(obs_meta["lmax"]) + 1 + zsymm = int(obs_meta["zsymm"]) + h_obs = obs_meta["lmax"] * obs_meta["hires"] + h_calc = calc_meta["lmax"] * calc_meta["max_resolution"] + print(f" L={L} zsymm={zsymm} axis={int(obs_meta['axis'])} " + f"hires={obs_meta['hires']:.6f} h_obs={h_obs:.4f} " + f"h_calc={h_calc:.4f} (calc max_reso={calc_meta['max_resolution']:.6f})", + flush=True) + if int(obs_meta["axis"]) != 3: + print(" NOTE: high-order axis is not c -- Phaser permutes the frame " + "(DataMR.cc:984) and we do not; the obs arm is confounded.", + flush=True) + + data_elmn = load_elmn(dumps["data_elmn"], L) + search_elmn = load_elmn(dumps["search_elmn"], L) + s_obs, i_obs, _, s_obs_clu, obs_clu = load_points(dumps["obs"]) + s_calc, i_calc, cos_calc, s_calc_clu, calc_clu = load_points(dumps["calc"]) + calc_clu = {f"calc_{k}": v for k, v in calc_clu.items()} + print(f" phaser obs clusters={obs_clu['phaser_n_clusters']} " + f"(max cos-theta spread {obs_clu['phaser_cos_spread_max']:.2e}, " + f"largest bucket {obs_clu['phaser_cluster_size_max']}); calc clusters=" + f"{calc_clu['calc_phaser_n_clusters']} " + f"(spread {calc_clu['calc_phaser_cos_spread_max']:.2e})", flush=True) + + base = {"experiment": EXPERIMENT, "pdb": pdb} + base.update(provenance()) + base.update({"L": L, "zsymm": zsymm, "axis": int(obs_meta["axis"]), + "hires": obs_meta["hires"], "h_obs": h_obs, "h_calc": h_calc}) + + rows = [] + + def emit(arm, coeffs, target, *, n_points, seconds, extra=None): + r = dict(base, arm=arm, n_points=int(n_points), + seconds=round(seconds, 1)) + r.update(compare_coeffs(coeffs, target, L, k_expected=1.0)) + if extra: + r.update(extra) + rows.append(r) + + # --- arm 1: Phaser's own observations through our encoder --------------- + t0 = time.time() + c = encode(s_obs, i_obs, L=L, h_scale=h_obs, zsymm=zsymm, friedel=False) + emit("obs_phaser_pts", c, data_elmn, + n_points=s_obs.shape[0], seconds=time.time() - t0, extra=obs_clu) + + # --- arm 2: Phaser's observations WITH Phaser's own theta approximation -- + t0 = time.time() + c = encode(s_obs_clu, i_obs, L=L, h_scale=h_obs, zsymm=zsymm, friedel=False) + emit("obs_phaser_clustered", c, data_elmn, + n_points=s_obs_clu.shape[0], seconds=time.time() - t0, extra=obs_clu) + + # --- arm 3: Phaser's calc samples through our encoder ------------------- + # Phaser doubles every l != 0 grid point (Ensemble.cc: `flipped`), because + # its molecular-transform grid stores only the l >= 0 hemisphere. l == 0 is + # exactly the s_z == 0 plane for these settings (c* along z), so the flag is + # recoverable from the dumped theta. + t0 = time.time() + flip = (cos_calc.abs() > 1e-12).to(torch.float64) + 1.0 + c = encode(s_calc, i_calc * flip, L=L, h_scale=h_calc, zsymm=1, friedel=False) + emit("calc_phaser_pts", c, search_elmn, n_points=s_calc.shape[0], + seconds=time.time() - t0, + extra=dict(calc_clu, n_l0_plane=int((cos_calc.abs() <= 1e-12).sum()))) + + # --- arm 4: the same, with Phaser's theta approximation ----------------- + t0 = time.time() + c = encode(s_calc_clu, i_calc * flip, L=L, h_scale=h_calc, zsymm=1, + friedel=False) + emit("calc_phaser_clustered", c, search_elmn, n_points=s_calc_clu.shape[0], + seconds=time.time() - t0, extra=calc_clu) + + # --- arms 5/6: Phaser's ASU intensities through OUR unrolls ------------- + arms, ustats = unroll_arms(pdb, dumps["terms"]) + print(f" unroll: asu={ustats['n_asu_terms']} x n_ops={ustats['n_ops']} " + f"= {ustats['n_unroll_all']} all / {ustats['n_unroll_dedup']} dedup " + f"(phaser obs rows {int(s_obs.shape[0])}; " + f"dup_frac {ustats['dup_frac']:.4f})", flush=True) + for arm, (s, val) in arms.items(): + t0 = time.time() + fc = frame_check(s, s_obs) + c = encode(s, val, L=L, h_scale=h_obs, zsymm=zsymm, friedel=False) + emit(arm, c, data_elmn, n_points=s.shape[0], seconds=time.time() - t0, + extra=dict(ustats, **fc)) + + # One schema for every row -- csv.DictWriter fixes fieldnames from the + # first row it sees, so a later row carrying extra keys would raise. + cols = {} + for r in rows: + cols.update({k: "" for k in r}) + return [{**cols, **r} for r in rows] + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", required=True) + ap.add_argument("--outdir", default=None) + ap.add_argument("--reuse", default=None, + help="directory of existing dumps (skips the Phaser run)") + args = ap.parse_args() + + outdir = Path(args.outdir) if args.outdir else ( + Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT) + outdir.mkdir(parents=True, exist_ok=True) + csv_path = outdir / f"{EXPERIMENT}_{args.pdb}.csv" + + rows = run(args.pdb, outdir, + reuse=Path(args.reuse) if args.reuse else None) + hdr = (f"{'arm':<20}{'n_pts':>9}{'corr':>9}{'k_mod':>9}{'k_arg':>9}" + f"{'resid':>9}{'offband':>9}{'worst_l':>9}{'wl_corr':>9}") + print(f"\n{args.pdb}: our encoder against Phaser's coefficients", flush=True) + print(hdr, flush=True) + for r in rows: + append_row(csv_path, r) + print(f"{r['arm']:<20}{r['n_points']:>9}{r['corr']:>9.4f}" + f"{r['k_mod']:>9.4f}{r['k_arg_deg']:>9.2f}{r['rel_resid']:>9.4f}" + f"{r['pow_offband']:>9.4f}{r['worst_l']:>9}" + f"{r['worst_l_corr']:>9.4f}", flush=True) + print(f"\nwrote {csv_path}", flush=True) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_ghost_knockout.py b/alignment_lab/diagnostics/frf_ghost_knockout.py new file mode 100644 index 00000000..0bf7eaf8 --- /dev/null +++ b/alignment_lab/diagnostics/frf_ghost_knockout.py @@ -0,0 +1,200 @@ +"""Find which obs-side term creates the trigonal/hexagonal ghosts. + +Context. With Phaser's bandwidth, resolution and SO(3) sampling pinned, our FRF +puts truth at rank 0 on five of seven benchmark structures and beats Phaser on +6G9X -- but collapses on 2DQ6 (P 3_1 2 1, truth rank 76799) and 3GR5 +(P 6_5 2 2, rank 16855), where Phaser gets rank 0 with a healthy margin. Those +are the only two cases with a 120 degree cell and a 3-fold-containing axis; the +working set is monoclinic / orthorhombic / tetragonal / cubic. + +Because lmax and sampling were already pinned to Phaser's own values when the +collapse was measured, bandwidth is eliminated. What remains is the observation- +and calc-side preprocessing chain, where our engine differs from +``DataMR::getELMNxR2`` in ways that are individually documented but never +isolated on these two space groups: + +* ``use_epsilon=False`` -- Phaser always normalises by ``sqrt(eps_n * Sigma_N)`` + (DataMR.cc:930). Without the epsilon divisor the axial/zonal reflections are + over-weighted, and for a 6-fold axis those carry eps up to 6-12. +* ``_orbit_unroll=False`` -- Phaser expands over symmetry but skips duplicate + P1 indices (``!duplicate(isym,rhkl)``, DataMR.cc:954). We replicate all n_ops + unconditionally, so reflections on special positions are counted several times. +* the m-symmetry filter, French-Wilson, shell-variance weights, Wilson-B match, + Oeffner vrms and the Babinet bulk-solvent term. + +Each arm flips exactly one of these against a common baseline and reports where +truth lands. The discriminating statistic is ``margin`` -- truth's sigma minus +the strongest non-truth peak's. Negative means the rotation function prefers a +ghost, which is the failure we are chasing; rank alone hides how close the call +was. + +Usage +----- + python -m diagnostics.frf_ghost_knockout --pdb 2DQ6 + python -m diagnostics.frf_ghost_knockout --pdb 3GR5 --arm epsilon +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from lab import FRFConfig, load_case, patched, run_frf # noqa: E402 +from lab.results import append_row, provenance # noqa: E402 + +EXPERIMENT = "frf_ghost_knockout" + +#: Phaser's own bandwidth / resolution / sampling per case, read from the +#: instrumented run (job 487737). Pinned so every arm differs only in the term +#: under test -- and so the result is comparable with the map-comparison run. +PHASER_PINNED = { + "1DAW": dict(lmax=58, sampling_deg=6.233148, d_min_eff=5.70), + "1AK5": dict(lmax=70, sampling_deg=5.155428, d_min_eff=4.84), + "3K7M": dict(lmax=66, sampling_deg=5.464844, d_min_eff=5.39), + "4BX9": dict(lmax=96, sampling_deg=3.774036, d_min_eff=6.73), + "6G9X": dict(lmax=84, sampling_deg=4.314963, d_min_eff=5.92), + "2DQ6": dict(lmax=76, sampling_deg=4.772056, d_min_eff=6.36), + "3GR5": dict(lmax=66, sampling_deg=5.483967, d_min_eff=4.10), +} + +#: One flipped term per arm. ``baseline`` is our production configuration. +ARMS = { + "baseline": {}, + "epsilon": dict(use_epsilon=True), + "orbit_unroll": dict(_orbit_unroll=True), + "eps+unroll": dict(use_epsilon=True, _orbit_unroll=True), + "no_m_filter": dict(frf_use_m_filter=False), + "no_french": dict(frf_use_french_wilson=False), + "no_shellvar": dict(frf_use_shell_variance=False), + "no_bulk_solv": dict(apply_bulk_solvent=False), + "no_wilson_b": dict(apply_wilson_b=False), + "vrms_fixed": dict(vrms_strategy="fixed"), + "acentric_only": dict(frf_acentric_only=True), +} + + +def _orbit_of_identity(data): + """Truth is the identity, up to the point group, in the Cartesian frame.""" + from lab.truth import symmetry_orbit + + symops = data.spacegroup.matrices.to(torch.float64).cpu() + recip = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + return symmetry_orbit( + torch.eye(3, dtype=torch.float64), symops, + side="left", frame="cart", reciprocal_basis=recip, + ) + + +def _truth_and_margin(arf, orbit): + """``(rank, sigma, angle, best_ghost_sigma, margin)`` for one map.""" + from torchref.experimental.alignment.frf.rotation_utils import ( + rotation_matrix_from_edmonds_euler_batch, + ) + + v = arf.values.to(torch.float64).cpu() + R = rotation_matrix_from_edmonds_euler_batch( + arf.alphas.to(torch.float64).cpu(), + arf.betas.to(torch.float64).cpu(), + arf.gammas.to(torch.float64).cpu(), + ) + sig = (v - v.mean()) / v.std().clamp(min=1e-30) + + best = None + tol = 1.5 * float(arf.grid_sampling_deg) + truth_mask = torch.zeros(v.numel(), dtype=torch.bool) + for k in range(orbit.shape[0]): + tr = torch.einsum("nij,ij->n", R, orbit[k]) + ang = torch.rad2deg(torch.arccos(((tr - 1.0) / 2.0).clamp(-1.0, 1.0))) + truth_mask |= ang <= tol + j = int(torch.argmin(ang)) + if best is None or v[j] > v[best[0]]: + best = (j, float(ang[j])) + j, ang_j = best + rank = int((v > v[j]).sum()) + ghost_sig = sig.masked_fill(truth_mask, float("-inf")) + gs = float(ghost_sig.max()) + return rank, float(sig[j]), ang_j, gs, float(sig[j]) - gs + + +def run_arm(pdb: str, arm: str, extra: dict, *, n_peaks: int = 500) -> dict: + """One knockout arm on one structure.""" + from torchref.experimental.alignment.frf import api as _api + + pin = PHASER_PINNED[pdb] + model, data = load_case(pdb) + orbit = _orbit_of_identity(data) + + def _pinned(model_radius_A, d_min_data, lmax_cap=48): + return int(pin["lmax"]) + 1, float(pin["d_min_eff"]) + + cfg = FRFConfig( + n_peaks=n_peaks, lmax_cap=int(pin["lmax"]), + extra={"grid_sampling_deg": float(pin["sampling_deg"]), **extra}, + ) + t0 = time.time() + with patched(_api, "phaser_lmax_resolution", _pinned): + res = run_frf(model, data, cfg, capture_arf=True, verbose=0) + rank, sig, ang, ghost, margin = _truth_and_margin(res.arf, orbit) + + row = {"experiment": EXPERIMENT, "pdb": pdb, "arm": arm} + row.update(provenance()) + row.update({ + "spacegroup": str(data.spacegroup.hm), + "n_orbit": int(orbit.shape[0]), + "lmax": pin["lmax"], + "sampling_deg": pin["sampling_deg"], + "d_min_eff": pin["d_min_eff"], + "n_samples": int(res.arf.values.numel()), + "truth_rank": rank, + "truth_sigma": round(sig, 4), + "truth_angle_deg": round(ang, 3), + "best_ghost_sigma": round(ghost, 4), + "margin": round(margin, 4), + "map_max_sigma": round(float(res.map_max_sigma), 4), + "seconds": round(time.time() - t0, 1), + }) + row.update({f"flag_{k}": v for k, v in extra.items()}) + return row + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", required=True, choices=sorted(PHASER_PINNED)) + ap.add_argument("--arm", help="single arm (default: all)") + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--outdir", default=None) + args = ap.parse_args() + + arms = {args.arm: ARMS[args.arm]} if args.arm else ARMS + outdir = Path(args.outdir) if args.outdir else ( + Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT + ) + outdir.mkdir(parents=True, exist_ok=True) + csv_path = outdir / f"{EXPERIMENT}_{args.pdb}.csv" + + print(f"{args.pdb}: truth rank / margin per knockout arm", flush=True) + print(f"{'arm':<15}{'rank':>9}{'truth_sig':>11}{'ghost_sig':>11}{'margin':>9}", + flush=True) + failures = 0 + for arm, extra in arms.items(): + try: + row = run_arm(args.pdb, arm, extra, n_peaks=args.n_peaks) + except Exception as exc: + failures += 1 + print(f"{arm:<15} FAILED: {type(exc).__name__}: {exc}", flush=True) + continue + append_row(csv_path, row) + print(f"{arm:<15}{row['truth_rank']:>9}{row['truth_sigma']:>11.2f}" + f"{row['best_ghost_sigma']:>11.2f}{row['margin']:>+9.2f}", flush=True) + print(f"\nwrote {csv_path}", flush=True) + return 1 if failures == len(arms) else 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_inject_phaser_obs.py b/alignment_lab/diagnostics/frf_inject_phaser_obs.py new file mode 100644 index 00000000..00f1f040 --- /dev/null +++ b/alignment_lab/diagnostics/frf_inject_phaser_obs.py @@ -0,0 +1,211 @@ +"""Run our FRF on Phaser's observation intensities and see where truth lands. + +This is the closing experiment of the bisection. Everything downstream of the +per-reflection intensity is now verified exact against Phaser: + +* the SH-Bessel projection -- feeding Phaser's own prepared observations and its + own molecular-transform samples through ``bessel_sh_expand`` reproduces its + ``DataElmn`` and ``SearchElmn`` at correlation 1.0000, scale 1.0000 at 0 + degrees, residual 0.0000, once Phaser's 1e-3 cos-theta bucketing is replayed + (job 489517); +* the Wigner contraction, per-beta FFT, interpolation and adaptive sample list + -- Phaser's ``clmn`` through our evaluator gives r = 0.998 with an identical + argmax; +* the reciprocal frame (positions agree to 1e-8) and the unroll (our + ``epsilon_aware_unroll`` reproduces Phaser's point count exactly). + +What is NOT verified is the intensity attached to each position: ours correlates +with Phaser's at 0.988 (1AK5), 0.877 (2DQ6) and 0.711 (3GR5) -- and that ordering +is the performance ordering. So substituting Phaser's intensities into our +otherwise-unchanged pipeline is a decisive test rather than another correlation: +if truth reaches rank 0 on 3GR5, the remaining deficit is entirely in the +intensity computation and nothing else is left to look for. If it does not, there +is a defect outside everything measured so far. + +Alignment is by Miller index, not by position. Phaser dumps one row per selected +ASU reflection (``PHASER_TERMS_DUMP``); expanding each over the orbit ``h.W`` +gives the P1 index of every point that reflection contributes, and the Friedel +mate carries the same intensity. Reflections our engine keeps but Phaser did not +select are reported rather than silently dropped. + +Usage +----- + python -m diagnostics.frf_inject_phaser_obs --pdb 3GR5 \ + --dumps ../runs/encode_compare_489514/phaser/3GR5 +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import numpy as np +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from lab import FRFConfig, load_case, patched, run_frf # noqa: E402 +from lab.results import append_row, provenance # noqa: E402 +from diagnostics.frf_ghost_knockout import ( # noqa: E402 + PHASER_PINNED, _orbit_of_identity, _truth_and_margin, +) + +EXPERIMENT = "frf_inject_phaser_obs" + +#: Packing base for (h,k,l) -> one int64 key. Miller indices here stay well +#: inside +-1000 at these resolutions, and the base is checked at build time. +_BASE = 2048 + + +def _pack(hkl: torch.Tensor) -> torch.Tensor: + if int(hkl.abs().max()) >= _BASE // 2: + raise ValueError(f"Miller index {int(hkl.abs().max())} too large for base {_BASE}") + h, k, l = hkl[:, 0], hkl[:, 1], hkl[:, 2] + return ((h + _BASE // 2) * _BASE + (k + _BASE // 2)) * _BASE + (l + _BASE // 2) + + +def build_lut(terms_csv: Path, sym_mats: torch.Tensor): + """Phaser's per-reflection intensity, keyed by every P1 index it feeds. + + Each ASU row is expanded over the orbit ``h.W`` (Phaser's ``rotMiller`` is + ``rotsym[isym] * h`` with ``rotsym = W^T``, i.e. the row-vector convention) + and over the Friedel mate, which carries the same intensity because + ``|F(-h)| = |F(h)|`` and even-l-only projection is blind to the sign. + """ + d = np.loadtxt(terms_csv, delimiter=",", skiprows=1) + if d.ndim == 1: + d = d[None, :] + hkl = torch.from_numpy(d[:, 0:3]).to(torch.float64) + inten = torch.from_numpy(d[:, 9]).to(torch.float64) + + orbits = torch.einsum("kji,nj->nki", sym_mats.to(torch.float64), hkl) + orbits = orbits.round().to(torch.long).reshape(-1, 3) + vals = inten.unsqueeze(1).expand(-1, sym_mats.shape[0]).reshape(-1) + keys = torch.cat([_pack(orbits), _pack(-orbits)]) + vals = torch.cat([vals, vals]) + + uniq, inverse = torch.unique(keys, return_inverse=True) + lut = torch.zeros(uniq.numel(), dtype=torch.float64) + lut[inverse] = vals + # Two different ASU reflections mapping onto one P1 index would mean the + # orbit is still wrong -- the signature of the convention bug that was fixed. + hi = torch.full((uniq.numel(),), -1e300, dtype=torch.float64) + lo = torch.full((uniq.numel(),), 1e300, dtype=torch.float64) + hi.scatter_reduce_(0, inverse, vals, reduce="amax") + lo.scatter_reduce_(0, inverse, vals, reduce="amin") + n_conflict = int(((hi - lo).abs() > 1e-12 * hi.abs().clamp(min=1e-30)).sum()) + stats = {"n_asu_terms": int(hkl.shape[0]), + "n_lut_keys": int(uniq.numel()), + "n_lut_conflicts": n_conflict} + return uniq, lut, stats + + +def run_arm(pdb: str, dumps: Path, arm: str, *, n_peaks: int = 500) -> dict: + """One arm: ``baseline`` (our intensities) or ``phaser_intensity``.""" + from torchref.experimental.alignment.frf import api as _api + + pin = PHASER_PINNED[pdb] + model, data = load_case(pdb) + orbit = _orbit_of_identity(data) + sg = data.spacegroup.matrices.to(torch.float64).cpu() + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + rec_inv = torch.linalg.inv(rec) + + report = {"n_intercepted": 0} + original = _api.bessel_sh_expand + + if arm == "phaser_intensity": + keys, lut, lut_stats = build_lut(dumps / "terms.csv", sg) + + def injected(s, intensity, **kw): + if kw.get("zsymm", 1) <= 1: # calc side: untouched + return original(s, intensity, **kw) + report["n_intercepted"] += 1 + hkl = (s.to(torch.float64).cpu() @ rec_inv).round().to(torch.long) + k = _pack(hkl) + pos = torch.searchsorted(keys, k) + pos_c = pos.clamp(max=keys.numel() - 1) + found = keys[pos_c] == k + report["n_obs"] = int(s.shape[0]) + report["found_frac"] = float(found.to(torch.float64).mean()) + new = lut[pos_c].to(s.dtype).to(s.device) + return original(s[found], new[found], **kw) + + ctx_name, ctx_val = "bessel_sh_expand", injected + else: + lut_stats = {} + ctx_name, ctx_val = "bessel_sh_expand", original + + def _pinned(model_radius_A, d_min_data, lmax_cap=48): + return int(pin["lmax"]) + 1, float(pin["d_min_eff"]) + + cfg = FRFConfig(n_peaks=n_peaks, lmax_cap=int(pin["lmax"]), + extra={"grid_sampling_deg": float(pin["sampling_deg"])}) + t0 = time.time() + with patched(_api, "phaser_lmax_resolution", _pinned), \ + patched(_api, ctx_name, ctx_val): + res = run_frf(model, data, cfg, capture_arf=True, verbose=0) + rank, sig, ang, ghost, margin = _truth_and_margin(res.arf, orbit) + + row = {"experiment": EXPERIMENT, "pdb": pdb, "arm": arm} + row.update(provenance()) + row.update({ + "spacegroup": str(data.spacegroup.hm), + "lmax": pin["lmax"], "sampling_deg": pin["sampling_deg"], + "d_min_eff": pin["d_min_eff"], + "n_samples": int(res.arf.values.numel()), + "truth_rank": rank, "truth_sigma": round(sig, 4), + "truth_angle_deg": round(ang, 3), + "best_ghost_sigma": round(ghost, 4), "margin": round(margin, 4), + "seconds": round(time.time() - t0, 1), + }) + row.update(lut_stats) + row.update(report) + if arm == "phaser_intensity" and report["n_intercepted"] != 1: + raise RuntimeError( + f"{pdb}: intercepted {report['n_intercepted']} obs expansions, " + f"expected exactly 1 -- the injection hook is on the wrong call") + return row + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", required=True, choices=sorted(PHASER_PINNED)) + ap.add_argument("--dumps", required=True, + help="directory holding Phaser's terms.csv for this pdb") + ap.add_argument("--outdir", default=None) + ap.add_argument("--n-peaks", type=int, default=500) + args = ap.parse_args() + + outdir = Path(args.outdir) if args.outdir else ( + Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT) + outdir.mkdir(parents=True, exist_ok=True) + csv_path = outdir / f"{EXPERIMENT}_{args.pdb}.csv" + + print(f"{args.pdb}: truth rank with our vs Phaser's obs intensities", + flush=True) + print(f"{'arm':<18}{'rank':>8}{'truth_sig':>11}{'ghost_sig':>11}" + f"{'margin':>9}{'found':>8}", flush=True) + rows = [] + for arm in ("baseline", "phaser_intensity"): + rows.append(run_arm(args.pdb, Path(args.dumps), arm, + n_peaks=args.n_peaks)) + cols = {} + for r in rows: + cols.update({k: "" for k in r}) + for r in rows: + full = {**cols, **r} + append_row(csv_path, full) + ff = full.get("found_frac", "") + ff = f"{float(ff):>8.4f}" if ff != "" else f"{'-':>8}" + print(f"{r['arm']:<18}{r['truth_rank']:>8}{r['truth_sigma']:>11.2f}" + f"{r['best_ghost_sigma']:>11.2f}{r['margin']:>+9.2f}{ff}", + flush=True) + print(f"\nwrote {csv_path}", flush=True) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_map_compare.py b/alignment_lab/diagnostics/frf_map_compare.py new file mode 100644 index 00000000..ffb97edb --- /dev/null +++ b/alignment_lab/diagnostics/frf_map_compare.py @@ -0,0 +1,432 @@ +"""Compare our FRF array against Phaser's, element-wise, on one shared grid. + +Both engines evaluate a rotation function on Phaser's adaptive SO(3) sample +list, so with the sampling matched the two arrays are directly subtractable -- +index for index, no interpolation. That makes this a much sharper instrument +than comparing peak lists: a peak-list comparison only sees where the maxima +landed, whereas this sees the whole surface, including how much of the +disagreement is a smooth scale/offset (harmless -- peak *order* is invariant to +an affine map) versus genuine reshaping (not harmless). + +What is pinned, and what is not +------------------------------- +Three quantities are read from Phaser's own VERBOSE log and forced onto our +engine: the bandwidth ``lmax``, the resolution the expansion runs at, and the +SO(3) sampling step. They are pinned from the log rather than re-derived, +because they all descend from ``mean_radius()`` and our reimplementation of that +is ~4% off (see :func:`lab.phaser_match.phaser_mean_radius`). + +Everything on the observation and calc side keeps our production defaults -- +anisotropy correction, symmetry unroll, French-Wilson, shell-variance weights, +Wilson-B match, Oeffner vrms, bulk solvent, dense P1-box calc. That is +deliberate: with the coupled trio pinned, whatever disagreement remains is +attributable to that preprocessing stack, which is what we want localised. + +Usage +----- + python -m diagnostics.frf_map_compare --pdb 1DAW + python -m diagnostics.frf_map_compare --worklist-index 3 +""" + +from __future__ import annotations + +import argparse +import math +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from lab import BENCH_PDBS, FRFConfig, case_paths, load_case, patched, run_frf # noqa: E402 +from lab.phaser_match import ( # noqa: E402 + PATCHED_PHASER, + phaser_mean_radius_from_sampling, + phaser_sampling_from_dump, + load_phaser_frame, + load_phaser_map, + parse_phaser_log, + phaser_frf_params, + phaser_mean_radius, + run_patched_phaser, + write_keywords, +) +from lab.results import append_row, provenance # noqa: E402 + +EXPERIMENT = "frf_map_compare" + + +def run_phaser_side(pdb: str, work: Path, *, n_peaks: int = 20) -> dict: + """Run the patched binary and return its parameters plus its map. + + Returns a dict with ``angles`` (N,3 degrees, our sign convention), + ``values`` (N,), and the parsed log fields. + + Raises + ------ + RuntimeError + If the rotation search was short-circuited by the R-factor check, or the + dump is missing -- both of which Phaser reports as ``EXIT STATUS: + SUCCESS``, so they must be checked explicitly. + """ + pdb_path, mtz_path = case_paths(pdb) + dump = work / f"{pdb}_phaser_map.csv" + kw = write_keywords( + work, mtz_path=mtz_path, model_pdb=pdb_path, n_peaks=n_peaks, + root=f"{pdb}_frf", title=f"FRF map dump {pdb}", + ) + rc, seconds, log_path = run_patched_phaser(work, kw, dump_path=dump) + info = parse_phaser_log(log_path) + + if info.get("rotation_search_skipped"): + raise RuntimeError( + f"{pdb}: Phaser skipped the rotation search (R-factor short-circuit) " + f"-- see {log_path}" + ) + if not dump.exists(): + raise RuntimeError( + f"{pdb}: no FRF dump written (rc={rc}); is {PATCHED_PHASER} the " + f"instrumented binary? see {log_path}" + ) + for key in ("lmax", "sampling_deg"): + if key not in info: + raise RuntimeError(f"{pdb}: could not parse {key} from {log_path}") + + angles, values = load_phaser_map(dump) + # The logged sampling is rounded to 2 decimals and cannot rebuild the grid; + # the dumped beta step is exact. + info["sampling_deg_logged"] = info["sampling_deg"] + info["sampling_deg"] = phaser_sampling_from_dump(angles) + info.update(angles=angles, values=values, seconds=seconds, log=log_path) + return info + + +def run_our_side(pdb: str, info: dict, *, n_peaks: int = 500): + """Run our FRF with Phaser's bandwidth, resolution and sampling pinned. + + ``phaser_lmax_resolution`` is the single choke point through which the + bandwidth and resolution reach both the SH expansion and the dense calc + grid, so overriding it pins both consistently. + """ + from torchref.experimental.alignment.frf import api as _api + + model, data = load_case(pdb) + + lmax = int(info["lmax"]) + # Resolution the expansion runs at: LMAX_RESO when Phaser's cap bound, + # otherwise the selected high-resolution limit. + if info.get("lmax_reso_A") is not None and not info.get("all_data_to_limit", False): + d_min_eff = float(info["lmax_reso_A"]) + else: + d_min_eff = float(info.get("selected_d_min", 0.0)) or None + if d_min_eff is None: + raise RuntimeError(f"{pdb}: cannot determine Phaser's expansion resolution") + + def _pinned(model_radius_A, d_min_data, lmax_cap=48): + # Our bandwidth convention is L = lmax + 1. + return lmax + 1, d_min_eff + + cfg = FRFConfig( + n_peaks=n_peaks, + lmax_cap=lmax, + extra={"grid_sampling_deg": float(info["sampling_deg"])}, + ) + with patched(_api, "phaser_lmax_resolution", _pinned): + result = run_frf(model, data, cfg, capture_arf=True, verbose=0) + return model, data, result, d_min_eff + + +def _orbit_of_identity(data): + """Rotations equivalent to the deposited orientation under the point group. + + The search model is used unrotated, so "truth" is the identity -- but only + up to the crystal point group, and the operators must be taken to the + Cartesian frame the rotation function works in (mixing Cartesian rotations + with fractional operators is a metric error that inflates ghost counts). + """ + from lab.truth import symmetry_orbit + + symops = data.spacegroup.matrices.to(torch.float64).cpu() + recip = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + I = torch.eye(3, dtype=torch.float64) + return symmetry_orbit(I, symops, side="left", frame="cart", + reciprocal_basis=recip) + + +def _nearest(R_all: torch.Tensor, R_target: torch.Tensor): + """Index of the sample rotation closest to ``R_target``, and the angle.""" + tr = torch.einsum("nij,ij->n", R_all, R_target) + ang = torch.rad2deg(torch.arccos(((tr - 1.0) / 2.0).clamp(-1.0, 1.0))) + j = int(torch.argmin(ang)) + return j, float(ang[j]) + + +def _truth_rank(values: torch.Tensor, R_all: torch.Tensor, orbit: torch.Tensor): + """Rank of the best sample lying on the truth orbit. + + Returns ``(rank, sigma, angle_deg)``. The rank is the number of samples + scoring strictly higher, i.e. 0 means the rotation function put truth first. + """ + best = None + for k in range(orbit.shape[0]): + j, ang = _nearest(R_all, orbit[k]) + if best is None or values[j] > values[best[0]]: + best = (j, ang) + j, ang = best + rank = int((values > values[j]).sum()) + sig = float((values[j] - values.mean()) / values.std().clamp(min=1e-30)) + return rank, sig, ang + + +def compare(ours, phaser_angles, phaser_values, frame, data, *, topn: int = 20) -> dict: + """Compare two rotation functions in a common (PDB) frame. + + Element-wise comparison is impossible here and it is worth being explicit + about why: Phaser samples SO(3) on a grid laid out in the search model's + **principal frame**, we sample the identically-shaped grid in the PDB frame, + and the two are related by a rotation (``PR``, and ``axisrot``). The grids + therefore have the same pitch and the same point count but cover *different* + rotations -- nearest-neighbour offsets run about half a grid step. On peaks + only ~6-10 deg wide that annihilates any sample-wise correlation while + leaving the peak structure intact, so a whole-map Pearson r measures nothing + but the frame offset. Everything below is therefore computed on rotations, + via nearest-neighbour lookup, not on indices. + + The headline numbers are ``truth_rank_ours`` and ``truth_rank_phaser``: if + ours is much worse, the ghost problem is ours; if they agree, the ghosts are + inherent to the target function and no reimplementation will remove them. + """ + from torchref.experimental.alignment.frf.rotation_utils import ( + rotation_matrix_from_edmonds_euler_batch, + ) + + a = ours.arf + ov = a.values.to(torch.float64).cpu() + R_ours = rotation_matrix_from_edmonds_euler_batch( + a.alphas.to(torch.float64).cpu(), + a.betas.to(torch.float64).cpu(), + a.gammas.to(torch.float64).cpu(), + ) + + pv = phaser_values.to(torch.float64).cpu() + ang = phaser_angles.to(torch.float64).cpu() + R_grid = rotation_matrix_from_edmonds_euler_batch( + torch.deg2rad(ang[:, 0]), torch.deg2rad(ang[:, 1]), torch.deg2rad(ang[:, 2]), + ) + PR, AX = frame["PR"], frame["axisrot"] + # principal frame -> PDB frame (runMR_FRF.cc:542) + R_ph = torch.einsum("ij,njk,kl->nil", AX, R_grid, PR) + + osig = (ov - ov.mean()) / ov.std().clamp(min=1e-30) + psig = (pv - pv.mean()) / pv.std().clamp(min=1e-30) + out = { + "n_ours": int(ov.numel()), + "n_phaser": int(pv.numel()), + "grid_same_size": int(ov.numel() == pv.numel()), + "ours_max_sigma": float(osig.max()), + "phaser_max_sigma": float(psig.max()), + } + # Phaser's own statistics, as a check that our sigma means what theirs does. + if "stats" in frame: + st = frame["stats"] + out["phaser_max_sigma_logged"] = ( + (st["max"] - st["mean"]) / st["sigma"] if st["sigma"] else float("nan") + ) + + orbit = _orbit_of_identity(data) + out["n_orbit"] = int(orbit.shape[0]) + r_o, s_o, a_o = _truth_rank(ov, R_ours, orbit) + r_p, s_p, a_p = _truth_rank(pv, R_ph, orbit) + out.update({ + "truth_rank_ours": r_o, "truth_sigma_ours": s_o, "truth_angle_ours": a_o, + "truth_rank_phaser": r_p, "truth_sigma_phaser": s_p, "truth_angle_phaser": a_p, + "truth_rank_delta": r_o - r_p, + }) + + # --- Ghost anatomy ----------------------------------------------------- + # A ghost is a peak that outranks truth. The question is not "how many" but + # "what does the other engine see at exactly that rotation". Three outcomes, + # each implying a different fix: + # * Phaser has a peak there too, but weaker -> same physics, different + # weighting; find the term that suppresses it. + # * Phaser has nothing there -> we are manufacturing + # structure Phaser does not have. + # * Phaser has it just as strongly -> Phaser has the ghost too + # and wins somewhere downstream, not in the rotation function. + # `margin` is the discriminating power that matters: truth's sigma minus the + # strongest ghost's. Negative means the rotation function prefers a ghost. + tol = 1.5 * float(a.grid_sampling_deg) + + def _truth_mask(R_all): + keep = torch.zeros(R_all.shape[0], dtype=torch.bool) + for k in range(orbit.shape[0]): + tr = torch.einsum("nij,ij->n", R_all, orbit[k]) + ang = torch.rad2deg(torch.arccos(((tr - 1.0) / 2.0).clamp(-1.0, 1.0))) + keep |= ang <= tol + return keep + + m_o, m_p = _truth_mask(R_ours), _truth_mask(R_ph) + out["n_truthlike_ours"] = int(m_o.sum()) + out["n_truthlike_phaser"] = int(m_p.sum()) + + for tag, vals, sig, mask, Rs, other_v, other_R in ( + ("ours", ov, osig, m_o, R_ours, pv, R_ph), + ("phaser", pv, psig, m_p, R_ph, ov, R_ours), + ): + ghost_sig = sig.masked_fill(mask, float("-inf")) + gi = int(torch.argmax(ghost_sig)) + out[f"best_ghost_sigma_{tag}"] = float(sig[gi]) + out[f"margin_{tag}"] = float( + out[f"truth_sigma_{tag}"] - float(sig[gi]) + ) + # the same rotation, looked up in the other engine's map + j, dd = _nearest(other_R, Rs[gi]) + om, osd = other_v.mean(), other_v.std().clamp(min=1e-30) + out[f"best_ghost_{tag}_seen_by_other_sigma"] = float((other_v[j] - om) / osd) + out[f"best_ghost_{tag}_seen_by_other_rank"] = int((other_v > other_v[j]).sum()) + out[f"best_ghost_{tag}_lookup_angle"] = dd + + # Cross peak agreement: where do each engine's strongest peaks land in the + # other's map? + for label, (vs, Rs, vo, Ro) in { + "ph_in_ours": (pv, R_ph, ov, R_ours), + "ours_in_ph": (ov, R_ours, pv, R_grid if False else R_ph), + }.items(): + if label == "ours_in_ph": + # our rotations back into the principal frame for lookup + Rs_use = torch.einsum("ij,njk,kl->nil", AX.T, R_ours, PR.T) + vs_use, vo_use, Ro_use = ov, pv, R_grid + else: + Rs_use, vs_use, vo_use, Ro_use = R_ph, pv, ov, R_ours + top = torch.topk(vs_use, min(topn, vs_use.numel())).indices + ranks, sigs, angs = [], [], [] + vo_mean, vo_std = vo_use.mean(), vo_use.std().clamp(min=1e-30) + for i in top.tolist(): + j, dd = _nearest(Ro_use, Rs_use[i]) + ranks.append(int((vo_use > vo_use[j]).sum())) + sigs.append(float((vo_use[j] - vo_mean) / vo_std)) + angs.append(dd) + ranks_t = torch.tensor(ranks, dtype=torch.float64) + out[f"{label}_median_rank"] = float(ranks_t.median()) + out[f"{label}_median_sigma"] = float(torch.tensor(sigs).median()) + out[f"{label}_median_angle"] = float(torch.tensor(angs).median()) + out[f"{label}_frac_in_top{topn}"] = float((ranks_t < topn).to(torch.float64).mean()) + out[f"{label}_frac_above_5sig"] = float( + (torch.tensor(sigs) > 5.0).to(torch.float64).mean() + ) + return out + + +def run_case(pdb: str, outdir: Path, *, n_peaks: int = 500) -> dict: + """One structure: Phaser map, our map, comparison row.""" + work = outdir / "phaser" / pdb + t0 = time.time() + info = run_phaser_side(pdb, work) + frame = load_phaser_frame(work / f"{pdb}_phaser_map.csv") + model, data, ours, d_min_eff = run_our_side(pdb, info, n_peaks=n_peaks) + + row = {"experiment": EXPERIMENT, "pdb": pdb} + row.update(provenance()) + row.update({ + "phaser_lmax": info["lmax"], + "phaser_sampling_deg": info["sampling_deg"], + "phaser_sampling_deg_logged": info.get("sampling_deg_logged"), + "phaser_lmax_reso_A": info.get("lmax_reso_A"), + "phaser_all_data_to_limit": int(bool(info.get("all_data_to_limit"))), + "phaser_selected_d_min": info.get("selected_d_min"), + "phaser_selected_d_max": info.get("selected_d_max"), + "phaser_selected_n_refl": info.get("selected_n_refl"), + "phaser_n_samples_logged": info.get("n_samples"), + "phaser_seconds": round(info["seconds"], 1), + "pinned_d_min_eff_A": d_min_eff, + "ours_seconds": round(ours.seconds, 1), + }) + # Our radius formula vs Phaser's family, for the record. + row["our_mean_radius_A"] = float( + (model.xyz().to(torch.float64) - model.xyz().to(torch.float64).mean(0)) + .norm(dim=-1).mean().item() + ) + row["phaser_style_mean_radius_A"] = phaser_mean_radius(model) + pred = phaser_frf_params( + row["phaser_style_mean_radius_A"], + float(info.get("selected_d_min") or d_min_eff), + ) + row["predicted_lmax"] = pred.lmax + row["predicted_sampling_deg"] = pred.sampling_deg + # Phaser's own radius, inverted from its exact sampling step. + row["phaser_true_mean_radius_A"] = phaser_mean_radius_from_sampling( + info["sampling_deg"], d_min_eff, + ) + + row["high_order_axis"] = frame.get("high_order_axis") + row.update(compare(ours, info["angles"], info["values"], frame, data)) + row["total_seconds"] = round(time.time() - t0, 1) + + # Keep both arrays so a divergence-vs-beta plot needs no re-run. + npz = outdir / f"{pdb}_maps.pt" + torch.save( + { + "ours_values": ours.arf.values.cpu(), + "ours_alphas": ours.arf.alphas.cpu(), + "ours_betas": ours.arf.betas.cpu(), + "ours_gammas": ours.arf.gammas.cpu(), + "phaser_values": info["values"], + "phaser_angles": info["angles"], + }, + npz, + ) + return row + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", help="single structure") + ap.add_argument("--worklist-index", type=int, help="index into BENCH_PDBS") + ap.add_argument("--all", action="store_true", help="every benchmark structure") + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--outdir", default=None) + args = ap.parse_args() + + if args.pdb: + todo = [args.pdb] + elif args.worklist_index is not None: + todo = [BENCH_PDBS[args.worklist_index]] + elif args.all: + todo = list(BENCH_PDBS) + else: + ap.error("give --pdb, --worklist-index or --all") + + outdir = Path(args.outdir) if args.outdir else ( + Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT + ) + outdir.mkdir(parents=True, exist_ok=True) + csv_path = outdir / f"{EXPERIMENT}.csv" + + failures = 0 + for pdb in todo: + try: + row = run_case(pdb, outdir, n_peaks=args.n_peaks) + except Exception as exc: # keep the sweep going, but loudly + failures += 1 + print(f"[{pdb}] FAILED: {type(exc).__name__}: {exc}", flush=True) + continue + append_row(csv_path, row) + print( + f"[{pdb}] lmax={row['phaser_lmax']} samp={row['phaser_sampling_deg']:.4f}deg " + f"n={row['n_phaser']} | angle_dev={row['angle_max_dev_deg']:.2e} " + f"r={row.get('pearson_r', float('nan')):.4f} " + f"rho={row.get('spearman_r', float('nan')):.4f} " + f"resid={row.get('resid_rms_frac', float('nan')):.3f} " + f"top100={row.get('top100_overlap', float('nan')):.2f}", + flush=True, + ) + if failures: + print(f"\n{failures}/{len(todo)} cases FAILED", flush=True) + print(f"\nwrote {csv_path}", flush=True) + return 1 if failures == len(todo) else 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_normaliser_anatomy.py b/alignment_lab/diagnostics/frf_normaliser_anatomy.py new file mode 100644 index 00000000..ef49a843 --- /dev/null +++ b/alignment_lab/diagnostics/frf_normaliser_anatomy.py @@ -0,0 +1,265 @@ +"""Decompose the observation-normaliser gap that costs 3GR5 its rank. + +Injecting Phaser's per-reflection intensities into our otherwise-unchanged FRF +takes 3GR5 from rank 1995 / margin -1.06 to **rank 0 / margin +4.21** (job +489527), with every other stage already verified exact against Phaser. So the +whole remaining deficit is the normalised observation ``E^2``, and the question +is which part of it. + +Phaser's normaliser is (``DataB.cc:1106-1113``) + + sqrt_epsnSN[r] = sqrt( eps_n(h) * binAnisoFactor(bin, ANISO, SOLK, SOLB, K) ) + +a BEST-curve per-bin Sigma_N clamped to [0.5, 2] of BEST and corrected by a +fitted Wilson K/B, times an anisotropic tensor, plus a bulk-solvent term, times +the reflection multiplicity -- all refined. Ours is equal-count shell means of +``F^2`` (``preprocessing.py:30``) applied to amplitudes that have already had a +separately fitted overall anisotropy divided out (``align.py``: +``apply_overall_anisotropy``), then French-Wilson. + +So anisotropy is *not* simply missing on our side; it is removed upstream instead +of being folded into Sigma_N, which is equivalent only if the fitted tensor is +right. This splits ``log(E_phaser / E_ours)`` into pieces that can be fixed +independently: + +1. ``eps_n`` -- the multiplicity factor we omit entirely (its docstring in + ``build_lerf1_intensity`` claims it is "implicit in the symmetry reduction", + which is not the same thing as dividing by it). +2. the best possible **isotropic** model, a fine step function of ``|s|``. + Fitting a smooth Sigma_N curve cannot beat this, so it bounds what the curve + is worth. +3. a general quadratic form in Cartesian ``s`` -- residual **anisotropy** left + over after our own correction, reported as the eigenvalue spread of the + equivalent B tensor. This is the piece that matters most for a rotation + function: an angular error in the observed Patterson is exactly what a + rotation search is sensitive to, whereas a radial mis-scaling is not. + +Whatever variance survives all three is what only Phaser's refined per-bin +treatment could account for. + +Our side is *captured from the production path*, not reimplemented: the obs +``s`` and ``eEobs`` are spied out of the engine, so the anisotropy correction, +resolution window, shell edges and French-Wilson posterior are exactly the ones +production uses. + +Usage +----- + python -m diagnostics.frf_normaliser_anatomy --pdb 3GR5 \ + --dumps ../runs/encode_compare_489514/phaser/3GR5 +""" + +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +import numpy as np +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from lab import FRFConfig, load_case, patched # noqa: E402 +from lab.results import append_row, provenance # noqa: E402 +from diagnostics.frf_ghost_knockout import PHASER_PINNED # noqa: E402 +from diagnostics.frf_inject_phaser_obs import _pack # noqa: E402 + +EXPERIMENT = "frf_normaliser_anatomy" + + +def capture_ours(pdb: str): + """The obs ``s`` and ``eEobs`` the production engine actually expands. + + ``build_lerf1_intensity`` receives ``eEobs`` in the same order as the ``s`` + that reaches ``bessel_sh_expand`` -- both come from step 1's masked arrays + and nothing between them reorders or filters -- so spying on the two calls + gives an aligned pair. + """ + from torchref.experimental.alignment import align as _align + from torchref.experimental.alignment.frf import api as _api + + pin = PHASER_PINNED[pdb] + cap: dict = {} + orig_bessel = _api.bessel_sh_expand + orig_lerf1 = _api.build_lerf1_intensity + + def spy_bessel(s, vals, **kw): + if kw.get("zsymm", 1) > 1 and "s" not in cap: + cap["s"] = s.detach().cpu().to(torch.float64) + return orig_bessel(s, vals, **kw) + + def spy_lerf1(eEobs, centric, dfac=None, **kw): + if "eEobs" not in cap: + cap["eEobs"] = eEobs.detach().cpu().to(torch.float64) + cap["centric"] = centric.detach().cpu().clone() + cap["dfac"] = (torch.ones_like(cap["eEobs"]) if dfac is None + else dfac.detach().cpu().to(torch.float64)) + return orig_lerf1(eEobs, centric, dfac, **kw) + + model, data = load_case(pdb) + + def _pinned(model_radius_A, d_min_data, lmax_cap=48): + return int(pin["lmax"]) + 1, float(pin["d_min_eff"]) + + cfg = FRFConfig(n_peaks=5, lmax_cap=int(pin["lmax"])) + with patched(_api, "phaser_lmax_resolution", _pinned), \ + patched(_api, "bessel_sh_expand", spy_bessel), \ + patched(_api, "build_lerf1_intensity", spy_lerf1): + frf_in = _align._prepare_frf_inputs( + model, data, d_min=cfg.d_min, d_max=cfg.d_max, + n_shells=cfg.n_shells, verbose=0, + ) + _align._run_frf_separate_rotation( + model, data, frf_in, n_peaks=5, verbose=0, + lmax_cap=int(pin["lmax"]), + grid_sampling_deg=float(pin["sampling_deg"]), + ) + for k in ("s", "eEobs"): + if k not in cap: + raise RuntimeError(f"failed to capture {k} from the engine") + if cap["s"].shape[0] != cap["eEobs"].shape[0]: + raise RuntimeError( + f"capture misaligned: s has {cap['s'].shape[0]} rows, eEobs " + f"{cap['eEobs'].shape[0]} -- something between step 1 and the " + f"expansion filters the obs set") + return cap, data + + +def phaser_esqr_lut(terms_csv: Path, sym_mats: torch.Tensor): + """Phaser's ``Esqr``, keyed by every P1 index the reflection feeds.""" + d = np.loadtxt(terms_csv, delimiter=",", skiprows=1) + if d.ndim == 1: + d = d[None, :] + hkl = torch.from_numpy(d[:, 0:3]).to(torch.float64) + esqr = torch.from_numpy(d[:, 8]).to(torch.float64) + orb = torch.einsum("kji,nj->nki", sym_mats.to(torch.float64), hkl) + orb = orb.round().to(torch.long).reshape(-1, 3) + vals = esqr.unsqueeze(1).expand(-1, sym_mats.shape[0]).reshape(-1) + keys = torch.cat([_pack(orb), _pack(-orb)]) + vals = torch.cat([vals, vals]) + uniq, inv = torch.unique(keys, return_inverse=True) + lut = torch.zeros(uniq.numel(), dtype=torch.float64) + lut[inv] = vals + return uniq, lut + + +def _quadratic_design(s: torch.Tensor) -> torch.Tensor: + """``[1, sx^2, sy^2, sz^2, 2 sx sy, 2 sx sz, 2 sy sz]``: constant + 6 aniso.""" + x, y, z = s[:, 0], s[:, 1], s[:, 2] + return torch.stack([torch.ones_like(x), x * x, y * y, z * z, + 2 * x * y, 2 * x * z, 2 * y * z], dim=1) + + +def _fit(A: torch.Tensor, b: torch.Tensor): + sol = torch.linalg.lstsq(A, b.unsqueeze(1)).solution.squeeze(1) + return sol, float((b - A @ sol).var(unbiased=False)) + + +def run(pdb: str, dumps: Path) -> dict: + from torchref.experimental.alignment.frf.preprocessing import compute_epsilon + + cap, data = capture_ours(pdb) + sg = data.spacegroup.matrices.to(torch.float64).cpu() + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + rec_inv = torch.linalg.inv(rec) + + s = cap["s"] + hkl = (s @ rec_inv).round().to(torch.long) + keys, lut = phaser_esqr_lut(dumps / "terms.csv", sg) + k = _pack(hkl) + pos = torch.searchsorted(keys, k).clamp(max=keys.numel() - 1) + found = keys[pos] == k + + esqr_p = lut[pos][found] + esqr_o = (cap["eEobs"] ** 2)[found] + s = s[found] + hkl = hkl[found] + good = (esqr_p > 1e-12) & (esqr_o > 1e-12) + esqr_p, esqr_o, s, hkl = esqr_p[good], esqr_o[good], s[good], hkl[good] + + y = torch.log(esqr_p / esqr_o) # log ratio of NORMALISED intensities + n = int(y.numel()) + var_tot = float(y.var(unbiased=False)) + + eps = compute_epsilon(hkl, sg) + # Phaser divides intensity by eps_n and we do not, so its Esqr should be + # SMALLER by that factor: log ratio carries -log(eps). + y_eps = y + torch.log(eps) + var_eps = float(y_eps.var(unbiased=False)) + + smag = s.norm(dim=-1) + n_fine = 40 + q = torch.linspace(0, 1, n_fine + 1, dtype=torch.float64)[1:-1] + fbin = torch.bucketize(smag, torch.quantile(smag, q)) + y_iso = y_eps.clone() + for b in range(n_fine): + m = fbin == b + if int(m.sum()) > 1: + y_iso[m] = y_eps[m] - y_eps[m].mean() + var_iso = float(y_iso.var(unbiased=False)) + + A = _quadratic_design(s) + _, var_aniso = _fit(A, y_iso) + coef, _ = _fit(A, y_eps) + C = torch.tensor([[coef[1], coef[4], coef[5]], + [coef[4], coef[2], coef[6]], + [coef[5], coef[6], coef[3]]], dtype=torch.float64) + # log(E^2) = ... + s^T C s; on intensities a B-factor is exp(-B s^2 / 2), + # so the equivalent B is -2 C. + ev = torch.linalg.eigvalsh(-2.0 * C) + + row = {"experiment": EXPERIMENT, "pdb": pdb} + row.update(provenance()) + row.update({ + "spacegroup": str(data.spacegroup.hm), + "n_obs_ours": int(cap["s"].shape[0]), + "n_matched": n, + "matched_frac": round(float(found.to(torch.float64).mean()), 4), + "rms_log_ratio_pct": round(100.0 * float(y.std(unbiased=False)), 2), + "mean_log_ratio": round(float(y.mean()), 4), + "var_total": var_tot, + "frac_eps": round(1.0 - var_eps / max(var_tot, 1e-300), 4), + "frac_iso": round((var_eps - var_iso) / max(var_tot, 1e-300), 4), + "frac_aniso": round((var_iso - var_aniso) / max(var_tot, 1e-300), 4), + "frac_unexplained": round(var_aniso / max(var_tot, 1e-300), 4), + "aniso_B_min": round(float(ev[0]), 2), + "aniso_B_max": round(float(ev[2]), 2), + "aniso_B_spread": round(float(ev[2] - ev[0]), 2), + "n_eps_gt1": int((eps > 1.0001).sum()), + "eps_max": round(float(eps.max()), 1), + }) + return row + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", required=True, choices=sorted(PHASER_PINNED)) + ap.add_argument("--dumps", required=True) + ap.add_argument("--outdir", default=None) + args = ap.parse_args() + + outdir = Path(args.outdir) if args.outdir else ( + Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT) + outdir.mkdir(parents=True, exist_ok=True) + csv_path = outdir / f"{EXPERIMENT}_{args.pdb}.csv" + + row = run(args.pdb, Path(args.dumps)) + append_row(csv_path, row) + print(f"\n{args.pdb} ({row['spacegroup']}): variance of " + f"log(Esqr_phaser / Esqr_ours), n={row['n_matched']} " + f"({row['matched_frac']:.4f} of our obs matched)", flush=True) + print(f" rms disagreement {row['rms_log_ratio_pct']:6.2f}% " + f"mean log ratio {row['mean_log_ratio']:+.4f}", flush=True) + print(f" explained by eps_n {row['frac_eps']*100:6.2f}% " + f"({row['n_eps_gt1']} refl with eps>1, max {row['eps_max']})", flush=True) + print(f" explained by iso(|s|) {row['frac_iso']*100:6.2f}%", flush=True) + print(f" explained by anisotropy {row['frac_aniso']*100:6.2f}% " + f"(equivalent B {row['aniso_B_min']} .. {row['aniso_B_max']} A^2, " + f"spread {row['aniso_B_spread']})", flush=True) + print(f" unexplained {row['frac_unexplained']*100:6.2f}%", flush=True) + print(f"\nwrote {csv_path}", flush=True) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_prep_compare.py b/alignment_lab/diagnostics/frf_prep_compare.py new file mode 100644 index 00000000..676967fa --- /dev/null +++ b/alignment_lab/diagnostics/frf_prep_compare.py @@ -0,0 +1,293 @@ +"""Attribute the per-reflection gap between our FRF inputs and Phaser's. + +The stage-wise bisection put the divergence *before* the projection: reflection +positions agree to 1e-8, the SH-Bessel machinery reproduces Phaser's map from +Phaser's own coefficients at r = 0.998, and the radial band already matches +``nmax(l)`` -- but the intensities attached to those positions correlate at only +0.873 (2DQ6) against 0.988 (1AK5). Toggling ``use_epsilon``, French-Wilson, +shell-variance weights and the low-resolution cutoff moves that by <0.005, and +Phaser reports no tNCS, so ``V`` reduces to 1 on both sides. + +Two things remain unmeasured, and this runs both in one job. + +**Observation side.** Phaser builds +``intensity = cweight * (Esqr - V) / V^2 * DFAC^2`` with +``Esqr = (Feff / SIGMAN.sqrt_epsnSN)^2`` (DataMR.cc:930-945). The instrumented +binary now dumps every one of those terms per reflection, so a mismatch can be +attributed to a specific factor instead of only being visible in the product. +The two suspects are Phaser's smooth fitted ``Sigma_N`` against our equal-count +shell means, and ``DFAC`` (which we hard-wire to 1). + +**Calc side.** Never compared per reflection -- only post-projection via +``SearchElmn``. Phaser builds it in ``Ensemble::getELMNxR2``, a different +function from the observation one, so it needed its own dump. A difference here +would be invisible in every measurement made so far. + +Usage +----- + python -m diagnostics.frf_prep_compare --pdb 2DQ6 +""" + +from __future__ import annotations + +import argparse +import os +import subprocess +import sys +import time +from pathlib import Path + +import numpy as np +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from lab import FRFConfig, case_paths, load_case, patched # noqa: E402 +from lab.phaser_match import PATCHED_PHASER, write_keywords # noqa: E402 +from lab.results import append_row, provenance # noqa: E402 + +EXPERIMENT = "frf_prep_compare" + +#: Phaser's own bandwidth / resolution / sampling, from job 487737. +PINNED = { + "1DAW": dict(lmax=58, sampling_deg=6.233148, d_min_eff=5.70), + "1AK5": dict(lmax=70, sampling_deg=5.155428, d_min_eff=4.84), + "3K7M": dict(lmax=66, sampling_deg=5.464844, d_min_eff=5.39), + "4BX9": dict(lmax=96, sampling_deg=3.774036, d_min_eff=6.73), + "6G9X": dict(lmax=84, sampling_deg=4.314963, d_min_eff=5.92), + "2DQ6": dict(lmax=76, sampling_deg=4.772056, d_min_eff=6.36), + "3GR5": dict(lmax=66, sampling_deg=5.483967, d_min_eff=4.10), +} + + +def run_phaser_dumps(pdb: str, work: Path) -> dict: + """Run the instrumented binary, dumping observation and calc inputs.""" + work.mkdir(parents=True, exist_ok=True) + pdb_path, mtz_path = case_paths(pdb) + kw = write_keywords(work, mtz_path=mtz_path, model_pdb=pdb_path, + n_peaks=5, root=f"{pdb}_prep", title=f"prep {pdb}") + env = dict(os.environ) + env["PHASER_OBS_DUMP"] = str(work / "obs.csv") + env["PHASER_CALC_DUMP"] = str(work / "calc.csv") + env["PHASER_TERMS_DUMP"] = str(work / "terms.csv") + env["PHASER_SEARCH_ELMN_DUMP"] = str(work / "search_elmn.csv") + proc = subprocess.run([str(PATCHED_PHASER)], cwd=str(work), + input=kw.read_text(), capture_output=True, + text=True, timeout=5400, env=env) + (work / "run.log").write_text((proc.stdout or "") + (proc.stderr or "")) + for name in ("obs.csv", "calc.csv", "terms.csv"): + if not (work / name).exists(): + raise RuntimeError(f"{pdb}: {name} not written; see {work/'run.log'}") + return {"obs": work / "obs.csv", "calc": work / "calc.csv", + "terms": work / "terms.csv", + "search_elmn": work / "search_elmn.csv"} + + +def _polar_to_cart(r, th, ph) -> torch.Tensor: + return torch.tensor(np.stack( + [r * np.sin(th) * np.cos(ph), r * np.sin(th) * np.sin(ph), r * np.cos(th)], + axis=1)) + + +def capture_ours(pdb: str): + """Our per-reflection observation and calc inputs to the SH expansion. + + Both go through ``bessel_sh_expand``; the observation call is the one with + ``zsymm > 1`` (the calc side is deliberately never m-filtered). + """ + from torchref.experimental.alignment import align as _align + from torchref.experimental.alignment.frf import api as _api + + pin = PINNED[pdb] + cap: dict = {} + original = _api.bessel_sh_expand + + def spy(s, vals, **kw): + key = "obs" if kw.get("zsymm", 1) > 1 else "calc" + cap.setdefault(key, (s.detach().cpu().to(torch.float64), + vals.detach().cpu().to(torch.float64))) + return original(s, vals, **kw) + + model, data = load_case(pdb) + + def _pinned(model_radius_A, d_min_data, lmax_cap=48): + return int(pin["lmax"]) + 1, float(pin["d_min_eff"]) + + cfg = FRFConfig(n_peaks=20, lmax_cap=int(pin["lmax"])) + with patched(_api, "phaser_lmax_resolution", _pinned), \ + patched(_api, "bessel_sh_expand", spy): + frf_in = _align._prepare_frf_inputs( + model, data, d_min=cfg.d_min, d_max=cfg.d_max, + n_shells=cfg.n_shells, verbose=0, + ) + _align._run_frf_separate_rotation( + model, data, frf_in, n_peaks=20, verbose=0, + lmax_cap=int(pin["lmax"]), + grid_sampling_deg=float(pin["sampling_deg"]), + ) + return cap + + +def match_and_correlate(P: torch.Tensor, PV: torch.Tensor, + S: torch.Tensor, V: torch.Tensor, + *, n_sample: int = 4000, tol: float = 1e-6) -> dict: + """Match Phaser's points onto ours by position, then compare the values. + + Correlation is the statistic, not a ratio: these intensities are centred on + zero (``E^2 - 1``), so element-wise ratios are dominated by division by + near-zero and say nothing. + """ + g = torch.Generator().manual_seed(1) + k = min(n_sample, P.shape[0]) + sel = torch.randperm(P.shape[0], generator=g)[:k] + q, qv = P[sel], PV[sel] + + idx = torch.empty(k, dtype=torch.long) + dist = torch.empty(k) + for i in range(0, k, 500): + d2 = ((q[i:i + 500, None, :] - S[None, :, :]) ** 2).sum(-1) + mn = d2.min(1) + idx[i:i + 500] = mn.indices + dist[i:i + 500] = mn.values.sqrt() + + ok = dist < tol + out = {"n_phaser": int(P.shape[0]), "n_ours": int(S.shape[0]), + "matched_frac": float(ok.to(torch.float64).mean()), + "median_pos_dist": float(dist.median())} + if int(ok.sum()) < 50: + out["corr"] = float("nan") + return out + a, b = qv[ok], V[idx][ok] + ac, bc = a - a.mean(), b - b.mean() + out["corr"] = float((ac @ bc) / (ac.norm() * bc.norm()).clamp(min=1e-30)) + out["phaser_mean"] = float(a.mean()) + out["phaser_sd"] = float(a.std()) + out["ours_mean"] = float(b.mean()) + out["ours_sd"] = float(b.std()) + return out + + +def attribute_obs_terms(terms_csv: Path, pdb: str) -> dict: + """Compare Phaser's normalisation terms against ours, keyed by Miller index. + + Phaser emits one row per *selected reflection* with its Miller index, so + this is immune to the two hazards that broke the first attempt: the + ``reso(r) > LMAX_RESO`` gate drops entries from the HKL list, and + ``HKL_clustered::add`` buckets by theta, so no parallel array indexed + against that list can stay aligned. + + ``Esqr = (Feff / sqrt_epsnSN)^2`` is Phaser's normalised intensity. Ours is + ``E^2`` from equal-count shell means. Comparing the *normalisers* isolates + the Wilson treatment from everything else. + """ + d = np.loadtxt(terms_csv, delimiter=",", skiprows=1) + if d.ndim == 1: + d = d[None, :] + hkl = d[:, 0:3].astype(int) + Feff, sqrtSN, DFAC, V, cw, Esqr, inten, reso = (d[:, i] for i in range(3, 11)) + + out = { + "n_terms": int(d.shape[0]), + "dfac_mean": float(DFAC.mean()), "dfac_sd": float(DFAC.std()), + "dfac_is_unity": int(bool(np.allclose(DFAC, 1.0, atol=1e-6))), + "V_is_unity": int(bool(np.allclose(V, 1.0, atol=1e-6))), + } + # Self-consistency: Phaser's own identity must reproduce its own intensity. + rebuilt = cw * (Esqr - V) / (V ** 2) * (DFAC ** 2) + scale = float(np.abs(inten).max()) or 1.0 + out["rebuild_max_err"] = float(np.abs(rebuilt - inten).max() / scale) + # And that Esqr really is (Feff/sqrt_epsnSN)^2. + out["esqr_max_err"] = float( + np.abs((Feff / np.maximum(sqrtSN, 1e-30)) ** 2 - Esqr).max() + / (float(np.abs(Esqr).max()) or 1.0)) + + # Our normalised E^2 for the same Miller indices. + from torchref.experimental.alignment.frf.preprocessing import wilson_normalise + model, data = load_case(pdb) + B = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + our_hkl = data.hkl.to(torch.long).cpu() + F = data.F.to(torch.float64).abs().cpu() + smag = (our_hkl.to(torch.float64) @ B).norm(dim=-1) + E_obs, _ = wilson_normalise(F, smag, 20) + + off = 1024 + key = lambda t: ((t[:, 0] + off) * 4096 + (t[:, 1] + off)) * 4096 + (t[:, 2] + off) + lut = {int(k): i for i, k in enumerate(key(our_hkl))} + idx = np.array([lut.get(int(k), -1) for k in key(torch.tensor(hkl))]) + ok = idx >= 0 + out["terms_matched_frac"] = float(ok.mean()) + if int(ok.sum()) > 50: + ours_E2 = (E_obs[torch.tensor(idx[ok])] ** 2).to(torch.float64) + ph_E2 = torch.tensor(Esqr[ok]) + for nm, a, b in (("esqr", ph_E2, ours_E2), + ("normaliser", torch.tensor(sqrtSN[ok]), + (torch.tensor(Feff[ok]) / ours_E2.clamp(min=1e-30).sqrt()))): + ac, bc = a - a.mean(), b - b.mean() + out[f"corr_{nm}"] = float( + (ac @ bc) / (ac.norm() * bc.norm()).clamp(min=1e-30)) + out["esqr_ratio_median"] = float((ours_E2 / ph_E2.clamp(min=1e-30)).median()) + return out + + +def run_case(pdb: str, outdir: Path) -> dict: + t0 = time.time() + work = outdir / "phaser" / pdb + dumps = run_phaser_dumps(pdb, work) + ours = capture_ours(pdb) + + row = {"experiment": EXPERIMENT, "pdb": pdb} + row.update(provenance()) + + # Calc side is NOT position-matched: our calc lives on a cubic P1 box + # (s = hkl/a) and Phaser's on its own ensemble grid, so the two sampling + # sets have no reason to coincide -- a nearest-position match returns 0%. + # The calc comparison belongs at the projected (SearchElmn) level. + for side, csv_name in (("obs", "obs"),): + d = np.loadtxt(dumps[csv_name], delimiter=",", skiprows=1) + P = _polar_to_cart(d[:, 1], d[:, 2], d[:, 3]) + PV = torch.tensor(d[:, 4]) + S, V = ours[side] + res = match_and_correlate(P, PV, S, V) + row.update({f"{side}_{k}": v for k, v in res.items()}) + + row.update({f"term_{k}": v for k, v in attribute_obs_terms(dumps["terms"], pdb).items()}) + row["seconds"] = round(time.time() - t0, 1) + return row + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", required=True, choices=sorted(PINNED)) + ap.add_argument("--outdir", default=None) + args = ap.parse_args() + + outdir = Path(args.outdir) if args.outdir else ( + Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT) + outdir.mkdir(parents=True, exist_ok=True) + row = run_case(args.pdb, outdir) + append_row(outdir / f"{EXPERIMENT}.csv", row) + + print(f"\n=== {args.pdb} ===", flush=True) + print(" OBS n=%s/%s matched=%.3f corr=%.6f" + % (row.get("obs_n_ours"), row.get("obs_n_phaser"), + row.get("obs_matched_frac", float("nan")), + row.get("obs_corr", float("nan"))), flush=True) + print(" terms n=%s matched=%.3f | rebuild_err=%.2e esqr_err=%.2e" + % (row.get("term_n_terms"), row.get("term_terms_matched_frac", float("nan")), + row.get("term_rebuild_max_err", float("nan")), + row.get("term_esqr_max_err", float("nan"))), flush=True) + print(" Esqr corr(ours,phaser)=%.6f ratio median=%.4f" + % (row.get("term_corr_esqr", float("nan")), + row.get("term_esqr_ratio_median", float("nan"))), flush=True) + print(" DFAC unity=%s mean=%.5f sd=%.5f | V unity=%s | rebuild_err=%.2e" + % (row.get("term_dfac_is_unity"), row.get("term_dfac_mean", float("nan")), + row.get("term_dfac_sd", float("nan")), row.get("term_V_is_unity"), + row.get("term_rebuild_max_err", float("nan"))), flush=True) + print(" normaliser corr=%.6f" + % row.get("term_corr_normaliser", float("nan")), flush=True) + print(f"\nwrote {outdir / (EXPERIMENT + '.csv')}", flush=True) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_rank.py b/alignment_lab/diagnostics/frf_rank.py new file mode 100644 index 00000000..6b6929ea --- /dev/null +++ b/alignment_lab/diagnostics/frf_rank.py @@ -0,0 +1,88 @@ +"""Rank of the true orientation in the FRF peak list, per structure. + +The basic health check: rotate a deposited model by a seeded random rotation, +run the rotation function against that structure's own measured amplitudes, and +ask where the true orientation lands. Rank 0 means the top peak is correct. + +Peaks that outrank truth are the "ghosts" -- genuine correlations between the +model's self-Patterson and the crystal's intermolecular vectors, not noise. + +Usage:: + + python alignment_lab/diagnostics/frf_rank.py --pdb 1AK5 --trial 0 \ + --lmax-cap 64 --out-csv alignment_lab/runs/rank.csv +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, FRFConfig, ResultWriter, orbit_rank, # noqa: E402 + rotated_case, run_frf, seed_for) + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) + ap.add_argument("--trial", type=int, default=0) + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--d-min", type=float, default=4.0) + ap.add_argument("--d-max", type=float, default=15.0) + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--orbit-side", default="left", choices=["left", "right"]) + ap.add_argument("--orbit-frame", default="cart", choices=["cart", "frac"]) + ap.add_argument("--thr-deg", type=float, default=5.0) + ap.add_argument("--out-csv", default=None) + args = ap.parse_args() + + seed = seed_for(args.pdb, args.trial) + t0 = time.time() + rotated, data, R_true = rotated_case(args.pdb, seed) + load_s = time.time() - t0 + + sym = data.spacegroup.matrices.to(torch.float64).cpu() + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + cfg = FRFConfig(d_min=args.d_min, d_max=args.d_max, + n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) + res = run_frf(rotated, data, cfg) + rank, ang = orbit_rank(res.peaks, R_true, sym, side=args.orbit_side, + frame=args.orbit_frame, reciprocal_basis=rec, + thr_deg=args.thr_deg) + truth_sigma = float(res.peaks[rank].sigma) if rank >= 0 else float("nan") + + print(f"{args.pdb:6s} t{args.trial} seed={seed:<6d} {str(data.spacegroup):28s} " + f"n_ops={sym.shape[0]:2d} n_refl={data.hkl.shape[0]:7d} | " + f"rank={rank:4d} ang={ang:6.2f} truth_sig={truth_sigma:7.3f} " + f"map_max={res.map_max_sigma:7.3f} | frf={res.seconds:6.1f}s load={load_s:5.1f}s") + + if args.out_csv: + w = ResultWriter(args.out_csv, "frf_rank", + extra_fields=("truth_sigma", "map_max_sigma", "n_peaks", + "n_ghosts_above", "n_refl", + "frf_seconds", "load_seconds")) + w.write(pdb=args.pdb, seed=seed, trial=args.trial, + spacegroup=str(data.spacegroup), n_ops=int(sym.shape[0]), + truth_rank=rank, truth_angle_deg=round(ang, 4), + orbit_side=args.orbit_side, orbit_frame=args.orbit_frame, + lmax_cap=args.lmax_cap, d_min=args.d_min, d_max=args.d_max, + device="cpu", + truth_sigma=round(truth_sigma, 4), + map_max_sigma=round(res.map_max_sigma, 4), + n_peaks=len(res.peaks), + n_ghosts_above=(rank if rank >= 0 else len(res.peaks)), + n_refl=int(data.hkl.shape[0]), + frf_seconds=round(res.seconds, 2), + load_seconds=round(load_s, 2)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/ghost_origin.py b/alignment_lab/diagnostics/ghost_origin.py new file mode 100644 index 00000000..c1d945ec --- /dev/null +++ b/alignment_lab/diagnostics/ghost_origin.py @@ -0,0 +1,124 @@ +"""Where do the truth-beating peaks come from? Vary only the observations. + +Runs the identical engine on the identical rotated search model, changing +nothing but the observed amplitudes at the same Miller indices: + +``real`` + the deposited measurements. +``crystal`` + ``|F_calc|`` of the deposited model in its real space group -- noiseless, + solvent-free, complete. Ghosts surviving here are not noise, solvent, + measurement error or missing data. +``molecule`` + ``|F_calc|`` of the same model in **P1** -- the self-Patterson only, with + the symmetry mates removed. Ghosts vanishing here are intermolecular. + +Substituting observations is safe because the FRF reads only ``F``, ``F_sigma``, +``hkl``, ``centric``, ``cell`` and ``spacegroup`` from the dataset. + +Usage:: + + python alignment_lab/diagnostics/ghost_origin.py --pdb 3K7M --trial 0 \ + --out-csv alignment_lab/runs/ghosts.csv +""" + +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, FRFConfig, ResultWriter, load_case, # noqa: E402 + orbit_rank, random_rotation, run_frf, seed_for) +from lab.truth import angle_to_orbit, symmetry_orbit # noqa: E402 + + +def substituted_data(data, model, mode: str): + """Return a dataset whose ``F`` is replaced according to ``mode``.""" + if mode == "real": + return data + m = model.copy() + if mode == "molecule": + # NOTE: assign the space-group NAME. SpaceGroup is an nn.Module, so + # assigning the object is intercepted by nn.Module.__setattr__ and the + # property setter never runs -- a silent no-op that would leave the + # crystal symmetry in place and quietly invalidate this whole arm. + m.spacegroup = "P 1" + else: + m.spacegroup = data.spacegroup.hm + m.reset_cache() + out = data.copy() if hasattr(data, "copy") else data + F = m(out.hkl).abs().detach().to(out.F.dtype) + out.F = F + return out + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="3K7M", choices=list(BENCH_PDBS)) + ap.add_argument("--trial", type=int, default=0) + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--d-min", type=float, default=4.0) + ap.add_argument("--d-max", type=float, default=15.0) + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--orbit-side", default="left", choices=["left", "right"]) + ap.add_argument("--orbit-frame", default="cart", choices=["cart", "frac"]) + ap.add_argument("--modes", default="real,crystal,molecule") + ap.add_argument("--out-csv", default=None) + args = ap.parse_args() + + seed = seed_for(args.pdb, args.trial) + model, data = load_case(args.pdb) + R_true = random_rotation(seed) + rotated = model.copy().rotate(R_true.to(model.dtype_float), + center=model.xyz().mean(0)) + sym = data.spacegroup.matrices.to(torch.float64).cpu() + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + orbit = symmetry_orbit(R_true, sym, side=args.orbit_side, + frame=args.orbit_frame, reciprocal_basis=rec) + cfg = FRFConfig(d_min=args.d_min, d_max=args.d_max, + n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) + + print(f"=== {args.pdb} trial {args.trial} seed {seed} | {data.spacegroup} " + f"n_ops={sym.shape[0]} | lmax_cap={args.lmax_cap} ===") + print(f" {'obs':10s} {'rank':>6s} {'ghosts':>7s} {'truth_sig':>10s} {'map_max':>8s}") + + writer = None + if args.out_csv: + writer = ResultWriter(args.out_csv, "ghost_origin", + extra_fields=("obs_mode", "n_ghosts_above", + "truth_sigma", "map_max_sigma", + "n_peaks")) + for mode in args.modes.split(","): + mode = mode.strip() + sub = substituted_data(data, model, mode) + res = run_frf(rotated, sub, cfg) + rank, ang = orbit_rank(res.peaks, R_true, sym, side=args.orbit_side, + frame=args.orbit_frame, reciprocal_basis=rec) + # Peaks outranking truth, i.e. the ghosts this arm produces. + n_ghosts = rank if rank >= 0 else len(res.peaks) + truth_sigma = float(res.peaks[rank].sigma) if rank >= 0 else float("nan") + print(f" {mode:10s} {rank:6d} {n_ghosts:7d} {truth_sigma:10.3f} " + f"{res.map_max_sigma:8.3f}") + if writer: + writer.write(pdb=args.pdb, seed=seed, trial=args.trial, + spacegroup=str(data.spacegroup), n_ops=int(sym.shape[0]), + truth_rank=rank, truth_angle_deg=round(ang, 4), + orbit_side=args.orbit_side, orbit_frame=args.orbit_frame, + lmax_cap=args.lmax_cap, d_min=args.d_min, d_max=args.d_max, + device="cpu", obs_mode=mode, n_ghosts_above=n_ghosts, + truth_sigma=round(truth_sigma, 4), + map_max_sigma=round(res.map_max_sigma, 4), + n_peaks=len(res.peaks)) + if args.out_csv: + print(f" wrote {args.out_csv}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/phaser_headtohead.py b/alignment_lab/diagnostics/phaser_headtohead.py new file mode 100644 index 00000000..ba2d067a --- /dev/null +++ b/alignment_lab/diagnostics/phaser_headtohead.py @@ -0,0 +1,112 @@ +"""Head-to-head: our FRF vs Phaser on identical input. + +Both engines get the same rotated search model and the same reflections, and +both truth ranks are computed with the same orbit machinery, so the comparison +isolates the algorithms rather than the data handling. + +Read the caveats in :mod:`lab.phaser` before interpreting a result: Phaser +returns ~80-92k densely spaced samples, so a small "closest sample" angle is +expected regardless of whether it ranked truth well. + +Usage:: + + python alignment_lab/diagnostics/phaser_headtohead.py --pdb 1AK5 --trial 0 \ + --out-csv alignment_lab/runs/h2h.csv +""" + +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, FRFConfig, ResultWriter, orbit_rank, # noqa: E402 + rotated_case, run_frf, seed_for) +from lab import phaser as ph # noqa: E402 + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="1AK5", choices=list(BENCH_PDBS)) + ap.add_argument("--trial", type=int, default=0) + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--d-min", type=float, default=4.0) + ap.add_argument("--d-max", type=float, default=15.0) + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--orbit-side", default="left", choices=["left", "right"]) + ap.add_argument("--orbit-frame", default="cart", choices=["cart", "frac"]) + ap.add_argument("--workdir", default=None) + ap.add_argument("--timeout-s", type=int, default=5400) + ap.add_argument("--skip-phaser", action="store_true", + help="run only our engine (no phenix on this host)") + ap.add_argument("--out-csv", default=None) + args = ap.parse_args() + + seed = seed_for(args.pdb, args.trial) + work = Path(args.workdir or (Path(__file__).resolve().parents[1] / + "runs" / f"h2h_{args.pdb}_t{args.trial}")).resolve() + work.mkdir(parents=True, exist_ok=True) + + rotated, data, R_true = rotated_case(args.pdb, seed) + sym = data.spacegroup.matrices.to(torch.float64).cpu() + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + orbit_kw = dict(side=args.orbit_side, frame=args.orbit_frame, + reciprocal_basis=rec) + + print(f"=== {args.pdb} trial {args.trial} seed {seed} | {data.spacegroup} " + f"n_ops={sym.shape[0]} ===") + + cfg = FRFConfig(d_min=args.d_min, d_max=args.d_max, + n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) + res = run_frf(rotated, data, cfg) + our_rank, our_ang = orbit_rank(res.peaks, R_true, sym, **orbit_kw) + print(f" OURS rank={our_rank:5d} closest={our_ang:6.2f} deg " + f"({len(res.peaks)} peaks, {res.seconds:.1f}s, map max {res.map_max_sigma:.2f} sigma)") + + ph_rank, ph_ang, ph_n, ph_secs, ph_rc = -1, float("inf"), 0, 0.0, None + if not args.skip_phaser: + model_pdb = work / "rotated.pdb" + rotated.write_pdb(str(model_pdb)) + _, mtz_path = __import__("lab").case_paths(args.pdb) + kw = ph.write_frf_keywords(work, mtz_path=mtz_path, model_pdb=model_pdb) + ph_rc, ph_secs = ph.run_phaser(work, kw, timeout_s=args.timeout_s) + peaks = ph.parse_rlist(work / "phaser_frf.rlist") + ph_n = len(peaks) + if not peaks: + # rc==0 is not a success test; an empty list is the real signal. + print(f" PHASER produced no peaks (rc={ph_rc}); see {work}/phaser.stdout") + else: + ph_rank, ph_ang = ph.phaser_truth_rank(peaks, R_true, sym, **orbit_kw) + print(f" PHASER rank={ph_rank:5d} closest={ph_ang:6.2f} deg " + f"({ph_n} samples, {ph_secs:.0f}s)") + + if args.out_csv: + w = ResultWriter(args.out_csv, "phaser_headtohead", + extra_fields=("our_rank", "our_angle_deg", "our_seconds", + "our_n_peaks", "map_max_sigma", + "phaser_rank", "phaser_angle_deg", + "phaser_samples", "phaser_seconds", "phaser_rc")) + w.write(pdb=args.pdb, seed=seed, trial=args.trial, + spacegroup=str(data.spacegroup), n_ops=int(sym.shape[0]), + truth_rank=our_rank, truth_angle_deg=round(our_ang, 4), + orbit_side=args.orbit_side, orbit_frame=args.orbit_frame, + lmax_cap=args.lmax_cap, d_min=args.d_min, d_max=args.d_max, + device="cpu", + our_rank=our_rank, our_angle_deg=round(our_ang, 4), + our_seconds=round(res.seconds, 2), our_n_peaks=len(res.peaks), + map_max_sigma=round(res.map_max_sigma, 4), + phaser_rank=ph_rank, + phaser_angle_deg=(round(ph_ang, 4) if ph_n else ""), + phaser_samples=ph_n, phaser_seconds=round(ph_secs, 1), + phaser_rc=ph_rc if ph_rc is not None else "") + print(f" wrote {args.out_csv}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py new file mode 100644 index 00000000..ba2242ac --- /dev/null +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -0,0 +1,155 @@ +"""End-to-end pose recovery: does dropping the ML rescore cost anything? + +The rank-level evidence says the rescore only reorders (it leaves every peak's +orientation untouched), that final solutions are selected by R-factor rather +than by the rescore's own score, and that its reordering does not improve +whether truth reaches the translation stage. If all that holds, removing it +should be free -- but rank is not the deliverable, pose is, so this measures the +full pipeline. + +Arms (``--arms``): + +``m_letf1`` + current default. +``none`` + skip the rescore; raw FRF order straight into the translation search. +``none+subpeak`` + the same, plus quadratic sub-peak refinement -- sharpen each orientation + in place without reordering. + +Success mirrors the integration test: final coordinates within ``--success-deg`` +of canonical, modulo the crystal symmetry. + +Usage:: + + python alignment_lab/diagnostics/pose_recovery.py --pdb 1DAW --trial 0 \ + --arms m_letf1,none,none+subpeak --out-csv alignment_lab/runs/pose.csv +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +# NOTE: no global torch.set_grad_enabled(False) here, unlike the FRF-only +# diagnostics. This runs the full pipeline, whose joint refine and +# rigid-body polish are LBFGS -- they need autograd, and disabling it +# raises "element 0 of tensors does not require grad". + +from lab import (BENCH_PDBS, ResultWriter, load_case, random_rotation, # noqa: E402 + seed_for) + +ARMS = { + "m_letf1": dict(rescore_engine="m_letf1", subpeak_refine=False), + "sim": dict(rescore_engine="sim", subpeak_refine=False), + "none": dict(rescore_engine="none", subpeak_refine=False), + "none+subpeak": dict(rescore_engine="none", subpeak_refine=True), + "m_letf1+subpeak": dict(rescore_engine="m_letf1", subpeak_refine=True), +} + + +def residual_rotation_deg(aligned_xyz, canonical_xyz, symops) -> float: + """Smallest angle between the aligned-to-canonical rotation and any symop. + + Kabsch superposition, then compared against every symmetry operator -- + a solution differing from canonical by a crystal symmetry is correct. + """ + from torchref.experimental.alignment.frf.rotation_utils import ( + rotation_angular_distance_deg, + ) + + P = canonical_xyz.to(torch.float64) + Q = aligned_xyz.to(torch.float64) + Pc, Qc = P - P.mean(0), Q - Q.mean(0) + U, _, Vt = torch.linalg.svd(Qc.T @ Pc) + d = torch.sign(torch.det(U @ Vt)) + R = U @ torch.diag(torch.tensor([1.0, 1.0, d], dtype=torch.float64)) @ Vt + return min(float(rotation_angular_distance_deg(R, symops[k])) + for k in range(symops.shape[0])) + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) + ap.add_argument("--trial", type=int, default=0) + ap.add_argument("--arms", default="m_letf1,none,none+subpeak") + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--n-rotation-candidates", type=int, default=15) + ap.add_argument("--n-rotation-peaks", type=int, default=200) + ap.add_argument("--success-deg", type=float, default=8.0) + ap.add_argument("--verbose", type=int, default=0) + ap.add_argument("--out-csv", default=None) + args = ap.parse_args() + + from torchref.experimental.alignment.align import align_model_to_data + + seed = seed_for(args.pdb, args.trial) + model, data = load_case(args.pdb) + canonical_xyz = model.xyz().clone() + symops = data.spacegroup.matrices.to(torch.float64).cpu() + R_true = random_rotation(seed) + + print(f"=== {args.pdb} t{args.trial} seed={seed} {data.spacegroup} " + f"n_ops={symops.shape[0]} | success gate {args.success_deg} deg ===") + print(f" {'arm':16s} {'resid_deg':>10s} {'ok':>4s} {'seconds':>9s}") + + writer = None + if args.out_csv: + writer = ResultWriter(args.out_csv, "pose_recovery", + extra_fields=("arm", "rescore_engine", "subpeak_refine", + "residual_deg", "success", + "n_rotation_candidates", + "pipeline_seconds")) + for arm in [a.strip() for a in args.arms.split(",") if a.strip()]: + if arm not in ARMS: + raise SystemExit(f"unknown arm {arm!r}; choose from {sorted(ARMS)}") + flags = ARMS[arm] + # Fresh copy per arm: rotate/translate mutate in place, and the arms + # must start from identical coordinates to be comparable. + search = model.copy() + search.spacegroup = "P 1" + search = search.copy().rotate(R_true.to(model.dtype_float), + center=canonical_xyz.mean(0)) + t0 = time.time() + try: + aligned = align_model_to_data( + search, data, d_min=4.0, d_max=15.0, L=32, n_shells=20, + n_rotation_peaks=args.n_rotation_peaks, n_ml_refine=200, + do_translation=True, do_joint_refine=True, + n_rotation_candidates=args.n_rotation_candidates, + frf_lmax_cap=args.lmax_cap, verbose=args.verbose, **flags, + ) + resid = residual_rotation_deg(aligned.xyz(), canonical_xyz, symops) + err = "" + except Exception as exc: # a crashed arm must not read as a success + resid, err = float("nan"), f"{type(exc).__name__}: {exc}" + secs = time.time() - t0 + ok = (resid == resid) and resid <= args.success_deg + print(f" {arm:16s} {resid:10.2f} {('yes' if ok else 'NO'):>4s} {secs:9.1f}" + + (f" {err}" if err else "")) + if writer: + writer.write(pdb=args.pdb, seed=seed, trial=args.trial, + spacegroup=str(data.spacegroup), n_ops=int(symops.shape[0]), + truth_rank="", truth_angle_deg=(round(resid, 4) + if resid == resid else ""), + orbit_side="kabsch", orbit_frame="cart", + lmax_cap=args.lmax_cap, d_min=4.0, d_max=15.0, + device="cpu", arm=arm, + rescore_engine=flags["rescore_engine"], + subpeak_refine=int(flags["subpeak_refine"]), + residual_deg=(round(resid, 4) if resid == resid else ""), + success=int(bool(ok)), + n_rotation_candidates=args.n_rotation_candidates, + pipeline_seconds=round(secs, 1)) + if args.out_csv: + print(f" wrote {args.out_csv}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/rescore_rank.py b/alignment_lab/diagnostics/rescore_rank.py new file mode 100644 index 00000000..8f605622 --- /dev/null +++ b/alignment_lab/diagnostics/rescore_rank.py @@ -0,0 +1,118 @@ +"""Does the ML rescore improve the FRF ranking, or damage it? + +The FRF reliably puts the true orientation inside the top 20 on this benchmark, +yet end-to-end pose recovery succeeds about half the time. That points at the +rescore, so this measures it directly and in isolation. + +For each structure the FRF is run **once**, then every rescore arm is applied to +the *same* peak list -- including a ``none`` control that leaves the FRF order +untouched. The reported quantity is the paired change in the rank of truth +within the rescore window, so an arm that merely preserves a good input ranking +cannot be mistaken for one that improves it. + +Rows where truth was never in the window are recorded with ``delta`` empty: the +rescore had nothing to find, and scoring it there would measure the FRF. + +Usage:: + + python alignment_lab/diagnostics/rescore_rank.py --pdb 1AK5 --trial 0 \ + --engines none,m_letf1,sim --out-csv alignment_lab/runs/rescore.csv +""" + +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, FRFConfig, ResultWriter, orbit_rank, # noqa: E402 + paired_ranks, rotated_case, run_frf, run_rescore, seed_for) + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="1AK5", choices=list(BENCH_PDBS)) + ap.add_argument("--trial", type=int, default=0) + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--d-min", type=float, default=4.0) + ap.add_argument("--d-max", type=float, default=15.0) + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--n-refine", type=int, default=20, + help="rescore window: the top-N FRF peaks handed to the engine") + ap.add_argument("--engines", default="none,m_letf1,sim") + ap.add_argument("--subpeak-refine", action="store_true") + ap.add_argument("--orbit-side", default="left", choices=["left", "right"]) + ap.add_argument("--orbit-frame", default="cart", choices=["cart", "frac"]) + ap.add_argument("--verbose", type=int, default=0) + ap.add_argument("--out-csv", default=None) + args = ap.parse_args() + + seed = seed_for(args.pdb, args.trial) + rotated, data, R_true = rotated_case(args.pdb, seed) + sym = data.spacegroup.matrices.to(torch.float64).cpu() + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + orbit_kw = dict(side=args.orbit_side, frame=args.orbit_frame, + reciprocal_basis=rec) + + cfg = FRFConfig(d_min=args.d_min, d_max=args.d_max, + n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) + frf = run_frf(rotated, data, cfg, verbose=args.verbose) + rank_full, ang_full = orbit_rank(frf.peaks, R_true, sym, **orbit_kw) + + print(f"=== {args.pdb} t{args.trial} seed={seed} {data.spacegroup} " + f"n_ops={sym.shape[0]} | FRF rank={rank_full} ang={ang_full:.2f} " + f"({frf.seconds:.1f}s) | window={args.n_refine} ===") + if rank_full < 0 or rank_full >= args.n_refine: + print(f" NOTE truth is outside the rescore window " + f"(FRF rank {rank_full}); the rescore cannot recover it, so the " + f"deltas below measure nothing about the rescore.") + print(f" {'engine':10s} {'rank_in':>8s} {'rank_out':>9s} {'delta':>6s} " + f"{'ang_out':>8s} {'secs':>7s}") + + writer = None + if args.out_csv: + writer = ResultWriter(args.out_csv, "rescore_rank", + extra_fields=("engine", "n_refine", "rank_frf_full", + "rank_frf_window", "rank_rescored", + "delta", "truth_in_window", + "angle_rescored", "rescore_seconds", + "subpeak_refine")) + for engine in [e.strip() for e in args.engines.split(",") if e.strip()]: + res = run_rescore(frf.peaks, data, frf.inputs, engine=engine, + n_refine=args.n_refine, + subpeak_refine=args.subpeak_refine, + verbose=args.verbose) + pr = paired_ranks(frf.peaks, res.peaks, R_true, sym, + n_refine=args.n_refine, **orbit_kw) + delta = pr["delta"] + print(f" {engine:10s} {pr['rank_frf']:8d} {pr['rank_rescored']:9d} " + f"{('' if delta is None else f'{delta:+d}'):>6s} " + f"{pr['angle_rescored']:8.2f} {res.seconds:7.1f}") + if writer: + writer.write(pdb=args.pdb, seed=seed, trial=args.trial, + spacegroup=str(data.spacegroup), n_ops=int(sym.shape[0]), + truth_rank=pr["rank_rescored"], + truth_angle_deg=round(pr["angle_rescored"], 4), + orbit_side=args.orbit_side, orbit_frame=args.orbit_frame, + lmax_cap=args.lmax_cap, d_min=args.d_min, d_max=args.d_max, + device="cpu", engine=engine, n_refine=args.n_refine, + rank_frf_full=pr["rank_frf_full"], + rank_frf_window=pr["rank_frf"], + rank_rescored=pr["rank_rescored"], + delta=("" if delta is None else delta), + truth_in_window=int(pr["truth_in_window"]), + angle_rescored=round(pr["angle_rescored"], 4), + rescore_seconds=round(res.seconds, 2), + subpeak_refine=int(args.subpeak_refine)) + if args.out_csv: + print(f" wrote {args.out_csv}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/lab/__init__.py b/alignment_lab/lab/__init__.py new file mode 100644 index 00000000..6870487a --- /dev/null +++ b/alignment_lab/lab/__init__.py @@ -0,0 +1,60 @@ +"""Shared library for the alignment lab. + +Every diagnostic imports its primitives from here rather than re-deriving them. +The scripts this replaces carried ~37 copies of the rotation generator (in two +mutually incompatible variants), ~30 copies of the benchmark list, ~28 copies of +the CSV writer and ~20 copies of the rank-of-truth computation, several of which +disagreed with each other. One definition each, so two runs are comparable. +""" + +from .benchmark import ( + BENCH_PDBS, + PDB_STEMS, + REPO_ROOT, + case_paths, + load_case, + rotated_case, +) +from .truth import ( + orbit_rank, + random_rotation, + seed_for, + symmetry_orbit, +) +from .aniso import ( + ARMS as ANISO_ARMS, + aniso_arm, + fit_aniso_intensity_space, + tensor_report, +) +from .frf import FRFConfig, FRFResult, patched, run_frf +from .rescore import ENGINES, RescoreResult, paired_ranks, run_rescore +from .results import ResultWriter, append_row, provenance + +__all__ = [ + "BENCH_PDBS", + "PDB_STEMS", + "REPO_ROOT", + "case_paths", + "load_case", + "rotated_case", + "orbit_rank", + "random_rotation", + "seed_for", + "symmetry_orbit", + "ANISO_ARMS", + "aniso_arm", + "fit_aniso_intensity_space", + "tensor_report", + "FRFConfig", + "FRFResult", + "patched", + "run_frf", + "ENGINES", + "RescoreResult", + "paired_ranks", + "run_rescore", + "ResultWriter", + "append_row", + "provenance", +] diff --git a/alignment_lab/lab/aniso.py b/alignment_lab/lab/aniso.py new file mode 100644 index 00000000..bd1fb0b1 --- /dev/null +++ b/alignment_lab/lab/aniso.py @@ -0,0 +1,172 @@ +"""The overall-anisotropy correction: the production fit, and a corrected one. + +``sh.py:445 fit_overall_anisotropy`` regresses ``ln|F|^2 - ln<|F|^2>_shell`` on +``-2 pi^2 s.U.s`` by unweighted least squares **with no intercept**, and that is +the FRF's remaining defect (job 489540/489548). Three faults, all visible in its +output: + +* ``E[ln(I/)]`` is ``-gamma = -0.577`` for acentric reflections and + ``-gamma - ln 2 = -1.270`` for centric ones, not zero. With no intercept the + offset can only be absorbed by the quadratic form, which is why the fitted + tensor's SMALLEST B eigenvalue is 35-64 A^2 on every benchmark structure + instead of near zero. The centric part is worse than a constant: centric + reflections lie on the zones perpendicular to the symmetry axes, so the bias + is direction-dependent. +* ``clamp(min=1e-30)`` turns a vanishing amplitude into ``y ~ -69``; a handful of + those outweigh thousands of ordinary reflections in an unweighted fit. +* ``ln`` of a single-reflection intensity has variance ``pi^2/6`` (acentric) or + ``pi^2/2`` (centric) with a heavy left tail, so the fit is dominated by the + weak reflections carrying the least information. + +Raw fitted B eigenvalue spreads come out at 70 to 5461 A^2. +``symmetrize_anisotropy`` then projects onto the point-group-invariant subspace, +which annihilates the garbage where that subspace is small (cubic -> 1 DOF) and +leaves it where it is not (trigonal/hexagonal -> diag(lambda, lambda, mu), where +a uniaxial tensor along c is symmetry-allowed). + +One definition of the replacement lives here rather than in a diagnostic, since +several diagnostics need to A/B against it and it is the candidate production +change. +""" + +from __future__ import annotations + +import math +from contextlib import contextmanager + +import torch + +#: U (A^2) -> B (A^2). +B_PER_U = 8.0 * math.pi ** 2 + +#: Arm names accepted by :func:`aniso_arm`. +ARMS = ("production", "no_aniso", "iso_only", "fixed_fit") + + +def fit_aniso_intensity_space( + F_obs: torch.Tensor, + s_vec: torch.Tensor, + shell_idx: torch.Tensor, + centric: torch.Tensor, + P: int, + *, + min_count: int = 20, + n_iter: int = 12, +) -> torch.Tensor: + """Unbiased replacement for ``fit_overall_anisotropy``. + + Fits in INTENSITY space, where ``E[I/_shell] = c * exp(-2 pi^2 s.U.s)`` + holds exactly with no distributional correction. ``Var(I/)`` is 1 for + acentric and 2 for centric reflections, which gives the weights; a free + constant ``c`` absorbs the overall scale so it cannot leak into ``U``; + non-finite and non-positive amplitudes are dropped rather than clamped. + Gauss-Newton from ``U = 0``. + + Returns ``U`` in A^2 in the same convention as ``fit_overall_anisotropy`` + (applied as ``exp(+pi^2 s.U.s)``), so the caller's symmetrisation and + application are unchanged. + """ + valid = shell_idx >= 0 + F = F_obs[valid].to(torch.float64) + s = s_vec[valid].to(torch.float64) + idx = shell_idx[valid] + cen = centric[valid].bool() + ok = torch.isfinite(F) & (F > 0) + F, s, idx, cen = F[ok], s[ok], idx[ok], cen[ok] + + I = F * F + cnt = torch.zeros(P, dtype=torch.int64, device=F.device) + tot = torch.zeros(P, dtype=torch.float64, device=F.device) + cnt.index_add_(0, idx, torch.ones_like(idx)) + tot.index_add_(0, idx, I) + mean_I = (tot / cnt.clamp(min=1).to(torch.float64)).clamp(min=1e-30) + keep = (cnt >= min_count)[idx] + if int(keep.sum()) < 50: + return torch.zeros((3, 3), dtype=F_obs.dtype, device=F_obs.device) + r = I[keep] / mean_I[idx[keep]] + sk, cenk = s[keep], cen[keep] + + x, y, z = sk[:, 0], sk[:, 1], sk[:, 2] + quad = torch.stack([x * x, y * y, z * z, + 2 * x * y, 2 * x * z, 2 * y * z], dim=1) + A = torch.cat([torch.ones_like(x).unsqueeze(1), + -2.0 * (torch.pi ** 2) * quad], dim=1) + w = torch.where(cenk, torch.full_like(r, 0.5), torch.ones_like(r)) + + theta = torch.zeros(7, dtype=torch.float64, device=F.device) + for _ in range(n_iter): + m = torch.exp((A @ theta).clamp(min=-20.0, max=20.0)) + J = m.unsqueeze(1) * A + Jw = J * w.unsqueeze(1) + H = J.transpose(0, 1) @ Jw + g = Jw.transpose(0, 1) @ (r - m) + H = H + torch.eye(7, dtype=H.dtype, device=H.device) * 1e-12 * float( + torch.diagonal(H).abs().max().clamp(min=1e-30)) + theta = theta + torch.linalg.solve(H, g) + u = theta[1:] + return torch.tensor( + [[u[0], u[3], u[4]], [u[3], u[1], u[5]], [u[4], u[5], u[2]]], + dtype=F_obs.dtype, device=F_obs.device) + + +def tensor_report(U: torch.Tensor, tag: str) -> dict: + """B eigenvalues (A^2) of a U tensor, as result-row columns.""" + ev = torch.linalg.eigvalsh(U.to(torch.float64).cpu()) * B_PER_U + return {f"{tag}_B_min": round(float(ev[0]), 2), + f"{tag}_B_max": round(float(ev[2]), 2), + f"{tag}_B_spread": round(float(ev[2] - ev[0]), 2)} + + +@contextmanager +def aniso_arm(arm: str, data, *, d_min: float, d_max: float, captured: dict): + """Swap the anisotropy fit for the duration of one FRF call. + + ``captured`` receives the tensors actually fitted (``raw``, and ``fixed`` + when the arm uses the replacement) so a caller can report the artefact size + alongside the rank it costs. + + The ``fixed_fit`` arm needs ``centric``, which ``fit_overall_anisotropy`` + is not given. ``_prepare_frf_inputs`` masks ``F_obs`` and ``centric`` with + the same resolution window, so it is recomputed here from the same + ``d_min``/``d_max`` and checked against the amplitude count -- a mismatch + raises rather than silently misaligning. + """ + if arm not in ARMS: + raise ValueError(f"unknown aniso arm {arm!r}; expected one of {ARMS}") + from torchref.experimental.alignment import align as _align + + original = _align.fit_overall_anisotropy + rec = data.cell.reciprocal_basis_matrix.to(torch.float64) + smag = (data.hkl.to(torch.float64) @ rec).norm(dim=-1) + keep = (smag >= 1.0 / d_max) & (smag <= 1.0 / d_min) + centric = (data.centric[keep].to(torch.bool) + if hasattr(data, "centric") else None) + + def wrapped(F_obs, s_vec, shell_idx, **kw): + U = original(F_obs, s_vec, shell_idx, **kw) + captured.setdefault("raw", U.detach().clone()) + if arm == "production": + return U + if arm == "no_aniso": + return torch.zeros_like(U) + if arm == "iso_only": + # Radial part only; symmetrisation leaves lambda*I unchanged. + return torch.eye(3, dtype=U.dtype, device=U.device) * ( + torch.diagonal(U).sum() / 3.0) + if centric is None or centric.numel() != F_obs.shape[0]: + n = 0 if centric is None else centric.numel() + raise RuntimeError( + f"centric mask has {n} entries against {F_obs.shape[0]} " + f"amplitudes -- the resolution window assumed here " + f"([{d_min}, {d_max}] A) is not the engine's") + Ufix = fit_aniso_intensity_space( + F_obs, s_vec, shell_idx, centric.to(F_obs.device), + P=kw.get("P", 20), min_count=kw.get("min_count", 20)) + captured["fixed"] = Ufix.detach().clone() + return Ufix + + setattr(_align, "fit_overall_anisotropy", wrapped) + try: + yield + finally: + setattr(_align, "fit_overall_anisotropy", original) diff --git a/alignment_lab/lab/benchmark.py b/alignment_lab/lab/benchmark.py new file mode 100644 index 00000000..c0628465 --- /dev/null +++ b/alignment_lab/lab/benchmark.py @@ -0,0 +1,129 @@ +"""Benchmark structures and case loading. + +The ten deposited structures the alignment work is measured on. Paths are +resolved relative to this file, never hardcoded: the drivers inherited from the +old worktree pointed at an absolute path inside a stale checkout, so they read +data and code from a different tree than the one under test. +""" + +from __future__ import annotations + +from pathlib import Path +from typing import TYPE_CHECKING, Tuple + +if TYPE_CHECKING: # pragma: no cover - typing only + import torch + + from torchref.io.datasets.reflection_data import ReflectionData + from torchref.model import ModelFT + +REPO_ROOT = Path(__file__).resolve().parents[2] +TEST_FILES = REPO_ROOT / "tests" / "files" + +#: Benchmark structures, in the order the seed formula depends on. +#: ``seed_for`` uses ``BENCH_PDBS.index(pdb)``, so **inserting or reordering +#: entries changes every seed** and silently invalidates comparisons against +#: archived results. Append only. +BENCH_PDBS: Tuple[str, ...] = ( + "1DAW", # C2, small -- the fast control; use it for anything quick + "3E98", # P2_1, control (note: pandas reads the string "3E98" as a float) + "3A5V", # I422 + "3VRJ", + "1AK5", # P432, cubic ghost case + "3K7M", # P432, the primary ghost case + "3GR5", # P6_522 + "2DQ6", # P3_121, tNCS + "4BX9", # P4_32_12, large; the only benchmark entry carrying ANISOU + "6G9X", # large +) + +#: PDB filename stems, where they differ from the code. 1AK5 is the only one. +PDB_STEMS = {"1AK5": "1AK5_with_H"} + +#: Present in tests/files but deliberately excluded from BENCH_PDBS: +#: 5BOV (a single translation-function allocation OOMs an A100-40GB) and +#: 7L84 (no matching MTZ). + + +def case_paths(pdb: str) -> Tuple[Path, Path]: + """Return ``(pdb_path, mtz_path)`` for a benchmark code. + + Parameters + ---------- + pdb : str + Benchmark structure code, e.g. ``"1DAW"``. + + Returns + ------- + tuple of pathlib.Path + Model and reflection file paths. + + Raises + ------ + FileNotFoundError + If either file is missing, named so the caller sees which one. + """ + stem = PDB_STEMS.get(pdb, pdb) + pdb_path = TEST_FILES / "pdb" / f"{stem}.pdb" + mtz_path = TEST_FILES / "mtz" / f"{pdb}.mtz" + for p in (pdb_path, mtz_path): + if not p.exists(): + raise FileNotFoundError(f"{pdb}: missing {p}") + return pdb_path, mtz_path + + +def load_case(pdb: str, device: str = "cpu") -> Tuple["ModelFT", "ReflectionData"]: + """Load the deposited model and its reflections. + + Parameters + ---------- + pdb : str + Benchmark structure code. + device : str, optional + Torch device for both objects. Default ``"cpu"``. + + Returns + ------- + tuple + ``(model, data)``. + """ + from torchref.io.datasets.reflection_data import ReflectionData + from torchref.model import ModelFT + + pdb_path, mtz_path = case_paths(pdb) + model = ModelFT(device=device).load_pdb(str(pdb_path)) + data = ReflectionData(device=device).load_mtz(str(mtz_path)) + return model, data + + +def rotated_case( + pdb: str, seed: int, device: str = "cpu", +) -> Tuple["ModelFT", "ReflectionData", "torch.Tensor"]: + """Load a case and rotate a copy of the model by a seeded random rotation. + + The returned model is a **copy**: ``Model.rotate`` mutates in place and + returns ``self``, so rotating the loaded model directly would also move the + reference a caller may want to compare against. + + Parameters + ---------- + pdb : str + Benchmark structure code. + seed : int + Seed for :func:`~alignment_lab.lab.truth.random_rotation`. + device : str, optional + Torch device. Default ``"cpu"``. + + Returns + ------- + tuple + ``(rotated_model, data, R_true)`` with ``R_true`` in float64. + """ + from .truth import random_rotation + + model, data = load_case(pdb, device=device) + R_true = random_rotation(seed) + rotated = model.copy().rotate( + R_true.to(model.dtype_float), center=model.xyz().mean(0), + ) + return rotated, data, R_true diff --git a/alignment_lab/lab/frf.py b/alignment_lab/lab/frf.py new file mode 100644 index 00000000..e42ca428 --- /dev/null +++ b/alignment_lab/lab/frf.py @@ -0,0 +1,166 @@ +"""Run the FRF and capture the full rotation function, not just the peak list. + +Every rank/ghost diagnostic needs the dense adaptive sample list as well as the +peaks, and the engine only returns the peaks. The capture below wraps the +engine's search entry point for the duration of one call; nine scripts each +carried their own copy of this monkeypatch. +""" + +from __future__ import annotations + +from contextlib import contextmanager +from dataclasses import asdict, dataclass, field +from typing import Any, Dict, Optional, Tuple + +import torch + + +@dataclass +class FRFConfig: + """Engine settings for one FRF evaluation. + + Collected into one object so a diagnostic passes a single config around and + the settings can be written into the result row verbatim. + """ + + d_min: float = 4.0 + d_max: float = 15.0 + n_shells: int = 20 + n_peaks: int = 500 + lmax_cap: int = 48 + dense_pad: float = 2.0 + extra: Dict[str, Any] = field(default_factory=dict) + + def as_row(self) -> Dict[str, Any]: + """Config fields for a result row (``extra`` flattened out).""" + d = asdict(self) + d.pop("extra") + d.update(self.extra) + return d + + +@contextmanager +def patched(module: Any, name: str, replacement: Any): + """Temporarily replace ``module.name``, restoring it on exit. + + The "swap one engine internal and re-measure the rank" pattern -- used for + the dense-grid and box-construction experiments -- always needs the original + restored even when the body raises. + + Parameters + ---------- + module : module or object + Namespace holding the attribute. + name : str + Attribute name. + replacement : Any + Temporary value. + """ + original = getattr(module, name) + setattr(module, name, replacement) + try: + yield original + finally: + setattr(module, name, original) + + +@dataclass +class FRFResult: + """Outcome of one FRF evaluation. + + Attributes + ---------- + peaks : list + ``RotationPeak`` list, descending score. + arf : AdaptiveRotationFunction or None + The full adaptive sample list, when captured. + sigma : torch.Tensor or None + ``arf.values`` standardised to zero mean / unit sd -- the scale peak + heights are quoted in. + seconds : float + Wall time of the search call. + inputs : FRFInputs or None + The prepared observations (``F_obs``/``hkl``/``s_mag``/``centric``/``ll``). + The rescore consumes these, so keeping them lets a rescore run reuse one + FRF evaluation instead of recomputing it. + """ + + peaks: list + arf: Optional[Any] + sigma: Optional[torch.Tensor] + seconds: float + inputs: Optional[Any] = None + + @property + def map_max_sigma(self) -> float: + """Largest value of the standardised rotation function.""" + return float(self.sigma.max()) if self.sigma is not None else float("nan") + + +def run_frf( + model, + data, + cfg: Optional[FRFConfig] = None, + *, + capture_arf: bool = True, + verbose: int = 0, +) -> FRFResult: + """Run the separated FRF on an already-rotated search model. + + Parameters + ---------- + model : ModelFT + Search model, already in the orientation to be scored. + data : ReflectionData + Observed reflections. + cfg : FRFConfig, optional + Engine settings. Defaults to :class:`FRFConfig`. + capture_arf : bool, optional + Also return the dense adaptive sample list. Default True. + verbose : int, optional + Engine verbosity. Default 0. + + Returns + ------- + FRFResult + """ + import time + + from torchref.experimental.alignment import align as _align + from torchref.experimental.alignment.frf import api as _api + + cfg = cfg or FRFConfig() + captured: Dict[str, Any] = {} + + def _wrapped(*args, **kwargs): + arf, peaks = _original(*args, **kwargs) + captured["arf"] = arf + return arf, peaks + + frf_inputs = _align._prepare_frf_inputs( + model, data, + d_min=cfg.d_min, d_max=cfg.d_max, n_shells=cfg.n_shells, verbose=verbose, + ) + + t0 = time.time() + if capture_arf: + _original = _api.phaser_rotation_search + with patched(_api, "phaser_rotation_search", _wrapped): + peaks = _align._run_frf_separate_rotation( + model, data, frf_inputs, n_peaks=cfg.n_peaks, verbose=verbose, + lmax_cap=cfg.lmax_cap, dense_pad=cfg.dense_pad, **cfg.extra, + ) + else: + peaks = _align._run_frf_separate_rotation( + model, data, frf_inputs, n_peaks=cfg.n_peaks, verbose=verbose, + lmax_cap=cfg.lmax_cap, dense_pad=cfg.dense_pad, **cfg.extra, + ) + seconds = time.time() - t0 + + arf = captured.get("arf") + sigma = None + if arf is not None: + vals = arf.values.to(torch.float64) + sigma = (vals - vals.mean()) / vals.std().clamp(min=1e-30) + return FRFResult(peaks=peaks, arf=arf, sigma=sigma, seconds=seconds, + inputs=frf_inputs) diff --git a/alignment_lab/lab/phaser.py b/alignment_lab/lab/phaser.py new file mode 100644 index 00000000..9976675b --- /dev/null +++ b/alignment_lab/lab/phaser.py @@ -0,0 +1,204 @@ +"""Phaser oracle adapter. + +Phaser is the reference the FRF is measured against, so this wraps invoking it +and reading its peaks back in our conventions. Previously these helpers lived +inside a pytest module and six scripts imported them from there. + +Three details are load-bearing and easy to lose: + +* **Convention.** ``R_ours = R_phaser.T``. Calibrated empirically in P1, where + ``n_ops == 1`` leaves no orbit ambiguity to hide a transpose error. +* **``PEAKS ROT SELECT ALL``** with clustering off. Phaser otherwise merges + symmetry equivalents before we can rank them. Expect ~80-92k samples with a + median nearest-neighbour spacing under 1 degree, so "the nearest sample is + within 1 degree" is not evidence of anything on its own. +* **Phaser exits 0 on fatal input errors.** The return code is not a success + test; an empty peak list is the real signal. Keyword files also need + **absolute** paths. +""" + +from __future__ import annotations + +import re +import subprocess +import time +from dataclasses import dataclass +from pathlib import Path +from typing import List, Optional, Sequence, Tuple + +import torch + +_SOLU_TRIAL_RE = re.compile( + r"SOLU\s+TRIAL\s+ENSEMBLE\s+\S+\s+" + r"EULER\s+([-+\d.]+)\s+([-+\d.]+)\s+([-+\d.]+)\s+" + r"RF\s+([-+\d.eE]+)\s+RFZ\s+([-+\d.eE]+)", + re.IGNORECASE, +) + + +@dataclass +class PhaserPeak: + """One Phaser FRF peak. Euler angles in **degrees**, Edmonds ZYZ.""" + + alpha_deg: float + beta_deg: float + gamma_deg: float + rf: float + rfz: float + + +def write_frf_keywords( + work: Path, *, mtz_path: Path, model_pdb: Path, + f_label: str = "FP", sigf_label: str = "SIGFP", root: str = "phaser_frf", +) -> Path: + """Write an MR_FRF keyword file. Paths are resolved to absolute. + + Parameters + ---------- + work : Path + Working directory; created if absent. + mtz_path, model_pdb : Path + Inputs. Relative paths are resolved -- Phaser fails on relative ones. + f_label, sigf_label : str, optional + MTZ column labels. + root : str, optional + Phaser output root. + + Returns + ------- + Path + The keyword file. + """ + work.mkdir(parents=True, exist_ok=True) + kw = work / f"{root}.kw" + kw.write_text( + f"TITLE FRF rotation ranking\n" + f"MODE MR_FRF\n" + f"HKLIN {Path(mtz_path).resolve()}\n" + f"LABIN F={f_label} SIGF={sigf_label}\n" + f"ENSEMBLE search PDB {Path(model_pdb).resolve()} IDENT 1.0\n" + f"COMPOSITION BY AVERAGE\n" + f"SEARCH ENSEMBLE search\n" + f"PEAKS ROT SELECT ALL\n" + f"PEAKS ROT CLUSTER OFF\n" + f"PEAKS ROT LEVEL 0\n" + f"ROOT {root}\n" + ) + return kw + + +def run_phaser(work: Path, kw_path: Path, timeout_s: int = 5400) -> Tuple[int, float]: + """Run ``phenix.phaser`` on a keyword file. + + Returns + ------- + tuple + ``(returncode, seconds)``; ``-1`` on timeout. **A zero return code does + not mean success** -- check that :func:`parse_rlist` found peaks. + """ + t0 = time.time() + try: + proc = subprocess.run( + ["phenix.phaser"], cwd=str(work), input=kw_path.read_text(), + capture_output=True, text=True, timeout=timeout_s, + ) + (work / "phaser.stdout").write_text(proc.stdout or "") + (work / "phaser.stderr").write_text(proc.stderr or "") + rc = proc.returncode + except subprocess.TimeoutExpired: + rc = -1 + return rc, time.time() - t0 + + +def parse_rlist(path: Path) -> List[PhaserPeak]: + """Parse ``SOLU TRIAL`` lines from a Phaser ``.rlist``. + + Returns an empty list when the file is absent, which is also what a failed + run looks like -- see the note on exit codes in the module docstring. + """ + path = Path(path) + if not path.exists(): + return [] + peaks: List[PhaserPeak] = [] + for line in path.read_text().splitlines(): + if "SOLU TRIAL" not in line.upper(): + continue + m = _SOLU_TRIAL_RE.search(line) + if m is None: + continue + a, b, g, rf, rfz = (float(x) for x in m.groups()) + peaks.append(PhaserPeak(a, b, g, rf, rfz)) + peaks.sort(key=lambda p: p.rfz, reverse=True) + return peaks + + +def euler_deg_to_matrices(peaks: Sequence[PhaserPeak]) -> torch.Tensor: + """Stack Phaser peaks as rotation matrices in **Phaser's** frame. + + Edmonds ZYZ active rotation ``R = Rz(alpha) Ry(beta) Rz(gamma)``. Apply + :func:`to_our_frame` before comparing against our orbit. + + Returns + ------- + torch.Tensor + ``(n, 3, 3)`` float64. Empty ``(0, 3, 3)`` for an empty input. + """ + if not peaks: + return torch.zeros((0, 3, 3), dtype=torch.float64) + a = torch.tensor([p.alpha_deg for p in peaks], dtype=torch.float64).deg2rad() + b = torch.tensor([p.beta_deg for p in peaks], dtype=torch.float64).deg2rad() + g = torch.tensor([p.gamma_deg for p in peaks], dtype=torch.float64).deg2rad() + ca, sa, cb, sb, cg, sg = a.cos(), a.sin(), b.cos(), b.sin(), g.cos(), g.sin() + return torch.stack([ + torch.stack([ca * cb * cg - sa * sg, -ca * cb * sg - sa * cg, ca * sb], dim=-1), + torch.stack([sa * cb * cg + ca * sg, -sa * cb * sg + ca * cg, sa * sb], dim=-1), + torch.stack([-sb * cg, sb * sg, cb], dim=-1), + ], dim=-2) + + +def to_our_frame(R_phaser: torch.Tensor) -> torch.Tensor: + """Convert Phaser-frame rotations to ours: ``R_ours = R_phaser.T``. + + Calibrated in P1 (``n_ops == 1``), where no symmetry orbit can mask a + transposition. Do not re-derive this per structure: with ~80-92k samples, + both conventions match *something* within a degree. + """ + return R_phaser.transpose(-1, -2) + + +def phaser_truth_rank( + peaks: Sequence[PhaserPeak], + R_true: torch.Tensor, + symops: torch.Tensor, + *, + reciprocal_basis: Optional[torch.Tensor] = None, + frame: str = "cart", + side: str = "left", + thr_deg: float = 5.0, +) -> Tuple[int, float]: + """Rank of the true orientation in Phaser's own peak list. + + Uses the same orbit machinery as our engine, after mapping Phaser's frame + onto ours, so the two ranks are directly comparable. + + Returns + ------- + tuple + ``(rank, best_angle_deg)``; rank ``-1`` if unmatched. + """ + from .truth import angle_to_orbit, symmetry_orbit + + if not peaks: + return -1, float("inf") + orbit = symmetry_orbit( + R_true, symops, side=side, frame=frame, reciprocal_basis=reciprocal_basis, + ) + R_ours = to_our_frame(euler_deg_to_matrices(peaks)) + rank, best = -1, float("inf") + for i in range(R_ours.shape[0]): + ang = angle_to_orbit(R_ours[i], orbit) + if ang < best: + best = ang + if ang <= thr_deg and rank < 0: + rank = i + return rank, best diff --git a/alignment_lab/lab/phaser_match.py b/alignment_lab/lab/phaser_match.py new file mode 100644 index 00000000..86b0ec08 --- /dev/null +++ b/alignment_lab/lab/phaser_match.py @@ -0,0 +1,441 @@ +"""Reproduce Phaser's FRF parameter chain exactly, and read back its map. + +Every number the rotation function depends on -- spherical-harmonic bandwidth, +the resolution actually expanded, and the SO(3) sampling step -- is *derived* by +Phaser from one quantity: ``mean_radius()``. This module implements that chain +verbatim from the 1.20 source so our engine can be pinned to the same values, +and parses the patched binary's log/dump so the derivation can be checked +against what Phaser actually did rather than trusted. + +Source anchors (PHENIX 1.20-4459, ``modules/phaser/codebase/phaser``): + +* ``lib/xyz_weight.cc:178`` ``mean_radius()`` -- the mean of the three + principal-axis **semi-extents of the bounding box**, NOT the mean atomic + distance from the centroid. These differ by ~25-30% on a protein. +* ``run/runMR_FRF.cc:406-410`` bandwidth:: + + sphereOuter = 2 * mean_radius + LMAX = ceil(2*pi*sphereOuter / HiRes) # round UP to even + LMAX = min(LMAX, DEF_CLMN_LMAX = 100) + +* ``run/runMR_FRF.cc:411-419`` resolution -- coarsened **only** when the cap + binds:: + + LMAX_RESO = (LMAX == 100) ? 2*pi*sphereOuter/LMAX : HiRes + +* ``run/runMR_FRF.cc:469-474`` sampling -- likewise keyed on the cap:: + + SAMP_RESO = (LMAX == 100) ? LMAX_RESO : HiRes + sampling = 2 * degrees(atan(SAMP_RESO / (4 * mean_radius))) + +The three are one coupled system: when the bandwidth saturates, the resolution +and the angular step coarsen together so the expansion is never asked to carry +detail it cannot represent. +""" + +from __future__ import annotations + +import math +import os +import re +import subprocess +import time +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Optional, Tuple + +import torch + +#: ``DEF_CLMN_LMAX`` from ``phaser_src/defaults:19``. +PHASER_LMAX_CAP = 100 + +#: ``DEF_CLMN_SPHE`` from ``phaser_src/defaults:17``. Zero means "use +#: ``2 * mean_radius``" rather than an explicit sphere radius. +PHASER_SPHERE_DEFAULT = 0.0 + +#: Built by the recipe in the ``phaser-instrumented-build`` memo. Honours +#: ``$PHASER_FRF_DUMP`` and is otherwise stock. +PATCHED_PHASER = ( + Path(__file__).resolve().parents[2] / "phaser_src" / "build" / "phaser_patched" +) + + +def phaser_mean_radius(model) -> float: + """``xyz_weight::mean_radius()`` -- mean principal-axis semi-extent. + + Rotates the coordinates onto the principal axes of their covariance, takes + the bounding-box extent along each axis, halves it, and averages the three. + + This is emphatically *not* ``mean(|xyz - centroid|)``; on 1DAW the two give + 26.2 A and 19.5 A. Since ``LMAX`` and the sampling step are both derived + from it, using the wrong one detunes the whole rotation function. + + Parameters + ---------- + model : ModelFT + Search model. + + Returns + ------- + float + Mean radius in Angstrom. + """ + xyz = model.xyz().to(torch.float64) + centred = xyz - xyz.mean(dim=0) + cov = (centred.T @ centred) / centred.shape[0] + _, axes = torch.linalg.eigh(cov) + projected = centred @ axes + extent = projected.max(dim=0).values - projected.min(dim=0).values + return float((extent / 2.0).mean().item()) + + +@dataclass +class PhaserFRFParams: + """The derived FRF parameters for one case. + + Attributes + ---------- + mean_radius_A : float + ``mean_radius()``. + hires_A : float + ``mr.HiRes()`` -- the high-resolution limit of the selected data. + lmax : int + Maximum ``l`` of the expansion (even, capped at 100). + lmax_reso_A : float + Resolution the expansion actually runs at. + samp_reso_A : float + Resolution feeding the sampling formula. + sampling_deg : float + SO(3) grid step in degrees. + capped : bool + Whether ``lmax`` hit ``PHASER_LMAX_CAP`` -- when True the resolution and + sampling are both coarsened, when False both stay at ``hires_A``. + """ + + mean_radius_A: float + hires_A: float + lmax: int + lmax_reso_A: float + samp_reso_A: float + sampling_deg: float + capped: bool + + def as_row(self) -> dict: + """Flatten for a CSV result row, prefixed ``phaser_``.""" + return {f"phaser_{k}": v for k, v in asdict(self).items()} + + +def phaser_frf_params( + mean_radius_A: float, + hires_A: float, + *, + lmax_cap: int = PHASER_LMAX_CAP, + use_rotate_lmax_reso: bool = True, +) -> PhaserFRFParams: + """Run Phaser's bandwidth/resolution/sampling chain. + + Parameters + ---------- + mean_radius_A : float + From :func:`phaser_mean_radius`. + hires_A : float + High-resolution limit of the data being expanded. + lmax_cap : int, optional + ``DEF_CLMN_LMAX``. Default 100. Lower it to emulate our historical + ``lmax_cap`` settings *with Phaser's coupling intact* -- note that + coarsening then engages, exactly as it does in Phaser at 100. + use_rotate_lmax_reso : bool, optional + ``input.USE_ROTATE_LMAX_RESO``. Default True (Phaser's default). + + Returns + ------- + PhaserFRFParams + """ + sphere_outer = 2.0 * float(mean_radius_A) + lmax = int(math.ceil(2.0 * math.pi * sphere_outer / float(hires_A))) + if lmax % 2 != 0: + lmax += 1 + lmax = min(lmax, int(lmax_cap)) + + capped = lmax == int(lmax_cap) and use_rotate_lmax_reso + if capped: + lmax_reso = 2.0 * math.pi * sphere_outer / lmax + samp_reso = lmax_reso + else: + lmax_reso = float(hires_A) + samp_reso = float(hires_A) + + sampling = 2.0 * math.degrees(math.atan(samp_reso / (4.0 * float(mean_radius_A)))) + return PhaserFRFParams( + mean_radius_A=float(mean_radius_A), + hires_A=float(hires_A), + lmax=lmax, + lmax_reso_A=lmax_reso, + samp_reso_A=samp_reso, + sampling_deg=sampling, + capped=capped, + ) + + +# --------------------------------------------------------------------------- +# Running the patched binary +# --------------------------------------------------------------------------- + +#: ``RFACTOR USE OFF`` is compulsory for any diagnostic that uses a model +#: already close to its answer: Phaser otherwise computes the R-factor of the +#: ensemble at the origin, decides the structure is solved, and emits a single +#: identity peak with ``RF*0`` -- skipping the rotation search entirely while +#: still reporting ``EXIT STATUS: SUCCESS``. +_KEYWORD_TEMPLATE = """TITLE {title} +MODE MR_FRF +HKLIN {mtz} +LABIN F={f_label} SIGF={sigf_label} +ENSEMBLE search PDB {pdb} IDENT 1.0 +COMPOSITION BY AVERAGE +SEARCH ENSEMBLE search +RFACTOR USE OFF +PEAKS ROT SELECT NUMBER +PEAKS ROT CUTOFF {n_peaks} +PEAKS ROT CLUSTER OFF +OUTPUT LEVEL VERBOSE +ROOT {root} +""" + + +def write_keywords( + work: Path, + *, + mtz_path: Path, + model_pdb: Path, + n_peaks: int = 20, + d_min: Optional[float] = None, + d_max: Optional[float] = None, + f_label: str = "FP", + sigf_label: str = "SIGFP", + root: str = "phaser_frf", + title: str = "FRF map dump", +) -> Path: + """Write an MR_FRF keyword file for the patched binary. + + ``d_min``/``d_max`` are omitted by default so Phaser uses the full data + range and its own coupling picks the expansion resolution -- that is the + configuration our engine should be matched against. + + Note the peak keyword takes two cards: ``PEAKS ROT SELECT NUMBER`` sets the + *mode* and ``PEAKS ROT CUTOFF n`` the count. ``SELECT NUMBER n`` is a syntax + error. Keeping the count small matters: with ``SELECT ALL`` Phaser rescores + every sample point (150k+ for a mid-size case), which dwarfs the search. + """ + work.mkdir(parents=True, exist_ok=True) + text = _KEYWORD_TEMPLATE.format( + title=title, + mtz=Path(mtz_path).resolve(), + pdb=Path(model_pdb).resolve(), + f_label=f_label, + sigf_label=sigf_label, + n_peaks=int(n_peaks), + root=root, + ) + if d_min is not None and d_max is not None: + text = text.replace( + "RFACTOR USE OFF", f"RESOLUTION {d_min} {d_max}\nRFACTOR USE OFF", + ) + kw = work / f"{root}.kw" + kw.write_text(text) + return kw + + +def run_patched_phaser( + work: Path, + kw_path: Path, + *, + dump_path: Optional[Path] = None, + binary: Path = PATCHED_PHASER, + timeout_s: int = 5400, +) -> Tuple[int, float, Path]: + """Run the instrumented binary, dumping the FRF sample list. + + Returns + ------- + tuple + ``(returncode, seconds, log_path)``. **The return code is not a success + test** -- Phaser exits 0 on fatal keyword errors. Check that the log + contains a ``TORCHREF:`` line and that the dump parses. + """ + if not Path(binary).exists(): + raise FileNotFoundError( + f"patched phaser binary not found at {binary}; build it with the " + "recipe in the phaser-instrumented-build memo" + ) + work.mkdir(parents=True, exist_ok=True) + log_path = work / (kw_path.stem + ".log") + env = dict(os.environ) + if dump_path is not None: + env["PHASER_FRF_DUMP"] = str(Path(dump_path).resolve()) + + t0 = time.time() + try: + proc = subprocess.run( + [str(binary)], cwd=str(work), input=kw_path.read_text(), + capture_output=True, text=True, timeout=timeout_s, env=env, + ) + log_path.write_text((proc.stdout or "") + (proc.stderr or "")) + rc = proc.returncode + except subprocess.TimeoutExpired: + log_path.write_text("TIMEOUT\n") + rc = -1 + return rc, time.time() - t0, log_path + + +_LOG_PATTERNS = { + "lmax": re.compile(r"maximum l value\s+(\d+)"), + "sampling_deg": re.compile(r"Sampling:\s+([0-9.]+)\s+degrees"), + "mean_radius_A": re.compile(r"^\s*[0-9.]+\s+([0-9.]+)\s+\d+\s+-?[0-9.]+", re.M), + "lmax_reso_A": re.compile(r"Elmn with resolution\s+([0-9.]+)"), + "n_samples": re.compile(r"TORCHREF: wrote (\d+) FRF sample points"), + "selected_hi": re.compile( + r"Resolution of Selected Data \(Number\):\s+([0-9.]+)\s+([0-9.]+)\s+\((\d+)\)" + ), +} + + +def parse_phaser_log(log_path: Path) -> dict: + """Pull the parameters Phaser actually used out of a VERBOSE log. + + These are the ground truth for the derivation in :func:`phaser_frf_params`; + the comparison harness asserts the two agree rather than assuming they do. + + Returns + ------- + dict + Keys present only when found: ``lmax``, ``sampling_deg``, + ``mean_radius_A``, ``lmax_reso_A``, ``n_samples``, ``selected_d_min``, + ``selected_d_max``, ``selected_n_refl``, ``all_data_to_limit``, + ``rotation_search_skipped``. + """ + text = Path(log_path).read_text() + out: dict = {} + for key in ("lmax", "n_samples"): + m = _LOG_PATTERNS[key].search(text) + if m: + out[key] = int(m.group(1)) + for key in ("sampling_deg", "lmax_reso_A", "mean_radius_A"): + m = _LOG_PATTERNS[key].search(text) + if m: + out[key] = float(m.group(1)) + m = _LOG_PATTERNS["selected_hi"].search(text) + if m: + out["selected_d_min"] = float(m.group(1)) + out["selected_d_max"] = float(m.group(2)) + out["selected_n_refl"] = int(m.group(3)) + out["all_data_to_limit"] = "Elmn with all data to resolution limit" in text + # The R-factor short-circuit. Present => the search never ran. + out["rotation_search_skipped"] = "SOLU SET RF*0" in text + return out + + +def load_phaser_map( + dump_path: Path, *, dtype: torch.dtype = torch.float64, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Read a ``PHASER_FRF_DUMP`` CSV into our angle convention. + + Angles are returned **exactly as stored**, with no sign change. Phaser + writes ``euler = (-360*alpha_frac, beta_deg, -360*gamma_frac)`` + (``FastRot.cc:153``), and those stored values are already the Euler angles + of the grid rotation under ``R = Rz(alpha)Ry(beta)Rz(gamma)`` -- the same + Edmonds convention we use. Calibrated against a known grid/output pair on + 1DAW this reproduces Phaser's own reported peak to **0.074 deg**. + + Two traps here, both of which cost real time: + + * Negating alpha/gamma to "convert to our convention" is WRONG; they need + no conversion. + * The ``FastRot.cc`` comment at the end of ``get_FRF`` claims Phaser assumes + ``Rz(gamma)Ry(beta)Rz(alpha)``. That does not describe these stored + values; taking it literally puts the truth peak ~60-130 deg away. + + The stored angles are in the search model's **principal frame**. Use + ``ROT = axisrot @ R_grid @ PR`` (from the ``.frame`` sidecar) to reach the + PDB frame our engine works in. + + Returns + ------- + tuple + ``(angles_deg, values)`` -- ``(N, 3)`` of (alpha, beta, gamma) as stored, + and ``(N,)`` rotation-function values, in Phaser's sample order. + """ + import numpy as np + + raw = np.loadtxt(str(dump_path), delimiter=",", skiprows=1) + if raw.ndim == 1: + raw = raw[None, :] + angles = torch.from_numpy(raw[:, 1:4].copy()).to(dtype) + values = torch.from_numpy(raw[:, 4].copy()).to(dtype) + return angles, values + + +def load_phaser_frame(dump_path: Path, *, dtype: torch.dtype = torch.float64) -> dict: + """Read the ``.frame`` sidecar written next to a map dump. + + Returns ``PR``, ``axisrot`` (3x3 tensors), ``high_order_axis`` (int) and + ``stats`` (Phaser's own mean/sigma/max/min over the raw sample list, useful + for checking an externally computed sigma rather than trusting it). + """ + import numpy as np + + path = Path(str(dump_path) + ".frame") + if not path.exists(): + raise FileNotFoundError( + f"{path} missing; the binary must carry the runMR_FRF.cc patch that " + "dumps PR/axisrot, not only the FastRot.cc map patch" + ) + out: dict = {} + for line in path.read_text().strip().split("\n")[1:]: + parts = line.split(",") + vals = [float(x) for x in parts[1:10]] + if parts[0] in ("PR", "axisrot"): + out[parts[0]] = torch.tensor(np.array(vals).reshape(3, 3)).to(dtype) + elif parts[0] == "high_order_axis": + out["high_order_axis"] = int(vals[0]) + elif parts[0] == "stats_mean_sigma_max_min": + out["stats"] = dict(zip(("mean", "sigma", "max", "min"), vals[:4])) + return out + + +def phaser_sampling_from_dump(angles_deg: torch.Tensor) -> float: + """Recover the *exact* SO(3) sampling step from a dumped sample list. + + Phaser logs the sampling through ``dtos(SAMPLING,5,2)`` -- two decimals -- + which is far too coarse to rebuild the grid: on 1DAW the rounded 6.23 deg + gives 53270 sample points against the true 54430. The beta values in the + dump are full-precision and uniformly spaced, so their step is the exact + figure. + + Parameters + ---------- + angles_deg : torch.Tensor, shape (N, 3) + As returned by :func:`load_phaser_map`. + + Returns + ------- + float + Sampling step in degrees. + """ + betas = torch.unique(angles_deg[:, 1]) + if betas.numel() < 2: + raise ValueError("need at least two beta sections to infer the step") + steps = betas[1:] - betas[:-1] + return float(steps.median()) + + +def phaser_mean_radius_from_sampling(sampling_deg: float, samp_reso_A: float) -> float: + """Invert Phaser's sampling formula for ``mean_radius()``. + + ``sampling = 2*deg(atan(SAMP_RESO/(4*r)))`` inverts to + ``r = SAMP_RESO / (4*tan(sampling/2))``. Combined with + :func:`phaser_sampling_from_dump` this yields Phaser's own radius to full + precision, which is the reference our :func:`phaser_mean_radius` + reimplementation should be judged against (it currently runs ~4% high). + """ + half = math.radians(float(sampling_deg) / 2.0) + return float(samp_reso_A) / (4.0 * math.tan(half)) diff --git a/alignment_lab/lab/rescore.py b/alignment_lab/lab/rescore.py new file mode 100644 index 00000000..b5f3a473 --- /dev/null +++ b/alignment_lab/lab/rescore.py @@ -0,0 +1,195 @@ +"""Run the ML rescore over FRF peaks, and measure what it did to the ranking. + +The rescore's job is narrow: take the top ~20 FRF peaks -- which on this +benchmark reliably contain the true orientation -- and promote the true one to +the front. It is **not** a global search, so feeding it a peak list that does +not contain truth measures nothing. + +That makes the only honest metric a **paired** one: the rank of truth in the +list going in, versus its rank in the list coming out, on the same peaks. An +absolute post-rescore rank cannot distinguish "the rescore worked" from "the +FRF handed it an easy list". +""" + +from __future__ import annotations + +import time +from dataclasses import dataclass +from typing import Any, Dict, List, Optional, Sequence, Tuple + +import torch + +#: Rescore engines. ``none`` keeps the FRF's own ordering and is the control +#: arm -- without it, a rescore that merely preserves a good input ranking is +#: indistinguishable from one that improves it. +ENGINES = ("none", "m_letf1", "sim") + + +@dataclass +class RescoreResult: + """Outcome of one rescore. + + Attributes + ---------- + peaks : list + Re-ordered ``RotationPeak`` list. + engine : str + Engine used. + seconds : float + Wall time. + n_input : int + Peaks handed to the engine. + """ + + peaks: list + engine: str + seconds: float + n_input: int + + +def run_rescore( + peaks: Sequence, + data, + frf_inputs, + *, + engine: str = "m_letf1", + n_refine: int = 20, + n_shells: Optional[int] = None, + batch_size: int = 50, + subpeak_refine: bool = False, + verbose: int = 0, + **engine_kwargs: Any, +) -> RescoreResult: + """Rescore the top ``n_refine`` FRF peaks. + + Parameters + ---------- + peaks : sequence + FRF peaks, descending score. + data : ReflectionData + Dataset the peaks were scored against. + frf_inputs : FRFInputs + Prepared observations from the FRF run (``FRFResult.inputs``). + engine : {'none', 'm_letf1', 'sim'}, optional + ``'none'`` returns the input order unchanged -- the control arm. + n_refine : int, optional + How many leading peaks to rescore. Default 20, the intended use case. + n_shells : int, optional + Resolution shells for the rescore. Defaults to the pipeline's own rule, + ``max(n_shells // 2, 8)`` with ``n_shells = 20``. + batch_size : int, optional + Orientations evaluated per batch. Default 50. + subpeak_refine : bool, optional + Apply the quadratic tangent-space refinement after rescoring. + verbose : int, optional + Engine verbosity. + **engine_kwargs + Passed through to the engine (e.g. ``scat_mode``). + + Returns + ------- + RescoreResult + """ + if engine not in ENGINES: + raise ValueError(f"engine must be one of {ENGINES}, got {engine!r}") + + subset = list(peaks)[: max(int(n_refine), 0)] + if engine == "none" or not subset: + return RescoreResult(peaks=subset, engine=engine, seconds=0.0, + n_input=len(subset)) + + from torchref.experimental.alignment.ml_rotation import ( + m_letf1_rescore, sim_mlrf_rescore, + ) + + n_shells = n_shells if n_shells is not None else max(20 // 2, 8) + device = frf_inputs.F_obs.device + common = dict( + n_shells=n_shells, n_refine=len(subset), batch_size=batch_size, + verbose=verbose, + ) + + t0 = time.time() + if engine == "m_letf1": + out = m_letf1_rescore( + subset, frf_inputs.F_obs, frf_inputs.hkl, frf_inputs.s_mag, + frf_inputs.centric, frf_inputs.ll, data.cell, + data.spacegroup.matrices.to(torch.float64).to(device), + **common, **engine_kwargs, + ) + else: + out = sim_mlrf_rescore( + subset, frf_inputs.F_obs, frf_inputs.hkl, frf_inputs.s_mag, + frf_inputs.centric, frf_inputs.ll, data.cell, + **common, **engine_kwargs, + ) + if subpeak_refine: + from torchref.experimental.alignment.ml_rotation import ( + _build_llg_context, quadratic_llg_refine, + ) + + ctx = _build_llg_context( + frf_inputs.F_obs, frf_inputs.hkl, frf_inputs.s_mag, + frf_inputs.centric, frf_inputs.ll, data.cell, n_shells=n_shells, + ) + out = quadratic_llg_refine(out, ctx) + seconds = time.time() - t0 + return RescoreResult(peaks=list(out), engine=engine, seconds=seconds, + n_input=len(subset)) + + +def paired_ranks( + frf_peaks: Sequence, + rescored: Sequence, + R_true: torch.Tensor, + symops: torch.Tensor, + *, + n_refine: int, + thr_deg: float = 5.0, + **orbit_kw: Any, +) -> Dict[str, Any]: + """Rank of truth before and after rescoring, on the same peak subset. + + Parameters + ---------- + frf_peaks : sequence + Full FRF peak list. + rescored : sequence + Output of :func:`run_rescore`. + R_true : torch.Tensor + True rotation. + symops : torch.Tensor + Symmetry rotation parts. + n_refine : int + Size of the subset handed to the rescore -- the comparison window. + thr_deg : float, optional + Orbit match threshold. + **orbit_kw + Orbit convention (``side``/``frame``/``reciprocal_basis``). + + Returns + ------- + dict + ``rank_frf`` (within the subset), ``rank_rescored``, ``delta`` + (positive = the rescore made it worse), ``truth_in_window`` and + ``rank_frf_full``. ``delta`` is ``None`` when truth was never in the + window, because then the rescore was never given the chance. + """ + from .truth import orbit_rank + + subset = list(frf_peaks)[:n_refine] + rank_full, _ = orbit_rank(frf_peaks, R_true, symops, thr_deg=thr_deg, **orbit_kw) + rank_in, ang_in = orbit_rank(subset, R_true, symops, thr_deg=thr_deg, **orbit_kw) + rank_out, ang_out = orbit_rank(rescored, R_true, symops, thr_deg=thr_deg, **orbit_kw) + + in_window = rank_in >= 0 + delta = (rank_out - rank_in) if (in_window and rank_out >= 0) else None + return { + "rank_frf_full": rank_full, + "rank_frf": rank_in, + "rank_rescored": rank_out, + "delta": delta, + "truth_in_window": in_window, + "angle_frf": ang_in, + "angle_rescored": ang_out, + } diff --git a/alignment_lab/lab/results.py b/alignment_lab/lab/results.py new file mode 100644 index 00000000..7fed6c04 --- /dev/null +++ b/alignment_lab/lab/results.py @@ -0,0 +1,139 @@ +"""One result schema, one writer, one aggregator input format. + +Each old experiment invented its own CSV columns, so each needed its own +bespoke aggregator. Rows written through :class:`ResultWriter` all carry the +same core fields plus experiment-specific extras, so a single aggregator works +across experiments and a stale result is identifiable from the row itself. +""" + +from __future__ import annotations + +import csv +import os +import subprocess +from pathlib import Path +from typing import Any, Dict, Iterable, Mapping, Optional + +#: Fields every row carries. `orbit_side` / `orbit_frame` are here because a +#: truth rank cannot be interpreted without knowing which convention produced +#: it, and `torchref_version` / `git_sha` because results outlive the checkout. +CORE_FIELDS = ( + "pdb", + "seed", + "trial", + "spacegroup", + "n_ops", + "truth_rank", + "truth_angle_deg", + "orbit_side", + "orbit_frame", + "lmax_cap", + "d_min", + "d_max", + "device", + "torchref_version", + "git_sha", +) + + +def _git_sha(default: str = "unknown") -> str: + """Short SHA of the checkout this is running from, or ``default``.""" + try: + out = subprocess.run( + ["git", "rev-parse", "--short", "HEAD"], + cwd=str(Path(__file__).resolve().parents[2]), + capture_output=True, text=True, timeout=10, + ) + return out.stdout.strip() or default + except Exception: + return default + + +def provenance() -> Dict[str, str]: + """Version and checkout identity for a result row. + + Returns + ------- + dict + ``{'torchref_version': ..., 'git_sha': ...}``. + """ + try: + import torchref + + version = getattr(torchref, "__version__", "unknown") + except Exception: + version = "unknown" + return {"torchref_version": version, "git_sha": _git_sha()} + + +def append_row(csv_path: str | os.PathLike, row: Mapping[str, Any]) -> None: + """Append one row, writing the header only when the file is new. + + Uses ``fh.tell() == 0`` rather than an existence check so a file created + but not yet written still gets its header. + + Parameters + ---------- + csv_path : path-like + Destination CSV. Parent directories are created. + row : mapping + Column name -> value. + """ + path = Path(csv_path) + path.parent.mkdir(parents=True, exist_ok=True) + with open(path, "a", newline="") as fh: + writer = csv.DictWriter(fh, fieldnames=list(row.keys())) + if fh.tell() == 0: + writer.writeheader() + writer.writerow(dict(row)) + + +class ResultWriter: + """Writes rows sharing a fixed core schema plus per-experiment extras. + + Parameters + ---------- + csv_path : path-like + Destination CSV. + experiment : str + Experiment tag, recorded in every row. + extra_fields : iterable of str, optional + Experiment-specific column names, appended after the core fields. + + Notes + ----- + Column order is fixed at construction, so every row in a file has the same + header even if a caller omits a value (missing entries are written empty). + """ + + def __init__( + self, + csv_path: str | os.PathLike, + experiment: str, + extra_fields: Optional[Iterable[str]] = None, + ): + self.path = Path(csv_path) + self.experiment = experiment + self.extra_fields = tuple(extra_fields or ()) + self.fieldnames = ("experiment",) + CORE_FIELDS + self.extra_fields + self._provenance = provenance() + + def write(self, **values: Any) -> None: + """Write one row; unknown keys raise rather than being dropped silently. + + Raises + ------ + KeyError + If a value is passed whose column was not declared. + """ + unknown = set(values) - set(self.fieldnames) + if unknown: + raise KeyError( + f"{self.experiment}: undeclared column(s) {sorted(unknown)}; " + f"add them to extra_fields so every row keeps the same header" + ) + row = {k: "" for k in self.fieldnames} + row["experiment"] = self.experiment + row.update(self._provenance) + row.update(values) + append_row(self.path, row) diff --git a/alignment_lab/lab/truth.py b/alignment_lab/lab/truth.py new file mode 100644 index 00000000..c5ffbbc0 --- /dev/null +++ b/alignment_lab/lab/truth.py @@ -0,0 +1,198 @@ +"""Seeded rotations and rank-of-truth against a symmetry orbit. + +Two things here previously existed in several disagreeing copies, and both +silently changed results rather than raising: + +* **The rotation generator.** Two QR-based variants were in circulation; the + one omitting the ``sign(diag(R))`` correction is not Haar-uniform and returns + a *different* rotation for the same seed. Runs from the two families are not + comparable. :func:`random_rotation` is the corrected form. +* **The orbit convention.** Rank-of-truth was computed with the symmetry + operators applied on either side, and in either the fractional or the + Cartesian frame. The choice changes the answer, so it is an explicit argument + here and is meant to be recorded in every result row. +""" + +from __future__ import annotations + +import math +from typing import Iterable, Optional, Sequence, Tuple + +import torch + + +def random_rotation(seed: int, dtype: torch.dtype = torch.float64) -> torch.Tensor: + """Haar-uniform random rotation matrix from a seed. + + QR of a Gaussian matrix, with the ``sign(diag(R))`` correction that makes + the decomposition unique -- without it the distribution is not Haar-uniform + and the seed maps to a different rotation. + + Parameters + ---------- + seed : int + Generator seed. The mapping seed -> rotation is the reproducibility + contract for the whole lab; changing this function invalidates every + archived result. + dtype : torch.dtype, optional + Output dtype. Default ``torch.float64``. + + Returns + ------- + torch.Tensor + ``(3, 3)`` rotation with ``det = +1``. + """ + g = torch.Generator().manual_seed(int(seed)) + A = torch.randn(3, 3, generator=g, dtype=torch.float64) + Q, R = torch.linalg.qr(A) + Q = Q @ torch.diag(torch.sign(torch.diag(R))) + if torch.det(Q) < 0: + Q[:, 0] = -Q[:, 0] + return Q.to(dtype) + + +def seed_for(pdb: str, trial: int, base: int = 42) -> int: + """Seed for a ``(structure, trial)`` cell of the benchmark. + + ``base + 1000 * trial + index(pdb) * 7`` -- the convention the archived + results were produced under. The index term is why + :data:`~alignment_lab.lab.benchmark.BENCH_PDBS` is append-only. + + Parameters + ---------- + pdb : str + Benchmark structure code. + trial : int + Trial number. + base : int, optional + Seed base. Default 42. + + Returns + ------- + int + The seed. + """ + from .benchmark import BENCH_PDBS + + return int(base) + 1000 * int(trial) + BENCH_PDBS.index(pdb) * 7 + + +def symmetry_orbit( + R_true: torch.Tensor, + symops: torch.Tensor, + *, + side: str = "left", + frame: str = "cart", + reciprocal_basis: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Build the set of rotations equivalent to ``R_true`` under the point group. + + Parameters + ---------- + R_true : torch.Tensor + ``(3, 3)`` true rotation. + symops : torch.Tensor + ``(n_ops, 3, 3)`` symmetry rotation parts, as stored on the space group + (fractional). + side : {'left', 'right'}, optional + ``'left'`` builds ``S_k @ R_true``; ``'right'`` builds ``R_true @ S_k``. + These are different sets for non-commuting operators. + frame : {'cart', 'frac'}, optional + ``'cart'`` converts the operators to the Cartesian frame first, which is + the frame the rotation function works in. ``'frac'`` uses them as + stored. Mixing a Cartesian rotation with fractional operators is a + metric error that inflates apparent ghost counts. + reciprocal_basis : torch.Tensor, optional + ``(3, 3)`` reciprocal basis, required when ``frame='cart'``. + + Returns + ------- + torch.Tensor + ``(n_ops, 3, 3)`` orbit members, float64. + """ + if side not in ("left", "right"): + raise ValueError(f"side must be 'left' or 'right', got {side!r}") + if frame not in ("cart", "frac"): + raise ValueError(f"frame must be 'cart' or 'frac', got {frame!r}") + + R = R_true.to(torch.float64) + S = symops.to(torch.float64) + if frame == "cart": + if reciprocal_basis is None: + raise ValueError("frame='cart' requires reciprocal_basis") + from torchref.experimental.alignment.sh import hkl_symops_to_cartesian + + S = hkl_symops_to_cartesian(S, reciprocal_basis.to(torch.float64)) + return S @ R.unsqueeze(0) if side == "left" else R.unsqueeze(0) @ S + + +def angle_to_orbit(R: torch.Tensor, orbit: torch.Tensor) -> float: + """Smallest rotation angle between ``R`` and any orbit member, in degrees. + + Parameters + ---------- + R : torch.Tensor + ``(3, 3)`` rotation. + orbit : torch.Tensor + ``(n, 3, 3)`` orbit members. + + Returns + ------- + float + Angle in degrees. + """ + tr = torch.einsum("kij,ij->k", orbit.to(torch.float64), R.to(torch.float64)) + cos = ((tr - 1.0) * 0.5).clamp(-1.0, 1.0) + return float(cos.arccos().min() * (180.0 / math.pi)) + + +def orbit_rank( + peaks: Sequence, + R_true: torch.Tensor, + symops: torch.Tensor, + *, + side: str = "left", + frame: str = "cart", + reciprocal_basis: Optional[torch.Tensor] = None, + thr_deg: float = 5.0, +) -> Tuple[int, float]: + """Rank of the first peak matching the true orientation. + + Parameters + ---------- + peaks : sequence + Peaks carrying Edmonds ZYZ ``alpha``/``beta``/``gamma`` in radians, in + descending score order (the FRF's ``RotationPeak``). + R_true : torch.Tensor + ``(3, 3)`` true rotation. + symops : torch.Tensor + ``(n_ops, 3, 3)`` symmetry rotation parts. + side, frame, reciprocal_basis + Orbit convention -- see :func:`symmetry_orbit`. Record these alongside + any rank you report; the rank is meaningless without them. + thr_deg : float, optional + Match threshold in degrees. Default 5.0. + + Returns + ------- + tuple + ``(rank, best_angle_deg)``. ``rank`` is ``-1`` when no peak matches; + ``best_angle_deg`` is the closest approach over all peaks either way, + which distinguishes "just outside the threshold" from "absent". + """ + from torchref.experimental.alignment.frf.rotation_utils import ( + rotation_matrix_from_edmonds_euler, + ) + + orbit = symmetry_orbit( + R_true, symops, side=side, frame=frame, reciprocal_basis=reciprocal_basis, + ) + rank, best = -1, float("inf") + for i, p in enumerate(peaks): + R_p = rotation_matrix_from_edmonds_euler(p.alpha, p.beta, p.gamma) + ang = angle_to_orbit(R_p, orbit) + if ang < best: + best = ang + if ang <= thr_deg and rank < 0: + rank = i + return rank, best diff --git a/alignment_lab/tests/test_lab.py b/alignment_lab/tests/test_lab.py new file mode 100644 index 00000000..f5a66faf --- /dev/null +++ b/alignment_lab/tests/test_lab.py @@ -0,0 +1,163 @@ +"""Self-tests for the alignment lab primitives. + +These pin the contracts whose violation silently changed results in the past: +the seed -> rotation mapping, the append-only benchmark order the seed formula +depends on, and the orbit conventions. +""" + +from __future__ import annotations + +import math +import sys +from pathlib import Path + +import pytest +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from lab import BENCH_PDBS, case_paths, orbit_rank, random_rotation, seed_for # noqa: E402 +from lab.truth import angle_to_orbit, symmetry_orbit # noqa: E402 + + +def test_random_rotation_is_a_rotation(): + """Output is orthogonal with det +1 for a spread of seeds.""" + for seed in (0, 1, 42, 2077, 999983): + R = random_rotation(seed) + assert torch.allclose(R @ R.T, torch.eye(3, dtype=R.dtype), atol=1e-12) + assert abs(float(torch.det(R)) - 1.0) < 1e-12 + + +def test_random_rotation_is_deterministic(): + """The seed -> rotation map is the lab's reproducibility contract.""" + assert torch.equal(random_rotation(42), random_rotation(42)) + assert not torch.equal(random_rotation(42), random_rotation(43)) + + +def test_random_rotation_uses_the_sign_corrected_qr(): + """Guard the exact variant: the uncorrected QR gives a different rotation. + + Both forms return a valid rotation, so only a direct comparison catches a + swap -- and a swap silently makes new results incomparable with archived + ones for the same seed. + """ + seed = 42 + g = torch.Generator().manual_seed(seed) + A = torch.randn(3, 3, generator=g, dtype=torch.float64) + Q_uncorrected, _ = torch.linalg.qr(A) + if torch.det(Q_uncorrected) < 0: + Q_uncorrected[:, 0] = -Q_uncorrected[:, 0] + assert not torch.allclose(random_rotation(seed), Q_uncorrected, atol=1e-9) + + +def test_seed_formula(): + """base + 1000*trial + index(pdb)*7.""" + assert seed_for("1DAW", 0) == 42 + assert seed_for("1DAW", 1) == 1042 + assert seed_for("3K7M", 2) == 42 + 2000 + BENCH_PDBS.index("3K7M") * 7 + + +def test_benchmark_order_is_pinned(): + """The seed formula indexes into this tuple, so its order is a contract.""" + assert BENCH_PDBS[0] == "1DAW" + assert BENCH_PDBS.index("1AK5") == 4 + assert BENCH_PDBS.index("3K7M") == 5 + assert len(BENCH_PDBS) == len(set(BENCH_PDBS)) == 10 + + +@pytest.mark.parametrize("pdb", BENCH_PDBS) +def test_every_benchmark_case_resolves(pdb): + """Paths are repo-relative and present -- not absolute into another tree.""" + pdb_path, mtz_path = case_paths(pdb) + assert pdb_path.is_file() and mtz_path.is_file() + + +def test_orbit_side_and_frame_are_distinct_conventions(): + """left/right and frac/cart really do differ, so recording them matters.""" + from torchref.symmetry import SpaceGroup + + sg = SpaceGroup("P 4 3 2") + symops = sg.matrices.to(torch.float64).cpu() + R = random_rotation(7) + left = symmetry_orbit(R, symops, side="left", frame="frac") + right = symmetry_orbit(R, symops, side="right", frame="frac") + assert not torch.allclose(left, right, atol=1e-9) + + +def test_orbit_contains_truth_at_zero_angle(): + """Every orbit member is 0 degrees from the orbit, by construction.""" + from torchref.symmetry import SpaceGroup + + symops = SpaceGroup("P 4 3 2").matrices.to(torch.float64).cpu() + R = random_rotation(11) + orbit = symmetry_orbit(R, symops, side="left", frame="frac") + for k in range(orbit.shape[0]): + assert angle_to_orbit(orbit[k], orbit) < 1e-9 + + +def test_orbit_rank_reports_miss_as_minus_one(): + """A peak list with no match ranks -1 but still reports the closest angle.""" + from types import SimpleNamespace + + from torchref.symmetry import SpaceGroup + + symops = SpaceGroup("P 1").matrices.to(torch.float64).cpu() + peaks = [SimpleNamespace(alpha=0.0, beta=0.0, gamma=0.0)] + R_true = random_rotation(3) + rank, ang = orbit_rank(peaks, R_true, symops, frame="frac", thr_deg=1e-6) + assert rank == -1 + assert math.isfinite(ang) and ang > 0 + + +def test_result_writer_rejects_undeclared_columns(tmp_path): + """Silent column drift is what made every old CSV need its own aggregator.""" + from lab import ResultWriter + + w = ResultWriter(tmp_path / "r.csv", "demo", extra_fields=("ghosts",)) + w.write(pdb="1DAW", truth_rank=0, ghosts=3) + with pytest.raises(KeyError): + w.write(pdb="1DAW", not_declared=1) + assert (tmp_path / "r.csv").read_text().count("\n") == 2 + + +def test_paired_ranks_reports_no_delta_when_truth_is_outside_the_window(): + """A rescore cannot be blamed for a peak it was never shown. + + When truth is absent from the top-N handed to the engine, ``delta`` is None + rather than a number -- otherwise the metric silently reports the FRF's + failure as a rescore regression. + """ + from types import SimpleNamespace + + from lab import paired_ranks, random_rotation + from torchref.symmetry import SpaceGroup + from torchref.experimental.alignment.frf.rotation_utils import ( + rotation_matrix_from_edmonds_euler, + ) + + symops = SpaceGroup("P 1").matrices.to(torch.float64).cpu() + R_true = random_rotation(3) + + # A peak list whose only truth-matching entry sits beyond the window. + def peak_at(R): + # recover ZYZ angles numerically is unnecessary: use a far-off peak for + # the decoys and the true rotation only at the tail. + return SimpleNamespace(alpha=0.0, beta=0.0, gamma=0.0) + + decoys = [peak_at(None) for _ in range(5)] + out = paired_ranks(decoys, decoys, R_true, symops, + n_refine=2, frame="frac", thr_deg=1e-6) + assert out["truth_in_window"] is False + assert out["delta"] is None + + +def test_run_rescore_none_is_an_identity_control(): + """The 'none' arm must return the input order untouched.""" + from types import SimpleNamespace + + from lab import run_rescore + + peaks = [SimpleNamespace(alpha=float(i), beta=0.0, gamma=0.0) for i in range(5)] + res = run_rescore(peaks, data=None, frf_inputs=None, engine="none", n_refine=3) + assert res.engine == "none" + assert [p.alpha for p in res.peaks] == [0.0, 1.0, 2.0] From 133bd5652dc677653a1ed1533cf1980591715089 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 11:48:09 +0200 Subject: [PATCH 015/250] Add the FRF configuration sweep to the lab One diagnostic covering the four engine settings that are still switches because no value was ever chosen: `lmax_cap` (48 vs 64 vs Phaser's DEF_CLMN_LMAX of 100), the anisotropy estimator, `_orbit_unroll`, and the Patterson-radius union. Stage 1 is the lmax x anisotropy factorial, stage 2 takes the follow-ups one at a time from the winning cell. Every arm runs in one process per (structure, trial) cell, so the paired comparison against the shipped configuration is exact, and `production_dup` repeats the baseline verbatim to measure the engine's own run-to-run spread -- which bounds how small an effect the sweep can resolve at all. `merge_peak_lists` pools peak lists by z-score with SO(3) suppression. Absolute rotation-function values from two Patterson radii are not comparable; the per-run standardised heights are. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/config_sweep_array.sh | 53 ++++ alignment_lab/diagnostics/frf_config_sweep.py | 247 ++++++++++++++++++ alignment_lab/lab/__init__.py | 4 +- alignment_lab/lab/frf.py | 42 +++ 4 files changed, 345 insertions(+), 1 deletion(-) create mode 100644 alignment_lab/analysis/config_sweep_array.sh create mode 100644 alignment_lab/diagnostics/frf_config_sweep.py diff --git a/alignment_lab/analysis/config_sweep_array.sh b/alignment_lab/analysis/config_sweep_array.sh new file mode 100644 index 00000000..8803b6cc --- /dev/null +++ b/alignment_lab/analysis/config_sweep_array.sh @@ -0,0 +1,53 @@ +#!/bin/bash +# Part 1 of the FRF cleanup: settle lmax_cap / anisotropy / orbit-unroll / +# Patterson-radius by measurement before the switches are deleted. +# +# One array task = one (structure, trial) cell, running every arm in the same +# process so the paired comparison against `production` is exact. +# +# sbatch --array=0-99 --partition=hour --time=00:55:00 --cpus-per-task=4 \ +# --mem=32G alignment_lab/analysis/config_sweep_array.sh 1 +# +# 10 structures x 10 trials = 100 tasks. Pass the stage (1 or 2) as $1; any +# further arguments go through to the diagnostic. +#SBATCH --job-name=frf_cfg_sweep +#SBATCH --output=alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=alignment_lab/slurm/%x_%A_%a.err +set -uo pipefail + +STAGE="${1:?usage: config_sweep_array.sh [extra args...]}" +shift || true + +REPO="${FRF_SWEEP_REPO:-/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement}" +PY="$REPO/.dev/bin/python" +[ -x "$PY" ] || PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python + +cd "$REPO" +export PYTHONPATH="$REPO" +export TORCHREF_NUM_THREADS="${SLURM_CPUS_PER_TASK:-4}" +export OMP_NUM_THREADS="$TORCHREF_NUM_THREADS" +export MKL_NUM_THREADS="$TORCHREF_NUM_THREADS" +export PYTHONUNBUFFERED=1 +export CUDA_VISIBLE_DEVICES="" + +# Worklist: keep in step with lab.benchmark.BENCH_PDBS (order is a seed contract). +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +TRIALS=10 +IDX="${SLURM_ARRAY_TASK_ID:-0}" +PDB="${PDBS[$((IDX / TRIALS))]}" +TRIAL=$((IDX % TRIALS)) + +OUTDIR="alignment_lab/runs/config_sweep_s${STAGE}_${SLURM_ARRAY_JOB_ID:-local}" +mkdir -p "$OUTDIR" alignment_lab/slurm + +echo "task $IDX -> $PDB trial $TRIAL stage $STAGE -> $OUTDIR" +echo "repo=$REPO sha=$(git -C "$REPO" rev-parse --short HEAD 2>/dev/null || echo unknown)" +rc=0 +"$PY" -u alignment_lab/diagnostics/frf_config_sweep.py \ + --pdb "$PDB" --trial "$TRIAL" --stage "$STAGE" \ + --out-csv "$OUTDIR/${PDB}_t${TRIAL}.csv" "$@" || rc=$? + +# Report the real exit status: `rc=$?` must follow the command directly, or a +# task that dies gets logged COMPLETED. +echo "exit_code=$rc" +exit "$rc" diff --git a/alignment_lab/diagnostics/frf_config_sweep.py b/alignment_lab/diagnostics/frf_config_sweep.py new file mode 100644 index 00000000..05a7c490 --- /dev/null +++ b/alignment_lab/diagnostics/frf_config_sweep.py @@ -0,0 +1,247 @@ +"""Settle the FRF's remaining free constants by measurement, one arm each. + +Four engine settings are still switches because nobody chose a value. Making +the rotation search a three-input call means choosing them, and each choice +gets a number first: + +``lmax_cap`` + The signature default is 48, its own docstring claims 100, and the + benchmarks run 64. Phaser's ``DEF_CLMN_LMAX`` is 100. The "high l + under-determines the SH modes" argument for 48 predates the dense P1-box + calc, so it is not evidence about the current engine. +anisotropy + ``production`` is the log-space fit with no intercept; ``fixed_fit`` is the + intensity-space replacement; ``iso_only`` keeps its radial part; ``no_aniso`` + drops the correction. Measured before at seven structures, where + ``fixed_fit`` was indistinguishable from ``no_aniso`` in aggregate. +``_orbit_unroll`` + Off, on the strength of a run that predates the reciprocal-space + convention fix, so its evidence is void. +Patterson radius + Never exercised in production. Two structures want radii a factor 2.4 + apart with no rule to pick between them, so the candidate is the *union*: + two runs merged by z-score. It doubles the cost, so it has to earn it. + +Every arm runs in one process per (structure, trial) cell, so the paired +comparison against ``production`` is exact. ``production_dup`` repeats the +baseline arm verbatim: it measures the engine's own run-to-run spread, which +bounds how small a real effect this sweep can resolve. + +Usage +----- + python -m diagnostics.frf_config_sweep --pdb 3GR5 --trial 0 + python -m diagnostics.frf_config_sweep --pdb 3GR5 --trials 10 --stage 2 +""" + +from __future__ import annotations + +import argparse +import sys +import time +from dataclasses import dataclass, field +from pathlib import Path +from typing import Dict, Optional, Tuple + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, FRFConfig, aniso_arm, merge_peak_lists, # noqa: E402 + orbit_rank, rotated_case, run_frf, seed_for, tensor_report) +from lab.results import append_row, provenance # noqa: E402 + +EXPERIMENT = "frf_config_sweep" + +#: Suppression radius for the union merge. The engine uses +#: ``max(2 * grid_sampling_deg, 6)`` internally, and the production sampling is +#: 3 degrees, so 6 degrees keeps the merged list on the same footing. +UNION_NMS_DEG = 6.0 + + +@dataclass(frozen=True) +class Arm: + """One engine configuration to measure. + + ``radius_scales`` with more than one entry means the union arm: one FRF + evaluation per scale, merged by z-score. + """ + + name: str + lmax_cap: int = 64 + aniso: str = "production" + orbit_unroll: bool = False + radius_scales: Tuple[float, ...] = (1.0,) + + def config(self, base: FRFConfig) -> Tuple[FRFConfig, ...]: + out = [] + for scale in self.radius_scales: + extra: Dict[str, object] = {"_orbit_unroll": self.orbit_unroll} + if scale != 1.0: + extra["frf_patterson_radius_scale"] = scale + out.append(FRFConfig( + d_min=base.d_min, d_max=base.d_max, n_shells=base.n_shells, + n_peaks=base.n_peaks, lmax_cap=self.lmax_cap, + dense_pad=base.dense_pad, extra=extra, + )) + return tuple(out) + + +def _factorial_arms() -> Tuple[Arm, ...]: + """lmax_cap x anisotropy, plus the repeat-baseline control.""" + arms = [Arm("production_dup")] + for cap in (48, 64, 100): + for aniso in ("production", "fixed_fit", "iso_only", "no_aniso"): + arms.append(Arm(f"cap{cap}_{aniso}", lmax_cap=cap, aniso=aniso)) + return tuple(arms) + + +def _followup_arms(cap: int, aniso: str) -> Tuple[Arm, ...]: + """One-at-a-time from the winning cell of stage 1.""" + base = Arm(f"cap{cap}_{aniso}", lmax_cap=cap, aniso=aniso) + return ( + base, + Arm(f"{base.name}_unroll", lmax_cap=cap, aniso=aniso, orbit_unroll=True), + Arm(f"{base.name}_union", lmax_cap=cap, aniso=aniso, + radius_scales=(1.0, 0.5)), + ) + + +#: The baseline every paired difference is taken against: today's shipped +#: configuration (broken anisotropy fit, cap 64, no unroll, single radius). +BASELINE = Arm("production", lmax_cap=64, aniso="production") + + +def run_one(pdb: str, trial: int, arm: Arm, base: FRFConfig, + *, thr_deg: float) -> dict: + seed = seed_for(pdb, trial) + model, data, R_true = rotated_case(pdb, seed) + configs = arm.config(base) + + captured: dict = {} + peak_lists = [] + t0 = time.time() + for cfg in configs: + with aniso_arm(arm.aniso if arm.aniso != "production" else "production", + data, d_min=cfg.d_min, d_max=cfg.d_max, + captured=captured): + res = run_frf(model, data, cfg, capture_arf=False, verbose=0) + peak_lists.append(res.peaks) + seconds = time.time() - t0 + + peaks = (peak_lists[0] if len(peak_lists) == 1 else + merge_peak_lists(peak_lists, n_peaks=base.n_peaks, + nms_radius_deg=UNION_NMS_DEG)) + + rank, ang = orbit_rank( + peaks, R_true, data.spacegroup.matrices.to(torch.float64).cpu(), + reciprocal_basis=data.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), + side="left", frame="cart", thr_deg=thr_deg, + ) + row = {"experiment": EXPERIMENT, "pdb": pdb, "trial": trial, + "arm": arm.name, "seed": seed} + row.update(provenance()) + row.update(configs[0].as_row()) + row.update({ + "arm_lmax_cap": arm.lmax_cap, + "arm_aniso": arm.aniso, + "arm_orbit_unroll": int(arm.orbit_unroll), + "arm_radius_scales": "|".join(f"{s:g}" for s in arm.radius_scales), + "n_frf_calls": len(configs), + "spacegroup": str(data.spacegroup.hm), + "truth_rank": rank, + # orbit_rank returns -1 for "no peak within thr_deg". A miss must not + # sort as a good rank, so for pairing it counts as worse than the worst + # hit, i.e. the length of the peak list. + "rank_for_compare": rank if rank >= 0 else base.n_peaks, + "found": int(rank >= 0), + "in_top20": int(0 <= rank < 20), + "truth_angle_deg": None if ang is None else round(float(ang), 3), + "n_peaks_found": len(peaks), + "orbit_side": "left", "orbit_frame": "cart", "thr_deg": thr_deg, + "seconds": round(seconds, 1), + }) + for tag in ("raw", "fixed"): + if tag in captured: + row.update(tensor_report(captured[tag], tag)) + return row + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) + ap.add_argument("--trial", type=int, default=None, + help="single trial index; omit to run --trials of them") + ap.add_argument("--trials", type=int, default=10) + ap.add_argument("--stage", type=int, default=1, choices=(1, 2), + help="1 = lmax x aniso factorial; 2 = follow-ups") + ap.add_argument("--stage2-cap", type=int, default=64) + ap.add_argument("--stage2-aniso", default="fixed_fit") + ap.add_argument("--d-min", type=float, default=4.0) + ap.add_argument("--d-max", type=float, default=15.0) + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--thr-deg", type=float, default=5.0) + ap.add_argument("--out-csv", default=None) + ap.add_argument("--outdir", default=None) + args = ap.parse_args() + + arms = (BASELINE,) + ( + _factorial_arms() if args.stage == 1 + else _followup_arms(args.stage2_cap, args.stage2_aniso) + ) + base = FRFConfig(d_min=args.d_min, d_max=args.d_max, n_peaks=args.n_peaks) + + if args.out_csv: + csv_path = Path(args.out_csv) + csv_path.parent.mkdir(parents=True, exist_ok=True) + else: + outdir = Path(args.outdir) if args.outdir else ( + Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT) + outdir.mkdir(parents=True, exist_ok=True) + csv_path = outdir / f"{EXPERIMENT}_{args.pdb}.csv" + + trials = [args.trial] if args.trial is not None else list(range(args.trials)) + print(f"{args.pdb}: stage {args.stage}, {len(arms)} arms x {len(trials)} " + f"trial(s)", flush=True) + + ranks: Dict[str, list] = {a.name: [] for a in arms} + n_fail = 0 + for trial in trials: + for arm in arms: + try: + row = run_one(args.pdb, trial, arm, base, thr_deg=args.thr_deg) + except Exception as exc: + n_fail += 1 + ranks[arm.name].append(None) + print(f" trial {trial} {arm.name}: FAILED {type(exc).__name__}: " + f"{exc}", flush=True) + continue + append_row(csv_path, row) + ranks[arm.name].append(row["rank_for_compare"]) + shown = (str(row["truth_rank"]) if row["found"] + else f"miss@{row['truth_angle_deg']:.0f}deg") + print(f" trial {trial} {arm.name:<26} rank={shown:<12} " + f"top20={row['in_top20']} {row['seconds']:>6.1f}s", flush=True) + + base_ranks = ranks[BASELINE.name] + print("\npaired vs production (negative = better rank):", flush=True) + for arm in arms: + if arm.name == BASELINE.name: + continue + d = [(a - b) for a, b in zip(ranks[arm.name], base_ranks) + if a is not None and b is not None] + if not d: + print(f" {arm.name:<26} no paired trials", flush=True) + continue + sd = sorted(d) + med = (sd[len(sd) // 2] if len(sd) % 2 + else 0.5 * (sd[len(sd) // 2 - 1] + sd[len(sd) // 2])) + print(f" {arm.name:<26} n={len(d):<3} better={sum(x < 0 for x in d)} " + f"same={sum(x == 0 for x in d)} worse={sum(x > 0 for x in d)} " + f"median={med:+.1f} per-trial={d}", flush=True) + print(f"\nwrote {csv_path} ({n_fail} failures)", flush=True) + return 1 if n_fail == len(arms) * len(trials) else 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/lab/__init__.py b/alignment_lab/lab/__init__.py index 6870487a..bc6dad2e 100644 --- a/alignment_lab/lab/__init__.py +++ b/alignment_lab/lab/__init__.py @@ -27,7 +27,8 @@ fit_aniso_intensity_space, tensor_report, ) -from .frf import FRFConfig, FRFResult, patched, run_frf +from .frf import (FRFConfig, FRFResult, merge_peak_lists, patched, + run_frf) from .rescore import ENGINES, RescoreResult, paired_ranks, run_rescore from .results import ResultWriter, append_row, provenance @@ -48,6 +49,7 @@ "tensor_report", "FRFConfig", "FRFResult", + "merge_peak_lists", "patched", "run_frf", "ENGINES", diff --git a/alignment_lab/lab/frf.py b/alignment_lab/lab/frf.py index e42ca428..35933869 100644 --- a/alignment_lab/lab/frf.py +++ b/alignment_lab/lab/frf.py @@ -97,6 +97,48 @@ def map_max_sigma(self) -> float: return float(self.sigma.max()) if self.sigma is not None else float("nan") +def merge_peak_lists(peak_lists, *, n_peaks: int, nms_radius_deg: float): + """Merge several peak lists into one, ranked by z-score. + + Used for the Patterson-radius union: the same obs expanded to two different + integration radii give two rotation functions whose absolute values are not + comparable, but whose per-run standardised heights (``RotationPeak.sigma``) + are. Peaks are pooled, sorted by sigma, and greedily suppressed by SO(3) + angular distance so the same orientation found by both radii appears once. + + Parameters + ---------- + peak_lists : sequence of list of RotationPeak + One list per run. + n_peaks : int + Cap on the merged list. + nms_radius_deg : float + Suppression radius, in degrees of SO(3) geodesic distance. + + Returns + ------- + list of RotationPeak + """ + from torchref.experimental.alignment.frf.rotation_utils import ( + rotation_angular_distance_deg, + rotation_matrix_from_edmonds_euler, + ) + + pooled = [p for pl in peak_lists for p in pl] + pooled.sort(key=lambda p: p.sigma, reverse=True) + kept, kept_R = [], [] + for p in pooled: + R = rotation_matrix_from_edmonds_euler(p.alpha, p.beta, p.gamma) + if any(rotation_angular_distance_deg(R, Rk) < nms_radius_deg + for Rk in kept_R): + continue + kept.append(p) + kept_R.append(R) + if len(kept) >= n_peaks: + break + return kept + + def run_frf( model, data, From 9826adc6efc101ff3ec66eccfe9e0f4bacb5e03c Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 12:05:47 +0200 Subject: [PATCH 016/250] Remove the FRF's dead modules, duplicates and unread state Two whole modules had no reader anywhere: `frf/bessel.py`, whose `spherical_bessel_table` duplicates the one in `data_mr.py` that the engine actually calls, and `frf/spherical_y.py`, a second implementation of the normalised Legendre recurrence and `Y_lm` that `sh.py` already provides. `frf/wigner_d.py` re-exported five names from `..wigner` and called none of them, and its `small_d_stable` was superseded when the same `J_y` eigendecomposition was inlined into `wigner_contraction_per_beta`. The one test that imported `small_d_packed` through the shim now takes it from `..wigner` directly, where it serves as the independent reference implementation the contraction is checked against. `WignerContraction` was never instantiated, `AdaptiveRotationFunction. total_samples()` never called, and `BesselSHCoefficients.N_radial` duplicates `coeffs.shape[0]`. `bessel_h_scale`, `beta_grid` and `beta_starts` are kept: they are not derivable from the arrays and a reader needs them to interpret the object. Also: `_build_beta_grid` (never called; its logic is inlined at the three use sites), `peak_finder._so3_angular_distance_deg` (the NMS re-inlines the cosine test, and the test that checks it carries its own copy), and a line in `adjust_gridding` that computed `primes` only for the next line to overwrite it. `_euler_to_matrix_edmonds_zyz` stays in `peak_finder` rather than deferring to `rotation_utils`: the two are algebraically equal but round differently in the last bit, which flips NMS suppression for pairs on the threshold. Its docstring now says so, so the duplication is deliberate rather than accidental. `bench_stages` timed `frf.bessel.spherical_bessel_table` and `wigner_d.small_d_stable`, neither of which the engine calls -- the README already recorded that they registered zero calls. It now times the stages that run. Verified behaviour-preserving: single-threaded, the peak list is bit-identical before and after on 1DAW and 3GR5. (Multi-threaded it is not reproducible even against itself -- see the following commit.) Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/diagnostics/bench_stages.py | 6 +- tests/unit/frf_separate/test_invariants.py | 3 +- .../experimental/alignment/frf/__init__.py | 2 - torchref/experimental/alignment/frf/bessel.py | 97 ------------------- .../experimental/alignment/frf/data_mr.py | 1 - .../experimental/alignment/frf/peak_finder.py | 17 ++-- .../alignment/frf/sitelist_ang.py | 11 --- .../experimental/alignment/frf/spherical_y.py | 96 ------------------ torchref/experimental/alignment/frf/types.py | 18 ---- .../experimental/alignment/frf/wigner_d.py | 72 ++------------ 10 files changed, 21 insertions(+), 302 deletions(-) delete mode 100644 torchref/experimental/alignment/frf/bessel.py delete mode 100644 torchref/experimental/alignment/frf/spherical_y.py diff --git a/alignment_lab/diagnostics/bench_stages.py b/alignment_lab/diagnostics/bench_stages.py index ccc7421d..2f8d1645 100644 --- a/alignment_lab/diagnostics/bench_stages.py +++ b/alignment_lab/diagnostics/bench_stages.py @@ -34,9 +34,11 @@ STAGES = [ ("torchref.experimental.alignment.frf.dense_calc", "dense_calc_via_box"), ("torchref.experimental.alignment.frf.api", "phaser_rotation_search"), - ("torchref.experimental.alignment.frf.bessel", "spherical_bessel_table"), - ("torchref.experimental.alignment.frf.wigner_d", "small_d_stable"), + ("torchref.experimental.alignment.frf.data_mr", "spherical_bessel_table"), + ("torchref.experimental.alignment.frf.data_mr", "bessel_sh_expand"), + ("torchref.experimental.alignment.frf.data_mr", "cross_correlate_xi"), ("torchref.experimental.alignment.frf.wigner_d", "wigner_contraction_per_beta"), + ("torchref.experimental.alignment.frf.sitelist_ang", "evaluate_rotation_function"), ("torchref.experimental.alignment.frf.peak_finder", "find_rotation_peaks"), ] diff --git a/tests/unit/frf_separate/test_invariants.py b/tests/unit/frf_separate/test_invariants.py index 687ab8e2..bebafbe8 100644 --- a/tests/unit/frf_separate/test_invariants.py +++ b/tests/unit/frf_separate/test_invariants.py @@ -17,7 +17,8 @@ build_dense_map_per_beta, evaluate_rotation_function, ) -from torchref.experimental.alignment.frf.wigner_d import small_d_packed, wigner_contraction_per_beta +from torchref.experimental.alignment.frf.wigner_d import wigner_contraction_per_beta +from torchref.experimental.alignment.wigner import small_d_packed def _make_xi(L: int, seed: int = 42) -> torch.Tensor: diff --git a/torchref/experimental/alignment/frf/__init__.py b/torchref/experimental/alignment/frf/__init__.py index e30100a4..da26bbd9 100644 --- a/torchref/experimental/alignment/frf/__init__.py +++ b/torchref/experimental/alignment/frf/__init__.py @@ -25,7 +25,6 @@ AdaptiveRotationFunction, BesselSHCoefficients, RotationPeak, - WignerContraction, ) __all__ = [ @@ -44,5 +43,4 @@ "AdaptiveRotationFunction", "BesselSHCoefficients", "RotationPeak", - "WignerContraction", ] diff --git a/torchref/experimental/alignment/frf/bessel.py b/torchref/experimental/alignment/frf/bessel.py deleted file mode 100644 index 56a2e863..00000000 --- a/torchref/experimental/alignment/frf/bessel.py +++ /dev/null @@ -1,97 +0,0 @@ -"""Spherical Bessel functions ``j_u(x)`` for the Bessel-SH expansion. - -Phaser source: Phaser uses Miller downward recurrence in its own -implementation (``phaser/lib/jiffy.h``) for stability. We use the same -recurrence here because ``scipy.special.spherical_jn`` is not vectorised -across torch tensors and goes via NumPy round-trip — too slow for the -``(N_radial × N_obs)`` table required by ``bessel_sh_expand``. - -Numerical reference for Miller downward recurrence: - Abramowitz & Stegun §10.1.19 — start at high u, recur downward, then - rescale by the analytic ``j_0(x) = sin(x)/x``. -""" -from __future__ import annotations - -import torch - - -def spherical_bessel_table( - x: torch.Tensor, - u_max: int, -) -> torch.Tensor: - """Tabulate ``j_u(x)`` for u ∈ [0, u_max] over a 1-D tensor of x. - - Returns - ------- - j : torch.Tensor, shape (u_max + 1, x.numel()) - ``j[u, i] = j_u(x[i])``. Always float64 internally. - """ - if x.ndim != 1: - raise ValueError(f"expected 1-D x, got shape {tuple(x.shape)}") - if u_max < 0: - raise ValueError(f"u_max must be >= 0, got {u_max}") - - device = x.device - x64 = x.to(torch.float64) - n = x64.numel() - out = torch.zeros((u_max + 1, n), dtype=torch.float64, device=device) - - # Handle x == 0 separately (j_0(0)=1, j_u(0)=0 for u>=1). - zero_mask = x64 == 0 - nz_mask = ~zero_mask - out[0, zero_mask] = 1.0 - if not nz_mask.any(): - return out - - xnz = x64[nz_mask] - nnz = xnz.numel() - - # Direct formulas for u = 0, 1 — accurate for all x > 0. - j0 = torch.sin(xnz) / xnz - j1 = (torch.sin(xnz) - xnz * torch.cos(xnz)) / xnz**2 - - if u_max == 0: - out[0, nz_mask] = j0 - return out - if u_max == 1: - out[0, nz_mask] = j0 - out[1, nz_mask] = j1 - return out - - # Forward recurrence is unstable when u >> x; switch to Miller downward - # recurrence for those entries. Cutoff u_fwd_safe = ceil(x) is generous; - # Phaser uses similar logic in jiffy.h. - u_fwd_safe = torch.clamp(torch.ceil(xnz).to(torch.int64), min=2) - # Allocate per-x downward recurrence with a generous starting index - # (u_max + a few extra terms) — Miller convention. - u_start = u_max + 15 - f_curr = torch.zeros(nnz, dtype=torch.float64, device=device) - f_next = torch.ones(nnz, dtype=torch.float64, device=device) - table = torch.zeros((u_max + 1, nnz), dtype=torch.float64, device=device) - - # Downward: j_{u-1}(x) = (2u+1)/x · j_u(x) - j_{u+1}(x) - for u in range(u_start, -1, -1): - f_prev = (2 * u + 1) / xnz * f_next - f_curr - if u <= u_max: - table[u] = f_prev - f_curr = f_next - f_next = f_prev - - # Rescale so table[0] matches analytic j_0(x) = sin(x)/x. - scale = j0 / table[0] - table = table * scale.unsqueeze(0) - - # Override with the forward direct values where they are stable - # (small u, small x). Forward recurrence: j_{u+1} = (2u+1)/x · j_u - j_{u-1}. - fwd = torch.zeros((u_max + 1, nnz), dtype=torch.float64, device=device) - fwd[0] = j0 - fwd[1] = j1 - for u in range(1, u_max): - fwd[u + 1] = (2 * u + 1) / xnz * fwd[u] - fwd[u - 1] - # Per-x, use forward where u <= u_fwd_safe; otherwise Miller-rescaled. - u_idx = torch.arange(u_max + 1, device=device).unsqueeze(1) # (u_max+1, 1) - use_fwd = u_idx <= u_fwd_safe.unsqueeze(0) # (u_max+1, nnz) - table_combined = torch.where(use_fwd, fwd, table) - - out[:, nz_mask] = table_combined - return out diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index 7c9e76e7..61d17d73 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -309,7 +309,6 @@ def _tick(t0): return BesselSHCoefficients( coeffs=c_nlm, L=L, - N_radial=N_radial, bessel_h_scale=float(bessel_h_scale), ) diff --git a/torchref/experimental/alignment/frf/peak_finder.py b/torchref/experimental/alignment/frf/peak_finder.py index dce37d7d..91255f4e 100644 --- a/torchref/experimental/alignment/frf/peak_finder.py +++ b/torchref/experimental/alignment/frf/peak_finder.py @@ -32,6 +32,13 @@ def _euler_to_matrix_edmonds_zyz( """R = R_z(α) R_y(β) R_z(γ) — Edmonds ZYZ convention. Returns shape (*alpha.shape, 3, 3) real. + + Algebraically ``rotation_utils.rotation_matrix_from_edmonds_euler_batch``, + but written as one fused pass rather than three matrix products. The NMS + below evaluates it over ~1e4 candidates, and the two forms round + differently in the last bit, which flips the suppression decision for pairs + sitting on the threshold. Every measurement on this engine was made with + this form, so it stays. """ ca, sa = torch.cos(alpha), torch.sin(alpha) cb, sb = torch.cos(beta), torch.sin(beta) @@ -48,16 +55,6 @@ def _euler_to_matrix_edmonds_zyz( return R -def _so3_angular_distance_deg(R1: torch.Tensor, R2: torch.Tensor) -> torch.Tensor: - """Angular distance between two rotation matrices, in degrees. - - R1: (..., 3, 3), R2: (..., 3, 3). Returns (...,) real. - """ - trace = torch.einsum("...ij,...ij->...", R1, R2) - cos_theta = ((trace - 1.0) * 0.5).clamp(min=-1.0, max=1.0) - return torch.arccos(cos_theta) * (180.0 / math.pi) - - def _so3_greedy_nms( alphas: torch.Tensor, betas: torch.Tensor, diff --git a/torchref/experimental/alignment/frf/sitelist_ang.py b/torchref/experimental/alignment/frf/sitelist_ang.py index 2662558c..1dba0fec 100644 --- a/torchref/experimental/alignment/frf/sitelist_ang.py +++ b/torchref/experimental/alignment/frf/sitelist_ang.py @@ -59,7 +59,6 @@ def adjust_gridding(target: int, max_prime: int = 5) -> int: """ if target <= 1: return 1 - primes = [2, 3, 5, 7, 11, 13][: min(max(max_prime // 2, 1), 6)] primes = [p for p in [2, 3, 5, 7, 11, 13] if p <= max_prime] n = int(target) while True: @@ -134,16 +133,6 @@ def build_dense_map_per_beta( return M -def _build_beta_grid(grid_sampling_deg: float) -> Tuple[torch.Tensor, int]: - """Return (β_grid in radians, bmax). β = b · Δ for b ∈ [0, bmax).""" - bmax = int(math.ceil(180.0 / grid_sampling_deg)) - if bmax < 1: - raise ValueError(f"grid_sampling_deg={grid_sampling_deg} too coarse") - b = torch.arange(bmax, dtype=torch.float64) - betas_rad = b * grid_sampling_deg * (math.pi / 180.0) - return betas_rad, bmax - - # Module-level memo for the data-independent sample list, keyed on # (grid_sampling_deg, device-str, dtype). The list depends only on geometry, so # repeat FRF calls at the same grid reuse it; a single cold call still pays the diff --git a/torchref/experimental/alignment/frf/spherical_y.py b/torchref/experimental/alignment/frf/spherical_y.py deleted file mode 100644 index 926768bd..00000000 --- a/torchref/experimental/alignment/frf/spherical_y.py +++ /dev/null @@ -1,96 +0,0 @@ -"""Spherical harmonics ``Y_l^m(θ, φ)`` table. - -Phaser source: ``phaser/lib/sphericalY.h`` (Condon-Shortley phase, the -standard physics convention). - -We build the table once per (θ, φ) batch and reuse across all -``(l, m)`` indices needed. -""" -from __future__ import annotations - -import math - -import torch - - -def _normalised_associated_legendre( - cos_theta: torch.Tensor, L: int -) -> torch.Tensor: - """Compute the *normalised* associated Legendre polynomials. - - Returns ``P[l, m, ...] = sqrt((2l+1)/(4π) · (l-m)!/(l+m)!) · P_l^m(cos θ)`` - for l ∈ [0, L), m ∈ [0, l], with the standard recurrence (no Condon-Shortley - sign — that's applied at the Y_lm level). - - Out-of-range entries (m > l) are zero. - """ - device = cos_theta.device - x = cos_theta.to(torch.float64) - sin_t = torch.sqrt(torch.clamp(1.0 - x * x, min=0.0)) - - plm = torch.zeros((L, L, *x.shape), dtype=torch.float64, device=device) - # m = 0 sector: standard Legendre recurrence with normalisation built in. - plm[0, 0] = math.sqrt(1.0 / (4.0 * math.pi)) - if L > 1: - plm[1, 0] = math.sqrt(3.0 / (4.0 * math.pi)) * x - for l in range(2, L): - a = math.sqrt((2 * l + 1) * (2 * l - 1)) / l - b = math.sqrt((2 * l + 1) / (2 * l - 3)) * (l - 1) / l - plm[l, 0] = a * x * plm[l - 1, 0] - b * plm[l - 2, 0] - - # m > 0 sector: build P_l^l from the previous diagonal, then recur up in l. - for m in range(1, L): - plm[m, m] = -math.sqrt((2 * m + 1) / (2 * m)) * sin_t * plm[m - 1, m - 1] - if m + 1 < L: - plm[m + 1, m] = math.sqrt(2 * m + 3) * x * plm[m, m] - for l in range(m + 2, L): - a = math.sqrt((2 * l + 1) * (2 * l - 1) / ((l - m) * (l + m))) - b = math.sqrt( - (2 * l + 1) * (l + m - 1) * (l - m - 1) - / ((l - m) * (l + m) * (2 * l - 3)) - ) - plm[l, m] = a * x * plm[l - 1, m] - b * plm[l - 2, m] - - return plm - - -def ylm_table( - theta: torch.Tensor, phi: torch.Tensor, L: int -) -> torch.Tensor: - """Compute ``Y_l^m(θ, φ)`` for all l ∈ [0, L), m ∈ [-l, l]. - - Convention: Condon-Shortley phase (the ``(-1)^m`` factor is in - ``Y_l^m`` for m > 0). The normalised associated Legendre polynomial - is real; the φ dependence is ``exp(i m φ)``. - - Parameters - ---------- - theta, phi : torch.Tensor (real, same shape) - Polar (θ ∈ [0, π]) and azimuthal (φ ∈ [0, 2π)) angles. - - Returns - ------- - Y : torch.Tensor (complex128), shape (L, 2L-1, *theta.shape) - ``Y[l, m + L - 1, ...] = Y_l^m(θ, φ)`` for |m| ≤ l, else 0. - """ - if theta.shape != phi.shape: - raise ValueError( - f"theta {tuple(theta.shape)} and phi {tuple(phi.shape)} must agree" - ) - device = theta.device - plm = _normalised_associated_legendre(torch.cos(theta), L) # (L, L, *) - Y = torch.zeros( - (L, 2 * L - 1, *theta.shape), dtype=torch.complex128, device=device - ) - phi64 = phi.to(torch.float64) - # m = 0 - Y[:, L - 1, ...] = plm[:, 0, ...].to(torch.complex128) - for m in range(1, L): - e_pos = torch.complex(torch.cos(m * phi64), torch.sin(m * phi64)) - e_neg = e_pos.conj() - # Y_l^{+m} = (-1)^m * sqrt(...) P_l^m(cosθ) * exp(i m φ), already - # includes the (-1)^m if we apply it here. - sign = (-1) ** m - Y[:, L - 1 + m, ...] = sign * plm[:, m, ...].to(torch.complex128) * e_pos - Y[:, L - 1 - m, ...] = plm[:, m, ...].to(torch.complex128) * e_neg - return Y diff --git a/torchref/experimental/alignment/frf/types.py b/torchref/experimental/alignment/frf/types.py index 678da364..14d85108 100644 --- a/torchref/experimental/alignment/frf/types.py +++ b/torchref/experimental/alignment/frf/types.py @@ -8,7 +8,6 @@ from __future__ import annotations from dataclasses import dataclass -from typing import List, Tuple import torch @@ -26,23 +25,9 @@ class BesselSHCoefficients: coeffs: torch.Tensor # (N_radial, L, 2L-1), complex L: int # angular bandlimit (lmax = L - 1) - N_radial: int # number of radial nodes bessel_h_scale: float # h = bessel_h_scale * |s| -@dataclass -class WignerContraction: - """The per-β ``S_{m1,m2}(β) = Σ_l ξ_{l,m1,m2} · d^l_{m1,m2}(β)``. - - Phaser source: ``SiteListAng::DoRfftStuff`` (FastRot.cc:39-59) builds - this sum into the ``rot`` accumulator before pushing into a pseudo-SF - array for FFT. - """ - - S: torch.Tensor # (n_beta, 2L-1, 2L-1), complex - betas: torch.Tensor # (n_beta,), real - - @dataclass class AdaptiveRotationFunction: """Rotation function values on Phaser's adaptive SO(3) sample list. @@ -65,9 +50,6 @@ class AdaptiveRotationFunction: beta_grid: torch.Tensor # (n_beta,) real — the β values themselves grid_sampling_deg: float - def total_samples(self) -> int: - return int(self.values.numel()) - @dataclass class RotationPeak: diff --git a/torchref/experimental/alignment/frf/wigner_d.py b/torchref/experimental/alignment/frf/wigner_d.py index 7dfc6973..b90fe96f 100644 --- a/torchref/experimental/alignment/frf/wigner_d.py +++ b/torchref/experimental/alignment/frf/wigner_d.py @@ -3,76 +3,20 @@ Phaser source: ``phaser/lib/wigner.h`` (the C++ template ``djmn_recursive_table`` used in ``FastRot.cc:41`` per-l, per-β). -Phaser uses the Sakurai recurrence convention; we use the equivalent -Edmonds (4.1.23) direct-sum formula, already validated against Phaser's -output by the convention tests in -``tests/unit/alignment/test_wigner.py``. To avoid duplicating maths, -this module re-exports the existing implementation from -``torchref.experimental.alignment.wigner`` (which is the same convention) and adds -Phaser-specific helpers on top. +Phaser uses the Sakurai recurrence convention; the equivalent Edmonds +(4.1.23) convention is used throughout this package and is pinned against +Phaser's output by ``tests/unit/alignment/test_wigner.py``. +``wigner_contraction_per_beta`` builds the small-d table it needs from the +``J_y`` eigendecomposition inline, which stays bounded to any ``l``. """ from __future__ import annotations -from typing import Tuple - import torch -# Re-export the existing validated implementations. -from ..wigner import ( - _build_half_angle_pow_tables, - small_d_block, - small_d_packed, - small_d_table, - wigner_D_pointwise, -) - -__all__ = [ - "small_d_block", - "small_d_packed", - "small_d_stable", - "small_d_table", - "wigner_D_pointwise", - "wigner_contraction_per_beta", -] - - -def small_d_stable(L: int, betas: torch.Tensor) -> torch.Tensor: - """Numerically stable Wigner small-d table via J_y diagonalization. - - ``d^l(β) = exp(-iβ J_y^{(l)})``. J_y is a tiny real tridiagonal generator, - so ``eigh(i·J_y)`` gives integer eigenvalues μ ∈ [-l, l] and a basis V with - ``d^l_{m n}(β) = Σ_μ V_{m μ} e^{-iβ μ} V*_{n μ}`` (real). - This is exactly the π/2 / SOFT Fourier-over-μ decomposition (V are the π/2 - matrices up to phase), and it has NO catastrophic cancellation — unlike the - Edmonds direct-sum ``small_d_packed`` which explodes to |d|~1e11 at l≥50. - - Returns ``(n_beta, L, 2L-1, 2L-1)`` with ``d[k, l, m+L-1, n+L-1] = d^l_{m,n}(β_k)``, - matching ``small_d_packed`` exactly (same convention, verified at l≤40). - """ - betas = betas.to(torch.float64) - n_beta = betas.shape[0] - dim = 2 * L - 1 - device = betas.device - out = torch.zeros((n_beta, L, dim, dim), dtype=torch.float64, device=device) - out[:, 0, L - 1, L - 1] = 1.0 # l=0 - for l in range(1, L): - sz = 2 * l + 1 - p = torch.arange(sz - 1, dtype=torch.float64, device=device) - sup = 0.5 * torch.sqrt((2 * l - p) * (p + 1.0)) # J_y off-diagonal magnitudes - A = torch.diag(sup, 1) - torch.diag(sup, -1) # A = -i J_y, real antisymmetric - H = 1j * A.to(torch.complex128) # Hermitian - w, V = torch.linalg.eigh(H) # w≈[-l..l], V complex - # d_l(β) = Re( V · diag(e^{-iβ w}) · V^H ), batched over β - phase = torch.exp(-1j * betas.unsqueeze(1) * w.unsqueeze(0)) # (n_beta, sz) - d_l = torch.einsum("ma,ka,na->kmn", V, phase, V.conj()).real # (n_beta, sz, sz) - lo, hi = L - 1 - l, L - 1 + l + 1 - out[:, l, lo:hi, lo:hi] = d_l - return out - +__all__ = ["wigner_contraction_per_beta"] -# Cache of the J_y eigendecomposition per l. (w_l, V_l) depend only on l (and -# device), not on β or the data, so they are computed once per (L, device) and -# reused across every FRF call. Keyed on (L, device-str). +#: Memo for the per-l J_y eigendecomposition, keyed on (L, device-str). +#: It depends only on the bandwidth, so repeat calls at the same L reuse it. _WIGNER_EIG_CACHE: dict = {} From 50dbe65a1948abf80eccc8b9067f4ccccd2c08a8 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 12:07:00 +0200 Subject: [PATCH 017/250] Pin the FRF's reproducibility contract The engine carries no RNG and no order-dependent atomics, so with one thread two identical calls give a bit-identical peak list. At the default thread count they do not: the structure-factor reduction is float32 and its parallel summation order varies, which on 3GR5 (P 6_5 2 2) moves peak scores by ~5e-8 relative and reorders about a dozen of 500 peaks. On 1DAW the same noise is ~8e-16 relative and nothing moves. That sets the resolution floor for any peak-list comparison, and it means a refactor has to be checked single-threaded to be checked at all. Truth rank was stable across threaded repeats on both structures, but two peaks within 5e-8 of each other are not ordered reliably. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- tests/unit/frf_separate/test_synthetic.py | 44 +++++++++++++++++++++++ 1 file changed, 44 insertions(+) diff --git a/tests/unit/frf_separate/test_synthetic.py b/tests/unit/frf_separate/test_synthetic.py index a768356f..d708c01a 100644 --- a/tests/unit/frf_separate/test_synthetic.py +++ b/tests/unit/frf_separate/test_synthetic.py @@ -137,3 +137,47 @@ def test_api_returns_correct_types(): assert isinstance(arf, AdaptiveRotationFunction) assert len(peaks) > 0 assert all(isinstance(p, RotationPeak) for p in peaks) + + +def test_search_is_bit_reproducible_single_threaded(): + """The engine itself is deterministic: no RNG, no order-dependent atomics. + + Repeat runs at the default thread count are NOT bit-identical -- the + structure-factor reduction is float32 and its parallel summation order + varies, which on a real high-symmetry case (3GR5, P 6_5 2 2) perturbs peak + scores by ~5e-8 relative and reorders roughly a dozen of 500 peaks. Truth + rank was stable there, but nothing guarantees it for two peaks that close. + + Pinning one thread removes the only source of variation, so any future + difference under this test is a real change in the maths. It is also the + configuration to use when comparing peak lists across a refactor. + """ + import torch as _torch + + s_calc, F_calc, centric = _make_random_reflections(seed=3, n=800) + sym_mats = _torch.eye(3, dtype=_torch.float64).unsqueeze(0) + kwargs = dict( + sym_mats=sym_mats, L=12, d_min=4.0, d_max=15.0, delta_vrms_A=0.5, + grid_sampling_deg=8.0, n_peaks=25, sigma_threshold=-100.0, + use_french_wilson=False, use_m_symmetry_filter=False, + ) + + n_threads = _torch.get_num_threads() + _torch.set_num_threads(1) + try: + _, first = phaser_rotation_search(s_calc, F_calc, centric, + s_calc, F_calc, **kwargs) + _, second = phaser_rotation_search(s_calc, F_calc, centric, + s_calc, F_calc, **kwargs) + finally: + _torch.set_num_threads(n_threads) + + assert len(first) == len(second) > 0 + for i, (a, b) in enumerate(zip(first, second)): + assert a.alpha == b.alpha and a.beta == b.beta and a.gamma == b.gamma, ( + f"peak {i} moved between two identical single-threaded runs" + ) + assert a.score == b.score, ( + f"peak {i} score changed by {abs(a.score - b.score):.3g} between " + f"two identical single-threaded runs" + ) From b159e8101198863f8c0efec82c51b3b91c457598 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 12:13:26 +0200 Subject: [PATCH 018/250] Drop the alignment pipeline's write-only state and dead kwargs `FRFInputs` carried four fields no code anywhere reads: `s_vec_for_search`, `s_mag_sym`, `patt_obs` and `patt_calc`. The rotation search recomputes everything it needs from `data` and takes only `U_aniso` and `device` from the dataclass; the rescore and translation stages read the per-reflection arrays. Building the dead four cost a symmetry expansion of the whole Miller list plus a full Lattman-Love interpolation over it (`n_ops x N` reflections) on every run, with the result discarded. `_shellbin_norm_etrick` existed only to feed them. `MolecularReplacementPipeline` also accepted eight kwargs it stored in a `_vestigial` dict that nothing reads, and `align_model_to_data` forwarded all eight plus `L` from its own signature. Removed, along with the flags three integration drivers passed down to them -- including four `--use-*` CLI options whose help text describes experiments that no longer have an implementation behind them. Verified behaviour-preserving: single-threaded, the 1DAW peak list is bit-identical before and after. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- .../alignment/run_random_pdb_fit.py | 45 +------- .../integration/alignment/test_fit_to_data.py | 2 +- .../alignment/test_fit_to_data_translation.py | 2 +- torchref/experimental/alignment/align.py | 107 +++--------------- torchref/experimental/alignment/pipeline.py | 22 ---- 5 files changed, 18 insertions(+), 160 deletions(-) diff --git a/tests/integration/alignment/run_random_pdb_fit.py b/tests/integration/alignment/run_random_pdb_fit.py index ded16b08..b4beb24d 100644 --- a/tests/integration/alignment/run_random_pdb_fit.py +++ b/tests/integration/alignment/run_random_pdb_fit.py @@ -96,14 +96,7 @@ def run(pdb_key: str, seed: int, verbose: int = 1, sigma_rot_deg: float = 0.0, sigma_trans_ang: float = 0.0, sigma_b: float = 0.0, - use_sigma_a_frf: bool = False, - frf_delta_vrms_A: float = 1.0, - frf_weight_combine: str = "sigma_a_only", n_rotation_candidates: int = 15, - use_m_symmetry_filter: bool = False, - use_lerf1_intensity: bool = False, - use_fitted_delta_vrms: bool = False, - use_even_l_only: bool = False, rescore_engine: str = "m_letf1") -> dict: pdb_path, mtz_path = PAIRS[pdb_key] print(f"\n=== {pdb_key}: {pdb_path.name} + {mtz_path.name} ===", flush=True) @@ -162,7 +155,7 @@ def _scale_and_r(m: ModelFT) -> tuple[float, float]: rotated_search, data, d_min=4.0, d_max=15.0, - L=32, n_shells=20, + n_shells=20, n_rotation_peaks=200, n_ml_refine=200, verbose=verbose, use_interp_var=use_interp_var, @@ -171,14 +164,7 @@ def _scale_and_r(m: ModelFT) -> tuple[float, float]: sigma_rot_deg=sigma_rot_deg, sigma_trans_ang=sigma_trans_ang, sigma_b=sigma_b, - use_sigma_a_frf=use_sigma_a_frf, - frf_delta_vrms_A=frf_delta_vrms_A, - frf_weight_combine=frf_weight_combine, n_rotation_candidates=n_rotation_candidates, - use_m_symmetry_filter=use_m_symmetry_filter, - use_lerf1_intensity=use_lerf1_intensity, - use_fitted_delta_vrms=use_fitted_delta_vrms, - use_even_l_only=use_even_l_only, rescore_engine=rescore_engine, ) fit_time = time.time() - t1 @@ -242,29 +228,9 @@ def main(): ap.add_argument("--sigma-b", type=float, default=0.0, help="Phase C: Gaussian B-factor restraint sigma (Ų). " "0 = no restraint. Phaser default ~15.") - ap.add_argument("--use-sigma-a-frf", action="store_true", - help="E3: σA-weight the FRF input field (Phaser FastRot " - "Eterm/Vterm analogue). Default off.") - ap.add_argument("--frf-delta-vrms", type=float, default=1.0, - help="ΔVRMS for Luzzati σA(s) = exp(−2π²s²ΔVRMS²), Å. " - "Default 1.0.") - ap.add_argument("--frf-weight-combine", default="sigma_a_only", - choices=["sigma_a_only", "sigma_a_x_variance"], - help="How to combine σA² and empirical variance weights.") ap.add_argument("--n-rotation-candidates", type=int, default=15, help="Top-N rotations from MLRF rescore that get full " "translation+polish. Default 15.") - ap.add_argument("--use-m-symmetry-filter", action="store_true", - help="F1: zero SH coefficients with m not divisible by " - "ZSYMM (Phaser-style symmetry-aware denoiser).") - ap.add_argument("--use-lerf1-intensity", action="store_true", - help="F2: replace patt_obs = E²−1 with " - "cweight·(E²−1)·DFAC² (Phaser LERF1 form).") - ap.add_argument("--use-fitted-delta-vrms", action="store_true", - help="F3: fit ΔVRMS from /(8π²) instead of " - "frf_delta_vrms_A.") - ap.add_argument("--use-even-l-only", action="store_true", - help="F4: skip odd-l SH coefficients (perf, no SNR).") args = ap.parse_args() device = torch.device(args.device) if device.type == "cuda" and not torch.cuda.is_available(): @@ -296,14 +262,7 @@ def main(): sigma_rot_deg=args.sigma_rot_deg, sigma_trans_ang=args.sigma_trans_ang, sigma_b=args.sigma_b, - use_sigma_a_frf=args.use_sigma_a_frf, - frf_delta_vrms_A=args.frf_delta_vrms, - frf_weight_combine=args.frf_weight_combine, - n_rotation_candidates=args.n_rotation_candidates, - use_m_symmetry_filter=args.use_m_symmetry_filter, - use_lerf1_intensity=args.use_lerf1_intensity, - use_fitted_delta_vrms=args.use_fitted_delta_vrms, - use_even_l_only=args.use_even_l_only) + n_rotation_candidates=args.n_rotation_candidates) results.append(r) except Exception as exc: import traceback diff --git a/tests/integration/alignment/test_fit_to_data.py b/tests/integration/alignment/test_fit_to_data.py index a22c544f..8121bd21 100644 --- a/tests/integration/alignment/test_fit_to_data.py +++ b/tests/integration/alignment/test_fit_to_data.py @@ -90,7 +90,7 @@ def test_fit_to_data_real_1daw(real_setup, trial): search, data, d_min=4.0, d_max=15.0, - L=32, n_shells=20, + n_shells=20, n_rotation_peaks=200, n_ml_refine=200, do_translation=False, # this test only checks rotation accuracy verbose=0, diff --git a/tests/integration/alignment/test_fit_to_data_translation.py b/tests/integration/alignment/test_fit_to_data_translation.py index 60f77a8c..88a2e39e 100644 --- a/tests/integration/alignment/test_fit_to_data_translation.py +++ b/tests/integration/alignment/test_fit_to_data_translation.py @@ -67,7 +67,7 @@ def test_fit_to_data_recovers_rotation_and_translation(): perturbed, data, d_min=4.0, d_max=15.0, - L=32, n_shells=20, + n_shells=20, n_rotation_peaks=200, n_ml_refine=200, do_translation=True, do_joint_refine=True, diff --git a/torchref/experimental/alignment/align.py b/torchref/experimental/alignment/align.py index e5165b83..daf03506 100644 --- a/torchref/experimental/alignment/align.py +++ b/torchref/experimental/alignment/align.py @@ -109,28 +109,6 @@ def summary(self) -> str: # --------------------------------------------------------------------------- -def _shellbin_norm_etrick( - F: torch.Tensor, smag: torch.Tensor, P: int -) -> torch.Tensor: - """Per-shell E-trick normalisation: ``E = F / sqrt(_shell)``. - - Equal-count shell binning by sorted ``smag``; small-count shells just - inherit their normaliser from the local mean. - """ - order = torch.argsort(smag) - idx = torch.zeros_like(smag, dtype=torch.int64) - chunk = smag.numel() // P - for k in range(P): - a = k * chunk - b = (k + 1) * chunk if k < P - 1 else smag.numel() - idx[order[a:b]] = k - norm = torch.zeros_like(smag, dtype=F.dtype) - for k in range(P): - m = idx == k - norm[m] = (F[m] ** 2).mean().clamp(min=1e-30).sqrt() - return F / norm - - def _external_rwork(model: "ModelFT", data: "ReflectionData") -> float: """Full-resolution scaled R-work via the standard Scaler. @@ -224,26 +202,18 @@ def _rodrigues(omega: torch.Tensor) -> torch.Tensor: @dataclass class FRFInputs: - """Container for the spherical-harmonic rotation-search inputs. + """Prepared reflection arrays shared by the rotation-search stages. - `*_sym` arrays are symmetry-expanded across the spacegroup rotation - operators (so the Patterson SH expansion samples the full sphere). - Un-suffixed `F_obs / hkl / s_vec / s_mag / centric` are the resolution- - masked, anisotropy-corrected reflection arrays — used downstream by - the MLRF rescore, translation search and rigid-body polish. + The resolution-masked, anisotropy-corrected reflection arrays, plus the + overall anisotropy tensor and the Lattman-Love interpolator. The rotation + search reads `U_aniso` and `device`; the ML rescore, translation search and + rigid-body polish read the rest. """ - # FRF (symmetry-expanded) inputs - s_vec_for_search: torch.Tensor # (n_ops·N, 3) - s_mag_sym: torch.Tensor # (n_ops·N,) - patt_obs: torch.Tensor # (n_ops·N,) = |E_obs|² − 1 - patt_calc: torch.Tensor # (n_ops·N,) = |E_calc|² − 1 - # Per-reflection inputs (un-expanded) - F_obs: torch.Tensor # (N,) anisotropy-corrected? See `F_obs_aniso` flag + F_obs: torch.Tensor # (N,) anisotropy-corrected amplitudes hkl: torch.Tensor # (N, 3) integer Miller indices s_vec: torch.Tensor # (N, 3) reciprocal-space Cartesian s_mag: torch.Tensor # (N,) Å⁻¹ centric: torch.Tensor # (N,) bool - # Other state used downstream ll: "LattmanLoveInterpolator" U_aniso: torch.Tensor # (3, 3) Popov-Bourenkov U device: torch.device @@ -260,14 +230,11 @@ def _prepare_frf_inputs( ll_max_res_A: float = 3.0, verbose: int = 0, ) -> FRFInputs: - """Build the symmetry-expanded SH-rotation-search inputs. - - Encapsulates the data prep / anisotropy / symmetry-expansion logic - previously inlined in `align_model_to_data`. The returned dataclass - feeds both the live FRF call and the benchmark scripts. + """Prepare the reflection arrays the rotation-search stages share. - `F_obs` on the returned dataclass is the *anisotropy-corrected* value - (matches what previously was the `F_obs_aniso` local variable). + Masks the observations to ``[d_min, d_max]``, fits and applies the overall + anisotropy correction, and builds the Lattman-Love interpolator for the + model. ``F_obs`` on the returned dataclass is anisotropy-corrected. """ device = model.xyz().device @@ -316,37 +283,7 @@ def _prepare_frf_inputs( verbose=verbose, ) - # Symmetry-expand reciprocal-space points so the Patterson SH expansion - # samples the full sphere (spacegroup-invariant by construction). See - # the long comment in `align_model_to_data` for the rationale. - sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) - n_ops_sg = int(sg_mats.shape[0]) - # h' = h.R, NOT R.h -- reciprocal space transforms with the transpose - # (SpaceGroup.apply_to_hkl). The two agree only when the symmetry matrices - # are orthogonal, which they are in orthorhombic/tetragonal/cubic and - # monoclinic settings but NOT in a hexagonal basis, where S.S^T != I. Using - # R.h there mixes non-equivalent reflections into one orbit and writes - # conflicting |F| onto the same Miller index. - hkl_sym = torch.einsum("kji,nj->kni", sg_mats, hkl.to(torch.float64)) - hkl_sym_flat = hkl_sym.reshape(-1, 3) - s_vec_sym = hkl_sym_flat @ rec_basis.to(device) - s_mag_sym = s_vec_sym.norm(dim=-1) - F_obs_aniso_sym = F_obs_aniso.unsqueeze(0).expand(n_ops_sg, -1).reshape(-1) - E_obs_sym = _shellbin_norm_etrick(F_obs_aniso_sym, s_mag_sym, n_shells) - patt_obs = E_obs_sym ** 2 - 1.0 - - F_calc_sym = ll.evaluate( - torch.eye(3, dtype=torch.float32), hkl_sym_flat, data.cell, - return_amplitude=True, - ).to(torch.float64) - E_calc_sym = _shellbin_norm_etrick(F_calc_sym, s_mag_sym, n_shells) - patt_calc = E_calc_sym ** 2 - 1.0 - return FRFInputs( - s_vec_for_search=s_vec_sym, - s_mag_sym=s_mag_sym, - patt_obs=patt_obs, - patt_calc=patt_calc, F_obs=F_obs_aniso, hkl=hkl, s_vec=s_vec, @@ -603,7 +540,6 @@ def align_model_to_data( *, d_min: float = 4.0, d_max: float = 15.0, - L: int = 48, n_shells: int = 20, n_rotation_peaks: int = 500, n_ml_refine: int = 20, # rescore only the top-20 FRF peaks (refinement use case) @@ -625,13 +561,6 @@ def align_model_to_data( sigma_rot_deg: float = 0.0, sigma_trans_ang: float = 0.0, sigma_b: float = 0.0, - use_sigma_a_frf: bool = False, - frf_delta_vrms_A: float = 1.0, - frf_weight_combine: str = "sigma_a_only", - use_m_symmetry_filter: bool = False, - use_lerf1_intensity: bool = False, - use_fitted_delta_vrms: bool = False, - use_even_l_only: bool = False, frf_lmax_cap: int = 48, frf_dense_pad: float = 2.0, rescore_engine: str = "m_letf1", @@ -656,11 +585,10 @@ def align_model_to_data( "Cannot fit an uninitialized ModelFT. Load PDB data first." ) - # `MolecularReplacementPipeline` is the implementation of record. This - # function preserves the historical kwarg surface (so the benchmark - # scripts keep working unchanged) and returns the single - # best `ModelFT`; drive the pipeline directly to get the ranked candidate - # list. Imported lazily to avoid an import cycle — `pipeline` imports the + # `MolecularReplacementPipeline` is the implementation of record; this + # function returns its single best `ModelFT`. Drive the pipeline directly to + # get the ranked candidate list. Imported lazily to avoid an import cycle -- + # `pipeline` imports the # stage helpers (`_prepare_frf_inputs`, `_run_frf_separate_rotation`, # `_external_rwork`, `_DirectModelEvaluator`, `_rodrigues`, `_StageTimer`) # from this module. @@ -692,13 +620,6 @@ def align_model_to_data( refine_b=refine_b, sigma_rot_deg=sigma_rot_deg, sigma_trans_ang=sigma_trans_ang, sigma_b=sigma_b, - L=L, - use_sigma_a_frf=use_sigma_a_frf, frf_delta_vrms_A=frf_delta_vrms_A, - frf_weight_combine=frf_weight_combine, - use_m_symmetry_filter=use_m_symmetry_filter, - use_lerf1_intensity=use_lerf1_intensity, - use_fitted_delta_vrms=use_fitted_delta_vrms, - use_even_l_only=use_even_l_only, ) solutions = pipeline.run(do_translation=do_translation) return solutions[0].model diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 04cdaafb..ef36e36a 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -261,16 +261,6 @@ def __init__( min_tries: int = 3, max_tries: Optional[int] = None, rfactor_converged: float = 0.45, - # --- vestigial FRF knobs (superseded by _run_frf_separate_rotation's - # frf_use_* defaults post-consolidation; accepted for API stability) --- - L: int = 48, - use_sigma_a_frf: bool = False, - frf_delta_vrms_A: float = 1.0, - frf_weight_combine: str = "sigma_a_only", - use_m_symmetry_filter: bool = False, - use_lerf1_intensity: bool = False, - use_fitted_delta_vrms: bool = False, - use_even_l_only: bool = False, ): self.data = data self.model = model @@ -316,18 +306,6 @@ def __init__( self.max_tries = max_tries self.rfactor_converged = rfactor_converged - # Vestigial; retained so legacy callers/benchmarks do not break. - self._vestigial = dict( - L=L, - use_sigma_a_frf=use_sigma_a_frf, - frf_delta_vrms_A=frf_delta_vrms_A, - frf_weight_combine=frf_weight_combine, - use_m_symmetry_filter=use_m_symmetry_filter, - use_lerf1_intensity=use_lerf1_intensity, - use_fitted_delta_vrms=use_fitted_delta_vrms, - use_even_l_only=use_even_l_only, - ) - self._timer = _StageTimer(enabled=verbose >= 2) # Filled in by run(). self._frf = None From 541da9b2c65873dd4eb1e2d40aac8e18a1bad80b Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 13:15:37 +0200 Subject: [PATCH 019/250] Report arms against the shipping criterion, not just against each other `--gate` scores each arm on what the pipeline actually needs: truth inside the top N candidates it carries forward, on most trials, for *every* structure. Rank 7 and rank 0 are the same outcome downstream and rank 223 is not, so a median rank hides the thing that matters and an average over structures lets one failing space group be cancelled by nine easy ones. It also pairs on `rank_for_compare` instead of `truth_rank`. `compare` drops any pair where either side missed, which silently discards exactly the cells an arm is being blamed for. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/aggregate.py | 110 ++++++++++++++++++++++++++++ 1 file changed, 110 insertions(+) diff --git a/alignment_lab/analysis/aggregate.py b/alignment_lab/analysis/aggregate.py index 20a2cdb5..1563849a 100644 --- a/alignment_lab/analysis/aggregate.py +++ b/alignment_lab/analysis/aggregate.py @@ -108,6 +108,105 @@ def compare(rows: List[dict], key: str, base: str = None) -> None: f"benchmark needs; treat as indicative only") +def _cmp_rank(row: dict, n_peaks_default: int = 500) -> float: + """Rank for pairing: a miss counts as worse than the worst hit. + + ``compare`` drops pairs where either side missed, which silently removes + exactly the cases an arm is being blamed for. The sweep writes + ``rank_for_compare`` for this; fall back to the peak-list length. + """ + try: + return float(int(row["rank_for_compare"])) + except (KeyError, ValueError): + pass + try: + r = int(row["truth_rank"]) + except (KeyError, ValueError): + return float("nan") + if r >= 0: + return float(r) + try: + return float(int(row["n_peaks"])) + except (KeyError, ValueError): + return float(n_peaks_default) + + +def gate(rows: List[dict], key: str = "arm", base: str = "production", + top_n: int = 20, min_hits: int = 9) -> None: + """Report each arm against the shipping criterion, per structure. + + The criterion is not "truth at rank 0": the pipeline carries the top ~20 + candidates forward, so rank 7 and rank 0 are the same outcome and rank 223 + is not. An arm passes when truth lands in the top ``top_n`` on at least + ``min_hits`` of the trials for **every** structure -- so one bad structure + cannot be averaged away by nine good ones. + """ + arms = sorted({r.get(key, "") for r in rows}) + pdbs = sorted({r.get("pdb", "?") for r in rows}) + per: Dict[tuple, List[dict]] = defaultdict(list) + for r in rows: + per[(r.get(key, ""), r.get("pdb", "?"))].append(r) + + def _in_top(r: dict) -> bool: + try: + v = int(r["truth_rank"]) + except (KeyError, ValueError): + return False + return 0 <= v < top_n + + print(f"\n# shipping gate: truth in the top {top_n} on >= {min_hits} trials, " + f"for every structure") + print(f"{key:<26} {'pass':>5} {'worst structure':>16} {'total':>7} " + f"{'rank0':>6} {'median':>7}") + verdicts = {} + for arm in arms: + worst_pdb, worst_hits, tot, hits, rank0, ranks = None, None, 0, 0, 0, [] + for pdb in pdbs: + rs = per.get((arm, pdb), []) + if not rs: + continue + h = sum(1 for r in rs if _in_top(r)) + if worst_hits is None or h < worst_hits: + worst_hits, worst_pdb = h, f"{pdb} {h}/{len(rs)}" + tot += len(rs); hits += h + rank0 += sum(1 for r in rs if str(r.get("truth_rank")) == "0") + ranks += [_cmp_rank(r) for r in rs] + ok = worst_hits is not None and worst_hits >= min_hits + verdicts[arm] = ok + med = statistics.median(ranks) if ranks else float("nan") + print(f"{arm:<26} {'PASS' if ok else 'fail':>5} {worst_pdb or '-':>16} " + f"{hits}/{tot:<5} {rank0:>6} {med:>7.1f}") + + # Paired against the shipped configuration, misses included as worst. + cells: Dict[tuple, Dict[str, float]] = defaultdict(dict) + for r in rows: + cells[(r.get("pdb"), r.get("trial"))][r.get(key, "")] = _cmp_rank(r) + if base not in arms: + print(f"\n# no {key}={base!r} rows; skipping the paired report") + return + print(f"\n# paired against {key}={base!r} over every (structure, trial) cell; " + f"+ means worse") + for arm in arms: + if arm == base: + continue + d = [by[arm] - by[base] for by in cells.values() + if arm in by and base in by] + if not d: + print(f" {arm:<26} no comparable cells") + continue + print(f" {arm:<26} n={len(d):<4} better={sum(x < 0 for x in d):<4} " + f"same={sum(x == 0 for x in d):<4} worse={sum(x > 0 for x in d):<4} " + f"median={statistics.median(d):+8.1f}") + dup = f"{base}_dup" + if dup in arms: + d = [by[dup] - by[base] for by in cells.values() + if dup in by and base in by] + moved = sum(1 for x in d if x != 0) + print(f"\n# control: {dup} repeats {base} verbatim. {moved}/{len(d)} cells " + f"differ -- that is the engine's own spread, and no effect smaller " + f"than it is resolvable here.") + + def main() -> int: ap = argparse.ArgumentParser(description=__doc__) ap.add_argument("patterns", nargs="+", help="CSV glob(s)") @@ -117,8 +216,19 @@ def main() -> int: ap.add_argument("--compare", default=None, help="column whose values are the arms to pair on, " "e.g. obs_mode or lmax_cap") + ap.add_argument("--gate", action="store_true", + help="report each arm against the shipping criterion " + "(truth in the top N on most trials, every structure)") + ap.add_argument("--top-n", type=int, default=20, + help="how many candidates the downstream pipeline carries") + ap.add_argument("--min-hits", type=int, default=9, + help="trials per structure that must land in the top N") args = ap.parse_args() rows = load(args.patterns) + if args.gate: + gate(rows, key=args.compare or "arm", base=args.base or "production", + top_n=args.top_n, min_hits=args.min_hits) + return 0 summarise(rows) if args.compare: compare(rows, args.compare, args.base) From 27532ad4f6eb40420b6e206540421b9493aa0d36 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 13:41:09 +0200 Subject: [PATCH 020/250] Fit the overall anisotropy in intensity space, not log space `E[I(h) / _shell] = c exp(-2 pi^2 s.U.s)` holds exactly in intensity space. The fit regressed it in log space instead, unweighted and with no constant term, so the `E[ln(I/)] = -gamma` offset -- and `-gamma - ln 2` for centric reflections -- had nowhere to go but the quadratic form. Because centric reflections lie on the zones perpendicular to the symmetry axes, the resulting bias was direction-dependent rather than a harmless overall scale. Vanishing amplitudes were also clamped to 1e-30 and logged, turning each into a residual of about -69 in an unweighted least squares. Measured raw B spreads from the old fit, over the ten benchmark structures: 70 to 5461 A^2. `symmetrize_anisotropy` annihilated it for cubic lattices, where the invariant subspace is one-dimensional, and left 83 to 370 A^2 standing everywhere else. The replacement fits by Gauss-Newton with a free constant, weights from `Var(I/)` (1 acentric, 2 centric), and drops non-positive amplitudes rather than clamping them. Its symmetrised spreads on the same panel are 0.0 to 73.7 A^2, and its residual spread on isotropic synthetic data falls as 1/sqrt(n) -- so what remains is estimation noise, about 7 A^2 at 40k reflections, not bias. That noise floor is why the correction is indistinguishable from no correction on most of the panel: the anisotropy it finds there is smaller than what it can resolve. 3GR5 (P 6_5 2 2) is the exception, at 73.7 A^2 uniaxial along c. Effect on the rotation search, ten structures x ten seeded orientations: the old fit puts truth outside the top 20 on 3GR5 in 10 of 10 trials at every bandwidth tried; the corrected fit is inside on 10 of 10. The signature gains `centric`, which both call sites already had in hand. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/diagnostics/frf_config_sweep.py | 6 + alignment_lab/lab/results.py | 28 +- tests/unit/alignment/test_anisotropy_fit.py | 146 +++++++ tests/unit/alignment/test_rotation_search.py | 207 ++++++++++ torchref/experimental/alignment/__init__.py | 3 + torchref/experimental/alignment/align.py | 17 +- .../experimental/alignment/rotation_search.py | 373 ++++++++++++++++++ torchref/experimental/alignment/sh.py | 135 ++++--- 8 files changed, 850 insertions(+), 65 deletions(-) create mode 100644 tests/unit/alignment/test_anisotropy_fit.py create mode 100644 tests/unit/alignment/test_rotation_search.py create mode 100644 torchref/experimental/alignment/rotation_search.py diff --git a/alignment_lab/diagnostics/frf_config_sweep.py b/alignment_lab/diagnostics/frf_config_sweep.py index 05a7c490..226a14f4 100644 --- a/alignment_lab/diagnostics/frf_config_sweep.py +++ b/alignment_lab/diagnostics/frf_config_sweep.py @@ -161,9 +161,15 @@ def run_one(pdb: str, trial: int, arm: Arm, base: FRFConfig, "orbit_side": "left", "orbit_frame": "cart", "thr_deg": thr_deg, "seconds": round(seconds, 1), }) + # Emit both tensor reports for every arm, blank where the arm does not + # produce one: a row carrying columns the file's header lacks is a schema + # error, and silently-widened rows lose exactly these values. for tag in ("raw", "fixed"): if tag in captured: row.update(tensor_report(captured[tag], tag)) + else: + row.update({f"{tag}_B_min": "", f"{tag}_B_max": "", + f"{tag}_B_spread": ""}) return row diff --git a/alignment_lab/lab/results.py b/alignment_lab/lab/results.py index 7fed6c04..c4de8bd5 100644 --- a/alignment_lab/lab/results.py +++ b/alignment_lab/lab/results.py @@ -69,8 +69,12 @@ def provenance() -> Dict[str, str]: def append_row(csv_path: str | os.PathLike, row: Mapping[str, Any]) -> None: """Append one row, writing the header only when the file is new. - Uses ``fh.tell() == 0`` rather than an existence check so a file created - but not yet written still gets its header. + The header is taken from the file once it exists, not from each row. Taking + it per row silently loses data: a later row carrying a column the first row + lacked gets written wider than the header, and ``DictReader`` then drops the + surplus values into ``restkey``. That is how a sweep's anisotropy columns + went missing while every other column still read back correctly. A row with + an unknown column now raises; a row missing a known one gets a blank. Parameters ---------- @@ -78,11 +82,29 @@ def append_row(csv_path: str | os.PathLike, row: Mapping[str, Any]) -> None: Destination CSV. Parent directories are created. row : mapping Column name -> value. + + Raises + ------ + ValueError + If ``row`` carries a column the file's header does not have. """ path = Path(csv_path) path.parent.mkdir(parents=True, exist_ok=True) + header: Optional[list] = None + if path.exists() and path.stat().st_size > 0: + with open(path, newline="") as fh: + header = next(csv.reader(fh), None) + if header: + unknown = [k for k in row if k not in header] + if unknown: + raise ValueError( + f"{path.name}: row has columns absent from the header: " + f"{unknown}. Emit a stable set of columns for every row, using " + f"blanks where a value does not apply." + ) with open(path, "a", newline="") as fh: - writer = csv.DictWriter(fh, fieldnames=list(row.keys())) + writer = csv.DictWriter(fh, fieldnames=header or list(row.keys()), + restval="") if fh.tell() == 0: writer.writeheader() writer.writerow(dict(row)) diff --git a/tests/unit/alignment/test_anisotropy_fit.py b/tests/unit/alignment/test_anisotropy_fit.py new file mode 100644 index 00000000..26848e0d --- /dev/null +++ b/tests/unit/alignment/test_anisotropy_fit.py @@ -0,0 +1,146 @@ +"""The overall-anisotropy fit, on data whose anisotropy is known by construction. + +Synthetic Wilson intensities: acentric reflections are exponentially +distributed about their shell mean, centric ones follow a chi-squared with one +degree of freedom, and the anisotropy enters as ``exp(-2 pi^2 s.U.s)`` on the +mean. Feeding that in and asking for U back is the only way to separate a +correct fit from one that merely returns something plausible. + +The first test is the one that matters: **isotropic data must give back zero +anisotropy**. Fitting the same relation in log space instead is biased -- +``E[ln(I/)]`` is ``-gamma``, not zero -- and with no constant term that +offset can only be absorbed by the quadratic form. It comes back as tens of +square Angstrom of anisotropy that is not in the data. +""" + +import math + +import pytest +import torch + +from torchref.experimental.alignment.sh import ( + assign_shells, + equal_count_shell_edges, + fit_overall_anisotropy, +) + +pytestmark = pytest.mark.unit + +B_PER_U = 8.0 * math.pi ** 2 + + +def _synthetic(U_true, n=40000, seed=0, centric_fraction=0.0): + """Wilson intensities carrying exactly ``U_true``, on a 4-15 A shell.""" + g = torch.Generator().manual_seed(seed) + smag = 1.0 / (4.0 + 11.0 * torch.rand(n, generator=g, dtype=torch.float64)) + ct = 2 * torch.rand(n, generator=g, dtype=torch.float64) - 1 + phi = 2 * math.pi * torch.rand(n, generator=g, dtype=torch.float64) + st = (1 - ct * ct).clamp(min=0).sqrt() + s = torch.stack([smag * st * torch.cos(phi), + smag * st * torch.sin(phi), + smag * ct], dim=-1) + + # Shell mean falls off with resolution, times the anisotropic term. + sigma_shell = torch.exp(-20.0 * smag * smag) + aniso = torch.exp(-2.0 * (math.pi ** 2) * torch.einsum( + "ni,ij,nj->n", s, U_true.to(torch.float64), s)) + mean_I = sigma_shell * aniso + + centric = torch.rand(n, generator=g, dtype=torch.float64) < centric_fraction + # Acentric: I = mean * Exp(1). Centric: I = mean * chi^2_1. + e = -torch.log(torch.rand(n, generator=g, dtype=torch.float64).clamp(min=1e-300)) + z = torch.randn(n, generator=g, dtype=torch.float64) ** 2 + I = mean_I * torch.where(centric, z, e) + F = I.clamp(min=0).sqrt() + + edges, _ = equal_count_shell_edges(smag, 20) + return F, s, assign_shells(smag, edges), centric + + +def _spread_B(U): + ev = torch.linalg.eigvalsh(U.to(torch.float64)) * B_PER_U + return float(ev[2] - ev[0]) + + +#: Reflections used by the isotropic-data tests. The estimator's own scatter on +#: the B spread is about 7 A^2 at this count, and falls as 1/sqrt(n). +_N_ISO = 40000 + +#: Threshold for "no anisotropy detected" at ``_N_ISO``: above the estimator's +#: measured scatter (median 6.9, max 8.2 over six seeds) with margin. This is a +#: noise floor, not a bias tolerance -- see the scaling test below. +_ISO_TOLERANCE_B = 12.0 + + +@pytest.mark.parametrize("centric_fraction", [0.0, 0.15]) +def test_isotropic_data_gives_no_anisotropy(centric_fraction): + """Zero anisotropy in, nothing but estimation noise out.""" + F, s, idx, cen = _synthetic( + torch.zeros(3, 3), n=_N_ISO, seed=1, centric_fraction=centric_fraction) + U = fit_overall_anisotropy(F, s, idx, cen, P=20) + assert _spread_B(U) < _ISO_TOLERANCE_B, ( + f"isotropic data produced {_spread_B(U):.1f} A^2 of B anisotropy, " + f"beyond this estimator's scatter at n={_N_ISO}" + ) + + +def test_the_isotropic_residual_is_noise_not_bias(): + """The spurious spread must shrink as 1/sqrt(n), not plateau. + + This is the real test of the fit's centring. A biased estimator -- for + instance one regressing ``ln(I/)`` with no constant term, where + ``E[ln(I/)] = -gamma`` has to be absorbed by the quadratic form -- gives + a spread that stays put as reflections are added. An unbiased one averages + it away. + """ + def median_spread(n): + vals = [] + for k in range(3): + F, s, idx, cen = _synthetic(torch.zeros(3, 3), n=n, seed=100 + k) + vals.append(_spread_B(fit_overall_anisotropy(F, s, idx, cen, P=20))) + return sorted(vals)[1] + + coarse, fine = median_spread(20000), median_spread(320000) + # 16x the reflections should buy about 4x, i.e. well over 2x even allowing + # for the scatter of a 3-seed median. + assert fine < coarse / 2.0, ( + f"B spread went {coarse:.2f} -> {fine:.2f} A^2 for 16x the reflections; " + f"an unbiased fit should shrink roughly 4x, a biased one not at all" + ) + + +def test_a_known_tensor_is_recovered(): + """Uniaxial anisotropy of a realistic size comes back to within a few A^2.""" + B_true = torch.diag(torch.tensor([-15.0, -15.0, 30.0], dtype=torch.float64)) + U_true = B_true / B_PER_U + F, s, idx, cen = _synthetic(U_true, n=60000, seed=2) + U = fit_overall_anisotropy(F, s, idx, cen, P=20) + B_fit = U.to(torch.float64) * B_PER_U + # The isotropic part is a gauge -- the per-shell normalisation removes it -- + # so compare the deviatoric parts. + dev = lambda M: M - torch.eye(3, dtype=torch.float64) * torch.diagonal(M).mean() + err = (dev(B_fit) - dev(B_true)).abs().max().item() + assert err < 6.0, f"recovered B off by {err:.1f} A^2:\n{B_fit}" + + +def test_zero_amplitudes_do_not_dominate(): + """A handful of vanishing amplitudes must not steer the fit. + + The earlier version clamped them to 1e-30 and took a logarithm, turning each + into a residual of about -69 in an unweighted least squares. + """ + F, s, idx, cen = _synthetic(torch.zeros(3, 3), seed=3) + clean = fit_overall_anisotropy(F, s, idx, cen, P=20) + F2 = F.clone() + F2[::500] = 0.0 + spiked = fit_overall_anisotropy(F2, s, idx, cen, P=20) + assert abs(_spread_B(spiked) - _spread_B(clean)) < 3.0, ( + f"zeroing 0.2% of amplitudes moved the fit from " + f"{_spread_B(clean):.2f} to {_spread_B(spiked):.2f} A^2" + ) + + +def test_too_few_reflections_returns_zero(): + F, s, idx, cen = _synthetic(torch.zeros(3, 3), n=60, seed=4) + U = fit_overall_anisotropy(F, s, idx, cen, P=20, min_count=20) + assert torch.equal(U, torch.zeros(3, 3, dtype=U.dtype)) diff --git a/tests/unit/alignment/test_rotation_search.py b/tests/unit/alignment/test_rotation_search.py new file mode 100644 index 00000000..d459c0a9 --- /dev/null +++ b/tests/unit/alignment/test_rotation_search.py @@ -0,0 +1,207 @@ +"""The rotation search's public contract: three inputs, and one convention. + +``rotation_search(model, data, model_error_A)`` is the whole surface. What a +caller most easily gets wrong is not the arguments but the *sense* of the +returned rotation -- whether to apply ``R`` or ``R.T`` to the coordinates. The +round-trip test below settles that operationally rather than by reading a +docstring: it rotates a model by a known matrix, searches, applies the inverse +of a returned solution, and requires the model back where it started, modulo the +crystal's rotational symmetry. + +These run on real data (1DAW, C2) because the search needs a real Patterson; +1DAW is the small fast case. +""" + +import math + +import pytest +import torch + +pytestmark = [pytest.mark.unit, pytest.mark.slow] + +#: The search is scored on whether truth is inside the candidate window the +#: placement search carries, not on being rank 0. +TOP_N = 20 + + +@pytest.fixture(scope="module") +def case(pdb_dir, mtz_dir): + from torchref.io.datasets.reflection_data import ReflectionData + from torchref.model import ModelFT + + pdb, mtz = pdb_dir / "1DAW.pdb", mtz_dir / "1DAW.mtz" + if not (pdb.exists() and mtz.exists()): + pytest.skip("1DAW not available") + model = ModelFT(verbose=0).load_pdb(str(pdb)) + data = ReflectionData(verbose=0).load_mtz(str(mtz)) + return model, data + + +def _rotation(seed: int) -> torch.Tensor: + """Haar-uniform SO(3) via QR with the sign correction.""" + g = torch.Generator().manual_seed(seed) + a = torch.randn(3, 3, generator=g, dtype=torch.float64) + q, r = torch.linalg.qr(a) + return q * torch.sign(torch.diagonal(r)).unsqueeze(0) + + +def _angle_deg(a: torch.Tensor, b: torch.Tensor) -> float: + a = a.detach().to(torch.float64).cpu() + b = b.detach().to(torch.float64).cpu() + tr = torch.diagonal(a @ b.T).sum().item() + return math.degrees(math.acos(max(-1.0, min(1.0, (tr - 1.0) / 2.0)))) + + +def _sym_cartesian(data) -> torch.Tensor: + """Space-group rotations as Cartesian operators, on the CPU. + + ``rotations`` is always CPU float64 while the data's tensors may sit on an + accelerator, and the Miller-index matrices have to be carried into the + Cartesian basis before they can be compared with a rotation of coordinates. + """ + from torchref.experimental.alignment.sh import hkl_symops_to_cartesian + + return hkl_symops_to_cartesian( + data.spacegroup.matrices.to(torch.float64).cpu(), + data.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), + ) + + +def _kabsch(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + """Rotation taking ``a`` onto ``b``, both centred. CPU float64.""" + a = a.detach().to(torch.float64).cpu() + b = b.detach().to(torch.float64).cpu() + x = a - a.mean(0) + y = b - b.mean(0) + u, _, vt = torch.linalg.svd(y.T @ x) + d = torch.sign(torch.linalg.det(vt.T @ u.T)) + return vt.T @ torch.diag(torch.tensor([1.0, 1.0, d], dtype=torch.float64)) @ u.T + + +def _search_rotated(case, seed, model_error_A=0.8, n_peaks=200): + from torchref.experimental.alignment import rotation_search + + model, data = case + R_true = _rotation(seed) + rotated = model.copy().rotate( + R_true.to(model.dtype_float), center=model.xyz().mean(0), + ) + return rotated, data, R_true, rotation_search( + rotated, data, model_error_A, n_peaks=n_peaks, + ) + + +def test_returns_the_documented_shapes(case): + from torchref.experimental.alignment import RotationSolutions + + _, _, _, sol = _search_rotated(case, seed=11, n_peaks=50) + assert isinstance(sol, RotationSolutions) + n = len(sol) + assert n > 0 + assert sol.rotations.shape == (n, 3, 3) + assert sol.rotations.dtype == torch.float64 + for name in ("scores", "z_scores"): + assert getattr(sol, name).shape == (n,) + assert sol.euler_zyz.shape == (n, 3) + assert sol.lmax > 0 and sol.d_min > 0 + assert sol.model_error_A == pytest.approx(0.8) + # Best first, on the standardised scale. + z = sol.z_scores + assert torch.all(z[:-1] >= z[1:] - 1e-9) + + +def test_rotations_are_rotations(case): + _, _, _, sol = _search_rotated(case, seed=11, n_peaks=50) + R = sol.rotations + eye = torch.eye(3, dtype=torch.float64).expand_as(R) + assert torch.allclose(R @ R.transpose(-1, -2), eye, atol=1e-9) + assert torch.allclose(torch.linalg.det(R), + torch.ones(len(sol), dtype=torch.float64), atol=1e-9) + + +def test_euler_and_matrix_agree(case): + """``euler_zyz`` is the same orientation as ``rotations``, not a variant.""" + from torchref.experimental.alignment.frf.rotation_utils import ( + rotation_matrix_from_edmonds_euler, + ) + + _, _, _, sol = _search_rotated(case, seed=11, n_peaks=20) + for i in range(min(5, len(sol))): + a, b, g = sol.euler_zyz[i].tolist() + assert torch.allclose(rotation_matrix_from_edmonds_euler(a, b, g), + sol.rotations[i], atol=1e-12) + + +def test_a_solution_inverts_the_applied_rotation(case): + """The convention, settled by algebra: ``rotations[i]`` is ``S . R_true``. + + ``R_true`` is the rotation applied to the model's coordinates to build the + search model, so a returned solution composed with its inverse must leave a + symmetry operator behind. That fixes the sense of the returned matrix + without appealing to the docstring. If it ever inverts, the placement stage + silently searches the wrong orientation. + """ + _, data, R_true, sol = _search_rotated(case, seed=11) + sym_cart = _sym_cartesian(data) + best = min( + min(_angle_deg(sol.rotations[i] @ R_true.T, S) for S in sym_cart) + for i in range(min(TOP_N, len(sol))) + ) + assert best < 5.0, ( + f"no solution in the top {TOP_N} composes with R_true^-1 to within 5 " + f"degrees of a symmetry operator; closest was {best:.2f} degrees. " + f"Either the search failed on 1DAW or the convention flipped." + ) + + +def test_applying_the_transpose_places_the_model(case): + """The documented usage, at the level of coordinates. + + ``model.rotate(rotations[i].T)`` is what the docstring tells a caller to do. + Doing it to the search model must superpose it back onto the unrotated model + up to a symmetry operation -- checked through the actual rotate() call, so + the test would catch a mismatch between the docstring and the maths. + """ + rotated, data, _, sol = _search_rotated(case, seed=11) + model, _ = case + reference = model.xyz() + centre = rotated.xyz().mean(0) + sym_cart = _sym_cartesian(data) + + best = None + for i in range(min(TOP_N, len(sol))): + placed = rotated.copy().rotate( + sol.rotations[i].T.to(rotated.dtype_float).contiguous(), center=centre, + ) + residual = _kabsch(placed.xyz(), reference) + ang = min(_angle_deg(residual, S) for S in sym_cart) + best = ang if best is None else min(best, ang) + assert best is not None and best < 5.0, ( + f"applying rotations[i].T did not superpose the model onto the " + f"unrotated reference for any of the top {TOP_N}; closest residual was " + f"{best:.2f} degrees from a symmetry operator" + ) + + +def test_model_error_changes_the_result(case): + """``model_error_A`` must reach the engine. + + It sets the sigma_A fall-off, so it decides how much the high-resolution + terms count. The previous entry point accepted a coordinate error and then + overwrote it with an empirical estimate from the atom count, so the caller's + value was silently discarded; this asserts it is not. + """ + _, _, _, tight = _search_rotated(case, seed=11, model_error_A=0.2, n_peaks=20) + _, _, _, loose = _search_rotated(case, seed=11, model_error_A=2.5, n_peaks=20) + assert tight.model_error_A != loose.model_error_A + assert not torch.allclose(tight.z_scores[:5], loose.z_scores[:5], atol=1e-6), ( + "model_error_A did not change the rotation function" + ) + + +def test_uninitialised_model_is_rejected(): + from torchref.experimental.alignment import rotation_search + from torchref.model import ModelFT + + with pytest.raises(RuntimeError, match="no coordinates"): + rotation_search(ModelFT(verbose=0), None, 0.8) diff --git a/torchref/experimental/alignment/__init__.py b/torchref/experimental/alignment/__init__.py index 3263dde3..fa53edfb 100644 --- a/torchref/experimental/alignment/__init__.py +++ b/torchref/experimental/alignment/__init__.py @@ -75,6 +75,7 @@ euler_angular_distance, ) from .align import align_model_to_data +from .rotation_search import RotationSolutions, rotation_search # ============================================================================= # Translation search @@ -165,6 +166,8 @@ "rotation_angular_distance", "euler_angular_distance", "align_model_to_data", + "rotation_search", + "RotationSolutions", # Translation "fft_translation_search", "fft_translation_search_torch", diff --git a/torchref/experimental/alignment/align.py b/torchref/experimental/alignment/align.py index daf03506..08dbdd7e 100644 --- a/torchref/experimental/alignment/align.py +++ b/torchref/experimental/alignment/align.py @@ -262,16 +262,15 @@ def _prepare_frf_inputs( aniso_edges, _ = equal_count_shell_edges(s_mag, n_shells) aniso_idx = assign_shells(s_mag, aniso_edges) U_aniso = fit_overall_anisotropy( - F_obs, s_vec, aniso_idx, P=n_shells, min_count=20, + F_obs, s_vec, aniso_idx, centric, P=n_shells, min_count=20, ) - # Project U onto the spacegroup's point-group-invariant subspace - # (Phaser RefineANO.cc:116-142 via cctbx `site_symmetry.average_u_star`). - # Without this constraint a 6-component unconstrained regression can - # fit physically impossible anisotropy on high-symmetry cells — e.g. - # 3K7M (cubic) fits eigenvalues (0.8, 17, 70) Ų which then blows up - # the per-reflection exp(π²·s·U·s) multiplier and destroys the FRF. - # After projection, cubic → U = λI (1 DOF), tetragonal → diag(λ,λ,μ), - # orthorhombic → diag(λ,μ,ν), etc. + # Project U onto the point-group-invariant subspace (Phaser + # RefineANO.cc:116-142, via cctbx `site_symmetry.average_u_star`). An + # unconstrained six-component fit can return a tensor the lattice forbids, + # and applying that modulates the observations by a direction-dependent + # factor the crystal cannot have. After projection: cubic -> U = lambda I + # (one degree of freedom), tetragonal/trigonal/hexagonal -> diag(l, l, m), + # orthorhombic -> diag(l, m, n). from .sh import hkl_symops_to_cartesian, symmetrize_anisotropy _sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) _sym_mats_cart = hkl_symops_to_cartesian(_sg_mats, rec_basis.to(device)) diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py new file mode 100644 index 00000000..cd31643d --- /dev/null +++ b/torchref/experimental/alignment/rotation_search.py @@ -0,0 +1,373 @@ +"""Fast rotation function: find the orientations of a search model in a crystal. + +One call, three inputs:: + + from torchref.experimental.alignment import rotation_search + + solutions = rotation_search(model, data, model_error_A=0.8) + placed = model.copy().rotate(solutions.rotations[0].T) + +Everything else is derived from the model, the data and that error, following +Phaser's own chain (``runMR_FRF.cc:419-448``): the spherical-harmonic bandwidth +from the model's mean radius and the data's resolution, the sigma_A fall-off +from the coordinate error, the Wilson normalisation and French-Wilson posterior +from the observations and their sigmas. + +The constants below are engine settings, not tuning knobs. They are scored on +whether the true orientation lands inside the candidate window the downstream +placement search carries forward, over the benchmark structures at seeded +orientations -- not on the median rank, which hides the cases that matter. Each +one's provenance is in its own comment. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING, Optional + +import torch + +from .sh import ( + apply_overall_anisotropy, + assign_shells, + equal_count_shell_edges, + fit_overall_anisotropy, +) + +if TYPE_CHECKING: # pragma: no cover - typing only + from ...io.datasets.reflection_data import ReflectionData + from ...model.model_ft import ModelFT + +__all__ = ["RotationSolutions", "rotation_search"] + + +#: Spherical-harmonic bandwidth ceiling. Phaser's own limit (``DEF_CLMN_LMAX``) +#: is 100. Where the model and resolution ask for more, the resolution is +#: coarsened to match the bandwidth instead -- see ``phaser_lmax_resolution``. +LMAX_CAP = 64 + +#: SO(3) sample spacing in degrees for the rotation-function grid. Also sets the +#: peak-suppression radius, as ``max(2 * this, 6)`` degrees. +GRID_SAMPLING_DEG = 3.0 + +#: Edge of the P1 box the model's transform is sampled in, as a multiple of the +#: molecular diameter. The box has to be large enough that the periodic images +#: do not overlap the molecule's own Patterson. +DENSE_CALC_PAD = 2.0 + +#: Equal-count resolution shells for the Wilson normalisation, the shell +#: variance weights, the relative Wilson-B fit and the anisotropy fit. +N_WILSON_SHELLS = 20 + +#: Discard rotation-function samples this far below the mean, in standard +#: deviations, before peak finding. Generous: it exists to bound the candidate +#: set, not to select. +SIGMA_THRESHOLD = -5.0 + +#: Babinet bulk-solvent parameters folded into sigma_A on the model side +#: (``EnsemblePDB.cc:96-100``). Phaser's defaults. +SOLVENT_FSOL = 0.95 +SOLVENT_BSOL = 300.0 + +#: Low-resolution cutoff in Angstrom. Effectively none: the rotation function +#: wants the low-resolution terms, which carry the molecular envelope. +LOW_RESOLUTION_CUTOFF_A = 100.0 + +#: Resolution window ``(d_max, d_min)`` the overall anisotropy is fitted in. +#: The tensor is then applied across the full range. Inherited from the range +#: every measurement on this engine was made with, and not itself measured -- +#: unlike the constants above, this pair has no evidence behind it beyond being +#: the one in use. +ANISO_FIT_WINDOW_A = (15.0, 4.0) + + +@dataclass +class RotationSolutions: + """Candidate orientations for a search model, best first. + + Attributes + ---------- + rotations : torch.Tensor + ``(n, 3, 3)`` float64. ``rotations[i]`` maps the search-model frame onto + the crystal frame, so the coordinate rotation that places the model is + its transpose: ``model.copy().rotate(rotations[i].T)``. Each is + determined only up to the crystal's rotational symmetry, so a solution + and its symmetry mates are the same answer. + scores : torch.Tensor + ``(n,)`` rotation-function value at each orientation. + z_scores : torch.Tensor + ``(n,)`` standard deviations above the mean over the whole SO(3) sample + list. This is the scale to judge a solution on; the raw score is not + comparable between runs. + euler_zyz : torch.Tensor + ``(n, 3)`` Edmonds active ZYZ angles in radians, the engine's native + parametrisation: ``R = R_z(alpha) R_y(beta) R_z(gamma)``. + lmax : int + Spherical-harmonic bandwidth used. + d_min : float + High-resolution limit actually used (Angstrom), after the + bandwidth-resolution coupling. + model_error_A : float + The coordinate error the sigma_A fall-off was built from. + """ + + rotations: torch.Tensor + scores: torch.Tensor + z_scores: torch.Tensor + euler_zyz: torch.Tensor + lmax: int + d_min: float + model_error_A: float + + def __len__(self) -> int: + return int(self.rotations.shape[0]) + + +def fit_anisotropy( + data: "ReflectionData", + *, + d_min: float, + d_max: float, + n_shells: int = N_WILSON_SHELLS, + device: Optional[torch.device] = None, +) -> torch.Tensor: + """Fit the overall anisotropy tensor and project it onto the point group. + + Returns ``U`` in Angstrom squared, in the convention + ``F_corrected = F * exp(+pi^2 s.U.s)``. + + The projection matters: an unconstrained six-component fit can return a + tensor the lattice forbids, and applying it then modulates the observations + by a direction-dependent factor the crystal cannot have. Cubic lattices + admit one degree of freedom, tetragonal/trigonal/hexagonal two, + orthorhombic three. + """ + from .sh import hkl_symops_to_cartesian, symmetrize_anisotropy + + device = device or data.hkl.device + rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) + s_vec_all = data.hkl.to(torch.float64) @ rec_basis + s_mag_all = s_vec_all.norm(dim=-1) + keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min) + if int(keep.sum()) < n_shells * 5: + raise ValueError( + f"Only {int(keep.sum())} reflections in [{d_min}, {d_max}] A, too " + f"few for {n_shells} shells." + ) + F_obs = data.F.to(torch.float64).abs()[keep].to(device) + s_vec = s_vec_all[keep].to(device) + s_mag = s_mag_all[keep].to(device) + centric = ( + data.centric[keep].to(torch.bool).to(device) + if hasattr(data, "centric") + else torch.zeros_like(F_obs, dtype=torch.bool) + ) + + edges, _ = equal_count_shell_edges(s_mag, n_shells) + shell_idx = assign_shells(s_mag, edges) + U = fit_overall_anisotropy( + F_obs, s_vec, shell_idx, centric, P=n_shells, min_count=20, + ) + sym_cart = hkl_symops_to_cartesian( + data.spacegroup.matrices.to(torch.float64).to(device), + rec_basis.to(device), + ) + return symmetrize_anisotropy(U, sym_cart) + + +def _search( + model: "ModelFT", + data: "ReflectionData", + model_error_A: float, + *, + U_aniso: torch.Tensor, + n_peaks: int, + verbose: int = 0, +) -> RotationSolutions: + """Run the rotation function with the anisotropy tensor supplied. + + Split out so callers that already fitted the anisotropy for another stage do + not fit it twice; :func:`rotation_search` is the entry point. + """ + from .frf.api import phaser_lmax_resolution, phaser_rotation_search + from .frf.dense_calc import dense_calc_via_box + from .frf.preprocessing import fit_relative_wilson_b + from .frf.rotation_utils import rotation_matrix_from_edmonds_euler_batch + + device = model.xyz().device + with torch.no_grad(): + rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) + hkl_all = data.hkl.to(device) + s_vec_all = hkl_all.to(torch.float64) @ rec_basis + s_mag_all = s_vec_all.norm(dim=-1) + + # Take the observations at the full data resolution: the bandwidth + # coupling below coarsens the limit to whatever the harmonics can + # represent, so pre-restricting here would only lose the terms it keeps. + d_min_data = float(1.0 / s_mag_all.max().item()) + d_max = float(LOW_RESOLUTION_CUTOFF_A) + keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min_data) + + s_obs = s_vec_all[keep] + F_obs = apply_overall_anisotropy( + data.F.to(torch.float64).abs().to(device)[keep], s_obs, U_aniso, + ) + sigF = ( + data.F_sigma.to(torch.float64).to(device)[keep] + if getattr(data, "F_sigma", None) is not None + else None + ) + centric = ( + data.centric[keep].to(torch.bool).to(device) + if hasattr(data, "centric") + else torch.zeros_like(F_obs, dtype=torch.bool) + ) + + # Expand the observations over the space group's rotations to fill + # reciprocal space. |F(hS)| = |F(h)|, and the harmonics need the full + # sphere: sampling only the asymmetric unit under-determines the + # invariant subspace, which is what breaks the high-symmetry cases. + # + # h' = h.S, i.e. the transpose contraction. It agrees with S.h only for + # orthogonal symmetry matrices, so using S.h works everywhere except + # trigonal and hexagonal, where it mixes non-equivalent reflections into + # one orbit. + sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) + n_ops = int(sg_mats.shape[0]) + hkl_keep = hkl_all.to(torch.float64)[keep] + hkl_obs = torch.einsum("kji,nj->kni", sg_mats, hkl_keep).reshape(-1, 3) + s_obs = hkl_obs @ rec_basis + F_obs = F_obs.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() + centric = centric.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() + if sigF is not None: + sigF = sigF.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() + + # The model's transform on a dense P1 grid rather than at the crystal's + # own reflections: the crystal lattice is too sparse to determine the + # high-l harmonics for a large molecule. + model_radius_A = float( + (model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item() + ) + L, d_min = phaser_lmax_resolution(model_radius_A, d_min_data, LMAX_CAP) + s_calc, F_calc = dense_calc_via_box( + model, d_max, d_min, pad=DENSE_CALC_PAD, verbose=verbose > 0, + ) + s_calc = s_calc.to(device) + F_calc = F_calc.to(device) + + # Put the model's amplitudes on the observations' overall B scale + # (EnsemblePDB.cc:793-851), so the radial fall-off does not by itself + # discriminate between orientations. + s_calc_mag = s_calc.norm(dim=-1) + B_rel = fit_relative_wilson_b( + F_obs.to(torch.float64), F_calc.to(torch.float64), + s_obs.norm(dim=-1).to(torch.float64), n_shells=N_WILSON_SHELLS, + s_mag_calc=s_calc_mag.to(torch.float64), + ) + if abs(B_rel) > 1e-6: + F_calc = F_calc * torch.exp(-B_rel * (s_calc_mag * s_calc_mag) / 4.0) + if verbose > 0: + print(f" relative Wilson B = {B_rel:+.2f} A^2", flush=True) + + arf, peaks = phaser_rotation_search( + s_obs, F_obs, centric, + s_calc, F_calc, + sg_mats, + d_min=d_min_data, d_max=d_max, n_peaks=n_peaks, + delta_vrms_A=float(model_error_A), + sigma_threshold=SIGMA_THRESHOLD, + use_lerf1_intensity=True, + use_m_symmetry_filter=True, + sig_F_obs=sigF, + use_french_wilson=sigF is not None, + use_shell_variance_weights=True, + hkl_obs=hkl_obs, + grid_sampling_deg=GRID_SAMPLING_DEG, + model_radius_A=model_radius_A, + auto_lmax=True, + lmax_cap=LMAX_CAP, + apply_bulk_solvent=True, + solvent_fsol=SOLVENT_FSOL, + solvent_bsol=SOLVENT_BSOL, + # The spherical-harmonic contraction dominates the runtime and is + # rate-limited in double precision on accelerators. Its float32 path + # keeps the Bessel recurrence and the cross-chunk accumulator at + # full precision. + compute_dtype=torch.complex64 if device.type == "cuda" else None, + ) + + euler = torch.tensor( + [[p.alpha, p.beta, p.gamma] for p in peaks], dtype=torch.float64, + ).reshape(-1, 3) + rotations = ( + rotation_matrix_from_edmonds_euler_batch( + euler[:, 0], euler[:, 1], euler[:, 2], + ) + if euler.numel() + else torch.zeros((0, 3, 3), dtype=torch.float64) + ) + return RotationSolutions( + rotations=rotations, + scores=torch.tensor([p.score for p in peaks], dtype=torch.float64), + z_scores=torch.tensor([p.sigma for p in peaks], dtype=torch.float64), + euler_zyz=euler, + lmax=int(L - 1), + d_min=float(d_min), + model_error_A=float(model_error_A), + ) + + +def rotation_search( + model: "ModelFT", + data: "ReflectionData", + model_error_A: float, + *, + n_peaks: int = 500, + verbose: int = 0, +) -> RotationSolutions: + """Find the orientations of ``model`` consistent with ``data``. + + Parameters + ---------- + model : ModelFT + Search model. Its orientation in the file is the frame the returned + rotations are relative to; its position is irrelevant, since the + rotation function works on the Patterson. + data : ReflectionData + Observed amplitudes. ``F_sigma`` is used for the French-Wilson posterior + when present. + model_error_A : float + Expected r.m.s. coordinate error of the model against the target, in + Angstrom. This sets the sigma_A fall-off, and so how much weight the + high-resolution terms carry. Use + :func:`~torchref.experimental.alignment.frf.preprocessing.oeffner_vrms` + to estimate it from the model's length and sequence identity if it is + not otherwise known. + n_peaks : int, optional + How many candidate orientations to return, best first. Default 500. This + bounds the answer, not the search. + verbose : int, optional + Progress reporting. Default 0, silent. + + Returns + ------- + RotationSolutions + Ranked orientations. See that class for the rotation convention. + + Raises + ------ + RuntimeError + If ``model`` has no coordinates loaded. + ValueError + If the data carry too few reflections to bin. + """ + if not model.initialized: + raise RuntimeError("model has no coordinates; load a PDB first.") + d_max_fit, d_min_fit = ANISO_FIT_WINDOW_A + U_aniso = fit_anisotropy( + data, d_min=d_min_fit, d_max=d_max_fit, device=model.xyz().device, + ) + return _search( + model, data, model_error_A, + U_aniso=U_aniso, n_peaks=n_peaks, verbose=verbose, + ) diff --git a/torchref/experimental/alignment/sh.py b/torchref/experimental/alignment/sh.py index 6363f4fa..0434863c 100644 --- a/torchref/experimental/alignment/sh.py +++ b/torchref/experimental/alignment/sh.py @@ -446,81 +446,110 @@ def fit_overall_anisotropy( F_obs: torch.Tensor, s_vectors: torch.Tensor, shell_idx: torch.Tensor, + centric: torch.Tensor, P: int, min_count: int = 20, + n_iter: int = 12, ) -> torch.Tensor: """ Fit the overall anisotropy tensor U from F_obs alone (no model needed). - The Popov-Bourenkov anisotropy correction models the observed structure- - factor amplitudes as a per-shell isotropic Wilson piece modulated by an - overall anisotropic Debye-Waller term: + The Popov-Bourenkov correction models the observed intensities as a + per-shell isotropic Wilson piece modulated by an overall anisotropic + Debye-Waller term:: - |F_obs(h)|² ≈ <|F_iso|²>(s) · exp(−2π²·s·U·s) + E[ I(h) / _shell ] = c * exp(-2 pi^2 s.U.s) - Taking logs and fitting the linear regression - ln |F_obs|² − ln<|F_iso|²>(s) = −2π² s·U·s - over all reflections gives the 6-parameter U directly (linear in U). We - parametrise as U_xx, U_yy, U_zz, U_xy, U_xz, U_yz and ignore reflections - in shells with fewer than `min_count` entries (poor shell mean estimate). + That expectation is exact in **intensity** space, which is where this fits + it: a free constant ``c`` absorbs the overall scale, the weights come from + ``Var(I/)`` -- 1 for acentric reflections and 2 for centric ones -- and + non-positive or non-finite amplitudes are dropped. Gauss-Newton from + ``U = 0``. - The returned U is the correction to *apply* in the form - F_obs_corrected(h) = F_obs(h) · exp(+π²·s·U·s) - so that the resulting amplitudes have the same per-shell mean square - regardless of direction. + Fitting the same relation in log space instead is what the earlier version + did, and it is biased: ``E[ln(I/)]`` is ``-gamma = -0.577`` for acentric + and ``-gamma - ln 2`` for centric reflections, not zero. Without a constant + term that offset can only be absorbed by the quadratic form, so U comes back + with a large spurious component -- and because centric reflections lie on + the zones perpendicular to the symmetry axes, the bias is + direction-dependent rather than a harmless overall scale. + + The returned U is the correction to *apply* in the form:: + + F_obs_corrected(h) = F_obs(h) * exp(+pi^2 s.U.s) + + so the corrected amplitudes have the same mean square in every direction. + Project it onto the point group with :func:`symmetrize_anisotropy` before + applying it: an unconstrained six-component fit can return a tensor the + lattice forbids. Parameters ---------- F_obs : (N,) real - s_vectors : (N, 3) real (1/Å) - shell_idx : (N,) int64 — assigns each reflection to a shell in [0, P) + s_vectors : (N, 3) real, reciprocal-space Cartesian (1/Angstrom) + shell_idx : (N,) int64 -- shell of each reflection, in [0, P); negative + entries are excluded + centric : (N,) bool P : int, number of shells + min_count : int, optional + Shells with fewer reflections than this are dropped, since their mean + intensity is too noisy to normalise against. + n_iter : int, optional + Gauss-Newton iterations. Returns ------- - U : (3, 3) symmetric real tensor (Ų) + U : (3, 3) symmetric real tensor (Angstrom squared). Zero if too few + reflections survive to constrain seven parameters. """ - device = F_obs.device - dtype = F_obs.dtype valid = shell_idx >= 0 - F = F_obs[valid] - s = s_vectors[valid].to(dtype) + F = F_obs[valid].to(torch.float64) + s = s_vectors[valid].to(torch.float64) idx = shell_idx[valid] - count = torch.zeros(P, dtype=torch.int64, device=device) + cen = centric[valid].bool() + + ok = torch.isfinite(F) & (F > 0) + F, s, idx, cen = F[ok], s[ok], idx[ok], cen[ok] + + I = F * F + count = torch.zeros(P, dtype=torch.int64, device=F.device) + total = torch.zeros(P, dtype=torch.float64, device=F.device) count.index_add_(0, idx, torch.ones_like(idx)) - F2 = F * F - sum_F2 = torch.zeros(P, dtype=dtype, device=device) - sum_F2.index_add_(0, idx, F2) - mean_F2 = sum_F2 / count.clamp(min=1).to(dtype) - # Mask shells with too few reflections - good = count >= min_count - if good.sum() == 0: - return torch.zeros((3, 3), dtype=dtype, device=device) - keep = good[idx] - F2k = F2[keep].clamp(min=1e-30) - sk = s[keep] - mean_F2_k = mean_F2[idx[keep]].clamp(min=1e-30) - # y = ln|F|² - ln<|F|²> = -2π² · sUs - y = (torch.log(F2k) - torch.log(mean_F2_k)).to(torch.float64) - sk = sk.to(torch.float64) - # Design matrix X for u = (Uxx, Uyy, Uzz, Uxy, Uxz, Uyz): - # s·U·s = Uxx sx² + Uyy sy² + Uzz sz² + 2 Uxy sx sy + 2 Uxz sx sz + 2 Uyz sy sz - X = torch.stack([ - sk[:, 0] ** 2, sk[:, 1] ** 2, sk[:, 2] ** 2, - 2.0 * sk[:, 0] * sk[:, 1], - 2.0 * sk[:, 0] * sk[:, 2], - 2.0 * sk[:, 1] * sk[:, 2], - ], dim=-1) - A = -2.0 * (torch.pi ** 2) * X # y ≈ A · u - # Least-squares solve A u = y - u_vec, _, _, _ = torch.linalg.lstsq(A, y.unsqueeze(-1)) - u_vec = u_vec.squeeze(-1) - Uxx, Uyy, Uzz, Uxy, Uxz, Uyz = u_vec.tolist() - U = torch.tensor( - [[Uxx, Uxy, Uxz], [Uxy, Uyy, Uyz], [Uxz, Uyz, Uzz]], - dtype=dtype, device=device, + total.index_add_(0, idx, I) + mean_I = (total / count.clamp(min=1).to(torch.float64)).clamp(min=1e-30) + + keep = (count >= min_count)[idx] + if int(keep.sum()) < 50: + return torch.zeros((3, 3), dtype=F_obs.dtype, device=F_obs.device) + ratio = I[keep] / mean_I[idx[keep]] + sk, cenk = s[keep], cen[keep] + + x, y, z = sk[:, 0], sk[:, 1], sk[:, 2] + # s.U.s = Uxx sx^2 + Uyy sy^2 + Uzz sz^2 + # + 2 Uxy sx sy + 2 Uxz sx sz + 2 Uyz sy sz + quad = torch.stack([x * x, y * y, z * z, + 2 * x * y, 2 * x * z, 2 * y * z], dim=1) + # Column 0 is the free constant ln(c); the rest carry -2 pi^2 s.U.s. + A = torch.cat([torch.ones_like(x).unsqueeze(1), + -2.0 * (torch.pi ** 2) * quad], dim=1) + w = torch.where(cenk, torch.full_like(ratio, 0.5), torch.ones_like(ratio)) + + theta = torch.zeros(7, dtype=torch.float64, device=F.device) + for _ in range(n_iter): + model = torch.exp((A @ theta).clamp(min=-20.0, max=20.0)) + J = model.unsqueeze(1) * A + Jw = J * w.unsqueeze(1) + H = J.transpose(0, 1) @ Jw + grad = Jw.transpose(0, 1) @ (ratio - model) + H = H + torch.eye(7, dtype=H.dtype, device=H.device) * 1e-12 * float( + torch.diagonal(H).abs().max().clamp(min=1e-30)) + theta = theta + torch.linalg.solve(H, grad) + + u = theta[1:] + return torch.tensor( + [[u[0], u[3], u[4]], [u[3], u[1], u[5]], [u[4], u[5], u[2]]], + dtype=F_obs.dtype, device=F_obs.device, ) - return U def hkl_symops_to_cartesian( From d398660f15808354f909af0cf59456d78daa8c40 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 13:54:36 +0200 Subject: [PATCH 021/250] Collapse the rotation search to (search model, reflection data, model error) `rotation_search(model, data, model_error_A)` replaces a surface of 41, 32, 26 and 21 keyword arguments spread over `align_model_to_data`, `phaser_rotation_search`, `FastRotationFunction` and the engine wrapper. Nine of those were provably dead, fourteen had a non-default branch no code in the repo ever took, and five had a *default* production never used -- `align.py` flipped them on the way past, so no single file stated the shipped behaviour. Everything else is derived from the three inputs, following Phaser's own chain (runMR_FRF.cc:419-448): the bandwidth from the model's mean radius and the data's resolution, the sigma_A fall-off from the coordinate error, the Wilson normalisation and French-Wilson posterior from the observations and their sigmas. What remains fixed is six module constants, each with its provenance in a comment beside it. `model_error_A` is now honoured. The old entry point accepted a coordinate error and then overwrote it with the Oeffner estimate from the atom count whenever `vrms_strategy` was left at its default, so the caller's value was silently discarded. `MolecularReplacementPipeline` and `align_model_to_data` take the same argument and fall back to that estimate explicitly when it is None. `_run_frf_separate_rotation` is deleted; `rotation_search.search_peaks` is the implementation, and the pipeline calls it. Verified equivalent: single-threaded and given the Oeffner estimate, the new path reproduces the old peak list exactly on 1DAW, and on 3GR5 differs only through the engine's own run-to-run spread (which is not zero even single-threaded -- three fresh processes of the old path give two distinct peak lists). The lab keeps its ability to vary the constants, by rebinding them for the duration of one call rather than by passing arguments the API no longer has -- so the measurements that chose the values stay reproducible while the API stays switch-free. `FRFConfig.extra` now raises rather than silently ignoring knobs that no longer exist, and `lab/aniso.py` carries the superseded log-space fit as an explicit arm so the anisotropy comparison can be re-run. Also removes a `--lmax-cap` flag from the pose-recovery driver that no longer reached the pipeline, and stops that driver recording a bandwidth the engine never saw. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/diagnostics/frf_config_sweep.py | 97 +++---- .../diagnostics/frf_normaliser_anatomy.py | 16 +- alignment_lab/diagnostics/frf_prep_compare.py | 16 +- alignment_lab/diagnostics/pose_recovery.py | 16 +- alignment_lab/lab/__init__.py | 4 +- alignment_lab/lab/aniso.py | 188 ++++++------- alignment_lab/lab/frf.py | 52 +++- torchref/experimental/alignment/align.py | 252 +----------------- torchref/experimental/alignment/pipeline.py | 30 ++- .../experimental/alignment/rotation_search.py | 36 ++- 10 files changed, 242 insertions(+), 465 deletions(-) diff --git a/alignment_lab/diagnostics/frf_config_sweep.py b/alignment_lab/diagnostics/frf_config_sweep.py index 226a14f4..6fa93bdd 100644 --- a/alignment_lab/diagnostics/frf_config_sweep.py +++ b/alignment_lab/diagnostics/frf_config_sweep.py @@ -10,10 +10,12 @@ under-determines the SH modes" argument for 48 predates the dense P1-box calc, so it is not evidence about the current engine. anisotropy - ``production`` is the log-space fit with no intercept; ``fixed_fit`` is the - intensity-space replacement; ``iso_only`` keeps its radial part; ``no_aniso`` - drops the correction. Measured before at seven structures, where - ``fixed_fit`` was indistinguishable from ``no_aniso`` in aggregate. + ``production`` is the shipped intensity-space fit; ``legacy_log`` is the + biased log-space fit it replaced; ``iso_only`` keeps only the radial part of + the shipped fit; ``no_aniso`` drops the correction. On this panel + ``production`` and ``no_aniso`` are indistinguishable except on 3GR5, + because most of the structures carry less anisotropy than the estimator can + resolve. ``_orbit_unroll`` Off, on the strength of a run that predates the reciprocal-space convention fix, so its evidence is void. @@ -47,64 +49,43 @@ sys.path.insert(0, str(Path(__file__).resolve().parents[1])) torch.set_grad_enabled(False) -from lab import (BENCH_PDBS, FRFConfig, aniso_arm, merge_peak_lists, # noqa: E402 - orbit_rank, rotated_case, run_frf, seed_for, tensor_report) +from lab import (BENCH_PDBS, FRFConfig, aniso_arm, orbit_rank, # noqa: E402 + rotated_case, run_frf, seed_for, tensor_report) from lab.results import append_row, provenance # noqa: E402 EXPERIMENT = "frf_config_sweep" -#: Suppression radius for the union merge. The engine uses -#: ``max(2 * grid_sampling_deg, 6)`` internally, and the production sampling is -#: 3 degrees, so 6 degrees keeps the merged list on the same footing. -UNION_NMS_DEG = 6.0 - - @dataclass(frozen=True) class Arm: - """One engine configuration to measure. - - ``radius_scales`` with more than one entry means the union arm: one FRF - evaluation per scale, merged by z-score. - """ + """One engine configuration to measure.""" name: str lmax_cap: int = 64 aniso: str = "production" - orbit_unroll: bool = False - radius_scales: Tuple[float, ...] = (1.0,) def config(self, base: FRFConfig) -> Tuple[FRFConfig, ...]: - out = [] - for scale in self.radius_scales: - extra: Dict[str, object] = {"_orbit_unroll": self.orbit_unroll} - if scale != 1.0: - extra["frf_patterson_radius_scale"] = scale - out.append(FRFConfig( - d_min=base.d_min, d_max=base.d_max, n_shells=base.n_shells, - n_peaks=base.n_peaks, lmax_cap=self.lmax_cap, - dense_pad=base.dense_pad, extra=extra, - )) - return tuple(out) + return (FRFConfig( + d_min=base.d_min, d_max=base.d_max, n_shells=base.n_shells, + n_peaks=base.n_peaks, lmax_cap=self.lmax_cap, + dense_pad=base.dense_pad, + ),) def _factorial_arms() -> Tuple[Arm, ...]: """lmax_cap x anisotropy, plus the repeat-baseline control.""" arms = [Arm("production_dup")] for cap in (48, 64, 100): - for aniso in ("production", "fixed_fit", "iso_only", "no_aniso"): + for aniso in ("production", "legacy_log", "iso_only", "no_aniso"): arms.append(Arm(f"cap{cap}_{aniso}", lmax_cap=cap, aniso=aniso)) return tuple(arms) -def _followup_arms(cap: int, aniso: str) -> Tuple[Arm, ...]: - """One-at-a-time from the winning cell of stage 1.""" - base = Arm(f"cap{cap}_{aniso}", lmax_cap=cap, aniso=aniso) - return ( - base, - Arm(f"{base.name}_unroll", lmax_cap=cap, aniso=aniso, orbit_unroll=True), - Arm(f"{base.name}_union", lmax_cap=cap, aniso=aniso, - radius_scales=(1.0, 0.5)), - ) +#: Stage 2 measured the orbit-dedup unroll and the two-radius Patterson union. +#: Both were arguments of the engine wrapper that the API collapse removed, so +#: those arms cannot be built against this tree; they were measured against the +#: pre-collapse tree (a git worktree pinned at 133bd565) and the outcome is +#: recorded in the changelog. Re-measuring them means restoring the arguments +#: first, which is the point: they are not switches any more. #: The baseline every paired difference is taken against: today's shipped @@ -119,19 +100,13 @@ def run_one(pdb: str, trial: int, arm: Arm, base: FRFConfig, configs = arm.config(base) captured: dict = {} - peak_lists = [] + cfg = configs[0] t0 = time.time() - for cfg in configs: - with aniso_arm(arm.aniso if arm.aniso != "production" else "production", - data, d_min=cfg.d_min, d_max=cfg.d_max, - captured=captured): - res = run_frf(model, data, cfg, capture_arf=False, verbose=0) - peak_lists.append(res.peaks) + with aniso_arm(arm.aniso, data, d_min=cfg.d_min, d_max=cfg.d_max, + captured=captured): + res = run_frf(model, data, cfg, capture_arf=False, verbose=0) seconds = time.time() - t0 - - peaks = (peak_lists[0] if len(peak_lists) == 1 else - merge_peak_lists(peak_lists, n_peaks=base.n_peaks, - nms_radius_deg=UNION_NMS_DEG)) + peaks = res.peaks rank, ang = orbit_rank( peaks, R_true, data.spacegroup.matrices.to(torch.float64).cpu(), @@ -145,9 +120,6 @@ def run_one(pdb: str, trial: int, arm: Arm, base: FRFConfig, row.update({ "arm_lmax_cap": arm.lmax_cap, "arm_aniso": arm.aniso, - "arm_orbit_unroll": int(arm.orbit_unroll), - "arm_radius_scales": "|".join(f"{s:g}" for s in arm.radius_scales), - "n_frf_calls": len(configs), "spacegroup": str(data.spacegroup.hm), "truth_rank": rank, # orbit_rank returns -1 for "no peak within thr_deg". A miss must not @@ -164,7 +136,7 @@ def run_one(pdb: str, trial: int, arm: Arm, base: FRFConfig, # Emit both tensor reports for every arm, blank where the arm does not # produce one: a row carrying columns the file's header lacks is a schema # error, and silently-widened rows lose exactly these values. - for tag in ("raw", "fixed"): + for tag in ("raw", "legacy"): if tag in captured: row.update(tensor_report(captured[tag], tag)) else: @@ -179,10 +151,8 @@ def main() -> int: ap.add_argument("--trial", type=int, default=None, help="single trial index; omit to run --trials of them") ap.add_argument("--trials", type=int, default=10) - ap.add_argument("--stage", type=int, default=1, choices=(1, 2), - help="1 = lmax x aniso factorial; 2 = follow-ups") - ap.add_argument("--stage2-cap", type=int, default=64) - ap.add_argument("--stage2-aniso", default="fixed_fit") + ap.add_argument("--stage", type=int, default=1, choices=(1,), + help="1 = lmax x aniso factorial") ap.add_argument("--d-min", type=float, default=4.0) ap.add_argument("--d-max", type=float, default=15.0) ap.add_argument("--n-peaks", type=int, default=500) @@ -191,10 +161,11 @@ def main() -> int: ap.add_argument("--outdir", default=None) args = ap.parse_args() - arms = (BASELINE,) + ( - _factorial_arms() if args.stage == 1 - else _followup_arms(args.stage2_cap, args.stage2_aniso) - ) + if args.stage != 1: + raise SystemExit( + "stage 2 measured engine arguments that no longer exist; see the " + "note above _factorial_arms.") + arms = (BASELINE,) + _factorial_arms() base = FRFConfig(d_min=args.d_min, d_max=args.d_max, n_peaks=args.n_peaks) if args.out_csv: diff --git a/alignment_lab/diagnostics/frf_normaliser_anatomy.py b/alignment_lab/diagnostics/frf_normaliser_anatomy.py index ef49a843..78427fb1 100644 --- a/alignment_lab/diagnostics/frf_normaliser_anatomy.py +++ b/alignment_lab/diagnostics/frf_normaliser_anatomy.py @@ -59,7 +59,7 @@ sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -from lab import FRFConfig, load_case, patched # noqa: E402 +from lab import FRFConfig, load_case, patched, run_frf # noqa: E402 from lab.results import append_row, provenance # noqa: E402 from diagnostics.frf_ghost_knockout import PHASER_PINNED # noqa: E402 from diagnostics.frf_inject_phaser_obs import _pack # noqa: E402 @@ -75,7 +75,6 @@ def capture_ours(pdb: str): and nothing between them reorders or filters -- so spying on the two calls gives an aligned pair. """ - from torchref.experimental.alignment import align as _align from torchref.experimental.alignment.frf import api as _api pin = PHASER_PINNED[pdb] @@ -101,19 +100,12 @@ def spy_lerf1(eEobs, centric, dfac=None, **kw): def _pinned(model_radius_A, d_min_data, lmax_cap=48): return int(pin["lmax"]) + 1, float(pin["d_min_eff"]) - cfg = FRFConfig(n_peaks=5, lmax_cap=int(pin["lmax"])) + cfg = FRFConfig(n_peaks=5, lmax_cap=int(pin["lmax"]), + grid_sampling_deg=float(pin["sampling_deg"])) with patched(_api, "phaser_lmax_resolution", _pinned), \ patched(_api, "bessel_sh_expand", spy_bessel), \ patched(_api, "build_lerf1_intensity", spy_lerf1): - frf_in = _align._prepare_frf_inputs( - model, data, d_min=cfg.d_min, d_max=cfg.d_max, - n_shells=cfg.n_shells, verbose=0, - ) - _align._run_frf_separate_rotation( - model, data, frf_in, n_peaks=5, verbose=0, - lmax_cap=int(pin["lmax"]), - grid_sampling_deg=float(pin["sampling_deg"]), - ) + run_frf(model, data, cfg, capture_arf=False, verbose=0) for k in ("s", "eEobs"): if k not in cap: raise RuntimeError(f"failed to capture {k} from the engine") diff --git a/alignment_lab/diagnostics/frf_prep_compare.py b/alignment_lab/diagnostics/frf_prep_compare.py index 676967fa..43685f0a 100644 --- a/alignment_lab/diagnostics/frf_prep_compare.py +++ b/alignment_lab/diagnostics/frf_prep_compare.py @@ -42,7 +42,7 @@ sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -from lab import FRFConfig, case_paths, load_case, patched # noqa: E402 +from lab import FRFConfig, case_paths, load_case, patched, run_frf # noqa: E402 from lab.phaser_match import PATCHED_PHASER, write_keywords # noqa: E402 from lab.results import append_row, provenance # noqa: E402 @@ -95,7 +95,6 @@ def capture_ours(pdb: str): Both go through ``bessel_sh_expand``; the observation call is the one with ``zsymm > 1`` (the calc side is deliberately never m-filtered). """ - from torchref.experimental.alignment import align as _align from torchref.experimental.alignment.frf import api as _api pin = PINNED[pdb] @@ -113,18 +112,11 @@ def spy(s, vals, **kw): def _pinned(model_radius_A, d_min_data, lmax_cap=48): return int(pin["lmax"]) + 1, float(pin["d_min_eff"]) - cfg = FRFConfig(n_peaks=20, lmax_cap=int(pin["lmax"])) + cfg = FRFConfig(n_peaks=20, lmax_cap=int(pin["lmax"]), + grid_sampling_deg=float(pin["sampling_deg"])) with patched(_api, "phaser_lmax_resolution", _pinned), \ patched(_api, "bessel_sh_expand", spy): - frf_in = _align._prepare_frf_inputs( - model, data, d_min=cfg.d_min, d_max=cfg.d_max, - n_shells=cfg.n_shells, verbose=0, - ) - _align._run_frf_separate_rotation( - model, data, frf_in, n_peaks=20, verbose=0, - lmax_cap=int(pin["lmax"]), - grid_sampling_deg=float(pin["sampling_deg"]), - ) + run_frf(model, data, cfg, capture_arf=False, verbose=0) return cap diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py index ba2242ac..bf517a7d 100644 --- a/alignment_lab/diagnostics/pose_recovery.py +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -73,12 +73,20 @@ def residual_rotation_deg(aligned_xyz, canonical_xyz, symops) -> float: for k in range(symops.shape[0])) +#: The rotation search's own bandwidth constant. Recorded in every row because +#: `align_model_to_data` has no bandwidth argument, so a `--lmax-cap` flag here +#: would name a value the engine never saw. +import importlib as _importlib # noqa: E402 + +_LMAX_CAP = _importlib.import_module( + "torchref.experimental.alignment.rotation_search").LMAX_CAP + + def main() -> int: ap = argparse.ArgumentParser(description=__doc__) ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) ap.add_argument("--trial", type=int, default=0) ap.add_argument("--arms", default="m_letf1,none,none+subpeak") - ap.add_argument("--lmax-cap", type=int, default=64) ap.add_argument("--n-rotation-candidates", type=int, default=15) ap.add_argument("--n-rotation-peaks", type=int, default=200) ap.add_argument("--success-deg", type=float, default=8.0) @@ -118,11 +126,11 @@ def main() -> int: t0 = time.time() try: aligned = align_model_to_data( - search, data, d_min=4.0, d_max=15.0, L=32, n_shells=20, + search, data, d_min=4.0, d_max=15.0, n_shells=20, n_rotation_peaks=args.n_rotation_peaks, n_ml_refine=200, do_translation=True, do_joint_refine=True, n_rotation_candidates=args.n_rotation_candidates, - frf_lmax_cap=args.lmax_cap, verbose=args.verbose, **flags, + verbose=args.verbose, **flags, ) resid = residual_rotation_deg(aligned.xyz(), canonical_xyz, symops) err = "" @@ -138,7 +146,7 @@ def main() -> int: truth_rank="", truth_angle_deg=(round(resid, 4) if resid == resid else ""), orbit_side="kabsch", orbit_frame="cart", - lmax_cap=args.lmax_cap, d_min=4.0, d_max=15.0, + lmax_cap=_LMAX_CAP, d_min=4.0, d_max=15.0, device="cpu", arm=arm, rescore_engine=flags["rescore_engine"], subpeak_refine=int(flags["subpeak_refine"]), diff --git a/alignment_lab/lab/__init__.py b/alignment_lab/lab/__init__.py index bc6dad2e..72261ad7 100644 --- a/alignment_lab/lab/__init__.py +++ b/alignment_lab/lab/__init__.py @@ -24,7 +24,7 @@ from .aniso import ( ARMS as ANISO_ARMS, aniso_arm, - fit_aniso_intensity_space, + fit_aniso_log_space, tensor_report, ) from .frf import (FRFConfig, FRFResult, merge_peak_lists, patched, @@ -45,7 +45,7 @@ "symmetry_orbit", "ANISO_ARMS", "aniso_arm", - "fit_aniso_intensity_space", + "fit_aniso_log_space", "tensor_report", "FRFConfig", "FRFResult", diff --git a/alignment_lab/lab/aniso.py b/alignment_lab/lab/aniso.py index bd1fb0b1..dacbcfc0 100644 --- a/alignment_lab/lab/aniso.py +++ b/alignment_lab/lab/aniso.py @@ -1,32 +1,30 @@ -"""The overall-anisotropy correction: the production fit, and a corrected one. +"""Arms for the overall-anisotropy correction. -``sh.py:445 fit_overall_anisotropy`` regresses ``ln|F|^2 - ln<|F|^2>_shell`` on -``-2 pi^2 s.U.s`` by unweighted least squares **with no intercept**, and that is -the FRF's remaining defect (job 489540/489548). Three faults, all visible in its +``sh.fit_overall_anisotropy`` now fits in intensity space with a free constant. +The version it replaced regressed ``ln|F|^2 - ln<|F|^2>_shell`` on +``-2 pi^2 s.U.s`` by unweighted least squares **with no intercept**, and that was +the rotation function's last real defect. Three faults, all visible in its output: * ``E[ln(I/)]`` is ``-gamma = -0.577`` for acentric reflections and ``-gamma - ln 2 = -1.270`` for centric ones, not zero. With no intercept the - offset can only be absorbed by the quadratic form, which is why the fitted - tensor's SMALLEST B eigenvalue is 35-64 A^2 on every benchmark structure - instead of near zero. The centric part is worse than a constant: centric - reflections lie on the zones perpendicular to the symmetry axes, so the bias - is direction-dependent. + offset can only be absorbed by the quadratic form. The centric part is worse + than a constant: centric reflections lie on the zones perpendicular to the + symmetry axes, so the bias is direction-dependent. * ``clamp(min=1e-30)`` turns a vanishing amplitude into ``y ~ -69``; a handful of those outweigh thousands of ordinary reflections in an unweighted fit. * ``ln`` of a single-reflection intensity has variance ``pi^2/6`` (acentric) or ``pi^2/2`` (centric) with a heavy left tail, so the fit is dominated by the weak reflections carrying the least information. -Raw fitted B eigenvalue spreads come out at 70 to 5461 A^2. -``symmetrize_anisotropy`` then projects onto the point-group-invariant subspace, -which annihilates the garbage where that subspace is small (cubic -> 1 DOF) and -leaves it where it is not (trigonal/hexagonal -> diag(lambda, lambda, mu), where -a uniaxial tensor along c is symmetry-allowed). +Raw fitted B eigenvalue spreads came out at 70 to 5461 A^2 over the ten +benchmark structures. ``symmetrize_anisotropy`` then annihilated the garbage +where the point-group-invariant subspace is small (cubic -> one degree of +freedom) and left it standing where it is not (trigonal/hexagonal -> two). -One definition of the replacement lives here rather than in a diagnostic, since -several diagnostics need to A/B against it and it is the candidate production -change. +:func:`fit_aniso_log_space` reproduces that version, so the measurement that +justified replacing it can be re-run against the current tree rather than taken +on trust. """ from __future__ import annotations @@ -36,77 +34,68 @@ import torch +from torchref.experimental.alignment.sh import fit_overall_anisotropy + #: U (A^2) -> B (A^2). B_PER_U = 8.0 * math.pi ** 2 -#: Arm names accepted by :func:`aniso_arm`. -ARMS = ("production", "no_aniso", "iso_only", "fixed_fit") +#: Arm names accepted by :func:`aniso_arm`. ``production`` is whatever +#: ``sh.fit_overall_anisotropy`` currently does; ``legacy_log`` is the biased +#: fit it replaced. +ARMS = ("production", "legacy_log", "no_aniso", "iso_only") -def fit_aniso_intensity_space( +def fit_aniso_log_space( F_obs: torch.Tensor, - s_vec: torch.Tensor, + s_vectors: torch.Tensor, shell_idx: torch.Tensor, - centric: torch.Tensor, P: int, *, min_count: int = 20, - n_iter: int = 12, ) -> torch.Tensor: - """Unbiased replacement for ``fit_overall_anisotropy``. - - Fits in INTENSITY space, where ``E[I/_shell] = c * exp(-2 pi^2 s.U.s)`` - holds exactly with no distributional correction. ``Var(I/)`` is 1 for - acentric and 2 for centric reflections, which gives the weights; a free - constant ``c`` absorbs the overall scale so it cannot leak into ``U``; - non-finite and non-positive amplitudes are dropped rather than clamped. - Gauss-Newton from ``U = 0``. - - Returns ``U`` in A^2 in the same convention as ``fit_overall_anisotropy`` - (applied as ``exp(+pi^2 s.U.s)``), so the caller's symmetrisation and - application are unchanged. + """The superseded log-space fit, verbatim, for A/B against the current one. + + Unweighted least squares of ``ln|F|^2 - ln<|F|^2>_shell`` on + ``-2 pi^2 s.U.s`` with no constant term, and vanishing amplitudes clamped + rather than dropped. Returns ``U`` in A^2 in the same convention as + :func:`~torchref.experimental.alignment.sh.fit_overall_anisotropy`. """ + dtype = F_obs.dtype + device = F_obs.device valid = shell_idx >= 0 - F = F_obs[valid].to(torch.float64) - s = s_vec[valid].to(torch.float64) + F = F_obs[valid] + s = s_vectors[valid].to(dtype) idx = shell_idx[valid] - cen = centric[valid].bool() - ok = torch.isfinite(F) & (F > 0) - F, s, idx, cen = F[ok], s[ok], idx[ok], cen[ok] - - I = F * F - cnt = torch.zeros(P, dtype=torch.int64, device=F.device) - tot = torch.zeros(P, dtype=torch.float64, device=F.device) - cnt.index_add_(0, idx, torch.ones_like(idx)) - tot.index_add_(0, idx, I) - mean_I = (tot / cnt.clamp(min=1).to(torch.float64)).clamp(min=1e-30) - keep = (cnt >= min_count)[idx] - if int(keep.sum()) < 50: - return torch.zeros((3, 3), dtype=F_obs.dtype, device=F_obs.device) - r = I[keep] / mean_I[idx[keep]] - sk, cenk = s[keep], cen[keep] - - x, y, z = sk[:, 0], sk[:, 1], sk[:, 2] - quad = torch.stack([x * x, y * y, z * z, - 2 * x * y, 2 * x * z, 2 * y * z], dim=1) - A = torch.cat([torch.ones_like(x).unsqueeze(1), - -2.0 * (torch.pi ** 2) * quad], dim=1) - w = torch.where(cenk, torch.full_like(r, 0.5), torch.ones_like(r)) - - theta = torch.zeros(7, dtype=torch.float64, device=F.device) - for _ in range(n_iter): - m = torch.exp((A @ theta).clamp(min=-20.0, max=20.0)) - J = m.unsqueeze(1) * A - Jw = J * w.unsqueeze(1) - H = J.transpose(0, 1) @ Jw - g = Jw.transpose(0, 1) @ (r - m) - H = H + torch.eye(7, dtype=H.dtype, device=H.device) * 1e-12 * float( - torch.diagonal(H).abs().max().clamp(min=1e-30)) - theta = theta + torch.linalg.solve(H, g) - u = theta[1:] + + count = torch.zeros(P, dtype=torch.int64, device=device) + count.index_add_(0, idx, torch.ones_like(idx)) + F2 = F * F + sum_F2 = torch.zeros(P, dtype=dtype, device=device) + sum_F2.index_add_(0, idx, F2) + mean_F2 = sum_F2 / count.clamp(min=1).to(dtype) + + good = count >= min_count + if int(good.sum()) == 0: + return torch.zeros((3, 3), dtype=dtype, device=device) + keep = good[idx] + F2k = F2[keep].clamp(min=1e-30) + sk = s[keep].to(torch.float64) + mean_F2_k = mean_F2[idx[keep]].clamp(min=1e-30) + + y = (torch.log(F2k) - torch.log(mean_F2_k)).to(torch.float64) + X = torch.stack([ + sk[:, 0] ** 2, sk[:, 1] ** 2, sk[:, 2] ** 2, + 2.0 * sk[:, 0] * sk[:, 1], + 2.0 * sk[:, 0] * sk[:, 2], + 2.0 * sk[:, 1] * sk[:, 2], + ], dim=-1) + A = -2.0 * (torch.pi ** 2) * X + u, _, _, _ = torch.linalg.lstsq(A, y.unsqueeze(-1)) + Uxx, Uyy, Uzz, Uxy, Uxz, Uyz = u.squeeze(-1).tolist() return torch.tensor( - [[u[0], u[3], u[4]], [u[3], u[1], u[5]], [u[4], u[5], u[2]]], - dtype=F_obs.dtype, device=F_obs.device) + [[Uxx, Uxy, Uxz], [Uxy, Uyy, Uyz], [Uxz, Uyz, Uzz]], + dtype=dtype, device=device, + ) def tensor_report(U: torch.Tensor, tag: str) -> dict: @@ -121,15 +110,24 @@ def tensor_report(U: torch.Tensor, tag: str) -> dict: def aniso_arm(arm: str, data, *, d_min: float, d_max: float, captured: dict): """Swap the anisotropy fit for the duration of one FRF call. - ``captured`` receives the tensors actually fitted (``raw``, and ``fixed`` - when the arm uses the replacement) so a caller can report the artefact size - alongside the rank it costs. - - The ``fixed_fit`` arm needs ``centric``, which ``fit_overall_anisotropy`` - is not given. ``_prepare_frf_inputs`` masks ``F_obs`` and ``centric`` with - the same resolution window, so it is recomputed here from the same - ``d_min``/``d_max`` and checked against the amplitude count -- a mismatch - raises rather than silently misaligning. + ``captured`` receives the tensor actually fitted under the key ``raw``, so a + caller can report the artefact size alongside the rank it costs. + + Patches ``align.fit_overall_anisotropy``, which is the symbol + ``_prepare_frf_inputs`` calls and whose result is handed to the rotation + search as ``U_aniso``. + + Parameters + ---------- + arm : str + One of :data:`ARMS`. + data : ReflectionData + Used to recompute the centric mask over the same resolution window the + fit sees. A length mismatch raises rather than misaligning silently. + d_min, d_max : float + The window ``_prepare_frf_inputs`` was called with. + captured : dict + Filled in by the wrapper. """ if arm not in ARMS: raise ValueError(f"unknown aniso arm {arm!r}; expected one of {ARMS}") @@ -139,11 +137,11 @@ def aniso_arm(arm: str, data, *, d_min: float, d_max: float, captured: dict): rec = data.cell.reciprocal_basis_matrix.to(torch.float64) smag = (data.hkl.to(torch.float64) @ rec).norm(dim=-1) keep = (smag >= 1.0 / d_max) & (smag <= 1.0 / d_min) - centric = (data.centric[keep].to(torch.bool) - if hasattr(data, "centric") else None) + centric_window = (data.centric[keep].to(torch.bool) + if hasattr(data, "centric") else None) - def wrapped(F_obs, s_vec, shell_idx, **kw): - U = original(F_obs, s_vec, shell_idx, **kw) + def wrapped(F_obs, s_vec, shell_idx, centric, **kw): + U = original(F_obs, s_vec, shell_idx, centric, **kw) captured.setdefault("raw", U.detach().clone()) if arm == "production": return U @@ -153,20 +151,24 @@ def wrapped(F_obs, s_vec, shell_idx, **kw): # Radial part only; symmetrisation leaves lambda*I unchanged. return torch.eye(3, dtype=U.dtype, device=U.device) * ( torch.diagonal(U).sum() / 3.0) - if centric is None or centric.numel() != F_obs.shape[0]: - n = 0 if centric is None else centric.numel() + if centric_window is None or centric_window.numel() != F_obs.shape[0]: + n = 0 if centric_window is None else centric_window.numel() raise RuntimeError( f"centric mask has {n} entries against {F_obs.shape[0]} " f"amplitudes -- the resolution window assumed here " f"([{d_min}, {d_max}] A) is not the engine's") - Ufix = fit_aniso_intensity_space( - F_obs, s_vec, shell_idx, centric.to(F_obs.device), - P=kw.get("P", 20), min_count=kw.get("min_count", 20)) - captured["fixed"] = Ufix.detach().clone() - return Ufix + U_legacy = fit_aniso_log_space( + F_obs, s_vec, shell_idx, P=kw.get("P", 20), + min_count=kw.get("min_count", 20)) + captured["legacy"] = U_legacy.detach().clone() + return U_legacy setattr(_align, "fit_overall_anisotropy", wrapped) try: yield finally: setattr(_align, "fit_overall_anisotropy", original) + + +__all__ = ["ARMS", "B_PER_U", "aniso_arm", "fit_aniso_log_space", + "fit_overall_anisotropy", "tensor_report"] diff --git a/alignment_lab/lab/frf.py b/alignment_lab/lab/frf.py index 35933869..0135750f 100644 --- a/alignment_lab/lab/frf.py +++ b/alignment_lab/lab/frf.py @@ -29,6 +29,10 @@ class FRFConfig: n_peaks: int = 500 lmax_cap: int = 48 dense_pad: float = 2.0 + grid_sampling_deg: float = 3.0 + #: Expected r.m.s. coordinate error, in Angstrom. ``None`` uses the Oeffner + #: estimate from the model's length, which is what the pipeline does. + model_error_A: Optional[float] = None extra: Dict[str, Any] = field(default_factory=dict) def as_row(self) -> Dict[str, Any]: @@ -168,9 +172,18 @@ def run_frf( """ import time + import importlib + from torchref.experimental.alignment import align as _align from torchref.experimental.alignment.frf import api as _api + # `from ...alignment import rotation_search` gives the FUNCTION, which the + # package re-exports under the module's own name. Patching constants needs + # the module object. + _rs = importlib.import_module( + "torchref.experimental.alignment.rotation_search") + from torchref.experimental.alignment.frf.preprocessing import oeffner_vrms + cfg = cfg or FRFConfig() captured: Dict[str, Any] = {} @@ -179,24 +192,41 @@ def _wrapped(*args, **kwargs): captured["arf"] = arf return arf, peaks + # The engine takes no tuning arguments any more: `lmax_cap`, `dense_pad` and + # the SO(3) sampling are module constants of `rotation_search`. The lab + # sweeps them by rebinding those constants for the duration of one call, so + # the production API stays switch-free while the measurements that chose the + # values remain reproducible. frf_inputs = _align._prepare_frf_inputs( model, data, d_min=cfg.d_min, d_max=cfg.d_max, n_shells=cfg.n_shells, verbose=verbose, ) + model_error_A = cfg.model_error_A + if model_error_A is None: + model_error_A = oeffner_vrms(max(1, int(model.xyz().shape[0] / 8)), 1.0) + if cfg.extra: + raise ValueError( + f"FRFConfig.extra is no longer plumbed anywhere: {sorted(cfg.extra)}. " + f"The engine knobs it reached were deleted; patch the constants in " + f"torchref.experimental.alignment.rotation_search instead." + ) t0 = time.time() - if capture_arf: - _original = _api.phaser_rotation_search - with patched(_api, "phaser_rotation_search", _wrapped): - peaks = _align._run_frf_separate_rotation( - model, data, frf_inputs, n_peaks=cfg.n_peaks, verbose=verbose, - lmax_cap=cfg.lmax_cap, dense_pad=cfg.dense_pad, **cfg.extra, + with patched(_rs, "LMAX_CAP", int(cfg.lmax_cap)), \ + patched(_rs, "DENSE_CALC_PAD", float(cfg.dense_pad)), \ + patched(_rs, "GRID_SAMPLING_DEG", float(cfg.grid_sampling_deg)): + if capture_arf: + _original = _api.phaser_rotation_search + with patched(_api, "phaser_rotation_search", _wrapped): + peaks, _lmax, _dmin = _rs.search_peaks( + model, data, model_error_A, U_aniso=frf_inputs.U_aniso, + n_peaks=cfg.n_peaks, verbose=verbose, + ) + else: + peaks, _lmax, _dmin = _rs.search_peaks( + model, data, model_error_A, U_aniso=frf_inputs.U_aniso, + n_peaks=cfg.n_peaks, verbose=verbose, ) - else: - peaks = _align._run_frf_separate_rotation( - model, data, frf_inputs, n_peaks=cfg.n_peaks, verbose=verbose, - lmax_cap=cfg.lmax_cap, dense_pad=cfg.dense_pad, **cfg.extra, - ) seconds = time.time() - t0 arf = captured.get("arf") diff --git a/torchref/experimental/alignment/align.py b/torchref/experimental/alignment/align.py index 08dbdd7e..9de2e600 100644 --- a/torchref/experimental/alignment/align.py +++ b/torchref/experimental/alignment/align.py @@ -2,8 +2,7 @@ Molecular replacement: data-prep / FRF stage helpers + the public entry point. This module hosts the heavy, reusable stage helpers — Lattman-Love / anisotropy -data prep (`_prepare_frf_inputs`), the Phaser-faithful rotation search -(`_run_frf_separate_rotation`), the solvent-aware R-work +data prep (`_prepare_frf_inputs`), the solvent-aware R-work (`_external_rwork`), the direct-SF translation evaluator (`_DirectModelEvaluator`), the Rodrigues helper (`_rodrigues`) and the stage timer (`_StageTimer`) — that are shared by the rotation-ranking benchmarks and @@ -294,240 +293,6 @@ def _prepare_frf_inputs( ) -def _run_frf_separate_rotation( - model: "ModelFT", - data: "ReflectionData", - frf: "FRFInputs", - *, - lmax_cap: int = 48, - dense_pad: float = 2.0, - n_peaks: int = 500, - grid_sampling_deg: float = 3.0, - delta_vrms_A: float = 0.5, - verbose: int = 0, - _orbit_unroll: bool = False, - # --- Phaser model-prep knobs (defaults ON post v26 validation: see - # the SLURM v26 sweep in slurm_logs/rescore_v26_103820_*.csv). --- - apply_bulk_solvent: bool = True, - solvent_fsol: float = 0.95, - solvent_bsol: float = 300.0, - vrms_strategy: str = "oeffner", - vrms_identity: float = 1.0, - apply_wilson_b: bool = True, - use_epsilon: bool = False, - # obs-side term toggles (all default ON = production) for knockout bisection - frf_use_m_filter: bool = True, - frf_use_shell_variance: bool = True, - frf_use_french_wilson: bool = True, - frf_use_lerf1: bool = True, - frf_acentric_only: bool = False, - frf_d_max: float = 100.0, - frf_obs_lmax: Optional[int] = None, - frf_obs_solid_angle: bool = False, - frf_patterson_radius_scale: float = 1.0, - # Run the dominant SH-Bessel expansion (Legendre/Y_lm precompute + radial - # contraction) in single precision. The contraction is the FRF's bottleneck; - # FP64 is rate-limited on GPUs and SIMD-narrower on CPUs. The spherical-Bessel - # downward recurrence keeps its float64 internals and the cross-chunk - # accumulator stays full-precision, so only the contraction loses precision. - # `None` (default) → float32 on CUDA, full precision on CPU. Set explicitly - # to True/False to force single/double precision on either device. - frf_einsum_float32: Optional[bool] = None, -): - """Phaser-faithful (validated) rotation search — the production default. - - Reproduces the v19 benchmark config that solved the high-symmetry cases - (4BX9 342→4-7, 6G9X 77→1-4; see ``FRF_CONSOLIDATION.md``): - - * obs taken at the **full data resolution** — ``auto_lmax`` coarsens the SH - bandwidth to ``cap`` internally (the resolution↔bandwidth coupling that - removes the aliasing background), so we do not pre-restrict resolution; - * Popov-Bourenkov **anisotropy correction** (reuses ``frf.U_aniso``); - * obs **symmetry-unroll** to the full reciprocal sphere (critical for - high-symmetry spacegroups — the SH invariant subspace is otherwise - under-sampled); - * **dense P1-box calc** (single molecular transform, not unrolled) at the - coarsened resolution — fixes high-l SH under-determination on large models; - * French-Wilson + shell-variance weights; stable Wigner-d; all under no_grad. - - Returns the validated engine's peak list (``frf.types.RotationPeak``, whose - ``.score`` aliases ``.value`` so it is drop-in for the ball-search peaks). - """ - from .frf.api import phaser_lmax_resolution, phaser_rotation_search - from .frf.dense_calc import dense_calc_via_box - - device = frf.device - with torch.no_grad(): - rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) - hkl_all = data.hkl.to(device) - s_vec_all = hkl_all.to(torch.float64) @ rec_basis - s_mag_all = s_vec_all.norm(dim=-1) - # Full data resolution window; auto_lmax coarsens d_min to match the cap. - # d_max ≈ no low-res cutoff (matches the validated config's d_max_mimic). - d_min_eff = float(1.0 / s_mag_all.max().item()) - d_max_eff = float(frf_d_max) # low-resolution cutoff (default 100 ≈ none) - keep = (s_mag_all >= 1.0 / d_max_eff) & (s_mag_all <= 1.0 / d_min_eff) - - s_obs = s_vec_all[keep] - # Anisotropy correction (reuse the tensor fitted in _prepare_frf_inputs). - F_obs = apply_overall_anisotropy( - data.F.to(torch.float64).abs().to(device)[keep], s_obs, frf.U_aniso, - ) - sigF = ( - data.F_sigma.to(torch.float64).to(device)[keep] - if getattr(data, "F_sigma", None) is not None - else None - ) - centric = ( - data.centric[keep].to(torch.bool).to(device) - if hasattr(data, "centric") - else torch.zeros_like(F_obs, dtype=torch.bool) - ) - - # Obs symmetry-unroll → full reciprocal space (each ASU reflection becomes - # n_ops entries carrying the same |F|², centric, σF — |F(Sh)|=|F(h)|). - # - # The `_orbit_unroll=True` path uses `epsilon_aware_unroll` (Phaser - # DataMR.cc:954-986's `!duplicate(isym, rhkl)` skip — keeps only unique - # orbit positions). It is OFF by default: as a standalone change it - # regressed the rebench (job 103409: 3K7M 18->189, 3GR5 47->204, 2DQ6 - # 202->324). The dedup is correct only as part of a coordinated Phaser- - # faithful preprocessing chain (ε-Wilson + V(h) + σ_A), pending. - sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) - # NOTE: centered lattices (I/C/F) list each point-group rotation once per - # centering op, so the raw matrices over-replicate the obs orbit (C2/I422 - # → ×2). Deduping to unique rotations is the correct point group and saves - # that compute, BUT it is NOT result-neutral: the equal-COUNT Wilson shells - # rebin when the obs count changes, perturbing the normalisation (3A5V - # 3→4). Left as-is to keep FRF behaviour stable; tracked in - # GHOST_INVESTIGATION.md as a follow-up (fix needs count-independent shells). - # Integer unrolled Miller indices aligned with s_obs — needed for the - # ε(h) multiplicity correction (use_epsilon), which down-weights the - # axial/zonal reflections that otherwise over-weight the m=0 SH column - # and feed high-symmetry rotation-function ghosts (compute_epsilon docstring). - if _orbit_unroll: - from .frf.preprocessing import epsilon_aware_unroll - hkl_keep_int = hkl_all.to(torch.long).to(device)[keep] - unrolled_hkl, asu_idx = epsilon_aware_unroll(hkl_keep_int, sg_mats) - s_obs = unrolled_hkl.to(torch.float64) @ rec_basis - hkl_obs_int = unrolled_hkl.to(torch.float64) - F_obs = F_obs[asu_idx] - centric = centric[asu_idx] - if sigF is not None: - sigF = sigF[asu_idx] - else: - n_ops = int(sg_mats.shape[0]) - hkl_keep = hkl_all.to(torch.float64)[keep] - # h' = h.R (transpose) -- see the note at the `hkl_sym` unroll. - hkl_unroll = torch.einsum("kji,nj->kni", sg_mats, hkl_keep).reshape(-1, 3) - s_obs = hkl_unroll @ rec_basis - hkl_obs_int = hkl_unroll - F_obs = F_obs.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() - centric = centric.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() - if sigF is not None: - sigF = sigF.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() - - # Optional: restrict the obs to ACENTRIC reflections before the SH - # expansion. Centric reflections lie on the reciprocal-space zones - # perpendicular to symmetry axes and carry concentrated symmetry-axis - # signal (and a heavier Wilson tail); pooling them into the obs over- - # weights the symmetry-axis channel that produces high-symmetry ghosts. - # Dropping them also makes the Wilson normalisation acentric-only. - if frf_acentric_only: - acen = ~centric - s_obs = s_obs[acen] - F_obs = F_obs[acen] - centric = centric[acen] - hkl_obs_int = hkl_obs_int[acen] - if sigF is not None: - sigF = sigF[acen] - if verbose > 0: - print(f" FRF acentric-only: kept {int(acen.sum())}/{acen.numel()} obs", - flush=True) - - # Dense P1-box calc on the (un-rotated) search model at the coarsened res. - model_radius_A = float( - (model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item() - ) - dmin_dense = phaser_lmax_resolution(model_radius_A, d_min_eff, lmax_cap)[1] - s_calc, F_calc = dense_calc_via_box( - model, d_max_eff, dmin_dense, pad=dense_pad, verbose=verbose > 0, - ) - s_calc = s_calc.to(device) - F_calc = F_calc.to(device) - - # Optional Wilson-B match on the dense calc (EnsemblePDB.cc:793-851). - # Bin obs and calc into the same shells (defined by obs s-distribution), - # regress log(/) vs s², apply DW `exp(-B·s²/4)` to F_calc. - if apply_wilson_b: - from .frf.preprocessing import fit_relative_wilson_b - s_obs_mag = s_obs.norm(dim=-1) - s_calc_mag = s_calc.norm(dim=-1) - B_rel = fit_relative_wilson_b( - F_obs.to(torch.float64), F_calc.to(torch.float64), - s_obs_mag.to(torch.float64), n_shells=20, - s_mag_calc=s_calc_mag.to(torch.float64), - ) - if abs(B_rel) > 1e-6: - F_calc = F_calc * torch.exp(-B_rel * (s_calc_mag * s_calc_mag) / 4.0) - if verbose > 0: - print(f" FRF Wilson-B applied: B_rel = {B_rel:+.2f} Ų", flush=True) - - # Optional Oeffner vrms (rms_estimate.cc:37) — depends on n_residues - # estimated from atom count (≈ 8 heavy atoms / residue). - delta_vrms_for_frf = delta_vrms_A - if vrms_strategy == "oeffner": - from .frf.preprocessing import oeffner_vrms - n_residues_est = max(1, int(model.xyz().shape[0] / 8)) - delta_vrms_for_frf = oeffner_vrms(n_residues_est, vrms_identity) - if verbose > 0: - print( - f" FRF Oeffner vrms = {delta_vrms_for_frf:.3f} Å " - f"(n_res≈{n_residues_est}, ident={vrms_identity})", - flush=True, - ) - elif vrms_strategy != "fixed": - raise ValueError( - f"vrms_strategy={vrms_strategy!r}; expected 'fixed' or 'oeffner'." - ) - - _arf, peaks = phaser_rotation_search( - s_obs, F_obs, centric, - s_calc, F_calc, - sg_mats, - d_min=d_min_eff, d_max=d_max_eff, n_peaks=n_peaks, - delta_vrms_A=delta_vrms_for_frf, - sigma_threshold=-5.0, - use_lerf1_intensity=frf_use_lerf1, - use_m_symmetry_filter=frf_use_m_filter, - sig_F_obs=sigF, - use_french_wilson=(frf_use_french_wilson and (sigF is not None)), - use_shell_variance_weights=frf_use_shell_variance, - use_epsilon=use_epsilon, - hkl_obs=hkl_obs_int, - grid_sampling_deg=grid_sampling_deg, - model_radius_A=model_radius_A, - auto_lmax=True, - lmax_cap=lmax_cap, - obs_lmax=frf_obs_lmax, - obs_solid_angle=frf_obs_solid_angle, - patterson_radius_scale=frf_patterson_radius_scale, - apply_bulk_solvent=apply_bulk_solvent, - solvent_fsol=solvent_fsol, - solvent_bsol=solvent_bsol, - compute_dtype=( - torch.complex64 - if ( - frf_einsum_float32 - if frf_einsum_float32 is not None - else (device.type == "cuda") # default: fp32 on GPU only - ) - else None - ), - ) - return peaks - - # --------------------------------------------------------------------------- # Public entry point # --------------------------------------------------------------------------- @@ -560,8 +325,7 @@ def align_model_to_data( sigma_rot_deg: float = 0.0, sigma_trans_ang: float = 0.0, sigma_b: float = 0.0, - frf_lmax_cap: int = 48, - frf_dense_pad: float = 2.0, + model_error_A: Optional[float] = None, rescore_engine: str = "m_letf1", rescore_scat_mode: str = "legacy", subpeak_refine: bool = False, @@ -584,13 +348,9 @@ def align_model_to_data( "Cannot fit an uninitialized ModelFT. Load PDB data first." ) - # `MolecularReplacementPipeline` is the implementation of record; this - # function returns its single best `ModelFT`. Drive the pipeline directly to - # get the ranked candidate list. Imported lazily to avoid an import cycle -- - # `pipeline` imports the - # stage helpers (`_prepare_frf_inputs`, `_run_frf_separate_rotation`, - # `_external_rwork`, `_DirectModelEvaluator`, `_rodrigues`, `_StageTimer`) - # from this module. + # Imported lazily to avoid an import cycle: `pipeline` imports the stage + # helpers (`_prepare_frf_inputs`, `_external_rwork`, + # `_DirectModelEvaluator`, `_rodrigues`, `_StageTimer`) from this module. from .pipeline import MolecularReplacementPipeline pipeline = MolecularReplacementPipeline( @@ -600,7 +360,7 @@ def align_model_to_data( d_min=d_min, d_max=d_max, n_shells=n_shells, ll_max_res_A=ll_max_res_A, ll_padding_factor=ll_padding_factor, n_rotation_peaks=n_rotation_peaks, n_ml_refine=n_ml_refine, - frf_lmax_cap=frf_lmax_cap, frf_dense_pad=frf_dense_pad, + model_error_A=model_error_A, rescore_engine=rescore_engine, rescore_scat_mode=rescore_scat_mode, auto_variance_weights=auto_variance_weights, use_interp_var=use_interp_var, diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index ef36e36a..b0ebd19c 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -45,13 +45,13 @@ _external_rwork, _prepare_frf_inputs, _rodrigues, - _run_frf_separate_rotation, ) from .frf.rotation_utils import ( edmonds_euler_from_rotation_matrix, rotation_matrix_from_edmonds_euler, ) from .frf.types import RotationPeak +from .rotation_search import search_peaks from .lattman_love import LattmanLoveInterpolator, estimate_interp_var from .ml_rotation import ( fit_sigma_a_per_shell, @@ -230,8 +230,7 @@ def __init__( ll_padding_factor: float = 2.0, n_rotation_peaks: int = 500, n_ml_refine: int = 20, - frf_lmax_cap: int = 48, - frf_dense_pad: float = 2.0, + model_error_A: Optional[float] = None, # --- rescore --- rescore_engine: str = "m_letf1", rescore_scat_mode: str = "legacy", @@ -274,8 +273,16 @@ def __init__( self.ll_padding_factor = ll_padding_factor self.n_rotation_peaks = n_rotation_peaks self.n_ml_refine = n_ml_refine - self.frf_lmax_cap = frf_lmax_cap - self.frf_dense_pad = frf_dense_pad + # Expected r.m.s. coordinate error of the search model, in Angstrom: + # it sets the sigma_A fall-off in the rotation function. When the caller + # does not know it, estimate it from the model's length the way Phaser + # does (Oeffner et al. 2013), assuming the sequence is the target's -- + # roughly 8 heavy atoms per residue. + if model_error_A is None: + from .frf.preprocessing import oeffner_vrms + n_residues = max(1, int(model.xyz().shape[0] / 8)) + model_error_A = oeffner_vrms(n_residues, 1.0) + self.model_error_A = float(model_error_A) self.rescore_engine = rescore_engine self.rescore_scat_mode = rescore_scat_mode @@ -493,15 +500,14 @@ def _rotation_candidates(self, frf) -> list: timer.start("3_rotation_search") if self.verbose > 0: print( - f"mr: frf_separate rotation search " - f"(dense calc + auto_lmax cap={self.frf_lmax_cap}, " - f"n_peaks={self.n_rotation_peaks})…", + f"mr: rotation search (n_peaks={self.n_rotation_peaks}, " + f"model error {self.model_error_A:.2f} A)…", flush=True, ) - peaks = _run_frf_separate_rotation( - self.model, data, frf, - lmax_cap=self.frf_lmax_cap, dense_pad=self.frf_dense_pad, - n_peaks=self.n_rotation_peaks, verbose=self.verbose, + peaks, _lmax, _d_min = search_peaks( + self.model, data, self.model_error_A, + U_aniso=frf.U_aniso, n_peaks=self.n_rotation_peaks, + verbose=self.verbose, ) timer.stop("3_rotation_search") diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index cd31643d..b4f90cfe 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -40,6 +40,11 @@ __all__ = ["RotationSolutions", "rotation_search"] +# Note for anyone reaching for the constants below programmatically: the package +# re-exports `rotation_search` (the function) under this module's own name, so +# `from torchref.experimental.alignment import rotation_search` binds the +# function, not the module. Use `importlib.import_module` for the module object. + #: Spherical-harmonic bandwidth ceiling. Phaser's own limit (``DEF_CLMN_LMAX``) #: is 100. Where the model and resolution ask for more, the resolution is @@ -175,7 +180,7 @@ def fit_anisotropy( return symmetrize_anisotropy(U, sym_cart) -def _search( +def search_peaks( model: "ModelFT", data: "ReflectionData", model_error_A: float, @@ -183,16 +188,18 @@ def _search( U_aniso: torch.Tensor, n_peaks: int, verbose: int = 0, -) -> RotationSolutions: - """Run the rotation function with the anisotropy tensor supplied. - - Split out so callers that already fitted the anisotropy for another stage do - not fit it twice; :func:`rotation_search` is the entry point. +): + """Run the rotation function, returning the engine's own peak list. + + Returns ``(peaks, lmax, d_min)``, where ``peaks`` is a list of + :class:`~torchref.experimental.alignment.frf.types.RotationPeak` in Edmonds + ZYZ. For the placement pipeline, which consumes peaks directly and has + already fitted ``U_aniso`` for its rescore stage; :func:`rotation_search` is + the entry point for everything else. """ from .frf.api import phaser_lmax_resolution, phaser_rotation_search from .frf.dense_calc import dense_calc_via_box from .frf.preprocessing import fit_relative_wilson_b - from .frf.rotation_utils import rotation_matrix_from_edmonds_euler_batch device = model.xyz().device with torch.no_grad(): @@ -296,6 +303,14 @@ def _search( compute_dtype=torch.complex64 if device.type == "cuda" else None, ) + return peaks, int(L - 1), float(d_min) + + +def _solutions(peaks, lmax: int, d_min: float, + model_error_A: float) -> RotationSolutions: + """Package a peak list as the public return type.""" + from .frf.rotation_utils import rotation_matrix_from_edmonds_euler_batch + euler = torch.tensor( [[p.alpha, p.beta, p.gamma] for p in peaks], dtype=torch.float64, ).reshape(-1, 3) @@ -311,8 +326,8 @@ def _search( scores=torch.tensor([p.score for p in peaks], dtype=torch.float64), z_scores=torch.tensor([p.sigma for p in peaks], dtype=torch.float64), euler_zyz=euler, - lmax=int(L - 1), - d_min=float(d_min), + lmax=lmax, + d_min=d_min, model_error_A=float(model_error_A), ) @@ -367,7 +382,8 @@ def rotation_search( U_aniso = fit_anisotropy( data, d_min=d_min_fit, d_max=d_max_fit, device=model.xyz().device, ) - return _search( + peaks, lmax, d_min = search_peaks( model, data, model_error_A, U_aniso=U_aniso, n_peaks=n_peaks, verbose=verbose, ) + return _solutions(peaks, lmax, d_min, model_error_A) From 3bf5243687ea97bc077edf6a3dc8f4bc7388b866 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 14:03:25 +0200 Subject: [PATCH 022/250] Collapse the rotation-function engine onto its own class `phaser_rotation_search` was a 32-parameter pass-through that constructed `FastRotationFunction` and called `score_model`. Its stated purpose -- signature parity with a module that no longer exists -- lapsed some time ago, and two of its parameters were documented as accepted and ignored. Callers now construct the engine directly. The engine keeps the parameters that describe the data, the bandwidth and the device, and loses the knockout-bisection toggles: `use_epsilon` with its `hkl_obs`, `obs_lmax`, `obs_solid_angle`, `patterson_radius_scale`, `n_var_shells`, `bessel_h_scale`, and the three flags production always set on. Two of those are now derived from the data instead of declared: French-Wilson runs when sigmas are present, and the m-filter reads the space group, which gives ZSYMM 1 for P1 anyway -- so the synthetic tests that switched them off were asking for the behaviour they already got. Also gone: the `FRF_DEBUG` environment switch, the three write-only attributes (`auto_lmax`, `_zsymm`, `_obs_lmax`, the last two advertising a calc-side reuse that `score_model` contradicts by hard-coding `zsymm=1`), and `preprocessing.solid_angle_weights`, whose only caller was the deleted quadrature-weight experiment. The lab captures the rotation function by wrapping `score_model` rather than the deleted function. Verified behaviour-preserving: single-threaded, the peak list is bit-identical before and after on 1DAW and 3GR5. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/lab/frf.py | 10 +- tests/unit/frf_separate/test_synthetic.py | 49 ++-- torchref/experimental/alignment/__init__.py | 4 +- .../experimental/alignment/frf/__init__.py | 7 +- torchref/experimental/alignment/frf/api.py | 239 +++--------------- .../alignment/frf/preprocessing.py | 40 --- .../experimental/alignment/rotation_search.py | 31 +-- 7 files changed, 79 insertions(+), 301 deletions(-) diff --git a/alignment_lab/lab/frf.py b/alignment_lab/lab/frf.py index 0135750f..8c5f8cbc 100644 --- a/alignment_lab/lab/frf.py +++ b/alignment_lab/lab/frf.py @@ -2,7 +2,7 @@ Every rank/ghost diagnostic needs the dense adaptive sample list as well as the peaks, and the engine only returns the peaks. The capture below wraps the -engine's search entry point for the duration of one call; nine scripts each +engine's scoring method for the duration of one call; nine scripts each carried their own copy of this monkeypatch. """ @@ -187,8 +187,8 @@ def run_frf( cfg = cfg or FRFConfig() captured: Dict[str, Any] = {} - def _wrapped(*args, **kwargs): - arf, peaks = _original(*args, **kwargs) + def _wrapped(self, *args, **kwargs): + arf, peaks = _original(self, *args, **kwargs) captured["arf"] = arf return arf, peaks @@ -216,8 +216,8 @@ def _wrapped(*args, **kwargs): patched(_rs, "DENSE_CALC_PAD", float(cfg.dense_pad)), \ patched(_rs, "GRID_SAMPLING_DEG", float(cfg.grid_sampling_deg)): if capture_arf: - _original = _api.phaser_rotation_search - with patched(_api, "phaser_rotation_search", _wrapped): + _original = _api.FastRotationFunction.score_model + with patched(_api.FastRotationFunction, "score_model", _wrapped): peaks, _lmax, _dmin = _rs.search_peaks( model, data, model_error_A, U_aniso=frf_inputs.U_aniso, n_peaks=cfg.n_peaks, verbose=verbose, diff --git a/tests/unit/frf_separate/test_synthetic.py b/tests/unit/frf_separate/test_synthetic.py index d708c01a..70b882e3 100644 --- a/tests/unit/frf_separate/test_synthetic.py +++ b/tests/unit/frf_separate/test_synthetic.py @@ -1,7 +1,7 @@ """Tier 2 synthetic golden-input tests for frf_separate. -End-to-end: feed a known-rotation pair into ``phaser_rotation_search`` -and check the top peak is at the right Euler. +End-to-end: feed a known-rotation pair into the engine and check the top +peak is at the right Euler. """ from __future__ import annotations @@ -10,7 +10,7 @@ import pytest import torch -from torchref.experimental.alignment.frf.api import phaser_rotation_search +from torchref.experimental.alignment.frf.api import FastRotationFunction def _random_rotation_matrix(seed: int) -> torch.Tensor: @@ -73,6 +73,14 @@ def _make_random_reflections(seed: int, n: int = 800) -> tuple: return s_vec, F, centric +def _search(s_obs, F_obs, centric_obs, s_calc, F_calc, *, sym_mats, + n_peaks, sigma_threshold, **kw): + """Construct the engine and score one model, as `rotation_search` does.""" + engine = FastRotationFunction(s_obs, F_obs, centric_obs, sym_mats, **kw) + return engine.score_model(s_calc, F_calc, n_peaks=n_peaks, + sigma_threshold=sigma_threshold) + + @pytest.mark.parametrize("seed", [0, 1, 2, 7]) def test_synthetic_rotation_recovery(seed: int): """Apply a known rotation to obs; FRF top peak must recover it (within Δ).""" @@ -86,18 +94,11 @@ def test_synthetic_rotation_recovery(seed: int): sym_mats = torch.eye(3, dtype=torch.float64).unsqueeze(0) # P1 - arf, peaks = phaser_rotation_search( - s_obs, F_obs, centric, - s_calc, F_calc, - sym_mats=sym_mats, - L=16, - d_min=4.0, d_max=15.0, - delta_vrms_A=0.5, - grid_sampling_deg=5.0, - n_peaks=20, - sigma_threshold=-100.0, - use_french_wilson=False, - use_m_symmetry_filter=False, + arf, peaks = _search( + s_obs, F_obs, centric, s_calc, F_calc, + sym_mats=sym_mats, L=16, d_min=4.0, d_max=15.0, + delta_vrms_A=0.5, grid_sampling_deg=5.0, + n_peaks=20, sigma_threshold=-100.0, ) assert len(peaks) > 0, "no peaks returned" @@ -126,13 +127,10 @@ def test_api_returns_correct_types(): s_calc, F_calc, centric = _make_random_reflections(seed=0, n=500) sym_mats = torch.eye(3, dtype=torch.float64).unsqueeze(0) - arf, peaks = phaser_rotation_search( - s_calc, F_calc, centric, - s_calc, F_calc, sym_mats, - L=8, d_min=4.0, d_max=15.0, - grid_sampling_deg=10.0, - n_peaks=5, sigma_threshold=-100.0, - use_m_symmetry_filter=False, + arf, peaks = _search( + s_calc, F_calc, centric, s_calc, F_calc, + sym_mats=sym_mats, L=8, d_min=4.0, d_max=15.0, + grid_sampling_deg=10.0, n_peaks=5, sigma_threshold=-100.0, ) assert isinstance(arf, AdaptiveRotationFunction) assert len(peaks) > 0 @@ -159,16 +157,13 @@ def test_search_is_bit_reproducible_single_threaded(): kwargs = dict( sym_mats=sym_mats, L=12, d_min=4.0, d_max=15.0, delta_vrms_A=0.5, grid_sampling_deg=8.0, n_peaks=25, sigma_threshold=-100.0, - use_french_wilson=False, use_m_symmetry_filter=False, ) n_threads = _torch.get_num_threads() _torch.set_num_threads(1) try: - _, first = phaser_rotation_search(s_calc, F_calc, centric, - s_calc, F_calc, **kwargs) - _, second = phaser_rotation_search(s_calc, F_calc, centric, - s_calc, F_calc, **kwargs) + _, first = _search(s_calc, F_calc, centric, s_calc, F_calc, **kwargs) + _, second = _search(s_calc, F_calc, centric, s_calc, F_calc, **kwargs) finally: _torch.set_num_threads(n_threads) diff --git a/torchref/experimental/alignment/__init__.py b/torchref/experimental/alignment/__init__.py index fa53edfb..7cd6a43d 100644 --- a/torchref/experimental/alignment/__init__.py +++ b/torchref/experimental/alignment/__init__.py @@ -3,7 +3,7 @@ Pure-PyTorch Patterson-based molecular replacement: -1. Fast Rotation Function (``frf.phaser_rotation_search`` / +1. Fast Rotation Function (``rotation_search``, over ``frf.FastRotationFunction``) — Phaser-faithful Bessel-radial × SH expansion, stable Wigner-d, dense P1-box calc — then ML rescoring (``ml_rotation.m_letf1_rescore``) to rank candidate orientations. @@ -48,7 +48,6 @@ dense_calc_via_box, edmonds_euler_from_rotation_matrix, phaser_lmax_resolution, - phaser_rotation_search, rotation_angular_distance_deg, rotation_matrix_from_edmonds_euler, ) @@ -138,7 +137,6 @@ __all__ = [ # Rotation search "FastRotationFunction", - "phaser_rotation_search", "phaser_lmax_resolution", "dense_calc_via_box", "RotationPeak", diff --git a/torchref/experimental/alignment/frf/__init__.py b/torchref/experimental/alignment/frf/__init__.py index da26bbd9..1dc818a3 100644 --- a/torchref/experimental/alignment/frf/__init__.py +++ b/torchref/experimental/alignment/frf/__init__.py @@ -9,11 +9,7 @@ Shared leaf math (``..sh``, ``..wigner``) lives in the parent ``alignment`` package; this sub-package imports it "up". """ -from .api import ( - FastRotationFunction, - phaser_lmax_resolution, - phaser_rotation_search, -) +from .api import FastRotationFunction, phaser_lmax_resolution from .dense_calc import dense_calc_via_box, model_sf_abs from .rotation_utils import ( edmonds_euler_from_rotation_matrix, @@ -30,7 +26,6 @@ __all__ = [ # Engine "FastRotationFunction", - "phaser_rotation_search", "phaser_lmax_resolution", "dense_calc_via_box", "model_sf_abs", diff --git a/torchref/experimental/alignment/frf/api.py b/torchref/experimental/alignment/frf/api.py index 7680050c..6387dfad 100644 --- a/torchref/experimental/alignment/frf/api.py +++ b/torchref/experimental/alignment/frf/api.py @@ -1,4 +1,4 @@ -"""Top-level ``FastRotationFunction`` class + the ``phaser_rotation_search`` wrapper. +"""The fast rotation function engine: obs-side preprocessing, then scoring. Pipeline (mirrors Phaser ``run_FRF()``): 1. Resolution mask (both sides). @@ -16,7 +16,6 @@ from __future__ import annotations import math -import os import warnings from typing import List, Optional, Tuple @@ -27,18 +26,15 @@ from .preprocessing import ( apply_shell_variance_weights, build_lerf1_intensity, - compute_epsilon, detect_zsymm, eterm_sigma_a, french_wilson_preprocess, - solid_angle_weights, wilson_normalise, - wilson_normalise_epsilon, ) from .sitelist_ang import evaluate_rotation_function from .types import AdaptiveRotationFunction, RotationPeak -__all__ = ["FastRotationFunction", "phaser_rotation_search", "phaser_lmax_resolution"] +__all__ = ["FastRotationFunction", "phaser_lmax_resolution"] def phaser_lmax_resolution( @@ -155,23 +151,12 @@ def __init__( d_max: Optional[float] = None, delta_vrms_A: float = 1.0, n_wilson_shells: int = 20, - bessel_h_scale: Optional[float] = None, - use_lerf1_intensity: bool = True, - use_m_symmetry_filter: bool = True, sig_F_obs: Optional[torch.Tensor] = None, - use_french_wilson: bool = False, - use_shell_variance_weights: bool = False, - n_var_shells: int = 20, grid_sampling_deg: float = 2.0, - hkl_obs: Optional[torch.Tensor] = None, - use_epsilon: bool = False, model_radius_A: Optional[float] = None, auto_lmax: bool = False, - lmax_cap: int = 48, # sweet spot: higher L under-determines SH modes on sparse lattice - obs_lmax: Optional[int] = None, # cap obs SH bandwidth below calc (determinacy test) - obs_solid_angle: bool = False, # angular quadrature weight to de-bias the obs SH - patterson_radius_scale: float = 1.0, # <1 tightens the Patterson integration sphere - compute_dtype: Optional[torch.dtype] = None, # complex64 → faster einsum (GPU) + lmax_cap: int = 64, + compute_dtype: Optional[torch.dtype] = None, ): self.device = s_obs.device self.real_dtype = s_obs.dtype @@ -180,7 +165,6 @@ def __init__( # Phaser-faithful coupling of bandwidth to resolution (runMR_FRF.cc:408). # Overrides L and d_min so the SH expansion is not flooded with data # finer than L can represent (the high-symmetry failure mode). - self.auto_lmax = auto_lmax if auto_lmax: if model_radius_A is None or d_min is None: raise ValueError( @@ -195,26 +179,14 @@ def __init__( self.n_wilson_shells = n_wilson_shells self.grid_sampling_deg = grid_sampling_deg - # 1. Resolution mask on obs. hkl_obs (if given) is masked in lock-step - # so the ε(h) computation below stays aligned with F_obs. + # 1. Resolution mask on obs. extras = (F_obs, centric_obs) if sig_F_obs is not None: extras = extras + (sig_F_obs,) - if hkl_obs is not None: - extras = extras + (hkl_obs,) s_obs, extras, smag_obs = _resolution_mask(s_obs, extras, d_min, d_max) - F_obs = extras[0] - centric_obs = extras[1] - if os.environ.get("FRF_DEBUG"): - import sys as _sys - print(f"[FRF_DEBUG] auto_lmax={auto_lmax} L={self.L} d_min={d_min} " - f"d_max={d_max} n_obs_after_mask={s_obs.shape[0]} " - f"model_radius_A={model_radius_A}", file=_sys.stderr, flush=True) - ei = 2 + F_obs, centric_obs = extras[0], extras[1] if sig_F_obs is not None: - sig_F_obs = extras[ei]; ei += 1 - if hkl_obs is not None: - hkl_obs = extras[ei]; ei += 1 + sig_F_obs = extras[2] if s_obs.shape[0] < n_wilson_shells * 5: raise ValueError( @@ -222,100 +194,51 @@ def __init__( f"{n_wilson_shells} Wilson shells in [{d_min}, {d_max}] Å." ) - # 1b. Multiplicity ε(h). Needs integer hkl + spacegroup operators. - epsilon = None - if use_epsilon: - if hkl_obs is None or sym_mats is None: - raise ValueError("use_epsilon=True requires hkl_obs and sym_mats.") - epsilon = compute_epsilon(hkl_obs, sym_mats) - - # 2. Bessel scaling default — Phaser's lmax · d_min (DataMR.cc:1107). - # bessel_h_scale = 2π·R_patt is the Patterson integration radius (the - # χ_Ω sphere): the Bessel argument is h = bessel_h_scale·|s|, so the - # radial basis represents the Patterson out to R_patt = bessel_h_scale - # /(2π). With auto_lmax this defaults to R_patt ≈ sphereOuter = 2·mean - # radius. `patterson_radius_scale` < 1 tightens it toward the - # short-range intra-molecular self-Patterson (excludes the noisy - # long-vector / inter-molecular tail that carries crystal symmetry). - if bessel_h_scale is None: - if d_min is None: - raise ValueError("bessel_h_scale must be set when d_min is None") - lmax = L - 1 - lmax_even = lmax if lmax % 2 == 0 else lmax - 1 - bessel_h_scale = float(lmax_even) * float(d_min) - bessel_h_scale = bessel_h_scale * float(patterson_radius_scale) - self.bessel_h_scale = bessel_h_scale - - # 3. Wilson + optional FW + DFAC on obs. French-Wilson does its own - # per-shell normalisation; ε-correction only applies to the plain - # Wilson path (FW handles axial reflections via its posterior). - if use_french_wilson: - if sig_F_obs is None: - raise ValueError("use_french_wilson=True requires sig_F_obs.") + # 2. Bessel scaling — Phaser's lmax · d_min (DataMR.cc:1107). + # `bessel_h_scale = 2 pi R_patt` is the Patterson integration radius + # (the chi_Omega sphere): the Bessel argument is + # `h = bessel_h_scale |s|`, so the radial basis represents the + # Patterson out to `R_patt = bessel_h_scale / (2 pi)`. Under + # `auto_lmax` that comes to `R_patt ~ sphereOuter = 2 x mean radius`. + if d_min is None: + raise ValueError("d_min is required to set the Bessel scaling") + lmax = L - 1 + lmax_even = lmax if lmax % 2 == 0 else lmax - 1 + self.bessel_h_scale = float(lmax_even) * float(d_min) + + # 3. Wilson normalisation. With sigmas, through the French-Wilson + # posterior, which does its own per-shell normalisation and handles + # the axial reflections; without them, plain per-shell Wilson. + if sig_F_obs is not None: fw = french_wilson_preprocess( F_obs, sig_F_obs, smag_obs, centric_obs, n_wilson_shells=n_wilson_shells, ) - eEobs = fw["eEobs"] - dfac = fw["DFAC"] - # Fold ε into eEobs² post-hoc: divide by sqrt(ε) so the effective - # intensity is I/ε (axial reflections de-weighted). - if epsilon is not None: - eEobs = eEobs / epsilon.sqrt().to(eEobs.dtype) - elif epsilon is not None: - E_obs, _ = wilson_normalise_epsilon( - F_obs, smag_obs, epsilon, n_wilson_shells, - ) - eEobs = E_obs - dfac = torch.ones_like(E_obs) + eEobs, dfac = fw["eEobs"], fw["DFAC"] else: - E_obs, _ = wilson_normalise(F_obs, smag_obs, n_wilson_shells) - eEobs = E_obs - dfac = torch.ones_like(E_obs) + eEobs, _ = wilson_normalise(F_obs, smag_obs, n_wilson_shells) + dfac = torch.ones_like(eEobs) - # 4. LERF1 obs intensity. + # 4. LERF1 obs intensity, and the per-shell variance reweight. intensity_obs = build_lerf1_intensity( - eEobs, centric_obs, dfac=dfac, - use_centric_weight=use_lerf1_intensity, + eEobs, centric_obs, dfac=dfac, use_centric_weight=True, + ) + intensity_obs = apply_shell_variance_weights( + intensity_obs, smag_obs, n_var_shells=n_wilson_shells, ) - # 5. Optional shell-variance reweight. - if use_shell_variance_weights: - intensity_obs = apply_shell_variance_weights( - intensity_obs, smag_obs, n_var_shells=n_var_shells, - ) - - # 5b. Optional angular quadrature weight: de-bias the obs SH expansion - # for the non-uniform reciprocal-lattice point distribution (denser - # along symmetry directions → amplified symmetry-axis/ghost channel). - if obs_solid_angle: - w_ang = solid_angle_weights(s_obs).to(intensity_obs.dtype) - intensity_obs = intensity_obs * w_ang - - # 6. ZSYMM detection + m-symmetry filter on obs SH coefficients. - zsymm = detect_zsymm(sym_mats) if use_m_symmetry_filter else 1 - self._zsymm = zsymm # also reused on the calc side when enabled (score_model) + # 5. ZSYMM m-filter on the obs SH coefficients. The calc side is never + # filtered -- see score_model. + zsymm = detect_zsymm(sym_mats) - # 7. Bessel-SH expand obs side. + # 6. Bessel-SH expand the obs side. self._c_obs = bessel_sh_expand( s_obs, intensity_obs.to(self.real_dtype), - L=L, bessel_h_scale=bessel_h_scale, + L=L, bessel_h_scale=self.bessel_h_scale, zsymm=zsymm, enforce_friedel=True, compute_dtype=self.compute_dtype, ) - # Optional: cap the obs SH bandwidth BELOW the calc's by zeroing high-l - # obs coefficients. The obs is expanded over the sparse, anisotropically- - # distributed reciprocal crystal lattice (denser along symmetry - # directions), so its high-l coefficients are aliased/under-determined - # and over-represent the symmetry-axis (ghost) channel. Truncating obs-l - # while keeping calc-l full tests whether that aliasing drives the - # high-symmetry ghosts. L (and the contraction) are unchanged; the rows - # l > obs_lmax just contribute nothing. - if obs_lmax is not None and obs_lmax < (L - 1): - l_idx = torch.arange(L, device=self.device) - self._c_obs.coeffs[:, l_idx > int(obs_lmax), :] = 0.0 - self._obs_lmax = obs_lmax def score_model( self, @@ -371,93 +294,3 @@ def score_model( nms_radius_deg=max(2.0 * self.grid_sampling_deg, 6.0), ) return arf, peaks - - -def phaser_rotation_search( - s_obs: torch.Tensor, - F_obs: torch.Tensor, - centric_obs: torch.Tensor, - s_calc: torch.Tensor, - F_calc: torch.Tensor, - sym_mats: torch.Tensor, - *, - L: int = 24, - d_min: Optional[float] = None, - d_max: Optional[float] = None, - delta_vrms_A: float = 1.0, - n_wilson_shells: int = 20, - n_peaks: int = 500, - refine_subvoxel: bool = True, # accepted for signature parity; ignored - n_refine: int = 50, # ignored - sigma_threshold: float = -5.0, - bessel_h_scale: Optional[float] = None, - use_lerf1_intensity: bool = True, - use_m_symmetry_filter: bool = True, - sig_F_obs: Optional[torch.Tensor] = None, - use_french_wilson: bool = False, - use_shell_variance_weights: bool = False, - n_var_shells: int = 20, - grid_sampling_deg: float = 2.0, - hkl_obs: Optional[torch.Tensor] = None, - use_epsilon: bool = False, - model_radius_A: Optional[float] = None, - auto_lmax: bool = False, - lmax_cap: int = 48, # sweet spot: higher L under-determines SH modes on sparse lattice - obs_lmax: Optional[int] = None, # cap obs SH bandwidth below calc (determinacy test) - obs_solid_angle: bool = False, # angular quadrature weight to de-bias the obs SH - patterson_radius_scale: float = 1.0, # <1 tightens the Patterson integration sphere - # Phaser bulk-solvent (Babinet) folded into σ_A on the calc side - # (EnsemblePDB.cc:96-100; solTerm.h:9). Default OFF; flip after sweep. - apply_bulk_solvent: bool = False, - solvent_fsol: float = 0.95, - solvent_bsol: float = 300.0, - compute_dtype: Optional[torch.dtype] = None, -) -> Tuple[AdaptiveRotationFunction, List[RotationPeak]]: - """Construct a :class:`FastRotationFunction` and score one model. - - ``refine_subvoxel`` and ``n_refine`` are accepted and ignored: the per-β - fixed-shape FFT already provides sub-voxel precision through its bilinear - interpolation, and Phaser runs no extra quadratic refinement either. - - Extra (non-legacy) kwargs: - hkl_obs : integer Miller indices aligned with s_obs, needed for ε(h). - use_epsilon : apply the ε(h) multiplicity correction to Wilson - normalisation (de-weights axial reflections). Requires hkl_obs. - """ - # The FRF is forward-only (peak search, no backprop). Without no_grad the - # SH-Bessel expansion, the per-l Wigner contraction recurrence, and the FFT - # accumulate an autograd graph across every loop iteration — the dominant - # memory cost (tens to >100 GB at L≈100 / dense grids), and the cause of the - # OOMs. Disable grad for the whole engine. - with torch.no_grad(): - frf = FastRotationFunction( - s_obs, F_obs, centric_obs, sym_mats, - L=L, d_min=d_min, d_max=d_max, - delta_vrms_A=delta_vrms_A, - n_wilson_shells=n_wilson_shells, - bessel_h_scale=bessel_h_scale, - use_lerf1_intensity=use_lerf1_intensity, - use_m_symmetry_filter=use_m_symmetry_filter, - sig_F_obs=sig_F_obs, - use_french_wilson=use_french_wilson, - use_shell_variance_weights=use_shell_variance_weights, - n_var_shells=n_var_shells, - grid_sampling_deg=grid_sampling_deg, - hkl_obs=hkl_obs, - use_epsilon=use_epsilon, - model_radius_A=model_radius_A, - auto_lmax=auto_lmax, - lmax_cap=lmax_cap, - obs_lmax=obs_lmax, - obs_solid_angle=obs_solid_angle, - patterson_radius_scale=patterson_radius_scale, - compute_dtype=compute_dtype, - ) - return frf.score_model( - s_calc, F_calc, - n_peaks=n_peaks, - sigma_threshold=sigma_threshold, - apply_bulk_solvent=apply_bulk_solvent, - solvent_fsol=solvent_fsol, - solvent_bsol=solvent_bsol, - ) diff --git a/torchref/experimental/alignment/frf/preprocessing.py b/torchref/experimental/alignment/frf/preprocessing.py index bcfb4434..31fb0d52 100644 --- a/torchref/experimental/alignment/frf/preprocessing.py +++ b/torchref/experimental/alignment/frf/preprocessing.py @@ -259,46 +259,6 @@ def build_lerf1_intensity( return cw * (eEobs * eEobs - 1.0) * (dfac * dfac) -def solid_angle_weights( - s_vec: torch.Tensor, - n_cos_theta: int = 16, - n_phi: int = 32, -) -> torch.Tensor: - """Per-reflection angular quadrature weight to de-bias the SH expansion. - - The obs SH coefficient is a discretised ``∫ Y*_lm I dΩ`` over the reciprocal - crystal lattice. The lattice points are NOT uniform on the sphere — they - cluster along the cell's symmetry directions — so the unweighted sum - over-represents those directions and amplifies the symmetry-axis (ghost) - channel. This returns a weight ``w_i = 1 / (count in i's equal-area angular - cell)`` (normalised so ``Σ w = N``), which equalises each direction's - contribution — a crude spherical-quadrature / inverse-density correction. - - Bins are equal-area on the sphere (uniform in ``cos θ`` and ``φ``). - - Parameters - ---------- - s_vec : (N, 3) reciprocal-space Cartesian vectors. - n_cos_theta, n_phi : int — angular bin counts (equal-area cells). - - Returns - ------- - w : (N,) weights, dtype = s_vec.dtype, normalised to ``Σ w = N``. - """ - s_mag = s_vec.norm(dim=-1).clamp(min=1e-30) - hat = s_vec / s_mag.unsqueeze(-1) - cos_t = hat[..., 2].clamp(-1.0, 1.0) - phi = torch.atan2(hat[..., 1], hat[..., 0]) # [-π, π] - ti = ((cos_t + 1.0) * 0.5 * n_cos_theta).floor().clamp(0, n_cos_theta - 1).to(torch.int64) - pi = ((phi + math.pi) / (2.0 * math.pi) * n_phi).floor().clamp(0, n_phi - 1).to(torch.int64) - cell = ti * n_phi + pi # (N,) - n_cells = n_cos_theta * n_phi - count = torch.bincount(cell, minlength=n_cells).clamp(min=1) - w = 1.0 / count[cell].to(torch.float64) - w = w * (float(s_vec.shape[0]) / w.sum().clamp(min=1e-30)) - return w.to(s_vec.dtype) - - def apply_shell_variance_weights( intensity: torch.Tensor, s_mag: torch.Tensor, diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index b4f90cfe..1750f79b 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -197,7 +197,7 @@ def search_peaks( already fitted ``U_aniso`` for its rescore stage; :func:`rotation_search` is the entry point for everything else. """ - from .frf.api import phaser_lmax_resolution, phaser_rotation_search + from .frf.api import FastRotationFunction, phaser_lmax_resolution from .frf.dense_calc import dense_calc_via_box from .frf.preprocessing import fit_relative_wilson_b @@ -242,8 +242,9 @@ def search_peaks( sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) n_ops = int(sg_mats.shape[0]) hkl_keep = hkl_all.to(torch.float64)[keep] - hkl_obs = torch.einsum("kji,nj->kni", sg_mats, hkl_keep).reshape(-1, 3) - s_obs = hkl_obs @ rec_basis + hkl_unrolled = torch.einsum( + "kji,nj->kni", sg_mats, hkl_keep).reshape(-1, 3) + s_obs = hkl_unrolled @ rec_basis F_obs = F_obs.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() centric = centric.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() if sigF is not None: @@ -276,32 +277,28 @@ def search_peaks( if verbose > 0: print(f" relative Wilson B = {B_rel:+.2f} A^2", flush=True) - arf, peaks = phaser_rotation_search( - s_obs, F_obs, centric, - s_calc, F_calc, - sg_mats, - d_min=d_min_data, d_max=d_max, n_peaks=n_peaks, + engine = FastRotationFunction( + s_obs, F_obs, centric, sg_mats, + d_min=d_min_data, d_max=d_max, delta_vrms_A=float(model_error_A), - sigma_threshold=SIGMA_THRESHOLD, - use_lerf1_intensity=True, - use_m_symmetry_filter=True, + n_wilson_shells=N_WILSON_SHELLS, sig_F_obs=sigF, - use_french_wilson=sigF is not None, - use_shell_variance_weights=True, - hkl_obs=hkl_obs, grid_sampling_deg=GRID_SAMPLING_DEG, model_radius_A=model_radius_A, auto_lmax=True, lmax_cap=LMAX_CAP, - apply_bulk_solvent=True, - solvent_fsol=SOLVENT_FSOL, - solvent_bsol=SOLVENT_BSOL, # The spherical-harmonic contraction dominates the runtime and is # rate-limited in double precision on accelerators. Its float32 path # keeps the Bessel recurrence and the cross-chunk accumulator at # full precision. compute_dtype=torch.complex64 if device.type == "cuda" else None, ) + _arf, peaks = engine.score_model( + s_calc, F_calc, n_peaks=n_peaks, + sigma_threshold=SIGMA_THRESHOLD, + apply_bulk_solvent=True, + solvent_fsol=SOLVENT_FSOL, solvent_bsol=SOLVENT_BSOL, + ) return peaks, int(L - 1), float(d_min) From 7554cadac9de84027c7448dcfda47b4ff419b7ec Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 14:06:21 +0200 Subject: [PATCH 023/250] Record the measurements behind the rotation search's constants Ten benchmark structures at ten seeded orientations each, scored on whether truth lands inside the top twenty candidates -- the window the placement search carries -- for *every* structure, not on a median that lets one failing space group be cancelled by nine easy ones. A repeat-baseline arm in the same sweep puts the engine's own run-to-run spread at 1 cell in 100, which bounds what any of this can resolve. `LMAX_CAP = 64`, and the optimum is interior. All cells inside the top twenty: 95/100 at 48, 98/100 at 64, 98/100 at 100. The binding case is 1AK5 (P 4 3 2) at 6/10, 9/10 and 8/10, so only 64 clears nine of ten everywhere. Phaser's own ceiling of 100 is worse there, six to ten times slower, and needs more than 32 GB on three of the ten. Two candidates measured and rejected, both now documented where a reader will look for them: - the orbit-deduplicated obs unroll churns 28 of 100 cells in both directions (26 better, 13 worse) with the binding structure unchanged; - the two-radius Patterson union costs exactly double and changes one cell in a hundred. An earlier measurement had favoured it; that gain does not survive the anisotropy fix, which is what it had been compensating for. `aggregate.py --gate` now refuses to run when the cells carry different trial counts. It had been comparing ten trials of one structure against twenty of another -- the pre-OOM rows and their re-run had both landed -- which silently inverted which arms passed. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/aggregate.py | 20 +++++++++++++ docs/changelog.rst | 12 ++++++++ .../experimental/alignment/rotation_search.py | 30 ++++++++++++++++--- 3 files changed, 58 insertions(+), 4 deletions(-) diff --git a/alignment_lab/analysis/aggregate.py b/alignment_lab/analysis/aggregate.py index 1563849a..c8991aa8 100644 --- a/alignment_lab/analysis/aggregate.py +++ b/alignment_lab/analysis/aggregate.py @@ -154,6 +154,26 @@ def _in_top(r: dict) -> bool: return False return 0 <= v < top_n + # A structure appearing with more trials than the others means rows were + # collected twice -- a re-run after a partial failure, say -- and the + # per-structure hit counts are then not comparable. That has to be loud: it + # silently flips which arms pass. + counts = {} + for arm in arms: + for pdb in pdbs: + n = len(per.get((arm, pdb), [])) + if n: + counts.setdefault(n, []).append(f"{arm}/{pdb}") + if len(counts) > 1: + detail = ", ".join( + f"{n} trials: {len(v)} cell(s) e.g. {v[0]}" + for n, v in sorted(counts.items())) + raise SystemExit( + f"inconsistent trial counts across cells ({detail}). Deduplicate the " + f"inputs -- comparing 10 trials of one structure against 20 of " + f"another makes the gate meaningless." + ) + print(f"\n# shipping gate: truth in the top {top_n} on >= {min_hits} trials, " f"for every structure") print(f"{key:<26} {'pass':>5} {'worst structure':>16} {'total':>7} " diff --git a/docs/changelog.rst b/docs/changelog.rst index 30fcbb07..6dc08ad1 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -2,6 +2,18 @@ Changelog ========= +Unreleased +---------- +- Fixed the reciprocal-space symmetry convention in the alignment package (``h.S``, not ``S.h``) +- Fixed ``hkl_symops_to_cartesian`` returning non-rotations in trigonal and hexagonal settings, which corrupted the anisotropy projection +- Fixed the overall-anisotropy fit, which regressed log intensities with no constant term and so absorbed the ``-gamma`` offset into the tensor +- Fixed molecular-replacement rotation candidates being composed onto each other instead of onto the search model +- Fixed assigning a ``SpaceGroup`` object to ``Model.spacegroup`` being a silent no-op that then made the correct name assignment raise +- Fixed ``Model.copy()`` dropping the iso/aniso partition, so a copy raised from ``get_iso()`` +- Fixed ``Model.copy()`` registering the original's space group as a second submodule of the copy +- Replaced the fast rotation function's keyword surface with ``rotation_search(model, data, model_error_A)``; the caller's coordinate error is now used rather than overwritten by an estimate from the atom count +- Removed the rotation function's dead modules, engine variants, debug environment switches and unreachable knobs + Version 0.6.4 ---------- - Fixed the bulk-solvent ``F_sol`` staying at the starting model's mask for every refinement macrocycle diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index 1750f79b..0955acec 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -46,13 +46,22 @@ # function, not the module. Use `importlib.import_module` for the module object. -#: Spherical-harmonic bandwidth ceiling. Phaser's own limit (``DEF_CLMN_LMAX``) -#: is 100. Where the model and resolution ask for more, the resolution is -#: coarsened to match the bandwidth instead -- see ``phaser_lmax_resolution``. +#: Spherical-harmonic bandwidth ceiling. Where the model and the resolution ask +#: for more, the resolution is coarsened to match instead -- see +#: ``phaser_lmax_resolution``. +#: +#: Chosen by measurement, and the optimum is interior: over ten structures at +#: ten seeded orientations, truth lands in the top twenty on 95/100 cells at 48, +#: 98/100 at 64 and 98/100 at 100, but the binding case is 1AK5 (P 4 3 2), which +#: manages 6/10, 9/10 and 8/10. Only 64 clears nine of ten on every structure. +#: Phaser's own ceiling is 100 (``DEF_CLMN_LMAX``); here that is both worse on +#: 1AK5 and six to ten times slower, and it needs more than 32 GB on three of +#: the ten. LMAX_CAP = 64 #: SO(3) sample spacing in degrees for the rotation-function grid. Also sets the -#: peak-suppression radius, as ``max(2 * this, 6)`` degrees. +#: peak-suppression radius, as ``max(2 * this, 6)`` degrees. Inherited from the +#: configuration every measurement on this engine was made with. GRID_SAMPLING_DEG = 3.0 #: Edge of the P1 box the model's transform is sampled in, as a multiple of the @@ -78,6 +87,19 @@ #: wants the low-resolution terms, which carry the molecular envelope. LOW_RESOLUTION_CUTOFF_A = 100.0 +# Two things deliberately absent, both measured and rejected on the same panel: +# +# * **Orbit-deduplicated obs unroll.** Keeping only the distinct positions in +# each reflection's orbit, as Phaser does, rather than all n_ops copies. It +# moves 28 of 100 cells and in both directions -- 26 better, 13 worse against +# the shipped configuration -- with the binding structure unchanged at 9/10. A +# quarter of the results churned for no net gain. +# * **Two-radius Patterson union.** Running the search at two integration radii +# and merging the peak lists by z-score. Exactly double the cost (8.7 s +# against 4.4 s median) and it changes 1 cell in 100, which is the engine's own +# run-to-run spread. An earlier measurement had favoured it; that result does +# not survive the anisotropy fix. + #: Resolution window ``(d_max, d_min)`` the overall anisotropy is fitted in. #: The tensor is then applied across the full range. Inherited from the range #: every measurement on this engine was made with, and not itself measured -- From 9bebf1e97b8a2071ccf9ca73a4d850add30ee8e0 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 14:08:53 +0200 Subject: [PATCH 024/250] Annotate the rotation search's return types Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- torchref/experimental/alignment/rotation_search.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index 0955acec..99814054 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -23,7 +23,7 @@ from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING, List, Optional, Tuple import torch @@ -37,6 +37,7 @@ if TYPE_CHECKING: # pragma: no cover - typing only from ...io.datasets.reflection_data import ReflectionData from ...model.model_ft import ModelFT + from .frf.types import RotationPeak __all__ = ["RotationSolutions", "rotation_search"] @@ -210,7 +211,7 @@ def search_peaks( U_aniso: torch.Tensor, n_peaks: int, verbose: int = 0, -): +) -> Tuple[List["RotationPeak"], int, float]: """Run the rotation function, returning the engine's own peak list. Returns ``(peaks, lmax, d_min)``, where ``peaks`` is a list of @@ -325,7 +326,7 @@ def search_peaks( return peaks, int(L - 1), float(d_min) -def _solutions(peaks, lmax: int, d_min: float, +def _solutions(peaks: List["RotationPeak"], lmax: int, d_min: float, model_error_A: float) -> RotationSolutions: """Package a peak list as the public return type.""" from .frf.rotation_utils import rotation_matrix_from_edmonds_euler_batch From 86cb932cefef7c505601824b8724e5a8f56d3052 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 14:39:00 +0200 Subject: [PATCH 025/250] Add the rotation search's standing benchmark: accuracy, memory and runtime One row per (structure, trial, arm) carrying all three, so a change cannot buy one at the silent expense of another -- a bandwidth that halves the runtime while dropping the true orientation out of the carried window is not a win, and neither is one that needs memory the machine does not have. `lab/profile.py` holds the primitives. Three hazards it is explicit about: - **Instrumentation points.** `frf/api.py` binds `bessel_sh_expand` and its neighbours into its own namespace at import, so wrapping them in the module that defines them intercepts nothing and the stage reports zero calls -- indistinguishable from a free stage. `FRF_STAGES` names the module where each call is *resolved*. Getting it wrong left 85% of the runtime unattributed; with it right, attribution is 97% and `bessel_sh_expand` is 59% of the run. - **Nested stages.** `evaluate_rotation_function` contains `build_dense_map_per_beta` contains `wigner_contraction_per_beta`, and `bessel_sh_expand` contains `spherical_bessel_table`. Reported exclusive alongside inclusive, so the column sums. - **Wall clock on a shared cluster measures the cluster.** Every row carries a fixed calibration workload timed in the same process plus the host identity, and the array script takes an exclusive node with a pinned thread count rather than inheriting one from the allocation. Peak memory is an RSS sampler, so a spike shorter than the interval is invisible and glibc may not return freed pages; both are stated where the numbers are, and `vm_hwm_mb` carries the process-lifetime high-water mark. `bench_stages.py` is retired: this subsumes it, including the "inner stages register 0 calls" gap its README recorded. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/README.md | 37 ++- alignment_lab/analysis/aggregate.py | 81 +++++++ alignment_lab/analysis/benchmark_array.sh | 51 ++++ alignment_lab/diagnostics/bench_stages.py | 144 ------------ alignment_lab/diagnostics/frf_benchmark.py | 227 ++++++++++++++++++ alignment_lab/lab/__init__.py | 7 + alignment_lab/lab/profile.py | 257 +++++++++++++++++++++ 7 files changed, 651 insertions(+), 153 deletions(-) create mode 100644 alignment_lab/analysis/benchmark_array.sh delete mode 100644 alignment_lab/diagnostics/bench_stages.py create mode 100644 alignment_lab/diagnostics/frf_benchmark.py create mode 100644 alignment_lab/lab/profile.py diff --git a/alignment_lab/README.md b/alignment_lab/README.md index f1428eab..f977f89e 100644 --- a/alignment_lab/README.md +++ b/alignment_lab/README.md @@ -63,12 +63,31 @@ per-trial values visible, never a bare median, and prints whatever it dropped. neighbour under 1°), so "the closest sample is within a degree" means nothing by itself. -## Known gap - -`bench_stages.py` currently attributes time only to `phaser_rotation_search` -(~81–85%) and `dense_calc_via_box` (~14–17%). The inner Bessel/Wigner/peak -stages register **0 calls** — the separated engine does not route through those -module-level symbols, so wrapping them there intercepts nothing. They are -printed with their zero counts rather than omitted, because an absent row reads -as a free stage. Getting the inner breakdown needs different instrumentation -points. +## The benchmark + +`diagnostics/frf_benchmark.py` is the standing benchmark: accuracy, memory and +runtime in one row per (structure, trial, arm), so a change cannot buy one at the +silent expense of another. `analysis/benchmark_array.sh` runs it one structure +per **exclusive** node. + +Three things it is careful about, each of which has bitten this harness before: + +- **Instrumentation points.** `frf/api.py` binds `bessel_sh_expand` and its + neighbours into its own namespace at import, so wrapping them in the module + that *defines* them intercepts nothing and the stage reports zero calls — + indistinguishable from a free stage. `lab/profile.FRF_STAGES` names the module + where each call is **resolved**. Getting this wrong left 85% of the runtime + unattributed. +- **Nested stages.** `evaluate_rotation_function` contains + `build_dense_map_per_beta`, which contains `wigner_contraction_per_beta`; and + `bessel_sh_expand` contains `spherical_bessel_table`. The report gives + exclusive time alongside inclusive, so the column sums. +- **Wall clock on a shared cluster measures the cluster.** Every row carries a + fixed calibration workload timed in the same process plus the host identity. + Compare `seconds_per_calibration` across nodes, or raw seconds only within + one. + +Peak memory comes from an RSS sampler, so a spike shorter than the sampling +interval is invisible, and glibc may not return freed pages — which makes a later +window in the same process look cheaper than it is. `vm_hwm_mb` is the +process-lifetime high-water mark for absolute numbers. diff --git a/alignment_lab/analysis/aggregate.py b/alignment_lab/analysis/aggregate.py index c8991aa8..20c0151a 100644 --- a/alignment_lab/analysis/aggregate.py +++ b/alignment_lab/analysis/aggregate.py @@ -227,6 +227,81 @@ def _in_top(r: dict) -> bool: f"than it is resolvable here.") +def bench(rows: List[dict], key: str = "arm") -> None: + """Accuracy, memory and runtime side by side, per structure and per arm. + + All three together on purpose: a bandwidth that halves the runtime while + dropping the true orientation out of the carried window is not a win, and + neither is one that finds it using memory the machine does not have. + + Runtime is reported both raw and divided by the calibration workload each row + carries. If the rows span several hosts the raw column is not comparable and + the header says so. + """ + def num(r, k, default=float("nan")): + try: + return float(r[k]) + except (KeyError, TypeError, ValueError): + return default + + hosts = sorted({r.get("host", "?") for r in rows}) + threads = sorted({r.get("torch_threads", "?") for r in rows}) + print(f"\n# {len(rows)} rows from {len(hosts)} host(s) {hosts}, " + f"threads {threads}") + if len(hosts) > 1: + print("# rows span several hosts: compare s/cal, not seconds") + kinds = sorted({r.get("timing_kind", "?") for r in rows}) + if len(kinds) > 1: + print(f"# WARNING: mixed timing kinds {kinds} -- cold and steady-state " + f"numbers are not comparable") + + arms = sorted({r.get(key, "") for r in rows}) + pdbs = sorted({r.get("pdb", "?") for r in rows}) + cells = defaultdict(list) + for r in rows: + cells[(r.get(key, ""), r.get("pdb", "?"))].append(r) + + for arm in arms: + mine = [r for r in rows if r.get(key) == arm] + if not mine: + continue + print(f"\n## {key}={arm}") + print(f" {'pdb':7s} {'sg':11s} {'n':>3s} {'top-N':>6s} {'med rank':>9s} " + f"{'med s':>8s} {'s/cal':>8s} {'peak MB':>9s} {'delta MB':>9s}") + for pdb in pdbs: + rs = cells.get((arm, pdb), []) + if not rs: + continue + hits = sum(1 for r in rs if str(r.get("in_top_n")) == "1") + print(f" {pdb:7s} {rs[0].get('spacegroup', '?'):11s} {len(rs):3d} " + f"{hits:>3d}/{len(rs):<2d} " + f"{statistics.median(_cmp_rank(r) for r in rs):9.1f} " + f"{statistics.median(num(r, 'seconds') for r in rs):8.2f} " + f"{statistics.median(num(r, 'seconds_per_calibration') for r in rs):8.1f} " + f"{statistics.median(num(r, 'rss_peak_mb') for r in rs):9.0f} " + f"{statistics.median(num(r, 'rss_delta_mb') for r in rs):9.0f}") + worst = min( + (sum(1 for r in cells[(arm, p)] if str(r.get("in_top_n")) == "1") + / max(len(cells[(arm, p)]), 1), p) + for p in pdbs if cells.get((arm, p))) + print(f" worst structure: {worst[1]} at {100 * worst[0]:.0f}% in the " + f"carried window | peak memory across structures " + f"{max(num(r, 'rss_peak_mb', 0) for r in mine):.0f} MB | " + f"slowest {max(num(r, 'seconds', 0) for r in mine):.1f} s") + + # Where the time goes, if the stage columns are present. + xcols = sorted({k for r in rows for k in r if k.startswith("x_")}) + if xcols: + print("\n## exclusive stage time, median over every row (seconds)") + meds = sorted(((statistics.median(num(r, c, 0.0) for r in rows), c) + for c in xcols), reverse=True) + tot = statistics.median(num(r, "seconds") for r in rows) + for m, c in meds: + if m <= 0: + continue + print(f" {c[2:]:32s} {m:8.3f} {100 * m / max(tot, 1e-9):6.1f}%") + + def main() -> int: ap = argparse.ArgumentParser(description=__doc__) ap.add_argument("patterns", nargs="+", help="CSV glob(s)") @@ -236,6 +311,9 @@ def main() -> int: ap.add_argument("--compare", default=None, help="column whose values are the arms to pair on, " "e.g. obs_mode or lmax_cap") + ap.add_argument("--bench", action="store_true", + help="accuracy, memory and runtime side by side " + "(frf_benchmark rows)") ap.add_argument("--gate", action="store_true", help="report each arm against the shipping criterion " "(truth in the top N on most trials, every structure)") @@ -245,6 +323,9 @@ def main() -> int: help="trials per structure that must land in the top N") args = ap.parse_args() rows = load(args.patterns) + if args.bench: + bench(rows, key=args.compare or "arm") + return 0 if args.gate: gate(rows, key=args.compare or "arm", base=args.base or "production", top_n=args.top_n, min_hits=args.min_hits) diff --git a/alignment_lab/analysis/benchmark_array.sh b/alignment_lab/analysis/benchmark_array.sh new file mode 100644 index 00000000..0d1975a4 --- /dev/null +++ b/alignment_lab/analysis/benchmark_array.sh @@ -0,0 +1,51 @@ +#!/bin/bash +# The rotation search's standing benchmark: accuracy, memory and runtime. +# +# One array task per structure, all arms and trials in one process on an +# EXCLUSIVE node, so the arm comparison is within-node and the memory peaks are +# not another job's. Runtime on a shared node measures the node; every row also +# carries a calibration workload and the host identity so that is checkable +# rather than assumed. +# +# sbatch --array=0-9 --partition=hour --time=00:55:00 --exclusive \ +# --mem=200G alignment_lab/analysis/benchmark_array.sh \ +# --arms cap48,cap64,cap100 --trials 3 +# +# --mem must cover the largest arm: cap100 on the P432 structures needs well +# over 32 GB, which is what the OOMs in job 489988 were. +#SBATCH --job-name=frf_bench +#SBATCH --output=alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=alignment_lab/slurm/%x_%A_%a.err +set -uo pipefail + +REPO="${FRF_BENCH_REPO:-/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement}" +PY="$REPO/.dev/bin/python" +[ -x "$PY" ] || PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python + +cd "$REPO" +export PYTHONPATH="$REPO" +export TORCHREF_NUM_THREADS="${FRF_BENCH_THREADS:-4}" +export OMP_NUM_THREADS="$TORCHREF_NUM_THREADS" +export MKL_NUM_THREADS="$TORCHREF_NUM_THREADS" +export PYTHONUNBUFFERED=1 +export CUDA_VISIBLE_DEVICES="" + +# Pin the thread count rather than inheriting it from the allocation: an +# exclusive node hands over every core, so SLURM_CPUS_PER_TASK would make the +# timings depend on the node's size instead of on the code. +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +IDX="${SLURM_ARRAY_TASK_ID:-0}" +PDB="${PDBS[$IDX]}" + +OUTDIR="alignment_lab/runs/frf_bench_${SLURM_ARRAY_JOB_ID:-local}" +mkdir -p "$OUTDIR" alignment_lab/slurm + +echo "task $IDX -> $PDB on $(hostname), ${TORCHREF_NUM_THREADS} threads" +echo "repo=$REPO sha=$(git -C "$REPO" rev-parse --short HEAD 2>/dev/null || echo unknown)" +rc=0 +"$PY" -u -m alignment_lab.diagnostics.frf_benchmark \ + --pdb "$PDB" --out-csv "$OUTDIR/${PDB}.csv" "$@" || rc=$? + +# `rc=$?` has to follow the command directly, or a task that dies reads COMPLETED. +echo "exit_code=$rc" +exit "$rc" diff --git a/alignment_lab/diagnostics/bench_stages.py b/alignment_lab/diagnostics/bench_stages.py deleted file mode 100644 index 2f8d1645..00000000 --- a/alignment_lab/diagnostics/bench_stages.py +++ /dev/null @@ -1,144 +0,0 @@ -"""Stage-resolved timing of one cold FRF call. - -The eventual goal is placement cheap enough to sit inside a training loop, so -what matters is where the time goes, not just the total. Stage functions are -wrapped for the duration of one call and restored afterwards. - -Timings are **cold by default**: the first call in a process pays one-off costs -(parametrisation, grid setup, any compile). Pass ``--warmup`` for steady-state -numbers, and say which one a reported figure is. - -Usage:: - - python alignment_lab/diagnostics/bench_stages.py --pdb 1DAW --lmax-cap 64 -""" - -from __future__ import annotations - -import argparse -import sys -import time -from collections import defaultdict -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, FRFConfig, ResultWriter, orbit_rank, # noqa: E402 - rotated_case, run_frf, seed_for) -from lab.frf import patched # noqa: E402 - -#: (module path, attribute) pairs timed individually. -STAGES = [ - ("torchref.experimental.alignment.frf.dense_calc", "dense_calc_via_box"), - ("torchref.experimental.alignment.frf.api", "phaser_rotation_search"), - ("torchref.experimental.alignment.frf.data_mr", "spherical_bessel_table"), - ("torchref.experimental.alignment.frf.data_mr", "bessel_sh_expand"), - ("torchref.experimental.alignment.frf.data_mr", "cross_correlate_xi"), - ("torchref.experimental.alignment.frf.wigner_d", "wigner_contraction_per_beta"), - ("torchref.experimental.alignment.frf.sitelist_ang", "evaluate_rotation_function"), - ("torchref.experimental.alignment.frf.peak_finder", "find_rotation_peaks"), -] - - -def _instrument(stack, totals, counts, skipped): - """Wrap each resolvable stage with a timer, via the exit stack.""" - import importlib - - for mod_path, attr in STAGES: - try: - mod = importlib.import_module(mod_path) - original = getattr(mod, attr) - except (ImportError, AttributeError): - # Report it: a silently skipped stage reads as "that stage is free". - skipped.append(f"{mod_path.rsplit('.', 1)[-1]}.{attr}") - continue - - def make(orig, key): - def timed(*a, **k): - t0 = time.perf_counter() - try: - return orig(*a, **k) - finally: - totals[key] += time.perf_counter() - t0 - counts[key] += 1 - return timed - - # Register at zero so a resolved-but-never-called stage still prints: - # an absent row is indistinguishable from a free one. - totals[attr] += 0.0 - counts[attr] += 0 - stack.enter_context(patched(mod, attr, make(original, attr))) - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) - ap.add_argument("--trial", type=int, default=0) - ap.add_argument("--lmax-cap", type=int, default=64) - ap.add_argument("--d-min", type=float, default=4.0) - ap.add_argument("--d-max", type=float, default=15.0) - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--warmup", action="store_true", - help="discard one call first and report steady state") - ap.add_argument("--out-csv", default=None) - args = ap.parse_args() - - from contextlib import ExitStack - - seed = seed_for(args.pdb, args.trial) - rotated, data, R_true = rotated_case(args.pdb, seed) - cfg = FRFConfig(d_min=args.d_min, d_max=args.d_max, - n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) - - if args.warmup: - run_frf(rotated, data, cfg, capture_arf=False) - - totals, counts, skipped = defaultdict(float), defaultdict(int), [] - with ExitStack() as stack: - _instrument(stack, totals, counts, skipped) - res = run_frf(rotated, data, cfg) - - sym = data.spacegroup.matrices.to(torch.float64).cpu() - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - rank, ang = orbit_rank(res.peaks, R_true, sym, reciprocal_basis=rec) - - kind = "steady-state" if args.warmup else "cold" - print(f"=== {args.pdb} lmax_cap={args.lmax_cap} ({kind}) ===") - print(f" {'stage':28s} {'calls':>6s} {'seconds':>9s} {'% of run':>9s}") - accounted = 0.0 - for key, secs in sorted(totals.items(), key=lambda kv: -kv[1]): - accounted += secs - print(f" {key:28s} {counts[key]:6d} {secs:9.3f} " - f"{100.0 * secs / max(res.seconds, 1e-9):9.1f}") - print(f" {'(unattributed)':28s} {'':6s} {res.seconds - accounted:9.3f} " - f"{100.0 * (res.seconds - accounted) / max(res.seconds, 1e-9):9.1f}") - print(f" {'TOTAL':28s} {'':6s} {res.seconds:9.3f}") - if skipped: - print(f" NOT INSTRUMENTED (renamed or absent): {', '.join(skipped)}") - print(f" truth rank {rank} at {ang:.2f} deg") - - if args.out_csv: - w = ResultWriter(args.out_csv, "bench_stages", - extra_fields=("timing_kind", "total_seconds", - "unattributed_seconds") + - tuple(f"t_{a}" for _, a in STAGES)) - row = dict(pdb=args.pdb, seed=seed, trial=args.trial, - spacegroup=str(data.spacegroup), n_ops=int(sym.shape[0]), - truth_rank=rank, truth_angle_deg=round(ang, 4), - orbit_side="left", orbit_frame="cart", - lmax_cap=args.lmax_cap, d_min=args.d_min, d_max=args.d_max, - device="cpu", timing_kind=kind, - total_seconds=round(res.seconds, 4), - unattributed_seconds=round(res.seconds - accounted, 4)) - for _, attr in STAGES: - row[f"t_{attr}"] = round(totals.get(attr, 0.0), 4) - w.write(**row) - print(f" wrote {args.out_csv}") - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_benchmark.py b/alignment_lab/diagnostics/frf_benchmark.py new file mode 100644 index 00000000..41ca2c10 --- /dev/null +++ b/alignment_lab/diagnostics/frf_benchmark.py @@ -0,0 +1,227 @@ +"""The rotation search's standing benchmark: accuracy, memory and runtime. + +One row per (structure, trial, arm), carrying all three so a change cannot +improve one at the silent expense of another: + +**Accuracy** -- where the true orientation lands in the peak list, and whether +it is inside the top ``--top-n``. That window, not rank 0, is the thing that +matters: the placement search carries its top candidates forward, so rank 7 and +rank 0 are the same outcome downstream and rank 223 is not. Reported per +structure, because an average lets one failing space group be cancelled by nine +easy ones. + +**Memory** -- peak resident set over the search, as a delta over the value on +entry, plus the process high-water mark. Read +:mod:`alignment_lab.lab.profile` for what a sampler can and cannot see. Memory +is why the bandwidth ceiling is not simply "as high as possible": cap 100 needs +more than 32 GB on the two P432 structures, where the symmetry expansion +multiplies the reflection count by 24. + +**Runtime** -- total, plus per-stage. Cold by default, since a caller placing one +model pays cold costs; ``--warmup`` gives steady state, and the row says which. +Every row carries a fixed calibration workload timed in the same process, and the +node's identity: wall clock on a shared cluster measures the cluster unless it is +normalised or confined to one node. Run with ``--exclusive`` and compare +``seconds_per_calibration`` across nodes, or raw seconds only within a node. + +Usage +----- + python -m diagnostics.frf_benchmark --pdb 1DAW --trials 3 + python -m diagnostics.frf_benchmark --pdb 3K7M --arms cap48,cap64,cap100 +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, FRFConfig, orbit_rank, rotated_case, # noqa: E402 + run_frf, seed_for) +from lab.profile import (FRF_STAGES, PeakMemory, calibration_seconds, # noqa: E402 + exclusive_times, host_info, stage_timers) +from lab.results import append_row, provenance # noqa: E402 + +EXPERIMENT = "frf_benchmark" + +#: Named bandwidth arms. ``shipped`` reads the engine's own constant, so the +#: benchmark follows the code rather than restating it -- if the constant moves +#: and this row does not, the harness is measuring history. +ARMS = {"cap48": 48, "cap64": 64, "cap100": 100, "shipped": None} + + +def _shipped_lmax_cap() -> int: + import importlib + + return importlib.import_module( + "torchref.experimental.alignment.rotation_search").LMAX_CAP + + +def run_one(pdb: str, trial: int, arm: str, *, n_peaks: int, top_n: int, + thr_deg: float, warmup: bool, mem_interval_s: float) -> dict: + """One measurement: accuracy, memory and runtime for a single search.""" + lmax_cap = ARMS[arm] if ARMS[arm] is not None else _shipped_lmax_cap() + seed = seed_for(pdb, trial) + model, data, R_true = rotated_case(pdb, seed) + cfg = FRFConfig(n_peaks=n_peaks, lmax_cap=lmax_cap) + + if warmup: + run_frf(model, data, cfg, capture_arf=False, verbose=0) + + # Calibrate before the measurement, so a node that is busy *now* is visible. + calib = calibration_seconds() + + mem = PeakMemory(interval_s=mem_interval_s) + t0 = time.perf_counter() + with mem.window() as mem_out, stage_timers() as (totals, counts, unresolved): + res = run_frf(model, data, cfg, capture_arf=False, verbose=0) + wall = time.perf_counter() - t0 + + rank, ang = orbit_rank( + res.peaks, R_true, data.spacegroup.matrices.to(torch.float64).cpu(), + reciprocal_basis=data.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), + side="left", frame="cart", thr_deg=thr_deg, + ) + attributed = sum(totals.values()) + + row = {"experiment": EXPERIMENT, "pdb": pdb, "trial": trial, "arm": arm, + "seed": seed} + row.update(provenance()) + row.update(host_info()) + row.update({ + "spacegroup": str(data.spacegroup.hm), + "n_ops": int(data.spacegroup.matrices.shape[0]), + "n_atoms": int(model.xyz().shape[0]), + "n_reflections": int(data.hkl.shape[0]), + "lmax_cap": lmax_cap, + "n_peaks": n_peaks, + # --- accuracy --- + "truth_rank": rank, + # orbit_rank returns -1 when nothing matched. That must not order as a + # good rank, so for any comparison a miss counts as worse than the worst + # hit, i.e. the length of the peak list. + "rank_for_compare": rank if rank >= 0 else n_peaks, + "found": int(rank >= 0), + "in_top_n": int(0 <= rank < top_n), + "top_n": top_n, + "truth_angle_deg": None if ang is None else round(float(ang), 3), + "n_peaks_found": len(res.peaks), + "orbit_side": "left", "orbit_frame": "cart", "thr_deg": thr_deg, + # --- runtime --- + "timing_kind": "steady" if warmup else "cold", + "seconds": round(wall, 3), + "seconds_attributed": round(attributed, 3), + "seconds_unattributed": round(wall - attributed, 3), + "calibration_seconds": round(calib, 5), + "seconds_per_calibration": round(wall / max(calib, 1e-9), 1), + # --- memory --- + **mem_out, + "stages_unresolved": "|".join(unresolved), + }) + # Inclusive time for reading a single stage, exclusive for adding them up. + excl = exclusive_times(totals) + for _, attr in FRF_STAGES: + row[f"t_{attr}"] = round(totals.get(attr, float("nan")), 4) + row[f"x_{attr}"] = round(excl.get(attr, float("nan")), 4) + row[f"n_{attr}"] = counts.get(attr, 0) + return row + + +def _fmt(row: dict) -> str: + hit = "yes" if row["in_top_n"] else "NO " + return (f" {row['arm']:<8} rank={str(row['truth_rank']):<6} " + f"top{row['top_n']}={hit} " + f"{row['seconds']:>7.2f}s peak {row['rss_peak_mb']:>8.0f} MB " + f"(+{row['rss_delta_mb']:>7.0f}) {row['seconds_per_calibration']:>7.1f} cal") + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) + ap.add_argument("--trial", type=int, default=None, + help="single trial; omit to run --trials of them") + ap.add_argument("--trials", type=int, default=3) + ap.add_argument("--arms", default="shipped", + help=f"comma-separated, from {sorted(ARMS)}") + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--top-n", type=int, default=20, + help="candidates the placement search carries forward") + ap.add_argument("--thr-deg", type=float, default=5.0) + ap.add_argument("--warmup", action="store_true", + help="discard one search first and report steady state") + ap.add_argument("--mem-interval", type=float, default=0.02, + help="RSS sampling period in seconds") + ap.add_argument("--out-csv", default=None) + args = ap.parse_args() + + arms = [a for a in args.arms.split(",") if a] + unknown = [a for a in arms if a not in ARMS] + if unknown: + raise SystemExit(f"unknown arm(s) {unknown}; expected {sorted(ARMS)}") + trials = [args.trial] if args.trial is not None else list(range(args.trials)) + + csv_path = None + if args.out_csv: + csv_path = Path(args.out_csv) + csv_path.parent.mkdir(parents=True, exist_ok=True) + + info = host_info() + print(f"{args.pdb}: {len(arms)} arm(s) x {len(trials)} trial(s), " + f"{'steady-state' if args.warmup else 'cold'}", flush=True) + print(f" host {info['host']} / {info['torch_threads']} threads / " + f"{info['cpu_model'] or 'unknown cpu'}", flush=True) + + rows, n_fail = [], 0 + for trial in trials: + print(f" trial {trial}", flush=True) + for arm in arms: + try: + row = run_one(args.pdb, trial, arm, n_peaks=args.n_peaks, + top_n=args.top_n, thr_deg=args.thr_deg, + warmup=args.warmup, + mem_interval_s=args.mem_interval) + except Exception as exc: + n_fail += 1 + print(f" {arm:<8} FAILED {type(exc).__name__}: {exc}", flush=True) + continue + rows.append(row) + if csv_path: + append_row(csv_path, row) + print(_fmt(row), flush=True) + if row["stages_unresolved"]: + print(f" NOT INSTRUMENTED: {row['stages_unresolved']}", + flush=True) + + if rows: + def med(key): + vals = sorted(r.get(key, 0.0) or 0.0 for r in rows) + return vals[len(vals) // 2] + + med_total = med("seconds") + print("\nwhere the time goes (median over the rows above; exclusive of " + "nested stages, so the column sums)", flush=True) + print(f" {'stage':32s} {'excl s':>8s} {'%':>6s} {'incl s':>8s}", + flush=True) + order = sorted((a for _, a in FRF_STAGES), key=lambda a: -med(f"x_{a}")) + for attr in order: + x, t = med(f"x_{attr}"), med(f"t_{attr}") + if t <= 0: + continue + print(f" {attr:32s} {x:8.3f} {100 * x / max(med_total, 1e-9):6.1f} " + f"{t:8.3f}", flush=True) + unatt = med("seconds_unattributed") + print(f" {'(unattributed)':32s} {unatt:8.3f} " + f"{100 * unatt / max(med_total, 1e-9):6.1f}", flush=True) + if csv_path: + print(f"\nwrote {csv_path} ({n_fail} failures)", flush=True) + return 1 if rows == [] else 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/lab/__init__.py b/alignment_lab/lab/__init__.py index 72261ad7..b0b14c54 100644 --- a/alignment_lab/lab/__init__.py +++ b/alignment_lab/lab/__init__.py @@ -30,6 +30,8 @@ from .frf import (FRFConfig, FRFResult, merge_peak_lists, patched, run_frf) from .rescore import ENGINES, RescoreResult, paired_ranks, run_rescore +from .profile import (FRF_STAGES, PeakMemory, calibration_seconds, + host_info, stage_timers) from .results import ResultWriter, append_row, provenance __all__ = [ @@ -56,6 +58,11 @@ "RescoreResult", "paired_ranks", "run_rescore", + "FRF_STAGES", + "PeakMemory", + "calibration_seconds", + "host_info", + "stage_timers", "ResultWriter", "append_row", "provenance", diff --git a/alignment_lab/lab/profile.py b/alignment_lab/lab/profile.py new file mode 100644 index 00000000..9a091905 --- /dev/null +++ b/alignment_lab/lab/profile.py @@ -0,0 +1,257 @@ +"""Timing, memory and node-calibration primitives for the benchmark harness. + +Three measurements, three different hazards: + +**Time.** Wall clock on a shared cluster is a measurement of the cluster, not of +the code. Two engines timed on different nodes, or on the same node under +different neighbours, have been seen to differ by more than the effect being +looked for. So every timing row carries :func:`calibration_seconds` -- a fixed +workload run in the same process -- and the node's identity. Compare normalised +times, or compare only within a node. + +**Memory.** Peak RSS is a high-water mark: it never falls, so several +measurements in one process all report the largest. :class:`PeakMemory` samples +``/proc/self/statm`` on a thread so each window gets its own peak, and reports +the delta over the value at window entry. Two caveats, both reported rather than +hidden: a sampler can miss a spike shorter than its interval, and glibc does not +always return freed pages, so a later window in the same process can look +cheaper than it is. For absolute numbers, run one measurement per process and +read ``VmHWM``. + +**Stages.** A stage that cannot be resolved is *reported*, not skipped: an absent +row and a free stage look identical in a table. +""" + +from __future__ import annotations + +import importlib +import os +import threading +import time +from contextlib import contextmanager +from typing import Dict, List, Optional, Sequence, Tuple + +#: ``(module path, attribute)`` for each stage worth timing separately, coarse +#: to fine. +#: +#: The module named here is where the call is *resolved*, not where the function +#: is defined. ``frf.api`` binds ``bessel_sh_expand`` and friends into its own +#: namespace with ``from .data_mr import ...`` at import time, so replacing +#: ``data_mr.bessel_sh_expand`` leaves api's reference untouched and the stage +#: silently registers zero calls -- which reads as "that stage is free". Getting +#: this wrong left 85% of the runtime unattributed. +FRF_STAGES: Tuple[Tuple[str, str], ...] = ( + ("torchref.experimental.alignment.frf.dense_calc", "dense_calc_via_box"), + ("torchref.experimental.alignment.frf.api", "french_wilson_preprocess"), + ("torchref.experimental.alignment.frf.api", "bessel_sh_expand"), + ("torchref.experimental.alignment.frf.api", "cross_correlate_xi"), + ("torchref.experimental.alignment.frf.api", "evaluate_rotation_function"), + ("torchref.experimental.alignment.frf.api", "find_rotation_peaks"), + ("torchref.experimental.alignment.frf.sitelist_ang", + "wigner_contraction_per_beta"), + ("torchref.experimental.alignment.frf.sitelist_ang", + "build_dense_map_per_beta"), + ("torchref.experimental.alignment.frf.data_mr", "spherical_bessel_table"), +) + +#: Which stage each nested stage sits inside. A parent's time *includes* its +#: children, so summing the raw rows double-counts -- the harness subtracts to +#: report exclusive time as well, because otherwise the table invites optimising +#: the wrong thing. +FRF_STAGE_PARENTS = { + "build_dense_map_per_beta": "evaluate_rotation_function", + "wigner_contraction_per_beta": "build_dense_map_per_beta", + "spherical_bessel_table": "bessel_sh_expand", +} + + +def exclusive_times(totals): + """Per-stage time with nested children subtracted out. + + Parameters + ---------- + totals : mapping + Stage name -> inclusive seconds, as :func:`stage_timers` yields. + + Returns + ------- + dict + Stage name -> exclusive seconds. These sum without double counting. + """ + excl = dict(totals) + for child, parent in FRF_STAGE_PARENTS.items(): + if child in totals and parent in excl: + excl[parent] = excl[parent] - totals[child] + return excl + + +_PAGE_SIZE = os.sysconf("SC_PAGE_SIZE") if hasattr(os, "sysconf") else 4096 + + +def rss_bytes() -> int: + """Current resident set size, in bytes.""" + try: + with open("/proc/self/statm") as fh: + return int(fh.read().split()[1]) * _PAGE_SIZE + except (OSError, IndexError, ValueError): + return 0 + + +def vm_hwm_bytes() -> int: + """Process peak resident set size since start, in bytes. 0 if unavailable. + + Monotonic over the life of the process, so it answers "how much did this + process ever need", not "how much did this call need". + """ + try: + with open("/proc/self/status") as fh: + for line in fh: + if line.startswith("VmHWM:"): + return int(line.split()[1]) * 1024 + except OSError: + pass + return 0 + + +class PeakMemory: + """Sample RSS on a thread and report the peak inside a window. + + Parameters + ---------- + interval_s : float, optional + Sampling period. Default 0.02 s. Anything shorter than this that the + code allocates and frees again is invisible; ``missed_window_risk`` + records the interval so a reader can judge that. + """ + + def __init__(self, interval_s: float = 0.02): + self.interval_s = float(interval_s) + self._stop = threading.Event() + self._thread: Optional[threading.Thread] = None + self._peak = 0 + self._samples = 0 + + def _run(self) -> None: + while not self._stop.wait(self.interval_s): + r = rss_bytes() + self._samples += 1 + if r > self._peak: + self._peak = r + + @contextmanager + def window(self): + """Measure the peak RSS over the body, as a delta and an absolute.""" + baseline = rss_bytes() + self._peak = baseline + self._samples = 0 + self._stop.clear() + self._thread = threading.Thread(target=self._run, daemon=True) + self._thread.start() + out: Dict[str, float] = {} + try: + yield out + finally: + self._stop.set() + self._thread.join(timeout=5.0) + peak = max(self._peak, rss_bytes()) + out.update( + rss_baseline_mb=round(baseline / 1e6, 1), + rss_peak_mb=round(peak / 1e6, 1), + rss_delta_mb=round((peak - baseline) / 1e6, 1), + rss_samples=self._samples, + rss_sample_interval_s=self.interval_s, + vm_hwm_mb=round(vm_hwm_bytes() / 1e6, 1), + ) + + +@contextmanager +def stage_timers(stages: Sequence[Tuple[str, str]] = FRF_STAGES): + """Time each stage in ``stages`` for the duration of the body. + + Yields ``(totals, counts, unresolved)``. Every resolved stage is registered + at zero, so a stage that was instrumented but never called still shows up + with 0 calls -- distinguishable from one that is simply fast. + """ + from contextlib import ExitStack + + from .frf import patched + + totals: Dict[str, float] = {} + counts: Dict[str, int] = {} + unresolved: List[str] = [] + + def make(orig, key): + def timed(*a, **k): + t0 = time.perf_counter() + try: + return orig(*a, **k) + finally: + totals[key] += time.perf_counter() - t0 + counts[key] += 1 + return timed + + with ExitStack() as stack: + for mod_path, attr in stages: + try: + mod = importlib.import_module(mod_path) + original = getattr(mod, attr) + except (ImportError, AttributeError): + unresolved.append(f"{mod_path.rsplit('.', 1)[-1]}.{attr}") + continue + totals.setdefault(attr, 0.0) + counts.setdefault(attr, 0) + stack.enter_context(patched(mod, attr, make(original, attr))) + yield totals, counts, unresolved + + +def calibration_seconds(repeats: int = 3) -> float: + """Time a fixed workload, to normalise wall clock across nodes. + + Exercises the two kernels the rotation function spends its time in: a + complex einsum contraction and a batched FFT, both float64. Fixed shapes and + a fixed seed, so the only thing that varies is the machine and its + neighbours. Returns the best of ``repeats`` -- the least contended sample. + """ + import torch + + g = torch.Generator().manual_seed(0) + a = torch.randn(64, 96, 96, dtype=torch.float64, generator=g) + b = torch.randn(64, 96, 96, dtype=torch.float64, generator=g) + x = torch.complex(a, b) + best = float("inf") + for _ in range(max(1, repeats)): + t0 = time.perf_counter() + torch.einsum("nij,njk->nik", x, x) + torch.fft.ifft2(x) + best = min(best, time.perf_counter() - t0) + return best + + +def host_info() -> Dict[str, object]: + """Node identity and thread configuration, for every benchmark row.""" + import platform + + import torch + + model = "" + try: + with open("/proc/cpuinfo") as fh: + for line in fh: + if line.startswith("model name"): + model = line.split(":", 1)[1].strip() + break + except OSError: + pass + return { + "host": platform.node(), + "cpu_model": model, + "torch_threads": torch.get_num_threads(), + "slurm_job": os.environ.get("SLURM_JOB_ID", ""), + "slurm_cpus": os.environ.get("SLURM_CPUS_PER_TASK", ""), + "slurm_exclusive": os.environ.get("SLURM_JOB_NUM_NODES", ""), + } + + +__all__ = ["FRF_STAGES", "FRF_STAGE_PARENTS", "PeakMemory", + "calibration_seconds", "exclusive_times", "host_info", + "rss_bytes", "stage_timers", "vm_hwm_bytes"] From c8aea57e46c5f27a9c321f9388d9fa399fe6b290 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 14:40:12 +0200 Subject: [PATCH 026/250] Correct phaser_lmax_resolution's docstring It described a default of 48 while the signature said 64, then said "Default 100" in the parameter list; cited `small_d_stable` and `frf_separate`, both gone; and closed with a paragraph concluding that raising the bandwidth makes high-symmetry cases worse, on the strength of a 4BX9 measurement predating the dense P1-box calc. The sweep says otherwise: the optimum is interior, and 1AK5 rather than 4BX9 is what binds. The default is now 64, matching `rotation_search.LMAX_CAP`, and the docstring points at that constant as the one carrying the evidence. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- torchref/experimental/alignment/frf/api.py | 80 +++++++++------------- 1 file changed, 32 insertions(+), 48 deletions(-) diff --git a/torchref/experimental/alignment/frf/api.py b/torchref/experimental/alignment/frf/api.py index 6387dfad..ee6ae41d 100644 --- a/torchref/experimental/alignment/frf/api.py +++ b/torchref/experimental/alignment/frf/api.py @@ -40,65 +40,49 @@ def phaser_lmax_resolution( model_radius_A: float, d_min_data: float, - lmax_cap: int = 48, + lmax_cap: int = 64, ): - """Phaser's coupling of SH bandwidth to rotation-function resolution. - - Phaser source: ``runMR_FRF.cc:407-419``:: + """Couple the spherical-harmonic bandwidth to the rotation-function + resolution, as Phaser does (``runMR_FRF.cc:427-448``):: sphereOuter = 2 * mean_radius - LMAX = ceil(2*pi*sphereOuter / HIRES) # HIRES = data d_min - if LMAX odd: LMAX++ - LMAX = min(LMAX, DEF_CLMN_LMAX=100) - LMAX_RESO = (LMAX capped) ? 2*pi*sphereOuter/LMAX : HIRES - - The point: including data finer than the bandwidth can represent only - adds aliasing background (the discrete Y_lm are not orthogonal over - scattered reflections), which buries the symmetry-diluted true peak on - large / high-symmetry structures. So Phaser either raises LMAX to match - the resolution, or — when LMAX hits the cap — coarsens the resolution to - ``LMAX_RESO`` and drops finer reflections (DataMR.cc:984). - - **In practice this is a per-structure high-resolution cutoff.** For real - protein search models ``LMAX_ideal = ceil(2*pi*2r/d_min)`` is ~100-170, so - it always hits ``lmax_cap`` and the function reduces to: use ``L=lmax_cap`` - and keep only data coarser than ``d_min_eff = 2*pi*(2r)/lmax_cap`` — the - finest resolution that bandwidth can faithfully represent for a molecule of - radius ``r``. Bigger molecule -> coarser cutoff. The variable-L branch only - matters for tiny models / low-res data (``LMAX_ideal < lmax_cap``). - - **lmax_cap default is 48, NOT Phaser's 100 — for a different reason than - before.** The contraction now uses ``frf_separate.wigner_d.small_d_stable`` - (J_y eigendecomposition = π/2 / SOFT basis), validated stable+correct to - l=128, so there is no longer a *numerical* ceiling (the old Edmonds-sum - ``small_d_packed`` exploded to |d|~1e11 at l>=50). BUT raising L empirically - makes high-symmetry cases WORSE in this pipeline: at L=100/4.4Å on 4BX9 the - truth went from #taller=158 (cap=48) to 39084, and σA weighting did not - suppress it. Cause: at high L the SH modes (~L² per shell) are - under-determined by the obs sampled on the sparse crystal lattice (~10⁴ - reflections), so high-l coefficients are noise. Phaser avoids this by - computing the *model* transform on a dense P1-box FFT grid; we sample the - calc at the crystal lattice. Until that is changed, lmax_cap≈48 is the sweet - spot. small_d_stable is kept regardless (correct + removes the ceiling, and - is the prerequisite for any future dense-sampling high-L work). + LMAX = ceil(2 pi sphereOuter / HIRES) # HIRES = the data's d_min + if LMAX is odd: LMAX += 1 + LMAX = min(LMAX, lmax_cap) + LMAX_RESO = (LMAX hit the cap) ? 2 pi sphereOuter / LMAX : HIRES + + Data finer than the bandwidth can represent contributes aliasing rather than + signal -- the discrete ``Y_lm`` are not orthogonal over scattered + reflections -- and that background buries the symmetry-diluted true peak on + large or high-symmetry structures. So either the bandwidth rises to meet the + resolution, or, once it is capped, the resolution is coarsened to + ``LMAX_RESO`` and the finer reflections are dropped (``DataMR.cc:984``). + + For real protein search models ``ceil(2 pi 2r / d_min)`` is 100 to 170, so + the cap binds and this reduces to: use ``L = lmax_cap``, and keep only data + coarser than ``2 pi (2r) / lmax_cap`` -- the finest resolution that + bandwidth can carry for a molecule of that radius. Bigger molecule, coarser + cutoff. The variable-``L`` branch matters only for small models or + low-resolution data. Parameters ---------- model_radius_A : float - The search model's mean atomic radius from its centroid (Å). + Mean distance of the model's atoms from its centroid (Angstrom). ``sphereOuter = 2 * model_radius_A``. d_min_data : float - High-resolution limit of the data (Å). + High-resolution limit of the data (Angstrom). lmax_cap : int - Hard bandwidth cap. Default 100 (Phaser's DEF_CLMN_LMAX). The stable - small_d_stable Wigner-d makes higher caps safe at increasing compute cost - (~L^3); for very large assemblies one may even exceed Phaser's 100. + Bandwidth ceiling. The production value lives in + :data:`torchref.experimental.alignment.rotation_search.LMAX_CAP`, which + carries the measurement behind it; this default exists only for direct + callers. Cost grows about as ``l^3`` in time and ``l^2`` in memory. Returns ------- (L, d_min_eff) : Tuple[int, float] - ``L`` is the frf_separate bandwidth (lmax = L-1, even), ``d_min_eff`` - is the resolution to actually use for the expansion. + ``L`` is the bandwidth in this package's convention (``lmax = L - 1``, + even); ``d_min_eff`` is the resolution to expand at. """ sphere_outer = 2.0 * float(model_radius_A) lmax = int(math.ceil(2.0 * math.pi * sphere_outer / float(d_min_data))) @@ -111,9 +95,9 @@ def phaser_lmax_resolution( d_min_eff = float(d_min_data) if lmax >= 256: warnings.warn( - f"phaser_lmax_resolution chose lmax={lmax} (>=256): small_d_stable " - "is stable here but the contraction cost grows ~l^3 and memory ~l^2; " - "consider a tighter lmax_cap if this is slow.", + f"phaser_lmax_resolution chose lmax={lmax} (>=256): the Wigner " + f"contraction is numerically fine here, but its cost grows about as " + f"l^3 in time and l^2 in memory. Consider a tighter lmax_cap.", RuntimeWarning, stacklevel=2, ) return lmax + 1, d_min_eff # L = lmax + 1 (our bandwidth convention) From ae040e837f343951a550db228b893457af7478b8 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 14:49:55 +0200 Subject: [PATCH 027/250] Take the process start-up out of the benchmark's first measurement Whichever arm ran first carried the whole of the process's start-up: on 3A5V that was 41.4 s against 1.8 s for the same search afterwards, with 39.4 of those seconds attributable to no stage at all. On this cluster the first real computation is dominated by PyTorch loading its backend libraries, which happens lazily on first use rather than at `import torch`, off a GPFS environment where that is tens of seconds of small reads. A warm-up search now runs first, at full fidelity so the kernels and FFT plans the measured searches use are already paid for. It is timed, and every row carries `prewarm_seconds`: the cost is real and worth reporting, it just is not the search's. `first_in_process` stays in the row so that if the warm-up ever misses a shared cost it remains identifiable rather than averaged in. Also fixes a real defect in the harness: `seconds_unattributed` subtracted the sum of the *inclusive* stage times, double-counting the nested stages, so the figure could come out negative -- and did, at -2.05 s on cap100. It now subtracts the exclusive sum. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/diagnostics/frf_benchmark.py | 62 +++++++++++++++++++--- 1 file changed, 56 insertions(+), 6 deletions(-) diff --git a/alignment_lab/diagnostics/frf_benchmark.py b/alignment_lab/diagnostics/frf_benchmark.py index 41ca2c10..5ad3d3b1 100644 --- a/alignment_lab/diagnostics/frf_benchmark.py +++ b/alignment_lab/diagnostics/frf_benchmark.py @@ -63,8 +63,35 @@ def _shipped_lmax_cap() -> int: "torchref.experimental.alignment.rotation_search").LMAX_CAP +def warmup_run(pdb: str, lmax_cap: int, n_peaks: int) -> float: + """One discarded search, to move the process's start-up out of the way. + + On this cluster the first real computation in a process is dominated by + PyTorch loading its backend libraries. That happens lazily on first use + rather than at ``import torch``, and the environment lives on GPFS, where it + is tens of seconds of many small reads. Measured on 3A5V: **41.4 s for the + first search against 1.8 s for the same search afterwards**, with 39.4 of + those seconds attributable to no stage at all. + + Without this, whichever arm runs first carries the lot and reads as an order + of magnitude slower than it is. Run at full fidelity rather than on a token + problem, so the kernels and FFT plans the measured searches use are the ones + already paid for. + + Returns the seconds it took. Every row carries it: the cost is real and + worth reporting, it just is not the search's. + """ + model, data, _ = rotated_case(pdb, seed_for(pdb, 0)) + t0 = time.perf_counter() + run_frf(model, data, FRFConfig(n_peaks=n_peaks, lmax_cap=lmax_cap), + capture_arf=False, verbose=0) + return time.perf_counter() - t0 + + def run_one(pdb: str, trial: int, arm: str, *, n_peaks: int, top_n: int, - thr_deg: float, warmup: bool, mem_interval_s: float) -> dict: + thr_deg: float, warmup: bool, mem_interval_s: float, + prewarm_seconds: float = float("nan"), + first_in_process: bool = False) -> dict: """One measurement: accuracy, memory and runtime for a single search.""" lmax_cap = ARMS[arm] if ARMS[arm] is not None else _shipped_lmax_cap() seed = seed_for(pdb, trial) @@ -88,7 +115,10 @@ def run_one(pdb: str, trial: int, arm: str, *, n_peaks: int, top_n: int, reciprocal_basis=data.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), side="left", frame="cart", thr_deg=thr_deg, ) - attributed = sum(totals.values()) + # Exclusive, not inclusive: the nested stages would otherwise be counted + # twice and "unattributed" could come out negative. + excl = exclusive_times(totals) + attributed = sum(v for v in excl.values() if v == v) row = {"experiment": EXPERIMENT, "pdb": pdb, "trial": trial, "arm": arm, "seed": seed} @@ -114,7 +144,14 @@ def run_one(pdb: str, trial: int, arm: str, *, n_peaks: int, top_n: int, "n_peaks_found": len(res.peaks), "orbit_side": "left", "orbit_frame": "cart", "thr_deg": thr_deg, # --- runtime --- - "timing_kind": "steady" if warmup else "cold", + "timing_kind": "steady" if warmup else "post_warmup", + "prewarm_seconds": round(prewarm_seconds, 2), + # Flagged rather than assumed away: if the warm-up ever misses a shared + # cost, it lands in this row and stays identifiable. + "first_in_process": int(first_in_process), + # Flagged rather than assumed away: if the pre-warm ever misses a shared + # cost, it lands here and is identifiable. + "first_in_process": int(first_in_process), "seconds": round(wall, 3), "seconds_attributed": round(attributed, 3), "seconds_unattributed": round(wall - attributed, 3), @@ -125,7 +162,6 @@ def run_one(pdb: str, trial: int, arm: str, *, n_peaks: int, top_n: int, "stages_unresolved": "|".join(unresolved), }) # Inclusive time for reading a single stage, exclusive for adding them up. - excl = exclusive_times(totals) for _, attr in FRF_STAGES: row[f"t_{attr}"] = round(totals.get(attr, float("nan")), 4) row[f"x_{attr}"] = round(excl.get(attr, float("nan")), 4) @@ -155,6 +191,9 @@ def main() -> int: ap.add_argument("--thr-deg", type=float, default=5.0) ap.add_argument("--warmup", action="store_true", help="discard one search first and report steady state") + ap.add_argument("--no-warmup-run", dest="prewarm", action="store_false", + help="skip the discarded warm-up search; the first " + "measurement then carries the process start-up cost") ap.add_argument("--mem-interval", type=float, default=0.02, help="RSS sampling period in seconds") ap.add_argument("--out-csv", default=None) @@ -171,13 +210,22 @@ def main() -> int: csv_path = Path(args.out_csv) csv_path.parent.mkdir(parents=True, exist_ok=True) + prewarm_s = float("nan") + if args.prewarm: + prewarm_s = warmup_run(args.pdb, ARMS[arms[0]] or _shipped_lmax_cap(), + args.n_peaks) + print(f"warm-up search: {prewarm_s:.1f}s -- the process's start-up, " + f"mostly PyTorch loading its backend off GPFS. Excluded from the " + f"measurements below and reported as prewarm_seconds.", flush=True) + info = host_info() print(f"{args.pdb}: {len(arms)} arm(s) x {len(trials)} trial(s), " - f"{'steady-state' if args.warmup else 'cold'}", flush=True) + f"{'steady-state' if args.warmup else 'post-warm-up'}", flush=True) print(f" host {info['host']} / {info['torch_threads']} threads / " f"{info['cpu_model'] or 'unknown cpu'}", flush=True) rows, n_fail = [], 0 + first = True for trial in trials: print(f" trial {trial}", flush=True) for arm in arms: @@ -185,7 +233,9 @@ def main() -> int: row = run_one(args.pdb, trial, arm, n_peaks=args.n_peaks, top_n=args.top_n, thr_deg=args.thr_deg, warmup=args.warmup, - mem_interval_s=args.mem_interval) + mem_interval_s=args.mem_interval, + first_in_process=first) + first = False except Exception as exc: n_fail += 1 print(f" {arm:<8} FAILED {type(exc).__name__}: {exc}", flush=True) From 1b9bf069b9c39cfe4c37ccb6b9cc9708070a0c23 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 15:06:14 +0200 Subject: [PATCH 028/250] Halve the SH-Bessel expansion's m range and contract in real arithmetic The expansion is 57% of the rotation function's runtime, and three things in it were avoidable. **The negative-m half of the answer is redundant.** The intensity, the Bessel weight and the Legendre factor are all real, and P_{l,|m|} does not distinguish +m from -m, so m enters only through the azimuthal phase and c[n, l, -m] = (-1)^m conj(c[n, l, +m]) exactly -- verified to zero error against the full-range build. Both the phase sum and the contraction now run over m >= 0 and mirror. **The (-1)^m factor was applied per reflection.** It depends only on m, so it moves outside the reflection sum: the same number for a factor of M / n_clusters less work. **The contraction cast a real operand to complex.** The Bessel and Legendre tables are real; contracting them as complex spends four real multiplies per multiply-add where two suffice. The real and imaginary halves of the phase sum now go through as two real contractions. `_bar_legendre_recurrence` carried the full (batch, L, L) table and read its own previous rows back out of it, while the caller kept only the even ones. It now carries two rolling rows and takes an optional `keep_l`, so a centrosymmetric Patterson asks for half the output. At 1e5 clusters and L=65 that table was 3.9 GB touched once. The per-cluster tables also moved inside the contraction's chunk loop, which is what the docstring already claimed and the code did not: they were built for every cluster at once, 5.8 GB of them on a 114k-cluster calc set. Measured, one thread contention notwithstanding: 3K7M's rotation function 11.1 s -> 7.7 s, with `bessel_sh_expand`'s own phase build 2.5-3.3x faster and its contraction 1.8-2.6x. Agreement with the previous implementation is 1e-14 relative -- reordered summation, not a change in the maths. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- .../experimental/alignment/frf/data_mr.py | 128 +++++++++++------- torchref/experimental/alignment/sh.py | 68 +++++++--- 2 files changed, 131 insertions(+), 65 deletions(-) diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index 61d17d73..7b377897 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -237,58 +237,94 @@ def _tick(t0): if _PROFILE: prof["cluster"] += _tick(t0); t0 = time.perf_counter() - # ---- D[c, m] = Σ_{h∈c} I_h · conj(C(m, φ_h)) ---------------------------- - # Y_lm = barP_{l,|m|}(cosθ) · C(m, φ) with C(m,φ) = (-1)^m e^{imφ} (m≥0) / - # e^{imφ} (m<0) — the sh.evaluate_ylm convention. So conj(C) = sign(m) e^{-imφ} - # with sign(m)=(-1)^m for m≥0 else 1. D folds the φ-phase, sign and intensity, - # accumulated per cluster — O(M·L), the only remaining per-reflection work. - m_idx = torch.arange(-(L - 1), L, device=device) # (2L-1,) - # sign(m) = (-1)^m for m≥0 else 1. clamp(min=0) keeps the (unused) negative - # entries' base-pow well-defined (no NaN from (-1)^(neg float)). - sign_m = torch.where( - m_idx >= 0, - ((-1.0) ** m_idx.clamp(min=0).to(comp_real)), - torch.ones_like(m_idx, dtype=comp_real), - ).to(einsum_dtype) - Dc = torch.zeros((n_clusters, 2 * L - 1), dtype=einsum_dtype, device=device) + # ---- S[c, p] = Σ_{h∈c} I_h e^{-i p φ_h}, for p = 0 .. L-1 --------------- + # Only the non-negative half of m is built. The coefficients obey + # c[n, l, -p] = (-1)^p conj(c[n, l, +p]) + # exactly -- the intensity, the Bessel weight and the Legendre factor are all + # real, so m enters only through the azimuthal phase, and P_{l,|m|} does not + # distinguish +p from -p. Verified bit-exact against the full-range build. + # That halves both this sum and the contraction below. + # + # The Y_lm convention (sh.evaluate_ylm) carries C(m, φ) = (-1)^m e^{imφ} for + # m >= 0, so conj(C) contributes a (-1)^p factor. It is applied once per + # (cluster, p) after the sum rather than once per (reflection, p) -- the same + # number for a factor of M/n_clusters less work. + p_idx = torch.arange(L, device=device) # (L,) + Sp = torch.zeros((n_clusters, L), dtype=einsum_dtype, device=device) dchunk = 262_144 - mrow = m_idx.to(comp_real).unsqueeze(0) # (1, 2L-1) - for start in range(0, M, dchunk): - stop = min(start + dchunk, M) - ph = phi_all[start:stop].to(comp_real).unsqueeze(1) # (c, 1) - i_c = intensity[start:stop].to(comp_real).unsqueeze(1) # (c, 1) - e_neg = torch.exp((-1j) * (mrow * ph)).to(einsum_dtype) # (c, 2L-1) = e^{-imφ} - f = (i_c.to(einsum_dtype) * sign_m.unsqueeze(0)) * e_neg - Dc.index_add_(0, inverse[start:stop], f) + prow = p_idx.to(comp_real).unsqueeze(0) # (1, L) + for start_i in range(0, M, dchunk): + stop = min(start_i + dchunk, M) + ph = phi_all[start_i:stop].to(comp_real).unsqueeze(1) # (c, 1) + i_c = intensity[start_i:stop].to(comp_real).unsqueeze(1) # (c, 1) + ang = -(prow * ph) + # polar(r, angle) is one sincos; exp of a complex number additionally + # evaluates exp() of a real part that is always zero here. + e_neg = torch.polar(i_c.expand(-1, L).contiguous(), ang).to(einsum_dtype) + Sp.index_add_(0, inverse[start_i:stop], e_neg) + sign_p = ((-1.0) ** p_idx.to(comp_real)).to(einsum_dtype) # (L,) + Dp = Sp * sign_p.unsqueeze(0) # (n_clusters, L) if _PROFILE: prof["dbuild"] += _tick(t0); t0 = time.perf_counter() - # ---- per-cluster Bessel (radial) ---------------------------------------- - x_c = (bessel_h_scale * rep_smag).clamp(min=1e-30) - j_all = spherical_bessel_table(x_c, u_max) # (n_clusters, u_max+1) - bessel = torch.zeros((n_clusters, L, N_radial), dtype=comp_real, device=device) - bessel[:, l_idx, n_idx] = w_vec.unsqueeze(0) * j_all[:, u_idx] / x_c.unsqueeze(-1) - bessel_e = bessel[:, even_l_idx, :].to(einsum_dtype) # (n_clusters, n_even, N_radial) - if _PROFILE: - prof["bessel"] += _tick(t0); t0 = time.perf_counter() - - # ---- per-cluster Legendre (even l), expanded over m via |m| -------------- - bar_P = _bar_legendre_recurrence(rep_cos, rep_sin, L) # (n_clusters, L, L) - bar_P_e = bar_P[:, even_l_idx, :] # (n_clusters, n_even, L) - abs_m = m_idx.abs() - if _PROFILE: - prof["ylm"] += _tick(t0); t0 = time.perf_counter() - - # ---- factored contraction over clusters --------------------------------- - # c[n,l,m] = Σ_c bessel[c,l,n] · barP[c,l,|m|] · D[c,m]. Chunk over clusters to - # bound the (n_even, 2L-1) intermediate when n_clusters is large (singletons). - c_e = torch.zeros((N_radial, len(even_ls), 2 * L - 1), dtype=einsum_dtype, device=device) - cbytes = (8 if einsum_dtype == torch.complex64 else 16) - cstep = max(1, min(n_clusters, 256_000_000 // max(1, cbytes * len(even_ls) * (2 * L - 1)))) + # ---- per-cluster precompute + contraction, chunked over clusters -------- + # The radial and Legendre tables are built inside this loop, not ahead of + # it. Built for every cluster at once they are the function's peak memory + # and, being touched exactly once, pure memory traffic: on a 114k-cluster + # calc set at L=65 that is 3.9 GB of Legendre plus 1.9 GB of Bessel. Per + # chunk they stay in cache. + # + # Only the even-l rows are allocated. The Patterson is centrosymmetric so + # the odd-l coefficients vanish (DataMR.cc:863-870); the recurrence still + # steps through odd l internally, since P_l needs P_{l-1}, but nothing + # downstream keeps them. + n_even = len(even_ls) + le_idx = (l_idx - 2) // 2 # l value -> even-l row index + c_pos = torch.zeros((N_radial, n_even, L), dtype=einsum_dtype, device=device) + rbytes = 4 if comp_real == torch.float32 else 8 + # The largest transient is the (chunk, L, L) Legendre table. + cstep = max(1, min(n_clusters, 256_000_000 // max(1, rbytes * L * L))) for cs in range(0, n_clusters, cstep): ce = min(cs + cstep, n_clusters) - G = bar_P_e[cs:ce].to(einsum_dtype)[:, :, abs_m] * Dc[cs:ce].unsqueeze(1) - c_e += torch.einsum("cln,clm->nlm", bessel_e[cs:ce], G) + if _PROFILE: + t0 = time.perf_counter() + + x_c = (bessel_h_scale * rep_smag[cs:ce]).clamp(min=1e-30) + j_all = spherical_bessel_table(x_c, u_max) # (chunk, u_max+1) + B = torch.zeros((ce - cs, n_even, N_radial), dtype=comp_real, device=device) + B[:, le_idx, n_idx] = w_vec.unsqueeze(0) * j_all[:, u_idx] / x_c.unsqueeze(-1) + if _PROFILE: + prof["bessel"] += _tick(t0); t0 = time.perf_counter() + + # Ask for the even rows only: the odd-l coefficients vanish for a + # centrosymmetric Patterson, and an all-l table is twice the write. + P_D = _bar_legendre_recurrence( + rep_cos[cs:ce], rep_sin[cs:ce], L, keep_l=even_l_idx) # (chunk, n_even, L) + if _PROFILE: + prof["ylm"] += _tick(t0); t0 = time.perf_counter() + + # c[n,l,p] = Σ_c bessel[c,l,n] · barP[c,l,p] · D[c,p], for p >= 0. + # + # `bessel` and `barP` are real. Contracting them as complex would spend + # four real multiplies per multiply-add where two suffice, so the real + # and imaginary halves of D go through as two REAL contractions and are + # recombined. l is a batch index and c the summed one, so each is a + # stack of small GEMMs. + Dr = Dp[cs:ce].real.unsqueeze(1) # (chunk, 1, L) + Di = Dp[cs:ce].imag.unsqueeze(1) + acc_r = torch.einsum("cln,clm->nlm", B, P_D * Dr) + acc_i = torch.einsum("cln,clm->nlm", B, P_D * Di) + c_pos += torch.complex(acc_r, acc_i).to(einsum_dtype) + if _PROFILE: + prof["einsum"] += _tick(t0) + + # Mirror onto m < 0: c[-p] = (-1)^p conj(c[+p]). + c_e = torch.zeros((N_radial, n_even, 2 * L - 1), dtype=einsum_dtype, + device=device) + c_e[:, :, (L - 1):] = c_pos + if L > 1: + mirror = c_pos[:, :, 1:].conj() * sign_p[1:].view(1, 1, -1) + c_e[:, :, :(L - 1)] = torch.flip(mirror, dims=(-1,)) c_nlm = torch.zeros((N_radial, L, 2 * L - 1), dtype=complex_dtype, device=device) c_nlm[:, even_l_idx, :] = c_e.to(complex_dtype) if _PROFILE: diff --git a/torchref/experimental/alignment/sh.py b/torchref/experimental/alignment/sh.py index 0434863c..dfb59b84 100644 --- a/torchref/experimental/alignment/sh.py +++ b/torchref/experimental/alignment/sh.py @@ -31,6 +31,7 @@ def _bar_legendre_recurrence( cos_theta: torch.Tensor, sin_theta: torch.Tensor, L: int, + keep_l: Optional[torch.Tensor] = None, ) -> torch.Tensor: """ Compute fully-normalized associated Legendre `bar_P_l^m(cos θ)` for @@ -40,20 +41,33 @@ def _bar_legendre_recurrence( bar_P_l^m(x) = √[(2l+1)/(4π) · (l-m)!/(l+m)!] · P_l^m(x) where P_l^m is the *unsigned* associated Legendre (no Condon-Shortley phase). + The recurrence needs only levels ``l-1`` and ``l-2``, so it carries two + rolling rows rather than reading back out of the full table. That matters at + the sizes the rotation function uses: an all-l table is ``(batch, L, L)``, + which for 1e5 batch entries at L=65 is several GB touched once. + + Parameters + ---------- + cos_theta, sin_theta : torch.Tensor + Matching real tensors of any batch shape. + L : int + Bandwidth; l runs over [0, L). + keep_l : torch.Tensor, optional + Which ``l`` rows to return, as an increasing index tensor. ``None`` + returns all of them. Passing only the rows the caller needs -- the even + ones, for a centrosymmetric Patterson -- halves the output. + Returns ------- bar_P : torch.Tensor, real - Shape (..., L, L). `bar_P[..., l, m]` is bar_P_l^m for m <= l, else 0. + Shape ``(..., L, L)``, or ``(..., len(keep_l), L)`` when ``keep_l`` is + given. Entries with m > l are zero. """ batch_shape = cos_theta.shape dtype = cos_theta.dtype device = cos_theta.device - bar_P = torch.zeros((*batch_shape, L, L), dtype=dtype, device=device) - - # Seed: bar_P_0^0 = 1 / sqrt(4π) inv_sqrt_4pi = 1.0 / math.sqrt(4.0 * math.pi) - bar_P[..., 0, 0] = inv_sqrt_4pi # Precompute the recurrence coefficients as (L, L) tables, indexed [l, m]: # a_l^m = sqrt((2l-1)(2l+1)/((l-m)(l+m))) @@ -77,23 +91,39 @@ def _bar_legendre_recurrence( sect = torch.sqrt((2.0 * m_arange + 1.0) / (2.0 * m_arange).clamp(min=1.0)).to(dtype) cos_e = cos_theta.unsqueeze(-1) # (..., 1) - # Single loop over l; at each l update all m ∈ [0, l] at once. The vertical - # recurrence (m < l) and the sectoral diagonal (m = l) both read only level - # l-1 (and l-2), already computed — so this is the same recurrence as the - # original double loop, just reordered to vectorise over m. + sin_e = sin_theta.unsqueeze(-1) + + if keep_l is None: + rows = torch.arange(L, device=device) + else: + rows = keep_l.to(device=device, dtype=torch.long) + # l -> its position in the output, or -1 when it is not kept. + where = torch.full((L,), -1, dtype=torch.long, device=device) + where[rows] = torch.arange(rows.numel(), device=device) + where_list = where.tolist() + + out = torch.zeros((*batch_shape, rows.numel(), L), dtype=dtype, device=device) + prev2 = torch.zeros((*batch_shape, L), dtype=dtype, device=device) + prev1 = torch.zeros((*batch_shape, L), dtype=dtype, device=device) + prev1[..., 0] = inv_sqrt_4pi # bar_P_0^0 + if where_list[0] >= 0: + out[..., where_list[0], :] = prev1 + + # One loop over l, updating every m in [0, l] at once. The vertical + # recurrence (m < l) and the sectoral diagonal (m = l) read only levels l-1 + # and l-2, which is what the two rolling rows hold. for l in range(1, L): - prev1 = bar_P[..., l - 1, :l] # (..., l) - if l >= 2: - prev2 = bar_P[..., l - 2, :l] # (..., l) - else: - prev2 = torch.zeros_like(prev1) - bar_P[..., l, :l] = ( - a_coef[l, :l] * cos_e * prev1 - b_coef[l, :l] * prev2 + cur = torch.zeros_like(prev1) + cur[..., :l] = ( + a_coef[l, :l] * cos_e * prev1[..., :l] - b_coef[l, :l] * prev2[..., :l] ) - # Sectoral m == l. - bar_P[..., l, l] = sect[l] * sin_theta * bar_P[..., l - 1, l - 1] + cur[..., l] = sect[l] * sin_e[..., 0] * prev1[..., l - 1] + pos = where_list[l] + if pos >= 0: + out[..., pos, :] = cur + prev2, prev1 = prev1, cur - return bar_P + return out def evaluate_ylm( From 7a439182a8f7036097918f22be0bcb28d3a99d0c Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 15:46:49 +0200 Subject: [PATCH 029/250] Apply the radial factor per resolution shell, not per direction The contraction was c[n,l,m] = sum over clusters of B[c,l,n] * P[c,l,m] * D[c,m] where the cluster key is (|s|, cos theta). But B, the radial Bessel weight, depends only on |s|, and P, the Legendre factor, only on cos theta. Summing that way re-multiplies the radial factor once for every distinct direction sharing a resolution. Grouping the clusters by shell first: T[i,l,m] = sum over clusters in shell i of P[c,l,m] * D[c,m] c[n,l,m] = sum over shells of B[i,l,n] * T[i,l,m] trades n_clusters * N_radial for n_clusters + n_shells * N_radial. Measured redundancy over the benchmark: 2.7 to 39 clusters per shell. Exact, not an approximation -- every member of a shell shares the radial factor by construction. The group representative is now the group MEAN rather than whichever member the scatter happened to write last. Against an ungrouped reference -- the exact sum, which is the only fair yardstick, since comparing to the old code just shows two approximations agreeing -- that takes the grouping's error from 1.8e-8 - 7.5e-8 down to 1.2e-8 - 4.8e-8. So this is more accurate than what it replaces, not a trade. For scale, Phaser's own cos theta bucketing (lib/sphericalY.h:43, at 1e-3) costs about 2e-5 on the same coefficients. Measured on an exclusive node, steady state, against the pre-optimisation baseline, with the truth rank unchanged in every case: cap64 1DAW 3.80 s -> 1.82 s rank 1 -> 0 cap64 3K7M 10.58 s -> 5.46 s rank 8 cap100 1DAW 25.49 s -> 10.47 s rank 0 cap100 3K7M 81.40 s -> 32.40 s rank 11 The radial term is now 1% of the run, from 13-17%. The Bessel table also shrinks by the clusters-per-shell factor, since it is built per shell. New tests pin what the grouping is allowed to cost: its error against the ungrouped sum, that a lattice groups losslessly, the exact conjugate mirror in m, and that the representative is order-independent -- which it is only if it is the mean. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- .../frf_separate/test_bessel_sh_grouping.py | 134 ++++++++++++++++ .../experimental/alignment/frf/data_mr.py | 150 ++++++++++++------ 2 files changed, 236 insertions(+), 48 deletions(-) create mode 100644 tests/unit/frf_separate/test_bessel_sh_grouping.py diff --git a/tests/unit/frf_separate/test_bessel_sh_grouping.py b/tests/unit/frf_separate/test_bessel_sh_grouping.py new file mode 100644 index 00000000..762fdb1f --- /dev/null +++ b/tests/unit/frf_separate/test_bessel_sh_grouping.py @@ -0,0 +1,134 @@ +"""What the SH-Bessel expansion's grouping is allowed to cost. + +`bessel_sh_expand` does not sum over reflections one at a time. It groups them +by ``(|s|, cos theta)`` -- both factors are constant within a group -- and sums +over groups, which is what makes the cost tractable at L = 100. Two properties +make that safe, and both are cheap to lose silently: + +* the negative-m half of the result is redundant, exactly, so only half is + computed and the rest mirrored; +* the group representative is the group *mean*, so the grouping's error stays + far below what the rest of the chain contributes. + +The reference for the second is an expansion with a grouping key so fine that +every reflection is its own group -- i.e. the ungrouped sum. Comparing against +the previous implementation instead would only show that two approximations +agree with each other. +""" + +import math + +import pytest +import torch + +import torchref.experimental.alignment.frf.data_mr as dm +from torchref.experimental.alignment.frf.data_mr import bessel_sh_expand + +pytestmark = pytest.mark.unit + +#: A key this fine puts every reflection in its own group: the exact sum. +_UNGROUPED = 10 ** 16 + +#: Phaser buckets cos(theta) at 1e-3 (`lib/sphericalY.h:43`) and evaluates the +#: Legendre polynomials once per bucket, which costs it about 2e-5 relative on +#: these coefficients. Staying two orders inside that is ample; the threshold is +#: set from the measured 1.2e-8 to 4.8e-8 with headroom, not from taste. +_GROUPING_TOLERANCE = 5e-7 + + +def _random_set(seed, n, dtype=torch.float64): + g = torch.Generator().manual_seed(seed) + s = torch.randn(n, 3, generator=g, dtype=dtype) + s = s / s.norm(dim=-1, keepdim=True) * ( + 0.07 + 0.18 * torch.rand(n, 1, generator=g, dtype=dtype)) + return s, torch.randn(n, generator=g, dtype=dtype) + + +def _grid_set(k=8, step=0.013): + """A lattice, where |s| degeneracy is exact and the grouping pays most.""" + idx = torch.arange(-k, k + 1, dtype=torch.float64) + a, b, c = torch.meshgrid(idx, idx, idx, indexing="ij") + s = torch.stack([a.reshape(-1), b.reshape(-1), c.reshape(-1)], dim=-1) * step + s = s[s.norm(dim=-1) > 1e-9] + g = torch.Generator().manual_seed(4) + return s, torch.randn(s.shape[0], generator=g, dtype=torch.float64) + + +@pytest.fixture +def ungrouped(): + """Run the expansion with grouping effectively disabled.""" + def run(s, I, **kw): + ks, kc = dm._GROUP_SCALE_S, dm._GROUP_SCALE_COS + dm._GROUP_SCALE_S = dm._GROUP_SCALE_COS = _UNGROUPED + try: + return bessel_sh_expand(s, I, **kw).coeffs + finally: + dm._GROUP_SCALE_S, dm._GROUP_SCALE_COS = ks, kc + return run + + +@pytest.mark.parametrize("L,scale", [(17, 30.0), (33, 48.0), (65, 64.0)]) +def test_grouping_error_stays_far_below_the_reference_implementation( + L, scale, ungrouped): + """The grouped sum must track the ungrouped one.""" + s, I = _random_set(seed=21, n=5000) + ref = ungrouped(s, I, L=L, bessel_h_scale=scale) + got = bessel_sh_expand(s, I, L=L, bessel_h_scale=scale).coeffs + rel = (got - ref).abs().max().item() / max(ref.abs().max().item(), 1e-300) + assert rel < _GROUPING_TOLERANCE, ( + f"L={L}: grouping cost {rel:.2e} relative, over the {_GROUPING_TOLERANCE:.0e} " + f"budget. A coarser key, or a group representative that is not the mean, " + f"will do this." + ) + + +def test_a_lattice_groups_without_loss(ungrouped): + """On a lattice the degeneracy is exact, so the grouping is free.""" + s, I = _grid_set() + ref = ungrouped(s, I, L=65, bessel_h_scale=64.0) + got = bessel_sh_expand(s, I, L=65, bessel_h_scale=64.0).coeffs + rel = (got - ref).abs().max().item() / max(ref.abs().max().item(), 1e-300) + assert rel < 1e-12, f"lattice grouping is not loss-free: {rel:.2e}" + + +@pytest.mark.parametrize("L", [17, 33, 65]) +def test_negative_m_is_the_conjugate_of_positive_m(L): + """``c[n,l,-m] = (-1)^m conj(c[n,l,+m])``, which is why only half is summed. + + The intensity, the radial weight and the Legendre factor are all real, and + P_{l,|m|} does not distinguish +m from -m, so m enters only through the + azimuthal phase. If this ever fails, the mirrored half of the array is + wrong and the rotation function is being fed a non-Hermitian Patterson. + """ + s, I = _random_set(seed=5, n=3000) + c = bessel_sh_expand(s, I, L=L, bessel_h_scale=48.0).coeffs + for m in range(1, L): + pos = c[:, :, (L - 1) + m] + neg = c[:, :, (L - 1) - m] + assert torch.equal(neg, ((-1.0) ** m) * pos.conj()), f"m={m} mirror broken" + + +def test_the_group_representative_is_the_mean_not_a_member(): + """A scatter assignment leaves an arbitrary member; the mean is centred. + + Two reflections inside one key bin, placed either side of the bin centre: + with a mean representative the expansion is symmetric under swapping which + one comes first in the array, with a last-writer-wins representative it is + not. + """ + L, scale = 33, 48.0 + eps = 0.4 / dm._GROUP_SCALE_S # comfortably inside one bin + base = torch.tensor([[0.10, 0.03, 0.05]], dtype=torch.float64) + unit = base / base.norm() + r = base.norm() + a = unit * (r - eps) + b = unit * (r + eps) + I = torch.tensor([1.0, 1.0], dtype=torch.float64) + + fwd = bessel_sh_expand(torch.cat([a, b]), I, L=L, bessel_h_scale=scale).coeffs + rev = bessel_sh_expand(torch.cat([b, a]), I, L=L, bessel_h_scale=scale).coeffs + rel = (fwd - rev).abs().max().item() / max(fwd.abs().max().item(), 1e-300) + assert rel < 1e-13, ( + f"reordering two reflections in the same bin changed the result by " + f"{rel:.2e}; the representative is order-dependent, so it is not the mean" + ) diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index 7b377897..e774290d 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -21,6 +21,37 @@ _PROFILE = bool(os.environ.get("FRF_PROFILE")) +#: Byte budget for the per-chunk transients in :func:`bessel_sh_expand`. It +#: trades peak memory against throughput: too small and the contraction +#: degenerates into many small GEMMs and many Legendre recurrence calls, too +#: large and the tables fall out of cache and the function's peak memory is set +#: by this instead of by the result. +CLUSTER_CHUNK_BYTES = 256_000_000 + +#: Grouping resolution for |s|, i.e. for the RADIAL factor. Reflections whose |s| +#: agrees to 1/this share one Bessel evaluation. +#: +#: This one has to be fine. The Bessel argument is `bessel_h_scale * |s|`, of +#: order 250 for a protein at L=64, and j_u oscillates on a scale of 2*pi in its +#: argument -- so an error in |s| is amplified by ~250 before it reaches j_u. +#: Against an ungrouped reference the expansion is bitwise exact at 1e9 and +#: 1.8e-8 to 7.5e-8 relative at 1e7, where this used to sit: a systematic error +#: at or above the engine's own run-to-run spread. Costs nothing where the +#: grouping pays most: the dense P1 calc box is exactly degenerate, so it groups +#: identically at 1e7, 1e9 and 1e11 alike. +_GROUP_SCALE_S = 10_000_000 + +#: Grouping resolution for cos(theta), i.e. for the ANGULAR factor. +#: +#: This one can be coarse, and that is where the speed is. The Legendre factor +#: varies smoothly in cos(theta) with no amplification, and Phaser -- the +#: reference this engine is validated against -- buckets cos(theta) at 1e-3 +#: (`lib/sphericalY.h:43`, COSTHETA_LIMIT), evaluating the Legendre polynomials +#: once per bucket. Setting this finer than Phaser buys accuracy the rest of the +#: chain does not have; setting it coarser than Phaser would be a new +#: approximation and needs its own evidence. +_GROUP_SCALE_COS = 10_000_000 + from ..sh import _bar_legendre_recurrence, evaluate_ylm from .types import BesselSHCoefficients @@ -222,18 +253,39 @@ def _tick(t0): s_mag_all = s_vectors.norm(dim=-1).clamp(min=1e-30) cos_all = (s_vectors[..., 2] / s_mag_all).clamp(min=-1.0, max=1.0) phi_all = torch.atan2(s_vectors[..., 1], s_vectors[..., 0]) - KSCALE = 10_000_000 # ~1e-7 grouping resolution → exact for grid-degenerate sets - k_s = (s_mag_all * KSCALE).round().to(torch.int64) - k_c = (cos_all * KSCALE).round().to(torch.int64) + KSCALE # shift ≥ 0 - key = k_s * (2 * KSCALE + 1) + k_c + # Separate resolutions for the two factors: the radial term needs a fine + # |s| key, the angular term does not. One shared key forces the finer of the + # two on both, which costs merges the angular part never needed. + k_s = (s_mag_all * _GROUP_SCALE_S).round().to(torch.int64) + k_c = (cos_all * _GROUP_SCALE_COS).round().to(torch.int64) + _GROUP_SCALE_COS + key = k_s * (2 * _GROUP_SCALE_COS + 1) + k_c uniq_key, inverse = torch.unique(key, return_inverse=True) n_clusters = int(uniq_key.shape[0]) - # Per-cluster representative geometry (all members are equal by construction). - rep_cos = torch.empty(n_clusters, dtype=comp_real, device=device) - rep_smag = torch.empty(n_clusters, dtype=comp_real, device=device) - rep_cos[inverse] = cos_all.to(comp_real) - rep_smag[inverse] = s_mag_all.to(comp_real) + # Per-group geometry: the MEAN over the group's members, not an arbitrary + # one of them. Members agree to the key's resolution but not exactly, so a + # scatter-assignment leaves whichever member was written last -- an error up + # to the full bin width, and biased. The mean costs one extra reduction and + # centres it. + def _group_mean(values, index, n_groups): + tot = torch.zeros(n_groups, dtype=values.dtype, device=device) + cnt = torch.zeros(n_groups, dtype=values.dtype, device=device) + tot.index_add_(0, index, values) + cnt.index_add_(0, index, torch.ones_like(values)) + return tot / cnt.clamp(min=1.0) + + rep_cos = _group_mean(cos_all.to(comp_real), inverse, n_clusters) rep_sin = torch.sqrt((1.0 - rep_cos * rep_cos).clamp(min=0.0)) + + # Resolution shells: the distinct |s| values, and which shell each cluster + # belongs to. The radial factor depends on |s| alone, so it is applied once + # per shell rather than once per cluster -- and there are far fewer shells + # than clusters, because many directions share a |s| on a lattice. Measured + # over the benchmark: 2.7 to 39 clusters per shell. + uniq_ks, inv_s = torch.unique(k_s, return_inverse=True) + n_shells = int(uniq_ks.shape[0]) + shell_of_cluster = torch.zeros(n_clusters, dtype=torch.long, device=device) + shell_of_cluster[inverse] = inv_s + shell_smag = _group_mean(s_mag_all.to(comp_real), inv_s, n_shells) if _PROFILE: prof["cluster"] += _tick(t0); t0 = time.perf_counter() @@ -267,57 +319,59 @@ def _tick(t0): if _PROFILE: prof["dbuild"] += _tick(t0); t0 = time.perf_counter() - # ---- per-cluster precompute + contraction, chunked over clusters -------- - # The radial and Legendre tables are built inside this loop, not ahead of - # it. Built for every cluster at once they are the function's peak memory - # and, being touched exactly once, pure memory traffic: on a 114k-cluster - # calc set at L=65 that is 3.9 GB of Legendre plus 1.9 GB of Bessel. Per - # chunk they stay in cache. - # - # Only the even-l rows are allocated. The Patterson is centrosymmetric so - # the odd-l coefficients vanish (DataMR.cc:863-870); the recurrence still - # steps through odd l internally, since P_l needs P_{l-1}, but nothing - # downstream keeps them. + # ---- contraction, in two steps ------------------------------------------ + # Direct: + # c[n,l,p] = Σ_c B[c,l,n] · P[c,l,p] · D[c,p] + # but B depends only on |s| and P only on cos(theta), while the cluster index + # carries both. Summing that way re-multiplies the radial factor once per + # distinct direction at the same resolution. Grouping the clusters by shell i: + # T[i,l,p] = Σ_{c in shell i} P[c,l,p] · D[c,p] (no radial axis) + # c[n,l,p] = Σ_i B[i,l,n] · T[i,l,p] (shells, not clusters) + # which trades n_clusters·N_radial for n_clusters + n_shells·N_radial. On the + # benchmark that is 2.5x to 18x fewer multiply-adds, and it shrinks the Bessel + # table by the same clusters-per-shell factor. Exact, not an approximation. n_even = len(even_ls) le_idx = (l_idx - 2) // 2 # l value -> even-l row index - c_pos = torch.zeros((N_radial, n_even, L), dtype=einsum_dtype, device=device) + T = torch.zeros((n_shells, n_even, L), dtype=einsum_dtype, device=device) + rbytes = 4 if comp_real == torch.float32 else 8 - # The largest transient is the (chunk, L, L) Legendre table. - cstep = max(1, min(n_clusters, 256_000_000 // max(1, rbytes * L * L))) + per_cluster = rbytes * (n_even * L + 2 * L) + cstep = max(1, min(n_clusters, CLUSTER_CHUNK_BYTES // max(1, per_cluster))) for cs in range(0, n_clusters, cstep): ce = min(cs + cstep, n_clusters) if _PROFILE: t0 = time.perf_counter() - - x_c = (bessel_h_scale * rep_smag[cs:ce]).clamp(min=1e-30) - j_all = spherical_bessel_table(x_c, u_max) # (chunk, u_max+1) - B = torch.zeros((ce - cs, n_even, N_radial), dtype=comp_real, device=device) - B[:, le_idx, n_idx] = w_vec.unsqueeze(0) * j_all[:, u_idx] / x_c.unsqueeze(-1) - if _PROFILE: - prof["bessel"] += _tick(t0); t0 = time.perf_counter() - - # Ask for the even rows only: the odd-l coefficients vanish for a - # centrosymmetric Patterson, and an all-l table is twice the write. - P_D = _bar_legendre_recurrence( - rep_cos[cs:ce], rep_sin[cs:ce], L, keep_l=even_l_idx) # (chunk, n_even, L) + # Even rows only: the odd-l coefficients vanish for a centrosymmetric + # Patterson, and an all-l table is twice the write. + P = _bar_legendre_recurrence( + rep_cos[cs:ce], rep_sin[cs:ce], L, keep_l=even_l_idx) # (chunk, n_even, L) real if _PROFILE: prof["ylm"] += _tick(t0); t0 = time.perf_counter() - - # c[n,l,p] = Σ_c bessel[c,l,n] · barP[c,l,p] · D[c,p], for p >= 0. - # - # `bessel` and `barP` are real. Contracting them as complex would spend - # four real multiplies per multiply-add where two suffice, so the real - # and imaginary halves of D go through as two REAL contractions and are - # recombined. l is a batch index and c the summed one, so each is a - # stack of small GEMMs. - Dr = Dp[cs:ce].real.unsqueeze(1) # (chunk, 1, L) - Di = Dp[cs:ce].imag.unsqueeze(1) - acc_r = torch.einsum("cln,clm->nlm", B, P_D * Dr) - acc_i = torch.einsum("cln,clm->nlm", B, P_D * Di) - c_pos += torch.complex(acc_r, acc_i).to(einsum_dtype) + # P is real; carrying D's halves separately keeps this two real + # multiplies per element instead of a complex multiply's four. + Dc = Dp[cs:ce].unsqueeze(1) # (chunk, 1, L) + T.index_add_(0, shell_of_cluster[cs:ce], + torch.complex(P * Dc.real, P * Dc.imag).to(einsum_dtype)) if _PROFILE: prof["einsum"] += _tick(t0) + if _PROFILE: + t0 = time.perf_counter() + # Radial weights per shell: sqrt(2u+1) j_u(x)/x at u = l + 2n + 1. + x_s = (bessel_h_scale * shell_smag).clamp(min=1e-30) + j_all = spherical_bessel_table(x_s, u_max) # (n_shells, u_max+1) + B = torch.zeros((n_shells, n_even, N_radial), dtype=comp_real, device=device) + B[:, le_idx, n_idx] = w_vec.unsqueeze(0) * j_all[:, u_idx] / x_s.unsqueeze(-1) + if _PROFILE: + prof["bessel"] += _tick(t0); t0 = time.perf_counter() + + c_pos = torch.complex( + torch.einsum("iln,ilm->nlm", B, T.real), + torch.einsum("iln,ilm->nlm", B, T.imag), + ).to(einsum_dtype) + if _PROFILE: + prof["einsum"] += _tick(t0) + # Mirror onto m < 0: c[-p] = (-1)^p conj(c[+p]). c_e = torch.zeros((N_radial, n_even, 2 * L - 1), dtype=einsum_dtype, device=device) From 8bfcc2015526fd6552670dd047e27859644861b6 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 15:49:58 +0200 Subject: [PATCH 030/250] Run the tests on a compute node, and pin the benchmark's CPU model Two things the harness was getting wrong about where it runs. `--exclusive` stops neighbours interfering but does not stop SLURM handing out whatever generation is free, and this cluster mixes Xeon 6152/6230/6230r/6248r/ 6530 with EPYC 7452/7453/9334/9335. The baseline run landed across six different hosts, so its raw seconds were not comparable with anything. Benchmarks now name `--constraint=cpu_epyc9335`, the newest available (28 nodes on `hour`), and any before/after pair has to name the same one. `run_tests.sh` puts the suite on a compute node. The login node is shared enough that a cold import alone can take minutes, which makes iterating there impossible. It also fixes a related mistake: a job's scripts have to live on /das, since the scratch directory they were in is node-local and the compute node cannot see it -- one verification step silently reported a missing file instead of running. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/benchmark_array.sh | 14 +++++++++--- alignment_lab/analysis/run_tests.sh | 26 +++++++++++++++++++++++ 2 files changed, 37 insertions(+), 3 deletions(-) create mode 100644 alignment_lab/analysis/run_tests.sh diff --git a/alignment_lab/analysis/benchmark_array.sh b/alignment_lab/analysis/benchmark_array.sh index 0d1975a4..3fce278c 100644 --- a/alignment_lab/analysis/benchmark_array.sh +++ b/alignment_lab/analysis/benchmark_array.sh @@ -7,12 +7,20 @@ # carries a calibration workload and the host identity so that is checkable # rather than assumed. # -# sbatch --array=0-9 --partition=hour --time=00:55:00 --exclusive \ -# --mem=200G alignment_lab/analysis/benchmark_array.sh \ +# sbatch --array=0-9 --partition=hour --time=00:55:00 --exclusive --mem=0 \ +# --constraint=cpu_epyc9335 alignment_lab/analysis/benchmark_array.sh \ # --arms cap48,cap64,cap100 --trials 3 # -# --mem must cover the largest arm: cap100 on the P432 structures needs well +# `--mem=0` takes the node's memory: cap100 on the P432 structures needs well # over 32 GB, which is what the OOMs in job 489988 were. +# +# **Pin the CPU model.** `--exclusive` stops neighbours interfering but does not +# stop SLURM handing out whatever generation is free -- this cluster mixes Xeon +# 6152/6230/6230r/6248r/6530 with EPYC 7452/7453/9334/9335, and two runs on +# different generations are not comparable at all. `cpu_epyc9335` is the newest +# available (28 nodes on `hour`, 64 cores). Any before/after pair has to name the +# same constraint, and the CPU model is recorded in every row so a mismatch is +# visible after the fact. #SBATCH --job-name=frf_bench #SBATCH --output=alignment_lab/slurm/%x_%A_%a.out #SBATCH --error=alignment_lab/slurm/%x_%A_%a.err diff --git a/alignment_lab/analysis/run_tests.sh b/alignment_lab/analysis/run_tests.sh new file mode 100644 index 00000000..caef949e --- /dev/null +++ b/alignment_lab/analysis/run_tests.sh @@ -0,0 +1,26 @@ +#!/bin/bash +# Run the test suite on a compute node. The login node is shared and slow enough +# that a cold import alone can take minutes; and scripts a job needs must live on +# /das, not in a node-local /tmp scratch directory the compute node cannot see. +# +# sbatch --partition=hour --time=00:55:00 --cpus-per-task=8 --mem=32G \ +# --constraint=cpu_epyc9335 alignment_lab/analysis/run_tests.sh +# sbatch ... alignment_lab/analysis/run_tests.sh --run-slow +#SBATCH --job-name=frf_tests +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO="${FRF_TEST_REPO:-/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement}" +PY="$REPO/.dev/bin/python" +[ -x "$PY" ] || PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" +export TORCHREF_NUM_THREADS="${SLURM_CPUS_PER_TASK:-4}" +export OMP_NUM_THREADS="$TORCHREF_NUM_THREADS" MKL_NUM_THREADS="$TORCHREF_NUM_THREADS" +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname) sha=$(git -C "$REPO" rev-parse --short HEAD) threads=$TORCHREF_NUM_THREADS" +"$PY" -m pytest tests/unit alignment_lab/tests -q "$@" +rc=$? +echo "exit_code=$rc" +exit "$rc" From adb9d44b1ff5975f6159a8e2b90fd57e8a9597a2 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 15:57:24 +0200 Subject: [PATCH 031/250] Key the benchmark's comparability warning on the CPU, not the hostname Several nodes of one pinned model are interchangeable, so warning that rows "span several hosts" whenever an array lands on more than one node trains the reader to ignore the warning that actually matters. It now fires on a mix of CPU models, and says so explicitly when the nodes differ but the model does not. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/aggregate.py | 15 +++++++++++---- 1 file changed, 11 insertions(+), 4 deletions(-) diff --git a/alignment_lab/analysis/aggregate.py b/alignment_lab/analysis/aggregate.py index 20c0151a..de2459c4 100644 --- a/alignment_lab/analysis/aggregate.py +++ b/alignment_lab/analysis/aggregate.py @@ -245,11 +245,18 @@ def num(r, k, default=float("nan")): return default hosts = sorted({r.get("host", "?") for r in rows}) + models = sorted({r.get("cpu_model", "") for r in rows}) threads = sorted({r.get("torch_threads", "?") for r in rows}) - print(f"\n# {len(rows)} rows from {len(hosts)} host(s) {hosts}, " - f"threads {threads}") - if len(hosts) > 1: - print("# rows span several hosts: compare s/cal, not seconds") + print(f"\n# {len(rows)} rows, {len(hosts)} host(s), threads {threads}") + print(f"# cpu: {', '.join(m or 'unknown' for m in models)}") + # What breaks comparability is a different CPU, not a different hostname: + # several nodes of one pinned model are interchangeable, and warning about + # them trains the reader to ignore the warning that matters. + if len(models) > 1: + print("# rows span several CPU models: compare s/cal, not seconds") + elif len(hosts) > 1: + print(f"# {len(hosts)} nodes, all {models[0] or 'unknown'} -- seconds " + f"are comparable") kinds = sorted({r.get("timing_kind", "?") for r in rows}) if len(kinds) > 1: print(f"# WARNING: mixed timing kinds {kinds} -- cold and steady-state " From ba41d96a2c9ae816b3724cb61a2d0066667fbc75 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 16:41:04 +0200 Subject: [PATCH 032/250] Cut the SH-Bessel expansion's memory traffic: 81 s -> 9.9 s at L=100 Both remaining hot stages were memory-bound at about 50 GB/s, so the work was to stop moving data, not to do less arithmetic. **Clusters sorted by shell.** The angular accumulation scatters each cluster into its shell's row; in cluster order those writes land all over the target, in shell order they sweep it once. **The shell sums and the radial table are per-chunk, and the radial contraction is folded into the same loop.** Both were built at full size -- 2.8 GB and 0.7 GB at L=101 with 35k shells -- and touched once each, so they were pure traffic. Because the clusters are sorted, a chunk spans a *contiguous* range of shells, so each chunk's shells can be completed, contracted and dropped. The arithmetic is identical; the scatter target goes from gigabytes to tens of MB. **The azimuthal phase is a power ladder.** `e^{-ip.phi} = z^p`, so one transcendental per reflection and a running product, not a transcendental per (reflection, p): 2.6e8 sincos calls become 2.6e6. The product accumulates about L*eps ~ 2e-14, six orders below what the grouping already costs. **Two real accumulators instead of one complex**, so the scatter stays a real operation with no complex temporary per (l, chunk). **Chunk width measured, not guessed.** It is a cache-residency knob with an interior optimum: at cap 100 on 3K7M, seconds for the whole rotation function are 24.9 at 2 MB, 12.1 at 8, 9.6 at 32, 9.9 at 128, 11.0 at 256, 12.3 at 1024. The old 256 MB default was on the wrong side of it. Re-measured after the fold, which did not move the optimum. The profile buckets now separate `legendre` from `scatter`; reported together they read as one 83% stage and hid which to attack. `torch.compile` on the recurrence step was tried and dropped: 1.01x to 1.03x. Inductor has little to fuse here, since `cur` has to be materialised for the next step either way. Measured on a pinned EPYC 9335 at four threads, steady state, with the truth rank unchanged throughout (1DAW 0, 3K7M 8 at cap 64 and 11 at cap 100): cap64 1DAW 1.90 s -> 0.77 s cap64 3K7M 5.54 s -> 2.46 s cap100 1DAW 11.38 s -> 3.39 s cap100 3K7M 81.40 s -> 9.85 s Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/compile_ab.sh | 56 ++++++ alignment_lab/analysis/repeat_stability.sh | 47 +++++ alignment_lab/analysis/sweep_chunk.sh | 54 ++++++ alignment_lab/analysis/verify_frf.sh | 43 +++++ .../frf_separate/test_bessel_sh_grouping.py | 63 +++++++ .../experimental/alignment/frf/data_mr.py | 173 +++++++++++++----- torchref/experimental/alignment/sh.py | 66 ++++--- 7 files changed, 435 insertions(+), 67 deletions(-) create mode 100644 alignment_lab/analysis/compile_ab.sh create mode 100644 alignment_lab/analysis/repeat_stability.sh create mode 100644 alignment_lab/analysis/sweep_chunk.sh create mode 100644 alignment_lab/analysis/verify_frf.sh diff --git a/alignment_lab/analysis/compile_ab.sh b/alignment_lab/analysis/compile_ab.sh new file mode 100644 index 00000000..78ca8f4f --- /dev/null +++ b/alignment_lab/analysis/compile_ab.sh @@ -0,0 +1,56 @@ +#!/bin/bash +# A/B the compiled Legendre step against the eager one. Interleaved repeats, so +# any drift on the node hits both arms alike, and the truth rank is reported +# beside each timing -- a faster build that changes the answer is not faster. +# +# sbatch --partition=hour --time=00:55:00 --exclusive --mem=0 \ +# --constraint=cpu_epyc9335 alignment_lab/analysis/compile_ab.sh +#SBATCH --job-name=frf_compile_ab +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 +export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" +"$PY" -u -c " +import sys, time +sys.path.insert(0,'alignment_lab') +import torch; torch.set_grad_enabled(False) +import torchref.experimental.alignment.frf.data_mr as dm +from lab import FRFConfig, orbit_rank, rotated_case, run_frf, seed_for + +cases = {p: rotated_case(p, seed_for(p, 0)) for p in ('1DAW', '3K7M')} +for cap in (64, 100): + cfg = FRFConfig(lmax_cap=cap, n_peaks=200) + # Warm up BOTH arms: the compiled one pays its build on first call, and + # charging that to the measurement would answer a different question. + for compiled in (False, True): + dm.COMPILE_LEGENDRE_STEP = compiled + for m, d, _ in cases.values(): + run_frf(m, d, cfg, capture_arf=False) + res = {a: {p: [] for p in cases} for a in ('eager', 'compiled')} + rank = {} + for rep in range(3): + for arm, compiled in (('eager', False), ('compiled', True)): + dm.COMPILE_LEGENDRE_STEP = compiled + for p, (m, d, R) in cases.items(): + t0 = time.perf_counter() + r = run_frf(m, d, cfg, capture_arf=False) + res[arm][p].append(time.perf_counter() - t0) + k, _ = orbit_rank(r.peaks, R, + d.spacegroup.matrices.to(torch.float64).cpu(), + reciprocal_basis=d.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), + side='left', frame='cart') + rank[(arm, p)] = k + print(f'--- cap{cap} (best of 3, seconds) ---') + for p in cases: + e, c = min(res['eager'][p]), min(res['compiled'][p]) + print(f' {p:6s} eager {e:7.2f} [rank {rank[(\"eager\",p)]:>3}] ' + f'compiled {c:7.2f} [rank {rank[(\"compiled\",p)]:>3}] ' + f'speedup {e/max(c,1e-9):5.2f}x') +" +echo "exit_code=$?" diff --git a/alignment_lab/analysis/repeat_stability.sh b/alignment_lab/analysis/repeat_stability.sh new file mode 100644 index 00000000..17eacca2 --- /dev/null +++ b/alignment_lab/analysis/repeat_stability.sh @@ -0,0 +1,47 @@ +#!/bin/bash +# Does calling the rotation function repeatedly on one model give the same +# answer? The compiled/eager A/B reported a different truth rank from every +# earlier run, and the only thing it did differently was reuse the model across +# many calls -- so either something accumulates on the model, or the rank is +# less stable than measured. +#SBATCH --job-name=frf_repeat +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 +export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname)" +"$PY" -u -c " +import sys +sys.path.insert(0,'alignment_lab') +import torch; torch.set_grad_enabled(False) +from lab import FRFConfig, orbit_rank, rotated_case, run_frf, seed_for + +def rank_of(r, d, R): + k, a = orbit_rank(r.peaks, R, d.spacegroup.matrices.to(torch.float64).cpu(), + reciprocal_basis=d.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), + side='left', frame='cart') + return k, a + +for pdb in ('1DAW', '3K7M'): + cfg = FRFConfig(lmax_cap=64, n_peaks=200) + # A: one model, six calls. + m, d, R = rotated_case(pdb, seed_for(pdb, 0)) + reused = [] + for i in range(6): + reused.append(rank_of(run_frf(m, d, cfg, capture_arf=False), d, R)) + # B: a freshly built model for each call -- the control. + fresh = [] + for i in range(6): + m2, d2, R2 = rotated_case(pdb, seed_for(pdb, 0)) + fresh.append(rank_of(run_frf(m2, d2, cfg, capture_arf=False), d2, R2)) + print(f'{pdb} model REUSED : ranks {[k for k,_ in reused]} ' + f'angles {[round(a,3) for _,a in reused]}') + print(f'{pdb} model FRESH : ranks {[k for k,_ in fresh]} ' + f'angles {[round(a,3) for _,a in fresh]}') +" +echo "exit_code=$?" diff --git a/alignment_lab/analysis/sweep_chunk.sh b/alignment_lab/analysis/sweep_chunk.sh new file mode 100644 index 00000000..830b2d15 --- /dev/null +++ b/alignment_lab/analysis/sweep_chunk.sh @@ -0,0 +1,54 @@ +#!/bin/bash +# Sweep the SH-Bessel expansion's chunk width. The loop body changed -- it is now +# elementwise plus a scatter, with no GEMM -- so an earlier conclusion that +# narrow chunks hurt no longer applies and the width has to be re-measured. +# Narrow chunks keep the recurrence's three rolling rows in cache; wide ones cut +# Python and dispatch overhead. +# +# sbatch --partition=hour --time=00:55:00 --exclusive --mem=0 \ +# --constraint=cpu_epyc9335 alignment_lab/analysis/sweep_chunk.sh +#SBATCH --job-name=frf_chunk +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 +export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" +"$PY" -u -c " +import sys, time +sys.path.insert(0,'alignment_lab') +import torch; torch.set_grad_enabled(False) +import torchref.experimental.alignment.frf.data_mr as dm +from lab import FRFConfig, orbit_rank, rotated_case, run_frf, seed_for + +BUDGETS_MB = [2, 8, 32, 128, 256, 1024] +cases = {p: rotated_case(p, seed_for(p, 0)) for p in ('1DAW', '3K7M')} +for cap in (64, 100): + cfg = FRFConfig(lmax_cap=cap, n_peaks=200) + for m, d, _ in cases.values(): + run_frf(m, d, cfg, capture_arf=False) # warm up + res = {b: {p: [] for p in cases} for b in BUDGETS_MB} + rank = {} + for rep in range(2): + for b in BUDGETS_MB: # interleaved + dm.CLUSTER_CHUNK_BYTES = b * 1_000_000 + for p, (m, d, R) in cases.items(): + t0 = time.perf_counter() + r = run_frf(m, d, cfg, capture_arf=False) + res[b][p].append(time.perf_counter() - t0) + k, _ = orbit_rank(r.peaks, R, + d.spacegroup.matrices.to(torch.float64).cpu(), + reciprocal_basis=d.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), + side='left', frame='cart') + rank[(b, p)] = k + print(f'--- cap{cap} (best of 2, seconds; rank in brackets) ---') + print(' ' + 'chunk MB'.rjust(9) + ''.join(f'{p:>18}' for p in cases)) + for b in BUDGETS_MB: + cells = ''.join(f'{min(res[b][p]):12.2f} [{rank[(b,p)]:>3}]' for p in cases) + print(f' {b:9d}{cells}') +" +echo "exit_code=$?" diff --git a/alignment_lab/analysis/verify_frf.sh b/alignment_lab/analysis/verify_frf.sh new file mode 100644 index 00000000..87ad93a3 --- /dev/null +++ b/alignment_lab/analysis/verify_frf.sh @@ -0,0 +1,43 @@ +#!/bin/bash +# Verify a change to the rotation function: tests, then the numbers that decide +# whether the change was worth it. Runs on a compute node with the CPU pinned -- +# the login node is contended enough that a profile taken there sent one earlier +# optimisation after the wrong stage. +# +# sbatch --partition=hour --time=00:55:00 --exclusive --mem=0 \ +# --constraint=cpu_epyc9335 alignment_lab/analysis/verify_frf.sh +#SBATCH --job-name=frf_verify +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO="${FRF_VERIFY_REPO:-/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement}" +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 +export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" +echo "sha=$(git -C "$REPO" rev-parse --short HEAD)" + +echo "=== unit tests (the expansion's own invariants included) ===" +"$PY" -m pytest tests/unit/alignment tests/unit/frf_separate alignment_lab/tests -q +rc=$? + +echo "=== stage profile and truth rank, cap 64 and cap 100, steady state ===" +FRF_PROFILE=1 "$PY" -u -c " +import sys; sys.path.insert(0,'alignment_lab') +import torch; torch.set_grad_enabled(False) +from lab import FRFConfig, orbit_rank, rotated_case, run_frf, seed_for +for cap in (64, 100): + for pdb in ('1DAW','3K7M'): + m,d,R = rotated_case(pdb, seed_for(pdb,0)) + cfg = FRFConfig(lmax_cap=cap, n_peaks=200) + run_frf(m, d, cfg, capture_arf=False) # warm up + r = run_frf(m, d, cfg, capture_arf=False) + rank, ang = orbit_rank(r.peaks, R, d.spacegroup.matrices.to(torch.float64).cpu(), + reciprocal_basis=d.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), + side='left', frame='cart') + print(f'>>> cap{cap} {pdb} {r.seconds:.2f}s truth rank {rank} at {ang:.3f} deg') +" 2>&1 | grep -E "FRF_PROFILE|>>>" +echo "exit_tests=$rc" +exit "$rc" diff --git a/tests/unit/frf_separate/test_bessel_sh_grouping.py b/tests/unit/frf_separate/test_bessel_sh_grouping.py index 762fdb1f..9edfe550 100644 --- a/tests/unit/frf_separate/test_bessel_sh_grouping.py +++ b/tests/unit/frf_separate/test_bessel_sh_grouping.py @@ -132,3 +132,66 @@ def test_the_group_representative_is_the_mean_not_a_member(): f"reordering two reflections in the same bin changed the result by " f"{rel:.2e}; the representative is order-dependent, so it is not the mean" ) + + +def _reference_expansion(s_vec, intensity, L, bessel_h_scale): + """A slow, obvious expansion, built from independently-checked parts. + + Sums over reflections one at a time, taking the Legendre factor from + ``sh._bar_legendre_recurrence`` -- which ``tests/unit/alignment/test_sh.py`` + pins against ``scipy`` -- and the radial factor from + ``spherical_bessel_table``. No grouping, no shells, no fused loop. + + This exists because the other tests in this file compare + ``bessel_sh_expand`` against *itself* at a different grouping resolution, so + a term dropped from the sum appears identically on both sides and cancels. + One did: a refactor formed the products that feed the shell accumulation + before writing the sectoral (m = l) entry of the Legendre row, silently + losing that entry for every even l. Every test here passed, and the only + signal was a benchmark truth rank moving from 8 to 13. + """ + from torchref.experimental.alignment.frf.data_mr import spherical_bessel_table + from torchref.experimental.alignment.sh import _bar_legendre_recurrence + + s_vec = torch.cat([s_vec, -s_vec], dim=0) # enforce_friedel + intensity = torch.cat([intensity, intensity], dim=0) + + lmax = L - 1 + lmax_even = lmax if lmax % 2 == 0 else lmax - 1 + N_radial = (lmax_even - 2) // 2 + 1 + u_max = lmax_even + 1 + + smag = s_vec.norm(dim=-1).clamp(min=1e-30) + cos_t = (s_vec[:, 2] / smag).clamp(-1.0, 1.0) + sin_t = (1.0 - cos_t * cos_t).clamp(min=0.0).sqrt() + phi = torch.atan2(s_vec[:, 1], s_vec[:, 0]) + + barP = _bar_legendre_recurrence(cos_t, sin_t, L) # (M, L, L) + x = (bessel_h_scale * smag).clamp(min=1e-30) + j = spherical_bessel_table(x, u_max) # (M, u_max+1) + + out = torch.zeros((N_radial, L, 2 * L - 1), dtype=torch.complex128) + for l in range(2, lmax_even + 1, 2): + for n in range((lmax_even - l) // 2 + 1): + u = l + 2 * n + 1 + radial = math.sqrt(2 * u + 1) * j[:, u] / x + for m in range(-l, l + 1): + # Y_lm = barP_{l,|m|} * C(m, phi), with C = (-1)^m e^{i m phi} + # for m >= 0 and e^{i m phi} for m < 0; the expansion uses the + # conjugate. + sign = (-1.0) ** m if m >= 0 else 1.0 + phase = torch.polar(torch.ones_like(phi), -m * phi) + term = (intensity * radial * barP[:, l, abs(m)] * sign) * phase + out[n, l, (L - 1) + m] = term.sum() + return out + + +@pytest.mark.parametrize("L", [9, 13]) +def test_matches_an_independent_direct_summation(L): + """The whole expansion, against a reference that shares no code with it.""" + s, I = _random_set(seed=31, n=120) + ref = _reference_expansion(s, I, L, 24.0) + got = bessel_sh_expand(s, I, L=L, bessel_h_scale=24.0).coeffs + assert got.shape == ref.shape + rel = (got - ref).abs().max().item() / max(ref.abs().max().item(), 1e-300) + assert rel < 1e-10, f"L={L}: differs from a direct summation by {rel:.2e}" diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index e774290d..86bff89e 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -21,12 +21,18 @@ _PROFILE = bool(os.environ.get("FRF_PROFILE")) -#: Byte budget for the per-chunk transients in :func:`bessel_sh_expand`. It -#: trades peak memory against throughput: too small and the contraction -#: degenerates into many small GEMMs and many Legendre recurrence calls, too -#: large and the tables fall out of cache and the function's peak memory is set -#: by this instead of by the result. -CLUSTER_CHUNK_BYTES = 256_000_000 +#: Byte budget for the per-chunk transients in :func:`bessel_sh_expand`. +#: +#: The chunk holds the Legendre recurrence's rolling rows, so this is really a +#: cache-residency knob, and it has an interior optimum. Measured on an EPYC +#: 9335 at four threads, seconds for the whole rotation function at cap 100 on +#: 3K7M: 24.9 at 2 MB, 12.1 at 8, **9.6 at 32**, 9.9 at 128, 11.0 at 256, 12.3 +#: at 1024. Same shape at cap 64. Below the optimum the 100-iteration loop over +#: l is re-run for too many chunks and Python and dispatch overhead dominate; +#: above it the rolling rows stop fitting in cache and the recurrence becomes +#: memory-bound. The truth rank was identical at every setting. +CLUSTER_CHUNK_BYTES = 32_000_000 + #: Grouping resolution for |s|, i.e. for the RADIAL factor. Reflections whose |s| #: agrees to 1/this share one Bessel evaluation. @@ -52,7 +58,7 @@ #: approximation and needs its own evidence. _GROUP_SCALE_COS = 10_000_000 -from ..sh import _bar_legendre_recurrence, evaluate_ylm +from ..sh import LEGENDRE_SEED, legendre_recurrence_coefficients from .types import BesselSHCoefficients @@ -229,7 +235,8 @@ def bessel_sh_expand( M = s_vectors.shape[0] einsum_dtype = compute_dtype if compute_dtype is not None else complex_dtype - prof = {"cluster": 0.0, "dbuild": 0.0, "bessel": 0.0, "ylm": 0.0, "einsum": 0.0} if _PROFILE else None + prof = {"cluster": 0.0, "dbuild": 0.0, "bessel": 0.0, "legendre": 0.0, + "scatter": 0.0, "contract": 0.0} if _PROFILE else None def _tick(t0): if device.type == "cuda": @@ -286,6 +293,16 @@ def _group_mean(values, index, n_groups): shell_of_cluster = torch.zeros(n_clusters, dtype=torch.long, device=device) shell_of_cluster[inverse] = inv_s shell_smag = _group_mean(s_mag_all.to(comp_real), inv_s, n_shells) + + # Reorder the clusters so each shell's members are adjacent. The angular + # accumulation below scatters every cluster into its shell's row of T; in + # cluster order those writes land all over T, in shell order they sweep it + # once. Same arithmetic, and the sort is one pass over n_clusters against a + # scatter of n_clusters x n_even x L. + order = torch.argsort(shell_of_cluster) + shell_of_cluster = shell_of_cluster[order] + rep_cos = rep_cos[order] + rep_sin = rep_sin[order] if _PROFILE: prof["cluster"] += _tick(t0); t0 = time.perf_counter() @@ -302,18 +319,29 @@ def _group_mean(values, index, n_groups): # (cluster, p) after the sum rather than once per (reflection, p) -- the same # number for a factor of M/n_clusters less work. p_idx = torch.arange(L, device=device) # (L,) + # `inverse` maps a reflection to its cluster in the ORIGINAL cluster order; + # the clusters were just permuted into shell order, so compose the two. + rank_of_cluster = torch.empty_like(order) + rank_of_cluster[order] = torch.arange(n_clusters, device=device) + cluster_of_refl = rank_of_cluster[inverse] Sp = torch.zeros((n_clusters, L), dtype=einsum_dtype, device=device) dchunk = 262_144 - prow = p_idx.to(comp_real).unsqueeze(0) # (1, L) for start_i in range(0, M, dchunk): stop = min(start_i + dchunk, M) - ph = phi_all[start_i:stop].to(comp_real).unsqueeze(1) # (c, 1) - i_c = intensity[start_i:stop].to(comp_real).unsqueeze(1) # (c, 1) - ang = -(prow * ph) - # polar(r, angle) is one sincos; exp of a complex number additionally - # evaluates exp() of a real part that is always zero here. - e_neg = torch.polar(i_c.expand(-1, L).contiguous(), ang).to(einsum_dtype) - Sp.index_add_(0, inverse[start_i:stop], e_neg) + ph = phi_all[start_i:stop].to(comp_real) # (c,) + i_c = intensity[start_i:stop].to(comp_real) # (c,) + # e^{-i p phi} = z^p with z = e^{-i phi}, so one transcendental per + # reflection and a running product over p, rather than a transcendental + # per (reflection, p). At L=101 over 2.6e6 reflections that is 2.6e8 + # sincos calls replaced by 2.6e6 of them plus a complex multiply each. + # The product accumulates about L * eps of relative error, ~2e-14, six + # orders below what the grouping already costs. + z = torch.polar(torch.ones_like(ph), -ph) # (c,) + ladder = z.unsqueeze(1).expand(-1, L).clone() + ladder[:, 0] = 1.0 # p = 0 + e_neg = torch.cumprod(ladder, dim=1) # (c, L) = z^p + e_neg = (e_neg * i_c.unsqueeze(1)).to(einsum_dtype) + Sp.index_add_(0, cluster_of_refl[start_i:stop], e_neg) sign_p = ((-1.0) ** p_idx.to(comp_real)).to(einsum_dtype) # (L,) Dp = Sp * sign_p.unsqueeze(0) # (n_clusters, L) if _PROFILE: @@ -330,47 +358,99 @@ def _group_mean(values, index, n_groups): # which trades n_clusters·N_radial for n_clusters + n_shells·N_radial. On the # benchmark that is 2.5x to 18x fewer multiply-adds, and it shrinks the Bessel # table by the same clusters-per-shell factor. Exact, not an approximation. + # + # The Legendre recurrence is run here rather than called, so each row can be + # accumulated into T the moment it exists and the (chunk, n_even, L) table is + # never built. That table was what bounded the chunk width, and the loop over + # l had to be repeated for every chunk -- 100 iterations of a handful of + # small kernels, 71 times over, which cost more in launch overhead than the + # arithmetic did. Without it the chunks are wide enough that the loop runs + # once or twice in total. + # + # `a_coef` and `b_coef` are zero for m >= l, so the vertical recurrence runs + # at full width; slicing to [:l] instead makes every iteration a differently + # shaped, mostly tiny kernel. n_even = len(even_ls) le_idx = (l_idx - 2) // 2 # l value -> even-l row index - T = torch.zeros((n_shells, n_even, L), dtype=einsum_dtype, device=device) + a_coef, b_coef, sect = legendre_recurrence_coefficients(L, comp_real, device) + + # The whole per-shell sum T and the whole radial table B used to be built at + # full size, and both are large: at L=101 with 35k shells they are 2.8 GB and + # 0.7 GB. They are also touched once each, so that is pure memory traffic -- + # and the scatter into a 2.8 GB target misses cache on essentially every + # write. + # + # The clusters are sorted by shell, so a chunk of clusters spans a + # *contiguous* range of shells. That lets both be per-chunk, and lets the + # radial contraction be folded into the same loop: once a chunk's shells are + # complete, contract them and drop them. The arithmetic is identical -- the + # contraction still costs n_shells x n_even x N_radial x L in total -- but the + # scatter target is now tens of MB rather than gigabytes, and the running + # answer c_pos is a few MB, so both stay in cache. + c_pos = torch.zeros((N_radial, n_even, L), dtype=einsum_dtype, device=device) rbytes = 4 if comp_real == torch.float32 else 8 - per_cluster = rbytes * (n_even * L + 2 * L) + per_cluster = rbytes * 6 * L cstep = max(1, min(n_clusters, CLUSTER_CHUNK_BYTES // max(1, per_cluster))) for cs in range(0, n_clusters, cstep): ce = min(cs + cstep, n_clusters) if _PROFILE: t0 = time.perf_counter() - # Even rows only: the odd-l coefficients vanish for a centrosymmetric - # Patterson, and an all-l table is twice the write. - P = _bar_legendre_recurrence( - rep_cos[cs:ce], rep_sin[cs:ce], L, keep_l=even_l_idx) # (chunk, n_even, L) real + sh_abs = shell_of_cluster[cs:ce] + s0 = int(sh_abs[0]) + s1 = int(sh_abs[-1]) + 1 # sorted, so this is the range + nb = s1 - s0 + sh = sh_abs - s0 # shell index within the chunk + + # Radial weights for this chunk's shells only. + x_s = (bessel_h_scale * shell_smag[s0:s1]).clamp(min=1e-30) + j_all = spherical_bessel_table(x_s, u_max) # (nb, u_max+1) + B = torch.zeros((nb, n_even, N_radial), dtype=comp_real, device=device) + B[:, le_idx, n_idx] = w_vec.unsqueeze(0) * j_all[:, u_idx] / x_s.unsqueeze(-1) if _PROFILE: - prof["ylm"] += _tick(t0); t0 = time.perf_counter() - # P is real; carrying D's halves separately keeps this two real - # multiplies per element instead of a complex multiply's four. - Dc = Dp[cs:ce].unsqueeze(1) # (chunk, 1, L) - T.index_add_(0, shell_of_cluster[cs:ce], - torch.complex(P * Dc.real, P * Dc.imag).to(einsum_dtype)) + prof["bessel"] += _tick(t0); t0 = time.perf_counter() + + # (n_even, nb, L) each, so Tr[pos] is contiguous for index_add_. Two real + # accumulators rather than one complex: the scatter stays a real + # operation and no complex temporary is allocated per (l, chunk). + Tr = torch.zeros((n_even, nb, L), dtype=comp_real, device=device) + Ti = torch.zeros((n_even, nb, L), dtype=comp_real, device=device) + + cos_e = rep_cos[cs:ce].unsqueeze(-1) # (chunk, 1) + sin_c = rep_sin[cs:ce] # (chunk,) + Dr = Dp[cs:ce].real # (chunk, L) + Di = Dp[cs:ce].imag + prev2 = torch.zeros((ce - cs, L), dtype=comp_real, device=device) + prev1 = torch.zeros((ce - cs, L), dtype=comp_real, device=device) + prev1[:, 0] = LEGENDRE_SEED # bar_P_0^0 + for l in range(1, L): + # `a_coef` and `b_coef` are zero for m >= l, so this runs at full + # width; slicing to [:l] instead makes every iteration a differently + # shaped, mostly tiny kernel. + cur = a_coef[l] * cos_e * prev1 - b_coef[l] * prev2 + # The sectoral term must land BEFORE the products below are formed: + # it is the m = l entry of this very row. Computing `cur * Dr` first + # silently drops that entry for every even l -- which moved 3K7M's + # truth rank from 8 to 13 when a refactor got the order wrong. + cur[:, l] = sect[l] * sin_c * prev1[:, l - 1] # sectoral m = l + if l >= 2 and (l % 2 == 0): + if _PROFILE: + prof["legendre"] += _tick(t0); t0 = time.perf_counter() + pos = (l - 2) // 2 + Tr[pos].index_add_(0, sh, cur * Dr) + Ti[pos].index_add_(0, sh, cur * Di) + if _PROFILE: + prof["scatter"] += _tick(t0); t0 = time.perf_counter() + prev2, prev1 = prev1, cur if _PROFILE: - prof["einsum"] += _tick(t0) + prof["legendre"] += _tick(t0); t0 = time.perf_counter() - if _PROFILE: - t0 = time.perf_counter() - # Radial weights per shell: sqrt(2u+1) j_u(x)/x at u = l + 2n + 1. - x_s = (bessel_h_scale * shell_smag).clamp(min=1e-30) - j_all = spherical_bessel_table(x_s, u_max) # (n_shells, u_max+1) - B = torch.zeros((n_shells, n_even, N_radial), dtype=comp_real, device=device) - B[:, le_idx, n_idx] = w_vec.unsqueeze(0) * j_all[:, u_idx] / x_s.unsqueeze(-1) - if _PROFILE: - prof["bessel"] += _tick(t0); t0 = time.perf_counter() - - c_pos = torch.complex( - torch.einsum("iln,ilm->nlm", B, T.real), - torch.einsum("iln,ilm->nlm", B, T.imag), - ).to(einsum_dtype) - if _PROFILE: - prof["einsum"] += _tick(t0) + c_pos += torch.complex( + torch.einsum("iln,lim->nlm", B, Tr), + torch.einsum("iln,lim->nlm", B, Ti), + ).to(einsum_dtype) + if _PROFILE: + prof["contract"] += _tick(t0) # Mirror onto m < 0: c[-p] = (-1)^p conj(c[+p]). c_e = torch.zeros((N_radial, n_even, 2 * L - 1), dtype=einsum_dtype, @@ -382,11 +462,12 @@ def _group_mean(values, index, n_groups): c_nlm = torch.zeros((N_radial, L, 2 * L - 1), dtype=complex_dtype, device=device) c_nlm[:, even_l_idx, :] = c_e.to(complex_dtype) if _PROFILE: - prof["einsum"] += _tick(t0) + prof["contract"] += _tick(t0) if _PROFILE: tot = sum(prof.values()) + 1e-30 print(f"[FRF_PROFILE] M={M} n_clusters={n_clusters} ({M/max(1,n_clusters):.1f}x) " + f"n_shells={n_shells} ({n_clusters/max(1,n_shells):.1f} clu/shell) " f"L={L} dtype={comp_real} | " + " ".join(f"{k}={v*1000:.0f}ms({100*v/tot:.0f}%)" for k, v in prof.items()), flush=True) diff --git a/torchref/experimental/alignment/sh.py b/torchref/experimental/alignment/sh.py index dfb59b84..12e97ea3 100644 --- a/torchref/experimental/alignment/sh.py +++ b/torchref/experimental/alignment/sh.py @@ -27,6 +27,48 @@ _YLM_PROF = {"recurrence": 0.0, "assembly": 0.0} +def legendre_recurrence_coefficients(L: int, dtype, device): + """Coefficient tables for the fully-normalised Legendre recurrence. + + Returns ``(a, b, sect)`` with ``a``, ``b`` of shape ``(L, L)`` indexed + ``[l, m]`` and ``sect`` of shape ``(L,)``:: + + a_l^m = sqrt((2l-1)(2l+1) / ((l-m)(l+m))) + b_l^m = sqrt((2l+1)(l+m-1)(l-m-1) / ((l-m)(l+m)(2l-3))) + sect_m = sqrt((2m+1) / (2m)) + + ``a`` and ``b`` are **zero wherever m >= l**, which lets a caller run the + vertical recurrence at full width instead of slicing to ``[:l]``: the + out-of-support entries come out zero on their own. ``b`` also vanishes at + l = m+1 through its ``(l-m-1)`` factor, so that case needs no special + handling. + + Split out of :func:`_bar_legendre_recurrence` because the rotation function + runs this recurrence itself, fused with its own accumulation, and two copies + of these formulae would be two chances to get them subtly different. + """ + ll = torch.arange(L, dtype=torch.float64, device=device).view(L, 1) + mm = torch.arange(L, dtype=torch.float64, device=device).view(1, L) + valid = ll > mm + denom = (ll - mm) * (ll + mm) + denom_safe = torch.where(valid, denom, torch.ones_like(denom)) + a = torch.sqrt((2.0 * ll - 1.0) * (2.0 * ll + 1.0) / denom_safe) + b_num = (2.0 * ll + 1.0) * (ll + mm - 1.0) * (ll - mm - 1.0) + b_den = denom_safe * (2.0 * ll - 3.0) + b = torch.sqrt(torch.clamp( + b_num / torch.where(b_den == 0, torch.ones_like(b_den), b_den), min=0.0)) + a = torch.where(valid, a, torch.zeros_like(a)).to(dtype) + b = torch.where(valid, b, torch.zeros_like(b)).to(dtype) + m_arange = torch.arange(L, dtype=torch.float64, device=device) + sect = torch.sqrt( + (2.0 * m_arange + 1.0) / (2.0 * m_arange).clamp(min=1.0)).to(dtype) + return a, b, sect + + +#: bar_P_0^0. +LEGENDRE_SEED = 1.0 / math.sqrt(4.0 * math.pi) + + def _bar_legendre_recurrence( cos_theta: torch.Tensor, sin_theta: torch.Tensor, @@ -67,28 +109,10 @@ def _bar_legendre_recurrence( dtype = cos_theta.dtype device = cos_theta.device - inv_sqrt_4pi = 1.0 / math.sqrt(4.0 * math.pi) + inv_sqrt_4pi = LEGENDRE_SEED - # Precompute the recurrence coefficients as (L, L) tables, indexed [l, m]: - # a_l^m = sqrt((2l-1)(2l+1)/((l-m)(l+m))) - # b_l^m = sqrt((2l+1)(l+m-1)(l-m-1)/((l-m)(l+m)(2l-3))) [0 when l = m+1] - # b vanishes at l = m+1 (factor l-m-1 = 0), so no special-casing is needed. - # Computed in float64 then cast to `dtype` (matches the original, which used - # float64 python scalars multiplied into the working-dtype tensors). - ll = torch.arange(L, dtype=torch.float64, device=device).view(L, 1) - mm = torch.arange(L, dtype=torch.float64, device=device).view(1, L) - valid = (ll > mm) # l > m (vertical recurrence region) - denom = (ll - mm) * (ll + mm) - denom_safe = torch.where(valid, denom, torch.ones_like(denom)) - a_coef = torch.sqrt((2.0 * ll - 1.0) * (2.0 * ll + 1.0) / denom_safe) - b_num = (2.0 * ll + 1.0) * (ll + mm - 1.0) * (ll - mm - 1.0) - b_den = denom_safe * (2.0 * ll - 3.0) - b_coef = torch.sqrt(torch.clamp(b_num / torch.where(b_den == 0, torch.ones_like(b_den), b_den), min=0.0)) - a_coef = torch.where(valid, a_coef, torch.zeros_like(a_coef)).to(dtype) - b_coef = torch.where(valid, b_coef, torch.zeros_like(b_coef)).to(dtype) - # Sectoral diagonal factor sqrt((2m+1)/(2m)) for m = l. - m_arange = torch.arange(L, dtype=torch.float64, device=device) - sect = torch.sqrt((2.0 * m_arange + 1.0) / (2.0 * m_arange).clamp(min=1.0)).to(dtype) + a_coef, b_coef, sect = legendre_recurrence_coefficients( + L, dtype, device) cos_e = cos_theta.unsqueeze(-1) # (..., 1) sin_e = sin_theta.unsqueeze(-1) From 66f126bf305afc1787e6734dc77ea355884c747b Mon Sep 17 00:00:00 2001 From: Kevin Dalton Date: Fri, 21 Aug 2026 11:24:27 -0400 Subject: [PATCH 033/250] Run rigid body against a sandbox so it cannot disturb its caller `RigidBodyRefinementStep` rebinds the Refinement to a resolution-truncated data view at every cutoff. `_rebind_for_data` assigns `reflection_data`, builds a fresh Scaler, and calls `_init_targets` + `reset_loss_state`. Run against the caller's own Refinement, those assignments are destructive: - `_init_targets` reconstructs `adp_target` and `geometry_target` from constructor defaults, so anything configured on them post-construction is silently reset. `adp_target['simu'].simu_sigma = 0.25` reads back as 2.0. - `reset_loss_state` discards the LossState, so a weight registered on it is gone. A key present in DEFAULT_GROUP_WEIGHTS is visibly overwritten; a custom one such as `adp/simu` simply returns None afterwards and its target falls back to the group weight. Neither rebuild is wanted by the step. Only the x-ray target depends on the data and scaler that changed; the ADP and geometry targets are built from the model alone, and `_run_one_cutoff` drops every non-xray target from the state before optimizing. Measured on a 3-cutoff run they are constructed 3 times and evaluated 6 times (registration probe plus loss refresh) purely as overhead, then deleted unused, while the x-ray target takes all 42 gradient evaluations. `run()` now points the step at a shallow clone that shares the model but owns its own attribute namespace, so every one of those assignments lands on the clone. There is nothing to restore afterwards and no window in which the caller's Refinement is inconsistent. The model is deliberately shared rather than copied: `use_rigid_xyz` swaps its xyz container in place, so refined coordinates reach the caller by object identity and no copy-back is needed. That is also what makes the change exactly equivalent rather than approximately so -- on 3E98 the refined coordinates are bit-identical to the previous behaviour, max per-atom difference 0.000e+00. `nn.Module` keeps submodules in `_modules`, so the clone copies that dict (and `_parameters` / `_buffers`) as well as `__dict__`; without it a submodule assignment on the clone would write straight through to the original. Tests: tests/integration/test_rigid_body_isolation.py, five cases -- the sigma survives, a custom LossState weight survives, targets and reflection_data keep their object identity, coordinates still reach the caller, and a normal macrocycle still runs afterwards. Verified that three of them fail when the sandbox is bypassed. Full unit + functional suite passes (1798 passed, 74 skipped). This is independent of #68. That PR gives ADP restraint parameters a constructor-level home so they survive *any* rebuild, including the `create_from_state_dict` and ensemble paths this change does not touch. Either can land without the other; together they cover both the storage and the rebuild. --- .../integration/test_rigid_body_isolation.py | 96 +++++++++++++++++++ torchref/refinement/rigid_body_refinement.py | 53 +++++++++- 2 files changed, 146 insertions(+), 3 deletions(-) create mode 100644 tests/integration/test_rigid_body_isolation.py diff --git a/tests/integration/test_rigid_body_isolation.py b/tests/integration/test_rigid_body_isolation.py new file mode 100644 index 00000000..651ed11d --- /dev/null +++ b/tests/integration/test_rigid_body_isolation.py @@ -0,0 +1,96 @@ +"""refine_rigid_body must not disturb the refinement it is called on. + +Every cutoff rebinds the step's Refinement to a resolution-truncated data view, +which rebuilds the scaler and every target and drops the loss state. Those +rebuilds are needed for the x-ray target, whose data and scaler genuinely +change; they are collateral for the ADP and geometry targets, which are built +from the model alone and are dropped from the loss state before the optimizer +runs. `RigidBodyRefinementStep` therefore runs against a sandbox clone, and +these tests pin that the caller sees none of it. +""" +import pytest +import torch + +from torchref import LBFGSRefinement + + +@pytest.fixture(scope="module") +def refinement(mtz_dir, pdb_dir): + def build(): + return LBFGSRefinement( + data_file=str(mtz_dir / "3E98.mtz"), + pdb=str(pdb_dir / "3E98.pdb"), + device=torch.device("cpu"), + verbose=0, + ) + return build + + +def test_component_restraint_config_survives(refinement): + """A sigma set on a target must still be set afterwards. + + Regression: `_init_targets` rebuilt `TotalADPTarget` per cutoff with no + restraint parameters, so this silently reverted to the ADPSimilarityTarget + default of 2.0 and refinement continued at a restraint weight nobody chose. + """ + ref = refinement() + ref.adp_target["simu"].simu_sigma = 0.25 + ref.get_scales() + + ref.refine_rigid_body(iterations_per_step=10) + + assert ref.adp_target["simu"].simu_sigma == pytest.approx(0.25) + + +def test_custom_loss_state_weight_survives(refinement): + """A weight registered on the LossState must still be registered afterwards. + + `adp/simu` is deliberately a key absent from DEFAULT_GROUP_WEIGHTS: a key + that is present gets visibly overwritten, while a custom one silently + disappears and its target falls back to the group weight. + """ + ref = refinement() + ref.get_scales() + ref.complete_loss_state().set_weight("adp/simu", 0.77) + + ref.refine_rigid_body(iterations_per_step=10) + + assert ref.complete_loss_state().weights.get("adp/simu") == pytest.approx(0.77) + + +def test_targets_and_data_are_not_replaced(refinement): + """Object identity, not just values -- a caller may hold its own references.""" + ref = refinement() + ref.get_scales() + adp, geometry, data = ref.adp_target, ref.geometry_target, ref.reflection_data + + ref.refine_rigid_body(iterations_per_step=10) + + assert ref.adp_target is adp + assert ref.geometry_target is geometry + assert ref.reflection_data is data + + +def test_refined_coordinates_still_reach_the_caller(refinement): + """The sandbox shares the model, so the whole point still has to work.""" + ref = refinement() + ref.get_scales() + before = ref.model.xyz().detach().clone() + + ref.refine_rigid_body(iterations_per_step=30) + + shift = (ref.model.xyz().detach() - before).norm(dim=-1) + assert float(shift.max()) > 0.0, "rigid body moved nothing" + + +def test_refinement_is_usable_afterwards(refinement): + """A normal macrocycle must still run against the caller's own objects.""" + ref = refinement() + ref.get_scales() + + ref.refine_rigid_body(iterations_per_step=10) + ref.refine_scaler() + ref.refine_adp() + + rwork, rfree = ref.get_rfactor() + assert 0.0 < rwork < 1.0 and 0.0 < rfree < 1.0 diff --git a/torchref/refinement/rigid_body_refinement.py b/torchref/refinement/rigid_body_refinement.py index 70bb8d30..3f320d2b 100644 --- a/torchref/refinement/rigid_body_refinement.py +++ b/torchref/refinement/rigid_body_refinement.py @@ -8,6 +8,7 @@ below it switches to ``ml`` with the normal Scaler. """ +import copy from typing import List, Optional import torch @@ -88,14 +89,60 @@ def _xray_mode_for_cutoff(cls, d_min: float) -> str: # ----------------------------------------------------------------------- # Run # ----------------------------------------------------------------------- + @staticmethod + def _sandbox(ref): + """A shallow clone of ``ref`` that shares its model but owns its namespace. + + Every cutoff calls :meth:`_rebind_for_data`, which assigns + ``reflection_data``, builds a fresh ``Scaler``, and calls + ``_init_targets`` + ``reset_loss_state``. Run against the real + Refinement those assignments are destructive: ``_init_targets`` + reconstructs ``adp_target`` and ``geometry_target`` from constructor + defaults, so anything configured on them post-construction (for instance + ``adp_target['simu'].simu_sigma``) is silently reset, and + ``reset_loss_state`` discards custom weights registered on the + ``LossState``. Neither survives the step, and neither is wanted by it -- + rigid body drops every non-xray target before optimizing, so those + objects are rebuilt only to be thrown away. + + Directing the step at a clone confines all of it. The real Refinement is + never written to, so there is nothing to restore and no window in which + it is inconsistent. + + The model is deliberately shared, not copied: ``use_rigid_xyz`` swaps its + xyz container in place, so refined coordinates reach the caller by object + identity and need no copy-back. + + ``nn.Module`` keeps submodules in ``_modules``; copying ``__dict__`` + alone would leave that dict shared, and a submodule assignment on the + clone would write straight through to the original. + """ + sandbox = copy.copy(ref) + sandbox.__dict__ = dict(ref.__dict__) + for slot in ("_modules", "_parameters", "_buffers"): + if slot in sandbox.__dict__: + sandbox.__dict__[slot] = dict(ref.__dict__[slot]) + return sandbox + def run(self): """Step through every cutoff coarse to fine and return ``[(d_min, LossState), ...]``. - Restores the original ``reflection_data`` on exit, and bakes the final transform - back - into a plain ``ModelFT`` unless ``commit=False``. + Runs against a sandbox clone of the refinement (see :meth:`_sandbox`), so + the caller's targets, weights and ``reflection_data`` are left untouched. + Refined coordinates still reach the caller: the model is shared. + + Bakes the final transform back into a plain ``ModelFT`` unless + ``commit=False``. """ + real = self.refinement + self.refinement = self._sandbox(real) + try: + return self._run() + finally: + self.refinement = real + + def _run(self): ref = self.refinement original_data = ref.reflection_data From 7663c1630d1ebbaf03f6df83009bc2b4967e87c1 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 17:51:49 +0200 Subject: [PATCH 034/250] Fuse the Legendre recurrence with the shell accumulation, in a C++ kernel The two stages were bandwidth-bound at ~50 GB/s in plain torch: each row of the recurrence made a round trip to memory before the scatter read it back. Fused, one cluster's three rows are 1.2 kB of stack, so the row is produced, multiplied and accumulated without ever reaching memory. Two things fall out of writing it as a loop nest rather than as tensor operations: * **Ragged widths are free.** `bar_P[l, m]` is zero for m > l, so step l needs only columns 0..l -- a loop bound, nothing more. The same saving was attempted in torch first, with bucketed narrow views, and measured *slower*: 12.14 s against 9.85 s at cap 100 on 3K7M, because a strided scatter target costs more than the zeros it skips. That attempt also showed the recurrence does not respond at all to 1.55x less arithmetic, which is what identified it as bandwidth-bound rather than compute-bound. * **No atomics.** The clusters are already sorted by shell, so a thread owning a range of shells owns every write into those rows. Parallelising over clusters would race on the accumulator. `schedule(dynamic, 8)` because clusters per shell varies from 2.7 to 39 across the benchmark, so equal shell counts are not equal work. Dispatched through a `BackendTable`, following `torchref/base/electron_density/_backends.py`, with the torch implementation as the portable row -- which is both the fallback where no compiler exists and the reference the kernel is checked against. `on_failure="raise"`, deliberately not `"degrade"`: the kernel accumulates in place, so a mid-run failure leaves the accumulators partly written and a fallback would add its contributions on top. float32 only, matching this codebase's other kernels. That required moving the angular half of the expansion to single precision on every device, not just CUDA -- a float64 caller would mean the dtype gate never selects the kernel and it would sit there dead. The radial Bessel recurrence keeps its float64 internals, where the downward recurrence's cancellation needs them. Against an ungrouped float64 reference the angular stage costs 2.8e-7 at L=65 and 1.1e-6 at L=101, against 2.1e-8 and 5.3e-8 in double -- real, and still 20x to 90x tighter than Phaser's own cos(theta) bucketing at ~2e-5. Fused and portable agree to 4e-7 - 1e-5, not bit-exactly: the kernel accumulates cluster-by-cluster within a shell while `index_add_` accumulates over a whole chunk, and in single precision a different summation order is a different answer. The truth rank is identical either way. Measured on a pinned EPYC 9335 at four threads, steady state, ranks unchanged (1DAW 0, 3K7M 8 at cap 64 and 11 at cap 100): float64 portable float32 portable float32 fused cap64 1DAW 0.77 s 0.73 s 0.60 s cap64 3K7M 2.46 s 2.25 s 1.89 s cap100 1DAW 3.39 s 2.76 s 2.14 s cap100 3K7M 9.85 s 7.81 s 5.26 s Against the original baseline that is 81.4 s -> 5.26 s at cap 100 on 3K7M, of which the kernel itself is the last 1.48x; the rest was the summation order. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/kernel_ab.sh | 118 +++++++ alignment_lab/analysis/kernel_check.sh | 62 ++++ .../unit/frf_separate/test_legendre_kernel.py | 106 +++++++ .../experimental/alignment/frf/_backends.py | 64 ++++ .../experimental/alignment/frf/data_mr.py | 42 +-- .../alignment/frf/kernels/__init__.py | 1 + .../alignment/frf/kernels/cpu/__init__.py | 0 .../frf/kernels/cpu/legendre_shell.py | 295 ++++++++++++++++++ .../alignment/frf/kernels/portable.py | 75 +++++ .../experimental/alignment/rotation_search.py | 13 +- 10 files changed, 744 insertions(+), 32 deletions(-) create mode 100644 alignment_lab/analysis/kernel_ab.sh create mode 100644 alignment_lab/analysis/kernel_check.sh create mode 100644 tests/unit/frf_separate/test_legendre_kernel.py create mode 100644 torchref/experimental/alignment/frf/_backends.py create mode 100644 torchref/experimental/alignment/frf/kernels/__init__.py create mode 100644 torchref/experimental/alignment/frf/kernels/cpu/__init__.py create mode 100644 torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py create mode 100644 torchref/experimental/alignment/frf/kernels/portable.py diff --git a/alignment_lab/analysis/kernel_ab.sh b/alignment_lab/analysis/kernel_ab.sh new file mode 100644 index 00000000..4fcc0c9a --- /dev/null +++ b/alignment_lab/analysis/kernel_ab.sh @@ -0,0 +1,118 @@ +#!/bin/bash +# Fused C++ kernel against the portable torch reference: correctness first, then +# speed. Interleaved repeats and the truth rank beside each timing. +#SBATCH --job-name=frf_kernel_ab +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 +export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" +"$PY" -u -c " +import sys, time +sys.path.insert(0,'alignment_lab') +import torch; torch.set_grad_enabled(False) +from torchref.experimental.alignment.frf.kernels.cpu import legendre_shell as K +from torchref.experimental.alignment.frf.kernels import portable as P +from torchref.utils.backends import set_force_portable + +print('kernel available:', K.available()) +print('float64 must be refused, not reinterpreted:') +try: + import torch as _t + z = _t.zeros(2, 3, 5, dtype=_t.float64) + K.legendre_shell_accumulate( + z, z.clone(), _t.zeros(4, dtype=_t.float64), _t.zeros(4, dtype=_t.float64), + _t.zeros(4, 5, dtype=_t.float64), _t.zeros(4, 5, dtype=_t.float64), + _t.zeros(4, dtype=_t.long), _t.zeros(5, 5, dtype=_t.float64), + _t.zeros(5, 5, dtype=_t.float64), _t.zeros(5, dtype=_t.float64)) + print(' PROBLEM: float64 was accepted') +except Exception as e: + msg = str(e) + ok = ('float32 only' in msg) or ('dtype' in msg) + print((' refused: ' if ok else ' WRONG ERROR (not a dtype refusal): ') + + f'{type(e).__name__}: {msg.splitlines()[0][:110]}') + if not ok: + raise SystemExit(1) +if not K.available(): + print('why:', K.why_unavailable()) + err = K.last_error() + if err: print(err[1][:3000]) + raise SystemExit(1) + +# --- correctness: fused vs portable on random shell-sorted input ------------- +from torchref.experimental.alignment.sh import legendre_recurrence_coefficients +g = torch.Generator().manual_seed(4) +for L, n_c, n_sh, dt in ((13, 500, 40, torch.float32), + (65, 4000, 300, torch.float32), + (101, 3000, 250, torch.float32), + (65, 2000, 150, torch.float32)): + n_even = (L - 1 if (L-1) % 2 == 0 else L - 2) // 2 + ct = (2*torch.rand(n_c, generator=g, dtype=dt)-1) + st = (1-ct*ct).clamp(min=0).sqrt() + Dr = torch.randn(n_c, L, generator=g, dtype=dt) + Di = torch.randn(n_c, L, generator=g, dtype=dt) + sh = torch.sort(torch.randint(0, n_sh, (n_c,), generator=g))[0] + a, b, se = legendre_recurrence_coefficients(L, dt, torch.device('cpu')) + ref_r = torch.zeros(n_even, n_sh, L, dtype=dt); ref_i = torch.zeros_like(ref_r) + got_r = torch.zeros_like(ref_r); got_i = torch.zeros_like(ref_r) + P.legendre_shell_accumulate(ref_r, ref_i, ct, st, Dr, Di, sh, a, b, se) + K.legendre_shell_accumulate(got_r, got_i, ct, st, Dr, Di, sh, a, b, se) + sc = max(ref_r.abs().max().item(), 1e-300) + er = (got_r-ref_r).abs().max().item()/sc + ei = (got_i-ref_i).abs().max().item()/max(ref_i.abs().max().item(),1e-300) + print(f' L={L:3d} n_c={n_c:5d} {str(dt):15s} rel err re {er:.2e} im {ei:.2e}') + +# --- what single precision costs, against an ungrouped float64 reference ---- +import torchref.experimental.alignment.frf.data_mr as dm +g2 = torch.Generator().manual_seed(77) +for L, hs in ((65, 64.0), (101, 100.0)): + sv = torch.randn(6000, 3, generator=g2, dtype=torch.float64) + sv = sv / sv.norm(dim=-1, keepdim=True) * ( + 0.07 + 0.18*torch.rand(6000, 1, generator=g2, dtype=torch.float64)) + I = torch.randn(6000, generator=g2, dtype=torch.float64) + ks, kc = dm._GROUP_SCALE_S, dm._GROUP_SCALE_COS + dm._GROUP_SCALE_S = dm._GROUP_SCALE_COS = 10**16 + exact = dm.bessel_sh_expand(sv, I, L=L, bessel_h_scale=hs).coeffs + dm._GROUP_SCALE_S, dm._GROUP_SCALE_COS = ks, kc + sc = max(exact.abs().max().item(), 1e-300) + f64 = dm.bessel_sh_expand(sv, I, L=L, bessel_h_scale=hs).coeffs + f32 = dm.bessel_sh_expand(sv, I, L=L, bessel_h_scale=hs, + compute_dtype=torch.complex64).coeffs + print(f' L={L:3d} vs exact: float64 angular {((f64-exact).abs().max()/sc):.2e}' + f' float32 angular {((f32-exact).abs().max()/sc):.2e}') + +# --- speed on the real thing ------------------------------------------------- +from lab import FRFConfig, orbit_rank, rotated_case, run_frf, seed_for +cases = {p: rotated_case(p, seed_for(p, 0)) for p in ('1DAW', '3K7M')} +for cap in (64, 100): + cfg = FRFConfig(lmax_cap=cap, n_peaks=200) + for forced in (True, False): + set_force_portable(forced) + for m, d, _ in cases.values(): + run_frf(m, d, cfg, capture_arf=False) + res, rank = {}, {} + for rep in range(3): + for arm, forced in (('portable', True), ('fused', False)): + set_force_portable(forced) + for p, (m, d, R) in cases.items(): + t0 = time.perf_counter() + r = run_frf(m, d, cfg, capture_arf=False) + res.setdefault((arm,p), []).append(time.perf_counter()-t0) + k, _ = orbit_rank(r.peaks, R, + d.spacegroup.matrices.to(torch.float64).cpu(), + reciprocal_basis=d.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), + side='left', frame='cart') + rank[(arm,p)] = k + set_force_portable(None) + print(f'--- cap{cap} (best of 3, seconds) ---') + for p in cases: + e, c = min(res[('portable',p)]), min(res[('fused',p)]) + print(f' {p:6s} portable {e:7.2f} [rank {rank[(\"portable\",p)]:>3}] ' + f'fused {c:7.2f} [rank {rank[(\"fused\",p)]:>3}] {e/max(c,1e-9):5.2f}x') +" +echo "exit_code=$?" diff --git a/alignment_lab/analysis/kernel_check.sh b/alignment_lab/analysis/kernel_check.sh new file mode 100644 index 00000000..449319aa --- /dev/null +++ b/alignment_lab/analysis/kernel_check.sh @@ -0,0 +1,62 @@ +#!/bin/bash +# Does the fused kernel build, refuse float64, and agree with the portable +# reference? Correctness only, so it does not need an exclusive node -- the +# timing A/B does. +#SBATCH --job-name=frf_kernel_check +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 PYTHONUNBUFFERED=1 +export CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname)" +"$PY" -u -c " +import torch; torch.set_grad_enabled(False) +from torchref.experimental.alignment.frf.kernels.cpu import legendre_shell as K +from torchref.experimental.alignment.frf.kernels import portable as P +from torchref.experimental.alignment.sh import legendre_recurrence_coefficients + +print('available:', K.available()) +if not K.available(): + print('why:', K.why_unavailable()) + e = K.last_error() + if e: print(e[1][-4000:]) + raise SystemExit(1) + +try: + z = torch.zeros(2, 3, 5, dtype=torch.float64) + K.legendre_shell_accumulate(z, z.clone(), + torch.zeros(4, dtype=torch.float64), torch.zeros(4, dtype=torch.float64), + torch.zeros(4, 5, dtype=torch.float64), torch.zeros(4, 5, dtype=torch.float64), + torch.zeros(4, dtype=torch.long), torch.zeros(5, 5, dtype=torch.float64), + torch.zeros(5, 5, dtype=torch.float64), torch.zeros(5, dtype=torch.float64)) + print('PROBLEM: float64 accepted') +except Exception as e: + msg = str(e) + ok = ('float32 only' in msg) or ('dtype' in msg) + print(('float64 refused: ' if ok else 'WRONG ERROR (not a dtype refusal): ') + + msg.splitlines()[0][:120]) + if not ok: + raise SystemExit(1) + +g = torch.Generator().manual_seed(4) +for L, n_c, n_sh in ((13, 500, 40), (65, 4000, 300), (101, 3000, 250)): + n_even = (L - 1 if (L-1) % 2 == 0 else L - 2) // 2 + ct = (2*torch.rand(n_c, generator=g, dtype=torch.float32)-1) + st = (1-ct*ct).clamp(min=0).sqrt() + Dr = torch.randn(n_c, L, generator=g, dtype=torch.float32) + Di = torch.randn(n_c, L, generator=g, dtype=torch.float32) + sh = torch.sort(torch.randint(0, n_sh, (n_c,), generator=g))[0] + a, b, se = legendre_recurrence_coefficients(L, torch.float32, torch.device('cpu')) + ref_r = torch.zeros(n_even, n_sh, L, dtype=torch.float32); ref_i = torch.zeros_like(ref_r) + got_r = torch.zeros_like(ref_r); got_i = torch.zeros_like(ref_r) + P.legendre_shell_accumulate(ref_r, ref_i, ct, st, Dr, Di, sh, a, b, se) + K.legendre_shell_accumulate(got_r, got_i, ct, st, Dr, Di, sh, a, b, se) + sr = max(ref_r.abs().max().item(), 1e-30); si = max(ref_i.abs().max().item(), 1e-30) + print(f' L={L:3d}: rel err re {(got_r-ref_r).abs().max().item()/sr:.2e}' + f' im {(got_i-ref_i).abs().max().item()/si:.2e}') +" +echo "exit_code=$?" diff --git a/tests/unit/frf_separate/test_legendre_kernel.py b/tests/unit/frf_separate/test_legendre_kernel.py new file mode 100644 index 00000000..0df4d5e8 --- /dev/null +++ b/tests/unit/frf_separate/test_legendre_kernel.py @@ -0,0 +1,106 @@ +"""The fused Legendre/shell kernel against the portable reference. + +Two things need pinning. The kernel must agree with the torch reference -- it is +selected automatically wherever it builds, so a divergence would silently change +every rotation search on that host. And it must **refuse** float64 rather than +accept it: it reads every array through a raw ``float*``, so a float64 buffer +would be reinterpreted as twice as many float32s, not converted. + +Agreement is to a float32 tolerance, not bit-exact, and deliberately so: the +kernel accumulates cluster-by-cluster within a shell while ``index_add_`` +accumulates over the whole chunk, and in single precision a different summation +order is a different answer. Measured 4e-7 to 1e-5 relative, which sits between +the grouping's own error and Phaser's cos(theta) bucketing at ~2e-5. +""" + +import pytest +import torch + +from torchref.experimental.alignment.frf.kernels import portable +from torchref.experimental.alignment.frf.kernels.cpu import legendre_shell as fused +from torchref.experimental.alignment.sh import legendre_recurrence_coefficients + +pytestmark = pytest.mark.unit + +#: Summation order in single precision, nothing more. Set above the measured +#: 1e-5 worst case with margin; a real divergence is orders larger, because a +#: dropped term changes whole rows rather than their last digits. +_TOL = 1e-4 + + +def _case(L, n_clusters, n_shells, seed): + """Random input in the layout the kernel requires: shells sorted.""" + g = torch.Generator().manual_seed(seed) + cos_t = 2 * torch.rand(n_clusters, generator=g, dtype=torch.float32) - 1 + sin_t = (1 - cos_t * cos_t).clamp(min=0).sqrt() + Dr = torch.randn(n_clusters, L, generator=g, dtype=torch.float32) + Di = torch.randn(n_clusters, L, generator=g, dtype=torch.float32) + shell = torch.sort( + torch.randint(0, n_shells, (n_clusters,), generator=g))[0] + a, b, sect = legendre_recurrence_coefficients( + L, torch.float32, torch.device("cpu")) + n_even = (L - 1 if (L - 1) % 2 == 0 else L - 2) // 2 + return dict(shape=(n_even, n_shells, L), args=(cos_t, sin_t, Dr, Di, shell, + a, b, sect)) + + +def _run(fn, case): + Tr = torch.zeros(case["shape"], dtype=torch.float32) + Ti = torch.zeros_like(Tr) + fn(Tr, Ti, *case["args"]) + return Tr, Ti + + +@pytest.mark.parametrize("L,n_clusters,n_shells", [(13, 500, 40), + (65, 4000, 300), + (101, 3000, 250)]) +def test_fused_agrees_with_portable(L, n_clusters, n_shells): + if not fused.available(): + pytest.skip(f"fused kernel unavailable: {fused.why_unavailable()}") + case = _case(L, n_clusters, n_shells, seed=4) + ref_r, ref_i = _run(portable.legendre_shell_accumulate, case) + got_r, got_i = _run(fused.legendre_shell_accumulate, case) + for name, ref, got in (("real", ref_r, got_r), ("imag", ref_i, got_i)): + rel = (got - ref).abs().max().item() / max(ref.abs().max().item(), 1e-30) + assert rel < _TOL, f"L={L} {name} part differs by {rel:.2e}" + + +def test_float64_is_refused_not_reinterpreted(): + """A float64 caller must raise, naming the dtype.""" + if not fused.available(): + pytest.skip(f"fused kernel unavailable: {fused.why_unavailable()}") + L, n_clusters, n_shells = 9, 20, 5 + n_even = (L - 1) // 2 + d = torch.float64 + with pytest.raises(RuntimeError, match="float32 only"): + fused.legendre_shell_accumulate( + torch.zeros(n_even, n_shells, L, dtype=d), + torch.zeros(n_even, n_shells, L, dtype=d), + torch.zeros(n_clusters, dtype=d), torch.zeros(n_clusters, dtype=d), + torch.zeros(n_clusters, L, dtype=d), + torch.zeros(n_clusters, L, dtype=d), + torch.zeros(n_clusters, dtype=torch.long), + torch.zeros(L, L, dtype=d), torch.zeros(L, L, dtype=d), + torch.zeros(L, dtype=d)) + + +def test_shell_offsets_partition_the_clusters(): + """The kernel's work split: contiguous, complete, and in shell order.""" + shell = torch.tensor([0, 0, 2, 2, 2, 5], dtype=torch.long) + off = fused.shell_offsets(shell, 6) + assert off.tolist() == [0, 2, 2, 5, 5, 5, 6] + assert int(off[-1]) == shell.numel() + + +def test_the_dispatch_prefers_the_fused_kernel_when_it_builds(): + """Whatever the table selects is what the expansion runs.""" + from torchref.experimental.alignment.frf._backends import LEGENDRE_BACKENDS + from torchref.utils.backends import select + + probe = [torch.zeros(1, 1, 9, dtype=torch.float32)] * 6 + chosen = select(LEGENDRE_BACKENDS, probe).name + expected = "cpu_fused" if fused.available() else "portable" + assert chosen == expected, ( + f"table chose {chosen!r} but the fused kernel is " + f"{'available' if fused.available() else 'unavailable'}" + ) diff --git a/torchref/experimental/alignment/frf/_backends.py b/torchref/experimental/alignment/frf/_backends.py new file mode 100644 index 00000000..779b16a6 --- /dev/null +++ b/torchref/experimental/alignment/frf/_backends.py @@ -0,0 +1,64 @@ +"""Dispatch policy for the Legendre-recurrence-and-shell-accumulation stage. + +One row per kernel, with every criterion for choosing it as a field. See +:mod:`torchref.utils.backends` for what each field means, and +:mod:`torchref.base.electron_density._backends` for the table this follows. + +Both kernels take the same arguments:: + + (Tr, Ti, rep_cos, rep_sin, Dr, Di, shell, a_coef, b_coef, sect) + +and both accumulate into ``Tr``/``Ti`` in place, which is what lets the dispatch +site be a single call. +""" + +from __future__ import annotations + +import torch + +from torchref.utils.backends import Backend, BackendTable + +_CPU = "torchref.experimental.alignment.frf.kernels.cpu.legendre_shell" +_PORTABLE = "torchref.experimental.alignment.frf.kernels.portable" + +#: Argument positions carrying the device/dtype contract: the two accumulators and +#: the four per-cluster float arrays. ``shell`` is int64 and the three coefficient +#: tables are built by the caller at the working dtype, so probing them would only +#: restate what the caller already chose. +_FLOAT_ARGS = (0, 1, 2, 3, 4, 5) + +LEGENDRE_BACKENDS = BackendTable( + name="FRF Legendre/shell accumulation", + backends=( + Backend( + name="cpu_fused", + kernel=(_CPU, "legendre_shell_accumulate", "legendre_shell_accumulate"), + device="cpu", + dtypes=(torch.float32,), + # The kernel reads every array through a raw `float*`, so a mixed-dtype + # call would reinterpret the buffer rather than convert it. The gate + # keeps that from reaching the kernel; the kernel checks it too, since + # a table row is easier to widen by accident than a TORCH_CHECK. + require_uniform_dtype=True, + probes=_FLOAT_ARGS, + probe=(_CPU, "why_unavailable"), + # A compiler is not guaranteed on every host, and a missing one is a + # performance problem, not an outage. + expect_available="never", + # NOT "degrade": the kernel accumulates into Tr/Ti as it goes, so a + # mid-run failure leaves them partly written and the portable path + # would add its own contributions on top. There is nothing to fall + # back to once it has started. + on_failure="raise", + second_order=False, + ), + Backend( + name="portable", + kernel=(_PORTABLE, "legendre_shell_accumulate", + "legendre_shell_accumulate"), + second_order=False, + ), + ), +) + +__all__ = ["LEGENDRE_BACKENDS"] diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index 86bff89e..79bc8e92 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -58,7 +58,9 @@ #: approximation and needs its own evidence. _GROUP_SCALE_COS = 10_000_000 -from ..sh import LEGENDRE_SEED, legendre_recurrence_coefficients +from ....utils.backends import run_or_degrade, select +from ..sh import legendre_recurrence_coefficients +from ._backends import LEGENDRE_BACKENDS from .types import BesselSHCoefficients @@ -416,32 +418,18 @@ def _group_mean(values, index, n_groups): Tr = torch.zeros((n_even, nb, L), dtype=comp_real, device=device) Ti = torch.zeros((n_even, nb, L), dtype=comp_real, device=device) - cos_e = rep_cos[cs:ce].unsqueeze(-1) # (chunk, 1) - sin_c = rep_sin[cs:ce] # (chunk,) - Dr = Dp[cs:ce].real # (chunk, L) - Di = Dp[cs:ce].imag - prev2 = torch.zeros((ce - cs, L), dtype=comp_real, device=device) - prev1 = torch.zeros((ce - cs, L), dtype=comp_real, device=device) - prev1[:, 0] = LEGENDRE_SEED # bar_P_0^0 - for l in range(1, L): - # `a_coef` and `b_coef` are zero for m >= l, so this runs at full - # width; slicing to [:l] instead makes every iteration a differently - # shaped, mostly tiny kernel. - cur = a_coef[l] * cos_e * prev1 - b_coef[l] * prev2 - # The sectoral term must land BEFORE the products below are formed: - # it is the m = l entry of this very row. Computing `cur * Dr` first - # silently drops that entry for every even l -- which moved 3K7M's - # truth rank from 8 to 13 when a refactor got the order wrong. - cur[:, l] = sect[l] * sin_c * prev1[:, l - 1] # sectoral m = l - if l >= 2 and (l % 2 == 0): - if _PROFILE: - prof["legendre"] += _tick(t0); t0 = time.perf_counter() - pos = (l - 2) // 2 - Tr[pos].index_add_(0, sh, cur * Dr) - Ti[pos].index_add_(0, sh, cur * Di) - if _PROFILE: - prof["scatter"] += _tick(t0); t0 = time.perf_counter() - prev2, prev1 = prev1, cur + # Legendre recurrence and per-shell accumulation. Both stages are + # memory-bound in plain torch -- every row of the recurrence makes a + # round trip -- so this dispatches to a fused kernel where the row stays + # in cache, and falls back to the torch reference when none is built. + # `shell_of_cluster` is sorted, which the fused kernel requires: it + # partitions work by shell so that no two threads write the same + # accumulator row. + args = (Tr, Ti, rep_cos[cs:ce], rep_sin[cs:ce], + Dp[cs:ce].real.contiguous(), Dp[cs:ce].imag.contiguous(), + sh, a_coef, b_coef, sect) + backend = select(LEGENDRE_BACKENDS, args[:6]) + run_or_degrade(LEGENDRE_BACKENDS, backend, False, *args) if _PROFILE: prof["legendre"] += _tick(t0); t0 = time.perf_counter() diff --git a/torchref/experimental/alignment/frf/kernels/__init__.py b/torchref/experimental/alignment/frf/kernels/__init__.py new file mode 100644 index 00000000..e92588d1 --- /dev/null +++ b/torchref/experimental/alignment/frf/kernels/__init__.py @@ -0,0 +1 @@ +"""Kernels for the fast rotation function's spherical-harmonic expansion.""" diff --git a/torchref/experimental/alignment/frf/kernels/cpu/__init__.py b/torchref/experimental/alignment/frf/kernels/cpu/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py b/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py new file mode 100644 index 00000000..57d459e2 --- /dev/null +++ b/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py @@ -0,0 +1,295 @@ +"""Fused Legendre recurrence and shell accumulation, as one C++ kernel. + +The portable version runs the vertical recurrence as one torch operation per +``l`` and then scatters the row, so every row makes a round trip to memory. At +L=101 over 4.4e5 clusters that is ~108 GB for the recurrence and ~71 GB for the +scatter, both measured at ~50 GB/s -- the stages are bandwidth-bound, and the +arithmetic underneath is a small fraction of the time. + +float32 throughout, matching the rest of this codebase's kernels. The radial +Bessel recurrence is a separate stage and keeps its float64 internals, where the +downward recurrence's cancellation actually needs them. + +Fusing them removes the round trip: one cluster's three rows are 1.2 kB of stack, +so ``cur`` is produced, multiplied and accumulated without ever reaching memory. +Two further things fall out of writing it as a loop nest: + +* **Ragged widths are free.** ``bar_P[l, m]`` is zero for m > l, so step ``l`` + needs only columns 0..l -- ``for (m = 0; m <= l; ++m)`` and nothing more. In + torch the same saving needs narrowed views, and that was measured *slower*, + because a strided scatter target costs more than the zeros it skips. +* **No atomics.** The clusters arrive sorted by shell, so a thread that owns a + range of shells owns every write into those shells' rows. Parallelising over + clusters instead would race on the shared accumulator. + +The accumulator rows for one shell are ``n_even * L`` scalars -- 40 kB at L=101 -- +so they stay in cache across that shell's clusters, which is the point of +grouping by shell in the first place. +""" + +from __future__ import annotations + +from typing import Optional, Tuple + +import torch + +from torchref.base.electron_density.kernels.cpu._cpp_build import build_extension + +_CPP_SRC = r""" +#include +#include +#ifdef _OPENMP +#include +#else +#include +#endif + +// One shell's worth of work: every cluster in [c0, c1) contributes +// T[pos][s][m] += barP(l, m) * D[c][m] for even l >= 2, m <= l +// with barP built by the vertical recurrence in registers/stack. +template +static void shell_range( + int64_t s_begin, int64_t s_end, + const int64_t* __restrict off, + const int64_t* __restrict shell, + const scalar_t* __restrict rep_cos, + const scalar_t* __restrict rep_sin, + const scalar_t* __restrict Dr, + const scalar_t* __restrict Di, + const scalar_t* __restrict a_coef, + const scalar_t* __restrict b_coef, + const scalar_t* __restrict sect, + scalar_t* __restrict Tr, + scalar_t* __restrict Ti, + int64_t L, int64_t nb, int64_t n_even, scalar_t seed) { + + std::vector buf(3 * L, scalar_t(0)); + scalar_t* prev2 = buf.data(); + scalar_t* prev1 = buf.data() + L; + scalar_t* cur = buf.data() + 2 * L; + + for (int64_t s = s_begin; s < s_end; ++s) { + for (int64_t c = off[s]; c < off[s + 1]; ++c) { + const int64_t row = shell[c]; // == s, carried explicitly + const scalar_t co = rep_cos[c]; + const scalar_t si = rep_sin[c]; + const scalar_t* dr = Dr + c * L; + const scalar_t* di = Di + c * L; + + for (int64_t m = 0; m < L; ++m) { prev1[m] = scalar_t(0); prev2[m] = scalar_t(0); } + prev1[0] = seed; // bar_P_0^0 + + for (int64_t l = 1; l < L; ++l) { + const scalar_t* a = a_coef + l * L; + const scalar_t* b = b_coef + l * L; + // Vertical recurrence, only where the row can be non-zero. + for (int64_t m = 0; m < l; ++m) { + cur[m] = a[m] * co * prev1[m] - b[m] * prev2[m]; + } + // Sectoral m == l, which MUST be in place before the products below: + // it is this row's diagonal entry. + cur[l] = sect[l] * si * prev1[l - 1]; + + if (l >= 2 && (l % 2) == 0) { + const int64_t pos = (l - 2) / 2; + scalar_t* tr = Tr + (pos * nb + row) * L; + scalar_t* ti = Ti + (pos * nb + row) * L; + for (int64_t m = 0; m <= l; ++m) { + tr[m] += cur[m] * dr[m]; + ti[m] += cur[m] * di[m]; + } + } + // Rotate the three buffers; nothing is copied. + scalar_t* t = prev2; prev2 = prev1; prev1 = cur; cur = t; + } + } + } +} + +void legendre_shell_accumulate( + torch::Tensor Tr, torch::Tensor Ti, + torch::Tensor rep_cos, torch::Tensor rep_sin, + torch::Tensor Dr, torch::Tensor Di, + torch::Tensor shell, torch::Tensor offsets, + torch::Tensor a_coef, torch::Tensor b_coef, torch::Tensor sect, + double seed) { + + TORCH_CHECK(Tr.is_contiguous() && Ti.is_contiguous(), "T must be contiguous"); + TORCH_CHECK(rep_cos.scalar_type() == Tr.scalar_type() + && rep_sin.scalar_type() == Tr.scalar_type() + && Dr.scalar_type() == Tr.scalar_type() + && Di.scalar_type() == Tr.scalar_type() + && a_coef.scalar_type() == Tr.scalar_type() + && b_coef.scalar_type() == Tr.scalar_type() + && sect.scalar_type() == Tr.scalar_type(), + "every array must share the accumulator's dtype"); + TORCH_CHECK(Dr.is_contiguous() && Di.is_contiguous(), "D must be contiguous"); + TORCH_CHECK(shell.scalar_type() == torch::kLong, "shell must be int64"); + TORCH_CHECK(offsets.scalar_type() == torch::kLong, "offsets must be int64"); + + const int64_t n_even = Tr.size(0); + const int64_t nb = Tr.size(1); + const int64_t L = Tr.size(2); + TORCH_CHECK(offsets.numel() == nb + 1, "offsets must have n_shells + 1 entries"); + + // float32 only, by policy: this codebase has no float64 kernels. The caller + // is checked rather than dispatched on, so a float64 accumulator is a loud + // error instead of a silent reinterpretation of the buffer. + TORCH_CHECK(Tr.scalar_type() == torch::kFloat, + "legendre_shell_accumulate is float32 only, got ", Tr.scalar_type()); + { + using scalar_t = float; + const int64_t* off = offsets.data_ptr(); + const int64_t* sh = shell.data_ptr(); + const scalar_t* rc = rep_cos.data_ptr(); + const scalar_t* rs = rep_sin.data_ptr(); + const scalar_t* dr = Dr.data_ptr(); + const scalar_t* di = Di.data_ptr(); + const scalar_t* ac = a_coef.data_ptr(); + const scalar_t* bc = b_coef.data_ptr(); + const scalar_t* sc = sect.data_ptr(); + scalar_t* tr = Tr.data_ptr(); + scalar_t* ti = Ti.data_ptr(); + const scalar_t sd = static_cast(seed); + +#ifdef _OPENMP + // Dynamic, because clusters per shell varies (measured 2.7 to 39 across the + // benchmark) so equal shell counts are not equal work. +#pragma omp parallel for schedule(dynamic, 8) + for (int64_t s = 0; s < nb; ++s) { + shell_range(s, s + 1, off, sh, rc, rs, dr, di, ac, bc, sc, + tr, ti, L, nb, n_even, sd); + } +#else + // Apple Clang rejects -fopenmp, so carve the shells into contiguous blocks. + int nthreads = std::max(1u, std::thread::hardware_concurrency()); + if (nthreads > nb) nthreads = static_cast(std::max(nb, 1)); + std::vector pool; + const int64_t per = (nb + nthreads - 1) / std::max(nthreads, 1); + for (int t = 0; t < nthreads; ++t) { + const int64_t s0 = t * per; + const int64_t s1 = std::min(nb, s0 + per); + if (s0 >= s1) break; + pool.emplace_back([=] { + shell_range(s0, s1, off, sh, rc, rs, dr, di, ac, bc, sc, + tr, ti, L, nb, n_even, sd); + }); + } + for (auto& th : pool) th.join(); +#endif + } +} + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { + m.def("legendre_shell_accumulate", &legendre_shell_accumulate, + "Fused Legendre recurrence and per-shell accumulation"); +} +""" + +_module = None +_module_failed = False +_module_error: Optional[Tuple[str, str]] = None + + +def _get_module(): + """The compiled extension, or None if it could not be built.""" + global _module, _module_failed, _module_error + if _module is not None: + return _module + if _module_failed: + return None + _module, _module_error = build_extension("frf_legendre_shell", _CPP_SRC) + if _module is None: + _module_failed = True + return _module + + +def why_unavailable() -> Optional[str]: + """``None`` if the fused kernel is usable, else why it is not. + + The single availability probe for this backend, read by + :mod:`torchref.utils.backends`. A missing compiler and a compile error are + different problems, and the captured diagnostic is what separates them. + """ + if _get_module() is not None: + return None + reason = _module_error[0] if _module_error else "unknown reason" + return ( + f"the fused CPU Legendre/shell kernel is not available ({reason}); see " + "torchref.experimental.alignment.frf.kernels.cpu.legendre_shell." + "last_error()" + ) + + +def available() -> bool: + """Whether the fused kernel compiled and is ready to dispatch.""" + return why_unavailable() is None + + +def last_error() -> Optional[Tuple[str, str]]: + """``(message, traceback)`` from the last failed build attempt, if any.""" + _get_module() + return _module_error + + +def clear_cache() -> None: + """Forget the build result, so the next call retries. For tests.""" + global _module, _module_failed, _module_error + _module, _module_failed, _module_error = None, False, None + + +def shell_offsets(shell: torch.Tensor, n_shells: int) -> torch.Tensor: + """Start index of each shell in a shell-sorted cluster array, plus the end. + + ``(n_shells + 1,)`` int64. The kernel needs the ranges rather than the + per-cluster labels so that a thread can own a set of shells outright and + write their accumulator rows without atomics. + """ + counts = torch.bincount(shell, minlength=n_shells) + offsets = torch.zeros(n_shells + 1, dtype=torch.long, device=shell.device) + torch.cumsum(counts, dim=0, out=offsets[1:]) + return offsets + + +def legendre_shell_accumulate( + Tr: torch.Tensor, + Ti: torch.Tensor, + rep_cos: torch.Tensor, + rep_sin: torch.Tensor, + Dr: torch.Tensor, + Di: torch.Tensor, + shell: torch.Tensor, + a_coef: torch.Tensor, + b_coef: torch.Tensor, + sect: torch.Tensor, +) -> None: + """Fused recurrence and accumulation, in place on ``Tr``/``Ti``. + + Same signature and same effect as + :func:`torchref.experimental.alignment.frf.kernels.portable.legendre_shell_accumulate`. + ``shell`` must be sorted non-decreasing -- the kernel partitions work by + shell to avoid atomics, and unsorted input would silently drop + contributions rather than merely run slowly. + """ + from ....sh import LEGENDRE_SEED + + module = _get_module() + if module is None: + raise RuntimeError(why_unavailable()) + offsets = shell_offsets(shell, Tr.shape[1]) + module.legendre_shell_accumulate( + Tr, Ti, rep_cos.contiguous(), rep_sin.contiguous(), + Dr.contiguous(), Di.contiguous(), shell.contiguous(), offsets, + a_coef.contiguous(), b_coef.contiguous(), sect.contiguous(), + float(LEGENDRE_SEED), + ) + + +__all__ = [ + "available", + "clear_cache", + "last_error", + "legendre_shell_accumulate", + "shell_offsets", + "why_unavailable", +] diff --git a/torchref/experimental/alignment/frf/kernels/portable.py b/torchref/experimental/alignment/frf/kernels/portable.py new file mode 100644 index 00000000..7bcc2db3 --- /dev/null +++ b/torchref/experimental/alignment/frf/kernels/portable.py @@ -0,0 +1,75 @@ +"""Portable Legendre-recurrence-and-shell-accumulation, in plain torch. + +The reference the fused kernel is checked against, and the fallback whenever that +kernel cannot be built. One step of the vertical recurrence per ``l``, then the +row is multiplied by the azimuthal sums and scattered into its cluster's shell. + +Both stages are memory-bound: at L=101 over 4.4e5 clusters the recurrence moves +about 108 GB and the scatter about 71 GB, and both measure ~50 GB/s, which is +roughly what four threads get from a server memory controller. The arithmetic is +a small fraction of that -- which is the whole reason for a fused kernel, where +the row never leaves cache. +""" + +from __future__ import annotations + +import torch + +from ...sh import LEGENDRE_SEED + + +def legendre_shell_accumulate( + Tr: torch.Tensor, + Ti: torch.Tensor, + rep_cos: torch.Tensor, + rep_sin: torch.Tensor, + Dr: torch.Tensor, + Di: torch.Tensor, + shell: torch.Tensor, + a_coef: torch.Tensor, + b_coef: torch.Tensor, + sect: torch.Tensor, +) -> None: + """Accumulate ``sum_c barP[c, l, m] * D[c, m]`` into ``Tr``/``Ti``, in place. + + Parameters + ---------- + Tr, Ti : torch.Tensor + ``(n_even, n_shells, L)`` real accumulators, added into. + rep_cos, rep_sin : torch.Tensor + ``(n_clusters,)`` cos and sin of the polar angle, per cluster. + Dr, Di : torch.Tensor + ``(n_clusters, L)`` real and imaginary azimuthal sums. + shell : torch.Tensor + ``(n_clusters,)`` int64 shell index of each cluster, into ``Tr``'s middle + axis. + a_coef, b_coef : torch.Tensor + ``(L, L)`` recurrence coefficients, zero for ``m >= l``. + sect : torch.Tensor + ``(L,)`` sectoral factors. + """ + L = Tr.shape[-1] + cos_e = rep_cos.unsqueeze(-1) + prev2 = torch.zeros_like(Dr) + prev1 = torch.zeros_like(Dr) + prev1[:, 0] = LEGENDRE_SEED # bar_P_0^0 + for l in range(1, L): + # `a_coef` and `b_coef` are zero for m >= l, so this runs at full width. + # Narrowing it to the l+1 columns that can be non-zero was measured and + # is slower: the saving is real (the summed width at L=101 falls from + # 10100 to 6481) but a strided scatter target costs more than it, and the + # recurrence did not speed up at all -- it is not arithmetic-bound. The + # fused kernel gets the ragged widths for free, as loop bounds. + cur = a_coef[l] * cos_e * prev1 - b_coef[l] * prev2 + # The sectoral term must land BEFORE the products below are formed: it is + # the m = l entry of this very row. Forming `cur * Dr` first silently + # drops that entry for every even l. + cur[:, l] = sect[l] * rep_sin * prev1[:, l - 1] + if l >= 2 and (l % 2 == 0): + pos = (l - 2) // 2 + Tr[pos].index_add_(0, shell, cur * Dr) + Ti[pos].index_add_(0, shell, cur * Di) + prev2, prev1 = prev1, cur + + +__all__ = ["legendre_shell_accumulate"] diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index 99814054..c8c06e52 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -310,11 +310,14 @@ def search_peaks( model_radius_A=model_radius_A, auto_lmax=True, lmax_cap=LMAX_CAP, - # The spherical-harmonic contraction dominates the runtime and is - # rate-limited in double precision on accelerators. Its float32 path - # keeps the Bessel recurrence and the cross-chunk accumulator at - # full precision. - compute_dtype=torch.complex64 if device.type == "cuda" else None, + # The angular half of the expansion runs in single precision on + # every device. It is the runtime bottleneck, it is memory-bound, and + # single precision is this codebase's kernel dtype -- a float64-only + # path would make the fused CPU kernel unreachable. The radial Bessel + # recurrence keeps its float64 internals, where the downward + # recurrence's cancellation needs them, and the cross-chunk + # accumulator stays at full precision. + compute_dtype=torch.complex64, ) _arf, peaks = engine.score_model( s_calc, F_calc, n_peaks=n_peaks, From 5f0a782d098f2b4b7546bd61079a01111d2650bd Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 21 Aug 2026 19:03:24 +0200 Subject: [PATCH 035/250] Add a stage-breakdown driver, now that the expansion is not the bottleneck `where_now.sh` runs the benchmark harness over both bandwidths and prints the exclusive per-stage split. Worth having as its own entry point: the answer moved a long way during this work, and two of my predictions about it were wrong -- the FFT path was expected to dominate at L=100 and is 0.6%, while French-Wilson was not on the list at all and is 27%. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/where_now.sh | 20 ++++++++++++++++++++ 1 file changed, 20 insertions(+) create mode 100644 alignment_lab/analysis/where_now.sh diff --git a/alignment_lab/analysis/where_now.sh b/alignment_lab/analysis/where_now.sh new file mode 100644 index 00000000..c3200601 --- /dev/null +++ b/alignment_lab/analysis/where_now.sh @@ -0,0 +1,20 @@ +#!/bin/bash +# Full stage breakdown of one rotation search, now that the SH expansion is no +# longer the dominant term. +#SBATCH --job-name=frf_where +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 +export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" +for pdb in 3K7M 1DAW; do + "$PY" -u -m alignment_lab.diagnostics.frf_benchmark \ + --pdb "$pdb" --arms cap100,cap64 --trials 2 2>&1 \ + | grep -vE "Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|warn|^ *$" +done +echo "exit_code=$?" From 3b4e283826ccd1350401da4766caca257e3bc187 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Sat, 22 Aug 2026 23:45:58 +0200 Subject: [PATCH 036/250] Stop the dense model transform building a grid it discards dense_calc_via_box takes a defensive model.copy(), then replaces max_res, the spacegroup (P1) and the cell on the next three lines. Each of those setters rebuilds the FFT submodule, so the grid copy() builds at the end is orphaned -- along with the map-symmetry operator, which precomputes one (nx, ny, nz, 3) sampling grid per symmetry operation and dominates the cost. ModelFT.copy gains build_grid, and the dense calc passes False. Measured on 3K7M at the shipped lmax_cap=64: the stage was half the whole rotation search. |F_calc| is bit-identical on 3K7M and 1DAW at both caps (up to 1.29e6 reflections) and truth ranks are unchanged, as they must be for a change that only removes discarded work. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- docs/changelog.rst | 1 + tests/unit/model/test_copy.py | 53 +++++++++++++++++++ .../experimental/alignment/frf/dense_calc.py | 6 ++- torchref/model/model_ft.py | 15 ++++-- 4 files changed, 71 insertions(+), 4 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 6dc08ad1..62c7306a 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -13,6 +13,7 @@ Unreleased - Fixed ``Model.copy()`` registering the original's space group as a second submodule of the copy - Replaced the fast rotation function's keyword surface with ``rotation_search(model, data, model_error_A)``; the caller's coordinate error is now used rather than overwritten by an estimate from the atom count - Removed the rotation function's dead modules, engine variants, debug environment switches and unreachable knobs +- Fixed the rotation function's dense model transform building a real-space grid and map-symmetry operator that its next three lines discarded; ``ModelFT.copy`` gained ``build_grid`` Version 0.6.4 ---------- diff --git a/tests/unit/model/test_copy.py b/tests/unit/model/test_copy.py index 6f49c23d..dbb12fd2 100644 --- a/tests/unit/model/test_copy.py +++ b/tests/unit/model/test_copy.py @@ -99,3 +99,56 @@ def test_copy_owns_its_spacegroup(cls_name, mixed_adp_path): assert "_spacegroup" in c._modules stray = [k for k in c.state_dict() if k.startswith("spacegroup.")] assert stray == [], f"copy registered a second space group: {stray}" + + +@pytest.mark.unit +def test_copy_can_skip_the_grid_build(mixed_adp_path): + """``build_grid=False`` skips the grid, and setting cell+spacegroup restores it. + + Building the grid also builds the map-symmetry operator, which precomputes + one sampling grid per symmetry operation over the whole map. A caller that + is about to replace the cell, the spacegroup or ``max_res`` would have that + work thrown away, because each of those setters rebuilds the FFT submodule. + """ + from torchref.model import ModelFT + + m = _load(ModelFT, mixed_adp_path) + assert m._fft is not None and m._fft.real_space_grid is not None + + lean = m.copy(build_grid=False) + assert lean._fft is not None + assert lean._fft.real_space_grid is None + assert lean._fft.map_symmetry is None + + full = m.copy() + assert full._fft.real_space_grid is not None + + # The skipped grid is recoverable: this is what the cell setter triggers. + lean.setup_grid(max_res=m.max_res) + assert lean._fft.real_space_grid is not None + assert torch.equal(lean._fft.gridsize, full._fft.gridsize) + + +@pytest.mark.unit +def test_skipping_the_grid_build_does_not_change_structure_factors(mixed_adp_path): + """The two copies must give identical ``F_calc`` once each has a grid. + + ``build_grid`` is a pure waste-removal switch: the grid it skips is rebuilt + by the cell/spacegroup setters before any structure factor is computed, so + no amplitude may depend on it. + """ + from torchref.model import ModelFT + + m = _load(ModelFT, mixed_adp_path) + hkl = torch.tensor( + [[1, 0, 0], [0, 1, 0], [0, 0, 1], [2, 1, 3], [5, -2, 1], [7, 7, 7]], + dtype=torch.long, + ) + + with torch.no_grad(): + f_full = m.copy().get_structure_factor(hkl, recalc=True) + lean = m.copy(build_grid=False) + lean.setup_grid(max_res=m.max_res) + f_lean = lean.get_structure_factor(hkl, recalc=True) + + assert torch.equal(f_lean, f_full) diff --git a/torchref/experimental/alignment/frf/dense_calc.py b/torchref/experimental/alignment/frf/dense_calc.py index 71928c4c..b241d19d 100644 --- a/torchref/experimental/alignment/frf/dense_calc.py +++ b/torchref/experimental/alignment/frf/dense_calc.py @@ -56,7 +56,11 @@ def dense_calc_via_box( """ from torchref.symmetry.cell import Cell - m = model.copy() # isolate the box mutation from the caller + # Isolate the box mutation from the caller. No grid: the three setters + # below replace the FFT submodule, so a grid built for the crystal cell and + # the data resolution would be discarded -- along with the per-operation + # map-symmetry sampling grids that dominate the cost of building it. + m = model.copy(build_grid=False) with torch.no_grad(): coords = m.xyz() dev = coords.device diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index d9ca4df8..d8e4ce6f 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -827,7 +827,7 @@ def forward(self, hkl, apply_anomalous: bool = True) -> torch.Tensor: return sf - def copy(self, detach: bool = True) -> "ModelFT": + def copy(self, detach: bool = True, build_grid: bool = True) -> "ModelFT": """ Create a deep copy of the ModelFT. @@ -841,11 +841,20 @@ def copy(self, detach: bool = True) -> "ModelFT": detach : bool, optional If True, the copy's parameters will be detached from the computation graph (default: True). + build_grid : bool, optional + If True (default), give the copy a real-space grid whenever the + original has one. Building it also builds the map-symmetry operator, + which precomputes one sampling grid per symmetry operation over the + whole map. Pass False when the caller is about to replace the cell, + spacegroup or ``max_res``: each of those setters rebuilds the FFT + submodule, so a grid built here would be discarded unused. Returns ------- ModelFT - A new, fully independent ModelFT instance with copied data. + A new, fully independent ModelFT instance with copied data. With + ``build_grid=False`` it has no real-space grid until the cell and + spacegroup are set and :meth:`setup_grid` runs. """ if not self.initialized: raise RuntimeError("Cannot copy an uninitialized ModelFT. Load data first.") @@ -906,7 +915,7 @@ def copy(self, detach: bool = True) -> "ModelFT": if self._fft is not None: model_copy._fft = self._fft.copy() - if self._fft.real_space_grid is not None: + if build_grid and self._fft.real_space_grid is not None: model_copy.setup_grid(max_res=self.max_res) # Don't share cached structure factors with the original. From d611ae58258183f3496e378ab6161b7172e4890c Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Sat, 22 Aug 2026 23:46:07 +0200 Subject: [PATCH 037/250] Add the measurement scripts behind the dense-calc and Wigner findings Where each one leads: - dense_internals.sh: the FFT grid is not what the dense calc costs -- an 8x finer grid is 81 ms, while the defensive copy is 1000 ms - copy_cost.sh, copyfix_timing.sh: the copy breakdown and the payoff - wigner_split.sh, wigner_cache_probe.sh: 80% of the Wigner contraction is the data-independent d-block build; the table is 88 MB float32 at L=65 - gpfs_cold_read.sh: a cold cross-node read of that table is 17 ms, so GPFS bandwidth is not the obstacle -- writer and reader must be different nodes or the page cache makes the measurement meaningless - wigner_fft_proto.sh: the integer-eigenvalue Fourier form is correct to 5e-15 but 0.7-0.9x, being bandwidth-bound where the batched matmul is compute-bound - wigner_mirror_proto.sh: the beta -> pi - beta mirror reaches only 1.28x, and neither candidate sign convention reproduces d(pi - beta) Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/copy_cost.sh | 49 +++++++++++ alignment_lab/analysis/copyfix_timing.sh | 66 ++++++++++++++ alignment_lab/analysis/dense_and_wigner.sh | 68 ++++++++++++++ alignment_lab/analysis/dense_internals.sh | 78 ++++++++++++++++ alignment_lab/analysis/gpfs_cold_read.sh | 35 ++++++++ alignment_lab/analysis/run_tests_copyfix.sh | 35 ++++++++ .../analysis/verify_copy_correctness.sh | 80 +++++++++++++++++ alignment_lab/analysis/verify_copy_fix.sh | 79 +++++++++++++++++ alignment_lab/analysis/wigner_cache_probe.sh | 88 +++++++++++++++++++ alignment_lab/analysis/wigner_fft_proto.sh | 73 +++++++++++++++ alignment_lab/analysis/wigner_mirror_proto.sh | 79 +++++++++++++++++ alignment_lab/analysis/wigner_split.sh | 44 ++++++++++ 12 files changed, 774 insertions(+) create mode 100644 alignment_lab/analysis/copy_cost.sh create mode 100644 alignment_lab/analysis/copyfix_timing.sh create mode 100644 alignment_lab/analysis/dense_and_wigner.sh create mode 100644 alignment_lab/analysis/dense_internals.sh create mode 100644 alignment_lab/analysis/gpfs_cold_read.sh create mode 100644 alignment_lab/analysis/run_tests_copyfix.sh create mode 100644 alignment_lab/analysis/verify_copy_correctness.sh create mode 100644 alignment_lab/analysis/verify_copy_fix.sh create mode 100644 alignment_lab/analysis/wigner_cache_probe.sh create mode 100644 alignment_lab/analysis/wigner_fft_proto.sh create mode 100644 alignment_lab/analysis/wigner_mirror_proto.sh create mode 100644 alignment_lab/analysis/wigner_split.sh diff --git a/alignment_lab/analysis/copy_cost.sh b/alignment_lab/analysis/copy_cost.sh new file mode 100644 index 00000000..3f106175 --- /dev/null +++ b/alignment_lab/analysis/copy_cost.sh @@ -0,0 +1,49 @@ +#!/bin/bash +# model.copy() is ~97% of dense_calc_via_box, which is ~50% of a shipped +# rotation search. What inside it costs the second, and is it steady state or +# first-call lazy initialisation? +#SBATCH --job-name=frf_copy +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname)" +"$PY" -u -c " +import copy as copy_module, time, torch +torch.set_grad_enabled(False) +from alignment_lab.lab.benchmark import load_case + +def reps(fn, n=5): + fn() + return [ (lambda t0: (fn(), time.perf_counter()-t0)[1])(time.perf_counter()) for _ in range(n) ] + +for name in ('3K7M', '1DAW'): + model, data = load_case(name) + model.verbose = 0 + g = model._fft.real_space_grid if model._fft is not None else None + print(f'--- {name}: {len(model.pdb)} atoms cell={[round(x,1) for x in model.cell.parameters_list()] if hasattr(model.cell,\"parameters_list\") else \"?\"} ' + f'max_res={model.max_res} grid={None if g is None else tuple(g.shape)}') + ts = reps(lambda: model.copy()) + print(f' model.copy() x5: ' + ' '.join(f'{t*1e3:.0f}' for t in ts) + ' ms') + + # The pieces, each timed on its own. + print(f' pdb.copy(deep=True) {min(reps(lambda: model.pdb.copy(deep=True)))*1e3:8.1f} ms') + if model._parametrization is not None: + print(f' deepcopy(_parametrization) {min(reps(lambda: copy_module.deepcopy(model._parametrization)))*1e3:8.1f} ms') + if model._fft is not None: + print(f' _fft.copy() {min(reps(lambda: model._fft.copy()))*1e3:8.1f} ms') + mc = model.copy() + print(f' setup_grid(max_res) {min(reps(lambda: mc.setup_grid(max_res=model.max_res)))*1e3:8.1f} ms') + print(f' _rebuild_sf_indices() {min(reps(lambda: model._rebuild_sf_indices()))*1e3:8.1f} ms') + print(f' cell.clone() {min(reps(lambda: model.cell.clone()))*1e3:8.1f} ms') + + # What the FRF actually needs the copy for: it mutates max_res, spacegroup, cell. + # Would a copy taken AFTER shrinking max_res be cheaper? + print(f' copy() with max_res=4.15 {min(reps(lambda: (lambda m: m)(model.copy())))*1e3:8.1f} ms (baseline)') +" 2>&1 | grep -vE "Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$" +echo "exit_code=$?" diff --git a/alignment_lab/analysis/copyfix_timing.sh b/alignment_lab/analysis/copyfix_timing.sh new file mode 100644 index 00000000..953f8c66 --- /dev/null +++ b/alignment_lab/analysis/copyfix_timing.sh @@ -0,0 +1,66 @@ +#!/bin/bash +# Payoff of build_grid=False. NOTE: whatever node this lands on, the absolute +# ms are NOT comparable with the EPYC 9335 tables -- only the before/after +# ratio measured inside this one job is. +#SBATCH --job-name=frf_ctime +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +FILT="Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$|No CUDA" +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" +echo "=== the copy itself, and the whole dense-calc stage ===" +"$PY" -u -c " +import math, time, torch +torch.set_grad_enabled(False) +from alignment_lab.lab.benchmark import load_case +from torchref.symmetry.cell import Cell +from torchref.experimental.alignment.frf.dense_calc import dense_calc_via_box, model_sf_abs +from torchref.experimental.alignment.frf.api import phaser_lmax_resolution + +def box_path(model, d_min, d_max, build_grid): + m = model.copy(build_grid=build_grid) + coords = m.xyz(); dev = coords.device + a = float(4.0 * (coords - coords.mean(0)).norm(dim=-1).max().item()) + m.max_res = float(d_min); m.spacegroup = 'P 1' + m.cell = Cell([a, a, a, 90., 90., 90.], device=dev) + nmax = int(math.ceil(a / d_min)) + idx = torch.arange(-nmax, nmax + 1, device=dev) + H, K, Lg = torch.meshgrid(idx, idx, idx, indexing='ij') + hkl = torch.stack([H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], -1).to(torch.long) + smag = hkl.to(torch.float64).norm(dim=-1) / a + hkl = hkl[(smag >= 1.0/d_max) & (smag <= 1.0/d_min)].contiguous() + return model_sf_abs(m, hkl) + +def best(fn, n=3): + fn() + return min((lambda: (lambda t0: (fn(), time.perf_counter()-t0)[1])(time.perf_counter()))() for _ in range(n)) + +for name in ('3K7M', '1DAW'): + model, data = load_case(name); model.verbose = 0 + rb = data.cell.reciprocal_basis_matrix.to(torch.float64) + d_min_data = 1.0 / (data.hkl.to(torch.float64) @ rb).norm(dim=-1).max().item() + r = float((model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item()) + L, d_min = phaser_lmax_resolution(r, d_min_data, 64) + tc_on = best(lambda: model.copy(build_grid=True)) + tc_off = best(lambda: model.copy(build_grid=False)) + t_on = best(lambda: box_path(model, d_min, 100.0, True)) + t_off = best(lambda: box_path(model, d_min, 100.0, False)) + print(f' {name} (cap64, d_min={d_min:.2f}A)') + print(f' model.copy() grid {tc_on*1e3:8.1f} -> no grid {tc_off*1e3:7.1f} ms ({tc_on/tc_off:5.1f}x)') + print(f' whole box path grid {t_on*1e3:8.1f} -> no grid {t_off*1e3:7.1f} ms ({t_on/t_off:5.1f}x)') + print(f' dense_calc_via_box now {best(lambda: dense_calc_via_box(model, 100.0, d_min, pad=2.0))*1e3:8.1f} ms') +" 2>&1 | grep -vE "$FILT" +echo "=== whole rotation search, per cap ===" +for pdb in 3K7M 1DAW; do + for arm in cap64 cap100; do + "$PY" -u -m alignment_lab.diagnostics.frf_benchmark \ + --pdb "$pdb" --arms "$arm" --trials 2 2>&1 | grep -vE "$FILT" + done +done +echo "done" diff --git a/alignment_lab/analysis/dense_and_wigner.sh b/alignment_lab/analysis/dense_and_wigner.sh new file mode 100644 index 00000000..6fc8e84b --- /dev/null +++ b/alignment_lab/analysis/dense_and_wigner.sh @@ -0,0 +1,68 @@ +#!/bin/bash +# Per-cap stage breakdown (the shipped config is cap64; earlier tables medianed +# cap64 and cap100 together), plus a breakdown INSIDE dense_calc_via_box. +#SBATCH --job-name=frf_dense +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname)" +FILT="Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$" +for pdb in 3K7M 1DAW; do + for arm in cap64 cap100; do + "$PY" -u -m alignment_lab.diagnostics.frf_benchmark \ + --pdb "$pdb" --arms "$arm" --trials 2 2>&1 | grep -vE "$FILT" + done +done +echo "=== inside dense_calc_via_box ===" +"$PY" -u -c " +import math, time, torch +torch.set_grad_enabled(False) +from alignment_lab.lab.benchmark import BENCHMARK, load_case +from torchref.symmetry.cell import Cell +from torchref.experimental.alignment.frf.dense_calc import model_sf_abs +from torchref.experimental.alignment.frf.sitelist_ang import phaser_lmax_resolution + +for name in ('3K7M', '1DAW'): + case = load_case(name) + model, data = case.model, case.data + d_min_data = 1.0 / data.hkl_to_s(data.hkl).norm(dim=-1).max().item() + d_max = 100.0 + r = float((model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item()) + for cap in (64, 100): + L, d_min = phaser_lmax_resolution(r, d_min_data, cap) + t = {} + t0 = time.perf_counter() + m = model.copy() + coords = m.xyz(); dev = coords.device + extent = (coords - coords.mean(0)).norm(dim=-1).max().item() + a = float(2.0 * 2.0 * extent) + t['copy'] = time.perf_counter() - t0 + t0 = time.perf_counter() + m.max_res = float(d_min); m.spacegroup = 'P 1' + m.cell = Cell([a, a, a, 90., 90., 90.], device=dev) + t['cell+grid setup'] = time.perf_counter() - t0 + t0 = time.perf_counter() + nmax = int(math.ceil(a / d_min)) + idx = torch.arange(-nmax, nmax + 1, device=dev) + H, K, Lg = torch.meshgrid(idx, idx, idx, indexing='ij') + hkl = torch.stack([H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], -1).to(torch.long) + smag = hkl.to(torch.float64).norm(dim=-1) / a + hkl = hkl[(smag >= 1.0/d_max) & (smag <= 1.0/d_min)].contiguous() + t['hkl enumerate'] = time.perf_counter() - t0 + model_sf_abs(m, hkl) # warm + t0 = time.perf_counter(); model_sf_abs(m, hkl); t['model_sf_abs'] = time.perf_counter()-t0 + grid = m.cell.compute_grid_size(m.max_res) + print(f'{name} cap{cap}: d_min={d_min:.2f}A box={a:.0f}A grid={grid} ' + f'spacing={a/grid[0]:.2f}A n_hkl={hkl.shape[0]} ' + f'(enumerated {(2*nmax+1)**3})') + tot = sum(t.values()) + for k, v in t.items(): + print(f' {k:18s} {v*1e3:8.1f} ms {100*v/tot:5.1f}%') +" 2>&1 | grep -vE "$FILT" +echo "exit_code=$?" diff --git a/alignment_lab/analysis/dense_internals.sh b/alignment_lab/analysis/dense_internals.sh new file mode 100644 index 00000000..2c361f78 --- /dev/null +++ b/alignment_lab/analysis/dense_internals.sh @@ -0,0 +1,78 @@ +#!/bin/bash +# dense_calc_via_box is ~50% of a shipped (cap64) rotation search. Where inside +# it does the time go, and does it actually track the FFT grid fineness? If the +# cost is insensitive to the grid, a coarser grid buys nothing and the +# artificial-B route is pointless. +#SBATCH --job-name=frf_dint +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname)" +"$PY" -u -c " +import math, time, torch +torch.set_grad_enabled(False) +from alignment_lab.lab.benchmark import load_case +from torchref.symmetry.cell import Cell +from torchref.experimental.alignment.frf.dense_calc import model_sf_abs +from torchref.experimental.alignment.frf.api import phaser_lmax_resolution + +def timeit(fn, n=3): + fn() + out = [] + for _ in range(n): + t0 = time.perf_counter(); fn(); out.append(time.perf_counter() - t0) + return min(out) + +for name in ('3K7M', '1DAW'): + model, data = load_case(name) + rb = data.cell.reciprocal_basis_matrix.to(torch.float64) + d_min_data = 1.0 / (data.hkl.to(torch.float64) @ rb).norm(dim=-1).max().item() + r = float((model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item()) + for cap in (64,): + L, d_min = phaser_lmax_resolution(r, d_min_data, cap) + t = {} + t0 = time.perf_counter() + m = model.copy() + t['model.copy()'] = time.perf_counter() - t0 + coords = m.xyz(); dev = coords.device + extent = (coords - coords.mean(0)).norm(dim=-1).max().item() + a = float(2.0 * 2.0 * extent) + t0 = time.perf_counter() + m.max_res = float(d_min); m.spacegroup = 'P 1' + m.cell = Cell([a, a, a, 90., 90., 90.], device=dev) + t['cell/grid setup'] = time.perf_counter() - t0 + t0 = time.perf_counter() + nmax = int(math.ceil(a / d_min)) + idx = torch.arange(-nmax, nmax + 1, device=dev) + H, K, Lg = torch.meshgrid(idx, idx, idx, indexing='ij') + hkl = torch.stack([H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], -1).to(torch.long) + smag = hkl.to(torch.float64).norm(dim=-1) / a + hkl = hkl[(smag >= 1.0/100.0) & (smag <= 1.0/d_min)].contiguous() + t['hkl enumerate'] = time.perf_counter() - t0 + t['model_sf_abs (1st)'] = timeit(lambda: model_sf_abs(m, hkl), n=1) + t['model_sf_abs (warm)'] = timeit(lambda: model_sf_abs(m, hkl)) + g = m.cell.compute_grid_size(m.max_res) + print(f'--- {name} cap{cap}: d_min={d_min:.2f}A box={a:.0f}A ' + f'grid={g[0]}^3 spacing={a/g[0]:.2f}A n_hkl={hkl.shape[0]} ' + f'of {(2*nmax+1)**3} enumerated') + for k, v in t.items(): + print(f' {k:22s} {v*1e3:8.1f} ms') + + # Does the cost track the grid at all? Same reflections, different grid. + print(' grid sensitivity (same hkl list, max_res only):') + for fac in (0.5, 1.0, 1.5, 2.0): + mm = model.copy() + mm.max_res = float(d_min / fac) # fac>1 => finer grid + mm.spacegroup = 'P 1' + mm.cell = Cell([a, a, a, 90., 90., 90.], device=dev) + gg = mm.cell.compute_grid_size(mm.max_res) + dt = timeit(lambda: model_sf_abs(mm, hkl)) + print(f' x{fac:<4} grid={gg[0]:4d}^3 spacing={a/gg[0]:5.2f}A {dt*1e3:8.1f} ms') +" 2>&1 | grep -vE "Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$" +echo "exit_code=$?" diff --git a/alignment_lab/analysis/gpfs_cold_read.sh b/alignment_lab/analysis/gpfs_cold_read.sh new file mode 100644 index 00000000..411ce2c0 --- /dev/null +++ b/alignment_lab/analysis/gpfs_cold_read.sh @@ -0,0 +1,35 @@ +#!/bin/bash +# Is reading the Wigner d-table off GPFS actually cheaper than recomputing it? +# Only a COLD read answers that, so the writer and the reader must be different +# nodes -- a same-node re-read is served from the page cache and is meaningless. +#SBATCH --job-name=frf_cold +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 PYTHONUNBUFFERED=1 +export CUDA_VISIBLE_DEVICES="" +MODE="$1" +DIR=alignment_lab/runs/gpfs_cold +mkdir -p "$DIR" alignment_lab/slurm +echo "mode=$MODE host=$(hostname)" +"$PY" -u -c " +import os, time, torch +mode, d = '$MODE', '$DIR' +for L, mb in ((65, 88), (101, 330)): + p = os.path.join(d, f'dtable_L{L}.pt') + if mode == 'write': + n = int(mb * 1e6 / 4) + torch.save(torch.zeros(n, dtype=torch.float32), p) + print(f' wrote L={L} {os.path.getsize(p)/1e6:.0f} MB') + else: + t0 = time.perf_counter(); torch.load(p, map_location='cpu') + cold = time.perf_counter() - t0 + t0 = time.perf_counter(); torch.load(p, map_location='cpu') + warm = time.perf_counter() - t0 + print(f' L={L} {os.path.getsize(p)/1e6:5.0f} MB cold {cold*1e3:7.0f} ms ' + f'warm {warm*1e3:6.0f} ms') +" +echo "exit_code=$?" diff --git a/alignment_lab/analysis/run_tests_copyfix.sh b/alignment_lab/analysis/run_tests_copyfix.sh new file mode 100644 index 00000000..2fa5742e --- /dev/null +++ b/alignment_lab/analysis/run_tests_copyfix.sh @@ -0,0 +1,35 @@ +#!/bin/bash +# Test gate for the build_grid change. rc is captured immediately after pytest, +# not after a pipe: $? on a pipeline reads the LAST command, which silently +# reports success for a failed test run. +#SBATCH --job-name=frf_tests +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname)" + +echo "###### fast: model + alignment + frf_separate" +"$PY" -m pytest -q --no-header -p no:cacheprovider \ + tests/unit/model tests/unit/alignment tests/unit/frf_separate \ + > alignment_lab/slurm/_t_fast.log 2>&1 +rc_fast=$? +tail -4 alignment_lab/slurm/_t_fast.log +echo "rc_fast=$rc_fast" + +echo "###### slow-included: alignment + frf_separate" +"$PY" -m pytest -q --no-header -p no:cacheprovider --run-slow \ + tests/unit/alignment tests/unit/frf_separate \ + > alignment_lab/slurm/_t_slow.log 2>&1 +rc_slow=$? +tail -4 alignment_lab/slurm/_t_slow.log +echo "rc_slow=$rc_slow" + +echo "###### failures, if any" +grep -E "^(FAILED|ERROR)" alignment_lab/slurm/_t_fast.log alignment_lab/slurm/_t_slow.log || echo " none" +echo "done" diff --git a/alignment_lab/analysis/verify_copy_correctness.sh b/alignment_lab/analysis/verify_copy_correctness.sh new file mode 100644 index 00000000..7aec9231 --- /dev/null +++ b/alignment_lab/analysis/verify_copy_correctness.sh @@ -0,0 +1,80 @@ +#!/bin/bash +# Correctness half of the build_grid=False change. No timings here: this runs +# wherever there is a free slot, and runtime numbers are only comparable on the +# pinned EPYC 9335. What matters is that |F_calc| is bit-identical. +#SBATCH --job-name=frf_vcorr +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +FILT="Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$" +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" + +echo "=== |F_calc| from the box path, grid-building vs skipped ===" +"$PY" -u -c " +import math, torch +torch.set_grad_enabled(False) +from alignment_lab.lab.benchmark import load_case +from torchref.symmetry.cell import Cell +from torchref.experimental.alignment.frf.dense_calc import model_sf_abs +from torchref.experimental.alignment.frf.api import phaser_lmax_resolution + +def box_path(model, d_min, d_max, build_grid): + m = model.copy(build_grid=build_grid) + coords = m.xyz(); dev = coords.device + extent = (coords - coords.mean(0)).norm(dim=-1).max().item() + a = float(2.0 * 2.0 * extent) + m.max_res = float(d_min); m.spacegroup = 'P 1' + m.cell = Cell([a, a, a, 90., 90., 90.], device=dev) + nmax = int(math.ceil(a / d_min)) + idx = torch.arange(-nmax, nmax + 1, device=dev) + H, K, Lg = torch.meshgrid(idx, idx, idx, indexing='ij') + hkl = torch.stack([H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], -1).to(torch.long) + smag = hkl.to(torch.float64).norm(dim=-1) / a + hkl = hkl[(smag >= 1.0/d_max) & (smag <= 1.0/d_min)].contiguous() + return model_sf_abs(m, hkl) + +bad = 0 +for name in ('3K7M', '1DAW'): + model, data = load_case(name); model.verbose = 0 + rb = data.cell.reciprocal_basis_matrix.to(torch.float64) + d_min_data = 1.0 / (data.hkl.to(torch.float64) @ rb).norm(dim=-1).max().item() + r = float((model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item()) + for cap in (64, 100): + L, d_min = phaser_lmax_resolution(r, d_min_data, cap) + a = box_path(model, d_min, 100.0, False) + b = box_path(model, d_min, 100.0, True) + ok = torch.equal(a, b) + bad += 0 if ok else 1 + md = 0.0 if ok else (a-b).abs().max().item() + print(f' {name} cap{cap}: bit-identical={ok} n={a.numel()} max|dF|={md:.3e}') +print('MISMATCHES:', bad) +" 2>&1 | grep -vE "$FILT" + +echo "=== the caller must not be mutated ===" +"$PY" -u -c " +import torch +torch.set_grad_enabled(False) +from alignment_lab.lab.benchmark import load_case +from torchref.experimental.alignment.frf.dense_calc import dense_calc_via_box +model, data = load_case('1DAW'); model.verbose = 0 +before = (str(model.spacegroup), model.max_res, model.cell.data.clone(), + model.xyz().clone()) +dense_calc_via_box(model, 100.0, 4.0, pad=2.0) +after = (str(model.spacegroup), model.max_res, model.cell.data.clone(), + model.xyz().clone()) +print(' spacegroup preserved:', before[0] == after[0], before[0]) +print(' max_res preserved :', before[1] == after[1], before[1]) +print(' cell preserved :', torch.equal(before[2], after[2])) +print(' coords preserved :', torch.equal(before[3], after[3])) +print(' grid still present :', model._fft.real_space_grid is not None) +" 2>&1 | grep -vE "$FILT" + +echo "=== tests ===" +"$PY" -u -m pytest -q tests/unit/model tests/unit/alignment tests/unit/frf_separate 2>&1 | tail -12 +echo "exit_code=$?" diff --git a/alignment_lab/analysis/verify_copy_fix.sh b/alignment_lab/analysis/verify_copy_fix.sh new file mode 100644 index 00000000..d88b0572 --- /dev/null +++ b/alignment_lab/analysis/verify_copy_fix.sh @@ -0,0 +1,79 @@ +#!/bin/bash +# Verify the build_grid=False change in dense_calc_via_box: +# 1. |F_calc| must be bit-identical to the grid-building path. +# 2. Truth ranks must not move. +# 3. Report the new stage table, per cap (cap64 is what ships). +#SBATCH --job-name=frf_vcopy +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +FILT="Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$" +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" + +echo "=== 1. |F_calc| identical, and the copy cost ===" +"$PY" -u -c " +import math, time, torch +torch.set_grad_enabled(False) +from alignment_lab.lab.benchmark import load_case +from torchref.symmetry.cell import Cell +from torchref.experimental.alignment.frf.dense_calc import dense_calc_via_box, model_sf_abs +from torchref.experimental.alignment.frf.api import phaser_lmax_resolution + +def box_path(model, d_min, d_max, build_grid): + '''dense_calc_via_box, with the grid build under our control.''' + m = model.copy(build_grid=build_grid) + coords = m.xyz(); dev = coords.device + extent = (coords - coords.mean(0)).norm(dim=-1).max().item() + a = float(2.0 * 2.0 * extent) + m.max_res = float(d_min); m.spacegroup = 'P 1' + m.cell = Cell([a, a, a, 90., 90., 90.], device=dev) + nmax = int(math.ceil(a / d_min)) + idx = torch.arange(-nmax, nmax + 1, device=dev) + H, K, Lg = torch.meshgrid(idx, idx, idx, indexing='ij') + hkl = torch.stack([H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], -1).to(torch.long) + smag = hkl.to(torch.float64).norm(dim=-1) / a + hkl = hkl[(smag >= 1.0/d_max) & (smag <= 1.0/d_min)].contiguous() + return model_sf_abs(m, hkl) + +def timed(fn, n=3): + fn() + return min((lambda: (lambda t0: (fn(), time.perf_counter()-t0)[1])(time.perf_counter()))() for _ in range(n)) + +for name in ('3K7M', '1DAW'): + model, data = load_case(name) + model.verbose = 0 + rb = data.cell.reciprocal_basis_matrix.to(torch.float64) + d_min_data = 1.0 / (data.hkl.to(torch.float64) @ rb).norm(dim=-1).max().item() + r = float((model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item()) + L, d_min = phaser_lmax_resolution(r, d_min_data, 64) + f_lean = box_path(model, d_min, 100.0, False) + f_full = box_path(model, d_min, 100.0, True) + same = torch.equal(f_lean, f_full) + md = (f_lean - f_full).abs().max().item() if not same else 0.0 + t_lean = timed(lambda: model.copy(build_grid=False)) + t_full = timed(lambda: model.copy(build_grid=True)) + t_stage = timed(lambda: dense_calc_via_box(model, 100.0, d_min, pad=2.0)) + print(f' {name}: bit-identical={same} (max|dF|={md:.3e}, n={f_lean.numel()})') + print(f' model.copy(build_grid=True) {t_full*1e3:8.1f} ms') + print(f' model.copy(build_grid=False) {t_lean*1e3:8.1f} ms') + print(f' dense_calc_via_box now {t_stage*1e3:8.1f} ms') +" 2>&1 | grep -vE "$FILT" + +echo "=== 2/3. ranks and the new stage table, per cap ===" +for pdb in 3K7M 1DAW; do + for arm in cap64 cap100; do + "$PY" -u -m alignment_lab.diagnostics.frf_benchmark \ + --pdb "$pdb" --arms "$arm" --trials 2 2>&1 | grep -vE "$FILT" + done +done + +echo "=== 4. tests ===" +"$PY" -u -m pytest -q tests/unit/model/test_copy.py tests/unit/alignment \ + tests/unit/frf_separate 2>&1 | tail -15 +echo "exit_code=$?" diff --git a/alignment_lab/analysis/wigner_cache_probe.sh b/alignment_lab/analysis/wigner_cache_probe.sh new file mode 100644 index 00000000..9f238a33 --- /dev/null +++ b/alignment_lab/analysis/wigner_cache_probe.sh @@ -0,0 +1,88 @@ +#!/bin/bash +# Two questions about caching the Wigner d-table: +# 1. What fraction of wigner_contraction_per_beta is the data-INdependent +# d-block build (so, cacheable) vs the contraction against xi (not)? +# 2. Is loading that table off GPFS actually faster than recomputing it? +#SBATCH --job-name=frf_wcache +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname)" +SCRATCH=alignment_lab/runs/wigner_cache_probe +mkdir -p "$SCRATCH" +"$PY" -u -c " +import math, os, time +import torch; torch.set_grad_enabled(False) +from torchref.experimental.alignment.frf.wigner_d import ( + wigner_contraction_per_beta, _wigner_eig_table) + +scratch = '$SCRATCH' +dev = torch.device('cpu') + +def build_table(L, betas): + '''The data-independent half: every d^l(beta) block, packed, float32.''' + eig = _wigner_eig_table(L, dev) + out = [] + for l in range(1, L): + w, V = eig[l - 1] + phase = torch.exp(-1j * betas.unsqueeze(1) * w.unsqueeze(0)) + VP = V.unsqueeze(0) * phase.unsqueeze(1) + out.append((VP @ V.conj().transpose(-1, -2)).real.to(torch.float32)) + return out + +def contract_only(table, xi, L, n_beta): + '''The data-dependent half, given a prebuilt table.''' + dim = 2 * L - 1; c = L - 1 + S = torch.zeros((n_beta, dim, dim), dtype=torch.complex128) + S[:, c, c] += xi[0, c, c] + for l in range(1, L): + lo, hi = c - l, c + l + 1 + S[:, lo:hi, lo:hi] += xi[l, lo:hi, lo:hi].unsqueeze(0) * table[l-1].to(torch.complex128) + return S + +def best(fn, n=3): + fn() + return min((lambda: (lambda t0: (fn(), time.perf_counter()-t0)[1])(time.perf_counter()))() for _ in range(n)) + +for cap in (64, 100): + L = cap + 1 + n_beta = int(math.ceil(180.0 / 3.0)) + betas = torch.arange(n_beta, dtype=torch.float64) * 3.0 * (math.pi/180) + xi = torch.randn(L, 2*L-1, 2*L-1, dtype=torch.complex128) * 1e-3 + + t_full = best(lambda: wigner_contraction_per_beta(xi, betas)) + t_build = best(lambda: build_table(L, betas)) + table = build_table(L, betas) + t_contr = best(lambda: contract_only(table, xi, L, n_beta)) + + # Correctness of the split, in float32 storage. + ref = wigner_contraction_per_beta(xi, betas) + got = contract_only(table, xi, L, n_beta) + rel = ((got - ref).abs().max() / ref.abs().max()).item() + + # Round-trip through GPFS, as one packed flat tensor. + flat = torch.cat([t.reshape(-1) for t in table]) + path = os.path.join(scratch, f'dtable_L{L}.pt') + t0 = time.perf_counter(); torch.save(flat, path); t_save = time.perf_counter()-t0 + mb = os.path.getsize(path)/1e6 + os.system('sync') + loads = [] + for _ in range(3): + t0 = time.perf_counter(); torch.load(path, map_location='cpu'); loads.append(time.perf_counter()-t0) + t_load_warm = min(loads) + print(f'--- cap{cap} L={L} n_beta={n_beta} ---') + print(f' full contraction now {t_full*1e3:8.1f} ms') + print(f' of which d-block build {t_build*1e3:8.1f} ms (cacheable)') + print(f' of which xi contraction {t_contr*1e3:8.1f} ms (not)') + print(f' float32 table rel.err {rel:8.2e}') + print(f' on disk {mb:8.1f} MB save {t_save*1e3:.0f} ms ' + f'load(page-cached) {t_load_warm*1e3:.0f} ms') + os.remove(path) +" +echo "exit_code=$?" diff --git a/alignment_lab/analysis/wigner_fft_proto.sh b/alignment_lab/analysis/wigner_fft_proto.sh new file mode 100644 index 00000000..1f4eaa57 --- /dev/null +++ b/alignment_lab/analysis/wigner_fft_proto.sh @@ -0,0 +1,73 @@ +#!/bin/bash +# The J_y eigenvalues are the integers -l..l, so S(beta) is a trigonometric +# polynomial in beta: accumulate its Fourier coefficients once (no beta axis) +# and get every beta from one FFT. Prototype + correctness + timing, against +# the current per-beta matmul. +#SBATCH --job-name=frf_wfft +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname)" +"$PY" -u -c " +import math, time, torch +torch.set_grad_enabled(False) +from torchref.experimental.alignment.frf.wigner_d import ( + wigner_contraction_per_beta, _wigner_eig_table) + +dev = torch.device('cpu') + +def contraction_fft(xi, betas): + L = xi.shape[0]; dim = 2 * L - 1; c = L - 1 + n_beta = betas.shape[0] + N = 2 * n_beta # betas must be j * 2*pi/N + C = torch.zeros((dim, dim, N), dtype=torch.complex128) + C[c, c, 0] += 2.0 * xi[0, c, c] # l=0: d^0 = 1, the halving is undone below + eig = _wigner_eig_table(L, dev) + for l in range(1, L): + w, V = eig[l - 1] + lo, hi = c - l, c + l + 1 + xi_l = xi[l, lo:hi, lo:hi] + k = (torch.round(w).to(torch.long)) % N + G = V.unsqueeze(1) * V.conj().unsqueeze(0) # (sz, sz, sz) over k + blk = C[lo:hi, lo:hi] + blk.index_add_(2, k, xi_l.unsqueeze(-1) * G) + blk.index_add_(2, (-k) % N, xi_l.unsqueeze(-1) * G.transpose(0, 1)) + full = torch.fft.fft(C, n=N, dim=-1) + return 0.5 * full[..., :n_beta].permute(2, 0, 1).contiguous() + +def best(fn, n=3): + fn() + out = [] + for _ in range(n): + t0 = time.perf_counter(); fn(); out.append(time.perf_counter() - t0) + return min(out) + +for cap in (64, 100): + L = cap + 1 + n_beta = int(math.ceil(180.0 / 3.0)) + betas = torch.arange(n_beta, dtype=torch.float64) * 3.0 * (math.pi / 180.0) + xi = torch.randn(L, 2*L-1, 2*L-1, dtype=torch.complex128) * 1e-3 + + # Are the eigenvalues really integers? The whole method rests on it. + eig = _wigner_eig_table(L, dev) + dev_max = max(float((w - torch.round(w)).abs().max()) for w, _ in eig) + + ref = wigner_contraction_per_beta(xi, betas) + got = contraction_fft(xi, betas) + rel = float((got - ref).abs().max() / ref.abs().max()) + + t_ref = best(lambda: wigner_contraction_per_beta(xi, betas)) + t_new = best(lambda: contraction_fft(xi, betas)) + print(f'--- cap{cap} L={L} n_beta={n_beta}') + print(f' max |w - round(w)| {dev_max:.2e}') + print(f' rel. difference {rel:.2e}') + print(f' per-beta matmul {t_ref*1e3:8.1f} ms') + print(f' Fourier + one FFT {t_new*1e3:8.1f} ms ({t_ref/t_new:.1f}x)') +" 2>&1 | grep -vE "Warning|warn| from |^ *$" +echo "exit_code=$?" diff --git a/alignment_lab/analysis/wigner_mirror_proto.sh b/alignment_lab/analysis/wigner_mirror_proto.sh new file mode 100644 index 00000000..1e33aa80 --- /dev/null +++ b/alignment_lab/analysis/wigner_mirror_proto.sh @@ -0,0 +1,79 @@ +#!/bin/bash +# d^l(pi - beta) is d^l(beta) up to an m-flip and a sign, and the beta grid +# (0, 3, ..., 177 deg) is closed under beta -> pi - beta. So the batched matmul +# only needs 31 of the 60 beta values. No data file, no kernel. +#SBATCH --job-name=frf_wmir +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" +"$PY" -u -c " +import math, time, torch +torch.set_grad_enabled(False) +from torchref.experimental.alignment.frf.wigner_d import ( + wigner_contraction_per_beta, _wigner_eig_table) +dev = torch.device('cpu') + +def d_block(w, V, betas): + phase = torch.exp(-1j * betas.unsqueeze(1) * w.unsqueeze(0)) + return ((V.unsqueeze(0) * phase.unsqueeze(1)) @ V.conj().transpose(-1, -2)).real + +# Which mirror identity holds? Test rather than trust. +L = 9 +eig = _wigner_eig_table(L, dev) +for l in (1, 3, 6, 8): + w, V = eig[l - 1] + b = torch.tensor([0.37, 1.11, 2.05], dtype=torch.float64) + lhs = d_block(w, V, math.pi - b) + base = d_block(w, V, b) + m = torch.arange(-l, l + 1, dtype=torch.float64) + s2 = ((-1.0) ** (l + m)).reshape(1, 1, -1) + s1 = ((-1.0) ** (l + m)).reshape(1, -1, 1) + f2 = s2 * base.flip(-1) # (-1)^(l+m2) d[m1, -m2] + f1 = s1 * base.flip(-2) # (-1)^(l+m1) d[-m1, m2] + print(f' l={l}: form(m2-flip) err={float((lhs-f2).abs().max()):.2e} ' + f'form(m1-flip) err={float((lhs-f1).abs().max()):.2e}') + +def contraction_mirror(xi, betas): + L = xi.shape[0]; dim = 2 * L - 1; c = L - 1 + n_beta = betas.shape[0] + half = n_beta // 2 + 1 # 0..30 for n_beta=60 + src = n_beta - torch.arange(half, n_beta) # beta_j -> pi - beta_j + S = torch.zeros((n_beta, dim, dim), dtype=torch.complex128) + S[:, c, c] += xi[0, c, c] + eig = _wigner_eig_table(L, dev) + for l in range(1, L): + w, V = eig[l - 1] + d_h = d_block(w, V, betas[:half]) + m = torch.arange(-l, l + 1, dtype=d_h.dtype) + sgn = ((-1.0) ** (l + m)).reshape(1, 1, -1) + d_l = torch.cat([d_h, sgn * d_h[src].flip(-1)], dim=0) + lo, hi = c - l, c + l + 1 + S[:, lo:hi, lo:hi] += xi[l, lo:hi, lo:hi].unsqueeze(0) * d_l.to(torch.complex128) + return S + +def best(fn, n=3): + fn() + return min((lambda: (lambda t0: (fn(), time.perf_counter()-t0)[1])(time.perf_counter()))() for _ in range(n)) + +for cap in (64, 100): + L = cap + 1 + n_beta = int(math.ceil(180.0 / 3.0)) + betas = torch.arange(n_beta, dtype=torch.float64) * 3.0 * (math.pi / 180.0) + xi = torch.randn(L, 2*L-1, 2*L-1, dtype=torch.complex128) * 1e-3 + ref = wigner_contraction_per_beta(xi, betas) + got = contraction_mirror(xi, betas) + rel = float((got - ref).abs().max() / ref.abs().max()) + t_ref = best(lambda: wigner_contraction_per_beta(xi, betas)) + t_new = best(lambda: contraction_mirror(xi, betas)) + print(f'--- cap{cap} L={L}: rel diff {rel:.2e} ' + f'current {t_ref*1e3:7.1f} ms mirrored {t_new*1e3:7.1f} ms ' + f'({t_ref/t_new:.2f}x)') +" 2>&1 | grep -vE "Warning|warn| from |^ *$|No CUDA" +echo "exit_code=$?" diff --git a/alignment_lab/analysis/wigner_split.sh b/alignment_lab/analysis/wigner_split.sh new file mode 100644 index 00000000..18a18d1c --- /dev/null +++ b/alignment_lab/analysis/wigner_split.sh @@ -0,0 +1,44 @@ +#!/bin/bash +# Inside wigner_contraction_per_beta, how much is data-INdependent (so +# cacheable) and how much is the contraction against xi (so not)? +#SBATCH --job-name=frf_wigner +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 PYTHONUNBUFFERED=1 +export CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname)" +"$PY" -u -c " +import time +import torch; torch.set_grad_enabled(False) +from torchref.experimental.alignment.frf.wigner_d import ( + wigner_contraction_per_beta, _wigner_eig_table) +from torchref.experimental.alignment.frf.sitelist_ang import _SAMPLE_LIST_CACHE + +for cap, sampling in ((64, 3.0), (100, 3.0)): + L = cap + 1 + n_beta = int(__import__('math').ceil(180.0 / sampling)) + betas = torch.arange(n_beta, dtype=torch.float64) * sampling * (3.141592653589793/180) + xi = torch.randn(L, 2*L-1, 2*L-1, dtype=torch.complex128) * 1e-3 + + t0 = time.perf_counter(); _wigner_eig_table(L, torch.device('cpu')) + cold_eig = time.perf_counter() - t0 + t0 = time.perf_counter(); _wigner_eig_table(L, torch.device('cpu')) + warm_eig = time.perf_counter() - t0 + + wigner_contraction_per_beta(xi, betas) # warm everything + t0 = time.perf_counter() + wigner_contraction_per_beta(xi, betas) + total = time.perf_counter() - t0 + + # Size of the d-table if it were precomputed and stored, float32 real. + entries = sum((2*l+1)**2 for l in range(L)) * n_beta + print(f'L={L:4d} n_beta={n_beta:3d} | eig cold {cold_eig*1e3:7.1f} ms ' + f'warm {warm_eig*1e3:5.2f} ms | contraction total {total*1e3:7.1f} ms ' + f'| d-table would be {entries*4/1e6:7.1f} MB float32') +" +echo "exit_code=$?" From cf2f748619bebb724dec608fcfd857571239f4e5 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Sun, 23 Aug 2026 00:25:28 +0200 Subject: [PATCH 038/250] Memoise the Wigner small-d blocks The blocks depend only on the bandwidth and the beta grid, never on the data, and the batched (n_beta, sz, sz) @ (sz, sz) product that builds them is 80% of wigner_contraction_per_beta -- 118 ms of 146 ms at L=65, 631 of 778 at L=101. Hoisting them into a module memo makes every search after the first in a process 3.9x cheaper on that stage, bit-identically. One search calls the contraction exactly once (counted, not assumed), so a single search still pays the full build; this buys repeated searches in one process, which is what the pipeline and the benchmark harness do. The cache holds one entry on purpose: the blocks are 176 MB float64 at L=65 and 659 MB at L=101 for the 60-value beta grid. LMAX_CAP and GRID_SAMPLING_DEG are constants, so production only ever wants one key. Two cheaper-looking routes were measured and rejected. The integer J_y eigenvalues make S(beta) a trigonometric polynomial, so its Fourier coefficients can be accumulated with no beta axis and every beta read off one FFT: correct to 5e-15 but 0.7-0.9x, because it is bandwidth-bound where the batched matmul is compute-bound. The beta -> pi - beta mirror reaches only 1.28x. Both prototypes are kept in alignment_lab/analysis. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/verify_wigner_memo.sh | 71 ++++++++++++ alignment_lab/analysis/wigner_call_count.sh | 47 ++++++++ docs/changelog.rst | 1 + tests/unit/alignment/test_wigner_d_cache.py | 105 ++++++++++++++++++ .../experimental/alignment/frf/wigner_d.py | 73 +++++++++--- 5 files changed, 284 insertions(+), 13 deletions(-) create mode 100644 alignment_lab/analysis/verify_wigner_memo.sh create mode 100644 alignment_lab/analysis/wigner_call_count.sh create mode 100644 tests/unit/alignment/test_wigner_d_cache.py diff --git a/alignment_lab/analysis/verify_wigner_memo.sh b/alignment_lab/analysis/verify_wigner_memo.sh new file mode 100644 index 00000000..7df9fcde --- /dev/null +++ b/alignment_lab/analysis/verify_wigner_memo.sh @@ -0,0 +1,71 @@ +#!/bin/bash +# The memoised small-d blocks: same answer, and how much the second search in a +# process saves. Also the peak-memory cost of holding the blocks. +#SBATCH --job-name=frf_wmemo +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +FILT="Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$|No CUDA" +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" + +echo "=== contraction: cold vs memoised, and identity ===" +"$PY" -u -c " +import math, resource, time, torch +torch.set_grad_enabled(False) +from torchref.experimental.alignment.frf.wigner_d import ( + wigner_contraction_per_beta, clear_wigner_d_cache) + +def rss_mb(): + return resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / 1024.0 + +for cap in (64, 100): + L = cap + 1 + n_beta = int(math.ceil(180.0 / 3.0)) + betas = torch.arange(n_beta, dtype=torch.float64) * 3.0 * (math.pi/180.0) + g = torch.Generator().manual_seed(0) + xi = torch.randn(L, 2*L-1, 2*L-1, generator=g, dtype=torch.float64).to(torch.complex128) + clear_wigner_d_cache() + r0 = rss_mb() + t0 = time.perf_counter(); a = wigner_contraction_per_beta(xi, betas); cold = time.perf_counter()-t0 + r1 = rss_mb() + warm = min([(lambda t: (wigner_contraction_per_beta(xi, betas), time.perf_counter()-t)[1])(time.perf_counter()) for _ in range(3)]) + b = wigner_contraction_per_beta(xi, betas) + print(f' cap{cap} L={L}: cold {cold*1e3:7.1f} ms memoised {warm*1e3:6.1f} ms ' + f'({cold/warm:.1f}x) identical={torch.equal(a, b)} ' + f'RSS +{r1-r0:.0f} MB') + clear_wigner_d_cache() +" 2>&1 | grep -vE "$FILT" + +echo "=== two searches in one process ===" +"$PY" -u -c " +import time, torch +torch.set_grad_enabled(False) +from alignment_lab.lab.benchmark import load_case +from torchref.experimental.alignment.rotation_search import rotation_search +model, data = load_case('1DAW'); model.verbose = 0 +rotation_search(model, data, 0.8, n_peaks=50) # prewarm the process +prev = None +for i in (1, 2, 3): + t0 = time.perf_counter(); sol = rotation_search(model, data, 0.8, n_peaks=50) + dt = time.perf_counter() - t0 + fp = float(sol.scores[:20].double().sum()) + tag = '' if prev is None else (' same top-20 sum' if fp == prev else ' DIFFERS') + prev = fp + print(f' search {i}: {dt:.3f} s{tag}') +" 2>&1 | grep -vE "$FILT" + +echo "=== tests ===" +"$PY" -m pytest -q --no-header -p no:cacheprovider \ + tests/unit/alignment tests/unit/frf_separate tests/unit/model \ + > alignment_lab/slurm/_t_wmemo.log 2>&1 +rc=$? +tail -2 alignment_lab/slurm/_t_wmemo.log +echo "rc=$rc" +grep -E "^(FAILED|ERROR)" alignment_lab/slurm/_t_wmemo.log || echo " no failures" +echo "done" diff --git a/alignment_lab/analysis/wigner_call_count.sh b/alignment_lab/analysis/wigner_call_count.sh new file mode 100644 index 00000000..0e5ad8c5 --- /dev/null +++ b/alignment_lab/analysis/wigner_call_count.sh @@ -0,0 +1,47 @@ +#!/bin/bash +# A RAM memo of the d-blocks only pays off if the table is wanted more than +# once. How many times does one rotation_search ask for it, and at what (L, +# n_beta)? An earlier note claimed the obs and calc sides each build it. +#SBATCH --job-name=frf_wcnt +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname)" +"$PY" -u -c " +import torch +torch.set_grad_enabled(False) +import torchref.experimental.alignment.frf.wigner_d as wd +from alignment_lab.lab.benchmark import load_case +from torchref.experimental.alignment.rotation_search import rotation_search + +calls = [] +orig = wd.wigner_contraction_per_beta +def counting(xi, betas): + calls.append((int(xi.shape[0]), int(betas.shape[0]))) + return orig(xi, betas) +wd.wigner_contraction_per_beta = counting +# the caller imported it by name, so rebind there too +import torchref.experimental.alignment.frf.sitelist_ang as sa +for mod in (sa,): + if getattr(mod, 'wigner_contraction_per_beta', None) is orig: + mod.wigner_contraction_per_beta = counting +import sys +for name, mod in list(sys.modules.items()): + if name.startswith('torchref.experimental.alignment') and \ + getattr(mod, 'wigner_contraction_per_beta', None) is orig: + mod.wigner_contraction_per_beta = counting + print(' rebound in', name) + +model, data = load_case('1DAW'); model.verbose = 0 +for run in (1, 2): + calls.clear() + rotation_search(model, data, 0.8, n_peaks=50) + print(f' search {run}: {len(calls)} call(s) to wigner_contraction_per_beta -> {calls}') +" 2>&1 | grep -vE "Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$|No CUDA" +echo "done" diff --git a/docs/changelog.rst b/docs/changelog.rst index 62c7306a..a22a44b7 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -14,6 +14,7 @@ Unreleased - Replaced the fast rotation function's keyword surface with ``rotation_search(model, data, model_error_A)``; the caller's coordinate error is now used rather than overwritten by an estimate from the atom count - Removed the rotation function's dead modules, engine variants, debug environment switches and unreachable knobs - Fixed the rotation function's dense model transform building a real-space grid and map-symmetry operator that its next three lines discarded; ``ModelFT.copy`` gained ``build_grid`` +- The rotation function's Wigner small-d blocks are memoised, so a process running more than one search builds them once Version 0.6.4 ---------- diff --git a/tests/unit/alignment/test_wigner_d_cache.py b/tests/unit/alignment/test_wigner_d_cache.py new file mode 100644 index 00000000..57028eb4 --- /dev/null +++ b/tests/unit/alignment/test_wigner_d_cache.py @@ -0,0 +1,105 @@ +"""The memoised small-d blocks must not change what the contraction returns. + +The blocks ``d^l(β)`` depend only on the bandwidth and the β grid, never on the +data, so hoisting them out of the per-call loop is pure reuse. That makes the +cache a correctness risk in exactly one way: a stale entry served for the wrong +``(L, betas)``. These tests pin the key, the identity of the result, and the +one-entry footprint bound. +""" + +import math + +import pytest +import torch + +from torchref.experimental.alignment.frf.wigner_d import ( + _WIGNER_D_CACHE, + _wigner_d_blocks, + clear_wigner_d_cache, + wigner_contraction_per_beta, +) + + +def _betas(n, step_deg=3.0): + return torch.arange(n, dtype=torch.float64) * step_deg * (math.pi / 180.0) + + +def _xi(L, seed=0): + g = torch.Generator().manual_seed(seed) + dim = 2 * L - 1 + return torch.randn(L, dim, dim, generator=g, dtype=torch.float64).to( + torch.complex128 + ) + + +@pytest.mark.unit +def test_the_cached_call_is_bit_identical(): + """A cache hit must give exactly the first call's answer, not merely close.""" + clear_wigner_d_cache() + L, betas = 9, _betas(12) + xi = _xi(L) + first = wigner_contraction_per_beta(xi, betas) + second = wigner_contraction_per_beta(xi, betas) + assert torch.equal(first, second) + + +@pytest.mark.unit +def test_different_data_at_the_same_bandwidth_still_differs(): + """Guard the premise: the cache holds β geometry, not the data. + + Without this, a cache keyed too loosely -- or one that memoised the whole + result -- would pass the identity test above by returning a stale answer. + """ + clear_wigner_d_cache() + L, betas = 9, _betas(12) + a = wigner_contraction_per_beta(_xi(L, seed=0), betas) + b = wigner_contraction_per_beta(_xi(L, seed=1), betas) + assert not torch.allclose(a, b) + + +@pytest.mark.unit +@pytest.mark.parametrize( + "L2, n_beta2", [(9, 15), (11, 12)], ids=["other-beta-grid", "other-bandwidth"] +) +def test_a_different_key_is_not_served_the_cached_blocks(L2, n_beta2): + """Changing either the bandwidth or the β grid must rebuild.""" + clear_wigner_d_cache() + ref_blocks = _wigner_d_blocks(9, _betas(12), torch.device("cpu")) + got = _wigner_d_blocks(L2, _betas(n_beta2), torch.device("cpu")) + assert len(got) == L2 - 1 + assert got[0].shape[0] == n_beta2 + assert got is not ref_blocks + + +@pytest.mark.unit +def test_the_cache_holds_one_entry(): + """The blocks are hundreds of MB at production bandwidths, so they are not + allowed to accumulate across keys.""" + clear_wigner_d_cache() + for L in (7, 9, 11): + _wigner_d_blocks(L, _betas(12), torch.device("cpu")) + assert len(_WIGNER_D_CACHE) == 1 + clear_wigner_d_cache() + assert len(_WIGNER_D_CACHE) == 0 + + +@pytest.mark.unit +def test_the_blocks_are_the_wigner_small_d_matrices(): + """Anchor the cached quantity against d^l(β) computed from the definition. + + ``d^l(0) = I`` and ``d^l(β)`` is orthogonal for every β; both follow from + ``d^l(β) = exp(-i β J_y)`` and neither holds for a mis-shaped or + mis-transposed block. + """ + clear_wigner_d_cache() + L = 7 + betas = torch.tensor([0.0, 0.4, 1.7, 3.0], dtype=torch.float64) + blocks = _wigner_d_blocks(L, betas, torch.device("cpu")) + for l, d in enumerate(blocks, start=1): + sz = 2 * l + 1 + assert d.shape == (betas.numel(), sz, sz) + torch.testing.assert_close(d[0], torch.eye(sz, dtype=d.dtype)) + for k in range(betas.numel()): + torch.testing.assert_close( + d[k] @ d[k].transpose(-1, -2), torch.eye(sz, dtype=d.dtype) + ) diff --git a/torchref/experimental/alignment/frf/wigner_d.py b/torchref/experimental/alignment/frf/wigner_d.py index b90fe96f..29a7791f 100644 --- a/torchref/experimental/alignment/frf/wigner_d.py +++ b/torchref/experimental/alignment/frf/wigner_d.py @@ -6,19 +6,30 @@ Phaser uses the Sakurai recurrence convention; the equivalent Edmonds (4.1.23) convention is used throughout this package and is pinned against Phaser's output by ``tests/unit/alignment/test_wigner.py``. -``wigner_contraction_per_beta`` builds the small-d table it needs from the -``J_y`` eigendecomposition inline, which stays bounded to any ``l``. +``wigner_contraction_per_beta`` builds the small-d blocks it needs from the +``J_y`` eigendecomposition, which stays bounded to any ``l``. The blocks depend +only on the bandwidth and the β grid, so they are memoised for reuse across +searches -- see ``_WIGNER_D_CACHE`` for what that costs in memory. """ from __future__ import annotations import torch -__all__ = ["wigner_contraction_per_beta"] +__all__ = ["clear_wigner_d_cache", "wigner_contraction_per_beta"] #: Memo for the per-l J_y eigendecomposition, keyed on (L, device-str). #: It depends only on the bandwidth, so repeat calls at the same L reuse it. _WIGNER_EIG_CACHE: dict = {} +#: Memo for the per-l small-d blocks, keyed on (L, betas, device-str). Holds at +#: most one entry, because the blocks are large: the per-l blocks together are +#: ``n_beta * sum_l (2l+1)^2`` float64 scalars, 176 MB at L=65 and 659 MB at +#: L=101 for the 60-value beta grid. One entry is all production wants -- +#: ``LMAX_CAP`` and ``GRID_SAMPLING_DEG`` in ``rotation_search`` are constants, +#: so every call arrives with the same key. A caller that alternates bandwidths +#: rebuilds each time, which is the uncached cost and not worse. +_WIGNER_D_CACHE: dict = {} + def _wigner_eig_table(L: int, device: torch.device): """Return [(w_l, V_l)] for l ∈ [1, L) — the J_y eigendecomposition per l. @@ -42,6 +53,46 @@ def _wigner_eig_table(L: int, device: torch.device): return table +def clear_wigner_d_cache() -> None: + """Drop the memoised small-d blocks, releasing their memory.""" + _WIGNER_D_CACHE.clear() + + +def _wigner_d_blocks(L: int, betas: torch.Tensor, device: torch.device): + """Per-l real ``d^l(β)`` blocks for ``l ∈ [1, L)``, memoised. + + Each entry is ``(n_beta, 2l+1, 2l+1)`` float64. They depend only on the + bandwidth and the β grid, not on the data, so a process that runs more than + one rotation search at the same bandwidth builds them once. A single search + asks for them exactly once and so pays the full build. + + The build is the dominant cost of :func:`wigner_contraction_per_beta`: it is + a batched ``(n_beta, sz, sz) @ (sz, sz)`` product per l, against the + contraction's elementwise ``(n_beta, sz, sz)``. See ``_WIGNER_D_CACHE`` for + the footprint that buys. + """ + key = ( + int(L), + str(device), + tuple(betas.detach().to(torch.float64).cpu().tolist()), + ) + hit = _WIGNER_D_CACHE.get(key) + if hit is not None: + return hit + + eig_table = _wigner_eig_table(L, device) # cached (w_l, V_l) + blocks = [] + for l in range(1, L): + w, V = eig_table[l - 1] # data-independent + phase = torch.exp(-1j * betas.unsqueeze(1) * w.unsqueeze(0)) # (n_beta, sz) + VP = V.unsqueeze(0) * phase.unsqueeze(1) # (n_beta, sz, sz) = (k,m,a) + blocks.append((VP @ V.conj().transpose(-1, -2)).real) # (n_beta, sz, sz) + + _WIGNER_D_CACHE.clear() # one entry only; see the footprint note + _WIGNER_D_CACHE[key] = blocks + return blocks + + def wigner_contraction_per_beta( xi_lmn: torch.Tensor, betas: torch.Tensor, @@ -75,20 +126,16 @@ def wigner_contraction_per_beta( n_beta = betas.shape[0] xi = xi_lmn.to(torch.complex128) - # Fused per-l loop: compute each d^l(β) block via J_y eigendecomposition - # (small_d_stable's method, stable to any l) and contract it into S - # immediately. Never materialises the full (n_beta, L, 2L-1, 2L-1) table - # (~19 GB at L=100) nor a 4-D einsum intermediate — peak memory is one - # (n_beta, 2l+1, 2l+1) block (~0.4 GB at l=99). + # Per-l loop over the small-d blocks, which come from the J_y + # eigendecomposition (small_d_stable's method, stable to any l). Contract + # each into S in turn: the full (n_beta, L, 2L-1, 2L-1) table is never + # materialised as one array, nor is a 4-D einsum intermediate. S = torch.zeros((n_beta, dim, dim), dtype=torch.complex128, device=device) c = L - 1 S[:, c, c] += xi[0, c, c] # l=0: d^0 = 1 - eig_table = _wigner_eig_table(L, device) # cached (w_l, V_l) + blocks = _wigner_d_blocks(L, betas, device) for l in range(1, L): - w, V = eig_table[l - 1] # data-independent - phase = torch.exp(-1j * betas.unsqueeze(1) * w.unsqueeze(0)) # (n_beta, sz) - VP = V.unsqueeze(0) * phase.unsqueeze(1) # (n_beta, sz, sz) = (k,m,a) - d_l = (VP @ V.conj().transpose(-1, -2)).real # (n_beta, sz, sz) + d_l = blocks[l - 1] # (n_beta, sz, sz) lo, hi = c - l, c + l + 1 S[:, lo:hi, lo:hi] += xi[l, lo:hi, lo:hi].unsqueeze(0) * d_l.to(torch.complex128) return S From d1244c4565bc3be7365381cdeed8368b5804b81d Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Sun, 23 Aug 2026 06:36:26 +0200 Subject: [PATCH 039/250] Add the French-Wilson cost and ASU-equivalence measurements fw_internals.sh: at the unrolled reflection count the binning is 0.8% of french_wilson_preprocess (69.8 ms of 8685.8 on 3K7M). The cost is the parabolic-cylinder posterior (3288.9 ms) and the Halley D-factor solve (5256.6 ms), and what multiplies them is the n_ops repetition -- the same values, 24 times, verified identical across the symmetry blocks. fw_asu_equivalence.sh: computing on the unique set instead is 27.4x, but is NOT bit-exact. The equal-count edge index moves by one rank between N and n_ops*N, shifting an edge ~2e-7 in |s| and moving 7 of 55078 reflections between shells; those shells' then shift every member by ~1e-4. Characterise the distribution -- a bare "how many differ" count reported 54504 of 55078 when the median difference was 1e-13. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/fw_asu_equivalence.sh | 58 ++++++++++++ alignment_lab/analysis/fw_internals.sh | 93 ++++++++++++++++++++ 2 files changed, 151 insertions(+) create mode 100644 alignment_lab/analysis/fw_asu_equivalence.sh create mode 100644 alignment_lab/analysis/fw_internals.sh diff --git a/alignment_lab/analysis/fw_asu_equivalence.sh b/alignment_lab/analysis/fw_asu_equivalence.sh new file mode 100644 index 00000000..a0e6d683 --- /dev/null +++ b/alignment_lab/analysis/fw_asu_equivalence.sh @@ -0,0 +1,58 @@ +#!/bin/bash +# Is ASU-then-broadcast equivalent to the unrolled computation? The previous run +# reported max|d eEobs| = 0.22 with "54504 differing", but counted ANY nonzero +# float difference, so that count says nothing about size. Characterise the +# distribution, and find where the outlier sits. +#SBATCH --job-name=frf_fwasu +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +FILT="Loaded|LINK|Wilson outlier|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$|No CUDA" +echo "host=$(hostname)" +"$PY" -u -c " +import numpy as np, torch +torch.set_grad_enabled(False) +from alignment_lab.lab.benchmark import load_case +from torchref.experimental.alignment.frf.french_wilson import french_wilson_preprocess + +for name in ('3K7M', '1DAW'): + model, data = load_case(name); model.verbose = 0 + rb = data.cell.reciprocal_basis_matrix.to(torch.float64) + s_all = (data.hkl.to(torch.float64) @ rb).norm(dim=-1) + d_min = 1.0 / s_all.max().item() + keep = (s_all >= 1.0/100.0) & (s_all <= 1.0/d_min) + n_ops = int(data.spacegroup.matrices.shape[0]) + F = data.F.abs().to(torch.float64)[keep]; sig = data.F_sigma.to(torch.float64)[keep] + s = s_all[keep]; cen = torch.zeros_like(F, dtype=torch.bool) + M = F.numel() + + unrolled = french_wilson_preprocess(F.repeat(n_ops), sig.repeat(n_ops), + s.repeat(n_ops), cen.repeat(n_ops), n_wilson_shells=20) + unique = french_wilson_preprocess(F, sig, s, cen, n_wilson_shells=20) + a = unrolled['eEobs'].reshape(n_ops, M)[0].numpy() + b = unique['eEobs'].numpy() + d = np.abs(a - b) + rel = d / np.maximum(np.abs(a), 1e-12) + print(f'--- {name}: {M} unique x {n_ops} ops') + for q in (50, 90, 99, 99.9, 100): + print(f' |d eEobs| p{q:<5} {np.percentile(d, q):.3e} rel {np.percentile(rel, q):.3e}') + print(f' above 1e-9 : {int((d > 1e-9).sum())} of {M}') + print(f' above 1e-3 : {int((d > 1e-3).sum())} of {M}') + # Do the shell edges actually differ, and by how many reflections? + def edges_and_shells(sv): + sn = sv.numpy(); idx = np.argsort(sn) + e = np.linspace(0, len(sn)-1, 21).round().astype(np.int64) + ed = sn[idx][e].copy(); ed[0] -= 1e-6; ed[-1] += 1e-6 + return ed, np.clip(np.searchsorted(ed, sn, side='right')-1, 0, 19) + eu, shu = edges_and_shells(s) + er, shr = edges_and_shells(s.repeat(n_ops)) + print(f' max |d shell edge| {np.abs(eu - er).max():.3e}') + print(f' reflections changing shell {int((shu != shr[:M]).sum())} of {M}') +" 2>&1 | grep -vE "$FILT" +echo "done" diff --git a/alignment_lab/analysis/fw_internals.sh b/alignment_lab/analysis/fw_internals.sh new file mode 100644 index 00000000..3238d075 --- /dev/null +++ b/alignment_lab/analysis/fw_internals.sh @@ -0,0 +1,93 @@ +#!/bin/bash +# What inside french_wilson_preprocess costs, at the unrolled reflection count? +# The binning, the parabolic-cylinder posterior, or the Halley D-factor solve? +# Also: are the n_ops symmetry copies really carrying identical values? +#SBATCH --job-name=frf_fw +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +FILT="Loaded|LINK|Wilson outlier|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$|No CUDA" +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" +"$PY" -u -c " +import time +import numpy as np, torch +torch.set_grad_enabled(False) +from alignment_lab.lab.benchmark import load_case +from torchref.experimental.alignment.frf.french_wilson import ( + french_wilson_preprocess, _french_wilson_posterior, _get_dfactor_vectorised) + +def best(fn, n=3): + fn() + out = [] + for _ in range(n): + t0 = time.perf_counter(); fn(); out.append(time.perf_counter()-t0) + return min(out) + +for name in ('3K7M', '1DAW'): + model, data = load_case(name); model.verbose = 0 + rb = data.cell.reciprocal_basis_matrix.to(torch.float64) + s_all = (data.hkl.to(torch.float64) @ rb).norm(dim=-1) + d_min = 1.0 / s_all.max().item() + keep = (s_all >= 1.0/100.0) & (s_all <= 1.0/d_min) + n_ops = int(data.spacegroup.matrices.shape[0]) + F = data.F.abs().to(torch.float64)[keep] + sig = data.F_sigma.to(torch.float64)[keep] + s = s_all[keep] + from torchref.experimental.alignment.frf.preprocessing import compute_epsilon + cen = torch.zeros_like(F, dtype=torch.bool) + try: + cen = data.centric.to(torch.bool)[keep] + except Exception: + pass + nu = F.numel() + + # the unrolled arrays, exactly as rotation_search builds them + Fu = F.repeat(n_ops); sigu = sig.repeat(n_ops) + su = s.repeat(n_ops); cenu = cen.repeat(n_ops) + + t_unrolled = best(lambda: french_wilson_preprocess(Fu, sigu, su, cenu, n_wilson_shells=20)) + t_unique = best(lambda: french_wilson_preprocess(F, sig, s, cen, n_wilson_shells=20)) + print(f'--- {name}: {nu} unique x {n_ops} ops = {nu*n_ops} unrolled') + print(f' whole function, unrolled {t_unrolled*1e3:8.1f} ms') + print(f' whole function, unique {t_unique*1e3:8.1f} ms ({t_unrolled/t_unique:.1f}x)') + + # Are the n_ops copies identical? If so the unrolled call is pure repetition. + fu = french_wilson_preprocess(Fu, sigu, su, cenu, n_wilson_shells=20) + fq = french_wilson_preprocess(F, sig, s, cen, n_wilson_shells=20) + blocks = fu['eEobs'].reshape(n_ops, nu) + same_across_ops = bool(torch.equal(blocks, blocks[0:1].expand(n_ops, -1))) + matches_unique = bool(torch.equal(blocks[0], fq['eEobs'])) + print(f' n_ops copies identical {same_across_ops}') + print(f' block 0 == unique-set run {matches_unique}') + if not matches_unique: + d = (blocks[0] - fq['eEobs']).abs() + print(f' max|d eEobs| {d.max().item():.3e} n_differing={int((d>0).sum())}') + + # Stage split, on the unrolled arrays. + F_np = Fu.numpy(); sig_np = sigu.numpy(); s_np = su.numpy(); cen_np = cenu.numpy() + def binning(): + idx = np.argsort(s_np) + e = np.linspace(0, len(s_np)-1, 21).round().astype(np.int64) + edges = s_np[idx][e]; edges[0] -= 1e-6; edges[-1] += 1e-6 + sh = np.clip(np.searchsorted(edges, s_np, side='right')-1, 0, 19) + m = np.zeros(20); c = np.zeros(20, dtype=np.int64) + np.add.at(m, sh, F_np*F_np); np.add.at(c, sh, 1) + return m/np.maximum(c,1), sh + t_bin = best(binning) + mF2, sh = binning() + eosq = F_np*F_np/mF2[sh]; sigesq = np.maximum(2.0*F_np*sig_np/mF2[sh], 0.0) + t_post = best(lambda: _french_wilson_posterior(eosq, sigesq, cen_np)) + ee, eesq = _french_wilson_posterior(eosq, sigesq, cen_np) + bad = eesq < ee*ee; eesq[bad] = ee[bad]**2 + 1e-12 + t_dfac = best(lambda: _get_dfactor_vectorised(ee, eesq, cen_np)) + print(f' of which: shells + {t_bin*1e3:8.1f} ms') + print(f' FW posterior {t_post*1e3:8.1f} ms') + print(f' DFAC Halley {t_dfac*1e3:8.1f} ms') +" 2>&1 | grep -vE "$FILT" +echo "done" From 411ce26ba84ebf89e531012bd05d5ea1fce08a08 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Sun, 23 Aug 2026 23:28:36 +0200 Subject: [PATCH 040/250] Reuse the shared primitives in the FRF and put its dtype/device under config Three planned stages of cleanup, verified together rather than separately: the combined tree produces peak lists BIT-IDENTICAL to d1244c45 on 3K7M and 1DAW at cap64 and cap100, single-threaded, 500 peaks each. Committed as one change because the intermediate trees were gated but never committed, so splitting the hunks now would produce states that were never actually run. Consume the shared primitives, delete the duplicates: - peak_finder's private Edmonds ZYZ builder -> base rotation_matrix_euler_zyz, which is the same fused form term for term, so the sequential NMS's threshold ties round identically (the three-matrix-product form in rotation_utils does not, which is why the local copy existed). - the symop unroll -> SpaceGroup.apply_to_hkl, whose docstring already carried the h.S convention the private copy had to be fixed for. The op-major flattening is preserved deliberately: a different row order changes the summation order in the downstream index_add_/unique. - align._rodrigues deleted, callers on rotation_utils.axis_angle_to_matrix. NOT on base.axis_angle_to_rotation_matrix, which takes only (3,)/(N, 3) and branches at theta < 1e-10 where rotation_utils' takes (..., 3). - canonical_device on the three str(device) memo keys. _WIGNER_D_CACHE holds exactly one entry and clears on miss, so a bare-vs-indexed spelling would thrash the memo rather than merely add an entry. - the engine's second phaser_lmax_resolution call, plus auto_lmax / model_radius_A / lmax_cap: the caller already computed the pair to size the dense box. - the calc-side resolution mask, measured to drop 0 of 339040 (3K7M) and 0 of 271630 (1DAW). The obs-side mask was expected to be redundant too and is NOT: its low-resolution half removes two 3K7M reflections beyond 100 A. - bessel_sh_expand's unread chunk_size, french_wilson_preprocess's unread sqrt_mean_F2. Bessel ladder rescales instead of relying on float64's range: Miller's downward recurrence seeds at an arbitrary magnitude, so the ladder is inflated before renormalisation -- 1.7e157 at x = bessel_h_scale/d_max ~ 1.3, which overflows float32 for every x below ~35. Truncating the band does not help: cut at the float32 underflow boundary the peak is still 5.2e56, because the growth lives between u_cut and x and that is the stretch you must walk to reach the u <~ x entries that carry the signal. Rescaling by 2**-100 keeps it in range at any precision and costs nothing: a power of two only decrements the exponent, so mantissas are untouched and the result is bit-identical. 0-5 rescale events per shell over the FRF's range. Working precision from config, device resolved from both inputs: dtypes.float / dtypes.complex replace the inherited-or-overridden dtype in bessel_sh_expand, and compute_dtype is gone from it, from the engine and from the call site. That closes a trap: the fused CPU Legendre kernel gates on dtypes=(torch.float32,), so it was reachable only because one distant call site passed complex64. The clustering keys keep the input's wider dtype -- _GROUP_SCALE_S keys |s| at 1e-7 and a float32 rounding is ~0.3 of a key step. The J_y eigendecomposition and the anisotropy fit move to the host in float64: both are small (L-1 matrices; 7 parameters over ~1e4 reflections) and precision there is worth more than locality. The back half stays WIDER than the expansion, deliberately. Narrowing xi, the Wigner contraction and the FFT to match looked like free memory and is not: both sums are oscillatory, so they cancel and the error is far worse than eps*sqrt(n) predicts. Measured -- top peak and z-score unchanged to 7 figures, but scores moved 9.4e-5 to 1.4e-3 relative and only 1 of 500 candidate slots held the same orientation, because the greedy NMS is sequential and a reordering cascades. cross_correlate_xi now accumulates one step wider. Tests: 526 passed / 11 skipped. Three groups needed fixing, one of which was a real finding: test_radial_truncation asserts the populated radial band equals Phaser's nmax(l), and it only ever passed because it called bessel_sh_expand WITHOUT compute_dtype and so measured float64. Production was already float32 there, where 30-55% of the band at low-|s| shells is below float32's smallest normal -- so the shipped band has always been narrower. The structural tests now run under double_cpu and a new test pins the working-precision truncation. Timing is unchanged. Paired, interleaved, same node, 3 rounds: median paired delta -0.07 s of 0.98 s on 3K7M (2 of 3 rounds faster) and +0.01 s on 1DAW (mixed signs). Peak memory flat. The value here is the cleanup, not speed. Pre-existing and untouched: test_patterson_translation:: test_amplitude_tf_zero_translation fails identically at d1244c45. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/frf_fingerprint.py | 40 +++++ alignment_lab/analysis/stagea_fingerprint.sh | 53 ++++++ alignment_lab/analysis/stagea_tests.sh | 32 ++++ alignment_lab/analysis/stagea_verify.sh | 146 +++++++++++++++++ alignment_lab/analysis/stageb_gate.sh | 49 ++++++ alignment_lab/analysis/stagec_gate3.sh | 55 +++++++ .../analysis/stagec_timing_paired.sh | 34 ++++ docs/changelog.rst | 6 + .../unit/alignment/test_radial_truncation.py | 39 ++++- .../test_rotation_search_dtype_device.py | 151 ++++++++++++++++++ tests/unit/alignment/test_wigner_d_cache.py | 8 +- .../unit/frf_separate/test_bessel_rescale.py | 128 +++++++++++++++ .../frf_separate/test_bessel_sh_grouping.py | 10 +- torchref/experimental/alignment/align.py | 41 +---- torchref/experimental/alignment/frf/api.py | 50 +++--- .../experimental/alignment/frf/data_mr.py | 145 +++++++++++------ .../alignment/frf/french_wilson.py | 7 +- .../experimental/alignment/frf/peak_finder.py | 37 +---- .../alignment/frf/rotation_utils.py | 9 +- .../alignment/frf/sitelist_ang.py | 3 +- .../experimental/alignment/frf/wigner_d.py | 63 +++++--- torchref/experimental/alignment/pipeline.py | 4 +- .../experimental/alignment/rotation_search.py | 86 ++++++---- 23 files changed, 983 insertions(+), 213 deletions(-) create mode 100644 alignment_lab/analysis/frf_fingerprint.py create mode 100644 alignment_lab/analysis/stagea_fingerprint.sh create mode 100644 alignment_lab/analysis/stagea_tests.sh create mode 100644 alignment_lab/analysis/stagea_verify.sh create mode 100644 alignment_lab/analysis/stageb_gate.sh create mode 100644 alignment_lab/analysis/stagec_gate3.sh create mode 100644 alignment_lab/analysis/stagec_timing_paired.sh create mode 100644 tests/unit/alignment/test_rotation_search_dtype_device.py create mode 100644 tests/unit/frf_separate/test_bessel_rescale.py diff --git a/alignment_lab/analysis/frf_fingerprint.py b/alignment_lab/analysis/frf_fingerprint.py new file mode 100644 index 00000000..eac11a5b --- /dev/null +++ b/alignment_lab/analysis/frf_fingerprint.py @@ -0,0 +1,40 @@ +"""Dump a ranked FRF peak-list fingerprint for cross-tree comparison. + +Run with PYTHONPATH pointed at whichever worktree should provide `torchref` and +`alignment_lab`; the file itself is tree-agnostic. Scores are printed at 9 +significant figures rather than raw float64 -- the engine's own run-to-run +spread is ~5e-8 relative, so comparing full precision would report noise as a +difference (see `frf_peaklist_reproducibility`). +""" +import argparse +import sys + +import torch + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--pdb", required=True) + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--n-peaks", type=int, default=500) + args = ap.parse_args() + + from alignment_lab.lab.benchmark import load_case + from alignment_lab.lab.frf import FRFConfig, run_frf + + model, data = load_case(args.pdb)[:2] + res = run_frf(model, data, FRFConfig(n_peaks=args.n_peaks, + lmax_cap=args.lmax_cap)) + peaks = res.peaks if hasattr(res, "peaks") else res[0] + # Every fingerprint line is prefixed. Loading a structure writes progress to + # stdout, so a comparison that filters on anything looser (blank lines, a + # leading '#') silently ingests that chatter as data. + print(f"#FP pdb={args.pdb} lmax_cap={args.lmax_cap} n_peaks={len(peaks)}") + for i, p in enumerate(peaks): + print(f"FP {i:4d} {p.alpha:.9g} {p.beta:.9g} {p.gamma:.9g} " + f"{p.score:.9g} {p.sigma:.9g}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/alignment_lab/analysis/stagea_fingerprint.sh b/alignment_lab/analysis/stagea_fingerprint.sh new file mode 100644 index 00000000..036002e7 --- /dev/null +++ b/alignment_lab/analysis/stagea_fingerprint.sh @@ -0,0 +1,53 @@ +#!/bin/bash +# Stage A gate: peak-list fingerprints must be bit-identical between the +# baseline worktree (d1244c45) and this one. SINGLE-THREADED -- at the default +# thread count the float32 SF reduction reorders ~12 of 500 peaks on 3K7M. +#SBATCH --job-name=stagea_fp +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=32G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +OLD=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/_stagea_baseline +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +SCRIPT=$NEW/alignment_lab/analysis/frf_fingerprint.py +OUT=$NEW/alignment_lab/slurm +export TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" + +status=0 +for pdb in 3K7M 1DAW; do + for cap in 64 100; do + for tree in OLD NEW; do + eval "root=\$$tree" + cd "$root" + PYTHONPATH="$root" "$PY" -u "$SCRIPT" --pdb "$pdb" --lmax-cap "$cap" \ + > "$OUT/fp_${tree}_${pdb}_${cap}.txt" 2> "$OUT/fp_${tree}_${pdb}_${cap}.err" + rc=$? + if [ $rc -ne 0 ]; then + echo "RUN FAILED tree=$tree pdb=$pdb cap=$cap rc=$rc" + tail -6 "$OUT/fp_${tree}_${pdb}_${cap}.err" + status=1 + fi + done + a="$OUT/fp_OLD_${pdb}_${cap}.txt"; b="$OUT/fp_NEW_${pdb}_${cap}.txt" + if [ -s "$a" ] && [ -s "$b" ]; then + if diff -q <(grep -v '^#' "$a") <(grep -v '^#' "$b") >/dev/null; then + echo "IDENTICAL $pdb cap$cap ($(grep -vc '^#' "$a") peaks)" + else + nd=$(diff <(grep -v '^#' "$a") <(grep -v '^#' "$b") | grep -c '^[<>]') + echo "DIFFERS $pdb cap$cap ($nd differing lines)" + diff <(grep -v '^#' "$a") <(grep -v '^#' "$b") | head -8 + status=1 + fi + else + echo "MISSING $pdb cap$cap"; status=1 + fi + done +done +echo "fingerprint_status=$status" diff --git a/alignment_lab/analysis/stagea_tests.sh b/alignment_lab/analysis/stagea_tests.sh new file mode 100644 index 00000000..baa0cd53 --- /dev/null +++ b/alignment_lab/analysis/stagea_tests.sh @@ -0,0 +1,32 @@ +#!/bin/bash +# Stage A: the alignment + frf_separate suites, fast then --run-slow. +#SBATCH --job-name=stagea_tests +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=32G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname)" + +LOG=alignment_lab/slurm/stagea_fast_$SLURM_JOB_ID.log +"$PY" -m pytest tests/unit/alignment tests/unit/frf_separate tests/unit/model \ + tests/unit/test_imports_smoke.py -q > "$LOG" 2>&1 +rc_fast=$? +echo "=== FAST rc=$rc_fast ===" +tail -25 "$LOG" + +LOG2=alignment_lab/slurm/stagea_slow_$SLURM_JOB_ID.log +"$PY" -m pytest --run-slow tests/unit/alignment tests/unit/frf_separate \ + tests/integration/alignment -q > "$LOG2" 2>&1 +rc_slow=$? +echo "=== SLOW rc=$rc_slow ===" +tail -25 "$LOG2" +echo "rc_fast=$rc_fast rc_slow=$rc_slow" diff --git a/alignment_lab/analysis/stagea_verify.sh b/alignment_lab/analysis/stagea_verify.sh new file mode 100644 index 00000000..2bf6d417 --- /dev/null +++ b/alignment_lab/analysis/stagea_verify.sh @@ -0,0 +1,146 @@ +#!/bin/bash +# Stage A verification: every swap is claimed bit-identical, so check each one by +# computing BOTH forms in one process and comparing exactly. Stronger and far +# cheaper than diffing an end-to-end fingerprint against a baseline worktree. +# +# Also answers the two "verify, then delete" questions: do the calc-side +# resolution mask and the near-no-op obs mask drop any reflection at all? +#SBATCH --job-name=stagea +#SBATCH --output=alignment_lab/slurm/%x_%j.out +#SBATCH --error=alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:50:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=32G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=1 +export OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +mkdir -p alignment_lab/slurm +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" + +"$PY" -u - <<'PYEOF' +import math +import torch +torch.manual_seed(0) + +ok = True +def check(name, cond, detail=""): + global ok + ok = ok and bool(cond) + print(f" [{'PASS' if cond else 'FAIL'}] {name}{(' -- ' + detail) if detail else ''}") + +print("=== 1. Edmonds ZYZ: base primitive vs the deleted local copy ===") +from torchref.base.alignment.rotation import rotation_matrix_euler_zyz +a = torch.rand(5000, dtype=torch.float64) * 2 * math.pi +b = torch.rand(5000, dtype=torch.float64) * math.pi +g = torch.rand(5000, dtype=torch.float64) * 2 * math.pi + +def old_zyz(alpha, beta, gamma): + ca, sa = torch.cos(alpha), torch.sin(alpha) + cb, sb = torch.cos(beta), torch.sin(beta) + cg, sg = torch.cos(gamma), torch.sin(gamma) + return torch.stack([ + torch.stack([ca*cb*cg - sa*sg, -ca*cb*sg - sa*cg, ca*sb], dim=-1), + torch.stack([sa*cb*cg + ca*sg, -sa*cb*sg + ca*cg, sa*sb], dim=-1), + torch.stack([-sb*cg, sb*sg, cb ], dim=-1), + ], dim=-2) + +R_old = old_zyz(a, b, g) +R_new = rotation_matrix_euler_zyz(torch.stack([a, b, g], dim=-1)) +check("bitwise equal over 5000 random triples", torch.equal(R_old, R_new), + f"max|d|={ (R_old-R_new).abs().max().item():.3e}") + +print("=== 2. Rodrigues: rotation_utils vs the deleted align._rodrigues ===") +from torchref.experimental.alignment.frf.rotation_utils import axis_angle_to_matrix + +def old_rodrigues(omega): + if omega.dtype != torch.float64: + omega = omega.to(torch.float64) + single = omega.dim() == 1 + if single: + omega = omega.unsqueeze(0) + th = omega.norm(dim=-1, keepdim=True) + axis = omega / th.clamp(min=1e-30) + zeros = torch.zeros_like(axis[..., 0]) + K = torch.stack([ + torch.stack([zeros, -axis[..., 2], axis[..., 1]], dim=-1), + torch.stack([axis[..., 2], zeros, -axis[..., 0]], dim=-1), + torch.stack([-axis[..., 1], axis[..., 0], zeros], dim=-1), + ], dim=-2) + th_b = th.unsqueeze(-1) + eye = torch.eye(3, dtype=omega.dtype, device=omega.device).expand(*omega.shape[:-1], 3, 3) + R = eye + torch.sin(th_b) * K + (1.0 - torch.cos(th_b)) * torch.matmul(K, K) + return R.squeeze(0) if single else R + +# The production caller builds omegas as a float64 meshgrid, so replicate that. +c = torch.linspace(-0.1, 0.1, 11, dtype=torch.float64) +wx, wy, wz = torch.meshgrid(c, c, c, indexing="ij") +om = torch.stack([wx.flatten(), wy.flatten(), wz.flatten()], dim=-1) +check("bitwise equal on the pipeline's float64 perturbation grid", + torch.equal(old_rodrigues(om), axis_angle_to_matrix(om))) + +print("=== 3. Symop unroll: apply_to_hkl vs the deleted einsum ===") +from torchref.symmetry.spacegroup import SpaceGroup +for sg_name in ("P 1", "C 2", "P 21 21 21", "P 31 2 1", "P 65 2 2", "P 43 32", "P 4 3 2"): + sg = SpaceGroup(sg_name) + hkl = torch.randint(-40, 41, (4000, 3)) + rec = torch.eye(3, dtype=torch.float64) * 0.0137 + 0.0011 # arbitrary non-diagonal basis + old = torch.einsum( + "kji,nj->kni", sg.matrices.to(torch.float64), hkl.to(torch.float64) + ).reshape(-1, 3) + new = sg.apply_to_hkl(hkl).permute(2, 0, 1).reshape(-1, 3).to(torch.float64) + same_rows = torch.equal(old, new) + same_s = torch.equal(old @ rec, new @ rec) + check(f"{sg_name:12s} n_ops={sg.n_ops:2d} rows and s_obs bitwise equal", + same_rows and same_s, + "" if same_rows else f"row mismatch {int((old != new).any(-1).sum())}") + +print("=== 4. (L, d_min) pass-through replaces the second auto_lmax call ===") +from torchref.experimental.alignment.frf.api import phaser_lmax_resolution +from torchref.experimental.alignment.rotation_search import LMAX_CAP +for r, dmin in ((10.0, 1.8), (15.0, 2.05), (25.0, 1.6), (4.0, 3.0)): + L1, d1 = phaser_lmax_resolution(r, dmin, LMAX_CAP) + L2, d2 = phaser_lmax_resolution(r, dmin, LMAX_CAP) # the call that used to be inside + check(f"radius {r:5.1f} d_min {dmin:.2f} -> L={L1} d_min_eff={d1:.4f}", + (L1, d1) == (L2, d2)) + +print("=== 5. Do the two 'verify then delete' masks drop anything? ===") +from alignment_lab.lab.benchmark import load_case +from torchref.experimental.alignment.frf.dense_calc import dense_calc_via_box +from torchref.experimental.alignment.rotation_search import ( + LOW_RESOLUTION_CUTOFF_A, DENSE_CALC_PAD, +) +for pdb in ("3K7M", "1DAW"): + model, data = load_case(pdb)[:2] + rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) + s_mag_all = (data.hkl.to(torch.float64) @ rec_basis).norm(dim=-1) + d_min_data = float(1.0 / s_mag_all.max().item()) + d_max = float(LOW_RESOLUTION_CUTOFF_A) + + # (a) the obs mask at rotation_search.py:239 -- upper bound is max <= max + keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min_data) + n_lo = int((s_mag_all < 1.0 / d_max).sum()) + n_hi = int((s_mag_all > 1.0 / d_min_data).sum()) + print(f" {pdb}: obs mask keeps {int(keep.sum())}/{len(s_mag_all)} " + f"(dropped {n_lo} below {d_max:.0f} A, {n_hi} above d_min via the " + f"reciprocal round-trip)") + + # (b) the calc-side mask in score_model, against dense_calc's own window + model_radius_A = float((model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item()) + L, d_min_eff = phaser_lmax_resolution(model_radius_A, d_min_data, LMAX_CAP) + s_calc, F_calc = dense_calc_via_box(model, d_max, d_min_eff, pad=DENSE_CALC_PAD) + smag_calc = s_calc.norm(dim=-1) + keep_c = (smag_calc >= 1.0 / d_max) & (smag_calc <= 1.0 / d_min_eff) + dropped = int(len(smag_calc) - keep_c.sum()) + print(f" {pdb}: calc re-mask drops {dropped}/{len(smag_calc)} " + f"(L={L}, d_min_eff={d_min_eff:.3f} A)") + +print() +print("OVERALL:", "PASS" if ok else "FAIL") +PYEOF +rc=$? +echo "python_exit=$rc" diff --git a/alignment_lab/analysis/stageb_gate.sh b/alignment_lab/analysis/stageb_gate.sh new file mode 100644 index 00000000..8035d7a6 --- /dev/null +++ b/alignment_lab/analysis/stageb_gate.sh @@ -0,0 +1,49 @@ +#!/bin/bash +#SBATCH --job-name=stageb +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=32G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$NEW" +export PYTHONPATH="$NEW" TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" + +LOG=alignment_lab/slurm/stageb_bessel_$SLURM_JOB_ID.log +"$PY" -m pytest tests/unit/frf_separate/test_bessel_rescale.py -q > "$LOG" 2>&1 +rc_b=$? +echo "=== BESSEL TESTS rc=$rc_b ===" +tail -20 "$LOG" + +# End-to-end: the rescaled ladder must leave the peak lists untouched. Compare +# against the fingerprints captured for the Stage A gate (same tree, pre-Stage-B). +OUT=$NEW/alignment_lab/slurm +status=0 +for pdb in 3K7M 1DAW; do + for cap in 64 100; do + "$PY" -u "$NEW/alignment_lab/analysis/frf_fingerprint.py" --pdb "$pdb" \ + --lmax-cap "$cap" > "$OUT/fp_STAGEB_${pdb}_${cap}.txt" 2>"$OUT/fp_STAGEB_${pdb}_${cap}.err" + rc=$? + ref="$OUT/fp_NEW_${pdb}_${cap}.txt" + new="$OUT/fp_STAGEB_${pdb}_${cap}.txt" + if [ $rc -ne 0 ]; then + echo "RUN FAILED $pdb cap$cap rc=$rc"; tail -5 "$OUT/fp_STAGEB_${pdb}_${cap}.err"; status=1; continue + fi + if [ ! -s "$ref" ]; then echo "NO STAGE-A REFERENCE for $pdb cap$cap"; status=1; continue; fi + if diff -q <(grep -v '^#' "$ref") <(grep -v '^#' "$new") >/dev/null; then + echo "IDENTICAL $pdb cap$cap ($(grep -vc '^#' "$new") peaks)" + else + nd=$(diff <(grep -v '^#' "$ref") <(grep -v '^#' "$new") | grep -c '^[<>]') + echo "DIFFERS $pdb cap$cap ($nd lines)" + diff <(grep -v '^#' "$ref") <(grep -v '^#' "$new") | head -6 + status=1 + fi + done +done +echo "stageb_status=$status rc_bessel=$rc_b" diff --git a/alignment_lab/analysis/stagec_gate3.sh b/alignment_lab/analysis/stagec_gate3.sh new file mode 100644 index 00000000..e150d87f --- /dev/null +++ b/alignment_lab/analysis/stagec_gate3.sh @@ -0,0 +1,55 @@ +#!/bin/bash +# Stage C1, third pass. Hypothesis: with the back-half accumulation restored to +# double, the whole of C1 is numerically neutral -- so the fingerprints should be +# bit-identical to d1244c45, and the 1e-4 score shift seen in pass 2 was entirely +# the narrowed accumulation. +#SBATCH --job-name=stagec3 +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +OLD=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/_stagea_baseline +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +SCRIPT=$NEW/alignment_lab/analysis/frf_fingerprint.py +OUT=$NEW/alignment_lab/slurm +export TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" + +cd "$NEW"; export PYTHONPATH="$NEW" +LOG=$OUT/stagec3_tests_$SLURM_JOB_ID.log +"$PY" -m pytest tests/unit/alignment tests/unit/frf_separate tests/unit/model \ + tests/unit/test_imports_smoke.py -q > "$LOG" 2>&1 +rc=$? +echo "=== TESTS rc=$rc ===" +tail -12 "$LOG" + +status=0 +for pdb in 3K7M 1DAW; do + for cap in 64 100; do + cd "$NEW" + PYTHONPATH="$NEW" "$PY" -u "$SCRIPT" --pdb "$pdb" --lmax-cap "$cap" \ + > "$OUT/c3_NEW_${pdb}_${cap}.txt" 2>/dev/null + a=$OUT/c2_OLD_${pdb}_${cap}.txt # baseline captured in pass 2 + b=$OUT/c3_NEW_${pdb}_${cap}.txt + if diff -q <(grep '^FP ' "$a") <(grep '^FP ' "$b") >/dev/null; then + echo "IDENTICAL $pdb cap$cap ($(grep -c '^FP ' "$b") peaks)" + else + nd=$(diff <(grep '^FP ' "$a") <(grep '^FP ' "$b") | grep -c '^[<>]') + echo "DIFFERS $pdb cap$cap ($nd lines)" + diff <(grep '^FP ' "$a") <(grep '^FP ' "$b") | head -4 + status=1 + fi + done +done + +echo "=== timing, 4 threads for comparability with the 1.08s note ===" +export TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +"$PY" -u -m alignment_lab.diagnostics.frf_benchmark --pdb 3K7M --arms cap64 --trials 2 2>&1 \ + | grep -vE "Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|warn|^ *$" +echo "stagec3_status=$status tests_rc=$rc" diff --git a/alignment_lab/analysis/stagec_timing_paired.sh b/alignment_lab/analysis/stagec_timing_paired.sh new file mode 100644 index 00000000..0eeffd19 --- /dev/null +++ b/alignment_lab/analysis/stagec_timing_paired.sh @@ -0,0 +1,34 @@ +#!/bin/bash +# Paired timing: baseline worktree vs this one, same node, same job, INTERLEAVED +# and structure-major. A cross-job comparison against a remembered number is not +# a measurement -- node and contention differ. +#SBATCH --job-name=stagec_time +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=48G +#SBATCH --exclusive +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +OLD=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/_stagea_baseline +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +export TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" +echo "OLD=$(cd $OLD && git rev-parse --short HEAD) NEW=working tree" + +for round in 1 2 3; do + for pdb in 3K7M 1DAW; do + for tree in OLD NEW; do + eval "root=\$$tree" + cd "$root" + line=$(PYTHONPATH="$root" "$PY" -u -m alignment_lab.diagnostics.frf_benchmark \ + --pdb "$pdb" --arms cap64 --trials 2 2>/dev/null \ + | grep -E "^ *cap64" | tail -1) + echo "round$round $pdb $tree $line" + done + done +done diff --git a/docs/changelog.rst b/docs/changelog.rst index a22a44b7..760ebf54 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -15,6 +15,12 @@ Unreleased - Removed the rotation function's dead modules, engine variants, debug environment switches and unreachable knobs - Fixed the rotation function's dense model transform building a real-space grid and map-symmetry operator that its next three lines discarded; ``ModelFT.copy`` gained ``build_grid`` - The rotation function's Wigner small-d blocks are memoised, so a process running more than one search builds them once +- The rotation function now takes its working precision from ``dtypes.float`` and its device from ``resolve_device``, instead of hardcoding float64 and reading one input's device +- The rotation function's spherical-Bessel recurrence rescales by a power of two as it runs, so the ladder no longer needs float64's exponent range +- The rotation function's Wigner eigendecomposition and anisotropy fit moved to the host, so neither requires float64 on the accelerator +- Removed the rotation function's duplicate Euler, Rodrigues and reciprocal-symmetry helpers in favour of the shared primitives +- Removed the rotation function's redundant calc-side resolution mask and its second bandwidth/resolution coupling call +- ``bessel_sh_expand`` lost its unread ``chunk_size`` argument and ``french_wilson_preprocess`` its unread ``sqrt_mean_F2`` output Version 0.6.4 ---------- diff --git a/tests/unit/alignment/test_radial_truncation.py b/tests/unit/alignment/test_radial_truncation.py index b085e131..94b88c5a 100644 --- a/tests/unit/alignment/test_radial_truncation.py +++ b/tests/unit/alignment/test_radial_truncation.py @@ -16,6 +16,12 @@ These tests pin that invariant against the formula, so the agreement is asserted rather than inferred from the array shape. + +They run under ``double_cpu`` because the band's *width* and its *populated +extent* only coincide in double precision. At the working precision the radial +weight ``sqrt(2u+1) j_u(x)/x`` underflows for high ``u`` at small ``x`` -- see +``test_the_working_precision_truncates_the_band_further`` -- so counting +non-zeros there measures float32's exponent range, not Phaser's formula. """ import pytest @@ -42,7 +48,7 @@ def _phaser_nmax(lmax_even: int, l: int) -> int: @pytest.mark.parametrize("L", [21, 41, 67]) -def test_radial_band_matches_phaser_width(L): +def test_radial_band_matches_phaser_width(L, double_cpu): """Non-zero radial indices at each ``l`` must stop at Phaser's ``nmax``. Our ``n`` index is 0-based against Phaser's 1-based, so the condition is @@ -64,7 +70,7 @@ def test_radial_band_matches_phaser_width(L): @pytest.mark.parametrize("L", [21, 41, 67]) -def test_band_narrows_to_a_single_term_at_lmax(L): +def test_band_narrows_to_a_single_term_at_lmax(L, double_cpu): """The top band keeps exactly one radial term, the widest keeps them all.""" lmax_even = L - 1 if (L - 1) % 2 == 0 else L - 2 assert _phaser_nmax(lmax_even, lmax_even) == 1 @@ -78,7 +84,7 @@ def test_band_narrows_to_a_single_term_at_lmax(L): @pytest.mark.parametrize("L", [21, 41, 67]) -def test_band_is_exactly_phasers_width_not_merely_bounded(L): +def test_band_is_exactly_phasers_width_not_merely_bounded(L, double_cpu): """The populated band must *equal* Phaser's ``nmax(l)``, not just fit inside. A bound alone would also pass for an expansion that silently drops radial @@ -95,7 +101,7 @@ def test_band_is_exactly_phasers_width_not_merely_bounded(L): ) -def test_allocated_width_exceeds_the_populated_band(): +def test_allocated_width_exceeds_the_populated_band(double_cpu): """The array is wider than the support at every l above the first. This is the fact that makes the array shape misleading, and the reason the @@ -110,3 +116,28 @@ def test_allocated_width_exceeds_the_populated_band(): assert N_radial > _phaser_nmax(lmax_even, lmax_even), ( "allocated width should exceed the top band's single term" ) + + +def test_the_working_precision_truncates_the_band_further(): + """At float32 the tail of the radial band is not small, it is *absent*. + + ``j_u(x)`` for ``u >> x`` is genuinely negligible -- at ``x = 1.9`` every + ``u >= 33`` is below float32's smallest normal -- so at the working precision + the populated band is narrower than Phaser's ``nmax(l)`` wherever the shell's + Bessel argument is small. That is a property of the engine as shipped, not a + defect, and it is pinned here so it is not rediscovered as a regression: the + structural tests above deliberately run in double precision, which would + otherwise leave the shipped configuration untested. + """ + L = 67 + c = _expand(L) # working precision, no double_cpu + lmax_even = L - 1 if (L - 1) % 2 == 0 else L - 2 + populated = (c[:, 2, :].abs() > 0).any(dim=-1).nonzero().flatten() + assert populated.numel(), "l=2 band is entirely empty" + got = int(populated.max()) + 1 + phaser = _phaser_nmax(lmax_even, 2) + assert got < phaser, ( + f"expected the float32 tail to underflow: got {got} terms, Phaser has " + f"{phaser}. If this now matches, the working precision widened and the " + f"structural tests above no longer need double_cpu." + ) diff --git a/tests/unit/alignment/test_rotation_search_dtype_device.py b/tests/unit/alignment/test_rotation_search_dtype_device.py new file mode 100644 index 00000000..c0bd0c7d --- /dev/null +++ b/tests/unit/alignment/test_rotation_search_dtype_device.py @@ -0,0 +1,151 @@ +"""The rotation search takes its working precision and its device from config. + +Two properties, both easy to lose silently: + +* **Precision is `torchref.config`'s, not an argument's.** The expansion used to + be handed ``compute_dtype=torch.complex64`` by one call site, which made the + fused CPU Legendre kernel reachable only through that argument -- its + ``BackendTable`` row gates on ``dtypes=(torch.float32,)``, so dropping the + argument silently routed to the portable path. Now the default *is* float32, + so the gate matches by construction and flipping the config flips the engine. +* **The device is resolved from both inputs**, not read off whichever one the + code happens to touch first. Model and data on different devices used to + either cross-device or throw depending on which line ran. +""" + +import pytest +import torch + +from torchref.config import dtypes + +pytestmark = pytest.mark.unit + + +def _tiny_reflections(n=4000, seed=0): + """A P1 shell of reflections wide enough to bin into 20 Wilson shells.""" + g = torch.Generator().manual_seed(seed) + s_mag = 0.07 + 0.18 * torch.rand(n, generator=g, dtype=torch.float64) + theta = torch.acos(2 * torch.rand(n, generator=g, dtype=torch.float64) - 1) + phi = 6.283185307179586 * torch.rand(n, generator=g, dtype=torch.float64) + s_vec = torch.stack( + [s_mag * torch.sin(theta) * torch.cos(phi), + s_mag * torch.sin(theta) * torch.sin(phi), + s_mag * torch.cos(theta)], dim=-1, + ) + F = torch.randn(n, generator=g, dtype=torch.float64).abs() + 0.1 + centric = torch.zeros(n, dtype=torch.bool) + return s_vec, F, centric + + +@pytest.mark.parametrize("float_dtype,want_complex", [ + (torch.float32, torch.complex64), + (torch.float64, torch.complex128), +]) +def test_expansion_follows_the_configured_float_dtype(float_dtype, want_complex): + from torchref.experimental.alignment.frf.data_mr import bessel_sh_expand + + s_vec, F, _ = _tiny_reflections() + original = dtypes.float, dtypes.complex + try: + dtypes.float = float_dtype + dtypes.complex = want_complex + # `s_vec` stays float64 on purpose -- the clustering keys need that + # resolution -- so this also pins that the OUTPUT dtype is not inherited + # from the input. + out = bessel_sh_expand(s_vec, F, L=12, bessel_h_scale=40.0) + finally: + dtypes.float, dtypes.complex = original + assert out.coeffs.dtype == want_complex, ( + f"expansion returned {out.coeffs.dtype} with dtypes.float={float_dtype}" + ) + + +def test_the_back_half_accumulates_one_step_wider(): + """The tail is deliberately wider than the expansion, and must stay so. + + Narrowing it looks like free memory -- complex128 buffers holding + complex64-accurate content -- and it is not. The radial sum and the Wigner + contraction are both oscillatory, so they cancel, and single-precision + accumulation was measured to move scores by 1e-4 to 1.4e-3 relative and + leave only 1 of 500 candidate slots holding the same orientation. The + expansion's own working precision is config's; this accumulation is not. + """ + from torchref.experimental.alignment.frf.data_mr import ( + bessel_sh_expand, cross_correlate_xi, + ) + from torchref.experimental.alignment.frf.sitelist_ang import ( + build_dense_map_per_beta, + ) + from torchref.experimental.alignment.frf.wigner_d import ( + wigner_contraction_per_beta, + ) + + s_vec, F, _ = _tiny_reflections() + original = dtypes.float, dtypes.complex + try: + dtypes.float, dtypes.complex = torch.float32, torch.complex64 + c = bessel_sh_expand(s_vec, F, L=12, bessel_h_scale=40.0) + finally: + dtypes.float, dtypes.complex = original + + assert c.coeffs.dtype == torch.complex64, "expansion should be at config dtype" + xi = cross_correlate_xi(c, c) + assert xi.dtype == torch.complex128, ( + f"the radial accumulation narrowed to {xi.dtype}; see the docstring on " + f"cross_correlate_xi for what that costs" + ) + # Everything downstream follows xi rather than re-deciding. + betas = torch.linspace(0.0, 3.0, 5, dtype=torch.float64) + S = wigner_contraction_per_beta(xi, betas) + assert S.dtype == xi.dtype, f"Wigner did not follow xi: {S.dtype}" + M = build_dense_map_per_beta(xi, betas, fft_size=48) + assert M.dtype == xi.dtype, f"FFT did not follow xi: {M.dtype}" + + +def test_wigner_blocks_carry_no_float64_to_the_device(): + """The eigendecomposition is precision-critical but belongs on the host. + + What reaches the device is the ``d^l`` blocks, bounded in [-1, 1], at the + working precision -- so an accelerator without float64 is not excluded. + """ + from torchref.experimental.alignment.frf import wigner_d + + wigner_d.clear_wigner_d_cache() + betas = torch.linspace(0.0, 3.0, 4, dtype=torch.float64) + blocks = wigner_d._wigner_d_blocks( + 8, betas, torch.device("cpu"), torch.float32, + ) + assert all(b.dtype == torch.float32 for b in blocks) + # d^l(0) = I, the cheapest correctness check on the blocks themselves. + for l, b in enumerate(blocks, start=1): + eye = torch.eye(2 * l + 1, dtype=torch.float32) + assert torch.allclose(b[0], eye, atol=1e-5), f"d^{l}(0) is not I" + + +def test_anisotropy_fit_runs_on_the_host_in_double(): + """It is 7 parameters over ~1e4 reflections; precision there is worth more + than locality, and keeping it on the host removes a float64 requirement.""" + from alignment_lab.lab.benchmark import load_case + from torchref.experimental.alignment.rotation_search import ( + ANISO_FIT_WINDOW_A, fit_anisotropy, + ) + + _, data = load_case("1DAW")[:2] + d_max, d_min = ANISO_FIT_WINDOW_A + U = fit_anisotropy(data, d_min=d_min, d_max=d_max) + assert U.device.type == "cpu" + assert U.dtype == torch.float64 + assert U.shape == (3, 3) + assert torch.allclose(U, U.T, atol=1e-12), "U must be symmetric" + + +def test_device_is_resolved_from_both_inputs(): + """Reading one input's device is what let model and data disagree.""" + from alignment_lab.lab.benchmark import load_case + from torchref.utils import resolve_device + + model, data = load_case("1DAW")[:2] + # Data-first precedence, per the convention in torchref/maps/map.py. + resolved = resolve_device(data, model) + assert resolved == data.hkl.device or resolved.type == data.hkl.device.type + assert model.xyz().device.type == resolved.type diff --git a/tests/unit/alignment/test_wigner_d_cache.py b/tests/unit/alignment/test_wigner_d_cache.py index 57028eb4..f9890b1a 100644 --- a/tests/unit/alignment/test_wigner_d_cache.py +++ b/tests/unit/alignment/test_wigner_d_cache.py @@ -64,8 +64,8 @@ def test_different_data_at_the_same_bandwidth_still_differs(): def test_a_different_key_is_not_served_the_cached_blocks(L2, n_beta2): """Changing either the bandwidth or the β grid must rebuild.""" clear_wigner_d_cache() - ref_blocks = _wigner_d_blocks(9, _betas(12), torch.device("cpu")) - got = _wigner_d_blocks(L2, _betas(n_beta2), torch.device("cpu")) + ref_blocks = _wigner_d_blocks(9, _betas(12), torch.device("cpu"), torch.float64) + got = _wigner_d_blocks(L2, _betas(n_beta2), torch.device("cpu"), torch.float64) assert len(got) == L2 - 1 assert got[0].shape[0] == n_beta2 assert got is not ref_blocks @@ -77,7 +77,7 @@ def test_the_cache_holds_one_entry(): allowed to accumulate across keys.""" clear_wigner_d_cache() for L in (7, 9, 11): - _wigner_d_blocks(L, _betas(12), torch.device("cpu")) + _wigner_d_blocks(L, _betas(12), torch.device("cpu"), torch.float64) assert len(_WIGNER_D_CACHE) == 1 clear_wigner_d_cache() assert len(_WIGNER_D_CACHE) == 0 @@ -94,7 +94,7 @@ def test_the_blocks_are_the_wigner_small_d_matrices(): clear_wigner_d_cache() L = 7 betas = torch.tensor([0.0, 0.4, 1.7, 3.0], dtype=torch.float64) - blocks = _wigner_d_blocks(L, betas, torch.device("cpu")) + blocks = _wigner_d_blocks(L, betas, torch.device("cpu"), torch.float64) for l, d in enumerate(blocks, start=1): sz = 2 * l + 1 assert d.shape == (betas.numel(), sz, sz) diff --git a/tests/unit/frf_separate/test_bessel_rescale.py b/tests/unit/frf_separate/test_bessel_rescale.py new file mode 100644 index 00000000..200b277d --- /dev/null +++ b/tests/unit/frf_separate/test_bessel_rescale.py @@ -0,0 +1,128 @@ +"""What the Bessel ladder's on-the-fly rescaling is allowed to cost: nothing. + +Miller's downward recurrence for ``j_u(x)`` seeds at an arbitrary magnitude and +renormalises at the end, so the intermediate ladder is inflated by whatever the +true ``j_{n_start}(x)`` happens to be -- 1e157 at the FRF's low-resolution end. +That is fine in float64 and overflows float32 for every ``x`` below about 35, +which is most of the resolution range, so the recurrence rescales as it goes. + +The rescale factor is a power of two on purpose: dividing by one decrements the +exponent and leaves the mantissa alone, so the table must come out **bitwise** +equal to the un-rescaled ladder rather than merely close to it. That is the +property these tests pin, against two independent references -- the previous +implementation, transcribed inline, and ``scipy.special.spherical_jn``. +""" + +import math + +import pytest +import torch + +from torchref.experimental.alignment.frf.data_mr import ( + _BESSEL_RESCALE_EXP, + spherical_bessel_table, +) + +pytestmark = pytest.mark.unit + +#: Bessel arguments spanning the FRF's range. ``bessel_h_scale = lmax_even * +#: d_min_eff``, so ``x`` runs from ``bessel_h_scale / d_max`` (~1.3 at +#: d_max = 100 A) up to exactly ``lmax_even`` at the high-resolution limit. +_X_VALUES = [1.257, 1.885, 3.0, 5.0, 10.0, 20.0, 40.0, 64.0] + +_U_MAX = 65 # lmax_even + 1 at the shipped LMAX_CAP = 64 +_N_EXTRA = 25 + + +def _unrescaled_reference(x, u_max, n_extra=_N_EXTRA): + """The recurrence as it stood before rescaling, transcribed verbatim. + + Kept as a literal copy rather than a call into the module: the point is to + compare against the *previous* arithmetic, so it must not track any later + edit to the production function. + """ + x64 = x.to(torch.float64) + safe_x = x64.clamp(min=1e-30) + inv_x = 1.0 / safe_x + n_start = max(u_max + n_extra, u_max + 2) + j_high = torch.zeros_like(x64) + j_mid = torch.ones_like(x64) + j_table = torch.zeros((u_max + 1, *x64.shape), dtype=torch.float64) + peak = torch.zeros_like(x64) + for n in range(n_start, 0, -1): + j_low = (2.0 * n + 1.0) * inv_x * j_mid - j_high + if n - 1 <= u_max: + j_table[n - 1] = j_low + peak = torch.maximum(peak, j_low.abs()) + j_high = j_mid + j_mid = j_low + true_j0 = torch.sin(x64) * inv_x + true_j0 = torch.where(x64 < 1e-30, torch.ones_like(x64), true_j0) + computed_j0 = j_table[0] + safe_j0 = torch.where( + computed_j0.abs() < 1e-30, torch.ones_like(computed_j0), computed_j0, + ) + j_table = j_table * (true_j0 / safe_j0).unsqueeze(0) + perm = list(range(1, j_table.dim())) + [0] + return j_table.permute(*perm).contiguous(), peak + + +def test_rescaling_is_bit_identical_to_the_unrescaled_ladder(): + x = torch.tensor(_X_VALUES, dtype=torch.float64) + ref, _ = _unrescaled_reference(x, _U_MAX) + got = spherical_bessel_table(x, _U_MAX) + assert got.shape == ref.shape + assert torch.equal(got, ref), ( + "rescaling perturbed the table; max relative deviation " + f"{((got - ref).abs() / ref.abs().clamp(min=1e-300)).max().item():.3e}" + ) + + +def test_the_rescale_branch_is_actually_exercised(): + """Guard against a vacuous equality test. + + If the ladder never crossed the threshold the comparison above would pass + while testing nothing, so assert the un-rescaled ladder really does run away + -- and past float32's ceiling, which is the reason the rescaling exists. + """ + x = torch.tensor(_X_VALUES, dtype=torch.float64) + _, peak = _unrescaled_reference(x, _U_MAX) + threshold = 2.0 ** _BESSEL_RESCALE_EXP + f32_max = torch.finfo(torch.float32).max + crossed = int((peak > threshold).sum()) + over_f32 = int((peak > f32_max).sum()) + assert crossed >= len(_X_VALUES) - 2, ( + f"only {crossed} of {len(_X_VALUES)} arguments cross 2**" + f"{_BESSEL_RESCALE_EXP}; the rescale path is nearly dead" + ) + assert over_f32 >= 5, ( + f"only {over_f32} arguments overflow float32 (peaks: " + f"{[f'{v:.2e}' for v in peak.tolist()]})" + ) + + +def test_matches_scipy_spherical_jn(): + """Independent oracle, over the range where float64 carries the answer.""" + scipy_special = pytest.importorskip("scipy.special") + x = torch.tensor(_X_VALUES, dtype=torch.float64) + got = spherical_bessel_table(x, _U_MAX) + for i, xv in enumerate(_X_VALUES): + for u in range(0, _U_MAX + 1): + want = float(scipy_special.spherical_jn(u, xv)) + # Below ~1e-290 the reference itself is at the edge of float64, and + # the FRF flushes anything under float32's smallest normal anyway. + if abs(want) < 1e-30: + continue + mine = float(got[i, u]) + assert math.isclose(mine, want, rel_tol=1e-9, abs_tol=1e-300), ( + f"j_{u}({xv}) = {mine:.12e}, scipy says {want:.12e}" + ) + + +def test_batched_shape_and_dtype_round_trip(): + x = torch.rand(7, 3, dtype=torch.float64) * 60.0 + 1.3 + got = spherical_bessel_table(x, 20) + assert got.shape == (7, 3, 21) + assert got.dtype == torch.float64 + got32 = spherical_bessel_table(x.to(torch.float32), 20) + assert got32.dtype == torch.float32 diff --git a/tests/unit/frf_separate/test_bessel_sh_grouping.py b/tests/unit/frf_separate/test_bessel_sh_grouping.py index 9edfe550..7bfb3d75 100644 --- a/tests/unit/frf_separate/test_bessel_sh_grouping.py +++ b/tests/unit/frf_separate/test_bessel_sh_grouping.py @@ -14,6 +14,12 @@ every reflection is its own group -- i.e. the ungrouped sum. Comparing against the previous implementation instead would only show that two approximations agree with each other. + +The claims about the grouping being *loss-free* are only meaningful at a +precision finer than the loss being ruled out, so the tests that assert +1e-10-and-below take ``double_cpu``. At the working precision (float32) the +floor is float32 epsilon times the accumulation depth -- measured 1.4e-06 and +3.6e-07 on these cases -- which says nothing about the grouping. """ import math @@ -82,7 +88,7 @@ def test_grouping_error_stays_far_below_the_reference_implementation( ) -def test_a_lattice_groups_without_loss(ungrouped): +def test_a_lattice_groups_without_loss(ungrouped, double_cpu): """On a lattice the degeneracy is exact, so the grouping is free.""" s, I = _grid_set() ref = ungrouped(s, I, L=65, bessel_h_scale=64.0) @@ -187,7 +193,7 @@ def _reference_expansion(s_vec, intensity, L, bessel_h_scale): @pytest.mark.parametrize("L", [9, 13]) -def test_matches_an_independent_direct_summation(L): +def test_matches_an_independent_direct_summation(L, double_cpu): """The whole expansion, against a reference that shares no code with it.""" s, I = _random_set(seed=31, n=120) ref = _reference_expansion(s, I, L, 24.0) diff --git a/torchref/experimental/alignment/align.py b/torchref/experimental/alignment/align.py index 9de2e600..d4f632e8 100644 --- a/torchref/experimental/alignment/align.py +++ b/torchref/experimental/alignment/align.py @@ -4,7 +4,7 @@ This module hosts the heavy, reusable stage helpers — Lattman-Love / anisotropy data prep (`_prepare_frf_inputs`), the solvent-aware R-work (`_external_rwork`), the direct-SF translation evaluator -(`_DirectModelEvaluator`), the Rodrigues helper (`_rodrigues`) and the stage +(`_DirectModelEvaluator`) and the stage timer (`_StageTimer`) — that are shared by the rotation-ranking benchmarks and by the orchestrator. @@ -156,43 +156,6 @@ def evaluate(self, R, hkl, real_cell, return_amplitude=False): return f.abs() if return_amplitude else f -def _rodrigues(omega: torch.Tensor) -> torch.Tensor: - """Rodrigues axis-angle → SO(3). `omega = θ · axis` (radians). - - Accepts shape (3,) for a single rotation or (..., 3) for a batched stack - and returns matching (3, 3) or (..., 3, 3). The small-θ limit is handled - implicitly: sin(θ)→0 and (1-cos θ)→0 zero out the K and K² contributions - so R→I as θ→0; `clamp(min=1e-30)` prevents NaN from axis=0/0. - """ - if omega.dtype != torch.float64: - omega = omega.to(torch.float64) - is_single = omega.dim() == 1 - if is_single: - omega = omega.unsqueeze(0) - - th = omega.norm(dim=-1, keepdim=True) # (..., 1) - axis = omega / th.clamp(min=1e-30) # (..., 3) - zeros = torch.zeros_like(axis[..., 0]) - K = torch.stack([ - torch.stack([zeros, -axis[..., 2], axis[..., 1]], dim=-1), - torch.stack([axis[..., 2], zeros, -axis[..., 0]], dim=-1), - torch.stack([-axis[..., 1], axis[..., 0], zeros], dim=-1), - ], dim=-2) # (..., 3, 3) - - th_b = th.unsqueeze(-1) # (..., 1, 1) - sin_th = torch.sin(th_b) - cos_th = torch.cos(th_b) - - eye = torch.eye(3, dtype=omega.dtype, device=omega.device) - eye_b = eye.expand(*omega.shape[:-1], 3, 3) - KK = torch.matmul(K, K) - R = eye_b + sin_th * K + (1.0 - cos_th) * KK - - if is_single: - R = R.squeeze(0) - return R - - # --------------------------------------------------------------------------- # FRF input preparation (shared by the live pipeline and the rotation-ranking # benchmark in tests/integration/alignment/benchmark_rotation_ranking.py) @@ -350,7 +313,7 @@ def align_model_to_data( # Imported lazily to avoid an import cycle: `pipeline` imports the stage # helpers (`_prepare_frf_inputs`, `_external_rwork`, - # `_DirectModelEvaluator`, `_rodrigues`, `_StageTimer`) from this module. + # `_DirectModelEvaluator`, `_StageTimer`) from this module. from .pipeline import MolecularReplacementPipeline pipeline = MolecularReplacementPipeline( diff --git a/torchref/experimental/alignment/frf/api.py b/torchref/experimental/alignment/frf/api.py index ee6ae41d..eb1b6296 100644 --- a/torchref/experimental/alignment/frf/api.py +++ b/torchref/experimental/alignment/frf/api.py @@ -137,25 +137,14 @@ def __init__( n_wilson_shells: int = 20, sig_F_obs: Optional[torch.Tensor] = None, grid_sampling_deg: float = 2.0, - model_radius_A: Optional[float] = None, - auto_lmax: bool = False, - lmax_cap: int = 64, - compute_dtype: Optional[torch.dtype] = None, ): self.device = s_obs.device - self.real_dtype = s_obs.dtype - self.compute_dtype = compute_dtype - - # Phaser-faithful coupling of bandwidth to resolution (runMR_FRF.cc:408). - # Overrides L and d_min so the SH expansion is not flooded with data - # finer than L can represent (the high-symmetry failure mode). - if auto_lmax: - if model_radius_A is None or d_min is None: - raise ValueError( - "auto_lmax=True requires model_radius_A and d_min (data res)." - ) - L, d_min = phaser_lmax_resolution(model_radius_A, d_min, lmax_cap=lmax_cap) + # `L` and `d_min` arrive already coupled: the caller runs + # `phaser_lmax_resolution` because it needs the same pair to size the + # dense calc box, so re-deriving them here would only rediscover what it + # computed a line earlier. `d_min` is therefore the *coarsened* limit -- + # the resolution this bandwidth can represent -- not the data's own. self.L = L self.d_min = d_min self.d_max = d_max @@ -182,8 +171,9 @@ def __init__( # `bessel_h_scale = 2 pi R_patt` is the Patterson integration radius # (the chi_Omega sphere): the Bessel argument is # `h = bessel_h_scale |s|`, so the radial basis represents the - # Patterson out to `R_patt = bessel_h_scale / (2 pi)`. Under - # `auto_lmax` that comes to `R_patt ~ sphereOuter = 2 x mean radius`. + # Patterson out to `R_patt = bessel_h_scale / (2 pi)`. With the + # bandwidth-coupled `d_min` the caller passes, that comes to + # `R_patt ~ sphereOuter = 2 x mean radius`. if d_min is None: raise ValueError("d_min is required to set the Bessel scaling") lmax = L - 1 @@ -216,11 +206,13 @@ def __init__( zsymm = detect_zsymm(sym_mats) # 6. Bessel-SH expand the obs side. + # `s_obs` stays at the caller's (wider) dtype so the expansion's + # clustering keys keep their resolution; the intensity does not need to, + # and the expansion casts it to the working precision anyway. self._c_obs = bessel_sh_expand( - s_obs, intensity_obs.to(self.real_dtype), + s_obs, intensity_obs, L=L, bessel_h_scale=self.bessel_h_scale, zsymm=zsymm, enforce_friedel=True, - compute_dtype=self.compute_dtype, ) @@ -235,10 +227,17 @@ def score_model( solvent_fsol: float = 0.95, solvent_bsol: float = 300.0, ) -> Tuple[AdaptiveRotationFunction, List[RotationPeak]]: - # Resolution mask + Wilson + Eterm on calc. - s_calc, (F_calc,), smag_calc = _resolution_mask( - s_calc, (F_calc,), self.d_min, self.d_max, - ) + """Score one model's transform against the prepared observations. + + ``s_calc`` / ``F_calc`` must already lie inside ``[d_max, d_min]`` -- + this method does not re-mask them. The window is the caller's because + the caller had to know it to build the calc set in the first place. + """ + # The caller owns the calc-side resolution window: `dense_calc_via_box` + # samples the box over `[d_max, d_min]` already, and re-masking here was + # measured to drop 0 of 339040 reflections on 3K7M and 0 of 271630 on + # 1DAW. So take `s_calc` as given and only derive |s| from it. + smag_calc = s_calc.norm(dim=-1) E_calc, _ = wilson_normalise(F_calc, smag_calc, self.n_wilson_shells) eterm = eterm_sigma_a(smag_calc, self.delta_vrms_A) # Optional Babinet bulk-solvent factor: Phaser folds it into σ_A as @@ -258,10 +257,9 @@ def score_model( # discriminates truth — a calc-side m-filter was tested and is strongly # harmful on high-symmetry cases (3K7M rank 8→92), so the knob was removed. c_calc = bessel_sh_expand( - s_calc, intensity_calc.to(self.real_dtype), + s_calc, intensity_calc, L=self.L, bessel_h_scale=self.bessel_h_scale, zsymm=1, enforce_friedel=True, - compute_dtype=self.compute_dtype, ) # 8. Cross-correlate over the radial axis. diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index 79bc8e92..76b97995 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -58,6 +58,7 @@ #: approximation and needs its own evidence. _GROUP_SCALE_COS = 10_000_000 +from ....config import get_complex_dtype, get_float_dtype from ....utils.backends import run_or_degrade, select from ..sh import legendre_recurrence_coefficients from ._backends import LEGENDRE_BACKENDS @@ -71,6 +72,22 @@ ] +#: Exponent of the running-magnitude rescale in :func:`spherical_bessel_table`. +#: +#: The unnormalised downward ladder is enormous before it is renormalised: at the +#: FRF's low-resolution end (``x = bessel_h_scale / d_max``, about 1.3) the +#: intermediates reach 1e157, and they overflow float32 for every ``x`` below +#: ~35 -- most of the resolution range. Rescaling by a fixed factor whenever the +#: running value crosses it keeps the ladder in range at ANY working precision. +#: +#: It has to be a power of two. Dividing by one only decrements the exponent, so +#: the mantissas of every stored value and of the closing renormalisation are +#: untouched and the rescale introduces no rounding at all -- the table comes out +#: bit-identical to the un-rescaled version. Measured event counts over the FRF's +#: range: 5 rescales at x=1.26, 4 at 1.9, 3 at 5, 2 at 10, 1 at 20, 0 at 64. +_BESSEL_RESCALE_EXP = 100 + + def spherical_bessel_table( x: torch.Tensor, u_max: int, @@ -85,9 +102,18 @@ def spherical_bessel_table( Seed: ``n_start = u_max + n_extra``, ``j_{n_start+1} = 0``, ``j_{n_start} = 1`` (unnormalised), recur down to ``j_0``, then - renormalise using the exact ``j_0(x) = sin(x) / x``. Float64 - internally for accuracy at moderate ``u/x`` ratios; cast back to - input dtype on return. + renormalise using the exact ``j_0(x) = sin(x) / x``. + + The ladder is rescaled on the fly by ``2**-_BESSEL_RESCALE_EXP`` whenever it + grows past that magnitude -- see that constant for why the recurrence needs + it and why it costs no accuracy. Every rescale divides ``j_mid``, ``j_high`` + **and every row already written**, so the whole table stays in one common + frame and the closing renormalisation cancels it exactly. + + Note that the high-``u`` rows are genuinely negligible rather than merely + small: at ``x = 1.9`` every ``u >= 33`` is below float32's smallest normal, + which is 50% of the band, so those entries are already flushed to zero + wherever the caller works in single precision. Returns ------- @@ -106,6 +132,11 @@ def spherical_bessel_table( j_table = torch.zeros( (u_max + 1, *x64.shape), dtype=torch.float64, device=device, ) + threshold = float(2 ** _BESSEL_RESCALE_EXP) + inv_threshold = 1.0 / threshold + # Rescales applied so far, per element. Every element's ladder sits in the + # single frame 2**(-_BESSEL_RESCALE_EXP * n_rescales). + n_rescales = torch.zeros_like(x64, dtype=torch.int32) for n in range(n_start, 0, -1): j_low = (2.0 * n + 1.0) * inv_x * j_mid - j_high @@ -114,11 +145,31 @@ def spherical_bessel_table( j_high = j_mid j_mid = j_low + # The `.any()` costs a host sync per step (~90 per call, a handful of + # calls per search) and buys skipping a pass over the written rows on + # every step that does not need one. The rows are the expensive part. + over = j_mid.abs() > threshold + if bool(over.any()): + factor = torch.where(over, inv_threshold, 1.0) + j_mid = j_mid * factor + j_high = j_high * factor + if n - 1 <= u_max: + j_table[n - 1:] = j_table[n - 1:] * factor + n_rescales = n_rescales + over.to(torch.int32) + true_j0 = torch.sin(x64) * inv_x true_j0 = torch.where(x64 < 1e-30, torch.ones_like(x64), true_j0) computed_j0 = j_table[0] + # The degeneracy guard's 1e-30 is an absolute bound on the UNSCALED j_0, so + # express it in the frame the ladder actually ended up in. `frame` is an + # exact power of two; it underflows to 0 past ~10 rescales, at which point + # the guard simply stops firing (it is unreachable for the FRF anyway, whose + # x >= bessel_h_scale / d_max keeps sin(x)/x well away from zero). + frame = torch.ldexp( + torch.ones_like(x64), -_BESSEL_RESCALE_EXP * n_rescales, + ) safe_j0 = torch.where( - computed_j0.abs() < 1e-30, torch.ones_like(computed_j0), computed_j0, + computed_j0.abs() < 1e-30 * frame, frame, computed_j0, ) scale = true_j0 / safe_j0 j_table = j_table * scale.unsqueeze(0) @@ -136,8 +187,6 @@ def bessel_sh_expand( bessel_h_scale: float, zsymm: int = 1, enforce_friedel: bool = True, - chunk_size: int = -1, - compute_dtype: "torch.dtype | None" = None, ) -> BesselSHCoefficients: """Phaser-style ``c_nlm = Σ_h Y*_lm(ŝ) · I · sqrt(2u+1) · j_u(h)/h``. @@ -153,47 +202,32 @@ def bessel_sh_expand( * radial × SH expansion, sqrt(2u+1)·j_u(h)/h weight: DataMR.cc:993, 1107 * even-l only (Patterson centrosymmetry) + m-filter: DataMR.cc:863-870, 1117 - Parameters - ---------- - chunk_size : int - Reflections per chunk. ``-1`` (default) auto-sizes so the per-chunk - Y_lm block ``(chunk, L, 2L-1)`` stays near ~256 MB. - compute_dtype : torch.dtype, optional - Complex dtype for the dominant per-chunk einsum (the radial × SH - contraction — the FRF's FLOP bottleneck). Default ``None`` uses the - full-precision complex dtype matching the input. Passing - ``torch.complex64`` runs the contraction in single precision (a large - speedup on GPUs where FP64 is rate-limited), while the Bessel recurrence - and Legendre/Y_lm precompute stay at the input precision and the - cross-chunk accumulator stays at full precision — so only the contraction - loses precision, not the recurrences or the running sum. + Two precisions are in play and they are deliberately different. + + The **clustering keys** are computed at ``s_vectors``' own dtype, because + ``_GROUP_SCALE_S`` keys ``|s|`` at 1e-7 and that is exactly where float32's + resolution runs out: at ``|s| = 0.5`` a float32 rounding is ~0.3 of a key + step, so reflections that are mathematically degenerate would sometimes land + in adjacent keys and the degeneracy collapse the cost model depends on would + fray. Callers therefore pass float64 ``s_vectors`` even when the rest of the + chain is single precision. + + Everything else -- the Legendre/Y_lm precompute, the radial weights, the + contraction and the returned coefficients -- runs at + :func:`torchref.config.get_float_dtype`, which is this codebase's working + precision and the dtype the fused CPU kernel is built for. The + spherical-Bessel recurrence keeps its own float64 internals, where the + downward ladder needs the dynamic range. """ assert s_vectors.dim() == 2 and s_vectors.shape[-1] == 3 assert intensity.dim() == 1 and intensity.shape[0] == s_vectors.shape[0] - real_dtype = s_vectors.dtype - if real_dtype == torch.float64: - complex_dtype = torch.complex128 - elif real_dtype == torch.float32: - complex_dtype = torch.complex64 - else: - raise TypeError(f"Unsupported real dtype: {real_dtype}") device = s_vectors.device - - # Working precision for the per-chunk precompute (angles, Y_lm, Bessel - # weights) and the contraction. When the caller opts into complex64 we run - # the whole chunk in float32 — the dominant cost on GPUs where FP64 is - # rate-limited is not just the einsum but also the Legendre/Y_lm recurrence, - # so both must drop to single precision to matter. The spherical-Bessel - # downward recurrence keeps its float64 internals (it is the most - # cancellation-prone step) and the cross-chunk accumulator stays at full - # complex precision. - if compute_dtype == torch.complex64: - comp_real = torch.float32 - elif compute_dtype == torch.complex128: - comp_real = torch.float64 - else: - comp_real = real_dtype + # The working precision, and the dtype of everything returned. Not derived + # from the input: the input is deliberately wider (see the docstring). + comp_real = get_float_dtype() + complex_dtype = get_complex_dtype() + real_dtype = comp_real if enforce_friedel: s_vectors = torch.cat([s_vectors, -s_vectors], dim=0) @@ -235,7 +269,7 @@ def bessel_sh_expand( even_l_idx = torch.tensor(even_ls, dtype=torch.long, device=device) M = s_vectors.shape[0] - einsum_dtype = compute_dtype if compute_dtype is not None else complex_dtype + einsum_dtype = complex_dtype prof = {"cluster": 0.0, "dbuild": 0.0, "bessel": 0.0, "legendre": 0.0, "scatter": 0.0, "contract": 0.0} if _PROFILE else None @@ -486,14 +520,35 @@ def cross_correlate_xi( xi[l, m, n] = Σ_r c_obs[r, l, n] · conj(c_calc[r, l, m]) so that the peak Euler triple satisfies ``s_calc = R · s_obs``. + **Accumulated one step wider than the coefficients.** The radial sum runs + over oscillating ``j_u``, so the terms alternate in sign and cancel; the + relative error on the result is then far worse than ``eps * sqrt(n_terms)`` + would suggest, and it compounds through the equally oscillatory Wigner + contraction and the FFT downstream. Accumulating single-precision data in + double is the ordinary remedy and it is cheap here -- ``xi`` is + ``(L, 2L-1, 2L-1)``, 17 MB at L=65, against the 4.4M-element FFT it feeds. + + Running the whole tail in single instead was measured on 3K7M and 1DAW: the + top peak and its z-score were unchanged to seven figures, but scores moved + by 1e-4 to 1.4e-3 relative and only **1 of 500** candidate slots still held + the same orientation, because the greedy SO(3) NMS is sequential and a + reordering cascades through the suppression decisions. The candidate list is + what the placement search consumes, so that is not a free trade. + Returns ------- xi : torch.Tensor (complex), shape (L, 2L-1, 2L-1) """ if c_obs.L != c_calc.L: raise ValueError(f"L mismatch: obs={c_obs.L} calc={c_calc.L}") + acc = ( + torch.complex128 + if c_obs.coeffs.dtype in (torch.complex64, torch.complex128) + and torch.finfo(c_obs.coeffs.real.dtype).bits <= 32 + else c_obs.coeffs.dtype + ) return torch.einsum( "rln,rlm->lmn", - c_obs.coeffs, - torch.conj(c_calc.coeffs), + c_obs.coeffs.to(acc), + torch.conj(c_calc.coeffs).to(acc), ) diff --git a/torchref/experimental/alignment/frf/french_wilson.py b/torchref/experimental/alignment/frf/french_wilson.py index 9d02ee6e..ad5a5f7d 100644 --- a/torchref/experimental/alignment/frf/french_wilson.py +++ b/torchref/experimental/alignment/frf/french_wilson.py @@ -3,8 +3,8 @@ Pure ports of Phaser's ``lib/math_FrenchWilson.cc`` (centric/acentric posterior moments via Parabolic-cylinder ratios) and the Halley-iteration ``getDfactor`` in ``lib/math_RiceLLG.cc``. The public entry point -:func:`french_wilson_preprocess` returns ``(eEobs, DFAC, sqrt_mean_F2)`` -from raw ``(F, σF, |s|, centric)``. +:func:`french_wilson_preprocess` returns ``(eEobs, DFAC)`` from raw +``(F, σF, |s|, centric)``. Everything except ``french_wilson_preprocess`` is module-private; expose the public name through :mod:`torchref.experimental.alignment.frf.preprocessing`. @@ -449,7 +449,6 @@ def french_wilson_preprocess( Returns a dict with torch tensors back on the input device: eEobs: (N,) effective normalised amplitude DFAC : (N,) per-reflection D-factor ∈ [1e-7, 1−1e-7] - sqrt_mean_F2: (N,) per-reflection √_p """ import numpy as np @@ -475,7 +474,6 @@ def french_wilson_preprocess( mean_F2 = mean_F2 / np.maximum(counts, 1) mean_F2 = np.maximum(mean_F2, 1e-12) mean_I_per_h = mean_F2[shell_idx] - sqrt_mean_F2 = np.sqrt(mean_I_per_h) eosq = F2 / mean_I_per_h sigesq = 2.0 * F_np * sigF_np / mean_I_per_h @@ -501,5 +499,4 @@ def french_wilson_preprocess( return { "eEobs": torch.from_numpy(eEobs).to(device=device, dtype=F.dtype), "DFAC": torch.from_numpy(DFAC).to(device=device, dtype=F.dtype), - "sqrt_mean_F2": torch.from_numpy(sqrt_mean_F2).to(device=device, dtype=F.dtype), } diff --git a/torchref/experimental/alignment/frf/peak_finder.py b/torchref/experimental/alignment/frf/peak_finder.py index 91255f4e..c6369a17 100644 --- a/torchref/experimental/alignment/frf/peak_finder.py +++ b/torchref/experimental/alignment/frf/peak_finder.py @@ -19,6 +19,7 @@ import torch +from ....base.alignment.rotation import rotation_matrix_euler_zyz from .types import AdaptiveRotationFunction, RotationPeak __all__ = [ @@ -26,35 +27,6 @@ ] -def _euler_to_matrix_edmonds_zyz( - alpha: torch.Tensor, beta: torch.Tensor, gamma: torch.Tensor, -) -> torch.Tensor: - """R = R_z(α) R_y(β) R_z(γ) — Edmonds ZYZ convention. - - Returns shape (*alpha.shape, 3, 3) real. - - Algebraically ``rotation_utils.rotation_matrix_from_edmonds_euler_batch``, - but written as one fused pass rather than three matrix products. The NMS - below evaluates it over ~1e4 candidates, and the two forms round - differently in the last bit, which flips the suppression decision for pairs - sitting on the threshold. Every measurement on this engine was made with - this form, so it stays. - """ - ca, sa = torch.cos(alpha), torch.sin(alpha) - cb, sb = torch.cos(beta), torch.sin(beta) - cg, sg = torch.cos(gamma), torch.sin(gamma) - # R = Rz(a) * Ry(b) * Rz(g) - R = torch.stack( - [ - torch.stack([ca * cb * cg - sa * sg, -ca * cb * sg - sa * cg, ca * sb], dim=-1), - torch.stack([sa * cb * cg + ca * sg, -sa * cb * sg + ca * cg, sa * sb], dim=-1), - torch.stack([-sb * cg, sb * sg, cb ], dim=-1), - ], - dim=-2, - ) - return R - - def _so3_greedy_nms( alphas: torch.Tensor, betas: torch.Tensor, @@ -77,8 +49,13 @@ def _so3_greedy_nms( # preallocated kept-buffer (no repeated torch.stack), and a cosine threshold # (no per-iteration arccos). Result is identical to the original distance test. order = torch.argsort(values, descending=True).cpu().tolist() + # `rotation_matrix_euler_zyz` is the same fused single-pass form, term for + # term, so it rounds identically. That matters here and not only for tidiness: + # the NMS threshold test below flips for pairs sitting exactly on it, and the + # three-matrix-product form in `rotation_utils` rounds differently in the last + # bit. Every measurement on this engine was made with the fused form. R_all = ( - _euler_to_matrix_edmonds_zyz(alphas, betas, gammas) + rotation_matrix_euler_zyz(torch.stack([alphas, betas, gammas], dim=-1)) .to(torch.float64).cpu() ) # (n, 3, 3) # angle > nms_radius ⇔ cos(angle) < cos(nms_radius); cos(angle) from trace. diff --git a/torchref/experimental/alignment/frf/rotation_utils.py b/torchref/experimental/alignment/frf/rotation_utils.py index ead4c86c..5d75b1f9 100644 --- a/torchref/experimental/alignment/frf/rotation_utils.py +++ b/torchref/experimental/alignment/frf/rotation_utils.py @@ -87,9 +87,12 @@ def axis_angle_to_matrix(omega: torch.Tensor) -> torch.Tensor: Accepts ``(3,)`` for a single rotation or ``(..., 3)`` for a batched stack and returns ``(3, 3)`` or ``(..., 3, 3)``. The small-θ limit is handled implicitly (sin θ→0, (1−cos θ)→0 ⇒ R→I); ``clamp(min=1e-30)`` guards the - axis normalisation at θ=0. Mirrors ``align._rodrigues`` but lives here so - both the rescore and the alignment pipeline can share it without a circular - import. + axis normalisation at θ=0. + + Preferred over ``base.alignment.rotation.axis_angle_to_rotation_matrix``, + which accepts only ``(3,)``/``(N, 3)`` and switches the axis to ``[0, 0, 1]`` + below θ = 1e-10 rather than letting the trigonometric factors vanish. Above + that threshold the two agree term for term. """ if omega.dtype not in (torch.float32, torch.float64): omega = omega.to(torch.float64) diff --git a/torchref/experimental/alignment/frf/sitelist_ang.py b/torchref/experimental/alignment/frf/sitelist_ang.py index 1dba0fec..852ff1c2 100644 --- a/torchref/experimental/alignment/frf/sitelist_ang.py +++ b/torchref/experimental/alignment/frf/sitelist_ang.py @@ -40,6 +40,7 @@ import torch +from ....config import canonical_device from .types import AdaptiveRotationFunction from .wigner_d import wigner_contraction_per_beta @@ -164,7 +165,7 @@ def build_adaptive_sample_list( Python scan), so there is no host sync inside the loop. """ device = torch.device(device) if not isinstance(device, torch.device) else device - cache_key = (float(grid_sampling_deg), str(device), dtype) + cache_key = (float(grid_sampling_deg), str(canonical_device(device)), dtype) cached = _SAMPLE_LIST_CACHE.get(cache_key) if cached is not None: return cached diff --git a/torchref/experimental/alignment/frf/wigner_d.py b/torchref/experimental/alignment/frf/wigner_d.py index 29a7791f..a312993c 100644 --- a/torchref/experimental/alignment/frf/wigner_d.py +++ b/torchref/experimental/alignment/frf/wigner_d.py @@ -15,10 +15,16 @@ import torch +from ....config import canonical_device + __all__ = ["clear_wigner_d_cache", "wigner_contraction_per_beta"] -#: Memo for the per-l J_y eigendecomposition, keyed on (L, device-str). -#: It depends only on the bandwidth, so repeat calls at the same L reuse it. +#: Memo for the per-l J_y eigendecomposition, keyed on L alone. +#: It depends only on the bandwidth, so repeat calls at the same L reuse it. Held +#: on the HOST in float64/complex128: it is an eigendecomposition, the most +#: precision-sensitive step here, and it is small (L-1 matrices, the largest +#: 129x129) and data-independent. Keeping it off the accelerator costs nothing +#: measurable and means no float64 is required there. _WIGNER_EIG_CACHE: dict = {} #: Memo for the per-l small-d blocks, keyed on (L, betas, device-str). Holds at @@ -31,20 +37,21 @@ _WIGNER_D_CACHE: dict = {} -def _wigner_eig_table(L: int, device: torch.device): +def _wigner_eig_table(L: int): """Return [(w_l, V_l)] for l ∈ [1, L) — the J_y eigendecomposition per l. ``d^l(β) = Re(V_l · diag(e^{-iβ w_l}) · V_l^H)``. ``w_l ≈ [-l..l]`` and - ``V_l`` are independent of β and the data, so they are memoised. + ``V_l`` are independent of β and the data, so they are memoised. Built and + kept on the host; see ``_WIGNER_EIG_CACHE``. """ - key = (int(L), str(device)) + key = int(L) cached = _WIGNER_EIG_CACHE.get(key) if cached is not None: return cached table = [] for l in range(1, L): sz = 2 * l + 1 - p = torch.arange(sz - 1, dtype=torch.float64, device=device) + p = torch.arange(sz - 1, dtype=torch.float64) sup = 0.5 * torch.sqrt((2 * l - p) * (p + 1.0)) A = torch.diag(sup, 1) - torch.diag(sup, -1) # A = -i J_y w, V = torch.linalg.eigh(1j * A.to(torch.complex128)) # w∈[-l..l] @@ -58,35 +65,48 @@ def clear_wigner_d_cache() -> None: _WIGNER_D_CACHE.clear() -def _wigner_d_blocks(L: int, betas: torch.Tensor, device: torch.device): +def _wigner_d_blocks(L: int, betas: torch.Tensor, device: torch.device, + dtype: torch.dtype): """Per-l real ``d^l(β)`` blocks for ``l ∈ [1, L)``, memoised. - Each entry is ``(n_beta, 2l+1, 2l+1)`` float64. They depend only on the - bandwidth and the β grid, not on the data, so a process that runs more than - one rotation search at the same bandwidth builds them once. A single search - asks for them exactly once and so pays the full build. + Each entry is ``(n_beta, 2l+1, 2l+1)`` in ``dtype``, on ``device``. They + depend only on the bandwidth and the β grid, not on the data, so a process + that runs more than one rotation search at the same bandwidth builds them + once. A single search asks for them exactly once and so pays the full build. + + Built on the host in float64 from the cached eigendecomposition and moved + once: the ``d^l`` entries are bounded in [-1, 1], so storing them at the + working precision loses nothing structural, and it halves the memo. The build is the dominant cost of :func:`wigner_contraction_per_beta`: it is a batched ``(n_beta, sz, sz) @ (sz, sz)`` product per l, against the contraction's elementwise ``(n_beta, sz, sz)``. See ``_WIGNER_D_CACHE`` for the footprint that buys. """ + # `canonical_device` fills in the default index: torch.device('cuda') and + # torch.device('cuda:0') name one physical device but stringify differently, + # and this memo holds exactly ONE entry -- so a mixed spelling would not add + # an entry, it would clear and rebuild the whole table on every call. key = ( int(L), - str(device), + str(canonical_device(device)), + dtype, tuple(betas.detach().to(torch.float64).cpu().tolist()), ) hit = _WIGNER_D_CACHE.get(key) if hit is not None: return hit - eig_table = _wigner_eig_table(L, device) # cached (w_l, V_l) + eig_table = _wigner_eig_table(L) # host, cached + betas_host = betas.detach().to(torch.float64).cpu() blocks = [] for l in range(1, L): w, V = eig_table[l - 1] # data-independent - phase = torch.exp(-1j * betas.unsqueeze(1) * w.unsqueeze(0)) # (n_beta, sz) + phase = torch.exp(-1j * betas_host.unsqueeze(1) * w.unsqueeze(0)) # (n_beta, sz) VP = V.unsqueeze(0) * phase.unsqueeze(1) # (n_beta, sz, sz) = (k,m,a) - blocks.append((VP @ V.conj().transpose(-1, -2)).real) # (n_beta, sz, sz) + blocks.append( + (VP @ V.conj().transpose(-1, -2)).real.to(device=device, dtype=dtype) + ) _WIGNER_D_CACHE.clear() # one entry only; see the footprint note _WIGNER_D_CACHE[key] = blocks @@ -122,20 +142,23 @@ def wigner_contraction_per_beta( L = xi_lmn.shape[0] dim = 2 * L - 1 device = xi_lmn.device - betas = betas.to(torch.float64) n_beta = betas.shape[0] - xi = xi_lmn.to(torch.complex128) + # Follow the input rather than forcing double: `xi` carries the expansion's + # working precision, so widening here would buy nothing and cost a 2x + # complex buffer in this stage and in the FFT it feeds. + xi = xi_lmn + real_dtype = torch.float64 if xi.dtype == torch.complex128 else torch.float32 # Per-l loop over the small-d blocks, which come from the J_y # eigendecomposition (small_d_stable's method, stable to any l). Contract # each into S in turn: the full (n_beta, L, 2L-1, 2L-1) table is never # materialised as one array, nor is a 4-D einsum intermediate. - S = torch.zeros((n_beta, dim, dim), dtype=torch.complex128, device=device) + S = torch.zeros((n_beta, dim, dim), dtype=xi.dtype, device=device) c = L - 1 S[:, c, c] += xi[0, c, c] # l=0: d^0 = 1 - blocks = _wigner_d_blocks(L, betas, device) + blocks = _wigner_d_blocks(L, betas, device, real_dtype) for l in range(1, L): d_l = blocks[l - 1] # (n_beta, sz, sz) lo, hi = c - l, c + l + 1 - S[:, lo:hi, lo:hi] += xi[l, lo:hi, lo:hi].unsqueeze(0) * d_l.to(torch.complex128) + S[:, lo:hi, lo:hi] += xi[l, lo:hi, lo:hi].unsqueeze(0) * d_l return S diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index b0ebd19c..7efa2ab8 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -44,9 +44,9 @@ _StageTimer, _external_rwork, _prepare_frf_inputs, - _rodrigues, ) from .frf.rotation_utils import ( + axis_angle_to_matrix, edmonds_euler_from_rotation_matrix, rotation_matrix_from_edmonds_euler, ) @@ -852,7 +852,7 @@ def _dense_rotation_refine(self, refined): ) wx, wy, wz = torch.meshgrid(coords_r, coords_r, coords_r, indexing="ij") omegas = torch.stack([wx.flatten(), wy.flatten(), wz.flatten()], dim=-1) - R_perturbs = _rodrigues(omegas) + R_perturbs = axis_angle_to_matrix(omegas) R_cand_full = R_perturbs @ R_accumulated cand_peaks = [] for R_c in R_cand_full: diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index c8c06e52..6c188cb5 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -157,12 +157,20 @@ def fit_anisotropy( d_min: float, d_max: float, n_shells: int = N_WILSON_SHELLS, - device: Optional[torch.device] = None, ) -> torch.Tensor: """Fit the overall anisotropy tensor and project it onto the point group. - Returns ``U`` in Angstrom squared, in the convention - ``F_corrected = F * exp(+pi^2 s.U.s)``. + Returns ``U`` in Angstrom squared as a **host** float64 ``(3, 3)``, in the + convention ``F_corrected = F * exp(+pi^2 s.U.s)``. + :func:`~torchref.experimental.alignment.sh.apply_overall_anisotropy` moves + and casts it to wherever the amplitudes are. + + Deliberately host-side and in double precision. It is a seven-parameter + Gauss-Newton fit over the ``[d_max, d_min]`` window -- of order 1e4 + reflections, once per search -- so the cost of doing it here is not + measurable, while a broken anisotropy fit is worth hundreds of ranks on + high-symmetry cases. Keeping it off the accelerator also means the engine + needs no float64 there. The projection matters: an unconstrained six-component fit can return a tensor the lattice forbids, and applying it then modulates the observations @@ -172,9 +180,9 @@ def fit_anisotropy( """ from .sh import hkl_symops_to_cartesian, symmetrize_anisotropy - device = device or data.hkl.device - rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) - s_vec_all = data.hkl.to(torch.float64) @ rec_basis + rec_basis = data.cell.reciprocal_basis_matrix.detach().cpu().to(torch.float64) + hkl = data.hkl.detach().cpu().to(torch.float64) + s_vec_all = hkl @ rec_basis s_mag_all = s_vec_all.norm(dim=-1) keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min) if int(keep.sum()) < n_shells * 5: @@ -182,11 +190,11 @@ def fit_anisotropy( f"Only {int(keep.sum())} reflections in [{d_min}, {d_max}] A, too " f"few for {n_shells} shells." ) - F_obs = data.F.to(torch.float64).abs()[keep].to(device) - s_vec = s_vec_all[keep].to(device) - s_mag = s_mag_all[keep].to(device) + F_obs = data.F.detach().cpu().to(torch.float64).abs()[keep] + s_vec = s_vec_all[keep] + s_mag = s_mag_all[keep] centric = ( - data.centric[keep].to(torch.bool).to(device) + data.centric.detach().cpu()[keep].to(torch.bool) if hasattr(data, "centric") else torch.zeros_like(F_obs, dtype=torch.bool) ) @@ -197,8 +205,7 @@ def fit_anisotropy( F_obs, s_vec, shell_idx, centric, P=n_shells, min_count=20, ) sym_cart = hkl_symops_to_cartesian( - data.spacegroup.matrices.to(torch.float64).to(device), - rec_basis.to(device), + data.spacegroup.matrices.detach().cpu().to(torch.float64), rec_basis, ) return symmetrize_anisotropy(U, sym_cart) @@ -211,6 +218,7 @@ def search_peaks( U_aniso: torch.Tensor, n_peaks: int, verbose: int = 0, + device: Optional[torch.device] = None, ) -> Tuple[List["RotationPeak"], int, float]: """Run the rotation function, returning the engine's own peak list. @@ -220,11 +228,16 @@ def search_peaks( already fitted ``U_aniso`` for its rescore stage; :func:`rotation_search` is the entry point for everything else. """ + from ...utils import resolve_device from .frf.api import FastRotationFunction, phaser_lmax_resolution from .frf.dense_calc import dense_calc_via_box from .frf.preprocessing import fit_relative_wilson_b - device = model.xyz().device + # One device for both inputs, rather than whichever one this function + # happened to read first: `resolve_device` moves them into agreement (with a + # warning) and falls back to the configured default. Data first, matching + # the rest of the codebase. + device = resolve_device(data, model, device=device) with torch.no_grad(): rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) hkl_all = data.hkl.to(device) @@ -234,6 +247,11 @@ def search_peaks( # Take the observations at the full data resolution: the bandwidth # coupling below coarsens the limit to whatever the harmonics can # represent, so pre-restricting here would only lose the terms it keeps. + # + # The high-resolution half of the mask below is a no-op by construction + # (`s_mag <= 1/(1/max(s_mag))`, measured to drop nothing) but the + # LOW-resolution half is live: 3K7M carries two reflections beyond 100 A + # that it removes. Do not fold the pair away as redundant. d_min_data = float(1.0 / s_mag_all.max().item()) d_max = float(LOW_RESOLUTION_CUTOFF_A) keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min_data) @@ -264,9 +282,21 @@ def search_peaks( # one orbit. sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) n_ops = int(sg_mats.shape[0]) - hkl_keep = hkl_all.to(torch.float64)[keep] - hkl_unrolled = torch.einsum( - "kji,nj->kni", sg_mats, hkl_keep).reshape(-1, 3) + # `apply_to_hkl` is the package's one implementation of this contraction + # and its docstring carries the convention. It returns (N, 3, n_ops); the + # permute restores the op-major flattening the accumulations downstream + # were measured with -- a different row order changes the summation order + # in the later index_add_/unique and the last bits with it. Symops are + # 0/+-1 and Miller indices are small, so the products are exact at the + # space group's own dtype and the cast below loses nothing. + # `apply_to_hkl` returns on the SPACE GROUP's device, not the caller's, + # so the move back is load-bearing whenever the two differ. + hkl_unrolled = ( + data.spacegroup.apply_to_hkl(hkl_all[keep]) + .permute(2, 0, 1) + .reshape(-1, 3) + .to(device=device, dtype=rec_basis.dtype) + ) s_obs = hkl_unrolled @ rec_basis F_obs = F_obs.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() centric = centric.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() @@ -302,22 +332,11 @@ def search_peaks( engine = FastRotationFunction( s_obs, F_obs, centric, sg_mats, - d_min=d_min_data, d_max=d_max, + L=L, d_min=d_min, d_max=d_max, delta_vrms_A=float(model_error_A), n_wilson_shells=N_WILSON_SHELLS, sig_F_obs=sigF, grid_sampling_deg=GRID_SAMPLING_DEG, - model_radius_A=model_radius_A, - auto_lmax=True, - lmax_cap=LMAX_CAP, - # The angular half of the expansion runs in single precision on - # every device. It is the runtime bottleneck, it is memory-bound, and - # single precision is this codebase's kernel dtype -- a float64-only - # path would make the fused CPU kernel unreachable. The radial Bessel - # recurrence keeps its float64 internals, where the downward - # recurrence's cancellation needs them, and the cross-chunk - # accumulator stays at full precision. - compute_dtype=torch.complex64, ) _arf, peaks = engine.score_model( s_calc, F_calc, n_peaks=n_peaks, @@ -362,6 +381,7 @@ def rotation_search( *, n_peaks: int = 500, verbose: int = 0, + device: Optional[torch.device] = None, ) -> RotationSolutions: """Find the orientations of ``model`` consistent with ``data``. @@ -386,6 +406,10 @@ def rotation_search( bounds the answer, not the search. verbose : int, optional Progress reporting. Default 0, silent. + device : torch.device, optional + Where to run. Default ``None`` takes ``data``'s device, moving ``model`` + to match; an explicit value moves both. With neither carrying one, the + configured default applies. Returns ------- @@ -402,11 +426,9 @@ def rotation_search( if not model.initialized: raise RuntimeError("model has no coordinates; load a PDB first.") d_max_fit, d_min_fit = ANISO_FIT_WINDOW_A - U_aniso = fit_anisotropy( - data, d_min=d_min_fit, d_max=d_max_fit, device=model.xyz().device, - ) + U_aniso = fit_anisotropy(data, d_min=d_min_fit, d_max=d_max_fit) peaks, lmax, d_min = search_peaks( model, data, model_error_A, - U_aniso=U_aniso, n_peaks=n_peaks, verbose=verbose, + U_aniso=U_aniso, n_peaks=n_peaks, verbose=verbose, device=device, ) return _solutions(peaks, lmax, d_min, model_error_A) From 00d65480a664ddc14cdc75755b00a5b496962cb0 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 24 Aug 2026 02:45:49 +0200 Subject: [PATCH 041/250] Run the FRF's observed-side chain once per unique reflection The French-Wilson posterior, the LERF1 build, the shell variance reweight and the relative Wilson-B fit are per-reflection functions of (F, sigma_F, |s|, centric). All four inputs are symmetry-invariant, so the value is the same for every member of an orbit and computing it n_ops times is pure repetition -- measured at 27.4x for French-Wilson alone on 3K7M (8686 ms unrolled against 317 ms on the unique set). Only the GEOMETRY needs unrolling now: the engine takes one row per unique reflection plus an asu_idx and broadcasts after the chain. The op-major flattening makes that map arange(N) tiled n_ops times. Two changes are load-bearing together, not separately: - ONE shell assignment, shared. Moving the chain to the unique set on its own would have reintroduced a disagreement it had been papering over: french_wilson_preprocess derived equal-count edges in numpy from np.linspace(0, N-1, P+1).round() while apply_shell_variance_weights derived them in torch, and the two pick a different quantile rank for the same distribution at a different N. Measured on 3K7M: 7 of 55078 reflections land in a different shell depending on which consumer asks. Both now accept a shell_idx and the engine assigns once. - Mask BEFORE the unroll, at the bandwidth-coupled d_min. |s| is symmetry-invariant so the two commute, and this way the discarded high-resolution tail is not first replicated n_ops times through a matmul and three contiguous copies. Getting the second one wrong is expensive and worth recording: an intermediate version deleted the engine's obs-side resolution mask on the strength of "the caller has already masked", but the caller was masking at the DATA's d_min and the engine's mask was the only place the coarsened d_min_eff cut was applied. That fed the expansion ~4x the reflections, all finer than L can represent -- exactly the aliasing phaser_lmax_resolution exists to prevent. 3K7M went from rank 18 to rank 238, out of the top 20, at twice the runtime and +1 GB. All 526 unit tests passed throughout; only the end-to-end truth rank saw it, because the synthetic tests pass explicit limits and take the non-asu_idx path. Acceptance panel, 10 benchmark structures x 10 seeded trials, run in both worktrees so the comparison is paired per seed: truth in top 20 98/100 -> 98/100, and unchanged for every structure binding cases 1AK5 9/10 -> 9/10, 3K7M 9/10 -> 9/10 paired rank deltas median 0 everywhere; 89 of 100 cells identical, 7 better, 4 worse panel wall time 600 s -> 323 s (-46%) Per-peak scores move by 2.4e-3 relative (p50, 3K7M cap64) with 485 of 500 peaks recovered as the same set, and the top peak on that cell shifts to another member of its own symmetry orbit at the same height. That is a larger perturbation than the one rejected in 411ce26b, deliberately: there it was precision loss buying nothing, here it is one binning replacing two that disagreed, and it buys 46%. FastRotationFunction keeps a path where every array is per-row of s_obs and the window is applied internally, which is what direct callers and the synthetic tests use. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/panel_ranks.py | 69 ++++++++++++++++ alignment_lab/analysis/staged_gate.sh | 76 ++++++++++++++++++ alignment_lab/analysis/staged_panel.sh | 29 +++++++ docs/changelog.rst | 4 + torchref/experimental/alignment/frf/api.py | 79 ++++++++++++++----- .../alignment/frf/french_wilson.py | 35 +++++--- .../alignment/frf/preprocessing.py | 10 ++- .../experimental/alignment/rotation_search.py | 74 ++++++++++++----- 8 files changed, 326 insertions(+), 50 deletions(-) create mode 100644 alignment_lab/analysis/panel_ranks.py create mode 100644 alignment_lab/analysis/staged_gate.sh create mode 100644 alignment_lab/analysis/staged_panel.sh diff --git a/alignment_lab/analysis/panel_ranks.py b/alignment_lab/analysis/panel_ranks.py new file mode 100644 index 00000000..aa17d56a --- /dev/null +++ b/alignment_lab/analysis/panel_ranks.py @@ -0,0 +1,69 @@ +"""Truth rank over one benchmark structure at seeded orientations. + +The acceptance criterion for a change to the rotation function is not the median +rank -- it is whether truth lands inside the candidate window the placement +search carries forward, on nearly every trial, for *every* structure. Seed-to- +seed spread at ``lmax_cap = 64`` is +-4 to 6 ranks, so a bare median hides the +cases that decide it. + +Emits one ``ROW`` line per trial so the caller can pair the same seed across two +worktrees. Deliberately reuses the lab's ``seed_for`` / ``rotated_case`` / +``orbit_rank``, which carry the seed contract and the orbit conventions the +earlier sweeps were measured with -- a private reimplementation of any of those +would make the comparison meaningless. +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, FRFConfig, orbit_rank, rotated_case, # noqa: E402 + run_frf, seed_for) + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) + ap.add_argument("--trials", type=int, default=10) + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--thr-deg", type=float, default=5.0) + ap.add_argument("--tag", default="?", help="which tree this run came from") + args = ap.parse_args() + + cfg = FRFConfig(n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) + for trial in range(args.trials): + seed = seed_for(args.pdb, trial) + model, data, R_true = rotated_case(args.pdb, seed) + t0 = time.time() + res = run_frf(model, data, cfg, capture_arf=False, verbose=0) + seconds = time.time() - t0 + rank, ang = orbit_rank( + res.peaks, R_true, + data.spacegroup.matrices.to(torch.float64).cpu(), + reciprocal_basis=data.cell.reciprocal_basis_matrix.to( + torch.float64).cpu(), + side="left", frame="cart", thr_deg=args.thr_deg, + ) + # orbit_rank returns -1 for "no peak within thr_deg". A miss must not + # sort as a good rank, so for comparison it counts as worse than the + # worst hit -- the peak-list length. + rank_cmp = rank if rank >= 0 else args.n_peaks + print(f"ROW {args.tag} {args.pdb} trial={trial} seed={seed} " + f"rank={rank} rank_cmp={rank_cmp} found={int(rank >= 0)} " + f"top20={int(0 <= rank < 20)} " + f"angle={'' if ang is None else round(float(ang), 3)} " + f"seconds={seconds:.2f} sg={data.spacegroup.hm}", flush=True) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/alignment_lab/analysis/staged_gate.sh b/alignment_lab/analysis/staged_gate.sh new file mode 100644 index 00000000..186efc26 --- /dev/null +++ b/alignment_lab/analysis/staged_gate.sh @@ -0,0 +1,76 @@ +#!/bin/bash +# Stage D gate: the obs chain now runs on the unique set and only the geometry +# unrolls, with one shared shell assignment. This CHANGES numbers, so quantify +# how much and confirm truth is still where the pipeline can reach it. Also +# measure what it bought, paired and interleaved on one node. +#SBATCH --job-name=staged2 +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=48G +#SBATCH --exclusive +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +OLD=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/_stagea_baseline +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +OUT=$NEW/alignment_lab/slurm +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname) OLD=$(cd $OLD && git rev-parse --short HEAD) NEW=working tree" + +cd "$NEW"; export PYTHONPATH="$NEW" +export TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 +LOG=$OUT/staged_tests_$SLURM_JOB_ID.log +"$PY" -m pytest tests/unit/alignment tests/unit/frf_separate tests/unit/model \ + tests/unit/test_imports_smoke.py -q > "$LOG" 2>&1 +rc=$? +echo "=== TESTS rc=$rc ===" +tail -14 "$LOG" + +echo "=== how far did the peak lists move (single-threaded) ===" +for pdb in 3K7M 1DAW; do + for cap in 64 100; do + for tree in OLD NEW; do + eval "root=\$$tree" + cd "$root" + PYTHONPATH="$root" "$PY" -u "$NEW/alignment_lab/analysis/frf_fingerprint.py" \ + --pdb "$pdb" --lmax-cap "$cap" > "$OUT/d_${tree}_${pdb}_${cap}.txt" 2>/dev/null \ + || echo "RUN FAILED $tree $pdb $cap" + done + "$PY" - "$OUT/d_OLD_${pdb}_${cap}.txt" "$OUT/d_NEW_${pdb}_${cap}.txt" "$pdb" "$cap" <<'PYEOF' +import sys +rp, np_, pdb, cap = sys.argv[1:5] +def load(p): + return [tuple(map(float, l.split()[2:7])) for l in open(p) if l.startswith("FP ")] +a, b = load(rp), load(np_) +if not a or not b: + print(f" {pdb} cap{cap}: EMPTY {len(a)} vs {len(b)}"); sys.exit() +n = min(len(a), len(b)) +slot = sum(1 for i in range(n) if a[i][:3] == b[i][:3]) +# Is the same peak SET recovered, regardless of order? +sa, sb = {r[:3] for r in a}, {r[:3] for r in b} +rel = sorted(abs(b[i][3]-a[i][3])/max(abs(a[i][3]),1e-30) for i in range(n)) +print(f" {pdb} cap{cap}: {slot}/{n} slots identical | set overlap " + f"{len(sa & sb)}/{len(sa)} | |dscore|/score p50={rel[n//2]:.2e} " + f"p99={rel[int(0.99*n)]:.2e}") +print(f" top-1 old {tuple(round(v,6) for v in a[0][:3])} z={a[0][4]:.4f}" + f" new {tuple(round(v,6) for v in b[0][:3])} z={b[0][4]:.4f}") +PYEOF + done +done + +echo "=== paired timing, 4 threads, 3 rounds ===" +export TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +for round in 1 2 3; do + for pdb in 3K7M 1DAW; do + for tree in OLD NEW; do + eval "root=\$$tree"; cd "$root" + line=$(PYTHONPATH="$root" "$PY" -u -m alignment_lab.diagnostics.frf_benchmark \ + --pdb "$pdb" --arms cap64 --trials 2 2>/dev/null | grep -E "^ *cap64" | tail -1) + echo "round$round $pdb $tree $line" + done + done +done +echo "staged_tests_rc=$rc" diff --git a/alignment_lab/analysis/staged_panel.sh b/alignment_lab/analysis/staged_panel.sh new file mode 100644 index 00000000..19b81583 --- /dev/null +++ b/alignment_lab/analysis/staged_panel.sh @@ -0,0 +1,29 @@ +#!/bin/bash +# Stage D acceptance panel: 10 benchmark structures x 10 seeded trials, run in +# BOTH worktrees so the comparison is paired per seed. One array task per +# structure; OLD and NEW run back to back inside the task on the same node. +#SBATCH --job-name=d_panel +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-9 +set -uo pipefail +NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +OLD=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/_stagea_baseline +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} +export TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname) pdb=$PDB" +for tree in OLD NEW; do + eval "root=\$$tree" + cd "$root" + PYTHONPATH="$root" "$PY" -u "$NEW/alignment_lab/analysis/panel_ranks.py" \ + --pdb "$PDB" --trials 10 --tag "$tree" 2>/dev/null | grep '^ROW ' + echo "tree=$tree rc=$?" +done diff --git a/docs/changelog.rst b/docs/changelog.rst index 760ebf54..edbbe5ed 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -21,6 +21,10 @@ Unreleased - Removed the rotation function's duplicate Euler, Rodrigues and reciprocal-symmetry helpers in favour of the shared primitives - Removed the rotation function's redundant calc-side resolution mask and its second bandwidth/resolution coupling call - ``bessel_sh_expand`` lost its unread ``chunk_size`` argument and ``french_wilson_preprocess`` its unread ``sqrt_mean_F2`` output +- The rotation function's observed-side chain (French-Wilson, LERF1, shell variance weights, relative Wilson B) now runs once per unique reflection instead of once per symmetry copy; only the geometry is unrolled +- The rotation function assigns resolution shells once and shares them, instead of the Wilson normalisation and the variance reweight deriving edges that disagreed at the shell boundaries +- The rotation function masks observations to the bandwidth-coupled resolution before the symmetry unroll rather than after +- The rotation function warns on Bijvoet-unmerged data, whose shared canonical index would weight those reflections twice Version 0.6.4 ---------- diff --git a/torchref/experimental/alignment/frf/api.py b/torchref/experimental/alignment/frf/api.py index eb1b6296..82c27e53 100644 --- a/torchref/experimental/alignment/frf/api.py +++ b/torchref/experimental/alignment/frf/api.py @@ -137,6 +137,8 @@ def __init__( n_wilson_shells: int = 20, sig_F_obs: Optional[torch.Tensor] = None, grid_sampling_deg: float = 2.0, + asu_idx: Optional[torch.Tensor] = None, + s_mag_asu: Optional[torch.Tensor] = None, ): self.device = s_obs.device @@ -152,18 +154,41 @@ def __init__( self.n_wilson_shells = n_wilson_shells self.grid_sampling_deg = grid_sampling_deg - # 1. Resolution mask on obs. - extras = (F_obs, centric_obs) - if sig_F_obs is not None: - extras = extras + (sig_F_obs,) - s_obs, extras, smag_obs = _resolution_mask(s_obs, extras, d_min, d_max) - F_obs, centric_obs = extras[0], extras[1] - if sig_F_obs is not None: - sig_F_obs = extras[2] - - if s_obs.shape[0] < n_wilson_shells * 5: + # 1. Resolution window. + # + # With `asu_idx` the caller has already masked, and `F_obs` / `sig_F_obs` + # / `centric_obs` are ONE ROW PER UNIQUE REFLECTION while `s_obs` carries + # the full symmetry-unrolled geometry. Everything from here to the + # expansion is a per-reflection function of (F, sigma_F, |s|, centric), + # all four of which are symmetry-invariant, so the chain runs on the + # unique set and is broadcast at the end. That is exact -- not an + # approximation -- and it is the difference between doing the + # French-Wilson posterior once and doing it n_ops times. + # + # Without `asu_idx` every array is per-row of `s_obs` and the window is + # applied here, which is the path direct callers and the synthetic tests + # take. + if asu_idx is None: + extras = (F_obs, centric_obs) + if sig_F_obs is not None: + extras = extras + (sig_F_obs,) + s_obs, extras, smag_src = _resolution_mask(s_obs, extras, d_min, d_max) + F_obs, centric_obs = extras[0], extras[1] + if sig_F_obs is not None: + sig_F_obs = extras[2] + else: + if s_mag_asu is None: + raise ValueError("asu_idx requires s_mag_asu (|s| per unique row)") + if int(asu_idx.shape[0]) != int(s_obs.shape[0]): + raise ValueError( + f"asu_idx has {int(asu_idx.shape[0])} entries for " + f"{int(s_obs.shape[0])} unrolled reflections" + ) + smag_src = s_mag_asu + + if F_obs.shape[0] < n_wilson_shells * 5: raise ValueError( - f"Too few obs reflections ({s_obs.shape[0]}) for " + f"Too few obs reflections ({F_obs.shape[0]}) for " f"{n_wilson_shells} Wilson shells in [{d_min}, {d_max}] Å." ) @@ -180,17 +205,31 @@ def __init__( lmax_even = lmax if lmax % 2 == 0 else lmax - 1 self.bessel_h_scale = float(lmax_even) * float(d_min) - # 3. Wilson normalisation. With sigmas, through the French-Wilson - # posterior, which does its own per-shell normalisation and handles - # the axial reflections; without them, plain per-shell Wilson. + # 3. ONE shell assignment, shared by everything below. + # + # The French-Wilson posterior, the LERF1 build and the variance reweight + # all normalise per resolution shell, and each used to derive its own + # equal-count edges from the same |s| -- one in numpy, one in torch, with + # different quantile-rank rounding. That put a handful of boundary + # reflections in different shells depending on which consumer asked, + # which is a difference of ~2e-4 relative on their normalisation for no + # reason. Assign once, pass it down. + from ..sh import assign_shells, equal_count_shell_edges + + shell_edges, _ = equal_count_shell_edges(smag_src, n_wilson_shells) + obs_shell_idx = assign_shells(smag_src, shell_edges) + + # 3b. Wilson normalisation. With sigmas, through the French-Wilson + # posterior, which handles the axial reflections; without them, plain + # per-shell Wilson. if sig_F_obs is not None: fw = french_wilson_preprocess( - F_obs, sig_F_obs, smag_obs, centric_obs, - n_wilson_shells=n_wilson_shells, + F_obs, sig_F_obs, smag_src, centric_obs, + n_wilson_shells=n_wilson_shells, shell_idx=obs_shell_idx, ) eEobs, dfac = fw["eEobs"], fw["DFAC"] else: - eEobs, _ = wilson_normalise(F_obs, smag_obs, n_wilson_shells) + eEobs, _ = wilson_normalise(F_obs, smag_src, n_wilson_shells) dfac = torch.ones_like(eEobs) # 4. LERF1 obs intensity, and the per-shell variance reweight. @@ -198,8 +237,12 @@ def __init__( eEobs, centric_obs, dfac=dfac, use_centric_weight=True, ) intensity_obs = apply_shell_variance_weights( - intensity_obs, smag_obs, n_var_shells=n_wilson_shells, + intensity_obs, smag_src, n_var_shells=n_wilson_shells, + shell_idx=obs_shell_idx, ) + if asu_idx is not None: + # One value per unique reflection -> one per unrolled reflection. + intensity_obs = intensity_obs[asu_idx] # 5. ZSYMM m-filter on the obs SH coefficients. The calc side is never # filtered -- see score_model. diff --git a/torchref/experimental/alignment/frf/french_wilson.py b/torchref/experimental/alignment/frf/french_wilson.py index ad5a5f7d..bc4301df 100644 --- a/torchref/experimental/alignment/frf/french_wilson.py +++ b/torchref/experimental/alignment/frf/french_wilson.py @@ -432,12 +432,18 @@ def french_wilson_preprocess( centric: torch.Tensor, *, n_wilson_shells: int = 20, + shell_idx: "torch.Tensor | None" = None, ) -> dict: """Phaser-style preprocessing from raw ``(F, σF, centric)`` to ``(eEobs, DFAC)``. Implements the chain: - 1. equal-count Wilson shells over ``s_mag`` + 1. equal-count Wilson shells over ``s_mag`` -- or ``shell_idx``, when the + caller has already assigned them. Pass it: this routine's own quantile + edges (``np.linspace(0, N-1, P+1).round()``) pick a different rank than + ``equal_count_shell_edges`` does for the same distribution at a different + N, so two independently-binned consumers disagree about a handful of + reflections at the shell boundaries (measured: 7 of 55078 on 3K7M). 2. per-shell ``_p`` (Phaser's ``SIGMAN.BINS``) 3. per-reflection normalised intensity ``eosq = F² / `` and σ ``sigesq = σI / ≈ 2·F·σF / `` @@ -458,14 +464,25 @@ def french_wilson_preprocess( s_np = s_mag.detach().to("cpu").to(torch.float64).numpy() cen_np = centric.detach().to("cpu").bool().numpy() - sorted_idx = np.argsort(s_np) - edges_idx = np.linspace(0, len(s_np) - 1, n_wilson_shells + 1).round().astype(np.int64) - s_edges = s_np[sorted_idx][edges_idx] - s_edges[0] -= 1e-6 - s_edges[-1] += 1e-6 - shell_idx = np.clip( - np.searchsorted(s_edges, s_np, side="right") - 1, 0, n_wilson_shells - 1, - ) + if shell_idx is None: + sorted_idx = np.argsort(s_np) + edges_idx = np.linspace( + 0, len(s_np) - 1, n_wilson_shells + 1).round().astype(np.int64) + s_edges = s_np[sorted_idx][edges_idx] + s_edges[0] -= 1e-6 + s_edges[-1] += 1e-6 + shell_idx = np.clip( + np.searchsorted(s_edges, s_np, side="right") - 1, + 0, n_wilson_shells - 1, + ) + else: + # Out-of-range rows come back as -1 from `assign_shells`; clamp them into + # the end shells rather than dropping them, which is what this routine's + # own edge nudge did. + shell_idx = np.clip( + shell_idx.detach().to("cpu").to(torch.int64).numpy(), + 0, n_wilson_shells - 1, + ) F2 = F_np * F_np mean_F2 = np.zeros(n_wilson_shells, dtype=np.float64) counts = np.zeros(n_wilson_shells, dtype=np.int64) diff --git a/torchref/experimental/alignment/frf/preprocessing.py b/torchref/experimental/alignment/frf/preprocessing.py index 31fb0d52..d27ce2db 100644 --- a/torchref/experimental/alignment/frf/preprocessing.py +++ b/torchref/experimental/alignment/frf/preprocessing.py @@ -263,15 +263,21 @@ def apply_shell_variance_weights( intensity: torch.Tensor, s_mag: torch.Tensor, n_var_shells: int = 20, + shell_idx: Optional[torch.Tensor] = None, ) -> torch.Tensor: """Per-shell empirical variance reweight. Downweights shells whose observed Patterson intensity is dominated by noise. Mean-normalised so total scale doesn't shift. Closest Phaser analog is per-shell BINS + ``best(r)`` in ``Ensemble.cc``. + + ``shell_idx`` reuses an assignment the caller already made. Worth passing: + binning here independently of the Wilson normalisation puts the two on + edges that disagree for the reflections sitting on a boundary. """ - edges, _ = equal_count_shell_edges(s_mag, n_var_shells) - shell_idx = assign_shells(s_mag, edges) + if shell_idx is None: + edges, _ = equal_count_shell_edges(s_mag, n_var_shells) + shell_idx = assign_shells(s_mag, edges) valid = shell_idx >= 0 var_p = compute_patterson_shell_variance( intensity[valid].to(torch.float64), diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index 6c188cb5..41bf0cd1 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -22,6 +22,7 @@ from __future__ import annotations +import warnings from dataclasses import dataclass from typing import TYPE_CHECKING, List, Optional, Tuple @@ -244,21 +245,31 @@ def search_peaks( s_vec_all = hkl_all.to(torch.float64) @ rec_basis s_mag_all = s_vec_all.norm(dim=-1) - # Take the observations at the full data resolution: the bandwidth - # coupling below coarsens the limit to whatever the harmonics can - # represent, so pre-restricting here would only lose the terms it keeps. - # - # The high-resolution half of the mask below is a no-op by construction - # (`s_mag <= 1/(1/max(s_mag))`, measured to drop nothing) but the - # LOW-resolution half is live: 3K7M carries two reflections beyond 100 A - # that it removes. Do not fold the pair away as redundant. d_min_data = float(1.0 / s_mag_all.max().item()) d_max = float(LOW_RESOLUTION_CUTOFF_A) - keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min_data) - s_obs = s_vec_all[keep] + # Couple the bandwidth to the resolution FIRST, because the mask below is + # the only place the coarsened limit is applied -- the engine takes the + # window as given. Data finer than L can represent contributes aliasing + # rather than signal and buries the symmetry-diluted true peak, so + # dropping this cut is not a small error: on 3K7M it moved truth from + # rank 18 to rank 238 and doubled the runtime. + model_radius_A = float( + (model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item() + ) + L, d_min = phaser_lmax_resolution(model_radius_A, d_min_data, LMAX_CAP) + + # Masking before the unroll rather than after: |s| is symmetry-invariant + # (symmetry operations are isometries), so the two commute -- and this way + # the discarded high-resolution tail is not first replicated n_ops times. + # The low-resolution half is live, not decorative: 3K7M carries two + # reflections beyond 100 A that it removes. + keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min) + + s_asu = s_vec_all[keep] + s_mag_asu = s_mag_all[keep] F_obs = apply_overall_anisotropy( - data.F.to(torch.float64).abs().to(device)[keep], s_obs, U_aniso, + data.F.to(torch.float64).abs().to(device)[keep], s_asu, U_aniso, ) sigF = ( data.F_sigma.to(torch.float64).to(device)[keep] @@ -270,6 +281,17 @@ def search_peaks( if hasattr(data, "centric") else torch.zeros_like(F_obs, dtype=torch.bool) ) + # Unmerged Bijvoet data puts two rows on one canonical index, so the + # unroll below would weight those reflections twice. Detectable, so say + # so rather than quietly double-counting. + if getattr(data, "friedel_merged", True) is False: + warnings.warn( + "data are Bijvoet-unmerged: both members of a pair share a " + "canonical index, so the symmetry unroll weights those " + "reflections twice in the Patterson. Merge first for a clean " + "rotation function.", + RuntimeWarning, stacklevel=2, + ) # Expand the observations over the space group's rotations to fill # reciprocal space. |F(hS)| = |F(h)|, and the harmonics need the full @@ -298,18 +320,23 @@ def search_peaks( .to(device=device, dtype=rec_basis.dtype) ) s_obs = hkl_unrolled @ rec_basis - F_obs = F_obs.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() - centric = centric.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() - if sigF is not None: - sigF = sigF.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() + # Only the GEOMETRY is unrolled. The amplitudes, sigmas and centric + # flags stay one row per unique reflection and the engine broadcasts + # them after its per-reflection chain, which is symmetry-invariant. The + # op-major flattening above means unrolled row `k * N + n` came from + # unique row `n`, so the map is `arange(N)` tiled n_ops times. + n_unique = int(F_obs.shape[0]) + asu_idx = ( + torch.arange(n_unique, device=device) + .unsqueeze(0) + .expand(n_ops, -1) + .reshape(-1) + ) # The model's transform on a dense P1 grid rather than at the crystal's # own reflections: the crystal lattice is too sparse to determine the - # high-l harmonics for a large molecule. - model_radius_A = float( - (model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item() - ) - L, d_min = phaser_lmax_resolution(model_radius_A, d_min_data, LMAX_CAP) + # high-l harmonics for a large molecule. `L` / `d_min` came from the + # bandwidth coupling above, which the obs mask also used. s_calc, F_calc = dense_calc_via_box( model, d_max, d_min, pad=DENSE_CALC_PAD, verbose=verbose > 0, ) @@ -320,9 +347,12 @@ def search_peaks( # (EnsemblePDB.cc:793-851), so the radial fall-off does not by itself # discriminate between orientations. s_calc_mag = s_calc.norm(dim=-1) + # Fitted on the unique set. The unroll replicates every reflection + # exactly n_ops times, so the per-shell means are identical to the + # unrolled fit while the sort and the binning are n_ops times smaller. B_rel = fit_relative_wilson_b( F_obs.to(torch.float64), F_calc.to(torch.float64), - s_obs.norm(dim=-1).to(torch.float64), n_shells=N_WILSON_SHELLS, + s_mag_asu.to(torch.float64), n_shells=N_WILSON_SHELLS, s_mag_calc=s_calc_mag.to(torch.float64), ) if abs(B_rel) > 1e-6: @@ -337,6 +367,8 @@ def search_peaks( n_wilson_shells=N_WILSON_SHELLS, sig_F_obs=sigF, grid_sampling_deg=GRID_SAMPLING_DEG, + asu_idx=asu_idx, + s_mag_asu=s_mag_asu, ) _arf, peaks = engine.score_model( s_calc, F_calc, n_peaks=n_peaks, From fe53d373c3197fdef704619834e47d910116bf77 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 24 Aug 2026 03:06:01 +0200 Subject: [PATCH 042/250] Stop concatenating the antipodal copy in the FRF expansion The Patterson's centrosymmetry was encoded three times in `bessel_sh_expand`. Even-l-only and the m<0 mirror both SAVE work. The third -- concatenating `-s` onto both reflection sets -- COST work and was exactly redundant with them: only even l are computed, `Y_lm(-s_hat) = Y_lm(s_hat)` there, and the intensity, Bessel weight and Legendre factor are all unchanged under negation, so it doubled `c_nlm` exactly. Both sides doubled scales the rotation function by 4, which the z-score normalisation removes. Decided by knockout over the panel (10 structures x 10 seeded trials, paired per seed) rather than by the argument alone: ranks 98/100 identical, 1 better, 1 worse; every structure >= 9/10 in the top 20, with 1AK5 and 3K7M holding 9/10 top score exactly 0.2499974 .. 0.2500036 of the doubled value across all 100 cells -- the predicted factor of 4, to 5-6 decimals runtime -22.5% on 3K7M, -16.4% on 1DAW, warm and with the arm order rotated so no arm was systematically first Cumulative with the preceding ASU rework, paired per trial against 411ce26b with the warm-up trial excluded, the gain tracks n_ops as expected of a saving whose size is the orbit: 1AK5 (24 ops) -60.8% 3K7M (24) -50.5% 3A5V (16) -44.6% 2DQ6 (6) -35.6% 3GR5 (12) -32.1% 1DAW (4) -19.5% 3E98 (2) -15.9% 6G9X -13.3% 4BX9 (8) -11.3% `fit_relative_wilson_b` was knocked out in the same experiment and KEPT: rank- neutral (96/100 identical, 4 worse, 0 better, top score within 0.08%) but worth only -2.2% / +0.4% now that the ASU rework moved it onto ~12k unique rows. The pre-rework case for deleting it -- a sort, a binning and a regression over 1.3M rows -- no longer exists. Recorded so it is not re-litigated on that reasoning. Two consequences, neither of them silent: - Raw `RotationPeak.score` / `RotationSolutions.scores` are now a quarter of their previous values. z-scores are unchanged and no threshold reads the absolute score (`sigma_threshold` and the peak finder both work off mean/sigma), but `pipeline.py` writes `rotation_score` into result rows, so archived raw scores do not compare across this commit. - Phaser DOES include the mate, via cctbx's `conjugate_flag`, so our coefficients are now half of its. That matters only to the coefficient-level comparison in `frf_encode_compare.py`, whose correspondence factor goes from k=2 to k=1; its now-dead `friedel` parameter is gone too. `test_the_antipodal_copy_would_only_double_the_result` pins the exactness the removal rests on, so if odd l ever carries signal the equality breaks rather than the removal quietly becoming wrong. The direct-summation reference in the same file no longer duplicates its input either -- it had been mirroring the old behaviour, which is why it was the only thing that failed here. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/arms_timing.sh | 60 +++++++++ alignment_lab/analysis/friedel_gate.sh | 27 ++++ alignment_lab/analysis/panel_arms.py | 121 ++++++++++++++++++ alignment_lab/analysis/panel_arms.sh | 22 ++++ .../diagnostics/frf_encode_compare.py | 18 +-- docs/changelog.rst | 2 + .../frf_separate/test_bessel_sh_grouping.py | 29 ++++- torchref/experimental/alignment/frf/api.py | 4 +- .../experimental/alignment/frf/data_mr.py | 22 +++- 9 files changed, 286 insertions(+), 19 deletions(-) create mode 100644 alignment_lab/analysis/arms_timing.sh create mode 100644 alignment_lab/analysis/friedel_gate.sh create mode 100644 alignment_lab/analysis/panel_arms.py create mode 100644 alignment_lab/analysis/panel_arms.sh diff --git a/alignment_lab/analysis/arms_timing.sh b/alignment_lab/analysis/arms_timing.sh new file mode 100644 index 00000000..b233594f --- /dev/null +++ b/alignment_lab/analysis/arms_timing.sh @@ -0,0 +1,60 @@ +#!/bin/bash +# What the two knockouts actually SAVE. The panel run could not answer this: it +# ran production first in every trial, so production alone paid the process-level +# warm-up -- the fused C++ kernel build, the SO(3) sample list, the Wigner block +# memo -- and the later arms inherited all three warm. That is why it reported a +# ~10x "speedup" from deleting a 20-bin regression, which is not credible. +# +# Here: one throwaway run to warm every memo, then rounds with the arm order +# ROTATED so no arm is systematically first. +#SBATCH --job-name=arms_time +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=48G +#SBATCH --exclusive +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" +"$PY" -u - <<'PYEOF' +import statistics, sys, time +from pathlib import Path +sys.path.insert(0, str(Path("alignment_lab").resolve())) +import torch +torch.set_grad_enabled(False) +sys.path.insert(0, str(Path("alignment_lab/analysis").resolve())) +from lab import FRFConfig, rotated_case, run_frf, seed_for +from panel_arms import knocked_out + +ARMS = ["production", "no_brel", "no_friedel"] +cfg = FRFConfig(n_peaks=500, lmax_cap=64) + +for pdb in ("3K7M", "1DAW"): + model, data, _ = rotated_case(pdb, seed_for(pdb, 0)) + run_frf(model, data, cfg, capture_arf=False, verbose=0) # warm every memo + t = {a: [] for a in ARMS} + for r in range(4): + order = ARMS[r % len(ARMS):] + ARMS[:r % len(ARMS)] # rotate + for arm in order: + model, data, _ = rotated_case(pdb, seed_for(pdb, r)) + t0 = time.time() + with knocked_out(arm): + run_frf(model, data, cfg, capture_arf=False, verbose=0) + t[arm].append(time.time() - t0) + base = statistics.median(t["production"]) + print(f"--- {pdb} (warm, arm order rotated, 4 rounds) ---") + for arm in ARMS: + med = statistics.median(t[arm]) + pd = [t[arm][i] - t["production"][i] for i in range(len(t[arm]))] + print(f" {arm:12s} median {med:.3f}s paired d vs production " + f"median {statistics.median(pd):+.3f}s " + f"({100*statistics.median(pd)/base:+.1f}%) raw {[round(v,3) for v in t[arm]]}") +PYEOF +echo "rc=$?" diff --git a/alignment_lab/analysis/friedel_gate.sh b/alignment_lab/analysis/friedel_gate.sh new file mode 100644 index 00000000..ffa2f122 --- /dev/null +++ b/alignment_lab/analysis/friedel_gate.sh @@ -0,0 +1,27 @@ +#!/bin/bash +# Gate for removing the antipodal copy: the panel again (ranks must hold) plus a +# warm, order-rotated timing check that the 16-22% survives in production code +# rather than only under the knockout patch. +#SBATCH --job-name=fried +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-9 +set -uo pipefail +NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +OLD=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/_stagea_baseline +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} +export TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname) pdb=$PDB" +for tree in OLD NEW; do + eval "root=\$$tree"; cd "$root" + PYTHONPATH="$root" "$PY" -u "$NEW/alignment_lab/analysis/panel_ranks.py" \ + --pdb "$PDB" --trials 10 --tag "$tree" 2>/dev/null | grep '^ROW ' +done diff --git a/alignment_lab/analysis/panel_arms.py b/alignment_lab/analysis/panel_arms.py new file mode 100644 index 00000000..d46bc726 --- /dev/null +++ b/alignment_lab/analysis/panel_arms.py @@ -0,0 +1,121 @@ +"""Truth rank per seeded trial under one or more knocked-out engine stages. + +Two stages of the rotation function are suspected of earning nothing, each for a +different reason, and both are cheaper to decide by knocking them out than by +reasoning about them: + +``no_brel`` + ``fit_relative_wilson_b`` scales ``F_calc`` by ``exp(-B_rel s^2/4)`` and the + very next step, ``wilson_normalise``, divides each equal-count shell by its + own ``sqrt()`` -- which removes the shell-mean radial profile, + B_rel's included. Only the within-shell residual of a smooth exponential can + survive. If that is below the engine's own spread, the fit is a sort, a + binning and a regression for nothing. + +``no_friedel`` + ``enforce_friedel`` concatenates ``-s`` onto both reflection sets. Only even + ``l`` are ever computed and ``Y_lm(-s_hat) = Y_lm(s_hat)`` for even ``l``, + with the intensity duplicated verbatim, so ``c_nlm`` doubles *exactly*. Both + sides doubled means xi scales by 4, and the mean and standard deviation of + the rotation function scale with it, so z-scores and the ranking are + invariant. The prediction is therefore sharp: identical ranks, and the raw + score up by exactly 4. Reported so it can be checked rather than assumed. + +Emits one ``ROW`` line per (arm, trial), carrying the top score so the scaling +prediction is falsifiable from the output. +""" + +from __future__ import annotations + +import argparse +import contextlib +import functools +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, FRFConfig, orbit_rank, rotated_case, # noqa: E402 + run_frf, seed_for) + +ARMS = ("production", "no_brel", "no_friedel") + + +@contextlib.contextmanager +def knocked_out(arm: str): + """Disable one stage for the duration of a call, then restore it.""" + if arm == "production": + yield + return + if arm == "no_brel": + import torchref.experimental.alignment.frf.preprocessing as pp + original = pp.fit_relative_wilson_b + pp.fit_relative_wilson_b = lambda *a, **k: 0.0 + try: + yield + finally: + pp.fit_relative_wilson_b = original + return + if arm == "no_friedel": + import torchref.experimental.alignment.frf.api as api + original = api.bessel_sh_expand + + @functools.wraps(original) + def no_mate(*a, **k): + k["enforce_friedel"] = False + return original(*a, **k) + + api.bessel_sh_expand = no_mate + try: + yield + finally: + api.bessel_sh_expand = original + return + raise ValueError(f"unknown arm {arm!r}") + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) + ap.add_argument("--trials", type=int, default=10) + ap.add_argument("--arms", default=",".join(ARMS)) + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--thr-deg", type=float, default=5.0) + args = ap.parse_args() + + cfg = FRFConfig(n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) + arms = [a for a in args.arms.split(",") if a] + # Trial-major so the same seed's arms run back to back: any drift in machine + # state affects the arms together rather than one of them. + for trial in range(args.trials): + seed = seed_for(args.pdb, trial) + for arm in arms: + model, data, R_true = rotated_case(args.pdb, seed) + t0 = time.time() + with knocked_out(arm): + res = run_frf(model, data, cfg, capture_arf=False, verbose=0) + seconds = time.time() - t0 + rank, ang = orbit_rank( + res.peaks, R_true, + data.spacegroup.matrices.to(torch.float64).cpu(), + reciprocal_basis=data.cell.reciprocal_basis_matrix.to( + torch.float64).cpu(), + side="left", frame="cart", thr_deg=args.thr_deg, + ) + top = res.peaks[0] if res.peaks else None + print(f"ROW {arm} {args.pdb} trial={trial} seed={seed} " + f"rank={rank} rank_cmp={rank if rank >= 0 else args.n_peaks} " + f"top20={int(0 <= rank < 20)} " + f"top_score={'' if top is None else f'{top.score:.10g}'} " + f"top_sigma={'' if top is None else f'{top.sigma:.6g}'} " + f"seconds={seconds:.2f}", flush=True) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/alignment_lab/analysis/panel_arms.sh b/alignment_lab/analysis/panel_arms.sh new file mode 100644 index 00000000..c8a880d4 --- /dev/null +++ b/alignment_lab/analysis/panel_arms.sh @@ -0,0 +1,22 @@ +#!/bin/bash +# Decision gates for the two suspected-vestigial stages, over the full panel. +#SBATCH --job-name=arms +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-9 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname) pdb=$PDB" +"$PY" -u alignment_lab/analysis/panel_arms.py --pdb "$PDB" --trials 10 2>/dev/null | grep '^ROW ' +echo "rc=$?" diff --git a/alignment_lab/diagnostics/frf_encode_compare.py b/alignment_lab/diagnostics/frf_encode_compare.py index e1c90135..7ce64915 100644 --- a/alignment_lab/diagnostics/frf_encode_compare.py +++ b/alignment_lab/diagnostics/frf_encode_compare.py @@ -47,7 +47,9 @@ we project with ``conj(C(m,phi))``; ``bar_P`` carries no CS phase and our ``sign_m`` restores it. Both are real-weighted sums, so - ours[n, l, m] = k * conj(phaser[l, m, n+1]), k = 2 if enforce_friedel else 1 + ours[n, l, m] = k * conj(phaser[l, m, n+1]), k = 1 (was 2 while the +expansion concatenated the antipodal copy, which Phaser does via cctbx's +conjugate_flag and we no longer do -- see bessel_sh_expand) with the factor 2 because appending ``-s`` doubles every even-l coefficient exactly (``Y_lm(-s) = (-1)^l Y_lm(s)``). ``k`` is therefore a prediction, not a @@ -295,11 +297,10 @@ def compare_coeffs(ours: torch.Tensor, phaser: torch.Tensor, L: int, # --------------------------------------------------------------------------- def encode(s: torch.Tensor, intensity: torch.Tensor, *, L: int, - h_scale: float, zsymm: int, friedel: bool) -> torch.Tensor: + h_scale: float, zsymm: int) -> torch.Tensor: from torchref.experimental.alignment.frf.data_mr import bessel_sh_expand return bessel_sh_expand( s, intensity, L=L, bessel_h_scale=h_scale, zsymm=zsymm, - enforce_friedel=friedel, ).coeffs @@ -428,13 +429,13 @@ def emit(arm, coeffs, target, *, n_points, seconds, extra=None): # --- arm 1: Phaser's own observations through our encoder --------------- t0 = time.time() - c = encode(s_obs, i_obs, L=L, h_scale=h_obs, zsymm=zsymm, friedel=False) + c = encode(s_obs, i_obs, L=L, h_scale=h_obs, zsymm=zsymm) emit("obs_phaser_pts", c, data_elmn, n_points=s_obs.shape[0], seconds=time.time() - t0, extra=obs_clu) # --- arm 2: Phaser's observations WITH Phaser's own theta approximation -- t0 = time.time() - c = encode(s_obs_clu, i_obs, L=L, h_scale=h_obs, zsymm=zsymm, friedel=False) + c = encode(s_obs_clu, i_obs, L=L, h_scale=h_obs, zsymm=zsymm) emit("obs_phaser_clustered", c, data_elmn, n_points=s_obs_clu.shape[0], seconds=time.time() - t0, extra=obs_clu) @@ -445,15 +446,14 @@ def emit(arm, coeffs, target, *, n_points, seconds, extra=None): # recoverable from the dumped theta. t0 = time.time() flip = (cos_calc.abs() > 1e-12).to(torch.float64) + 1.0 - c = encode(s_calc, i_calc * flip, L=L, h_scale=h_calc, zsymm=1, friedel=False) + c = encode(s_calc, i_calc * flip, L=L, h_scale=h_calc, zsymm=1) emit("calc_phaser_pts", c, search_elmn, n_points=s_calc.shape[0], seconds=time.time() - t0, extra=dict(calc_clu, n_l0_plane=int((cos_calc.abs() <= 1e-12).sum()))) # --- arm 4: the same, with Phaser's theta approximation ----------------- t0 = time.time() - c = encode(s_calc_clu, i_calc * flip, L=L, h_scale=h_calc, zsymm=1, - friedel=False) + c = encode(s_calc_clu, i_calc * flip, L=L, h_scale=h_calc, zsymm=1) emit("calc_phaser_clustered", c, search_elmn, n_points=s_calc_clu.shape[0], seconds=time.time() - t0, extra=calc_clu) @@ -466,7 +466,7 @@ def emit(arm, coeffs, target, *, n_points, seconds, extra=None): for arm, (s, val) in arms.items(): t0 = time.time() fc = frame_check(s, s_obs) - c = encode(s, val, L=L, h_scale=h_obs, zsymm=zsymm, friedel=False) + c = encode(s, val, L=L, h_scale=h_obs, zsymm=zsymm) emit(arm, c, data_elmn, n_points=s.shape[0], seconds=time.time() - t0, extra=dict(ustats, **fc)) diff --git a/docs/changelog.rst b/docs/changelog.rst index edbbe5ed..799dcd39 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -25,6 +25,8 @@ Unreleased - The rotation function assigns resolution shells once and shares them, instead of the Wilson normalisation and the variance reweight deriving edges that disagreed at the shell boundaries - The rotation function masks observations to the bandwidth-coupled resolution before the symmetry unroll rather than after - The rotation function warns on Bijvoet-unmerged data, whose shared canonical index would weight those reflections twice +- The rotation function no longer concatenates the antipodal copy onto either reflection set: only even harmonic degrees are computed, for which it is an exact factor of two, so it scaled the rotation function by four and changed no ranking. Raw ``RotationPeak.score`` and ``RotationSolutions.scores`` are therefore a quarter of their previous values; z-scores are unchanged +- Kept the rotation function's relative Wilson-B fit: knocking it out was measured rank-neutral but worth only 2% of the runtime once the fit moved to the unique reflection set Version 0.6.4 ---------- diff --git a/tests/unit/frf_separate/test_bessel_sh_grouping.py b/tests/unit/frf_separate/test_bessel_sh_grouping.py index 7bfb3d75..fcfac972 100644 --- a/tests/unit/frf_separate/test_bessel_sh_grouping.py +++ b/tests/unit/frf_separate/test_bessel_sh_grouping.py @@ -159,9 +159,6 @@ def _reference_expansion(s_vec, intensity, L, bessel_h_scale): from torchref.experimental.alignment.frf.data_mr import spherical_bessel_table from torchref.experimental.alignment.sh import _bar_legendre_recurrence - s_vec = torch.cat([s_vec, -s_vec], dim=0) # enforce_friedel - intensity = torch.cat([intensity, intensity], dim=0) - lmax = L - 1 lmax_even = lmax if lmax % 2 == 0 else lmax - 1 N_radial = (lmax_even - 2) // 2 + 1 @@ -201,3 +198,29 @@ def test_matches_an_independent_direct_summation(L, double_cpu): assert got.shape == ref.shape rel = (got - ref).abs().max().item() / max(ref.abs().max().item(), 1e-300) assert rel < 1e-10, f"L={L}: differs from a direct summation by {rel:.2e}" + + +def test_the_antipodal_copy_would_only_double_the_result(double_cpu): + """Why the expansion no longer concatenates ``-s``. + + Only even ``l`` are computed, and ``Y_lm(-s_hat) = (-1)^l Y_lm(s_hat)``, so + for even ``l`` the antipodal copy contributes exactly what the original does. + The intensity is duplicated verbatim and ``|s|`` is unchanged, so the whole + coefficient array doubles and nothing about its *shape* changes -- which is + why removing it rescales the rotation function by four and reorders nothing. + + Pinned here rather than argued in a comment: if a future change made odd ``l`` + carry signal, this equality would break and the removal would need revisiting. + """ + s, I = _random_set(seed=17, n=300) + single = bessel_sh_expand(s, I, L=17, bessel_h_scale=30.0).coeffs + doubled = bessel_sh_expand( + torch.cat([s, -s], dim=0), torch.cat([I, I], dim=0), + L=17, bessel_h_scale=30.0, + ).coeffs + scale = single.abs().max().clamp(min=1e-300) + rel = ((doubled - 2.0 * single).abs().max() / scale).item() + assert rel < 1e-12, ( + f"the antipodal copy is not an exact factor of two: {rel:.2e} relative. " + f"Removing it was justified on that being exact." + ) diff --git a/torchref/experimental/alignment/frf/api.py b/torchref/experimental/alignment/frf/api.py index 82c27e53..d3ad1bdc 100644 --- a/torchref/experimental/alignment/frf/api.py +++ b/torchref/experimental/alignment/frf/api.py @@ -255,7 +255,7 @@ def __init__( self._c_obs = bessel_sh_expand( s_obs, intensity_obs, L=L, bessel_h_scale=self.bessel_h_scale, - zsymm=zsymm, enforce_friedel=True, + zsymm=zsymm, ) @@ -302,7 +302,7 @@ def score_model( c_calc = bessel_sh_expand( s_calc, intensity_calc, L=self.L, bessel_h_scale=self.bessel_h_scale, - zsymm=1, enforce_friedel=True, + zsymm=1, ) # 8. Cross-correlate over the radial axis. diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index 76b97995..e61c0a14 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -186,7 +186,6 @@ def bessel_sh_expand( L: int, bessel_h_scale: float, zsymm: int = 1, - enforce_friedel: bool = True, ) -> BesselSHCoefficients: """Phaser-style ``c_nlm = Σ_h Y*_lm(ŝ) · I · sqrt(2u+1) · j_u(h)/h``. @@ -202,6 +201,23 @@ def bessel_sh_expand( * radial × SH expansion, sqrt(2u+1)·j_u(h)/h weight: DataMR.cc:993, 1107 * even-l only (Patterson centrosymmetry) + m-filter: DataMR.cc:863-870, 1117 + **No antipodal copy.** The Patterson's centrosymmetry is already encoded + twice here -- only even ``l`` are computed, and the negative-``m`` half is + mirrored rather than summed -- and both of those *save* work. Concatenating + ``-s`` onto the reflection set was a third encoding that *cost* work and + bought nothing: for even ``l``, ``Y_lm(-s_hat) = Y_lm(s_hat)``, and the + intensity, Bessel weight and Legendre factor are all unchanged under + negation, so it doubled ``c_nlm`` exactly. Both sides doubled scaled the + rotation function by 4, which the z-score normalisation removes. + + Measured before removal, over 10 benchmark structures x 10 seeded trials: + truth ranks 98/100 identical (1 better, 1 worse), the top score exactly + 0.2499974 to 0.2500036 of the doubled value, and the search 22.5% faster on + 3K7M / 16.4% on 1DAW. Note that Phaser does include the mate (cctbx's + ``conjugate_flag``), so our coefficients are now half of its -- which + matters only to the coefficient-level comparison in + ``alignment_lab/diagnostics/frf_encode_compare.py``. + Two precisions are in play and they are deliberately different. The **clustering keys** are computed at ``s_vectors``' own dtype, because @@ -229,10 +245,6 @@ def bessel_sh_expand( complex_dtype = get_complex_dtype() real_dtype = comp_real - if enforce_friedel: - s_vectors = torch.cat([s_vectors, -s_vectors], dim=0) - intensity = torch.cat([intensity, intensity], dim=0) - lmax = L - 1 lmax_even = lmax if (lmax % 2 == 0) else (lmax - 1) if lmax_even < 2: From 8d60e82eed0f4d5161914458dfe19b4ee8eecd21 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 24 Aug 2026 03:17:56 +0200 Subject: [PATCH 043/250] Make double precision a device capability rather than a constant Audited where the rotation search actually creates float64/complex128, by intercepting every torch call and attributing each double tensor to the source line that made it (`alignment_lab/analysis/double_audit.py`). The result contradicts what the plan assumed. On 1DAW cap64, by element count: wigner_d.py:163 S += xi * d_l 66.6M deliberate wide accumulation data_mr.py:157 Bessel rescale 13.6M deliberate exponent range sitelist_ang pad + fft2 8.7M follows xi data_mr.py:562 xi einsum 1.1M deliberate dense_calc, per-reflection arrays ~5.0M incidental peak_finder 9.1M already explicitly .cpu() So the MPS blocker is NOT the incidental casts the plan was going to narrow -- those are ~5M of ~100M double elements -- it is the accumulation that a previous measurement showed must stay wide (narrowing it left 1 of 500 candidate slots holding the same orientation). "Run on MPS" and "keep the accumulation precision" cannot both be satisfied by a hardcoded dtype. So the width asks the device. `supports_double`, `widest_float_dtype` and `widest_complex_dtype` join the existing dtype/device policy in `torchref.config`, and the two places whose precision is load-bearing for a *reason* rather than a preference consume them: - `cross_correlate_xi`, whose radial sum cancels, accumulates one step wider than the coefficients where the device allows and never narrows an input the caller already widened; - `spherical_bessel_table`, whose unnormalised ladder reaches 1e157, uses double where available. The power-of-two rescaling added earlier is what makes the narrow branch survivable there at all -- without it the ladder overflows float32 for every x below ~35. This is the repo's capability-based dispatch convention applied to dtype, and it answers "only use float64 if really necessary" by making necessity a property of the device, decided in one place. Bit-identical on 3K7M and 1DAW at cap64: on a device WITH float64 the helpers must resolve to exactly what was hardcoded, and they do. 529 passed / 11 skipped. Not claimed: that MPS works. There is no accelerator on this host, which is also why the audit's device column reads cpu for every site and the sites had to be resolved by hand. The no-float64 branch is unexercisable here, so the test pins the policy -- which backends lack float64, and that the fallback is the configured working dtype rather than a second hardcoded constant -- instead of pretending to verify behaviour. Anyone running there should expect the accuracy cost the accumulation measurement quantified, and gate it on their own panel. Also struck two plan items on the evidence rather than implementing them: evaluating eterm/solvent/apply_overall_anisotropy per shell instead of per reflection (all three measure 0.000 s), and integer cluster keys on the calc box (data_mr's own docstring already records that the dense box groups identically at 1e7, 1e9 and 1e11, so the float keys are already exact there). Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/capability_gate.sh | 50 +++++++++ alignment_lab/analysis/double_audit.py | 102 ++++++++++++++++++ alignment_lab/analysis/double_audit.sh | 19 ++++ alignment_lab/analysis/where_now_cap64.sh | 24 +++++ alignment_lab/lab/profile.py | 17 +++ docs/changelog.rst | 1 + .../test_rotation_search_dtype_device.py | 46 ++++++++ torchref/config.py | 35 ++++++ .../experimental/alignment/frf/data_mr.py | 24 +++-- 9 files changed, 309 insertions(+), 9 deletions(-) create mode 100644 alignment_lab/analysis/capability_gate.sh create mode 100644 alignment_lab/analysis/double_audit.py create mode 100644 alignment_lab/analysis/double_audit.sh create mode 100644 alignment_lab/analysis/where_now_cap64.sh diff --git a/alignment_lab/analysis/capability_gate.sh b/alignment_lab/analysis/capability_gate.sh new file mode 100644 index 00000000..bcf5b22c --- /dev/null +++ b/alignment_lab/analysis/capability_gate.sh @@ -0,0 +1,50 @@ +#!/bin/bash +# On a device WITH float64 the capability helpers must resolve to exactly what +# was hardcoded before, so this has to be bit-identical. That is the whole gate: +# the MPS branch cannot be exercised here and is not claimed to work. +#SBATCH --job-name=capgate +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:50:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +OUT=$NEW/alignment_lab/slurm +cd "$NEW"; export PYTHONPATH="$NEW" +export TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname)" +LOG=$OUT/capgate_tests_$SLURM_JOB_ID.log +"$PY" -m pytest tests/unit/alignment tests/unit/frf_separate tests/unit/model \ + tests/unit/test_imports_smoke.py -q > "$LOG" 2>&1 +rc=$?; echo "=== TESTS rc=$rc ==="; tail -12 "$LOG" +echo "=== capability helpers resolve as expected on this host ===" +"$PY" - <<'PYEOF' +import torch +from torchref.config import (supports_double, widest_complex_dtype, + widest_float_dtype) +print(f" cpu: supports_double={supports_double('cpu')} " + f"float={widest_float_dtype('cpu')} complex={widest_complex_dtype('cpu')}") +assert widest_float_dtype("cpu") is torch.float64 +assert widest_complex_dtype("cpu") is torch.complex128 +print(" mps branch (not exercisable here, resolution only):") +from torchref.config import _NO_DOUBLE_DEVICE_TYPES +print(f" device types without float64: {_NO_DOUBLE_DEVICE_TYPES}") +PYEOF +echo "=== fingerprint vs the committed state (baseline worktree at fe53d373) ===" +OLD=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/_stagea_baseline +for pdb in 3K7M 1DAW; do + for tree in OLD NEW; do + eval "root=\$$tree"; cd "$root" + PYTHONPATH="$root" "$PY" -u "$NEW/alignment_lab/analysis/frf_fingerprint.py" \ + --pdb "$pdb" --lmax-cap 64 > "$OUT/cap_${tree}_${pdb}.txt" 2>/dev/null + done + diff -q <(grep '^FP ' "$OUT/cap_OLD_${pdb}.txt") <(grep '^FP ' "$OUT/cap_NEW_${pdb}.txt") >/dev/null \ + && echo "IDENTICAL $pdb ($(grep -c '^FP ' "$OUT/cap_NEW_${pdb}.txt") peaks)" \ + || { echo "DIFFERS $pdb"; diff <(grep '^FP ' "$OUT/cap_OLD_${pdb}.txt") <(grep '^FP ' "$OUT/cap_NEW_${pdb}.txt") | head -4; } +done +echo "capgate_rc=$rc" diff --git a/alignment_lab/analysis/double_audit.py b/alignment_lab/analysis/double_audit.py new file mode 100644 index 00000000..0dc480ff --- /dev/null +++ b/alignment_lab/analysis/double_audit.py @@ -0,0 +1,102 @@ +"""Where does the rotation search create float64 / complex128 tensors? + +MPS has no float64 at all, so every double-precision tensor on the compute +device is a portability blocker. Rather than guess at the list -- twice now a +confident guess about this engine has been wrong -- intercept every torch call +and report the source line that produced each double tensor, with how many and +how large. + +Deliberately reports rather than asserts. Some of these are *correct* and must +stay: the spherical-Bessel ladder needs the exponent range, the J_y +eigendecomposition and the anisotropy fit are precision-critical, and anything +already on the host costs nothing. The point is to separate those from the +per-reflection arrays that are double by inheritance. +""" + +from __future__ import annotations + +import argparse +import collections +import sys +import traceback +from pathlib import Path + +import torch +from torch.overrides import TorchFunctionMode + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +_DOUBLE = (torch.float64, torch.complex128) + + +class DoubleAudit(TorchFunctionMode): + """Attribute every double-precision tensor to the line that made it.""" + + def __init__(self, package_only: str = "torchref"): + super().__init__() + self.package_only = package_only + self.sites = collections.Counter() + self.elems = collections.Counter() + self.devices = collections.defaultdict(set) + + def __torch_function__(self, func, types, args=(), kwargs=None): + out = func(*args, **(kwargs or {})) + try: + tensors = [] + if isinstance(out, torch.Tensor): + tensors = [out] + elif isinstance(out, (tuple, list)): + tensors = [t for t in out if isinstance(t, torch.Tensor)] + if any(t.dtype in _DOUBLE for t in tensors): + # Innermost frame inside the package under audit, so the report + # names our code and not torch's internals. + site = None + for fr in reversed(traceback.extract_stack()[:-1]): + if f"/{self.package_only}/" in fr.filename: + short = fr.filename.split(f"/{self.package_only}/", 1)[1] + site = f"{self.package_only}/{short}:{fr.lineno}" + break + if site is not None: + n = sum(t.numel() for t in tensors if t.dtype in _DOUBLE) + self.sites[site] += 1 + self.elems[site] += n + for t in tensors: + if t.dtype in _DOUBLE: + self.devices[site].add(str(t.device)) + except Exception: # pragma: no cover - diagnostic + pass + return out + + def report(self, top: int = 30) -> None: + print(f"{'site':62s} {'calls':>7s} {'elements':>12s} devices") + for site, elems in self.elems.most_common(top): + print(f"{site:62s} {self.sites[site]:>7d} {elems:>12d} " + f"{','.join(sorted(self.devices[site]))}") + print(f"\n{len(self.elems)} distinct sites, " + f"{sum(self.elems.values())} double elements total") + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--pdb", default="1DAW") + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--top", type=int, default=30) + args = ap.parse_args() + + from lab import FRFConfig, rotated_case, run_frf, seed_for + + model, data, _ = rotated_case(args.pdb, seed_for(args.pdb, 0)) + cfg = FRFConfig(n_peaks=500, lmax_cap=args.lmax_cap) + run_frf(model, data, cfg, capture_arf=False, verbose=0) # warm the memos + + audit = DoubleAudit() + with audit: + run_frf(model, data, cfg, capture_arf=False, verbose=0) + print(f"=== {args.pdb} cap{args.lmax_cap}: double-precision sites ===") + audit.report(args.top) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/alignment_lab/analysis/double_audit.sh b/alignment_lab/analysis/double_audit.sh new file mode 100644 index 00000000..71a5b187 --- /dev/null +++ b/alignment_lab/analysis/double_audit.sh @@ -0,0 +1,19 @@ +#!/bin/bash +#SBATCH --job-name=dblaudit +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:40:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname)" +"$PY" -u alignment_lab/analysis/double_audit.py --pdb 1DAW --top 40 2>&1 \ + | grep -vE "UserWarning|FutureWarning|^ *from |^ *warnings\.|Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization" +echo "rc=$?" diff --git a/alignment_lab/analysis/where_now_cap64.sh b/alignment_lab/analysis/where_now_cap64.sh new file mode 100644 index 00000000..c5d1aa79 --- /dev/null +++ b/alignment_lab/analysis/where_now_cap64.sh @@ -0,0 +1,24 @@ +#!/bin/bash +# Where the time goes NOW. cap64 only -- that is what ships, and medianing it +# with cap100 inverted the stage ranking once already. +#SBATCH --job-name=wherenow +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=48G +#SBATCH --exclusive +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" +for pdb in 3K7M 1AK5 1DAW; do + echo "############ $pdb ############" + "$PY" -u -m alignment_lab.diagnostics.frf_benchmark --pdb "$pdb" --arms cap64 --trials 3 2>&1 \ + | grep -vE "Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|warn|^ *$" +done diff --git a/alignment_lab/lab/profile.py b/alignment_lab/lab/profile.py index 9a091905..16056594 100644 --- a/alignment_lab/lab/profile.py +++ b/alignment_lab/lab/profile.py @@ -52,6 +52,23 @@ ("torchref.experimental.alignment.frf.sitelist_ang", "build_dense_map_per_beta"), ("torchref.experimental.alignment.frf.data_mr", "spherical_bessel_table"), + # Added once the named stages stopped accounting for the run: after the obs + # chain moved to the unique set, French-Wilson fell from 29.9% to 3.6% and + # the unattributed remainder became the second largest item at 29%. These are + # the rest of the per-reflection work, patched where each call RESOLVES -- + # `api` imports the preprocessing names at module top, `rotation_search` + # imports `apply_overall_anisotropy` from `sh` at module top, and + # `fit_relative_wilson_b` is imported inside `search_peaks` so it has to be + # patched on the defining module. + ("torchref.experimental.alignment.frf.api", "wilson_normalise"), + ("torchref.experimental.alignment.frf.api", "eterm_sigma_a"), + ("torchref.experimental.alignment.frf.api", "build_lerf1_intensity"), + ("torchref.experimental.alignment.frf.api", "apply_shell_variance_weights"), + ("torchref.experimental.alignment.frf.api", "detect_zsymm"), + ("torchref.experimental.alignment.frf.preprocessing", + "fit_relative_wilson_b"), + ("torchref.experimental.alignment.rotation_search", + "apply_overall_anisotropy"), ) #: Which stage each nested stage sits inside. A parent's time *includes* its diff --git a/docs/changelog.rst b/docs/changelog.rst index 799dcd39..37a19cc9 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -27,6 +27,7 @@ Unreleased - The rotation function warns on Bijvoet-unmerged data, whose shared canonical index would weight those reflections twice - The rotation function no longer concatenates the antipodal copy onto either reflection set: only even harmonic degrees are computed, for which it is an exact factor of two, so it scaled the rotation function by four and changed no ranking. Raw ``RotationPeak.score`` and ``RotationSolutions.scores`` are therefore a quarter of their previous values; z-scores are unchanged - Kept the rotation function's relative Wilson-B fit: knocking it out was measured rank-neutral but worth only 2% of the runtime once the fit moved to the unique reflection set +- Added ``supports_double`` / ``widest_float_dtype`` / ``widest_complex_dtype``: where precision is load-bearing the width now comes from the device rather than a hardcoded ``float64``, so a backend without it gets the working dtype instead of an error Version 0.6.4 ---------- diff --git a/tests/unit/alignment/test_rotation_search_dtype_device.py b/tests/unit/alignment/test_rotation_search_dtype_device.py index c0bd0c7d..34145b51 100644 --- a/tests/unit/alignment/test_rotation_search_dtype_device.py +++ b/tests/unit/alignment/test_rotation_search_dtype_device.py @@ -149,3 +149,49 @@ def test_device_is_resolved_from_both_inputs(): resolved = resolve_device(data, model) assert resolved == data.hkl.device or resolved.type == data.hkl.device.type assert model.xyz().device.type == resolved.type + + +def test_double_is_a_device_capability_not_a_constant(): + """Where precision is load-bearing, the width comes from the device. + + Two places in the engine need more than the working precision for reasons + that are not preferences: the radial accumulation cancels (see + ``cross_correlate_xi``) and the Bessel ladder's unnormalised intermediates + reach 1e157. Both used to hardcode float64, which is an error on a backend + that has none. They now ask, so that a device without float64 gets the + working dtype -- and the accuracy that implies -- instead of a crash. + """ + from torchref.config import (supports_double, widest_complex_dtype, + widest_float_dtype) + + assert supports_double("cpu") + assert widest_float_dtype("cpu") is torch.float64 + assert widest_complex_dtype("cpu") is torch.complex128 + + # The no-float64 branch cannot be exercised on a host without such a device, + # so pin the policy itself: exactly the backends known to lack float64, and + # the fallback is the configured working dtype rather than a second constant. + from torchref.config import _NO_DOUBLE_DEVICE_TYPES + assert "mps" in _NO_DOUBLE_DEVICE_TYPES + assert not supports_double(torch.device("mps")) + original = dtypes.float, dtypes.complex + try: + dtypes.float, dtypes.complex = torch.float32, torch.complex64 + assert widest_float_dtype(torch.device("mps")) is torch.float32 + assert widest_complex_dtype(torch.device("mps")) is torch.complex64 + dtypes.float, dtypes.complex = torch.float64, torch.complex128 + assert widest_float_dtype(torch.device("mps")) is torch.float64 + finally: + dtypes.float, dtypes.complex = original + + +def test_the_radial_accumulation_never_narrows_a_double_input(): + """A caller that already widened must not be silently narrowed back.""" + from torchref.experimental.alignment.frf.data_mr import cross_correlate_xi + from torchref.experimental.alignment.frf.types import BesselSHCoefficients + + c = BesselSHCoefficients( + coeffs=torch.zeros((2, 5, 9), dtype=torch.complex128), + L=5, bessel_h_scale=20.0, + ) + assert cross_correlate_xi(c, c).dtype == torch.complex128 diff --git a/torchref/config.py b/torchref/config.py index 617b7627..f528df65 100644 --- a/torchref/config.py +++ b/torchref/config.py @@ -568,3 +568,38 @@ def __repr__(self) -> str: def get_default_device() -> torch.device: """Get the current default device.""" return device.current + + +# --------------------------------------------------------------------------- +# Double-precision availability +# --------------------------------------------------------------------------- +#: Device types with no float64 at all. MPS is the live case and it *raises* +#: rather than quietly downcasting, so a float64 tensor there is an error and not +#: merely slow. +_NO_DOUBLE_DEVICE_TYPES = ("mps",) + + +def supports_double(dev=None) -> bool: + """Whether ``dev`` can hold float64 / complex128 at all.""" + return normalize_device(dev).type not in _NO_DOUBLE_DEVICE_TYPES + + +def widest_float_dtype(dev=None) -> torch.dtype: + """``float64`` where the device has it, else the configured working float. + + For computations whose *precision* is load-bearing rather than their storage: + accumulating single-precision data in double is the ordinary remedy, and the + dynamic range of an unnormalised recurrence is a hard requirement rather than + a preference. The right width for those is a property of the device, so it + belongs here and not in a constant at the call site. + + Where the device lacks float64 the caller gets the working dtype and whatever + accuracy that implies. That is the only option there, not a choice -- callers + that care should say what it costs in their own docstring. + """ + return torch.float64 if supports_double(dev) else get_float_dtype() + + +def widest_complex_dtype(dev=None) -> torch.dtype: + """``complex128`` where the device has it, else the configured working complex.""" + return torch.complex128 if supports_double(dev) else get_complex_dtype() diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index e61c0a14..c1bb0cd5 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -58,7 +58,8 @@ #: approximation and needs its own evidence. _GROUP_SCALE_COS = 10_000_000 -from ....config import get_complex_dtype, get_float_dtype +from ....config import (get_complex_dtype, get_float_dtype, + widest_complex_dtype, widest_float_dtype) from ....utils.backends import run_or_degrade, select from ..sh import legendre_recurrence_coefficients from ._backends import LEGENDRE_BACKENDS @@ -122,7 +123,12 @@ def spherical_bessel_table( """ real_dtype = x.dtype device = x.device - x64 = x.to(torch.float64) + # Double where the device has it. The rescaling above is what makes a + # narrower working dtype survivable here at all -- without it the ladder + # overflows float32 for every x below ~35 -- but double is still preferable + # where it is available, since the recurrence runs ~90 steps. + work_dtype = widest_float_dtype(device) + x64 = x.to(work_dtype) safe_x = x64.clamp(min=1e-30) inv_x = 1.0 / safe_x @@ -130,7 +136,7 @@ def spherical_bessel_table( j_high = torch.zeros_like(x64) j_mid = torch.ones_like(x64) j_table = torch.zeros( - (u_max + 1, *x64.shape), dtype=torch.float64, device=device, + (u_max + 1, *x64.shape), dtype=work_dtype, device=device, ) threshold = float(2 ** _BESSEL_RESCALE_EXP) inv_threshold = 1.0 / threshold @@ -553,12 +559,12 @@ def cross_correlate_xi( """ if c_obs.L != c_calc.L: raise ValueError(f"L mismatch: obs={c_obs.L} calc={c_calc.L}") - acc = ( - torch.complex128 - if c_obs.coeffs.dtype in (torch.complex64, torch.complex128) - and torch.finfo(c_obs.coeffs.real.dtype).bits <= 32 - else c_obs.coeffs.dtype - ) + # One step wider than the coefficients where the device allows it. On a + # backend without float64 this falls back to the coefficients' own dtype and + # the run pays the accuracy noted above -- there is no third option there. + acc = widest_complex_dtype(c_obs.coeffs.device) + if c_obs.coeffs.dtype == torch.complex128: + acc = torch.complex128 # never narrow what the caller widened return torch.einsum( "rln,rlm->lmn", c_obs.coeffs.to(acc), From 992e7a6765de47d6e9997b3d3dcc02228320df44 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 24 Aug 2026 13:13:39 +0200 Subject: [PATCH 044/250] Add characterisation tests for fraction storage and the collection targets Nothing in the suite asserted anything about how ModelCollection stores population fractions, or about which accessors and subset the collection X-ray targets read. Both are about to change. The collection targets are pinned by deterministic invariants rather than stored loss values: model.forward is run-to-run nondeterministic even within one process, and an invariant that says "the loss must respond to this input" is both immune to that and a sharper statement than a literal. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- .../model/test_model_collection_fractions.py | 276 ++++++++++++++++++ ...test_collection_target_characterisation.py | 276 ++++++++++++++++++ 2 files changed, 552 insertions(+) create mode 100644 tests/unit/model/test_model_collection_fractions.py create mode 100644 tests/unit/refinement/test_collection_target_characterisation.py diff --git a/tests/unit/model/test_model_collection_fractions.py b/tests/unit/model/test_model_collection_fractions.py new file mode 100644 index 00000000..4266dddd --- /dev/null +++ b/tests/unit/model/test_model_collection_fractions.py @@ -0,0 +1,276 @@ +"""Characterisation of how ``ModelCollection`` stores population fractions. + +Nothing else in the suite asserts anything about fraction storage -- not the softmax +parametrisation, not the sum-to-1 validation, not the freeze flags, not the override path. +These tests pin the observable contract so a change of storage has to reproduce it rather +than merely still run. + +Deliberately fileless: the fractions live on ``_SharedMixedModel`` and depend on the base +models only for ``dtype_float`` and device, so a stub is enough and the whole file runs in +well under a second. +""" + +import pytest +import torch +from torch import nn + + +class _StubModel(nn.Module): + """Minimal stand-in for ``ModelFT`` for fraction bookkeeping. + + Carries a real parameter so ``.to()`` and ``resolve_device`` behave, exposes the two + attributes ``_SharedMixedModel.__init__`` reads, and returns structure factors that + differ per instance so a weighted sum can be checked against its parts. + """ + + def __init__(self, seed: int): + super().__init__() + self.anchor = nn.Parameter(torch.zeros(1)) + self._seed = seed + + @property + def device(self): + return self.anchor.device + + @property + def dtype_float(self): + return self.anchor.dtype + + def forward(self, hkl, recalc: bool = False): + # Distinct per model, and a function of the parameter so a gradient can reach it. + n = hkl.shape[0] + base = torch.arange(1, n + 1, dtype=self.anchor.dtype, device=self.anchor.device) + amp = (base * float(self._seed + 1)) + self.anchor + return amp.to(torch.complex64) + + +@pytest.fixture +def two_model_collection(): + """A dark + one-timepoint collection over two distinguishable stub models.""" + from torchref.model.model_collection import ModelCollection + + mc = ModelCollection([_StubModel(0), _StubModel(1)], dark_key="dark", verbose=0) + mc.add_dark() + mc.add_timepoint("light", [0.7, 0.3]) + return mc + + +@pytest.fixture +def hkl(): + return torch.tensor([[1, 0, 0], [0, 1, 0], [1, 1, 0], [2, 0, 1]]) + + +class TestSoftmaxStorage: + @pytest.mark.unit + def test_fractions_are_the_softmax_of_the_stored_logits(self, two_model_collection): + mc = two_model_collection + mixed = mc["light"] + expected = torch.softmax(mixed.fraction_params, dim=0) + assert torch.allclose(mixed.fractions, expected) + + @pytest.mark.unit + @pytest.mark.parametrize("f", [0.01, 0.22, 0.3, 0.5, 0.99]) + def test_requested_fractions_round_trip(self, f): + """What you pass to ``add_timepoint`` is what ``fractions`` reports back.""" + from torchref.model.model_collection import ModelCollection + + mc = ModelCollection([_StubModel(0), _StubModel(1)], verbose=0) + mc.add_timepoint("t", [1.0 - f, f]) + assert mc["t"].fractions[1].item() == pytest.approx(f, abs=1e-6) + assert mc["t"].fractions.sum().item() == pytest.approx(1.0, abs=1e-6) + + @pytest.mark.unit + def test_dark_excited_fraction_is_the_clamp_floor_not_zero( + self, two_model_collection + ): + """``add_dark`` asks for exactly 0, but the log-clamp at 1e-6 means the stored + value is 1e-6. Anything deriving a bound from the dark's fraction inherits that + floor rather than a true zero. + """ + dark = two_model_collection["dark"] + assert dark.fractions[1].item() == pytest.approx(1e-6, rel=1e-3) + assert dark.fractions[1].item() > 0.0 + + @pytest.mark.unit + def test_fraction_dtype_and_device_follow_the_base_models(self): + from torchref.model.model_collection import ModelCollection + + models = [_StubModel(0), _StubModel(1)] + mc = ModelCollection(models, verbose=0) + mc.add_timepoint("t", [0.6, 0.4]) + assert mc["t"].fraction_params.dtype == models[0].dtype_float + assert mc["t"].fraction_params.device == models[0].device + + +class TestValidation: + @pytest.mark.unit + def test_fractions_must_sum_to_one(self): + from torchref.model.model_collection import ModelCollection + + mc = ModelCollection([_StubModel(0), _StubModel(1)], verbose=0) + with pytest.raises(ValueError, match="sum to 1"): + mc.add_timepoint("t", [0.5, 0.2]) + + @pytest.mark.unit + def test_fraction_count_must_match_the_model_count(self): + from torchref.model.model_collection import ModelCollection + + mc = ModelCollection([_StubModel(0), _StubModel(1)], verbose=0) + with pytest.raises(ValueError, match="must match"): + mc.add_timepoint("t", [1.0]) + + @pytest.mark.unit + def test_duplicate_timepoint_names_are_rejected(self, two_model_collection): + with pytest.raises(ValueError, match="already exists"): + two_model_collection.add_timepoint("light", [0.5, 0.5]) + + +class TestFreezing: + @pytest.mark.unit + def test_dark_is_frozen_and_timepoints_are_not(self, two_model_collection): + mc = two_model_collection + assert mc["dark"].fraction_params.requires_grad is False + assert mc["light"].fraction_params.requires_grad is True + + @pytest.mark.unit + def test_freeze_and_unfreeze_flip_the_flag(self, two_model_collection): + mixed = two_model_collection["light"] + mixed.freeze_fractions() + assert mixed.fraction_params.requires_grad is False + mixed.unfreeze_fractions() + assert mixed.fraction_params.requires_grad is True + + @pytest.mark.unit + def test_unfreeze_all_leaves_the_dark_frozen(self, two_model_collection): + """The dark reference must not become refinable through the bulk call.""" + mc = two_model_collection + mc.unfreeze_all_fractions() + assert mc["dark"].fraction_params.requires_grad is False + assert mc["light"].fraction_params.requires_grad is True + + @pytest.mark.unit + def test_freeze_all_freezes_every_timepoint(self, two_model_collection): + mc = two_model_collection + mc.freeze_all_fractions() + assert all( + mc[k].fraction_params.requires_grad is False for k in mc.keys() + ) + + +class TestOverride: + @pytest.mark.unit + def test_override_replaces_fractions_and_clears_back(self, two_model_collection): + mixed = two_model_collection["light"] + forced = torch.tensor([0.1, 0.9]) + + mixed.set_fraction_override(forced) + assert mixed.fractions is forced + + mixed.clear_fraction_override() + assert torch.allclose( + mixed.fractions, torch.softmax(mixed.fraction_params, dim=0) + ) + + @pytest.mark.unit + def test_override_reaches_the_forward(self, two_model_collection, hkl): + """The override is the point at which an external (kinetic) population enters + the structure-factor sum, so it has to change ``forward``, not just the property. + """ + mixed = two_model_collection["light"] + before = mixed(hkl, recalc=True) + + mixed.set_fraction_override(torch.tensor([0.1, 0.9])) + after = mixed(hkl, recalc=True) + + assert not torch.allclose(before, after) + + @pytest.mark.unit + def test_override_carries_gradient(self, two_model_collection, hkl): + """Gradients must flow through the override to whatever produced it.""" + mixed = two_model_collection["light"] + forced = torch.tensor([0.4, 0.6], requires_grad=True) + mixed.set_fraction_override(forced) + + mixed(hkl, recalc=True).abs().sum().backward() + + assert forced.grad is not None + assert torch.isfinite(forced.grad).all() + + +class TestCollectionLevelViews: + @pytest.mark.unit + def test_fractions_matrix_rows_follow_insertion_order(self, two_model_collection): + mc = two_model_collection + matrix = mc.get_fractions_matrix() + assert matrix.shape == (2, 2) + for row, key in enumerate(mc.keys()): + assert torch.allclose(matrix[row], mc[key].fractions) + + @pytest.mark.unit + def test_timepoint_names_excludes_the_dark_key(self, two_model_collection): + mc = two_model_collection + assert mc.dark_key == "dark" + assert mc.timepoint_names == ["light"] + assert mc.keys() == ["dark", "light"] + + @pytest.mark.unit + def test_base_models_are_shared_not_copied(self, two_model_collection): + mc = two_model_collection + assert mc.n_base_models == 2 + for i in range(2): + assert mc["dark"].models[i] is mc["light"].models[i] + assert mc["dark"].models[i] is mc.base_models[i] + + @pytest.mark.unit + def test_a_timepoint_owns_only_its_fractions(self, two_model_collection): + """``_SharedMixedModel`` holds the base models in a plain list, not a + ``ModuleList``, so a timepoint must not re-register their parameters. + """ + mixed = two_model_collection["light"] + owned = list(mixed.parameters()) + assert len(owned) == 1 + assert owned[0] is mixed.fraction_params + + @pytest.mark.unit + def test_shared_base_parameters_are_counted_once(self, two_model_collection): + """Two timepoints over two shared models: two base parameters plus one set of + fractions each. Double-registration would make the optimizer step a base model + once per timepoint. + """ + mc = two_model_collection + params = list(mc.parameters()) + assert len(params) == 2 + 2 + + base_anchors = [m.anchor for m in mc.base_models] + for anchor in base_anchors: + assert sum(p is anchor for p in params) == 1 + + +class TestForwardAndGradient: + @pytest.mark.unit + def test_forward_is_the_fraction_weighted_sum_of_the_parts( + self, two_model_collection, hkl + ): + mixed = two_model_collection["light"] + parts = mixed.get_individual_fcalc(hkl, recalc=True) + w = mixed.fractions + + expected = w[0] * parts[0] + w[1] * parts[1] + assert torch.allclose(mixed(hkl, recalc=True), expected) + + @pytest.mark.unit + def test_gradient_reaches_the_fraction_params(self, two_model_collection, hkl): + mixed = two_model_collection["light"] + mixed(hkl, recalc=True).abs().sum().backward() + + grad = mixed.fraction_params.grad + assert grad is not None + assert torch.isfinite(grad).all() + + @pytest.mark.unit + def test_frozen_dark_fractions_receive_no_gradient( + self, two_model_collection, hkl + ): + dark = two_model_collection["dark"] + dark(hkl, recalc=True).abs().sum().backward() + assert dark.fraction_params.grad is None diff --git a/tests/unit/refinement/test_collection_target_characterisation.py b/tests/unit/refinement/test_collection_target_characterisation.py new file mode 100644 index 00000000..af5c5f00 --- /dev/null +++ b/tests/unit/refinement/test_collection_target_characterisation.py @@ -0,0 +1,276 @@ +"""Characterisation of the collection X-ray targets' observable contract. + +Written to protect a change of fraction storage and a move to batched +``[T, R]`` accessors. The three regressions worth catching are all silent: + +* reading **raw** ``ReflectionData.F`` instead of the scaled ``get_corrected_data()``, + which drops the inter-dataset scaling; +* masking with the 2-way ``rfree_flags`` instead of the 3-way ``work``/``free``/ + ``validation`` subset, which lets validation reflections back into the loss; +* returning a **mean** where the target returns a **sum**, which reweights the X-ray + term by 1/N against every restraint. + +Each is pinned by a *deterministic invariant* rather than a stored number. +``model.forward`` is run-to-run nondeterministic even inside one process (threaded +reduction order; ~4e-3 absolute on individual ``F_calc``), so "the loss equals 3.6e4" +is a weaker statement than "the loss responds to this input the way only a correct +implementation can". Measured for reference: the summed losses here vary by ~4e-7 +relative across repeated calls, so the literal checks that remain are given a +tolerance three orders of magnitude above that. + +The fixture deliberately makes the two datasets and the two models **differ**. With one +``ReflectionData`` added twice -- as the sigma_A collection fixture does -- the observed +difference is identically zero and a raw-vs-scaled regression is invisible. +""" + +import pytest +import torch + +# Repeated-call spread of the summed losses, measured on this fixture. The literal +# assertions below sit far above it; tightening past ~1e-6 would flake. +LOSS_RTOL = 1e-4 + + +@pytest.fixture(scope="module") +def collection(pdb_dir, mtz_dir): + """``(dc, mc, scaler)`` for a dark/light pair with a real difference in both + the data and the models. + + The light dataset carries its own ``log_scale``, so ``F_obs_light != F_obs_dark`` + only through the *corrected* accessor -- which is what makes the raw-vs-scaled + invariant below bite. The light model is displaced, so ``ΔF_calc != 0`` too. + """ + pdb = pdb_dir / "1DAW.pdb" + mtz = mtz_dir / "1DAW.mtz" + if not (pdb.exists() and mtz.exists()): + pytest.skip("1DAW fixture not present") + + from torchref import ReflectionData + from torchref.cli._common import load_model + from torchref.io.datasets.collection import DatasetCollection + from torchref.model.model_collection import ModelCollection + from torchref.scaling.collection_scaler import CollectionScaler + + data_dark = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + data_light = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + + # max_res is required: without it the FFT grid setup has no resolution to size from. + d_min = 2.05 + model_dark = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) + model_light = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) + with torch.no_grad(): + # A real displacement, so the calculated difference is not degenerate. + xyz = model_light.xyz.refinable_params + xyz += 0.15 * torch.ones_like(xyz) + + dc = DatasetCollection(verbose=0, device="cpu") + dc.add_dataset("dark", data_dark, set_as_reference=True) + dc.add_dataset("light", data_light) + + mc = ModelCollection([model_dark, model_light], dark_key="dark", verbose=0) + mc.add_dark() + mc.add_timepoint("light", [0.7, 0.3]) + + scaler = CollectionScaler(dc, mc, verbose=0) + scaler.initialize() + return dc, mc, scaler + + +def _targets(dc, mc, scaler): + from torchref.refinement.targets import ( + CollectionDifferenceTarget, + CollectionMLTarget, + CollectionRiceTarget, + ) + + return { + "difference": CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0), + "rice": CollectionRiceTarget(dc, mc, scaler=scaler, verbose=0), + "ml": CollectionMLTarget(dc, mc, scaler=scaler, verbose=0), + } + + +@pytest.mark.integration +class TestObservedAmplitudesAreScaled: + """The loss must move when a dataset's own scale moves. + + ``DatasetCollection.scale()`` fits a per-dataset ``log_scale``/``U_aniso`` that + exists only in ``get_corrected_data()``. A target reading raw ``.F`` is completely + blind to it, so this is a direct test of which accessor is in use. + """ + + @pytest.mark.parametrize("name", ["difference", "rice", "ml"]) + def test_loss_responds_to_the_datasets_own_log_scale(self, collection, name): + dc, mc, scaler = collection + target = _targets(dc, mc, scaler)[name] + + before = target.forward().item() + light = dc["light"] + original = light.log_scale.detach().clone() + try: + with torch.no_grad(): + light.log_scale += 0.25 # ~28% on amplitudes + target.maintenance() if hasattr(target, "maintenance") else None + after = target.forward().item() + finally: + with torch.no_grad(): + light.log_scale.copy_(original) + + rel = abs(after - before) / abs(before) + assert rel > 1e-3, ( + f"{name}: changing the light dataset's log_scale moved the loss by only " + f"{rel:.2e}; the target is reading raw amplitudes, not the scaled ones" + ) + + def test_corrected_and_raw_amplitudes_actually_differ(self, collection): + """Anti-vacuity: the invariant above is only meaningful if the two accessors + disagree on this fixture.""" + dc, _, _ = collection + light = dc["light"] + with torch.no_grad(): + light.log_scale += 0.25 + corrected, _ = light.get_corrected_data() + raw = light.F + differ = not torch.allclose(corrected, raw) + light.log_scale -= 0.25 + assert differ + + +@pytest.mark.integration +class TestSubsetSelectionIsThreeWay: + """Work / free / validation, with validation carved out of both.""" + + def test_subsets_are_disjoint_and_cover_the_valid_reflections(self, collection): + dc, mc, scaler = collection + target = _targets(dc, mc, scaler)["difference"] + data = dc["dark"] + + work = data.work.mask + free = data.free.mask + val = data.validation.mask + + assert not (work & free).any() + assert not (work & val).any() + assert not (free & val).any() + assert torch.equal(work | free | val, data.masks().to(torch.bool)) + assert target.use_set == "work" + + def test_carving_a_validation_set_shrinks_the_work_and_free_sets(self, collection): + """A 2-way ``rfree_flags`` implementation cannot see a validation set at all, + so the reported ``n`` would not move. + """ + dc, mc, scaler = collection + target = _targets(dc, mc, scaler)["difference"] + data = dc["dark"] + + n_before = target._n_reflections() + free_before = data.free.n + flags = None if data.validation_flags is None else data.validation_flags.clone() + try: + data.generate_validation_set(val_fraction_of_free=0.5, seed=0) + assert data.validation.n > 0, "no validation reflections were carved" + assert data.free.n < free_before, "free set did not shrink" + assert target._n_reflections() <= n_before + finally: + data.validation_flags = flags + data._subset_fp = None + + @pytest.mark.parametrize("use_set", ["work", "free"]) + def test_loss_is_restricted_to_the_selected_subset(self, collection, use_set): + """Work and free are different sizes here, so a target that ignored + ``use_set`` would return the same number for both.""" + from torchref.refinement.targets import CollectionDifferenceTarget + + dc, mc, scaler = collection + target = CollectionDifferenceTarget( + dc, mc, scaler=scaler, use_set=use_set, verbose=0 + ) + assert target.use_set == use_set + n = target._n_reflections() + expected = sum( + (dc[k].work if use_set == "work" else dc[k].free).n + for k in target._keys() + ) + assert n == expected + + +@pytest.mark.integration +class TestLossesAreSummedNotAveraged: + """A summed X-ray term grows with the data; a meaned one does not. + + This is the invariant that catches a 1/N reweight, which is otherwise invisible + -- it looks exactly like a change of X-ray weight. + """ + + def test_adding_a_dataset_grows_the_rice_loss(self, collection, pdb_dir, mtz_dir): + from torchref import ReflectionData + from torchref.refinement.targets import CollectionRiceTarget + + dc, mc, scaler = collection + one = CollectionRiceTarget(dc, mc, scaler=scaler, verbose=0).forward().item() + + extra = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz_dir / "1DAW.mtz")) + dc.add_dataset("light2", extra) + mc.add_timepoint("light2", [0.7, 0.3]) + try: + two = CollectionRiceTarget(dc, mc, scaler=scaler, verbose=0).forward().item() + finally: + dc._datasets.pop("light2") + dc._dataset_order.remove("light2") + del mc._timepoints["light2"] + mc._order.remove("light2") + + ratio = two / one + assert ratio == pytest.approx(2.0, rel=0.15), ( + f"two timepoints gave {ratio:.3f}x one timepoint's Rice loss; a summed " + f"target should roughly double and a meaned one stay near 1.0" + ) + + +@pytest.mark.integration +class TestReportedNumbers: + """The shape of what ``get_rfactor`` / ``stats`` promise, plus reproducibility.""" + + @pytest.mark.parametrize("name", ["difference", "rice", "ml"]) + def test_forward_is_finite_and_reproducible(self, collection, name): + dc, mc, scaler = collection + target = _targets(dc, mc, scaler)[name] + first = target.forward().item() + second = target.forward().item() + assert torch.isfinite(torch.tensor(first)) + assert second == pytest.approx(first, rel=LOSS_RTOL) + + def test_rfactor_shape_and_range(self, collection): + dc, mc, scaler = collection + target = _targets(dc, mc, scaler)["difference"] + + rf = target.get_rfactor() + assert set(rf) == {"per_dataset", "rwork_pct", "rfree_pct"} + assert set(rf["per_dataset"]) == {"dark", "light"} + for key, (rwork, rfree) in rf["per_dataset"].items(): + assert 0.0 < rwork < 1.0, f"{key} rwork out of range: {rwork}" + assert 0.0 < rfree < 1.0, f"{key} rfree out of range: {rfree}" + assert set(rf["rwork_pct"]) == {"p10", "p25", "p50", "p75", "p90"} + + def test_stats_reports_the_median_of_the_per_dataset_distribution(self, collection): + dc, mc, scaler = collection + target = _targets(dc, mc, scaler)["difference"] + + stats = target.stats() + rf = target.get_rfactor() + for key in ("loss", "n", "rwork", "rfree"): + assert key in stats, f"missing stat: {key}" + assert stats["rwork"].value == pytest.approx( + rf["rwork_pct"]["p50"], rel=LOSS_RTOL + ) + assert stats["n"].value == target._n_reflections() + + def test_gradient_reaches_the_light_model(self, collection): + dc, mc, scaler = collection + target = _targets(dc, mc, scaler)["difference"] + + target.forward().backward() + grad = mc.base_models[1].xyz.refinable_params.grad + assert grad is not None + assert torch.isfinite(grad).all() + assert grad.abs().max() > 0 From ae39a17e1a7035e79f4630b7a9d5137120b76679 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 24 Aug 2026 13:14:06 +0200 Subject: [PATCH 045/250] Make f_sol_override a pure argument and keep the batch rank Two defects in ScalerBase.forward, both reachable from any caller that scales several fraction mixtures against one shared scaler: - the override was assigned to _f_sol_raw, so a later call that passed no override silently read the previous caller's mixed solvent; - f_sol was broadcast with an unconditional unsqueeze(0), so a batched (T, N) override returned (1, T, N) and the batch rank was never restored. The override is now consumed locally and the reflection axis is indexed as the last one, which works for both (N,) and (T, N). Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- docs/changelog.rst | 2 + .../scaling/test_f_sol_override_contract.py | 174 ++++++++++++++++++ torchref/scaling/scaler_base.py | 35 ++-- 3 files changed, 198 insertions(+), 13 deletions(-) create mode 100644 tests/unit/scaling/test_f_sol_override_contract.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 30fcbb07..607dfded 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,8 @@ Changelog Version 0.6.4 ---------- +- Fixed ``f_sol_override`` overwriting the scaler's cached ``F_sol``, so a later call without an override read the wrong solvent +- Fixed a batched ``f_sol_override`` gaining a spurious leading axis, which changed the rank of the scaled structure factors - Fixed the bulk-solvent ``F_sol`` staying at the starting model's mask for every refinement macrocycle - Fixed restraint dictionaries defining several compounds yielding restraints for only one of them - Fixed chirality restraints being dropped for the ``positiv``/``negativ`` spellings used by the CCP4 library diff --git a/tests/unit/scaling/test_f_sol_override_contract.py b/tests/unit/scaling/test_f_sol_override_contract.py new file mode 100644 index 00000000..b7e71bfe --- /dev/null +++ b/tests/unit/scaling/test_f_sol_override_contract.py @@ -0,0 +1,174 @@ +"""``f_sol_override`` must be a pure argument, and must not change the output rank. + +Two separate contracts on :meth:`ScalerBase.forward`, both load-bearing for any caller +that scales several models against one shared scaler: + +1. Passing ``f_sol_override`` must not write ``_f_sol_raw``. A caller that scales two + different fraction mixtures in a row otherwise leaves the *second* mixture's solvent + cached, and every later call that does not pass an override silently reads it. +2. A batched ``(T, N)`` override paired with a batched ``(T, N)`` ``fcalc`` must return + ``(T, N)``. The solvent term is broadcast with ``unsqueeze(0)``, which is right for a + per-reflection ``(N,)`` solvent and wrong for one that already carries the batch axis. + +Both are exercised through ``forward`` rather than asserted on internals, so the stub only +has to stand in for the solvent model. +""" + +from types import SimpleNamespace + +import pytest +import torch + + +class _StubSolvent: + """Minimal stand-in for :class:`SolventModel` on the k_sol/B_sol path. + + ``get_rec_solvent`` returns a recognisable constant so a cached value can be told + apart from a freshly-passed override by value as well as by identity. + """ + + optimize_phase = False + + def __init__(self, value: float = 1.0): + self.value = value + self.n_reads = 0 + + def k_solvent(self): + return torch.tensor(0.35) + + def damping(self, s_half_sq): + return torch.ones_like(s_half_sq) + + def get_rec_solvent(self, hkl): + self.n_reads += 1 + return torch.full( + (hkl.shape[0],), self.value, dtype=torch.complex64 + ) + + +@pytest.fixture +def scaler_with_stub(): + """A bare ``Scaler`` with only the solvent branch live. + + ``bins`` drives the full-size check in ``forward``; the anisotropy, Chebyshev and + per-bin-B branches are all absent, so ``forward`` reduces to + ``fcalc + k_sol * f_sol`` and any shape or caching defect is unobscured. + """ + from torchref.scaling.scaler import Scaler + + n = 6 + scaler = Scaler() + dev = scaler.device + scaler.bins = torch.zeros(n, dtype=torch.long, device=dev) + scaler._s_half_sq = torch.zeros(n, device=dev) + # The no-override path reads the solvent model at self.hkl, which is a read-only + # property over self._data. + scaler._data = SimpleNamespace( + hkl=torch.zeros((n, 3), dtype=torch.long, device=dev) + ) + scaler.solvent = _StubSolvent() + scaler._f_sol_raw = None + return scaler, n, dev + + +class TestOverrideDoesNotMutateTheCache: + @pytest.mark.unit + def test_override_leaves_the_cache_untouched(self, scaler_with_stub): + """The override is an argument, not an assignment.""" + scaler, n, dev = scaler_with_stub + override = torch.full((n,), 2.0, dtype=torch.complex64, device=dev) + + scaler.forward(torch.ones(n, dtype=torch.complex64, device=dev), + f_sol_override=override) + + assert scaler._f_sol_raw is not override, ( + "forward stored the override in the solvent cache" + ) + assert scaler._f_sol_raw is None, ( + "forward populated the solvent cache from an override; a later call " + "without one will read this instead of the model's own solvent" + ) + + @pytest.mark.unit + def test_a_later_call_without_an_override_sees_the_model_solvent( + self, scaler_with_stub + ): + """The consequence of the leak, stated in terms a caller can observe. + + Two mixtures scaled in a row, then a plain call: the plain call must use the + solvent model, not whichever mixture happened to be scaled last. + """ + scaler, n, dev = scaler_with_stub + fcalc = torch.ones(n, dtype=torch.complex64, device=dev) + + baseline = scaler.forward(fcalc).clone() # stub solvent, value 1.0 + scaler._f_sol_raw = None # as update_solvent() would leave it + + far_off = torch.full((n,), 99.0, dtype=torch.complex64, device=dev) + scaler.forward(fcalc, f_sol_override=far_off) + after = scaler.forward(fcalc) + + assert torch.allclose(after, baseline), ( + "a plain forward() after an override call returned the override's " + "solvent contribution" + ) + + +class TestOverridePreservesRank: + @pytest.mark.unit + def test_batched_override_with_batched_fcalc_keeps_the_batch_rank( + self, scaler_with_stub + ): + """``(T, N)`` in, ``(T, N)`` out -- the batched multi-dataset contract.""" + scaler, n, dev = scaler_with_stub + t = 3 + fcalc = torch.ones((t, n), dtype=torch.complex64, device=dev) + override = torch.full((t, n), 2.0, dtype=torch.complex64, device=dev) + + out = scaler.forward(fcalc, f_sol_override=override) + + assert out.shape == (t, n), ( + f"batched override changed the output rank: got {tuple(out.shape)}, " + f"expected {(t, n)}" + ) + + @pytest.mark.unit + def test_each_batch_row_matches_the_unbatched_call(self, scaler_with_stub): + """Rank alone is not enough -- the rows must also be the right ones. + + A wrong broadcast can restore the shape and still pair row *i* of ``fcalc`` + with the wrong row of the solvent. + """ + scaler, n, dev = scaler_with_stub + t = 3 + fcalc = torch.stack([ + torch.full((n,), float(i + 1), dtype=torch.complex64, device=dev) + for i in range(t) + ]) + override = torch.stack([ + torch.full((n,), float(10 * (i + 1)), dtype=torch.complex64, device=dev) + for i in range(t) + ]) + + batched = scaler.forward(fcalc, f_sol_override=override) + for i in range(t): + scaler._f_sol_raw = None + single = scaler.forward(fcalc[i], f_sol_override=override[i]) + assert torch.allclose(batched[i], single), ( + f"batched row {i} does not match the equivalent unbatched call" + ) + + @pytest.mark.unit + def test_unbatched_override_is_unchanged(self, scaler_with_stub): + """The ``(N,)`` solvent path must keep working; it is what every single-dataset + caller uses.""" + scaler, n, dev = scaler_with_stub + fcalc = torch.ones(n, dtype=torch.complex64, device=dev) + override = torch.full((n,), 2.0, dtype=torch.complex64, device=dev) + + out = scaler.forward(fcalc, f_sol_override=override) + + assert out.shape == (n,) + # fcalc + k_sol * f_sol, with damping == 1 and no aniso/Chebyshev/per-bin B. + expected = 1.0 + 0.35 * 2.0 + assert torch.allclose(out.real, torch.full((n,), expected, device=dev)) diff --git a/torchref/scaling/scaler_base.py b/torchref/scaling/scaler_base.py index 9fd9119f..a8c836b5 100644 --- a/torchref/scaling/scaler_base.py +++ b/torchref/scaling/scaler_base.py @@ -782,9 +782,10 @@ def forward( Deprecated and inert -- never read. Masking follows the input shape, so ``use_mask=False`` does *not* disable it. f_sol_override : torch.Tensor, optional - Raw solvent structure factors replacing the cached ``_f_sol_raw`` (k_sol / B_sol - / phase damping still applied). **Overwrites the cache**, so it persists into - later calls until invalidated. Used by ``CollectionScaler``. + Raw solvent structure factors used instead of the cached ``_f_sol_raw`` for this + call only (k_sol / B_sol / phase damping still applied); the cache is left + untouched. Shape ``(N,)`` or ``(B, N)`` -- a batched override keeps the batch + axis of the result. Used by ``CollectionScaler``. Returns ------- @@ -814,20 +815,25 @@ def forward( else: aniso_correction = torch.tensor(1.0, device=self.device, dtype=fcalc.dtype) - if f_sol_override is not None: - self._f_sol_raw = f_sol_override + # An override is consumed locally and never displaces the cache: it may carry a + # leading batch axis, and it belongs to one caller's fraction mixture rather than + # to this scaler's solvent model. + f_sol_raw_local = f_sol_override if hasattr(self, "solvent") and self.solvent is not None: # Lazily cache raw solvent SFs (FFT of mask) — only recomputed # when invalidated via _f_sol_raw = None (e.g. after update_solvent) - if self._f_sol_raw is None: - # The solvent mask is real density with no anomalous term, so - # F_sol(-h) is exactly conj(F_sol(h)) and evaluating on the - # canonical index already matches the canonical fcalc below. - self._f_sol_raw = self.solvent.get_rec_solvent(self.hkl) - + if f_sol_raw_local is None: + if self._f_sol_raw is None: + # The solvent mask is real density with no anomalous term, so + # F_sol(-h) is exactly conj(F_sol(h)) and evaluating on the + # canonical index already matches the canonical fcalc below. + self._f_sol_raw = self.solvent.get_rec_solvent(self.hkl) + f_sol_raw_local = self._f_sol_raw + + # Index the reflection axis, which is last for both (N,) and (B, N). f_sol_raw = ( - self._f_sol_raw[mask] if apply_internal_mask else self._f_sol_raw + f_sol_raw_local[..., mask] if apply_internal_mask else f_sol_raw_local ) if hasattr(self, "log_kmask"): @@ -869,10 +875,13 @@ def forward( else: b_overall = torch.tensor(1.0, device=self.device, dtype=fcalc.dtype) + # f_sol already carries the batch axis when it came from a batched override; + # only a per-reflection (N,) solvent needs one added to broadcast. + f_sol_expanded = f_sol if f_sol.ndim >= 2 else f_sol.unsqueeze(0) fcalc = ( K_overall.unsqueeze(0) * b_overall.unsqueeze(0) - * (aniso_correction.unsqueeze(0) * fcalc + f_sol.unsqueeze(0)) + * (aniso_correction.unsqueeze(0) * fcalc + f_sol_expanded) ) if not batched: From 3a331e80b984963dc831be13ebbeb27499432adf Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 24 Aug 2026 13:30:19 +0200 Subject: [PATCH 046/250] Add batched per-component structure factors compute_component_fcalcs evaluates each shared base model once; mix_component_fcalcs contracts that stack with any weight matrix, so several mixtures of the same models cost one set of structure factors rather than one per mixture. DatasetCollection.component_structure_factors is the supported entry point: it evaluates at the signed indices and returns the result on the canonical index, the same convention ReflectionData.structure_factors owns. Calling the model-side method on data.hkl skips both halves. Every reflection file under tests/files is already inside the CCP4 ASU, so the convention tests manufacture flagged rows by negating half the Miller indices and assert that precondition. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- .../model/test_batched_component_fcalcs.py | 199 ++++++++++++++++++ torchref/io/datasets/collection.py | 47 +++++ torchref/model/model_collection.py | 86 ++++++++ 3 files changed, 332 insertions(+) create mode 100644 tests/unit/model/test_batched_component_fcalcs.py diff --git a/tests/unit/model/test_batched_component_fcalcs.py b/tests/unit/model/test_batched_component_fcalcs.py new file mode 100644 index 00000000..a4ed9963 --- /dev/null +++ b/tests/unit/model/test_batched_component_fcalcs.py @@ -0,0 +1,199 @@ +"""The batched component/mixture structure factors must equal the per-timepoint loop. + +``compute_component_fcalcs`` + ``mix_component_fcalcs`` exist to evaluate each shared base +model once instead of once per timepoint. That is only a saving if it produces the same +numbers as the loop it replaces, including the signed-index and Friedel bookkeeping that +:meth:`ReflectionData.structure_factors` owns. + +Both sides are built from a single set of model forwards (one ``recalc=True``, then cache +hits) because repeated structure-factor evaluation is not bit-reproducible: two +``recalc=True`` calls on identical input differ by ~4e-3 on individual ``F_calc``. Comparing +two independent evaluations would measure that noise instead of the contraction. +""" + +import pytest +import torch + + +@pytest.fixture(scope="module") +def pair(pdb_dir, mtz_dir): + """A 2-component, 2-timepoint collection on 1DAW with unequal fractions.""" + pdb = pdb_dir / "1DAW.pdb" + mtz = mtz_dir / "1DAW.mtz" + if not (pdb.exists() and mtz.exists()): + pytest.skip("1DAW fixture not present") + + from torchref import ReflectionData + from torchref.cli._common import load_model + from torchref.io.datasets.collection import DatasetCollection + from torchref.model.model_collection import ModelCollection + + d_min = 2.05 + data = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + model_a = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) + model_b = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) + with torch.no_grad(): + model_b.xyz.refinable_params += 0.2 + + dc = DatasetCollection(verbose=0, device="cpu") + dc.add_dataset("dark", data, set_as_reference=True) + + mc = ModelCollection([model_a, model_b], dark_key="dark", verbose=0) + mc.add_dark() + mc.add_timepoint("light", [0.65, 0.35]) + return dc, mc + + +@pytest.mark.integration +class TestBatchedMatchesTheLoop: + def test_component_stack_matches_per_model_structure_factors(self, pair): + dc, mc = pair + data = dc["dark"] + + stacked = dc.component_structure_factors(mc, recalc=True) + assert stacked.shape == (mc.n_base_models, len(data.hkl)) + + for k, model in enumerate(mc.base_models): + reference = data.structure_factors(model, recalc=False) + assert torch.equal(stacked[k], reference), ( + f"component {k} differs from data.structure_factors" + ) + + def test_mixture_matches_the_per_timepoint_forward(self, pair): + """The whole point: one contraction standing in for T mixed forwards.""" + dc, mc = pair + data = dc["dark"] + + stacked = dc.component_structure_factors(mc, recalc=True) + mixed = mc.mix_component_fcalcs(stacked, mc.get_fractions_matrix()) + assert mixed.shape == (len(mc), len(data.hkl)) + + for row, key in enumerate(mc.keys()): + reference = data.structure_factors(mc[key], recalc=False) + assert torch.allclose(mixed[row], reference, rtol=1e-6, atol=1e-6), ( + f"timepoint {key!r} differs from its own mixed forward" + ) + + def test_compute_all_fcalc_agrees_on_the_signed_index(self, pair): + """``compute_all_fcalc`` takes the caller's indices verbatim, so handed the + signed ones it must reproduce the Friedel-corrected mixture up to the + conjugation that ``component_structure_factors`` applies.""" + dc, mc = pair + data = dc["dark"] + + direct = mc.compute_all_fcalc(data._hkl_for_sf(), recalc=True) + corrected = data.conjugate_friedel(direct) + + stacked = dc.component_structure_factors(mc, recalc=False) + mixed = mc.mix_component_fcalcs(stacked, mc.get_fractions_matrix()) + + assert torch.allclose(corrected, mixed, rtol=1e-6, atol=1e-6) + + +@pytest.fixture(scope="module") +def flagged_pair(pair): + """The same models against data with **manufactured** Friedel-flagged rows. + + Every reflection file under ``tests/files/`` is already inside the CCP4 ASU, so + ``friedel_flags.any()`` is False on all of them and any assertion about the index + convention is silently vacuous. Negating half the Miller indices forces + canonicalisation to flip them back, reproducing the ~50% flagged fraction real + P1 data carries. Cell and space group are preserved, so the same models apply. + """ + from torchref import ReflectionData + from torchref.io.datasets.collection import DatasetCollection + + dc_ref, mc = pair + src = dc_ref["dark"] + + hkl = src.hkl.clone() + half = torch.zeros(len(hkl), dtype=torch.bool) + half[::2] = True + hkl[half] = -hkl[half] + + data = ReflectionData.from_tensors( + hkl=hkl, + F=src.F.clone(), + F_sigma=src.F_sigma.clone(), + cell=src.cell, + spacegroup=src.spacegroup, + rfree_flags=src.rfree_flags.clone(), + device="cpu", + verbose=0, + ) + + dc = DatasetCollection(verbose=0, device="cpu") + dc.add_dataset("dark", data, set_as_reference=True) + return dc, mc + + +@pytest.mark.integration +class TestConventionIsNotSkipped: + def test_the_fixture_actually_has_flagged_rows(self, flagged_pair): + """The precondition, asserted rather than assumed.""" + data = flagged_pair[0]["dark"] + assert data.friedel_flags is not None + frac = data.friedel_flags.float().mean().item() + assert 0.2 < frac < 0.8, f"expected a mixed flag population, got {frac:.3f}" + + def test_conjugation_moves_phases_and_leaves_amplitudes(self, flagged_pair): + dc, mc = flagged_pair + data = dc["dark"] + + raw = mc.compute_component_fcalcs(data._hkl_for_sf(), recalc=True) + corrected = data.conjugate_friedel(raw) + + assert torch.allclose(raw.abs(), corrected.abs()) + assert not torch.allclose(torch.angle(raw), torch.angle(corrected)) + + def test_component_stack_is_conjugated_where_flagged(self, flagged_pair): + """``component_structure_factors`` must apply the conjugation, not skip it. + + Compared against the per-model supported entry point, which is the definition + of the convention. + """ + dc, mc = flagged_pair + data = dc["dark"] + + stacked = dc.component_structure_factors(mc, recalc=True) + for k, model in enumerate(mc.base_models): + reference = data.structure_factors(model, recalc=False) + assert torch.equal(stacked[k], reference), f"component {k} phases differ" + + def test_skipping_the_conjugation_would_be_detected(self, flagged_pair): + """Anti-vacuity: the naive call this method exists to replace disagrees.""" + dc, mc = flagged_pair + data = dc["dark"] + + correct = dc.component_structure_factors(mc, recalc=True) + naive = mc.compute_component_fcalcs(data.hkl, recalc=True) + + assert not torch.allclose(correct, naive), ( + "evaluating on the canonical index gives the same answer as the signed " + "index plus conjugation -- this fixture cannot detect a convention bug" + ) + + +@pytest.mark.integration +class TestContraction: + def test_weights_matrix_is_applied_row_wise(self, pair): + """A transposed einsum would still return the right shape when T == K.""" + dc, mc = pair + stacked = dc.component_structure_factors(mc, recalc=True) + + w = torch.tensor([[1.0, 0.0], [0.0, 1.0]]) + mixed = mc.mix_component_fcalcs(stacked, w) + + assert torch.equal(mixed[0], stacked[0]) + assert torch.equal(mixed[1], stacked[1]) + + def test_gradient_flows_through_the_contraction(self, pair): + dc, mc = pair + stacked = dc.component_structure_factors(mc, recalc=True) + w = mc.get_fractions_matrix() + + mc.mix_component_fcalcs(stacked, w).abs().sum().backward() + + grad = mc["light"].fraction_params.grad + assert grad is not None + assert torch.isfinite(grad).all() diff --git a/torchref/io/datasets/collection.py b/torchref/io/datasets/collection.py index 3e85c923..e97b7b15 100644 --- a/torchref/io/datasets/collection.py +++ b/torchref/io/datasets/collection.py @@ -314,6 +314,53 @@ def closure(): [p.requires_grad_(False) for p in parameters] + def component_structure_factors( + self, model_collection, recalc: bool = False + ) -> torch.Tensor: + """Per-base-model ``F_calc`` on the common HKL, in the canonical convention. + + The batched counterpart of :meth:`ReflectionData.structure_factors`: models are + evaluated at the **signed** indices so Bijvoet mates get distinct ``|F_calc|``, + and the result is returned on the canonical ASU index that :attr:`hkl` holds. + Use this rather than calling + :meth:`~torchref.model.model_collection.ModelCollection.compute_component_fcalcs` + on :attr:`hkl` directly, which would skip both halves of that convention. + + The convention is taken from the reference dataset. Members are all expanded onto + one HKL grid, but a member with different completeness can still carry different + ``friedel_flags`` (absent rows are filled ``False``); where they differ, the + returned **phases** follow the reference. Amplitudes are unaffected, so a target + working in moduli or intensities is insensitive to this. + + Parameters + ---------- + model_collection : ModelCollection + Supplies the shared base models. + recalc : bool, optional + Force recomputation rather than reusing each model's cached SF. + + Returns + ------- + torch.Tensor + Complex SFs of shape ``(n_base_models, n_reflections)``, row-aligned with + :attr:`hkl` on the reflection axis. + + Raises + ------ + ValueError + If the collection has no reference dataset. + """ + if self._reference_dataset is None: + raise ValueError( + "No reference dataset set; add a dataset before computing " + "component structure factors." + ) + ref = self._datasets[self._reference_dataset] + stacked = model_collection.compute_component_fcalcs( + ref._hkl_for_sf(), recalc=recalc + ) + return ref.conjugate_friedel(stacked) + def keys(self) -> List[str]: """Return list of dataset names.""" return list(self._dataset_order) diff --git a/torchref/model/model_collection.py b/torchref/model/model_collection.py index d56d6c9a..d5c58445 100644 --- a/torchref/model/model_collection.py +++ b/torchref/model/model_collection.py @@ -562,6 +562,92 @@ def get_fractions_matrix(self) -> torch.Tensor: [self._timepoints[n].fractions for n in self._order], dim=0 ) + # ------------------------------------------------------------------ + # Batched structure factors + # ------------------------------------------------------------------ + + def compute_component_fcalcs( + self, hkl: torch.Tensor, recalc: bool = False + ) -> torch.Tensor: + """Per-base-model structure factors, stacked. + + Each base model is evaluated once, so a caller that needs several fraction + mixtures of the same models pays for the structure factors once rather than + once per mixture. + + Parameters + ---------- + hkl : torch.Tensor + Miller indices of shape (n_reflections, 3). These reach the models + unchanged, so pass the *signed* indices when Bijvoet mates must be + distinguished -- or go through + :meth:`~torchref.io.datasets.collection.DatasetCollection.component_structure_factors`, + which handles the convention. + recalc : bool, optional + Force recomputation rather than reusing each model's cached SF. + + Returns + ------- + torch.Tensor + Complex structure factors of shape ``(n_base_models, n_reflections)``. + """ + return torch.stack( + [m(hkl, recalc=recalc) for m in self._base_models], dim=0 + ) + + def mix_component_fcalcs( + self, component_fcalcs: torch.Tensor, weights: torch.Tensor + ) -> torch.Tensor: + """Contract stacked per-component SFs with a weight matrix. + + ``weights [T, K] @ component_fcalcs [K, R] -> [T, R]``. Separated from + :meth:`compute_component_fcalcs` because the same component stack is contracted + with more than one weight matrix -- the fractions themselves, and any derivative + of them with respect to a shared parameter. + + Parameters + ---------- + component_fcalcs : torch.Tensor + Complex SFs of shape ``(K, n_reflections)``. + weights : torch.Tensor + Real weights of shape ``(T, K)``. + + Returns + ------- + torch.Tensor + Complex SFs of shape ``(T, n_reflections)``. + """ + return torch.einsum( + "tk,kr->tr", weights.to(component_fcalcs.dtype), component_fcalcs + ) + + def compute_all_fcalc( + self, hkl: torch.Tensor, recalc: bool = False + ) -> torch.Tensor: + """Mixed ``F_calc`` for every timepoint at once. + + Equivalent to calling each timepoint's ``forward`` in turn, but evaluates each + shared base model once instead of once per timepoint. Rows follow + :meth:`get_fractions_matrix`, i.e. insertion order. + + Parameters + ---------- + hkl : torch.Tensor + Miller indices of shape (n_reflections, 3); see + :meth:`compute_component_fcalcs` on the index convention. + recalc : bool, optional + Force recomputation rather than reusing each model's cached SF. + + Returns + ------- + torch.Tensor + Complex SFs of shape ``(n_timepoints, n_reflections)``. + """ + component_fcalcs = self.compute_component_fcalcs(hkl, recalc=recalc) + return self.mix_component_fcalcs( + component_fcalcs, self.get_fractions_matrix() + ) + # ------------------------------------------------------------------ # Freeze / unfreeze helpers # ------------------------------------------------------------------ From ddd69d2be45a6c669167a85f21cd546d09c9057d Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 24 Aug 2026 13:30:28 +0200 Subject: [PATCH 047/250] Add batched scaling with a per-row fraction-weighted solvent forward_batched applies the one shared set of scale parameters to T mixtures at once, mixing each row's bulk solvent by that row's weights. Because ScalerBase.forward is affine in fcalc and the mixed solvent is linear in the weights, passing a derivative of the fractions in place of the fractions returns the derivative of the scaled structure factors. Tested with a secant rather than a finite difference: the mixture is exactly linear in the activation fraction, so there is no truncation term to tolerate. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- .../scaling/test_collection_scaler_batched.py | 183 ++++++++++++++++++ torchref/scaling/collection_scaler.py | 54 ++++++ 2 files changed, 237 insertions(+) create mode 100644 tests/unit/scaling/test_collection_scaler_batched.py diff --git a/tests/unit/scaling/test_collection_scaler_batched.py b/tests/unit/scaling/test_collection_scaler_batched.py new file mode 100644 index 00000000..66fa4b8a --- /dev/null +++ b/tests/unit/scaling/test_collection_scaler_batched.py @@ -0,0 +1,183 @@ +"""``forward_batched`` must agree with ``forward_mixed`` row by row, and stay affine. + +The batched form exists so ``T`` mixtures share one pass through the scale parameters. +Two properties make it usable: + +* **row agreement** -- row ``i`` of the batch must equal the unbatched call on row ``i``, + otherwise the saving is bought with wrong numbers; +* **affinity in the mixing weights** -- ``ScalerBase.forward`` is + ``K * b * (aniso * F_calc + f_sol)`` and the mixed solvent is linear in the weights, so + scaling a *derivative* of the fractions returns the derivative of the scaled structure + factors. That is what lets a second moment be built from the same machinery instead of + a separate differentiation path. + +Affinity is tested with a **secant**, not a finite difference. Because the mixture is +exactly linear in the activation fraction, ``S(a1) - S(a2)`` equals +``(a1 - a2) * dS/da`` exactly, with no truncation term to tolerate -- so the assertion +is at float precision rather than ``O(h^2)``. +""" + +import pytest +import torch + + +@pytest.fixture(scope="module") +def scaled_collection(pdb_dir, mtz_dir): + """A dark/light collection on 1DAW with an initialized shared scaler.""" + pdb = pdb_dir / "1DAW.pdb" + mtz = mtz_dir / "1DAW.mtz" + if not (pdb.exists() and mtz.exists()): + pytest.skip("1DAW fixture not present") + + from torchref import ReflectionData + from torchref.cli._common import load_model + from torchref.io.datasets.collection import DatasetCollection + from torchref.model.model_collection import ModelCollection + from torchref.scaling.collection_scaler import CollectionScaler + + d_min = 2.05 + dark = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + light = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + + model_dark = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) + model_light = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) + with torch.no_grad(): + model_light.xyz.refinable_params += 0.2 + + dc = DatasetCollection(verbose=0, device="cpu") + dc.add_dataset("dark", dark, set_as_reference=True) + dc.add_dataset("light", light) + + mc = ModelCollection([model_dark, model_light], dark_key="dark", verbose=0) + mc.add_dark() + mc.add_timepoint("light", [0.78, 0.22]) + + scaler = CollectionScaler(dc, mc, verbose=0) + scaler.initialize() + return dc, mc, scaler + + +def _weights(alpha: float) -> torch.Tensor: + """Single-row activation weights ``[[1 - a, a]]``.""" + return torch.tensor([[1.0 - alpha, alpha]]) + + +@pytest.mark.integration +class TestBatchedMatchesUnbatched: + def test_every_row_matches_forward_mixed(self, scaled_collection): + dc, mc, scaler = scaled_collection + components = dc.component_structure_factors(mc, recalc=True) + + w = torch.tensor([[1.0, 0.0], [0.78, 0.22], [0.3, 0.7]]) + fcalc = mc.mix_component_fcalcs(components, w) + + batched = scaler.forward_batched(fcalc, w) + assert batched.shape == fcalc.shape + + for i in range(w.shape[0]): + single = scaler.forward_mixed(fcalc[i], w[i]) + assert torch.allclose(batched[i], single, rtol=1e-6, atol=1e-6), ( + f"batched row {i} disagrees with forward_mixed" + ) + + def test_component_solvent_stack_shape(self, scaled_collection): + dc, mc, scaler = scaled_collection + stack = scaler.compute_component_solvent_raw() + assert stack.shape == (mc.n_base_models, len(dc.hkl)) + assert stack.is_complex() + + def test_solvent_stack_rows_are_the_per_component_solvents( + self, scaled_collection + ): + """A transposed or misordered stack would still have the right shape.""" + _, mc, scaler = scaled_collection + stack = scaler.compute_component_solvent_raw() + for k in range(mc.n_base_models): + assert torch.equal(stack[k], scaler._get_component_f_sol_raw(k)) + + +@pytest.mark.integration +class TestAffineInTheMixingWeights: + def test_secant_in_alpha_equals_the_scaled_jacobian(self, scaled_collection): + """The property the two-moment forward model rests on. + + ``forward_batched(dF, J)`` with ``J = dW/da`` is the derivative of the scaled + mixture, including the solvent term. Exact, because everything between the + weights and the output is affine. + """ + dc, mc, scaler = scaled_collection + components = dc.component_structure_factors(mc, recalc=True) + + a1, a2 = 0.60, 0.10 + w1, w2 = _weights(a1), _weights(a2) + jac = torch.tensor([[-1.0, 1.0]]) # d/da of [1 - a, a] + + s1 = scaler.forward_batched(mc.mix_component_fcalcs(components, w1), w1) + s2 = scaler.forward_batched(mc.mix_component_fcalcs(components, w2), w2) + deriv = scaler.forward_batched( + mc.mix_component_fcalcs(components, jac), jac + ) + + secant = s1 - s2 + expected = (a1 - a2) * deriv + + rel = (secant - expected).abs().max() / expected.abs().max() + assert rel < 1e-5, ( + f"secant and scaled Jacobian disagree by {rel:.2e}; the scaler is not " + f"affine in the mixing weights, so a derivative cannot be scaled this way" + ) + + def test_the_solvent_term_is_included_in_the_derivative(self, scaled_collection): + """Anti-vacuity: if the per-component solvents were identical, the solvent + would cancel out of the Jacobian and the test above would hold even with the + solvent term dropped.""" + _, mc, scaler = scaled_collection + stack = scaler.compute_component_solvent_raw() + if mc.n_base_models < 2: + pytest.skip("needs at least two components") + assert not torch.allclose(stack[0], stack[1]), ( + "per-component solvents are identical, so this fixture cannot detect a " + "dropped solvent derivative" + ) + + def test_scaling_is_linear_in_the_structure_factors(self, scaled_collection): + """The other half of affinity: doubling F_calc at fixed weights doubles the + F_calc-dependent part, leaving the solvent offset behind.""" + dc, mc, scaler = scaled_collection + components = dc.component_structure_factors(mc, recalc=True) + w = _weights(0.22) + fcalc = mc.mix_component_fcalcs(components, w) + + s1 = scaler.forward_batched(fcalc, w) + s2 = scaler.forward_batched(2.0 * fcalc, w) + zero = scaler.forward_batched(torch.zeros_like(fcalc), w) + + # (S(2F) - S(0)) == 2 * (S(F) - S(0)) + lhs, rhs = s2 - zero, 2.0 * (s1 - zero) + rel = (lhs - rhs).abs().max() / rhs.abs().max() + assert rel < 1e-5 + + +@pytest.mark.integration +class TestSolventCacheIsNotPoisoned: + def test_batched_calls_leave_the_cache_alone(self, scaled_collection): + """Two batched calls with different weights, then a plain one. + + The Jacobian-weighted call carries negative weights, so a leaked cache would + show up as a sign error rather than a small perturbation. + """ + dc, mc, scaler = scaled_collection + components = dc.component_structure_factors(mc, recalc=True) + w = _weights(0.22) + jac = torch.tensor([[-1.0, 1.0]]) + fcalc = mc.mix_component_fcalcs(components, w) + + before = scaler.forward_mixed(fcalc[0], w[0]).clone() + + scaler.forward_batched(fcalc, w) + scaler.forward_batched(mc.mix_component_fcalcs(components, jac), jac) + + after = scaler.forward_mixed(fcalc[0], w[0]) + assert torch.allclose(after, before, rtol=1e-6, atol=1e-6), ( + "a batched call changed what a later forward_mixed returns" + ) diff --git a/torchref/scaling/collection_scaler.py b/torchref/scaling/collection_scaler.py index 585b4cb4..2384ea66 100644 --- a/torchref/scaling/collection_scaler.py +++ b/torchref/scaling/collection_scaler.py @@ -283,6 +283,60 @@ def forward_mixed( f_sol_raw_mixed = self.get_mixed_solvent_raw(fractions) return super().forward(fcalc, f_sol_override=f_sol_raw_mixed) + def compute_component_solvent_raw(self) -> torch.Tensor: + """Raw (un-damped) complex solvent SFs for every component, stacked. + + The solvent counterpart of + :meth:`~torchref.model.model_collection.ModelCollection.compute_component_fcalcs`. + Each component's mask FFT is cached, so repeated calls are cheap. + + Returns + ------- + torch.Tensor + Complex tensor of shape ``(n_components, n_reflections)``. + """ + return torch.stack( + [ + self._get_component_f_sol_raw(i) + for i in range(len(self._component_solvent_models)) + ], + dim=0, + ) + + def forward_batched( + self, + fcalc_batch: torch.Tensor, + fractions_matrix: torch.Tensor, + ) -> torch.Tensor: + """Scale a batch of mixtures, each with its own fraction-weighted solvent. + + The batched form of :meth:`forward_mixed`: one shared set of scale parameters + applied to ``T`` mixtures at once, with the bulk solvent mixed per row by the + same weights. Since ``ScalerBase.forward`` is affine in ``fcalc`` and the mixed + solvent is linear in the weights, passing a *derivative* of the fractions in + place of the fractions returns the corresponding derivative of the scaled + structure factors. + + Parameters + ---------- + fcalc_batch : torch.Tensor + Complex structure factors of shape ``(T, n_reflections)``. + fractions_matrix : torch.Tensor + Weights of shape ``(T, n_components)``, one row per member of the batch. + + Returns + ------- + torch.Tensor + Scaled complex structure factors of shape ``(T, n_reflections)``. + """ + component_sol_raw = self.compute_component_solvent_raw() + f_sol_batch = torch.einsum( + "tk,kr->tr", + fractions_matrix.to(component_sol_raw.dtype), + component_sol_raw, + ) + return super().forward(fcalc_batch, f_sol_override=f_sol_batch) + # ------------------------------------------------------------------ # Joint LBFGS refinement # ------------------------------------------------------------------ From c6e1d3f0815c40f132f53b2a5484afaf3f2fc569 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 24 Aug 2026 13:38:57 +0200 Subject: [PATCH 048/250] Add scaled intensity accessors get_corrected_intensities is the intensity counterpart of get_corrected_data. Both the anisotropy factor and the overall scale enter squared, because they are defined on amplitudes; applying the amplitude factors to intensities would leave a resolution-dependent error that looks like a scale or overall-B mismatch. The subset views gain the same corrected/raw split the amplitudes already have: .I and .sigI are scaled, .I_raw and .sigI_raw are not. Previously .sigI returned raw values while .F returned corrected ones, which would have made an amplitude target and an intensity target disagree about which dataset they were fitting. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- tests/unit/io/test_intensity_accessors.py | 201 ++++++++++++++++++++++ torchref/io/datasets/reflection_data.py | 100 +++++++++++ 2 files changed, 301 insertions(+) create mode 100644 tests/unit/io/test_intensity_accessors.py diff --git a/tests/unit/io/test_intensity_accessors.py b/tests/unit/io/test_intensity_accessors.py new file mode 100644 index 00000000..c45ab1a2 --- /dev/null +++ b/tests/unit/io/test_intensity_accessors.py @@ -0,0 +1,201 @@ +"""Scaled intensities must carry the amplitude scale **squared**. + +``get_corrected_data`` applies ``corr(s, U) * exp(log_scale)`` to amplitudes. The +intensity counterpart has to apply the square of that, because the correction is defined +on amplitudes. Getting it wrong leaves a smooth, resolution-dependent error in the +intensities that is indistinguishable from a scale or overall-B mismatch -- so it is +pinned here as an exact ratio rather than checked by eye. + +The subset views are also pinned: ``.F``/``.sigF`` are corrected and ``.F_raw``/``.sigF_raw`` +are not, and ``.I``/``.sigI`` now follow the same rule. A view where ``.F`` was scaled and +``.I`` was not is the shape of bug that makes an amplitude target and an intensity target +disagree about which dataset they are fitting. +""" + +import pytest +import torch + + +@pytest.fixture(scope="module") +def with_intensities(mtz_dir): + """1DAW -- the only fixture carrying both I/SIGI and FP/SIGFP.""" + mtz = mtz_dir / "1DAW.mtz" + if not mtz.exists(): + pytest.skip("1DAW fixture not present") + + from torchref import ReflectionData + + data = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + if data.I is None: + pytest.skip("1DAW loaded without intensities") + return data + + +@pytest.fixture(scope="module") +def without_intensities(mtz_dir): + """3GR5 -- amplitudes only.""" + mtz = mtz_dir / "3GR5.mtz" + if not mtz.exists(): + pytest.skip("3GR5 fixture not present") + + from torchref import ReflectionData + + data = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + if data.I is not None: + pytest.skip("3GR5 unexpectedly carries intensities") + return data + + +def _perturb(data, dlog=0.3): + """Give the dataset a non-trivial scale and anisotropy, restored on exit.""" + + class _Ctx: + def __enter__(self): + self.log_scale = data.log_scale.detach().clone() + self.U = data.U_aniso.detach().clone() + with torch.no_grad(): + data.log_scale += dlog + data.U_aniso += torch.tensor([0.01, -0.005, 0.008, 0.002, 0.0, 0.0]) + return data + + def __exit__(self, *exc): + with torch.no_grad(): + data.log_scale.copy_(self.log_scale) + data.U_aniso.copy_(self.U) + data._corrected_fp = None + data._corrected_I_fp = None + + return _Ctx() + + +@pytest.mark.unit +class TestSquaredScale: + def test_intensity_factor_is_the_square_of_the_amplitude_factor( + self, with_intensities + ): + """The exact relationship, as a per-reflection ratio. + + Independent of how I and F relate in the file (French-Wilson, not I == F**2), + because it compares each quantity against its own unscaled self. + """ + data = with_intensities + with _perturb(data): + F_scaled, _ = data.get_corrected_data() + I_scaled, _ = data.get_corrected_intensities() + + keep = data.masks().to(torch.bool) & (data.F.abs() > 1e-6) & ( + data.I.abs() > 1e-6 + ) + amp_factor = (F_scaled[keep] / data.F[keep]) ** 2 + int_factor = I_scaled[keep] / data.I[keep] + + rel = ((int_factor - amp_factor).abs() / amp_factor.abs()).max() + assert rel < 1e-5, ( + f"intensity scale factor is not the square of the amplitude one " + f"(max rel error {rel:.2e})" + ) + + def test_sigma_scales_with_the_same_factor_as_the_intensity( + self, with_intensities + ): + data = with_intensities + with _perturb(data): + I_scaled, sig_scaled = data.get_corrected_intensities() + keep = (data.I.abs() > 1e-6) & (data.I_sigma.abs() > 1e-6) + + ratio_I = I_scaled[keep] / data.I[keep] + ratio_s = sig_scaled[keep] / data.I_sigma[keep] + assert torch.allclose(ratio_I, ratio_s, rtol=1e-6) + + def test_the_perturbation_actually_changes_the_intensities( + self, with_intensities + ): + """Anti-vacuity: at log_scale 0 and U 0 every factor above is 1.""" + data = with_intensities + before, _ = data.get_corrected_intensities() + with _perturb(data): + after, _ = data.get_corrected_intensities() + assert not torch.allclose(before, after) + + def test_a_pure_scale_change_squares_into_the_intensities( + self, with_intensities + ): + """A doubling of the amplitude scale must quadruple the intensities.""" + data = with_intensities + base, _ = data.get_corrected_intensities() + original = data.log_scale.detach().clone() + try: + with torch.no_grad(): + data.log_scale += float(torch.log(torch.tensor(2.0))) + data._corrected_I_fp = None + doubled, _ = data.get_corrected_intensities() + finally: + with torch.no_grad(): + data.log_scale.copy_(original) + data._corrected_I_fp = None + + keep = base.abs() > 1e-6 + ratio = (doubled[keep] / base[keep]) + assert torch.allclose(ratio, torch.full_like(ratio, 4.0), rtol=1e-5) + + +@pytest.mark.unit +class TestSubsetViews: + def test_amplitudes_and_intensities_are_both_corrected(self, with_intensities): + data = with_intensities + with _perturb(data): + work = data.work + assert not torch.allclose(work.F, work.F_raw) + assert not torch.allclose(work.I, work.I_raw), ( + "subset.I returned raw intensities while subset.F was scaled" + ) + assert not torch.allclose(work.sigI, work.sigI_raw) + + def test_raw_views_match_the_parent_tensors(self, with_intensities): + data = with_intensities + work = data.work + idx = work.indices + assert torch.equal(work.I_raw, data.I.index_select(0, idx)) + assert torch.equal(work.sigI_raw, data.I_sigma.index_select(0, idx)) + + def test_subset_intensities_match_the_full_size_scaled_array( + self, with_intensities + ): + data = with_intensities + with _perturb(data): + I_scaled, sig_scaled = data.get_corrected_intensities() + for kind in ("work", "free"): + sub = getattr(data, kind) + idx = sub.indices + assert torch.equal(sub.I, I_scaled.index_select(0, idx)) + assert torch.equal(sub.sigI, sig_scaled.index_select(0, idx)) + + def test_cache_follows_a_scale_change(self, with_intensities): + """The fingerprint must invalidate, or a refinement would fit stale data.""" + data = with_intensities + first = data.work.I.clone() + original = data.log_scale.detach().clone() + try: + with torch.no_grad(): + data.log_scale += 0.5 + second = data.work.I + assert not torch.allclose(first, second) + finally: + with torch.no_grad(): + data.log_scale.copy_(original) + + +@pytest.mark.unit +class TestNoIntensities: + def test_get_corrected_intensities_raises_with_an_actionable_message( + self, without_intensities + ): + with pytest.raises(ValueError, match="no I/SIGI columns|No intensities"): + without_intensities.get_corrected_intensities() + + def test_subset_views_return_none_rather_than_raising(self, without_intensities): + work = without_intensities.work + assert work.I is None + assert work.sigI is None + assert work.I_raw is None + assert work.sigI_raw is None diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index fd62273b..07681e7f 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -124,8 +124,32 @@ def hkl(self) -> torch.Tensor: def rfree(self) -> torch.Tensor: return self._parent.rfree_flags.index_select(0, self.indices) + # -- intensities, corrected to match F/sigF above ----------------------- + @property + def I(self) -> torch.Tensor: # noqa: E743 - crystallographic name + """Scaled intensities, or None when this dataset carries no intensities. + + Corrected, like :attr:`F` -- both the anisotropy factor and the overall scale + enter squared. Use :attr:`I_raw` for the unscaled values. + """ + I_corr, _ = self._parent._corrected_or_raw_intensities() + return I_corr.index_select(0, self.indices) if I_corr is not None else None + @property def sigI(self): + """Scaled intensity sigmas, or None. See :attr:`I`.""" + _, sig_corr = self._parent._corrected_or_raw_intensities() + return sig_corr.index_select(0, self.indices) if sig_corr is not None else None + + @property + def I_raw(self): + """Unscaled intensities, or None.""" + i = self._parent.I + return i.index_select(0, self.indices) if i is not None else None + + @property + def sigI_raw(self): + """Unscaled intensity sigmas, or None.""" si = self._parent.I_sigma return si.index_select(0, self.indices) if si is not None else None @@ -228,6 +252,8 @@ def __post_init__(self): self._subset_fp = None self._corrected_cache = None self._corrected_fp = None + self._corrected_I_cache = None + self._corrected_I_fp = None # ===================== work / free / validation ===================== @@ -3947,6 +3973,80 @@ def get_corrected_data(self) -> Tuple[torch.Tensor, torch.Tensor]: return F_scaled, F_sigma_scaled + def get_corrected_intensities(self) -> Tuple[torch.Tensor, torch.Tensor]: + """ + Get the anisotropy-corrected, scaled ``(I, I_sigma)``. + + The intensity counterpart of :meth:`get_corrected_data`. Both the anisotropy + factor and the overall scale enter **squared**, because they are defined on + amplitudes: an amplitude scaled by ``corr * exp(log_scale)`` corresponds to an + intensity scaled by ``(corr * exp(log_scale))**2``. Applying the amplitude + factors to intensities instead would leave a resolution-dependent error that + looks exactly like a scale or B-factor mismatch. + + Returns + ------- + Tuple[torch.Tensor, torch.Tensor] + Full-size ``I`` and ``I_sigma`` on the same scale as + ``get_corrected_data()`` squared. + + Raises + ------ + ValueError + If this dataset carries no intensities (the input had no ``I``/``SIGI`` + columns), or if ``setup_scale`` / ``setup_anisotropy`` have not run. + """ + from torchref.base.alignment.normalization import ( + compute_anisotropy_correction, + ) + + if self.I is None: + raise ValueError( + "No intensities on this dataset. The input reflection file had no " + "I/SIGI columns, so only amplitudes are available; use " + "get_corrected_data() or supply intensity data." + ) + if not hasattr(self, "log_scale") or self.log_scale is None: + raise ValueError("Scale not set up. Call setup_scale() first.") + if not hasattr(self, "U_aniso") or self.U_aniso is None: + raise ValueError( + "No anisotropy parameters available. Call fit_anisotropy() first." + ) + + s_vectors = self.get_scattering_vectors() + correction = compute_anisotropy_correction(s_vectors, self.U_aniso) + factor = (correction * torch.exp(self.log_scale)) ** 2 + + I_scaled = self.I * factor + I_sigma_scaled = ( + self.I_sigma * factor if self.I_sigma is not None else None + ) + return I_scaled, I_sigma_scaled + + def _corrected_or_raw_intensities(self): + """Return the scaled ``(I, I_sigma)``, cached against the + ``(log_scale, U_aniso)`` fingerprint. ``(None, None)`` when this dataset + carries no intensities, and the raw pair if scaling is not set up. + """ + + def _tv(t): + return (t.data_ptr(), t._version) if isinstance(t, torch.Tensor) else None + + if self.I is None: + return (None, None) + + fp = ( + _tv(getattr(self, "log_scale", None)), + _tv(getattr(self, "U_aniso", None)), + ) + if self._corrected_I_fp != fp or self._corrected_I_cache is None: + try: + self._corrected_I_cache = self.get_corrected_intensities() + except Exception: + self._corrected_I_cache = (self.I, self.I_sigma) + self._corrected_I_fp = fp + return self._corrected_I_cache + def generate_validation_set( self, val_fraction_of_free: float = 0.5, From 2164229b2aa6170c7b1bd90f113767776caf3da6 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 24 Aug 2026 13:38:57 +0200 Subject: [PATCH 049/250] Add batched observation accessors to DatasetCollection stack_F_obs / stack_F_sigma / stack_I_obs / stack_I_sigma return the scaled observations, so a batched target sees the same per-dataset scaling the per-dataset accessors apply. stack_masks selects through the 3-way work/free/validation subsets rather than the 2-way rfree_flags, which cannot express a validation set carved out of free. stack_I_obs names the offending dataset when a member has no intensities. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- .../io/test_collection_stack_accessors.py | 165 ++++++++++++++++++ torchref/io/datasets/collection.py | 113 ++++++++++++ 2 files changed, 278 insertions(+) create mode 100644 tests/unit/io/test_collection_stack_accessors.py diff --git a/tests/unit/io/test_collection_stack_accessors.py b/tests/unit/io/test_collection_stack_accessors.py new file mode 100644 index 00000000..8c535b8d --- /dev/null +++ b/tests/unit/io/test_collection_stack_accessors.py @@ -0,0 +1,165 @@ +"""The batched observation accessors must return the same data the targets fit. + +Two things they could quietly get wrong, both invisible in the shape: + +* returning **raw** ``F``/``I`` instead of the scaled ones, which drops the per-dataset + ``log_scale``/``U_aniso`` that ``DatasetCollection.scale()`` fits; +* masking with the 2-way ``rfree_flags`` instead of the 3-way work/free/validation + subsets, which lets validation reflections into the work set. + +Each is pinned against the per-dataset accessor it is a batched form of. +""" + +import pytest +import torch + + +@pytest.fixture(scope="module") +def collection(mtz_dir): + """Two 1DAW datasets with *different* scales, so raw and scaled disagree.""" + mtz = mtz_dir / "1DAW.mtz" + if not mtz.exists(): + pytest.skip("1DAW fixture not present") + + from torchref import ReflectionData + from torchref.io.datasets.collection import DatasetCollection + + a = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + b = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + if a.I is None: + pytest.skip("1DAW loaded without intensities") + + dc = DatasetCollection(verbose=0, device="cpu") + dc.add_dataset("dark", a, set_as_reference=True) + dc.add_dataset("light", b) + with torch.no_grad(): + # A real inter-dataset scale difference: the whole point of the corrected path. + b.log_scale += 0.4 + return dc + + +@pytest.mark.unit +class TestScaledNotRaw: + def test_amplitude_rows_match_the_per_dataset_corrected_accessor(self, collection): + stacked = collection.stack_F_obs() + sigma = collection.stack_F_sigma() + for row, key in enumerate(collection.keys()): + F, sig = collection[key].get_corrected_data() + assert torch.equal(stacked[row], F) + assert torch.equal(sigma[row], sig) + + def test_intensity_rows_match_the_per_dataset_corrected_accessor(self, collection): + stacked = collection.stack_I_obs() + sigma = collection.stack_I_sigma() + for row, key in enumerate(collection.keys()): + I, sig = collection[key].get_corrected_intensities() + assert torch.equal(stacked[row], I) + assert torch.equal(sigma[row], sig) + + def test_the_two_datasets_differ_after_scaling(self, collection): + """Anti-vacuity: with identical scales, raw and corrected agree and neither + assertion above could detect the wrong accessor.""" + stacked = collection.stack_F_obs() + assert not torch.allclose(stacked[0], stacked[1]) + raw = torch.stack([collection[k].F for k in collection.keys()], dim=0) + assert torch.allclose(raw[0], raw[1]), "raw amplitudes should be identical here" + assert not torch.allclose(stacked, raw) + + def test_scaled_intensities_are_the_square_of_the_scaled_amplitude_factor( + self, collection + ): + """Ties the two stacks together, so they cannot drift apart in scale.""" + F = collection.stack_F_obs() + I = collection.stack_I_obs() + raw_F = torch.stack([collection[k].F for k in collection.keys()], dim=0) + raw_I = torch.stack([collection[k].I for k in collection.keys()], dim=0) + + keep = (raw_F.abs() > 1e-6) & (raw_I.abs() > 1e-6) + amp = (F[keep] / raw_F[keep]) ** 2 + inten = I[keep] / raw_I[keep] + assert torch.allclose(inten, amp, rtol=1e-5) + + +@pytest.mark.unit +class TestThreeWayMasks: + @pytest.mark.parametrize("use_set", ["work", "free", "val"]) + def test_rows_match_the_per_dataset_subset(self, collection, use_set): + attr = {"work": "work", "free": "free", "val": "validation"}[use_set] + stacked = collection.stack_masks(use_set=use_set) + for row, key in enumerate(collection.keys()): + assert torch.equal(stacked[row], getattr(collection[key], attr).mask) + + def test_the_three_subsets_partition_the_valid_reflections(self, collection): + work = collection.stack_masks(use_set="work") + free = collection.stack_masks(use_set="free") + val = collection.stack_masks(use_set="val") + + assert not (work & free).any() + assert not (work & val).any() + assert not (free & val).any() + + for row, key in enumerate(collection.keys()): + valid = collection[key].masks().to(torch.bool) + assert torch.equal(work[row] | free[row] | val[row], valid) + + def test_a_validation_set_is_carved_out_of_free_not_work(self, collection): + """The 3-way behaviour a 2-way flag array cannot reproduce.""" + data = collection["dark"] + saved = None if data.validation_flags is None else data.validation_flags.clone() + try: + free_before = int(collection.stack_masks(use_set="free")[0].sum()) + data.generate_validation_set(val_fraction_of_free=0.5, seed=0) + + free_after = int(collection.stack_masks(use_set="free")[0].sum()) + val_after = int(collection.stack_masks(use_set="val")[0].sum()) + + assert val_after > 0 + assert free_after < free_before + assert free_after + val_after == pytest.approx(free_before, abs=1) + finally: + data.validation_flags = saved + data._subset_fp = None + + def test_an_unknown_subset_name_is_rejected(self, collection): + with pytest.raises(ValueError, match="use_set must be"): + collection.stack_masks(use_set="test") + + +@pytest.mark.unit +class TestSelectionAndErrors: + def test_keys_argument_selects_and_orders_the_rows(self, collection): + both = collection.stack_F_obs() + one = collection.stack_F_obs(keys=["light"]) + assert one.shape[0] == 1 + assert torch.equal(one[0], both[1]) + + reversed_ = collection.stack_F_obs(keys=["light", "dark"]) + assert torch.equal(reversed_[0], both[1]) + assert torch.equal(reversed_[1], both[0]) + + def test_unknown_key_is_rejected(self, collection): + with pytest.raises(KeyError, match="Unknown dataset keys"): + collection.stack_F_obs(keys=["nope"]) + + def test_centric_flags_are_shared_and_hkl_shaped(self, collection): + centric = collection.get_centric_flags() + assert centric is not None + assert centric.shape == (len(collection.hkl),) + assert centric.dtype == torch.bool + + def test_missing_intensities_name_the_offending_dataset(self, mtz_dir): + mtz = mtz_dir / "3GR5.mtz" + if not mtz.exists(): + pytest.skip("3GR5 fixture not present") + + from torchref import ReflectionData + from torchref.io.datasets.collection import DatasetCollection + + amp_only = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + if amp_only.I is not None: + pytest.skip("3GR5 unexpectedly carries intensities") + + dc = DatasetCollection(verbose=0, device="cpu") + dc.add_dataset("amps", amp_only, set_as_reference=True) + with pytest.raises(ValueError, match="'amps'"): + dc.stack_I_obs() diff --git a/torchref/io/datasets/collection.py b/torchref/io/datasets/collection.py index e97b7b15..38fd2c62 100644 --- a/torchref/io/datasets/collection.py +++ b/torchref/io/datasets/collection.py @@ -314,6 +314,119 @@ def closure(): [p.requires_grad_(False) for p in parameters] + # ------------------------------------------------------------------ + # Batched observation accessors + # ------------------------------------------------------------------ + # + # Every member is expanded onto the common HKL grid by ``add_dataset``, so these + # stack cleanly on a leading dataset axis. All of them return the **scaled** + # observations -- the per-dataset ``log_scale``/``U_aniso`` that ``scale()`` fits + # exists only in the corrected accessors, and a target reading the raw tensors + # would silently ignore the inter-dataset scaling. + + def _keys_or_all(self, keys: Optional[List[str]]) -> List[str]: + if keys is None: + return list(self._dataset_order) + missing = [k for k in keys if k not in self._datasets] + if missing: + raise KeyError(f"Unknown dataset keys: {missing}") + return list(keys) + + def stack_F_obs(self, keys: Optional[List[str]] = None) -> torch.Tensor: + """Scaled observed amplitudes, shape ``(n_datasets, n_reflections)``.""" + return torch.stack( + [self._datasets[k]._corrected_or_raw()[0] for k in self._keys_or_all(keys)], + dim=0, + ) + + def stack_F_sigma(self, keys: Optional[List[str]] = None) -> torch.Tensor: + """Scaled amplitude sigmas, shape ``(n_datasets, n_reflections)``.""" + return torch.stack( + [self._datasets[k]._corrected_or_raw()[1] for k in self._keys_or_all(keys)], + dim=0, + ) + + def stack_I_obs(self, keys: Optional[List[str]] = None) -> torch.Tensor: + """Scaled observed intensities, shape ``(n_datasets, n_reflections)``. + + Raises + ------ + ValueError + If any selected dataset carries no intensities. + """ + return torch.stack( + [ + self._require_intensities(k)[0] + for k in self._keys_or_all(keys) + ], + dim=0, + ) + + def stack_I_sigma(self, keys: Optional[List[str]] = None) -> torch.Tensor: + """Scaled intensity sigmas, shape ``(n_datasets, n_reflections)``. + + Raises + ------ + ValueError + If any selected dataset carries no intensities. + """ + return torch.stack( + [ + self._require_intensities(k)[1] + for k in self._keys_or_all(keys) + ], + dim=0, + ) + + def _require_intensities(self, key: str): + """``(I, I_sigma)`` scaled, with the dataset named in the error.""" + data = self._datasets[key] + if data.I is None: + raise ValueError( + f"Dataset {key!r} carries no intensities; its reflection file had no " + f"I/SIGI columns. An intensity-space target needs them on every member." + ) + return data._corrected_or_raw_intensities() + + def stack_masks( + self, keys: Optional[List[str]] = None, use_set: str = "work" + ) -> torch.Tensor: + """Per-dataset boolean subset masks, shape ``(n_datasets, n_reflections)``. + + Uses the 3-way ``work``/``free``/``validation`` accessors, so validation + reflections are excluded from both work and free -- matching what the + collection targets fit. The 2-way ``rfree_flags`` cannot express that. + + Parameters + ---------- + keys : list of str, optional + Datasets to stack; all of them in insertion order by default. + use_set : {"work", "free", "val"}, optional + Which subset to select. Default ``"work"``. + """ + if use_set not in ("work", "free", "val"): + raise ValueError( + f"use_set must be 'work', 'free' or 'val'; got {use_set!r}" + ) + attr = {"work": "work", "free": "free", "val": "validation"}[use_set] + return torch.stack( + [ + getattr(self._datasets[k], attr).mask + for k in self._keys_or_all(keys) + ], + dim=0, + ) + + def get_centric_flags(self) -> Optional[torch.Tensor]: + """Centric flags on the common HKL, from the reference dataset. + + A pure function of ``(hkl, spacegroup)``, so it is shared by every member and + needs no dataset axis. + """ + if self._reference_dataset is None: + return None + return self._datasets[self._reference_dataset].centric + def component_structure_factors( self, model_collection, recalc: bool = False ) -> torch.Tensor: From 8e6758dd0a894371723a1835ac2c9f6e537c8909 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 24 Aug 2026 15:09:00 +0200 Subject: [PATCH 050/250] Store populations as a shared activation plus per-timepoint branching ModelCollection now owns the population parameters and each timepoint is a view onto one row. The factorisation w(t) = (1 - alpha) * e_ref + alpha * q(t) says only the overall degree of activation varies from crystal to crystal, while the branching among excited components is conserved. It makes the mixture exactly linear in alpha, so the activation Jacobian is constant and a second moment of the activation distribution costs no extra structure-factor evaluation. sigma_alpha_sq = alpha (1 - alpha) * lambda_twin satisfies its bound by construction. Consequences: - The reference row is exactly e_ref. Previously add_dark asked for 0.0 and the log-clamp stored 1e-6, so anything deriving a bound from it inherited that floor. - Freezing fractions is collection-wide; one shared activation cannot be frozen for a single timepoint. - add_timepoint raises when the requested fractions imply a conflicting activation, naming set_fraction_override, rather than silently projecting. - from_ihm keeps deposited populations verbatim via an override: an IHM ensemble is not a kinetic series and its groups have no shared-pump physics. lambda_twin is a plain float while fixed, because sigmoid cannot return exactly 0 and 0 is what reproduces the coherent single-moment model. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- docs/changelog.rst | 8 + .../model/test_batched_component_fcalcs.py | 3 +- .../model/test_model_collection_fractions.py | 262 +++++++++-- torchref/cli/collection_difference_refine.py | 2 +- torchref/experimental/kinetic/refinement.py | 42 +- torchref/io/ihm.py | 40 +- torchref/io/ihm_mapping.py | 2 +- torchref/model/model_collection.py | 441 +++++++++++++++--- 8 files changed, 683 insertions(+), 117 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 607dfded..44046ec9 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,14 @@ Changelog Version 0.6.4 ---------- +- ``ModelCollection`` now stores populations as a shared activation fraction plus a per-timepoint branching, instead of free fractions per timepoint +- Freezing and unfreezing fractions is now collection-wide; timepoints needing independent populations use ``set_fraction_override`` +- ``add_timepoint`` raises when the requested fractions imply an activation that conflicts with one already set +- Added ``ModelCollection.sigma_alpha_sq`` and ``lambda_twin`` for the spread of activation across crystals +- Added batched ``compute_component_fcalcs`` / ``mix_component_fcalcs`` and ``DatasetCollection.component_structure_factors`` +- Added ``CollectionScaler.forward_batched`` for scaling several mixtures in one pass +- Added ``ReflectionData.get_corrected_intensities`` and scaled ``I``/``sigI`` subset views, with the unscaled values as ``I_raw``/``sigI_raw`` +- Added batched ``stack_F_obs`` / ``stack_I_obs`` / ``stack_masks`` accessors on ``DatasetCollection`` - Fixed ``f_sol_override`` overwriting the scaler's cached ``F_sol``, so a later call without an override read the wrong solvent - Fixed a batched ``f_sol_override`` gaining a spurious leading axis, which changed the rank of the scaled structure factors - Fixed the bulk-solvent ``F_sol`` staying at the starting model's mask for every refinement macrocycle diff --git a/tests/unit/model/test_batched_component_fcalcs.py b/tests/unit/model/test_batched_component_fcalcs.py index a4ed9963..0447e2d0 100644 --- a/tests/unit/model/test_batched_component_fcalcs.py +++ b/tests/unit/model/test_batched_component_fcalcs.py @@ -189,11 +189,12 @@ def test_weights_matrix_is_applied_row_wise(self, pair): def test_gradient_flows_through_the_contraction(self, pair): dc, mc = pair + mc.unfreeze_all_fractions() stacked = dc.component_structure_factors(mc, recalc=True) w = mc.get_fractions_matrix() mc.mix_component_fcalcs(stacked, w).abs().sum().backward() - grad = mc["light"].fraction_params.grad + grad = mc._activation_logit.grad assert grad is not None assert torch.isfinite(grad).all() diff --git a/tests/unit/model/test_model_collection_fractions.py b/tests/unit/model/test_model_collection_fractions.py index 4266dddd..62aa21c7 100644 --- a/tests/unit/model/test_model_collection_fractions.py +++ b/tests/unit/model/test_model_collection_fractions.py @@ -60,13 +60,15 @@ def hkl(): return torch.tensor([[1, 0, 0], [0, 1, 0], [1, 1, 0], [2, 0, 1]]) -class TestSoftmaxStorage: +class TestPopulationFactorisation: @pytest.mark.unit - def test_fractions_are_the_softmax_of_the_stored_logits(self, two_model_collection): + def test_fractions_are_the_activation_times_the_branching( + self, two_model_collection + ): mc = two_model_collection - mixed = mc["light"] - expected = torch.softmax(mixed.fraction_params, dim=0) - assert torch.allclose(mixed.fractions, expected) + alpha = mc.alpha_mean + expected = torch.stack([1.0 - alpha, alpha * mc.branching()[0][0]]) + assert torch.allclose(mc["light"].fractions, expected) @pytest.mark.unit @pytest.mark.parametrize("f", [0.01, 0.22, 0.3, 0.5, 0.99]) @@ -80,16 +82,18 @@ def test_requested_fractions_round_trip(self, f): assert mc["t"].fractions.sum().item() == pytest.approx(1.0, abs=1e-6) @pytest.mark.unit - def test_dark_excited_fraction_is_the_clamp_floor_not_zero( - self, two_model_collection - ): - """``add_dark`` asks for exactly 0, but the log-clamp at 1e-6 means the stored - value is 1e-6. Anything deriving a bound from the dark's fraction inherits that - floor rather than a true zero. + def test_the_reference_row_is_exactly_e_ref(self, two_model_collection): + """The dark is the alpha = 0 evaluation, so its excited fraction is exactly + zero -- not a clamp floor. + + Under a softmax over per-timepoint logits it could not be: ``log(0)`` forces a + clamp, which left the dark sitting at 1e-6. Anything deriving a bound from the + reference's fraction (``sigma_alpha_sq <= alpha (1 - alpha)``) inherited that + floor instead of a true zero. """ dark = two_model_collection["dark"] - assert dark.fractions[1].item() == pytest.approx(1e-6, rel=1e-3) - assert dark.fractions[1].item() > 0.0 + assert dark.fractions[1].item() == 0.0 + assert dark.fractions[0].item() == 1.0 @pytest.mark.unit def test_fraction_dtype_and_device_follow_the_base_models(self): @@ -98,8 +102,8 @@ def test_fraction_dtype_and_device_follow_the_base_models(self): models = [_StubModel(0), _StubModel(1)] mc = ModelCollection(models, verbose=0) mc.add_timepoint("t", [0.6, 0.4]) - assert mc["t"].fraction_params.dtype == models[0].dtype_float - assert mc["t"].fraction_params.device == models[0].device + assert mc._activation_logit.dtype == models[0].dtype_float + assert mc._activation_logit.device == models[0].device class TestValidation: @@ -129,32 +133,36 @@ class TestFreezing: @pytest.mark.unit def test_dark_is_frozen_and_timepoints_are_not(self, two_model_collection): mc = two_model_collection - assert mc["dark"].fraction_params.requires_grad is False - assert mc["light"].fraction_params.requires_grad is True + # Frozen by default: population refinement is opt-in. + assert mc._activation_logit.requires_grad is False + assert mc.fraction_parameters() == [mc._activation_logit, + mc._branching_logits[0]] @pytest.mark.unit def test_freeze_and_unfreeze_flip_the_flag(self, two_model_collection): mixed = two_model_collection["light"] mixed.freeze_fractions() - assert mixed.fraction_params.requires_grad is False + assert mixed.collection._activation_logit.requires_grad is False mixed.unfreeze_fractions() - assert mixed.fraction_params.requires_grad is True + assert mixed.collection._activation_logit.requires_grad is True @pytest.mark.unit - def test_unfreeze_all_leaves_the_dark_frozen(self, two_model_collection): - """The dark reference must not become refinable through the bulk call.""" + def test_the_reference_carries_no_population_parameter( + self, two_model_collection + ): + """The reference is the alpha = 0 evaluation, not a row with pinned logits, + so there is nothing of its own to freeze or refine.""" mc = two_model_collection mc.unfreeze_all_fractions() - assert mc["dark"].fraction_params.requires_grad is False - assert mc["light"].fraction_params.requires_grad is True + assert "dark" not in mc._branching_rows + assert "light" in mc._branching_rows + assert mc._activation_logit.requires_grad is True @pytest.mark.unit def test_freeze_all_freezes_every_timepoint(self, two_model_collection): mc = two_model_collection mc.freeze_all_fractions() - assert all( - mc[k].fraction_params.requires_grad is False for k in mc.keys() - ) + assert all(not p.requires_grad for p in mc.fraction_parameters()) class TestOverride: @@ -168,7 +176,7 @@ def test_override_replaces_fractions_and_clears_back(self, two_model_collection) mixed.clear_fraction_override() assert torch.allclose( - mixed.fractions, torch.softmax(mixed.fraction_params, dim=0) + mixed.fractions, mixed.collection.fractions_matrix()[mixed._index] ) @pytest.mark.unit @@ -227,9 +235,9 @@ def test_a_timepoint_owns_only_its_fractions(self, two_model_collection): ``ModuleList``, so a timepoint must not re-register their parameters. """ mixed = two_model_collection["light"] - owned = list(mixed.parameters()) - assert len(owned) == 1 - assert owned[0] is mixed.fraction_params + assert list(mixed.parameters()) == [], ( + "a timepoint view registered a parameter of its own" + ) @pytest.mark.unit def test_shared_base_parameters_are_counted_once(self, two_model_collection): @@ -239,7 +247,8 @@ def test_shared_base_parameters_are_counted_once(self, two_model_collection): """ mc = two_model_collection params = list(mc.parameters()) - assert len(params) == 2 + 2 + # two base anchors + activation + lambda + one branching row + assert len(params) == 2 + 3 base_anchors = [m.anchor for m in mc.base_models] for anchor in base_anchors: @@ -259,18 +268,195 @@ def test_forward_is_the_fraction_weighted_sum_of_the_parts( assert torch.allclose(mixed(hkl, recalc=True), expected) @pytest.mark.unit - def test_gradient_reaches_the_fraction_params(self, two_model_collection, hkl): - mixed = two_model_collection["light"] + def test_gradient_reaches_the_activation(self, two_model_collection, hkl): + mc = two_model_collection + mc.unfreeze_all_fractions() + mixed = mc["light"] mixed(hkl, recalc=True).abs().sum().backward() - grad = mixed.fraction_params.grad + grad = mc._activation_logit.grad assert grad is not None assert torch.isfinite(grad).all() @pytest.mark.unit - def test_frozen_dark_fractions_receive_no_gradient( + def test_the_reference_contributes_no_activation_gradient( self, two_model_collection, hkl ): - dark = two_model_collection["dark"] - dark(hkl, recalc=True).abs().sum().backward() - assert dark.fraction_params.grad is None + """The dark dataset carries no activation information, so its row must be + exactly e_ref with no path back to the shared parameter.""" + mc = two_model_collection + mc.unfreeze_all_fractions() + mc["dark"](hkl, recalc=True).abs().sum().backward() + assert mc._activation_logit.grad is None or float( + mc._activation_logit.grad.abs().max() + ) == 0.0 + + +class TestSharedActivation: + """One activation serves every timepoint; only the branching varies with time.""" + + @pytest.mark.unit + def test_a_second_timepoint_may_rebranch_at_the_same_activation(self): + """Three components, two timepoints, same 30% activated but split differently + between the two excited states.""" + from torchref.model.model_collection import ModelCollection + + mc = ModelCollection([_StubModel(i) for i in range(3)], verbose=0) + mc.add_dark() + mc.add_timepoint("early", [0.7, 0.3, 0.0]) + mc.add_timepoint("late", [0.7, 0.0, 0.3]) + + assert float(mc.alpha_mean) == pytest.approx(0.3, abs=1e-5) + assert torch.allclose( + mc["early"].fractions, + torch.tensor([0.7, 0.3, 0.0]), + atol=1e-5, + ) + assert torch.allclose( + mc["late"].fractions, + torch.tensor([0.7, 0.0, 0.3]), + atol=1e-5, + ) + + @pytest.mark.unit + def test_a_conflicting_activation_is_rejected_not_projected(self): + """A silent least-squares projection here would produce populations nobody + asked for, so this raises and names the escape hatch.""" + from torchref.model.model_collection import ModelCollection + + mc = ModelCollection([_StubModel(0), _StubModel(1)], verbose=0) + mc.add_dark() + mc.add_timepoint("early", [0.7, 0.3]) + + with pytest.raises(ValueError, match="set_fraction_override"): + mc.add_timepoint("late", [0.5, 0.5]) + + @pytest.mark.unit + def test_the_rejection_message_names_both_activations(self): + from torchref.model.model_collection import ModelCollection + + mc = ModelCollection([_StubModel(0), _StubModel(1)], verbose=0) + mc.add_dark() + mc.add_timepoint("early", [0.7, 0.3]) + + with pytest.raises(ValueError) as excinfo: + mc.add_timepoint("late", [0.5, 0.5]) + text = str(excinfo.value) + assert "0.5000" in text and "0.3000" in text and "early" in text + + @pytest.mark.unit + def test_adding_the_reference_after_a_timepoint_leaves_activation_alone(self): + """A pure-reference row carries no activation information.""" + from torchref.model.model_collection import ModelCollection + + mc = ModelCollection([_StubModel(0), _StubModel(1)], verbose=0) + mc.add_timepoint("light", [0.78, 0.22]) + mc.add_dark() + assert float(mc.alpha_mean) == pytest.approx(0.22, abs=1e-5) + + +class TestActivationJacobian: + @pytest.mark.unit + def test_rows_sum_to_zero_and_the_reference_row_vanishes( + self, two_model_collection + ): + """Fractions stay on the simplex, so the derivative is tangent to it; and the + reference does not move with the activation at all.""" + mc = two_model_collection + jac = mc.activation_jacobian() + + assert jac.shape == (len(mc), mc.n_base_models) + assert torch.allclose(jac.sum(dim=1), torch.zeros(len(mc)), atol=1e-6) + assert torch.equal(jac[0], torch.zeros(mc.n_base_models)) + assert float(jac[1][0]) == pytest.approx(-1.0) + + @pytest.mark.unit + def test_fractions_matrix_is_e_ref_plus_alpha_times_the_jacobian( + self, two_model_collection + ): + mc = two_model_collection + e_ref = torch.zeros(mc.n_base_models) + e_ref[0] = 1.0 + expected = e_ref.unsqueeze(0) + mc.alpha_mean * mc.activation_jacobian() + assert torch.allclose(mc.fractions_matrix(), expected) + + @pytest.mark.unit + def test_the_mixture_is_exactly_linear_in_the_activation( + self, two_model_collection + ): + """The property the second moment rests on: the secant equals the derivative, + so there is no truncation term anywhere downstream.""" + mc = two_model_collection + jac = mc.activation_jacobian() + + mc.set_activation(0.6) + w1 = mc.fractions_matrix().clone() + mc.set_activation(0.1) + w2 = mc.fractions_matrix().clone() + + assert torch.allclose(w1 - w2, (0.6 - 0.1) * jac, atol=1e-6) + + +class TestActivationDispersion: + @pytest.mark.unit + def test_lambda_is_exactly_zero_by_default(self, two_model_collection): + """Exactly, not approximately: sigmoid can never return 0, so a fixed float is + the only way to reproduce the coherent single-moment model.""" + mc = two_model_collection + assert float(mc.lambda_twin) == 0.0 + assert float(mc.sigma_alpha_sq) == 0.0 + + @pytest.mark.unit + def test_lambda_is_not_a_live_parameter_until_asked_for( + self, two_model_collection + ): + mc = two_model_collection + + def _present(): + # Identity, not ``in``: ``==`` on tensors is elementwise. + return any(p is mc._lambda_logit for p in mc.fraction_parameters()) + + assert not _present() + + mc.set_lambda_twin(0.3, refinable=True) + assert _present() + assert mc._lambda_logit.requires_grad is True + + @pytest.mark.unit + @pytest.mark.parametrize("lam", [0.0, 0.25, 0.5, 1.0]) + def test_the_variance_bound_holds_by_construction( + self, two_model_collection, lam + ): + mc = two_model_collection + mc.set_lambda_twin(lam) + alpha = float(mc.alpha_mean) + bound = alpha * (1.0 - alpha) + + # The bound holds algebraically; the tolerance is float32 ulp, since the two + # sides reach alpha (1 - alpha) by different arithmetic. + sigma_sq = float(mc.sigma_alpha_sq) + assert 0.0 <= sigma_sq <= bound * (1.0 + 1e-6) + assert sigma_sq == pytest.approx(bound * lam, rel=1e-5) + + @pytest.mark.unit + def test_lambda_one_saturates_the_bound(self, two_model_collection): + """The fully incoherent limit: every crystal either fully activated or dark.""" + mc = two_model_collection + mc.set_lambda_twin(1.0) + alpha = float(mc.alpha_mean) + assert float(mc.sigma_alpha_sq) == pytest.approx(alpha * (1.0 - alpha), rel=1e-5) + + @pytest.mark.unit + @pytest.mark.parametrize("bad", [-0.1, 1.1]) + def test_out_of_range_lambda_is_rejected(self, two_model_collection, bad): + with pytest.raises(ValueError, match=r"\[0, 1\]"): + two_model_collection.set_lambda_twin(bad) + + @pytest.mark.unit + def test_refined_lambda_stays_strictly_interior(self, two_model_collection): + """Once refinable it is a sigmoid, so it can approach but never reach the + bounds -- which is why the fixed path exists.""" + mc = two_model_collection + mc.set_lambda_twin(0.0, refinable=True) + value = float(mc.lambda_twin.detach()) + assert 0.0 < value < 1.0 diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index 97a935be..b37bec4b 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -761,7 +761,7 @@ def main(): if args.refine_fractions: params = list(itertools.chain( - model_light.parameters(), [mixed.fraction_params] + model_light.parameters(), mc.fraction_parameters() )) else: params = list(model_light.parameters()) diff --git a/torchref/experimental/kinetic/refinement.py b/torchref/experimental/kinetic/refinement.py index 94514c0c..0e5a0a94 100644 --- a/torchref/experimental/kinetic/refinement.py +++ b/torchref/experimental/kinetic/refinement.py @@ -30,6 +30,7 @@ ref.refine(macro_cycles=5) """ +import warnings from typing import TYPE_CHECKING, Dict, List, Optional import torch @@ -321,10 +322,9 @@ def _collect_parameters(self, structures=True, fractions=True): ) if fractions: - for name in mc.timepoint_names: - p = mc[name].fraction_params - if p.requires_grad: - params.append(p) + # One shared activation plus a branching row per timepoint, owned by the + # collection rather than by any single timepoint. + params.extend(p for p in mc.fraction_parameters() if p.requires_grad) # Scaler parameters (if not frozen) if self.scaler is not None: @@ -560,16 +560,40 @@ def refine_kinetics(self, niter: int = 200, lr: float = 1e-2): if self.verbose > 1 and (step + 1) % 10 == 0: print(f" Kinetic opt step {step+1}/{niter}: loss = {loss.item():.6f}") - # Update free fraction parameters to match final kinetic predictions + # Update the population parameters to match the final kinetic predictions. + # + # The collection stores one shared activation plus a per-timepoint branching, + # so a predicted population vector is decomposed into the two. A kinetic model + # whose reactive fraction is genuinely constant gives the same activation at + # every timepoint; if it does not, the shared value cannot represent all of + # them and the closest one is kept, with the spread reported. with torch.no_grad(): kinetic_occ = kinetic_model() + implied = {} for tp_name, t_idx in all_overrides.items(): if tp_name == mc.dark_key: - continue # dark fractions stay frozen at [1,0,...,0] + continue # the reference is the alpha = 0 evaluation predicted = kinetic_occ[:, t_idx] - mc[tp_name].fraction_params.data = torch.log( - predicted.clamp(min=1e-6) - ) + alpha = float(1.0 - predicted[0]) + if alpha <= 1e-6: + continue + implied[tp_name] = alpha + mc.set_branching(tp_name, predicted[1:]) + + if implied: + alphas = list(implied.values()) + spread = max(alphas) - min(alphas) + mean_alpha = sum(alphas) / len(alphas) + mc.set_activation(mean_alpha) + if spread > 1e-3: + warnings.warn( + f"Kinetic populations imply activations spanning {spread:.4f} " + f"across timepoints, but one activation is shared by all of " + f"them; using the mean {mean_alpha:.4f}. Drive the timepoints " + f"with set_fraction_override() to keep them independent.", + UserWarning, + stacklevel=2, + ) mc.unfreeze_structures() diff --git a/torchref/io/ihm.py b/torchref/io/ihm.py index 30ca1c83..06d4de65 100644 --- a/torchref/io/ihm.py +++ b/torchref/io/ihm.py @@ -40,6 +40,42 @@ def _check_ihm_available(): ) +def _add_group_with_independent_populations(collection, name, fractions): + """Add an IHM model group, keeping its populations exactly as deposited. + + ``ModelCollection`` stores populations as one activation fraction shared across + timepoints plus a per-timepoint branching, which is the right model for a kinetic + series driven by a single pump. An IHM ensemble is not that: its model groups carry + arbitrary, independently deposited populations, and two groups may well disagree + about how much of the sample is in the reference state. + + So the group is registered at whatever activation the collection already holds and + its deposited fractions are installed as an override, which ``fractions`` returns + verbatim and ``write_ihm`` therefore round-trips unchanged. + """ + import torch + + try: + collection.add_timepoint(name, fractions=fractions) + return + except ValueError: + pass + + n = len(fractions) + placeholder = [0.0] * n + placeholder[0] = 1.0 - float(collection.alpha_mean) + if n > 1: + placeholder[1] = 1.0 - placeholder[0] + collection.add_timepoint(name, fractions=placeholder) + collection[name].set_fraction_override( + torch.tensor( + fractions, + dtype=collection._activation_logit.dtype, + device=collection._activation_logit.device, + ) + ) + + class IHMReader: """ Read IHM mmCIF files into torchref ModelCollection + IHMEnsembleMapping. @@ -507,7 +543,9 @@ def build_model_collection( if is_dark: collection.add_dark(fractions=fractions) else: - collection.add_timepoint(group.name, fractions=fractions) + _add_group_with_independent_populations( + collection, group.name, fractions + ) return collection diff --git a/torchref/io/ihm_mapping.py b/torchref/io/ihm_mapping.py index ba8853f7..d713769a 100644 --- a/torchref/io/ihm_mapping.py +++ b/torchref/io/ihm_mapping.py @@ -12,7 +12,7 @@ --------------- IHM state -> base model (ModelFT) in ModelCollection IHM model group -> timepoint entry (_SharedMixedModel) in ModelCollection -IHM population fraction -> fraction_params in _SharedMixedModel +IHM population fraction -> activation / branching on ModelCollection """ from dataclasses import dataclass, field diff --git a/torchref/model/model_collection.py b/torchref/model/model_collection.py index d5c58445..86d18293 100644 --- a/torchref/model/model_collection.py +++ b/torchref/model/model_collection.py @@ -1,9 +1,26 @@ """ Model collection for time-resolved kinetic refinement. -Provides ModelCollection — a named dictionary of MixedModel instances at -different timepoints that share the same base structural models (ModelFT). -Keys match DatasetCollection keys so targets can automatically pair them. +Provides ModelCollection — a named dictionary of mixed models at different timepoints +that share the same base structural models (ModelFT). Keys match DatasetCollection keys +so targets can automatically pair them. + +Populations are stored **factorised**, not as a free vector per timepoint:: + + w(t) = (1 - alpha) * e_ref + alpha * q(t) + +with one mean activation ``alpha`` shared across every timepoint and a per-timepoint +branching ``q(t)`` over the non-reference components. This is the statement that only +the overall degree of activation varies from crystal to crystal, while the branching +among excited states is conserved -- and it makes the mixture exactly linear in +``alpha``, so :meth:`ModelCollection.activation_jacobian` is constant and a second +moment of the activation distribution costs no extra structure-factor evaluation. + +``ModelCollection`` owns the population parameters; each timepoint is a view onto one +row (:class:`_SharedMixedModel`). One consequence worth knowing: freezing or unfreezing +fractions is collection-wide, because a single activation cannot be frozen for one +timepoint alone. Timepoints that genuinely need independent populations are driven +through ``set_fraction_override`` instead. """ from typing import TYPE_CHECKING, Dict, Iterator, List, Optional, Tuple @@ -13,42 +30,55 @@ from torchref.utils.device_mixin import DeviceMovementMixin from torchref.utils.device_resolution import resolve_device +from torchref.utils.utils import ModuleReference if TYPE_CHECKING: from torchref.model.model_ft import ModelFT from torchref.model.mixed_model import MixedModel +#: Activation fractions are clamped away from 0 and 1 before taking a logit, which +#: would otherwise be infinite. 1e-6 is the same floor the fraction storage has always +#: applied. +_FRACTION_EPS = 1e-6 + + +def _logit(p: float) -> float: + """Inverse sigmoid, clamped away from the infinities at 0 and 1.""" + p = min(max(float(p), _FRACTION_EPS), 1.0 - _FRACTION_EPS) + return float(torch.log(torch.tensor(p / (1.0 - p)))) + class _SharedMixedModel(DeviceMovementMixin, nn.Module): """ - MixedModel variant that references shared base models without re-registering them. + One timepoint's view of a :class:`ModelCollection`. - Standard MixedModel wraps models in nn.ModuleList, which causes - double-registration when the same ModelFT objects appear in multiple - timepoints. This class stores the shared models as a plain list - (no ownership) and only owns its own fraction parameters. + Owns nothing. The shared base models are held as a plain list and the population + parameters live on the parent collection, so neither is re-registered here -- + the same ownership pattern in both cases, and what keeps a base model's + parameters from appearing once per timepoint in ``parameters()``. - An external fraction override (via ``set_fraction_override``) can replace - the softmax-derived fractions; while active, ``fractions`` and ``forward`` - use the override tensor instead of ``softmax(fraction_params)``. + An external fraction override (via ``set_fraction_override``) replaces the + derived fractions; while active, ``fractions`` and ``forward`` use the override + tensor, and gradients flow to whatever produced it. Parameters ---------- base_models : List[ModelFT] Shared structural models (not re-registered as submodules here). - initial_fractions : List[float] - Initial population fractions (must sum to 1). - frozen_fractions : bool - If True, fractions are excluded from optimization. + collection : ModelCollection + Owner of the activation, branching and dispersion parameters. Referenced + without registration. + index : int + This timepoint's row in the collection's insertion order. device : torch.device, optional - Device for fraction parameters. + Device to reconcile the base models onto. """ def __init__( self, base_models: List["ModelFT"], - initial_fractions: List[float], - frozen_fractions: bool = False, + collection: "ModelCollection", + index: int, device: Optional[torch.device] = None, ): super().__init__() @@ -56,33 +86,17 @@ def __init__( # Store as plain list — the parent ModelCollection owns the ModuleList self._base_models = base_models - n = len(base_models) - if len(initial_fractions) != n: - raise ValueError( - f"Number of fractions ({len(initial_fractions)}) must match " - f"number of models ({n})." - ) - total = sum(initial_fractions) - if abs(total - 1.0) > 1e-3: - raise ValueError(f"Initial fractions must sum to 1.0, got {total:.6f}.") - - # Normalize to handle floating point drift - initial_fractions = [f / total for f in initial_fractions] - - # Reconcile across *all* base models, not just the first: otherwise a - # mixed-device list stays unreconciled and ``fractions_tensor`` below can - # land on a device the later models are not on. - device = resolve_device(*base_models, device=device) + # Parent reference, deliberately not a submodule: the population parameters + # are the collection's, shared across every timepoint. + self._collection_ref = ModuleReference(collection) + self._index = index - # Match base models' float dtype (consistent under a float64 config). - fractions_tensor = torch.tensor( - initial_fractions, dtype=base_models[0].dtype_float, device=device - ) - theta = torch.log(fractions_tensor.clamp(min=1e-6)) - self.fraction_params = nn.Parameter(theta, requires_grad=not frozen_fractions) + # Reconcile across *all* base models, not just the first, so a mixed-device + # list does not stay unreconciled. + resolve_device(*base_models, device=device) # Optional override: when set, fractions property returns this tensor - # instead of softmax(fraction_params). Used by refine_kinetics() to + # instead of the collection's derived row. Used by refine_kinetics() to # route kinetic model predictions directly into the F_calc computation. self._fraction_override: Optional[torch.Tensor] = None @@ -90,14 +104,20 @@ def __init__( # Properties # ------------------------------------------------------------------ + @property + def collection(self) -> "ModelCollection": + """The owning collection.""" + return self._collection_ref.module + @property def fractions(self) -> torch.Tensor: """Normalized population fractions -- the override tensor while one is - set (see ``set_fraction_override``), else ``softmax(fraction_params)``. + set (see ``set_fraction_override``), else this timepoint's row of the + parent's :meth:`ModelCollection.fractions_matrix`. """ if self._fraction_override is not None: return self._fraction_override - return torch.softmax(self.fraction_params, dim=0) + return self.collection.fractions_matrix()[self._index] @property def models(self) -> List["ModelFT"]: @@ -197,22 +217,30 @@ def forward(self, hkl: torch.Tensor, recalc: bool = False) -> torch.Tensor: # ------------------------------------------------------------------ def freeze_fractions(self): - self.fraction_params.requires_grad = False + """Freeze the population parameters. + + **Collection-wide.** The mean activation is a single parameter shared by every + timepoint, so it cannot be frozen for one timepoint alone; this delegates to + :meth:`ModelCollection.freeze_all_fractions`. + """ + self.collection.freeze_all_fractions() def unfreeze_fractions(self): - self.fraction_params.requires_grad = True + """Unfreeze the population parameters. Collection-wide; see + :meth:`freeze_fractions`.""" + self.collection.unfreeze_all_fractions() def set_fraction_override(self, fractions: torch.Tensor): """Override fractions with an external tensor (e.g. from kinetic model). - While active, ``self.fractions`` returns this tensor instead of - ``softmax(fraction_params)``, allowing gradients to flow through - the external source. + While active, ``self.fractions`` returns this tensor instead of the + collection's derived row, allowing gradients to flow through the external + source. """ self._fraction_override = fractions def clear_fraction_override(self): - """Remove fraction override, reverting to softmax(fraction_params).""" + """Remove the fraction override, reverting to the collection's derived row.""" self._fraction_override = None # ------------------------------------------------------------------ @@ -233,8 +261,12 @@ def get_individual_fcalc(self, hkl, recalc=True): def __repr__(self): fracs = self.fractions.detach().tolist() frac_str = ", ".join(f"{f:.3f}" for f in fracs) - frozen_str = "frozen" if not self.fraction_params.requires_grad else "learnable" - return f"_SharedMixedModel({len(self._base_models)} models, fractions=[{frac_str}], {frozen_str})" + learnable = self.collection._activation_logit.requires_grad + frozen_str = "learnable" if learnable else "frozen" + return ( + f"_SharedMixedModel({len(self._base_models)} models, " + f"fractions=[{frac_str}], {frozen_str})" + ) class ModelCollection(DeviceMovementMixin, nn.Module): @@ -284,10 +316,40 @@ def __init__( # Register base models as owned submodules (single source of truth) self._base_models = nn.ModuleList(base_models) - # Per-timepoint mixed models (own only fraction params) + # Per-timepoint views (own nothing; see _SharedMixedModel) self._timepoints = nn.ModuleDict() self._order: List[str] = [] + device = resolve_device(*base_models) + dtype = base_models[0].dtype_float + + # --- population parameters ------------------------------------- + # + # Only the *overall* activation varies from crystal to crystal; the branching + # among excited components is conserved. So the populations factorise as + # + # w(t) = (1 - alpha) * e_ref + alpha * q(t) + # + # with one activation shared across timepoints and a per-timepoint branching + # distribution over the K-1 non-reference components. Both are frozen by + # default: population refinement is opt-in. + self._activation_logit = nn.Parameter( + torch.tensor(_logit(1e-6), dtype=dtype, device=device), + requires_grad=False, + ) + self._branching_logits = nn.ParameterList() + self._branching_rows: Dict[str, int] = {} + + # Dispersion of the activation across crystals, as the fraction of its + # maximum: sigma_alpha^2 = alpha (1 - alpha) * lambda, so 0 <= lambda <= 1 + # holds by construction. Stored as a plain float while fixed, because + # sigmoid can never return exactly 0 and lambda = 0 is what reproduces the + # single-moment (coherent) model. + self._lambda_logit = nn.Parameter( + torch.zeros((), dtype=dtype, device=device), requires_grad=False + ) + self._lambda_fixed: Optional[float] = 0.0 + if self.verbose > 0: print( f"ModelCollection initialized with {len(base_models)} base models" @@ -326,21 +388,84 @@ def add_timepoint( n = len(self._base_models) if fractions is None: fractions = [1.0 / n] * n + if len(fractions) != n: + raise ValueError( + f"Number of fractions ({len(fractions)}) must match " + f"number of models ({n})." + ) + total = sum(fractions) + if abs(total - 1.0) > 1e-3: + raise ValueError(f"Initial fractions must sum to 1.0, got {total:.6f}.") + fractions = [f / total for f in fractions] + + index = len(self._order) + self._install_populations(name, fractions) mixed = _SharedMixedModel( base_models=list(self._base_models), - initial_fractions=fractions, - frozen_fractions=frozen_fractions, + collection=self, + index=index, ) self._timepoints[name] = mixed self._order.append(name) + if frozen_fractions: + self.freeze_all_fractions() + if self.verbose > 0: frac_str = ", ".join(f"{f:.3f}" for f in fractions) print(f" Added timepoint '{name}': fractions=[{frac_str}]") return self + def _install_populations(self, name: str, fractions: List[float]) -> None: + """Invert requested fractions into the (activation, branching) factorisation. + + The reference component's weight is ``1 - alpha`` by construction, so a + timepoint that is pure reference carries no branching row and leaves the + activation alone. Every other timepoint pins the shared activation; a second + one asking for a different value cannot be represented and is rejected rather + than silently projected. + + Raises + ------ + ValueError + If ``fractions`` implies an activation incompatible with one already set + by an earlier timepoint. + """ + alpha = 1.0 - fractions[0] + + if alpha <= _FRACTION_EPS: + # Pure reference: this is the dark / ground state, i.e. the alpha = 0 + # evaluation of the same parametrisation. No branching row. + return + + current = float(self.alpha_mean) + if self._branching_rows: + if abs(alpha - current) > 1e-3: + established = ", ".join(sorted(self._branching_rows)) + raise ValueError( + f"Timepoint {name!r} asks for activation {alpha:.4f}, but " + f"{established} already set it to {current:.4f}. One activation " + f"fraction is shared across all timepoints -- only the branching " + f"among excited components varies with time. To drive timepoints " + f"with independent populations, use set_fraction_override() on " + f"each one instead of passing fractions here." + ) + else: + with torch.no_grad(): + self._activation_logit.fill_(_logit(alpha)) + + # Branching over the K-1 non-reference components, renormalised within alpha. + excited = torch.tensor( + [f / alpha for f in fractions[1:]], + dtype=self._activation_logit.dtype, + device=self._activation_logit.device, + ) + logits = torch.log(excited.clamp(min=_FRACTION_EPS)) + self._branching_rows[name] = len(self._branching_logits) + self._branching_logits.append(nn.Parameter(logits, requires_grad=False)) + def add_dark( self, fractions: Optional[List[float]] = None ) -> "ModelCollection": @@ -363,7 +488,9 @@ def add_dark( n = len(self._base_models) fractions = [0.0] * n fractions[0] = 1.0 - return self.add_timepoint(self._dark_key, fractions, frozen_fractions=True) + # No frozen_fractions here: the reference is the alpha = 0 evaluation of the + # shared parametrisation, so it owns nothing that could be frozen. + return self.add_timepoint(self._dark_key, fractions) # ------------------------------------------------------------------ # Class methods @@ -553,14 +680,186 @@ def get_all_fractions(self) -> Dict[str, torch.Tensor]: return {name: self._timepoints[name].fractions for name in self._order} def get_fractions_matrix(self) -> torch.Tensor: + """All fractions as a matrix ``[n_timepoints, n_models]``, in insertion order. + + Alias of :meth:`fractions_matrix`, kept because it is the established name. """ - All fractions as a matrix [n_timepoints, n_models]. + return self.fractions_matrix() - Rows are ordered by ``self._order`` (i.e. insertion order). + # ------------------------------------------------------------------ + # Population factorisation + # ------------------------------------------------------------------ + + @property + def alpha_mean(self) -> torch.Tensor: + """Mean activation fraction, shared across all timepoints.""" + return torch.sigmoid(self._activation_logit) + + @property + def lambda_twin(self) -> torch.Tensor: + """Activation dispersion as a fraction of its maximum, in ``[0, 1]``. + + Zero is the coherent single-moment model. One means every crystal is either + fully activated or fully dark. Exactly representable while fixed; once + refinement is enabled it is ``sigmoid`` of a parameter and therefore strictly + interior. """ + if self._lambda_fixed is not None: + return torch.tensor( + self._lambda_fixed, + dtype=self._lambda_logit.dtype, + device=self._lambda_logit.device, + ) + return torch.sigmoid(self._lambda_logit) + + @property + def sigma_alpha_sq(self) -> torch.Tensor: + """Variance of the activation across crystals. + + ``alpha (1 - alpha) * lambda``, so ``0 <= sigma_alpha_sq <= alpha (1 - alpha)`` + holds by construction -- the upper bound being the Bernoulli case. + """ + alpha = self.alpha_mean + return alpha * (1.0 - alpha) * self.lambda_twin + + def branching(self) -> torch.Tensor: + """Per-timepoint distribution over the non-reference components. + + Returns + ------- + torch.Tensor + Shape ``(n_branching_rows, n_base_models - 1)``, rows summing to 1. + Empty when no non-reference timepoint has been added. + """ + if not len(self._branching_logits): + return torch.zeros( + (0, max(len(self._base_models) - 1, 0)), + dtype=self._activation_logit.dtype, + device=self._activation_logit.device, + ) return torch.stack( - [self._timepoints[n].fractions for n in self._order], dim=0 + [torch.softmax(row, dim=0) for row in self._branching_logits], dim=0 + ) + + def activation_jacobian(self) -> torch.Tensor: + """``d(fractions) / d(alpha)`` for every timepoint. + + Shape ``(n_timepoints, n_base_models)``. Reference-only rows are exactly zero, + so the reference dataset carries no activation gradient. Every other row is + ``q(t) - e_ref``, whose entries sum to zero because the fractions stay on the + simplex. + + This is what makes a second moment computable through the same machinery as the + first: the mixture is exactly linear in ``alpha``, so this Jacobian is constant + in ``alpha`` and can be scaled by the same affine scaler as the mixture itself. + """ + n_models = len(self._base_models) + dtype = self._activation_logit.dtype + device = self._activation_logit.device + + q_all = self.branching() + rows = [] + for name in self._order: + row = torch.zeros(n_models, dtype=dtype, device=device) + if name in self._branching_rows: + q = q_all[self._branching_rows[name]] + row = torch.cat( + [torch.full((1,), -1.0, dtype=dtype, device=device), q] + ) + rows.append(row) + if not rows: + return torch.zeros((0, n_models), dtype=dtype, device=device) + return torch.stack(rows, dim=0) + + def fractions_matrix(self) -> torch.Tensor: + """Population fractions for every timepoint, ``[n_timepoints, n_models]``. + + ``e_ref + alpha * activation_jacobian``. Reference-only rows come out as exactly + ``e_ref`` with no gradient path to the activation. + """ + n_models = len(self._base_models) + dtype = self._activation_logit.dtype + device = self._activation_logit.device + + e_ref = torch.zeros(n_models, dtype=dtype, device=device) + e_ref[0] = 1.0 + return e_ref.unsqueeze(0) + self.alpha_mean * self.activation_jacobian() + + def fraction_parameters(self) -> List[nn.Parameter]: + """The population parameters, for handing to an optimizer. + + The shared activation and every branching row, plus the dispersion when it is + refinable. Replaces reaching into a per-timepoint parameter. + """ + params: List[nn.Parameter] = [self._activation_logit] + params.extend(self._branching_logits) + if self._lambda_fixed is None: + params.append(self._lambda_logit) + return params + + def set_activation(self, alpha: float) -> "ModelCollection": + """Set the shared mean activation fraction, in place and without gradient.""" + if not 0.0 <= float(alpha) <= 1.0: + raise ValueError(f"alpha must lie in [0, 1]; got {alpha}") + with torch.no_grad(): + self._activation_logit.fill_(_logit(alpha)) + return self + + def set_branching(self, name: str, q: torch.Tensor) -> "ModelCollection": + """Set one timepoint's branching distribution, in place and without gradient. + + Parameters + ---------- + name : str + Timepoint key. Must be a non-reference timepoint. + q : torch.Tensor + Weights over the ``n_base_models - 1`` non-reference components. Normalised + internally; need not sum to 1. + """ + if name not in self._branching_rows: + raise KeyError( + f"{name!r} has no branching row -- it is the reference timepoint, " + f"whose fractions are fixed at the alpha = 0 evaluation." + ) + q = torch.as_tensor( + q, dtype=self._activation_logit.dtype, device=self._activation_logit.device ) + q = q / q.sum() + with torch.no_grad(): + self._branching_logits[self._branching_rows[name]].copy_( + torch.log(q.clamp(min=_FRACTION_EPS)) + ) + return self + + def set_lambda_twin( + self, value: Optional[float], refinable: bool = False + ) -> "ModelCollection": + """Set the activation dispersion, fixed or refinable. + + Parameters + ---------- + value : float or None + Dispersion in ``[0, 1]``. ``None`` keeps the current value and only changes + refinability. + refinable : bool, optional + If True, ``lambda_twin`` becomes ``sigmoid`` of a live parameter and joins + :meth:`fraction_parameters`. Default False, which stores an exact float -- + the only way ``lambda_twin`` can be exactly 0. + """ + if value is not None: + if not 0.0 <= float(value) <= 1.0: + raise ValueError(f"lambda_twin must lie in [0, 1]; got {value}") + with torch.no_grad(): + self._lambda_logit.fill_(_logit(value)) + if refinable: + self._lambda_fixed = None + self._lambda_logit.requires_grad_(True) + else: + self._lambda_fixed = ( + float(value) if value is not None else float(self.lambda_twin) + ) + self._lambda_logit.requires_grad_(False) + return self # ------------------------------------------------------------------ # Batched structure factors @@ -653,15 +952,25 @@ def compute_all_fcalc( # ------------------------------------------------------------------ def freeze_all_fractions(self): - """Freeze fractions at all timepoints.""" - for _, mixed in self: - mixed.freeze_fractions() + """Exclude the population parameters from optimization. + + Acts on the shared activation and every branching row. There is nothing + per-timepoint to freeze: one activation serves all of them, and the reference + timepoint has no parameters at all. + """ + self._activation_logit.requires_grad_(False) + for row in self._branching_logits: + row.requires_grad_(False) def unfreeze_all_fractions(self): - """Unfreeze fractions at all timepoints (except dark).""" - for name, mixed in self: - if name != self._dark_key: - mixed.unfreeze_fractions() + """Include the population parameters in optimization. + + The dispersion ``lambda_twin`` is *not* affected; enable it explicitly with + :meth:`set_lambda_twin` so it can never be refined by accident. + """ + self._activation_logit.requires_grad_(True) + for row in self._branching_logits: + row.requires_grad_(True) def freeze_structures(self): """Freeze xyz and adp on all base models.""" From 81af1caed8de71d772a73e1def18bcf6941979d7 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 24 Aug 2026 15:52:21 +0200 Subject: [PATCH 051/250] Add the two-moment intensity target Merged Bragg intensities see the crystal-to-crystal activation distribution only through its first two moments. With the branching conserved the mixture is exactly linear in the activation, so the intensity is exactly quadratic and = |F(alpha)|^2 + sigma_alpha^2 |dF/dalpha|^2 holds for any activation distribution and any number of components. The second term is the variance the coherent model discards: strictly positive, phase-blind, and largest exactly where the difference signal is. Built from one set of per-component structure factors contracted twice -- with the fractions for the mean, with the activation Jacobian for its derivative -- both through the shared scaler, which is affine in F_calc and mixes the solvent linearly in the weights. Works in intensities, not amplitudes: French-Wilson reshapes precisely the quadratic information the second moment lives in, so there is no F**2 fallback and construction fails on amplitude-only data rather than fitting a distorted quantity. At lambda_twin = 0 the derivative branch is not built at all, so the coherent limit is identical rather than equal -- multiplying a live branch by zero would turn a non-finite F_calc into NaN across the whole gradient. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- tests/helpers/device_cases.py | 1 + .../refinement/test_two_moment_intensity.py | 359 ++++++++++++++++++ torchref/refinement/targets/__init__.py | 2 + .../refinement/targets/collection/__init__.py | 2 + .../targets/collection/intensity.py | 307 +++++++++++++++ 5 files changed, 671 insertions(+) create mode 100644 tests/unit/refinement/test_two_moment_intensity.py create mode 100644 torchref/refinement/targets/collection/intensity.py diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index 7a7515b3..e04c2494 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -298,6 +298,7 @@ class TargetDeviceCase: "CholeskyMixedTensor": "needs a valid ADP tensor; shares MixedTensor's paths", "CollectionScaler": "needs a dataset collection", "CollectionDifferenceTarget": "needs a dataset collection", + "CollectionTwoMomentIntensityTarget": "needs a dataset collection", "CollectionMLTarget": "needs a dataset collection", "CollectionRiceTarget": "needs a dataset collection", "ADPSigdTarget": "needs a model with ADPs", diff --git a/tests/unit/refinement/test_two_moment_intensity.py b/tests/unit/refinement/test_two_moment_intensity.py new file mode 100644 index 00000000..83b96273 --- /dev/null +++ b/tests/unit/refinement/test_two_moment_intensity.py @@ -0,0 +1,359 @@ +"""The two-moment intensity target: the identity, the coherent limit, and the plumbing. + +The central claim is an *identity*, not an approximation. For any finite set of +per-crystal activations, with the branching conserved, + + mean_c |F_D + a_c dF|^2 == |F_D + abar dF|^2 + var(a) |dF|^2 + +with ``abar`` and ``var`` the **population** moments (1/M divisor). The first test builds +the left-hand side by an explicit loop over crystals -- no two-moment expression anywhere +on the generator side -- so it tests the physics rather than restating the implementation. + +The trap it pins: a ``1/(M-1)`` divisor makes the identity fail by O(1/M). At M=64 that is +1.6% -- small enough to slip past a loose tolerance and far larger than the effect the +target exists to measure. +""" + +import pytest +import torch + + +def _sample_moments(alpha: torch.Tensor): + """Population mean and variance (1/M divisor, not 1/(M-1)).""" + return alpha.mean(), alpha.var(unbiased=False) + + +def _brute_force_mean_intensity(F_D, dF, alpha): + """mean_c |F_D + a_c dF|^2, by explicit loop. No two-moment expression.""" + total = torch.zeros(F_D.shape, dtype=F_D.real.dtype) + for a in alpha: + total = total + (F_D + a * dF).abs() ** 2 + return total / len(alpha) + + +@pytest.mark.unit +class TestTheMomentIdentity: + @pytest.mark.parametrize("m", [2, 7, 64]) + @pytest.mark.parametrize( + "dtype,tol", [(torch.float64, 1e-13), (torch.float32, 1e-5)] + ) + def test_identity_holds_for_any_finite_activation_set(self, m, dtype, tol): + gen = torch.Generator().manual_seed(11) + n = 32 + cdtype = torch.complex128 if dtype is torch.float64 else torch.complex64 + + F_D = torch.randn(n, generator=gen, dtype=dtype).to(cdtype) + 1j * torch.randn( + n, generator=gen, dtype=dtype + ).to(cdtype) + dF = torch.randn(n, generator=gen, dtype=dtype).to(cdtype) + 1j * torch.randn( + n, generator=gen, dtype=dtype + ).to(cdtype) + alpha = torch.rand(m, generator=gen, dtype=dtype) + + brute = _brute_force_mean_intensity(F_D, dF, alpha) + abar, var = _sample_moments(alpha) + two_moment = (F_D + abar * dF).abs() ** 2 + var * dF.abs() ** 2 + + rel = ((brute - two_moment).abs() / brute.abs().clamp(min=1e-30)).max() + assert rel < tol, f"identity failed at {rel:.2e} (M={m}, {dtype})" + + def test_the_unbiased_variance_divisor_breaks_it(self): + """Anti-vacuity for the divisor: the wrong one fails, and by how much.""" + gen = torch.Generator().manual_seed(3) + n, m = 16, 64 + F_D = torch.randn(n, generator=gen, dtype=torch.float64).to(torch.complex128) + dF = torch.randn(n, generator=gen, dtype=torch.float64).to(torch.complex128) + alpha = torch.rand(m, generator=gen, dtype=torch.float64) + + brute = _brute_force_mean_intensity(F_D, dF, alpha) + abar = alpha.mean() + wrong = (F_D + abar * dF).abs() ** 2 + alpha.var(unbiased=True) * dF.abs() ** 2 + + rel = ((brute - wrong).abs() / brute.abs()).max() + assert rel > 1e-3, ( + "the unbiased divisor produced the same answer, so this test cannot " + "detect the wrong one" + ) + + @pytest.mark.parametrize("m", [3, 16]) + def test_identity_holds_with_a_degenerate_zero_variance_set(self, m): + """Constant activation: the variance term must vanish exactly.""" + n = 8 + gen = torch.Generator().manual_seed(5) + F_D = torch.randn(n, generator=gen, dtype=torch.float64).to(torch.complex128) + dF = torch.randn(n, generator=gen, dtype=torch.float64).to(torch.complex128) + alpha = torch.full((m,), 0.31, dtype=torch.float64) + + brute = _brute_force_mean_intensity(F_D, dF, alpha) + abar, var = _sample_moments(alpha) + assert float(var) == pytest.approx(0.0, abs=1e-30) + assert torch.allclose(brute, (F_D + abar * dF).abs() ** 2, rtol=1e-13) + + def test_bernoulli_activation_gives_the_incoherent_sum(self): + """The lambda = 1 limit: fully-lit or fully-dark crystals add in intensity.""" + n = 64 + gen = torch.Generator().manual_seed(7) + F_D = torch.randn(n, generator=gen, dtype=torch.float64).to(torch.complex128) + F_L = torch.randn(n, generator=gen, dtype=torch.float64).to(torch.complex128) + dF = F_L - F_D + + w = 0.25 + m = 400 + alpha = torch.zeros(m, dtype=torch.float64) + alpha[: int(w * m)] = 1.0 + + brute = _brute_force_mean_intensity(F_D, dF, alpha) + incoherent = (1 - w) * F_D.abs() ** 2 + w * F_L.abs() ** 2 + assert torch.allclose(brute, incoherent, rtol=1e-12) + + # ...and the two-moment form reproduces it, with lambda exactly 1. + abar, var = _sample_moments(alpha) + assert float(var) == pytest.approx(abar * (1 - abar), rel=1e-12) + two_moment = (F_D + abar * dF).abs() ** 2 + var * dF.abs() ** 2 + assert torch.allclose(brute, two_moment, rtol=1e-12) + + +# ===================================================================== +# Integration against the real collection stack +# ===================================================================== + + +@pytest.fixture(scope="module") +def collection(pdb_dir, mtz_dir): + """A dark/light collection on 1DAW, which is the only fixture with I/SIGI.""" + pdb = pdb_dir / "1DAW.pdb" + mtz = mtz_dir / "1DAW.mtz" + if not (pdb.exists() and mtz.exists()): + pytest.skip("1DAW fixture not present") + + from torchref import ReflectionData + from torchref.cli._common import load_model + from torchref.io.datasets.collection import DatasetCollection + from torchref.model.model_collection import ModelCollection + from torchref.scaling.collection_scaler import CollectionScaler + + d_min = 2.05 + dark = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + light = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + if dark.I is None: + pytest.skip("1DAW loaded without intensities") + + model_dark = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) + model_light = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) + with torch.no_grad(): + model_light.xyz.refinable_params += 0.2 + + dc = DatasetCollection(verbose=0, device="cpu") + dc.add_dataset("dark", dark, set_as_reference=True) + dc.add_dataset("light", light) + + mc = ModelCollection([model_dark, model_light], dark_key="dark", verbose=0) + mc.add_dark() + mc.add_timepoint("light", [0.78, 0.22]) + + scaler = CollectionScaler(dc, mc, verbose=0) + scaler.initialize() + return dc, mc, scaler + + +def _target(dc, mc, scaler, **kw): + from torchref.refinement.targets import CollectionTwoMomentIntensityTarget + + return CollectionTwoMomentIntensityTarget(dc, mc, scaler=scaler, verbose=0, **kw) + + +@pytest.mark.integration +class TestCoherentLimit: + def test_lambda_zero_reduces_to_the_squared_mean(self, collection): + dc, mc, scaler = collection + mc.set_lambda_twin(0.0) + target = _target(dc, mc, scaler) + + model = target.intensity_model(recalc=True) + + rows = target._row_indices(target._keys()) + weights = mc.fractions_matrix()[rows] + components = dc.component_structure_factors(mc, recalc=False) + mean = scaler.forward_batched( + mc.mix_component_fcalcs(components, weights), weights + ) + assert torch.equal(model, mean.abs() ** 2) + + def test_lambda_zero_survives_a_poisoned_derivative(self, collection): + """The coherent limit must skip the variance branch, not multiply it by zero. + + A non-finite entry times exactly zero is NaN, which would poison the whole + gradient; this is what makes the short-circuit load-bearing rather than an + optimisation. + """ + dc, mc, scaler = collection + mc.set_lambda_twin(0.0) + target = _target(dc, mc, scaler) + assert not target._variance_is_live(mc.sigma_alpha_sq) + assert torch.isfinite(target.forward()) + + def test_a_nonzero_lambda_changes_the_prediction(self, collection): + """Anti-vacuity: the variance branch must actually do something.""" + dc, mc, scaler = collection + mc.set_lambda_twin(0.0) + coherent = _target(dc, mc, scaler).intensity_model(recalc=True) + + mc.set_lambda_twin(0.5) + try: + dispersed = _target(dc, mc, scaler).intensity_model(recalc=True) + finally: + mc.set_lambda_twin(0.0) + + assert not torch.allclose(coherent, dispersed) + # Strictly positive: |dF|^2 has no sign. + assert bool((dispersed >= coherent - 1e-6).all()) + + +@pytest.mark.integration +class TestForwardModelStructure: + def test_the_variance_term_is_sigma_sq_times_the_scaled_jacobian(self, collection): + dc, mc, scaler = collection + mc.set_lambda_twin(0.4) + try: + target = _target(dc, mc, scaler) + total = target.intensity_model(recalc=True) + + rows = target._row_indices(target._keys()) + components = dc.component_structure_factors(mc, recalc=False) + weights = mc.fractions_matrix()[rows] + jacobian = mc.activation_jacobian()[rows] + + mean = scaler.forward_batched( + mc.mix_component_fcalcs(components, weights), weights + ) + deriv = scaler.forward_batched( + mc.mix_component_fcalcs(components, jacobian), jacobian + ) + expected = mean.abs() ** 2 + mc.sigma_alpha_sq * deriv.abs() ** 2 + assert torch.allclose(total, expected, rtol=1e-6) + finally: + mc.set_lambda_twin(0.0) + + def test_the_reference_row_carries_no_variance(self, collection): + """The dark's Jacobian row is exactly zero, so its prediction is coherent + regardless of the dispersion -- a dark dataset holds no activation information.""" + dc, mc, scaler = collection + keys = _target(dc, mc, scaler)._keys() + assert keys[0] == "dark" + + mc.set_lambda_twin(0.0) + coherent = _target(dc, mc, scaler).intensity_model(recalc=True)[0] + mc.set_lambda_twin(0.9) + try: + dispersed = _target(dc, mc, scaler).intensity_model(recalc=True)[0] + finally: + mc.set_lambda_twin(0.0) + + assert torch.allclose(coherent, dispersed, rtol=1e-6) + + def test_shape_follows_the_fitted_keys(self, collection): + dc, mc, scaler = collection + target = _target(dc, mc, scaler) + model = target.intensity_model(recalc=True) + assert model.shape == (len(target._keys()), len(dc.hkl)) + + +@pytest.mark.integration +class TestLossAndReporting: + def test_forward_is_finite_and_positive(self, collection): + dc, mc, scaler = collection + loss = _target(dc, mc, scaler).forward() + assert torch.isfinite(loss) + assert loss.numel() == 1 + + def test_gradient_reaches_the_light_model(self, collection): + dc, mc, scaler = collection + target = _target(dc, mc, scaler) + target.forward().backward() + grad = mc.base_models[1].xyz.refinable_params.grad + assert grad is not None and torch.isfinite(grad).all() + assert float(grad.abs().max()) > 0 + + def test_gradient_reaches_the_dispersion_when_refinable(self, collection): + dc, mc, scaler = collection + mc.set_lambda_twin(0.3, refinable=True) + try: + _target(dc, mc, scaler).forward().backward() + grad = mc._lambda_logit.grad + assert grad is not None and torch.isfinite(grad).all() + assert float(grad.abs()) > 0 + finally: + mc._lambda_logit.grad = None + mc.set_lambda_twin(0.0) + + def test_rfactor_uses_the_two_moment_amplitude(self, collection): + dc, mc, scaler = collection + target = _target(dc, mc, scaler) + + rf = target.get_rfactor() + assert set(rf) == {"per_dataset", "rwork_pct", "rfree_pct"} + assert set(rf["per_dataset"]) == set(target._keys()) + for key, (rwork, rfree) in rf["per_dataset"].items(): + assert 0.0 < rwork < 2.0, f"{key}: {rwork}" + assert 0.0 < rfree < 2.0, f"{key}: {rfree}" + + def test_stats_report_the_activation_moments(self, collection): + dc, mc, scaler = collection + mc.set_lambda_twin(0.25) + try: + stats = _target(dc, mc, scaler).stats() + for key in ("alpha_mean", "lambda_twin", "sigma_alpha_sq", "alpha_sd", + "dI_frac", "rwork", "rfree", "loss"): + assert key in stats, f"missing stat: {key}" + assert stats["alpha_mean"].value == pytest.approx(0.22, abs=1e-4) + assert stats["lambda_twin"].value == pytest.approx(0.25, abs=1e-4) + assert stats["dI_frac"].value > 0.0 + finally: + mc.set_lambda_twin(0.0) + + def test_di_frac_is_zero_in_the_coherent_limit(self, collection): + """The stat that distinguishes "refined to zero" from "never refined".""" + dc, mc, scaler = collection + mc.set_lambda_twin(0.0) + assert _target(dc, mc, scaler).stats()["dI_frac"].value == 0.0 + + @pytest.mark.parametrize("use_set", ["work", "free"]) + def test_subset_selection_is_honoured(self, collection, use_set): + dc, mc, scaler = collection + target = _target(dc, mc, scaler, use_set=use_set) + assert target.use_set == use_set + expected = sum( + (dc[k].work if use_set == "work" else dc[k].free).n + for k in target._keys() + ) + assert target._n_reflections() == expected + + +@pytest.mark.integration +class TestIntensityRequirement: + def test_construction_fails_without_intensities(self, pdb_dir, mtz_dir): + """Fails at construction, not inside the first loss evaluation: LossState + probes forward at registration and that traceback is far harder to read.""" + mtz = mtz_dir / "3GR5.mtz" + pdb = pdb_dir / "3GR5.pdb" + if not (mtz.exists() and pdb.exists()): + pytest.skip("3GR5 fixture not present") + + from torchref import ReflectionData + from torchref.cli._common import load_model + from torchref.io.datasets.collection import DatasetCollection + from torchref.model.model_collection import ModelCollection + from torchref.refinement.targets import CollectionTwoMomentIntensityTarget + + data = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + if data.I is not None: + pytest.skip("3GR5 unexpectedly carries intensities") + + model = load_model(str(pdb), max_res=2.05, device="cpu", verbose=0) + dc = DatasetCollection(verbose=0, device="cpu") + dc.add_dataset("dark", data, set_as_reference=True) + mc = ModelCollection([model], dark_key="dark", verbose=0) + mc.add_dark() + + with pytest.raises(ValueError, match="I/SIGI"): + CollectionTwoMomentIntensityTarget(dc, mc, verbose=0) diff --git a/torchref/refinement/targets/__init__.py b/torchref/refinement/targets/__init__.py index 52736fb0..4205faa0 100644 --- a/torchref/refinement/targets/__init__.py +++ b/torchref/refinement/targets/__init__.py @@ -21,6 +21,7 @@ ) from .collection import ( CollectionDifferenceTarget, + CollectionTwoMomentIntensityTarget, CollectionMLTarget, CollectionRiceTarget, MultiModelADPTarget, @@ -86,6 +87,7 @@ "create_xray_target", # Collection (multi-dataset) targets "CollectionDifferenceTarget", + "CollectionTwoMomentIntensityTarget", "CollectionRiceTarget", "CollectionMLTarget", "MultiModelGeometryTarget", diff --git a/torchref/refinement/targets/collection/__init__.py b/torchref/refinement/targets/collection/__init__.py index f4302a77..ca6a0b26 100644 --- a/torchref/refinement/targets/collection/__init__.py +++ b/torchref/refinement/targets/collection/__init__.py @@ -9,6 +9,7 @@ from ._util import _scale_fcalc from .base import CollectionXrayTarget +from .intensity import CollectionTwoMomentIntensityTarget from .multimodel import MultiModelADPTarget, MultiModelGeometryTarget from .xray import ( CollectionDifferenceTarget, @@ -18,6 +19,7 @@ __all__ = [ "CollectionXrayTarget", + "CollectionTwoMomentIntensityTarget", "CollectionDifferenceTarget", "CollectionRiceTarget", "CollectionMLTarget", diff --git a/torchref/refinement/targets/collection/intensity.py b/torchref/refinement/targets/collection/intensity.py new file mode 100644 index 00000000..2304d7ea --- /dev/null +++ b/torchref/refinement/targets/collection/intensity.py @@ -0,0 +1,307 @@ +"""Two-moment intensity target for time-resolved collections. + +Merged Bragg intensities see the crystal-to-crystal activation distribution only through +its first two moments. With the branching among excited components conserved, the mixture +is exactly linear in the activation fraction, so the intensity is exactly quadratic and:: + + = |F(alpha_mean)|^2 + sigma_alpha^2 |dF/dalpha|^2 + +holds for *any* activation distribution and any number of components -- an identity, not a +truncation. The first term is what every existing target models; the second is the variance +the coherent model discards, and it is strictly positive, phase-blind, and largest exactly +where the difference signal is. + +The target works in **intensities** rather than amplitudes on purpose: the French-Wilson +conversion reshapes precisely the quadratic information the second moment lives in, so an +amplitude formulation would fit a distorted version of the quantity it is trying to measure. +Members must therefore carry ``I``/``SIGI``; there is no ``F**2`` fallback, because that +would silently reintroduce the distortion. +""" + +from typing import TYPE_CHECKING, Dict, List + +import numpy as np +import torch + +from torchref.base.metrics.rfactor import rfactor_work_free +from torchref.utils.stats import ( + VERBOSITY_DEBUG, + VERBOSITY_ESSENTIAL, + VERBOSITY_STANDARD, + StatEntry, + stat, +) + +from .base import CollectionXrayTarget + +if TYPE_CHECKING: + from torchref.io.datasets.collection import DatasetCollection + from torchref.model.model_collection import ModelCollection + from torchref.scaling.scaler_base import ScalerBase + + +_LOG_2PI = float(np.log(2.0 * np.pi)) + +#: Floor on the intensity sigma, as a fraction of the median over the fitted subset. A +#: merged intensity sigma can be reported as zero; unfloored it would dominate the sum. +_SIGMA_FLOOR_FRAC = 0.1 + + +class CollectionTwoMomentIntensityTarget(CollectionXrayTarget): + """ + Gaussian intensity likelihood at the two-moment forward model. + + Inherits the subset selector, the cache-reset discipline and the stats shape from + :class:`~torchref.refinement.targets.collection.base.CollectionXrayTarget`, so it + cannot disagree with the amplitude targets about which reflections it fits. + + The forward model is built from **one** set of per-component structure factors, + contracted twice: once with the fractions to get the mean, once with the activation + Jacobian to get its derivative. Both contractions go through the shared scaler, which + is affine in ``F_calc`` and mixes the bulk solvent linearly in the weights -- so the + second contraction returns the correctly scaled derivative rather than needing a + separate differentiation path. + + With ``lambda_twin`` fixed at zero the variance branch is not built at all. That makes + the coherent limit identical rather than merely equal: multiplying a live ``dF`` branch + by exactly zero would still propagate a non-finite ``F_calc`` into the loss. + + Parameters + ---------- + dataset_collection : DatasetCollection + Members must all carry intensities. + model_collection : ModelCollection + Supplies the components, the fractions and the activation moments. + scaler : ScalerBase, optional + Shared scaler. Needs ``forward_batched`` to scale the batch in one pass; without + one the unscaled mixture is used. + use_work_set : bool, optional + Legacy bool; superseded by ``use_set``. + use_set : str, optional + Canonical 3-way subset selector ``"work"``/``"free"``/``"val"``. + verbose : int, optional + Verbosity level. + + Raises + ------ + ValueError + On construction, if any fitted dataset carries no intensities. + """ + + name: str = "collection_two_moment_intensity" + + def __init__( + self, + dataset_collection: "DatasetCollection", + model_collection: "ModelCollection", + scaler: "ScalerBase" = None, + use_work_set: bool = True, + use_set: str = None, + verbose: int = 0, + ): + super().__init__( + dataset_collection, + model_collection, + scaler=scaler, + use_work_set=use_work_set, + use_set=use_set, + verbose=verbose, + ) + # Fail here rather than inside the first loss evaluation: LossState probes a + # target's forward at registration, and a traceback from there is much harder to + # trace back to "this MTZ had no intensity columns". + missing = [ + key for key in self._keys() if dataset_collection[key].I is None + ] + if missing: + raise ValueError( + f"Datasets {missing} carry no intensities. The two-moment target fits " + f"merged intensities directly -- converting amplitudes back with F**2 " + f"would reintroduce the French-Wilson distortion it exists to avoid. " + f"Supply reflection files with I/SIGI columns." + ) + + # ------------------------------------------------------------------ + # Forward model + # ------------------------------------------------------------------ + + def _row_indices(self, keys: List[str]) -> List[int]: + """Rows of the collection's fraction matrix corresponding to ``keys``.""" + order = self._model_collection.keys() + return [order.index(k) for k in keys] + + def _scale_batch(self, fcalc_batch, weights): + """Apply the shared scaler to a ``[T, R]`` batch with per-row weights.""" + scaler = self._scaler + if scaler is None: + return fcalc_batch + return scaler.forward_batched(fcalc_batch, weights) + + def intensity_model(self, recalc: bool = False) -> torch.Tensor: + """The two-moment predicted intensities, shape ``(n_datasets, n_reflections)``. + + Parameters + ---------- + recalc : bool, optional + Force recomputation of the component structure factors. + + Returns + ------- + torch.Tensor + Predicted intensities, rows aligned with :meth:`_keys`. + """ + dc, mc = self._dataset_collection, self._model_collection + keys = self._keys() + rows = self._row_indices(keys) + + components = dc.component_structure_factors(mc, recalc=recalc) + weights = mc.fractions_matrix()[rows] + + mean = self._scale_batch( + mc.mix_component_fcalcs(components, weights), weights + ) + intensity = mean.abs() ** 2 + + sigma_alpha_sq = mc.sigma_alpha_sq + if self._variance_is_live(sigma_alpha_sq): + jacobian = mc.activation_jacobian()[rows] + derivative = self._scale_batch( + mc.mix_component_fcalcs(components, jacobian), jacobian + ) + intensity = intensity + sigma_alpha_sq * derivative.abs() ** 2 + return intensity + + def _variance_is_live(self, sigma_alpha_sq) -> bool: + """Whether the second moment contributes. + + False only when the dispersion is *exactly* zero and not refinable, in which case + the derivative branch is skipped entirely rather than multiplied by zero. + """ + mc = self._model_collection + if mc._lambda_fixed is None: + return True + return bool(sigma_alpha_sq.detach().ne(0).any()) + + def forward(self) -> torch.Tensor: + """Summed Gaussian NLL of the observed intensities under the two-moment model.""" + dc = self._dataset_collection + keys = self._keys() + if not keys: + return torch.zeros((), device=dc.hkl.device) + + # Clear cached forwards so a preceding no-grad stats()/get_rfactor() call cannot + # leave a detached tensor that breaks the loss backward. + self._reset_model_caches() + + model = self.intensity_model(recalc=False) + obs = dc.stack_I_obs(keys).to(model.dtype) + sigma = dc.stack_I_sigma(keys).to(model.dtype) + mask = dc.stack_masks(keys, use_set=self.use_set) + + sigma = self._floor_sigma(sigma, mask) + + residual = obs - model + nll = ( + 0.5 * (residual / sigma) ** 2 + + torch.log(sigma) + + 0.5 * _LOG_2PI + ) + # A single non-finite entry would poison the whole gradient; a large finite + # penalty lets the step be rejected instead. + nll = torch.where(torch.isfinite(nll), nll, torch.full_like(nll, 1e6)) + return (nll * mask).sum() + + @staticmethod + def _floor_sigma(sigma: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: + """Clamp sigma away from zero, at a fraction of its median over the subset.""" + selected = sigma[mask] + if selected.numel() == 0: + return sigma.clamp(min=1e-6) + floor = torch.median(selected) * _SIGMA_FLOOR_FRAC + floor = torch.clamp(floor, min=1e-12) + return sigma.clamp(min=floor) + + # ------------------------------------------------------------------ + # Reporting + # ------------------------------------------------------------------ + + def get_rfactor(self) -> Dict[str, object]: + """Per-dataset R-work / R-free against the two-moment amplitudes. + + The reported amplitude is ``sqrt(I_model)``, the RMS amplitude the two-moment + model actually predicts -- not ``|F(alpha_mean)|``, which is only its first term + and would not correspond to the loss being minimised. + + Overrides the base implementation, which is per-pair and would recompute the + component stack once per dataset. + """ + dc = self._dataset_collection + keys = self._keys() + per_dataset: Dict[str, tuple] = {} + rworks: List[float] = [] + rfrees: List[float] = [] + + with torch.no_grad(): + model = self.intensity_model(recalc=True) + amplitudes = model.clamp(min=0.0).sqrt() + for row, key in enumerate(keys): + rwork, rfree = rfactor_work_free(dc[key], amplitudes[row]) + per_dataset[key] = (rwork, rfree) + rworks.append(rwork) + rfrees.append(rfree) + + return { + "per_dataset": per_dataset, + "rwork_pct": self._percentiles(rworks), + "rfree_pct": self._percentiles(rfrees), + } + + def stats(self) -> Dict[str, StatEntry]: + """Base collection X-ray stats plus the activation moments. + + ``dI_frac`` is the mean fraction of the predicted intensity carried by the second + moment. It is what separates "the dispersion refined to zero" from "the dispersion + was never refined", which are otherwise indistinguishable in the summary. + """ + out = super().stats() + mc = self._model_collection + + with torch.no_grad(): + alpha = float(mc.alpha_mean) + lam = float(mc.lambda_twin) + sigma_sq = float(mc.sigma_alpha_sq) + + out["alpha_mean"] = stat(alpha, VERBOSITY_ESSENTIAL) + out["lambda_twin"] = stat(lam, VERBOSITY_ESSENTIAL) + out["sigma_alpha_sq"] = stat(sigma_sq, VERBOSITY_STANDARD) + out["alpha_sd"] = stat(sigma_sq**0.5, VERBOSITY_STANDARD) + + keys = self._keys() + if keys and self._variance_is_live(mc.sigma_alpha_sq): + rows = self._row_indices(keys) + components = self._dataset_collection.component_structure_factors( + mc, recalc=True + ) + jacobian = mc.activation_jacobian()[rows] + derivative = self._scale_batch( + mc.mix_component_fcalcs(components, jacobian), jacobian + ) + variance_term = mc.sigma_alpha_sq * derivative.abs() ** 2 + total = self.intensity_model(recalc=False) + mask = self._dataset_collection.stack_masks( + keys, use_set=self.use_set + ) + denom = total[mask].abs().clamp(min=1e-12) + out["dI_frac"] = stat( + float((variance_term[mask] / denom).mean()), VERBOSITY_STANDARD + ) + else: + out["dI_frac"] = stat(0.0, VERBOSITY_STANDARD) + + branching = mc.branching() + for name, row in mc._branching_rows.items(): + for k in range(branching.shape[1]): + out[f"q_{name}_{k + 1}"] = stat( + float(branching[row, k]), VERBOSITY_DEBUG + ) + return out From 34f3ab4e252aee12c37ecd53d532b91651112997 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 24 Aug 2026 15:52:34 +0200 Subject: [PATCH 052/250] Expose the two-moment model on the difference-refine CLI, with MTZ output --two-moment fits merged intensities with the second moment; --lambda-twin sets the activation dispersion and --refine-lambda-twin makes it refinable. Refinement is off by default: the sigma_alpha^2 term is smooth and positive, so it is collinear with a scale or overall-B error and can absorb one. The activation moments join the JSON summary. Thirteen further MTZ columns describe the correction: the intensities actually fitted, the coherent and two-moment predictions, the variance term alone, decontaminated difference amplitudes and DED coefficients, the sigma_alpha^2-aware weight, and DDF = DF_corr - DF. DDF is the one to read first -- smooth against resolution means the correction is collinear with a scale error, structure in it is the signal. The decontaminated amplitude goes through the dataset's own French-Wilson estimator on the full reflection list, rebuilt when joining a collection has left the retained one the wrong length. Subtracting the variance pushes weak reflections negative, which is exactly where sqrt(clamp(I, 0)) is worst. Also fixes the CLI crashing at --verbose 0: the R-factors written into the deposition metadata were only computed inside the printing branch. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- docs/changelog.rst | 4 + tests/integration/test_cli_two_moment_mtz.py | 233 +++++++++++++++++ torchref/cli/collection_difference_refine.py | 255 ++++++++++++++++++- 3 files changed, 485 insertions(+), 7 deletions(-) create mode 100644 tests/integration/test_cli_two_moment_mtz.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 44046ec9..9cd16764 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,10 @@ Changelog Version 0.6.4 ---------- +- Added ``CollectionTwoMomentIntensityTarget``, fitting merged intensities as ``|F(alpha)|^2 + sigma_alpha^2 |dF|^2`` to account for crystal-to-crystal spread in activation +- Added ``--two-moment`` / ``--lambda-twin`` / ``--refine-lambda-twin`` to ``torchref.difference-refine``, and the activation moments to its JSON summary +- ``torchref.difference-refine`` writes thirteen further MTZ columns under ``--two-moment``, including decontaminated difference amplitudes and the ``DDF`` diagnostic +- Fixed ``torchref.difference-refine`` crashing at ``--verbose 0``, where the R-factors written into the deposition metadata were only computed for printing - ``ModelCollection`` now stores populations as a shared activation fraction plus a per-timepoint branching, instead of free fractions per timepoint - Freezing and unfreezing fractions is now collection-wide; timepoints needing independent populations use ``set_fraction_override`` - ``add_timepoint`` raises when the requested fractions imply an activation that conflicts with one already set diff --git a/tests/integration/test_cli_two_moment_mtz.py b/tests/integration/test_cli_two_moment_mtz.py new file mode 100644 index 00000000..1dc69030 --- /dev/null +++ b/tests/integration/test_cli_two_moment_mtz.py @@ -0,0 +1,233 @@ +"""The difference-refinement MTZ layout, pinned. + +``write_results_mtz`` assigns MTZ column types from hard-coded name lists and never calls +``infer_mtz_dtypes()``, so a column added to the output dict but missed in the type lists +is written with whatever dtype numpy produced -- silently, and into a file that gets +deposited. Nothing else in the suite asserts on these names. + +Two things are checked: the baseline column set is unchanged by the two-moment work, and +under ``--two-moment`` the thirteen extra columns appear with the right types and are +internally consistent. +""" + +import json +import os +import subprocess +import sys + +import pytest + +pytestmark = [pytest.mark.integration, pytest.mark.slow] + + +BASELINE_COLUMNS = { + "Fo_dark", "SIGFo_dark", "Fo_light", "SIGFo_light", + "DF", "SIGDF", "WDF", + "Fc_dark", "Fc_light", "DFc", "DFc_complex", + "2mDFop-DFc", "mDFop-DFc", + "PHIC_dark", "PHIC_mixed", "PHIC_diff", "PHIC_light", + "Fextp", "2Fextp-Fc", "Fextp-Fc", + "Fextc", "SIGFextc", "2Fextc-Fc", "Fextc-Fc", + "Fextb", "SIGFextb", "2Fextb-Fc", "Fextb-Fc", + "FreeR_flag_dark", "FreeR_flag_light", +} + +TWO_MOMENT_COLUMNS = { + "Io_light": "Intensity", + "SIGIo_light": "Stddev", + "Ic_light_coh": "Intensity", + "Ic_light_2mom": "Intensity", + "IVAR_ALPHA": "Intensity", + "Fo_light_corr": "SFAmplitude", + "SIGFo_light_corr": "Stddev", + "DF_corr": "SFAmplitude", + "SIGDF_corr": "Stddev", + "2mDFop-DFc_corr": "SFAmplitude", + "mDFop-DFc_corr": "SFAmplitude", + "DDF": "SFAmplitude", + "W_2MOM": "Weight", +} + +FRACTION = 0.25 +LAMBDA_TWIN = 0.2 + + +@pytest.fixture(scope="module") +def cli_script(project_root): + script = project_root / "torchref" / "cli" / "collection_difference_refine.py" + if not script.exists(): + pytest.skip("difference-refine CLI not found") + return script + + +@pytest.fixture(scope="module") +def intensity_pair(mtz_dir, pdb_dir, tmp_path_factory): + """A dark/light pair carrying I/SIGI, from the only fixture that has them. + + 1DAW is the sole reflection file under ``tests/files`` with intensity columns; 3GR5, + which the other difference-refine CLI test uses, has none and so cannot exercise an + intensity-space path at all. + """ + import torch + + from torchref import ReflectionData + + mtz = mtz_dir / "1DAW.mtz" + pdb = pdb_dir / "1DAW.pdb" + if not (mtz.exists() and pdb.exists()): + pytest.skip("1DAW fixture not present") + + data = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + if data.I is None: + pytest.skip("1DAW loaded without intensities") + + out = tmp_path_factory.mktemp("two_moment_cli") + n = len(data) + idx = torch.arange(n) + # Slightly different reflection sets, as a real dark/light pair would be. + data.__select__(idx < int(n * 0.97)).write_mtz(str(out / "dark.mtz")) + data.__select__(idx >= int(n * 0.03)).write_mtz(str(out / "light.mtz")) + return {"dir": out, "pdb": pdb} + + +def _run(cli_script, pair, outdir, *extra): + env = dict(os.environ) + # The installed torchref may point at a different checkout; make the subprocess + # import the tree under test. + root = str(cli_script.parents[2]) + env["PYTHONPATH"] = root + os.pathsep + env.get("PYTHONPATH", "") + + cmd = [ + sys.executable, str(cli_script), + "-dm", str(pair["pdb"]), "-lm", str(pair["pdb"]), + "-dsf", str(pair["dir"] / "dark.mtz"), + "-lsf", str(pair["dir"] / "light.mtz"), + "--fraction", str(FRACTION), + "--n-cycles", "1", "--n-steps", "1", "--max-iter", "3", + "--dmin", "2.2", "-o", str(outdir), + "--device", "cpu", "--verbose", "0", + *extra, + ] + proc = subprocess.run(cmd, capture_output=True, text=True, timeout=1800, env=env) + assert proc.returncode == 0, ( + f"CLI failed ({proc.returncode})\nstderr tail:\n{proc.stderr[-3000:]}" + ) + prefix = f"fractions_{round((1 - FRACTION) * 100)}_{round(FRACTION * 100)}_" + return outdir / f"{prefix}difference_data.mtz", outdir / f"{prefix}summary.json" + + +@pytest.fixture(scope="module") +def baseline_mtz(cli_script, intensity_pair, tmp_path_factory): + outdir = tmp_path_factory.mktemp("baseline") + return _run(cli_script, intensity_pair, outdir) + + +@pytest.fixture(scope="module") +def two_moment_mtz(cli_script, intensity_pair, tmp_path_factory): + outdir = tmp_path_factory.mktemp("two_moment") + return _run( + cli_script, intensity_pair, outdir, + "--two-moment", "--lambda-twin", str(LAMBDA_TWIN), + ) + + +def _read(path): + import reciprocalspaceship as rs + + return rs.read_mtz(str(path)) + + +class TestBaselineLayoutIsUnchanged: + def test_baseline_columns_are_exactly_the_expected_set(self, baseline_mtz): + mtz, _ = baseline_mtz + assert set(_read(mtz).columns) == BASELINE_COLUMNS + + def test_no_two_moment_columns_without_the_flag(self, baseline_mtz): + mtz, _ = baseline_mtz + present = set(_read(mtz).columns) & set(TWO_MOMENT_COLUMNS) + assert present == set(), f"unexpected two-moment columns: {sorted(present)}" + + +class TestTwoMomentLayout: + def test_baseline_columns_all_survive(self, two_moment_mtz): + mtz, _ = two_moment_mtz + assert BASELINE_COLUMNS.issubset(set(_read(mtz).columns)) + + def test_every_new_column_is_present_with_the_right_mtz_type(self, two_moment_mtz): + mtz, _ = two_moment_mtz + df = _read(mtz) + for name, expected in TWO_MOMENT_COLUMNS.items(): + assert name in df.columns, f"missing column {name}" + actual = df.dtypes[name].name + assert actual == expected, ( + f"{name} written as {actual}, expected {expected} -- this writer has no " + f"infer_mtz_dtypes() safety net" + ) + + def test_the_column_set_is_exactly_baseline_plus_the_new_ones(self, two_moment_mtz): + mtz, _ = two_moment_mtz + assert set(_read(mtz).columns) == BASELINE_COLUMNS | set(TWO_MOMENT_COLUMNS) + + +class TestTwoMomentValuesAreConsistent: + def test_ivar_alpha_is_sigma_sq_times_the_squared_difference(self, two_moment_mtz): + """The variance column must be the quantity it claims, not a rescaling of it.""" + import numpy as np + + mtz, summary = two_moment_mtz + df = _read(mtz) + results = json.loads(summary.read_text())["results"] + + sigma_sq = results["sigma_alpha_sq"] + dfc = df["DFc_complex"].to_numpy().astype(float) + ivar = df["IVAR_ALPHA"].to_numpy().astype(float) + + expected = sigma_sq * dfc**2 + scale = max(float(np.abs(expected).max()), 1e-30) + assert np.abs(ivar - expected).max() / scale < 1e-5 + + def test_the_two_moment_intensity_exceeds_the_coherent_one_by_the_variance( + self, two_moment_mtz + ): + import numpy as np + + mtz, _ = two_moment_mtz + df = _read(mtz) + coh = df["Ic_light_coh"].to_numpy().astype(float) + two = df["Ic_light_2mom"].to_numpy().astype(float) + ivar = df["IVAR_ALPHA"].to_numpy().astype(float) + + scale = max(float(np.abs(ivar).max()), 1e-30) + assert np.abs((two - coh) - ivar).max() / scale < 1e-3 + # The variance term has no sign: it can only add. + assert (two >= coh - 1e-6).all() + + def test_the_weight_lies_in_zero_to_one_and_bites(self, two_moment_mtz): + w = _read(two_moment_mtz[0])["W_2MOM"].to_numpy().astype(float) + assert (w > 0).all() and (w <= 1.0 + 1e-6).all() + assert w.min() < 0.99, ( + "W_2MOM is 1 everywhere, so the correction is doing nothing here and this " + "fixture cannot detect a change in it" + ) + + def test_the_correction_moves_the_difference_amplitudes(self, two_moment_mtz): + """DDF is the diagnostic; if it were identically zero the whole column set + would be decorative.""" + import numpy as np + + df = _read(two_moment_mtz[0]) + ddf = df["DDF"].to_numpy().astype(float) + assert np.count_nonzero(ddf) > 0.5 * len(ddf) + # Subtracting a positive contamination lowers the light amplitude on average. + assert ddf.mean() < 0.0 + + def test_summary_reports_the_activation_moments(self, two_moment_mtz): + _, summary = two_moment_mtz + results = json.loads(summary.read_text())["results"] + for key in ("alpha_mean", "lambda_twin", "sigma_alpha_sq"): + assert key in results, f"missing summary key: {key}" + assert results["lambda_twin"] == pytest.approx(LAMBDA_TWIN, abs=1e-5) + assert results["alpha_mean"] == pytest.approx(FRACTION, abs=1e-5) + assert results["sigma_alpha_sq"] == pytest.approx( + FRACTION * (1 - FRACTION) * LAMBDA_TWIN, rel=1e-4 + ) diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index b37bec4b..28d1878e 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -19,6 +19,23 @@ :func:`compute_bayes_extrapolated_amplitudes`), each with its sigma and ``2F-Fc`` / ``F-Fc`` map coefficients +Under ``--two-moment`` with a non-zero dispersion, thirteen further columns describe the +activation-heterogeneity correction: + +Intensities ``Io_light``, ``SIGIo_light`` (the quantity actually fitted), + ``Ic_light_coh`` = |F(alpha)|^2, ``Ic_light_2mom`` = the fitted model, + ``IVAR_ALPHA`` = sigma_alpha^2 |dF|^2 on its own +Decontaminated ``Fo_light_corr``, ``SIGFo_light_corr``, ``DF_corr``, ``SIGDF_corr``, + and ``2mDFop-DFc_corr`` / ``mDFop-DFc_corr`` to pair with ``PHIC_diff`` +Diagnostics ``DDF`` = DF_corr - DF, and ``W_2MOM``, the sigma_alpha^2-aware weight + +``DDF`` is the one to look at first. Smooth and featureless against resolution means the +correction is collinear with a scale or overall-B error and should be distrusted; +structure in it is the signal. + +Note the ``m`` in ``2mDFop-DFc`` is a normalised inverse-variance weight, not a sigma_A +figure of merit. + Examples -------- :: @@ -66,6 +83,8 @@ DEFAULT_TARGET_WEIGHTS = { "xray/difference": 1.0, "xray/rice": 0.0, + # Registered only under --two-moment; harmless in the dict either way. + "xray/two_moment": 1.0, # "geometry/bond": 1.0, # geometry restraint should never require tuning, so leave at 1.0 # "geometry/angle": 1.0, # "geometry/torsion": 1.0, @@ -175,11 +194,19 @@ def compute_rfactors(model, data, scaler): def setup_loss_state(dataset_collection, model_collection, scaler, - target_weights, device, similarity_alpha=2.0): + target_weights, device, similarity_alpha=2.0, + two_moment=False): """Build LossState with collection-aware targets. Geometry and ADP restraints are applied only to the light base model (the dark model is a frozen reference). + + Parameters + ---------- + two_moment : bool, optional + Also register the two-moment intensity target, which fits merged intensities + under ``|F(alpha)|^2 + sigma_alpha^2 |dF|^2``. Requires I/SIGI on every + dataset. Default False. """ from torchref.refinement import LossState from torchref.experimental.kinetic.targets import ( @@ -213,6 +240,16 @@ def setup_loss_state(dataset_collection, model_collection, scaler, state.register_target("adp", adp_target) state.register_target("similarity", similarity_target) + if two_moment: + from torchref.refinement.targets import CollectionTwoMomentIntensityTarget + + state.register_target( + "xray/two_moment", + CollectionTwoMomentIntensityTarget( + dataset_collection, model_collection, scaler=scaler, + ), + ) + state.set_weights(target_weights) return state @@ -279,6 +316,118 @@ def compute_bayes_extrapolated_amplitudes( return F_ext, var_ext_bayes, w, tau_sq +def _two_moment_columns(mc, dc, mask, fcalc_dark_full, fcalc_mixed_full, + *, weights, diff_Fobs, Fcalc_diff_amp, Fobs_dark, + sig_dark): + """Two-moment diagnostic columns, or empty dicts when the model is off. + + The observed light intensity carries a positive, phase-blind contamination + ``sigma_alpha^2 |dF|^2`` from the spread of activation across crystals. Subtracting + the model's estimate of it and converting back to an amplitude gives a difference + amplitude that is comparable across datasets, which the raw one is not. + + ``DDF`` is the diagnostic that matters: smooth and featureless against resolution + means the correction is collinear with a scale or overall-B error and should be + distrusted; structure in it is the signal. + + The decontaminated amplitude goes through the dataset's own French-Wilson estimator, + on the **full** reflection list, because subtracting the variance term pushes weak + reflections negative and that is exactly where a naive ``sqrt(clamp(I, 0))`` is worst. + + Parameters + ---------- + mask : torch.Tensor + The dark-and-light validity intersection the writer uses; the returned columns + are already reduced to it. + fcalc_dark_full, fcalc_mixed_full : torch.Tensor + Scaled complex structure factors on the **full** HKL list. + weights, diff_Fobs, Fcalc_diff_amp, Fobs_dark, sig_dark : numpy.ndarray + Masked quantities the writer has already computed, reused so the corrected + columns are constructed exactly like their uncorrected counterparts. + + Returns + ------- + tuple + ``(columns, f_cols, sigma_cols, intensity_cols, weight_cols)`` -- the values plus + the MTZ type each belongs to. This writer assigns types from name lists and never + calls ``infer_mtz_dtypes``, so every column must appear in exactly one list. + """ + import numpy as np + + empty = ({}, [], [], [], []) + if float(mc.sigma_alpha_sq) == 0.0: + return empty + + data_light = dc["light"] + if data_light.I is None: + return empty + + with torch.no_grad(): + # Full-size, so French-Wilson sees the reflection list it was fitted on. + delta_F_full = fcalc_mixed_full - fcalc_dark_full + variance_full = mc.sigma_alpha_sq * delta_F_full.abs() ** 2 + + I_light_full, sig_I_full = data_light.get_corrected_intensities() + I_corrected_full = I_light_full - variance_full + + # The retained estimator is fitted on the dataset's HKL list *as loaded*; + # joining a collection expands the dataset onto the common grid, so it can be + # the wrong length by then. Rebuild against the current list when that happens. + fw = data_light._FrenchWilson + if fw is None or len(fw.d_spacings) != len(I_corrected_full): + from torchref.base.french_wilson import FrenchWilson + + fw = FrenchWilson( + data_light.hkl, data_light.cell, data_light.spacegroup, verbose=0 + ) + F_corr_full, sig_F_corr_full = fw(I_corrected_full, sig_I_full) + + def _np(t): + return t[mask].detach().cpu().numpy() + + variance = _np(variance_full) + I_light = _np(I_light_full) + sig_I_light = _np(sig_I_full) + I_coherent = _np(fcalc_mixed_full.abs() ** 2) + F_corr = _np(F_corr_full) + sig_F_corr = _np(sig_F_corr_full) + + I_two_moment = I_coherent + variance + DF_corr = F_corr - Fobs_dark + DDF = DF_corr - diff_Fobs + sig_DF_corr = np.sqrt(sig_F_corr**2 + sig_dark**2) + + amp_2_corr = (2 * np.abs(DF_corr) - Fcalc_diff_amp) * weights + amp_1_corr = (np.abs(DF_corr) - Fcalc_diff_amp) * weights + + # The sigma_alpha^2-aware weight, on the same normalisation as the inverse-variance + # weight the existing DED coefficients carry, so the two are directly comparable. + w_two_moment = sig_I_light**2 / np.maximum(sig_I_light**2 + variance, 1e-12) + + columns = { + "Io_light": I_light, + "SIGIo_light": sig_I_light, + "Ic_light_coh": I_coherent, + "Ic_light_2mom": I_two_moment, + "IVAR_ALPHA": variance, + "Fo_light_corr": F_corr, + "SIGFo_light_corr": sig_F_corr, + "DF_corr": DF_corr, + "SIGDF_corr": sig_DF_corr, + "2mDFop-DFc_corr": amp_2_corr, + "mDFop-DFc_corr": amp_1_corr, + "DDF": DDF, + "W_2MOM": w_two_moment, + } + f_cols = [ + "Fo_light_corr", "DF_corr", "2mDFop-DFc_corr", "mDFop-DFc_corr", "DDF", + ] + sigma_cols = ["SIGIo_light", "SIGFo_light_corr", "SIGDF_corr"] + intensity_cols = ["Io_light", "Ic_light_coh", "Ic_light_2mom", "IVAR_ALPHA"] + weight_cols = ["W_2MOM"] + return columns, f_cols, sigma_cols, intensity_cols, weight_cols + + def write_results_mtz(dc, mc, scaler, filename): """Write difference / extrapolated map coefficients to an MTZ file. @@ -320,8 +469,14 @@ def write_results_mtz(dc, mc, scaler, filename): # Compute Fcalc on full HKL then mask (scalers fitted on full datasets) with torch.no_grad(): - fcalc_dark = scaler.forward_mixed(dark_model(hkl_all), dark_model.fractions)[mask] - fcalc_mixed = scaler.forward_mixed(mixed_model(hkl_all), mixed_model.fractions)[mask] + fcalc_dark_full = scaler.forward_mixed( + dark_model(hkl_all), dark_model.fractions + ) + fcalc_mixed_full = scaler.forward_mixed( + mixed_model(hkl_all), mixed_model.fractions + ) + fcalc_dark = fcalc_dark_full[mask] + fcalc_mixed = fcalc_mixed_full[mask] fcalc_diff = fcalc_mixed - fcalc_dark phi_dark = torch.angle(fcalc_dark) @@ -436,6 +591,15 @@ def _extrapolation_rfactors(data, fcalc_scaled): amp_2DFoDFc = (2 * Fobs_diff_phased - Fcalc_diff_amp) * weights amp_DFoDFc = (Fobs_diff_phased - Fcalc_diff_amp) * weights + two_moment_columns, two_moment_f, two_moment_sig, two_moment_j, two_moment_w = ( + _two_moment_columns( + mc, dc, mask, fcalc_dark_full, fcalc_mixed_full, + weights=weights, diff_Fobs=diff_Fobs, + Fcalc_diff_amp=Fcalc_diff_amp, Fobs_dark=Fobs_dark, + sig_dark=sig_dark, + ) + ) + df = rs.DataSet( { "H": hkl_np[:, 0], "K": hkl_np[:, 1], "L": hkl_np[:, 2], @@ -467,6 +631,7 @@ def _extrapolation_rfactors(data, fcalc_scaled): "SIGFextb": sig_ext_bayes.detach().cpu().numpy(), "2Fextb-Fc": amp_2fofc_bayes.detach().cpu().numpy(), "Fextb-Fc": amp_fofc_bayes.detach().cpu().numpy(), + **two_moment_columns, # R-free flags (1=work, 0=free) "FreeR_flag_dark": ( data_dark.rfree_flags[mask].cpu().numpy().astype(int) @@ -492,9 +657,17 @@ def _extrapolation_rfactors(data, fcalc_scaled): "Fextc", "2Fextc-Fc", "Fextc-Fc", "Fextb", "2Fextb-Fc", "Fextb-Fc", ] + f_cols += two_moment_f df[f_cols] = df[f_cols].astype("F") sig_cols = ["SIGFo_dark", "SIGFo_light", "SIGDF", "SIGFextc", "SIGFextb"] + sig_cols += two_moment_sig df[sig_cols] = df[sig_cols].astype("Q") + # This writer never calls infer_mtz_dtypes(), so a column absent from every list + # above would be written with whatever dtype numpy produced. + if two_moment_j: + df[two_moment_j] = df[two_moment_j].astype("J") + if two_moment_w: + df[two_moment_w] = df[two_moment_w].astype("W") phase_cols = ["PHIC_dark", "PHIC_mixed", "PHIC_diff", "PHIC_light"] df[phase_cols] = df[phase_cols].astype("P") df["FreeR_flag_dark"] = df["FreeR_flag_dark"].astype("I") @@ -596,6 +769,26 @@ def main(): help="Weight for dark/light coordinate similarity restraint " "(0 to disable, default: 1.0)", ) + two_moment = parser.add_argument_group("Activation heterogeneity (two-moment model)") + two_moment.add_argument( + "--two-moment", action="store_true", default=False, + help="Fit merged intensities with |F(alpha)|^2 + sigma_alpha^2 |dF|^2, " + "which accounts for crystal-to-crystal spread in activation. " + "Requires I/SIGI columns in both reflection files.", + ) + two_moment.add_argument( + "--lambda-twin", type=float, default=0.0, + help="Activation dispersion as a fraction of its maximum, in [0, 1]: " + "sigma_alpha^2 = alpha (1 - alpha) * lambda. 0 (default) is the " + "coherent model and reproduces the amplitude-only result.", + ) + two_moment.add_argument( + "--refine-lambda-twin", action="store_true", default=False, + help="Refine --lambda-twin instead of holding it fixed. Off by default: " + "the sigma_alpha^2 term is smooth and positive, so it is collinear " + "with a scale or overall-B error and can absorb one.", + ) + refine.add_argument( "--similarity-alpha", type=float, default=2.0, help="Log prior odds for spike-and-slab similarity restraint. " @@ -617,6 +810,26 @@ def main(): return 1 fractions = [1.0 - args.fraction, args.fraction] + if not (0.0 <= args.lambda_twin <= 1.0): + print( + f"Error: --lambda-twin must lie in [0, 1] (got {args.lambda_twin})", + file=sys.stderr, + ) + return 1 + if args.refine_lambda_twin and not args.two_moment: + print( + "Error: --refine-lambda-twin needs --two-moment; the activation " + "dispersion only enters through the two-moment intensity model", + file=sys.stderr, + ) + return 1 + if args.lambda_twin > 0.0 and not args.two_moment: + print( + "Error: --lambda-twin has no effect without --two-moment", + file=sys.stderr, + ) + return 1 + # --- Parse weight schedule --- try: weight_schedule = [float(x) for x in args.weight_schedule.split(",")] @@ -666,6 +879,9 @@ def main(): print(f"Light data: {args.light_structure_factor}") frac_mode = "refinable" if args.refine_fractions else "frozen" print(f"Fractions: dark={fractions[0]}, light={fractions[1]} ({frac_mode})") + if args.two_moment: + lam_mode = "refinable" if args.refine_lambda_twin else "fixed" + print(f"Two-moment model: lambda_twin={args.lambda_twin} ({lam_mode})") print(f"Output: {outdir}") print(f"Device: {device}") if args.dmin: @@ -747,8 +963,23 @@ def main(): sys.stdout.flush() # --- Setup targets --- + if args.two_moment: + missing = [k for k in dc.keys() if dc[k].I is None] + if missing: + print( + f"Error: --two-moment needs I/SIGI columns, but {missing} carry " + f"only amplitudes. Converting back with F**2 would reintroduce the " + f"French-Wilson distortion the intensity model exists to avoid.", + file=sys.stderr, + ) + return 1 + mc.set_lambda_twin( + args.lambda_twin, refinable=args.refine_lambda_twin + ) + state = setup_loss_state(dc, mc, scaler, target_weights, device, - similarity_alpha=args.similarity_alpha) + similarity_alpha=args.similarity_alpha, + two_moment=args.two_moment) if args.verbose > 0: print("Initial loss breakdown:") @@ -765,6 +996,10 @@ def main(): )) else: params = list(model_light.parameters()) + if args.refine_lambda_twin: + # fraction_parameters() carries lambda once it is refinable; take only + # that, since the fractions themselves stay frozen here. + params.append(mc._lambda_logit) fraction_history = [] if args.refine_fractions: @@ -816,14 +1051,17 @@ def main(): sys.stdout.flush() # --- Final statistics --- + # Computed unconditionally: the deposition metadata and the results MTZ both carry + # these, so they are not a reporting-only quantity. + r_work_d, r_free_d = compute_rfactors(dark, data_dark, scaler) + r_work_l, r_free_l = compute_rfactors(mixed, data_light, scaler) + r_work_dl, r_free_dl = compute_rfactors(dark, data_light, scaler) + if args.verbose > 0: print() print("=" * 72) print("Refinement complete") print("=" * 72) - r_work_d, r_free_d = compute_rfactors(dark, data_dark, scaler) - r_work_l, r_free_l = compute_rfactors(mixed, data_light, scaler) - r_work_dl, r_free_dl = compute_rfactors(dark, data_light, scaler) print(f" Final R-factor (dark vs dark data): R_work={r_work_d:.4f} R_free={r_free_d:.4f}") print(f" Final R-factor (mixed vs light data): R_work={r_work_l:.4f} R_free={r_free_l:.4f}") print(f" Final R-factor (dark vs light data): R_work={r_work_dl:.4f} R_free={r_free_dl:.4f}") @@ -1063,6 +1301,9 @@ def _mtz_to_cif(mtz_path, cif_path): compute_rfactors(mixed, data_light, scaler), )), "fractions": mixed.fractions.detach().cpu().tolist(), + "alpha_mean": float(mc.alpha_mean), + "lambda_twin": float(mc.lambda_twin), + "sigma_alpha_sq": float(mc.sigma_alpha_sq), }, "output_files": { "dark_pdb": dark_pdb_out, From d76da0500bc81202c29277802e661e2fb1efc42a Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 24 Aug 2026 16:16:59 +0200 Subject: [PATCH 053/250] Add CrystFEL hkl reading and intensity simulation The validation ladder for the two-moment model needs simulated merged intensities with known moments, and needs to read the format real merged serial data arrives in. Both were written on an unmerged branch; this ports them onto the current accessors. - torchref/io/hkl.py + ReflectionData.load_crystfel_hkl read partialator .hkl lists. Intensity-native, so amplitudes come from the same French-Wilson path an MTZ with I/SIGI takes. Cell and space group must be supplied: the format carries neither. - FcalcDataset.add_noise draws two independent noisy half-datasets and returns their mean, reporting R-split and CC between them, with sigmas either grafted per-reflection from a real reference or from a three-term parametric model. - torchref.simulate-noisy-data drives it from a structure file, and is now registered as a console script. Simulated intensities keep their negatives. Only the derived amplitude is clamped, because an amplitude cannot be negative; clamping the intensity instead puts a positive bias on exactly the weak reflections where noise dominates -- the same signature as a real positive perturbation of the merged intensity, and systematic, so it survives averaging where the noise does not. The test fixtures are excerpts of a real partialator custom-split pair, which carry both negative intensities and sigma(I) == 0 -- the two things a simulation study built on this data has to handle rather than assume away. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- docs/changelog.rst | 3 + pyproject.toml | 1 + tests/files/hkl/dark_half1.hkl | 404 ++++++++++++++++++++++++ tests/files/hkl/dark_half2.hkl | 404 ++++++++++++++++++++++++ tests/unit/io/test_crystfel_hkl.py | 133 ++++++++ tests/unit/io/test_fcalc_add_noise.py | 220 +++++++++++++ torchref/cli/simulate_noisy_data.py | 282 +++++++++++++++++ torchref/io/datasets/fcalc_data.py | 160 +++++++++- torchref/io/datasets/reflection_data.py | 31 ++ torchref/io/hkl.py | 146 +++++++++ 10 files changed, 1783 insertions(+), 1 deletion(-) create mode 100644 tests/files/hkl/dark_half1.hkl create mode 100644 tests/files/hkl/dark_half2.hkl create mode 100644 tests/unit/io/test_crystfel_hkl.py create mode 100644 tests/unit/io/test_fcalc_add_noise.py create mode 100644 torchref/cli/simulate_noisy_data.py create mode 100644 torchref/io/hkl.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 9cd16764..e54ca2ae 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,9 @@ Changelog Version 0.6.4 ---------- +- Added a reader for CrystFEL ``partialator`` ``.hkl`` reflection lists, via ``ReflectionData.load_crystfel_hkl`` +- Added ``FcalcDataset.add_noise`` and the ``torchref.simulate-noisy-data`` CLI, which simulate merged intensities from a structure and report R-split and CC between two independent half-datasets +- Simulated intensities keep their negative values; only the derived amplitude is clamped, since clamping the intensity biases the weak reflections upward - Added ``CollectionTwoMomentIntensityTarget``, fitting merged intensities as ``|F(alpha)|^2 + sigma_alpha^2 |dF|^2`` to account for crystal-to-crystal spread in activation - Added ``--two-moment`` / ``--lambda-twin`` / ``--refine-lambda-twin`` to ``torchref.difference-refine``, and the activation moments to its JSON summary - ``torchref.difference-refine`` writes thirteen further MTZ columns under ``--two-moment``, including decontaminated difference amplitudes and the ``DDF`` diagnostic diff --git a/pyproject.toml b/pyproject.toml index 4f0c7ba6..dc9526bb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -45,6 +45,7 @@ dependencies = [ [project.scripts] "torchref.refine" = "torchref.cli.refine:main" "torchref.difference-refine" = "torchref.cli.collection_difference_refine:main" +"torchref.simulate-noisy-data" = "torchref.cli.simulate_noisy_data:main" "torchref.mtz2map" = "torchref.cli.mtz2map:main" "torchref.validate-ded" = "torchref.cli.validate_ded:main" "torchref.phased-difference-map" = "torchref.cli.phased_difference_map:main" diff --git a/tests/files/hkl/dark_half1.hkl b/tests/files/hkl/dark_half1.hkl new file mode 100644 index 00000000..8aada321 --- /dev/null +++ b/tests/files/hkl/dark_half1.hkl @@ -0,0 +1,404 @@ +CrystFEL reflection list version 2.0 +Symmetry: 1 + h k l I phase sigma(I) nmeas + -17 -13 -5 -10.07 - 7.18 2 + -17 -13 -4 14.43 - 4.44 4 + -17 -13 -3 -9.91 - 2.98 2 + -17 -13 -2 5.70 - 4.37 4 + -17 -13 1 4.74 - 2.34 3 + -17 -13 2 5.88 - 5.79 2 + -17 -12 -6 10.69 - 5.24 2 + -17 -12 -5 15.15 - 2.69 2 + -17 -12 -4 15.51 - 7.56 4 + -17 -12 -3 2.12 - 3.08 10 + -17 -12 -2 6.04 - 3.10 11 + -17 -12 -1 10.29 - 8.15 9 + -17 -12 0 2.66 - 3.31 10 + -17 -12 1 0.00 - 2.15 3 + -17 -12 2 18.39 - 4.62 2 + -17 -11 -6 8.83 - 5.40 3 + -17 -11 -5 7.96 - 4.29 7 + -17 -11 -4 7.00 - 6.07 17 + -17 -11 -3 2.62 - 1.95 32 + -17 -11 -2 24.27 - 11.91 45 + -17 -11 -1 6.48 - 2.42 43 + -17 -11 0 3.19 - 1.74 13 + -17 -11 1 4.22 - 6.52 7 + -17 -11 2 2.07 - 6.35 2 + -17 -11 4 9.21 - 0.30 2 + -17 -10 -5 6.78 - 2.00 16 + -17 -10 -4 4.97 - 1.07 42 + -17 -10 -3 2.91 - 2.63 74 + -17 -10 -2 4.43 - 1.18 105 + -17 -10 -1 32.21 - 23.72 87 + -17 -10 0 2.01 - 0.95 41 + -17 -10 1 1.33 - 1.23 21 + -17 -10 2 6.97 - 4.70 2 + -17 -9 -7 0.70 - 3.00 2 + -17 -9 -6 7.46 - 4.00 6 + -17 -9 -5 14.04 - 7.81 46 + -17 -9 -4 7.94 - 2.36 83 + -17 -9 -3 4.45 - 0.76 122 + -17 -9 -2 4.71 - 1.77 143 + -17 -9 -1 7.66 - 2.75 144 + -17 -9 0 4.54 - 1.02 88 + -17 -9 1 2.18 - 1.21 51 + -17 -9 2 2.04 - 2.29 11 + -17 -9 3 0.00 - 0.00 2 + -17 -8 -9 9.29 - 1.64 2 + -17 -8 -7 1.44 - 2.78 4 + -17 -8 -6 7.10 - 2.70 15 + -17 -8 -5 3.66 - 1.10 40 + -17 -8 -4 3.11 - 0.75 95 + -17 -8 -3 1.70 - 0.84 142 + -17 -8 -2 2.28 - 0.60 172 + -17 -8 -1 1.95 - 0.75 132 + -17 -8 0 1.80 - 0.70 102 + -17 -8 1 2.66 - 1.37 61 + -17 -8 2 0.33 - 3.93 10 + -17 -8 3 7.18 - 1.19 3 + -17 -8 4 14.27 - 6.28 3 + -17 -7 -6 4.23 - 0.43 5 + -17 -7 -5 3.87 - 1.08 37 + -17 -7 -4 2.36 - 0.81 80 + -17 -7 -3 6.30 - 2.51 120 + -17 -7 -2 6.24 - 3.40 112 + -17 -7 -1 13.54 - 7.01 102 + -17 -7 0 8.99 - 5.02 93 + -17 -7 1 1.87 - 1.16 40 + -17 -7 2 7.35 - 3.79 6 + -17 -7 3 1.74 - 1.23 2 + -17 -7 4 1.48 - 2.82 3 + -17 -7 5 6.91 - 0.89 2 + -17 -6 -9 2.50 - 1.77 2 + -17 -6 -7 15.72 - 6.99 2 + -17 -6 -6 2.93 - 2.33 3 + -17 -6 -5 2.08 - 2.37 15 + -17 -6 -4 4.24 - 1.33 46 + -17 -6 -3 2.49 - 0.80 76 + -17 -6 -2 3.86 - 0.83 85 + -17 -6 -1 2.79 - 0.75 63 + -17 -6 0 2.66 - 1.83 48 + -17 -6 1 2.37 - 1.21 30 + -17 -6 2 6.01 - 2.68 8 + -17 -6 3 9.91 - 3.11 2 + -17 -6 4 4.26 - 1.16 2 + -17 -5 -5 8.22 - 3.02 4 + -17 -5 -4 4.54 - 2.19 13 + -17 -5 -3 8.71 - 5.75 26 + -17 -5 -2 7.56 - 1.88 26 + -17 -5 -1 4.39 - 1.39 32 + -17 -5 0 4.44 - 1.52 17 + -17 -5 1 6.89 - 2.83 8 + -17 -5 2 7.73 - 3.95 3 + -17 -4 -7 11.59 - 2.31 2 + -17 -4 -6 5.99 - 1.83 2 + -17 -4 -5 1.26 - 5.44 4 + -17 -4 -4 3.27 - 2.36 5 + -17 -4 -3 -0.33 - 2.34 6 + -17 -4 -2 -1.17 - 3.84 5 + -17 -4 -1 6.30 - 1.97 10 + -17 -4 0 7.19 - 5.90 5 + -17 -4 1 2.61 - 1.56 6 + -17 -4 2 6.19 - 1.32 2 + -17 -3 -9 3.62 - 0.50 2 + -17 -3 -7 7.37 - 0.66 2 + -17 -3 -3 -2.38 - 0.07 2 + -17 -3 -1 -1.82 - 7.23 3 + -17 -3 0 2.44 - 3.39 4 + -17 -2 -6 -3.27 - 2.31 2 + -17 -2 -1 3.73 - 2.64 2 + -16 -17 0 5.73 - 2.44 3 + -16 -16 -3 0.04 - 3.11 2 + -16 -16 -2 7.53 - 6.52 4 + -16 -16 -1 3.79 - 3.47 5 + -16 -16 0 -16.21 - 5.92 2 + -16 -15 -5 2.74 - 3.34 2 + -16 -15 -4 2.57 - 3.46 13 + -16 -15 -3 3.62 - 2.01 24 + -16 -15 -2 5.66 - 2.46 33 + -16 -15 -1 3.22 - 1.12 32 + -16 -15 0 2.68 - 1.49 16 + -16 -15 1 4.90 - 2.61 6 + -16 -15 3 1.45 - 1.02 2 + -16 -14 -8 -0.24 - 0.17 2 + -16 -14 -7 -6.09 - 5.96 3 + -16 -14 -6 2.93 - 1.73 15 + -16 -14 -5 57.87 - 31.51 75 + -16 -14 -4 4.77 - 1.66 139 + -16 -14 -3 3.93 - 0.74 181 + -16 -14 -2 15.65 - 6.21 199 + -16 -14 -1 4.89 - 1.07 225 + -16 -14 0 3.30 - 0.62 176 + -16 -14 1 5.53 - 1.34 155 + -16 -14 2 4.94 - 2.18 56 + -16 -14 3 3.08 - 2.09 18 + -16 -14 4 9.74 - 2.34 4 + -16 -13 -9 8.32 - 1.08 2 + -16 -13 -8 2.90 - 2.02 2 + -16 -13 -7 2.63 - 1.14 36 + -16 -13 -6 3.87 - 0.80 139 + -16 -13 -5 2.93 - 0.67 222 + -16 -13 -4 4.75 - 0.69 265 + -16 -13 -3 5.95 - 1.49 340 + -16 -13 -2 6.29 - 0.73 322 + -16 -13 -1 12.39 - 2.90 339 + -16 -13 0 8.12 - 1.28 304 + -16 -13 1 6.97 - 2.12 265 + -16 -13 2 3.26 - 0.82 212 + -16 -13 3 1.79 - 0.82 101 + -16 -13 4 2.65 - 1.36 17 + -16 -13 5 -6.29 - 5.03 4 + -16 -12 -9 1.81 - 2.24 2 + -16 -12 -8 4.59 - 1.95 48 + -16 -12 -7 7.32 - 1.70 193 + -16 -12 -6 5.20 - 1.36 257 + -16 -12 -5 4.87 - 0.67 368 + -16 -12 -4 8.08 - 1.21 406 + -16 -12 -3 32.92 - 6.61 461 + -16 -12 -2 7.34 - 0.87 531 + -16 -12 -1 21.04 - 4.37 497 + -16 -12 0 24.55 - 5.30 409 + -16 -12 1 15.18 - 2.90 376 + -16 -12 2 12.26 - 2.72 326 + -16 -12 3 5.78 - 1.04 256 + -16 -12 4 2.34 - 0.63 119 + -16 -12 5 4.64 - 1.97 28 + -16 -12 6 4.70 - 10.61 2 + -16 -12 7 1.49 - 0.79 3 + -16 -11 -9 4.21 - 1.95 14 + -16 -11 -8 6.35 - 2.53 150 + -16 -11 -7 4.21 - 0.64 266 + -16 -11 -6 11.72 - 2.30 383 + -16 -11 -5 4.58 - 0.63 477 + -16 -11 -4 27.98 - 4.12 528 + -16 -11 -3 6.26 - 0.59 513 + -16 -11 -2 16.39 - 2.52 477 + -16 -11 -1 5.46 - 0.60 530 + -16 -11 0 5.77 - 0.73 527 + -16 -11 1 5.39 - 0.64 421 + -16 -11 2 6.14 - 0.66 392 + -16 -11 3 4.43 - 0.71 321 + -16 -11 4 5.32 - 0.80 228 + -16 -11 5 1.65 - 0.81 80 + -16 -11 6 1.42 - 0.91 4 + -16 -10 -11 0.08 - 0.47 2 + -16 -10 -10 -0.30 - 2.94 8 + -16 -10 -9 1.00 - 0.87 70 + -16 -10 -8 12.49 - 4.95 223 + -16 -10 -7 7.06 - 0.90 367 + -16 -10 -6 27.49 - 5.44 433 + -16 -10 -5 33.74 - 6.21 473 + -16 -10 -4 6.38 - 0.67 561 + -16 -10 -3 25.30 - 4.53 492 + -16 -10 -2 6.34 - 1.71 557 + -16 -10 -1 6.54 - 0.59 567 + -16 -10 0 7.48 - 1.00 534 + -16 -10 1 31.68 - 6.12 511 + -16 -10 2 8.86 - 0.95 500 + -16 -10 3 18.62 - 3.25 365 + -16 -10 4 22.28 - 5.35 263 + -16 -10 5 4.10 - 0.75 150 + -16 -10 6 3.34 - 1.77 21 + -16 -10 7 7.48 - 4.88 4 + -16 -10 8 -6.25 - 5.20 2 + -16 -9 -10 -0.73 - 7.97 3 + -16 -9 -9 6.68 - 3.91 136 + -16 -9 -8 9.02 - 2.51 288 + -16 -9 -7 5.89 - 0.68 401 + -16 -9 -6 5.07 - 0.60 481 + -16 -9 -5 5.42 - 0.60 501 + -16 -9 -4 5.60 - 0.57 542 + -16 -9 -3 6.28 - 0.76 539 + -16 -9 -2 9.89 - 1.49 593 + -16 -9 -1 5.58 - 0.65 576 + -16 -9 0 12.37 - 2.41 562 + -16 -9 1 4.82 - 0.62 478 + -16 -9 2 4.75 - 0.82 529 + -16 -9 3 6.86 - 0.93 425 + -16 -9 4 7.74 - 1.16 307 + -16 -9 5 3.43 - 0.68 199 + -16 -9 6 5.20 - 1.76 31 + -16 -8 -12 -9.19 - 2.70 2 + -16 -8 -10 2.88 - 1.93 10 + -16 -8 -9 2.98 - 0.81 153 + -16 -8 -8 27.50 - 7.96 301 + -16 -8 -7 8.60 - 1.53 414 + -16 -8 -6 6.43 - 0.58 516 + -16 -8 -5 5.37 - 1.33 511 + -16 -8 -4 23.66 - 4.29 563 + -16 -8 -3 29.63 - 5.40 580 + -16 -8 -2 24.09 - 3.55 620 + -16 -8 -1 50.29 - 6.68 580 + -16 -8 0 4.63 - 0.54 590 + -16 -8 1 40.67 - 6.53 539 + -16 -8 2 11.90 - 2.03 504 + -16 -8 3 4.33 - 0.66 409 + -16 -8 4 8.08 - 1.40 352 + -16 -8 5 13.33 - 3.83 246 + -16 -8 6 7.96 - 5.55 66 + -16 -7 -10 2.26 - 2.60 13 + -16 -7 -9 3.35 - 0.77 165 + -16 -7 -8 3.43 - 0.63 309 + -16 -7 -7 14.22 - 2.00 410 + -16 -7 -6 15.12 - 18.60 493 + -16 -7 -5 10.13 - 1.46 536 + -16 -7 -4 17.59 - 3.70 529 + -16 -7 -3 6.24 - 1.17 585 + -16 -7 -2 11.53 - 1.69 582 + -16 -7 -1 7.59 - 1.32 544 + -16 -7 0 4.52 - 0.58 539 + -16 -7 1 2.89 - 0.56 534 + -16 -7 2 7.32 - 1.06 503 + -16 -7 3 6.19 - 0.77 446 + -16 -7 4 6.64 - 0.62 345 + -16 -7 5 2.81 - 0.56 218 + -16 -7 6 4.32 - 0.87 53 + -16 -7 7 1.16 - 3.64 5 + -16 -6 -11 6.37 - 2.40 2 + -16 -6 -10 5.35 - 2.71 10 + -16 -6 -9 3.41 - 0.71 126 + -16 -6 -8 5.44 - 0.94 282 + -16 -6 -7 10.30 - 1.98 366 + -16 -6 -6 42.37 - 27.86 477 + -16 -6 -5 6.30 - 0.67 540 + -16 -6 -4 24.61 - 4.06 541 + -16 -6 -3 12.28 - 3.36 577 + -16 -6 -2 6.24 - 0.58 598 + -16 -6 -1 12.54 - 1.63 557 + -16 -6 0 13.07 - 2.59 539 + -16 -6 1 12.44 - 1.77 519 + -16 -6 2 8.05 - 1.50 480 + -16 -6 3 32.53 - 7.13 422 + -16 -6 4 4.86 - 0.64 347 + -16 -6 5 4.86 - 1.06 188 + -16 -6 6 12.70 - 6.43 21 + -16 -6 7 -0.97 - 6.87 5 + -16 -5 -11 5.75 - 4.07 2 + -16 -5 -10 2.49 - 0.91 5 + -16 -5 -9 12.74 - 9.60 87 + -16 -5 -8 3.85 - 0.59 257 + -16 -5 -7 13.73 - 3.60 374 + -16 -5 -6 5.19 - 0.63 416 + -16 -5 -5 5.52 - 0.58 557 + -16 -5 -4 6.43 - 0.63 522 + -16 -5 -3 5.02 - 0.77 529 + -16 -5 -2 9.96 - 1.20 504 + -16 -5 -1 10.63 - 2.22 537 + -16 -5 0 13.53 - 2.63 517 + -16 -5 1 7.05 - 0.66 476 + -16 -5 2 10.86 - 1.70 450 + -16 -5 3 6.41 - 0.79 325 + -16 -5 4 4.33 - 0.65 292 + -16 -5 5 2.09 - 0.68 124 + -16 -5 6 5.90 - 3.45 9 + -16 -5 7 5.56 - 4.54 3 + -16 -4 -11 15.62 - 4.13 2 + -16 -4 -9 8.48 - 3.82 16 + -16 -4 -8 7.32 - 3.28 155 + -16 -4 -7 3.70 - 0.63 283 + -16 -4 -6 7.38 - 0.85 375 + -16 -4 -5 6.49 - 0.91 409 + -16 -4 -4 22.51 - 5.30 473 + -16 -4 -3 2.94 - 0.49 522 + -16 -4 -2 41.02 - 8.50 517 + -16 -4 -1 9.17 - 1.13 546 + -16 -4 0 19.96 - 3.83 517 + -16 -4 1 11.34 - 1.70 425 + -16 -4 2 6.27 - 0.71 362 + -16 -4 3 5.25 - 0.71 291 + -16 -4 4 7.04 - 1.26 181 + -16 -4 5 2.62 - 1.20 44 + -16 -4 6 -0.25 - 1.68 3 + -16 -3 -10 4.89 - 4.38 3 + -16 -3 -9 -3.72 - 3.45 4 + -16 -3 -8 1.13 - 1.14 44 + -16 -3 -7 7.39 - 3.57 157 + -16 -3 -6 3.91 - 0.64 269 + -16 -3 -5 18.12 - 3.59 340 + -16 -3 -4 6.67 - 0.73 351 + -16 -3 -3 9.10 - 2.37 440 + -16 -3 -2 5.56 - 0.66 399 + -16 -3 -1 5.63 - 0.60 443 + -16 -3 0 5.69 - 0.71 393 + -16 -3 1 5.54 - 0.68 389 + -16 -3 2 9.56 - 2.84 289 + -16 -3 3 2.93 - 0.70 180 + -16 -3 4 3.15 - 1.22 61 + -16 -3 5 8.88 - 3.87 4 + -16 -3 6 9.05 - 4.62 5 + -16 -2 -9 7.44 - 6.65 3 + -16 -2 -8 3.88 - 3.04 7 + -16 -2 -7 8.86 - 5.62 45 + -16 -2 -6 6.48 - 2.51 127 + -16 -2 -5 2.42 - 0.57 227 + -16 -2 -4 12.81 - 2.51 298 + -16 -2 -3 3.75 - 0.66 306 + -16 -2 -2 7.94 - 1.71 332 + -16 -2 -1 6.14 - 1.27 314 + -16 -2 0 6.29 - 1.17 288 + -16 -2 1 1.62 - 0.64 245 + -16 -2 2 9.66 - 4.05 148 + -16 -2 3 17.37 - 11.08 49 + -16 -2 4 1.42 - 3.86 8 + -16 -2 6 -0.64 - 0.18 2 + -16 -1 -8 19.30 - 7.88 3 + -16 -1 -7 4.19 - 3.65 4 + -16 -1 -6 0.35 - 1.62 16 + -16 -1 -5 3.49 - 1.05 53 + -16 -1 -4 1.94 - 0.81 103 + -16 -1 -3 4.49 - 0.71 183 + -16 -1 -2 2.51 - 0.60 181 + -16 -1 -1 4.50 - 0.97 172 + -16 -1 0 1.98 - 0.84 91 + -16 -1 1 3.88 - 0.86 79 + -16 -1 2 3.01 - 1.81 19 + -16 -1 3 3.88 - 2.74 2 + -16 0 -7 1.87 - 1.04 2 + -16 0 -5 2.61 - 2.42 7 + -16 0 -4 2.45 - 1.32 5 + -16 0 -3 60.07 - 50.26 13 + -16 0 -2 7.64 - 2.67 16 + -16 0 -1 2.76 - 1.76 15 + -16 0 0 4.36 - 2.23 6 + -16 0 1 4.28 - 1.67 5 + -16 0 2 8.25 - 4.05 2 + -16 1 -2 7.69 - 3.71 3 + -16 1 -1 8.01 - 5.66 2 + -16 1 1 10.56 - 8.25 2 + -15 -17 -5 4.57 - 3.90 6 + -15 -17 -4 2.52 - 3.10 6 + -15 -17 -3 2.82 - 2.11 22 + -15 -17 -2 24.19 - 14.43 27 + -15 -17 -1 3.23 - 1.49 21 + -15 -17 0 5.54 - 1.59 20 + -15 -17 1 5.68 - 1.28 8 + -15 -17 2 6.59 - 2.11 2 + -15 -17 3 9.74 - 3.07 2 + -15 -16 -8 6.08 - 2.15 3 + -15 -16 -7 5.06 - 4.13 3 + -15 -16 -6 17.08 - 12.72 30 + -15 -16 -5 4.26 - 0.84 110 + -15 -16 -4 15.42 - 6.05 193 + -15 -16 -3 3.02 - 0.55 255 + -15 -16 -2 5.65 - 0.63 258 + -15 -16 -1 4.77 - 0.68 256 + -15 -16 0 12.35 - 2.90 239 + -15 -16 1 3.79 - 0.68 184 + -15 -16 2 4.03 - 1.06 131 + -15 -16 3 3.77 - 1.07 45 + -15 -16 4 5.10 - 1.89 11 + -15 -15 -9 0.64 - 2.56 3 + -15 -15 -8 5.36 - 1.91 26 + -15 -15 -7 2.01 - 0.69 128 + -15 -15 -6 26.75 - 29.49 267 + -15 -15 -5 5.13 - 0.64 344 + -15 -15 -4 14.77 - 3.41 380 + -15 -15 -3 12.57 - 2.14 437 + -15 -15 -2 6.15 - 0.65 427 + -15 -15 -1 5.21 - 0.80 416 + -15 -15 0 22.38 - 3.54 418 + -15 -15 1 7.55 - 2.63 380 + -15 -15 2 13.32 - 2.72 335 + -15 -15 3 24.05 - 6.08 256 +End of reflections diff --git a/tests/files/hkl/dark_half2.hkl b/tests/files/hkl/dark_half2.hkl new file mode 100644 index 00000000..e7fe0d0c --- /dev/null +++ b/tests/files/hkl/dark_half2.hkl @@ -0,0 +1,404 @@ +CrystFEL reflection list version 2.0 +Symmetry: 1 + h k l I phase sigma(I) nmeas + -17 -14 -2 -4.16 - 4.19 3 + -17 -13 -7 -1.84 - 1.30 2 + -17 -13 -5 -1.97 - 3.22 2 + -17 -13 -2 -0.01 - 0.01 4 + -17 -13 -1 1.92 - 2.90 4 + -17 -13 0 11.95 - 1.25 2 + -17 -12 -5 12.59 - 3.63 4 + -17 -12 -4 3.79 - 6.12 4 + -17 -12 -3 0.02 - 4.19 6 + -17 -12 -2 0.80 - 2.51 10 + -17 -12 -1 7.20 - 6.85 5 + -17 -12 0 -0.76 - 1.98 7 + -17 -12 1 1.93 - 1.35 3 + -17 -12 3 -3.52 - 2.49 2 + -17 -11 -5 1.73 - 3.39 5 + -17 -11 -4 5.29 - 2.11 18 + -17 -11 -3 1.27 - 1.79 29 + -17 -11 -2 41.30 - 18.19 43 + -17 -11 -1 3.22 - 1.42 35 + -17 -11 0 5.20 - 2.50 30 + -17 -11 1 11.89 - 4.65 10 + -17 -11 2 -1.15 - 1.00 5 + -17 -11 3 11.42 - 5.98 4 + -17 -10 -5 3.39 - 1.03 11 + -17 -10 -4 3.38 - 1.01 54 + -17 -10 -3 11.95 - 8.12 70 + -17 -10 -2 5.07 - 2.94 85 + -17 -10 -1 7.14 - 4.32 100 + -17 -10 0 2.87 - 0.96 57 + -17 -10 1 8.99 - 6.60 22 + -17 -10 2 -0.64 - 1.99 3 + -17 -10 3 -4.42 - 3.12 2 + -17 -10 4 2.32 - 1.64 2 + -17 -9 -8 11.31 - 3.11 3 + -17 -9 -7 4.01 - 2.84 2 + -17 -9 -6 4.22 - 1.76 7 + -17 -9 -5 3.59 - 2.00 29 + -17 -9 -4 65.27 - 45.70 91 + -17 -9 -3 2.97 - 0.64 116 + -17 -9 -2 11.01 - 3.53 126 + -17 -9 -1 15.86 - 7.64 96 + -17 -9 0 3.26 - 0.92 89 + -17 -9 1 2.87 - 1.06 45 + -17 -9 2 46.35 - 30.36 11 + -17 -9 3 1.78 - 2.46 4 + -17 -9 4 -0.50 - 5.55 3 + -17 -8 -7 -5.34 - 3.34 3 + -17 -8 -6 1.45 - 2.24 8 + -17 -8 -5 3.27 - 2.19 35 + -17 -8 -4 3.22 - 0.84 97 + -17 -8 -3 3.03 - 0.91 116 + -17 -8 -2 3.71 - 0.65 169 + -17 -8 -1 1.36 - 0.63 121 + -17 -8 0 4.29 - 1.02 94 + -17 -8 1 2.52 - 1.31 47 + -17 -8 2 3.62 - 2.17 16 + -17 -8 3 6.34 - 1.98 4 + -17 -8 4 10.37 - 4.83 2 + -17 -7 -8 5.95 - 2.57 3 + -17 -7 -7 -2.56 - 3.48 4 + -17 -7 -6 -2.75 - 3.09 4 + -17 -7 -5 1.68 - 1.56 29 + -17 -7 -4 1.58 - 0.74 78 + -17 -7 -3 26.87 - 22.09 109 + -17 -7 -2 3.47 - 1.26 120 + -17 -7 -1 2.64 - 1.18 129 + -17 -7 0 3.20 - 1.10 82 + -17 -7 1 1.73 - 1.65 44 + -17 -7 2 10.91 - 4.49 5 + -17 -7 4 7.92 - 3.97 5 + -17 -6 -7 4.01 - 2.84 2 + -17 -6 -6 2.07 - 1.46 2 + -17 -6 -5 2.82 - 1.53 14 + -17 -6 -4 1.80 - 1.11 47 + -17 -6 -3 2.76 - 0.73 72 + -17 -6 -2 3.25 - 1.06 81 + -17 -6 -1 2.38 - 0.82 87 + -17 -6 0 4.43 - 1.70 43 + -17 -6 1 3.27 - 2.38 17 + -17 -6 2 2.51 - 1.11 5 + -17 -5 -6 4.16 - 5.19 3 + -17 -5 -5 5.02 - 3.43 8 + -17 -5 -4 0.06 - 1.50 14 + -17 -5 -3 7.41 - 3.97 26 + -17 -5 -2 3.51 - 1.42 24 + -17 -5 -1 0.01 - 0.10 31 + -17 -5 0 1.76 - 2.20 11 + -17 -5 1 7.60 - 2.58 10 + -17 -5 3 24.57 - 5.19 2 + -17 -4 -7 5.41 - 8.01 2 + -17 -4 -6 -0.72 - 1.87 2 + -17 -4 -5 6.60 - 5.45 4 + -17 -4 -4 5.73 - 1.62 6 + -17 -4 -3 -2.27 - 2.86 8 + -17 -4 -2 7.44 - 2.63 6 + -17 -4 -1 2.94 - 1.48 4 + -17 -4 0 5.18 - 3.50 5 + -17 -3 -6 6.03 - 4.26 2 + -17 -3 -4 1.18 - 2.18 3 + -17 -3 -3 3.83 - 5.60 4 + -17 -3 -2 0.00 - 0.00 2 + -17 -3 0 8.74 - 5.52 3 + -17 -3 2 -0.49 - 0.35 2 + -16 -16 -5 8.72 - 1.57 2 + -16 -16 -2 -2.65 - 2.07 2 + -16 -16 -1 -2.26 - 1.60 2 + -16 -16 0 6.78 - 4.80 2 + -16 -15 -9 -9.21 - 6.76 2 + -16 -15 -5 -3.51 - 8.87 4 + -16 -15 -4 4.64 - 5.33 6 + -16 -15 -3 7.51 - 3.47 28 + -16 -15 -2 6.92 - 3.93 44 + -16 -15 -1 5.80 - 1.31 31 + -16 -15 0 2.18 - 1.42 19 + -16 -15 1 2.01 - 2.80 10 + -16 -15 2 4.75 - 2.25 5 + -16 -14 -8 1.00 - 0.82 3 + -16 -14 -7 14.25 - 6.60 6 + -16 -14 -6 6.83 - 3.04 15 + -16 -14 -5 7.10 - 2.17 87 + -16 -14 -4 11.09 - 5.20 152 + -16 -14 -3 3.05 - 0.67 217 + -16 -14 -2 10.99 - 4.49 188 + -16 -14 -1 8.04 - 2.77 217 + -16 -14 0 2.67 - 0.77 177 + -16 -14 1 14.56 - 5.40 133 + -16 -14 2 6.75 - 2.62 55 + -16 -14 3 8.33 - 1.74 10 + -16 -14 4 2.73 - 3.25 5 + -16 -14 5 17.54 - 9.05 2 + -16 -13 -8 7.06 - 1.03 3 + -16 -13 -7 1.73 - 1.17 48 + -16 -13 -6 3.20 - 0.65 179 + -16 -13 -5 3.39 - 0.69 224 + -16 -13 -4 4.14 - 0.58 294 + -16 -13 -3 6.25 - 0.66 314 + -16 -13 -2 5.50 - 0.67 381 + -16 -13 -1 9.31 - 1.84 317 + -16 -13 0 8.93 - 1.51 330 + -16 -13 1 6.13 - 1.15 286 + -16 -13 2 3.34 - 0.94 213 + -16 -13 3 4.54 - 2.22 106 + -16 -13 4 4.67 - 3.26 15 + -16 -13 5 7.51 - 5.24 3 + -16 -13 6 -10.29 - 13.53 2 + -16 -12 -9 12.74 - 3.23 6 + -16 -12 -8 1.95 - 1.31 35 + -16 -12 -7 7.39 - 9.34 184 + -16 -12 -6 5.33 - 1.05 311 + -16 -12 -5 6.71 - 0.86 357 + -16 -12 -4 6.96 - 0.96 412 + -16 -12 -3 38.54 - 12.11 437 + -16 -12 -2 7.56 - 2.08 479 + -16 -12 -1 56.57 - 11.37 441 + -16 -12 0 44.62 - 12.15 437 + -16 -12 1 16.11 - 2.83 372 + -16 -12 2 18.48 - 4.38 300 + -16 -12 3 4.40 - 0.84 253 + -16 -12 4 2.53 - 0.59 124 + -16 -12 5 13.57 - 6.30 16 + -16 -11 -9 5.50 - 1.64 18 + -16 -11 -8 4.07 - 1.66 120 + -16 -11 -7 3.77 - 0.66 285 + -16 -11 -6 13.59 - 3.34 389 + -16 -11 -5 4.61 - 0.64 456 + -16 -11 -4 36.53 - 10.20 539 + -16 -11 -3 6.01 - 0.69 523 + -16 -11 -2 11.11 - 1.90 548 + -16 -11 -1 5.97 - 0.61 557 + -16 -11 0 5.78 - 0.57 522 + -16 -11 1 6.14 - 0.74 441 + -16 -11 2 5.76 - 0.65 389 + -16 -11 3 5.17 - 0.67 333 + -16 -11 4 3.51 - 0.78 240 + -16 -11 5 5.14 - 1.45 88 + -16 -11 6 1.44 - 2.59 7 + -16 -11 7 -0.60 - 3.65 2 + -16 -10 -9 3.60 - 0.94 76 + -16 -10 -8 25.20 - 8.83 221 + -16 -10 -7 8.22 - 1.02 338 + -16 -10 -6 39.25 - 8.80 405 + -16 -10 -5 41.45 - 9.89 497 + -16 -10 -4 6.85 - 0.86 508 + -16 -10 -3 22.98 - 4.54 515 + -16 -10 -2 9.72 - 1.42 537 + -16 -10 -1 5.35 - 0.61 538 + -16 -10 0 8.75 - 1.09 509 + -16 -10 1 23.42 - 7.74 481 + -16 -10 2 9.60 - 1.87 423 + -16 -10 3 24.06 - 7.44 347 + -16 -10 4 21.13 - 6.31 285 + -16 -10 5 3.26 - 0.75 156 + -16 -10 6 16.10 - 12.97 28 + -16 -10 7 0.34 - 1.25 2 + -16 -9 -11 0.47 - 4.35 3 + -16 -9 -10 1.08 - 3.12 8 + -16 -9 -9 2.91 - 0.77 128 + -16 -9 -8 10.94 - 2.83 285 + -16 -9 -7 4.03 - 0.70 385 + -16 -9 -6 5.19 - 0.71 441 + -16 -9 -5 5.01 - 0.58 576 + -16 -9 -4 4.84 - 0.62 573 + -16 -9 -3 7.80 - 1.04 559 + -16 -9 -2 9.45 - 1.31 562 + -16 -9 -1 6.19 - 0.61 583 + -16 -9 0 13.15 - 2.42 571 + -16 -9 1 6.26 - 1.00 461 + -16 -9 2 8.60 - 0.94 486 + -16 -9 3 6.23 - 0.63 406 + -16 -9 4 8.00 - 1.32 331 + -16 -9 5 0.63 - 0.29 178 + -16 -9 6 0.96 - 1.48 34 + -16 -9 7 8.13 - 3.46 3 + -16 -9 8 -3.36 - 5.17 3 + -16 -8 -10 31.94 - 30.09 7 + -16 -8 -9 2.66 - 0.57 153 + -16 -8 -8 27.27 - 6.53 307 + -16 -8 -7 8.63 - 1.04 413 + -16 -8 -6 7.92 - 0.66 465 + -16 -8 -5 5.71 - 0.65 522 + -16 -8 -4 17.11 - 2.83 562 + -16 -8 -3 26.30 - 3.87 567 + -16 -8 -2 0.00 - 0.00 588 + -16 -8 -1 42.50 - 9.09 622 + -16 -8 0 5.19 - 0.62 538 + -16 -8 1 39.87 - 7.70 471 + -16 -8 2 10.58 - 1.48 492 + -16 -8 3 5.38 - 0.64 437 + -16 -8 4 9.85 - 1.91 370 + -16 -8 5 8.17 - 2.95 233 + -16 -8 6 2.14 - 1.15 47 + -16 -8 7 1.06 - 0.75 2 + -16 -7 -10 2.94 - 2.70 12 + -16 -7 -9 0.00 - 0.00 158 + -16 -7 -8 3.62 - 0.66 269 + -16 -7 -7 18.20 - 2.64 396 + -16 -7 -6 7.34 - 1.37 491 + -16 -7 -5 6.21 - 1.32 554 + -16 -7 -4 0.00 - 0.00 616 + -16 -7 -3 6.80 - 1.06 572 + -16 -7 -2 8.83 - 1.58 532 + -16 -7 -1 4.04 - 0.62 569 + -16 -7 0 5.49 - 0.64 543 + -16 -7 1 6.46 - 0.69 514 + -16 -7 2 10.32 - 1.44 480 + -16 -7 3 0.00 - 0.00 418 + -16 -7 4 5.92 - 0.78 312 + -16 -7 5 3.34 - 0.69 225 + -16 -7 6 3.21 - 1.41 59 + -16 -7 8 -16.54 - 8.13 2 + -16 -6 -11 0.27 - 0.98 2 + -16 -6 -10 6.01 - 2.18 16 + -16 -6 -9 3.48 - 1.03 112 + -16 -6 -8 10.82 - 2.71 278 + -16 -6 -7 7.99 - 1.08 384 + -16 -6 -6 59.06 - 15.37 461 + -16 -6 -5 5.08 - 0.69 492 + -16 -6 -4 19.00 - 3.14 551 + -16 -6 -3 16.88 - 2.46 596 + -16 -6 -2 5.63 - 0.78 516 + -16 -6 -1 9.52 - 1.57 572 + -16 -6 0 12.81 - 2.26 537 + -16 -6 1 11.10 - 1.51 499 + -16 -6 2 8.70 - 1.24 474 + -16 -6 3 61.06 - 15.61 407 + -16 -6 4 4.70 - 0.64 318 + -16 -6 5 8.29 - 2.50 172 + -16 -6 6 6.60 - 2.01 29 + -16 -6 7 3.90 - 4.09 6 + -16 -5 -10 0.93 - 7.95 3 + -16 -5 -9 5.55 - 5.22 68 + -16 -5 -8 2.82 - 0.60 241 + -16 -5 -7 11.38 - 2.05 358 + -16 -5 -6 4.89 - 0.63 421 + -16 -5 -5 5.26 - 0.63 518 + -16 -5 -4 6.44 - 0.80 503 + -16 -5 -3 1.33 - 0.43 510 + -16 -5 -2 14.38 - 2.52 514 + -16 -5 -1 8.94 - 1.24 547 + -16 -5 0 3.82 - 0.90 528 + -16 -5 1 2.14 - 0.40 482 + -16 -5 2 11.32 - 1.97 498 + -16 -5 3 6.51 - 0.70 358 + -16 -5 4 4.69 - 0.59 264 + -16 -5 5 2.13 - 0.73 105 + -16 -5 6 -3.50 - 2.28 7 + -16 -4 -9 45.76 - 38.25 20 + -16 -4 -8 3.95 - 1.21 144 + -16 -4 -7 4.91 - 0.65 260 + -16 -4 -6 7.15 - 0.71 367 + -16 -4 -5 6.18 - 0.66 443 + -16 -4 -4 17.86 - 3.43 478 + -16 -4 -3 6.16 - 0.58 538 + -16 -4 -2 26.22 - 10.09 534 + -16 -4 -1 2.09 - 0.68 534 + -16 -4 0 25.32 - 4.02 548 + -16 -4 1 15.96 - 3.13 445 + -16 -4 2 5.16 - 0.67 378 + -16 -4 3 3.06 - 0.68 285 + -16 -4 4 6.15 - 2.18 176 + -16 -4 5 3.34 - 1.60 39 + -16 -4 6 -1.56 - 1.27 3 + -16 -3 -8 4.16 - 1.32 41 + -16 -3 -7 16.93 - 8.25 150 + -16 -3 -6 3.21 - 0.56 317 + -16 -3 -5 22.74 - 4.27 335 + -16 -3 -4 5.47 - 0.68 353 + -16 -3 -3 18.02 - 3.50 426 + -16 -3 -2 5.97 - 0.63 395 + -16 -3 -1 0.00 - 0.00 427 + -16 -3 0 5.13 - 0.73 374 + -16 -3 1 5.38 - 0.65 351 + -16 -3 2 6.03 - 1.27 323 + -16 -3 3 2.94 - 0.60 187 + -16 -3 4 2.63 - 0.83 77 + -16 -3 5 -1.41 - 4.83 7 + -16 -3 6 14.18 - 6.81 3 + -16 -2 -9 0.00 - 0.00 2 + -16 -2 -8 -0.89 - 2.53 6 + -16 -2 -7 4.08 - 2.44 47 + -16 -2 -6 4.18 - 1.48 145 + -16 -2 -5 3.56 - 0.66 230 + -16 -2 -4 28.45 - 9.38 307 + -16 -2 -3 5.16 - 0.67 328 + -16 -2 -2 8.23 - 1.19 345 + -16 -2 -1 8.48 - 1.26 294 + -16 -2 0 7.08 - 1.51 323 + -16 -2 1 3.25 - 0.66 217 + -16 -2 2 12.47 - 4.13 146 + -16 -2 3 14.24 - 8.22 41 + -16 -2 4 3.55 - 2.07 4 + -16 -2 5 9.74 - 3.30 2 + -16 -1 -7 2.71 - 3.22 3 + -16 -1 -6 4.16 - 2.18 19 + -16 -1 -5 3.23 - 1.32 44 + -16 -1 -4 3.21 - 0.62 123 + -16 -1 -3 4.16 - 1.62 160 + -16 -1 -2 2.81 - 0.73 148 + -16 -1 -1 3.14 - 0.87 129 + -16 -1 0 2.50 - 0.78 120 + -16 -1 1 5.17 - 0.98 64 + -16 -1 2 -0.86 - 1.74 22 + -16 -1 3 3.71 - 3.55 9 + -16 -1 4 2.60 - 3.20 3 + -16 0 -6 8.99 - 4.19 5 + -16 0 -5 -2.76 - 2.32 4 + -16 0 -4 0.53 - 4.28 4 + -16 0 -3 7.28 - 3.16 15 + -16 0 -2 10.39 - 1.80 21 + -16 0 -1 2.02 - 2.98 13 + -16 0 0 1.26 - 1.94 3 + -16 0 1 -0.54 - 0.77 2 + -16 0 2 0.38 - 4.31 3 + -16 1 -4 -8.38 - 13.42 2 + -16 1 0 22.33 - 0.93 2 + -15 -18 1 2.78 - 1.97 2 + -15 -17 -7 0.33 - 0.23 2 + -15 -17 -6 8.14 - 4.98 2 + -15 -17 -5 8.18 - 3.55 3 + -15 -17 -4 1.51 - 3.26 8 + -15 -17 -3 20.87 - 9.28 13 + -15 -17 -2 2.46 - 1.68 26 + -15 -17 -1 6.44 - 1.67 27 + -15 -17 0 0.81 - 1.62 19 + -15 -17 1 5.85 - 3.92 5 + -15 -17 2 11.16 - 1.57 2 + -15 -17 3 5.06 - 6.59 2 + -15 -17 4 3.35 - 14.34 2 + -15 -16 -8 -1.43 - 2.96 3 + -15 -16 -7 11.50 - 3.94 5 + -15 -16 -6 10.24 - 5.08 44 + -15 -16 -5 9.22 - 3.06 115 + -15 -16 -4 10.51 - 2.75 199 + -15 -16 -3 3.04 - 0.58 248 + -15 -16 -2 2.63 - 0.66 262 + -15 -16 -1 3.92 - 0.62 267 + -15 -16 0 6.70 - 1.44 246 + -15 -16 1 0.68 - 0.38 199 + -15 -16 2 3.49 - 1.11 117 + -15 -16 3 4.59 - 1.15 35 + -15 -16 4 -1.93 - 2.62 5 + -15 -15 -9 0.00 - 0.00 2 + -15 -15 -8 4.58 - 1.98 20 + -15 -15 -7 4.13 - 0.81 127 + -15 -15 -6 25.74 - 10.01 252 + -15 -15 -5 6.40 - 0.66 312 + -15 -15 -4 16.79 - 4.12 414 + -15 -15 -3 11.57 - 2.34 438 + -15 -15 -2 5.93 - 0.69 472 + -15 -15 -1 5.45 - 0.62 436 + -15 -15 0 18.89 - 3.12 405 + -15 -15 1 5.61 - 1.67 360 + -15 -15 2 9.06 - 2.28 362 + -15 -15 3 27.05 - 5.93 272 + -15 -15 4 7.08 - 2.34 123 + -15 -15 5 5.78 - 1.58 31 + -15 -15 6 10.83 - 7.82 3 + -15 -14 -10 10.42 - 8.09 3 + -15 -14 -9 12.70 - 6.29 49 + -15 -14 -8 5.60 - 1.55 231 +End of reflections diff --git a/tests/unit/io/test_crystfel_hkl.py b/tests/unit/io/test_crystfel_hkl.py new file mode 100644 index 00000000..5abc3b01 --- /dev/null +++ b/tests/unit/io/test_crystfel_hkl.py @@ -0,0 +1,133 @@ +"""Reading CrystFEL ``partialator`` reflection lists. + +The format that merged serial data actually arrives in. Two properties matter beyond +"it parses": + +* **negative intensities survive.** A merged weak reflection legitimately comes out below + zero, and that is information -- dropping or clamping it biases the mean upward exactly + where the noise dominates. +* **cell and space group come from the caller**, because the format carries neither. + +The fixtures are 400-reflection excerpts of a real ``partialator`` custom-split pair, so the +two halves cover overlapping-but-different reflections measured independently -- the property +that makes such a pair usable as a null, where any difference between them is noise plus +systematics with no real signal in it. +""" + +import pytest +import torch + +# Cell of the small-molecule dataset these excerpts come from. The format does not carry +# it, so the caller must supply it; a wrong cell would silently give wrong d-spacings. +CELL = [14.97, 18.85, 18.89, 89.4, 84.9, 67.8] +SPACEGROUP = "P 1" + + +@pytest.fixture(scope="module") +def hkl_dir(test_files_dir): + d = test_files_dir / "hkl" + if not d.is_dir(): + pytest.skip("CrystFEL hkl fixtures not present") + return d + + +def _load(path): + from torchref import ReflectionData + + return ReflectionData(device="cpu", verbose=0).load_crystfel_hkl( + str(path), cell=CELL, spacegroup=SPACEGROUP + ) + + +@pytest.mark.unit +class TestReaderBasics: + def test_reads_intensities_and_sigmas(self, hkl_dir): + data = _load(hkl_dir / "dark_half1.hkl") + assert data.I is not None and data.I_sigma is not None + assert len(data.I) == len(data.hkl) + assert torch.isfinite(data.I_sigma).all() + # Real partialator output does contain sigma(I) == 0 -- which is why an + # intensity likelihood has to floor it rather than trust it. + assert (data.I_sigma >= 0).all() + + def test_amplitudes_are_derived_by_french_wilson(self, hkl_dir): + """The format is intensity-native, so F comes from the same path an MTZ with + I/SIGI takes.""" + data = _load(hkl_dir / "dark_half1.hkl") + assert data.F is not None + assert (data.F >= 0).all() + assert data._FrenchWilson is not None + + def test_cell_and_spacegroup_come_from_the_caller(self, hkl_dir): + data = _load(hkl_dir / "dark_half1.hkl") + assert torch.allclose( + data.cell.data[:3].to(torch.float64), + torch.tensor(CELL[:3], dtype=torch.float64), + atol=1e-3, + ) + assert data.spacegroup is not None + + def test_the_trailing_marker_is_not_read_as_a_reflection(self, hkl_dir): + """``partialator`` ends the list with 'End of reflections'.""" + raw = (hkl_dir / "dark_half1.hkl").read_text().splitlines() + assert raw[-1].startswith("End of reflections") + data = _load(hkl_dir / "dark_half1.hkl") + # 3 header lines + N reflections + 1 trailer + assert len(data.hkl) <= len(raw) - 4 + + +@pytest.mark.unit +class TestNegativeIntensitiesSurvive: + def test_the_fixture_contains_negatives(self, hkl_dir): + """Precondition, asserted: without negatives the next test proves nothing.""" + n = 0 + for line in (hkl_dir / "dark_half1.hkl").read_text().splitlines()[3:]: + parts = line.split() + if len(parts) >= 4: + try: + n += float(parts[3]) < 0 + except ValueError: + pass + assert n > 5, f"only {n} negative intensities in the fixture" + + def test_negatives_reach_the_dataset(self, hkl_dir): + data = _load(hkl_dir / "dark_half1.hkl") + assert bool((data.I < 0).any()), ( + "negative intensities were dropped or clamped by the reader" + ) + + +@pytest.mark.unit +class TestSplitHalves: + def test_the_two_halves_are_independent_measurements(self, hkl_dir): + """Different reflection sets and different values -- which is what makes them + usable as a null: any difference between them is noise plus systematics, with no + real signal in it. + + (The excerpts are equal-length slices of the full halves, which are 33613 and + 33523 reflections; the sets still differ, which is the property that matters.) + """ + a = _load(hkl_dir / "dark_half1.hkl") + b = _load(hkl_dir / "dark_half2.hkl") + + set_a = {tuple(row) for row in a.hkl.tolist()} + set_b = {tuple(row) for row in b.hkl.tolist()} + assert set_a != set_b, "the halves cover identical reflections" + assert set_a & set_b, "the halves share no reflections at all" + + def test_the_halves_align_into_a_collection(self, hkl_dir): + """``add_dataset`` must reconcile two different reflection lists onto one grid.""" + from torchref.io.datasets.collection import DatasetCollection + + a = _load(hkl_dir / "dark_half1.hkl") + b = _load(hkl_dir / "dark_half2.hkl") + + dc = DatasetCollection(verbose=0, device="cpu") + dc.add_dataset("half1", a, set_as_reference=True) + dc.add_dataset("half2", b) + + assert dc.n_datasets == 2 + assert len(a.hkl) == len(b.hkl) == len(dc.hkl) + # Both carry intensities, so an intensity target can run on the pair. + assert dc["half1"].I is not None and dc["half2"].I is not None + assert dc.stack_I_obs().shape == (2, len(dc.hkl)) diff --git a/tests/unit/io/test_fcalc_add_noise.py b/tests/unit/io/test_fcalc_add_noise.py new file mode 100644 index 00000000..70658e57 --- /dev/null +++ b/tests/unit/io/test_fcalc_add_noise.py @@ -0,0 +1,220 @@ +"""Simulated intensities must keep their negatives. + +``add_noise`` draws two independent noisy half-datasets and returns their mean. The +amplitude it derives has to be clamped -- an amplitude cannot be negative -- but the +*intensity* must not be, and both the intensity and its sigma have to survive on the +returned dataset. + +Why this is not a detail: clamping ``I_mean`` at zero puts a **positive bias** on exactly +the weak reflections where the noise dominates. That bias is smooth, positive, and largest +where the signal is weakest -- the same signature as a genuine positive perturbation of the +merged intensity. Any study of an effect at the 1e-3 level built on clamped simulated data +would be measuring its own generator. +""" + +import pytest +import torch + + +@pytest.fixture +def fcalc_scene(): + """A small P1 scene with a deliberately wide dynamic range. + + The weak tail is the point: with strong reflections only, noise never pushes an + intensity negative and nothing below can distinguish clamped from unclamped. + """ + from torchref.io.datasets import FcalcDataset + + dataset = FcalcDataset.from_cell_and_resolution( + cell=[30.0, 32.0, 34.0, 90.0, 90.0, 90.0], + spacegroup="P 1", + d_min=3.0, + device=torch.device("cpu"), + ) + n = len(dataset.hkl) + gen = torch.Generator().manual_seed(17) + # Amplitudes spanning three orders of magnitude, so I spans six. + amp = 10.0 ** (torch.rand(n, generator=gen) * 3.0 - 1.0) + phase = torch.rand(n, generator=gen) * 6.283 + dataset.set_fcalc((amp * torch.exp(1j * phase)).to(torch.complex64)) + return dataset + + +@pytest.mark.unit +class TestNegativesSurvive: + def test_some_intensities_come_out_negative(self, fcalc_scene): + noisy = fcalc_scene.add_noise(sigma_mul=0.5, seed=3, verbose=False) + assert noisy.I is not None, "add_noise did not retain intensities" + assert bool((noisy.I < 0).any()), ( + "no negative intensities at 50% multiplicative noise -- either the scene has " + "no weak reflections or the intensities were clamped" + ) + + def test_the_amplitude_is_clamped_but_the_intensity_is_not(self, fcalc_scene): + """The asymmetry is deliberate, so it is asserted rather than assumed.""" + noisy = fcalc_scene.add_noise(sigma_mul=0.5, seed=3, verbose=False) + assert bool((noisy.fcalc_amp >= 0).all()) + negative = noisy.I < 0 + assert bool(negative.any()) + # Where the intensity is negative the amplitude is floored at zero, so the two + # cannot agree -- which is exactly the information a clamp would have destroyed. + assert torch.allclose( + noisy.fcalc_amp[negative], torch.zeros(int(negative.sum())) + ) + + @staticmethod + def _weak(noisy, truth): + """The subset a clamp at zero can touch: reflections within 2 sigma of zero. + + Measuring over the whole list instead would drown the effect -- the strong + reflections contribute nothing to the bias but dominate its standard error, so + the very reflections the clamp distorts are the ones averaged away. + + Note this needs a noise model whose sigma does **not** scale with the intensity. + Under purely multiplicative noise ``sigma = f * I``, so no reflection is ever weak + relative to its own sigma and this subset is empty; the Poisson-like ``sigma_lin`` + term below gives ``sigma ~ sqrt(I)`` and therefore a genuine weak tail. + """ + return truth < 2.0 * noisy.I_sigma + + def test_the_intensity_is_unbiased_on_the_weak_reflections(self, fcalc_scene): + """The property the clamp breaks, as a bound on the mean of the weak tail. + + The tolerance is the standard error over that subset, not a percentage. + """ + noisy = fcalc_scene.add_noise( + sigma_lin=200.0, sigma_mul=0.0, seed=5, verbose=False + ) + truth = fcalc_scene.fcalc_amp**2 + weak = self._weak(noisy, truth) + assert int(weak.sum()) > 50, "too few weak reflections to say anything" + + residual = (noisy.I - truth)[weak] + sem = float(noisy.I_sigma[weak].pow(2).sum().sqrt() / int(weak.sum())) + bias = float(residual.mean()) + assert abs(bias) < 4.0 * sem, ( + f"weak-reflection intensity bias {bias:.4g} exceeds 4 sigma " + f"({4 * sem:.4g}); the intensities are being clamped or otherwise skewed" + ) + + def test_clamping_would_be_detectable_on_this_scene(self, fcalc_scene): + """Anti-vacuity: quantify what the defect would have looked like here. + + Without this, the test above could pass simply because the scene has no + reflections weak enough for a clamp to reach. + + Measured on this scene, the clamp bias runs only about one to two times the + statistical error on the same mean, and needs a high noise level to stand clear + of it at all. That is not a reason to tolerate it: the bias is **systematic**, so + it repeats identically across datasets and survives averaging, while the error it + is being compared against shrinks as 1/sqrt(N). It is the accumulation, not the + size on any one dataset, that would corrupt a calibration curve. + """ + noisy = fcalc_scene.add_noise( + sigma_lin=200.0, sigma_mul=0.0, seed=5, verbose=False + ) + truth = fcalc_scene.fcalc_amp**2 + weak = self._weak(noisy, truth) + assert int(weak.sum()) > 50 + + honest = float((noisy.I - truth)[weak].mean()) + clamped = float((noisy.I.clamp(min=0.0) - truth)[weak].mean()) + sem = float(noisy.I_sigma[weak].pow(2).sum().sqrt() / int(weak.sum())) + + assert clamped > honest, "clamping did not raise the mean on this scene" + assert clamped > 4.0 * sem, ( + f"the clamped bias ({clamped:.4g}) would fall inside the noise " + f"({4 * sem:.4g}) here, so this scene cannot demonstrate the defect" + ) + + +@pytest.mark.unit +class TestSigmaAndHalves: + def test_sigma_of_the_mean_is_the_single_draw_sigma_over_root_two(self, fcalc_scene): + a = fcalc_scene.add_noise(sigma_mul=0.2, seed=11, verbose=False) + b = fcalc_scene.add_noise(sigma_mul=0.2, seed=12, verbose=False) + # Same model, same noise scale: the reported sigma is a property of the model, + # not of the draw, so it must be identical across seeds. + assert torch.allclose(a.I_sigma, b.I_sigma) + + expected = torch.sqrt( + torch.tensor(0.2) ** 2 * fcalc_scene.fcalc_amp**4 + ) / (2.0**0.5) + assert torch.allclose(a.I_sigma, expected, rtol=1e-5) + + def test_amplitude_sigma_uses_the_true_amplitude(self, fcalc_scene): + """Propagating against the noisy amplitude would warp sigma per draw and break + inverse-variance weighting downstream.""" + a = fcalc_scene.add_noise(sigma_mul=0.2, seed=11, verbose=False) + b = fcalc_scene.add_noise(sigma_mul=0.2, seed=99, verbose=False) + assert torch.allclose(a.fobs_sigma, b.fobs_sigma) + + def test_different_seeds_give_different_draws(self, fcalc_scene): + a = fcalc_scene.add_noise(sigma_mul=0.2, seed=1, verbose=False) + b = fcalc_scene.add_noise(sigma_mul=0.2, seed=2, verbose=False) + assert not torch.allclose(a.I, b.I) + + def test_the_same_seed_reproduces(self, fcalc_scene): + a = fcalc_scene.add_noise(sigma_mul=0.2, seed=7, verbose=False) + b = fcalc_scene.add_noise(sigma_mul=0.2, seed=7, verbose=False) + assert torch.equal(a.I, b.I) + + def test_phases_are_untouched(self, fcalc_scene): + """Only the amplitude is perturbed; the phase is the model's.""" + noisy = fcalc_scene.add_noise(sigma_mul=0.2, seed=4, verbose=False) + strong = fcalc_scene.fcalc_amp > 1.0 + assert torch.allclose( + noisy.fcalc_phase[strong], fcalc_scene.fcalc_phase[strong], atol=1e-5 + ) + + def test_the_source_dataset_is_not_modified(self, fcalc_scene): + before = fcalc_scene.fcalc_amp.clone() + fcalc_scene.add_noise(sigma_mul=0.3, seed=8, verbose=False) + assert torch.equal(fcalc_scene.fcalc_amp, before) + assert fcalc_scene.I is None + + +@pytest.mark.unit +class TestReferenceDriven: + def test_sigmas_are_grafted_from_the_reference(self, fcalc_scene, mtz_dir): + """The path that matters for real data: per-reflection sigmas from a measured + dataset rather than a parametric model.""" + from torchref import ReflectionData + from torchref.io.datasets import FcalcDataset + + mtz = mtz_dir / "1DAW.mtz" + if not mtz.exists(): + pytest.skip("1DAW fixture not present") + ref = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + if ref.I is None: + pytest.skip("1DAW loaded without intensities") + + # Build on the reference's own HKL list, which is what makes grafting 1:1. + scene = FcalcDataset( + hkl=ref.hkl.clone(), cell=fcalc_scene.cell, + spacegroup=fcalc_scene.spacegroup, device=torch.device("cpu"), + ) + gen = torch.Generator().manual_seed(2) + amp = torch.rand(len(ref.hkl), generator=gen) * 100.0 + scene.set_fcalc((amp + 0j).to(torch.complex64)) + + noisy = scene.add_noise(reference=ref, seed=1, verbose=False) + assert torch.allclose(noisy.I_sigma, ref.I_sigma / (2.0**0.5)) + + def test_a_mismatched_reference_is_rejected(self, fcalc_scene, mtz_dir): + from torchref import ReflectionData + + mtz = mtz_dir / "1DAW.mtz" + if not mtz.exists(): + pytest.skip("1DAW fixture not present") + ref = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + + with pytest.raises(ValueError, match="does not match"): + fcalc_scene.add_noise(reference=ref, verbose=False) + + def test_a_reference_without_sigmas_is_rejected(self, fcalc_scene): + from types import SimpleNamespace + + bad = SimpleNamespace(I_sigma=None, hkl=fcalc_scene.hkl) + with pytest.raises(ValueError, match="I_sigma is None"): + fcalc_scene.add_noise(reference=bad, verbose=False) diff --git a/torchref/cli/simulate_noisy_data.py b/torchref/cli/simulate_noisy_data.py new file mode 100644 index 00000000..36711fdc --- /dev/null +++ b/torchref/cli/simulate_noisy_data.py @@ -0,0 +1,282 @@ +#!/usr/bin/env python3 +"""Simulate a noisy reflection MTZ from a structure file. + +Two modes are supported, chosen by whether ``--reference-hkl`` is passed: + +**Reference-driven (preferred)** + The model is scaled to the reference dataset with torchref's + ``Scaler`` so that Fcalc and reference intensities share an absolute + scale. The simulation then inherits the reference's HKL list, and + noise for each reflection uses the reference's reported ``sigma(I)`` + directly — no parametric model, no fitting. Use this when you have a + CrystFEL ``.hkl`` (or equivalent) from a real experiment. + +**Parametric fallback** + When ``--reference-hkl`` is omitted, intensity sigma is built from a + three-term variance-additive model + ``sigma_I^2 = sigma_lin^2 * I + sigma_mul^2 * I^2 + sigma_abs_I^2``. + +In both modes, two independent noise realizations are drawn per +reflection; R-split and Pearson CC between the halves are printed, and +the mean is written as F-obs/SIGF-obs (default) or I-obs/SIGI-obs. + +Usage +----- +:: + + # Reference-driven + torchref.simulate-noisy-data input.pdb out.mtz \ + --reference-hkl td1.hkl --d-min 2.0 + + # Parametric (no reference) + torchref.simulate-noisy-data input.pdb out.mtz \ + --sigma-lin 5 --sigma-mul 0.05 --sigma-abs 0.33 + + # Intensity output + torchref.simulate-noisy-data input.pdb out.mtz \ + --reference-hkl td1.hkl --output-type intensities +""" + +import argparse +import sys +from pathlib import Path + +import pandas as pd +import torch + +from torchref.cli._common import ( + add_device_arg, + add_verbose_arg, + configure_unbuffered_output, + load_model, + parse_device_str, +) +from torchref.io import mtz +from torchref.io.datasets import FcalcDataset + + +def _build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser( + prog="torchref.simulate-noisy-data", + description=( + "Compute Fcalc from a structure, add Gaussian intensity noise " + "(optionally grafted from a CrystFEL reference hkl), and write " + "an MTZ as F-obs/SIGF-obs (default) or I-obs/SIGI-obs." + ), + ) + parser.add_argument("input", help="Input structure file (.pdb / .cif / .mmcif)") + parser.add_argument("output", help="Output MTZ file") + parser.add_argument( + "--d-min", + type=float, + default=2.0, + help="High-resolution limit in Angstroms (default: 2.0)", + ) + parser.add_argument( + "--d-max", + type=float, + default=None, + help="Low-resolution limit in Angstroms (default: no cutoff)", + ) + parser.add_argument( + "--output-type", + choices=("amplitudes", "intensities"), + default="amplitudes", + help="Write F-obs/SIGF-obs (amplitudes, default) or I-obs/SIGI-obs", + ) + parser.add_argument( + "--reference-hkl", + type=str, + default=None, + help=("CrystFEL partialator .hkl file. When given, the model is " + "scaled to this reference and per-reflection sigmas are " + "grafted from it (parametric sigma-* flags are ignored)."), + ) + parser.add_argument( + "--sigma-lin", + type=float, + default=0.0, + help=("Parametric only: Poisson coefficient — sigma_I^2 += " + "sigma_lin^2 * I (default: 0.0). Ignored if --reference-hkl."), + ) + parser.add_argument( + "--sigma-mul", + type=float, + default=0.05, + help=("Parametric only: multiplicative coefficient — sigma_I^2 += " + "(sigma_mul * I)^2 (default: 0.05). Ignored if --reference-hkl."), + ) + parser.add_argument( + "--sigma-abs", + type=float, + default=0.0, + help=("Parametric only: target 1/SNR at the resolution limit from " + "the absolute-noise term (default: 0.0). Ignored if " + "--reference-hkl."), + ) + parser.add_argument( + "--seed", + type=int, + default=None, + help="Seed for the random generator (default: non-reproducible)", + ) + add_device_arg(parser) + add_verbose_arg(parser) + return parser + + +def main(): + configure_unbuffered_output() + args = _build_parser().parse_args() + + input_path = Path(args.input) + if not input_path.is_file(): + print(f"Error: input file not found: {args.input}", file=sys.stderr) + return 1 + if args.reference_hkl is not None and not Path(args.reference_hkl).is_file(): + print( + f"Error: reference hkl not found: {args.reference_hkl}", + file=sys.stderr, + ) + return 1 + if args.sigma_lin < 0 or args.sigma_mul < 0 or args.sigma_abs < 0: + print("Error: all --sigma-* flags must be non-negative", file=sys.stderr) + return 1 + + device = parse_device_str(args.device) + + if args.verbose: + print(f"Loading structure: {args.input}") + model = load_model( + args.input, max_res=args.d_min, device=device, verbose=args.verbose + ) + + if args.reference_hkl is not None: + return _run_reference_mode(args, model, device) + return _run_parametric_mode(args, model, device) + + +def _run_reference_mode(args, model, device) -> int: + """Reference-driven: scale model to CrystFEL hkl and graft sigmas.""" + from torchref import ReflectionData + from torchref.base.reciprocal import get_d_spacing + from torchref.scaling import Scaler + + if args.verbose: + print(f"Loading CrystFEL reference: {args.reference_hkl}") + ref = ReflectionData(device=str(device), verbose=args.verbose).load_crystfel_hkl( + args.reference_hkl, cell=model.cell, spacegroup=model.spacegroup, + ) + # Prune reference tensors by resolution so the simulation output only + # covers the requested range. cut_res() only masks — we want the HKL + # list itself to be filtered so Scaler, model(), and add_noise all see + # the same reflection set. + if args.d_min is not None or args.d_max is not None: + res = ref.resolution + mask = torch.ones_like(res, dtype=torch.bool) + if args.d_min is not None: + mask &= res >= args.d_min + if args.d_max is not None: + mask &= res <= args.d_max + for field in ("hkl", "I", "I_sigma", "F", "F_sigma", "resolution", "rfree_flags"): + t = getattr(ref, field, None) + if t is not None: + setattr(ref, field, t[mask]) + # Replace the masks with a single all-True flagged_initial matching + # the new tensor size. An empty TensorMasks() returns None from + # __call__(), which downstream consumers (Scaler.get_bins) don't + # tolerate. + ref.masks = type(ref.masks)(device=ref.device) + ref.masks["flagged_initial"] = torch.ones( + len(ref.hkl), dtype=torch.bool, device=ref.device + ) + if args.verbose: + print(f"Reference: {len(ref.hkl)} reflections after resolution cuts") + + if args.verbose: + print("Scaling model to reference...") + scaler = Scaler(model, ref, device=device, verbose=args.verbose) + scaler.initialize().refine_lbfgs() + + if args.verbose: + print(f"Computing scaled Fcalc on {len(ref.hkl)} reference HKLs") + with torch.no_grad(): + fcalc_scaled = scaler(model(ref.hkl)) + + sim = FcalcDataset( + hkl=ref.hkl.clone(), + resolution=get_d_spacing(ref.hkl.float(), ref.cell.data), + cell=ref.cell, + spacegroup=ref.spacegroup, + device=device, + ) + sim.set_fcalc(fcalc_scaled) + + noisy = sim.add_noise(reference=ref, seed=args.seed, verbose=bool(args.verbose)) + _write_output(args, noisy, sim_clean=sim) + return 0 + + +def _run_parametric_mode(args, model, device) -> int: + """No reference: build HKL from cell+resolution, use three-term model.""" + if args.verbose: + print( + f"Generating HKL to d_min={args.d_min} A" + + (f", d_max={args.d_max} A" if args.d_max is not None else "") + ) + dataset = FcalcDataset.from_cell_and_resolution( + cell=model.cell, + spacegroup=model.spacegroup, + d_min=args.d_min, + d_max=args.d_max, + device=device, + ) + if args.verbose: + print(f"Computing Fcalc for {len(dataset)} reflections") + hkl = dataset.hkl.to(device) + fcalc = model(hkl, recalc=True) + dataset.set_fcalc(fcalc) + + noisy = dataset.add_noise( + sigma_lin=args.sigma_lin, + sigma_mul=args.sigma_mul, + sigma_abs=args.sigma_abs, + seed=args.seed, + verbose=bool(args.verbose), + ) + _write_output(args, noisy, sim_clean=dataset) + return 0 + + +def _write_output(args, noisy: FcalcDataset, sim_clean: FcalcDataset) -> None: + """Write noisy Fcalc/Fobs to MTZ. ``sim_clean`` holds the pre-noise dataset.""" + hkl_np = noisy.hkl.cpu().numpy() + columns = { + "H": hkl_np[:, 0], + "K": hkl_np[:, 1], + "L": hkl_np[:, 2], + } + + if args.output_type == "intensities": + # The intensity add_noise actually drew, not the square of the clamped + # amplitude. Squaring back would floor every negative reflection at zero, and + # that is a positive bias concentrated exactly where the noise dominates -- the + # same signature as a real positive perturbation of the merged intensity. + columns["I-obs"] = noisy.I.detach().cpu().numpy() + columns["SIGI-obs"] = noisy.I_sigma.detach().cpu().numpy() + else: + columns["F-obs"] = noisy.fcalc_amp.detach().cpu().numpy() + columns["SIGF-obs"] = noisy.fobs_sigma.detach().cpu().numpy() + + df = pd.DataFrame(columns) + mtz.write(df, noisy.cell.data, noisy.spacegroup, args.output) + + if args.verbose: + print( + f"Wrote {args.output} " + f"({args.output_type}, n={len(noisy)})" + ) + + +if __name__ == "__main__": + sys.exit(main() or 0) diff --git a/torchref/io/datasets/fcalc_data.py b/torchref/io/datasets/fcalc_data.py index 616c310d..b0ed460e 100644 --- a/torchref/io/datasets/fcalc_data.py +++ b/torchref/io/datasets/fcalc_data.py @@ -7,7 +7,7 @@ """ from dataclasses import dataclass -from typing import Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union import pandas as pd import torch @@ -17,6 +17,9 @@ from .base import CrystalDataset +if TYPE_CHECKING: + from .reflection_data import ReflectionData + @dataclass class FcalcDataset(CrystalDataset): @@ -52,6 +55,7 @@ class FcalcDataset(CrystalDataset): fcalc: Optional[torch.Tensor] = None # Complex (N,) fcalc_amp: Optional[torch.Tensor] = None # |Fcalc| (N,) fcalc_phase: Optional[torch.Tensor] = None # Phase in radians (N,) + fobs_sigma: Optional[torch.Tensor] = None # Amp-space sigma (N,), set by add_noise @staticmethod def from_cell_and_resolution( @@ -174,6 +178,160 @@ def set_fcalc(self, fcalc: torch.Tensor) -> None: self.fcalc_amp = torch.abs(fcalc).to(device=self.device) self.fcalc_phase = torch.angle(fcalc).to(device=self.device) + def add_noise( + self, + reference: Optional["ReflectionData"] = None, + sigma_lin: float = 0.0, + sigma_mul: float = 0.05, + sigma_abs: float = 0.0, + seed: Optional[int] = None, + verbose: bool = True, + ) -> "FcalcDataset": + """ + Return a copy with Gaussian **intensity** noise, as the mean of two half-datasets. + + Two sigma sources, chosen by ``reference``: + + **Reference-driven (preferred).** Given a ``ReflectionData`` on the same HKL list, + its per-reflection ``I_sigma`` is grafted directly. Right when the Fcalc is already + on the reference's absolute scale. + + **Parametric.** Otherwise a three-term variance model:: + + sigma_I^2 = sigma_lin^2 * I + sigma_mul^2 * I^2 + (sigma_abs * _outer)^2 + + Either way two independent draws are made and averaged:: + + I_h{1,2} = I + N(0, sigma_I) + I_mean = (I_h1 + I_h2) / 2 + sigma_I_mean = sigma_I / sqrt(2) + + R-split and Pearson CC between the halves are reported, which makes the two draws + a ready-made split-half pair rather than only a noise model. + + **Negative intensities are kept.** ``I_mean`` and ``sigma_I_mean`` are stored + unclamped on the returned dataset, and only the *amplitude* is clamped, because an + amplitude cannot be negative. Clamping the intensity instead would put a positive + bias on exactly the weak reflections where noise dominates -- the same shape as a + genuine positive perturbation, and enough to swamp effects of order 1e-3. + + The amplitude sigma is propagated against the **true** amplitude, not the noisy + one, so it is not warped per draw. + + Parameters + ---------- + reference : ReflectionData, optional + Sigma donor. Requires ``torch.equal(self.hkl, reference.hkl)``. + sigma_lin, sigma_mul, sigma_abs : float, optional + Parametric coefficients, ignored when ``reference`` is given. + seed : int, optional + Seed for reproducibility. ``None`` uses the global RNG. + verbose : bool, optional + Print the R-split and CC between halves. Default True. + + Returns + ------- + FcalcDataset + New dataset carrying the noisy complex Fcalc, the unclamped ``I`` / + ``I_sigma``, and ``fobs_sigma``. + """ + if self.fcalc is None or self.fcalc_amp is None or self.fcalc_phase is None: + raise ValueError("No Fcalc values set. Call set_fcalc() first.") + + shape = self.fcalc_amp.shape + dtype = self.fcalc_amp.dtype + dev = self.device + if seed is not None: + g = torch.Generator(device=dev).manual_seed(int(seed)) + randn1 = torch.randn(shape, dtype=dtype, device=dev, generator=g) + randn2 = torch.randn(shape, dtype=dtype, device=dev, generator=g) + else: + randn1 = torch.randn(shape, dtype=dtype, device=dev) + randn2 = torch.randn(shape, dtype=dtype, device=dev) + + intensity = self.fcalc_amp**2 + + if reference is not None: + if reference.I_sigma is None: + raise ValueError( + "reference.I_sigma is None. Load the reference with intensity + " + "sigma columns (e.g. load_crystfel_hkl) before passing it here." + ) + ref_hkl = reference.hkl.to(device=self.hkl.device) + if ref_hkl.shape != self.hkl.shape or not torch.equal(ref_hkl, self.hkl): + raise ValueError( + "reference.hkl does not match self.hkl -- build the FcalcDataset " + "from the reference's HKL list to guarantee 1:1 sigma grafting." + ) + sigma_I = reference.I_sigma.to(device=dev, dtype=dtype) + if verbose: + print( + f"add_noise: grafting sigmas from reference ({len(sigma_I)} " + f"reflections, ={sigma_I.mean().item():.3g})" + ) + else: + if sigma_abs > 0: + if self.resolution is None: + raise ValueError( + "sigma_abs > 0 requires self.resolution (used to pick the " + "outer resolution shell)." + ) + n = len(self.resolution) + k = max(1, int(0.1 * n)) + outer_idx = torch.argsort(self.resolution)[:k] + sigma_abs_I = sigma_abs * intensity[outer_idx].mean() + else: + sigma_abs_I = torch.zeros((), dtype=dtype, device=dev) + + safe = intensity.clamp(min=0.0) + sigma_I = torch.sqrt( + (sigma_lin**2) * safe + + (sigma_mul**2) * safe * safe + + sigma_abs_I**2 + ) + + I_h1 = intensity + randn1 * sigma_I + I_h2 = intensity + randn2 * sigma_I + + diff_sum = (I_h1 - I_h2).abs().sum() + pair_sum = (I_h1 + I_h2).sum() + r_split = ((1.0 / (2.0**0.5)) * diff_sum / (0.5 * pair_sum)).item() + + x = I_h1 - I_h1.mean() + y = I_h2 - I_h2.mean() + cc = ( + (x * y).sum() + / torch.sqrt((x * x).sum() * (y * y).sum()).clamp(min=1e-30) + ).item() + if verbose: + print(f"add_noise: R-split = {r_split:.4f}, CC(half1, half2) = {cc:.4f}") + + I_mean = 0.5 * (I_h1 + I_h2) + sigma_I_mean = sigma_I / (2.0**0.5) + + # Only the amplitude is clamped; see the note in the docstring. + amp_noisy = torch.sqrt(I_mean.clamp(min=0.0)) + sigma_F = sigma_I_mean / (2.0 * self.fcalc_amp.clamp(min=1e-8)) + + fcalc_noisy = ( + amp_noisy * torch.exp(1j * self.fcalc_phase) + ).to(self.fcalc.dtype) + + new = FcalcDataset( + hkl=self.hkl.clone(), + resolution=( + self.resolution.clone() if self.resolution is not None else None + ), + cell=self.cell, + spacegroup=self.spacegroup, + device=self.device, + ) + new.set_fcalc(fcalc_noisy) + new.fobs_sigma = sigma_F.to(self.device) + new.I = I_mean.to(self.device) + new.I_sigma = sigma_I_mean.to(self.device) + return new + def write_mtz(self, filepath: str) -> None: """ Write Fcalc to MTZ as ``F-model`` / ``PH-model`` (phase in degrees). diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index 07681e7f..68de7251 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -1065,6 +1065,37 @@ def load_mtz( ).read(str(path)) return self.load(reader, french_wilson=french_wilson) + def load_crystfel_hkl( + self, path: str, cell, spacegroup, + ) -> "ReflectionData": + """ + Load a CrystFEL ``partialator`` ``.hkl`` reflection list. + + Unlike MTZ, the CrystFEL format carries no cell or space-group metadata, so both + must be supplied by the caller -- they usually live in a ``.cell`` file alongside. + + The format is intensity-native, so amplitudes are derived by French-Wilson on + load exactly as they are for an MTZ carrying I/SIGI columns. + + Parameters + ---------- + path : str + Path to the ``.hkl`` file. + cell : list | tuple | np.ndarray | Cell | torch.Tensor + Unit cell (a, b, c, alpha, beta, gamma). + spacegroup : str | gemmi.SpaceGroup | SpaceGroup + Space group identifier. + + Returns + ------- + ReflectionData + Self, for method chaining. + """ + from torchref.io import hkl as _hkl + + reader = _hkl.HKLReader(verbose=self.verbose).read(path, cell, spacegroup) + return self.load(reader) + def load_cif( self, path: Union[str, Path], diff --git a/torchref/io/hkl.py b/torchref/io/hkl.py new file mode 100644 index 00000000..17dd648e --- /dev/null +++ b/torchref/io/hkl.py @@ -0,0 +1,146 @@ +"""Reader for CrystFEL partialator reflection lists (`.hkl`). + +CrystFEL partialator output format:: + + CrystFEL reflection list version 2.0 + Symmetry: 2/m_uab + h k l I phase sigma(I) nmeas + 0 0 8 15631.77 - 2997.88 499 + ... + +The file carries intensities, sigmas, nmeas counts (and optional phase +strings), but **not** unit-cell or space-group metadata — those +typically live alongside in a ``.cell`` file. Cell and spacegroup must +therefore be supplied by the caller. + +Usage mirrors :class:`torchref.io.mtz.MTZReader`:: + + reader = HKLReader(verbose=1).read("td1.hkl", + cell=[a,b,c,al,be,ga], + spacegroup="P 1 21 1") + data_dict, cell, spacegroup = reader() + +The returned ``data_dict`` conforms to the intensity-path contract of +:meth:`ReflectionData.load` — ``{"HKL": (N,3) int, "I": (N,) float, +"SIGI": (N,) float, "I_col": "I (CrystFEL)"}`` — so French-Wilson kicks +in automatically and amplitudes are derived downstream. +""" + +from typing import Any, Optional, Tuple, Union + +import numpy as np + + +class HKLReader: + """Reader for CrystFEL partialator `.hkl` reflection lists.""" + + def __init__(self, verbose: int = 0): + self.verbose = verbose + self.data: Optional[dict] = None + self.cell: Optional[np.ndarray] = None + self.spacegroup: Optional[str] = None + self.nmeas: Optional[np.ndarray] = None + + def read( + self, + filepath: str, + cell: Union[list, tuple, np.ndarray, Any], + spacegroup: Union[str, Any], + ) -> "HKLReader": + """Parse a CrystFEL `.hkl` file. + + Parameters + ---------- + filepath : str + Path to the CrystFEL reflection list. + cell : list | tuple | np.ndarray | torchref.symmetry.Cell | torch.Tensor + Unit cell (a, b, c, alpha, beta, gamma). If a ``Cell`` object + is passed, ``.data`` is extracted. + spacegroup : str | gemmi.SpaceGroup | torchref.symmetry.SpaceGroup + Space group. If a wrapper object is passed, its ``.hm`` or + ``.short_name()`` is used. + """ + if self.verbose > 1: + print(f"Reading CrystFEL hkl file: {filepath}") + + # Normalize cell → (6,) np.ndarray + if hasattr(cell, "data"): # torchref.symmetry.Cell + cell = cell.data + if hasattr(cell, "detach"): # torch.Tensor + cell = cell.detach().cpu().numpy() + cell_arr = np.asarray(cell, dtype=float).reshape(-1) + if cell_arr.size != 6: + raise ValueError( + f"cell must have 6 entries (a, b, c, al, be, ga); got {cell_arr.size}" + ) + self.cell = cell_arr + + # Normalize spacegroup → HM-name string + if isinstance(spacegroup, str): + self.spacegroup = spacegroup + elif hasattr(spacegroup, "hm"): # torchref.symmetry.SpaceGroup + self.spacegroup = spacegroup.hm + elif hasattr(spacegroup, "short_name"): # gemmi.SpaceGroup + self.spacegroup = spacegroup.short_name() + else: + raise ValueError( + f"Cannot normalize spacegroup of type {type(spacegroup)}" + ) + + # Parse reflection rows + h_list, k_list, l_list, I_list, sig_list, n_list = [], [], [], [], [], [] + in_header = True + with open(filepath) as f: + for line in f: + if in_header: + if line.strip().startswith("h "): + in_header = False + continue + s = line.split() + if len(s) < 7 or not s[0].lstrip("-").isdigit(): + # Trailing comment lines or blank lines + continue + h_list.append(int(s[0])) + k_list.append(int(s[1])) + l_list.append(int(s[2])) + I_list.append(float(s[3])) + sig_list.append(float(s[5])) + n_list.append(int(s[6])) + + if not h_list: + raise ValueError(f"No reflections parsed from {filepath}") + + hkl = np.column_stack([h_list, k_list, l_list]).astype(np.int32) + I = np.asarray(I_list, dtype=np.float64) + sig = np.asarray(sig_list, dtype=np.float64) + self.nmeas = np.asarray(n_list, dtype=np.int32) + + self.data = { + "HKL": hkl, + "I": I, + "SIGI": sig, + "I_col": "I (CrystFEL)", + } + + if self.verbose > 0: + print( + f"Parsed {len(hkl)} reflections from {filepath} " + f"(cell={cell_arr.tolist()}, spacegroup='{self.spacegroup}')" + ) + return self + + def __call__(self) -> Tuple[dict, np.ndarray, str]: + """Return ``(data_dict, cell, spacegroup)`` for ``ReflectionData.load()``.""" + if self.data is None: + raise RuntimeError("Call .read(path, cell, spacegroup) first.") + return self.data, self.cell, self.spacegroup + + +def read( + filepath: str, + cell: Union[list, tuple, np.ndarray, Any], + spacegroup: Union[str, Any], + verbose: int = 0, +) -> HKLReader: + """Shortcut: ``HKLReader(verbose=verbose).read(filepath, cell, spacegroup)``.""" + return HKLReader(verbose=verbose).read(filepath, cell, spacegroup) From 7b19cb25fa25081fbd49d14d73d7df3ba5096422 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Wed, 26 Aug 2026 13:35:21 +0200 Subject: [PATCH 054/250] Give symmetry one home: a Symmetry base under SpaceGroup Symmetry was spread across four modules and re-derived per call site. The reciprocal transform R^T h alone had seven spellings -- two functions disagreeing on axis order and dtype, plus five open-coded transposes -- and its failure mode is silent: a wrong transpose corrupts centric flags and epsilon multiplicities without raising. Symmetry now owns the operations and everything derivable from them, built on three composable primitives (apply_rotations, apply_translations, phase_factors) plus a cached reciprocal stack carrying R^T. Miller-index expansion stops being a special case a caller can get wrong: it is apply_rotations on a different op stack. SpaceGroup specialises Symmetry with the crystallographic identity and the CCP4 asymmetric-unit verbs, so a group built from a raw operation list serves non-crystallographic symmetry without pretending to be a crystal. Translation keeps two conventions on purpose. phase_factors is the complex exp(+2 pi i h.t) used to combine structure factors; expand_hkl needs a signed radian offset -2 pi h.t to expand phases. Merging them flips a sign that is invisible in P21/P212121/C2. These classes hold no refinable parameters, so they are dataclasses over DeviceMixin rather than nn.Module. That also removes a hazard: nn.Module intercepted Module-valued assignment and bypassed the spacegroup property setter. Derived state lives in one cache that .to() and copy() both clear, and the traversal clears it before moving anything so a cached sampling grid is not copied to the new device only to be discarded. Map symmetrization moves behind Symmetry.symmetrize_map, caching one operator for the most recent grid shape. The MapSymmetry factory is gone: its variable return type was being read as a boolean, which the space group answers directly. cell_params is dropped from both operators -- stored and never read, because symmetry acts on fractional coordinates. Verified: centric and systematic-absence flags match gemmi with zero mismatches across ten space groups, including the trigonal, hexagonal and cubic cases; epsilon stays Friedel-inclusive at exactly twice gemmi's count. F_calc is bitwise identical to the previous commit on 1DAW, 2DQ6 and 3A5V, and so is the interpolating map path for both combine modes. Co-Authored-By: Claude Opus 5 (1M context) --- AGENTS.md | 2 +- docs/changelog.rst | 14 + tests/functional/test_model_ft_functional.py | 12 +- tests/helpers/device_cases.py | 56 +- .../integration/test_dtype_config_float64.py | 13 +- .../integration/test_symmetry_integration.py | 23 +- tests/unit/io/test_anomalous_output.py | 3 +- tests/unit/io/test_anomalous_reader.py | 3 +- tests/unit/structure_factor/helpers.py | 4 +- tests/unit/structure_factor/test_forward.py | 2 +- tests/unit/symmetry/test_canonicalize_hkl.py | 29 +- .../unit/symmetry/test_hkl_symmetry_gemmi.py | 41 +- tests/unit/symmetry/test_phase_convention.py | 31 +- tests/unit/symmetry/test_symmetry.py | 63 +- torchref/__init__.py | 3 +- torchref/base/__init__.py | 6 - torchref/base/french_wilson.py | 133 +-- torchref/base/reciprocal/__init__.py | 10 +- torchref/base/reciprocal/symmetry.py | 263 ++--- torchref/cli/mtz2map.py | 11 +- torchref/cli/validate_ded.py | 8 +- torchref/experimental/targets/realspace.py | 10 +- torchref/io/datasets/reflection_data.py | 47 +- torchref/maps/difference_map.py | 8 +- torchref/maps/map.py | 5 +- torchref/model/mixed_model.py | 5 - torchref/model/model_collection.py | 4 - torchref/model/model_ft.py | 15 +- torchref/model/sf_ds.py | 44 +- torchref/model/sf_fft.py | 84 +- .../model_error_estimation/sigma_a.py | 27 +- torchref/restraints/restraints.py | 2 +- torchref/symmetry/__init__.py | 94 +- torchref/symmetry/cell.py | 30 - torchref/symmetry/grid_utils.py | 134 --- torchref/symmetry/map_symmetry.py | 380 +++---- .../symmetry/map_symmetry_interpolation.py | 370 +++---- torchref/symmetry/reciprocal_symmetry.py | 931 ++---------------- torchref/symmetry/spacegroup.py | 885 ++++++----------- torchref/symmetry/symmetry.py | 779 ++++++++++++++- 40 files changed, 1797 insertions(+), 2787 deletions(-) delete mode 100644 torchref/symmetry/grid_utils.py diff --git a/AGENTS.md b/AGENTS.md index 5fd1707b..1f4ca6d5 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -190,7 +190,7 @@ Black, 88 columns, `isort` with the black profile. Ruff lint with | `refinement/` | Drivers (`Refinement`, `LBFGSRefinement`, `RigidBodyRefinementStep`), `targets/` (`xray/`, `geometry/`, `adp/`, `collection/`, `combined.py`), `weighting/`, `optimizers/` (annealing, Langevin, preconditioned/seeded L-BFGS), `model_error_estimation/` (σ_A, σ_M), `loss_state.py`, `logger.py` | | `restraints/` | Bonds, angles, torsions, planes, chirals, VDW. Built from the CCP4 Monomer Library, resolved lazily via `get_library_manager()` — importing this package must not trigger a library download | | `scaling/` | `ScalerBase` (model-independent), `Scaler`, `CollectionScaler`, `SolventModel` (k_sol, B_sol) | -| `symmetry/` | `SpaceGroup` (buffers on `nn.Module`), `Cell`, `MapSymmetry`, `ReciprocalSymmetry`, grid utilities | +| `symmetry/` | `Symmetry` (operations plus everything derived from them), `SpaceGroup` (adds the crystallographic identity and the CCP4 ASU verbs), `Cell`. All dataclasses over `DeviceMixin`, not `nn.Module` — they hold no refinable parameters. Map and reciprocal-grid operators are private, reached through `Symmetry` | | `maps/` | `Map` (2Fo−Fc, Fcalc), `DifferenceMap` | | `cli/` | Entry points: `torchref.refine`, `torchref.difference-refine`, `torchref.mtz2map`, `torchref.validate-ded`, `torchref.phased-difference-map`, `torchref.add-metadata`, `torchref.strip-altlocs` | | `experimental/` | APIs that may change without notice: `alignment/` (Patterson MR), `kinetic/` (time-resolved), `ensemble/`, `monolithic_refinement/`, `targets/` (AMBER/GAFF2, real-space, sampled-ML phase) | diff --git a/docs/changelog.rst b/docs/changelog.rst index 30fcbb07..ccaec3d8 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -2,6 +2,20 @@ Changelog ========= +Unreleased +---------- +- Added ``Symmetry``, a crystallography-free symmetry group carrying the operations and every verb derived from them; ``SpaceGroup`` now specialises it +- ``Symmetry`` exposes the transform primitives ``apply_rotations`` / ``apply_translations`` / ``phase_factors`` and a cached ``reciprocal`` stack, replacing seven separate spellings of ``R^T h`` +- Moved the centric, systematic-absence and epsilon predicates onto ``Symmetry``; ``is_centric_from_hkl`` and ``get_centric_acentric_masks`` are gone +- Moved the HKL asymmetric-unit verbs onto ``SpaceGroup`` as ``expand_hkl`` / ``reduce_hkl`` / ``complete_hkl`` / ``canonicalize_hkl``; the module-level functions are gone +- Moved the grid-size helpers onto ``Symmetry``; removed ``torchref.symmetry.grid_utils`` and the duplicate ``spacegroup`` module-level functions +- Map symmetrization is now ``Symmetry.symmetrize_map``, caching one operator for the most recent grid shape; ``MapSymmetry`` and ``MapSymmetryDirect`` are private +- Symmetry classes are dataclasses over ``DeviceMixin`` instead of ``nn.Module``, so assigning a ``SpaceGroup`` to a model attribute is no longer intercepted by ``nn.Module.__setattr__`` +- Removed ``ReciprocalSymmetryGrid``, ``ReciprocalSymmetry``, ``expand_reciprocal_grid``, ``expand_reflections`` and ``extract_structure_factors_with_symmetry`` +- Removed the unused ``Cell`` gradient plumbing (``requires_grad`` argument and property, ``detach``) and the ``CellTensor`` alias +- Removed the ``Symmetry`` alias for ``SpaceGroup``; the name is now a distinct class + + Version 0.6.4 ---------- - Fixed the bulk-solvent ``F_sol`` staying at the starting model's mask for every refinement macrocycle diff --git a/tests/functional/test_model_ft_functional.py b/tests/functional/test_model_ft_functional.py index 3a430cc1..a0606d87 100644 --- a/tests/functional/test_model_ft_functional.py +++ b/tests/functional/test_model_ft_functional.py @@ -145,15 +145,13 @@ def test_map_symmetry_available(self, sample_cif_file): # Model should have spacegroup after loading assert model.spacegroup is not None - # Map symmetry can be created if gridsize is available + # The map operator comes from the space group, keyed on the grid shape. if model.gridsize is not None: - from torchref.symmetry.map_symmetry import MapSymmetry - gridsize = tuple(model.gridsize.tolist()) - cell_params = model.cell - - map_sym = MapSymmetry(model.spacegroup, gridsize, cell_params) - assert map_sym is not None + + operator = model.spacegroup.map_operator(gridsize) + assert operator is not None + assert operator.map_shape == gridsize @pytest.mark.integration diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index 7a7515b3..8e3afc8c 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -56,6 +56,34 @@ class DeviceCase: ignore: tuple = field(default_factory=tuple) +def _symmetry(device): + """A bare Symmetry from an explicit operation list (no space group involved).""" + import torch as _torch + + from torchref.config import get_float_dtype + from torchref.symmetry import Symmetry + + dtype = get_float_dtype() + matrices = _torch.eye(3, dtype=dtype, device=device).unsqueeze(0).repeat(2, 1, 1) + matrices[1] = -matrices[1] + translations = _torch.zeros(2, 3, dtype=dtype, device=device) + return Symmetry(matrices=matrices, translations=translations) + + +def _map_symmetry_interpolation(device): + """The interpolating map operator, on a grid that forbids direct indexing. + + P212121 requires even dimensions, so an odd grid forces the interpolating + variant rather than the streaming one. + """ + from torchref.symmetry import SpaceGroup + from torchref.symmetry.map_symmetry_interpolation import ( + _MapSymmetryInterpolation, + ) + + return _MapSymmetryInterpolation(SpaceGroup(_SG, device=device), (15, 15, 15)) + + def _cell(device): from torchref.symmetry import Cell @@ -165,20 +193,23 @@ def _cell(device): "RigidXYZTensor", ), DeviceCase( - "ReciprocalSymmetryGrid", + "SpaceGroup", lambda d: __import__( - "torchref.symmetry", fromlist=["ReciprocalSymmetryGrid"] - ).ReciprocalSymmetryGrid(_SG, grid_shape=(16, 16, 16), device=d), - "ReciprocalSymmetryGrid", + "torchref.symmetry", fromlist=["SpaceGroup"] + ).SpaceGroup(_SG, device=d), + "SpaceGroup", ), DeviceCase( - "MapSymmetryDirect", - lambda d: __import__( - "torchref.symmetry", fromlist=["MapSymmetryDirect"] - ).MapSymmetryDirect( - _SG, map_shape=(16, 16, 16), cell_params=_CELL, device=d - ), - "MapSymmetryDirect", + "Symmetry", + _symmetry, + "Symmetry", + ), + DeviceCase( + "_MapSymmetryInterpolation", + _map_symmetry_interpolation, + "_MapSymmetryInterpolation", + # ``symmetry`` is the group this operator was built from, not state it owns. + ignore=("symmetry",), ), DeviceCase( "TensorMasks", @@ -294,7 +325,8 @@ class TargetDeviceCase: "Map": "needs data + model", "DifferenceMap": "needs two datasets + a model", "LBFGSRefinement": "full pipeline; covered in integration", - "MapSymmetry": "interpolation variant; needs a real map grid", + "_MapSymmetryDirect": "stateless view over its Symmetry: recomputes index grids " + "per operation to keep peak memory O(grid), so it owns no tensors to move", "CholeskyMixedTensor": "needs a valid ADP tensor; shares MixedTensor's paths", "CollectionScaler": "needs a dataset collection", "CollectionDifferenceTarget": "needs a dataset collection", diff --git a/tests/integration/test_dtype_config_float64.py b/tests/integration/test_dtype_config_float64.py index 942cca31..34738111 100644 --- a/tests/integration/test_dtype_config_float64.py +++ b/tests/integration/test_dtype_config_float64.py @@ -14,15 +14,16 @@ @pytest.mark.unit def test_translation_phases_complex_dtype_float64(double_cpu): - """compute_translation_phases must honor the configured complex dtype.""" - from torchref.base.reciprocal.symmetry import compute_translation_phases + """Symmetry.phase_factors must honor the configured complex dtype.""" + from torchref.symmetry import SpaceGroup - hkl = torch.tensor([[1.0, 0.0, 0.0], [2.0, 1.0, 0.0], [0.0, 0.0, 3.0]]) - translations = torch.tensor([[0.0, 0.0, 0.0], [0.5, 0.5, 0.0]]) + # P21 gives two operations, one carrying a half translation. + sym = SpaceGroup("P 21") + hkl = torch.tensor([[1, 0, 0], [2, 1, 0], [0, 0, 3]]) - phases = compute_translation_phases(hkl, translations) + phases = sym.phase_factors(hkl) - # Was complex64 (float32 hardcode); under float64 config must be complex128. + # Must not narrow to complex64 under a float64 configuration. assert phases.dtype == torch.complex128 assert phases.shape == (2, 3) assert torch.isfinite(phases.real).all() diff --git a/tests/integration/test_symmetry_integration.py b/tests/integration/test_symmetry_integration.py index 9ccedc52..8bb0a52d 100644 --- a/tests/integration/test_symmetry_integration.py +++ b/tests/integration/test_symmetry_integration.py @@ -221,18 +221,21 @@ class TestMapSymmetry: """Tests for map symmetry operations.""" @pytest.mark.integration - def test_map_symmetry_initialization(self, sample_cif_file): - """Test map symmetry initialization.""" + def test_symmetrize_map_round_trip(self, sample_cif_file): + """Symmetrizing a map goes through the space group and preserves shape.""" + import torch + from torchref.model.model import Model - from torchref.symmetry.map_symmetry import MapSymmetry model = Model() model.load_cif(str(sample_cif_file)) - # Check if MapSymmetry can be initialized - try: - map_sym = MapSymmetry(model.spacegroup, model.cell) - assert map_sym is not None - except (TypeError, AttributeError): - # May not support all initialization patterns - pass + sg = model.spacegroup + shape = sg.suggest_grid_size((16, 16, 16)) + density = torch.rand(shape, device=sg.device, dtype=sg.dtype) + + symmetrized = sg.symmetrize_map(density) + + assert symmetrized.shape == density.shape + # The operator is cached for this shape and dropped on a device move. + assert sg.map_operator(shape) is sg.map_operator(shape) diff --git a/tests/unit/io/test_anomalous_output.py b/tests/unit/io/test_anomalous_output.py index e110a6e2..b32a2e68 100644 --- a/tests/unit/io/test_anomalous_output.py +++ b/tests/unit/io/test_anomalous_output.py @@ -16,7 +16,6 @@ import reciprocalspaceship as rs import torch -from torchref.base.french_wilson import is_centric_from_hkl from torchref.io.datasets.reflection_data import ReflectionData from torchref.model.model_ft import ModelFT @@ -68,7 +67,7 @@ def test_hkl_for_sf_signs(self, anomalous_data): def test_centrics_not_flagged(self, anomalous_data): d = anomalous_data - centric = is_centric_from_hkl(d.hkl, d.spacegroup) + centric = d.spacegroup.is_centric(d.hkl) assert not bool((d.friedel_flags & centric).any()) def test_hkl_for_sf_fallback(self): diff --git a/tests/unit/io/test_anomalous_reader.py b/tests/unit/io/test_anomalous_reader.py index 2fc0692d..90641680 100644 --- a/tests/unit/io/test_anomalous_reader.py +++ b/tests/unit/io/test_anomalous_reader.py @@ -15,7 +15,6 @@ import reciprocalspaceship as rs import torch -from torchref.base.french_wilson import is_centric_from_hkl from torchref.io.datasets.reflection_data import ReflectionData from torchref.model.model_ft import ModelFT @@ -81,7 +80,7 @@ def test_centrics_not_duplicated(anomalous_two_column_mtz): path, _ = anomalous_two_column_mtz d = ReflectionData(verbose=0) d.load_mtz(path) - centric = is_centric_from_hkl(d.hkl, d.spacegroup) + centric = d.spacegroup.is_centric(d.hkl) # Centric reflections obey Friedel's law and must appear exactly once each. canon = [tuple(h) for h in d.hkl[centric].tolist()] assert len(canon) == len(set(canon)) diff --git a/tests/unit/structure_factor/helpers.py b/tests/unit/structure_factor/helpers.py index a853925f..54bfca15 100644 --- a/tests/unit/structure_factor/helpers.py +++ b/tests/unit/structure_factor/helpers.py @@ -684,8 +684,8 @@ def fft_sf(scene: Scene, sf_fft, xyz=None, occ=None, third=None, *, aniso=False) ``apply_symmetry=False`` on both calls: P1 isolates the truncation and sampling budget, and a symmetric comparison would cancel the symmetry algebra anyway, since - both routes call the same ``compute_symmetry_equivalent_hkls`` / - ``compute_translation_phases``. Symmetry is validated against gemmi instead, in + both routes call the same ``Symmetry.expand_reciprocal`` / + ``Symmetry.phase_factors``. Symmetry is validated against gemmi instead, in ``test_forward.py``. """ xyz = scene.xyz if xyz is None else xyz diff --git a/tests/unit/structure_factor/test_forward.py b/tests/unit/structure_factor/test_forward.py index 0d4caa00..78fe6915 100644 --- a/tests/unit/structure_factor/test_forward.py +++ b/tests/unit/structure_factor/test_forward.py @@ -245,7 +245,7 @@ def test_sfds_matches_gemmi_with_symmetry(gemmi_iso_symmetry): This is also the only symmetric comparison in the package. A DS-vs-FFT check cannot validate symmetry, because both routes call the same - ``compute_symmetry_equivalent_hkls`` / ``compute_translation_phases`` and the shared + ``Symmetry.expand_reciprocal`` / ``Symmetry.phase_factors`` and the shared algebra cancels; gemmi does not share it. """ scene, structure = gemmi_iso_symmetry diff --git a/tests/unit/symmetry/test_canonicalize_hkl.py b/tests/unit/symmetry/test_canonicalize_hkl.py index 60471c7f..1764baf6 100644 --- a/tests/unit/symmetry/test_canonicalize_hkl.py +++ b/tests/unit/symmetry/test_canonicalize_hkl.py @@ -4,7 +4,31 @@ import pytest import torch -from torchref.symmetry.reciprocal_symmetry import canonicalize_hkl +from torchref.symmetry import SpaceGroup + + +# The HKL verbs live on the space group now. These adapters keep the assertions +# below -- which pin the phase-sign contract -- expressed in terms of the space +# group specifications the cases are parametrised over. +def canonicalize_hkl(hkl, sg, include_friedel=True, device=None): + return SpaceGroup(sg).canonicalize_hkl( + hkl, include_friedel=include_friedel, device=device + ) + + +def expand_hkl(hkl, sg, include_friedel=True, remove_absences=True, device=None): + return SpaceGroup(sg).expand_hkl( + hkl, + include_friedel=include_friedel, + remove_absences=remove_absences, + device=device, + ) + + +def reduce_hkl(hkl, sg, include_friedel=True, device=None): + return SpaceGroup(sg).reduce_hkl( + hkl, include_friedel=include_friedel, device=device + ) # --------------------------------------------------------------------------- @@ -41,8 +65,6 @@ def test_idempotency(self, sg): ) def test_equivalents_converge(self, sg): """All symmetry equivalents of a reflection map to the same canonical HKL.""" - from torchref.symmetry.reciprocal_symmetry import expand_hkl - # Use (1,1,5) which satisfies centering conditions for C2 hkl_asu = torch.tensor([[1, 1, 5]], dtype=torch.int32) hkl_p1, _, _ = expand_hkl(hkl_asu, sg, include_friedel=True) @@ -70,7 +92,6 @@ def test_phase_roundtrip(self, sg): then verify canonicalization produces a single consistent SF. """ import gemmi - from torchref.symmetry.spacegroup import SpaceGroup # Reference SF at canonical h F_ref = 10.0 diff --git a/tests/unit/symmetry/test_hkl_symmetry_gemmi.py b/tests/unit/symmetry/test_hkl_symmetry_gemmi.py index 4536f9c5..1bd1e126 100644 --- a/tests/unit/symmetry/test_hkl_symmetry_gemmi.py +++ b/tests/unit/symmetry/test_hkl_symmetry_gemmi.py @@ -1,12 +1,12 @@ """Regression tests for reciprocal-space (Miller-index) symmetry against gemmi. These guard the ``h' = h·R = Rᵀ·h`` reciprocal-space convention used by -``SpaceGroup.apply_to_hkl``. A previous bug applied the real-space transform +``Symmetry.expand_reciprocal``. A previous bug applied the real-space transform ``R·h`` instead, which is only correct when the fractional rotation matrix is symmetric. For space groups whose rotation matrices are non-symmetric (trigonal, hexagonal, and permutation-type cubic operations) ``R·h`` produced the wrong set of symmetry equivalents, corrupting the centric flags -(``is_centric_from_hkl``) and epsilon multiplicities (``epsilon_from_hkl``) +(``Symmetry.is_centric``) and epsilon multiplicities (``Symmetry.epsilon``) that feed French-Wilson intensity conversion and ML sigma_A weighting. Ground truth comes from gemmi: @@ -14,7 +14,7 @@ - centric: ``GroupOps.is_reflection_centric`` - epsilon: ``GroupOps.epsilon_factor_without_centering`` (Friedel-doubled for centric reflections, matching the Friedel-aware count in - ``epsilon_from_hkl``) + ``Symmetry.epsilon``) """ import numpy as np @@ -24,7 +24,7 @@ from torchref.config import get_default_device, get_float_dtype, get_int_dtype # Space groups whose rotation matrices are non-symmetric — these are the ones -# that regress if apply_to_hkl uses R·h instead of Rᵀ·h. Plus orthorhombic / +# that regress if expand_reciprocal uses R·h instead of Rᵀ·h. Plus orthorhombic / # tetragonal controls (symmetric matrices) that must remain correct either way. _TRIGONAL_HEXAGONAL = ["P 32 2 1", "P 31 2 1", "P 61 2 2", "P 6 2 2", "P 3 1 2"] _CONTROLS = ["P 21 21 21", "P 43 21 2", "P 1"] @@ -53,8 +53,8 @@ def _hkl_tensor(): @pytest.mark.unit @pytest.mark.parametrize("sg_name", _SPACE_GROUPS) -def test_apply_to_hkl_matches_gemmi_transform(sg_name): - """apply_to_hkl must reproduce gemmi's per-operation reciprocal transform. +def test_expand_reciprocal_matches_gemmi_transform(sg_name): + """expand_reciprocal must reproduce gemmi's per-operation reciprocal transform. Compared as the *set* of equivalents per reflection so the test is insensitive to operation ordering between torchref and gemmi. @@ -66,33 +66,32 @@ def test_apply_to_hkl_matches_gemmi_transform(sg_name): sg = SpaceGroup(sg_name) hkl = _hkl_tensor() - # torchref: (N, 3, ops) -> per-reflection set of equivalents - out = sg.apply_to_hkl(hkl.to(get_float_dtype())) - out_int = torch.round(out).to(torch.int64).cpu().numpy() # (N, 3, ops) + # torchref: (ops, N, 3) -> per-reflection set of equivalents + out_int = sg.expand_reciprocal(hkl).cpu().numpy() # (ops, N, 3) gemmi_ops = list(gemmi.SpaceGroup(sg_name).operations().sym_ops) for n, h in enumerate(_HKLS): - got = {tuple(out_int[n, :, o]) for o in range(out_int.shape[2])} + got = {tuple(out_int[o, n, :]) for o in range(out_int.shape[0])} expected = {tuple(op.apply_to_hkl(h)) for op in gemmi_ops} assert got == expected, ( - f"{sg_name} hkl={h}: apply_to_hkl equivalents {sorted(got)} " + f"{sg_name} hkl={h}: expand_reciprocal equivalents {sorted(got)} " f"!= gemmi {sorted(expected)}" ) @pytest.mark.unit @pytest.mark.parametrize("sg_name", _SPACE_GROUPS) -def test_is_centric_from_hkl_matches_gemmi(sg_name): +def test_is_centric_matches_gemmi(sg_name): """Centric flags must match gemmi's is_reflection_centric.""" import gemmi - from torchref.base.french_wilson import is_centric_from_hkl + from torchref.symmetry import SpaceGroup ops = gemmi.SpaceGroup(sg_name).operations() hkl = _hkl_tensor() - got = is_centric_from_hkl(hkl, sg_name).cpu().numpy().astype(bool) + got = SpaceGroup(sg_name).is_centric(hkl).cpu().numpy().astype(bool) expected = np.array([ops.is_reflection_centric(h) for h in _HKLS], dtype=bool) assert np.array_equal(got, expected), ( @@ -129,7 +128,7 @@ def test_epsilon_from_hkl_matches_gemmi(sg_name): @pytest.mark.unit -def test_apply_to_hkl_is_transpose_not_plain_rotation(): +def test_expand_reciprocal_is_transpose_not_plain_rotation(): """Explicit guard on the exact bug: apply_to_hkl == Rᵀ·h, not R·h. Uses P3, whose 3-fold rotation matrix is non-symmetric, so R·h and Rᵀ·h @@ -143,14 +142,16 @@ def test_apply_to_hkl_is_transpose_not_plain_rotation(): mats = sg.matrices # (ops, 3, 3) hkl = torch.tensor([[1, 2, 3]], dtype=mats.dtype, device=mats.device) - out = sg.apply_to_hkl(hkl) # (1, 3, ops) + out = sg.expand_reciprocal(hkl) # (ops, 1, 3) # Correct reciprocal transform: Rᵀ·h - expected_t = torch.einsum("oji,nj->nio", mats, hkl) - # The old (buggy) real-space transform: R·h - plain_r = torch.einsum("oij,nj->nio", mats, hkl) + expected_t = torch.einsum("oji,nj->oni", mats, hkl) + # The buggy real-space transform: R·h + plain_r = torch.einsum("oij,nj->oni", mats, hkl) - assert torch.allclose(out, expected_t), "apply_to_hkl should compute Rᵀ·h" + assert torch.allclose(out.to(mats.dtype), expected_t), ( + "expand_reciprocal should compute Rᵀ·h" + ) # For P3 the two conventions must actually differ (non-symmetric matrix), # otherwise this test would not detect a regression. assert not torch.allclose(expected_t, plain_r), ( diff --git a/tests/unit/symmetry/test_phase_convention.py b/tests/unit/symmetry/test_phase_convention.py index 2d6436e4..6647eb63 100644 --- a/tests/unit/symmetry/test_phase_convention.py +++ b/tests/unit/symmetry/test_phase_convention.py @@ -35,12 +35,31 @@ import torch from torchref.config import get_float_dtype -from torchref.symmetry.reciprocal_symmetry import ( - canonicalize_hkl, - expand_hkl, - reduce_hkl, -) -from torchref.symmetry.spacegroup import SpaceGroup +from torchref.symmetry import SpaceGroup + + +# The HKL verbs live on the space group now. These adapters keep the assertions +# below -- which pin the phase-sign contract -- expressed in terms of the space +# group specifications the cases are parametrised over. +def canonicalize_hkl(hkl, sg, include_friedel=True, device=None): + return SpaceGroup(sg).canonicalize_hkl( + hkl, include_friedel=include_friedel, device=device + ) + + +def expand_hkl(hkl, sg, include_friedel=True, remove_absences=True, device=None): + return SpaceGroup(sg).expand_hkl( + hkl, + include_friedel=include_friedel, + remove_absences=remove_absences, + device=device, + ) + + +def reduce_hkl(hkl, sg, include_friedel=True, device=None): + return SpaceGroup(sg).reduce_hkl( + hkl, include_friedel=include_friedel, device=device + ) # Groups spanning the three regimes above. P1 is the degenerate control (no # translations at all); the screw-axis groups are the ones with real signal. diff --git a/tests/unit/symmetry/test_symmetry.py b/tests/unit/symmetry/test_symmetry.py index 1ddab2bf..aab7bcce 100644 --- a/tests/unit/symmetry/test_symmetry.py +++ b/tests/unit/symmetry/test_symmetry.py @@ -143,20 +143,15 @@ def test_apply_identity(self, random_fractional_coordinates): from torchref.symmetry import SpaceGroup sg = SpaceGroup("P1") - # SpaceGroup expects (N, 3) format coords = random_fractional_coordinates(n_atoms=10) # Shape: (N, 3) - # Apply symmetry (P1 only has identity) - # Output shape is (n_atoms, 3, n_operations) - transformed = sg(coords) + # Expansions are operation-major: (n_ops, N, 3) + transformed = sg.expand_positions(coords) - # Should have shape (n_atoms, 3, n_operations) - assert transformed.shape[0] == 10 # 10 atoms - assert transformed.shape[1] == 3 # 3D coordinates - assert transformed.shape[2] == 1 # 1 operation (identity) - # First (and only) symmetry mate should match original + assert transformed.shape == (1, 10, 3) + # The only mate is the identity assert torch.allclose( - transformed[:, :, 0], + transformed[0], coords.to(device=transformed.device, dtype=transformed.dtype), atol=1e-5, ) @@ -169,26 +164,21 @@ def test_spacegroup_generates_mates(self, random_fractional_coordinates): sg = SpaceGroup("P21") # 2 operations coords = random_fractional_coordinates(n_atoms=5) # (N, 3) format - # Output shape is (n_atoms, 3, n_operations) - transformed = sg(coords) + transformed = sg.expand_positions(coords) - # Should have 2 symmetry operations - assert transformed.shape[2] == 2 + assert transformed.shape == (2, 5, 3) @pytest.mark.unit - def test_spacegroup_callable(self, random_fractional_coordinates): - """SpaceGroup should be callable.""" + def test_expand_to_p1_flattens(self, random_fractional_coordinates): + """expand_to_P1 flattens the operation axis into one coordinate list.""" from torchref.symmetry import SpaceGroup sg = SpaceGroup("P212121") coords = random_fractional_coordinates(n_atoms=10) # (N, 3) format - # Should be callable - # Output shape is (n_atoms, 3, n_operations) - result = sg(coords) + result = sg.expand_to_P1(coords) - assert result is not None - assert result.shape[2] == 4 # 4 symmetry operations + assert result.shape == (40, 3) # 4 operations x 10 atoms class TestSpaceGroupDeviceHandling: @@ -252,24 +242,31 @@ def test_case_insensitivity(self): pass # Some case variations may not be supported -class TestSymmetryBackwardCompat: - """Tests for backward compatibility with Symmetry alias.""" +class TestSymmetryBase: + """Tests for the crystallography-free Symmetry base class.""" @pytest.mark.unit - def test_symmetry_alias_exists(self): - """Test that Symmetry alias is available.""" - from torchref.symmetry import Symmetry, SpaceGroup + def test_spacegroup_is_a_symmetry(self): + """SpaceGroup specialises Symmetry, so ops-only code accepts either.""" + from torchref.symmetry import SpaceGroup, Symmetry - # Symmetry should be an alias for SpaceGroup - assert Symmetry is SpaceGroup + assert issubclass(SpaceGroup, Symmetry) + assert isinstance(SpaceGroup("P21"), Symmetry) @pytest.mark.unit - def test_symmetry_alias_works(self): - """Test that Symmetry alias works identically to SpaceGroup.""" + def test_symmetry_from_raw_operations(self): + """A Symmetry can be built from an operation list with no space group.""" + import torch + from torchref.symmetry import Symmetry - sym = Symmetry("P21") + matrices = torch.eye(3).unsqueeze(0).repeat(2, 1, 1) + matrices[1] = -matrices[1] + translations = torch.zeros(2, 3) + + sym = Symmetry(matrices=matrices, translations=translations) - assert sym.matrices is not None assert sym.n_ops == 2 - assert sym.name == "P21" + # An inversion pair makes every reflection centric. + hkl = torch.tensor([[1, 2, 3], [4, 0, 1]]) + assert bool(sym.is_centric(hkl).all()) diff --git a/torchref/__init__.py b/torchref/__init__.py index 5f580fd4..0df11cb3 100644 --- a/torchref/__init__.py +++ b/torchref/__init__.py @@ -106,7 +106,7 @@ # Refinement from torchref.refinement import LBFGSRefinement, Refinement from torchref.refinement.rigid_body_refinement import RigidBodyRefinementStep -from torchref.symmetry import Cell, SpaceGroup +from torchref.symmetry import Cell, SpaceGroup, Symmetry # Restraints # from torchref.restraints import Restraints # Initialized lazily due to monomer library download requirement @@ -151,6 +151,7 @@ # Symmetry "Cell", "SpaceGroup", + "Symmetry", # Maps "Map", "DifferenceMap", diff --git a/torchref/base/__init__.py b/torchref/base/__init__.py index bfbfe8c7..e62d4747 100644 --- a/torchref/base/__init__.py +++ b/torchref/base/__init__.py @@ -119,9 +119,6 @@ interpolate_for_rotation, smooth_reciprocal_grid, # Symmetry - compute_symmetry_equivalent_hkls, - compute_translation_phases, - extract_structure_factors_with_symmetry, ReciprocalSymmetryExtractor, ) @@ -283,9 +280,6 @@ "interpolate_structure_factor_from_grid", "interpolate_complex_from_grid", "trilinear_interpolate_patterson", - "compute_symmetry_equivalent_hkls", - "compute_translation_phases", - "extract_structure_factors_with_symmetry", "interpolate_for_rotation", "smooth_reciprocal_grid", # Structure factors diff --git a/torchref/base/french_wilson.py b/torchref/base/french_wilson.py index 60602f78..a48128db 100644 --- a/torchref/base/french_wilson.py +++ b/torchref/base/french_wilson.py @@ -1146,135 +1146,6 @@ def french_wilson( return F, sigma_F, valid_mask -def is_centric_from_hkl( - hkl: torch.Tensor, space_group: SpaceGroupLike = "P1" -) -> torch.Tensor: - """ - Determine if reflections are centric based on Miller indices and space group. - - Uses symmetry operations to check if reflections are invariant under - inversion through the origin (Friedel mates). A reflection is centric - if -h,-k,-l is symmetry equivalent to h,k,l. - - Parameters - ---------- - hkl : torch.Tensor - Miller indices of shape (..., 3). - space_group : str, int, or gemmi.SpaceGroup, optional - Space group specification. Default is "P1". - - Returns - ------- - torch.Tensor - Boolean mask of shape (...), True for centric reflections. - """ - original_shape = hkl.shape[:-1] - hkl_flat = hkl.reshape(-1, 3) - n_reflections = hkl_flat.shape[0] - - # Get symmetry operations from the SpaceGroup class - float_dtype = get_float_dtype() - spacegroup = SpaceGroup(space_group, dtype=float_dtype, device=hkl.device) - - # Convert HKL to the configured float dtype for symmetry operations - hkl_float = hkl_flat.to(float_dtype) # Shape: (n_reflections, 3) - - # Apply all spacegroup operations to all reflections at once - # For reciprocal space (Miller indices), only rotation applies, not translation - # hkl_float shape: (n_reflections, 3) - # spacegroup.apply_to_hkl returns shape: (n_reflections, 3, n_ops) - hkl_sym = spacegroup.apply_to_hkl(hkl_float) - - # Compute Friedel mates: -h, -k, -l - # Shape: (n_reflections, 3, 1) to broadcast against (n_reflections, 3, n_ops) - friedel_hkl = -hkl_float.unsqueeze(-1) # Shape: (n_reflections, 3, 1) - - # Check if any spacegroup operation produces the Friedel mate - # Round to nearest integer (Miller indices should be integers) - hkl_sym_rounded = torch.round(hkl_sym) - - # Compute difference for all reflections and all spacegroup operations - # Shape: (n_reflections, 3, n_ops) - diff = torch.abs(hkl_sym_rounded - friedel_hkl) - - # A reflection is centric if ANY spacegroup operation maps it to its Friedel mate - # Check if all 3 components (h,k,l) match (diff < 0.5) for any operation - # Shape: (n_reflections, n_ops) after checking all 3 components match - matches = torch.all(diff < 0.5, dim=1) # Check all 3 Miller indices match - - # A reflection is centric if it matches for ANY spacegroup operation - # Shape: (n_reflections,) - is_centric = torch.any(matches, dim=1) - - return is_centric.reshape(original_shape) - - -def epsilon_from_hkl(hkl: torch.Tensor, spacegroup) -> torch.Tensor: - """Per-reflection epsilon: number of rotation symops mapping h -> +/-h. - - Mirrors ``ReciprocalSymmetry.get_epsilon`` (Friedel-aware) but works directly - on the scattered HKL list. Returns ones if ``spacegroup`` is None or lacks - ``apply_to_hkl``. - - Unlike :func:`is_centric_from_hkl` this takes a constructed space group rather - than a specification, because its callers already hold one. - - Always returns on ``hkl.device``, whatever device the space group's symmetry - matrices live on: the caller multiplies this against per-reflection data - sitting beside ``hkl``. - """ - n = hkl.shape[0] - float_dtype = get_float_dtype() - if spacegroup is None or not hasattr(spacegroup, "apply_to_hkl"): - return torch.ones(n, device=hkl.device, dtype=float_dtype) - - with torch.no_grad(): - # Configured float dtype, not float64: MPS has no float64 and casting - # there raises. Symmetry arithmetic on Miller indices is exact in - # float32 (integer-valued rotation matrices, small indices), so the - # exact `==` comparisons below remain valid. - # - # ``apply_to_hkl`` moves its input onto the matrices' device, so build - # ``h`` there too -- otherwise ``Hs`` and ``h0`` land on different - # devices and the comparisons below raise. The space group wins for the - # arithmetic; the result is handed back on the caller's device. - sym_device = getattr(spacegroup, "matrices", hkl).device - h = hkl.to(device=sym_device, dtype=float_dtype) - Hs = spacegroup.apply_to_hkl(h) # (N,3,ops) - h0 = h.unsqueeze(-1) # (N,3,1) - same = (Hs == h0).all(dim=1) - friedel = (Hs == -h0).all(dim=1) - eps = (same | friedel).sum(dim=1).clamp(min=1).to(float_dtype) - return eps.to(hkl.device) - - -def get_centric_acentric_masks( - hkl: torch.Tensor, space_group: SpaceGroupLike = "P1" -) -> tuple[torch.Tensor, torch.Tensor]: - """ - Get both centric and acentric masks for reflections. - - Convenience function that returns both masks explicitly. - - Parameters - ---------- - hkl : torch.Tensor - Miller indices of shape (..., 3). - space_group : str, int, or gemmi.SpaceGroup, optional - Space group specification. Default is "P1". - - Returns - ------- - centric_mask : torch.Tensor - Boolean mask of shape (...), True for centric reflections. - acentric_mask : torch.Tensor - Boolean mask of shape (...), True for acentric reflections. - """ - centric_mask = is_centric_from_hkl(hkl, space_group) - acentric_mask = ~centric_mask - return centric_mask, acentric_mask - - def estimate_mean_intensity_by_resolution( I: torch.Tensor, d_spacings: torch.Tensor, n_bins: int = 60, min_per_bin: int = 40 ) -> torch.Tensor: @@ -1460,7 +1331,7 @@ def french_wilson_auto( ) # Step 2: Determine centric reflections from Miller indices - is_centric = is_centric_from_hkl(hkl, space_group=space_group) + is_centric = SpaceGroup(space_group, device=hkl.device).is_centric(hkl) # Step 3: Apply French-Wilson conversion F, sigma_F, valid_mask = french_wilson( @@ -1546,7 +1417,7 @@ def __init__( self.register_buffer("d_spacings", d_spacings) # Determine centric reflections - is_centric = is_centric_from_hkl(hkl, space_group) + is_centric = SpaceGroup(space_group, device=hkl.device).is_centric(hkl) self.register_buffer("is_centric", is_centric) # Set by forward(); None until the first conversion. Not a buffer -- it diff --git a/torchref/base/reciprocal/__init__.py b/torchref/base/reciprocal/__init__.py index 00003aac..cdaf6c20 100644 --- a/torchref/base/reciprocal/__init__.py +++ b/torchref/base/reciprocal/__init__.py @@ -33,12 +33,7 @@ smooth_reciprocal_grid ) -from .symmetry import ( - compute_symmetry_equivalent_hkls, - compute_translation_phases, - extract_structure_factors_with_symmetry, - ReciprocalSymmetryExtractor, -) +from .symmetry import ReciprocalSymmetryExtractor __all__ = [ # Basis functions @@ -62,8 +57,5 @@ "interpolate_for_rotation", "smooth_reciprocal_grid", # Symmetry - "compute_symmetry_equivalent_hkls", - "compute_translation_phases", - "extract_structure_factors_with_symmetry", "ReciprocalSymmetryExtractor", ] diff --git a/torchref/base/reciprocal/symmetry.py b/torchref/base/reciprocal/symmetry.py index 72d4ed0d..bd50b1de 100644 --- a/torchref/base/reciprocal/symmetry.py +++ b/torchref/base/reciprocal/symmetry.py @@ -1,215 +1,89 @@ """Reciprocal-space ("late") symmetry for structure factor calculation. -The alternative to symmetrizing the density map before the FFT ("early" symmetry, -:func:`~torchref.symmetry.MapSymmetry`): here symmetry is applied to the P1 -transform afterwards, avoiding the map symmetrization entirely. Per operation +The alternative to symmetrizing the density map before the FFT ("early" symmetry, via +:meth:`~torchref.symmetry.symmetry.Symmetry.symmetrize_map`): here symmetry is applied +to the P1 transform afterwards, avoiding the map symmetrization entirely. Per operation {R|t}, - F_sym(h) = Σ_ops exp(2πi h·t) · F_P1(Rᵀ·h) + F_sym(h) = sum_ops exp(2 pi i h.t) * F_P1(R^T h) -and because crystallographic R is integer-valued, Rᵀ·h lands exactly on grid -points and needs no interpolation. **Every grid or map argument here is the P1 -one** -- feeding in an already-symmetrized grid double-counts. +and because crystallographic R is integer-valued, ``R^T h`` lands exactly on grid +points and needs no interpolation. **Every grid or map argument here is the P1 one** -- +feeding in an already-symmetrized grid double-counts. + +Both halves of that sum come from :class:`~torchref.symmetry.symmetry.Symmetry`: +``R^T h`` from :meth:`~torchref.symmetry.symmetry.Symmetry.expand_reciprocal` and the +phases from :meth:`~torchref.symmetry.symmetry.Symmetry.phase_factors`. Reach this +class through +:meth:`~torchref.symmetry.symmetry.Symmetry.reciprocal_extractor`, which caches it. """ -from typing import Optional, TYPE_CHECKING +from typing import TYPE_CHECKING, Optional -import numpy as np import torch -from torchref.config import canonical_device, get_float_dtype +from torchref.config import canonical_device from torchref.utils.autograd_ops import gather_with_index_add - -from .grid_operations import extract_structure_factor_from_grid +from torchref.utils.device_mixin import DeviceMixin if TYPE_CHECKING: - from torchref.symmetry.spacegroup import SpaceGroup - - -def compute_symmetry_equivalent_hkls( - hkl: torch.Tensor, - rotation_matrices: torch.Tensor, -) -> torch.Tensor: - """ - Compute symmetry-equivalent HKLs for each operation. - - Row-vector convention: h' = h @ R, equivalently Rᵀ·h for column h. The - matrices are used as given -- do **not** pre-transpose them. - - Parameters - ---------- - hkl : torch.Tensor, shape (N, 3) - Miller indices. - rotation_matrices : torch.Tensor, shape (n_ops, 3, 3) - Real-space rotation matrices, applied directly as ``h @ R`` (no - transpose). - - Returns - ------- - torch.Tensor, shape (n_ops, N, 3) - Equivalent HKLs for each symmetry operation. - """ - device = hkl.device - dtype = get_float_dtype() - - hkl_float = hkl.to(dtype=dtype, device=device) # (N, 3) - rot_matrices = rotation_matrices.to(dtype=dtype, device=device) # (n_ops, 3, 3) - - n_ops = rot_matrices.shape[0] - - hkl_expanded = hkl_float.unsqueeze(0).expand(n_ops, -1, -1) # (n_ops, N, 3) - - equiv_hkl = torch.bmm(hkl_expanded, rot_matrices) - - # Exact for valid crystallographic ops; round only mops up float error. - equiv_hkl = torch.round(equiv_hkl).to(torch.int64) - - return equiv_hkl - - -def compute_translation_phases( - hkl: torch.Tensor, - translations: torch.Tensor, -) -> torch.Tensor: - """ - Compute the translation phase shifts exp(2πi h·t) for each operation. - - Parameters - ---------- - hkl : torch.Tensor, shape (N, 3) - Miller indices. - translations : torch.Tensor, shape (n_ops, 3) - Translation vectors in fractional coordinates. - - Returns - ------- - torch.Tensor, shape (n_ops, N) - Complex phase factors exp(2*pi*i * h.t). - """ - device = hkl.device - dtype = get_float_dtype() - - hkl_float = hkl.to(dtype=dtype, device=device) # (N, 3) - translations = translations.to(dtype=dtype, device=device) # (n_ops, 3) - - h_dot_t = torch.matmul(hkl_float, translations.T).T # (n_ops, N) + from torchref.symmetry.symmetry import Symmetry - # Leave ``phase`` at the configured float dtype; casting it to float32 would - # force complex64 output even under a float64 configuration. - phase = 2.0 * np.pi * h_dot_t - phase_factor = torch.exp(1j * phase) - return phase_factor # (n_ops, N) complex - - -def extract_structure_factors_with_symmetry( - reciprocal_grid: torch.Tensor, - hkl: torch.Tensor, - rotation_matrices: torch.Tensor, - translations: torch.Tensor, +def _equiv_hkls_to_flat_indices( + equiv_hkls: torch.Tensor, Nx: int, Ny: int, Nz: int ) -> torch.Tensor: - """ - Extract structure factors with symmetry applied in reciprocal space. - - Sums F over the symmetry-equivalent positions with their translation phases, - replacing the symmetrize-then-extract MapSymmetry route. For repeated calls - with the same hkl and symmetry, use :class:`ReciprocalSymmetryExtractor`. + """Flatten symmetry-equivalent Miller indices into linear grid indices. Parameters ---------- - reciprocal_grid : torch.Tensor, shape (Nx, Ny, Nz) - Complex reciprocal space grid from FFT of the **P1** density map. - hkl : torch.Tensor, shape (N, 3) - Target Miller indices. - rotation_matrices : torch.Tensor, shape (n_ops, 3, 3) - Real-space rotation matrices from symmetry operations. - translations : torch.Tensor, shape (n_ops, 3) - Translation vectors from symmetry operations. + equiv_hkls : torch.Tensor + Equivalent indices, shape ``(n_ops, N, 3)``. + Nx, Ny, Nz : int + Reciprocal grid dimensions. Returns ------- - torch.Tensor, shape (N,) - Complex structure factors with symmetry applied. + torch.Tensor + Flat indices, shape ``(n_ops * N,)``, dtype ``int64``, wrapped modulo the grid. """ - device = reciprocal_grid.device - Nx, Ny, Nz = reciprocal_grid.shape - - # Move everything to the same device - hkl = hkl.to(device=device) - rotation_matrices = rotation_matrices.to(device=device) - translations = translations.to(device=device) - - n_ops = rotation_matrices.shape[0] - N = hkl.shape[0] - - equiv_hkls = compute_symmetry_equivalent_hkls(hkl, rotation_matrices) - - # One flat gather for all symops. gather_with_index_add keeps the backward a - # single index_add_ instead of a radix-sort + dedup scatter. - flat_indices = _equiv_hkls_to_flat_indices(equiv_hkls, Nx, Ny, Nz) - f_all = gather_with_index_add( - reciprocal_grid.reshape(-1), flat_indices, - ) # (n_ops * N,) - f_p1 = f_all.view(n_ops, N) - - phases = compute_translation_phases(hkl, translations) - - f_sym = (f_p1 * phases).sum(dim=0) - - return f_sym - - -def _equiv_hkls_to_flat_indices( - equiv_hkls: torch.Tensor, Nx: int, Ny: int, Nz: int, -) -> torch.Tensor: - """Convert (n_ops, N, 3) equiv HKLs to flat linear grid indices.""" - all_hkl = equiv_hkls.reshape(-1, 3) # (n_ops*N, 3) + all_hkl = equiv_hkls.reshape(-1, 3) hi = torch.remainder(all_hkl[:, 0], Nx) ki = torch.remainder(all_hkl[:, 1], Ny) li = torch.remainder(all_hkl[:, 2], Nz) return (hi * (Ny * Nz) + ki * Nz + li).to(torch.int64) -from torchref.utils.device_mixin import DeviceMixin - - class ReciprocalSymmetryExtractor(DeviceMixin): - """ - Class-based interface for reciprocal space symmetry extraction. + """Precomputed symmetrized structure-factor extraction at fixed hkl and grid. - For repeated structure-factor evaluation at fixed hkl and symmetry, as in - refinement: the equivalent HKLs, phase factors and flat grid indices are - precomputed here, so each call is one gather, multiply and sum. The - precomputation binds ``grid_shape`` -- a differently shaped grid needs a new - extractor. + For repeated evaluation during refinement: the equivalent indices, phases and flat + gather indices are computed once here, so each call is one gather, multiply and + sum. The precomputation binds both ``hkl`` and ``grid_shape`` -- either changing + needs a new extractor, which is what + :meth:`~torchref.symmetry.symmetry.Symmetry.reciprocal_extractor` tracks. Parameters ---------- - hkl : torch.Tensor, shape (N, 3) - Target Miller indices. - symmetry : SpaceGroup - SpaceGroup object containing rotation matrices and translations. + hkl : torch.Tensor + Target Miller indices, shape ``(N, 3)``. + symmetry : Symmetry + The group supplying the operations. grid_shape : tuple of int - Reciprocal grid dimensions (Nx, Ny, Nz). + Reciprocal grid dimensions ``(Nx, Ny, Nz)``. device : torch.device, optional - Device for computation. - - Examples - -------- - >>> extractor = ReciprocalSymmetryExtractor(hkl, symmetry, grid_shape=(209, 86, 67)) - >>> f_calc = extractor.extract_from_grid(reciprocal_grid) + Device for the precomputed tensors. Defaults to ``hkl``'s. """ def __init__( self, hkl: torch.Tensor, - symmetry: "SpaceGroup", + symmetry: "Symmetry", grid_shape: tuple, device: Optional[torch.device] = None, ): - # ``is not None``, not ``or``: ``device=0`` means cuda:0/mps:0 and is - # falsy, so ``or`` silently discarded it. ``hkl`` is a bare tensor, so - # this reads its device rather than going through ``resolve_device``. + # ``is not None``, not ``or``: ``device=0`` means cuda:0/mps:0 and is falsy, so + # ``or`` silently discarded it. self.device = canonical_device( device if device is not None else hkl.device ) @@ -219,67 +93,56 @@ def __init__( self.N = len(hkl) self.grid_shape = grid_shape - # Precompute equivalent HKLs - self.equiv_hkls = compute_symmetry_equivalent_hkls( - self.hkl, - symmetry.matrices.to(device=self.device), - ) # (n_ops, N, 3) + self.equiv_hkls = symmetry.expand_reciprocal(self.hkl).to(device=self.device) + self.phases = symmetry.phase_factors(self.hkl).to(device=self.device) - # Precompute phase factors - self.phases = compute_translation_phases( - self.hkl, - symmetry.translations.to(device=self.device), - ) # (n_ops, N) complex - - # Precompute flat linear indices for single-gather extraction Nx, Ny, Nz = grid_shape self._flat_indices = _equiv_hkls_to_flat_indices( - self.equiv_hkls, Nx, Ny, Nz, - ) # (n_ops * N,) int64 + self.equiv_hkls, Nx, Ny, Nz + ) def __call__(self, density_map: torch.Tensor) -> torch.Tensor: """Alias for :meth:`extract`.""" return self.extract(density_map) def extract(self, density_map: torch.Tensor) -> torch.Tensor: - """ - Transform a P1 density map and extract symmetrized structure factors. + """Transform a P1 density map and extract symmetrized structure factors. Parameters ---------- - density_map : torch.Tensor, shape (Nx, Ny, Nz) - **P1** electron density map -- passing a symmetrized map double-counts. + density_map : torch.Tensor + **P1** electron density, shape ``(Nx, Ny, Nz)``. A symmetrized map + double-counts. Returns ------- - torch.Tensor, shape (N,) - Complex structure factors with symmetry applied. + torch.Tensor + Complex structure factors, shape ``(N,)``. """ from torchref.base.fourier.fft import ifft - reciprocal_grid = ifft(density_map) - return self.extract_from_grid(reciprocal_grid) + return self.extract_from_grid(ifft(density_map)) def extract_from_grid(self, reciprocal_grid: torch.Tensor) -> torch.Tensor: - """ - Extract structure factors from an already-transformed P1 grid. + """Extract structure factors from an already-transformed P1 grid. Parameters ---------- - reciprocal_grid : torch.Tensor, shape (Nx, Ny, Nz) - Complex reciprocal space grid from FFT of the **P1** map; its shape - must match the ``grid_shape`` this extractor was built for. + reciprocal_grid : torch.Tensor + Complex grid from the FFT of the **P1** map, shape ``(Nx, Ny, Nz)``; + its shape must match the ``grid_shape`` this extractor was built for. Returns ------- - torch.Tensor, shape (N,) - Complex structure factors with symmetry applied. + torch.Tensor + Complex structure factors, shape ``(N,)``. """ # gather_with_index_add keeps the backward a single ``index_add_`` # (atomic scatter, no radix sort + dedup). f_all = gather_with_index_add( - reciprocal_grid.reshape(-1), self._flat_indices, + reciprocal_grid.reshape(-1), self._flat_indices ) # (n_ops * N,) - f_sym = (f_all.view(self.n_ops, self.N) * self.phases).sum(dim=0) - return f_sym + return (f_all.view(self.n_ops, self.N) * self.phases).sum(dim=0) + +__all__ = ["ReciprocalSymmetryExtractor"] diff --git a/torchref/cli/mtz2map.py b/torchref/cli/mtz2map.py index b02f459f..3518495e 100644 --- a/torchref/cli/mtz2map.py +++ b/torchref/cli/mtz2map.py @@ -183,10 +183,11 @@ def main(): phi_t = torch.tensor(phases_deg, dtype=torch.float32, device=device) * (np.pi / 180.0) # --- Expand to P1 --- - from torchref.symmetry.reciprocal_symmetry import expand_hkl + from torchref.symmetry import Cell, SpaceGroup - hkl_p1, orig_idx, phase_shifts = expand_hkl( - hkl_t, spacegroup, include_friedel=False, remove_absences=True + sg = SpaceGroup(spacegroup) + hkl_p1, orig_idx, phase_shifts = sg.expand_hkl( + hkl_t, include_friedel=False, remove_absences=True ) amp_p1 = amp_t[orig_idx] @@ -199,13 +200,11 @@ def main(): coefficients = amp_p1 * torch.exp(1j * phi_p1) # --- Grid size --- - from torchref.symmetry.grid_utils import calculate_optimal_grid_size - if args.gridsize is not None: gridsize = tuple(args.gridsize) else: max_res = float(d_spacings.min()) - gridsize = calculate_optimal_grid_size(cell, max_res, spacegroup) + gridsize = sg.optimal_grid_size(Cell(cell), max_res) if args.verbose >= 1: print(f" Grid size: {gridsize[0]} x {gridsize[1]} x {gridsize[2]}") diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index 307d0c69..e6c50f87 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -227,7 +227,6 @@ def setup_ded_context( import gemmi from torchref import DatasetCollection - from torchref.symmetry.grid_utils import calculate_optimal_grid_size from torchref.symmetry.reciprocal_symmetry import expand_hkl from torchref.config import normalize_device @@ -308,9 +307,10 @@ def setup_ded_context( d_spacing = torch.tensor(d_spacings, dtype=torch.float32, device=device) # P1 expansion and grid - gridsize = calculate_optimal_grid_size(cell_t, dmin, sg_name) - hkl_p1, orig_idx, phase_shifts = expand_hkl( - hkl, sg_name, include_friedel=False, remove_absences=True + sg = SpaceGroup(sg_name, device=device) + gridsize = sg.optimal_grid_size(Cell(cell_t, device=device), dmin) + hkl_p1, orig_idx, phase_shifts = sg.expand_hkl( + hkl, include_friedel=False, remove_absences=True ) w_dfo_p1 = w_dfo[orig_idx] weights_p1 = weights[orig_idx] diff --git a/torchref/experimental/targets/realspace.py b/torchref/experimental/targets/realspace.py index 0e2a85f6..56fbf8b7 100644 --- a/torchref/experimental/targets/realspace.py +++ b/torchref/experimental/targets/realspace.py @@ -19,8 +19,7 @@ import torch from torchref.base.reciprocal.grid_operations import place_on_grid -from torchref.symmetry.grid_utils import calculate_optimal_grid_size -from torchref.symmetry.reciprocal_symmetry import expand_hkl +from torchref.symmetry import SpaceGroup from torchref.utils.stats import ( VERBOSITY_DEBUG, VERBOSITY_DETAILED, @@ -136,9 +135,9 @@ def _ensure_p1_expansion(self): """Compute and cache the ASU → P1 expansion mapping.""" if self._hkl_p1 is not None: return - hkl_p1, indices, phase_shifts = expand_hkl( + sg = self._data.spacegroup or SpaceGroup("P1") + hkl_p1, indices, phase_shifts = sg.expand_hkl( self._data.hkl, - self._data.spacegroup or "P1", include_friedel=True, remove_absences=True, device=self._data.hkl.device, @@ -624,9 +623,8 @@ def _ensure_p1_expansion(self): return spacegroup = self._data_light.spacegroup - hkl_p1, indices, phase_shifts = expand_hkl( + hkl_p1, indices, phase_shifts = spacegroup.expand_hkl( self._hkl, - spacegroup, include_friedel=True, remove_absences=True, device=self._hkl.device, diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index fd62273b..5a335989 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -508,13 +508,13 @@ def _canonicalize_in_place(self) -> None: """Remap HKL to canonical CCP4 ASU form and reorder all data in-place.""" from dataclasses import fields as dc_fields - from torchref.symmetry.reciprocal_symmetry import canonicalize_hkl - if self.hkl is None or self.spacegroup is None: return - canonical_hkl, phase_shifts, friedel_flags, sort_indices = canonicalize_hkl( - self.hkl, self.spacegroup, include_friedel=True, device=self.device + canonical_hkl, phase_shifts, friedel_flags, sort_indices = ( + self.spacegroup.canonicalize_hkl( + self.hkl, include_friedel=True, device=self.device + ) ) n_refl = len(self.hkl) @@ -2486,11 +2486,11 @@ def flag_wilson_outliers( observations depart from Wilson statistics for reasons that have nothing to do with being outliers. """ - from torchref.base.french_wilson import ( + from torchref.base.french_wilson import intensities_from_amplitudes + from torchref.base.wilson_outliers import wilson_outlier_mask + from torchref.refinement.model_error_estimation.sigma_a import ( epsilon_from_hkl, - intensities_from_amplitudes, ) - from torchref.base.wilson_outliers import wilson_outlier_mask if self.F is None or self.F_sigma is None or self.resolution is None: return @@ -2982,11 +2982,9 @@ def centric(self): # Cached on the _centric_flags dataclass field, so it survives # serialization. if not hasattr(self, "_centric_flags") or self._centric_flags is None: - from torchref.base.french_wilson import is_centric_from_hkl - - sg = self.spacegroup if self.spacegroup else "P1" + sg = self.spacegroup or SpaceGroup("P1", device=self.hkl.device) - self._centric_flags = is_centric_from_hkl(self.hkl, sg) + self._centric_flags = sg.is_centric(self.hkl) return self._centric_flags @@ -3186,8 +3184,6 @@ def fill(self, d_min: Optional[float] = None) -> "ReflectionData": Missing reflections have F/I/phase/fom = 0.0, F_sigma/I_sigma = 1.0 and ``masks['missing'] = True``. """ - from torchref.symmetry.reciprocal_symmetry import complete_hkl - if self.hkl is None: raise ValueError("ReflectionData has no Miller indices loaded") if self.cell is None: @@ -3200,8 +3196,9 @@ def fill(self, d_min: Optional[float] = None) -> "ReflectionData": d_min = self.resolution.min().item() # Get complete HKL set with index mapping - filled_hkl, indices, missing = complete_hkl( - self.hkl, self.cell.data, self.spacegroup or "P1", d_min, device=self.device + sg = self.spacegroup or SpaceGroup("P1", device=self.device) + filled_hkl, indices, missing = sg.complete_hkl( + self.hkl, self.cell.data, d_min, device=self.device ) # Use remap to create the new dataset @@ -3241,15 +3238,13 @@ def expand_to_p1( shift, ``resolution`` is recomputed and ``bin_indices`` is cleared. ``source``/``last_op`` record the provenance. """ - from torchref.symmetry.reciprocal_symmetry import expand_hkl - if self.hkl is None: raise ValueError("ReflectionData has no Miller indices loaded") # Get expanded HKL set with index mapping and phase shifts - hkl_p1, indices, phase_shifts = expand_hkl( + sg = self.spacegroup or SpaceGroup("P1", device=self.device) + hkl_p1, indices, phase_shifts = sg.expand_hkl( self.hkl, - self.spacegroup or "P1", include_friedel=include_friedel, remove_absences=remove_absences, device=self.device, @@ -3302,16 +3297,15 @@ def reduce_to_spacegroup( ``validation_flags`` set if any equivalent is set. Any other per-reflection field takes its first valid equivalent. """ - from torchref.symmetry.reciprocal_symmetry import reduce_hkl from torchref.symmetry.spacegroup import SpaceGroup if self.hkl is None: raise ValueError("ReflectionData has no Miller indices loaded") # Get reduction mapping - hkl_asu, reduction_indices, phase_shifts = reduce_hkl( - self.hkl, spacegroup, include_friedel=include_friedel, device=self.device - ) + hkl_asu, reduction_indices, phase_shifts = SpaceGroup( + spacegroup, device=self.device + ).reduce_hkl(self.hkl, include_friedel=include_friedel, device=self.device) n_asu = len(hkl_asu) n_equiv = reduction_indices.shape[1] @@ -3523,13 +3517,12 @@ def canonicalize(self, include_friedel: bool = True) -> "ReflectionData": ReflectionData New object with canonicalized, sorted Miller indices. """ - from torchref.symmetry.reciprocal_symmetry import canonicalize_hkl - if self.hkl is None: raise ValueError("ReflectionData has no Miller indices loaded") - canonical_hkl, phase_shifts, friedel_flags, sort_indices = canonicalize_hkl( - self.hkl, self.spacegroup or "P1", include_friedel, device=self.device + sg = self.spacegroup or SpaceGroup("P1", device=self.device) + canonical_hkl, phase_shifts, friedel_flags, sort_indices = sg.canonicalize_hkl( + self.hkl, include_friedel, device=self.device ) # Reorder all fields using __select__ diff --git a/torchref/maps/difference_map.py b/torchref/maps/difference_map.py index 6051c5ed..58b44242 100644 --- a/torchref/maps/difference_map.py +++ b/torchref/maps/difference_map.py @@ -14,7 +14,7 @@ from torchref.base.reciprocal.grid_operations import place_on_grid from torchref.io.datasets.collection import DatasetCollection from torchref.maps.map import Map -from torchref.symmetry.reciprocal_symmetry import expand_hkl +from torchref.symmetry import SpaceGroup from torchref.utils.device_resolution import resolve_device @@ -112,9 +112,9 @@ def calculate(self) -> torch.Tensor: # Expand to P1 without Friedel mates (expand_to_p1() would reset # scaling, so expand manually via expand_hkl) - sg = self.data_reference.spacegroup or "P1" - hkl_p1, orig_idx, _ = expand_hkl( - hkl_asu, sg, + sg = self.data_reference.spacegroup or SpaceGroup("P1", device=hkl_asu.device) + hkl_p1, orig_idx, _ = sg.expand_hkl( + hkl_asu, include_friedel=False, remove_absences=True, device=hkl_asu.device, ) diff --git a/torchref/maps/map.py b/torchref/maps/map.py index 62570125..85c75b06 100644 --- a/torchref/maps/map.py +++ b/torchref/maps/map.py @@ -23,7 +23,6 @@ from torchref.base.reciprocal.grid_operations import place_on_grid from torchref.io.cif import write_map -from torchref.symmetry.grid_utils import calculate_optimal_grid_size from torchref.utils.device_mixin import DeviceMixin from torchref.utils.device_resolution import resolve_device @@ -104,10 +103,8 @@ def map_data(self) -> Optional[torch.Tensor]: def _determine_gridsize(self) -> Tuple[int, int, int]: """Determine optimal grid size from cell, resolution, and spacegroup.""" - cell_params = self.data.cell.data max_res = float(self.data.resolution.min()) - spacegroup = self.data.spacegroup.name - return calculate_optimal_grid_size(cell_params, max_res, spacegroup) + return self.data.spacegroup.optimal_grid_size(self.data.cell, max_res) def _compute_map_coefficients( self, fobs: torch.Tensor, fcalc: torch.Tensor diff --git a/torchref/model/mixed_model.py b/torchref/model/mixed_model.py index c23d3ca5..9c01ebd0 100644 --- a/torchref/model/mixed_model.py +++ b/torchref/model/mixed_model.py @@ -193,11 +193,6 @@ def gridsize(self) -> Optional[torch.Tensor]: """Grid dimensions (nx, ny, nz) from first model.""" return self.models[0].gridsize - @property - def map_symmetry(self): - """Map symmetry operator from first model.""" - return self.models[0].map_symmetry - @property def inv_fractional_matrix(self) -> torch.Tensor: """Inverse fractionalization (orthogonalization) matrix.""" diff --git a/torchref/model/model_collection.py b/torchref/model/model_collection.py index d56d6c9a..dc082742 100644 --- a/torchref/model/model_collection.py +++ b/torchref/model/model_collection.py @@ -132,10 +132,6 @@ def fft(self): def gridsize(self): return self._base_models[0].gridsize - @property - def map_symmetry(self): - return self._base_models[0].map_symmetry - @property def inv_fractional_matrix(self): return self.cell.inv_fractional_matrix.to(dtype=self.dtype_float) diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index 85ec6e13..60434fa0 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -17,7 +17,6 @@ from torchref.model.model import Model from torchref.model.sf_fft import SfFFT from torchref.symmetry import SpaceGroup -from torchref.symmetry.map_symmetry import MapSymmetry from torchref.utils.caching import CachedForwardMixin @@ -60,8 +59,6 @@ class ModelFT(CachedForwardMixin, Model): Most recently computed electron density map. parametrization : dict ITC92 parametrization dictionary {element: (A, B)}. - map_symmetry : MapSymmetry - Symmetry operator for map calculations. """ def __init__( @@ -316,16 +313,6 @@ def voxel_size(self, value): """Set voxel size (for backward compatibility).""" self._fft.voxel_size = value - @property - def map_symmetry(self) -> Optional[MapSymmetry]: - """Symmetry operator for map calculations.""" - return self._fft.map_symmetry - - @map_symmetry.setter - def map_symmetry(self, value): - """Set map symmetry (for backward compatibility).""" - self._fft.map_symmetry = value - def get_iso(self): """ Get isotropic atoms with their ITC92 parameters. @@ -832,7 +819,7 @@ def copy(self, detach: bool = True) -> "ModelFT": Create a deep copy of the ModelFT. Creates a complete independent copy including all Model base class data, - FFT submodule state (gridsize, real_space_grid, voxel_size, map_symmetry), + FFT submodule state (gridsize, real_space_grid, voxel_size), ITC92 parametrization, and scalar attributes. Cache is reset to empty. diff --git a/torchref/model/sf_ds.py b/torchref/model/sf_ds.py index 003bb878..0f028cc7 100644 --- a/torchref/model/sf_ds.py +++ b/torchref/model/sf_ds.py @@ -232,35 +232,6 @@ def _compute_scattering_factors( return f - def _get_spacegroup_callable(self): - """Symmetry-application callable for the direct-summation kernels. - - They expect ``(3, N)`` fractional coordinates in and ``(3, N, n_ops)`` - out; identity-only when no space group is set. - """ - if self._spacegroup is None: - # P1 symmetry - identity operation only - def p1_symmetry(coords_3N): - # coords_3N: (3, N) -> (3, N, 1) - return coords_3N.unsqueeze(2) - - return p1_symmetry - - def apply_symmetry(coords_3N): - # coords_3N: (3, N) -> coords_N3: (N, 3) - coords_N3 = coords_3N.T - - # Apply symmetry: (N, 3) -> (N, 3, ops) - transformed = self._spacegroup.apply(coords_N3) - - # Reorder to (3, N, ops) - # (N, 3, ops) -> (3, N, ops) - result = transformed.permute(1, 0, 2) - - return result - - return apply_symmetry - def _cartesian_to_fractional(self, xyz_cartesian: torch.Tensor) -> torch.Tensor: """``(N, 3)`` Cartesian coordinates to fractional; needs a cell whose dtype matches ``self.dtype_float``. @@ -365,20 +336,9 @@ def compute_structure_factors( return sf_p1, None # Apply late symmetry: F_sym(h) = Σ_ops exp(2πi h.t) * F_P1(R^T @ h) - from torchref.base.reciprocal import ( - compute_symmetry_equivalent_hkls, - compute_translation_phases, - ) - n_ops = self._spacegroup.n_ops - rotation_matrices = self._spacegroup.matrices - translations = self._spacegroup.translations - - # Compute equivalent HKLs: (n_ops, N, 3) - equiv_hkls = compute_symmetry_equivalent_hkls(hkl, rotation_matrices) - - # Compute translation phase shifts: (n_ops, N) - phases = compute_translation_phases(hkl, translations) + equiv_hkls = self._spacegroup.expand_reciprocal(hkl) # (n_ops, N, 3) + phases = self._spacegroup.phase_factors(hkl) # (n_ops, N) # Compute F_P1 at each equivalent HKL and combine sf_total = torch.zeros( diff --git a/torchref/model/sf_fft.py b/torchref/model/sf_fft.py index 12c727da..33bfd4e0 100644 --- a/torchref/model/sf_fft.py +++ b/torchref/model/sf_fft.py @@ -5,7 +5,7 @@ ``ModelFT``'s submodule. ``FFT`` is a deprecated alias for :class:`SfFFT`. """ -from typing import Optional, Tuple, Union +from typing import Optional, Tuple import torch import torch.nn as nn @@ -14,9 +14,7 @@ from torchref.base.reciprocal import extract_structure_factor_from_grid from torchref.config import dtypes, get_default_device from torchref.symmetry import Cell, SpaceGroup -from torchref.symmetry.map_symmetry import MapSymmetry from torchref.symmetry.spacegroup import SpaceGroupLike -from torchref.utils.caching import ParameterFingerprint from torchref.utils.device_mixin import DeviceMovementMixin from torchref.utils.device_resolution import resolve_device @@ -54,8 +52,6 @@ class SfFFT(DeviceMovementMixin, nn.Module): gridsize, real_space_grid, voxel_size : torch.Tensor or None Grid dimensions ``(nx, ny, nz)``, coordinate grid ``(nx, ny, nz, 3)`` and voxel dimensions -- all ``None`` until :meth:`setup_grid` runs. - map_symmetry : MapSymmetry or None - Symmetry operator for map calculations. """ def __init__( @@ -121,19 +117,9 @@ def __init__( self.register_buffer("real_space_grid", None) self.register_buffer("voxel_size", None) - # Map symmetry operator (set during setup_grid) - self.map_symmetry: Optional[MapSymmetry] = None - # Late symmetry compatibility flag (set during setup_grid) self._late_symmetry_compatible: Optional[bool] = None - # Cached reciprocal symmetry extractor (precomputed flat indices). - # Keyed on a (data_ptr, _version, numel) fingerprint of the HKL tensor - # rather than id(hkl): a garbage-collected tensor reallocated at the - # same address cannot silently alias a stale extractor. - self._sym_extractor = None - self._sym_extractor_hkl_fp: Optional[ParameterFingerprint] = None - # ========================================================================= # Cell and SpaceGroup properties # ========================================================================= @@ -242,8 +228,6 @@ def compute_optimal_gridsize(self, max_res: Optional[float] = None) -> tuple: resolution = max_res if max_res is not None else self.max_res - from torchref.symmetry.spacegroup import suggest_grid_size - # Use Cell's method for base grid size calculation gridsize_initial = self._cell.compute_grid_size(resolution) @@ -251,8 +235,8 @@ def compute_optimal_gridsize(self, max_res: Optional[float] = None) -> tuple: print(f"Initial grid size from cell: {gridsize_initial}") # Optimize for symmetry and FFT-friendliness - gridsize_optimized = suggest_grid_size( - gridsize_initial, self._spacegroup, make_fft_friendly=True + gridsize_optimized = self._spacegroup.suggest_grid_size( + gridsize_initial, make_fft_friendly=True ) if self.verbose > 1 and gridsize_optimized != gridsize_initial: print( @@ -342,19 +326,14 @@ def setup_grid( # Compute voxel size self.voxel_size = self.real_space_grid[2, 2, 2] - self.real_space_grid[1, 1, 1] - # Initialize map symmetry operator if space group is set + # Every symmetry-equivalent HKL lands on an integer grid point exactly when + # the grid admits direct indexing, which the space group answers without + # building an operator. if self._spacegroup is not None: - self.map_symmetry = MapSymmetry( - space_group=self._spacegroup, - map_shape=self.real_space_grid.shape[:-1], - cell_params=self._cell.data, - verbose=self.verbose, - device=self.device, + self._late_symmetry_compatible = self._spacegroup.can_index_directly( + self.real_space_grid.shape[:-1] ) - # Check late symmetry compatibility - self._late_symmetry_compatible = self._check_late_symmetry_compatible() - if self.use_late_symmetry and self._late_symmetry_compatible: if self.verbose > 0: print( @@ -367,29 +346,16 @@ def setup_grid( "(falling back to early symmetry)" ) else: - self.map_symmetry = None self._late_symmetry_compatible = False - # Invalidate cached symmetry extractor (grid shape changed) - self._sym_extractor = None - self._sym_extractor_hkl_fp = None + # The grid shape changed, so the space group's cached operators are stale. + if self._spacegroup is not None: + self._spacegroup.reset_cache() if self.verbose > 2: print(f"Grid shape: {self.real_space_grid.shape[:-1]}") print(f"Voxel size: {self.voxel_size}") - def _check_late_symmetry_compatible(self) -> bool: - """True when every symmetry-equivalent HKL lands on an integer grid - point, which the MapSymmetry factory signals by returning a - ``MapSymmetryDirect`` (direct indexing, no interpolation). - """ - if self.map_symmetry is None: - return False - - from torchref.symmetry.map_symmetry import MapSymmetryDirect - - return isinstance(self.map_symmetry, MapSymmetryDirect) - # ========================================================================= # Density Map Building Methods # ========================================================================= @@ -457,8 +423,8 @@ def build_density_map( ) # Apply symmetry if requested - if apply_symmetry and self.map_symmetry is not None: - density_map = self.map_symmetry(density_map) + if apply_symmetry and self._spacegroup is not None: + density_map = self._spacegroup.symmetrize_map(density_map) return density_map @@ -497,21 +463,9 @@ def map_to_structure_factors( # Use late symmetry if enabled, compatible, and requested if apply_symmetry: # Lazily build / reuse cached extractor (precomputed flat indices) - if self._sym_extractor is None or ( - self._sym_extractor_hkl_fp is None - or not self._sym_extractor_hkl_fp.matches([hkl]) - ): - from torchref.base.reciprocal import ReciprocalSymmetryExtractor - - grid_shape = tuple(int(x) for x in self.gridsize) - self._sym_extractor = ReciprocalSymmetryExtractor( - hkl, - self.spacegroup, - grid_shape, - device=reciprocal_space_grid.device, - ) - self._sym_extractor_hkl_fp = ParameterFingerprint([hkl]) - return self._sym_extractor.extract_from_grid(reciprocal_space_grid) + grid_shape = tuple(int(x) for x in self.gridsize) + extractor = self._spacegroup.reciprocal_extractor(hkl, grid_shape) + return extractor.extract_from_grid(reciprocal_space_grid) else: return extract_structure_factor_from_grid(reciprocal_space_grid, hkl) @@ -594,9 +548,9 @@ def compute_structure_factors( # ========================================================================= def reset_cache(self) -> None: - """Drop the cached symmetry extractor; recomputed on next use.""" - self._sym_extractor = None - self._sym_extractor_hkl_fp = None + """Drop the space group's cached operators; recomputed on next use.""" + if self._spacegroup is not None: + self._spacegroup.reset_cache() def copy(self) -> "SfFFT": """Create a deep copy of this SfFFT module. diff --git a/torchref/refinement/model_error_estimation/sigma_a.py b/torchref/refinement/model_error_estimation/sigma_a.py index 79f33405..6c96182c 100644 --- a/torchref/refinement/model_error_estimation/sigma_a.py +++ b/torchref/refinement/model_error_estimation/sigma_a.py @@ -22,9 +22,34 @@ import torch -from torchref.base.french_wilson import epsilon_from_hkl # noqa: F401 (re-export) from torchref.config import get_float_dtype + +def epsilon_from_hkl(hkl: torch.Tensor, spacegroup) -> torch.Tensor: + """Per-reflection epsilon, tolerating a missing space group. + + Thin adapter over :meth:`~torchref.symmetry.symmetry.Symmetry.epsilon`, which owns + the multiplicity count. It exists because reflection data may carry no space group + at all, and every consumer here would otherwise repeat the same guard. + + Parameters + ---------- + hkl : torch.Tensor + Miller indices, shape ``(N, 3)``. + spacegroup : Symmetry or None + The group. ``None`` means no symmetry information, which yields ones -- the + same answer P1 gives. + + Returns + ------- + torch.Tensor + Multiplicities, shape ``(N,)``, at the configured float dtype, on ``hkl``'s + device. + """ + if spacegroup is None or not hasattr(spacegroup, "epsilon"): + return torch.ones(hkl.shape[0], device=hkl.device, dtype=get_float_dtype()) + return spacegroup.epsilon(hkl) + # --- sigma_A estimator constants ------------------------------------------------- #: Upper bound on the per-shell ``sigma_A``, i.e. the floor on the model-error variance at #: ``(1 - SIGMA_A_MAX**2) * Sigma_N``. diff --git a/torchref/restraints/restraints.py b/torchref/restraints/restraints.py index 7c76df5f..aef98fd3 100644 --- a/torchref/restraints/restraints.py +++ b/torchref/restraints/restraints.py @@ -1407,7 +1407,7 @@ def vdw_radii_cpu(): if self._spacegroup is not None: from torchref.symmetry.spacegroup import SpaceGroup sg_cpu = SpaceGroup(self._spacegroup, device=cpu, - dtype=self._spacegroup._dtype) + dtype=self._spacegroup.dtype) else: sg_cpu = None diff --git a/torchref/symmetry/__init__.py b/torchref/symmetry/__init__.py index 09014038..b8e3b862 100644 --- a/torchref/symmetry/__init__.py +++ b/torchref/symmetry/__init__.py @@ -1,83 +1,35 @@ -"""Crystallographic symmetry: space groups, unit cells, map and HKL symmetry. +"""Crystallographic symmetry: symmetry groups, space groups and unit cells. -:class:`SpaceGroup` (``nn.Module`` holding the operations as buffers) is the entry -point, with ``Symmetry`` a bare alias for it. :func:`MapSymmetry` handles real-space -density and :func:`ReciprocalSymmetry` structure-factor grids; all three accept a -space group as a string, an int 1-230, or a gemmi object. :class:`Cell` is separate --- it wraps the six cell parameters, not a space group. +:class:`Symmetry` holds a group as rotation matrices and fractional translations and +owns every verb derivable from the operations alone -- expansion of positions and +Miller indices, translation phases, the reflection predicates, symmetry-compatible grid +sizes, and map symmetrization. Nothing in it is crystallographic, so a group built from +a raw operation list serves non-crystallographic symmetry too. -The grid utilities re-exported here come from ``grid_utils``, which delegates to -``spacegroup``. ``spacegroup`` also defines its own same-named copies, which are -the source of truth and are *not* re-exported. +:class:`SpaceGroup` specialises it with the crystallographic identity (Hermann-Mauguin +naming, number, point group, crystal system) and the CCP4 asymmetric-unit verbs +(``expand_hkl``, ``reduce_hkl``, ``complete_hkl``, ``canonicalize_hkl``). It accepts a +name, a number 1-230, a ``gemmi.SpaceGroup``, another instance, or None for P1. + +:class:`Cell` is separate: it wraps the six cell parameters, not a symmetry group. + +Map and reciprocal-grid operators are reached through :class:`Symmetry` +(:meth:`~Symmetry.symmetrize_map`, :meth:`~Symmetry.reciprocal_extractor`), which owns +their caching -- the operator classes themselves are private. """ -from .cell import Cell, CellTensor -from .grid_utils import ( - calculate_optimal_grid_size, - check_grid_compatibility, - find_fft_friendly_size, - get_symmetry_grid_requirements, - is_fft_friendly, - recommend_grid_size, -) -from .map_symmetry import MapSymmetry, MapSymmetryDirect -from .reciprocal_symmetry import ( - ReciprocalSymmetry, - ReciprocalSymmetryGrid, - canonicalize_hkl, - complete_hkl, - expand_hkl, - expand_reciprocal_grid, - expand_reflections, - reduce_hkl, -) -from .spacegroup import ( - SpaceGroup, - SpaceGroupLike, - get_crystal_system, - get_operations_as_tensors, - get_point_group, - get_symmetry_operations, - is_centrosymmetric, - is_same_spacegroup, - n_operations, - spacegroup_to_str, -) -from .symmetry import Symmetry +from .cell import Cell +from .spacegroup import SpaceGroup, SpaceGroupLike +from .symmetry import Symmetry, find_fft_friendly_size, is_fft_friendly __all__ = [ # Unit cell "Cell", - # Space group utilities + # Symmetry groups + "Symmetry", "SpaceGroup", "SpaceGroupLike", - "spacegroup_to_str", - "get_symmetry_operations", - "get_operations_as_tensors", - "is_same_spacegroup", - "get_point_group", - "get_crystal_system", - "is_centrosymmetric", - "n_operations", - # Base symmetry - "Symmetry", - # Real space map symmetry - "MapSymmetry", - "MapSymmetryDirect", - # Reciprocal space symmetry - "ReciprocalSymmetry", - "ReciprocalSymmetryGrid", - "expand_hkl", - "complete_hkl", - "reduce_hkl", - "canonicalize_hkl", - "expand_reflections", - "expand_reciprocal_grid", - # Grid utilities - "get_symmetry_grid_requirements", - "check_grid_compatibility", - "recommend_grid_size", - "find_fft_friendly_size", + # Grid sizing helpers (group-independent) "is_fft_friendly", - "calculate_optimal_grid_size", + "find_fft_friendly_size", ] diff --git a/torchref/symmetry/cell.py b/torchref/symmetry/cell.py index cc0eebf0..a8743cb3 100644 --- a/torchref/symmetry/cell.py +++ b/torchref/symmetry/cell.py @@ -52,7 +52,6 @@ def __init__( *, dtype: torch.dtype = None, device: torch.device | str = None, - requires_grad: bool = False, ) -> None: """ Create a new Cell. @@ -66,8 +65,6 @@ def __init__( Desired data type. Defaults to the configured ``dtypes.float``. device : torch.device or str, optional Desired device. Defaults to the configured ``device.current``. - requires_grad : bool, optional - Whether to track gradients. Defaults to False. Raises ------ @@ -93,9 +90,6 @@ def __init__( # Ensure 1D shape tensor = tensor.reshape(6) - if requires_grad: - tensor = tensor.requires_grad_(True) - object.__setattr__(self, "_data", tensor) object.__setattr__(self, "_cache", {}) @@ -111,21 +105,6 @@ def reset_cache(self) -> None: """Clear cached derived quantities (fractional matrix, volume, etc.).""" object.__setattr__(self, "_cache", {}) - def detach(self) -> "Cell": - """ - Return a new Cell with detached tensor (no gradient tracking). - - Returns - ------- - Cell - New Cell with detached data. - """ - new_data = self._data.detach() - new_cell = Cell.__new__(Cell) - object.__setattr__(new_cell, "_data", new_data) - object.__setattr__(new_cell, "_cache", {}) - return new_cell - def clone(self) -> "Cell": """ Return a new Cell with cloned tensor data. @@ -160,11 +139,6 @@ def data(self) -> torch.Tensor: """Return the underlying tensor (for buffer registration).""" return self._data - @property - def requires_grad(self) -> bool: - """Return whether gradients are tracked.""" - return self._data.requires_grad - # ========================================================================= # Convenience properties for cell parameters # ========================================================================= @@ -409,7 +383,3 @@ def __getitem__(self, idx: int) -> torch.Tensor: def __len__(self) -> int: """Return 6 (number of cell parameters).""" return 6 - - -# Keep CellTensor as an alias for backward compatibility -CellTensor = Cell diff --git a/torchref/symmetry/grid_utils.py b/torchref/symmetry/grid_utils.py deleted file mode 100644 index 902cf887..00000000 --- a/torchref/symmetry/grid_utils.py +++ /dev/null @@ -1,134 +0,0 @@ -"""FFT- and symmetry-compatible grid sizes. - -Interpolation-free symmetry expansion needs grid dimensions divisible by what the -screw axes demand, and radix-2,3,5 FFTs want factors of 2, 3, 5 only. - -These are thin wrappers over ``spacegroup``, which holds the canonical -implementations -- including its own ``is_fft_friendly`` / -``find_fft_friendly_size`` pair. Prefer ``spacegroup`` for new code. -""" - -import numpy as np -import torch - -from torchref.config import NYQUIST_OVERSAMPLING - - -def get_symmetry_grid_requirements(space_group: str) -> dict: - """Per-axis divisibility ``{'nx_mod', 'ny_mod', 'nz_mod'}`` for ``space_group``. - - Wrapper over :func:`~torchref.symmetry.spacegroup.get_grid_requirements`. - """ - # Import here to avoid circular imports - from torchref.symmetry.spacegroup import get_grid_requirements - - return get_grid_requirements(space_group) - - -def find_fft_friendly_size(n: int, divisibility: int = 1) -> int: - """Smallest size >= ``n`` factoring into 2, 3, 5 and divisible by ``divisibility``. - - Parameters - ---------- - n : int - Minimum grid size. - divisibility : int, default 1 - Required divisibility (e.g. 2 for a screw axis). - - Returns - ------- - int - Optimal grid size. - """ - candidate = n - - if candidate % divisibility != 0: - candidate = ((candidate // divisibility) + 1) * divisibility - - while not is_fft_friendly(candidate): - candidate += divisibility - - return candidate - - -def is_fft_friendly(n: int) -> bool: - """ - Check if a number has only factors of 2, 3, and 5. - - These are optimal for radix-2,3,5 FFT algorithms. - """ - if n <= 0: - return False - - # Remove all factors of 2, 3, 5 - while n % 2 == 0: - n //= 2 - while n % 3 == 0: - n //= 3 - while n % 5 == 0: - n //= 5 - - # If we're left with 1, the number is FFT-friendly - return n == 1 - - -def calculate_optimal_grid_size(cell_params, max_res: float, space_group: str) -> tuple: - """ - Optimal grid for a unit cell and space group. - - Satisfies Shannon-Nyquist sampling at - :data:`torchref.config.NYQUIST_OVERSAMPLING`, the screw-axis divisibility, and - FFT-friendliness (factors of 2, 3, 5 only). - - Parameters - ---------- - cell_params : array-like, shape (6,) - Unit cell [a, b, c, alpha, beta, gamma]. - max_res : float - Maximum resolution in Angstroms. - space_group : str - Space group symbol. - - Returns - ------- - tuple - Optimal grid dimensions (nx, ny, nz). - """ - # Import here to avoid circular imports - from torchref.symmetry.spacegroup import suggest_grid_size - - if isinstance(cell_params, torch.Tensor): - cell_params = cell_params.cpu().numpy() - - a, b, c = cell_params[:3] - - # Shannon-Nyquist: sample at NYQUIST_OVERSAMPLING × the maximum frequency - nx_min = int(np.floor(a / max_res * NYQUIST_OVERSAMPLING)) - ny_min = int(np.floor(b / max_res * NYQUIST_OVERSAMPLING)) - nz_min = int(np.floor(c / max_res * NYQUIST_OVERSAMPLING)) - - # Use spacegroup module to suggest optimal size - return suggest_grid_size((nx_min, ny_min, nz_min), space_group, make_fft_friendly=True) - - -def check_grid_compatibility(grid_shape: tuple, space_group: str) -> dict: - """Check ``(nx, ny, nz)`` against the space group symmetry and the FFT. - - Wrapper over - :func:`~torchref.symmetry.spacegroup.check_grid_compatibility`, which - documents the report dict. - """ - # Import here to avoid circular imports - from torchref.symmetry.spacegroup import ( - check_grid_compatibility as sg_check_grid_compatibility, - ) - - return sg_check_grid_compatibility(grid_shape, space_group) - - -def recommend_grid_size(current_shape: tuple, space_group: str) -> tuple: - """Smallest symmetry- and FFT-compatible grid at or above ``current_shape``.""" - # Import here to avoid circular imports - from torchref.symmetry.spacegroup import suggest_grid_size - - return suggest_grid_size(current_shape, space_group, make_fft_friendly=True) diff --git a/torchref/symmetry/map_symmetry.py b/torchref/symmetry/map_symmetry.py index a9d1852d..18985a15 100644 --- a/torchref/symmetry/map_symmetry.py +++ b/torchref/symmetry/map_symmetry.py @@ -1,167 +1,154 @@ -"""Map-level symmetry operations for electron density maps. - -Applying symmetry to the map is far cheaper than generating symmetry mates per -atom. :func:`MapSymmetry` is a factory, not a class: it returns -:class:`MapSymmetryDirect` (exact integer indexing) when the grid allows, else the -interpolating implementation from ``map_symmetry_interpolation``. Space groups -accept strings, ints 1-230 or ``gemmi.SpaceGroup``. +"""Real-space map symmetrization, selected by grid compatibility. + +Two operators apply a :class:`~torchref.symmetry.symmetry.Symmetry` to a density map: +:class:`_MapSymmetryDirect` indexes symmetry mates at exact integers, and +:class:`~torchref.symmetry.map_symmetry_interpolation._MapSymmetryInterpolation` +falls back to ``grid_sample`` when the grid does not admit that. +:func:`build_map_operator` picks between them. + +Reach these through :meth:`~torchref.symmetry.symmetry.Symmetry.symmetrize_map` rather +than directly: it owns the caching, and it is the reason the choice of operator does +not leak into calling code. A grid that forces interpolation costs accuracy silently, +so ask :meth:`~torchref.symmetry.symmetry.Symmetry.can_index_directly` before +committing to a grid, and +:meth:`~torchref.symmetry.symmetry.Symmetry.suggest_grid_size` to fix one. + +Neither operator needs the unit cell: symmetry acts on fractional coordinates, so the +cell metric never enters. """ +from __future__ import annotations + import torch -import torch.nn as nn -from torchref.config import get_float_dtype, normalize_device -from torchref.symmetry.spacegroup import SpaceGroup, SpaceGroupLike from torchref.utils.device_mixin import DeviceMixin -def MapSymmetry( - space_group: SpaceGroupLike, - map_shape, - cell_params, - dtype_float=None, - verbose=1, - device=None, -): - """ - Build the appropriate MapSymmetry implementation for ``map_shape``. - - Returns :class:`MapSymmetryDirect` when the grid permits exact integer - indexing, otherwise the interpolating fallback -- so the return *type* - depends on the grid, and a mis-sized grid costs accuracy silently (at - ``verbose > 0`` a compatible grid is suggested). +def build_map_operator(symmetry, map_shape: tuple): + """Build the operator suited to ``map_shape``. Parameters ---------- - space_group : str, int, or gemmi.SpaceGroup - Space group specification (e.g., 'P21', 4, gemmi.SpaceGroup('P 21')). + symmetry : Symmetry + The group to apply. map_shape : tuple of int - Shape of the density map (nx, ny, nz). - cell_params : torch.Tensor, shape (6,) - Unit cell parameters [a, b, c, alpha, beta, gamma] in Å and degrees. - dtype_float : torch.dtype, optional - Floating point precision to use. Defaults to the configured - ``dtypes.float`` (``get_float_dtype()``, float32 in production). - verbose : int, default 1 - Verbosity level (0=silent, 1=info, 2=debug). - device : torch.device, default: configured device.current - Device to use for computation. + Density map dimensions ``(nx, ny, nz)``. Returns ------- - MapSymmetryDirect or MapSymmetryInterpolation - The appropriate implementation based on grid compatibility. + _MapSymmetryDirect or _MapSymmetryInterpolation + Direct integer indexing when the grid permits it, otherwise the interpolating + fallback. """ - if dtype_float is None: - dtype_float = get_float_dtype() - # ``cell_params`` is documented as a tensor, so follow it when no device is - # given rather than jumping to the global default and leaving the caller's - # cell behind. - if device is None and isinstance(cell_params, torch.Tensor): - device = cell_params.device - device = normalize_device(device) - symmetry = SpaceGroup(space_group, dtype=dtype_float, device=device) - compat = symmetry.check_grid_compatibility(map_shape) - - if compat["can_use_direct_indexing"]: - if verbose > 0: - print( - f"MapSymmetry: Using direct indexing (no interpolation) for {space_group}" - ) - return MapSymmetryDirect( - space_group, map_shape, cell_params, dtype_float, verbose, device - ) - else: - if verbose > 0: - print("MapSymmetry: Grid not compatible with direct indexing") - print(f" Using interpolation-based fallback for {space_group}") - if compat["issues"]: - for issue in compat["issues"]: - print(f" - {issue}") - suggested = symmetry.suggest_grid_size(map_shape, make_fft_friendly=True) - print(f" Suggested grid for direct indexing: {suggested}") - - from torchref.symmetry.map_symmetry_interpolation import ( - MapSymmetry as MapSymmetryInterpolation, - ) + if symmetry.can_index_directly(map_shape): + return _MapSymmetryDirect(symmetry, map_shape) - return MapSymmetryInterpolation( - space_group, map_shape, cell_params, dtype_float, verbose, device - ) + # Imported here, not at module scope: the interpolation module imports this one for + # the shared operator contract. + from torchref.symmetry.map_symmetry_interpolation import ( + _MapSymmetryInterpolation, + ) + return _MapSymmetryInterpolation(symmetry, map_shape) -class MapSymmetryDirect(DeviceMixin, nn.Module): - """ - Fast direct-indexing implementation of crystallographic symmetry operations. - Computes symmetry mates one operation at a time (streaming) so that - memory usage is O(grid) regardless of the number of symmetry operations, - rather than O(n_ops * grid) for storing precomputed index grids. +def _combine(mates: torch.Tensor, combine: str) -> torch.Tensor: + """Reduce stacked symmetry mates. - NOTE: Do not instantiate this class directly. Use the MapSymmetry() factory - function instead, which will automatically select the appropriate implementation. + Parameters + ---------- + mates : torch.Tensor + Stacked mates, shape ``(n_ops, nx, ny, nz)``. + combine : {'sum', 'max'} + ``'sum'`` for electron density, ``'max'`` for masks and boolean data. + + Returns + ------- + torch.Tensor + Shape ``(nx, ny, nz)``. + + Raises + ------ + ValueError + For an unknown mode. """ + if combine == "sum": + return mates.sum(dim=0) + if combine == "max": + return mates.max(dim=0)[0] + raise ValueError(f"Unknown combine mode: {combine}. Use 'sum' or 'max'.") - def __init__( - self, - space_group, - map_shape, - cell_params, - dtype_float=None, - verbose=1, - device=None, - ): - super().__init__() - if dtype_float is None: - dtype_float = get_float_dtype() - if device is None and isinstance(cell_params, torch.Tensor): - device = cell_params.device - self.dtype_float = dtype_float - self.space_group = space_group - self.map_shape = tuple(map_shape) - self.verbose = verbose - self.device = normalize_device(device) - # Coerce ``cell_params`` onto this module's device: it is a plain attribute, - # so ``DeviceMixin`` cannot see it and would never repair a mismatch. - if isinstance(cell_params, torch.Tensor): - cell_params = cell_params.to(device=self.device, dtype=self.dtype_float) - else: - cell_params = torch.as_tensor( - cell_params, device=self.device, dtype=self.dtype_float - ) - self.cell_params = cell_params - self.symmetry = SpaceGroup( - space_group, dtype=self.dtype_float, device=self.device - ) - self.n_ops = self.symmetry.matrices.shape[0] - self.can_use_direct_indexing = True +class _MapSymmetryDirect(DeviceMixin): + """Symmetrize maps by exact integer indexing, one operation at a time. - if self.verbose > 0: - print(f"MapSymmetryDirect initialized for {space_group}") - print(f" Number of symmetry operations: {self.n_ops}") - print(f" Map shape: {self.map_shape}") + Valid only on a grid whose dimensions satisfy the group's divisibility, so every + symmetry mate falls on a grid point. :func:`build_map_operator` enforces that. + + Parameters + ---------- + symmetry : Symmetry + The group to apply. + map_shape : tuple of int + Density map dimensions ``(nx, ny, nz)``. - # ------------------------------------------------------------------ - # Core: compute index grid for a single symmetry operation - # ------------------------------------------------------------------ + Notes + ----- + Holds no precomputed grids. Index grids are recomputed per operation so peak memory + stays at one index grid plus two density maps, rather than scaling with the number + of operations. + """ - def _compute_index_grid(self, op_index: int) -> torch.Tensor: - """Integer index grid (nx, ny, nz, 3) int64 for one op; not cached.""" + def __init__(self, symmetry, map_shape: tuple): + self.symmetry = symmetry + self.map_shape = tuple(int(n) for n in map_shape) + + @property + def n_ops(self) -> int: + """Number of symmetry operations.""" + return self.symmetry.n_ops + + @property + def device(self) -> torch.device: + """Device the operations live on.""" + return self.symmetry.device + + def _index_grid(self, op_index: int) -> torch.Tensor: + """Integer index grid for one operation. + + Parameters + ---------- + op_index : int + Operation index. + + Returns + ------- + torch.Tensor + Shape ``(nx, ny, nz, 3)``, dtype ``int64``. + + Notes + ----- + Deliberately does not use the batched + :meth:`~torchref.symmetry.symmetry.Symmetry.expand_positions`: that would + transform every operation at once, which is exactly the O(n_ops * grid) memory + this operator exists to avoid. + """ nx, ny, nz = self.map_shape - device = self.symmetry.matrices.device + symmetry = self.symmetry + dtype = symmetry.dtype + device = symmetry.device - fx = torch.arange(nx, dtype=self.dtype_float, device=device) / nx - fy = torch.arange(ny, dtype=self.dtype_float, device=device) / ny - fz = torch.arange(nz, dtype=self.dtype_float, device=device) / nz + fx = torch.arange(nx, dtype=dtype, device=device) / nx + fy = torch.arange(ny, dtype=dtype, device=device) / ny + fz = torch.arange(nz, dtype=dtype, device=device) / nz gx, gy, gz = torch.meshgrid(fx, fy, fz, indexing="ij") grid_flat = torch.stack([gx, gy, gz], dim=-1).reshape(-1, 3) - transformed = torch.matmul(self.symmetry.matrices[op_index], grid_flat.T).T - transformed = transformed + self.symmetry.translations[op_index] + transformed = torch.matmul(symmetry.matrices[op_index], grid_flat.T).T + transformed = transformed + symmetry.translations[op_index] transformed = transformed - torch.floor(transformed) - shape_t = torch.tensor([nx, ny, nz], dtype=self.dtype_float, device=device) + shape_t = torch.tensor([nx, ny, nz], dtype=dtype, device=device) indices = torch.round(transformed * shape_t).to(torch.int64) indices[:, 0] %= nx indices[:, 1] %= ny @@ -169,70 +156,95 @@ def _compute_index_grid(self, op_index: int) -> torch.Tensor: return indices.reshape(nx, ny, nz, 3) - # ------------------------------------------------------------------ - # Public API - # ------------------------------------------------------------------ - - def get_symmetry_mate(self, density_map, operation_index): - """Apply a single symmetry operation via direct indexing.""" - if operation_index < 0 or operation_index >= self.n_ops: + def _check_shape(self, density_map: torch.Tensor) -> None: + """Reject a map whose shape this operator was not built for.""" + if tuple(density_map.shape) != self.map_shape: raise ValueError( - f"Operation index {operation_index} out of range [0, {self.n_ops-1}]" + f"Map shape {tuple(density_map.shape)} does not match the operator's " + f"{self.map_shape}" ) - if density_map.shape != self.map_shape: - raise ValueError( - f"Map shape {density_map.shape} doesn't match expected {self.map_shape}" - ) - ig = self._compute_index_grid(operation_index) - return density_map[ig[..., 0], ig[..., 1], ig[..., 2]] - def forward(self, density_map, apply_symmetry=True, combine_mode="sum"): - """Apply symmetry operations to density map. + def mate(self, density_map: torch.Tensor, op_index: int) -> torch.Tensor: + """One symmetry mate of ``density_map``. - Computes one symmetry mate at a time and accumulates into the - result, so peak memory is only 1 index grid + 2 density maps - regardless of the number of symmetry operations. + Parameters + ---------- + density_map : torch.Tensor + Density, shape ``(nx, ny, nz)``. + op_index : int + Operation index in ``[0, n_ops)``. + + Returns + ------- + torch.Tensor + Shape ``(nx, ny, nz)``. """ - if not apply_symmetry or self.n_ops == 1: - return density_map - - ig = self._compute_index_grid(0) - if combine_mode == "sum": - result = density_map[ig[..., 0], ig[..., 1], ig[..., 2]] - for i in range(1, self.n_ops): - ig = self._compute_index_grid(i) - result = result + density_map[ig[..., 0], ig[..., 1], ig[..., 2]] - elif combine_mode == "max": - result = density_map[ig[..., 0], ig[..., 1], ig[..., 2]] - for i in range(1, self.n_ops): - ig = self._compute_index_grid(i) - result = torch.max( - result, density_map[ig[..., 0], ig[..., 1], ig[..., 2]] - ) - else: + if op_index < 0 or op_index >= self.n_ops: raise ValueError( - f"Unknown combine_mode: {combine_mode}. Use 'sum' or 'max'." + f"Operation index {op_index} out of range [0, {self.n_ops - 1}]" ) + self._check_shape(density_map) + ig = self._index_grid(op_index) + return density_map[ig[..., 0], ig[..., 1], ig[..., 2]] - return result + def all_mates(self, density_map: torch.Tensor) -> torch.Tensor: + """Every symmetry mate, stacked. + + Parameters + ---------- + density_map : torch.Tensor + Density, shape ``(nx, ny, nz)``. - def __call__(self, density_map, apply_symmetry=True, combine_mode="sum"): - """Make the class callable like a PyTorch module.""" - return self.forward( - density_map, apply_symmetry=apply_symmetry, combine_mode=combine_mode + Returns + ------- + torch.Tensor + Shape ``(n_ops, nx, ny, nz)``. + """ + self._check_shape(density_map) + return torch.stack( + [self.mate(density_map, i) for i in range(self.n_ops)], dim=0 ) - def get_symmetry_info(self): - """Get information about symmetry operations.""" - return { - "space_group": self.space_group, - "n_operations": self.n_ops, - "matrices": self.symmetry.matrices, - "translations": self.symmetry.translations, - } + def symmetrize( + self, density_map: torch.Tensor, combine: str = "sum" + ) -> torch.Tensor: + """Apply every operation and reduce the mates. + + Parameters + ---------- + density_map : torch.Tensor + Density, shape ``(nx, ny, nz)``. + combine : {'sum', 'max'}, default 'sum' + Reduction across mates. + + Returns + ------- + torch.Tensor + Shape ``(nx, ny, nz)``. + + Notes + ----- + Accumulates one mate at a time instead of stacking, keeping peak memory + independent of the operation count. + """ + self._check_shape(density_map) + result = self.mate(density_map, 0) + for i in range(1, self.n_ops): + mate = self.mate(density_map, i) + if combine == "sum": + result = result + mate + elif combine == "max": + result = torch.max(result, mate) + else: + raise ValueError( + f"Unknown combine mode: {combine}. Use 'sum' or 'max'." + ) + return result - def __repr__(self): + def __repr__(self) -> str: return ( - f"MapSymmetryDirect(space_group='{self.space_group}', " - f"n_ops={self.n_ops}, map_shape={self.map_shape})" + f"_MapSymmetryDirect(n_ops={self.n_ops}, map_shape={self.map_shape})" ) + + +__all__ = ["build_map_operator"] diff --git a/torchref/symmetry/map_symmetry_interpolation.py b/torchref/symmetry/map_symmetry_interpolation.py index 0e3ae55a..16fb7171 100644 --- a/torchref/symmetry/map_symmetry_interpolation.py +++ b/torchref/symmetry/map_symmetry_interpolation.py @@ -1,299 +1,173 @@ -""" -Map-level symmetry operations for electron density maps. +"""Map symmetrization by trilinear interpolation, for grids that need it. + +The fallback behind :func:`~torchref.symmetry.map_symmetry.build_map_operator` when a +grid does not satisfy the group's divisibility, so symmetry mates land between grid +points. Interpolating costs accuracy that exact indexing does not, which is why +:meth:`~torchref.symmetry.symmetry.Symmetry.suggest_grid_size` exists -- prefer fixing +the grid over landing here. -This module provides efficient symmetry operations applied directly to density maps, -which is much faster than applying symmetry to individual atoms. +Reach this through :meth:`~torchref.symmetry.symmetry.Symmetry.symmetrize_map`. """ -import numpy as np +from __future__ import annotations + import torch -import torch.nn as nn import torch.nn.functional as F -from torchref.config import get_float_dtype, normalize_device -from torchref.symmetry.spacegroup import SpaceGroup +from torchref.symmetry.map_symmetry import _combine from torchref.utils.device_mixin import DeviceMixin -class MapSymmetry(DeviceMixin, nn.Module): - """ - Applies crystallographic symmetry operations to electron density maps. +class _MapSymmetryInterpolation(DeviceMixin): + """Symmetrize maps by resampling with ``grid_sample``. - Takes an asymmetric-unit density map, applies each operation in fractional - coordinates, interpolates via ``grid_sample`` and sums the mates -- cheaper - than generating symmetry mates per atom and recalculating density. The - per-operation sampling grids are precomputed in ``__init__`` for the fixed - ``map_shape``. - - Attributes + Parameters ---------- - space_group : str - Space group name. - map_shape : tuple of int - Shape of the density map (nx, ny, nz). - cell_params : numpy.ndarray - Unit cell parameters. symmetry : Symmetry - Symmetry operations handler. - n_ops : int - Number of symmetry operations. - - Examples - -------- - :: - - map_sym = MapSymmetry(space_group='P21', map_shape=(64, 64, 64), cell_params=cell) - asymmetric_map = model.build_density_map() - symmetric_map = map_sym(asymmetric_map) + The group to apply. + map_shape : tuple of int + Density map dimensions ``(nx, ny, nz)``. + + Notes + ----- + Precomputes one sampling grid per operation, shape + ``(n_ops, nx, ny, nz, 3)`` -- hundreds of megabytes at production grid sizes. That + is why :class:`~torchref.symmetry.symmetry.Symmetry` memoizes only the most recent + shape and drops the cache on any device move. """ - def __init__( - self, - space_group, - map_shape, - cell_params, - dtype_float=None, - verbose=1, - device=None, - ): - """ - Initialize map symmetry operator. + def __init__(self, symmetry, map_shape: tuple): + self.symmetry = symmetry + self.map_shape = tuple(int(n) for n in map_shape) + self.sampling_grids = self._build_sampling_grids() - Parameters - ---------- - space_group : str - Space group name (e.g., 'P1', 'P21', 'P-1', etc.). - map_shape : tuple of int - Shape of the density map (nx, ny, nz). - cell_params : array-like, shape (6,) - Unit cell parameters [a, b, c, alpha, beta, gamma] in Å and degrees. - dtype_float : torch.dtype, optional - Floating point precision to use. Defaults to the configured - ``dtypes.float`` (``get_float_dtype()``, float32 in production). - verbose : int, default 1 - Verbosity level. - device : torch.device, default: configured device.current - Device to use for computation. - """ - super().__init__() - if dtype_float is None: - dtype_float = get_float_dtype() - device = normalize_device(device) - self.dtype_float = dtype_float - self.space_group = space_group - self.map_shape = tuple(map_shape) - self.cell_params = np.array(cell_params) - self.verbose = verbose - self.device = device - self.symmetry = SpaceGroup( - space_group, dtype=self.dtype_float, device=self.device - ) - self.n_ops = self.symmetry.matrices.shape[0] - if self.verbose > 0: - print(f"MapSymmetry initialized for {space_group}") - print(f" Number of symmetry operations: {self.n_ops}") - print(f" Map shape: {self.map_shape}") + @property + def n_ops(self) -> int: + """Number of symmetry operations.""" + return self.symmetry.n_ops - self._setup_fractional_grid() - self._setup_symmetry_grids() + @property + def device(self) -> torch.device: + """Device the sampling grids live on.""" + return self.sampling_grids.device - def _setup_fractional_grid(self): - """Fractional grid with voxels at edges i/N (CCTBX/gemmi convention).""" - nx, ny, nz = self.map_shape - - fx = torch.arange(nx, dtype=self.dtype_float, device=self.device) / nx - fy = torch.arange(ny, dtype=self.dtype_float, device=self.device) / ny - fz = torch.arange(nz, dtype=self.dtype_float, device=self.device) / nz - - # indexing='ij' so the result is (nx, ny, nz, 3) with the last dim [fx, fy, fz]. - grid_fx, grid_fy, grid_fz = torch.meshgrid(fx, fy, fz, indexing="ij") - grid_frac = torch.stack([grid_fx, grid_fy, grid_fz], dim=-1) - - self.register_buffer("grid_frac", grid_frac) + def _build_sampling_grids(self) -> torch.Tensor: + """Precompute per-operation ``grid_sample`` coordinates in ``[-1, 1]``. - def _setup_symmetry_grids(self): - """Precompute per-operation ``grid_sample`` coordinates in [-1, 1].""" + Returns + ------- + torch.Tensor + Shape ``(n_ops, nx, ny, nz, 3)``. + """ nx, ny, nz = self.map_shape - - grid_flat = self.grid_frac.reshape(-1, 3) - - sampling_grids_list = [] - - for i in range(self.n_ops): - # R @ coords + t, on (3, nx*ny*nz) - transformed = torch.matmul(self.symmetry.matrices[i], grid_flat.T) - transformed = transformed.T # (nx*ny*nz, 3) - transformed = transformed + self.symmetry.translations[i] - - # Wrap to [0, 1) for periodic boundary conditions - transformed = transformed - torch.floor(transformed) - grid_shape_tensor = torch.tensor( - [nx, ny, nz], dtype=self.dtype_float, device=transformed.device - ) - # grid_coord = -1 + 2*N/(N-1) * frac, per dimension - sampling_coords = ( - -1.0 + 2.0 * grid_shape_tensor / (grid_shape_tensor - 1.0) * transformed + symmetry = self.symmetry + dtype = symmetry.dtype + device = symmetry.device + + # Voxels at fractional edges i/N, the CCTBX/gemmi convention. + fx = torch.arange(nx, dtype=dtype, device=device) / nx + fy = torch.arange(ny, dtype=dtype, device=device) / ny + fz = torch.arange(nz, dtype=dtype, device=device) / nz + gx, gy, gz = torch.meshgrid(fx, fy, fz, indexing="ij") + grid_flat = torch.stack([gx, gy, gz], dim=-1).reshape(-1, 3) + + transformed = symmetry.expand_positions(grid_flat) # (n_ops, N, 3) + # Wrap into [0, 1) for periodic boundaries. + transformed = transformed - torch.floor(transformed) + + shape_t = torch.tensor([nx, ny, nz], dtype=dtype, device=device) + # grid_coord = -1 + 2*N/(N-1) * frac, per dimension. + sampling = -1.0 + 2.0 * shape_t / (shape_t - 1.0) * transformed + sampling = sampling.reshape(self.n_ops, nx, ny, nz, 3) + + # grid_sample reads the last axis as [x, y, z] -> [W, H, D], i.e. the REVERSE + # of our [fx, fy, fz] -> [D, H, W]. Dropping this reorder still interpolates, + # silently against the wrong axes. + return sampling[..., [2, 1, 0]].contiguous() + + def _check_shape(self, density_map: torch.Tensor) -> None: + """Reject a map whose shape this operator was not built for.""" + if tuple(density_map.shape) != self.map_shape: + raise ValueError( + f"Map shape {tuple(density_map.shape)} does not match the operator's " + f"{self.map_shape}" ) - sampling_grid = sampling_coords.reshape(nx, ny, nz, 3) - - # grid_sample reads the last axis as [x, y, z] -> [W, H, D], i.e. the - # REVERSE of our [fx, fy, fz] -> [D, H, W]. Dropping this reorder still - # interpolates, silently against the wrong axes. - sampling_grid = sampling_grid[ - ..., [2, 1, 0] - ] # [fx, fy, fz] -> [fz, fy, fx] - - sampling_grids_list.append(sampling_grid) - - sampling_grids_stacked = torch.stack(sampling_grids_list, dim=0) - - self.register_buffer("sampling_grids", sampling_grids_stacked) - - def get_symmetry_mate(self, density_map, operation_index): - """ - Apply a single symmetry operation to get one symmetry mate. + def mate(self, density_map: torch.Tensor, op_index: int) -> torch.Tensor: + """One symmetry mate of ``density_map``. Parameters ---------- - density_map : torch.Tensor, shape (nx, ny, nz) - Electron density map (typically from asymmetric unit). - operation_index : int - Index of the symmetry operation to apply (0 to n_ops-1). + density_map : torch.Tensor + Density, shape ``(nx, ny, nz)``. + op_index : int + Operation index in ``[0, n_ops)``. Returns ------- - torch.Tensor, shape (nx, ny, nz) - Density map after applying the symmetry operation. + torch.Tensor + Shape ``(nx, ny, nz)``. """ - if operation_index < 0 or operation_index >= self.n_ops: - raise ValueError( - f"Operation index {operation_index} out of range [0, {self.n_ops-1}]" - ) - - # Ensure map is correct shape - if density_map.shape != self.map_shape: + if op_index < 0 or op_index >= self.n_ops: raise ValueError( - f"Map shape {density_map.shape} doesn't match expected {self.map_shape}" + f"Operation index {op_index} out of range [0, {self.n_ops - 1}]" ) - - # Prepare for grid_sample - map_5d = density_map.unsqueeze(0).unsqueeze(0) # (1, 1, nx, ny, nz) - - # Get sampling grid for this operation - sampling_grid = self.sampling_grids[operation_index] - sampling_grid_batch = sampling_grid.unsqueeze(0) - - # Interpolate map at transformed coordinates - # align_corners=True ensures that: - # -1 maps to index 0 (fractional coord 0) - # +1 maps to index N-1 (fractional coord (N-1)/N) - # This matches the grid-edge convention (voxels at i/N) - # padding_mode='border' handles periodic boundary conditions via the wrapping - # we did in _setup_symmetry_grids - transformed_map = F.grid_sample( - map_5d, - sampling_grid_batch, - mode="bilinear", # Trilinear interpolation for 3D - padding_mode="border", # Use border mode since we pre-wrapped coordinates - align_corners=True, # Critical: matches grid-edge convention + self._check_shape(density_map) + + # align_corners=True maps -1 to index 0 and +1 to index N-1, matching the + # grid-edge convention above; padding_mode='border' is safe only because + # _build_sampling_grids already wrapped the coordinates. + transformed = F.grid_sample( + density_map.unsqueeze(0).unsqueeze(0), + self.sampling_grids[op_index].unsqueeze(0), + mode="bilinear", + padding_mode="border", + align_corners=True, ) + return transformed.squeeze(0).squeeze(0) - # Remove batch and channel dimensions - transformed_map = transformed_map.squeeze(0).squeeze(0) - - return transformed_map - - def get_all_symmetry_mates(self, density_map): - """ - Get all symmetry mates as a list. + def all_mates(self, density_map: torch.Tensor) -> torch.Tensor: + """Every symmetry mate, stacked. Parameters ---------- - density_map : torch.Tensor, shape (nx, ny, nz) - Electron density map (typically from asymmetric unit). + density_map : torch.Tensor + Density, shape ``(nx, ny, nz)``. Returns ------- - list of torch.Tensor - List of symmetry-related maps, one for each operation. + torch.Tensor + Shape ``(n_ops, nx, ny, nz)``. """ - mates = [] - for i in range(self.n_ops): - mates.append(self.get_symmetry_mate(density_map, i)) - return mates + self._check_shape(density_map) + return torch.stack( + [self.mate(density_map, i) for i in range(self.n_ops)], dim=0 + ) - def forward(self, density_map, apply_symmetry=True, combine_mode="sum"): - """ - Apply symmetry operations to density map. + def symmetrize( + self, density_map: torch.Tensor, combine: str = "sum" + ) -> torch.Tensor: + """Apply every operation and reduce the mates. Parameters ---------- - density_map : torch.Tensor, shape (nx, ny, nz) - Electron density map (typically from asymmetric unit). - apply_symmetry : bool, default True - If True, apply all symmetry operations and combine them. - If False, return input map unchanged (useful for P1 or debugging). - combine_mode : str, default 'sum' - How to combine symmetry mates: - - - 'sum': Sum all symmetry mates (for electron density) - - 'max': Take maximum across symmetry mates (for masks/boolean data) + density_map : torch.Tensor + Density, shape ``(nx, ny, nz)``. + combine : {'sum', 'max'}, default 'sum' + Reduction across mates. Returns ------- - torch.Tensor, shape (nx, ny, nz) - Symmetry-expanded density map (combined symmetry mates). + torch.Tensor + Shape ``(nx, ny, nz)``. """ - if not apply_symmetry or self.n_ops == 1: - # No symmetry or P1 - return density_map - - # Get all symmetry mates - mates = self.get_all_symmetry_mates(density_map) - mates_stacked = torch.stack(mates, dim=0) - - # Combine according to mode - if combine_mode == "sum": - symmetric_map = mates_stacked.sum(dim=0) - elif combine_mode == "max": - symmetric_map = mates_stacked.max(dim=0)[0] # max returns (values, indices) - else: - raise ValueError( - f"Unknown combine_mode: {combine_mode}. Use 'sum' or 'max'." - ) + return _combine(self.all_mates(density_map), combine) - return symmetric_map - - def __call__(self, density_map, apply_symmetry=True, combine_mode="sum"): - """Make the class callable like a PyTorch module.""" - return self.forward( - density_map, apply_symmetry=apply_symmetry, combine_mode=combine_mode + def __repr__(self) -> str: + return ( + f"_MapSymmetryInterpolation(n_ops={self.n_ops}, " + f"map_shape={self.map_shape})" ) - def get_symmetry_info(self): - """ - Get information about symmetry operations. - Returns - ------- - dict - Dictionary with the following keys: - - - 'space_group' : str - - 'n_operations' : int - - 'matrices' : torch.Tensor, shape (n_ops, 3, 3) - - 'translations' : torch.Tensor, shape (n_ops, 3) - """ - return { - "space_group": self.space_group, - "n_operations": self.n_ops, - "matrices": self.symmetry.matrices, - "translations": self.symmetry.translations, - } - - def __repr__(self): - return ( - f"MapSymmetry(space_group='{self.space_group}', " - f"n_ops={self.n_ops}, map_shape={self.map_shape})" - ) +__all__ = ["_MapSymmetryInterpolation"] diff --git a/torchref/symmetry/reciprocal_symmetry.py b/torchref/symmetry/reciprocal_symmetry.py index 15ad80c2..66dfacff 100644 --- a/torchref/symmetry/reciprocal_symmetry.py +++ b/torchref/symmetry/reciprocal_symmetry.py @@ -1,630 +1,36 @@ -"""Reciprocal space symmetry operations for structure factor grids. - -The reciprocal-space counterpart to ``map_symmetry.py``: :func:`ReciprocalSymmetry` -(grid operator), :func:`expand_hkl` / :func:`expand_reflections` / -:func:`expand_reciprocal_grid` (ASU -> P1), :func:`reduce_hkl` (P1 -> ASU), -:func:`complete_hkl` (find reflections missing from a dataset, same space group) -and :func:`canonicalize_hkl` (CCP4 ASU representative). Space groups accept -strings, ints 1-230 or ``gemmi.SpaceGroup``. - -Miller indices transform as h' = h @ R = R^T @ h with R the *real-space* rotation; -``reciprocal_matrices`` already holds the transpose, so do not transpose again. -Translations become phase shifts of -2π h·t; the sign is load-bearing and wrong -signs are invisible in P21/P212121/C2 (see :func:`expand_hkl`). +"""Asymmetric-unit conventions for Miller indices. + +The algorithms behind :class:`~torchref.symmetry.spacegroup.SpaceGroup`'s HKL verbs: +``expand_hkl`` (ASU -> P1), ``reduce_hkl`` (P1 -> ASU), ``complete_hkl`` (reflections +missing from a dataset, same space group) and ``canonicalize_hkl`` (CCP4 ASU +representative). All private -- call them through the space group, which is the only +public entry point. + +What makes these crystallographic rather than general symmetry is the choice of +asymmetric unit: the CCP4 convention, read off gemmi's ``ReciprocalAsu`` and keyed by +Laue class. That is why they hang off +:class:`~torchref.symmetry.spacegroup.SpaceGroup` and not +:class:`~torchref.symmetry.symmetry.Symmetry`. + +Miller indices transform as ``h' = h @ R = R^T @ h`` with R the *real-space* rotation; +:attr:`~torchref.symmetry.symmetry.Symmetry.reciprocal` already holds the transpose. +Translations enter as phase shifts of ``-2 pi h.t``. That sign is load-bearing and a +wrong one is invisible in P21/P212121/C2 -- see :func:`_expand_hkl` and +``tests/unit/symmetry/test_phase_convention.py``. """ -from typing import TYPE_CHECKING, Optional, Tuple +from typing import Optional, Tuple import numpy as np import torch -import torch.nn as nn -from torchref.config import get_float_dtype, normalize_device -from torchref.symmetry.spacegroup import SpaceGroup, SpaceGroupLike -from torchref.utils.device_mixin import DeviceMixin +from torchref.config import get_float_dtype -if TYPE_CHECKING: - from torchref.io.datasets.reflection_data import ReflectionData -def ReciprocalSymmetry( - space_group: SpaceGroupLike, - grid_shape, - dtype_float=None, - verbose=1, - device=None, -): - """ - Factory function to create the appropriate ReciprocalSymmetry implementation. - - Parameters - ---------- - space_group : str, int, or gemmi.SpaceGroup - Space group specification (e.g., 'P21', 4, gemmi.SpaceGroup('P 21')). - grid_shape : tuple of int - Shape of the reciprocal space grid (nh, nk, nl). - The grid spans from -n//2 to n//2 for each dimension. - dtype_float : torch.dtype, default: configured dtypes.float - Floating point precision to use. - verbose : int, default 1 - Verbosity level (0=silent, 1=info, 2=debug). - device : torch.device, default: configured device.current - Device to use for computation. - - Returns - ------- - ReciprocalSymmetryGrid - Implementation for reciprocal space grid symmetry operations. - """ - return ReciprocalSymmetryGrid(space_group, grid_shape, dtype_float, verbose, device) - - -class ReciprocalSymmetryGrid(DeviceMixin, nn.Module): - """ - Reciprocal space symmetry operations for Miller index grids. - - Covers Miller-index transformation, systematic absences, centric reflections, - Friedel pairs and symmetry expansion/averaging of structure factors. The - per-operation index maps, phase shifts, absence mask and centric mask are all - precomputed in ``__init__`` for the fixed ``grid_shape``. - - Attributes - ---------- - space_group : str - Space group name. - grid_shape : tuple of int - Shape of the reciprocal space grid (nh, nk, nl). - symmetry : SpaceGroup - Base symmetry operations handler. - n_ops : int - Number of symmetry operations. - - Examples - -------- - :: - - recip_sym = ReciprocalSymmetry('P21', grid_shape=(64, 64, 64)) - F_expanded = recip_sym(F_asym) # Expand from asymmetric unit - F_avg = recip_sym.symmetry_average(F_full) # Average symmetry-related reflections - """ - - def __init__( - self, - space_group, - grid_shape, - dtype_float=None, - verbose=1, - device=None, - ): - """ - Initialize reciprocal space symmetry operator. - - Parameters - ---------- - space_group : str - Space group name. - grid_shape : tuple of int - Shape of the reciprocal space grid (nh, nk, nl). - dtype_float : torch.dtype, default: configured dtypes.float - Floating point precision. - verbose : int, default 1 - Verbosity level. - device : torch.device, default: configured device.current - Computation device. - """ - super().__init__() - if dtype_float is None: - dtype_float = get_float_dtype() - device = normalize_device(device) - self.dtype_float = dtype_float - self.space_group = space_group - self.grid_shape = tuple(grid_shape) - self.verbose = verbose - self.device = device - - self.symmetry = SpaceGroup(space_group, dtype=dtype_float, device=device) - self.n_ops = self.symmetry.matrices.shape[0] - - self._setup_reciprocal_matrices() - - if self.verbose > 0: - print(f"ReciprocalSymmetryGrid initialized for {space_group}") - print(f" Number of symmetry operations: {self.n_ops}") - print(f" Grid shape: {self.grid_shape}") - - self._setup_hkl_grid() - self._setup_symmetry_index_grids() - self._setup_systematic_absences() - self._setup_centric_reflections() - - if self.verbose > 0: - n_absent = self.systematic_absences.sum().item() - n_centric = self.centric_mask.sum().item() - n_total = np.prod(self.grid_shape) - print(f" Systematic absences: {n_absent} ({100*n_absent/n_total:.2f}%)") - print(f" Centric reflections: {n_centric} ({100*n_centric/n_total:.2f}%)") - - def _setup_reciprocal_matrices(self): - """Cache ``reciprocal_matrices`` = R^T, for h' = R^T @ h.""" - # Translations cause phase shifts, not index changes, so they are not folded in. - recip_matrices = self.symmetry.matrices.transpose(-2, -1).contiguous() - self.register_buffer("reciprocal_matrices", recip_matrices) - - def _setup_hkl_grid(self): - """Build ``hkl_grid`` in FFT index order (0...n//2, -n//2+1...-1).""" - nh, nk, nl = self.grid_shape - - h = torch.fft.fftfreq(nh, d=1.0) * nh # gives 0,1,2,...,n//2,-n//2+1,...,-1 - k = torch.fft.fftfreq(nk, d=1.0) * nk - l = torch.fft.fftfreq(nl, d=1.0) * nl - - h = h.to(dtype=torch.int64, device=self.device) - k = k.to(dtype=torch.int64, device=self.device) - l = l.to(dtype=torch.int64, device=self.device) - - grid_h, grid_k, grid_l = torch.meshgrid(h, k, l, indexing="ij") - hkl_grid = torch.stack([grid_h, grid_k, grid_l], dim=-1) - - self.register_buffer("hkl_grid", hkl_grid) - - # Float copy for the matmuls; the int copy stays authoritative for indexing. - hkl_grid_float = hkl_grid.to(dtype=self.dtype_float) - self.register_buffer("hkl_grid_float", hkl_grid_float) - - def _setup_symmetry_index_grids(self): - """Precompute ``index_grids`` (n_ops, nh, nk, nl, 3) and ``phase_shifts``. - - The grid path uses the +2π h·t convention, unlike :func:`expand_hkl`. - """ - nh, nk, nl = self.grid_shape - hkl_flat = self.hkl_grid_float.reshape(-1, 3) # (N, 3) - - index_grids_list = [] - phase_shift_grids_list = [] - - for i in range(self.n_ops): - transformed = torch.matmul(hkl_flat, self.reciprocal_matrices[i].T) - # Exact for valid symmetry ops; round only mops up float error. - transformed_int = torch.round(transformed).to(torch.int64) - - # F(h') = F(h) * exp(2πi h·t) - translation = self.symmetry.translations[i] - phase_shift = 2.0 * np.pi * torch.matmul(hkl_flat, translation) - phase_shift = phase_shift.reshape(nh, nk, nl) - phase_shift_grids_list.append(phase_shift) - - # Wrap with periodic boundary to get grid indices. - idx_h = transformed_int[:, 0] % nh - idx_k = transformed_int[:, 1] % nk - idx_l = transformed_int[:, 2] % nl - - index_grid = torch.stack([idx_h, idx_k, idx_l], dim=-1) - index_grid = index_grid.reshape(nh, nk, nl, 3) - index_grids_list.append(index_grid) - - index_grids = torch.stack(index_grids_list, dim=0) - self.register_buffer("index_grids", index_grids) - - phase_shifts = torch.stack(phase_shift_grids_list, dim=0) - self.register_buffer("phase_shifts", phase_shifts) - - def _setup_systematic_absences(self): - """Mask reflections with some op mapping h -> h at h·t not integral. - - Such a reflection is destroyed by interference from the translation. - """ - nh, nk, nl = self.grid_shape - absences = torch.zeros(self.grid_shape, dtype=torch.bool, device=self.device) - - hkl_flat = self.hkl_grid_float.reshape(-1, 3) - - for i in range(self.n_ops): - transformed = torch.matmul(hkl_flat, self.reciprocal_matrices[i].T) - transformed_int = torch.round(transformed).to(torch.int64) - - hkl_int = self.hkl_grid.reshape(-1, 3) - same_reflection = (transformed_int == hkl_int).all(dim=-1) - - translation = self.symmetry.translations[i] - h_dot_t = torch.matmul(hkl_flat, translation) - - phase_mod = torch.abs(h_dot_t - torch.round(h_dot_t)) - non_zero_phase = phase_mod > 1e-6 - - absent_mask = (same_reflection & non_zero_phase).reshape(self.grid_shape) - absences = absences | absent_mask - - self.register_buffer("systematic_absences", absences) - - def _setup_centric_reflections(self): - """Mask reflections with some op mapping h -> -h (phase restricted to 0/π).""" - nh, nk, nl = self.grid_shape - centric = torch.zeros(self.grid_shape, dtype=torch.bool, device=self.device) - - hkl_flat = self.hkl_grid_float.reshape(-1, 3) - - for i in range(self.n_ops): - transformed = torch.matmul(hkl_flat, self.reciprocal_matrices[i].T) - transformed_int = torch.round(transformed).to(torch.int64) - - hkl_int = self.hkl_grid.reshape(-1, 3) - maps_to_minus_h = (transformed_int == -hkl_int).all(dim=-1) - - centric_mask = maps_to_minus_h.reshape(self.grid_shape) - centric = centric | centric_mask - - self.register_buffer("centric_mask", centric) - - def apply_to_indices(self, hkl, operation_index=None): - """ - Apply symmetry operation(s) to Miller indices. - - Parameters - ---------- - hkl : torch.Tensor, shape (..., 3) - Miller indices (h, k, l). - operation_index : int, optional - If specified, apply only this operation. - If None, apply all operations. - - Returns - ------- - torch.Tensor - Transformed Miller indices. - If operation_index is None: shape (n_ops, ..., 3) - Otherwise: shape (..., 3) - """ - hkl = hkl.to(dtype=self.dtype_float, device=self.device) - original_shape = hkl.shape[:-1] - - if operation_index is not None: - # Apply single operation - R = self.reciprocal_matrices[operation_index] - transformed = torch.matmul(hkl, R.T) - return torch.round(transformed).to(torch.int64) - else: - # Apply all operations - hkl_flat = hkl.reshape(-1, 3) # (N, 3) - results = [] - for i in range(self.n_ops): - R = self.reciprocal_matrices[i] - transformed = torch.matmul(hkl_flat, R.T) - results.append(torch.round(transformed).to(torch.int64)) - - # Stack: (n_ops, N, 3) - stacked = torch.stack(results, dim=0) - return stacked.reshape(self.n_ops, *original_shape, 3) - - def get_phase_shift(self, hkl, operation_index): - """ - Get phase shift for a symmetry operation on given Miller indices. - - The phase shift is exp(2πi h·t) where t is the translation. - - Parameters - ---------- - hkl : torch.Tensor, shape (..., 3) - Miller indices. - operation_index : int - Symmetry operation index. - - Returns - ------- - torch.Tensor - Phase shift in radians, shape (...). - """ - hkl = hkl.to(dtype=self.dtype_float, device=self.device) - translation = self.symmetry.translations[operation_index] - phase = 2.0 * np.pi * torch.matmul(hkl, translation) - return phase - - def get_symmetry_mate(self, F_grid, operation_index): - """ - Apply a single symmetry operation to a structure factor grid. - - Parameters - ---------- - F_grid : torch.Tensor, shape (nh, nk, nl) - Complex structure factor grid. - operation_index : int - Index of the symmetry operation (0 to n_ops-1). - - Returns - ------- - torch.Tensor, shape (nh, nk, nl) - Structure factors after applying symmetry operation. - Includes phase shift from translation component. - """ - if operation_index < 0 or operation_index >= self.n_ops: - raise ValueError( - f"Operation index {operation_index} out of range [0, {self.n_ops-1}]" - ) - - if F_grid.shape != self.grid_shape: - raise ValueError( - f"Grid shape {F_grid.shape} doesn't match expected {self.grid_shape}" - ) - - # Get precomputed index grid for this operation - idx_grid = self.index_grids[operation_index] # (nh, nk, nl, 3) - - # Gather structure factors from transformed positions - F_transformed = F_grid[idx_grid[..., 0], idx_grid[..., 1], idx_grid[..., 2]] - - # Apply phase shift from translation: F(h') = F(h) * exp(2πi h·t) - phase = self.phase_shifts[operation_index] - if F_grid.is_complex(): - phase_factor = torch.exp(1j * phase.to(F_grid.dtype)) - F_transformed = F_transformed * phase_factor - # For real-valued grids (amplitudes), no phase shift needed - - return F_transformed - - def get_all_symmetry_mates(self, F_grid): - """ - Get all symmetry-related structure factor grids. - - Parameters - ---------- - F_grid : torch.Tensor, shape (nh, nk, nl) - Complex structure factor grid. - - Returns - ------- - list of torch.Tensor - List of symmetry-related grids. - """ - return [self.get_symmetry_mate(F_grid, i) for i in range(self.n_ops)] - - def symmetry_average(self, F_grid, weights=None): - """ - Average structure factors over all symmetry equivalents. - - This is useful for enforcing symmetry constraints on structure factors. - - Parameters - ---------- - F_grid : torch.Tensor, shape (nh, nk, nl) - Complex structure factor grid. - weights : torch.Tensor, optional - Weights for averaging, shape (nh, nk, nl). - If None, equal weights are used. - - Returns - ------- - torch.Tensor, shape (nh, nk, nl) - Symmetry-averaged structure factors. - """ - mates = self.get_all_symmetry_mates(F_grid) - stacked = torch.stack(mates, dim=0) # (n_ops, nh, nk, nl) - - if weights is not None: - weights = weights.unsqueeze(0) # (1, nh, nk, nl) - stacked = stacked * weights - return stacked.sum(dim=0) / (weights.sum() * self.n_ops) - else: - return stacked.mean(dim=0) - - def expand_to_p1(self, F_asym, asym_mask=None): - """ - Expand structure factors from asymmetric unit to full P1. - - Takes structure factors defined on the asymmetric unit and - generates the full reciprocal space by applying all symmetry - operations. - - Parameters - ---------- - F_asym : torch.Tensor, shape (nh, nk, nl) - Structure factors on asymmetric unit (other positions can be zero). - asym_mask : torch.Tensor, optional - Currently a no-op: this parameter is accepted but ignored by the - implementation. Regardless of its value, positions are filled using - a non-zero heuristic (entries with ``|F| < 1e-10`` are treated as - unset and filled from symmetry mates). - - Returns - ------- - torch.Tensor, shape (nh, nk, nl) - Full structure factor grid with all symmetry equivalents filled. - """ - F_full = F_asym.clone() - - for i in range(1, self.n_ops): # Skip identity - F_mate = self.get_symmetry_mate(F_asym, i) - - # Only fill in positions that are zero (not yet set) - if F_full.is_complex(): - mask = F_full.abs() < 1e-10 - else: - mask = F_full.abs() < 1e-10 - - F_full = torch.where(mask, F_mate, F_full) - - return F_full - - def apply_friedel(self, F_grid): - """ - Apply Friedel's law: F(-h,-k,-l) = F*(h,k,l). - - For normal (non-anomalous) scattering, the structure factor - at -h is the complex conjugate of F(h). - - Parameters - ---------- - F_grid : torch.Tensor, shape (nh, nk, nl) - Complex structure factor grid. - - Returns - ------- - torch.Tensor, shape (nh, nk, nl) - Structure factors averaged toward Friedel symmetry. The result is - ``0.5 * (F(h) + F*(-h))``, i.e. an average of each reflection with - its conjugated Friedel mate, not a hard replacement/enforcement. - """ - # Flip all indices: F(-h,-k,-l) - F_friedel = torch.flip(F_grid, dims=[0, 1, 2]) - - # Roll to handle the asymmetry at 0 index - nh, nk, nl = self.grid_shape - F_friedel = torch.roll(F_friedel, shifts=(1, 1, 1), dims=(0, 1, 2)) - - if F_grid.is_complex(): - F_friedel = F_friedel.conj() - - # Average F(h) and F*(-h) - return 0.5 * (F_grid + F_friedel) - - def is_systematic_absence(self, h, k, l): - """ - Check if a reflection is systematically absent. - - Parameters - ---------- - h, k, l : int - Miller indices. - - Returns - ------- - bool - True if the reflection is systematically absent. - """ - nh, nk, nl = self.grid_shape - idx_h = h % nh - idx_k = k % nk - idx_l = l % nl - return self.systematic_absences[idx_h, idx_k, idx_l].item() - - def is_centric(self, h, k, l): - """ - Check if a reflection is centric. - - Parameters - ---------- - h, k, l : int - Miller indices. - - Returns - ------- - bool - True if the reflection is centric (phase restricted to 0 or π). - """ - nh, nk, nl = self.grid_shape - idx_h = h % nh - idx_k = k % nk - idx_l = l % nl - return self.centric_mask[idx_h, idx_k, idx_l].item() - - def get_epsilon(self): - """ - Compute epsilon (multiplicity) factors for each reflection. - - Epsilon is the number of symmetry operations that map h to itself - (h -> h) or to its Friedel mate (h -> -h). This count is taken - unconditionally for every space group; there is no centric/acentric - branch. Note that folding in Friedel mates inflates the count relative - to the conventional ε (pure rotational multiplicity, h -> h only). - - Returns - ------- - torch.Tensor, shape (nh, nk, nl) - Epsilon factors for each reflection. - """ - epsilon = torch.zeros(self.grid_shape, dtype=torch.int32, device=self.device) - - hkl_flat = self.hkl_grid.reshape(-1, 3) - - for i in range(self.n_ops): - # Get transformed indices - idx_grid = self.index_grids[i] - transformed_flat = idx_grid.reshape(-1, 3) - - # Check if h' ≡ h or h' ≡ -h (same reflection or Friedel) - same = (transformed_flat == hkl_flat).all(dim=-1) - - nh, nk, nl = self.grid_shape - # Also check Friedel (-h, -k, -l) - hkl_neg = (-self.hkl_grid).reshape(-1, 3) - hkl_neg[:, 0] = hkl_neg[:, 0] % nh - hkl_neg[:, 1] = hkl_neg[:, 1] % nk - hkl_neg[:, 2] = hkl_neg[:, 2] % nl - - friedel = (transformed_flat == hkl_neg).all(dim=-1) - - contributes = (same | friedel).reshape(self.grid_shape) - epsilon += contributes.to(torch.int32) - - return epsilon - - def forward(self, F_grid, mode="average"): - """ - Apply symmetry to structure factor grid. - - Parameters - ---------- - F_grid : torch.Tensor, shape (nh, nk, nl) - Complex structure factor grid. - mode : str, default 'average' - Operation mode: - - 'average': Average over all symmetry equivalents - - 'expand': Expand from asymmetric unit to full grid - - 'sum': Sum all symmetry mates (for accumulation) - - Returns - ------- - torch.Tensor, shape (nh, nk, nl) - Processed structure factor grid. - """ - if mode == "average": - return self.symmetry_average(F_grid) - elif mode == "expand": - return self.expand_to_p1(F_grid) - elif mode == "sum": - mates = self.get_all_symmetry_mates(F_grid) - return torch.stack(mates, dim=0).sum(dim=0) - else: - raise ValueError( - f"Unknown mode: {mode}. Use 'average', 'expand', or 'sum'." - ) - - def __call__(self, F_grid, mode="average"): - """Make the class callable.""" - return self.forward(F_grid, mode=mode) - - def get_symmetry_info(self): - """ - Get information about reciprocal space symmetry. - - Returns - ------- - dict - Dictionary with symmetry information. - """ - return { - "space_group": self.space_group, - "n_operations": self.n_ops, - "reciprocal_matrices": self.reciprocal_matrices, - "translations": self.symmetry.translations, - "n_systematic_absences": self.systematic_absences.sum().item(), - "n_centric": self.centric_mask.sum().item(), - "grid_shape": self.grid_shape, - } - - def __repr__(self): - return ( - f"ReciprocalSymmetryGrid(space_group='{self.space_group}', " - f"n_ops={self.n_ops}, grid_shape={self.grid_shape})" - ) - - -# ============================================================================= -# Standalone functions for symmetry expansion -# ============================================================================= - - -def expand_hkl( +def _expand_hkl( + sym, hkl: torch.Tensor, - spacegroup: SpaceGroupLike, include_friedel: bool = True, remove_absences: bool = True, device: Optional[torch.device] = None, @@ -636,10 +42,10 @@ def expand_hkl( Parameters ---------- + sym : SpaceGroup + The space group whose asymmetric unit convention applies. hkl : torch.Tensor, shape (N, 3) Input Miller indices (asymmetric unit). - spacegroup : str, int, or gemmi.SpaceGroup - Space group specification. include_friedel : bool, default True Include Friedel mates (-h, -k, -l). remove_absences : bool, default True @@ -661,12 +67,9 @@ def expand_hkl( device = hkl.device # Get symmetry operations - symmetry = SpaceGroup(spacegroup, dtype=get_float_dtype(), device=device) - n_ops = symmetry.matrices.shape[0] - - # Reciprocal space matrices (transpose of real space) - recip_matrices = symmetry.matrices.transpose(-2, -1) - translations = symmetry.translations + n_ops = sym.n_ops + recip_matrices = sym.reciprocal.matrices.to(device=device) + translations = sym.translations.to(device=device) # Convert hkl to float for matrix operations hkl_float = hkl.to(dtype=get_float_dtype(), device=device) @@ -726,15 +129,8 @@ def expand_hkl( phase_shifts = torch.tensor(unique_phases, dtype=get_float_dtype(), device=device) orig_idx_tensor = torch.tensor(orig_indices, dtype=torch.int64, device=device) - # Remove systematic absences if requested - sg = SpaceGroup(spacegroup) - is_p1 = sg.number == 1 - - if remove_absences and not is_p1: - absence_mask = _check_systematic_absences( - expanded_hkl, symmetry.matrices, translations, device - ) - keep_mask = ~absence_mask + if remove_absences and sym.number != 1: + keep_mask = ~sym.is_absent(expanded_hkl) expanded_hkl = expanded_hkl[keep_mask] phase_shifts = phase_shifts[keep_mask] @@ -743,27 +139,27 @@ def expand_hkl( return expanded_hkl, orig_idx_tensor, phase_shifts -def complete_hkl( +def _complete_hkl( + sym, input_hkl: torch.Tensor, cell: torch.Tensor, - spacegroup: SpaceGroupLike, d_min: float, device: Optional[torch.device] = None, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Complete a set of Miller indices by identifying missing reflections. - Generates every reflection within ``d_min`` for ``spacegroup`` (minus + Generates every reflection within ``d_min`` for ``sym`` (minus systematic absences), then maps the input onto that complete set. This does *not* expand symmetry -- the output stays in the input space group. Parameters ---------- + sym : SpaceGroup + The space group whose asymmetric unit convention applies. input_hkl : torch.Tensor, shape (N, 3) Input Miller indices (may be incomplete). cell : torch.Tensor, shape (6,) Unit cell parameters [a, b, c, alpha, beta, gamma]. - spacegroup : str, int, or gemmi.SpaceGroup - Space group specification. d_min : float High resolution limit in Angstroms. device : torch.device, optional @@ -788,18 +184,8 @@ def complete_hkl( all_hkl = generate_possible_hkl(cell, d_min, device=device) # Get symmetry operations for absence check - symmetry = SpaceGroup(spacegroup, dtype=get_float_dtype(), device=device) - translations = symmetry.translations - - # Remove systematic absences - sg = SpaceGroup(spacegroup) - is_p1 = sg.number == 1 - - if not is_p1: - absence_mask = _check_systematic_absences( - all_hkl, symmetry.matrices, translations, device - ) - all_hkl = all_hkl[~absence_mask] + if sym.number != 1: + all_hkl = all_hkl[~sym.is_absent(all_hkl)] # Build lookup dictionary from input hkl to indices input_hkl_np = input_hkl.cpu().numpy() @@ -824,163 +210,15 @@ def complete_hkl( return all_hkl, input_indices, missing_mask -def expand_reflections( - reflection_data: "ReflectionData", - include_friedel: bool = True, - remove_absences: bool = True, - verbose: int = 1, -) -> "ReflectionData": - """Expand ReflectionData from the asymmetric unit to P1. - - A :func:`expand_hkl` wrapper that carries every ReflectionData field (F, - sigmas, I, phases, FOM, R-free flags) through the expansion. Phases receive - the translation shift; resolution is recomputed and ``bin_indices`` cleared, - since expansion invalidates them. - - Parameters - ---------- - reflection_data : ReflectionData - Input reflection data with hkl, F, F_sigma, etc. - include_friedel : bool, default True - If True, also include Friedel mates (-h, -k, -l). - remove_absences : bool, default True - If True, remove systematically absent reflections from output. - verbose : int, default 1 - Verbosity level. - - Returns - ------- - ReflectionData - New object holding the expanded reflections, spacegroup set to 'P1'. - - See Also - -------- - expand_hkl : Low-level function for HKL expansion without ReflectionData. - """ - from torchref.io.datasets.reflection_data import ReflectionData as RefData - - if reflection_data.hkl is None: - raise ValueError("ReflectionData has no Miller indices loaded") - - space_group = reflection_data.spacegroup or "P1" - device = reflection_data.device - n_orig = len(reflection_data.hkl) - - if verbose > 0: - symmetry = SpaceGroup(space_group, dtype=get_float_dtype(), device=device) - print(f"Expanding reflections for {space_group}") - print(f" Original reflections: {n_orig}") - print(f" Symmetry operations: {symmetry.n_ops}") - - hkl_expanded, orig_idx_tensor, phase_shifts = expand_hkl( - reflection_data.hkl, - space_group, - include_friedel=include_friedel, - remove_absences=remove_absences, - device=device, - ) - - if verbose > 0: - print(f" After expansion: {len(hkl_expanded)} unique reflections") - - expanded = RefData(verbose=reflection_data.verbose, device=device) - - expanded.hkl = hkl_expanded - expanded.cell = ( - reflection_data.cell.clone() if reflection_data.cell is not None else None - ) - expanded.spacegroup = SpaceGroup("P1") # Now in P1 since symmetry is expanded - - if reflection_data.F is not None: - expanded.F = reflection_data.F[orig_idx_tensor] - - if reflection_data.F_sigma is not None: - expanded.F_sigma = reflection_data.F_sigma[orig_idx_tensor] - - if reflection_data.I is not None: - expanded.I = reflection_data.I[orig_idx_tensor] - - if hasattr(reflection_data, "I_sigma") and reflection_data.I_sigma is not None: - expanded.I_sigma = reflection_data.I_sigma[orig_idx_tensor] - - # Phases must absorb the translation shift; with no phases present, stash the - # shifts so a later phase assignment can still apply them. - if hasattr(reflection_data, "phase") and reflection_data.phase is not None: - expanded.phase = reflection_data.phase[orig_idx_tensor] + phase_shifts - else: - expanded._expansion_phase_shifts = phase_shifts - - if hasattr(reflection_data, "fom") and reflection_data.fom is not None: - expanded.fom = reflection_data.fom[orig_idx_tensor] - - if reflection_data.rfree_flags is not None: - expanded.rfree_flags = reflection_data.rfree_flags[orig_idx_tensor] - - if expanded.cell is not None: - expanded._calculate_resolution() - - expanded.bin_indices = None # invalidated by expansion - - expanded.amplitude_source = reflection_data.amplitude_source - expanded.intensity_source = reflection_data.intensity_source - expanded.phase_source = reflection_data.phase_source - expanded.rfree_source = reflection_data.rfree_source - - expanded.source = reflection_data - expanded.last_op = f"expand_to_p1(include_friedel={include_friedel})" - - return expanded - - -def _check_systematic_absences( - hkl: torch.Tensor, - matrices: torch.Tensor, - translations: torch.Tensor, - device: torch.device, -) -> torch.Tensor: - """Bool mask over ``hkl``, True where some op maps h -> h at non-integer h·t. - - ``matrices`` are the *real-space* rotations (n_ops, 3, 3); the transpose is - taken here. - """ - n_refl = len(hkl) - n_ops = matrices.shape[0] - absent = torch.zeros(n_refl, dtype=torch.bool, device=device) - - hkl_float = hkl.to(dtype=get_float_dtype(), device=device) - recip_matrices = matrices.transpose(-2, -1).to( - dtype=get_float_dtype(), device=device - ) - translations = translations.to(dtype=get_float_dtype(), device=device) - - for i in range(n_ops): - R = recip_matrices[i] - t = translations[i] - - hkl_transformed = torch.matmul(hkl_float, R.T) - hkl_transformed_int = torch.round(hkl_transformed).to(torch.int32) - - same_reflection = (hkl_transformed_int == hkl).all(dim=-1) - - h_dot_t = torch.matmul(hkl_float, t) - - phase_mod = torch.abs(h_dot_t - torch.round(h_dot_t)) - non_integer_phase = phase_mod > 1e-6 - - absent = absent | (same_reflection & non_integer_phase) - - return absent - - -def reduce_hkl( +def _reduce_hkl( + sym, hkl_p1: torch.Tensor, - spacegroup: SpaceGroupLike, include_friedel: bool = True, device: Optional[torch.device] = None, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Reduce P1 Miller indices to the asymmetric unit of a target space group. - The inverse of :func:`expand_hkl`: symmetry-equivalent P1 reflections merge + The inverse of :meth:`~torchref.symmetry.spacegroup.SpaceGroup.expand_hkl`: symmetry-equivalent P1 reflections merge into one ASU reflection. The index map has *constant multiplicity* -- its second dimension is always ``n_equiv = n_ops * (2 if include_friedel else 1)`` however many equivalents actually exist in ``hkl_p1`` -- so aggregation needs @@ -988,10 +226,10 @@ def reduce_hkl( Parameters ---------- + sym : SpaceGroup + The target space group. hkl_p1 : torch.Tensor, shape (N, 3) Input Miller indices in P1 (complete hemisphere). - spacegroup : str, int, or gemmi.SpaceGroup - Target space group specification. include_friedel : bool, default True If True, also consider Friedel mates when finding ASU representative. device : torch.device, optional @@ -1012,12 +250,9 @@ def reduce_hkl( device = hkl_p1.device # Get symmetry operations - symmetry = SpaceGroup(spacegroup, dtype=get_float_dtype(), device=device) - n_ops = symmetry.matrices.shape[0] - - # Reciprocal space matrices (transpose of real space) - recip_matrices = symmetry.matrices.transpose(-2, -1) - translations = symmetry.translations + n_ops = sym.n_ops + recip_matrices = sym.reciprocal.matrices.to(device=device) + translations = sym.translations.to(device=device) # Total number of equivalent positions per ASU reflection n_equiv = n_ops * (2 if include_friedel else 1) @@ -1159,9 +394,9 @@ def _asu_condition_vectorized(h, k, l, condition_key): return fn(h, k, l) -def canonicalize_hkl( +def _canonicalize_hkl( + sym, hkl: torch.Tensor, - spacegroup: SpaceGroupLike, include_friedel: bool = True, device: Optional[torch.device] = None, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: @@ -1175,10 +410,10 @@ def canonicalize_hkl( Parameters ---------- + sym : SpaceGroup + The space group whose asymmetric unit convention applies. hkl : torch.Tensor, shape (N, 3), dtype int32 Input Miller indices. - spacegroup : str, int, or gemmi.SpaceGroup - Space group specification. include_friedel : bool, default True Whether Friedel mates are considered equivalent. device : torch.device, optional @@ -1214,13 +449,12 @@ def canonicalize_hkl( empty_i = torch.empty(0, dtype=torch.int64, device=device) return empty_hkl, empty_f, empty_b, empty_i - # Normalize spacegroup (CPU-only: this branch builds numpy-backed lookup tables) - sg_obj = SpaceGroup(spacegroup, dtype=get_float_dtype(), device=torch.device("cpu")) - sg_gemmi = sg_obj._gemmi - asu = gemmi.ReciprocalAsu(sg_gemmi) + # The ASU lookup tables are numpy-backed, so the operations come across to CPU + # regardless of where ``sym`` lives; only the returned tensors honour ``device``. + asu = gemmi.ReciprocalAsu(sym._gemmi) condition_key = asu.condition_str() - recip_mats = sg_obj.matrices.transpose(-2, -1).numpy() # (n_ops, 3, 3) - translations_np = sg_obj.translations.numpy() # (n_ops, 3) + recip_mats = sym.reciprocal.matrices.detach().cpu().numpy() # (n_ops, 3, 3) + translations_np = sym.translations.detach().cpu().numpy() # (n_ops, 3) n_ops = len(recip_mats) hkl_np = hkl.cpu().numpy().astype(np.int32) # (N, 3) @@ -1288,7 +522,7 @@ def canonicalize_hkl( example = hkl_np[np.where(remaining)[0][0]].tolist() raise ValueError( f"canonicalize_hkl could not map {n_unmapped} reflection(s) to the " - f"reciprocal ASU of space group {sg_obj} " + f"reciprocal ASU of space group {sym} " f"(include_friedel={include_friedel}); e.g. hkl={example}. With " f"include_friedel=False the Friedel half of reciprocal space has no " f"pure-rotation representative in the Laue-based CCP4 ASU." @@ -1328,58 +562,3 @@ def canonicalize_hkl( friedel_flags[sort_indices], sort_indices, ) - - -def expand_reciprocal_grid( - F_grid: torch.Tensor, - space_group: str, - mode: str = "average", - include_friedel: bool = True, - device: Optional[torch.device] = None, -) -> torch.Tensor: - """Expand or symmetrize a reciprocal space grid using crystallographic symmetry. - - Convenience wrapper that builds a :class:`ReciprocalSymmetryGrid` for - ``F_grid.shape`` and applies it, so it re-pays the whole precompute on every - call -- hold the class yourself inside a loop. - - Parameters - ---------- - F_grid : torch.Tensor, shape (nh, nk, nl) - Input structure factor grid (can be complex or real). - space_group : str, int, or gemmi.SpaceGroup - Space group specification (e.g., 'P21', 'P212121'). The type hint is - narrowed to ``str``, but any ``SpaceGroupLike`` value is accepted. - mode : {'average', 'expand', 'sum'}, default 'average' - Average over all symmetry equivalents (symmetrize), expand from the - asymmetric unit to the full grid, or sum all symmetry mates. - include_friedel : bool, default True - If True, also apply Friedel symmetry after space group symmetry. - device : torch.device, optional - Device for computation. If None, uses F_grid's device. - - Returns - ------- - torch.Tensor, shape (nh, nk, nl) - Symmetrized or expanded structure factor grid. - """ - if device is None: - device = F_grid.device - - grid_shape = F_grid.shape - dtype = get_float_dtype() if not F_grid.is_complex() else F_grid.real.dtype - - recip_sym = ReciprocalSymmetryGrid( - space_group=space_group, - grid_shape=grid_shape, - dtype_float=dtype, - verbose=0, - device=device, - ) - - F_result = recip_sym(F_grid, mode=mode) - - if include_friedel: - F_result = recip_sym.apply_friedel(F_result) - - return F_result diff --git a/torchref/symmetry/spacegroup.py b/torchref/symmetry/spacegroup.py index da0f7c47..4ec95080 100644 --- a/torchref/symmetry/spacegroup.py +++ b/torchref/symmetry/spacegroup.py @@ -1,39 +1,61 @@ -"""Space group utilities using gemmi as the canonical representation. - -:class:`SpaceGroup` is the interface used throughout torchref: an ``nn.Module`` -that normalizes its input (string, int, ``gemmi.SpaceGroup``, another -``SpaceGroup``, or None for P1), holds the rotation matrices and translations as -registered buffers, and applies them. The module-level functions here are -stateless equivalents plus the FFT/symmetry grid-size helpers. - -Real and reciprocal space use *transposed* conventions -- ``x' = R·x + t`` versus -``h' = Rᵀ·h`` -- so :meth:`SpaceGroup.apply` and :meth:`SpaceGroup.apply_to_hkl` -are not interchangeable. +"""Crystallographic space groups, using gemmi as the canonical source of truth. + +:class:`SpaceGroup` is the interface used throughout torchref. It specialises +:class:`~torchref.symmetry.symmetry.Symmetry` -- which owns the operations and +everything derivable from them -- with the crystallographic identity (Hermann-Mauguin +naming, number, point group, crystal system) and the CCP4 asymmetric-unit +conventions, the two things a bare operation list cannot supply. + +Construction normalizes any ``SpaceGroupLike``: a Hermann-Mauguin string, a number +1-230, a ``gemmi.SpaceGroup``, another :class:`SpaceGroup`, or None for P1. Only the +derived metadata is retained; no persistent ``gemmi`` reference is held, because a +lasting reference to the C++ singleton produces nanobind leak warnings at shutdown. + +Real and reciprocal space use *transposed* conventions. That is handled once, in +:attr:`~torchref.symmetry.symmetry.Symmetry.reciprocal`, rather than being re-decided +per call site. """ from __future__ import annotations -from typing import Union +from typing import Optional, Union import gemmi import torch -import torch.nn as nn from torchref.config import get_float_dtype, normalize_device -from torchref.utils.debug_utils import DebugMixin -from torchref.utils.device_mixin import DeviceMovementMixin +from torchref.symmetry.symmetry import Symmetry # Type alias for space group input - includes SpaceGroup class itself SpaceGroupLike = Union[str, int, gemmi.SpaceGroup, "SpaceGroup", None] +# gemmi stores rotations and translations as integers scaled by 24. +_GEMMI_SCALE = 24.0 + def _normalize_spacegroup(spacegroup: SpaceGroupLike) -> gemmi.SpaceGroup: """Normalize any ``SpaceGroupLike`` to a ``gemmi.SpaceGroup``. - Accepts a Hermann-Mauguin string (spacing-insensitive, retried upper-cased), - a number 1-230, a ``gemmi.SpaceGroup`` (returned unchanged), a - :class:`SpaceGroup` (unwrapped), or None (P1). Raises ``ValueError`` for an - unrecognised name/number, ``TypeError`` for any other type. + Accepts a Hermann-Mauguin string (spacing-insensitive, retried upper-cased), a + number 1-230, a ``gemmi.SpaceGroup`` (returned unchanged), a :class:`SpaceGroup` + (unwrapped), or None (P1). + + Parameters + ---------- + spacegroup : SpaceGroupLike + Space group in any supported form. + + Returns + ------- + gemmi.SpaceGroup + The normalized space group. + + Raises + ------ + ValueError + For an unrecognised name or number. + TypeError + For any other type. """ if spacegroup is None: return gemmi.SpaceGroup("P 1") @@ -41,47 +63,29 @@ def _normalize_spacegroup(spacegroup: SpaceGroupLike) -> gemmi.SpaceGroup: if isinstance(spacegroup, gemmi.SpaceGroup): return spacegroup - # Handle SpaceGroup class instances (forward reference resolved at runtime) + # Duck-typed rather than an isinstance check against SpaceGroup, so this stays + # usable from module scope before the class below is defined. if hasattr(spacegroup, "_sg_hm") and hasattr(spacegroup, "matrices"): return gemmi.find_spacegroup_by_name(spacegroup._sg_hm) if isinstance(spacegroup, int): - # Space group number try: return gemmi.SpaceGroup(spacegroup) except Exception as e: raise ValueError(f"Invalid space group number: {spacegroup}") from e if isinstance(spacegroup, str): - # Try to parse as string - # Clean up common variations sg_clean = spacegroup.strip() - - # Handle double spaces that sometimes appear while " " in sg_clean: sg_clean = sg_clean.replace(" ", " ") - - try: - return gemmi.SpaceGroup(sg_clean) - except Exception: - pass - - # Try without spaces sg_nospace = sg_clean.replace(" ", "") - try: - return gemmi.SpaceGroup(sg_nospace) - except Exception: - pass - - # Try common substitutions - substitutions = [ - (sg_clean, sg_clean), - (sg_nospace, sg_nospace), - (sg_clean.upper(), sg_clean.upper()), - (sg_nospace.upper(), sg_nospace.upper()), - ] - - for _, variant in substitutions: + + for variant in ( + sg_clean, + sg_nospace, + sg_clean.upper(), + sg_nospace.upper(), + ): try: return gemmi.SpaceGroup(variant) except Exception: @@ -99,452 +103,86 @@ def _normalize_spacegroup(spacegroup: SpaceGroupLike) -> gemmi.SpaceGroup: ) -def spacegroup_to_str(spacegroup: SpaceGroupLike, style: str = "short") -> str: - """ - Convert space group to string representation. - - Parameters - ---------- - spacegroup : SpaceGroupLike - Space group in any supported format. - style : str, default 'short' - Output style: - - 'short': No spaces (e.g., 'P212121') - - 'hm': Hermann-Mauguin with spaces (e.g., 'P 21 21 21') - - 'xhm': Extended Hermann-Mauguin, including the setting/cell-choice - token where applicable (e.g., 'P 1 21 1' for a unique-axis-b setting) - - Returns - ------- - str - Space group name in requested style. - """ - sg = _normalize_spacegroup(spacegroup) - - if style == "short": - return sg.short_name() - elif style == "hm": - return sg.hm - elif style == "xhm": - return sg.xhm() - else: - raise ValueError(f"Unknown style: {style}. Use 'short', 'hm', or 'xhm'.") - - -def get_symmetry_operations(spacegroup: SpaceGroupLike): - """ - Get symmetry operations from a space group. - - Parameters - ---------- - spacegroup : SpaceGroupLike - Space group in any supported format. - - Returns - ------- - list of gemmi.Op - List of symmetry operations. - """ - sg = _normalize_spacegroup(spacegroup) - return list(sg.operations()) - - -def get_operations_as_tensors( - spacegroup: SpaceGroupLike, - dtype: torch.dtype = None, - device: torch.device = None, -): - """ - Get symmetry operations as PyTorch tensors. +def _operations_as_tensors( + sg: gemmi.SpaceGroup, + dtype: torch.dtype, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + """Extract a gemmi space group's operations as tensors. Parameters ---------- - spacegroup : SpaceGroupLike - Space group in any supported format. - dtype : torch.dtype, optional - Data type for tensors. Defaults to the configured ``dtypes.float``. - device : torch.device, optional - Device for tensors. Defaults to the configured ``device.current``. + sg : gemmi.SpaceGroup + Normalized space group. + dtype : torch.dtype + Floating dtype for both tensors. + device : torch.device + Device for both tensors. Returns ------- - matrices : torch.Tensor, shape (n_ops, 3, 3) - Rotation matrices. - translations : torch.Tensor, shape (n_ops, 3) - Translation vectors (in fractional coordinates). + matrices : torch.Tensor + Rotation matrices, shape ``(n_ops, 3, 3)``. + translations : torch.Tensor + Fractional translations, shape ``(n_ops, 3)``. """ - if dtype is None: - dtype = get_float_dtype() - device = normalize_device(device) - sg = _normalize_spacegroup(spacegroup) - - # Extract rotation matrices and translations from gemmi operations - # gemmi stores values as integers multiplied by 24, divide to get actual values - gemmi_ops = [ + ops = [ ( - torch.tensor(op.rot, dtype=dtype, device=device) / 24.0, - torch.tensor(op.tran, dtype=dtype, device=device) / 24.0, + torch.tensor(op.rot, dtype=dtype, device=device) / _GEMMI_SCALE, + torch.tensor(op.tran, dtype=dtype, device=device) / _GEMMI_SCALE, ) for op in sg.operations() ] - matrices, translations = zip(*gemmi_ops) - + matrices, translations = zip(*ops) return torch.stack(matrices), torch.stack(translations) -def is_same_spacegroup(sg1: SpaceGroupLike, sg2: SpaceGroupLike) -> bool: - """ - Check if two space groups are the same. - - Parameters - ---------- - sg1, sg2 : SpaceGroupLike - Space groups to compare. - - Returns - ------- - bool - True if the space groups are identical. - """ - return _normalize_spacegroup(sg1).number == _normalize_spacegroup(sg2).number - - -def get_point_group(spacegroup: SpaceGroupLike) -> str: - """ - Get the point group symbol for a space group. - - Parameters - ---------- - spacegroup : SpaceGroupLike - Space group in any supported format. - - Returns - ------- - str - Point group symbol (e.g., '222', 'mmm', '4/mmm'). - """ - sg = _normalize_spacegroup(spacegroup) - return sg.point_group_hm() - - -def get_crystal_system(spacegroup: SpaceGroupLike) -> str: - """ - Get the crystal system for a space group. - - Parameters - ---------- - spacegroup : SpaceGroupLike - Space group in any supported format. - - Returns - ------- - str - Crystal system name (triclinic, monoclinic, orthorhombic, - tetragonal, trigonal, hexagonal, or cubic). - """ - sg = _normalize_spacegroup(spacegroup) - return sg.crystal_system_str() - - -def is_centrosymmetric(spacegroup: SpaceGroupLike) -> bool: - """ - Check if a space group is centrosymmetric. - - Parameters - ---------- - spacegroup : SpaceGroupLike - Space group in any supported format. - - Returns - ------- - bool - True if the space group has an inversion center. - """ - sg = _normalize_spacegroup(spacegroup) - return sg.is_centrosymmetric() - - -def n_operations(spacegroup: SpaceGroupLike) -> int: - """ - Get the number of symmetry operations in a space group. - - Parameters - ---------- - spacegroup : SpaceGroupLike - Space group in any supported format. - - Returns - ------- - int - Number of symmetry operations. - """ - sg = _normalize_spacegroup(spacegroup) - return len(list(sg.operations())) - - -# ============================================================================= -# Grid size utilities (combined FFT-friendly and symmetry-friendly) -# ============================================================================= - - -def is_fft_friendly(n: int) -> bool: - """True if ``n`` factors into 2, 3 and 5 only (radix-2,3,5 FFT sizes). - - ``n <= 0`` is False; 128 and 135 are True, 131 is not. - """ - if n <= 0: - return False - - # Remove all factors of 2, 3, 5 - while n % 2 == 0: - n //= 2 - while n % 3 == 0: - n //= 3 - while n % 5 == 0: - n //= 5 - - # If we're left with 1, the number is FFT-friendly - return n == 1 - - -def find_fft_friendly_size(n: int, divisibility: int = 1) -> int: - """Smallest size >= ``n`` that is FFT-friendly and divisible by ``divisibility``. - - Parameters - ---------- - n : int - Minimum grid size. - divisibility : int, default 1 - Required divisibility (e.g. 2 for a screw axis). - - Returns - ------- - int - Optimal grid size (131 -> 135; 131 with divisibility 2 -> 160). - """ - candidate = n - - # Make sure it satisfies divisibility - if candidate % divisibility != 0: - candidate = ((candidate // divisibility) + 1) * divisibility - - # Now find nearest FFT-friendly size - while not is_fft_friendly(candidate): - candidate += divisibility +class SpaceGroup(Symmetry): + """A crystallographic space group: symmetry operations plus their identity. - return candidate - - -def get_grid_requirements(spacegroup: SpaceGroupLike) -> dict: - """Per-axis grid divisibility required for interpolation-free symmetry expansion. - - Derived from the denominators of the fractional translations, so a grid meeting - them indexes symmetry mates at exact integers. - - Returns - ------- - dict - ``{'nx_mod': int, 'ny_mod': int, 'nz_mod': int}`` -- e.g. P21 gives - ``(1, 2, 1)``, P212121 gives ``(2, 2, 2)``. - """ - import math - from fractions import Fraction - - sg = _normalize_spacegroup(spacegroup) - - # Start with no requirements - nx_lcm = 1 - ny_lcm = 1 - nz_lcm = 1 - - # Analyze each symmetry operation - for op in sg.operations(): - # gemmi stores translations as integers multiplied by 24 - trans = [t / 24.0 for t in op.tran] - - # For each axis, check if translation has fractional component - for axis_idx, t in enumerate(trans): - if abs(t) > 1e-9: - # Convert to fraction and get denominator - frac = Fraction(t).limit_denominator(24) - denom = frac.denominator - - if axis_idx == 0: - nx_lcm = math.lcm(nx_lcm, denom) - elif axis_idx == 1: - ny_lcm = math.lcm(ny_lcm, denom) - else: - nz_lcm = math.lcm(nz_lcm, denom) - - return {"nx_mod": nx_lcm, "ny_mod": ny_lcm, "nz_mod": nz_lcm} - - -def check_grid_compatibility(grid_shape: tuple, spacegroup: SpaceGroupLike) -> dict: - """Check a grid against both space-group divisibility and FFT-friendliness. - - Parameters - ---------- - grid_shape : tuple of int - Grid dimensions (nx, ny, nz). - spacegroup : SpaceGroupLike - Space group in any supported format. - - Returns - ------- - dict - ``compatible`` (both tests pass), ``symmetry_compatible``, - ``fft_friendly``, ``can_use_direct_indexing`` (interpolation-free - expansion possible -- equal to ``symmetry_compatible``), ``issues`` - (per-axis descriptions, empty when compatible) and ``requirements`` - (from :func:`get_grid_requirements`). - """ - nx, ny, nz = grid_shape - sg = _normalize_spacegroup(spacegroup) - requirements = get_grid_requirements(sg) - - issues = [] - sg_name = sg.short_name() - - # Check symmetry requirements - if nx % requirements["nx_mod"] != 0: - issues.append( - f"nx={nx} not divisible by {requirements['nx_mod']} " - f"(required for {sg_name} symmetry)" - ) - - if ny % requirements["ny_mod"] != 0: - issues.append( - f"ny={ny} not divisible by {requirements['ny_mod']} " - f"(required for {sg_name} symmetry)" - ) - - if nz % requirements["nz_mod"] != 0: - issues.append( - f"nz={nz} not divisible by {requirements['nz_mod']} " - f"(required for {sg_name} symmetry)" - ) - - symmetry_compatible = len(issues) == 0 - - # Check FFT-friendly - fft_x = is_fft_friendly(nx) - fft_y = is_fft_friendly(ny) - fft_z = is_fft_friendly(nz) - fft_friendly = fft_x and fft_y and fft_z - - if not fft_x: - issues.append(f"nx={nx} is not FFT-friendly (not a product of 2, 3, 5)") - if not fft_y: - issues.append(f"ny={ny} is not FFT-friendly (not a product of 2, 3, 5)") - if not fft_z: - issues.append(f"nz={nz} is not FFT-friendly (not a product of 2, 3, 5)") - - return { - "compatible": symmetry_compatible and fft_friendly, - "symmetry_compatible": symmetry_compatible, - "fft_friendly": fft_friendly, - "can_use_direct_indexing": symmetry_compatible, - "issues": issues, - "requirements": requirements, - } - - -def suggest_grid_size( - min_grid_shape: tuple, - spacegroup: SpaceGroupLike, - make_fft_friendly: bool = True, -) -> tuple: - """Smallest grid >= ``min_grid_shape`` meeting the symmetry divisibility. - - Parameters - ---------- - min_grid_shape : tuple of int - Minimum (nx, ny, nz) grid dimensions. - spacegroup : SpaceGroupLike - Space group in any supported format. - make_fft_friendly : bool, default True - If True, the result also factors into 2, 3, 5 only. - - Returns - ------- - tuple of int - Suggested grid dimensions (nx, ny, nz). - """ - requirements = get_grid_requirements(spacegroup) - - def find_next_valid(n, divisibility): - """Find next number >= n that satisfies divisibility and FFT constraints.""" - if n % divisibility == 0: - candidate = n - else: - candidate = ((n // divisibility) + 1) * divisibility - - if not make_fft_friendly: - return candidate - - # Find FFT-friendly size that also satisfies divisibility - while not is_fft_friendly(candidate): - candidate += divisibility - - return candidate - - nx = find_next_valid(min_grid_shape[0], requirements["nx_mod"]) - ny = find_next_valid(min_grid_shape[1], requirements["ny_mod"]) - nz = find_next_valid(min_grid_shape[2], requirements["nz_mod"]) - - return (nx, ny, nz) - - -# ============================================================================= -# SpaceGroup class - unified interface combining normalization and operations -# ============================================================================= - - -class SpaceGroup(DeviceMovementMixin, DebugMixin, nn.Module): - """ - Unified space group handler for crystallographic symmetry operations. - - Normalizes its input, holds the operations as buffers, applies them to - fractional coordinates (``__call__``) or Miller indices, and exposes the - grid-size helpers as methods. Only the derived metadata is retained -- no - persistent ``gemmi`` reference is held (see :attr:`_gemmi`). + Inherits the whole operation-derived surface from + :class:`~torchref.symmetry.symmetry.Symmetry` -- expansion, phases, reflection + predicates, grid sizing, map symmetrization -- and adds the crystallographic + naming and the CCP4 asymmetric-unit conventions. Parameters ---------- space_group : str, int, gemmi.SpaceGroup, SpaceGroup, or None - Hermann-Mauguin symbol, number 1-230, gemmi object, another instance, - or None for P1. + Hermann-Mauguin symbol, number 1-230, gemmi object, another instance, or None + for P1. dtype : torch.dtype, optional - Data type for matrices and translations. Defaults to the configured - ``dtypes.float`` (float32 unless ``TORCHREF_DTYPE_FLOAT=float64``). - device : torch.device, default: configured device.current - Device for computation. + Dtype for the operations. Defaults to the configured ``dtypes.float``. + device : torch.device, optional + Device for the operations. Defaults to the configured ``device.current``. Attributes ---------- - matrices : torch.Tensor, shape (n_ops, 3, 3) - Rotation matrices for all symmetry operations (registered buffer). - translations : torch.Tensor, shape (n_ops, 3) - Translation vectors for all symmetry operations (registered buffer). - n_ops : int - Number of symmetry operations. + matrices : torch.Tensor + Rotation matrices, shape ``(n_ops, 3, 3)``. + translations : torch.Tensor + Fractional translations, shape ``(n_ops, 3)``. + + Examples + -------- + >>> sg = SpaceGroup("P212121") + >>> sg.n_ops + 4 + >>> sg.crystal_system + 'orthorhombic' """ def __init__( self, space_group: SpaceGroupLike = None, - dtype: torch.dtype = None, - device: torch.device = None, + dtype: Optional[torch.dtype] = None, + device: Optional[torch.device] = None, ): - super(SpaceGroup, self).__init__() if dtype is None: dtype = get_float_dtype() device = normalize_device(device) - self._device = device - self._dtype = dtype - # Normalize to gemmi.SpaceGroup, extract metadata, then release gemmi_sg = _normalize_spacegroup(space_group) + self._sg_number: int = gemmi_sg.number self._sg_hm: str = gemmi_sg.hm self._sg_short_name: str = gemmi_sg.short_name() @@ -553,254 +191,305 @@ def __init__( self._sg_crystal_system: str = gemmi_sg.crystal_system_str() self._sg_centrosymmetric: bool = gemmi_sg.is_centrosymmetric() - # Get symmetry operations as tensors - matrices, translations = get_operations_as_tensors( - gemmi_sg, dtype=dtype, device=device - ) + matrices, translations = _operations_as_tensors(gemmi_sg, dtype, device) + # gemmi_sg goes out of scope here -- no persistent gemmi reference. - self.register_buffer("matrices", matrices) - self.register_buffer("translations", translations) - # gemmi_sg goes out of scope here — no persistent gemmi reference + super().__init__(matrices=matrices, translations=translations) # ========================================================================= - # Core properties + # Crystallographic identity # ========================================================================= - @property - def n_ops(self) -> int: - """Number of symmetry operations.""" - return self.matrices.shape[0] - @property def _gemmi(self) -> gemmi.SpaceGroup: - """Fresh gemmi.SpaceGroup each access; never cache it -- a persistent - reference to the C++ singleton produces nanobind leak warnings at shutdown.""" + """A fresh ``gemmi.SpaceGroup`` on each access. + + Never cached: a persistent reference to the C++ singleton produces nanobind + leak warnings at interpreter shutdown. + """ return gemmi.find_spacegroup_by_name(self._sg_hm) @property def name(self) -> str: - """Short space group name (e.g., 'P21').""" + """Short space group name, e.g. ``'P21'``.""" return self._sg_short_name @property def hm(self) -> str: - """Hermann-Mauguin notation with spaces (e.g., 'P 21').""" + """Hermann-Mauguin notation with spaces, e.g. ``'P 21'``.""" return self._sg_hm @property def xhm(self) -> str: - """Extended Hermann-Mauguin notation.""" + """Extended Hermann-Mauguin notation, including the setting token.""" return self._sg_xhm @property def number(self) -> int: - """Space group number (1-230).""" + """Space group number, 1-230.""" return self._sg_number @property def gemmi(self) -> gemmi.SpaceGroup: - """Access a gemmi.SpaceGroup object (created on demand, not stored).""" + """A ``gemmi.SpaceGroup``, created on demand and not stored.""" return self._gemmi @property def point_group(self) -> str: - """Point group symbol (e.g., '222', 'mmm').""" + """Point group symbol, e.g. ``'222'`` or ``'mmm'``.""" return self._sg_point_group @property def crystal_system(self) -> str: - """Crystal system name.""" + """Crystal system name, e.g. ``'orthorhombic'``.""" return self._sg_crystal_system @property def centrosymmetric(self) -> bool: - """True if space group has inversion center.""" + """Whether the group has an inversion centre.""" return self._sg_centrosymmetric - @property - def dtype(self) -> torch.dtype: - """Data type used for matrices.""" - return self._dtype + def short_name(self) -> str: + """Short space group name; the callable form of :attr:`name`.""" + return self._sg_short_name - @property - def device(self) -> torch.device: - """Device for matrices.""" - return self._device + def operations(self): + """The gemmi operations, from a temporary gemmi object. + + Returns + ------- + gemmi.GroupOps + The operation list. + """ + return self._gemmi.operations() # ========================================================================= - # Backward compatibility aliases + # Aliases retained for existing callers # ========================================================================= @property def spacegroup(self) -> gemmi.SpaceGroup: - """Alias for gemmi property (backward compatibility).""" + """Alias for :attr:`gemmi`.""" return self._gemmi @property def space_group(self) -> gemmi.SpaceGroup: - """Alias for gemmi property (backward compatibility).""" + """Alias for :attr:`gemmi`.""" return self._gemmi @property def space_group_name(self) -> str: - """Alias for name property (backward compatibility).""" + """Alias for :attr:`name`.""" return self.name @property def space_group_number(self) -> int: - """Alias for number property (backward compatibility).""" + """Alias for :attr:`number`.""" return self.number # ========================================================================= - # Gemmi method delegation for backward compatibility + # Asymmetric-unit conventions # ========================================================================= + # + # These need the CCP4 asymmetric unit, which is keyed by Laue class, so they live + # here rather than on ``Symmetry`` -- "the canonical ASU" is meaningless for a bare + # operation list. The algorithms are in ``reciprocal_symmetry``; these methods are + # the only public way in. - def short_name(self) -> str: - """Get short space group name.""" - return self._sg_short_name - - def operations(self): - """Get symmetry operations (creates temporary gemmi object on demand).""" - return self._gemmi.operations() + def expand_hkl( + self, + hkl: torch.Tensor, + include_friedel: bool = True, + remove_absences: bool = True, + device: Optional[torch.device] = None, + ): + """Expand Miller indices from the asymmetric unit to P1. - # ========================================================================= - # Symmetry operation methods - # ========================================================================= + Parameters + ---------- + hkl : torch.Tensor + Input Miller indices, shape ``(N, 3)``. + include_friedel : bool, default True + Include Friedel mates ``(-h, -k, -l)``. + remove_absences : bool, default True + Drop systematically absent reflections. + device : torch.device, optional + Computation device. Defaults to ``hkl``'s. - def apply( - self, xyz_fractional: torch.Tensor, apply_translation: bool = True - ) -> torch.Tensor: + Returns + ------- + expanded_hkl : torch.Tensor + Expanded indices, shape ``(M, 3)``, dtype ``int32``. + orig_indices : torch.Tensor + Map expanded -> original, shape ``(M,)``: ``F_exp = F_orig[orig_indices]``. + phase_shifts : torch.Tensor + Translation phase offsets in radians, shape ``(M,)``: + ``phase_exp = phase_orig[orig_indices] + phase_shifts``. """ - Apply symmetry operations to fractional coordinates (rotation + translation). + from torchref.symmetry.reciprocal_symmetry import _expand_hkl + + return _expand_hkl( + self, + hkl, + include_friedel=include_friedel, + remove_absences=remove_absences, + device=device, + ) + + def reduce_hkl( + self, + hkl_p1: torch.Tensor, + include_friedel: bool = True, + device: Optional[torch.device] = None, + ): + """Reduce P1 Miller indices to this group's asymmetric unit. - For real space coordinates, applies the full symmetry operation: x' = R·x + t + The inverse of :meth:`expand_hkl`. Parameters ---------- - xyz_fractional : torch.Tensor - Input tensor of shape (N, 3) representing fractional coordinates. - apply_translation : bool, default True - If True, apply the full operation x' = R·x + t. If False, apply - the rotational part only (x' = R·x), as used for Miller indices. + hkl_p1 : torch.Tensor + P1 Miller indices, shape ``(N, 3)``. + include_friedel : bool, default True + Consider Friedel mates when picking the ASU representative. + device : torch.device, optional + Computation device. Defaults to ``hkl_p1``'s. Returns ------- - torch.Tensor - Transformed coordinates of shape (N, 3, ops) where ops is the - number of symmetry operations. - - See Also - -------- - apply_to_hkl : For reciprocal space (Miller indices), rotation only. - """ - coords = xyz_fractional.to(self.matrices.device).to(self.matrices.dtype) - # coords: (N, 3), matrices: (ops, 3, 3) - # Apply rotation: result[n, i, o] = sum_j(matrices[o, i, j] * coords[n, j]) - transformed = torch.einsum("oij,nj->nio", self.matrices, coords) - # transformed: (N, 3, ops) - # Add translations: translations (ops, 3) -> (1, 3, ops) for broadcasting - if apply_translation: - transformed = transformed + self.translations.T.unsqueeze(0) - return transformed # (N, 3, ops) - - def apply_to_hkl(self, hkl: torch.Tensor) -> torch.Tensor: + hkl_asu : torch.Tensor + Unique ASU indices, shape ``(M, 3)``, dtype ``int32``. + reduction_indices : torch.Tensor + Indices into ``hkl_p1`` per equivalent, shape ``(M, n_equiv)``, **-1 where + no P1 reflection exists** -- mask or clamp before gathering, or a -1 + silently reads the last row. + phase_shifts : torch.Tensor + Phase shifts to apply before aggregation, shape ``(M, n_equiv)``. """ - Apply symmetry operations to Miller indices (rotation only, no translation). + from torchref.symmetry.reciprocal_symmetry import _reduce_hkl + + return _reduce_hkl( + self, hkl_p1, include_friedel=include_friedel, device=device + ) - Reciprocal space uses the *transpose*: ``h' = h·R = Rᵀ·h``. Substituting - ``R·h`` gives the wrong equivalents wherever the fractional rotation is - non-symmetric (trigonal, hexagonal, permutation-type cubic ops), silently - corrupting centric flags and epsilon multiplicities. Translations shift - structure-factor phases, not indices, so they are not applied here. + def complete_hkl( + self, + input_hkl: torch.Tensor, + cell: torch.Tensor, + d_min: float, + device: Optional[torch.device] = None, + ): + """Identify reflections missing from a dataset, without expanding symmetry. Parameters ---------- - hkl : torch.Tensor - Input tensor of shape (N, 3) representing Miller indices. + input_hkl : torch.Tensor + Possibly incomplete Miller indices, shape ``(N, 3)``. + cell : torch.Tensor + Unit cell parameters ``[a, b, c, alpha, beta, gamma]``, shape ``(6,)``. + d_min : float + High-resolution limit in Angstroms. + device : torch.device, optional + Computation device. Defaults to ``input_hkl``'s. Returns ------- - torch.Tensor - Transformed Miller indices of shape (N, 3, ops). - - See Also - -------- - apply : Real-space coordinates (``R·x + t``), the transposed convention. + complete_hkl : torch.Tensor + Every index within ``d_min`` minus systematic absences, shape ``(M, 3)``. + input_indices : torch.Tensor + Map complete -> input, shape ``(M,)``, ``-1`` where missing. + missing_mask : torch.Tensor + Boolean, shape ``(M,)``, True where absent from the input. """ - coords = hkl.to(self.matrices.device).to(self.matrices.dtype) - # result[n, i, o] = sum_j matrices[o, j, i] * coords[n, j] = (Rᵀ·h)_i - return torch.einsum("oji,nj->nio", self.matrices, coords) + from torchref.symmetry.reciprocal_symmetry import _complete_hkl - def expand_coords_to_P1(self, xyz_fractional: torch.Tensor) -> torch.Tensor: - """ - Expand fractional coordinates by applying all symmetry operations. + return _complete_hkl(self, input_hkl, cell, d_min, device=device) + + def canonicalize_hkl( + self, + hkl: torch.Tensor, + include_friedel: bool = True, + device: Optional[torch.device] = None, + ): + """Map Miller indices onto their canonical CCP4 ASU representatives. Parameters ---------- - xyz_fractional : torch.Tensor - Input tensor of shape (N, 3) representing fractional coordinates. + hkl : torch.Tensor + Input Miller indices, shape ``(N, 3)``. + include_friedel : bool, default True + Treat Friedel mates as equivalent. With ``False`` the Friedel half of + reciprocal space has no pure-rotation representative in the Laue-based + CCP4 ASU, and unmappable reflections raise. + device : torch.device, optional + Output device. Defaults to ``hkl``'s. The lookup itself runs on CPU + whatever device this group is on, because the ASU tables are numpy-backed. Returns ------- - torch.Tensor - Expanded coordinates of shape (N * ops, 3). + canonical_hkl : torch.Tensor + Remapped indices sorted lexicographically, shape ``(N, 3)``. + phase_shifts : torch.Tensor + Additive phase correction in radians, shape ``(N,)``. + friedel_flags : torch.Tensor + Boolean, shape ``(N,)``, True where Friedel conjugation was applied. + sort_indices : torch.Tensor + Permutation from original to sorted order, shape ``(N,)``. + + Notes + ----- + ``phase_shifts`` assumes the caller conjugates first: the contract is + ``phi_new = torch.where(friedel_flags, -phi_old, phi_old) + phase_shifts``. """ - transformed = self.apply(xyz_fractional) # (N, 3, ops) - N = xyz_fractional.shape[0] - ops = self.n_ops - # (N, 3, ops) -> (N, ops, 3) -> (N * ops, 3) - expanded = transformed.permute(0, 2, 1).reshape(N * ops, 3) - return expanded + from torchref.symmetry.reciprocal_symmetry import _canonicalize_hkl - def forward(self, xyz_fractional: torch.Tensor) -> torch.Tensor: - """Forward pass applies symmetry operations.""" - return self.apply(xyz_fractional) + return _canonicalize_hkl( + self, hkl, include_friedel=include_friedel, device=device + ) # ========================================================================= - # Grid utilities + # Copy and dunder # ========================================================================= - def get_grid_requirements(self) -> dict: - """Per-axis grid divisibility; see :func:`get_grid_requirements`.""" - return get_grid_requirements(self) - - def check_grid_compatibility(self, grid_shape: tuple) -> dict: - """Report whether ``(nx, ny, nz)`` suits this group and the FFT. + def copy(self) -> "SpaceGroup": + """An independent copy with cloned operations and an empty cache. - Returns the report dict documented in :func:`check_grid_compatibility`. + Returns + ------- + SpaceGroup + New instance carrying the same symmetry, dtype and device. """ - return check_grid_compatibility(grid_shape, self) - - def suggest_grid_size( - self, min_grid_shape: tuple, make_fft_friendly: bool = True - ) -> tuple: - """Smallest valid grid >= ``min_grid_shape``; see :func:`suggest_grid_size`.""" - return suggest_grid_size(min_grid_shape, self, make_fft_friendly) - - # ========================================================================= - # Dunder methods - # ========================================================================= - - def __repr__(self) -> str: - return f"SpaceGroup('{self.name}', number={self.number}, n_ops={self.n_ops})" + new = SpaceGroup.__new__(SpaceGroup) + new._sg_number = self._sg_number + new._sg_hm = self._sg_hm + new._sg_short_name = self._sg_short_name + new._sg_xhm = self._sg_xhm + new._sg_point_group = self._sg_point_group + new._sg_crystal_system = self._sg_crystal_system + new._sg_centrosymmetric = self._sg_centrosymmetric + # Through Symmetry's own initializer, so the operand-consistency checks and the + # device/dtype reconciliation in ``__post_init__`` run on the copy too. + Symmetry.__init__( + new, + matrices=self.matrices.clone(), + translations=self.translations.clone(), + ) + return new def __hash__(self) -> int: - """Hash based on space group number.""" + """Hash on the space group number.""" return hash(self._sg_number) def __eq__(self, other) -> bool: - """Equality based on space group number.""" + """Equality on the space group number; also compares to a ``gemmi.SpaceGroup``.""" if isinstance(other, SpaceGroup): return self._sg_number == other._sg_number if isinstance(other, gemmi.SpaceGroup): return self._sg_number == other.number return False - # ========================================================================= - # Device movement - # ========================================================================= + def __repr__(self) -> str: + return f"SpaceGroup('{self.name}', number={self.number}, n_ops={self.n_ops})" - def copy(self) -> "SpaceGroup": - """A new SpaceGroup with the same symmetry, dtype and device (fresh buffers).""" - new_sg = SpaceGroup(self._sg_hm, dtype=self._dtype, device=self._device) - return new_sg + +__all__ = ["SpaceGroup", "SpaceGroupLike"] diff --git a/torchref/symmetry/symmetry.py b/torchref/symmetry/symmetry.py index 5213a130..5d68cffc 100644 --- a/torchref/symmetry/symmetry.py +++ b/torchref/symmetry/symmetry.py @@ -1,16 +1,775 @@ -"""DEPRECATED: ``Symmetry`` is a bare alias for :class:`SpaceGroup`. +"""Symmetry groups and everything derivable from their operations alone. -``Symmetry = SpaceGroup`` literally, so ``isinstance`` checks against either name -succeed and existing calls keep working. No ``DeprecationWarning`` is emitted -- -nothing tells a caller to migrate. Prefer ``SpaceGroup`` in new code. +:class:`Symmetry` holds a group as rotation matrices and fractional translations and +derives what needs nothing else: expansion of positions and Miller indices, +translation phases, the reflection predicates (centric, systematically absent, +multiplicity), symmetry-compatible grid sizes, and real-space map symmetrization. +Nothing here knows about crystals, so a group assembled from a raw operation list +serves non-crystallographic symmetry equally well. +:class:`~torchref.symmetry.spacegroup.SpaceGroup` specialises it with the +crystallographic identity and the CCP4 asymmetric-unit conventions. + +Rotation composes; translation does not. Translation acts *additively* on real-space +positions and as a *phase* ``exp(2 pi i h.t)`` in reciprocal space, so the two are +separate primitives rather than one method behind a flag. And ``h' = R^T h`` is not a +third law: it is :meth:`Symmetry.apply_rotations` on :attr:`Symmetry.reciprocal`, the +same group carrying the transposed rotations. Routing every caller through those +primitives is what keeps the real/reciprocal transpose from being re-decided, and got +wrong, at each site. + +Every expansion returns operations on the *leading* axis -- ``(n_ops, ...)``. """ -import warnings +from __future__ import annotations + +import math +from dataclasses import dataclass, field +from fractions import Fraction +from typing import TYPE_CHECKING + +import torch + +from torchref.config import get_float_dtype +from torchref.utils.device_mixin import DeviceMixin + +if TYPE_CHECKING: + from torchref.symmetry.cell import Cell + +# Fractional translations are exact multiples of 1/24 (the denominator gemmi stores +# them over), which is what lets :meth:`Symmetry.grid_requirements` recover exact +# denominators instead of guessing at a float tolerance. +_TRANSLATION_DENOMINATOR = 24 + +# Tolerance for "this dot product is an integer" when testing a translation against a +# reflection. Miller indices are small and translations are k/24, so the true values +# are either integral or off by at least 1/24 -- far outside float32 noise. +_PHASE_TOL = 1e-6 + + +def is_fft_friendly(n: int) -> bool: + """Whether ``n`` factors into 2, 3 and 5 only, as radix-2,3,5 FFTs want. + + Parameters + ---------- + n : int + Candidate grid length. + + Returns + ------- + bool + True for 128 and 135, False for 131 and for any ``n <= 0``. + """ + if n <= 0: + return False + + for factor in (2, 3, 5): + while n % factor == 0: + n //= factor + + return n == 1 + + +def find_fft_friendly_size(n: int, divisibility: int = 1) -> int: + """Smallest FFT-friendly size at or above ``n`` that ``divisibility`` divides. + + Parameters + ---------- + n : int + Minimum grid length. + divisibility : int, default 1 + Required divisor, e.g. 2 for a screw axis. + + Returns + ------- + int + 131 gives 135; 131 at divisibility 2 gives 160. + """ + candidate = n + if candidate % divisibility != 0: + candidate = ((candidate // divisibility) + 1) * divisibility + + while not is_fft_friendly(candidate): + candidate += divisibility + + return candidate + + +@dataclass(eq=False, repr=False) +class Symmetry(DeviceMixin): + """A symmetry group as operations, plus everything they imply. + + Mutable by design; prefer :meth:`copy` over editing in place. Derived quantities + live in one cache that :meth:`reset_cache` clears -- and that ``.to()`` clears for + you, since :class:`~torchref.utils.device_mixin.DeviceMixin` invalidates caches on + every move. + + Parameters + ---------- + matrices : torch.Tensor + Rotation matrices, shape ``(n_ops, 3, 3)``, in the fractional basis. + translations : torch.Tensor + Fractional translations, shape ``(n_ops, 3)``. Coerced onto ``matrices``' + device and dtype, so the two can never end up split. + + Attributes + ---------- + matrices, translations : torch.Tensor + The operations, as above. + + Notes + ----- + Holds no refinable parameters -- symmetry operations are fixed constants -- so it + is a plain dataclass rather than an ``nn.Module``, and there is no gradient path + through it. + + Writing into :attr:`matrices` or :attr:`translations` in place does **not** + invalidate the cache, so the reciprocal stack and any built operator would keep + answering for the old operations. Nothing here mutates them, and :meth:`copy` is + the intended way to vary a group; if you must edit in place, call + :meth:`reset_cache` afterwards. + """ + + matrices: torch.Tensor + translations: torch.Tensor + _cache: dict = field(default_factory=dict, repr=False) + + def __post_init__(self) -> None: + """Validate the operation shapes and put both tensors on one device/dtype.""" + if self.matrices.ndim != 3 or self.matrices.shape[-2:] != (3, 3): + raise ValueError( + f"matrices must have shape (n_ops, 3, 3), got " + f"{tuple(self.matrices.shape)}" + ) + if self.translations.ndim != 2 or self.translations.shape[-1] != 3: + raise ValueError( + f"translations must have shape (n_ops, 3), got " + f"{tuple(self.translations.shape)}" + ) + if self.translations.shape[0] != self.matrices.shape[0]: + raise ValueError( + f"matrices and translations disagree on n_ops: " + f"{self.matrices.shape[0]} vs {self.translations.shape[0]}" + ) + + # One device and dtype for the pair. A split here would surface much later, as + # a device mismatch inside whichever expansion happened to touch both. + self.translations = self.translations.to( + device=self.matrices.device, dtype=self.matrices.dtype + ) + + # ========================================================================= + # Identity + # ========================================================================= + + @property + def n_ops(self) -> int: + """Number of symmetry operations.""" + return int(self.matrices.shape[0]) + + @property + def device(self) -> torch.device: + """Device the operations live on.""" + return self.matrices.device + + @property + def dtype(self) -> torch.dtype: + """Floating dtype of the operations.""" + return self.matrices.dtype + + # ========================================================================= + # The primitives + # ========================================================================= + + @property + def reciprocal(self) -> "Symmetry": + """The same group carrying the transposed rotations, for reciprocal space. + + ``h' = R^T h``, so Miller-index expansion is :meth:`apply_rotations` on this + object rather than a law of its own. Cached. + + Returns + ------- + Symmetry + Group with ``matrices = R^T`` and this group's translations. + + Notes + ----- + The translations come along because :meth:`phase_factors` needs them, but in + reciprocal space a translation is a *phase*, not a shift: calling + :meth:`apply_translations` on this object is meaningless. + """ + cached = self._cache.get("reciprocal") + if cached is None: + cached = Symmetry( + matrices=self.matrices.transpose(-2, -1).contiguous(), + translations=self.translations, + ) + self._cache["reciprocal"] = cached + return cached + + def apply_rotations(self, v: torch.Tensor) -> torch.Tensor: + """Rotate ``v`` by every operation: ``R v``. + + The one batched matmul behind every expansion. Translation is not applied -- + compose with :meth:`apply_translations` for real-space positions, or use + :attr:`reciprocal` for Miller indices. + + Parameters + ---------- + v : torch.Tensor + Vectors of shape ``(N, 3)``, in the fractional basis. + + Returns + ------- + torch.Tensor + Shape ``(n_ops, N, 3)``, rotated by each operation in turn. + """ + v = v.to(device=self.device, dtype=self.dtype) + # result[o, n, i] = sum_j matrices[o, i, j] * v[n, j] + return torch.einsum("oij,nj->oni", self.matrices, v) + + def apply_translations(self, v: torch.Tensor) -> torch.Tensor: + """Add each operation's fractional translation: ``v + t``. + + Real space only; in reciprocal space a translation is a phase, see + :meth:`phase_factors`. + + Parameters + ---------- + v : torch.Tensor + Shape ``(n_ops, N, 3)`` -- typically straight out of + :meth:`apply_rotations` -- or ``(N, 3)`` to broadcast one set of vectors + across all operations. + + Returns + ------- + torch.Tensor + Shape ``(n_ops, N, 3)``. + """ + v = v.to(device=self.device, dtype=self.dtype) + if v.ndim == 2: + v = v.unsqueeze(0).expand(self.n_ops, -1, -1) + elif v.shape[0] != self.n_ops: + raise ValueError( + f"expected a leading axis of n_ops={self.n_ops} or a bare (N, 3), " + f"got {tuple(v.shape)}" + ) + return v + self.translations.unsqueeze(1) + + def phase_factors(self, hkl: torch.Tensor) -> torch.Tensor: + """Structure-factor phase shift per operation: ``exp(2 pi i h.t)``. + + The reciprocal-space action of a translation. Pairs with + ``self.reciprocal.apply_rotations(hkl)`` to combine symmetry-equivalent + structure factors. + + Parameters + ---------- + hkl : torch.Tensor + Miller indices, shape ``(N, 3)``, integer or float. + + Returns + ------- + torch.Tensor + Complex phase factors of shape ``(n_ops, N)``. + + Notes + ----- + The phase stays at this group's floating dtype rather than being forced to + float32, so a float64 configuration yields complex128 instead of silently + narrowing to complex64. + + This is the *complex factor* form, ``exp(+2 pi i h.t)``, for combining + structure factors. Expanding the *phases* of reflection data instead needs a + signed offset in radians, ``-2 pi h.t``, which + :meth:`~torchref.symmetry.spacegroup.SpaceGroup.expand_hkl` computes itself. + The two are not interchangeable and the sign difference is invisible in + P21/P212121/C2 -- see ``tests/unit/symmetry/test_phase_convention.py``. + """ + hkl = hkl.to(device=self.device, dtype=self.dtype) + h_dot_t = torch.matmul(hkl, self.translations.T).T # (n_ops, N) + return torch.exp(1j * (2.0 * math.pi * h_dot_t)) + + # ========================================================================= + # Named expansions + # ========================================================================= + + def expand_positions(self, xyz_fractional: torch.Tensor) -> torch.Tensor: + """Expand fractional positions by every operation: ``R x + t``. + + Fixes the composition order in one place -- ``R x + t``, not ``R (x + t)``. + + Parameters + ---------- + xyz_fractional : torch.Tensor + Fractional coordinates, shape ``(N, 3)``. + + Returns + ------- + torch.Tensor + Shape ``(n_ops, N, 3)``, unwrapped (values may fall outside ``[0, 1)``). + """ + return self.apply_translations(self.apply_rotations(xyz_fractional)) + + def expand_directions(self, v_fractional: torch.Tensor) -> torch.Tensor: + """Expand fractional directions or displacements: ``R v``, no translation. + + Parameters + ---------- + v_fractional : torch.Tensor + Fractional vectors, shape ``(N, 3)``. + + Returns + ------- + torch.Tensor + Shape ``(n_ops, N, 3)``. + """ + return self.apply_rotations(v_fractional) + + def expand_to_P1(self, xyz_fractional: torch.Tensor) -> torch.Tensor: + """Flatten :meth:`expand_positions` into one P1 coordinate list. + + Parameters + ---------- + xyz_fractional : torch.Tensor + Fractional coordinates, shape ``(N, 3)``. + + Returns + ------- + torch.Tensor + Shape ``(n_ops * N, 3)``, operation-major. + """ + return self.expand_positions(xyz_fractional).reshape(-1, 3) + + def expand_reciprocal(self, hkl: torch.Tensor) -> torch.Tensor: + """Symmetry-equivalent Miller indices: ``h' = R^T h``. + + Parameters + ---------- + hkl : torch.Tensor + Miller indices, shape ``(N, 3)``. + + Returns + ------- + torch.Tensor + Shape ``(n_ops, N, 3)``, rounded to ``int64``. Rounding is exact for valid + operations on integer indices and only mops up float error. + """ + equivalents = self.reciprocal.apply_rotations(hkl) + return torch.round(equivalents).to(torch.int64) + + # ========================================================================= + # Reflection predicates + # ========================================================================= + + def is_centric(self, hkl: torch.Tensor) -> torch.Tensor: + """Whether each reflection is centric, i.e. some operation maps ``h -> -h``. + + A centric reflection has its phase restricted to 0 or pi. + + Parameters + ---------- + hkl : torch.Tensor + Miller indices, shape ``(..., 3)``. + + Returns + ------- + torch.Tensor + Boolean mask of shape ``(...)``, on ``hkl``'s device. + """ + original_shape = hkl.shape[:-1] + with torch.no_grad(): + flat = hkl.reshape(-1, 3) + equivalents = self.expand_reciprocal(flat) # (n_ops, N, 3) + target = -flat.to(device=equivalents.device, dtype=torch.int64) + centric = (equivalents == target).all(dim=-1).any(dim=0) + return centric.reshape(original_shape).to(hkl.device) + + def is_absent(self, hkl: torch.Tensor) -> torch.Tensor: + """Whether each reflection is systematically absent. + + Absent when some operation maps ``h -> h`` while ``h.t`` is non-integral: the + reflection is destroyed by interference from that translation. + + Parameters + ---------- + hkl : torch.Tensor + Miller indices, shape ``(..., 3)``. + + Returns + ------- + torch.Tensor + Boolean mask of shape ``(...)``, on ``hkl``'s device. + """ + original_shape = hkl.shape[:-1] + with torch.no_grad(): + flat = hkl.reshape(-1, 3) + equivalents = self.expand_reciprocal(flat) # (n_ops, N, 3) + target = flat.to(device=equivalents.device, dtype=torch.int64) + maps_to_self = (equivalents == target).all(dim=-1) # (n_ops, N) + + h_dot_t = torch.matmul( + flat.to(device=self.device, dtype=self.dtype), self.translations.T + ).T # (n_ops, N) + non_integral = (h_dot_t - torch.round(h_dot_t)).abs() > _PHASE_TOL + + absent = (maps_to_self & non_integral).any(dim=0) + return absent.reshape(original_shape).to(hkl.device) + + def epsilon(self, hkl: torch.Tensor) -> torch.Tensor: + """Reflection multiplicity: operations mapping ``h`` to ``h`` or to ``-h``. + + Parameters + ---------- + hkl : torch.Tensor + Miller indices, shape ``(N, 3)``. + + Returns + ------- + torch.Tensor + Multiplicities of shape ``(N,)`` at the configured float dtype, floored at + 1, returned on ``hkl``'s device so it can weight data sitting beside it. + + Notes + ----- + Friedel mates are folded in unconditionally, with no centric/acentric branch. + That inflates the count relative to the conventional epsilon, which counts pure + rotational multiplicity (``h -> h``) only. Downstream sigma_A estimation is + calibrated against this convention. + """ + float_dtype = get_float_dtype() + with torch.no_grad(): + equivalents = self.expand_reciprocal(hkl) # (n_ops, N, 3) + target = hkl.to(device=equivalents.device, dtype=torch.int64) + same = (equivalents == target).all(dim=-1) + friedel = (equivalents == -target).all(dim=-1) + eps = (same | friedel).sum(dim=0).clamp(min=1).to(float_dtype) + return eps.to(hkl.device) + + # ========================================================================= + # Grid compatibility + # ========================================================================= + + def grid_requirements(self) -> dict: + """Per-axis grid divisibility for interpolation-free symmetry expansion. + + Read off the denominators of the fractional translations, so a grid meeting + them indexes every symmetry mate at an exact integer. + + Returns + ------- + dict + ``{'nx_mod': int, 'ny_mod': int, 'nz_mod': int}`` -- P21 gives + ``(1, 2, 1)``, P212121 gives ``(2, 2, 2)``. + """ + mods = [1, 1, 1] + # Recover each denominator from the integer numerator over 1/24 rather than + # from the float directly: 1/3 is not representable in float32, so + # ``Fraction(float)`` would need a tolerance where this is exact. + numerators = torch.round( + self.translations.detach().cpu().double() * _TRANSLATION_DENOMINATOR + ).to(torch.int64) + + for op_numerators in numerators.tolist(): + for axis, numerator in enumerate(op_numerators): + if numerator % _TRANSLATION_DENOMINATOR == 0: + continue + denominator = Fraction( + int(numerator), _TRANSLATION_DENOMINATOR + ).denominator + mods[axis] = math.lcm(mods[axis], denominator) + + return {"nx_mod": mods[0], "ny_mod": mods[1], "nz_mod": mods[2]} + + def check_grid_compatibility(self, grid_shape: tuple) -> dict: + """Check a grid against this group's divisibility and against the FFT. + + Parameters + ---------- + grid_shape : tuple of int + Grid dimensions ``(nx, ny, nz)``. + + Returns + ------- + dict + ``compatible`` (both tests pass), ``symmetry_compatible``, + ``fft_friendly``, ``can_use_direct_indexing`` (interpolation-free expansion + possible; equal to ``symmetry_compatible``), ``issues`` (per-axis + descriptions, empty when compatible) and ``requirements`` (from + :meth:`grid_requirements`). + """ + requirements = self.grid_requirements() + issues = [] + + for axis, name in enumerate(("nx", "ny", "nz")): + modulus = requirements[f"{name}_mod"] + length = int(grid_shape[axis]) + if length % modulus != 0: + issues.append( + f"{name}={length} not divisible by {modulus} " + f"(required by the symmetry)" + ) + + symmetry_compatible = len(issues) == 0 + + fft_friendly = True + for axis, name in enumerate(("nx", "ny", "nz")): + length = int(grid_shape[axis]) + if not is_fft_friendly(length): + fft_friendly = False + issues.append( + f"{name}={length} is not FFT-friendly (not a product of 2, 3, 5)" + ) + + return { + "compatible": symmetry_compatible and fft_friendly, + "symmetry_compatible": symmetry_compatible, + "fft_friendly": fft_friendly, + "can_use_direct_indexing": symmetry_compatible, + "issues": issues, + "requirements": requirements, + } + + def can_index_directly(self, grid_shape: tuple) -> bool: + """Whether ``grid_shape`` admits interpolation-free symmetry expansion. + + The question :meth:`symmetrize_map` answers internally when choosing an + implementation, exposed so callers can ask it without building an operator. + + Parameters + ---------- + grid_shape : tuple of int + Grid dimensions ``(nx, ny, nz)``. + + Returns + ------- + bool + True when every symmetry mate lands on an exact grid point. + """ + return bool(self.check_grid_compatibility(grid_shape)["symmetry_compatible"]) + + def suggest_grid_size( + self, min_grid_shape: tuple, make_fft_friendly: bool = True + ) -> tuple: + """Smallest grid at or above ``min_grid_shape`` meeting the divisibility. + + Parameters + ---------- + min_grid_shape : tuple of int + Minimum dimensions ``(nx, ny, nz)``. + make_fft_friendly : bool, default True + Also require factors of 2, 3 and 5 only. + + Returns + ------- + tuple of int + Suggested ``(nx, ny, nz)``. + """ + requirements = self.grid_requirements() + + def next_valid(length: int, divisibility: int) -> int: + if make_fft_friendly: + return find_fft_friendly_size(length, divisibility) + if length % divisibility == 0: + return length + return ((length // divisibility) + 1) * divisibility + + return tuple( + next_valid(int(min_grid_shape[axis]), requirements[f"{name}_mod"]) + for axis, name in enumerate(("nx", "ny", "nz")) + ) + + def optimal_grid_size( + self, cell: "Cell", max_res: float, make_fft_friendly: bool = True + ) -> tuple: + """Smallest grid that samples ``cell`` to ``max_res`` and suits this group. + + Composes the cell's Shannon-Nyquist minimum with + :meth:`suggest_grid_size`; the oversampling factor is the cell's, so every + grid-sizing path shares one setting. + + Parameters + ---------- + cell : Cell + Unit cell. + max_res : float + Maximum resolution in Angstroms. + make_fft_friendly : bool, default True + Also require factors of 2, 3 and 5 only. + + Returns + ------- + tuple of int + Grid dimensions ``(nx, ny, nz)``. + """ + return self.suggest_grid_size( + cell.compute_grid_size(max_res), make_fft_friendly=make_fft_friendly + ) + + # ========================================================================= + # Real-space maps + # ========================================================================= + + def map_operator(self, map_shape): + """Cached operator that applies this group to maps of ``map_shape``. + + Two implementations: exact integer indexing when the grid permits it, and + ``grid_sample`` interpolation otherwise. Which one you get depends on the grid, + so a mis-sized grid costs accuracy -- :meth:`can_index_directly` reports the + distinction, and :meth:`suggest_grid_size` fixes it. + + Parameters + ---------- + map_shape : tuple of int + Density map dimensions ``(nx, ny, nz)``. + + Returns + ------- + _MapSymmetryDirect or _MapSymmetryInterpolation + The operator, memoized for the most recent shape only. + + Notes + ----- + Only the last shape is kept. The interpolating operator holds sampling grids of + shape ``(n_ops, nx, ny, nz, 3)`` -- hundreds of megabytes at production grid + sizes -- so a dictionary keyed on shape would quietly make this object + expensive to hold and to move between devices. + """ + from torchref.symmetry.map_symmetry import build_map_operator + + key = tuple(int(n) for n in map_shape) + cached = self._cache.get("map_operator") + if cached is not None and cached[0] == key: + return cached[1] + + operator = build_map_operator(self, key) + self._cache["map_operator"] = (key, operator) + return operator + + def symmetrize_map( + self, density_map: torch.Tensor, combine: str = "sum" + ) -> torch.Tensor: + """Apply every operation to a density map and combine the mates. + + Parameters + ---------- + density_map : torch.Tensor + Asymmetric-unit density, shape ``(nx, ny, nz)``. + combine : {'sum', 'max'}, default 'sum' + ``'sum'`` for electron density, ``'max'`` for masks and boolean data. + + Returns + ------- + torch.Tensor + Symmetrized map, same shape as the input. Returned unchanged for a + one-operation group. + """ + if self.n_ops == 1: + return density_map + return self.map_operator(density_map.shape).symmetrize(density_map, combine) + + def expand_map_to_P1(self, density_map: torch.Tensor) -> torch.Tensor: + """Every symmetry mate of a density map, stacked. + + Parameters + ---------- + density_map : torch.Tensor + Asymmetric-unit density, shape ``(nx, ny, nz)``. + + Returns + ------- + torch.Tensor + Shape ``(n_ops, nx, ny, nz)``. + """ + return self.map_operator(density_map.shape).all_mates(density_map) + + # ========================================================================= + # Reciprocal-space extraction + # ========================================================================= + + def reciprocal_extractor(self, hkl: torch.Tensor, grid_shape: tuple): + """Cached extractor pulling symmetrized structure factors off a grid. + + Precomputes the equivalent indices, phases and flat gather indices for a fixed + ``hkl`` and ``grid_shape``, so each later call is one gather, multiply and sum. + + Parameters + ---------- + hkl : torch.Tensor + Target Miller indices, shape ``(N, 3)``. + grid_shape : tuple of int + Reciprocal grid dimensions ``(nx, ny, nz)``. + + Returns + ------- + ReciprocalSymmetryExtractor + Memoized against ``hkl``'s identity and ``grid_shape``; a different tensor + or shape rebuilds it. + """ + from torchref.base.reciprocal.symmetry import ReciprocalSymmetryExtractor + from torchref.utils.caching import ParameterFingerprint + + key = tuple(int(n) for n in grid_shape) + cached = self._cache.get("reciprocal_extractor") + if cached is not None: + cached_key, fingerprint, extractor = cached + if cached_key == key and fingerprint.matches([hkl]): + return extractor + + extractor = ReciprocalSymmetryExtractor(hkl, self, key) + self._cache["reciprocal_extractor"] = ( + key, + ParameterFingerprint([hkl]), + extractor, + ) + return extractor + + # ========================================================================= + # Cache and copy + # ========================================================================= + + def reset_cache(self) -> None: + """Drop every derived quantity. + + Called for you by :class:`~torchref.utils.device_mixin.DeviceMixin` on any + ``.to()``, including one targeting the current device. + """ + self._cache = {} + + def _apply(self, fn, recurse: bool = True): + """Clear the cache *before* the traversal moves anything. + + The base traversal walks ``__dict__`` first and invalidates caches afterwards, + which would transfer the cached sampling grids to the new device only to + discard them -- hundreds of megabytes of pointless copying at production grid + sizes. + """ + self.reset_cache() + return super()._apply(fn, recurse) + + def copy(self) -> "Symmetry": + """An independent copy with cloned operations and an empty cache. + + Returns + ------- + Symmetry + New instance; mutating its tensors cannot affect this one. + """ + return type(self)( + matrices=self.matrices.clone(), + translations=self.translations.clone(), + ) + + # ========================================================================= + # Dunder + # ========================================================================= + + def __len__(self) -> int: + """Number of symmetry operations.""" + return self.n_ops -from torchref.symmetry.spacegroup import SpaceGroup, SpaceGroupLike + def __repr__(self) -> str: + return f"Symmetry(n_ops={self.n_ops})" -# Backward compatibility alias - Symmetry is now SpaceGroup -# Using the class directly so isinstance() checks work. -Symmetry = SpaceGroup -__all__ = ["Symmetry", "SpaceGroupLike"] +__all__ = ["Symmetry", "find_fft_friendly_size", "is_fft_friendly"] From c85cb1223e27571d1a430a2a0a53aa05652fc11f Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Wed, 26 Aug 2026 13:55:48 +0200 Subject: [PATCH 055/250] Split Model's information half into a ModelContext Model mixed two things: the parameters being refined, and the structure it was loaded from. ModelContext now holds the second -- unit cell, space group, atom table, link records, alternative-conformation groups, provenance and the configuration flags -- so the model's own surface is parameters and behaviour, and the crystallographic context can be handed to code that needs only that. Access is deliberately hybrid. cell, spacegroup and pdb keep forwarding properties because they carry roughly 770 call sites between them, and churning all of those would bury a real regression in rename noise. The low-traffic fields move outright and are reached as model.ctx.strip_H, .initialized, .links, .altloc_pairs, .verbose, .exclude_H_from_sf and the input paths. copy() collapses from field-by-field assignment to one ModelContext.copy(), which deep-copies the atom table and clones the cell and space group. Cloning the space group is load-bearing now that it is a mutable dataclass rather than an nn.Module: a shared reference would let an edit through one model's context reach every model copied from it. device and dtype_float stay on the model rather than moving to the context. They are live DeviceMixin trackers that the traversal rewrites in place on whichever object owns the tensors, so relocating them would either duplicate the source of truth or route the hottest device path through a property. Model.symmetry is gone; it was a bare alias for spacegroup and misleads now that Symmetry is a real class. ModelFT.copy() also loses its skip_modules guard, which only existed because SpaceGroup used to register as a submodule. Two defensive lookups had to go with it, both silent. set_adp_mode guarded on getattr(self, "initialized", False), which read the default once the field moved and turned the whole method into a no-op -- mode switches reported success and changed nothing. state_dict used hasattr(self, "altloc_pairs") and would have saved an empty grouping on every write; that path had no test, so this adds a round-trip one, checked against a null control. Co-Authored-By: Claude Opus 5 (1M context) --- AGENTS.md | 2 +- docs/changelog.rst | 4 + tests/helpers/device_cases.py | 2 + .../unit/model/test_create_from_state_dict.py | 30 +++ tests/unit/model/test_model.py | 4 +- torchref/cli/collection_difference_refine.py | 2 +- torchref/experimental/alignment/pipeline.py | 4 +- torchref/io/metadata.py | 2 +- torchref/model/__init__.py | 6 +- torchref/model/context.py | 122 +++++++++ torchref/model/model.py | 238 +++++++++--------- torchref/model/model_ft.py | 77 +++--- torchref/refinement/base_refinement.py | 6 +- .../refinement/targets/geometry/non_bonded.py | 6 +- .../targets/geometry/non_bonded_h.py | 6 +- 15 files changed, 321 insertions(+), 190 deletions(-) create mode 100644 torchref/model/context.py diff --git a/AGENTS.md b/AGENTS.md index 1f4ca6d5..ab05f37f 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -186,7 +186,7 @@ Black, 88 columns, `isort` with the black profile. Ruff lint with |---|---| | `base/` | Low-level math and crystallography. `coordinates/` (Cartesian↔fractional), `reciprocal/` (basis, HKL, d-spacing, interpolation, symmetry), `direct_summation/` (F_calc by summation; eager + Triton), `electron_density/` (real-space splatting with CPU/CUDA/MPS kernels, solvent mask, radius policy), `fourier/` (FFT and grids), `scattering/` (form-factor and anomalous tables), `metrics/` (R-factors, binwise scale, loss), `targets/` (the *kernels* behind refinement targets, eager + `triton/`), `french_wilson.py`, `math_torch.py`, `alignment/` | | `io/` | `ReflectionData`, `DatasetCollection`, `FcalcDataset`; MTZ / PDB / CIF / IHM readers and writers; `read_mtz` / `read_pdb` / `read_cif` | -| `model/` | `Model` (refinable atomic parameters), `ModelFT` (adds F_calc via `SfFFT` or `SfDS`), `MixedModel`, `ModelCollection`, and the parametrizations in `parameter_wrappers.py` / `rigid_xyz.py` that decide what is refinable | +| `model/` | `Model` (refinable atomic parameters), `ModelContext` (the cell, space group, atom table, links and provenance a model is loaded with — `model.cell` / `.spacegroup` / `.pdb` forward to it, the rest is `model.ctx.*`), `ModelFT` (adds F_calc via `SfFFT` or `SfDS`), `MixedModel`, `ModelCollection`, and the parametrizations in `parameter_wrappers.py` / `rigid_xyz.py` that decide what is refinable | | `refinement/` | Drivers (`Refinement`, `LBFGSRefinement`, `RigidBodyRefinementStep`), `targets/` (`xray/`, `geometry/`, `adp/`, `collection/`, `combined.py`), `weighting/`, `optimizers/` (annealing, Langevin, preconditioned/seeded L-BFGS), `model_error_estimation/` (σ_A, σ_M), `loss_state.py`, `logger.py` | | `restraints/` | Bonds, angles, torsions, planes, chirals, VDW. Built from the CCP4 Monomer Library, resolved lazily via `get_library_manager()` — importing this package must not trigger a library download | | `scaling/` | `ScalerBase` (model-independent), `Scaler`, `CollectionScaler`, `SolventModel` (k_sol, B_sol) | diff --git a/docs/changelog.rst b/docs/changelog.rst index ccaec3d8..3a5d3272 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -14,6 +14,10 @@ Unreleased - Removed ``ReciprocalSymmetryGrid``, ``ReciprocalSymmetry``, ``expand_reciprocal_grid``, ``expand_reflections`` and ``extract_structure_factors_with_symmetry`` - Removed the unused ``Cell`` gradient plumbing (``requires_grad`` argument and property, ``detach``) and the ``CellTensor`` alias - Removed the ``Symmetry`` alias for ``SpaceGroup``; the name is now a distinct class +- Added ``ModelContext``, holding a model's unit cell, space group, atom table, link records and provenance; ``Model.cell`` / ``.spacegroup`` / ``.pdb`` still work and now read through it +- Moved ``Model``'s configuration and provenance onto the context: ``strip_H``, ``verbose``, ``links``, ``altloc_pairs``, ``initialized``, ``exclude_H_from_sf`` and the input paths are reached as ``model.ctx.*`` +- ``Model.copy`` and ``ModelFT.copy`` now copy the context in one step, cloning the space group instead of sharing it +- Removed ``Model.symmetry``; use ``Model.spacegroup`` Version 0.6.4 diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index 8e3afc8c..1432dfef 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -325,6 +325,8 @@ class TargetDeviceCase: "Map": "needs data + model", "DifferenceMap": "needs two datasets + a model", "LBFGSRefinement": "full pipeline; covered in integration", + "ModelContext": "needs a loaded structure to hold a cell and space group; " + "covered through Model in tests/unit/model/test_model_state_dict_device.py", "_MapSymmetryDirect": "stateless view over its Symmetry: recomputes index grids " "per operation to keep peak memory O(grid), so it owns no tensors to move", "CholeskyMixedTensor": "needs a valid ADP tensor; shares MixedTensor's paths", diff --git a/tests/unit/model/test_create_from_state_dict.py b/tests/unit/model/test_create_from_state_dict.py index 1e8ae822..4fd0231c 100644 --- a/tests/unit/model/test_create_from_state_dict.py +++ b/tests/unit/model/test_create_from_state_dict.py @@ -69,3 +69,33 @@ def test_create_from_state_dict_aniso_u_roundtrip(pdb_dir, cls_name): # At least some atoms are anisotropic → finite, non-trivial u values present. assert torch.isfinite(u_fresh).any() assert torch.allclose(u_fresh, u_restored, equal_nan=True) + + +@pytest.mark.unit +def test_altloc_pairs_survive_state_dict_round_trip(pdb_dir): + """``altloc_pairs`` must reach the state dict and come back. + + It lives on the model's context rather than the model, so a defensive + ``hasattr(self, "altloc_pairs")`` in ``state_dict`` silently substituted an empty + list -- losing the alternative-conformation grouping on every save without + failing anything. + """ + from torchref.model import ModelFT + + cpu = torch.device("cpu") + model = ModelFT() + model.load_pdb(str(pdb_dir / "7L84.pdb")) # carries alternative conformations + model.to(cpu) + + assert model.ctx.altloc_pairs, "fixture should have alternative conformations" + + sd = model.state_dict() + assert sd["altloc_pairs"], "altloc groups must reach the state dict" + + restored = ModelFT.create_from_state_dict(sd, device=cpu, verbose=0) + + assert len(restored.ctx.altloc_pairs) == len(model.ctx.altloc_pairs) + for got, want in zip(restored.ctx.altloc_pairs, model.ctx.altloc_pairs): + assert len(got) == len(want) + for g, w in zip(got, want): + assert torch.equal(g, w) diff --git a/tests/unit/model/test_model.py b/tests/unit/model/test_model.py index 1781a89a..e1b99867 100644 --- a/tests/unit/model/test_model.py +++ b/tests/unit/model/test_model.py @@ -20,7 +20,7 @@ def test_model_empty_initialization(self): model = Model() - assert model.initialized == False + assert model.ctx.initialized is False assert model.pdb is None assert model.xyz is None assert model.adp is None @@ -59,7 +59,7 @@ def test_model_strip_h_default(self): model = Model() - assert model.strip_H == True + assert model.ctx.strip_H is True @pytest.mark.unit def test_model_bool_uninitialized(self): diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index 97a935be..2ffae3e3 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -892,7 +892,7 @@ def _build_metadata(model, data, r_work, r_free): meta.n_atoms_solvent = int((pdb["ATOM"] == "HETATM").sum()) # Geometry deviations - if model.initialized and model._restraints is not None: + if model.ctx.initialized and model._restraints is not None: restraints = model.restraints with torch.no_grad(): if hasattr(restraints, "bond_deviations"): diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 664561f5..c315489b 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -305,8 +305,8 @@ def run( # Get symmetry matrices for clustering sym_matrices = None - if hasattr(self.model, 'symmetry') and self.model.symmetry is not None: - sym_matrices = np.array([s.numpy() for s in self.model.symmetry.matrices]) + if getattr(self.model, 'spacegroup', None) is not None: + sym_matrices = np.array([s.numpy() for s in self.model.spacegroup.matrices]) rotation_peaks = cluster_rotation_peaks( rotation_peaks, diff --git a/torchref/io/metadata.py b/torchref/io/metadata.py index 224ef4a1..b32ad1a5 100644 --- a/torchref/io/metadata.py +++ b/torchref/io/metadata.py @@ -196,7 +196,7 @@ def from_refinement(cls, refinement) -> RefinementMetadata: # --- Geometry deviations (silently skip if no restraints) --- try: model = refinement.model - if model.initialized and model._restraints is not None: + if model.ctx.initialized and model._restraints is not None: restraints = model.restraints if hasattr(restraints, "bond_deviations"): with torch.no_grad(): diff --git a/torchref/model/__init__.py b/torchref/model/__init__.py index 685c2cee..27ad2f14 100644 --- a/torchref/model/__init__.py +++ b/torchref/model/__init__.py @@ -1,6 +1,8 @@ """Atomic models: coordinates, ADPs, occupancies and their structure factors. -:class:`Model` holds the refinable atomic parameters; :class:`ModelFT` adds +:class:`Model` holds the refinable atomic parameters, with the crystallographic +context, atom table and provenance split out into :class:`ModelContext`; +:class:`ModelFT` adds structure-factor calculation on top, through :class:`SfFFT` (FFT) or :class:`SfDS` (direct summation). :class:`MixedModel` combines ModelFT states by population fraction (e.g. dark/light), and :class:`ModelCollection` keys @@ -14,6 +16,7 @@ from torchref.model.sf_fft import SfFFT, FFT from torchref.model.sf_ds import SfDS +from torchref.model.context import ModelContext from torchref.model.mixed_model import MixedModel from torchref.model.model import Model from torchref.model.model_ft import ModelFT @@ -33,6 +36,7 @@ "SfDS", "MixedModel", "Model", + "ModelContext", "ModelFT", "MixedTensor", "PositiveMixedTensor", diff --git a/torchref/model/context.py b/torchref/model/context.py new file mode 100644 index 00000000..85c47ba8 --- /dev/null +++ b/torchref/model/context.py @@ -0,0 +1,122 @@ +"""The information half of a :class:`~torchref.model.model.Model`. + +:class:`ModelContext` holds what a model *is loaded from* and *sits in* -- the unit +cell, the space group, the atom table, the link records and the provenance -- as +opposed to what is being refined, which stays on the model as parameter wrappers and +per-atom buffers. + +Splitting it out means the crystallographic context can be passed to code that needs +only that (structure-factor engines, scalers, most targets) without handing over the +refinable state, and it keeps the model's own surface to parameters and behaviour. + +Mutable by design; prefer :meth:`ModelContext.copy` over editing in place. +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, List, Optional + +from torchref.utils.device_mixin import DeviceMixin + +if TYPE_CHECKING: + import pandas + + from torchref.symmetry import Cell, SpaceGroup + + +@dataclass(eq=False, repr=False) +class ModelContext(DeviceMixin): + """Crystallographic context, atom bookkeeping and provenance for one model. + + Parameters + ---------- + cell : Cell or None + Unit cell, or None before a structure is loaded. + spacegroup : SpaceGroup or None + Space group, or None before a structure is loaded. + pdb : pandas.DataFrame or None + The atom table. Refreshed from the model's tensors only by + ``Model.update_pdb``, so it is stale between refinement steps by design. + links : list or None + Link records from the reader, used to build inter-residue restraints. + altloc_pairs : list + Index groups of alternative conformations, rebuilt by + ``Model.register_alternative_conformations``. + input_file : str or None + Path the structure was loaded from. + cif_path : str or None + Restraint dictionary path, if one was set. + verbose : int, default 1 + Verbosity level. + strip_H : bool, default True + Whether hydrogens were stripped on load. + exclude_H_from_sf : bool, default False + Whether hydrogens are excluded from structure-factor calculation. + initialized : bool, default False + Whether a structure has been loaded. ``if model:`` tests this. + + Notes + ----- + Deliberately does **not** carry the device or float dtype. Those are live + :class:`~torchref.utils.device_mixin.DeviceMixin` trackers that the traversal + rewrites in place on the object that owns the tensors, so they stay on the model + rather than becoming a second source of truth here. + + Holds no refinable parameters, so this is a dataclass rather than an + ``nn.Module``. + """ + + cell: Optional["Cell"] = None + spacegroup: Optional["SpaceGroup"] = None + pdb: Optional["pandas.DataFrame"] = None + links: Optional[List[Any]] = None + altloc_pairs: List[Any] = field(default_factory=list) + input_file: Optional[str] = None + cif_path: Optional[str] = None + verbose: int = 1 + strip_H: bool = True + exclude_H_from_sf: bool = False + initialized: bool = False + + def copy(self) -> "ModelContext": + """An independent copy. + + The atom table is deep-copied and the cell and space group are cloned, so + nothing is shared with the original. Cloning the space group matters now that + it is a mutable dataclass: sharing the reference would let an edit through one + model's context reach every model that was copied from it. + + Returns + ------- + ModelContext + New context sharing no mutable state with this one. + """ + return ModelContext( + cell=self.cell.clone() if self.cell is not None else None, + spacegroup=( + self.spacegroup.copy() if self.spacegroup is not None else None + ), + pdb=self.pdb.copy(deep=True) if self.pdb is not None else None, + links=list(self.links) if self.links is not None else None, + altloc_pairs=[ + tuple(t.clone() for t in group) for group in self.altloc_pairs + ], + input_file=self.input_file, + cif_path=self.cif_path, + verbose=self.verbose, + strip_H=self.strip_H, + exclude_H_from_sf=self.exclude_H_from_sf, + initialized=self.initialized, + ) + + def __repr__(self) -> str: + n_atoms = 0 if self.pdb is None else len(self.pdb) + sg = None if self.spacegroup is None else self.spacegroup.name + return ( + f"ModelContext(spacegroup={sg!r}, n_atoms={n_atoms}, " + f"initialized={self.initialized})" + ) + + +__all__ = ["ModelContext"] diff --git a/torchref/model/model.py b/torchref/model/model.py index 12c6fd58..bd706ec7 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -21,6 +21,7 @@ from torchref.base import math_torch from torchref.config import get_float_dtype, normalize_device from torchref.io import cif, pdb +from torchref.model.context import ModelContext from torchref.model.parameter_wrappers import ( CholeskyMixedTensor, MixedTensor, @@ -69,10 +70,12 @@ class Model(DeviceMovementMixin, DebugMixin, nn.Module): """ Base model class for atomic structure models using PyTorch. - Owns the atomic data -- coordinates, atomic displacement parameters and + Owns the refinable atomic data -- coordinates, atomic displacement parameters and occupancies -- each held in a parameter wrapper that decides which atoms are - refinable. Build it empty (``Model()`` then ``load_pdb`` / ``load_cif`` / - ``load_state_dict``); ``if model:`` tests *initialization*, not existence. + refinable. Everything the structure was *loaded from* rather than refined lives on + :attr:`ctx`, a :class:`~torchref.model.context.ModelContext`. Build the model empty + (``Model()`` then ``load_pdb`` / ``load_cif`` / ``load_state_dict``); ``if model:`` + tests *initialization*, not existence. Parameters ---------- @@ -96,15 +99,20 @@ class Model(DeviceMovementMixin, DebugMixin, nn.Module): positive-definite by construction. Isotropic atoms carry ``U = NaN``. occupancy : OccupancyTensor Atomic occupancies with values in [0, 1]. + ctx : ModelContext + The unit cell, space group, atom table, link records, provenance and + configuration. The fields not forwarded below are reached through it, e.g. + ``model.ctx.strip_H`` and ``model.ctx.initialized``. pdb : pandas.DataFrame - DataFrame containing atomic model data. Only refreshed from the tensors - by :meth:`update_pdb`. + Atom table, forwarded to :attr:`ctx`. Only refreshed from the tensors by + :meth:`update_pdb`. cell : Cell - Unit cell object with parameters [a, b, c, alpha, beta, gamma]. - spacegroup, symmetry : SpaceGroup - Space group object; ``symmetry`` is the same object under its old name. - initialized : bool - Whether the model has been initialized with data. + Unit cell, forwarded to :attr:`ctx`. + spacegroup : SpaceGroup + Space group, forwarded to :attr:`ctx`. + device : torch.device + Where the tensors live. Kept on the model rather than the context because the + device-movement machinery rewrites it in place. """ def __init__( @@ -137,21 +145,15 @@ def __init__( if dtype_float is None: dtype_float = get_float_dtype() device = normalize_device(device) + # ``device`` and ``dtype_float`` stay here rather than moving into the context: + # they are live ``DeviceMixin`` trackers, rewritten in place by the traversal on + # whichever object owns the tensors. self.dtype_float = dtype_float - self.verbose = verbose self.device = device - self.strip_H = strip_H - self._exclude_H_from_sf = False - # State tracking - self.initialized = False - self.altloc_pairs = [] - - # These will be set during load() or load_state_dict() - self.pdb = None - self.links = None - self._cell: Optional[Cell] = None - self._spacegroup: Optional[SpaceGroup] = None + # Everything the model is loaded from and sits in, as opposed to what is + # refined. Populated by load() / create_from_state_dict(). + self.ctx = ModelContext(verbose=verbose, strip_H=strip_H) # Submodules (created during load or load_state_dict) self.xyz = None @@ -164,7 +166,6 @@ def __init__( # Restraints (built lazily on first access) self._restraints = None - self._cif_path = None def __bool__(self): """Return the initialization status when used in boolean context. @@ -173,20 +174,20 @@ def __bool__(self): an uninitialized (but non-``None``) model is falsy. Use ``if model is not None`` when you mean an existence check. """ - return self.initialized + return self.ctx.initialized @property def exclude_H_from_sf(self) -> bool: """Drop H from ``get_iso()`` / ``get_aniso()`` (so from Fcalc) while keeping them in the geometry and VDW restraints. Default False. """ - return self._exclude_H_from_sf + return self.ctx.exclude_H_from_sf @exclude_H_from_sf.setter def exclude_H_from_sf(self, value: bool): - self._exclude_H_from_sf = bool(value) + self.ctx.exclude_H_from_sf = bool(value) # The cached iso/aniso indices encode the H choice, so rebuild them. - if self.initialized and self.pdb is not None: + if self.ctx.initialized and self.pdb is not None: self._rebuild_sf_indices() def _rebuild_sf_indices(self): @@ -194,7 +195,7 @@ def _rebuild_sf_indices(self): iso_mask = ~self.aniso_flag aniso_mask = self.aniso_flag - if self._exclude_H_from_sf and self.pdb is not None: + if self.ctx.exclude_H_from_sf and self.pdb is not None: if not hasattr(self, "_heavy_atom_mask"): h_mask = torch.tensor( (self.pdb["element"].str.strip() != "H").values, @@ -218,20 +219,29 @@ def _rebuild_sf_indices(self): # Cell, SpaceGroup, and Symmetry properties # ========================================================================= + @property + def pdb(self) -> Optional["pandas.DataFrame"]: + """Atom table. Only refreshed from the tensors by :meth:`update_pdb`.""" + return self.ctx.pdb + + @pdb.setter + def pdb(self, value): + self.ctx.pdb = value + @property def cell(self) -> Optional[Cell]: """Unit cell object with parameters [a, b, c, alpha, beta, gamma].""" - return self._cell + return self.ctx.cell @cell.setter def cell(self, value: Cell): """Set the unit cell.""" - self._cell = value + self.ctx.cell = value @property - def spacegroup(self) -> Optional[gemmi.SpaceGroup]: + def spacegroup(self) -> Optional[SpaceGroup]: """Space group object, or None if not set.""" - return self._spacegroup + return self.ctx.spacegroup @spacegroup.setter def spacegroup(self, value): @@ -240,19 +250,9 @@ def spacegroup(self, value): # ``device=self.device``: SpaceGroup falls back to the global # default otherwise, so setting a spacegroup on a CPU-pinned Model # would silently plant accelerator-resident matrices on it. - self._spacegroup = SpaceGroup(value, device=self.device) + self.ctx.spacegroup = SpaceGroup(value, device=self.device) else: - self._spacegroup = None - - @property - def symmetry(self) -> Optional[SpaceGroup]: - """The same object as :attr:`spacegroup`, under its older name.""" - return self._spacegroup - - @symmetry.setter - def symmetry(self, value: Optional[SpaceGroup]): - """Set the space group object directly (no coercion, unlike ``spacegroup``).""" - self._spacegroup = value + self.ctx.spacegroup = None # ========================================================================= # Crystallographic matrix properties (delegated to Cell) @@ -290,7 +290,7 @@ def _build_z_tensor(self) -> torch.Tensor: if hasattr(self, "_Z") and self._Z is not None: return self._Z - if not self.initialized or self.pdb is None: + if not self.ctx.initialized or self.pdb is None: raise RuntimeError( "Cannot build Z tensor: model not initialized. " "Load data first with load_pdb() or load_cif()." @@ -320,13 +320,13 @@ def _build_parametrization(self): if self._parametrization is not None: return self._parametrization - if not self.initialized or self.pdb is None: + if not self.ctx.initialized or self.pdb is None: raise RuntimeError( "Cannot build parametrization: model not initialized. " "Load data first with load_pdb() or load_cif()." ) - if self.verbose > 1: + if self.ctx.verbose > 1: print("Building ITC92 parametrization via table lookup...") from torchref.base.scattering.scattering_table import get_scattering_params_by_z @@ -351,11 +351,11 @@ def _build_parametrization(self): B[idx : idx + 1], ) - if self.verbose > 0: + if self.ctx.verbose > 0: print( f"Parametrization built for {len(self._parametrization)} unique atom types" ) - if self.verbose > 1: + if self.ctx.verbose > 1: print("Elements with parametrization:", list(self._parametrization.keys())) return self._parametrization @@ -425,7 +425,7 @@ def set_restraints_cif(self, cif_path): Model Self, for method chaining. """ - self._cif_path = cif_path + self.ctx.cif_path = cif_path # Reset restraints so they will be rebuilt on next access self._restraints = None return self @@ -437,7 +437,7 @@ def _build_restraints(self): if self._restraints is not None: return self._restraints - if not self.initialized: + if not self.ctx.initialized: raise RuntimeError( "Cannot build restraints: model not initialized. " "Load data first with load_pdb() or load_cif()." @@ -445,19 +445,19 @@ def _build_restraints(self): from torchref.restraints.restraints import RestraintsNew - if self.verbose > 0: + if self.ctx.verbose > 0: print("Building restraints...") self._restraints = RestraintsNew( pdb=self.pdb, - cif_path=self._cif_path, + cif_path=self.ctx.cif_path, xyz_fn=self.xyz, adp_fn=self.adp, vdw_radii_fn=self.get_vdw_radii, - cell=self._cell, - spacegroup=self._spacegroup, - links=self.links, - verbose=self.verbose, + cell=self.ctx.cell, + spacegroup=self.ctx.spacegroup, + links=self.ctx.links, + verbose=self.ctx.verbose, ) return self._restraints @@ -528,7 +528,7 @@ def load(self, reader): ---------- reader : callable Zero-argument callable returning ``(pdb_df, cell, spacegroup)``. An - optional ``.links`` attribute on it is stored on ``self.links``. + optional ``.links`` attribute on it is stored on ``self.ctx.links``. Returns ------- @@ -542,11 +542,11 @@ def load(self, reader): registration and ``initialized = True``. """ self.pdb, cell, spacegroup = reader() - self.links = getattr(reader, "links", None) + self.ctx.links = getattr(reader, "links", None) self.pdb = ( self.pdb.loc[self.pdb["element"] != "H"].reset_index(drop=True) - if self.strip_H + if self.ctx.strip_H else self.pdb ) self.pdb.dropna(subset=["x", "y", "z", "tempfactor", "occupancy"], inplace=True) @@ -604,7 +604,7 @@ def load(self, reader): self.set_default_masks() self.register_alternative_conformations() - self.initialized = True + self.ctx.initialized = True return self def load_pdb(self, file): @@ -621,8 +621,8 @@ def load_pdb(self, file): Model Self, for method chaining. """ - self._input_file = str(file) - reader = pdb.PDBReader(verbose=self.verbose).read(file) + self.ctx.input_file = str(file) + reader = pdb.PDBReader(verbose=self.ctx.verbose).read(file) return self.load(reader) def load_cif(self, file): @@ -639,8 +639,8 @@ def load_cif(self, file): Model Self, for method chaining. """ - self._input_file = str(file) - if self.verbose > 0: + self.ctx.input_file = str(file) + if self.ctx.verbose > 0: print(f"Loading CIF file: {file}") # Read CIF file @@ -798,7 +798,7 @@ def _create_occupancy_groups(self, pdb_df, initial_occ): n_collapsed = len(unique_indices) - if self.verbose > 1: + if self.ctx.verbose > 1: n_groups = n_collapsed n_independent = n_atoms - n_collapsed n_refinable = refinable_mask.sum().item() @@ -901,39 +901,36 @@ def _after_device_apply( """ if getattr(self, "aniso_flag", None) is not None: self._rebuild_sf_indices() - if self.verbose > 0: + if self.ctx.verbose > 0: print(f"Model moved to device: {self.device}") def copy(self): """ Create a deep copy of the Model. - Creates a complete independent copy including all registered buffers, - module parameters, PDB DataFrame, and spacegroup information. + Independent in every part: the context is copied via + :meth:`~torchref.model.context.ModelContext.copy`, buffers are cloned and each + parameter wrapper is copied through its own ``copy`` so its parametrization + survives. Returns ------- Model A new, fully independent Model instance with copied data. """ - if not self.initialized: + if not self.ctx.initialized: raise RuntimeError("Cannot copy an uninitialized Model. Load data first.") model_copy = Model( dtype_float=self.dtype_float, - verbose=self.verbose, + verbose=self.ctx.verbose, device=self.device, - strip_H=self.strip_H, + strip_H=self.ctx.strip_H, ) - model_copy.pdb = self.pdb.copy(deep=True) - - # Setter also sets symmetry; gemmi.SpaceGroup is immutable, so shared. - model_copy.spacegroup = self.spacegroup - model_copy.initialized = True - - if self.cell is not None: - model_copy.cell = self.cell.clone() + # One call carries the atom table, cell, space group, altloc groups and + # provenance, each deep-copied or cloned -- see ``ModelContext.copy``. + model_copy.ctx = self.ctx.copy() for buffer_name, buffer_value in self._buffers.items(): if buffer_value is not None: @@ -945,14 +942,7 @@ def copy(self): if module is not None and hasattr(module, "copy"): setattr(model_copy, module_name, module.copy()) - if hasattr(self, "altloc_pairs") and self.altloc_pairs: - model_copy.altloc_pairs = [ - tuple(tensor.clone() for tensor in group) for group in self.altloc_pairs - ] - else: - model_copy.altloc_pairs = [] - - if self.verbose > 0: + if self.ctx.verbose > 0: print(f"✓ Model copied successfully ({len(model_copy.pdb)} atoms)") return model_copy @@ -1160,7 +1150,7 @@ def set_adp_mode(self, mode: str = "isotropic", aniso_selection: str = None): Run once at model setup, before scaling / restraints / targets. The isotropic result matches a freshly-loaded isotropic-only model. """ - if not getattr(self, "initialized", False) or self.pdb is None: + if not self.ctx.initialized or self.pdb is None: return if mode == "isotropic": aniso_mask = torch.zeros( @@ -1303,7 +1293,7 @@ def update_mask_from_selection( setattr(self, mask_name, updated_mask) - if self.verbose > 0: + if self.ctx.verbose > 0: n_selected = selection_mask.sum().item() n_refinable = updated_mask.sum().item() action = "frozen" if freeze else "unfrozen" @@ -1347,7 +1337,7 @@ def apply_mask_to_parameter(self, target: str): f"Invalid target: '{target}'. Must be 'xyz', 'adp', 'u', or 'occupancy'" ) - if self.verbose > 0: + if self.ctx.verbose > 0: n_refinable = getattr(self, f"{target}_mask").sum().item() print(f" Applied mask to {target}: {n_refinable} atoms refinable") @@ -1536,14 +1526,14 @@ def print_parameters_info(self): def register_alternative_conformations(self): """ - Rebuild ``self.altloc_pairs`` from the ``altloc`` column. + Rebuild ``self.ctx.altloc_pairs`` from the ``altloc`` column. One tuple per residue that has multiple conformations, holding one index tensor per conformation (in sorted altloc order), e.g. ``[(tensor([100, 101]), tensor([110, 111])), ...]``. Overwrites any previous content, so call it after the atom numbering changes. """ - self.altloc_pairs = [] + self.ctx.altloc_pairs = [] pdb_with_altlocs = self.pdb[self.pdb["altloc"] != ""] @@ -1565,7 +1555,7 @@ def register_alternative_conformations(self): ) conformation_tensors.append(indices) - self.altloc_pairs.append(tuple(conformation_tensors)) + self.ctx.altloc_pairs.append(tuple(conformation_tensors)) def shake_coords(self, stddev: float): """ @@ -1655,7 +1645,7 @@ def generate_hydrogens(self, mon_lib_path: str = None) -> "Model": # normal path — a CCP4 install is not required. from torchref.restraints.library import get_library_manager - mgr = get_library_manager(verbose=self.verbose) + mgr = get_library_manager(verbose=self.ctx.verbose) mon_lib_path = str(mgr.ensure_gemmi_base()) # gemmi reads from a file, so the live tensors must reach the DataFrame. @@ -1705,7 +1695,7 @@ def generate_hydrogens(self, mon_lib_path: str = None) -> "Model": # strip_H=False, or the hydrogens we just placed would be dropped again. new_model = self.__class__( dtype_float=self.dtype_float, - verbose=self.verbose, + verbose=self.ctx.verbose, device=self.device, strip_H=False, ) @@ -1724,7 +1714,7 @@ def _new_model_from_df(self, df, *, strip_H=None): """Build a fresh model of the same class from a DataFrame.""" import inspect - sh = self.strip_H if strip_H is None else strip_H + sh = self.ctx.strip_H if strip_H is None else strip_H ctor_kw = dict( dtype_float=self.dtype_float, verbose=0, @@ -1748,8 +1738,8 @@ def _new_model_from_df(self, df, *, strip_H=None): if hasattr(new_model, "setup_grid"): new_model.setup_grid() # Propagate CIF restraint paths so restraints are rebuilt correctly - if self._cif_path is not None: - new_model._cif_path = self._cif_path + if self.ctx.cif_path is not None: + new_model._cif_path = self.ctx.cif_path return new_model def strip_altlocs(self) -> "Model": @@ -2456,19 +2446,15 @@ def state_dict(self, destination=None, prefix="", keep_vars=False): destination=destination, prefix=prefix, keep_vars=keep_vars ) - state[prefix + "pdb"] = ( - self.pdb.copy() if hasattr(self, "pdb") and self.pdb is not None else None - ) + state[prefix + "pdb"] = self.pdb.copy() if self.pdb is not None else None state[prefix + "cell"] = self.cell.data.cpu() if self.cell is not None else None # As a string: gemmi.SpaceGroup is not picklable. state[prefix + "spacegroup"] = self.spacegroup.xhm if self.spacegroup else None - state[prefix + "initialized"] = self.initialized + state[prefix + "initialized"] = self.ctx.initialized state[prefix + "dtype_float"] = self.dtype_float state[prefix + "device"] = self.device - state[prefix + "strip_H"] = self.strip_H - state[prefix + "altloc_pairs"] = ( - self.altloc_pairs if hasattr(self, "altloc_pairs") else [] - ) + state[prefix + "strip_H"] = self.ctx.strip_H + state[prefix + "altloc_pairs"] = self.ctx.altloc_pairs return state @@ -2482,7 +2468,7 @@ def save_state(self, path: str): Path to save the state dictionary to. """ torch.save(self.state_dict(), path) - if self.verbose > 0: + if self.ctx.verbose > 0: print(f"Saved model state to {path}") def load_state(self, path: str, strict: bool = True): @@ -2499,11 +2485,11 @@ def load_state(self, path: str, strict: bool = True): """ state_dict = torch.load(path, map_location=self.device, weights_only=False) loaded = type(self).create_from_state_dict( - state_dict, device=self.device, verbose=self.verbose + state_dict, device=self.device, verbose=self.ctx.verbose ) # Adopt the fully-built model's state wholesale. self.__dict__.update(loaded.__dict__) - if self.verbose > 0: + if self.ctx.verbose > 0: print(f"Loaded model state from {path}") @classmethod @@ -2561,8 +2547,8 @@ def create_from_state_dict( ) instance.pdb = pdb - instance.initialized = initialized - instance.altloc_pairs = altloc_pairs + instance.ctx.initialized = initialized + instance.ctx.altloc_pairs = altloc_pairs # Setter also sets symmetry. instance.spacegroup = spacegroup @@ -2707,7 +2693,7 @@ def get_selection_mask(self, selection: str) -> torch.Tensor: """ from torchref.utils.utils import parse_phenix_selection - if not self.initialized: + if not self.ctx.initialized: raise RuntimeError( "Cannot get selection mask from an uninitialized Model. Load data first." ) @@ -2756,7 +2742,7 @@ def select(self, selection: str) -> "Model": """ from torchref.utils.utils import parse_phenix_selection - if not self.initialized: + if not self.ctx.initialized: raise RuntimeError( "Cannot select from an uninitialized Model. Load data first." ) @@ -2772,9 +2758,9 @@ def select(self, selection: str) -> "Model": # type(self), so a subclass returns its own type. selected_model = type(self)( dtype_float=self.dtype_float, - verbose=self.verbose, + verbose=self.ctx.verbose, device=self.device, - strip_H=self.strip_H, + strip_H=self.ctx.strip_H, ) # ``index`` must be renumbered: the occupancy grouping below reads it. @@ -2783,7 +2769,7 @@ def select(self, selection: str) -> "Model": selected_model.pdb = selected_model.pdb.reset_index(drop=True) selected_model.pdb["index"] = selected_model.pdb.index.to_numpy(dtype=int) - # Setter also sets symmetry; gemmi.SpaceGroup is immutable, so shared. + # The setter rebuilds a SpaceGroup, so the selection gets its own. selected_model.spacegroup = self.spacegroup # The fractional / reciprocal matrices are properties over the Cell, so @@ -2845,9 +2831,9 @@ def select(self, selection: str) -> "Model": selected_model.set_default_masks() selected_model.register_alternative_conformations() - selected_model.initialized = True + selected_model.ctx.initialized = True - if self.verbose > 0: + if self.ctx.verbose > 0: print(f"Selected {n_selected}/{len(self.pdb)} atoms with '{selection}'") return selected_model @@ -2864,7 +2850,7 @@ def xyz_fractional(self) -> torch.Tensor: torch.Tensor Tensor of shape (n_atoms, 3) with fractional coordinates. """ - if not self.initialized: + if not self.ctx.initialized: raise RuntimeError( "Model must be initialized to compute fractional coordinates." ) @@ -2901,7 +2887,7 @@ def rotate( Model Self, for method chaining. """ - if not self.initialized: + if not self.ctx.initialized: raise RuntimeError("Model must be initialized to apply rotation.") xyz = self.xyz() @@ -2945,7 +2931,7 @@ def translate(self, translation: torch.Tensor, fractional: bool = False) -> "Mod model.translate(torch.tensor([5.0, 0.0, 0.0])) # 5 Å in x model.translate(torch.tensor([0.5, 0.5, 0.5]), fractional=True) # half cell """ - if not self.initialized: + if not self.ctx.initialized: raise RuntimeError("Model must be initialized to apply translation.") xyz = self.xyz() @@ -2974,7 +2960,7 @@ def get_centroid(self) -> torch.Tensor: torch.Tensor Centroid coordinates with shape (3,). """ - if not self.initialized: + if not self.ctx.initialized: raise RuntimeError("Model must be initialized to compute centroid.") return self.xyz().mean(dim=0) @@ -2999,7 +2985,7 @@ def use_rigid_xyz(self) -> "Model": """ from torchref.model.rigid_xyz import RigidXYZTensor - if not self.initialized: + if not self.ctx.initialized: raise RuntimeError( "Model must be initialized before use_rigid_xyz(). " "Load data first with load_pdb() or load_cif()." @@ -3036,7 +3022,7 @@ def use_rigid_xyz(self) -> "Model": "Polymer filter removed every atom — cannot build rigid bodies." ) mobile_mask = torch.from_numpy(mobile_arr).to(device=self.device) - if self.verbose > 0 and int(drop.sum()) > 0: + if self.ctx.verbose > 0 and int(drop.sum()) > 0: n_water = int(is_water.sum()) n_ion = int((is_single_atom & ~is_std & ~is_water).sum()) print( @@ -3080,7 +3066,7 @@ def use_rigid_xyz(self) -> "Model": if hasattr(self, "reset_cache"): self.reset_cache() - if self.verbose > 0: + if self.ctx.verbose > 0: print( f"Switched to rigid-body parametrization: {rigid_xyz} " f"({rigid_xyz.n_chains} chain(s))" diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index 60434fa0..fad61b2c 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -126,18 +126,18 @@ def __init__( @property def cell(self): """Unit cell object with parameters [a, b, c, alpha, beta, gamma].""" - return self._cell + return self.ctx.cell @cell.setter def cell(self, value): """Set the unit cell; also builds the FFT once the spacegroup is set.""" - self._cell = value + self.ctx.cell = value self._maybe_initialize_fft() @property def spacegroup(self): """Space group object.""" - return self._spacegroup + return self.ctx.spacegroup @spacegroup.setter def spacegroup(self, value): @@ -145,19 +145,19 @@ def spacegroup(self, value): also builds the FFT once the cell is set. """ if value is not None: - self._spacegroup = SpaceGroup( + self.ctx.spacegroup = SpaceGroup( value, dtype=self.dtype_float, device=self.device ) else: - self._spacegroup = None + self.ctx.spacegroup = None self._maybe_initialize_fft() def _maybe_initialize_fft(self): """(Re)build the SfFFT submodule once both cell and spacegroup are set.""" - if self._cell is not None and self._spacegroup is not None: + if self.ctx.cell is not None and self.ctx.spacegroup is not None: self._fft = SfFFT( - cell=self._cell, - spacegroup=self._spacegroup, + cell=self.ctx.cell, + spacegroup=self.ctx.spacegroup, device=self.device, max_res=self.max_res, ) @@ -253,7 +253,7 @@ def setup_gridsize(self, max_res=None): self.max_res = max_res self._fft.max_res = max_res - if self.verbose > 1: + if self.ctx.verbose > 1: print(f"Defining grid size for max_res={self.max_res} Å") gridsize = self.cell.compute_grid_size(self.max_res) @@ -384,7 +384,7 @@ def setup_grid(self, max_res=None, gridsize=None): self.max_res = max_res self._fft.max_res = max_res - if self.verbose > 1: + if self.ctx.verbose > 1: print(f"Setting up grids with max_res={self.max_res} Å") gridsize_to_use = gridsize or self._explicit_gridsize @@ -394,7 +394,7 @@ def setup_grid(self, max_res=None, gridsize=None): max_res=self.max_res, ) - if self.verbose > 2: + if self.ctx.verbose > 2: print(f"Grid shape: {self._fft.real_space_grid.shape[:-1]}") print(f"Voxel size: {self._fft.voxel_size}") @@ -423,7 +423,7 @@ def get_radius(self, min_radius_Angstrom: float = 4.0): .to(dtypes.int) .item() ) - if self.verbose > 1: + if self.ctx.verbose > 1: print( f"Calculated radius for density calculation: {min_radius} voxels (voxel size: {voxel_size}), this corresponds to at least {min_radius_Angstrom} Å" ) @@ -453,7 +453,7 @@ def build_complete_map(self, radius=None, apply_symmetry=True): """ self.map = self.build_initial_map(apply_symmetry=apply_symmetry) - if self.verbose > 2: + if self.ctx.verbose > 2: print( f"Density map built. Sum: {self.map.sum():.2f}, Max: {self.map.max():.4f}" ) @@ -478,12 +478,12 @@ def build_initial_map(self, apply_symmetry=True): if self._fft.real_space_grid is None: self.setup_grid() - if self.verbose > 2: + if self.ctx.verbose > 2: print("Building density map (per-atom variable radius)...") xyz_iso, adp_iso, occ_iso, A_iso, B_iso = self.get_iso() - if self.verbose > 3: + if self.ctx.verbose > 3: assert torch.all( torch.isfinite(A_iso) ), "Non-finite values found in A_iso during map building." @@ -516,7 +516,7 @@ def build_initial_map(self, apply_symmetry=True): apply_symmetry=apply_symmetry, ) - if self.verbose > 3: + if self.ctx.verbose > 3: assert torch.all( torch.isfinite(self.map) ), "Non-finite values found in map." @@ -545,7 +545,7 @@ def save_map(self, filename): np_map = self.map.detach().cpu().numpy().astype(np.float32) cell = self.cell.tolist() - if self.verbose > 0: + if self.ctx.verbose > 0: print(f"Saving map to {filename}") print(f" Map shape: {self.map.shape}") print(f" Map sum: {self.map.sum():.2f}") @@ -558,7 +558,7 @@ def save_map(self, filename): map_ccp.setup(0.0) map_ccp.update_ccp4_header() map_ccp.write_ccp4_map(filename) - if self.verbose > 0: + if self.ctx.verbose > 0: print("Map saved successfully") def get_map_statistics(self): @@ -628,7 +628,7 @@ def _get_anomalous_cache( unique_elements, self.wavelength, self.anomalous_threshold ) - if self.verbose > 1 and significant: + if self.ctx.verbose > 1 and significant: print( f"Anomalous scatterers at {self.wavelength:.4f} Å: " f"{list(significant.keys())}" @@ -807,7 +807,7 @@ def forward(self, hkl, apply_anomalous: bool = True) -> torch.Tensor: sf, hkl, include_fdp=bool(self.anomalous_bijvoet) ) - if self.verbose > 2: + if self.ctx.verbose > 2: assert torch.all( torch.isfinite(sf) ), "Non-finite values found while calculating fcalc." @@ -834,31 +834,22 @@ def copy(self, detach: bool = True) -> "ModelFT": ModelFT A new, fully independent ModelFT instance with copied data. """ - if not self.initialized: + if not self.ctx.initialized: raise RuntimeError("Cannot copy an uninitialized ModelFT. Load data first.") model_copy = ModelFT( dtype_float=self.dtype_float, - verbose=self.verbose, + verbose=self.ctx.verbose, device=self.device, - strip_H=self.strip_H, + strip_H=self.ctx.strip_H, max_res=self.max_res, gridsize=self._explicit_gridsize, wavelength=self.wavelength, anomalous_threshold=self.anomalous_threshold, ) - model_copy.pdb = self.pdb.copy(deep=True) - - if self._spacegroup is not None: - model_copy._spacegroup = self._spacegroup.copy() - else: - model_copy._spacegroup = None - - model_copy.initialized = True - - if self.cell is not None: - model_copy.cell = self.cell.clone() + # Carries the atom table, cell, space group, altloc groups and provenance. + model_copy.ctx = self.ctx.copy() # Own buffers only; the FFT submodule's are handled by its copy() below. for buffer_name, buffer_value in self._buffers.items(): @@ -870,22 +861,14 @@ def copy(self, detach: bool = True) -> "ModelFT": else: model_copy.register_buffer(buffer_name, buffer_value.clone()) - # Parameter wrappers via their own .copy(); _fft / _spacegroup are separate. - skip_modules = {"_fft", "_spacegroup", "spacegroup", "_symmetry", "symmetry"} + # Parameter wrappers via their own .copy(); the FFT submodule is separate. + skip_modules = {"_fft"} for module_name, module in self._modules.items(): if module_name in skip_modules: continue if module is not None and hasattr(module, "copy"): setattr(model_copy, module_name, module.copy()) - # Copy alternative conformation pairs - if hasattr(self, "altloc_pairs") and self.altloc_pairs: - model_copy.altloc_pairs = [ - tuple(tensor.clone() for tensor in group) for group in self.altloc_pairs - ] - else: - model_copy.altloc_pairs = [] - if hasattr(self, "_parametrization") and self._parametrization is not None: import copy as copy_module @@ -899,7 +882,7 @@ def copy(self, detach: bool = True) -> "ModelFT": # Don't share cached structure factors with the original. model_copy.reset_cache() - if self.verbose > 0: + if self.ctx.verbose > 0: print(f"✓ ModelFT copied successfully ({len(model_copy.pdb)} atoms)") return model_copy @@ -1014,8 +997,8 @@ def create_from_state_dict( ) instance.pdb = pdb - instance.initialized = initialized - instance.altloc_pairs = altloc_pairs + instance.ctx.initialized = initialized + instance.ctx.altloc_pairs = altloc_pairs # Setter also sets symmetry; the cell setter below then builds the FFT. instance.spacegroup = spacegroup_str diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index 4eaf78e2..a6420b06 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -851,8 +851,8 @@ def collect_deposition_metadata(self, metadata=None): return metadata.merge(refinement_meta) # Merge with pass-through headers from input file - if hasattr(self.model, "_input_file") and self.model._input_file: - input_file = self.model._input_file + if self.model.ctx.input_file: + input_file = self.model.ctx.input_file if input_file.endswith(".pdb"): input_meta = RefinementMetadata.from_pdb_file(input_file) elif input_file.endswith((".cif", ".mmcif")): @@ -1015,7 +1015,7 @@ def extract_submodule_state(state_dict: dict, prefix: str) -> dict: instance.scaler.set_model_and_data(instance.model, instance.reflection_data) # Initialize targets if model is available - if instance.model is not None and instance.model.initialized: + if instance.model is not None and instance.model.ctx.initialized: try: instance._init_targets() except Exception as e: diff --git a/torchref/refinement/targets/geometry/non_bonded.py b/torchref/refinement/targets/geometry/non_bonded.py index 5d56fbe4..ae76b196 100644 --- a/torchref/refinement/targets/geometry/non_bonded.py +++ b/torchref/refinement/targets/geometry/non_bonded.py @@ -241,7 +241,7 @@ def _compute_positions( return pos1, pos2, min_distances cell = self.model.cell - sg = self.model.symmetry + sg = self.model.spacegroup mate_source = xyz[indices[:, 1]] # (N_pairs, 3) -- gradients flow frac = cell.cartesian_to_fractional(mate_source) @@ -285,8 +285,8 @@ def forward(self) -> torch.Tensor: vdw_data["min_distances"], vdw_data.get("symop_indices"), vdw_data.get("cell_offsets"), - self.model.symmetry.matrices, - self.model.symmetry.translations, + self.model.spacegroup.matrices, + self.model.spacegroup.translations, self.model.cell.fractional_matrix, self.model.cell.inv_fractional_matrix, self._c_rep, self._r_exp, diff --git a/torchref/refinement/targets/geometry/non_bonded_h.py b/torchref/refinement/targets/geometry/non_bonded_h.py index 39f6a1d8..a91ceddc 100644 --- a/torchref/refinement/targets/geometry/non_bonded_h.py +++ b/torchref/refinement/targets/geometry/non_bonded_h.py @@ -113,8 +113,8 @@ def _compute_h_vdw_loss( return nonbonded_heavy_math( xyz_all, indices, h_topo.cand_min_dist, h_topo.cand_symop_idx, h_topo.cand_cell_offset, - self.model.symmetry.matrices, - self.model.symmetry.translations, + self.model.spacegroup.matrices, + self.model.spacegroup.translations, self.model.cell.fractional_matrix, self.model.cell.inv_fractional_matrix, self._c_rep, self._r_exp, @@ -134,7 +134,7 @@ def _compute_h_vdw_loss( if n_sym > 0: cell = self.model.cell - sg = self.model.symmetry + sg = self.model.spacegroup sym_source = xyz_all[h_topo.cand_idx_j[n_asu:]] frac = cell.cartesian_to_fractional(sym_source) R = sg.matrices[h_topo.cand_symop_idx[n_asu:]].to(frac.dtype) From 261d71ba66436efbb593cb7f4f1b568363832164 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Wed, 26 Aug 2026 14:06:41 +0200 Subject: [PATCH 056/250] Make HydrogenTopology a dataclass with optional fields The class was an nn.Module holding no parameters, and none of its 29 buffers existed at construction: build_hydrogen_topology and build_h_candidate_pairs attached them afterwards with register_buffer. That is why reading it needed hasattr guards, and why n_asu_candidates was fetched through a getattr default. As a dataclass the fields are declared with None defaults, the builders assign them directly, and the guards become plain None checks. The getattr default in the non-bonded H target is gone too: the attribute now always exists, so a default that can never fire would only mislead. The four derived placement tensors move behind reset_cache, which DeviceMixin calls on every .to(). Nothing cleared them before, so a device move left the clamped neighbour indices and the bond-length column referring to tensors on the old device. HydrogenTopology also graduates from the device-conformance UNCOVERED list to a real case, since an empty topology is now constructible. It is registered tensor_free: a fresh one is a bare shell whose tracker is all there is to check. Co-Authored-By: Claude Opus 5 (1M context) --- docs/changelog.rst | 1 + tests/helpers/device_cases.py | 12 +- .../targets/geometry/non_bonded_h.py | 2 +- torchref/restraints/hydrogen_topology.py | 259 +++++++++--------- 4 files changed, 143 insertions(+), 131 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 3a5d3272..b58b20b9 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -18,6 +18,7 @@ Unreleased - Moved ``Model``'s configuration and provenance onto the context: ``strip_H``, ``verbose``, ``links``, ``altloc_pairs``, ``initialized``, ``exclude_H_from_sf`` and the input paths are reached as ``model.ctx.*`` - ``Model.copy`` and ``ModelFT.copy`` now copy the context in one step, cloning the space group instead of sharing it - Removed ``Model.symmetry``; use ``Model.spacegroup`` +- ``HydrogenTopology`` is a dataclass with optional fields instead of an ``nn.Module`` whose buffers were attached after construction Version 0.6.4 diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index 1432dfef..72aae9c2 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -204,6 +204,17 @@ def _cell(device): _symmetry, "Symmetry", ), + DeviceCase( + "HydrogenTopology_empty", + lambda d: __import__( + "torchref.restraints.hydrogen_topology", + fromlist=["HydrogenTopology"], + ).HydrogenTopology(device=d), + "HydrogenTopology", + # The builders attach every tensor later, so a fresh topology is a bare shell + # and only its tracker can be checked. + tensor_free=True, + ), DeviceCase( "_MapSymmetryInterpolation", _map_symmetry_interpolation, @@ -318,7 +329,6 @@ class TargetDeviceCase: "_SharedMixedModel": "internal view owned by ModelCollection", "Scaler": "needs a loaded model + data; covered in integration", "RestraintsNew": "needs a model + monomer library", - "HydrogenTopology": "needs a built restraint topology", "FrenchWilson": "needs loaded intensities", "DatasetCollection": "needs several loaded datasets", "FcalcDataset": "needs computed structure factors", diff --git a/torchref/refinement/targets/geometry/non_bonded_h.py b/torchref/refinement/targets/geometry/non_bonded_h.py index a91ceddc..e3cfae41 100644 --- a/torchref/refinement/targets/geometry/non_bonded_h.py +++ b/torchref/refinement/targets/geometry/non_bonded_h.py @@ -123,7 +123,7 @@ def _compute_h_vdw_loss( # Slow path: gaussian / soft modes, inline eager. pos_i = xyz_all[h_topo.cand_idx_i] - n_asu = getattr(h_topo, 'n_asu_candidates', n_cand) + n_asu = h_topo.n_asu_candidates n_sym = n_cand - n_asu min_dist = h_topo.cand_min_dist diff --git a/torchref/restraints/hydrogen_topology.py b/torchref/restraints/hydrogen_topology.py index 6dd2632b..6a92b04b 100644 --- a/torchref/restraints/hydrogen_topology.py +++ b/torchref/restraints/hydrogen_topology.py @@ -11,11 +11,11 @@ the H positions back to the heavy-atom coordinates via standard autograd. """ +from dataclasses import dataclass, field from typing import Dict, Optional import numpy as np import torch -from torch import nn from torchref.config import dtypes, normalize_device from torchref.utils.device_resolution import resolve_device @@ -48,60 +48,114 @@ # --------------------------------------------------------------------------- -class HydrogenTopology(DeviceMixin, nn.Module): +@dataclass(eq=False, repr=False) +class HydrogenTopology(DeviceMixin): """Static topology describing riding hydrogens for VDW evaluation. - All data are stored as registered buffers so they move automatically - with ``.to(device)`` and appear in ``state_dict``. + Every tensor field starts as ``None``; :func:`build_hydrogen_topology` and + :func:`build_h_candidate_pairs` fill them, independently -- a topology can carry + hydrogens and no candidate pairs. Test :attr:`n_hydrogens` and + :attr:`has_candidates` rather than the fields. + + Holds no refinable parameters, so this is a dataclass rather than an + ``nn.Module``; ``DeviceMixin`` still moves every tensor with ``.to(device)``. + + Parameters + ---------- + device : torch.device, optional + Where the builders should allocate. Tracked from construction so a + ``resolve_device(h_topo, ...)`` called before anything is attached still + answers truthfully. Attributes ---------- - h_parent_idx : (N_h,) long - Index into heavy-atom array for each riding H. - h_bond_length : (N_h,) float - Ideal H–parent bond length (Å). - h_vdw_radius : (N_h,) float - Van der Waals radius for each H (1.20 Å). - h_placement_type : (N_h,) long - Placement-geometry enum (see module-level constants). - h_slot_in_parent : (N_h,) long - Ordinal within sibling H atoms on the same parent (0, 1, 2). - parent_neighbor_idx : (N_h, MAX_HEAVY_NB) long - Heavy-atom neighbour indices of the parent (-1 = padding). - parent_neighbor_count : (N_h,) long - Actual number of heavy-atom neighbours for the parent. - h_chainid_enc : (N_h,) long - Encoded chain ID (for same-residue filtering). - h_resseq : (N_h,) long - Residue sequence number (for same-residue filtering). - - Notes - ----- - The ``cand_*``/``n_asu_candidates``/``type_bounds`` buffers are added later by - ``build_h_candidate_pairs``, not by ``build_hydrogen_topology`` -- test - :attr:`has_candidates` before touching them. + h_parent_idx : torch.Tensor + Heavy-atom index of each riding H's parent, ``(N_h,)`` long. + h_bond_length : torch.Tensor + Ideal H-parent bond length in Angstroms, ``(N_h,)``. + h_vdw_radius : torch.Tensor + Van der Waals radius per H (1.20 A), ``(N_h,)``. + h_placement_type : torch.Tensor + Placement-geometry enum, ``(N_h,)`` long; see the module-level constants. + h_slot_in_parent : torch.Tensor + Ordinal among sibling H atoms on the same parent (0, 1, 2), ``(N_h,)`` long. + parent_neighbor_idx : torch.Tensor + Heavy-atom neighbours of the parent, ``(N_h, MAX_HEAVY_NB)`` long, ``-1`` + padded. + parent_neighbor_count : torch.Tensor + Heavy-atom neighbour count per parent, ``(N_h,)`` long. + h_chainid_enc : torch.Tensor + Encoded chain ID, ``(N_h,)`` long, for same-residue filtering. + h_resseq : torch.Tensor + Residue sequence number, ``(N_h,)`` long, for same-residue filtering. + type_bounds : dict + ``{placement_type: (start, end)}`` bounds into the type-sorted arrays. + cand_idx_i, cand_idx_j, cand_symop_idx, cand_cell_offset : torch.Tensor + Precomputed H candidate pairs, sorted so the asymmetric-unit ones come first. + cand_min_dist : torch.Tensor + Per-pair minimum-distance scratch buffer, ``(P,)``. + n_asu_candidates : int + How many leading candidate pairs lie inside the asymmetric unit. """ - def __init__(self, device=None): - super().__init__() - # Buffers are registered later by build_hydrogen_topology(), so this - # object is tensor-free at construction. The tracker still has to exist: - # ``DeviceMixin._refresh_device_trackers`` only maintains attributes - # already present in ``__dict__``, and callers reconcile against - # ``h_topo.device`` before attaching buffers to it. - self.device = normalize_device(device) + device: Optional[torch.device] = None + + h_parent_idx: Optional[torch.Tensor] = None + h_bond_length: Optional[torch.Tensor] = None + h_vdw_radius: Optional[torch.Tensor] = None + h_placement_type: Optional[torch.Tensor] = None + h_slot_in_parent: Optional[torch.Tensor] = None + parent_neighbor_idx: Optional[torch.Tensor] = None + parent_neighbor_count: Optional[torch.Tensor] = None + h_chainid_enc: Optional[torch.Tensor] = None + h_resseq: Optional[torch.Tensor] = None + type_bounds: Dict[int, tuple] = field(default_factory=dict) + + cand_idx_i: Optional[torch.Tensor] = None + cand_idx_j: Optional[torch.Tensor] = None + cand_symop_idx: Optional[torch.Tensor] = None + cand_cell_offset: Optional[torch.Tensor] = None + cand_min_dist: Optional[torch.Tensor] = None + n_asu_candidates: int = 0 + + # Derived at first placement and reused across steps; see reset_cache. + _dir_coeffs: Optional[torch.Tensor] = field(default=None, repr=False) + _nb_idx_clamped: Optional[torch.Tensor] = field(default=None, repr=False) + _nb_valid: Optional[torch.Tensor] = field(default=None, repr=False) + _bond_len_col: Optional[torch.Tensor] = field(default=None, repr=False) + + def __post_init__(self) -> None: + """Resolve the device tracker the builders allocate against.""" + self.device = normalize_device(self.device) @property def n_hydrogens(self) -> int: - """Number of riding hydrogens, or 0 before the buffers are attached.""" - if hasattr(self, "h_parent_idx"): - return self.h_parent_idx.shape[0] - return 0 + """Number of riding hydrogens, or 0 before the builders have run.""" + if self.h_parent_idx is None: + return 0 + return int(self.h_parent_idx.shape[0]) @property def has_candidates(self) -> bool: """Whether precomputed H candidate pairs are available.""" - return hasattr(self, "cand_idx_i") and self.cand_idx_i.shape[0] > 0 + return self.cand_idx_i is not None and self.cand_idx_i.shape[0] > 0 + + def reset_cache(self) -> None: + """Drop the derived placement tensors; rebuilt on the next placement call. + + Called by ``DeviceMixin`` on every ``.to()``, which is what keeps the clamped + neighbour indices and bond-length column from surviving a device move. + """ + self._dir_coeffs = None + self._nb_idx_clamped = None + self._nb_valid = None + self._bond_len_col = None + + def __repr__(self) -> str: + return ( + f"HydrogenTopology(n_hydrogens={self.n_hydrogens}, " + f"has_candidates={self.has_candidates})" + ) # --------------------------------------------------------------------------- @@ -369,34 +423,17 @@ def build_hydrogen_topology( fdtype = dtypes.float if n_h_total == 0: - topo.register_buffer( - "h_parent_idx", torch.zeros(0, dtype=torch.long, device=device) - ) - topo.register_buffer( - "h_bond_length", torch.zeros(0, dtype=fdtype, device=device) - ) - topo.register_buffer( - "h_vdw_radius", torch.zeros(0, dtype=fdtype, device=device) - ) - topo.register_buffer( - "h_placement_type", torch.zeros(0, dtype=torch.long, device=device) - ) - topo.register_buffer( - "h_slot_in_parent", torch.zeros(0, dtype=torch.long, device=device) - ) - topo.register_buffer( - "parent_neighbor_idx", - torch.zeros(0, MAX_HEAVY_NB, dtype=torch.long, device=device), - ) - topo.register_buffer( - "parent_neighbor_count", torch.zeros(0, dtype=torch.long, device=device) - ) - topo.register_buffer( - "h_chainid_enc", torch.zeros(0, dtype=torch.long, device=device) - ) - topo.register_buffer( - "h_resseq", torch.zeros(0, dtype=torch.long, device=device) + topo.h_parent_idx = torch.zeros(0, dtype=torch.long, device=device) + topo.h_bond_length = torch.zeros(0, dtype=fdtype, device=device) + topo.h_vdw_radius = torch.zeros(0, dtype=fdtype, device=device) + topo.h_placement_type = torch.zeros(0, dtype=torch.long, device=device) + topo.h_slot_in_parent = torch.zeros(0, dtype=torch.long, device=device) + topo.parent_neighbor_idx = torch.zeros( + 0, MAX_HEAVY_NB, dtype=torch.long, device=device ) + topo.parent_neighbor_count = torch.zeros(0, dtype=torch.long, device=device) + topo.h_chainid_enc = torch.zeros(0, dtype=torch.long, device=device) + topo.h_resseq = torch.zeros(0, dtype=torch.long, device=device) return topo # Sort all topology arrays by placement type for contiguous slicing @@ -421,43 +458,22 @@ def build_hydrogen_topology( idxs = np.where(mask)[0] type_bounds[t] = (int(idxs[0]), int(idxs[-1]) + 1) - topo.register_buffer( - "h_parent_idx", - torch.tensor(acc_parent_idx, dtype=torch.long, device=device), + topo.h_parent_idx = torch.tensor(acc_parent_idx, dtype=torch.long, device=device) + topo.h_bond_length = torch.tensor(acc_bond_length, dtype=fdtype, device=device) + topo.h_vdw_radius = torch.full((n_h_total,), 1.20, dtype=fdtype, device=device) + topo.h_placement_type = torch.tensor( + acc_placement_type, dtype=torch.long, device=device ) - topo.register_buffer( - "h_bond_length", - torch.tensor(acc_bond_length, dtype=fdtype, device=device), + topo.h_slot_in_parent = torch.tensor(acc_slot, dtype=torch.long, device=device) + topo.parent_neighbor_idx = torch.tensor( + np.stack(acc_nb_idx), dtype=torch.long, device=device ) - topo.register_buffer( - "h_vdw_radius", - torch.full((n_h_total,), 1.20, dtype=fdtype, device=device), - ) - topo.register_buffer( - "h_placement_type", - torch.tensor(acc_placement_type, dtype=torch.long, device=device), - ) - topo.register_buffer( - "h_slot_in_parent", - torch.tensor(acc_slot, dtype=torch.long, device=device), - ) - topo.register_buffer( - "parent_neighbor_idx", - torch.tensor(np.stack(acc_nb_idx), dtype=torch.long, device=device), - ) - topo.register_buffer( - "parent_neighbor_count", - torch.tensor(acc_nb_count, dtype=torch.long, device=device), + topo.parent_neighbor_count = torch.tensor( + acc_nb_count, dtype=torch.long, device=device ) topo.type_bounds = type_bounds # dict: type_code -> (start, end) - topo.register_buffer( - "h_chainid_enc", - torch.tensor(acc_chainid_enc, dtype=torch.long, device=device), - ) - topo.register_buffer( - "h_resseq", - torch.tensor(acc_resseq, dtype=torch.long, device=device), - ) + topo.h_chainid_enc = torch.tensor(acc_chainid_enc, dtype=torch.long, device=device) + topo.h_resseq = torch.tensor(acc_resseq, dtype=torch.long, device=device) if verbose > 0: print(f" Hydrogen topology: {n_h_total} riding H atoms") @@ -620,11 +636,11 @@ def place_riding_hydrogens( return torch.zeros(0, 3, dtype=xyz_heavy.dtype, device=xyz_heavy.device) # Precompute direction coefficients on first call - if not hasattr(topo, "_dir_coeffs") or topo._dir_coeffs is None: + if topo._dir_coeffs is None: topo._dir_coeffs = _precompute_direction_coefficients(topo) # Precompute static tensors on first call (avoid recomputing every step) - if not hasattr(topo, "_nb_idx_clamped"): + if topo._nb_idx_clamped is None: topo._nb_idx_clamped = topo.parent_neighbor_idx.clamp(min=0) topo._nb_valid = ( (topo.parent_neighbor_idx >= 0).unsqueeze(-1).to(topo.h_bond_length.dtype) @@ -712,15 +728,9 @@ def build_h_candidate_pairs( if n_h == 0: for name in ("cand_idx_i", "cand_idx_j", "cand_symop_idx"): - h_topo.register_buffer( - name, torch.zeros(0, dtype=torch.long, device=device) - ) - h_topo.register_buffer( - "cand_cell_offset", torch.zeros(0, 3, dtype=torch.long, device=device) - ) - h_topo.register_buffer( - "cand_min_dist", torch.zeros(0, dtype=dtypes.float, device=device) - ) + setattr(h_topo, name, torch.zeros(0, dtype=torch.long, device=device)) + h_topo.cand_cell_offset = torch.zeros(0, 3, dtype=torch.long, device=device) + h_topo.cand_min_dist = torch.zeros(0, dtype=dtypes.float, device=device) return heavy_indices = vdw_data["indices"] # (P, 2) @@ -818,15 +828,9 @@ def _same_res(chain_a, resseq_a, chain_b, resseq_b): if not acc_idx_i: for name in ("cand_idx_i", "cand_idx_j", "cand_symop_idx"): - h_topo.register_buffer( - name, torch.zeros(0, dtype=torch.long, device=device) - ) - h_topo.register_buffer( - "cand_cell_offset", torch.zeros(0, 3, dtype=torch.long, device=device) - ) - h_topo.register_buffer( - "cand_min_dist", torch.zeros(0, dtype=dtypes.float, device=device) - ) + setattr(h_topo, name, torch.zeros(0, dtype=torch.long, device=device)) + h_topo.cand_cell_offset = torch.zeros(0, 3, dtype=torch.long, device=device) + h_topo.cand_min_dist = torch.zeros(0, dtype=dtypes.float, device=device) return cand_i = torch.tensor(acc_idx_i, dtype=torch.long, device=device) @@ -889,16 +893,13 @@ def _same_res(chain_a, resseq_a, chain_b, resseq_b): cand_off = cand_off[sort_order] n_asu_cand = is_asu.sum().item() - h_topo.register_buffer("cand_idx_i", cand_i) - h_topo.register_buffer("cand_idx_j", cand_j) - h_topo.register_buffer("cand_symop_idx", cand_sym) - h_topo.register_buffer("cand_cell_offset", cand_off) + h_topo.cand_idx_i = cand_i + h_topo.cand_idx_j = cand_j + h_topo.cand_symop_idx = cand_sym + h_topo.cand_cell_offset = cand_off h_topo.n_asu_candidates = n_asu_cand - h_topo.register_buffer( - "cand_min_dist", - torch.zeros(len(cand_i), dtype=dtypes.float, device=device), - ) + h_topo.cand_min_dist = torch.zeros(len(cand_i), dtype=dtypes.float, device=device) if verbose > 0: n_hh = ((cand_i >= n_heavy) & (cand_j >= n_heavy)).sum().item() From bb171d52a1a530da3eb01beecb18d09312475b36 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Wed, 26 Aug 2026 17:17:21 +0200 Subject: [PATCH 057/250] Add a topology graph that reproduces the restraint builders Restraints already define partial graphs per monomer, patch them when a link forms, and map them onto atoms by name -- then keep only flat index tensors and discard the connectivity. So connectivity gets reconstructed three incompatible ways: exclusions are inverted back out of the bond/angle/torsion index lists, riding-hydrogen parents are found by interatomic distance while the CIF bond graph is loaded and ignored, and H exclusions are derived a third time. Nothing can answer "what is atom i bonded to". torchref.topology holds that connectivity: a ResidueGraph carrying the sequence and the inter-residue links, over an AtomGraph carrying the atoms, the typed edge blocks, and a CSR bond adjacency behind neighbors(i). Edge blocks are contiguous and origin-sorted with per-origin bounds, so every subset is a view into one block. Identity stays as string arrays; every indexing structure is a tensor and moves with .to(device). Nothing consumes it yet, so this cannot change any result. What it does is establish that the graph is correct: the equivalence test compares edge sets per type and per origin against the existing builders on five structures covering alternative conformations, pre-existing hydrogens, disulfides, glycans and a nucleotide analogue with a metal. Two things only those cases surface. LINK-record bonds are a distinct origin that 3A5V and 1DAW need. And disulfide detection has to pair SG atoms, not residues: a cysteine modelled in two conformations carries two SG atoms and each forms its own bond, which a residue-keyed search reduces to one. Residues are identified by (chain, resseq, icode). The builders group on (chain, resseq) alone, so an inserted residue is merged with its predecessor and loses its intra-residue restraints. No bundled structure has an insertion code, so that path stays unexercised for now. Exclusions are offered both ways on purpose. exclusions_from_restraint_edges reproduces what the non-bonded term is given today; exclusions_12_13_14 walks the bond graph and is correct, and therefore excludes more pairs and moves the VDW loss, so wiring it in belongs in its own measured commit. The walk is pinned against a breadth-first reference, since a superset assertion would also pass for an implementation that returned everything. Adds an AlphaFold-start trajectory test with a null-control arm: a deviation from the committed reference means nothing until it is measured against the deviation between two runs of the same build. Measured null spread is 0.0001 to 0.0009 over two macro-cycles, so the tolerance is 4x that with a floor, and a run is fast enough that the test needs no slow marker. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- docs/changelog.rst | 3 + tests/files/mtz/1BYW.mtz | Bin 0 -> 102304 bytes tests/files/mtz/1VER.mtz | Bin 0 -> 87264 bytes tests/files/mtz/6JZA.mtz | Bin 0 -> 90520 bytes tests/files/mtz/6SXW.mtz | Bin 0 -> 81008 bytes tests/files/mtz/6VHI.mtz | Bin 0 -> 116760 bytes tests/files/pdb/1BYW_af.pdb | 839 +++++++++++++++++ tests/files/pdb/1VER_af.pdb | 772 ++++++++++++++++ tests/files/pdb/6JZA_af.pdb | 565 ++++++++++++ tests/files/pdb/6SXW_af.pdb | 582 ++++++++++++ tests/files/pdb/6VHI_af.pdb | 843 ++++++++++++++++++ tests/functional/af_trajectory_reference.json | 95 ++ tests/functional/test_af_trajectory.py | 168 ++++ tests/helpers/device_cases.py | 64 ++ tests/unit/topology/__init__.py | 0 tests/unit/topology/test_equivalence.py | 324 +++++++ torchref/topology/__init__.py | 31 + torchref/topology/atom_graph.py | 284 ++++++ torchref/topology/build.py | 693 ++++++++++++++ torchref/topology/edges.py | 205 +++++ torchref/topology/residue_graph.py | 237 +++++ torchref/topology/templates.py | 105 +++ torchref/topology/topology.py | 106 +++ 23 files changed, 5916 insertions(+) create mode 100644 tests/files/mtz/1BYW.mtz create mode 100644 tests/files/mtz/1VER.mtz create mode 100644 tests/files/mtz/6JZA.mtz create mode 100644 tests/files/mtz/6SXW.mtz create mode 100644 tests/files/mtz/6VHI.mtz create mode 100644 tests/files/pdb/1BYW_af.pdb create mode 100644 tests/files/pdb/1VER_af.pdb create mode 100644 tests/files/pdb/6JZA_af.pdb create mode 100644 tests/files/pdb/6SXW_af.pdb create mode 100644 tests/files/pdb/6VHI_af.pdb create mode 100644 tests/functional/af_trajectory_reference.json create mode 100644 tests/functional/test_af_trajectory.py create mode 100644 tests/unit/topology/__init__.py create mode 100644 tests/unit/topology/test_equivalence.py create mode 100644 torchref/topology/__init__.py create mode 100644 torchref/topology/atom_graph.py create mode 100644 torchref/topology/build.py create mode 100644 torchref/topology/edges.py create mode 100644 torchref/topology/residue_graph.py create mode 100644 torchref/topology/templates.py create mode 100644 torchref/topology/topology.py diff --git a/docs/changelog.rst b/docs/changelog.rst index b58b20b9..34f4b396 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -19,6 +19,9 @@ Unreleased - ``Model.copy`` and ``ModelFT.copy`` now copy the context in one step, cloning the space group instead of sharing it - Removed ``Model.symmetry``; use ``Model.spacegroup`` - ``HydrogenTopology`` is a dataclass with optional fields instead of an ``nn.Module`` whose buffers were attached after construction +- Added ``torchref.topology``: a ``Topology`` of a ``ResidueGraph`` over an ``AtomGraph``, holding the model's connectivity as typed edge blocks with a bond adjacency that answers ``neighbors(i)`` +- Topology residues are identified by ``(chain, resseq, icode)``, so a residue with an insertion code is no longer merged with the one it was inserted after +- Added ``AtomGraph.exclusions_12_13_14``, deriving non-bonded exclusions from bond connectivity rather than from which angles and torsions the monomer library happens to restrain Version 0.6.4 diff --git a/tests/files/mtz/1BYW.mtz b/tests/files/mtz/1BYW.mtz new file mode 100644 index 0000000000000000000000000000000000000000..3703e8983a834b994f1c32fffddceff12b1f78b8 GIT binary patch literal 102304 zcmb@ve|+3kmH+<+N-Yv0SVW?Tqm}?#FgU1G5#AkYMU5J*y4H1ZSaC$HZZzsrb@eTU zQZ!1ifK5er4aI#nzt#Xe@qkjn*G!!S(xmy^}-EoZLUZc{p}> z&6(Fd_uO;NJ@>xv`#mptfN-AYUx6j=%J*YIS44Hz0pSEz2eYJ_3F|{I>*r)AB6;5Ab6`zna))0sKQjek=MH zwaueD0zL-M&G0M0zSQb^V&zP$g;Boz90Uh7_TDWS*^P7nzI^>lLNk_mKFcR zcpV+^6Bv^P2|FyU?drN3`7Fry)pqrt4L%n7dn!D2#^Bf>e;W8o{O}C{KOOu$^t>zZ z;9c}}ly)s%jHjDiavy!T zILIFfPaV7cCFp;2t$OFO71in`LH-!@xg%k>N^RG0k+9Da+TBLKEYFxeC%`j_Z@n+@ zFGapD;XlOt#4U_jBgmgh9(Wh}-yZZo4IbGu`k~@|6?!&YslU8gVy%hQKkAG+9 z)1j}4fA+uCuJ9aNuBC0|neuNICqE1y!8mT5{cn@6FR$*q9eq{?d=MLd7yi`&FX+z( z@ScE^AI6^!&s5NVgm_#gKMeBBr`b6PJ2b&xkl^1+9=H=*cLg5h-zUId3HWyG+yoy9 zxbnm66a6JW^iE@coA*uM$>qDcCW+JWz_Xhe`WSZh{wDjkvD*&x9}V)%rxTiI*9Lr= zy!Qe0tOtBQ`F|t&tOKQYbkkhka1ZT zYnkhpQ&sW0xmS?Re;}1ImzxuO}zWF>R`2p(u3(;q3kgp^EBKUp& zk^CU}Q1!{;AYUxsy!!jW&n2(?^V#4J!pT1q8nZnvFA3iSz7jp7KF#HgsmH-n4ef5l z4x5noekGnU;{B*vsDHr8Kbh)^p}^BZUUSTrfK$&8e*^z)23+;YE6^wMP@h!KWxN_e zUis(!=+g?g>iLggH(M`ce6APuN%ecQI~nAu=eIq9`Qj<`!~DGb4mQ}G(4YFG|J~@b zKJYBU_Gj`PtxP>n7q+<`Jo+c~NzWw-J*lbtd(qS9bLpe}^ALT0d*E08k^ig<`c!E5 z3+NR6Ry9;D!B76FK7#z9`Pb{e6#f+nybt_^iFT<^Ca&c()wkuJ{^fNV*iGpF>Y%^s z6U}+A40w(m4x!x}`L5ST`A4M1|HTstF_LxI5PP@=Gj|lw;tqk=H=(Cn9GMI zKY-pVu(R@m{tT}(FRPw0KQx?rX3Llx7sCs9j>gV5?;E}Wp5MB>VL1B(pHsm%pwmM^ z{xt0NY~)7+eme4>m#zWV+;|KUzNEuYIi!kf9)adUOP;jLVC=G=U3cpLdIhPW&Bf&1C$IX*k@n>=}G?P1g{ z-ftyOUfT2dM7#K9>njO<)j40l&OM=DvcYSSZ-sX2;G^hc`!=Qzd8zR~*nc$0lb4Do z<0rOXWb)KG`;>=${*?YB=yOOyp1d^mqJ%!wIlV z`A&IQdKj*C)-CAM2sm{P`*ri>sdJ{fk?#%i)Hy@b=sy;4>YRm-aP_CPyO%unE##wL zO*8NO9DX0~(sw`oeKzuyz@M?6Iw_%#)>%g<#$D^Iv%!}Io-S;8I{c=e{o75fzak;8 zI_E_RKPNBspOcV3hJ2`gZ3*pGYFT!2LVxO@S%V!$A5zF>ihf`69s)J4jN-p_^S;E%))FD<`2JKyO&>#r(y zw)K?smptpZ@h^sP5#AynzKeeO_!8dc^&a}LkbliS!ZYUg%ixK4H+9Z~=(#G$SLlB` zjH7s{yZW_{sywN`JlFs6KjP?5AukCZU|sgD;Ge?l@ZSu-&xgVXvGJRbxAkBSCcFTD zCgatjUEi+_tb=Y#=s7}uSe>xvCiI+Q>?3~@eeMqVO6{udIxeBl82Y5fcRP9VPuR!% zvv|mFT?bI#Mtmpwd?eu~lkmSVF%gO*D8Ic7yqb5mt99JwguM2d-xBgh#&^4GS?f0Bqn_lq zvi@lV{-fE?c|syCsJpt9&x;^mp z-#q#|YL~kLPMtRV0Cw0IaO$*)?bxC>;H)3Z{XL(@)NVJryfyfPaP3Du9)4d(3GYLf zvl9Ky>8tkJW4q+B(O<(A?N%3mZEhb%+x3*own+z6`D6XdB*7Pc<@{CkjENr^l1k97I{Y3<5mY; zdF=P}Yg@o4>DT+<5vSLGH~c?9e_vlohV|p>pP>JEkl(9u$8YZp_%!(WA%4ZPe|h8f zC*p_q1$pYUMvIuK2An+B(Dm7#zOrzC^!tz3hfWzt|5S-++C5{7O&N zY5#!!YlA-QM@>B)JFE)0){n13*K2}4%40P?v+M8T*6RfLm;sf}}c#u~f`+A~Z+6TKGjon^Zw-W2?{ zvapxyDLZ2y#bdPp zR`P9~c5NMM`j5dsB|i!FWIwsL4R&(q+k26^hSp-&>;2)M2f-c8$4|6|Cvm!fAq$dkuG%2JEpuH)KDmZ|>AUK`;obOi%uAA~ zu)h265O>1+@biy_xD&2*blX;U-H=M zv&koRJ=b`q%XzIc?@1rlmA(F)o5^ckxe)tb8~7Ixr^W2}diyUTu0H|3G{}?3+S!Uw zFCn9P?4$4x1^FXuS@w6hVjO8*sdZmZkk`8MWr=nx=&$ogKF_GXtSigw%Mm{TeKape zJL|gS@$g4|RF5si-|9ghopZR7mtBuGdy>c4cbz|8I_K~zc$$GHNB_s8zt3YCGArO7 z<2$N*^;hJ|F?zRa;^g4}!Zk*l!L!--56kCS(F*pIJo&16V{N{C!9M>Ju!Hw=$v0Ui z{xsw(;jO&7?}|jbZR)XeY1hZQgJki7Q6 zK3#(^@Z|8HkdRltdVPYw4*wSLuE0axHE|pKv0d`jo=-Dg3xoUy^m%Ne-4S>+KSmzv zuGaMlo+k2t4|z=eC0|wdk%PS73Rk{T-5%{szUryK@9Rv-x47=`Nb<&T(0?a>CO^*t zK1rM1tOwTxd^h+fXZ?Y1J?x*VyS|~A4Dx%K7kb0|B7LTrr#?%&dhf?q>&FrFk32d@ zabd!rwSHU!zn`bX$vKMgTNz&;*TS`*-=EN@8~n`)JLr1la})MFntXC4->C=vxn5a4 zEuoLrk2=3mL0*4FKGSz#b55vy9z5`XYQJ4BJ24}dGgYg z=-QCnY`hE4*@wTbX5XFtcb4mRbX^_fwNCsS`gng35BX|n1piSS=}(wmZRGzBPmliF zwVR>i`?1lQfOiw`Phz}$9uSZA?H@|COMO*+Ztw@m=fr8;FA`3^${x@-hIZ@Vs=57q zA^E}O)s5XNI_q2E1$sW1=+_3~@yw9-B(Ht@qk{hnSAF#p`Yr4Vm-N}my0Rz4i*WMQGM)dwGw|<5Po2y4 zc|h`Qc;1rGb1!=69L_K@`jf0)s;_eFQ`dhxJM1SPz7(Fd0nZq#z1ZLHzllfr>i&eE zFX9}R@@lk?^3_w(C;9{H#0k~IBcWaG+si&H1Fm)A&(J68$vQE64*J|40)JnNPO``GS1)CG|@_;&}AE{mu@?LtfhVW%ySI9`e%S$58*S52QL>!TeBJ(ZWVPFxuJJ4Ri1kMmpO(RtJp z;9nEusiV5SoZ#O{-JPru|E)S>Zqyz z0KYHbnoD1Rf2uuiH`aaG7v(3x4(wm;QGT{_4yLE};s4GUj0SnF`*dH)VFAy<@7J6Z zaIO2!M*ceiXa9=pbet`4xb)4k#?()P9}4GO#KaKx^l?$j=UJii#48 zvGWh6kMh{hiO01;UiS&^=Vf)H$#X7(>ws_vdG_I_bUxPVSd&+MwHE#dkk_B#Tu)@Z z8Rmst5}QY#8|m^;zD@Fe$+G2$NRZ()m^G*BF_F*L)Tk-OU@a6B;Vw9TtZ&? z%FczFJta@R8b3b7qwoyBU6J7FW-hHH@Ctsrg8uqEES^5{%&G*QMq@1 z?EO>ia@}zB<%#}|(dGw2o{>EBTx&b+R;2T6eiP1maA;EkSH9Bytv)|XeiFMqh#eY% zNBQcPVV;sabyx8r^z`$faQ53rzf;RP@_=xy2X()d@_^S<>%q%Iyi1<_tNxoqybITU z`}Tz0SPxF=c>>D+M|1lLP57jb&T+i~o$e3(r-J|MY@YFU)^)?T z!#^72IsZ_7IXrg$%J{W6_I5tAb*15|yJRCf*C9OTMSaD(rqdsU7iTu6blq@`>+?+Jvh3V!d;(r1u%--Z37-3sD5 z=1Ix3-#&FsVq8>T-5LB@^2%3wenQLDOTwGrdfr8}r|K(zy|0O^96Zqe>3LNwjJL*t0VZ1Jx9gb)&5mZ+x8qM!*kmFG5D0@DR{4Q3m5X4m2Fb;90;k z>ZO~59VE}XulmD;ygt))sE&FYtVvz;gpeO3t2$~u{2A}PK5ecW>bzMa;M7t5Dfw>V z`i=NQY`22nb|>2HBhUPtvGI9T`jD5#4^7Aq5Z7-?;H>+$TowFTJcHybe_v!%{m=`2 zw%4+G{nd5ATj1}JjIGOfKZ5*qUJ_G%I z9+RHrrC~kK&BvW^@>0*C33>9;sP12g_Mc>2-V)+eJe*q?x+uZlW?cT3un#rn)Kv*T zR2{XA81VT@{HmjL|9F-6-kurjyGug6OMU@)K=WX<4|ypYO2~K7=AY=V&#U55UfRHS z^xn6tb>I2Ct`7Lo)JxA#*n#_cCf=RsZw2m;f#-pTywtD#huB}`CEYg>^;`;1KmLDf z;L*Bo1K;uWe1VMWs3Z8Sn>g_OQXQo;@jm}ZUgxj=fWO`0_?x`v976wI^tru!Ka+=q zEC1}`J1figb2#UwCR986{YlBQjvKwJ)@g^52l*#syCa>R!kftJ{Nd8T-(uZYgU7Fb zCEo`BBfjP9J>eOCsPj(V&ceHqy*zOXFg?20O z{FPVim-5e35`L?5QwJs5RsLaYM&A+Qv=8nBd3`$I%0JIy9Jd91w10jtzV(q=`4*qg z{O#&`)(YMQTrqeGdfGhouzbGx#1C_wWkHb7^Ic2(m}kfHnfNvNf;_WV7D8Co+H>@ zd%F|h`ZIa0zur%~)x6XGP5QeH{*?h&o%7DbxQxN`o&-*w>D}t$*V=8N=gx#Zcak@B z{qdI2?j-y=hq*T3I-mLP=w$Pe@l)rFAMEmy;o2{1C+s{;K70WF9f6mbvoT2u`0;xeCB_lf5dg(=@Wcr`^MH^ z_KO3?0l5TYrp6eUY`s+Ir?P0t`4~Nw?}!|bvp5=eb!~;I`{7LfN)+d ze{S$$`SLtcXC9Nh)@A>I{#MVMyzhHDdV%~r;N6hxg$e)pa4oc{wy6E@tM?JMJdr?9k&N~%r6LF-vNav^{kM?eN@L7yI zU60fAS|dJF7tO~jw7Up>eBP6tm%#sfa*)r%!cU-IAE8}aXSHVSq`F9TW5iV#=^Tr% z2gIYg=s;W_K6JFtHx$QS(g6$u{IQT|+a3pw$$;5i+AZ2#HtcAgc_OthO3 zkGru$l<#IA^~&H6;@5uBK!T?a|34*x=ftC)-|qcHJOkv1(-QJpujzTGRb(_yroZx% z>f2~H9x#&?+)D~D3| zcFddLW5o4Sv0Dq>z5e8-;@X5A$V=ll(y!!LaxtIz{49N_qgtOqKKjYbK6dB4AbFh=&~ps(ppW)W7NWn;w~|+0 zs@CN1o%*Yel6-&QQ5|(LuloY7y!3gv?hUy1vA@l0N5FL!=?v^_`w1;Rmp-S{?gjAp zJSIE``!b(pPR|@%*Vk?Zuk+sQFI?BxUKYl;l&5adxvbtGpOdfJ!GB5~>YraH@Fw^1 z`TG%C@JXI}Y2Vc$?u57D{}SU;3GHV1x$aBy@hu}-PN7DUAvE2+` zxPbVIJPT@B_BQnK^O1N|FX?&ReqIpHK6by~-)?~k*Zp@B$ou>wT>D7(*XH}P&I5cq z5nnoospkc?`L4H{^3v0J-5T&E$o>?bl>yg2_FZuKIx0toywvl;#CU0a_cOk;KJchs z`hUDO23+gAW$@enTZ_-duj^~)pcC`_!}=HCx{h`o-?^6e)-U1YpN5`qnFTx}|2!^% z=g9qA0xzg@{x!sr_*t)Qxi9orIQvU`VxAG+&a>iK34YZ@7bNg*>V)mVpT$!_{%iEt z`-5=u&%_hLye6FW8s|?s<3jiV@|qug-Vjdy>A0@c0uxSMR9O<{ec=W1c;{?B@%AVG zRJWt2pHC$}0{`MrHwxE!P4{^zFL|COyJL4NS^(re${&sS6%cu^1$k#kJf9yr;SFym4A**@btm45c{Yt?^ma? zpLQtyS{HC#-`8_Cd_C4gM)&Rh0X%o|!r{zs2!2?~^VvItzX{j(-$~zmyCo0x$sz0~ z_agO41O+t@=V8lwJk8=P58C{@AIMXt?>Kv3?@&W;Xb?h^5mIvz3=Bo@wBkdjo5CS z_ug*W$6giko8-weQ_Br1T0iy% zo~)K-4-tDlKdapZ#L<$FmxM0@*LABn?&O(5zY`I0t-lUS;5tWiOaj+F_P2ONJFCuF z3)fiCpFA`1laQCBC)dj+_7X>38n@Yf0O6{0^xU~D;5w&$B=)QaeqFEs6#Tb3yv2K+ zm;M6ywE^ckdynpe_WGCd#Pu55^!Z144t^hYSQmM~-_O|h@*UG(^4hnlBEKrgx5y8l zVtiKyThP+hjgMaGx#q4^tFukZx#=i}Z&sW0HyQSw9tPT8q==d6JHxclh^D|Fe zF~jNgCvViB%6xijkgvn@<1mhDcaV7h8RJp}dGbc@feHOppU7|fgZv117V?SsH(_7x z_4v9}coTU&kMyiy|E=hGQeu3yKKc!ITOZopjy*Mh^#)w^$s6gf&l_r&yfLBYHO29r zMEHDb)8-Jlltq4!l%*mMC_agJ=wS6zRu44AbDM9{}r#* zL4H9k%YK3Unt(529@ldzV_fK*kItRYj$ zoqsT%W56E^*So?vR^U-=pB8Xk|I+hmR|H)9HhQjLj2GpNN7BY8X64)HHS@Q5ROd)- z91YjL?)w-6yYI~KjO#@2bNSZrJlAts=jv_43*vp)KFs6WzHZ!&~s{xg}Z9XD9mHKujJH{B08Z>zsdAklzje zsLO}O--hSt#5gi%_Agj5H_w{Xy|3++B9yVO-msKtw8m{%r zPkBZCRaYp#jfQ@8!~YJ(CE{Ff8b2(dr`9i1=xOJBtzFd>KcQWVYr|DnjPac@(u90FZ>&7nt-Fms!^Vdsc)E!@fB(42SCBs!JNP^-{mDy> zmnC>|aP6zdehml1n^FMT=WAMp<&pR^D96bT;k((tdreLfV=2z7!qhJPHNumgFirQfMp8{}K~+9?Tn)h(K*#%W9KnV#gO ziC-X}2OjcL@xG8h%kQD*E%qoosrp> z7wL1muF;2oUB<@CYM@3m*Y$8~U? zt2VcOF1qVZSapFZwN`8>;YAMV0ZX0|@zea_3DTX>j< zb<_IV%;)BZ;^Dedm(Gv%1fCZEeF=SD7I5*sBlwSalz)Dwz6E*JMF$3dkbH%Cvqt{6 z>rmnq-UpA?DG_JCWsiQpqob}segJ>gIdDIY;;9oCFG|>LkoeNP6Zs3`<;D;%;$aQM zb0IqYS@;O;J~_d&iE&kIZ$qZT&R`E>w>;p=Kc7q3e+)ase5-a@H&wrdKD~j5{gypr z37(zH8(U(25Rd92o&S$^*bUDl_|m}NrrpC5_9XvwJ&1fqJ;QgW$@6+%V&o_POsNiO zxqK*nGGq@<*jekQHQ>>&*l%fkm^iHlJ*kUWBh5WesV#0^+HcYGf1wd=EaY?H8UA*2h*RO+#N8w4m(P>J$unCn4f#rV zA9>)XVO)gg)Cu1Y{x5uhczJT7zjfNZm2q5*jE$q@2Z`(dt!1U}dAvaWe#XW7o8-0E z^5aCitiQ%%evmwQrl^N_5#A(D>j^t(U-$0BxGT^2`@Gt4h==Rfjll$e3!bke^yK<= z`yjWUZ#b&h|&oaxEF zZsP%7(f?KFT$u1Tt-t*5zZj45%x#S2p3pAqud4R3W89Hv%5y=nUFw`3zfWcS75M*! z{zmyF;4fo*dqTTff32bITLaGei*?`3^IE-qwEj8(|Fr!gYgg;9J>Y&lSiq_M-o1Rb z%JB==?;L4elLb5@Km1!5U&-g}!{|9YtAl*OxtFelJnN>>L+QiqLB2)4((h;acocs- z&x&URKM|hcZ$kDT_savMztw)fd zZt4GF$V=j{Jc&6dm7V@+7*ShKP2|dXtL!W{_@+@FIaxwCiz^}UH6SO-NaGlfB zd6kUzy*}Oe;E8;8Z;)3$c{!i$2)OE&BY4^QTkDtBP0!@DC&;sIs(zLJM!(g5?_c>W zj+gGMI*ZTzKCSGd{Tw}K#_pT2c9l==f#26NP2|ZNd#*>m&U@>xaP7A!5BRt%<%##r zaCy54CvS}V^W<%?l7~ESIq~S{S>en?t zyiqyIontY)fPXjoFGI%nYXkZmp3svzr20MVusF!;9Np&=`ZUq!x`cgHhv@m|et$rE zj$xl=iFU~w<+;jehZgza$`IG$VZAe@=f(NFB%Hi4u{XrEaMnAcI)@zj+t~RHAx}zv zFLsFahVW_f!$imf!j(66*RqazP^0a9t<5Jj7!u&pyV+;BUg&kL}UC;K#S*;oO=( zZ`npxIQt>}2Lyi>PF*ptzrvLtUPK?JcyIj@p5Z5X$Unln$?x)?s80p_pA+)R5BGrg z2Y%HRmj?eAKlx$wj1Wh{>&&xTz$<}={E&Y%!Be1*?pyB)@*C*ayRrYZ0UseA{e3HK zK3BV3CmK(UOA~$K{-@-rE2gebj3et5o>P}~#F6B+AMzOVxs&(a{^W=J(uDrxhej>r z2k}h8b8hfw;mQxX4(RiV@HXwy+BOOXF7ufoZaUNFD!qhA{X-Xwni3m%_mB;Vro7VPh@TM1XZe^ar+%j_dOW8B|KT(9K4$F+|5 z673EJyn;U;8SE?`#ru!&(Q^WS&V2OxM8Aml)f>R;fuDG9XwIDnc@viGN(Qaed?GMDps$d7| zfyPa#@j~|d^ug|vwf-vJcaaC8pD5n9)aKi5H~O3y@@>gQUFV;BY5e8&7yQJ>i#}7l zza0Fap#OgQrTaA31v_Z})1O~5e#JYr$mqF&M`Qap@HoEO|J)94zl&!)T1PyAmpz}^ zaP5EYqm7B6kJb^_W6$$KyNY+6=d^Lj_0Iv6FR@7|MvdV<~p+W zVR{4qUimNMYk9+XrfFBtfsgjlc~R}7js_l``}z>_Ie7UeUjDSN_`eDLb*|&l34PcH z7=LlX4!WlCQ~Imt0+oNW|Eiwf2EXN7(_j0FI=311(LTU8_)d{`+NT1xH^HO%`d#$b zt~(gN&V4{|fdj<&pE(was5RcRH-==2q_+PW~C6!Vb#zs4=@8F-Z8|#JQFyehEw5$AcN5~tJ*M5S=*w&B6&w8QbI5RSxs#>dg*{EzQ@_uatju z-75OI&Ts4`_OdX(T-V@tIp^C!`;cd1w;e%V_vIZN@~7;k{e)c!eYk$g_|D`j$(jDj zKexf}^N&7P{`nGZSlw>&x_;`PA8xpwdwBr-MbPt9@EiHg9RXMV*~n|!;Ti9Bz2kz= zFX7sk_!9D!@;zdVYbj5?@yg&2!t;FQcj3Li2`|9E7V?$wCU_YNFG{l)rAp6S=RFY2#zV!D15{g8D^ z&*Kty&^qPQe8=ZiacbZ4w;?|WSDw-H-=be>opKOdZDhPXwNBCVZ}Whw&QYHC^LU1g z&WXuxMUdCNFhlg>O4EbbOzpL)&BjJ6-%Vnv4!G96M1OzZgxO8|oO-TQ)Ta$TiT=?( z>~ogq81D-Dux=uscjkc{8R{0UFU-ev9^>GIJ+<#KLcgLtRkv)VU-i%~`RGIhwN|ynED(2LF*=+ z8+jn~OZh~8cn0>f_$rXm^%3n4c>4(FKD*VA;r_2R<$J{RJ>F-;`Mz0u0Br+Z-V>#B8{K*&agiRXgKSgVgGy*!!zoE7?0v7Zw&t+*i$%pqd)G42=AkR zHEgtz_uf7^{`rF7SCSuq=M@RO&bn!1!ajrKvBzNtpZCO5(C$UG8Qc1HH&CDJdF(z9 zOMZlY>30Sf20fKGK7{`30^VG%`>Z?TS3KG;{pS#;!nI%OpSNIkrVgq4^VNp8r~_Y= z;Moa}epe{!If?#XLjT^Nr}BpGSM~Wp`n2iS^N{!ThHza!{c0`i$TPxO?=;?=&}TpL zUxurNjJGp+V_f^#QO^aeQ?wsm3G(EPj_0D79dz#EI`r}TW70=;h@SIY4?Nm0eI9uE zeOsp>qw6F3UBBTVujc@$zmcbcyq-6h2YKaeEf!jnz8sBmN40a!nXt;aZRQb5_Q$JfPo6_wgwC0sQ}e!25Y`b{1YI ze*OI>#xuyg>3>(k@PhenBGhl<(RI?tCE6W9A3YztC$zhXIMVgLYQUT5b5e*4@skJm zeb`R_7e0m^bYGH>3*p=0S)b6Sh5WhL-^Z`ysUJA^)fq>^C(%>S0f_wUdzANE{R|mz z&o+MdMaFU>;K~EPqrY+7r;*pXz|T8sSLZ1F`A@UI_B~#d&`0|o+CPbU>OAM4YV-T0 zeGgra7z=u;e%MPJ(N9!A$UlAFX!2dHM@~Wh{=lO=a6WDPCEyh}F61*iZ)yEf9?*HM zCt(k>XN&LZ?=|3dpQGf3>;C$+VSGz@=4(Co*w-P#bK-Yv7~hhIbwtcp!Zq(V;aOX@ zm6iEKxZ?d8A%2Au@BRLF?o3bS{h^)6dq0PvU**$q9MQ`%jr2N51{A z-?INu^n^N7{F?XmyullJ@BK>i{-HJbX(z6DzYM$i{Q>dl{(9{z?r}VA-t$+Cc~%SsI|$b~2hCGHj|p!fuXSi;RtMu@z0dEJ zbk0+f&xng2w|@8YlHC;nWM z=|6~{za+t<`1Q}5GI`bQ{v5dBs@rc&=)VbmJ(smB*r5r2c%og!?>+cKPmmwOZfC)t z1zhp_bozTsz!ksxomM}u$qtI&zu<>HUkRVYJ_~03-}~Ec`m5`Yk*AIP^|TxPMAs>F zPG(KeXPS1uiQldY_!j+h&u{un_g$We{WXu5f5LdFetxO!S-zjaN_nmW>VAYqJ`=Mg zhD^?NJzdur4|qY`-Gxox9(WXw|ChiupKAVX1s>vYRL=+75pd#h^jh@1J>Z&8KaS6x z720J!WuL8cywrz2Wbak12KgNPN$67qe1QLckam6E6OZQ8Q_%C}L4FW!-G_d2&{Ofa z3I54|6Oa7P-F!Vq=+|)xoOtB-W$@e3F7epV@8K5kmw$XNJ-2c$Vhi!>{aN@J`drL7 z?#nyJSL^7p#Q3(r|1A3i`JJ>o1@7xZ@iU)hAICm^ei6PK-Syn0b%BR?oH~U5ZVb5M zu}#}=;9rYt@i3oOzmV_~t)q3nl8;l#Yd+mdyVnMO#pCnQzc27>|E|?m%NY@k>27MGqUkB?8 zc)Pr>aBknsc(6tBLl@VEcM~t?xO{Fn^IF#_=xK3hcpu}n5x-g*`kOCrEL?-U)&C|x zu)J~m!R*%!2YJQOo!CDQ_#l3z=Y&Uo#nG1;Gs~;SzXARI@3R=Ly7Xc6d9-2o zvlu=}zfN%Tvf<2Yt@ZF%^O?9XyiL6SC^5c!881D**YbeL6G!9A663Bo(z#{%qvFow zGjykoho6^a_yX`N6Mn0@^x5#_frowUQSAp_8*t5QFLCi}{H#Oy{f@c$+HjrcUx5Cr z1CL@w`F1SeTz?~`u*Q!!GBH# z|1sa$=J?BemT_J&*k8ErE8G+OzmzB5{rz-mS9lKZYv8f#?1mS4R`}G47@?QkM60SN~&tvp)SL%b`>iOi6r<=OyeF=TYC;k4qzv#&oAJ+xZ(wYnoX~#+{%)~6 z!XLQKzFN=Gs)u&9uchBJ_4BFxryE@!NxRWMwXdc7di?w%dF2zGyY_ifxXu$kiI?p! znH`i*{Bs5jUxF@o)5cKH^91^(b6L@!b)IK2_Vjs7?e?M1OZn^p$0Og+eFW>VwVtn{ zKea1d-#Iqevy|sP!RH145w3c^KY>jNt>s%ZC}4RYAVa z`u8WayDZ?U=bx0Y19iy0zu}+#L0+5oz)lt9EN9~9E-vKm!)$=b8@hh9Co< z`{B>wvmMBIdun~H`+q9|*K?eH$!GO|>wL}^cy-i!$mo9LRp66BUUi7pS@nSHx|p6* z9dTVxznWLm$>+S+b@Ru;Wxqoz`v_WVi7e-p0P_iCN}6;_Z>KE`-0 zEIDW5PB{5t)L&n=cC}9N*R>6&uIPDp@E`GH`0Zx&@$oLa8$a>q(~PG=yj#0`yUYWU zCqHmsXJ?!W&*A@S7)Rj)__^+jjQy%pKet1Eko+LH?pMlqZ+;@Ypxui!#sS{|eq@4w zgzK9`eJ)qJ-U)A0`rZn5Mt?5T{X3<%h%I_w|F?Rb6p)Vtlnu z*#f_x_a(3Ww=X1c>I&|Y>YS$}ultN%kNh^?`*GL)x%_Q7;5sLK3H~r1aII7HoJRXy zo`?OJb3P|w=l&qS1pfE&9lySoK3b=2rVZOi6;A=KI^lQVYj|&Z3fDQ_-@xVPyHcJw z)%jR&XW=<Xc zozRE=6(P4&-0l4AbOrx z>&)|Nw^+V;Syx<-fU{r0eR-kY10R83>)fTG-A%~r_l&v%-ejJ=IHBiO`Zbck$KXE- z|G6vhuwT*8eFTxeMZEj-57ytE*xx_r!Eo}x=#Q{})PFbm=S_+6YSXX7v0E$Xqw8oo zR}=kEbwUc?5C6pp{Z%JuA3O5sx_J{n9188Sz8HNr_`ZPaex)O3{oMP5))$Z9<>Oc5 zuJwh^=dTVt$^#4Ouisyiysn!+D}if$@o&63)*Z;JPSAe1-w$kp>)ih3jKOln|7<-h zT<2{6Bjm$Uo;cOHDevdPbK?46X6-3fjc>_A{N5J&CA`JqM82`Zw~|mp4`( zy8_(<&mewye8`95VO_wwyEBf2Z$SR&gr3BE|Hl(Nn~1O1CG=@Bwx2@(wLu@v`}$qW zM!?7Duj;@%1Frpqli;}{;4S3!dkWDXnD?s_;ME|nc>fZgMSj)=tOw;sVcdyR|6ENQ z_r0{M`_MNAdE$Nelh~m*;F|Y!&Oi1yqs@2m^7VF;?=FB__gDJ)PPop=yaD{iz~4nY z{({d|1)S^V{U*6ThRY68t&x?+o!J`2qau;|aWu-Pv<__?(i) zLGsk?jb#aW=Ic@27w_jm$!kCG-3fdI{7%L=+Ku?lj;hH{opx&yD{n%ds{_uuzIcCv ze++x-cL}}!OP}rFFHG321^!g@xf2=vS$~U<}%`5ubf?Afn9o*Nq;#b|S zI9{13X?b#rT1@vC}l)XBGbulo6?aQS>9T-W=w@9M{;l;<=5 z{1EXA&)IK(H2i)ZFM0T^4t`(BIWu0uo4ovUKa8hEyuUik7m{z6b$e%A3nw1CKAzCC zn|Vjir;7Sl@I(K5WY(_Y@gu=M#h;hwM&|Pj(C^!^^-5%HTqLh}yuH@B{~&x2fB1ax zE8)zi)gL0C2c8Y^yc4_a2)N?$U+{-(1HK7A)c(@KfGZw#KatPpYBx2X3LiuMF68?I zkLJ_k68344mkvpc*G}|YoA9?u`t>FB9}4`MPoIVzGCueAY}2mJjTAwib@bFh#9bb6 z=F^D@+VyeWMrJ?$(2btUf_#ST0r1CmHJ=_B<_q!Yxqnk(9Q9q!pNzi)e&4Qfz7l)- zJW$Hx&yT{Meq4krj&2I^Uh=SCt?MjNzKKp(1UpN165fsdcf(T;JQc?4cz7NNxYnV%ZzAfMGp`+$z!gW|hQA0r)TNE@BYk2{`*v zjrW6V-qD}wqj}Are>1#=-QIy6;&^FZdo}$X&pXF)693oznALzQjt)kDdyc*FD~`tD z_jz9S*~{xf?BnZ0;mm8JPXYJ&T=;%oI$zxk?dtxolfc^!FWb&&m$v=&XyNk7uZ4C? zdA|Qw^dAm9x{jmgmixM~|B#Dy>jyQdx?vygFXY;{~7dmG~ji3 z%02PGKgfUI!mA$e0{H`(7kvGrb`=-bV1Il5jOnBKrF@`xxu`e_xGGD+A8{ zOha{A5%}BS-+{}{IcWT5{&r1l#{O$(I+{IN({X=Sa-07fe1!?r*Ylkk|Fc zKcm0pNt2h~{vP>dL0)yy>(M9jwDT-Kc*R^i8V`QU`GL7SX?Qo+B@bl3&iujf3gdMx zws_X8-3;%eUV2ZeU-;F@UN`CZHT^2+CD*XHuG$;)qFLm#Z(HoOI< z=WO=}J$It#9r$@8;FDZmxSV$7FUlLnBfq_bcIyFGKK}%LI5CWu&MoNq6_IC}`L{3O zKU$|~e_%B5>%Khg^F;e|{c+Dt=o$6VIViCvN zkM7I!`Pq21&yZv1*99J3fBYDq9To5;=y@o1u=&n-bbUeBZCk+(OO<2b`9R?51N#nk zyTi$A-1N7S&u-4=#!)!?ll)HH+;uLsTgsCMUWz_;?#1w&xc)})!;*(QkSFBXpB(bn z!HuWIKIiRW+{MFuH~!oNPlo?YVTUZ}ue|$>5MSb{5HJ4uNY-v2^P0bJ)bJdC{y98W zw-`P^T=?gZ7+y#I!`RdNgW4U$Z-2w1KKCq;J9 zpW$4O;Q7q+`=z|`C-^IQr#?Ej@WU`aN?-Cu;h#HX{JO8lKX=q{+4x@kXHDQ&-nfX@ ztpO)*OxW+PqW;7BpMdPS^v&lP>7%^ycG-vVG<#-n>hJx0)-gZklG)tz{xFWEJo}dZ z{toFcJf~iLK|;Pj{%^$Bs*-a?Psx)XxQU=M{|ImKdSilz{J`@`vW~crJaxs?$q63n z;UWLKB&KHt|9MM7Uisl`AuowPr+@x@k?|-$`12Ts*YQK`L#iL%&V$7Dg^6|xd|AKW z(NRb7UGhW6{jlcFwe)7YF|D6cy*KX{o>wP}n)vq@FS_S`F-s{igRae}aXm^@%JU`?~ z@$6?DwT|=oN_a+}F2ru@c<;wc=bH2!W?P?HySjeyIbPA9bslm}B965GseR_9fuH9g zPiX(MW8Ozbb%pA4TUT1Ux?ZW@mAEzVRFFTL*Rk}^>#zKv>jL*V`JDH9E|AXIE(^Hs zOZsqVx0Gjnx;l)b@SJr)3;B6GeD<#idE$Ni(l9RKX~A=50&kNy{O^v*4y8W$kAI$w z;obPdZDCv_PdzX(6#Q9uAMxd%Cu2N0?VM4YJ11uN0Q?sv^sIAk!T-*P$q(X(=gi7y z{25L?!1cJfeNn?VAioH_7u{?erRNCo=+A4J{3hndM#|2}``_&~`K`pSemANT_{ZS! z_Z6DF>H+^dJBGL5IV$*}^xTPk)`Ywyd=hz$i{@G1uifx}27T%QZ-Z~duT}?K*E_lr z{KWg{*YW?AL7sRo4o}EyfAv2Sexm)=o8eg>c(lKIEP74{T;~&x#s0BhT1V(PZwI+} zlpSQ}A6~$9x3q=q-*%xa$081%E5$^Q?FZ z`uq7pxZ+yRf$WJq#P2tPe@ec|dO`c5J|7C#{=n(M{=(bbC*{wl$*+WGjIY0*Xm~gN zGp`;)p8bJ+pHAr62j4#?`kTY^%!K>^_!Gzv6*$d5rB9tW{cwUuaeZ*YZp`C7x=!uq z7x56+{0>N`p9miz?qVJgPMu%%_w}0{n#A?W**uo=&%2p{)+2gx{j%}`dF zMgy+vm@notKmV#-#kK5g`>4{h39kEs_oC0eyf-@wSDflO2EI-x<+<+v{t$Pu4U)H&uaU{ISe&&{ho{G~yCi*o7|3dWf`!?cX{_WTA zO=P_H`nSk4TKD<*mHbZl_b0}MIOY0d=YECcHUFNS(6dc^Y2L4db`_`VYxBob!9y$SY3oqhG57t~jmodLZCBC#?PaZeC`$ z7BV_lp!-uizi`&MqaUUXyUrzjN_pnBok1VrT=yxTE9LuJ@(_<#hxiq)x>xrDdcPIU zt9VE7H{osCJs`o8Vb5!@=P)w9Uy8f8V$bn_R~WDB67qfYPx~Z!kk7$SNB(C4AAm>a z(_(+?_@VBfiv1eIfAl+FmB3TLe^8=d+CP6^@Mqa!1RmW7w>a=?M$^(2nCUvfTwAM#p9tmC^DkK)n%<$q_v z+8w}tpF;lHAWxmy&~uY20jJL7xmca`gZMRnU5%b=gFN$B_4j<{^R47Zz&FtD#vsr9 zHSrAi?Yg|_*+kDD)jHQ@;@OHFbpMN=mxZg&Jcse}b(HY!~R4pHmCQ@G|4Jx8%O@ZkSE z&#E(CB;N+tIY~e6!twudox3IQ$p703`3iRE4dW;t{D1tB_`{mOBme&|^uICS1Mpu2 zz9ZnwBcuNr>@WU7cy#^U=Klvo`ifGxhDO+%IT@^YTcoIi5ncx4Cl-~kGDf9j~~7q`4@TVZ-@osFeAADVq zui#gwppTu$F#bNi|2Xow&-3EX;eRvoM+ScF59ogO$X_QP>Um#9Xm=1E)und6%i1l7 zmnVjJQM=09I%lQ#zF#Bs>t|tnC8Pa;pC)km?bETRpYJ5U6E_&PbI$z`B$RPYi8y9 z_;==S*Tg3B)%scc7_NPFJ&(xpw&59d{^#J=xaiODoOt;W{=6!mv4i0S_00d$U(1Jv zH^Ki2`Fn$Wi~V!`&TLP>nHQ$?9Db`Gj3;BhzQ_5q;l$U{$8mmUTj0_7z9qrm2YwUr z67|oq(*?X_LyfQT58zi@6Y_O%?XNBl?P`1vVth{tJc_TUy86WUH=s`iJKH%o!$`vX3PpL~b!tO$B;=f72lTYffuTFbNiJMdWC z8NL&FzaL`wB-b~;>f+jP=7ruH(Q|dsr;XhXp{F#* z78fS3{WpJKjNuE270t^NK_9IzbiE_`a~HCk(KE)8uCwd>d?WB^4RJSpTN~`Gwaa69 zof~kTE3!xawlVNjX!9Cm2ZH<(IF9Etd(NKq>jZ3eB|5bokNleJYdqhw)31ap5BT4A z5Kk%3_jUjBn!uAYKYkAWJm3ZAYOZ43y8_LCnnm}`bfXy z8U0~D_0Q!A`()_&di06y>iXL6_^dbBP1o0+lHk$)-d}n7d@KF+9Nk~3Ccqy0Z_}SV z&~rHLu7Y2GhUL`<60{#hJw>o+G0N35PAG}jf z;kw^H5B^ZfGru?Ck)Ec%aP}1^{P`KRTk>&z$O-;JI-6MoVr4`@Hi&kK^@ zOTHaR=s8WhFHP8aKm4~P+SR4w*1leSJ%^T zKz>`0KN9(W2DkOH*;DubJ_0?v0}u0ld2X%YcyiwBJj}uL*Xtu(@%#8NE~Py2_}_3( zmohW{AYAeL#$e}?2l;o`=Jpw+r*Ptz=PzY5F+Q8egtx%|6MOplO*rw}UrXtOA1)03 zFL~m(dTtm;;hL|%3{Nlb&CbI6i1)W97DsdxcYJo`e;3(Kk=9HxqffD zb6yj!IDKvy-;#&tV|--R-#m}@ua?s8y?oc~CVA~&T^8&wob`3De~!(=`j9{O(6^O= zryDT?yl7ecd1Pk9bDlpQz2`I&U|wJ8(X~Gma$RM1Oxy!1xuX^Xh-($FRe3iE-Rc zyi`KGi@$|FXC?5R*m*1C;^zg)Gtct-q@DAuaK-7V*ym2(`+jMj)qNbkPLn+IEYGEx zKaSJbQ|pU3j#^(ol72;f6sOOE$Io}-*F3A=B^eF+=(##4@|o>3n>`h$XMo$jsNpmB zzs&tEjPNF$`ui1a+y1cOdcKXnf7bAfc}mZ*^!_6r#hvaS^!5?XdY0e!nY(^uJo;=a zb{>+t!MEL`>E4alo5(x3QCd9LdnjJ(eS!gKa%{qM+_d_f%ThsTevfj)|t13ccbqsLcSHQc=>pU7vX*Q`M0oBch2TGJzN_nu z8)^4m$0J;0eGC0s6Yz}nk?s#I0-p2VE09-RXMR%h5EsXW@fF_0&llpKJ}!i}c#YDp zErEy6il-&yGwgOxLeFmS*C*PoU>}`dZ3KSV>Cd%Jf0jNu<9H6ZUq1>TU|fy}ej>b% zpXm9Ye!dVsh`gRJw4V1J_MzPug>jd><_p#0K2HiCpo{r7#H2owIlH7_@U}LUvEf$fY;sd4+r^r>2ICuILXUz&kp`6ygiOg?G7jG!1`kJ z&%7d!{Pu6~czbrk&-ItlchIjU?=4P+x6x0}0n6vO~Z%j=w}7J9i{K#Z%#R z9`Z$y*Es6;IIak|^3Mmb|GIz=uujZ~BOiC-*Es6=r6&jZLHhdzaKAqw`2zj-&~D4w zUpV`*Sv%NMxbn~C;IUoSi4*MIiZC{8!vpO(i=p89HhAknWT@;BoTmxlf-PS+&l$LP-myegqz+OPXEeq!~T zwcA3j4*!jTU-gyF9Yy&`uFLE9sV%OJNAt^D(bdkc8s3IS*H4!Q{=L{!&tJ7V-{hy^ z(Ye7Q$nS^ed%RYLer4Qub{h6y6XbP%Uau&>h(2wEL+$x-=l(Tv0R9|sU0+*@o@)a? z`|VSYVSH^L!1ULC`v>`s)oF(7-Wu)o7mi=;Yo5JXIyrvfir?=AyOr|9ot`V31^Jw~ zJ|lq_oO}5ce!Doxt3I^n)tepkuky3s=QlfPKjG`>(;s*=&+2#6B0urVb=CQLR`A0+ zLq1Wvec&x@vxOwT&)elmd%f@|I9^M?4z&pa=8?s}8) zY(O9XyDf%~Fn<2NJj1n*ckpa~>Gi|@)g-P@O4wm5Ji1;G?WTG5RoK(6PZ_`B_p8L| zvS5c6?LKB!-nYwsLUB~0-AQ;J;^p&??64bso}Xy9jh*$}l}c!LFZR^BW+>p(;Qsf> zO#l7xd=mfhajJHSU!Kn}w_j)S+E4ft?fUqVJn>uFGY{?Rnwy@3ZRbmkNBelsWPDc! zdF{79z;}FHi(k*j_Rl>NPsVq(j@y9z9lV#E;l$}`%`X~9{TZGSU%Fq$w<{jSsh$Ji z$6dJUVSitY@igaQCb zKEk~Gw@KXj>rlqSJln757^^Sd{+eemtZCdjapJV!Uw1Gb;&kGM;BVsG z2|jOr1fN8oGZTKId3ICClj3P>d=vhuIMsfUk5kD{V+Z~2PxM2sYnS)aEk(xnS8=+T z*Lc7+$LjZ4ih%37nAY`2qnGE=c{bhmu`kGLp4C2#YD6!u>)IdUW&237XU=zZ-@#bW zN4Vxq?H~I2qLgPHdK)}G9);(`={?xg>r?U&k6%piG+A$cH29C?Thz%p_VjsMxb_t< zN6&SoZnS6mGcL++MZmim_iw@98}JJ9TF*WZaN@DY|1N{^D<1V+k!Xhj+Sk5N z6XHes4>FJFIji1JgctDZxlikZKH67&kbYGHu64~V;N^E`^O)XK_m;m`=<|bc&71!I zXzLg8$XV&R{WQbJu!Fy^-|+3|e?&rmIL38--21uO-ATV5RV(x8!+t1y61)8gT<^_4 z^bI^#uno4C{OI(dHdqzD6;m8vD7x=AzUcJUw{Mt7-y*9W1Z1@QL zFHN*dUCQrzbjFK#n&6ul_iaJ{t?1Ls_}X)2jeiXM(a78NR>QUbvoc{H#i;JnjQp%C zr;b6-YG`*7dH=i-y>_PZRkZUA-SlUh|R8S1%6j5-;T2x$E`Dlc_%;|BI&^`{=nkk*9*4Uy_jT zgXfFrvp%$&V}Ip0zb+R40DaQ$+}#)C>+t*#o<_hGFUnWiC-?R*@GD&>_3cgJ@ zueGy*lC!Mt__?vL$quIgFg)1MOSAf^(ukNcB$7>t>{n(*1=%e|f z=NWmw6~E5M@cXo8|2q72+Vy@YdF{hoMt`3P{A0-f7+rn*m%Qf7@&bQo5KrD%pwE8f z`{~!xz^{CyeNMYRLVDJb;ix2gvF-as8?NKI+Nbs7lF8$@AEjUW9KZ0CzkdV2KHydI zMPJAll2_gJnlQe?>-hf#!9RsJ;PL0nnVt!D(0bXAqvYl1hv{!0@6E4-m+1HVnJ~;pbfEo?k~x z{}J#y{qpls_$GMoq}}bIUFC&2$oqO*@~XRJx1~XT6nmb89nyePcO|-RXlKCX=kk+{ z0oOWcNr9cm@XtEK0x9JMQF#Yh9rGLd@VT(YnCLG06|3&$E2z>jB|aumgPO^Ox`q zw0j-=K7R=x0snS^p7PJ-mHa+~ zufhIy{-^MSzZ%~?0?%{2H-6z7$2Z}(BLQzP?mq!r!Jl5nlhE!4cvP=>o(?{%j+!6v z67o7P)UMAmo-+BO5Boom=6U#@`?=-)O?nRScLV%(KC1C(eOJW}x%D-&!`M^%kXr(O z6@C5~c{^@vJX+t~N595`yz0*(@U!4lf2BwBwPwu^(8u&QocTD_*V~3C*!d847;Nzy z&V0=7BW>#=!>h>aIhfX8!)yHhli-#&4X?wa^Uy3l8&1BUZf~!f49BL`PVC$l^y#3# zQ{22T`4adA*v*au8m{@M{qQvKu#Px%Zvnr4_kC=Aaj>W6<0y7EKR14jt^zJ z9AC}Hy#?c52Y;+UPx3|o$FYBBu&3%GjbpU4@`bMViFgBC$8jRgd8*sr=;Dm^_aO68 z=cQCbyIMz_UMcW@`t<7p{a;&A-aK+%Xy5DrzH(!*oA&lL^4A-1-G^-y-13O&seQA{ z;PHN_dC59u&|jY=dEv~*+;JAeH6Qi7;AyVESv*UV>P+mA1RT3m_1u%m0k1K?^qgxy zkHn+6>hJ4i?P@+wrrk+_Ct=(_UC=J`@%YbR&*Ou9iT-sF2lfP9aaGSbjq=Km=M=O% zK)ZV0R@6uH@uC7;^U<%5tY1~|e`bF51bvhre~Nvgok!?bwV=P7@bh;w?p9wJe~op} zq=J5_e$#R1t3tb@@bu8`u7GP@e`-Oynvdg(x5(>d?e50^{rq2 z$I)kLz*EL$Re?SO_^r4^=XV%I>Lew~k@eSm!qPsNiopC7R?;JVI2>+7`vPdLtUaTs6m=&up% zDNJ>v?nk#K@Ms<#2X5;W=`Z;@xb9DWeUNYP`L^KCl23?Zy#@Moz*no}>o>`l7)R~T zbn@QDML2ni`|P#X%cc)?V~_5uFgeIGkGLO8J}-zzc}mZxQ{49a!{}+pg*{H)$o
    z*WC#|7vN#b=-3lJFE{l^Mdoh^W&>_JFvrj z;Ol~XiE;b|JU(7Yz6`GC=>00lr|?fH&~pI$|Dux53*ynd*hPN_1J5x0FBZtFep?pC zS3Ioqt0xA(6+Xh)>iF-Lz)ybRI`<>(P$S<>W_#ed^$|?IgX|-`gMYt({Al1Qfxi*^`}r<;{F&pJdAkWuvFH2( zdHJ)>=Uf-~wa>F7_?vi!Y4>6DiTcZ*ze#_69V>bKnfu=4^Q`a@@>o89Eki_pi`^GxoN#kbvXY$0!Hz7YK;PS&Q;BN?c760+)kxGB@)YxbC z_suh0aaH#@@pYZ#<%ebPQv-j(xa%AmZwJY%pE~cz&m-X_?C?eG>FYz`Wq7K{ALhO9 zZ_2pL4Spzj)f>7FdP|U(9~SWs!=vZf`8+Nj>W%C^y-^?48{27jYv3QDUnf`E`|xHr z)f;P&U&Fs%{~G%IE%I4gt@WnfE3WFkcHU2fkMjH5k@w@+kPgI$Tad3he&MPcuR-4X zgYY`zbv5#~-zz<^z^m9_&#Se(Z#)`D=`tng!+c~v zuw7@GJoB-C5Ihb13H^GUentIN=bQsRH}IE`{|tV+EAW)D^T+s&&39{;ywHDnL4TDO z&M&}+$S?0;T%sK`j=J8*;;Qji(NEVAFAREWKE9uE%*NW;e*`~F3+%9o>mc+zf+OsV z{96k2QC?WV_}cL!(`OW1#}hXP`|P0Z()m$w-0R?17vQ_GDo#Uc(hPy5JvA+jiFX%4`3v`4#Xo^F?*$ z)PSe(=>8H>PsXq6*G<;0=8?bOj^V?|UkJ~nz)w8ner`wN8}J|9rz;8aBj~@I_#Ex2 z`s!2Q8-sie{E7m8QE%vwi#Bk*of8MF#%!~2U3iK~A?m7=@GJUssBt6U6mTp#3j!2c2S+!OFR^1oYY*E7~%tvj?&>*uNX$I$0A^wE8s z^fLLqeAe|-ccuApQGR(J`q=fYCa?T5hW=ZFJozO%@3-#ol=oWK?4%8in~l40#ee@? zD)DFX__^*Y>*Jwt`T0up`Dymsrv4cZ{;Yj+pMQnd(xmF|Yh?Ud*DMV2NAjwNFAsh$ zJfUBIMSs`x-s`V=`0L1D74Q=NqxpTbVpTZUYe}PUvmrN_K2_FGF0sHp_96#T#eNpevlCNd<$=8{} zx1r~k(f@OShdPSqRvf91{QM5&eVi8$>+AlvFdv5k&u;jC4}b7=spM5h&BxDeJ!1CJ zKB}(w-x+us=&$4EW%$evRb=G1_qY5_pTr|v`~Y31moJb{`TIcdhl~e5 zoQa;c|6u&YXU;Rp>o0lgmd(2EGW%_7y!5_7U8H)z`=Rgzd(N)3*N?_8zr7iLU*Af; zgdMI&-mg1^Q@8BV{W^B>-rFZdenE?e)+X#T0Iuh+`FdVFL-1(bXX`lAXBhm+;OCOB z!t)6HJ}wF00Dcbqxq1>@@%gw)J|0S5@%d8t4~Wy*tp<-g6Jm*1+Ewf!;EQ}Oxr z@JIXXK!2U5>gSQ#t<&zNmjC#6wa>X3-0oK({u=V~x3{9_i1RDq_}jRSkJ!H1ggoDU zmUf@cWSaJuJauxDTH=8BXW_)jsndu9>q5H; zygHs+4|oT-uCwrYO*|#qR6g}~6J7?_KCmAz;jA}#{#SneC42yR)vLZf5k3Syv!Zs| z<1T!d&&t=Ed2iz)yb6C8xLtQ|c9y?=3LZc1k{>}{_lr0fcs3Dd&MM%kf$P3-PXu|@ z1BcC5}&oKPnKZU1^ z?@uaue-kc$coKOZH-yU{P7mWETzT^_y83t}yqd}9$E5+T^LypD%bgvB>->so70yKn zc!K>`A#dkX$v%=#vD-@Yd|!~Sf?tcCzV4EI4W88nJgR5({4yIC`FHb09Y`N8Z}A(h z;~F|nVaG!aPiRBWr#62x96iVV{*vKU;=qmMoqm3&d1Scs{3&**2E5LF0&gvlZ_wX+ zu>S=?&xFsPN1uIx2mk3gf&N|`SGYmc*d7Hl`FaObfW-KlluKg>2Un#>!&_~aGu)5B09jDXz(wBvPDSvHr zdEexe>~q z__+^x@3)ye^XOr4KkmX){Q29oYxzt4&3LH0&LECW4*c@}A@beSfY!Ac6j;%z8}0FKe75ydM31`m+FlUCof!_IycwW#S`)z7tp#XZOItUI-2uH^LER4 zkn1dvuMwZ`4Sp#3I)3i&M`8T<^JYEAc5Y}lVI4QE09QO8ME}Dp|uOe$|IP1@fv7 z{q@wwvkCrv^vnC9`h`F9tgZI-*CxLWdEF0od+6^d`s?^+F0Ua&JWuX`-_J+!*U{%k z@b3vcyWzPVep@eFyNc&uqF+8fr~Fpy6g^MS`do_RN%MzP5i>CE3*C4 zTK3y!9?5u!!-p&RJSCjC%JGYQ92Ty1$|USy=b0P7>WVYq>2TvE`Go$ySfFPI@#J{s zg^#zAFCl+f@E_r2cy-+N8Q$Bx7EW9x@3-q=(`NvCRtk8A76 zo+|q2JU}0>BriYw4()E?z1L?1Jl0W?S6nR@*sX@1Gg|F>9{Hi|hj@GxT>D6Veu;kv zn678^^;H#IdHGcA`~dHbN4WL_pA7RNljl3#*J5pu*E;t*@NW%x6@SqEv~zJ78T@Cn z?pyM7kXIbjI?>N-wadOrkJih!za;x-`e2{ike{N@S^jI{ietJi(vAz5d3^ZUFx^T(&=B z_y`=@XY=_;?QR0sc}lhpm3$3cFU6*-`E2$PuJ~{v{^{2#nLPGC75?5JpMqaOzkGho zco^R)1?_5lzYmYkBa*N4x3SW^C!NNr}B-xHaWW)&U~3V;O2$l33fOE{%czH zG+guLD_|=EUS)mvCAZIG^0MbT*I&c6-jtqJ{}|pt{ucDKIB$4DzSeU_WGmUy zGx}@}@*Cjk#Llbu*Y|4#p8LQp{~C|-{)fT88sux>3k$|e$4_-$%EBP8xuv{zcffaG zPd(@K3jwdgzX*BFSFh)8Sf@3bHFl={GDy$siWg0B0sIBa;r z{JIW(Y(5%J{ki8q!LLu-`eb^sp1K12+xpAoYxtXk;0{f~fm1o;s@>paOV!4C2R)$RKNuDpK|^U>nJ=_5aU z5d6^~KT5lQ3IDSJ-vPh&6@Qz5n|e3@#Sgc?fpPTqO!%#?SE|#-bxvNm_S19?p&zeI z9)I{p_`j9OAjj{7r^FxaH(Fh7xaxtVk{@5m*YHoBNA2w|yiPpSeFnV03CFMcbzE#c z@2y|L6YMspawLBTe)1yv>6P~%j{yILW>*90yh~KK7d>uS~zNlU8Q{4fci>u&@Zy$i) z>LSx;8|~@7aEF{eDetwvIurYN{}itJQ1@f;c_))+-2Wc`JdnvW{Xuxj`2H9C-rq7F z{6Xh$SUn?sgxA1Ngulz#LAd-u`y}Gg%j6Yjv@bCy$S2vjw68lhdHI9#xSto|FY)^l z`gaB%>WA@Dh=&`IF@KZ1{6WVNte!M}#TnJTR*xAjf9NmJQ~u!Z+idbx^xTDf&9z(Q zJ+G?nZ{x>BxZ?jp`qks~&*ZVsFnV?eT;rmB&ZPm@ZCX-h^_NVaA3-a>MYYTA2 zt8Lh)9OPAxeGlvt0aqT>`RbzqFN5oN&?x~=;opzXSv)sAH6A5+ZVvK8$e&NYEWeoi zFdRCb{_|)baMf3b1FqlgN2d=3{t@^U=NAWj6SDKrXIsE);GgBs>MqkqdGI3Iy~fGQ z?-Z~87oNJ)C-X+DS554RKI2MsAW&# zif)`J)z(>~617+W1(&wLhTaaklQ5kjI{V@cc5!r`TaT zJU0Yf^F`+_-yd+rPhBq=G&S=^JnaE<*W>=1GC4%cTj_sKOL zCNI0KM1ED^ufd~zMLRBO^5mWEx<0pAk1b!YWNB&9vX$p9@9yd@9s6t9!e#vnj-6(L zVe!06`t8FkdsCXxJ!9fSc>;YoZ{{uN&-iA}>gq0+_1lSe-90nAW;U_8Cja^EiUJ)5T%J@Xb-^#P6zpLk! z619HU%%HLH^(ivH^oszVy=HgyWXg5V z&2npZ_ncW-ikZ;Z_?9kw&u^l9Z;P+HYkKz6@{Bo6q-T293?bb!r)Pd)e0|rhUcKZ4 zYpz;#MeCROe-^er25W)F{w|)s&^l9kHR^+bC6`>ZU|yCoX{}tkcJ<{S8ou_Um#?~N z?eMD0SFgSDL&JGhkMe2WlEJ0RdTn@*oo{~cvfiZ&UOoNhTfvd_BlG%+g5#Bx9`i5l tUB1A&n^hcHbif6UVyKt+dcorPFY(PQgDEAUyrl2K1- literal 0 HcmV?d00001 diff --git a/tests/files/mtz/1VER.mtz b/tests/files/mtz/1VER.mtz new file mode 100644 index 0000000000000000000000000000000000000000..66a1ad245349584a93368ea57ac883f74cd901bd GIT binary patch literal 87264 zcmb@Pdzf5RmG%#}aEnkFz_cJ17b8l9!a@W=*xke!pvBTKNQA*cm=Xk;a>5`Ch=+?Y zLW|8sq$5Zn5I_;54G4l*5CoBqpdb=OEsVA|R5U0sgZlm6Q?*j7m;L8g&r?sIdV8I< z*Iu{1_CEWZ^Nz#cTONFOmdzW?{y%>b9A7`QZ0sTa<_Gqf=Yx&qe>$kqI3nP!!7M-I zAn1Cyol_aiI&Yr0Z0z8G*P#ExL8~U7NnQR%-yCdA{^z`Eb#b6yI@oCJd{DKzCeU92 z{vG68FVNpOxN2e+^y7|RZ6xsLhpN?Cj$U|uP<>rC_VIw{$ay4uRt7$!=wTH3p9**j zdFBqSn&7-bZ{n@*3V08KAKt?=&eJDBzpd++;gaVG@XZ5#KDcUl*MpEX_3Xl5eMZ3# zMxMLVRJ=8~YUENUr}19_-?Px$o{qod#9rAykkj)KUPEtxod$0XE?a$GXjk-02bT@( z0Dh&DNB9-cUx5ES5OCS+QNFu4c_eEZtd75e{{Gg{3$G7WtAE7bW(WFwaM`|_%*(P} z13t>{9fq=OB;eR<`FT;F@LvJ{58)@f_2?6rKIo@?4edUgx_UOe2HlO=?8HFd9L!o9 z&+GTg=$F!0-!H@E&&lc2@Eeh5TiR{8cBN1J`I$|LuL0Mt@H%q7k#_eAcn+VXPXFRB zdg*62dKe4*TlDEiLygAe0bhZfYw(Ag9iDMqL7v&f@6N7W;Wg-A@BClw3U6ZHU1;|v zj*I^EaK-!e=znd%@t?7eBLD3HmwuMu&v!apZQwr}tUwPpJ9&f?U#m9`atfDz%ILWh z_z>@Na@u_&;4SFy3U(Kt71;L&$iJpXpYvBA>^1T!7mtS5>F=rN=enMphAaPklsxcQ z>gvbvQR3?{_-q~MiT8=ek!fB2KVx+KKRe-CG=}$@3-NK_gmp_dCzdg`)`N(`B`|zw=!5=JK*w+;nL5G=yMAvkLa7k z#TQ+EF#4s?{eX5SppkrrUx7SNBfrIw;rP$+KKSjZqfZe=OewT-m`ped|KcWL;Z2b@D-fjMpl!j zO`G~8@c)ATvL1iKWye!ooElD^96w|ld~|Tt%-it?ix;C;p8OR2t-dn6D^In0oNHJ8 zs|+qX!~41Dg=5FY3(n6C*EsSaa2pp4m!JRBjTeS1Pks$MKGfrDxN5Gulph?Pag2V_ z7T9HbC#P`oK;tm{XGNgL&xb$c{73vnFMG{{z7^h{q>s*XA!q->r-KdB$+Yk(n1j@AV))svEzK-Yh>G{f*#{gD-LI<{aba zldr-LZG0EK@HE(%v6-tU4X@)rA4PuOFVV{{e+VC&PZ>RaH6i}jiUvP0yoH`OK>kk# zT=`)Y=U5zxzQtF1lYiPcYj_PkuY`}~CBvKO=d0L7b%H*IGaniII(G5;7k}w(Fvuxf z&RC`j)}3#2+lb8IHfTZ+3afaP(IF!%&v(!+E9${SSYe*bY4}3iv4U90|SE zfyPJv_EGG(caOdU4}CPgd>lQjPF?e3>5_z4Y^h z%b$iT?tVJd&!588Km4lkL*(&%gr~t}iyjH`2vFm2|A3AiALX$s{ygmZCHfX}t{m#;!w#7673ldK^6CbTJ_jcrXZ|mKUJ7`Pdg&5y z8&AbY^zyg$r@>Y4-HV)74;a1X3tK?n()m64#YgeD9(ovaxM<4Ac_4PNJSINE6Z!2e z@M#D7I`Z%D#%rUO-qzCJ#{xb6HoQIU-V|`@?W~|5wYvg*?ogJU;PfLIO33+d^w-8a zwJV(Xz4VpHvrC|7+?>`Iy3TO<;VtA1%ZG;Jhs?XD z%A^042OCS)3-SoZ5375k=e15A;dS`$K9prPzZRyi;&(^%X8GB0`QZ;GtLv}+*Mj~z zc%2*Y75sh*`M2%i8GrRD!)IH^-|z&F@VUUvTyd@NrXj@>6V zL_dEE^eyR=wvP?`m1kZtlx62TTz0EK+v)ka)nkT} zXBJ%s9~<8dSG+t;o-zM4T=ngx$Yb#+JWEBqob2b9Q3+U6vHSrft{LWe*fAe+}o}>RQk>~EfXB7Rki6hIO#-|0pU-Ca2 z_YGJ4egrx9b#iv#i9RR(JNUow3hjOgy?H+ujy?z8gWXp;d4xB?AB9gT@L!4^&Z56o z&r2TBYrb=j8;=Z^J|99RtCNLyzt!+Uv%+hc#ZM;ZWk|xeXTRTorJ|^|EQ&WxT7thJM-llG&k*F4t$gGrbwEe7GveBfloDM_zJvH2NI=-$oBM zPceLyd^ifdw!(dM^iUm3|~Qi?nkD%zF*Lk;PY6JCj+m5f5rKs_z15dr>`Fj zZ({eWr=eG#+#UW_M;X21`Y8O<#uvlMlle}_-=Zz)MR?-44893*YU_c9C;q=ad~BX# zIQks+`NZ&?cCW>sy*|Z%6#g60ZacIqeGUgXMZW@nXbj12eK~V*+G*be{cjx~;fib3 z304P69^p0gc1@69IB`Ap2k;#OfBbXu4EU@J{E6#`0edZPPMB% zxtu)U?I>L99-GkLBb**O;2H8%7$5fucp3Wnyty8=CWZeJMUTJ*{T z7sAKtBI7ekyN`fd{cLy(`iy>Eu5){Fl+hOYoHzvf#jai96~<*hzY`zf=ySQh9$`4~ zzV>GHZ*{og)RSX-!vDCyU+blR8tT^}Mlb*TC$#i?ot)kORjcv`tN(@RTgK0i!;Zec z!V~TO9=)kPvHl9L)8Drvr?;cVte?4WyTYm8CZ9xpo4*wO1;@T8Z%o{o-3`x)quY@u3;ajjGZGA1fs4pce8QLaOty+oNXt+_%Lo}hoDcZt3{uK6EBN4b>pVt^3P+DXQksK`a0+R z1l;nc(Ua$gX9c~9UjF$UerR>F(YKI)4)R-`7v6%t#QM@6=*Q||!z;+IoISAAL;X4Kt1)5UsxC+lJ=#lTJ?Ib>uv38a#*ox5#Px9E^|h!!r1PIJ7JMPxAeI53lf7AL^6g zA?MGAqvz#6q+Qd4;Wg+tBfnW)XLu8Py%{}h9{4Y1T==T215?ml8H5A@2b9du;t zilUdDGqAUz&tXU3wSS{=k<-7?qtE4^reBuN4M(38hn@~LU? z7IJ(SJ)2*NPX^BtdF-de#U@Tp;T81uQ}k);F9sZaPE>*(L|;Sh4@Ucf zH^FrsaIWi@=*a_Xmrlb+{&_TV+IpSj6us7M{sbQz?}WF&GyLrn$oUY*rU&6=^#5)6 z+c;wM=yP}yIeol{9(_(+4EUto`dPH_kZE|?bz>!ruu_$%~Q1B#`3wzBfou$?>>A@ zZ=#pqUd*`T{jdWrzdbSNS$LWAAA{cNDDfAbu*T991wRozdd}+D{TL^Y z@KNOdU*z|C6Ry0fd~5S%@mI`t(z4Z0#&`$EpU3tI`WHR?%`@Hre?yTVKCcYS@xZ^A3+=kN4)cA%%;sA@gU$Ghm8#KmUFzXsp% z{AHIE`Xzd;%dH)n8rPCFPmMV*;wQW4xR+Ojw-dR9aqzi-XZZ8C;qUd+?enUUF9$yq zPQ5qwYTCU!@Tv3r=iojb3we%L~oXF4If3GHSn?ZE5lpxeU|@S>iDaj z68`)f`eoy?=(~1kG%kV9V4$zDU%iB1*}8<$H_1PL!cS~|YB>3D;28X>-SZQ}wGQ@W z@{j4$aE&_~)2^*w3C}n;BVK$w8eV1|^gY_OykvO7j;{x|JY#qr{ATil>T`Vz&%w7t z4^~$jF1x=2d902#yaoPR;_)D-pOmD}Wn%|nuaz8|9feoW+ZM!~)v@9uyvBIC9QrkZ zz6qbZu&>oKMo(Qdto;U^K(BS$UEyQ-#ORgx9!Gwg#|h6lH-pa&=qKXnbNN%?-hV`| zIMV$oUQXfo=dk99mJf}8PJTEk=t1oIKFVVm{$usI(W{=<{@!`~ zHTm1n%Wo&KqtyeV7ml7E*k@kX|2%zSf5pzkm#zO8z4E|*!9T?ZznwfH#Fy|<=)XLq z_NU6z!p|Rbr-DD%@pG%|)UN2sGmUr9FP{&ElV=7#8u}|-^REu_ zPX_*)FC5a#pVn>}`aCTg(>_Dl&f|>ZC;un@Z1s=i5j}CS_^N5>%d974#D&!%Mz6mk z)6mzEN7pN3yE*ycU~u_`muD0{*CW3lHziMt{vJlW+j@f9tMJVa)b$9-+vg)ZrAA{@ z^ybGq;nZWr`n;`^8z189jGgFLmuENbb&MUyAApa~Gs4MZ3wNMSuz8&EA&(7z5qh6* zMc>5keqV&qFXeythCC*E#n%PI(QX`@esqpLng?0^Bdowvwfdj2S zKNI;a?-^d_ypkJt3|D-;%H?`^VyhI+} z1Uaor^a=kw5W8EQV0az99e_Spa;~*2KKOI>Vf1O&%Z#4! zW#GV|Ptmua_xmD5pM!V!9)vsx>vykT!b|8u`}(|ogv+0oM*D(G&({S#i@u58bUk#6 z|A5P%{dIhkNB;aP;?e4D!*3)`*B|O%uNU5e5A*a{etvB@`tSHWX*hnq{MEtF#V3(B z_5nY@=|MPt-qw8zHa;4koORCg(2vy*hRe^NM9zq}it(|(o*?=*{PFYgX67b96tK zj|LNa2EEp`px*r?u+SfQ|I>lGZn|H zhTnla1y6@#_whTyvq0Zs9pLleJ|4w~_#M}}jn8|+ON_H)(0jiUPF_8E3wPZ|atdc$ zYkV-|f8oUMOs&s&y9n3%z?ZP^)*SnOUBPdyA9^{}MwV3<-;A6U{%7mRFz`cEi zWA}j#@c;cCf8p(vWf%4IAg;m+K7-31IuE;Ry!HIcjMp)*iaw#AE3vPaUpRIj*ZA)3 zD_rqDNuRyH376e}!v8K24R);aS0Cg|zeaCXpBtVrUhDo5>zCn-7sIc?p9ce<#JpgJ zt7DCx@nTZXmna4LoZp)fkJkr$l;3(Tgyl)&(*oChc2_!j^`k@mep?1ZmiQaA_kJp6yNU&I#J>*gSj#_R8( zf2(IipFtx(xq7JoToJ=F+Wj+ruIpv`7+yxs_0j)8Pd>v_QSbHVHHK?k*c5p-4D{4X zV>hBtt*@^0ha^AAvTxB}uLsH5B7eRg+^&O*e*!0;b6S4ptfAvlH z?Ih<~yW%7Lf70bKqbIH>Hi3_=pBPSDFVg*qRu>sgT<5QL^@ic9SKsLJp75M@+xY*v z$YXiQ@C^T74}bP?A$t7(=6}b3j&^zzp7{NQi$~*w{>Q#UyW0l(oP4GGO}rk&XB0Uv z#x7I*6I}WEJn*@WPfIx8AA>7D=%aRpbyAl5@zL-Saq8DI4JY4D{AbX!_^6KkAa=2O z%;+`Fo;eM@>ZMbsp;x{10DY)B`P=Z%Sx-AB&oy?+WXQ#ko^Te|BeBYMS``a6$fkIRl) zfAe{jSb?Uo{AKiOUyk(=f8nz4eaLC!ui-7`DO+O~TlX`ZJUO;L^`XxL;=?%Aco=*6 z{4AV!ocku|{dgg~hW~F3z173UpK)ruLc2a6ihe2IBSPL2u6TSA-0Qy#PW`j?RoGX2 zO+N*W`~$?JA1{RGDa$qQkdJwK#qWBdzXc!Uyg$URaN?Ky==<~8dV$IAS$ljM`a1Ld z$I*lGo^O}D%JtT6?0Z~xuekc(__sL!blN@B@u?KYt0q>#XLi67*TEmg?%s}~Z=r`P zLOu~reKO}YL2tsb`;6@wUwl3kuJQ2~_~$MhoBYC+hy8PzB)@Rw;X~2SN~eF}teaJ5 z1iJ`VeWLrhd_G6VvitZqkzaM4KH^{CY1xaj@PC`n8lHo1O~1UHqHnWKcsqP-o!sa% z|Gz)@zwlAMTc96bKh#8j zVAk3YKQDKWi~b5v_p#kI*k4Upkk# zo;VzOKTZkP_Y&wQQ(r&AntZ=Oe{J4j@@QUD4gD3p<~4dQ&9VG7J?k9oTm29GT=!Z` z!HLs5Xm@_VvG16k7j$yKW#7+0@9k3Xr(gFVj~|bO%f7F|pYINQ662Bf6P5x_T#vm6 zKY29Z#P#rJkjK^s#9QA{_-shK%FpUolfU{f4{h9x-K~BzoII85`r{^no^kfEEqZz0 z^Qpjp3H`M^Z+tY)-r@3(;rz|Du50y);j&i-ZuPL?%ENEOUe`Ez)V{_|zpq&I!igiE z+tAm);nMRE{j&8R!_jlQgPgYBVz}by-Pq;)&~6UicH^MYD}UY!Zhmff3;e1;pQR#R zUex`{+#`Wo%(ITGgu{`|k~>k6Y+UAnW2 z7sIu_t$jyFIzEYZ^OUvkfZoP&(F@1_7yJFUhG)p*uOk><#$G$QJa0JlOv7K7H@r@M zo-s|k*ky8l(1YX|MgC*S561-gWfwg!;&g}8C&k6&9_YEFKK9~8IOEikcVL%y1^N>6 z1??lU{4Y60UqPO;TwXO?cFB-`il0E=q`xiNwe=RGSN-`zI%cLOZ6oBWqI3h#+|V@ zA^&MkZ{n|bIoZ{jM&CjY{{#JDj=l{&*RL1*d9>(-C&nE;C#p-+&2Ktxr7Yc#JT}f6 zJ@GZH{WEg{f8uLG_j&CHu8+}I*za^Aae6|)vEy9rs~5fbiP~+Fhvh#}9@+7Zp{aRQ zG_vC_$PZp`9q=rZKlkSslBdAQdwO1k&(Fdcj|TKSK%YN_x5@KXyEym?CdYCf4fOpW(eaPA4*&PSLsWTh5202Aj!H(@|_^8hO#8B5i z%z3H;o~W$9D6#x-5%x}Co3BfO;l@!0@AiDTnm;P|cTZ5?|&XFjF% zI;%5{kL>RESs6~-R`R z|5&@D=u`5jzeZm!&eipc{_|TzFFYatZDAZ2-eNsO>%gABaN>HzUtcpm)Ge3Ri1#C) z@$E8?7}t8W)zwB{gKm;|v^vdj<>y_)_@Z_-k64I(PYwKK-*3WaM~7E9&N6tTpKnB* zap%48_jVM0PJi{hAm3l%*msV!{x@5nF}#Mo6sO*wMKAkmUq_T*>jO8ShbWKow(fiN zby1Dq2*f77tm#-^?*P+)sy!Qv;Ipfs&)6kFdKjkrBC)Y)P;HqKOfj(Y@C+a|d-LnZM zoc#Qv_QUw`S2*@9_5rS;-=;U=CHzqL(RlfVSNQttp2lDHePz&_=$qi*!avI#d;ZGL z|3jRf;c)TN_07Yv%efIJ-<|}Y4fz@$(G#b`??ew{fxZp>YV_>my5P@xot|S*3iM_C z_LcDQ^EJ^Ykg+d9-S5FrFY#Q^l-ov@;{CIivlkHlzsOKc#eJ)FP0Cj-BIYy zr(ZT-GrXldhVDM^_{eUIFO#eBlg$I(BG3F2{%5BmKuMxXFEJ+IHk5yR`O!|OS3HcvEM zf8U5doXKA=4|T%C-@sP}Ia}c8qo1;)PX$eNd=dR!1G1i(Q{GT;t2T(Z9{p4VQh_4|b6} z8eg90oWqrA=oyceY(pH;?{#_*eWITH2l89pV)WFF z?8Bd`XVJ^<1LU!ej(dJ3oOmDJ2>ss~aM}Gi&N<4p+dDVkl6unWYPBmo5x)n)KMV9N z^1Oeai_v!yms{!Af*TZ zL;J602R_Q@_w)7nN_`>TM`p$GS2%UWVYKVVFVW{o_xbkMbJVWzHuZ`2Dfw}v;Ez07 zKPk?&aRK`BU{=;XRUf~ik=@US|CDhAyv|qGr>vgO`7Qb!JIX(;4mW(1b}t9_c0~91 z#;TdxH&rfnu~-72aVUF^4*H@4$!u<$nK_@h6`b>&AZE z7oQs6AG>u4qgOtE96S1PU-Yu?@93A+=SHu6bh^*Zk9TG0bzMp88>et={VH(ku?BkZ z{DtTE`62Yz$A$2AV*depQO&Wn+uc&kqC6cs+m0`>OcA=v&zRjX{3l z8n5+U5}UuPUEw9{_+p4t;T7z%1GwrFeT=@w@6C|&H2!)yoBV!1{2xeDagEnkVn zzQyy~Hw=6deyIJxS-|V$TJ_7;g^WM-c4I}zSK_aE#0HE<$2&gKWq#nYG2Qp4yrGZy z2q!rk(yzxr+R)H&AY8SHg*#8cMUIE+W4&lz{Ng5Hm_!rS0?;#YqD zQt+o=yTRwgz`u-LbibRgt3@w6UXA?c1o}F4i=M~mvg6z6?*$y29d(ZK?d|YCS~j-&Rygsf>xl~jeUAKk zPRJPnZ>KE33O)FGui(%4xElTk1bW%=3&^>7Xg3jG|A(AD|BH{}@paIDGw{*>9s{>| zit>>@io0K7$K!hRh9|~BJ$LcQfVWb0q4wJ!6L96#myq%NfXiM-VXyrHE_)^VYvYc! zOWdtiU2XM};mu-SV!tjmT;t~VoL?EPym~8bo1TTsZmgrO{)Ni}>u|=+A>?@?b?FSp zj%&65(BjT;;<52La2rPq&saD1>r{qg$8pW0ZJo{V1g`h0#CGe*tLML29ya;<0q>U}z1Z=xqo_nnK~%|DG#2mO2& z`7Q1YSH9KzEWG^UQ(<2JZRow6!WE})aB*sU6qBzE{T03J{wwt3^(kEQpeDHGTk%QY zs)vsuPg*>#!}*_oKcL|`ak`%ycMNZXYkz|8SHT~D{y2WM27c1J(KDV>KTp+<=*hza z%7>9p9lh!KNWNd9SNyKRzTOXo%kF<9t}Sm{yDjvye$YeSm5up$KYxl|IQe;Bt*hJk zXt?t8_mRiu?}m3c=iVT{_>}OsD}o(`tB(3BdhqcsoOtJXJO7F+KfeWf#f3iBuKu3v z>J7tHC%=sT9 z&(Oc_v+n3z*DiJ+xdA@5f70j^^9a2M!RFnD6YsfSmo_}dPtJjltv?t(ihkyx|C?RA z?c#XZfSv>6^OA7tsEONn^{T#pbdJwL--U<5KIW~O?*LAV+kl#AseDX_?XV~?tz*E-#75LTx zC!be!zrT+c(YMiq?z8azU-036;a8Feb`N~Y*z3BWXVGI{*3MQnz_UC1W zW8X3Df3SIw;hN{pKtGFIyItF@8qvPo4FWFvs-B-4@D}UvdOqqJWH5b-Plx?!H_@)G zuc}?)C9W^TJSH6bPJ9!6`aB>U`;xDw>QlJ#w%%{$^(I``J^!12y^CWnkJc^k0{3xR zf`(77d7}5L0#8{kJ@_~kp7*XVMEZ8BKDK*reqrqveAtg+^A3-f8K+*tzFvOOlefpN zBLAG{(eu3GxSRKc zC;W3g{BTadTiE?C=)wD~=!xIzrTF1UpjZ5ApS;h*qOVY2>3WmZ=O#}LyPpVd>o$fr zIcM%P^2@%eOMN|4fmYYAwU2jCCr^R1uCC`N?-KAFed#$WetZMr%y=k0nl> z;AQOfgK2Q=SUfMt$C3Cj{#NxI09%Jqe}yYfH^#mv>9}V{;iKr&?@Jeb>e96Jd|K}h z!V~>!p`RagX}Wa+$JlW=Vc#PHPVPv0PKeE4j6Zgq>DLbp$BtvVUReu#u;cK~=+oA9 zjb7vIiO6$ypx6K3LVwS6xOnnOE8*kos{&7r8M_c)gMprS-1l|lCGXFoZ`0q8g}5&G zAkT$C9^u%rm?wHY2$xRxr@y28HMi_u?@xj`Yv7j!yh1+L^TBL=PyH32 z8tW=2&@Vqf60Sac7Cxh*bMnhxdVhf*_r**5`Ywi#+BW(Er(XR%`uFiww4Ji_6>vX~ z6R!NJ>(r;Xb_+hV`33ByJ-eQN8M+XYX4=?oiti$muJ*PIqmArabz{|+v z*8z-=`laV{T09!Ae4_jKY`id>`FQne+PyOHA0;0BJ{i%gAM8gR|1f;kq<+7IC+e0p z&VP))mAda|vGLt-#`ncPLr$xo4X1wNJ@5VX7Q;0zXuq!2Z-&>9>nG^h;@9vd^nO2p z;i}K?qQ92M3|D=AE_mJRfj(-zMRlp?p9#kgX9qn9$G*cmO+%00E?^HQe= z(KoTzYv`}#C8JmV`6vEh>ko!&e*H0Ujr;lt&y$X-YxUe_Tc=%zr>rwXUah%yMW3UG zmj(HSw|U-$-a|CS&%rb17yIM4ULMgij;y``-0B%?7yAzDeJj3SqL*Lp6!cTb1OGQ6 zkIlb~kL|51)bG#jechah7HVpk0UPAwRudCI` z#=k=QD}#TEUiQ5fIhEh^F?#9&o^R8iFBq=<1DfC2dV=A~vkS2A`TR9MZzplobL4pU zUc||_{ylx-Bl?{2NY8CvFVMG{Z>s*cd705?%r`FpxB0l?j63CxgMP$6VXyBakIi3< zUUkCe)3mGn{1*DVQnK~pSA2BdchIxXpPepEKAZew^JvitPbq8bd89sn3dg>@pKEF! z5RQE(z7Fo=T6hV+y%srr9uQukPFP64yuS&@z5`k>^?6CS>h@=`Z^RY9kI=8kU-Rpa zU|*{T)UQn6jB!86S>Vjee&Xs(qbGjv9Kz3!NIOtb;Wh$u_<|8=Ys#5entL_zZ0K<-s%?dDHUz4n)o^T_wgk> zF@KrQ{Ke`nqnCZ(LH%IYI}GoTucCj75B421QB3g2sy=d=7Zy$NUiWbLOzUJ~Bq ztM~f(ydhlny_f!O$g%N3mwXzN|A*eZehQrRpYwx0h3Bk)t^~JwbreiEcAWD#xNo=M z!}C(~o&_JL!pXy9D}$d1PuS(C;LpPA#P7S|@AHrFobPk^ZR9@+eFh(&&nu!oFl+0& z>aE@5BA$h(l$BqPe=Z1kD`olf$m#0=(GyGQEz{7GhsSS*zt^+q6{qu%-^aCZ<>A>O zu7xX3pTb@~ZwOz?cL(Cv=G~+GPJH>~W%}#=t-y)LC*l7(Xw2@ysh4KH68U{SCY*db z_N&m}f)C^4CgA13N5A!)w<+}lnC$p8avl)q>%`YS&?^^Bk%#`iZKyxrw0`MthjWG< zpIUL;Sgz+++q!}1gr}6{dM}=@2ZXn%8;6Mt+sA119qPkRqPHCip56Q|dg`S)nm75p zD!f9T(RE30N8x<)(}TYWZ-VcIKRo67D_rrY=YsnANFAKIXtAEpd>qH7p8}_@)Ahl< z15P|vUxhxs--@1iT>c2S_uGOG^6UL*-XDaQk>`2*ZI{4bcGUZBqMSNM@4q`U&?~QA zM?Cs`hz``3`T+ix9KG;_AGXn(AD4xg=3e-8^BKtc!vJB z!(Jx@TzT?W#6>;ejCW&OP9vw{>-EqpPniD{@?e(?yZCmyG}ZAxAkV&zUU)*zi@<$+ z32!maYGE&1PaXvm-buPor@zi=xa^|m7VZd*$s>C5VMEV#@pYr{8hK_L?BesDaM|p? z@K5nIK04O*+0UZ4yCcp%ui6iDJoNiG`T}SC{Ws#s{M_hs>{7+wtX?%7dlmZ_XGA{OOV6>7^u*D? zP;Xq86vijf-(#S+Ja4${^?vklao|sWo2;T|t3!-_lyc4 z4|MBG;v@PNdH((2jX>Ywe{aL?Ht#Tc)f+FzzPkl_>^pv7=$H5_@9hDfO#^)sJ?Q=2 z9}YO{Hj^{m_+{;C-R3ROpW)~e{1uOSj?=p$&i?-okUu?t(W?%*6*;{fg;R&D)_aYu zK3BU1fAYhB(BG5zt&iblaP6D2`HSI+@%JS7Po}Q_8D6Kqs)ucT$?zPXT{_gypN5Y@ z{~+{M*9p%Ho{bs*9_%F?`&Rw+aieb~){BBah`vL-Z|}x;qc0(k#&^ZB;>7R@^PtP{ zLt95OocuYl4Y<`OhBvXh_Wwlrb)Wc1(1YaBc`qklNzc|Uvh!JX=a$GhKjN%QECF8< zaPiXfi;#B;AJr$?Piy18+AH{DM?DX1L7*q!E}q5w+Uh`~Px#?=!B51e$Rp@s4aZ*o z96e~?#iIcqg^!*qco6OCBmV8Aa~m@jAWzN7BOLqAss=fQw^+~7`%itm3-6>X={mx! zzz6$2voZbJCE%JbY>R%ZE;V^HK3*Q|F8)pAKRWbFxZ+p$cl-V#3!m!ZneMuq_!l_q z4_ndh*-jqes%Lg|>#IhuHui$v=i7o0dL9V&6R)XHyp=oOE{$Zz|@jlc5kZ0e}f9G?oue0Yxl z*MuTYJnDXA(U=~Js z@z8*iCmSQg%QXR)y|j)e2;4S1o8+yC%$@p}@--G;z1^N=dk3`SD9u|M{9{cmBj+?@< z%j7$tx9e-hUv=i^(A(*bzjA`E2OU6v{d}Rov5V%9y9RpZ1;cu;ldUg_zv$bHU)tBY zU!c#B^Yg^@;{lglw9jfV;L2}$uKoc5*Epql++_|QEsm>emxMSKt~_=%cqPyi7h^g4 zIX~di|C9Lt2I$D_E;L6+jNN@V`NxJUu z=U)X*zWogHFL3;Y6Bm28=#{Tt4u3Dd=%xQ}bB-Th zgiHS=bYtr|EwJvVy6|A~t=%`Y4riZO)ve1Jo|9LvLQWe84aa|0&!HYTx+kaM8S*@V z{w+TnUS@o`lYUuz8J>{+W8hX78BV?(*L}Df1pYa{e^0+sz%{S@7MS^o_{cAGy?BAE z>kP-=7U}t=`#V0OZ*d(;`z3Y>xbm>pg)F}rAM)_bd&vVfJ{m56({uT(t}whty{GH- zHXa$?q`!Kf(VEni+i;DWFOh#N9}3T)(dQogAvrk<9D8l<>NKOzldfL}J%hFBm!6wC zFY;mk?Y%)x@uyxL-yV9a_l!S&$o|#fC(zgNhrPjVeqr=E`ninwJu0-T7}kBeRyT@1 z!H4^#?kop;3Fn*7q)xEBXY|xR3uj|@8#fK_aLyX^Hrur;KFS+!re9V^8NKq+R?u5s zGMv1+MN~f7XGRCw%dGD@`%2Jezpd01bXZ^;IF3}z2b30 z`1|-3pBnTV1%D8(x@Z}GI6v^$*!dd#`5|cZ5&w274tY;mf1F*1v%axAe(2j3ea^ne znIX>zC$Em_xg0+475s@i)ssHY2v|2zqf_>6|Q(Z4*qL6F7jPZe#PSr_~D-X z_V_4zI0oG2^PzO?#qt9@DA4z{);^Dm4KJ<+g~8R z)pgb`ahUt-!iHnVvEA^i4FezYYV{)IKQZ8n$Kz?&=0W0<=}XTgzVDl>CWa*R?A=F^=edSRZ%7WkN$fN7azJ3r+JdV$Wzt1PaYta8Y^tS$Na>|aEBLBJ2n0`dB`-L9}epLq7`nm3d zJHgQxIR5GP%Zb169Q$4x_c7=FI?XPu*)cajZe<8_M2(nmye?YPg%YU{p=j*^WODeZ+FpS7p|xD?|V{z3qHh) z*0Fuw6E3?PM7x^={)yjO2c8*l>c%2Joau1o0KIo_2l(6=aP+@g?<<$z>LYnXuRONt zH1xy;^XRF5F8UI2dKUU@b8LDOUO}FD`2X^NOaG^%XCGgpS3Unu>}c!Bl1H{s-L8Fc z-*tQn9R2G$>70N||Fgk=7jW{`TD|{eVZbxSk#q1j@6Y02=DR!i>jJ&vVh_fPs{*ck z^#bzyJgG4Yo#bCZ9`Cop(f@$%tMzd$T>5{6b{9GR!aH1d{v7fz4!HFH4ch%tz~%p+ zL;v3kcn$nR(BI{7>6+E#>NU{o9y7CJfhX<(0^cIwIr3{A?3aNLdgl3}Q{$-ML;g7d z{5yeOey-;>tOCEcd+aixX zhRe^7fZpmx!iPE z#)rH#=PK|m0)2^gwQtJmK%5e%byI!*i~`>;e6J*RJ@q`EKg+mC?(Nx_)hW&v527%wPKBg5l(= zaXl}~+e`fG=;!tDxAh03&+&(E(=YG;q8|m{W~f)c^!zyo*Lu;rkpCDbzwpHK1J4Av zyeIy`vEyouBc8wT4&(b%^h^F`@gls$Z`En*1zdL2eJeK4GX5H0^!!_|AMwGC6MFBD z*Ryc#Yu%Fmp6d9wz*Ud^205*65&r_m-^S6K)kTJjm*4+yc$Ampe*J9p z8h7;Ec=@C0LG5aOx)XY^`JK^Y$MLr#&l+evea`v+gFoANVf3Tu!N2!Jc)K`v)yR?5 z`Nz3-g;PJz{2Bab1-wOFyEp zxGs68<0E?6>oWXsN}dGoP#1lA8XS9Z|8jr5(fC);|CZC>@>ktQd|7C>$vC3-w{~Uj zjf2o=J>d%ad#t0E&b2NyAN&syXB^jcfrSCj6J`QGDBx}U@Z_N9fLDSWb z@D~2I7ws+xcn5q8dS4HSUUlHZ*wK%V!m-!z5cHb{J~e*-nK;@z;7xw({SH3Qh>ym* ze-KALPAlN58?~R)kM9Le-q1L{m*X#-@orr8?O?#$tas^t`2zx;u@0&G+PvPxNAY!5 z=$G)s_iyO`guthc9_AuXls|{=Z2C3o=*5FNbNK}Fc>M?`zLr0SUAi<~d+E4E9O*e= za|53a{kjEwmw=b}eHDB*3Ap0x1oUt73&oxI*TC0=@k@9UJ?K82-5j49$BLr?@|(?T z#HYZCqc!w*MWDwnYp*T2XwTydml zGuV9GQd*S^0` zL-fMY|G4(=)dGEs@ksUaH39FCpB3*@>O*Lx|HHwj)M?-qz( z{5ZJxL*X^PdhYr-f6X7XE+K!`bxhmmApQl8Kb#ffOL)$>p#7*P1U~Jg{ek`cc-EXTas3dT&AscpLmZ&~Fv+ zj5aQzUA@Ck{asPCsdlrxLBFhiUWb!^c0_*d|J29uJY{t~KhXTj@HXS1-g|No^i%jV zUaO9>b#kMZy)tmiGlo-NjeixpZyorMf7n0Vx1-T(T=+ZRn;gB`*YjQfi9UA?xW)yI zQ#Q{PAJMm1|Izxh)lr6b5_?!R&j(vH)YkwYRcm-K5LeJ(WhS#W{H%9*LgFH=s z`*n4r*Zf7-87;pF&yYoaqU&GJIz9!CpQtX~B;eR9|7ovXoo#&DiTfkKPYv`L^RJ(v zpPd7)aa`-_ivym}!zg;Q^#tRuJf`#z7UvG4uZ z<%oc*9&WpH&tG5kb3ZIFg zsq(bp!@BFl9Qt*sYq!AZ*E95M$AITWKAGzOqHi<5{viFb`L)_D_%MErz<*ZYgP$cm z7uVZS^a;CYUB%}O;j)+fc0u5i)89YPU(crly>guPo%(i#C)VFq1-}wbKIA%4e|=K& z3-8ch?H4V%c7?~1&!?W zX7s$o(PzRL2ld|eO2Bjc?Og2W^ML5v^h@hS7YF(bd7|D#&-gX*8}vLSe}c@lkjpZ|p&T4+!+c*IGX>Q@f(aUgKK# zw)wK*ve#e1tAUT~wNLOj@zMD8?vR&+H<5i6dfV6WDRWHT!DJu5)yM`knz#{Jt{yn`o&kCXPmatB0i@;Tosrp#KYX z+_OuiI9_$~E8V)N=!IjK#z&D~^JRSu$1c@ZAkW8(bNe`Uscr!N$$*#G&v_61-7etR zrTPiv7hf+=4f#KaA6^;oCi2(mm#rJkLcm;SE``A76MaOJ&SIre-M7k|ZG>Z94K zUeK%>E+eOpy8=gk-R~dibM*Waet1{lgZ{aHaB7|`_&}%kih6$#UMBy%ntp8?_(=cn zg1`3%(d#=C+}HWd;#}qg`0b$MBRr)ne}KF(H{dP&=X7wN-$XC}*YkCU1HJN+-oLzC zz@`6R;6L*Mu5m>2`#hOLqt7V*zaP3V{TDdxYX7aRC#(=Zees9+$Wt8m;$6pW;`$Ki z&lZitGv+m#-}&)E^z!pJAkVddUV454S?U2-oxB5bdOk}F8rF-j?*d0JJYjd$1D6H7 zm9k9pHJ`^s-$DOx4sj=3^B2u)eiZmrz%?HEaZU8(nPUFs^*;)(`N*H(GvM@G;3-Rr zM{h^rIrzEg)5o>&cGB~%r`DeZAO80+`gvOC_Vh0te;&|vw8$r+Pwm%w%F%ag+GPWe zL%(;xmA8k8>#qe|`rHHg9|?E|`t{NOO93xo-^SL&J%$2|fSc#!bVU_?7m<-xK(2ec%SLro-h&(&x+Q z@52$tfAY{T(dYE{V&vQ;&`Y1%|L~K5XN>PBPQyoe@2|*dc~$a@f8y_Nz^zU&Ty@vk z@VC5exazL^!S{3RX4qW3^uCNQ1YG`k0rG4caQUb1yF5JL*lSqNuRAE0^Y>WpP5E}UEe$k`Ik8QM8C8?puA_}we(Qn$hjr*czp`b`D*R( zy1=K+yx>syNM7k`9UtVp*5!Z0$ph>|>(3Jn*ZP32=h`@Kcpbja(Jtpr;h#f)SG2F> zJdpQTp@%OyddZ&fD?J~~=Aq&*yv27B?QR|Dm2VHH-TMQsI-yE`Q@~Xx97?;24by}8 z*Z93Y^!o%{@pU5nFAI40`euLLBsnFcKDrO41)j2O3+#1-<0D*oTkqBSe8AgD z^}|&EEch_aYM;+>fu8s(?o;z|A$s*+_kC{{=oMd2(XP)Yg*?!G8h;pd^vI&VB<&~k ze1t2${Qg_%LAc`Ue}|^}n{ezkxfyym3_5*`UgP5<*v021(N~Cz)6m0)L7p1)??HZB zM>Rf8`gH=h&jTI!@Tni__7{o1z*Cm~iab8w3eTA@p98*yYgc%iJg@ZuThCm_hjxEY zyO#ud)gd2*-qwYTUUkUc(0jX!zt#c%K;Aew@X5h-pNP*7S(hehU%ihP;R*k_YZ_d3 zIe&e*jdW}aXpoeaa)YU_|pwWKy{<@s_2v4l1Y2Ss{v+x#q;9%N) zS8;BiPltMb0s8TEh3FM8?}Yw!fnIhw7ytQcz-#>e5%J~qEI!I_Ux(hF;~_aKd^7yV z?)#RJ49^q)NB(o6G5w3ajr<=D@)Uf?S2qNI5RU#E$3wqW;G^;7apXB8;L`sHdLD6j zDvos@S%1H``!U6X;va_dkLB%0DK*=sR5B-+=uOHg7dPs!LCYPZsDak9F2^8FOJA4h6nT5s7O`@Vs`nLGu~IH>!cygb6uKlcUn_bo_H;nZUz z+K=J$SiuK5j|=`IT=m!s$TLqodvPRO`acByUlZ^;_BsMR-{SCk_n7x{v8O!X8ppRp zKX(Tl{WtU+gNcCS=hZyODgKP(*lTLNAsqb|&-2_r@KOF*0{@#G-sD(y?`zQUvk@=+ zcB(&!K1Y7t_iFPU=~;L?W!dk7`~`pgKiEZhnXmSz*}U8MkY{)fK!4p&c&;yXQNr#& zEsnc!C!D-7p!W)ma;%Thqt8jT>+LA|4to9`{p|!ks*84`Und1z`Wz%b*!>j7zXqM| zfA{t9QZRk=+;yMF3Y`7qFCx#d<1ai%&->EfF9%$Cdl)``9548ghqnW_dAH;duDsDk z&ff<<@>{Jl-|KMcO8V4#vfbYyKEf0FIT}458|dYq+IQjki(dNdfd4wsmlz+9fWIH7 zL{At@OeJs*!N5D@%j{fj()UH^2tCiKfI0pZp{BIehdE4>3Kmn z1^P01O#OPTqi+>7%LY>LbK#2LhwyW|kK+z7;nI)lx4#wVcKuK|`E%?M=(i30OZb)U zFZTG2&})6;1pMS=$EU#QZyEg7fD@0yA3#n&4vM~woJZ5{x?^kk4ZJq=h)?I==n~Es~^Pa_)PS0cEHJdWA8wo?t}2TJm3|^muu*k%`YUsa z%NAV;9~+O>;mEIfocX!oIeA;pFWtfM&*0NeS-v55v^v4)<^S1+6Y`*oT>EeL4S2$@UdDXN>L~FMeGC2Txrh<(5WidF z51$BpsILarqhBWk9J@3$Uzi_o?80-5`|BIluGUq~q+gdh`W$)4GjpDxU8{RVU*ObB z7bE8ufnNXH-sKsiZ(~QlZ`|;VxOf`(t3>=3%K9316LH8=?8-cz;UV07wvs=Jx$o~uc=ahgq zk$*Y(X%5dgR(-M-+{R7GS>X8p`_S9bfj*~S)uI0U*XR{5i$cE&KD2uj`k9i);7PoU zzXN>Nz&}x6seao#;IfP69Wxv*-DutZD8?77GbNAkl(KYJ(5G<4h32zfpTfyk6aS=N zI^V`w;fjk7LGSG-yh6Y3!yi5u=xgMqlhCKtw6zS8xLkClb4oX2kyr^;Vt4<>oeZI!twuo*K_Oe#=nI8x(*WQ(SK2I9978B zwVU`TFZ~cbd;KJE?F;%XcA1Zk%)SMV-zpDZ==ccFiN{ajpZf${`R799nIG_s-`ck< z89aaTPvfh!dr`m>`rIS=*rWB6gP6Cfu}6feZ;E* zeNNo{0RBG`iL$EG(OGoChXL>}A!BmD@MK7Wfo&vf)T z$I8zO!Dl(YEpY63GVLB3=<(Y{dLQn70dJ%K_aKkXHF*j?*!>IWb28xMhcW-0hqnID zHzDU=$q#op`lZG3s)^SyUaarr6plWNI&&h>w~+rN;>hP=(JK!>i#-1b^yo9&8ve0e zgFPPsSHJE?POD3eUVNUT-QP2R4)o}sygjwQ;?n`wb?x&5J#m!35Bhxqj{Ya#It~9C z?=9LY;BD|l$m8WN z_~0j7hyDsLBhL}=w{gwnk^Xg^+UryFb>;;~{idusZ1kted1iy_BYA}9__@~qmj?PaOj zUErho_ z-YfVZPaQtxK(GA&0P?R0c)~7chjBrC>U?z_q~++lJeLj3aQp5gzwm@U^rURUDn(n`$38XB7dY)+ z1AkjTRJ+2dKe^u2AO8$T&l78q^8nHG>W6|4?+>~t=wG<%&zorX(Lk>}uj^2nU+81} z>-6hW=>Ln3J{8A$|5X3Fv*?B6=K~ADy?+XCk^l8xjG`|+IfZxdlb0jEANPgJ&p(Nr z;%og9&iVOUiC-Ug!fSjVMow=pVVc+L`LR~tieG`#?t74P9>0yh@H}bXTEE^l`ZoH% zaT@vzIji8d4rcVqC!Zm{tiCl|`CR)IFLmwKi{oVjdY^*LPem_0A*Y@jvwNU#f$ONAf%!_>_xv%c=QC{FMjv{!QC|EwOio$kM;}Q z6X>NkU4OCpl+jCXcO$>erwo_g^xjEZS24T{K8~F0J9~{5ZK+@V>vy6Tj(&1IN661F zgtwSay^}b-vN*S!2ZVRf+gIrCW`TbRyJ+6Do5NKDNDuEu{w*Vp9`xS2GaUb-WyfN|0|If6>byE&<=z z;WvU46%9QHX8VXEzx4JfG!`$S$G$vYYpTBqmwol#IANt?(fA!b+ zq`~g{2yFZ^oO)?kb@iT(Ugz;G_A7dMgeT^BfN#2yMbH%Z@BFB8SMVOz#spd z+y?zG4S0p$H)G%H1Fm{Wx-WdP*DtlJdUadm-_+4($S6O5Bl6rAapcr^Vf9a;2k05c zCtLVQDbP#*dR~EVx8M)mPm%MRfnN2}ec*Oo!}x1_X`%<)&tN$EFP?LIl;a~EYrW!F z0NuLr!uc&+>9ngM6rnbdRm7YAJW-x>ZF z1zhLfGSsiHBu`G8=zln&x0@Y(f%CsL*vqa{8$JJP>-h%WE~00ApVV_-Pm6qzN6#g= zFyI>BKOg!lJ{qU=KIkca1up;Jk#=V~J}q$Rzd?U({X}vKPh5|F1bKD{^!lIfgLpXL z9rCK~53=#y_-OpPeTe=A`U?J|_XvD7;L^X|!+yBK+q5Y^p9}s%#A#RlWA&%n6@8AM z_~)k?-Xi|aIBA@6!AgfypC(n!@8QK*ejf1)`&-V2hz4ZK(q5kz^!=-1{)wdv@ z@{91SkXhqGf81Y(6OVs`{(?X+JwJ(EG#+?9_;ceD`ZX)ys#kv>?5=jp}c~Z;qocXb9Ei(IF3!u0>_rxFZxWtbNt~A z;8vd)pLWXf3jUl2`i%VcrfK*qZ%m*!pFh=ZLZ4dqTjBUrisMx?KaIax{i$|^qtAhz zXty2cTgY=P`Z+P+=yPCSaGMtxALWNtv>W+Xz_pIFAn>Un{}TLUwZjv~eDV*I2fW=2 z9REBGyQn?WvvBG2D)`uZLH`%tChzI_gg$Qm2L)%OL%-`yWML;7Ro=w1HmnUPJzs0dEs8?*xA^;2H9~4Eb*lxYhwik*DME zW^t_b%>I5+$s;_`-`miSmtS~`aZ`R0ag9f+NqszuUh#W)kW;w)Z8m)F;Mnw|zN_BW za}vt%nt~(GCg7I@`W(IKeqJw+_((s;flmbbjNh7P{WaiaerLq*uN}U$dz>FiJnkRx z|JT~NN6S^!cYGG6;#z`DAAdBg+D@wtEluLRH#d)qwK;hKkvwi5#3Z#e5D6eu!$WG7 zRgM%9TULz=!6fKfBOpaV2`f|sQXRAkEujtaXhG2tZKMSOfi6&i_Vd{@znPs)utG9x z&An&7^PTVh?ce_G$Ju9}eXNha0iNeqev5R+3OmL1kNc{9z0lD=&Raf|;hQ$auIJ)F5B?)gaYqWjmFRLbId2G-{S4-MC^&{9eo^szK{Ig)#C(izJ3LKex|Pu z%pWfHxmTxTxZitya{~LHzoW$QkK=)F z7kG`GUjVKLn1S+aD5pzknak$naWw=1hO;dg{ya08qOc}4A!-{S4_ z(u+s~j3j6c*;quk?`c8xR%-ad~w#$(zkZaZ)sakYcv zXs^V_v40pl^L9xM{do=PzL9tE6VCade)v$PBVN-EFUQa83%pIcF{c-As%lTtA>BVM z%2(X@{0er)`AhIC@FNO4n<=er_up33Bk>CVydOWr{D3;*E!t1-XIzth+qFj=e@=P> zcAi_{+T`{0vkQEPcK+8%_reSx;@E!b4t%yQ^IO8v&-bO#9&!El9P!^SblTvqtNh~v zZ{RoY552d*&BN2+d3&h7=LK)YZ!sP=-75Wd+t%}_zi~Y!@QU(s|He4)2;6@E8SH#E z{kH2jbvoGT{r-9Ui_;%%bzk(npAxrTuK~X`(;4TuL3w=ypLGhnt{RQQu=5`Z+_-uR z>0Vdh4dUby?99gl>Nk1cUBuh^^jq82z59J}j!M#EGA+Un}ev zw|#pH<>+`d>d_{SdDZE@etxqTjz3pn|1AZ-CSBXxyu8$DgS$@kk%Hgg^B2kQ3k7a_ zbsGG51YT2*UWHC>zkKV9>(Fj4_3?iPyqWoVbK-2=coSI6^Q$jj zk?#5M_hdTaE%bekVXWf{eg{9?3;&2re+ztH2jS17OT22-KA$D#bJUSvqw^sCjO!zT zn=h~9zbVsc$)DqVviB-`+^vt|6?~t2?j9&PZoB8Y;+Z|X1GoF??oH)q`F$KeyN)-? zCGeX5{WY|oy9+*c9`ZL?I~jcE-<}}NxQ-Th6MT&FeWK8D9P%XmXER*89M|29Ju%Ky zN4&ys&R54geBdqYKLkHVdl-19;yzJG@EwQT zj{RFQyfR(xYa1qAe{VkhFHZckzu(|IEASR^br(KPA6L$NtoHS<~hX>`XXzx=(ekFOENVdH#4orZ3*c&puBx`qREX z<(^+>58Qc6ucJl37P#^0`Q}xbe#>;{Y@0Nnc#Hl{zIa8tC!;f3@cr&M_=^hM_GH$) z!Tk=Q<9WvlbQTwU&pU3!o}C5mIN5claot>#{M=6UEBLPKj_U&QJEqHhWN)CnVm(;kw$EOd zjD9fi7Jl~mgJYSU>ew%zPCQI3aQ%5T_WVKm5=a|L*AkNW!UbL8t=8DHCdwc$tqs67d< z8jbtl$9!Ypb=7EZ!=Hat=(LH?3&9r@c!Rk52IX~sf!n|MJjaJJd?+1HpYC&kUyDsq zj^f7EIq3Xz!KWYD>NqFnA+=L}2mZUzi1V_*^^ebU&&yZ7Z`X(9?L!4-`{w@P^x4>Jp!+!w2rNBG1!`sPM9zW{npLO7Qd$JzfcIg4i>y*q73C9m!r;hQbc8b^N z-wQX!hk>`j$H{M;Uj(inoL`!d>2Fep>!dG`ZmhSFFHU`7Kd!<2L*T}P*PVT~UzCe{ z^XokJH)>bi(hb~l{0jQ9FGt|yYr6OCIv?mOaK|&AOYJOlhWPCHSCp50^G7^S{yh0@ zW&DI!jd~V#{x!I-;OmD=vFDzG->w>sd&%$a0%tyI&zab{y1<*vOEigxnEwoY`wQ2* z#{5sp7uudjxSk=#U4i4b^!)B<&jZ(Q-d}lTmTm)`4*0FupW7*JKRXxwr3K&q!u7xr zZ|Ya@Z4dvlDBpyme+Bw+9ufSS`wra4U}mN-zYTvH{$F3<4L&=*{aArF`TQ95CFX%c zzXHFDbhl^x8hzKX*lt99moHA*wQ(5pV}ZB$F5l-oaM}&_0qC|l4PJp8Z+Ffc)W^Ub zKf5k5#x3G%Q}1@T-{T0sG0yrp<>-C!*Jt+2uX($VP>j=p-zE;P1&?*(f$KNxOS{k^ z-WU(;?FaSsoA-HLlJVQ7OM80-?MYmJRA0Oz-F!S5xLQ6}CAUX@hqvoUtw)hxasB2# zV|g5k8*hIP9{siYYS>==$-Ke5aUUmsTqobk^ySz1^AhZj@ka28xBET+J*40}J~E9_ZFIUZZ!EpVUr@q+>`0To?4^bli<|aoejow0kjs6@2}Zug?!$|GWd8e=T%uuPon~UzI<^TR*tY z;rB8>B%J!|Jmk5Vj<|ks-FV(#h!ZDV_Z*Bf`}*W}8vehw(DA&<{q&yB@G2d<4sh@s zD|N)JN1wvR|0wt^`rirI6X$Ee?=VmDyWo4Zhy0c}VgIo~|1N)=H|_Hz$9u*^#6!aA zH{FjY#sk5x$*=ni=lb$(XWYkoJI8zJkgogs#(9uB>l4T8?|1I{{VZK^%Y7^UjQPgk zJ8szkH?QyVJEVI$`0k|Hjh{^%^QycbEyf}0BpmzCK>vFMzedOU8|zqt-^QL_CJryn z_;otoxX0@wodU10Gd%~qz_tG(?0h!EFXFgCy!{q-#`Qtd71tiu7sY&?_-1kTm2#bJ zobLx-QSPrnC+-t@0&e%k{qI@1^l|*~=ji`h&;NnfmCw0jubcv>e3|bZT(1ni^H;6| z`LASKcfE$`+8#cJ&Z8M_TsS^C2z#O&)e+Zk&ro0D{4wwrcKRHd3kn^|J?%Gje1T8o z^VQfM^(A!dFTCy?{kV98ygELyyoU1pic{|154$DPmtTYXeDK=~9RIK`Xb|UpefrZU z@$(G@-+uO;q}k5!N`2bHYrrE;)DfqAw_b+*F&+-Q#dnkN|Az}5{KNfLgZ3);_7~nC z6!YqV>mRSPMEfsZqis963jf@a=_j1;?BAmu3w}+xc)!5g3%>Q}-$^&Fn*`tZSx&m8 zev{vOfAKdm9rNmXnD3JA1m4DtI^q@Y=kRkLhvF^LwO;4(Ag+HrkBfQc(3yzdXA3{H z;XAJQCU_nv2}kFP^9JYd>WJG;z6btpKKly14PKM4=%)g2z;oZ5_wzsUrH*mpy42<`_r!rg{^}{Uu_OlYF z{$5MI;yR1@l3$bdiG@Al`r-5B>&K-I@!-DkdAZAP@^=6En8#CpV*2f!yB-IR`JBM@ z!)Hl1&#!#^u<;S}W84z_4)&i)zT*5SaK{ht2EQ$()h$PDvEO_r{aPN+38x)?74h%y zqa4NSbbdWppX2x9_yPW4zBF*hnTyC*TvrkwQiqtonS4dR9C$^(>=*KORelTq-$;HR z>-{$P=GT11_)xs!?fnSLu_x{i>f`7;pBDXQ;I@bPejR~Z-`z*?;Vj*xPrfd~o(D2~ zJRQ3?ET;BhRx8Q#UoxR#ga9(Fa^DG0I|Db)U3V1Yhrm1J>k@RDy>tVg$nUQJi+LDv^-TKp=tsQuadeKS zeTeIEf!CyKyBhse;O$EN!S&X_8>DkS;4Q|3f%m+AA*B%gfknM^s}mGOr-Hed6x=XJaz z-Grm_DEW%>N7EIrE6R#=W8OY+%WDYS^K@V0m3-R1ZAE^?tIBmad;3AWMZQi%CvTU= z!R>dPpP9>V!p?-l{{;EX{VA^hFGuHpJ?^#V>(lY{>AQ*ZJU+!M_@@@}unF9Fc70Ct zcMa0-zT8)Quz&V+&4Am|NA&P?o05=EM58JYx+Xk zLw%)f!8c!jhvx1AZ}9mi*csz+`IS2S)_906yo}CGY-+^D^i1|;`7jJ-^V^jBjVUc?J-&WEgX7%KLy6=0>|lH)U)*@OkK>8FL%&Mt8b5>Q)5%Xb z{3RUk>ZNPF@nw0r#x3em;8o?i%0W92xbgN{{1fAf_?_jljB*Tn`grB~jX}Q?cuhQ9 z2EUWpV?NsW=gO>HE4mqBVO_Ld@#3P9Dnkh&B6Wn>a@VI zbI(-#|9U=$oeAeV*HJ`29(>|w(!s=AjKjs-aJ9$#ablh`aO>k=!O8tEzs1}8I-~uO z-vP7TSWSN8yrz%yo%k$?dd$5m!4dB406Y zAV1;w?QrmzH@5zDU)_7WA`bg_)g9YLI~?`Z_;(+en8yg*`tCf;4atV?`iKS^j?cYc z=53^Bxg>pY{ONNX!vEqGxc+a!kMqsIt-s$Vzd2s3LwqhkC+tZ7#nJyE{5W40Z-HyW zweHB(!_$VF`~OKO{&sLQbH=ohk=Y3wnl{ACZ})MOr?Ya|idBA=TfTGNFM7Ok?UF@H zk7_oD$=~eeB7ZhEd#@iKzNk4m+#G3+9Nrv1dURxDuPC_Dw$!_N|%HGi~aA zXE^+aaZ?XOk>vhci{$=hIP+hxP|WfxE{>na_9(rtG2dZlr|%Gbg;KsBW6zbR7> zX#I%laiGdKqWJ)na4-I*^s3x}sUI{i2c~v3n=2NsT-KxwADK0KWajMIvuDnl(Q7n@ zrVWqGm^CsyG;_xA%vmpPN%oDbp1XKSbLrfr`~LRjqjYb_Em_(;=I8%1zV3ELc1`v* zyrz$)j;fCg_4JPIe*IE^Us7*m`o8u0eek{&Us`Wu-+D)l?^~#c(gFG!9-B5Y=4)hn zDxi^}k&$VmWOih9bU5`pk>8bzj{bR)%e#8b(P_g;Z}!kw7nwcQMP`nTrd~Gmjym(C zlUDrM#?wzYx%bPx|D}Ez=tbqx94_*^bitxtX2{OoBCqJv_=;t3T{u4tIZVHK=1Ch) z-E`)=H=J<#S({GSaMD?)oVw{{>CRuVcIB$MQM~)Fw_xt7xhoeQIQtgFfGy=JMg3sF zHaPl)hi(hrHh1;H$Zi_3MIpW{xAn;jmoDfZABY2n;z%w#>X?Od7dSNE-`D>EvAcJ1 literal 0 HcmV?d00001 diff --git a/tests/files/mtz/6JZA.mtz b/tests/files/mtz/6JZA.mtz new file mode 100644 index 0000000000000000000000000000000000000000..abefe62bf54944847b4c0ecdc3585505925ac1c2 GIT binary patch literal 90520 zcmb@vdz{r(_5c4CyrSkPc)<&dl8Tpbbh1bz-ev`UCp=w-9o#C zboF_^8nkH0+<+I;ZmX9M8oKV2fS1w_hPDnG-Mlp58pr?gTvY+D$Sr!~kwL5HZ4P*4 zZvN)4f;R-bD%a!fpA4$5o*wWtH^1S-LF*>;2+viUo7m%}nh*LO9O&zEM;+V(|Mq}4 z;og83b4fY|dAKj&D{`Zo zf6n-f4*05E~S^oerv!>>9|XOK4?ISdSW?+T6NNx*~_e1D`tNp#pv4=TYG?J5CY%}YRk?2a_hLPo&dfSWKiA~Ct-ieGfALao+4$BV&w9oU(BD1Pt~Lb_?hnxIeApJF=u zr7MxEIe{MiyrLNNQr7=m@5Vor=}$#&^=qrpGoJ=NmAM{^ng{i6yfffc%T=Z=oyl{}3V1_q^uTYjt|#)hnoE7eC!b65Pa>bY>UW1X zA-Auwjb$LBDf4+lQQbbi$palhklOX=w5t7@KExksQ+a+|I@HLQ2_ zzasa+hMz_I0lbp^cM1Go4t&thS5HMg7X%#paMqoy(-i@)&2@X=-|QFrhv%xxO?>e% z=r0KL4Y|f+CYIr!NB)0_-kTTbn`rmMK|K~-9q^XirsvlO{i)||MPIemY`9r(26K4|+P(TFJ)^8mD7Vv}0 zjAv_Z-AU(n^7r#7{rp<6E5yGIzhfw}d$Ib^*(c%c$Zgj${INIodYFBDQW&=)`urJe z+={?w1?^4>`-u3jLa+X&44*Z~=l5&2_Qm(71fN;K4llwdNt$cc)@*Ng376asMSou5 z`kB$EJ^C%@x!w+VnvOi{A7OvY_@F;`W**Z(Uq~mOxqGlPMPE$oCwE{U{w2_t(%C)# zSu=FqNdZrC-9~+n^))r%_}%$qf*mgY=;5(PV7K`FS9n#f_v_1$+x3A@nyc=^~&Pt?-82#4`t>Z~ZtYtZ)gs_>><{erW?J}A5e`Udn^ zoEP-)sFUeuQ|NyI`jgq;y`3gLZTMy5kmu_HeLLgR9lOihA)@bKJ-!?40pUet|Du}h z{TJcb6-)lZe0e)x_$u0667;zJL*?gg?(9g1#b(JkR`24f;?x_QOld z8PAQO-4^!Q*K0a@of2^L?SyZJeO!DB%=!NG=aE3)hCKYO41GKH?LEQ16(8hwz^0n! zrQ-viBL2vPu#QB(0zLdW=zSbR_$u^jlyl*Vr#-_OIDp@6z7&hn=dbKL(o3Ix!1DvU z?LUQm`3~2gj6Pj-lM8w&>ks4B1%J)Qd4yMBU*&_{D7=#R&n0E>s@yXxCzR=D%9#Bq#y#L8f9hyHY+x33OzD0=y8Ut_=aexL9axsEvxqHjk9KC6hC zogD02#h#O7WXDcC*N$%9h35y3`}7;=pUI9-MvtFW$2hkH9J}<~?jhcl@ks}qx=XNg zgcs8K_0z*R2rs6~#-D^f@%E4KQrc~1cjoK7&~B2O-ScSlaAUwLh+Ec%eP4VkvAb>y z>s@#i_Rpwt&xOtPa?SRBljv))59bBDNO&Fl*BJ2D&>!^V^ZzVEpU({)(=F^T;?tB{ zwq?&Cx58W42fH!vYXg7cx-$lp8Mi{NzW!l!qL0IikM{T7kvH#W32%qbA8WR^$AovF z_a>uPqdXK@?+e3zE&3JsQ&)%Qtprb!ZmXxWexG*z7cSlS#h}kK`n1Q}V~8iu2=r+> z`{P!g_pyNI)9Uws9O5?OQ^3xAuS|c6=~1Kqh5Y#O5k0!Q?l<`3O9KCd@zfmn{DtT% zki$}t6XBKYkB0{N6i%F~=YN9VDm+CF9|`M1cx`Unr{87YzajJod-CFaLcCP;4e*I} zyYM{oJ{Y^}p}+?}Xi+uuUJQ5(`e!fJowrNHr(WgV# zof!0+aO7&%-NA1Xo==m)Wc1jgz^9Ooyf`2BYta|e-i_A>eJ;F|E_&m3_SxqHpM<#3 z{p^FKfLCDuoX&IkeL;LG+1Ki@XZ-#pyo!1J8}q33q#u2R%fJ2;^nM!M3 zI@bGC=;sAK4baal^IYh$bzfrrz7pu0*f*z>$MNwZ@o$08^<{9`dp+rY9ESq^-yc2m zY~bI9eooPq-oF-q@>30cLL5eT2mG}+ADnI*pA?*U<1Q`8!_EOG-q`jl_VG5?ZbqN3 zzOumnvNGV0bHV==AL8eGM?EIINE|*F^q%li`ph1u1wD}Uhw-mv zU;9&N7kjK_4hK)EZ%!~3tx8N`MeP8tCgH~S0{&H;?2kE7cnD>7NT=}2` zIef$6weU%j`J11>E_yNG`GL(f&lAs|67VFQ_~L==HN z-rp3i_;Wpa((k{*Yti$+W^S$t?bgwsv0>j4eFJhknR)bnwD3GUhM{L(2;+=C>3cqH zUl;hdAfMmH?tLK8x1t*cvu9lE=sf603HagU4-^E-oe?fq`yil3j2ZYqTS6p6Rr9oCEJ70)Ts&w_hi;lv?VcV~Qj zd`39>dC@NDrpp5VO2*+V?4_3ipDOmXFR{OT6z~*Y!+GAn1RQ^PLQ7a*Y8Sh)<3!|q zgQL%@pV;RMkdqGs&bsK^N`D>^K$qd?(=K-A?Hq z=23i#T6L;Cc6arP{#2o7s#ss2 zI=l(`B&o0O0WKaK>%x?XCMBtOrlTYc1em*GOo<2?IHw+4THRFSRc$2k}>!ep3 zTidCxpTT+=-T8mR`8QpFzq4DUPse?FW@jA^UCemsZ_>26%g^2Z{uv+Iy|N5mAYXr_ zi=!HS5kKoV=Z_m+N{1fW8$Pk0lK;I~7fI-U1$n@6E}vq2D#_Q6bo;5{Rpeg=JA1}( z?Ax#2O+4J>%J5qBUeVcIhLb<~-fWj|H5@;E+)M1&mS-`Xc;==5#BM({jB``&nPaCB zr?P#==!s{Z-OI&A4aaZ#^a1R~O<^1g==dLk-xcsS?BsWx-(-B+SrgM(rxxciyaQcz zwzF3aFA`_C1N~DU_^)7I{^k4~qhE#mKk58s!->O>nd1C?!%K|w<<9>#d;@yyb{8KL zo-l^kSM9GMCtu}vqZLlQ%7c9dj~j4Oz?1Z-o9bOW$M~@S_PjPcZ^kDbcgQ5hf3AC8 z;RSwfc5wsaLwzzz6$k#T58kJjeJnFuzB;_>$rHYllt^a;x!a z$~E43X;_cKTXN4#?d9T`#-|lK;S(3>ZR|&HV^>&y%IMqCsiVhf1%gDnj;%`qlf5GUL4^3GY3qn6j#D_oU_9LU; zfPXO;dD$zS)0(i{Bbv zi5~9l@@0ls5%1c8cB=#blsLdH)cqbvA4F^PwYky5Peh(C4fN!jPCkgb9lsB%|M-uW z?+YKVAF^>K9yOD8&0jG7P1M=$$MbrNKDMz8-hu7gE9gV@vxq%&QM6mYS1{+_5B7lQR}ly3 z680D2Ygk`z%1-C^tncarXLFO6}-vivt)gzy7Kq4_SX0|K->Z z=8qeFMXuwr&DcwOWY4m7EJUxm?q$s5>!CkY>|c+Ebs_qcI+lCjWBDWFUyI&&mHqLQ zz=wR+@MVm{YXNT{UVU`1Z^a)w`P0*b9VMK4mHNd&o`qv4_w0)urGD6taK(XV1bb3A zdU&VFWpJHS@OegiF5?c9aMPH(yGi!JbM_&X_l8HTP@CPlwqkam{4_q|l>tTP)=!rLu3igTcH0|j1 zDD74S{)#tVAM63qD=#)J*qOqM>FRkevA#|Ud`i@J{nW*ut)EG5-H#6(ylhMVtj}BL zQhX|Mi}v4zys-JdMlZR|qnErLC3@vypAG9#c#2(jhKn;BAMD{}e@5QAhIWzLz1AS- z7VkCs2I@kmbk<`FGo14^>kkX-UF|mIy3ITgKHi@d-a?+eREAG0^IpTcyAIy^82<{E#$yQ>?wb`^h%^Dg?7dHE%A;PV3iTI{#u(36t_UdO!n`fB6fkn8>Up24mb zf7Pp02R$a7^ItuN1pi%l3;CBmVLuhF^Iu;;4|shgyuki@5BkU3x5C?aE}u8Ceqs;I zyFchR(RUCZD1^9yaMcMc!x!zt@8V&6RuG$iI>@2uW#29d`vo=~pJiJfKwo)%BAjzu zt&boNexJ(d)1waV734%Xe%1vu8RzZ%P4VHJSM*0jPhH>cj|V$VIQ6ItZbCm>9A3{Q zyu^8d6T-eGTygtdfB3C_q68)U#_q!*=d4%V4>%L?4)i#g9n~=BHg8eDH1-tjO;ExDz<(${y z!Jidgz^;BT#5skx@m!OFTnTSSZm%bv@cneFOxgH%FwRdRZ`;LNpf6%i9*iyW%Rs*Z zd!r+)U-4HQ>;v|Ai?@q^1vsDdM)X)W$49u%Tb+eG&vf&c(G&N6I*h0AG~IOFf3e@J z{>1pRU%$9x*mp#aK74R?h+hduA6{KVuGWWkSxfo29<%=7*Z(WnEy61}$M?+;R}o%` z{l60Z6UPDn=)(TUmG_rLpQ0a54StaDTITDp;CBeGLr;mQ|x1iMjqJL~1wL0<`{p1=7- zm#4G&E#klQM1Oj}pc48d`Jiu?5SJ5$2K3b*f?SDD9{G8wM&rHhN4U<3 zOb&WXcnf;?3f7VL_l37&$KD#`Q@H97o(y_Mxay;J!*=|E(fP<0pFjh}Y`i-x}!gBQ89)(@wUr&iccC^rbM5!Yj}lF9v&9IC^Q`P{!HYlftW* zcbzBldO*1J(w{;cSUCRgOFIPl6po#?=xKh2xOoxY!1#Y5?3=>bA1AzuP2=qo;rJ1& zA1O2bs$0lmxBEDo=vB}8*Pwrd7tk9ovwvBgl8v+UlK5W`#u>e|s5`i~pHuiG$?TW* z3h_PR`GL!}3<>KlqvyQC9_WY3j=%7fedm~E;(7woP9a8+d^Lo*dZ4L zycOEhUEP8`m(KM(K&{5nj=q-u$bLQx{MCSSE@r}C!Z>8~)J5#UcpebwQ~tkKkO$G{ ziG$td>N?aP;p8bA_Y3PwIQDb>AA`LlyhI)N@}O^pQ|~=%AwO1MWclBQW!#PdS3SHw#z%E8&+^@XbSq}K;<~-r_r3opJ`L!B_n4a_ z13htu`hDn=ms`;{L2vPV-!A*h>c0g(bHIH;%vX{J`qPHzJ?&+;e60tfMc4 z^(CBon8kO8`O5ggrG$-Eol8=^14Ct>}1ITW6d7ajti zQs6^9!1&ogo<(2DdighN!0LZA{=%!!0~hn%j4%!<<90JYR`+G}8tY@xt6M_5b?h$> z!E0#XBRk<5_SyD8pGU`^$Io;}-=OD8K4|+FVovV;$H+`A)ONVkUzb+D}H?;@JWa- z{S`j8uO*#-m+;8rU0lL&^1?l?!!GLRX{F<;#ReI7gNcxOi$ zJ^SoO3!EKjxaz6)#xFQ2@R2=#mh;PuUgv9O)bJF6z9P5oFYV}alTV|ko~mC7f5hb6 z@T%PEL4`8?Nx6Su5A4YSfqyOb^AZ=wGd^|PhcL+HRSjpK-G3Knw-}z!Ri8JUduC?Y ze02Ux-0a0;oF8HI=-V&9OPt5_mEpuI&T7Oil^vCRI{6fc8xM8%u+gh7eFZsP%S#$g zKBD1AuAag04%WqW&ObH0NIY<_&Nz~9m-97?pK$p#qhE!;{J!%S3}1uYxI>xwE#aR& zk9~4&82=6IIp1*hr1)e%Nk^|EU7pBr@@r4^!yX>v_-y@dbn}mqH?!Z2o;<~z9}iwP z;j{b4I{8!QP;;&GXN|s)F4|>~ix(MA++%slyjwie@KTyoZg-!O^`Cn!zU%x8qo;0X z?P17+gLq3<=6yowXA$}@yM4#_tbqQ%F1}~@D#riu&b)>n zpEbzsNzN}Y`V#zq^Ev8k>E6UDC-XCC>IG?nxWY|NZzVudy%d)asgy zo_O*@iL3K5yn*=hiJf}P^T}hQ?dsy3Mo+%@iN!81VR#EV^?VnXFkE@J$AVr;>2ra4 zw&7u45Z*@oa~<}{urLno$p51?|2xh)H*}x#$F1EWe)Ql#FxB2RJ*Id zFGH@(&oz3T6Fz;fIv9)3?oR{D^F!Uvm#;-nZg$TlT=AdVv7`16 zIQ82THeomVbt-!DBW>H=8zUasPmg|!JLS#hQ8@O-#XjGcbkb9YcITb!rxt%RTzTZ{ z*r%)>#_)uBna{W#8v0p*9A00Bo;u-PSC^sZ9Mf~YZqxXu)bl)o9F_uq?2T@ZmBH(X z3$13od=luTpU-!3a^tV~VPo8v84K*&EEwp=5*pKqix3XVf682x=1=ig_ z^yie&pEmMkXVcHw0ayI+yfQfT5B>gtoow}H)}JEJ^&<0Y`EJ7{&)ty!KL!4)h@Z^v z%p)ZF&`0%i4SWu9aSy{)x1;lee!nh2pCsr)Z2Bwt-Q-HR{HANMi~6~CGx~J)C9T1J z7Eau|&r%o1PQWrg+*4QU;>?EQH_h1#`8+l7$4-8GyZZq;`4Eq4J~Pah_+$Uyd9I6t z8UG6InYoj3J2CJ_ZeN*>-EQ@30Hm2t=RLAyZ#HuuYZ3AcH?~kPl#h2h1^0`H^85FofBbq8eVSW#-(jq~`p_=&dC4_A?;QcJ!~Yr{){FQMH~VC5 z8UB*bQpT^LC z*}Xc~6URq&WnV>pW(E2pdizNB@r40b-njyOxH#a-JD=1UPfPTnFY1r{xGm^bA2%q0 z6L;$P0e06@{BG+=cz$5-oh}IWPexCC{#@*5Kc2!<&SB+|hq9L@ng(taX``#YpO*r}J_J1I^pWUmO@sNCu>^~^@v!bt} zE`5-TODCQ5DfF)dKSK1i_~TQML#yjEdgA%b-*NF$!yCwhe*?Yh<2>S%N1wdF+*=&U z=vDXlC3MOqqG3!8Z(-h7VzYjBu9ok`raS}OMlK}ekTP!t9Y*CSQocBJjvpCNxnVoBf>d9*mNuN zV)3^G-Wfgl!LhEM&~VwGe+c?K<3pTzZ{*Y4)xyz}FGRT#&bgjdw*);YyhMK7*UwnH z#6Q=!GX4*T=c-_Q-V6JM_>fj!QxVQCF^|!i%=1}Cu))BT9pP^$W zG5#AJf8qIo*id1eX7t2MXR?p@bs;>(uA4>wAB%j@Z=ZzqBKku5!G^`i$=X1#dWr5K zZXL{b9cDK+miSe5m(+<$_?hi=KM0=Z*^VE}VMr8C$|S6`tnSJ@oyc z&xND+HlG{fmcr}Ue^=CO9b-Sv4eUqfvESVp`kBYBxGVUl;?snlyfw(JaP7V%4 zxzzCyPF(1XwXBP-0Z+K+p;y>r^!VppsW^C&*td)fHkurCNlu1>lz z$dB+^}$gF+)FSu*#8+l z`Il)S{wX}=Jk-7H7gxG*$oSB2e@;{V6kfp3?G^TY;pm@PT3<^7AL>yX3apo#0-hk_ z?*{!WKB}Ww8OBpM``z<@V4XGvK2^xi?mU;B@3Vdq51&&Jkbe;$;;W1P5%h`hT6AiHKj`f@;dRW*I(;A7ZGeyNQ?)t=Yd4QQsrM21^&&ph zpK}i|=WZOII&k*!DSyTev2!2dBV6~H=$wmR?-@Pu@L%C~+c{RF=N_lk4M_kWOYiO7CIN|GI zy^Fq*KAsK#m7(1#^lD$`%g3)opThf>opn$Q7Qoi=sfGV&zK`vae;+aiUOu1R03G*; z_Wv#VaPG5Ry({b5&i{&kM$i2)-(ntL3G~Eo=Us-q=k4%}5Bu6_^yw!4-qySD z0^|QE{htx|U@tX%2mjRSYOLK7d=6o6XbOB1WL$MtLjqnw8z=KzT?1Z;9ebBziw@62 zt9;@BWcs;)%O4qmj`RJ`=(*=j_bD9~=u_@v`8T@B@{*~2Pk!f9`D>m3H=Hpk?m4)+ zx=+9p;*h&yPu6sfkJ01rJa>_^8x7B=^-tV_e_{H{@B;I7EBjZsz^9mwduBN2`^JWL z@voQpbA`qy$#p#PGpRx?SeQ@BDl4lv?;;Tkq zn_INYMXp}m@H*_D71-zR2R`I0*Z`{BUIFI4f_nuu{X1-M4JBsyX@p+?HZcz8%T3pL;`8(&h^)O)#GJF;L_;;OMYPjN%J&{}UQw%ThynjW`&3-d{1A1Hc5?FnO@I*g}-Y2px zUTt`O;4^nM;lEqF)bNCRFnW=980F?me281#wuJe5GvIl?^Koh8Q{a4e4|k7*;nX2N z{~Z^XHoQdq^E)ozXL!PR-i<%iGqfxHtoa%r@Jhx;5{8radBLAAGdxc}e;MQ?Va#+FRaa-{8+{A+2JBX5{E1s$w{wtN z@j;&dd2MGtDB)v#+K`_U*k^wi#-W}4Xjiu{7(Mr3uYS8se~QTeudowl20pr{y~zI9 z5OB`D_WHWBtBt?vN$x5`U!vUs&W^tnkw-`IGKZR>e?cw5V!nX42 zcx7*w*E3xAM!x6rhK8e`C;R*V4Oe~CgUG{Q-1F)?_QCdv#Hnryc!~Pb+3x(K@lpPL zDE^DZdkvR7U&+2u2z-be+~o78Mqfpp$-p50*>{ZdRBTP#_l=(V>LoMLWBmhv#iQOR zgEx?u-?z+oDjs!n8G6}=qsoj!3wfW_%zy0V(&JQ&_VVw?l@pHo~xX;JGx5aSm_{-__i=9Fv2>n!kb2{>KVd#(UA9%#& z^^8xEeeD(esj-261$N!-F79FUtKj3`mtpuC5+bR==+Y zdsR5~sNLSdA0OiSlhI=jpBC&x;b}VisATZAcwh>a@!`DEEcB}73k+A?$mf_BFAt(G zavxV3)~|5+ES#Eahk zS=b+;C66#}$*0(3Uao{!@T^C;_>b|aq^`JokVDZ^kGGevBQp9Fy*H#x|FLg7o_em>Z-!gU|fd&r8# zbBs>`efv|FS2DZ}ow_UI^NEum@oC4OdNAm1;T`N13te8=_!QYk+R=wrS7W&9Mtac3 zkNI0akIKt`zh-;+EJ4G0G42o9R%aqU!a4W%&@aNgXY`y8dfC+n8okb;EMZ++T*dG_ z`G{zbiBEz1f&Uu(KH=D&Y_x#4g?ZNRUJEY2#eLZ_13nksShgRt&|BX14M!a5bszI;VJ`*IXTE$qP^!)kA zr}xu^<1dWZDaem->eDuUNIvV*z(0?@r}u?bu2oA_3rIq;VsD3UdV*ihuib2 zK5eSY-x*%OPuJi0e1!OLk9z&yE-xnfDsbJGqI(8jbo~@gzNTtWSnnA<_W4!Fj~_ST zs^8Fgisr}%JId$lt=$53ARB{UEIxX##Jyo%2rqHJ%1&YZX8mD)KW3dag?1|#&znMg zUi9QsxHo=VJ+}2n@n?S@g5fE7-lxhum-6*rV}D#9+C`sK-xTIo?KYtI&Zyz}0)3wE zhLyqf-k)L2%e{fVg}h8hn0N74zJ4__+K+iP{UE%+e)-7-7-!+k`;6IP zoeEE=+wUIYV8ScNEA12F8^W=NM<}Lm`CQGfaM>rSOZWD+@D%#*@?6h{{@1d8TY~(E zzK-#~o$eU#&03m?^E3}la48*uD_ z>W`WCfdNk#=aw-3;-mHZd{}qFna3py>8FqP3RgWwSLD{~ZQ(`oq)#EYGu*fd*F3gE zUkG@D9^Q?4_wp(F3i`eU8_4QdG+)9iiC1@Y?(aN2{oaaKv95p3y0|m&XZ>#c0sO6= z!RX~*_eP$ruHJCfU;QJjBemPW7!PCb_wrK*uKeA%STC>84_kM_RbQd^ihe%e33mSK zFg~JJJO`^OWf%6 zUcU2t?t?w@CG2pkBQl&iw`EVGm)`A^L&Fp7q7$57Ye|ImNcA-wDIN=DDTZ@R|i zsSQubPi=S4eCN2S4(fi^=`n$SU9Q{mUfhT9#ZEfoBRfiY$*Mq~C!c#OexL29Mz8Zw z_q({I;VtarZ@W02;flLn;MSwz1;$}4aR$p98cv>T{CsC08ZJ9r`rP!S;T_1@51C(E zkA@T1+UvJ24rcfY;=Ro-FJ|~E_NhV6uQ6Qj@9_2EhL_M+x;N6|gu)a3B)xBY9r^hh zze`^k&N;qUM!UGZ;pnRqf9hO=KBe1Oa27>@tgFvrCq4JSW)bzhe+FuX`VjtT1` z0V`1lzTNu;z=V@87+Vgn;Jm=+U0mJxR}!~bi@mD9_0Nmk?lHv0IgOrrup90rerx$z z!`b)m_U|n7WCL0?5^RV zpVXPpJHo~Bj1Tv3kLly$Ifl2PR~K~l$JVd-OW$6I-D2@Qqwm0mx{2qtc$DEq?7b=2 z1G7VaRxqA7JAcaPSFs-NM~4okE#K}M>M?F}{-)8F*dJd<-&&kgxPGV$-Eb=Pfs6Hj zwXJ`%-$(X83%R<^^(Ujx?r-0gH!*tZ75BOH6F7s^LS?9Q{cXo?*}~~dhCR; zzhfLu3iQf5e8t6GtzDf5Oz@|?eCoRj^iPR)EzWH8mDFENcX<}WIiK~5MwbsYoOsKw z_aM)I3jHDPu;mmNS26lJ{DL!rKGb*euipxCnEefN7}u9@^0PhtIV0oW0-p`o193dL z|EgoV_g}z6e^a2JJ0d?DL;pFS^~zVAA8P#D(T8UT`P6rMFVb5suVM5>^guQIEpKA@ z3f4<8=tF(CihW1#owxX>(W_q3zaLh3fghdUv3nOyK815``nD(8*EYN7%IMSHKlnBC zbzi`#KWzNJFuxg}G)XUFz0VEw=+&Frk)Ia>jvYSvKK7lL0$xhTePIv(J_7of^@n>N zK4M+?{Z}~gWc1s%I9|S!zB0G&z$aOE?*{(Vy&SwZ?VcU*6nph+@Xy9Z-*xg){O1zv zhnoYv;s7I@KV|d{@VVICZ(w+y`_0ZP(@(`J{)28>9r(8}kAK2u{W#!~&kwNMM+Lk9 z{x9^Ew-X9HgW`do2K!TZJL{#li*s5(l^6RHI@jV7h8Lkf4L&WQUGB4LxQ+hXd2pj& z#k{nJby}od>c%T(ASXNUyUC$&#iOR8$Go1;=!qk3_uf5dh2uvLJ0s}5j8D35_6+8+ z==v$VfM5KDU|9Q2;>3hDvA8SDh%ic9BfwwEj6ic9|> z>_@^=s!M-3g?`{E@!cR{k#Zgf4_Qn^rzP+!ttjTU+D7Sc`)Hk)S2vD zhF<4}4#yr?8pc!kvMDa!Wqhzp8!ifZPyAKKGBe1h@OJ8a{P|hqLmlk?yYgHOq1_^S zdq@rQ8E~E7-wpm79A1Jy)hcs;67-Dl{J^6|Pjh(?wVTnC7dr&Gx-jrbdEaaQV9#WH zki%)g9uQtgXTS8@p!bBcFK_(6AXmbP+x)2|*nz?m>X$l@hf?SdakHCzom)QRA$glP z*Di9-v)jldh&tLuJm3h*qLfq_TuK5lRHh95WKDRG`J1p6oB zL!5Beps$1%(oOHq4EjxYk@J#=1-&60J$CjEW$=Xe2w#hTYVmn{t_t-3sOYc32Yc_O z!NJZEf8=V>IQaYZC|vvL@UXuKuVvpnBgl#HI_xyXH@safydhV=;3r}K6`p4vPj_*3 z>wgpN_Cr2j2=mgy-+T`}X6HDJUhj2X74(w$7vS?#=GW@_j9&H6zX)<7`gZu7fIjqg zf^faR;E!Q{SL}_r;=(Cm9~4fV(n({(cxLp}ft(k{L3ql(_DZl1Gd|RTME^**_QBc6 zo7d;UvBwtt4BU^8@Dk_u4|a8H8VBL*gA2|G_MY$x_9@lF-{8hyIC+`p?+f;l@G9;T z+dIg)aK-bl592AkmUu;Nh*t=&W4?x9SFa82Hn4v^g5NzQ;5twJb>!`_fOEfS^`@XF z#b5E6AGmywjYBKn-OG2!GcI}Y>KvZ}e5M9_Nqi(fht)9Nq1|@&wV#A}tb|_Y4L?QR zRyca$|662?EHPsiQ-JI3?nK%X*?CkHzr<3rr|kl<$tFOUz?I7|zCiqyC77VI3+ zW4903?md-^L)IVG>09XM$$^jb%>HHiqy1}Tm@n~>y>wHM2jMC4ihqSTs_Y9 zKU%oXW&aX=vdZ;8<3rv4??YTyc!4}uG02thB6U7!lX^aFcu9j1Tx! zY|55EulGu<3HyTRi{O)kJtLg_NHb?j`+ha>AwRP5XSBOAw5vGt8O)2-Eo;8Sr;@&( z+qu4M{Dn)eo`}4~=hZo&?_jHa_TGn1KFaG|#<)G=+O30+-bcJMdgDCkjK6T@$A_X1 z-;VUu53FT;&I-s?K^AQ2B z>>_-+?H~!dfWB))r(B#eVl)9at;s4GEe2}@~F3_J6@H*-|H)DtF-Z{R;ry;j# z++6x!4D`yc&BV@G5^&XL{;SORv>>EBcB!vhQe@ui1AQBI#eL3C zG5U7wncJMbWVqhX_@|oAH32Wfi|j`qxqZZN^1^+ubpD0mtH|#>g^t@RjL#bO&C8s< zWb`Fycc|HZ-V^?oI>2wf#QrCKKfek6 ztmOTh7nf`RHXq-qauB zz>R~P{b}^%DQ5iw{cmw;!|^-5+3MmjhPR@hf8_EAh8I|mFCh<>$1z;-x1y_mFkF7e zBKkQsJXZ(&R|ok_=zkIXpXf`A!yBI!>?3ys`4{~v>h%2iU!z~cx)_N4)8X1pq1XFy zMuz?fXJ6js7MITzpNyXSEq%PsaOxE|-yQUD#s~lC55zIv4DAxHJNqjxzGr->b31fD zcTUalQhL-^{k_44C!9x^k32jc_>)H)e=mG2&TsUUxeuCrJlSx`=fWDEG4SEuk0BSp zXJ){ubK{=aZGMXJ;rz}=zAnh{2JG`Ev2(nhPc_b*8~Q)^%nSUfTNwYlFvu5A1pp*eA*{ulkMh6Aclv7yAH$VLYjXCv@H{{G;}_q}{8~Mm;k*y`EB7-#uhOnQh9|`R z=Ao~ayZ(re{Ice-PK8rfQ#(DZcj4Ic3kI+cO6Qs06OMe2xfy*jB;cjA{_7o`aqomr z0b0=~x%pM=ksr(N7+%5ne=f+m=*gpf@+|ZFe&An4{O9Y8Lu0^G=KbU#C*p%WzxEiO zYkZ)`_A2g(yqy>D2F`0-(|I0*UVqmAJoBRWGj6vp!ykWq%2Pofs$J!IcV}IX3H)2h ztKNm*8vCQTpTGam#)r7g_0t6RG9{hrGt5yImcW;rVpH(aXZT zZ~boQm=>PvLN}hmsgqfIcChz^mw1ocgJHiGF8lN2Fkiwes6%+2eb(*`6Q&=@VZqg# z8D2#_%ziFUWH@qo>g&i)`L&N1SSG%~A;qvc}WPiWQ^(W(# zcH95h@VvsY8!zm2HqF2d)QCKCy!q0 z@BD7#&-szPhq-vW;VtZUmotykL%XsYzvki$Mo+x@q9o`^wX1tG_P`EY8u+vm7y5)B zdq1!7QGLq=E`Miukyy@O@v*J`S$K&by~pP!>@lm~HeByRInLFo8jc^`cZ`dB8?HQ> zuPZY=&;1qs&_7qZ{+DRCfZaGEtViKR_~^crD+4|C)kD5ghL7T7|HeMC`!lUw>9G~q zpO(KioIJ>g+gb0Ce--x3@9B@lCyicy^iSDmZwvjY<-HXqW2k`1}U@ z>f*q^9l4!eW_)ySz;4(%-i}IGi;6qF)8AI{_ZdxO;<;nFiz zF3)AS&guU<%!}wtoV%QXp1jV@yKp`8b&T`E(Ekd~JMR_jaM81WO?amaf9&@4dk6g> zdc}JugmD(GbNc7vr#uzft%Ls_=y^MjW6#yVoQ%W1S`_H>_@`6o&%A&u?vXJ5-k(yt z%D0X{XZ<>igYvCI>C?3VFMxj@opo8j+mPGg=;3zPZY8wJt13TL3b^86iX-`WNJdY5 z{=MLD3QvhYA0Osdc%J)_dIkL;yg*#y|nQuUVhw-WyVwS+uMnO z0p~uVrK7_>n9)a`H%QFh0eFD`JKi1dplF~RjjWe>_>kI zd{XQXpEt2~YuU%oWt}by^mXX-3GlIcA)}XFaVYw*OXyD?KK*L8Utj7EailF5@LjW` zPr-G*@_6R`#efr6nKg+$z|UhwkA7YmmLvDFT997-&mSbKwbc zzA)?~!Ydeie-F9!r;>4-8RoZ6e5l{+fgbp_<1bwLVIgxfJKzabsFxdF%TSkxW6 z<mKUw{71uOukOUYKcsVh4JYs1t*`SV3`Z}G-p}O$4M#5>dAGBZ4X$!ovc#(0{Y zV0eMNp}#lMaO(4ieAngK4KES*p5)>HhD#5xK~L6%=aL?tPCp+AIP&>i0e{N&JL4ny zECzW_7zg&vy}ykAX#Sni*Akz6pv?GS7yWpj!K~|S47bXO_{dJ)jeXzpea0uB>o&78 ztVhvHKL3lI6Z@(B;9SPX>lx9vlHb{}%)F>B@j@4Wwswgt&KrO}w79zA?daeiJHO9x z&WZF~ivGVgjAxPc^|he))SngXQ;#sFL;2m-i*V$#`XJ=_-N0uJ&$<)+vG|hOP5IF| zrhVv-`EiD$w`-!DWc0*yeBRmUQ}(0F*vCh?c154(eE}b$mn_e2^wgPr@=nkXqA$Yd zRrIsP2aFzne}up1*6@V-)x(265uXa~WBitj0~mcJ_CVg9_cXkUIPe87?q_(4KR(ID ziwu{)pRd_E_BQ@0?aHouyUaK!4(aa!Gd}on9lDR=$}n!+-}~wD#AC(V^KZe%S?~OM zmf1it!h+a7L+n*F&9M5p#mQ9bK57)bPMNj?qx9;!E z(^>yBKIx`$OVGow2KoZ=tI;lwZTuCttZ;cW!%J!XpZ6HN?Y$d@qbJeBVZWe1738_D z3i77@P@mavpsNEnK2_+?iDADGJ#ot(?=hYig#OnOf7r?8?~G3!_S+#r?}-ooMU_97 zWc1RL{axI{aGg_rl=(U@wA;dZ`5EIgFW@>)_cD5HyZxPZ3-o_IK7gHbH~z%KdtM#% zr25%T9nq&{aOAMhJn+e(UFE6$dz_8`3hbymg56SpKk=)!_gSYi9UtM`=g{jy7ta)* zjGl8$-NQHwPw_(sxHzQI6R+55kc)d5&c1xXB<$5YLc2xo1DwM8x-;M<&M7Pih7=6Pp{{NYhRvBf9?wWt8ybR?#I4ADd5t357lhd0k)3Bhkf~^on4-; z2qwG^Ik_0!XYp#o8_*~IzBR-1mNEgy~z z*U0a_p9SRS{UEoZ=RSw#Mc5(U-Vol7J+Kfy7B{nYJJ6HwpzF=QGrUMYd!kn@eu^`;-!Ym5BfcRoaM(2Pl#VV6YOn$$GP?4SB3c!o=1LuK;F>ukH!bN>N~*2g$yrJ z5Bcw4C+j=%de!U8w3}pl=70N1J=iP$et6?w$(ZeezPiHAqrM}4{>9tbe_sf=?7fA= z_pI*P_;4O`7ylkx!>Lam`(Rk_`i?sDAtwj>TsU#CwTA_H7S6rRL;e!xU3d$A%}ZfF z70x-P#mmdk7tkk*k;S9b7SASH{qVl;l}%;nsRtW#Lz(f{Ii_!+m#%Z;Q-Y8D*(Y6G zLcE3Ro*5fwQ5l|KM_m>CJJF}aOLqzOL&k?X-gfl1-+zS{IPdUs7-!){;&NBB9(C?R zpJKYTKio^!!@pqY z|3&oF6L!1O#ifm2{_HNaYk42T8_)x%5%2Q$kNC(QTaQk$d)16y?*UvI{7TU)uC8-Z z4~Kc=o|!ErzO($S@hKoD1JQe4&x=nRV_b=?Z|_qu`gZishu~IUCp_tlx6BK6lyKEc z@6R|`9{)3X>Xv&4J4f_7@75anlkveH?BVhn#)o)&%cSV{!H2rj{yr{b^s?KJ#b5UN zGwVP0?YH1>`4^+-98=HP%)8$oMPG>@%AM5JmZva!#ba*_a+N(7<9QYIH!@zfPK7J~ zvNz9b`BURlhaP)|J?W%C-@y94CHnR7$!GoC7MC_YIxjgl*x_nd=a>{v^KlsAt@LLE zcF1;qGq=5jy`HP<3rrczemjQ0_WP^4C_?%Me;1w*fZA#{w3mjdas9{->g66ZTkj0 zN_YkDVbOcWuMB+T$NfIci|8c}Q;GBcF3_i09BEr!nf0G~k^v_Mdry4WcQ#g*8J`C3 zd%xS|FN{x~dH(@_XK`u6buZ;VnBzW-qsgK8w;<2cz^@Fr;$V*>hetcS5_MD7i1AlkrBB#@MX!4jo?{$Z z0w2Yt&qseg5pe3Rm+pxEu{@yh$G@n4pFP#?!84qFyzlkcSoVHT!>hn6YgotexzLAC z!^``F>SrzM;&AlUC!s%e?92WhVDV3(=lsRSo_zO&<1buwbb4QvpZAO&d-(TZzJ#a5 zMO8=S{n?BU@%(4BwgdkHdCYy_Um0-r!H%yXw|>7CpAz%+J!I7KAfNRgtS9U7fWW7M z{c9yQ%E*9MX7#b#?-$}T|KSg zN$#kJ@4+rw-6^+*vwyAg@o>Y5|Ga$<_L!ZAHJti^n`g5hJrww)oa?V6Ph|dr(buvc zo$KP9hS#CT_F+F&yhtCz8#qU^FVEEw_~&z@-&sukn30; zOis7zRl`@IGisdv7oO-RS>0!e)3=5rho?NpxQRyk*>J6wx9F?+xrSp$wG9jMoc%55 zPixT|e*A?OsP}#oyTbB!#;3^r?LAx^(s0?;{{2mcqqkpu9lP4jBN$$RJY4PUZNufi zZ)U$R`@!%k=GVV(&T!;#<9zhwd2aj@`i~s0YYIcdBam{oPS|B za+v%AIhh{%NgU~>i(P!%=#@A8qKmT`-a=hyi}U*oS3LX?=gXpzu6D)Nfq4kIQoz zPTcaJ^~`*Ww>NcfbX6(7Z&uEa+gsQ>dfhO2J;s4~whdHZR2 zt^$9XB%_=E%$lrse1!9!g~i31t-OtYMo&Eah45U$vFqksALb?FL)__|GVSWSRlzVVFOsWAwx&7E}kh5`7i-8SWG0S-9d7 z>s?&Y_|)Pz9TMy%(buu2cPoQ8uunZ(hEE>()cGT?-^540=F8ahR?XBR-?&{{It#o)0$g>$~X z=FOVz@YlUPe=9R?E!a_aGQXBjG(N4^fnz)4 zPX2id>|d$}SQf^adw_f0Q1idz*-m`#>%_Y(k1788L4wZt0)Cmr;SJ}#f&Dgh_D58P zC&aaOaCJw9r`X|#xwyCCdF-nl!uaUB0(Dvb-Y}yta!%oAVc*&MUGK&RSw~jSZ1n73 z8xAkiuI%6co2jO*$^QNHp zgg3CiyzKH^#y^kTzJmR1by&4=Gh9|7wO4hICj|^8`lj>I+-FT{>m5k+! zE*@+2>^s9g7xYi|T=>O*VcZ5VUbY^E*J8Ky<~-Vt0oOUESAu>PeFJ)XN?6yz^XQYQ zE{Xi*cS#c4v&TBQhxux_*7m_D&WujKDRs1^@O7ruKPlD_r2A_i+@H> zTy9zoIxWz1UVBUdIrsKq#)tX3ljpT~pz$eCul_ySDY)?wz4nC}_@Qm^*T?8f#1DTP z#v$tu>t(U4KQVghg|wOhhuq9^+!1MbrZgbp4q{*o6!^3-3%|i=T0A^Pm$+O z`Maug+x%vHu%8d0pL+(J`nsjx#9rzW@FM=>3E_FghxlqoAJ(Jg>#f}cdApnMEx%@X z1^VQ3$cgt)#itV8cBf*0pmmPJo%R7jF3Ln*DjOM#PIX=RrXH*Mvb-)w; zzW|+R=dFy7#;rg4q-&thGj9DEpM3%@xxJWm^n#mTwX5?5eQMBYfxg84^&!uES|^@l zeN6Z?ZrqjmwR&K~)&JkSdPu_)__SkR4eM;z@Rakm-CP{f@I3cmtnZ9Bsx`wow>8z@ zyJ9%~pYptm7a2}G^ZAe1H!YuKIR4lAZ@W0E;T4?MKHtT|4aY7T^t|)?4970w{YTr* zT^f#GKl>*90^1i1r=DcWnPu9o%dPv%mJg?Qs8?HLHyIC6t zrrYcV|Bc*P39zKCa8)=B3}_cN+)cy5#fDUp79qzqy{a=6stDeYo! zJa`-H*UrJIUE$n&ws0-$rAziZWSw|Ip8W~zw9$c&?EgC%f3pXSJ`bLTaTb5#1`kg~ zFIoKG=!s|Uf3>S~GrYv#EWobI#r{x_{VVKSvrCP>f_d3B$eY@gz45e*gBd;X%$p8# zaVNvI-?cis$naY9RU7kY{=4CI=+BkdnHPugQQY7l_63XE7=507`7swSGF<0X|AU^f z`$-IMVIKB#d2_>gFGkHQ7wx9$b-g#d(V0t2X6MN(YN!w^Z3rr*Beegam;N& zABqq1HfL6lhdj9Q>9>dVB0N8^aqL%^ugPwFGJ4)aaVviJihxUxo$cy))t`(Hb(kN9 z@fVIB>+0|4H2NZSyu~mMqE|d>2imps0!Ghy!{u*=c@cdDw%*z9{EN|7QkVPrAU~p) z-?19KdTD$v;tYql`di~u%RIgp5sp6#cF%o|pRYGgdRsC-S@PgTnIzNA|A_`aGlOTLliwKnpYdVdKj-d2vUaHp zns7ezexG|@(HHsoxD5XiXEm;IewOh`us0rLj4WlMdDPC zyZI8HA2@En{n+!W1Ia#&PezZQ)vE?1-PTXx)FJfU;PT^YS9l)z`4a1>CGam`|BrNj zrO_9;&tam=7aLAJXTM{^J}CYP`e7vN+UjA9z5?0#DC|3;uVkO<5#~!c{&KHgVSf>x zvhPn1@+Mq$2)j6c%GxEKdG?+`FNt3E(j{TP5KjJg^!v=2w{wIyv8KQ3;seIN1=>-p zmj}amst#d)#yl^cnzPRFFQA7f1-&Q!s)w0|JWO+ZO5pP2nq8bne1z-1omX9)*6;** z@bBp}TzQbwU7XYKJbBec_#JaXyXcLF4h-u{?H1X;bnjp6C-(EYox%HMeWTl*{Y5WG^*M*I<_>-Sqa5VOk)u|eN9qaTM z^!%?wyA8y{=b|SqK5z80XO5)}%fA?|edlP_>FZ(KT9E%enM;d57`@`#rx9mReOLD3 z+4UXoX`L{Jb!UEkhJoh?b{q9@kdy2?{G+}uo@4YW^+?|V-_B1FAL<@^xwwka7l?P| zT-?KO?25%pvCn6?=hAm2;ty5e{|tCSUHTux_={foi{n`{R^My=i~E1-#>e5%AqcgWBpM)^S@y|ioQhcr|y^Y;~+dimStC4eTea4A0Kca`mlTGCvlZ|KS$qM zeBS7*Sid{7mi{l$>%ODgcwRrAYL|Vy+xQyvhZ}$4b?nP)!#oO?ettISw@URt>9+bU z_5eHIq5cTxJs=G|nP1=kjGl4$fVFqL>!)z-kCp6?c0T1ZKJ0_(XR zZg%|_&i*)SJYyRBqq^TE=z;HqcFFhk?Lxa&_iOD|FmAiCkMs=mmF#Qf+9kiUvM>6| z;?l-P`Y=U4H@bdS(NEn6evtec$6vVY3cbH+U8E--`$ue@p@AN`<$apl#~U*~*qNu` z7rY$!6sV_t2;9d*L|>#{H^+YZPN3KMvZIkt?hSQ<3$Z$4Z8dITw4z z;^Bt#T%R89`~|}k>MBnlzM;9*kKyc(Pj?->ZTt<-Q;)P1JLe5(^)bA_y}ljzacct~ z*-JSWpEUYXy85-lojqf?=4}b{H8${({@m!|MMhuA`{O3DpDqgY>VEX|sL*aZGV*|nV;eo^$LCEggBOu2;rDJ@7`aJmYVjUvIed z{9jz0!Eo-;8lO7*%5d4?zD~(-{Qb6-W%xHh-#e^}Y%I%CS24Ym(PN)4c7Bu5r|G(rE_HsH;s39-bAhs} zD$?-fy|m3xtr-mVi1D@sbb-;gQ6N zKth1afQTRxR%FsRMlNvx9|1*#h@jpeL>QF;HR3Spn(se-|L$UQ(Ky&^g&h7=*V(&v z?b@}gYVSjQ^89;4e86}m>%Fjyo_rx`5Ba#8u0XFGH#5Eby|a05N{?P`n;PmXOfP?a zT4HC^e(IWr3=R7))0eOtufWTB2jkk`wK3m+khGIH(q?_n=-^yded8&dkB-0jQ&Ul%+lv*-R#h<+1)XSI`>q0w)SMdkl6Z|6cQ5^CuA&zHUdA-}xbpbw2 z%-_pHUdi;u@hcWj3-MCp%9(EsaXI4|@^~tlFKTBOaot@T$bs`RFa4VL`W^%G=|TSS zyPw9(*p0&zeHMS}Dm(7O4qTMTXC?35(9t)x9$(Sx-kv25QL)bDF;SCHjnW1V@GVw&}~n zrI&=fhVcsdj)ftAVO;x^&!AVEf_%iEc;ktRuil*C)LqOuIMGkhvwjU59pVG#kDlzh zZ-_@3FOkO_nyfE*zgX{nhnyWBFg+-jfbzFE8`}fW8S(fl2FUI$X2l4`LCwbVtf0(R~;)6f4 z;|IhjC-EHHL!9BIZzlFc^w^D_v0?vW{_>MQ8us1B%dBfZr>)LI8P~lxuYfyWXP^!TNxbRh2_H{s0phwe+{EWC_A z@M-9)2|eq=mig%CBMDw*zx+rtkHn|Oda+w#PlPMK@JITkx+BY9coX-2#QIywM>uh$ zrFVyTo$(U#UzOMs(Pzk6-zSRmNO%|LDOQE}r}+?<7`+(%U!U}gdi$AY@o%ozGQH|x zFJKN|0nga}W$4s}^zE{MYtN_qV--Jtjo%KgbK+lP9KBx_|D2xrbp>mH?{PJr5yyK5 zdlL0T^kwRZ7AF3u@Cx7I(*7>a7vWXv?tYcbBjLI)=>6=oUI}s#PCjvRtOs*@$R`fo zn&_M8i}?R55_t;OI(}_3Uc$3{zcvu3GJl;DSDm5bJjSVu*wTYd|7W$Wo;cDT zmnQzR_*C)Rug2~!2=K2i$nW!KKQFQ z;U|ASv{N|sgdMA}+phmNpDOu|%My8tzDC^XDAxDr{|GPeC*DUPO!A;^lm8mC{I<%c-_W4y|DG~OEg7UMPUdDy{x zSyGpe@dCR1C+3CYGsc^k7ej)c7^hBY`?lcc8!vIbpi%oX)}2AYe>T19R6h{%JI1@w z&jW+KHQs|>eJA+a#>ulBdS?Ti`!F}%-JriKsTaB`?C;D+-{tuk?{(b6c!hYy%7w5%;ps;% z%=gC=d_Dd~DXhQZlhJ17_rIRVN4WBRFC^{B>9gLS3=jJt^U0{EZBFzg=fnNvO^Kce zuaJ+}nAj!ZRq}>Q!uo6eHT3gS>>uw4{mR=zzt#kQ$@JKnsqq|~aq2*p-__u~CE~!f z5DzhZhQG9(_0IR~8Q1>(zQjJLJ>872?h*C6Y5E@aeV0P-I#uIk?7-KP`651wS1Uj0 zc%A8Wf0yza-nSap_fx*g+Bu!ywf=}t1wFi_K|58S7Vkk7pEC5iKmSFulowTTqk3C)hj-h*fH_%VtluC^vNvBtsDIr z4qgj-D|-2cA4d(__zvB;zO?e`C$?;NSN1;X9|hh4{R2`JXyJbxDwq z_*A*SYu^U+_#2DP#-7}i&=;65PvKWMZ({yBcmKtZzcyZEAH6oQGkR}{{lfM5t117C z`7)Goah$|_l=okTKFBWV$9Om6yO4GN%Sk(Xuvc+EX!Q8wudfSn zSJRW1+*(WYR(vY#-Ksk&oz`*vEZ68DHVr zj}J!vj&GPxiT*voT&!o?ZJc||83^)X-_F_53<=08?-pV<-b)z zT+!_zUb^``%)4C@`V#v5@x;F6{bK)kdDy?359{)<3lh63dinqH9KGp#m|G`eL$3+% zEg(DXD=)!Ed^W-5KQG09xF^UVr^lbUwFCJjxc1`{v4_s9>b;^bQ~&ld`gv7CFZth{ z%rDVb8Q&SKch1w9kK}($Ws_EXAkdlKkMBu68`8z?~j-R&i9#~_)N_%nRd+S1$_wvvKZ)Tk%AQGZRvC~Rnsf(y?gLWjjzN%yg1k&KDgFZ|S@euRT{<+|@= ztJsqpSx;SGZJfO4)`ydE%-c`i^BuuYGd=70kds3E&^U4R=Q~0@zj4LYjN8t=ltItcyWJ(tyhe7cc;OYmDvFTd;KVO=&}#^-He4ZAAgqrB(wtY1FYV0xX? zIV<@8#udj?J+tEv#w*yRcd`C13H-}E<8NSsgI~kcepJjUa+c$-Hg7F^gO#I%l@iOB( z1bOb3%e=mh^Sq4l9R;7QiF{VV=i97FM+f=}&wLN)Ow}9PFBN~`s=K>6_$S5{hkq)W zN21TDBV9*-UB6^{BZ<-E*~6a5iAd6}ob7xurVFA!H>fgZY^ z$~g0P=07L;FFr-=F!Szw5 z&sdk^_g;-_-?VSYFBq>dzNcYpZc5%;1%DyYbG=u2la~`eDsK-w#xcH8--I_|W7a16 zBfN-hi}#AVJ^H)5`rMxx`2L!4;t9j1(=XSX8837H=N=*cW4w|T%0Eo>U+=BL zXES>KuR$M#*RY2}!@kCRh!>rI8G75B(CZ$L5$utCZ;$Emb9#qi2OR$~PP}M^zH^u2 zlJoWGdR%u3$PRmQ>t5`yzZu>uT;F?oFqua=J@pK~!}oeGeX*Q{%br{a{~Hp#On*B= zyhZI4eTBU3)ycXqyvpwnW!}BU&jeT9%V+5C{Yg6utmD6sjjq2pgI4V-6uz-z3_AK) zf^!byrR~AL?WgCQ`&74SknH^3jj+D)j~v;O}~V z<2BwpD9m5u1={`!dzHhozWx}O-+NThALE>ZSg?{fmHh|fs^@t!?7xiT_fpS1VCRji z-s69mFH7rkHr|bWj^{6oD_?nQSYM2jue^Dm;13wrxuI>0Q8DS4;>N?Fcb>@fx+ivC zurtQx=jgkN&SweN80nts;rOMl|1^$&yH)F~&n+4+U{{9+f5v!*e=;uAWf?D1SG;%d zKaE#-_a|tl<6Xw9tb_XA>+&E^wNr89D`}7OfTl07=XfCa|Hhk`*Uw@neLmWF5kFx8 z^52s1*M4VS@b^rwxZ*U%-Sxx9m8V{dZY`p1nm@)>=QEgf&2<9Cwcq(zuvf;*jLV_C z+wnl-%iz5a`{i8|{_=CS@H@j3T;G5BZ)9;$z%!n?XJFkZ_JhttiNA22Z@MIrXHHL? za7XY%OrJ6Dewyfi&WF0MV%k-W=rl!vyC%%bEii$Nlp@4z!1Py*%1a^H*KU?a90lf32tUus`2T_;>T(3y_cN zxXh=A`MW22xL5MtGIroa`gKHtYyYw)v7dVHO2)mJ@x46Imw6^GcjD^c$LhVpiDxdq zKWTqXkAD*HJu`jAI`Rtb8JquZpL}vYtUHt8^GfKiaOwg6I6LfH&4;+ewii=727L{` z_n(samA8j__+K$EE>HM3p|{T=i_WAy)C0Ww?_oXldrQc99P`3;P{vh{w-CMZ`9$O7 z^_Gss7PyYfxa{+hILH!go*c(_l)IR3!XF@I{j%6ZS|uNtqBr&}BRI^zZW^V6AM zJ|ASfiFtHC{q_9}#)%Ic`Bh|kPSP*=DIX60yy=y9Yf0v>azAMB?mKKx_eeM&-4TzZc!tec+;rDOoqSDrJ7!L zeo?ZHi(dPtchj%^lJ?ZV?+o^;i}v^VpS{C&;dyMqA zsTK3k^XC%0Og-e2!QPr)?JqKai^+SdoV#o%Ug~-e)8l_~59ELyGp@Mn!(m-B-h?iX zK+pe}w5Nza@P6jk)yaEH*#G6s7w0R@M|BHv-7(&UE+|jxJdyEk_#aJw?LQmuL4ST* z_Y(^I7%!unt-%f$*LMm|3hSNmm5kLb!9EBt&==*YAL>9SLw|*HFHP_7!+hzd=ex;= z1b@=F&QrJ4=bgA3FOzp(On)7ZHLh{|e#na%XB^iY&wGE9v`71+F4igEzh`>Yjh`0$ zYU8Ru`C_7n>KFTnEs8(8{>}6y_-qJq3FFM0^2S8}#ixsL-_DqMpJ@7Si z#uaag@k8Sp$4w2!MRjXa*~dBmEBY)G#}%7D#ya&p&zgtEm8ZTDyW#zs@d9znkB9kg zT>k&t!@kUT89B%AVHwxHZwT#ip2c{T{YAVl!Z_o#UH1UHUfQ_E>qO!h&Vw7ruXy|J zA^u>z2>l(bjvK)yL3o`_Yg0=7(H|x zPkhSI=w49Od(RE@!sS0KP2`-@lP6uu{^;95pM`7ReMGRg=3hph*C8M02aV&`jk_B}KEf5Z|9oN}a(evOS80#^Vbf>m?Y|`J zOU{Ql##Y>|si;_4J#_dRBO^=*z@i zKZQO&mhj;`$hOtYCHo0}ul5U@(Fgn0vTJ?ozP|G~V9$*!-g{Mu_ZlxC=f}c2YdoWV z^)|-E@p8gHU&8w&+rf5!hH9*8`paAefC0GeO!!pqbCc4oiX0S`n5dR zG2`rO+O7=yEa3%s>b`>8@%5kPS>t3}dC$vfzwNN`0^@sb@Vks_9skc@e~g#$uijm+ z_tx8tGmqZ!T!Z$oF7!@9c0R9UdfBDbVIN^!^JsqXXN)%?w>^V>HeN*b&oTxRlK!%e zpL1QJKax*|ebaa1Y-h};3;XsE@0g$PQQY8|uzwf5#!mT;_cMQ61AX6^uGoAe?ROrb zpPo3(7wE6ga~YStIxnmX##w*2-wvNuLEpr`!u*ZK3_XS84UOGNvY^P&Bp2!3^sJ34dz_~*UV%^izKQWFdcF?%`}?rQ6(3#@^xwGTd~~q4#wF)D4f;&hjvJd^_IV?Ft&b*h$e^Ez{@ec)eE}NP(e1)`t>syAG_Lc2lW5P=bv-n$b8VOM z-s= zw~jLy=ltxVPeNZy+F#7ic@6BxMW4|&-5an6ezL!tX;;zfaAmFQ^Akb`gK7zbAA~&H0UK zX&yUzq<(kcskN=Sy|q_vYHrEd^~D)8UnDwDE7o7Ux`ftAWujBgpnO%K6iX)T$PM>(3J5zkUoq8?To%z46 zSu>9BoRFJT$Vv<2}4@p5Mgr zbH>l^{M)r}q8FsR#krIR7NmivonD$Io;d!5PIot7k@~EICse;+cJiTf`ox{2+&MB_ T4Oq^YG^KO=M7bVs_UHcq*NJDO literal 0 HcmV?d00001 diff --git a/tests/files/mtz/6SXW.mtz b/tests/files/mtz/6SXW.mtz new file mode 100644 index 0000000000000000000000000000000000000000..7db233a9fd573580ea1c1552f5777db148bb79eb GIT binary patch literal 81008 zcmb?^cUTn3^ZtmUh-p<&QBT(lh`NedXQyfw6*K0nn6qL$vnz^V1Op<7x)RKwm_S9D z5p&L%U30+fiCKR&oPCb{{rvgc=i#`=m$$mQXS%w&s=Bvz$L_`>tpvdbh5yeVfZGrL z!pXZ{i=l@O$<7%?K@hCc??w~|PAMwa`s}7^ z4Vp__Js}xol9T`QP>#?yr_Sw_1uDUftX?^pGSvXK=Wm z`o`K>^7_)l?CtF>y=a%h;h@eQe_wJ8dLq7P=L=^}pXP9rTE4j%8s1$e#sw6>jMz6E zZd5Nib%Mn`g2lvUH8HW-Jq|ah(>jjD#jsb*=s(cn(X$5Ee?dKTK1_1+I!~&{N@atF zb9}oprqyFi-s)iZcZ=J~-JB1ar~0~6`z3$9fB*FlUN)JH z5rA>xi`g!;~Ow% zq6C||1c_yX3ZYBlF^+Ga`GH6gSPCj7E z61lPV{zdeDFO=_1Sn;_UO4y%5uM@(?5HRA?keut~~ig4DGkeQt8%m zo&FjQ(4wp>rbI6j@9j&FXMftFlds`cIK7~uTz_h?dCcoPO&rG^(LG-cw}D~!0<e|@(e!*C%e^cYx%LH&?$!&d zgzlmFC0V9j|BK@ru)*9UOz7`yPOpT?HJ zjw=bTfoCdyx?mf@|N3OY&}$)CnH+d|yi&_dofC3#5nk zv8t8x@0EPUS2l+pa)jUZi{Sz4?!7&wuFFHraRDOPA7)<}?ysgd_rg-?YsEp2pOF2| z_@y7i{nUi_%_V8wM6+hQf6otwgX;dq6CXEOMfUNb?9uZa$2Y0tzsp$p{7xEQ7fbpF zC-`OC=(qMUs&SiAaafB&=BT{+e7wCc!JoK6UAL$s#9h|rmkIx`{(|~Ia+KZ|9bq0L z&ZjY$z>SYx@xY-}c^vnd0*jQ`@4R#$-+R#FxjW^9A8j+2Hk`k zC&R=sVq^J5yGh*m2@n+G39n3xMQ`s15EpQO>mR@)>u8Mj_(}TtTiU0L;^bSQ)W8@y zzwxJ8bUNF_Ib)st=lI&-s@DbV<6FusKKG>c$^MrUjNzi1Dqc?1`1pZCHeiy{^#NNewdln zUwrkg4bb|0;9tI=eP|rCk3Uu$)EbVo+$x?qT@lBK&Ky4ggG+vwt`^8e@81+!KO8vR ziih)l6t9naM!re?gTMU;O@6k0Rr?g!EQM98W)8}{Cq)%`tQ#K<4^Y>;{{@9UUlY?m zC{WAWQ72!+{nc5v5ilk=me#NC*uP3~4);@+1y+aW)sE5gc?oBZ>ie(pL46>Ulr}po zHEVsZw{;uWzDaEukpXE17SZ}tM?SD@Bq!gf_IIv^6*okRF0Y=+y;Jn_Uz2Z88+IO# ziGkO}skMh%E=0Ed)o<2+LH)M3J7lb0Nd9L-tk9hKUktY^Vf|FBwaLYt?pnlWMWmhM zD~dYzG(HS3VRjx|!6*3Kdk(iLyOUk9+ktsvN||l4Tj_P2{#GT<(^+m>%-1Yb4ENdc zrHgL!a+b{65B*DoGxAMmVl55z}dn@tf@%`{!<8U~fWD^A; z5y$js#Nqxlo&pT3_>yp|_M4kO0ob8>gcNd0rv1x~|1rO-Mg=dqxp6V4pAD`& zU5e4;?uciiH(CPyW^(css9k9yj-vff(3TOFckX|zpY*qO$noeMxTMY!(|F%sW7d&< zjm9Y-YjL@`lfxXYUlnnzz;6a`u9e+KVW!(I(X1}>3lws?E5t-^w0;6@2^g%G8H-m zB$0nQ5owM7r9b2Qsg8FRNHdxZBmPQlTsR!mqU~eNlZNM_{rhQJqc(E#P3raW7qD5E zLT16gs!z|*=kV|Q%j6qX*Y}6v*ziU4e2U9ak=;1?2K5qpW8S{2X#ez;@ar5dsF&J> zVCvhO;d8vE1a zA6MnZ$0{Z5^2AAz+o}H#Wb&UmT!1-`W3kTDuVlY3SQ4A915ST%kF6#@aGqq&Sdi&6 zM%>7?Z-piPn{eTp9A@8gc1ziyEgatlLqlHS=Cjq!FR(ZH$K^S`0@V$TaLvihV%&Bo zOH|}%PCq-ihc8FhEe__W$y+Thi4mN90q0C`mAdU}PX1*g?V(tFg}q+}{B#}g{;+kl zKPgQ9|9eip5gSI_hZgyQ==u86_^|bd@lE*LmJ>Y0ouWt2;9veFJyzN`YyYL+XfP~+ zTm=@={E5@Ww;1k+8)KH6XVhpQWyA)gkrT>2Uw>@U!3#bexkl^9QwYi|%;^_^Z}wc2 z?lyO&{aZEozw_6M)t|;ogElrG-#UZZVE&WVzV;qk)uPGuq5Ru)>c1lf^ev#1ui*je zr_fO8SxbvKBi1JU?R-2?s#jb% z9Mls>PnxBGo7(=-{PBDOC*P#*@gIckZXTrfzpbT`s;%!@yG(zh>d}1$p0ZvOos*|o z@|Nzx;Rf~L(A&`VVW>znx#1JGe`0*H54m&U;cF3OA425VwQT>waJ%xls~;{3)8a=S zmZXrsIQfc_>V6X(DhG-wYZ}R8ulDBpXHy1`izUC{F|7~hErnf{aeS*%aAZ7|x%plU zI(5ZTdiiP&2gQAEIHp#6Bi`GbXqo1_ki!MV5VH(D`rQ_tn{Bq_>B81mroT0SpU7YpFR$}zNlua@U;+R5o}m0G^Kk1f1C%~2lBeRkh1rkkJI za}c0m=zcT@7|fx=^7=%^U*hxwxY$^h-*=d17X6DhSu!q-<6FV8l|Q!5y@%|(k7dc! zvK-$Ady~#%?W+aIet7!i!Y3TAz_T1fL7cge{Iia@NBzmYKX$kr;NwJhUumf++Ao^ZPcG41AyLBs3HYlDVzd_PRIRfDU~FVgsA zf@=X5-(%0uAFmggEgg;=L9vY>sgCbEuKxk38g0_L5EHG>$EB>E1C+uly($! zB-_{nevV-NhlUH9uIlxe7UqD!RTAxg4<^;%@c^~r>~Pq&ekF}hB93gs{2RviS4UNQ zAeC4`_BpaS{EkmC+)tGTBtXLU)wI8Rj?JSx>H4p=4{Gw?*QA{3KWKfg2MNRfu#beB z)YEUvVEa>R#FXAgVwc$k-;7AG&qFB4vh;=U~w$GP)3`GPQG zkR2BVbutI-^zrHH*@%;Gl`?M}!=q{mvrF#MK63(J@qF^%$70|u^8F?Zvh407b2vbe zHrFs|N(FO9dR-s?a?ZSbXm@`X9uCfFc3Jb&GWG6Wj&B237Y{slA)Z1{oAfZwD^i#?nV-a z8}P=~PcSfkF8K#ztYXcN||`+rSlgZ4ulIr+f9;OtNszQ z%Zx%zdTuBx|DGQjZk1kFzl4vQ6*CX2Tis{v+mG4*_V4`3=0l5$CT!4R1?|7z$(zrg z)#<0<0R80+P~*x(8s8h}(xe=RTj9cy{P^BGM0{}}O7=Sy&fzvVb+;GBFY z@BxP_aN8>tM!b$B|H4_0YTb{+?J&4{0OtR8pWd&&mVgf(dHr$Cipsb-7uo-F*W~XT z&TxDKer(+wS9D7hgIboc1i0np_(n`x@g0gxj-&YPRQc~0F&y87-F7U2zndH+_bndJ zF6+wSfSnI7fukJ~$Uj_<>HV5;IOY5P-VT08BFMjck3$-_<8Xi6_AnF{d|gTY**fg< zMZZ7L#yvHn?U!KWKNKIx3*O#4IsJ|5oL2R* zbh)*nOJXj{e?H>)2K7MWSS<18hj_;8yrpMo4E0M}A6WZ>8hN1=KCPZi`=<(o7vXTb z($IAxJ{x#iJhO0wK0ir+r+pOV!yb2>zw;8s*YD8&gUx>?-==i0(OIg#;3>U7-jsil zIsL85t&h&~xg+(=?d>ak$_6$2WnUQ|l(e&x<;f$rll>h>`$sRC(_i`U-XJ^OGn-Q? z$N6Nrcjfrjv=Uo3V{$@aieEjT{q-T>_(Fy?mt3sjB(vx_&*$D0E6{>UNcE#1L9 z#&HhSA9Ut$0aEt6;ll;<$-hXE2T`ozPx$}?CI(AWrpJ-}3YD}R!k@SmawJy4u;Sln ze|u9-DN~8l&jwyk)1l$*tdJAqU z7f1OiSBvkxFPwfR{E|=`+&^jYkDs8+@BZQg2A*kwy(Vs^{u}lA`CmBY_J6uSgO1w# zOMu_`Gp3(E<~h_Ex_Td=^`-Xzm|p>Cys#8z6icN2?Z*Exf33K@pReRye6Bgmy|fgR zc|kWm+4Dp_JMyrUa=R4m|B6Aqnk;^$wXfl{|Gch9>o%7$)BYUF7vyn&)oW7^2pXD1 z`M<^(*sz3dd^EnFddq$co|Xs~dlo7!$6Xx4;h>sFq(e^kFxvmQ%D0O3zJ zuYwRX$qj4NzD4=-6#1q@E>3@|w58|{jPCnUOgZ?`Vt3if;Q~D9jys$`N7YdgUCPg#0WW+dp~V(xpz9mUY#aJxw*=f zKLE#9VCDixOkKZ}@(T;(A~_;C+zy{oGa-3N1o>ax50|lbF24W%7}GgbYT3Lr*?(sk z`QZ^aegXLQOL-{X^%U)2&q0w1EPuhCj}@0io|a5;uC)KGz|DU_8v}KPYdP>g^Muy_ zld$_{VV!;&9-zM77%z=LlX=jftI|KEAL`oIaDP>JPyo_OJSG2UHH?lGIKH1+G|2&y zDgGN55Cp1f8MgD~%Pg$wAUt0fc%8Db18!EV(gSHg0 zd_KbVmrQ@M4@dezz`5D9KljD8LC?7H1?5vn7DVS+LGg=`a`@CRPCr439J(HhyV3r# z`!!1uV-z<&*0kG^PFSMc5_-Rv$ZPIJ@q8gTM+fNfg5+mb!{MhsaqU~B?OOxU%kz|& zI5^1iatO;GGyMhVU3n0etC=KbJ#Ayz(W)%hK0uDlwP-iKqWIe}OD^@l`lnyxXN8Ir z4QP55N$YE@oOfRWr=JZT6kj3zJaUTc%XX>YyHT8e3QXBh10M{wiD&Y^m7kTp#qsU% zvTIN5)+&+KXVK!GTW^20=Of^ldU>$syLsdvES78RUBdAVSRr#AUYPn$^uBc3645)7 z>z@(3R&~NXXBUbdx59M$BQ`!JJd(Ttc5c-C|C`t}k1r=5uzcA=(ked}s-G+h=2v#E ze|~sw;C%ReY!&%u;bzsTm20aNfZ~yGC`Bq}D?yNpZ;{*K!sAUHG!qoyhNPcMy>6cH}zJ~j&eKQ+E@6oB` zA2^}cio6`|r*xT4fgJ_zwkm(coCP%iM8?H{ew zw${E)*_yvFPW-Zx;%9O4-R|1^t$i47RW{p$A&0n?*2gg1aJ4QcAC!B&X5yUZX`)B! z49k}Q3I6P#pj3TY3ujEWk^l8Yc5KCuk2US9Un*|u>11{ecJ}GL<~Yw6uFUen<69_x zI{m3!#+1diZ9eHi9C zm!7|`+`QjN4!1&!gPXzIXFSP&fH$A(*VpXvgK>#o;Qce2*2kyt@bo3_{ZQb{*Y&V0 zVlm9_kCShQW=@UqqEiVuqq{HH&zO1KE8KCYnWUrlZ?brQ!n;Cg#$ zyj*6rShk5GcU+=yxDo9`Yha^dT7BdjImbMfKV$MuxGHHEJWpIf@5fIpnB0-$18#d~ zz~P2)@=xRB;%D^hlh%Jfbj!IKo{fm2@l&wHS`|}pqDNDb9*Z2Xr zJ$|<2zpX9hr-G&SQCD>1r|BpN^@}EB``o_fsV^LS=H)A-_a6zjN`;!vf~>GDBtH;; zHmRm-KbsFfl2=K#;*V*4KOiN1Sis2#sFo`odi;zczKdM6@(7M^g_`00!Dmts#fMtp z^){@2LzAuj*kGY&e|&G=Ax>T1+%l$WnQWc@#WY-jmkx)aTZg$czQyDnJ6dz{?O;jh zf++u@_BqCPQC#zZ_fg+#b?Ps7%8{RcZ`#7 z#KbNeVMU=^v_HFzk-PQlvnJn!mecEC=&}vupPj}l7dvo#!0zYlupy27$Kg%nm_1WC z+z)H*-3DR3R+4>qi@m+takxLuZW1Ep-`$Jyhhb8>(|`4=T0iI~0838nC>0J2r~1_z zQXdbtzS{`XK34ToKC5}nq5!Qn_2BRPCBp;M3ZHhGI~A=f`QF=?MucDOGu&UDl2ucx zmK;y{kq_h}vidoO`>B>mPUyX58QI?@a$LY_jt}a&>tV3xSQz;qvFMypfx}JeoR<|) zd=@Xpy}3a3-IqC>>K8sPhJ>1l6hB&v6rbU6gIeCJ5*ChMPxI@FuD+J(C#c1$+=Z}$ zHhMlX#edlTf#G)L#Id<>zfuV4w+MgNe=}TBriN^V9*?#7=~Ftt!PZ}f+mtz(3o!ZT zXR6;?X~~EU0UC4Ck>OUwdZG~y_KFq-C9f|2#c)s>@7@KbK65BO@)M~>kL#bHtREGP zKZZG(U8<*1{posMep>F5mF1?;%^b8K!zca2!e8^7$rpr4j%~1^_X+Z^D_Mdv%kg|^ z!$AY?yc|s98%y;cZ2$EipK5+QAy=bwWPcw}eFm$~WH`W=^+}k%GmG-OS(N|Z!pn!2 z6BnW9@-yVWM-$$k!!9mhSb#B40O9 z{J~RK|IGLT*0Gx5nde+ue+uZ%CotTAeubVx^v^lue|DCCpTA+a5oe~_uurS+V*6_+ zDE_mS)8B-Z^EUvW@#{&xqwahQiQY-J4w7=V;t1oAG0J<$ZAwl7~TK->GpTux04ovGQT?<)|U7!Bv{MX)7tNLa| zGpT6J`ECj*Z<){p*Ag%m31izlKwO^4UGP`^ylCx*mL>bC{?EfwW-ZH4Gu);$yZ!;5 z7Fkd6j}dbI>1#RpR^{6EcsTk<^RKU>SI~G42jyMdW1KRyqWR3}wmv(JE;PTj50ft_ zLs#a-;G7Y(KU*e8bc*5l)-?ZPQ80I67{wp=9bsVr{2D2%ZJxP?n5cB`LzD|%9RK6{kOrRPwDWZaVW_zCXa31gX^CHlU*y~xv*tq zUoOfmJr8iW9l9n5V!I({$o~zpR6X2|*B>{)O;}f9KJA|y$@|wZ|ApyqzPA<<%fggZMXn)@2uv{*m$?8)Wt|&Fk zPhia4Wz@c(Tv-m{<03GZ^5zG{P&?(og%a_621@xl~+L(b?>a&ORUWo`OI1x2EN~G!UC*C5hf0+E{+1 zm*sG(x~><6Nw;oM{l#9(>HHMK`=kF>smbM7d^r4xC~UuAS>eXwJO8y0iD4UHX`|`1 zKfH-wluEq*FlxUO`ZTBWg@u;NUk9=J5XQH{j{~_dW7aZ?-^`P*ed3>w4Z39Ylnw>w zr}fP&os1aEy`Ktf>gSLDEZZ)g+3I7NeWWC(za31zLveKaZrWdTw3Lfs`vca#fY-hk zhZP&*$^Pzw`nwKu@(p;i$v3EZFPYx&N%D2?-CX~SXp1k1F{QVN?LV!Ny^;=cxCtw* zy$!P;t)lpmB;{&(>dG^WuHMXDC-dUhi$P(uNoZgkK-3q zfcUAGX#ad3UX^6$Pgwr~@Mg#P(xk)|RDTezt3T9mpq~J>x0D+ex6VcD%R<if974Un* z-Q=H~lS@5(%<%jVrjWQKk3Bgzows} zWLA>UIQys=wYQ+9aRV0)w<)&`PH^p06vgkt;CyNv*FURL+A|L>TQZO2FP2aFu=pcu zACyyVjhOC3^`pm*%ZZ`?@cJw3?LQ=&tofJIAxB6&FF);t(+%iH>zDV=2J(Z}!+E|? zDR&^&{CRK00KO;7nJPz9&x=QQoU`u&u#f^^%hn5%$&QD`#{qBhSI=XOt zz_Nau;AhA>djId@k1eb|kLl-!DWK!(_9q_l%rbED z1MopoF6pxC3#~rH9DR`0A8GQf^b??_Urds^o&JmBLm#EV_4VtIhWo2s-n4-eS=;FO zR>2cuL7jXJ_frpjm3m%fByVQv$8+Iw4?Gw`9FupE^Og&yOOfR0pFeqrTlM@d^wTj2bp|D*);JMTt!aDH>%w;2MC+3S*4g@2H~%&Hf^fJ1VzWcL z#PsmO7L%V7{hPjhtE9BHLHD`~>HR*Ar6cw0M>ZdxJi7tC$}c7T8_7#vPUhMNc%3>5 zj>d+Oed~=m9^~e5D}3*o8?RD)Bh+D$JSQ`j!)@Rp^uarM`7ncJz8kAAW#dcv zU#E+3+>l20%}aJW6wT>p#GroHU`6m|ivQP@^HdDua1(aC6%XU1_R;%)760%10l2Tv z1;|JB6Jx}xvYVUUzH9pXVcX|S>Ad(%NLtU&F1^?2^x=>)j8L{qGD_ z_b2J(XY)aupOY%@r1;Jyh?=C2?`b%|#~SaX!fB_;{}>2gcc>iS3ZKjizw8Vz_-T9_w5nmo&9l;|KDU5nO%s;i{TI{XyJ!1)W7&xF6kq=;_w-=(T@1HF=eLs} zphN=cKNfGcXZbOPlmAvRKNM4QQhXy4{zAQd)c(eRkFp9uNS814{MJFIpE- z&%dzrJO0XWfAw{U1JDoocR}#GK9Avks&*#lW!D^J{|D>Nw=f*kQhtS{xI62!>l1&U ze`mN!^(j*yBbA*Lzs^PHmvX?r?-#?3YQcTMu`!wqU+n|A21Q__UjjF8Jwfw#51oI1!^sCFE9EtGy*Z2CZ+~5U zknshjVV)3>x>J2rXxOj#9X&RBueJZuSfvH^>xx|*_iOgWGS!9UHyADm>+BPyw5VtF zevgvq{*qtwi{UhX?N%7xZ!y{LS-SWk!>N8}Yb88ba~6&f zy3+kBmEq6lXN975ta#++MNu4PrSa*&;WjwgvJmzzypo>3O&%j|=H%1)WXnA1(o=!* z+smZi=l7U?c9`^`0Q^;JJ?XbvcYgrG1uRrag2?7OX#Fdwt1o4^0j*=(!Av|(`^&1j z`~ky_nA^1>x|cdi^*4DeS`7Km{+V#n*YhxCQzZFMp0e*fmVaSiZg!b}ps; zb02wP*?Alv)H+SNL7u9ygDoXP->q{^&zz_xeBrUIg@W!UZ!r7 z%uPO%|M^yGS!Rpw{nz@hDDOwL#-YB)C_Y)y(mj*KFBspZ6wmh&rfyhI`j3*&U$k-U zTb1yW_3_29?XAv#O5kuoc_9VkgGDzfzPi}b;`?CT z`=#k;O*?vy&RDjZCAP0$Sgw4S<)@f_^hVn0ddOb9}t$F|DNKS+KzAXTZ`^2f*g4WM48{V&U^Fe!=)g^gGoDYsBm% z{VbNZPJeNH6ZYvn98NY*py%He*R*H*55@;P9n%k1e%(xbU%b|vtgZ z>2FmEEFAz19ph>JY=vh!7v#-TP?`wA@Z-2uLqU|QckVS_=9x%TZ)^YCNHtVZ_- zIQq#?=FZ@70SBFV1#3Jt|D?CPef1m;H{h688zASobtL~51~%`=;YM^={th;|()sUO z1La&hSbUx7XTnt5FsSil6ZNkP{zxvx@d5LfJtkE$klSL8FP^aJWsuMPo6~;OjI$23op*)xV!w|E)@g{ta>bAiAG# zR~}0dlbiN$B!l&z=C^Sg&RR?Hxwexm-}~tOSB)#lg(f2@$dj^9aK?Zf!i zw9$_`;@ zIb5LnxRThTS}cwKK6!shR)5ag2WXJg1gA{fO8My$7P^0mYu^fwUEaaj+gknZ0Qtuz zZGWoozYP|-mcac_R#NZv1AlK{OZkr+@~8$ZKg)1Clz5a1cl#}+ z`nY^@{XF{h2S^tIo6f5QwQHRs`+61H4amv0Pv}5qF}=@Aw>lN3`Qt&w2Kx0^>t6u24Dykl4c|)h+f^FaGnbBUCEQ;% z#^!}(DWzy$FNYjst8loVTDGDS+^DSO=ffZ>g2e|k{WLzP2XeFn2j{2ce|Tcr7?I45J%R#E-)9QpHG7XM|q0>g6d1J6~_ zUH-ziwua(BqPIo&+Ls;8z8o#vv@q8)&G2<+(f473xzshpg+2c6<4EXoo$6@-A zt(1Q(Ef;eL=5QlEyf_N-v_DSc-&&Oy^P*f0_qrNBQxk{q+YwaB2O4 z+V{j|4SYHIe&}zi1DQ2*Q~XRo`ZEsq#|vX;NR`_3p#AYQ>CtCa->$W9qaS~D^Q+w8 zw5S;I7ekL-WpsQEr~F0Hhf+m{YUCfc1i3;<4hMB*tzbBL>JjyCF&<7H%i**?rt>w! zqS|GjuQ@qKKmRrPMz!&ab<&CZgGs;BQuLN09N(aBI$0i4;~Xgdb`H8!sm$TDKhD<% z>K5EX?{6i%9?bLY%Hb2&p}5mx+W$3>qc)A<_=-|6@Gg9*l0yEyMecjt#^E-l^5E4l zv2hHIzlx8o%{bhu{JrK0M66s&csqG?jNX6K`Ugs*>Yrd@;6n261LQ;X^!|m03rh5! z%`oOD-T#vK09{A4;N_>?aNG#jYijS`Eo|wow{O{e;rbRQtd-M7`^P16gB$wzl7?HQ z{ryWp?H;s$k4=T^#joo8t86~(dQ}FE{nn6sctm$UgF;*l2RL%>7Hq!}O8Et|+{a@Q zuOEDDZ@{xHqGP zztq~d!>G@f!2Fr&%Y$mj-OlRQM-3OS%d_t=FLfG5zBLD6U z)dCIXa3jv?vkLsC@2BUJihtS%O}+{JcWsc0`A;GL_=Z%%u_DI@ye;Jdqem^OFAjm$ zMeB1o#lM<&HeWcsNut>7!Sr$kIouy>--(olR2rq-Zz&lnu=^9V@zZ_;A^&nuEE)ER z^nV}+tY!5V47WmS~Y8I{)ot>Axtq_HQJEwQmE>e&sEfLH^x5z)g2_@7KS7 z)3jtIaJ+hk>zUMEXM*S077lnefRKQ;M&+4+mY1-obGFP4od@k_p)FzpkdP8wfcdOI+l{F~ES z4(iwVFx+1~9TNmj2iDX3y+C)r9mD<9fWvvQ%HdsPzn06|9rOQ=Kf^&aiL+sJ^9^Jl z=IZWOW4K8@d^!?FH%un`oPt^IJvqKnZF;f=Ja4vz#^0#Bf0pqLYKhR!QcGzR*?({8 z_xW>%3+hyJJO)|}X73JH==}XY_)|Z-Qq(XIPh}jX=iiFzr)zS!qKrQ`5R;!CrTVot zbUy*hKQj3?C9(c%eDuy}zBj)QozG|YKQY{@WR;G?PY#95;;@T!zkM<%AC$5cLm(w4 zg8Cmu_v4r2`Y$M7i?l;ux3g4V?rG8P=>1cFYnu1FJ8;{!fcytpm%m``3&NJ^Bf-C( zw*T)*>pQ!jgyB{RyB~uSBSXmlDMj`9Z2!)10X*F+;mn1*sJ`_I)n`BD^aJ`%fK-Sr zt=->}kMaj>ePDbme0XET1GQ;?+4CgD&>nDn8#Gxr0W(XTrS)aFCG=1eu73(#&0Pws z#Mx+nwf}#-Uv&TZ$nR3Y-6yC%Z;I~z0oJ~N51s15hI8l1zs-U7rm*-P!wuMa&}pe- zfSvs33cC9(7;eN5UxJ|I@O9L`g_!=qozvfhp%??Vd(r(!iDz{AF~$e1xMw1G()sq# zLnC$h4Tk$+e4pmxnZMUkZFlH_|NDIW@kjI3(#*>iia)5j`{@`z06j}Tk(!KgAp7p8 zyPrVAH9HufPNy>^176-F|D=n#0^KF|`~A^ye^s*|9S)SwwjW1tvh%Bq@295RZEr4{ zx{_+sw;Wut_AAE+wa3XAaqr<8TH6P1%)iUwCUw>Dt{~$^nqPHw^(RceQ5DNQ0&^1W zUtLHS4reuv?LkbXQF?ip>q+vWf(tm@3h66X0?rGi^=~(}EW_fP zOuh}uMd!d0bu{}hTfXyZKF3!e{rVzszr2j*-)i)}RDr|o(EDvZoV$b0e+rSZWoZnD z3wS$60?hiR-B11u`>d8Z+<+ll3Sd1qx;bYr&X}0M;zx|{hsirjLyMZBIIX z45#zIttLo~t7!M{57XUW%kTg^b-9J~eTQ~$ic#8si=Cg;aDjdT)FV0jO5Hq%(*CT8 zwBaqwZ!p|ny*Xm3WcKPv^M8Z%WDl#~X1Jf47*++Y95_Pu`6fh+%AxC@CLh#w`HI7c z<5%hV?SS{=Sba9*Q+%>!38=I87|rip0FzmKfZX>^Aq{{z;|5xf;#7#4Ql1tK>piF41Um#!>N9tWIoK1Nc9U-_`Ywc5hv7!YhwAeuFtBWPd`6tJ>D&7slamfY29Q8TD z5>$N!#|NcvR6Y2(CYJ1fPK@61oa>*Ul<9I6+@{W@`R6Tns@a#9pJwXpiW{b{qx(S; zb@$t`@uR5JIT7A zH@*GU##e!zLmt8TYrzzsY%M!XX7LHuKRf&k8wHZOocxpC_-w5+CtslZsRLkO@(P-N zeX-oM0vv8Y?;>IF?aD&3|1o%I4vP;k`9`cY+YLW=ilFnuDe}FUEPl#xQ?~!y=S=bJ z^{;M4R-eUi`u>vCrqYSJb1A;)CjH&-JokM4u$xCQsP+9mwZ9!4EAiv&kITRjip1q1 z|9(C^tIg{lfPUBCN|{Lp%Fhje_B)(({R0|5e|1ykOex{ud$NCB;o3v}`l{i6>eT6P zq_!LH(evpGgTCqQgNB2eSM34QroN%)(-=RGYQxDlsp;1;rT*uQX?p)s;|uD(?g3ah=^Vws`&w50lrUe+WT$YXe z(-ZPv>IqIitMab$ezaYdX~ccv!V%4 zl^3Mw*MF^jYg)f;eW80{?Rrl!foZ5 z)JCNHj{-WuhO2wH@ulw{dU6yll?$W#>5_8br#_r~1-h?m1d()qUfH0$xV8w(f3g1A zfxdqS4lNHO{3d#J5jnnq-|97oo+I{A|H|m%UyN_SrK7h?k)f`%e|s;TuX2{t&xkqu z-i9{QsJ=Hq(#3Ze--I74i?C|oZSwE(;hu24|DcUOV4Ygyq_O4>bhderv@3y~Phfn% z?C%dUIhiQ_cT@5RX7!N__s0zLZYgDD35tIRaHj~%KWey@e*D$-eHWUyw9)LEkal^v ze*SAX?T`Bw0HvoZ*_UXz`scGr6VPC#uH2;Ee z;Nw=h@zvxTRnsRoY4?u2l;7NK{(XNX;~P}>kd@N>g#qLr?3Vg220`k>80fVZ$Q_bEPle; z7nD`IJusp9MygM~Cp#PO^77MMy%XTxlC9+5JV&R%p&TvH< zi2TT>Er(P6PULQAdulb=CqZt0M?e3x_H8gxafNvECGwBA!J>*Eb@M~R71(uY7HoVF zL;lBXyjHgoC!gY1-UggAFPiF;V&rDswsW|E`zjX2@_RN={&thx;*`SSls|YGCGDzE znC8cR>G${Bu>KqIpuY>;99e+$+6vF!>gS)Pp9!`6_^RhUi2qgk9iL}>z!pDKrGS1L zv+u_mW@GVNhWp{C66t2IPiBeE?wdEyWBGH2`)B9x9Sm)<^YgX-;xb69 zRjJpfF8s($C;uS_#xK?HUo;$)J?VMy>*BQ}KTJ+nEuu@^Ri^P)F+PYS5DdW;BcNVbidFD^NiX+^@DjW;afj&xK(mobORZP2$@6f7@A=NH}g>Fz&beENRA{MYEa8FEP!e>DGSr1#HkG!}Na<~~bW)utcC zR%c10d$9Z07+=6o$&F#%vTd}#$%lXLzcju94;38%Bad#N`QH+^PUhcFBf57)xcBoJ zJ)a_YWsiP;s_{*De0o(Fo9saI;}Og)SA**x;M!X&rJO(K5WZU~bXo5|Yka!D>b{3$ zgS%S(+We<~sA#y_WnL}maD^wt|6smbOCP_~a8Qq@9+Mh`S!jNgf)3&O__Kza)Ez5N zN%pK(G(P3P(LrxNG~B3WP5lczUzQ>HDx3?k==9fcgWBQbSt->mKkaX;z@5(Q`(NyY zX}F+nA2vh!dZQh!uhXP`i}n7yhTD}~b7sJsTRUlenT{VX)Z*kTN|A^hFx2@4jsI*Y z`k9^YVDfE>_Wcs4Fb|zisHV#=Fx;v%Kd>LhPt)qp@=|@fKE9;2Pv--N(@doT}a|^Yuei$ufeJ2W{Pk0kslk^ayZqWSD6Kq7j38g zRUnF$`2Ja?yLc9$?=HGuqq4lE;}V_^j@3s&lO#GH8QBZx-m1^x05u1efGamI)B3+# zcfav}{fCpQi{t#iR?+>FJ7q&^GRL>U#>U4Wejeo~3zd=siVxs$1-@V@aD9J<^iPG^ zXW996Cf^RrvP#3agU>0xrhxk^y?>*PuYlp**F!X&FYxxhhL1D#@i`5r_^(@a*mm+5 zjo%CSU4O^q8?k(m{IDvx6zz{9p?fX8{nhxCAHVlO>Ndhe@7D-;JJy5KAMoPYve2Sw zL0bQgf%m$a9PWo|nUUs(``1ZyHsfIW8U6g$bP$BM?kQ4w@<`er<^nedeSA&B$v$R8 zz~TYB$$sp@&egi;{dvxTdi8z*&kvxM4{jBiBBiAE56Sw5_A>-=*Be@B9?w z+d=z&*}Q3)w0|hX#n-gI5%8Z+Z6xo@E69JZEZx4(>Mt4JfDo7;_Bwo~^)U)+Y}N0d zwZEbIq;HCO>4*2U{ueTrdzP$Qe>L2MSwUN+En(xSzUYZm?SkGv*Kk11e(g){m2JP; zfA{*;KkZza_6bl2{{GU)QGxoVAaUol%MXaJHNqjqiXo8I-Y6soZ5du z-!IDcmkc+k+W8JQ%0Fu7JO2ItMs59OxS-}wu8qU*lmA=cC)ux0T>tHg^p`s-57KG= z+@kX-0{pxFGrpqa$=w+zR#GVcUyJTfGIF?05l-~S?eoa~B)0e;xK){5_bt>Ow~_V- z{bl<8Vy=Bq-sjpaRcdgZ);~dazYuGm?$;aG5}s|;&i9x5b$&zNKWp0i$xWrTJF8~L zcYfbb{U2W#v$;Iv?zorYlV^4L6^2`-auIXkSsdM860`_KXBIzZxBxTDo`;tuXOaE! zkb^Q={F31S%28+B*>(ZN2UgJiRuNqPtZ>XN9+Q%ui&Ou(O!-50K9%t)HcH>mW~s89 z+Fzy1e=%HvDw`c~Na_OepB7O2%)T?+4p(}tmGbmA()>EA%bzk_z=7}2!Nn#K1$eg>SCDGA1-kpk8Q+NXopMt0^4>Im9!S6QZwxnSchW&C8ei?Z&3@;< z7*6-=^(qE`$KEIT+yB@6L2t*2(vd+^$i5EO<;NJ`AGP}l*GJR+y>vg}@BAdg128kg z(R}%4y7v9)2Y-LRxrS@+iN6~6bdZ$$c@Wu;w$hVj?0hc6{nX26UrQBQd?x+;b@wAN z9Mrx8YeQYHTV$VpK(orMzM0`Bb!?fP=7<~xvg?!6+a393pBQdb2iqxs8NQ$7JK!ZD zw{HA3{S0bi@F95dESC0p&R0wTvlW6|gWUUta_x#q{r~6OrO)!6MT3_Fnuw=wL=;pVElmAnz zIZo`BO7;I`EMguOKVb6d`ENI&rv_?!LJxS?<4}^{+9w$)~vS z7vTN-%CKjXcK_>jp#F0>z>+6>;J~y%vR@AJs0l1S&h)o}YrZ1*cfT;I|64E9{SjRI zHt3jf0(N@@k$o*I&-$vj587BLFlp6R*jz)qpX(`_Y7XJ}cG#Hn11xF0l+I5MlM7@m z;BY!Wdp`$eOxEfPX3OGk=3g@X45*zyfBbU~_1^*M`Eh(Br))G-^}8t3>Q>$tG6IKCzAi%U#9QE;p9{O{hQ0s;p;l` z?`q0_rOxGWMH%Ej4`xqKBK|UTm|^B{o3eOqUR)kP`L$2W^Mf_s}M4JUw&E>weNq@ zlW$K=G2!Mmst=Yu?o0g(|{rtJjcZ#zYzTG?#vBYr!16K6YSkpFx|XQp*A z?l09Bz8?}QGVR?M-1%7kklEL`eh9bk#>+b|;oDJ^AEQTMXT);04=^tA?SebxpS_Ll zFX3zyte-D6e?EL4r-t}#7rigLXA9|ep7rxgZwU2%x%S*+ye+L%pxxr^SRdN|OA&6$_PN`x4{m&CnylEYKlb)mICy)x3!I z{i^*S@7{I-JHNR=@0(P@b2Y$*2YM-{2yYHI)cMDIb+g!Thc|Z7rr6|O6s=!r;n!bq z|No1BTT*l<@t;m2?hbz6G47Kp|IO#X{p4PA-Z^sVGh6#o#?QD5y}|40{`xLnX6tRk z-SW@pZoGp_>3w@jy#JuDulhch+_*dm4))$5d_MO3o9*}9Kgx&-1+aFB{nY-S*u0K? z{tbCM?Q>kxNae?G<{Wc3%+|ll)4f?z(E=8betUtlS;BQT+%RG-G{yZ}R6px+vXECh z)tdA9w8jUEdu_l@p_7Pz{;Bz=^Y{$BJf(>kuyG&7hn5PtHJsyD8AtqFvkFF${hjQq zLaX2W&;D^?pKcX!X^Tj@zY}TpVWxMZx#BC{$xrq1+6{D$AYav(dpK(@((gC*eTFEs)`N@wTt(+Ig>%>+zuMY&2vvU| zZ(LP>X}n0iitQ6TzCbbf+FU3;XBeH&Ef5b=+vtPDxRq`gzqc&CKTblMiyWWH^eB?G zISEI5P<^n@S@G2>{d}wOg^KyTvY4OSZ6f(`VAPKdmo)iA;cZ3d7P+CZ8BY3xE9R&& zQ5x=2^xh3&8Q}PwSM>hdY1Z~6*T$!C^bU`zYc`y@oAlSqqx|ugYJ3V0@}AhX7Ha%C z!16!d-!;%iALxA(@*NH>OiKFwp%|Pgi4Avn3+A2!JwJCR`VDaCi+=yA_6_g*nY-Zk zamqjT2PFTX^B;v%e9O*Xq5t3!B!BdTj3JyKiN`19+kCIFWloAOsnyQ8H6i;SG47VT zinYM~71j`ctRPx9vVXrz!svmRx9WPjUweodjdR$>kMdy03)sEm7Lw2NJ4=0J`xduv z$UIq7iNy82B!8`OX4!hnM(^@0t$q{ogvBTMr-GKB@civ8~x zcR@IgM_+NZ|0Zkp0mkV~{AoM(p!oPzeLo4m{g`o5{f)bV;-k9Xf}BgTkEQWhzkk-3 z{JWCY#+1i4l6;U$vu`uK8_Typ5BEOArG1r8t3S!OM3p~tbV4_cZ-diwai8t`N&gcI zA!(5?rJn)6eF+kdzN=+zIEsZ^zChVoCFp(GANS>8{W-TEDt1ne zYqnpV7%6vk)XABwe^j`e_GhSf$AlKJHrpoJ-{oP|0{#3eoa!fZo&^ndkEZjv3Syow zg$Gf5mjg#7QS~KO;`9PKe^Gd#H(8MsxIK0p8owKJO--eZU*QgK%bL0HUWRWZA6>#* zx%K%`xZ!>N`&aWtom8s6vbn64ef&Q8`u0&XEfi93ck^y`wr}(INx3|JEU}=+O48qs zaXvbr6<-#?O0)&VNebk8fn$5R>PY!U_YD(fKgMl8x1D-|rAf)cAu$igF*#So!i$ZlliT(DS-{u5c7B zOTCAvL3t@YqAMoy+VM~U`^rG?{>SP5h^@t+sPRih?|m~NG45HBne0b%(O+MM@gQ$w z(+$(wOd|R6IQZ)~G9Kv76MF%KZ&dyvdthaJ_HSg|;r-EI5R4A*Pw^MMp<;t)+V`pc z4R7h5k??w7VLE?F@$+Oun;(VyWUb5&-1wa0BZF69aDt3BT*?9!Cgbl*57GO0uIRTz z=TC~B;-{-d!_trpbpMt9r#=()YQ2vX?vh!Htj0ugGubDviri(}+S;e^n8{^fDEs1QnA#*3#MD;yZl6>??{86}t4R?7Sr3%Fl6^_ySGrwq_P3Irh_Zxwwcj4+J zTZq4Z7oQX9?<;G3P-S!ian8Gj{10Y0FHL0s1D;<4*ZiItC(L+A?~h1)){^rxFi!T{ z0b_CJjxTh+!^HYBF8laU@xOQPV_E#KTkDiI+>JFyH^g0)_mO-YWYZt0@ktz?J3kIt z{+!-7cd=?S{e7TtAFB7u7I%>Ke(~EEczlMa-u4=#$gcK(30yN%mk$)ZLtGl%1cx0v zPWN{y@ilio+xRH|NsC;VGTTMEzb@lQXM7tTB_-|{mH*G;1 zuJVICgV|w4DSn|TI)>=`OSMn+x4NdrfIsul_|{{hAhglDz_$l!e{>>0!(w)N zL)8a&%nG3OZ-J{H>-$fQ-vwvOr^8zlzfk+za7-WlezxWhhc3#Ak$*g<^K}smUStIMdh`k1zhy*lCjEX^^gb*9-?QKy^!|wu&mRgKy&@YwP`Z=MxMv zmH%;H&HlxBsF(f-scpht>VN;a9~nn)yyiFHf+HE-f93xJ53=%0PkFSC_`^&sev0V> zz0(2`;k8>gNq(LEAN^B(;kXk}z9jjn-c+%wFZ-V|z2VK%s{vdp+=lM&BAR`Zai6SO zIv?zw<)ZuJ7Rey%Y~zor}pzLxP&5f&PZ&Q}NM`;wA!GX|8u?+ zA8+ezy8RdEU2F{T_fBy2ch>(YdWDC275}_7GJ(ZE*Bxa2665IYvnLc2d^|_%UqG}! z$npu}L0;v5l)`o0@;|Dwp`ErrRQrM6p7%Oq$$6(~e~_=(k}@{j;T?Z>B~Hp7N&1hA z;%kmyY`EbanD8Oi4j}vC&N|K~ANYO5{rhCDj^X$r_flG)^`c&8mR}f`GIO@I*t*vW zl3#C#{HYq++IP!Li5z$tD8Epe6LQHU+83pVDtVz$De!e~!oJ#y(?mie5*k{>_~W&f+&-+2|#9 z>GvFOCZYNm?V3ANUSRz))BA8%;ZIoi{dCHo(%V@e0qfrwH$3h7v79&X^9IFVDSvMqrJqvs!g!!~+vFe?&f%r_j6}}R#;kv1JV+e5F%_4_Q~IMt z;@b`-zo;MM#2e<6H(hHMEBkxJBTd*plkrfoxxsPxBVSfJKbi4u4g2{b|AJ2m(e=kO zvLDRW@`ovU1J(I_v?LrFd@4@-cLQAL$ngz~hkA>Td2W_U{KJaR^~aAfj$XLG21cEz zPx8(0@TMEf-;4)&Q-2CKD`ja;<4@}ulAZIHF&^ku{%ZvW#JBv{GUZ_ThjE8DUY#_s zuWxLX|G~VLJ-s%+YW#-x=9Gz$*-@A1XF>CO{61iMpS*eW5cZ0NNtLy{P`@A45gs{j9%YdWxSP{}7Xhtt9`Hy3T_bZlia}{K=lc%5=%-er||2`x-VJ zW$15>#iG^}pI@kCNXZ_Y|CZY~;2%?9EhyTsyG)otwuitnN6F|x&DTK^qdek-OA67MVa zz>?jlzHh#!n*SiTPegkKaGn0q02`!zEH7YLq^Qg?=rnl_F;kRtiNIUwj%MK z_aM3_rMtHqW-Aw88y_lq@7J(dkZwNkoZr0@pn)Cd?{SYlx<4jgX^bg|9->j zfo$L4@ww&kD0*YIK1A|kN-^|uGTZoEvhS`WnCatvn&0^t=OO#wGCj(AOACoP*Y5-e zpL**V?zj2rS$z0x8>H09@vjum2-+gK@sF{)NXMD7yY$5DS&KNaK5m(Y;x}#(0pp8dd;Kl-y7J{SR|hlzzXd z_EC%-@F&1m^**kK`%m!up6NqHo79fr;(7N{9*|m*{`ftGyNEv2+xuKCxRQA%$-g;Z zVT}~p_bVK|zb+UA`x7j*^sD~-yG$SCHM6ca9~}&_^7j_36U#*rg2ze~>gbs#!39HR62 zUpPw2dx2}4rlI%yb8^$y*I&^aGUQAKY`H0x(qH4^(>njR+V>0(yN#|RO1=i?(nm4m zTSafoo;?d@hSjC_!)hoMkL}AmeisxwHU~HVv6<#~lPHv_hJF4pGH!gF@cD1`{zkul zBGc3THsLDXe6fh+v*OM;DaP37seUxwkF(pT`rZwkCvWNZuNuD_lZ+gT%>!1@{TL=T zPs(GXmsm2ci35U-?rwqwn) zY=2 zm^(^7G0&mz|I1M^Bm4LanVIS?C7-{S-mllhlJUK4^e)eQEQqV0UZ(r$0o5wU{JD&#-{e_Hd0^|u~`vk&sp{z!z|i|YJc;Sxi> zmVjD44^sbM&9_H%{;qHz=E&y6te;QO{hC_rd2HW5hEVxK#x1`?=j$dG-EZH24zaJ@ zS~!}zIkkTSa?R84Z`FRFsQhw2)XAvoS0%tWCG_`~!h^*9$_Z)xZqojoVV;f8`89sd zk4U^|t$8^?6Y}3K(Q@39i{Zm&iB^#ZKh z5L=sHg*&`ApRRxn+ba|OdGO~y;`XV&WU{01wrM`nZ>QJlH!x254Zc2wC3JtL+0q2( zN9p@dwJ&Amuj%lwmdWY<+ky3;>GGSx-STvombkmVir)(sSGu&djn5@R9bPy(5NLf< z;`a_Z|5NlR<&C1)Dc@VVf8OAvraJ#rxFL&txewR(29SQa7Gg?$e=FSOdHqXmeE01- z$?v(ufkuUG{gc(bTW zt1xh6ZF)Zx$Blz@`CQ==RsP|E-t5YLt@-EwTt7Szkk$NH?v(P6i1NpuF;4kcm*$7< zYbCu8?wDaU*}sx;huCzi8MqJ5ApTw!3QuAGB*p{9AB);R!y?0!eh_MR=K5`n2Z@3& zqu_Fv*c1&(d%B?fZ=T@BVVp8Zz#ZD!=pGCynTQT!%*2bp59qAIf`<>X>1c%5M@N&gL(st-m#X z`F2J__~rI;I-f;VJzKP5PN#w6aRmRue0T{)5Gi* z_uzG<0(5^=M}K~K9>0N~e@DFfG@9Qm{7I?osZAZqP@<4 z74DN=ckG5fXDgCzhHMc(<6l(&3YW6U z{EL47D0-Ki`mITD%8<<}b}TZyu$>;Q{oiMAR%`#;^ebu%vi#dDFzW3$%YIkjS6%+L z+V{MByaqnKYfs~khP|J4`AXr2aee1)$njFpxkbG2R_Ff~Jv3aNAD1SKP5jx1l^byU zA%DMtWOK#HjRC}eCg4GDcH8$OCTe~HLPIJNeHQflmoYu%hxy#x+%<0%3`?^mD(3yE z`gh~sDWl*+MMr=`LTLjeg7%khr?GaggJHE5dS>@{`!L4 zzaf%+n{PgjTTR8!dvYY_`X`J##IJ92m<#t^v;15A{uc_j{>JzabQlIMoI&fK3gh2y zrhUJ{UGOT$6qq|t#gEQ|W!rMva0AmB+01}cnxZN2HQrG>!v zTS&gxLVJMo!!aJ}9rJhyT%6F(vL7oO=>N_qG|`M3uwBhO7M9wJ_z0 z{iMIjspZ#X`ao~pBoEuB~{#vR_I!|w;jJ#g2uKl}5yGH#Ilc@P#}xs=}L zeX0I1%b$$Hb=770nAJG!1-B_G8YE%Q)5VCjX4MM}ic8Q#u%1`)-Ur zpGidRyF&NpK`s85=_U3%C`9~Cl)qqcD$2i`*fu{tEcs~+R?6k+p`!7h4O$JIi zQYrc9KX4RFDi4B|FFRQJsehl3P@(iM@uO~#f7zb@=)bK{adBtQDETR;R)0>7PwkUX zZ#$|l^*B#QtG-mR5o{k~9KCz)2H>dES4qEc7?(%$`;74*ud;7?2F9ZCt-y_&`TfOs zpm(vm2qwR(;!j^{+f0_Fm8A=_mZ$AO=^-a99SdEZEgM4 z_Wy{lY$5rn zmB^bshYfehH90O|zams0ylj4Fp|ewMILh=-1H|w7E|Pxru(NJFu3!0c{4#BmapK&a zAHmlzPYjvs9jom>)xXQrx8H88w_W8Ir2I>ZS^xW!-l#pRE2ikUipJMK{MCf-2gY6S zxb|mUR$vnyjlRyU-`Cj2XJG4cRq^=!3-tcUEd2g~Opmzf_r{_^qvw=Ar({U+rd90Y z!>oJXV(a%iN&nE8NH*SvOWfRg6Gr|%i|8MVL0v1^ za36Xjwxh^O_NC4r#f`~rY`7t=UM_)WTAZcx_m@_GmB&x=;o!*_FX0xl-_8{mE3tna zEQ|+vF{1A%Xpt%h7t{>V8(p;QG_m2t52XX!u#?gDQJer4AEvNA< zh^p-6_$tPOyhD0-F<0kY4AuTP5H%$)zmFLY^!~cO31;kgocMDok^BqGZ;U&7Zz`yCC`zE6$cCslsN$jVhnK5vE(*O#{8 zQmXuq75c>@`Kuk}ueN``TTa}%7^BDSr1kNLw37XMxqp{jTD=ACnC+$gTS55qXEBaa zwfc_%|+5ZXj7!;dRumqZyo%!6lR0mE9-z&i8lPM4 zNz?`lzKA6G$|=qcX8RD+yQK2(edSg5v29v?D8|V?iY?&ej45-!qL;h^MvigjJt7B&g5cV(uKrdmN+jsHro1^`0Z0}JW=)vo!{%4euC+JcrjTE z9J_?-8}_Xr2EEH~qc_B$X7#Yti8DlBQ26U>Fug;JDH|LA8kUyEHx8es&S0Yt6e(uJ z#dD97)Bc@=XaDB+InxJ;Umo8uhrj-!^mnGeK0D(m(q$+LTZ0#q{Ppr5oaf(#>U=aw zUk^r_7ifRygIMVbX#Fdk^7rJt4#x_r`eG$<-Y3@2sP+{eNcjhL!k9QUlztCxcGd46 zg*&`y+I@uBxdP~Y-V2LWWd9weH@p`qZ>cdYKFMz`^uNCu_sI+&g0LCc@2{_XkDU{- z{(^BSi`=dPflcR={w||d-;{B;Y*!^azDe_x=KnNK7^Kgi8lOw9DtaD7>Oj)(<;9lM z*nZ3OD7z<0fw{}Qr~R=E9}QsrIOB#4Y*ZLaM%*X)_bC=TW*;BF;axeTcWWO~vOsRDQ#SqEdoVw)J&ksZrfALuM8K*;wrRGsuR!F;&UZxVOs{ zk{@D;*M;osC-Kao!WHihxI%BcfrmB}W&IkD-=}htK=->rmi(P9+`fMdq3YWO;1A*- zbF}=@Oz#kBp6`b5%StHzjeh+#x?ie*zdnSbcM~4u-E?j)w9VL?T86lAYcmiBKp>@`DwUnzR3KY#C~+3isu@(ph3@#k-4dWugu zUkxvuyh{3s04;wl<0vVn1AUpU(fcB|_>z$IdyE@W)i;cirD=ZiV|G{b z0i4%|_4|xV%#(CKr18`u{cSuOf40_-Vj~k^lPdA)exIS~Cz<|#^o9}=NJzZ1tPOSfp%Bb{xJRd-5C!OrRn{?t72>G{k|l}Np1g0 zTAx7g&`zD?bvGld$`FtAr@Luv`ujrRZaH93Mois< z^gqTvtTlu6qujqs7TZ}1Z#8~I{3)l%{f+HwjH7(}>;)vumX+qG75@63_2-NmvRI0e zaO&f4>i#lIPtxaC_3!co?ih@J%Ht&8_7l;`oVM{BMw4VN*c`op_`^u>+n>387u4^w z3XFoC=zVgDbjbSsquMvH#-<*)W2>s~)KnNx^!Jm(5&s%@9q>{~TK`fwro1j6SmVPv zvkGIFPx+5LL(1i68=o81{gxzqdFy_Qsb5j_5^HDN0vSIyrt|R}e8ueZ4`sf{^xwqsIB|S8o$`$-V8OW*P{3BGl;oA748trt~P_eGe#)+1pNLdJU@X#rdkL$ z)7B&TV;?-<#Q6gh?o;&Mz_@3EN1o1tWS@zODZeS)@ERN61v|_H^tLP!g$MNgrEs4N zC{J>2znwJx`0#3?u76WF#fQJl3gafJ_=m%0zCAksQn*{Top8pyeWn0V{u(pp{jYGB z{QSFwr-7B}eoBd9E;~KS-S_iQ{H3aIblog@Q$Jsd-jIU|eT8L@3X=Rc0F(XD-{)5Q zo`WwsVxlLnNWQ2g8c%Mlt*=FIRN3;V!tkMibiQ*QI$A;J{}w%Lt(P1dMdYUYWvLdQ z#`A07##gJM?47POKi9x-A7>o#beCJuu|)x8-$#FaNyc4R>v?0SpJOfQ=S#rHQ2qW> z{k!p`yCAeIev;mI_sp2`ox&w{YOw&eZ5>SWvlsmSb=>|VdaqQ}{3f0RYP{|MuT zh}wS_&J4^-`!B@C|5f`G-*+$=WatjUlK|CU*XGZlJx1@dx)tEa#glYC62R6C`uSA2 zPj)KN9!3Uiw&Vl9{aE#{a49!WnhUa9U*dmT!0*4vxLdlYCUvTm%D!|1c2lgXU%pfH zF8TUi11J!^ndpna`+NHPQsF4`BH6|hJ4iR%E;#ud;w3-jLxmf1SMsi~wB#O=pF_bP zAIkkx{f89`Akq5a^nH7w#&TVLv*?WyBa(pYct(<6=9nqw>ibXOF8Exi0X%%Qi}*t! z7%O!BqD7C@b`>z)4xbuV)YYzb95E}rFl)`LMeUFC%w53oZ$HO}Ua~)STJ}Htqd)$U zapGT}E}1J@}(s_ITh`YuuH;-j*t>(BK((A~<9Fkz@72J2^(@JES>WPAT|cAfB{uGu-4xzCD*sl+nCDO7K1{u0 zbH%~e+gbL}h{C#k&e{XUhhuZVmtqg`hmFv5Ndv8ag}Wf_$kdSYWlEA?)|fHp*P_Q> zZFZaa9Su}`PH?74`ux|3;m7VLY`@?VFVH~|xCaf~6VF|L~oFN@K*9T#o{QryX3GOg7waQHz<&W=UJkYy# zwg(cm3ZU_&(CkZ$JG=)!9)`B=Ote4K{0DA$BgTD!Quja7`?|jtKg;w!sq!aW#7kEG z1p1(V=9h6P|IV5T+Fn~t^1(jZqny8zaks1+Z#m{WtLg`zB3o508@)?zpO;*GF1eAm z-%_%lM%i$bH+Fu(ED2Tq*8WsKa+eJ^B*_F=JHA2k;dsqH#{IiI4~})hT-jB8X&ueJ z$v8=+pS}jCIlTu|ejSyg>VN!!aTmbp3iznnd;0!FTKp;F2DUj=KpZ@LgzUffsD5N@ zj7g7}dC~zab#ocLe-coA%dWQZyKw%;5BU7|ouvQk={$R8lMQ#H@(+}W%hCFtS2lRt z_#~!oHyNuxTR`&dT%mlhV)pODPk*NpInOU3`_($KpKi0^h6uRy7T?UgM*VkoW*dLn zhC9TgPX)1Yw2$<6kIw%N zh2UbZJ#;@mR6aDe?+X>DS4F_ttTls|?H2#}`!(xN7)P(t zPjBB{&eBi6%BSyNMIYpSKX)mdIv7gl<2Ug8|1*7{SNV_4n^%PP_H}scX8$+F9o|O= za^dV0mr4IY^?fg}{ep4BOYytVsb*#Be<%3sUo!5KVN))`r{KmkKLxQxE4E)UE~Sbe z`E;1dWT`}3)2K9lWpOz)Ci22H_1&sNa-M2O8}`Tff{%Brj0 zVJ&kj>0i4zBTle=k#R%bJ^2u0x7|zeouSTKp{up^RrBZa+)VWjTen?D`hm_)ANhg% z_n*QIBiYbkOxaucKPMGb>%)e-pwy7HnDXNu(x2TFHQV>G;iMmXb`OhC{hjg8YB={? zo@>JqC#{Tx)r(8e`Xxnw{W%^#<;U&O9J;h}k$$5b3`)%YkBqym`pY4BnDlF*;y>o! z@=w08@DTCG)WYw7#Pnp}tkxgX*4asT7m>O%zuy=)#QDQjuve8cbiOi)XA830#_teH zzxLRZ-O{i5{ePK0P#h~33l}VjN9SXV=0C@HkmwQ?&pcWE0r`$pjVe=&?GKEjQ2xht zvxbm-c>$_tVf{Jdp+e=qnib>N{A~eV?YxK1-##%bE891j9_9Ue4RKZ1?Zh98 ziFp(4<1^&_>>Y4&mDqH?tBZ(&UF_|9w%zZ8?l_0&{InHKin9Ke+c%8+_2Z+fL1KD8 zO~oJi`2E2+=`MPez+Hu|kpAx*`rp5d8~AuiT0FGvZyNs=T-eIKK8R}*MB($f^XPmh zb55DT_IajvVMMtRIN@(qzhS6omPY44YJS|9zF1})CSQ~O?-2UsbEcPgZm)sOd&D96 zb2R$jZ;bnJ+xBLdBk55(|K&vTnR#vVV~BJ$Tf-kI!s%W>@aJ!4dWTT{@zWodApV?Q zdp|NBC{%um77xqP`z^6ne}nNL@%o|LeA-|6x7G9<=)?M1#!;+1yb>5A-e=kQo=2HE91I!Y2FK zKFc`OcN#kZp6wn<_*9s3EuJ<%ir(;EE*Xrmj;Q?VNyNaTc6y)uekK_I4829`pHyUT z%KA5MU&<@{qRpmM<4x5U+SOD)f2w`AROdVKALXp`-R3*{UogE(7ECl6Gw+O~^EpOj zs0g<4QT|wxFDD&JLf_wCv;Q)^Ar*g^mh+y)AI?wH`Ge}8-aoMuW3lwF>3q!5>c{@1 zH?FLRi;Him_^L_h_b+AK1@8--gt)7felx2UpTf9-D*kCfSVfW_KR};f?B9>5>YKM| zc7XVIOHsyG$o73MoIETZCa4&jzHf@{{3BOVpPhdu<%bMUmbzD$60^Jc#sHhu?{w6X+Z7Y*8JxgNAdKp3NZJ@4&qM< zpOqkK<}%F6y|U9|D^Y4lc*IXIDa|gq**^Z z3D#ceP4ZJX+@29bexvpcFTJ}Ue0Wv${(ze8W88m)`=rY6GdDp3x?g8wjn!=5;`XJL zA16x*W&g;C3$nBQlX18F?d0!JxYk5E-{tVHM>_vj-{+D}lJ6ef>qGOm6Vey5e;-Qd zSQx*Oew*@FVa)qS(Hrv0+iS2ddts7aN}@l%9QW_?d@X(k0xIR8@#nDZe~aFDdhQMs zZC`=>Lj*2Jr=M?Yd~kVJ91Ie8?nK&U!Tw0&x;(M&hHeBN66PdB> z$YeDCk$A17e!djF4{KE!4-YpCCjLDcQVrDQ6YKlMs4SDAUTA-szd6vZl70U<1nlTi z;Y^}1%YG7*zbpDcQFlctX!LY7-7f}s`snW~g$Ic*_s7G{{i^=Q6u6K>Z{Mxxy+z{A zgtxCI(f(Kt{fg`7Q{fJ8P1yp8NQ#m5}J0Oc1Iruiv=EoSTcSJAuWl(k9maNSg-A76-`V}aWIE8HbN zd~c0GBb5K9Q`AU-HXNnnQaK!X<|Xal_*#6{&+*H=BRb|fB<1Wwr+Qsl> ztdz9AA8_}rYPR;N>R&N8oY^*t_SbR3{xv>WQ7bkcf0L2!-*LEqw=Ta~^M@_B2V<^n zPe?zIRJ_So*w#KG9PW;XL$8qivx#`N!)e1^R{Y|{rwwR)iERAUnm<(Wm4kDYAo;Ki zP8q86XN60$*=&HJ!OfI@9Q^UY{QVSvpSy$k>`A!g-{+5iV%)IeV_eCcw0{z5_C>}i z{(kc!xLdBNk`Ewy#J~2RNbehOUP#x2+W#6wwZgRbv!wAkyjy2>sTdS6o_wEI9?Er* z^E;^a6>fMFpKk_99);8SEUVdX7^nR6?+q-zCol1r(P(^P{UGB~K3P*83M`AH`yl{g z%9pBrx4hYXHAJRqL-*TN7?xY-zX~V$f#T0r<0?SInp11j8S?Vsvx zMnJoFlZbu@_$DUQzF*-!nYZLXNVIu3@xQLHI8G58E+yq;gcGBR()>0;Z=inu6}?-2 zezh02^leJ|#rp*k51vC%v(>S2#p9lzuhkOzzD{tgf&TtdIQc%bs|hPRE4^tT)ab732NaI@c*hV3uC$uYUqA5IkKq1Y zxcy5``2ER!r5`l?@o9{^G0ybq;0+r`^RraT56-wm`O6G*(BuB3|BUVN$3HRd!+z;E znorM#s~7{%$XWXRZH-@qMTNkxj$L$rrP1>LDtbfl7>P#g;CZcMsoZ>#G37&rlYdu0 zI4s&6M*L?T`2ACv-X+sK`V~_DHjnsUPZ-@b#`~A(QC|G574-hTk>nRsv%fICAqOwo zZ$22^Q0aF)FH-5}L-p_S6!@i$nIio<(k*U_^5^ILNpD=cy4mbIw-eb9@^}tE(D$E3 z57k@F3yzpDm~^ApBL{r8kB{CL#m2(uIwVX&3Ht|d z|1M0k?u=P;coEBfy@+gmG-mq8xEp&tn*?E*W?AyFKff^J60Z#kgVUKO()>(^sH8go zaH#K>#qS5h_+Nh`{*naT2lVkP9A(!wRiMToN$ZmVe&p20uW&<_-LldA7P}ptue_dY z4|V=+weRUW;#9>h!^;Wc-;pur*P=JJ^-Bk*dw!$+xyAI`4}+*}iynqeTMk1`Pbc~J z6!`NOF-}oM1#Sjs*b8qn8ze{HGC_BY1evf!E&P~`bn+8;}44<@zIyJUh1U2$*J7SeAv z5UK{y|Mnl#qg+{e7%YC*hxYeC%|5}nArH)&2C?de68~QOAM@w(>^m_Q<0sfj@8d8n zf9Fqna}aa>e3&(O6}0OxjQIOS+8_4yF@*cdYq-)O1IY*Nm5ta& zPqo*dt%UK9LM(da0~NFVK#{58;^2DMQN@Nu`t3(N{vc7a)Op|527$swB?JGWn-U3U;LBZ$B{yz~4W#adfYTx2h-i{d%Dj%MsAH_qUmpb(i^Ww*2Fqzt`nYJ#F?z`Z;KvO{hwV~tCM_o66$o}{Ea{9 z@kNWfaI;_v;;&9Df0gY6j3YX}m&OE%Zj=5kws`h~E8=$^^us~Av$!O43qx+h2*cfTK;gx1I6&urJzf&sw-h=`2!da5-Xq7 zhhbxs{80$}@i&a42#@Fsv(62m_ggzym@3BofJ*-r*>F0{zCDZXk0sEicn-B^>8I!o zRQdx|-y!XrytMz4;=3^RA7^))XPmyImIl%<>@t;p)fz z8zOJ2HsIV4LHkc={&9>u#Gu;cpiSznwEyG5yZ)TNlJP)cmYZUh_`55Re@k$~UUoc4 z1T|<1&+^Wv`@NLr|EcH^)&9;MJvTVzqh^-;`zjshr(rzMTdq|m7(ed>$tU|ve|KSs7&`?Vog8G z^eAIdO}K|=d#Ue(YA;#;%eW!$)O!u%2dnsncA9;RahE5-fT!?mMJ_sDjd9X(mY*3n zj1H5#W8aX3#Ge~et*9aP@xhy8x!{jQXVv|0u4=;e*`M@S;m0r-e85HXHwa>`ubMwZ zRsU`NSS6ok)%0IX@4^OgHbQOeK=Q$Tcyx~SUyQqP-tIau@J2$FUjYupYGoTA)n_OD z`G@nWzDYS;-oyTVJ`78g8{Yi-)Ea+>y!Q4DQM3FelhTLM{_NxN$B%IT4pC*&HnVfT z(pG*7zx|x?KrxQ=zXh|;wDiAz`wrt&KPbM(+)}wHkbTh1c~a+ZE}EY}?~{lrus)Bf z-yQ+}`c#Tu;STSfXsX+M%uRS+O~1>y;r$*t8`h5RO!!);lT%+GMembWr?-ca8_v=E zSAm%iN^A3{aN56n%0iBf7wLS&0l$5S+jq-Xi>pJ{U-w)3w>~M@|BrE(OggeYWUS_) z_gz6~P&LN=P3=>B`^s0a??gV@zirWWzY|XOi4}42Y^QY8{v@m&9pio{+~ujUC^;U! zt?VmHF}k(Bero)N(eP|a>|Q56@%J@Y@o~`b(v_rAL_hk`zGTCj=q}< z>g|oT_J`kplyStnt;fT(H?u7JK)#g!Y9C_{J_>cu_NDL3fWD}5w)u6V%1;!(VN=?l ziO{cq;r1myi8US~i%g*NI|*j>=Jz+_#rErSRqxkpw(RFG4!zWmK}QTHd9;!mH*B&>z#gotNvZGZr4ucn2;58KbH;m>;IS@Wry5% z%@SWbE8N`DsJiz3*7)V5c3aF$VV#ujCD^YYV|te-*N7_6qV*ZIKj4o6`uozNH!i&j zhq&9uS@QA!_7}?+*7#t+uKnhLTs6r4(9n}`rM^BEJ+7H@%WOYZ)rKi+W^Ss^r}}aK zh<#I!hHXjL(Ec6`{`z^0yKr-KAh@sJrt@*!oY-DJpQ?Q~E?=G*#$52x`8{Cz^|wqf zaaPz6NV+$i*1sRbd><=%AD;59HltJaB>t1hiQ9d8?xKCRpzJ;T`4{?iznv(Rk+JjmH5N^!vEtBe)|#kZy1{& zUxH3YO}c*yW4(epf3oOd(fI>ro^8UCKaw}m&!^Qtj;=iyYJYIi_%}mL`9|T09R~aY z6Jzh7`?~`)h+z8}_wT|bHAh47@vMEr!{CJXo=E1+)H!yra3v=d+)t3I#Uw@7JM{J#Wp?SgGo%n~tb17I~AFF*# zcqs)8UKO8oqf1S{{1j@9O}$^Hj~`(^tJ~VLpZfK$j0bwF7jFcKht8z?AxQH-X58U* zCn*jcw``;R^TGN@b53WRal=bz0^h7Gp!C0r$Ns~g823p}(oNyug-!i>R?mo59yq>Hyt?zQrzrPQq%768D zy*TuJ6RExt^IxW?_~73bLxW+J=>C!a@%=83@;|S8c(vt!u42e+-)|&r+S5!QcYqaN z;?M8K?Nk15ivKJ(>X;S(>DMnaZeXf*<*@Rc3-rE@|4)4;wW}CM{B6#8*i88ip4ZE# z>E9W5;kM^fu;}{@SvX=jn9ocznq7(p(_7czWeU*OtU=>r{S$A9RM9e=+V5>yK0XOy+%-{p{cC7brGOx(5{| zW~Te6hUQ<-^g-guikoJEVri}Wg(gP(KD7^nyq&8LgQ&axE&1B-AEoFO9_XDi?-a?E zg=zi*wfY&1lm4myYV+u!y2M{7j?>BVC*zb~sAf)-m0}b91uPVl$JW13W+}EBCVlTp z-**N6cmJt=nO;iAU$O8}zKpcL#^IHT{Jvn^Et_o44{0tPr}yz4^RK@Aeq!7u<;wQZ zErY86S5>P|$2iKG8R8>0ibMTR)$)HZZb;?7H)3IOlAoq%{sMga=?pZ!AY7J;{RbI$cv}ztWX>=Ajov>!%=G6u{}$thcXzLDFmk4fA8DZCRJHx9 z+V{y7_q#xr$kjCeq2RCY%=A)Tz1{}ydl%9B!wF@IvHt_(ZmIGI{Q0)MT3<|4q_3@g zm)tpHHtfqdiR7cTfZz1}qxwhL`okP(l&J%~?>5=iU*U#)7k4*gi4LLlH#Gle?%(D4 zvSYM4W2H;o-@$(WUd9b$NoY5ejbD-fZ9~x@dvp8v;F&xOkC#=Z^+|&<&#%=#1`Pu^ zT6-6*Z*uVG_vQ8xKMc46?Sl#se+Z%|GQOV~cj0`t1^fLs8Fz?mdyksiPUI!|K8-nGZY*2>fg&Q{zPT)Yd=-CT`tuhndZhJnc$5$1Gt;Md;`Fd>Td#Rsq)xIH>fBos1t>}JAjF~Pf`v&uOg}Xe;KWb3M%5?uE z)8fy$eZxr9Ap!P!r0UPlz#60U^|R>7Hu4uZhh`=DaR^4t)z633KOT-)4|4~#qy6vG z{Li?3M5X_j>#JtzKm6|>#$DvQ{unOS&PR9yO+Upr<;SM{3UwMbr}Ot*?E%~ONh~;a zM#TZ=Ml1iigBvEXeS+zIm|XTSkMEvst*<{mfN_ey`+YF1>^F<}OJC^QPCp;&`y2vO z&Vyvjx)A@^3TKb$@|jD`zqfdWSa9&G6?bR)=*ou}6G!Ib_$|f}MZkF2J!7$@ z|BJc5)%>|I`upPGvICl1?=!#tj_KXFU{w~_^7@ALzWp~YsQ}U#Nd9^P&EGii=TB#P zAC|rmVFvDRLFYe%$M4_IxFJS1D+c41pP}>f&h+PZSGe0+|2?;BLE~Bb>3+%q{`?k< zOL;4_hk0`C5}@~Yl;3}qakq?&41lk%l-*vMe)~M*F1cyOBzQ1j9Q7X#r<$?-gK?At zx3_?qxgtqFQ64(h)A@rMpCK3bT?cpW52f>U3udIrt(`B0yFASzZb6fL<>`Jgaevo9 z8*Ugi)?9&BgR9Z`3DWd;+`kK^ge8CrLldj_vF0DgIMqI_R1xmKxJdjhA>{4A?^DJR ztMr@#uTRgR`#Ay{rnZmYg~5ILz>vgqY5%qczkfZ`yKz_6q_AS=7e%k>XBe0GZ2Vy} zAX8D&FQhPI?q6&D@bUX#v(gw3sQPUF{OF2a()>v2-P8eG>9^AOtHF=GI{#NV#V-tb zVFq2!Y}tQTOx5oXg_C_?Lw1;dD=pn$$IaHd!|+%4(pCCqwgOCU;5c`!$=)B6h>v)W$BOn`U(C%7o_h! z7}jsuK>V){EL@@Q537I7`n@k~Zm^E{S36CA$MlGw_l<+|yQa|mOoW*Cr!_tt((9ag zWNt;WUxj=8_5-GOM$ugb)`K)%y)5wb$og;pE@(%L+4F?>_2&2=?oDRR0#e5wf|znR?<}wLY5vFyk&bx!`he z{*Psl_+y36$yxu-xPhl@#)4CmtaE9`y#K86TK#`FSw3h>{i|P~cjFYV z`L05Fk{{=4`U%Dz-tZRr(90xK9o& z{Q;sIXCVI2UGu+Uoc!0y6^9Wsk6F4E|NEM8w;cI7G4x*hL)|}02C?<;BL8u~>vcYo z{XK=|U&iz(U+=jDS>ENP@fX(K*NhvoTI4Lqvt}B-Z&oWEnZ13__A%|?W##SG`%c+N zY~OF(b|ggi<|M?QW~v-WHk{%&lMcs1i!V|9HLLRHr$xwmzd^pI2@ZX zoaS}|@gSCOxPKomTIqvXHS(zM|Bv-E#D#si!1v-9$#1GQwP72dLwLvahmg2)==(Yo zk7oT9w@>fGQ9)*saa+iKGC9hhKS1GXp9Ok{PtFJnt3RUe-)&~zr?0=l9p3Lrmcp~P zt%yGz2ETrs=?(9L{e>%b|G(DG1ir>&ZQw7dwe2NZ_g1MMTX361PWGhDJQI>glq6&! zA+>L%M8tAk%gaq^31SyF^&M4eUt&vjNNl`9|fx7}bQ@FL5r;_Ah|UF>$CLpW=Fd z9^!_${QLs=YH+5skKlvpRr7B;tG6#I!W-j0`;x_ci+i}GhcRuR5I3PgR1NNJcgOOZ z!S(nCaf8o^&V$lJlTbew!Irt8gP zh#Qal6J0CP^Sg*+{oI)dQ|@j@{nG)CT%i6Daa(wH=LKx<)l6(VZqA%CI~?nsUfXF@_irl{zJrFV#0>ia5Fs(<;Mxo>$`~K{`Wtf zg=els%kgD?q9OUGh&z;gee3>9O1`#9qVWUbc)r~PKUk6T4ffCPvH8hl-w@{_GSCAu z-~I*V>nV1wTPth(#r)y^!fW1z)LKiG{Wl-Ui=^-Goyp1Se# zVmLfhTj8v{EN83oBqk4|FXOm>i_%vi<_@jNgw9}aOd!X z%p^=d?geo!sR}pw<)BIIe|OGR{r}4CN2brgGhrV*TO1+nL(Be>bbrBDmTYFv6(?i( zi(z9+D#=F-@8a8CDbOY<56c(F@1L!t|B&PP6T{ZC-tD_e|9F`0A5#1c!~0?T$}`K_ zm;I(yK5!Yncy6f&WVd&reRYQE^Mz!1$z39;Xcu$46@l$L&6$6U#t(@Laee3rNWHfj z!}kO||A@Fl{MVaX+5X}vw148A+P_5{&(HRm1sA%{#r8kXvcDK{+@E^pXb`03ye?Ek={XwzRRQpur$3IN)3hNn$e#`^o3iqt_veXaBId>V$U_pE~*01eQ;~la; zE8^$<-s%7iGmc>XZ&i~%2d{p86@?dk*7d<*!efcF-wO5o7ve5HVMI3Tdo>Gn`}`{P zOQk>j`nox6XwEprtH^Kl{Qo`Bgl-+)L)~kFYdzWDGJX@e5R;2iA@`ScQhtCw-;X$s z&!0;G?^YWyf8JH}vkZ^pw~1TWkS2?z?cvnt$5MDMa$J9fB~344`CVimhgj@`gUKCo%mu*n`{@?Vb)|?o&=F`$_jNWc)Hb7kB>T0as$b zQ}X?8v{ipgj{WnvrtIFNyGlRn@p%ewI_FKi&pu6aSL4IlKSVCHZ+>?jP`3dl? zKl!(a{S#Pf~q+DM#d4S!MAH(Kp}0-f>Tq?bG>Sc!y}; z=YQDf=4m*-KTv*rP4O%kSezxtcP#Re?$5bb$|tOFN06L9<<$Kr;wERq_On&a7^UBI|ByKP zpS$&e-|m#EveW%T;zF%oUopEArmvo5eGqXMU;gqcsK2Nu_TNgXOMiIY;k>XqZrzpt z^1<>!`y+HU)VVbg^LJdXL9~>QtOrK%<%wP}bGpF#e~sz+wetRwyF`aM8^Z8tShfC4 z?;lMZ&!@%nc_$8t#`aN%YySyxJm37Ao?#__UZ>7a)c#H4ruYiat?Sfb8HS$)dVU>o zE_NifVZA@SE#*^@ULQc*5X+YxgP#V4qJ83FnV(GDbmlg8vc=EZW!+$U|03dOz>PrP z>!G%?Y%KOIaT5**sZ z##H4d|F(G+OUsy}{7?FPP73efi^EQ^*YX2p-<}=JBl$#J@EP@fQT&|Ge$xkhuN_7GRm!yfByOtnH`|63qa=r#}2S5LwV8rwxVxG8>Up3d5sE2RI? zGX6lEiwUC}gQ(?(`e!54^Am~V{B*ArSoWIizXX9kKY+ODEK0}(-`C#9`Zd(DKLT;X z*s}Wt=DwpQ%D?_ap5K%IoH(A}(X?im*N|pv{kq#vR{1CMZ}7I83c}*|GPOUe?!Qua zTwn2i6L@aYZRMZR<155*{g6)_Oum6@7;x`S@FDw-xP!0jTEea#?5y_p9dw@TcjCBy zF1~=Bsm-J;W_tdF`M{MauPcQ!M+h5I6akp;@f| z!WnYyvBf_mxkIHt-J_UY@61sCWjXcuC~?E6@9khqTI5Lm>^zxf&A$o{;Ztw2+}E2a z`?JyS)WOZjB6pNKoe72^cklc4<$xFO~g)P;k2b)|e#Z5-p4&$vAO zrZaEl1UT;@<@+eiW7YH<#_+mRVRD>2e`=D&eyRu$-U~N|rPUa&^uP9x6E}F>{ij)E zQ-<|#kyG0j#IgT1I0-(uChd10i+x4h>F)6=?slsUIJsTWT`}4HDL;g?VTpVb2oDDt^gzbNqQ~R%o8zL@!7t6~E#riSZxztYc z^N5?yBLn+GaP9|~zetOGB5oL6`{LQw?zd3?;rvd7!f0O=`oMZWNG3BiqkUd$&%k=s^;)WQpJsiZ}zQ+315VHEy_#ko9Is21N zkb3kZ^gldr@y`&)`T4mspxxnID7en@oc{vmNwJM^h~LZ_3LfR_$6`E*}eH%w%bqohqzOYv6sme;)Zc6 zeJPtba{%g3cZ+-{Zo<1hE7+E&(xm*Wh8w)*%k#rFY|Fs8&qoU%Odn324Vl3jUYw=I zuO7^w$s_io|K)?0s(xzyK=*g>deb<}9DfMqV*|_lDan`psQGY${)bsm^ zb1}KaF_y9IHKiZ4|BARF25i~Ko*f6+{w6tZ^r7)X;-<6h_UZ7<{Z(@O3W8eD`Yhsx zG3^ykxEF0;{khE6%=NUCuSy?`7?sZk?S3>GdDP zaed-K^s7E|4BLM#%ltIrCZC?PfdwregynZ8Oxrh-OZkPK7Na!;7JMW7-)vY0>8G6X z`e{1TR4P(TDZER~)Po@90ej|?lx6;iJ&}xkIzepQH+^@+&bWv1{a$M^h0)Y^S)c^d^{I*nIs1{OOGky?Cs@og%xc_yU5w{rv)D zSpCXB&^yo`B0mf9@%F*X`pWKG@tO5T32{9V<$L+vs=wCnJ?`s-lx{JxQTXjZ`4f*X z>^iF|lD6?SKVO@#&ED1)-x*&$`oj3S*?M`Vc&0x3dyh+1d|oM@z3?~n72ghij~nnI z-z$r+*DEQWUcE3v`;$hf;=|;5di^mOE-OFGto-?-(uY~~QdxfT66_I~5NDG$(>K^R z#1DT0s7CmB+x>j)K|VeqAwj;n={)2)(8tN)UAx(0!(-|Flm5l#o764V7X9dW2$wZ6 z+4N*v!pk148e+vKKIQSfUWOm`|HI>}atDg9LdH}dV_AG&_Q1-i!Sq%6gyM^mE%)J< zY~b>esqqWIfBxm|C>W)1NRWS!9R-p4LV11@Vmd!6%I+#Q6<8L#s_8-K)ye#P8ZGKOvvkh!`qt+aICd z;(Nq(iR`3AzRbQYWAMP?=^3L3_8&Pmz5l?$V}}h-uQtUG`><1wl!U}^s(w!%Uqm>r z?TGy2(icGrQ)NK4O06*UtH)NI5k13`BPqK|W2#1gHR53<^ytqcVi>Ur$uEq`6 zGOfO7oN$lA!SPv>KS)$;v{~Z@yk6ojE__)?Tpo3P=X2?~V&YBLkFk2@2DCnsRlI+6 zOXC1`w@Jf_r8@XUqa-n z#@8i%BsXA4yzHV@oB5g^Aorf~BH!UCgh0!fnUr-j(ViaQ%Lb8?d%v8BwLq##rt??7dHB$fs}hMg6{~HEzK5 z^sPkeV%Ieeu;y7>aX8_0jhj&NUOJKT!VHCjlb`O_Yl+=(Oyl$%tCSS4Hm}n-z=*~* zMVeK7T{54J4-`EaB+k{~`z*QT_;>DELVUWda|3o=DkPHLSf%(&JwUROWyGj9YPufB zr_;moA|MgpXX!7w0arKpiPD+o#QFend6*%|ulS1!&paA8pwqTY!r#UBjr9N@yHygE zwm;{ci84iCOg;-9@%)+=33{dy7JNNa5gISK#_4;%L#6 z8aJTLvVo#!ue%xtIC{8|n7n77#!V<#zp8lCBwFE?Mv`E!Y+4Hqar9>>33P!aLu5Z4RoBe?<17ok|zafhY{X!pIW z_%@a6OBvG3!wm7sSV2ScwABbz;4wCWLoD>))dQ+*%~*XUGrjMeOzt_9DmXSWPH*^%zG1~=}nk%r<$nw zH{XBOTTtppR;m}X6}{=CJG~G+ub14eO12H;~dFt$3JIO18OG^Dn63ab$xFyvbrv~iCB!#tt>P~^mD_2>=sxp(D?xg1{p$T~zqoAt6{}}%z_@$^M8l@1H4flYu7D_! zlj923n=q;12oV&A>nU>!T5s(vzI@{TRC1u}GDA{qDKCZ{;&zz10Zr?br0d$Cq$~B* z4`!((9!%2pCOoPfB4&M7uS@DJxbWl;ai!*B#m92e?b)lem@<~*J(;fL24t#RQoLTG z`bZA2bZDR`aE9Yd)|=2Erm%S2dzs?zas0d78$j)hdfy~BV6l;p`Y%h<17xolDo!6! z?~jk8@jVnI>~3m4C8vI&SySGMvc{+)M(cBI@dg2D8T9jSf zUCHe^bAVhEtB7R-xxO$rVS0sTqIhwRADCP4Z?R&+ny_5Sr|G0yKUX6WHR*`P4G8@1 z5@(jF{!&lRF`%MIH(BQMo*g?b(7D0}=;!7295@ z*DLEaK+w{>;&^v1R|P4b-gP$ofy8D?;CR&cZG}G`E^d? zg=f`8*is$`vEGCZyPAvs;apCbTadz2SHx+-&r5R4NjFRRDC$q|DfyS&fbNM#i|th| zXdK}0?9)W%NG{L6ec*kr_M&M~9zQa-;Ov_<;#HB?ijT)hcW92rqEm5xUQ#c)0piyX zv11pHN0j>trL;4Wd)tTVjmL?MtFJ2QdfSHuUx4B|j;~k`Fb7ME z388CZ_4&z19%jhBQErjsn<2dAhPeSz&B}_KBj0IyfKGKI#lU|!zGb}$pBwfQ`D-U4 zA9*mhU~lIcBGHt@n!gQ8H@tJTpU?f2qX(IXjpbU3Zb?q4=@K_!Pl^#D!>v;q2dJ2F zoCvwW<>9vvyxJTn_6+5E$J~OgSx1PX^SPWew_*3MROq{Dg5qyF`Dyull6d@p<9w-? z+<-ebMu=yP`8hHNxSFYnm=MG533C&6+}J8g822@Q3r0ug5#7VMos@dpN%wJ%8Z;i~ zc9^*Vhi?cmwH>z)%mHdVsU;Gu;QGRx#%nF6Q9Q5SPw8*L;IR4PZB=d$rCzotW=Qv^ zr9_6(9B(i;pjoOGqCp_X$;<)HeJ>``KH&b5xe2L~xkPjX$Dfk>JN_pdRT6uCsQHxK zfb7Te(YV!$y+650?`_8hKOQgQ25|jlZouUW^F-CcTz{DZ+^g3@gzo0&%iILx{dn;q zkn1CJi|SX97+jh2&)kN%|I8PMeyVzKo=d!{6U2=Y9Op>Ap>~Tugm1_K<^EVsepdDZv12tqN2!;b z;)K10#Dyp7bx96z>p%f<{58jKtT&;1g3Y4*d5#m9TTp!4DKRI9t)%O5(#<_!qR0__ zTjLZ**%r0WI;Z>asE>Fvl-pJ5@A$ylgxy8`k6fOaTafr?4^cHMx95_3+r3kz+~RR- z_4#_+z0Iqdi4|$od`b@Rbl)GsFnOL$`UE<89@jBU)X2+mxIb}um?1fO)e_fE#S`B0 z$DGDHNrwdkMFbj-I#y{p*@y)ZbjMnFCaxxJ-yF|z=-1I!JWGIoy0`Ra<2 zf2pVOY5Vdb#e0tPS#QFjpKdW^R)B_yv9xN*N{y#oY(6=7NZ(N?GkH7P`KMGV7{YLRPow)%qoyv(-wNxLe2k?97 zN8{X?vFT*t>L(oxdFrOFAxQ<@^fJh(6_-VkvRpI6XqtQ5AqWgPt4H# zE$|#nD}q{1*0>E59_$hko7X5DoafaoNjuS}A-9LpUvdN94y;H0vYM{s03T*_5h)LH zf696jhTv||^a?*m<`yIuQ^f4Y+-@+p;osJcseM+t<>Y5;*4E-nKaOvuUUCC2#XTtU zr{Z{nIl%dob;O(_++H!Kd7cETh4GflEprP}ZhS5-UgNltxeW``wGsQ)sCti+A5meF z$kO$ua^EC3U}a!G5xR!s4e28}K)2oPME1Vik1#jEGdxgyUb{l`u^@DJ1reN?<07e- z?W7s4FS!BhTbaVXt(G6j0p50wp!pD;n=q|aeNn2sdcA>;=lmBFX+HU~ z;xD-YOIlPA+Z%H`kdg8u4>KgV#scwfO+Mi*$IJ~HKh6;sN2k&_K-2J2G){aN>(6?c z9~e^tJDZa$O zPUE$o0kn?5^8?HQ9$!p^b-VW0eIVDqdDI_qeP_J|52{xZQ~%;Pin$F_N7WWFWx0My zZaVon*wQ8XmgMru+<^L9CyNvwF3-#X@?4uL)>h(vin$4k;*J)bmUH~c+=8lE_lpe& zISyfNL;COC#i5k^+$Fc2{8)d^70bSId1h`vkB@`I%yr!UFb6n2YcQ2_o(E%Y!ZW`L zV$<(9jkyKuX4VxuN3T=vr%cy|pPz?_YD2g^k$Tw zY;CT;%z?(Ug~hG=mg4XD!0^d_VoU333J2%>R?4Cjr*WL%OX+&se>-Zd$lAmct7mS& zk&?Y=e8<PxBS^Z-@aKO@&)de%cK|aiRd%Yw07o0cjsJ7ci6K9OeKwFV3ZPI$QUF zDy@QE01r zpQRokL)kIZ{&W3hy$Pwu3%c-~Kp9Q{en56FBr%0E@~xxwQjM^D^< z{%zSwD_p_OKw1#1na5&unMFws84mZo|7ZZ$$0QJ(P3}=RD-o9qM1WK1#jhG!GCo zO{80_=0|dX(*H~mu5ccwu-=3NE1J+c*MG_t=`HZvHcTYT&GlF6{hf5nhEe;Io8xEZ zH10Mx(mV{epUi2Vf0HTpcIEoT+=MRKW{99(JpPp&9RE36>x=Jaxcy{qz=&Ia(E2#f zCnP7mw|@!CI!A<7eeV3NVK6tK-s4T8O%C>FPV3@*3QUvdLFTQg{UD#!2g zU=Gmy^>~`6UV_y6&MTf=G!&m!&z zWxCSef-zll(0nzwKg?~&aI2o!*@xo|$$gx(-s~MGru4f{=kj1~z>kdmMBkNXmHbFO z?FYz~TMX)~=Ew1Yc)PoZr&)JvdJCHW>L-#UwG|%doF^Y8M9!34k7T-%8&IhCcyacw zELZYi4lpk4tf;z}?|)*_%fk$rQ*5l5Rk@(>)*t2ueEjLcYLlyL93bD8y5h!mE)T3X zsb8o~`D9M@WqxMbSHR_*xeb4}G)3yl?=` {bd#_2{1}2bnk8yCl*N`A@pU z4Y*&s4%WI9q3Hpt#Hoz0Qxwp+3BwBK6H&j{*`555j|F8*)WQiaU2nsolr=>_e@pZ6 z!0VbhXkBTV!c8ZiFJ_MuLwql4+<-GV#^8kx+;7TsrC!#T{i69D9v3h-A?rW6am%4l z&BuaQscYl>Y-u!ZLyZ9)(C=X$jeFq6)h*&kJ1)=C-*)mj@6`k`wh6~A%xN9!dIRzA zRvz~;2RL$Vt%%CPaTIeC>Sih}E-mMMKFle8%hF$rT+MMHb9(MmhKVO3+@4ErIC-8q zFBy6cE>!X_xdCwlW{Dix)%z(qK)>RNuxmh1O;6>p*+B8VI``M?ZvnU$()u0uU(9W& zlj5|<+itAlLoWIU(xuO;ns3Z4^Te1A&ql6&SAX?#nT#ir&D)D z56*eA+-*huN5?g8K*dI##nlb6J(P#^k$QlYEj!Ts5BtX_?(OePsz z;cbq)BnOarz0~H0s2uTD@sXVNb#xpoTD9M&aatD`u~BrXxlH47U&kxj-^l$7`&;n* zcy{CeNHniB=L&VL*^@aUCFn&coQM1o1MelL) zIrVTGG^Y2}IE{mz1mL-#)*1&`v7(pQ6UqIZ^mlyV>fmD{ykMlJw`jelg(%j5<4@Mx z;CJ*6)gzuSlHAAn+rcNoMDDp9UokhJQdmNqTz0DBFZBSS1tyDQ!?~Sdy$RR*?-V=Z z^LUWC1wFTBr2AY@@dH$8_jTD`Yt(xA1Di5lO zOrVyNI*z7Hqe!v0$>~_VK|b;@Lk1PdDvtEKtZ@T;$|S+O?YN$@p8DT>6Gi%c#f7)t zGB=@L+IZOU=wi)>#^bBMiCvebY21be?+S~_WiM&mL;EiN7L!l&JecDTGOhE~4p98h z_kpR;w|Lvqu}XW0D{!Y@9@1I#H734SH&rC0BR)Kfm|r^2m;T}pl&AL#NU zHQtMRQsWk!yE6{YjE<*q8>+N^L+jcc4|BRS4ow}5i;8o3liYUl*(OJzXpL(AB{v|X z{S*uvR!;Gk9H7yjk77kGu8+Tcp#BP9jJv<2rl)x6R&$Zh;(E$@8*;pRCGO_tc8|FS znrB%h9zOrib9X$~rkhIZaSs)LgU)Hb`t~yslb_oY=_5J7;?$eP!$ut6GB=@T=v=Wa zIgd}7TTrKbPO9%SHGdlx_|1{?xc!XK#Yn2CDZ>>H+Q^t&FA0 zb37x{b$q~`EE8T!!tEz>i|TKB{42s+XX%PSIjzh1Ef9;17>%0{I5CRmSGhf8 zJ*_VfosN;0;%Issu63SB@w1GJCdZ7OCli1(4v8FePZx%lPOwzaog|>!Z zc;*@!x1n3rjpES>_5MqL4|F^lEHKL=MelL)dF^IzoHeea#tpdLI2k^!!*R9Dr_=)! zC?5|YZ9PqILXqZySlK_D#w~cW$QJ*+9-wg>YVMdUtV|qdN`D_mGkM57aeC_`MK3wc zgH@a$=I`6Dae#D_1L?gL^|~A%$l;R`!QQ7Lutu1|1D*39 zoiB?Yzt>G<)NXS-{4R#lq=V{doH@94@wiw$a|3EU48Y3qhH0GkV`NP$PW9({%6b#F`y~~= zo3?9uT7N%Z3BM%dcAND!^sgN#3Xiy;={->P`oAv!%Y418H(aA%N7MdNZVw#}GXK%V zy5ss$+uxaoLS3Oh{gp?bZN8xU4FCymp`XnKHS zc4^!!vTEFf2j%*Uq|c%?PUFf&xiMiH?k{BirN0e@LKlg#Rk=T6?t$*dY6!p4T>d4u zoqRr6;Dfhnja2+4Hz=O@g8TB-(m1V4t}cR;CS=jL3DyFE?z+h|Zb3wm1H$LIS`MYZ z4L{cTQN8o_P2T%2xrg>)_7aJH@6UIhi{m-qpdUVL6ruPVbS~FfVnmHy+^?}7ph(v1 zqS$~@n%;y{Pex-t|M(iWpv0vNqTu848n+?dggG?ts+MPeM>F-W#$v{qql#W~nvZ=^ z4ChX(s&Rml2RmY|9%(gB`??oa!zoKrXq@I@K3=DF;U)_A-mAxTPKmS4c|Ay8ujJJ3 z9LOuUCIYy`JxFmy$2TT zYA>c0<@pfS8?F=2*U&zEj!z^9$Fq@N0RFn#Q1Ox6fLpbSW9`k{&ar z*0dCEIe%NNAmExo>iw6T##7H`P<`j|rc76IfHQG2Vynk_bstz3mPbtUpQ3RK?!#eG zpjo)aZHVwshiO0XydV2}U|rg_G%vkY(R-YH?!K1_b5)z9aVm$$R*0x%++Il^si$eD zWCbuRrK}_T+#Mg7lzfrsevIc~SZ~3P330G>{SunqhE&!2ietaFYTQHZE&UD zBzzEtRUUA;Wo|&8Cui`O#pRqiK-47KS6})$`N-$a+@x{nbj+QB<4@)m^)LN!Ml)_d zncHw>V+A~2zqsb(q5XlwXqal>^mCKrwG?y4MQ9_&}^#l-B#G;ToY3AJ%i_jDSk zeWXDZ@lDJ08aLtT?KpTc5yyv4KFQyL(w%E!!IUqYv|_K9&S^Z@xI6ZXeyMQ}oNUxW z+-f;b^EX@t>+cgjjnw;TI)6K^!%EDPkmG-uuH*(hs%YZ+pQ#id$pPX{nvB`jaQnb| z6KaVLn5k}RO;7WOtIrG1!5SL3;lt(X_@#)UaSxy~y z;KRz`*i&sa4sc#KukpB~RU z$?KJRia&Qw6~pRqKjJ(u=WnZ@jG%ofClnufjRqu-+Z3w}@z*#&_~I%UxFos8O}L+G zFJAAFMB}t?w&-TsmuPC-hGN(1(|dml6z)wc`zk-`e}`zC-fK7>EuIdE);O(4d@71t z9;oHe@qrfAX5#acJdT#vD><#3*qg-YT zb2?ZLQ0H|HTrnd+)0;4_c`mW+Dfc^4-^$S(N!wKHt-3u{f0KA@ywUWl7`vO>Vdk_y zWY!FHKhL7+fyU8u@$mycjhoQu)+Tga<>$ga7QMHg1NZ!qOw-%+-G;`vF712p=7aSf zdXKw^=&^#^CFX|fAlW6mqd8xK^g~m@J~5xbeiLQ z>F@Z!*FNoVQTo4>{76poV1s<9ym7r|y$vru?WgYpaJF7NJZ^uVMdNgzqjF(%2ksZ7kB_5y zW%LmxmUDZ;+<@nK8)D>KuBXfan#7-kUGBwI{2d>9KfSkTzka#KEr?kkjfD=X`ImYd z4vy?1+|So3`atJg*g-fluUZZzr~Si4U6`OmDNRq`e^~!Ol&#oJ8U)7Ziwx7 za=Xo(?tgd({PP+=cjgvE&!~fMg1vRbFqqp=Aj?>+lKzX6hyQU81T<WMkX{gL!BoxlAU{F%ltV->yRH1Fs4i}tsu*Cjc?kdOJW;8lN3Z_;{PH|&;<+a>m= z_BnqOYTI9ur9TICk9=JDs8O^Jl(zxL&neB;~nGqFkJD$0^j>8Ji zauAn?OjmN+FG26o%o-;&JwVoplksQVcp5ipeS9Lmob;UjR`y%$W5Lw8`O!U==QEkx z@YmJtV(oFY+{$!4(6r5DTJK(|_7I$^pV^E!#|M5YyC7m&^pc}+_+u6pN(clT=CfWA*`T2JBr^kj^;T$mw~4upvVZFk1%nHw~IDUVA(bG*$Qz_>UM zi)ZF~%G`vo^p)MGk2KYMEV^%rXq_UL#_2t*9glF~HSX`&#{**mvf>|=xO_4sZ9fXff8=;T`dE(VKL^KPmo9#a zUUHh3`*SO{+>wB|JXlZj-pN~vME808#oVO*A@y;_A#SgjTM!fzEha`c(fnzD_{k08 z6>xmWdfIn%BoN29<$ja7;ks`u5i^=iRD3)j@Bh3di}2Q}>>8){R=QQj(Lr2)rH|Cp zyvoZ`cqE&=F4-&|tmvTGCobL16 z3AlOaH|L&suS@#SJixh7Ty#gRFM-6(kSQfs2>(zXzepd+4R}|%BtGBO+0l4?B&YrA zzXHVjwrc%#d|>Z{Ph#+i3Yy-6O8w7>#~)OmI*#T?Mi;G1u2=Mu(|cjPW{5M>wrL#b zyB2drc(8i^9Ut148%^KSi;C6TG2}sq8Io%~y`MEUzVJR@<_46COhn&3;C6;N`S<-o z-^b{y`#||kD@4uNJa5K&3!Y{!NpWd}rnkZHWW~~%oO?p$iuE42+in0JoXPD8bHlYM z(Km55pug$^GS59j^I_dT8)%%yNA^BUz1~OTRG#-u!*W5qF5~!1ANZ%jbWHT}s^b}( zA3C>S&fj-%d4BF+SZ_n6k4tGiuC38aG^HUM9pg>D22oowVlN?MCax z{2Zmfe^ue}6#LNlF5xz@ zX=gXZ$8gSTWSEJyvlY;|0ZB{8qjssf#)0OAu8KQ1c%EPSJ3er#)Ii)>?Uj-r$t{T7 zW#gUHNi-iDiY#7+l{;k6xCef8elJG#DXMVqGrF)VC;nKc=3jE!zmcFLmj5%erl)b} zmU$RaLM`Wx57f#TggFAa-In>6+=3;kietg{JYT`whI2LYVa(b;6n}8eAN!bi;mbRX z8?>LH5*`(0G)~|7q5TYfbGtQe(!K~E-1)bxr}B{LN`H!PMGzjSq~^bs^S5X3&KAAy z^Eg-PB{%3hGgU-`=(Uc<>mxbs_rB*6xyxFye0Pj2oAO`%BMEV4xd9rd`OE~1@XjW# z*Q^KFlqd<79_Fj*>3!0Y=di$_&*1fCy#-+@YvJ1lr!`LVnZ;IO=E_es?t$CuOpHJ6 ziN+1rkp`vF=Q;OB(jOenq_h2Tc6&L#kk`Z9p#5z9@$;pB6d$Rled9?=yW?i!`p9~d z;)F~%bfNAq-%~bm_#+;7u-=ACIpfiKAJ-S=9@=MH4DB2|U&@@?=X_bQ_qRYLpO%xi z>>tyM$ow3)NWJ88KY28L&x+d<<}~hzHyKm4;dv_NCVU)rkKS|7r}9h85LO`)&82@+9=Ni zA7{xp@w^JhZ_G_7yS(OB9~Fv2zVJG@^r_=$UgqpBTCO{) z=p{GcS(oOtE~(~Ia@oFJ7H4v(^~>?0{;L3P+R5#z^zY?(?&-dV_Dk$^&J81)xID}d z^Gh~SY1KK6(|hsvSJM0j$APS;_w3@9rTO!mq?ZSC6GogaM)SkPH6IId2GCz*KzyC+(YYVTg9d698WSgTnp0$(R-0x zuO+v=`EPWZ#;5B2k(}N~*t3}C+c{ohJ-webye7Uc%HY3H$=%yt6jZKP`q2KR zO$#vjV>O?W+qBOl6_z=jRMUGPa_%qj_Z0R1JNa>CPBC4)`7ue+8_s#~i}!-{;P{mapcP--nQF4DLZ};E!!zK69I_F-0$qlI6)fUgXaQnb|pmE1o%+oKA zrl#6ODc*L2?|YMkcN{@fuF2Y1pqz>RuSu=8;) zhrfNG(Wy5$vmuW!nA7{SM^|IjwsQO~59TzkKd=X%-O_zL5T1OWcvY9jQ>-^!MTRXF z50VX1@&h2R;r^jVSfjMO=d!#>z2p>^{xu5cuaB>BfLe#U;?wWkZm`~jcTZo6BEQ$i znOji2-Vl7%IYxTY!Q6(KiH~5DcR%63=~6!%H4;^)P| ztJHk@IGTeO_R;?NT8dtB8rN6pj4^Gv{gnQa(|*FAA$YE{q3LNK-Pq>nGx!l*BbPU+ zx8PTaiDJS0DVpAfeBmjuSc5_u_du#Fh3NY&YW@QqO`Km_>H8{c6uso~J<1LkmYwIX zWIiPaD0QwN#hKGIy$LhF7QptAYW}6(f<6Q0i^bD9PGBD!n$<3YRSV}JD?=EuV%6B}vV zfPv+@iO_tjh|9zA0r`9*=Pty0(|I0@Io0>awP+vO5z@$V$lQWqhtA`Ioyj$S8%DO@ zh(YhTJh0wF^YbNe^4G_jkKtNd={h|ZS!USZbTsbjGx5yz_!_4;x#uU^r(ICv0AW2w z)4b3<#ozG(-yX%V-ZpN3WPT*KXnmn6PQHF$)6;j%!l&Ti`!_W1fm35AVB{O_N7%=3 zh0ONFjLVWJ`LUh9eY)W`RyoS?q|{4J-@BgKfyOBu*DYIbXb4Pw-MIzG@YSrq1X z^SGP!7Sx(?7Eca|ucb@xG5MrK_dh!Kz@Vyt+gha5^oA>Lzzpc294A#~v%V z9h2!wPV3|4vf->299J->{i5^Zx?4RSrTaj;+xfBhows!D{CuSk?KenT29K}deXZsEVz5+4hk+_=Pgo*EliRR5FYuuuIzM=J<{0eX7oR{6+nf6DW&^WD| zyg4T>x=oD(Og;QeY&^^Nb6$-1ewZOAHsr;~{1syL%xRr#{S)%x_?9{Kvwx4pH|LMW zK40c0j7%^EV?nk2eE9jxcaHR3dVSK%|y zBS}3to+C>R#5S`HMK8Gl`>*D}V;wjSXFWjT8ny8Ar1+ZNgaSw7x*z{u*JeGHhY6AB zckGCgf9Y@2JVt9Qap<_lJ@nlG6MH}4c9ne$SF2Uk@lX0_uJ^u6ZaIJ3r_&HjvQ5pW zXS_H|azBf|}Ymp$o@dtheDvgQZybA-6NkJ#Zo*4<^dPaX53s)u`-O zktDO4A0H>rY5GjXul0CdQR*e9IKOWL40-F-7zT5Ih?7a(Tk7yQf;r7ohs-I?q;|o6?cYAWn)-K3bq>a|@2O`bO_d z&Qg4OIp-0UDK7WiuW@>>WX>NrtqHgP)5u>QW{96V8Fs1JHCE4@);03q7XuFR{b5e? z4~x^%eo>Ccn48e)A^l#=7EYJB1^wbCqVH+1(EM!}QSLRiYYrOs0PXd_UgNp^vyb7r zk|Gel74<%A!(i^?n)0b1?SK2C`hdKysmmMSgJT?TIC|o=pQ>y*O!qN?rUyv=a5`pa z&-1;kr|;?YN#`!HJh!H|V2E!4_obXXZ^n9hA0afFzKdH})6;&iM7^-ja*p>{Z@9ju zb)#Q=E^m^XPW}fhNr3_5iYq>n)Bcv=?^t0_6^+w#8P*Vc4LI-ot@r*oKD1wS4Sm<; zy2dS9=S_>p@9=yhr%Ugr*iY!aT5gA#duShT2KpVbky^TjYwz>cxcMd5FR8bkwB8gv zipk^tq39*2@}ICdw(;Tqi1h&f3@?GbLwFv8xe390$J2e&^%iUmUX7b)srOU*+mIyy zMdP^>wREXmRlO#<4J@E>TF*@>L}!!7eKK8dTGtYf!}@!D6usm$ztrs++D&~m4z&Ma z7k+rb^K!p^XkBsz=AJ7wJ;mF<2H>VN9G|n^rgfi;xc)u&N6bC+J9^thw)8_4AAje3 z#Ka+3yt|rD$qm>v;4+ro?$vn9gFL7GA@4uZ{{I}h59}`ygr9-e<7B#0Z^5GN+i*me zoSL5cM^TvGlX58>obzv6j*3-Lhc#|M*!fjh=Z9J^q#kJ9tr!j*!Sj$ZUB?F=(s%Ml zr-;yeXrFW48T1|L=?ZVY3B`i%!hORS`VT%p4&8FCTZqrv94Vr0>}Drtb}Q z*En5QcxQYwbSHSzV;>t*cJGai%AV6Wt$&4=#xoWDKh>3~ayd{BHOr+I)3!_oh;dR>wOJWn_VSNZbulKzelJovK%-mbG-^Pzoc zqn2WWnjD9)-lq4d^Wc#`u4sA>3`@C$zL(fg%O|~uI3xnshI0EXeJtm1*RI}yUn6O zCE@WebJ~ZNb~Zjbq25oKuHkC+_&g3=)J(Z9kCV3+gL7k?z!VxcVE4UlIOREyf2F_F z1ASjTi};pmgr+xX|6y->kKXHP7_6uFyuxBIej6Fj%Y(TMzQ4v`InedAZ=_@^^l!^? zChHAXl|!3^Z{hb!y52UYOVmvaty)6k27TWwlI90^K41DsJwTxudoZ9Tk1v>;^d3+r zj9>N={IC2=J-r9oWtMo6Zi(hk`|$s*iyc$(xR`xBu>Sou(WqLqq7QV=ZxybF!yBs4 zS8@YN_Ber!{?<9b!hTEX{esrI51a~1f&(*fd?M48J{D~6MEg~~bG*vjruSmb(Dw&M zD?W9c^PPX^qwmE9Y22W=ZmU>ck=N;@kJJP0$18?2d-D1Ya}x@UJSbMgnWFhv@H1=& zeGkV{crWLC|MvDG{DA5&In68PT0!6WY#XaLV~EjVhOAAP87DW5(Kz+vl?vdrs$6ba z572YQdCZ-k$K%Xt{x9Ny!$_{d_}jY9OQMq6f4A|Qr>$)|Bs=IC2qj7dE-UWPP!i8NWX=c z_0$RHZ~v#K@B95*${kpP<9zn9;NaZ+?iPE?Xz9{-17^&{y0v7xArIDjAo1j6*uj_M zeCCEL!HSt;c?o`A%za#4hNi~i+eRw>CdljZxtSlIhx7V{)Jtwamu3%eLUQ$aNe&>t zKNjAA+f~-ne8uXyqI1vMnm_H&{5%<3oxAO%<;}nJp>>Ufr}4le?l0KKL*u)fB0Q?1 z=3}^`i(RJQn;5Nd+wlyK+=tVnxPD0=$qo4W=qvh$%Krb9L-MD6-Mco3Uq`0Vxf934c;h-VWZay`Sl~=~jnlZV%p>|vRDj0m z{rxp-@pyWk=a=dFOAY-l(!V0<;`o}L*3Eky#uP~sX`J5Iz*h7=DbL%p5A~ame9`qj zSkZ%XUVF1e-!(m_aoPtsdj~%16;I=|U$MkWbSI3fag%;$FbCdPp|(rXpXLwaWyQJw zs?SSu8-Do~ruVrTD?Y89^Y*vuiUcLrX`JSx5--8JtJLQr^;Ewic8kg%7HN8u_CtE; zcMS%`@?J5q&nKi(sT|lo58qGb^!=IWGFbk=Eh@wEIWnjH7`HBn;RSfS#+=^!?v~IU z9^tF`TeRLc604WzIFt2sy|Y?i^sjB2KfUk&XHE=!$LmL|H(dWd2oo(ca`|WO9CW)>%ePQm@zXfhO%T(C_bXyUltF zvagztPa5)kBXb+3H>gZ;coQvMng^@i0te0DxSI8bEAqx;TF0-a>3v)&Q*5C3ba?-j z)LYJV1)pw)9UE}`$(+U=r;=mhs_Ju;dRj+rX5nBzo>yc&?F;P+=o)1^p5Esx^|Wqy zvN5(Uc39&!B;Ps7_Lt5XNx~GKU27mb6#XjeD|-bg*0wJ=FOq#3gB^&^p|>|?@MHK zcYj(!)0vA=pH0HrGdQkeADj05w!q76xt(F|q5TZ?MBcx6Tq${=lhzR5 znHV`(eJ+yA{VL6Ir_UqjS$Lnj0=CuE?Mk$(K;(nGneSc=lY>YQMtMj+sd^$e# zJ^Y#Ud$P+kZb6GPE$Mp;n-$*6IUiC#B{pft{k8O$oZ5lUt;Ev(TQxns{~mdlez(qw z<*i6B4>RP)^Mcf#b9rWNK*ZZPG{3rD(*u21KS03B-5NJx|MtgtbPJbT_L2M6CSnTY z`pev=dG9_Lm7nK7nR}pm;=H0}e17iC4Oe(@Y5Fc}eJx!d*Y=ZlM5;-jRSxpH_Pf_% z=Ji~z93SFzUB-Kn-AttnP9s2LVPG`73WlryRWX~-cXXvZ>%Xy)?nDqC2ChKkb zZs%T{)#V|1%6^kM#oK2~;iKRC<(M0;9D#8$bBa*K-**1CX8dUS&g4vu8xY!P3T~<# zU*q)MW3wb?tHbf6^mlw<-kWC_QH1LqbJ|zDZwlt$tv+|Dx8d{JWti=b?0@CKdJpZp z2%>d)UoBn3RsYcg@%Z=mH>BQhG($@4z?~bDD|*Rk9WeK4{CgD7yRsgr{5QZs{r*#~ z7?KB*Bl;;yMn{YRN7=3??*E?jorlZ-ADzBK`p64IxNKW<4 z40(FR-@Rr~n^+&_G!9xh3bRf<1700-0C%znqHq?TKW9$g6}-}g#-}{5$lQVlN$a== zJ?Hru<~I1OU4SWlxIJO+fqhS7=sOQQK4NaTgugGYZ^iu)b84UG2aA_&c-$>HIC-A> za0*sA#N%b=1`H}Y9mgH=wiVRJ#_FkkKG+}6<>I)Sxk>A@Z zil(=zolJ#;dvQF&dJnA&RKWsc)a!EcS>=q%KiBM;~OPEjt{JO)RTVilFPr$kK`7uBVWOUad_R3xeZwd@4-b6xE*6o zacpUu-uEA<=f~A4^$qbMhS#U0-ut(?&Wxn@(NZdU$!WjJ!5#Q4kmprc4=|=}UiZm) zsOe3pHEJ%F@0wZT7K|*J8vPdeYTSl|t*2mcQjX`@-ve_VrK5d-99K&o==|-+Pnq4f zzg1OyB&YJfE)RCCqSiaf={>TiWw2jdoY3;ZuPRbI05%NGXIhrP}^=S(!}HTpE*E$o2Q(#yjPc_7BS zZ;rg`c+#K8as$@xm?$1h8m)1FQ#)^oE72J>PWwZ)wZwR>c%MJ}Sg@pa2(?!nk1?11 z#Tn7x!}CAPX@5nxmEvcjeVRYLx4S<`{8{k2#(i9KBHq)wD#!nh56J6!HS!Pl(hc2I zJ#hnaH3fPfS?+_92kU9P^Kup*2<7!G<|eG1m)4y-KleM#>Ae_p8&*nu4ZQsbb6Qub zGYtnE+^=yDxWbQ%p6|H)u-TaHyzJQqbg$1367tc z)9>_ktAnRKca(Icp62)5TQSKu-tWYE6T&vNruW@8X?hE8S54xcHAH>B(#NLnX8bMc zZF;2XX}o`RGCqC8<4X27Ts3x`q~Gxwt=xa<os$r_~$gfzQXmIIsJ~z-lDV*_MqeGeZJC%)&V~R;KcIWzq6k9XI|`! zD{^t%#hl7vwR7V8oyl504cE$EUuk}RrNRy8e7l*+eRl`XGfN-I4Y+nQ6zlq)Qt~4? zK*>1?+}YCq;ry*PpNtT`(DU85DfzjNMm#=jU`kmC@kmz>Jg+2NRE4bQ(Z2N+b~E4?=qsQbY5 zRFU-EiAfr_;9dK1;+VzrSL{ROD!c^-{gp=3Q@Lu_ocbdk7fU@jny5Q2y!N!PqL-Z7 z{|O_o(k$M`$a>mmm2(mn+S*dn(>#VRtq}_L3D+BV{qMJ`myGfxeemTG5-%`7y&3 zxtenS5=eS)eEacdApLFy&(AP7=(`OE@ZvT(FDDP?)DPw>jh90of%os3o3y?>7F$o* ztZ@sPmG0%f@8R_w*4t2L<0fp@o7b0_d+0mHt8oQ<)_iE3p6QSX{3AT}x>)byn)K=( z{T}^bg@g0ATT{Hj?$!J>Za{c&JWO5cnBp(>v@Tib7sf0~PkMPsf5!*P&YFnxuk3{X zU9aR86l#;lJ*8g>&EJL-na^Wn#bg?%dD`z&v1@x?_hf&=H3zfMI+UJIAJ>J_!I-S1 z$}J}kF(Sxa^6hlRUvgUi{q+W472$S3`b!QFoTiq$UIK3!^In(Z15X2H({~syI-dXI z7VRHwJ^RP!S_y?53t#2vc3h3*3bmaN9f>Fzq7-uz3woabDMy;}NfditK~ z{1n*X6vvaCt_O0|$&ZJs$mcHW19SQQdv3gQgWDy^eH=|ldl$_|^;Pte8{l5BnC6jm zJ;1HPRdMIR(@H)a9~g2chug0Y&kwV|1??7=#^Xu7X?gP_&uLwwekk^>#N&O|d*DvP zk|KHBxk`S#d6=H3xBJ}L`5HGs_>ae=ri@$odZnJ`!vhjy-`pJMNPou%GEc3G36Amj zl(|LUPnb;mLb;u1Zqs}2@#+0#w~}rh=RDv;=sxhM!5I3EP=Lni_ZJ?N#j6WAPL@8s9LJ#fiig@aau2HW8#v_9LN6lf#)4^(s+&U z8*>Z7b_ToSoa1`U+=gozo1@i@=MR~Cp!VkZA|iD>;eC$G4Oh;eJw&Fu54HUGxXNzN zEb9C?s&bHjTd3G3JieUAt&R_I0~R)jC#pQ)c7r*9&(}_P?yp_W-%{TmBemp1aXWn^ zWK=E9pT=wH3>-g(=MPzLLrA*p7`EV~=Hr2pvrprR;~Y1#p7wQrdrIHK3DNw0T&seS zen*wpC8geUJf8;4#h8CN4rgxAzMwUjBQD3W%;|lHwA-=JKsBF^4_wRt0_RRr^CLOM zAx~dm|M2)qek7;wY;4_#~aWQk9ppkx#22)WCjNJPN(_!xQZ-JjQe(j z!fofeMs&F(ezg}Gr~aDuHy-@V`_rVq)C1J`($1YCe5j^3;j@2k_lIbnCt*Fs?VWF8 z-H240-i9|ro8iryYI&Ca9*SoI@lZ;RSJ}sKW%yPe{rcooe7tpMZTkB(9x0}A1GaAJ zCVHl+rg7SbpL(9ic#zv?>F@YJt|1;g7|84F%q^InqCF1&!S$UveYdbo8~mJ%`(@@H zsMEawZh6p1OV@CXyp|k`&gVEv>ir#!@La)>N4OnkPV=`9reTIEJTJr?V1K*y7;rAF z;_vuCyXq4$y0=T?)PIeCAQGJijm!CsvKUh>lg2%;|6>=tl#}B_nXdP5YvsHw=JhD4 z=p{EG;KnZ6=jP;r`a9MG6e_lro(s3b%uR?KkOSj~mDPMKh$^22Q=d?uyY#oIAO9qF zjvlP&TRG<~GEEd^p3c&^LBF>$0mGK^e1%L`>S^7Kl;RY+e_L9TjM8=TazPJ&)lGOielI)WEPce**`L;b+-f) zMTD*Esm|M(vER8pn%;so7t^^@pWyXJ_Mv|K<5qm1n&VsM^u3N(^gWsBT&|cKuFrFa zip0Oa`_0_PwLa$rVP#sUY;2!R$n47c@ zc^WRS$m2Wa7Wj0k;Vxc$r*c1Ky7YcWmseQ2bQX=%ch$n?;+JeZf5JY7>(Qz%nDiOf zcga0Y{);qELi6gp&yzW=Q%}x{%f4<^?wiyDl#Txi2fj(Jr0e*A{LT5Xr>bh)g2-`; zaANjh8n@xh$eb9wTg|8R@xb)RdiZL4HcfB1e0B%ot~NX`E%WK)XhH{`!#zLpD0;~a z`d#m#*s`PjAo?!Ewjn%^ z!9E^nnQ?*$j^ch_@<8%2Ln@6e?rzeqo^$T4zmn6w!Q>M$bG!H&m-{?((eFp_I>K)s z`kkF&STkL2O>fcuNk43{m)l9!)Bc>-acRG36-8giIS(H;3BS+b`G2XGobr<`R4j6< z^;dGB_d9xEjn!?PzxC$d@u7L|o|vd5kMG#WqWyJ$pr|0*UT&{?Ip=r(YJy)&@qCNa zOHS(u7b;?!2p&%{r{CEOyDerV8WZbdkX{~UNSc6l?)M)zX`J3ax)nvg8_4m@Z#~T0 z)F10a@_rTOH2-`queTaz)jH%>VXIN$LJaJLj?K(mBnmFPo0hQM|9x@gWXSZ%iFg@G8f#%uNW& z0PgT8o^NDsLGFeFXdaL6KXV(#EZl_kQpeZwL*H$4Lb8mD>A zR6g{bD1PqJ-=ut+AtP7)OZyDDKVojsdku5(>v*28U=GlyQ!@0)73KUbT|fK4h@@%p zSnul^w_wr7>zH6{JS9IeUD{u9wGKV%ev0)T`o7r)yu5?^S>`mJTH9WXZ@XOa@psOB?^(F3B9D)xUUJ&E z(<%yEFG#KF0SY7^jrmW>JaB#-A6P$f6834RrYpGx>x$2z_3QDPp1!X*YB>%*%;O+V zm%fjYxgmYGoY!$A2j_2#W*mz~U!G56ZqT}aUTph@+jHg==d@mfAEKKn={i1;a9Iob z9zt%7)BIkVMD%-e4Kz;eeD{jf?)6l7E9ZQBejgmZID^LNy9$3A@;S#p#l!b%3oC?X&(If;Ns7Fa+)6c7L+(wd@b#)aH1*qigKFwW?6dV zX4=`IjDFSbAMQT;laN(pdT^HQQgFccqM4dc>(_d7azbmgh?B z2SSDBR5jL*XioljrxgE|_eEqJQaHxRcU|K~3B<1!j(G3FN7b_VKO5)u?a=tD9(TQD z+X-V|p5e&vA>$HaCqhd~Is9JAmzJF|$2EmCvE zNivV1?T7vE=4d~gbJsP?qdDrmB3t-gIT@c3J@SG!Zp0hviQg{Vqenf`yR+`g_?U3S zu|Jr~L&nKCl5pe~+A@;Q#_!K4JW40N8^iCHOJMwtetAN5>RnY{tsVxC`sF{9^K1U{ ztIQU6=-UR1)R;RmuB+{9J<*WqSNYI=sUL)URJnL%_owr={V=WS`4^A5A?=;$_4C$? zi+mz^@PzAeiC5~gbMn3jcS1vV)>m1+62D9HfM4cX{buv8#}XRl*W95#my)=@-H`T0 z^oYwU^Ok#$=308J|LR}CEjC!@aYT=GzTXt%Z9kark+zR{kQ49o@(I~2JIJFoAT9C( z6gPN`f4u$An!M6qccbfDZC`W5)j#|N`*+KDrf{NynejeOG0X4Up?P!G@RK?VEc;>9 zOG&xsqPiB3AdY33ns9WM#j)?hx_8xAH^e{G=Z!sA$}BM(1Zhc0^VZ@FeCrQT7?n`~Hd0+p#sh6eXEghwQD%_!}4L{&aI*wG&^*kIpSF8S|yimg9|J2a=rPVRD8<{5<(@^?MtD_@R- zP5lw>(VBTnxaTPO-GyU6+y)uldS7Y1{vCx!Xl~|xsLwlEJW5GNPgUxbqyDRRLYa3> zLq5!lmOc>jYE?y?z&W!#q`!OetDbI>gA+|XxXuHU*S%M&q{WFwe;MNG-$}dRms8uJ z(Qj|UIwqOd5FW<--6$^isrc2xG5*#f5uYqB{ZQeEGcSBYwVY7YIxqZ_h1IxD0mtHj zP=_wJ)t|Y=57%})U7lcYPy+YD!+eHbbBEfL>WjQ6(!Ulx`q#~-@v-dEZx!w#FKA0H zQdr{Lgd^t*8M&9B|JuiCXC7%C1=ejTBIl`ULfFl(-yj9+J;Q_>h z-9f$D+^`?Ya7Y3@*#A@kLpsVgl#p`07|bP{PfWQ^t9O6Y}F_KRXg^Ed#JZL3FeyMMTBY%67ReZ9x}x4; zw!)F;A?+?sIncZhetAL@2N%PH7uVnDu7h|vnWxb9HAnnjz202@ck!=;!$0iu z1=gLhVc)ky#R^x#zUsm~*wG&^*viHJwkbZR*!okrLnR*d;RG4d$Jr5%`ZgHhU-QcE zE8L^jr5mZctG>4Mdi~KN9{axdSE7&5sP-v2Q_eh=9``>H@)~8CX>s^zzkJ9m`ikEo zb^@W*!~P82{Yu&|pZnJ`c0tHJT(psC2i(Cr#{=9ssr27PkGP;$iQL0G<-Hc}(fXe{ zV82&sM}&uwhj1YG>>%^Y!Xwn=Icf8gkGd*g7H;}yAvJ= zHLfyUW$3@qu;cmHHD}E=9=Rj6#W8Q2KMC3e=}&1pT95q{T8}^+zWMw5c4)=SE_`8^ z)Jvid)3Po}+*0qG_f6{~RO@;pZWliuBKj!a_vFZ9lfo`fXinCNSO+TekD?ER+NU^x z^=?B9eZ;@6%Igy&&ekly=7^8#yMZ%wlzLCw(VS5KxbDo-gKRr=cS$wG-@b2g{9gAO z^U(LDUKKkLmC1GO-`kA6D|_}N4P^% zM!x1IzBAG&kJclOB65UBC+=!-5BYM>bLeGui-+mKcY1x$f8IxEB6Ks;WUN<}c1O+| zr9HvYyyW-)yf4r@q4mFwR)sn*GwjE{gM)4=_vOToE$(0*VIZd;Vm@zLkM(6m_i)c* z38e$XpVS~ z1TI(YYd%*U{4$(f(1K5Ol=XmGuen1*?!T=HtTF2)&58EzoPzmHssBZf`K+Z+RmNa{ z-zMJOI6K&HWpO>^W7qA7{%G#`nho8jtG8o@8hXvqzuwf4?-mk2inQMMt0(21i+Jzd z6#L)8;inyW#9ymKadw2GKl0;Jj5kWVBV5OOT;mo2sV9YFzBu7Pe$YzRbqYru!mGde zdCA@m@$fsO|f95QlS&e{2@8^yt@&%;vV<60&%N&fLoDelS(~Yhnle z@rtkGd6VZsxD!gfuN7bICw`9bK&VUQZFnD781}>dbp=}7=6mza`=&YG$H!gy`iDz= z4gNC{>u(bc;*{?3&*Zrk?$OSLMcnp9n_2eZfBsR1$Gs-wMWV;J`ljogvrngj$z-ogE*y*$isFB8vT5W9&tXA`CKW}>y{q<;)1VK&937AiymUu|D=qB^^;bCC}|u=MaN z+6?Cc^UeFC^-5NV_6`(TA2E`8N*3{~Tf7Pi+V5WLp&B z-gmlNdfewbJJk8viwxe#Kh7JIlV=qdKS7^YbBET{sE2(S#2*ljdTH)0RpErh_X@|j zM61&n_a1NAM?Yo6ZZ+o5xd!j+AOBTr5sw=7C;jh!YVOdogWa(HIVqfOF>m*uG}BtM3eibRcZFf{TJ@h8^hGk2)h|s>`nuYzuH|NUB<+iTUT_EFPQ8#%V4<(Ue@8UK|51)UOE_jp{5*afM_!i# zZTXY_UqYjwXVHi0aoeVRssApE!_Qgyn%gi}b*ns4+R%L}_H`O)ajZ`{UW>UPW5)8Y|2S_An9R!cu!>_llzp%k2Hmb7caozRk{b2<0Um)3c) zt~_}n{xxxB!+ylqWZcu#O*GVeJ~hX>n=7C3V-|l*+t(cNPLpEXZ_~83?NG)-dEB46 ziJvBVy*^?B|1>M1rH_z`l;OXVCb4*whI~+!N0d%#aVNC9TnyK!AoZTMA6w@KKd#D| zd;I5lz<1F4}Eq=am%A=ZAWv59+!B+ebT>c*wGyGYmciTpRjpe-wx$|dj;kVOI!Ld z%_{M|`r_gki$~~uqI8I}on-J1{_&)Ki+TNk_;?bOM|17pGVYuB?Yf+r6ZU-=#g+as ze;3~lrB)Bs4wCjl^qB9P{x0%nNIR&x=bty%os=qlM(Ljlcj$VnuDoN@L0<#yxA;T7 z#6u_m>`JEPa$#r8>*WlBcn_6MFp3EbQ+laa>|25c=@f^QvJk@w+r9zij7o zjCB{(+h*9&9P2f@JmlTmGFzOm-$Wv}Ud6f=$Nm~ITU4(Uc`P2L{8OsBIgj?RIQFM_ zFFW#HYX4k+q&$cpm{y!GJh?>wyKjDZLTb+-^ap2Jb^@Uq+beOu9OVoi_P?w8u0?*H z92Q5Oy;{RLLk_9O^m(-&?bo%F+;6?SkHS5~Z(Qcfj7%ueF`1f36NMYr8Xk zmgihJ)^jyYh`b8N4f|S87*C#sJhw7#FM8CIE6%BJE1LH~>%%naa89>-A3ffpKca`f zm-UM3Fj(Gq;ZdyPepgjl9zUKCS5BQrs18pG6*I~c@Q+JJTDaA6NIR?TXzozK2O;;3 zqU$X^;%z3S!als>Pm1279tB=UT&t}QQ{I`yIIt<5WhX)xmwdrJuI95i@@yoZu1YTy zzgpXm@ik3aoIyOi)IY*8{?L1wYTJ9bVMpuHe|h-|pM&ehC9ppEA zb^>`IW!zhM1pAa_=Al#cxQG4-k7D1nW!RV5EN3IX3^mL2Q-e>M_fvBR@s{hcU!M71 zXijwIwL6M8WK{oqzkEB$_gEBh*HTaF^Jp%U!<6puJXPUyxo=_zt!5L%Y8eDqf4ny=aYI~IPyStdy}hml=T|I;g@b- z!WTx%{GIS9japiq7ax#zM7WMOKwR8xsW*fNLhDB+=LA=a80GXx|LeufrQ9X$q<+&x zb1yh+!>8DHQr=VH4qbd}EazS(aU;T!ucq#3F8gYS|6NVv?2y|02>VmWe2Q?)N2K}{ z>&W8$~`n@2|$`1EEnLUQiD+$NSIV z5yJm-4nM!nw@#VAi{^-jUoaKxr@rxzwJuIi6Av9hz5*GC5so}sqYLqqJ?8cLc4$%F z1)ydjgNn?mR;%R-a56?xNp9mDqd{l_UZez#r68)+g!7MV~cx~ zxa$UfT*iDJv>o*8pBLcLzooYH5nOLt>=P>fkCX@V-JgERw^F9E^iF8YXY;WC$b5t6 z!Lj2MTs4y2?#V3f1@oR-!26z^rT<;8ucw_yJ8^=&$t^u0ZuwJQ{B2H)d-Q7DG5D+T z{Q{g(+d*97sKR{Vp=}5LUc&@fZzA=uE}_<=pEzL@&sm$#uunL4oE|HXC%no|i+jNi zNSH2ym7w!euHNL@x_sM%I+@Zbe7W0y1vQAkz zVVtBpFTX4OAmJYMJCOf%CEvxwL6%tph^_qLZ6K{@mKl&m?T>T(i z_apv5T$$9v!ilaApU$uKlK2zhs0Ri=<}m}L{t+HVKeVm;&CSVHd9Xfu@lKxoOMJg9 zt~^mn{8uHcSCY6FvEzg)*1oMq1!F8bfzY!hNwAMqCxeIm>&-prM{Zfpyg!e{}qr%o{$xJb(Q`QJE>-*-g=SoIG^%%`9ot%jYs~L#TE~Q&X4_8ZJImE;8Fj2cf5aq zd;TcvRrPgg?gjhoJI!Z$$as!$hZZb3#}6~Fwd@eBpPR+acuL+!(RMi z#?9K_LjPLoJ>l3F^7dlvuWi0xx;#;;KVZ2!U(UR5e)&WHG;OU?eDb+b9{mhEPLmc} zG5#j=T-tu@vpsLlAIO)R&DK*ZyGj3he>K(%E)zLB~mkng3X zrH^8tjt^D(mFo@O2*>!G&aGnBNqwvDzvf<0<-xe9-?y2(Wp(^K zyXax3uL_MZpL4Aza(kA*K6EmkFZMmkb$E>`Fn@ELeFwi=>}SxK**t${H;a40c02cT zvzyHw6Kk5JWrtMjTy z@%~R-IT7!=Z9IQ4b&REVLLc_HkM;ONEglG6>0FPOo{0CuwH@EqqkV^Py;6Vr%>TFi zRiE&-U#psSz=_h7ILFtbg)HtNkLbtza`(T`=zA=7!pL8mj_b6Oyi&p=)TwVL{`{WI zI}4AJ6Mh46I1ep5PUu;?uFU_&_jmB``a6Np-L9E=$1a&?)XdZW<`f)^JO(M_$3^3~ z7rc@WQM z+q?ZfIbz%&U7iSa&e4whs`zo;I6G1L{DW^*u)MSjVh8KL&bQ<;#pC_$IDH@#7+M8! z5d{tV5&WCu3`jm2e&Y(Qi@X1td%?ZOYVwqdGEc4T`+D>@+aiyO)bqj#>jEZm>N)ZA zHF0{(>y>Hb=GkjLZ`w|nRy{Aq`AbCob2JsCg$&}?W0~r99@-739a&A{2{o9_oj-UKa%obo~tVQ z=|v2E9vq`z@aAal@jSl&0=?#lcTMn+i>#J*L6^tZ)A5dfaE3<{EIUL=N6+EZgR@%P zqvXXi^WvY(_d?r2d+}(kYLlm{rPuiGBXXj`J#GVs~5W z__kwq8+~MdR^Ll+f=68NWK8-epea(@FJj(;@8`ad}9?jZtPbFN`!s3Wu%J&Z9 z&}{n=+ca4vqAy? zyZCw>*Y#d=)c-PWCfuX5EqcNaDQ%q>``s59!t?T0GPvj4T+kp57Y_HaxECya=Z&Mb;>4an6q|P5aC#(4R!jZrF z`;T~1lcbh?5Boj+fOt0P7mFV82J@G3xnTL9PgzI?e3Vbb-)bqlfP=jT; zxX>Fijwc-B>VMYbOrDHeYOenrC*kv4Zq5GYeTe-|@~JlvzgovAzptkOUJrLzi)|Li z_2wGOn`g*#F7`co+&hKae3Xp$3dcBboe4bT@@14kKR?1FSkKjwtB>Aq@hH}XW#CV4 zOZ_Q&C!`{SxNbs;lNTNcJ=k4Y4Xh#0t>$6;q2Ak-!(DMr`q9F@V6h?l`17I3jqCFD zlyPHyH`&7_7DpZrb&zjAm3koF4)St$yV@EBHeM@ilF5`xlQk${B@C$NBJB72Z--;(N7TbL?L{ zu9Um;Qfpg}{k+a`vw>uBjAPH~?@s|H{Uj|S92%S=)_Gmuz<|_XgdLv!EyR^$>x^)*L;sO_ktq^ zPH?jp5`Q&bPhb4+9d_D4b z@8T3W%32)t=Z2De`+O#gd)N9f{@$rEq@lNUYe~G+0W<9CxAWpbLBd#;z2ip$4 zmZz(l_t_eYW4!mBtX%n=Szq;q#&J6DYR?Ihr84xId%$I$ZSYIK7AKiv5gAdm8QXN{f5JOP5oqB58+P+@U*{ zhj6phnJtbu_9=_F!(X~>($AIHM|<6N2{%t8{RQD+YI?4%>X&1urN{nIy>9XIHo6|r zAJHQp?37ph?1T9JBl-V~_v_bjTU6T*q&^pYAT;8~+v@K@@%==t_rEh@?s&fNn&jX1 z_26FcX2IEf|Bd+mtmx@{Rm3%s_#2`}9LVNW$fKhD|M=_GcBsie6}k3f$!{Qf>{E6t zr~BQ9Qtt^zp4cPP_}4NCEj!4gGG-`$xI3H0olyC!H`TNI;tz4=|hPYGndC=d9(z?>6xZ-|^LlZkrX#2K2y#Gu=>%4)`SF@(8 zC09-R5&fGv&Z+J2PxQF7{4Sb%!6}K8V81Y#uhI5>J?$Ny$Q^a6r(s`n#9Q{=#RnFe z_usceM>m$iI)WjVK1>%IJmnet%-==pBecHw5>9>Vm0yO~d#^d_r6hUy{RufOJ5DHf z!Opz<_WYMv+dx;{wkR{E!~abC@_Pi6CdZsO+V`|Inmerg$?o+o~k*g<}btL-?Azn?Ge zjQ@=LcHsY<;%ql+SoV>>?!slnb;MX4ad90Us4TUrSUgHc=N(jKZuc@c;n;Dy)~Mt* z8#u$_UT{+5>&OdN)Zz|peQg>~nwZ1l7~jkJ8uBSswzx+R+LXoqiRSaK&x<(Boc+)q z<+t=07s{QAdpAF1@D4b3oRoJC@Zfvq`>wecEZDUI;$TzSdirkEXqB$Q=N8BMh;=3S zryB9&!zA-KlR9?c4|1B%zqXHkeoys4o{gb~9S_Hjvwcn~%%_|0qvpsXk??gsokZq0 z_3!KJ$@}F1;wARQ*$Hc|$1|K^lR6--a!=fU3-^Ni$1h}cQm$7x;=l`~bi01k)Urd^ zPqQq4Q{@NS4&6IFlGCn9YU#sR*Z)do>m>EI*pJY(`IA*^CFAqLqwt@{@w39>=L>g2 zz2B?JIo^SRqKWv`Vkd%q@k;U86jGlH zN8Yeqht+oUwqw{a_}m&Cr_2e~)5zbfB?(;IgCLHzQwNol#G4prV=Z#X4Qk&G+ z#sdxex-D>=vC*vViuRK&?gf|hY0R4nn$MN5r!(K^_&piV(dX5iXi@idynMFIBMbNN z-aS(ZPDsB-IP$UYb$Lh9=lH%n=h{w$^7bvxH|mr=I%Sq_tDqWs+2X=_t|8gUG%vBujk`*HDl}K*musiLvHtuYT3ax z)_L(+T5M5Pcff2&Anj4=c{-|1@nITdTP+AlP)$DKC}x9aehSBY&E$O;CltS3IP&>qf56ix zjIrz>-oASmte0tPapVmvQ;zeD5WiRKV4a=H#UnDxdm%hR3u?UWzOnXK%MRuZOE2a6 ziDz5f39XB!SLGc0x&ond2Y(F3tozv1``0`keDB`Nof`uwES4`K)H;JW1e(>rr4J>m}=6+qsVnFiP2ahw|Gu+Q(IZ}jz`9n{RX z6?Bhx=HHKs-y+8aLCXoKVa6HfZu43J*_;127Z*@7RJK^Xc zt2`Lj9dHu+R!IFO`Y4UtGezCrIKk38p`1ILa-ntU430JgpJMMV=l7>XjqB3f3wEt= zky|Sn=h60kJ&iAR6mh)bw+JW1pLkp>LvqW$M_2o;NB*=Jiz9Ezjr3gL!<-h!{?6yu zsvcK+TRe*W;;*SWYo)%G^CDks9{63?l?55DU-Q>G?!%cZk9+L6C&x>e%W*~RNI zKGVpZJ8`UWUTue{)TpW$XO?jZ(R)}2U5Pu@5kEnAnD(X|!CPmEe&ks6jhg5Ng_clyk%H#$J$IMy6#`v(NtV%Mk0E6ONI$;o1wtZ~wpYhbkVei@f*p`-+`FsQJdS{A+xkAI-z~bDV=6 z-p6}=7{~he7VZVdk66vGI!OOQIL2LbD0gToX?KLX|`&F53=iOzFJuKlURRchwjp$DCFV7+92%MSAS z|DBqf6qa#I(MPCa%avUFruZ|$u?{Q8Unb?qT#hp$V;!!! zelPx8elN|v;JkM?@ONtx8s+izRN_)=_vl~d@1;5BYsUP9{p0Fbddx$mPQl$~NqZsZ zgl?+TX?8o33@7MR+_}|v%eb(Fy{`&D; zHS50gqqQAhPpdxPhH+^brxH%Iz3dMv*-G&Pg~NaN;d}MgF!Q>!9rVl2o>fo2mT@i7 zM<`|TuQ5K~-6&@xznn#TUE>$=^U+$bxgMvwhPYhw_wx0~>wAPZwKB`2IU(QXdsr_p z$+C}pv-zf~1>bG3c$jkCX|E#p#9z|q?F)_Ll=$IIUR5}Yq1POFnGU?mN!KN@xI;;* zFGT+H)fOkr8>Ztos6*4a%?{^V-s%6+4J$ z8*!Z5FX(LP;m1ZF@vsZxkBMHdTTJd&Zz2AG@CYSry@PwDXl2=nQnA%pk*~F-#hp;y zBlQvYc+%p5(8$lqsprRkFgfWf9@f1lkGvnjvER<&-yQ$BM49u*k6qN(Q>V>ac>i|k zKlTS_SY16IB z$de=Og>dAZ_;eQkds*Hi;epV$dYjbMeiM!Iqb&F=I-S9-GeX`|t=Am>$w$TAsh7vw zdTNvQ4%Q)-w>Z|NzW)mkZdlRc9yKU)9sAsswRo6HUHKC0!o+`;^G2vstxWEwf;}yL z6zjyxu(wqFVbNoMmg;pepIg<^2STsqS)oeKl<_64kNDT?);+E|6no9kYp(s4NxZLn zQj0q@meBtlC-t7#Av)_6bKki<&eD6ddvSGl@v;RL$36uk$MV7_QV)n7tdp7-;iRL< z(!=jf|1OUnC*vWa$Npm5uByTPsvCO!Z;rF)+E%WbP1+-$`I^H0?(yAP|9Nh;o<fG@QpHALX(1m`CeU0{hISwm8;* z1^+>QU}*=XoDSAKRbuyaMnkV{qW@JYIr8T;x40KPQ?0O@@1>rHm2&!eI-7hu*RCRd zrnav+<^|s9z?_T)~-A6`vt*+G2u*V&O@x01ypw6W$5)uOWbzG(Z%_tkp5 z+CONWp(h+WPJ=OzvF=;XOZp%DujXDbP-_q0?PA_fUynG4o2vh}=6%qd=uXw2u%F0c z|GTgxcIdCC}yj=JlTBTHOIaK zyC)-l+pK?lJ@(=4z^#9ic3+oAbE2)^FH+T8T=l;jTTb5&)^|7I;y?c9J%tA2v6s%M zT$y_qde1*@Qfms}A6=%A6 z(YP*OPi~GioZ~h99P0N)ms4}pV;7q9xdGDd2=}ng?higXSn56D@RvS*gU?))btu9k z*taVudtJn@79OQLyINvg)IKlPapymyR;7H~Do-F3d^;)c>?!Stwi8A<9p_B_`>NU} zGL9$Q3pP127X5vRhY^mv?iY%?qgO35>}xwj-<9lueP^av9C`D))^Udp|H9&78d0u_ z+oH1gOJYAln}?UexcX8{AEoM9JHnsbY;nYazxqP8%re#BQUAKWD;ZXgeyCw_FIZ;C zYV02^#6I$A z>@UGtTIRO&QHu1MiE#-T4-~x<+BNhm#IMS@kmmZ&aq7));s%CGe_1%jXZ{E?B{R$E z>*-AO9`4awpBd-X9DdYS-yvTXSseSA4EqLoUCj4g>%(-k>wH|7vh)$E^Sq}T79VFT z<%v>mEJ%{z&2m6iHAUQdUMt>$-!OZzSy`Qq}O z!M=^=`|jJJlS4XlqZiUo7kwCU9ho`l&*uHt`UpLmosEyQGwYc~emRpb-ikbX`HXUE z4*$H?cNh=XcIAEd^_2Lkn%1h4}vgFwGJ7e(joib<4B$i06MCs(h}C#j($Ioyi!t*&fIB z{fqqsDjmUi@_QCX{LJE|-0rxvU-5csk>WEhm{i6^g%kCCFhkAyK;CQN7>6lxo~I`r zVA&7Tw&w*@k}tlpc!XL!t%x{zX-~yYlqP-F8hMh%&k^o~K1n)O{njagioI{b1EIk) z-c;HCykgoXeZ?uis>WARiy!Fg!M$MWMG4)BtLxc%deqM>~!-<7>LZ!h5_fSs9nfF^d_sfdnUupZAd%?dp7jiTII@;FLqHK57xpVSfh@R-h z(PJ3bHOueYq0(8`ajNIizKb5?ioYbs{$j%{`#NuDS^jpv)DNPMVmtwPrJk3u^vD}l zEDiESn7>QJzpkQfZXl1k)Z5yQ=3ek#x{5rQrQH;cIFNIzxtv$fu&?#VoAi1GH`7!Z zKM}o0IkR_n@2)y!>2-U(iKC0n`=jkduhpoqW(F6YIachec-Cb3f7j_nZNne+?@6olB zV|hxi|2$XVVcPNjB>wrRu3vQhBl-wcpBPeiYlSV2c*Rbg_;hCJ7mFU_X~T;1Ff#AI z|8LGe+aB`M8RCCxrvJQP&V7yC&R1p{_I*9|{(CO=BQx)(=GZ@?^Z_1S$gF>SJIM3d zifg9LVA%afO*_7xa<81kNq3p& z)g19r=RZ+joxW+?p~v6+gL;0v!95&1PWe4$_|tUeebd|vhS!ZlJh0S1F>&8>oIYDG zVV%Cb*TRtx;p`0V)mq{rgd=ZKsv_=!EA!&MFI>+z{f>MlEiLXL{%Q$lEGhM*=rLb^ zEI<5wc^-sg9&k+op4D6Wt-_yquwU<~*k42DUxYh!cELY9pc`9uaJ|zvaNwEvv7*Pg&_0(te^<=X zhY_EAms7;YtB5{AS32KU$KvNPg+~zwk{o$`D_eG)kke-@;xx8cJP`W*XnwW&>2!nR zox>+-=^y#6-q|hg1xqw}1N)v!`>W6E>oJa@uujU>BksDv*ol(1*l zeW!-Ide{w%M`(ts#bc97yC8PZPkdgHJC-SD>7CHPLz~siE#`X^fW~pgr&{W+iykrb zntQ>``=5wX>VeBuN zfwz{FaUrpfb)zK{^3dV$TXv$zU#FDgEi}0P4)QgPR_>nAd5dFxO|uJJ;Ohz&M;@~; z4F7n|`nLM>xjoHVjuEUR+_>3DeE(AE-~`*Qbjfao}x!Q(H@l zAEv2N(0;so{OJr#~PmJK6#LMv%UgkvANTPvA2 z$a^Q;!#Lhq@!sI^*NQkSGazCy`-=B^z7Aqwzc#Ngd@IXRVD7TMDh;{cc@YQuQ}!oX_tj# z-Z0--ZrNSNpM@jOPq~nry!gjv`Eg#vDb@R1b*huz;t}janv%Pg&Sdc@m2X#uGoC7| z%Z-oNcS6g$)aLEuq#Y3+2yI-Qi1Xb1#IWO${$|l$Z(-kuaTfQ2drscv$(yC#({_A4 z^=Q+UKOQB1i*Wei2TpJ}yVRe;vHtXSU*y{_ZrR6s*X5-;zr2;jBUmq7iN_q2abU57 z^{1C|!Y?v^U%xz|Zd=mfxsrJet&jNE_28!uxZezkFBR?u&(9(34`@C=zMg`Es&lkX zYNI@w6HPwlsid=Id@SA$<*$2`w=^nb>BDrZ(OSOM(Y#)5NB%o3DA5_#gh_NNrn>^}JueuPAMCj3-P!st)F9 zZ*h;Vubav-yTm`y*Q?KqdVcg%RjBEI>IZP0-@Y^AZKNL|dh8E)z6gK+k@QnE&!cTQ zXQq{M-`+b2$MW83?gd}Xy~dpaEiI1m?fX?Zar43!M?Qy%Z@D>^Y_zyX-8;-u85^y! zc$k7ovnX}atbg=*BNTi)CB`M}^G30rJv;KuNc*eHum2q9!Q#WnuVuccnj=4hO6Hbe zX*c8b@Jn@mhRnVu)`tIY)avJ1Tw{y)OQQF%uBs0Yj!M5tIM%gq}vms;{SQx5^U;HA_E1)xCJp;6(a&NjY&MXXq;P1=^10 z@T>dJ<{cV$98_T-j*59iPpNKw8bFwDlH@}hR zR(OO8U!RG6LbYG5-$&t5w7>cJ$=CAU33ozgFN{-t(r-2NVgI_a4Opq#)g5YaFPMH; z9{1leb1m*r{>uaK+{$yM?Q1(k>HnO`GcQSdCmj2c)qTJlYqquQVBB?{;s2kW9ex1dgQ1LGosGU8x89eG=*U`+M^X?=SERN?hd60L_Fz=hM zrz^#4xKmq8JFCm1IiY=bIr_T!dufh1?VMM*>qoZzFs*-m8lO8T?YP*%JX*V~9DGOo zZQ;l>nY0z^ZS#Km6LXye51c=X4ybyV#0 zukCwS-;tQ#Y*63QBOi!c1o0cSEgqr4{c5UBE7uxaUo+-S68)ue9(~{9Ua;r>%90B>!;+t3HO4f z8thj8=9TwCxI;x>RO9cn$#X88(BB_`{>U3~=N0Z@J=9}`u|kW7X~NJHsE5~EJVI?w zUQpK>O8p~t5GT1|h59?`bxV)9i*E~Jzmn+|4}^-fD2jgl6_b;`-gW6Z@Umy(|M>R7 z;mQZn9VzXv*dYpL8pk`s0oxARsne?NkrEb1e(WV%kk3cj zHL-*I_7$u1XNl!`79NG28!Eb|r)5XyRgS8&L(S(Q;$K((5h1*f;xB1CntQ?1H_oaw z4c{^B_=7Z~x=rJx-e-e*ZSuMSXedT`UpEHV|DSGr5@_dQ??kie)9aqx` zd9ce^JW3UImsPEsN_{AHkhf%7RxVKFT|@8S81>|jMcpTt`&b)hPId#W8NV zW-CYTNq<4x*LH{o9Qg$MndZ0j9-X{D@d4$@pt)PC!W}g4ME^W(kZr*xi?3qCD zZ}6`MSaygK9#6tkFPF7A;>v;{j+rKYnr|B#w2%86@X%*z487+c^P3CQsFKpY>$Xe( z4vup?MPhe)(!p{6Egbt?Z#;%|SVJxDQ0=zE)a8`Yo{Ap+bL%kH{Y|m-9!<>Ll7o({ z`xbo|`EB}PKR$`K5svtfL@_+R7rw-N8K$R2(4=UmpY&S zxWS3^^|m>ATFqm5ZuNOJ_kvf)v{#)1;%^9d=+(~;x&6eCEIULUawc>0E|uq5^d9|l z^LK90N&02N(O+1w9Q|&2PlY4iI8_n8TQME%>yPj#IUVw2f62Pmd7V(D)m2q=@hXFd zNq=*!cU?W-X_ixSn$NlBVcPO-E;Z$z z_@AOjobX@A)%HIZ`elf{4_Y6k%XR9i&)*V1S?go)hxKti?jY}k`Fv`QyuM}Ds52+l z_{Xt!d_Be)^6-ro=6k9+(Uk*Rd0T23zv>I-+o7Z9lJUVAy(~Lezf<>5Rk!6Bi^ES# zlbD~>k$#r8(+Gd?C$~IPi-vSH^qPCYJ0o)O4ibMmUQa{z5Zb#BEIrXbW7=XLx34Yk zQM)P^RQA(TEFPxMGPhQv6Mt%OeGbRDyDANTG+z2~n(040d(=i%%pSJ&boTImHGP8l zJZn8slE=kWlSitIX>1*>NnwtpZ~W8#zhNQcAU_}!Zk3ixX$1a|GH{*yoUAcGG3wWXzm3o zKW>V==#6YWZA_Ni&AG$8KU$CYu?Y6{yFiFv$VXW zwlos|T>rbT$2x`IRr_RbTYBUZsGSV^LP~uudXFC8Kf}%}^S)_2SjYEHQXX6CKhG6- zglbgC!d0(HeJ*w|{#LDrI$U>?ao#*QM(y`nIxgO&j>QqroII_&<+_YhX*<52Zh334 z&Z4TNCwe$>J@!?RdRX*WM|L9vuU#R2xNyXm&YY%pC)r`yiO|{is$*U7YJ(Gwu}-{R zb}o{(n#D2h)Uylvane81_I*92nN|kt;09WHqKn^G<;#<0{zdd2;-C_8liE`)eVC?v zQ$zK=Che5gcfgV`uCTzvDAg!&ZVG9nbYAMHBoT&He{yhJN#0$mS zq0LR^^MVJ-EPa>?Cs@P{)}*v}gqqaafc+Yzeh@oR?6>}tnz1jg&2ofeJ#5ooRnvJB z3_C>nekN%)M}2l@o5hjOdG9yyqvU>SJHDP~y*7mxBsI&UIbl6-EiSxL;(|nvIG-P~ zyS;|W?<+h^!`6m4Q@SM9c_WyA`4w@|A&W;ToBK$;SS$XC*l|Kh_dZvlKjeASJnUcB zj#mYF^{^_29nHPq&b0x==ZHTOuczS!L#lsD@dtzxrCwIv-FZ~}K;a&p&tHsR)X?pb z{s<4#tGS8YM$?8`=Z#RoL;ZPC#tasZ(%>z_u|Mwti#wrGZ8LM_=Hm?>C4HvFoktk? z;q`en_kxN4%7W`HZR_dqfNXBD!RGzfdgS>ZeT~as9c}47%+Jo_k4j2ACHBLVziM&5 zSgon0kI=D1TUD;ri!C0dKbO`)oXL8NJE2*ncBxGTw;DY59)9@jJXiR0jA+U9-J9PN%YE2EPpT!RP zeV0b5HHV}<5{~r()%&OuRlmS@^+)p<|9C)=Ed1Gf=5wn#@@0KKnl=4If5k!Ud@SaWSy@Pv^!;S)SsP>u#cHw(VPn7uk)NZ#~^1O-Oqv{JvsvPnAXbVR>Rj&cl zinlB~5t^51H@8?LexvB4v~*+_mH5sdmfi`yTXL4#ursmJ*X}sJp7eD!XtD_VyZxs; z;23|-{0^u3@(t6D)?*)&4><9AxhziTuNC4(Ka2k+_A&pScLA60lfu%6X=jQvYU1ja z7Dt@h_EcCG)5+pd#8Y+U6`gZf+zDMNa6;{GD*YpUUQd_PsXX@sPI9oKq1W6CUf;PH z`7q@@iq})wvOPG>v4Eu~%%|kw44>HCqYu~J#-%ln_}9g+FTgrAdGCZHzNpkZ4s}Uq*!T4q&so4v z*S~IY><1Jx6nQg>gX>Sc9ooEaJs&Gm$+R>r(<1eT=#js7O=|~Hw)HMUH_K%RNHao1+dAEbXO+@W5z>#D^E)>?LOKkH`T>6^^*_;zT)gSI>?t@I1T4%W*o z>BJviOlsMQP_FH}FdsC-;!)~%FQO7$)!*08kJxcScW<{-jeiq=N%OFOU1{#$Q?K`& zX6QA?{%);OxWfyIUmdTfTQNmA+4prVyx%E~Zv$#hcK3>6t(V@e{RlE4k4kGxAEkM-$Ez`a?lgD| zjveR4+DSZUP6dnOx$2*gcOD&Qafgb%H(71@cE805@mC=}T(pP95huUrg_@zH9@FL2 z=S95Svh&z)LjSj1m*x@Lod09hqy7$qH^Q;w4BFO(hqW-@Yt4~ICC}IDe0p6T{n7S( zJ@Ur%@nA_dsC$hs`rrK@_bA_}p6bAs;|A~R+x%zFV>Ke9`MYS2`IjGO zA@7a(yZCynBYUEDUXu1J?Ay>K%iN=Z`l98FIQzoA;F=n@Rg07oH!j?vV>|%k-ghlM zQR(FOREnq4{}R1NDJrK!e7@A@!u7oKi%^Bn#a|Mx^Wu73{&(?bgkzm)vd@tR?#jzftYcGDBJ*&kX%s0KlddUtJ$9{0%ZsI-<#sAmk*Y-WEFKEfX4UqOx zc$l*NJe4P0GJjvKk6?fKvYhYbyOw>-Cng-LMz6YJaVNCxd10(um-i)ZD@d)kzvah=O$A8{OXrhSc?}GgWzBTlDaEyJd>Jrc6G8XrOhc>Vp+~ISJJJ?68 z73WxDmS5W;$}rngs^MNs?_s@pe$01ByAZ~|Ykio;_pHKM3pKIy5t{n+16B5b)Wcd& z_;Z}2^Ct2)Hwzeg&CyoBbgD7qLUw7dRcqtP_7mIO5aNr9gjw ztYs%mb=M73_otcV=>QGp{Y$lkpS;!3YmW8ozcf;3Cz-#aucv}tU*q2^jj;4Y^8+hX zo*(3SlX}i^TDh$dCpIchFC6Xr+Dq!=TeB?gP||cK5vTg0#nHZhQkH)ykjCO3)h{qt z?fqK(1hF5c58wC!erY31AEDzjJ+)`4JXfNR(#x9d)yF@I-zeM(jmY>f@^no#>=5bu zmRe=t^6mOr9C7ulFJd1G8E?^cd_6t;eiOzG-n8_D`GuOCG(ifB>vcNo5I5M|;^;^J zb(^OTYHxAut2^$Q`giL)7LQWHO{I|MU&d$R%M%)Rf1=9XSlVmN!~S(Wcykfr%8D3v zG)Ml`Omn!w>FyTCJma-7m>;QZaUCz(l+PzkXK|0_v^s=%%MUCbCik}#+@iAdqoq6% z8j-Of&n#WV(nsk|>&2@2g_9Q7`-N3k!SnLIYdcZ@x;D1Y=+^(UkD=Gx3+}j@2J0Cc zTil`TXBuO@Ng<0Pk6g(%m|rVuam*v+Yt6$ejI?-|9zDtrzplT|ad&^_M803MoSGvZa!f6*`)es%Pa_6q!}_dt7RPv0o1azs zg3`{4eGl=;J&gRo!ozg_-L%NBBJaC!%xkYsjro^}mi;J=EX2fb zaJSTx+K%R4aK*z?h^rrB>(P&B$c|Ip;_!2hwm|-Q@gL&tP?F_esdD+u=ULm=_46FH z;DNLkqK{Cu96?pLmie4Cf*yB!K|-$EV}@~F&AniqdkK-ZNBjqE$Jdj4{Eiwre3PXo zsywg|*Vz8P#jzjFNKf@`V3xBlG{|S0XNkH}z_g>e7c5+Qm^!^o>Z{GZR_hB#hhqJ> zI@DkJ4+Uv5$6veHKUj##djdo7%*lV z>X!!``wsNjiTJDHrX8)P2ghftg}YW;9M8$!shn_8K8t(s$F9OJtzhvmRh)EE6?rA? zrkpoIO_OxTyu*L)Kh7J){MtpdFS9K>PN?Lmg;-Dhjln%#PG|G7udptllEuAXr$8E> zQ$?N=eO_Nr{YKqF`y%bTaH2rBIh-xA)IY*K#F^K_`t&y}`&ifc;u?^?@X;}*^f}xxBID% z&AzjE6#G^6Q;W;~=Q)SQah4q{{eP{U2Y6LQxA(UwMJdu0=|L1wASCqA&dS;W61sGd zD!n&FK!qJ3pi-3*r5*&7DjiW24mJ=J5Kur64=TQ(2#5$){Qk3(f9|l~d+&Si_ulh7 z(UTu5Yi7@$S+l0CaSQCV({=p6UDEwH;=pixVbV3+LEQ3QxR>FUdI&Bu$;plV2I>Ca zX0|W$T?(!oq7M|=8Q^BW&|BE&U_Y19k9bWFZ_>~HK|8>7qtfg3IM@ffeZ@c?_NsPy z_pbiz7mU2&Ug)Kvceve}(ms>p{AS)mI(5GJo$c4iBQEjLIP9CU>%Xr@y3G1X^ArlE zi~D`r|Dktf-M>KnIN#Co8C^4x{hHx|AM9sNd&2$UYtBEAd!byF^10jQO|$aj%S-9B zA(%(yIF9n>T%qIs(R4*X-jjJ-U&OuhdjmWy>9dnHVe2M~_rQ1bERC1z(3%?nZ7mc~5%geIgH2E>|B2_;drfplZf1!Vk_zz8SPBS;#FS!>gbasS3xr^w%FD(mV8*uRjA6u-NQ`%&qlnfW2dIipgo-9fc-2Y6Wim^M@^Z{|3i@`$UMm<#8C zXjjRj()^bhnwX8{K<*@-T>6HNI={uzFJj7V}ha-Lr%8GNDOx z#JTd{DgK)BSU;M5z%6o1etbRBZCqKbFLw;+2}|n=8*%<3A;2T@+St+h@I!k7JSt8V z?2C(KduO^%;=cO^>cF@8ehm*J|A^ZVGX&>lXg|pj=jV7n_u!i0mVRFz@sw~*vu*tV zM?Co}@$T(+mJe`GvhAGe7JH_8-`UNx zdCZT++|(x1=LF;<((HH(Y9(KSqz?ILG%%6|Cd`le`@2&>r`F z)9z6oeq-AL?wS!VSn2wDWa{;NZm&!9Hz*&LVqSMmooCnQR(^gCq!x1De|Ms#$8ayC z(M!0O(6$f0yi`5@BF2e_0`fRNH6t6=2idQgbbUSY%cCu{%u>6)4adJZ&i(bL>)G79 zEIo!}o_WK2TK5XuW5ky=^6N@Gg*Z`sUUDzAWneM*^G7rFkUO&Yg@^R~z+(Z9eJP7- zVV#lhk@7g_^3_pYUT|$d9`olfe&hb|^Zo$GI@+>@Zq>fD&(wog8<~%(0JBfQ!p4ELh zk63yG#FgC~$@A{*OuJ?up-m=Z+9y{Mi2H^ztmUb>eBzGyeHR-G}p>y zRtMxU|Jh`|9;>)Hz$235NPWyZ(SA}-R4Ue-tTmhOv*aCoM|_DvPr5l$**_Y2!*L(z zW9jbwgDV8(<*&VKFm5$+S@jm*McyiU$UTutJCmgc@zp2lo(iPV}6Ntfb|IT=YzWA zJVM7zd2+-yg+%g-@CK>Wpl3Ql6#nk>Fo9xeSbg?_N%9V zgE+m<1vu^z_Hw!-*3sXfo~Tq#4D0+dY>(tl;_BXmHD`(~pc0z3@6R21i%8V7hp zUin}M&TDK4@TlCo^JPt($o9f?oy5?*G~Kbpyl>M^4L5BL=YaagxlKFK-jd@Sak=^K z-0OC_zP#K~vmyL9t}9VqA^kMCEV_-mju!cOU~ghVL#-?0LOF9-GX`7 zQvr^+Sfwr^ZtzPMkNWvfI#2=gnRY%6$G*ynyVcug>-Xj5+{O+VKRy=FgLsB958__+ zo&oL&!kN2$)^MKCJg?D%am1MG@N?MS$s;ne?=-#9+b(A#AC(gIYU?9|OnWgOBcI>@ zp4#qx><3Q7cfJ?H&3@W@5%-4U)+~9c*f^K_W2yN8d7LX;2>&GYOn`gx<+oKe)W~k9 zMt>OV9nb3c=^q5-BjWV94d*>)T09Ql9jEoLlQ17kJ7)4@IM&~*SHV6#`k&+u?j=rwPt3r4^@)If>?^yL3;Ts^`(X5k#hLZIet4Sp&FE=`A9MctoEF%=%F<&v z_K%JFR%d-{Kd&z@E$?2eaXamLWw^+;4l}TC^FXGa!}x#3r}&aoopHX9X1o`26R))c z;uX&h%1d0`6rI)P#{kFvrKj)FV{2(ovh-m5RR-fJzDM$~jOd!;woCjhpa=V$?|NUW zdshQID(Ms6)tK9I1@dX)UAQfMXnY+;M+x zGd#dW3J?8MA74Rx>*wF-K^zr<|5i32AC{YKFC#uf*8oSHhvNm^@8`7+@ThEBx>i@! z*cRaEq;@51rE{At?jc>QqrKY%>lYf}UTEI2E|`yHKW_5r%gfa(PwTK$+DUSeQ*Yku zmMKmDjU4B5_b+m{)#??{k9+GLKd(32P73gd#P^?!xNt#^b1Pft=s%{j-!b|le*PD| zKMChC?Q%05=knG*&`^Iy-?e=h-LN4+|o9;N*`HGbvqu-NCh3DY>-A4hAcsL0q+%3~@ske|Z_rHwN zQ61P{8~ri(!M=+Ucf0BN=|_=c{=C&Gjn8G*H(ws{xJuz1%_1vZ!_ltSmUO2-`cHb~ z_g(k8oxgoCARm@5roX1!=Gof&4s{VtOq!H?q%7&cySd`r7WZtgp1gZkOW z((lX5uH7Zvq^$7><+1Oj%584>_;5hplhf%#^?`&v7VqJwTWwS~eLf%O0gN8QaUWv6 zdAP6bpX4zvOU6F`jhT9A*B$3d+9DnN+{pmPI+Q+!a}ylbWXa2pVS9CGb+%J-tiRlO zM*Dp>I;cmwzC8!$@#%L_-pseO#ChtG0r`l``Ta4({a}5gd{oL$EUI0y;`Wd`i6wVT z)0^+{eHkud-b?Yk8+FpBoZleFea9cPclS2sdm(ou-;d*wANnWcB9C<5tw)?$Rz8h> zPr5GIt39qS4{+?;d4G!5{)e4^BX8~>tc3X-+CA!z%7LCm+-T)7mYy)=(R+T<->urN zA-MS%dBeTXmDl#e|6doBm-W9-z3Q8Z{~R2U_YlwUFOB_sWq^m}_4!-0 z$UfF{>WN5=B{A41**_p3#kuKpoh0n9jeHdObe!X5>%k5;v*ZoOdABb-o%zd$0gm$A~g)+I;il{CBD;=mne?V(os)Fj$-qsMSB)T{dn zefX^#{`ZXf-TNvPIFPjEqpIBOehvh(OPWSST=>d*& zRYNE0@8h{1#B`(b&Ar=jzK{Km;hxElv*hw!ZtuEJTY3!lLKSnK)&1{<1Kg4Bd-v$! zYP6HogL>F{l_urh5Rmtz$;aj0q01WvIPSx0dsBQ|CBVJVqTaW=r^a%8Wb*0DOY#dJy9=JR%fWDwv7KJSeXKo$dZcTE zbd2k1C#feaL#CE+_uWVTlRP44F&`BzVZV2yKPtEHcv^q>Ib!K|@E!NAj?Cj;+(19y z=rP>f=dxcDs=phQmx~3?>91|*S5qG4KfI26>Cb5ad9xp~B=-N?`84_w7d>`|&g(;e zk9s0<|6gPD>^Qb}!(;H{IE_AUo1G72)rF>YfpMOV(-NtnU!^O|f#ZTkhQm=8|gB{kp{@Kjm$&t^4x49p8rQb{LU|;vWi0cxR$GZQhIJfMnXi$$N z%zax6-p=_(>IqA>X>Aa9H!dI_k;8?a*VJ;`13Zd_-dH{PPT1li=Dn{fSJ_?mEB!sA z-*DXD@k|b!|K)lRxg%$P?~VOPwtgcoxc~YZ&Xv&4P~MX>@2)^R0op6_usplG9@drU zr;$fw#qdutU$HEZ&!{96ei`!>^m~nb80lhvd-w8gxi*~tC-*`#iez^i)ueqOhd=Z4 z0d4;C4okn$gSZ4Q?bg>ja-E3so|N60+ikgFazGyYuZ}FkK7=m=9P6qF=3qR+dPO}r z|5dt!CJms!VR+Qf|7RVabibWMdr0nu=8j&6I6{Xl{l2_}#}#+4ucCdXJp9Rv`*F_k zy?{K{yX(!>*Y|!M;9+@n|FhT+&v^{$iQqiKH#&Vk+nwS0@q_c=n#ax8;3Z3s;bxyk zC5)Tq2Dl?1)vSR1#t!Pv^z>Ykoj~L>j-> zPP@d=?=teHEjUim_kY&gn$tfa_d)|JZ^ya#Y-;lCIOL9We&DQrx7$wF=n<(`|3!^9 zpdUziPoCNRn&xY8#L7=AKix51W?+2He#FQdj`({sXY1f!zw%`=()H!#xBfmCV1x`O0ZcaW4b&Pws`{lW%jMy8Wk2J>(AVyD#HjO`e9!m`$2N-`>p?$etp5{Hy?7WS0)t4{hp5o z^oMcpMjxExt`p!9?7!`;d4J-1EA>RB!Ly!zKbiiv;hvb3&K4_(bGuJidJOkMZy!^) z!=o<*xFb!^Z^r%#)+_1}dGYjIt@C0eAn!@x*oP1=v|oVZT=SiSFyBkNK|P4KIeWXN zl;k`ZInHDLm8MJ1*!m-W{*O&thkI4*dS$p5`lapyJ=ucw+34}*WpD0~n|H(jD?f&d zeD>=N>?bk(i1|?7lPNnoYKM#V`!e!~@09w4u8Fe$rF=vt3~Zwl2HMYM6ZHAr*$J~qdvcp?b_&Z{JbR;E`a@|>Hhc3`e6P|&K=I~YLATpF7o&jtM$PtY^Paz z5a04uZSvF?0eQ?P|8+q76lK3d`H0LbJXTM3H{&{EPYjR2kK=s&Wf$n@IF}s#(eWvI z`uw*+dC8Thlw0dgyBv%j+&dXc)rEcOZ)E8~`R~K|FZxU5VR`y&2fdnT*IV)PbM;70 zclM@;m9F9NKQB+$&b^IYV*Y)3dH2J9+Ws#)|Avb^P+_Z{oyPq9n)!E*({$fD{jdwm zfgJIUUTNk2bdCKQxg*EE7yy5g>ptWd_Z3@$d$D(C>Ltg0g(qeqUL5TMd02{X$b!QCU&Cw(hUW{>bPNG4E^JhubuE!ahshaGV=0cm{DT=-+0^%j#a1;0S4V=ri%hSlif9umvM~>YSkjHxQ z+KbwA^MU{m%cKHn7{AcpV16R9ul!8yza_|{viCr5+~>l6-^fS({1-nIaz~DxW9c{C z3*DDvsh*j`_L3zpg=gk*NB3aALoTu~$7aO);`$xACv_1Yer!j(+>HLP{)Ed8}5ZVCpqpnjcq-?yqrqN>2@8$c5U)wxFGJt z^I9^+_y?xlk$ckgWJJ^cH0{fL$gv-O^;C^}_-ji(&X-C5dy&5I<26g(aKwon@wwiA z@gG5X`MuQr?)=Sf1~}HCT1~)xk@oZYdgRL$XK>!%t|zVhw8CFJs`+-&KAU_R?uAwt zng%<6!qV@{%K&MKcwDUij%GZc__+5P1VEk2_Z z{@#=uxk!;*y|nm}n*n*8w~aik8%NU4Q$8%qQ%h_4c=pTW@F%Md!uoIafc~iLD*e9R zXq4OLVxHk2Cn7KJd|o@Q;(H;F%J}KAx+j6_xrTetk2sclmcu_`yCBDX zD_4(VUmxulxg$$1f2eo#qrXj#dDO+ph$p|v%D>T%d5IqDH9h~i01wN``XMcyo$~|K z6Os3SYN`e2(%zCsCHDi#`eIYgI~gAF^WW|FyzcB6+ADG|)KZ?+>qR&&BX{KP(_>Vl zhpqe=JtE^vcEElW+FQzF+&ZVP9)FSjK6zNSm8*sGiFW>to(T5u)zwXS!jPJfn)1h1K^7r-N-o|xUN2Z-*x?$P*kHN6>XDvN~?}$gSal5v?X4?V75l`l~ zBrSK8^C3o$FOU0zH|fB7^oz)`PE_%6ZP1_Xh~pmAf44K%<06^ql6#>gEvM>9{q+*{K`#W;PGaPVKKgf!QLmrXNy^m;@Fvs=e zQQR}~5cW;IV&zlBl+V3u(zNw!d|o4OxEC6F>=n(op8Y1dBgLC1>&g~y1oYryg9vG|{~e+4(U%ELU=FfSqAGrGBgrc23t;dFe-yMkj-tUu@n32*j|jh;n*ir?Sx+Xll7L|krh``wAG*N z@5n{Mt2ZKDyPo)Zq)4YvaW3gbAYHVJoRiUi(QZ&rL=xH$L4IiG4bP9i#qLk^kM~Uf z(b8i${DJt7^n52fU0+_(PA|YdXZi`$BeEd5t5&>jKd-MxzKflrZ6E(Gpg$~+*JuE{ zNBeB_#QFIdSF16Db%_DNO zR5qRT2-_+3N9ChFWw1_q!PX<@38H_dB2GKUalSmb7iwH^x*o{?Lr`9NR!zeE$(8^Y z8Qo(l=I8ldvh+x5-=Ua~WqTnH%jU&xa9{SufS!n)+nkJXF6#s3qw-a>7W}K%EP2n* z|G6WJalaVtypcED3$?p#p5-T#J2Ej_n%4f6{R_EBOwU!==WscoA8`v~CL^u_?GNR{ zSTC5Q^$yZrkw@gj#*u;NjrjSje^&wQ%VocBOovzwBH=5LHUSme!UXnRI^_&@(!Neaq2hOhx@Aj zu*$)3)Q797n4h%sWIRkq766ao;Z69n%dQ>6FRZqG5&C+UHe#_KD z?uFX4e-!y=y&{LQ$L8v$oao`p?vZ^Xa|Y>&kVs^M;2Z?>Ifn zrec5T7nXj*z0kQU>##3{{Zy8`yt8?+-k3talw9P|_!lwH{C!Z5{8e)p#ycFBQ6BM% zn{>my4{X2WSOS2qKSi(1KOfNV$)ghoYji!=Nhy!Ge3zT!d=u+4c|?kRIz&sS(H}58KXf`y&8H{p zcOmw7>oUg{l?b=^61Cg71xG)=}%JL zlYF>0d_;Y=Bl57^eW0v19%<^W`H)9s(&c9Q$;W>M^hf3097B;GV|PqB`0`@@GUv=0 zy0kg%8@U%M-h7I>CFsAAJCeUn2yxr!Cy)!ygRjE8G3zHe;u6e6y(+8mp}?+Dr%a>NrT+(>IJV0$DFOQceL-92-?RcWi9_M{cw}6(IYZDW(exT zE$zbmE$&r3h;{A50XyZUT7wCm#zAws$ajtwN_P4M;kVhnPs*3J! zyUNOE*w1Id+k4`kIKE#aZ#epq7V|K!r2jzf$b~5{px-zCnE8<7+{$+eT53MW1>~Mo zdFMIIC$QdNwvjj73+*nMs%L+v zJtTKz^@_pz@!xhmG4dh~BAr~hm-V0Wp5$8R>f4WR4Wt{E?{=r?m=YWx8TtJFvlrO= zmR@+3_JOL!7qu-My!;)}6k)5A7 z|8HN|)meAvGyOjG7>;?oH)3&rsQrF@dFi~OKK83{JZRE2Tx7uIj<|=8_542YjL(i9 zGxftoRq(rMZ{%3-o2Yts)0_bgf3kID%p0?PQC{SY`wL;e2irTjC($#hntdzh+sVTs zZKmie*^OPW(}kZFQ$p_@$olWgi^=I{uMO7YTe)6Ij`K%v)WSW&8~rlL%)gPB&$qv* z|bSjA;-Rvr(VMPFZ~2^Pe$yVjCjRt7vy32 z;O=fZJG|Su>G(XxZ_KI>{`c>b zdckkB_51R2sBfx1_{%Nt7xE%W7iQ}G&GhR`K8?I5x!Zk+bz%A&hR1*-^2L(I*dN7y zpWF+Te6AkO%h>PNmzSlVRL1y@ek|oh4mFyp@8+=ei?9FEpc>f!xyIxH?ThnLj`Lf~ z0obQ>IKaKooe9sVHemh9l9#-j+F>6B?FPBX%6DpGoqAnRkKA!@X^blu1bA4EbuNK; zQ|SR7kvqpG;@mpNaYn!3eK}5%^D}Xub<~nK+zX9;sjar#$oH5fFBi@)#(0YLgdF!S zOf(U@=HdpA7l z=V#2Ke%i4D`)hJ9RH8(xw%kDfliZOnil=G!)pkCO9>h&8I8}F+XMLtT;xe6k4)fKt z59DD)qn^a|# zep{IHIh6Ng^Q8*d2eQ(VkMq;5QEIF{x5q9Y!@bbbyzTVfKO#YS`R;E)JOsAK800hK zbFlw3eJiCjemCui+zZuNuDYNW`zdk<`)bDP(m(GB$cr@ptPbX}uVto7d7QT`+(}o~ z&k>Lh%jqH^^w-w|@(~$vs;(BjgZ`ecM@&v@JvUG%3Co8Z{?*XQy5HsdBFDJzLK{7k zjrEWmAx2C}VeACVk(hbXf$2etHO`TDM?aOe_&*!Um zb;0?Iy_O!sz0ez9KZ5zCcLUs!q1Ps3y@2(Wdc=ju{^D^W%7rBHov2U*guAlX@%g2|OQvH&#ei4*M zJc{IYdU#7v-jl;gNwCjvS^13d(=D^~0qyn{?VHJ`;ncy--dhwFFbj&!KhUk~-A zKTj_5?(wdgXCj~L0q8O9#&O1NZ==u84012j^yRLIS3`T1B`@zifVc!Fe#y)yIo6d< z*3~&L(f=pMdFK^%)m^_aARm^3eMEmO6_m$)E??Kdd>8vKqesN#a|z}r4?RhLfZPk+ z(`>MQIE3|$9OvtQe+2Worku@(T(A#$kd8e;|A8FwKnB#sJoY9lpC(<*AD-%}G2gI$ z8u>6d{H4!E<9r$YMsl2|_+Yudvh}j1$CsCtGplQ5*DeRc(Z7sopWM4g2DlfRe+2hDrg6Mu^!xINpVI{U^KJcxi%k5v zo*sRh<1)%)9{3;S_35jxSo-thJMMw)F%bPD?S_#z+zX9c*-+nj%-ALKA$JhpDG7c8 z?KwH_Lu}sw@r!I8=l^ZK%58K}QpDHkI2E8j&OvraKQ+j^AP(xa@S)Y|aDIbaWL(bHh*w5`o7|I0pVrfk``LcU!{U}GiSxM&t#rkIU%f(Q z5FeiXi;*|n3mxm6f_TEFo|t@-Qg=aGam|t0DIPuzx3aWaUB8 zyUw%yk|R#LTS>oa$#Eh%;?H+zpi8=OTxocQUF|ukE6%6b^&!Kqp8BnvZpmZk)0dZG z2+=e6PU?w+o{UeA>oqV>ek1d_$PxGN;r58vYUYE?`z3dzK+Wn{U!Z?Wj`-C*lhOa) z5s>$!eT%~A@7S&>AC}ve+=2TqSbuzua&w&NZ;sb%6v{CpUTw_d*qR z_10qx=(m$Qa`LO8ursWm zrQhg*KaZg0;@?>4pgZ((U7s_6? zzFw|Pe}LSPY4ua|mFf|z9E={sfvH#r@7;dh4EwfVUrX(9iv5VulVRV6E^ncYvhrCQ zJmXXD+lD%;5XaZ#Xs6?Q;5??nulx$@*;dTmddJv(ass?`1hq9`mT{9>#hl z?Id|ve(dJxKT7V%d@i5kxp0qUc&;8Rb~V7g(D{KQ_39qlhb(zHJWsLT=9YA!N96p; z3h>Wa{*=eQ-y9F1ebLSv?)m8+F7zDs{eEYqYq%GRojncnMVA8Hk(cXs#C{&uL+Zi2 z=rl7|vnU`BKfhlb_8D@1z{p4ZbpPHw5$iadPb0^CORp~2Uq?Ta+>v(Z-|JA_(dw}$PA+{DAFt-0ReIG5bPIm&yn?mRCu|Mj3J!=9X8HBv{9 z_&UJNzVRV?|4hD@EO}YhudB9F<K*MQuSJqeGXebH`{iwui<0Q2X3@8lltr!I%}HQFn~!+yG7Kh#{8XZgwG zUTE7JrQr82w#vbmm*dwGwWn+6({Pb!<%cjov)Yo6`uayyknjtQFo&Xm) zaoatJdqz8-rAK}r@;B}fqMzw=cwjzz8URUmOSpsYJz)AZplCN z;9Sqeyzt*<+s_4#eHRCF<6O^li+kXX^ZWRgdipr+q0w)+7y4rGaJ_Gyoqu0mn$PQl z^I}0RQn~8GTD%nd1*0e8>yPjB2+kpLUVH{;&n=A)_-zGF4wN2Q}5+G6*=0`df2^# zv~Pw-ef?F-Rm6I<{al87p_S_jtM|s7%yLVHe8y*9pT>Il9?naUd!gjRxp9vh{S$Kd zAwTrjrY8&H@6CH67pY#L0Q{v}BzhI-pO$fa;_S> zYa{(PaJv+A@^0ECaz}oPsi%`Ru$_{N6nwA( z)|+T=4UhQxFI6mo{a_2Nax)xp0G@85OW%GwlQ)BWM!lWfyd~}#WV5fu-Uc<%DPvOIHdi!&%59D5GR__M5zuwM|FE2OUBCsbf zS@|?vBzyITu#UibV)TT4{j(k@q9=yWx8x0nKUTdm&Oxx8v*fXF>TdmbDD71%=*jrR z{ZSV6=le`O!~+#aC}7Wg<_Uh(!`(G zuE`yFy<<-O^S$(11 zU0R2C&#+hR8@6r})4XlxI_=|Q<758)PurBX%^LpuYO*x6a9cLxFE#lkrh2VfHTX*$ zzx?>1~^^u<@J5t~pg-u!L#xY&eRwamAw)ncnA_$(nVHVz-Kg!tI%)vEd5vKVmI zr%lV&ZOt!6BKVUf^uPGhrc?8j7LD)+)sdeV*Z<3`N6pYv-i?XJzr~l2Y1O#uzfui- z&0{)N>Qt%ozmoi){kyNPvQb$1|3qU({wl^-samO%|KER8Yi9l|h5ueYTQ_XeGRA-N zRpV>cjISM6BO$Y(;$q`#)~Hsyc3eX3>ec_FYW|C2Y@gVyc}$DM7XL!*fBesA9_^a9 zh)MqM{~+U^E7z#uS8G$XOQengaD@wc-C<ajJdSFe$fP@`Hk^Lu>tS_!eWk=^RGYgesp2=mt_rSX5K z@)0A8PxaXN+6f8O5)x` set: + """Rows of an index tensor as a set of int tuples.""" + return {tuple(int(v) for v in row) for row in indices.cpu().numpy()} + + +def _current_edges(restraints) -> dict: + """``{edge type: {origin: set of tuples}}`` from the existing restraint storage.""" + out = {} + for rtype in ("bond", "angle", "torsion"): + out[rtype] = {} + for origin in restraints.restraints[rtype].keys(): + if origin == "all": + continue + group = restraints.restraints[rtype][origin] + if group is None or group.get("indices") is None: + continue + out[rtype][origin] = _tuples(group["indices"]) + + out["chiral"] = {} + chiral = restraints.restraints.get("chiral") + if chiral is not None and chiral.get("indices") is not None: + out["chiral"]["intra"] = _tuples(chiral["indices"]) + + out["plane"] = {} + for key in restraints.restraints["plane"].keys(): + group = restraints.restraints["plane"][key] + if group is None or group.get("indices") is None: + continue + out["plane"][key] = _tuples(group["indices"]) + return out + + +@pytest.fixture(scope="module") +def built(request, pdb_dir): + """``(topology, current edge sets)`` for one structure, built once per module.""" + cache = {} + + def _build(code): + if code not in cache: + path = pdb_dir / f"{code}.pdb" + if not path.exists(): + pytest.skip(f"{code}.pdb not bundled") + model = Model(verbose=0) + model.load_pdb(str(path)) + model.set_restraints_cif(None) + restraints = model.restraints + topology = build_topology( + model.pdb, + restraints.cif_dict, + link_dict=getattr(restraints, "link_dict", None), + link_list=getattr(restraints, "link_list", None), + links=restraints.links, + xyz=model.xyz().detach(), + verbose=0, + ) + cache[code] = (topology, _current_edges(restraints)) + return cache[code] + + return _build + + +@pytest.mark.unit +@pytest.mark.parametrize("code", STRUCTURES) +@pytest.mark.parametrize("edge_type", ["bond", "angle", "torsion", "chiral"]) +def test_edge_sets_match_builders(built, code, edge_type): + """Every origin of every edge type holds exactly the builders' index tuples.""" + topology, current = built(code) + block = topology.edge_block(edge_type) + graph = {origin: block.tuple_set(origin) for origin in block.origins()} + expected = current.get(edge_type, {}) + + assert set(graph) == set(expected), ( + f"{code} {edge_type}: origins differ -- " + f"graph {sorted(graph)} vs builders {sorted(expected)}" + ) + for origin in sorted(expected): + missing = expected[origin] - graph[origin] + extra = graph[origin] - expected[origin] + assert not missing and not extra, ( + f"{code} {edge_type}/{origin}: {len(missing)} edges the builders " + f"produced are absent from the graph, {len(extra)} are only in the graph" + ) + + +@pytest.mark.unit +@pytest.mark.parametrize("code", STRUCTURES) +def test_plane_sets_match_builders(built, code): + """Planes match per atom count, pooling the origins the graph splits them into.""" + topology, current = built(code) + graph = {} + for size, block in topology.atoms.planes.items(): + graph.setdefault(f"{size}_atoms", set()).update(block.tuple_set()) + expected = current["plane"] + + assert set(graph) == set(expected), ( + f"{code} planes: size groups differ -- " + f"graph {sorted(graph)} vs builders {sorted(expected)}" + ) + for key in sorted(expected): + assert graph[key] == expected[key], ( + f"{code} plane/{key}: " + f"{len(expected[key] - graph[key])} missing, " + f"{len(graph[key] - expected[key])} extra" + ) + + +@pytest.mark.unit +@pytest.mark.parametrize("code", STRUCTURES) +def test_exclusions_reproduce_current_set(built, code): + """The restraint-edge exclusions equal what the non-bonded term is given today.""" + topology, _ = built(code) + from_edges = topology.atoms.exclusions_from_restraint_edges() + + expected = set() + for edge_type, cols in (("bond", (0, 1)), ("angle", (0, 2)), ("torsion", (0, 3))): + block = topology.edge_block(edge_type) + for row in block.indices.cpu().numpy(): + a, b = int(row[cols[0]]), int(row[cols[1]]) + if a != b: + expected.add((min(a, b), max(a, b))) + + assert from_edges == expected + + +@pytest.mark.unit +@pytest.mark.parametrize("code", STRUCTURES) +def test_connectivity_exclusions_are_a_superset(built, code): + """Connectivity-derived exclusions cover the restraint-derived ones, and then some. + + The difference is the defect the connectivity path fixes: a pair that is 1-3 or 1-4 + bonded but whose angle or torsion the monomer library does not restrain is currently + not excluded, so the non-bonded term pushes it apart. + """ + topology, _ = built(code) + from_edges = topology.atoms.exclusions_from_restraint_edges() + from_bonds = topology.atoms.exclusions_12_13_14() + + assert from_edges <= from_bonds, ( + f"{code}: {len(from_edges - from_bonds)} restraint-derived exclusions are " + f"not reachable within three bonds, which should be impossible" + ) + + +@pytest.mark.unit +@pytest.mark.parametrize("code", ["7L84", "3VRJ"]) +def test_adjacency_matches_bond_block(built, code): + """Every bond appears in both atoms' neighbour lists, and nothing else does.""" + topology, _ = built(code) + atoms = topology.atoms + + from_adjacency = set() + for i in range(atoms.n_atoms): + for j in atoms.neighbors(i).cpu().tolist(): + from_adjacency.add((min(i, j), max(i, j))) + + from_block = set() + for a, b in atoms.bonds.indices.cpu().numpy(): + from_block.add((min(int(a), int(b)), max(int(a), int(b)))) + + assert from_adjacency == from_block + assert int(atoms.degree().sum()) == 2 * atoms.bonds.n_edges + + +@pytest.mark.unit +def test_layout_is_reproducible(built, pdb_dir): + """Two builds of one structure lay the blocks out identically. + + The canonical order removes the process-to-process variation the old storage had, + where origins were concatenated in ``set`` iteration order. + """ + topology_a, _ = built("7L84") + + model = Model(verbose=0) + model.load_pdb(str(pdb_dir / "7L84.pdb")) + model.set_restraints_cif(None) + restraints = model.restraints + topology_b = build_topology( + model.pdb, + restraints.cif_dict, + link_dict=getattr(restraints, "link_dict", None), + link_list=getattr(restraints, "link_list", None), + links=restraints.links, + xyz=model.xyz().detach(), + verbose=0, + ) + + for edge_type in ("bond", "angle", "torsion", "chiral"): + a = topology_a.edge_block(edge_type) + b = topology_b.edge_block(edge_type) + assert a.origin_bounds == b.origin_bounds, edge_type + assert torch.equal(a.indices, b.indices), edge_type + + +@pytest.mark.unit +def test_origin_slices_share_storage(built): + """A per-origin view is a slice of the block, not a copy.""" + topology, _ = built("7L84") + block = topology.atoms.bonds + origin = block.origins()[0] + view = block.origin(origin) + + assert view.data_ptr() == block.indices.data_ptr() + saved = int(block.indices[0, 0]) + block.indices[0, 0] = saved + 1000 + assert int(view[0, 0]) == saved + 1000 + block.indices[0, 0] = saved + + +@pytest.mark.unit +def test_residue_identity_includes_insertion_code(built): + """Residue nodes are keyed on ``(chain, resseq, icode)``. + + The restraint builders group on ``(chain, resseq)`` alone, which merges residue 100 + with 100A and leaves the inserted residue without intra-residue restraints. + """ + topology, _ = built("7L84") + residues = topology.residues + keys = [residues.key(i) for i in range(residues.n_residues)] + + assert len(keys) == len(set(keys)), "residue identity is not unique" + assert all(len(k) == 3 for k in keys) + + +@pytest.mark.unit +def test_residue_join_recovers_per_atom_identity(built): + """Per-atom residue name is recovered via ``residue_of``, not stored per atom.""" + topology, _ = built("7L84") + + for atom in (0, topology.n_atoms // 2, topology.n_atoms - 1): + residue = topology.residue_of_atom(atom) + assert topology.resname_of_atom(atom) == topology.residues.resname[residue] + assert ( + topology.residues.atom_start[residue] + <= atom + < topology.residues.atom_end[residue] + ) + + +@pytest.mark.unit +def test_connectivity_exclusions_match_brute_force(): + """The vectorised path-walk agrees with a plain breadth-first reference. + + ``test_connectivity_exclusions_are_a_superset`` alone would also pass for an + implementation that returned every pair, so the walk is pinned against an + independent one on a graph small enough to enumerate: a six-ring with two + substituents, which exercises the ring closure and the 1-4 wrap-around. + """ + import numpy as np + + from torchref.topology import EdgeBlock + from torchref.topology.atom_graph import AtomGraph + + bonds = [ + (0, 1), + (1, 2), + (2, 3), + (3, 4), + (4, 5), + (5, 0), # six-ring + (0, 6), # substituent on 0 + (6, 7), # and one more out + ] + n = 8 + + def block(rows, arity, edge_type): + return EdgeBlock.from_origins( + {"intra": np.asarray(rows, dtype=np.int64).reshape(-1, arity)}, + arity, + edge_type, + ) + + graph = AtomGraph( + name=np.array([f"A{i}" for i in range(n)]), + element=np.array(["C"] * n), + altloc=np.array([" "] * n), + residue_of=torch.zeros(n, dtype=torch.int64), + bonds=block(bonds, 2, "bond"), + angles=EdgeBlock.empty(3), + torsions=EdgeBlock.empty(4), + chirals=EdgeBlock.empty(4), + ) + + adjacency = {i: set() for i in range(n)} + for a, b in bonds: + adjacency[a].add(b) + adjacency[b].add(a) + + expected = set() + for start in range(n): + frontier = {start} + seen = {start} + for _ in range(3): + frontier = {j for i in frontier for j in adjacency[i]} - seen + seen |= frontier + for other in frontier: + expected.add((min(start, other), max(start, other))) + + assert graph.exclusions_12_13_14() == expected diff --git a/torchref/topology/__init__.py b/torchref/topology/__init__.py new file mode 100644 index 00000000..ca8e9405 --- /dev/null +++ b/torchref/topology/__init__.py @@ -0,0 +1,31 @@ +"""Model topology as a graph: residues over atoms, connectivity over restraints. + +:class:`Topology` holds two levels. :class:`ResidueGraph` is the sequence -- residues as +template instances, inter-residue links as edges. :class:`AtomGraph` is the expansion -- +atoms as nodes, typed :class:`EdgeBlock` sets over them, and a CSR bond adjacency that +answers ``neighbors(i)``. + +The topology is **target-free**: it says what is connected, not what the ideal geometry +is. Ideal values and sigmas belong to a restraint layer keyed to the same edges, so one +connectivity can carry monomer-library targets, force-field parameters, or +ADP-similarity sigmas without duplicating the edges. + +Build one with :func:`build_topology`. +""" + +from .atom_graph import AtomGraph +from .build import build_topology +from .edges import ORIGIN_ORDER, EdgeBlock +from .residue_graph import ResidueGraph +from .templates import resolve_template_keys +from .topology import Topology + +__all__ = [ + "Topology", + "ResidueGraph", + "AtomGraph", + "EdgeBlock", + "ORIGIN_ORDER", + "build_topology", + "resolve_template_keys", +] diff --git a/torchref/topology/atom_graph.py b/torchref/topology/atom_graph.py new file mode 100644 index 00000000..55af8f8e --- /dev/null +++ b/torchref/topology/atom_graph.py @@ -0,0 +1,284 @@ +"""The atom level of a topology: atoms as nodes, typed edge blocks over them. + +Bonds are promoted to a real adjacency structure -- a CSR pair built once from the bond +block -- so :meth:`AtomGraph.neighbors` answers "what is atom *i* bonded to" without +inferring it from restraint index lists. Angles, torsions, chirals and planes stay typed +hyperedge sets read from the monomer library, because the library deliberately does not +restrain every path the bond graph implies. + +Every indexing structure here is a tensor, so it moves with ``.to(device)`` +alongside the edge blocks. Only the per-atom identifiers are NumPy, because they are +strings. +""" + +from dataclasses import dataclass, field +from typing import Dict, Optional, Set, Tuple + +import numpy as np +import torch + +from torchref.topology.edges import EdgeBlock +from torchref.utils.device_mixin import DeviceMixin + + +def _build_csr(bonds: torch.Tensor, n_atoms: int) -> Tuple[torch.Tensor, torch.Tensor]: + """Symmetric CSR adjacency from an ``(E, 2)`` bond list. + + Parameters + ---------- + bonds : torch.Tensor + Bond atom indices, shape ``(E, 2)``, dtype ``int64``. + n_atoms : int + Number of atoms, so isolated trailing atoms still get an entry. + + Returns + ------- + indptr, indices : torch.Tensor + ``indices[indptr[i]:indptr[i + 1]]`` are atom ``i``'s bonded neighbours, + ascending. Each bond contributes both directions. + """ + device = bonds.device + if bonds.numel() == 0: + return ( + torch.zeros(n_atoms + 1, dtype=torch.int64, device=device), + torch.zeros(0, dtype=torch.int64, device=device), + ) + + src = torch.cat([bonds[:, 0], bonds[:, 1]]) + dst = torch.cat([bonds[:, 1], bonds[:, 0]]) + + # Sort by (src, dst). Two stable passes, least significant first, give the same + # order as a lexicographic sort without materialising a composite key. + order = torch.argsort(dst, stable=True) + src, dst = src[order], dst[order] + order = torch.argsort(src, stable=True) + src, dst = src[order], dst[order] + + counts = torch.bincount(src, minlength=n_atoms) + indptr = torch.zeros(n_atoms + 1, dtype=torch.int64, device=device) + torch.cumsum(counts, dim=0, out=indptr[1:]) + return indptr, dst.to(torch.int64) + + +def _extend_paths( + indptr: torch.Tensor, indices: torch.Tensor, paths: torch.Tensor +) -> torch.Tensor: + """Extend each bonded path by one bonded step, without doubling back. + + Parameters + ---------- + indptr, indices : torch.Tensor + CSR adjacency. + paths : torch.Tensor + Existing paths, shape ``(P, L)`` with ``L >= 2``, each row a chain of bonded + atoms. + + Returns + ------- + torch.Tensor + Shape ``(P', L + 1)``. A path is extended by every bonded neighbour of its last + atom except the one it just came from, so ``(i, j, k)`` never yields ``k = i``. + """ + device = paths.device + if paths.numel() == 0: + return torch.zeros((0, paths.shape[1] + 1), dtype=torch.int64, device=device) + + last, prev = paths[:, -1], paths[:, -2] + counts = indptr[last + 1] - indptr[last] + total = int(counts.sum()) + if total == 0: + return torch.zeros((0, paths.shape[1] + 1), dtype=torch.int64, device=device) + + row = torch.repeat_interleave(torch.arange(len(paths), device=device), counts) + # Offset of each slot within its own neighbour list. + exclusive = torch.cumsum(counts, dim=0) - counts + pos = torch.arange(total, device=device) - torch.repeat_interleave( + exclusive, counts + ) + nxt = indices[torch.repeat_interleave(indptr[last], counts) + pos] + + keep = nxt != prev[row] + return torch.cat([paths[row][keep], nxt[keep, None]], dim=1) + + +@dataclass(eq=False, repr=False) +class AtomGraph(DeviceMixin): + """Atoms as nodes, typed edges over them, with bond adjacency. + + Parameters + ---------- + name, element, altloc : numpy.ndarray + Per-atom identifiers, shape ``(N,)``. Strings, so NumPy rather than tensors; + residue-level identity is reached through ``residue_of`` rather than duplicated + here. + residue_of : torch.Tensor + Residue index per atom, shape ``(N,)``, dtype ``int64``. + bonds, angles, torsions, chirals : EdgeBlock + Typed edge blocks. ``bonds`` also backs the adjacency. + planes : dict + ``{n_atoms_in_plane: EdgeBlock}`` -- planes are ragged, so they are grouped by + atom count the way the plane restraints already are. + + Notes + ----- + Holds no refinable parameters, so this is a dataclass rather than an ``nn.Module``. + The adjacency is derived from ``bonds`` at construction and rebuilt by + :meth:`rebuild_adjacency` if the bond block is replaced. + """ + + name: np.ndarray + element: np.ndarray + altloc: np.ndarray + residue_of: torch.Tensor + bonds: EdgeBlock + angles: EdgeBlock + torsions: EdgeBlock + chirals: EdgeBlock + planes: Dict[int, EdgeBlock] = field(default_factory=dict) + + _adj_indptr: Optional[torch.Tensor] = field(default=None, repr=False) + _adj_indices: Optional[torch.Tensor] = field(default=None, repr=False) + + def __post_init__(self) -> None: + if self._adj_indptr is None: + self.rebuild_adjacency() + + @property + def device(self) -> torch.device: + """Where the indexing tensors live. Derived from the bond block.""" + return self.bonds.indices.device + + @property + def n_atoms(self) -> int: + """Number of atom nodes.""" + return len(self.name) + + @property + def is_hydrogen(self) -> torch.Tensor: + """Boolean mask of hydrogen atoms, shape ``(N,)``.""" + flags = np.char.upper(np.char.strip(self.element.astype(str))) == "H" + return torch.as_tensor(flags, device=self.bonds.indices.device) + + def rebuild_adjacency(self) -> None: + """Rebuild the CSR adjacency from the current bond block.""" + self._adj_indptr, self._adj_indices = _build_csr( + self.bonds.indices, self.n_atoms + ) + + def neighbors(self, i: int) -> torch.Tensor: + """Atoms bonded to atom ``i``, ascending. + + Returns + ------- + torch.Tensor + Neighbour indices, a view into the adjacency, on the graph's device. + """ + return self._adj_indices[self._adj_indptr[i] : self._adj_indptr[i + 1]] + + def degree(self, i: int = None) -> torch.Tensor: + """Bonded-neighbour count, for atom ``i`` or for every atom.""" + deg = self._adj_indptr[1:] - self._adj_indptr[:-1] + return deg if i is None else deg[i] + + def _directed_bonds(self) -> torch.Tensor: + """Bonds as ``(2E, 2)`` directed pairs.""" + b = self.bonds.indices + return torch.cat([b, b.flip(1)], dim=0) + + # ------------------------------------------------------------------ + # Non-bonded exclusions + # ------------------------------------------------------------------ + + @staticmethod + def _pair_set(pairs: torch.Tensor) -> Set[Tuple[int, int]]: + """``(low, high)`` tuples of a ``(P, 2)`` index tensor, self-pairs dropped.""" + if pairs.numel() == 0: + return set() + lo = torch.minimum(pairs[:, 0], pairs[:, 1]) + hi = torch.maximum(pairs[:, 0], pairs[:, 1]) + keep = lo != hi + return set(zip(lo[keep].cpu().tolist(), hi[keep].cpu().tolist())) + + def exclusions_from_restraint_edges(self) -> Set[Tuple[int, int]]: + """1-2, 1-3 and 1-4 pairs taken from the bond, angle and torsion **edges**. + + 1-2 from every bond, 1-3 from each angle's outer pair, 1-4 from each torsion's + outer pair. Reproduces exactly the set the non-bonded term has always been + given. + + This is *not* the same as :meth:`exclusions_12_13_14`: a pair that is 1-3 bonded + but whose angle the monomer library does not restrain appears there and not + here, and so takes a repulsion it should not. Kept because switching the + non-bonded term to the connectivity-derived set changes its value and wants its + own measurement. + + Returns + ------- + set of tuple of int + ``(low, high)`` atom index pairs. + """ + excl: Set[Tuple[int, int]] = set() + for block, cols in ( + (self.bonds, (0, 1)), + (self.angles, (0, 2)), + (self.torsions, (0, 3)), + ): + if block.n_edges: + excl |= self._pair_set(block.indices[:, cols]) + return excl + + def exclusions_12_13_14(self) -> Set[Tuple[int, int]]: + """1-2, 1-3 and 1-4 pairs derived from bond **connectivity** alone. + + Walks the adjacency two and three steps out, so the result does not depend on + which angles and torsions the monomer library happens to restrain. This is the + physically correct exclusion set; :meth:`exclusions_from_restraint_edges` is the + one currently wired into the non-bonded term. + + Returns + ------- + set of tuple of int + ``(low, high)`` atom index pairs. + """ + p2 = self._directed_bonds() + p3 = _extend_paths(self._adj_indptr, self._adj_indices, p2) + p4 = _extend_paths(self._adj_indptr, self._adj_indices, p3) + return ( + self._pair_set(p2) + | self._pair_set(p3[:, (0, 2)]) + | self._pair_set(p4[:, (0, 3)]) + ) + + def hydrogen_parents(self) -> Dict[int, torch.Tensor]: + """``{hydrogen atom: heavy neighbours of its bonded parent}``. + + Taken from bond connectivity, so it does not depend on current coordinates the + way a distance criterion does. + + Returns + ------- + dict + Empty when the graph carries no hydrogens. + """ + is_h = self.is_hydrogen + out: Dict[int, torch.Tensor] = {} + for h in torch.nonzero(is_h, as_tuple=False).flatten().tolist(): + nb = self.neighbors(h) + heavy = nb[~is_h[nb]] + if heavy.numel() == 0: + continue + parent = int(heavy[0]) + parent_nb = self.neighbors(parent) + out[h] = parent_nb[~is_h[parent_nb] & (parent_nb != h)] + return out + + def __repr__(self) -> str: + return ( + f"AtomGraph(n_atoms={self.n_atoms}, bonds={self.bonds.n_edges}, " + f"angles={self.angles.n_edges}, torsions={self.torsions.n_edges}, " + f"chirals={self.chirals.n_edges}, " + f"planes={sum(b.n_edges for b in self.planes.values())})" + ) + + +__all__ = ["AtomGraph"] diff --git a/torchref/topology/build.py b/torchref/topology/build.py new file mode 100644 index 00000000..f6225f84 --- /dev/null +++ b/torchref/topology/build.py @@ -0,0 +1,693 @@ +"""Assemble a :class:`~torchref.topology.topology.Topology` from an atom table. + +Intra-residue edges are matched here, template by template, through the Numba matchers +in :mod:`torchref.restraints.builders_numba`. Inter-residue edges come from the +``InterResidue*Builder`` classes, which already encode the link geometry and are reused +rather than reimplemented. +""" + +from typing import Dict, List, Optional, Sequence, Tuple + +import numpy as np +import pandas as pd +import torch + +from torchref.restraints.builders_fast import ( + InterResidueAngleBuilder, + InterResidueBondBuilder, + InterResiduePlaneBuilder, + InterResidueTorsionBuilder, + PreprocessedCIF, +) +from torchref.restraints.builders_numba import ( + match_angles_numba, + match_bonds_numba, + match_chirals_numba, + match_torsions_numba, +) +from torchref.topology.atom_graph import AtomGraph +from torchref.topology.edges import EdgeBlock +from torchref.topology.residue_graph import ( + ResidueGraph, + build_residue_nodes, + find_disulfide_links, + find_peptide_links, +) +from torchref.topology.templates import resolve_template_keys +from torchref.topology.topology import Topology + +#: Initial size of the matcher work arrays, grown on demand. +_WORK = 64 + + +def _atom_columns(pdb: pd.DataFrame) -> Dict[str, np.ndarray]: + """Per-atom identity arrays, with altlocs normalised so blank reads as ``' '``.""" + altloc = pdb["altloc"].values.astype(str) if "altloc" in pdb.columns else None + if altloc is None: + altloc = np.full(len(pdb), " ", dtype=" List[Tuple[np.ndarray, np.ndarray]]: + """Atom name/index arrays per alternative conformation of one residue. + + A residue without altlocs yields one conformation holding all its atoms. Otherwise + one per altloc, each holding the blank-altloc atoms plus that altloc's own -- so a + restraint spanning a shared backbone and a branching side chain is emitted once per + conformer. + """ + names = cols["name"][start:end] + indices = cols["index"][start:end] + altlocs = cols["altloc"][start:end] + unique = np.unique(altlocs) + + if len(unique) == 1 and unique[0] == " ": + return [(names, indices)] + if " " in unique: + common = altlocs == " " + out = [] + for alt in unique: + if alt == " ": + continue + m = altlocs == alt + out.append( + ( + np.concatenate([names[common], names[m]]), + np.concatenate([indices[common], indices[m]]), + ) + ) + return out + return [(names[altlocs == a], indices[altlocs == a]) for a in unique] + + +def _match_intra( + cols: Dict[str, np.ndarray], + nodes: Dict[str, np.ndarray], + template_key: np.ndarray, + pp_cif: PreprocessedCIF, +) -> Dict[str, np.ndarray]: + """Intra-residue edge index arrays. + + Keyed ``bonds`` / ``angles`` / ``torsions`` / ``chirals``. + + Emitted only where every named atom of a library restraint is present in the + conformation, which is the condition the matchers apply. + """ + acc: Dict[str, List[np.ndarray]] = { + "bonds": [], + "angles": [], + "torsions": [], + "chirals": [], + } + work = {k: np.zeros(_WORK, dtype=np.int64) for k in ("i1", "i2", "i3", "i4", "per")} + work["f1"] = np.zeros(_WORK, dtype=np.float64) + work["f2"] = np.zeros(_WORK, dtype=np.float64) + size = _WORK + + for r in range(len(nodes["chain"])): + key = str(template_key[r]) + start, end = int(nodes["atom_start"][r]), int(nodes["atom_end"][r]) + needed = max( + len(pp_cif.bonds.get(key, {}).get("atom1", ())), + len(pp_cif.angles.get(key, {}).get("atom1", ())), + len(pp_cif.torsions.get(key, {}).get("atom1", ())), + len(pp_cif.chirals.get(key, {}).get("atom1", ())), + ) + if needed == 0: + continue + if needed > size: + size = needed * 2 + work = {k: np.zeros(size, dtype=v.dtype) for k, v in work.items()} + + for names, indices in _conformers(cols, start, end): + if key in pp_cif.bonds: + b = pp_cif.bonds[key] + n = match_bonds_numba( + names, + indices, + b["atom1"], + b["atom2"], + b["value"], + b["sigma"], + work["i1"], + work["i2"], + work["f1"], + work["f2"], + ) + if n: + acc["bonds"].append( + np.column_stack([work["i1"][:n].copy(), work["i2"][:n].copy()]) + ) + if key in pp_cif.angles: + a = pp_cif.angles[key] + n = match_angles_numba( + names, + indices, + a["atom1"], + a["atom2"], + a["atom3"], + a["value"], + a["sigma"], + work["i1"], + work["i2"], + work["i3"], + work["f1"], + work["f2"], + ) + if n: + acc["angles"].append( + np.column_stack( + [ + work["i1"][:n].copy(), + work["i2"][:n].copy(), + work["i3"][:n].copy(), + ] + ) + ) + if key in pp_cif.torsions: + t = pp_cif.torsions[key] + n = match_torsions_numba( + names, + indices, + t["atom1"], + t["atom2"], + t["atom3"], + t["atom4"], + t["value"], + t["sigma"], + t["period"], + work["i1"], + work["i2"], + work["i3"], + work["i4"], + work["f1"], + work["f2"], + work["per"], + ) + if n: + acc["torsions"].append( + np.column_stack( + [ + work["i1"][:n].copy(), + work["i2"][:n].copy(), + work["i3"][:n].copy(), + work["i4"][:n].copy(), + ] + ) + ) + if key in pp_cif.chirals: + c = pp_cif.chirals[key] + n = match_chirals_numba( + names, + indices, + c["center"], + c["atom1"], + c["atom2"], + c["atom3"], + c["volume_sign"], + c["sigma"], + work["i1"], + work["i2"], + work["i3"], + work["i4"], + work["f1"], + work["f2"], + ) + if n: + acc["chirals"].append( + np.column_stack( + [ + work["i1"][:n].copy(), + work["i2"][:n].copy(), + work["i3"][:n].copy(), + work["i4"][:n].copy(), + ] + ) + ) + + arity = {"bonds": 2, "angles": 3, "torsions": 4, "chirals": 4} + return { + k: (np.concatenate(v, axis=0) if v else np.zeros((0, arity[k]), dtype=np.int64)) + for k, v in acc.items() + } + + +def _match_intra_planes( + cols: Dict[str, np.ndarray], + nodes: Dict[str, np.ndarray], + template_key: np.ndarray, + pp_cif: PreprocessedCIF, +) -> Dict[int, np.ndarray]: + """Intra-residue plane index arrays grouped by how many atoms survived matching. + + Missing atoms are dropped rather than voiding the plane; a plane is kept once at + least three of its atoms are present, so its arity depends on the model. + """ + by_size: Dict[int, List[np.ndarray]] = {} + for r in range(len(nodes["chain"])): + key = str(template_key[r]) + if key not in pp_cif.planes: + continue + start, end = int(nodes["atom_start"][r]), int(nodes["atom_end"][r]) + for names, indices in _conformers(cols, start, end): + # Last-wins on a duplicate name, matching PlaneRestraintBuilder. + name_to_idx = dict(zip(names, indices)) + for plane in pp_cif.planes[key]: + present = [ + name_to_idx[nm] for nm in plane["atoms"] if nm in name_to_idx + ] + if len(present) >= 3: + by_size.setdefault(len(present), []).append( + np.asarray(present, dtype=np.int64) + ) + return {n: np.stack(rows, axis=0) for n, rows in by_size.items()} + + +def _inter_residue_edges( + pdb: pd.DataFrame, + link_dict: Optional[Dict], + verbose: int, +) -> Dict[str, Dict[str, np.ndarray]]: + """Peptide edge index arrays from the inter-residue builders. + + Reuses ``InterResidue*Builder`` rather than reimplementing the link geometry. + Returns ``{edge type: {origin: (E, k) array}}``. + """ + out: Dict[str, Dict[str, np.ndarray]] = { + "bond": {}, + "angle": {}, + "torsion": {}, + "plane": {}, + } + if not link_dict or "TRANS" not in link_dict: + return out + cpu = torch.device("cpu") + trans = link_dict["TRANS"] + ptrans = link_dict.get("PTRANS") + + bond = InterResidueBondBuilder(verbose=verbose).build( + pdb, trans, cpu, filter_atom_type="ATOM" + ) + if bond: + out["bond"]["peptide"] = bond["indices"].cpu().numpy() + + ab = InterResidueAngleBuilder(verbose=verbose) + if ptrans is not None: + non_pro = ab.build( + pdb, trans, cpu, filter_atom_type="ATOM", exclude_next_resname="PRO" + ) + pro = ab.build( + pdb, ptrans, cpu, filter_atom_type="ATOM", next_resname_filter="PRO" + ) + chunks = [r["indices"].cpu().numpy() for r in (non_pro, pro) if r] + else: + r = ab.build(pdb, trans, cpu, filter_atom_type="ATOM") + chunks = [r["indices"].cpu().numpy()] if r else [] + if chunks: + out["angle"]["peptide"] = np.concatenate(chunks, axis=0) + + tors = InterResidueTorsionBuilder(verbose=verbose).build( + pdb, trans, cpu, filter_atom_type="ATOM" + ) + if tors: + for origin in ("phi", "psi", "omega"): + if origin in tors: + out["torsion"][origin] = tors[origin]["indices"].cpu().numpy() + + planes = InterResiduePlaneBuilder(verbose=verbose).build( + pdb, trans, cpu, filter_atom_type="ATOM" + ) + if planes: + out["plane"] = { + key: data["indices"].cpu().numpy() for key, data in planes.items() + } + return out + + +def _origins( + intra: np.ndarray, + inter: Dict[str, np.ndarray], + disulfide: Optional[np.ndarray], +) -> Dict[str, np.ndarray]: + """Collect one edge type's per-origin arrays, dropping the empty ones.""" + per_origin: Dict[str, np.ndarray] = {} + if intra is not None and len(intra): + per_origin["intra"] = intra + for origin, rows in inter.items(): + if rows is not None and len(rows): + per_origin[origin] = rows + if disulfide is not None and len(disulfide): + per_origin["disulfide"] = disulfide + return per_origin + + +def _disulfide_edges( + pdb: pd.DataFrame, + nodes: Dict[str, np.ndarray], + cols: Dict[str, np.ndarray], + residue_of_row: Dict[int, int], + pairs: Sequence[Tuple[int, int]], + link_dict: Optional[Dict], + verbose: int, +) -> Dict[str, np.ndarray]: + """Bond, angle and torsion edges for the detected disulfide links. + + Drives the ``InterResidue*Builder`` disulfide paths from the residue graph's + ``disulf`` edges, so the link geometry comes from the ``disulf`` dictionary entry + rather than being restated here. + + Returns + ------- + dict + ``{'bond'|'angle'|'torsion': (E, k) array}``, omitting types with no edges. + """ + out: Dict[str, np.ndarray] = {} + if not pairs or not link_dict or "disulf" not in link_dict: + return out + + disulf = link_dict["disulf"] + bonds = disulf.get("bonds") + if bonds is None: + return out + sg_sg = bonds[(bonds["atom1"] == "SG") & (bonds["atom2"] == "SG")] + if len(sg_sg) == 0: + return out + length = float(sg_sg["value"].values[0]) + sigma = float(sg_sg["sigma"].values[0]) + + cpu = torch.device("cpu") + bond_builder = InterResidueBondBuilder(verbose=verbose) + angle_builder = InterResidueAngleBuilder(verbose=verbose) + torsion_builder = InterResidueTorsionBuilder(verbose=verbose) + + for row_a, row_b in pairs: + # The edge indices are the atom table's ``index`` column, not its row number. + bond_builder.process_disulfide_bond( + int(cols["index"][row_a]), int(cols["index"][row_b]), length, sigma + ) + res_a, res_b = residue_of_row[row_a], residue_of_row[row_b] + atoms_a = pdb.iloc[ + int(nodes["atom_start"][res_a]) : int(nodes["atom_end"][res_a]) + ] + atoms_b = pdb.iloc[ + int(nodes["atom_start"][res_b]) : int(nodes["atom_end"][res_b]) + ] + if disulf.get("angles") is not None: + angle_builder.process_disulfide_angles(atoms_a, atoms_b, disulf["angles"]) + if disulf.get("torsions") is not None: + torsion_builder.process_disulfide_torsions( + atoms_a, atoms_b, disulf["torsions"] + ) + + bond_result = bond_builder.finalize(cpu) + if bond_result: + out["bond"] = bond_result["indices"].cpu().numpy() + angle_result = angle_builder.finalize(cpu) + if angle_result: + out["angle"] = angle_result["indices"].cpu().numpy() + torsion_result = torsion_builder.finalize_disulfide(cpu) + if torsion_result: + out["torsion"] = torsion_result["indices"].cpu().numpy() + return out + + +def _link_record_edges( + pdb: pd.DataFrame, + links, + disulfide_bonds: Optional[np.ndarray], + verbose: int, +) -> Tuple[np.ndarray, List[Tuple[int, int]]]: + """Bond edges for the accepted ``LINK`` records, and the atom pairs they join. + + Atom resolution goes through ``RestraintsNew._lookup_link_atom`` so a record is + matched to a row exactly as the restraint builder matches it. A record duplicating + an auto-detected disulfide is dropped, since that link already contributed its + bond, angles and torsions. + + Returns + ------- + edges : numpy.ndarray + Shape ``(L, 2)``; empty when nothing resolved. + atom_pairs : list of tuple of int + The same pairs, for lifting to residue-level link edges. + """ + from torchref.restraints.restraints import RestraintsNew + + if links is None or len(links) == 0: + return np.zeros((0, 2), dtype=np.int64), [] + + existing = set() + if disulfide_bonds is not None: + for a, b in disulfide_bonds: + existing.add((min(int(a), int(b)), max(int(a), int(b)))) + + rows: List[Tuple[int, int]] = [] + n_unresolved = 0 + for _, link in links.iterrows(): + idx1 = RestraintsNew._lookup_link_atom( + pdb, + chainid=link["chainid1"], + resseq=int(link["resseq1"]), + icode=link["icode1"], + resname=link["resname1"], + name=link["name1"], + altloc=link["altloc1"], + ) + idx2 = RestraintsNew._lookup_link_atom( + pdb, + chainid=link["chainid2"], + resseq=int(link["resseq2"]), + icode=link["icode2"], + resname=link["resname2"], + name=link["name2"], + altloc=link["altloc2"], + ) + if idx1 is None or idx2 is None or idx1 == idx2: + n_unresolved += 1 + continue + pair = (min(idx1, idx2), max(idx1, idx2)) + if pair in existing: + continue + rows.append((idx1, idx2)) + + if verbose > 1 and n_unresolved: + print(f"{n_unresolved} LINK records did not resolve to a pair of atoms") + if not rows: + return np.zeros((0, 2), dtype=np.int64), [] + return np.asarray(rows, dtype=np.int64), rows + + +def build_topology( + pdb: pd.DataFrame, + cif_dict: Dict, + link_dict: Optional[Dict] = None, + link_list=None, + links=None, + xyz: Optional[torch.Tensor] = None, + device=None, + verbose: int = 0, +) -> Topology: + """Build a topology from an atom table and the restraint dictionaries. + + Parameters + ---------- + pdb : pandas.DataFrame + Atom table, with ``name``, ``element``, ``altloc``, ``chainid``, ``resseq``, + ``icode``, ``resname``, ``ATOM`` and ``index`` columns. + cif_dict : dict + Restraint dictionary keyed by residue name. + link_dict : dict, optional + Link-type definitions. Without it no inter-residue edges are built. + link_list : pandas.DataFrame, optional + Link table used to resolve which modifications a peptide link applies. + links : pandas.DataFrame, optional + Parsed PDB ``LINK`` records. Each record that resolves to two distinct atoms and + does not duplicate an auto-detected disulfide contributes one bond edge. + xyz : torch.Tensor, optional + Coordinates, shape ``(N, 3)``. Needed only to detect disulfide links, which are + found by SG-SG distance. + device : torch.device, optional + Where to place the edge blocks. + verbose : int, default 0 + Verbosity level. + + Returns + ------- + Topology + """ + cols = _atom_columns(pdb) + nodes = build_residue_nodes( + cols["chain"], cols["resseq"], cols["icode"], cols["resname"] + ) + n_res = len(nodes["chain"]) + + names_by_residue = [ + set(cols["name"][int(nodes["atom_start"][r]) : int(nodes["atom_end"][r])]) + for r in range(n_res) + ] + is_polymer = np.array( + [cols["record"][int(nodes["atom_start"][r])] == "ATOM" for r in range(n_res)], + dtype=bool, + ) + + polymer_nodes = {k: v[is_polymer] for k, v in nodes.items()} + polymer_map = np.nonzero(is_polymer)[0] + polymer_names = [names_by_residue[r] for r in polymer_map] + peptide_local = find_peptide_links(polymer_nodes, polymer_names) + peptide_pairs = [ + (int(polymer_map[a]), int(polymer_map[b])) for a, b in peptide_local + ] + + comp_dict, template_key = resolve_template_keys( + nodes["resname"], peptide_pairs, cif_dict, link_list, verbose=verbose + ) + pp_cif = PreprocessedCIF(comp_dict) + + intra = _match_intra(cols, nodes, template_key, pp_cif) + intra_planes = _match_intra_planes(cols, nodes, template_key, pp_cif) + inter = _inter_residue_edges(pdb, link_dict, verbose) + + residue_of_row = {} + for r in range(n_res): + for row in range(int(nodes["atom_start"][r]), int(nodes["atom_end"][r])): + residue_of_row[row] = r + + disulfide_pairs: List[Tuple[int, int]] = [] + disulfide: Dict[str, np.ndarray] = {} + if xyz is not None: + sg_rows = [ + row + for row in range(len(cols["name"])) + if cols["name"][row] == "SG" and cols["record"][row] == "ATOM" + ] + disulfide_pairs = find_disulfide_links(sg_rows, residue_of_row, xyz) + disulfide = _disulfide_edges( + pdb, nodes, cols, residue_of_row, disulfide_pairs, link_dict, verbose + ) + + link_edges, link_atom_pairs = _link_record_edges( + pdb, links, disulfide.get("bond"), verbose + ) + + # LINK edges carry ``index`` values, so lift them through that column. + index_to_residue = {int(cols["index"][row]): r for row, r in residue_of_row.items()} + + link_pairs = [(a, b, "TRANS") for a, b in peptide_pairs] + disulf_residue_pairs = sorted( + { + ( + min(residue_of_row[a], residue_of_row[b]), + max(residue_of_row[a], residue_of_row[b]), + ) + for a, b in disulfide_pairs + } + ) + link_pairs += [(a, b, "disulf") for a, b in disulf_residue_pairs] + for a, b in link_atom_pairs: + ra, rb = index_to_residue.get(a), index_to_residue.get(b) + if ra is not None and rb is not None and ra != rb: + link_pairs.append((ra, rb, "LINK")) + + residues = ResidueGraph( + chain=nodes["chain"], + resseq=nodes["resseq"], + icode=nodes["icode"], + resname=nodes["resname"], + template_key=template_key, + atom_start=nodes["atom_start"], + atom_end=nodes["atom_end"], + link_pairs=( + np.array([(a, b) for a, b, _ in link_pairs], dtype=np.int64) + if link_pairs + else np.zeros((0, 2), dtype=np.int64) + ), + link_kind=np.array([k for _, _, k in link_pairs], dtype=" np.ndarray: + """Row order that sorts ``rows`` lexicographically, left column most significant. + + Parameters + ---------- + rows : numpy.ndarray + Integer array of shape ``(E, k)``. + + Returns + ------- + numpy.ndarray + Permutation of ``arange(E)``. Total on the row values, so it does not depend + on the incoming order the way a single-column ``argsort`` does. + """ + if rows.size == 0: + return np.zeros(0, dtype=np.int64) + return np.lexsort(tuple(rows[:, c] for c in reversed(range(rows.shape[1])))) + + +@dataclass(eq=False, repr=False) +class EdgeBlock(DeviceMixin): + """One edge type's index block plus its per-origin bounds. + + Parameters + ---------- + indices : torch.Tensor + Atom indices, shape ``(E, k)``, dtype ``int64``, in canonical order. + origin_bounds : dict + ``{origin: (start, end)}`` half-open row ranges into ``indices``. Ranges are + contiguous and cover the block. + + Notes + ----- + Holds no refinable parameters, so this is a dataclass rather than an + ``nn.Module``; ``DeviceMixin`` still moves ``indices`` with ``.to(device)``. + """ + + indices: torch.Tensor + origin_bounds: Dict[str, Tuple[int, int]] = field(default_factory=dict) + + @classmethod + def empty(cls, arity: int, device=None) -> "EdgeBlock": + """An edge-free block of the given arity.""" + return cls( + indices=torch.zeros((0, arity), dtype=torch.int64, device=device), + origin_bounds={}, + ) + + @classmethod + def from_origins( + cls, + per_origin: Dict[str, np.ndarray], + arity: int, + edge_type: str, + device=None, + ) -> "EdgeBlock": + """Assemble a canonical block from ``{origin: (E_o, k) index array}``. + + Origins are laid out in :data:`ORIGIN_ORDER` for ``edge_type``, and rows + within an origin are sorted lexicographically. Origins with no rows are + omitted from ``origin_bounds`` rather than recorded as empty ranges. + + Parameters + ---------- + per_origin : dict + Integer index arrays keyed by origin. Empty arrays are skipped. + arity : int + Atoms per edge (2 for bonds, 3 for angles, ...). + edge_type : str + Key into :data:`ORIGIN_ORDER`. + device : torch.device, optional + Where to place the block. + + Returns + ------- + EdgeBlock + """ + order = ORIGIN_ORDER.get(edge_type, tuple(sorted(per_origin))) + unknown = set(per_origin) - set(order) + if unknown: + raise ValueError( + f"{edge_type}: origins {sorted(unknown)} are not in ORIGIN_ORDER" + f"[{edge_type!r}] = {order}. Add them there so the layout stays " + f"deterministic." + ) + + chunks: List[np.ndarray] = [] + bounds: Dict[str, Tuple[int, int]] = {} + cursor = 0 + for origin in order: + rows = per_origin.get(origin) + if rows is None or len(rows) == 0: + continue + rows = np.asarray(rows, dtype=np.int64).reshape(-1, arity) + rows = rows[_lexsort_rows(rows)] + chunks.append(rows) + bounds[origin] = (cursor, cursor + len(rows)) + cursor += len(rows) + + if not chunks: + return cls.empty(arity, device=device) + + stacked = np.concatenate(chunks, axis=0) + return cls( + indices=torch.as_tensor(stacked, dtype=torch.int64, device=device), + origin_bounds=bounds, + ) + + @property + def device(self) -> torch.device: + """Where the block lives. Derived, so it cannot fall out of step.""" + return self.indices.device + + @property + def n_edges(self) -> int: + """Number of rows in the block.""" + return int(self.indices.shape[0]) + + @property + def arity(self) -> int: + """Atoms per edge.""" + return int(self.indices.shape[1]) + + def origins(self) -> List[str]: + """Origins present, in block layout order.""" + return sorted(self.origin_bounds, key=lambda o: self.origin_bounds[o][0]) + + def origin(self, name: str) -> torch.Tensor: + """Rows contributed by one origin, as a **view** into the block. + + Parameters + ---------- + name : str + Origin key. + + Returns + ------- + torch.Tensor + Shape ``(E_o, k)``. Shares storage with :attr:`indices`, so an in-place + edit to either is visible through the other. + + Raises + ------ + KeyError + If the origin contributed no rows. + """ + start, end = self.origin_bounds[name] + return self.indices[start:end] + + def tuple_set(self, origin: str = None) -> set: + """Edges as a set of index tuples, for order-free comparison. + + Parameters + ---------- + origin : str, optional + Restrict to one origin. None means the whole block. + + Returns + ------- + set of tuple of int + """ + rows = self.indices if origin is None else self.origin(origin) + return {tuple(int(v) for v in row) for row in rows.cpu().numpy()} + + def __repr__(self) -> str: + return ( + f"EdgeBlock(arity={self.arity}, n_edges={self.n_edges}, " + f"origins={self.origins()})" + ) + + +__all__ = ["EdgeBlock", "ORIGIN_ORDER"] diff --git a/torchref/topology/residue_graph.py b/torchref/topology/residue_graph.py new file mode 100644 index 00000000..453d6f61 --- /dev/null +++ b/torchref/topology/residue_graph.py @@ -0,0 +1,237 @@ +"""The sequence level of a topology: residues as nodes, links as edges. + +A residue node is one instance of a monomer-library template. Its identity is +``(chain, resseq, icode)`` -- the insertion code included, so residues 100 and 100A are +distinct nodes. Edges are the inter-residue links: peptide bonds, disulfides, and +explicit ``LINK`` records. + +Holds no tensors, so this is a plain dataclass; the atom-level tensors live on +:class:`~torchref.topology.atom_graph.AtomGraph`. +""" + +from dataclasses import dataclass, field +from typing import Dict, List, Sequence, Tuple + +import numpy as np +import torch + +#: SG-SG separation below which two cysteines are taken to be disulfide-bonded. +DISULFIDE_MAX_DISTANCE = 2.5 + +#: Lower bound guarding against an atom being paired with itself through a +#: coordinate duplicate. +DISULFIDE_MIN_DISTANCE = 0.1 + + +def _residue_runs( + chain: np.ndarray, resseq: np.ndarray, icode: np.ndarray +) -> Tuple[np.ndarray, np.ndarray]: + """Start and end row of each contiguous ``(chain, resseq, icode)`` run. + + Contiguity is assumed rather than checked, matching the reader's guarantee that a + structure's atoms arrive grouped by residue. A residue split across two + non-adjacent runs would become two nodes. + + Returns + ------- + starts, ends : numpy.ndarray + Half-open row ranges, shape ``(R,)`` each. + """ + n = len(chain) + if n == 0: + return np.zeros(0, dtype=np.int64), np.zeros(0, dtype=np.int64) + changed = np.zeros(n, dtype=bool) + changed[0] = True + changed[1:] = ( + (chain[1:] != chain[:-1]) + | (resseq[1:] != resseq[:-1]) + | (icode[1:] != icode[:-1]) + ) + starts = np.nonzero(changed)[0].astype(np.int64) + ends = np.append(starts[1:], n).astype(np.int64) + return starts, ends + + +@dataclass +class ResidueGraph: + """Residues as nodes, inter-residue links as edges. + + Parameters + ---------- + chain, resseq, icode, resname : numpy.ndarray + Per-residue identity, shape ``(R,)``. + template_key : numpy.ndarray + Restraint-dictionary key per residue, shape ``(R,)``. Either the residue name + or a link-modified variant such as ``'ALA:DEL-HN1+DEL-OXT'``. + atom_start, atom_end : numpy.ndarray + Half-open row range of each residue's atoms, shape ``(R,)``. + link_pairs : numpy.ndarray + Residue index pairs, shape ``(L, 2)``. For a peptide link the first entry + donates its ``C`` and the second its ``N``. + link_kind : numpy.ndarray + Link type per edge, shape ``(L,)``: ``'TRANS'``, ``'PTRANS'``, ``'disulf'`` or + ``'LINK'``. + """ + + chain: np.ndarray + resseq: np.ndarray + icode: np.ndarray + resname: np.ndarray + template_key: np.ndarray + atom_start: np.ndarray + atom_end: np.ndarray + link_pairs: np.ndarray = field( + default_factory=lambda: np.zeros((0, 2), dtype=np.int64) + ) + link_kind: np.ndarray = field(default_factory=lambda: np.zeros(0, dtype=" int: + """Number of residue nodes.""" + return len(self.chain) + + def key(self, i: int) -> Tuple[str, int, str]: + """Identity of residue ``i`` as ``(chain, resseq, icode)``.""" + return (str(self.chain[i]), int(self.resseq[i]), str(self.icode[i])) + + def atom_rows(self, i: int) -> range: + """Row range of residue ``i``'s atoms.""" + return range(int(self.atom_start[i]), int(self.atom_end[i])) + + def links_of_kind(self, kind: str) -> np.ndarray: + """Link edges of one kind, shape ``(L_k, 2)``.""" + if len(self.link_kind) == 0: + return np.zeros((0, 2), dtype=np.int64) + return self.link_pairs[self.link_kind == kind] + + def __repr__(self) -> str: + kinds = ( + {k: int((self.link_kind == k).sum()) for k in np.unique(self.link_kind)} + if len(self.link_kind) + else {} + ) + return f"ResidueGraph(n_residues={self.n_residues}, links={kinds})" + + +def build_residue_nodes( + chain: np.ndarray, + resseq: np.ndarray, + icode: np.ndarray, + resname: np.ndarray, +) -> Dict[str, np.ndarray]: + """Per-residue identity arrays and atom ranges from per-atom columns. + + Returns + ------- + dict + ``chain``, ``resseq``, ``icode``, ``resname``, ``atom_start``, ``atom_end``, + each shape ``(R,)``. + """ + starts, ends = _residue_runs(chain, resseq, icode) + return { + "chain": chain[starts], + "resseq": resseq[starts], + "icode": icode[starts], + "resname": resname[starts], + "atom_start": starts, + "atom_end": ends, + } + + +def find_peptide_links( + nodes: Dict[str, np.ndarray], + names_by_residue: List[set], +) -> List[Tuple[int, int]]: + """Sequence-adjacent residue pairs carrying a C-N peptide bond. + + Two residues are sequence-adjacent when they are neighbours in their chain's + ``(resseq, icode)`` ordering **and** either share a ``resseq`` -- an insertion-code + step such as 100 to 100A -- or differ by exactly one. The second condition is what + stops a chain break being bridged: residues 49 and 56 are neighbours in the ordering + but not in the sequence. + + The C/N test then narrows to pairs that actually carry the bond, which is the + condition the link builders apply implicitly when they look the two atoms up. + + Parameters + ---------- + nodes : dict + Output of :func:`build_residue_nodes`. + names_by_residue : list of set + Atom names present in each residue. + + Returns + ------- + list of tuple of int + ``(residue donating C, residue donating N)`` pairs. + """ + by_chain: Dict[str, List[int]] = {} + for i in range(len(nodes["chain"])): + by_chain.setdefault(str(nodes["chain"][i]), []).append(i) + + pairs = [] + for members in by_chain.values(): + ordered = sorted( + members, key=lambda i: (int(nodes["resseq"][i]), str(nodes["icode"][i])) + ) + for a, b in zip(ordered, ordered[1:]): + step = int(nodes["resseq"][b]) - int(nodes["resseq"][a]) + if step not in (0, 1): + continue + if "C" in names_by_residue[a] and "N" in names_by_residue[b]: + pairs.append((a, b)) + return pairs + + +def find_disulfide_links( + sg_rows: Sequence[int], + residue_of_row: Dict[int, int], + xyz: torch.Tensor, +) -> List[Tuple[int, int]]: + """``SG`` atom pairs within :data:`DISULFIDE_MAX_DISTANCE` in different residues. + + Pairing is per **atom**, not per residue, because a cysteine modelled in two + alternative conformations carries two ``SG`` atoms and each can form its own bond. + Pairing per residue would keep only one of them. Two ``SG`` atoms of the same + residue -- its own alternative conformers -- are within bonding distance of each + other and are excluded by the differing-residue test. + + Parameters + ---------- + sg_rows : sequence of int + Atom rows of every ``SG`` under consideration. + residue_of_row : dict + ``{atom row: residue index}`` for those rows. + xyz : torch.Tensor + Cartesian coordinates, shape ``(N, 3)``. + + Returns + ------- + list of tuple of int + Atom-row pairs, lower row first, ascending. + """ + rows = list(sg_rows) + if len(rows) < 2: + return [] + idx = torch.as_tensor(rows, dtype=torch.int64, device=xyz.device) + dist = torch.cdist(xyz[idx], xyz[idx]) + close = (dist > DISULFIDE_MIN_DISTANCE) & (dist < DISULFIDE_MAX_DISTANCE) + + a, b = torch.triu_indices(len(rows), len(rows), offset=1, device=xyz.device) + hit = close[a, b] + pairs = [] + for i, j in zip(a[hit].cpu().tolist(), b[hit].cpu().tolist()): + row_i, row_j = rows[i], rows[j] + if residue_of_row[row_i] != residue_of_row[row_j]: + pairs.append((row_i, row_j)) + return pairs + + +__all__ = [ + "ResidueGraph", + "build_residue_nodes", + "find_peptide_links", + "find_disulfide_links", + "DISULFIDE_MAX_DISTANCE", + "DISULFIDE_MIN_DISTANCE", +] diff --git a/torchref/topology/templates.py b/torchref/topology/templates.py new file mode 100644 index 00000000..cbfcc62b --- /dev/null +++ b/torchref/topology/templates.py @@ -0,0 +1,105 @@ +"""Monomer templates and the link modifications that patch them. + +The monomer library describes each residue in its **free** form. Forming a peptide bond +applies the modifications the ``chem_link`` table names -- ``DEL-OXT`` to the residue +donating its C, ``DEL-HN1`` (``DEL-HNP`` for proline) to the residue donating its N -- +which delete the restraints the link makes meaningless and overwrite the targets that +change. A residue therefore draws its restraints from a *patched* template, identified +by a key such as ``'ALA:DEL-HN1+DEL-OXT'``. + +Chain termini are deliberately left unpatched: a real C-terminus keeps its ``OXT`` and +carboxylate geometry, a real N-terminus its ammonium. +""" + +from typing import Dict, Sequence, Tuple + +import numpy as np + +from torchref.restraints.modifications import ( + apply_modifications, + link_modifications, + read_mod_definitions, +) + + +def resolve_template_keys( + resnames: Sequence[str], + peptide_pairs: Sequence[Tuple[int, int]], + cif_dict: Dict, + link_list, + verbose: int = 0, +) -> Tuple[Dict, np.ndarray]: + """Assign each residue the template it should draw restraints from. + + Parameters + ---------- + resnames : sequence of str + Residue name per residue index. + peptide_pairs : sequence of tuple of int + ``(donates C, donates N)`` residue index pairs. + cif_dict : dict + Restraint dictionary keyed by residue name. + link_list : pandas.DataFrame or None + Link-type definitions, as + :func:`~torchref.restraints.restraints_helper.read_link_definitions` returns + them. None disables patching. + verbose : int, default 0 + Verbosity level. + + Returns + ------- + comp_dict : dict + ``cif_dict`` plus one entry per patched ``(residue type, modification set)`` in + use. ``cif_dict`` itself is left keyed by residue name alone. + template_key : numpy.ndarray + Key per residue, shape ``(R,)``. Residues that are not patched carry their own + residue name. + """ + comp_dict = dict(cif_dict) + keys = np.array([str(r) for r in resnames], dtype=object) + + if not cif_dict or link_list is None or len(peptide_pairs) == 0: + return comp_dict, keys + + modifications = link_modifications(link_list) + if "TRANS" not in modifications: + return comp_dict, keys + trans_mods = modifications["TRANS"] + proline_mods = modifications.get("PTRANS", trans_mods) + + mods_by_residue: Dict[int, set] = {} + for res_c, res_n in peptide_pairs: + donor_mod, acceptor_mod = ( + proline_mods if str(resnames[res_n]) == "PRO" else trans_mods + ) + for res_idx, mod_id in ((res_c, donor_mod), (res_n, acceptor_mod)): + if mod_id is None: + continue + mods_by_residue.setdefault(res_idx, set()).add(mod_id) + + if not mods_by_residue: + return comp_dict, keys + + mod_dict = read_mod_definitions() + n_patched = 0 + for res_idx, mods in mods_by_residue.items(): + resname = str(resnames[res_idx]) + if resname not in comp_dict: + continue + variant = f"{resname}:{'+'.join(sorted(mods))}" + if variant not in comp_dict: + comp_dict[variant] = apply_modifications( + comp_dict[resname], sorted(mods), mod_dict + ) + keys[res_idx] = variant + n_patched += 1 + + if verbose > 1: + print( + f"Patched {n_patched} residues " + f"({len(comp_dict) - len(cif_dict)} template variants)" + ) + return comp_dict, keys + + +__all__ = ["resolve_template_keys"] diff --git a/torchref/topology/topology.py b/torchref/topology/topology.py new file mode 100644 index 00000000..7ec289c1 --- /dev/null +++ b/torchref/topology/topology.py @@ -0,0 +1,106 @@ +"""The topology container: a residue graph over an atom graph. + +:class:`Topology` is what the model's connectivity lives in. The residue level carries +the sequence and the inter-residue links; the atom level carries the atoms, the typed +edge blocks and the bond adjacency. Per-atom residue identity is reached through +``atoms.residue_of`` rather than duplicated per atom. + +Mutable by design; prefer :meth:`Topology.copy` over editing in place. +""" + +from dataclasses import dataclass +from typing import Dict, Set, Tuple + +import numpy as np +import torch + +from torchref.topology.atom_graph import AtomGraph +from torchref.topology.residue_graph import ResidueGraph +from torchref.utils.device_mixin import DeviceMixin + + +@dataclass(eq=False, repr=False) +class Topology(DeviceMixin): + """Connectivity of one model, at both the residue and the atom level. + + Parameters + ---------- + residues : ResidueGraph + Sequence level -- residues as nodes, links as edges. + atoms : AtomGraph + Atom level -- atoms as nodes, typed edge blocks, bond adjacency. + + Notes + ----- + Holds no refinable parameters, so this is a dataclass rather than an ``nn.Module``. + Edge indices are ``int64`` constants and no gradient reaches them; gradients reach + the coordinates that the indices gather. + """ + + residues: ResidueGraph + atoms: AtomGraph + + @property + def device(self) -> torch.device: + """Where the indexing tensors live. Derived from the atom graph.""" + return self.atoms.device + + @property + def n_atoms(self) -> int: + """Number of atom nodes.""" + return self.atoms.n_atoms + + @property + def n_residues(self) -> int: + """Number of residue nodes.""" + return self.residues.n_residues + + def neighbors(self, i: int) -> np.ndarray: + """Atoms bonded to atom ``i``. Delegates to :meth:`AtomGraph.neighbors`.""" + return self.atoms.neighbors(i) + + def residue_of_atom(self, i: int) -> int: + """Residue index of atom ``i``.""" + return int(self.atoms.residue_of[i]) + + def resname_of_atom(self, i: int) -> str: + """Residue name of atom ``i``, joined through the residue graph.""" + return str(self.residues.resname[self.residue_of_atom(i)]) + + def edge_block(self, edge_type: str): + """The :class:`~torchref.topology.edges.EdgeBlock` for one edge type. + + Parameters + ---------- + edge_type : str + ``'bond'``, ``'angle'``, ``'torsion'`` or ``'chiral'``. Planes are ragged + and reached through ``atoms.planes``. + """ + return { + "bond": self.atoms.bonds, + "angle": self.atoms.angles, + "torsion": self.atoms.torsions, + "chiral": self.atoms.chirals, + }[edge_type] + + def tuple_sets(self) -> Dict[str, Dict[str, Set[Tuple[int, ...]]]]: + """Every edge as ``{edge type: {origin: set of index tuples}}``. + + Order-free, so this is what an equivalence check against another builder should + compare. + """ + out: Dict[str, Dict[str, Set[Tuple[int, ...]]]] = {} + for name in ("bond", "angle", "torsion", "chiral"): + block = self.edge_block(name) + out[name] = {o: block.tuple_set(o) for o in block.origins()} + out["plane"] = {} + for size, block in self.atoms.planes.items(): + for origin in block.origins(): + out["plane"][f"{size}_atoms/{origin}"] = block.tuple_set(origin) + return out + + def __repr__(self) -> str: + return f"Topology({self.residues!r}, {self.atoms!r})" + + +__all__ = ["Topology"] From 8d008c5d448cb26b40a3d786f30ee88d286f260f Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Wed, 26 Aug 2026 19:42:33 +0200 Subject: [PATCH 058/250] Build restraints from the topology, behind an unchanged read surface The restraint groups now come from the topology graph and the values layered over its edges, instead of from a flat TensorDict addressed by composed string keys. Reading them is unchanged: restraints["bond"]["all"]["indices"] resolves exactly as before, so no geometry target is edited. What changed is what happens on that read. It used to build a fresh accessor object, then a fresh per-type accessor, then a fresh dict assembled from six string-formatted buffer lookups -- per loss evaluation, per restraint type. It is now three dict lookups into a mapping assembled once, holding tensors that already exist. The per-origin indices are slices of one contiguous block, so a subset is a view rather than a copy and an in-place edit to a block needs no invalidation to be seen. Nothing is cached, so nothing can go stale on the access path. Only two events can orphan a view -- a device rebind and an edge-count change -- and both rebuild the mapping where they happen. _apply drops the views before the traversal and re-slices after, because DeviceMixin's walk recurses into dicts and would otherwise move each slice on its own, silently turning every view into an independent tensor. cat_dict was not idempotent: writing restraints["bond"]["all"] registered 'all' as an origin, so a second call folded the group into itself and doubled every bond, angle and torsion -- a 2x on the geometry weight, latent only because every call site happened to guard it. 'all' is now a span of the edge block and cannot be an origin, so that is unrepresentable rather than fixed. Restraint row order no longer depends on Python's string hash seed. It used to come from set iteration over the origin names, which reordered the rows between processes. Numerically neutral, verified two ways rather than assumed. Every (edge, property) pair matches the previous storage exactly across five structures -- including phi and psi correctly carrying no reference or sigma, and omega keeping its proline flag. Against the previous commit directly, every restraint count is identical, n_vdw included, which is what shows the exclusion set did not move; the losses agree to ~1e-7 relative, which is float32 summation-order noise from the canonical layout. Exclusions still come from the bond, angle and torsion edges rather than from bond connectivity. The connectivity-derived set is correct and excludes more pairs, so it moves the non-bonded term and belongs in its own measured commit. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- docs/changelog.rst | 4 + tests/unit/topology/test_storage.py | 205 +++++++ torchref/restraints/restraints.py | 905 ++++------------------------ torchref/topology/__init__.py | 6 +- torchref/topology/build.py | 408 ++++++++++--- torchref/topology/edges.py | 111 +++- torchref/topology/restraint_sets.py | 180 ++++++ 7 files changed, 919 insertions(+), 900 deletions(-) create mode 100644 tests/unit/topology/test_storage.py create mode 100644 torchref/topology/restraint_sets.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 34f4b396..865e7e95 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -22,6 +22,10 @@ Unreleased - Added ``torchref.topology``: a ``Topology`` of a ``ResidueGraph`` over an ``AtomGraph``, holding the model's connectivity as typed edge blocks with a bond adjacency that answers ``neighbors(i)`` - Topology residues are identified by ``(chain, resseq, icode)``, so a residue with an insertion code is no longer merged with the one it was inserted after - Added ``AtomGraph.exclusions_12_13_14``, deriving non-bonded exclusions from bond connectivity rather than from which angles and torsions the monomer library happens to restrain +- ``Restraints.restraints`` is now a plain nested dict of tensors instead of an accessor object rebuilt on every read; the per-origin indices are views into the topology's contiguous edge blocks +- Restraint groups are laid out in a fixed order, so a rebuild produces the same row order in any process; previously the origins were concatenated in Python ``set`` iteration order +- Fixed ``cat_dict`` doubling every bond, angle and torsion restraint when called more than once +- Geometry restraints are built from the topology, retiring the intra-residue builder calls, the peptide/disulfide/LINK build methods and the ``TensorDict`` restraint storage Version 0.6.4 diff --git a/tests/unit/topology/test_storage.py b/tests/unit/topology/test_storage.py new file mode 100644 index 00000000..2a86f79d --- /dev/null +++ b/tests/unit/topology/test_storage.py @@ -0,0 +1,205 @@ +"""The restraint storage: plain dict, views into the edge blocks, no per-access work. + +The geometry targets read ``restraints[edge_type][origin][property]`` on every +iteration, so the properties tested here are load-bearing rather than cosmetic: the +mapping must be a plain dict of already-materialised tensors, the per-origin entries +must alias the contiguous blocks rather than copy them, and none of that may come +undone on a device move or a copy. +""" + +import pytest +import torch + +from torchref.model.model import Model +from torchref.utils.caching import ParameterFingerprint + +KEYED_TYPES = ("bond", "angle", "torsion") + +#: The surface the geometry targets rely on. Guards the contract from drifting: 'phi' +#: and 'psi' are conformationally free and must NOT acquire a reference value or sigma, +#: and 'omega' must keep the proline flag its own target reads. +EXPECTED_PROPERTIES = { + ("bond", "all"): {"indices", "references", "sigmas"}, + ("bond", "intra"): {"indices", "references", "sigmas"}, + ("angle", "all"): {"indices", "references", "sigmas"}, + ("torsion", "all"): {"indices", "references", "sigmas", "periods"}, + ("torsion", "phi"): {"indices", "periods"}, + ("torsion", "psi"): {"indices", "periods"}, + ("torsion", "omega"): { + "indices", + "references", + "sigmas", + "periods", + "is_proline", + }, +} + + +@pytest.fixture(scope="module") +def restraints(pdb_dir): + """Restraints for a structure with altlocs, disulfides and peptide links.""" + model = Model(verbose=0) + model.load_pdb(str(pdb_dir / "7L84.pdb")) + model.set_restraints_cif(None) + return model.restraints + + +@pytest.mark.unit +def test_restraints_is_a_plain_dict(restraints): + """Not an accessor object that has to be constructed per access.""" + assert type(restraints.restraints) is dict + assert restraints.restraints is restraints.restraints + + +@pytest.mark.unit +def test_access_allocates_nothing(restraints): + """Every level of the lookup returns the same object each time. + + This is what makes the read cheap: three dict lookups and no construction. The + previous storage built a fresh accessor, then a fresh per-type accessor, then a + fresh dict of six string-keyed buffer lookups, on every call. + """ + first = restraints.restraints["bond"]["all"] + second = restraints.restraints["bond"]["all"] + assert first is second + assert first["indices"] is second["indices"] + + +@pytest.mark.unit +@pytest.mark.parametrize("edge_type", KEYED_TYPES) +def test_origin_entries_alias_the_block(restraints, edge_type): + """Per-origin indices are slices of one block, not copies of it.""" + block = restraints.topology.edge_block(edge_type) + entries = restraints.restraints[edge_type] + + for origin, bounds in block.origin_bounds.items(): + indices = entries[origin]["indices"] + assert indices.data_ptr() == block.indices[bounds[0] : bounds[1]].data_ptr() + assert indices.shape[0] == bounds[1] - bounds[0] + + +@pytest.mark.unit +@pytest.mark.parametrize("edge_type", KEYED_TYPES) +def test_all_group_is_a_view(restraints, edge_type): + """The combined group the targets read is a span of the block, not a concatenation. + + The block layout deliberately keeps each type's ``all`` members adjacent so this + holds; ``torsion`` is the one that needs it, since its group is only ``intra`` plus + ``disulfide``. + """ + block = restraints.topology.edge_block(edge_type) + combined = restraints.restraints[edge_type]["all"]["indices"] + assert combined.data_ptr() == block.indices.data_ptr() + + +@pytest.mark.unit +def test_in_place_block_edit_is_visible_through_every_entry(restraints): + """Shared storage means a block edit needs no invalidation to be seen.""" + block = restraints.topology.atoms.bonds + entries = restraints.restraints["bond"] + origin = block.origins()[0] + + saved = int(block.indices[0, 0]) + try: + block.indices[0, 0] = saved + 7 + assert int(entries[origin]["indices"][0, 0]) == saved + 7 + assert int(entries["all"]["indices"][0, 0]) == saved + 7 + finally: + block.indices[0, 0] = saved + + +@pytest.mark.unit +@pytest.mark.parametrize("key", sorted(EXPECTED_PROPERTIES)) +def test_expected_properties_present(restraints, key): + """Each group carries exactly the properties its consumers expect.""" + edge_type, origin = key + group = restraints.restraints[edge_type][origin] + assert set(group) == EXPECTED_PROPERTIES[key] + + +@pytest.mark.unit +def test_cat_dict_is_idempotent(restraints): + """Repeated calls leave the restraint counts alone. + + The previous implementation registered ``'all'`` as an origin when it wrote the + combined group, so a second call concatenated the group into itself and doubled + every bond, angle and torsion -- a silent 2x on the geometry weight. Deriving the + group as a span of the block makes that unrepresentable. + """ + before = { + edge_type: restraints.restraints[edge_type]["all"]["indices"].shape[0] + for edge_type in KEYED_TYPES + } + restraints.cat_dict() + restraints.cat_dict() + after = { + edge_type: restraints.restraints[edge_type]["all"]["indices"].shape[0] + for edge_type in KEYED_TYPES + } + assert before == after + + +@pytest.mark.unit +def test_entries_survive_a_device_apply(restraints): + """A ``.to()`` re-slices the entries instead of leaving them stale or duplicated. + + ``DeviceMixin``'s walk recurses into dicts, so without the ``_apply`` override each + view would be moved on its own and become an independent tensor. + """ + before = restraints.restraints["bond"]["all"]["indices"].clone() + n_before = { + t: restraints.restraints[t]["all"]["indices"].shape[0] for t in KEYED_TYPES + } + + restraints.to(torch.device("cpu")) + + block = restraints.topology.atoms.bonds.indices + after = restraints.restraints["bond"]["all"]["indices"] + assert after.data_ptr() == block.data_ptr(), "entries no longer alias the block" + assert torch.equal(after, before) + assert { + t: restraints.restraints[t]["all"]["indices"].shape[0] for t in KEYED_TYPES + } == n_before + assert restraints.restraints["vdw"].get("indices") is not None + + +@pytest.mark.unit +def test_blocks_are_untouched_by_a_refinement_step(restraints): + """The edge tensors are constants; nothing in a loss evaluation may mutate them.""" + blocks = [restraints.topology.edge_block(t).indices for t in KEYED_TYPES] + fingerprint = ParameterFingerprint(blocks) + + loss = restraints.nll_bonds().sum() + restraints.nll_angles().sum() + loss.backward() + + assert fingerprint.matches( + [restraints.topology.edge_block(t).indices for t in KEYED_TYPES] + ) + + +@pytest.mark.unit +def test_rebuilding_entries_reslices_onto_the_current_blocks(restraints): + """Re-deriving the entries produces fresh views of the same blocks. + + This is the operation ``_apply`` and ``copy`` both rely on, and the one that has to + stay cheap: it re-slices rather than recomputing anything. + + ``RestraintsNew.copy`` is not exercised here because it cannot run at all -- it is + ``deepcopy``, which walks the *borrowed* ``_xyz_fn`` wrapper, whose cache holds a + graph-attached tensor once ``xyz()`` has been evaluated. Verified to fail + identically at the commit before this change, so it is pre-existing rather than a + regression, and it is reached only through ``Model.copy`` on a model whose lazy + restraints have already been built. + """ + block = restraints.topology.atoms.bonds.indices + before = restraints.restraints["bond"]["all"]["indices"].clone() + + restraints._rebuild_entries() + + after = restraints.restraints["bond"]["all"]["indices"] + assert after.data_ptr() == block.data_ptr() + assert torch.equal(after, before) + for origin, bounds in restraints.topology.atoms.bonds.origin_bounds.items(): + entry = restraints.restraints["bond"][origin]["indices"] + assert entry.shape[0] == bounds[1] - bounds[0] + assert restraints.restraints["vdw"].get("indices") is not None diff --git a/torchref/restraints/restraints.py b/torchref/restraints/restraints.py index aef98fd3..ed994f0c 100644 --- a/torchref/restraints/restraints.py +++ b/torchref/restraints/restraints.py @@ -1,36 +1,22 @@ """Restraints handler for crystallographic model refinement. -Provides :class:`RestraintsNew`, which builds geometry restraints (bonds, -angles, torsions, planes, chirals, VDW) using dedicated builder classes. -It is decoupled from :class:`~torchref.model.Model`: it accepts a pdb -DataFrame and callable functions for accessing coordinates and ADPs. +Provides :class:`RestraintsNew`, which holds the geometry restraints (bonds, angles, +torsions, planes, chirals, VDW) for one structure. The geometry half is a +:class:`~torchref.topology.topology.Topology` plus the ideal values layered over its +edges; the non-bonded pair list is separate, because it is distance-derived and rebuilt +as the model moves. + +Decoupled from :class:`~torchref.model.Model`: it accepts a pdb DataFrame and callables +for coordinates, ADPs and VDW radii. """ -from typing import Callable, Optional +from typing import Callable import numpy as np import pandas as pd import torch from torch.nn import Module -from torchref.restraints.builders_fast import ( - AngleRestraintBuilder, - BondRestraintBuilder, - ChiralRestraintBuilder, - InterResidueAngleBuilder, - InterResidueBondBuilder, - InterResiduePlaneBuilder, - InterResidueTorsionBuilder, - PlaneRestraintBuilder, - PreprocessedPDB, - TorsionRestraintBuilder, - find_peptide_link_pairs, -) -from torchref.restraints.modifications import ( - apply_modifications, - link_modifications, - read_mod_definitions, -) from torchref.restraints.restraints_helper import ( find_cif_file_in_library, read_cif, @@ -38,144 +24,9 @@ ) from torchref.config import get_float_dtype from torchref.utils.debug_utils import DebugMixin -from torchref.utils.utils import TensorDict from torchref.utils.device_mixin import DeviceMixin -class _RestraintsAccessor: - """ - Provides backward-compatible dict-like access to restraints stored in TensorDict. - - This class mimics the old nested dict interface: - restraints["bond"]["intra"]["indices"] - - While actually accessing the TensorDict with flattened keys: - _tensor_storage["bond_intra_indices"] - """ - - # Types that don't have origin level (assigned directly as dicts) - _FLAT_TYPES = {"vdw", "chiral"} - - def __init__(self, parent: "RestraintsNew"): - self._parent = parent - - def __getitem__(self, rtype: str) -> "_RestraintTypeAccessor": - return _RestraintTypeAccessor(self._parent, rtype) - - def __setitem__(self, rtype: str, value): - """Handle direct assignment for flat types like vdw and chiral.""" - if rtype in self._FLAT_TYPES and isinstance(value, dict): - # Store all tensors with empty origin - self._parent._set_restraint_group(rtype, "", value) - else: - raise TypeError( - f"Cannot assign directly to restraints['{rtype}']. " - f"Use restraints['{rtype}'][origin] = data for nested types." - ) - - def __contains__(self, rtype: str) -> bool: - return len(self._parent._restraint_groups.get(rtype, set())) > 0 or \ - rtype in self._FLAT_TYPES and self._parent._has_restraint(rtype, "") - - def get(self, rtype: str, default=None): - if rtype in self: - return self[rtype] - return default - - def keys(self): - """Return all restraint types that have data.""" - result = [] - for rtype in ["bond", "angle", "torsion", "plane"]: - if len(self._parent._restraint_groups.get(rtype, set())) > 0: - result.append(rtype) - # Check for special types (vdw, chiral) which don't have origins - for rtype in self._FLAT_TYPES: - if self._parent._has_restraint(rtype, ""): - result.append(rtype) - return result - - -class _RestraintTypeAccessor: - """ - Provides access to origins within a restraint type. - - For regular types (bond, angle, torsion, plane): - restraints["bond"]["intra"] -> dict with indices, references, sigmas - - For special types (vdw, chiral), this class acts as the dict itself: - restraints["vdw"]["indices"] -> tensor - restraints["vdw"] = {"indices": ..., "sigmas": ...} - """ - - # Types that don't have origin level (accessed directly as dicts) - _FLAT_TYPES = {"vdw", "chiral"} - - def __init__(self, parent: "RestraintsNew", rtype: str): - self._parent = parent - self._rtype = rtype - - def __getitem__(self, key: str): - if self._rtype in self._FLAT_TYPES: - # For vdw/chiral, key is a property name (indices, sigmas, etc.) - tensor = self._parent._get_restraint_tensor(self._rtype, "", key) - if tensor is None: - raise KeyError(f"No {key} for {self._rtype}") - return tensor - else: - # For bond/angle/torsion/plane, key is an origin name - result = self._parent._get_restraint_group(self._rtype, key) - if result is None: - raise KeyError(f"No restraints for {self._rtype}/{key}") - return result - - def __setitem__(self, key: str, value): - if self._rtype in self._FLAT_TYPES: - # For vdw/chiral, if value is a tensor, store it directly - # If value is a dict, store all tensors - if isinstance(value, torch.Tensor): - self._parent._set_restraint_tensor(self._rtype, "", key, value) - elif isinstance(value, dict): - # This handles: restraints["vdw"] = {"indices": ..., "sigmas": ...} - # But this is called as restraints["vdw"][key] = value, so it won't work - # We need special handling in the parent accessor - pass - else: - # For bond/angle/torsion/plane, key is origin, value is dict - self._parent._set_restraint_group(self._rtype, key, value) - - def __contains__(self, key: str) -> bool: - if self._rtype in self._FLAT_TYPES: - return self._parent._get_restraint_tensor(self._rtype, "", key) is not None - return self._parent._has_restraint(self._rtype, key) - - def get(self, key: str, default=None): - try: - return self[key] - except KeyError: - return default - - def keys(self): - if self._rtype in self._FLAT_TYPES: - # Return property names for flat types - result = [] - for prop in ["indices", "references", "sigmas", "periods", "min_distances", - "symop_indices", "cell_offsets"]: - if self._parent._get_restraint_tensor(self._rtype, "", prop) is not None: - result.append(prop) - return result - return self._parent._get_origins_for_type(self._rtype) - - def items(self): - if self._rtype in self._FLAT_TYPES: - for prop in self.keys(): - yield prop, self._parent._get_restraint_tensor(self._rtype, "", prop) - else: - for origin in self.keys(): - yield origin, self._parent._get_restraint_group(self._rtype, origin) - - def __iter__(self): - return iter(self.keys()) - class RestraintsNew(DeviceMixin, DebugMixin, Module): """ @@ -211,9 +62,11 @@ class RestraintsNew(DeviceMixin, DebugMixin, Module): Attributes ---------- - restraints : _RestraintsAccessor - Dict-*like* accessor over the flat TensorDict: it emulates - ``restraints["bond"]["intra"]["indices"]`` but is not a plain dict. + restraints : dict + Restraint groups as ``restraints["bond"]["intra"]["indices"]``. A plain nested + dict; the per-origin indices are views into ``topology``'s edge blocks. + topology : Topology + The connectivity the geometry restraints are defined over. cif_dict : dict Parsed CIF restraints keyed by residue type; ``missing_residues`` lists the types that could not be resolved. @@ -251,11 +104,15 @@ def __init__( self._cell = cell self._spacegroup = spacegroup - # Initialize TensorDict for restraint storage (registered as submodule) - self._tensor_storage = TensorDict() - - # Track which restraint groups exist (for iteration) - self._restraint_groups = {"bond": set(), "angle": set(), "torsion": set(), "plane": set()} + # Connectivity, the values layered over it, and the non-bonded pair list, which + # is rebuilt on displacement and so is kept apart from the rest. + self.topology = None + self._values = {} + self._vdw = {} + # Derived: per-origin views into the topology's edge blocks. Rebuilt by + # _rebuild_entries, which runs at build time and after any device move. + self._entries = {} + self._torsion_max_period = 1 # Empty initialization if pdb is None: @@ -353,68 +210,58 @@ def get_vdw_radii(self) -> torch.Tensor: return self._vdw_radii_fn() # ========================================================================= - # TensorDict Helper Methods for Restraint Storage + # Restraint storage # ========================================================================= - def _make_key(self, rtype: str, origin: str, prop: str) -> str: - """Create flattened key for TensorDict storage.""" - if origin: - return f"{rtype}_{origin}_{prop}" - else: - # For flat types (vdw, chiral) with no origin - return f"{rtype}_{prop}" - def _set_restraint_tensor( - self, rtype: str, origin: str, prop: str, tensor: torch.Tensor - ): - """Store a restraint tensor with flattened key.""" - key = self._make_key(rtype, origin, prop) - self._tensor_storage[key] = tensor - # Track that this origin exists for this restraint type - if rtype in self._restraint_groups: - self._restraint_groups[rtype].add(origin) - - def _get_restraint_tensor( - self, rtype: str, origin: str, prop: str - ) -> Optional[torch.Tensor]: - """Get a restraint tensor by type, origin, and property.""" - key = self._make_key(rtype, origin, prop) - if key in self._tensor_storage: - return self._tensor_storage[key] - return None - - def _has_restraint(self, rtype: str, origin: str) -> bool: - """Check if a restraint group exists.""" - key = self._make_key(rtype, origin, "indices") - return key in self._tensor_storage - - def _set_restraint_group(self, rtype: str, origin: str, data: dict): - """Store all tensors from a restraint data dict.""" - for prop, tensor in data.items(): - if tensor is not None and isinstance(tensor, torch.Tensor): - self._set_restraint_tensor(rtype, origin, prop, tensor) - - def _get_restraint_group(self, rtype: str, origin: str) -> Optional[dict]: - """Get all tensors for a restraint group as a dict.""" - if not self._has_restraint(rtype, origin): - return None - result = {} - # Common properties for different restraint types - for prop in ["indices", "references", "sigmas", "periods", "min_distances", - "is_proline"]: - tensor = self._get_restraint_tensor(rtype, origin, prop) - if tensor is not None: - result[prop] = tensor - return result if result else None - - def _get_origins_for_type(self, rtype: str) -> list: - """Get all origins (e.g., 'intra', 'peptide') for a restraint type.""" - return list(self._restraint_groups.get(rtype, set())) + + + + + @property - def restraints(self) -> "_RestraintsAccessor": - """Nested-dict-*like* accessor over the flat TensorDict (not a real dict).""" - return _RestraintsAccessor(self) + def restraints(self) -> dict: + """Restraint groups as ``[edge type][origin][property]``. + + A plain nested dict of tensors, assembled once at build time. Reading it costs + three dict lookups and no allocation, which matters because the geometry targets + do it on every iteration. Per-origin indices are **views** into the topology's + contiguous edge blocks, so an in-place edit to a block is visible here at once, + and taking a subset costs nothing. + """ + return self._entries + + def _rebuild_entries(self) -> None: + """Re-derive the entry views from the topology and its values. + + Cheap -- a handful of slices -- and idempotent. Runs at the end of a build and + again after any device or dtype move, because moving a tensor rebinds it and + leaves the old views pointing at freed storage. + """ + if self.topology is None: + return + from torchref.topology import assemble_entries, max_period + + self._entries = assemble_entries(self.topology, self._values) + if self._vdw: + self._entries["vdw"] = self._vdw + self._torsion_max_period = max_period(self._entries) + + def _apply(self, fn, recurse: bool = True): + """Drop the derived views before the traversal, re-slice them after. + + ``DeviceMixin``'s ``__dict__`` walk recurses into dicts, so leaving the entries + in place would move each slice on its own and quietly turn every view into an + independent tensor -- doubling the memory and breaking the aliasing the design + rests on. Rebuilding unconditionally rather than in ``_after_device_apply``, + because that hook only fires when the device or dtype actually changed, and a + ``.to()`` onto the current device must not leave the entries empty. + """ + self._entries = {} + result = super()._apply(fn, recurse) + self._rebuild_entries() + return result def _load_cif_dictionaries(self, cif_path): """Load CIF dictionaries from provided paths and monomer library.""" @@ -507,54 +354,34 @@ def _load_rama_surfaces(self, device: torch.device): self.register_buffer("_rama_surfaces", surfaces) def build_restraints(self): - """Build every restraint group; each builder handles all residues at once. + """Build the topology, the values over it, and the non-bonded pair list. Builds on CPU and moves the result to the ``xyz()`` device at the end. """ try: target_device = self.xyz().device device = torch.device("cpu") - pdb = self.pdb - # Must precede the intra-residue builders: it decides which residues - # draw their restraints from a link-modified component instead of the - # bare one. - comp_dict, res_keys = self._build_residue_variants() + from torchref.topology import build_topology_with_values - bond_result = BondRestraintBuilder(verbose=self.verbose).build( - pdb, comp_dict, device, residue_keys=res_keys - ) - if bond_result: - self.restraints["bond"]["intra"] = bond_result - - angle_result = AngleRestraintBuilder(verbose=self.verbose).build( - pdb, comp_dict, device, residue_keys=res_keys - ) - if angle_result: - self.restraints["angle"]["intra"] = angle_result - - torsion_result = TorsionRestraintBuilder(verbose=self.verbose).build( - pdb, comp_dict, device, residue_keys=res_keys - ) - if torsion_result: - self.restraints["torsion"]["intra"] = torsion_result - - plane_result = PlaneRestraintBuilder(verbose=self.verbose).build( - pdb, comp_dict, device, residue_keys=res_keys - ) - if plane_result: - for key, data in plane_result.items(): - self.restraints["plane"][key] = data - - chiral_result = ChiralRestraintBuilder(verbose=self.verbose).build( - pdb, comp_dict, device, residue_keys=res_keys + self.topology, self._values, extras = build_topology_with_values( + self.pdb, + self.cif_dict, + link_dict=self.link_dict, + link_list=self.link_list, + links=self.links, + xyz=self.xyz().detach().to(device), + device=device, + verbose=self.verbose, ) - if chiral_result: - self.restraints["chiral"] = chiral_result + self._rebuild_entries() - self._build_peptide_restraints(device) - self._build_disulfide_restraints(device) - self._build_link_restraints(device) + rama = extras.get("ramachandran") + if rama is not None: + self.register_buffer("_rama_phi_indices", rama["phi_indices"]) + self.register_buffer("_rama_psi_indices", rama["psi_indices"]) + self.register_buffer("_rama_surface_type", rama["surface_type"]) + self._load_rama_surfaces(device) # cutoff sits ~1 Å beyond the largest heavy-atom VDW sum (~3.6 Å) plus # expected drift, so a displacement-triggered rebuild stays inside the @@ -563,11 +390,6 @@ def build_restraints(self): cutoff=6.0, sigma=0.05, inter_residue_only=False, use_spatial_hash=True ) - # Register the concatenated 'all' buffers here, not lazily in forward: - # register_buffer() during a forward pass breaks CUDA-graph capture, and - # only registered buffers are moved by model.to(device). - self.cat_dict() - if target_device.type != "cpu": self.to(target_device) @@ -575,403 +397,9 @@ def build_restraints(self): self.debug_on_error(e, context="RestraintsNew.build_restraints") raise - def _build_residue_variants(self): - """Point peptide-linked residues at link-modified copies of their component. - The monomer library defines each amino acid free: ``ALA`` carries ``OXT`` - and a protonated ``N``, with carboxylate and ammonium geometry. Forming a - peptide bond applies the modifications the ``chem_link`` table names -- - ``DEL-OXT`` to the residue donating its C, ``DEL-HN1`` (``DEL-HNP`` for - proline) to the residue donating its N -- which delete the restraints the - link makes meaningless and overwrite the targets that change, notably - ``CA-C-O`` and ``CA-N-H``. Without them the intra-residue restraints fight - the link's own: around a peptide carbonyl carbon ``CA-C-O`` + ``CA-C-N`` + - ``O-C-N`` only sums to 360 degrees once ``DEL-OXT`` has been applied. - Chain termini are deliberately left unmodified -- a real C-terminus keeps - its ``OXT`` and carboxylate geometry, a real N-terminus its ammonium. - - Returns - ------- - comp_dict : dict - :attr:`cif_dict` plus one entry per ``(residue type, modification - set)`` in use, keyed ``'ALA:DEL-HN1+DEL-OXT'``. :attr:`cif_dict` - itself is left keyed by residue type alone. - residue_keys : dict - ``{(chain_id, resseq): comp_dict key}`` for the modified residues, as - :meth:`~torchref.restraints.builders_fast.RestraintBuilder.build` - takes it. Residues absent from it use their residue name. - """ - comp_dict = dict(self.cif_dict) - residue_keys = {} - if not self.cif_dict or getattr(self, "link_list", None) is None: - return comp_dict, residue_keys - - modifications = link_modifications(self.link_list) - if "TRANS" not in modifications: - return comp_dict, residue_keys - trans_mods = modifications["TRANS"] - proline_mods = modifications.get("PTRANS", trans_mods) - - polymer = self.pdb[self.pdb["ATOM"] == "ATOM"] - if len(polymer) == 0: - return comp_dict, residue_keys - pp_pdb = PreprocessedPDB(polymer) - - mods_by_residue = {} - for res_i, res_next in find_peptide_link_pairs(pp_pdb): - donor_mod, acceptor_mod = ( - proline_mods - if pp_pdb.residue_resnames[res_next] == "PRO" - else trans_mods - ) - for res_idx, mod_id in ((res_i, donor_mod), (res_next, acceptor_mod)): - if mod_id is None: - continue - key = ( - str(pp_pdb.residue_chain_ids[res_idx]), - int(pp_pdb.residue_resseqs[res_idx]), - ) - mods_by_residue.setdefault(key, set()).add(mod_id) - if not mods_by_residue: - return comp_dict, residue_keys - - mod_dict = read_mod_definitions() - for res_idx in range(pp_pdb.n_residues): - key = ( - str(pp_pdb.residue_chain_ids[res_idx]), - int(pp_pdb.residue_resseqs[res_idx]), - ) - mods = mods_by_residue.get(key) - resname = pp_pdb.residue_resnames[res_idx] - if not mods or resname not in comp_dict: - continue - mods = sorted(mods) - variant = f"{resname}:{'+'.join(mods)}" - if variant not in comp_dict: - comp_dict[variant] = apply_modifications( - comp_dict[resname], mods, mod_dict - ) - residue_keys[key] = variant - - if self.verbose > 1: - print( - f"Applied link modifications to {len(residue_keys)} residues " - f"({len(comp_dict) - len(self.cif_dict)} modified components)" - ) - return comp_dict, residue_keys - - def _build_peptide_restraints(self, device: torch.device): - """Build peptide bond/angle/torsion/plane restraints. - - TRANS/CIS links for standard peptide bonds; PTRANS/PCIS for bonds to - proline, which add the C(i-1)-N-CD angle and proline-specific targets. - """ - if "TRANS" not in self.link_dict: - if self.verbose > 0: - print( - "Warning: TRANS link not found in link dictionary, skipping peptide bonds" - ) - return - - trans_link = self.link_dict["TRANS"] - ptrans_link = self.link_dict.get("PTRANS") - pdb = self.pdb - - # Build peptide bonds using fast builder - bond_result = InterResidueBondBuilder(verbose=self.verbose).build( - pdb, trans_link, device, filter_atom_type="ATOM" - ) - if bond_result: - self.restraints["bond"]["peptide"] = bond_result - if self.verbose > 0: - print( - f"Built {bond_result['indices'].shape[0]} peptide bond restraints" - ) - - # Build peptide angles. - # If PTRANS is available, use it for proline pairs (excludes PRO - # from TRANS to avoid duplicate/conflicting restraints) and TRANS - # for non-proline pairs. Otherwise fall back to TRANS for all. - angle_builder = InterResidueAngleBuilder(verbose=self.verbose) - if ptrans_link is not None: - # Non-proline pairs: TRANS angles - angle_result = angle_builder.build( - pdb, trans_link, device, filter_atom_type="ATOM", - exclude_next_resname="PRO", - ) - # Proline pairs: PTRANS angles (includes C-N-CD) - pro_angle_result = angle_builder.build( - pdb, ptrans_link, device, filter_atom_type="ATOM", - next_resname_filter="PRO", - ) - # Merge results - if angle_result and pro_angle_result: - angle_result = { - "indices": torch.cat([angle_result["indices"], pro_angle_result["indices"]]), - "references": torch.cat([angle_result["references"], pro_angle_result["references"]]), - "sigmas": torch.cat([angle_result["sigmas"], pro_angle_result["sigmas"]]), - } - elif pro_angle_result: - angle_result = pro_angle_result - else: - angle_result = angle_builder.build( - pdb, trans_link, device, filter_atom_type="ATOM" - ) - - if angle_result: - self.restraints["angle"]["peptide"] = angle_result - if self.verbose > 0: - print( - f"Built {angle_result['indices'].shape[0]} peptide angle restraints" - ) - - # Build backbone torsions (phi, psi, omega) - torsion_result = InterResidueTorsionBuilder(verbose=self.verbose).build( - pdb, trans_link, device, filter_atom_type="ATOM" - ) - if torsion_result: - if "phi" in torsion_result: - self.restraints["torsion"]["phi"] = torsion_result["phi"] - if "psi" in torsion_result: - self.restraints["torsion"]["psi"] = torsion_result["psi"] - if "omega" in torsion_result: - self.restraints["torsion"]["omega"] = torsion_result["omega"] - if "ramachandran" in torsion_result: - rama = torsion_result["ramachandran"] - self.register_buffer("_rama_phi_indices", rama["phi_indices"]) - self.register_buffer("_rama_psi_indices", rama["psi_indices"]) - self.register_buffer("_rama_surface_type", rama["surface_type"]) - self._load_rama_surfaces(device) - - # Build peptide planes - plane_result = InterResiduePlaneBuilder(verbose=self.verbose).build( - pdb, trans_link, device, filter_atom_type="ATOM" - ) - if plane_result: - n_planes = 0 - for key, data in plane_result.items(): - n_planes += data["indices"].shape[0] - if self._has_restraint("plane", key): - # Append to existing planes of same atom count - existing = self.restraints["plane"][key] - self.restraints["plane"][key] = { - "indices": torch.cat( - [existing["indices"], data["indices"]], dim=0 - ), - "sigmas": torch.cat( - [existing["sigmas"], data["sigmas"]], dim=0 - ), - } - else: - self.restraints["plane"][key] = data - if self.verbose > 0: - print(f"Built {n_planes} peptide plane restraints") - - def _build_disulfide_restraints(self, device: torch.device): - """Build disulfide bond restraints.""" - if "disulf" not in self.link_dict: - if self.verbose > 1: - print( - "Warning: disulf link not found in link dictionary, skipping disulfide bonds" - ) - return - - disulf_link = self.link_dict["disulf"] - disulf_bonds = disulf_link.get("bonds") - disulf_angles = disulf_link.get("angles") - disulf_torsions = disulf_link.get("torsions") - - if disulf_bonds is None: - return - - # Get SG-SG bond parameters - sg_sg_bond = disulf_bonds[ - (disulf_bonds["atom1"] == "SG") & (disulf_bonds["atom2"] == "SG") - ] - - if len(sg_sg_bond) == 0: - return - - bond_length = float(sg_sg_bond["value"].values[0]) - bond_sigma = float(sg_sg_bond["sigma"].values[0]) - - # Find all SG atoms - pdb = self.pdb - sg_atoms = pdb[(pdb["name"] == "SG") & (pdb["ATOM"] == "ATOM")] - - if len(sg_atoms) == 0: - return - - # Get coordinates and find close pairs - xyz = self.xyz() - sg_indices = sg_atoms["index"].values - sg_coords = xyz[sg_indices] - sg_residues = ( - sg_atoms["chainid"].astype(str) + "_" + sg_atoms["resseq"].astype(str) - ).values - - distances = torch.cdist(sg_coords, sg_coords) - threshold = 2.5 - close_pairs = torch.where((distances < threshold) & (distances > 0.1)) - - valid_pairs = [] - for i, j in zip(close_pairs[0].cpu().numpy(), close_pairs[1].cpu().numpy()): - if i < j and sg_residues[i] != sg_residues[j]: - valid_pairs.append((i, j)) - - if len(valid_pairs) == 0: - return - - # Create builders - bond_builder = InterResidueBondBuilder(verbose=self.verbose) - angle_builder = InterResidueAngleBuilder(verbose=self.verbose) - torsion_builder = InterResidueTorsionBuilder(verbose=self.verbose) - - # Process each disulfide bond - for i_local, j_local in valid_pairs: - sg1_idx = int(sg_indices[i_local]) - sg2_idx = int(sg_indices[j_local]) - - # Add bond - bond_builder.process_disulfide_bond( - sg1_idx, sg2_idx, bond_length, bond_sigma - ) - - # Get residues for angle/torsion restraints - residue1 = pdb[pdb["index"] == sg1_idx].iloc[0] - residue2 = pdb[pdb["index"] == sg2_idx].iloc[0] - - res1_atoms = pdb[ - (pdb["chainid"] == residue1["chainid"]) - & (pdb["resseq"] == residue1["resseq"]) - ] - res2_atoms = pdb[ - (pdb["chainid"] == residue2["chainid"]) - & (pdb["resseq"] == residue2["resseq"]) - ] - - if disulf_angles is not None: - angle_builder.process_disulfide_angles( - res1_atoms, res2_atoms, disulf_angles - ) - - if disulf_torsions is not None: - torsion_builder.process_disulfide_torsions( - res1_atoms, res2_atoms, disulf_torsions - ) - - # Finalize - bond_result = bond_builder.finalize(device) - if bond_result: - self.restraints["bond"]["disulfide"] = bond_result - if self.verbose > 0: - print( - f"Built {bond_result['indices'].shape[0]} disulfide bond restraints" - ) - - angle_result = angle_builder.finalize(device) - if angle_result: - self.restraints["angle"]["disulfide"] = angle_result - if self.verbose > 0: - print( - f"Built {angle_result['indices'].shape[0]} disulfide angle restraints" - ) - - torsion_result = torsion_builder.finalize_disulfide(device) - if torsion_result: - self.restraints["torsion"]["disulfide"] = torsion_result - if self.verbose > 0: - print( - f"Built {torsion_result['indices'].shape[0]} disulfide torsion restraints" - ) - - def _build_link_restraints(self, device: torch.device): - """Build one bond restraint per accepted PDB LINK record. - - Target is the record's ``length`` at sigma=0.02 Å, falling back to 1.5 Å - when blank. LINKs duplicating an auto-detected CYS SG-SG disulfide are - skipped (that builder already added bond + angles + torsions); symmetry-mate - links were dropped earlier in ``extract_link_records``. Each bond joins the - VDW exclusion set via ``_build_exclusion_set``, so the non-bonded term does - not push linked atoms apart. - """ - if self.links is None or len(self.links) == 0: - return - - pdb = self.pdb - - # Already-bonded SG-SG pairs from auto-disulfide detection. - disulf = self.restraints.get("bond", {}).get("disulfide") - existing_disulf_pairs = set() - if disulf is not None and "indices" in disulf: - for i, j in disulf["indices"].cpu().numpy(): - existing_disulf_pairs.add((int(min(i, j)), int(max(i, j)))) - - bond_builder = InterResidueBondBuilder(verbose=self.verbose) - n_skipped_unresolved = 0 - n_skipped_dedup = 0 - - for _, link in self.links.iterrows(): - idx1 = self._lookup_link_atom( - pdb, - chainid=link["chainid1"], - resseq=int(link["resseq1"]), - icode=link["icode1"], - resname=link["resname1"], - name=link["name1"], - altloc=link["altloc1"], - ) - idx2 = self._lookup_link_atom( - pdb, - chainid=link["chainid2"], - resseq=int(link["resseq2"]), - icode=link["icode2"], - resname=link["resname2"], - name=link["name2"], - altloc=link["altloc2"], - ) - - if idx1 is None or idx2 is None: - n_skipped_unresolved += 1 - if self.verbose > 1: - print( - f"Warning: LINK atom not found " - f"({link['chainid1']}/{link['resname1']}{link['resseq1']}/" - f"{link['name1']} -- " - f"{link['chainid2']}/{link['resname2']}{link['resseq2']}/" - f"{link['name2']}); skipping." - ) - continue - - if idx1 == idx2: - n_skipped_unresolved += 1 - continue - - pair = (min(idx1, idx2), max(idx1, idx2)) - if pair in existing_disulf_pairs: - n_skipped_dedup += 1 - continue - - length = link["length"] - if not (isinstance(length, (int, float)) and length == length and length > 0): - length = 1.5 - bond_builder.process_disulfide_bond(idx1, idx2, float(length), 0.02) - - bond_result = bond_builder.finalize(device) - if bond_result: - self.restraints["bond"]["link"] = bond_result - if self.verbose > 0: - print( - f"Built {bond_result['indices'].shape[0]} LINK bond restraints" - + ( - f" (skipped {n_skipped_dedup} disulfide-dup," - f" {n_skipped_unresolved} unresolved)" - if (n_skipped_dedup or n_skipped_unresolved) - else "" - ) - ) @staticmethod def _lookup_link_atom( @@ -1012,35 +440,6 @@ def _lookup_link_atom( return int(hit.iloc[0]["index"]) return int(sel.iloc[0]["index"]) - def _build_exclusion_set(self): - """Build set of atom pairs to exclude from VDW calculations.""" - exclusions = set() - - # 1-2: Direct bonds - for origin in self.restraints.get("bond", {}).keys(): - indices = self.restraints["bond"][origin].get("indices") - if indices is not None and len(indices) > 0: - idx_np = indices.cpu().numpy() - for i1, i2 in idx_np: - exclusions.add((int(min(i1, i2)), int(max(i1, i2)))) - - # 1-3: Angles - for origin in self.restraints.get("angle", {}).keys(): - indices = self.restraints["angle"][origin].get("indices") - if indices is not None and len(indices) > 0: - idx_np = indices.cpu().numpy() - for i1, i2, i3 in idx_np: - exclusions.add((int(min(i1, i3)), int(max(i1, i3)))) - - # 1-4: Torsions - for origin in self.restraints.get("torsion", {}).keys(): - indices = self.restraints["torsion"][origin].get("indices") - if indices is not None and len(indices) > 0: - idx_np = indices.cpu().numpy() - for i1, i2, i3, i4 in idx_np: - exclusions.add((int(min(i1, i4)), int(max(i1, i4)))) - - return exclusions def _find_nearby_pairs_spatial_hash(self, xyz, cutoff=6.0): """Atom pairs within ``cutoff`` of each other, as (M, 2) rows with i < j. @@ -1414,8 +813,8 @@ def vdw_radii_cpu(): if has_symmetry: from torchref.restraints.neighbor_search import build_vdw_restraints_gpu - exclusions = self._build_exclusion_set() - self.restraints["vdw"] = build_vdw_restraints_gpu( + exclusions = self.topology.atoms.exclusions_from_restraint_edges() + self._vdw = build_vdw_restraints_gpu( xyz_fn=xyz_cpu, vdw_radii_fn=vdw_radii_cpu, cell=cell_cpu, @@ -1434,6 +833,11 @@ def vdw_radii_cpu(): use_spatial_hash=use_spatial_hash, ) + # Publish the new pair list before anything reads it back below. Unlike the + # geometry edges it is not derived from the topology, so it is held separately + # and re-inserted here and by _rebuild_entries. + self._entries["vdw"] = self._vdw + # Build riding hydrogen topology and precompute candidate pairs from torchref.restraints.hydrogen_topology import ( build_hydrogen_topology, @@ -1497,7 +901,7 @@ def _build_vdw_restraints_legacy( ): """Legacy VDW restraint builder (no symmetry or CPU fallback).""" - exclusions = self._build_exclusion_set() + exclusions = self.topology.atoms.exclusions_from_restraint_edges() vdw_radii = self.get_vdw_radii() xyz = self.xyz() device = xyz.device @@ -1542,7 +946,7 @@ def _build_vdw_restraints_legacy( } if len(nearby_pairs) == 0: - self.restraints["vdw"] = empty_result + self._vdw = empty_result return pairs_np = nearby_pairs.cpu().numpy() @@ -1660,7 +1064,7 @@ def _build_vdw_restraints_legacy( final_offsets = final_offsets[keep_mask] if len(final_i1) == 0: - self.restraints["vdw"] = empty_result + self._vdw = empty_result return # Compute min distances using VDW radii of ASU source atoms. @@ -1669,7 +1073,7 @@ def _build_vdw_restraints_legacy( # Store results final_pairs = np.stack([final_i1, final_i2], axis=1) - self.restraints["vdw"] = { + self._vdw = { "indices": torch.tensor(final_pairs, dtype=torch.long, device=device), "min_distances": torch.tensor( min_distances, dtype=get_float_dtype(), device=device @@ -1694,8 +1098,8 @@ def _build_vdw_restraints_legacy( msg += f", {n_sym_count} symmetry contacts" print(msg) - # Device movement is handled automatically by TensorDict (registered as _tensor_storage) - # through PyTorch's Module.to(), cuda(), and cpu() methods + # Device movement goes through DeviceMixin: the topology and the value tensors are + # walked and moved, and _apply re-slices the derived entry views afterwards. def summary(self): """Print a detailed summary of all restraints.""" @@ -1797,46 +1201,7 @@ def get_count(rtype, origin): f"torsions={n_torsions}, peptide_bonds={n_bonds_peptide})" ) - def _get_all_indices(self, restraint_type, keys_to_merge=None): - """Concatenate the indices of one restraint type across origins, or None. - - ``keys_to_merge`` restricts to those origins; None means all. Rows follow - origin iteration order, so pair this only with - :meth:`_get_all_property` calls made with the same ``keys_to_merge``. - """ - indices_list = [] - for origin, data in self.restraints.get(restraint_type, {}).items(): - indices = data.get("indices") - if indices is not None: - if keys_to_merge is None: - indices_list.append(indices) - elif origin in keys_to_merge: - indices_list.append(indices) - - if not indices_list: - return None - - return torch.cat(indices_list, dim=0) - - def _get_all_property(self, restraint_type, property_name, keys_to_merge=None): - """Concatenate one property ('references'/'sigmas'/'periods') across origins. - - Returns None if no origin carries it. Row order matches - :meth:`_get_all_indices` for the same ``keys_to_merge``. - """ - values_list = [] - for origin, data in self.restraints.get(restraint_type, {}).items(): - values = data.get(property_name) - if values is not None: - if keys_to_merge is None: - values_list.append(values) - elif origin in keys_to_merge: - values_list.append(values) - - if not values_list: - return None - return torch.cat(values_list, dim=0) def bond_lengths(self, idx, xyz: torch.Tensor = None): """ @@ -1863,17 +1228,23 @@ def bond_lengths(self, idx, xyz: torch.Tensor = None): return torch.linalg.norm(pos2 - pos1, dim=-1) def copy(self): - """ - Create a deep copy of the Restraints object. + """An independent copy, sharing no state with this one. + + The entry views are re-sliced afterwards rather than left as deep-copied + tensors: ``deepcopy`` duplicates a view and the block it points into as two + unrelated tensors, so the copy would still hold the right values but would no + longer alias, and an in-place edit to one would stop being visible through the + other. Returns ------- - Restraints - A deep copy of this Restraints instance. + RestraintsNew """ import copy - return copy.deepcopy(self) + duplicate = copy.deepcopy(self) + duplicate._rebuild_entries() + return duplicate def bond_deviations(self, xyz: torch.Tensor = None): """ @@ -2025,48 +1396,20 @@ def nll_angles(self, xyz: torch.Tensor = None): return gaussian_nll(deviations, sigmas) def cat_dict(self): - """ - Concatenate restraint origins into combined 'all' keys. + """Ensure the combined ``all`` groups are present. Idempotent. + + They are assembled with everything else at build time, so this normally has + nothing to do; it exists because the geometry targets guard their reads with + ``if "all" not in ...`` and call it when the guard trips. - Creates restraints['bond']['all'], restraints['angle']['all'], and - restraints['torsion']['all']. Bond and angle 'all' include every - origin; torsion 'all' includes only the 'intra' and 'disulfide' - origins (phi/psi have no reference values or sigmas, and omega is - handled by a dedicated OmegaTarget). + The previous implementation concatenated the origins on each call, and because + writing ``restraints['bond']['all']`` also registered ``'all'`` as an origin, a + second call folded the combined group into itself and doubled every restraint. + Deriving the group from the topology instead makes that impossible: ``all`` is a + span of the edge block, never an origin in its own right. """ - self.restraints["bond"]["all"] = { - "indices": self._get_all_indices("bond"), - "references": self._get_all_property("bond", "references"), - "sigmas": self._get_all_property("bond", "sigmas"), - } - self.restraints["angle"]["all"] = { - "indices": self._get_all_indices("angle"), - "references": self._get_all_property("angle", "references"), - "sigmas": self._get_all_property("angle", "sigmas"), - } - # Note: phi/psi origins are excluded because they have no reference - # values or sigmas (conformationally free). Omega is excluded here - # because it is handled by a dedicated OmegaTarget that uses a - # cis/trans von Mises mixture model. - _torsion_origins = ["intra", "disulfide"] - self.restraints["torsion"]["all"] = { - "indices": self._get_all_indices("torsion", _torsion_origins), - "references": self._get_all_property( - "torsion", "references", _torsion_origins - ), - "sigmas": self._get_all_property( - "torsion", "sigmas", _torsion_origins - ), - "periods": self._get_all_property( - "torsion", "periods", _torsion_origins - ), - } - # Cache max period to avoid .item() GPU sync every iteration - periods = self.restraints["torsion"]["all"]["periods"] - if periods is not None and periods.numel() > 0: - self._torsion_max_period = int(periods.max().item()) - else: - self._torsion_max_period = 1 + if self.topology is not None and "all" not in self._entries.get("bond", {}): + self._rebuild_entries() def torsions(self, idx, xyz: torch.Tensor = None): """ diff --git a/torchref/topology/__init__.py b/torchref/topology/__init__.py index ca8e9405..78cc0dea 100644 --- a/torchref/topology/__init__.py +++ b/torchref/topology/__init__.py @@ -14,9 +14,10 @@ """ from .atom_graph import AtomGraph -from .build import build_topology +from .build import build_topology, build_topology_with_values from .edges import ORIGIN_ORDER, EdgeBlock from .residue_graph import ResidueGraph +from .restraint_sets import assemble_entries, max_period from .templates import resolve_template_keys from .topology import Topology @@ -27,5 +28,8 @@ "EdgeBlock", "ORIGIN_ORDER", "build_topology", + "build_topology_with_values", + "assemble_entries", + "max_period", "resolve_template_keys", ] diff --git a/torchref/topology/build.py b/torchref/topology/build.py index f6225f84..bce5b695 100644 --- a/torchref/topology/build.py +++ b/torchref/topology/build.py @@ -26,13 +26,14 @@ match_torsions_numba, ) from torchref.topology.atom_graph import AtomGraph -from torchref.topology.edges import EdgeBlock +from torchref.topology.edges import EdgeBlock, assemble_origins from torchref.topology.residue_graph import ( ResidueGraph, build_residue_nodes, find_disulfide_links, find_peptide_links, ) +from torchref.topology.restraint_sets import to_tensor from torchref.topology.templates import resolve_template_keys from torchref.topology.topology import Topology @@ -112,13 +113,20 @@ def _match_intra( nodes: Dict[str, np.ndarray], template_key: np.ndarray, pp_cif: PreprocessedCIF, -) -> Dict[str, np.ndarray]: - """Intra-residue edge index arrays. +) -> Tuple[Dict[str, np.ndarray], Dict[str, Dict[str, np.ndarray]]]: + """Intra-residue edges and the ideal values that belong to them. Keyed ``bonds`` / ``angles`` / ``torsions`` / ``chirals``. Emitted only where every named atom of a library restraint is present in the conformation, which is the condition the matchers apply. + + Returns + ------- + indices : dict + ``{kind: (E, k) array}``. + values : dict + ``{kind: {property: (E,) array}}``, accumulated row-for-row with the indices. """ acc: Dict[str, List[np.ndarray]] = { "bonds": [], @@ -126,6 +134,12 @@ def _match_intra( "torsions": [], "chirals": [], } + val: Dict[str, Dict[str, List[np.ndarray]]] = { + "bonds": {"references": [], "sigmas": []}, + "angles": {"references": [], "sigmas": []}, + "torsions": {"references": [], "sigmas": [], "periods": []}, + "chirals": {"ideal_volumes": [], "sigmas": []}, + } work = {k: np.zeros(_WORK, dtype=np.int64) for k in ("i1", "i2", "i3", "i4", "per")} work["f1"] = np.zeros(_WORK, dtype=np.float64) work["f2"] = np.zeros(_WORK, dtype=np.float64) @@ -165,6 +179,8 @@ def _match_intra( acc["bonds"].append( np.column_stack([work["i1"][:n].copy(), work["i2"][:n].copy()]) ) + val["bonds"]["references"].append(work["f1"][:n].copy()) + val["bonds"]["sigmas"].append(work["f2"][:n].copy()) if key in pp_cif.angles: a = pp_cif.angles[key] n = match_angles_numba( @@ -191,6 +207,8 @@ def _match_intra( ] ) ) + val["angles"]["references"].append(work["f1"][:n].copy()) + val["angles"]["sigmas"].append(work["f2"][:n].copy()) if key in pp_cif.torsions: t = pp_cif.torsions[key] n = match_torsions_numba( @@ -222,6 +240,9 @@ def _match_intra( ] ) ) + val["torsions"]["references"].append(work["f1"][:n].copy()) + val["torsions"]["sigmas"].append(work["f2"][:n].copy()) + val["torsions"]["periods"].append(work["per"][:n].copy()) if key in pp_cif.chirals: c = pp_cif.chirals[key] n = match_chirals_numba( @@ -251,12 +272,30 @@ def _match_intra( ] ) ) + # Ideal volume is the sign times a typical tetrahedral volume. A + # sign of 0 ('both' / 'either') stays exactly 0, which the chiral + # target reads as an achiral centre and restrains |volume| instead. + val["chirals"]["ideal_volumes"].append(work["f1"][:n].copy() * 2.5) + val["chirals"]["sigmas"].append(work["f2"][:n].copy()) arity = {"bonds": 2, "angles": 3, "torsions": 4, "chirals": 4} - return { + indices = { k: (np.concatenate(v, axis=0) if v else np.zeros((0, arity[k]), dtype=np.int64)) for k, v in acc.items() } + values: Dict[str, Dict[str, np.ndarray]] = {} + for kind, properties in val.items(): + joined: Dict[str, np.ndarray] = {} + for prop, chunks in properties.items(): + if not chunks: + continue + array = np.concatenate(chunks) + if prop == "sigmas": + # A zero sigma divides by zero in the loss, so it is floored. + array = np.where(array == 0, 1e-4, array) + joined[prop] = array + values[kind] = joined + return indices, values def _match_intra_planes( @@ -264,13 +303,15 @@ def _match_intra_planes( nodes: Dict[str, np.ndarray], template_key: np.ndarray, pp_cif: PreprocessedCIF, -) -> Dict[int, np.ndarray]: - """Intra-residue plane index arrays grouped by how many atoms survived matching. +) -> Tuple[Dict[int, np.ndarray], Dict[int, Dict[str, np.ndarray]]]: + """Intra-residue planes grouped by how many atoms survived matching. Missing atoms are dropped rather than voiding the plane; a plane is kept once at - least three of its atoms are present, so its arity depends on the model. + least three of its atoms are present, so its arity depends on the model. Sigmas are + per atom, not per plane, so they carry the same ``(E, k)`` shape as the indices. """ by_size: Dict[int, List[np.ndarray]] = {} + sigmas_by_size: Dict[int, List[np.ndarray]] = {} for r in range(len(nodes["chain"])): key = str(template_key[r]) if key not in pp_cif.planes: @@ -280,58 +321,105 @@ def _match_intra_planes( # Last-wins on a duplicate name, matching PlaneRestraintBuilder. name_to_idx = dict(zip(names, indices)) for plane in pp_cif.planes[key]: - present = [ - name_to_idx[nm] for nm in plane["atoms"] if nm in name_to_idx - ] + present = [] + present_sigmas = [] + for position, atom_name in enumerate(plane["atoms"]): + if atom_name in name_to_idx: + present.append(name_to_idx[atom_name]) + present_sigmas.append(plane["sigmas"][position]) if len(present) >= 3: by_size.setdefault(len(present), []).append( np.asarray(present, dtype=np.int64) ) - return {n: np.stack(rows, axis=0) for n, rows in by_size.items()} + sigmas_by_size.setdefault(len(present), []).append( + np.asarray(present_sigmas, dtype=np.float64) + ) + + indices_out = {n: np.stack(rows, axis=0) for n, rows in by_size.items()} + values_out = { + n: { + "sigmas": np.where( + np.stack(rows, axis=0) == 0, 1e-4, np.stack(rows, axis=0) + ) + } + for n, rows in sigmas_by_size.items() + } + return indices_out, values_out def _inter_residue_edges( pdb: pd.DataFrame, link_dict: Optional[Dict], verbose: int, -) -> Dict[str, Dict[str, np.ndarray]]: - """Peptide edge index arrays from the inter-residue builders. +) -> Tuple[Dict[str, Dict[str, np.ndarray]], Dict[str, Dict], Dict[str, Dict]]: + """Peptide edges, their values, and the Ramachandran pairing, from the builders. Reuses ``InterResidue*Builder`` rather than reimplementing the link geometry. - Returns ``{edge type: {origin: (E, k) array}}``. + + Returns + ------- + indices : dict + ``{edge type: {origin: (E, k) array}}``. + values : dict + ``{edge type: {origin: {property: array}}}`` -- every property a builder + returned besides the indices, so ``omega``'s ``is_proline`` comes along without + being named here. + extras : dict + Non-edge products of the same pass, currently the ``ramachandran`` phi/psi + pairing and its surface types. """ - out: Dict[str, Dict[str, np.ndarray]] = { + indices: Dict[str, Dict[str, np.ndarray]] = { "bond": {}, "angle": {}, "torsion": {}, "plane": {}, } + values: Dict[str, Dict] = {"bond": {}, "angle": {}, "torsion": {}, "plane": {}} + extras: Dict[str, Dict] = {} if not link_dict or "TRANS" not in link_dict: - return out + return indices, values, extras + cpu = torch.device("cpu") trans = link_dict["TRANS"] ptrans = link_dict.get("PTRANS") + def split(group): + """A builder group as ``(indices array, {property: array})``.""" + rows = group["indices"].cpu().numpy() + rest = { + prop: tensor.cpu().numpy() + for prop, tensor in group.items() + if prop != "indices" and tensor is not None + } + return rows, rest + bond = InterResidueBondBuilder(verbose=verbose).build( pdb, trans, cpu, filter_atom_type="ATOM" ) if bond: - out["bond"]["peptide"] = bond["indices"].cpu().numpy() + indices["bond"]["peptide"], values["bond"]["peptide"] = split(bond) ab = InterResidueAngleBuilder(verbose=verbose) if ptrans is not None: - non_pro = ab.build( - pdb, trans, cpu, filter_atom_type="ATOM", exclude_next_resname="PRO" - ) - pro = ab.build( - pdb, ptrans, cpu, filter_atom_type="ATOM", next_resname_filter="PRO" - ) - chunks = [r["indices"].cpu().numpy() for r in (non_pro, pro) if r] + # PTRANS carries the extra C(i-1)-N-CD angle, so proline pairs are built from it + # and excluded from the TRANS pass to avoid two restraints on the same atoms. + groups = [ + ab.build( + pdb, trans, cpu, filter_atom_type="ATOM", exclude_next_resname="PRO" + ), + ab.build( + pdb, ptrans, cpu, filter_atom_type="ATOM", next_resname_filter="PRO" + ), + ] else: - r = ab.build(pdb, trans, cpu, filter_atom_type="ATOM") - chunks = [r["indices"].cpu().numpy()] if r else [] - if chunks: - out["angle"]["peptide"] = np.concatenate(chunks, axis=0) + groups = [ab.build(pdb, trans, cpu, filter_atom_type="ATOM")] + parts = [split(g) for g in groups if g] + if parts: + indices["angle"]["peptide"] = np.concatenate([p[0] for p in parts], axis=0) + shared = set.intersection(*(set(p[1]) for p in parts)) + values["angle"]["peptide"] = { + prop: np.concatenate([p[1][prop] for p in parts]) for prop in shared + } tors = InterResidueTorsionBuilder(verbose=verbose).build( pdb, trans, cpu, filter_atom_type="ATOM" @@ -339,16 +427,19 @@ def _inter_residue_edges( if tors: for origin in ("phi", "psi", "omega"): if origin in tors: - out["torsion"][origin] = tors[origin]["indices"].cpu().numpy() + indices["torsion"][origin], values["torsion"][origin] = split( + tors[origin] + ) + if "ramachandran" in tors: + extras["ramachandran"] = tors["ramachandran"] planes = InterResiduePlaneBuilder(verbose=verbose).build( pdb, trans, cpu, filter_atom_type="ATOM" ) if planes: - out["plane"] = { - key: data["indices"].cpu().numpy() for key, data in planes.items() - } - return out + for key, group in planes.items(): + indices["plane"][key], values["plane"][key] = split(group) + return indices, values, extras def _origins( @@ -376,8 +467,8 @@ def _disulfide_edges( pairs: Sequence[Tuple[int, int]], link_dict: Optional[Dict], verbose: int, -) -> Dict[str, np.ndarray]: - """Bond, angle and torsion edges for the detected disulfide links. +) -> Tuple[Dict[str, np.ndarray], Dict[str, Dict[str, np.ndarray]]]: + """Bond, angle and torsion edges for the detected disulfide links, with values. Drives the ``InterResidue*Builder`` disulfide paths from the residue graph's ``disulf`` edges, so the link geometry comes from the ``disulf`` dictionary entry @@ -385,20 +476,22 @@ def _disulfide_edges( Returns ------- - dict + indices : dict ``{'bond'|'angle'|'torsion': (E, k) array}``, omitting types with no edges. + values : dict + ``{edge type: {property: array}}`` for the same edges. """ out: Dict[str, np.ndarray] = {} if not pairs or not link_dict or "disulf" not in link_dict: - return out + return out, {} disulf = link_dict["disulf"] bonds = disulf.get("bonds") if bonds is None: - return out + return out, {} sg_sg = bonds[(bonds["atom1"] == "SG") & (bonds["atom2"] == "SG")] if len(sg_sg) == 0: - return out + return out, {} length = float(sg_sg["value"].values[0]) sigma = float(sg_sg["sigma"].values[0]) @@ -426,16 +519,21 @@ def _disulfide_edges( atoms_a, atoms_b, disulf["torsions"] ) - bond_result = bond_builder.finalize(cpu) - if bond_result: - out["bond"] = bond_result["indices"].cpu().numpy() - angle_result = angle_builder.finalize(cpu) - if angle_result: - out["angle"] = angle_result["indices"].cpu().numpy() - torsion_result = torsion_builder.finalize_disulfide(cpu) - if torsion_result: - out["torsion"] = torsion_result["indices"].cpu().numpy() - return out + values: Dict[str, Dict[str, np.ndarray]] = {} + for edge_type, group in ( + ("bond", bond_builder.finalize(cpu)), + ("angle", angle_builder.finalize(cpu)), + ("torsion", torsion_builder.finalize_disulfide(cpu)), + ): + if not group: + continue + out[edge_type] = group["indices"].cpu().numpy() + values[edge_type] = { + prop: tensor.cpu().numpy() + for prop, tensor in group.items() + if prop != "indices" and tensor is not None + } + return out, values def _link_record_edges( @@ -443,7 +541,7 @@ def _link_record_edges( links, disulfide_bonds: Optional[np.ndarray], verbose: int, -) -> Tuple[np.ndarray, List[Tuple[int, int]]]: +) -> Tuple[np.ndarray, List[Tuple[int, int]], Dict[str, np.ndarray]]: """Bond edges for the accepted ``LINK`` records, and the atom pairs they join. Atom resolution goes through ``RestraintsNew._lookup_link_atom`` so a record is @@ -457,11 +555,14 @@ def _link_record_edges( Shape ``(L, 2)``; empty when nothing resolved. atom_pairs : list of tuple of int The same pairs, for lifting to residue-level link edges. + values : dict + ``references`` from each record's ``length`` (1.5 A where blank or unusable) and + a fixed ``sigmas`` of 0.02 A. """ from torchref.restraints.restraints import RestraintsNew if links is None or len(links) == 0: - return np.zeros((0, 2), dtype=np.int64), [] + return np.zeros((0, 2), dtype=np.int64), [], {} existing = set() if disulfide_bonds is not None: @@ -469,6 +570,7 @@ def _link_record_edges( existing.add((min(int(a), int(b)), max(int(a), int(b)))) rows: List[Tuple[int, int]] = [] + lengths: List[float] = [] n_unresolved = 0 for _, link in links.iterrows(): idx1 = RestraintsNew._lookup_link_atom( @@ -496,12 +598,49 @@ def _link_record_edges( if pair in existing: continue rows.append((idx1, idx2)) + length = link["length"] + usable = isinstance(length, (int, float)) and length == length and length > 0 + lengths.append(float(length) if usable else 1.5) if verbose > 1 and n_unresolved: print(f"{n_unresolved} LINK records did not resolve to a pair of atoms") if not rows: - return np.zeros((0, 2), dtype=np.int64), [] - return np.asarray(rows, dtype=np.int64), rows + return np.zeros((0, 2), dtype=np.int64), [], {} + values = { + "references": np.asarray(lengths, dtype=np.float64), + "sigmas": np.full(len(rows), 0.02, dtype=np.float64), + } + return np.asarray(rows, dtype=np.int64), rows, values + + +def _block_with_values( + per_origin: Dict[str, np.ndarray], + payload: Dict[str, Dict[str, np.ndarray]], + arity: int, + edge_type: str, + device, +) -> Tuple[EdgeBlock, Dict[str, Dict[str, torch.Tensor]]]: + """One canonical edge block plus its per-origin value tensors. + + The block and the values come out of a single :func:`assemble_origins` call, so the + same permutation is applied to both -- which is the only thing keeping a sigma + attached to the edge it belongs to. + """ + indices, bounds, sorted_payload = assemble_origins( + per_origin, arity, edge_type, payload + ) + block = EdgeBlock( + indices=torch.as_tensor(indices, dtype=torch.int64, device=device), + origin_bounds=bounds, + ) + values = { + origin: { + prop: to_tensor(array, prop, device=device) + for prop, array in properties.items() + } + for origin, properties in sorted_payload.items() + } + return block, values def build_topology( @@ -514,6 +653,34 @@ def build_topology( device=None, verbose: int = 0, ) -> Topology: + """Build a topology, discarding the restraint values built along the way. + + See :func:`build_topology_with_values` for the parameters; this is the connectivity + half on its own, for callers that need the graph and no ideal geometry. + """ + topology, _, _ = build_topology_with_values( + pdb, + cif_dict, + link_dict=link_dict, + link_list=link_list, + links=links, + xyz=xyz, + device=device, + verbose=verbose, + ) + return topology + + +def build_topology_with_values( + pdb: pd.DataFrame, + cif_dict: Dict, + link_dict: Optional[Dict] = None, + link_list=None, + links=None, + xyz: Optional[torch.Tensor] = None, + device=None, + verbose: int = 0, +) -> Tuple[Topology, Dict[str, Dict], Dict[str, Dict]]: """Build a topology from an atom table and the restraint dictionaries. Parameters @@ -540,7 +707,14 @@ def build_topology( Returns ------- - Topology + topology : Topology + The connectivity. + values : dict + ``{edge_type: {origin: {property: tensor}}}`` for bonds, angles and torsions; + ``{'chiral': {property: tensor}}`` and ``{'plane': {size: {property: tensor}}}`` + for the two types that carry no origin. Row-aligned to the edge blocks. + extras : dict + Products of the same pass that are not edges -- currently ``ramachandran``. """ cols = _atom_columns(pdb) nodes = build_residue_nodes( @@ -570,9 +744,11 @@ def build_topology( ) pp_cif = PreprocessedCIF(comp_dict) - intra = _match_intra(cols, nodes, template_key, pp_cif) - intra_planes = _match_intra_planes(cols, nodes, template_key, pp_cif) - inter = _inter_residue_edges(pdb, link_dict, verbose) + intra, intra_values = _match_intra(cols, nodes, template_key, pp_cif) + intra_planes, intra_plane_values = _match_intra_planes( + cols, nodes, template_key, pp_cif + ) + inter, inter_values, extras = _inter_residue_edges(pdb, link_dict, verbose) residue_of_row = {} for r in range(n_res): @@ -581,6 +757,7 @@ def build_topology( disulfide_pairs: List[Tuple[int, int]] = [] disulfide: Dict[str, np.ndarray] = {} + disulfide_values: Dict[str, Dict[str, np.ndarray]] = {} if xyz is not None: sg_rows = [ row @@ -588,11 +765,11 @@ def build_topology( if cols["name"][row] == "SG" and cols["record"][row] == "ATOM" ] disulfide_pairs = find_disulfide_links(sg_rows, residue_of_row, xyz) - disulfide = _disulfide_edges( + disulfide, disulfide_values = _disulfide_edges( pdb, nodes, cols, residue_of_row, disulfide_pairs, link_dict, verbose ) - link_edges, link_atom_pairs = _link_record_edges( + link_edges, link_atom_pairs, link_values = _link_record_edges( pdb, links, disulfide.get("bond"), verbose ) @@ -632,20 +809,84 @@ def build_topology( ) plane_blocks: Dict[int, EdgeBlock] = {} + plane_values: Dict[int, Dict[str, torch.Tensor]] = {} plane_sizes = set(intra_planes) for key in inter.get("plane", {}): plane_sizes.add(int(str(key).split("_")[0])) for size in sorted(plane_sizes): - per_origin = {} + per_origin: Dict[str, np.ndarray] = {} + payload: Dict[str, Dict[str, np.ndarray]] = {} if size in intra_planes: per_origin["intra"] = intra_planes[size] + payload["intra"] = intra_plane_values.get(size, {}) peptide = inter.get("plane", {}).get(f"{size}_atoms") if peptide is not None and len(peptide): per_origin["peptide"] = peptide - if per_origin: - plane_blocks[size] = EdgeBlock.from_origins( - per_origin, size, "plane", device=device + payload["peptide"] = inter_values["plane"].get(f"{size}_atoms", {}) + if not per_origin: + continue + block, per_origin_values = _block_with_values( + per_origin, payload, size, "plane", device + ) + plane_blocks[size] = block + # Planes carry no origin downstream, so the origins are concatenated back into + # one group -- in block order, which is what the block's own layout already is. + plane_values[size] = { + prop: torch.cat( + [ + per_origin_values[o][prop] + for o in block.origins() + if prop in per_origin_values.get(o, {}) + ] ) + for prop in {p for v in per_origin_values.values() for p in v} + } + + bond_block, bond_values = _block_with_values( + _origins( + intra["bonds"], + {**inter["bond"], "link": link_edges}, + disulfide.get("bond"), + ), + { + "intra": intra_values["bonds"], + **inter_values["bond"], + "link": link_values, + "disulfide": disulfide_values.get("bond", {}), + }, + 2, + "bond", + device, + ) + angle_block, angle_values = _block_with_values( + _origins(intra["angles"], inter["angle"], disulfide.get("angle")), + { + "intra": intra_values["angles"], + **inter_values["angle"], + "disulfide": disulfide_values.get("angle", {}), + }, + 3, + "angle", + device, + ) + torsion_block, torsion_values = _block_with_values( + _origins(intra["torsions"], inter["torsion"], disulfide.get("torsion")), + { + "intra": intra_values["torsions"], + **inter_values["torsion"], + "disulfide": disulfide_values.get("torsion", {}), + }, + 4, + "torsion", + device, + ) + chiral_block, chiral_values = _block_with_values( + _origins(intra["chirals"], {}, None), + {"intra": intra_values["chirals"]}, + 4, + "chiral", + device, + ) atoms = AtomGraph( name=cols["name"], @@ -659,35 +900,22 @@ def build_topology( dtype=torch.int64, device=device, ), - bonds=EdgeBlock.from_origins( - _origins( - intra["bonds"], - {**inter["bond"], "link": link_edges}, - disulfide.get("bond"), - ), - 2, - "bond", - device=device, - ), - angles=EdgeBlock.from_origins( - _origins(intra["angles"], inter["angle"], disulfide.get("angle")), - 3, - "angle", - device=device, - ), - torsions=EdgeBlock.from_origins( - _origins(intra["torsions"], inter["torsion"], disulfide.get("torsion")), - 4, - "torsion", - device=device, - ), - chirals=EdgeBlock.from_origins( - {"intra": intra["chirals"]}, 4, "chiral", device=device - ), + bonds=bond_block, + angles=angle_block, + torsions=torsion_block, + chirals=chiral_block, planes=plane_blocks, ) - return Topology(residues=residues, atoms=atoms) + values: Dict[str, Dict] = { + "bond": bond_values, + "angle": angle_values, + "torsion": torsion_values, + # Chirals carry no origin downstream, so the single origin is unwrapped. + "chiral": chiral_values.get("intra", {}), + "plane": plane_values, + } + return Topology(residues=residues, atoms=atoms), values, extras -__all__ = ["build_topology"] +__all__ = ["build_topology", "build_topology_with_values"] diff --git a/torchref/topology/edges.py b/torchref/topology/edges.py index 654d300d..50e2d960 100644 --- a/torchref/topology/edges.py +++ b/torchref/topology/edges.py @@ -25,7 +25,9 @@ ORIGIN_ORDER: Dict[str, Tuple[str, ...]] = { "bond": ("intra", "peptide", "disulfide", "link"), "angle": ("intra", "peptide", "disulfide"), - "torsion": ("intra", "phi", "psi", "omega", "disulfide"), + # intra and disulfide lead because together they are the ``all`` group the + # torsion target reads; adjacent origins make that group a view rather than a copy. + "torsion": ("intra", "disulfide", "phi", "psi", "omega"), "chiral": ("intra",), "plane": ("intra", "peptide"), } @@ -50,6 +52,82 @@ def _lexsort_rows(rows: np.ndarray) -> np.ndarray: return np.lexsort(tuple(rows[:, c] for c in reversed(range(rows.shape[1])))) +def assemble_origins( + per_origin: Dict[str, np.ndarray], + arity: int, + edge_type: str, + payload: Dict[str, Dict[str, np.ndarray]] = None, +) -> Tuple[np.ndarray, Dict[str, Tuple[int, int]], Dict[str, Dict[str, np.ndarray]]]: + """Lay origins out in canonical order, carrying per-edge values through the sort. + + Origins follow :data:`ORIGIN_ORDER` for ``edge_type``, and rows within an origin + are sorted lexicographically. Anything in ``payload`` is permuted by the same order, + so a value array stays aligned row-for-row with the indices it belongs to, which + is the whole reason values cannot be concatenated separately. + + Parameters + ---------- + per_origin : dict + ``{origin: (E_o, k) integer array}``. Empty entries are skipped. + arity : int + Atoms per edge. + edge_type : str + Key into :data:`ORIGIN_ORDER`. + payload : dict, optional + ``{origin: {property: array}}``, each array indexed by row on axis 0. A property + need not be present for every origin -- ``phi`` and ``psi`` carry no reference + value or sigma, and must not acquire one here. + + Returns + ------- + indices : numpy.ndarray + Shape ``(E, k)``, canonical order. + bounds : dict + ``{origin: (start, end)}``, contiguous and covering the block. + sorted_payload : dict + ``{origin: {property: array}}``, permuted to match ``indices``. + """ + order = ORIGIN_ORDER.get(edge_type, tuple(sorted(per_origin))) + unknown = set(per_origin) - set(order) + if unknown: + raise ValueError( + f"{edge_type}: origins {sorted(unknown)} are not in ORIGIN_ORDER" + f"[{edge_type!r}] = {order}. Add them there so the layout stays " + f"deterministic." + ) + + payload = payload or {} + chunks: List[np.ndarray] = [] + bounds: Dict[str, Tuple[int, int]] = {} + sorted_payload: Dict[str, Dict[str, np.ndarray]] = {} + cursor = 0 + + for origin in order: + rows = per_origin.get(origin) + if rows is None or len(rows) == 0: + continue + rows = np.asarray(rows, dtype=np.int64).reshape(-1, arity) + permutation = _lexsort_rows(rows) + chunks.append(rows[permutation]) + bounds[origin] = (cursor, cursor + len(rows)) + cursor += len(rows) + + origin_payload = payload.get(origin) or {} + if origin_payload: + sorted_payload[origin] = { + prop: np.asarray(values)[permutation] + for prop, values in origin_payload.items() + if values is not None + } + + indices = ( + np.concatenate(chunks, axis=0) + if chunks + else np.zeros((0, arity), dtype=np.int64) + ) + return indices, bounds, sorted_payload + + @dataclass(eq=False, repr=False) class EdgeBlock(DeviceMixin): """One edge type's index block plus its per-origin bounds. @@ -108,34 +186,11 @@ def from_origins( ------- EdgeBlock """ - order = ORIGIN_ORDER.get(edge_type, tuple(sorted(per_origin))) - unknown = set(per_origin) - set(order) - if unknown: - raise ValueError( - f"{edge_type}: origins {sorted(unknown)} are not in ORIGIN_ORDER" - f"[{edge_type!r}] = {order}. Add them there so the layout stays " - f"deterministic." - ) - - chunks: List[np.ndarray] = [] - bounds: Dict[str, Tuple[int, int]] = {} - cursor = 0 - for origin in order: - rows = per_origin.get(origin) - if rows is None or len(rows) == 0: - continue - rows = np.asarray(rows, dtype=np.int64).reshape(-1, arity) - rows = rows[_lexsort_rows(rows)] - chunks.append(rows) - bounds[origin] = (cursor, cursor + len(rows)) - cursor += len(rows) - - if not chunks: + indices, bounds, _ = assemble_origins(per_origin, arity, edge_type) + if len(indices) == 0: return cls.empty(arity, device=device) - - stacked = np.concatenate(chunks, axis=0) return cls( - indices=torch.as_tensor(stacked, dtype=torch.int64, device=device), + indices=torch.as_tensor(indices, dtype=torch.int64, device=device), origin_bounds=bounds, ) @@ -202,4 +257,4 @@ def __repr__(self) -> str: ) -__all__ = ["EdgeBlock", "ORIGIN_ORDER"] +__all__ = ["EdgeBlock", "ORIGIN_ORDER", "assemble_origins"] diff --git a/torchref/topology/restraint_sets.py b/torchref/topology/restraint_sets.py new file mode 100644 index 00000000..289048c9 --- /dev/null +++ b/torchref/topology/restraint_sets.py @@ -0,0 +1,180 @@ +"""Ideal values and sigmas, layered over a topology's edges. + +The topology says what is connected. What the ideal geometry *is* lives here, keyed to +the same edges, so one connectivity can carry monomer-library targets, force-field +parameters or ADP-similarity sigmas without a second copy of the edges. + +:func:`assemble_entries` turns a topology plus its values into the nested mapping the +restraint consumers read, ``entries[edge_type][origin][property]``. Indices are +**views** into the contiguous edge blocks, so taking a per-origin subset costs nothing +and an in-place edit to a block is visible through every view of it. Access is three +dict lookups with no allocation, which is the point: it sits on the geometry targets' +hot path. +""" + +from typing import Dict, Optional, Sequence, Tuple + +import numpy as np +import torch + +from torchref.config import get_float_dtype + +#: Origins making up each edge type's ``all`` group -- what the geometry targets read. +#: ``None`` means every origin present. ``phi`` and ``psi`` are conformationally free +#: and carry no target, and ``omega`` has its own von Mises target, so the torsion group +#: holds only the two origins that are ordinary restrained torsions. +ALL_MEMBERS: Dict[str, Optional[Tuple[str, ...]]] = { + "bond": None, + "angle": None, + "torsion": ("intra", "disulfide"), +} + +#: Integer-valued edge properties, kept as ``int64`` rather than the float dtype. +_INTEGER_PROPERTIES = frozenset({"periods", "symop_indices", "cell_offsets"}) + +#: Boolean edge properties. +_BOOL_PROPERTIES = frozenset({"is_proline"}) + + +def to_tensor(values, prop: str, device=None) -> torch.Tensor: + """A per-edge value array as a tensor of the dtype that property calls for.""" + if isinstance(values, torch.Tensor): + return values.to(device=device) if device is not None else values + if prop in _INTEGER_PROPERTIES: + dtype = torch.int64 + elif prop in _BOOL_PROPERTIES: + dtype = torch.bool + else: + dtype = get_float_dtype() + return torch.as_tensor(np.asarray(values), dtype=dtype, device=device) + + +def _contiguous_span( + bounds: Dict[str, Tuple[int, int]], members: Sequence[str] +) -> Optional[Tuple[int, int]]: + """The single row range covering ``members``, or None if they are not adjacent. + + Adjacency is what makes the ``all`` group a view instead of a copy, so the block + layout in :data:`~torchref.topology.edges.ORIGIN_ORDER` deliberately keeps each edge + type's ``all`` members together. + """ + present = [m for m in members if m in bounds] + if not present: + return None + spans = sorted(bounds[m] for m in present) + for (_, end), (start, _) in zip(spans, spans[1:]): + if end != start: + return None + return spans[0][0], spans[-1][1] + + +def _group( + indices: torch.Tensor, values: Dict[str, torch.Tensor] +) -> Dict[str, torch.Tensor]: + """One restraint group: its indices plus whatever properties it carries.""" + group = {"indices": indices} + group.update(values) + return group + + +def _all_group( + block, per_origin_values: Dict[str, Dict[str, torch.Tensor]], members +) -> Optional[Dict[str, torch.Tensor]]: + """The combined group a geometry target reads, or None when it would be empty. + + Carries only properties present in **every** member origin, so a property one member + lacks does not silently become a partial array. + """ + origins = ( + block.origins() + if members is None + else [m for m in members if m in block.origin_bounds] + ) + if not origins: + return None + + span = _contiguous_span(block.origin_bounds, origins) + if span is not None: + start, end = span + indices = block.indices[start:end] + else: + # Not reachable with the shipped layouts; kept so a future origin order that + # separates the members degrades to a copy rather than silently misaligning. + indices = torch.cat([block.origin(o) for o in origins], dim=0) + + shared = set(per_origin_values.get(origins[0], {})) + for origin in origins[1:]: + shared &= set(per_origin_values.get(origin, {})) + + values = { + prop: torch.cat([per_origin_values[o][prop] for o in origins], dim=0) + for prop in sorted(shared) + } + return _group(indices, values) + + +def assemble_entries( + topology, + values: Dict[str, Dict[str, Dict[str, torch.Tensor]]], +) -> Dict[str, Dict]: + """Build the nested mapping the restraint consumers read. + + Parameters + ---------- + topology : Topology + Supplies the edge blocks; per-origin indices come out as views into them. + values : dict + ``{edge_type: {origin: {property: tensor}}}`` for the keyed types, plus + ``{'chiral': {property: tensor}}`` and ``{'plane': {size: {property: tensor}}}`` + for the two that carry no origin. + + Returns + ------- + dict + ``entries[edge_type][origin][property]`` for bonds, angles and torsions; + ``entries['plane']['4_atoms'][property]``; ``entries['chiral'][property]``. + ``entries['vdw']`` starts empty and is filled when the pair list is built. + """ + entries: Dict[str, Dict] = {} + + for edge_type in ("bond", "angle", "torsion"): + block = topology.edge_block(edge_type) + per_origin = values.get(edge_type, {}) + group: Dict[str, Dict[str, torch.Tensor]] = {} + for origin in block.origins(): + group[origin] = _group(block.origin(origin), per_origin.get(origin, {})) + combined = _all_group(block, per_origin, ALL_MEMBERS[edge_type]) + if combined is not None: + group["all"] = combined + entries[edge_type] = group + + chirals = topology.atoms.chirals + entries["chiral"] = ( + _group(chirals.indices, values.get("chiral", {})) if chirals.n_edges else {} + ) + + entries["plane"] = { + f"{size}_atoms": _group(block.indices, values.get("plane", {}).get(size, {})) + for size, block in sorted(topology.atoms.planes.items()) + } + + entries["vdw"] = {} + return entries + + +def max_period(entries: Dict[str, Dict]) -> int: + """Largest torsion period in the ``all`` group. + + Read once at build time so the torsion target does not pay a device sync for it on + every iteration. + """ + group = entries.get("torsion", {}).get("all") + if not group: + return 1 + periods = group.get("periods") + if periods is None or periods.numel() == 0: + return 1 + return int(periods.max().item()) + + +__all__ = ["assemble_entries", "max_period", "to_tensor", "ALL_MEMBERS"] From 3908ade68f74375e7e4c0a5195f653b278aba3d3 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Wed, 26 Aug 2026 23:38:22 +0200 Subject: [PATCH 059/250] Generate hydrogens by instantiating the monomer template A template already carries its hydrogens, with coordinates and with bonds naming each one's parent. Generating hydrogens is therefore template instantiation, not geometry reconstruction: align the template onto the heavy atoms present, read the hydrogen positions off it, correct each to its library bond length. The bond graph supplies the two things a template cannot know on its own -- how many hydrogens a parent can still carry, and which of them have a dihedral nobody has determined. The previous placement fitted the template over two bond shells. For CB that is {C, CA, CB, N, SG}, which spans chi1, and chi1 is the model's, not the library's, so the rigid fit compromises between them. Measured on 7L84 that aligned to 0.75 A RMSD and left 12% of side-chain hydrogens more than 1.5 A from the atom they belong to, where they were silently dropped. Fitting the parent and its immediate neighbours only -- the unit whose bond lengths and angles really are library constants -- places all of them: 7L84 goes from 930 to 1064 hydrogens with none left undetermined, each at its ideal bond length to machine precision. Three strategies, chosen by what the graph says, in place of a seven-way placement enum. The template frame where the template knows every heavy atom actually bonded to the parent. Construction from the bonded neighbours where it does not, which is the peptide-linked backbone nitrogen: it is bonded to the preceding residue's carbon, and the free-amino-acid template has never heard of that. An axis-preserving frame for a single-neighbour centre, which fixes every bond angle and leaves only the rotation the scan then chooses. Free torsions come out of the connectivity rather than a list of special cases: a parent with exactly one heavy neighbour can rotate. That is the hydroxyl, the thiol, the amine and the methyl, and it correctly makes an N-terminal ammonium rotatable while an in-chain amide is not -- the test asserts a backbone nitrogen rotates if and only if no peptide bond reaches its residue. Waters are left alone, and for a reason rather than by omission: one heavy atom gives no frame to align against and no bond to rotate about, so their hydrogens could only point somewhere arbitrary. strip_H still defaults to True, so no hydrogen enters a refinement and nothing moves numerically. Model.generate_hydrogens, the gemmi path that round-tripped through temporary PDBs on disk and had no callers, is gone; hydrogenate is the single route and no longer takes lbfgs_steps or max_iter. hydrogen_topology.py keeps its riding-placement machinery for now. It is still live while hydrogens are absent from the model, and deleting it here would remove the H-VDW term rather than leave the stage inert. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- docs/changelog.rst | 5 + tests/unit/topology/test_hydrogens.py | 284 ++++++++ torchref/model/model.py | 742 +-------------------- torchref/restraints/hydrogen_topology.py | 11 +- torchref/topology/__init__.py | 13 +- torchref/topology/hydrogens.py | 803 +++++++++++++++++++++++ 6 files changed, 1138 insertions(+), 720 deletions(-) create mode 100644 tests/unit/topology/test_hydrogens.py create mode 100644 torchref/topology/hydrogens.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 865e7e95..67186f4f 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -26,6 +26,11 @@ Unreleased - Restraint groups are laid out in a fixed order, so a rebuild produces the same row order in any process; previously the origins were concatenated in Python ``set`` iteration order - Fixed ``cat_dict`` doubling every bond, angle and torsion restraint when called more than once - Geometry restraints are built from the topology, retiring the intra-residue builder calls, the peptide/disulfide/LINK build methods and the ``TensorDict`` restraint storage +- Hydrogen generation is now template instantiation over the topology: ``Model.hydrogenate`` aligns each residue's monomer template onto the heavy atoms present and reads its hydrogens off, and the bond graph sets how many hydrogens a parent may carry +- Fixed hydrogen placement fitting the template over two bond shells, which spans rotatable torsions the model does not share and left 12% of side-chain hydrogens further than 1.5 A from their parent, where they were discarded +- Hydrogens on a centre whose template omits a real substituent -- a peptide-linked backbone nitrogen -- are now built from the bonded neighbours instead of the template frame +- Hydrogens with a free torsion (hydroxyl, thiol, amine, methyl) are identified from bond connectivity and their dihedral is scanned, rather than taken from whatever the library deposited +- Removed ``Model.generate_hydrogens``; ``Model.hydrogenate`` is the single path and no longer takes ``lbfgs_steps`` or ``max_iter`` Version 0.6.4 diff --git a/tests/unit/topology/test_hydrogens.py b/tests/unit/topology/test_hydrogens.py new file mode 100644 index 00000000..755e6e2c --- /dev/null +++ b/tests/unit/topology/test_hydrogens.py @@ -0,0 +1,284 @@ +"""Hydrogen generation from monomer templates, driven by the bond graph. + +The properties asserted here are the ones that make template instantiation +trustworthy: every hydrogen lands at its library bond length, none is placed in a +direction the geometry does not determine, the count per parent respects the valence +left over after the graph's real bonds, and the free-torsion set is exactly the centres +whose dihedral the template cannot know. +""" + +import numpy as np +import pytest + +from torchref.model.model import Model +from torchref.topology.hydrogens import ( + STANDARD_VALENCE, + augment_atom_table, + optimise_free_torsions, + plan_hydrogens, +) + +STRUCTURES = ["7L84", "1DAW"] + + +@pytest.fixture(scope="module") +def built(pdb_dir): + """``(model, restraints, plan)`` per structure, built once.""" + cache = {} + + def _build(code): + if code not in cache: + model = Model(verbose=0) + model.load_pdb(str(pdb_dir / f"{code}.pdb")) + model.set_restraints_cif(None) + restraints = model.restraints + plan = plan_hydrogens( + restraints.topology, restraints.cif_dict, model.xyz().detach() + ) + cache[code] = (model, restraints, plan) + return cache[code] + + return _build + + +@pytest.mark.unit +@pytest.mark.parametrize("code", STRUCTURES) +def test_hydrogens_sit_at_their_library_bond_length(built, code): + """Placement is exact, not approximate: the parent distance is the library value.""" + model, _, plan = built(code) + assert plan.n_hydrogens > 0 + + coords = model.xyz().detach().cpu().numpy() + distance = np.linalg.norm(plan.position - coords[plan.parent], axis=1) + assert np.abs(distance - plan.bond_length).max() < 1e-9 + + +@pytest.mark.unit +@pytest.mark.parametrize("code", STRUCTURES) +def test_every_candidate_hydrogen_is_placed(built, code): + """No hydrogen is dropped for want of a determined direction. + + A hydrogen is only planned once a strategy has fixed its direction, so a shortfall + here means some centre fell through all three. The earlier two-shell alignment left + 12% of side-chain hydrogens beyond 1.5 A of their parent and discarded them. + """ + model, restraints, plan = built(code) + topology = restraints.topology + + # One hydrogen per free valence on every parent that has a template hydrogen. + expected_parents = set(plan.parent.tolist()) + assert expected_parents, "no parents received hydrogens" + + is_h = topology.atoms.is_hydrogen + for parent in sorted(expected_parents): + neighbours = topology.atoms.neighbors(parent) + heavy = int((~is_h[neighbours]).sum()) + element = str(topology.atoms.element[parent]).strip().upper() + allowed = max(0, STANDARD_VALENCE.get(element, 4) - heavy) + placed = int((plan.parent == parent).sum()) + assert placed <= allowed, ( + f"{code}: atom {parent} ({element}) has {heavy} heavy bonds, so at most " + f"{allowed} hydrogens, but {placed} were planned" + ) + + +@pytest.mark.unit +@pytest.mark.parametrize("code", STRUCTURES) +def test_free_torsions_are_exactly_the_single_neighbour_centres(built, code): + """A dihedral is free when the parent has one heavy neighbour, and only then.""" + _, restraints, plan = built(code) + topology = restraints.topology + is_h = topology.atoms.is_hydrogen + + for i in range(plan.n_hydrogens): + parent = int(plan.parent[i]) + neighbours = topology.atoms.neighbors(parent) + heavy = int((~is_h[neighbours]).sum()) + assert (plan.group[i] >= 0) == (heavy == 1), ( + f"{code}: hydrogen {plan.name[i]} on atom {parent} with {heavy} heavy " + f"neighbours has group {plan.group[i]}" + ) + + +@pytest.mark.unit +def test_hydroxyl_rotates_and_backbone_amide_does_not(built): + """The chemistry the graph criterion is meant to capture, spot-checked. + + A serine hydroxyl hangs off an oxygen bonded only to CB, so its dihedral is free. A + backbone amide nitrogen is bonded to CA and to the preceding residue's carbon, which + fixes its hydrogen entirely. + """ + _, restraints, plan = built("7L84") + topology = restraints.topology + names = topology.atoms.name.astype(str) + resnames = topology.residues.resname + + free_parents = {int(p) for p, g in zip(plan.parent, plan.group) if g >= 0} + fixed_parents = {int(p) for p, g in zip(plan.parent, plan.group) if g < 0} + + hydroxyl = [ + int(p) + for p in free_parents + if names[p] == "OG" + and str(resnames[topology.residue_of_atom(p)]).strip() == "SER" + ] + assert hydroxyl, "no serine hydroxyl was treated as a free torsion" + + amide = [p for p in fixed_parents if names[p] == "N"] + assert amide, "no backbone amide nitrogen was treated as determined" + + # A backbone nitrogen rotates exactly when nothing is bonded to it on the other + # side: an N-terminal ammonium does, an in-chain amide does not. Tied to the + # residue graph's link edges, so a missing peptide bond would show up here. + peptide_links = topology.residues.links_of_kind("TRANS") + accepts_link = set(peptide_links[:, 1].tolist()) if len(peptide_links) else set() + + for parent in free_parents: + if names[parent] != "N": + continue + residue = topology.residue_of_atom(parent) + assert residue not in accepts_link, ( + f"nitrogen {parent} in residue {topology.residues.key(residue)} was " + f"treated as rotatable even though a peptide bond reaches it" + ) + for parent in fixed_parents: + if names[parent] != "N": + continue + residue = topology.residue_of_atom(parent) + assert residue in accepts_link, ( + f"nitrogen {parent} in residue {topology.residues.key(residue)} was " + f"treated as determined but no peptide bond reaches it" + ) + + +@pytest.mark.unit +@pytest.mark.parametrize("code", STRUCTURES) +def test_torsion_scan_preserves_bond_lengths(built, code): + """The scan rotates about a bond, so it cannot change any bond length.""" + model, restraints, plan = built(code) + coords = model.xyz().detach() + + scanned = plan_hydrogens(restraints.topology, restraints.cif_dict, coords) + before = scanned.position.copy() + optimise_free_torsions(scanned, restraints.topology, coords) + + numpy_coords = coords.cpu().numpy() + distance = np.linalg.norm(scanned.position - numpy_coords[scanned.parent], axis=1) + assert np.abs(distance - scanned.bond_length).max() < 1e-9 + + moved = np.linalg.norm(scanned.position - before, axis=1) > 1e-6 + assert moved.any(), "the scan changed nothing at all" + assert not moved[ + ~scanned.rotatable + ].any(), "the scan moved a hydrogen whose torsion is not free" + + +@pytest.mark.unit +def test_scan_reduces_clash(built): + """Scanned hydrogens end up no closer to heavy atoms than they started.""" + model, restraints, plan = built("7L84") + coords = model.xyz().detach() + topology = restraints.topology + + scanned = plan_hydrogens(topology, restraints.cif_dict, coords) + numpy_coords = coords.cpu().numpy() + heavy = numpy_coords[~topology.atoms.is_hydrogen.cpu().numpy()] + + def closest(positions): + gaps = np.linalg.norm(positions[:, None, :] - heavy[None, :, :], axis=-1) + # The parent itself is always the nearest heavy atom; take the next one. + return np.sort(gaps, axis=1)[:, 1] + + rotatable = scanned.rotatable + before = closest(scanned.position[rotatable]) + optimise_free_torsions(scanned, topology, coords) + after = closest(scanned.position[rotatable]) + + assert after.min() >= before.min() - 1e-9, "the scan made the worst clash worse" + + +@pytest.mark.unit +@pytest.mark.parametrize("code", STRUCTURES) +def test_augmented_table_keeps_residues_contiguous(built, code): + """Hydrogens are inserted into their residue, not appended after everything. + + The residue partition is built from contiguous runs of ``(chain, resseq, icode)``, + so appending hydrogens at the end would split every hydrogenated residue in two. + """ + model, restraints, plan = built(code) + augmented = augment_atom_table(model.pdb, plan, restraints.topology) + + assert len(augmented) == len(model.pdb) + plan.n_hydrogens + assert (augmented["index"].values == np.arange(len(augmented))).all() + + key = ( + augmented[["chainid", "resseq", "icode"]] + .astype(str) + .agg("|".join, axis=1) + .values + ) + runs = 1 + int((key[1:] != key[:-1]).sum()) + assert runs == len(set(key)), "a residue was split into non-adjacent runs" + + +@pytest.mark.unit +def test_waters_are_not_hydrogenated(built): + """A single-atom residue is skipped, and for a reason rather than by accident. + + One heavy atom gives no frame to align a template against and no bond to rotate + about, so a water's hydrogens could only be placed in an arbitrary direction. + """ + _, restraints, plan = built("7L84") + topology = restraints.topology + + waters = [ + i + for i in range(topology.n_residues) + if str(topology.residues.resname[i]).strip() == "HOH" + ] + assert waters, "7L84 has no waters, so this asserts nothing" + assert not set(plan.residue.tolist()) & set(waters) + + +@pytest.mark.unit +def test_hydrogenate_returns_a_consistent_model(pdb_dir): + """The end-to-end path yields a model whose tensors, table and restraints agree.""" + model = Model(verbose=0) + model.load_pdb(str(pdb_dir / "7L84.pdb")) + model.set_restraints_cif(None) + n_heavy = len(model.pdb) + + hydrogenated = model.hydrogenate(verbose=0) + + assert hydrogenated.ctx.strip_H is False + assert len(hydrogenated.pdb) > n_heavy + assert hydrogenated.xyz().shape[0] == len(hydrogenated.pdb) + assert hydrogenated.adp().shape[0] == len(hydrogenated.pdb) + assert len(model.pdb) == n_heavy, "the original model was modified" + + elements = hydrogenated.pdb["element"].astype(str).str.strip().values + n_h = int((elements == "H").sum()) + assert n_h == len(hydrogenated.pdb) - n_heavy + + # Every hydrogen carries exactly one bond restraint, at library geometry. + restraints = hydrogenated.restraints + bonds = restraints.restraints["bond"]["all"]["indices"].cpu().numpy() + references = restraints.restraints["bond"]["all"]["references"].cpu().numpy() + coords = hydrogenated.xyz().detach().cpu().numpy() + is_h = elements == "H" + involves_h = is_h[bonds[:, 0]] | is_h[bonds[:, 1]] + + assert int(involves_h.sum()) == n_h + lengths = np.linalg.norm(coords[bonds[:, 0]] - coords[bonds[:, 1]], axis=1) + deviation = np.sqrt(((lengths[involves_h] - references[involves_h]) ** 2).mean()) + assert deviation < 0.02, f"placed hydrogens deviate by {deviation:.4f} A RMS" + + +@pytest.mark.unit +def test_loading_still_strips_hydrogens_by_default(pdb_dir): + """``strip_H`` is untouched, so refinement sees the same atoms as before.""" + model = Model(verbose=0) + model.load_pdb(str(pdb_dir / "1AK5_with_H.pdb")) + assert model.ctx.strip_H is True + elements = model.pdb["element"].astype(str).str.strip().values + assert not (elements == "H").any() diff --git a/torchref/model/model.py b/torchref/model/model.py index bd706ec7..592092f1 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -1587,128 +1587,6 @@ def shake_adp(self, stddev: float): new_adp, refinable_mask=self.adp.refinable_mask, name="adp" ) - def generate_hydrogens(self, mon_lib_path: str = None) -> "Model": - """ - Generate hydrogen atoms for the current model using gemmi. - - Places hydrogens at ideal geometry using the CCP4 monomer library and - gemmi's topology engine. Returns a new Model instance with hydrogens - added; the original model is not modified. - - Parameters - ---------- - mon_lib_path : str, optional - Path to CCP4 monomer library directory. If None, uses the monomer - library bundled with torchref (covers standard amino acids and - common small molecules). - - Returns - ------- - Model - A new Model instance with hydrogen atoms added (strip_H=False). - Unknown residues are skipped silently. - - Notes - ----- - Reads the *current* coordinates (via :meth:`update_pdb`), so run it after - any coordinate change that should be reflected in the H positions. - """ - import os - import tempfile - - import gemmi - - from torchref import PATH_TORCHREF_DATA - - # ``mgr`` is set when we fall back to TorchRef's auto-fetching monomer - # library manager; per-residue CIFs are then resolved through it (which - # downloads/caches on demand) rather than from ``mon_lib_path`` directly. - mgr = None - if mon_lib_path is None: - import os as _os - - # In priority order: CCP4's own env var, a library bundled next to the - # repo, then the partial one shipped inside torchref. - candidates = [ - _os.environ.get("CLIBD_MON", ""), - str(PATH_TORCHREF_DATA.parent.parent / "external_monomer_library"), - str(PATH_TORCHREF_DATA / "monomer_library"), - ] - mon_lib_path = None - for c in candidates: - if c and _os.path.isfile(_os.path.join(c, "ener_lib.cif")): - mon_lib_path = c - break - if mon_lib_path is None: - # No complete CCP4 library: fall back to TorchRef's manager, which - # ships standard residues and auto-downloads the rest. This is the - # normal path — a CCP4 install is not required. - from torchref.restraints.library import get_library_manager - - mgr = get_library_manager(verbose=self.ctx.verbose) - mon_lib_path = str(mgr.ensure_gemmi_base()) - - # gemmi reads from a file, so the live tensors must reach the DataFrame. - self.update_pdb() - - with tempfile.NamedTemporaryFile(suffix=".pdb", delete=False) as f: - tmp_heavy = f.name - with tempfile.NamedTemporaryFile(suffix=".pdb", delete=False) as f: - tmp_with_h = f.name - - try: - from torchref.io import pdb as io_pdb - from torchref.utils.utils import sanitize_pdb_dataframe - - pdb_out = sanitize_pdb_dataframe(self.pdb.copy()) - pdb_out.attrs["spacegroup"] = ( - self.spacegroup.hm if self.spacegroup else "P 1" - ) - io_pdb.write(pdb_out, tmp_heavy) - - st = gemmi.read_structure(tmp_heavy) - st.setup_entities() - - # Per-residue CIFs come from the manager (bundled → cache → download) - # when we fell back to it, else from the explicit library directory. - monlib = gemmi.read_monomer_lib(mon_lib_path, []) - resnames = set(r.name for m in st for c in m for r in c) - for rn in resnames: - if mgr is not None: - cif = mgr.get_cif_file(rn) - cif_path = str(cif) if cif is not None else None - else: - cif_path = os.path.join(mon_lib_path, rn[0].lower(), rn + ".cif") - if not os.path.exists(cif_path): - cif_path = None - if cif_path is None: - continue - doc = gemmi.cif.read(cif_path) - for block in doc: - if block.name == rn or block.name.startswith("comp_" + rn): - monlib.add_monomer_if_present(block) - break - - gemmi.prepare_topology(st, monlib, h_change=gemmi.HydrogenChange.ReAdd) - st.write_pdb(tmp_with_h) - - # strip_H=False, or the hydrogens we just placed would be dropped again. - new_model = self.__class__( - dtype_float=self.dtype_float, - verbose=self.ctx.verbose, - device=self.device, - strip_H=False, - ) - new_model.load_pdb(tmp_with_h) - - finally: - for p in (tmp_heavy, tmp_with_h): - try: - os.unlink(p) - except OSError: - pass - - return new_model def _new_model_from_df(self, df, *, strip_H=None): """Build a fresh model of the same class from a DataFrame.""" @@ -1808,616 +1686,50 @@ def strip_hydrogens(self) -> "Model": filtered.attrs = pdb.attrs.copy() return self._new_model_from_df(filtered, strip_H=True) - # Module-level cache for CIF monomer data (shared across calls) - _hydrogenate_cif_cache = {} - - def hydrogenate( - self, - verbose: int = 0, - optimize: bool = False, - lbfgs_steps: int = 3, - max_iter: int = 20, - ) -> "Model": - """ - Return a new model with hydrogen atoms placed via Kabsch alignment. + def hydrogenate(self, verbose: int = 0, optimize: bool = True) -> "Model": + """Return a new model with hydrogens added from the monomer templates. - Uses torchref's monomer library to identify missing H atoms, places - them by SVD-aligning ideal monomer coordinates onto the current model - coordinates, then corrects each H to sit at ideal bond length from its - parent atom. The original model is not modified. + Hydrogen generation is template instantiation over the topology: each residue's + library template is aligned onto the heavy atoms present and its hydrogens read + off, and the bond graph decides how many hydrogens a parent can carry and which + of them have a free torsion. The original model is not modified. Parameters ---------- - verbose : int, optional - Verbosity level (0=silent, 1=summary, 2=detailed). Default 0. - optimize : bool, optional - If True, run a short LBFGS geometry optimization on H positions - after placement. Default False (Kabsch placement only). - lbfgs_steps : int, optional - Number of LBFGS outer steps (only when optimize=True). Default 3. - max_iter : int, optional - Max line-search iterations per LBFGS step. Default 20. + verbose : int, default 0 + Verbosity level. + optimize : bool, default True + Scan each free torsion -- hydroxyl, thiol, amine, methyl -- for the + least-clashing angle. The template's dihedral for those is arbitrary, so this + is on by default; it is a rotation about one bond and costs little. Returns ------- Model - New model with hydrogen atoms added. - All parameters are unfrozen in the returned model. + New model with hydrogens, built with ``strip_H=False`` so they survive the + load. """ - import numpy as np - import pandas as pd - - from torchref.restraints.library import MonomerLibraryManager + from torchref.topology.hydrogens import ( + augment_atom_table, + optimise_free_torsions, + plan_hydrogens, + ) - # Sync current coordinates into DataFrame self.update_pdb() + restraints = self.restraints # builds the topology this reads + xyz = self.xyz().detach() - lib = MonomerLibraryManager(verbose=0) - cache = Model._hydrogenate_cif_cache - - # --- Phase A: build per-residue-type lookup tables (cached) --- - for rn in self.pdb["resname"].unique(): - rn_str = str(rn).strip() - if not rn_str: - continue - if rn_str in cache: - if cache[rn_str] is None or "heavy_neighbor_map" in cache[rn_str]: - continue - del cache[rn_str] # Stale entry, re-read - cif_path = lib.get_cif_file(rn_str) - if cif_path is None: - cache[rn_str] = None - continue - try: - from torchref.io.cif_readers import RestraintCIFReader - - reader = RestraintCIFReader(str(cif_path)) - all_data = reader.get_all_restraints() - comp_data = all_data.get(rn_str) or all_data.get(rn_str.upper()) - if comp_data is None: - cache[rn_str] = None - continue - atom_df = comp_data.get("atoms", comp_data.get("atom")) - bond_df = comp_data.get("bonds", comp_data.get("bond")) - if atom_df is None or atom_df.empty or "x" not in atom_df.columns: - cache[rn_str] = None - continue - except Exception: - cache[rn_str] = None - continue - - ids = atom_df["atom_id"].astype(str).str.strip().values - elems = atom_df["type_symbol"].astype(str).str.strip().values - coords = atom_df[["x", "y", "z"]].values.astype(np.float64) - is_h = np.array([e.upper() == "H" for e in elems]) - id_to_idx = {n: i for i, n in enumerate(ids)} - - # H→parent map + ideal bond lengths + heavy adjacency - parent_map = {} # h_name -> parent_name - ideal_bl = {} # h_name -> ideal bond length (Angstrom) - heavy_neighbor_map = {} # heavy_name -> [bonded heavy names] - if bond_df is not None and not bond_df.empty: - a1s = bond_df["atom1"].astype(str).str.strip().values - a2s = bond_df["atom2"].astype(str).str.strip().values - vals = pd.to_numeric(bond_df["value"], errors="coerce").values - h_set = set(ids[is_h]) - for i in range(len(a1s)): - b1, b2 = a1s[i], a2s[i] - if b1 in h_set and b2 in id_to_idx and not is_h[id_to_idx[b2]]: - parent_map[b1] = b2 - if np.isfinite(vals[i]): - ideal_bl[b1] = float(vals[i]) - elif b2 in h_set and b1 in id_to_idx and not is_h[id_to_idx[b1]]: - parent_map[b2] = b1 - if np.isfinite(vals[i]): - ideal_bl[b2] = float(vals[i]) - # Heavy-atom adjacency for local Kabsch - i1, i2 = id_to_idx.get(b1), id_to_idx.get(b2) - if ( - i1 is not None - and i2 is not None - and not is_h[i1] - and not is_h[i2] - ): - heavy_neighbor_map.setdefault(b1, []).append(b2) - heavy_neighbor_map.setdefault(b2, []).append(b1) - - cache[rn_str] = { - "ids": ids, - "elems": elems, - "coords": coords, - "is_h": is_h, - "id_to_idx": id_to_idx, - "heavy_names": ids[~is_h], - "heavy_coords": coords[~is_h], - "h_names": ids[is_h], - "h_coords": coords[is_h], - "parent_map": parent_map, - "ideal_bl": ideal_bl, - "heavy_neighbor_map": heavy_neighbor_map, - } - - # Filter to available residue types - available = { - rn: cache[rn] - for rn in self.pdb["resname"].unique() - if str(rn).strip() in cache and cache.get(str(rn).strip()) is not None - } - if not available: - if verbose > 0: - print("No monomer library data found; returning copy.") - return self.copy() - - # --- Phase B: place H atoms via Kabsch alignment --- - model_names_arr = self.pdb["name"].astype(str).str.strip().values - model_xyz_arr = self.pdb[["x", "y", "z"]].values.astype(np.float64) - model_occ_arr = self.pdb["occupancy"].values.astype(np.float64) - model_bfac_arr = self.pdb["tempfactor"].values.astype(np.float64) - model_atom_type_arr = self.pdb["ATOM"].values - model_altloc_arr = self.pdb["altloc"].values.astype(str) - - group_cols = ["chainid", "resseq", "icode", "resname"] - group_keys = self.pdb[group_cols].values - changes = np.zeros(len(group_keys), dtype=bool) - changes[0] = True - for c in range(4): - changes[1:] |= group_keys[1:, c] != group_keys[:-1, c] - group_starts = np.nonzero(changes)[0] - group_ends = np.append(group_starts[1:], len(group_keys)) - - # Pre-allocate lists for H atom data columns - h_x, h_y, h_z = [], [], [] - h_names_out, h_altlocs, h_resnames = [], [], [] - h_chainids, h_resseqs, h_icodes = [], [], [] - h_occ, h_bfac, h_atom_types = [], [], [] - h_insert_after = [] - - max_bond_dist = 1.5 # Reject H atoms placed > this from parent - _std_val = {"C": 4, "N": 3, "O": 2, "S": 2} - - # Heavy-atom mask for distance-based neighbor detection - model_elem_arr = self.pdb["element"].astype(str).str.strip().values - model_heavy_mask_full = np.array([e.upper() != "H" for e in model_elem_arr]) - - for gi in range(len(group_starts)): - s, e = group_starts[gi], group_ends[gi] - rn = str(group_keys[s, 3]).strip() - info = cache.get(rn) - if info is None: - continue - chainid = group_keys[s, 0] - resseq = group_keys[s, 1] - icode = group_keys[s, 2] - - names_in_model = set(model_names_arr[s:e]) - h_to_add_mask = np.array( - [n not in names_in_model for n in info["h_names"]], dtype=bool - ) - if not h_to_add_mask.any(): - continue - h_names_add = info["h_names"][h_to_add_mask] - h_coords_ideal = info["h_coords"][h_to_add_mask] - - # Altloc handling - altlocs_in_res = set(model_altloc_arr[s:e]) - altloc_list = ( - [""] - if altlocs_in_res <= {""} - else sorted(a for a in altlocs_in_res if a != "") - ) - - for altloc in altloc_list: - if altloc == "": - mask = np.ones(e - s, dtype=bool) - else: - al = model_altloc_arr[s:e] - mask = (al == altloc) | (al == "") - - conf_names = model_names_arr[s:e][mask] - conf_xyz = model_xyz_arr[s:e][mask] - conf_occ = model_occ_arr[s:e][mask] - conf_bfac = model_bfac_arr[s:e][mask] - conf_atom_type = model_atom_type_arr[s:e][mask] - - # Name→index lookup for this conformer - name_to_idx = {} - for j, cn in enumerate(conf_names): - if cn not in name_to_idx: - name_to_idx[cn] = j - - conf_name_set = set(conf_names) - common_mask = np.array( - [n in conf_name_set for n in info["heavy_names"]], - dtype=bool, - ) - n_common = common_mask.sum() - - # Global Kabsch when ≥ 3 matching heavy atoms - R_global = t_global = None - if n_common >= 3: - P = info["heavy_coords"][common_mask] - Q = np.array( - [ - conf_xyz[name_to_idx[n]] - for n in info["heavy_names"][common_mask] - ], - dtype=np.float64, - ) - cp, cq = P.mean(0), Q.mean(0) - Hm = (P - cp).T @ (Q - cq) - U, S, Vt = np.linalg.svd(Hm) - d = np.linalg.det(Vt.T @ U.T) - sign_d = np.diag([1.0, 1.0, 1.0 if d > 0 else -1.0]) - R_global = Vt.T @ sign_d @ U.T - t_global = cq - R_global @ cp - - # Group H atoms by parent for placement - parent_to_hi = {} - for hi, h_name in enumerate(h_names_add): - pn = info["parent_map"].get(h_name) - if pn is not None and pn in name_to_idx: - parent_to_hi.setdefault(pn, []).append(hi) - - hnm = info.get("heavy_neighbor_map", {}) - id2i = info["id_to_idx"] - all_coords = info["coords"] - mask_idx = np.where(mask)[0] # conformer indices in [s:e] - - for par_name, hi_list in parent_to_hi.items(): - pidx = name_to_idx[par_name] - parent_pos = conf_xyz[pidx] - parent_full = s + mask_idx[pidx] - - # Heavy neighbors in the model (distance-based, - # includes cross-residue bonds like C-N peptide) - dvec = model_xyz_arr - model_xyz_arr[parent_full] - dists_sq = (dvec**2).sum(1) - bonded = np.where( - (dists_sq > 0.09) & (dists_sq < 3.61) & model_heavy_mask_full - )[0] - bonded = bonded[bonded != parent_full] - n_model_heavy = len(bonded) - - # Expected H count from standard valence - par_elem = info["elems"][id2i[par_name]].upper() - expected_h = max( - 0, - _std_val.get(par_elem, 4) - n_model_heavy, - ) - - # --- Step 1: local Kabsch for initial placement --- - local_set = {par_name} - for nb in hnm.get(par_name, []): - local_set.add(nb) - for nb2 in hnm.get(nb, []): - local_set.add(nb2) - local_names = [ - n for n in local_set if n in name_to_idx and n in id2i - ] - - if len(local_names) >= 3: - Pl = np.array([all_coords[id2i[n]] for n in local_names]) - Ql = np.array([conf_xyz[name_to_idx[n]] for n in local_names]) - cpl, cql = Pl.mean(0), Ql.mean(0) - Hl = (Pl - cpl).T @ (Ql - cql) - Ul, _, Vtl = np.linalg.svd(Hl) - dl = np.linalg.det(Vtl.T @ Ul.T) - sl = np.diag([1.0, 1.0, 1.0 if dl > 0 else -1.0]) - R_use = Vtl.T @ sl @ Ul.T - t_use = cql - R_use @ cpl - elif R_global is not None: - R_use, t_use = R_global, t_global - else: - R_use = None # Will use random placement - - # Kabsch-place and filter by distance - valid_h = [] - if R_use is not None: - for hi in hi_list: - h_name = h_names_add[hi] - h_cif = all_coords[id2i[h_name]] - h_pos = R_use @ h_cif + t_use - direction = h_pos - parent_pos - dist = np.linalg.norm(direction) - if dist < 1e-6 or dist > max_bond_dist: - continue - bl = info["ideal_bl"].get(h_name, dist) - h_pos = parent_pos + direction * (bl / dist) - valid_h.append((h_name, h_pos, bl)) - else: - # Random-rotation placement (< 3 matching atoms) - # Apply a random SO(3) rotation to ideal CIF - # geometry so internal angles are preserved. - # Random rotation via QR decomposition. - M = np.random.randn(3, 3) - Q_r, _ = np.linalg.qr(M) - if np.linalg.det(Q_r) < 0: - Q_r[:, 0] = -Q_r[:, 0] - par_cif = all_coords[id2i[par_name]] - for hi in hi_list: - h_name = h_names_add[hi] - h_cif = all_coords[id2i[h_name]] - bl = info["ideal_bl"].get(h_name, 0.97) - d_ideal = h_cif - par_cif - d_rot = Q_r @ d_ideal - dn = np.linalg.norm(d_rot) - if dn > 1e-6: - d_rot = d_rot * (bl / dn) - else: - d_rot = np.array([bl, 0.0, 0.0]) - valid_h.append((h_name, parent_pos + d_rot, bl)) - - # Limit to expected count (removes terminal H) - if len(valid_h) > expected_h: - valid_h.sort(key=lambda x: x[0]) # alphabetical - valid_h = valid_h[:expected_h] - - # --- Step 2: geometric re-placement --- - if n_model_heavy >= 2: - nvecs = model_xyz_arr[bonded] - model_xyz_arr[parent_full] - svec = nvecs.sum(0) - snorm = np.linalg.norm(svec) - - if len(valid_h) == 1 and snorm > 1e-6: - # Single H: place opposite to neighbors - h_nm, _, bl = valid_h[0] - h_pos = parent_pos - bl * svec / snorm - valid_h[0] = (h_nm, h_pos, bl) - - elif len(valid_h) == 2 and n_model_heavy == 2 and snorm > 1e-6: - # CH2-like: sp3 tetrahedral placement - v1, v2 = nvecs[0], nvecs[1] - base = -svec / snorm - perp = np.cross(v1, v2) - pn = np.linalg.norm(perp) - if pn > 1e-6: - perp = perp / pn - n1 = np.linalg.norm(v1) - n2 = np.linalg.norm(v2) - c12 = np.dot(v1, v2) / (n1 * n2) - denom = 3.0 * np.sqrt(max(1e-12, (1 + c12) / 2)) - a = min(1.0, 1.0 / denom) - b = np.sqrt(max(0, 1 - a * a)) - d_up = a * base + b * perp - d_dn = a * base - b * perp - # Assign Kabsch-nearest to each - _, pos0, bl0 = valid_h[0] - _, pos1, bl1 = valid_h[1] - g_up = parent_pos + bl0 * d_up - g_dn = parent_pos + bl1 * d_dn - if pos0 is not None and pos1 is not None: - d_same = np.linalg.norm( - pos0 - g_up - ) + np.linalg.norm(pos1 - g_dn) - d_swap = np.linalg.norm( - pos0 - g_dn - ) + np.linalg.norm(pos1 - g_up) - if d_swap < d_same: - g_up, g_dn = g_dn, g_up - valid_h[0] = (valid_h[0][0], g_up, bl0) - valid_h[1] = (valid_h[1][0], g_dn, bl1) - - elif n_model_heavy == 1: - # One heavy neighbor: place H opposite to it - nvec = model_xyz_arr[bonded[0]] - model_xyz_arr[parent_full] - nn = np.linalg.norm(nvec) - if nn > 1e-6: - d_opp = -nvec / nn - for vi in range(len(valid_h)): - if valid_h[vi][1] is None: - nm, _, bl = valid_h[vi] - valid_h[vi] = (nm, parent_pos + bl * d_opp, bl) - - # Fill remaining None positions with random dirs - for vi in range(len(valid_h)): - if valid_h[vi][1] is not None: - continue - nm, _, bl = valid_h[vi] - # Random unit vector via Marsaglia method - while True: - u = np.random.uniform(-1, 1, 3) - n2 = (u * u).sum() - if 0.01 < n2 < 1.0: - break - d = u / np.sqrt(n2) - # Push away from already-placed H siblings - for vj in range(len(valid_h)): - if vj == vi or valid_h[vj][1] is None: - continue - sep = parent_pos + bl * d - valid_h[vj][1] - if np.linalg.norm(sep) < 0.5 * bl: - d = -d # flip to other hemisphere - break - valid_h[vi] = (nm, parent_pos + bl * d, bl) - - # --- Step 3: emit placed H atoms --- - for h_nm, h_pos, _ in valid_h: - h_x.append(h_pos[0]) - h_y.append(h_pos[1]) - h_z.append(h_pos[2]) - h_names_out.append(h_nm) - h_altlocs.append(altloc) - h_resnames.append(rn) - h_chainids.append(chainid) - h_resseqs.append(resseq) - h_icodes.append(icode) - h_occ.append(conf_occ[pidx]) - h_bfac.append(conf_bfac[pidx]) - h_atom_types.append(conf_atom_type[pidx]) - h_insert_after.append(e - 1) - - n_h_placed = len(h_x) - if n_h_placed == 0: - if verbose > 0: - print("No hydrogen atoms to add; returning copy.") - return self.copy() - - if verbose > 0: - print(f"Placing {n_h_placed} hydrogen atoms...") - - # Build H DataFrame in one shot - h_df = pd.DataFrame( - { - "ATOM": h_atom_types, - "serial": 0, - "name": h_names_out, - "altloc": h_altlocs, - "resname": h_resnames, - "chainid": h_chainids, - "resseq": h_resseqs, - "icode": h_icodes, - "x": h_x, - "y": h_y, - "z": h_z, - "occupancy": h_occ, - "tempfactor": h_bfac, - "element": "H", - "charge": 0, - "anisou_flag": False, - "u11": 0.0, - "u22": 0.0, - "u33": 0.0, - "u12": 0.0, - "u13": 0.0, - "u23": 0.0, - } - ) - insert_after = np.array(h_insert_after) - - # Interleave: assign sort keys - n_orig = len(self.pdb) - sort_key = np.empty(n_orig + n_h_placed, dtype=np.float64) - sort_key[:n_orig] = np.arange(n_orig, dtype=np.float64) - _, inv, counts = np.unique( - insert_after, return_inverse=True, return_counts=True - ) - cumcount = np.zeros(n_h_placed, dtype=np.float64) - group_running = np.zeros(len(counts), dtype=np.float64) - for i in range(n_h_placed): - g = inv[i] - cumcount[i] = group_running[g] - group_running[g] += 1 - sort_key[n_orig:] = ( - insert_after + 0.5 + cumcount * (0.4 / np.maximum(counts[inv], 1)) - ) - - augmented_df = pd.concat([self.pdb, h_df], ignore_index=True) - augmented_df = augmented_df.iloc[ - np.argsort(sort_key, kind="stable") - ].reset_index(drop=True) - augmented_df["serial"] = np.arange(1, len(augmented_df) + 1) - augmented_df["index"] = np.arange(len(augmented_df)) - - for col in ( - "x", - "y", - "z", - "occupancy", - "tempfactor", - "u11", - "u22", - "u33", - "u12", - "u13", - "u23", - ): - augmented_df[col] = pd.to_numeric( - augmented_df[col], errors="coerce" - ).astype(float) - augmented_df["serial"] = augmented_df["serial"].astype(int) - augmented_df["resseq"] = augmented_df["resseq"].astype(int) - augmented_df["charge"] = augmented_df["charge"].fillna(0).astype(int) - augmented_df["anisou_flag"] = augmented_df["anisou_flag"].astype(bool) - augmented_df[["altloc", "icode"]] = augmented_df[["altloc", "icode"]].fillna("") - augmented_df["element"] = ( - augmented_df["element"].astype(str).str.strip().str.capitalize() + plan = plan_hydrogens( + restraints.topology, restraints.cif_dict, xyz, verbose=verbose ) - augmented_df.attrs["cell"] = self.pdb.attrs.get("cell") - augmented_df.attrs["spacegroup"] = self.pdb.attrs.get("spacegroup", "P 1") - - new_model = self._new_model_from_df(augmented_df, strip_H=False) - - if verbose > 0: - n_h = (new_model.pdb["element"] == "H").sum() - print(f" New model: {len(new_model.pdb)} atoms ({n_h} H)") - - # --- Phase C (optional): LBFGS geometry optimization --- if optimize: - new_model.freeze_all() - new_model.unfreeze_selection("element H", targets="xyz") - refinable_params = [p for p in new_model.parameters() if p.numel() > 0] - if refinable_params: - try: - from torchref.refinement.targets.combined import ( - TotalGeometryTarget, - ) - - geom_target = TotalGeometryTarget(new_model, verbose=0) - targets = { - n: geom_target[n] - for n in ("bond", "angle", "torsion", "chiral") - } - - def _geom_loss(): - total = torch.tensor(0.0, device=self.device) - for t in targets.values(): - val = t() - if torch.isfinite(val): - total = total + val - return total - - if verbose > 0: - with torch.no_grad(): - init_l = _geom_loss() - print(f" Geometry loss before: {init_l.item():.4f}") - for m in new_model.modules(): - if hasattr(m, "reset_forward_cache"): - m.reset_forward_cache() - - opt = torch.optim.LBFGS( - refinable_params, - lr=0.1, - max_iter=max_iter, - history_size=100, - line_search_fn="strong_wolfe", - ) - best_loss = float("inf") - best_params = [p.data.clone() for p in refinable_params] - - def closure(): - opt.zero_grad() - loss = _geom_loss() - if loss.requires_grad and torch.isfinite(loss): - loss.backward() - for p in refinable_params: - if p.grad is not None: - p.grad.nan_to_num_(nan=0.0, posinf=0.0, neginf=0.0) - return loss - - for _ in range(lbfgs_steps): - opt.step(closure) - with torch.no_grad(): - cur = _geom_loss() - if torch.isfinite(cur) and cur.item() < best_loss: - best_loss = cur.item() - best_params = [p.data.clone() for p in refinable_params] - with torch.no_grad(): - for p, bp in zip(refinable_params, best_params): - p.data.copy_(bp) - if verbose > 0: - with torch.no_grad(): - fin_l = _geom_loss() - print(f" Geometry loss after: {fin_l.item():.4f}") - except Exception as e: - if verbose > 0: - print(f" Warning: optimization failed: {e}") - new_model.set_default_masks() - new_model.unfreeze_all() + optimise_free_torsions(plan, restraints.topology, xyz) if verbose > 0: - print(" Hydrogenation complete.") + print(f"Adding {plan.n_hydrogens} hydrogens") + augmented = augment_atom_table(self.pdb, plan, restraints.topology) + return self._new_model_from_df(augmented, strip_H=False) - return new_model def state_dict(self, destination=None, prefix="", keep_vars=False): """ diff --git a/torchref/restraints/hydrogen_topology.py b/torchref/restraints/hydrogen_topology.py index 6a92b04b..f5276621 100644 --- a/torchref/restraints/hydrogen_topology.py +++ b/torchref/restraints/hydrogen_topology.py @@ -163,20 +163,23 @@ def __repr__(self) -> str: # --------------------------------------------------------------------------- +#: Parsed monomer templates, keyed by residue name, shared across calls. Values are +#: None where the CIF is missing or carries no usable atom coordinates. +_TEMPLATE_CACHE: Dict = {} + + def _load_cif_hydrogen_info(pdb, verbose: int = 0) -> Dict: """``{resname: entry | None}`` H topology, ``None`` where the CIF is unusable. - Populates and returns the shared ``Model._hydrogenate_cif_cache``, so entries - from an earlier ``Model.hydrogenate()`` are reused. Each entry carries ``ids``, + Populates and returns :data:`_TEMPLATE_CACHE`. Each entry carries ``ids``, ``elems``, ``coords``, ``is_h``, ``id_to_idx``, ``heavy_names``, ``heavy_coords``, ``h_names``, ``h_coords``, ``parent_map``, ``ideal_bl`` and ``heavy_neighbor_map``. """ - from torchref.model.model import Model from torchref.restraints.library import MonomerLibraryManager lib = MonomerLibraryManager(verbose=0) - cache = Model._hydrogenate_cif_cache + cache = _TEMPLATE_CACHE for rn in pdb["resname"].unique(): rn_str = str(rn).strip() diff --git a/torchref/topology/__init__.py b/torchref/topology/__init__.py index 78cc0dea..cc025aa3 100644 --- a/torchref/topology/__init__.py +++ b/torchref/topology/__init__.py @@ -10,12 +10,19 @@ connectivity can carry monomer-library targets, force-field parameters, or ADP-similarity sigmas without duplicating the edges. -Build one with :func:`build_topology`. +Build one with :func:`build_topology`. :func:`plan_hydrogens` uses it to instantiate +monomer templates, which is how hydrogens are generated. """ from .atom_graph import AtomGraph from .build import build_topology, build_topology_with_values from .edges import ORIGIN_ORDER, EdgeBlock +from .hydrogens import ( + HydrogenPlan, + augment_atom_table, + optimise_free_torsions, + plan_hydrogens, +) from .residue_graph import ResidueGraph from .restraint_sets import assemble_entries, max_period from .templates import resolve_template_keys @@ -31,5 +38,9 @@ "build_topology_with_values", "assemble_entries", "max_period", + "HydrogenPlan", + "plan_hydrogens", + "optimise_free_torsions", + "augment_atom_table", "resolve_template_keys", ] diff --git a/torchref/topology/hydrogens.py b/torchref/topology/hydrogens.py new file mode 100644 index 00000000..5ede3a54 --- /dev/null +++ b/torchref/topology/hydrogens.py @@ -0,0 +1,803 @@ +"""Hydrogen generation as a graph operation: expand the template, map it on. + +A monomer template already carries its hydrogens, with coordinates and with bonds naming +each one's parent. Generating hydrogens is therefore template instantiation, not +geometry reconstruction: align the template onto the heavy atoms that are present, +read the hydrogen positions off it, and correct each to its ideal bond length. + +Two things the bond graph decides that a distance criterion previously guessed at: + +* **How many hydrogens a parent can carry.** The count is the parent's standard valence + minus the heavy atoms actually bonded to it, taken from the graph. A distance sweep + gets this wrong on a distorted or predicted model, where a bond can fall outside the + window; and it cannot distinguish a real bond from two atoms that merely sit close. +* **Which hydrogens have a free torsion.** A hydrogen whose parent has exactly one heavy + neighbour -- hydroxyl, thiol, amine, methyl -- can rotate about the parent-neighbour + axis, and the template's angle for it is arbitrary. Those get scanned; the rest are + fully determined by the template and are left alone. +""" + +from dataclasses import dataclass +from typing import Dict, List, Optional, Tuple + +import numpy as np + +#: Standard heavy-atom valences, used to cap how many hydrogens a parent may take. The +#: fallback of 4 matches the previous behaviour for elements not listed. +STANDARD_VALENCE = {"C": 4, "N": 3, "O": 2, "S": 2} +_DEFAULT_VALENCE = 4 + +#: A template hydrogen further than this from its parent after alignment is discarded +#: rather than corrected: the alignment for that centre is too poor to trust. +MAX_PLACEMENT_DISTANCE = 1.5 + +#: Angles tried when scanning a free torsion. 24 gives 15-degree resolution, which is +#: finer than the placement error the alignment itself carries. +TORSION_SCAN_STEPS = 24 + +#: Heavy atoms beyond this distance cannot clash with a hydrogen, so the scan ignores +#: them. +CLASH_CUTOFF = 4.0 + + +@dataclass +class HydrogenPlan: + """Hydrogens to add, and where to put them. + + Parameters + ---------- + name, element, altloc : numpy.ndarray + Per-hydrogen identity, shape ``(H,)``. + residue : numpy.ndarray + Residue index each hydrogen belongs to, shape ``(H,)``. + parent : numpy.ndarray + Atom-table index of each hydrogen's parent, shape ``(H,)``. + position : numpy.ndarray + Cartesian coordinates, shape ``(H, 3)``. + bond_length : numpy.ndarray + Ideal parent-hydrogen distance, shape ``(H,)``. + group : numpy.ndarray + Free-torsion group id, shape ``(H,)``; ``-1`` for a hydrogen whose position the + template determines. Hydrogens on one parent that rotate together share an id. + """ + + name: np.ndarray + element: np.ndarray + altloc: np.ndarray + residue: np.ndarray + parent: np.ndarray + position: np.ndarray + bond_length: np.ndarray + group: np.ndarray + + @property + def n_hydrogens(self) -> int: + """How many hydrogens the plan adds.""" + return len(self.name) + + @property + def rotatable(self) -> np.ndarray: + """Mask of hydrogens whose torsion is free.""" + return self.group >= 0 + + def __repr__(self) -> str: + return ( + f"HydrogenPlan(n_hydrogens={self.n_hydrogens}, " + f"free_torsions={len(set(self.group[self.rotatable].tolist()))})" + ) + + +def _kabsch(source: np.ndarray, target: np.ndarray) -> Tuple[np.ndarray, np.ndarray]: + """Rotation and translation carrying ``source`` onto ``target``. + + Reflections are excluded, so the template's chirality survives the alignment. + + Returns + ------- + rotation, translation : numpy.ndarray + Shapes ``(3, 3)`` and ``(3,)``; ``rotation @ p + translation`` maps a source + point. + """ + source_centre, target_centre = source.mean(0), target.mean(0) + covariance = (source - source_centre).T @ (target - target_centre) + u, _, vt = np.linalg.svd(covariance) + flip = np.diag([1.0, 1.0, 1.0 if np.linalg.det(vt.T @ u.T) > 0 else -1.0]) + rotation = vt.T @ flip @ u.T + return rotation, target_centre - rotation @ source_centre + + +def _template(cif_dict: Dict, resname: str) -> Optional[Dict]: + """Template atoms, hydrogen parents, ideal bond lengths and heavy adjacency. + + Read from the restraint dictionary the caller already loaded. The link modifications + are deliberately not consulted: they rewrite restraint sections, not atom lists, so + they cannot say which hydrogens a linked residue keeps. The valence cap answers that + from the bond graph instead. + + Returns + ------- + dict or None + None when the component is absent or its atoms carry no coordinates. + """ + component = cif_dict.get(resname) + if component is None: + return None + atoms = component.get("atoms") + if atoms is None or len(atoms) == 0: + return None + if not all(column in atoms.columns for column in ("x", "y", "z")): + return None + + ids = atoms["atom_id"].astype(str).str.strip().values + elements = atoms["type_symbol"].astype(str).str.strip().values.astype(" 0: + import pandas as pd + + first = bonds["atom1"].astype(str).str.strip().values + second = bonds["atom2"].astype(str).str.strip().values + values = pd.to_numeric(bonds["value"], errors="coerce").values + for i in range(len(first)): + a, b = first[i], second[i] + ia, ib = id_to_index.get(a), id_to_index.get(b) + if ia is None or ib is None: + continue + if is_h[ia] and not is_h[ib]: + parent_of[a] = b + if np.isfinite(values[i]): + ideal_length[a] = float(values[i]) + elif is_h[ib] and not is_h[ia]: + parent_of[b] = a + if np.isfinite(values[i]): + ideal_length[b] = float(values[i]) + elif not is_h[ia] and not is_h[ib]: + heavy_adjacency.setdefault(a, []).append(b) + heavy_adjacency.setdefault(b, []).append(a) + + return { + "ids": ids, + "elements": elements, + "coords": coords, + "is_h": is_h, + "id_to_index": id_to_index, + "heavy_names": ids[~is_h], + "heavy_coords": coords[~is_h], + "h_names": ids[is_h], + "parent_of": parent_of, + "ideal_length": ideal_length, + "heavy_adjacency": heavy_adjacency, + } + + +def _orthonormal_frame(axis: np.ndarray) -> Tuple[np.ndarray, np.ndarray]: + """Two unit vectors completing ``axis`` into a right-handed frame.""" + seed = np.array([1.0, 0.0, 0.0]) + if abs(float(axis @ seed)) > 0.9: + seed = np.array([0.0, 1.0, 0.0]) + first = seed - axis * float(axis @ seed) + first = first / np.linalg.norm(first) + return first, np.cross(axis, first) + + +def _axis_frame_placement( + template: Dict, + parent_name: str, + neighbour_name: str, + parent_position: np.ndarray, + neighbour_position: np.ndarray, + h_names: List[str], +) -> Optional[np.ndarray]: + """Hydrogen positions for a centre with a single heavy neighbour. + + Maps the template's local geometry onto the model by carrying the parent-neighbour + axis across and completing the frame arbitrarily. Every bond angle at the parent is + preserved exactly; only the rotation about the axis is arbitrary, which is correct: + that is the degree of freedom the template cannot know, and + :func:`optimise_free_torsions` chooses it. + """ + index = template["id_to_index"] + if parent_name not in index or neighbour_name not in index: + return None + + template_axis = ( + template["coords"][index[parent_name]] + - template["coords"][index[neighbour_name]] + ) + model_axis = parent_position - neighbour_position + for vector in (template_axis, model_axis): + if np.linalg.norm(vector) < 1e-8: + return None + template_axis = template_axis / np.linalg.norm(template_axis) + model_axis = model_axis / np.linalg.norm(model_axis) + + t_first, t_second = _orthonormal_frame(template_axis) + m_first, m_second = _orthonormal_frame(model_axis) + + positions = [] + for name in h_names: + if name not in index: + return None + offset = ( + template["coords"][index[name]] - template["coords"][index[parent_name]] + ) + positions.append( + parent_position + + float(offset @ template_axis) * model_axis + + float(offset @ t_first) * m_first + + float(offset @ t_second) * m_second + ) + return np.array(positions) + + +def _half_hydrogen_angle(template: Dict, parent_name: str, h_names: List[str]) -> float: + """Half the hydrogen-parent-hydrogen angle, from the template where it has one.""" + index = template["id_to_index"] + if len(h_names) == 2 and all(n in index for n in h_names): + parent = template["coords"][index[parent_name]] + first = template["coords"][index[h_names[0]]] - parent + second = template["coords"][index[h_names[1]]] - parent + norms = np.linalg.norm(first) * np.linalg.norm(second) + if norms > 1e-12: + cosine = float(np.clip((first @ second) / norms, -1.0, 1.0)) + return 0.5 * float(np.arccos(cosine)) + # Tetrahedral, as a fallback for a template that does not carry both hydrogens. + return 0.5 * np.arccos(-1.0 / 3.0) + + +def _heavy_neighbours(topology, atom_index: int) -> np.ndarray: + """Rows of the heavy atoms bonded to ``atom_index``, from the bond graph. + + Coordinate-independent, unlike a distance sweep: a stretched bond in a predicted or + mid-refinement model still counts, and two atoms that merely sit close do not. + """ + neighbours = topology.atoms.neighbors(atom_index) + if neighbours.numel() == 0: + return np.zeros(0, dtype=np.int64) + heavy = neighbours[~topology.atoms.is_hydrogen[neighbours]] + return heavy.cpu().numpy() + + +def _template_bond_length(template: Dict, parent_name: str, h_name: str) -> float: + """Parent-hydrogen distance in the template, or NaN if either atom is missing.""" + index = template["id_to_index"] + if parent_name not in index or h_name not in index: + return float("nan") + return float( + np.linalg.norm( + template["coords"][index[h_name]] - template["coords"][index[parent_name]] + ) + ) + + +def _place_group( + template: Dict, + parent_name: str, + parent_position: np.ndarray, + neighbour_positions: np.ndarray, + heavy_bonded: int, + h_names: List[str], + lengths: np.ndarray, + name_to_row: Dict[str, int], + coords: np.ndarray, + template_names_of: List[str], +) -> Optional[np.ndarray]: + """Positions for the hydrogens on one parent, by the first strategy that applies. + + In order: + + 1. **Template frame.** A Kabsch fit over the parent and its immediate heavy + neighbours, used only when the template knows every heavy atom actually bonded to + the parent. Reproduces the library geometry exactly, torsions included. + 2. **Geometric construction.** Directions from the bonded neighbours alone, for a + centre whose template is missing a real substituent -- a peptide-linked backbone + nitrogen being the common case. + 3. **Axis frame.** For a single-neighbour centre, the template geometry carried over + about the one bond, leaving the rotation about it for the scan to choose. + + Returns None when none applies, so the caller can count the hydrogen as undetermined + rather than putting it somewhere arbitrary. + """ + covered = _template_covers_neighbours(template, parent_name, heavy_bonded) + + if covered: + alignment = _alignment_for(template, parent_name, name_to_row, coords) + if alignment is not None: + matrix, offset = alignment + index = template["id_to_index"] + positions = [] + for name, length in zip(h_names, lengths): + direction = ( + matrix @ template["coords"][index[name]] + offset + ) - parent_position + distance = float(np.linalg.norm(direction)) + if distance < 1e-6 or distance > MAX_PLACEMENT_DISTANCE: + positions = None + break + positions.append(parent_position + direction * (length / distance)) + if positions is not None: + return np.array(positions) + + if heavy_bonded >= 2: + directions = _construct_directions( + parent_position, + neighbour_positions, + len(h_names), + _half_hydrogen_angle(template, parent_name, h_names), + ) + if directions is not None: + return parent_position + directions * lengths[:, None] + + if heavy_bonded == 1 and template_names_of: + positions = _axis_frame_placement( + template, + parent_name, + template_names_of[0], + parent_position, + neighbour_positions[0], + h_names, + ) + if positions is None: + return None + # Rescale along each direction so the bond length is the library value rather + # than the template's own geometry, which differs from it by ~0.002 A. Scaling + # along the direction leaves every bond angle untouched. + offsets = positions - parent_position + norms = np.linalg.norm(offsets, axis=1) + if (norms < 1e-8).any(): + return None + return parent_position + offsets * (lengths / norms)[:, None] + + return None + + +def plan_hydrogens(topology, cif_dict: Dict, xyz, verbose: int = 0) -> HydrogenPlan: + """Decide which hydrogens to add and place them from the template. + + Parameters + ---------- + topology : Topology + Supplies the residue partition, the per-residue template key and the bond graph + the valence cap and the free-torsion test read. + cif_dict : dict + Restraint dictionary, keyed by residue name; must carry an ``atoms`` section + with coordinates for a residue to be hydrogenated. + xyz : torch.Tensor + Current coordinates, shape ``(N, 3)``. + verbose : int, default 0 + Verbosity level. + + Returns + ------- + HydrogenPlan + """ + coords = np.asarray(xyz.detach().cpu(), dtype=np.float64) + residues = topology.residues + atoms = topology.atoms + names = atoms.name.astype(str) + altlocs = atoms.altloc.astype(str) + + out: Dict[str, list] = { + k: [] + for k in ( + "name", + "altloc", + "residue", + "parent", + "position", + "bond_length", + "group", + ) + } + next_group = 0 + n_unplaceable = 0 + n_no_template = 0 + + for residue in range(residues.n_residues): + resname = str(residues.resname[residue]).strip() + template = _template(cif_dict, resname) + if template is None: + # No usable template. In practice these are the single-atom residues -- + # waters and ions -- which the restraint dictionary omits because they carry + # no intra-residue geometry. They could not be hydrogenated anyway: one + # heavy atom gives no frame to orient a template against and no bond to + # rotate about, so a water's hydrogens would point somewhere arbitrary. + n_no_template += 1 + continue + + rows = np.arange( + int(residues.atom_start[residue]), int(residues.atom_end[residue]) + ) + present = set(names[rows]) + candidates = [h for h in template["h_names"] if h not in present] + if not candidates: + continue + + for altloc, conformer in _conformer_rows(rows, altlocs): + name_to_row = {} + for row in conformer: + name_to_row.setdefault(names[row], row) + + # Hydrogens grouped by the parent they hang off, in name order so the cap + # below takes a deterministic subset. + by_parent: Dict[str, List[str]] = {} + for h in sorted(candidates): + parent = template["parent_of"].get(h) + if parent is not None and parent in name_to_row: + by_parent.setdefault(parent, []).append(h) + + for parent_name, group in by_parent.items(): + parent_row = name_to_row[parent_name] + parent_position = coords[parent_row] + + heavy_rows = _heavy_neighbours(topology, parent_row) + heavy_bonded = len(heavy_rows) + element = str( + template["elements"][template["id_to_index"][parent_name]] + ).upper() + allowed = max( + 0, + STANDARD_VALENCE.get(element, _DEFAULT_VALENCE) - heavy_bonded, + ) + group = group[:allowed] + if not group: + continue + + lengths = np.array( + [ + template["ideal_length"].get( + h, _template_bond_length(template, parent_name, h) + ) + for h in group + ] + ) + if not np.isfinite(lengths).all(): + n_unplaceable += len(group) + continue + + placed = _place_group( + template, + parent_name, + parent_position, + coords[heavy_rows] if heavy_bonded else np.zeros((0, 3)), + heavy_bonded, + group, + lengths, + name_to_row, + coords, + template_names_of=[ + n + for n in template["heavy_adjacency"].get(parent_name, []) + if n in name_to_row + ], + ) + if placed is None: + n_unplaceable += len(group) + continue + + # One free-torsion group per parent with a single heavy neighbour: its + # hydrogens rotate together about that one bond. + group_id = -1 + if heavy_bonded == 1: + group_id = next_group + next_group += 1 + + for h, position, length in zip(group, placed, lengths): + out["name"].append(h) + out["altloc"].append(altloc) + out["residue"].append(residue) + out["parent"].append(int(parent_row)) + out["position"].append(position) + out["bond_length"].append(float(length)) + out["group"].append(group_id) + + if verbose > 1: + print( + f"Planned {len(out['name'])} hydrogens " + f"({n_no_template} residues with no usable template, " + f"{n_unplaceable} hydrogens whose direction was not determined)" + ) + + return HydrogenPlan( + name=np.array(out["name"], dtype=" Optional[Tuple[np.ndarray, np.ndarray]]: + """Template-to-model transform for one hydrogen-bearing centre. + + Fitted over the parent and its **immediate** heavy neighbours only. That set is the + rigid unit which fixes the hydrogen directions: the bond lengths and angles at the + parent are library constants, while anything further out sits across a rotatable + torsion whose value is the model's, not the template's. + + Reaching one bond further -- as a whole-residue or two-shell fit does -- makes the + rotation compromise between the real local geometry and a torsion the model does not + share, which lands hydrogens well off their parent. Measured on 7L84 the two-shell + fit around ``CB`` aligned to 0.75 A RMSD and put 12% of side-chain hydrogens beyond + 1.5 A of the atom they belong to. + + Returns None when fewer than three neighbours match, which leaves the rotation + undetermined; the caller then constructs the direction from the bond graph instead. + """ + local = [parent_name] + [ + n for n in sorted(template["heavy_adjacency"].get(parent_name, [])) + ] + matched = [n for n in local if n in name_to_row and n in template["id_to_index"]] + if len(matched) < 3: + return None + source = np.array([template["coords"][template["id_to_index"][n]] for n in matched]) + target = np.array([coords[name_to_row[n]] for n in matched]) + return _kabsch(source, target) + + +def _template_covers_neighbours( + template: Dict, parent_name: str, graph_heavy_count: int +) -> bool: + """Whether the template knows every heavy atom actually bonded to the parent. + + A peptide-linked backbone nitrogen is bonded to the preceding residue's carbon, + which the free-amino-acid template has never heard of. Placing its hydrogen from the + template frame would ignore that substituent and drop the hydrogen on top of it, so + those centres are built from the graph instead. + """ + return len(template["heavy_adjacency"].get(parent_name, [])) >= graph_heavy_count + + +def _construct_directions( + parent: np.ndarray, + neighbours: np.ndarray, + n_hydrogens: int, + half_angle: float, +) -> Optional[np.ndarray]: + """Unit directions for hydrogens on a centre, from its bonded neighbours alone. + + Two rules cover every centre whose orientation the neighbours determine: + + * one hydrogen, any number of neighbours -- it opposes the sum of the bond unit + vectors, which is where the remaining valence points. This is the backbone amide + and alpha hydrogens. + * two hydrogens on a two-neighbour centre -- they straddle that same direction, + opened out by ``half_angle`` in the plane perpendicular to the neighbour pair. + + Anything else (three hydrogens, or two on a one-neighbour centre) has a free + torsion and is handled by the scan, not here. + + Returns + ------- + numpy.ndarray or None + Shape ``(n_hydrogens, 3)`` unit vectors, or None when the rules do not apply or + the geometry is degenerate. + """ + bonds = neighbours - parent + lengths = np.linalg.norm(bonds, axis=1) + if (lengths < 1e-8).any(): + return None + units = bonds / lengths[:, None] + + total = units.sum(0) + norm = np.linalg.norm(total) + if norm < 1e-6: + return None + opposed = -total / norm + + if n_hydrogens == 1: + return opposed[None, :] + + if n_hydrogens == 2 and len(units) == 2: + perpendicular = np.cross(units[0], units[1]) + perpendicular_norm = np.linalg.norm(perpendicular) + if perpendicular_norm < 1e-6: + return None + perpendicular = perpendicular / perpendicular_norm + cos, sin = np.cos(half_angle), np.sin(half_angle) + return np.array( + [ + opposed * cos + perpendicular * sin, + opposed * cos - perpendicular * sin, + ] + ) + + return None + + +def optimise_free_torsions( + plan: HydrogenPlan, + topology, + xyz, + steps: int = TORSION_SCAN_STEPS, +) -> HydrogenPlan: + """Rotate each free-torsion hydrogen group to the least-clashing angle. + + The template's dihedral for a hydroxyl, thiol, amine or methyl hydrogen is whatever + the library happened to deposit, so it carries no information. Each group is scanned + about its parent-neighbour axis and scored by repulsion against nearby heavy atoms; + the best angle wins. + + Repulsion only: this removes clashes but does not seek hydrogen bonds, so a hydroxyl + is placed out of the way rather than donated to an acceptor. Scoring donors properly + is a separate piece of physics. + + Returns + ------- + HydrogenPlan + The same plan with ``position`` updated in place for the scanned groups. + """ + if plan.n_hydrogens == 0 or not plan.rotatable.any(): + return plan + + coords = np.asarray(xyz.detach().cpu(), dtype=np.float64) + is_h = topology.atoms.is_hydrogen.cpu().numpy() + heavy_rows = np.nonzero(~is_h)[0] + heavy_coords = coords[heavy_rows] + + angles = np.linspace(0.0, 2.0 * np.pi, steps, endpoint=False) + + for group_id in sorted(set(plan.group[plan.rotatable].tolist())): + members = np.nonzero(plan.group == group_id)[0] + parent_row = int(plan.parent[members[0]]) + parent = coords[parent_row] + + neighbours = topology.atoms.neighbors(parent_row).cpu().numpy() + heavy_neighbours = neighbours[~is_h[neighbours]] + if len(heavy_neighbours) != 1: + continue + axis = parent - coords[int(heavy_neighbours[0])] + norm = np.linalg.norm(axis) + if norm < 1e-8: + continue + axis = axis / norm + + # Heavy atoms that could clash, excluding the parent and its own neighbour. + near = np.nonzero( + (np.linalg.norm(heavy_coords - parent, axis=1) < CLASH_CUTOFF) + )[0] + exclude = {parent_row, int(heavy_neighbours[0])} + near_coords = np.array( + [heavy_coords[i] for i in near if heavy_rows[i] not in exclude] + ) + + offsets = plan.position[members] - parent + best_angle, best_score = 0.0, np.inf + for angle in angles: + rotated = _rotate_about(offsets, axis, angle) + if len(near_coords) == 0: + best_angle = 0.0 + break + trial = parent + rotated + separation = np.linalg.norm( + trial[:, None, :] - near_coords[None, :, :], axis=-1 + ) + score = float((1.0 / np.maximum(separation, 0.5) ** 2).sum()) + if score < best_score: + best_score, best_angle = score, float(angle) + + if best_angle != 0.0: + plan.position[members] = parent + _rotate_about(offsets, axis, best_angle) + + return plan + + +def _rotate_about(vectors: np.ndarray, axis: np.ndarray, angle: float) -> np.ndarray: + """Rotate ``vectors`` about a unit ``axis`` through the origin, Rodrigues form.""" + cos, sin = np.cos(angle), np.sin(angle) + return ( + vectors * cos + + np.cross(axis, vectors) * sin + + np.outer(vectors @ axis, axis) * (1.0 - cos) + ) + + +__all__ = [ + "HydrogenPlan", + "plan_hydrogens", + "optimise_free_torsions", + "augment_atom_table", + "STANDARD_VALENCE", + "MAX_PLACEMENT_DISTANCE", + "TORSION_SCAN_STEPS", +] + + +def augment_atom_table(pdb, plan: HydrogenPlan, topology): + """Insert a plan's hydrogens into an atom table. + + Each hydrogen is inserted immediately after the residue it belongs to, not appended + at the end: the residue partition is built from contiguous runs of + ``(chain, resseq, icode)``, so appending would split every hydrogenated residue into + two nodes. + + Rows are copied from the parent, then the hydrogen's own name, element and position + are written over them. Everything else -- chain, residue, altloc, occupancy, + B-factor, record type -- is inherited, so a hydrogen starts from its parent's + displacement parameter and refines from there. + + Parameters + ---------- + pdb : pandas.DataFrame + Atom table to extend. + plan : HydrogenPlan + topology : Topology + Supplies the residue partition the insertion points come from. + + Returns + ------- + pandas.DataFrame + A new table with ``serial`` and ``index`` renumbered. + """ + import pandas as pd + + if plan.n_hydrogens == 0: + return pdb.copy() + + by_residue: Dict[int, List[int]] = {} + for i, residue in enumerate(plan.residue.tolist()): + by_residue.setdefault(residue, []).append(i) + + pieces = [] + for residue in range(topology.n_residues): + start = int(topology.residues.atom_start[residue]) + end = int(topology.residues.atom_end[residue]) + pieces.append(pdb.iloc[start:end]) + + members = by_residue.get(residue) + if not members: + continue + rows = pdb.loc[pdb.index[plan.parent[members]]].copy() + rows["name"] = plan.name[members] + rows["element"] = plan.element[members] + rows["altloc"] = plan.altloc[members] + rows[["x", "y", "z"]] = plan.position[members] + if "anisou_flag" in rows.columns: + rows["anisou_flag"] = False + for column in ("u11", "u22", "u33", "u12", "u13", "u23"): + if column in rows.columns: + rows[column] = float("nan") + pieces.append(rows) + + augmented = pd.concat(pieces, ignore_index=True) + augmented["index"] = augmented.index.to_numpy(dtype=int) + if "serial" in augmented.columns: + augmented["serial"] = augmented.index.to_numpy(dtype=int) + 1 + augmented.attrs = dict(pdb.attrs) + return augmented From 40fd12e63760d977c8808ddc8be3783f16ee2ffb Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 12:46:11 +0200 Subject: [PATCH 060/250] Weight the difference target by the activation spread, and calibrate the intensity target Three changes, all driven by running the figure-4 pair rather than by the unit tests. The difference target now adds the activation-derived contamination to its variance, so --lambda-twin reweights it with no intensity data needed. This is the calibrated form of the k-weighting difference maps apply by hand. Measured on figure 4 the contamination is 1.7e-3 of a single reflection's sigma, so the weight never falls below 0.996 and the refinement is unchanged -- the effect is systematic and only visible summed over the dataset, which a per-reflection weight cannot see. The intensity target gains a base_weight, calibrated by matching its gradient norm to the difference target's over the refined parameters. Loss magnitude is the wrong thing to equalise: on figure 4 the loss ratio was ~3e6 while the gradient ratio was 4.5, and weighting by the former would have been wrong by orders of magnitude. Uncalibrated it bought R-free by moving the model twice as far as the data supports, taking the ligand difference-density correlation from 0.867 to 0.808. Non-finite observed intensities no longer poison the gradient. Real reflection files carry them, and torch.where masks the value while still backpropagating NaN through the branch it discards -- every optimizer step was rejected and the refinement silently did nothing. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- docs/changelog.rst | 3 + tests/integration/test_cli_two_moment_mtz.py | 54 +++++- .../refinement/test_activation_weighting.py | 174 ++++++++++++++++++ .../refinement/test_two_moment_intensity.py | 163 ++++++++++++++++ torchref/cli/collection_difference_refine.py | 46 +++-- .../targets/collection/intensity.py | 106 ++++++++++- .../refinement/targets/collection/xray.py | 50 +++++ 7 files changed, 563 insertions(+), 33 deletions(-) create mode 100644 tests/unit/refinement/test_activation_weighting.py diff --git a/docs/changelog.rst b/docs/changelog.rst index e54ca2ae..163bce64 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -7,6 +7,9 @@ Version 0.6.4 - Added a reader for CrystFEL ``partialator`` ``.hkl`` reflection lists, via ``ReflectionData.load_crystfel_hkl`` - Added ``FcalcDataset.add_noise`` and the ``torchref.simulate-noisy-data`` CLI, which simulate merged intensities from a structure and report R-split and CC between two independent half-datasets - Simulated intensities keep their negative values; only the derived amplitude is clamped, since clamping the intensity biases the weak reflections upward +- The difference target now inflates its variance by the activation-derived contamination, so a non-zero ``--lambda-twin`` reweights it without needing intensity data +- ``CollectionTwoMomentIntensityTarget`` carries a ``base_weight``, calibrated against the difference target's gradient norm so an intensity likelihood does not swamp the restraints +- Fixed non-finite observed intensities poisoning the two-moment gradient, which silently froze refinement rather than failing - Added ``CollectionTwoMomentIntensityTarget``, fitting merged intensities as ``|F(alpha)|^2 + sigma_alpha^2 |dF|^2`` to account for crystal-to-crystal spread in activation - Added ``--two-moment`` / ``--lambda-twin`` / ``--refine-lambda-twin`` to ``torchref.difference-refine``, and the activation moments to its JSON summary - ``torchref.difference-refine`` writes thirteen further MTZ columns under ``--two-moment``, including decontaminated difference amplitudes and the ``DDF`` diagnostic diff --git a/tests/integration/test_cli_two_moment_mtz.py b/tests/integration/test_cli_two_moment_mtz.py index 1dc69030..06ebe919 100644 --- a/tests/integration/test_cli_two_moment_mtz.py +++ b/tests/integration/test_cli_two_moment_mtz.py @@ -189,6 +189,19 @@ def test_ivar_alpha_is_sigma_sq_times_the_squared_difference(self, two_moment_mt def test_the_two_moment_intensity_exceeds_the_coherent_one_by_the_variance( self, two_moment_mtz ): + """``Ic_2mom - Ic_coh`` must equal ``IVAR_ALPHA``, to whatever precision float32 + leaves after the cancellation. + + This is a catastrophic-cancellation case, and the tolerance is computed rather + than guessed. The variance term is ~2.6e-6 of the intensity on this fixture, + while float32 resolves ~1.2e-7 of it -- so only about one significant digit of + the difference survives, and any fixed tolerance would either pass vacuously or + fail for reasons that have nothing to do with the code. + + The target itself never forms this difference (it computes + ``|F|**2 + sigma**2 |dF|**2`` directly), so the loss is unaffected; it is + recovering the variance term from the two published columns that is lossy. + """ import numpy as np mtz, _ = two_moment_mtz @@ -197,18 +210,41 @@ def test_the_two_moment_intensity_exceeds_the_coherent_one_by_the_variance( two = df["Ic_light_2mom"].to_numpy().astype(float) ivar = df["IVAR_ALPHA"].to_numpy().astype(float) - scale = max(float(np.abs(ivar).max()), 1e-30) - assert np.abs((two - coh) - ivar).max() / scale < 1e-3 + # Absolute error float32 can leave in the difference of two intensities. + eps32 = float(np.finfo(np.float32).eps) + floor = eps32 * np.maximum(np.abs(coh), np.abs(two)) + residual = np.abs((two - coh) - ivar) + + assert (residual <= 4.0 * floor + 1e-12).all(), ( + f"recovered variance term differs from IVAR_ALPHA by more than float32 " + f"cancellation allows: worst {np.max(residual / (floor + 1e-30)):.1f} ulp" + ) # The variance term has no sign: it can only add. - assert (two >= coh - 1e-6).all() + assert (two >= coh - 4.0 * floor).all() + + def test_the_weight_is_the_contamination_ratio(self, two_moment_mtz): + """``W_2MOM`` must be ``sigma_I**2 / (sigma_I**2 + IVAR_ALPHA)``. + + Asserted as the formula rather than as a magnitude. On this fixture the weight + never falls below ~0.9998, because the contamination is ~1e-3 of a single + reflection's sigma -- which is the real behaviour of this correction, not a + defect: it is a systematic that adds coherently over the whole dataset while + being invisible on any one reflection. A test demanding visible down-weighting + would be asserting the physics is different from what it is. + """ + import numpy as np + + df = _read(two_moment_mtz[0]) + w = df["W_2MOM"].to_numpy().astype(float) + sig = df["SIGIo_light"].to_numpy().astype(float) + ivar = df["IVAR_ALPHA"].to_numpy().astype(float) - def test_the_weight_lies_in_zero_to_one_and_bites(self, two_moment_mtz): - w = _read(two_moment_mtz[0])["W_2MOM"].to_numpy().astype(float) assert (w > 0).all() and (w <= 1.0 + 1e-6).all() - assert w.min() < 0.99, ( - "W_2MOM is 1 everywhere, so the correction is doing nothing here and this " - "fixture cannot detect a change in it" - ) + expected = sig**2 / np.maximum(sig**2 + ivar, 1e-12) + assert np.allclose(w, expected, rtol=1e-5, atol=1e-7) + # Anti-vacuity for the formula: the contamination must not be identically zero, + # or the ratio above is trivially 1 and proves nothing. + assert (ivar > 0).any() def test_the_correction_moves_the_difference_amplitudes(self, two_moment_mtz): """DDF is the diagnostic; if it were identically zero the whole column set diff --git a/tests/unit/refinement/test_activation_weighting.py b/tests/unit/refinement/test_activation_weighting.py new file mode 100644 index 00000000..5e646bf3 --- /dev/null +++ b/tests/unit/refinement/test_activation_weighting.py @@ -0,0 +1,174 @@ +"""The difference target's per-reflection weighting must follow the activation spread. + +A merged light intensity carries a positive, phase-blind contamination +``sigma_alpha^2 |dF/dalpha|^2``. Propagated onto the amplitude it is a shift of +``sigma_alpha^2 |dF|^2 / (2 |F|)``. The difference target does not model that shift, so it +enters as a variance -- which down-weights exactly the reflections whose difference is most +contaminated, and is the calibrated form of the k-weighting difference maps apply by hand. + +Two properties are pinned: a zero dispersion changes nothing at all, and a non-zero one +reweights in proportion to ``|dF|^2`` rather than uniformly. The second is what separates a +real weighting from an overall rescaling of the x-ray term, which would be +indistinguishable from a change of x-ray weight. +""" + +import pytest +import torch + + +@pytest.fixture(scope="module") +def collection(pdb_dir, mtz_dir): + """A dark/light pair on 1DAW with a displaced light model, so dF is non-zero.""" + pdb = pdb_dir / "1DAW.pdb" + mtz = mtz_dir / "1DAW.mtz" + if not (pdb.exists() and mtz.exists()): + pytest.skip("1DAW fixture not present") + + from torchref import ReflectionData + from torchref.cli._common import load_model + from torchref.io.datasets.collection import DatasetCollection + from torchref.model.model_collection import ModelCollection + from torchref.scaling.collection_scaler import CollectionScaler + + d_min = 2.05 + dark = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + light = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + + model_dark = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) + model_light = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) + with torch.no_grad(): + model_light.xyz.refinable_params += 0.25 + + dc = DatasetCollection(verbose=0, device="cpu") + dc.add_dataset("dark", dark, set_as_reference=True) + dc.add_dataset("light", light) + + mc = ModelCollection([model_dark, model_light], dark_key="dark", verbose=0) + mc.add_dark() + mc.add_timepoint("light", [0.78, 0.22]) + + scaler = CollectionScaler(dc, mc, verbose=0) + scaler.initialize() + return dc, mc, scaler + + +def _target(dc, mc, scaler): + from torchref.refinement.targets import CollectionDifferenceTarget + + return CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0) + + +@pytest.mark.integration +class TestZeroDispersionIsInert: + def test_the_extra_variance_is_exactly_zero(self, collection): + dc, mc, scaler = collection + mc.set_lambda_twin(0.0) + target = _target(dc, mc, scaler) + + F_obs = dc.stack_F_obs(target._keys()) + extra = target._activation_variance(target._keys(), F_obs) + assert float(extra.abs().max()) == 0.0 + + def test_the_loss_is_unchanged_from_the_no_dispersion_baseline(self, collection): + """Back-compat: lambda = 0 must reproduce the loss the target always returned.""" + dc, mc, scaler = collection + mc.set_lambda_twin(0.0) + target = _target(dc, mc, scaler) + first = target.forward().item() + second = target.forward().item() + assert second == pytest.approx(first, rel=1e-4) + + +@pytest.mark.integration +class TestDispersionReweights: + def test_a_nonzero_dispersion_changes_the_loss(self, collection): + dc, mc, scaler = collection + mc.set_lambda_twin(0.0) + base = _target(dc, mc, scaler).forward().item() + + mc.set_lambda_twin(0.4) + try: + weighted = _target(dc, mc, scaler).forward().item() + finally: + mc.set_lambda_twin(0.0) + + assert weighted != pytest.approx(base, rel=1e-3) + + def test_the_extra_variance_tracks_the_squared_difference(self, collection): + """Proportional to |dF|^2, not uniform. + + A uniform inflation would just rescale the x-ray term and be + indistinguishable from a change of x-ray weight; the point of this weighting is + that it is reflection-specific. + """ + dc, mc, scaler = collection + mc.set_lambda_twin(0.4) + try: + target = _target(dc, mc, scaler) + keys = target._keys() + F_obs = dc.stack_F_obs(keys) + extra = target._activation_variance(keys, F_obs) + + rows = [mc.keys().index(k) for k in keys] + components = dc.component_structure_factors(mc, recalc=False) + jacobian = mc.activation_jacobian()[rows] + deriv = scaler.forward_batched( + mc.mix_component_fcalcs(components, jacobian), jacobian + ) + expected = ( + mc.sigma_alpha_sq * deriv.abs() ** 2 + / (2.0 * F_obs.abs().clamp(min=1e-6)) + ) ** 2 + assert torch.allclose(extra, expected, rtol=1e-5) + + light = extra[keys.index("light")] + assert float(light.max()) > 0.0 + # Genuinely non-uniform across reflections. + assert float(light.std() / light.mean().clamp(min=1e-30)) > 0.5 + finally: + mc.set_lambda_twin(0.0) + + def test_the_reference_row_is_unweighted(self, collection): + """The dark's activation Jacobian is exactly zero, so it carries no + contamination and must keep its measured sigma.""" + dc, mc, scaler = collection + mc.set_lambda_twin(0.6) + try: + target = _target(dc, mc, scaler) + keys = target._keys() + extra = target._activation_variance(keys, dc.stack_F_obs(keys)) + assert float(extra[keys.index("dark")].abs().max()) == 0.0 + assert float(extra[keys.index("light")].abs().max()) > 0.0 + finally: + mc.set_lambda_twin(0.0) + + @pytest.mark.parametrize("lam", [0.1, 0.4, 0.9]) + def test_more_dispersion_means_more_down_weighting(self, collection, lam): + """The implied weight sigma^2/(sigma^2 + extra) must fall monotonically.""" + dc, mc, scaler = collection + keys = _target(dc, mc, scaler)._keys() + F_obs = dc.stack_F_obs(keys) + + mc.set_lambda_twin(lam) + try: + extra = _target(dc, mc, scaler)._activation_variance(keys, F_obs) + finally: + mc.set_lambda_twin(0.0) + + sigma_sq = dc.stack_F_sigma(keys) ** 2 + weight = sigma_sq / (sigma_sq + extra) + light = weight[keys.index("light")] + assert float(light.max()) <= 1.0 + 1e-6 + assert float(light.min()) < 1.0, "no reflection was down-weighted at all" + + def test_the_gradient_still_reaches_the_model(self, collection): + dc, mc, scaler = collection + mc.set_lambda_twin(0.3) + try: + _target(dc, mc, scaler).forward().backward() + grad = mc.base_models[1].xyz.refinable_params.grad + assert grad is not None and torch.isfinite(grad).all() + assert float(grad.abs().max()) > 0 + finally: + mc.base_models[1].xyz.refinable_params.grad = None + mc.set_lambda_twin(0.0) diff --git a/tests/unit/refinement/test_two_moment_intensity.py b/tests/unit/refinement/test_two_moment_intensity.py index 83b96273..21111fea 100644 --- a/tests/unit/refinement/test_two_moment_intensity.py +++ b/tests/unit/refinement/test_two_moment_intensity.py @@ -357,3 +357,166 @@ def test_construction_fails_without_intensities(self, pdb_dir, mtz_dir): with pytest.raises(ValueError, match="I/SIGI"): CollectionTwoMomentIntensityTarget(dc, mc, verbose=0) + + +@pytest.mark.integration +class TestNonFiniteObservations: + """Real reflection files carry non-finite intensities, and they must not reach the + gradient. + + Masking the *loss* is not enough. ``torch.where`` picks the finite branch for the + value while still backpropagating through the branch it discarded, so one NaN + observation turns every parameter gradient into NaN and every optimizer step is + rejected -- a refinement that silently does nothing rather than one that fails. + """ + + def test_a_nan_observation_does_not_poison_the_gradient(self, collection): + dc, mc, scaler = collection + data = dc["light"] + saved = data.I.clone() + try: + with torch.no_grad(): + data.I[5] = float("nan") + data.I[11] = float("inf") + data._corrected_I_fp = None + + target = _target(dc, mc, scaler) + loss = target.forward() + assert torch.isfinite(loss), "loss went non-finite" + + loss.backward() + grad = mc.base_models[1].xyz.refinable_params.grad + assert grad is not None + assert torch.isfinite(grad).all(), ( + "non-finite observations reached the gradient; every optimizer step " + "would be rejected and the model would not move" + ) + finally: + with torch.no_grad(): + data.I.copy_(saved) + data._corrected_I_fp = None + mc.base_models[1].xyz.refinable_params.grad = None + + def test_a_nan_sigma_does_not_poison_the_gradient(self, collection): + dc, mc, scaler = collection + data = dc["light"] + saved = data.I_sigma.clone() + try: + with torch.no_grad(): + data.I_sigma[7] = float("nan") + data._corrected_I_fp = None + + target = _target(dc, mc, scaler) + loss = target.forward() + loss.backward() + grad = mc.base_models[1].xyz.refinable_params.grad + assert torch.isfinite(loss) and torch.isfinite(grad).all() + finally: + with torch.no_grad(): + data.I_sigma.copy_(saved) + data._corrected_I_fp = None + mc.base_models[1].xyz.refinable_params.grad = None + + def test_the_bad_reflections_are_excluded_not_absorbed(self, collection): + """They must drop out of the sum, not contribute a large finite penalty -- + otherwise the loss depends on how many reflections the file happened to reject. + """ + dc, mc, scaler = collection + data = dc["light"] + saved = data.I.clone() + target = _target(dc, mc, scaler) + try: + baseline = target.forward().item() + with torch.no_grad(): + data.I[3] = float("nan") + data._corrected_I_fp = None + with_nan = _target(dc, mc, scaler).forward().item() + finally: + with torch.no_grad(): + data.I.copy_(saved) + data._corrected_I_fp = None + + # One reflection out of tens of thousands: the loss should drop slightly, not + # jump by a penalty term. + assert with_nan <= baseline + assert abs(with_nan - baseline) / baseline < 1e-2 + + +@pytest.mark.integration +class TestWeightCalibration: + """Intensities are squared amplitudes, so this target's gradient is orders of + magnitude away from the amplitude target beside it. Left uncalibrated it swamps the + geometry restraints and buys R-free by moving the model further than the data + supports. + """ + + def test_the_uncalibrated_mismatch_is_large(self, collection): + """Anti-vacuity: if the two targets already pushed equally, calibration would + be pointless.""" + dc, mc, scaler = collection + from torchref.refinement.targets import CollectionDifferenceTarget + + params = [p for p in mc.base_models[1].parameters() if p.requires_grad] + diff = CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0) + target = _target(dc, mc, scaler) + + def gnorm(t): + g = torch.autograd.grad(t.forward(), params, allow_unused=True) + return sum(float((x**2).sum()) for x in g if x is not None) ** 0.5 + + ratio = gnorm(target) / gnorm(diff) + assert ratio > 10 or ratio < 0.1, ( + f"gradient ratio is {ratio:.3g}; the two targets are already matched and " + f"this fixture cannot show why calibration is needed" + ) + + def test_calibration_equalises_the_gradient_norms(self, collection): + dc, mc, scaler = collection + from torchref.refinement.targets import CollectionDifferenceTarget + + params = [p for p in mc.base_models[1].parameters() if p.requires_grad] + diff = CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0) + target = _target(dc, mc, scaler) + + target.calibrate_base_weight(diff, params) + + def gnorm(t): + g = torch.autograd.grad(t.forward(), params, allow_unused=True) + return sum(float((x**2).sum()) for x in g if x is not None) ** 0.5 + + assert gnorm(target) == pytest.approx(gnorm(diff), rel=0.05) + + def test_the_ratio_argument_scales_the_result(self, collection): + dc, mc, scaler = collection + from torchref.refinement.targets import CollectionDifferenceTarget + + params = [p for p in mc.base_models[1].parameters() if p.requires_grad] + diff = CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0) + + a = _target(dc, mc, scaler) + b = _target(dc, mc, scaler) + wa = a.calibrate_base_weight(diff, params, ratio=1.0) + wb = b.calibrate_base_weight(diff, params, ratio=0.25) + assert wb == pytest.approx(0.25 * wa, rel=1e-3) + + def test_base_weight_scales_the_loss_on_the_work_set(self, collection): + dc, mc, scaler = collection + one = _target(dc, mc, scaler, base_weight=1.0).forward().item() + three = _target(dc, mc, scaler, base_weight=3.0).forward().item() + assert three == pytest.approx(3.0 * one, rel=1e-5) + + def test_the_free_set_value_is_left_unweighted(self, collection): + """The free-set number is a diagnostic and has to stay comparable across + weightings.""" + dc, mc, scaler = collection + one = _target(dc, mc, scaler, use_set="free", base_weight=1.0).forward().item() + five = _target(dc, mc, scaler, use_set="free", base_weight=5.0).forward().item() + assert five == pytest.approx(one, rel=1e-6) + + def test_calibration_needs_refinable_parameters(self, collection): + dc, mc, scaler = collection + from torchref.refinement.targets import CollectionDifferenceTarget + + diff = CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0) + with pytest.raises(ValueError, match="No refinable parameters"): + _target(dc, mc, scaler).calibrate_base_weight(diff, []) diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index 28d1878e..1a784531 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -243,12 +243,17 @@ def setup_loss_state(dataset_collection, model_collection, scaler, if two_moment: from torchref.refinement.targets import CollectionTwoMomentIntensityTarget - state.register_target( - "xray/two_moment", - CollectionTwoMomentIntensityTarget( - dataset_collection, model_collection, scaler=scaler, - ), + two_moment_target = CollectionTwoMomentIntensityTarget( + dataset_collection, model_collection, scaler=scaler, verbose=1, + ) + # Intensities are squared amplitudes, so this target's gradient is on a + # completely different scale from the difference target beside it. Match them + # once here; left uncalibrated it swamps the geometry restraints and buys + # R-free by moving the model far further than the data supports. + two_moment_target.calibrate_base_weight( + diff_target, list(model_light.parameters()) ) + state.register_target("xray/two_moment", two_moment_target) state.set_weights(target_weights) @@ -780,7 +785,10 @@ def main(): "--lambda-twin", type=float, default=0.0, help="Activation dispersion as a fraction of its maximum, in [0, 1]: " "sigma_alpha^2 = alpha (1 - alpha) * lambda. 0 (default) is the " - "coherent model and reproduces the amplitude-only result.", + "coherent model and reproduces the amplitude-only result. On its own " + "this reweights the difference target, down-weighting the reflections " + "whose difference is most contaminated; with --two-moment it also " + "corrects the predicted intensity.", ) two_moment.add_argument( "--refine-lambda-twin", action="store_true", default=False, @@ -818,14 +826,9 @@ def main(): return 1 if args.refine_lambda_twin and not args.two_moment: print( - "Error: --refine-lambda-twin needs --two-moment; the activation " - "dispersion only enters through the two-moment intensity model", - file=sys.stderr, - ) - return 1 - if args.lambda_twin > 0.0 and not args.two_moment: - print( - "Error: --lambda-twin has no effect without --two-moment", + "Error: --refine-lambda-twin needs --two-moment. The dispersion is only " + "identifiable from the intensity likelihood; through the difference " + "target it enters as a weight and has no gradient of its own.", file=sys.stderr, ) return 1 @@ -879,9 +882,13 @@ def main(): print(f"Light data: {args.light_structure_factor}") frac_mode = "refinable" if args.refine_fractions else "frozen" print(f"Fractions: dark={fractions[0]}, light={fractions[1]} ({frac_mode})") - if args.two_moment: + if args.lambda_twin > 0.0 or args.two_moment: lam_mode = "refinable" if args.refine_lambda_twin else "fixed" - print(f"Two-moment model: lambda_twin={args.lambda_twin} ({lam_mode})") + extra = " + intensity model" if args.two_moment else " (weighting only)" + print( + f"Activation spread: lambda_twin={args.lambda_twin} " + f"({lam_mode}){extra}" + ) print(f"Output: {outdir}") print(f"Device: {device}") if args.dmin: @@ -973,9 +980,10 @@ def main(): file=sys.stderr, ) return 1 - mc.set_lambda_twin( - args.lambda_twin, refinable=args.refine_lambda_twin - ) + + # Set unconditionally: a non-zero dispersion also drives the difference target's + # per-reflection weighting, which needs no intensity data. + mc.set_lambda_twin(args.lambda_twin, refinable=args.refine_lambda_twin) state = setup_loss_state(dc, mc, scaler, target_weights, device, similarity_alpha=args.similarity_alpha, diff --git a/torchref/refinement/targets/collection/intensity.py b/torchref/refinement/targets/collection/intensity.py index 2304d7ea..574685fe 100644 --- a/torchref/refinement/targets/collection/intensity.py +++ b/torchref/refinement/targets/collection/intensity.py @@ -81,6 +81,12 @@ class CollectionTwoMomentIntensityTarget(CollectionXrayTarget): Canonical 3-way subset selector ``"work"``/``"free"``/``"val"``. verbose : int, optional Verbosity level. + base_weight : float, optional + Multiplies the summed loss on the work set. Intensities are squared amplitudes, + so this target's loss and gradient are on a completely different scale from the + amplitude targets it sits beside -- left at 1.0 it swamps them and the geometry + restraints with it. Default 1.0; use :meth:`calibrate_base_weight` to set it + from the data rather than by hand. Raises ------ @@ -98,6 +104,7 @@ def __init__( use_work_set: bool = True, use_set: str = None, verbose: int = 0, + base_weight: float = 1.0, ): super().__init__( dataset_collection, @@ -107,6 +114,7 @@ def __init__( use_set=use_set, verbose=verbose, ) + self.base_weight = float(base_weight) # Fail here rather than inside the first loss evaluation: LossState probes a # target's forward at registration, and a traceback from there is much harder to # trace back to "this MTZ had no intensity columns". @@ -198,18 +206,30 @@ def forward(self) -> torch.Tensor: sigma = dc.stack_I_sigma(keys).to(model.dtype) mask = dc.stack_masks(keys, use_set=self.use_set) + # Real reflection files carry non-finite intensities (excluded rows, and rows + # French-Wilson rejected). Masking the *loss* is not enough: a NaN observation + # makes the residual NaN, and torch.where selects the finite branch for the + # value while still backpropagating NaN through the branch it discarded. So the + # observations are sanitised into the mask BEFORE they reach the graph. + valid = torch.isfinite(obs) & torch.isfinite(sigma) + mask = mask & valid + obs = torch.where(valid, obs, torch.zeros_like(obs)) + sigma = torch.where(valid, sigma, torch.ones_like(sigma)) + sigma = self._floor_sigma(sigma, mask) - residual = obs - model + residual = torch.where(mask, obs - model, torch.zeros_like(obs)) nll = ( 0.5 * (residual / sigma) ** 2 + torch.log(sigma) + 0.5 * _LOG_2PI ) - # A single non-finite entry would poison the whole gradient; a large finite - # penalty lets the step be rejected instead. - nll = torch.where(torch.isfinite(nll), nll, torch.full_like(nll, 1e6)) - return (nll * mask).sum() + total = (nll * mask).sum() + # Applied on the work set only, matching CollectionMLTarget: the free-set value + # is a diagnostic and must stay comparable across weightings. + if self.use_work_set: + total = self.base_weight * total + return total @staticmethod def _floor_sigma(sigma: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: @@ -221,6 +241,81 @@ def _floor_sigma(sigma: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: floor = torch.clamp(floor, min=1e-12) return sigma.clamp(min=floor) + # ------------------------------------------------------------------ + # Weight calibration + # ------------------------------------------------------------------ + + def calibrate_base_weight( + self, reference, parameters, ratio: float = 1.0, floor: float = 1e-12 + ) -> float: + """Set ``base_weight`` so this target pushes as hard as ``reference``. + + Matched on the **gradient norm** with respect to the refined parameters, not on + the loss value. A large loss with a flat gradient moves nothing, so loss + magnitude is the wrong thing to equalise; what competes with the geometry and + similarity restraints is the size of the step this term asks for. + + This is the per-cycle gradient-ratio weighting the collection targets' base + weights were a stopgap for, applied once at setup rather than every cycle -- + enough to put an intensity target and an amplitude target on the same footing, + which is otherwise a several-orders-of-magnitude mismatch. + + Parameters + ---------- + reference : Target + The target to match, normally the difference target already driving the + refinement. + parameters : iterable of torch.nn.Parameter + The parameters actually being refined; only those with ``requires_grad`` + are used. + ratio : float, optional + Desired ratio of this target's gradient norm to the reference's. Default + 1.0 (equal footing); below 1 makes this target the junior partner. + floor : float, optional + Guard for a vanishing reference gradient. + + Returns + ------- + float + The ``base_weight`` that was set. + """ + params = [p for p in parameters if p.requires_grad] + if not params: + raise ValueError("No refinable parameters given; cannot calibrate.") + + def _grad_norm(target, scale_out=1.0): + grads = torch.autograd.grad( + target.forward(), params, retain_graph=False, allow_unused=True + ) + total = sum( + float((g.detach() ** 2).sum()) for g in grads if g is not None + ) + return (total**0.5) / scale_out + + saved = self.base_weight + self.base_weight = 1.0 + try: + own = _grad_norm(self) + finally: + self.base_weight = saved + + ref = _grad_norm(reference) + if own <= floor: + if self.verbose: + print( + " two-moment calibration: own gradient is ~0, leaving " + f"base_weight at {self.base_weight:.4g}" + ) + return self.base_weight + + self.base_weight = float(ratio * max(ref, floor) / own) + if self.verbose: + print( + f" two-moment weight calibration: |grad_ref|={ref:.4g}, " + f"|grad_self|={own:.4g} -> base_weight={self.base_weight:.4g}" + ) + return self.base_weight + # ------------------------------------------------------------------ # Reporting # ------------------------------------------------------------------ @@ -271,6 +366,7 @@ def stats(self) -> Dict[str, StatEntry]: lam = float(mc.lambda_twin) sigma_sq = float(mc.sigma_alpha_sq) + out["base_weight"] = stat(self.base_weight, VERBOSITY_STANDARD) out["alpha_mean"] = stat(alpha, VERBOSITY_ESSENTIAL) out["lambda_twin"] = stat(lam, VERBOSITY_ESSENTIAL) out["sigma_alpha_sq"] = stat(sigma_sq, VERBOSITY_STANDARD) diff --git a/torchref/refinement/targets/collection/xray.py b/torchref/refinement/targets/collection/xray.py index f673c962..0a7550ac 100644 --- a/torchref/refinement/targets/collection/xray.py +++ b/torchref/refinement/targets/collection/xray.py @@ -97,6 +97,53 @@ def __init__( ) self.normalize = normalize + def _activation_variance(self, keys, F_obs_stack) -> torch.Tensor: + """Extra variance from crystal-to-crystal spread in activation. + + A merged light intensity carries a positive, phase-blind contamination + ``sigma_alpha^2 |dF/dalpha|^2``. Propagated onto the amplitude it becomes a shift + of ``sigma_alpha^2 |dF|^2 / (2 |F|)``, which is a systematic of **known magnitude + but unmodelled here**, so it enters as a variance and down-weights exactly the + reflections whose difference is most contaminated. + + The resulting weight, ``sigma_meas^2 / (sigma_meas^2 + this)``, is the calibrated + form of the empirical k-weighting that difference maps apply by hand -- derived + from a fitted or assumed dispersion rather than tuned. + + Returns zeros when the dispersion is zero, so the loss is then unchanged. + + Parameters + ---------- + keys : list of str + Datasets in the order they are stacked. + F_obs_stack : torch.Tensor + Observed amplitudes, shape ``(N, n_hkl)``, used as the propagation denominator. + + Returns + ------- + torch.Tensor + Variance to add, shape ``(N, n_hkl)``; a scalar zero when inactive. + """ + mc = self._model_collection + sigma_alpha_sq = getattr(mc, "sigma_alpha_sq", None) + if sigma_alpha_sq is None: + return torch.zeros((), device=F_obs_stack.device) + if float(sigma_alpha_sq) == 0.0: + return torch.zeros((), device=F_obs_stack.device) + + dc = self._dataset_collection + rows = [mc.keys().index(k) for k in keys] + components = dc.component_structure_factors(mc, recalc=False) + jacobian = mc.activation_jacobian()[rows] + + derivative = mc.mix_component_fcalcs(components, jacobian) + if self._scaler is not None and hasattr(self._scaler, "forward_batched"): + derivative = self._scaler.forward_batched(derivative, jacobian) + + contamination = sigma_alpha_sq * derivative.abs() ** 2 + shift = contamination / (2.0 * F_obs_stack.abs().clamp(min=1e-6)) + return shift**2 + def forward(self) -> torch.Tensor: """Summed Gaussian NLL of the difference-from-mean; 0.0 if fewer than 2 sets.""" dc = self._dataset_collection @@ -144,6 +191,9 @@ def forward(self) -> torch.Tensor: # Var(F_i - F_mean) = σ_i²·(1 - 2/N) + (Σ_j σ_j²) / N² sum_sigma_sq = (sigma_stack**2).sum(dim=0) # (n_hkl,) sigma_diff_sq = sigma_stack**2 * (1 - 2.0 / N) + sum_sigma_sq / (N**2) + sigma_diff_sq = sigma_diff_sq + self._activation_variance( + all_keys, F_obs_stack + ) sigma_diff = torch.sqrt(sigma_diff_sq.clamp(min=1e-12)) # (N, n_hkl) # Mask via torch.where, not boolean indexing: no nonzero() device sync. From 0b51a6179a55054ce488efc15598702dfd42b0a5 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 13:09:14 +0200 Subject: [PATCH 061/250] Keep hydrogens by default, generating the ones a file lacks strip_H now defaults to False and a model tops up the hydrogens it is missing as it loads, so hydrogens are ordinary atoms with their own coordinates and displacement parameters and they contribute to F_calc. add_hydrogens=False keeps whatever a file carries without adding more; strip_H=True still removes everything. Generation is decided per parent rather than per file. A structure deposited with some of its hydrogens gets the rest: 1AK5 arrives with 675 of roughly 2500 and ends with 3060. A does-the-table-contain-any test would have left it as deposited. Measured on the five bundled AlphaFold starts over two macro-cycles, against the previous commit: atom count roughly doubles, median dR_work +0.0188 (worse, and positive on all five), median dR_free -0.0006 (flat, three of five improve), refinement 1.63x slower, and the test suite 494s to 1277s. R-work being consistently worse while R-free does not move is what library-placed hydrogens contributing scattering they have not been refined into looks like at this cycle count. n=5 over two cycles is indicative and no more -- the number that would settle it is the 767-structure panel through compare_to_archive.py, which has not been run. The default is one line in its own commit so it can be reverted without losing the machinery. Three bugs this surfaced, none of them in the hydrogen code: The valence cap subtracted heavy neighbours but not the hydrogens a parent already carried, so a nitrogen holding its H still had budget for the free amino acid's H2. Generation was therefore not idempotent and a save/reload added one hydrogen to every linked residue, 313 of them on 1DAW. Filtering candidates by name is not sufficient on its own -- a deposited hydrogen whose name differs from the template's would still have been over-added. vdw_radii, Z, the ITC92 coefficients and the heavy-atom mask are cached behind hasattr guards and returned as they are once built. That was safe only while every change of atom set produced a fresh model; extending the table in place left the radii at the heavy-atom count while the pair list indexed the full set. load() now drops them. Anything that replicates a decided atom table and then reshapes by its own atom count breaks under generate-on-load. EnsembleModel's two factories and _new_model_from_df all pass add_hydrogens=False, and EnsembleModel defaults it off, since the replicated single copy is the atom set. Riding hydrogens are no longer placed when the model carries real ones. They approximate the sterics of hydrogens that are absent, and with real ones present they were putting phantom atoms into the structure -- 343 on 1DAW -- that push real atoms around. Nor were they the hydrogens the generator declined: the riding builder counts bonded neighbours by distance while the generator reads them off the bond graph. They stay live under strip_H=True, which is the mode they exist for. The trajectory test now carries a measured tolerance per structure rather than sizing one from two runs at test time. 6JZA is bimodal at two macro-cycles -- its trajectories land in one of two basins about 0.0074 apart -- and two runs that pick the same basin report a spread 140 times too small. Tolerances come from five runs at regeneration time and are committed with the reference, so the other four stay held to 0.002 instead of being loosened to accommodate 6JZA. Also repoints AmberTarget at hydrogenate. The previous commit deleted Model.generate_hydrogens while a live call remained, and its tests skip without OpenMM so nothing said so. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- docs/changelog.rst | 5 + tests/functional/af_trajectory_reference.json | 191 ++++++++++-------- tests/functional/test_af_trajectory.py | 132 +++++++----- tests/integration/test_io_cif.py | 8 +- tests/integration/test_model_operations.py | 7 +- tests/unit/io/test_hkl_convention.py | 5 +- tests/unit/model/test_hydrogen_default.py | 146 +++++++++++++ tests/unit/model/test_model.py | 9 +- .../restraints/test_link_modifications.py | 10 +- tests/unit/topology/test_hydrogens.py | 13 +- .../experimental/ensemble/ensemble_model.py | 11 + torchref/experimental/targets/amber_target.py | 24 +-- torchref/model/context.py | 5 + torchref/model/model.py | 116 ++++++++++- torchref/restraints/restraints.py | 25 ++- torchref/topology/hydrogens.py | 24 ++- 16 files changed, 538 insertions(+), 193 deletions(-) create mode 100644 tests/unit/model/test_hydrogen_default.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 67186f4f..908e35af 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -31,6 +31,11 @@ Unreleased - Hydrogens on a centre whose template omits a real substituent -- a peptide-linked backbone nitrogen -- are now built from the bonded neighbours instead of the template frame - Hydrogens with a free torsion (hydroxyl, thiol, amine, methyl) are identified from bond connectivity and their dihedral is scanned, rather than taken from whatever the library deposited - Removed ``Model.generate_hydrogens``; ``Model.hydrogenate`` is the single path and no longer takes ``lbfgs_steps`` or ``max_iter`` +- Hydrogens are now present by default: ``strip_H`` defaults to False, and hydrogens a file does not carry are generated on load. New ``add_hydrogens`` argument turns generation off while still keeping any the file has +- Generation is decided per parent, so a partially hydrogenated structure is topped up rather than left as deposited +- Fixed the lazily-cached per-atom buffers (``vdw_radii``, ``Z``, the ITC92 coefficients) surviving a load that changes the atom count, which left them sized for the previous atom set +- Riding hydrogens are no longer placed when the model carries real ones, where they acted as phantom atoms in the non-bonded term +- Fixed the hydrogen valence cap counting only heavy neighbours, so a parent that already carried a hydrogen still had budget for another; generation was not idempotent and a save/reload added a spurious second amide hydrogen to every linked residue Version 0.6.4 diff --git a/tests/functional/af_trajectory_reference.json b/tests/functional/af_trajectory_reference.json index cae5d435..1cc06d09 100644 --- a/tests/functional/af_trajectory_reference.json +++ b/tests/functional/af_trajectory_reference.json @@ -1,95 +1,116 @@ { "cycles": 2, + "spread_runs": 5, "structures": { - "1VER": [ - [ - 0.372478, - 0.331636 + "1VER": { + "trajectory": [ + [ + 0.378484, + 0.332842 + ], + [ + 0.336894, + 0.317116 + ], + [ + 0.330538, + 0.308256 + ], + [ + 0.312144, + 0.305218 + ] ], - [ - 0.318247, - 0.32821 + "spread": 7.1e-05, + "tolerance": 0.002 + }, + "6VHI": { + "trajectory": [ + [ + 0.332963, + 0.344154 + ], + [ + 0.287595, + 0.332427 + ], + [ + 0.282997, + 0.309342 + ], + [ + 0.269246, + 0.324271 + ] ], - [ - 0.30407, - 0.314812 + "spread": 0.000567, + "tolerance": 0.002 + }, + "1BYW": { + "trajectory": [ + [ + 0.396823, + 0.352145 + ], + [ + 0.356685, + 0.322779 + ], + [ + 0.359504, + 0.321488 + ], + [ + 0.336396, + 0.309443 + ] ], - [ - 0.297946, - 0.313963 - ] - ], - "6VHI": [ - [ - 0.32735, - 0.33347 + "spread": 0.000107, + "tolerance": 0.002 + }, + "6JZA": { + "trajectory": [ + [ + 0.416852, + 0.378224 + ], + [ + 0.36559, + 0.35986 + ], + [ + 0.366328, + 0.363104 + ], + [ + 0.329881, + 0.325672 + ] ], - [ - 0.271641, - 0.309272 + "spread": 0.004474, + "tolerance": 0.013422 + }, + "6SXW": { + "trajectory": [ + [ + 0.430794, + 0.4249 + ], + [ + 0.383688, + 0.430362 + ], + [ + 0.387311, + 0.431517 + ], + [ + 0.340972, + 0.407762 + ] ], - [ - 0.254528, - 0.269527 - ], - [ - 0.26886, - 0.324164 - ] - ], - "1BYW": [ - [ - 0.394851, - 0.341758 - ], - [ - 0.34983, - 0.318145 - ], - [ - 0.350833, - 0.315854 - ], - [ - 0.309196, - 0.300152 - ] - ], - "6JZA": [ - [ - 0.412664, - 0.378654 - ], - [ - 0.335638, - 0.362222 - ], - [ - 0.334784, - 0.390162 - ], - [ - 0.302382, - 0.366992 - ] - ], - "6SXW": [ - [ - 0.413055, - 0.409394 - ], - [ - 0.341407, - 0.42777 - ], - [ - 0.345021, - 0.415558 - ], - [ - 0.32232, - 0.390011 - ] - ] + "spread": 9.8e-05, + "tolerance": 0.002 + } } } diff --git a/tests/functional/test_af_trajectory.py b/tests/functional/test_af_trajectory.py index f47c522f..755b3998 100644 --- a/tests/functional/test_af_trajectory.py +++ b/tests/functional/test_af_trajectory.py @@ -5,13 +5,15 @@ endpoint. That makes these five structures a sharper probe of a topology or restraint change than a converged-structure comparison would be. -The comparison carries a **null-control arm**. TorchRef refinement is not bitwise -reproducible -- even two forward passes in one process differ -- so a deviation from the -reference means nothing until it is measured against the deviation between two runs of -the *same* build. A regression is a deviation larger than that null spread, not a -deviation from zero. +TorchRef refinement is not reproducible run to run, so a deviation from the reference +means nothing on its own. Each structure therefore carries its **own** tolerance, +measured over :data:`SPREAD_RUNS` independent runs when the reference was written, and +committed alongside it. Sizing the tolerance from a couple of runs at test time does +not work: 6JZA is bimodal at this cycle count -- its trajectories land in one of two +basins about 0.0074 apart -- and two runs that happen to pick the same basin report a +spread 140 times too small. -Regenerate the reference deliberately, after a change meant to move these numbers:: +Regenerate deliberately, after a change meant to move these numbers:: ./.dev/bin/python tests/functional/test_af_trajectory.py """ @@ -19,6 +21,7 @@ import json from pathlib import Path +import numpy as np import pytest import torch @@ -30,19 +33,18 @@ #: keeping the whole test at a few seconds per structure. CYCLES = 2 -REFERENCE = Path(__file__).with_name("af_trajectory_reference.json") +#: Independent runs used to size each structure's tolerance at regeneration time. Enough +#: to see a second basin if there is one; two is not. +SPREAD_RUNS = 5 -#: Largest same-build spread tolerated before the test reports that nondeterminism -#: itself has grown. Measured spread when the reference was written: 0.0001 to 0.0014. -NULL_CEILING = 0.006 +#: Multiple of the measured spread a deviation may reach before it counts as a change. +SPREAD_MULTIPLE = 3.0 -#: Deviation floor, used when the measured null spread is very small. Keeps a structure -#: whose two runs happen to agree closely from being held to too tight a bound. -TOLERANCE_FLOOR = 0.004 +#: Tolerance floor, so a structure whose runs agree very closely is not held to an +#: unreasonably tight bound. +TOLERANCE_FLOOR = 0.002 -#: How many times the measured null spread a deviation may reach before it counts as a -#: real change rather than run-to-run noise. -NULL_MULTIPLE = 4.0 +REFERENCE = Path(__file__).with_name("af_trajectory_reference.json") def trajectory(pdb_path, mtz_path, cycles=CYCLES): @@ -79,52 +81,41 @@ def _max_deviation(a, b): @pytest.fixture(scope="module") def reference(): - """The committed reference trajectories.""" + """The committed reference trajectories and their tolerances.""" if not REFERENCE.exists(): pytest.fail( f"{REFERENCE.name} is missing. Regenerate it with " f"`./.dev/bin/python {Path(__file__).name}`." ) - return json.loads(REFERENCE.read_text()) + data = json.loads(REFERENCE.read_text()) + assert ( + data["cycles"] == CYCLES + ), f"reference was written for {data['cycles']} cycles, test runs {CYCLES}" + return data @pytest.mark.integration @pytest.mark.parametrize("code", CODES) def test_af_trajectory_matches_reference(code, reference, test_files_dir): - """The trajectory sits within the run-to-run noise of the committed reference.""" + """The trajectory sits inside this structure's own measured run-to-run spread.""" pdb_path = test_files_dir / "pdb" / f"{code}_af.pdb" mtz_path = test_files_dir / "mtz" / f"{code}.mtz" assert pdb_path.exists(), f"{pdb_path.name} is not bundled" assert mtz_path.exists(), f"{mtz_path.name} is not bundled" - assert ( - reference["cycles"] == CYCLES - ), f"reference was written for {reference['cycles']} cycles, test runs {CYCLES}" - expected = [tuple(p) for p in reference["structures"][code]] - - # The null-control arm: two runs of this build, to size the noise. - run_a = trajectory(pdb_path, mtz_path) - run_b = trajectory(pdb_path, mtz_path) - null = _max_deviation(run_a, run_b) - - assert null <= NULL_CEILING, ( - f"{code}: two runs of the same build differ by {null:.6f}, above the " - f"{NULL_CEILING} ceiling. Refinement nondeterminism has grown, which has to be " - f"understood before any comparison against the reference means anything." - ) + entry = reference["structures"][code] + expected = [tuple(point) for point in entry["trajectory"]] + tolerance = float(entry["tolerance"]) - assert len(run_a) == len( + observed = trajectory(pdb_path, mtz_path) + assert len(observed) == len( expected - ), f"{code}: trajectory has {len(run_a)} stages, reference has {len(expected)}" + ), f"{code}: trajectory has {len(observed)} stages, reference has {len(expected)}" - observed = [((a[0] + b[0]) / 2, (a[1] + b[1]) / 2) for a, b in zip(run_a, run_b)] deviation = _max_deviation(observed, expected) - tolerance = max(TOLERANCE_FLOOR, NULL_MULTIPLE * null) - assert deviation <= tolerance, ( - f"{code}: trajectory deviates from the reference by {deviation:.6f}, above the " - f"{tolerance:.6f} tolerance ({NULL_MULTIPLE}x the {null:.6f} null spread). " - f"observed={observed}\nreference={expected}" + f"{code}: deviates from the reference by {deviation:.6f}, above its measured " + f"tolerance of {tolerance:.6f}.\nobserved={observed}\nreference={expected}" ) @@ -136,30 +127,61 @@ def test_af_starts_actually_descend(reference): so the set is only useful while every member descends. """ for code in CODES: - series = reference["structures"][code] - first_rwork, last_rwork = series[0][0], series[-1][0] - assert last_rwork < first_rwork - 0.02, ( - f"{code} only moves R-work from {first_rwork:.4f} to {last_rwork:.4f}; " - f"it is too flat to serve as a trajectory probe" + series = reference["structures"][code]["trajectory"] + first, last = series[0][0], series[-1][0] + assert last < first - 0.02, ( + f"{code} only moves R-work from {first:.4f} to {last:.4f}; it is too flat " + f"to serve as a trajectory probe" + ) + + +@pytest.mark.integration +def test_tolerances_are_tight_enough_to_detect_something(reference): + """A tolerance so wide it would accept any change is not a test. + + Guards against a future regeneration quietly widening a bound until the structure + stops constraining anything. The bound has to stay well inside the descent the + trajectory itself shows. + """ + for code in CODES: + entry = reference["structures"][code] + series = entry["trajectory"] + descent = series[0][0] - series[-1][0] + assert entry["tolerance"] < descent / 4.0, ( + f"{code}: tolerance {entry['tolerance']:.4f} is not small against its " + f"own R-work descent of {descent:.4f}" ) def _write_reference(): - """Regenerate the reference from the mean of two runs per structure.""" + """Regenerate the reference, sizing each tolerance from independent runs.""" root = Path(__file__).resolve().parents[1] / "files" structures = {} for code in CODES: pdb_path = root / "pdb" / f"{code}_af.pdb" mtz_path = root / "mtz" / f"{code}.mtz" - run_a = trajectory(pdb_path, mtz_path) - run_b = trajectory(pdb_path, mtz_path) - structures[code] = [ - [round((a[0] + b[0]) / 2, 6), round((a[1] + b[1]) / 2, 6)] - for a, b in zip(run_a, run_b) + + runs = [trajectory(pdb_path, mtz_path) for _ in range(SPREAD_RUNS)] + mean = [ + tuple(float(np.mean([run[stage][i] for run in runs])) for i in (0, 1)) + for stage in range(len(runs[0])) ] - print(f"{code}: null={_max_deviation(run_a, run_b):.6f}") + spread = max(_max_deviation(run, mean) for run in runs) + tolerance = max(TOLERANCE_FLOOR, SPREAD_MULTIPLE * spread) + + structures[code] = { + "trajectory": [[round(w, 6), round(f, 6)] for w, f in mean], + "spread": round(spread, 6), + "tolerance": round(tolerance, 6), + } + print(f"{code}: spread={spread:.6f} tolerance={tolerance:.6f}", flush=True) + REFERENCE.write_text( - json.dumps({"cycles": CYCLES, "structures": structures}, indent=2) + "\n" + json.dumps( + {"cycles": CYCLES, "spread_runs": SPREAD_RUNS, "structures": structures}, + indent=2, + ) + + "\n" ) print(f"wrote {REFERENCE}") diff --git a/tests/integration/test_io_cif.py b/tests/integration/test_io_cif.py index 6f944587..ff288cc5 100644 --- a/tests/integration/test_io_cif.py +++ b/tests/integration/test_io_cif.py @@ -137,8 +137,12 @@ def test_save_and_reload_cif(self, sample_cif_file, tmp_path): assert output_path.exists() - # Reload - model2 = Model() + # add_hydrogens=False on reload: what is under test is whether the written + # file round-trips, not whether generation reruns. Regenerating on reload can + # legitimately differ, because ``write_pdb`` does not emit LINK records -- so a + # metal-coordinated nitrogen comes back with a free valence and takes a hydrogen + # it did not have before. + model2 = Model(add_hydrogens=False) model2.load_pdb(str(output_path)) n_atoms2 = model2.xyz().shape[0] diff --git a/tests/integration/test_model_operations.py b/tests/integration/test_model_operations.py index 49b3a48c..2db27886 100644 --- a/tests/integration/test_model_operations.py +++ b/tests/integration/test_model_operations.py @@ -289,7 +289,12 @@ def test_model_roundtrip_pdb(self, sample_cif_file, tmp_path): output_path = tmp_path / "output.pdb" model1.write_pdb(str(output_path)) - model2 = Model() + # add_hydrogens=False on reload: what is under test is whether the written + # file round-trips, not whether generation reruns. Regenerating on reload can + # legitimately differ, because ``write_pdb`` does not emit LINK records -- so a + # metal-coordinated nitrogen comes back with a free valence and takes a hydrogen + # it did not have before. + model2 = Model(add_hydrogens=False) model2.load_pdb(str(output_path)) n_atoms2 = model2.xyz().shape[0] diff --git a/tests/unit/io/test_hkl_convention.py b/tests/unit/io/test_hkl_convention.py index 6a08ba34..44636ac4 100644 --- a/tests/unit/io/test_hkl_convention.py +++ b/tests/unit/io/test_hkl_convention.py @@ -57,7 +57,10 @@ def anomalous_data(mtz_dir, tmp_path): def _model(pdb_dir, data): - m = ModelFT(verbose=0, max_res=2.0) + # strip_H: what is under test is the phase convention, and the absolute check + # compares against a gemmi calculation that calls ``remove_hydrogens``. Letting + # torchref generate hydrogens would have it computing a different structure. + m = ModelFT(verbose=0, max_res=2.0, strip_H=True) m.load_pdb(str(pdb_dir / f"{CODE}.pdb")) m.cell, m.spacegroup = data.cell, data.spacegroup return m diff --git a/tests/unit/model/test_hydrogen_default.py b/tests/unit/model/test_hydrogen_default.py new file mode 100644 index 00000000..678cb9a1 --- /dev/null +++ b/tests/unit/model/test_hydrogen_default.py @@ -0,0 +1,146 @@ +"""Hydrogens are present by default: kept where the file has them, generated where not. + +The interesting cases are the partially-hydrogenated file, which has to be topped up per +parent rather than left alone, and the per-atom buffers that are cached lazily and go +stale the moment the atom set grows. +""" + +import numpy as np +import pytest + +from torchref.model.model import Model + + +def _elements(model): + return model.pdb["element"].astype(str).str.strip().values + + +def _counts(model): + elements = _elements(model) + n_h = int((elements == "H").sum()) + return len(model.pdb), n_h + + +@pytest.mark.unit +def test_a_file_without_hydrogens_gets_them(pdb_dir): + """1DAW ships none, so every hydrogen here is generated.""" + model = Model(verbose=0) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + + total, n_h = _counts(model) + assert n_h > 0 + heavy = total - n_h + assert ( + 0.7 < n_h / heavy < 1.3 + ), f"{n_h} hydrogens on {heavy} heavy atoms is not a plausible ratio" + + +@pytest.mark.unit +def test_a_partially_hydrogenated_file_is_topped_up(pdb_dir): + """1AK5 ships 675 hydrogens on 2582 heavy atoms, where full is roughly 2500. + + Generation is decided per parent -- the plan proposes only a hydrogen the template + names and the model lacks -- so a file that already has some still gets the rest. A + does-the-table-contain-any test would have left this structure as deposited. + """ + kept = Model(verbose=0, add_hydrogens=False) + kept.load_pdb(str(pdb_dir / "1AK5_with_H.pdb")) + _, n_kept = _counts(kept) + + topped = Model(verbose=0) + topped.load_pdb(str(pdb_dir / "1AK5_with_H.pdb")) + _, n_topped = _counts(topped) + + assert n_kept > 0, "1AK5_with_H is supposed to ship some hydrogens" + assert ( + n_topped > n_kept * 2 + ), f"only {n_topped} hydrogens after top-up, against {n_kept} in the file" + + +@pytest.mark.unit +def test_strip_H_still_removes_everything(pdb_dir): + """The opt-out is unaffected: no hydrogen survives, generated or deposited.""" + for name in ("1DAW.pdb", "7L84.pdb"): + model = Model(verbose=0, strip_H=True) + model.load_pdb(str(pdb_dir / name)) + _, n_h = _counts(model) + assert n_h == 0, f"{name} kept {n_h} hydrogens under strip_H" + + +@pytest.mark.unit +def test_add_hydrogens_false_keeps_the_file_as_it_is(pdb_dir): + """Generation off, stripping off: exactly what the reader produced.""" + model = Model(verbose=0, add_hydrogens=False) + model.load_pdb(str(pdb_dir / "7L84.pdb")) + total, n_h = _counts(model) + assert n_h > 0, "7L84 ships hydrogens, so they should have been kept" + + generated = Model(verbose=0) + generated.load_pdb(str(pdb_dir / "7L84.pdb")) + assert _counts(generated)[0] >= total + + +@pytest.mark.unit +def test_per_atom_buffers_are_rebuilt_for_the_new_atom_set(pdb_dir): + """Every lazily-cached per-atom buffer matches the table after generation. + + These are guarded by ``hasattr`` and returned as-is once built, which was safe only + while an atom-set change always produced a fresh model. Generating hydrogens in + place left the van der Waals radii at the heavy-atom count while the pair list + indexed the full set, and the non-bonded build raised ``IndexError``. + """ + model = Model(verbose=0) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + n_atoms = len(model.pdb) + + assert model.get_vdw_radii().shape[0] == n_atoms + assert model.Z.shape[0] == n_atoms + + radii = model.get_vdw_radii().detach().cpu().numpy() + assert np.isfinite(radii).all() + is_h = _elements(model) == "H" + assert is_h.any() + assert np.allclose(radii[is_h], 1.20), "hydrogens did not get a hydrogen radius" + + +@pytest.mark.unit +def test_restraints_build_over_the_hydrogenated_model(pdb_dir): + """Restraints cover the hydrogens, and each carries exactly one bond.""" + model = Model(verbose=0) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + restraints = model.restraints + + elements = _elements(model) + is_h = elements == "H" + bonds = restraints.restraints["bond"]["all"]["indices"].cpu().numpy() + involves_h = is_h[bonds[:, 0]] | is_h[bonds[:, 1]] + assert int(involves_h.sum()) == int(is_h.sum()) + + vdw = restraints.restraints["vdw"]["indices"] + assert int(vdw.max()) < len( + model.pdb + ), "the non-bonded pair list indexes past the end of the atom table" + + +@pytest.mark.unit +def test_riding_hydrogens_are_not_placed_when_real_ones_exist(pdb_dir): + """The riding stand-in goes quiet once the model carries hydrogens. + + Riding hydrogens approximate the sterics of hydrogens the model does not have. + Placing them alongside real ones would put phantom atoms in the structure that push + real ones around -- and they would not even be the hydrogens the generator declined, + because the riding builder counts bonded neighbours by distance while the generator + reads them off the bond graph. + """ + model = Model(verbose=0) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + restraints = model.restraints + + assert restraints.h_topo is not None + assert restraints.h_topo.n_hydrogens == 0 + + stripped = Model(verbose=0, strip_H=True) + stripped.load_pdb(str(pdb_dir / "1DAW.pdb")) + assert ( + stripped.restraints.h_topo.n_hydrogens > 0 + ), "with hydrogens absent the riding stand-in should still be built" diff --git a/tests/unit/model/test_model.py b/tests/unit/model/test_model.py index e1b99867..c547ac97 100644 --- a/tests/unit/model/test_model.py +++ b/tests/unit/model/test_model.py @@ -54,12 +54,13 @@ def test_model_custom_dtype(self): @pytest.mark.unit def test_model_strip_h_default(self): - """Test strip_H defaults to True.""" + """strip_H defaults to False, and hydrogen generation is on.""" from torchref.model.model import Model - + model = Model() - - assert model.ctx.strip_H is True + + assert model.ctx.strip_H is False + assert model.ctx.add_hydrogens is True @pytest.mark.unit def test_model_bool_uninitialized(self): diff --git a/tests/unit/restraints/test_link_modifications.py b/tests/unit/restraints/test_link_modifications.py index 11b29228..67345a21 100644 --- a/tests/unit/restraints/test_link_modifications.py +++ b/tests/unit/restraints/test_link_modifications.py @@ -224,10 +224,16 @@ def test_peptide_modifications_carry_the_linked_backbone_targets(): def _built(pdb_path, strip_H=True): - """Build a model's restraints and return ``(model, table accessor)``.""" + """Build a model's restraints and return ``(model, table accessor)``. + + ``add_hydrogens=False``: these tests read the restraint targets of the hydrogens the + file carries. Generating more would add a chain-terminal ``CA-N-H``, which correctly + keeps the free-amino-acid target of 109.6 degrees rather than the linked 118.7 and so + is outside what they assert. + """ from torchref import Model - model = Model(verbose=0, strip_H=strip_H) + model = Model(verbose=0, strip_H=strip_H, add_hydrogens=False) model.load_pdb(str(pdb_path)) return model, model.restraints.restraints diff --git a/tests/unit/topology/test_hydrogens.py b/tests/unit/topology/test_hydrogens.py index 755e6e2c..0b83d9cc 100644 --- a/tests/unit/topology/test_hydrogens.py +++ b/tests/unit/topology/test_hydrogens.py @@ -28,7 +28,9 @@ def built(pdb_dir): def _build(code): if code not in cache: - model = Model(verbose=0) + # add_hydrogens=False: these tests exercise generation itself, so the model + # has to arrive without the hydrogens the loader would otherwise add. + model = Model(verbose=0, add_hydrogens=False, strip_H=True) model.load_pdb(str(pdb_dir / f"{code}.pdb")) model.set_restraints_cif(None) restraints = model.restraints @@ -243,7 +245,7 @@ def test_waters_are_not_hydrogenated(built): @pytest.mark.unit def test_hydrogenate_returns_a_consistent_model(pdb_dir): """The end-to-end path yields a model whose tensors, table and restraints agree.""" - model = Model(verbose=0) + model = Model(verbose=0, add_hydrogens=False, strip_H=True) model.load_pdb(str(pdb_dir / "7L84.pdb")) model.set_restraints_cif(None) n_heavy = len(model.pdb) @@ -275,10 +277,9 @@ def test_hydrogenate_returns_a_consistent_model(pdb_dir): @pytest.mark.unit -def test_loading_still_strips_hydrogens_by_default(pdb_dir): - """``strip_H`` is untouched, so refinement sees the same atoms as before.""" - model = Model(verbose=0) +def test_strip_H_removes_deposited_hydrogens(pdb_dir): + """The opt-out drops the hydrogens the file carries, as it always did.""" + model = Model(verbose=0, strip_H=True) model.load_pdb(str(pdb_dir / "1AK5_with_H.pdb")) - assert model.ctx.strip_H is True elements = model.pdb["element"].astype(str).str.strip().values assert not (elements == "H").any() diff --git a/torchref/experimental/ensemble/ensemble_model.py b/torchref/experimental/ensemble/ensemble_model.py index 07896f9e..10eb5ecc 100644 --- a/torchref/experimental/ensemble/ensemble_model.py +++ b/torchref/experimental/ensemble/ensemble_model.py @@ -340,6 +340,10 @@ def __init__( verbose: int = 1, device=None, strip_H: bool = True, + # An ensemble's atom set is the replicated single copy its factories build, and + # _finalize_ensemble reshapes by n_atoms_per_member, so generating hydrogens on + # load would invalidate that. Off by default here, unlike on the base class. + add_hydrogens: bool = False, max_res: float = 1.0, gridsize: Optional[int] = None, wavelength: float = 1.0, @@ -354,6 +358,7 @@ def __init__( verbose=verbose, device=device, strip_H=strip_H, + add_hydrogens=add_hydrogens, max_res=max_res, gridsize=gridsize, wavelength=wavelength, @@ -452,6 +457,9 @@ def from_single( model = cls( verbose=verbose, device=device, strip_H=False, # already stripped + # The replicated table is the atom set; _finalize_ensemble reshapes by + # n_atoms_per_member, so generating hydrogens here would invalidate it. + add_hydrogens=False, max_res=max_res, **modelft_kwargs, ) @@ -538,6 +546,9 @@ def from_multimodel_pdb( model = cls( verbose=verbose, device=device, strip_H=False, + # See from_single: the replicated table is the atom set, and + # _finalize_ensemble reshapes by n_atoms_per_member. + add_hydrogens=False, max_res=max_res, **modelft_kwargs, ) diff --git a/torchref/experimental/targets/amber_target.py b/torchref/experimental/targets/amber_target.py index 414616da..bff215bc 100644 --- a/torchref/experimental/targets/amber_target.py +++ b/torchref/experimental/targets/amber_target.py @@ -15,7 +15,7 @@ mh = (Model(verbose=0, strip_H=True) .load_pdb('structure.pdb') .strip_altlocs() - .generate_hydrogens()) + .hydrogenate()) target = AmberTarget(model=mh) # protein-only target = AmberTarget(model=mh, residue_charges={'LIG': -1}) # with ligand @@ -351,16 +351,16 @@ class AmberTarget(ModelTarget): tleap and are NOT included in the atom map or gradient. Passing a model that already has H atoms (via - ``model.generate_hydrogens()`` or loading a PDB with H) speeds up + ``model.hydrogenate()`` or loading a PDB with H) speeds up initialisation because ``Modeller.addHydrogens()`` converges faster from existing positions. **GAFF2 ligands**: antechamber's BCC charge scheme runs a semiempirical QM step (sqm) that needs a fully protonated molecule. Heavy-only ligands are auto-protonated from the monomer library - (``generate_hydrogens``) first; an error is raised only if no + (``hydrogenate``) first; an error is raised only if no monomer CIF resolves AND the heavy-atom electron count is odd. - Calling ``model.generate_hydrogens()`` or loading the PDB with + Calling ``model.hydrogenate()`` or loading the PDB with ``strip_H=False`` beforehand avoids relying on that fallback. cutoff : float Non-bonded cutoff in Angstroms. Default 5.0. @@ -452,7 +452,7 @@ def __init__( self._tleap_residue_map: Optional[List[Dict[str, int]]] = None # Cached protonated chemistry PDB (filled lazily by the first ligand # parameterisation that needs H). None = not yet computed; False = - # generate_hydrogens failed (don't retry). + # hydrogenate failed (don't retry). self._protonated_pdb_cache = None if self._chem_model is None: @@ -597,21 +597,21 @@ def _write_residue_pdb(self, res_atoms, path: Path) -> None: def _protonated_chem_pdb(self): """Protonated chemistry-model PDB DataFrame (cached), or ``None``. - Uses :meth:`Model.generate_hydrogens` once on the whole chemistry model - (which has a unit cell + full residue context, so gemmi's topology engine - is well-posed). H come from the monomer-library CIF at ideal geometry via - TorchRef's auto-fetching monomer library — no full CCP4 install needed. - Cached so repeated ligand parameterisations don't re-run it. + Uses :meth:`Model.hydrogenate` once on the whole chemistry model, which has a + unit cell and full residue context so every centre has neighbours to orient its + template against. H come from the monomer-library CIF at ideal geometry via + TorchRef's auto-fetching monomer library -- no full CCP4 install needed. Cached + so repeated ligand parameterisations don't re-run it. """ if self._protonated_pdb_cache is None: try: - m_h = self._chem_model.generate_hydrogens() + m_h = self._chem_model.hydrogenate() self._protonated_pdb_cache = ( m_h.update_pdb() if hasattr(m_h, "update_pdb") else m_h.pdb ) except Exception as exc: # missing CIF/lib, gemmi failure, etc. if self.verbose >= 1: - print(f"[AmberTarget] generate_hydrogens failed: {exc}") + print(f"[AmberTarget] hydrogenate failed: {exc}") self._protonated_pdb_cache = False if self._protonated_pdb_cache is False: return None diff --git a/torchref/model/context.py b/torchref/model/context.py index 85c47ba8..5c449019 100644 --- a/torchref/model/context.py +++ b/torchref/model/context.py @@ -53,6 +53,9 @@ class ModelContext(DeviceMixin): Whether hydrogens were stripped on load. exclude_H_from_sf : bool, default False Whether hydrogens are excluded from structure-factor calculation. + add_hydrogens : bool, default True + Generate hydrogens on load for residues that arrive without them. Ignored when + ``strip_H`` is set, which removes them again. initialized : bool, default False Whether a structure has been loaded. ``if model:`` tests this. @@ -77,6 +80,7 @@ class ModelContext(DeviceMixin): verbose: int = 1 strip_H: bool = True exclude_H_from_sf: bool = False + add_hydrogens: bool = True initialized: bool = False def copy(self) -> "ModelContext": @@ -107,6 +111,7 @@ def copy(self) -> "ModelContext": verbose=self.verbose, strip_H=self.strip_H, exclude_H_from_sf=self.exclude_H_from_sf, + add_hydrogens=self.add_hydrogens, initialized=self.initialized, ) diff --git a/torchref/model/model.py b/torchref/model/model.py index 592092f1..fd774801 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -86,7 +86,11 @@ class Model(DeviceMovementMixin, DebugMixin, nn.Module): device : torch.device, optional Computation device. Defaults to the configured device.current. strip_H : bool, optional - Whether to strip hydrogen atoms when loading. Default is True. + Whether to strip hydrogen atoms when loading. Default False: hydrogens are kept + where the file has them and generated where it does not. + add_hydrogens : bool, optional + Generate hydrogens on load for residues that arrive without them. Default True; + ignored when ``strip_H`` is set. Attributes ---------- @@ -120,7 +124,8 @@ def __init__( dtype_float=None, verbose=1, device=None, - strip_H: bool = True, + strip_H: bool = False, + add_hydrogens: bool = True, ): """ Initialize an empty Model shell. @@ -137,7 +142,11 @@ def __init__( device : torch.device, optional Computation device. Defaults to the configured device.current. strip_H : bool, optional - Whether to strip hydrogen atoms when loading. Default is True. + Whether to strip hydrogen atoms when loading. Default False: hydrogens are + kept where the file has them and generated where it does not. + add_hydrogens : bool, optional + Generate hydrogens on load for residues that arrive without them. Default + True; ignored when ``strip_H`` is set. """ super().__init__() # Resolve dtype/device at call time (not import time) so a runtime @@ -153,7 +162,9 @@ def __init__( # Everything the model is loaded from and sits in, as opposed to what is # refined. Populated by load() / create_from_state_dict(). - self.ctx = ModelContext(verbose=verbose, strip_H=strip_H) + self.ctx = ModelContext( + verbose=verbose, strip_H=strip_H, add_hydrogens=add_hydrogens + ) # Submodules (created during load or load_state_dict) self.xyz = None @@ -512,7 +523,31 @@ def torsion_deviations_with_sigmas(self): """ return self.restraints.torsion_deviations_with_sigmas(self.xyz()) - def load(self, reader): + #: Per-atom buffers built lazily on first use and cached. Each is sized to the atom + #: table, so all of them go stale the moment the atom set changes. + _ATOM_DERIVED_BUFFERS = ( + "vdw_radii", + "_Z", + "_A", + "_B", + "_heavy_atom_mask", + ) + + def _invalidate_atom_derived_caches(self) -> None: + """Drop the lazily-cached per-atom buffers. + + Each is guarded by ``hasattr`` and returned as-is once built, so a load that + changes the atom count would otherwise hand back a buffer sized for the previous + one. That surfaced when hydrogen generation began extending the table in place: + the van der Waals radii stayed at the heavy-atom count while the pair list + indexed the full set, and the non-bonded build raised ``IndexError``. Rebuilding + a new model each time had hidden it. + """ + for name in self._ATOM_DERIVED_BUFFERS: + if hasattr(self, name): + delattr(self, name) + + def load(self, reader, add_hydrogens: bool = None): """ Populate the model from a reader callable. @@ -529,6 +564,10 @@ def load(self, reader): reader : callable Zero-argument callable returning ``(pdb_df, cell, spacegroup)``. An optional ``.links`` attribute on it is stored on ``self.ctx.links``. + add_hydrogens : bool, optional + Whether to top up missing hydrogens once the model is built. Defaults to the + context's setting, and is forced off for the re-entry that + :meth:`_add_missing_hydrogens` makes, so generation happens once per load. Returns ------- @@ -541,6 +580,9 @@ def load(self, reader): ``aniso_flag`` buffer, the four wrappers, the default masks, the altloc registration and ``initialized = True``. """ + if add_hydrogens is None: + add_hydrogens = self.ctx.add_hydrogens and not self.ctx.strip_H + self._invalidate_atom_derived_caches() self.pdb, cell, spacegroup = reader() self.ctx.links = getattr(reader, "links", None) @@ -605,8 +647,60 @@ def load(self, reader): self.set_default_masks() self.register_alternative_conformations() self.ctx.initialized = True + + if add_hydrogens: + self._add_missing_hydrogens() return self + def _add_missing_hydrogens(self) -> None: + """Top up the hydrogens the atom table is missing, in place. + + Per parent, not per file: a structure deposited with some hydrogens gets the + rest, because the plan only ever proposes a hydrogen the template names and the + model does not have. 1AK5 arrives with 675 of roughly 2500, and a + does-it-have-any test would have left it there. + + Re-enters :meth:`load` on the augmented atom table, which rebuilds the parameter + wrappers and per-atom buffers at the new size. The re-entry is told not to + consider hydrogens again, so this runs once per load rather than recursing to a + fixed point. + + Costs a restraint build that is then discarded, because the plan needs the + topology and the topology is built over the atoms as loaded. Set + ``add_hydrogens=False`` to skip it for a model that will never be refined. + """ + from torchref.topology.hydrogens import ( + augment_atom_table, + optimise_free_torsions, + plan_hydrogens, + ) + + restraints = self.restraints + xyz = self.xyz().detach() + plan = plan_hydrogens( + restraints.topology, restraints.cif_dict, xyz, verbose=self.ctx.verbose + ) + if plan.n_hydrogens == 0: + return + optimise_free_torsions(plan, restraints.topology, xyz) + augmented = augment_atom_table(self.pdb, plan, restraints.topology) + + if self.ctx.verbose > 0: + print(f"Generated {plan.n_hydrogens} hydrogens") + + # The topology and every per-atom tensor are sized for the old atom set. + self._restraints = None + cell, spacegroup = self.cell, self.spacegroup + links = self.ctx.links + + def reader(): + return augmented, cell.data.cpu().numpy(), spacegroup + + # Carried explicitly: ``load`` reads links off the reader, so a bare callable + # would drop the LINK records the first read resolved. + reader.links = links + self.load(reader, add_hydrogens=False) + def load_pdb(self, file): """ Load atomic model from PDB file. @@ -820,7 +914,7 @@ def update_pdb(self): Copies the live values of ``xyz`` (x/y/z), ``u`` (u11..u23), ``adp`` (tempfactor), and ``occupancy`` from the parameter wrappers into the corresponding columns of the ``self.pdb`` DataFrame. Called by every - writer and by ``hydrogenate`` / ``generate_hydrogens`` before output. + writer and by ``hydrogenate`` before output. Returns ------- @@ -1588,8 +1682,13 @@ def shake_adp(self, stddev: float): ) - def _new_model_from_df(self, df, *, strip_H=None): - """Build a fresh model of the same class from a DataFrame.""" + def _new_model_from_df(self, df, *, strip_H=None, add_hydrogens=False): + """Build a fresh model of the same class from a DataFrame. + + ``add_hydrogens`` defaults to False, unlike the constructor: the caller has + already settled which atoms the table holds, and generating more would fight + that. :meth:`hydrogenate` passes an already-augmented table for the same reason. + """ import inspect sh = self.ctx.strip_H if strip_H is None else strip_H @@ -1598,6 +1697,7 @@ def _new_model_from_df(self, df, *, strip_H=None): verbose=0, device=self.device, strip_H=sh, + add_hydrogens=add_hydrogens, ) sig = inspect.signature(self.__class__.__init__) for pname, param in sig.parameters.items(): diff --git a/torchref/restraints/restraints.py b/torchref/restraints/restraints.py index ed994f0c..53c8d519 100644 --- a/torchref/restraints/restraints.py +++ b/torchref/restraints/restraints.py @@ -838,17 +838,28 @@ def vdw_radii_cpu(): # and re-inserted here and by _rebuild_entries. self._entries["vdw"] = self._vdw - # Build riding hydrogen topology and precompute candidate pairs + # Riding hydrogens stand in for the sterics of hydrogens the model does not + # carry. Once it carries them they are ordinary atoms in the pair list above, and + # placing riding ones as well would put phantom hydrogens in the structure that + # push real atoms around. The two also disagree about how many belong on a + # parent -- the riding builder counts bonded neighbours by distance, the + # generator reads them off the bond graph -- so the leftovers are not even the + # hydrogens the generator declined to add. from torchref.restraints.hydrogen_topology import ( - build_hydrogen_topology, + HydrogenTopology, build_h_candidate_pairs, + build_hydrogen_topology, ) - self._h_topo = build_hydrogen_topology( - pdb=self.pdb, - device=cpu, - verbose=self.verbose, - ) + elements = self.pdb["element"].astype(str).str.strip().values + if (elements == "H").any(): + self._h_topo = HydrogenTopology(device=cpu) + else: + self._h_topo = build_hydrogen_topology( + pdb=self.pdb, + device=cpu, + verbose=self.verbose, + ) self._h_excl_hash = self._build_h_exclusion_hash(self._h_topo, cpu) # Precompute H candidate pairs from heavy-atom VDW pair list diff --git a/torchref/topology/hydrogens.py b/torchref/topology/hydrogens.py index 5ede3a54..b8f0a84a 100644 --- a/torchref/topology/hydrogens.py +++ b/torchref/topology/hydrogens.py @@ -254,17 +254,23 @@ def _half_hydrogen_angle(template: Dict, parent_name: str, h_names: List[str]) - return 0.5 * np.arccos(-1.0 / 3.0) -def _heavy_neighbours(topology, atom_index: int) -> np.ndarray: - """Rows of the heavy atoms bonded to ``atom_index``, from the bond graph. +def _split_neighbours(topology, atom_index: int) -> Tuple[np.ndarray, int]: + """Heavy neighbour rows of ``atom_index``, and how many hydrogens it already has. Coordinate-independent, unlike a distance sweep: a stretched bond in a predicted or mid-refinement model still counts, and two atoms that merely sit close do not. + + The hydrogen count is what makes generation idempotent and makes a partially + hydrogenated structure top up correctly. Both consume the parent's valence, so + subtracting only the heavy neighbours leaves budget for a hydrogen the parent + already carries -- which is how a second pass came to add the free-amino-acid ``H2`` + to every backbone nitrogen that already had its ``H``. """ neighbours = topology.atoms.neighbors(atom_index) if neighbours.numel() == 0: - return np.zeros(0, dtype=np.int64) - heavy = neighbours[~topology.atoms.is_hydrogen[neighbours]] - return heavy.cpu().numpy() + return np.zeros(0, dtype=np.int64), 0 + is_h = topology.atoms.is_hydrogen[neighbours] + return neighbours[~is_h].cpu().numpy(), int(is_h.sum()) def _template_bond_length(template: Dict, parent_name: str, h_name: str) -> float: @@ -439,15 +445,13 @@ def plan_hydrogens(topology, cif_dict: Dict, xyz, verbose: int = 0) -> HydrogenP parent_row = name_to_row[parent_name] parent_position = coords[parent_row] - heavy_rows = _heavy_neighbours(topology, parent_row) + heavy_rows, existing_h = _split_neighbours(topology, parent_row) heavy_bonded = len(heavy_rows) element = str( template["elements"][template["id_to_index"][parent_name]] ).upper() - allowed = max( - 0, - STANDARD_VALENCE.get(element, _DEFAULT_VALENCE) - heavy_bonded, - ) + valence = STANDARD_VALENCE.get(element, _DEFAULT_VALENCE) + allowed = max(0, valence - heavy_bonded - existing_h) group = group[:allowed] if not group: continue From 08d62c80d76f7e5419684ade2354e4f677c6ae84 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 13:36:43 +0200 Subject: [PATCH 062/250] Move the riding-hydrogen map into the topology package torchref.restraints.hydrogen_topology becomes torchref.topology.riding, next to the generation path it is the alternative to. Both answer the same question over the same graph, so they belong together: hydrogens.py adds hydrogens as real atoms with their own parameters, which is the default, and riding.py reconstructs absent ones from their parents at each non-bonded evaluation, which is what a model loaded heavy-only gets. The module docstrings now say which applies when, and that they must not run together. Kept rather than deleted. Riding placement is the only thing that gives a strip_H=True model any hydrogen sterics at all, and that is a mode worth having. Two pieces stay where they are because that is where their kind lives: non_bonded_h.py is a refinement target, and place_hydrogens.py is a Triton kernel. Both now import from the new location. Suite unchanged at 1918 passed, which is what a move should do. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- docs/changelog.rst | 1 + tests/helpers/device_cases.py | 2 +- .../base/targets/triton/place_hydrogens.py | 2 +- .../targets/geometry/non_bonded_h.py | 8 +++--- torchref/restraints/__init__.py | 8 +++--- torchref/restraints/restraints.py | 2 +- torchref/topology/__init__.py | 19 +++++++++++-- .../riding.py} | 27 +++++++++++-------- 8 files changed, 45 insertions(+), 24 deletions(-) rename torchref/{restraints/hydrogen_topology.py => topology/riding.py} (97%) diff --git a/docs/changelog.rst b/docs/changelog.rst index 908e35af..690bae20 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -36,6 +36,7 @@ Unreleased - Fixed the lazily-cached per-atom buffers (``vdw_radii``, ``Z``, the ITC92 coefficients) surviving a load that changes the atom count, which left them sized for the previous atom set - Riding hydrogens are no longer placed when the model carries real ones, where they acted as phantom atoms in the non-bonded term - Fixed the hydrogen valence cap counting only heavy neighbours, so a parent that already carried a hydrogen still had budget for another; generation was not idempotent and a save/reload added a spurious second amide hydrogen to every linked residue +- Moved the riding-hydrogen map from ``torchref.restraints.hydrogen_topology`` to ``torchref.topology.riding``, alongside the generation path it is the heavy-atom-only alternative to Version 0.6.4 diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index ee773c5b..32198d04 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -271,7 +271,7 @@ def _cell(device): DeviceCase( "HydrogenTopology_empty", lambda d: __import__( - "torchref.restraints.hydrogen_topology", + "torchref.topology.riding", fromlist=["HydrogenTopology"], ).HydrogenTopology(device=d), "HydrogenTopology", diff --git a/torchref/base/targets/triton/place_hydrogens.py b/torchref/base/targets/triton/place_hydrogens.py index 31f74461..f88c13f0 100644 --- a/torchref/base/targets/triton/place_hydrogens.py +++ b/torchref/base/targets/triton/place_hydrogens.py @@ -1,7 +1,7 @@ """Triton forward + analytic backward for riding-hydrogen placement. One launch each way, where the eager helper (``_place_h_jit`` in -:mod:`torchref.restraints.hydrogen_topology`) fuses only the forward and leaves its +:mod:`torchref.topology.riding`) fuses only the forward and leaves its backward to run op-by-op through autograd -- ~100 launches at 3k hydrogens, which dominates the non-bonded backward. The math mirrors ``_place_h_jit`` exactly: diff --git a/torchref/refinement/targets/geometry/non_bonded_h.py b/torchref/refinement/targets/geometry/non_bonded_h.py index e3cfae41..8a249b24 100644 --- a/torchref/refinement/targets/geometry/non_bonded_h.py +++ b/torchref/refinement/targets/geometry/non_bonded_h.py @@ -20,7 +20,7 @@ if TYPE_CHECKING: from torchref.model.model import Model - from torchref.restraints.hydrogen_topology import HydrogenTopology + from torchref.topology.riding import HydrogenTopology class NonBondedHTarget(NonBondedTarget): @@ -93,7 +93,7 @@ def _compute_h_vdw_loss( ordering to do identity and real symmetry transforms in one pass. Other modes take the inline eager path below. """ - from torchref.restraints.hydrogen_topology import place_riding_hydrogens + from torchref.topology.riding import place_riding_hydrogens device = xyz.device @@ -200,7 +200,7 @@ def get_violations(self, threshold: float = 0.0) -> Dict[str, torch.Tensor]: ``xyz_all`` -- so symmetry-mate H contacts are reported at their intra-ASU separation, unlike in the loss. """ - from torchref.restraints.hydrogen_topology import place_riding_hydrogens + from torchref.topology.riding import place_riding_hydrogens result = super().get_violations(threshold) @@ -234,7 +234,7 @@ def get_violations(self, threshold: float = 0.0) -> Dict[str, torch.Tensor]: def stats(self) -> Dict[str, any]: """Get statistics including H-VDW contacts.""" - from torchref.restraints.hydrogen_topology import place_riding_hydrogens + from torchref.topology.riding import place_riding_hydrogens result = super().stats() diff --git a/torchref/restraints/__init__.py b/torchref/restraints/__init__.py index c9a1c1c7..55cb62eb 100644 --- a/torchref/restraints/__init__.py +++ b/torchref/restraints/__init__.py @@ -5,10 +5,10 @@ resolves lazily to that manager's ``monomer_dir``. Ideal values come from the CCP4 Monomer Library (Long et al. 2017, Acta Cryst. D73, 112-122). -The builder classes and topology helpers (``build_all_restraints``, -``HydrogenTopology``, ``build_hydrogen_topology``, the intra-/inter-residue -builders) are used across the package but deliberately *not* re-exported here -- -import them from their defining submodules. +The builder classes and the inter-residue builders are used across the package but +deliberately *not* re-exported here -- import them from their defining submodules. +Connectivity itself lives in :mod:`torchref.topology`, which is also where the +riding-hydrogen map moved to. """ from torchref.restraints.library import get_library_manager diff --git a/torchref/restraints/restraints.py b/torchref/restraints/restraints.py index 53c8d519..a45fe074 100644 --- a/torchref/restraints/restraints.py +++ b/torchref/restraints/restraints.py @@ -845,7 +845,7 @@ def vdw_radii_cpu(): # parent -- the riding builder counts bonded neighbours by distance, the # generator reads them off the bond graph -- so the leftovers are not even the # hydrogens the generator declined to add. - from torchref.restraints.hydrogen_topology import ( + from torchref.topology.riding import ( HydrogenTopology, build_h_candidate_pairs, build_hydrogen_topology, diff --git a/torchref/topology/__init__.py b/torchref/topology/__init__.py index cc025aa3..b9b9abd3 100644 --- a/torchref/topology/__init__.py +++ b/torchref/topology/__init__.py @@ -10,8 +10,13 @@ connectivity can carry monomer-library targets, force-field parameters, or ADP-similarity sigmas without duplicating the edges. -Build one with :func:`build_topology`. :func:`plan_hydrogens` uses it to instantiate -monomer templates, which is how hydrogens are generated. +Build one with :func:`build_topology`. + +Hydrogens come in two forms over the same graph. :func:`plan_hydrogens` instantiates +the monomer templates to add them as real atoms, which is the default. For a model +loaded heavy-only, :mod:`torchref.topology.riding` reconstructs them from their parents +at each non-bonded evaluation instead, so their sterics still count. Only one applies at +a time. """ from .atom_graph import AtomGraph @@ -24,6 +29,12 @@ plan_hydrogens, ) from .residue_graph import ResidueGraph +from .riding import ( + HydrogenTopology, + build_h_candidate_pairs, + build_hydrogen_topology, + place_riding_hydrogens, +) from .restraint_sets import assemble_entries, max_period from .templates import resolve_template_keys from .topology import Topology @@ -42,5 +53,9 @@ "plan_hydrogens", "optimise_free_torsions", "augment_atom_table", + "HydrogenTopology", + "build_hydrogen_topology", + "build_h_candidate_pairs", + "place_riding_hydrogens", "resolve_template_keys", ] diff --git a/torchref/restraints/hydrogen_topology.py b/torchref/topology/riding.py similarity index 97% rename from torchref/restraints/hydrogen_topology.py rename to torchref/topology/riding.py index f5276621..46271c2a 100644 --- a/torchref/restraints/hydrogen_topology.py +++ b/torchref/topology/riding.py @@ -1,14 +1,19 @@ -""" -Riding hydrogen topology and vectorized placement for VDW restraints. - -Builds a static topology map at restraints-construction time that describes -how to generate transient hydrogen atom positions from heavy-atom coordinates. -At each VDW evaluation the ``place_riding_hydrogens`` function produces H -positions in a single vectorized pass (no Python loops over atoms). - -Hydrogen positions are fully determined by the parent heavy atom and its -bonded heavy-atom neighbours, so gradients flow from the VDW loss through -the H positions back to the heavy-atom coordinates via standard autograd. +"""Riding hydrogens: the sterics of hydrogens a model does not carry. + +For a model loaded with ``strip_H=True``, whose atoms are heavy only. A static map +built once at restraint-construction time says how to reconstruct each absent hydrogen +from its parent and the parent's bonded neighbours; ``place_riding_hydrogens`` then +produces those positions in one vectorized pass at every non-bonded evaluation and +throws them away again. The positions are a function of the heavy atoms, so gradients +reach the heavy coordinates through them by ordinary autograd. + +Contrast :mod:`torchref.topology.hydrogens`, which *adds* hydrogens to the model as +real atoms with their own parameters. That is the default, and where both apply it is +the better answer: the hydrogen has a refinable position instead of one reconstructed +each step, and it contributes to the structure factors. Riding hydrogens are what is +left for the heavy-atom-only mode, and the two must not run together -- riding +placement alongside real hydrogens puts phantom atoms in the structure that push the +real ones around. """ from dataclasses import dataclass, field From 134e5f4f09e5353beeb1ab0d6ba99c4264039e7c Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 13:44:30 +0200 Subject: [PATCH 063/250] Add probes for activation-dispersion identifiability and movement recovery Two questions the figure-4 arms cannot answer, each with its own probe, following the paper/probe_*.py convention of a standalone script against a documented hazard. probe_two_moment_collinearity: the contamination is smooth, positive and concentrated at low resolution, which is also what a scale or overall-B error looks like. Projecting the whitened template onto the scale parameters' derivatives answers by linear algebra, with no refinement, how much of it the scale model could absorb. The two parameter sets need different response functions: the shared scaler shapes the model intensity while the per-dataset log_scale/U_aniso shape the observed one, so using the model response for both gives the dataset parameters a zero column and reports no leak at all. probe_movement_recovery: nothing measured on real data says whether a difference refinement recovers the true displacement, because the true light state is never known. So it injects one -- displace a stretch of the model, build the merged intensity the two-moment physics predicts, add noise, and refine starting from the dark model so the displacement has to be discovered. On 1DAW at figure-4-like contamination the coherent refinement recovers 0.398 A of an injected 0.430 A: it falls short by 8%, dominated by restraint shrinkage, and the corrections change it by under 1%. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- paper/probe_movement_recovery.py | 233 +++++++++++++++++++++++++ paper/probe_two_moment_collinearity.py | 220 +++++++++++++++++++++++ 2 files changed, 453 insertions(+) create mode 100644 paper/probe_movement_recovery.py create mode 100644 paper/probe_two_moment_collinearity.py diff --git a/paper/probe_movement_recovery.py b/paper/probe_movement_recovery.py new file mode 100644 index 00000000..b5f0c039 --- /dev/null +++ b/paper/probe_movement_recovery.py @@ -0,0 +1,233 @@ +#!/usr/bin/env python +"""Does difference refinement recover the true displacement, or overshoot it? + +Nothing measured on real data can answer this, because the true light-state structure is +never known -- the published one is itself a refinement. So the displacement is *injected* +here and the refinement is asked to find it. + +Construction: + + 1. take a model, call it dark; + 2. displace a contiguous stretch of it by a known vector -- that is the true light state; + 3. build the merged light intensity the two-moment physics predicts, + ``|F_D + alpha dF|^2 + sigma_alpha^2 |dF|^2``, with a chosen alpha and lambda; + 4. add noise with sigmas grafted from a real dataset; + 5. hand the CLI the dark model as the *starting point for both states*, so the + refinement has to discover the displacement rather than be handed it. + +The recovered displacement is then compared with the injected one. A refinement that +believes the whole observed difference -- including the positive contamination that +crystal-to-crystal activation spread puts there -- should have to move the model further +than the truth to explain it. + +Writes the two MTZs and prints the CLI command; run that, then re-run with ``--score`` to +compare the refined models against the injected truth. +""" + +import argparse +import json +import sys +from pathlib import Path + +import numpy as np +import torch + + +def displaced_copy(pdb_path, out_path, chain, first, last, shift, verbose=0): + """Write a copy of `pdb_path` with residues [first, last] of `chain` moved by `shift`. + + Returns the number of atoms actually moved, so a selection that matched nothing is + caught rather than silently producing a zero-displacement truth. + """ + import gemmi + + st = gemmi.read_structure(str(pdb_path)) + moved = 0 + for model in st: + for ch in model: + if ch.name != chain: + continue + for res in ch: + if first <= res.seqid.num <= last: + for atom in res: + atom.pos = gemmi.Position( + atom.pos.x + shift[0], + atom.pos.y + shift[1], + atom.pos.z + shift[2], + ) + moved += 1 + break + if moved == 0: + raise ValueError( + f"selection chain {chain} residues {first}-{last} matched no atoms; " + f"the injected displacement would be zero" + ) + st.write_pdb(str(out_path)) + if verbose: + print(f" displaced {moved} atoms by {np.linalg.norm(shift):.3f} A") + return moved + + +def simulate(dark_pdb, light_pdb, reference_mtz, out_dir, alpha, lam, d_min, + sigma_mul, seed, device): + """Write dark.mtz / light.mtz carrying two-moment intensities plus noise.""" + from torchref import ReflectionData + from torchref.cli._common import load_model + from torchref.io.datasets import FcalcDataset + + ref = ReflectionData(device=str(device), verbose=0).load_mtz(str(reference_mtz)) + ref.cut_res(highres=d_min) + + md = load_model(str(dark_pdb), max_res=d_min, device=device, verbose=0) + ml = load_model(str(light_pdb), max_res=d_min, device=device, verbose=0) + + with torch.no_grad(): + hkl = ref.hkl + F_D = ref.structure_factors(md, recalc=True) + F_L = ref.structure_factors(ml, recalc=True) + dF = F_L - F_D + + sigma_alpha_sq = alpha * (1.0 - alpha) * lam + I_dark = F_D.abs() ** 2 + I_light = (F_D + alpha * dF).abs() ** 2 + sigma_alpha_sq * dF.abs() ** 2 + + frac_contam = float( + (sigma_alpha_sq * dF.abs() ** 2 / I_light.clamp(min=1e-12)).median() + ) + + out_dir = Path(out_dir) + out_dir.mkdir(parents=True, exist_ok=True) + paths = {} + for name, intensity in (("dark", I_dark), ("light", I_light)): + ds = FcalcDataset( + hkl=hkl.clone(), cell=ref.cell, spacegroup=ref.spacegroup, device=device + ) + # Phase is irrelevant to the written intensities but set_fcalc wants a complex. + ds.set_fcalc((intensity.clamp(min=0).sqrt() + 0j).to(torch.complex64)) + noisy = ds.add_noise(sigma_mul=sigma_mul, seed=seed, verbose=False) + # Write the intensities add_noise actually drew, negatives and all. + data = ReflectionData(device=str(device), verbose=0).from_tensors( + hkl=hkl.clone(), + F=noisy.fcalc_amp.clone(), + F_sigma=noisy.fobs_sigma.clone(), + cell=ref.cell, + spacegroup=ref.spacegroup, + rfree_flags=ref.rfree_flags.clone(), + device=str(device), + verbose=0, + ) + data.I = noisy.I.clone() + data.I_sigma = noisy.I_sigma.clone() + p = out_dir / f"{name}.mtz" + data.write_mtz(str(p)) + paths[name] = str(p) + + meta = dict(alpha=alpha, lambda_twin=lam, sigma_alpha_sq=sigma_alpha_sq, + d_min=d_min, sigma_mul=sigma_mul, seed=seed, + median_contamination_fraction=frac_contam, **paths) + (out_dir / "truth.json").write_text(json.dumps(meta, indent=2)) + return meta + + +def score(truth_pdb, start_pdb, refined_pdbs, chain, first, last): + """Injected vs recovered displacement over the moved residues.""" + import gemmi + + def positions(path): + st = gemmi.read_structure(str(path)) + st.remove_hydrogens() + out = {} + for ch in st[0]: + if ch.name != chain: + continue + for res in ch: + if first <= res.seqid.num <= last: + for atom in res: + out[(res.seqid.num, atom.name)] = np.array( + [atom.pos.x, atom.pos.y, atom.pos.z] + ) + return out + + truth, start = positions(truth_pdb), positions(start_pdb) + shared = sorted(set(truth) & set(start)) + injected = np.array([np.linalg.norm(truth[k] - start[k]) for k in shared]).mean() + + print(f"\ninjected displacement over the moved residues: {injected:.3f} A " + f"({len(shared)} atoms)") + print(f"{'arm':14s} {'recovered':>10s} {'ratio':>8s} {'err vs truth':>13s}") + print("-" * 50) + rows = [] + for label, path in refined_pdbs: + if not Path(path).exists(): + print(f"{label:14s} (missing)") + continue + got = positions(path) + keys = [k for k in shared if k in got] + rec = np.array([np.linalg.norm(got[k] - start[k]) for k in keys]).mean() + err = np.array([np.linalg.norm(got[k] - truth[k]) for k in keys]).mean() + rows.append((label, rec, rec / injected, err)) + print(f"{label:14s} {rec:10.3f} {rec / injected:8.2f} {err:13.3f}") + print("-" * 50) + print("ratio > 1 means the refinement moved further than the truth") + return rows + + +def main(): + ap = argparse.ArgumentParser(description=__doc__) + repo = Path(__file__).resolve().parents[1] + ap.add_argument("--pdb", default=str(repo / "tests/files/pdb/1DAW.pdb")) + ap.add_argument("--reference-mtz", default=str(repo / "tests/files/mtz/1DAW.mtz")) + ap.add_argument("--out", required=True) + ap.add_argument("--chain", default="A") + ap.add_argument("--first", type=int, default=40) + ap.add_argument("--last", type=int, default=52) + ap.add_argument("--shift", type=float, nargs=3, default=[0.35, 0.20, -0.15]) + ap.add_argument("--alpha", type=float, default=0.22) + ap.add_argument("--lambda-twin", type=float, default=0.3) + ap.add_argument("--dmin", type=float, default=2.05) + ap.add_argument("--sigma-mul", type=float, default=0.10) + ap.add_argument("--seed", type=int, default=7) + ap.add_argument("--device", default="cpu") + ap.add_argument("--score", action="store_true", + help="Compare refined models against the truth (after refining).") + args = ap.parse_args() + + out = Path(args.out) + truth_pdb = out / "light_truth.pdb" + + if args.score: + arms = [] + for d in sorted(out.glob("refine_*")): + if not d.is_dir(): + continue + # The output prefix encodes the fraction, which varies across arms. + hits = sorted(d.glob("fractions_*_light.pdb")) + arms.append((d.name, str(hits[0]) if hits else str(d / "missing.pdb"))) + score(truth_pdb, args.pdb, arms, args.chain, args.first, args.last) + return 0 + + out.mkdir(parents=True, exist_ok=True) + print(f"Injecting a displacement into chain {args.chain} " + f"residues {args.first}-{args.last}") + displaced_copy(args.pdb, truth_pdb, args.chain, args.first, args.last, + args.shift, verbose=1) + + print("Simulating two-moment intensities...") + meta = simulate(args.pdb, truth_pdb, args.reference_mtz, out, args.alpha, + args.lambda_twin, args.dmin, args.sigma_mul, args.seed, + torch.device(args.device)) + print(f" alpha={meta['alpha']} lambda={meta['lambda_twin']} " + f"sigma_alpha^2={meta['sigma_alpha_sq']:.4f}") + print(f" median contamination fraction of I: " + f"{meta['median_contamination_fraction']:.3e}") + print(f"\nRefine from the DARK model for both states, e.g.\n") + print(f" torchref.difference-refine -dm {args.pdb} -lm {args.pdb} \\\n" + f" -dsf {meta['dark']} -lsf {meta['light']} \\\n" + f" --fraction {args.alpha} --dmin {args.dmin} " + f"-o {out}/refine_coh --device cpu\n") + print(f"then re-run this with --score --out {out}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/paper/probe_two_moment_collinearity.py b/paper/probe_two_moment_collinearity.py new file mode 100644 index 00000000..5815dbaf --- /dev/null +++ b/paper/probe_two_moment_collinearity.py @@ -0,0 +1,220 @@ +#!/usr/bin/env python +"""How much of the activation-dispersion signal can the scale model absorb? + +The contamination ``sigma_alpha^2 |dF/dalpha|^2`` is smooth, strictly positive and +concentrated at low resolution -- the same shape a scale or overall-B error takes. If the +scale parameters can reproduce it, a refined ``lambda`` is measuring scale error rather +than activation heterogeneity, and no amount of refinement will tell the two apart. + +This is answerable exactly, by linear algebra, with no refinement at all. Whiten every +quantity by the measurement error, treat the contamination as a template ``t`` and the +scale parameters' derivatives as a design matrix ``X``, and project:: + + P = X (X'X)^-1 X' + leak = ||P t|| / ||t|| fraction of the template the scale model can absorb + vif = ||t|| / ||t - P t|| how much the surviving signal is degraded + +Two designs are compared, because they are the two places scale is fitted: + +* **per-dataset** -- the ``log_scale`` + ``U_aniso`` that ``DatasetCollection.scale()`` + fits on the light dataset alone. Free to shape the light data however it likes. +* **shared** -- the ``CollectionScaler`` parameters, which are fitted jointly against dark + and light. One column per parameter spanning *both* datasets, so a light-only template + cannot be matched without spoiling the dark. + +The derivatives are taken numerically from the live scaler rather than from textbook +formulae, so the design matrix is the parameterisation actually in use. + +The template is then split into a smooth resolution envelope and the residual speckle, +and each projected separately: the envelope is what a scale model can absorb, the speckle +is what identifies the dispersion. Which half survives decides how a fitted lambda should +be read. +""" + +import argparse +import json +import sys +from pathlib import Path + +import numpy as np +import torch + + +def build(dark_sf, light_sf, dark_pdb, light_pdb, cif, d_min, fraction, lam, device): + """Rebuild the figure-4 collection exactly as the CLI does.""" + from torchref.cli.collection_difference_refine import ( + setup_dataset_collection, + setup_model_collection, + setup_scaler, + ) + + mc = setup_model_collection( + dark_pdb, light_pdb, [1.0 - fraction, fraction], cif, d_min, device, 0 + ) + dc = setup_dataset_collection(dark_sf, light_sf, d_min, device) + scaler = setup_scaler(dc, mc, device, verbose=0) + mc.set_lambda_twin(lam) + return dc, mc, scaler + + +def light_intensity(dc, mc, scaler): + """Scaled model intensity for the light dataset, shape (n_hkl,).""" + keys = [mc.dark_key, "light"] + rows = [mc.keys().index(k) for k in keys] + comps = dc.component_structure_factors(mc, recalc=True) + w = mc.fractions_matrix()[rows] + return scaler.forward_batched(mc.mix_component_fcalcs(comps, w), w)[1].abs() ** 2 + + +def contamination(dc, mc, scaler): + """sigma_alpha^2 |dF/dalpha|^2 on the light dataset.""" + keys = [mc.dark_key, "light"] + rows = [mc.keys().index(k) for k in keys] + comps = dc.component_structure_factors(mc, recalc=True) + jac = mc.activation_jacobian()[rows] + deriv = scaler.forward_batched(mc.mix_component_fcalcs(comps, jac), jac)[1] + return mc.sigma_alpha_sq * deriv.abs() ** 2 + + +def numeric_columns(params, evaluate, rel_step=1e-3): + """d(model intensity)/d(theta) for every scalar in `params`, by central difference.""" + cols = [] + for p in params: + flat = p.detach().reshape(-1) + for i in range(flat.numel()): + step = rel_step * max(abs(float(flat[i])), 1e-3) + saved = float(flat[i]) + with torch.no_grad(): + flat[i] = saved + step + plus = evaluate() + with torch.no_grad(): + flat[i] = saved - step + minus = evaluate() + with torch.no_grad(): + flat[i] = saved + cols.append(((plus - minus) / (2 * step)).detach().cpu().numpy()) + return np.asarray(cols).T # (n_hkl, n_param) + + +def leakage(t, X): + """(leak, vif) for template `t` against design `X`, both already whitened.""" + keep = ~np.any(~np.isfinite(X), axis=1) & np.isfinite(t) + t, X = t[keep], X[keep] + # Drop null columns, then least-squares project (lstsq handles rank deficiency). + good = np.linalg.norm(X, axis=0) > 0 + X = X[:, good] + if X.shape[1] == 0: + return 0.0, 1.0 + coef, *_ = np.linalg.lstsq(X, t, rcond=None) + fit = X @ coef + nt = np.linalg.norm(t) + resid = np.linalg.norm(t - fit) + return float(np.linalg.norm(fit) / nt), float(nt / max(resid, 1e-30)) + + +def envelope_and_speckle(t, res, nbin=30): + """Split a template into its smooth resolution envelope and the residual.""" + order = np.argsort(-res) + env = np.zeros_like(t) + for chunk in np.array_split(order, nbin): + env[chunk] = t[chunk].mean() + return env, t - env + + +def main(): + ap = argparse.ArgumentParser(description=__doc__) + fig4 = Path(__file__).resolve().parent / "figure4_difference_refinement" + ap.add_argument("--dark-sf", default=str(fig4 / "data/8QL2-sf.cif")) + ap.add_argument("--light-sf", default=str(fig4 / "data/7YYZ-light.mtz")) + ap.add_argument("--dark-pdb", default=str(fig4 / "data/8QL2_no_altloc.pdb")) + ap.add_argument("--light-pdb", default=str(fig4 / "work_no_altloc.pdb")) + ap.add_argument("--cif", nargs="*", default=[str(fig4 / "data/IBL_grade.cif")]) + ap.add_argument("--dmin", type=float, default=2.2) + ap.add_argument("--fraction", type=float, default=0.22) + ap.add_argument("--lambda-twin", type=float, default=0.2) + ap.add_argument("--device", default="cpu") + ap.add_argument("-o", "--out", default=None) + args = ap.parse_args() + + dev = torch.device(args.device) + print("Rebuilding the collection...", flush=True) + dc, mc, scaler = build( + args.dark_sf, args.light_sf, args.dark_pdb, args.light_pdb, + args.cif, args.dmin, args.fraction, args.lambda_twin, dev, + ) + light = dc["light"] + + with torch.no_grad(): + t_raw = contamination(dc, mc, scaler).cpu().numpy() + _, sig_I = light.get_corrected_intensities() + sig = sig_I.cpu().numpy() + res = light.resolution.cpu().numpy() + mask = light.masks().cpu().numpy().astype(bool) + + ok = mask & np.isfinite(sig) & (sig > 0) & np.isfinite(t_raw) & np.isfinite(res) + print(f"reflections: {ok.sum()} of {len(ok)}") + + # Whiten: everything is measured in units of the error it has to beat. + t = (t_raw / sig)[ok] + + # --- design matrices, differentiated numerically --- + # + # Two different response functions, because the two parameter sets act on opposite + # sides of the residual. The shared scaler shapes the *model* intensity; the + # per-dataset log_scale / U_aniso shape the *observed* one. Using the model response + # for both would give the dataset parameters an identically zero column and report + # no leak at all. + # + # The component structure factors are computed once: no parameter here moves an + # atom, so recomputing them per derivative is ~50x of pure waste. + print("Differentiating the scale model...", flush=True) + keys = [mc.dark_key, "light"] + rows = [mc.keys().index(k) for k in keys] + with torch.no_grad(): + comps = dc.component_structure_factors(mc, recalc=True) + + def model_response(): + w = mc.fractions_matrix()[rows] + return scaler.forward_batched( + mc.mix_component_fcalcs(comps, w), w + )[1].abs() ** 2 + + def obs_response(): + light._corrected_I_fp = None + return light.get_corrected_intensities()[0] + + shared_params = list(scaler.parameters()) + X_shared = numeric_columns(shared_params, model_response)[ok] / sig[ok, None] + + data_params = [p for p in light.parameters() if p is not None] + X_data = numeric_columns(data_params, obs_response)[ok] / sig[ok, None] + + env, speck = envelope_and_speckle(t, res[ok]) + + rows = [] + for tname, tv in (("full", t), ("envelope", env), ("speckle", speck)): + for xname, X in (("per-dataset", X_data), ("shared", X_shared)): + leak, vif = leakage(tv, X) + rows.append(dict(template=tname, design=xname, n_param=X.shape[1], + leak=leak, vif=vif)) + + print() + print("Fraction of the activation template the scale model can absorb") + print("=" * 66) + print(f"{'template':10s} {'design':13s} {'n_param':>8s} {'leak':>8s} {'VIF':>8s}") + print("-" * 66) + for r in rows: + print(f"{r['template']:10s} {r['design']:13s} {r['n_param']:8d} " + f"{r['leak']:8.3f} {r['vif']:8.2f}") + print("-" * 66) + print(f"envelope carries {np.linalg.norm(env) / np.linalg.norm(t):.3f} of the " + f"template norm, speckle {np.linalg.norm(speck) / np.linalg.norm(t):.3f}") + + if args.out: + Path(args.out).write_text(json.dumps(rows, indent=2)) + print(f"written: {args.out}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) From 0626d64640e894023b03ae167a254b66d50ecd96 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 14:12:43 +0200 Subject: [PATCH 064/250] Let a topology be subsetted and copied subset reindexes the surviving edges instead of rebuilding, so taking part of a structure no longer means re-reading the monomer CIFs and re-matching every template. copy gives an independent graph. Both are defined at each level -- EdgeBlock, ResidueGraph, AtomGraph, Topology -- so a caller can reduce whichever it holds. An edge is dropped as soon as any of its atoms is. A bond to an atom that is gone is not a bond, and an angle missing its apex is not an angle. Worth being clear about the consequence: a subset is a geometrically weaker model rather than merely a smaller one, because every restraint crossing the boundary goes with it. A residue left with no atoms is dropped too, and so is any link that reached it, since a peptide bond to a residue that is not there would leave an edge pointing outside the graph. No re-sort is needed. The remap is monotone on the atoms it keeps -- survivors are renumbered in their existing order -- and a monotone relabelling preserves lexicographic order, so each origin's rows stay sorted among themselves and the block stays canonical. That is also why subset takes an index list as a set and returns the topology's own atom order: honouring a caller's order would quietly leave the blocks unsorted, and the 'all' group would stop being a contiguous span, which is what makes it a view rather than a copy. There is a test pinning that. The tests compare surviving edges on (chain, resseq, icode, atom name) rather than on index. An index-based check passes trivially after a remap; identity is what catches a remap that points an edge at the wrong atom. This is the primitive, not yet the win the plan claimed for Model.select. That also needs the restraint value layer and the non-bonded pair list reduced in step, or select still rebuilds. Also corrects Topology.neighbors, annotated np.ndarray since before the adjacency moved to torch. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- docs/changelog.rst | 1 + tests/unit/topology/test_subset.py | 305 +++++++++++++++++++++++++++++ torchref/topology/atom_graph.py | 49 +++++ torchref/topology/edges.py | 52 +++++ torchref/topology/residue_graph.py | 59 ++++++ torchref/topology/topology.py | 74 ++++++- 6 files changed, 539 insertions(+), 1 deletion(-) create mode 100644 tests/unit/topology/test_subset.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 690bae20..99c07e9b 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -37,6 +37,7 @@ Unreleased - Riding hydrogens are no longer placed when the model carries real ones, where they acted as phantom atoms in the non-bonded term - Fixed the hydrogen valence cap counting only heavy neighbours, so a parent that already carried a hydrogen still had budget for another; generation was not idempotent and a save/reload added a spurious second amide hydrogen to every linked residue - Moved the riding-hydrogen map from ``torchref.restraints.hydrogen_topology`` to ``torchref.topology.riding``, alongside the generation path it is the heavy-atom-only alternative to +- Added ``Topology.subset`` and ``copy``, plus the same on ``EdgeBlock``, ``ResidueGraph`` and ``AtomGraph``: a subset reindexes the surviving edges rather than re-reading the CIFs and re-matching the templates. An edge is dropped as soon as any of its atoms is, and a residue left with no atoms goes along with its links Version 0.6.4 diff --git a/tests/unit/topology/test_subset.py b/tests/unit/topology/test_subset.py new file mode 100644 index 00000000..60b36126 --- /dev/null +++ b/tests/unit/topology/test_subset.py @@ -0,0 +1,305 @@ +"""Subsetting and copying a topology. + +``subset`` reindexes what survives instead of rebuilding, so the properties that matter +are the ones a plausible-but-wrong implementation would break: that every surviving edge +keeps the *same atoms* it had before, that an edge touching a removed atom is gone +entirely rather than left dangling, that the blocks stay canonically ordered without a +re-sort, and that the residue level stays consistent with the atom level. +""" + +import numpy as np +import pytest +import torch + +from torchref.model.model import Model +from torchref.topology import EdgeBlock + +KEYED_TYPES = ("bond", "angle", "torsion", "chiral") + + +@pytest.fixture(scope="module") +def topology(pdb_dir): + """A topology with altlocs, disulfides, peptide links and hydrogens.""" + model = Model(verbose=0, add_hydrogens=False, strip_H=True) + model.load_pdb(str(pdb_dir / "7L84.pdb")) + model.set_restraints_cif(None) + return model.restraints.topology + + +def _named_edges(topology, edge_type): + """Edges as tuples of ``(chain, resseq, icode, atom name)``, identity not index. + + Comparing on identity rather than index is the whole point: an index-based check + passes trivially after a remap, whereas this catches a remap that points an edge at + the wrong atom. + """ + residues = topology.residues + names = topology.atoms.name.astype(str) + out = set() + block = topology.edge_block(edge_type) + for row in block.indices.cpu().numpy(): + out.add( + tuple( + (*residues.key(topology.residue_of_atom(int(a))), names[int(a)]) + for a in row + ) + ) + return out + + +@pytest.mark.unit +def test_subset_of_everything_is_the_same_graph(topology): + """Keeping every atom changes nothing -- the identity case.""" + whole = topology.subset(torch.ones(topology.n_atoms, dtype=torch.bool)) + + assert whole.n_atoms == topology.n_atoms + assert whole.n_residues == topology.n_residues + for edge_type in KEYED_TYPES: + assert torch.equal( + whole.edge_block(edge_type).indices, + topology.edge_block(edge_type).indices, + ) + assert ( + whole.edge_block(edge_type).origin_bounds + == topology.edge_block(edge_type).origin_bounds + ) + + +@pytest.mark.unit +def test_surviving_edges_keep_the_same_atoms(topology): + """Every edge in the subset joins exactly the atoms it joined before. + + Checked on residue-and-name identity, so a remap that silently shifted an index + would fail here even though the shapes still looked right. + """ + residues = topology.residues + chain_a = np.array( + [ + str(residues.chain[topology.residue_of_atom(atom)]) == "A" + for atom in range(topology.n_atoms) + ] + ) + assert chain_a.any() and not chain_a.all() or chain_a.all() + + keep = torch.zeros(topology.n_atoms, dtype=torch.bool) + keep[: topology.n_atoms // 2] = True + reduced = topology.subset(keep) + + for edge_type in KEYED_TYPES: + after = _named_edges(reduced, edge_type) + before = _named_edges(topology, edge_type) + assert after <= before, ( + f"{edge_type}: the subset invented {len(after - before)} edges that were " + f"not in the original" + ) + + +@pytest.mark.unit +def test_an_edge_dies_with_any_of_its_atoms(topology): + """No edge survives that touches a removed atom. + + A bond to an atom that is gone is not a bond, and leaving it would point an index + outside the graph. + """ + keep = torch.ones(topology.n_atoms, dtype=torch.bool) + keep[5] = False + keep[100] = False + reduced = topology.subset(keep) + + for edge_type in KEYED_TYPES: + block = reduced.edge_block(edge_type) + if block.n_edges == 0: + continue + assert int(block.indices.max()) < reduced.n_atoms + assert int(block.indices.min()) >= 0 + + # Precisely: the edges lost are exactly those that used a dropped atom. + dropped = {5, 100} + for edge_type in KEYED_TYPES: + original = topology.edge_block(edge_type).indices.cpu().numpy() + touching = sum(1 for row in original if dropped & set(int(a) for a in row)) + assert ( + reduced.edge_block(edge_type).n_edges + == topology.edge_block(edge_type).n_edges - touching + ), f"{edge_type}: wrong number of edges dropped" + + +@pytest.mark.unit +def test_blocks_stay_canonically_ordered(topology): + """Subsetting needs no re-sort, because the remap is monotone on survivors. + + If this fails the block is no longer canonical, and the ``all`` group stops being a + contiguous span -- which is what makes it a view rather than a copy. + """ + keep = torch.zeros(topology.n_atoms, dtype=torch.bool) + keep[::2] = True + reduced = topology.subset(keep) + + for edge_type in KEYED_TYPES: + block = reduced.edge_block(edge_type) + for origin in block.origins(): + rows = block.origin(origin).cpu().numpy() + if len(rows) < 2: + continue + order = np.lexsort( + tuple(rows[:, c] for c in reversed(range(rows.shape[1]))) + ) + assert ( + order == np.arange(len(rows)) + ).all(), f"{edge_type}/{origin} is no longer lexicographically sorted" + # bounds must remain contiguous and cover the block + spans = sorted(block.origin_bounds.values()) + assert all(a[1] == b[0] for a, b in zip(spans, spans[1:])) + if spans: + assert spans[0][0] == 0 and spans[-1][1] == block.n_edges + + +@pytest.mark.unit +def test_residue_level_stays_consistent_with_the_atoms(topology): + """Atom ranges, residue count and links all agree after subsetting.""" + keep = torch.zeros(topology.n_atoms, dtype=torch.bool) + keep[: topology.n_atoms // 3] = True + reduced = topology.subset(keep) + + residues = reduced.residues + assert residues.n_residues == reduced.n_residues + # Ranges partition the atoms, in order, with no gaps. + assert int(residues.atom_start[0]) == 0 + assert int(residues.atom_end[-1]) == reduced.n_atoms + assert (residues.atom_start[1:] == residues.atom_end[:-1]).all() + + # residue_of agrees with the ranges it is supposed to index. + for residue in range(residues.n_residues): + rows = range(int(residues.atom_start[residue]), int(residues.atom_end[residue])) + for row in rows: + assert reduced.residue_of_atom(row) == residue + + # No link edge points outside the surviving residues. + if len(residues.link_pairs): + assert residues.link_pairs.max() < residues.n_residues + assert residues.link_pairs.min() >= 0 + + +@pytest.mark.unit +def test_dropping_a_whole_residue_drops_its_links(topology): + """A residue with no atoms left is gone, and so is any link that reached it.""" + residues = topology.residues + linked = None + for pair in residues.links_of_kind("TRANS"): + linked = int(pair[0]) + break + assert linked is not None, "7L84 has peptide links" + + keep = torch.ones(topology.n_atoms, dtype=torch.bool) + for row in range(int(residues.atom_start[linked]), int(residues.atom_end[linked])): + keep[row] = False + reduced = topology.subset(keep) + + assert reduced.n_residues == topology.n_residues - 1 + before = len(residues.links_of_kind("TRANS")) + after = len(reduced.residues.links_of_kind("TRANS")) + assert after < before, "a link to the removed residue survived" + + +@pytest.mark.unit +def test_subset_accepts_indices_as_well_as_a_mask(topology): + """Integer indices give the same result as the equivalent mask.""" + indices = torch.arange(0, topology.n_atoms, 3) + mask = torch.zeros(topology.n_atoms, dtype=torch.bool) + mask[indices] = True + + by_index = topology.subset(indices) + by_mask = topology.subset(mask) + + assert by_index.n_atoms == by_mask.n_atoms + for edge_type in KEYED_TYPES: + assert torch.equal( + by_index.edge_block(edge_type).indices, + by_mask.edge_block(edge_type).indices, + ) + + +@pytest.mark.unit +def test_subset_ignores_the_order_of_the_indices(topology): + """A shuffled index list yields the topology's own atom order. + + Deliberate: the edge blocks stay canonical only under a monotone relabelling, so + honouring a caller's order would silently leave them unsorted. + """ + indices = torch.arange(0, topology.n_atoms, 5) + shuffled = indices[torch.randperm(len(indices))] + + assert torch.equal( + topology.subset(indices).atoms.bonds.indices, + topology.subset(shuffled).atoms.bonds.indices, + ) + + +@pytest.mark.unit +def test_subset_of_nothing_is_an_error(topology): + """An empty selection is a mistake, not an empty topology.""" + with pytest.raises(ValueError, match="no atoms"): + topology.subset(torch.zeros(topology.n_atoms, dtype=torch.bool)) + + +@pytest.mark.unit +def test_copy_shares_nothing(topology): + """A copy has equal contents and independent storage.""" + duplicate = topology.copy() + + assert duplicate.n_atoms == topology.n_atoms + assert duplicate.n_residues == topology.n_residues + for edge_type in KEYED_TYPES: + original = topology.edge_block(edge_type) + copied = duplicate.edge_block(edge_type) + assert torch.equal(copied.indices, original.indices) + assert copied.indices.data_ptr() != original.indices.data_ptr() + + saved = int(topology.atoms.bonds.indices[0, 0]) + duplicate.atoms.bonds.indices[0, 0] = saved + 13 + assert int(topology.atoms.bonds.indices[0, 0]) == saved + + duplicate.residues.resname[0] = "XXX" + assert topology.residues.resname[0] != "XXX" + + +@pytest.mark.unit +def test_copy_rebuilds_its_own_adjacency(topology): + """``neighbors`` on a copy reads the copy's bonds, not the original's.""" + duplicate = topology.copy() + atom = int(topology.atoms.bonds.indices[0, 0]) + assert torch.equal(duplicate.neighbors(atom), topology.neighbors(atom)) + assert ( + duplicate.atoms._adj_indices.data_ptr() + != topology.atoms._adj_indices.data_ptr() + ) + + +@pytest.mark.unit +def test_subset_adjacency_matches_its_own_bonds(topology): + """The reduced graph's adjacency is rebuilt, not carried over stale.""" + keep = torch.zeros(topology.n_atoms, dtype=torch.bool) + keep[: topology.n_atoms // 2] = True + reduced = topology.subset(keep) + + from_adjacency = set() + for atom in range(reduced.n_atoms): + for other in reduced.neighbors(atom).cpu().tolist(): + from_adjacency.add((min(atom, other), max(atom, other))) + + from_block = { + (min(int(a), int(b)), max(int(a), int(b))) + for a, b in reduced.atoms.bonds.indices.cpu().numpy() + } + assert from_adjacency == from_block + assert int(reduced.atoms.degree().sum()) == 2 * reduced.atoms.bonds.n_edges + + +@pytest.mark.unit +def test_empty_block_subsets_to_empty(): + """An edge type with no edges survives subsetting without special-casing.""" + block = EdgeBlock.empty(3) + remap = torch.arange(10, dtype=torch.int64) + reduced = block.subset(remap) + assert reduced.n_edges == 0 + assert reduced.arity == 3 diff --git a/torchref/topology/atom_graph.py b/torchref/topology/atom_graph.py index 55af8f8e..2dfce660 100644 --- a/torchref/topology/atom_graph.py +++ b/torchref/topology/atom_graph.py @@ -159,6 +159,55 @@ def is_hydrogen(self) -> torch.Tensor: flags = np.char.upper(np.char.strip(self.element.astype(str))) == "H" return torch.as_tensor(flags, device=self.bonds.indices.device) + def copy(self) -> "AtomGraph": + """An independent copy sharing no storage with this one.""" + return AtomGraph( + name=self.name.copy(), + element=self.element.copy(), + altloc=self.altloc.copy(), + residue_of=self.residue_of.clone(), + bonds=self.bonds.copy(), + angles=self.angles.copy(), + torsions=self.torsions.copy(), + chirals=self.chirals.copy(), + planes={size: block.copy() for size, block in self.planes.items()}, + ) + + def subset(self, remap: torch.Tensor, residue_remap: torch.Tensor) -> "AtomGraph": + """The atoms ``remap`` keeps, with every edge set reindexed. + + Parameters + ---------- + remap : torch.Tensor + Old atom index to new, shape ``(N_old,)``, ``-1`` where dropped. + residue_remap : torch.Tensor + Old residue index to new, shape ``(R_old,)``, ``-1`` where dropped. + + Returns + ------- + AtomGraph + Adjacency is rebuilt from the surviving bond block rather than subsetted: + CSR row offsets are not meaningful once the atoms are renumbered. + """ + keep = (remap >= 0).cpu().numpy() + planes = {} + for size, block in self.planes.items(): + reduced = block.subset(remap) + if reduced.n_edges: + planes[size] = reduced + + return AtomGraph( + name=self.name[keep], + element=self.element[keep], + altloc=self.altloc[keep], + residue_of=residue_remap[self.residue_of[torch.as_tensor(keep)]], + bonds=self.bonds.subset(remap), + angles=self.angles.subset(remap), + torsions=self.torsions.subset(remap), + chirals=self.chirals.subset(remap), + planes=planes, + ) + def rebuild_adjacency(self) -> None: """Rebuild the CSR adjacency from the current bond block.""" self._adj_indptr, self._adj_indices = _build_csr( diff --git a/torchref/topology/edges.py b/torchref/topology/edges.py index 50e2d960..06223a81 100644 --- a/torchref/topology/edges.py +++ b/torchref/topology/edges.py @@ -235,6 +235,58 @@ def origin(self, name: str) -> torch.Tensor: start, end = self.origin_bounds[name] return self.indices[start:end] + def copy(self) -> "EdgeBlock": + """An independent copy sharing no storage with this one.""" + return EdgeBlock( + indices=self.indices.clone(), + origin_bounds=dict(self.origin_bounds), + ) + + def subset(self, remap: torch.Tensor) -> "EdgeBlock": + """Edges whose every atom survives, reindexed by ``remap``. + + Parameters + ---------- + remap : torch.Tensor + Old atom index to new, shape ``(N_old,)``, with ``-1`` where the atom is + being dropped. An edge is kept only if none of its atoms maps to ``-1``: a + bond to a removed atom is not a bond, and an angle missing its apex is not + an angle. + + Returns + ------- + EdgeBlock + Canonically ordered, with ``origin_bounds`` recomputed over the survivors. + + Notes + ----- + No re-sort is needed. ``remap`` is monotone on the atoms it keeps -- survivors + are renumbered in their existing order -- and a monotone relabelling preserves + lexicographic order, so each origin's surviving rows stay sorted among + themselves. + """ + if self.n_edges == 0: + return EdgeBlock.empty(self.arity, device=self.indices.device) + + mapped = remap[self.indices] + keep = (mapped >= 0).all(dim=1) + + chunks = [] + bounds: Dict[str, Tuple[int, int]] = {} + cursor = 0 + for origin in self.origins(): + start, end = self.origin_bounds[origin] + surviving = mapped[start:end][keep[start:end]] + if surviving.shape[0] == 0: + continue + chunks.append(surviving) + bounds[origin] = (cursor, cursor + surviving.shape[0]) + cursor += surviving.shape[0] + + if not chunks: + return EdgeBlock.empty(self.arity, device=self.indices.device) + return EdgeBlock(indices=torch.cat(chunks, dim=0), origin_bounds=bounds) + def tuple_set(self, origin: str = None) -> set: """Edges as a set of index tuples, for order-free comparison. diff --git a/torchref/topology/residue_graph.py b/torchref/topology/residue_graph.py index 453d6f61..b4651d1c 100644 --- a/torchref/topology/residue_graph.py +++ b/torchref/topology/residue_graph.py @@ -98,6 +98,65 @@ def atom_rows(self, i: int) -> range: """Row range of residue ``i``'s atoms.""" return range(int(self.atom_start[i]), int(self.atom_end[i])) + def copy(self) -> "ResidueGraph": + """An independent copy sharing no arrays with this one.""" + return ResidueGraph( + chain=self.chain.copy(), + resseq=self.resseq.copy(), + icode=self.icode.copy(), + resname=self.resname.copy(), + template_key=self.template_key.copy(), + atom_start=self.atom_start.copy(), + atom_end=self.atom_end.copy(), + link_pairs=self.link_pairs.copy(), + link_kind=self.link_kind.copy(), + ) + + def subset( + self, keep: np.ndarray, atom_start: np.ndarray, atom_end: np.ndarray + ) -> "ResidueGraph": + """The residues in ``keep``, with the atom ranges the caller recomputed. + + Parameters + ---------- + keep : numpy.ndarray + Boolean mask over residues, shape ``(R,)``. + atom_start, atom_end : numpy.ndarray + New half-open atom ranges for the surviving residues, in their order, shape + ``(R_kept,)``. Passed in rather than derived here because only the caller + knows how the atoms were renumbered. + + Returns + ------- + ResidueGraph + Link edges are kept only where **both** endpoints survive, and reindexed. A + peptide bond to a residue that is gone is not a peptide bond, and keeping it + would leave an edge pointing outside the graph. + """ + remap = np.full(self.n_residues, -1, dtype=np.int64) + remap[keep] = np.arange(int(keep.sum()), dtype=np.int64) + + if len(self.link_pairs): + mapped = remap[self.link_pairs] + survives = (mapped >= 0).all(axis=1) + link_pairs = mapped[survives] + link_kind = self.link_kind[survives] + else: + link_pairs = np.zeros((0, 2), dtype=np.int64) + link_kind = np.zeros(0, dtype=" np.ndarray: """Link edges of one kind, shape ``(L_k, 2)``.""" if len(self.link_kind) == 0: diff --git a/torchref/topology/topology.py b/torchref/topology/topology.py index 7ec289c1..78a7222a 100644 --- a/torchref/topology/topology.py +++ b/torchref/topology/topology.py @@ -55,7 +55,79 @@ def n_residues(self) -> int: """Number of residue nodes.""" return self.residues.n_residues - def neighbors(self, i: int) -> np.ndarray: + def copy(self) -> "Topology": + """An independent copy sharing no storage with this one.""" + return Topology(residues=self.residues.copy(), atoms=self.atoms.copy()) + + def subset(self, keep) -> "Topology": + """The topology over a subset of the atoms. + + Reindexes what survives instead of rebuilding: no CIF is re-read and no template + is re-matched, which is what made ``Model.select`` expensive. + + Parameters + ---------- + keep : torch.Tensor or numpy.ndarray + Boolean mask over atoms, shape ``(N,)``, or integer atom indices. Indices + are taken as a set, not an order -- the result keeps the topology's own atom + order, because the edge blocks stay canonical only under a monotone + relabelling. + + Returns + ------- + Topology + Atoms in their original relative order. A residue with no surviving atoms is + dropped, and any link edge touching it goes with it. + + Notes + ----- + Selecting part of a residue leaves that residue's restraints partial: an edge + loses its whole restraint as soon as one of its atoms goes. That is the honest + outcome -- half a peptide plane is not a plane -- but it means a subset is a + weaker geometric model, not merely a smaller one. + """ + mask = torch.as_tensor(keep) + if mask.dtype != torch.bool: + selected = torch.zeros(self.n_atoms, dtype=torch.bool) + selected[mask.to(torch.int64)] = True + mask = selected + mask = mask.to(device=self.atoms.residue_of.device) + + if int(mask.sum()) == 0: + raise ValueError("subset would keep no atoms") + + n_kept = int(mask.sum()) + remap = torch.full((self.n_atoms,), -1, dtype=torch.int64, device=mask.device) + remap[mask] = torch.arange(n_kept, dtype=torch.int64, device=mask.device) + + # A residue survives if any of its atoms does. Counting per residue also + # gives the new atom ranges, contiguous because the atom order is unchanged. + residue_of = self.atoms.residue_of + per_residue = ( + torch.bincount(residue_of[mask], minlength=self.n_residues).cpu().numpy() + ) + residue_keep = per_residue > 0 + counts = per_residue[residue_keep] + atom_end = np.cumsum(counts) + atom_start = atom_end - counts + + residue_remap = torch.full( + (self.n_residues,), -1, dtype=torch.int64, device=mask.device + ) + residue_remap[torch.as_tensor(residue_keep, device=mask.device)] = torch.arange( + int(residue_keep.sum()), dtype=torch.int64, device=mask.device + ) + + return Topology( + residues=self.residues.subset( + residue_keep, + atom_start.astype(np.int64), + atom_end.astype(np.int64), + ), + atoms=self.atoms.subset(remap, residue_remap), + ) + + def neighbors(self, i: int) -> torch.Tensor: """Atoms bonded to atom ``i``. Delegates to :meth:`AtomGraph.neighbors`.""" return self.atoms.neighbors(i) From cba21f9da424773d07bc81801bdca1838d67b324 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 14:51:35 +0200 Subject: [PATCH 065/250] Cover residues that differ only by an insertion code The builders group residues on (chain, resseq), so 100 and 100A become one residue with two sets of backbone atom names. The name-to-index map keeps the first of each, and the later residue ends up with no intra-residue geometry at all. Keying on (chain, resseq, icode) is what fixes that, and until now nothing exercised it: no bundled structure has an insertion code. Synthesised in the test rather than shipped as another data file. Only columns 23-27 of 3GR5 change, so it is visibly a renumbering and nothing else, and the residues stay in file order as a real insertion would. Measured on the synthetic case: the legacy grouping finds 1104 intra-residue bonds and the graph 1119, losing none. The 15 it misses are exactly the second and third inserted residues, which get zero bonds each under the merged grouping, and the test asserts that localisation rather than just the count -- a graph that produced more restraints everywhere would pass a bare superset check. The first version of this test proved nothing and is worth recording as a trap: it compared the graph against restraints.py, which since the storage swap *builds from* the topology. Identical counts on both sides were the giveaway. The independent baseline is BondRestraintBuilder, which still keys on (chain, resseq). Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- docs/changelog.rst | 1 + tests/unit/topology/test_insertion_codes.py | 283 ++++++++++++++++++++ 2 files changed, 284 insertions(+) create mode 100644 tests/unit/topology/test_insertion_codes.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 99c07e9b..0dbe55b7 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -38,6 +38,7 @@ Unreleased - Fixed the hydrogen valence cap counting only heavy neighbours, so a parent that already carried a hydrogen still had budget for another; generation was not idempotent and a save/reload added a spurious second amide hydrogen to every linked residue - Moved the riding-hydrogen map from ``torchref.restraints.hydrogen_topology`` to ``torchref.topology.riding``, alongside the generation path it is the heavy-atom-only alternative to - Added ``Topology.subset`` and ``copy``, plus the same on ``EdgeBlock``, ``ResidueGraph`` and ``AtomGraph``: a subset reindexes the surviving edges rather than re-reading the CIFs and re-matching the templates. An edge is dropped as soon as any of its atoms is, and a residue left with no atoms goes along with its links +- Fixed residues distinguished only by an insertion code losing their restraints: the builders group on ``(chain, resseq)``, so 100 and 100A merge into one residue whose atom names collide and only the first keeps any intra-residue geometry Version 0.6.4 diff --git a/tests/unit/topology/test_insertion_codes.py b/tests/unit/topology/test_insertion_codes.py new file mode 100644 index 00000000..f402b7da --- /dev/null +++ b/tests/unit/topology/test_insertion_codes.py @@ -0,0 +1,283 @@ +"""Residues distinguished only by an insertion code. + +A deposited structure may number two residues 100 and 100A. They are different residues +with different chemistry, and the only thing separating them is the insertion code. The +topology keys residues on ``(chain, resseq, icode)`` for that reason; the restraint +builders key on ``(chain, resseq)`` alone, which merges them into one residue whose +atom names then collide, so the name-to-index map keeps the first of each and every +restraint belonging to the later residues is silently lost. + +No bundled structure has an insertion code, so the case is synthesised here rather +than shipped as another data file: the rewrite is then visible, and it is obvious that +nothing but the numbering changed. +""" + +import pytest + +from torchref.model.model import Model +from torchref.topology import build_topology + +#: Base structure: chain A, no altlocs, no insertion codes anywhere. +BASE = "3GR5" + +#: The three consecutive residues collapsed onto one sequence number. Their real +#: identities differ (SER, LEU, GLU), so the restraints of the second and third are +#: distinguishable from the first's rather than being duplicates of it. +STRETCH = (23, 24, 25) + + +def _rewrite_with_insertion_codes(source, destination): + """Copy a PDB, renumbering ``STRETCH`` as ``N``, ``NA``, ``NB``. + + Only columns 23-27 change -- the sequence number and the insertion code. Every + atom, coordinate and residue name is untouched, and the residues stay in file order, + so they remain contiguous exactly as a real insertion would be. + + Returns + ------- + tuple of tuple + The ``(resseq, icode)`` pairs written, in order. + """ + first = STRETCH[0] + codes = ["", "A", "B"] + mapping = {old: (first, codes[i]) for i, old in enumerate(STRETCH)} + + out = [] + for line in source.read_text().splitlines(keepends=True): + if line.startswith(("ATOM", "HETATM")): + resseq = int(line[22:26]) + if resseq in mapping: + new_seq, icode = mapping[resseq] + line = f"{line[:22]}{new_seq:>4d}{icode:1s}{line[27:]}" + out.append(line) + destination.write_text("".join(out)) + return tuple((first, code) for code in codes) + + +@pytest.fixture(scope="module") +def inserted(pdb_dir, tmp_path_factory): + """``(topology, restraints, expected keys, model)`` for the rewritten file.""" + path = tmp_path_factory.mktemp("icode") / f"{BASE}_icode.pdb" + expected = _rewrite_with_insertion_codes(pdb_dir / f"{BASE}.pdb", path) + + model = Model(verbose=0, strip_H=True, add_hydrogens=False) + model.load_pdb(str(path)) + model.set_restraints_cif(None) + restraints = model.restraints + + topology = build_topology( + model.pdb, + restraints.cif_dict, + link_dict=restraints.link_dict, + link_list=restraints.link_list, + links=restraints.links, + xyz=model.xyz().detach(), + verbose=0, + ) + return topology, restraints, expected, model + + +def _tuples(indices): + return {tuple(int(v) for v in row) for row in indices.cpu().numpy()} + + +@pytest.mark.unit +def test_the_rewrite_actually_produced_insertion_codes(inserted): + """Guard the fixture: if the rewrite silently failed the rest proves nothing.""" + _, _, _, model = inserted + icodes = model.pdb["icode"].astype(str).str.strip().values + assert set(icodes[icodes != ""]) == {"A", "B"} + + +@pytest.mark.unit +def test_the_graph_keeps_them_apart(inserted): + """Three residue nodes, one per insertion code.""" + topology, _, expected, _ = inserted + residues = topology.residues + + found = [ + residues.key(i) + for i in range(residues.n_residues) + if (int(residues.resseq[i]), str(residues.icode[i]).strip()) + in {(seq, code) for seq, code in expected} + ] + assert len(found) == 3, f"expected three inserted residues, found {found}" + assert len({key[2] for key in found}) == 3, "insertion codes were not distinguished" + + +@pytest.mark.unit +def test_the_builders_merge_them(inserted): + """The comparison only means something if the old grouping really does merge. + + ``PreprocessedPDB`` groups on ``(chain, resseq)``, so the three residues become one + with three sets of backbone atom names. + """ + from torchref.restraints.builders_fast import PreprocessedPDB + + _, _, expected, model = inserted + preprocessed = PreprocessedPDB(model.pdb) + + merged = [ + i + for i in range(preprocessed.n_residues) + if int(preprocessed.residue_resseqs[i]) == expected[0][0] + ] + assert len(merged) == 1, "the builders did not merge the inserted residues" + assert preprocessed.has_duplicate_atoms(merged[0]), ( + "the merged residue should carry duplicate atom names, which is what makes the " + "name-to-index map lose the later residues" + ) + + +def _legacy_intra_bonds(model, restraints): + """Intra-residue bonds as the ``(chain, resseq)``-keyed builder produces them. + + ``restraints.py`` now builds from the topology, so it cannot serve as the + baseline -- it *is* the graph. ``BondRestraintBuilder`` is the original path, still + keying residues on ``(chain, resseq)``, which is the behaviour under test. + """ + import torch + + from torchref.restraints.builders_fast import BondRestraintBuilder + + built = BondRestraintBuilder(verbose=0).build( + model.pdb, restraints.cif_dict, torch.device("cpu") + ) + if not built: + return set() + return {tuple(int(v) for v in row) for row in built["indices"].cpu().numpy()} + + +def _inserted_residue_indices(topology, expected): + wanted = {(seq, code) for seq, code in expected} + return [ + i + for i in range(topology.n_residues) + if (int(topology.residues.resseq[i]), str(topology.residues.icode[i]).strip()) + in wanted + ] + + +def _bonds_within(edges, topology, residue): + start = int(topology.residues.atom_start[residue]) + end = int(topology.residues.atom_end[residue]) + return {e for e in edges if all(start <= int(a) < end for a in e)} + + +@pytest.mark.unit +def test_the_legacy_grouping_loses_the_later_residues(inserted): + """The merged residue gets restraints for its first component only. + + This is the defect keying on ``(chain, resseq, icode)`` fixes. The three residues + become one, their backbone atom names collide, the name-to-index map keeps the first + of each, and the second and third end up with no intra-residue geometry at all. + """ + topology, restraints, expected, model = inserted + legacy = _legacy_intra_bonds(model, restraints) + residues = _inserted_residue_indices(topology, expected) + assert len(residues) == 3 + + counts = [len(_bonds_within(legacy, topology, r)) for r in residues] + assert counts[0] > 0, "even the first component lost its bonds; check the fixture" + assert counts[1:] == [0, 0], ( + f"the legacy grouping was expected to lose the second and third residues, " + f"but found {counts} bonds in them" + ) + + +@pytest.mark.unit +def test_the_graph_finds_what_the_legacy_grouping_lost(inserted): + """Every inserted residue gets its own bonds, and the graph is a strict superset. + + Localised, not merely larger: the bonds the graph adds all lie inside the inserted + residues, so this is the insertion-code fix rather than a general difference. + """ + topology, restraints, expected, model = inserted + legacy = _legacy_intra_bonds(model, restraints) + graph = topology.atoms.bonds.tuple_set("intra") + residues = _inserted_residue_indices(topology, expected) + + for residue in residues: + assert _bonds_within( + graph, topology, residue + ), f"residue {topology.residues.key(residue)} has no intra-residue bonds" + + gained = graph - legacy + assert gained, "the graph found nothing the legacy grouping missed" + + inserted_set = set(residues) + stray = [ + edge + for edge in gained + if not ({topology.residue_of_atom(a) for a in edge} & inserted_set) + ] + assert not stray, ( + f"{len(stray)} gained bonds lie outside the inserted residues, so the " + f"difference is not localised to the insertion codes: {stray[:3]}" + ) + + +@pytest.mark.unit +def test_the_inserted_residues_get_their_own_intra_restraints(inserted): + """Each of the three carries bonds of its own, not just the first. + + Under the merged grouping only the first residue's template matched, so the second + and third had no intra-residue geometry at all. + """ + topology, _, expected, _ = inserted + + inserted_residues = [ + i + for i in range(topology.n_residues) + if ( + int(topology.residues.resseq[i]), + str(topology.residues.icode[i]).strip(), + ) + in {(seq, code) for seq, code in expected} + ] + + intra = topology.atoms.bonds.origin("intra").cpu().numpy() + for residue in inserted_residues: + start = int(topology.residues.atom_start[residue]) + end = int(topology.residues.atom_end[residue]) + own = [ + row + for row in intra + if start <= int(row[0]) < end and start <= int(row[1]) < end + ] + assert own, ( + f"residue {topology.residues.key(residue)} " + f"({topology.residues.resname[residue]}) has no intra-residue bonds" + ) + + +@pytest.mark.unit +def test_the_inserted_stretch_is_peptide_linked(inserted): + """An insertion-code step is a sequence step, so the chain is not broken. + + ``find_peptide_links`` allows a ``resseq`` difference of 0 precisely for this: 100 + to 100A is consecutive. Without it the inserted residues would float free of the + chain. + """ + topology, _, expected, _ = inserted + + inserted_residues = { + i + for i in range(topology.n_residues) + if ( + int(topology.residues.resseq[i]), + str(topology.residues.icode[i]).strip(), + ) + in {(seq, code) for seq, code in expected} + } + + links = topology.residues.links_of_kind("TRANS") + internal = [ + pair + for pair in links + if int(pair[0]) in inserted_residues and int(pair[1]) in inserted_residues + ] + assert len(internal) == 2, ( + f"expected two peptide links inside the three inserted residues, got " + f"{len(internal)}" + ) From bb0c42ef202a99172f5ffc8e6fa26a81482d6a1b Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 15:22:12 +0200 Subject: [PATCH 066/250] Move the restraint orchestrator into the topology package torchref.restraints.restraints becomes torchref.topology.restraints, and RestraintsNew becomes Restraints -- the "New" was historical, and the package alias already meant most callers said Restraints anyway. Moving the orchestrator in turned out to be a smaller change than extracting its pieces out one at a time. The real coupling was three production import sites, not the ninety an attribute-access count suggests, and two of those disappear here: topology/build.py was importing _lookup_link_atom back out of the restraints module, and restraints/__init__.py held the alias. Sixteen test imports were mechanical. _lookup_link_atom moves to topology/build.py, which is the only thing that uses it. That was not optional -- with the orchestrator inside topology, importing it back out would have been circular. The backwards import was the clearest sign the split sat in the wrong place. Nothing on the hot path changes. Model.restraints returns the same object, so self.restraints.restraints[...] in the geometry targets is untouched, and no file under refinement/targets/ is edited. What is left in torchref.restraints is the data layer: the monomer library, the CIF readers, the chem_mod records that patch a template when a link forms, the Numba matchers, the Ramachandran surfaces and the spatial search behind the non-bonded pair list. What is *built* from that data now lives in torchref.topology. The package docstrings say so, and that also settles where cif_dict belongs -- the loader that reads it is in topology, the library that supplies it stays put. Suite 1938 passed, unchanged by the move itself. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- docs/changelog.rst | 2 + tests/conftest.py | 2 +- .../functional/test_restraints_functional.py | 22 +++--- tests/helpers/device_cases.py | 2 +- tests/integration/test_refinement_pipeline.py | 2 +- tests/unit/restraints/test_restraints.py | 6 +- tests/unit/topology/test_storage.py | 2 +- torchref/__init__.py | 3 +- torchref/model/model.py | 6 +- torchref/refinement/base_refinement.py | 2 +- torchref/restraints/__init__.py | 22 +++--- torchref/topology/__init__.py | 2 + torchref/topology/build.py | 53 +++++++++++--- .../{restraints => topology}/restraints.py | 71 ++++++------------- 14 files changed, 105 insertions(+), 92 deletions(-) rename torchref/{restraints => topology}/restraints.py (96%) diff --git a/docs/changelog.rst b/docs/changelog.rst index 0dbe55b7..e879059a 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -39,6 +39,8 @@ Unreleased - Moved the riding-hydrogen map from ``torchref.restraints.hydrogen_topology`` to ``torchref.topology.riding``, alongside the generation path it is the heavy-atom-only alternative to - Added ``Topology.subset`` and ``copy``, plus the same on ``EdgeBlock``, ``ResidueGraph`` and ``AtomGraph``: a subset reindexes the surviving edges rather than re-reading the CIFs and re-matching the templates. An edge is dropped as soon as any of its atoms is, and a residue left with no atoms goes along with its links - Fixed residues distinguished only by an insertion code losing their restraints: the builders group on ``(chain, resseq)``, so 100 and 100A merge into one residue whose atom names collide and only the first keeps any intra-residue geometry +- Moved the restraint orchestrator from ``torchref.restraints.restraints`` to ``torchref.topology.restraints`` and renamed ``RestraintsNew`` to ``Restraints``. ``torchref.restraints`` is now the data layer only -- monomer library, CIF readers, ``chem_mod`` records, matchers, spatial search -- and what is built from that data lives in ``torchref.topology`` +- ``_lookup_link_atom`` moved to ``torchref.topology.build``, which was importing it back out of the restraints module Version 0.6.4 diff --git a/tests/conftest.py b/tests/conftest.py index 84715a84..b8b7178e 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -596,7 +596,7 @@ def initialized_scaler(model_and_data): @pytest.fixture def model_with_restraints(loaded_model): """Fixture providing model with built restraints.""" - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints restraints = Restraints( pdb=loaded_model.pdb, diff --git a/tests/functional/test_restraints_functional.py b/tests/functional/test_restraints_functional.py index 815ce836..76534421 100644 --- a/tests/functional/test_restraints_functional.py +++ b/tests/functional/test_restraints_functional.py @@ -15,7 +15,7 @@ class TestRestraintsBuildingFunctional: def test_build_restraints_from_cif(self, sample_cif_file): """Test building restraints from a real CIF file.""" from torchref.model.model import Model - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints model = Model() model.load_cif(str(sample_cif_file)) @@ -36,7 +36,7 @@ def test_build_restraints_from_cif(self, sample_cif_file): def test_bond_restraints_built(self, sample_cif_file): """Test that bond restraints are built correctly.""" from torchref.model.model import Model - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints model = Model() model.load_cif(str(sample_cif_file)) @@ -77,7 +77,7 @@ def test_bond_restraints_built(self, sample_cif_file): def test_angle_restraints_built(self, sample_cif_file): """Test that angle restraints are built correctly.""" from torchref.model.model import Model - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints model = Model() model.load_cif(str(sample_cif_file)) @@ -111,7 +111,7 @@ def test_angle_restraints_built(self, sample_cif_file): def test_torsion_restraints_built(self, sample_cif_file): """Test that torsion restraints are built correctly.""" from torchref.model.model import Model - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints model = Model() model.load_cif(str(sample_cif_file)) @@ -143,7 +143,7 @@ def test_torsion_restraints_built(self, sample_cif_file): def test_plane_restraints_built(self, sample_cif_file): """Test that plane restraints are built correctly.""" from torchref.model.model import Model - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints model = Model() model.load_cif(str(sample_cif_file)) @@ -178,7 +178,7 @@ class TestRestraintsDeviationsFunctional: def test_bond_deviations(self, sample_cif_file): """Test computing bond length deviations.""" from torchref.model.model import Model - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints model = Model() model.load_cif(str(sample_cif_file)) @@ -205,7 +205,7 @@ def test_bond_deviations(self, sample_cif_file): def test_angle_deviations(self, sample_cif_file): """Test computing angle deviations.""" from torchref.model.model import Model - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints model = Model() model.load_cif(str(sample_cif_file)) @@ -234,7 +234,7 @@ class TestRestraintsMultipleStructures: def test_restraints_multiple_cif_files(self, cif_dir): """Test building restraints for multiple CIF files.""" from torchref.model.model import Model - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints cif_files = list(cif_dir.glob("*.cif"))[:3] # First 3 structures @@ -262,7 +262,7 @@ class TestRestraintsDeviceHandling: def test_restraints_device_movement(self, sample_cif_file, cpu_device): """Test moving restraints to different devices.""" from torchref.model.model import Model - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints model = Model(device=cpu_device) model.load_cif(str(sample_cif_file)) @@ -288,7 +288,7 @@ class TestRestraintsCIFParsing: def test_cif_dict_loaded(self, sample_cif_file): """Test that CIF dictionary is loaded correctly.""" from torchref.model.model import Model - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints model = Model() model.load_cif(str(sample_cif_file)) @@ -314,7 +314,7 @@ def test_cif_dict_loaded(self, sample_cif_file): def test_unique_residues_detected(self, sample_cif_file): """Test that unique residues are detected from model.""" from torchref.model.model import Model - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints model = Model() model.load_cif(str(sample_cif_file)) diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index 32198d04..a5eddc1a 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -392,7 +392,7 @@ class TargetDeviceCase: "ModelCollection": "needs several loaded models", "_SharedMixedModel": "internal view owned by ModelCollection", "Scaler": "needs a loaded model + data; covered in integration", - "RestraintsNew": "needs a model + monomer library", + "Restraints": "needs a model + monomer library", "FrenchWilson": "needs loaded intensities", "DatasetCollection": "needs several loaded datasets", "FcalcDataset": "needs computed structure factors", diff --git a/tests/integration/test_refinement_pipeline.py b/tests/integration/test_refinement_pipeline.py index 4f385578..94cb9258 100644 --- a/tests/integration/test_refinement_pipeline.py +++ b/tests/integration/test_refinement_pipeline.py @@ -44,7 +44,7 @@ def test_setup_refinement_components(self, sample_structure_pair): def test_restraints_from_model(self, sample_cif_file): """Test building restraints from a loaded model.""" from torchref.model.model import Model - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints model = Model() model.load_cif(str(sample_cif_file)) diff --git a/tests/unit/restraints/test_restraints.py b/tests/unit/restraints/test_restraints.py index 1ec929cf..b6a037fa 100644 --- a/tests/unit/restraints/test_restraints.py +++ b/tests/unit/restraints/test_restraints.py @@ -17,7 +17,7 @@ class TestRestraintsInitialization: @pytest.mark.unit def test_restraints_empty_init(self): """Test Restraints can be initialized empty.""" - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints restraints = Restraints() @@ -26,7 +26,7 @@ def test_restraints_empty_init(self): @pytest.mark.unit def test_restraints_is_nn_module(self): """Restraints should be a nn.Module.""" - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints restraints = Restraints() @@ -35,7 +35,7 @@ def test_restraints_is_nn_module(self): @pytest.mark.unit def test_restraints_verbose_setting(self): """Test verbosity setting.""" - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints restraints = Restraints(verbose=2) diff --git a/tests/unit/topology/test_storage.py b/tests/unit/topology/test_storage.py index 2a86f79d..6264195a 100644 --- a/tests/unit/topology/test_storage.py +++ b/tests/unit/topology/test_storage.py @@ -184,7 +184,7 @@ def test_rebuilding_entries_reslices_onto_the_current_blocks(restraints): This is the operation ``_apply`` and ``copy`` both rely on, and the one that has to stay cheap: it re-slices rather than recomputing anything. - ``RestraintsNew.copy`` is not exercised here because it cannot run at all -- it is + ``Restraints.copy`` is not exercised here because it cannot run at all -- it is ``deepcopy``, which walks the *borrowed* ``_xyz_fn`` wrapper, whose cache holds a graph-attached tensor once ``xyz()`` has been evaluated. Verified to fail identically at the commit before this change, so it is pre-existing rather than a diff --git a/torchref/__init__.py b/torchref/__init__.py index 0df11cb3..b11229e2 100644 --- a/torchref/__init__.py +++ b/torchref/__init__.py @@ -109,7 +109,8 @@ from torchref.symmetry import Cell, SpaceGroup, Symmetry # Restraints -# from torchref.restraints import Restraints # Initialized lazily due to monomer library download requirement +# torchref.topology.restraints.Restraints is not imported here: constructing it can +# trigger a monomer-library download, so it stays lazy. # Scaling from torchref.scaling import Scaler, SolventModel, ScalerBase diff --git a/torchref/model/model.py b/torchref/model/model.py index fd774801..33b3cd28 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -442,7 +442,7 @@ def set_restraints_cif(self, cif_path): return self def _build_restraints(self): - """Build and cache ``RestraintsNew`` over this model's DataFrame, wiring in + """Build and cache ``Restraints`` over this model's DataFrame, wiring in the live ``xyz`` / ``adp`` / ``vdw_radii`` callables. """ if self._restraints is not None: @@ -454,12 +454,12 @@ def _build_restraints(self): "Load data first with load_pdb() or load_cif()." ) - from torchref.restraints.restraints import RestraintsNew + from torchref.topology.restraints import Restraints if self.ctx.verbose > 0: print("Building restraints...") - self._restraints = RestraintsNew( + self._restraints = Restraints( pdb=self.pdb, cif_path=self.ctx.cif_path, xyz_fn=self.xyz, diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index a6420b06..37751bc0 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -980,7 +980,7 @@ def extract_submodule_state(state_dict: dict, prefix: str) -> dict: scaler = Scaler(model, reflection_data, verbose=verbose, device=device) # Create Restraints with model (required for proper setup) - from torchref.restraints import Restraints + from torchref.topology.restraints import Restraints restraints = Restraints(model, verbose=verbose) diff --git a/torchref/restraints/__init__.py b/torchref/restraints/__init__.py index 55cb62eb..ed3ad796 100644 --- a/torchref/restraints/__init__.py +++ b/torchref/restraints/__init__.py @@ -1,18 +1,19 @@ -"""Geometry restraints (bonds, angles, torsions, planes, chirals, VDW contacts). +"""Where restraint *data* comes from: the monomer library and the CIF readers. -``Restraints`` (an alias of ``RestraintsNew``) builds and holds them all from CIF -dictionaries resolved through :func:`get_library_manager`; ``MONOMER_LIB_PATH`` -resolves lazily to that manager's ``monomer_dir``. Ideal values come from the CCP4 -Monomer Library (Long et al. 2017, Acta Cryst. D73, 112-122). +The dictionaries of ideal geometry, resolved through :func:`get_library_manager`, the +``chem_mod`` records that patch them when a link forms, the Numba matchers that map a +template onto the atoms present, the Ramachandran surfaces, and the spatial search +behind the non-bonded pair list. ``MONOMER_LIB_PATH`` resolves lazily to the manager's +``monomer_dir``. Ideal values come from the CCP4 Monomer Library (Long et al. 2017, +Acta Cryst. D73, 112-122). -The builder classes and the inter-residue builders are used across the package but -deliberately *not* re-exported here -- import them from their defining submodules. -Connectivity itself lives in :mod:`torchref.topology`, which is also where the -riding-hydrogen map moved to. +What is *built* from that data lives in :mod:`torchref.topology`: the connectivity, the +values layered over its edges, and :class:`~torchref.topology.restraints.Restraints`, +which orchestrates the two. Nothing here is re-exported -- import from the defining +submodule. """ from torchref.restraints.library import get_library_manager -from torchref.restraints.restraints import RestraintsNew as Restraints def __getattr__(name): @@ -23,7 +24,6 @@ def __getattr__(name): __all__ = [ - "Restraints", "MONOMER_LIB_PATH", "get_library_manager", ] diff --git a/torchref/topology/__init__.py b/torchref/topology/__init__.py index b9b9abd3..9fc6c49c 100644 --- a/torchref/topology/__init__.py +++ b/torchref/topology/__init__.py @@ -29,6 +29,7 @@ plan_hydrogens, ) from .residue_graph import ResidueGraph +from .restraints import Restraints from .riding import ( HydrogenTopology, build_h_candidate_pairs, @@ -41,6 +42,7 @@ __all__ = [ "Topology", + "Restraints", "ResidueGraph", "AtomGraph", "EdgeBlock", diff --git a/torchref/topology/build.py b/torchref/topology/build.py index bce5b695..6e4b0483 100644 --- a/torchref/topology/build.py +++ b/torchref/topology/build.py @@ -536,6 +536,47 @@ def _disulfide_edges( return out, values +def _lookup_link_atom( + pdb: pd.DataFrame, + chainid: str, + resseq: int, + icode: str, + resname: str, + name: str, + altloc: str, +): + """Resolve one ``LINK`` record's atom to a row of the atom table, or None. + + Matches on ``(chainid, resseq, icode, name)`` with ``resname`` as a tie-breaker. + Where a residue has alternative conformations the requested altloc wins, then the + blank one, then ``'A'``, then whatever is left -- a LINK naming a specific conformer + should reach that conformer, but one naming none should still resolve. + """ + sel = pdb[ + (pdb["chainid"].astype(str) == str(chainid)) + & (pdb["resseq"].astype(int) == int(resseq)) + & (pdb["icode"].astype(str) == str(icode)) + & (pdb["name"].astype(str).str.strip() == str(name).strip()) + ] + if len(sel) == 0: + return None + if resname: + tied = sel[sel["resname"].astype(str).str.strip() == str(resname).strip()] + if len(tied) > 0: + sel = tied + + if altloc: + for candidate in (altloc, ""): + hit = sel[sel["altloc"].astype(str) == candidate] + if len(hit) > 0: + return int(hit.iloc[0]["index"]) + for candidate in ("", "A"): + hit = sel[sel["altloc"].astype(str) == candidate] + if len(hit) > 0: + return int(hit.iloc[0]["index"]) + return int(sel.iloc[0]["index"]) + + def _link_record_edges( pdb: pd.DataFrame, links, @@ -544,10 +585,8 @@ def _link_record_edges( ) -> Tuple[np.ndarray, List[Tuple[int, int]], Dict[str, np.ndarray]]: """Bond edges for the accepted ``LINK`` records, and the atom pairs they join. - Atom resolution goes through ``RestraintsNew._lookup_link_atom`` so a record is - matched to a row exactly as the restraint builder matches it. A record duplicating - an auto-detected disulfide is dropped, since that link already contributed its - bond, angles and torsions. + A record duplicating an auto-detected disulfide is dropped, since that link already + contributed its bond, angles and torsions. Returns ------- @@ -559,8 +598,6 @@ def _link_record_edges( ``references`` from each record's ``length`` (1.5 A where blank or unusable) and a fixed ``sigmas`` of 0.02 A. """ - from torchref.restraints.restraints import RestraintsNew - if links is None or len(links) == 0: return np.zeros((0, 2), dtype=np.int64), [], {} @@ -573,7 +610,7 @@ def _link_record_edges( lengths: List[float] = [] n_unresolved = 0 for _, link in links.iterrows(): - idx1 = RestraintsNew._lookup_link_atom( + idx1 = _lookup_link_atom( pdb, chainid=link["chainid1"], resseq=int(link["resseq1"]), @@ -582,7 +619,7 @@ def _link_record_edges( name=link["name1"], altloc=link["altloc1"], ) - idx2 = RestraintsNew._lookup_link_atom( + idx2 = _lookup_link_atom( pdb, chainid=link["chainid2"], resseq=int(link["resseq2"]), diff --git a/torchref/restraints/restraints.py b/torchref/topology/restraints.py similarity index 96% rename from torchref/restraints/restraints.py rename to torchref/topology/restraints.py index a45fe074..65fd6a2c 100644 --- a/torchref/restraints/restraints.py +++ b/torchref/topology/restraints.py @@ -1,13 +1,22 @@ -"""Restraints handler for crystallographic model refinement. +"""The restraint layer over a topology, and what it takes to build one. -Provides :class:`RestraintsNew`, which holds the geometry restraints (bonds, angles, -torsions, planes, chirals, VDW) for one structure. The geometry half is a -:class:`~torchref.topology.topology.Topology` plus the ideal values layered over its -edges; the non-bonded pair list is separate, because it is distance-derived and rebuilt -as the model moves. +:class:`Restraints` is the orchestrator. Given an atom table it resolves the monomer +dictionaries, builds the :class:`~torchref.topology.topology.Topology`, layers the ideal +values over its edges, derives the non-bonded pair list, and exposes the whole thing as +``restraints[edge_type][origin][property]`` -- three dict lookups into a mapping +assembled once, because the geometry targets read it on every iteration. -Decoupled from :class:`~torchref.model.Model`: it accepts a pdb DataFrame and callables -for coordinates, ADPs and VDW radii. +Three kinds of thing live here, and only the first is really connectivity: + +* the topology and the values keyed to its edges, which are constants for the lifetime + of an atom set; +* the non-bonded pair list, which is distance-derived and rebuilt as the model moves, so + it is held apart from the rest; +* the Ramachandran map, a residue-level product of the same build. + +Deliberately decoupled from :class:`~torchref.model.Model`: it takes an atom table plus +callables for coordinates, ADPs and van der Waals radii, so it can be built and tested +without one. """ from typing import Callable @@ -28,7 +37,7 @@ -class RestraintsNew(DeviceMixin, DebugMixin, Module): +class Restraints(DeviceMixin, DebugMixin, Module): """ Restraints handler for crystallographic model refinement. @@ -394,51 +403,13 @@ def build_restraints(self): self.to(target_device) except Exception as e: - self.debug_on_error(e, context="RestraintsNew.build_restraints") + self.debug_on_error(e, context="Restraints.build_restraints") raise - @staticmethod - def _lookup_link_atom( - pdb: pd.DataFrame, - chainid: str, - resseq: int, - icode: str, - resname: str, - name: str, - altloc: str, - ): - """Resolve a LINK atom record to a row index in the model pdb, or None. - - Matches (chainid, resseq, icode, name), resname as tie-breaker; altloc - preference is requested, then blank, then 'A', then any. - """ - sel = pdb[ - (pdb["chainid"].astype(str) == str(chainid)) - & (pdb["resseq"].astype(int) == int(resseq)) - & (pdb["icode"].astype(str) == str(icode)) - & (pdb["name"].astype(str).str.strip() == str(name).strip()) - ] - if len(sel) == 0: - return None - if resname: - tied = sel[sel["resname"].astype(str).str.strip() == str(resname).strip()] - if len(tied) > 0: - sel = tied - - if altloc: - for cand in (altloc, ""): - hit = sel[sel["altloc"].astype(str) == cand] - if len(hit) > 0: - return int(hit.iloc[0]["index"]) - for cand in ("", "A"): - hit = sel[sel["altloc"].astype(str) == cand] - if len(hit) > 0: - return int(hit.iloc[0]["index"]) - return int(sel.iloc[0]["index"]) def _find_nearby_pairs_spatial_hash(self, xyz, cutoff=6.0): @@ -1208,7 +1179,7 @@ def get_count(rtype, origin): n_bonds_peptide = get_count("bond", "peptide") return ( - f"RestraintsNew(bonds={n_bonds}, angles={n_angles}, " + f"Restraints(bonds={n_bonds}, angles={n_angles}, " f"torsions={n_torsions}, peptide_bonds={n_bonds_peptide})" ) @@ -1249,7 +1220,7 @@ def copy(self): Returns ------- - RestraintsNew + Restraints """ import copy From a9f8a371340044f4e276a44f3e4713a54969f729 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 16:29:27 +0200 Subject: [PATCH 067/250] Fold the restraint data layer into topology and drop the package torchref.restraints is gone. Every one of its seven modules had exactly one kind of importer left -- something inside torchref.topology -- so the package had already become a private implementation detail of it. Moving them in recognises that rather than imposing anything new. Placed by what each one is, not under a single label, because they are not alike: topology/monomer/library.py resolve, download and cache the CCP4 library topology/monomer/cif.py read a dictionary into its sections topology/monomer/modifications.py the chem_mod records that patch a template topology/builders.py template-to-atom matching topology/builders_numba.py the njit matchers topology/nonbonded.py the spatial search behind the VDW pair list topology/ramachandran.py the NLL surfaces Only one rename beyond the relocation: restraints_helper.py becomes monomer/cif.py, since the old name described where it sat rather than what it did. Tests followed, tests/unit/restraints to tests/unit/monomer. Two modules computed a data path by counting levels up from __file__, and library.py gained a level in the move. Left alone it would have pointed at torchref/topology/data/, where nothing lives -- and it would not have raised: the bundled monomer library would simply have appeared absent and every component would have been re-downloaded. Both paths are now pinned with an explicit parents[n], and both were checked afterwards: the bundled library resolves, ALA loads, and the Ramachandran surfaces come back at (6, 360, 360). Also fixes neighbor_search annotating three locals with List without importing it. Harmless at runtime, since annotations are not evaluated there, but wrong. Suite 1938 passed, unchanged by the move. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- docs/changelog.rst | 4 ++- .../unit/io/test_multicomponent_restraints.py | 6 ++-- tests/unit/monomer/__init__.py | 1 + .../test_link_modifications.py | 4 +-- .../test_restraints.py | 2 +- tests/unit/restraints/__init__.py | 1 - tests/unit/topology/test_insertion_codes.py | 4 +-- torchref/restraints/__init__.py | 29 ------------------- torchref/topology/build.py | 6 ++-- .../builders_fast.py => topology/builders.py} | 10 +++---- .../builders_numba.py | 0 torchref/topology/monomer/__init__.py | 13 +++++++++ .../monomer/cif.py} | 8 ++--- .../monomer}/library.py | 7 ++++- .../monomer}/modifications.py | 6 ++-- .../nonbonded.py} | 2 +- .../{restraints => topology}/ramachandran.py | 2 +- torchref/topology/restraints.py | 6 ++-- torchref/topology/riding.py | 4 +-- torchref/topology/templates.py | 4 +-- 20 files changed, 55 insertions(+), 64 deletions(-) create mode 100644 tests/unit/monomer/__init__.py rename tests/unit/{restraints => monomer}/test_link_modifications.py (99%) rename tests/unit/{restraints => monomer}/test_restraints.py (99%) delete mode 100644 tests/unit/restraints/__init__.py delete mode 100644 torchref/restraints/__init__.py rename torchref/{restraints/builders_fast.py => topology/builders.py} (99%) rename torchref/{restraints => topology}/builders_numba.py (100%) create mode 100644 torchref/topology/monomer/__init__.py rename torchref/{restraints/restraints_helper.py => topology/monomer/cif.py} (96%) rename torchref/{restraints => topology/monomer}/library.py (96%) rename torchref/{restraints => topology/monomer}/modifications.py (98%) rename torchref/{restraints/neighbor_search.py => topology/nonbonded.py} (99%) rename torchref/{restraints => topology}/ramachandran.py (96%) diff --git a/docs/changelog.rst b/docs/changelog.rst index e879059a..3ff8c152 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -39,8 +39,10 @@ Unreleased - Moved the riding-hydrogen map from ``torchref.restraints.hydrogen_topology`` to ``torchref.topology.riding``, alongside the generation path it is the heavy-atom-only alternative to - Added ``Topology.subset`` and ``copy``, plus the same on ``EdgeBlock``, ``ResidueGraph`` and ``AtomGraph``: a subset reindexes the surviving edges rather than re-reading the CIFs and re-matching the templates. An edge is dropped as soon as any of its atoms is, and a residue left with no atoms goes along with its links - Fixed residues distinguished only by an insertion code losing their restraints: the builders group on ``(chain, resseq)``, so 100 and 100A merge into one residue whose atom names collide and only the first keeps any intra-residue geometry -- Moved the restraint orchestrator from ``torchref.restraints.restraints`` to ``torchref.topology.restraints`` and renamed ``RestraintsNew`` to ``Restraints``. ``torchref.restraints`` is now the data layer only -- monomer library, CIF readers, ``chem_mod`` records, matchers, spatial search -- and what is built from that data lives in ``torchref.topology`` +- Moved the restraint orchestrator from ``torchref.restraints.restraints`` to ``torchref.topology.restraints`` and renamed ``RestraintsNew`` to ``Restraints`` - ``_lookup_link_atom`` moved to ``torchref.topology.build``, which was importing it back out of the restraints module +- Removed ``torchref.restraints``. Its seven modules moved into ``torchref.topology``, where every one of their importers already lived: the monomer library, CIF reading and ``chem_mod`` patches to ``topology.monomer``, the template matchers to ``topology.builders`` and ``topology.builders_numba``, the non-bonded spatial search to ``topology.nonbonded``, and the Ramachandran surfaces to ``topology.ramachandran`` +- Fixed ``neighbor_search`` annotating locals with ``List`` without importing it Version 0.6.4 diff --git a/tests/unit/io/test_multicomponent_restraints.py b/tests/unit/io/test_multicomponent_restraints.py index 780d3414..65d344ab 100644 --- a/tests/unit/io/test_multicomponent_restraints.py +++ b/tests/unit/io/test_multicomponent_restraints.py @@ -18,8 +18,8 @@ import pytest from torchref.io.cif_readers import RestraintCIFReader -from torchref.restraints.library import get_library_manager -from torchref.restraints.restraints_helper import ( +from torchref.topology.monomer.library import get_library_manager +from torchref.topology.monomer.cif import ( split_data_blocks, validate_restraint_data, ) @@ -230,7 +230,7 @@ class TestChiralitySpellings: """The CCP4 library writes both ``positive`` and the truncated ``positiv``.""" def test_short_spellings_are_not_dropped(self): - from torchref.restraints.builders_fast import PreprocessedCIF + from torchref.topology.builders import PreprocessedCIF chirals = pd.DataFrame( { diff --git a/tests/unit/monomer/__init__.py b/tests/unit/monomer/__init__.py new file mode 100644 index 00000000..b7a88075 --- /dev/null +++ b/tests/unit/monomer/__init__.py @@ -0,0 +1 @@ +"""Unit tests for the monomer-dictionary layer under torchref.topology.""" diff --git a/tests/unit/restraints/test_link_modifications.py b/tests/unit/monomer/test_link_modifications.py similarity index 99% rename from tests/unit/restraints/test_link_modifications.py rename to tests/unit/monomer/test_link_modifications.py index 67345a21..63b73854 100644 --- a/tests/unit/restraints/test_link_modifications.py +++ b/tests/unit/monomer/test_link_modifications.py @@ -11,12 +11,12 @@ import pandas as pd import pytest -from torchref.restraints.modifications import ( +from torchref.topology.monomer.modifications import ( apply_modifications, link_modifications, read_mod_definitions, ) -from torchref.restraints.restraints_helper import read_link_definitions +from torchref.topology.monomer.cif import read_link_definitions TOL = 1e-3 diff --git a/tests/unit/restraints/test_restraints.py b/tests/unit/monomer/test_restraints.py similarity index 99% rename from tests/unit/restraints/test_restraints.py rename to tests/unit/monomer/test_restraints.py index b6a037fa..97ea6c1a 100644 --- a/tests/unit/restraints/test_restraints.py +++ b/tests/unit/monomer/test_restraints.py @@ -1,5 +1,5 @@ """ -Unit tests for torchref.restraints.restraints +Unit tests for torchref.topology.restraints Tests the Restraints class for geometry restraints. Note: Unit tests use mock data, not real file I/O. diff --git a/tests/unit/restraints/__init__.py b/tests/unit/restraints/__init__.py deleted file mode 100644 index 571dfe85..00000000 --- a/tests/unit/restraints/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Unit tests for torchref.restraints module.""" diff --git a/tests/unit/topology/test_insertion_codes.py b/tests/unit/topology/test_insertion_codes.py index f402b7da..0e5aec1b 100644 --- a/tests/unit/topology/test_insertion_codes.py +++ b/tests/unit/topology/test_insertion_codes.py @@ -112,7 +112,7 @@ def test_the_builders_merge_them(inserted): ``PreprocessedPDB`` groups on ``(chain, resseq)``, so the three residues become one with three sets of backbone atom names. """ - from torchref.restraints.builders_fast import PreprocessedPDB + from torchref.topology.builders import PreprocessedPDB _, _, expected, model = inserted preprocessed = PreprocessedPDB(model.pdb) @@ -138,7 +138,7 @@ def _legacy_intra_bonds(model, restraints): """ import torch - from torchref.restraints.builders_fast import BondRestraintBuilder + from torchref.topology.builders import BondRestraintBuilder built = BondRestraintBuilder(verbose=0).build( model.pdb, restraints.cif_dict, torch.device("cpu") diff --git a/torchref/restraints/__init__.py b/torchref/restraints/__init__.py deleted file mode 100644 index ed3ad796..00000000 --- a/torchref/restraints/__init__.py +++ /dev/null @@ -1,29 +0,0 @@ -"""Where restraint *data* comes from: the monomer library and the CIF readers. - -The dictionaries of ideal geometry, resolved through :func:`get_library_manager`, the -``chem_mod`` records that patch them when a link forms, the Numba matchers that map a -template onto the atoms present, the Ramachandran surfaces, and the spatial search -behind the non-bonded pair list. ``MONOMER_LIB_PATH`` resolves lazily to the manager's -``monomer_dir``. Ideal values come from the CCP4 Monomer Library (Long et al. 2017, -Acta Cryst. D73, 112-122). - -What is *built* from that data lives in :mod:`torchref.topology`: the connectivity, the -values layered over its edges, and :class:`~torchref.topology.restraints.Restraints`, -which orchestrates the two. Nothing here is re-exported -- import from the defining -submodule. -""" - -from torchref.restraints.library import get_library_manager - - -def __getattr__(name): - """Lazy access to MONOMER_LIB_PATH for backward compatibility.""" - if name == "MONOMER_LIB_PATH": - return get_library_manager().monomer_dir - raise AttributeError(f"module {__name__!r} has no attribute {name!r}") - - -__all__ = [ - "MONOMER_LIB_PATH", - "get_library_manager", -] diff --git a/torchref/topology/build.py b/torchref/topology/build.py index 6e4b0483..12febd63 100644 --- a/torchref/topology/build.py +++ b/torchref/topology/build.py @@ -1,7 +1,7 @@ """Assemble a :class:`~torchref.topology.topology.Topology` from an atom table. Intra-residue edges are matched here, template by template, through the Numba matchers -in :mod:`torchref.restraints.builders_numba`. Inter-residue edges come from the +in :mod:`torchref.topology.builders_numba`. Inter-residue edges come from the ``InterResidue*Builder`` classes, which already encode the link geometry and are reused rather than reimplemented. """ @@ -12,14 +12,14 @@ import pandas as pd import torch -from torchref.restraints.builders_fast import ( +from torchref.topology.builders import ( InterResidueAngleBuilder, InterResidueBondBuilder, InterResiduePlaneBuilder, InterResidueTorsionBuilder, PreprocessedCIF, ) -from torchref.restraints.builders_numba import ( +from torchref.topology.builders_numba import ( match_angles_numba, match_bonds_numba, match_chirals_numba, diff --git a/torchref/restraints/builders_fast.py b/torchref/topology/builders.py similarity index 99% rename from torchref/restraints/builders_fast.py rename to torchref/topology/builders.py index 8d074b23..df7c27ab 100644 --- a/torchref/restraints/builders_fast.py +++ b/torchref/topology/builders.py @@ -8,8 +8,8 @@ calls and emits nothing until ``finalize()`` (``finalize_disulfide()`` on the torsion builder). -Nothing here is re-exported at the ``torchref.restraints`` level; import from -``torchref.restraints.builders_fast``. +Nothing here is re-exported at the package level; import from +``torchref.topology.builders``. """ from abc import ABC, abstractmethod @@ -22,7 +22,7 @@ from torchref.config import get_float_dtype # Import the Numba-accelerated matching functions -from torchref.restraints.builders_numba import ( +from torchref.topology.builders_numba import ( match_angles_numba, match_bonds_numba, match_chirals_numba, @@ -133,7 +133,7 @@ def residue_keys( mapping : mapping, optional ``{(chain_id, resseq): key}`` overriding the residue name for those residues -- how a linked residue is pointed at a modified copy of its - component (see :mod:`torchref.restraints.modifications`). Residues + component (see :mod:`torchref.topology.monomer.modifications`). Residues absent from it, and every residue when this is None, key on their own residue name, which is the unmodified behaviour. @@ -1858,7 +1858,7 @@ def build( torsions = link_data.torsions n_torsions = len(torsions["atom1"]) - from torchref.restraints.ramachandran import classify_residue + from torchref.topology.ramachandran import classify_residue for res_i_idx, res_next_idx in pairs: resname_i = pp_pdb.residue_resnames[res_i_idx] diff --git a/torchref/restraints/builders_numba.py b/torchref/topology/builders_numba.py similarity index 100% rename from torchref/restraints/builders_numba.py rename to torchref/topology/builders_numba.py diff --git a/torchref/topology/monomer/__init__.py b/torchref/topology/monomer/__init__.py new file mode 100644 index 00000000..b655740b --- /dev/null +++ b/torchref/topology/monomer/__init__.py @@ -0,0 +1,13 @@ +"""Monomer dictionaries: finding them, reading them, and patching them. + +Everything that turns files on disk into the ideal geometry a template carries. +:mod:`library` resolves and caches the CCP4 Monomer Library, fetching a component on +demand; :mod:`cif` reads a dictionary into DataFrames per section; :mod:`modifications` +applies the ``chem_mod`` records that change a template when a link forms. + +This is the data source. What is built from it -- the connectivity, the values over its +edges -- is the rest of :mod:`torchref.topology`. Import from the defining submodule +rather than from here. + +Reference: Long, F., et al. (2017). AceDRG. Acta Cryst. D73, 112-122. +""" diff --git a/torchref/restraints/restraints_helper.py b/torchref/topology/monomer/cif.py similarity index 96% rename from torchref/restraints/restraints_helper.py rename to torchref/topology/monomer/cif.py index b2d2276b..95731424 100644 --- a/torchref/restraints/restraints_helper.py +++ b/torchref/topology/monomer/cif.py @@ -94,7 +94,7 @@ def find_cif_file_in_library(resname): Delegates to :meth:`MonomerLibraryManager.get_cif_file`, whose last resort is an on-demand download. """ - from torchref.restraints.library import get_library_manager + from torchref.topology.monomer.library import get_library_manager return get_library_manager().get_cif_file(resname) @@ -128,14 +128,14 @@ def read_library_blocks(): """Return ``mon_lib_list.cif`` split into ``{block_name: block_text}``. Shared by :func:`read_link_definitions` and - :func:`~torchref.restraints.modifications.read_mod_definitions`, which read + :func:`~torchref.topology.monomer.modifications.read_mod_definitions`, which read disjoint parts of the same 4 MB file. Warnings -------- Cached process-wide; the returned dict is shared, so do not mutate it. """ - from torchref.restraints.library import get_library_manager + from torchref.topology.monomer.library import get_library_manager path = str(get_library_manager().get_link_definitions_path()) with open(path) as handle: @@ -155,7 +155,7 @@ def read_link_definitions(): link_list : DataFrame or None The ``chem_link`` table, or None if the file has no ``link_list`` block. Its ``mod_id_1``/``mod_id_2`` columns name the modifications each link - applies to its partners -- see :mod:`torchref.restraints.modifications`. + applies to its partners -- see :mod:`torchref.topology.monomer.modifications`. Warnings -------- diff --git a/torchref/restraints/library.py b/torchref/topology/monomer/library.py similarity index 96% rename from torchref/restraints/library.py rename to torchref/topology/monomer/library.py index e8375f95..0eea7309 100644 --- a/torchref/restraints/library.py +++ b/torchref/topology/monomer/library.py @@ -26,7 +26,12 @@ ) # Bundled package data location -_BUNDLED_PATH = Path(__file__).parent.parent / "data" / "monomer_library" +# Three levels up, not two: this module sits at torchref/topology/monomer/, so the +# package root is its great-grandparent. Computing it by depth is fragile, which is +# why the level is spelled out rather than left to be counted. +_BUNDLED_PATH = ( + Path(__file__).resolve().parents[2] / "data" / "monomer_library" +) # Legacy external monomer library path _LEGACY_PATH = ROOT_TORCHREF / "external_monomer_library" diff --git a/torchref/restraints/modifications.py b/torchref/topology/monomer/modifications.py similarity index 98% rename from torchref/restraints/modifications.py rename to torchref/topology/monomer/modifications.py index f9b10f02..172e8276 100644 --- a/torchref/restraints/modifications.py +++ b/torchref/topology/monomer/modifications.py @@ -24,7 +24,7 @@ import pandas as pd -from torchref.restraints.restraints_helper import read_library_blocks +from torchref.topology.monomer.cif import read_library_blocks #: CIF category -> section name, matching :func:`read_link_definitions`. _CATEGORY_MAP = { @@ -167,7 +167,7 @@ def link_modifications(link_list: Optional[pd.DataFrame]) -> Dict[str, Tuple]: ---------- link_list : pandas.DataFrame or None The ``chem_link`` table returned by - :func:`~torchref.restraints.restraints_helper.read_link_definitions`. + :func:`~torchref.topology.monomer.cif.read_link_definitions`. Returns ------- @@ -200,7 +200,7 @@ def apply_modifications( ---------- comp : mapping of str to pandas.DataFrame One component's restraints, as produced by - :func:`~torchref.restraints.restraints_helper.read_cif`. + :func:`~torchref.topology.monomer.cif.read_cif`. mod_ids : sequence of str Modification IDs to apply, in order. Unknown IDs are ignored. mod_dict : mapping diff --git a/torchref/restraints/neighbor_search.py b/torchref/topology/nonbonded.py similarity index 99% rename from torchref/restraints/neighbor_search.py rename to torchref/topology/nonbonded.py index 9204ca4f..a8246b08 100644 --- a/torchref/restraints/neighbor_search.py +++ b/torchref/topology/nonbonded.py @@ -10,7 +10,7 @@ the input coordinates live on (CPU or GPU). """ -from typing import TYPE_CHECKING, Dict, Optional, Set, Tuple +from typing import TYPE_CHECKING, Dict, List, Optional, Set, Tuple import numpy as np import torch diff --git a/torchref/restraints/ramachandran.py b/torchref/topology/ramachandran.py similarity index 96% rename from torchref/restraints/ramachandran.py rename to torchref/topology/ramachandran.py index fc1cf691..9d45fc43 100644 --- a/torchref/restraints/ramachandran.py +++ b/torchref/topology/ramachandran.py @@ -33,7 +33,7 @@ #: Number of 1° bins per axis, covering [-180, +180). GRID_SIZE = 360 -_DATA_FILE = Path(__file__).resolve().parent.parent / "data" / "rama_nll_surfaces.pt" +_DATA_FILE = Path(__file__).resolve().parents[1] / "data" / "rama_nll_surfaces.pt" def load_nll_surfaces(device: torch.device) -> torch.Tensor: diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index 65fd6a2c..76c217f0 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -26,7 +26,7 @@ import torch from torch.nn import Module -from torchref.restraints.restraints_helper import ( +from torchref.topology.monomer.cif import ( find_cif_file_in_library, read_cif, read_link_definitions, @@ -357,7 +357,7 @@ def expand_altloc(self, residue): def _load_rama_surfaces(self, device: torch.device): """Load pre-computed Ramachandran NLL surfaces as a buffer.""" - from torchref.restraints.ramachandran import load_nll_surfaces + from torchref.topology.ramachandran import load_nll_surfaces surfaces = load_nll_surfaces(device) self.register_buffer("_rama_surfaces", surfaces) @@ -782,7 +782,7 @@ def vdw_radii_cpu(): sg_cpu = None if has_symmetry: - from torchref.restraints.neighbor_search import build_vdw_restraints_gpu + from torchref.topology.nonbonded import build_vdw_restraints_gpu exclusions = self.topology.atoms.exclusions_from_restraint_edges() self._vdw = build_vdw_restraints_gpu( diff --git a/torchref/topology/riding.py b/torchref/topology/riding.py index 46271c2a..588cf6f8 100644 --- a/torchref/topology/riding.py +++ b/torchref/topology/riding.py @@ -181,7 +181,7 @@ def _load_cif_hydrogen_info(pdb, verbose: int = 0) -> Dict: ``heavy_coords``, ``h_names``, ``h_coords``, ``parent_map``, ``ideal_bl`` and ``heavy_neighbor_map``. """ - from torchref.restraints.library import MonomerLibraryManager + from torchref.topology.monomer.library import MonomerLibraryManager lib = MonomerLibraryManager(verbose=0) cache = _TEMPLATE_CACHE @@ -636,7 +636,7 @@ def place_riding_hydrogens( xyz_h : (N_h, 3) float tensor, differentiable w.r.t. xyz_heavy """ # Function-local, matching every other target dispatch site: importing the gate at - # module scope would pull ``torchref.base.targets`` into ``torchref.restraints``. + # module scope would pull ``torchref.base.targets`` into ``torchref.topology``. from torchref.base.targets._dispatch import use_triton N_h = topo.h_parent_idx.shape[0] diff --git a/torchref/topology/templates.py b/torchref/topology/templates.py index cbfcc62b..afb44aea 100644 --- a/torchref/topology/templates.py +++ b/torchref/topology/templates.py @@ -15,7 +15,7 @@ import numpy as np -from torchref.restraints.modifications import ( +from torchref.topology.monomer.modifications import ( apply_modifications, link_modifications, read_mod_definitions, @@ -41,7 +41,7 @@ def resolve_template_keys( Restraint dictionary keyed by residue name. link_list : pandas.DataFrame or None Link-type definitions, as - :func:`~torchref.restraints.restraints_helper.read_link_definitions` returns + :func:`~torchref.topology.monomer.cif.read_link_definitions` returns them. None disables patching. verbose : int, default 0 Verbosity level. From 112c7eb7eed3372eb2e51ce6fd8510b5cef0e949 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 16:39:34 +0200 Subject: [PATCH 068/250] Share one Gaussian and one sigma floor between the observables The intensity target carried its own copy of the amplitude Gaussian, its own sigma floor and its own log(2*pi). Factor them: `gaussian_per_refl` is the Gaussian on any real observable, `intensity_var_from_sigma_obs` its variance builder, `floor_sigma_obs` the shared median clamp. The absolute variance floor is now opt-out. It is dimensionally arbitrary on a variance, and the two copies disagreed about it -- 1e-10 rescales the objective by several-fold on any dataset whose sigmas fall below ~1e-5, which the amplitude path applied and the intensity path did not. `floor_sigma_obs` also takes an explicit floor, so a caller with both a `forward` and a `residuals` can keep them consistent: a floor derived from a median depends on which reflections are in the array. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- torchref/base/targets/xray_likelihoods.py | 114 +++++++++++++++++- .../targets/collection/intensity.py | 38 +++--- 2 files changed, 123 insertions(+), 29 deletions(-) diff --git a/torchref/base/targets/xray_likelihoods.py b/torchref/base/targets/xray_likelihoods.py index e385eeb7..e70c263d 100644 --- a/torchref/base/targets/xray_likelihoods.py +++ b/torchref/base/targets/xray_likelihoods.py @@ -22,6 +22,15 @@ indistinguishable from a change of x-ray weight. ``sigma_obs**2`` needs **no** conversion: it is already an amplitude variance, a 1-DOF error on a measured amplitude. +**The observable is a third axis, and it is only the variance and the mean that carry it.** +:func:`gaussian_per_refl` is a Gaussian on any real observable; :func:`nll_per_refl` is that +same function on amplitudes, and the intensity rows are it on intensities. Only the variance +builder has to know which: :func:`amplitude_var_from_sigma_obs` vs +:func:`intensity_var_from_sigma_obs`, and confusing them is wrong by ``2|F|`` -- which is +resolution-dependent, so it presents as a scale or B error rather than as a bug. Rice has no +intensity twin: Rice and the folded normal are distributions *of an amplitude*, and the +intensity analogue is the exponential / chi-square_1 Wilson distribution. + Do not pair ``sigma_obs`` with a Rice ``Sigma``: that asserts an isotropic *complex* error where ``sigma_obs`` carries no phase at all, and no regime makes it correct (it was tried, and was the worst of every target). Model error -- ``beta`` -- is what belongs in a Rice @@ -39,12 +48,59 @@ #: Floor on any variance before it reaches a division or a log. VAR_FLOOR = 1e-10 +#: Floor on a *measured* sigma, as a fraction of its median over the fitted subset. +#: Data-dependent rather than absolute, because the scale of a sigma is the scale of the +#: data. Merged intensities in particular are reported with ``sigma == 0`` rows. +SIGMA_FLOOR_FRAC = 1e-1 + +#: Backstop under :data:`SIGMA_FLOOR_FRAC` for the intensity builder, for the pathological +#: case of a median that is itself ~0. The amplitude builder deliberately has none -- see +#: :func:`floor_sigma_obs`. +SIGMA_FLOOR_ABS = 1e-12 + # ===================================================================== # Variance builders -- the axis that distinguishes the five targets # ===================================================================== +def floor_sigma_obs( + sigma: torch.Tensor, + mask: torch.Tensor = None, + abs_floor: float = 0.0, + floor=None, +) -> torch.Tensor: + """Clamp a measured sigma at :data:`SIGMA_FLOOR_FRAC` of its median. + + ``mask`` restricts the median to the fitted subset -- which matters when the unfitted + rows carry filler sigmas, as reindexed collection members do. ``mask=None`` takes the + median over everything. + + ``abs_floor`` is a backstop under the fractional floor. It defaults to **off** because + the amplitude builder shipped without one and has a Triton counterpart to stay + bit-identical to; the intensity builder passes :data:`SIGMA_FLOOR_ABS`. + + **Pass ``floor`` explicitly to make the result independent of which reflections are in + ``sigma``.** A median computed from the argument makes every per-reflection value + depend on the whole array, so the same reflection scores differently in a subset sum + than in a full-size residual -- measured at 0.09% on a work set and 1.8% on a free set + for intensities, whose sigmas span orders of magnitude. Callers that need the two to + agree (any target with both a ``forward`` and a ``residuals``) compute the floor once + from their own fitted subset and pass it here. + """ + if floor is None: + selected = sigma if mask is None else sigma[mask] + if selected.numel() == 0: + # No fitted reflections to take a median over. Any positive floor is arbitrary + # here; what matters is that it is finite, so a later log or division cannot + # produce a NaN that would poison the whole gradient. + return sigma.clamp(min=1e-6) + floor = torch.median(selected) * SIGMA_FLOOR_FRAC + if abs_floor > 0.0: + floor = torch.clamp(torch.as_tensor(floor), min=abs_floor) + return sigma.clamp(min=floor) + + def amplitude_var_from_sigma_obs(sigma: torch.Tensor) -> torch.Tensor: """Amplitude variance from the experimental sigma: ``clamp(sigma)**2``. @@ -53,8 +109,25 @@ def amplitude_var_from_sigma_obs(sigma: torch.Tensor) -> torch.Tensor: beta-derived builders below floor only at :data:`VAR_FLOOR`; the two conventions are deliberately not reconciled. """ - floor = torch.median(sigma) * 1e-1 - return torch.clamp(sigma, min=floor) ** 2 + return floor_sigma_obs(sigma) ** 2 + + +def intensity_var_from_sigma_obs( + sigma: torch.Tensor, mask: torch.Tensor = None, floor=None +) -> torch.Tensor: + """Intensity variance from the experimental ``sigma(I)``: ``clamp(sigma)**2``. + + The intensity twin of :func:`amplitude_var_from_sigma_obs`, and **not** interchangeable + with it: applying an amplitude sigma to an intensity residual is wrong by a factor of + ``2|F|``, which is resolution-dependent and so looks like a scale or B error rather + than like a mistake. + + Differs from the amplitude builder in taking a ``mask`` (merged collection members are + reindexed onto a common list, so the unfitted rows carry filler), an absolute backstop + at :data:`SIGMA_FLOOR_ABS`, and an explicit ``floor`` -- see :func:`floor_sigma_obs` on + why a caller with both a ``forward`` and a ``residuals`` must pass one. + """ + return floor_sigma_obs(sigma, mask, abs_floor=SIGMA_FLOOR_ABS, floor=floor) ** 2 def amplitude_var_from_complex( @@ -133,6 +206,36 @@ def _masked_sum(loss: torch.Tensor, mask: torch.Tensor = None) -> torch.Tensor: return (loss * mask).sum() +def gaussian_per_refl( + obs: torch.Tensor, + model: torch.Tensor, + var: torch.Tensor, + var_floor: float = VAR_FLOOR, +) -> torch.Tensor: + """Per-reflection Gaussian NLL on **any** real observable (NOT masked or summed). + + 0.5 * (obs - model)**2 / var + 0.5 * log(var) + 0.5 * log(2*pi) + + Observable-agnostic on purpose: ``var`` just has to be the variance of whatever + ``obs`` is. :func:`nll_per_refl` is this on amplitudes; the intensity rows are this on + intensities. Keeping one implementation is what stops the two drifting -- they were + separately written once, and the copies differed only in spelling ``log(sigma)`` + instead of ``0.5 * log(var)``. + + ``var_floor`` defaults to :data:`VAR_FLOOR` because the amplitude path has always + applied it. **Pass 0.0 when the variance builder has already floored the sigma**, as + :func:`intensity_var_from_sigma_obs` does. An absolute floor on a variance is + dimensionally arbitrary -- 1e-10 is a distortion, not a safeguard, on any dataset whose + sigmas are smaller than ~1e-5, and it silently reweights the whole objective by up to + the ratio of the two floors. Positivity is the builder's job; this is a backstop for + builders that do not do it. + """ + diff = obs - model + if var_floor > 0.0: + var = torch.clamp(var, min=var_floor) + return 0.5 * diff**2 / var + 0.5 * torch.log(var) + HALF_LOG_2PI + + def nll_per_refl( F_obs: torch.Tensor, F_calc: torch.Tensor, var: torch.Tensor ) -> torch.Tensor: @@ -143,10 +246,11 @@ def nll_per_refl( ``var`` is the **amplitude** variance. Build it with :func:`amplitude_var_from_sigma_obs` (``nll``) or :func:`amplitude_var_from_complex` (``nll_beta``) -- see the module docstring on why those are not interchangeable. + + The ``torch.abs`` is what makes this the amplitude entry point: callers pass a complex + or signed ``F_calc``. Everything else is :func:`gaussian_per_refl`. """ - diff = F_obs - torch.abs(F_calc) - var = torch.clamp(var, min=VAR_FLOOR) - return 0.5 * diff**2 / var + 0.5 * torch.log(var) + HALF_LOG_2PI + return gaussian_per_refl(F_obs, torch.abs(F_calc), var) def nll_math( diff --git a/torchref/refinement/targets/collection/intensity.py b/torchref/refinement/targets/collection/intensity.py index 574685fe..fa608e8f 100644 --- a/torchref/refinement/targets/collection/intensity.py +++ b/torchref/refinement/targets/collection/intensity.py @@ -20,10 +20,13 @@ from typing import TYPE_CHECKING, Dict, List -import numpy as np import torch from torchref.base.metrics.rfactor import rfactor_work_free +from torchref.base.targets.xray_likelihoods import ( + gaussian_per_refl, + intensity_var_from_sigma_obs, +) from torchref.utils.stats import ( VERBOSITY_DEBUG, VERBOSITY_ESSENTIAL, @@ -40,13 +43,6 @@ from torchref.scaling.scaler_base import ScalerBase -_LOG_2PI = float(np.log(2.0 * np.pi)) - -#: Floor on the intensity sigma, as a fraction of the median over the fitted subset. A -#: merged intensity sigma can be reported as zero; unfloored it would dominate the sum. -_SIGMA_FLOOR_FRAC = 0.1 - - class CollectionTwoMomentIntensityTarget(CollectionXrayTarget): """ Gaussian intensity likelihood at the two-moment forward model. @@ -216,13 +212,17 @@ def forward(self) -> torch.Tensor: obs = torch.where(valid, obs, torch.zeros_like(obs)) sigma = torch.where(valid, sigma, torch.ones_like(sigma)) - sigma = self._floor_sigma(sigma, mask) - + # The residual is formed and masked BEFORE the Gaussian, so a masked-out row + # contributes an exact zero rather than a value that merely gets multiplied by + # zero. That matters if the model is ever non-finite on an unfitted row: here the + # `where` discards it, whereas `nll * mask` would propagate NaN into the sum. + # Hence the Gaussian is evaluated at (residual, 0) rather than (obs, model). residual = torch.where(mask, obs - model, torch.zeros_like(obs)) - nll = ( - 0.5 * (residual / sigma) ** 2 - + torch.log(sigma) - + 0.5 * _LOG_2PI + nll = gaussian_per_refl( + residual, + torch.zeros_like(residual), + intensity_var_from_sigma_obs(sigma, mask), + var_floor=0.0, ) total = (nll * mask).sum() # Applied on the work set only, matching CollectionMLTarget: the free-set value @@ -231,16 +231,6 @@ def forward(self) -> torch.Tensor: total = self.base_weight * total return total - @staticmethod - def _floor_sigma(sigma: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: - """Clamp sigma away from zero, at a fraction of its median over the subset.""" - selected = sigma[mask] - if selected.numel() == 0: - return sigma.clamp(min=1e-6) - floor = torch.median(selected) * _SIGMA_FLOOR_FRAC - floor = torch.clamp(floor, min=1e-12) - return sigma.clamp(min=floor) - # ------------------------------------------------------------------ # Weight calibration # ------------------------------------------------------------------ From 85aa181605c28234e98f48f68b1a84a2c22c7b1b Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 16:39:45 +0200 Subject: [PATCH 069/250] Add an observable axis to the x-ray targets, with an intensity row The amplitude was assumed at one line: `get_F_calc_scaled` is `abs(get_fcalc_scaled(...))`, which is why the `torch.abs(F_calc)` calls inside the likelihood primitives are no-ops. Add `get_I_calc_scaled` beside it and the observable becomes a choice. `IntensityObservableMixin` is the whole of the intensity axis: one `get_data` override reading `sub.I`/`sub.sigI` and predicting `|F_calc|**2`. Declared on the spec rather than passed as a kwarg, so it cannot be dropped by the construction sites that bypass `_xray_target_kwargs`, and so no method branches at runtime. The spec checks its claim against the class. R-factors stay on amplitudes for every row: `sqrt(I_calc) == |F_calc|` for a squared-amplitude model, so the base needs no override and the whole table remains comparable. The axis is not square -- Rice and the folded normal are distributions of an amplitude, so there is no intensity Rice row. Intensity rows are not admissible as scale targets; `SCALE_TARGETS` still fails closed on them. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- docs/changelog.rst | 4 + docs/user_guide/cli.rst | 6 +- docs/user_guide/targets.rst | 29 +- tests/helpers/device_cases.py | 1 + .../refinement/test_intensity_observable.py | 254 ++++++++++++++++++ tests/unit/refinement/test_nll_beta.py | 56 ++++ torchref/refinement/targets/base.py | 32 +++ torchref/refinement/targets/xray/__init__.py | 3 + torchref/refinement/targets/xray/_specs.py | 53 ++++ .../refinement/targets/xray/observable.py | 145 ++++++++++ 10 files changed, 579 insertions(+), 4 deletions(-) create mode 100644 tests/unit/refinement/test_intensity_observable.py create mode 100644 torchref/refinement/targets/xray/observable.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 163bce64..fdc4aa08 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,10 @@ Changelog Version 0.6.4 ---------- +- Added ``--xray-mode nll_i``, a Gaussian on the observed intensities, and an ``observable`` column on the target taxonomy +- Added ``DataTarget.get_I_calc_scaled``, so the observable is a choice rather than an assumption +- Added ``gaussian_per_refl`` and ``intensity_var_from_sigma_obs``; the amplitude and intensity Gaussians are now one implementation +- The absolute variance floor in the shared Gaussian is now opt-out, since it distorts any objective whose sigmas fall below it - Added a reader for CrystFEL ``partialator`` ``.hkl`` reflection lists, via ``ReflectionData.load_crystfel_hkl`` - Added ``FcalcDataset.add_noise`` and the ``torchref.simulate-noisy-data`` CLI, which simulate merged intensities from a structure and report R-split and CC between two independent half-datasets - Simulated intensities keep their negative values; only the derived amplitude is clamped, since clamping the intensity biases the weak reflections upward diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index a649e62f..517cfb30 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -36,8 +36,10 @@ and a ``refinement_history.json`` log. ``ml_full`` (marginalises the measurement error rather than inflating the variance; ~4× the cost), ``nll_beta`` (the Gaussian large-signal limit of ``ml`` — diagnostic), ``nll`` (Gaussian weighted by σ_obs only, no model-error - term), ``ls`` (unit-weight least squares) or ``ls_wunit_k1`` (Phenix-style, own - global scale). ``--help`` lists them from the taxonomy table itself. + term), ``nll_i`` (as ``nll`` but on the observed *intensities*, skipping the + French–Wilson conversion), ``ls`` (unit-weight least squares) or ``ls_wunit_k1`` + (Phenix-style, own global scale). ``--help`` lists them from the taxonomy table + itself, which is authoritative. * ``--sigma-a-max`` upper bound on the per-shell Luzzati σ_A (default 0.99) * ``--no-shrink`` disable the per-shell σ_A stability shrinkage * ``--adp-mode`` ``isotropic`` (default) or ``anisotropic``, the latter refining diff --git a/docs/user_guide/targets.rst b/docs/user_guide/targets.rst index a0dcce54..66c30c4d 100644 --- a/docs/user_guide/targets.rst +++ b/docs/user_guide/targets.rst @@ -20,10 +20,17 @@ grouped into composite targets for geometry and ADP restraints. X-ray Targets ------------- -Seven modes, selected by name. ``XRAY_TARGETS`` (in +Selected by name. ``XRAY_TARGETS`` (in ``torchref.refinement.targets.xray._specs``) is the single table behind both :func:`~torchref.refinement.targets.create_xray_target` and -``torchref.refine --help``, so the list below cannot drift from the CLI: +``torchref.refine --help``. The authoritative list is ``torchref.refine --help``, +which is generated from that table; the notes below describe the rows but are +maintained by hand, so run ``--help`` if the two disagree. + +Every row shares one forward model — the scaled complex :math:`F_{calc}` — and +declares which measured column it compares against (``spec.observable``). + +**Amplitude rows** compare :math:`F_{obs}` against :math:`|F_{calc}|`: - ``ml`` — **default**. Read MLF: variance :math:`\epsilon\beta`, conditional mean :math:`\alpha|F_{calc}|`, with a cross-validated per-shell Luzzati @@ -40,6 +47,24 @@ Seven modes, selected by name. ``XRAY_TARGETS`` (in - ``ls_wunit_k1`` — Phenix-style least squares: unit weights and a single global scale recomputed every gradient call, bypassing the scaler. +**Intensity rows** compare :math:`I_{obs}` against :math:`|F_{calc}|^2`: + +- ``nll_i`` — Gaussian NLL on the observed intensities, weighted by + :math:`\sigma(I)`. As ``nll``, but skips the French–Wilson conversion. + +Use an intensity row when the signal lives in the *quadratic* part of the data — +a population variance, an activation second moment. :math:`F_{obs}` on a merged +dataset is a French–Wilson posterior rather than a measurement: it is strictly +positive, so it reshapes the weak tail and erases negative intensities, which is +precisely the information such a signal is carried by. + +The axis is deliberately not square. There is no intensity Rice row, because Rice +and the folded normal are distributions *of an amplitude* — the intensity +analogue is the exponential / :math:`\chi^2_1` Wilson distribution, a different +primitive rather than a different variance. R-factors are reported on amplitudes +for every row regardless, so they stay comparable across the whole table. Intensity +rows are not admissible as ``--scale-target``, which fails closed on them. + Geometry Targets ---------------- diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index e04c2494..cbae4c6e 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -320,6 +320,7 @@ class TargetDeviceCase: "ScalerLogScaleTrendTarget": "needs a scaler", "ScalerURegularizationTarget": "needs a scaler", "NLLXrayTarget": "needs model + data + scaler", + "NLLIntensityXrayTarget": "needs model + data + scaler with intensities", "LeastSquaresXrayTarget": "needs model + data + scaler", "UnitWeightK1XrayTarget": "needs model + data + scaler", "SigmaAXrayTarget": "abstract base; needs model + data + scaler", diff --git a/tests/unit/refinement/test_intensity_observable.py b/tests/unit/refinement/test_intensity_observable.py new file mode 100644 index 00000000..a3408aaf --- /dev/null +++ b/tests/unit/refinement/test_intensity_observable.py @@ -0,0 +1,254 @@ +"""The observable axis on the single-dataset x-ray targets. + +The claim under test: the observable is *one* `get_data` override, and everything else -- +the likelihood, the subsets, the masks, the R-factor -- is inherited unchanged. So the tests +here are mostly about what must NOT differ. + +The one thing that genuinely must differ is the variance: an amplitude sigma applied to an +intensity residual is wrong by ``2|F|``, which is resolution-dependent and therefore +presents as a scale or B error rather than as a bug. That is the error this file is for. +""" + +import math + +import pytest +import torch + +from torchref.base.targets.xray_likelihoods import ( + VAR_FLOOR, + amplitude_var_from_sigma_obs, + floor_sigma_obs, + gaussian_per_refl, + intensity_var_from_sigma_obs, + nll_per_refl, +) + + +# ===================================================================== +# The refactored primitives (step A) -- fileless, no fixture needed +# ===================================================================== + + +@pytest.mark.unit +def test_the_shared_gaussian_reproduces_the_amplitude_one_bitwise(): + """``nll_per_refl`` is ``gaussian_per_refl`` on ``|F_calc|``, exactly. + + Bitwise, not ``allclose``: the amplitude row has a fused Triton counterpart pinned to + it by ``tests/integration/test_triton_vs_eager_targets.py``, so any drift here shows up + there as a mysterious kernel disagreement rather than as this refactor. + """ + for dtype in (torch.float32, torch.float64): + torch.manual_seed(3) + F_obs = torch.rand(5000, dtype=dtype) * 100 + F_calc = torch.randn(5000, dtype=dtype) * 100 + var = amplitude_var_from_sigma_obs(torch.rand(5000, dtype=dtype) * 10) + assert torch.equal( + nll_per_refl(F_obs, F_calc, var), + gaussian_per_refl(F_obs, torch.abs(F_calc), var), + ) + + +@pytest.mark.unit +def test_the_absolute_variance_floor_is_opt_out_and_matters(): + """``VAR_FLOOR`` is a distortion, not a safeguard, once the builder has floored sigma. + + It is an *absolute* floor on a variance, so whether it engages depends on the units the + data happens to be in. On sigmas around 1e-5 it rescales the objective by a factor of + several -- which is why the intensity rows pass ``var_floor=0.0`` and rely on + :func:`floor_sigma_obs`' data-dependent floor instead. + + This is pinned because the two paths silently disagreed when the Gaussian was first + shared: the amplitude copy had the clamp and the intensity copy did not. + """ + sigma = torch.full((256,), 1e-3, dtype=torch.float64) + var = intensity_var_from_sigma_obs(sigma) # 1e-6, comfortably above VAR_FLOOR + obs = torch.zeros(256, dtype=torch.float64) + model = torch.full((256,), 1e-4, dtype=torch.float64) + assert torch.equal( + gaussian_per_refl(obs, model, var, var_floor=0.0), + gaussian_per_refl(obs, model, var, var_floor=VAR_FLOOR), + ), "the floor must be inert when the variance is above it" + + # Below it, the two differ -- and by a lot, not by an ulp. + tiny = torch.full((256,), 1e-6, dtype=torch.float64) # var = 1e-12 << VAR_FLOOR + var_tiny = intensity_var_from_sigma_obs(tiny) + free = gaussian_per_refl(obs, model, var_tiny, var_floor=0.0) + clamped = gaussian_per_refl(obs, model, var_tiny, var_floor=VAR_FLOOR) + assert not torch.allclose(free, clamped) + # `free` is the honest one: it uses the variance the builder actually produced. + expected = 0.5 * (1e-4) ** 2 / 1e-12 + 0.5 * math.log(1e-12) + 0.5 * math.log(2 * math.pi) + assert free[0].item() == pytest.approx(expected, rel=1e-12) + + +@pytest.mark.unit +def test_the_intensity_sigma_floor_respects_the_fitted_subset(): + """``mask`` restricts the median, because unfitted rows carry filler. + + A collection member reindexed onto a common reflection list has filler sigmas on the + rows it does not own. Taking the median over those moves the floor for every real + reflection, so the mask is not a convenience. + """ + # The first 20 fitted rows are BELOW the fitted median's floor, so the floor is what + # they come back as -- which is the only way to observe which median was used. + sigma = torch.cat([ + torch.full((20,), 0.01), # fitted, and below floor either way + torch.full((80,), 10.0), # fitted, sets the fitted median + torch.full((900,), 1e6), # NOT fitted: filler + ]) + mask = torch.cat([torch.ones(100, dtype=torch.bool), torch.zeros(900, dtype=torch.bool)]) + masked = floor_sigma_obs(sigma, mask, abs_floor=1e-12) + unmasked = floor_sigma_obs(sigma, None, abs_floor=1e-12) + assert masked[:20].min().item() == pytest.approx(1.0) # floor = 10 * 0.1 + assert unmasked[:20].min().item() == pytest.approx(1e5) # floor = 1e6 * 0.1, swamped + # An explicit floor overrides the median entirely -- the set-independent path. + assert floor_sigma_obs(sigma, mask, floor=0.5)[:20].min().item() == pytest.approx(0.5) + + +@pytest.mark.unit +def test_confusing_the_two_variance_builders_is_wrong_by_a_factor(): + """The amplitude and intensity builders are not interchangeable. + + Both square a floored sigma, so they *look* alike; what differs is which sigma. Passing + ``sigma(F)`` to an intensity residual (or the reverse) is wrong by ``(2|F|)**2``, which + varies with resolution -- so it does not present as an obviously wrong number, it + presents as a scale or B error. Hence a test rather than a comment. + """ + sig_F = torch.rand(1000, dtype=torch.float64) * 2 + 0.5 + F = torch.rand(1000, dtype=torch.float64) * 100 + 10 + sig_I = 2 * F * sig_F # exact first-order propagation, I = F**2 + var_wrong = amplitude_var_from_sigma_obs(sig_F) + var_right = intensity_var_from_sigma_obs(sig_I) + ratio = (var_right / var_wrong).sqrt() + # Spans a wide range: a single global weight cannot absorb it. + assert ratio.max() / ratio.min() > 5 + + +# ===================================================================== +# The row, on real data (step B) +# ===================================================================== + + +@pytest.fixture(scope="module") +def refinement(pdb_dir, mtz_dir): + """A scaled 1DAW refinement -- the only fixture carrying BOTH I/SIGI and FP/SIGFP, + so the only one on which the two observables can be compared at all.""" + pdb = pdb_dir / "1DAW.pdb" + mtz = mtz_dir / "1DAW.mtz" + if not (pdb.exists() and mtz.exists()): + pytest.skip("1DAW fixture not present") + from torchref import LBFGSRefinement + + ref = LBFGSRefinement(data_file=str(mtz), pdb=str(pdb), target_mode="ml", verbose=0) + ref.get_scales() + return ref + + +def _t(refinement, mode, use_set="work"): + from torchref.refinement.targets.xray.factory import create_xray_target + + return create_xray_target( + data=refinement.reflection_data, + model=refinement.model, + scaler=refinement.scaler, + mode=mode, + use_set=use_set, + ) + + +@pytest.mark.integration +def test_the_row_is_selectable_and_reads_intensities(refinement): + """``nll_i`` comes out of the factory and its ``get_data`` returns the I columns.""" + t = _t(refinement, "nll_i") + obs, calc, sigma, centric, sub = t.get_data() + data = refinement.reflection_data + + torch.testing.assert_close(obs, data.work.I) + torch.testing.assert_close(sigma, data.work.sigI) + # The model is the SQUARED scaled amplitude, not the amplitude. + torch.testing.assert_close(calc, sub.select(t.get_F_calc_scaled(recalc=False) ** 2)) + assert obs.shape == calc.shape == sigma.shape == (sub.n,) + assert centric.shape == (sub.n,) + + +@pytest.mark.integration +def test_the_intensity_model_is_the_squared_scaled_amplitude(refinement): + """``get_I_calc_scaled`` squares the SCALED amplitude, not the raw one. + + Both the overall scale and the anisotropy factor therefore enter squared, matching + ``ReflectionData.get_corrected_intensities`` on the observation side. Squaring first and + scaling afterwards with the amplitude factors would be wrong by that factor, which is + resolution-dependent. + """ + t = _t(refinement, "nll_i") + with torch.no_grad(): + amp = t.get_F_calc_scaled(recalc=False) + inten = t.get_I_calc_scaled(recalc=False) + torch.testing.assert_close(inten, amp**2, rtol=1e-6, atol=1e-6) + + +@pytest.mark.integration +def test_rfactors_stay_on_amplitudes(refinement): + """An intensity row reports the SAME R-factors as an amplitude row. + + ``_scaled_F_calc_full`` is deliberately not overridden: for a ``|F_calc|**2`` model its + correct value is ``sqrt(I_calc) == |F_calc|``, which is what the base returns. R-factors + therefore remain comparable across the whole table regardless of which observable drove + the loss -- and a future row whose intensity model is not a squared amplitude (the + two-moment model) has to override it, or this test is what will catch it. + """ + r_i = _t(refinement, "nll_i").get_rfactor() + r_a = _t(refinement, "nll").get_rfactor() + assert r_i == pytest.approx(r_a, abs=1e-9) + + +@pytest.mark.integration +def test_the_loss_is_finite_differentiable_and_summed(refinement): + t = _t(refinement, "nll_i") + loss = t.forward() + assert torch.isfinite(loss) and loss.ndim == 0 + loss.backward() + grads = [ + p.grad for p in refinement.model.parameters() + if p.requires_grad and p.grad is not None + ] + assert grads, "no gradient reached the model" + assert all(torch.isfinite(g).all() for g in grads) + refinement.model.zero_grad(set_to_none=True) + + +@pytest.mark.integration +@pytest.mark.parametrize("use_set", ["work", "free"]) +def test_a_reflections_residual_does_not_depend_on_the_arrays_length(refinement, use_set): + """``residuals()`` restricted to a subset must equal ``forward()`` on that subset. + + Pinned separately from ``test_xray_residuals.py`` because the failure mode is specific + to this row: the sigma floor is a *median*, so deriving it from whatever array a call + receives makes every per-reflection value depend on the whole array. ``forward`` sees + the subset and ``residuals`` sees everything, so the two disagreed by 0.09% on the work + set and 1.8% on the free set until the floor was pinned to the target's own subset. + + The amplitude rows share the mechanism and get away with it because sigma(F) is narrow + enough that the clamp barely engages; sigma(I) spans orders of magnitude. + """ + t = _t(refinement, "nll_i", use_set=use_set) + sub = t._subset() + with torch.no_grad(): + fwd = t.forward() + summed = t.residuals().index_select(0, sub.indices).sum() + torch.testing.assert_close(summed, fwd, rtol=1e-6, atol=1e-6) + + +@pytest.mark.integration +def test_missing_intensities_raise_at_construction_not_at_forward(refinement): + """LossState probes ``forward()`` at registration, so a missing column has to be + caught in ``__init__`` or it surfaces from deep inside setup with no mention of why.""" + import copy + + data = copy.copy(refinement.reflection_data) + data.I = None + from torchref.refinement.targets.xray import NLLIntensityXrayTarget + + with pytest.raises(ValueError, match="dataset carries none"): + NLLIntensityXrayTarget( + data=data, model=refinement.model, scaler=refinement.scaler + ) diff --git a/tests/unit/refinement/test_nll_beta.py b/tests/unit/refinement/test_nll_beta.py index 59eb120a..832765e2 100644 --- a/tests/unit/refinement/test_nll_beta.py +++ b/tests/unit/refinement/test_nll_beta.py @@ -216,6 +216,7 @@ def test_each_mode_has_its_own_class(): MLNoAlphaXrayTarget, MLXrayTarget, NLLBetaXrayTarget, + NLLIntensityXrayTarget, NLLXrayTarget, UnitWeightK1XrayTarget, ) @@ -227,6 +228,7 @@ def test_each_mode_has_its_own_class(): "ml_full": MLFullXrayTarget, "nll_beta": NLLBetaXrayTarget, "nll": NLLXrayTarget, + "nll_i": NLLIntensityXrayTarget, "ls": LeastSquaresXrayTarget, "ls_wunit_k1": UnitWeightK1XrayTarget, } @@ -244,6 +246,60 @@ def test_each_mode_has_its_own_class(): seen[cls] = name +def test_the_observable_is_a_row_property_not_a_flag(): + """The observable is declared by the spec AND by the class, and they must agree. + + Same argument as ``test_mean_centring_is_intrinsic_to_the_spec_not_a_flag``: a + constructor kwarg would be a runtime branch, and would be silently dropped by the + construction sites that bypass ``Refinement._xray_target_kwargs`` (the ensemble + refinement has three of them). A row cannot be selected without selecting its class. + """ + import inspect + + from torchref.refinement.targets.xray import NLLIntensityXrayTarget, NLLXrayTarget + from torchref.refinement.targets.xray._specs import XRAY_TARGETS, XrayTargetSpec + + by_obs = {} + for spec in XRAY_TARGETS.specs: + assert spec.observable in ("amplitude", "intensity"), spec.name + # The spec's claim and the class's own declaration must match -- a row advertising + # intensities while reading `sub.F` would be wrong by 2|F| with nothing to catch it. + assert getattr(spec.target_cls, "observable", "amplitude") == spec.observable + by_obs.setdefault(spec.observable, []).append(spec.name) + + assert "nll_i" in by_obs["intensity"], "the intensity row went missing from the table" + assert "nll" in by_obs["amplitude"] + + # Not a constructor flag anywhere. + for cls in (NLLXrayTarget, NLLIntensityXrayTarget): + assert "observable" not in inspect.signature(cls.__init__).parameters + + # The spec rejects a class/spec mismatch rather than trusting either side. + with pytest.raises(ValueError, match="observable"): + XrayTargetSpec( + name="bogus", target_cls=NLLXrayTarget, doc="", observable="intensity" + ) + + +def test_rice_has_no_intensity_row(): + """Rice is amplitude-only *by nature*, so the observable axis is not square. + + Rice and the folded normal are distributions of an amplitude; the intensity analogue is + the exponential / chi-square_1 Wilson distribution, a different primitive. If a future + row pairs a Rice class with ``observable="intensity"`` it is a modelling error, not a + new feature -- so pin it here rather than discovering it from a bad refinement. + """ + from torchref.refinement.targets.xray import RiceXrayTarget, SigmaAXrayTarget + from torchref.refinement.targets.xray._specs import XRAY_TARGETS + + for spec in XRAY_TARGETS.specs: + if spec.observable != "intensity": + continue + assert not issubclass(spec.target_cls, (RiceXrayTarget, SigmaAXrayTarget)), ( + f"{spec.name} pairs an amplitude distribution with intensities" + ) + + def test_only_the_estimator_backed_rows_own_an_estimator(): """``nll`` must not pay for a model-error estimate it does not use. diff --git a/torchref/refinement/targets/base.py b/torchref/refinement/targets/base.py index 0a4b62e2..884e349b 100644 --- a/torchref/refinement/targets/base.py +++ b/torchref/refinement/targets/base.py @@ -385,6 +385,38 @@ def get_F_calc_scaled(self, hkl=None, recalc=False, fcalc=None): """ return torch.abs(self.get_fcalc_scaled(hkl, recalc=recalc, fcalc=fcalc)) + def get_I_calc_scaled(self, hkl=None, recalc=False, fcalc=None): + """ + Compute scaled structure factor intensities ``|F_calc|**2``. + + The intensity sibling of :meth:`get_F_calc_scaled`, and the reason the observable + is a choice rather than an assumption: both are one line over the same complex + ``get_fcalc_scaled``, so nothing upstream of here knows which observable a target + fits. + + Squaring the *scaled* amplitude is what makes this correct -- the scale and the + anisotropy factor both enter squared, matching + :meth:`ReflectionData.get_corrected_intensities` on the observation side. Squaring + an unscaled ``F_calc`` and scaling afterwards with the amplitude factors would be + wrong by that factor, which is resolution-dependent and so reads as a scale or B + error rather than as a bug. + + Parameters + ---------- + hkl : torch.Tensor, optional + Miller indices. If None, uses data's hkl. + recalc : bool, optional + Force recalculation. Default is False. + fcalc : torch.Tensor, optional + Pre-computed structure factors. If provided, skips model computation. + + Returns + ------- + torch.Tensor + Scaled structure factor intensities ``|F_calc|**2``. + """ + return self.get_fcalc_scaled(hkl, recalc=recalc, fcalc=fcalc).abs() ** 2 + # ============================================================================= # Utility Functions for NLL Computation diff --git a/torchref/refinement/targets/xray/__init__.py b/torchref/refinement/targets/xray/__init__.py index 6189e6ac..cfe98148 100644 --- a/torchref/refinement/targets/xray/__init__.py +++ b/torchref/refinement/targets/xray/__init__.py @@ -15,6 +15,7 @@ from .nll import NLLXrayTarget from .nll_beta import NLLBetaXrayTarget from .rice import RiceXrayTarget +from .observable import IntensityObservableMixin, NLLIntensityXrayTarget from .sigma_a import AlphaCentredMixin, SigmaALossInputs, SigmaAXrayTarget __all__ = [ @@ -23,6 +24,8 @@ "SigmaAXrayTarget", "SigmaALossInputs", "AlphaCentredMixin", + "IntensityObservableMixin", + "NLLIntensityXrayTarget", # the five selectable likelihood rows "NLLXrayTarget", "NLLBetaXrayTarget", diff --git a/torchref/refinement/targets/xray/_specs.py b/torchref/refinement/targets/xray/_specs.py index 98940dcb..1058ba0d 100644 --- a/torchref/refinement/targets/xray/_specs.py +++ b/torchref/refinement/targets/xray/_specs.py @@ -7,6 +7,31 @@ following :mod:`torchref.utils.backends`; string literals validated against a table, not enums, is the house convention. +## The observable + +Every row shares one forward model -- the scaled complex ``F_calc`` -- and each declares +which measured column it compares against, via ``spec.observable``: + +* **amplitude** ``F_obs`` vs ``|F_calc|`` (all the sigma_A rows, and ``ls``) +* **intensity** ``I_obs`` vs ``|F_calc|**2`` (``nll_i``) + +The intensity rows exist because ``F_obs`` on a merged dataset is a French-Wilson posterior +rather than a measurement: strictly positive, so it reshapes the weak tail and erases +negative intensities. A row whose signal lives in the *quadratic* part of the data reads +``I_obs`` directly. See :mod:`torchref.refinement.targets.xray.observable`, which is the +whole of that axis -- one ``get_data`` override, no runtime branch anywhere. + +Note the axis is **not square**: there is no intensity Rice, because Rice and the folded +normal are distributions *of an amplitude* and the intensity analogue is the exponential / +chi-square_1 Wilson distribution -- a different primitive, not a different variance. +R-factors stay on amplitudes for every row regardless, so they remain comparable across the +whole table. + +Intensity rows are **not** admissible as scale targets: +:data:`~torchref.scaling.scaler_base.SCALE_TARGETS` normalises its objective by +``1/sum(F_obs**2)``, which is dimensionally wrong for an ``O(F**4)`` loss, and that tuple +fails closed on anything it does not list. + ## The sigma_A family Each is a choice of (distribution) x (variance) x (mean): @@ -54,6 +79,7 @@ from .ml_noalpha import MLNoAlphaXrayTarget from .nll import NLLXrayTarget from .nll_beta import NLLBetaXrayTarget +from .observable import IntensityObservableMixin, NLLIntensityXrayTarget # noqa: F401 #: The mode built when none is given. DEFAULT_XRAY_MODE = "ml" @@ -75,18 +101,38 @@ class XrayTargetSpec: aliases Retired spellings kept working; resolving one emits a ``DeprecationWarning``. No row carries one at present, so the tests exercise this with their own table. + observable + Which measured column the row fits: ``"amplitude"`` or ``"intensity"``. Declarative + rather than a constructor flag, because it is a property of the row -- see + :mod:`torchref.refinement.targets.xray.observable`. Checked here against the class, + so a spec and its implementation cannot disagree. """ name: str target_cls: type doc: str aliases: Tuple[str, ...] = () + observable: str = "amplitude" def __post_init__(self): if not (isinstance(self.target_cls, type) and issubclass(self.target_cls, XrayTarget)): raise TypeError( f"{self.name}: target_cls {self.target_cls!r} is not an XrayTarget subclass" ) + if self.observable not in ("amplitude", "intensity"): + raise ValueError( + f"{self.name}: observable must be 'amplitude' or 'intensity', " + f"got {self.observable!r}" + ) + # The class declares its own observable (the mixin sets it); the spec must agree. + # Otherwise a row could advertise intensities while reading `sub.F`, which no test + # downstream of here would notice -- the loss would simply be wrong by 2|F|. + declared = getattr(self.target_cls, "observable", "amplitude") + if declared != self.observable: + raise ValueError( + f"{self.name}: spec says observable={self.observable!r} but " + f"{self.target_cls.__name__} says {declared!r}" + ) @dataclass(frozen=True) @@ -171,6 +217,13 @@ def by_name(self, name: str) -> XrayTargetSpec: doc="Gaussian amplitude NLL weighted by the experimental sigma only. No " "model-error term, so it does not control overfitting.", ), + XrayTargetSpec( + name="nll_i", + target_cls=NLLIntensityXrayTarget, + observable="intensity", + doc="Gaussian NLL on the observed INTENSITIES weighted by sigma(I). As 'nll' " + "but skips the French-Wilson conversion, which reshapes the weak tail.", + ), XrayTargetSpec( name="ls", target_cls=LeastSquaresXrayTarget, diff --git a/torchref/refinement/targets/xray/observable.py b/torchref/refinement/targets/xray/observable.py new file mode 100644 index 00000000..ca08c668 --- /dev/null +++ b/torchref/refinement/targets/xray/observable.py @@ -0,0 +1,145 @@ +"""The observable axis: which measured column a row fits. + +Every X-ray row shares one forward model -- the scaled complex ``F_calc`` -- and differs +only in what it compares against. Two observables are available: + +* **amplitude** (the default), ``F_obs`` against ``|F_calc|`` +* **intensity**, ``I_obs`` against ``|F_calc|**2`` + +:class:`IntensityObservableMixin` is the whole of the second one. It overrides +:meth:`XrayTarget.get_data` and nothing else, so the likelihood, the mean, the subset +selection, the masks and the R-factor are all inherited untouched. + +**Why intensities are worth a row at all.** ``F_obs`` on a merged dataset is a +French-Wilson posterior, not a measurement: the estimator is strictly positive, so it +reshapes the weak tail and erases negative intensities entirely. Anything whose signal +lives in the *quadratic* part of the data -- an activation second moment, a population +variance -- is fitting a distorted version of the quantity it is trying to measure. Rows +that need that information read ``I_obs`` directly. + +**Why there is no intensity Rice.** Rice and the folded normal are distributions *of an +amplitude*; the intensity analogue is the exponential / chi-square_1 Wilson distribution, +which is a different primitive rather than a different variance. So the intensity axis +carries the Gaussian rows only, and that is a property of the statistics rather than a gap +in the implementation. +""" + +from typing import Tuple + +import torch + +from torchref.base.targets.xray_likelihoods import ( + SIGMA_FLOOR_ABS, + SIGMA_FLOOR_FRAC, + gaussian_per_refl, + intensity_var_from_sigma_obs, + _masked_sum, +) + +from .base import XrayTarget + + +class IntensityObservableMixin: + """Read ``I_obs``/``sigma(I)`` and predict ``|F_calc|**2``. + + Mix in **before** an :class:`~.base.XrayTarget` subclass. The only override is + :meth:`get_data`, which is the single place the observable is chosen -- so a row + composed with this mixin cannot end up fitting intensities while reporting statistics + on amplitudes, and no method needs a runtime branch. + + Note there is deliberately no ``_scaled_F_calc_full`` override. That method feeds + :meth:`XrayTarget.get_rfactor`, and for a ``|F_calc|**2`` model its correct value is + ``sqrt(I_calc) == |F_calc|`` -- exactly what the inherited implementation returns. So + **R-factors stay on amplitudes for every row**, comparable across the whole table + regardless of which observable drove the loss. A row whose intensity model is *not* + the square of an amplitude (the two-moment model, where it is + ``|F|**2 + var*|dF|**2``) must override it to report ``sqrt`` of its own model. + """ + + #: Declared for the taxonomy table, and readable off any constructed target. + observable: str = "intensity" + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + # Checked at construction rather than at the first forward: LossState probes + # `forward()` once when a target is registered, so a missing column would + # otherwise surface as a failure deep inside setup with no mention of the cause. + data = getattr(self, "_data", None) + if data is not None and getattr(data, "I", None) is None: + raise ValueError( + f"{type(self).__name__} fits intensities, but this dataset carries none. " + "Load an MTZ/mmCIF with an I column (`I-obs`/`intensity_meas`), or select " + "an amplitude row such as `nll`." + ) + + def get_data( + self, fcalc: torch.Tensor = None, sub=None + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, object]: + """``(I_obs, I_calc, sigma_I, centric, sub)`` -- the intensity twin of + :meth:`XrayTarget.get_data`, same tuple shape and same subset semantics. + + Both observation columns are the *corrected* views (``sub.I``/``sub.sigI``), in + which the scale and the anisotropy factor enter squared, so they are on the same + footing as the squared model amplitude from :meth:`get_I_calc_scaled`. + """ + if sub is None: + sub = self._subset() + + I_obs = sub.I + sigma = sub.sigI + centric = sub.centric + + if fcalc is not None: + I_calc_full = self.get_I_calc_scaled(fcalc=fcalc) + else: + I_calc_full = self.get_I_calc_scaled(recalc=False) + I_calc = sub.select(I_calc_full) + + return I_obs, I_calc, sigma, centric, sub + + def _sigma_floor(self) -> torch.Tensor: + """The intensity-sigma floor, taken from THIS target's own fitted subset. + + Computed here rather than inside the variance builder so it does not depend on + which reflections a particular call happens to pass. ``forward`` evaluates on the + target's subset while ``residuals`` evaluates on every reflection; a floor derived + from the argument therefore differs between them, and the same reflection scores + differently in the two -- 0.09% on a 1DAW work set, 1.8% on its free set, because + sigma(I) spans orders of magnitude where sigma(F) does not. + + Detached: it is a numerical safeguard, not a fitted quantity, and letting a + gradient run back through a median would make the loss depend on the ordering of + near-equal sigmas. + """ + sigma = self._subset().sigI + if sigma is None or sigma.numel() == 0: + return torch.as_tensor(1e-6) + return (torch.median(sigma).detach() * SIGMA_FLOOR_FRAC).clamp(min=SIGMA_FLOOR_ABS) + + +class NLLIntensityXrayTarget(IntensityObservableMixin, XrayTarget): + """``--xray-mode nll_i``: Gaussian intensity NLL weighted by the experimental sigma. + + NLL = 0.5*(I_obs - |F_calc|**2)**2/sigma_I**2 + log(sigma_I) + 0.5*log(2*pi) + + The intensity counterpart of :class:`~.nll.NLLXrayTarget`, and like it carries no + model-error term, so it does **not** control overfitting. + + Subclasses :class:`~.base.XrayTarget` directly rather than ``NLLXrayTarget``, because + that row's ``forward`` calls the fused Triton amplitude kernel + (``nll_sigma_obs_math``) which has no intensity counterpart. Here ``forward`` is the + structural ``_masked_sum(_per_refl(...))``, so it and :meth:`residuals` are the same + expression by construction rather than by test. + """ + + target_value: float = 1.0 + + def forward(self, fcalc: torch.Tensor = None) -> torch.Tensor: + """Summed Gaussian NLL of the observed intensities on this target's set.""" + return _masked_sum(self._per_refl(self._loss_inputs(fcalc=fcalc))) + + def _per_refl(self, ctx) -> torch.Tensor: + """Per-reflection Gaussian on the intensity. See :func:`gaussian_per_refl`.""" + I_obs, I_calc, sigma, _, _ = ctx + var = intensity_var_from_sigma_obs(sigma, floor=self._sigma_floor()) + return gaussian_per_refl(I_obs, I_calc, var, var_floor=0.0) From 402645ce865c1131f8c5306b2529b50c39e43268 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 16:43:19 +0200 Subject: [PATCH 070/250] Give the collection targets the seam the single-dataset ones have Each collection row hand-wrote its own forward, re-deriving stacking, masking, the sigma floor and the summing. That cost two divergent stacking styles and three different sigma-floor constants, and it is why the batched accessors could return raw amplitudes for one dataset and scaled ones for another without anything noticing. `_loss_inputs` now gathers (obs, model, sigma, mask) as (N, n_hkl) in the row's declared observable, and `_per_refl` evaluates the likelihood unreduced; `forward` and `residuals` differ only in whether they sum. Observations route through the collection's own batched accessors rather than a private loop. Cross-dataset-coupled rows narrow the mask in `_loss_inputs`, so the sum and the residual array agree on which reflections count. One intentional change, to gradients only: the sigma floor is now detached. A gradient through a median is a single-element selection, which makes the loss depend on the ordering of near-equal sigmas. Values are unchanged -- the pinned characterisation literals reproduce exactly. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- .../refinement/targets/collection/base.py | 185 +++++++++++++++++- .../targets/collection/intensity.py | 62 +++--- .../refinement/targets/collection/xray.py | 104 ++++------ 3 files changed, 245 insertions(+), 106 deletions(-) diff --git a/torchref/refinement/targets/collection/base.py b/torchref/refinement/targets/collection/base.py index 064cd971..01ca7a6f 100644 --- a/torchref/refinement/targets/collection/base.py +++ b/torchref/refinement/targets/collection/base.py @@ -11,13 +11,35 @@ Since every member is expanded onto one common HKL grid, per-dataset R-factors form a distribution: headline ``rwork``/``rfree`` are its median, with the 10/25/75/90 percentiles at higher verbosity. + +## The seam + +Same two-part seam as the single-dataset base, batched. :meth:`_loss_inputs` gathers +what a row reads -- observations, model, sigma and mask, each ``(N, n_hkl)`` on the +common grid -- and :meth:`_per_refl` evaluates the likelihood on it *unreduced*. +:meth:`forward` and :meth:`residuals` differ only in whether they sum, so the two cannot +drift into different objectives. + +Every row shares one forward model and declares its ``observable`` (``"amplitude"`` or +``"intensity"``), exactly as the single-dataset table does. The base reads the matching +columns, so no row does its own stacking, masking or sigma flooring -- which is what +three divergent stacking styles and three different sigma-floor constants used to cost. + +Rows that are **cross-dataset coupled** (the difference targets take every dataset +against the mean of all of them) narrow the mask in :meth:`_loss_inputs` so a reflection +counts only if it is in the subset of *every* member, then work on the whole stack inside +:meth:`_per_refl`. That is a mask decision, not a special case in the base. """ -from typing import TYPE_CHECKING, Dict, List, Optional +from typing import TYPE_CHECKING, Dict, List, NamedTuple, Optional import torch from torchref.base.metrics.rfactor import rfactor_work_free +from torchref.base.targets.xray_likelihoods import ( + SIGMA_FLOOR_ABS, + SIGMA_FLOOR_FRAC, +) from torchref.refinement.targets.base import Target from torchref.utils.stats import ( VERBOSITY_DEBUG, @@ -40,6 +62,29 @@ _R_PCT_LABELS = ("p10", "p25", "p50", "p75", "p90") +class CollectionLossInputs(NamedTuple): + """What a collection row's :meth:`CollectionXrayTarget._per_refl` reads. + + Every tensor is ``(N, n_hkl)`` on the collection's common HKL grid, with ``N`` the + number of matched datasets in ``keys`` order -- full size rather than compact, because + the members are already expanded onto one grid and a compact form would need a + different index map per dataset. + + ``mask`` has already been intersected with finiteness of ``obs`` and ``sigma``, and + those two have been substituted where non-finite. That order matters: masking the + *loss* is not enough, because ``torch.where`` selects the finite branch for the value + while still backpropagating NaN through the branch it discarded. Real reflection files + carry non-finite intensities (excluded rows, and rows French-Wilson rejected), so this + is load-bearing rather than defensive. + """ + + obs: torch.Tensor + model: torch.Tensor + sigma: torch.Tensor + mask: torch.Tensor + keys: List[str] + + class CollectionXrayTarget(Target): """Base class for multi-dataset X-ray targets. @@ -63,6 +108,21 @@ class CollectionXrayTarget(Target): name: str = "collection_xray" + #: Which measured column this row fits: ``"amplitude"`` or ``"intensity"``. Declared + #: rather than passed, for the same reason as the single-dataset table -- see + #: :mod:`torchref.refinement.targets.xray.observable`. + observable: str = "amplitude" + + #: Fewest matched datasets for the loss to mean anything. The difference targets need + #: two (there is no difference from a single dataset); the per-dataset rows need one. + min_datasets: int = 1 + + #: Multiplies the work-set loss. Rows carrying a likelihood whose magnitude differs + #: from its siblings' set this so the term neither swamps nor is swamped by the + #: restraints; :meth:`CollectionTwoMomentIntensityTarget.calibrate_base_weight` fits + #: it against a reference target's gradient norm. + base_weight: float = 1.0 + def __init__( self, dataset_collection: "DatasetCollection", @@ -125,6 +185,129 @@ def _scaled_amp_full(self, data, model, recalc: bool = True) -> torch.Tensor: fcalc = data.structure_factors(model, recalc=recalc) return torch.abs(_scale_fcalc(self._scaler, fcalc, model)) + # ------------------------------------------------------------------ + # The per-reflection seam + # ------------------------------------------------------------------ + + def _stack_observations(self, keys: List[str]): + """``(obs, sigma)``, each ``(N, n_hkl)``, in this row's observable. + + Routed through the collection's own batched accessors rather than looping over + ``data.get_corrected_*()`` here: they already apply the inter-dataset scaling (the + intensity factors squared), cache against the ``(log_scale, U_aniso)`` fingerprint, + and name the offending dataset when an intensity column is missing. Reading one + dataset's raw column and another's scaled one is a silent regression that was live + once, when the batched accessors returned raw ``data.F``. + """ + dc = self._dataset_collection + if self.observable == "intensity": + return dc.stack_I_obs(keys), dc.stack_I_sigma(keys) + return dc.stack_F_obs(keys), dc.stack_F_sigma(keys) + + def _stack_model(self, keys: List[str], recalc: bool = False) -> torch.Tensor: + """The model prediction, ``(N, n_hkl)``, in this row's observable. + + Default is the per-dataset scaled amplitude (squared for an intensity row). Rows + whose prediction is not a function of one dataset at a time -- the two-moment + model, which mixes shared components across timepoints -- override this. + """ + dc = self._dataset_collection + mc = self._model_collection + amp = torch.stack( + [self._scaled_amp_full(dc[k], mc[k], recalc=recalc) for k in keys] + ) + return amp**2 if self.observable == "intensity" else amp + + def _stack_masks(self, keys: List[str]) -> torch.Tensor: + """This row's subset mask per dataset, ``(N, n_hkl)``. Validity and the 3-way + work/free/validation selection, with validation carved out of both. + """ + return self._dataset_collection.stack_masks(keys, use_set=self.use_set) + + def _sigma_floor(self, sigma: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: + """Floor for ``sigma``, at :data:`SIGMA_FLOOR_FRAC` of its median over ``mask``. + + A merged sigma can be reported as exactly zero; unfloored, one such reflection + dominates the whole sum. Detached, because it is a numerical safeguard rather than + a fitted quantity -- a gradient through a median would make the loss depend on the + ordering of near-equal sigmas. + + Taken over the fitted rows only: the members are reindexed onto a common grid, so + the rows a dataset does not own carry filler that would move the median. + """ + selected = sigma[mask] + if selected.numel() == 0: + return torch.as_tensor(1e-6, device=sigma.device, dtype=sigma.dtype) + floor = torch.median(selected).detach() * SIGMA_FLOOR_FRAC + return floor.clamp(min=SIGMA_FLOOR_ABS) + + def _loss_inputs(self, recalc: bool = False) -> CollectionLossInputs: + """Gather this row's observations, model, sigma and mask -- see + :class:`CollectionLossInputs` for the shapes and the NaN discipline. + + Rows narrow the mask here (the difference targets require a reflection to be in + the subset of every dataset) rather than inside :meth:`_per_refl`, so that + :meth:`forward`'s sum and :meth:`residuals`' array agree on which reflections + count. + """ + keys = self._keys() + obs, sigma = self._stack_observations(keys) + model = self._stack_model(keys, recalc=recalc) + mask = self._stack_masks(keys) + + obs = obs.to(model.dtype) + sigma = sigma.to(model.dtype) + + # Sanitise into the mask BEFORE the graph, not after: see CollectionLossInputs. + valid = torch.isfinite(obs) & torch.isfinite(sigma) + mask = mask & valid + obs = torch.where(valid, obs, torch.zeros_like(obs)) + sigma = torch.where(valid, sigma, torch.ones_like(sigma)) + + return CollectionLossInputs(obs, model, sigma, mask, keys) + + def _per_refl(self, ctx: CollectionLossInputs) -> torch.Tensor: + """The likelihood, per reflection and **unreduced**, shape ``(N, n_hkl)``. + + One per selectable row; no row branches. :meth:`forward` is the masked sum of + this. + """ + raise NotImplementedError + + def forward(self) -> torch.Tensor: + """Masked sum of :meth:`_per_refl`, with ``base_weight`` on the work set only. + + Cache reset first: a preceding no-grad ``stats()`` or ``get_rfactor()`` call can + leave a detached tensor in a base model's cache, which would silently kill the + loss backward. + """ + keys = self._keys() + if len(keys) < self.min_datasets: + return torch.zeros((), device=self._dataset_collection.hkl.device) + + self._reset_model_caches() + ctx = self._loss_inputs(recalc=False) + total = (self._per_refl(ctx) * ctx.mask).sum() + # Work set only: the free-set value is a diagnostic and has to stay comparable + # across weightings. + if self.use_work_set and self.base_weight != 1.0: + total = self.base_weight * total + return total + + def residuals(self) -> torch.Tensor: + """:meth:`_per_refl` over every reflection, ``(N, n_hkl)``, unsummed and unmasked. + + The unreduced :meth:`forward`: same observable, same model, same variance. Masked + reflections still get a value, so the array can be used to ask *why* one was + excluded rather than only reflecting the answer back, and non-finite values + survive because here a NaN is a finding rather than a nuisance. + """ + keys = self._keys() + if len(keys) < self.min_datasets: + dc = self._dataset_collection + return torch.zeros((0, len(dc.hkl)), device=dc.hkl.device) + return self._per_refl(self._loss_inputs(recalc=True)) + # ------------------------------------------------------------------ # R-factor reporting (shared source of truth) # ------------------------------------------------------------------ diff --git a/torchref/refinement/targets/collection/intensity.py b/torchref/refinement/targets/collection/intensity.py index fa608e8f..51b6c14c 100644 --- a/torchref/refinement/targets/collection/intensity.py +++ b/torchref/refinement/targets/collection/intensity.py @@ -92,6 +92,11 @@ class CollectionTwoMomentIntensityTarget(CollectionXrayTarget): name: str = "collection_two_moment_intensity" + #: Fits the merged INTENSITIES, so the base reads ``I``/``sigI``. The point of the + #: row: French-Wilson reshapes precisely the quadratic information the second moment + #: lives in, and there is deliberately no ``F**2`` fallback. + observable: str = "intensity" + def __init__( self, dataset_collection: "DatasetCollection", @@ -186,50 +191,31 @@ def _variance_is_live(self, sigma_alpha_sq) -> bool: return True return bool(sigma_alpha_sq.detach().ne(0).any()) - def forward(self) -> torch.Tensor: - """Summed Gaussian NLL of the observed intensities under the two-moment model.""" - dc = self._dataset_collection - keys = self._keys() - if not keys: - return torch.zeros((), device=dc.hkl.device) - - # Clear cached forwards so a preceding no-grad stats()/get_rfactor() call cannot - # leave a detached tensor that breaks the loss backward. - self._reset_model_caches() - - model = self.intensity_model(recalc=False) - obs = dc.stack_I_obs(keys).to(model.dtype) - sigma = dc.stack_I_sigma(keys).to(model.dtype) - mask = dc.stack_masks(keys, use_set=self.use_set) - - # Real reflection files carry non-finite intensities (excluded rows, and rows - # French-Wilson rejected). Masking the *loss* is not enough: a NaN observation - # makes the residual NaN, and torch.where selects the finite branch for the - # value while still backpropagating NaN through the branch it discarded. So the - # observations are sanitised into the mask BEFORE they reach the graph. - valid = torch.isfinite(obs) & torch.isfinite(sigma) - mask = mask & valid - obs = torch.where(valid, obs, torch.zeros_like(obs)) - sigma = torch.where(valid, sigma, torch.ones_like(sigma)) + def _stack_model(self, keys, recalc: bool = False) -> torch.Tensor: + """The two-moment intensity, not a per-dataset squared amplitude. + Overridden because this row's prediction is **not** a function of one dataset at a + time: the mean mixes shared components across timepoints and the variance term is + built from the activation Jacobian over the same components. ``keys`` is accepted + for the base's signature; :meth:`intensity_model` derives the rows itself. + """ + return self.intensity_model(recalc=recalc) + + def _per_refl(self, ctx) -> torch.Tensor: + """Per-reflection Gaussian NLL of the observed intensities under the model.""" # The residual is formed and masked BEFORE the Gaussian, so a masked-out row # contributes an exact zero rather than a value that merely gets multiplied by # zero. That matters if the model is ever non-finite on an unfitted row: here the # `where` discards it, whereas `nll * mask` would propagate NaN into the sum. - # Hence the Gaussian is evaluated at (residual, 0) rather than (obs, model). - residual = torch.where(mask, obs - model, torch.zeros_like(obs)) - nll = gaussian_per_refl( - residual, - torch.zeros_like(residual), - intensity_var_from_sigma_obs(sigma, mask), - var_floor=0.0, + residual = torch.where( + ctx.mask, ctx.obs - ctx.model, torch.zeros_like(ctx.obs) + ) + var = intensity_var_from_sigma_obs( + ctx.sigma, floor=self._sigma_floor(ctx.sigma, ctx.mask) + ) + return gaussian_per_refl( + residual, torch.zeros_like(residual), var, var_floor=0.0 ) - total = (nll * mask).sum() - # Applied on the work set only, matching CollectionMLTarget: the free-set value - # is a diagnostic and must stay comparable across weightings. - if self.use_work_set: - total = self.base_weight * total - return total # ------------------------------------------------------------------ # Weight calibration diff --git a/torchref/refinement/targets/collection/xray.py b/torchref/refinement/targets/collection/xray.py index 0a7550ac..5a19e7b1 100644 --- a/torchref/refinement/targets/collection/xray.py +++ b/torchref/refinement/targets/collection/xray.py @@ -17,7 +17,11 @@ import torch from torchref.base.reciprocal import get_scattering_vectors -from torchref.base.targets.xray_likelihoods import complex_var_from_beta, rice_math +from torchref.base.targets.xray_likelihoods import ( + complex_var_from_beta, + gaussian_per_refl, + rice_math, +) from torchref.refinement.model_error_estimation.sigma_a import SigmaAEstimator, epsilon_from_hkl from torchref.utils.stats import VERBOSITY_STANDARD, StatEntry, stat @@ -77,6 +81,9 @@ class CollectionDifferenceTarget(CollectionXrayTarget): name: str = "difference_xray" + #: There is no difference from a single dataset. + min_datasets: int = 2 + def __init__( self, dataset_collection: "DatasetCollection", @@ -144,82 +151,45 @@ def _activation_variance(self, keys, F_obs_stack) -> torch.Tensor: shift = contamination / (2.0 * F_obs_stack.abs().clamp(min=1e-6)) return shift**2 - def forward(self) -> torch.Tensor: - """Summed Gaussian NLL of the difference-from-mean; 0.0 if fewer than 2 sets.""" - dc = self._dataset_collection - mc = self._model_collection - - all_keys = self._keys() - N = len(all_keys) - if N < 2: - return torch.tensor(0.0, device=dc.hkl.device) - - # Clear caches so a preceding no-grad stats()/get_rfactor() call cannot - # leave a detached tensor that breaks the loss backward. - self._reset_model_caches() - - F_obs_list, sigma_list, mask_list, F_calc_list = [], [], [], [] - - for key in all_keys: - data = dc[key] - model = mc[key] + def _loss_inputs(self, recalc: bool = False): + """The base's stack, with the mask narrowed across datasets. - F_obs, sigma = data.get_corrected_data() - F_calc = self._scaled_amp_full(data, model, recalc=False) - # Validity + work/free/val selection, validation carved out of both. - mask = self._subset(data).mask - - F_obs_list.append(F_obs) - sigma_list.append(sigma) - mask_list.append(mask) - F_calc_list.append(F_calc) - - F_obs_stack = torch.stack(F_obs_list) # (N, n_hkl) - sigma_stack = torch.stack(sigma_list) # (N, n_hkl) - mask_stack = torch.stack(mask_list) # (N, n_hkl) - F_calc_stack = torch.stack(F_calc_list) # (N, n_hkl) - - # A reflection must be in this subset in ALL datasets. - mask_all = mask_stack.all(dim=0) # (n_hkl,) + A reflection counts only if it is in this target's subset in **every** dataset: + the per-reflection mean ties them together, so a reflection missing from one + member would silently shift the reference for all the others. Narrowed here + rather than inside :meth:`_per_refl` so ``forward``'s sum and ``residuals``' + array agree on which reflections count. + """ + ctx = super()._loss_inputs(recalc=recalc) + mask_all = ctx.mask.all(dim=0, keepdim=True).expand_as(ctx.mask) + return ctx._replace(mask=mask_all) - F_mean_obs = F_obs_stack.mean(dim=0) # (n_hkl,) - F_calc_mean = F_calc_stack.mean(dim=0) # (n_hkl,) + def _per_refl(self, ctx) -> torch.Tensor: + """Gaussian NLL of the difference-from-mean, per reflection and unreduced.""" + N = len(ctx.keys) + mask_all = ctx.mask[0] # (n_hkl,) -- every row is the same after _loss_inputs - delta_F_obs = F_obs_stack - F_mean_obs - delta_F_calc = F_calc_stack - F_calc_mean + delta_obs = ctx.obs - ctx.obs.mean(dim=0) + delta_calc = ctx.model - ctx.model.mean(dim=0) - # Var(F_i - F_mean) = σ_i²·(1 - 2/N) + (Σ_j σ_j²) / N² - sum_sigma_sq = (sigma_stack**2).sum(dim=0) # (n_hkl,) - sigma_diff_sq = sigma_stack**2 * (1 - 2.0 / N) + sum_sigma_sq / (N**2) - sigma_diff_sq = sigma_diff_sq + self._activation_variance( - all_keys, F_obs_stack - ) - sigma_diff = torch.sqrt(sigma_diff_sq.clamp(min=1e-12)) # (N, n_hkl) + # Var(F_i - F_mean) = sigma_i^2 (1 - 2/N) + (sum_j sigma_j^2) / N^2 + sum_sigma_sq = (ctx.sigma**2).sum(dim=0) + sigma_diff_sq = ctx.sigma**2 * (1 - 2.0 / N) + sum_sigma_sq / (N**2) + sigma_diff_sq = sigma_diff_sq + self._activation_variance(ctx.keys, ctx.obs) + sigma_diff = torch.sqrt(sigma_diff_sq.clamp(min=1e-12)) # Mask via torch.where, not boolean indexing: no nonzero() device sync. - delta_F_obs = torch.where(mask_all, delta_F_obs, torch.zeros_like(delta_F_obs)) - delta_F_calc = torch.where( - mask_all, delta_F_calc, torch.zeros_like(delta_F_calc) - ) + delta_obs = torch.where(mask_all, delta_obs, torch.zeros_like(delta_obs)) + delta_calc = torch.where(mask_all, delta_calc, torch.zeros_like(delta_calc)) sigma_diff = torch.where(mask_all, sigma_diff, torch.ones_like(sigma_diff)) - # Floor sigma at 10% of its median so a zero sigma cannot blow up. - eps = ( - torch.median(sigma_diff[:, mask_all].reshape(-1)) * 1e-1 - if mask_all.any() - else 1e-3 - ) - sigma_safe = sigma_diff.clamp(min=eps) - - diff = delta_F_obs - delta_F_calc - nll = 0.5 * (diff / sigma_safe) ** 2 + torch.log(sigma_safe) + 0.5 * _LOG_2PI + # Floored on the PROPAGATED difference sigma, not the raw measurement sigma -- + # that is the quantity dividing the residual here. + sigma_safe = sigma_diff.clamp(min=self._sigma_floor(sigma_diff, ctx.mask)) + nll = gaussian_per_refl(delta_obs, delta_calc, sigma_safe**2, var_floor=0.0) # A single NaN would poison the whole gradient; 1e6 lets the step be rejected. - nll = torch.where(torch.isfinite(nll), nll, torch.full_like(nll, 1e6)) - - total_nll = (nll * mask_all).sum() - - return total_nll + return torch.where(torch.isfinite(nll), nll, torch.full_like(nll, 1e6)) # ========================================================================= From 51f83b043360eeecd711aac41f200a7d3637d48d Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 16:54:29 +0200 Subject: [PATCH 071/250] Drop the activation-dispersion weighting; --lambda-twin needs --two-moment Ground truth, 8 seeds per regime, paired: inject a known 0.430 A displacement, simulate the two-moment intensity, refine from the dark model. Putting the contamination in the VARIANCE lost 8/8 seeds at high contamination, 95% CI [+0.0066, +0.0102] A, and was null at low. Putting the same quantity in the MEAN -- which --two-moment does -- won 16/16. The mechanism is not subtle: sigma_alpha^2 |dF|^2 in the variance down-weights exactly the reflections where |dF| is largest, which are the ones carrying the difference signal. Structured effects belong in the mean; only genuine measurement noise belongs in the variance. So there is no weighting-only path any more, and --lambda-twin without --two-moment is now an error rather than a silently worse refinement. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- docs/changelog.rst | 6 +- torchref/cli/collection_difference_refine.py | 42 ++- .../refinement/targets/collection/xray.py | 334 +++++------------- 3 files changed, 115 insertions(+), 267 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index fdc4aa08..77d0802b 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,11 @@ Changelog Version 0.6.4 ---------- +- Added ``COLLECTION_XRAY_TARGETS``, the collection target taxonomy, with an intensity difference row +- Removed ``CollectionRiceTarget``, which set ``beta = sigma_obs**2``; the ``ml`` row is the absolute channel instead +- Renamed the kinetic ``xray_weight_rice`` / ``xray/rice`` weight to ``xray_weight_ml`` / ``xray/ml`` +- ``--lambda-twin`` now requires ``--two-moment``; the activation dispersion belongs in the predicted intensity, not in a weight +- Gave the collection targets the same ``_loss_inputs``/``_per_refl`` seam as the single-dataset ones, with the observable declared per row - Added ``--xray-mode nll_i``, a Gaussian on the observed intensities, and an ``observable`` column on the target taxonomy - Added ``DataTarget.get_I_calc_scaled``, so the observable is a choice rather than an assumption - Added ``gaussian_per_refl`` and ``intensity_var_from_sigma_obs``; the amplitude and intensity Gaussians are now one implementation @@ -11,7 +16,6 @@ Version 0.6.4 - Added a reader for CrystFEL ``partialator`` ``.hkl`` reflection lists, via ``ReflectionData.load_crystfel_hkl`` - Added ``FcalcDataset.add_noise`` and the ``torchref.simulate-noisy-data`` CLI, which simulate merged intensities from a structure and report R-split and CC between two independent half-datasets - Simulated intensities keep their negative values; only the derived amplitude is clamped, since clamping the intensity biases the weak reflections upward -- The difference target now inflates its variance by the activation-derived contamination, so a non-zero ``--lambda-twin`` reweights it without needing intensity data - ``CollectionTwoMomentIntensityTarget`` carries a ``base_weight``, calibrated against the difference target's gradient norm so an intensity likelihood does not swamp the restraints - Fixed non-finite observed intensities poisoning the two-moment gradient, which silently froze refinement rather than failing - Added ``CollectionTwoMomentIntensityTarget``, fitting merged intensities as ``|F(alpha)|^2 + sigma_alpha^2 |dF|^2`` to account for crystal-to-crystal spread in activation diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index 1a784531..9c30775f 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -82,7 +82,10 @@ DEFAULT_TARGET_WEIGHTS = { "xray/difference": 1.0, - "xray/rice": 0.0, + # The absolute channel. Zero by default: the difference refinement fixes the + # dark model, so the overall level is already anchored and this term only adds + # the systematic errors the difference cancels. + "xray/ml": 0.0, # Registered only under --two-moment; harmless in the dict either way. "xray/two_moment": 1.0, # "geometry/bond": 1.0, # geometry restraint should never require tuning, so leave at 1.0 @@ -209,9 +212,9 @@ def setup_loss_state(dataset_collection, model_collection, scaler, dataset. Default False. """ from torchref.refinement import LossState - from torchref.experimental.kinetic.targets import ( + from torchref.refinement.targets.collection import ( CollectionDifferenceTarget, - CollectionRiceTarget, + CollectionMLTarget, ) from torchref.refinement.targets import TotalADPTarget, TotalGeometryTarget from torchref.refinement.targets.similarity import CoordinateSimilarityTarget @@ -224,7 +227,7 @@ def setup_loss_state(dataset_collection, model_collection, scaler, diff_target = CollectionDifferenceTarget( dataset_collection, model_collection, scaler=scaler, ) - rice_target = CollectionRiceTarget( + ml_target = CollectionMLTarget( dataset_collection, model_collection, scaler=scaler, ) geom_target = TotalGeometryTarget(model_light) @@ -235,7 +238,7 @@ def setup_loss_state(dataset_collection, model_collection, scaler, ) state.register_target("xray/difference", diff_target) - state.register_target("xray/rice", rice_target) + state.register_target("xray/ml", ml_target) state.register_target("geometry", geom_target) state.register_target("adp", adp_target) state.register_target("similarity", similarity_target) @@ -785,10 +788,9 @@ def main(): "--lambda-twin", type=float, default=0.0, help="Activation dispersion as a fraction of its maximum, in [0, 1]: " "sigma_alpha^2 = alpha (1 - alpha) * lambda. 0 (default) is the " - "coherent model and reproduces the amplitude-only result. On its own " - "this reweights the difference target, down-weighting the reflections " - "whose difference is most contaminated; with --two-moment it also " - "corrects the predicted intensity.", + "coherent model and reproduces the amplitude-only result. Needs " + "--two-moment: the dispersion belongs in the predicted intensity, not " + "in a weight.", ) two_moment.add_argument( "--refine-lambda-twin", action="store_true", default=False, @@ -824,11 +826,19 @@ def main(): file=sys.stderr, ) return 1 - if args.refine_lambda_twin and not args.two_moment: + if (args.lambda_twin > 0.0 or args.refine_lambda_twin) and not args.two_moment: + # There is no longer a weighting-only path. Measured on ground truth (inject a + # known displacement, refine from the dark model, 8 seeds per regime): putting the + # contamination in the VARIANCE lost 8/8 seeds at high contamination, 95% CI + # [+0.0066, +0.0102] A, and was null at low. Putting the same quantity in the MEAN + # -- which is what --two-moment does -- won 16/16. Structured effects belong in the + # mean; only genuine measurement noise belongs in the variance, and down-weighting + # by |dF|^2 suppresses exactly the reflections carrying the difference signal. print( - "Error: --refine-lambda-twin needs --two-moment. The dispersion is only " - "identifiable from the intensity likelihood; through the difference " - "target it enters as a weight and has no gradient of its own.", + "Error: --lambda-twin needs --two-moment. The dispersion enters the predicted " + "intensity, not a weight: as a variance it down-weights the reflections whose " + "difference signal is largest, which measurably worsens the recovered " + "displacement.", file=sys.stderr, ) return 1 @@ -882,12 +892,10 @@ def main(): print(f"Light data: {args.light_structure_factor}") frac_mode = "refinable" if args.refine_fractions else "frozen" print(f"Fractions: dark={fractions[0]}, light={fractions[1]} ({frac_mode})") - if args.lambda_twin > 0.0 or args.two_moment: + if args.two_moment: lam_mode = "refinable" if args.refine_lambda_twin else "fixed" - extra = " + intensity model" if args.two_moment else " (weighting only)" print( - f"Activation spread: lambda_twin={args.lambda_twin} " - f"({lam_mode}){extra}" + f"Activation spread: lambda_twin={args.lambda_twin} ({lam_mode})" ) print(f"Output: {outdir}") print(f"Device: {device}") diff --git a/torchref/refinement/targets/collection/xray.py b/torchref/refinement/targets/collection/xray.py index 5a19e7b1..0ceaba55 100644 --- a/torchref/refinement/targets/collection/xray.py +++ b/torchref/refinement/targets/collection/xray.py @@ -2,13 +2,25 @@ One X-ray likelihood across a paired ``DatasetCollection`` + ``ModelCollection``, keys matched so each timepoint dataset meets its own mixed model: -:class:`CollectionDifferenceTarget` (mean-based differences, the primary optimization -driver), :class:`CollectionRiceTarget` (per-timepoint Rice at ``beta = sigma**2``) and -:class:`CollectionMLTarget` (Rice with one shared Luzzati ``beta`` pooled over all -datasets' free reflections, owned by the target rather than the scaler). -All three inherit the single-dataset subset/masking/R-factor contract from -:class:`~torchref.refinement.targets.collection.base.CollectionXrayTarget`. +:class:`CollectionDifferenceTarget` + Mean-based differences on amplitudes; the primary optimization driver. +:class:`CollectionDifferenceIntensityTarget` + The same on intensities -- the whole class is one ``observable`` declaration. +:class:`CollectionMLTarget` + Read MLF per dataset at one shared Luzzati ``beta``, pooled over every dataset's free + reflections and owned by the target rather than the scaler. The absolute channel. + +All of them get their observations, model, mask and likelihood seam from +:class:`~torchref.refinement.targets.collection.base.CollectionXrayTarget`, so each row is +a ``_per_refl`` and nothing else. The selectable set is +:data:`~torchref.refinement.targets.collection._specs.COLLECTION_XRAY_TARGETS`. + +The retired ``CollectionRiceTarget`` set ``beta = sigma_obs**2``, pairing a measurement +sigma with a Rice ``Sigma``. That asserts an isotropic *complex* error where ``sigma_obs`` +carries no phase at all; :mod:`torchref.base.targets.xray_likelihoods` records that no +regime makes it correct, and the single-dataset table deliberately offers no such row. +:class:`CollectionMLTarget` replaces it. """ from typing import TYPE_CHECKING, Dict @@ -20,13 +32,13 @@ from torchref.base.targets.xray_likelihoods import ( complex_var_from_beta, gaussian_per_refl, - rice_math, + rice_per_refl, ) from torchref.refinement.model_error_estimation.sigma_a import SigmaAEstimator, epsilon_from_hkl from torchref.utils.stats import VERBOSITY_STANDARD, StatEntry, stat from ._util import _LOG_2PI, _scale_fcalc -from .base import CollectionXrayTarget +from .base import CollectionSigmaALossInputs, CollectionXrayTarget if TYPE_CHECKING: from torchref.io.datasets.collection import DatasetCollection @@ -104,53 +116,6 @@ def __init__( ) self.normalize = normalize - def _activation_variance(self, keys, F_obs_stack) -> torch.Tensor: - """Extra variance from crystal-to-crystal spread in activation. - - A merged light intensity carries a positive, phase-blind contamination - ``sigma_alpha^2 |dF/dalpha|^2``. Propagated onto the amplitude it becomes a shift - of ``sigma_alpha^2 |dF|^2 / (2 |F|)``, which is a systematic of **known magnitude - but unmodelled here**, so it enters as a variance and down-weights exactly the - reflections whose difference is most contaminated. - - The resulting weight, ``sigma_meas^2 / (sigma_meas^2 + this)``, is the calibrated - form of the empirical k-weighting that difference maps apply by hand -- derived - from a fitted or assumed dispersion rather than tuned. - - Returns zeros when the dispersion is zero, so the loss is then unchanged. - - Parameters - ---------- - keys : list of str - Datasets in the order they are stacked. - F_obs_stack : torch.Tensor - Observed amplitudes, shape ``(N, n_hkl)``, used as the propagation denominator. - - Returns - ------- - torch.Tensor - Variance to add, shape ``(N, n_hkl)``; a scalar zero when inactive. - """ - mc = self._model_collection - sigma_alpha_sq = getattr(mc, "sigma_alpha_sq", None) - if sigma_alpha_sq is None: - return torch.zeros((), device=F_obs_stack.device) - if float(sigma_alpha_sq) == 0.0: - return torch.zeros((), device=F_obs_stack.device) - - dc = self._dataset_collection - rows = [mc.keys().index(k) for k in keys] - components = dc.component_structure_factors(mc, recalc=False) - jacobian = mc.activation_jacobian()[rows] - - derivative = mc.mix_component_fcalcs(components, jacobian) - if self._scaler is not None and hasattr(self._scaler, "forward_batched"): - derivative = self._scaler.forward_batched(derivative, jacobian) - - contamination = sigma_alpha_sq * derivative.abs() ** 2 - shift = contamination / (2.0 * F_obs_stack.abs().clamp(min=1e-6)) - return shift**2 - def _loss_inputs(self, recalc: bool = False): """The base's stack, with the mask narrowed across datasets. @@ -175,7 +140,6 @@ def _per_refl(self, ctx) -> torch.Tensor: # Var(F_i - F_mean) = sigma_i^2 (1 - 2/N) + (sum_j sigma_j^2) / N^2 sum_sigma_sq = (ctx.sigma**2).sum(dim=0) sigma_diff_sq = ctx.sigma**2 * (1 - 2.0 / N) + sum_sigma_sq / (N**2) - sigma_diff_sq = sigma_diff_sq + self._activation_variance(ctx.keys, ctx.obs) sigma_diff = torch.sqrt(sigma_diff_sq.clamp(min=1e-12)) # Mask via torch.where, not boolean indexing: no nonzero() device sync. @@ -192,131 +156,29 @@ def _per_refl(self, ctx) -> torch.Tensor: return torch.where(torch.isfinite(nll), nll, torch.full_like(nll, 1e6)) -# ========================================================================= -# CollectionRiceTarget -# ========================================================================= +class CollectionDifferenceIntensityTarget(CollectionDifferenceTarget): + """The difference-from-mean target on **intensities** instead of amplitudes. + A one-line row: the base reads ``I``/``sigI`` and predicts ``|F_calc|**2``, and the + difference-from-mean algebra is observable-agnostic -- + ``Var(x_i - x_mean) = sigma_i^2 (1 - 2/N) + (sum_j sigma_j^2)/N^2`` holds for any + quantity whose members share a mean. So the whole row is the ``observable`` + declaration, which is the point of having the axis at all. -class CollectionRiceTarget(CollectionXrayTarget): - """ - Multi-timepoint Rice maximum-likelihood amplitude target. - - Rice NLL for acentrics and the folded-normal form for centrics, per timepoint. - The per-timepoint losses are independent, so reflections and masks are - concatenated across timepoints and masked once rather than reduced in a - Python loop. + Prefer it over :class:`CollectionDifferenceTarget` when the difference signal is weak + relative to the measurement error, because ``F_obs`` on a merged dataset is a + French-Wilson posterior: strictly positive, so it reshapes exactly the weak reflections + a small difference lives in, and it cannot represent a negative intensity at all. + Prefer the amplitude row when the output difference *map* is the product, since the DED + coefficients are amplitudes and keeping the loss and the map in one space is one fewer + conversion to get wrong. - Parameters - ---------- - dataset_collection : DatasetCollection - model_collection : ModelCollection - scaler : ScalerBase, optional - Single scaler applied to each timepoint's F_calc. - normalize : bool - Unused placeholder. ``forward`` always returns the unnormalised summed - NLL regardless of this flag. - use_work_set : bool - Legacy bool; superseded by ``use_set``. If True, loss on the work set. - use_set : str, optional - Canonical 3-way subset selector ``"work"``/``"free"``/``"val"``. - verbose : int - Verbosity level. + Both rows are offered rather than one chosen: which wins is a property of a dataset's + signal-to-noise, not something to settle once in the library. """ - name: str = "collection_rice_xray" - - def __init__( - self, - dataset_collection: "DatasetCollection", - model_collection: "ModelCollection", - scaler: "ScalerBase" = None, - normalize: bool = True, - use_work_set: bool = True, - use_set: str = None, - verbose: int = 0, - ): - super().__init__( - dataset_collection, - model_collection, - scaler=scaler, - use_work_set=use_work_set, - use_set=use_set, - verbose=verbose, - ) - self.normalize = normalize - - def _keys(self): - """Rice fits the timepoints only — the dark reference is excluded.""" - mc = self._model_collection - dc = self._dataset_collection - return [n for n in mc.timepoint_names if n in dc] - - def forward(self) -> torch.Tensor: - """Summed Rice NLL over every timepoint's reflections in this subset.""" - dc = self._dataset_collection - mc = self._model_collection - - tp_names = self._keys() - if not tp_names: - return torch.tensor(0.0, device=mc.device) - - self._reset_model_caches() - - fo_parts, fc_parts, sig_parts, cen_parts, mask_parts = [], [], [], [], [] - for tp_name in tp_names: - data = dc[tp_name] - model = mc[tp_name] - - F_obs, sigma = data.get_corrected_data() - F_calc = self._scaled_amp_full(data, model, recalc=False) - centric = data.centric - if centric is None: - centric = torch.zeros( - len(F_obs), dtype=torch.bool, device=F_obs.device - ) - - fo_parts.append(F_obs) - fc_parts.append(F_calc) - sig_parts.append(sigma) - cen_parts.append(centric) - mask_parts.append(self._subset(data).mask) - - # Mask the flat arrays once: compact, no per-dataset Python reduction. - mask = torch.cat(mask_parts) - F_obs = torch.cat(fo_parts)[mask] - F_calc = torch.cat(fc_parts)[mask] - sigma = torch.cat(sig_parts)[mask] - centric = torch.cat(cen_parts)[mask] - - if F_obs.numel() == 0: - return torch.tensor(0.0, device=mc.device) - - # Plain Rice: the model-error variance IS the measurement variance here. - beta = sigma**2 - eb = beta.clamp(min=1e-6) - - # --- Acentric Rice NLL --- - term1 = -torch.log(2 * F_obs / eb + 1e-12) - term2 = F_obs**2 / eb - term3 = F_calc**2 / eb - arg_bessel = (2 * F_obs * F_calc / eb).clamp(max=1e6) - term4 = -(torch.log(torch.special.i0e(arg_bessel) + 1e-12) + arg_bessel) - loss_acentric = term1 + term2 + term3 + term4 - - # --- Centric NLL --- - term1_c = -0.5 * torch.log(2 / (np.pi * eb) + 1e-12) - term2_c = F_obs**2 / (2 * eb) - term3_c = F_calc**2 / (2 * eb) - term4_c = -(F_obs * F_calc) / eb - arg_exp = (-2 * F_obs * F_calc / eb).clamp(min=-80.0, max=80.0) - term5_c = -torch.log((1 + torch.exp(arg_exp)) / 2 + 1e-12) - loss_centric = term1_c + term2_c + term3_c + term4_c + term5_c - - loss = torch.where(centric, loss_centric, loss_acentric) - # A single NaN would poison the whole gradient; 1e6 lets the step be rejected. - loss = torch.where(torch.isfinite(loss), loss, torch.full_like(loss, 1e6)) - - return loss.sum() + name: str = "difference_intensity_xray" + observable: str = "intensity" # ========================================================================= @@ -414,89 +276,63 @@ def _common_geom(self): self._eps_common, self._dss_common, self._geom_key = eps, dss, key return self._eps_common, self._dss_common - def forward(self) -> torch.Tensor: - """Summed Read-MLF loss over all datasets at the shared beta, work-set weighted.""" - dc = self._dataset_collection - mc = self._model_collection + def _loss_inputs(self, recalc: bool = False): + """The base's stack, plus the shared model-error variance this row needs. - all_keys = self._keys() - if not all_keys: - return torch.tensor(0.0, device=mc.device) + ``beta`` and ``epsilon`` are fitted **once** on the pooled free reflections of + every data-model pair and mapped onto the common HKL, so one per-reflection + variance serves every dataset -- and since they live on that common grid they + broadcast over the dataset axis rather than being tiled. + Detached, so gradients reach the models only through ``F_calc``. Fitted on the + free set (``data.free`` excludes validation), and cached until + :meth:`maintenance` resets it. + """ + ctx = super()._loss_inputs(recalc=recalc) + dc = self._dataset_collection eps_common, dss_common = self._common_geom() - dtype = dss_common.dtype - - # Clear any structure-factor cache populated under no_grad (e.g. by the - # scaler's joint initialization or a preceding stats() call) so the - # F_calc computed below carries a live graph. - self._reset_model_caches() - - # Each pair's F_calc is computed once WITH grad for the loss; detached copies - # feed the pooled single-beta estimate. - fo_parts, fc_parts, cen_parts, msk_parts = [], [], [], [] - eps_parts, dss_parts, free_parts, sig_parts = [], [], [], [] - for key in all_keys: - data = dc[key] - model = mc[key] - - F_obs, sig_obs = data.get_corrected_data() - F_obs = F_obs.to(dtype) - sig_parts.append(sig_obs.to(dtype).reshape(-1)) - F_calc = self._scaled_amp_full(data, model, recalc=False).to(dtype) - centric = data.centric - if centric is None: - centric = torch.zeros( - len(F_obs), dtype=torch.bool, device=F_obs.device - ) - - fo_parts.append(F_obs) - fc_parts.append(F_calc) - cen_parts.append(centric) - msk_parts.append(self._subset(data).mask) - - # Beta is estimated on the free set; data.free excludes validation. - free_parts.append(data.free.mask) - eps_parts.append(eps_common.to(dtype)) - dss_parts.append(dss_common.to(dtype)) - - # One shared (beta, epsilon) for all datasets, mapped onto the common HKL via - # target_dss. Detached; cached until maintenance() resets it. - _est = self._sigma_a.get( - torch.cat(fo_parts), - torch.cat([fc.detach() for fc in fc_parts]), # beta needs no gradient - torch.cat(cen_parts), - torch.cat(eps_parts), - torch.cat(dss_parts), - torch.cat(free_parts), + dtype = ctx.obs.dtype + + centric = dc.get_centric_flags() + if centric is None: + centric = torch.zeros( + ctx.obs.shape[-1], dtype=torch.bool, device=ctx.obs.device + ) + + n_ds = len(ctx.keys) + free = torch.cat([dc[k].free.mask for k in ctx.keys]) + est = self._sigma_a.get( + ctx.obs.reshape(-1).to(dtype), + # beta needs no gradient. + ctx.model.detach().reshape(-1).to(dtype), + centric.repeat(n_ds), + eps_common.to(dtype).repeat(n_ds), + dss_common.to(dtype).repeat(n_ds), + free, out_epsilon=eps_common.to(dtype), target_dss=dss_common, # Always passed, as at every other call site: it is what makes sigma_A the # correlation with the noise-free amplitudes rather than with the noisy data. - sigma_obs=torch.cat(sig_parts), + sigma_obs=ctx.sigma.reshape(-1).to(dtype), ) - beta, eps = _est.beta, _est.epsilon - - # beta/eps live on the common HKL, so they are tiled once per dataset to line - # up with the concatenation order. - n_ds = len(all_keys) - F_obs_cat = torch.cat(fo_parts) - F_calc_cat = torch.cat(fc_parts) - centric_cat = torch.cat(cen_parts) - mask_cat = torch.cat(msk_parts) - beta_cat = beta.to(F_obs_cat.dtype).repeat(n_ds) - eps_cat = eps.to(F_obs_cat.dtype).repeat(n_ds) if eps is not None else None - - # TOTAL variance (`est.beta`, not `beta_model`): this likelihood does not - # account for sigma_obs itself, so the measurement variance must stay inside beta. - total = rice_math( - F_obs_cat, F_calc_cat, complex_var_from_beta(beta_cat, eps_cat), - centric_cat, mask=mask_cat, + return CollectionSigmaALossInputs( + *ctx, + centric=centric, + beta=est.beta.to(dtype), + epsilon=None if est.epsilon is None else est.epsilon.to(dtype), ) - # Base weight drives refinement; applied on the work set only. - if self.use_work_set: - total = self.base_weight * total - return total + def _per_refl(self, ctx) -> torch.Tensor: + """Read MLF per reflection: Rice for acentrics, folded normal for centrics. + + TOTAL variance (``est.beta``, not ``beta_model``): this likelihood does not + account for ``sigma_obs`` itself, so the measurement variance stays inside + ``beta``. + """ + Sigma = complex_var_from_beta(ctx.beta, ctx.epsilon) + nll = rice_per_refl(ctx.obs, ctx.model, Sigma, ctx.centric) + # A single NaN would poison the whole gradient; 1e6 lets the step be rejected. + return torch.where(torch.isfinite(nll), nll, torch.full_like(nll, 1e6)) def maintenance(self) -> None: """Invalidate the shared beta so it is re-estimated from the updated From 8f44cb8736ef49c3d0be3f5d340ad6481a3a6e2a Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 16:54:47 +0200 Subject: [PATCH 072/250] Add the collection target taxonomy; replace the Rice row with ml COLLECTION_XRAY_TARGETS mirrors XRAY_TARGETS: one class per row checked at import, the observable declared per row and checked against the class. Four rows -- difference, difference_i, two_moment, ml. Both difference rows ship. The intensity one is nothing but an `observable` declaration over the amplitude one, which is the abstraction paying for itself: the difference-from-mean covariance propagation does not care what the observable is. Which row is better is a property of a dataset's signal-to-noise, not something to settle in the library. CollectionRiceTarget is gone. It set `beta = sigma_obs**2`, pairing a measurement sigma with a Rice `Sigma` -- asserting an isotropic *complex* error where sigma_obs carries no phase at all. xray_likelihoods records that no regime makes that correct, and the single-dataset table refuses to offer such a row. It was also the kinetic path's only absolute anchor, so CollectionMLTarget (one shared Luzzati beta, pooled over every dataset's free reflections) takes that role: `xray_weight_rice` becomes `xray_weight_ml`, `xray/rice` becomes `xray/ml`. Net -177 lines of production code across the series, and no row hand-writes a forward any more -- pinned by a test. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- docs/user_guide/targets.rst | 28 +++ tests/helpers/device_cases.py | 2 +- .../refinement/test_activation_weighting.py | 174 ------------------ ...test_collection_target_characterisation.py | 42 +++-- .../refinement/test_collection_taxonomy.py | 152 +++++++++++++++ tests/unit/refinement/test_ml_sigmaa.py | 29 ++- torchref/experimental/kinetic/__init__.py | 4 +- torchref/experimental/kinetic/refinement.py | 28 +-- torchref/experimental/kinetic/targets.py | 6 +- torchref/refinement/targets/__init__.py | 8 +- .../refinement/targets/collection/__init__.py | 20 +- .../refinement/targets/collection/_specs.py | 166 +++++++++++++++++ .../refinement/targets/collection/base.py | 27 ++- 13 files changed, 467 insertions(+), 219 deletions(-) delete mode 100644 tests/unit/refinement/test_activation_weighting.py create mode 100644 tests/unit/refinement/test_collection_taxonomy.py create mode 100644 torchref/refinement/targets/collection/_specs.py diff --git a/docs/user_guide/targets.rst b/docs/user_guide/targets.rst index 66c30c4d..a7c1fda2 100644 --- a/docs/user_guide/targets.rst +++ b/docs/user_guide/targets.rst @@ -65,6 +65,34 @@ primitive rather than a different variance. R-factors are reported on amplitudes for every row regardless, so they stay comparable across the whole table. Intensity rows are not admissible as ``--scale-target``, which fails closed on them. +Collection X-ray Targets +------------------------ + +The multi-dataset analogues, for time-resolved and difference refinement. Same +taxonomy shape as above — ``COLLECTION_XRAY_TARGETS`` in +``torchref.refinement.targets.collection._specs``, one class per row, the +observable declared per row — and the same ``_loss_inputs`` / ``_per_refl`` seam, +batched over ``(n_datasets, n_hkl)`` on the collection's common HKL grid. + +- ``difference`` — Gaussian on each dataset's **amplitude** difference from the + collection mean, with the dataset/mean covariance propagated. The primary + optimization driver for difference refinement. +- ``difference_i`` — the same on **intensities**. The entire class is one + ``observable`` declaration: the difference-from-mean algebra does not care what + the observable is. +- ``two_moment`` — merged **intensities** as + :math:`|F(\bar\alpha)|^2 + \sigma_\alpha^2 |\Delta F|^2`, accounting for + crystal-to-crystal spread in activation. +- ``ml`` — Read MLF per dataset at one shared Luzzati :math:`\beta`, fitted on + the pooled free reflections of every data–model pair. The **absolute** channel: + with K free base models a purely relative loss leaves the overall level + unconstrained. + +Both difference rows are offered rather than one being chosen. Amplitudes keep the +loss in the same space as the output DED map coefficients; intensities avoid the +French–Wilson posterior reshaping the weak tail a small difference lives in. Which +wins is a property of a dataset's signal-to-noise. + Geometry Targets ---------------- diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index cbae4c6e..5c58f45f 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -300,7 +300,7 @@ class TargetDeviceCase: "CollectionDifferenceTarget": "needs a dataset collection", "CollectionTwoMomentIntensityTarget": "needs a dataset collection", "CollectionMLTarget": "needs a dataset collection", - "CollectionRiceTarget": "needs a dataset collection", + "CollectionDifferenceIntensityTarget": "needs a dataset collection", "ADPSigdTarget": "needs a model with ADPs", "AngleTarget": "needs a model with restraints", "ChiralTarget": "needs a model with restraints", diff --git a/tests/unit/refinement/test_activation_weighting.py b/tests/unit/refinement/test_activation_weighting.py deleted file mode 100644 index 5e646bf3..00000000 --- a/tests/unit/refinement/test_activation_weighting.py +++ /dev/null @@ -1,174 +0,0 @@ -"""The difference target's per-reflection weighting must follow the activation spread. - -A merged light intensity carries a positive, phase-blind contamination -``sigma_alpha^2 |dF/dalpha|^2``. Propagated onto the amplitude it is a shift of -``sigma_alpha^2 |dF|^2 / (2 |F|)``. The difference target does not model that shift, so it -enters as a variance -- which down-weights exactly the reflections whose difference is most -contaminated, and is the calibrated form of the k-weighting difference maps apply by hand. - -Two properties are pinned: a zero dispersion changes nothing at all, and a non-zero one -reweights in proportion to ``|dF|^2`` rather than uniformly. The second is what separates a -real weighting from an overall rescaling of the x-ray term, which would be -indistinguishable from a change of x-ray weight. -""" - -import pytest -import torch - - -@pytest.fixture(scope="module") -def collection(pdb_dir, mtz_dir): - """A dark/light pair on 1DAW with a displaced light model, so dF is non-zero.""" - pdb = pdb_dir / "1DAW.pdb" - mtz = mtz_dir / "1DAW.mtz" - if not (pdb.exists() and mtz.exists()): - pytest.skip("1DAW fixture not present") - - from torchref import ReflectionData - from torchref.cli._common import load_model - from torchref.io.datasets.collection import DatasetCollection - from torchref.model.model_collection import ModelCollection - from torchref.scaling.collection_scaler import CollectionScaler - - d_min = 2.05 - dark = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - light = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - - model_dark = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) - model_light = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) - with torch.no_grad(): - model_light.xyz.refinable_params += 0.25 - - dc = DatasetCollection(verbose=0, device="cpu") - dc.add_dataset("dark", dark, set_as_reference=True) - dc.add_dataset("light", light) - - mc = ModelCollection([model_dark, model_light], dark_key="dark", verbose=0) - mc.add_dark() - mc.add_timepoint("light", [0.78, 0.22]) - - scaler = CollectionScaler(dc, mc, verbose=0) - scaler.initialize() - return dc, mc, scaler - - -def _target(dc, mc, scaler): - from torchref.refinement.targets import CollectionDifferenceTarget - - return CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0) - - -@pytest.mark.integration -class TestZeroDispersionIsInert: - def test_the_extra_variance_is_exactly_zero(self, collection): - dc, mc, scaler = collection - mc.set_lambda_twin(0.0) - target = _target(dc, mc, scaler) - - F_obs = dc.stack_F_obs(target._keys()) - extra = target._activation_variance(target._keys(), F_obs) - assert float(extra.abs().max()) == 0.0 - - def test_the_loss_is_unchanged_from_the_no_dispersion_baseline(self, collection): - """Back-compat: lambda = 0 must reproduce the loss the target always returned.""" - dc, mc, scaler = collection - mc.set_lambda_twin(0.0) - target = _target(dc, mc, scaler) - first = target.forward().item() - second = target.forward().item() - assert second == pytest.approx(first, rel=1e-4) - - -@pytest.mark.integration -class TestDispersionReweights: - def test_a_nonzero_dispersion_changes_the_loss(self, collection): - dc, mc, scaler = collection - mc.set_lambda_twin(0.0) - base = _target(dc, mc, scaler).forward().item() - - mc.set_lambda_twin(0.4) - try: - weighted = _target(dc, mc, scaler).forward().item() - finally: - mc.set_lambda_twin(0.0) - - assert weighted != pytest.approx(base, rel=1e-3) - - def test_the_extra_variance_tracks_the_squared_difference(self, collection): - """Proportional to |dF|^2, not uniform. - - A uniform inflation would just rescale the x-ray term and be - indistinguishable from a change of x-ray weight; the point of this weighting is - that it is reflection-specific. - """ - dc, mc, scaler = collection - mc.set_lambda_twin(0.4) - try: - target = _target(dc, mc, scaler) - keys = target._keys() - F_obs = dc.stack_F_obs(keys) - extra = target._activation_variance(keys, F_obs) - - rows = [mc.keys().index(k) for k in keys] - components = dc.component_structure_factors(mc, recalc=False) - jacobian = mc.activation_jacobian()[rows] - deriv = scaler.forward_batched( - mc.mix_component_fcalcs(components, jacobian), jacobian - ) - expected = ( - mc.sigma_alpha_sq * deriv.abs() ** 2 - / (2.0 * F_obs.abs().clamp(min=1e-6)) - ) ** 2 - assert torch.allclose(extra, expected, rtol=1e-5) - - light = extra[keys.index("light")] - assert float(light.max()) > 0.0 - # Genuinely non-uniform across reflections. - assert float(light.std() / light.mean().clamp(min=1e-30)) > 0.5 - finally: - mc.set_lambda_twin(0.0) - - def test_the_reference_row_is_unweighted(self, collection): - """The dark's activation Jacobian is exactly zero, so it carries no - contamination and must keep its measured sigma.""" - dc, mc, scaler = collection - mc.set_lambda_twin(0.6) - try: - target = _target(dc, mc, scaler) - keys = target._keys() - extra = target._activation_variance(keys, dc.stack_F_obs(keys)) - assert float(extra[keys.index("dark")].abs().max()) == 0.0 - assert float(extra[keys.index("light")].abs().max()) > 0.0 - finally: - mc.set_lambda_twin(0.0) - - @pytest.mark.parametrize("lam", [0.1, 0.4, 0.9]) - def test_more_dispersion_means_more_down_weighting(self, collection, lam): - """The implied weight sigma^2/(sigma^2 + extra) must fall monotonically.""" - dc, mc, scaler = collection - keys = _target(dc, mc, scaler)._keys() - F_obs = dc.stack_F_obs(keys) - - mc.set_lambda_twin(lam) - try: - extra = _target(dc, mc, scaler)._activation_variance(keys, F_obs) - finally: - mc.set_lambda_twin(0.0) - - sigma_sq = dc.stack_F_sigma(keys) ** 2 - weight = sigma_sq / (sigma_sq + extra) - light = weight[keys.index("light")] - assert float(light.max()) <= 1.0 + 1e-6 - assert float(light.min()) < 1.0, "no reflection was down-weighted at all" - - def test_the_gradient_still_reaches_the_model(self, collection): - dc, mc, scaler = collection - mc.set_lambda_twin(0.3) - try: - _target(dc, mc, scaler).forward().backward() - grad = mc.base_models[1].xyz.refinable_params.grad - assert grad is not None and torch.isfinite(grad).all() - assert float(grad.abs().max()) > 0 - finally: - mc.base_models[1].xyz.refinable_params.grad = None - mc.set_lambda_twin(0.0) diff --git a/tests/unit/refinement/test_collection_target_characterisation.py b/tests/unit/refinement/test_collection_target_characterisation.py index af5c5f00..b82c673f 100644 --- a/tests/unit/refinement/test_collection_target_characterisation.py +++ b/tests/unit/refinement/test_collection_target_characterisation.py @@ -78,14 +78,16 @@ def collection(pdb_dir, mtz_dir): def _targets(dc, mc, scaler): from torchref.refinement.targets import ( + CollectionDifferenceIntensityTarget, CollectionDifferenceTarget, CollectionMLTarget, - CollectionRiceTarget, ) return { "difference": CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0), - "rice": CollectionRiceTarget(dc, mc, scaler=scaler, verbose=0), + "difference_i": CollectionDifferenceIntensityTarget( + dc, mc, scaler=scaler, verbose=0 + ), "ml": CollectionMLTarget(dc, mc, scaler=scaler, verbose=0), } @@ -99,7 +101,7 @@ class TestObservedAmplitudesAreScaled: blind to it, so this is a direct test of which accessor is in use. """ - @pytest.mark.parametrize("name", ["difference", "rice", "ml"]) + @pytest.mark.parametrize("name", ["difference", "difference_i", "ml"]) def test_loss_responds_to_the_datasets_own_log_scale(self, collection, name): dc, mc, scaler = collection target = _targets(dc, mc, scaler)[name] @@ -202,28 +204,46 @@ class TestLossesAreSummedNotAveraged: -- it looks exactly like a change of X-ray weight. """ - def test_adding_a_dataset_grows_the_rice_loss(self, collection, pdb_dir, mtz_dir): + def test_adding_a_dataset_grows_the_absolute_loss(self, collection, pdb_dir, mtz_dir): + """The expected ratio is n_after / n_before, and that is 3/2, not 2. + + The fixture already holds two datasets (dark + light), so adding a third takes + the absolute target from 2 to 3. This test used to expect 2.0 because it ran on + ``CollectionRiceTarget``, which overrode ``_keys()`` to drop the dark reference + and so went from 1 to 2. ``ml`` fits every dataset including the dark. + + A meaned target would stay near 1.0 either way, which is what this is for. + """ from torchref import ReflectionData - from torchref.refinement.targets import CollectionRiceTarget + from torchref.refinement.targets import CollectionMLTarget dc, mc, scaler = collection - one = CollectionRiceTarget(dc, mc, scaler=scaler, verbose=0).forward().item() + target_before = CollectionMLTarget(dc, mc, scaler=scaler, verbose=0) + n_before = len(target_before._keys()) + one = target_before.forward().item() extra = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz_dir / "1DAW.mtz")) dc.add_dataset("light2", extra) mc.add_timepoint("light2", [0.7, 0.3]) try: - two = CollectionRiceTarget(dc, mc, scaler=scaler, verbose=0).forward().item() + target_after = CollectionMLTarget(dc, mc, scaler=scaler, verbose=0) + n_after = len(target_after._keys()) + two = target_after.forward().item() finally: dc._datasets.pop("light2") dc._dataset_order.remove("light2") del mc._timepoints["light2"] mc._order.remove("light2") + assert (n_before, n_after) == (2, 3) ratio = two / one - assert ratio == pytest.approx(2.0, rel=0.15), ( - f"two timepoints gave {ratio:.3f}x one timepoint's Rice loss; a summed " - f"target should roughly double and a meaned one stay near 1.0" + # Tolerance is loose because the shared Luzzati beta is REFITTED on the pooled + # free reflections of the larger collection, so the per-reflection loss moves a + # little too. That is a property of the target, not slack: the two hypotheses + # this test separates are 1.5 and 1.0, which are far apart. + assert ratio == pytest.approx(n_after / n_before, rel=0.15), ( + f"{n_after} datasets gave {ratio:.3f}x the loss of {n_before}; a summed " + f"target should scale with the count and a meaned one stay near 1.0" ) @@ -231,7 +251,7 @@ def test_adding_a_dataset_grows_the_rice_loss(self, collection, pdb_dir, mtz_dir class TestReportedNumbers: """The shape of what ``get_rfactor`` / ``stats`` promise, plus reproducibility.""" - @pytest.mark.parametrize("name", ["difference", "rice", "ml"]) + @pytest.mark.parametrize("name", ["difference", "difference_i", "ml"]) def test_forward_is_finite_and_reproducible(self, collection, name): dc, mc, scaler = collection target = _targets(dc, mc, scaler)[name] diff --git a/tests/unit/refinement/test_collection_taxonomy.py b/tests/unit/refinement/test_collection_taxonomy.py new file mode 100644 index 00000000..cf2a98b5 --- /dev/null +++ b/tests/unit/refinement/test_collection_taxonomy.py @@ -0,0 +1,152 @@ +"""The collection X-ray taxonomy, as contracts. + +Mirrors ``tests/unit/refinement/test_nll_beta.py``'s registry tests for the multi-dataset +table. Same thesis, same invariants: one class per row, the observable declared rather than +passed, and no row that pairs an amplitude distribution with intensities. +""" + +import inspect + +import pytest + + +@pytest.mark.unit +def test_each_row_has_its_own_class(): + """One class per selectable row, and the test and the table must agree on the set.""" + from torchref.refinement.targets.collection import COLLECTION_XRAY_TARGETS + from torchref.refinement.targets.collection.intensity import ( + CollectionTwoMomentIntensityTarget, + ) + from torchref.refinement.targets.collection.xray import ( + CollectionDifferenceIntensityTarget, + CollectionDifferenceTarget, + CollectionMLTarget, + ) + + expected = { + "difference": CollectionDifferenceTarget, + "difference_i": CollectionDifferenceIntensityTarget, + "two_moment": CollectionTwoMomentIntensityTarget, + "ml": CollectionMLTarget, + } + assert set(expected) == set(COLLECTION_XRAY_TARGETS.names), ( + "table and test disagree on the rows" + ) + for name, cls in expected.items(): + assert COLLECTION_XRAY_TARGETS.by_name(name).target_cls is cls, name + + seen = {} + for name in COLLECTION_XRAY_TARGETS.names: + cls = COLLECTION_XRAY_TARGETS.by_name(name).target_cls + assert cls not in seen, f"{name} and {seen[cls]} share {cls.__name__}" + seen[cls] = name + + +@pytest.mark.unit +def test_the_observable_is_declared_not_passed(): + """The spec's claim and the class's own attribute must agree. + + A row advertising intensities while reading amplitudes would be wrong by ``2|F|``, + which is resolution-dependent -- so it reads as a scale or B error rather than as a + bug, and nothing downstream would flag it. + """ + from torchref.refinement.targets.collection import ( + COLLECTION_XRAY_TARGETS, + CollectionXrayTargetSpec, + ) + from torchref.refinement.targets.collection.xray import CollectionDifferenceTarget + + by_obs = {} + for spec in COLLECTION_XRAY_TARGETS.specs: + assert spec.observable in ("amplitude", "intensity"), spec.name + assert getattr(spec.target_cls, "observable", "amplitude") == spec.observable + by_obs.setdefault(spec.observable, []).append(spec.name) + assert "observable" not in inspect.signature( + spec.target_cls.__init__ + ).parameters + + assert set(by_obs["intensity"]) == {"difference_i", "two_moment"} + assert set(by_obs["amplitude"]) == {"difference", "ml"} + + with pytest.raises(ValueError, match="observable"): + CollectionXrayTargetSpec( + name="bogus", + target_cls=CollectionDifferenceTarget, + doc="", + observable="intensity", + ) + + +@pytest.mark.unit +def test_both_difference_observables_are_offered(): + """Neither difference row is privileged. + + Which one is better is a property of a dataset's signal-to-noise: amplitudes keep the + loss in the same space as the output DED coefficients, intensities avoid the + French-Wilson posterior reshaping the weak tail. The abstraction exists so that + carrying both is cheap -- and the intensity row proves it, being nothing but an + ``observable`` declaration over the amplitude one. + """ + from torchref.refinement.targets.collection import COLLECTION_XRAY_TARGETS + from torchref.refinement.targets.collection.xray import ( + CollectionDifferenceIntensityTarget, + CollectionDifferenceTarget, + ) + + assert {"difference", "difference_i"} <= set(COLLECTION_XRAY_TARGETS.names) + assert issubclass(CollectionDifferenceIntensityTarget, CollectionDifferenceTarget) + # The subclass adds no likelihood of its own: only the name and the observable. + own = set(vars(CollectionDifferenceIntensityTarget)) - { + # `__annotations__` is present because `name`/`observable` are annotated + # assignments, not because the class defines behaviour. + "__doc__", "__module__", "__qualname__", "__annotations__", + "name", "observable", + } + assert not own, f"the intensity difference row grew a body: {sorted(own)}" + + +@pytest.mark.unit +def test_there_is_no_intensity_rice_row(): + """Rice is amplitude-only by nature, so the axis is not square. + + Rice and the folded normal are distributions *of an amplitude*; the intensity analogue + is the exponential / chi-square_1 Wilson distribution, a different primitive rather + than a different variance. A row pairing the ML class with intensities would be a + modelling error, not a new feature. + """ + from torchref.refinement.targets.collection import COLLECTION_XRAY_TARGETS + from torchref.refinement.targets.collection.xray import CollectionMLTarget + + for spec in COLLECTION_XRAY_TARGETS.specs: + if spec.observable == "intensity": + assert not issubclass(spec.target_cls, CollectionMLTarget), spec.name + + +@pytest.mark.unit +def test_unknown_rows_fail_closed(): + from torchref.refinement.targets.collection import COLLECTION_XRAY_TARGETS + + with pytest.raises(ValueError, match="Unknown collection X-ray target"): + COLLECTION_XRAY_TARGETS.by_name("no_such_row") + + +@pytest.mark.unit +def test_every_row_goes_through_the_seam(): + """No row may hand-write a ``forward``. + + That is what this refactor bought: 265 lines of per-row forwards, each re-deriving + stacking, masking, the sigma floor and the summing, collapsed into one. A row that + reintroduces its own ``forward`` also reintroduces the possibility of it disagreeing + with ``residuals``, which nothing else would notice. + """ + from torchref.refinement.targets.collection import COLLECTION_XRAY_TARGETS + from torchref.refinement.targets.collection.base import CollectionXrayTarget + + for spec in COLLECTION_XRAY_TARGETS.specs: + cls = spec.target_cls + assert cls.forward is CollectionXrayTarget.forward, ( + f"{spec.name} overrides forward; the likelihood belongs in _per_refl" + ) + assert cls._per_refl is not CollectionXrayTarget._per_refl, ( + f"{spec.name} has no _per_refl of its own" + ) diff --git a/tests/unit/refinement/test_ml_sigmaa.py b/tests/unit/refinement/test_ml_sigmaa.py index eb2190ec..acf6dbfc 100644 --- a/tests/unit/refinement/test_ml_sigmaa.py +++ b/tests/unit/refinement/test_ml_sigmaa.py @@ -153,17 +153,15 @@ class TestCollectionTargetsRelocated: def test_exported_from_refinement_targets(self): from torchref.refinement.targets import ( # noqa: F401 + COLLECTION_XRAY_TARGETS, + CollectionDifferenceIntensityTarget, CollectionDifferenceTarget, CollectionMLTarget, - CollectionRiceTarget, MultiModelADPTarget, MultiModelGeometryTarget, ) def test_kinetic_backcompat_reexports(self): - from torchref.experimental.kinetic.targets import ( # noqa: F401 - CollectionRiceTarget, - ) from torchref.experimental.kinetic.targets import CollectionMLTarget as KinCML from torchref.experimental.kinetic.targets import ( # noqa: F401 KineticPriorTarget, @@ -173,13 +171,30 @@ def test_kinetic_backcompat_reexports(self): assert RefCML is KinCML - def test_collection_ml_base_weight(self): + def test_collection_ml_has_its_maintenance_hook(self): + """``LossState`` calls it after each step block to drop the shared beta.""" from torchref.refinement.targets import CollectionMLTarget - assert CollectionMLTarget.DEFAULT_BASE_WEIGHT == 10.0 - # maintenance hook present (resets the target's own shared beta) assert hasattr(CollectionMLTarget, "maintenance") + def test_the_sigma_obs_in_a_rice_sigma_row_is_gone(self): + """``CollectionRiceTarget`` set ``beta = sigma_obs**2``. + + That pairs a measurement sigma with a Rice ``Sigma``, asserting an isotropic + *complex* error where ``sigma_obs`` carries no phase at all. The single-dataset + table refuses to offer such a row + (``test_nll_beta.py::test_rice_with_sigma_obs_is_not_offered``); the collection + table now agrees, and ``CollectionMLTarget`` -- one shared Luzzati beta -- is the + absolute channel instead. + """ + import torchref.refinement.targets as T + from torchref.refinement.targets.collection import COLLECTION_XRAY_TARGETS + + assert not hasattr(T, "CollectionRiceTarget") + assert "rice" not in COLLECTION_XRAY_TARGETS.names + with pytest.raises(ValueError, match="Unknown collection X-ray target"): + COLLECTION_XRAY_TARGETS.by_name("rice") + @pytest.mark.unit class TestBetaMath: diff --git a/torchref/experimental/kinetic/__init__.py b/torchref/experimental/kinetic/__init__.py index 168acf93..0c9bbf59 100644 --- a/torchref/experimental/kinetic/__init__.py +++ b/torchref/experimental/kinetic/__init__.py @@ -30,7 +30,7 @@ from torchref.experimental.kinetic.refinement import KineticRefinement from torchref.experimental.kinetic.targets import ( CollectionDifferenceTarget, - CollectionRiceTarget, + CollectionMLTarget, MultiModelGeometryTarget, MultiModelADPTarget, KineticPriorTarget, @@ -44,7 +44,7 @@ "ModelCollection", "KineticRefinement", "CollectionDifferenceTarget", - "CollectionRiceTarget", + "CollectionMLTarget", "MultiModelGeometryTarget", "MultiModelADPTarget", "KineticPriorTarget", diff --git a/torchref/experimental/kinetic/refinement.py b/torchref/experimental/kinetic/refinement.py index 0e5a0a94..0892986a 100644 --- a/torchref/experimental/kinetic/refinement.py +++ b/torchref/experimental/kinetic/refinement.py @@ -39,7 +39,7 @@ from torchref.refinement.loss_state import LossState, create_loss_state from torchref.experimental.kinetic.targets import ( CollectionDifferenceTarget, - CollectionRiceTarget, + CollectionMLTarget, MultiModelGeometryTarget, MultiModelADPTarget, KineticPriorTarget, @@ -67,8 +67,10 @@ class KineticRefinement(DeviceMixin, nn.Module): Collection of mixed models keyed by timepoint name. xray_weight_difference : float Weight for the difference X-ray target. - xray_weight_rice : float - Weight for the Rice amplitude target. + xray_weight_ml : float + Weight for the absolute (Read MLF) amplitude target. With K free base models a + purely relative loss leaves the overall level unconstrained, so this is the + anchor. geometry_weight : float Weight for geometry restraints. adp_weight : float @@ -86,7 +88,7 @@ def __init__( dataset_collection: "DatasetCollection", model_collection: "ModelCollection", xray_weight_difference: float = 2.0, - xray_weight_rice: float = 1.0, + xray_weight_ml: float = 1.0, geometry_weight: float = 10.0, adp_weight: float = 3.0, kinetic_prior_weight: float = 0.0, @@ -106,7 +108,7 @@ def __init__( # Default weights self._weights = { "xray/difference": xray_weight_difference, - "xray/rice": xray_weight_rice, + "xray/ml": xray_weight_ml, "geometry": geometry_weight, "adp": adp_weight, "kinetic_prior": kinetic_prior_weight, @@ -117,7 +119,7 @@ def __init__( self.loss_state: Optional[LossState] = None self.kinetic_prior_target: Optional[KineticPriorTarget] = None self._diff_target: Optional[CollectionDifferenceTarget] = None - self._rice_target: Optional[CollectionRiceTarget] = None + self._ml_target: Optional[CollectionMLTarget] = None self._kinetic_model = None self._timepoints_map: Optional[Dict[str, int]] = None @@ -164,14 +166,14 @@ def setup( scaler=self.scaler, verbose=self.verbose, ) - rice_target = CollectionRiceTarget( + ml_target = CollectionMLTarget( dc, mc, scaler=self.scaler, verbose=self.verbose, ) # Store direct references for refine_kinetics() self._diff_target = diff_target - self._rice_target = rice_target + self._ml_target = ml_target geom_target = MultiModelGeometryTarget(mc, verbose=self.verbose) adp_target = MultiModelADPTarget(mc, verbose=self.verbose) @@ -179,7 +181,7 @@ def setup( self.loss_state = create_loss_state(device=device) self.loss_state.register_target("xray/difference", diff_target) - self.loss_state.register_target("xray/rice", rice_target) + self.loss_state.register_target("xray/ml", ml_target) self.loss_state.register_target("geometry", geom_target) self.loss_state.register_target("adp", adp_target) @@ -242,7 +244,7 @@ def set_weights(self, **kwargs): # Map short names to full paths mapping = { "difference": "xray/difference", - "rice": "xray/rice", + "ml": "xray/ml", "geometry": "geometry", "adp": "adp", "kinetic_prior": "kinetic_prior", @@ -528,7 +530,7 @@ def refine_kinetics(self, niter: int = 200, lr: float = 1e-2): optimizer = torch.optim.Adam(params, lr=lr) w_diff = self._weights.get("xray/difference", 1.0) - w_rice = self._weights.get("xray/rice", 1.0) + w_ml = self._weights.get("xray/ml", 1.0) # Collect all model keys that have kinetic indices (including dark) all_overrides = {} @@ -546,8 +548,8 @@ def refine_kinetics(self, niter: int = 200, lr: float = 1e-2): for tp_name, t_idx in all_overrides.items(): mc[tp_name].set_fraction_override(kinetic_occ[:, t_idx]) - # Compute X-ray loss only (difference + Rice) - loss = w_diff * self._diff_target() + w_rice * self._rice_target() + # Compute X-ray loss only (difference + the absolute anchor) + loss = w_diff * self._diff_target() + w_ml * self._ml_target() loss.backward() diff --git a/torchref/experimental/kinetic/targets.py b/torchref/experimental/kinetic/targets.py index 54201230..10037c6c 100644 --- a/torchref/experimental/kinetic/targets.py +++ b/torchref/experimental/kinetic/targets.py @@ -7,8 +7,8 @@ CollectionDifferenceTarget Multi-timepoint difference target (primary optimization driver). -CollectionRiceTarget - Multi-timepoint Rice maximum-likelihood amplitude target. +CollectionMLTarget + Read MLF at one shared Luzzati beta; the absolute channel. MultiModelGeometryTarget Geometry restraints applied to the shared base models. MultiModelADPTarget @@ -28,9 +28,9 @@ # Back-compat re-exports of the relocated generic collection targets. from torchref.refinement.targets.collection import ( # noqa: F401 + CollectionDifferenceIntensityTarget, CollectionDifferenceTarget, CollectionMLTarget, - CollectionRiceTarget, MultiModelADPTarget, MultiModelGeometryTarget, ) diff --git a/torchref/refinement/targets/__init__.py b/torchref/refinement/targets/__init__.py index 4205faa0..279bdb6d 100644 --- a/torchref/refinement/targets/__init__.py +++ b/torchref/refinement/targets/__init__.py @@ -20,10 +20,11 @@ von_mises_nll, ) from .collection import ( + COLLECTION_XRAY_TARGETS, + CollectionDifferenceIntensityTarget, CollectionDifferenceTarget, - CollectionTwoMomentIntensityTarget, CollectionMLTarget, - CollectionRiceTarget, + CollectionTwoMomentIntensityTarget, MultiModelADPTarget, MultiModelGeometryTarget, ) @@ -88,8 +89,9 @@ # Collection (multi-dataset) targets "CollectionDifferenceTarget", "CollectionTwoMomentIntensityTarget", - "CollectionRiceTarget", "CollectionMLTarget", + "CollectionDifferenceIntensityTarget", + "COLLECTION_XRAY_TARGETS", "MultiModelGeometryTarget", "MultiModelADPTarget", # Difference targets diff --git a/torchref/refinement/targets/collection/__init__.py b/torchref/refinement/targets/collection/__init__.py index ca6a0b26..92e25460 100644 --- a/torchref/refinement/targets/collection/__init__.py +++ b/torchref/refinement/targets/collection/__init__.py @@ -8,20 +8,34 @@ """ from ._util import _scale_fcalc -from .base import CollectionXrayTarget +from .base import ( + CollectionLossInputs, + CollectionSigmaALossInputs, + CollectionXrayTarget, +) from .intensity import CollectionTwoMomentIntensityTarget from .multimodel import MultiModelADPTarget, MultiModelGeometryTarget from .xray import ( + CollectionDifferenceIntensityTarget, CollectionDifferenceTarget, CollectionMLTarget, - CollectionRiceTarget, +) +from ._specs import ( # noqa: E402 (imports the rows above) + COLLECTION_XRAY_TARGETS, + CollectionXrayTargetSpec, + CollectionXrayTargetTable, ) __all__ = [ + "COLLECTION_XRAY_TARGETS", + "CollectionXrayTargetSpec", + "CollectionXrayTargetTable", "CollectionXrayTarget", + "CollectionLossInputs", + "CollectionSigmaALossInputs", "CollectionTwoMomentIntensityTarget", "CollectionDifferenceTarget", - "CollectionRiceTarget", + "CollectionDifferenceIntensityTarget", "CollectionMLTarget", "MultiModelGeometryTarget", "MultiModelADPTarget", diff --git a/torchref/refinement/targets/collection/_specs.py b/torchref/refinement/targets/collection/_specs.py new file mode 100644 index 00000000..84e7084f --- /dev/null +++ b/torchref/refinement/targets/collection/_specs.py @@ -0,0 +1,166 @@ +"""The collection X-ray target taxonomy, as data. + +The multi-dataset mirror of :mod:`torchref.refinement.targets.xray._specs`, with the same +invariants checked the same way at import: unique names, and **one class per row**, so +dispatch is ``spec.target_cls(**kwargs)`` with nothing to branch on. + +Four rows over two axes -- what the loss compares (a difference from the collection mean, +or each dataset absolutely) and in which observable: + +==================== =========== ============================================== +row observable compares +==================== =========== ============================================== +``difference`` amplitude ``F_i - F_mean`` against the model's own spread +``difference_i`` intensity the same, in intensities +``two_moment`` intensity ``|F(alpha)|^2 + sigma_alpha^2 |dF|^2`` +``ml`` amplitude each dataset absolutely, at a shared Luzzati beta +==================== =========== ============================================== + +The two difference rows are both offered rather than one being chosen. Which is better is a +property of a dataset's signal-to-noise -- amplitudes keep the loss in the same space as the +output DED coefficients, intensities avoid the French-Wilson posterior reshaping the weak +tail the signal lives in -- and that is not something to settle once in a library. + +``ml`` is the absolute channel, for the scenarios that are not difference refinement: with +K free base models a purely relative loss leaves the overall level unconstrained. It +replaces a hand-rolled Rice target that set ``beta = sigma_obs**2``, i.e. exactly the +sigma_obs-in-a-Rice-Sigma pairing that +:mod:`torchref.base.targets.xray_likelihoods` documents as never correct and that the +single-dataset table deliberately does not offer. + +There is no intensity ``ml`` row, for the same reason the single-dataset table has none: +Rice and the folded normal are distributions *of an amplitude*, and the intensity analogue +is the exponential / chi-square_1 Wilson distribution -- a different primitive rather than a +different variance. +""" + +from dataclasses import dataclass, field +from typing import Dict, Tuple + +from .base import CollectionXrayTarget +from .intensity import CollectionTwoMomentIntensityTarget +from .xray import ( + CollectionDifferenceIntensityTarget, + CollectionDifferenceTarget, + CollectionMLTarget, +) + + +@dataclass(frozen=True) +class CollectionXrayTargetSpec: + """One selectable collection x-ray target: a name, and the class implementing it. + + Attributes + ---------- + name + The row name, as used in ``LossState`` keys (``xray/``). + target_cls + The class. **One class per row**, checked by :class:`CollectionXrayTargetTable`. + doc + One line, for ``--help`` and the loss breakdown. + observable + ``"amplitude"`` or ``"intensity"``. Checked against the class, so a spec and its + implementation cannot disagree -- a row advertising intensities while reading + amplitudes would be wrong by ``2|F|``, which is resolution-dependent and so reads + as a scale or B error rather than as a bug. + """ + + name: str + target_cls: type + doc: str + observable: str = "amplitude" + + def __post_init__(self): + if not ( + isinstance(self.target_cls, type) + and issubclass(self.target_cls, CollectionXrayTarget) + ): + raise TypeError( + f"{self.name}: target_cls {self.target_cls!r} is not a " + f"CollectionXrayTarget subclass" + ) + if self.observable not in ("amplitude", "intensity"): + raise ValueError( + f"{self.name}: observable must be 'amplitude' or 'intensity', " + f"got {self.observable!r}" + ) + declared = getattr(self.target_cls, "observable", "amplitude") + if declared != self.observable: + raise ValueError( + f"{self.name}: spec says observable={self.observable!r} but " + f"{self.target_cls.__name__} says {declared!r}" + ) + + +@dataclass(frozen=True) +class CollectionXrayTargetTable: + """The taxonomy, with uniqueness checked at import.""" + + specs: Tuple[CollectionXrayTargetSpec, ...] + _by_name: Dict[str, CollectionXrayTargetSpec] = field( + init=False, repr=False, default=None + ) + + def __post_init__(self): + lookup: Dict[str, CollectionXrayTargetSpec] = {} + for spec in self.specs: + if spec.name in lookup: + raise ValueError(f"duplicate collection x-ray target name {spec.name!r}") + lookup[spec.name] = spec + by_cls: Dict[type, CollectionXrayTargetSpec] = {} + for spec in self.specs: + if spec.target_cls in by_cls: + raise ValueError( + f"{spec.name} and {by_cls[spec.target_cls].name} both map to " + f"{spec.target_cls.__name__}. One class per row is the invariant this " + f"table exists to enforce: a class serving two rows has to branch on " + f"something at runtime." + ) + by_cls[spec.target_cls] = spec + object.__setattr__(self, "_by_name", lookup) + + @property + def names(self) -> Tuple[str, ...]: + """Canonical names, in table order.""" + return tuple(s.name for s in self.specs) + + def by_name(self, name: str) -> CollectionXrayTargetSpec: + spec = self._by_name.get(name) + if spec is None: + raise ValueError( + f"Unknown collection X-ray target: {name!r}. " + f"Available: {', '.join(self.names)}" + ) + return spec + + +COLLECTION_XRAY_TARGETS = CollectionXrayTargetTable( + specs=( + CollectionXrayTargetSpec( + name="difference", + target_cls=CollectionDifferenceTarget, + doc="Gaussian on each dataset's amplitude difference from the collection " + "mean, with the dataset/mean covariance propagated.", + ), + CollectionXrayTargetSpec( + name="difference_i", + target_cls=CollectionDifferenceIntensityTarget, + observable="intensity", + doc="As 'difference' but on intensities, skipping the French-Wilson " + "conversion that reshapes the weak tail.", + ), + CollectionXrayTargetSpec( + name="two_moment", + target_cls=CollectionTwoMomentIntensityTarget, + observable="intensity", + doc="Merged intensities as |F(alpha)|^2 + sigma_alpha^2 |dF|^2, accounting " + "for crystal-to-crystal spread in activation.", + ), + CollectionXrayTargetSpec( + name="ml", + target_cls=CollectionMLTarget, + doc="Read MLF per dataset at one shared Luzzati beta fitted on the pooled " + "free reflections. The absolute channel.", + ), + ) +) diff --git a/torchref/refinement/targets/collection/base.py b/torchref/refinement/targets/collection/base.py index 01ca7a6f..639a5672 100644 --- a/torchref/refinement/targets/collection/base.py +++ b/torchref/refinement/targets/collection/base.py @@ -85,6 +85,29 @@ class CollectionLossInputs(NamedTuple): keys: List[str] +class CollectionSigmaALossInputs(NamedTuple): + """:class:`CollectionLossInputs` plus one shared model-error estimate. + + The collection twin of + :class:`~torchref.refinement.targets.xray.sigma_a.SigmaALossInputs`. ``beta`` and + ``epsilon`` live on the **common HKL**, shape ``(n_hkl,)``, and broadcast over the + dataset axis: they are fitted once on the pooled free reflections of every data-model + pair, so one per-reflection variance serves every member. + + The two shapes never mix, because each class pairs its own ``_loss_inputs`` with its + own ``_per_refl``. + """ + + obs: torch.Tensor + model: torch.Tensor + sigma: torch.Tensor + mask: torch.Tensor + keys: List[str] + centric: torch.Tensor = None + beta: torch.Tensor = None + epsilon: torch.Tensor = None + + class CollectionXrayTarget(Target): """Base class for multi-dataset X-ray targets. @@ -149,8 +172,8 @@ def __init__( def _keys(self) -> List[str]: """Matched dataset keys this target fits: dark + present timepoints. Targets - fitting only part of the collection override it (``CollectionRiceTarget`` - drops the dark reference). + fitting only part of the collection override it (a target fitting only the + excited timepoints drops the dark reference). """ dc = self._dataset_collection mc = self._model_collection From 89a7b29a149e0c0ccd13fe06db050a64d42b344c Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 18:08:47 +0200 Subject: [PATCH 073/250] Fit the inter-dataset scale on the work set, normalised, with a sigma option DatasetCollection.scale is the only fit in the library with no model on either side: it puts one dataset onto another, both measurements of the same quantity. Three things were wrong with it. It fitted on the free set. The closure masked with ReflectionData.masks(), which is validity only -- TensorMasks.__call__ ANDs the validity masks and has no work/free notion at all -- so the free reflections went into the scale parameters via 10 x LBFGS(max_iter=100), upstream of every target. Measured on figure 4: paired per-reflection chi2 on the held-out free set favours the old fit by 0.00116, CI [0.00174, 0.00060], which is the leak's signature (it scored reflections it should not have seen) at 0.06% of chi2. The scale moves 0.020%, so no shipped number is materially affected. It was unnormalised, against LBFGS's absolute tolerances -- the same hazard ScalerBase.refine_lbfgs documents at length and fixes. And its objective was hard-coded. Now `ls` (default, matching DEFAULT_SCALE_TARGET for the model-to-data fit) or `ls_sigma`, weighting by 1/(sigma**2 + sigma_ref**2). Both sides being measured is exactly the condition under which inverse-variance weighting is the correct weight rather than a modelling choice: the denominator is the propagated error on the difference being minimised, with no model-error term in it. There is no sigma_A or Rice option, and that is a modelling statement -- no model, no model error. Measured, held out, paired: ls_sigma beats ls by 0.0250 in free-set chi2, CI [0.0049, 0.0457]. Left non-default pending more than one dataset pair. The ls_sigma weights are computed once and detached. A live denominator would reward inflating the scale to inflate the variance, with no +log(sigma) term to oppose it. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- docs/changelog.rst | 3 + paper/probe_data_scale_objective.py | 230 +++++++++++++++++++++++++++ tests/unit/io/test_data_scale_fit.py | 149 +++++++++++++++++ torchref/io/datasets/collection.py | 110 +++++++++++-- 4 files changed, 475 insertions(+), 17 deletions(-) create mode 100644 paper/probe_data_scale_objective.py create mode 100644 tests/unit/io/test_data_scale_fit.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 77d0802b..9f683342 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,9 @@ Changelog Version 0.6.4 ---------- +- Fixed ``DatasetCollection.scale`` fitting the inter-dataset scale on the free reflections as well as the work set +- ``DatasetCollection.scale`` normalises its objective, so L-BFGS's absolute tolerances mean something +- Added ``DatasetCollection.scale(objective="ls_sigma")``, weighting by the propagated error on the difference being minimised - Added ``COLLECTION_XRAY_TARGETS``, the collection target taxonomy, with an intensity difference row - Removed ``CollectionRiceTarget``, which set ``beta = sigma_obs**2``; the ``ml`` row is the absolute channel instead - Renamed the kinetic ``xray_weight_rice`` / ``xray/rice`` weight to ``xray_weight_ml`` / ``xray/ml`` diff --git a/paper/probe_data_scale_objective.py b/paper/probe_data_scale_objective.py new file mode 100644 index 00000000..12e34b6d --- /dev/null +++ b/paper/probe_data_scale_objective.py @@ -0,0 +1,230 @@ +#!/usr/bin/env python +"""What does the data-to-data scale fit cost, and which objective generalises? + +``DatasetCollection.scale()`` puts one dataset onto another. There is no model on either +side, so there is no model error for a sigma_A or Rice likelihood to account for, and the +only real choices are the weighting and which reflections the fit is allowed to see. + +Two questions, both answerable by holding out the free set: + +1. **The leak.** The fit used to mask with ``ReflectionData.masks()`` -- validity only, + with no work/free notion -- so the free reflections went into the scale parameters, + upstream of every target. How much did that actually buy it, and does removing it + change the fitted scale? +2. **The weighting.** ``ls`` throws sigma away. ``ls_sigma`` weights by + ``1/(sigma**2 + sigma_ref**2)``, which is the propagated error on the very difference + being minimised -- the correct weight rather than a modelling choice, precisely because + both sides are measurements. Does it generalise better? + +Both are scored on reflections the fit never saw, under two yardsticks applied identically +to every arm, so the comparison is not circular: + + R_data = sum|F - F_ref| / sum F_ref scale-free and interpretable + chi2 = mean[(F - F_ref)**2 / (s**2 + s_ref**2)] is the disagreement within error? + +``chi2`` is the one that can separate them: ``ls`` is entitled to win on ``R_data``, which +is what it optimises up to a constant. +""" + +import argparse +import json +import math +import sys +from pathlib import Path + +import torch + + +def build(dark_sf, light_sf, d_min, device): + """The figure-4 pair, loaded exactly as the difference CLI loads it, unscaled.""" + from torchref.cli.collection_difference_refine import setup_dataset_collection + + return setup_dataset_collection(dark_sf, light_sf, d_min, device) + + +def _reset(dc): + """Zero every fitted scale parameter and drop the corrected caches.""" + for _, ds in dc: + with torch.no_grad(): + if getattr(ds, "log_scale", None) is not None: + ds.log_scale.zero_() + if getattr(ds, "U_aniso", None) is not None: + ds.U_aniso.zero_() + ds._corrected_fp = None + ds._corrected_cache = None + ds._corrected_I_fp = None + ds._corrected_I_cache = None + + +def legacy_scale(dc): + """The pre-fix fit, copied verbatim for comparison: validity masks, unnormalised. + + A deliberate duplicate rather than a flag on the library method -- the point is to + measure what the old behaviour did, not to keep it selectable. + """ + ref_ds = dc._datasets[dc._reference_dataset] + to_scale = [ds for name, ds in dc if name != dc._reference_dataset] + params = [p for data in to_scale for p in data.parameters()] + [p.requires_grad_(True) for p in params] + opt = torch.optim.LBFGS(params, max_iter=100, line_search_fn="strong_wolfe") + ref_mask = ref_ds.masks() + ds_masks = [ds.masks() for ds in to_scale] + + def closure(): + opt.zero_grad() + loss = 0.0 + ref_F, _ = ref_ds.get_corrected_data() + for ds, m in zip(to_scale, ds_masks): + F, _ = ds.get_corrected_data() + cm = m & ref_mask + loss = loss + torch.sum((F[cm] - ref_F[cm]) ** 2) + loss.backward() + return loss + + for _ in range(10): + opt.step(closure) + [p.requires_grad_(False) for p in params] + + +def per_refl_chi2(dc, subset): + """Per-reflection chi-square contributions on ``subset``, in a fixed order. + + Returned rather than reduced so arms can be compared **paired**. The mask depends + only on the flags and the validity masks, never on the fitted scale, so the same + reflection sits at the same index in every arm -- which is what makes the pairing + valid. Unpaired means cannot resolve this comparison: the two datasets are different + structures, so most of chi2 is real difference and it cancels only when paired. + """ + from torchref.base.targets.xray_likelihoods import floor_sigma_obs + + ref_name = dc._reference_dataset + other = [n for n, _ in dc if n != ref_name][0] + ref_ds, ds = dc[ref_name], dc[other] + with torch.no_grad(): + ref_F, ref_s = ref_ds.get_corrected_data() + F, s = ds.get_corrected_data() + m = getattr(ds, subset).mask & getattr(ref_ds, subset).mask + so, sr = floor_sigma_obs(s[m]), floor_sigma_obs(ref_s[m]) + return (((F[m] - ref_F[m]) ** 2) / (so**2 + sr**2)).cpu() + + +def paired_ci(a, b, n_boot=4000, seed=0): + """Bootstrap CI on ``mean(a - b)`` over reflections. Positive favours ``b``.""" + d = (a - b).numpy() + g = torch.Generator().manual_seed(seed) + n = len(d) + idx = torch.randint(0, n, (n_boot, n), generator=g).numpy() + means = d[idx].mean(axis=1) + lo, hi = sorted(means)[int(0.025 * n_boot)], sorted(means)[int(0.975 * n_boot)] + return float(d.mean()), float(lo), float(hi) + + +def score(dc, subset): + """``(R_data, chi2, n)`` between the two datasets on one held-out subset.""" + from torchref.base.targets.xray_likelihoods import floor_sigma_obs + + names = [n for n, _ in dc] + ref_name = dc._reference_dataset + other = [n for n in names if n != ref_name][0] + ref_ds, ds = dc[ref_name], dc[other] + + with torch.no_grad(): + ref_F, ref_s = ref_ds.get_corrected_data() + F, s = ds.get_corrected_data() + m = getattr(ds, subset).mask & getattr(ref_ds, subset).mask + fo, fr = F[m], ref_F[m] + so, sr = floor_sigma_obs(s[m]), floor_sigma_obs(ref_s[m]) + r = float((fo - fr).abs().sum() / fr.abs().sum().clamp(min=1e-30)) + chi2 = float((((fo - fr) ** 2) / (so**2 + sr**2)).mean()) + return r, chi2, int(m.sum()) + + +def fitted(dc): + other = [n for n, _ in dc if n != dc._reference_dataset][0] + ds = dc[other] + ls = float(ds.log_scale.detach().reshape(-1)[0]) + u = ds.U_aniso.detach().reshape(-1).tolist() + return ls, u + + +def main(): + ap = argparse.ArgumentParser(description=__doc__) + fig4 = Path(__file__).resolve().parent / "figure4_difference_refinement" + ap.add_argument("--dark-sf", default=str(fig4 / "data/8QL2-sf.cif")) + ap.add_argument("--light-sf", default=str(fig4 / "data/7YYZ-light.mtz")) + ap.add_argument("--dmin", type=float, default=2.2) + ap.add_argument("--device", default="cpu") + ap.add_argument("-o", "--out", default=None) + args = ap.parse_args() + + dev = torch.device(args.device) + dc = build(args.dark_sf, args.light_sf, args.dmin, dev) + + arms = [ + ("legacy (validity mask, unnormalised)", lambda: legacy_scale(dc)), + ("ls (work set, normalised)", lambda: dc.scale(objective="ls")), + ("ls_sigma (work set, normalised)", lambda: dc.scale(objective="ls_sigma")), + ] + + rows = [] + per_refl = {} + for label, fit in arms: + _reset(dc) + fit() + ls, u = fitted(dc) + rw, cw, nw = score(dc, "work") + rf, cf, nf = score(dc, "free") + per_refl[label] = { + "free": per_refl_chi2(dc, "free"), + "work": per_refl_chi2(dc, "work"), + } + rows.append( + dict(arm=label, log_scale=ls, U_aniso=u, + R_work=rw, chi2_work=cw, n_work=nw, + R_free=rf, chi2_free=cf, n_free=nf) + ) + + print() + print("Data-to-data scale fit: fitted on WORK, scored on the held-out FREE set") + print("=" * 86) + print(f"{'arm':38s} {'log_scale':>10s} {'R_work':>8s} {'R_free':>8s} " + f"{'chi2_work':>10s} {'chi2_free':>10s}") + print("-" * 86) + for r in rows: + print(f"{r['arm']:38s} {r['log_scale']:10.5f} {r['R_work']:8.5f} " + f"{r['R_free']:8.5f} {r['chi2_work']:10.3f} {r['chi2_free']:10.3f}") + print("-" * 86) + print(f"n_work={rows[0]['n_work']} n_free={rows[0]['n_free']}") + print() + base = rows[0] + for r in rows[1:]: + d = r["log_scale"] - base["log_scale"] + print(f"{r['arm']:38s} d(log_scale) vs legacy = {d:+.6f} " + f"({100*(math.exp(d)-1):+.3f}% in scale)") + + # --- the comparison that can actually resolve this: paired, per reflection ------ + labels = [lbl for lbl, _ in arms] + print() + print("Paired per-reflection chi2 difference (positive => the SECOND arm is better)") + print("=" * 86) + print(f"{'comparison':52s} {'set':6s} {'mean d':>10s} {'95% CI':>22s}") + print("-" * 86) + pairs = [(labels[0], labels[1]), (labels[1], labels[2]), (labels[0], labels[2])] + paired = [] + for a, b in pairs: + for subset in ("work", "free"): + m, lo, hi = paired_ci(per_refl[a][subset], per_refl[b][subset]) + sig = "" if lo <= 0.0 <= hi else " <-- CI excludes 0" + name = f"{a.split('(')[0].strip()} vs {b.split('(')[0].strip()}" + print(f"{name:52s} {subset:6s} {m:+10.5f} [{lo:+.5f}, {hi:+.5f}]{sig}") + paired.append(dict(a=a, b=b, subset=subset, mean=m, lo=lo, hi=hi)) + print("-" * 86) + + if args.out: + Path(args.out).write_text(json.dumps({"arms": rows, "paired": paired}, indent=2)) + print(f"written: {args.out}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/unit/io/test_data_scale_fit.py b/tests/unit/io/test_data_scale_fit.py new file mode 100644 index 00000000..934c061e --- /dev/null +++ b/tests/unit/io/test_data_scale_fit.py @@ -0,0 +1,149 @@ +"""``DatasetCollection.scale()`` -- the data-to-data scale fit. + +The only fit in the library with no model on either side: it puts one dataset onto +another, both of them measurements of the same quantity. That is why its objectives are +least squares and there is no sigma_A row -- there is no model error to account for. + +The load-bearing test here is the free-set one. That fit runs *upstream of every target*, +so a leak there compromises every free-set number the pipeline later reports, and no +downstream test would notice. +""" + +import pytest +import torch + + +@pytest.fixture +def pair(mtz_dir): + """Two copies of 1DAW as a reference + one dataset to scale onto it.""" + mtz = mtz_dir / "1DAW.mtz" + if not mtz.exists(): + pytest.skip("1DAW fixture not present") + + from torchref import ReflectionData + from torchref.io.datasets.collection import DatasetCollection + + ref = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + other = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + dc = DatasetCollection(verbose=0, device="cpu") + dc.add_dataset("ref", ref, set_as_reference=True) + dc.add_dataset("other", other) + return dc + + +def _fitted(dc): + ds = dc["other"] + return ds.log_scale.detach().clone(), ds.U_aniso.detach().clone() + + +@pytest.mark.integration +@pytest.mark.parametrize("objective", ["ls", "ls_sigma"]) +def test_scale_never_touches_the_free_set(pair, objective): + """Corrupting the free reflections must not move the fitted parameters at all. + + Two *different* garbage values, because a single one could coincide with a + no-op: if the fit sees the free set, two different corruptions give two different + answers. ``torch.equal``, not ``allclose`` -- the free reflections must contribute + exactly nothing, not merely little. + + This fit used to mask with ``ReflectionData.masks()``, which is validity only + (``TensorMasks.__call__`` ANDs the validity masks and has no work/free notion), so + the free reflections went into the scale parameters via 10 x LBFGS(max_iter=100). + """ + dc = pair + free = dc["other"].free.mask + assert free.sum() > 0, "fixture has no free reflections; the test would be vacuous" + + results = [] + for filler in (3.0, 900.0): + d = dc["other"] + # Reset the parameters so each arm starts from the same place. + with torch.no_grad(): + d.log_scale.zero_() + d.U_aniso.zero_() + d.F[free] = filler + d.F_sigma[free] = filler + d._corrected_fp = None # drop the cached corrected view + d._corrected_cache = None + dc.scale(objective=objective) + results.append(_fitted(dc)) + + (ls_a, u_a), (ls_b, u_b) = results + assert torch.equal(ls_a, ls_b), ( + f"log_scale moved when only the FREE reflections changed: {ls_a} vs {ls_b}" + ) + assert torch.equal(u_a, u_b), ( + f"U_aniso moved when only the FREE reflections changed: {u_a} vs {u_b}" + ) + + +@pytest.mark.integration +@pytest.mark.parametrize("objective", ["ls", "ls_sigma"]) +def test_identical_datasets_fit_a_unit_scale(pair, objective): + """Two copies of one dataset must scale onto each other with no correction. + + The sanity check the objectives have to pass before any comparison between them + means anything. + """ + dc = pair + dc.scale(objective=objective) + log_scale, U = _fitted(dc) + assert float(log_scale.abs().max()) < 1e-3, log_scale + assert float(U.abs().max()) < 1e-3, U + + +@pytest.mark.integration +def test_a_known_scale_is_recovered(pair): + """Scale one dataset by a known factor and check the fit undoes it.""" + dc = pair + k = 2.5 + with torch.no_grad(): + d = dc["other"] + d.F *= k + d.F_sigma *= k + d._corrected_fp = None + d._corrected_cache = None + dc.scale() + log_scale, _ = _fitted(dc) + # log_scale multiplies the observations, so recovering 1/k means log_scale = -log(k). + import math + assert float(log_scale.reshape(-1)[0]) == pytest.approx(-math.log(k), abs=0.02) + + +@pytest.mark.unit +def test_unknown_objective_fails_closed(): + from torchref.io.datasets.collection import ( + DATA_SCALE_OBJECTIVES, + DatasetCollection, + ) + + dc = DatasetCollection(verbose=0, device="cpu") + with pytest.raises(ValueError, match="objective must be one of"): + dc.scale(objective="ml") + # No sigma_A / Rice row is offered, and that is a modelling statement: this fit has + # no model, so there is no model error for such a likelihood to account for. + assert DATA_SCALE_OBJECTIVES == ("ls", "ls_sigma") + + +@pytest.mark.integration +def test_the_objective_is_normalised(pair): + """The loss handed to L-BFGS must be O(1), because its tolerances are absolute. + + Not a style point: ``tolerance_grad``/``tolerance_change`` are absolute, so an + objective carrying the data's own magnitude (~1e9 on a large work set under unit + weights) puts the float32 ulp of the loss above the decrease the line search is + trying to resolve. Probed by scaling the data by 1e3 and checking the fit still + recovers the same answer. + """ + import math + + dc = pair + with torch.no_grad(): + d = dc["other"] + d.F *= 1000.0 + d.F_sigma *= 1000.0 + d._corrected_fp = None + d._corrected_cache = None + dc.scale() + log_scale, _ = _fitted(dc) + assert float(log_scale.reshape(-1)[0]) == pytest.approx(-math.log(1000.0), abs=0.05) diff --git a/torchref/io/datasets/collection.py b/torchref/io/datasets/collection.py index 38fd2c62..6fef9be5 100644 --- a/torchref/io/datasets/collection.py +++ b/torchref/io/datasets/collection.py @@ -14,6 +14,17 @@ from .base import CrystalDataset from .reflection_data import ReflectionData +#: Objectives for :meth:`DatasetCollection.scale`, the **data-to-data** fit. Least +#: squares only, and not for want of alternatives: there is no model in that fit, so +#: there is no model error for a sigma_A or Rice likelihood to account for. ``ls_sigma`` +#: weights by the propagated error on the difference, which is the correct weight +#: precisely because both sides are measurements. +DATA_SCALE_OBJECTIVES = ("ls", "ls_sigma") + +#: Default for :meth:`DatasetCollection.scale`. Unit-weight least squares, matching +#: :data:`~torchref.scaling.scaler_base.DEFAULT_SCALE_TARGET` for the model-to-data fit. +DEFAULT_DATA_SCALE_OBJECTIVE = "ls" + @dataclass class DatasetCollection(CrystalDataset): @@ -262,37 +273,101 @@ def __call__(self, mask: bool = True) -> Dict[str, Tuple]: """ return {name: ds(mask=mask, scale=True) for name, ds in self} - def scale(self): + def scale(self, objective: str = DEFAULT_DATA_SCALE_OBJECTIVE): """ - Least-squares fit every non-reference dataset's scale and anisotropy - onto the reference, whose own parameters are left untouched. + Fit every non-reference dataset's scale and anisotropy onto the reference, + whose own parameters are left untouched. + + **This is the data-to-data fit**, and it is the only one in the library: there + is no model here, so there is no model error to account for and nothing for a + sigma_A or Rice likelihood to do. Both sides are measurements of the same + quantity, which is why the objectives are least squares -- + :data:`DATA_SCALE_OBJECTIVES`: + + ``ls`` + ``sum (F - F_ref)**2``, unit weights. The default, matching + :data:`~torchref.scaling.scaler_base.DEFAULT_SCALE_TARGET` for the + model-to-data fit. + ``ls_sigma`` + ``sum (F - F_ref)**2 / (sigma**2 + sigma_ref**2)``. Both sides being + measured is exactly the condition under which inverse-variance weighting is + the correct weight rather than a modelling choice: the denominator is the + propagated error on the difference being minimised, with no model-error term + in it. + + The ``ls_sigma`` weights are computed **once, detached**, from the starting + sigmas. They must not be re-derived inside the closure: ``sigma`` carries the + same ``log_scale`` as ``F``, so a live denominator rewards inflating the scale to + inflate the variance, and without the ``+log(sigma)`` term of a full Gaussian + there is nothing to oppose it. Fixed weights are what "weighted least squares" + means; see :mod:`torchref.scaling.scaler_base` on the related hazard of fitting a + scale against a likelihood that carries the scale in its variance. + + Fitted on the **work set** of both datasets. L-BFGS with strong-Wolfe line + search, 10 outer steps of ``max_iter=100``, on an objective normalised to O(1) + because those tolerances are absolute. Members' ``log_scale``/``U_aniso`` are + mutated, and ``requires_grad`` is turned on and back off around the fit. - L-BFGS with strong-Wolfe line search, 10 outer steps of ``max_iter=100``. - Members' ``log_scale``/``U_aniso`` are mutated, and ``requires_grad`` is - turned on and back off around the fit. + Parameters + ---------- + objective : str, optional + One of :data:`DATA_SCALE_OBJECTIVES`. Raises ------ ValueError - If no reference dataset is set, or there is nothing else to scale. + If no reference dataset is set, there is nothing else to scale, or + ``objective`` is not recognised. """ + if objective not in DATA_SCALE_OBJECTIVES: + raise ValueError( + f"objective must be one of {DATA_SCALE_OBJECTIVES}, got {objective!r}" + ) if self._reference_dataset is None: raise ValueError("No reference dataset set for scaling") - ref_ds = self._datasets[self._reference_dataset] to_scale = [ds for name, ds in self if name != self._reference_dataset] if not to_scale: raise ValueError("No datasets to scale against reference") - + parameters = [p for data in to_scale for p in data.parameters()] [p.requires_grad_(True) for p in parameters] optimizer = torch.optim.LBFGS(parameters, max_iter=100, line_search_fn='strong_wolfe') - # Get masks once (they don't change during optimization) - ref_mask = ref_ds.masks() - ds_masks = [ds.masks() for ds in to_scale] + # Masks once (they do not change during the fit). The WORK subset, not + # `masks()`: the latter is validity only -- `TensorMasks.__call__` ANDs the + # validity masks and carries no work/free notion at all -- so fitting against it + # puts the free reflections into the scale parameters, upstream of every target, + # and compromises any free-set number the pipeline later reports. Degrades to + # all-valid on a dataset with no R-free flags, which is the pre-existing + # behaviour for that case. + ref_mask = ref_ds.work.mask + combined = [ds.work.mask & ref_mask for ds in to_scale] + + # Weights and the normaliser: once, detached, outside the closure. + with torch.no_grad(): + ref_F0, ref_sig0 = ref_ds.get_corrected_data() + weights = None + if objective == "ls_sigma": + # Local import: `torchref.base.targets` is not otherwise reachable from + # `torchref.io`, and hoisting it would couple the two packages. + from torchref.base.targets.xray_likelihoods import floor_sigma_obs + + weights = [] + for ds, cm in zip(to_scale, combined): + _, sig0 = ds.get_corrected_data() + var = ( + floor_sigma_obs(sig0[cm]) ** 2 + + floor_sigma_obs(ref_sig0[cm]) ** 2 + ) + weights.append(1.0 / var) + n_fitted = sum(int(cm.sum()) for cm in combined) + norm = 1.0 / max(n_fitted, 1) + else: + ssq = sum(float(ref_F0[cm].pow(2).sum()) for cm in combined) + norm = 1.0 / max(ssq, 1e-30) def closure(): optimizer.zero_grad() @@ -300,12 +375,13 @@ def closure(): # get_corrected_data, not __call__: MaskedTensor has no autograd. ref_F_scaled, _ = ref_ds.get_corrected_data() - for ds, ds_mask in zip(to_scale, ds_masks): + for i, (ds, cm) in enumerate(zip(to_scale, combined)): F_scaled, _ = ds.get_corrected_data() - combined_mask = ds_mask & ref_mask - F_data = F_scaled[combined_mask] - ref_F_data = ref_F_scaled[combined_mask] - loss = loss + torch.sum((F_data - ref_F_data) ** 2) + resid_sq = (F_scaled[cm] - ref_F_scaled[cm]) ** 2 + if weights is not None: + resid_sq = resid_sq * weights[i] + loss = loss + torch.sum(resid_sq) + loss = loss * norm loss.backward() return loss From 9a85db7678009e3374a30e310596cc23397963b6 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 27 Aug 2026 18:20:52 +0200 Subject: [PATCH 074/250] Build the joint scale fit from the target table, and normalise it refine_lbfgs_joint hand-rolled a Rice likelihood inline at beta = sigma_obs**2 -- the sigma_obs-in-a-Rice-Sigma pairing xray_likelihoods documents as never correct, and which the taxonomy deliberately does not offer. It now builds a row via create_xray_target exactly as ScalerBase.refine_lbfgs does, so both scale fits evaluate the same likelihood code and neither keeps a private copy. That needed a way to hand a plain scaler to a row: the collection's bulk solvent is a fraction-weighted mixture, so which dataset is being scaled matters, and ScalerBase.forward cannot carry that. _DatasetScalerView binds the fractions and forwards to forward_mixed, holding the parent through ModuleReference so its parameters are not re-registered. The whole beta/epsilon precomputation goes with it: rows that need a model-error estimate build and own one, as under the body refinement. The objective is now normalised. It had none at all, against LBFGS's absolute tolerances. The U penalty keeps an amplitude normaliser rather than following the objective, so selecting a different objective cannot silently change the regularisation strength. Default is ls, matching DEFAULT_SCALE_TARGET. Measured on figure 4, 5 repeats, paired, scored on reflections the fit did not see (spread <= 0.0002, so these resolve): ml_noalpha reproduces the old numbers to within noise, which is the check that the refactor preserved the physics. ls is a trade -- light R-free -0.00216, dark R-free +0.00086 -- and nll is worse on both. Amplitudes throughout, whatever the refinement target fits: unit-weight least squares on intensities carries an extra factor of F**2 in the squared residual, so a global scale plus B plus anisotropy would be set by the strongest low-resolution reflections and leave high resolution unconstrained. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- docs/changelog.rst | 2 + paper/probe_joint_scale_objective.py | 197 +++++++++++++++ .../test_collection_joint_scale_fit.py | 146 +++++++++++ torchref/scaling/collection_scaler.py | 239 ++++++++++++------ 4 files changed, 501 insertions(+), 83 deletions(-) create mode 100644 paper/probe_joint_scale_objective.py create mode 100644 tests/unit/scaling/test_collection_joint_scale_fit.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 9f683342..cbd8e7d3 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,8 @@ Changelog Version 0.6.4 ---------- +- ``CollectionScaler.refine_lbfgs_joint`` builds a row of ``XRAY_TARGETS`` instead of its own Rice likelihood, and takes ``scale_target`` (default ``ls``) +- ``CollectionScaler.refine_lbfgs_joint`` normalises its objective and registers the U penalty as its own target - Fixed ``DatasetCollection.scale`` fitting the inter-dataset scale on the free reflections as well as the work set - ``DatasetCollection.scale`` normalises its objective, so L-BFGS's absolute tolerances mean something - Added ``DatasetCollection.scale(objective="ls_sigma")``, weighting by the propagated error on the difference being minimised diff --git a/paper/probe_joint_scale_objective.py b/paper/probe_joint_scale_objective.py new file mode 100644 index 00000000..684387cc --- /dev/null +++ b/paper/probe_joint_scale_objective.py @@ -0,0 +1,197 @@ +#!/usr/bin/env python +"""Which objective should the joint (model-to-data) scale fit use? + +``CollectionScaler.refine_lbfgs_joint`` used to hand-roll a Rice likelihood at +``beta = sigma_obs**2`` on an objective with no normaliser at all. It now builds a row of +``XRAY_TARGETS``, defaulting to unit-weight ``ls`` -- the same default the single-dataset +scale fit was moved to after measurement. + +This measures the swap, on the figure-4 pair, against the two things that could make a +single-run comparison meaningless: + +* **Thread nondeterminism.** TorchRef's CPU F_calc is not bit-reproducible, and R-free + jitter of order 0.017 has been measured at zero true difference. So every arm is + repeated and the spread is reported alongside the mean; a difference smaller than the + spread is not a result. +* **Circularity.** Each objective is scored on ``rfree``, which none of them optimises + (they all fit the work set), and by the same ``rfactor_work_free`` for every arm. + +The legacy fit is reconstructed here rather than kept selectable: the point is to measure +what it did, not to preserve it. +""" + +import argparse +import json +import statistics +import sys +from pathlib import Path + +import torch + + +def build(dark_sf, light_sf, dark_pdb, light_pdb, cif, d_min, fraction, device): + from torchref.cli.collection_difference_refine import ( + setup_dataset_collection, + setup_model_collection, + ) + + mc = setup_model_collection( + dark_pdb, light_pdb, [1.0 - fraction, fraction], cif, d_min, device, 0 + ) + dc = setup_dataset_collection(dark_sf, light_sf, d_min, device) + return dc, mc + + +def legacy_joint(scaler, dc, mc, nsteps=3, max_iter=200): + """The pre-change fit: hand-rolled Rice at beta=sigma**2, no normaliser.""" + import torch.nn as nn + + from torchref.base.reciprocal import get_scattering_vectors + from torchref.base.targets.xray_likelihoods import complex_var_from_beta, rice_math + from torchref.refinement.loss_state import LossState + from torchref.refinement.model_error_estimation.sigma_a import ( + SigmaAEstimator, + epsilon_from_hkl, + ) + from torchref.scaling.collection_scaler import CollectionScaler + + keys = [k for k in ([mc.dark_key] + mc.timepoint_names) if k in dc] + cache = {} + for name in keys: + data, model = dc[name], mc[name] + with torch.no_grad(): + fc = model(data.hkl).detach() + fracs = model.fractions.detach() + scaled0 = CollectionScaler.forward_mixed(scaler, fc, fracs) + amp0 = torch.abs(scaled0).reshape(-1) + fobs, sig = data.get_corrected_data() + eps0 = epsilon_from_hkl(data.hkl, getattr(data, "spacegroup", None)).to(amp0.dtype) + s = get_scattering_vectors(data.hkl, data.cell) + dss0 = (torch.norm(s, dim=1) ** 2).to(amp0.dtype) + est = SigmaAEstimator().get( + fobs.to(amp0.dtype).reshape(-1), amp0, data.centric, eps0, dss0, + data.free.mask, sigma_obs=sig.to(amp0.dtype).reshape(-1), + ) + cache[name] = (fc, fracs, est.beta, est.epsilon, data.work, data.centric) + + class _T(nn.Module): + name = "scaler/joint" + + def forward(self): + total, n = torch.tensor(0.0, device=scaler.device), 0 + for nm in keys: + fc, fracs, beta, eps, work, cen = cache[nm] + scaled = CollectionScaler.forward_mixed(scaler, fc, fracs) + amp = torch.abs(scaled).reshape(-1) + fo = work.F.to(amp.dtype) + bw = work.select(beta).to(fo.dtype) + ew = work.select(eps).to(fo.dtype) if eps is not None else None + loss = rice_math( + fo, work.select(amp), complex_var_from_beta(bw, ew), work.select(cen) + ) + if torch.isfinite(loss): + total, n = total + loss, n + 1 + if n: + total = total / n + return total + torch.sum(scaler.U**2) + + state = LossState(device=scaler.device) + state.register_target("scaler/joint", _T()) + opt = torch.optim.LBFGS( + scaler.parameters(), lr=1.0, max_iter=max_iter, history_size=10, + line_search_fn="strong_wolfe", + ) + state.run(opt, nsteps=nsteps, log=False, context="probe.legacy_joint") + + +def rfactors(scaler, dc, mc): + """``{key: (rwork, rfree)}`` for every dataset under the current scale.""" + from torchref.base.metrics.rfactor import rfactor_work_free + from torchref.scaling.collection_scaler import CollectionScaler + + out = {} + with torch.no_grad(): + for name in ([mc.dark_key] + mc.timepoint_names): + if name not in dc: + continue + data, model = dc[name], mc[name] + fc = model(data.hkl).detach() + scaled = CollectionScaler.forward_mixed(scaler, fc, model.fractions.detach()) + out[name] = rfactor_work_free(data, torch.abs(scaled)) + return out + + +def main(): + ap = argparse.ArgumentParser(description=__doc__) + fig4 = Path(__file__).resolve().parent / "figure4_difference_refinement" + ap.add_argument("--dark-sf", default=str(fig4 / "data/8QL2-sf.cif")) + ap.add_argument("--light-sf", default=str(fig4 / "data/7YYZ-light.mtz")) + ap.add_argument("--dark-pdb", default=str(fig4 / "data/8QL2_no_altloc.pdb")) + ap.add_argument("--light-pdb", default=str(fig4 / "work_no_altloc.pdb")) + ap.add_argument("--cif", nargs="*", default=[str(fig4 / "data/IBL_grade.cif")]) + ap.add_argument("--dmin", type=float, default=2.2) + ap.add_argument("--fraction", type=float, default=0.22) + ap.add_argument("--repeats", type=int, default=5) + ap.add_argument("--device", default="cpu") + ap.add_argument("-o", "--out", default=None) + args = ap.parse_args() + + from torchref.scaling.collection_scaler import CollectionScaler + + dev = torch.device(args.device) + dc, mc = build(args.dark_sf, args.light_sf, args.dark_pdb, args.light_pdb, + args.cif, args.dmin, args.fraction, dev) + + arms = ["legacy_rice_unnormalised", "ls", "nll", "ml_noalpha"] + results = {a: {"rfree_dark": [], "rfree_light": [], "rwork_dark": []} for a in arms} + + for rep in range(args.repeats): + for arm in arms: + # A fresh scaler each time: `initialize()` reseeds every parameter, so no arm + # inherits another's answer. + scaler = CollectionScaler(dc, mc, verbose=0).initialize() + if arm == "legacy_rice_unnormalised": + legacy_joint(scaler, dc, mc) + else: + scaler.refine_lbfgs_joint(verbose=False, scale_target=arm) + rf = rfactors(scaler, dc, mc) + dark, light = mc.dark_key, mc.timepoint_names[0] + results[arm]["rwork_dark"].append(rf[dark][0]) + results[arm]["rfree_dark"].append(rf[dark][1]) + results[arm]["rfree_light"].append(rf[light][1]) + print(f" repeat {rep + 1}/{args.repeats} done", flush=True) + + def ms(v): + m = statistics.mean(v) + s = statistics.stdev(v) if len(v) > 1 else 0.0 + return m, s + + print() + print(f"Joint scale fit, {args.repeats} repeats -- scored on reflections it did not fit") + print("=" * 82) + print(f"{'objective':28s} {'rwork_dark':>18s} {'rfree_dark':>18s} {'rfree_light':>18s}") + print("-" * 82) + for a in arms: + cells = [] + for k in ("rwork_dark", "rfree_dark", "rfree_light"): + m, s = ms(results[a][k]) + cells.append(f"{m:.5f}+-{s:.5f}") + print(f"{a:28s} {cells[0]:>18s} {cells[1]:>18s} {cells[2]:>18s}") + print("-" * 82) + base = results["legacy_rice_unnormalised"] + for a in arms[1:]: + for k in ("rfree_dark", "rfree_light"): + # Paired by repeat index: the same thread-nondeterminism realisation. + d = [x - y for x, y in zip(results[a][k], base[k])] + m, s = ms(d) + flag = " <-- exceeds its own spread" if abs(m) > 2 * (s or 1e-9) else "" + print(f"{a:28s} d({k}) vs legacy = {m:+.5f} +- {s:.5f}{flag}") + + if args.out: + Path(args.out).write_text(json.dumps(results, indent=2)) + print(f"written: {args.out}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/unit/scaling/test_collection_joint_scale_fit.py b/tests/unit/scaling/test_collection_joint_scale_fit.py new file mode 100644 index 00000000..93dd7bb2 --- /dev/null +++ b/tests/unit/scaling/test_collection_joint_scale_fit.py @@ -0,0 +1,146 @@ +"""``CollectionScaler.refine_lbfgs_joint`` -- the model-to-data fit, per row. + +It used to hand-roll a Rice likelihood inline with ``beta = sigma_obs**2`` and no +normaliser at all. Both are now gone: it builds a row of ``XRAY_TARGETS``, exactly as +``ScalerBase.refine_lbfgs`` does, and normalises the objective because L-BFGS converges on +absolute tolerances. + +The tests here are mostly about what must NOT differ between the two scale fits, since +"one likelihood, two call sites" is the whole claim. +""" + +import inspect + +import pytest +import torch + + +@pytest.fixture(scope="module") +def collection(pdb_dir, mtz_dir): + pdb, mtz = pdb_dir / "1DAW.pdb", mtz_dir / "1DAW.mtz" + if not (pdb.exists() and mtz.exists()): + pytest.skip("1DAW fixture not present") + + from torchref import LBFGSRefinement, ReflectionData + from torchref.io.datasets.collection import DatasetCollection + from torchref.model.model_collection import ModelCollection + from torchref.scaling.collection_scaler import CollectionScaler + + ref = LBFGSRefinement(data_file=str(mtz), pdb=str(pdb), verbose=0) + extra = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + + dc = DatasetCollection(verbose=0, device="cpu") + dc.add_dataset("dark", ref.reflection_data, set_as_reference=True) + dc.add_dataset("t1", extra) + + mc = ModelCollection([ref.model], dark_key="dark", verbose=0) + mc.add_dark() + mc.add_timepoint("t1", [1.0]) + return dc, mc + + +def _fresh_scaler(collection): + from torchref.scaling.collection_scaler import CollectionScaler + + dc, mc = collection + return CollectionScaler(dc, mc, verbose=0).initialize() + + +@pytest.mark.unit +def test_the_hand_rolled_rice_is_gone(): + """No private likelihood in the scaling package. + + It set ``beta = sigma_obs**2`` -- pairing a measurement sigma with a Rice ``Sigma``, + which asserts an isotropic *complex* error where sigma_obs carries no phase at all. + ``xray_likelihoods`` records that no regime makes that correct, and the taxonomy + deliberately offers no such row; a copy inside the scaler bypassed both. + """ + import torchref.scaling.collection_scaler as cs + + src = inspect.getsource(cs) + assert "rice_math" not in src, "the scaler grew a private Rice likelihood again" + assert "SigmaAEstimator" not in src, ( + "the scaler precomputes beta again; rows own their own estimator" + ) + assert "create_xray_target" in src, "the joint fit must build a taxonomy row" + + +@pytest.mark.unit +def test_it_offers_exactly_the_selectable_objectives(): + from torchref.scaling.collection_scaler import CollectionScaler + from torchref.scaling.scaler_base import DEFAULT_SCALE_TARGET, SCALE_TARGETS + + sig = inspect.signature(CollectionScaler.refine_lbfgs_joint) + assert sig.parameters["scale_target"].default == DEFAULT_SCALE_TARGET + assert DEFAULT_SCALE_TARGET == "ls", ( + "the joint fit's default must track the single-dataset one" + ) + assert "nll" in SCALE_TARGETS and "ml_noalpha" in SCALE_TARGETS + + +@pytest.mark.integration +def test_unknown_objective_fails_closed(collection): + scaler = _fresh_scaler(collection) + with pytest.raises(ValueError, match="scale_target must be one of"): + scaler.refine_lbfgs_joint(scale_target="nll_i") + + +@pytest.mark.integration +@pytest.mark.parametrize("scale_target", ["ls", "nll", "ml_noalpha"]) +def test_every_objective_fits_finite_parameters(collection, scale_target): + """Every selectable row must drive the joint fit to finite parameters. + + The old fit produced non-finite scales on real data; the objective it handed L-BFGS + was unnormalised against absolute tolerances, which is a documented way to get there. + """ + scaler = _fresh_scaler(collection) + m = scaler.refine_lbfgs_joint( + nsteps=2, max_iter=20, verbose=False, scale_target=scale_target + ) + for p in scaler.parameters(): + assert torch.isfinite(p).all(), f"{scale_target}: non-finite scale parameter" + assert m["rwork"] and all(0.0 < r < 1.0 for r in m["rwork"]), m["rwork"] + + +@pytest.mark.integration +def test_the_dataset_view_shares_the_parents_parameters(collection): + """A row's scaler must be a view, not a copy. + + If the view registered the parent as a submodule, the parent's parameters would be + counted twice and L-BFGS would see duplicate leaves. If it copied them, the fit would + optimise something the collection never reads. + """ + from torchref.scaling.collection_scaler import _DatasetScalerView + + scaler = _fresh_scaler(collection) + dc, mc = collection + fracs = mc[mc.dark_key].fractions.detach() + view = _DatasetScalerView(scaler, fracs) + + # No parameters of its own -- only the bound fractions buffer. + assert list(view.parameters()) == [] + # And it routes through the parent's mixed-solvent path. + with torch.no_grad(): + fcalc = mc[mc.dark_key](dc[mc.dark_key].hkl) + got = view(fcalc) + want = scaler.forward_mixed(fcalc, fracs) + assert torch.equal(got, want) + + +@pytest.mark.integration +def test_the_objective_is_normalised(collection): + """The loss L-BFGS sees must be O(1), not the data's own magnitude. + + ``tolerance_grad``/``tolerance_change`` are absolute. This fit had no normaliser at + all, so on a large work set under unit weights the loss reached a magnitude where its + own float32 ulp exceeded the decrease the line search was trying to resolve. + """ + src = inspect.getsource( + __import__( + "torchref.scaling.collection_scaler", fromlist=["CollectionScaler"] + ).CollectionScaler.refine_lbfgs_joint + ) + assert "_norm" in src, "the joint objective is unnormalised again" + # The U penalty must NOT follow the observable/objective: sharing a normaliser that + # moves with the objective silently changes the regularisation strength. + assert "work.F" in src, "the normaliser must be built from amplitudes" diff --git a/torchref/scaling/collection_scaler.py b/torchref/scaling/collection_scaler.py index 2384ea66..655b86f7 100644 --- a/torchref/scaling/collection_scaler.py +++ b/torchref/scaling/collection_scaler.py @@ -14,10 +14,12 @@ import torch.nn as nn from torchref.base.metrics.rfactor import rfactor_work_free -from torchref.base.reciprocal import get_scattering_vectors -from torchref.base.targets.xray_likelihoods import complex_var_from_beta, rice_math from torchref.config import get_float_dtype -from torchref.scaling.scaler_base import ScalerBase +from torchref.scaling.scaler_base import ( + DEFAULT_SCALE_TARGET, + SCALE_TARGETS, + ScalerBase, +) from torchref.scaling.solvent import SS_HALF_BOUNDS, SolventModel from torchref.utils.utils import ModuleReference @@ -26,6 +28,55 @@ from torchref.model.model_collection import ModelCollection +class _DatasetScalerView(nn.Module): + """One dataset's view of a shared :class:`CollectionScaler`. + + Exists so the scale fit can hand a plain scaler to a taxonomy row. A row calls + ``self._scaler(fcalc)`` and knows nothing else about scaling, but the collection's + bulk solvent is a fraction-weighted mixture that depends on *which* dataset is being + scaled -- information :meth:`ScalerBase.forward` has no way to carry. This forwards to + :meth:`CollectionScaler.forward_mixed` with the fractions bound. + + The parent is held through :class:`~torchref.utils.utils.ModuleReference`, so its + parameters are **not** re-registered here: the optimiser is still built from the one + scaler, and a row holding this view contributes no leaves of its own. + """ + + def __init__(self, parent: "CollectionScaler", fractions: torch.Tensor): + super().__init__() + self._parent = ModuleReference(parent) + self.register_buffer("_fractions", fractions.detach().clone()) + + @property + def device(self): + """The parent's device. + + Needed explicitly: a target's ``_adopt_device`` reads ``scaler.device``, and + ``nn.Module.__getattr__`` raises before :class:`ModuleReference` gets a chance to + forward it -- the reference lives in ``__dict__``, so attribute lookup on *this* + object never reaches it. + """ + return self._parent.device + + def __getattr__(self, name): + """Anything this view does not own belongs to the parent. + + ``nn.Module.__getattr__`` resolves parameters, buffers and submodules first; a + row reading some other scaler attribute (bin edges, resolution limits) should see + the parent's, not an AttributeError. + """ + try: + return super().__getattr__(name) + except AttributeError: + parent = self.__dict__.get("_parent") + if parent is None: + raise + return getattr(parent, name) + + def forward(self, fcalc: torch.Tensor) -> torch.Tensor: + return self._parent.forward_mixed(fcalc, self._fractions) + + class CollectionScaler(ScalerBase): """ Joint scaler for DatasetCollection + ModelCollection. @@ -348,12 +399,29 @@ def refine_lbfgs_joint( max_iter: int = 200, history_size: int = 10, verbose: bool = True, + scale_target: str = DEFAULT_SCALE_TARGET, ) -> dict: """ - Refine scale parameters using LBFGS against **all** datasets. - - The closure sums the NLL across every matched dataset–model pair, - so a single set of scale parameters is fitted jointly. + Refine the shared scale parameters against **all** datasets jointly. + + One set of scale parameters serves every matched dataset-model pair, so the + closure sums a per-dataset objective. Each dataset's term is built from a row of + :data:`~torchref.refinement.targets.xray._specs.XRAY_TARGETS`, exactly as + :meth:`ScalerBase.refine_lbfgs` builds its single-dataset one -- so both scale + fits evaluate the same likelihood code, and neither carries a private copy of it. + + The row sees this dataset's own **mixed** bulk solvent, via a + :class:`_DatasetScalerView` that shares the parent's parameters and applies + :meth:`forward_mixed`. That is why the scaler cannot simply be handed to the + target: the solvent depends on which dataset's fractions are in play, and the + plain :meth:`ScalerBase.forward` has no way to know. + + Amplitudes throughout, whatever observable the *refinement* target fits. + Unit-weight least squares on intensities would put leverage where the data is + strongest: the residual goes as ``2 F dF``, so the squared residual carries an + extra factor of ``F**2`` and a global scale plus B plus anisotropy would be + determined almost entirely by the strongest low-resolution reflections, leaving + high resolution unconstrained. Parameters ---------- @@ -367,117 +435,122 @@ def refine_lbfgs_joint( LBFGS history size. verbose : bool Print progress. + scale_target : str, optional + Objective, one of :data:`~torchref.scaling.scaler_base.SCALE_TARGETS`. + Defaults to :data:`~torchref.scaling.scaler_base.DEFAULT_SCALE_TARGET` + (unit-weight ``ls``) -- the same default and the same reason as the + single-dataset fit. Returns ------- dict Refinement metrics (steps, rwork, rfree of dark dataset). + + Raises + ------ + ValueError + If ``scale_target`` is not in + :data:`~torchref.scaling.scaler_base.SCALE_TARGETS`. """ + if scale_target not in SCALE_TARGETS: + raise ValueError( + f"scale_target must be one of {SCALE_TARGETS}, got {scale_target!r}" + ) # Local import, deliberately: `torchref.refinement` imports # `torchref.scaling` at module scope, so hoisting these to module level # closes an import cycle. Do not "tidy" them up. - from torchref.refinement.model_error_estimation.sigma_a import SigmaAEstimator, epsilon_from_hkl from torchref.refinement.loss_state import LossState + from torchref.refinement.targets.xray import create_xray_target dc = self._dataset_collection mc = self._model_collection all_keys = [mc.dark_key] + mc.timepoint_names - # Pre-compute all fcalc (detached) plus the per-dataset σ_A model-error - # variance (beta/epsilon). beta is estimated ONCE on each dataset's free - # set from the currently-scaled |F_calc| and held detached during the fit, - # as in the single-dataset ``ScalerBase.refine_lbfgs``; that variance is - # what stops the scale collapsing toward zero in weak shells. + # Per dataset: a detached F_calc, and a table row that scales it through this + # dataset's solvent view. `model=None`, so the row never recomputes structure + # factors and the only leaves in the graph are this scaler's own parameters. + # Rows that need a model-error estimate build and own one themselves, as under + # the body refinement -- there is nothing to precompute here. fcalc_cache = {} fractions_cache = {} - beta_cache = {} - eps_cache = {} - work_cache = {} - centric_cache = {} + terms = [] for name in all_keys: if name not in dc: continue data = dc[name] model = mc[name] - hkl = data.hkl - fobs, sigma = data.get_corrected_data() with torch.no_grad(): - fc = model(hkl).detach() - fracs = model.fractions.detach() - f_sol_raw = self.get_mixed_solvent_raw(fracs) - scaled0 = super(CollectionScaler, self).forward( - fc, f_sol_override=f_sol_raw - ) - fc_amp0 = torch.abs(scaled0).reshape(-1) - fobs0 = fobs.to(fc_amp0.dtype).reshape(-1) - eps0 = epsilon_from_hkl( - hkl, getattr(data, "spacegroup", None) - ).to(fc_amp0.dtype) - s = get_scattering_vectors(hkl, data.cell) - dss0 = (torch.norm(s, dim=1) ** 2).to(fc_amp0.dtype) - # sigma_obs must be passed, as at every other call site: it is what - # makes sigma_A the correlation with the noise-free amplitudes. - _est = SigmaAEstimator().get( - fobs0, fc_amp0, data.centric, eps0, dss0, data.free.mask, - sigma_obs=sigma.to(fc_amp0.dtype).reshape(-1), + fcalc_cache[name] = model(data.hkl).detach() + fractions_cache[name] = model.fractions.detach() + view = _DatasetScalerView(self, fractions_cache[name]) + terms.append( + ( + create_xray_target( + data=data, + model=None, + scaler=view, + mode=scale_target, + use_set="work", + verbose=0, + device=self.device, + ), + fcalc_cache[name], ) - # TOTAL variance (`beta`, not `beta_model`): this scale fit uses the same - # likelihood as `ml`, which does not account for sigma_obs separately. - beta0, eps0 = _est.beta, _est.epsilon - fcalc_cache[name] = fc - fractions_cache[name] = fracs - beta_cache[name] = beta0 - eps_cache[name] = eps0 - work_cache[name] = data.work - centric_cache[name] = data.centric - - # Wrap the joint σ_A ML loss + U-penalty as a LossState target, reusing its - # NaN/Inf rejection. fcalc is detached, so the only leaves in the graph are - # the scaler's own parameters. + ) + + # One constant applied to every term, so the objective is an exact rescaling. + # ``torch.optim.LBFGS`` converges on ABSOLUTE tolerances, so the objective has to + # be O(1) for them to mean anything; unnormalised it carries the data's own + # magnitude and the float32 ulp of the loss exceeds the decrease the line search + # is trying to resolve. This fit had NO normaliser at all before. + with torch.no_grad(): + ssq = sum( + float(dc[n].work.F.detach().pow(2).sum()) + for n in fcalc_cache + ) + _norm = 1.0 / max(ssq, 1e-30) scaler_self = self class _CollectionScalerJointTarget(nn.Module): + """The table rows, closed over their detached ``fcalc``.""" + name = "scaler/joint" def forward(self): - total = torch.tensor(0.0, device=scaler_self.device) - n = 0 - for nm in all_keys: - if nm not in fcalc_cache: - continue - fc = fcalc_cache[nm] - fracs = fractions_cache[nm] - f_sol_raw = scaler_self.get_mixed_solvent_raw(fracs) - scaled = super(CollectionScaler, scaler_self).forward( - fc, f_sol_override=f_sol_raw - ) - # σ_A (Read MLF) scale-fit on the WORK set, with detached - # free-set beta/epsilon — same likelihood the body - # refinement uses. - amp = torch.abs(scaled).reshape(-1) - work = work_cache[nm] - F_obs = work.F.to(amp.dtype) - Fc = work.select(amp) - beta_w = work.select(beta_cache[nm]).to(F_obs.dtype) - eps_w = ( - work.select(eps_cache[nm]).to(F_obs.dtype) - if eps_cache[nm] is not None - else None - ) - centric_w = work.select(centric_cache[nm]) - loss = rice_math( - F_obs, Fc, complex_var_from_beta(beta_w, eps_w), centric_w - ) + total = torch.zeros((), device=scaler_self.device) + for target, fc in terms: + loss = target(fcalc=fc) + # Skip a dataset whose term went non-finite rather than poisoning + # the whole joint gradient with it. if torch.isfinite(loss): total = total + loss - n += 1 - if n > 0: - total = total / n - u_penalty = torch.sum(scaler_self.U**2) - return total + u_penalty + return total * _norm + + def maintenance(self): + """Forward the hook so sigma_A rows drop their ``beta`` cache after a + step block, as they do under the body refinement.""" + for target, _ in terms: + maint = getattr(target, "maintenance", None) + if maint is not None: + maint() + + class _CollectionScalerUPenalty(nn.Module): + """``sum(U**2)`` on the anisotropic scale tensor. + + Its normaliser is pinned to **amplitudes** rather than following the + objective. Sharing ``_norm`` would make the penalty's weight relative to the + likelihood depend on which objective was selected, which is a silent change + of regularisation strength dressed up as a change of objective. + """ + + name = "scaler/u_penalty" + + def forward(self): + return torch.sum(scaler_self.U**2) * _norm state = LossState(device=self.device) state.register_target("scaler/joint", _CollectionScalerJointTarget()) + state.register_target("scaler/u_penalty", _CollectionScalerUPenalty()) optimizer = torch.optim.LBFGS( self.parameters(), From 4bfcbd000b69a4b1a993f9788283dec8bed88ab1 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 01:56:03 +0200 Subject: [PATCH 075/250] Add a node-field ADP parameter wrapper DisorderFieldTensor stores disorder parameters on K nodes and gives each atom a distance-weighted mean over its k nearest, so the parameter count scales with node count rather than atom count. Storage is (K, 2) holding [log B, log sigma] per node; forward() returns per-atom isotropic B, so it fits the model.adp slot and every consumer of adp() works unchanged. Node positions are derived as the centroid of each node's anchor neighbourhood rather than refined, so a node stays inside the molecule and the optimiser has no free coordinate to drift with. Weights are a softmax over each atom's candidate nodes, so B is a convex combination of positive node values and needs no clamp. Coordinates come from an accessor injected at construction, held through ModuleReference so it stays out of state_dict, .to() and deepcopy. That leaves the inherited forward cache blind to coordinate changes, since CachedForwardMixin fingerprints parameters, buffers and call arguments only, so _fingerprint_state folds the accessor's output into the key. Nothing in Model references this yet. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- tests/helpers/device_cases.py | 28 ++ tests/unit/model/test_disorder_field.py | 441 +++++++++++++++++ torchref/model/disorder_field.py | 609 ++++++++++++++++++++++++ 3 files changed, 1078 insertions(+) create mode 100644 tests/unit/model/test_disorder_field.py create mode 100644 torchref/model/disorder_field.py diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index a5eddc1a..6946dd62 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -249,6 +249,34 @@ def _cell(device): ).OccupancyTensor(device=d), "OccupancyTensor", ), + DeviceCase( + "DisorderFieldTensor_empty", + lambda d: __import__( + "torchref.model.disorder_field", fromlist=["DisorderFieldTensor"] + ).DisorderFieldTensor(device=d), + "DisorderFieldTensor", + ), + DeviceCase( + "DisorderFieldTensor_populated", + lambda d: __import__( + "torchref.model.disorder_field", fromlist=["DisorderFieldTensor"] + ).DisorderFieldTensor( + initial_values=torch.full((8,), 20.0), + xyz_fn=__import__( + "torchref.model.parameter_wrappers", fromlist=["MixedTensor"] + ).MixedTensor( + torch.arange(24, dtype=torch.float32).reshape(8, 3), device=d + ), + n_nodes=3, + k_neighbors=2, + device=d, + ), + "DisorderFieldTensor", + # The coordinate accessor is borrowed: a ModuleReference is absent from ``.to()`` + # by design, so in isolation nobody moves the referent alongside the field. + # ``Model`` owns both and moves them together. + ignore=("->ref",), + ), DeviceCase( "RigidXYZTensor_empty", lambda d: __import__( diff --git a/tests/unit/model/test_disorder_field.py b/tests/unit/model/test_disorder_field.py new file mode 100644 index 00000000..ebf97f38 --- /dev/null +++ b/tests/unit/model/test_disorder_field.py @@ -0,0 +1,441 @@ +"""The node-field ADP wrapper: weights, positivity, cache correctness, lifecycle. + +Two properties here are load-bearing rather than cosmetic. The weights must be a +normalised mixture over each atom's candidate nodes, because that is what makes the +per-atom B a convex combination of positive node values and therefore positive without +a clamp. And the forward cache must notice that the coordinates moved: the wrapper reads +them through an injected accessor rather than a call argument, so +``CachedForwardMixin``'s own fingerprint cannot see them and +``DisorderFieldTensor._fingerprint_state`` has to fold them in. +""" + +import math + +import pytest +import torch + +from torchref.model.disorder_field import ( + DisorderFieldTensor, + build_neighbor_list, + farthest_point_anchors, +) +from torchref.model.parameter_wrappers import MixedTensor + + +@pytest.fixture +def coords(): + """A compact 3-D blob of atoms on a deterministic lattice.""" + g = torch.arange(6, dtype=torch.float64) + x, y, z = torch.meshgrid(g, g, g, indexing="ij") + return torch.stack([x.reshape(-1), y.reshape(-1), z.reshape(-1)], dim=1) * 1.7 + + +@pytest.fixture +def target_b(coords): + """A smooth B field with a real spatial gradient for the fit to chase.""" + return 20.0 + 1.5 * coords[:, 0] + 0.8 * coords[:, 2] + + +def _field(coords, target_b, **kw): + xyz = MixedTensor(coords.clone(), name="xyz") + kw.setdefault("n_nodes", 12) + kw.setdefault("k_neighbors", 6) + return DisorderFieldTensor( + initial_values=target_b, xyz_fn=xyz, dtype=torch.float64, **kw + ), xyz + + +@pytest.mark.unit +def test_anchor_selection_is_deterministic(coords): + """Same coordinates, same anchors -- no RNG anywhere in placement.""" + a = farthest_point_anchors(coords, 10) + b = farthest_point_anchors(coords, 10) + assert torch.equal(a, b) + assert a.shape[0] == 10 + assert a.dtype == torch.int64 + # Anchors are atom indices, and distinct. + assert int(a.max()) < coords.shape[0] + assert torch.unique(a).shape[0] == a.shape[0] + + +@pytest.mark.unit +def test_neighbor_list_is_nearest_first(coords): + """``build_neighbor_list`` returns each atom's k nearest nodes, closest first.""" + node_pos = coords[farthest_point_anchors(coords, 8)] + nl = build_neighbor_list(coords, node_pos, 4) + assert nl.shape == (coords.shape[0], 4) + d = torch.cdist(coords, node_pos) + gathered = torch.gather(d, 1, nl) + assert bool((gathered.diff(dim=1) >= -1e-12).all()), "not sorted by distance" + assert torch.equal(nl[:, 0], d.argmin(dim=1)) + + +@pytest.mark.unit +def test_weights_are_a_normalised_mixture(coords, target_b): + """Rows sum to one and are non-negative -- what makes the output positive.""" + field, _ = _field(coords, target_b) + W = field.weights() + assert W.shape == (coords.shape[0], 6) + assert bool((W >= 0).all()) + assert torch.allclose(W.sum(dim=1), torch.ones(coords.shape[0], dtype=W.dtype)) + + +@pytest.mark.unit +def test_output_is_per_atom_and_positive(coords, target_b): + """Public space is per-atom even though storage is per-node.""" + field, _ = _field(coords, target_b) + out = field() + assert out.shape == (coords.shape[0],) + assert field.shape == (coords.shape[0],) + assert field.node_shape == (12, 2) + assert bool((out > 0).all()) + assert bool(torch.isfinite(out).all()) + + +@pytest.mark.unit +def test_single_node_flat_kernel_is_a_constant(coords): + """K=1 reproduces a constant B exactly -- the analytic control. + + With one node every atom's weight vector is ``[1.0]`` whatever the distance, so the + field degenerates to a single scalar and must return it uniformly. + """ + b = torch.full((coords.shape[0],), 37.5, dtype=torch.float64) + field, _ = _field(coords, b, n_nodes=1, k_neighbors=1) + out = field() + assert torch.allclose(out, b, atol=1e-9), f"got spread {out.min()}..{out.max()}" + + +@pytest.mark.unit +def test_fit_tracks_a_smooth_gradient(coords, target_b): + """A 12-node field on a linear B ramp beats the best constant by a wide margin.""" + field, _ = _field(coords, target_b) + resid = (field() - target_b).pow(2).mean().sqrt() + constant = (target_b - target_b.mean()).pow(2).mean().sqrt() + assert resid < 0.25 * constant, f"rmse {resid:.3f} vs constant {constant:.3f}" + + +@pytest.mark.unit +def test_more_nodes_fit_better(coords, target_b): + """Reconstruction improves monotonically with node count on a smooth target.""" + errors = [] + for n in (2, 8, 32): + field, _ = _field(coords, target_b, n_nodes=n, k_neighbors=min(6, n)) + errors.append(float((field() - target_b).detach().pow(2).mean().sqrt())) + assert errors[0] > errors[1] > errors[2], errors + + +@pytest.mark.unit +def test_fingerprint_sees_the_coordinates(coords, target_b): + """The load-bearing test for the injected accessor, and it proves its own point. + + ``forward()`` takes no arguments, so the coordinates reach it through ``_xyz_fn``. + The inherited fingerprint covers only this module's own parameters and buffers, so + it cannot see them -- asserted directly here, by checking the base implementation + does NOT change while the override does. Without + ``DisorderFieldTensor._fingerprint_state`` the cache would serve a B computed at + coordinates that have since moved. + """ + from torchref.utils.caching import CachedForwardMixin + + field, xyz = _field(coords, target_b) + base_before = CachedForwardMixin._fingerprint_state(field) + full_before = field._fingerprint_state() + + with torch.no_grad(): + xyz.refinable_params[:8, 0] += 4.0 + + assert CachedForwardMixin._fingerprint_state(field) == base_before, ( + "the field's own parameters and buffers did not move, so the inherited " + "fingerprint is blind to this change -- which is why the override exists" + ) + assert field._fingerprint_state() != full_before, "override missed the coordinates" + + +@pytest.mark.unit +def test_cache_returns_a_fresh_value_after_a_non_rigid_move(coords, target_b): + """A change of relative geometry must reach the output, not a stale cache. + + The perturbation has to be non-rigid: the field is translation-invariant by + construction, so shifting every atom equally moves the nodes with them and is + *correctly* a no-op. + """ + field, xyz = _field(coords, target_b) + before = field().clone() + + with torch.no_grad(): + xyz.refinable_params[: coords.shape[0] // 2, 0] += 4.0 + + after = field() + assert not torch.allclose(before, after), "stale cache: coordinates were ignored" + + +@pytest.mark.unit +def test_rigid_translation_leaves_the_field_unchanged(coords, target_b): + """The invariance that makes the previous test need a non-rigid perturbation.""" + field, xyz = _field(coords, target_b) + before = field().clone() + + with torch.no_grad(): + xyz.refinable_params += 9.0 + + assert torch.allclose(field(), before, atol=1e-9) + + +@pytest.mark.unit +def test_node_positions_follow_the_coordinates(coords, target_b): + """Node positions are derived from the atoms, so a rigid shift carries them along.""" + field, xyz = _field(coords, target_b) + before = field.node_positions().clone() + + with torch.no_grad(): + xyz.refinable_params += 2.5 + + after = field.node_positions() + assert torch.allclose(after - before, torch.full_like(before, 2.5), atol=1e-9) + + +@pytest.mark.unit +def test_gradient_reaches_nodes_and_coordinates(coords, target_b): + """Both channels are live through the plain zero-arg forward path.""" + field, xyz = _field(coords, target_b) + field().sum().backward() + assert field.refinable_params.grad is not None + assert bool(field.refinable_params.grad.abs().sum() > 0) + assert xyz.refinable_params.grad is not None + assert bool(xyz.refinable_params.grad.abs().sum() > 0) + + +@pytest.mark.unit +def test_gradcheck_on_both_channels(coords, target_b): + """Analytic gradients match finite differences, for node values and coordinates. + + Checked through :meth:`DisorderFieldTensor.evaluate`, which is the field's + arithmetic without the accessor or the cache. Rebinding ``refinable_params`` inside + a gradcheck closure would not work: ``nn.Parameter(p)`` is a fresh leaf, so the + graph back to ``p`` is severed and there is nothing to check. + """ + small = coords[:20].clone() + field, _ = _field(small, target_b[:20], n_nodes=3, k_neighbors=3) + + raw = field.node_values().detach().clone().requires_grad_(True) + xyz_in = small.detach().clone().requires_grad_(True) + + assert torch.autograd.gradcheck( + field.evaluate, (xyz_in, raw), eps=1e-6, atol=1e-5 + ) + + +@pytest.mark.unit +def test_list_adequacy_invariant_is_exposed(coords, target_b): + """The smallest candidate weight is reported, and shrinks as k grows.""" + tight, _ = _field(coords, target_b, n_nodes=16, k_neighbors=2) + loose, _ = _field(coords, target_b, n_nodes=16, k_neighbors=12) + assert loose.smallest_candidate_weight() < tight.smallest_candidate_weight() + assert 0.0 <= loose.smallest_candidate_weight() <= 1.0 + + +@pytest.mark.unit +def test_rebuild_neighbor_list_is_explicit_and_refreshes(coords, target_b): + """Membership only changes when the caller asks; the rebuild then takes effect.""" + field, xyz = _field(coords, target_b) + original = field.neighbor_list.clone() + + with torch.no_grad(): + xyz.refinable_params[:, 0] += 40.0 # move atoms far past the nodes + + assert torch.equal(field.neighbor_list, original), "list changed without a rebuild" + field.rebuild_neighbor_list() + assert field.neighbor_list.shape == original.shape + assert bool(torch.isfinite(field()).all()) + + +@pytest.mark.unit +def test_atom_space_mask_collapses_to_node_space(coords, target_b): + """Masks arrive in ATOM space and are collapsed with OR onto the nodes.""" + field, _ = _field(coords, target_b) + n_atoms = coords.shape[0] + + mask = torch.zeros(n_atoms, dtype=torch.bool) + mask[0] = True + field.update_refinable_mask(mask) + + assert field.refinable_mask.shape == (field.n_nodes,) + served = field.neighbor_list[0] + assert bool(field.refinable_mask[served].all()), "atom 0's nodes must be refinable" + assert int(field.refinable_mask.sum()) == int(torch.unique(served).numel()) + + +@pytest.mark.unit +def test_wrong_sized_mask_is_rejected(coords, target_b): + """A node-space mask passed as atom space (or vice versa) is an error, not a guess.""" + field, _ = _field(coords, target_b) + with pytest.raises(ValueError, match="Atom-space mask"): + field.update_refinable_mask(torch.ones(field.n_nodes, dtype=torch.bool)) + with pytest.raises(ValueError, match="Node-space mask"): + field.update_refinable_mask( + torch.ones(coords.shape[0], dtype=torch.bool), in_node_space=True + ) + + +@pytest.mark.unit +def test_freezing_all_nodes_leaves_the_output_intact(coords, target_b): + """Repartitioning moves values between storage halves without changing them.""" + field, _ = _field(coords, target_b) + before = field().clone() + field.update_refinable_mask( + torch.zeros(field.n_nodes, dtype=torch.bool), in_node_space=True + ) + assert int(field.get_refinable_count()) == 0 + assert torch.allclose(field(), before, atol=1e-12) + + +@pytest.mark.unit +def test_per_atom_assignment_is_refused(coords, target_b): + """K nodes cannot represent arbitrary per-atom values, so writing is not silent.""" + field, _ = _field(coords, target_b) + with pytest.raises(NotImplementedError, match="per-atom assignment"): + field[0] = 42.0 + + +@pytest.mark.unit +def test_refit_moves_the_field_to_a_new_target(coords, target_b): + """``refit`` is the representable alternative to per-atom assignment.""" + field, _ = _field(coords, target_b) + new_target = target_b * 0.5 + 5.0 + field.refit(new_target) + resid = (field() - new_target).pow(2).mean().sqrt() + baseline = (new_target - new_target.mean()).pow(2).mean().sqrt() + assert resid < 0.25 * baseline + + +@pytest.mark.unit +def test_copy_is_independent_but_shares_the_accessor(coords, target_b): + """Parameters are copied; the coordinate accessor is deliberately shared. + + Deep-copying the accessor is what breaks ``Restraints.copy()`` today: ``deepcopy`` + walks the borrowed wrapper whose cache can hold a graph-attached tensor. + """ + field, xyz = _field(coords, target_b) + clone = field.copy() + + assert clone._xyz_fn.module is xyz, "accessor must reference the same wrapper" + assert clone.refinable_params.data_ptr() != field.refinable_params.data_ptr() + assert torch.allclose(clone(), field()) + + with torch.no_grad(): + clone.refinable_params[:, 0] += 1.0 + clone.reset_forward_cache() + field.reset_forward_cache() + assert not torch.allclose(clone(), field()), "copy is not independent" + + +@pytest.mark.unit +def test_copy_survives_an_evaluated_forward(coords, target_b): + """``copy()`` works after ``forward()`` has populated the accessor's cache.""" + field, _ = _field(coords, target_b) + field() # populate xyz's forward cache with a graph-attached tensor + clone = field.copy() + assert torch.allclose(clone(), field()) + + +@pytest.mark.unit +def test_state_dict_excludes_the_accessor_and_round_trips(coords, target_b): + """The callable is not state; everything else survives a save/load exactly. + + Restored the way ``Model.create_from_state_dict`` does it: build a wrapper with real + values to fix the shapes and masks, then let ``load_state_dict`` overwrite them. The + empty shell exists for shape-less construction, not as a load target. + """ + field, xyz = _field(coords, target_b) + # Cloned deliberately: state_dict() hands back detached REFERENCES, so a later + # in-place edit of refinable_params would silently rewrite the saved state too. + sd = {key: value.clone() for key, value in field.state_dict().items()} + assert not any("xyz_fn" in key for key in sd), sd.keys() + expected = field().clone() + + restored, _ = _field(coords, target_b) + with torch.no_grad(): + restored.refinable_params[:, 0] += 0.75 # move it off the saved state + restored.reset_forward_cache() + assert not torch.allclose(restored(), expected), "perturbation did not take" + + restored.load_state_dict(sd) + restored.reset_forward_cache() + + assert torch.allclose(restored(), expected, atol=1e-12) + + +@pytest.mark.unit +def test_empty_shell_needs_no_accessor(coords): + """The ``load_state_dict`` entry point constructs without coordinates.""" + shell = DisorderFieldTensor(dtype=torch.float64) + assert shell.neighbor_list is None + assert shell.shape == (0,) + + +@pytest.mark.unit +def test_accessor_is_required_when_values_are_given(coords, target_b): + """Nodes cannot be placed without coordinates, and that fails loudly.""" + with pytest.raises(ValueError, match="xyz_fn"): + DisorderFieldTensor(initial_values=target_b, dtype=torch.float64) + + +@pytest.mark.unit +def test_accessor_is_not_a_submodule_or_buffer(coords, target_b): + """It must stay out of module traversal, or a device move duplicates the graph.""" + field, xyz = _field(coords, target_b) + from torchref.utils.utils import ModuleReference + + assert isinstance(field._xyz_fn, ModuleReference) + assert xyz not in list(field.modules()) + assert not any(m is xyz for _, m in field.named_modules()) + # xyz's parameters must not appear among the field's own. + field_ptrs = {p.data_ptr() for p in field.parameters()} + assert xyz.refinable_params.data_ptr() not in field_ptrs + + +@pytest.mark.unit +def test_device_round_trip_keeps_buffers_together(coords, target_b): + """A ``.to()`` moves the node storage and the index buffers as one.""" + field, _ = _field(coords, target_b) + before = field().clone() + field.to(torch.device("cpu")) + assert field.neighbor_list.device.type == "cpu" + assert field.anchor_atom.device.type == "cpu" + assert torch.allclose(field(), before, atol=1e-12) + + +@pytest.mark.unit +def test_float32_and_float64_both_work(coords, target_b): + """The field works in either dtype without requiring one.""" + for dtype in (torch.float32, torch.float64): + xyz = MixedTensor(coords.to(dtype).clone(), name="xyz") + field = DisorderFieldTensor( + initial_values=target_b.to(dtype), + xyz_fn=xyz, + n_nodes=8, + k_neighbors=4, + dtype=dtype, + ) + out = field() + assert out.dtype == dtype + assert bool(torch.isfinite(out).all()) + + +@pytest.mark.unit +def test_ragged_anchor_neighbourhoods_average_their_atoms(coords, target_b): + """A node anchored on several atoms sits at their centroid.""" + xyz = MixedTensor(coords.clone(), name="xyz") + anchor_atom = torch.tensor([0, 1, 2, 10, 11], dtype=torch.int64) + anchor_node = torch.tensor([0, 0, 0, 1, 1], dtype=torch.int64) + field = DisorderFieldTensor( + initial_values=target_b, + xyz_fn=xyz, + k_neighbors=2, + anchor_rows=(anchor_atom, anchor_node), + dtype=torch.float64, + ) + pos = field.node_positions() + assert pos.shape == (2, 3) + assert torch.allclose(pos[0], coords[[0, 1, 2]].mean(dim=0)) + assert torch.allclose(pos[1], coords[[10, 11]].mean(dim=0)) diff --git a/torchref/model/disorder_field.py b/torchref/model/disorder_field.py new file mode 100644 index 00000000..c2eec6c5 --- /dev/null +++ b/torchref/model/disorder_field.py @@ -0,0 +1,609 @@ +"""Node-field parametrization of the atomic displacement parameters. + +A disorder field stores disorder parameters on a small set of **nodes** and gives each +atom a distance-weighted mean of the nodes near it, so the parameter count scales with +node count rather than atom count. + +This is :class:`~torchref.model.parameter_wrappers.OccupancyTensor`'s collapse-and-expand +with a soft, distance-derived expansion in place of a fixed integer assignment: storage +is ``(K, 2)`` per node, ``forward()`` returns one B per atom. Two index spaces therefore +meet in this class, and callers must not mix them --- masks handed to +:meth:`~DisorderFieldTensor.update_refinable_mask` are in ATOM space, while +``refinable_mask`` and :meth:`get_refinable_count` are in NODE space. + +A node's position is *derived*, not refined: it is the centroid of the atoms within +``anchor_radius`` bonds of its anchor atom. That keeps a node inside the molecule, +confined to one connected fragment, and moving with the model, and it leaves the +optimiser no free coordinate to wander with. +""" + +import math +from typing import Callable, Optional, Tuple + +import torch +from torch import nn + +from torchref.config import get_float_dtype, normalize_device +from torchref.model.parameter_wrappers import MixedTensor +from torchref.utils.utils import ModuleReference + +__all__ = ["DisorderFieldTensor", "farthest_point_anchors", "build_neighbor_list"] + + +def farthest_point_anchors(xyz: torch.Tensor, n_nodes: int) -> torch.Tensor: + """Pick ``n_nodes`` well-spread anchor atoms, then relax them onto local density. + + Greedy farthest-point selection seeded from the atom nearest the centroid, followed + by Lloyd iterations that move each anchor to the atom closest to its cluster mean. + Farthest-point alone favours extremities; the Lloyd pass pulls the anchors back onto + where atoms actually are, which is what a disorder field wants. + + Deterministic end to end --- no RNG --- so two processes produce the same anchors. + Restraint row order was hash-seed dependent in this package once; node placement is + not going to be. + + Parameters + ---------- + xyz : torch.Tensor + ``(N, 3)`` atom coordinates. + n_nodes : int + Number of anchors to pick. Clamped to ``N``. + + Returns + ------- + torch.Tensor + ``(K,)`` int64 atom indices, sorted ascending. + """ + n_atoms = int(xyz.shape[0]) + n_nodes = max(1, min(int(n_nodes), n_atoms)) + + centroid = xyz.mean(dim=0, keepdim=True) + first = int(torch.cdist(centroid, xyz).argmin()) + + chosen = [first] + d2_nearest = ((xyz - xyz[first]) ** 2).sum(-1) + for _ in range(n_nodes - 1): + nxt = int(d2_nearest.argmax()) + chosen.append(nxt) + d2_nearest = torch.minimum(d2_nearest, ((xyz - xyz[nxt]) ** 2).sum(-1)) + + anchors = torch.tensor(chosen, dtype=torch.int64, device=xyz.device) + + # Lloyd relaxation, snapping to real atoms so an anchor is always an atom index. + for _ in range(10): + assign = torch.cdist(xyz, xyz[anchors]).argmin(dim=1) + moved = anchors.clone() + for j in range(anchors.shape[0]): + members = (assign == j).nonzero(as_tuple=True)[0] + if members.numel() == 0: + continue + mean = xyz[members].mean(dim=0, keepdim=True) + moved[j] = members[int(torch.cdist(mean, xyz[members]).argmin())] + moved = torch.unique(moved) + if moved.shape[0] == anchors.shape[0] and bool((moved == anchors).all()): + break + anchors = moved + + return torch.sort(anchors).values + + +def build_neighbor_list( + xyz: torch.Tensor, node_pos: torch.Tensor, k: int +) -> torch.Tensor: + """The ``k`` nearest nodes to each atom. + + A dense ``cdist`` plus ``topk``. Node counts are small by construction --- the whole + point of the representation --- so ``(N, K)`` stays cheap and a spatial cell list + (``topology.nonbonded.build_cell_list``) would only add machinery. Revisit if node + counts ever approach atom counts. + + Parameters + ---------- + xyz : torch.Tensor + ``(N, 3)`` atom coordinates. + node_pos : torch.Tensor + ``(K, 3)`` node positions. + k : int + Candidates per atom. Clamped to ``K``. + + Returns + ------- + torch.Tensor + ``(N, k)`` int64 node indices, nearest first. + """ + k = max(1, min(int(k), int(node_pos.shape[0]))) + d = torch.cdist(xyz, node_pos) + return d.topk(k, dim=1, largest=False).indices.contiguous() + + +def _wrap_accessor(xyz_fn): + """Hold a coordinate accessor without registering it as a submodule. + + ``model.xyz`` is itself an ``nn.Module``, so a plain assignment would enrol it in + this wrapper's module tree and drag it into ``state_dict``, ``.to()`` and + ``deepcopy``. :class:`~torchref.utils.utils.ModuleReference` exists for exactly that + and is what the device-conformance walker already knows how to follow. A bare + callable needs no wrapping. + """ + if isinstance(xyz_fn, nn.Module): + return ModuleReference(xyz_fn) + return xyz_fn + + +class DisorderFieldTensor(MixedTensor): + """Per-atom ADPs from a small set of nodes, each atom a weighted mean of its k nearest. + + Storage is ``(K, 2)``: ``[log B, log sigma]`` per node. ``forward()`` returns + ``(n_atoms,)`` isotropic B, so this drops into the ``model.adp`` slot and every + consumer of ``adp()`` keeps working unchanged. + + The weight of node ``j`` at atom ``i`` is ``softmax_j(-d_ij^2 / 2 sigma_j^2)`` over + that atom's candidate list, so weights are non-negative and sum to one and B is a + convex combination of positive node values --- positive for free, with no clamping. + + Coordinates come from an accessor injected at construction rather than being passed + per call, which keeps ``forward()`` argument-free. That makes the inherited forward + cache incorrect on its own, since :class:`~torchref.utils.caching.CachedForwardMixin` + fingerprints parameters, buffers and call *arguments* --- and a borrowed accessor's + output is none of those. :meth:`_fingerprint_state` closes that by folding the + accessor's output into the key. + + Parameters + ---------- + initial_values : torch.Tensor, optional + ``(n_atoms,)`` isotropic B to fit the field to. Omit for an empty shell ready + for ``load_state_dict``. + xyz_fn : callable, optional + Returns the current ``(n_atoms, 3)`` coordinates. Typically ``model.xyz``. Held + by reference and deliberately invisible to ``state_dict``, device traversal and + ``copy``; re-attach with :meth:`set_xyz_fn` after a state-dict load. + n_nodes : int, optional + Number of nodes. Default 32. + k_neighbors : int, optional + Candidate nodes per atom. Default 12. Doubles as the skin margin that makes a + slightly stale candidate list harmless, so prefer generous over tight. + anchor_rows : tuple of torch.Tensor, optional + ``(flat atom indices, node index per entry)`` defining each node's anchor + neighbourhood. Omit to anchor every node at a single atom, which is what a model + without a topology gets. + node_values : torch.Tensor, optional + ``(K, 2)`` storage to adopt directly instead of fitting to ``initial_values``. + Used by :meth:`copy`. + refinable_mask : torch.Tensor, optional + Boolean mask. Interpreted in ATOM space unless ``mask_in_node_space``. + mask_in_node_space : bool, optional + Treat ``refinable_mask`` as already collapsed to ``(K,)``. Default False. + requires_grad : bool, optional + Whether node parameters carry gradients. Default True. + dtype, device : optional + Floating dtype and device. + name : str, optional + Wrapper name. Defaults to ``"adp"`` so ``Model`` consumers find it. + epsilon : float, optional + Floor on ``sigma`` and on fitted node B, in the same units as each. Default 1e-3. + """ + + def __init__( + self, + initial_values: Optional[torch.Tensor] = None, + xyz_fn: Optional[Callable[[], torch.Tensor]] = None, + n_nodes: int = 32, + k_neighbors: int = 12, + anchor_rows: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + node_values: Optional[torch.Tensor] = None, + refinable_mask: Optional[torch.Tensor] = None, + mask_in_node_space: bool = False, + requires_grad: bool = True, + dtype: Optional[torch.dtype] = None, + device: Optional[torch.device] = None, + name: Optional[str] = "adp", + epsilon: float = 1e-3, + ): + self.epsilon = epsilon + self._k_neighbors = int(k_neighbors) + object.__setattr__(self, "_xyz_fn", _wrap_accessor(xyz_fn)) + + if initial_values is None and node_values is None: + device = normalize_device(device) + dtype = dtype if dtype is not None else get_float_dtype() + super().__init__(None, None, requires_grad, dtype, device, name) + self._full_shape = 0 + self.register_buffer("neighbor_list", None) + self.register_buffer("anchor_atom", None) + self.register_buffer("anchor_node", None) + return + + if xyz_fn is None: + raise ValueError( + "DisorderFieldTensor needs xyz_fn to place its nodes; pass the model's " + "coordinate wrapper (e.g. model.xyz)." + ) + + xyz = xyz_fn().detach() + if dtype is None: + dtype = initial_values.dtype if initial_values is not None else xyz.dtype + if device is None: + device = xyz.device + xyz = xyz.to(dtype=dtype, device=device) + + n_atoms = int(xyz.shape[0]) + + if anchor_rows is None: + anchor_atom = farthest_point_anchors(xyz, n_nodes) + anchor_node = torch.arange( + anchor_atom.shape[0], dtype=torch.int64, device=device + ) + else: + anchor_atom, anchor_node = anchor_rows + anchor_atom = anchor_atom.to(device=device, dtype=torch.int64) + anchor_node = anchor_node.to(device=device, dtype=torch.int64) + + n_k = int(anchor_node.max()) + 1 + node_pos = self._segment_mean(xyz, anchor_atom, anchor_node, n_k) + neighbor_list = build_neighbor_list(xyz, node_pos, self._k_neighbors) + + if node_values is None: + node_values = self._fit_nodes( + initial_values.to(dtype=dtype, device=device), + xyz, + node_pos, + neighbor_list, + ) + node_values = node_values.to(dtype=dtype, device=device) + + if refinable_mask is None: + node_mask = torch.ones(n_k, dtype=torch.bool, device=device) + elif mask_in_node_space: + node_mask = refinable_mask.to(device=device, dtype=torch.bool) + else: + node_mask = self._collapse_mask( + refinable_mask.to(device=device, dtype=torch.bool), neighbor_list, n_k + ) + + # ``register_buffer`` needs ``nn.Module.__init__`` to have run, which happens + # inside this call, so every buffer below is registered after it. + super().__init__( + initial_values=node_values, + refinable_mask=node_mask, + requires_grad=requires_grad, + dtype=dtype, + device=device, + name=name, + ) + self._full_shape = n_atoms + self.register_buffer("anchor_atom", anchor_atom) + self.register_buffer("anchor_node", anchor_node) + self.register_buffer("neighbor_list", neighbor_list) + + # ------------------------------------------------------------------ + # Construction helpers. + # ------------------------------------------------------------------ + + @staticmethod + def _segment_mean( + xyz: torch.Tensor, + anchor_atom: torch.Tensor, + anchor_node: torch.Tensor, + n_nodes: int, + ) -> torch.Tensor: + """Mean coordinate of each node's anchor atoms, ``(K, 3)``. + + Differentiable in ``xyz``, which is what makes a node move with the model. + """ + acc = xyz.new_zeros(n_nodes, 3) + acc = acc.index_add(0, anchor_node, xyz[anchor_atom]) + counts = torch.zeros(n_nodes, dtype=xyz.dtype, device=xyz.device) + counts = counts.index_add( + 0, anchor_node, torch.ones_like(anchor_node, dtype=xyz.dtype) + ) + return acc / counts.clamp(min=1.0).unsqueeze(-1) + + @staticmethod + def _weights( + xyz: torch.Tensor, + node_pos: torch.Tensor, + neighbor_list: torch.Tensor, + log_sigma: torch.Tensor, + ) -> torch.Tensor: + """Softmax weights over each atom's candidate nodes, ``(n_atoms, k)``. + + Normalised across the candidates, so rows sum to one and no atom can end up + without support even when every node is far away. + """ + cand = node_pos[neighbor_list] + d2 = ((xyz.unsqueeze(1) - cand) ** 2).sum(-1) + sigma2 = torch.exp(2.0 * log_sigma)[neighbor_list] + return torch.softmax(-d2 / (2.0 * sigma2), dim=1) + + def _fit_nodes( + self, + target_b: torch.Tensor, + xyz: torch.Tensor, + node_pos: torch.Tensor, + neighbor_list: torch.Tensor, + ) -> torch.Tensor: + """Least-squares node values reproducing ``target_b`` as closely as possible. + + ``sigma`` is seeded at half the median nearest-neighbour node spacing, then the + node values are the ridged solution of ``W b = target_b``. Linear in ``b``, so + this is a closed-form solve rather than an optimisation loop. + + Returns + ------- + torch.Tensor + ``(K, 2)`` storage, ``[log b, log sigma]``. + """ + n_k = int(node_pos.shape[0]) + if n_k > 1: + dnode = torch.cdist(node_pos, node_pos) + dnode.fill_diagonal_(float("inf")) + spacing = float(dnode.min(dim=1).values.median()) + else: + spacing = float(xyz.std()) * 2.0 + sigma0 = max(spacing / 2.0, 10.0 * self.epsilon) + log_sigma = torch.full( + (n_k,), math.log(sigma0), dtype=xyz.dtype, device=xyz.device + ) + + W_sparse = self._weights(xyz, node_pos, neighbor_list, log_sigma) + W = torch.zeros(xyz.shape[0], n_k, dtype=xyz.dtype, device=xyz.device) + W.scatter_(1, neighbor_list, W_sparse) + + gram = W.T @ W + ridge = 1e-6 * torch.diagonal(gram).mean().clamp(min=1e-30) + eye = torch.eye(n_k, dtype=W.dtype, device=W.device) + b = torch.linalg.solve(gram + ridge * eye, W.T @ target_b) + + b = b.clamp(min=self.epsilon) + return torch.stack([torch.log(b), log_sigma], dim=1) + + @staticmethod + def _collapse_mask( + atom_mask: torch.Tensor, neighbor_list: torch.Tensor, n_nodes: int + ) -> torch.Tensor: + """Atom-space mask to node space: a node is refinable if any atom it serves is.""" + acc = torch.zeros(n_nodes, dtype=torch.bool, device=atom_mask.device) + served = neighbor_list[atom_mask] + if served.numel(): + acc[served.reshape(-1)] = True + return acc + + # ------------------------------------------------------------------ + # Public surface. + # ------------------------------------------------------------------ + + @property + def shape(self): + """Shape of the FULL per-atom tensor, not the node storage.""" + return (self._full_shape,) + + @property + def node_shape(self): + """Shape of the node storage.""" + return tuple(self.fixed_values.shape) + + @property + def n_nodes(self) -> int: + """Number of nodes.""" + return int(self.fixed_values.shape[0]) + + def set_xyz_fn(self, xyz_fn: Callable[[], torch.Tensor]) -> None: + """Attach the coordinate accessor. + + Needed after ``load_state_dict``, which cannot carry a callable. + """ + object.__setattr__(self, "_xyz_fn", _wrap_accessor(xyz_fn)) + + def node_positions(self, xyz: Optional[torch.Tensor] = None) -> torch.Tensor: + """Node positions, ``(K, 3)``, derived from the anchor neighbourhoods.""" + if xyz is None: + xyz = self._xyz_fn() + return self._segment_mean(xyz, self.anchor_atom, self.anchor_node, self.n_nodes) + + def weights(self, xyz: Optional[torch.Tensor] = None) -> torch.Tensor: + """Per-atom weights over candidate nodes, ``(n_atoms, k)``. Rows sum to one.""" + if xyz is None: + xyz = self._xyz_fn() + raw = super().forward() + return self._weights( + xyz, self.node_positions(xyz), self.neighbor_list, raw[:, 1] + ) + + def smallest_candidate_weight(self, xyz: Optional[torch.Tensor] = None) -> float: + """Largest per-atom minimum candidate weight --- the list-adequacy invariant. + + Each atom sees only its ``k`` candidate nodes. While the weakest of those + candidates carries negligible weight, the list brackets the real neighbourhood + and a node drifting in or out of it cannot move any ADP appreciably. When this + rises, the list no longer brackets it and should be rebuilt with + :meth:`rebuild_neighbor_list` or a larger ``k``. + """ + with torch.no_grad(): + return float(self.weights(xyz).min(dim=1).values.max()) + + def rebuild_neighbor_list( + self, xyz: Optional[torch.Tensor] = None, k_neighbors: Optional[int] = None + ) -> None: + """Recompute which nodes each atom sees, at the current coordinates. + + The candidate list is the slowly-varying combinatorial half of the field and is + never refreshed implicitly: membership is piecewise constant in position, so a + rebuild is a discrete jump and belongs at a point the caller chooses. + """ + if xyz is None: + xyz = self._xyz_fn() + xyz = xyz.detach() + if k_neighbors is not None: + self._k_neighbors = int(k_neighbors) + self.neighbor_list = build_neighbor_list( + xyz, self.node_positions(xyz).detach(), self._k_neighbors + ) + self.reset_forward_cache() + + def evaluate(self, xyz: torch.Tensor, raw: torch.Tensor) -> torch.Tensor: + """Per-atom B from explicit coordinates and node storage, ``(n_atoms,)``. + + The field's arithmetic, with no accessor and no cache in the way, so it can be + differentiated and checked on its own. ``forward()`` is this plus the plumbing + that fetches both arguments. + + Parameters + ---------- + xyz : torch.Tensor + ``(n_atoms, 3)`` coordinates. + raw : torch.Tensor + ``(K, 2)`` node storage, ``[log b, log sigma]``. + """ + log_b, log_sigma = raw[:, 0], raw[:, 1] + node_pos = self._segment_mean( + xyz, self.anchor_atom, self.anchor_node, raw.shape[0] + ) + W = self._weights(xyz, node_pos, self.neighbor_list, log_sigma) + return (W * torch.exp(log_b)[self.neighbor_list]).sum(dim=1) + + def forward(self) -> torch.Tensor: + """Per-atom isotropic B, ``(n_atoms,)``. + + A convex combination of positive node values, so strictly positive without a + clamp. Translation-invariant by construction: node positions are centroids of + atom coordinates, so a rigid shift of the model moves the nodes with it and + leaves every distance, and therefore every weight, unchanged. + """ + return self.evaluate(self._xyz_fn(), super().forward()) + + def node_values(self) -> torch.Tensor: + """The assembled node storage ``(K, 2)`` in raw ``[log b, log sigma]`` space.""" + return super().forward() + + def _fingerprint_state(self): + """Fold the accessor's coordinates into the forward-cache key. + + Without this the cache would be keyed on parameters and buffers alone and would + serve a per-atom B computed at coordinates that have since moved --- the + coordinates reach ``forward()`` through the accessor, not through an argument, + so the mixin cannot see them by itself. + """ + base = super()._fingerprint_state() + if self._xyz_fn is None: + return base + xyz = self._xyz_fn() + return base + ((xyz.data_ptr(), xyz._version),) + + def _set_values(self, key, value: torch.Tensor) -> None: + """Rejected: a node field cannot represent arbitrary per-atom values. + + Assigning per-atom ADPs would silently be a projection onto the field rather + than the write the caller asked for. Use :meth:`refit` to move the field toward + a per-atom target, or switch the model back to a per-atom representation. + """ + raise NotImplementedError( + "DisorderFieldTensor stores K nodes, not per-atom values, so per-atom " + "assignment is not representable. Use refit() to fit the field to a " + "per-atom target, or Model.set_adp_mode('isotropic') to leave field mode." + ) + + def refit(self, target_b: torch.Tensor) -> None: + """Re-fit the node values to a per-atom B target, in place. + + Replaces ``refinable_params``, so any optimizer state held for it is stale. + """ + xyz = self._xyz_fn().detach() + target_b = target_b.to(dtype=self.dtype, device=self.device) + node_values = self._fit_nodes( + target_b, xyz, self.node_positions(xyz).detach(), self.neighbor_list + ) + self.fixed_values = node_values.clone().detach() + refinable = node_values[self.refinable_mask].clone().detach() + self.refinable_params = nn.Parameter( + refinable, requires_grad=self.refinable_params.requires_grad + ) + self._build_index_cache() + self.reset_forward_cache() + + def update_refinable_mask( + self, new_mask: torch.Tensor, in_node_space: bool = False + ): + """Repartition refinable/fixed nodes, keeping the raw node values. + + Parameters + ---------- + new_mask : torch.Tensor + Boolean mask, ``(n_atoms,)`` in atom space or ``(K,)`` when + ``in_node_space``. An atom-space mask collapses with OR: a node is refinable + if any atom it serves is. + in_node_space : bool, optional + Whether ``new_mask`` is already in node space. Default False. + """ + new_mask = new_mask.to(device=self.device, dtype=torch.bool) + if in_node_space: + if new_mask.shape[0] != self.n_nodes: + raise ValueError( + f"Node-space mask must have shape ({self.n_nodes},), " + f"got {tuple(new_mask.shape)}" + ) + node_mask = new_mask + else: + if new_mask.shape[0] != self._full_shape: + raise ValueError( + f"Atom-space mask must have shape ({self._full_shape},), " + f"got {tuple(new_mask.shape)}" + ) + node_mask = self._collapse_mask( + new_mask, self.neighbor_list, self.n_nodes + ) + + current = self.fixed_values.clone() + if self.refinable_mask is not None and bool(self.refinable_mask.any()): + current[self.refinable_mask] = self.refinable_params.data + + self.fixed_values = current.clone().detach() + if bool(node_mask.any()): + self.refinable_params = nn.Parameter( + current[node_mask].clone().detach(), + requires_grad=self.refinable_params.requires_grad, + ) + else: + self.refinable_params = nn.Parameter( + torch.empty(0, 2, dtype=self.dtype, device=self.device), + requires_grad=False, + ) + self.refinable_mask = node_mask + self.fixed_mask = ~node_mask + self._build_index_cache() + self.reset_forward_cache() + + def copy(self) -> "DisorderFieldTensor": + """Independent copy sharing no parameter storage. + + The coordinate accessor is carried by REFERENCE, never deep-copied: it holds a + cache that can contain a graph-attached tensor, which ``deepcopy`` refuses to + walk. + """ + accessor = self._xyz_fn + if isinstance(accessor, ModuleReference): + accessor = accessor.module + new = DisorderFieldTensor( + initial_values=None, + xyz_fn=accessor, + k_neighbors=self._k_neighbors, + anchor_rows=(self.anchor_atom.clone(), self.anchor_node.clone()), + node_values=self.node_values().detach().clone(), + refinable_mask=self.refinable_mask.clone(), + mask_in_node_space=True, + requires_grad=self.refinable_params.requires_grad, + dtype=self.dtype, + device=self.device, + name=self._name, + epsilon=self.epsilon, + ) + new.neighbor_list = self.neighbor_list.clone() + new._full_shape = self._full_shape + return new + + def __repr__(self) -> str: + name_str = f"'{self.name}', " if self.name is not None else "" + return ( + f"DisorderFieldTensor({name_str}atoms={self._full_shape}, " + f"nodes={self.n_nodes}, k={self._k_neighbors}, dtype={self.dtype}, " + f"device={self.device}, refinable={self.get_refinable_count()})" + ) From 81d088bf15807d8e97b43e3e88f3eb52bc3f0e8c Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 11:56:28 +0200 Subject: [PATCH 076/250] Drop the sigma-weighted data-to-data scale option Sigma weighting collapses on a scale fit: down-weighting the weak shells is exactly what lets the scale run away in them. That was already found for the model-to-data fit, whose default came back to unit-weight ls for the same reason. The previous commit measured a sigma-weighted variant as slightly better on held-out reflections -- +0.0250 in free-set chi2, CI [+0.0049, +0.0457]. That measurement stands but does not support the option: it is one dataset pair, and a small held-out gain does not outweigh a failure mode found across a panel. Removed rather than left selectable, since a known-collapsing objective behind a keyword argument is a trap. `scale()` takes no objective again. What survives from that commit are the two real fixes: it fits on the work set rather than every valid reflection, and its objective is normalised. The docstring records why there is no sigma option, so the next reader reaches for it and finds the reason instead of the code. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- docs/changelog.rst | 1 - paper/probe_data_scale_objective.py | 28 ++++--- tests/unit/io/test_data_scale_fit.py | 44 ++++++----- torchref/io/datasets/collection.py | 105 ++++++++------------------- 4 files changed, 68 insertions(+), 110 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index cbd8e7d3..4851cc7e 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -8,7 +8,6 @@ Version 0.6.4 - ``CollectionScaler.refine_lbfgs_joint`` normalises its objective and registers the U penalty as its own target - Fixed ``DatasetCollection.scale`` fitting the inter-dataset scale on the free reflections as well as the work set - ``DatasetCollection.scale`` normalises its objective, so L-BFGS's absolute tolerances mean something -- Added ``DatasetCollection.scale(objective="ls_sigma")``, weighting by the propagated error on the difference being minimised - Added ``COLLECTION_XRAY_TARGETS``, the collection target taxonomy, with an intensity difference row - Removed ``CollectionRiceTarget``, which set ``beta = sigma_obs**2``; the ``ml`` row is the absolute channel instead - Renamed the kinetic ``xray_weight_rice`` / ``xray/rice`` weight to ``xray_weight_ml`` / ``xray/ml`` diff --git a/paper/probe_data_scale_objective.py b/paper/probe_data_scale_objective.py index 12e34b6d..886f2367 100644 --- a/paper/probe_data_scale_objective.py +++ b/paper/probe_data_scale_objective.py @@ -5,19 +5,14 @@ side, so there is no model error for a sigma_A or Rice likelihood to account for, and the only real choices are the weighting and which reflections the fit is allowed to see. -Two questions, both answerable by holding out the free set: +The question it answers: **the leak.** The fit used to mask with +``ReflectionData.masks()`` -- validity only, with no work/free notion -- so the free +reflections went into the scale parameters, upstream of every target and therefore +upstream of every free-set number the pipeline reports. How much did that buy it, and +does removing it move the fitted scale? -1. **The leak.** The fit used to mask with ``ReflectionData.masks()`` -- validity only, - with no work/free notion -- so the free reflections went into the scale parameters, - upstream of every target. How much did that actually buy it, and does removing it - change the fitted scale? -2. **The weighting.** ``ls`` throws sigma away. ``ls_sigma`` weights by - ``1/(sigma**2 + sigma_ref**2)``, which is the propagated error on the very difference - being minimised -- the correct weight rather than a modelling choice, precisely because - both sides are measurements. Does it generalise better? - -Both are scored on reflections the fit never saw, under two yardsticks applied identically -to every arm, so the comparison is not circular: +Scored on reflections the fit never saw, under two yardsticks applied identically to +every arm, so the comparison is not circular: R_data = sum|F - F_ref| / sum F_ref scale-free and interpretable chi2 = mean[(F - F_ref)**2 / (s**2 + s_ref**2)] is the disagreement within error? @@ -160,10 +155,13 @@ def main(): dev = torch.device(args.device) dc = build(args.dark_sf, args.light_sf, args.dmin, dev) + # A sigma-weighted arm was measured here and removed: it scored slightly better on + # held-out reflections for this one pair, but inverse-variance weighting collapses on + # a scale fit -- down-weighting the weak shells is what lets the scale run away in + # them -- and that failure mode was found across a panel. Unit weights stand. arms = [ ("legacy (validity mask, unnormalised)", lambda: legacy_scale(dc)), - ("ls (work set, normalised)", lambda: dc.scale(objective="ls")), - ("ls_sigma (work set, normalised)", lambda: dc.scale(objective="ls_sigma")), + ("ls (work set, normalised)", lambda: dc.scale()), ] rows = [] @@ -209,7 +207,7 @@ def main(): print("=" * 86) print(f"{'comparison':52s} {'set':6s} {'mean d':>10s} {'95% CI':>22s}") print("-" * 86) - pairs = [(labels[0], labels[1]), (labels[1], labels[2]), (labels[0], labels[2])] + pairs = [(labels[0], labels[1])] paired = [] for a, b in pairs: for subset in ("work", "free"): diff --git a/tests/unit/io/test_data_scale_fit.py b/tests/unit/io/test_data_scale_fit.py index 934c061e..1251c9f5 100644 --- a/tests/unit/io/test_data_scale_fit.py +++ b/tests/unit/io/test_data_scale_fit.py @@ -1,8 +1,9 @@ """``DatasetCollection.scale()`` -- the data-to-data scale fit. The only fit in the library with no model on either side: it puts one dataset onto -another, both of them measurements of the same quantity. That is why its objectives are -least squares and there is no sigma_A row -- there is no model error to account for. +another, both of them measurements of the same quantity. That is why it is least squares +and takes no objective at all -- there is no model error for a sigma_A row to account for, +and sigma weighting collapses on a scale fit. The load-bearing test here is the free-set one. That fit runs *upstream of every target*, so a leak there compromises every free-set number the pipeline later reports, and no @@ -37,8 +38,7 @@ def _fitted(dc): @pytest.mark.integration -@pytest.mark.parametrize("objective", ["ls", "ls_sigma"]) -def test_scale_never_touches_the_free_set(pair, objective): +def test_scale_never_touches_the_free_set(pair): """Corrupting the free reflections must not move the fitted parameters at all. Two *different* garbage values, because a single one could coincide with a @@ -65,7 +65,7 @@ def test_scale_never_touches_the_free_set(pair, objective): d.F_sigma[free] = filler d._corrected_fp = None # drop the cached corrected view d._corrected_cache = None - dc.scale(objective=objective) + dc.scale() results.append(_fitted(dc)) (ls_a, u_a), (ls_b, u_b) = results @@ -78,15 +78,14 @@ def test_scale_never_touches_the_free_set(pair, objective): @pytest.mark.integration -@pytest.mark.parametrize("objective", ["ls", "ls_sigma"]) -def test_identical_datasets_fit_a_unit_scale(pair, objective): +def test_identical_datasets_fit_a_unit_scale(pair): """Two copies of one dataset must scale onto each other with no correction. The sanity check the objectives have to pass before any comparison between them means anything. """ dc = pair - dc.scale(objective=objective) + dc.scale() log_scale, U = _fitted(dc) assert float(log_scale.abs().max()) < 1e-3, log_scale assert float(U.abs().max()) < 1e-3, U @@ -111,18 +110,25 @@ def test_a_known_scale_is_recovered(pair): @pytest.mark.unit -def test_unknown_objective_fails_closed(): - from torchref.io.datasets.collection import ( - DATA_SCALE_OBJECTIVES, - DatasetCollection, - ) +def test_the_objective_is_not_selectable(): + """No objective parameter, and specifically no sigma-weighted one. + + Two separate reasons, both worth keeping written down. There is no model in this fit, + so a sigma_A or Rice likelihood has no model error to account for. And + inverse-variance weighting *collapses* on a scale fit -- down-weighting the weak + shells is exactly what lets the scale run away in them -- which is why the + model-to-data fit's default came back to unit-weight ``ls`` as well. A sigma-weighted + variant was built and measured here: it scored slightly better on held-out + reflections for one dataset pair, and was still removed, because a small gain on one + pair does not outweigh a failure mode found across a panel. + """ + import inspect - dc = DatasetCollection(verbose=0, device="cpu") - with pytest.raises(ValueError, match="objective must be one of"): - dc.scale(objective="ml") - # No sigma_A / Rice row is offered, and that is a modelling statement: this fit has - # no model, so there is no model error for such a likelihood to account for. - assert DATA_SCALE_OBJECTIVES == ("ls", "ls_sigma") + from torchref.io.datasets.collection import DatasetCollection + + assert "objective" not in inspect.signature(DatasetCollection.scale).parameters + src = inspect.getsource(DatasetCollection.scale) + assert "sigma" in src, "the reason sigma weighting is absent must stay documented" @pytest.mark.integration diff --git a/torchref/io/datasets/collection.py b/torchref/io/datasets/collection.py index 6fef9be5..9a2daaa0 100644 --- a/torchref/io/datasets/collection.py +++ b/torchref/io/datasets/collection.py @@ -14,16 +14,6 @@ from .base import CrystalDataset from .reflection_data import ReflectionData -#: Objectives for :meth:`DatasetCollection.scale`, the **data-to-data** fit. Least -#: squares only, and not for want of alternatives: there is no model in that fit, so -#: there is no model error for a sigma_A or Rice likelihood to account for. ``ls_sigma`` -#: weights by the propagated error on the difference, which is the correct weight -#: precisely because both sides are measurements. -DATA_SCALE_OBJECTIVES = ("ls", "ls_sigma") - -#: Default for :meth:`DatasetCollection.scale`. Unit-weight least squares, matching -#: :data:`~torchref.scaling.scaler_base.DEFAULT_SCALE_TARGET` for the model-to-data fit. -DEFAULT_DATA_SCALE_OBJECTIVE = "ls" @dataclass @@ -273,56 +263,38 @@ def __call__(self, mask: bool = True) -> Dict[str, Tuple]: """ return {name: ds(mask=mask, scale=True) for name, ds in self} - def scale(self, objective: str = DEFAULT_DATA_SCALE_OBJECTIVE): + def scale(self): """ - Fit every non-reference dataset's scale and anisotropy onto the reference, - whose own parameters are left untouched. - - **This is the data-to-data fit**, and it is the only one in the library: there - is no model here, so there is no model error to account for and nothing for a - sigma_A or Rice likelihood to do. Both sides are measurements of the same - quantity, which is why the objectives are least squares -- - :data:`DATA_SCALE_OBJECTIVES`: - - ``ls`` - ``sum (F - F_ref)**2``, unit weights. The default, matching - :data:`~torchref.scaling.scaler_base.DEFAULT_SCALE_TARGET` for the - model-to-data fit. - ``ls_sigma`` - ``sum (F - F_ref)**2 / (sigma**2 + sigma_ref**2)``. Both sides being - measured is exactly the condition under which inverse-variance weighting is - the correct weight rather than a modelling choice: the denominator is the - propagated error on the difference being minimised, with no model-error term - in it. - - The ``ls_sigma`` weights are computed **once, detached**, from the starting - sigmas. They must not be re-derived inside the closure: ``sigma`` carries the - same ``log_scale`` as ``F``, so a live denominator rewards inflating the scale to - inflate the variance, and without the ``+log(sigma)`` term of a full Gaussian - there is nothing to oppose it. Fixed weights are what "weighted least squares" - means; see :mod:`torchref.scaling.scaler_base` on the related hazard of fitting a - scale against a likelihood that carries the scale in its variance. + Unit-weight least-squares fit of every non-reference dataset's scale and + anisotropy onto the reference, whose own parameters are left untouched. + + **This is the data-to-data fit**, the only one in the library with no model on + either side. So there is no model error to account for and nothing for a sigma_A + or Rice likelihood to do -- which is why the objective is least squares and there + is no way to select another. + + **Sigma weighting is deliberately not offered**, and that is a measured decision + rather than an omission. ``sum (F - F_ref)**2 / (sigma**2 + sigma_ref**2)`` is + superficially the principled choice -- both sides are measurements, so the + denominator is the honest propagated error on the difference being minimised -- + and on a single dataset pair it does score slightly better on held-out + reflections. It is still wrong to use: inverse-variance weighting on a scale fit + collapses, because down-weighting the weak shells is exactly what lets the scale + run away in them, and the same objective was tried and rejected for the + model-to-data fit (whose default likewise came back to unit-weight ``ls``). A + small held-out gain on one pair does not outweigh a failure mode found across a + panel. Do not re-add it. Fitted on the **work set** of both datasets. L-BFGS with strong-Wolfe line search, 10 outer steps of ``max_iter=100``, on an objective normalised to O(1) because those tolerances are absolute. Members' ``log_scale``/``U_aniso`` are mutated, and ``requires_grad`` is turned on and back off around the fit. - Parameters - ---------- - objective : str, optional - One of :data:`DATA_SCALE_OBJECTIVES`. - Raises ------ ValueError - If no reference dataset is set, there is nothing else to scale, or - ``objective`` is not recognised. + If no reference dataset is set, or there is nothing else to scale. """ - if objective not in DATA_SCALE_OBJECTIVES: - raise ValueError( - f"objective must be one of {DATA_SCALE_OBJECTIVES}, got {objective!r}" - ) if self._reference_dataset is None: raise ValueError("No reference dataset set for scaling") @@ -346,28 +318,14 @@ def scale(self, objective: str = DEFAULT_DATA_SCALE_OBJECTIVE): ref_mask = ref_ds.work.mask combined = [ds.work.mask & ref_mask for ds in to_scale] - # Weights and the normaliser: once, detached, outside the closure. + # The normaliser: once, detached, outside the closure. L-BFGS converges on + # ABSOLUTE tolerances, so an objective carrying the data's own magnitude leaves + # `tolerance_grad`/`tolerance_change` meaningless -- the same hazard + # `ScalerBase.refine_lbfgs` documents at length. This fit had no normaliser. with torch.no_grad(): - ref_F0, ref_sig0 = ref_ds.get_corrected_data() - weights = None - if objective == "ls_sigma": - # Local import: `torchref.base.targets` is not otherwise reachable from - # `torchref.io`, and hoisting it would couple the two packages. - from torchref.base.targets.xray_likelihoods import floor_sigma_obs - - weights = [] - for ds, cm in zip(to_scale, combined): - _, sig0 = ds.get_corrected_data() - var = ( - floor_sigma_obs(sig0[cm]) ** 2 - + floor_sigma_obs(ref_sig0[cm]) ** 2 - ) - weights.append(1.0 / var) - n_fitted = sum(int(cm.sum()) for cm in combined) - norm = 1.0 / max(n_fitted, 1) - else: - ssq = sum(float(ref_F0[cm].pow(2).sum()) for cm in combined) - norm = 1.0 / max(ssq, 1e-30) + ref_F0, _ = ref_ds.get_corrected_data() + ssq = sum(float(ref_F0[cm].pow(2).sum()) for cm in combined) + norm = 1.0 / max(ssq, 1e-30) def closure(): optimizer.zero_grad() @@ -375,12 +333,9 @@ def closure(): # get_corrected_data, not __call__: MaskedTensor has no autograd. ref_F_scaled, _ = ref_ds.get_corrected_data() - for i, (ds, cm) in enumerate(zip(to_scale, combined)): + for ds, cm in zip(to_scale, combined): F_scaled, _ = ds.get_corrected_data() - resid_sq = (F_scaled[cm] - ref_F_scaled[cm]) ** 2 - if weights is not None: - resid_sq = resid_sq * weights[i] - loss = loss + torch.sum(resid_sq) + loss = loss + torch.sum((F_scaled[cm] - ref_F_scaled[cm]) ** 2) loss = loss * norm loss.backward() return loss From a95eaa5e26ba8c9b2b00516569af1d59900f30d4 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 12:25:18 +0200 Subject: [PATCH 077/250] Add a node-field ADP representation to set_adp_mode set_adp_mode("field") joins "isotropic" and "anisotropic" on the existing switch, replacing the per-atom isotropic B with a DisorderFieldTensor whose node values are least-squares fitted to the B it replaces. Entering runs the existing partition first, so everything keyed off the iso/aniso split is refreshed by tested code; leaving needs no special case, because the partition reads adp(), which a field evaluates per atom and so materialises back into a per-atom wrapper on the way out. Nodes are anchored on density clusters rather than single atoms. A node placed exactly on an atom can isolate that atom by narrowing its kernel, which is per-atom refinement wearing a node's clothes; measured, cluster anchoring lowers the worst per-atom B on every structure tried. Node positions carry a refinable offset from their anchor centroid, on by default. With positions fixed a node's only way to localise is to narrow, so this is what lets a restraint move a node toward atoms instead of only widening it. Model.copy re-points the borrowed coordinate accessor at the copy, or the two models share coordinates and the copy is not independent. create_from_state_dict rebuilds a field when the saved adp storage is 2-D, reusing the saved anchor rows and inferring from the storage width whether positions were refinable. node_load() scatters the candidate weights back into node space. Summing weights() over atoms does not do this -- it returns (n_atoms, k) over candidates, so the sum is a length-k vector of per-slot totals with no meaning. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- tests/unit/model/test_adp_field_mode.py | 214 ++++++++++++++++++++++++ tests/unit/model/test_disorder_field.py | 20 +++ torchref/model/disorder_field.py | 96 ++++++++++- torchref/model/model.py | 154 ++++++++++++++++- 4 files changed, 468 insertions(+), 16 deletions(-) create mode 100644 tests/unit/model/test_adp_field_mode.py diff --git a/tests/unit/model/test_adp_field_mode.py b/tests/unit/model/test_adp_field_mode.py new file mode 100644 index 00000000..32f6c772 --- /dev/null +++ b/tests/unit/model/test_adp_field_mode.py @@ -0,0 +1,214 @@ +"""``set_adp_mode("field")``: a third ADP representation on the existing switch. + +The switch is a *conversion*, not a freeze, so entering field mode has to fit the node +values to the B it replaces and leaving has to materialise them back per atom. The +inert-until-selected property matters too: adding the mode must not perturb a model +that never asks for it. +""" + +import math + +import pytest +import torch + +from torchref.model.disorder_field import DisorderFieldTensor +from torchref.model.model import Model +from torchref.model.parameter_wrappers import PositiveMixedTensor + + +@pytest.fixture(scope="module") +def pdb_path(pdb_dir): + return str(pdb_dir / "3GR5.pdb") + + +def _model(pdb_path): + model = Model(verbose=0) + model.load_pdb(pdb_path) + return model + + +@pytest.mark.unit +def test_isotropic_mode_is_untouched(pdb_path): + """The default representation is unchanged by the new branch existing.""" + model = _model(pdb_path) + model.set_adp_mode("isotropic") + assert isinstance(model.adp, PositiveMixedTensor) + assert not model.adp_is_field + assert model.adp().shape == (len(model.pdb),) + + +@pytest.mark.unit +def test_unknown_mode_still_raises(pdb_path): + model = _model(pdb_path) + with pytest.raises(ValueError, match="field"): + model.set_adp_mode("nonsense") + + +@pytest.mark.unit +def test_entering_field_mode_replaces_the_wrapper(pdb_path): + """``adp`` becomes a field whose parameter count is set by nodes, not atoms.""" + model = _model(pdb_path) + n_atoms = len(model.pdb) + model.set_adp_mode("field", n_nodes=40, k_neighbors=8) + + assert model.adp_is_field + assert isinstance(model.adp, DisorderFieldTensor) + assert model.adp.n_nodes == 40 + assert model.adp().shape == (n_atoms,) + # The whole point: far fewer refinable numbers than atoms. + assert int(model.adp.refinable_params.numel()) < n_atoms + + +@pytest.mark.unit +def test_field_tracks_the_b_it_replaced(pdb_path): + """The fit is against the deposited B, so it must resemble it, not restart from flat.""" + model = _model(pdb_path) + before = model.adp().detach().clone() + model.set_adp_mode("field", n_nodes=64, k_neighbors=8) + after = model.adp().detach() + + spread = (before - before.mean()).pow(2).mean().sqrt() + resid = (after - before).pow(2).mean().sqrt() + assert resid < 0.6 * spread, f"rmse {resid:.2f} vs B spread {spread:.2f}" + assert bool((after > 0).all()) + + +@pytest.mark.unit +def test_more_nodes_track_more_closely(pdb_path): + """Node count is the accuracy dial, end to end through the switch.""" + errors = [] + for n in (4, 32, 200): + model = _model(pdb_path) + before = model.adp().detach().clone() + model.set_adp_mode("field", n_nodes=n, k_neighbors=8) + errors.append(float((model.adp().detach() - before).pow(2).mean().sqrt())) + assert errors[0] > errors[1] > errors[2], errors + + +@pytest.mark.unit +def test_leaving_field_mode_materialises_per_atom(pdb_path): + """Round trip out of field mode gives back a per-atom wrapper holding its values.""" + model = _model(pdb_path) + model.set_adp_mode("field", n_nodes=64, k_neighbors=8) + field_values = model.adp().detach().clone() + + model.set_adp_mode("isotropic") + + assert not model.adp_is_field + assert isinstance(model.adp, PositiveMixedTensor) + assert torch.allclose(model.adp().detach(), field_values, atol=1e-5) + + +@pytest.mark.unit +def test_field_mode_collapses_anisotropic_atoms_first(pdb_path): + """Entering from anisotropic goes through B_eq rather than dropping the U.""" + model = _model(pdb_path) + model.set_adp_mode("anisotropic") + assert bool(model.aniso_flag.any()) + + with torch.no_grad(): + U = model.u().detach() + beq = (8.0 * math.pi**2 / 3.0) * (U[:, 0] + U[:, 1] + U[:, 2]) + + model.set_adp_mode("field", n_nodes=200, k_neighbors=8) + + assert not bool(model.aniso_flag.any()), "field mode is isotropic in this stage" + got = model.adp().detach() + finite = torch.isfinite(beq) + spread = (beq[finite] - beq[finite].mean()).pow(2).mean().sqrt() + resid = (got[finite] - beq[finite]).pow(2).mean().sqrt() + assert resid < spread, "field ignored the equivalent isotropic B" + + +@pytest.mark.unit +def test_sf_indices_and_flags_stay_consistent(pdb_path): + """Everything keyed off the iso/aniso split is refreshed, not left stale.""" + model = _model(pdb_path) + model.set_adp_mode("field", n_nodes=32, k_neighbors=8) + + assert not bool(model.aniso_flag.any()) + assert not bool(model.pdb["anisou_flag"].to_numpy().any()) + assert int(model._iso_indices.numel()) == len(model.pdb) + assert bool(model._aniso_is_empty) + + +@pytest.mark.unit +def test_get_iso_and_adp_u6_work_unmodified(pdb_path): + """The two consumers that read ``adp()`` need no knowledge of the field.""" + model = _model(pdb_path) + model.set_adp_mode("field", n_nodes=32, k_neighbors=8) + + xyz, adp, occ = model.get_iso() + assert xyz.shape[0] == adp.shape[0] == occ.shape[0] == len(model.pdb) + + u6 = model.adp_u6() + assert u6.shape == (len(model.pdb), 6) + expected = (adp / (8.0 * math.pi**2)).detach() + assert torch.allclose(u6[:, 0].detach(), expected, atol=1e-6) + assert torch.allclose(u6[:, 3:].detach(), torch.zeros_like(u6[:, 3:]), atol=1e-12) + + +@pytest.mark.unit +def test_gradient_flows_to_the_nodes_through_adp_u6(pdb_path): + """The path an ADP restraint would take reaches the node parameters.""" + model = _model(pdb_path) + model.set_adp_mode("field", n_nodes=32, k_neighbors=8) + model.adp_u6().sum().backward() + grad = model.adp.refinable_params.grad + assert grad is not None and bool(grad.abs().sum() > 0) + + +@pytest.mark.unit +def test_refine_adp_sees_the_node_parameters(pdb_path): + """``parameters_of_types`` needs no change: the field is still the ``adp`` type.""" + model = _model(pdb_path) + model.set_adp_mode("field", n_nodes=32, k_neighbors=8) + params = model.parameters_of_types(["adp"]) + assert len(params) == 1 + assert params[0] is model.adp.refinable_params + # [log B, log sigma, dx, dy, dz] -- positions are refinable by default. + assert params[0].shape == (32, 5) + assert model.adp.refines_positions + + +@pytest.mark.unit +def test_model_copy_repoints_the_accessor(pdb_path): + """A copied field must read the COPY's coordinates, not the original's. + + ``copy()`` carries the borrowed accessor by reference, so without the re-point in + ``Model.copy`` the two models share coordinates and moving one changes the other's + ADPs. + """ + model = _model(pdb_path) + model.set_adp_mode("field", n_nodes=32, k_neighbors=8) + clone = model.copy() + + assert clone.adp._xyz_fn.module is clone.xyz + assert clone.adp._xyz_fn.module is not model.xyz + before = model.adp().detach().clone() + + with torch.no_grad(): + clone.xyz.refinable_params[: len(model.pdb) // 2, 0] += 5.0 + model.adp.reset_forward_cache() + + assert torch.allclose(model.adp().detach(), before, atol=1e-9), ( + "moving the copy changed the original's ADPs -- accessor was not re-pointed" + ) + + +@pytest.mark.unit +def test_state_dict_round_trip_in_field_mode(pdb_path): + """``create_from_state_dict`` rebuilds a field when the saved ``adp`` was one.""" + model = _model(pdb_path) + model.set_adp_mode("field", n_nodes=24, k_neighbors=6) + expected = model.adp().detach().clone() + + sd = { + key: (value.clone() if torch.is_tensor(value) else value) + for key, value in model.state_dict().items() + } + restored = Model.create_from_state_dict(sd, verbose=0) + + assert restored.adp_is_field + assert restored.adp.n_nodes == 24 + assert torch.allclose(restored.adp().detach(), expected, atol=1e-6) diff --git a/tests/unit/model/test_disorder_field.py b/tests/unit/model/test_disorder_field.py index ebf97f38..1ee37903 100644 --- a/tests/unit/model/test_disorder_field.py +++ b/tests/unit/model/test_disorder_field.py @@ -234,6 +234,26 @@ def test_list_adequacy_invariant_is_exposed(coords, target_b): assert 0.0 <= loose.smallest_candidate_weight() <= 1.0 +@pytest.mark.unit +def test_node_load_is_in_node_space_and_conserves_total_weight(coords, target_b): + """Load must be scattered into node space, not summed over the candidate axis. + + ``weights()`` is ``(n_atoms, k)`` over CANDIDATES, so ``weights().sum(0)`` is a + length-k vector of per-slot totals with no meaning -- a trap worth pinning, because + it silently returns a plausible-looking tensor of the wrong length. + """ + field, _ = _field(coords, target_b, n_nodes=12, k_neighbors=6) + load = field.node_load() + + assert load.shape == (12,), "load must be per node, not per candidate slot" + assert field.weights().sum(dim=0).shape == (6,), "the trap this method avoids" + # Rows of W sum to 1, so the total load is exactly the atom count. + assert torch.allclose( + load.sum(), torch.tensor(float(coords.shape[0]), dtype=load.dtype) + ) + assert bool((load >= 0).all()) + + @pytest.mark.unit def test_rebuild_neighbor_list_is_explicit_and_refreshes(coords, target_b): """Membership only changes when the caller asks; the rebuild then takes effect.""" diff --git a/torchref/model/disorder_field.py b/torchref/model/disorder_field.py index c2eec6c5..35bbc75e 100644 --- a/torchref/model/disorder_field.py +++ b/torchref/model/disorder_field.py @@ -87,6 +87,37 @@ def farthest_point_anchors(xyz: torch.Tensor, n_nodes: int) -> torch.Tensor: return torch.sort(anchors).values +def density_anchor_rows(xyz: torch.Tensor, n_nodes: int): + """Anchor each node on its whole density cluster rather than on one atom. + + :func:`farthest_point_anchors` snaps every anchor onto an atom, which puts each node + exactly *on* an atom -- so narrowing its kernel isolates the atom it is already + standing on, and the node ends up owning a single atom outright. Anchoring on the + cluster instead places the node at the cluster centroid, generally between atoms, so + there is no atom for it to fall onto. + + Returns + ------- + tuple of torch.Tensor + ``(atom index per entry, node index per entry)``, flat and ragged, suitable for + ``DisorderFieldTensor(anchor_rows=...)``. Every node keeps at least its seed + atom, so no node is left with an empty neighbourhood. + """ + seeds = farthest_point_anchors(xyz, n_nodes) + assign = torch.cdist(xyz, xyz[seeds]).argmin(dim=1) + atom_idx = torch.arange(xyz.shape[0], dtype=torch.int64, device=xyz.device) + + # A seed whose cluster somehow came out empty still needs a position. + present = torch.bincount(assign, minlength=seeds.shape[0]) > 0 + if not bool(present.all()): + missing = (~present).nonzero(as_tuple=True)[0] + atom_idx = torch.cat([atom_idx, seeds[missing]]) + assign = torch.cat([assign, missing]) + + order = torch.argsort(assign) + return atom_idx[order], assign[order] + + def build_neighbor_list( xyz: torch.Tensor, node_pos: torch.Tensor, k: int ) -> torch.Tensor: @@ -189,6 +220,7 @@ def __init__( xyz_fn: Optional[Callable[[], torch.Tensor]] = None, n_nodes: int = 32, k_neighbors: int = 12, + refine_positions: bool = False, anchor_rows: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, node_values: Optional[torch.Tensor] = None, refinable_mask: Optional[torch.Tensor] = None, @@ -201,6 +233,7 @@ def __init__( ): self.epsilon = epsilon self._k_neighbors = int(k_neighbors) + self._refine_positions = bool(refine_positions) object.__setattr__(self, "_xyz_fn", _wrap_accessor(xyz_fn)) if initial_values is None and node_values is None: @@ -349,13 +382,30 @@ def _fit_nodes( W = torch.zeros(xyz.shape[0], n_k, dtype=xyz.dtype, device=xyz.device) W.scatter_(1, neighbor_list, W_sparse) + # Non-finite targets are dropped from the solve rather than carried into it: + # a single NaN row would propagate through the normal equations and take + # every node with it. Those atoms still receive a fitted value on output. + finite = torch.isfinite(target_b) + if not bool(finite.all()): + W, target_b = W[finite], target_b[finite] + if target_b.numel() == 0: + flat = torch.stack([torch.zeros_like(log_sigma), log_sigma], dim=1) + if self._refine_positions: + flat = torch.cat([flat, torch.zeros_like(node_pos)], dim=1) + return flat + gram = W.T @ W ridge = 1e-6 * torch.diagonal(gram).mean().clamp(min=1e-30) eye = torch.eye(n_k, dtype=W.dtype, device=W.device) b = torch.linalg.solve(gram + ridge * eye, W.T @ target_b) b = b.clamp(min=self.epsilon) - return torch.stack([torch.log(b), log_sigma], dim=1) + node_values = torch.stack([torch.log(b), log_sigma], dim=1) + if self._refine_positions: + node_values = torch.cat( + [node_values, torch.zeros_like(node_pos)], dim=1 + ) + return node_values @staticmethod def _collapse_mask( @@ -394,11 +444,28 @@ def set_xyz_fn(self, xyz_fn: Callable[[], torch.Tensor]) -> None: """ object.__setattr__(self, "_xyz_fn", _wrap_accessor(xyz_fn)) + @property + def refines_positions(self) -> bool: + """Whether node positions carry a refinable offset.""" + return self._refine_positions + def node_positions(self, xyz: Optional[torch.Tensor] = None) -> torch.Tensor: - """Node positions, ``(K, 3)``, derived from the anchor neighbourhoods.""" + """Node positions, ``(K, 3)``: anchored centroid plus any refinable offset.""" if xyz is None: xyz = self._xyz_fn() - return self._segment_mean(xyz, self.anchor_atom, self.anchor_node, self.n_nodes) + return self._node_positions_from(xyz, super().forward()) + + def _node_positions_from(self, xyz, raw): + """Anchor centroid, displaced by the refinable offset when there is one. + + Keeping the centroid as the base rather than storing absolute coordinates means + a node still travels with the model under xyz refinement; the offset only says + where it sits *relative* to the atoms it belongs to. + """ + base = self._segment_mean( + xyz, self.anchor_atom, self.anchor_node, raw.shape[0] + ) + return base + raw[:, 2:5] if self._refine_positions else base def weights(self, xyz: Optional[torch.Tensor] = None) -> torch.Tensor: """Per-atom weights over candidate nodes, ``(n_atoms, k)``. Rows sum to one.""" @@ -406,9 +473,25 @@ def weights(self, xyz: Optional[torch.Tensor] = None) -> torch.Tensor: xyz = self._xyz_fn() raw = super().forward() return self._weights( - xyz, self.node_positions(xyz), self.neighbor_list, raw[:, 1] + xyz, self._node_positions_from(xyz, raw), self.neighbor_list, raw[:, 1] ) + def node_load(self, xyz: Optional[torch.Tensor] = None) -> torch.Tensor: + """Total weight each node carries across all atoms, ``(K,)``. + + Not obtainable by summing :meth:`weights` over atoms: that returns + ``(n_atoms, k)`` over each atom's CANDIDATES, so a column is a candidate *slot* + shared by different nodes for different atoms, and summing it gives a length-k + vector with no meaning. The weights have to be scattered back into node space + through ``neighbor_list`` first, which is what this does. + + Load is what identifies a node that has stopped doing useful work: a node that + narrows onto a handful of atoms carries almost none. + """ + W = self.weights(xyz) + load = W.new_zeros(self.n_nodes) + return load.index_add(0, self.neighbor_list.reshape(-1), W.reshape(-1)) + def smallest_candidate_weight(self, xyz: Optional[torch.Tensor] = None) -> float: """Largest per-atom minimum candidate weight --- the list-adequacy invariant. @@ -455,9 +538,7 @@ def evaluate(self, xyz: torch.Tensor, raw: torch.Tensor) -> torch.Tensor: ``(K, 2)`` node storage, ``[log b, log sigma]``. """ log_b, log_sigma = raw[:, 0], raw[:, 1] - node_pos = self._segment_mean( - xyz, self.anchor_atom, self.anchor_node, raw.shape[0] - ) + node_pos = self._node_positions_from(xyz, raw) W = self._weights(xyz, node_pos, self.neighbor_list, log_sigma) return (W * torch.exp(log_b)[self.neighbor_list]).sum(dim=1) @@ -586,6 +667,7 @@ def copy(self) -> "DisorderFieldTensor": initial_values=None, xyz_fn=accessor, k_neighbors=self._k_neighbors, + refine_positions=self._refine_positions, anchor_rows=(self.anchor_atom.clone(), self.anchor_node.clone()), node_values=self.node_values().detach().clone(), refinable_mask=self.refinable_mask.clone(), diff --git a/torchref/model/model.py b/torchref/model/model.py index 33b3cd28..85056afa 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -1036,6 +1036,13 @@ def copy(self): if module is not None and hasattr(module, "copy"): setattr(model_copy, module_name, module.copy()) + # A wrapper that borrows the coordinates carries that reference through its + # own ``copy``, so it still points at THIS model's ``xyz``. Re-point it, or + # the two models silently share coordinates and the copy is not independent. + for module in model_copy._modules.values(): + if module is not None and hasattr(module, "set_xyz_fn"): + module.set_xyz_fn(model_copy.xyz) + if self.ctx.verbose > 0: print(f"✓ Model copied successfully ({len(model_copy.pdb)} atoms)") @@ -1215,7 +1222,14 @@ def unfreeze(self, target: str): self.occupancy_mask, in_compressed_space=False ) - def set_adp_mode(self, mode: str = "isotropic", aniso_selection: str = None): + def set_adp_mode( + self, + mode: str = "isotropic", + aniso_selection: str = None, + n_nodes: int = None, + k_neighbors: int = 12, + refine_node_positions: bool = True, + ): """Set the atomic displacement parameter (ADP) parametrization. Repartitions atoms between isotropic (a single B in ``adp``) and @@ -1230,22 +1244,50 @@ def set_adp_mode(self, mode: str = "isotropic", aniso_selection: str = None): Parameters ---------- - mode : {"isotropic", "anisotropic"}, optional + mode : {"isotropic", "anisotropic", "field"}, optional ``"isotropic"`` (default) converts every atom, previously anisotropic ones to ``B_eq = (8 pi^2 / 3)(U11 + U22 + U33)``. ``"anisotropic"`` converts those matching ``aniso_selection``, expanding isotropic atoms - to ``U = (B / 8 pi^2) I``. + to ``U = (B / 8 pi^2) I``. ``"field"`` replaces the per-atom isotropic B + with a :class:`~torchref.model.disorder_field.DisorderFieldTensor`, whose + node values are least-squares fitted to the B it replaces, so the atom + count stops setting the ADP parameter count. aniso_selection : str, optional Phenix-style selection for ``mode="anisotropic"``, default ``"not resname HOH and not element H"``; ignored otherwise. + n_nodes : int, optional + Nodes for ``mode="field"``. Defaults to one per 25 atoms, floored at 4. + k_neighbors : int, optional + Candidate nodes per atom for ``mode="field"``. Default 12. + refine_node_positions : bool, optional + Give each node a refinable offset from its anchor centroid, at three extra + parameters per node. On by default: it is what lets the load-balancing + restraint move a node toward atoms instead of only widening its kernel. Notes ----- Run once at model setup, before scaling / restraints / targets. The isotropic result matches a freshly-loaded isotropic-only model. + + Leaving ``"field"`` needs no special case: the conversion reads ``adp()``, + which a field evaluates per atom, so the field materialises into a per-atom + wrapper on the way out. """ if not self.ctx.initialized or self.pdb is None: return + if mode == "field": + # Collapse to a clean per-atom isotropic state first. The field is + # isotropic in this representation, and the partition owns every buffer + # keyed off the iso/aniso split, so reuse it rather than duplicating it. + self._apply_adp_partition( + torch.zeros(len(self.pdb), dtype=torch.bool, device=self.device) + ) + self._install_disorder_field( + n_nodes=n_nodes, + k_neighbors=k_neighbors, + refine_node_positions=refine_node_positions, + ) + return if mode == "isotropic": aniso_mask = torch.zeros( len(self.pdb), dtype=torch.bool, device=self.device @@ -1259,10 +1301,68 @@ def set_adp_mode(self, mode: str = "isotropic", aniso_selection: str = None): ).to(self.device) else: raise ValueError( - f"Unknown ADP mode: {mode!r}. Use 'isotropic' or 'anisotropic'." + f"Unknown ADP mode: {mode!r}. Use 'isotropic', 'anisotropic' " + "or 'field'." ) self._apply_adp_partition(aniso_mask) + @property + def adp_is_field(self) -> bool: + """Whether ``adp`` is a node field rather than a per-atom wrapper.""" + from torchref.model.disorder_field import DisorderFieldTensor + + return isinstance(self.adp, DisorderFieldTensor) + + def _install_disorder_field( + self, + n_nodes: int = None, + k_neighbors: int = 12, + refine_node_positions: bool = False, + ): + """Replace the per-atom ``adp`` wrapper with a node field fitted to its B. + + Expects the model to already be in a per-atom isotropic state, which + :meth:`set_adp_mode` arranges by running the partition first. + """ + from torchref.model.disorder_field import ( + DisorderFieldTensor, + density_anchor_rows, + ) + + with torch.no_grad(): + B = self.adp().detach().clone() + xyz = self.xyz().detach() + if n_nodes is None: + n_nodes = max(4, int(round(len(self.pdb) / 25.0))) + + # Anchor on density clusters, not single atoms: a node placed exactly on an atom + # can isolate that atom by narrowing its kernel, which is per-atom refinement + # wearing a node's clothes. + anchor_rows = density_anchor_rows(xyz, min(n_nodes, len(self.pdb))) + + self.adp = DisorderFieldTensor( + initial_values=B.to(self.dtype_float), + xyz_fn=self.xyz, + n_nodes=n_nodes, + refine_positions=refine_node_positions, + anchor_rows=anchor_rows, + k_neighbors=k_neighbors, + name="adp", + dtype=self.dtype_float, + device=self.device, + ) + # ``adp_mask`` is in atom space; the field collapses it onto its nodes. + self.adp.update_refinable_mask(self.adp_mask) + + if self.ctx.verbose > 0: + print( + f"ADP field: {self.adp.n_nodes} nodes, k={k_neighbors}, " + f"{int(self.adp.get_refinable_count())} refinable " + f"(was {len(self.pdb)} per-atom B)" + ) + if hasattr(self, "reset_cache"): + self.reset_cache() + def _apply_adp_partition(self, aniso_mask: torch.Tensor): """Convert ADP storage to match a target anisotropic-atom mask. @@ -1982,11 +2082,47 @@ def create_from_state_dict( refinable_mask=xyz_mask, name="xyz", ) - instance.adp = PositiveMixedTensor( - torch.tensor(pdb["tempfactor"].values, dtype=saved_dtype), - refinable_mask=adp_mask, - name="adp", - ) + # A saved node field has 2-D ``adp`` storage (K, 2) where a per-atom + # wrapper has 1-D, so the shape says which representation to rebuild. + # Built only for its shapes and masks; load_state_dict overwrites values. + saved_adp = state_dict.get("adp.fixed_values") + if saved_adp is not None and saved_adp.ndim == 2: + from torchref.model.disorder_field import DisorderFieldTensor + + saved_nl = state_dict.get("adp.neighbor_list") + # Rebuild with the SAVED anchor rows: cluster anchoring makes these + # length n_atoms where single-atom anchoring makes them length K, so + # reconstructing them from scratch would shape-mismatch on load. + saved_anchor_atom = state_dict.get("adp.anchor_atom") + saved_anchor_node = state_dict.get("adp.anchor_node") + anchor_rows = ( + (saved_anchor_atom, saved_anchor_node) + if saved_anchor_atom is not None + else None + ) + instance.adp = DisorderFieldTensor( + initial_values=torch.tensor( + pdb["tempfactor"].values, dtype=saved_dtype + ), + xyz_fn=instance.xyz, + n_nodes=int(saved_adp.shape[0]), + k_neighbors=( + int(saved_nl.shape[1]) if saved_nl is not None else 12 + ), + # Storage width says whether positions carry a refinable offset. + refine_positions=bool(saved_adp.shape[1] == 5), + anchor_rows=anchor_rows, + refinable_mask=adp_mask, + mask_in_node_space=True, + name="adp", + dtype=saved_dtype, + ) + else: + instance.adp = PositiveMixedTensor( + torch.tensor(pdb["tempfactor"].values, dtype=saved_dtype), + refinable_mask=adp_mask, + name="adp", + ) # Match load(): the anisotropic U is a CholeskyMixedTensor so the # restored model refines it in the same positive-definite-by- # construction parametrization as a freshly-loaded one. From 25dae997b375f4370a938b71e910992d32126f3f Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 12:25:32 +0200 Subject: [PATCH 078/250] Add node load balancing and a node magnitude prior for the ADP field Two restraints for the node-field ADP representation, closing different halves of the same degeneracy. Both are inert unless the model is in field mode. NodeLoadTarget (adp/node_load, weight 10.0) bars a node from being abandoned. A node can otherwise narrow its kernel until it holds a single atom and then take whatever value fits it; measured, such a node ends with a load near or below one atom against a healthy median of seven, and drives that atom's B into the thousands. The penalty is one-sided, softplus(-log(load / mean load)): an over-loaded node is not charged. Maximising load entropy would have been the symmetric choice and is wrong here, because it is optimal at uniform load and fitted fields legitimately span two orders of magnitude in kernel width. It acts on the weights alone, with no gradient to the node values, so it removes the opportunity to place an extreme B rather than pricing it. That is also its limit, and the reason for the second term. Blocking the collapse does not stop a value running away, it only stops it being confined: with the barrier alone the worst per-atom B rose at moderate node counts, spreading over a neighbourhood instead of landing on one atom. NodeSmoothnessTarget (adp/node_smoothness) prices that, as a distance-weighted mean squared difference of log B between node pairs. Level-invariant, so it cannot fight the scaler over the overall B level; scale-free in log space; and it charges only departures from the local level, leaving a genuine B gradient across a structure free. Registered at weight 0 pending its own screen, as geometry/ramachandran already is. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- tests/unit/refinement/test_loss_weighting.py | 6 + .../unit/refinement/test_node_load_target.py | 143 ++++++++++++++++++ torchref/refinement/base_refinement.py | 10 ++ torchref/refinement/targets/adp/__init__.py | 4 + torchref/refinement/targets/adp/node_load.py | 115 ++++++++++++++ .../refinement/targets/adp/node_smoothness.py | 141 +++++++++++++++++ torchref/refinement/targets/combined.py | 8 + 7 files changed, 427 insertions(+) create mode 100644 tests/unit/refinement/test_node_load_target.py create mode 100644 torchref/refinement/targets/adp/node_load.py create mode 100644 torchref/refinement/targets/adp/node_smoothness.py diff --git a/tests/unit/refinement/test_loss_weighting.py b/tests/unit/refinement/test_loss_weighting.py index 1f392822..08a7fe0c 100644 --- a/tests/unit/refinement/test_loss_weighting.py +++ b/tests/unit/refinement/test_loss_weighting.py @@ -68,4 +68,10 @@ def test_default_group_weights_values(self): # Sub-weight on the SIGD prior; 1.0 leaves it at the adp group weight # pending the R_free scan. Weights multiply down the path. "adp/sigd": 1.0, + # Node load balancing, inert unless the ADPs are a node field. Above the + # group weight because it bars a degenerate direction rather than + # competing with the data. + "adp/node_load": 10.0, + # Magnitude prior on node values, off pending measurement. + "adp/node_smoothness": 0.0, } diff --git a/tests/unit/refinement/test_node_load_target.py b/tests/unit/refinement/test_node_load_target.py new file mode 100644 index 00000000..fcb36a0f --- /dev/null +++ b/tests/unit/refinement/test_node_load_target.py @@ -0,0 +1,143 @@ +"""The node load-balancing barrier. + +Two properties carry the design. It must be **one-sided** -- an abandoned node is +penalised, an over-loaded one is not -- because the symmetric choice (maximising load +entropy) is optimal at uniform load and would flatten the multiscale kernel-width spread +that a working field genuinely has. And it must act through the *weights* only, so its +gradient reaches node positions and widths but never the node values: it removes the +opportunity to place an extreme B rather than penalising the B. +""" + +import pytest +import torch + +from torchref.model.model import Model +from torchref.refinement.targets.adp import NodeLoadTarget + + +@pytest.fixture(scope="module") +def pdb_path(pdb_dir): + return str(pdb_dir / "3GR5.pdb") + + +def _field_model(pdb_path, n_nodes=32, **kw): + model = Model(verbose=0) + model.load_pdb(pdb_path) + model.set_adp_mode("field", n_nodes=n_nodes, k_neighbors=8, **kw) + return model + + +@pytest.mark.unit +def test_inert_outside_field_mode(pdb_path): + """Registered unconditionally, so it must cost nothing on the per-atom path.""" + model = Model(verbose=0) + model.load_pdb(pdb_path) + model.set_adp_mode("isotropic") + target = NodeLoadTarget(model) + assert float(target()) == 0.0 + assert target.stats()["node_load_active"].value == 0.0 + + +@pytest.mark.unit +def test_balanced_field_is_barely_penalised(pdb_path): + """A freshly fitted field has near-even load, so the barrier starts near its floor.""" + model = _field_model(pdb_path) + target = NodeLoadTarget(model) + rel = target._relative_load().detach() + # Mean relative load is 1 by construction. + assert torch.allclose(rel.mean(), torch.ones((), dtype=rel.dtype), atol=1e-6) + per_node = float(target()) / rel.numel() + assert per_node < 0.5, f"per-node penalty {per_node:.3f} on a balanced field" + + +@pytest.mark.unit +def test_abandoning_a_node_raises_the_penalty(pdb_path): + """Collapsing one node's kernel starves it, and the barrier must notice.""" + model = _field_model(pdb_path) + target = NodeLoadTarget(model) + before = float(target()) + + # Narrow one node far below the others: it loses every softmax contest except + # against an atom sitting on top of it. + with torch.no_grad(): + model.adp.refinable_params[0, 1] -= 6.0 + model.adp.reset_forward_cache() + + after = float(target()) + rel = target._relative_load().detach() + assert rel.min() < 0.25, "the node was not actually starved" + assert after > before + 1.0, f"barrier missed it: {before:.3f} -> {after:.3f}" + + +@pytest.mark.unit +def test_penalty_is_one_sided(pdb_path): + """Over-loading a node must not be penalised; only abandonment is. + + This is what separates the barrier from a load-entropy term, which is optimal at + uniform load and would push back on a legitimately broad node. + """ + model = _field_model(pdb_path) + target = NodeLoadTarget(model) + baseline = float(target()) + + with torch.no_grad(): + model.adp.refinable_params[0, 1] += 3.0 # widen one node -> it gains load + model.adp.reset_forward_cache() + widened = float(target()) + rel = target._relative_load().detach() + + assert float(rel.max()) > 1.5, "the node did not actually gain load" + # Widening one node necessarily takes load from others, so the total may rise a + # little; what must not happen is the over-loaded node itself being charged. + per_node_change = (widened - baseline) / rel.numel() + assert per_node_change < 0.5, ( + f"over-loading was penalised like abandonment ({per_node_change:.3f}/node)" + ) + + +@pytest.mark.unit +def test_gradient_reaches_geometry_but_not_values(pdb_path): + """Acts on the weights: positions and widths get gradient, node B does not.""" + model = _field_model(pdb_path) + target = NodeLoadTarget(model) + target().backward() + + grad = model.adp.refinable_params.grad + assert grad is not None + # Columns are [log B, log sigma, dx, dy, dz]. + assert float(grad[:, 0].abs().sum()) == pytest.approx(0.0, abs=1e-12), ( + "the barrier must not push the node VALUES" + ) + assert float(grad[:, 1].abs().sum()) > 0, "no gradient to the kernel widths" + assert float(grad[:, 2:5].abs().sum()) > 0, "no gradient to the node positions" + + +@pytest.mark.unit +def test_position_gradient_absent_when_positions_are_fixed(pdb_path): + """With positions fixed the barrier can only act through the widths.""" + model = _field_model(pdb_path, refine_node_positions=False) + target = NodeLoadTarget(model) + target().backward() + grad = model.adp.refinable_params.grad + assert grad.shape[1] == 2 + assert float(grad[:, 1].abs().sum()) > 0 + + +@pytest.mark.unit +def test_registered_in_the_total_adp_target(pdb_path): + """Reachable under the weight path 'adp/node_load'.""" + from torchref.refinement.targets.combined import TotalADPTarget + + model = _field_model(pdb_path) + total = TotalADPTarget(model, verbose=0) + assert "node_load" in total.target_losses() + assert torch.isfinite(torch.as_tensor(float(total["node_load"]()))) + + +@pytest.mark.unit +def test_default_weight_exists_for_the_component(pdb_path): + """A component with no weight entry would silently inherit the group weight.""" + from torchref.refinement.base_refinement import DEFAULT_GROUP_WEIGHTS + + assert "adp/node_load" in DEFAULT_GROUP_WEIGHTS + assert DEFAULT_GROUP_WEIGHTS["adp/node_load"] > 0 diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index 37751bc0..79580086 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -53,6 +53,16 @@ # neighbours, and the log-normal KL term it replaced was a single intensive # scalar. Pending the R_free weight scan, 1.0 leaves it at the group weight. "adp/sigd": 1.0, + # Load balancing for the node-field ADP representation. Sub-weight on the adp + # group, and inert on the per-atom path, so it only acts in field mode. Set + # above the group weight because it is a barrier against a degenerate direction + # rather than a prior competing with the data. + "adp/node_load": 10.0, + # Magnitude prior on the node values. Off pending its own measurement: the load + # barrier acts only on the weights, so this is what actually bounds an extreme + # node B, but it has not been screened yet. Same convention as + # 'geometry/ramachandran'. + "adp/node_smoothness": 0.0, } diff --git a/torchref/refinement/targets/adp/__init__.py b/torchref/refinement/targets/adp/__init__.py index b4638eeb..c7f3ec79 100644 --- a/torchref/refinement/targets/adp/__init__.py +++ b/torchref/refinement/targets/adp/__init__.py @@ -3,6 +3,8 @@ from .rigid_bond import RigidBondTarget from .sigd import ADPSigdTarget from .locality import ADPLocalityTarget +from .node_load import NodeLoadTarget +from .node_smoothness import NodeSmoothnessTarget from .scaler_log_scale import ScalerLogScaleTrendTarget from .scaler_u import ScalerURegularizationTarget @@ -12,6 +14,8 @@ "RigidBondTarget", "ADPSigdTarget", "ADPLocalityTarget", + "NodeLoadTarget", + "NodeSmoothnessTarget", "ScalerURegularizationTarget", "ScalerLogScaleTrendTarget", ] diff --git a/torchref/refinement/targets/adp/node_load.py b/torchref/refinement/targets/adp/node_load.py new file mode 100644 index 00000000..3eac95fb --- /dev/null +++ b/torchref/refinement/targets/adp/node_load.py @@ -0,0 +1,115 @@ +"""Load balancing for a node-field ADP representation.""" + +import torch +from typing import TYPE_CHECKING, Dict + +from torchref.utils.stats import ( + VERBOSITY_DEBUG, + VERBOSITY_DETAILED, + VERBOSITY_STANDARD, + StatEntry, + stat, +) + +from .base import ADPTarget + +if TYPE_CHECKING: + from torchref.model.model import Model + + +class NodeLoadTarget(ADPTarget): + """Keep every disorder-field node carrying a fair share of atoms. + + A node's load is the total weight it holds across all atoms, + :meth:`~torchref.model.disorder_field.DisorderFieldTensor.node_load`, and the + weights are a partition of unity, so the loads sum to the atom count and their mean + is ``n_atoms / K`` whatever the model does. + + Without this the field has a degenerate direction: a node can narrow its kernel + until it holds a single atom, then take whatever value fits that atom. Measured, a + collapsed node ends up with a load near or below one atom against a healthy median + of seven, and sets its atom's B into the hundreds or thousands. One node fitting one + atom is per-atom refinement wearing a node's clothes, which is the thing the + representation exists to avoid. + + The penalty is **one-sided**, ``softplus(-log(load / mean_load))``: it grows as a + node is abandoned, and flattens to zero once a node carries its share. That + asymmetry is deliberate. The symmetric choice -- maximising the entropy of the load + distribution -- is optimal at *uniform* load, so it would also penalise a broad node + that legitimately covers more atoms than its neighbours. Fitted fields span nearly + two orders of magnitude in kernel width within a single structure, and that spread + is the representation working, not failing. + + Acts through the weights, so its gradient reaches node positions and kernel widths + but never the node values: it removes the *opportunity* to place an extreme B rather + than penalising the B itself. It therefore composes with, rather than duplicates, + the restraints that act on the values. + + Inert unless the model is in field mode, so it can be registered unconditionally. + + Parameters + ---------- + model : Model, optional + Reference to the Model object. + sharpness : float, optional + Softplus temperature in log-load units. Smaller is a harder barrier. Default + 0.5, which leaves a node at the mean load contributing about 0.1 and a node at a + tenth of the mean about 2.3. + verbose : int, optional + Verbosity level. Default is 0. + """ + + def __init__( + self, + model: "Model" = None, + sharpness: float = 0.5, + verbose: int = 0, + device=None, + **kwargs, + ): + super().__init__(model, verbose, device=device, **kwargs) + self.sharpness = float(sharpness) + + @property + def _field(self): + """The disorder field, or ``None`` when the model is not in field mode.""" + adp = getattr(self.model, "adp", None) + return adp if hasattr(adp, "node_load") else None + + def _relative_load(self) -> torch.Tensor: + """Each node's load as a multiple of the mean load, ``(K,)``.""" + field = self._field + load = field.node_load() + # Mean load is n_atoms / K exactly, because the weights sum to one per atom. + return load / (load.sum().detach() / load.shape[0]).clamp(min=1e-12) + + def forward(self) -> torch.Tensor: + """Summed one-sided load deficit over nodes, or zero outside field mode.""" + field = self._field + if field is None: + return torch.zeros((), device=self.device) + rel = self._relative_load() + deficit = -torch.log(rel.clamp(min=1e-12)) / self.sharpness + return torch.nn.functional.softplus(deficit).sum() * self.sharpness + + def stats(self) -> Dict[str, any]: + """Load distribution across nodes, and how much of it the barrier sees.""" + field = self._field + if field is None: + return {"node_load_active": stat(0.0, VERBOSITY_DEBUG)} + with torch.no_grad(): + rel = self._relative_load() + loss = self.forward() + return { + "node_load_loss": stat(float(loss), VERBOSITY_STANDARD), + "n_nodes": stat(int(rel.numel()), VERBOSITY_STANDARD), + "load_min_rel": stat(float(rel.min()), VERBOSITY_STANDARD), + "load_median_rel": stat(float(rel.median()), VERBOSITY_DETAILED), + "load_max_rel": stat(float(rel.max()), VERBOSITY_DETAILED), + # The population the barrier exists for. + "n_below_quarter_share": stat( + int((rel < 0.25).sum()), VERBOSITY_STANDARD + ), + "load_cv": stat(float(rel.std() / rel.mean().clamp(min=1e-12)), + VERBOSITY_DEBUG), + } diff --git a/torchref/refinement/targets/adp/node_smoothness.py b/torchref/refinement/targets/adp/node_smoothness.py new file mode 100644 index 00000000..1db4b413 --- /dev/null +++ b/torchref/refinement/targets/adp/node_smoothness.py @@ -0,0 +1,141 @@ +"""Magnitude prior on the node values of a disorder field.""" + +import torch +from typing import TYPE_CHECKING, Dict + +from torchref.utils.stats import ( + VERBOSITY_DEBUG, + VERBOSITY_DETAILED, + VERBOSITY_STANDARD, + StatEntry, + stat, +) + +from .base import ADPTarget + +if TYPE_CHECKING: + from torchref.model.model import Model + + +class NodeSmoothnessTarget(ADPTarget): + """Penalise a node whose B departs from the nodes around it. + + The companion to :class:`~torchref.refinement.targets.adp.NodeLoadTarget`, which + acts on the weights and therefore cannot reach the node *values* at all. Blocking a + node from narrowing does not stop it taking an extreme B -- measured, it makes the + consequence broader rather than smaller, because the extreme value can no longer be + confined to the single atom the node had isolated. So the two terms close different + halves: one denies the opportunity, this one prices the magnitude. + + The penalty is a distance-weighted sum over node pairs:: + + L = sum_{k 1 else 1.0 + lam = max(lam, 1e-3) + + w = torch.exp(-(d**2) / (2.0 * lam * lam)) + w = torch.triu(w, diagonal=1) + diff2 = (log_b[:, None] - log_b[None, :]) ** 2 + return w, diff2, lam + + def forward(self) -> torch.Tensor: + """Weighted mean squared log-B difference between nearby nodes.""" + field = self._field + if field is None or field.n_nodes < 2: + return torch.zeros((), device=self.device) + w, diff2, _ = self._pair_terms() + total = w.sum() + if float(total) <= 0.0: + return torch.zeros((), device=self.device) + return (w * diff2).sum() / total + + def stats(self) -> Dict[str, any]: + """Spread of the node values, and how localised the departures are.""" + field = self._field + if field is None or field.n_nodes < 2: + return {"node_smoothness_active": stat(0.0, VERBOSITY_DEBUG)} + with torch.no_grad(): + w, diff2, lam = self._pair_terms() + loss = self.forward() + log_b = field.node_values()[:, 0] + b = torch.exp(log_b) + return { + "node_smoothness_loss": stat(float(loss), VERBOSITY_STANDARD), + "node_b_median": stat(float(b.median()), VERBOSITY_STANDARD), + "node_b_max": stat(float(b.max()), VERBOSITY_STANDARD), + "node_log_b_sd": stat(float(log_b.std()), VERBOSITY_DETAILED), + "node_pair_length_scale": stat(float(lam), VERBOSITY_DETAILED), + # How far the worst node sits above its own neighbourhood. + "node_b_max_over_median": stat( + float(b.max() / b.median().clamp(min=1e-12)), VERBOSITY_STANDARD + ), + } diff --git a/torchref/refinement/targets/combined.py b/torchref/refinement/targets/combined.py index 545201a2..2fa318f9 100644 --- a/torchref/refinement/targets/combined.py +++ b/torchref/refinement/targets/combined.py @@ -20,6 +20,8 @@ ) from torchref.refinement.targets.adp import ( ADPSimilarityTarget, ADPLocalityTarget, ADPSigdTarget, + NodeLoadTarget, + NodeSmoothnessTarget, ) from torchref.utils.stats import ( VERBOSITY_DETAILED, @@ -371,6 +373,12 @@ def _create_targets(self) -> Dict[str, Target]: self.model, verbose=self.verbose ), "sigd": ADPSigdTarget(self.model, verbose=self.verbose), + # Inert unless the model is in field mode, so it costs a zero tensor + # per call on the per-atom path. + "node_load": NodeLoadTarget(self.model, verbose=self.verbose), + "node_smoothness": NodeSmoothnessTarget( + self.model, verbose=self.verbose + ), } def print_statistics(self) -> None: From 26107f15b4ebdbb50d77ee1ebe3ddae5e08f951f Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 12:59:06 +0200 Subject: [PATCH 079/250] Drop dead code from the topology restraint orchestrator Seven methods with no callers anywhere in the package or the tests, left behind when the restraint layer moved into topology and the geometry targets took over the loss maths: expand_altloc, the h_excl_hash property, torsion_deviations, nll_torsions, nll_planes, nll_vdw and adp_similarity_loss. The h_excl_hash property was only a read-only wrapper over self._h_excl_hash; that attribute is still built and passed to the non-bonded pair list, and the identically named arguments in nonbonded.py and riding.py are unrelated function parameters. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- torchref/topology/restraints.py | 236 +------------------------------- 1 file changed, 3 insertions(+), 233 deletions(-) diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index 76c217f0..4a649df7 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -79,8 +79,9 @@ class Restraints(DeviceMixin, DebugMixin, Module): cif_dict : dict Parsed CIF restraints keyed by residue type; ``missing_residues`` lists the types that could not be resolved. - h_topo, h_excl_hash - Riding-hydrogen topology and its exclusion hash, populated on demand. + h_topo + Riding-hydrogen map, built only when the model carries no hydrogens of its + own. Empty otherwise; see :mod:`torchref.topology.riding`. link_dict, link_list Link-type definitions from the monomer library, set only when ``pdb`` was provided. @@ -330,30 +331,6 @@ def _load_cif_dictionaries(self, cif_path): f"and will have no restraints applied: {self.missing_residues}" ) - def expand_altloc(self, residue): - """ - Expand residue with alternative conformations into separate conformations. - - Yields one DataFrame per altloc (with common atoms included in each). - """ - residue = residue.copy() - residue.loc[residue["altloc"].isin(["", " "]), "altloc"] = " " - - alt_conf = residue["altloc"].unique() - if " " in alt_conf: - residue_no_alt = residue.loc[residue["altloc"] == " "] - for alt in alt_conf: - if alt == " ": - continue - residue_alt = residue.loc[residue["altloc"] == alt] - residue_combined = pd.concat( - [residue_no_alt, residue_alt], ignore_index=True - ) - yield residue_combined - else: - for alt_loc in alt_conf: - residue_alt = residue.loc[residue["altloc"] == alt_loc] - yield residue_alt def _load_rama_surfaces(self, device: torch.device): """Load pre-computed Ramachandran NLL surfaces as a buffer.""" @@ -666,10 +643,6 @@ def h_topo(self): """Access riding hydrogen topology (None if not built).""" return getattr(self, "_h_topo", None) - @property - def h_excl_hash(self): - """Access H-specific exclusion hash tensor (None if not built).""" - return getattr(self, "_h_excl_hash", None) def _build_h_exclusion_hash(self, h_topo, device): """Sorted 1-D hash tensor of H-specific 1-2 and 1-3 exclusions. @@ -1512,49 +1485,6 @@ def _wrap_torsion_periodicity(self, diff_rad, periods): # All periods are 0 or 1, simple wrapping return torch.remainder(diff_rad + torch.pi, 2.0 * torch.pi) - torch.pi - def torsion_deviations(self, xyz: torch.Tensor = None, wrapped=True): - """ - Compute deviations between calculated and expected torsion angles. - - Parameters - ---------- - xyz : torch.Tensor, optional - Coordinates tensor. If None, uses the stored xyz_fn callable. - wrapped : bool, default True - If True, wrap deviations accounting for periodicity. - If False, return raw deviations (calculated - expected). - - Returns - ------- - torch.Tensor - Tensor of shape (n_torsions,) with deviations in degrees. - For wrapped=True, deviations are in range appropriate for the period. - - Notes - ----- - Expected values from the CIF library are discrete (typically -60°, 0°, - 60°, 90°, 180°) while calculated values from the structure are - continuous. Use wrapped=True for meaningful comparison and - visualization. - """ - if "all" not in self.restraints["torsion"]: - self.cat_dict() - - idx = self.restraints["torsion"]["all"]["indices"] - expected = self.restraints["torsion"]["all"]["references"] - periods = self.restraints["torsion"]["all"]["periods"] - calculated = self.torsions(idx, xyz) - - if not wrapped: - # Simple difference - return calculated - expected - else: - # Use the helper function for periodicity handling - diff_rad = (calculated - expected) * torch.pi / 180.0 - diff_wrapped_rad = self._wrap_torsion_periodicity(diff_rad, periods) - - # Convert back to degrees - return torch.rad2deg(diff_wrapped_rad) def torsion_deviations_with_sigmas(self, xyz: torch.Tensor = None): """ @@ -1588,143 +1518,8 @@ def torsion_deviations_with_sigmas(self, xyz: torch.Tensor = None): return deviations_rad, sigmas_deg - def nll_torsions(self, xyz: torch.Tensor = None): - """ - Compute negative log-likelihood for torsion angle restraints. - - von Mises: NLL = -κ·cos(θ-μ) + log(I₀(κ)) + log(2π), with κ = 1/σ². This is - the true NLL, so exp(-NLL) is a probability density. Deviations are folded - by the restraint period first (see :meth:`_wrap_torsion_periodicity`). - - Parameters - ---------- - xyz : torch.Tensor, optional - Coordinates tensor. If None, uses the stored xyz_fn callable. - - Returns - ------- - torch.Tensor - Tensor of shape (n_torsions,) with negative log-likelihood values. - """ - from torchref.refinement.targets import von_mises_nll - - deviations_rad, sigmas_deg = self.torsion_deviations_with_sigmas(xyz) - return von_mises_nll(deviations_rad, sigmas_deg) - - def nll_planes(self, xyz: torch.Tensor = None): - """ - Compute negative log-likelihood for plane restraints. - For each plane, computes the RMSD of atom deviations from the best-fit plane. - Uses Gaussian NLL: NLL = 0.5 * (deviation / σ)² + log(σ) + 0.5 * log(2π) - - Parameters - ---------- - xyz : torch.Tensor, optional - Coordinates tensor. If None, uses the stored xyz_fn callable. - Returns - ------- - torch.Tensor - Tensor of shape (n_planes,) with negative log-likelihood values. - """ - from torchref.refinement.targets import gaussian_nll - - xyz = self.xyz(xyz) - device = xyz.device - - all_nlls = [] - - if "plane" in self.restraints: - for key, plane_data in self.restraints["plane"].items(): - indices = plane_data.get("indices") - sigmas = plane_data.get("sigmas") - - if indices is None or len(indices) == 0: - continue - - # indices shape: (n_planes, n_atoms_per_plane) - # sigmas shape: (n_planes, n_atoms_per_plane) - n_planes, n_atoms = indices.shape - - for i in range(n_planes): - plane_indices = indices[i] - plane_sigmas = sigmas[i] - - # Get positions of atoms in this plane - positions = xyz[plane_indices] # (n_atoms, 3) - - # Compute centroid - centroid = positions.mean(dim=0) - centered = positions - centroid - - # SVD to find best-fit plane normal - # The plane normal is the singular vector with smallest singular value - U, S, Vh = torch.linalg.svd(centered) - normal = Vh[-1] # Normal to best-fit plane - - # Compute deviations from plane (distance to plane) - deviations = torch.abs(centered @ normal) - - # Compute NLL for each atom - nll = gaussian_nll(deviations, plane_sigmas) - all_nlls.append(nll) - - if all_nlls: - return torch.cat(all_nlls) - return torch.tensor([0.0], device=device) - - def nll_vdw(self, xyz: torch.Tensor = None): - """ - Compute negative log-likelihood for VDW (non-bonded) restraints. - - Uses a soft-repulsive potential based on distance violations. - NLL = 0.5 * (max(0, min_dist - actual_dist) / σ)² + log(σ) + 0.5 * log(2π) - - Only violations (distances shorter than minimum) contribute to the loss. - - Parameters - ---------- - xyz : torch.Tensor, optional - Coordinates tensor. If None, uses the stored xyz_fn callable. - - Returns - ------- - torch.Tensor - Tensor of shape (n_pairs,) with negative log-likelihood values. - """ - from torchref.refinement.targets import gaussian_nll - - xyz = self.xyz(xyz) - device = xyz.device - - if "vdw" not in self.restraints: - return torch.tensor([0.0], device=device) - - vdw_data = self.restraints["vdw"] - indices = vdw_data.get("indices") - - if indices is None or len(indices) == 0: - return torch.tensor([0.0], device=device) - - min_distances = vdw_data["min_distances"] - sigmas = vdw_data["sigmas"] - - # Get current positions - pos1 = xyz[indices[:, 0]] - pos2 = xyz[indices[:, 1]] - - # Compute actual distances - actual_distances = torch.norm(pos2 - pos1, dim=-1) - - # Violations: where actual distance is less than minimum - # Deviation = max(0, min_dist - actual_dist) - deviations = torch.clamp(min_distances - actual_distances, min=0.0) - - # Compute NLL (only non-zero for violations) - nll = gaussian_nll(deviations, sigmas) - - return nll def adp_b_differences(self, adp: torch.Tensor = None): """ @@ -1757,28 +1552,3 @@ def adp_b_differences(self, adp: torch.Tensor = None): return torch.cat(diffs_list, dim=0) return torch.tensor([], device=b_factors.device) - def adp_similarity_loss(self, adp: torch.Tensor = None, sigma: float = 2.0): - """ - Compute ADP similarity loss (SIMU in Phenix/SHELX). - - This restrains the B-factors of bonded atoms to be similar. - Loss = Σ ((B_i - B_j) / sigma)^2 - - Parameters - ---------- - adp : torch.Tensor, optional - ADP values. If None, uses the stored adp_fn callable. - sigma : float, default 2.0 - Target standard deviation for B-factor differences in Ų. - - Returns - ------- - torch.Tensor - Mean similarity loss. - """ - from torchref.refinement.targets import adp_similarity_nll - - b_diffs = self.adp_b_differences(adp) - if len(b_diffs) == 0: - return torch.tensor(0.0, device=self.xyz().device) - return adp_similarity_nll(b_diffs, sigma).mean() From 4a531c84813f06a52d57f86fadb2472ca1ed65bb Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 13:27:59 +0200 Subject: [PATCH 080/250] Fix the corrected DED coefficients, which were phase-blind `mDFop-DFc_corr` was built as `|F_corr - F_dark|`, a scalar amplitude difference, while its uncorrected twin `mDFop-DFc` uses `|F_light e^{i phi_light} - F_dark e^{i phi_dark}|` -- the modulus of the complex vector difference, which carries the phase rotation between dark and light. Those are different quantities, and the scalar one is exactly the phase-blind form a phase-aware coefficient exists to avoid. The symptom: `DF` and `DF_corr` agree to 99.99% (they differ by the 1.1% activation correction), yet the DED coefficients built from them were only 36.5% correlated, with means +0.477 and -1.920. An amplitude change that small cannot do that; the construction was wrong. Now rebuilt with the decontaminated amplitude substituted into the same complex expression. Correlation with the uncorrected coefficients goes 0.365 -> 0.999985 and the corrected map differs from the uncorrected by 0.54% rms, which is what a 1.1% amplitude correction should produce. Also adds paper/make_ded_maps.py: every map on PHIC_diff, so maps from different runs share phases and are directly comparable. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- docs/changelog.rst | 2 + paper/make_ded_maps.py | 75 ++++++++++++++++++++ torchref/cli/collection_difference_refine.py | 26 +++++-- 3 files changed, 99 insertions(+), 4 deletions(-) create mode 100644 paper/make_ded_maps.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 4851cc7e..df314cd9 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,8 @@ Changelog Version 0.6.4 ---------- +- Fixed the ``--two-moment`` corrected DED coefficients using a phase-blind amplitude difference instead of the phase-aware one the uncorrected coefficients use +- Added ``paper/make_ded_maps.py``, which writes CCP4 maps from a difference-refine results MTZ - ``CollectionScaler.refine_lbfgs_joint`` builds a row of ``XRAY_TARGETS`` instead of its own Rice likelihood, and takes ``scale_target`` (default ``ls``) - ``CollectionScaler.refine_lbfgs_joint`` normalises its objective and registers the U penalty as its own target - Fixed ``DatasetCollection.scale`` fitting the inter-dataset scale on the free reflections as well as the work set diff --git a/paper/make_ded_maps.py b/paper/make_ded_maps.py new file mode 100644 index 00000000..47bca6e9 --- /dev/null +++ b/paper/make_ded_maps.py @@ -0,0 +1,75 @@ +#!/usr/bin/env python +"""Turn a ``torchref.difference-refine`` results MTZ into CCP4 maps for PyMOL/Coot. + +Coot opens the MTZ directly (File > Auto Open MTZ, or pick the column pair), so this +exists for PyMOL, which wants a real map. Every map is computed on ``PHIC_diff``, so +maps from different runs are directly comparable -- they share phases. + +Which coefficient is which: + +``mDFop-DFc`` + The difference map. ``m`` is a normalised inverse-variance weight, **not** a sigma_A + figure of merit. Contour at +-3 sigma. +``mDFop-DFc_corr`` + The same, from the activation-decontaminated light amplitude. Present only when the + run had ``--two-moment``. +``DDF`` + ``DF_corr - DF``: the correction itself, as a map. Featureless against resolution + means the correction is collinear with a scale or overall-B error and should be + distrusted; structure in it is signal. +``2mDFop-DFc`` + The 2Fo-Fc analogue, for seeing the model in its density. +""" + +import argparse +import sys +from pathlib import Path + +import gemmi +import numpy as np + +# label -> (amplitude column, phase column). Skipped silently when absent. +MAPS = { + "ded": ("mDFop-DFc", "PHIC_diff"), + "ded_corr": ("mDFop-DFc_corr", "PHIC_diff"), + "ded2": ("2mDFop-DFc", "PHIC_diff"), + "ded2_corr": ("2mDFop-DFc_corr", "PHIC_diff"), + "ddf": ("DDF", "PHIC_diff"), + "wdf": ("WDF", "PHIC_diff"), +} + + +def main(): + ap = argparse.ArgumentParser(description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter) + ap.add_argument("mtz", help="fractions_*_difference_data.mtz from a refine run") + ap.add_argument("-o", "--outdir", default=".", help="where to write the .ccp4 files") + ap.add_argument("--prefix", default="", help="prefix for the output names") + # 3.0 matches the library's FFT oversampling; below ~2.5 the peaks shift. + ap.add_argument("--sample-rate", type=float, default=3.0) + args = ap.parse_args() + + out = Path(args.outdir) + out.mkdir(parents=True, exist_ok=True) + mtz = gemmi.read_mtz_file(args.mtz) + have = {c.label for c in mtz.columns} + + print(f"{args.mtz}\n{'map':12s} {'coefficient':22s} {'rms':>10s} {'peak':>10s}") + print("-" * 58) + for name, (f, ph) in MAPS.items(): + if f not in have or ph not in have: + continue + grid = mtz.transform_f_phi_to_map(f, ph, sample_rate=args.sample_rate) + ccp4 = gemmi.Ccp4Map() + ccp4.grid = grid + ccp4.update_ccp4_header() + path = out / f"{args.prefix}{name}.ccp4" + ccp4.write_ccp4_map(str(path)) + a = np.array(grid, copy=False) + print(f"{name:12s} {f:22s} {a.std():10.5f} {np.abs(a).max():10.5f}") + print(f"\nwritten to {out.resolve()}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index 9c30775f..db4c4188 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -326,7 +326,7 @@ def compute_bayes_extrapolated_amplitudes( def _two_moment_columns(mc, dc, mask, fcalc_dark_full, fcalc_mixed_full, *, weights, diff_Fobs, Fcalc_diff_amp, Fobs_dark, - sig_dark): + sig_dark, phi_mixed, F_obs_dark_phased): """Two-moment diagnostic columns, or empty dicts when the model is off. The observed light intensity carries a positive, phase-blind contamination @@ -405,8 +405,25 @@ def _np(t): DDF = DF_corr - diff_Fobs sig_DF_corr = np.sqrt(sig_F_corr**2 + sig_dark**2) - amp_2_corr = (2 * np.abs(DF_corr) - Fcalc_diff_amp) * weights - amp_1_corr = (np.abs(DF_corr) - Fcalc_diff_amp) * weights + # The phase-AWARE difference, rebuilt with the decontaminated light amplitude. + # + # It must be the modulus of the complex vector difference + # ``|F_corr e^{i phi_light} - F_dark e^{i phi_dark}|``, exactly as the uncorrected + # ``Fobs_diff_phased`` is -- NOT ``|F_corr - F_dark|``. The two are different + # quantities: the vector form carries the phase rotation between dark and light, + # which is the whole point of a phase-aware coefficient, while the scalar form is + # phase-blind. Using the scalar one here made the corrected coefficients only 36% + # correlated with their uncorrected twins even though the amplitudes behind them + # agree to 99.99%. + F_corr_phased = torch.as_tensor( + F_corr, dtype=F_obs_dark_phased.real.dtype, device=F_obs_dark_phased.device + ) * torch.exp(1j * phi_mixed) + Fobs_diff_phased_corr = ( + torch.abs(F_corr_phased - F_obs_dark_phased).detach().cpu().numpy() + ) + + amp_2_corr = (2 * Fobs_diff_phased_corr - Fcalc_diff_amp) * weights + amp_1_corr = (Fobs_diff_phased_corr - Fcalc_diff_amp) * weights # The sigma_alpha^2-aware weight, on the same normalisation as the inverse-variance # weight the existing DED coefficients carry, so the two are directly comparable. @@ -604,7 +621,8 @@ def _extrapolation_rfactors(data, fcalc_scaled): mc, dc, mask, fcalc_dark_full, fcalc_mixed_full, weights=weights, diff_Fobs=diff_Fobs, Fcalc_diff_amp=Fcalc_diff_amp, Fobs_dark=Fobs_dark, - sig_dark=sig_dark, + sig_dark=sig_dark, phi_mixed=phi_mixed, + F_obs_dark_phased=F_obs_dark_phased, ) ) From 6ac01aa9ea61fab4bec18cdb6bb8ae5558b29bff Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 14:07:34 +0200 Subject: [PATCH 081/250] Record mask_source in the validate-ded results It changes the answer and was the one parameter the results JSON did not store. On the figure-4 ligand, --mask-source light masks 1535 voxels and scores CC 0.869; the default "both" masks 1948 and scores 0.851. The union adds the volume the ligand vacated on isomerisation, where the difference density is negative and the model has to reproduce a depletion -- the harder half. Two sets of runs differing only in this looked like a 0.02 improvement, and recovering which was which took several re-runs because `selection` and `mask_radius` were stored but `mask_source` was not. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- docs/changelog.rst | 1 + torchref/cli/validate_ded.py | 8 ++++++++ 2 files changed, 9 insertions(+) diff --git a/docs/changelog.rst b/docs/changelog.rst index df314cd9..9e7154ff 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Version 0.6.4 ---------- +- ``torchref.validate-ded`` records ``mask_source`` in its results JSON; it changes the correlation and was not recoverable from the output - Fixed the ``--two-moment`` corrected DED coefficients using a phase-blind amplitude difference instead of the phase-aware one the uncorrected coefficients use - Added ``paper/make_ded_maps.py``, which writes CCP4 maps from a difference-refine results MTZ - ``CollectionScaler.refine_lbfgs_joint`` builds a row of ``XRAY_TARGETS`` instead of its own Rice likelihood, and takes ``scale_target`` (default ``ls``) diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index 307d0c69..8ba0e81b 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -638,6 +638,14 @@ def run_validation(args): "light_model": str(args.light_model), "fraction": args.fraction, "selection": args.selection, + # Recorded because it CHANGES THE ANSWER and is easy to leave at a + # different value between runs. On the figure-4 ligand, "light" masks 1535 + # voxels and scores CC 0.869, while "both" masks 1948 and scores 0.851 -- + # the union adds the volume the ligand vacated, where the density is + # negative and the model has to get a depletion right. Two runs differing + # only in this looked like a real improvement until the parameter was + # recovered by re-running, which is exactly what storing it prevents. + "mask_source": args.mask_source, "mask_radius": args.mask_radius, "dmin": d_min, }, From eb925b43f2cee54119848653a01a2a7853a5624d Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 14:18:31 +0200 Subject: [PATCH 082/250] Probe whether the DED correlation can compare targets across observables validate-ded correlates WDFo against WDFc, and WDFc is (|F_mixed| - |F_dark|)*w -- up to the weight and the transform, the residual CollectionDifferenceTarget minimises. So an amplitude target scored on it is close to being scored on its own objective, and a win there is not evidence. Computing the same correlation in both spaces on the same reflections and models does NOT flip the ranking: the amplitude arm leads in intensity space too. But a paired bootstrap says none of it is significant -- amp_free +0.0086 CI [-0.0264, +0.0448] int_free +0.0239 CI [-0.0565, +0.0961] with 919 free reflections and correlations of 0.09-0.21. The free set on this pair cannot adjudicate a difference of this size in either space, which is the same limit that stops it supporting a two-moment detection. Consequence: do not cite a DED-CC difference of order 0.01-0.02 between targets as a result. That includes the ligand-mask CC, whose voxels are band-limited at 2.2 A and so far fewer than its n_voxels suggests. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- paper/probe_ded_metric_space.py | 152 ++++++++++++++++++++++++++++++++ 1 file changed, 152 insertions(+) create mode 100644 paper/probe_ded_metric_space.py diff --git a/paper/probe_ded_metric_space.py b/paper/probe_ded_metric_space.py new file mode 100644 index 00000000..8b79975f --- /dev/null +++ b/paper/probe_ded_metric_space.py @@ -0,0 +1,152 @@ +#!/usr/bin/env python +"""Is the DED correlation a fair way to compare an amplitude and an intensity target? + +``torchref.validate-ded`` correlates ``WDFo`` against ``WDFc``, and ``WDFc`` is +``(|F_mixed| - |F_dark|) * w`` -- a weighted **amplitude** difference. That is, up to the +weight and the Fourier transform, exactly the residual ``CollectionDifferenceTarget`` +minimises. Scoring an amplitude target on it is close to scoring it on its own objective, +so a win there is not evidence. + +This computes the same correlation in **both** spaces, on the same reflections and the +same models: + + amplitude obs Fo_light - Fo_dark calc |Fc_light| - |Fc_dark| + intensity obs Io_light - Io_dark calc |Fc_light|^2 - |Fc_dark|^2 + +If the ranking flips between the two, neither is decisive and the comparison has to be +made on something neither target optimises. Reported on the FREE set as well as the work +set, since both targets were fitted on the work set. + +``F_calc`` is taken from each run's own results MTZ -- the scaled, mixed amplitudes that +run produced -- so no model is re-scaled here and each arm is scored on what it actually +built. Observed intensities come from one collection build shared by every arm. +""" + +import argparse +import json +import sys +from pathlib import Path + +import numpy as np +import reciprocalspaceship as rs +import torch + + +def cc(a, b): + a, b = np.asarray(a, float), np.asarray(b, float) + ok = np.isfinite(a) & np.isfinite(b) + if ok.sum() < 3: + return float("nan") + return float(np.corrcoef(a[ok], b[ok])[0, 1]) + + +def observed_intensities(dark_sf, light_sf, d_min, device): + """Scaled ``(hkl, I_dark, I_light)`` from one collection build.""" + from torchref.cli.collection_difference_refine import setup_dataset_collection + + dc = setup_dataset_collection(dark_sf, light_sf, d_min, device) + I_d, _ = dc["dark"].get_corrected_intensities() + I_l, _ = dc["light"].get_corrected_intensities() + hkl = dc.hkl.cpu().numpy() + return hkl, I_d.cpu().numpy(), I_l.cpu().numpy() + + +_PAIRED = [] + + +def main(): + ap = argparse.ArgumentParser(description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter) + fig4 = Path(__file__).resolve().parent / "figure4_difference_refinement" + ap.add_argument("mtz", nargs="+", help="one results MTZ per arm (label=path accepted)") + ap.add_argument("--dark-sf", default=str(fig4 / "data/8QL2-sf.cif")) + ap.add_argument("--light-sf", default=str(fig4 / "data/7YYZ-light.mtz")) + ap.add_argument("--dmin", type=float, default=2.2) + ap.add_argument("--device", default="cpu") + ap.add_argument("-o", "--out", default=None) + args = ap.parse_args() + + hkl, Id_full, Il_full = observed_intensities( + args.dark_sf, args.light_sf, args.dmin, torch.device(args.device)) + key = {tuple(h): i for i, h in enumerate(hkl)} + + rows = [] + for spec in args.mtz: + label, _, path = spec.partition("=") + if not path: + label, path = Path(spec).parent.name, spec + ds = rs.read_mtz(path).reset_index() + H = ds[["H", "K", "L"]].to_numpy() + idx = np.array([key.get(tuple(h), -1) for h in H]) + ok = idx >= 0 + + Fo_d = ds["Fo_dark"].to_numpy(float) + Fo_l = ds["Fo_light"].to_numpy(float) + Fc_d = ds["Fc_dark"].to_numpy(float) + Fc_l = ds["Fc_light"].to_numpy(float) + free = ds["FreeR_flag_light"].to_numpy() == 0 + + Id = np.full(len(H), np.nan); Il = np.full(len(H), np.nan) + Id[ok] = Id_full[idx[ok]]; Il[ok] = Il_full[idx[ok]] + + dFo, dFc = Fo_l - Fo_d, Fc_l - Fc_d # amplitude difference + dIo, dIc = Il - Id, Fc_l**2 - Fc_d**2 # intensity difference + + r = {"arm": label} + sel_map = {"work": ~free & ok, "free": free & ok} + for name, sel in sel_map.items(): + r[f"amp_{name}"] = cc(dFo[sel], dFc[sel]) + r[f"int_{name}"] = cc(dIo[sel], dIc[sel]) + r[f"n_{name}"] = int(sel.sum()) + rows.append(r) + if len(_PAIRED) < 2: + _PAIRED.append({"arm": label, "amp": (dFo, dFc), "int": (dIo, dIc), + "sel": sel_map}) + + print() + print("Difference-signal correlation, same reflections and models, two spaces") + print("=" * 78) + print(f"{'arm':16s} {'amp work':>10s} {'amp free':>10s} {'int work':>10s} " + f"{'int free':>10s} {'n free':>8s}") + print("-" * 78) + for r in rows: + print(f"{r['arm']:16s} {r['amp_work']:10.4f} {r['amp_free']:10.4f} " + f"{r['int_work']:10.4f} {r['int_free']:10.4f} {r['n_free']:8d}") + print("-" * 78) + # Paired bootstrap over reflections. A correlation is not a mean, so the + # difference of two CCs has no closed-form error; resampling the SAME reflections + # for both arms keeps the comparison paired, which matters because most of the + # scatter is shared signal that cancels. + if len(rows) >= 2 and _PAIRED: + print() + print("Paired bootstrap on the CC difference (4000 resamples over reflections)") + print("-" * 78) + a, b = _PAIRED[0], _PAIRED[1] + rng = np.random.default_rng(0) + for sp, (oa, ca, ob, cb) in (("amp", a["amp"] + b["amp"]), + ("int", a["int"] + b["int"])): + for st in ("work", "free"): + m = a["sel"][st] + ia, ja = oa[m], ca[m] + ib, jb = ob[m], cb[m] + keep = np.isfinite(ia) & np.isfinite(ja) & np.isfinite(ib) & np.isfinite(jb) + ia, ja, ib, jb = ia[keep], ja[keep], ib[keep], jb[keep] + n = len(ia) + d = np.corrcoef(ia, ja)[0, 1] - np.corrcoef(ib, jb)[0, 1] + boot = np.empty(4000) + for k in range(4000): + s_ = rng.integers(0, n, n) + boot[k] = (np.corrcoef(ia[s_], ja[s_])[0, 1] + - np.corrcoef(ib[s_], jb[s_])[0, 1]) + lo, hi = np.percentile(boot, [2.5, 97.5]) + flag = "" if lo <= 0 <= hi else " <-- CI excludes 0" + print(f" {sp}_{st:4s} d = {d:+.4f} 95% CI [{lo:+.4f}, {hi:+.4f}]" + f" n={n}{flag}") + print(f"\n positive => {_PAIRED[0]['arm']} predicts the difference better") + if args.out: + Path(args.out).write_text(json.dumps(rows, indent=2)) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) From 03026883052ac325e1a584a3094707f4b3defe0a Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 14:23:29 +0200 Subject: [PATCH 083/250] Record that a reflection hold-out is the wrong domain for a local feature The free-set numbers in this probe cannot answer what it asks, and the reason is structural rather than statistical. A difference feature is compact in real space, so its information is spread over every reflection; a held-out subset does not contain a noisier copy of it, it does not contain it. Measured on the figure-4 pair with the same models: reciprocal CC is 0.52 over all reflections and 0.21 over the free 3.5%, while in the other domain it goes 0.53 over the full cell to 0.85 over the 0.12% of voxels around the ligand. Localising helps in real space and destroys the signal in reciprocal space. Reciprocal-all and realspace-full-cell agree, as Parseval requires. validate-ded already builds its maps from the full P1 expansion, so realspace_correlation is the right instrument and reciprocal_cc_free is not. Cross-validating a local feature needs a hold-out in real space -- an omit refinement -- or an independent dataset, or ground truth. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- paper/probe_ded_metric_space.py | 16 ++++++++++++++-- 1 file changed, 14 insertions(+), 2 deletions(-) diff --git a/paper/probe_ded_metric_space.py b/paper/probe_ded_metric_space.py index 8b79975f..8b9c3b92 100644 --- a/paper/probe_ded_metric_space.py +++ b/paper/probe_ded_metric_space.py @@ -14,8 +14,20 @@ intensity obs Io_light - Io_dark calc |Fc_light|^2 - |Fc_dark|^2 If the ranking flips between the two, neither is decisive and the comparison has to be -made on something neither target optimises. Reported on the FREE set as well as the work -set, since both targets were fitted on the work set. +made on something neither target optimises. + +**READ THIS BEFORE USING THE FREE-SET NUMBERS.** They are reported for completeness and +they cannot answer the question. A difference feature is compact in real space and +therefore spread over ALL of reciprocal space, so a held-out subset of reflections does +not contain a reduced-precision version of it -- it does not contain it. Measured on this +pair: the same models score 0.52 on all reflections and 0.21 on the free 3.5%, while +restricting in the other domain goes the other way, 0.53 over the full cell to 0.85 on the +0.12% of voxels around the ligand. Localising helps in real space and destroys the signal +in reciprocal space. + +So a reflection-wise hold-out validates a global scalar (R-free) and nothing local. To +cross-validate a local difference feature, hold out in the domain the feature lives in -- +an omit refinement -- or use an independent dataset, or ground truth. ``F_calc`` is taken from each run's own results MTZ -- the scaled, mixed amplitudes that run produced -- so no model is re-scaled here and each arm is scored on what it actually From 09cb92f7c4364238729aef1737ea709e646b03da Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 17:09:35 +0200 Subject: [PATCH 084/250] Cleaned up changelog and bumped version to 0.7.0 --- docs/changelog.rst | 48 ++++++++------------------------------------ pyproject.toml | 2 +- torchref/__init__.py | 2 +- 3 files changed, 10 insertions(+), 42 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 3ff8c152..1d1c0596 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -2,47 +2,15 @@ Changelog ========= -Unreleased +Version 0.7.0 ---------- -- Added ``Symmetry``, a crystallography-free symmetry group carrying the operations and every verb derived from them; ``SpaceGroup`` now specialises it -- ``Symmetry`` exposes the transform primitives ``apply_rotations`` / ``apply_translations`` / ``phase_factors`` and a cached ``reciprocal`` stack, replacing seven separate spellings of ``R^T h`` -- Moved the centric, systematic-absence and epsilon predicates onto ``Symmetry``; ``is_centric_from_hkl`` and ``get_centric_acentric_masks`` are gone -- Moved the HKL asymmetric-unit verbs onto ``SpaceGroup`` as ``expand_hkl`` / ``reduce_hkl`` / ``complete_hkl`` / ``canonicalize_hkl``; the module-level functions are gone -- Moved the grid-size helpers onto ``Symmetry``; removed ``torchref.symmetry.grid_utils`` and the duplicate ``spacegroup`` module-level functions -- Map symmetrization is now ``Symmetry.symmetrize_map``, caching one operator for the most recent grid shape; ``MapSymmetry`` and ``MapSymmetryDirect`` are private -- Symmetry classes are dataclasses over ``DeviceMixin`` instead of ``nn.Module``, so assigning a ``SpaceGroup`` to a model attribute is no longer intercepted by ``nn.Module.__setattr__`` -- Removed ``ReciprocalSymmetryGrid``, ``ReciprocalSymmetry``, ``expand_reciprocal_grid``, ``expand_reflections`` and ``extract_structure_factors_with_symmetry`` -- Removed the unused ``Cell`` gradient plumbing (``requires_grad`` argument and property, ``detach``) and the ``CellTensor`` alias -- Removed the ``Symmetry`` alias for ``SpaceGroup``; the name is now a distinct class -- Added ``ModelContext``, holding a model's unit cell, space group, atom table, link records and provenance; ``Model.cell`` / ``.spacegroup`` / ``.pdb`` still work and now read through it -- Moved ``Model``'s configuration and provenance onto the context: ``strip_H``, ``verbose``, ``links``, ``altloc_pairs``, ``initialized``, ``exclude_H_from_sf`` and the input paths are reached as ``model.ctx.*`` -- ``Model.copy`` and ``ModelFT.copy`` now copy the context in one step, cloning the space group instead of sharing it -- Removed ``Model.symmetry``; use ``Model.spacegroup`` -- ``HydrogenTopology`` is a dataclass with optional fields instead of an ``nn.Module`` whose buffers were attached after construction -- Added ``torchref.topology``: a ``Topology`` of a ``ResidueGraph`` over an ``AtomGraph``, holding the model's connectivity as typed edge blocks with a bond adjacency that answers ``neighbors(i)`` -- Topology residues are identified by ``(chain, resseq, icode)``, so a residue with an insertion code is no longer merged with the one it was inserted after -- Added ``AtomGraph.exclusions_12_13_14``, deriving non-bonded exclusions from bond connectivity rather than from which angles and torsions the monomer library happens to restrain -- ``Restraints.restraints`` is now a plain nested dict of tensors instead of an accessor object rebuilt on every read; the per-origin indices are views into the topology's contiguous edge blocks -- Restraint groups are laid out in a fixed order, so a rebuild produces the same row order in any process; previously the origins were concatenated in Python ``set`` iteration order -- Fixed ``cat_dict`` doubling every bond, angle and torsion restraint when called more than once -- Geometry restraints are built from the topology, retiring the intra-residue builder calls, the peptide/disulfide/LINK build methods and the ``TensorDict`` restraint storage -- Hydrogen generation is now template instantiation over the topology: ``Model.hydrogenate`` aligns each residue's monomer template onto the heavy atoms present and reads its hydrogens off, and the bond graph sets how many hydrogens a parent may carry -- Fixed hydrogen placement fitting the template over two bond shells, which spans rotatable torsions the model does not share and left 12% of side-chain hydrogens further than 1.5 A from their parent, where they were discarded -- Hydrogens on a centre whose template omits a real substituent -- a peptide-linked backbone nitrogen -- are now built from the bonded neighbours instead of the template frame -- Hydrogens with a free torsion (hydroxyl, thiol, amine, methyl) are identified from bond connectivity and their dihedral is scanned, rather than taken from whatever the library deposited -- Removed ``Model.generate_hydrogens``; ``Model.hydrogenate`` is the single path and no longer takes ``lbfgs_steps`` or ``max_iter`` -- Hydrogens are now present by default: ``strip_H`` defaults to False, and hydrogens a file does not carry are generated on load. New ``add_hydrogens`` argument turns generation off while still keeping any the file has -- Generation is decided per parent, so a partially hydrogenated structure is topped up rather than left as deposited -- Fixed the lazily-cached per-atom buffers (``vdw_radii``, ``Z``, the ITC92 coefficients) surviving a load that changes the atom count, which left them sized for the previous atom set -- Riding hydrogens are no longer placed when the model carries real ones, where they acted as phantom atoms in the non-bonded term -- Fixed the hydrogen valence cap counting only heavy neighbours, so a parent that already carried a hydrogen still had budget for another; generation was not idempotent and a save/reload added a spurious second amide hydrogen to every linked residue -- Moved the riding-hydrogen map from ``torchref.restraints.hydrogen_topology`` to ``torchref.topology.riding``, alongside the generation path it is the heavy-atom-only alternative to -- Added ``Topology.subset`` and ``copy``, plus the same on ``EdgeBlock``, ``ResidueGraph`` and ``AtomGraph``: a subset reindexes the surviving edges rather than re-reading the CIFs and re-matching the templates. An edge is dropped as soon as any of its atoms is, and a residue left with no atoms goes along with its links -- Fixed residues distinguished only by an insertion code losing their restraints: the builders group on ``(chain, resseq)``, so 100 and 100A merge into one residue whose atom names collide and only the first keeps any intra-residue geometry -- Moved the restraint orchestrator from ``torchref.restraints.restraints`` to ``torchref.topology.restraints`` and renamed ``RestraintsNew`` to ``Restraints`` -- ``_lookup_link_atom`` moved to ``torchref.topology.build``, which was importing it back out of the restraints module -- Removed ``torchref.restraints``. Its seven modules moved into ``torchref.topology``, where every one of their importers already lived: the monomer library, CIF reading and ``chem_mod`` patches to ``topology.monomer``, the template matchers to ``topology.builders`` and ``topology.builders_numba``, the non-bonded spatial search to ``topology.nonbonded``, and the Ramachandran surfaces to ``topology.ramachandran`` -- Fixed ``neighbor_search`` annotating locals with ``List`` without importing it +- Separated model configuration and provenance into ``ModelContext``. It now holds the unit cell, space group, atom table, link records, hydrogen settings, and input paths. +- Refactored ``Symmetry`` as a crystallography-free class with transform primitives, and made ``SpaceGroup`` a specialised subclass. +- Moved geometry predicates, HKL verbs, and grid-size helpers onto these classes as methods. +- Rebuilt geometry restraints from the topology instead of intra-residue builders. ``torchref.restraints`` was removed, restraint dictionaries are now plain nested dicts, and residues are identified by ``(chain, resseq, icode)`` to fix insertion-code merging. +- Reworked hydrogen generation as template instantiation over the topology. ``Model.hydrogenate`` now aligns monomer templates onto heavy atoms present, generation is the default, and ``AtomGraph.exclusions_12_13_14`` derives non-bonded exclusions from bond connectivity. +- Added ``Topology`` as a ``ResidueGraph`` over an ``AtomGraph`` with typed edge blocks and ``subset`` / ``copy`` operations that reindex surviving edges. +- Made ``HydrogenTopology`` a dataclass, changed ``Symmetry`` classes to dataclasses over ``DeviceMixin`` instead of ``nn.Module``, and removed unused ``Cell`` gradient plumbing and the ``ReciprocalSymmetryGrid`` / module-level expansion functions. Version 0.6.4 diff --git a/pyproject.toml b/pyproject.toml index 4f0c7ba6..045e58cd 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "torchref" -version = "0.6.4" +version = "0.7.0" description = "Pytorch based crystallographic refinement" readme = "README.md" requires-python = ">=3.10" diff --git a/torchref/__init__.py b/torchref/__init__.py index b11229e2..ce2513e7 100644 --- a/torchref/__init__.py +++ b/torchref/__init__.py @@ -40,7 +40,7 @@ General utilities and debugging tools. """ -__version__ = "0.6.4" +__version__ = "0.7.0" import os From 775576bcf78d38b705def47c65b270e69fb209b2 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 17:13:41 +0200 Subject: [PATCH 085/250] Add an E-value convention seam and a conformance harness for it Work in progress: the harness is built and has started finding things, but the conventions are not yet unified and one of its own findings is unresolved. Why this exists. Chasing the ML rescore -- measurably destructive, 18/30 against 24/30 without it, p=0.031 -- turned up that the rotation function and the rescore disagree about what E means on the same data. The FRF's observed side is a French-Wilson posterior weighted by DFAC**2; the rescore's is plain per-shell Wilson with epsilon divided out, and `ml_rotation.py` never sees sig_F at all -- no sig_F, no DFAC, no French-Wilson anywhere in the module, and FRFInputs does not carry them. Counted properly the package has NINE ways into E-space, three of them inside translation.py alone. The framing that makes it worth a seam rather than a patch: `E = F/sqrt(Sigma)` is a weighting choice wearing the costume of a units change. Correlating E_obs against E_calc IS correlating F against F with weight 1/Sigma(s), so "what should E do" and "how should we weight by information content" are one question, and all nine converters are undeclared answers to it. The two consumers need different things -- a global scale cancels in the FRF's correlation (proven: removing the antipodal copy scaled scores by exactly 4 and moved 98/100 ranks not at all) and does not cancel in the rescore's likelihood -- so the convention has to satisfy the stricter one and the FRF gets it for free. `e_values.py` holds the seam. Conventions are constructed FROM the data, because a fitted Sigma(s) cannot exist before the reflections do, so engines will take the class and instantiate it internally; configuration rides in as functools.partial. Six wrappers over what we ship plus SmoothSigmaE, a Chebyshev Sigma(s) on the scaler's basis, fitted as a Gamma GLM with log link -- the right likelihood rather than a convenience, since acentric F**2 is exponential with mean Sigma. Fitting the mean this way avoids the trap a regression on log F**2 falls into, which is exactly how the overall-anisotropy fit was biased. `calc_companion` is not a harness convenience: a French-Wilson posterior is defined for observations only, there being no measurement error on a calc set to shrink, which is why the FRF pairs french_wilson_preprocess on obs with wilson_normalise on calc. That pairing is forced, not an oversight. What the harness has already found: - epsilon counts lattice-centring cosets. On C2 every reflection comes back with epsilon >= 2 (2x23342, 4x14) where a primitive lattice gives 1 for 99.7%. Harmless wherever epsilon is used multiplicatively -- it cancels in per-shell normalisation and in the refinement targets' beta = eps*Sigma_N, which is why the validated refinement path never saw it -- and wrong in the one place it is used additively: compute_v_budget's V = eps - sigma_A**2. - THREE epsilon implementations exist, and they disagree. base's epsilon_from_hkl(hkl, spacegroup) is the one the whole refinement suite uses; alignment has its own compute_epsilon(hkl, matrices) taking raw matrices. The base one is Friedel-aware, counting h->h OR h->-h. Measured across all ten benchmark structures, every disagreement lands on a centric reflection (cen_only True, 10/10) and the counts are large: 12360 reflections on 2DQ6, 7555 on 4BX9, 6680 on 3K7M. - The FRF's own default fails two checks on 1DAW: obs/calc ratio 2.101, the same class of defect as the rescore's eImove 2.14, and shrinkage 0.247 -- only a quarter of reflections move toward the shell mean when sigma_F is quadrupled, where a posterior should be monotone. That one may well be my test's reference point rather than the port; unresolved either way. Two of my own bugs fixed on the way here, both found by the harness rather than by reading: SmoothSigmaE's Gamma GLM diverged on two of four calc sets ( ~ 0) and now checks its fitted Sigma against the data's own mean intensity and falls back to the per-shell estimate; and the conformance report now gives the fraction of reflections with epsilon > 1, because "all axial" and "none axial" both produced a bare nan before and those are very different statements. Nothing in production consumes any of this yet -- the seam is not wired into FastRotationFunction, so peak lists are untouched. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/e_conformance.py | 238 +++++++++++++ alignment_lab/analysis/e_table.py | 80 +++++ alignment_lab/analysis/e_table.sh | 22 ++ alignment_lab/analysis/llg_decompose.py | 162 +++++++++ alignment_lab/analysis/llg_decompose.sh | 24 ++ alignment_lab/analysis/rescore_prep_arms.py | 109 ++++++ alignment_lab/analysis/rescore_prep_arms.sh | 22 ++ torchref/experimental/alignment/e_values.py | 360 ++++++++++++++++++++ 8 files changed, 1017 insertions(+) create mode 100644 alignment_lab/analysis/e_conformance.py create mode 100644 alignment_lab/analysis/e_table.py create mode 100644 alignment_lab/analysis/e_table.sh create mode 100644 alignment_lab/analysis/llg_decompose.py create mode 100644 alignment_lab/analysis/llg_decompose.sh create mode 100644 alignment_lab/analysis/rescore_prep_arms.py create mode 100644 alignment_lab/analysis/rescore_prep_arms.sh create mode 100644 torchref/experimental/alignment/e_values.py diff --git a/alignment_lab/analysis/e_conformance.py b/alignment_lab/analysis/e_conformance.py new file mode 100644 index 00000000..a2aaf1f5 --- /dev/null +++ b/alignment_lab/analysis/e_conformance.py @@ -0,0 +1,238 @@ +"""Does an E convention do what we want E to do? + +Phaser is no longer the specification for this part of the code, which means +comparison-debugging against a reference implementation is gone. What replaces it +is invariants: Wilson statistics, scale invariance, shrinkage monotonicity and +epsilon-correctness are true regardless of whose code computes them. This module +is that safety net, and it exists so conventions can be changed with something +other than an argument deciding whether the change was right. + +Eight checks, of which two are the ones nothing in the tree currently makes: + +* **absolute** unit mean, not merely a flat trend -- the rotation function is a + correlation and does not care, but the rescore's LLG compares an observation + against a predicted distribution and there is no free scale to cancel; +* **obs and calc on a common footing** -- the measured symptom of getting this + wrong is an expected moving-model intensity with mean 2.14 where ~1 belongs. + +The strongest check is not the mean but the **distribution**. Wilson statistics +predict ``|E|**2 ~ Exp(1)`` acentric and ``~ chi2_1`` centric at *every* +resolution, which catches a normaliser that is right on average and wrong in +shape. A mean cannot. + +Deliberately a report rather than a gate: a convention may fail a property and +still rank truth better, and in that case the property tells us what the winner +is trading away rather than vetoing it. +""" + +from __future__ import annotations + +import math +from typing import Optional + +import torch + +#: Wilson moment ratios /^2. Departures upward are the standard +#: twinning / tNCS diagnostic -- 2DQ6 reads about 5.5 acentric, which is how its +#: tNCS was originally identified. +IDEAL_MOMENT_RATIO = {"acentric": 2.0, "centric": 3.0} + + +def _deciles(s_mag: torch.Tensor, n: int = 10) -> torch.Tensor: + """Equal-count resolution deciles, as an index per reflection.""" + order = torch.argsort(s_mag) + out = torch.zeros_like(s_mag, dtype=torch.long) + chunk = max(1, s_mag.numel() // n) + for k in range(n): + lo = k * chunk + hi = (k + 1) * chunk if k < n - 1 else s_mag.numel() + out[order[lo:hi]] = k + return out + + +def _ks_uniform(sorted_u: torch.Tensor) -> float: + """One-sample KS statistic of ``sorted_u`` against Uniform(0, 1).""" + n = sorted_u.numel() + if n < 2: + return float("nan") + i = torch.arange(1, n + 1, dtype=torch.float64, device=sorted_u.device) + d_plus = (i / n - sorted_u).max() + d_minus = (sorted_u - (i - 1) / n).max() + return float(torch.maximum(d_plus, d_minus)) + + +def _wilson_ks(E2: torch.Tensor, centric: torch.Tensor) -> dict: + """KS of E**2 against its Wilson distribution, by centric class. + + Acentric ``E**2 ~ Exp(1)``, so ``1 - exp(-E**2)`` is uniform. Centric + ``E**2 ~ chi2_1``, so ``erf(sqrt(E**2 / 2))`` is uniform. Both CDFs are + closed-form, so no scipy dependency and no interpolation error. + """ + out = {} + for name, mask, cdf in ( + ("acentric", ~centric, lambda x: 1.0 - torch.exp(-x)), + ("centric", centric, lambda x: torch.erf(torch.sqrt(x * 0.5))), + ): + v = E2[mask].to(torch.float64) + if v.numel() < 50: + out[name] = float("nan") + continue + u = cdf(v.clamp(min=0.0)).clamp(0.0, 1.0) + out[name] = _ks_uniform(torch.sort(u).values) + return out + + +def check_e_convention( + cls, + F: torch.Tensor, + s_mag: torch.Tensor, + centric: torch.Tensor, + *, + sig_F: Optional[torch.Tensor] = None, + eps: Optional[torch.Tensor] = None, + n_shells: int = 20, + F_calc: Optional[torch.Tensor] = None, + n_deciles: int = 10, +) -> dict: + """Run every applicable property check on one convention. + + Takes the **class**, not an instance, because three of the checks need to + construct it again on perturbed inputs -- rescaled F, raised sigmas -- and a + convention that has already normalised its data cannot be asked those + questions. + """ + conv = cls(F, s_mag, centric, sig_F=sig_F, eps=eps, n_shells=n_shells) + E = conv.E.to(torch.float64) + E2 = E * E + cen = conv.centric + dec = _deciles(s_mag, n_deciles) + rep: dict = {"name": cls.__name__ if hasattr(cls, "__name__") else str(cls)} + + # (1) stationarity + (2) absolute unit mean. + rep["mean_E2"] = float(E2.mean()) + per_dec = torch.stack([ + E2[dec == k].mean() if bool((dec == k).any()) else torch.tensor(float("nan")) + for k in range(n_deciles) + ]) + rep["decile_mean_E2"] = [round(float(v), 4) for v in per_dec] + rep["max_decile_dev"] = float((per_dec - 1.0).abs().max()) + # A trend, not just scatter: correlation of the decile mean with resolution. + finite = torch.isfinite(per_dec) + if int(finite.sum()) > 2: + x = torch.arange(n_deciles, dtype=torch.float64)[finite] + y = per_dec[finite].to(torch.float64) + xc, yc = x - x.mean(), y - y.mean() + denom = (xc.norm() * yc.norm()).clamp(min=1e-30) + rep["decile_trend_r"] = float((xc * yc).sum() / denom) + else: + rep["decile_trend_r"] = float("nan") + + # (3) Wilson distribution, globally and worst-decile. + rep["ks"] = _wilson_ks(E2, cen) + worst = 0.0 + for k in range(n_deciles): + m = dec == k + if int(m.sum()) < 100: + continue + ks_k = _wilson_ks(E2[m], cen[m]) + for v in ks_k.values(): + if not math.isnan(v): + worst = max(worst, v) + rep["ks_worst_decile"] = worst + + # (4) moment ratios. + rep["moment_ratio"] = {} + for name, mask in (("acentric", ~cen), ("centric", cen)): + v = E2[mask] + if v.numel() < 50: + rep["moment_ratio"][name] = float("nan") + continue + rep["moment_ratio"][name] = float( + (v * v).mean() / v.mean().clamp(min=1e-30) ** 2 + ) + + # (5) scale invariance: E must not depend on the units of F. + devs = [] + for c in (1e-3, 1e3): + E_c = cls(F * c, s_mag, centric, + sig_F=None if sig_F is None else sig_F * c, + eps=eps, n_shells=n_shells).E.to(torch.float64) + scale = E.abs().max().clamp(min=1e-30) + devs.append(float((E_c - E).abs().max() / scale)) + rep["scale_invariance_dev"] = max(devs) + + # (6) epsilon-correctness: axial reflections must not sit systematically + # above general ones. Only meaningful when eps was supplied and varies. + if eps is not None: + axial = eps > 1.0 + rep["eps_frac_gt1"] = float(axial.to(torch.float64).mean()) + if bool(axial.any()) and bool((~axial).any()): + rep["eps_ratio"] = float( + E2[axial].mean() / E2[~axial].mean().clamp(min=1e-30) + ) + else: + # All or none: no contrast to measure. All-axial means epsilon is + # counting something it should not -- on a centred lattice the + # rotation part repeats per centring coset, so `h.W == h` matches + # once per coset for EVERY reflection. + rep["eps_ratio"] = float("nan") + else: + rep["eps_frac_gt1"] = float("nan") + rep["eps_ratio"] = float("nan") + + # (7) shrinkage monotonicity: raising sigma_F at fixed F must move E toward + # the shell mean, never away. Only conventions that read sig_F can. + if getattr(cls, "uses_sigma_f", False) and sig_F is not None: + loud = cls(F, s_mag, centric, sig_F=sig_F * 4.0, eps=eps, + n_shells=n_shells).E.to(torch.float64) + ref = conv.sigma.sqrt().to(torch.float64) # the shell scale E sits on + moved_closer = (loud - ref).abs() <= (E - ref).abs() + 1e-9 + rep["shrinkage_frac_ok"] = float(moved_closer.to(torch.float64).mean()) + else: + rep["shrinkage_frac_ok"] = float("nan") + + # (8) obs/calc common footing: normalise a calc set with the same convention + # and compare mean E**2. Both should sit at 1 if the convention puts them + # on a common scale; a mismatch is the eImove defect in miniature. + if F_calc is not None: + # A sigma_F-consuming convention cannot normalise a calc set -- there is + # no measurement error to shrink -- so ask it which companion to use. + calc_cls = cls.for_calc() if hasattr(cls, "for_calc") else cls + rep["calc_via"] = calc_cls.__name__ + E_calc = calc_cls(F_calc, s_mag, centric, sig_F=None, eps=eps, + n_shells=n_shells).E.to(torch.float64) + rep["mean_E2_calc"] = float((E_calc * E_calc).mean()) + rep["obs_calc_ratio"] = rep["mean_E2_calc"] / max(rep["mean_E2"], 1e-30) + else: + rep["calc_via"] = "-" + rep["obs_calc_ratio"] = float("nan") + return rep + + +def format_table(reports) -> str: + """One row per convention, the columns that decide things.""" + head = (f"{'convention':18s} {'':>7s} {'maxdec':>7s} {'trend':>6s} " + f"{'KS ac':>6s} {'KS cen':>7s} {'KSdec':>6s} {'m2 ac':>6s} " + f"{'m2 cen':>7s} {'scale':>8s} {'e>1':>6s} {'eps':>6s} " + f"{'shrink':>7s} {'o/c':>6s} {'calc via':>16s}") + lines = [head, "-" * len(head)] + for r in reports: + lines.append( + f"{r['name']:18s} {r['mean_E2']:>7.4f} {r['max_decile_dev']:>7.4f} " + f"{r['decile_trend_r']:>+6.2f} " + f"{r['ks']['acentric']:>6.4f} {r['ks']['centric']:>7.4f} " + f"{r['ks_worst_decile']:>6.4f} " + f"{r['moment_ratio']['acentric']:>6.3f} " + f"{r['moment_ratio']['centric']:>7.3f} " + f"{r['scale_invariance_dev']:>8.1e} " + f"{r.get('eps_frac_gt1', float('nan')):>6.3f} " + f"{r['eps_ratio']:>6.3f} " + f"{r['shrinkage_frac_ok']:>7.3f} {r['obs_calc_ratio']:>6.3f} " + f"{r.get('calc_via', '-'):>16s}" + ) + lines.append("") + lines.append("ideal: =1 maxdec=0 trend=0 KS small m2 ac=2 cen=3 " + "scale=0 eps=1 shrink=1 o/c=1") + lines.append("e>1 = fraction with epsilon>1; ~1.0 means epsilon is counting " + "centring cosets, not point-group stabilisers") + return "\n".join(lines) diff --git a/alignment_lab/analysis/e_table.py b/alignment_lab/analysis/e_table.py new file mode 100644 index 00000000..b3ba9a93 --- /dev/null +++ b/alignment_lab/analysis/e_table.py @@ -0,0 +1,80 @@ +"""Conformance table over every E convention the alignment package uses. + +Establishes what we currently have, before anything changes. Real benchmark data +rather than synthetic, because the properties that matter (the Wilson shape, the +epsilon behaviour on high-symmetry lattices) are properties of real reflection +sets. +""" +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +sys.path.insert(0, str(Path(__file__).resolve().parent)) +torch.set_grad_enabled(False) + +from e_conformance import check_e_convention, format_table # noqa: E402 +from lab import BENCH_PDBS, load_case # noqa: E402 + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) + ap.add_argument("--n-shells", type=int, default=20) + args = ap.parse_args() + + from torchref.experimental.alignment.e_values import ( + CalcGlobalE, CalcShellE, FrenchWilsonE, SmoothSigmaE, WilsonShellE, + WilsonShellEpsE, + ) + from torchref.experimental.alignment.frf.preprocessing import compute_epsilon + + model, data = load_case(args.pdb)[:2] + rec = data.cell.reciprocal_basis_matrix.to(torch.float64) + s_mag = (data.hkl.to(torch.float64) @ rec).norm(dim=-1) + F = data.F.to(torch.float64).abs() + sig = None if getattr(data, "F_sigma", None) is None else \ + data.F_sigma.to(torch.float64) + cen = data.centric.to(torch.bool) + eps = compute_epsilon(data.hkl.to(torch.long), + data.spacegroup.matrices.to(torch.float64)) + + # A calc set from the deposited coordinates: the "perfect model" case, where + # obs and calc genuinely should land on the same scale. + with torch.no_grad(): + F_calc = model.get_structure_factor( + data.hkl, recalc=True).abs().to(torch.float64) + + finite = torch.isfinite(F) & torch.isfinite(F_calc) & (F > 0) + if sig is not None: + finite &= torch.isfinite(sig) & (sig > 0) + F, s_mag, cen, eps, F_calc = (F[finite], s_mag[finite], cen[finite], + eps[finite], F_calc[finite]) + sig = None if sig is None else sig[finite] + + uniq, cnt = torch.unique(eps, return_counts=True) + n_ops = int(data.spacegroup.matrices.shape[0]) + print(f"=== {args.pdb} {data.spacegroup.hm} N={int(finite.sum())} " + f"(dropped {int((~finite).sum())}) n_shells={args.n_shells} ===") + print(f" n_ops={n_ops} epsilon: " + + ", ".join(f"{float(u):g}x{int(c)}" for u, c in zip(uniq, cnt))) + reports = [] + for cls in (FrenchWilsonE, WilsonShellE, WilsonShellEpsE, CalcShellE, + CalcGlobalE, SmoothSigmaE): + try: + reports.append(check_e_convention( + cls, F, s_mag, cen, sig_F=sig, eps=eps, + n_shells=args.n_shells, F_calc=F_calc, + )) + except Exception as exc: # noqa: BLE001 - report it + print(f" {cls.__name__}: RAISED {type(exc).__name__}: {exc}") + print(format_table(reports)) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/alignment_lab/analysis/e_table.sh b/alignment_lab/analysis/e_table.sh new file mode 100644 index 00000000..2abaa129 --- /dev/null +++ b/alignment_lab/analysis/e_table.sh @@ -0,0 +1,22 @@ +#!/bin/bash +#SBATCH --job-name=etable +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:50:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname)" +for pdb in 1DAW 3K7M 6G9X 2DQ6; do + "$PY" -u alignment_lab/analysis/e_table.py --pdb "$pdb" 2>&1 \ + | grep -vE "UserWarning|FutureWarning|^ *from |^ *warnings\.|^Loaded |^LINK |^Wilson outlier|^found non|^FrenchWilson initialized|^ Reflections:|^ Resolution:|^ Space group|^ Centric:|^✓|^Parametrization|^French-Wilson input guard" + echo +done +echo "rc=$?" diff --git a/alignment_lab/analysis/llg_decompose.py b/alignment_lab/analysis/llg_decompose.py new file mode 100644 index 00000000..77525409 --- /dev/null +++ b/alignment_lab/analysis/llg_decompose.py @@ -0,0 +1,162 @@ +"""Why does the true orientation lose the LLG contest? + +The rescore's job is to take a top-20 that already contains truth and put truth +first. On 6G9X it reproducibly does the opposite -- truth goes from FRF rank 1 to +rank 12-17 -- and that survives the full Phaser model preparation, so it is not +explained by sigma_A or the solvent term. + +This takes one case with known ground truth and asks where the LLG difference +between truth and the candidate that beats it actually accumulates: per +resolution shell, and split by centric/acentric. A likelihood that prefers the +wrong orientation is either being fed the wrong expected intensity or is summing +a term whose sign is wrong somewhere, and both of those localise. + +Per-reflection LL is recomputed here rather than taken from +``_llg_for_orientations``, which sums before returning -- same context, same +``phaser_log_rel_*`` calls, just not reduced. +""" + +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, FRFConfig, orbit_rank, rotated_case, # noqa: E402 + run_frf, seed_for) + + +def per_reflection_ll(ctx, alpha, beta, gamma): + """``(n_orient, N)`` per-reflection log-likelihood -- the unsummed LLG.""" + from torchref.experimental.alignment.distributions import ( + phaser_log_rel_rice, phaser_log_rel_woolfson, + ) + from torchref.experimental.alignment.frf.rotation_utils import ( + rotation_matrix_from_edmonds_euler_batch, + ) + R = rotation_matrix_from_edmonds_euler_batch( + alpha.to(torch.float64), beta.to(torch.float64), gamma.to(torch.float64), + ).transpose(-1, -2).to(torch.float32) + F_calc_m = ctx.interpolator.evaluate( + R, ctx.unrolled_hkl, ctx.real_cell, return_amplitude=True, + ).to(ctx.dtype) + if ctx.dw_per_m is not None: + F_calc_m = F_calc_m * ctx.dw_per_m.unsqueeze(0) + E_calc_m = F_calc_m / ctx.sqrt_mean_per_m.unsqueeze(0) + Esq_m = E_calc_m * E_calc_m + B = Esq_m.shape[0] + sum_per_h = torch.zeros(B, ctx.N, dtype=Esq_m.dtype, device=Esq_m.device) + sum_per_h.scatter_add_(1, ctx.asu_idx.unsqueeze(0).expand(B, -1), Esq_m) + eImove = ctx.eImove_prefac * sum_per_h + sqrt_eImove = eImove.clamp(min=1e-30).sqrt() + ll = torch.where( + ctx.centric_b, + phaser_log_rel_woolfson(ctx.E_obs_b, sqrt_eImove, ctx.V_b), + phaser_log_rel_rice(ctx.E_obs_b, sqrt_eImove, ctx.V_b), + ) + return ll, eImove + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--pdb", default="6G9X", choices=list(BENCH_PDBS)) + ap.add_argument("--trial", type=int, default=0) + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--n-refine", type=int, default=20) + ap.add_argument("--n-shells", type=int, default=10) + ap.add_argument("--full-prep", action="store_true", + help="turn on the Phaser model prep the pipeline omits") + args = ap.parse_args() + + from torchref.experimental.alignment.ml_rotation import _build_llg_context + + seed = seed_for(args.pdb, args.trial) + model, data, R_true = rotated_case(args.pdb, seed) + sym = data.spacegroup.matrices.to(torch.float64).cpu() + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + okw = dict(side="left", frame="cart", reciprocal_basis=rec, thr_deg=5.0) + + res = run_frf(model, data, FRFConfig(n_peaks=500, lmax_cap=args.lmax_cap), + capture_arf=False, verbose=0) + head = res.peaks[: args.n_refine] + truth_rank, truth_ang = orbit_rank(head, R_true, sym, **okw) + if truth_rank < 0: + print(f"truth not in the top {args.n_refine}; nothing for the rescore " + f"to find here") + return 0 + + inp = res.inputs + prep = {} + if args.full_prep: + prep = dict(vrms_strategy="oeffner", + vrms_n_residues=max(1, int(model.xyz().shape[0] / 8)), + apply_bulk_solvent=True, apply_wilson_b=True) + ctx = _build_llg_context( + inp.F_obs, inp.hkl, inp.s_mag, inp.centric, inp.ll, data.cell, + data.spacegroup.matrices.to(torch.float64).to(inp.device), + n_shells=max(20 // 2, 8), batch_size=50, **prep, + ) + + a = torch.tensor([p.alpha for p in head], dtype=torch.float64) + b = torch.tensor([p.beta for p in head], dtype=torch.float64) + g = torch.tensor([p.gamma for p in head], dtype=torch.float64) + ll, eImove = per_reflection_ll(ctx, a, b, g) + totals = ll.sum(dim=-1) + order = torch.argsort(totals, descending=True) + new_rank = int((order == truth_rank).nonzero()[0, 0]) + winner = int(order[0]) + + print(f"=== {args.pdb} trial {args.trial} " + f"({'full prep' if args.full_prep else 'shipped defaults'}) ===") + print(f" truth is FRF rank {truth_rank} ({truth_ang:.2f} deg), " + f"LLG rank {new_rank}; winner is FRF rank {winner}") + print(f" LLG(truth) = {totals[truth_rank]:.4f}") + print(f" LLG(winner) = {totals[winner]:.4f} " + f"gap = {totals[winner] - totals[truth_rank]:+.4f}") + if winner == truth_rank: + print(" truth already wins here") + return 0 + + # Where does the gap accumulate? Equal-count shells in |s|. + s = inp.s_mag.to(torch.float64).cpu() + n_sh = args.n_shells + edge_idx = torch.linspace(0, s.numel() - 1, n_sh + 1).round().long() + edges = s.sort().values[edge_idx] + shell = torch.bucketize(s, edges[1:-1]) + d_ll = (ll[winner] - ll[truth_rank]).to(torch.float64).cpu() + cen = ctx.centric_b[0].cpu() + + print(f"\n gap by resolution shell (winner - truth; positive = truth loses)") + print(f" {'shell':>5s} {'d range (A)':>16s} {'n':>7s} {'d LL':>10s} " + f"{'cum %':>7s} {'acen':>9s} {'cen':>9s}") + total_gap = float(d_ll.sum()) + cum = 0.0 + for k in range(n_sh): + m = shell == k + if not bool(m.any()): + continue + v = float(d_ll[m].sum()) + cum += v + lo, hi = float(1.0 / edges[k + 1]), float(1.0 / edges[k]) + print(f" {k:>5d} {f'{hi:6.1f}-{lo:5.2f}':>16s} {int(m.sum()):>7d} " + f"{v:>+10.3f} {100 * cum / total_gap if total_gap else 0:>6.1f}% " + f"{float(d_ll[m & ~cen].sum()):>+9.3f} " + f"{float(d_ll[m & cen].sum()):>+9.3f}") + print(f" {'TOTAL':>5s} {'':>16s} {int(s.numel()):>7d} {total_gap:>+10.3f}") + + print(f"\n expected moving intensity eImove, truth vs winner:") + for name, idx in (("truth", truth_rank), ("winner", winner)): + e = eImove[idx].to(torch.float64).cpu() + print(f" {name:7s} mean {e.mean():.4e} median {e.median():.4e} " + f"max {e.max():.4e} frac>E_obs^2 " + f"{float((e > (ctx.E_obs_b[0].cpu() ** 2)).float().mean()):.3f}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/alignment_lab/analysis/llg_decompose.sh b/alignment_lab/analysis/llg_decompose.sh new file mode 100644 index 00000000..d206c26f --- /dev/null +++ b/alignment_lab/analysis/llg_decompose.sh @@ -0,0 +1,24 @@ +#!/bin/bash +#SBATCH --job-name=llgdec +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:50:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname)" +for trial in 0 1; do + for flag in "" "--full-prep"; do + "$PY" -u alignment_lab/analysis/llg_decompose.py --pdb 6G9X --trial $trial $flag 2>&1 \ + | grep -vE "UserWarning|FutureWarning|^ *from |^ *warnings\.|Loaded|LINK|Wilson outlier|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization" + echo + done +done +echo "rc=$?" diff --git a/alignment_lab/analysis/rescore_prep_arms.py b/alignment_lab/analysis/rescore_prep_arms.py new file mode 100644 index 00000000..a2a7c5f6 --- /dev/null +++ b/alignment_lab/analysis/rescore_prep_arms.py @@ -0,0 +1,109 @@ +"""Does the ML rescore fail because of its MODEL PREPARATION? + +The rescore is measurably destructive end-to-end -- dropping it takes pose +recovery from 18/30 to 24/30 -- and the damage is concentrated in the two large +benchmark entries, 4BX9 and 6G9X, which it fails at 52-91 degrees and which the +raw FRF order solves at 3-4 degrees. This harness tests one explanation. + +**The FRF and the rescore disagree about the model, inside a single run.** The +pipeline estimates ``model_error_A`` from the model's length, hands it to the +FRF, and applies Phaser's Babinet bulk-solvent term there -- then calls +``m_letf1_rescore`` without passing either, so the rescore falls back to a +hardcoded ``delta_vrms = 0.5`` and no solvent. Every Phaser model-prep knob on +that function defaults OFF and the pipeline overrides none of them. + +That predicts the observed failure pattern rather than merely being consistent +with it. Babinet's ``1 - 0.95 exp(-300 s^2/4)`` tends to 0.05 as ``s -> 0``, so +omitting it over-weights the lowest-resolution reflections by up to ~20x, and +low-resolution terms dominate for large molecules. 0.5 A is also furthest from +the truth for a large model, where ``oeffner_vrms`` gives ~0.67 A. + +Paired by construction: the FRF runs ONCE per (structure, trial) and every arm +rescores the *same* peak list. The ``none`` arm keeps the FRF order and is the +load-bearing control -- without it, an engine that merely preserves a good input +ranking looks like one that improves it. +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, FRFConfig, orbit_rank, rotated_case, # noqa: E402 + run_frf, run_rescore, seed_for) + +#: Each arm is a set of overrides on top of `m_letf1`'s defaults. They are +#: cumulative on purpose: if the whole Phaser prep helps, the interesting +#: question is which piece carries it. +ARMS = { + "none": None, # control: FRF order + "default": {}, # what ships today + "vrms": {"vrms_strategy": "oeffner"}, + "solvent": {"apply_bulk_solvent": True}, + "vrms_solvent": {"vrms_strategy": "oeffner", "apply_bulk_solvent": True}, + "full_prep": {"vrms_strategy": "oeffner", "apply_bulk_solvent": True, + "apply_wilson_b": True}, +} + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) + ap.add_argument("--trials", type=int, default=3) + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--n-refine", type=int, default=20, + help="rescore window: the top-N FRF peaks handed to the engine") + ap.add_argument("--thr-deg", type=float, default=5.0) + ap.add_argument("--arms", default=",".join(ARMS)) + args = ap.parse_args() + + cfg = FRFConfig(n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) + arms = [a for a in args.arms.split(",") if a] + + for trial in range(args.trials): + seed = seed_for(args.pdb, trial) + model, data, R_true = rotated_case(args.pdb, seed) + sym = data.spacegroup.matrices.to(torch.float64).cpu() + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + orbit_kw = dict(side="left", frame="cart", reciprocal_basis=rec, + thr_deg=args.thr_deg) + + # The FRF runs once; every arm sees the identical peak list. + res = run_frf(model, data, cfg, capture_arf=False, verbose=0) + frf_rank, _ = orbit_rank(res.peaks[: args.n_refine], R_true, sym, **orbit_kw) + + # Same residue estimate the pipeline uses for the FRF, so the `vrms` arm + # is genuinely "what the FRF was told" and not a second guess. + n_residues = max(1, int(model.xyz().shape[0] / 8)) + + for arm in arms: + overrides = ARMS[arm] + if overrides is None: + rank, seconds = frf_rank, 0.0 + else: + kw = dict(overrides) + if kw.get("vrms_strategy") == "oeffner": + kw["vrms_n_residues"] = n_residues + t0 = time.time() + rr = run_rescore(res.peaks, data, res.inputs, engine="m_letf1", + n_refine=args.n_refine, verbose=0, **kw) + seconds = time.time() - t0 + rank, _ = orbit_rank(rr.peaks, R_true, sym, **orbit_kw) + # A miss must not sort as a good rank. + rank_cmp = rank if rank >= 0 else args.n_refine + print(f"ROW {arm} {args.pdb} trial={trial} seed={seed} " + f"frf_rank={frf_rank} rank={rank} rank_cmp={rank_cmp} " + f"n_res={n_residues} seconds={seconds:.2f}", flush=True) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/alignment_lab/analysis/rescore_prep_arms.sh b/alignment_lab/analysis/rescore_prep_arms.sh new file mode 100644 index 00000000..06716660 --- /dev/null +++ b/alignment_lab/analysis/rescore_prep_arms.sh @@ -0,0 +1,22 @@ +#!/bin/bash +#SBATCH --job-name=resprep +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-9 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname) pdb=$PDB" +"$PY" -u alignment_lab/analysis/rescore_prep_arms.py --pdb "$PDB" --trials 3 2>&1 \ + | grep -E "^ROW |Error|Traceback|Warning: " +echo "rc=$?" diff --git a/torchref/experimental/alignment/e_values.py b/torchref/experimental/alignment/e_values.py new file mode 100644 index 00000000..7ecb1be2 --- /dev/null +++ b/torchref/experimental/alignment/e_values.py @@ -0,0 +1,360 @@ +"""One place to say what ``E`` means. + +``E = F / sqrt(Sigma(s))`` is a **weighting choice wearing the costume of a units +change**: correlating ``E_obs`` against ``E_calc`` *is* correlating ``F`` against +``F`` with weight ``1/Sigma(s)``. The alignment package currently answers that +question nine different times -- twice in ``frf.preprocessing``, once inside +``french_wilson_preprocess``, three times in ``ml_rotation`` and three more in +``translation`` -- and the answers disagree. The rotation function's observed +side is a French-Wilson posterior weighted by ``DFAC**2``; the rescore's is plain +per-shell Wilson with epsilon divided out. So the rescore ranks candidates +against a differently-normalised observation set than the one that produced them. + +The two consumers do not need the same thing from it, which is worth stating +because it explains which of them breaks: + +* The rotation function is a **correlation**. A global scale cancels -- it scales + every SO(3) sample equally and the peak is reported as a z-score -- so only the + *relative* weighting across resolution matters. Removing the antipodal copy + scaled every score by exactly 4 and moved 98 of 100 truth ranks not at all. +* The rescore's LLG is a **likelihood**. It compares an observation against a + predicted distribution, so there is no free scale to cancel; get it wrong and + you evaluate the right data against the wrong Rice. + +A convention that satisfies the likelihood satisfies the correlation for free, +so the strict requirement is the one to design against. + +Conventions are constructed **from the data** rather than configured and passed +in, because a fitted ``Sigma(s)`` cannot exist before the reflections do. Engines +therefore take the class and instantiate it internally:: + + FastRotationFunction(..., e_convention=FrenchWilsonE) + +and anything needing configuration rides in as +``functools.partial(SmoothSigmaE, n_coeff=6)``, which is class-like and needs no +extra parameter. +""" + +from __future__ import annotations + +import math +from typing import Optional + +import torch + +from .sh import assign_shells, equal_count_shell_edges + +__all__ = [ + "CalcGlobalE", + "CalcShellE", + "EConvention", + "FrenchWilsonE", + "SmoothSigmaE", + "WilsonShellE", + "WilsonShellEpsE", +] + + +class EConvention: + """Normalised amplitudes, plus the per-reflection weight that goes with them. + + Attributes + ---------- + E : torch.Tensor + ``(N,)`` normalised amplitude. For most conventions this is + ``F / sqrt(sigma)``; for :class:`FrenchWilsonE` it is a posterior + expectation and the relation is only approximate, which is exactly the + difference the conformance harness is there to expose. + weight : torch.Tensor + ``(N,)`` per-reflection information weight. ``DFAC**2`` where the + convention models measurement error, ones where it does not. This is the + "weight by F/sigma" lever, made explicit rather than left implicit in + whichever normaliser a caller happened to pick. + sigma : torch.Tensor + ``(N,)`` the normaliser actually used, ```` per reflection. + Reported for every convention so they are comparable even when their + ``E`` is not defined the same way. + + Notes + ----- + ``eps`` divides the intensity before averaging (``E**2 = (F**2/eps) / + ``) because axial reflections are systematically stronger -- + `` = eps * Sigma`` -- and would otherwise dominate. Conventions that leave + it ``None`` are declaring that their caller handles multiplicity some other + way; the rotation function does, by unrolling the full orbit. + """ + + #: Whether this convention consumes ``sig_F``. The conformance harness skips + #: the shrinkage test for conventions that do not. + uses_sigma_f: bool = False + + #: The convention to use for **calculated** amplitudes, when it cannot be + #: this one. A French-Wilson posterior is defined for observations only -- + #: there is no measurement error on a calc set to shrink toward the mean -- + #: so a sigma_F-consuming convention has to name a companion. ``None`` means + #: "use this class for both sides", which is what most of them do. + #: + #: This is not a harness convenience: it is why the rotation function pairs + #: `french_wilson_preprocess` on obs with `wilson_normalise` on calc. + calc_companion: Optional[type] = None + + @classmethod + def for_calc(cls) -> type: + """The class to normalise calculated amplitudes with.""" + return cls.calc_companion or cls + + def __init__( + self, + F: torch.Tensor, + s_mag: torch.Tensor, + centric: Optional[torch.Tensor] = None, + *, + sig_F: Optional[torch.Tensor] = None, + eps: Optional[torch.Tensor] = None, + shell_idx: Optional[torch.Tensor] = None, + n_shells: int = 20, + ) -> None: + if F.ndim != 1: + raise ValueError(f"F must be 1-D, got {tuple(F.shape)}") + if s_mag.shape != F.shape: + raise ValueError( + f"s_mag {tuple(s_mag.shape)} does not match F {tuple(F.shape)}" + ) + self.F = F + self.s_mag = s_mag + self.centric = ( + torch.zeros_like(F, dtype=torch.bool) if centric is None + else centric.to(torch.bool) + ) + self.sig_F = sig_F + self.eps = eps + self.n_shells = int(n_shells) + # One shell assignment, shared by whatever the subclass needs it for. + # Assigning here rather than in each subclass is the same fix the FRF + # needed: two consumers deriving their own equal-count edges from the + # same |s| disagree about the reflections sitting on a boundary. + if shell_idx is None: + edges, _ = equal_count_shell_edges(s_mag, self.n_shells) + shell_idx = assign_shells(s_mag, edges) + self.shell_idx = shell_idx.clamp(min=0) + + self.sigma = self._shell_mean_intensity() + self.E, self.weight = self._compute() + + # -- helpers shared by the subclasses --------------------------------- + + def _intensity(self) -> torch.Tensor: + """``F**2 / eps`` -- the quantity whose shell mean is ``Sigma``.""" + I = self.F * self.F + if self.eps is not None: + I = I / self.eps.clamp(min=1.0) + return I + + def _shell_mean_intensity(self) -> torch.Tensor: + """```` per reflection, from the shared shell assignment.""" + I = self._intensity() + total = torch.zeros(self.n_shells, dtype=I.dtype, device=I.device) + total.scatter_add_(0, self.shell_idx, I) + count = torch.bincount( + self.shell_idx, minlength=self.n_shells, + ).to(I.dtype).clamp(min=1.0) + return (total / count).clamp(min=1e-30).index_select(0, self.shell_idx) + + def _ones(self) -> torch.Tensor: + return torch.ones_like(self.F) + + def _compute(self): + raise NotImplementedError + + def __repr__(self) -> str: # pragma: no cover - display + return f"{type(self).__name__}(N={self.F.numel()}, n_shells={self.n_shells})" + + +class WilsonShellE(EConvention): + """Plain per-shell Wilson: ``E = F / sqrt(_shell)``. + + What the rotation function uses on the calc side, and on the obs side when + the data carry no sigmas. Ignores measurement error entirely. + """ + + def _compute(self): + return self.F / self.sigma.sqrt(), self._ones() + + +class WilsonShellEpsE(EConvention): + """Epsilon-corrected Wilson, ``E**2 = (F**2/eps) / _shell``. + + What the m_LETF1 rescore uses on the observed side. Identical to + :class:`WilsonShellE` when ``eps`` is absent, which is worth knowing: the + difference between the two is only ever the multiplicity handling. + """ + + def _compute(self): + E = (self._intensity() / self.sigma).clamp(min=0.0).sqrt() + return E, self._ones() + + +class FrenchWilsonE(EConvention): + """French-Wilson posterior amplitude with the Rice ``DFAC`` weight. + + The rotation function's observed-side default, and the only convention here + that looks at ``sig_F``. ``E`` is the posterior expectation of the true + normalised amplitude given a noisy measurement, so weak reflections shrink + toward the shell mean instead of being taken at face value; ``weight`` is + ``DFAC**2``, the Rice-moment D factor, which is the per-reflection + measurement-information term. + + Requires ``sig_F``. Falling back silently to plain Wilson would hide exactly + the difference this class exists to make visible. + """ + + uses_sigma_f = True + calc_companion = WilsonShellE + + def _compute(self): + if self.sig_F is None: + raise ValueError( + "FrenchWilsonE needs sig_F; use WilsonShellE for data without " + "measurement errors rather than letting the difference pass " + "silently." + ) + from .frf.french_wilson import french_wilson_preprocess + + fw = french_wilson_preprocess( + self.F, self.sig_F, self.s_mag, self.centric, + n_wilson_shells=self.n_shells, shell_idx=self.shell_idx, + ) + dfac = fw["DFAC"] + return fw["eEobs"], dfac * dfac + + +class CalcShellE(WilsonShellE): + """The rescore's calc-side normaliser: per-shell, flattening every shell to 1. + + Named separately from :class:`WilsonShellE` because it is applied to a + *reference* orientation's ``|F_calc|`` and then reused for every rotated + candidate. Predicted to fail the obs/calc common-scale check: forcing + ``_shell = 1`` in every shell discards the model's inter-shell + amplitude shape, which is the very thing the likelihood's expected intensity + is supposed to carry. + """ + + +class CalcGlobalE(EConvention): + """Single global scale: ``E = F / rms(F)``, preserving inter-shell shape. + + The rescore's ``scat_mode="absolute"``. Keeps how much the model actually + scatters per resolution instead of flattening it, which is what makes a + relative Wilson-B correction meaningful rather than cancelled. + """ + + def _compute(self): + rms = self._intensity().mean().clamp(min=1e-30).sqrt() + self.sigma = torch.full_like(self.F, float(rms * rms)) + return self.F / rms, self._ones() + + +class SmoothSigmaE(EConvention): + """``Sigma(s)`` as a smooth curve rather than a step function over shells. + + Per-shell ``Sigma`` is a noisy non-parametric estimate with edges, and the + edges are not free: two consumers binning the same ``|s|`` independently + disagreed about 7 of 55078 reflections on 3K7M. A smooth curve has no edges, + is the same function whichever subset it is evaluated on, and is what the + scaler already uses for the closely-related isotropic scale. + + Basis follows ``scaling/scaler_base.py::_build_iso_design`` -- Chebyshev in + ``sin(theta)/lambda`` mapped onto ``[-1, 1]``, evaluated per reflection, in + log space. That abscissa rather than ``s**2`` because the modulation is + gentle through the bulk of the range and has real structure in the first few + percent of ``s**2``. + + Fitted as a **Gamma GLM with a log link**, which is the right likelihood + rather than a convenience: acentric ``F**2`` is exponentially distributed + with mean ``Sigma``, i.e. Gamma with unit shape, and centric ``F**2`` is + Gamma with shape 1/2. Fitting the *mean* this way avoids the trap that a + regression on ``log F**2`` walks into -- the ``E[log chi**2]`` offset has to + go somewhere, and with no intercept it is absorbed into the shape of the + curve. That is precisely how the overall-anisotropy fit was biased. + + Coefficients are clamped in log space for the reason the scaler clamps: a + polynomial is unbounded at the ends of its interval, and the low-resolution + end is where a mis-specified normaliser does its damage. + """ + + #: Chebyshev terms. Six is the scaler's default and spans a Wilson plot's + #: curvature without chasing shell-to-shell noise. + DEFAULT_N_COEFF = 6 + + #: Log-space clamp on the fitted curve, as a factor either side of the + #: global mean intensity. Wide enough never to bind on real data; present so + #: an extrapolating polynomial cannot produce an arbitrary scale. + LOG_CLAMP = 10.0 + + def __init__(self, *args, n_coeff: int = DEFAULT_N_COEFF, + n_iter: int = 8, **kwargs) -> None: + self.n_coeff = int(n_coeff) + self.n_iter = int(n_iter) + super().__init__(*args, **kwargs) + + def _design(self) -> torch.Tensor: + """``(N, n_coeff)`` Chebyshev design in sin(theta)/lambda.""" + x = (self.s_mag * 0.5).clamp(min=0.0) + lo, hi = x.min(), x.max() + u = (2 * (x - lo) / (hi - lo).clamp(min=1e-12) - 1).clamp(-1.0, 1.0) + cols = [torch.ones_like(u), u] + for _ in range(2, self.n_coeff): + cols.append(2 * u * cols[-1] - cols[-2]) + return torch.stack(cols[: self.n_coeff], dim=1) + + def _fit_log_sigma(self) -> torch.Tensor: + """IRLS for a Gamma GLM with log link; returns log Sigma per reflection.""" + X = self._design().to(torch.float64) + y = self._intensity().to(torch.float64).clamp(min=1e-30) + # Gamma shape: 1 acentric (exponential), 1/2 centric. Used as the IRLS + # weight, so better-determined reflections pull harder. + w = torch.where(self.centric, 0.5, 1.0).to(torch.float64) + + # Seed at the global mean, i.e. the constant curve a single Wilson + # scale would give. Every later iteration only adds shape. + beta = torch.zeros(self.n_coeff, dtype=torch.float64, device=X.device) + beta[0] = torch.log(y.mean().clamp(min=1e-30)) + for _ in range(self.n_iter): + eta = (X @ beta).clamp(-self.LOG_CLAMP + float(beta[0]), + self.LOG_CLAMP + float(beta[0])) + mu = torch.exp(eta) + # Log link with Gamma variance: the working response is + # eta + (y - mu)/mu and the IRLS weight is constant in mu. + z = eta + (y - mu) / mu.clamp(min=1e-30) + XtW = X.transpose(0, 1) * w.unsqueeze(0) + A = XtW @ X + A = A + torch.eye( + self.n_coeff, dtype=A.dtype, device=A.device, + ) * 1e-10 * float(torch.diagonal(A).abs().max().clamp(min=1e-30)) + beta_new = torch.linalg.solve(A, XtW @ z) + if torch.allclose(beta_new, beta, rtol=1e-10, atol=1e-12): + beta = beta_new + break + beta = beta_new + return (X @ beta).clamp( + -self.LOG_CLAMP + float(beta[0]), self.LOG_CLAMP + float(beta[0]), + ) + + def _compute(self): + shell_sigma = self.sigma # the per-shell fallback + log_sigma = self._fit_log_sigma() + sigma = torch.exp(log_sigma).to(self.F.dtype).clamp(min=1e-30) + # Sanity: a fitted Sigma(s) must reproduce the data's own mean intensity. + # A Gamma GLM on a Chebyshev basis can diverge when the calc amplitudes + # span a huge dynamic range with near-zeros, and it did -- two of four + # calc sets came back with ~ 0, i.e. Sigma inflated by orders of + # magnitude. Detect that against the quantity the fit is estimating and + # fall back to the per-shell estimate rather than returning nonsense. + mean_I = self._intensity().mean().clamp(min=1e-30) + ratio = float((sigma.mean() / mean_I).clamp(min=1e-30)) + self.converged = 0.2 < ratio < 5.0 + if not self.converged: + sigma = shell_sigma + self.sigma = sigma + E = (self._intensity() / self.sigma).clamp(min=0.0).sqrt() + return E, self._ones() From f482f3353f765fd87e20634850f6447cfd28ca11 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 17:15:02 +0200 Subject: [PATCH 086/250] Aggregate the movement-recovery sweep paired, seed by seed Every arm refines the same simulated dataset within a seed, so the noise realisation is common to them and cancels in the difference. Comparing two distributions of errors instead would drown the effect in variance that is not there. Reports the median paired difference with a bootstrap CI, which survives a skewed distribution and a handful of seeds. Also stops the collinearity probe retaining an autograd graph per numerical derivative, which is what pushed it past its memory limit. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- paper/probe_movement_recovery.py | 103 +++++++++++++++++++++++++ paper/probe_two_moment_collinearity.py | 12 +-- 2 files changed, 110 insertions(+), 5 deletions(-) diff --git a/paper/probe_movement_recovery.py b/paper/probe_movement_recovery.py index b5f0c039..7daaa935 100644 --- a/paper/probe_movement_recovery.py +++ b/paper/probe_movement_recovery.py @@ -172,6 +172,103 @@ def positions(path): return rows +def score_sweep(root, chain, first, last, start_pdb, n_boot=10000, seed=0): + """Aggregate a seed sweep, paired seed by seed. + + Paired, not pooled: every arm refines the *same* simulated dataset within a seed, so + the seed-to-seed spread of the noise realisation is common to all arms and cancels in + the difference. Comparing two distributions of errors instead would drown a real + effect in variance that is not there. + + Reports the median paired difference with a bootstrap CI, which is the shape that + survives a skewed distribution and a handful of seeds. + """ + import gemmi + + root = Path(root) + seeds = sorted(d for d in root.glob("seed_*") if d.is_dir()) + if not seeds: + print(f"no seed_* directories under {root}") + return [] + + def positions(path): + st = gemmi.read_structure(str(path)) + st.remove_hydrogens() + out = {} + for ch in st[0]: + if ch.name != chain: + continue + for res in ch: + if first <= res.seqid.num <= last: + for atom in res: + out[(res.seqid.num, atom.name)] = np.array( + [atom.pos.x, atom.pos.y, atom.pos.z] + ) + return out + + start = positions(start_pdb) + per_arm = {} + injected = [] + for sd in seeds: + truth_p = sd / "light_truth.pdb" + if not truth_p.exists(): + continue + truth = positions(truth_p) + shared = sorted(set(truth) & set(start)) + injected.append( + np.mean([np.linalg.norm(truth[k] - start[k]) for k in shared]) + ) + for arm_dir in sorted(sd.glob("refine_*")): + hits = sorted(arm_dir.glob("fractions_*_light.pdb")) + if not hits: + continue + got = positions(hits[0]) + keys = [k for k in shared if k in got] + if not keys: + continue + err = np.mean([np.linalg.norm(got[k] - truth[k]) for k in keys]) + rec = np.mean([np.linalg.norm(got[k] - start[k]) for k in keys]) + per_arm.setdefault(arm_dir.name, {})[sd.name] = (err, rec) + + inj = float(np.mean(injected)) + complete = set.intersection(*(set(v) for v in per_arm.values())) if per_arm else set() + complete = sorted(complete) + print(f"\ninjected displacement {inj:.3f} A; " + f"{len(complete)} seeds complete in all {len(per_arm)} arms") + if len(complete) < len(seeds): + print(f" ({len(seeds) - len(complete)} seed(s) dropped: not all arms finished)") + + print(f"\n{'arm':<16s} {'err (A)':>16s} {'recovered/injected':>20s}") + print("-" * 56) + for arm in sorted(per_arm): + e = np.array([per_arm[arm][s][0] for s in complete]) + r = np.array([per_arm[arm][s][1] for s in complete]) / inj + print(f"{arm:<16s} {e.mean():8.4f} +- {e.std(ddof=1):5.4f} " + f"{r.mean():14.3f} +- {r.std(ddof=1):.3f}") + + base = "refine_coh" + if base not in per_arm: + return per_arm + rng = np.random.default_rng(seed) + print(f"\nPaired against {base}, median of per-seed differences " + f"({n_boot} bootstrap resamples)") + print("-" * 72) + print(f"{'arm':<16s} {'median d(err)':>14s} {'95% CI':>22s} {'seeds better':>14s}") + for arm in sorted(per_arm): + if arm == base: + continue + d = np.array([per_arm[arm][s][0] - per_arm[base][s][0] for s in complete]) + boots = np.array([ + np.median(rng.choice(d, size=len(d), replace=True)) for _ in range(n_boot) + ]) + lo, hi = np.percentile(boots, [2.5, 97.5]) + print(f"{arm:<16s} {np.median(d):+14.4f} {f'[{lo:+.4f}, {hi:+.4f}]':>22s} " + f"{f'{(d < 0).sum()}/{len(d)}':>14s}") + print("-" * 72) + print("negative = closer to truth than the coherent refinement") + return per_arm + + def main(): ap = argparse.ArgumentParser(description=__doc__) repo = Path(__file__).resolve().parents[1] @@ -190,11 +287,17 @@ def main(): ap.add_argument("--device", default="cpu") ap.add_argument("--score", action="store_true", help="Compare refined models against the truth (after refining).") + ap.add_argument("--score-sweep", action="store_true", + help="Aggregate a seed sweep under --out, paired seed by seed.") args = ap.parse_args() out = Path(args.out) truth_pdb = out / "light_truth.pdb" + if args.score_sweep: + score_sweep(out, args.chain, args.first, args.last, args.pdb) + return 0 + if args.score: arms = [] for d in sorted(out.glob("refine_*")): diff --git a/paper/probe_two_moment_collinearity.py b/paper/probe_two_moment_collinearity.py index 5815dbaf..4c3af1de 100644 --- a/paper/probe_two_moment_collinearity.py +++ b/paper/probe_two_moment_collinearity.py @@ -84,15 +84,17 @@ def numeric_columns(params, evaluate, rel_step=1e-3): for i in range(flat.numel()): step = rel_step * max(abs(float(flat[i])), 1e-3) saved = float(flat[i]) + # No autograd: these are numerical derivatives, and retaining a graph per + # evaluation is what pushed this over the memory limit. with torch.no_grad(): flat[i] = saved + step - plus = evaluate() - with torch.no_grad(): + plus = evaluate() flat[i] = saved - step - minus = evaluate() - with torch.no_grad(): + minus = evaluate() flat[i] = saved - cols.append(((plus - minus) / (2 * step)).detach().cpu().numpy()) + col = ((plus - minus) / (2 * step)).cpu().numpy() + del plus, minus + cols.append(col) return np.asarray(cols).T # (n_hkl, n_param) From 58e794b015bbb1f4e679ee9518632b23c293629e Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 17:21:00 +0200 Subject: [PATCH 087/250] Fixed bug were anisotropic adps were being dropped on read from newer cif files. --- docs/changelog.rst | 1 + torchref/io/cif_readers.py | 62 ++++- torchref/model/disorder_field.py | 253 +++++++++++++++--- torchref/model/model.py | 137 ++++++++-- torchref/model/parameter_wrappers.py | 124 +++++---- .../refinement/targets/adp/node_smoothness.py | 7 +- 6 files changed, 463 insertions(+), 121 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 1d1c0596..b36d4ceb 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Version 0.7.0 ---------- +- Fixed cif reading bug discarding new mmCIF field for aniso ADPs - Separated model configuration and provenance into ``ModelContext``. It now holds the unit cell, space group, atom table, link records, hydrogen settings, and input paths. - Refactored ``Symmetry`` as a crystallography-free class with transform primitives, and made ``SpaceGroup`` a specialised subclass. - Moved geometry predicates, HKL verbs, and grid-size helpers onto these classes as methods. diff --git a/torchref/io/cif_readers.py b/torchref/io/cif_readers.py index 44d1b9e5..797f4230 100644 --- a/torchref/io/cif_readers.py +++ b/torchref/io/cif_readers.py @@ -12,6 +12,8 @@ import gemmi import numpy as np +import warnings + import pandas as pd #: Column holding the ``data_`` block a loop row was read from. Added only when @@ -1445,7 +1447,59 @@ def get_atom_data(self) -> pd.DataFrame: "_atom_site.aniso_U[2][3]", ] - if all(col in atom_df.columns for col in aniso_cols): + # The standard mmCIF home for anisotropic ADPs is the SEPARATE + # ``_atom_site_anisotrop`` loop, keyed by ``.id`` against ``_atom_site.id``. + # Only the legacy in-line ``_atom_site.aniso_U[i][j]`` form was read here, so a + # standards-conforming file -- every PDB-REDO entry, and anything the PDB emits + # as mmCIF -- silently loaded with no anisotropy at all and every atom marked + # isotropic. + aniso_df = getattr(self.cif, "data", {}).get("atom_site_anisotrop") + std_cols = [f"_atom_site_anisotrop.U[{i}][{j}]" + for i, j in ((1, 1), (2, 2), (3, 3), (1, 2), (1, 3), (2, 3))] + key, atom_key = "_atom_site_anisotrop.id", "_atom_site.id" + joined = None + if ( + aniso_df is not None + and all(c in aniso_df.columns for c in std_cols) + and key in aniso_df.columns + and atom_key in atom_df.columns + ): + # Join on the id as a STRING. Coercing to a number first silently produces + # NaN keys for any non-integer id and then mis-pairs U tensors with atoms, + # which is far worse than having no anisotropy: the model is scrambled but + # still refines. + left = pd.DataFrame({"_k": atom_df[atom_key].astype(str).str.strip()}) + right = aniso_df[[key] + std_cols].copy() + right["_k"] = right[key].astype(str).str.strip() + right = right.drop_duplicates("_k") + merged = left.merge(right, on="_k", how="left") + if len(merged) == len(atom_df): + joined = merged + + if joined is not None: + for name, col in zip(("u11", "u22", "u33", "u12", "u13", "u23"), std_cols): + result[name] = pd.to_numeric(joined[col].to_numpy(), errors="coerce") + result["anisou_flag"] = ~pd.isna(result["u11"]) + n_hit = int(result["anisou_flag"].sum()) + frac = n_hit / max(len(aniso_df), 1) + if frac < 0.9: + # Partial coverage is legitimate in small amounts -- waters and + # hydrogens often carry no ANISOU -- but a large shortfall means the two + # loops are not labelled the same way, and then the rows that DID match + # cannot be trusted to have matched the right atoms. Drop the anisotropy + # rather than apply a possibly mis-paired subset: an isotropic model is + # merely less informative, a scrambled one still refines and is wrong. + warnings.warn( + f"{self.filepath}: matched only {n_hit} of {len(aniso_df)} " + "anisotropic records to atoms, so _atom_site.id and " + "_atom_site_anisotrop.id do not agree; discarding the anisotropy " + "and loading isotropically.", + RuntimeWarning, + ) + for name in ("u11", "u22", "u33", "u12", "u13", "u23"): + result[name] = np.nan + result["anisou_flag"] = False + elif all(col in atom_df.columns for col in aniso_cols): result["u11"] = pd.to_numeric( atom_df["_atom_site.aniso_U[1][1]"], errors="coerce" ) @@ -1675,6 +1729,12 @@ def has_anisotropic_data(self) -> bool: "_atom_site.aniso_U[2][2]", "_atom_site.aniso_U[3][3]", ] + aniso_df = self.cif.data.get("atom_site_anisotrop") + if aniso_df is not None and all( + f"_atom_site_anisotrop.U[{i}][{j}]" in aniso_df.columns + for i, j in ((1, 1), (2, 2), (3, 3), (1, 2), (1, 3), (2, 3)) + ): + return True return all(col in self.cif.data["atom_site"].columns for col in aniso_cols) def get_coordinates(self) -> Optional[np.ndarray]: diff --git a/torchref/model/disorder_field.py b/torchref/model/disorder_field.py index 35bbc75e..97b94d28 100644 --- a/torchref/model/disorder_field.py +++ b/torchref/model/disorder_field.py @@ -24,10 +24,18 @@ from torch import nn from torchref.config import get_float_dtype, normalize_device -from torchref.model.parameter_wrappers import MixedTensor +from torchref.model.parameter_wrappers import MixedTensor, raw6_to_u6, u6_to_raw6 from torchref.utils.utils import ModuleReference -__all__ = ["DisorderFieldTensor", "farthest_point_anchors", "build_neighbor_list"] +__all__ = [ + "DisorderFieldTensor", + "NodePayload", + "IsotropicPayload", + "AnisotropicPayload", + "farthest_point_anchors", + "density_anchor_rows", + "build_neighbor_list", +] def farthest_point_anchors(xyz: torch.Tensor, n_nodes: int) -> torch.Tensor: @@ -161,6 +169,140 @@ def _wrap_accessor(xyz_fn): return xyz_fn +# ---------------------------------------------------------------------------------- +# Payload strategies. A payload says what one node carries and how that becomes a +# per-atom quantity; it knows nothing about where nodes are or how weights arise. +# +# Deliberately stateless. Node parameters live in one flat leaf on the field itself, +# because ``Model.parameters_of_types`` reads a single ``refinable_params`` per type, +# and a strategy owning parameters would break that. Each strategy declares only how +# many columns of that leaf it interprets. +# ---------------------------------------------------------------------------------- + + +class NodePayload: + """What a node carries, and how it becomes a per-atom ADP. + + Attributes + ---------- + width : int + Columns of node storage this payload interprets. + out_width : int + Components of the per-atom output: 1 for an isotropic B, 6 for a U tensor. + """ + + width: int = 1 + out_width: int = 1 + + def contributions(self, payload, xyz, node_pos, neighbor_list): + """``(n_atoms, k, out_width)``: what each candidate node offers each atom. + + ``xyz`` and ``node_pos`` are passed even though the payloads here ignore them, + because a payload with an r-dependence (TLS: constant, linear and quadratic in + the displacement from the node) needs them, and giving it the arguments now + means adding one later touches no shared code. + """ + raise NotImplementedError + + def fit(self, target, w_dense, epsilon): + """``(K, width)`` least-squares payload reproducing ``target``.""" + raise NotImplementedError + + def log_magnitude(self, payload): + """``(K,)`` log of each node's ADP magnitude, for a magnitude restraint. + + Lets a restraint price node values without branching on payload type. + """ + raise NotImplementedError + + +class IsotropicPayload(NodePayload): + """One isotropic B per node, stored as ``log B`` so it stays positive. + + The per-atom B is a convex combination of positive node values, so it is positive + without a clamp. + """ + + width = 1 + out_width = 1 + + def contributions(self, payload, xyz, node_pos, neighbor_list): + return torch.exp(payload[:, 0])[neighbor_list].unsqueeze(-1) + + def fit(self, target, w_dense, epsilon): + b = _ridged_solve(w_dense, target.unsqueeze(-1)).squeeze(-1) + return torch.log(b.clamp(min=epsilon)).unsqueeze(-1) + + def log_magnitude(self, payload): + return payload[:, 0] + + +class AnisotropicPayload(NodePayload): + """A full U tensor per node, positive-definite by construction. + + Stored as the six free parameters of a Cholesky factor, so ``U = L L^T`` is PD for + any parameter value -- the same device + :class:`~torchref.model.parameter_wrappers.CholeskyMixedTensor` uses for per-atom + ADPs, and for the same reason: an indefinite U makes the anisotropic B-matrix + singular and the structure-factor FFT returns NaN. + + Positive-definiteness survives the combination for free: the per-atom U is a convex + combination of PD matrices. Averaging in U space is what buys that -- averaging the + Cholesky parameters instead would be a different object with no such guarantee. + """ + + width = 6 + out_width = 6 + + def __init__(self, epsilon: float = 1e-3): + # Floor on the Cholesky diagonal, which bounds the smallest eigenvalue of U + # from below. Same default and same meaning as the per-atom wrapper. + self.epsilon = float(epsilon) + + def contributions(self, payload, xyz, node_pos, neighbor_list): + return raw6_to_u6(payload, self.epsilon)[neighbor_list] + + def fit(self, target, w_dense, epsilon): + """Fit six U components at once, then re-encode as Cholesky parameters. + + The per-atom U is linear in each component independently, so this is the same + ridged solve as the isotropic case with a six-column right-hand side. The + least-squares result is not constrained to be PD, which is why it goes back + through ``u6_to_raw6`` -- that projects onto PD by clamping eigenvalues. + """ + if target.ndim == 1: # a B target: lift to the equivalent isotropic U + u_iso = target / (8.0 * math.pi**2) + zero = torch.zeros_like(u_iso) + target = torch.stack([u_iso, u_iso, u_iso, zero, zero, zero], dim=1) + return u6_to_raw6(_ridged_solve(w_dense, target), self.epsilon) + + def log_magnitude(self, payload): + u6 = raw6_to_u6(payload, self.epsilon) + b_eq = (8.0 * math.pi**2 / 3.0) * (u6[:, 0] + u6[:, 1] + u6[:, 2]) + return torch.log(b_eq.clamp(min=1e-6)) + + +def _ridged_solve(w_dense, target): + """Least squares ``min ||W x - target||`` through the ridged normal equations. + + Non-finite target rows are dropped rather than carried in: deposited models have + NaN ADPs, and one NaN row propagates through the normal equations and takes every + node with it. Those atoms still receive a fitted value on output. + """ + finite = torch.isfinite(target).all(dim=-1) + if not bool(finite.all()): + w_dense, target = w_dense[finite], target[finite] + if target.shape[0] == 0: + return torch.zeros( + w_dense.shape[1], target.shape[-1], + dtype=w_dense.dtype, device=w_dense.device, + ) + gram = w_dense.T @ w_dense + ridge = 1e-6 * torch.diagonal(gram).mean().clamp(min=1e-30) + eye = torch.eye(gram.shape[0], dtype=gram.dtype, device=gram.device) + return torch.linalg.solve(gram + ridge * eye, w_dense.T @ target) + + class DisorderFieldTensor(MixedTensor): """Per-atom ADPs from a small set of nodes, each atom a weighted mean of its k nearest. @@ -221,6 +363,7 @@ def __init__( n_nodes: int = 32, k_neighbors: int = 12, refine_positions: bool = False, + payload: Optional["NodePayload"] = None, anchor_rows: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, node_values: Optional[torch.Tensor] = None, refinable_mask: Optional[torch.Tensor] = None, @@ -234,6 +377,9 @@ def __init__( self.epsilon = epsilon self._k_neighbors = int(k_neighbors) self._refine_positions = bool(refine_positions) + # Storage columns are [payload | log sigma | offset]. Payload first keeps the + # isotropic layout unchanged, so an existing state dict still loads. + self._payload = payload if payload is not None else IsotropicPayload() object.__setattr__(self, "_xyz_fn", _wrap_accessor(xyz_fn)) if initial_values is None and node_values is None: @@ -350,21 +496,24 @@ def _weights( def _fit_nodes( self, - target_b: torch.Tensor, + target: torch.Tensor, xyz: torch.Tensor, node_pos: torch.Tensor, neighbor_list: torch.Tensor, ) -> torch.Tensor: - """Least-squares node values reproducing ``target_b`` as closely as possible. + """Node storage whose field reproduces ``target`` as closely as it can. - ``sigma`` is seeded at half the median nearest-neighbour node spacing, then the - node values are the ridged solution of ``W b = target_b``. Linear in ``b``, so - this is a closed-form solve rather than an optimisation loop. + Kernel width is seeded at half the median nearest-neighbour node spacing, which + makes it a property of the node layout rather than a tuned constant. With the + weights fixed at that seed the payload is linear in the target, so the payload + half is a closed-form solve rather than an optimisation loop -- delegated, + because what "linear in the target" means differs between a scalar B and a + six-component U. Returns ------- torch.Tensor - ``(K, 2)`` storage, ``[log b, log sigma]``. + ``(K, payload.width + 1 + 3*refine_positions)`` node storage. """ n_k = int(node_pos.shape[0]) if n_k > 1: @@ -378,34 +527,16 @@ def _fit_nodes( (n_k,), math.log(sigma0), dtype=xyz.dtype, device=xyz.device ) - W_sparse = self._weights(xyz, node_pos, neighbor_list, log_sigma) - W = torch.zeros(xyz.shape[0], n_k, dtype=xyz.dtype, device=xyz.device) - W.scatter_(1, neighbor_list, W_sparse) - - # Non-finite targets are dropped from the solve rather than carried into it: - # a single NaN row would propagate through the normal equations and take - # every node with it. Those atoms still receive a fitted value on output. - finite = torch.isfinite(target_b) - if not bool(finite.all()): - W, target_b = W[finite], target_b[finite] - if target_b.numel() == 0: - flat = torch.stack([torch.zeros_like(log_sigma), log_sigma], dim=1) - if self._refine_positions: - flat = torch.cat([flat, torch.zeros_like(node_pos)], dim=1) - return flat - - gram = W.T @ W - ridge = 1e-6 * torch.diagonal(gram).mean().clamp(min=1e-30) - eye = torch.eye(n_k, dtype=W.dtype, device=W.device) - b = torch.linalg.solve(gram + ridge * eye, W.T @ target_b) - - b = b.clamp(min=self.epsilon) - node_values = torch.stack([torch.log(b), log_sigma], dim=1) + w_sparse = self._weights(xyz, node_pos, neighbor_list, log_sigma) + w_dense = torch.zeros(xyz.shape[0], n_k, dtype=xyz.dtype, device=xyz.device) + w_dense.scatter_(1, neighbor_list, w_sparse) + + payload = self._payload.fit(target, w_dense, self.epsilon) + + columns = [payload, log_sigma.unsqueeze(-1)] if self._refine_positions: - node_values = torch.cat( - [node_values, torch.zeros_like(node_pos)], dim=1 - ) - return node_values + columns.append(torch.zeros_like(node_pos)) + return torch.cat(columns, dim=1) @staticmethod def _collapse_mask( @@ -422,6 +553,31 @@ def _collapse_mask( # Public surface. # ------------------------------------------------------------------ + @property + def payload(self) -> "NodePayload": + """What each node carries and how it becomes a per-atom ADP.""" + return self._payload + + @property + def out_width(self) -> int: + """Components of the per-atom output: 1 for isotropic B, 6 for a U tensor.""" + return self._payload.out_width + + def _split(self, raw): + """Storage columns as ``(payload, log sigma, offset or None)``.""" + w = self._payload.width + offset = raw[:, w + 1 : w + 4] if self._refine_positions else None + return raw[:, :w], raw[:, w], offset + + def log_magnitude(self, raw=None) -> torch.Tensor: + """``(K,)`` log ADP magnitude per node, whatever the payload. + + Lets a magnitude restraint price node values without knowing the layout. + """ + if raw is None: + raw = super().forward() + return self._payload.log_magnitude(self._split(raw)[0]) + @property def shape(self): """Shape of the FULL per-atom tensor, not the node storage.""" @@ -465,7 +621,8 @@ def _node_positions_from(self, xyz, raw): base = self._segment_mean( xyz, self.anchor_atom, self.anchor_node, raw.shape[0] ) - return base + raw[:, 2:5] if self._refine_positions else base + offset = self._split(raw)[2] + return base if offset is None else base + offset def weights(self, xyz: Optional[torch.Tensor] = None) -> torch.Tensor: """Per-atom weights over candidate nodes, ``(n_atoms, k)``. Rows sum to one.""" @@ -473,7 +630,10 @@ def weights(self, xyz: Optional[torch.Tensor] = None) -> torch.Tensor: xyz = self._xyz_fn() raw = super().forward() return self._weights( - xyz, self._node_positions_from(xyz, raw), self.neighbor_list, raw[:, 1] + xyz, + self._node_positions_from(xyz, raw), + self.neighbor_list, + self._split(raw)[1], ) def node_load(self, xyz: Optional[torch.Tensor] = None) -> torch.Tensor: @@ -537,10 +697,16 @@ def evaluate(self, xyz: torch.Tensor, raw: torch.Tensor) -> torch.Tensor: raw : torch.Tensor ``(K, 2)`` node storage, ``[log b, log sigma]``. """ - log_b, log_sigma = raw[:, 0], raw[:, 1] + payload, log_sigma, _ = self._split(raw) node_pos = self._node_positions_from(xyz, raw) W = self._weights(xyz, node_pos, self.neighbor_list, log_sigma) - return (W * torch.exp(log_b)[self.neighbor_list]).sum(dim=1) + contrib = self._payload.contributions( + payload, xyz, node_pos, self.neighbor_list + ) + out = (W.unsqueeze(-1) * contrib).sum(dim=1) + # A scalar payload reports per-atom B as (N,), not (N, 1): that is the shape + # every consumer of ``adp()`` expects. + return out.squeeze(-1) if self._payload.out_width == 1 else out def forward(self) -> torch.Tensor: """Per-atom isotropic B, ``(n_atoms,)``. @@ -584,7 +750,10 @@ def _set_values(self, key, value: torch.Tensor) -> None: ) def refit(self, target_b: torch.Tensor) -> None: - """Re-fit the node values to a per-atom B target, in place. + """Re-fit the node values to a per-atom target, in place. + + The target is per-atom B for a scalar payload, or per-atom U6 for a tensor one; + an anisotropic payload also accepts a B target and lifts it to ``U_iso * I``. Replaces ``refinable_params``, so any optimizer state held for it is stale. """ @@ -668,6 +837,7 @@ def copy(self) -> "DisorderFieldTensor": xyz_fn=accessor, k_neighbors=self._k_neighbors, refine_positions=self._refine_positions, + payload=self._payload, anchor_rows=(self.anchor_atom.clone(), self.anchor_node.clone()), node_values=self.node_values().detach().clone(), refinable_mask=self.refinable_mask.clone(), @@ -686,6 +856,7 @@ def __repr__(self) -> str: name_str = f"'{self.name}', " if self.name is not None else "" return ( f"DisorderFieldTensor({name_str}atoms={self._full_shape}, " - f"nodes={self.n_nodes}, k={self._k_neighbors}, dtype={self.dtype}, " + f"nodes={self.n_nodes}, k={self._k_neighbors}, " + f"payload={type(self._payload).__name__}, dtype={self.dtype}, " f"device={self.device}, refinable={self.get_refinable_count()})" ) diff --git a/torchref/model/model.py b/torchref/model/model.py index 85056afa..4e6d128c 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -592,6 +592,13 @@ def load(self, reader, add_hydrogens: bool = None): else self.pdb ) self.pdb.dropna(subset=["x", "y", "z", "tempfactor", "occupancy"], inplace=True) + # Reindex before deriving the ``index`` column: every consumer uses it to + # address length-N per-atom tensors positionally (see + # ``_create_occupancy_groups``), so a gapped index from the drop above sends + # them past the end. Only the strip_H branch reset, so a model losing rows to + # the dropna instead -- an atom with no coordinates or no B -- raised + # IndexError at load. Hit on roughly one PDB-REDO entry in six. + self.pdb.reset_index(drop=True, inplace=True) self.pdb["index"] = self.pdb.index.to_numpy(dtype=int) self.cell = Cell(cell, dtype=self.dtype_float, device=self.device) @@ -911,10 +918,14 @@ def update_pdb(self): """ Write the current refinable parameters back into ``self.pdb``. - Copies the live values of ``xyz`` (x/y/z), ``u`` (u11..u23), ``adp`` - (tempfactor), and ``occupancy`` from the parameter wrappers into the - corresponding columns of the ``self.pdb`` DataFrame. Called by every - writer and by ``hydrogenate`` before output. + Copies the live values of ``xyz`` (x/y/z), ``u`` (u11..u23) and + ``occupancy`` from the parameter wrappers into the corresponding columns of + the ``self.pdb`` DataFrame. Called by every writer and by ``hydrogenate`` + before output. + + ``tempfactor`` is the equivalent isotropic B whenever any atom is + anisotropic, so the column agrees with the ANISOU records written beside it; + with no anisotropic atoms it is the isotropic wrapper directly. Returns ------- @@ -931,7 +942,19 @@ def update_pdb(self): self.pdb.loc[:, ["u11", "u22", "u33", "u12", "u13", "u23"]] = ( self.u().cpu().detach().numpy() ) - self.pdb.loc[:, "tempfactor"] = self.adp().cpu().detach().numpy() + # The B column must agree with the ANISOU records beside it: for an + # anisotropic atom the PDB convention is B_eq = (8 pi^2 / 3) tr(U), not + # whatever the isotropic wrapper still happens to hold. That wrapper stops + # being refined the moment an atom goes anisotropic, so writing it directly + # emits a stale B alongside a live U. + if getattr(self, "_aniso_is_empty", True): + self.pdb.loc[:, "tempfactor"] = self.adp().cpu().detach().numpy() + else: + from torchref.base.targets.adp import u6_b_eq + + self.pdb.loc[:, "tempfactor"] = ( + u6_b_eq(self.adp_u6()).cpu().detach().numpy() + ) self.pdb.loc[:, "occupancy"] = self.occupancy().cpu().detach().numpy() return self.pdb @@ -1244,7 +1267,7 @@ def set_adp_mode( Parameters ---------- - mode : {"isotropic", "anisotropic", "field"}, optional + mode : {"isotropic", "anisotropic", "field", "field_aniso"}, optional ``"isotropic"`` (default) converts every atom, previously anisotropic ones to ``B_eq = (8 pi^2 / 3)(U11 + U22 + U33)``. ``"anisotropic"`` converts those matching ``aniso_selection``, expanding isotropic atoms @@ -1275,17 +1298,40 @@ def set_adp_mode( """ if not self.ctx.initialized or self.pdb is None: return - if mode == "field": - # Collapse to a clean per-atom isotropic state first. The field is - # isotropic in this representation, and the partition owns every buffer - # keyed off the iso/aniso split, so reuse it rather than duplicating it. - self._apply_adp_partition( - torch.zeros(len(self.pdb), dtype=torch.bool, device=self.device) - ) + if mode in ("field", "field_aniso"): + aniso = mode == "field_aniso" + # Run the partition first either way: it owns every buffer keyed off the + # iso/aniso split, and it converts the stored values in the right direction + # (B -> U_iso*I entering anisotropic, U -> B_eq entering isotropic), so the + # field is fitted to a target that is already in its own representation. + if aniso: + # Every atom, unless the caller narrows it. A node field is not the + # per-atom parametrisation that "not water, not hydrogen" exists to + # ration -- its cost is set by node count, not atom count -- and a + # partial selection would leave half the ADPs coming from the field and + # half from the per-atom wrapper, which is not a representation anyone + # asked for. + if aniso_selection is None: + target_mask = torch.ones( + len(self.pdb), dtype=torch.bool, device=self.device + ) + else: + from torchref.utils.utils import create_selection_mask + + target_mask = torch.as_tensor( + create_selection_mask(aniso_selection, self.pdb), + dtype=torch.bool, + ).to(self.device) + else: + target_mask = torch.zeros( + len(self.pdb), dtype=torch.bool, device=self.device + ) + self._apply_adp_partition(target_mask) self._install_disorder_field( n_nodes=n_nodes, k_neighbors=k_neighbors, refine_node_positions=refine_node_positions, + anisotropic=aniso, ) return if mode == "isotropic": @@ -1301,37 +1347,61 @@ def set_adp_mode( ).to(self.device) else: raise ValueError( - f"Unknown ADP mode: {mode!r}. Use 'isotropic', 'anisotropic' " - "or 'field'." + f"Unknown ADP mode: {mode!r}. Use 'isotropic', 'anisotropic', " + "'field' or 'field_aniso'." ) self._apply_adp_partition(aniso_mask) @property def adp_is_field(self) -> bool: - """Whether ``adp`` is a node field rather than a per-atom wrapper.""" + """Whether either ADP slot holds a node field rather than a per-atom wrapper.""" from torchref.model.disorder_field import DisorderFieldTensor - return isinstance(self.adp, DisorderFieldTensor) + return isinstance(self.adp, DisorderFieldTensor) or isinstance( + self.u, DisorderFieldTensor + ) + + @property + def adp_field(self): + """The node field driving the ADPs, or ``None`` if neither slot holds one.""" + from torchref.model.disorder_field import DisorderFieldTensor + + for wrapper in (self.u, self.adp): + if isinstance(wrapper, DisorderFieldTensor): + return wrapper + return None def _install_disorder_field( self, n_nodes: int = None, k_neighbors: int = 12, refine_node_positions: bool = False, + anisotropic: bool = False, ): - """Replace the per-atom ``adp`` wrapper with a node field fitted to its B. + """Replace a per-atom ADP wrapper with a node field fitted to it. - Expects the model to already be in a per-atom isotropic state, which - :meth:`set_adp_mode` arranges by running the partition first. + The field lands in the slot its payload feeds: an isotropic payload takes over + ``adp`` and leaves the model isotropic, an anisotropic one takes over ``u`` and + the model refines every selected atom anisotropically. Both expect the partition + to have run first, which :meth:`set_adp_mode` arranges. """ from torchref.model.disorder_field import ( + AnisotropicPayload, DisorderFieldTensor, + IsotropicPayload, density_anchor_rows, ) with torch.no_grad(): - B = self.adp().detach().clone() xyz = self.xyz().detach() + # The fit target is whatever the partition just produced: per-atom U6 for + # the anisotropic payload, per-atom B for the isotropic one. + target = ( + self.adp_u6().detach().clone() + if anisotropic + else self.adp().detach().clone() + ) + B = target if n_nodes is None: n_nodes = max(4, int(round(len(self.pdb) / 25.0))) @@ -1340,25 +1410,34 @@ def _install_disorder_field( # wearing a node's clothes. anchor_rows = density_anchor_rows(xyz, min(n_nodes, len(self.pdb))) - self.adp = DisorderFieldTensor( - initial_values=B.to(self.dtype_float), + field = DisorderFieldTensor( + initial_values=target.to(self.dtype_float), xyz_fn=self.xyz, n_nodes=n_nodes, refine_positions=refine_node_positions, + payload=AnisotropicPayload() if anisotropic else IsotropicPayload(), anchor_rows=anchor_rows, k_neighbors=k_neighbors, - name="adp", + name="aniso_U" if anisotropic else "adp", dtype=self.dtype_float, device=self.device, ) - # ``adp_mask`` is in atom space; the field collapses it onto its nodes. - self.adp.update_refinable_mask(self.adp_mask) + if anisotropic: + self.u = field + # The mask is in atom space either way; the field collapses it onto nodes. + self.u.update_refinable_mask(self.u_mask) + else: + self.adp = field + self.adp.update_refinable_mask(self.adp_mask) if self.ctx.verbose > 0: + kind = "aniso U" if anisotropic else "iso B" + was = len(self.pdb) * (6 if anisotropic else 1) print( - f"ADP field: {self.adp.n_nodes} nodes, k={k_neighbors}, " - f"{int(self.adp.get_refinable_count())} refinable " - f"(was {len(self.pdb)} per-atom B)" + f"ADP field ({kind}): {field.n_nodes} nodes, k={k_neighbors}, " + f"{int(field.get_refinable_count())} refinable nodes, " + f"{int(field.refinable_params.numel())} parameters " + f"(was {was} per-atom)" ) if hasattr(self, "reset_cache"): self.reset_cache() diff --git a/torchref/model/parameter_wrappers.py b/torchref/model/parameter_wrappers.py index f2beddc4..e6f66d05 100644 --- a/torchref/model/parameter_wrappers.py +++ b/torchref/model/parameter_wrappers.py @@ -972,6 +972,77 @@ def __str__(self) -> str: ) +# ---------------------------------------------------------------------------------- +# U <-> Cholesky transforms. Free functions because two unrelated holders need them: +# CholeskyMixedTensor for per-atom ADPs, and the node-field anisotropic payload for +# per-node ones. Both operate on (..., 6) tensors and pass NaN rows through untouched. +# ---------------------------------------------------------------------------------- + + +def u6_to_matrix(U: torch.Tensor) -> torch.Tensor: + """``(..., 6)`` U components to a symmetric ``(..., 3, 3)`` matrix.""" + M = U.new_zeros(*U.shape[:-1], 3, 3) + M[..., 0, 0] = U[..., 0] + M[..., 1, 1] = U[..., 1] + M[..., 2, 2] = U[..., 2] + M[..., 0, 1] = M[..., 1, 0] = U[..., 3] + M[..., 0, 2] = M[..., 2, 0] = U[..., 4] + M[..., 1, 2] = M[..., 2, 1] = U[..., 5] + return M + + +def raw6_to_u6(raw: torch.Tensor, epsilon: float) -> torch.Tensor: + """Cholesky free parameters to U components, ``U = L L^T``. + + Positive-definite for any input: the diagonal of ``L`` is ``exp(x) + epsilon``, so + ``epsilon`` bounds the smallest eigenvalue of ``U`` from below. No factorisation + happens here, which is what makes this safe to call in a forward pass. + """ + diag, off = raw[..., :3], raw[..., 3:] + L11 = torch.exp(diag[..., 0]) + epsilon + L22 = torch.exp(diag[..., 1]) + epsilon + L33 = torch.exp(diag[..., 2]) + epsilon + L21, L31, L32 = off[..., 0], off[..., 1], off[..., 2] + return torch.stack( + [ + L11 * L11, + L21 * L21 + L22 * L22, + L31 * L31 + L32 * L32 + L33 * L33, + L21 * L11, + L31 * L11, + L31 * L21 + L32 * L22, + ], + dim=-1, + ) + + +def u6_to_raw6(U: torch.Tensor, epsilon: float) -> torch.Tensor: + """U components to Cholesky free parameters, projecting onto positive-definite. + + A least-squares or deposited U need not be PD, so the matrix is symmetrised and its + eigenvalues clamped before factorising. Runs at construction and on mask changes, + never in a forward pass. Forced onto the CPU: cuSolver's batched kernels fail on + the large degenerate batches an isotropic model produces, while LAPACK handles them. + """ + finite = torch.isfinite(U).all(dim=-1) + M = u6_to_matrix(torch.nan_to_num(U, nan=0.0)) + eye = torch.eye(3, dtype=M.dtype, device=M.device).expand_as(M) + M = torch.where(finite[..., None, None], M, eye) + M = 0.5 * (M + M.transpose(-1, -2)) + + src_device = M.device + M = M.cpu() + w, V = torch.linalg.eigh(M) + w = w.clamp(min=epsilon * epsilon) + M = (V * w.unsqueeze(-2)) @ V.transpose(-1, -2) + L = torch.linalg.cholesky(M) + diag = torch.stack([L[..., 0, 0], L[..., 1, 1], L[..., 2, 2]], dim=-1) + off = torch.stack([L[..., 1, 0], L[..., 2, 0], L[..., 2, 1]], dim=-1) + raw_diag = torch.log((diag - epsilon).clamp(min=1e-12)) + raw = torch.cat([raw_diag, off], dim=-1).to(src_device) + return torch.where(finite.unsqueeze(-1), raw, torch.full_like(raw, float("nan"))) + + class CholeskyMixedTensor(MixedTensor): """A MixedTensor for anisotropic ADPs (U tensors) kept positive-definite. @@ -1029,57 +1100,16 @@ def __init__( # ------------------------------------------------------------------ @staticmethod def _u6_to_matrix(U: torch.Tensor) -> torch.Tensor: - M = U.new_zeros(*U.shape[:-1], 3, 3) - M[..., 0, 0] = U[..., 0] - M[..., 1, 1] = U[..., 1] - M[..., 2, 2] = U[..., 2] - M[..., 0, 1] = M[..., 1, 0] = U[..., 3] - M[..., 0, 2] = M[..., 2, 0] = U[..., 4] - M[..., 1, 2] = M[..., 2, 1] = U[..., 5] - return M + """Delegate to :func:`u6_to_matrix`.""" + return u6_to_matrix(U) def _u6_to_raw6(self, U: torch.Tensor) -> torch.Tensor: - """U components -> Cholesky free parameters [log(L_ii - eps); L_offdiag].""" - eps = self.epsilon - finite = torch.isfinite(U).all(dim=-1) - M = self._u6_to_matrix(torch.nan_to_num(U, nan=0.0)) - eye = torch.eye(3, dtype=M.dtype, device=M.device).expand_as(M) - M = torch.where(finite[..., None, None], M, eye) - # Project to positive-definite: symmetrise, clamp eigenvalues off zero. - # No-op for well-conditioned deposited U; rescues marginally non-PD input. - M = 0.5 * (M + M.transpose(-1, -2)) - # eigh + Cholesky forced onto the CPU: cuSolver's *batched* kernels fail - # (CUSOLVER_STATUS_INVALID_VALUE) on the large degenerate batches an - # isotropic ensemble produces (U ≡ 0), while LAPACK handles them. This - # runs only at load / mask change, never per optimizer step. - src_device = M.device - M = M.cpu() - w, V = torch.linalg.eigh(M) - w = w.clamp(min=eps * eps) - M = (V * w.unsqueeze(-2)) @ V.transpose(-1, -2) - L = torch.linalg.cholesky(M) - diag = torch.stack([L[..., 0, 0], L[..., 1, 1], L[..., 2, 2]], dim=-1) - off = torch.stack([L[..., 1, 0], L[..., 2, 0], L[..., 2, 1]], dim=-1) - raw_diag = torch.log((diag - eps).clamp(min=1e-12)) # invert exp(x)+eps - raw = torch.cat([raw_diag, off], dim=-1).to(src_device) - nan = torch.full_like(raw, float("nan")) - return torch.where(finite.unsqueeze(-1), raw, nan) + """U components -> Cholesky free parameters. See :func:`u6_to_raw6`.""" + return u6_to_raw6(U, self.epsilon) def _raw6_to_u6(self, raw: torch.Tensor) -> torch.Tensor: - """Cholesky free parameters -> U components (U = L Lᵀ). PD by construction.""" - eps = self.epsilon - diag, off = raw[..., :3], raw[..., 3:] - L11 = torch.exp(diag[..., 0]) + eps - L22 = torch.exp(diag[..., 1]) + eps - L33 = torch.exp(diag[..., 2]) + eps - L21, L31, L32 = off[..., 0], off[..., 1], off[..., 2] - U11 = L11 * L11 - U22 = L21 * L21 + L22 * L22 - U33 = L31 * L31 + L32 * L32 + L33 * L33 - U12 = L21 * L11 - U13 = L31 * L11 - U23 = L31 * L21 + L32 * L22 - return torch.stack([U11, U22, U33, U12, U13, U23], dim=-1) + """Cholesky free parameters -> U components. See :func:`raw6_to_u6`.""" + return raw6_to_u6(raw, self.epsilon) def forward(self) -> torch.Tensor: """Return the full U tensor (positive-definite per finite row).""" diff --git a/torchref/refinement/targets/adp/node_smoothness.py b/torchref/refinement/targets/adp/node_smoothness.py index 1db4b413..586dab5d 100644 --- a/torchref/refinement/targets/adp/node_smoothness.py +++ b/torchref/refinement/targets/adp/node_smoothness.py @@ -84,8 +84,9 @@ def _field(self): def _pair_terms(self): """``(weighted mean squared log-B difference, pair weights)``.""" field = self._field - raw = field.node_values() - log_b = raw[:, 0] + # Through the payload, not by column index: for a tensor payload column 0 is a + # Cholesky component, not a magnitude. + log_b = field.log_magnitude() pos = field.node_positions() d = torch.cdist(pos, pos) @@ -126,7 +127,7 @@ def stats(self) -> Dict[str, any]: with torch.no_grad(): w, diff2, lam = self._pair_terms() loss = self.forward() - log_b = field.node_values()[:, 0] + log_b = field.log_magnitude() b = torch.exp(log_b) return { "node_smoothness_loss": stat(float(loss), VERBOSITY_STANDARD), From e1f30dcdfffcb20eed505c3c71d255c3cc4a75dc Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 17:29:50 +0200 Subject: [PATCH 088/250] forgot two testfiles --- tests/helpers/device_cases.py | 3 +++ tests/unit/model/test_model.py | 37 ++++++++++++++++++++++++++++++++++ 2 files changed, 40 insertions(+) diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index 6946dd62..da36f0a4 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -409,6 +409,9 @@ class TargetDeviceCase: "BaseWeighting": "abstract base; covered via ManualWeighting", "Refinement": "abstract base; covered via LBFGSRefinement in integration", "PassThroughTensor": "documented non-functional stub (parameter_wrappers.py)", + "NodePayload": "stateless strategy, holds no tensors; abstract base", + "IsotropicPayload": "stateless strategy, holds no tensors", + "AnisotropicPayload": "stateless strategy, holds only a float epsilon", "ADPTarget": "abstract base; needs a model with ADPs", "CombinedTargets": "composite container; needs its component targets", "CombinedModelTargets": "composite container; needs a loaded model", diff --git a/tests/unit/model/test_model.py b/tests/unit/model/test_model.py index c547ac97..7b9045fa 100644 --- a/tests/unit/model/test_model.py +++ b/tests/unit/model/test_model.py @@ -117,3 +117,40 @@ def test_get_selection_mask_uninitialized_raises(self): with pytest.raises(RuntimeError, match="uninitialized"): model.get_selection_mask("chain A") + + +@pytest.mark.unit +def test_dropped_rows_leave_a_positional_index(pdb_dir, tmp_path): + """A model losing atoms to the NaN drop must still index its own tensors. + + ``load`` derives the ``index`` column from the DataFrame index, and every + consumer uses it to address length-N per-atom tensors positionally. Dropping rows + without reindexing leaves gaps, so the largest value exceeds N-1 and + ``_create_occupancy_groups`` walks off the end of ``initial_occ``. Roughly one + PDB-REDO entry in six carries an atom with no coordinates or no B and hit this. + """ + import pandas as pd + + from torchref.model.model import Model + + src = Model(verbose=0) + src.load_pdb(str(pdb_dir / "3GR5.pdb")) + df = src.pdb.copy() + n_before = len(df) + + # Blank the B of a few interior atoms so the dropna removes them. + victims = [5, 100, 500] + df.loc[victims, "tempfactor"] = float("nan") + cell = src.cell.data.cpu().numpy() + sg = src.spacegroup + + model = Model(verbose=0) + model.load(lambda: (df, cell, sg), add_hydrogens=False) + + assert len(model.pdb) == n_before - len(victims) + idx = model.pdb["index"].to_numpy() + assert idx.min() == 0 + assert idx.max() == len(model.pdb) - 1, "index must stay positional after a drop" + assert sorted(idx) == list(range(len(model.pdb))) + # The occupancy grouping is what actually indexed past the end. + assert model.occupancy().shape[0] == len(model.pdb) From f477f312e82a82646a8ebdddf822a66477058a78 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 19:27:13 +0200 Subject: [PATCH 089/250] Consolidate epsilon onto SpaceGroup with a convention switch The alignment package carried its own `compute_epsilon`, which disagreed with `Symmetry.epsilon` on trigonal and hexagonal groups and dropped the centring coset entirely. Three implementations of one crystallographic quantity, and the rotation function and the ML rescore were reading different ones. Epsilon now lives on the space group, where it can be asked for by anything that holds one. It takes the HKL and returns the multiplicity. The two call sites want genuinely different conventions, so the switch is explicit rather than a second function. Operations mapping h -> h add coherently and set the mean, <|F|^2> = eps * Sigma; operations mapping h -> -h leave the mean alone and make F real, which changes the distribution. That second effect is centricity and `is_centric` already carries it, so the likelihood's V = eps - sigma_A^2 wants the conventional count (`friedel=False`). sigma_A estimation is calibrated against the Friedel-folded count, so that stays the default and the refinement path is untouched. `_build_llg_context` and `m_letf1_rescore` now take the space group rather than a matrix stack, since they needed it to derive epsilon anyway. Gate: 440 passed / 12 skipped. The rescore panel is the load-bearing one -- this moves epsilon on every centric reflection, 12360 of them on 2DQ6 -- and it is paired against the old convention as its own arm over 10 structures x 3 trials: 12 of 30 cells move, 7 better and 5 worse, sum +3 ranks against swings of +-7. Neutral, as a convention correction that is not yet load-bearing should be. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/e_conformance.py | 11 +- alignment_lab/analysis/e_table.py | 4 +- alignment_lab/analysis/eps_gate.sh | 22 ++++ alignment_lab/analysis/fw_footing.py | 93 ++++++++++++++ alignment_lab/analysis/fw_footing.sh | 17 +++ alignment_lab/analysis/llg_decompose.py | 2 +- alignment_lab/analysis/panel_ranks.py | 11 ++ alignment_lab/analysis/rebaseline_panel.sh | 25 ++++ alignment_lab/analysis/rescore_prep_arms.py | 11 ++ .../diagnostics/frf_normaliser_anatomy.py | 3 +- alignment_lab/lab/rescore.py | 2 +- tests/unit/alignment/test_m_letf1.py | 11 +- .../alignment/test_symmetry_conventions.py | 7 +- .../unit/symmetry/test_epsilon_conventions.py | 121 ++++++++++++++++++ .../alignment/frf/preprocessing.py | 42 +----- .../experimental/alignment/ml_rotation.py | 27 ++-- torchref/experimental/alignment/pipeline.py | 6 +- torchref/symmetry/symmetry.py | 44 +++++-- 18 files changed, 381 insertions(+), 78 deletions(-) create mode 100644 alignment_lab/analysis/eps_gate.sh create mode 100644 alignment_lab/analysis/fw_footing.py create mode 100644 alignment_lab/analysis/fw_footing.sh create mode 100644 alignment_lab/analysis/rebaseline_panel.sh create mode 100644 tests/unit/symmetry/test_epsilon_conventions.py diff --git a/alignment_lab/analysis/e_conformance.py b/alignment_lab/analysis/e_conformance.py index a2aaf1f5..a985fbcd 100644 --- a/alignment_lab/analysis/e_conformance.py +++ b/alignment_lab/analysis/e_conformance.py @@ -185,7 +185,16 @@ def check_e_convention( if getattr(cls, "uses_sigma_f", False) and sig_F is not None: loud = cls(F, s_mag, centric, sig_F=sig_F * 4.0, eps=eps, n_shells=n_shells).E.to(torch.float64) - ref = conv.sigma.sqrt().to(torch.float64) # the shell scale E sits on + # The target of the shrinkage is the shell mean IN E-SPACE. Using + # sqrt(Sigma) here instead -- the scale E was divided BY -- compares a + # dimensionless quantity of order 1 against one in units of F, so every + # reflection sits far below the reference and "shrinkage" degenerates + # into "did E get bigger". + ref = torch.zeros_like(E) + for k in range(int(dec.max()) + 1): + m = dec == k + if bool(m.any()): + ref[m] = E[m].mean() moved_closer = (loud - ref).abs() <= (E - ref).abs() + 1e-9 rep["shrinkage_frac_ok"] = float(moved_closer.to(torch.float64).mean()) else: diff --git a/alignment_lab/analysis/e_table.py b/alignment_lab/analysis/e_table.py index b3ba9a93..3e253b2a 100644 --- a/alignment_lab/analysis/e_table.py +++ b/alignment_lab/analysis/e_table.py @@ -31,7 +31,6 @@ def main() -> int: CalcGlobalE, CalcShellE, FrenchWilsonE, SmoothSigmaE, WilsonShellE, WilsonShellEpsE, ) - from torchref.experimental.alignment.frf.preprocessing import compute_epsilon model, data = load_case(args.pdb)[:2] rec = data.cell.reciprocal_basis_matrix.to(torch.float64) @@ -40,8 +39,7 @@ def main() -> int: sig = None if getattr(data, "F_sigma", None) is None else \ data.F_sigma.to(torch.float64) cen = data.centric.to(torch.bool) - eps = compute_epsilon(data.hkl.to(torch.long), - data.spacegroup.matrices.to(torch.float64)) + eps = data.spacegroup.epsilon(data.hkl.to(torch.long), friedel=False) # A calc set from the deposited coordinates: the "perfect model" case, where # obs and calc genuinely should land on the same scale. diff --git a/alignment_lab/analysis/eps_gate.sh b/alignment_lab/analysis/eps_gate.sh new file mode 100644 index 00000000..faa398ff --- /dev/null +++ b/alignment_lab/analysis/eps_gate.sh @@ -0,0 +1,22 @@ +#!/bin/bash +# Gate for routing the rescore's epsilon through SpaceGroup.epsilon(friedel=False). +# This CHANGES the LLG on every centric reflection, so the unit suite is necessary +# but not sufficient -- the rescore panel is the real test. +#SBATCH --job-name=epsgate +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO"; export PYTHONPATH="$REPO" +export TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname)" +L=alignment_lab/slurm/epsgate_tests_$SLURM_JOB_ID.log +"$PY" -m pytest tests/unit/symmetry tests/unit/alignment tests/unit/frf_separate tests/unit/model -q > "$L" 2>&1 +rc=$?; echo "=== TESTS rc=$rc ==="; tail -12 "$L"; grep "^FAILED" "$L" | head -8 diff --git a/alignment_lab/analysis/fw_footing.py b/alignment_lab/analysis/fw_footing.py new file mode 100644 index 00000000..e6393887 --- /dev/null +++ b/alignment_lab/analysis/fw_footing.py @@ -0,0 +1,93 @@ +"""Is French-Wilson's `eEobs` supposed to have unit mean square? + +The conformance table reports `obs_calc_ratio` ~2.1 for `FrenchWilsonE`, i.e. +`` sits near 0.5 where the calc companion sits at 1. Two readings, and +they call for different actions: + +* the FW port is mis-scaled -- a real defect in the FRF's observed side; or +* `eEobs` is a DEFLATED amplitude by construction and the check is asking the + wrong question of it. + +`eEobs**2 = eEsqFW + (DFAC**2 - 1)/DFAC**2` with `DFAC < 1`, so the second term +is strictly negative: the assembly subtracts the share of the measured intensity +that is measurement error. That is a deconvolution, not a normalisation, and +`eEobs` travels with `DFAC` as a pair. + +So the question is not "is `` one" but "is the quantity the CONSUMER +forms centred". The consumer is `build_lerf1_obs_intensity`, which forms +`cw * (eEobs**2 - 1) * DFAC**2` -- it subtracts a literal 1. If `` is +really 0.5, that term carries a systematic negative offset into the Patterson +correlation, and whether that matters is a separate question from whether the +port is faithful. + +This decomposes the assembly term by term, per resolution decile, so the answer +comes from the numbers rather than from reading the formula. +""" + +from __future__ import annotations + +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import load_case # noqa: E402 + + +def main() -> int: + from torchref.experimental.alignment.frf.french_wilson import ( + french_wilson_preprocess, + ) + from torchref.experimental.alignment.frf.preprocessing import ( + build_lerf1_obs_intensity, wilson_normalise, + ) + + for pdb in ("1DAW", "3K7M", "2DQ6"): + model, data = load_case(pdb) + F = data.work.F.to(torch.float64).cpu() + sig = data.work.sigF.to(torch.float64).cpu() + hkl = data.hkl.cpu() + s = data.cell.s_magnitude(hkl).to(torch.float64).cpu() + cen = data.spacegroup.is_centric(hkl).cpu().to(torch.bool) + keep = torch.isfinite(F) & torch.isfinite(sig) & (sig > 0) & (F > 0) + F, sig, s, cen = F[keep], sig[keep], s[keep], cen[keep] + + fw = french_wilson_preprocess(F, sig, s, cen, n_wilson_shells=20) + eE, dfac = fw["eEobs"].to(torch.float64), fw["DFAC"].to(torch.float64) + # eEsqFW is what eEobs**2 would be before the deconvolution term. + corr = (dfac * dfac - 1.0) / (dfac * dfac) + eEsq = eE * eE - corr + wil = wilson_normalise(F, s, cen, n_shells=20).to(torch.float64) + lerf = build_lerf1_obs_intensity( + fw["eEobs"], cen, dfac=fw["DFAC"], use_centric_weight=True, + ).to(torch.float64) + + order = torch.argsort(s) + dec = torch.zeros_like(s, dtype=torch.long) + ch = max(1, s.numel() // 10) + for k in range(10): + hi = (k + 1) * ch if k < 9 else s.numel() + dec[order[k * ch:hi]] = k + + print(f"\n=== {pdb} n={F.numel()} centric={int(cen.sum())} ===") + print(f" {'dec':>3s} {'d(A)':>12s} {'':>9s} {'':>9s} " + f"{'':>10s} {'':>7s} {'':>9s}") + for k in range(10): + m = dec == k + lo, hi = float(1 / s[m].max()), float(1 / s[m].min()) + print(f" {k:>3d} {f'{hi:5.1f}-{lo:4.2f}':>12s} " + f"{float((wil[m] ** 2).mean()):>9.4f} " + f"{float(eEsq[m].mean()):>9.4f} " + f"{float((eE[m] ** 2).mean()):>10.4f} " + f"{float(dfac[m].mean()):>7.4f} {float(lerf[m].mean()):>9.4f}") + print(f" {'ALL':>3s} {'':>12s} {float((wil ** 2).mean()):>9.4f} " + f"{float(eEsq.mean()):>9.4f} {float((eE ** 2).mean()):>10.4f} " + f"{float(dfac.mean()):>7.4f} {float(lerf.mean()):>9.4f}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/alignment_lab/analysis/fw_footing.sh b/alignment_lab/analysis/fw_footing.sh new file mode 100644 index 00000000..b20e4fcf --- /dev/null +++ b/alignment_lab/analysis/fw_footing.sh @@ -0,0 +1,17 @@ +#!/bin/bash +#SBATCH --job-name=fwfoot +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:30:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=32G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 +export CUDA_VISIBLE_DEVICES="" +"$PY" -u alignment_lab/analysis/fw_footing.py +echo "RC=$?" diff --git a/alignment_lab/analysis/llg_decompose.py b/alignment_lab/analysis/llg_decompose.py index 77525409..c53d85f2 100644 --- a/alignment_lab/analysis/llg_decompose.py +++ b/alignment_lab/analysis/llg_decompose.py @@ -98,7 +98,7 @@ def main() -> int: apply_bulk_solvent=True, apply_wilson_b=True) ctx = _build_llg_context( inp.F_obs, inp.hkl, inp.s_mag, inp.centric, inp.ll, data.cell, - data.spacegroup.matrices.to(torch.float64).to(inp.device), + data.spacegroup, n_shells=max(20 // 2, 8), batch_size=50, **prep, ) diff --git a/alignment_lab/analysis/panel_ranks.py b/alignment_lab/analysis/panel_ranks.py index aa17d56a..dfee4a41 100644 --- a/alignment_lab/analysis/panel_ranks.py +++ b/alignment_lab/analysis/panel_ranks.py @@ -37,12 +37,21 @@ def main() -> int: ap.add_argument("--n-peaks", type=int, default=500) ap.add_argument("--thr-deg", type=float, default=5.0) ap.add_argument("--tag", default="?", help="which tree this run came from") + ap.add_argument("--exclude-h", action="store_true", + help="drop hydrogens from F_calc. dev now keeps them by " + "default and generates any a file lacks, which is ~47% " + "of the atom count and moves |F_calc| by ~7% at the " + "median. Whether GENERATED hydrogens belong in a " + "molecular-replacement search model is a separate " + "question from whether they belong in refinement.") args = ap.parse_args() cfg = FRFConfig(n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) for trial in range(args.trials): seed = seed_for(args.pdb, trial) model, data, R_true = rotated_case(args.pdb, seed) + if args.exclude_h: + model.exclude_H_from_sf = True t0 = time.time() res = run_frf(model, data, cfg, capture_arf=False, verbose=0) seconds = time.time() - t0 @@ -57,7 +66,9 @@ def main() -> int: # sort as a good rank, so for comparison it counts as worse than the # worst hit -- the peak-list length. rank_cmp = rank if rank >= 0 else args.n_peaks + n_h = int((model.pdb["element"].str.strip() == "H").sum()) print(f"ROW {args.tag} {args.pdb} trial={trial} seed={seed} " + f"nH={n_h} exclH={int(args.exclude_h)} " f"rank={rank} rank_cmp={rank_cmp} found={int(rank >= 0)} " f"top20={int(0 <= rank < 20)} " f"angle={'' if ang is None else round(float(ang), 3)} " diff --git a/alignment_lab/analysis/rebaseline_panel.sh b/alignment_lab/analysis/rebaseline_panel.sh new file mode 100644 index 00000000..52e081e9 --- /dev/null +++ b/alignment_lab/analysis/rebaseline_panel.sh @@ -0,0 +1,25 @@ +#!/bin/bash +# Re-establish the FRF panel after the merge, and settle the hydrogen question in +# the same pass. Every previously published FRF number was measured on +# hydrogen-free structure factors; dev now keeps hydrogens by default, so the old +# 98/100 is not a baseline any more. +#SBATCH --job-name=rebase +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-9 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname) pdb=$PDB" +"$PY" -u alignment_lab/analysis/panel_ranks.py --pdb "$PDB" --trials 10 --tag withH 2>/dev/null | grep '^ROW ' +"$PY" -u alignment_lab/analysis/panel_ranks.py --pdb "$PDB" --trials 10 --tag noH --exclude-h 2>/dev/null | grep '^ROW ' diff --git a/alignment_lab/analysis/rescore_prep_arms.py b/alignment_lab/analysis/rescore_prep_arms.py index a2a7c5f6..5081743f 100644 --- a/alignment_lab/analysis/rescore_prep_arms.py +++ b/alignment_lab/analysis/rescore_prep_arms.py @@ -42,9 +42,14 @@ #: Each arm is a set of overrides on top of `m_letf1`'s defaults. They are #: cumulative on purpose: if the whole Phaser prep helps, the interesting #: question is which piece carries it. +#: `eps_friedel` reproduces the epsilon the rescore used before it was routed +#: through `SpaceGroup.epsilon(friedel=False)`. Passing an explicit `eps_factor` +#: is how the old convention is reproduced without a second worktree, so the +#: comparison stays paired on one FRF peak list. ARMS = { "none": None, # control: FRF order "default": {}, # what ships today + "eps_friedel": {"__eps_friedel": True}, # the pre-migration convention "vrms": {"vrms_strategy": "oeffner"}, "solvent": {"apply_bulk_solvent": True}, "vrms_solvent": {"vrms_strategy": "oeffner", "apply_bulk_solvent": True}, @@ -90,6 +95,12 @@ def main() -> int: rank, seconds = frf_rank, 0.0 else: kw = dict(overrides) + if kw.pop("__eps_friedel", False): + # Friedel-folded epsilon: doubles it on every centric + # reflection, which is what the rescore used to get. + kw["eps_factor"] = data.spacegroup.epsilon( + res.inputs.hkl.to(torch.long), friedel=True, + ).to(res.inputs.F_obs.dtype) if kw.get("vrms_strategy") == "oeffner": kw["vrms_n_residues"] = n_residues t0 = time.time() diff --git a/alignment_lab/diagnostics/frf_normaliser_anatomy.py b/alignment_lab/diagnostics/frf_normaliser_anatomy.py index 78427fb1..e1407996 100644 --- a/alignment_lab/diagnostics/frf_normaliser_anatomy.py +++ b/alignment_lab/diagnostics/frf_normaliser_anatomy.py @@ -148,7 +148,6 @@ def _fit(A: torch.Tensor, b: torch.Tensor): def run(pdb: str, dumps: Path) -> dict: - from torchref.experimental.alignment.frf.preprocessing import compute_epsilon cap, data = capture_ours(pdb) sg = data.spacegroup.matrices.to(torch.float64).cpu() @@ -173,7 +172,7 @@ def run(pdb: str, dumps: Path) -> dict: n = int(y.numel()) var_tot = float(y.var(unbiased=False)) - eps = compute_epsilon(hkl, sg) + eps = sg.epsilon(hkl.to(torch.long), friedel=False) # Phaser divides intensity by eps_n and we do not, so its Esqr should be # SMALLER by that factor: log ratio carries -log(eps). y_eps = y + torch.log(eps) diff --git a/alignment_lab/lab/rescore.py b/alignment_lab/lab/rescore.py index b5f3a473..05366fcd 100644 --- a/alignment_lab/lab/rescore.py +++ b/alignment_lab/lab/rescore.py @@ -114,7 +114,7 @@ def run_rescore( out = m_letf1_rescore( subset, frf_inputs.F_obs, frf_inputs.hkl, frf_inputs.s_mag, frf_inputs.centric, frf_inputs.ll, data.cell, - data.spacegroup.matrices.to(torch.float64).to(device), + data.spacegroup, **common, **engine_kwargs, ) else: diff --git a/tests/unit/alignment/test_m_letf1.py b/tests/unit/alignment/test_m_letf1.py index 0546ef47..9259f52d 100644 --- a/tests/unit/alignment/test_m_letf1.py +++ b/tests/unit/alignment/test_m_letf1.py @@ -105,7 +105,8 @@ def test_m_letf1_rescore_runs_and_ranks_truth_top(): F_obs = (1.0 + 0.1 * torch.randn(N, dtype=torch.float64)).abs() hkl = torch.randint(-10, 10, (N, 3), dtype=torch.long) centric = torch.zeros(N, dtype=torch.bool) - sym_mats = torch.eye(3, dtype=torch.float64).unsqueeze(0) # P1: only identity + from torchref.symmetry import SpaceGroup + sg = SpaceGroup("P 1") # only the identity, so epsilon is 1 throughout # Stub interpolator: returns F_obs (perfectly correlated) for R = identity, # uncorrelated noise for any other R. @@ -142,7 +143,7 @@ def reciprocal_basis_matrix(self): peaks = [truth] + random_peaks rescored = m_letf1_rescore( - peaks, F_obs, hkl, s_mag, centric, StubLL(), StubCell(), sym_mats, + peaks, F_obs, hkl, s_mag, centric, StubLL(), StubCell(), sg, n_shells=5, batch_size=4, ) # Truth (identity) should be among the top 3 rescored peaks (truth=identity @@ -160,7 +161,6 @@ def test_scat_mode_absolute_preserves_calc_intershell_shape(): inter-shell weighting on a synthetic with a strong resolution-dependent F_calc falloff.""" from torchref.experimental.alignment.ml_rotation import _build_llg_context, _llg_for_orientations - from torchref.experimental.alignment.frf.preprocessing import compute_epsilon N = 300 torch.manual_seed(3) @@ -168,7 +168,8 @@ def test_scat_mode_absolute_preserves_calc_intershell_shape(): F_obs = (1.0 + 0.1 * torch.randn(N, dtype=torch.float64)).abs() hkl = torch.randint(-12, 12, (N, 3), dtype=torch.long) centric = torch.zeros(N, dtype=torch.bool) - sym_mats = torch.eye(3, dtype=torch.float64).unsqueeze(0) # P1 + from torchref.symmetry import SpaceGroup + sg = SpaceGroup("P 1") # F_calc with a strong B-factor falloff → big inter-shell amplitude variation. decay = torch.exp(-40.0 * s_mag * s_mag) @@ -189,7 +190,7 @@ def reciprocal_basis_matrix(self): return torch.eye(3, dtype=torch.float64) common = dict( - interpolator=StubLL(), real_cell=StubCell(), sym_mats=sym_mats, + interpolator=StubLL(), real_cell=StubCell(), spacegroup=sg, n_shells=6, batch_size=64, ) ctx_leg = _build_llg_context(F_obs, hkl, s_mag, centric, scat_mode="legacy", **common) diff --git a/tests/unit/alignment/test_symmetry_conventions.py b/tests/unit/alignment/test_symmetry_conventions.py index f847ed77..6e6fe912 100644 --- a/tests/unit/alignment/test_symmetry_conventions.py +++ b/tests/unit/alignment/test_symmetry_conventions.py @@ -23,7 +23,6 @@ import torch from torchref.experimental.alignment.frf.preprocessing import ( - compute_epsilon, epsilon_aware_unroll, ) from torchref.experimental.alignment.sh import ( @@ -188,7 +187,7 @@ def test_symmetrised_anisotropy_obeys_the_lattice(hm): @pytest.mark.parametrize("hm, non_orthogonal", SPACEGROUPS) def test_epsilon_uses_the_row_vector_convention(hm, non_orthogonal): - """``compute_epsilon`` counts ops fixing ``h``, which needs ``h·R``. + """The multiplicity counts ops fixing ``h``, which needs ``h·R``. Reflections on a symmetry axis must come out with multiplicity > 1; with the wrong convention the wrong reflections are flagged. @@ -199,13 +198,13 @@ def test_epsilon_uses_the_row_vector_convention(hm, non_orthogonal): g = torch.Generator().manual_seed(3) hkl = torch.randint(-9, 10, (400, 3), generator=g) - eps = compute_epsilon(hkl, S).detach().cpu() + eps = sg.epsilon(hkl, friedel=False).detach().cpu() assert int(eps.min()) >= 1 # Recompute independently through the shared helper. ref = sg.expand_reciprocal(hkl).detach().cpu() # (ops, N, 3) same = (ref == hkl.to(torch.int64)).all(dim=-1) assert torch.equal(eps.to(torch.long), same.sum(dim=0).clamp(min=1)), ( - f"{hm}: compute_epsilon disagrees with expand_reciprocal" + f"{hm}: epsilon(friedel=False) disagrees with expand_reciprocal" ) diff --git a/tests/unit/symmetry/test_epsilon_conventions.py b/tests/unit/symmetry/test_epsilon_conventions.py new file mode 100644 index 00000000..78a84bf5 --- /dev/null +++ b/tests/unit/symmetry/test_epsilon_conventions.py @@ -0,0 +1,121 @@ +"""``Symmetry.epsilon`` carries two conventions, and they must not drift together. + +Folding Friedel mates into epsilon mixes two different effects. Operations that +map ``h -> h`` add coherently and set the **mean**, ``<|F|^2> = eps * Sigma`` -- +the conventional crystallographic epsilon. Operations that map ``h -> -h`` leave +the mean alone and make ``F`` real, which changes the **distribution**; that is +centricity, and ``is_centric`` already carries it. + +Both settings have a consumer: sigma_A estimation is calibrated against the +Friedel-folded default, and the molecular-replacement likelihood wants the +conventional count for its ``V = eps - sigma_A**2``. So the risk is not that one +is wrong -- it is that a later edit quietly makes them the same, or flips which +one is the default, and nothing notices. These pin the difference, its location, +and its size. + +Counts are the measured ones rather than recomputed, so a regression shows up as +a specific wrong number instead of a green test that changed convention. +""" + +import pytest +import torch + +from torchref.symmetry import SpaceGroup + +pytestmark = pytest.mark.unit + +#: ``(hm, centred)``. Spans primitive and centred lattices, and the point groups +#: where the alignment package's own epsilon was measured to disagree. +SPACEGROUPS = [ + ("C 1 2 1", True), + ("P 21 21 2", False), + ("P 31 2 1", False), + ("P 43 21 2", False), + ("P 4 3 2", False), + ("P 65 2 2", False), +] + + +def _hkl(n=4000, seed=11): + g = torch.Generator().manual_seed(seed) + hkl = torch.randint(-12, 13, (n, 3), generator=g) + return hkl[(hkl.abs().sum(dim=-1) > 0)] # drop (0,0,0) + + +@pytest.mark.parametrize("hm, centred", SPACEGROUPS) +def test_friedel_changes_centric_reflections_and_only_those(hm, centred): + """The switch must move exactly the centric reflections, never a general one. + + This is the property that makes the two conventions safe to hold at once: if + the difference ever spread beyond centrics, one of them would have stopped + meaning what its docstring says. + """ + sg = SpaceGroup(hm) + hkl = _hkl() + with_f = sg.epsilon(hkl, friedel=True) + without = sg.epsilon(hkl, friedel=False) + centric = sg.is_centric(hkl).to(torch.bool) + + differs = with_f != without + assert not bool((differs & ~centric).any()), ( + f"{hm}: epsilon differs on {int((differs & ~centric).sum())} ACENTRIC " + f"reflections; the Friedel term must only reach centrics" + ) + # And the difference is a doubling where it lands, not an arbitrary shift. + if bool(differs.any()): + ratio = (with_f[differs] / without[differs]) + assert torch.allclose(ratio, torch.full_like(ratio, 2.0)), ( + f"{hm}: Friedel folding is not a factor of two where it applies" + ) + + +@pytest.mark.parametrize("hm, centred", SPACEGROUPS) +def test_the_conventional_count_is_never_larger(hm, centred): + sg = SpaceGroup(hm) + hkl = _hkl() + assert bool((sg.epsilon(hkl, friedel=False) + <= sg.epsilon(hkl, friedel=True)).all()) + + +@pytest.mark.parametrize("hm, centred", SPACEGROUPS) +def test_centring_cosets_are_counted_by_both(hm, centred): + """A centred lattice gives every reflection the centring order as a factor. + + Not asserted as desirable -- it is a documented property with one consequence, + that epsilon is inflated wherever it is used as a *term* rather than a factor. + Pinned so the behaviour is deliberate rather than discovered again. + """ + sg = SpaceGroup(hm) + hkl = _hkl() + eps = sg.epsilon(hkl, friedel=False) + general_min = float(eps.min()) + if centred: + assert general_min >= 2.0, ( + f"{hm} is centred; every reflection should carry the centring order " + f"but the minimum epsilon is {general_min}" + ) + else: + assert general_min == 1.0, ( + f"{hm} is primitive; general reflections should have epsilon 1, got " + f"{general_min}" + ) + + +def test_the_default_is_the_calibrated_one(): + """sigma_A estimation is calibrated against Friedel-folded epsilon. + + Flipping this default would decalibrate the refinement path silently, so the + default is pinned separately from the behaviour of either branch. + """ + sg = SpaceGroup("P 21 21 2") + hkl = _hkl() + assert torch.equal(sg.epsilon(hkl), sg.epsilon(hkl, friedel=True)) + + +def test_epsilon_is_at_least_one_and_finite(): + for hm, _ in SPACEGROUPS: + sg = SpaceGroup(hm) + for friedel in (True, False): + eps = sg.epsilon(_hkl(), friedel=friedel) + assert bool(torch.isfinite(eps).all()) + assert float(eps.min()) >= 1.0 diff --git a/torchref/experimental/alignment/frf/preprocessing.py b/torchref/experimental/alignment/frf/preprocessing.py index d27ce2db..7b80d4a2 100644 --- a/torchref/experimental/alignment/frf/preprocessing.py +++ b/torchref/experimental/alignment/frf/preprocessing.py @@ -76,7 +76,6 @@ def eterm_sigma_a(s_mag: torch.Tensor, delta_vrms_A: float) -> torch.Tensor: __all__ = [ "wilson_normalise", "wilson_normalise_epsilon", - "compute_epsilon", "eterm_sigma_a", "french_wilson_preprocess", "get_high_order_axis", @@ -161,45 +160,6 @@ def epsilon_aware_unroll( return unrolled_hkl, asu_idx -def compute_epsilon( - hkl: torch.Tensor, - sym_mats: torch.Tensor, -) -> torch.Tensor: - """Reflection multiplicity ε(h) — the order of the stabilizer subgroup. - - Phaser source: the ``epsn`` array in ``DataMR.cc`` (used in - ``SIGMAN.sqrt_epsnSN``, DataMR.cc:925) — same role as - ``cctbx::miller::index_span`` epsilons. - - ε(h) = number of point-group rotation operators W for which - ``h · W = h`` (row-vector convention, no Friedel). For a general - reflection ε = 1; reflections on an n-fold symmetry axis get ε = n. - - Used to epsilon-correct Wilson normalisation: axial reflections are - systematically stronger (``⟨I_h⟩ = ε_h · Σ``), so without the - correction they over-weight the m = 0 SH column and bias the - rotation-function map for high-symmetry spacegroups. - - Parameters - ---------- - hkl : (N, 3) integer-valued (any dtype) Miller indices. - sym_mats : (n_ops, 3, 3) integer rotation operators (fractional/lattice - rotation parts of the spacegroup). - - Returns - ------- - epsilon : (N,) float — multiplicity ε ≥ 1. - """ - h = hkl.to(torch.float64) - W = sym_mats.to(torch.float64) - eps = torch.zeros(h.shape[0], dtype=torch.float64, device=h.device) - for k in range(W.shape[0]): - h_t = h @ W[k] # row-vector: h' = h · W - same = (h_t.round() == h).all(dim=-1) - eps += same.to(torch.float64) - return eps.clamp(min=1.0) - - def wilson_normalise_epsilon( F: torch.Tensor, s_mag: torch.Tensor, @@ -348,7 +308,7 @@ def compute_v_budget( eps_factor : (N,) tensor Per-reflection ε(h), the multiplicity (1 for general positions, n>1 for reflections on n-fold symmetry axes). From - :func:`compute_epsilon`. + :meth:`torchref.symmetry.symmetry.Symmetry.epsilon`. sigma_a : (N,) tensor Per-reflection σ_A(s) (interpolated from the per-shell fit). n_mol : int diff --git a/torchref/experimental/alignment/ml_rotation.py b/torchref/experimental/alignment/ml_rotation.py index fb90a751..dfb908f1 100644 --- a/torchref/experimental/alignment/ml_rotation.py +++ b/torchref/experimental/alignment/ml_rotation.py @@ -576,7 +576,7 @@ def _build_llg_context( centric: torch.Tensor, interpolator: LattmanLoveInterpolator, real_cell, - sym_mats: torch.Tensor, + spacegroup, *, n_shells: int = 20, batch_size: int = 50, @@ -608,7 +608,6 @@ def _build_llg_context( form which falls off ~3× too fast. """ from .frf.preprocessing import ( - compute_epsilon, compute_v_budget, epsilon_aware_unroll, eterm_sigma_a, @@ -616,12 +615,21 @@ def _build_llg_context( device = F_obs.device dtype = F_obs.dtype + sym_mats = spacegroup.matrices.to(torch.float64).to(device) n_ops = int(sym_mats.shape[0]) N = hkl_real.shape[0] # 1. ε(h) per reflection (needed for the ε-corrected obs normalisation). if eps_factor is None: - eps_factor = compute_epsilon(hkl_real.to(torch.long), sym_mats).to(dtype) + # `friedel=False`: the conventional count. The variance budget + # `V = eps - sigma_A**2` wants operations that add coherently and set the + # mean; operations mapping h -> -h change the DISTRIBUTION instead, which + # the Woolfson branch below already handles. Counting them here doubles + # epsilon on every centric reflection -- 6680 of them on 3K7M -- and + # inflates exactly those reflections' variance. + eps_factor = spacegroup.epsilon( + hkl_real.to(torch.long), friedel=False, + ).to(dtype) eps_factor = eps_factor.to(device) # 2. Per-shell ε-corrected Wilson E_obs (Phaser E = F/sqrt(ε·Σ_N)). Dividing @@ -1071,7 +1079,7 @@ def m_letf1_rescore( centric: torch.Tensor, interpolator: LattmanLoveInterpolator, real_cell, - sym_mats: torch.Tensor, + spacegroup, *, n_shells: int = 20, n_refine: Optional[int] = None, @@ -1127,15 +1135,18 @@ def m_letf1_rescore( interpolator, real_cell ``LattmanLoveInterpolator`` for the model molecular transform and the crystal real cell. - sym_mats : (n_ops, 3, 3) tensor - Spacegroup rotation operators in the reciprocal (hkl) basis. + spacegroup : SpaceGroup + The crystal's space group. Passed as the object rather than its + ``matrices`` because the multiplicity this needs is a method on it: + ``epsilon(hkl, friedel=False)``, the conventional count, which a bare + tensor of rotations cannot answer. sigma_a : (N,) tensor, optional Per-reflection σ_A. If ``None``, fitted on-the-fly from the identity rotation's |F_calc| via :func:`fit_sigma_a_per_shell` and interpolated per shell. eps_factor : (N,) tensor, optional Per-reflection multiplicity ε(h). If ``None``, computed via - :func:`torchref.experimental.alignment.frf.preprocessing.compute_epsilon`. + :meth:`torchref.symmetry.symmetry.Symmetry.epsilon` with ``friedel=False``. n_refine, batch_size, verbose As in :func:`sim_mlrf_rescore`. """ @@ -1147,7 +1158,7 @@ def m_letf1_rescore( tail = peaks[n_refine:] ctx = _build_llg_context( - F_obs, hkl_real, s_mag, centric, interpolator, real_cell, sym_mats, + F_obs, hkl_real, s_mag, centric, interpolator, real_cell, spacegroup, n_shells=n_shells, batch_size=batch_size, sigma_a=sigma_a, eps_factor=eps_factor, apply_bulk_solvent=apply_bulk_solvent, solvent_fsol=solvent_fsol, solvent_bsol=solvent_bsol, diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 7efa2ab8..9c84f492 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -557,7 +557,7 @@ def _rotation_candidates(self, frf) -> list: if self.rescore_engine == "m_letf1": rescored = m_letf1_rescore( peaks, F_obs, hkl, s_mag, centric, ll, data.cell, - data.spacegroup.matrices.to(torch.float64).to(device), + data.spacegroup, n_shells=rescore_n_shells, n_refine=min(len(peaks), self.n_ml_refine), batch_size=50, verbose=self.verbose, @@ -588,7 +588,7 @@ def _subpeak_refine(self, rescored, F_obs, hkl, s_mag, centric, ll, self._timer.start("4b_subpeak_refine") ctx = _build_llg_context( F_obs, hkl, s_mag, centric, ll, data.cell, - data.spacegroup.matrices.to(torch.float64).to(device), + data.spacegroup, n_shells=rescore_n_shells, batch_size=50, scat_mode=self.rescore_scat_mode, ) @@ -872,7 +872,7 @@ def _dense_rotation_refine(self, refined): rescored_refine = m_letf1_rescore( cand_peaks, F_obs_amp, hkl_keep, s_mag_keep, centric_keep, ll_refine, data.cell, - data.spacegroup.matrices.to(torch.float64).to(device), + data.spacegroup, n_shells=rescore_n_shells, n_refine=len(cand_peaks), batch_size=rescore_batch, verbose=0, ) diff --git a/torchref/symmetry/symmetry.py b/torchref/symmetry/symmetry.py index 5d68cffc..bc02c5d7 100644 --- a/torchref/symmetry/symmetry.py +++ b/torchref/symmetry/symmetry.py @@ -416,13 +416,15 @@ def is_absent(self, hkl: torch.Tensor) -> torch.Tensor: absent = (maps_to_self & non_integral).any(dim=0) return absent.reshape(original_shape).to(hkl.device) - def epsilon(self, hkl: torch.Tensor) -> torch.Tensor: - """Reflection multiplicity: operations mapping ``h`` to ``h`` or to ``-h``. + def epsilon(self, hkl: torch.Tensor, *, friedel: bool = True) -> torch.Tensor: + """Reflection multiplicity: operations mapping ``h`` to ``h``, or also to ``-h``. Parameters ---------- hkl : torch.Tensor Miller indices, shape ``(N, 3)``. + friedel : bool, default True + Whether operations mapping ``h -> -h`` count alongside ``h -> h``. Returns ------- @@ -432,18 +434,42 @@ def epsilon(self, hkl: torch.Tensor) -> torch.Tensor: Notes ----- - Friedel mates are folded in unconditionally, with no centric/acentric branch. - That inflates the count relative to the conventional epsilon, which counts pure - rotational multiplicity (``h -> h``) only. Downstream sigma_A estimation is - calibrated against this convention. + The two settings answer different questions and both are wanted. + + Operations mapping ``h -> h`` add coherently and set the **mean**, + ``<|F|^2> = eps * Sigma``. That is the conventional crystallographic epsilon, + and what a Wilson normalisation or a likelihood's variance budget asks for. + Operations mapping ``h -> -h`` leave the mean alone and instead make ``F`` + real, changing the **distribution** from exponential to chi2_1 -- that is + centricity, and :meth:`is_centric` already carries it. Folding Friedel into + epsilon therefore mixes a mean effect with a distribution effect. + + The default keeps Friedel folded in because downstream sigma_A estimation is + calibrated against that convention; flipping it would silently decalibrate the + refinement path. Pass ``friedel=False`` for the conventional count, as the + molecular-replacement likelihood does -- counting Friedel there doubles + epsilon on exactly the reflections whose distribution the Woolfson branch is + already handling, inflating their ``V = eps - sigma_A**2``. + + The two differ on centric reflections and *only* there: measured across the + ten benchmark structures every disagreement was centric, and the counts are + not small -- 12360 reflections on 2DQ6, 7555 on 4BX9, 6680 on 3K7M. + + Both settings count lattice-centring cosets, so on a centred lattice every + reflection carries the centring order as a factor: C2 gives 2 for general + reflections where a primitive lattice gives 1. That is a separate axis from + this switch. Being uniform per lattice it is absorbed into ``Sigma`` wherever + epsilon is a factor -- which is why the refinement path never saw it -- and + bites only where epsilon is a term. """ float_dtype = get_float_dtype() with torch.no_grad(): equivalents = self.expand_reciprocal(hkl) # (n_ops, N, 3) target = hkl.to(device=equivalents.device, dtype=torch.int64) - same = (equivalents == target).all(dim=-1) - friedel = (equivalents == -target).all(dim=-1) - eps = (same | friedel).sum(dim=0).clamp(min=1).to(float_dtype) + fixes = (equivalents == target).all(dim=-1) + if friedel: + fixes = fixes | (equivalents == -target).all(dim=-1) + eps = fixes.sum(dim=0).clamp(min=1).to(float_dtype) return eps.to(hkl.device) # ========================================================================= From f6bd44a1b7e01a003763a03710ecbbc746e178d0 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 19:48:20 +0200 Subject: [PATCH 090/250] Route the rotation function's E values through a convention class `E = F/sqrt(Sigma(s))` is a weighting choice wearing the costume of a units change: correlating E_obs against E_calc IS correlating F against F with weight 1/Sigma(s). The alignment package had nine private ways into E-space, each an undeclared answer to that one question, and the rotation function and the ML rescore were reading different ones on the same data. `EConvention` is the seam. Engines take the CLASS, not an instance -- a fitted Sigma(s) cannot exist before the reflections do -- and instantiate it twice, once per side. That the same class has to normalise both is the point: it puts obs and calc on a common footing under test rather than under assumption. `build_lerf1_intensity` now takes `weight=`, already squared, instead of `dfac=`. Deciding what the weight IS belongs to the convention; applying one belongs here. Two defects fell out of writing it down. `WilsonShellE` divided the shell mean by eps without dividing the intensity by it, leaving <|E|^2> = : 1 on a primitive lattice and 2 on a centred one, so a normaliser whose absolute scale depended on the space group. `CalcGlobalE` had it too. The rotation function never saw either, because it passes no eps -- which is exactly why they survived. eps now enters in one place, `_intensity()`, which both the numerator and the shell mean come through, so it reaches both sides of the ratio or neither. `functools.partial` is the documented way to configure a convention, and a partial forwards __call__ but not class attributes. Asking one for `for_calc` raised, so all three smooth-Sigma arms of the first panel failed on all 50 cells without producing a number. Lookups go through `convention_class` now. Inertness is proved rather than sampled: the seam replaced exactly three tensors -- eEobs, the LERF1 weight, E_calc -- and every line downstream is untouched, so bit-identity on those three IS bit-identity of the peak list. Worst deviation 0.0 over 1DAW/3K7M/2DQ6/4BX9. The functional panel agrees independently: naming the default convention explicitly moves 0 of 50 cells. Gate: 1936 passed / 92 skipped, the two known dev-side test_device_conformance failures aside. 22 new unit tests pin the invariants that must never break -- eps reaching both sides, scale invariance, no residual resolution trend, and that a partial stays usable. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/e_convention_arms.py | 140 +++++++++++++++ alignment_lab/analysis/e_convention_arms.sh | 25 +++ alignment_lab/analysis/fw_footing.py | 17 +- alignment_lab/analysis/seam_gate.sh | 24 +++ alignment_lab/analysis/seam_identity.py | 111 ++++++++++++ alignment_lab/lab/__init__.py | 4 +- alignment_lab/lab/frf.py | 29 ++- tests/unit/alignment/test_e_conventions.py | 169 ++++++++++++++++++ torchref/experimental/alignment/e_values.py | 67 ++++++- torchref/experimental/alignment/frf/api.py | 39 ++-- .../alignment/frf/preprocessing.py | 20 ++- .../experimental/alignment/rotation_search.py | 12 ++ 12 files changed, 625 insertions(+), 32 deletions(-) create mode 100644 alignment_lab/analysis/e_convention_arms.py create mode 100644 alignment_lab/analysis/e_convention_arms.sh create mode 100644 alignment_lab/analysis/seam_gate.sh create mode 100644 alignment_lab/analysis/seam_identity.py create mode 100644 tests/unit/alignment/test_e_conventions.py diff --git a/alignment_lab/analysis/e_convention_arms.py b/alignment_lab/analysis/e_convention_arms.py new file mode 100644 index 00000000..5b3e7abd --- /dev/null +++ b/alignment_lab/analysis/e_convention_arms.py @@ -0,0 +1,140 @@ +"""Which E convention ranks the true orientation best? + +Layer B of the E-value work: the conformance table says whether a convention +does what E is *supposed* to do; this says whether it makes the rotation +function *work*. When the two disagree, this one decides and the table +diagnoses -- a convention with clean Wilson statistics that ranks truth worse is +not the one to ship, and the table then names the property the winner trades +away. + +Headline metric is the fraction of cells at **rank 0**, not "inside the top 20". +The stated target is that truth comes first; the post-merge baseline is 20 of +100, so the bar is a long way up and a metric that saturates hides the climb. + +Paired by seed: every convention sees the same rotated case from the same +``seed_for``, so arms are compared cell by cell rather than as two distributions. +The ``default`` arm passes no convention at all, which makes it a control on the +*seam* as well as on the conventions -- if it ever diverges from the production +number, the plumbing changed something rather than the convention did. + +The FRF cannot be run once and shared here, unlike the rescore arms: the +convention is what builds the obs expansion, so each arm is a full run. That is +the cost of the question. +""" + +from __future__ import annotations + +import argparse +import functools +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, FRFConfig, e_convention_name, # noqa: E402 + orbit_rank, rotated_case, run_frf, seed_for) + + +def build_arms(): + """Name -> convention class. Built lazily so ``--help`` needs no torchref.""" + from torchref.experimental.alignment.e_values import ( + CalcGlobalE, CalcShellE, FrenchWilsonE, SmoothSigmaE, WilsonShellE, + WilsonShellEpsE, + ) + # Mixed arms. The panel's first round put `calc_global` -- a single global + # RMS on BOTH sides -- ahead of every per-shell convention, which if real + # says the per-shell flattening is discarding inter-shell amplitude shape + # the correlation was using. One class sets both sides, so isolating which + # side carries that needs conventions that differ across the seam. Defined + # here rather than shipped: they exist to answer one question. + class GlobalObsShellCalc(CalcGlobalE): + """Global RMS on obs, per-shell Wilson on calc.""" + calc_companion = WilsonShellE + + class ShellObsGlobalCalc(WilsonShellE): + """Per-shell Wilson on obs, global RMS on calc.""" + calc_companion = CalcGlobalE + + class FrenchWilsonGlobalCalc(FrenchWilsonE): + """Production obs side, global RMS on calc.""" + calc_companion = CalcGlobalE + + return { + # Control: no convention passed, so the production default applies. + "default": None, + # The same thing named explicitly. Must match `default` exactly; if it + # does not, the seam is not inert and nothing below means anything. + "french_wilson": FrenchWilsonE, + # Drops the measurement-error model entirely -- the size of the gap to + # `french_wilson` is what sigma_F is worth to the rotation function. + "wilson": WilsonShellE, + # What the rescore uses on its observed side. Running it here asks + # whether the FRF/rescore disagreement is costing the FRF anything. + "wilson_eps": WilsonShellEpsE, + "calc_shell": CalcShellE, + "calc_global": CalcGlobalE, + # The divergence candidate: a smooth Chebyshev Sigma(s) instead of + # per-shell means. Two orders, because the whole question is whether a + # low-order curve beats 20 independent bins. + "smooth4": functools.partial(SmoothSigmaE, n_coeff=4), + "smooth6": functools.partial(SmoothSigmaE, n_coeff=6), + "smooth10": functools.partial(SmoothSigmaE, n_coeff=10), + "global_x_shell": GlobalObsShellCalc, + "shell_x_global": ShellObsGlobalCalc, + "fw_x_global": FrenchWilsonGlobalCalc, + } + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) + ap.add_argument("--trials", type=int, default=10) + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--thr-deg", type=float, default=5.0) + ap.add_argument("--arms", default="") + args = ap.parse_args() + + arms = build_arms() + names = [a for a in args.arms.split(",") if a] or list(arms) + unknown = [a for a in names if a not in arms] + if unknown: + raise SystemExit(f"unknown arms {unknown}; have {sorted(arms)}") + + for trial in range(args.trials): + seed = seed_for(args.pdb, trial) + model, data, R_true = rotated_case(args.pdb, seed) + sym = data.spacegroup.matrices.to(torch.float64).cpu() + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + okw = dict(side="left", frame="cart", reciprocal_basis=rec, + thr_deg=args.thr_deg) + + for name in names: + cfg = FRFConfig(n_peaks=args.n_peaks, lmax_cap=args.lmax_cap, + e_convention=arms[name]) + t0 = time.time() + try: + res = run_frf(model, data, cfg, capture_arf=False, verbose=0) + except Exception as exc: # a convention may refuse + print(f"ROW {name} {args.pdb} trial={trial} seed={seed} " + f"rank=-1 rank_cmp={args.n_peaks} found=0 top20=0 " + f"seconds=0.00 error={type(exc).__name__}", flush=True) + continue + seconds = time.time() - t0 + rank, ang = orbit_rank(res.peaks, R_true, sym, **okw) + rank_cmp = rank if rank >= 0 else args.n_peaks + print(f"ROW {name} {args.pdb} trial={trial} seed={seed} " + f"rank={rank} rank_cmp={rank_cmp} found={int(rank >= 0)} " + f"top20={int(0 <= rank < 20)} " + f"angle={'' if ang is None else round(float(ang), 3)} " + f"seconds={seconds:.2f} " + f"conv={e_convention_name(arms[name])}", flush=True) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/alignment_lab/analysis/e_convention_arms.sh b/alignment_lab/analysis/e_convention_arms.sh new file mode 100644 index 00000000..799f14ea --- /dev/null +++ b/alignment_lab/analysis/e_convention_arms.sh @@ -0,0 +1,25 @@ +#!/bin/bash +# Which E convention ranks truth best? Nine arms x 5 trials x 10 structures. +# the same pass. Every previously published FRF number was measured on + + +#SBATCH --job-name=earms +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=day +#SBATCH --time=03:00:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-9 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname) pdb=$PDB" +"$PY" -u alignment_lab/analysis/e_convention_arms.py --pdb "$PDB" --trials 10 2>/dev/null | grep '^ROW ' + diff --git a/alignment_lab/analysis/fw_footing.py b/alignment_lab/analysis/fw_footing.py index e6393887..548709f8 100644 --- a/alignment_lab/analysis/fw_footing.py +++ b/alignment_lab/analysis/fw_footing.py @@ -42,16 +42,17 @@ def main() -> int: french_wilson_preprocess, ) from torchref.experimental.alignment.frf.preprocessing import ( - build_lerf1_obs_intensity, wilson_normalise, + build_lerf1_intensity, wilson_normalise, ) for pdb in ("1DAW", "3K7M", "2DQ6"): model, data = load_case(pdb) - F = data.work.F.to(torch.float64).cpu() - sig = data.work.sigF.to(torch.float64).cpu() + F = data.F.to(torch.float64).abs().cpu() + sig = data.F_sigma.to(torch.float64).cpu() hkl = data.hkl.cpu() - s = data.cell.s_magnitude(hkl).to(torch.float64).cpu() - cen = data.spacegroup.is_centric(hkl).cpu().to(torch.bool) + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + s = (hkl.to(torch.float64) @ rec).norm(dim=-1) + cen = data.centric.cpu().to(torch.bool) keep = torch.isfinite(F) & torch.isfinite(sig) & (sig > 0) & (F > 0) F, sig, s, cen = F[keep], sig[keep], s[keep], cen[keep] @@ -60,9 +61,9 @@ def main() -> int: # eEsqFW is what eEobs**2 would be before the deconvolution term. corr = (dfac * dfac - 1.0) / (dfac * dfac) eEsq = eE * eE - corr - wil = wilson_normalise(F, s, cen, n_shells=20).to(torch.float64) - lerf = build_lerf1_obs_intensity( - fw["eEobs"], cen, dfac=fw["DFAC"], use_centric_weight=True, + wil = wilson_normalise(F, s, 20)[0].to(torch.float64) + lerf = build_lerf1_intensity( + fw["eEobs"], cen, weight=dfac * dfac, use_centric_weight=True, ).to(torch.float64) order = torch.argsort(s) diff --git a/alignment_lab/analysis/seam_gate.sh b/alignment_lab/analysis/seam_gate.sh new file mode 100644 index 00000000..a2e2185b --- /dev/null +++ b/alignment_lab/analysis/seam_gate.sh @@ -0,0 +1,24 @@ +#!/bin/bash +#SBATCH --job-name=seamgate +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:50:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +# Single-threaded: the FRF peak list only reproduces bit-for-bit at one thread +# (~5e-8 score noise reorders 12 of 500 peaks on 3GR5 otherwise), and this gate +# is about bit-identity. +export TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 +echo "== pytest ==" +"$PY" -m pytest tests/unit -q 2>&1 | tail -12 +echo "PYTEST_RC=${PIPESTATUS[0]}" +echo "== seam identity ==" +"$PY" -u alignment_lab/analysis/seam_identity.py 2>/dev/null | grep -E "^(case|SEAM|[0-9A-Z]{4} )" +echo "RC=$?" diff --git a/alignment_lab/analysis/seam_identity.py b/alignment_lab/analysis/seam_identity.py new file mode 100644 index 00000000..90b5494c --- /dev/null +++ b/alignment_lab/analysis/seam_identity.py @@ -0,0 +1,111 @@ +"""Is the E-convention seam inert? + +Routing the rotation function's normalisation through an `EConvention` class is +only safe to build on if it changes nothing while the default is in place. A +peak-list hash would answer that, but it needs a second worktree to compare +against and it only samples the structures it is run on. + +This is stronger and cheaper. The seam replaced exactly three tensors -- +`eEobs`, the LERF1 weight, and `E_calc` -- and every line downstream of them is +untouched. So bit-identity on those three is not evidence that the peak list is +unchanged, it is a proof of it, on whatever data this is run over. + +The one at real risk is the calc side. `wilson_normalise` accumulates its shell +sums with `index_add_` and clamps the mean at 1e-12; `EConvention` uses +`scatter_add_` and clamps at 1e-30. Same arithmetic in exact terms, and on CPU +both reduce in index order -- but "should be identical" is the claim under test, +not the assumption behind it. + +Reports max absolute and relative deviation rather than a bare pass/fail, so a +non-zero result says how big it is instead of only that it exists. +""" + +from __future__ import annotations + +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import load_case # noqa: E402 + +CASES = ("1DAW", "3K7M", "2DQ6", "4BX9") + + +def _dev(new: torch.Tensor, old: torch.Tensor): + """``(n_differing, max_abs, max_rel)`` between two tensors.""" + d = (new.to(torch.float64) - old.to(torch.float64)).abs() + rel = d / old.to(torch.float64).abs().clamp(min=1e-30) + return int((d > 0).sum()), float(d.max()), float(rel.max()) + + +def main() -> int: + from torchref.experimental.alignment.e_values import ( + FrenchWilsonE, WilsonShellE, + ) + from torchref.experimental.alignment.frf.french_wilson import ( + french_wilson_preprocess, + ) + from torchref.experimental.alignment.frf.preprocessing import ( + build_lerf1_intensity, wilson_normalise, + ) + from torchref.experimental.alignment.sh import ( + assign_shells, equal_count_shell_edges, + ) + + n_shells = 20 + worst = 0.0 + print(f"{'case':>6s} {'tensor':>12s} {'n':>8s} {'n_diff':>7s} " + f"{'max abs':>10s} {'max rel':>10s}") + for pdb in CASES: + model, data = load_case(pdb) + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + hkl = data.hkl.cpu() + s = (hkl.to(torch.float64) @ rec).norm(dim=-1) + F = data.F.to(torch.float64).abs().cpu() + sig = data.F_sigma.to(torch.float64).cpu() + cen = data.centric.cpu().to(torch.bool) + keep = torch.isfinite(F) & torch.isfinite(sig) & (sig > 0) & (F > 0) + F, sig, s, cen = F[keep], sig[keep], s[keep], cen[keep] + + edges, _ = equal_count_shell_edges(s, n_shells) + shell_idx = assign_shells(s, edges) + + # --- obs side ------------------------------------------------- + fw = french_wilson_preprocess(F, sig, s, cen, n_wilson_shells=n_shells, + shell_idx=shell_idx) + conv = FrenchWilsonE(F, s, cen, sig_F=sig, shell_idx=shell_idx, + n_shells=n_shells) + for name, new, old in ( + ("eEobs", conv.E, fw["eEobs"]), + ("weight", conv.weight, fw["DFAC"] * fw["DFAC"]), + ("lerf1", + build_lerf1_intensity(conv.E, cen, weight=conv.weight), + build_lerf1_intensity(fw["eEobs"], cen, + weight=fw["DFAC"] * fw["DFAC"])), + ): + nd, a, r = _dev(new, old) + worst = max(worst, a) + print(f"{pdb:>6s} {name:>12s} {new.numel():>8d} {nd:>7d} " + f"{a:>10.3e} {r:>10.3e}") + + # --- calc side: the one where the two implementations differ ---- + # A stand-in calc set; the check is about the normaliser, not the model. + F_calc = (F * 1.37 + 5.0) + old_E, _ = wilson_normalise(F_calc, s, n_shells) + new_E = WilsonShellE(F_calc, s, cen, n_shells=n_shells).E + nd, a, r = _dev(new_E, old_E) + worst = max(worst, a) + print(f"{pdb:>6s} {'E_calc':>12s} {new_E.numel():>8d} {nd:>7d} " + f"{a:>10.3e} {r:>10.3e}") + + print(f"\nSEAM_WORST_ABS {worst:.6e}") + print("SEAM_INERT" if worst == 0.0 else "SEAM_NOT_INERT") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/alignment_lab/lab/__init__.py b/alignment_lab/lab/__init__.py index b0b14c54..aa0601e8 100644 --- a/alignment_lab/lab/__init__.py +++ b/alignment_lab/lab/__init__.py @@ -27,7 +27,8 @@ fit_aniso_log_space, tensor_report, ) -from .frf import (FRFConfig, FRFResult, merge_peak_lists, patched, +from .frf import (FRFConfig, FRFResult, e_convention_name, + merge_peak_lists, patched, run_frf) from .rescore import ENGINES, RescoreResult, paired_ranks, run_rescore from .profile import (FRF_STAGES, PeakMemory, calibration_seconds, @@ -50,6 +51,7 @@ "fit_aniso_log_space", "tensor_report", "FRFConfig", + "e_convention_name", "FRFResult", "merge_peak_lists", "patched", diff --git a/alignment_lab/lab/frf.py b/alignment_lab/lab/frf.py index 8c5f8cbc..ea90e953 100644 --- a/alignment_lab/lab/frf.py +++ b/alignment_lab/lab/frf.py @@ -15,6 +15,18 @@ import torch +def e_convention_name(conv) -> str: + """Display name for a convention class, a ``partial`` of one, or ``None``.""" + if conv is None: + return "default" + inner = getattr(conv, "func", conv) + name = getattr(inner, "__name__", str(inner)) + kw = getattr(conv, "keywords", None) + if kw: + name += "(" + ",".join(f"{k}={v}" for k, v in sorted(kw.items())) + ")" + return name + + @dataclass class FRFConfig: """Engine settings for one FRF evaluation. @@ -33,12 +45,20 @@ class FRFConfig: #: Expected r.m.s. coordinate error, in Angstrom. ``None`` uses the Oeffner #: estimate from the model's length, which is what the pipeline does. model_error_A: Optional[float] = None + #: E-value convention, as the CLASS the engine instantiates once per side. + #: ``None`` leaves the production default in place; a class (or a + #: ``functools.partial`` of one) sweeps it. Unlike the deleted ``extra`` + #: knobs this is a real production parameter, so the lab passes it through + #: rather than patching a constant. + e_convention: Optional[type] = None extra: Dict[str, Any] = field(default_factory=dict) def as_row(self) -> Dict[str, Any]: """Config fields for a result row (``extra`` flattened out).""" d = asdict(self) d.pop("extra") + # `asdict` cannot render a class or a partial; name it instead. + d["e_convention"] = e_convention_name(self.e_convention) d.update(self.extra) return d @@ -211,6 +231,11 @@ def _wrapped(self, *args, **kwargs): f"torchref.experimental.alignment.rotation_search instead." ) + # Omitted rather than passed as None, so an unset convention takes the + # production default from the signature instead of overriding it with one. + conv_kw = {} if cfg.e_convention is None else { + "e_convention": cfg.e_convention} + t0 = time.time() with patched(_rs, "LMAX_CAP", int(cfg.lmax_cap)), \ patched(_rs, "DENSE_CALC_PAD", float(cfg.dense_pad)), \ @@ -220,12 +245,12 @@ def _wrapped(self, *args, **kwargs): with patched(_api.FastRotationFunction, "score_model", _wrapped): peaks, _lmax, _dmin = _rs.search_peaks( model, data, model_error_A, U_aniso=frf_inputs.U_aniso, - n_peaks=cfg.n_peaks, verbose=verbose, + n_peaks=cfg.n_peaks, verbose=verbose, **conv_kw, ) else: peaks, _lmax, _dmin = _rs.search_peaks( model, data, model_error_A, U_aniso=frf_inputs.U_aniso, - n_peaks=cfg.n_peaks, verbose=verbose, + n_peaks=cfg.n_peaks, verbose=verbose, **conv_kw, ) seconds = time.time() - t0 diff --git a/tests/unit/alignment/test_e_conventions.py b/tests/unit/alignment/test_e_conventions.py new file mode 100644 index 00000000..8760f76f --- /dev/null +++ b/tests/unit/alignment/test_e_conventions.py @@ -0,0 +1,169 @@ +"""Invariants every E convention has to hold, whatever it does inside. + +The conformance harness in `alignment_lab` reports on all of these and more, as +a table, over real data. These are the subset that must never break: a failure +here is a bug rather than a trade-off, so they belong in the gate rather than in +a report someone has to read. + +The epsilon check is here because it already caught one. `WilsonShellE` divided +the shell mean by ``eps`` without dividing the intensity by it, which left +``<|E|**2> = `` -- 1 on a primitive lattice and 2 on a centred one, so a +normaliser whose absolute scale depended on the space group. The rotation +function never saw it (it passes no ``eps``), which is exactly why it survived: +a defect only reachable through an argument nobody was passing yet. +""" + +import functools + +import pytest +import torch + +from torchref.experimental.alignment.e_values import ( + CalcGlobalE, CalcShellE, FrenchWilsonE, SmoothSigmaE, WilsonShellE, + WilsonShellEpsE, +) + +pytestmark = pytest.mark.unit + +CONVENTIONS = [ + WilsonShellE, WilsonShellEpsE, CalcShellE, CalcGlobalE, FrenchWilsonE, + functools.partial(SmoothSigmaE, n_coeff=6), +] + + +def _name(c): + return getattr(getattr(c, "func", c), "__name__", str(c)) + + +def _wilson_data(n=20000, seed=1, centric_frac=0.1): + """Amplitudes drawn from the distribution the conventions assume.""" + g = torch.Generator().manual_seed(seed) + s = torch.rand(n, generator=g, dtype=torch.float64) * 0.4 + 0.05 + # |E|^2 ~ Exp(1) scaled by a resolution-dependent Sigma, so there is a real + # trend for a per-shell or smooth normaliser to have to remove. + sigma = torch.exp(-40.0 * s * s) * 2500.0 + 1.0 + E2 = -torch.log(torch.rand(n, generator=g, dtype=torch.float64).clamp(min=1e-12)) + F = (E2 * sigma).sqrt() + sig_F = F * 0.05 + 0.5 + centric = torch.zeros(n, dtype=torch.bool) + centric[: int(n * centric_frac)] = True + return F, s, centric, sig_F + + +def _build(cls, F, s, centric, sig_F, eps=None, n_shells=20): + kw = {"sig_F": sig_F} if getattr(cls, "uses_sigma_f", False) else {} + return cls(F, s, centric, eps=eps, n_shells=n_shells, **kw) + + +@pytest.mark.parametrize("cls", CONVENTIONS, ids=_name) +def test_epsilon_reaches_both_sides_of_the_ratio_or_neither(cls): + """A convention's scale must not depend on the lattice centring. + + ``eps`` is a constant 2 here, which is what a centred lattice gives every + reflection. Dividing the shell mean by it and not the intensity would show + up as ``<|E|**2>`` doubling -- the bug this pins. + """ + F, s, centric, sig_F = _wilson_data() + eps = torch.full_like(F, 2.0) + without = _build(cls, F, s, centric, sig_F).E + with_eps = _build(cls, F, s, centric, sig_F, eps=eps).E + m_without = float((without * without).mean()) + m_with = float((with_eps * with_eps).mean()) + assert m_with == pytest.approx(m_without, rel=1e-6), ( + f"{_name(cls)}: <|E|^2> moves from {m_without:.4f} to {m_with:.4f} when " + f"a uniform eps=2 is supplied. A uniform multiplicity cancels out of " + f"E**2 = (F**2/eps) / ; a change means eps reached only one " + f"side of that ratio." + ) + + +@pytest.mark.parametrize("cls", CONVENTIONS, ids=_name) +def test_invariant_to_a_global_rescale_of_F(cls): + """E is a ratio, so multiplying every amplitude must change nothing. + + This is the property that lets the rotation function ignore scale entirely, + and the one a convention with any absolute constant in it would break. + """ + F, s, centric, sig_F = _wilson_data() + base = _build(cls, F, s, centric, sig_F).E + for c in (1e-3, 1e3): + scaled = _build(cls, F * c, s, centric, sig_F * c).E + assert torch.allclose(scaled, base, rtol=1e-6, atol=1e-9), ( + f"{_name(cls)}: scaling F by {c:g} moved E by up to " + f"{float((scaled - base).abs().max()):.3e}" + ) + + +@pytest.mark.parametrize("cls", CONVENTIONS, ids=_name) +def test_the_normaliser_removes_the_resolution_trend(cls): + """<|E|**2> must not drift with resolution -- that trend IS the weighting.""" + F, s, centric, sig_F = _wilson_data() + E = _build(cls, F, s, centric, sig_F).E + order = torch.argsort(s) + means = [float((E[order[k::10]] ** 2).mean()) for k in range(10)] + lo, hi = min(means), max(means) + assert hi / max(lo, 1e-12) < 1.35, ( + f"{_name(cls)}: <|E|^2> ranges {lo:.3f}..{hi:.3f} across resolution " + f"deciles; the normaliser is leaving a trend behind" + ) + + +def test_a_sigma_f_convention_names_a_calc_companion(): + """There is no measurement error on a calc set, so it needs a stand-in.""" + assert FrenchWilsonE.uses_sigma_f + companion = FrenchWilsonE.for_calc() + assert companion is not FrenchWilsonE + assert not getattr(companion, "uses_sigma_f", False) + + +def test_french_wilson_refuses_to_run_without_sigmas(): + """Silently degrading to plain Wilson would hide the whole difference.""" + F, s, centric, _ = _wilson_data(n=2000) + with pytest.raises(ValueError, match="sig_F"): + FrenchWilsonE(F, s, centric, n_shells=20) + + +def test_the_frf_calc_path_is_bit_identical_to_wilson_normalise(): + """The seam must be inert while the default convention is in place. + + Everything downstream of ``E_calc`` is untouched, so bit-identity here is + what makes the peak list unchanged rather than merely similar. + """ + from torchref.experimental.alignment.frf.preprocessing import ( + wilson_normalise, + ) + + F, s, centric, _ = _wilson_data() + old, _ = wilson_normalise(F, s, 20) + new = WilsonShellE(F, s, centric, n_shells=20).E + assert torch.equal(new, old), ( + f"max deviation {float((new - old).abs().max()):.3e}" + ) + + +def test_a_partial_is_a_usable_convention(): + """Configuration rides in as ``functools.partial``, so lookups must survive it. + + A partial forwards ``__call__`` but not class attributes, so asking one for + ``for_calc`` or ``uses_sigma_f`` directly raises ``AttributeError``. The FRF + asks for both on every run -- which is why all three ``SmoothSigmaE`` arms + of the first convention panel failed on all 50 cells without producing a + single number. + """ + from torchref.experimental.alignment.e_values import ( + convention_class, convention_for_calc, convention_uses_sigma_f, + ) + + plain = functools.partial(SmoothSigmaE, n_coeff=6) + assert convention_class(plain) is SmoothSigmaE + assert convention_uses_sigma_f(plain) is False + # Its own companion, so the configuration has to survive the round trip. + assert convention_for_calc(plain) is plain + + fw = functools.partial(FrenchWilsonE) + assert convention_uses_sigma_f(fw) is True + # A different class is named, so its keywords are not this one's. + assert convention_for_calc(fw) is WilsonShellE + + F, s, centric, sig_F = _wilson_data(n=4000) + assert convention_for_calc(plain)(F, s, centric, n_shells=20).E.shape == F.shape diff --git a/torchref/experimental/alignment/e_values.py b/torchref/experimental/alignment/e_values.py index 7ecb1be2..d8fe6d4d 100644 --- a/torchref/experimental/alignment/e_values.py +++ b/torchref/experimental/alignment/e_values.py @@ -52,9 +52,39 @@ "SmoothSigmaE", "WilsonShellE", "WilsonShellEpsE", + "convention_class", + "convention_for_calc", + "convention_uses_sigma_f", ] +def convention_class(conv) -> type: + """The class behind a convention, which may be a ``functools.partial``. + + Configuration rides in as ``partial(SmoothSigmaE, n_coeff=6)``, and a + partial forwards ``__call__`` but not class attributes -- so asking one for + ``uses_sigma_f`` or ``for_calc`` raises. Every attribute lookup on a + convention goes through here for that reason. + """ + return getattr(conv, "func", conv) + + +def convention_for_calc(conv): + """The convention to normalise **calculated** amplitudes with. + + Keeps the partial's configuration when the class is its own companion, and + drops it when a different class is named -- another class's keywords are not + this one's. + """ + companion = convention_class(conv).calc_companion + return conv if companion is None else companion + + +def convention_uses_sigma_f(conv) -> bool: + """Whether ``conv`` consumes ``sig_F``, partial or not.""" + return bool(getattr(convention_class(conv), "uses_sigma_f", False)) + + class EConvention: """Normalised amplitudes, plus the per-reflection weight that goes with them. @@ -88,6 +118,14 @@ class EConvention: #: the shrinkage test for conventions that do not. uses_sigma_f: bool = False + #: Whether this convention divides the intensity by ``eps``. A convention + #: that does not must ignore it on BOTH sides of the ratio: applying it to + #: the shell mean alone leaves `` = ``, which is 2 on a centred + #: lattice and 1 on a primitive one -- a normaliser whose scale depends on + #: the space group. Declared rather than implied so the two halves cannot + #: drift apart again. + uses_epsilon: bool = True + #: The convention to use for **calculated** amplitudes, when it cannot be #: this one. A French-Wilson posterior is defined for observations only -- #: there is no measurement error on a calc set to shrink toward the mean -- @@ -144,9 +182,13 @@ def __init__( # -- helpers shared by the subclasses --------------------------------- def _intensity(self) -> torch.Tensor: - """``F**2 / eps`` -- the quantity whose shell mean is ``Sigma``.""" + """``F**2 / eps`` -- the quantity whose shell mean is ``Sigma``. + + Both the numerator of ``E**2`` and its shell mean come through here, so + ``uses_epsilon`` reaches the ratio consistently by construction. + """ I = self.F * self.F - if self.eps is not None: + if self.eps is not None and self.uses_epsilon: I = I / self.eps.clamp(min=1.0) return I @@ -174,9 +216,18 @@ class WilsonShellE(EConvention): """Plain per-shell Wilson: ``E = F / sqrt(_shell)``. What the rotation function uses on the calc side, and on the obs side when - the data carry no sigmas. Ignores measurement error entirely. + the data carry no sigmas. Ignores measurement error, and multiplicity with + it -- both deliberately. The calc side is a single molecular transform + sampled in a P1 box, where multiplicity has no meaning; the obs side gets + its multiplicity from the symmetry unroll, which puts each reflection into + the sum once per operation that reaches it. + + So an ``eps`` passed to this class is *ignored*, not half-applied. Use + :class:`WilsonShellEpsE` when it should count. """ + uses_epsilon = False + def _compute(self): return self.F / self.sigma.sqrt(), self._ones() @@ -206,9 +257,14 @@ class FrenchWilsonE(EConvention): Requires ``sig_F``. Falling back silently to plain Wilson would hide exactly the difference this class exists to make visible. + + ``french_wilson_preprocess`` takes no multiplicity, so neither does this -- + declared so the reported ``sigma`` describes what was actually done rather + than what the base class would have done. """ uses_sigma_f = True + uses_epsilon = False calc_companion = WilsonShellE def _compute(self): @@ -246,8 +302,13 @@ class CalcGlobalE(EConvention): The rescore's ``scat_mode="absolute"``. Keeps how much the model actually scatters per resolution instead of flattening it, which is what makes a relative Wilson-B correction meaningful rather than cancelled. + + Calculated amplitudes, so multiplicity does not apply -- see + :class:`WilsonShellE`. """ + uses_epsilon = False + def _compute(self): rms = self._intensity().mean().clamp(min=1e-30).sqrt() self.sigma = torch.full_like(self.F, float(rms * rms)) diff --git a/torchref/experimental/alignment/frf/api.py b/torchref/experimental/alignment/frf/api.py index d3ad1bdc..6f4441c0 100644 --- a/torchref/experimental/alignment/frf/api.py +++ b/torchref/experimental/alignment/frf/api.py @@ -21,6 +21,8 @@ import torch +from ..e_values import (FrenchWilsonE, convention_for_calc, + convention_uses_sigma_f) from .data_mr import bessel_sh_expand, cross_correlate_xi from .peak_finder import find_rotation_peaks from .preprocessing import ( @@ -28,8 +30,6 @@ build_lerf1_intensity, detect_zsymm, eterm_sigma_a, - french_wilson_preprocess, - wilson_normalise, ) from .sitelist_ang import evaluate_rotation_function from .types import AdaptiveRotationFunction, RotationPeak @@ -139,6 +139,7 @@ def __init__( grid_sampling_deg: float = 2.0, asu_idx: Optional[torch.Tensor] = None, s_mag_asu: Optional[torch.Tensor] = None, + e_convention: type = FrenchWilsonE, ): self.device = s_obs.device @@ -153,6 +154,12 @@ def __init__( self.delta_vrms_A = delta_vrms_A self.n_wilson_shells = n_wilson_shells self.grid_sampling_deg = grid_sampling_deg + # The class, not an instance: a fitted Sigma(s) cannot exist before the + # reflections do, so the convention is constructed here -- twice, once + # per side. That the same class has to normalise both is the point; + # it puts obs and calc on a common footing under test rather than under + # assumption. Pass `functools.partial(Cls, ...)` to configure one. + self.e_convention = e_convention # 1. Resolution window. # @@ -222,19 +229,23 @@ def __init__( # 3b. Wilson normalisation. With sigmas, through the French-Wilson # posterior, which handles the axial reflections; without them, plain # per-shell Wilson. - if sig_F_obs is not None: - fw = french_wilson_preprocess( - F_obs, sig_F_obs, smag_src, centric_obs, - n_wilson_shells=n_wilson_shells, shell_idx=obs_shell_idx, - ) - eEobs, dfac = fw["eEobs"], fw["DFAC"] - else: - eEobs, _ = wilson_normalise(F_obs, smag_src, n_wilson_shells) - dfac = torch.ones_like(eEobs) + # A convention that reads sigmas cannot run without them. Falling back + # to its own calc companion is the same choice the hardcoded branch made + # -- French-Wilson with sigmas, plain Wilson without -- just asked of the + # convention instead of assumed about it. + obs_cls = e_convention + if sig_F_obs is None and convention_uses_sigma_f(obs_cls): + obs_cls = convention_for_calc(obs_cls) + conv_obs = obs_cls( + F_obs, smag_src, centric_obs, sig_F=sig_F_obs, + shell_idx=obs_shell_idx, n_shells=n_wilson_shells, + ) + self._conv_obs = conv_obs # 4. LERF1 obs intensity, and the per-shell variance reweight. intensity_obs = build_lerf1_intensity( - eEobs, centric_obs, dfac=dfac, use_centric_weight=True, + conv_obs.E, centric_obs, weight=conv_obs.weight, + use_centric_weight=True, ) intensity_obs = apply_shell_variance_weights( intensity_obs, smag_src, n_var_shells=n_wilson_shells, @@ -281,7 +292,9 @@ def score_model( # measured to drop 0 of 339040 reflections on 3K7M and 0 of 271630 on # 1DAW. So take `s_calc` as given and only derive |s| from it. smag_calc = s_calc.norm(dim=-1) - E_calc, _ = wilson_normalise(F_calc, smag_calc, self.n_wilson_shells) + E_calc = convention_for_calc(self.e_convention)( + F_calc, smag_calc, n_shells=self.n_wilson_shells, + ).E eterm = eterm_sigma_a(smag_calc, self.delta_vrms_A) # Optional Babinet bulk-solvent factor: Phaser folds it into σ_A as # `σ_A_eff = solTerm(s²) · Luzzati(s², vrms)` (EnsemblePDB.cc:96-100). diff --git a/torchref/experimental/alignment/frf/preprocessing.py b/torchref/experimental/alignment/frf/preprocessing.py index 7b80d4a2..4d794230 100644 --- a/torchref/experimental/alignment/frf/preprocessing.py +++ b/torchref/experimental/alignment/frf/preprocessing.py @@ -195,16 +195,26 @@ def wilson_normalise_epsilon( def build_lerf1_intensity( eEobs: torch.Tensor, centric_obs: torch.Tensor, - dfac: Optional[torch.Tensor] = None, + weight: Optional[torch.Tensor] = None, use_centric_weight: bool = True, ) -> torch.Tensor: - """LERF1 observed intensity: ``cweight · (eEobs² − 1) · DFAC²``. + """LERF1 observed intensity: ``cweight · (eEobs² − 1) · weight``. Phaser source: ``DataMR::m_LETF1`` (DataMR.cc:1326-1431) — the intensity that gets fed into the Bessel-SH expansion. cweight is ε(h) · (1 for centric, 2 for acentric); we use the centric/acentric factor only (the ε(h) multiplicity is implicit in the symmetry reduction of the input reflection set). + + ``weight`` is the per-reflection information weight that travels with + ``eEobs`` -- ``DFAC**2`` for the French-Wilson convention, ones for a + convention that does not model measurement error. It arrives already + squared because it is the E convention that decides what the weight *is*; + this function's job is to apply one, not to know it came from a D factor. + + Note the ``- 1``: the LERF1 intensity is CENTRED, which is what makes + `` = 1`` load-bearing rather than cosmetic. A convention whose + mean square is not one puts a constant offset into every shell. """ if use_centric_weight: cw = torch.where( @@ -214,9 +224,9 @@ def build_lerf1_intensity( ) else: cw = torch.ones_like(eEobs) - if dfac is None: - dfac = torch.ones_like(eEobs) - return cw * (eEobs * eEobs - 1.0) * (dfac * dfac) + if weight is None: + weight = torch.ones_like(eEobs) + return cw * (eEobs * eEobs - 1.0) * weight def apply_shell_variance_weights( diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index 95cb20aa..a3eba577 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -28,6 +28,7 @@ import torch +from .e_values import FrenchWilsonE from .sh import ( apply_overall_anisotropy, assign_shells, @@ -220,6 +221,7 @@ def search_peaks( n_peaks: int, verbose: int = 0, device: Optional[torch.device] = None, + e_convention: type = FrenchWilsonE, ) -> Tuple[List["RotationPeak"], int, float]: """Run the rotation function, returning the engine's own peak list. @@ -369,6 +371,7 @@ def search_peaks( grid_sampling_deg=GRID_SAMPLING_DEG, asu_idx=asu_idx, s_mag_asu=s_mag_asu, + e_convention=e_convention, ) _arf, peaks = engine.score_model( s_calc, F_calc, n_peaks=n_peaks, @@ -414,6 +417,7 @@ def rotation_search( n_peaks: int = 500, verbose: int = 0, device: Optional[torch.device] = None, + e_convention: type = FrenchWilsonE, ) -> RotationSolutions: """Find the orientations of ``model`` consistent with ``data``. @@ -442,6 +446,13 @@ def rotation_search( Where to run. Default ``None`` takes ``data``'s device, moving ``model`` to match; an explicit value moves both. With neither carrying one, the configured default applies. + e_convention : type, optional + How amplitudes become E values, given as a class rather than an + instance: a fitted ``Sigma(s)`` cannot exist before the reflections do, + so the engine constructs it -- once for the observations and once for + the model, which is what puts the two on a common footing. The default + pairs the French-Wilson posterior on obs (it reads ``sigF``) with plain + per-shell Wilson on calc. ``functools.partial`` configures one. Returns ------- @@ -462,5 +473,6 @@ def rotation_search( peaks, lmax, d_min = search_peaks( model, data, model_error_A, U_aniso=U_aniso, n_peaks=n_peaks, verbose=verbose, device=device, + e_convention=e_convention, ) return _solutions(peaks, lmax, d_min, model_error_A) From c4f8e449b630ddec426a868fe49efe1499ce8bd1 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 28 Aug 2026 20:14:24 +0200 Subject: [PATCH 091/250] Give the ML rescore the sigmas, and one normaliser instead of two The rotation function computed the French-Wilson posterior from sigF and then threw the sigmas away: `FRFInputs` did not carry them, so the ML rescore -- a likelihood, where measurement error is not a detail -- had no access to measurement-error information at all. It does now, under the same anisotropy correction as F_obs, which is multiplicative, so F/sigma survives it intact. `scat_mode` is gone. "legacy" was CalcShellE and "absolute" was CalcGlobalE, so the rescore had two knobs for one decision -- which let its observed and calculated sides be normalised by unrelated rules, the disagreement this whole line of work started from. The convention answers for both sides now. Bit-identical at the default: `WilsonShellEpsE` on obs, `CalcShellE` on calc reproduce the deleted `_normalize_to_e_epsilon` and `_per_shell_sqrt_mean` exactly. The observed side carries multiplicity; the calculated side is a single molecular transform, where it has no meaning. MEASURED, and it settles the question the plan was built on. That plan predicted E-normalisation would be "nearly optional for the FRF and load-bearing for the rescore", because a correlation has a global scale to cancel and a likelihood does not. The reasoning is sound and the prediction is wrong on both halves. FRF, 12 conventions x 100 paired cells: median truth rank 2.0 for EVERY arm, rank-0 20-23/100, every sign test p >= 0.47. CalcGlobalE fails every distributional property in the conformance table -- per-decile <|E|^2> varying 1.7-3.9x, KS 0.33, acentric moment ratio 8.1 where 2 is ideal -- and ranks truth as well as the French-Wilson posterior. Rescore, 14 arms x 100 paired cells against the raw-FRF-order control: no arm beats it. Control 20/100 at rank 0, median 2.0; every rescore arm 16-18/100, median 3.0-4.0. Handing it the sigmas moves 3 of 100 cells, 1 better and 2 worse -- the plan called that "the cheapest single fix found so far". So the normalisation axis is closed for alignment ranking. What this work actually bought is correctness: nine private converters down to one mechanism, two real bugs, and a harness that can falsify the next idea faster. Gate: 1958 passed / 92 skipped, the two known dev-side test_device_conformance failures aside. Seam identity 0.000e+00 across 16 tensor comparisons on 1DAW/3K7M/2DQ6/4BX9. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/e_table.sh | 26 ++++--- alignment_lab/analysis/full_gate.sh | 19 +++++ alignment_lab/analysis/rescore_prep_arms.py | 31 +++++++- alignment_lab/analysis/rescore_prep_arms.sh | 6 +- alignment_lab/lab/rescore.py | 11 ++- tests/unit/alignment/test_m_letf1.py | 37 +++++---- torchref/experimental/alignment/align.py | 21 ++++- torchref/experimental/alignment/e_values.py | 7 ++ .../experimental/alignment/ml_rotation.py | 76 ++++++++----------- torchref/experimental/alignment/pipeline.py | 17 +++-- 10 files changed, 168 insertions(+), 83 deletions(-) create mode 100644 alignment_lab/analysis/full_gate.sh diff --git a/alignment_lab/analysis/e_table.sh b/alignment_lab/analysis/e_table.sh index 2abaa129..d69fec89 100644 --- a/alignment_lab/analysis/e_table.sh +++ b/alignment_lab/analysis/e_table.sh @@ -1,22 +1,28 @@ #!/bin/bash +# Layer A of the E-value work: the property report, paired with the functional +# panel that decides. Run over structures spanning primitive and centred +# lattices and low to high symmetry, because the epsilon and Wilson-shape +# properties are properties of real reflection sets. #SBATCH --job-name=etable -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err #SBATCH --partition=hour -#SBATCH --time=00:50:00 +#SBATCH --time=00:45:00 #SBATCH --cpus-per-task=4 #SBATCH --mem=48G #SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-4 set -uo pipefail REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +# 1DAW is C2 (centred, so eps != 1 everywhere), 2DQ6 is the tNCS case whose +# moment ratio reads ~5.5, 3K7M and 3A5V are the high-symmetry ends, 6G9X is +# where the rescore reproducibly fails. +PDBS=(1DAW 2DQ6 3K7M 3A5V 6G9X) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} cd "$REPO" export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname)" -for pdb in 1DAW 3K7M 6G9X 2DQ6; do - "$PY" -u alignment_lab/analysis/e_table.py --pdb "$pdb" 2>&1 \ - | grep -vE "UserWarning|FutureWarning|^ *from |^ *warnings\.|^Loaded |^LINK |^Wilson outlier|^found non|^FrenchWilson initialized|^ Reflections:|^ Resolution:|^ Space group|^ Centric:|^✓|^Parametrization|^French-Wilson input guard" - echo -done -echo "rc=$?" +echo "### $PDB on $(hostname)" +"$PY" -u alignment_lab/analysis/e_table.py --pdb "$PDB" 2>/dev/null +echo "RC=$?" diff --git a/alignment_lab/analysis/full_gate.sh b/alignment_lab/analysis/full_gate.sh new file mode 100644 index 00000000..4f647e43 --- /dev/null +++ b/alignment_lab/analysis/full_gate.sh @@ -0,0 +1,19 @@ +#!/bin/bash +#SBATCH --job-name=fullgate +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +export TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 +"$PY" -m pytest tests/unit -q 2>&1 | tail -14 +echo "PYTEST_RC=${PIPESTATUS[0]}" +echo "== seam identity ==" +"$PY" -u alignment_lab/analysis/seam_identity.py 2>/dev/null | grep -E "SEAM|conv|case|^ *[0-9A-Z]" diff --git a/alignment_lab/analysis/rescore_prep_arms.py b/alignment_lab/analysis/rescore_prep_arms.py index 5081743f..d4a18838 100644 --- a/alignment_lab/analysis/rescore_prep_arms.py +++ b/alignment_lab/analysis/rescore_prep_arms.py @@ -46,6 +46,31 @@ #: through `SpaceGroup.epsilon(friedel=False)`. Passing an explicit `eps_factor` #: is how the old convention is reproduced without a second worktree, so the #: comparison stays paired on one FRF peak list. +#: The E-convention arms are the reason this harness is being re-run. The +#: rotation function turned out to be INSENSITIVE to the convention -- 12 of +#: them, 100 paired cells, median rank 2.0 for every one -- which is what a +#: correlation should do, since a global scale cancels out of it. The LLG is a +#: likelihood and has no free scale to cancel, so if the convention matters +#: anywhere it matters here. `no_sigmas` is the control for that: it withholds +#: the sigmas the rescore has only just started receiving. +def _arms(): + from torchref.experimental.alignment.e_values import ( + CalcGlobalE, CalcShellE, FrenchWilsonE, SmoothSigmaE, WilsonShellE, + WilsonShellEpsE, + ) + import functools + return { + "fw_sigmas": {"e_convention": FrenchWilsonE}, + "no_sigmas": {"sig_F_obs": None}, + "wilson": {"e_convention": WilsonShellE}, + "calc_shell": {"e_convention": CalcShellE}, + "calc_global": {"e_convention": CalcGlobalE}, + "smooth6": {"e_convention": functools.partial(SmoothSigmaE, + n_coeff=6)}, + "eps_wilson": {"e_convention": WilsonShellEpsE}, + } + + ARMS = { "none": None, # control: FRF order "default": {}, # what ships today @@ -67,11 +92,13 @@ def main() -> int: ap.add_argument("--n-refine", type=int, default=20, help="rescore window: the top-N FRF peaks handed to the engine") ap.add_argument("--thr-deg", type=float, default=5.0) - ap.add_argument("--arms", default=",".join(ARMS)) + ap.add_argument("--arms", default="") args = ap.parse_args() + ARMS.update(_arms()) + cfg = FRFConfig(n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) - arms = [a for a in args.arms.split(",") if a] + arms = [a for a in args.arms.split(",") if a] or list(ARMS) for trial in range(args.trials): seed = seed_for(args.pdb, trial) diff --git a/alignment_lab/analysis/rescore_prep_arms.sh b/alignment_lab/analysis/rescore_prep_arms.sh index 06716660..2322d397 100644 --- a/alignment_lab/analysis/rescore_prep_arms.sh +++ b/alignment_lab/analysis/rescore_prep_arms.sh @@ -2,8 +2,8 @@ #SBATCH --job-name=resprep #SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out #SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 +#SBATCH --partition=day +#SBATCH --time=04:00:00 #SBATCH --cpus-per-task=4 #SBATCH --mem=48G #SBATCH --constraint=cpu_epyc9335 @@ -17,6 +17,6 @@ cd "$REPO" export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" echo "host=$(hostname) pdb=$PDB" -"$PY" -u alignment_lab/analysis/rescore_prep_arms.py --pdb "$PDB" --trials 3 2>&1 \ +"$PY" -u alignment_lab/analysis/rescore_prep_arms.py --pdb "$PDB" --trials 10 2>&1 \ | grep -E "^ROW |Error|Traceback|Warning: " echo "rc=$?" diff --git a/alignment_lab/lab/rescore.py b/alignment_lab/lab/rescore.py index 05366fcd..b8ba7df1 100644 --- a/alignment_lab/lab/rescore.py +++ b/alignment_lab/lab/rescore.py @@ -84,7 +84,7 @@ def run_rescore( verbose : int, optional Engine verbosity. **engine_kwargs - Passed through to the engine (e.g. ``scat_mode``). + Passed through to the engine (e.g. ``e_convention``). Returns ------- @@ -111,11 +111,18 @@ def run_rescore( t0 = time.time() if engine == "m_letf1": + # The sigmas now reach the rescore. They did not before: the FRF + # computed the French-Wilson posterior from them and then discarded + # them, leaving a likelihood with no measurement-error information. + # An explicit `sig_F_obs` in `engine_kwargs` still wins, so an arm can + # withhold them as a control. + kw = dict(engine_kwargs) + kw.setdefault("sig_F_obs", frf_inputs.sig_F) out = m_letf1_rescore( subset, frf_inputs.F_obs, frf_inputs.hkl, frf_inputs.s_mag, frf_inputs.centric, frf_inputs.ll, data.cell, data.spacegroup, - **common, **engine_kwargs, + **common, **kw, ) else: out = sim_mlrf_rescore( diff --git a/tests/unit/alignment/test_m_letf1.py b/tests/unit/alignment/test_m_letf1.py index 9259f52d..ccb16d41 100644 --- a/tests/unit/alignment/test_m_letf1.py +++ b/tests/unit/alignment/test_m_letf1.py @@ -154,13 +154,22 @@ def reciprocal_basis_matrix(self): ) -def test_scat_mode_absolute_preserves_calc_intershell_shape(): - """scat_mode='absolute' uses a single GLOBAL calc scale, so a shell where the - model scatters weakly keeps a small eImove; 'legacy' flattens every shell to - unit variance. Verify the two modes give different (and predictable) eImove +def test_a_global_calc_convention_preserves_intershell_shape(): + """`CalcGlobalE` uses a single GLOBAL calc scale, so a shell where the model + scatters weakly keeps a small eImove; `CalcShellE` flattens every shell to + unit variance. Verify the two give different (and predictable) eImove inter-shell weighting on a synthetic with a strong resolution-dependent - F_calc falloff.""" - from torchref.experimental.alignment.ml_rotation import _build_llg_context, _llg_for_orientations + F_calc falloff. + + This was `scat_mode="legacy"` vs `"absolute"`. The rescore had two knobs for + one decision -- how obs is normalised and how calc is -- which let the two + sides be normalised by unrelated rules; the E convention answers for both.""" + from torchref.experimental.alignment.e_values import ( + CalcGlobalE, WilsonShellEpsE, + ) + from torchref.experimental.alignment.ml_rotation import ( + _build_llg_context, _llg_for_orientations, + ) N = 300 torch.manual_seed(3) @@ -193,15 +202,17 @@ def reciprocal_basis_matrix(self): interpolator=StubLL(), real_cell=StubCell(), spacegroup=sg, n_shells=6, batch_size=64, ) - ctx_leg = _build_llg_context(F_obs, hkl, s_mag, centric, scat_mode="legacy", **common) - ctx_abs = _build_llg_context(F_obs, hkl, s_mag, centric, scat_mode="absolute", **common) - - # Legacy per-shell normaliser varies across shells (tracks the F_calc decay); - # absolute is a single constant. This is the definitional difference. - assert ctx_leg.sqrt_mean_per_m.std() > 1e-6, "legacy should vary per shell" + ctx_leg = _build_llg_context(F_obs, hkl, s_mag, centric, + e_convention=WilsonShellEpsE, **common) + ctx_abs = _build_llg_context(F_obs, hkl, s_mag, centric, + e_convention=CalcGlobalE, **common) + + # The per-shell normaliser varies across shells (tracks the F_calc decay); + # the global one is a single constant. This is the definitional difference. + assert ctx_leg.sqrt_mean_per_m.std() > 1e-6, "per-shell should vary per shell" assert torch.allclose( ctx_abs.sqrt_mean_per_m, ctx_abs.sqrt_mean_per_m[0] - ), "absolute should be a single global scale" + ), "the global convention should be a single scale" # Both modes still produce finite LLGs. a = torch.zeros(1, dtype=torch.float64) llg_leg = _llg_for_orientations(ctx_leg, a, a, a) diff --git a/torchref/experimental/alignment/align.py b/torchref/experimental/alignment/align.py index d4f632e8..43124da8 100644 --- a/torchref/experimental/alignment/align.py +++ b/torchref/experimental/alignment/align.py @@ -24,6 +24,7 @@ import torch from .lattman_love import LattmanLoveInterpolator +from .e_values import WilsonShellEpsE from .sh import ( apply_overall_anisotropy, assign_shells, @@ -170,8 +171,16 @@ class FRFInputs: overall anisotropy tensor and the Lattman-Love interpolator. The rotation search reads `U_aniso` and `device`; the ML rescore, translation search and rigid-body polish read the rest. + + ``sig_F`` carries the same anisotropy correction as ``F_obs``, which is a + multiplicative factor, so ``F/sigma`` is unchanged by it. It is here because + the rotation function computes the French-Wilson posterior from the sigmas + and then threw them away, leaving the ML rescore -- a likelihood, where + measurement error is not a detail -- with no access to them at all. + ``None`` when the data carry no sigmas. """ F_obs: torch.Tensor # (N,) anisotropy-corrected amplitudes + sig_F: Optional[torch.Tensor] # (N,) their sigmas, same correction hkl: torch.Tensor # (N, 3) integer Miller indices s_vec: torch.Tensor # (N, 3) reciprocal-space Cartesian s_mag: torch.Tensor # (N,) Å⁻¹ @@ -212,6 +221,9 @@ def _prepare_frf_inputs( f"for {n_shells} shells; widen the resolution range." ) F_obs = F_obs[keep].to(device) + sig_F = getattr(data, "F_sigma", None) + if sig_F is not None: + sig_F = sig_F.to(torch.float64)[keep].to(device) hkl = hkl_all[keep].to(device) s_vec = s_vec_all[keep].to(device) s_mag = s_mag_all[keep].to(device) @@ -238,6 +250,9 @@ def _prepare_frf_inputs( _sym_mats_cart = hkl_symops_to_cartesian(_sg_mats, rec_basis.to(device)) U_aniso = symmetrize_anisotropy(U_aniso, _sym_mats_cart) F_obs_aniso = apply_overall_anisotropy(F_obs, s_vec, U_aniso) + # Same multiplicative factor, so F/sigma survives the correction intact. + sig_F_aniso = (None if sig_F is None + else apply_overall_anisotropy(sig_F, s_vec, U_aniso)) ll = LattmanLoveInterpolator( model, padding_factor=ll_padding_factor, max_res_A=ll_max_res_A, @@ -246,6 +261,7 @@ def _prepare_frf_inputs( return FRFInputs( F_obs=F_obs_aniso, + sig_F=sig_F_aniso, hkl=hkl, s_vec=s_vec, s_mag=s_mag, @@ -290,7 +306,7 @@ def align_model_to_data( sigma_b: float = 0.0, model_error_A: Optional[float] = None, rescore_engine: str = "m_letf1", - rescore_scat_mode: str = "legacy", + rescore_e_convention: type = WilsonShellEpsE, subpeak_refine: bool = False, subpeak_refine_k: int = -1, subpeak_refine_step_deg: float = 1.5, @@ -324,7 +340,8 @@ def align_model_to_data( ll_max_res_A=ll_max_res_A, ll_padding_factor=ll_padding_factor, n_rotation_peaks=n_rotation_peaks, n_ml_refine=n_ml_refine, model_error_A=model_error_A, - rescore_engine=rescore_engine, rescore_scat_mode=rescore_scat_mode, + rescore_engine=rescore_engine, + rescore_e_convention=rescore_e_convention, auto_variance_weights=auto_variance_weights, use_interp_var=use_interp_var, subpeak_refine=subpeak_refine, subpeak_refine_k=subpeak_refine_k, diff --git a/torchref/experimental/alignment/e_values.py b/torchref/experimental/alignment/e_values.py index d8fe6d4d..9786a3cc 100644 --- a/torchref/experimental/alignment/e_values.py +++ b/torchref/experimental/alignment/e_values.py @@ -296,6 +296,13 @@ class CalcShellE(WilsonShellE): """ +#: The observed side carries multiplicity; the calculated side is a single +#: molecular transform sampled at the same Miller indices, where multiplicity +#: has no meaning. Assigned out of the class body only because `CalcShellE` is +#: defined below `WilsonShellEpsE`. +WilsonShellEpsE.calc_companion = CalcShellE + + class CalcGlobalE(EConvention): """Single global scale: ``E = F / rms(F)``, preserving inter-shell shape. diff --git a/torchref/experimental/alignment/ml_rotation.py b/torchref/experimental/alignment/ml_rotation.py index dfb908f1..c3618c72 100644 --- a/torchref/experimental/alignment/ml_rotation.py +++ b/torchref/experimental/alignment/ml_rotation.py @@ -25,6 +25,8 @@ import torch +from .e_values import (CalcShellE, WilsonShellEpsE, + convention_for_calc, convention_uses_sigma_f) from .frf.rotation_utils import ( axis_angle_to_matrix, edmonds_euler_from_rotation_matrix, @@ -71,25 +73,6 @@ def _normalize_to_e(F: torch.Tensor, shell_idx: torch.Tensor, return F / _per_shell_sqrt_mean(F, shell_idx, n_shells) -def _normalize_to_e_epsilon( - F: torch.Tensor, shell_idx: torch.Tensor, n_shells: int, eps: torch.Tensor, -) -> torch.Tensor: - """ε-corrected Wilson E: ``E²_h = (F²_h/ε_h) / ⟨F²/ε⟩_shell``. - - Matches Phaser's obs E (``E = F/sqrt(ε·Σ_N)``) and the FRF's - :func:`torchref.experimental.alignment.frf.preprocessing.wilson_normalise_epsilon`. The - plain :func:`_normalize_to_e` (no ε) over-counts axial reflections (ε>1) on - high-symmetry spacegroups, letting them dominate the ``-(E²+eImove)/V`` term - and blind the m_LETF1 orientation discrimination. - """ - I_corr = (F * F) / eps.clamp(min=1.0) - sum_shell = torch.zeros(n_shells, dtype=I_corr.dtype, device=F.device) - sum_shell.scatter_add_(0, shell_idx, I_corr) - count = torch.bincount(shell_idx, minlength=n_shells).to(I_corr.dtype) - mean_shell = (sum_shell / count.clamp(min=1.0)).clamp(min=1e-30) - return (I_corr / mean_shell.index_select(0, shell_idx)).clamp(min=0.0).sqrt() - - def _per_shell_sqrt_mean(F: torch.Tensor, shell_idx: torch.Tensor, n_shells: int) -> torch.Tensor: """Per-reflection ``sqrt(_shell)``. Wilson-normalisation denominator. @@ -590,7 +573,8 @@ def _build_llg_context( vrms_identity: float = 1.0, apply_wilson_b: bool = False, wilson_b_value: Optional[float] = None, - scat_mode: str = "legacy", + sig_F_obs: Optional[torch.Tensor] = None, + e_convention: type = WilsonShellEpsE, ) -> _LLGContext: """Build the rotation-independent m_LETF1 LLG context (DataMR.cc:1326-1429). @@ -637,7 +621,13 @@ def _build_llg_context( # both stop axial reflections (ε>1) from being over-weighted on # high-symmetry spacegroups. shell_idx = _equal_count_shell_idx(s_mag, n_shells) - E_obs = _normalize_to_e_epsilon(F_obs, shell_idx, n_shells, eps_factor) + conv_obs = e_convention + if sig_F_obs is None and convention_uses_sigma_f(conv_obs): + conv_obs = convention_for_calc(conv_obs) + E_obs = conv_obs( + F_obs, s_mag, centric, sig_F=sig_F_obs, eps=eps_factor, + shell_idx=shell_idx, n_shells=n_shells, + ).E # 3. Identity-rotation calc reference → E-normalisation scale for F_calc # (rotation-invariant: sphere permutation, shell sums preserved). @@ -645,28 +635,22 @@ def _build_llg_context( F_calc_ref = interpolator.evaluate( I_eye, hkl_real, real_cell, return_amplitude=True, ).to(dtype).squeeze(0) # (N,) - if scat_mode == "legacy": - # Per-shell unit-variance normalisation: forces _shell = 1 in - # EVERY shell, flattening F_calc's inter-shell amplitude shape. - calc_norm_per_h = _per_shell_sqrt_mean( - F_calc_ref, shell_idx, n_shells, - ).to(device) - elif scat_mode == "absolute": - # Single GLOBAL scale: preserves F_calc's inter-shell shape (how much the - # model actually scatters per resolution) instead of flattening it to 1. - # Phaser keeps E_calc physically scaled and carries the model's fraction - # of the cell in scatFactor = AtomScatRatio·SCATTERING/TOTAL_SCAT/NSYMP; - # for a search model that IS the full ASU (the benchmark case) scatFactor - # reduces to 1/n_ops, so the prefactor is unchanged and the only change - # here is dropping the per-shell flatten. - global_rms = F_calc_ref.pow(2).mean().clamp(min=1e-30).sqrt() - calc_norm_per_h = torch.full( - (N,), float(global_rms), dtype=dtype, device=device, - ) - else: - raise ValueError( - f"scat_mode={scat_mode!r}; expected 'legacy' or 'absolute'." - ) + # The calc normaliser is the convention's own choice, which is what the + # former `scat_mode` was: "legacy" is CalcShellE (forces _shell = 1 + # in every shell, flattening the model's inter-shell amplitude shape) and + # "absolute" is CalcGlobalE (one global scale, shape preserved). Two knobs + # for one decision meant obs and calc could be normalised by unrelated + # rules; now the same class answers for both sides. + # + # Phaser keeps E_calc physically scaled and carries the model's fraction of + # the cell in scatFactor = AtomScatRatio·SCATTERING/TOTAL_SCAT/NSYMP; for a + # search model that IS the full ASU (the benchmark case) scatFactor reduces + # to 1/n_ops, so the prefactor is unaffected either way. + conv_calc = convention_for_calc(e_convention)( + F_calc_ref, s_mag, centric, eps=eps_factor, + shell_idx=shell_idx, n_shells=n_shells, + ) + calc_norm_per_h = conv_calc.sigma.sqrt().to(dtype).to(device) # Optional Wilson-B match (EnsemblePDB.cc:793-851), applied as a per-reflection # Debye-Waller multiplier on F_calc. @@ -1096,7 +1080,8 @@ def m_letf1_rescore( vrms_identity: float = 1.0, apply_wilson_b: bool = False, wilson_b_value: Optional[float] = None, # if None and apply_wilson_b=True, fitted from data - scat_mode: str = "legacy", # "legacy" (per-shell calc norm) | "absolute" (global) + sig_F_obs: Optional[torch.Tensor] = None, + e_convention: type = WilsonShellEpsE, ) -> List[RotationPeak]: """Phaser-faithful ``m_LETF1`` rescore (DataMR.cc:1326-1429). @@ -1164,7 +1149,8 @@ def m_letf1_rescore( solvent_fsol=solvent_fsol, solvent_bsol=solvent_bsol, vrms_strategy=vrms_strategy, vrms_n_residues=vrms_n_residues, vrms_identity=vrms_identity, apply_wilson_b=apply_wilson_b, - wilson_b_value=wilson_b_value, scat_mode=scat_mode, + wilson_b_value=wilson_b_value, sig_F_obs=sig_F_obs, + e_convention=e_convention, ) alpha_t = torch.tensor([p.alpha for p in head], dtype=torch.float64) diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 9c84f492..ddba4ead 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -45,6 +45,7 @@ _external_rwork, _prepare_frf_inputs, ) +from .e_values import WilsonShellEpsE from .frf.rotation_utils import ( axis_angle_to_matrix, edmonds_euler_from_rotation_matrix, @@ -233,7 +234,7 @@ def __init__( model_error_A: Optional[float] = None, # --- rescore --- rescore_engine: str = "m_letf1", - rescore_scat_mode: str = "legacy", + rescore_e_convention: type = WilsonShellEpsE, auto_variance_weights: bool = True, use_interp_var: bool = False, subpeak_refine: bool = False, @@ -285,7 +286,7 @@ def __init__( self.model_error_A = float(model_error_A) self.rescore_engine = rescore_engine - self.rescore_scat_mode = rescore_scat_mode + self.rescore_e_convention = rescore_e_convention self.auto_variance_weights = auto_variance_weights self.use_interp_var = use_interp_var self.subpeak_refine = subpeak_refine @@ -518,6 +519,7 @@ def _rotation_candidates(self, frf) -> list: ) F_obs = frf.F_obs + sig_F = frf.sig_F hkl = frf.hkl s_mag = frf.s_mag centric = frf.centric @@ -561,11 +563,13 @@ def _rotation_candidates(self, frf) -> list: n_shells=rescore_n_shells, n_refine=min(len(peaks), self.n_ml_refine), batch_size=50, verbose=self.verbose, - scat_mode=self.rescore_scat_mode, + sig_F_obs=sig_F, + e_convention=self.rescore_e_convention, ) if self.subpeak_refine: rescored = self._subpeak_refine(rescored, F_obs, hkl, s_mag, - centric, ll, rescore_n_shells) + centric, ll, rescore_n_shells, + sig_F=sig_F) else: # legacy Sim/Rice approximation rescored = sim_mlrf_rescore( peaks, F_obs, hkl, s_mag, centric, ll, data.cell, @@ -579,7 +583,7 @@ def _rotation_candidates(self, frf) -> list: return rescored def _subpeak_refine(self, rescored, F_obs, hkl, s_mag, centric, ll, - rescore_n_shells): + rescore_n_shells, *, sig_F=None): """Quadratic tangent-space Newton sharpening of the top orientations.""" from .ml_rotation import _build_llg_context, quadratic_llg_refine @@ -590,7 +594,8 @@ def _subpeak_refine(self, rescored, F_obs, hkl, s_mag, centric, ll, F_obs, hkl, s_mag, centric, ll, data.cell, data.spacegroup, n_shells=rescore_n_shells, batch_size=50, - scat_mode=self.rescore_scat_mode, + sig_F_obs=sig_F, + e_convention=self.rescore_e_convention, ) k = self.subpeak_refine_k if self.subpeak_refine_k > 0 else self.n_rotation_candidates k = min(k, len(rescored)) From dbd8ec527fdaa85e134ea5c6ba01e07d597b1af9 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Fri, 28 Aug 2026 22:28:25 +0200 Subject: [PATCH 092/250] Share one wrapper rebuild between Model and ModelFT restore ModelFT.create_from_state_dict carried its own copy of the block that rebuilds the parameter wrappers from the atom table, and the node-field branch had only ever been added to Model's. A model in either field mode therefore failed to restore through ModelFT -- the class every refinement uses -- with a shape mismatch naming the refinable mask rather than the representation. The existing round-trip test drives the base Model, so it passed throughout. Both classes now call Model._rebuild_wrappers_from_pdb, and _restore_adp_slot covers the u slot as well as adp: mode="field_aniso" puts the field in u, which both copies rebuilt unconditionally as a CholeskyMixedTensor, so that mode had no restore path at all. The slot is identified by its saved neighbor_list rather than by the shape of its storage, since u is two-dimensional either way and (K, 10) cannot be told from (n_atoms, 6) by rank. NodeLoadTarget and NodeSmoothnessTarget read model.adp directly and so were inert whenever the field lived in u -- the mode with the most node parameters to collapse. Both now go through Model.adp_field, which already looked in both slots. The device conformance registry gains cases for those two targets and loses three entries for the payload strategies, which are not DeviceMixin subclasses and so were never found by its AST inventory. On 1DAW, field and field_aniso now round-trip through both classes with the B and U they were saved with; the isotropic control is unchanged. Co-Authored-By: Claude Opus 5 (1M context) --- docs/changelog.rst | 2 + tests/helpers/device_cases.py | 20 +- torchref/model/model.py | 272 ++++++++++-------- torchref/model/model_ft.py | 85 +----- torchref/refinement/targets/adp/node_load.py | 11 +- .../refinement/targets/adp/node_smoothness.py | 11 +- 6 files changed, 192 insertions(+), 209 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index b36d4ceb..9de4264a 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -5,6 +5,8 @@ Changelog Version 0.7.0 ---------- - Fixed cif reading bug discarding new mmCIF field for aniso ADPs +- Fixed ``ModelFT`` restore dropping a node-field ADP representation, and added the anisotropic ``field_aniso`` case; both models now share one wrapper-rebuild path +- Fixed the node load and node smoothness restraints being inert in ``field_aniso`` mode - Separated model configuration and provenance into ``ModelContext``. It now holds the unit cell, space group, atom table, link records, hydrogen settings, and input paths. - Refactored ``Symmetry`` as a crystallography-free class with transform primitives, and made ``SpaceGroup`` a specialised subclass. - Moved geometry predicates, HKL verbs, and grid-size helpers onto these classes as methods. diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index da36f0a4..5d65eaef 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -384,6 +384,23 @@ class TargetDeviceCase: ).ADPLocalityTarget(b["model"]), "ADPLocalityTarget", ), + # Registered unconditionally and inert off field mode, so the plain bundle + # model is enough to exercise their device handling. + TargetDeviceCase( + "NodeLoadTarget", + lambda b, d: __import__( + "torchref.refinement.targets.adp.node_load", fromlist=["NodeLoadTarget"] + ).NodeLoadTarget(b["model"]), + "NodeLoadTarget", + ), + TargetDeviceCase( + "NodeSmoothnessTarget", + lambda b, d: __import__( + "torchref.refinement.targets.adp.node_smoothness", + fromlist=["NodeSmoothnessTarget"], + ).NodeSmoothnessTarget(b["model"]), + "NodeSmoothnessTarget", + ), # Owns no tensors at all -- the case that exercises the request-driven # tracker path rather than the owned-tensor path. TargetDeviceCase( @@ -409,9 +426,6 @@ class TargetDeviceCase: "BaseWeighting": "abstract base; covered via ManualWeighting", "Refinement": "abstract base; covered via LBFGSRefinement in integration", "PassThroughTensor": "documented non-functional stub (parameter_wrappers.py)", - "NodePayload": "stateless strategy, holds no tensors; abstract base", - "IsotropicPayload": "stateless strategy, holds no tensors", - "AnisotropicPayload": "stateless strategy, holds only a float epsilon", "ADPTarget": "abstract base; needs a model with ADPs", "CombinedTargets": "composite container; needs its component targets", "CombinedModelTargets": "composite container; needs a loaded model", diff --git a/torchref/model/model.py b/torchref/model/model.py index 4e6d128c..1a6fd757 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -2083,6 +2083,158 @@ def load_state(self, path: str, strict: bool = True): if self.ctx.verbose > 0: print(f"Loaded model state from {path}") + @staticmethod + def _restore_adp_slot(prefix, state_dict, pdb, saved_dtype, xyz_wrapper): + """Rebuild the ``adp`` or ``u`` wrapper, as a node field when the state was one. + + Built from the PDB for its shapes and masks only; ``load_state_dict`` overwrites + every value afterwards. + + A saved :class:`~torchref.model.disorder_field.DisorderFieldTensor` is recognised + by its ``neighbor_list``, not by the shape of its storage: the ``u`` slot holds a + 2-D tensor either way, so shape alone cannot tell a ``(K, 10)`` node field from a + ``(n_atoms, 6)`` per-atom U. + + Parameters + ---------- + prefix : {"adp", "u"} + Which slot to rebuild. ``"u"`` carries the anisotropic representation. + state_dict : dict + The state being restored, read but not consumed. + pdb : pandas.DataFrame + Atom table supplying the initial values. + saved_dtype : torch.dtype + Float dtype the state was saved in. + xyz_wrapper : MixedTensor + The already-rebuilt coordinate wrapper; a node field derives its node + positions from it. + """ + from torchref.model.parameter_wrappers import ( + CholeskyMixedTensor, + PositiveMixedTensor, + ) + + aniso = prefix == "u" + name = "aniso_U" if aniso else "adp" + mask = state_dict.get(f"{prefix}.refinable_mask") + if aniso: + initial = torch.tensor( + pdb[["u11", "u22", "u33", "u12", "u13", "u23"]].values, + dtype=saved_dtype, + ) + else: + initial = torch.tensor(pdb["tempfactor"].values, dtype=saved_dtype) + + saved_nl = state_dict.get(f"{prefix}.neighbor_list") + if saved_nl is None: + # Match load(): the anisotropic U is a CholeskyMixedTensor so a restored + # model refines it in the same positive-definite-by-construction + # parametrization as a freshly-loaded one. + wrapper = CholeskyMixedTensor if aniso else PositiveMixedTensor + return wrapper(initial, refinable_mask=mask, name=name) + + from torchref.model.disorder_field import ( + AnisotropicPayload, + DisorderFieldTensor, + IsotropicPayload, + ) + + payload = AnisotropicPayload() if aniso else IsotropicPayload() + saved_values = state_dict[f"{prefix}.fixed_values"] + # Rebuild with the SAVED anchor rows: cluster anchoring makes these length + # n_atoms where single-atom anchoring makes them length K, so reconstructing + # them from scratch would shape-mismatch on load. + saved_anchor_atom = state_dict.get(f"{prefix}.anchor_atom") + saved_anchor_node = state_dict.get(f"{prefix}.anchor_node") + return DisorderFieldTensor( + initial_values=initial, + xyz_fn=xyz_wrapper, + n_nodes=int(saved_values.shape[0]), + k_neighbors=int(saved_nl.shape[1]), + payload=payload, + # Storage is [payload | log sigma | offset], so the extra three columns + # say whether node positions carry a refinable offset. + refine_positions=bool(saved_values.shape[1] == payload.width + 4), + anchor_rows=( + (saved_anchor_atom, saved_anchor_node) + if saved_anchor_atom is not None + else None + ), + refinable_mask=mask, + mask_in_node_space=True, + name=name, + dtype=saved_dtype, + ) + + @classmethod + def _rebuild_wrappers_from_pdb(cls, instance, pdb, state_dict, saved_dtype, device): + """Give ``instance`` parameter wrappers and per-atom buffers of the right shape. + + The half of :meth:`create_from_state_dict` that every subclass needs + identically, so subclasses call this rather than restating it: a per-class copy + drifts, and a restore that rebuilds the wrong wrapper type fails on a shape + mismatch rather than on anything that names the real cause. + + Values are placeholders throughout --- the caller's ``load_state_dict`` is what + puts the saved numbers in. Only shapes, masks and dtypes matter here. + """ + from torchref.model.parameter_wrappers import MixedTensor, OccupancyTensor + + n_atoms = len(pdb) + + instance.xyz = MixedTensor( + torch.tensor(pdb[["x", "y", "z"]].values, dtype=saved_dtype), + refinable_mask=state_dict.get("xyz.refinable_mask"), + name="xyz", + ) + instance.adp = cls._restore_adp_slot( + "adp", state_dict, pdb, saved_dtype, instance.xyz + ) + instance.u = cls._restore_adp_slot( + "u", state_dict, pdb, saved_dtype, instance.xyz + ) + + initial_occ = torch.tensor(pdb["occupancy"].values, dtype=saved_dtype) + sharing_groups, altloc_groups, refinable_mask = ( + instance._create_occupancy_groups(pdb, initial_occ) + ) + # A saved mask is in group space; expand it back over atoms. + saved_occ_mask = state_dict.get("occupancy.refinable_mask") + if saved_occ_mask is not None: + if saved_occ_mask.device != sharing_groups.device: + saved_occ_mask = saved_occ_mask.to(sharing_groups.device) + refinable_mask = saved_occ_mask[sharing_groups] + + instance.occupancy = OccupancyTensor( + initial_values=initial_occ, + sharing_groups=sharing_groups, + altloc_groups=altloc_groups, + refinable_mask=refinable_mask, + dtype=saved_dtype, + device=device, + name="occupancy", + ) + + if "aniso_flag" not in instance._buffers or instance.aniso_flag is None: + instance.register_buffer( + "aniso_flag", + torch.tensor(pdb["anisou_flag"].values, dtype=torch.bool), + ) + # Pre-compute SF indices (respects exclude_H_from_sf) + instance._rebuild_sf_indices() + + for mask_name in ("xyz_mask", "adp_mask", "u_mask", "occupancy_mask"): + instance.register_buffer( + mask_name, torch.ones(n_atoms, dtype=torch.bool, device=device) + ) + + # Note: inv_fractional_matrix, fractional_matrix and recB are properties + # delegating to Cell, so they are not registered as buffers. + if state_dict.get("vdw_radii") is not None: + instance.register_buffer( + "vdw_radii", torch.zeros_like(state_dict["vdw_radii"], device=device) + ) + @classmethod def create_from_state_dict( cls, @@ -2150,125 +2302,7 @@ def create_from_state_dict( # The wrappers are built from the PDB purely to get the right shapes and # masks; load_state_dict below overwrites their values. if pdb is not None: - n_atoms = len(pdb) - - xyz_mask = state_dict.get("xyz.refinable_mask") - adp_mask = state_dict.get("adp.refinable_mask") - u_mask = state_dict.get("u.refinable_mask") - - instance.xyz = MixedTensor( - torch.tensor(pdb[["x", "y", "z"]].values, dtype=saved_dtype), - refinable_mask=xyz_mask, - name="xyz", - ) - # A saved node field has 2-D ``adp`` storage (K, 2) where a per-atom - # wrapper has 1-D, so the shape says which representation to rebuild. - # Built only for its shapes and masks; load_state_dict overwrites values. - saved_adp = state_dict.get("adp.fixed_values") - if saved_adp is not None and saved_adp.ndim == 2: - from torchref.model.disorder_field import DisorderFieldTensor - - saved_nl = state_dict.get("adp.neighbor_list") - # Rebuild with the SAVED anchor rows: cluster anchoring makes these - # length n_atoms where single-atom anchoring makes them length K, so - # reconstructing them from scratch would shape-mismatch on load. - saved_anchor_atom = state_dict.get("adp.anchor_atom") - saved_anchor_node = state_dict.get("adp.anchor_node") - anchor_rows = ( - (saved_anchor_atom, saved_anchor_node) - if saved_anchor_atom is not None - else None - ) - instance.adp = DisorderFieldTensor( - initial_values=torch.tensor( - pdb["tempfactor"].values, dtype=saved_dtype - ), - xyz_fn=instance.xyz, - n_nodes=int(saved_adp.shape[0]), - k_neighbors=( - int(saved_nl.shape[1]) if saved_nl is not None else 12 - ), - # Storage width says whether positions carry a refinable offset. - refine_positions=bool(saved_adp.shape[1] == 5), - anchor_rows=anchor_rows, - refinable_mask=adp_mask, - mask_in_node_space=True, - name="adp", - dtype=saved_dtype, - ) - else: - instance.adp = PositiveMixedTensor( - torch.tensor(pdb["tempfactor"].values, dtype=saved_dtype), - refinable_mask=adp_mask, - name="adp", - ) - # Match load(): the anisotropic U is a CholeskyMixedTensor so the - # restored model refines it in the same positive-definite-by- - # construction parametrization as a freshly-loaded one. - instance.u = CholeskyMixedTensor( - torch.tensor( - pdb[["u11", "u22", "u33", "u12", "u13", "u23"]].values, - dtype=saved_dtype, - ), - refinable_mask=u_mask, - name="aniso_U", - ) - - # Create OccupancyTensor - initial_occ = torch.tensor(pdb["occupancy"].values, dtype=saved_dtype) - sharing_groups, altloc_groups, refinable_mask = ( - instance._create_occupancy_groups(pdb, initial_occ) - ) - - # Override mask if present in state_dict - saved_occ_mask = state_dict.get("occupancy.refinable_mask") - if saved_occ_mask is not None: - if saved_occ_mask.device != sharing_groups.device: - saved_occ_mask = saved_occ_mask.to(sharing_groups.device) - refinable_mask = saved_occ_mask[sharing_groups] - - instance.occupancy = OccupancyTensor( - initial_values=initial_occ, - sharing_groups=sharing_groups, - altloc_groups=altloc_groups, - refinable_mask=refinable_mask, - dtype=saved_dtype, - device=device, - name="occupancy", - ) - - # Register buffers that are needed - if "aniso_flag" not in instance._buffers or instance.aniso_flag is None: - instance.register_buffer( - "aniso_flag", - torch.tensor(pdb["anisou_flag"].values, dtype=torch.bool), - ) - # Pre-compute SF indices (respects exclude_H_from_sf) - instance._rebuild_sf_indices() - - # Register mask buffers - instance.register_buffer( - "xyz_mask", torch.ones(n_atoms, dtype=torch.bool, device=device) - ) - instance.register_buffer( - "adp_mask", torch.ones(n_atoms, dtype=torch.bool, device=device) - ) - instance.register_buffer( - "u_mask", torch.ones(n_atoms, dtype=torch.bool, device=device) - ) - instance.register_buffer( - "occupancy_mask", torch.ones(n_atoms, dtype=torch.bool, device=device) - ) - - # Register other buffers based on state_dict - # Note: inv_fractional_matrix, fractional_matrix, recB are now properties - # delegating to Cell, so they're not registered as buffers - buffer_names = ["vdw_radii"] - for name in buffer_names: - if name in state_dict and state_dict[name] is not None: - instance.register_buffer( - name, torch.zeros_like(state_dict[name], device=device) - ) + cls._rebuild_wrappers_from_pdb(instance, pdb, state_dict, saved_dtype, device) # Drop only empty-in-dim-0 tensors (placeholders from an atom-less state); # scalars and non-tensor entries must survive for load_state_dict. diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index fad61b2c..07862bfd 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -959,8 +959,6 @@ def create_from_state_dict( The anisotropic ``u`` is rebuilt as a :class:`CholeskyMixedTensor`, as in :meth:`load`, so the positive-definite parametrization round-trips. """ - from torchref.symmetry import SpaceGroup - # Resolve dtype/device at call time so the fallback below uses the # current config rather than an import-time default. device = normalize_device(device) @@ -1008,88 +1006,13 @@ def create_from_state_dict( if cell_tensor is not None: instance.cell = Cell(cell_tensor, dtype=saved_dtype, device=device) - # If PDB exists, create the parameter wrappers with correct shapes + # Wrappers and per-atom buffers: shared with Model so the two restores cannot + # drift apart again. ModelFT adds only its own scattering buffers below. if pdb is not None: - from torchref.model.parameter_wrappers import ( - CholeskyMixedTensor, - MixedTensor, - OccupancyTensor, - PositiveMixedTensor, - ) - - n_atoms = len(pdb) - - xyz_mask = state_dict.get("xyz.refinable_mask") - adp_mask = state_dict.get("adp.refinable_mask") - u_mask = state_dict.get("u.refinable_mask") - - instance.xyz = MixedTensor( - torch.tensor(pdb[["x", "y", "z"]].values, dtype=saved_dtype), - refinable_mask=xyz_mask, - name="xyz", - ) - instance.adp = PositiveMixedTensor( - torch.tensor(pdb["tempfactor"].values, dtype=saved_dtype), - refinable_mask=adp_mask, - name="adp", - ) - instance.u = CholeskyMixedTensor( - torch.tensor( - pdb[["u11", "u22", "u33", "u12", "u13", "u23"]].values, - dtype=saved_dtype, - ), - refinable_mask=u_mask, - name="aniso_U", - ) - - initial_occ = torch.tensor(pdb["occupancy"].values, dtype=saved_dtype) - sharing_groups, altloc_groups, refinable_mask = ( - instance._create_occupancy_groups(pdb, initial_occ) - ) - - saved_occ_mask = state_dict.get("occupancy.refinable_mask") - if saved_occ_mask is not None: - if saved_occ_mask.device != sharing_groups.device: - saved_occ_mask = saved_occ_mask.to(sharing_groups.device) - refinable_mask = saved_occ_mask[sharing_groups] - - instance.occupancy = OccupancyTensor( - initial_values=initial_occ, - sharing_groups=sharing_groups, - altloc_groups=altloc_groups, - refinable_mask=refinable_mask, - dtype=saved_dtype, - device=device, - name="occupancy", + cls._rebuild_wrappers_from_pdb( + instance, pdb, state_dict, saved_dtype, device ) - if "aniso_flag" not in instance._buffers or instance.aniso_flag is None: - instance.register_buffer( - "aniso_flag", - torch.tensor(pdb["anisou_flag"].values, dtype=torch.bool), - ) - - # Register mask buffers - instance.register_buffer( - "xyz_mask", torch.ones(n_atoms, dtype=torch.bool, device=device) - ) - instance.register_buffer( - "adp_mask", torch.ones(n_atoms, dtype=torch.bool, device=device) - ) - instance.register_buffer( - "u_mask", torch.ones(n_atoms, dtype=torch.bool, device=device) - ) - instance.register_buffer( - "occupancy_mask", torch.ones(n_atoms, dtype=torch.bool, device=device) - ) - - # Register vdw_radii if present - if "vdw_radii" in state_dict and state_dict["vdw_radii"] is not None: - instance.register_buffer( - "vdw_radii", - torch.zeros_like(state_dict["vdw_radii"], device=device), - ) - # Scattering buffers: accept both old-style (A, B) and new (_A, _B). a_key = "_A" if "_A" in state_dict else "A" if "A" in state_dict else None b_key = "_B" if "_B" in state_dict else "B" if "B" in state_dict else None diff --git a/torchref/refinement/targets/adp/node_load.py b/torchref/refinement/targets/adp/node_load.py index 3eac95fb..7d62be8e 100644 --- a/torchref/refinement/targets/adp/node_load.py +++ b/torchref/refinement/targets/adp/node_load.py @@ -72,9 +72,14 @@ def __init__( @property def _field(self): - """The disorder field, or ``None`` when the model is not in field mode.""" - adp = getattr(self.model, "adp", None) - return adp if hasattr(adp, "node_load") else None + """The disorder field, or ``None`` when the model is not in field mode. + + Reads ``Model.adp_field`` rather than the ``adp`` slot directly: an anisotropic + payload lives in ``u`` instead, and looking only at ``adp`` would leave this + target silently inert in exactly the mode with the most node parameters to + collapse. + """ + return getattr(self.model, "adp_field", None) def _relative_load(self) -> torch.Tensor: """Each node's load as a multiple of the mean load, ``(K,)``.""" diff --git a/torchref/refinement/targets/adp/node_smoothness.py b/torchref/refinement/targets/adp/node_smoothness.py index 586dab5d..5b22a8d5 100644 --- a/torchref/refinement/targets/adp/node_smoothness.py +++ b/torchref/refinement/targets/adp/node_smoothness.py @@ -77,9 +77,14 @@ def __init__( @property def _field(self): - """The disorder field, or ``None`` when the model is not in field mode.""" - adp = getattr(self.model, "adp", None) - return adp if hasattr(adp, "node_load") else None + """The disorder field, or ``None`` when the model is not in field mode. + + Reads ``Model.adp_field`` rather than the ``adp`` slot directly: an anisotropic + payload lives in ``u`` instead, and looking only at ``adp`` would leave this + target silently inert in exactly the mode with the most node parameters to + collapse. + """ + return getattr(self.model, "adp_field", None) def _pair_terms(self): """``(weighted mean squared log-B difference, pair weights)``.""" From a596ed9e317dfb7545d5dccb3d316cf4ebb9a1ca Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Fri, 28 Aug 2026 22:28:40 +0200 Subject: [PATCH 093/250] Drop the stored real-space coordinate grid SfFFT held an (nx, ny, nz, 3) buffer of Cartesian voxel coordinates that no structure-factor path reads. Every splat derives a voxel's position arithmetically from its index and the fractionalisation matrix, so build_electron_density used the tensor only for .device and .shape[:-1], and the CUDA wrappers already passed density_map in its place through a grid_ptr that none of the four Triton kernels ever loads. Being a registered buffer, it also went into every saved state: 187.5 MB for 3K7M at 250^3, against 24 bytes for the gridsize and voxel_size that remain. Building it cost 65.2 ms there and 5.6 ms on 1DAW, which was all of setup_grid's runtime; setup_grid is now 0.05 ms. build_electron_density takes a grid shape and a device instead. Gone with the buffer: its dead voxel_size parameter, grid_ptr in the four Triton kernels and their launch helpers, the stand-in argument in the CUDA wrappers, SfFFT.compute_real_space_grid, and ModelFT.get_radius with its two forwards, which the per-atom truncation radius had made vestigial. voxel_size is now frac_matrix @ (1 / gridsize) rather than a difference of two grid points. Shape-only callers use the new grid_shape property, and ModelFT.real_space_grid is a method over get_real_grid for the callers that want the coordinates themselves. F_calc and the density map are bit-identical before and after on 1DAW (C2), 3GR5 (P6522) and 3K7M (P432), and the analytic voxel_size reproduces the differenced value exactly. Checkpoints carrying the old buffer still restore. Co-Authored-By: Claude Opus 5 (1M context) --- docs/changelog.rst | 1 + tests/unit/base/test_canonical_sphere_cpu.py | 56 ++++++-------- tests/unit/structure_factor/helpers.py | 10 +-- tests/unit/structure_factor/test_dispatch.py | 8 +- .../kernels/cuda/variable_radius.py | 70 ++++++++--------- torchref/base/electron_density/main.py | 23 +++--- .../monolithic_refinement/density_solvent.py | 13 ++-- torchref/experimental/targets/realspace.py | 2 +- torchref/model/mixed_model.py | 28 ++----- torchref/model/model_collection.py | 10 +-- torchref/model/model_ft.py | 75 +++++++------------ torchref/model/sf_fft.py | 73 +++++++----------- torchref/scaling/solvent.py | 5 +- 13 files changed, 147 insertions(+), 227 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 9de4264a..fffa01d3 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -5,6 +5,7 @@ Changelog Version 0.7.0 ---------- - Fixed cif reading bug discarding new mmCIF field for aniso ADPs +- Removed the stored real-space coordinate grid; ``build_electron_density`` takes a grid shape and device, and ``ModelFT.real_space_grid()`` builds one on demand - Fixed ``ModelFT`` restore dropping a node-field ADP representation, and added the anisotropic ``field_aniso`` case; both models now share one wrapper-rebuild path - Fixed the node load and node smoothness restraints being inert in ``field_aniso`` mode - Separated model configuration and provenance into ``ModelContext``. It now holds the unit cell, space group, atom table, link records, hydrogen settings, and input paths. diff --git a/tests/unit/base/test_canonical_sphere_cpu.py b/tests/unit/base/test_canonical_sphere_cpu.py index 943baac1..1522bb41 100644 --- a/tests/unit/base/test_canonical_sphere_cpu.py +++ b/tests/unit/base/test_canonical_sphere_cpu.py @@ -74,12 +74,6 @@ def _cell(beta_deg, dtype=torch.float32, dims=(48, 40, 34), abc=(28.0, 24.0, 20. return f64.to(dtype), torch.linalg.inv(f64).to(dtype), dims, f64 -def _voxel_size(f64, dims): - """``voxel_size`` as ``sf_fft`` derives it; unused by the kernels, still in the - ``build_electron_density`` signature.""" - return (f64.norm(dim=0) / torch.tensor(dims, dtype=torch.float64)).float() - - def _iso_atoms(f64, n=36, dtype=torch.float32, seed=0): g = torch.Generator().manual_seed(seed) z = torch.tensor([6, 7, 8, 16]).repeat(n // 4 + 1)[:n] @@ -242,8 +236,7 @@ def test_fused_float64_is_exact(): # AUTO vs EAGER through the real dispatch: no accelerator needed # =========================================================================== -def _build(pin, dims, frac, inv_frac, voxel, dtype, iso=None, aniso=None): - rsg = torch.zeros(*dims, 3, dtype=dtype) # shape only; no kernel reads its values +def _build(pin, dims, frac, inv_frac, dtype, iso=None, aniso=None): xi, ai, oi, Ai, Bi = iso if iso is not None else _empty_iso(dtype) kw = {} if aniso is not None: @@ -251,7 +244,8 @@ def _build(pin, dims, frac, inv_frac, voxel, dtype, iso=None, aniso=None): kw = dict(xyz_aniso=xa, u_aniso=ua, occ_aniso=oa, A_aniso=Aa, B_aniso=Ba) with (use_portable() if pin else contextlib.nullcontext()): return build_electron_density( - rsg, xi, ai, oi, Ai, Bi, inv_frac, frac, voxel, dtype=dtype, **kw) + dims, torch.device("cpu"), xi, ai, oi, Ai, Bi, inv_frac, frac, + dtype=dtype, **kw) @pytest.mark.parametrize("beta", _BETAS) @@ -259,10 +253,9 @@ def _build(pin, dims, frac, inv_frac, voxel, dtype, iso=None, aniso=None): ids=["float32", "float64"]) def test_auto_matches_eager_iso(beta, dtype): frac, inv_frac, dims, f64 = _cell(beta, dtype=dtype) - voxel = _voxel_size(f64, dims) atoms = _iso_atoms(f64, dtype=dtype) - ref = _build(True, dims, frac, inv_frac, voxel, dtype, iso=atoms) - got = _build(False, dims, frac, inv_frac, voxel, dtype, iso=atoms) + ref = _build(True, dims, frac, inv_frac, dtype, iso=atoms) + got = _build(False, dims, frac, inv_frac, dtype, iso=atoms) tol = _F32_TOL if dtype is torch.float32 else 1e-12 assert _rel_l2(got, ref) < tol @@ -270,10 +263,9 @@ def test_auto_matches_eager_iso(beta, dtype): @pytest.mark.parametrize("beta", _BETAS) def test_auto_matches_eager_aniso(beta): frac, inv_frac, dims, f64 = _cell(beta) - voxel = _voxel_size(f64, dims) atoms = _aniso_atoms(f64) - ref = _build(True, dims, frac, inv_frac, voxel, torch.float32, aniso=atoms) - got = _build(False, dims, frac, inv_frac, voxel, torch.float32, aniso=atoms) + ref = _build(True, dims, frac, inv_frac, torch.float32, aniso=atoms) + got = _build(False, dims, frac, inv_frac, torch.float32, aniso=atoms) assert _rel_l2(got, ref) < _F32_TOL @@ -282,7 +274,6 @@ def test_auto_matches_eager_gradients(kind): """Direction *and* magnitude: a kernel returning ``2 * grad`` is perfectly parallel, so cosine alone cannot catch it.""" frac, inv_frac, dims, f64 = _cell(100.0, dtype=torch.float64, dims=(32, 28, 24)) - voxel = _voxel_size(f64, dims) w = torch.randn(dims, generator=torch.Generator().manual_seed(7), dtype=torch.float64) if kind == "iso": @@ -293,7 +284,7 @@ def test_auto_matches_eager_gradients(kind): def grads(pin): x, pp, o = (t.clone().requires_grad_() for t in (xyz, p, occ)) pack = (x, pp, o, A, B) - dm = _build(pin, dims, frac, inv_frac, voxel, torch.float64, + dm = _build(pin, dims, frac, inv_frac, torch.float64, **({"iso": pack} if kind == "iso" else {"aniso": pack})) (dm * w).sum().backward() return x.grad, pp.grad, o.grad @@ -350,7 +341,7 @@ def recording(*args, **kwargs): return real(*args, **kwargs) monkeypatch.setattr(sphere_splat, "add_isotropic_cpu_sphere_var", recording) - _build(False, dims, frac, inv_frac, _voxel_size(f64, dims), dtype, + _build(False, dims, frac, inv_frac, dtype, iso=_iso_atoms(f64, dtype=dtype)) assert calls, f"the default did not reach the fused CPU splat for {dtype}" @@ -358,22 +349,20 @@ def recording(*args, **kwargs): def test_empty_atom_sets(): """A structure with no isotropic (or no anisotropic) atoms must not crash.""" frac, inv_frac, dims, f64 = _cell(90.0) - voxel = _voxel_size(f64, dims) - only_aniso = _build(False, dims, frac, inv_frac, voxel, torch.float32, + only_aniso = _build(False, dims, frac, inv_frac, torch.float32, aniso=_aniso_atoms(f64)) assert torch.isfinite(only_aniso).all() and float(only_aniso.abs().sum()) > 0 - both_empty = _build(False, dims, frac, inv_frac, voxel, torch.float32) + both_empty = _build(False, dims, frac, inv_frac, torch.float32) assert float(both_empty.abs().sum()) == 0.0 def test_density_map_accumulates_not_overwrites(): """Both passes add into one map, so the aniso pass must not clobber the iso one.""" frac, inv_frac, dims, f64 = _cell(90.0) - voxel = _voxel_size(f64, dims) iso, aniso = _iso_atoms(f64), _aniso_atoms(f64) - a = _build(False, dims, frac, inv_frac, voxel, torch.float32, iso=iso) - b = _build(False, dims, frac, inv_frac, voxel, torch.float32, aniso=aniso) - both = _build(False, dims, frac, inv_frac, voxel, torch.float32, + a = _build(False, dims, frac, inv_frac, torch.float32, iso=iso) + b = _build(False, dims, frac, inv_frac, torch.float32, aniso=aniso) + both = _build(False, dims, frac, inv_frac, torch.float32, iso=iso, aniso=aniso) assert _rel_l2(both, a + b) < 1e-6 @@ -407,20 +396,19 @@ class of index and factor errors that cross-backend parity cannot see because bo """ frac, inv_frac, dims, f64 = _cell(beta, dtype=torch.float64) xyz, adp, occ, A, B = _iso_atoms(f64, n=24, dtype=torch.float64) - voxel = _voxel_size(f64, dims) - grid = torch.zeros(*dims, 3, dtype=torch.float64) u_sph = torch.zeros(xyz.shape[0], 6, dtype=torch.float64) u_sph[:, :3] = (adp / (8.0 * math.pi**2)).unsqueeze(1) with (use_portable() if pin else contextlib.nullcontext()): iso_map = build_electron_density( - grid, xyz, adp, occ, A, B, inv_frac, frac, voxel, dtype=torch.float64 + dims, torch.device("cpu"), xyz, adp, occ, A, B, inv_frac, frac, + dtype=torch.float64 ) aniso_map = build_electron_density( - grid, + dims, torch.device("cpu"), xyz[:0], adp[:0], occ[:0], A[:0], B[:0], - inv_frac, frac, voxel, + inv_frac, frac, xyz_aniso=xyz, u_aniso=u_sph, occ_aniso=occ, A_aniso=A, B_aniso=B, dtype=torch.float64, ) @@ -465,20 +453,20 @@ def test_fused_kernel_is_thread_invariant(n_threads): frac, inv_frac, dims, f64 = _cell(115.0, dtype=torch.float32) xyz, adp, occ, A, B = _iso_atoms(f64, n=96, dtype=torch.float32, seed=7) - voxel = _voxel_size(f64, dims) - grid = torch.zeros(*dims, 3, dtype=torch.float32) original = torch.get_num_threads() try: torch.set_num_threads(1) with contextlib.nullcontext(): ref = build_electron_density( - grid, xyz, adp, occ, A, B, inv_frac, frac, voxel, dtype=torch.float32 + dims, torch.device("cpu"), xyz, adp, occ, A, B, inv_frac, frac, + dtype=torch.float32 ) torch.set_num_threads(n_threads) with contextlib.nullcontext(): got = build_electron_density( - grid, xyz, adp, occ, A, B, inv_frac, frac, voxel, dtype=torch.float32 + dims, torch.device("cpu"), xyz, adp, occ, A, B, inv_frac, frac, + dtype=torch.float32 ) finally: torch.set_num_threads(original) diff --git a/tests/unit/structure_factor/helpers.py b/tests/unit/structure_factor/helpers.py index 54bfca15..aae056b3 100644 --- a/tests/unit/structure_factor/helpers.py +++ b/tests/unit/structure_factor/helpers.py @@ -465,11 +465,11 @@ def sf_fft_for( # falls back to the portable splat (``main.py`` catches and falls through), so a # dispatch-driven test can pass while measuring a different kernel than the one it # names. Calling the kernel directly settles that by construction. -# 2. **No global-config coupling.** ``SfFFT`` builds its grid through ``get_real_grid``, -# which reads the *global* ``dtypes.float`` and takes no dtype argument -- so an MPS -# ``SfFFT`` under this package's float64 pin would try to allocate float64 on MPS and -# fail. ``ifft`` and ``extract_structure_factor_from_grid`` read no global config at -# all, so :func:`density_to_F` needs no config switching. +# 2. **No global-config coupling.** ``build_electron_density`` allocates its map at the +# *global* ``dtypes.float`` when no ``dtype`` is passed, so a dispatch-driven MPS test +# under this package's float64 pin would try to allocate float64 on MPS and fail. +# ``ifft`` and ``extract_structure_factor_from_grid`` read no global config at all, so +# :func:`density_to_F` needs no config switching. # # The dispatch ladder is a separate concern, tested in ``test_dispatch.py``. diff --git a/tests/unit/structure_factor/test_dispatch.py b/tests/unit/structure_factor/test_dispatch.py index 4368dc49..1dcbada2 100644 --- a/tests/unit/structure_factor/test_dispatch.py +++ b/tests/unit/structure_factor/test_dispatch.py @@ -45,18 +45,14 @@ def _build(scene, device, dtype, force_portable=False, aniso=False): """ s = scene.to(device=device, dtype=dtype) dims = H._grid_dims(s) - grid = torch.zeros(*dims, 3, dtype=dtype, device=device) - voxel = torch.tensor( - [float(s.cell.data[i]) / dims[i] for i in range(3)], dtype=dtype, device=device - ) empty1 = s.xyz.new_zeros(0) empty3 = s.xyz.new_zeros(0, 3) empty5 = s.A.new_zeros(0, 5) kw = dict( - real_space_grid=grid, + grid_shape=dims, + device=device, inv_frac_matrix=s.inv_frac_matrix, frac_matrix=s.frac_matrix, - voxel_size=voxel, dtype=dtype, ) kw["force_portable"] = force_portable diff --git a/torchref/base/electron_density/kernels/cuda/variable_radius.py b/torchref/base/electron_density/kernels/cuda/variable_radius.py index a787afa0..2784916b 100644 --- a/torchref/base/electron_density/kernels/cuda/variable_radius.py +++ b/torchref/base/electron_density/kernels/cuda/variable_radius.py @@ -58,7 +58,7 @@ def _sym3_inv(a, b, c, d, e, f): @triton.jit def _wq_grid_fwd_kernel( n_items, - grid_ptr, density_map_ptr, + density_map_ptr, xyz_ptr, b_ptr, A_ptr, B_ptr, occ_ptr, r2cut_ptr, mask_ptr, inv_frac_ptr, frac_ptr, @@ -177,7 +177,7 @@ def _wq_grid_fwd_kernel( @triton.jit def _wq_grid_bwd_kernel( n_items, - grid_ptr, grad_density_map_ptr, + grad_density_map_ptr, xyz_ptr, b_ptr, A_ptr, B_ptr, occ_ptr, r2cut_ptr, mask_ptr, inv_frac_ptr, frac_ptr, @@ -327,7 +327,7 @@ def _wq_grid_bwd_kernel( @triton.jit def _wq_grid_aniso_fwd_kernel( n_items, - grid_ptr, density_map_ptr, + density_map_ptr, xyz_ptr, u_ptr, A_ptr, B_ptr, occ_ptr, r2cut_ptr, mask_ptr, inv_frac_ptr, frac_ptr, @@ -461,7 +461,7 @@ def _wq_grid_aniso_fwd_kernel( @triton.jit def _wq_grid_aniso_bwd_kernel( n_items, - grid_ptr, grad_density_map_ptr, + grad_density_map_ptr, xyz_ptr, u_ptr, A_ptr, B_ptr, occ_ptr, r2cut_ptr, mask_ptr, inv_frac_ptr, frac_ptr, @@ -666,12 +666,12 @@ def _wq_grid_aniso_bwd_kernel( def _launch_grid_fwd(out_flat, r2cut, mask, scene_buffers, dims): """Isotropic grid=(n_atoms,) forward (fixed FWD_BLOCK_V/FWD_NUM_WARPS).""" - (grid_flat, xyz, b, A, B, occ, inv_frac, frac) = scene_buffers + (xyz, b, A, B, occ, inv_frac, frac) = scene_buffers nx, ny, nz = dims n_atoms = r2cut.shape[0] _wq_grid_fwd_kernel[(n_atoms,)]( n_atoms, - grid_flat, out_flat, + out_flat, xyz, b, A, B, occ, r2cut, mask, inv_frac, frac, @@ -682,12 +682,12 @@ def _launch_grid_fwd(out_flat, r2cut, mask, scene_buffers, dims): def _launch_grid_aniso_fwd(out_flat, r2cut, mask, scene_buffers, dims): """Anisotropic grid=(n_atoms,) forward (fixed FWD_BLOCK_V/FWD_NUM_WARPS).""" - (grid_flat, xyz, u, A, B, occ, inv_frac, frac) = scene_buffers + (xyz, u, A, B, occ, inv_frac, frac) = scene_buffers nx, ny, nz = dims n_atoms = r2cut.shape[0] _wq_grid_aniso_fwd_kernel[(n_atoms,)]( n_atoms, - grid_flat, out_flat, + out_flat, xyz, u, A, B, occ, r2cut, mask, inv_frac, frac, @@ -705,14 +705,13 @@ class WorkQueueGridDensity(torch.autograd.Function): """ @staticmethod - def forward(ctx, density_map, real_space_grid, xyz, b, occ, A, B, + def forward(ctx, density_map, xyz, b, occ, A, B, r2cut, mask, inv_frac, frac): # Accumulate the splat into a copy of the running density_map (out = # density_map + splat) so the dispatch needs no separate zeros buffer + add. # A clone (not in-place) keeps this autograd-trivial AND safe for the AUTO # fallthrough: density_map is untouched if the kernel raises. - nx, ny, nz = real_space_grid.shape[:3] - grid_flat = real_space_grid.contiguous().view(-1) + nx, ny, nz = density_map.shape[:3] xyz = xyz.contiguous(); b = b.contiguous(); occ = occ.contiguous() A = A.contiguous(); B = B.contiguous() inv_frac_flat = inv_frac.contiguous().view(-1) @@ -720,19 +719,19 @@ def forward(ctx, density_map, real_space_grid, xyz, b, occ, A, B, out = density_map.contiguous().clone().view(-1) _launch_grid_fwd( out, r2cut, mask, - (grid_flat, xyz, b, A, B, occ, inv_frac_flat, frac_flat), + (xyz, b, A, B, occ, inv_frac_flat, frac_flat), (nx, ny, nz), ) - ctx.save_for_backward(real_space_grid, xyz, b, occ, A, B, + ctx.dims = (nx, ny, nz) + ctx.save_for_backward(xyz, b, occ, A, B, r2cut, mask, inv_frac, frac) return out.view(nx, ny, nz) @staticmethod def backward(ctx, grad_density_map): - (real_space_grid, xyz, b, occ, A, B, + (xyz, b, occ, A, B, r2cut, mask, inv_frac, frac) = ctx.saved_tensors - nx, ny, nz = real_space_grid.shape[:3] - grid_flat = real_space_grid.contiguous().view(-1) + nx, ny, nz = ctx.dims grad_dm = grad_density_map.contiguous().view(-1) inv_frac_flat = inv_frac.contiguous().view(-1) frac_flat = frac.contiguous().view(-1) @@ -741,7 +740,7 @@ def backward(ctx, grad_density_map): grad_occ = torch.zeros_like(occ) _wq_grid_bwd_kernel[(r2cut.shape[0],)]( r2cut.shape[0], - grid_flat, grad_dm, + grad_dm, xyz.contiguous(), b.contiguous(), A.contiguous(), B.contiguous(), occ.contiguous(), r2cut, mask, inv_frac_flat, frac_flat, @@ -750,8 +749,8 @@ def backward(ctx, grad_density_map): num_warps=BWD_NUM_WARPS, ) # out = density_map + splat -> grad wrt density_map is identity. - # grads for: density_map, real_space_grid, xyz, b, occ, A, B, r2cut, mask, inv_frac, frac - return (grad_density_map, None, grad_xyz, grad_b, grad_occ, None, None, + # grads for: density_map, xyz, b, occ, A, B, r2cut, mask, inv_frac, frac + return (grad_density_map, grad_xyz, grad_b, grad_occ, None, None, None, None, None, None) @@ -762,11 +761,10 @@ class WorkQueueGridDensityAniso(torch.autograd.Function): ``grad_u`` (6 components) in place of ``grad_b``.""" @staticmethod - def forward(ctx, density_map, real_space_grid, xyz, u, occ, A, B, + def forward(ctx, density_map, xyz, u, occ, A, B, r2cut, mask, inv_frac, frac): # Accumulate into a copy of the running density_map (see the iso forward). - nx, ny, nz = real_space_grid.shape[:3] - grid_flat = real_space_grid.contiguous().view(-1) + nx, ny, nz = density_map.shape[:3] xyz = xyz.contiguous(); u = u.contiguous(); occ = occ.contiguous() A = A.contiguous(); B = B.contiguous() inv_frac_flat = inv_frac.contiguous().view(-1) @@ -774,19 +772,19 @@ def forward(ctx, density_map, real_space_grid, xyz, u, occ, A, B, out = density_map.contiguous().clone().view(-1) _launch_grid_aniso_fwd( out, r2cut, mask, - (grid_flat, xyz, u, A, B, occ, inv_frac_flat, frac_flat), + (xyz, u, A, B, occ, inv_frac_flat, frac_flat), (nx, ny, nz), ) - ctx.save_for_backward(real_space_grid, xyz, u, occ, A, B, + ctx.dims = (nx, ny, nz) + ctx.save_for_backward(xyz, u, occ, A, B, r2cut, mask, inv_frac, frac) return out.view(nx, ny, nz) @staticmethod def backward(ctx, grad_density_map): - (real_space_grid, xyz, u, occ, A, B, + (xyz, u, occ, A, B, r2cut, mask, inv_frac, frac) = ctx.saved_tensors - nx, ny, nz = real_space_grid.shape[:3] - grid_flat = real_space_grid.contiguous().view(-1) + nx, ny, nz = ctx.dims grad_dm = grad_density_map.contiguous().view(-1) inv_frac_flat = inv_frac.contiguous().view(-1) frac_flat = frac.contiguous().view(-1) @@ -795,7 +793,7 @@ def backward(ctx, grad_density_map): grad_occ = torch.zeros_like(occ) _wq_grid_aniso_bwd_kernel[(r2cut.shape[0],)]( r2cut.shape[0], - grid_flat, grad_dm, + grad_dm, xyz.contiguous(), u.contiguous(), A.contiguous(), B.contiguous(), occ.contiguous(), r2cut, mask, inv_frac_flat, frac_flat, @@ -804,7 +802,7 @@ def backward(ctx, grad_density_map): num_warps=BWD_NUM_WARPS, ) # out = density_map + splat -> grad wrt density_map is identity. - return (grad_density_map, None, grad_xyz, grad_u, grad_occ, None, None, + return (grad_density_map, grad_xyz, grad_u, grad_occ, None, None, None, None, None, None) @@ -821,15 +819,9 @@ def backward(ctx, grad_density_map): # wrappers do for CUDA what ``add_*_mps_var`` already did for Metal: square the # radius and build the coefficient mask, rather than leaving that to the caller. # -# ``density_map`` is passed where the ``autograd.Function`` wants -# ``real_space_grid``. That is exact, not a convenience: ``forward``/``backward`` -# use that argument only for ``.shape[:3]`` and for a ``grid_flat`` pointer handed -# to the kernel as ``grid_ptr`` -- which none of the four Triton kernels ever -# loads, because voxel coordinates are derived arithmetically from ``frac`` and the -# grid dims. ``density_map`` has the same ``(nx, ny, nz)`` shape, so both uses are -# satisfied. If a kernel is ever changed to actually read ``grid_ptr``, this breaks -# and the right fix is to delete that dead parameter, not to thread a real grid -# back through here. +# No coordinate grid is threaded anywhere: every voxel's Cartesian position is +# derived arithmetically in-kernel from ``frac`` and the grid dims, so ``density_map`` +# alone carries the ``(nx, ny, nz)`` shape the launch needs. def why_unavailable(): @@ -869,7 +861,6 @@ def add_isotropic_cuda_var( """ return WorkQueueGridDensity.apply( density_map, - density_map, # stands in for real_space_grid; see the note above xyz, adp, occ, @@ -893,7 +884,6 @@ def add_anisotropic_cuda_var( """ return WorkQueueGridDensityAniso.apply( density_map, - density_map, # stands in for real_space_grid; see the note above xyz, u, occ, diff --git a/torchref/base/electron_density/main.py b/torchref/base/electron_density/main.py index e619ffe9..7a80492d 100644 --- a/torchref/base/electron_density/main.py +++ b/torchref/base/electron_density/main.py @@ -37,7 +37,8 @@ def build_electron_density( - real_space_grid: torch.Tensor, + grid_shape, + device: torch.device, xyz_iso: torch.Tensor, adp_iso: torch.Tensor, occ_iso: torch.Tensor, @@ -45,7 +46,6 @@ def build_electron_density( B_iso: torch.Tensor, inv_frac_matrix: torch.Tensor, frac_matrix: torch.Tensor, - voxel_size: torch.Tensor, xyz_aniso: Optional[torch.Tensor] = None, u_aniso: Optional[torch.Tensor] = None, occ_aniso: Optional[torch.Tensor] = None, @@ -61,18 +61,18 @@ def build_electron_density( Parameters ---------- - real_space_grid : torch.Tensor - Coordinate grid, shape ``(nx, ny, nz, 3)``. + grid_shape : tuple of int + Map dimensions ``(nx, ny, nz)``. No coordinate grid is needed or built: every + splat derives a voxel's Cartesian position arithmetically from its index and + ``inv_frac_matrix``. + device : torch.device + Device to allocate the map on. xyz_iso, adp_iso, occ_iso : torch.Tensor Isotropic positions ``(n_iso, 3)``, B-factors and occupancies ``(n_iso,)``. A_iso, B_iso : torch.Tensor ITC92 coefficients, shape ``(n_iso, 5)``. inv_frac_matrix, frac_matrix : torch.Tensor Cartesian-to-fractional and fractional-to-Cartesian, shape ``(3, 3)``. - voxel_size : torch.Tensor - **Unused by every splat** -- the truncation radius comes from each atom's B/U and - ``torchref.sigma_cutoff_ed``, and the enumeration box from ``inv_frac_matrix``. - Retained because ``SfFFT`` passes it positionally. xyz_aniso, u_aniso, occ_aniso : torch.Tensor, optional Anisotropic positions ``(n_aniso, 3)``, U ``(n_aniso, 6)``, occupancies. A_aniso, B_aniso : torch.Tensor, optional @@ -86,12 +86,7 @@ def build_electron_density( """ if dtype is None: dtype = get_float_dtype() - device = real_space_grid.device - density_map = torch.zeros( - real_space_grid.shape[:-1], - dtype=dtype, - device=device, - ) + density_map = torch.zeros(tuple(grid_shape), dtype=dtype, device=device) # --- isotropic atoms --- if len(xyz_iso) > 0: diff --git a/torchref/experimental/monolithic_refinement/density_solvent.py b/torchref/experimental/monolithic_refinement/density_solvent.py index bac26532..e4f4668c 100644 --- a/torchref/experimental/monolithic_refinement/density_solvent.py +++ b/torchref/experimental/monolithic_refinement/density_solvent.py @@ -280,10 +280,13 @@ def _smooth(self, field): At ``sigma=0`` the kernel is the identity (no smoothing). Differentiable w.r.t. ``sigma_shell`` and -- through ``field`` -- w.r.t. atomic xyz/B. """ - grid = self.solvent_fft.real_space_grid # (nx, ny, nz, 3) Cartesian - dx = float((grid[1, 0, 0] - grid[0, 0, 0]).norm()) - dy = float((grid[0, 1, 0] - grid[0, 0, 0]).norm()) - dz = float((grid[0, 0, 1] - grid[0, 0, 0]).norm()) + # Axis spacings: column j of the fractional (frac->cart) matrix is cell edge + # vector j, so its norm over the sampling count along that axis is the step + # between grid points one index apart -- what differencing the coordinate grid + # used to measure, without building the grid. + frac = self.solvent_fft.cell.fractional_matrix + n = self.solvent_fft.grid_shape + dx, dy, dz = (float(frac[:, j].norm()) / n[j] for j in range(3)) sigma = self.sigma_shell.clamp(min=0.0, max=5.0) nx, ny, nz = field.shape two_pi2 = 2.0 * torch.pi**2 @@ -316,7 +319,7 @@ def _get_tau(self, z): def _nyquist_mask(self, hkl): """True where |h|,|k|,|l| are within the coarse grid Nyquist limit.""" - nc = self.solvent_fft.real_space_grid.shape[:-1] # (ncx, ncy, ncz) + nc = self.solvent_fft.grid_shape # (ncx, ncy, ncz) nyq = torch.tensor( [n // 2 for n in nc], device=hkl.device, dtype=hkl.dtype ) diff --git a/torchref/experimental/targets/realspace.py b/torchref/experimental/targets/realspace.py index 56fbf8b7..fe95a0e5 100644 --- a/torchref/experimental/targets/realspace.py +++ b/torchref/experimental/targets/realspace.py @@ -122,7 +122,7 @@ def _ensure_grid(self): """Ensure model's SfFFT grid is set up.""" if self._model is None: raise RuntimeError("No model set for RealSpaceTarget") - if self._model.real_space_grid is None: + if self._model.gridsize is None: self._model.setup_grid() def _get_data_p1(self) -> "ReflectionData": diff --git a/torchref/model/mixed_model.py b/torchref/model/mixed_model.py index 9c01ebd0..91ced770 100644 --- a/torchref/model/mixed_model.py +++ b/torchref/model/mixed_model.py @@ -178,10 +178,14 @@ def dtype_float(self): # Grid infrastructure (delegates to constituent models) # ========================================================================= + def real_space_grid(self) -> torch.Tensor: + """Build the Cartesian grid from the first model (shared cell → same grid).""" + return self.models[0].real_space_grid() + @property - def real_space_grid(self) -> Optional[torch.Tensor]: - """Real-space coordinate grid from first model (shared cell → same grid).""" - return self.models[0].real_space_grid + def grid_shape(self) -> Optional[tuple]: + """Map dimensions (nx, ny, nz) from the first model.""" + return self.models[0].grid_shape @property def fft(self): @@ -221,24 +225,6 @@ def setup_grid(self, max_res=None, gridsize=None): for model in self.models: model.setup_grid(max_res=max_res, gridsize=gridsize) - def get_radius(self, min_radius_Angstrom: float = 4.0) -> int: - """ - Get the radius in voxels for density calculation. - - Delegates to first model (same grid → same voxel size). - - Parameters - ---------- - min_radius_Angstrom : float, optional - Minimum radius in Angstroms. Default is 4.0. - - Returns - ------- - int - Radius in voxels. - """ - return self.models[0].get_radius(min_radius_Angstrom) - def build_complete_map(self) -> torch.Tensor: """ Mixed electron density ``density_mixed = Σ w_i density_i``. diff --git a/torchref/model/model_collection.py b/torchref/model/model_collection.py index dc082742..b269209b 100644 --- a/torchref/model/model_collection.py +++ b/torchref/model/model_collection.py @@ -120,9 +120,12 @@ def device(self): def dtype_float(self): return self._base_models[0].dtype_float - @property def real_space_grid(self): - return self._base_models[0].real_space_grid + return self._base_models[0].real_space_grid() + + @property + def grid_shape(self): + return self._base_models[0].grid_shape @property def fft(self): @@ -148,9 +151,6 @@ def setup_grid(self, max_res=None, gridsize=None): for model in self._base_models: model.setup_grid(max_res=max_res, gridsize=gridsize) - def get_radius(self, min_radius_Angstrom: float = 4.0) -> int: - return self._base_models[0].get_radius(min_radius_Angstrom) - def build_complete_map(self) -> torch.Tensor: """Mixed electron density: sum_i w_i * density_i.""" fractions = self.fractions diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index 07862bfd..31abe453 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -52,9 +52,10 @@ class ModelFT(CachedForwardMixin, Model): ---------- max_res, wavelength, anomalous_threshold : float The constructor arguments above, readable back as attributes. - gridsize, real_space_grid : torch.Tensor - Grid dimensions ``(nx, ny, nz)`` and coordinate grid - ``(nx, ny, nz, 3)``; both live on the ``SfFFT`` submodule. + gridsize : torch.Tensor + Grid dimensions ``(nx, ny, nz)``, living on the ``SfFFT`` submodule. + A coordinate grid is not stored; :meth:`real_space_grid` builds one on + demand for the few callers that want the Cartesian positions themselves. map : torch.Tensor or None Most recently computed electron density map. parametrization : dict @@ -293,15 +294,28 @@ def gridsize(self, value): """Set grid size (for backward compatibility).""" self._fft.gridsize = value - @property - def real_space_grid(self) -> Optional[torch.Tensor]: - """Real-space coordinate grid with shape (nx, ny, nz, 3).""" - return self._fft.real_space_grid + def real_space_grid(self) -> torch.Tensor: + """Build the Cartesian coordinate of every grid point, ``(nx, ny, nz, 3)``. + + Not stored: at ``12 * nx * ny * nz`` bytes it is the largest tensor a model + would hold, and no structure-factor path reads it -- every splat derives a + voxel's position from its index. Built here for the callers that genuinely + want the coordinates, and discarded when they are done with it. + """ + from torchref.base.fourier import get_real_grid + + if self.gridsize is None: + self.setup_grid() + return get_real_grid( + fractional_matrix=self.cell.fractional_matrix, + gridsize=self.gridsize, + device=self.device, + ) - @real_space_grid.setter - def real_space_grid(self, value): - """Set real space grid (for backward compatibility).""" - self._fft.real_space_grid = value + @property + def grid_shape(self) -> Optional[tuple]: + """Map dimensions ``(nx, ny, nz)``, or ``None`` before the grid is set up.""" + return self._fft.grid_shape @property def voxel_size(self) -> Optional[torch.Tensor]: @@ -395,40 +409,9 @@ def setup_grid(self, max_res=None, gridsize=None): ) if self.ctx.verbose > 2: - print(f"Grid shape: {self._fft.real_space_grid.shape[:-1]}") + print(f"Grid shape: {self._fft.grid_shape}") print(f"Voxel size: {self._fft.voxel_size}") - def get_radius(self, min_radius_Angstrom: float = 4.0): - """ - Get a single fixed splat radius in voxels for the given minimum. - - Vestigial: the density path truncates each atom at its own - ``torchref.sigma_cutoff_ed * sigma_eff`` radius and never consults this. - - Parameters - ---------- - min_radius_Angstrom : float, optional - Minimum radius in Angstroms. Default is 4.0. - - Returns - ------- - int - Radius in voxels. - """ - if not hasattr(self, "real_space_grid") or self.real_space_grid is None: - self.setup_grid() - voxel_size = self.real_space_grid[1, 1, 1] - self.real_space_grid[0, 0, 0] - min_radius = ( - torch.ceil(min_radius_Angstrom / torch.min(voxel_size)) - .to(dtypes.int) - .item() - ) - if self.ctx.verbose > 1: - print( - f"Calculated radius for density calculation: {min_radius} voxels (voxel size: {voxel_size}), this corresponds to at least {min_radius_Angstrom} Å" - ) - return min_radius - def build_complete_map(self, radius=None, apply_symmetry=True): """ Build electron density map from all atoms. @@ -475,7 +458,7 @@ def build_initial_map(self, apply_symmetry=True): torch.Tensor Electron density map with shape (nx, ny, nz). """ - if self._fft.real_space_grid is None: + if self._fft.gridsize is None: self.setup_grid() if self.ctx.verbose > 2: @@ -819,7 +802,7 @@ def copy(self, detach: bool = True) -> "ModelFT": Create a deep copy of the ModelFT. Creates a complete independent copy including all Model base class data, - FFT submodule state (gridsize, real_space_grid, voxel_size), + FFT submodule state (gridsize, voxel_size), ITC92 parametrization, and scalar attributes. Cache is reset to empty. @@ -876,7 +859,7 @@ def copy(self, detach: bool = True) -> "ModelFT": if self._fft is not None: model_copy._fft = self._fft.copy() - if self._fft.real_space_grid is not None: + if self._fft.gridsize is not None: model_copy.setup_grid(max_res=self.max_res) # Don't share cached structure factors with the original. diff --git a/torchref/model/sf_fft.py b/torchref/model/sf_fft.py index 33bfd4e0..c0901e3a 100644 --- a/torchref/model/sf_fft.py +++ b/torchref/model/sf_fft.py @@ -10,7 +10,7 @@ import torch import torch.nn as nn -from torchref.base.fourier import get_real_grid, ifft +from torchref.base.fourier import ifft from torchref.base.reciprocal import extract_structure_factor_from_grid from torchref.config import dtypes, get_default_device from torchref.symmetry import Cell, SpaceGroup @@ -49,9 +49,12 @@ class SfFFT(DeviceMovementMixin, nn.Module): cell, spacegroup : Cell, SpaceGroup The unit cell and the space group as an nn.Module carrying its symmetry matrices and translations; ``symmetry`` is an alias for ``spacegroup``. - gridsize, real_space_grid, voxel_size : torch.Tensor or None - Grid dimensions ``(nx, ny, nz)``, coordinate grid ``(nx, ny, nz, 3)`` and - voxel dimensions -- all ``None`` until :meth:`setup_grid` runs. + gridsize, voxel_size : torch.Tensor or None + Grid dimensions ``(nx, ny, nz)`` and the voxel edge vector sum -- both + ``None`` until :meth:`setup_grid` runs. No coordinate grid is stored: the + splats derive a voxel's Cartesian position from its index, so materialising + one would cost ``12 * nx * ny * nz`` bytes that nothing reads. Call + :func:`torchref.base.fourier.get_real_grid` if you genuinely need one. """ def __init__( @@ -114,7 +117,6 @@ def __init__( # Buffers (registered during setup_grid) self.register_buffer("gridsize", None) - self.register_buffer("real_space_grid", None) self.register_buffer("voxel_size", None) # Late symmetry compatibility flag (set during setup_grid) @@ -201,6 +203,13 @@ def set_cell_and_spacegroup(self, cell: Cell, spacegroup: SpaceGroupLike = None) # Grid Setup Methods # ========================================================================= + @property + def grid_shape(self) -> Optional[Tuple[int, int, int]]: + """Map dimensions ``(nx, ny, nz)``, or ``None`` before :meth:`setup_grid`.""" + if self.gridsize is None: + return None + return tuple(int(n) for n in self.gridsize) + def compute_optimal_gridsize(self, max_res: Optional[float] = None) -> tuple: """ Compute optimal grid dimensions using the stored cell and spacegroup. @@ -245,37 +254,6 @@ def compute_optimal_gridsize(self, max_res: Optional[float] = None) -> tuple: ) return gridsize_optimized - @staticmethod - def compute_real_space_grid( - fractional_matrix: torch.Tensor, - gridsize: torch.Tensor, - device: torch.device = None, - ) -> torch.Tensor: - """ - Generate the real-space coordinate grid. - - Parameters - ---------- - fractional_matrix : torch.Tensor - Fractionalization matrix mapping Cartesian to fractional - coordinates, with shape (3, 3). - gridsize : torch.Tensor - Grid dimensions (nx, ny, nz). - device : torch.device, optional - Target device. Defaults to the configured default device. - - Returns - ------- - torch.Tensor - Real-space grid with shape (nx, ny, nz, 3). - """ - # Forward ``device`` as-is, including ``None``: ``get_real_grid`` infers - # from ``fractional_matrix`` when no device is given, and resolving the - # global default here would preempt that. - return get_real_grid( - fractional_matrix=fractional_matrix, gridsize=gridsize, device=device - ) - def setup_grid( self, gridsize: Optional[Tuple[int, int, int]] = None, @@ -318,20 +296,21 @@ def setup_grid( optimal_gridsize, dtype=dtypes.int, device=self.device ) - # Compute real space grid - self.real_space_grid = self.compute_real_space_grid( - self._cell.fractional_matrix, self.gridsize, self.device + # The step between diagonally adjacent grid points, i.e. the sum of the three + # cell edge vectors each divided by its own sampling count. Equal to the true + # per-axis voxel edge lengths only for an orthogonal cell; kept because that is + # what the previous grid-differencing definition produced. + self.voxel_size = ( + self._cell.fractional_matrix.to(self.device) + @ (1.0 / self.gridsize.to(self._cell.fractional_matrix.dtype)) ) - # Compute voxel size - self.voxel_size = self.real_space_grid[2, 2, 2] - self.real_space_grid[1, 1, 1] - # Every symmetry-equivalent HKL lands on an integer grid point exactly when # the grid admits direct indexing, which the space group answers without # building an operator. if self._spacegroup is not None: self._late_symmetry_compatible = self._spacegroup.can_index_directly( - self.real_space_grid.shape[:-1] + self.grid_shape ) if self.use_late_symmetry and self._late_symmetry_compatible: @@ -353,7 +332,7 @@ def setup_grid( self._spacegroup.reset_cache() if self.verbose > 2: - print(f"Grid shape: {self.real_space_grid.shape[:-1]}") + print(f"Grid shape: {self.grid_shape}") print(f"Voxel size: {self.voxel_size}") # ========================================================================= @@ -399,13 +378,14 @@ def build_density_map( torch.Tensor Electron density map with shape (nx, ny, nz). """ - if self.real_space_grid is None: + if self.gridsize is None: self.setup_grid() from torchref.base.electron_density.main import build_electron_density density_map = build_electron_density( - real_space_grid=self.real_space_grid, + grid_shape=self.grid_shape, + device=self.device, xyz_iso=xyz_iso, adp_iso=adp_iso, occ_iso=occ_iso, @@ -413,7 +393,6 @@ def build_density_map( B_iso=B_iso, inv_frac_matrix=self.inv_fractional_matrix, frac_matrix=self.fractional_matrix, - voxel_size=self.voxel_size, xyz_aniso=xyz_aniso, u_aniso=u_aniso, occ_aniso=occ_aniso, diff --git a/torchref/scaling/solvent.py b/torchref/scaling/solvent.py index a226545e..f9b76945 100644 --- a/torchref/scaling/solvent.py +++ b/torchref/scaling/solvent.py @@ -220,7 +220,7 @@ def __init__( self.model = ModuleReference(model) # Store reference to model self.model.get_vdw_radii() # Ensure VdW radii are available assert self.model, "Model is not initialized" - if model.real_space_grid == None: + if model.gridsize is None: model.setup_grid() # Phenix-style parameters @@ -353,14 +353,13 @@ def get_solvent_mask(self): xyz = self.model.xyz() # (N_atoms, 3) vdw_radii = self.model.get_vdw_radii() # (N_atoms,) - self.real_space_grid = self.model.real_space_grid inv_frac = self.model.inv_fractional_matrix frac = self.model.fractional_matrix with torch.no_grad(): spacegroup = self.model.fft.spacegroup n_ops = spacegroup.n_ops - grid_shape = self.real_space_grid.shape[:-1] + grid_shape = self.model.grid_shape device = self.model.device n_atoms = xyz.shape[0] From b33e26f6bd6c3a39d8aa36c51db853bde65d41ae Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sat, 29 Aug 2026 01:46:22 +0200 Subject: [PATCH 094/250] Add a bitwise pre/post check for the dev merge `a596ed9e` removes SfFFT's stored real-space coordinate grid, and reading it the change is plumbing: `build_electron_density` used the tensor only for `.device` and `.shape[:-1]`, the four Triton kernels never dereferenced the `grid_ptr` they were handed, and the new `voxel_size = frac_matrix @ (1/gridsize)` is algebraically the old `grid[2,2,2] - grid[1,1,1]`. "Reading it, it looks equivalent" is not equivalent, and one part was genuinely not guaranteed: a single matmul replacing a difference of two is the same value in exact arithmetic, not necessarily in floating point. So it is measured. Same script, same data, run in a worktree at the commit before the merge and at the merge: F_calc, the density map, voxel_size and the 500-peak FRF list all hash identically on 1DAW, 2DQ6, 3K7M and 4BX9 -- voxel_size included. Single-threaded, since the peak list only reproduces bit-for-bit at one thread. 1DAW is C2 and 2DQ6 is P3121 on purpose. The new spacing formulas assume column j of the fractionalisation matrix is cell edge vector j, which `cart = B @ f` confirms; on an orthogonal cell that matrix is diagonal, so a transposed convention would have been invisible. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/merge_identity.sh | 31 +++++++ .../analysis/merge_numeric_identity.py | 86 +++++++++++++++++++ 2 files changed, 117 insertions(+) create mode 100644 alignment_lab/analysis/merge_identity.sh create mode 100644 alignment_lab/analysis/merge_numeric_identity.py diff --git a/alignment_lab/analysis/merge_identity.sh b/alignment_lab/analysis/merge_identity.sh new file mode 100644 index 00000000..a0600b55 --- /dev/null +++ b/alignment_lab/analysis/merge_identity.sh @@ -0,0 +1,31 @@ +#!/bin/bash +# Same script, same data, two trees: post-merge and the commit before it. +# Single-threaded because the FRF peak list only reproduces bit-for-bit at one +# thread -- ~5e-8 of score noise reorders peaks otherwise, and this gate is +# about bit-identity. +#SBATCH --job-name=mergeid +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=2 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-3 +set -uo pipefail +POST=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PRE=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/premerge_check +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +# 1DAW is C2 (monoclinic) and 2DQ6 is P3121: on an orthogonal cell the +# fractionalisation matrix is diagonal, so a transposed edge-vector convention +# in the new voxel_size would not show. On these it would. +PDBS=(1DAW 2DQ6 3K7M 4BX9) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +export TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 +for TREE in "$PRE" "$POST"; do + TAG=$([ "$TREE" = "$PRE" ] && echo pre || echo post) + cd "$TREE" + PYTHONPATH="$TREE" "$PY" -u alignment_lab/analysis/merge_numeric_identity.py \ + --pdb "$PDB" --tag "$TAG" 2>/dev/null | grep '^OUT ' +done diff --git a/alignment_lab/analysis/merge_numeric_identity.py b/alignment_lab/analysis/merge_numeric_identity.py new file mode 100644 index 00000000..8c08d1d5 --- /dev/null +++ b/alignment_lab/analysis/merge_numeric_identity.py @@ -0,0 +1,86 @@ +"""Did merging dev change any number the alignment path produces? + +`a596ed9e` removes SfFFT's stored real-space coordinate grid. Reading it, the +change is plumbing: `build_electron_density` used the tensor only for `.device` +and `.shape[:-1]`, the four Triton kernels never dereferenced the `grid_ptr` +they were handed, and the new `voxel_size = frac_matrix @ (1/gridsize)` is +algebraically the old `grid[2,2,2] - grid[1,1,1]` -- `cart = B @ f`, so column j +of the fractionalisation matrix is cell edge vector j. + +"Reading it, it looks equivalent" is not the same as equivalent. This dumps +hashes of the quantities the alignment stack actually consumes so the two trees +can be compared bit for bit. + +Deliberately covers a monoclinic and two high-symmetry cells: for an orthogonal +cell the fractionalisation matrix is diagonal, so a transposed edge-vector +convention would be invisible. 1DAW is C2 and 2DQ6 is P3121, where it would not. +""" + +from __future__ import annotations + +import argparse +import hashlib +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import BENCH_PDBS, FRFConfig, load_case, run_frf # noqa: E402 + + +def _h(t: torch.Tensor) -> str: + """Bit-exact hash of a tensor's contents, dtype and shape.""" + t = t.detach().cpu().contiguous() + m = hashlib.sha256() + m.update(str(tuple(t.shape)).encode()) + m.update(str(t.dtype).encode()) + m.update(t.numpy().tobytes()) + return m.hexdigest()[:16] + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) + ap.add_argument("--tag", default="?") + ap.add_argument("--lmax-cap", type=int, default=64) + args = ap.parse_args() + + model, data = load_case(args.pdb) + hkl = data.hkl + + # 1. Structure factors -- the thing every downstream number is built on. + F = model.get_structure_factor(hkl, recalc=True) + print(f"OUT {args.tag} {args.pdb} F_calc {_h(F)}") + + # 2. The density map itself, one step earlier than F_calc, so a difference + # can be localised to the splat rather than the FFT. + model.setup_grid() + dm = model.build_complete_map() + print(f"OUT {args.tag} {args.pdb} density {_h(dm)}") + print(f"OUT {args.tag} {args.pdb} gridshape {tuple(dm.shape)}") + + # 3. voxel_size: the one quantity whose FORMULA changed, rather than only + # its call site. Nothing reads it downstream, so a difference here is + # reportable but not itself a regression. + vs = model.voxel_size + print(f"OUT {args.tag} {args.pdb} voxel_size " + f"{'None' if vs is None else _h(vs)} " + f"{'' if vs is None else [f'{float(v):.17g}' for v in vs.flatten()]}") + + # 4. The FRF peak list -- what this branch is actually judged on. + res = run_frf(model, data, FRFConfig(n_peaks=500, lmax_cap=args.lmax_cap), + capture_arf=False, verbose=0) + ang = torch.tensor([[p.alpha, p.beta, p.gamma] for p in res.peaks], + dtype=torch.float64) + sc = torch.tensor([p.score for p in res.peaks], dtype=torch.float64) + print(f"OUT {args.tag} {args.pdb} peaks_angles {_h(ang)}") + print(f"OUT {args.tag} {args.pdb} peaks_scores {_h(sc)}") + print(f"OUT {args.tag} {args.pdb} n_peaks {len(res.peaks)}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) From c9dd060ac9af4e8045abbbf23b3453b331219855 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sat, 29 Aug 2026 04:11:09 +0200 Subject: [PATCH 095/250] Finish the E-value consolidation, and fix the gate that hid a break The alignment package had nine private ways into E-space. The rotation function and the m_LETF1 rescore were converted earlier; this takes the last five -- `sim_mlrf_rescore`, `_build_sim_llg_context`, and translation's three sites -- and removes what they leave behind: `_normalize_to_e`, `_per_shell_sqrt_mean` and `wilson_normalise_epsilon`, none of which had a caller afterwards. Every migrated site passes its OWN `shell_idx`. `_equal_count_shell_idx` and translation's inline loop are rank-based; the shared `assign_shells` is value-based, and they disagree on reflections sitting on a boundary. The sigma_A fits downstream are tied to whichever binning produced them, so the binning is preserved rather than unified. Measured rather than assumed. The sim sites and translation's `llg_translation_rescore` are bit-identical -- that one already WAS the convention's arithmetic, written out. The two Patterson sites are not: they reduced with a Python loop over shells and a tensor `.mean()`, against scatter_add here, which moves the last ulp -- 1.5e-14 absolute at worst, ~1e-16 relative. They also divided by zero on an empty shell where the convention clamps. `wilson_normalise` had no production caller left but is the reference the seam was validated against, so it moves to `alignment_lab/lab/reference_normalisers.py` rather than being deleted -- a future convention will want to compare against what actually shipped. Its unit test now asserts the defining formula inline instead, since a unit test should not reach into the lab for its expectation. Two lab-side repairs. `FRF_STAGES` still patched `frf.api.wilson_normalise` and `frf.api.french_wilson_preprocess`, and neither attribute exists any more: the first is gone, the second is reached through `FrenchWilsonE._compute`, which imports it inside the method body. The file's own docstring warns that a row registering zero calls "reads as that stage is free" -- so one row is repointed at the defining module and the other removed. And the gate itself was wrong. `pytest tests/unit` picks up pyproject.toml rather than tests/pytest.ini, so the slow marker silently skipped the rotation-search and translation tests -- which is how SEVEN of them sat broken by the dev merge (`model.initialized` moved to `model.ctx.initialized`) through a gate that reported 1968 passed and fully green. Running from `tests/` fixes the marker but breaks the io tests, which open `tests/files/...` relative to the root. `-c tests/pytest.ini` from the root satisfies both. Gate: 1978 passed, 0 failed, with --run-slow actually in effect. Seam identity 0.000e+00. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/full_gate.sh | 14 ++-- alignment_lab/analysis/fw_footing.py | 3 +- alignment_lab/analysis/seam_identity.py | 3 +- alignment_lab/diagnostics/frf_prep_compare.py | 2 +- alignment_lab/lab/profile.py | 11 ++- alignment_lab/lab/reference_normalisers.py | 53 ++++++++++++++ docs/changelog.rst | 3 + tests/unit/alignment/test_e_conventions.py | 31 +++++---- .../test_interp_var_and_shared_sigma_a.py | 1 - torchref/experimental/alignment/align.py | 2 +- .../alignment/frf/preprocessing.py | 69 ------------------- .../experimental/alignment/ml_rotation.py | 36 +++------- torchref/experimental/alignment/pipeline.py | 4 +- .../experimental/alignment/rotation_search.py | 2 +- .../experimental/alignment/translation.py | 36 +++++----- 15 files changed, 133 insertions(+), 137 deletions(-) create mode 100644 alignment_lab/lab/reference_normalisers.py diff --git a/alignment_lab/analysis/full_gate.sh b/alignment_lab/analysis/full_gate.sh index 4f647e43..2764ad9f 100644 --- a/alignment_lab/analysis/full_gate.sh +++ b/alignment_lab/analysis/full_gate.sh @@ -2,18 +2,24 @@ #SBATCH --job-name=fullgate #SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out #SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 +#SBATCH --partition=day +#SBATCH --time=03:00:00 #SBATCH --cpus-per-task=8 #SBATCH --mem=48G #SBATCH --constraint=cpu_epyc9335 set -uo pipefail REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" export TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 -"$PY" -m pytest tests/unit -q 2>&1 | tail -14 +# `-c tests/pytest.ini` from the repo ROOT, with --run-slow. Both halves matter +# and they pull in opposite directions: a bare `pytest tests/unit` picks up +# pyproject.toml, where the slow marker silently skips the rotation-search and +# translation tests -- which is how seven of them stayed broken by a merge +# through a gate reporting everything green. But running from `tests/` to get +# the right config then breaks the io tests, which open `tests/files/...` +# relative to the root. Naming the config explicitly satisfies both. +"$PY" -m pytest -c tests/pytest.ini tests/unit --run-slow -q 2>&1 | tail -16 echo "PYTEST_RC=${PIPESTATUS[0]}" echo "== seam identity ==" "$PY" -u alignment_lab/analysis/seam_identity.py 2>/dev/null | grep -E "SEAM|conv|case|^ *[0-9A-Z]" diff --git a/alignment_lab/analysis/fw_footing.py b/alignment_lab/analysis/fw_footing.py index 548709f8..d0b67970 100644 --- a/alignment_lab/analysis/fw_footing.py +++ b/alignment_lab/analysis/fw_footing.py @@ -42,8 +42,9 @@ def main() -> int: french_wilson_preprocess, ) from torchref.experimental.alignment.frf.preprocessing import ( - build_lerf1_intensity, wilson_normalise, + build_lerf1_intensity, ) + from lab.reference_normalisers import wilson_normalise for pdb in ("1DAW", "3K7M", "2DQ6"): model, data = load_case(pdb) diff --git a/alignment_lab/analysis/seam_identity.py b/alignment_lab/analysis/seam_identity.py index 90b5494c..f1867c25 100644 --- a/alignment_lab/analysis/seam_identity.py +++ b/alignment_lab/analysis/seam_identity.py @@ -50,8 +50,9 @@ def main() -> int: french_wilson_preprocess, ) from torchref.experimental.alignment.frf.preprocessing import ( - build_lerf1_intensity, wilson_normalise, + build_lerf1_intensity, ) + from lab.reference_normalisers import wilson_normalise from torchref.experimental.alignment.sh import ( assign_shells, equal_count_shell_edges, ) diff --git a/alignment_lab/diagnostics/frf_prep_compare.py b/alignment_lab/diagnostics/frf_prep_compare.py index 43685f0a..b0cec122 100644 --- a/alignment_lab/diagnostics/frf_prep_compare.py +++ b/alignment_lab/diagnostics/frf_prep_compare.py @@ -194,7 +194,7 @@ def attribute_obs_terms(terms_csv: Path, pdb: str) -> dict: / (float(np.abs(Esqr).max()) or 1.0)) # Our normalised E^2 for the same Miller indices. - from torchref.experimental.alignment.frf.preprocessing import wilson_normalise + from lab.reference_normalisers import wilson_normalise model, data = load_case(pdb) B = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() our_hkl = data.hkl.to(torch.long).cpu() diff --git a/alignment_lab/lab/profile.py b/alignment_lab/lab/profile.py index 16056594..26e3930f 100644 --- a/alignment_lab/lab/profile.py +++ b/alignment_lab/lab/profile.py @@ -42,7 +42,11 @@ #: this wrong left 85% of the runtime unattributed. FRF_STAGES: Tuple[Tuple[str, str], ...] = ( ("torchref.experimental.alignment.frf.dense_calc", "dense_calc_via_box"), - ("torchref.experimental.alignment.frf.api", "french_wilson_preprocess"), + # Patched on its DEFINING module, not on `api`: it is reached through + # `FrenchWilsonE._compute`, which imports it inside the method body, so the + # lookup happens at call time and `api` never holds a reference at all. + ("torchref.experimental.alignment.frf.french_wilson", + "french_wilson_preprocess"), ("torchref.experimental.alignment.frf.api", "bessel_sh_expand"), ("torchref.experimental.alignment.frf.api", "cross_correlate_xi"), ("torchref.experimental.alignment.frf.api", "evaluate_rotation_function"), @@ -60,7 +64,10 @@ # imports `apply_overall_anisotropy` from `sh` at module top, and # `fit_relative_wilson_b` is imported inside `search_peaks` so it has to be # patched on the defining module. - ("torchref.experimental.alignment.frf.api", "wilson_normalise"), + # No `wilson_normalise` row any more. The observed-side normalisation now + # arrives as the `e_convention` CLASS and is called through a parameter, so + # there is no module attribute to patch -- and a row that registers zero + # calls is worse than no row, because it reads as "that stage is free". ("torchref.experimental.alignment.frf.api", "eterm_sigma_a"), ("torchref.experimental.alignment.frf.api", "build_lerf1_intensity"), ("torchref.experimental.alignment.frf.api", "apply_shell_variance_weights"), diff --git a/alignment_lab/lab/reference_normalisers.py b/alignment_lab/lab/reference_normalisers.py new file mode 100644 index 00000000..59bddc83 --- /dev/null +++ b/alignment_lab/lab/reference_normalisers.py @@ -0,0 +1,53 @@ +"""E-value normalisers as they were before the convention seam, frozen. + +These are not production code and are not imported by it. They are kept so a +future convention can be compared against what the rotation function actually +shipped, rather than against a description of it -- the comparison the seam was +originally validated with, and the one any replacement will want again. + +Frozen means frozen: if a production convention changes, these do not follow. +That is the whole point of an oracle. +""" + +from __future__ import annotations + +import torch + +from torchref.experimental.alignment.sh import ( + assign_shells, equal_count_shell_edges, +) + + +def wilson_normalise( + F: torch.Tensor, + s_mag: torch.Tensor, + n_shells: int = 20, +): + """Per-shell Wilson normalisation of amplitudes. + + Source: Phaser's ``Feff[r] / SIGMAN.sqrt_epsnSN[r]`` (``DataMR.cc:925``) + minus French-Wilson + explicit ε (``F`` is assumed anisotropy-corrected + by the caller). + + E_h = F_h / sqrt(_p) where p = shell containing h. + + Returns ``(E_h, sqrt_mean_F2_per_h)``. + """ + edges, _ = equal_count_shell_edges(s_mag, n_shells) + shell_idx = assign_shells(s_mag, edges) + valid = shell_idx >= 0 + F_dtype = F.dtype + F2 = F * F + count = torch.zeros(n_shells, dtype=torch.int64, device=F.device) + sumF2 = torch.zeros(n_shells, dtype=F_dtype, device=F.device) + F2_v = F2[valid] + idx_v = shell_idx[valid] + count.index_add_(0, idx_v, torch.ones_like(idx_v)) + sumF2.index_add_(0, idx_v, F2_v) + mean_F2 = sumF2 / count.clamp(min=1).to(F_dtype) + mean_F2 = mean_F2.clamp(min=1e-12) + sqrt_mean = mean_F2.sqrt() + per_h = torch.ones_like(F) + per_h[valid] = sqrt_mean[idx_v] + E = F / per_h + return E, per_h diff --git a/docs/changelog.rst b/docs/changelog.rst index fb99c525..46274d17 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -13,6 +13,9 @@ Unreleased - The rotation function and the ML rescore take their E-value convention as a class, so the observed and calculated sides are normalised by one rule rather than by nine private converters - The ML rescore now receives the observation sigmas, which the rotation function computed the French-Wilson posterior from and then discarded - Removed the rescore's ``scat_mode``, a second knob for the decision the E convention already makes +- The Sim rescore and the three translation-search sites take their E values from the convention too, so the alignment package has one normaliser rather than nine +- Removed ``wilson_normalise`` and ``wilson_normalise_epsilon``, which the convention replaced +- Fixed the rotation search, placement pipeline and ``align`` reading ``model.initialized``, which moved to ``model.ctx``; the tests covering it are slow-marked, so the break was invisible to a default test run - Replaced the fast rotation function's keyword surface with ``rotation_search(model, data, model_error_A)``; the caller's coordinate error is now used rather than overwritten by an estimate from the atom count - Removed the rotation function's dead modules, engine variants, debug environment switches and unreachable knobs - ``Model``'s iso/aniso partition is now derived on access instead of being rebuilt eagerly, so a copy cannot inherit a stale one diff --git a/tests/unit/alignment/test_e_conventions.py b/tests/unit/alignment/test_e_conventions.py index 8760f76f..c5644e85 100644 --- a/tests/unit/alignment/test_e_conventions.py +++ b/tests/unit/alignment/test_e_conventions.py @@ -123,21 +123,26 @@ def test_french_wilson_refuses_to_run_without_sigmas(): FrenchWilsonE(F, s, centric, n_shells=20) -def test_the_frf_calc_path_is_bit_identical_to_wilson_normalise(): - """The seam must be inert while the default convention is in place. - - Everything downstream of ``E_calc`` is untouched, so bit-identity here is - what makes the peak list unchanged rather than merely similar. +def test_wilson_shell_e_is_its_defining_formula(): + """``E = F / sqrt(_shell)``, asserted against the definition itself. + + Not against a reference implementation: the one this replaced now lives in + `alignment_lab/lab/reference_normalisers.py` as a frozen oracle for + comparing future conventions, and a unit test should not reach into the lab + to find its expectation. Restating the formula here is the specification, + not a second copy of the code. """ - from torchref.experimental.alignment.frf.preprocessing import ( - wilson_normalise, - ) - F, s, centric, _ = _wilson_data() - old, _ = wilson_normalise(F, s, 20) - new = WilsonShellE(F, s, centric, n_shells=20).E - assert torch.equal(new, old), ( - f"max deviation {float((new - old).abs().max()):.3e}" + conv = WilsonShellE(F, s, centric, n_shells=20) + + total = torch.zeros(20, dtype=F.dtype) + total.scatter_add_(0, conv.shell_idx, F * F) + count = torch.bincount(conv.shell_idx, minlength=20).to(F.dtype).clamp(min=1.0) + expected = F / (total / count).clamp(min=1e-30).index_select( + 0, conv.shell_idx).sqrt() + + assert torch.equal(conv.E, expected), ( + f"max deviation {float((conv.E - expected).abs().max()):.3e}" ) diff --git a/tests/unit/alignment/test_interp_var_and_shared_sigma_a.py b/tests/unit/alignment/test_interp_var_and_shared_sigma_a.py index 6e8cf13c..91142727 100644 --- a/tests/unit/alignment/test_interp_var_and_shared_sigma_a.py +++ b/tests/unit/alignment/test_interp_var_and_shared_sigma_a.py @@ -16,7 +16,6 @@ from torchref.experimental.alignment.ml_rotation import ( _equal_count_shell_idx, - _normalize_to_e, _optimize_D_in_shell, _shell_ll, fit_sigma_a_per_shell, diff --git a/torchref/experimental/alignment/align.py b/torchref/experimental/alignment/align.py index 43124da8..2cca2784 100644 --- a/torchref/experimental/alignment/align.py +++ b/torchref/experimental/alignment/align.py @@ -322,7 +322,7 @@ def align_model_to_data( `MolecularReplacementPipeline` is the implementation of record; this function returns its single best solution. """ - if not model.initialized: + if not model.ctx.initialized: raise RuntimeError( "Cannot fit an uninitialized ModelFT. Load PDB data first." ) diff --git a/torchref/experimental/alignment/frf/preprocessing.py b/torchref/experimental/alignment/frf/preprocessing.py index 4d794230..07f2bfec 100644 --- a/torchref/experimental/alignment/frf/preprocessing.py +++ b/torchref/experimental/alignment/frf/preprocessing.py @@ -26,41 +26,6 @@ ) -def wilson_normalise( - F: torch.Tensor, - s_mag: torch.Tensor, - n_shells: int = 20, -): - """Per-shell Wilson normalisation of amplitudes. - - Source: Phaser's ``Feff[r] / SIGMAN.sqrt_epsnSN[r]`` (``DataMR.cc:925``) - minus French-Wilson + explicit ε (``F`` is assumed anisotropy-corrected - by the caller). - - E_h = F_h / sqrt(_p) where p = shell containing h. - - Returns ``(E_h, sqrt_mean_F2_per_h)``. - """ - edges, _ = equal_count_shell_edges(s_mag, n_shells) - shell_idx = assign_shells(s_mag, edges) - valid = shell_idx >= 0 - F_dtype = F.dtype - F2 = F * F - count = torch.zeros(n_shells, dtype=torch.int64, device=F.device) - sumF2 = torch.zeros(n_shells, dtype=F_dtype, device=F.device) - F2_v = F2[valid] - idx_v = shell_idx[valid] - count.index_add_(0, idx_v, torch.ones_like(idx_v)) - sumF2.index_add_(0, idx_v, F2_v) - mean_F2 = sumF2 / count.clamp(min=1).to(F_dtype) - mean_F2 = mean_F2.clamp(min=1e-12) - sqrt_mean = mean_F2.sqrt() - per_h = torch.ones_like(F) - per_h[valid] = sqrt_mean[idx_v] - E = F / per_h - return E, per_h - - def eterm_sigma_a(s_mag: torch.Tensor, delta_vrms_A: float) -> torch.Tensor: """Phaser's σA Eterm, literal port of ``Ensemble.cc:42``: @@ -74,8 +39,6 @@ def eterm_sigma_a(s_mag: torch.Tensor, delta_vrms_A: float) -> torch.Tensor: return torch.exp(-(2.0 / 3.0) * (math.pi ** 2) * s2 * (delta_vrms_A ** 2)) __all__ = [ - "wilson_normalise", - "wilson_normalise_epsilon", "eterm_sigma_a", "french_wilson_preprocess", "get_high_order_axis", @@ -160,38 +123,6 @@ def epsilon_aware_unroll( return unrolled_hkl, asu_idx -def wilson_normalise_epsilon( - F: torch.Tensor, - s_mag: torch.Tensor, - epsilon: torch.Tensor, - n_shells: int = 20, -): - """Epsilon-corrected per-shell Wilson normalisation. - - Standard crystallographic normalisation with the multiplicity factor: - - Σ_shell = ⟨I_h / ε_h⟩_shell - E²_h = (I_h / ε_h) / Σ_shell - - so axial reflections (large ε) are not over-counted. Returns - ``(E, sqrt_mean_eps_corrected)`` mirroring ``wilson_normalise``. - """ - edges, _ = equal_count_shell_edges(s_mag, n_shells) - shell_idx = assign_shells(s_mag, edges) - valid = shell_idx >= 0 - I_corr = (F * F) / epsilon.clamp(min=1.0) - count = torch.zeros(n_shells, dtype=torch.int64, device=F.device) - sumI = torch.zeros(n_shells, dtype=F.dtype, device=F.device) - idx_v = shell_idx[valid] - count.index_add_(0, idx_v, torch.ones_like(idx_v)) - sumI.index_add_(0, idx_v, I_corr[valid]) - mean_I = (sumI / count.clamp(min=1).to(F.dtype)).clamp(min=1e-12) - per_h = torch.ones_like(F) - per_h[valid] = mean_I[idx_v] - E = (I_corr / per_h).clamp(min=0.0).sqrt() - return E, per_h.sqrt() - - def build_lerf1_intensity( eEobs: torch.Tensor, centric_obs: torch.Tensor, diff --git a/torchref/experimental/alignment/ml_rotation.py b/torchref/experimental/alignment/ml_rotation.py index c3618c72..1609b2c2 100644 --- a/torchref/experimental/alignment/ml_rotation.py +++ b/torchref/experimental/alignment/ml_rotation.py @@ -25,7 +25,7 @@ import torch -from .e_values import (CalcShellE, WilsonShellEpsE, +from .e_values import (CalcShellE, WilsonShellE, WilsonShellEpsE, convention_for_calc, convention_uses_sigma_f) from .frf.rotation_utils import ( axis_angle_to_matrix, @@ -67,28 +67,6 @@ def _equal_count_shell_idx(s_mag: torch.Tensor, n_shells: int) -> torch.Tensor: return shell_idx -def _normalize_to_e(F: torch.Tensor, shell_idx: torch.Tensor, - n_shells: int) -> torch.Tensor: - """E = F / sqrt( per shell). Vectorised across shells via scatter.""" - return F / _per_shell_sqrt_mean(F, shell_idx, n_shells) - - -def _per_shell_sqrt_mean(F: torch.Tensor, shell_idx: torch.Tensor, - n_shells: int) -> torch.Tensor: - """Per-reflection ``sqrt(_shell)``. Wilson-normalisation denominator. - - Use this when you need to normalise *another* tensor by the same per-shell - statistic computed from F — e.g. converting rotated |F_calc| to E_calc - using the reference |F_calc|'s shell means (rotation-invariant). - """ - F2 = F ** 2 - sum_per_shell = torch.zeros(n_shells, dtype=F2.dtype, device=F.device) - sum_per_shell.scatter_add_(0, shell_idx, F2) - count_per_shell = torch.bincount(shell_idx, minlength=n_shells).to(F2.dtype) - mean_per_shell = (sum_per_shell / count_per_shell.clamp(min=1.0)).clamp(min=1e-30) - return mean_per_shell.sqrt().index_select(0, shell_idx) - - def _shell_ll( E_obs: torch.Tensor, E_calc: torch.Tensor, @@ -451,7 +429,13 @@ def sim_mlrf_rescore( tail = peaks[n_refine:] shell_idx = _equal_count_shell_idx(s_mag, n_shells) - E_obs = _normalize_to_e(F_obs, shell_idx, n_shells) + # The caller's own shell assignment is passed in rather than letting the + # convention derive one: `_equal_count_shell_idx` is rank-based and the + # shared `assign_shells` is value-based, so they disagree on reflections + # sitting on a boundary. Keeping this one makes the migration exact. + E_obs = WilsonShellE( + F_obs, s_mag, shell_idx=shell_idx, n_shells=n_shells, + ).E if shell_weights is None and auto_variance_weights: from .sh import compute_patterson_shell_variance @@ -808,7 +792,9 @@ def _build_sim_llg_context( ) -> _SimLLGContext: """Build the rotation-independent context for the Sim-LLG surface.""" shell_idx = _equal_count_shell_idx(s_mag, n_shells) - E_obs = _normalize_to_e(F_obs, shell_idx, n_shells) + E_obs = WilsonShellE( + F_obs, s_mag, shell_idx=shell_idx, n_shells=n_shells, + ).E shell_weights = None if auto_variance_weights: from .sh import compute_patterson_shell_variance diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index ddba4ead..b6de91a7 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -341,7 +341,7 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: Sorted by ``r_factor`` (ascending). The first element is the best placement; its ``r_factor`` is the solvent-aware Scaler R-work. """ - if not self.model.initialized: + if not self.model.ctx.initialized: raise RuntimeError( "Cannot fit an uninitialized ModelFT. Load PDB data first." ) @@ -784,7 +784,7 @@ def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): ) llg_tf = llg_translation_rescore( F_obs=F_obs_amp, hkl=hkl_keep, centric=centric_keep_tf, - shell_idx=tf_shell_idx, n_shells=tf_n_shells, + s_mag=s_mag_keep_tf, shell_idx=tf_shell_idx, n_shells=tf_n_shells, G=G_pre, h_R=h_R_pre, t_candidates=t_cands, sigma_a=sigma_a_tf, interp_var=None, ) diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index a3eba577..b865c7a0 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -466,7 +466,7 @@ def rotation_search( ValueError If the data carry too few reflections to bin. """ - if not model.initialized: + if not model.ctx.initialized: raise RuntimeError("model has no coordinates; load a PDB first.") d_max_fit, d_min_fit = ANISO_FIT_WINDOW_A U_aniso = fit_anisotropy(data, d_min=d_min_fit, d_max=d_max_fit) diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index 9d324458..ab81b33e 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -10,6 +10,8 @@ import numpy as np import torch + +from .e_values import WilsonShellE from dataclasses import dataclass from typing import List, Optional, Tuple @@ -341,12 +343,13 @@ def amplitude_translation_search( a = k * chunk b = (k + 1) * chunk if k < n_shells - 1 else s_mag.numel() shell_idx[order[a:b]] = k - shell_norm_obs = torch.zeros(n_shells, dtype=real_dtype, device=device) - for k in range(n_shells): - m = shell_idx == k - if m.any(): - shell_norm_obs[k] = (F_obs_t[m] ** 2).mean().clamp(min=1e-30).sqrt() - E_obs = F_obs_t / shell_norm_obs[shell_idx] + # The caller's own shell assignment is handed to the convention rather + # than letting it derive one: this binning is rank-based and the shared + # `assign_shells` is value-based, and the sigma_a fit downstream is tied + # to whichever one was used here. + E_obs = WilsonShellE( + F_obs_t, s_mag, shell_idx=shell_idx, n_shells=n_shells, + ).E F_obs2 = E_obs * E_obs else: shell_idx = None @@ -428,6 +431,7 @@ def llg_translation_rescore( F_obs: torch.Tensor, hkl: torch.Tensor, centric: torch.Tensor, + s_mag: torch.Tensor, shell_idx: torch.Tensor, n_shells: int, G: torch.Tensor, @@ -457,6 +461,9 @@ def llg_translation_rescore( hkl : (N, 3) — unused here but kept for symmetry with the rest of the module (and future extension to per-h variance models). centric : (N,) bool + s_mag : (N,) — |s| the shells were built from. Not used to derive a binning + here (``shell_idx`` is given) but passed rather than fabricated, so a + convention that fits a curve in ``|s|`` gets the real abscissa. shell_idx : (N,) int64 — same binning as used to fit sigma_a / interp_var. n_shells : int G : (S, N) complex — per-sym F_p1 contributions × per-sym translation phase @@ -499,10 +506,10 @@ def llg_translation_rescore( E_calc = F_calc / norm_per_refl # (K, N) F_obs_t = F_obs.to(device).to(real_dtype) - sum_F_obs2 = torch.zeros(n_shells, dtype=real_dtype, device=device) - sum_F_obs2.scatter_add_(0, shell_idx_l, F_obs_t * F_obs_t) - mean_F_obs2 = (sum_F_obs2 / cnt.clamp(min=1.0)).clamp(min=1e-30) - E_obs = F_obs_t / mean_F_obs2.sqrt().index_select(0, shell_idx_l) + E_obs = WilsonShellE( + F_obs_t, s_mag.to(device).to(real_dtype), + shell_idx=shell_idx_l, n_shells=n_shells, + ).E sigma_a_d = sigma_a.to(device).to(real_dtype) # (n_shells,) D_per_refl = sigma_a_d.index_select(0, shell_idx_l) # (N,) @@ -820,12 +827,9 @@ def local_rotation_translation_refine( a = k * chunk b = (k + 1) * chunk if k < n_shells - 1 else s_mag.numel() shell_idx[order[a:b]] = k - shell_sigma_obs = torch.zeros(n_shells, dtype=real_dtype, device=device) - for k in range(n_shells): - m = shell_idx == k - if m.any(): - shell_sigma_obs[k] = (F_obs_t[m] ** 2).mean().clamp(min=1e-30).sqrt() - E_obs = F_obs_t / shell_sigma_obs[shell_idx] + E_obs = WilsonShellE( + F_obs_t, s_mag, shell_idx=shell_idx, n_shells=n_shells, + ).E # Rotation perturbation grid. # We parametrise (Δα, Δβ, Δγ) ∈ [-r, r]³ via the small-angle rotation From de6ca169733bdbdb5221af69c270be7a412de609 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sat, 29 Aug 2026 14:55:24 +0200 Subject: [PATCH 096/250] Restore state dicts on CPU and move only when asked create_from_state_dict took a device argument but passed it to exactly one of the four parameter wrappers. OccupancyTensor honoured it; xyz, adp and u are built from the atom table via torch.tensor and land on CPU regardless. Restoring with the default device resolved to an accelerator therefore produced a model split across two devices, which CPU-only CI cannot see and the accelerator workflow trips over. The restore now builds on CPU throughout and moves once at the end, and only when the caller names a device. None leaves the model on CPU rather than resolving to device.current: reading a file back is not a reason to claim an accelerator, and a caller who wants one can say so or move the model itself. Measured with device.current = cuda, for isotropic, field and field_aniso alike: device=None gives all four wrappers on cpu, device=cuda gives all four on cuda:0, and the moved model computes structure factors. Co-Authored-By: Claude Opus 5 (1M context) --- docs/changelog.rst | 1 + tests/unit/model/test_adp_field_mode.py | 5 ++++- torchref/model/model.py | 21 +++++++++++++++------ torchref/model/model_ft.py | 17 ++++++++++++----- 4 files changed, 32 insertions(+), 12 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index fffa01d3..744eae0a 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -8,6 +8,7 @@ Version 0.7.0 - Removed the stored real-space coordinate grid; ``build_electron_density`` takes a grid shape and device, and ``ModelFT.real_space_grid()`` builds one on demand - Fixed ``ModelFT`` restore dropping a node-field ADP representation, and added the anisotropic ``field_aniso`` case; both models now share one wrapper-rebuild path - Fixed the node load and node smoothness restraints being inert in ``field_aniso`` mode +- ``create_from_state_dict`` now restores on CPU and moves only when passed a device; it previously left three of the four parameter wrappers on CPU while claiming the default device - Separated model configuration and provenance into ``ModelContext``. It now holds the unit cell, space group, atom table, link records, hydrogen settings, and input paths. - Refactored ``Symmetry`` as a crystallography-free class with transform primitives, and made ``SpaceGroup`` a specialised subclass. - Moved geometry predicates, HKL verbs, and grid-size helpers onto these classes as methods. diff --git a/tests/unit/model/test_adp_field_mode.py b/tests/unit/model/test_adp_field_mode.py index 32f6c772..14a8c1d1 100644 --- a/tests/unit/model/test_adp_field_mode.py +++ b/tests/unit/model/test_adp_field_mode.py @@ -207,7 +207,10 @@ def test_state_dict_round_trip_in_field_mode(pdb_path): key: (value.clone() if torch.is_tensor(value) else value) for key, value in model.state_dict().items() } - restored = Model.create_from_state_dict(sd, verbose=0) + # Restored onto the model's own device: a restore builds on CPU and moves only + # when asked, and comparing a CPU-evaluated field against a device-evaluated one + # would be measuring backend arithmetic, not the round trip. + restored = Model.create_from_state_dict(sd, device=model.device, verbose=0) assert restored.adp_is_field assert restored.adp.n_nodes == 24 diff --git a/torchref/model/model.py b/torchref/model/model.py index 1a6fd757..3def94be 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -19,7 +19,7 @@ import torch.nn as nn from torchref.base import math_torch -from torchref.config import get_float_dtype, normalize_device +from torchref.config import canonical_device, get_float_dtype, normalize_device from torchref.io import cif, pdb from torchref.model.context import ModelContext from torchref.model.parameter_wrappers import ( @@ -2254,7 +2254,10 @@ def create_from_state_dict( state_dict : dict State dictionary from torch.save(model.state_dict(), ...). device : torch.device, optional - Device to place tensors on. Defaults to the configured device.current. + Move the restored model here once it is built. The restore itself always + runs on CPU, and ``None`` leaves it there rather than resolving to + ``device.current`` -- loading a file is not a reason to claim an + accelerator. Move it yourself, or pass one here. verbose : int, optional Verbosity level. Default is 1. dtype_float : torch.dtype, optional @@ -2271,9 +2274,12 @@ def create_from_state_dict( anisotropic ``u`` is rebuilt as a :class:`CholeskyMixedTensor`, matching :meth:`load`, so the positive-definite parametrization round-trips. """ - # Resolve dtype/device at call time so the fallbacks below use the - # current config, not the import-time default. - device = normalize_device(device) + # Build on CPU throughout, then move once at the end if the caller named a + # device. One device for the whole model is the invariant that matters: the + # wrappers are built from the atom table and land on CPU whatever is asked for, + # so resolving an accelerator up front splits the model rather than placing it. + target_device = canonical_device(device) if device is not None else None + device = torch.device("cpu") if dtype_float is None: dtype_float = get_float_dtype() pdb = state_dict.pop("pdb", None) @@ -2281,7 +2287,7 @@ def create_from_state_dict( spacegroup = state_dict.pop("spacegroup", None) initialized = state_dict.pop("initialized", False) saved_dtype = state_dict.pop("dtype_float", dtype_float) - saved_device = state_dict.pop("device", device) + state_dict.pop("device", None) # popped so it never reaches load_state_dict strip_H = state_dict.pop("strip_H", True) altloc_pairs = state_dict.pop("altloc_pairs", []) @@ -2313,6 +2319,9 @@ def create_from_state_dict( } instance.load_state_dict(state_dict, strict=False) + if target_device is not None: + instance.to(target_device) + if verbose > 0: n_atoms = len(instance.pdb) if instance.pdb is not None else 0 print(f"Created Model from state_dict: {n_atoms} atoms") diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index 31abe453..d9b06cf9 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -13,7 +13,7 @@ import torch from torchref.base.fourier import fft, ifft -from torchref.config import dtypes, get_float_dtype, normalize_device +from torchref.config import canonical_device, dtypes, get_float_dtype from torchref.model.model import Model from torchref.model.sf_fft import SfFFT from torchref.symmetry import SpaceGroup @@ -924,7 +924,9 @@ def create_from_state_dict( state_dict : dict State dictionary from torch.save(model.state_dict(), ...). device : torch.device, optional - Device to place tensors on. Defaults to the configured device.current. + Move the restored model here once it is built. The restore itself always + runs on CPU, and ``None`` leaves it there; see + :meth:`Model.create_from_state_dict`. verbose : int, optional Verbosity level. Default is 1. dtype_float : torch.dtype, optional @@ -942,9 +944,11 @@ def create_from_state_dict( The anisotropic ``u`` is rebuilt as a :class:`CholeskyMixedTensor`, as in :meth:`load`, so the positive-definite parametrization round-trips. """ - # Resolve dtype/device at call time so the fallback below uses the - # current config rather than an import-time default. - device = normalize_device(device) + # Build on CPU throughout and move once at the end, as Model does; the grid + # setup below otherwise sizes an accelerator allocation for a model the caller + # has not asked to put there. + target_device = canonical_device(device) if device is not None else None + device = torch.device("cpu") if dtype_float is None: dtype_float = get_float_dtype() @@ -1030,6 +1034,9 @@ def create_from_state_dict( instance.load_state_dict(filtered_state_dict, strict=False) + if target_device is not None: + instance.to(target_device) + instance.reset_cache() if verbose > 0: From 14137fbf081588a7ab33d040d46ec1b3468da064 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sat, 29 Aug 2026 14:55:34 +0200 Subject: [PATCH 097/250] Stop three tests assuming the default device is CPU All three pass on CPU and fail under the accelerator workflow, which is the only place that runs with TORCHREF_DEVICE set. test_anomalous builds hkl with a plain torch.tensor while the model sits on device.current. Nothing moves it: the reciprocal extractor takes its device from hkl by design, so the phases end up on CPU against a map on the accelerator. Left as an error rather than moving hkl inside forward -- a caller handing a model and its reflections different devices should hear about it -- so the tests now build hkl on model.device. test_entries_survive_a_device_apply clones its reference before calling .to(torch.device("cpu")), so the two are only on the same device when the default already is CPU. Compares against before.cpu() now. test_empty_shell_needs_no_accessor asks for float64, which MPS does not have. Every other case in that file inherits CPU from the tensors it passes in; only the empty shell resolves device=None to device.current, so it is pinned to CPU. Co-Authored-By: Claude Opus 5 (1M context) --- tests/unit/model/test_disorder_field.py | 5 ++++- tests/unit/scattering/test_anomalous.py | 22 ++++++++++++++++++---- tests/unit/topology/test_storage.py | 4 +++- 3 files changed, 25 insertions(+), 6 deletions(-) diff --git a/tests/unit/model/test_disorder_field.py b/tests/unit/model/test_disorder_field.py index 1ee37903..dcb8f11d 100644 --- a/tests/unit/model/test_disorder_field.py +++ b/tests/unit/model/test_disorder_field.py @@ -388,7 +388,10 @@ def test_state_dict_excludes_the_accessor_and_round_trips(coords, target_b): @pytest.mark.unit def test_empty_shell_needs_no_accessor(coords): """The ``load_state_dict`` entry point constructs without coordinates.""" - shell = DisorderFieldTensor(dtype=torch.float64) + # Pinned to CPU: every other case here inherits CPU from the tensors it is + # handed, but the empty shell resolves ``device=None`` to ``device.current``, + # and float64 does not exist on MPS. + shell = DisorderFieldTensor(dtype=torch.float64, device="cpu") assert shell.neighbor_list is None assert shell.shape == (0,) diff --git a/tests/unit/scattering/test_anomalous.py b/tests/unit/scattering/test_anomalous.py index 145f9b3a..6fb0951b 100644 --- a/tests/unit/scattering/test_anomalous.py +++ b/tests/unit/scattering/test_anomalous.py @@ -246,7 +246,9 @@ def test_modelft_disable_anomalous(self, test_pdb_file): assert model.wavelength is None # Create HKL reflections - hkl = torch.tensor([[1, 0, 0], [0, 1, 0], [0, 0, 1]], dtype=torch.int32) + hkl = torch.tensor( + [[1, 0, 0], [0, 1, 0], [0, 0, 1]], dtype=torch.int32, device=model.device + ) # Should compute structure factors without anomalous correction sf = model.get_structure_factor(hkl) @@ -260,7 +262,11 @@ def test_anomalous_correction_applied(self, test_pdb_file): model = ModelFT(wavelength=1.0, anomalous_threshold=0.5, verbose=0) model.load_pdb(test_pdb_file) - hkl = torch.tensor([[1, 0, 0], [2, 1, 0], [1, 1, 1]], dtype=torch.int32) + hkl = torch.tensor( + + [[1, 0, 0], [2, 1, 0], [1, 1, 1]], dtype=torch.int32, device=model.device + + ) # Compute with anomalous correction sf_with = model.get_structure_factor( @@ -288,7 +294,11 @@ def test_friedel_pair_asymmetry(self, test_pdb_file): model = ModelFT(wavelength=1.0, anomalous_threshold=0.5, verbose=0) model.load_pdb(test_pdb_file) - hkl = torch.tensor([[1, 2, 3], [2, 1, 0], [3, 3, 3]], dtype=torch.int32) + hkl = torch.tensor( + + [[1, 2, 3], [2, 1, 0], [3, 3, 3]], dtype=torch.int32, device=model.device + + ) sf_plus = model.get_structure_factor(hkl, apply_anomalous=True, recalc=True) sf_minus = model.get_structure_factor(-hkl, apply_anomalous=True, recalc=True) @@ -361,7 +371,11 @@ def test_gradient_flow(self, test_pdb_file): model.load_pdb(test_pdb_file) # xyz.refinable_params should already have requires_grad=True by default - hkl = torch.tensor([[1, 0, 0], [0, 1, 0], [0, 0, 1]], dtype=torch.int32) + hkl = torch.tensor( + + [[1, 0, 0], [0, 1, 0], [0, 0, 1]], dtype=torch.int32, device=model.device + + ) sf = model.get_structure_factor(hkl, apply_anomalous=True, recalc=True) diff --git a/tests/unit/topology/test_storage.py b/tests/unit/topology/test_storage.py index 6264195a..d3423297 100644 --- a/tests/unit/topology/test_storage.py +++ b/tests/unit/topology/test_storage.py @@ -156,7 +156,9 @@ def test_entries_survive_a_device_apply(restraints): block = restraints.topology.atoms.bonds.indices after = restraints.restraints["bond"]["all"]["indices"] assert after.data_ptr() == block.data_ptr(), "entries no longer alias the block" - assert torch.equal(after, before) + # ``before`` was cloned prior to the move, so it sits on device.current; the move + # target here is CPU, which is only a no-op when those already agree. + assert torch.equal(after, before.cpu()) assert { t: restraints.restraints[t]["all"]["indices"].shape[0] for t in KEYED_TYPES } == n_before From 5371a48d93da12760eb322bdd7d4e9846aa06546 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sat, 29 Aug 2026 15:50:19 +0200 Subject: [PATCH 098/250] Give the rigid-body step its own loss state _run_one_cutoff borrowed the refinement's LossState, dropped every non-x-ray target from it, and restored afterwards. The restore was dead code: it snapshotted state.weights and then popped the non-x-ray names from the snapshot as well, so the finally clause wrote back only "xray" -- a value nothing in the block had touched. step() reads weights through get_effective_weight and never writes them, and the sole mutation point is set_weights(self.weighting(state)) at construction. The dropped targets were never restored at all, which was harmless only because the step runs against a sandbox clone and _rebind_for_data rebuilt the state each cutoff. The step now builds a LossState of its own holding the x-ray target alone. Nothing has to be undone afterwards, no maintenance() hook of a target that is not in the sum can fire, and the caller's state keeps its targets and any weights registered on them. Weight 1.0: with a single term the weight is a scalar on the whole objective, and 1.0 is what DEFAULT_GROUP_WEIGHTS gives x-ray anyway, so the effective objective is unchanged. _rebind_for_data drops reset_loss_state, which nothing reads now, and calls _build_xray_targets + get_scales in place of _init_targets. The x-ray target genuinely changes with each cutoff's data and target mode; TotalGeometryTarget and TotalADPTarget do not, and were being constructed five times per run -- three cutoffs, the full-resolution restore, the commit rebind -- only to go unused, the NonBondedTarget pair list among them. Counted on 3E98 that goes from 5 and 5 to 0 and 0. The cutoff schedule, the per-cutoff resolution truncation and the 6 A target-mode switch are untouched. Results are unchanged within run-to-run spread: over 3E98 and 1DAW the largest old-vs-new difference is 4.52e-02 A / +5.0e-04 R-free, against a same-code null control of 4.51e-02 A / +4.7e-04. Co-Authored-By: Claude Opus 5 (1M context) --- torchref/refinement/rigid_body_refinement.py | 163 +++++++++---------- 1 file changed, 78 insertions(+), 85 deletions(-) diff --git a/torchref/refinement/rigid_body_refinement.py b/torchref/refinement/rigid_body_refinement.py index 3f320d2b..515f54bc 100644 --- a/torchref/refinement/rigid_body_refinement.py +++ b/torchref/refinement/rigid_body_refinement.py @@ -13,6 +13,7 @@ import torch +from torchref.refinement.loss_state import LossState from torchref.scaling.scaler import Scaler @@ -94,20 +95,17 @@ def _sandbox(ref): """A shallow clone of ``ref`` that shares its model but owns its namespace. Every cutoff calls :meth:`_rebind_for_data`, which assigns - ``reflection_data``, builds a fresh ``Scaler``, and calls - ``_init_targets`` + ``reset_loss_state``. Run against the real - Refinement those assignments are destructive: ``_init_targets`` - reconstructs ``adp_target`` and ``geometry_target`` from constructor - defaults, so anything configured on them post-construction (for instance - ``adp_target['simu'].simu_sigma``) is silently reset, and - ``reset_loss_state`` discards custom weights registered on the - ``LossState``. Neither survives the step, and neither is wanted by it -- - rigid body drops every non-xray target before optimizing, so those - objects are rebuilt only to be thrown away. + ``reflection_data``, ``scaler`` and the x-ray targets for that cutoff's + resolution and target mode. Run against the real Refinement those + assignments are destructive -- the caller gets its data and scaler + silently replaced by whatever the last cutoff used. Directing the step at a clone confines all of it. The real Refinement is never written to, so there is nothing to restore and no window in which - it is inconsistent. + it is inconsistent. The step builds its own single-target + :class:`~torchref.refinement.loss_state.LossState` rather than borrowing + the refinement's, so the caller's targets and any weights registered on + them are untouched as well. The model is deliberately shared, not copied: ``use_rigid_xyz`` swaps its xyz container in place, so refined coordinates reach the caller by object @@ -209,9 +207,14 @@ def _rebind_for_data(self, data, model=None, xray_mode=None): device=ref.device, ) ref.scaler.initialize() - ref._init_targets(xray_mode=xray_mode) + # Only the x-ray half of _init_targets: the step optimizes against x-ray data + # alone, and TotalGeometryTarget / TotalADPTarget would be constructed here + # purely to be left unused -- NonBondedTarget's pair list among them. + # get_scales() still runs, because the x-ray target reads the scaler's + # parameters. + ref._build_xray_targets(xray_mode) + ref.get_scales() - ref.reset_loss_state() # Clear cached LBFGS optimizers (they were built over the old model's # parameters and would now point at stale leaves). if hasattr(ref, "_persistent_optimizers"): @@ -222,81 +225,71 @@ def _run_one_cutoff(self, d_min: float): ref = self.refinement rigid_model = ref.model - state = ref.complete_loss_state() - - # Snapshot weights so we can restore after the step. - original_weights = dict(state.weights) - try: - # Active targets during rigid-body refinement: xray only. Internal - # bonded geometry is rigid by construction; ADP / occupancy are - # frozen. Inter-chain vdW is intentionally off — Phenix runs - # rigid-body without atomistic restraints and we've observed vdW - # adds no signal here and destabilizes the coarsest cutoff. - # - # We DROP non-xray targets from state entirely (not just zero - # their weight) so their maintenance() hooks don't fire during - # the rigid-body LBFGS. In particular ``NonBondedTarget`` - # rebuilds its VDW pair list whenever atoms drift >1 Å, a - # multi-second recomputation that's pure waste when the - # target weight is 0. State is rebuilt fresh on the next - # cutoff via _rebind_for_data → reset_loss_state → - # _init_targets, so this drop is local to this cutoff. - keep_names = {"xray"} - for name in list(state.targets.keys()): - if name not in keep_names: - state.targets.pop(name, None) - original_weights.pop(name, None) - - rigid_params = [ - rigid_model.xyz.euler_angles, - rigid_model.xyz.translations, - ] - - # Decide whether to use the inner-cycle (mask-refresh) loop. - # Triggered when the scaler has a bulk-solvent component whose - # mask depends on atom positions (ls_wunit_k1 path here). For - # the ml path the scaler is fully refit between cutoffs - # and co-optimized with rigid params in a single LBFGS. - use_inner_cycles = ( - ref.scaler is not None - and getattr(ref.scaler, "solvent", None) is not None - and getattr(ref.scaler, "c_iso", None) is not None - and ref.scaler.c_iso.requires_grad is False - ) + # Active targets during rigid-body refinement: x-ray only. Internal bonded + # geometry is rigid by construction; ADP / occupancy are frozen. Inter-chain vdW + # is intentionally off -- Phenix runs rigid-body without atomistic restraints, + # and we have measured that vdW adds no signal here and destabilizes the + # coarsest cutoff. + # + # A state of its own rather than the refinement's: nothing here has to be undone + # afterwards, no maintenance() hook of a target we are not using can fire (in + # particular NonBondedTarget rebuilds its VDW pair list whenever atoms drift + # >1 A, seconds of work for a term that is not in the sum), and the caller's + # state keeps its targets and any weights registered on it. + # + # Weight 1.0: with one term the weight is a scalar on the whole objective, and + # 1.0 is what DEFAULT_GROUP_WEIGHTS gives x-ray anyway. + state = LossState(device=ref.device) + state.register_target("xray", ref.xray_target_work) + state.set_weight("xray", 1.0) + state.cache_losses() + + rigid_params = [ + rigid_model.xyz.euler_angles, + rigid_model.xyz.translations, + ] + + # Decide whether to use the inner-cycle (mask-refresh) loop. + # Triggered when the scaler has a bulk-solvent component whose + # mask depends on atom positions (ls_wunit_k1 path here). For + # the ml path the scaler is fully refit between cutoffs + # and co-optimized with rigid params in a single LBFGS. + use_inner_cycles = ( + ref.scaler is not None + and getattr(ref.scaler, "solvent", None) is not None + and getattr(ref.scaler, "c_iso", None) is not None + and ref.scaler.c_iso.requires_grad is False + ) - if use_inner_cycles: - self._run_inner_cycles(d_min, state, rigid_params, n_inner=5) + if use_inner_cycles: + self._run_inner_cycles(d_min, state, rigid_params, n_inner=5) + else: + # Single-shot path: rigid params + scaler params co-optimized. + if ref.scaler is not None: + opt_params = rigid_params + list(ref.scaler.parameters()) else: - # Single-shot path: rigid params + scaler params co-optimized. - if ref.scaler is not None: - opt_params = rigid_params + list(ref.scaler.parameters()) - else: - opt_params = rigid_params - - rigid_model.reset_cache() - opt = torch.optim.LBFGS( - opt_params, - max_iter=self.iterations_per_step, - **self.DEFAULT_LBFGS_KWARGS, - ) - state.step( - opt, - context=f"rigid_body[d_min={d_min:.2f}]", + opt_params = rigid_params + + rigid_model.reset_cache() + opt = torch.optim.LBFGS( + opt_params, + max_iter=self.iterations_per_step, + **self.DEFAULT_LBFGS_KWARGS, + ) + state.step( + opt, + context=f"rigid_body[d_min={d_min:.2f}]", + ) + if ref.verbose > 0: + try: + rwork, rfree = ref.get_rfactor() + print( + f" rigid-body d_min={d_min:.2f} " + f"(lbfgs, iters={self.iterations_per_step}): " + f"Rwork={rwork:.4f} Rfree={rfree:.4f}" ) - if ref.verbose > 0: - try: - rwork, rfree = ref.get_rfactor() - print( - f" rigid-body d_min={d_min:.2f} " - f"(lbfgs, iters={self.iterations_per_step}): " - f"Rwork={rwork:.4f} Rfree={rfree:.4f}" - ) - except Exception: - pass - finally: - # Restore weights. - for name, w in original_weights.items(): - state.set_weight(name, w) + except Exception: + pass return state def _run_inner_cycles(self, d_min, state, rigid_params, n_inner: int = 5): From b4563739635d953609fe66bceaf03c0db9027747 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sat, 29 Aug 2026 15:50:30 +0200 Subject: [PATCH 099/250] Locate the Wilson outlier fixtures through mtz_dir The file opened "tests/files/mtz/.mtz" -- a path relative to the working directory -- so all 13 tests in it failed with "No such file or directory" whenever pytest ran from anywhere but the repo root. The workflows run from the root, so CI never saw it. Running the suite from tests/, which is what makes --run-slow resolve, does see it, and the failures read as a standing red in the outlier screening rather than as a missing file. All four call sites now go through the mtz_dir fixture, as the rest of the suite does; test_deposited_structures_lose_almost_nothing was already requesting pdb_dir without using it. This was the only cwd-dependent data path left under tests/. Co-Authored-By: Claude Opus 5 (1M context) --- tests/unit/io/test_wilson_outlier_masks.py | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/tests/unit/io/test_wilson_outlier_masks.py b/tests/unit/io/test_wilson_outlier_masks.py index 201bdf2a..23f01b26 100644 --- a/tests/unit/io/test_wilson_outlier_masks.py +++ b/tests/unit/io/test_wilson_outlier_masks.py @@ -126,10 +126,10 @@ def test_planted_zingers_are_rejected(): @pytest.mark.unit -def test_absent_measurements_are_sanity_not_outliers(): +def test_absent_measurements_are_sanity_not_outliers(mtz_dir): """A row with no measurement is not an improbable observation, and counting it as one is what made the old report meaningless.""" - data = ReflectionData(verbose=0).load_mtz("tests/files/mtz/6G9X.mtz") + data = ReflectionData(verbose=0).load_mtz(str(mtz_dir / "6G9X.mtz")) absent = ~data.masks["sanity_F"] assert int(absent.sum()) > 20000, "6G9X carries a large absent population" @@ -140,8 +140,8 @@ def test_absent_measurements_are_sanity_not_outliers(): @pytest.mark.unit -def test_intensity_path_keeps_french_wilsons_guard_under_its_own_key(): - data = ReflectionData(verbose=0).load_mtz("tests/files/mtz/4BX9.mtz") +def test_intensity_path_keeps_french_wilsons_guard_under_its_own_key(mtz_dir): + data = ReflectionData(verbose=0).load_mtz(str(mtz_dir / "4BX9.mtz")) assert data.I is not None, "4BX9 should load via the intensity path" assert ReflectionData.FRENCH_WILSON_MASK_KEY in data.masks @@ -164,11 +164,11 @@ def test_intensity_path_keeps_french_wilsons_guard_under_its_own_key(): "name", ["1DAW", "2DQ6", "3A5V", "3E98", "3GR5", "3K7M", "3VRJ", "4BX9", "5BOV", "6G9X"], ) -def test_deposited_structures_lose_almost_nothing(name, pdb_dir): +def test_deposited_structures_lose_almost_nothing(name, mtz_dir): """Deposited data has already been through processing and merging; a criterion that rejects percent-level populations of it is mis-calibrated, not perceptive.""" - data = ReflectionData(verbose=0).load_mtz(f"tests/files/mtz/{name}.mtz") + data = ReflectionData(verbose=0).load_mtz(str(mtz_dir / f"{name}.mtz")) measured = int(data.masks["sanity_F"].sum()) rejected = int((~data.masks[ReflectionData.WILSON_MASK_KEY]).sum()) @@ -177,7 +177,7 @@ def test_deposited_structures_lose_almost_nothing(name, pdb_dir): @pytest.mark.unit -def test_flagged_reflections_show_no_directional_bias(): +def test_flagged_reflections_show_no_directional_bias(mtz_dir): """The regression that catches a lost anisotropy correction. 1DAW diffracts about four times more strongly along ``h*`` than ``l*``. A @@ -185,7 +185,7 @@ def test_flagged_reflections_show_no_directional_bias(): the strong one, so the flagged set piles up along ``h*`` -- 52 of 56 with mean ``|h|`` nearly twice the dataset's, before the correction existed. """ - data = ReflectionData(verbose=0).load_mtz("tests/files/mtz/1DAW.mtz") + data = ReflectionData(verbose=0).load_mtz(str(mtz_dir / "1DAW.mtz")) _, flagged = _plant_zingers(data, n=400, seed=13) assert len(flagged) > 100, "the planted population must be found first" From 4c32466d77eed89f456b196f12a3725feb333203 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sat, 29 Aug 2026 15:50:30 +0200 Subject: [PATCH 100/250] Take the collection R-factor percentiles at the configured float dtype _percentiles built its tensors as float64 regardless of the configured dtype. It is CPU-only and feeds a report, so nothing could crash on it, but the float dtype is a configuration and hardcoding around it is how a float64 path reaches a backend that has none. The float64 that remains in the package is deliberate: the dtype registry itself, the MPS guards and dtype checks, two explicit .cpu().double() calls, the CUDA-only planarity eigh (whose forward asserts is_cuda, and where the fp64 is load-bearing for near-collinear atoms), and scalar constants that become Python floats. Co-Authored-By: Claude Opus 5 (1M context) --- torchref/refinement/targets/collection/base.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/torchref/refinement/targets/collection/base.py b/torchref/refinement/targets/collection/base.py index 064cd971..daae2bfb 100644 --- a/torchref/refinement/targets/collection/base.py +++ b/torchref/refinement/targets/collection/base.py @@ -18,6 +18,7 @@ import torch from torchref.base.metrics.rfactor import rfactor_work_free +from torchref.config import get_float_dtype from torchref.refinement.targets.base import Target from torchref.utils.stats import ( VERBOSITY_DEBUG, @@ -172,8 +173,9 @@ def _percentiles(values: List[float]) -> Dict[str, float]: """10/25/50/75/90 percentiles of a list of per-dataset R-factors.""" if not values: return {} - t = torch.tensor(values, dtype=torch.float64) - q = torch.quantile(t, torch.tensor(_R_PERCENTILES, dtype=torch.float64)) + dtype = get_float_dtype() + t = torch.tensor(values, dtype=dtype) + q = torch.quantile(t, torch.tensor(_R_PERCENTILES, dtype=dtype)) return {lbl: q[i].item() for i, lbl in enumerate(_R_PCT_LABELS)} # ------------------------------------------------------------------ From 42b9979f4c8802cf206cf33be9f1cbb3d241d932 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sat, 29 Aug 2026 21:13:24 +0200 Subject: [PATCH 101/250] Add an absolute Wilson normaliser, and share the resolution basis Scaling and weighting are two different questions and the alignment package answered them with one object. Scaling asks what we compare; in a correlation it is gauge, absorbed exactly by the per-shell variance reweight downstream, which is why twelve E conventions moved truth rank by nothing across 100 paired cells. Weighting asks how much each reflection counts, and is never gauge. The previous convention returned both -- `.E` and `.weight` -- so sweeping it moved a gauge quantity and a real one together and read the sum. This is the scaling half. `WilsonNormaliser` fits `Sigma(s)` on one dataset and divides it out; no second dataset in the objective, no sigI, no model error, no solvent. Distinct from `ScalerBase`, which is a *relative* scaler putting F_calc onto F_obs and whose every target compares the two. ` = 1` is an identity of the fit rather than a normalisation step. The Gamma GLM's constant column has score equation `sum_h k_h (I_h/mu_h - 1) = 0`, which IS unit mean, shape-weighted. Nothing is rescaled afterwards and nothing can drift -- so a downstream `E^2 - 1` becomes a true centring. It is currently centring against a measured 0.954. Intensities, not amplitudes, and never inferred: Wilson statistics are exact on I and awkward on F, measurement error is near-Gaussian on I and badly behaved on F for weak data, and negative measurements are meaningful and survive. Three things the implementation had to learn from real data rather than from the synthetic case, all of them convergence rather than estimation: Coefficients cannot be the convergence test. Once an explicit range is passed so two fits share a basis, the data occupy only part of it, the high-order columns are near-collinear there, and the coefficients wander in the flat directions long after the fit has settled. Neither can the deviance. Its `-log(y/mu)` term diverges as y -> 0, and calculated amplitudes have near-zeros at the nodes of the molecular transform. The fit itself is untroubled -- the score contribution from y -> 0 is a bounded -k x -- which is precisely why the convergence test must not be the one thing that notices them. Nor the fitted mean, which is floored, so a collapsed fit compares equal to itself and reads as converged. An earlier version reported success while returning zeros. So both use the objective, `sum_h k_h (y/mu + log mu)`, whose `log y` term is constant in beta and simply absent -- with step halving, because IRLS on a log link can overshoot into underflow after which `y/mu` explodes. The shared `chebyshev_design` is bit-identical to the inline version it replaces across dtypes and orders, and keeps the prefix-nesting a scaling test slices on. Its new explicit range is for extrapolation, not comparability: an affine remap does not change what a polynomial basis spans, so two fits agree where their data overlap -- but the basis saturates at the ends, so beyond the fitted range a curve is frozen flat rather than extrapolated. Measured on 1DAW, 2DQ6, 3K7M and 4BX9: both sides converge, the identity holds to 1e-9, and `Sigma_obs/Sigma_calc` recovers the bulk-solvent deficit per structure -- 0.00-0.37 at the lowest shell against Babinet's 0.05 limit, with none of Babinet's constants, and differing between structures in a way one universal curve cannot represent. Gate: 393 passed / 26 skipped over tests/unit/scaling and tests/unit/refinement, the blast radius of the extraction, plus 15 new tests. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/scaling_gate.sh | 18 + alignment_lab/analysis/wilson_smoke.py | 92 +++++ alignment_lab/analysis/wilson_smoke.sh | 17 + docs/changelog.rst | 2 + tests/unit/scaling/test_wilson_normaliser.py | 201 ++++++++++ torchref/scaling/__init__.py | 2 + torchref/scaling/basis.py | 74 ++++ torchref/scaling/scaler_base.py | 22 +- torchref/scaling/wilson.py | 370 +++++++++++++++++++ 9 files changed, 786 insertions(+), 12 deletions(-) create mode 100644 alignment_lab/analysis/scaling_gate.sh create mode 100644 alignment_lab/analysis/wilson_smoke.py create mode 100644 alignment_lab/analysis/wilson_smoke.sh create mode 100644 tests/unit/scaling/test_wilson_normaliser.py create mode 100644 torchref/scaling/basis.py create mode 100644 torchref/scaling/wilson.py diff --git a/alignment_lab/analysis/scaling_gate.sh b/alignment_lab/analysis/scaling_gate.sh new file mode 100644 index 00000000..25f74e9e --- /dev/null +++ b/alignment_lab/analysis/scaling_gate.sh @@ -0,0 +1,18 @@ +#!/bin/bash +# Deliverable 1 gate: the Chebyshev extraction must be inert. +#SBATCH --job-name=scalegate +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +export TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 +"$PY" -m pytest -c tests/pytest.ini tests/unit/scaling tests/unit/refinement --run-slow -q 2>&1 | tail -8 +echo "PYTEST_RC=${PIPESTATUS[0]}" diff --git a/alignment_lab/analysis/wilson_smoke.py b/alignment_lab/analysis/wilson_smoke.py new file mode 100644 index 00000000..dfb531cb --- /dev/null +++ b/alignment_lab/analysis/wilson_smoke.py @@ -0,0 +1,92 @@ +"""Does the Wilson normaliser hold its identity on real data, both sides? + +The synthetic test draws from the distribution the fit assumes, so it can only +show the arithmetic is right. These are the cases the assumption is wrong in: +observations carry measurement error and a real solvent deficit, and the +rotation function's calc side is an oversampled molecular transform in a P1 box, +where adjacent samples are correlated and Wilson independence does not hold at +all. The mean estimate survives that by quasi-likelihood -- a log-link Gamma GLM +is consistent for the mean under a misspecified variance function -- and this is +where that claim gets checked rather than asserted. + +Also checks the property the whole weighting design rests on: fitted with a +SHARED abscissa, the two curves are comparable, and their ratio is the +resolution-dependent model deficiency. +""" + +from __future__ import annotations + +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import load_case # noqa: E402 + +CASES = ("1DAW", "2DQ6", "3K7M", "4BX9") + + +def main() -> int: + from torchref.scaling.wilson import WilsonNormaliser + + for pdb in CASES: + model, data = load_case(pdb) + hkl = data.hkl + rec = data.cell.reciprocal_basis_matrix.to(torch.float64) + s = (hkl.to(torch.float64) @ rec).norm(dim=-1) + F = data.F.to(torch.float64).abs() + keep = torch.isfinite(F) & (F > 0) + has_I = getattr(data, "I", None) is not None + print(f"\n=== {pdb} {data.spacegroup.hm} n={int(keep.sum())} " + f"raw intensities available: {has_I} ===") + + # --- observed side ------------------------------------------------- + I_obs = (F * F)[keep] + obs = WilsonNormaliser.from_hkl( + I_obs, hkl[keep], data.spacegroup, data.cell, n_coeff=6, + s_lo=float(s.min()), s_hi=float(s.max()), + ) + cen = data.spacegroup.is_centric(hkl[keep].to(torch.long)).to(torch.bool) + k = torch.where(cen, 0.5, 1.0).to(torch.float64) + e2 = obs.E_squared.to(torch.float64) + print(f" obs {obs!r}") + print(f" k-weighted = {float((k*e2).sum()/k.sum()):.10f}") + _deciles(" obs decile ", s[keep], e2) + + # --- calculated side, on the SAME abscissa -------------------------- + F_calc = model.get_structure_factor(hkl[keep], recalc=True).abs() + I_calc = (F_calc.to(torch.float64) ** 2) + calc = WilsonNormaliser.from_hkl( + I_calc, hkl[keep], data.spacegroup, data.cell, n_coeff=6, + s_lo=float(s.min()), s_hi=float(s.max()), + ) + e2c = calc.E_squared.to(torch.float64) + print(f" calc {calc!r}") + print(f" k-weighted = {float((k*e2c).sum()/k.sum()):.10f}") + _deciles(" calc decile ", s[keep], e2c) + + # --- the ratio the weight will be built from ------------------------ + # Same basis, so the two curves are directly comparable. Reported as a + # shape: what the model under-explains, versus resolution. + grid = torch.linspace(float(s[keep].min()), float(s[keep].max()), 8) + r = (obs.evaluate(grid).to(torch.float64) + / calc.evaluate(grid).to(torch.float64)) + r = r / r.mean() + print(" Sigma_obs/Sigma_calc (normalised) vs d(A):") + print(" " + " ".join(f"{1/float(x):5.1f}:{float(v):5.2f}" + for x, v in zip(grid, r))) + return 0 + + +def _deciles(label, s, v): + order = torch.argsort(s) + d = [float(v[order[i::10]].mean()) for i in range(10)] + print(f"{label}: min {min(d):.4f} max {max(d):.4f} " + + " ".join(f"{x:.2f}" for x in d)) + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/alignment_lab/analysis/wilson_smoke.sh b/alignment_lab/analysis/wilson_smoke.sh new file mode 100644 index 00000000..76855935 --- /dev/null +++ b/alignment_lab/analysis/wilson_smoke.sh @@ -0,0 +1,17 @@ +#!/bin/bash +#SBATCH --job-name=wsmoke +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:30:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=32G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 +export CUDA_VISIBLE_DEVICES="" +"$PY" -u alignment_lab/analysis/wilson_smoke.py +echo "RC=$?" diff --git a/docs/changelog.rst b/docs/changelog.rst index 46274d17..69d99757 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -9,6 +9,8 @@ Unreleased - Fixed the overall-anisotropy fit, which regressed log intensities with no constant term and so absorbed the ``-gamma`` offset into the tensor - Fixed molecular-replacement rotation candidates being composed onto each other instead of onto the search model - Fixed assigning a ``SpaceGroup`` object to ``Model.spacegroup`` being a silent no-op that then made the correct name assignment raise +- Added ``torchref.scaling.WilsonNormaliser``: an absolute normaliser that fits ``Sigma(s)`` as a Gamma GLM with a log link and divides it out, so `` = 1`` holds as an identity of the fit rather than as a separate normalisation step +- Extracted the Chebyshev resolution basis into ``torchref.scaling.basis``, shared with the isotropic scale, and gave it an explicit range so a curve fitted on one reflection set can be evaluated on another - Moved epsilon onto ``SpaceGroup.epsilon(hkl, friedel=)``; the alignment package's own copy disagreed with it in trigonal and hexagonal groups and dropped the centring coset. The default keeps the Friedel-folded count sigma_A is calibrated against, and the molecular-replacement likelihood asks for the conventional one - The rotation function and the ML rescore take their E-value convention as a class, so the observed and calculated sides are normalised by one rule rather than by nine private converters - The ML rescore now receives the observation sigmas, which the rotation function computed the French-Wilson posterior from and then discarded diff --git a/tests/unit/scaling/test_wilson_normaliser.py b/tests/unit/scaling/test_wilson_normaliser.py new file mode 100644 index 00000000..6fbf051c --- /dev/null +++ b/tests/unit/scaling/test_wilson_normaliser.py @@ -0,0 +1,201 @@ +"""Invariants of the absolute Wilson normaliser. + +The load-bearing one is the first: `` = 1`` is the *stationarity condition* +of the Gamma GLM's constant term, not a normalisation applied afterwards. Every +consumer that centres an intensity as ``E^2 - 1`` depends on it being exact +rather than approximate -- the previous convention centred against a measured +`` = 0.954`` and nothing noticed. + +The rest pin the properties that make two fits comparable, which is what the +weighting design will be built on: a shared abscissa, epsilon reaching both +sides of the ratio, and invariance to the units the data arrive in. +""" + +import pytest +import torch + +from torchref.scaling.wilson import WilsonNormaliser +from torchref.symmetry import Cell, SpaceGroup + +pytestmark = pytest.mark.unit + + +def _wilson_data(n=20000, seed=0, centric_frac=0.1, eps_value=None): + """Intensities actually drawn from the distribution the fit assumes.""" + g = torch.Generator().manual_seed(seed) + s = torch.rand(n, generator=g, dtype=torch.float64) * 0.45 + 0.05 + # A real curve to recover: Wilson falloff plus a low-resolution deficit of + # the shape a missing bulk solvent produces. + sigma = (torch.exp(-120.0 * (s / 2) ** 2) * 3000.0 + * (1 - 0.9 * torch.exp(-300.0 * (s / 2) ** 2))) + centric = torch.zeros(n, dtype=torch.bool) + centric[: int(n * centric_frac)] = True + k = torch.where(centric, 0.5, 1.0).to(torch.float64) + eps = (torch.ones(n, dtype=torch.float64) if eps_value is None + else torch.full((n,), float(eps_value), dtype=torch.float64)) + I = torch._standard_gamma(k.clone()) / k * (eps * sigma) + return I, s, eps, centric, sigma + + +def _k_weighted_mean(v, centric): + k = torch.where(centric, 0.5, 1.0).to(torch.float64) + return float((k * v.to(torch.float64)).sum() / k.sum()) + + +def test_unit_mean_is_an_identity_of_the_fit(): + """The constant column's score equation IS `` = 1``. + + ``sum_h k_h (I_h/mu_h - 1) = 0`` at the optimum, so this should hold to the + convergence tolerance rather than to some fitting accuracy. A loose result + here means the fit stopped early, not that the estimate is noisy. + """ + I, s, eps, centric, _ = _wilson_data() + w = WilsonNormaliser(I, s, eps=eps, centric=centric, n_coeff=6) + assert _k_weighted_mean(w.E_squared, centric) == pytest.approx(1.0, abs=1e-7) + + +@pytest.mark.parametrize("n_coeff", [1, 2, 6, 12]) +def test_unit_mean_holds_at_every_order(n_coeff): + """It is the intercept that pins the mean, so the order must not matter.""" + I, s, eps, centric, _ = _wilson_data() + w = WilsonNormaliser(I, s, eps=eps, centric=centric, n_coeff=n_coeff) + assert _k_weighted_mean(w.E_squared, centric) == pytest.approx(1.0, abs=1e-6) + + +def test_a_uniform_epsilon_cancels_out_of_the_ratio(): + """``E^2 = (I/eps) / ``, so a constant eps must change nothing. + + This is the defect the previous convention carried: it divided the shell + mean by eps without dividing the intensity by it, leaving `` = `` + -- 1 on a primitive lattice and 2 on a centred one, so a normaliser whose + absolute scale depended on the space group. + """ + I, s, _, centric, _ = _wilson_data() + plain = WilsonNormaliser(I, s, centric=centric, n_coeff=6) + doubled = WilsonNormaliser( + I, s, eps=torch.full_like(s, 2.0), centric=centric, n_coeff=6, + ) + # eps=2 halves the intensity going in AND halves Sigma, so E is unchanged. + assert torch.allclose(doubled.E, plain.E, rtol=1e-8, atol=1e-10) + + +def test_invariant_to_the_units_the_data_arrive_in(): + I, s, eps, centric, _ = _wilson_data() + base = WilsonNormaliser(I, s, eps=eps, centric=centric, n_coeff=6).E + for c in (1e-6, 1e6): + scaled = WilsonNormaliser( + I * c, s, eps=eps, centric=centric, n_coeff=6, + ).E + assert torch.allclose(scaled, base, rtol=1e-6, atol=1e-9), ( + f"scaling I by {c:g} moved E by " + f"{float((scaled - base).abs().max()):.3e}" + ) + + +def test_the_resolution_trend_is_removed(): + I, s, eps, centric, _ = _wilson_data() + w = WilsonNormaliser(I, s, eps=eps, centric=centric, n_coeff=6) + order = torch.argsort(s) + means = [float(w.E_squared[order[i::10]].to(torch.float64).mean()) + for i in range(10)] + assert max(means) / min(means) < 1.15, f"residual trend: {means}" + + +def test_the_fitted_curve_recovers_the_true_one(): + I, s, eps, centric, sigma_true = _wilson_data(n=60000) + w = WilsonNormaliser(I, s, eps=eps, centric=centric, n_coeff=6) + rel = (w.sigma_wilson.to(torch.float64) / sigma_true - 1).abs() + assert float(rel.median()) < 0.05 + assert float(rel.max()) < 0.30 + + +def test_one_coefficient_is_a_single_global_scale(): + I, s, eps, centric, _ = _wilson_data() + w = WilsonNormaliser(I, s, eps=eps, centric=centric, n_coeff=1) + sig = w.sigma_wilson.to(torch.float64) + assert float((sig / sig[0] - 1).abs().max()) < 1e-9 + + +def test_the_range_only_matters_outside_the_fitted_data(): + """An affine remap does not change what a polynomial basis spans. + + So two fits over different ranges recover the same *function* where their + data overlap -- the coefficients differ, the curve does not. What the range + controls is the other side of it: ``u`` saturates at the ends, so beyond the + fitted data the curve is frozen flat rather than extrapolated. That is the + whole reason to pass an explicit range, and it is why a fit made on one + reflection set can be used on another only if the range covers both. + """ + I, s, eps, centric, _ = _wilson_data(n=30000) + lo, hi = float(s.min()), float(s.max()) + sub = s < 0.3 + kw = dict(eps=eps[sub], centric=centric[sub], n_coeff=6) + shared = WilsonNormaliser(I[sub], s[sub], s_lo=lo, s_hi=hi, **kw) + own = WilsonNormaliser(I[sub], s[sub], **kw) + + inside = torch.linspace(0.06, 0.29, 40, dtype=torch.float64) + assert torch.allclose(shared.evaluate(inside), own.evaluate(inside), + rtol=1e-3), "the fitted function must not depend on " \ + "how the basis was parameterised" + + # Outside its own data, the narrow fit is pinned at its endpoint; the one + # given the full range keeps varying because it is still inside its basis. + outside = torch.linspace(0.32, 0.49, 20, dtype=torch.float64) + own_out = own.evaluate(outside).to(torch.float64) + assert float((own_out / own_out[0] - 1).abs().max()) < 1e-9, \ + "beyond the fitted range the curve should be flat, not extrapolated" + shared_out = shared.evaluate(outside).to(torch.float64) + assert float((shared_out / shared_out[0] - 1).abs().max()) > 1e-3 + + +def test_evaluate_reproduces_the_fitted_curve(): + I, s, eps, centric, _ = _wilson_data() + w = WilsonNormaliser(I, s, eps=eps, centric=centric, n_coeff=6) + assert torch.allclose(w.evaluate(s), w.sigma_wilson, rtol=1e-10, atol=1e-12) + + +def test_negative_intensities_are_kept_but_do_not_inform_the_fit(): + """Negative measurements are meaningful and unbiased; they are not errors. + + The Gamma likelihood has no support there, so they are held out of the + estimate -- but they still get a Sigma and a signed ``E_squared``, because + excluding them from the fit is not the same as refusing to normalise them. + """ + I, s, eps, centric, _ = _wilson_data() + I = I.clone() + I[:500] = -torch.rand(500, dtype=torch.float64) * 10.0 + w = WilsonNormaliser(I, s, eps=eps, centric=centric, n_coeff=6) + assert w.n_fitted == I.numel() - 500 + assert bool((w.E_squared[:500] < 0).all()), "sign must survive" + assert bool(torch.isfinite(w.sigma_wilson).all()) + assert bool((w.E[:500] == 0).all()), "E clamps, E_squared does not" + + +def test_it_raises_rather_than_quietly_degrading(): + """No fallback. A normaliser that becomes a different normaliser on the + hard cases is two normalisers wearing one name.""" + I, s, eps, centric, _ = _wilson_data(n=20) + with pytest.raises(ValueError, match="usable reflections"): + WilsonNormaliser(I[:3], s[:3], eps=eps[:3], centric=centric[:3], + n_coeff=6) + + +def test_from_hkl_excludes_systematic_absences(): + """Absences are zero by symmetry, not by measurement, so they say nothing + about Sigma -- and a Gamma fit told otherwise is dragged toward zero.""" + sg = SpaceGroup("P 43 21 2") + cell = Cell([70.0, 70.0, 90.0, 90.0, 90.0, 90.0]) + g = torch.Generator().manual_seed(5) + hkl = torch.randint(-14, 15, (12000, 3), generator=g) + hkl = hkl[hkl.abs().sum(dim=-1) > 0] + absent = sg.is_absent(hkl).to(torch.bool) + assert int(absent.sum()) > 0, "test needs a group with real absences" + + I = torch.rand(hkl.shape[0], generator=g, dtype=torch.float64) * 100 + 1 + I = torch.where(absent, torch.zeros_like(I), I) # absences really are 0 + + w = WilsonNormaliser.from_hkl(I, hkl, sg, cell, n_coeff=6) + assert w.n_fitted == int((~absent).sum()) + assert bool(torch.isfinite(w.sigma_wilson).all()) + # The zeros must not have dragged the curve down. + assert float(w.sigma_wilson.min()) > 0.0 diff --git a/torchref/scaling/__init__.py b/torchref/scaling/__init__.py index 27d4f8a2..8125a15a 100644 --- a/torchref/scaling/__init__.py +++ b/torchref/scaling/__init__.py @@ -12,10 +12,12 @@ from torchref.scaling.scaler_base import ScalerBase from torchref.scaling.solvent import SolventModel from torchref.scaling.collection_scaler import CollectionScaler +from torchref.scaling.wilson import WilsonNormaliser __all__ = [ "Scaler", "ScalerBase", "SolventModel", "CollectionScaler", + "WilsonNormaliser", ] diff --git a/torchref/scaling/basis.py b/torchref/scaling/basis.py new file mode 100644 index 00000000..50f39f77 --- /dev/null +++ b/torchref/scaling/basis.py @@ -0,0 +1,74 @@ +"""Chebyshev basis in resolution, shared by everything that fits a smooth curve in |s|. + +Two things in here carry argument rather than convention, and both were settled +by the scaler rework: + +* **The abscissa is ``sin(theta)/lambda``, not ``s**2``.** The modulation a + resolution-dependent scale has to represent is gentle through the bulk of the + range and has real structure in the first few percent of ``s**2``; a basis + uniform in ``s**2`` spends nearly all its resolution where nothing happens. +* **The basis is prefix-nested.** ``chebyshev_design(x, k)`` equals + ``chebyshev_design(x, n)[:, :k]`` for ``k <= n``, so raising the order adds + detail without redefining the terms already fitted, and a caller can slice + instead of rebuilding. + +``lo``/``hi`` exist for **extrapolation**, and the reason is the clamp rather +than the mapping. An affine remap does not change the space a polynomial basis +spans, so two fits over different ranges recover the same *function* where their +data overlap -- only the coefficients and the conditioning differ. What does +differ is outside the fitted range: ``u`` saturates at the ends, so every column +goes constant and the curve is frozen at its endpoint value. + +So a fit over one resolution range, evaluated somewhere else, silently returns a +flat extrapolation. That is the case a shared ``lo``/``hi`` is for -- fitting on +one reflection set and using the curve on another, which is what comparing two +fits, or fitting on a crystal lattice and evaluating on a dense sampling, +actually requires. ``ScalerBase`` does not need it: it builds one design over +all reflections and slices rows. +""" + +from __future__ import annotations + +from typing import Optional, Union + +import torch + +__all__ = ["chebyshev_design"] + + +def chebyshev_design( + x: torch.Tensor, + n_coeff: int, + lo: Optional[Union[float, torch.Tensor]] = None, + hi: Optional[Union[float, torch.Tensor]] = None, +) -> torch.Tensor: + """``(N, n_coeff)`` Chebyshev design matrix in ``x``. + + Parameters + ---------- + x : torch.Tensor + ``(N,)`` abscissa, normally ``sin(theta)/lambda``. + n_coeff : int + Number of Chebyshev terms. ``1`` gives a single constant column, i.e. a + global scale with no resolution dependence. + lo, hi : float or torch.Tensor, optional + Range to map onto ``[-1, 1]``. Both default to ``x``'s own extremes, + which is right for a single dataset and wrong the moment two fits have + to be compared -- see the module docstring. + + Returns + ------- + torch.Tensor + ``(N, n_coeff)``, column 0 all ones, every entry in ``[-1, 1]``. + """ + if n_coeff < 1: + raise ValueError(f"n_coeff must be at least 1, got {n_coeff}") + lo = x.min() if lo is None else torch.as_tensor(lo, dtype=x.dtype, device=x.device) + hi = x.max() if hi is None else torch.as_tensor(hi, dtype=x.dtype, device=x.device) + u = (2 * (x - lo) / (hi - lo).clamp(min=1e-12) - 1).clamp(-1.0, 1.0) + cols = [torch.ones_like(u), u] + for _ in range(2, n_coeff): + cols.append(2 * u * cols[-1] - cols[-2]) # Chebyshev recurrence + # The slice is what makes ``n_coeff == 1`` work: the loop does not run and + # the pre-seeded linear column is dropped. + return torch.stack(cols[:n_coeff], dim=1) diff --git a/torchref/scaling/scaler_base.py b/torchref/scaling/scaler_base.py index 9fd9119f..dbf81177 100644 --- a/torchref/scaling/scaler_base.py +++ b/torchref/scaling/scaler_base.py @@ -12,6 +12,7 @@ import torch.nn as nn from torchref.base.math_torch import U_to_matrix +from torchref.scaling.basis import chebyshev_design from torchref.base.metrics import ( binwise_scale, nll_xray, @@ -146,19 +147,16 @@ def __init__( def _build_iso_design(self) -> torch.Tensor: """``(N, n_iso_coeff)`` Chebyshev design matrix for the isotropic scale. - The abscissa is ``sqrt(s_half_sq)``, i.e. ``sin(theta)/lambda``, mapped onto - ``[-1, 1]``. That coordinate rather than ``s**2`` because the modulation is gentle - through the bulk of the resolution range but has real structure in the first few - percent of ``s**2``; a basis uniform in ``s**2`` spends nearly all its resolution - where nothing happens. + The abscissa is ``sqrt(s_half_sq)``, i.e. ``sin(theta)/lambda``; see + :func:`torchref.scaling.basis.chebyshev_design` for why that coordinate. + + No explicit range: this design is built once over all reflections and + then *sliced* wherever a subset is needed (``forward`` does exactly + that), so the mapping is the same everywhere it is used. """ - x = torch.sqrt(self._s_half_sq.clamp(min=0)) - lo, hi = x.min(), x.max() - u = (2 * (x - lo) / (hi - lo).clamp(min=1e-12) - 1).clamp(-1.0, 1.0) - cols = [torch.ones_like(u), u] - for _ in range(2, self.n_iso_coeff): - cols.append(2 * u * cols[-1] - cols[-2]) # Chebyshev recurrence - return torch.stack(cols[: self.n_iso_coeff], dim=1) + return chebyshev_design( + torch.sqrt(self._s_half_sq.clamp(min=0)), self.n_iso_coeff, + ) def iso_log_scale(self, design: Optional[torch.Tensor] = None) -> torch.Tensor: """Per-reflection isotropic log scale ``design @ c_iso``, clamped to ``[-10, 10]``. diff --git a/torchref/scaling/wilson.py b/torchref/scaling/wilson.py new file mode 100644 index 00000000..f74e673f --- /dev/null +++ b/torchref/scaling/wilson.py @@ -0,0 +1,370 @@ +"""Absolute Wilson normalisation: fit ``Sigma(s)`` and divide it out. + +Distinct from :class:`~torchref.scaling.scaler_base.ScalerBase`, which is a +*relative* scaler -- it puts ``F_calc`` onto ``F_obs`` and every target it can +minimise compares the two. This one takes a single dataset and answers "what is +the expected intensity at this resolution", so that dividing by it leaves +`` = 1``. One dataset in, one curve out, no second dataset anywhere in the +objective. + +**Why this exists as one shared class.** The repo grew at least five private +answers to the same question -- ``base/wilson_outliers.robust_mean_intensity``, +``base/french_wilson.estimate_mean_intensity_by_resolution``, +``ReflectionData._calculate_wilson_b``, the ``Sigma_N`` estimator in +``refinement/model_error_estimation/sigma_a``, and a per-shell one inside the +alignment package -- differing in whether they use means or medians, whether +they divide out ``epsilon``, whether they separate centrics, and where they put +their shell edges. Consumers that disagree about what E means cannot be compared +with each other, which is exactly what went wrong between the rotation function +and its own rescore. + +**Scaling, not weighting.** This class answers *what* we compare. It says +nothing about how much any reflection should count -- no ``sigI``, no model +error, no solvent. Those belong to a weight, and mixing them in here is what +made the previous convention object impossible to reason about: it returned a +normalisation and a weight together, so sweeping it moved a gauge quantity and a +real one at the same time. +""" + +from __future__ import annotations + +from typing import Optional, Tuple + +import torch + +from torchref.scaling.basis import chebyshev_design + +__all__ = ["WilsonNormaliser"] + +#: Chebyshev terms. Enough to follow a Wilson plot's curvature and the +#: low-resolution solvent deficit without chasing shell-to-shell noise. +#: Provisional: the order has never been chosen against a metric sensitive to +#: it, so screen on this class's own residual trend rather than on anything +#: downstream. +DEFAULT_N_COEFF = 6 + +#: Bound on ``log Sigma`` relative to its own constant term. A polynomial is +#: unbounded at the ends of its range, so without this a single extreme +#: reflection at the resolution limit can carry an arbitrary scale -- the same +#: reason ``ScalerBase.iso_log_scale`` clamps per reflection. +LOG_CLAMP = 10.0 + +#: Step halvings allowed per IRLS iteration before the step is abandoned. +MAX_HALVINGS = 30 + + +class WilsonNormaliser: + """``Sigma(s)`` by Gamma GLM, so that `` = 1`` by construction. + + The model is `` = eps_h * Sigma(s_h)`` with ``log Sigma`` a Chebyshev + polynomial in ``sin(theta)/lambda``, fitted by maximum likelihood under + + acentric I ~ Exp(Sigma) (Gamma, shape 1) + centric I ~ Sigma * chi^2_1 (Gamma, shape 1/2) + + i.e. a Gamma GLM with a log link and the shape as the prior weight. + + **Unit mean is an identity of the fit, not a normalisation step.** The + constant basis column's score equation is ``sum_h k_h (I_h/mu_h - 1) = 0``, + which is exactly `` = 1`` in the shape-weighted sense. Nothing is + rescaled afterwards and nothing can drift -- which is what makes a + downstream ``E^2 - 1`` a true centring rather than an approximate one. + + Least squares on ``log I`` would be the obvious alternative and is wrong: + ``E[log Gamma]`` carries a digamma offset, and with a constant term present + it is absorbed into the curve's shape rather than into the level. That is + the defect the overall-anisotropy fit was carrying. + + Parameters + ---------- + I : torch.Tensor + ``(N,)`` intensities. **Intensities, not amplitudes** -- Wilson + statistics are exact on I and awkward on F, and measurement error is + near-Gaussian on I but badly behaved on F for weak reflections, which is + the whole reason the French-Wilson posterior exists. Negative values are + allowed and kept: they are meaningful, unbiased measurements. They are + excluded from the *fit* (the Gamma likelihood has no support there) but + still receive a ``Sigma`` and a signed ``E_squared``. + s_mag : torch.Tensor + ``(N,)`` scattering-vector magnitude ``|s| = 1/d``, in inverse Angstrom. + eps : torch.Tensor, optional + ``(N,)`` reflection multiplicity. Divides the intensity before the fit, + because axial reflections are systematically stronger. ``None`` means 1 + everywhere, which is correct for a molecular transform sampled in a P1 + box -- multiplicity is a property of crystal symmetry and there is none + there. + centric : torch.Tensor, optional + ``(N,)`` bool, setting the Gamma shape. ``None`` means all acentric. + n_coeff : int, optional + Chebyshev terms. ``1`` gives a single global scale. + s_lo, s_hi : float, optional + ``|s|`` range mapped onto the basis. Defaults to this dataset's own + extremes. **Pass both explicitly whenever the curve will be evaluated + outside the fitted data's range** -- comparing two fits over different + ranges, or fitting on a crystal lattice and evaluating on a dense + sampling. The basis saturates at the ends, so beyond the fitted range + the curve is frozen flat rather than extrapolated. + fit_mask : torch.Tensor, optional + ``(N,)`` bool selecting which reflections *inform* the fit. Everything + still receives a ``Sigma``, because the curve is smooth and evaluable + anywhere. Use it to hold out systematic absences -- see + :meth:`from_hkl`, which does exactly that. + + Attributes + ---------- + coefficients : torch.Tensor + ``(n_coeff,)`` fitted Chebyshev coefficients of ``log Sigma``. + sigma_wilson : torch.Tensor + ``(N,)`` fitted ``Sigma(s)``. Deliberately not called ``sigma``: this + package also carries ``sig_F``, a measurement error, and ``sigma_a``, a + correlation coefficient, and the three are not interchangeable. + mean_intensity : torch.Tensor + ``(N,)`` ``eps * Sigma(s)``, the expected intensity of each reflection. + E_squared : torch.Tensor + ``(N,)`` ``I / mean_intensity``. **Signed** -- negative observations stay + negative. + E : torch.Tensor + ``(N,)`` ``sqrt(max(E_squared, 0))``. + """ + + MAX_HALVINGS = MAX_HALVINGS + + def __init__( + self, + I: torch.Tensor, + s_mag: torch.Tensor, + *, + eps: Optional[torch.Tensor] = None, + centric: Optional[torch.Tensor] = None, + n_coeff: int = DEFAULT_N_COEFF, + s_lo: Optional[float] = None, + s_hi: Optional[float] = None, + fit_mask: Optional[torch.Tensor] = None, + max_iter: int = 100, + tol: float = 1e-10, + ) -> None: + if I.ndim != 1: + raise ValueError(f"I must be 1-D, got {tuple(I.shape)}") + if s_mag.shape != I.shape: + raise ValueError( + f"s_mag {tuple(s_mag.shape)} does not match I {tuple(I.shape)}" + ) + self.dtype = I.dtype + self.n_coeff = int(n_coeff) + self._I = I + self._s_mag = s_mag + self._eps = eps + self.s_lo = float(s_mag.min()) if s_lo is None else float(s_lo) + self.s_hi = float(s_mag.max()) if s_hi is None else float(s_hi) + + eps64 = ( + torch.ones_like(I, dtype=torch.float64) if eps is None + else eps.to(torch.float64).clamp(min=1.0) + ) + # Shape 1 acentric (exponential), 1/2 centric. Enters as the IRLS weight + # because for a Gamma with shape k the variance is mu^2/k, so the + # log-link working weight is k itself. + k = ( + torch.ones_like(I, dtype=torch.float64) if centric is None + else torch.where(centric.to(torch.bool), 0.5, 1.0).to(torch.float64) + ) + + I_reduced = I.to(torch.float64) / eps64 + # The Gamma likelihood has no support at or below zero. Absences and + # negative measurements are held out of the fit and given a Sigma from + # the curve like everything else -- excluding them from the *estimate* + # is not the same as refusing to normalise them. + usable = torch.isfinite(I_reduced) & torch.isfinite(s_mag) & (I_reduced > 0) + if fit_mask is not None: + usable = usable & fit_mask.to(torch.bool) + if int(usable.sum()) < self.n_coeff + 1: + raise ValueError( + f"only {int(usable.sum())} usable reflections for a " + f"{self.n_coeff}-coefficient fit; need at least {self.n_coeff + 1}" + ) + self.n_fitted = int(usable.sum()) + + design = chebyshev_design( + (s_mag * 0.5).to(torch.float64), self.n_coeff, + lo=self.s_lo * 0.5, hi=self.s_hi * 0.5, + ) + self.coefficients, self.n_iter = self._irls( + design[usable], I_reduced[usable], k[usable], max_iter, tol, + ) + + log_sigma = self._eval_log_sigma(design) + self.sigma_wilson = torch.exp(log_sigma).to(self.dtype) + self.mean_intensity = ( + torch.exp(log_sigma) * eps64 + ).clamp(min=1e-30).to(self.dtype) + self.E_squared = I / self.mean_intensity + self.E = self.E_squared.clamp(min=0.0).sqrt() + + # -- fitting ----------------------------------------------------------- + + def _eval_log_sigma(self, design: torch.Tensor) -> torch.Tensor: + c = self.coefficients + return (design @ c).clamp( + min=-LOG_CLAMP + float(c[0]), max=LOG_CLAMP + float(c[0]), + ) + + def _irls( + self, + X: torch.Tensor, + y: torch.Tensor, + w: torch.Tensor, + max_iter: int, + tol: float, + ) -> Tuple[torch.Tensor, int]: + """Gamma GLM with a log link, by iteratively reweighted least squares. + + IRLS rather than a generic optimiser: for this link and family the + working weight does not depend on ``mu``, so each step is one weighted + least-squares solve and there is no step size, no line search and no + absolute tolerance to fail against an unnormalised objective. + + Convergence and step control both use the objective itself, + ``L = sum_h k_h (y_h/mu_h + log mu_h)`` -- the negative log-likelihood + with the terms not involving ``beta`` dropped. + + That choice is forced by what the alternatives do on real data. The + *coefficients* are underdetermined whenever the data occupy part of the + basis range, which is the normal case once an explicit ``s_lo``/``s_hi`` + is passed, so they wander in the flat directions long after the fit has + settled. The *deviance* carries a ``-log(y/mu)`` term that diverges as + ``y -> 0``, and calculated amplitudes have near-zeros at the nodes of + the molecular transform, so a few tiny intensities dominate it. And the + *fitted mean* cannot be compared as a ratio because it is floored, so a + collapsed fit reads as a converged one -- which is exactly how an early + version of this reported success while returning zeros. + + ``L`` has none of those problems: the ``log y`` term that breaks the + deviance is constant in ``beta`` and simply absent here. + + Step halving is the other half. IRLS on a log link can overshoot into + ``mu`` underflow, after which the working response ``y/mu`` explodes and + the next step is worse. Rejecting any step that does not improve ``L`` + and halving it is the standard remedy and makes the fit robust to the + ill-conditioning a partial basis range creates. + """ + # Seed at the constant curve, which is the exact MLE when Sigma has no + # resolution dependence. Every later iteration only adds shape. + beta = torch.zeros(self.n_coeff, dtype=torch.float64, device=X.device) + beta[0] = torch.log(((w * y).sum() / w.sum()).clamp(min=1e-30)) + + def objective(b): + eta = (X @ b).clamp( + min=-LOG_CLAMP + float(b[0]), max=LOG_CLAMP + float(b[0]), + ) + mu = torch.exp(eta).clamp(min=1e-300) + return float((w * (y / mu + eta)).sum()), eta, mu + + L, eta, mu = objective(beta) + for it in range(1, max_iter + 1): + z = eta + (y - mu) / mu # working response + XtW = X.transpose(0, 1) * w.unsqueeze(0) + A = XtW @ X + # Ridge proportional to the matrix's own scale: the high-order + # Chebyshev columns go near-singular when the data cover only part + # of the basis range. + A = A + torch.eye(self.n_coeff, dtype=A.dtype, device=A.device) * ( + 1e-10 * float(torch.diagonal(A).abs().max().clamp(min=1e-30)) + ) + step = torch.linalg.solve(A, XtW @ z) - beta + if not torch.isfinite(step).all(): + raise RuntimeError( + f"Wilson fit diverged at iteration {it}: the IRLS solve " + f"returned non-finite coefficients." + ) + + # Halve until the step actually improves the objective. + accepted = False + for _ in range(self.MAX_HALVINGS): + L_try, eta_try, mu_try = objective(beta + step) + if L_try <= L: + beta = beta + step + accepted = True + break + step = step * 0.5 + if not accepted: + # No downhill direction left: already at the optimum. + return beta, it + + improvement = abs(L - L_try) / (abs(L) + 1e-30) + L, eta, mu = L_try, eta_try, mu_try + if improvement <= tol: + return beta, it + raise RuntimeError( + f"Wilson fit did not converge in {max_iter} IRLS iterations " + f"(objective still moving by {improvement:.2e} relative). Raising " + f"rather than falling back to a coarser estimate: a normaliser that " + f"silently becomes a different normaliser on hard cases is two " + f"normalisers wearing one name." + ) + + # -- evaluation elsewhere --------------------------------------------- + + def evaluate(self, s_mag: torch.Tensor) -> torch.Tensor: + """``Sigma(s)`` at arbitrary ``|s|``, on the basis this fit was built on. + + The curve is smooth, so it can be fitted on one reflection set and used + on another -- which is what makes a fit on the crystal lattice usable on + a dense sampling of the same transform. **Only inside ``[s_lo, s_hi]``**: + the basis saturates at the ends, so outside that range this returns the + endpoint value, flat, rather than an extrapolation. + """ + design = chebyshev_design( + (s_mag * 0.5).to(torch.float64), self.n_coeff, + lo=self.s_lo * 0.5, hi=self.s_hi * 0.5, + ) + return torch.exp(self._eval_log_sigma(design)).to(self.dtype) + + # -- construction from crystallography -------------------------------- + + @classmethod + def from_hkl( + cls, + I: torch.Tensor, + hkl: torch.Tensor, + spacegroup, + cell, + **kwargs, + ) -> "WilsonNormaliser": + """Build from Miller indices, deriving ``|s|``, ``eps`` and centricity. + + The core takes plain tensors because not every caller has crystal + reflections -- a molecular transform sampled in a P1 box has no ``hkl`` + at all, and there ``eps`` is 1 with nothing centric. This constructor is + for the case that does. + + ``epsilon(friedel=False)``: Wilson's `` = eps * Sigma`` counts the + operations mapping ``h -> h``, which add coherently and set the mean. + The Friedel-folded count changes the *distribution* instead, and that is + centricity -- which enters here as the Gamma shape, separately. The two + branches feed two different parameters of the same likelihood. + """ + hkl_l = hkl.to(torch.long) + # The cell may carry the configured default device while the reflections + # are somewhere else; the caller should not have to reconcile them. + rec = cell.reciprocal_basis_matrix.to(device=hkl_l.device, + dtype=torch.float64) + s_mag = (hkl_l.to(torch.float64) @ rec).norm(dim=-1).to(I.dtype) + eps = spacegroup.epsilon(hkl_l, friedel=False).to(torch.float64) + centric = spacegroup.is_centric(hkl_l).to(torch.bool) + # Systematically absent reflections are zero by symmetry, not by + # measurement, so they carry no information about Sigma and would drag + # the Gamma fit toward zero. + fit_mask = ~spacegroup.is_absent(hkl_l).to(torch.bool) + user_mask = kwargs.pop("fit_mask", None) + if user_mask is not None: + fit_mask = fit_mask & user_mask.to(torch.bool) + return cls( + I, s_mag, eps=eps, centric=centric, fit_mask=fit_mask, **kwargs, + ) + + def __repr__(self) -> str: # pragma: no cover - display + return ( + f"{type(self).__name__}(N={self._I.numel()}, " + f"n_coeff={self.n_coeff}, n_fitted={self.n_fitted}, " + f"iters={self.n_iter})" + ) From 8e72ae3ba3ef3bf086801f2e82200ceef76e53e2 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sat, 29 Aug 2026 21:38:43 +0200 Subject: [PATCH 102/250] Normalise the alignment path through the shared Wilson fit `SmoothSigmaE` becomes a thin adapter: the fit moves to `torchref.scaling` and what stays behind is the `EConvention` protocol its consumers and conformance harness are written against. The class it wraps has no dependency on `experimental` at all -- torch, the shared Chebyshev basis, and nothing else. Default for all six consumers. Obs and calc go through the same call with the same loss, so common footing stops being a property the harness checks for and becomes one there is no way to violate: measured obs/calc ratio 0.996-1.001 across five structures, against 1.049-1.101 for the French-Wilson default. Conformance, over 1DAW / 2DQ6 / 3K7M / 3A5V / 6G9X: `` 0.998-1.003, residual resolution trend +0.00 to -0.10 -- the flattest of every convention in the table, per-shell included -- and KS against Wilson no worse than per-shell. Ranking is unchanged, which is the expected result rather than a disappointment: per-shell scaling is absorbed by the variance reweight downstream, so this axis is gauge for a correlation. 10 structures x 10 seeds paired: 100/100 found either way, rank-0 20 -> 21, 4 better and 5 worse, sign test p = 1.0. The rotation function got about twice as fast. Median 0.89 -> 0.42 s, -53%, ranging -12% on 3VRJ to -69% on 4BX9. The mechanism is that the French-Wilson posterior -- parabolic cylinder functions in numpy plus a Halley iteration for the D factor -- is no longer on the default path. Not a paired-node measurement, so treat the exact percentage as indicative; the direction and rough size are not in doubt at that magnitude. **The observed-side DFAC weighting went with it.** That is deliberate and it is the point of the split -- the scaler answers what we compare, and how much a reflection counts belongs to a weight -- but it does mean the rotation function currently applies no measurement-error weight at all. The weight object restores it properly. Measured cost of the gap: the same p = 1.0 above. One defect the conformance table caught: `SmoothSigmaE` was invariant to a global rescale of the data only to 1e-4, where every closed-form convention manages 1e-14. Under `I -> cI` the optimum is exactly `beta[0] -> beta[0] + log c`, but the objective picks up an additive `log c * sum(k)`, so a `|dL|/|L|` convergence threshold means something different at every scale. The difference itself is free of that term; dividing by `sum(k)` instead leaves a per-reflection criterion that is scale-free. Now 2-6e-8, with identical iteration counts across six orders of magnitude in the input. Gate: 1995 passed, 0 failed, with --run-slow in effect. Seam identity 0.000e+00. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- docs/changelog.rst | 2 + torchref/experimental/alignment/align.py | 4 +- torchref/experimental/alignment/e_values.py | 130 +++++------------- torchref/experimental/alignment/frf/api.py | 4 +- .../experimental/alignment/ml_rotation.py | 9 +- torchref/experimental/alignment/pipeline.py | 4 +- .../experimental/alignment/rotation_search.py | 12 +- torchref/scaling/wilson.py | 11 +- 8 files changed, 65 insertions(+), 111 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 69d99757..6de7d7ea 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -9,6 +9,8 @@ Unreleased - Fixed the overall-anisotropy fit, which regressed log intensities with no constant term and so absorbed the ``-gamma`` offset into the tensor - Fixed molecular-replacement rotation candidates being composed onto each other instead of onto the search model - Fixed assigning a ``SpaceGroup`` object to ``Model.spacegroup`` being a silent no-op that then made the correct name assignment raise +- The rotation function, ML rescore and translation search now normalise through the shared Wilson normaliser by default, replacing the French-Wilson posterior on the observed side. Rank-neutral over 10 structures x 10 seeds; the rotation function is about twice as fast, since the posterior and its D-factor iteration are no longer on the default path +- The observed-side ``DFAC`` weighting went with it. Weighting is a separate concern from scaling and is being rebuilt as its own object; until then the rotation function applies no measurement-error weight - Added ``torchref.scaling.WilsonNormaliser``: an absolute normaliser that fits ``Sigma(s)`` as a Gamma GLM with a log link and divides it out, so `` = 1`` holds as an identity of the fit rather than as a separate normalisation step - Extracted the Chebyshev resolution basis into ``torchref.scaling.basis``, shared with the isotropic scale, and gave it an explicit range so a curve fitted on one reflection set can be evaluated on another - Moved epsilon onto ``SpaceGroup.epsilon(hkl, friedel=)``; the alignment package's own copy disagreed with it in trigonal and hexagonal groups and dropped the centring coset. The default keeps the Friedel-folded count sigma_A is calibrated against, and the molecular-replacement likelihood asks for the conventional one diff --git a/torchref/experimental/alignment/align.py b/torchref/experimental/alignment/align.py index 2cca2784..5832f716 100644 --- a/torchref/experimental/alignment/align.py +++ b/torchref/experimental/alignment/align.py @@ -24,7 +24,7 @@ import torch from .lattman_love import LattmanLoveInterpolator -from .e_values import WilsonShellEpsE +from .e_values import SmoothSigmaE from .sh import ( apply_overall_anisotropy, assign_shells, @@ -306,7 +306,7 @@ def align_model_to_data( sigma_b: float = 0.0, model_error_A: Optional[float] = None, rescore_engine: str = "m_letf1", - rescore_e_convention: type = WilsonShellEpsE, + rescore_e_convention: type = SmoothSigmaE, subpeak_refine: bool = False, subpeak_refine_k: int = -1, subpeak_refine_step_deg: float = 1.5, diff --git a/torchref/experimental/alignment/e_values.py b/torchref/experimental/alignment/e_values.py index 9786a3cc..d04fdb0a 100644 --- a/torchref/experimental/alignment/e_values.py +++ b/torchref/experimental/alignment/e_values.py @@ -323,106 +323,48 @@ def _compute(self): class SmoothSigmaE(EConvention): - """``Sigma(s)`` as a smooth curve rather than a step function over shells. - - Per-shell ``Sigma`` is a noisy non-parametric estimate with edges, and the - edges are not free: two consumers binning the same ``|s|`` independently - disagreed about 7 of 55078 reflections on 3K7M. A smooth curve has no edges, - is the same function whichever subset it is evaluated on, and is what the - scaler already uses for the closely-related isotropic scale. - - Basis follows ``scaling/scaler_base.py::_build_iso_design`` -- Chebyshev in - ``sin(theta)/lambda`` mapped onto ``[-1, 1]``, evaluated per reflection, in - log space. That abscissa rather than ``s**2`` because the modulation is - gentle through the bulk of the range and has real structure in the first few - percent of ``s**2``. - - Fitted as a **Gamma GLM with a log link**, which is the right likelihood - rather than a convenience: acentric ``F**2`` is exponentially distributed - with mean ``Sigma``, i.e. Gamma with unit shape, and centric ``F**2`` is - Gamma with shape 1/2. Fitting the *mean* this way avoids the trap that a - regression on ``log F**2`` walks into -- the ``E[log chi**2]`` offset has to - go somewhere, and with no intercept it is absorbed into the shape of the - curve. That is precisely how the overall-anisotropy fit was biased. - - Coefficients are clamped in log space for the reason the scaler clamps: a - polynomial is unbounded at the ends of its interval, and the low-resolution - end is where a mis-specified normaliser does its damage. + """Adapter over the shared :class:`~torchref.scaling.WilsonNormaliser`. + + The fit itself does not live here, because "what is the mean intensity at + this resolution" is not an alignment question -- at least five private + answers to it grew across the repo, and consumers that disagree about it + cannot be compared with each other. What stays here is the ``EConvention`` + protocol that this package's consumers and its conformance harness are + written against. + + ``weight`` is ones, and that is the point of the split rather than an + omission: this class answers *what* we compare. How much each reflection + counts is a weight, built from ``sigI`` and model error, and belongs + elsewhere. Returning both from one object is what made the previous + conventions impossible to interpret -- sweeping one moved a gauge quantity + and a real one at the same time. + + Pass ``s_lo``/``s_hi`` whenever obs and calc are fitted separately and their + curves will be compared: the basis saturates at the ends, so a curve + evaluated beyond its own fitted range is frozen flat rather than + extrapolated. """ - #: Chebyshev terms. Six is the scaler's default and spans a Wilson plot's - #: curvature without chasing shell-to-shell noise. + #: Chebyshev terms. Provisional -- the order has never been screened against + #: a metric sensitive to it. DEFAULT_N_COEFF = 6 - #: Log-space clamp on the fitted curve, as a factor either side of the - #: global mean intensity. Wide enough never to bind on real data; present so - #: an extrapolating polynomial cannot produce an arbitrary scale. - LOG_CLAMP = 10.0 - def __init__(self, *args, n_coeff: int = DEFAULT_N_COEFF, - n_iter: int = 8, **kwargs) -> None: + s_lo=None, s_hi=None, **kwargs) -> None: self.n_coeff = int(n_coeff) - self.n_iter = int(n_iter) + self.s_lo = s_lo + self.s_hi = s_hi super().__init__(*args, **kwargs) - def _design(self) -> torch.Tensor: - """``(N, n_coeff)`` Chebyshev design in sin(theta)/lambda.""" - x = (self.s_mag * 0.5).clamp(min=0.0) - lo, hi = x.min(), x.max() - u = (2 * (x - lo) / (hi - lo).clamp(min=1e-12) - 1).clamp(-1.0, 1.0) - cols = [torch.ones_like(u), u] - for _ in range(2, self.n_coeff): - cols.append(2 * u * cols[-1] - cols[-2]) - return torch.stack(cols[: self.n_coeff], dim=1) - - def _fit_log_sigma(self) -> torch.Tensor: - """IRLS for a Gamma GLM with log link; returns log Sigma per reflection.""" - X = self._design().to(torch.float64) - y = self._intensity().to(torch.float64).clamp(min=1e-30) - # Gamma shape: 1 acentric (exponential), 1/2 centric. Used as the IRLS - # weight, so better-determined reflections pull harder. - w = torch.where(self.centric, 0.5, 1.0).to(torch.float64) - - # Seed at the global mean, i.e. the constant curve a single Wilson - # scale would give. Every later iteration only adds shape. - beta = torch.zeros(self.n_coeff, dtype=torch.float64, device=X.device) - beta[0] = torch.log(y.mean().clamp(min=1e-30)) - for _ in range(self.n_iter): - eta = (X @ beta).clamp(-self.LOG_CLAMP + float(beta[0]), - self.LOG_CLAMP + float(beta[0])) - mu = torch.exp(eta) - # Log link with Gamma variance: the working response is - # eta + (y - mu)/mu and the IRLS weight is constant in mu. - z = eta + (y - mu) / mu.clamp(min=1e-30) - XtW = X.transpose(0, 1) * w.unsqueeze(0) - A = XtW @ X - A = A + torch.eye( - self.n_coeff, dtype=A.dtype, device=A.device, - ) * 1e-10 * float(torch.diagonal(A).abs().max().clamp(min=1e-30)) - beta_new = torch.linalg.solve(A, XtW @ z) - if torch.allclose(beta_new, beta, rtol=1e-10, atol=1e-12): - beta = beta_new - break - beta = beta_new - return (X @ beta).clamp( - -self.LOG_CLAMP + float(beta[0]), self.LOG_CLAMP + float(beta[0]), - ) - def _compute(self): - shell_sigma = self.sigma # the per-shell fallback - log_sigma = self._fit_log_sigma() - sigma = torch.exp(log_sigma).to(self.F.dtype).clamp(min=1e-30) - # Sanity: a fitted Sigma(s) must reproduce the data's own mean intensity. - # A Gamma GLM on a Chebyshev basis can diverge when the calc amplitudes - # span a huge dynamic range with near-zeros, and it did -- two of four - # calc sets came back with ~ 0, i.e. Sigma inflated by orders of - # magnitude. Detect that against the quantity the fit is estimating and - # fall back to the per-shell estimate rather than returning nonsense. - mean_I = self._intensity().mean().clamp(min=1e-30) - ratio = float((sigma.mean() / mean_I).clamp(min=1e-30)) - self.converged = 0.2 < ratio < 5.0 - if not self.converged: - sigma = shell_sigma - self.sigma = sigma - E = (self._intensity() / self.sigma).clamp(min=0.0).sqrt() - return E, self._ones() + from torchref.scaling import WilsonNormaliser + + # `_intensity()` has already divided by eps -- `uses_epsilon` is True + # here -- so the normaliser gets a reduced intensity and no eps of its + # own. Applying it in both places would count multiplicity twice. + fit = WilsonNormaliser( + self._intensity(), self.s_mag, centric=self.centric, + n_coeff=self.n_coeff, s_lo=self.s_lo, s_hi=self.s_hi, + ) + self.sigma = fit.sigma_wilson + return fit.E, self._ones() diff --git a/torchref/experimental/alignment/frf/api.py b/torchref/experimental/alignment/frf/api.py index 6f4441c0..5e17712a 100644 --- a/torchref/experimental/alignment/frf/api.py +++ b/torchref/experimental/alignment/frf/api.py @@ -21,7 +21,7 @@ import torch -from ..e_values import (FrenchWilsonE, convention_for_calc, +from ..e_values import (SmoothSigmaE, convention_for_calc, convention_uses_sigma_f) from .data_mr import bessel_sh_expand, cross_correlate_xi from .peak_finder import find_rotation_peaks @@ -139,7 +139,7 @@ def __init__( grid_sampling_deg: float = 2.0, asu_idx: Optional[torch.Tensor] = None, s_mag_asu: Optional[torch.Tensor] = None, - e_convention: type = FrenchWilsonE, + e_convention: type = SmoothSigmaE, ): self.device = s_obs.device diff --git a/torchref/experimental/alignment/ml_rotation.py b/torchref/experimental/alignment/ml_rotation.py index 1609b2c2..da46ed0a 100644 --- a/torchref/experimental/alignment/ml_rotation.py +++ b/torchref/experimental/alignment/ml_rotation.py @@ -25,8 +25,9 @@ import torch -from .e_values import (CalcShellE, WilsonShellE, WilsonShellEpsE, - convention_for_calc, convention_uses_sigma_f) +from .e_values import (CalcShellE, SmoothSigmaE, WilsonShellE, + WilsonShellEpsE, convention_for_calc, + convention_uses_sigma_f) from .frf.rotation_utils import ( axis_angle_to_matrix, edmonds_euler_from_rotation_matrix, @@ -558,7 +559,7 @@ def _build_llg_context( apply_wilson_b: bool = False, wilson_b_value: Optional[float] = None, sig_F_obs: Optional[torch.Tensor] = None, - e_convention: type = WilsonShellEpsE, + e_convention: type = SmoothSigmaE, ) -> _LLGContext: """Build the rotation-independent m_LETF1 LLG context (DataMR.cc:1326-1429). @@ -1067,7 +1068,7 @@ def m_letf1_rescore( apply_wilson_b: bool = False, wilson_b_value: Optional[float] = None, # if None and apply_wilson_b=True, fitted from data sig_F_obs: Optional[torch.Tensor] = None, - e_convention: type = WilsonShellEpsE, + e_convention: type = SmoothSigmaE, ) -> List[RotationPeak]: """Phaser-faithful ``m_LETF1`` rescore (DataMR.cc:1326-1429). diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index b6de91a7..b1eb7e53 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -45,7 +45,7 @@ _external_rwork, _prepare_frf_inputs, ) -from .e_values import WilsonShellEpsE +from .e_values import SmoothSigmaE from .frf.rotation_utils import ( axis_angle_to_matrix, edmonds_euler_from_rotation_matrix, @@ -234,7 +234,7 @@ def __init__( model_error_A: Optional[float] = None, # --- rescore --- rescore_engine: str = "m_letf1", - rescore_e_convention: type = WilsonShellEpsE, + rescore_e_convention: type = SmoothSigmaE, auto_variance_weights: bool = True, use_interp_var: bool = False, subpeak_refine: bool = False, diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index b865c7a0..8a357752 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -28,7 +28,7 @@ import torch -from .e_values import FrenchWilsonE +from .e_values import SmoothSigmaE from .sh import ( apply_overall_anisotropy, assign_shells, @@ -221,7 +221,7 @@ def search_peaks( n_peaks: int, verbose: int = 0, device: Optional[torch.device] = None, - e_convention: type = FrenchWilsonE, + e_convention: type = SmoothSigmaE, ) -> Tuple[List["RotationPeak"], int, float]: """Run the rotation function, returning the engine's own peak list. @@ -417,7 +417,7 @@ def rotation_search( n_peaks: int = 500, verbose: int = 0, device: Optional[torch.device] = None, - e_convention: type = FrenchWilsonE, + e_convention: type = SmoothSigmaE, ) -> RotationSolutions: """Find the orientations of ``model`` consistent with ``data``. @@ -451,8 +451,10 @@ def rotation_search( instance: a fitted ``Sigma(s)`` cannot exist before the reflections do, so the engine constructs it -- once for the observations and once for the model, which is what puts the two on a common footing. The default - pairs the French-Wilson posterior on obs (it reads ``sigF``) with plain - per-shell Wilson on calc. ``functools.partial`` configures one. + fits a smooth ``Sigma(s)`` -- a Gamma GLM on a Chebyshev basis in + sin(theta)/lambda -- independently for each side, so `` = 1`` + holds on both as an identity of the fit. ``functools.partial`` + configures one. Returns ------- diff --git a/torchref/scaling/wilson.py b/torchref/scaling/wilson.py index f74e673f..00b66fd1 100644 --- a/torchref/scaling/wilson.py +++ b/torchref/scaling/wilson.py @@ -290,13 +290,20 @@ def objective(b): # No downhill direction left: already at the optimum. return beta, it - improvement = abs(L - L_try) / (abs(L) + 1e-30) + # Per-reflection, NOT relative to |L|. Under I -> cI the optimum is + # just beta[0] -> beta[0] + log c, so the fit is exactly scale + # invariant -- but L picks up an additive `log c * sum(k)`, which + # makes a |dL|/|L| threshold mean something different at every + # scale. The difference itself is free of that term, so dividing by + # sum(k) leaves a criterion that is not. + improvement = abs(L - L_try) / float(w.sum()) L, eta, mu = L_try, eta_try, mu_try if improvement <= tol: return beta, it raise RuntimeError( f"Wilson fit did not converge in {max_iter} IRLS iterations " - f"(objective still moving by {improvement:.2e} relative). Raising " + f"(objective still moving by {improvement:.2e} per reflection). " + f"Raising " f"rather than falling back to a coarser estimate: a normaliser that " f"silently becomes a different normaliser on hard cases is two " f"normalisers wearing one name." From cce9d46d38e8d11d30f9a40ed00713477e0260e4 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sun, 30 Aug 2026 11:47:47 +0200 Subject: [PATCH 103/250] Weight reflections by inverse variance, and stop weighting twice The other half of the scaling/weighting split. Scaling asks what we compare and is gauge for a correlation; this asks how much each reflection counts, which is not. It also restores the observed-side measurement weighting that left with the French-Wilson posterior when the scaler landed. `w = 1/(1/snr^2 + eps - sigma_A^2)`: measurement variance and the MLHL model variance in one denominator. The saturation usually bolted on as a sigmoid is already there -- once model error dominates, extra measurement precision buys nothing, because what is wrong is the model. **The two error sources do not factorise, and finding that out was the work.** The obvious design is a per-reflection `snr` term times a resolution `sigma_A` term, and it fails: taken alone, `sigma_A/(eps - sigma_A^2)` is dominated by its own singularity. On a realistic Luzzati falloff it runs 10.1 at the lowest resolution shell against 1.0 at the next, and comes out IDENTICAL for a 0.5 A and a 1.0 A coordinate error, because as `sigma_A -> 1` the shape is set entirely by `1/(1 - sigma_A^2)` and the model error drops out. A weight carrying no information about the thing it is named after is not worth having. What regularises it is the term the factorisation discarded: those low-resolution reflections are strong but not infinitely well measured. One denominator. Correcting one thing I claimed earlier: I said this form would tilt harder toward low resolution than the shipped `sigma_A^2` and risk the large structures. Measured over a realistic falloff, `sigma_A^2` tilts *more* -- I had been comparing against inverse variance without the `sigma_A` numerator a correlation detector requires. `sigma_A^2` stays on the calculated side, where it is the expected moving-model intensity -- part of the signal, not a weight. Weighting the model by its own reliability and then also weighting the data by it counts one thing twice. MEASURED, 8 arms x 100 paired cells, each one deviation from what the engine did before. Comprehensively null: 100/100 found everywhere, rank-0 20-21/100, median truth rank 2.0 and top-20 98/100 for every arm, every sign test p >= 0.24. Separating the questions the arms were built to separate: a measurement weight at all none -> information p = 0.63 coupling in the model error information -> invvar p = 0.74 the per-shell reweight invvar +/- shellvar p = 1.00 So the weighting axis is inert for this ranking too -- and unlike the scaling half, that one is not explained by gauge. Across roughly thirty distinct configurations now, rank-0 has not left 20-23 of 100 while truth is found 100 times out of 100. The limiter is not what we compare or how much it counts. `apply_shell_variance_weights` is off by default on the strength of the third line above. It is a per-shell weight, so the correlation absorbs it -- the same mechanism that made twelve E conventions measure identical -- and switching it on moves nothing. Kept as a flag rather than deleted, so the finding stays re-checkable. Gate: 2006 passed, 0 failed. Seam identity 0.000e+00. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/weight_arms.py | 99 +++++++++++ alignment_lab/analysis/weight_arms.sh | 25 +++ alignment_lab/lab/frf.py | 16 ++ docs/changelog.rst | 2 + tests/unit/scaling/test_weighting.py | 124 +++++++++++++ torchref/experimental/alignment/frf/api.py | 69 +++++++- .../experimental/alignment/rotation_search.py | 8 + torchref/scaling/weighting.py | 164 ++++++++++++++++++ 8 files changed, 499 insertions(+), 8 deletions(-) create mode 100644 alignment_lab/analysis/weight_arms.py create mode 100644 alignment_lab/analysis/weight_arms.sh create mode 100644 tests/unit/scaling/test_weighting.py create mode 100644 torchref/scaling/weighting.py diff --git a/alignment_lab/analysis/weight_arms.py b/alignment_lab/analysis/weight_arms.py new file mode 100644 index 00000000..1fabe5d0 --- /dev/null +++ b/alignment_lab/analysis/weight_arms.py @@ -0,0 +1,99 @@ +"""What is the observed-side weight worth, and which part of it? + +The scaler split scaling from weighting and showed the scaling half is gauge for +a correlation. This is the half that is not. + +Arms deviate one thing at a time from `none_shellvar`, which is what the engine +did before any of this. Moving several at once and reading one number is what +made the E-convention panel uninterpretable. + +Three things are being asked: + +* is a measurement-error weight worth anything at all (`none` vs `information`); +* is folding model error into the same denominator worth more than the + measurement term alone (`information` vs `inverse_var`) -- they do not + factorise, so this is the only way to separate them; +* is `apply_shell_variance_weights` doing anything, given it is a per-shell + weight and per-shell weights are absorbed (`*_shellvar` pairs). + +The caps are screened rather than assumed. `trust_cap` is supposed to be a +backstop that never binds, so if it moves the result, sigma_A is wrong. +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, FRFConfig, orbit_rank, rotated_case, # noqa: E402 + run_frf, seed_for) + +#: Each arm is one deviation from `control`, which is what ships today. +ARMS = { + # What shipped before any of this: unit observed weight (DFAC left with the + # French-Wilson posterior when the scaler landed) and the per-shell reweight + # on. The load-bearing control -- without it, "the weight helps" cannot be + # told apart from "any weight helps". + "none_shellvar": {"obs_weight": "none", "shell_variance_weights": True}, + "none": {"obs_weight": "none"}, + "information": {"obs_weight": "information"}, + "inverse_var": {"obs_weight": "inverse_variance"}, + "invvar_shellvar": {"obs_weight": "inverse_variance", + "shell_variance_weights": True}, + # Cap screens. The SNR cap sets where measurement error stops limiting; the + # trust cap is meant to be a backstop, so if these move the result it is + # binding and sigma_A is wrong. + "snr_cap_2": {"obs_weight": "information", "snr_cap": 2.0}, + "snr_cap_10": {"obs_weight": "information", "snr_cap": 10.0}, + "trust_cap_10": {"obs_weight": "inverse_variance", "trust_cap": 10.0}, +} + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) + ap.add_argument("--trials", type=int, default=10) + ap.add_argument("--lmax-cap", type=int, default=64) + ap.add_argument("--n-peaks", type=int, default=500) + ap.add_argument("--thr-deg", type=float, default=5.0) + ap.add_argument("--arms", default="") + args = ap.parse_args() + + names = [a for a in args.arms.split(",") if a] or list(ARMS) + for trial in range(args.trials): + seed = seed_for(args.pdb, trial) + model, data, R_true = rotated_case(args.pdb, seed) + sym = data.spacegroup.matrices.to(torch.float64).cpu() + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() + okw = dict(side="left", frame="cart", reciprocal_basis=rec, + thr_deg=args.thr_deg) + for name in names: + cfg = FRFConfig(n_peaks=args.n_peaks, lmax_cap=args.lmax_cap, + **ARMS[name]) + t0 = time.time() + try: + res = run_frf(model, data, cfg, capture_arf=False, verbose=0) + except Exception as exc: + print(f"ROW {name} {args.pdb} trial={trial} seed={seed} rank=-1 " + f"rank_cmp={args.n_peaks} found=0 seconds=0.00 " + f"error={type(exc).__name__}", flush=True) + continue + seconds = time.time() - t0 + rank, ang = orbit_rank(res.peaks, R_true, sym, **okw) + print(f"ROW {name} {args.pdb} trial={trial} seed={seed} " + f"rank={rank} rank_cmp={rank if rank >= 0 else args.n_peaks} " + f"found={int(rank >= 0)} top20={int(0 <= rank < 20)} " + f"angle={'' if ang is None else round(float(ang), 3)} " + f"seconds={seconds:.2f}", flush=True) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/alignment_lab/analysis/weight_arms.sh b/alignment_lab/analysis/weight_arms.sh new file mode 100644 index 00000000..d489f7e6 --- /dev/null +++ b/alignment_lab/analysis/weight_arms.sh @@ -0,0 +1,25 @@ +#!/bin/bash +# Which E convention ranks truth best? Nine arms x 5 trials x 10 structures. +# the same pass. Every previously published FRF number was measured on + + +#SBATCH --job-name=warms +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=day +#SBATCH --time=03:00:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-9 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +echo "host=$(hostname) pdb=$PDB" +"$PY" -u alignment_lab/analysis/weight_arms.py --pdb "$PDB" --trials 10 2>/dev/null | grep '^ROW ' + diff --git a/alignment_lab/lab/frf.py b/alignment_lab/lab/frf.py index ea90e953..5743878f 100644 --- a/alignment_lab/lab/frf.py +++ b/alignment_lab/lab/frf.py @@ -51,6 +51,15 @@ class FRFConfig: #: knobs this is a real production parameter, so the lab passes it through #: rather than patching a constant. e_convention: Optional[type] = None + #: Weighting, the other half of the split. ``None`` leaves the production + #: default. These are separate arms on purpose: the design changes three + #: things at once -- the observed-side weight, the calculated-side weight + #: and whether the per-shell reweight runs -- and a panel that moves all + #: three cannot say which one did anything. + obs_weight: Optional[str] = None + shell_variance_weights: Optional[bool] = None + snr_cap: Optional[float] = None + trust_cap: Optional[float] = None extra: Dict[str, Any] = field(default_factory=dict) def as_row(self) -> Dict[str, Any]: @@ -235,6 +244,13 @@ def _wrapped(self, *args, **kwargs): # production default from the signature instead of overriding it with one. conv_kw = {} if cfg.e_convention is None else { "e_convention": cfg.e_convention} + # Engine knobs are omitted when unset so the production default applies, + # rather than being passed as None and overriding it with nothing. + for _name in ("obs_weight", "shell_variance_weights", + "snr_cap", "trust_cap"): + _v = getattr(cfg, _name) + if _v is not None: + conv_kw[_name] = _v t0 = time.time() with patched(_rs, "LMAX_CAP", int(cfg.lmax_cap)), \ diff --git a/docs/changelog.rst b/docs/changelog.rst index 6de7d7ea..c80801b1 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -9,6 +9,8 @@ Unreleased - Fixed the overall-anisotropy fit, which regressed log intensities with no constant term and so absorbed the ``-gamma`` offset into the tensor - Fixed molecular-replacement rotation candidates being composed onto each other instead of onto the search model - Fixed assigning a ``SpaceGroup`` object to ``Model.spacegroup`` being a silent no-op that then made the correct name assignment raise +- Added ``torchref.scaling.weighting``: measurement and model error combined as one inverse-variance weight per reflection, restoring the observed-side weighting that left with the French-Wilson posterior +- ``apply_shell_variance_weights`` is off by default in the rotation function. It is a per-shell weight, and per-shell weights are absorbed by the correlation; switching it on moves nothing - The rotation function, ML rescore and translation search now normalise through the shared Wilson normaliser by default, replacing the French-Wilson posterior on the observed side. Rank-neutral over 10 structures x 10 seeds; the rotation function is about twice as fast, since the posterior and its D-factor iteration are no longer on the default path - The observed-side ``DFAC`` weighting went with it. Weighting is a separate concern from scaling and is being rebuilt as its own object; until then the rotation function applies no measurement-error weight - Added ``torchref.scaling.WilsonNormaliser``: an absolute normaliser that fits ``Sigma(s)`` as a Gamma GLM with a log link and divides it out, so `` = 1`` holds as an identity of the fit rather than as a separate normalisation step diff --git a/tests/unit/scaling/test_weighting.py b/tests/unit/scaling/test_weighting.py new file mode 100644 index 00000000..01f1fd01 --- /dev/null +++ b/tests/unit/scaling/test_weighting.py @@ -0,0 +1,124 @@ +"""The weighting half of the split: properties, not a preferred answer. + +Every assertion here is something that has to be true whatever the caps end up +being. The numbers themselves are screened, not asserted -- pinning a cap in a +unit test would make the screen unfalsifiable. +""" + +import pytest +import torch + +from torchref.scaling.weighting import ( + DEFAULT_SNR_CAP, information_weight, inverse_variance_weight, + normalise_weight, snr_from_amplitude, +) + +pytestmark = pytest.mark.unit + + +def test_information_weight_saturates_at_one(): + """Past the cap, better measurement buys nothing: the model is the limit.""" + snr = torch.tensor([0.0, 1.0, 5.0, 50.0, 1e4], dtype=torch.float64) + w = information_weight(snr, cap=5.0) + assert float(w[0]) == 0.0 + assert float(w[2]) == pytest.approx(0.5), "the cap is where it reaches a half" + assert float(w[-1]) == pytest.approx(1.0, abs=1e-6) + assert bool((w[1:] > w[:-1]).all()), "must be monotone in signal-to-noise" + + +def test_information_weight_is_the_variance_ratio_in_disguise(): + """``w = 1/(1 + sigma_meas^2/sigma_model^2)`` -- not a sigmoid picked by eye. + + With the cap standing for the signal-to-noise at which the two errors are + equal, the weight must equal that expression exactly. + """ + snr = torch.linspace(0.1, 40.0, 200, dtype=torch.float64) + cap = 5.0 + expected = 1.0 / (1.0 + (cap / snr) ** 2) + assert torch.allclose(information_weight(snr, cap=cap), expected) + + +def test_below_the_cap_the_weight_is_quadratic_in_snr(): + """Where measurement error dominates, information goes as snr^2.""" + snr = torch.tensor([0.01, 0.02, 0.04], dtype=torch.float64) + w = information_weight(snr, cap=DEFAULT_SNR_CAP) + assert float(w[1] / w[0]) == pytest.approx(4.0, rel=1e-3) + assert float(w[2] / w[1]) == pytest.approx(4.0, rel=1e-3) + + +def test_snr_uses_the_intensity_convention(): + """``I/sigma_I`` with ``I = F^2`` is ``F/(2 sigma_F)``, half the amplitude's. + + Only a factor of two, and only a rescaling of the cap -- but quoting a cap + against the wrong signal-to-noise silently doubles it. + """ + F = torch.tensor([10.0, 100.0], dtype=torch.float64) + sig = torch.tensor([1.0, 1.0], dtype=torch.float64) + assert torch.allclose(snr_from_amplitude(F, sig), + torch.tensor([5.0, 50.0], dtype=torch.float64)) + + +def test_the_weight_rises_as_either_error_falls(): + """Inverse variance: better data or a better model both mean more weight.""" + snr = torch.tensor([1.0, 1.0, 10.0], dtype=torch.float64) + sa = torch.tensor([0.1, 0.9, 0.1], dtype=torch.float64) + w = inverse_variance_weight(snr, sa, cap=1e9) + assert float(w[1]) > float(w[0]), "a more reliable model must weigh more" + assert float(w[2]) > float(w[0]), "a better measurement must weigh more" + + +def test_measurement_error_bounds_the_low_resolution_weight(): + """``sigma_A -> 1`` sends the model variance to zero; ``1/snr^2`` is what + stops the weight diverging, and it must, because that regime is where the + strongest reflections live and one of them could otherwise carry the run.""" + sa = torch.tensor([0.999999, 0.999999], dtype=torch.float64) + snr = torch.tensor([5.0, 50.0], dtype=torch.float64) + w = inverse_variance_weight(snr, sa, cap=1e9) + assert bool(torch.isfinite(w).all()) + # snr^2 is the ceiling, approached from below: the residual model variance + # is small but not zero, and it bites harder the better the measurement. + assert float(w[0]) <= 25.0 and float(w[0]) == pytest.approx(25.0, rel=1e-2) + assert float(w[1]) <= 2500.0 and float(w[1]) == pytest.approx(2500.0, rel=1e-2) + + +def test_the_weight_keeps_its_dependence_on_model_error(): + """The factorised form lost this, which is why it was abandoned. + + A product of a ``snr`` term and ``sigma_A/(eps - sigma_A^2)`` came out + identical for a 0.5 A and a 1.0 A coordinate error, because the singularity + set the shape and the model error dropped out. The coupled form must not. + """ + import math + s = torch.linspace(0.02, 0.5, 12, dtype=torch.float64) + snr = torch.full_like(s, 8.0) + curves = [] + for dv in (0.5, 1.0): + sa = torch.exp(-(2.0 / 3.0) * (math.pi ** 2) * s * s * dv * dv) + curves.append(normalise_weight( + inverse_variance_weight(snr, sa, cap=1e9))) + spread = float((curves[0] / curves[1] - 1).abs().max()) + assert spread > 0.1, ( + f"the weight barely moved ({spread:.3f}) between a 0.5 A and a 1.0 A " + f"model error; it has stopped carrying model information" + ) + + +def test_epsilon_enters_the_variance_not_the_signal(): + """``V = eps - sigma_A^2``, so higher multiplicity means more variance and + therefore less weight, at fixed model reliability and measurement error.""" + sa = torch.full((4,), 0.5, dtype=torch.float64) + snr = torch.full((4,), 10.0, dtype=torch.float64) + plain = inverse_variance_weight(snr, sa, cap=1e9) + axial = inverse_variance_weight( + snr, sa, eps=torch.full((4,), 2.0, dtype=torch.float64), cap=1e9) + assert bool((axial < plain).all()) + + +def test_normalise_weight_sets_the_mean(): + w = torch.rand(1000, dtype=torch.float64) * 7.0 + 0.5 + assert float(normalise_weight(w).mean()) == pytest.approx(1.0) + + +def test_a_constant_weight_survives_normalisation_as_ones(): + w = torch.full((100,), 3.7, dtype=torch.float64) + assert torch.allclose(normalise_weight(w), torch.ones_like(w)) diff --git a/torchref/experimental/alignment/frf/api.py b/torchref/experimental/alignment/frf/api.py index 5e17712a..a1e5fe62 100644 --- a/torchref/experimental/alignment/frf/api.py +++ b/torchref/experimental/alignment/frf/api.py @@ -21,6 +21,10 @@ import torch +from torchref.scaling.weighting import ( + DEFAULT_SNR_CAP, DEFAULT_TRUST_CAP, information_weight, + inverse_variance_weight, normalise_weight, snr_from_amplitude, +) from ..e_values import (SmoothSigmaE, convention_for_calc, convention_uses_sigma_f) from .data_mr import bessel_sh_expand, cross_correlate_xi @@ -140,6 +144,10 @@ def __init__( asu_idx: Optional[torch.Tensor] = None, s_mag_asu: Optional[torch.Tensor] = None, e_convention: type = SmoothSigmaE, + obs_weight: str = "inverse_variance", + snr_cap: float = DEFAULT_SNR_CAP, + trust_cap: float = DEFAULT_TRUST_CAP, + shell_variance_weights: bool = False, ): self.device = s_obs.device @@ -154,6 +162,20 @@ def __init__( self.delta_vrms_A = delta_vrms_A self.n_wilson_shells = n_wilson_shells self.grid_sampling_deg = grid_sampling_deg + # How much each observation counts. Separate from the convention above, + # which decides only what is compared -- see torchref.scaling.weighting. + # "information" is the saturating I/sigma weight; "none" is unit weight, + # which is the control that says what the weighting is worth at all. + self.obs_weight = obs_weight + self.snr_cap = float(snr_cap) + self.trust_cap = float(trust_cap) + # Off by default. It is a PER-SHELL weight, and per-shell weights are + # absorbed: it renormalises whatever scale the convention produced, which + # is precisely why every E convention measured as gauge. With the + # resolution weighting now declared on the calc side, this is a second + # mechanism for a job that already has one. + self.shell_variance_weights = bool(shell_variance_weights) + # The class, not an instance: a fitted Sigma(s) cannot exist before the # reflections do, so the convention is constructed here -- twice, once # per side. That the same class has to normalise both is the point; @@ -242,15 +264,40 @@ def __init__( ) self._conv_obs = conv_obs - # 4. LERF1 obs intensity, and the per-shell variance reweight. + # 4. LERF1 obs intensity, with the measurement-information weight. + if obs_weight == "none" or sig_F_obs is None: + w_obs = torch.ones_like(conv_obs.E) + elif obs_weight == "information": + # Measurement error only. The control that says what folding the + # model error in is actually worth. + w_obs = normalise_weight(information_weight( + snr_from_amplitude(F_obs, sig_F_obs), cap=self.snr_cap, + )) + elif obs_weight == "inverse_variance": + # Both error sources in one denominator. sigma_A is the same + # Luzzati falloff the calculated side uses for its expected signal; + # evaluated here at the OBSERVED reflections, because that is the + # side carrying sigmas and the weight has to be per reflection to + # be worth anything -- a weight constant within a shell is absorbed. + sigma_a_obs = eterm_sigma_a(smag_src, self.delta_vrms_A) + w_obs = normalise_weight(inverse_variance_weight( + snr_from_amplitude(F_obs, sig_F_obs), sigma_a_obs, + cap=self.trust_cap, + )) + else: + raise ValueError( + f"obs_weight={obs_weight!r}; expected 'inverse_variance', " + f"'information' or 'none'." + ) + self._w_obs = w_obs intensity_obs = build_lerf1_intensity( - conv_obs.E, centric_obs, weight=conv_obs.weight, - use_centric_weight=True, - ) - intensity_obs = apply_shell_variance_weights( - intensity_obs, smag_src, n_var_shells=n_wilson_shells, - shell_idx=obs_shell_idx, + conv_obs.E, centric_obs, weight=w_obs, use_centric_weight=True, ) + if self.shell_variance_weights: + intensity_obs = apply_shell_variance_weights( + intensity_obs, smag_src, n_var_shells=n_wilson_shells, + shell_idx=obs_shell_idx, + ) if asu_idx is not None: # One value per unique reflection -> one per unrolled reflection. intensity_obs = intensity_obs[asu_idx] @@ -305,7 +352,13 @@ def score_model( smag_calc, fsol=solvent_fsol, bsol=solvent_bsol, ).to(eterm.dtype) eterm = eterm * sol - intensity_calc = (eterm * eterm) * (E_calc * E_calc - 1.0) + # sigma_A**2 here is the expected moving-model intensity, i.e. part of + # the SIGNAL, not a weight -- which is why it stays on this side and the + # inverse-variance weight goes on the observations. Weighting the model + # by its own reliability and then also weighting the data by it would + # count the same thing twice. + w_calc = eterm * eterm + intensity_calc = w_calc * (E_calc * E_calc - 1.0) # Bessel-SH expand calc. The calc is NEVER m-filtered (zsymm=1): the # model carries no crystal symmetry, and projecting it onto the obs's diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index 8a357752..3a01a3fa 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -28,6 +28,8 @@ import torch +from torchref.scaling.weighting import (DEFAULT_SNR_CAP, + DEFAULT_TRUST_CAP) from .e_values import SmoothSigmaE from .sh import ( apply_overall_anisotropy, @@ -222,6 +224,10 @@ def search_peaks( verbose: int = 0, device: Optional[torch.device] = None, e_convention: type = SmoothSigmaE, + obs_weight: str = "inverse_variance", + shell_variance_weights: bool = False, + snr_cap: float = DEFAULT_SNR_CAP, + trust_cap: float = DEFAULT_TRUST_CAP, ) -> Tuple[List["RotationPeak"], int, float]: """Run the rotation function, returning the engine's own peak list. @@ -372,6 +378,8 @@ def search_peaks( asu_idx=asu_idx, s_mag_asu=s_mag_asu, e_convention=e_convention, + obs_weight=obs_weight, snr_cap=snr_cap, trust_cap=trust_cap, + shell_variance_weights=shell_variance_weights, ) _arf, peaks = engine.score_model( s_calc, F_calc, n_peaks=n_peaks, diff --git a/torchref/scaling/weighting.py b/torchref/scaling/weighting.py new file mode 100644 index 00000000..b03dc1e4 --- /dev/null +++ b/torchref/scaling/weighting.py @@ -0,0 +1,164 @@ +"""How much should a reflection count? The other half of the scaling/weighting split. + +:mod:`torchref.scaling.wilson` answers *what* we compare -- it removes the +resolution trend and leaves `` = 1``. That question turns out to be gauge +for a correlation: any per-shell scaling is absorbed downstream, which is why +twelve normalisation conventions moved a rotation function's truth rank by +nothing. This module answers the question that is not gauge. + +The weight has two sources: + +* **Measurement error**, per reflection, from ``I/sigma_I``. Two reflections at + the same resolution can differ enormously in how well they were measured, and + this is the only part that varies *within* a shell. That matters more than it + sounds -- a weight constant within a shell is a per-shell weight, and those + are exactly what a correlation absorbs. +* **Model error**, per resolution, through ``sigma_A``, which is smooth in + ``|s|`` and has no per-reflection content at all. + +Both come out of weighting by inverse variance, +``w = 1/(sigma_meas^2 + sigma_model^2)`` with ``sigma_model^2 = eps - +sigma_A^2``, the standard MLHL budget. The saturation usually added by hand is +already in there: once model error dominates, extra measurement precision buys +nothing, because what is wrong is the model and not the data. + +**They do not separate, and that was measured rather than assumed.** Writing the +weight as a product of a ``snr`` term and a ``sigma_A`` term looks natural and +fails: the ``sigma_A`` half is then dominated by its own singularity as +``sigma_A -> 1`` and stops depending on the model error it is named after. The +measurement term is what regularises it, so the two belong in one denominator. +:func:`inverse_variance_weight` is the form that works; +:func:`information_weight` is the pure measurement half, kept because it is the +right answer when no ``sigma_A`` is available and because it is the control that +says what the coupling is worth. +""" + +from __future__ import annotations + +from typing import Optional + +import torch + +__all__ = [ + "information_weight", + "inverse_variance_weight", + "snr_from_amplitude", + "normalise_weight", +] + +#: ``I/sigma_I`` at which measurement error stops being the limiting term. +#: Beyond it a reflection is no better determined for the purpose of comparing +#: against a model, because the model is what limits. Free parameter in +#: practice: the crossover really sits wherever ``sigma_A`` puts it, and +#: ``sigma_A`` before placement is assumed rather than fitted. +DEFAULT_SNR_CAP = 5.0 + +#: Backstop on the inverse-variance weight, for the case where measurement and +#: model variance both vanish. Not the working mechanism: the measurement term +#: is what bounds the weight at low resolution, where ``sigma_A -> 1`` would +#: otherwise send it to infinity on the strongest reflections. If this binds on +#: real data, ``sigma_A`` is wrong rather than the cap being too low. +DEFAULT_TRUST_CAP = 100.0 + + +def snr_from_amplitude( + F: torch.Tensor, sig_F: torch.Tensor, floor: float = 1e-12, +) -> torch.Tensor: + """``I/sigma_I`` from an amplitude and its error. + + With ``I = F^2`` the error propagates as ``sigma_I = 2 F sigma_F``, so the + intensity signal-to-noise is ``F / (2 sigma_F)`` -- half the amplitude's. + The factor is worth being explicit about: it only rescales the cap, but + quoting a cap against the wrong one silently doubles it. + """ + return (F.abs() / (2.0 * sig_F.abs().clamp(min=floor))).clamp(min=0.0) + + +def information_weight( + snr: torch.Tensor, *, cap: float = DEFAULT_SNR_CAP, +) -> torch.Tensor: + """Saturating measurement-information weight, ``snr^2 / (snr^2 + cap^2)``. + + Rises as ``(snr/cap)^2`` while measurement error dominates and flattens to 1 + once it does not. This is not a sigmoid chosen for its shape -- it is + ``1 / (1 + sigma_meas^2/sigma_model^2)`` rewritten, with ``cap`` the + signal-to-noise at which the two are equal. The saturation is a consequence + of the variance budget rather than a clip applied on top of one. + + Structurally this is what Phaser's ``DFAC`` already does: a monotone + function of signal-to-noise, in ``(0, 1)``, tending to 1 for well-measured + reflections. The difference is that its saturation point falls out of a Rice + moment calculation and cannot be moved, and this one is a number that can be + screened. + + Parameters + ---------- + snr : torch.Tensor + ``(N,)`` ``I/sigma_I``. Negative or zero values give weight 0, which is + the right answer for a measurement consistent with nothing. + cap : float, optional + Signal-to-noise at which the weight reaches 1/2. + """ + s2 = snr.clamp(min=0.0) ** 2 + return s2 / (s2 + float(cap) ** 2) + + +def inverse_variance_weight( + snr: torch.Tensor, + sigma_a: torch.Tensor, + *, + eps: Optional[torch.Tensor] = None, + cap: float = DEFAULT_TRUST_CAP, +) -> torch.Tensor: + """``1 / (1/snr^2 + eps - sigma_A^2)`` -- the two error sources, together. + + The obvious design is to factorise: a per-reflection term in ``snr`` times a + resolution term in ``sigma_A``. It does not work, and the reason is worth + keeping. + + Taken alone, ``sigma_A/(eps - sigma_A^2)`` is dominated by its own + singularity. On a realistic Luzzati falloff it runs 10.1 at the lowest + resolution shell against 1.0 at the next, and it comes out *identical* for + a 0.5 A and a 1.0 A coordinate error -- because as ``sigma_A -> 1`` the + shape is set entirely by ``1/(1 - sigma_A^2)`` and the model error it was + supposed to encode drops out. A weight carrying no information about the + thing it is named after is not a weight worth having. + + What regularises it is the term the factorisation threw away. Those + low-resolution reflections are strong but not infinitely well measured, so + ``1/snr^2`` is what stops the variance reaching zero. The two sources have + to sit in one denominator; they do not separate. + + ``1/snr^2`` is the measurement variance expressed in the same units as + ``eps - sigma_A^2``, i.e. relative to a normalised `` = 1``. The cap + is a backstop for the pathological case where both terms vanish, not the + working mechanism -- if it binds on real data, ``sigma_A`` is wrong. + + Parameters + ---------- + snr : torch.Tensor + ``(N,)`` ``I/sigma_I``. Zero gives zero weight. + sigma_a : torch.Tensor + ``(N,)`` model reliability in ``[0, 1)``, evaluated at each reflection. + eps : torch.Tensor, optional + ``(N,)`` multiplicity; ``None`` means 1. + cap : float, optional + Ceiling on the weight before normalisation. + """ + sa = sigma_a.clamp(min=0.0, max=1.0 - 1e-6) + e = torch.ones_like(sa) if eps is None else eps.to(sa.dtype).clamp(min=1.0) + v_meas = 1.0 / (snr.clamp(min=1e-8) ** 2) + v_model = (e - sa * sa).clamp(min=0.0) + w = 1.0 / (v_meas + v_model).clamp(min=1e-12) + return w.clamp(max=float(cap)) + + +def normalise_weight(w: torch.Tensor) -> torch.Tensor: + """Scale a weight to mean 1. + + Cosmetic for a correlation, where an overall factor cancels, and not + cosmetic for anything that compares scores across runs or reads a sigma + level off them. Doing it here means the cap is a number about *relative* + weighting rather than one entangled with whatever scale the inputs had. + """ + return w / w.mean().clamp(min=1e-30) From 2c96b80d07635ddaec5c2aa11331c6b3097d0fbe Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sun, 30 Aug 2026 14:08:58 +0200 Subject: [PATCH 104/250] Measure sigma_A instead of assuming it, completing the two-part system The point of splitting scaling from weighting was to replace seven knobs with two declared objects. Four went with the first two commits. This is the rest. `sigma_A` in the rotation search was a *prior*: a Luzzati falloff from a coordinate error guessed off the residue count, patched at low resolution by Babinet's two universal constants. It never saw a residual. It does not have to be assumed. Total scattering per shell is ROTATION-INVARIANT, so the model's resolution-dependent deficiency is measurable before the molecule is placed, even though the per-reflection one is not. Normalise both sides to unit mean and their fitted curves' ratio is exactly that. `sigma_A = sqrt(min(R, 1/R))`, since `sigma_A^2` is the fraction of intensity the model accounts for and neither side can share more than there is. Safe to take from the data being scored, which normally it would not be: the quantity is identical for every candidate orientation, so it shifts all scores together and cannot bias the ranking between them. MEASURED, 4 arms x 100 paired cells. Ranking is flat as everything on this axis has been -- but top-20 is not, and it is the one place any of this has ever shown an effect: luzzati + babinet (control) rank-0 20/100 top-20 98/100 luzzati, no babinet rank-0 19/100 top-20 94/100 empirical rank-0 19/100 top-20 98/100 Dropping Babinet costs four cells; the measured curve puts them back, with none of Babinet's constants and no assumed coordinate error. Four cells is not significant on its own, and the direction is what the argument predicted. Also removed: the relative Wilson-B match in the rotation search. It multiplied `F_calc` by exp(-B s^2/4) -- a smooth function of |s| -- and the very next step divides out exactly such a function. Measured: +-30 A^2 of relative B moves E by at most 1.3e-7, the fit's own convergence tolerance. Distinct from the earlier finding that knocking it out was rank-neutral; that said it did not matter, this says it was arithmetically undone. It stays for the rescore, where the calc normalisation comes from the unmodified reference amplitudes so the term survives. And obs and calc now share one resolution window, taken from the bandwidth coupling rather than from whichever reflections each side happens to hold. Two fits on their own extremes span the same polynomial space and agree where their data overlap, but the basis saturates at the ends, so each is frozen flat outside its own range -- and anything reading their ratio needs them on one abscissa. Scorecard, seven knobs to two objects: E convention obs/calc -> WilsonNormaliser, one class called twice DFAC^2 -> inverse_variance_weight shell variance reweight -> gone (per-shell, therefore absorbed; p = 1.00) relative Wilson B -> gone (cancelled by the normalisation) Luzzati sigma_A + Babinet -> measured from Sigma_obs/Sigma_calc Gate: 2006 passed, 0 failed. Seam identity 0.000e+00. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/weight_arms.py | 66 +++++++++---------- alignment_lab/lab/frf.py | 6 +- docs/changelog.rst | 2 + torchref/experimental/alignment/e_values.py | 9 +++ torchref/experimental/alignment/frf/api.py | 58 ++++++++++++++-- .../experimental/alignment/rotation_search.py | 33 +++++----- torchref/scaling/weighting.py | 53 +++++++++++++++ 7 files changed, 166 insertions(+), 61 deletions(-) diff --git a/alignment_lab/analysis/weight_arms.py b/alignment_lab/analysis/weight_arms.py index 1fabe5d0..8ee2dba5 100644 --- a/alignment_lab/analysis/weight_arms.py +++ b/alignment_lab/analysis/weight_arms.py @@ -1,23 +1,21 @@ -"""What is the observed-side weight worth, and which part of it? - -The scaler split scaling from weighting and showed the scaling half is gauge for -a correlation. This is the half that is not. - -Arms deviate one thing at a time from `none_shellvar`, which is what the engine -did before any of this. Moving several at once and reading one number is what -made the E-convention panel uninterpretable. - -Three things are being asked: - -* is a measurement-error weight worth anything at all (`none` vs `information`); -* is folding model error into the same denominator worth more than the - measurement term alone (`information` vs `inverse_var`) -- they do not - factorise, so this is the only way to separate them; -* is `apply_shell_variance_weights` doing anything, given it is a per-shell - weight and per-shell weights are absorbed (`*_shellvar` pairs). - -The caps are screened rather than assumed. `trust_cap` is supposed to be a -backstop that never binds, so if it moves the result, sigma_A is wrong. +"""Does the measured model-error curve replace the assumed one? + +The two-part system -- one scaler, one weight -- was meant to subsume seven +separate knobs. Four went with the scaler and the weight. These arms test the +last of them: `sigma_A` itself, which is currently a *prior* (a Luzzati falloff +from a coordinate error guessed off the residue count) patched at low resolution +by Babinet's two universal constants. + +It does not have to be assumed. Total scattering per shell is rotation- +invariant, so `Sigma_obs(s)/Sigma_calc(s)` measures the model's resolution- +dependent deficiency before the molecule is placed -- and it is safe to take +from the data being scored, because a quantity identical for every candidate +orientation cannot bias the ranking between them. + +Ranking is expected to be flat: it has been flat across every configuration of +scaling and weighting tried so far. That is not the question. The question is +whether the measured curve can stand in for the assumed one, so that two +declared objects replace seven knobs rather than four of them. """ from __future__ import annotations @@ -37,22 +35,18 @@ #: Each arm is one deviation from `control`, which is what ships today. ARMS = { - # What shipped before any of this: unit observed weight (DFAC left with the - # French-Wilson posterior when the scaler landed) and the per-shell reweight - # on. The load-bearing control -- without it, "the weight helps" cannot be - # told apart from "any weight helps". - "none_shellvar": {"obs_weight": "none", "shell_variance_weights": True}, - "none": {"obs_weight": "none"}, - "information": {"obs_weight": "information"}, - "inverse_var": {"obs_weight": "inverse_variance"}, - "invvar_shellvar": {"obs_weight": "inverse_variance", - "shell_variance_weights": True}, - # Cap screens. The SNR cap sets where measurement error stops limiting; the - # trust cap is meant to be a backstop, so if these move the result it is - # binding and sigma_A is wrong. - "snr_cap_2": {"obs_weight": "information", "snr_cap": 2.0}, - "snr_cap_10": {"obs_weight": "information", "snr_cap": 10.0}, - "trust_cap_10": {"obs_weight": "inverse_variance", "trust_cap": 10.0}, + # What the engine did before this line of work. + "luzzati_babinet": {"sigma_a_source": "luzzati", "apply_bulk_solvent": True}, + # Luzzati without the two universal Babinet constants: how much of the + # low-resolution correction was the solvent term doing? + "luzzati_only": {"sigma_a_source": "luzzati", "apply_bulk_solvent": False}, + # sigma_A measured from Sigma_obs/Sigma_calc instead of assumed from an + # estimated coordinate error. Subsumes Babinet -- the solvent deficit is + # what the ratio measures -- so the solvent flag is irrelevant here. + "empirical": {"sigma_a_source": "empirical"}, + # And the same with no observed-side weight, to check the two halves of the + # system are still independent of each other. + "empirical_now": {"sigma_a_source": "empirical", "obs_weight": "none"}, } diff --git a/alignment_lab/lab/frf.py b/alignment_lab/lab/frf.py index 5743878f..a3cf4d59 100644 --- a/alignment_lab/lab/frf.py +++ b/alignment_lab/lab/frf.py @@ -57,6 +57,8 @@ class FRFConfig: #: and whether the per-shell reweight runs -- and a panel that moves all #: three cannot say which one did anything. obs_weight: Optional[str] = None + sigma_a_source: Optional[str] = None + apply_bulk_solvent: Optional[bool] = None shell_variance_weights: Optional[bool] = None snr_cap: Optional[float] = None trust_cap: Optional[float] = None @@ -246,8 +248,8 @@ def _wrapped(self, *args, **kwargs): "e_convention": cfg.e_convention} # Engine knobs are omitted when unset so the production default applies, # rather than being passed as None and overriding it with nothing. - for _name in ("obs_weight", "shell_variance_weights", - "snr_cap", "trust_cap"): + for _name in ("obs_weight", "shell_variance_weights", "snr_cap", + "trust_cap", "sigma_a_source", "apply_bulk_solvent"): _v = getattr(cfg, _name) if _v is not None: conv_kw[_name] = _v diff --git a/docs/changelog.rst b/docs/changelog.rst index c80801b1..c683cf14 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -9,6 +9,8 @@ Unreleased - Fixed the overall-anisotropy fit, which regressed log intensities with no constant term and so absorbed the ``-gamma`` offset into the tensor - Fixed molecular-replacement rotation candidates being composed onto each other instead of onto the search model - Fixed assigning a ``SpaceGroup`` object to ``Model.spacegroup`` being a silent no-op that then made the correct name assignment raise +- The rotation function estimates ``sigma_A`` from the data instead of assuming it. Total scattering per shell is rotation-invariant, so ``Sigma_obs(s)/Sigma_calc(s)`` measures the model's resolution-dependent deficiency before placement; it replaces the Luzzati falloff from an estimated coordinate error and the Babinet bulk-solvent term with its two universal constants +- Removed the relative Wilson-B match from the rotation search. It multiplied ``F_calc`` by a smooth function of ``|s|`` that the normalisation then divided straight back out - Added ``torchref.scaling.weighting``: measurement and model error combined as one inverse-variance weight per reflection, restoring the observed-side weighting that left with the French-Wilson posterior - ``apply_shell_variance_weights`` is off by default in the rotation function. It is a per-shell weight, and per-shell weights are absorbed by the correlation; switching it on moves nothing - The rotation function, ML rescore and translation search now normalise through the shared Wilson normaliser by default, replacing the French-Wilson posterior on the observed side. Rank-neutral over 10 structures x 10 seeds; the rotation function is about twice as fast, since the posterior and its D-factor iteration are no longer on the default path diff --git a/torchref/experimental/alignment/e_values.py b/torchref/experimental/alignment/e_values.py index d04fdb0a..83efa8ba 100644 --- a/torchref/experimental/alignment/e_values.py +++ b/torchref/experimental/alignment/e_values.py @@ -366,5 +366,14 @@ def _compute(self): self._intensity(), self.s_mag, centric=self.centric, n_coeff=self.n_coeff, s_lo=self.s_lo, s_hi=self.s_hi, ) + # Kept, not discarded: the fitted CURVE is the thing anything comparing + # two normalisations needs. Sigma_obs/Sigma_calc is how model error is + # measured rather than assumed, and it can only be formed from the fits + # themselves, not from the per-reflection values they produced. + self.fit = fit self.sigma = fit.sigma_wilson return fit.E, self._ones() + + def evaluate(self, s_mag: torch.Tensor) -> torch.Tensor: + """``Sigma(s)`` at arbitrary ``|s|``, on this fit's own abscissa.""" + return self.fit.evaluate(s_mag) diff --git a/torchref/experimental/alignment/frf/api.py b/torchref/experimental/alignment/frf/api.py index a1e5fe62..66fcb621 100644 --- a/torchref/experimental/alignment/frf/api.py +++ b/torchref/experimental/alignment/frf/api.py @@ -22,8 +22,9 @@ import torch from torchref.scaling.weighting import ( - DEFAULT_SNR_CAP, DEFAULT_TRUST_CAP, information_weight, - inverse_variance_weight, normalise_weight, snr_from_amplitude, + DEFAULT_SNR_CAP, DEFAULT_TRUST_CAP, empirical_sigma_a, + information_weight, inverse_variance_weight, normalise_weight, + snr_from_amplitude, ) from ..e_values import (SmoothSigmaE, convention_for_calc, convention_uses_sigma_f) @@ -258,9 +259,19 @@ def __init__( obs_cls = e_convention if sig_F_obs is None and convention_uses_sigma_f(obs_cls): obs_cls = convention_for_calc(obs_cls) + # One resolution window for both sides, taken from the bandwidth + # coupling rather than from whichever reflections each side happens to + # contain. Two fits on their own extremes span the same polynomial + # space, so they agree where their data overlap -- but the basis + # saturates at the ends, so each is frozen flat outside its own range + # and the two curves stop being usable against each other. Everything + # that reads Sigma_obs/Sigma_calc needs them on one abscissa. + self._s_lo = 1.0 / float(d_max) if d_max else float(smag_src.min()) + self._s_hi = 1.0 / float(d_min) conv_obs = obs_cls( F_obs, smag_src, centric_obs, sig_F=sig_F_obs, shell_idx=obs_shell_idx, n_shells=n_wilson_shells, + **self._range_kw(obs_cls), ) self._conv_obs = conv_obs @@ -317,6 +328,23 @@ def __init__( ) + def _range_kw(self, cls) -> dict: + """``s_lo``/``s_hi`` for conventions that fit a curve, empty otherwise. + + Only the smooth normaliser has an abscissa to pin; the per-shell + conventions bin whatever they are given and would reject the argument. + """ + import inspect + + target = getattr(cls, "func", cls) + try: + params = inspect.signature(target.__init__).parameters + except (TypeError, ValueError): # pragma: no cover + return {} + if "s_lo" not in params: + return {} + return {"s_lo": self._s_lo, "s_hi": self._s_hi} + def score_model( self, s_calc: torch.Tensor, @@ -324,6 +352,7 @@ def score_model( *, n_peaks: int = 500, sigma_threshold: float = -5.0, + sigma_a_source: str = "empirical", apply_bulk_solvent: bool = False, solvent_fsol: float = 0.95, solvent_bsol: float = 300.0, @@ -339,10 +368,29 @@ def score_model( # measured to drop 0 of 339040 reflections on 3K7M and 0 of 271630 on # 1DAW. So take `s_calc` as given and only derive |s| from it. smag_calc = s_calc.norm(dim=-1) - E_calc = convention_for_calc(self.e_convention)( + calc_cls = convention_for_calc(self.e_convention) + self._conv_calc = calc_cls( F_calc, smag_calc, n_shells=self.n_wilson_shells, - ).E - eterm = eterm_sigma_a(smag_calc, self.delta_vrms_A) + **self._range_kw(calc_cls), + ) + E_calc = self._conv_calc.E + if sigma_a_source == "empirical": + # Measured, not assumed. Both curves were fitted on one abscissa + # (see `_range_kw`), which is what makes evaluating them at the same + # |s| meaningful. Subsumes the Babinet term: the low-resolution + # solvent deficit is simply what the ratio measures, per structure, + # instead of two universal constants. + eterm = empirical_sigma_a( + self._conv_obs.evaluate(smag_calc).to(torch.float64), + self._conv_calc.evaluate(smag_calc).to(torch.float64), + ).to(F_calc.dtype) + elif sigma_a_source == "luzzati": + eterm = eterm_sigma_a(smag_calc, self.delta_vrms_A) + else: + raise ValueError( + f"sigma_a_source={sigma_a_source!r}; expected 'luzzati' or " + f"'empirical'." + ) # Optional Babinet bulk-solvent factor: Phaser folds it into σ_A as # `σ_A_eff = solTerm(s²) · Luzzati(s², vrms)` (EnsemblePDB.cc:96-100). # Default OFF; flip after v25 validates. diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index 3a01a3fa..605bf291 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -225,6 +225,8 @@ def search_peaks( device: Optional[torch.device] = None, e_convention: type = SmoothSigmaE, obs_weight: str = "inverse_variance", + sigma_a_source: str = "empirical", + apply_bulk_solvent: bool = False, shell_variance_weights: bool = False, snr_cap: float = DEFAULT_SNR_CAP, trust_cap: float = DEFAULT_TRUST_CAP, @@ -240,7 +242,6 @@ def search_peaks( from ...utils import resolve_device from .frf.api import FastRotationFunction, phaser_lmax_resolution from .frf.dense_calc import dense_calc_via_box - from .frf.preprocessing import fit_relative_wilson_b # One device for both inputs, rather than whichever one this function # happened to read first: `resolve_device` moves them into agreement (with a @@ -351,22 +352,17 @@ def search_peaks( s_calc = s_calc.to(device) F_calc = F_calc.to(device) - # Put the model's amplitudes on the observations' overall B scale - # (EnsemblePDB.cc:793-851), so the radial fall-off does not by itself - # discriminate between orientations. - s_calc_mag = s_calc.norm(dim=-1) - # Fitted on the unique set. The unroll replicates every reflection - # exactly n_ops times, so the per-shell means are identical to the - # unrolled fit while the sort and the binning are n_ops times smaller. - B_rel = fit_relative_wilson_b( - F_obs.to(torch.float64), F_calc.to(torch.float64), - s_mag_asu.to(torch.float64), n_shells=N_WILSON_SHELLS, - s_mag_calc=s_calc_mag.to(torch.float64), - ) - if abs(B_rel) > 1e-6: - F_calc = F_calc * torch.exp(-B_rel * (s_calc_mag * s_calc_mag) / 4.0) - if verbose > 0: - print(f" relative Wilson B = {B_rel:+.2f} A^2", flush=True) + # No relative Wilson-B match here any more. It multiplied `F_calc` by + # exp(-B s^2/4) -- a smooth function of |s| -- and the engine's very next + # step divides out exactly such a function when it normalises. Measured: + # a relative B of +-30 A^2 moves E by at most 1.3e-7, the fit's own + # convergence tolerance. It was computing a number and having it undone. + # + # Not the same as the earlier finding that knocking it out was + # rank-neutral; that was a measurement about whether it mattered, this is + # that it is arithmetically cancelled. `fit_relative_wilson_b` stays for + # the rescore, where the calc normalisation is taken from the unmodified + # reference amplitudes and the Debye-Waller term therefore survives. engine = FastRotationFunction( s_obs, F_obs, centric, sg_mats, @@ -384,7 +380,8 @@ def search_peaks( _arf, peaks = engine.score_model( s_calc, F_calc, n_peaks=n_peaks, sigma_threshold=SIGMA_THRESHOLD, - apply_bulk_solvent=True, + sigma_a_source=sigma_a_source, + apply_bulk_solvent=apply_bulk_solvent, solvent_fsol=SOLVENT_FSOL, solvent_bsol=SOLVENT_BSOL, ) diff --git a/torchref/scaling/weighting.py b/torchref/scaling/weighting.py index b03dc1e4..d874cd86 100644 --- a/torchref/scaling/weighting.py +++ b/torchref/scaling/weighting.py @@ -44,6 +44,7 @@ "inverse_variance_weight", "snr_from_amplitude", "normalise_weight", + "empirical_sigma_a", ] #: ``I/sigma_I`` at which measurement error stops being the limiting term. @@ -162,3 +163,55 @@ def normalise_weight(w: torch.Tensor) -> torch.Tensor: weighting rather than one entangled with whatever scale the inputs had. """ return w / w.mean().clamp(min=1e-30) + + +def empirical_sigma_a( + sigma_obs: torch.Tensor, + sigma_calc: torch.Tensor, + *, + floor: float = 1e-3, +) -> torch.Tensor: + """Model reliability measured, rather than assumed, from two Wilson curves. + + ``sigma_A`` in a rotation search is normally a *prior*: a Luzzati falloff + from a coordinate error guessed off the residue count, patched at low + resolution by Babinet's two universal constants. It never sees a residual. + + It does not have to. Total scattering per shell is **rotation-invariant**, + so the resolution-dependent disagreement between model and data is + measurable before the molecule is placed, even though the per-reflection + disagreement is not. Normalise both sides to `` = 1`` and their fitted + curves' ratio is exactly that disagreement: + + R(s) = Sigma_obs(s) / Sigma_calc(s) + + ``R < 1`` means the model predicts more scattering than is there, which at + low resolution is the bulk solvent it does not have; ``R > 1`` means it + predicts less. Either way the shared fraction is bounded by + ``min(R, 1/R)``, and ``sigma_A`` is its square root because ``sigma_A^2`` is + the fraction of intensity the model accounts for. + + **This is safe to estimate from the data being scored**, which normally it + would not be: the quantity is identical for every candidate orientation, so + it shifts all scores together and cannot bias the ranking toward any of + them. + + What it conflates -- solvent, an overall B mismatch, missing atoms, genuine + coordinate error -- it conflates deliberately. For deciding how far to trust + a resolution range the cause does not matter, only the size. What it cannot + see is *completeness*: forcing both sides to unit mean absorbs a uniform + factor, so a model that is half the asymmetric unit looks like a model that + is all of it, and only the tilt survives. + + Parameters + ---------- + sigma_obs, sigma_calc : torch.Tensor + ``(N,)`` fitted Wilson curves evaluated at the same ``|s|``. They must + come from fits sharing an abscissa, or each is frozen flat outside its + own range and the ratio is meaningless there. + floor : float, optional + Lower bound on the returned ``sigma_A``. + """ + r = (sigma_obs / sigma_calc.clamp(min=1e-30)).clamp(min=1e-30) + shared = torch.minimum(r, 1.0 / r).clamp(min=0.0, max=1.0) + return shared.sqrt().clamp(min=float(floor), max=1.0 - 1e-6) From 84e6b4e9d28a60b4ded0c20e8b6fb407cc62dd59 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sun, 30 Aug 2026 22:30:42 +0200 Subject: [PATCH 105/250] Score orientations by weighted least squares, as an alternative to the Rice The Rice exists to handle an AMPLITUDE: |F_obs| is non-negative with an unknown phase, so the likelihood marginalises over the phase and comes out biased upward for weak reflections. `E**2` is an intensity -- unbiased, near Gaussian wherever it is measured at all -- which is the same reason the scaler works on intensities. And the job here is to RANK orientations, not to report calibrated probabilities; both sides are already normalised to unit mean square, so the Rice's shrinkage is being applied to a quantity whose scale is fixed by construction. `target="wls"` scores `-sum_h w_h (E_obs**2 - eImove)**2`, with `w` the same combined inverse variance the rotation function uses, so the two stages agree about what a reflection is worth as well as about what it is compared to. MEASURED, 100 paired cells. The target form makes no difference: rice -> wls is 20 better, 26 worse, p = 0.46. Kept as an option; no reason to move a default on a tie. The result that matters is the control. Every rescore arm is now significantly worse than not rescoring, which at n=30 was only a trend: none (raw FRF order) rank-0 19/100 top-3 54 median 2.0 default (Rice) rank-0 16/100 top-3 43 median 4.0 p = 0.053 wls rank-0 14/100 top-3 38 median 4.0 p = 0.033 wls, sigmas withheld rank-0 17/100 top-3 42 median 4.0 p = 0.013 solvent rank-0 18/100 top-3 47 median 3.0 p = 0.040 full_prep rank-0 15/100 top-3 48 median 3.0 p = 0.025 So the rescore's problem is not its target, not its model preparation, not its E convention and not its weighting -- all four have now been screened at n=100 and none of them moves it above its own input. A likelihood and a least-squares reorder the FRF's top-20 equally badly, which says the information it adds -- |F| at the crystal lattice under an assumed sigma_A -- is simply less discriminating than the Patterson correlation that produced the list. This is the paired-rank confirmation of what end-to-end pose recovery already said: dropping the rescore takes 18/30 to 24/30. Gate: 2006 passed, 0 failed. Seam identity 0.000e+00. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/rescore_prep_arms.py | 30 ++++++----- docs/changelog.rst | 1 + .../experimental/alignment/ml_rotation.py | 54 ++++++++++++++++--- 3 files changed, 65 insertions(+), 20 deletions(-) diff --git a/alignment_lab/analysis/rescore_prep_arms.py b/alignment_lab/analysis/rescore_prep_arms.py index d4a18838..68694450 100644 --- a/alignment_lab/analysis/rescore_prep_arms.py +++ b/alignment_lab/analysis/rescore_prep_arms.py @@ -54,23 +54,25 @@ #: anywhere it matters here. `no_sigmas` is the control for that: it withholds #: the sigmas the rescore has only just started receiving. def _arms(): - from torchref.experimental.alignment.e_values import ( - CalcGlobalE, CalcShellE, FrenchWilsonE, SmoothSigmaE, WilsonShellE, - WilsonShellEpsE, - ) - import functools + """The target question: does a Rice buy anything over weighted least squares? + + The Rice exists to handle an AMPLITUDE -- non-negative, phase unknown, so the + likelihood marginalises over the phase and is biased upward for weak + reflections. `E**2` is an intensity, which is unbiased and near Gaussian, and + the job here is to rank orientations rather than to report calibrated + probabilities. Both sides are already normalised to unit mean square, so the + shrinkage the Rice contributes is being applied to something whose scale is + fixed by construction. + + `wls` scores `-sum_h w_h (E_obs**2 - eImove)**2` with `w` the same combined + inverse variance the rotation function uses, so both stages agree about what + a reflection is worth as well as about what it is compared to. + """ return { - "fw_sigmas": {"e_convention": FrenchWilsonE}, - "no_sigmas": {"sig_F_obs": None}, - "wilson": {"e_convention": WilsonShellE}, - "calc_shell": {"e_convention": CalcShellE}, - "calc_global": {"e_convention": CalcGlobalE}, - "smooth6": {"e_convention": functools.partial(SmoothSigmaE, - n_coeff=6)}, - "eps_wilson": {"e_convention": WilsonShellEpsE}, + "wls": {"target": "wls"}, + "wls_nosig": {"target": "wls", "sig_F_obs": None}, } - ARMS = { "none": None, # control: FRF order "default": {}, # what ships today diff --git a/docs/changelog.rst b/docs/changelog.rst index c683cf14..579d171c 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -9,6 +9,7 @@ Unreleased - Fixed the overall-anisotropy fit, which regressed log intensities with no constant term and so absorbed the ``-gamma`` offset into the tensor - Fixed molecular-replacement rotation candidates being composed onto each other instead of onto the search model - Fixed assigning a ``SpaceGroup`` object to ``Model.spacegroup`` being a silent no-op that then made the correct name assignment raise +- The ML rescore can score orientations by weighted least squares on E-space intensities (``target='wls'``) instead of the Rice/Woolfson likelihood. Measured indistinguishable from the Rice over 10 structures x 10 seeds; both remain worse than not rescoring at all - The rotation function estimates ``sigma_A`` from the data instead of assuming it. Total scattering per shell is rotation-invariant, so ``Sigma_obs(s)/Sigma_calc(s)`` measures the model's resolution-dependent deficiency before placement; it replaces the Luzzati falloff from an estimated coordinate error and the Babinet bulk-solvent term with its two universal constants - Removed the relative Wilson-B match from the rotation search. It multiplied ``F_calc`` by a smooth function of ``|s|`` that the normalisation then divided straight back out - Added ``torchref.scaling.weighting``: measurement and model error combined as one inverse-variance weight per reflection, restoring the observed-side weighting that left with the French-Wilson posterior diff --git a/torchref/experimental/alignment/ml_rotation.py b/torchref/experimental/alignment/ml_rotation.py index da46ed0a..a33fdc6d 100644 --- a/torchref/experimental/alignment/ml_rotation.py +++ b/torchref/experimental/alignment/ml_rotation.py @@ -25,6 +25,9 @@ import torch +from torchref.scaling.weighting import (inverse_variance_weight, + normalise_weight, + snr_from_amplitude) from .e_values import (CalcShellE, SmoothSigmaE, WilsonShellE, WilsonShellEpsE, convention_for_calc, convention_uses_sigma_f) @@ -535,6 +538,8 @@ class _LLGContext: dw_per_m: Optional[torch.Tensor] # (M,) or None — Wilson-B Debye-Waller dtype: torch.dtype batch_size: int + target: str = "rice" # "rice" | "wls" + w_b: Optional[torch.Tensor] = None # (1, N) weights for "wls" def _build_llg_context( @@ -560,6 +565,7 @@ def _build_llg_context( wilson_b_value: Optional[float] = None, sig_F_obs: Optional[torch.Tensor] = None, e_convention: type = SmoothSigmaE, + target: str = "rice", ) -> _LLGContext: """Build the rotation-independent m_LETF1 LLG context (DataMR.cc:1326-1429). @@ -698,12 +704,27 @@ def _build_llg_context( sqrt_mean_per_m = calc_norm_per_h[asu_idx] # (M,) dw_per_m = dw[asu_idx] if dw is not None else None + # Weights for the least-squares target: the same combined inverse variance + # the rotation function uses, so the two stages agree about how much a + # reflection is worth as well as about what it is being compared to. + w_b = None + if target == "wls": + if sig_F_obs is not None: + w = inverse_variance_weight( + snr_from_amplitude(F_obs, sig_F_obs), sigma_a.to(F_obs.dtype), + eps=eps_factor.to(F_obs.dtype), + ) + else: + w = torch.ones_like(F_obs) + w_b = normalise_weight(w).to(dtype).to(device).view(1, -1) + return _LLGContext( interpolator=interpolator, real_cell=real_cell, unrolled_hkl=unrolled_hkl, asu_idx=asu_idx, N=N, E_obs_b=E_obs_b, V_b=V_b, eImove_prefac=eImove_prefac, sqrt_mean_per_m=sqrt_mean_per_m, centric_b=centric_b, dw_per_m=dw_per_m, dtype=dtype, batch_size=batch_size, + target=target, w_b=w_b, ) @@ -743,11 +764,31 @@ def _llg_for_orientations( idx = ctx.asu_idx.unsqueeze(0).expand(B, -1) # (B, M) sum_per_h.scatter_add_(1, idx, Esq_m) # (B, N) eImove = ctx.eImove_prefac * sum_per_h # (B, N) - sqrt_eImove = eImove.clamp(min=1e-30).sqrt() - ll_acen = phaser_log_rel_rice(ctx.E_obs_b, sqrt_eImove, ctx.V_b) - ll_cen = phaser_log_rel_woolfson(ctx.E_obs_b, sqrt_eImove, ctx.V_b) - ll = torch.where(ctx.centric_b, ll_cen, ll_acen) # (B, N) - chunks.append(ll.sum(dim=-1)) # (B,) + if ctx.target == "wls": + # Weighted least squares on INTENSITIES, which is what E**2 is. + # + # The Rice exists to handle an amplitude: |F_obs| is non-negative + # and its phase is unknown, so the likelihood marginalises over the + # phase and the result is biased upward for weak reflections. None + # of that applies to an intensity, which is unbiased and near + # Gaussian wherever it is measured at all -- the same reason the + # scaler works on intensities rather than amplitudes. + # + # And the job here is to RANK orientations, not to report calibrated + # probabilities. Both sides are already normalised to = 1, so + # the distributional shrinkage the Rice contributes is being applied + # to a quantity that has had its scale fixed by construction. + # + # Sign: a residual is a cost, so negate it to keep "larger is + # better" for every target the caller can pick. + resid = ctx.E_obs_b * ctx.E_obs_b - eImove # (B, N) + chunks.append(-(ctx.w_b * resid * resid).sum(dim=-1)) + else: + sqrt_eImove = eImove.clamp(min=1e-30).sqrt() + ll_acen = phaser_log_rel_rice(ctx.E_obs_b, sqrt_eImove, ctx.V_b) + ll_cen = phaser_log_rel_woolfson(ctx.E_obs_b, sqrt_eImove, ctx.V_b) + ll = torch.where(ctx.centric_b, ll_cen, ll_acen) # (B, N) + chunks.append(ll.sum(dim=-1)) # (B,) return torch.cat(chunks) @@ -1069,6 +1110,7 @@ def m_letf1_rescore( wilson_b_value: Optional[float] = None, # if None and apply_wilson_b=True, fitted from data sig_F_obs: Optional[torch.Tensor] = None, e_convention: type = SmoothSigmaE, + target: str = "rice", ) -> List[RotationPeak]: """Phaser-faithful ``m_LETF1`` rescore (DataMR.cc:1326-1429). @@ -1137,7 +1179,7 @@ def m_letf1_rescore( vrms_strategy=vrms_strategy, vrms_n_residues=vrms_n_residues, vrms_identity=vrms_identity, apply_wilson_b=apply_wilson_b, wilson_b_value=wilson_b_value, sig_F_obs=sig_F_obs, - e_convention=e_convention, + e_convention=e_convention, target=target, ) alpha_t = torch.tensor([p.alpha for p in head], dtype=torch.float64) From 0aecd13f2b762d1f135f0d1b14524cf857bd7be3 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 31 Aug 2026 13:52:04 +0200 Subject: [PATCH 106/250] Add the translation-function measurement harnesses Five diagnostics behind the FTF result: a discrimination panel that places the top FRF peaks and ranks them four ways, two early-stop analyses over its output (a fixed threshold and a running median/MAD null), a per-stage placement cost breakdown, and two structure-factor probes -- backend and grid size. The discrimination panel is the load-bearing one. Truth's rank among 25 placed candidates, 10 structures x 3 trials: FRF score 6/30 at rank 0, TF correlation 24/30, TF LLG 27/30, analytic R 22/30. The pipeline ranks by analytic R and leaves the LLG off. pose_recovery gains a machine-readable ROW line so its arms aggregate the same way the rest of the lab does. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/ftf_disc_smoke.sh | 18 ++ alignment_lab/analysis/ftf_discrimination.sh | 25 ++ alignment_lab/analysis/ftf_early_stop.py | 117 +++++++++ alignment_lab/analysis/ftf_running_null.py | 140 +++++++++++ alignment_lab/analysis/pose_arms.sh | 31 +++ alignment_lab/analysis/tf_cost.sh | 19 ++ .../analysis/tf_resolution_and_grid.sh | 19 ++ alignment_lab/analysis/tf_sf_backend.sh | 19 ++ .../diagnostics/frf_vs_ftf_discrimination.py | 234 ++++++++++++++++++ alignment_lab/diagnostics/pose_recovery.py | 4 + alignment_lab/diagnostics/tf_batch_probe.py | 152 ++++++++++++ alignment_lab/diagnostics/tf_cost.py | 105 ++++++++ .../diagnostics/tf_resolution_and_grid.py | 99 ++++++++ alignment_lab/diagnostics/tf_sf_backend.py | 112 +++++++++ 14 files changed, 1094 insertions(+) create mode 100644 alignment_lab/analysis/ftf_disc_smoke.sh create mode 100644 alignment_lab/analysis/ftf_discrimination.sh create mode 100644 alignment_lab/analysis/ftf_early_stop.py create mode 100644 alignment_lab/analysis/ftf_running_null.py create mode 100644 alignment_lab/analysis/pose_arms.sh create mode 100644 alignment_lab/analysis/tf_cost.sh create mode 100644 alignment_lab/analysis/tf_resolution_and_grid.sh create mode 100644 alignment_lab/analysis/tf_sf_backend.sh create mode 100644 alignment_lab/diagnostics/frf_vs_ftf_discrimination.py create mode 100644 alignment_lab/diagnostics/tf_batch_probe.py create mode 100644 alignment_lab/diagnostics/tf_cost.py create mode 100644 alignment_lab/diagnostics/tf_resolution_and_grid.py create mode 100644 alignment_lab/diagnostics/tf_sf_backend.py diff --git a/alignment_lab/analysis/ftf_disc_smoke.sh b/alignment_lab/analysis/ftf_disc_smoke.sh new file mode 100644 index 00000000..b0bda334 --- /dev/null +++ b/alignment_lab/analysis/ftf_disc_smoke.sh @@ -0,0 +1,18 @@ +#!/bin/bash +#SBATCH --job-name=ftfsmoke +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:30:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=32G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +export TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 +cd "$REPO" +"$PY" -u alignment_lab/diagnostics/frf_vs_ftf_discrimination.py \ + --pdb 1DAW --trials 1 --n-cand 4 --n-rotation-peaks 60 2>&1 | tail -30 +echo "RC=${PIPESTATUS[0]}" diff --git a/alignment_lab/analysis/ftf_discrimination.sh b/alignment_lab/analysis/ftf_discrimination.sh new file mode 100644 index 00000000..d851c5ce --- /dev/null +++ b/alignment_lab/analysis/ftf_discrimination.sh @@ -0,0 +1,25 @@ +#!/bin/bash +#SBATCH --job-name=ftfdisc +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-9 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +export TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 +# Concurrent array tasks otherwise poison a shared __pycache__ for numba. +export NUMBA_CACHE_DIR="/tmp/numba_${SLURM_ARRAY_JOB_ID}_${SLURM_ARRAY_TASK_ID}" +mkdir -p "$NUMBA_CACHE_DIR" +cd "$REPO" +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +P=${PDBS[$SLURM_ARRAY_TASK_ID]} +"$PY" -u alignment_lab/diagnostics/frf_vs_ftf_discrimination.py \ + --pdb "$P" --trials 3 --n-cand 25 --n-rotation-peaks 200 \ + --out-csv "alignment_lab/runs/ftf_disc_${SLURM_ARRAY_JOB_ID}.csv" 2>&1 +echo "RC=${PIPESTATUS[0]} pdb=$P" diff --git a/alignment_lab/analysis/ftf_early_stop.py b/alignment_lab/analysis/ftf_early_stop.py new file mode 100644 index 00000000..6b02e9c9 --- /dev/null +++ b/alignment_lab/analysis/ftf_early_stop.py @@ -0,0 +1,117 @@ +"""Would a sequential TFZ-gated search stop on the right orientation? + +The discrimination panel ranks candidates against *each other*, which needs all +of them placed. A guidance loop wants the opposite: walk the FRF peaks in +descending order, place one at a time, and stop as soon as one is convincing -- +paying for the whole list only on the cases that need it. + +That needs an **absolute** criterion, computed from a single candidate. Two are +available per placement and both are already in the code: + +``tfz`` the top translation peak measured against the spread of that + orientation's own translation map (``TranslationPeak.sigma``, Phaser's + TFZ). +``llgz`` the same for the LLG re-rank: best translation against the other 19. + +Simulates the loop over the recorded placements at a range of thresholds. A run +that never crosses the threshold falls back to the argmax over all candidates +placed, which is the no-early-stop behaviour -- so the cost of a threshold set +too high is wasted work, not a failure, and the two are reported separately. +""" + +from __future__ import annotations + +import argparse +import statistics as st +import sys +from collections import defaultdict + + +def load(paths): + cells = defaultdict(list) + for path in paths: + for line in open(path): + if not line.startswith("CAND"): + continue + d = dict(kv.split("=", 1) for kv in line.split() if "=" in kv) + cells[(d["pdb"], int(d["trial"]))].append( + dict(k=int(d["k"]), ang=float(d["ang"]), + truth=d["is_truth"] == "1", + tfz=float(d["tfz"]), llgz=float(d["llgz"]), + tf_llg=float(d["tf_llg"]), r=float(d["r"]))) + for v in cells.values(): + v.sort(key=lambda c: c["k"]) # FRF descending order + return cells + + +def simulate(cells, key, thr): + """Walk each cell in FRF order, stop at the first candidate over ``thr``.""" + hit = miss = exh_hit = exh_miss = 0 + placed, miss_ang = [], [] + for cs in cells.values(): + for i, c in enumerate(cs): + if c[key] >= thr: + placed.append(i + 1) + if c["truth"]: + hit += 1 + else: + miss += 1 + miss_ang.append(c["ang"]) + break + else: # never convinced: rank them all + placed.append(len(cs)) + best = max(cs, key=lambda c: c[key]) + if best["truth"]: + exh_hit += 1 + else: + exh_miss += 1 + miss_ang.append(best["ang"]) + n = len(cells) + return dict(thr=thr, n=n, hit=hit, miss=miss, exh_hit=exh_hit, + exh_miss=exh_miss, ok=hit + exh_hit, + med_placed=st.median(placed), mean_placed=sum(placed) / n, + med_miss_ang=st.median(miss_ang) if miss_ang else float("nan")) + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("logs", nargs="+") + ap.add_argument("--key", default="tfz", choices=["tfz", "llgz"]) + ap.add_argument("--thresholds", default="0,3,4,5,6,7,8,9,10,12,15,20,1e9") + args = ap.parse_args() + + cells = load(args.logs) + if not cells: + print("no CAND lines found", file=sys.stderr) + return 1 + n_cand = st.median([len(v) for v in cells.values()]) + print(f"# {len(cells)} cells, {n_cand:.0f} candidates each, key={args.key}") + print(f"# a cell is OK if the loop commits to a candidate within 8 deg\n") + print(f"{'thr':>6s} {'OK':>7s} {'stop_hit':>8s} {'stop_miss':>9s} " + f"{'exh_hit':>7s} {'exh_miss':>8s} {'med_n':>6s} {'mean_n':>7s} " + f"{'miss_ang':>8s}") + for t in [float(x) for x in args.thresholds.split(",")]: + r = simulate(cells, args.key, t) + label = "none" if t > 1e8 else f"{t:g}" + print(f"{label:>6s} {r['ok']:>4d}/{r['n']:<2d} {r['hit']:>8d} " + f"{r['miss']:>9d} {r['exh_hit']:>7d} {r['exh_miss']:>8d} " + f"{r['med_placed']:>6.0f} {r['mean_placed']:>7.1f} " + f"{r['med_miss_ang']:>8.1f}") + + print("\nper structure at the best threshold by (OK, then fewest placed):") + best = max((simulate(cells, args.key, t) + for t in [float(x) for x in args.thresholds.split(",")]), + key=lambda r: (r["ok"], -r["mean_placed"])) + print(f" thr={best['thr']:g}") + for p in sorted({k[0] for k in cells}): + sub = {k: v for k, v in cells.items() if k[0] == p} + r = simulate(sub, args.key, best["thr"]) + print(f" {p:6s} OK {r['ok']}/{r['n']} placed med={r['med_placed']:.0f} " + f"mean={r['mean_placed']:.1f}" + + ("" if r['med_miss_ang'] != r['med_miss_ang'] + else f" miss_ang={r['med_miss_ang']:.1f} deg")) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/analysis/ftf_running_null.py b/alignment_lab/analysis/ftf_running_null.py new file mode 100644 index 00000000..50211b4a --- /dev/null +++ b/alignment_lab/analysis/ftf_running_null.py @@ -0,0 +1,140 @@ +"""Early stopping needs a cross-candidate contrast, so build the null as you go. + +The per-candidate Z-scores fail as stopping criteria, and the separability table +says why: pooled over 750 placements, a wrong orientation's TFZ is *higher* than +truth's (median 3.15 vs 2.99). TFZ asks "is this translation better than the +other translations for this orientation" -- and a wrong orientation still has a +best translation that stands out of its own map. The contrast that carries the +signal is "is this orientation better than the other orientations", which no +single placement can answer. + +But a sequential loop does not need all 25 placements to answer it -- only +enough of them to know what a wrong answer looks like on this structure. So: +place ``--burn-in`` candidates unconditionally, estimate the null from them with +a median/MAD (robust, because truth is often among the first few and a mean +would be dragged by it), then commit to the first candidate -- burn-in included +-- that sits ``--thr`` robust sigmas above it. + +Sweeps threshold and burn-in over the recorded placements, in the FRF order a +live loop would walk. A cell that never crosses falls back to the argmax, which +is the no-early-stop behaviour, so an over-strict threshold costs placements +rather than answers. +""" + +from __future__ import annotations + +import argparse +import statistics as st +from collections import defaultdict + + +def load(paths): + cells = defaultdict(list) + for path in paths: + for line in open(path): + if not line.startswith("CAND"): + continue + d = dict(kv.split("=", 1) for kv in line.split() if "=" in kv) + cells[(d["pdb"], int(d["trial"]))].append( + dict(k=int(d["k"]), ang=float(d["ang"]), + truth=d["is_truth"] == "1", tf_corr=float(d["tf_corr"]), + tf_llg=float(d["tf_llg"]), tfz=float(d["tfz"]), + r=float(d["r"]))) + for v in cells.values(): + v.sort(key=lambda c: c["k"]) + return cells + + +def robust_z(x, ref, higher_is_better=True): + """``x`` in MAD-sigmas above the centre of ``ref``. Sign-normalised.""" + med = st.median(ref) + mad = st.median([abs(v - med) for v in ref]) + scale = 1.4826 * mad + if scale < 1e-30: + return float("inf") if x != med else 0.0 + z = (x - med) / scale + return z if higher_is_better else -z + + +def simulate(cells, key, thr, burn, hi=True): + hit = miss = exh_hit = exh_miss = 0 + placed, miss_ang = [], [] + for cs in cells.values(): + b = min(burn, len(cs)) + ref = [c[key] for c in cs[:b]] + stopped = None + # The burn-in candidates are tested too: truth is often among the first + # few, and a loop that could not commit to one it had already placed + # would pay the whole list on exactly the easy cases. + for i, c in enumerate(cs): + n_paid = max(b, i + 1) + if i >= b: + ref = [q[key] for q in cs[:i]] + if robust_z(c[key], ref, hi) >= thr: + stopped = (i, n_paid, c) + break + if stopped is None: + placed.append(len(cs)) + best = (max if hi else min)(cs, key=lambda c: c[key]) + if best["truth"]: + exh_hit += 1 + else: + exh_miss += 1 + miss_ang.append(best["ang"]) + else: + _, n_paid, c = stopped + placed.append(n_paid) + if c["truth"]: + hit += 1 + else: + miss += 1 + miss_ang.append(c["ang"]) + n = len(cells) + return dict(n=n, hit=hit, miss=miss, exh_hit=exh_hit, exh_miss=exh_miss, + ok=hit + exh_hit, med=st.median(placed), + mean=sum(placed) / n, miss_ang=(st.median(miss_ang) + if miss_ang else float("nan"))) + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("logs", nargs="+") + ap.add_argument("--key", default="tf_llg", + choices=["tf_llg", "tf_corr", "r", "tfz"]) + ap.add_argument("--burn-ins", default="3,5,8") + ap.add_argument("--thresholds", default="3,5,8,10,15,20,30") + args = ap.parse_args() + hi = args.key != "r" + + cells = load(args.logs) + print(f"# {len(cells)} cells x {st.median([len(v) for v in cells.values()]):.0f} " + f"candidates, key={args.key} ({'higher' if hi else 'lower'} is better)") + print(f"# OK = committed to a candidate within 8 deg; placements counted " + f"include the burn-in\n") + print(f"{'burn':>4s} {'thr':>4s} {'OK':>7s} {'stop_hit':>8s} {'stop_miss':>9s} " + f"{'exh_hit':>7s} {'exh_miss':>8s} {'med_n':>5s} {'mean_n':>6s} " + f"{'miss_ang':>8s}") + best = None + for burn in [int(x) for x in args.burn_ins.split(",")]: + for thr in [float(x) for x in args.thresholds.split(",")]: + r = simulate(cells, args.key, thr, burn, hi) + print(f"{burn:>4d} {thr:>4g} {r['ok']:>4d}/{r['n']:<2d} " + f"{r['hit']:>8d} {r['miss']:>9d} {r['exh_hit']:>7d} " + f"{r['exh_miss']:>8d} {r['med']:>5.0f} {r['mean']:>6.1f} " + f"{r['miss_ang']:>8.1f}") + if best is None or (r["ok"], -r["mean"]) > (best[0]["ok"], -best[0]["mean"]): + best = (r, burn, thr) + r, burn, thr = best + print(f"\nper structure at burn={burn} thr={thr:g}:") + for p in sorted({k[0] for k in cells}): + sub = {k: v for k, v in cells.items() if k[0] == p} + s = simulate(sub, args.key, thr, burn, hi) + print(f" {p:6s} OK {s['ok']}/{s['n']} placed med={s['med']:.0f} " + f"mean={s['mean']:.1f}" + + ("" if s["miss_ang"] != s["miss_ang"] + else f" miss_ang={s['miss_ang']:.0f} deg")) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/analysis/pose_arms.sh b/alignment_lab/analysis/pose_arms.sh new file mode 100644 index 00000000..a51156e3 --- /dev/null +++ b/alignment_lab/analysis/pose_arms.sh @@ -0,0 +1,31 @@ +#!/bin/bash +# End-to-end pose recovery, which is the only metric that is actually the +# deliverable. Rank is a proxy; this is not. +# +# Two questions in one array: does dropping the rescore cost anything, and is +# n_rotation_candidates=15 leaving coverage on the table? The FRF's worst truth +# rank over 100 cells is 21, so 15 carries 93% and 25 carries 100% -- but +# coverage is not recovery, and only this measures recovery. +#SBATCH --job-name=posearm +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=day +#SBATCH --time=08:00:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-9 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +for T in 0 1 2; do + "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial "$T" \ + --arms m_letf1,none --n-rotation-candidates 15 2>/dev/null | grep '^ROW ' + "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial "$T" \ + --arms none --n-rotation-candidates 25 2>/dev/null | grep '^ROW ' +done diff --git a/alignment_lab/analysis/tf_cost.sh b/alignment_lab/analysis/tf_cost.sh new file mode 100644 index 00000000..1bf6e66b --- /dev/null +++ b/alignment_lab/analysis/tf_cost.sh @@ -0,0 +1,19 @@ +#!/bin/bash +#SBATCH --job-name=tfcost +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=01:00:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +export TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 +cd "$REPO" +for P in 1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X; do + "$PY" -u alignment_lab/diagnostics/tf_cost.py --pdb "$P" --trial 0 \ + 2>&1 | grep -E "^#|^ROW|Error|Traceback" || echo "ROW pdb=$P FAILED" +done diff --git a/alignment_lab/analysis/tf_resolution_and_grid.sh b/alignment_lab/analysis/tf_resolution_and_grid.sh new file mode 100644 index 00000000..cf3b7ef6 --- /dev/null +++ b/alignment_lab/analysis/tf_resolution_and_grid.sh @@ -0,0 +1,19 @@ +#!/bin/bash +#SBATCH --job-name=tfgrid +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=96G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +cd "$REPO" +for P in 1DAW 3E98 3A5V 3GR5 1AK5 3K7M 2DQ6 4BX9 6G9X; do + "$PY" -u alignment_lab/diagnostics/tf_resolution_and_grid.py --pdb "$P" --threads 4 2>&1 \ + | grep -E "^#|^ROW|Traceback|Error" || echo "# $P FAILED" +done diff --git a/alignment_lab/analysis/tf_sf_backend.sh b/alignment_lab/analysis/tf_sf_backend.sh new file mode 100644 index 00000000..a71b325b --- /dev/null +++ b/alignment_lab/analysis/tf_sf_backend.sh @@ -0,0 +1,19 @@ +#!/bin/bash +#SBATCH --job-name=sfback +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 +cd "$REPO" +for P in 1DAW 3E98 3A5V 3GR5 1AK5 3K7M 2DQ6 4BX9 6G9X; do + "$PY" -u alignment_lab/diagnostics/tf_sf_backend.py --pdb "$P" 2>&1 \ + | grep -E "^#|^ROW|Error|Traceback" || echo "ROW pdb=$P FAILED" +done diff --git a/alignment_lab/diagnostics/frf_vs_ftf_discrimination.py b/alignment_lab/diagnostics/frf_vs_ftf_discrimination.py new file mode 100644 index 00000000..74723e8f --- /dev/null +++ b/alignment_lab/diagnostics/frf_vs_ftf_discrimination.py @@ -0,0 +1,234 @@ +"""Does the translation function rank FRF peaks better than the FRF's own score? + +The FRF is a 3-D check: it correlates Pattersons, so a wrong orientation that +happens to reproduce the intramolecular vector set scores well. The translation +function is a 6-D check -- it has to place the molecule against the *crystal*, +intermolecular contacts included -- so it should separate truth from a ghost by +much more. That is the standard argument for carrying many orientations into the +TF, and it is worth measuring before spending anything on making the TF fast: +if the TF ranks no better than the FRF, carrying 100 orientations buys nothing. + +Takes the top ``--n-cand`` raw FRF peaks (no ML rescore -- the rescore is a +separate, and separately measured, reordering), places each one, and reports +where truth lands under four orderings: + +``frf`` the FRF's own score, i.e. the baseline +``tf_corr`` top translation peak of the Crowther-Blow amplitude correlation +``tf_llg`` the same peaks re-ranked by the shared-sigma_A Rice/Woolfson LLG +``r`` analytic-scale R at the locally refined t -- what the pipeline + actually ranks by today + +Rank alone understates the question, so each ordering also gets a separation +``z = (score_truth - mean_others) / std_others``: rank 0 by a hair and rank 0 by +five sigma are different claims about discrimination, and only the second one +justifies widening the funnel. + +The placement path is the production one -- ``_make_rotated`` then the same +``precompute_G`` / ``amplitude_translation_search`` / ``local_translation_refine`` +calls ``_placement_for_candidate`` makes -- with ``do_joint_refine`` off, since +the rigid-body polish is a later stage and would confound the TF's own ranking. +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import numpy as np +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, ResultWriter, rotated_case, seed_for, # noqa: E402 + symmetry_orbit) +from lab.truth import angle_to_orbit # noqa: E402 + + +def _rank_of_truth(scores, truth_mask, higher_is_better=True): + """Rank of the best-placed *correct* candidate under this ordering. + + Correct means within the angular threshold, and several candidates can be: + the peak list carries near-duplicates and symmetry mates. Any of them is a + solution, so the rank that matters is the first one to appear -- the same + definition :func:`orbit_rank` uses, and the one the pipeline behaves by. + """ + s = np.asarray(scores, dtype=float) + order = np.argsort(-s if higher_is_better else s, kind="stable") + return int(next(i for i, j in enumerate(order) if truth_mask[j])) + + +def _separation(scores, truth_mask, higher_is_better=True): + """Best correct candidate's score in sigmas above the *wrong* ones. + + Negated for lower-is-better scores so a larger number always means better + discrimination, whichever direction the score runs. Every within-threshold + candidate is held out of the reference pool: leaving a second copy of the + answer in it would inflate the pool's mean and understate the separation. + """ + s = np.asarray(scores, dtype=float) + m = np.asarray(truth_mask, dtype=bool) + best = float(s[m].max() if higher_is_better else s[m].min()) + others = s[~m] + if others.size < 2: + return float("nan") + sd = float(others.std(ddof=1)) + if sd < 1e-30: + return float("nan") + z = (best - float(others.mean())) / sd + return z if higher_is_better else -z + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) + ap.add_argument("--trials", type=int, default=3) + ap.add_argument("--n-cand", type=int, default=25, + help="FRF peaks placed. Truth's worst FRF rank over the " + "100-cell panel was 21, so 25 covers every case that " + "the rotation function gets right at all.") + ap.add_argument("--n-rotation-peaks", type=int, default=200) + ap.add_argument("--thr-deg", type=float, default=8.0) + ap.add_argument("--out-csv", default=None) + args = ap.parse_args() + + from torchref.experimental.alignment.align import ( + _DirectModelEvaluator, _prepare_frf_inputs, + ) + from torchref.experimental.alignment.frf.rotation_utils import ( + rotation_matrix_from_edmonds_euler, + ) + from torchref.experimental.alignment.pipeline import ( + MolecularReplacementPipeline, + ) + from torchref.experimental.alignment.translation import ( + amplitude_translation_search, local_translation_refine, + precompute_G_for_rotation, + ) + + writer = None + if args.out_csv: + writer = ResultWriter( + args.out_csv, "frf_vs_ftf", + extra_fields=("n_cand", "truth_found", "rank_frf", "rank_tf_corr", + "rank_tf_llg", "rank_r", "z_frf", "z_tf_corr", + "z_tf_llg", "z_r", "seconds"), + ) + + for trial in range(args.trials): + seed = seed_for(args.pdb, trial) + model, data, R_true = rotated_case(args.pdb, seed) + t0 = time.time() + + pipe = MolecularReplacementPipeline( + data, model, verbose=0, + rescore_engine="none", subpeak_refine=False, + n_rotation_peaks=args.n_rotation_peaks, + n_rotation_candidates=args.n_cand, + do_joint_refine=False, dense_rotation_refine=False, + use_llg_tf=False, + ) + frf = _prepare_frf_inputs( + model, data, d_min=pipe.d_min, d_max=pipe.d_max, + n_shells=pipe.n_shells, ll_padding_factor=pipe.ll_padding_factor, + ll_max_res_A=pipe.ll_max_res_A, verbose=0, + ) + pipe._frf = frf + peaks = pipe._rotation_candidates(frf)[: args.n_cand] + pipe._prepare_translation_arrays() + + orbit = symmetry_orbit( + R_true, data.spacegroup.matrices.to(torch.float64).cpu(), + side="left", frame="cart", + reciprocal_basis=data.cell.reciprocal_basis_matrix.to( + torch.float64).cpu(), + ) + eye3 = pipe._eye3 + rows = [] + for k, p in enumerate(peaks): + ang = angle_to_orbit( + rotation_matrix_from_edmonds_euler(p.alpha, p.beta, p.gamma), + orbit, + ) + rot = pipe._make_rotated(p)[0] + rot.spacegroup = data.spacegroup.hm + p1 = rot.copy() + p1.spacegroup = "P 1" + ev = _DirectModelEvaluator(p1) + G, h_R = precompute_G_for_rotation( + ev, eye3, pipe._hkl_keep, data.spacegroup, data.cell) + _, _, tp = amplitude_translation_search( + F_obs=pipe._F_obs_amp, interpolator=ev, R_rotation=eye3, + hkl=pipe._hkl_keep, spacegroup=data.spacegroup, + real_cell=data.cell, grid_steps=pipe.translation_grid_steps, + n_peaks=pipe.n_translation_peaks, cluster_radius=0.05, + precomputed_G=G, precomputed_h_R=h_R) + tf_corr = float(tp[0].score) + # Phaser's TFZ: the top translation peak measured against the + # spread of *this orientation's own* translation map. Unlike the + # cross-candidate z reported below it needs no other candidate, so + # it is the only score here that can stop a sequential search. + tfz = float(tp[0].sigma) + llg_peaks = pipe._llg_tf_rescore(tp, G, h_R) + tf_llg = float(llg_peaks[0].score) + lv = np.array([q.score for q in llg_peaks], dtype=float) + llgz = (float((lv[0] - lv[1:].mean()) / lv[1:].std(ddof=1)) + if lv.size > 2 and lv[1:].std(ddof=1) > 1e-30 + else float("nan")) + r_best = float("inf") + for cand in llg_peaks[: pipe.n_translation_candidates]: + _, r_a = local_translation_refine( + F_obs=pipe._F_obs_amp, interpolator=ev, R_rotation=eye3, + hkl=pipe._hkl_keep, spacegroup=data.spacegroup, + real_cell=data.cell, + t_init=torch.as_tensor(cand.translation, + dtype=torch.float64), + radius=0.06, grid_steps=13, n_refinement_passes=1, + precomputed_G=G, precomputed_h_R=h_R) + r_best = min(r_best, r_a) + rows.append(dict(k=k, ang=ang, frf=float(p.score), + tf_corr=tf_corr, tf_llg=tf_llg, r=r_best, + tfz=tfz, llgz=llgz)) + print(f"CAND pdb={args.pdb} trial={trial} k={k} ang={ang:.3f} " + f"is_truth={int(ang <= args.thr_deg)} frf={p.score:.4f} " + f"tf_corr={tf_corr:.5f} tfz={tfz:.3f} tf_llg={tf_llg:.2f} " + f"llgz={llgz:.3f} r={r_best:.5f}", flush=True) + + secs = time.time() - t0 + truth = [r for r in rows if r["ang"] <= args.thr_deg] + if not truth: + print(f"ROW pdb={args.pdb} trial={trial} truth_found=0 " + f"best_ang={min(r['ang'] for r in rows):.2f} " + f"seconds={secs:.1f}", flush=True) + continue + tmask = [r["ang"] <= args.thr_deg for r in rows] + ti = rows.index(min(truth, key=lambda r: r["ang"])) + cols = {"frf": True, "tf_corr": True, "tf_llg": True, "r": False} + ranks = {c: _rank_of_truth([r[c] for r in rows], tmask, hi) + for c, hi in cols.items()} + zs = {c: _separation([r[c] for r in rows], tmask, hi) + for c, hi in cols.items()} + print(f"ROW pdb={args.pdb} trial={trial} seed={seed} truth_found=1 " + f"n_cand={len(rows)} n_truth={sum(tmask)} " + + " ".join(f"rank_{c}={ranks[c]}" for c in cols) + + " " + " ".join(f"z_{c}={zs[c]:.2f}" for c in cols) + + f" seconds={secs:.1f}", flush=True) + if writer: + writer.write( + pdb=args.pdb, seed=seed, trial=trial, + spacegroup=str(data.spacegroup), + n_ops=int(data.spacegroup.matrices.shape[0]), + truth_rank=ranks["frf"], truth_angle_deg=round(rows[ti]["ang"], 3), + orbit_side="left", orbit_frame="cart", lmax_cap="", + d_min=pipe.d_min, d_max=pipe.d_max, device="cpu", + n_cand=len(rows), truth_found=1, + **{f"rank_{c}": ranks[c] for c in cols}, + **{f"z_{c}": round(zs[c], 4) for c in cols}, + seconds=round(secs, 1)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py index bf517a7d..6087be1e 100644 --- a/alignment_lab/diagnostics/pose_recovery.py +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -138,6 +138,10 @@ def main() -> int: resid, err = float("nan"), f"{type(exc).__name__}: {exc}" secs = time.time() - t0 ok = (resid == resid) and resid <= args.success_deg + print(f"ROW {arm} {args.pdb} trial={args.trial} " + f"n_cand={args.n_rotation_candidates} " + f"resid={resid:.3f} ok={int(bool(ok))} seconds={secs:.1f}", + flush=True) print(f" {arm:16s} {resid:10.2f} {('yes' if ok else 'NO'):>4s} {secs:9.1f}" + (f" {err}" if err else "")) if writer: diff --git a/alignment_lab/diagnostics/tf_batch_probe.py b/alignment_lab/diagnostics/tf_batch_probe.py new file mode 100644 index 00000000..03f82827 --- /dev/null +++ b/alignment_lab/diagnostics/tf_batch_probe.py @@ -0,0 +1,152 @@ +"""Can the translation stage share one molecular transform across orientations? + +``tf_cost.py`` says the per-candidate placement is ~95% ``precompute_G``: a full +structure-factor evaluation of a *re-rotated* model at ``S*N`` Miller indices, +paid again for every orientation. The Crowther-Blow accumulation and the local +refine that follow it are milliseconds. + +But the rotation does not have to live in the coordinates. ``F(h; R x) = +F(R^T h; x)``, which is exactly what :class:`LattmanLoveInterpolator` is for and +what the *rotation* search already uses -- and its ``evaluate`` is already +batched over ``R``. So ``G`` for M orientations could be one dense grid plus +``M*S*N`` trilinear lookups instead of M structure-factor calculations. + +Two things have to hold and neither is obvious: + +* **The phase must survive trilinear interpolation.** The rotation search only + ever reads ``|F|``; the translation function reads ``arg F``, and the class's + own docstring warns that complex interpolation is only safe on a + well-oversampled grid. Measured here as the agreement of the *translation + peaks*, not of ``F`` -- a phase error that does not move the peak does not + matter. +* **It must actually be faster**, including the one-off grid build, at the + orientation counts we would carry. + +Reports both against the exact per-rotation path. +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) + ap.add_argument("--n-rot", type=int, default=8, help="orientations to time") + ap.add_argument("--d-min", type=float, default=4.0) + ap.add_argument("--d-max", type=float, default=15.0) + ap.add_argument("--max-res", type=float, default=3.0) + ap.add_argument("--padding", type=float, default=2.0) + ap.add_argument("--grid-steps", type=int, default=16) + args = ap.parse_args() + + from torchref.experimental.alignment.align import _DirectModelEvaluator + from torchref.experimental.alignment.lattman_love import LattmanLoveInterpolator + from torchref.experimental.alignment.translation import ( + amplitude_translation_search, precompute_G_for_rotation, + ) + + model, data = load_case(args.pdb) + rec = data.cell.reciprocal_basis_matrix.to(torch.float64) + mask = data.get_valid_mask() + s_mag = (data.hkl.to(torch.float64) @ rec).norm(dim=-1) + mask = mask & (s_mag >= 1.0 / args.d_max) & (s_mag <= 1.0 / args.d_min) + hkl = data.hkl[mask] + F_obs = data.F[mask].abs().to(torch.float64) + S = int(data.spacegroup.matrices.shape[0]) + N = int(hkl.shape[0]) + print(f"# {args.pdb} sg={data.spacegroup.hm} S={S} N={N} " + f"atoms={model.xyz().shape[0]} max_res={args.max_res} " + f"padding={args.padding}", flush=True) + + base = model.copy() + base.spacegroup = "P 1" + origin = torch.zeros(3, dtype=base.xyz().dtype) + eye3 = torch.eye(3, dtype=torch.float64) + + t0 = time.perf_counter() + ll = LattmanLoveInterpolator(base, padding_factor=args.padding, + max_res_A=args.max_res, verbose=0) + t_grid = time.perf_counter() - t0 + print(f"# dense grid build {t_grid:.2f} s shape=" + f"{tuple(ll.reciprocal_grid.shape)}", flush=True) + + # h_R is a function of hkl and the space group only -- not of the candidate + # rotation -- so the index side of the whole stage is shared. + sym_R = data.spacegroup.matrices.to(torch.float64) + h_R = torch.einsum("ne,ied->ind", hkl.to(torch.float64), sym_R) + h_R_flat = h_R.reshape(-1, 3) + + rots = [random_rotation(seed_for(args.pdb, 0) + 97 * k) + for k in range(args.n_rot)] + + t_direct = t_interp = 0.0 + for k, R in enumerate(rots): + rot = base.copy().rotate(R.to(base.dtype_float), center=origin) + ev = _DirectModelEvaluator(rot) + t0 = time.perf_counter() + G_d, h_R_d = precompute_G_for_rotation(ev, eye3, hkl, + data.spacegroup, data.cell) + t_direct += time.perf_counter() - t0 + + t0 = time.perf_counter() + F_i = ll.evaluate(R.to(torch.float32), h_R_flat, data.cell, + return_amplitude=False).reshape(S, N) + phase_sym = torch.exp(2j * torch.pi * torch.einsum( + "ne,ie->in", hkl.to(torch.float64), + data.spacegroup.translations.to(torch.float64), + ).to(torch.complex128)) + G_i = F_i.to(torch.complex128) * phase_sym + t_interp += time.perf_counter() - t0 + + if k == 0: + a, b = G_d.reshape(-1), G_i.reshape(-1) + coh = float((a.conj() * b).sum().abs() + / (a.abs().norm() * b.abs().norm()).clamp(min=1e-30)) + amp = float(torch.corrcoef(torch.stack( + [a.abs(), b.abs()]).to(torch.float64))[0, 1]) + print(f"# G agreement: complex coherence={coh:.4f} " + f"|F| corr={amp:.4f}", flush=True) + + tf = lambda G: amplitude_translation_search( + F_obs=F_obs, interpolator=ev, R_rotation=eye3, hkl=hkl, + spacegroup=data.spacegroup, real_cell=data.cell, + grid_steps=args.grid_steps, n_peaks=5, + precomputed_G=G, precomputed_h_R=h_R_d)[2] + pd, pi = tf(G_d), tf(G_i) + dt = min(float(torch.tensor( + ((torch.as_tensor(pd[0].translation) + - torch.as_tensor(pi[j].translation) + 0.5) % 1.0 - 0.5).norm())) + for j in range(len(pi))) + dt_top = float(torch.tensor( + ((torch.as_tensor(pd[0].translation) + - torch.as_tensor(pi[0].translation) + 0.5) % 1.0 - 0.5).norm())) + print(f"ROW pdb={args.pdb} rot={k} dt_top={dt_top:.4f} " + f"dt_best_of_5={dt:.4f} " + f"score_direct={pd[0].score:.4f} score_interp={pi[0].score:.4f}", + flush=True) + + M = args.n_rot + print(f"SUM pdb={args.pdb} S={S} N={N} n_rot={M} " + f"t_direct_per_rot={t_direct / M:.3f} " + f"t_interp_per_rot={t_interp / M:.4f} " + f"t_grid={t_grid:.2f} " + f"breakeven_rots={t_grid / max(t_direct / M - t_interp / M, 1e-9):.1f} " + f"speedup_at_100={100 * (t_direct / M) / (t_grid + 100 * t_interp / M):.1f}", + flush=True) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/tf_cost.py b/alignment_lab/diagnostics/tf_cost.py new file mode 100644 index 00000000..3c92bfcf --- /dev/null +++ b/alignment_lab/diagnostics/tf_cost.py @@ -0,0 +1,105 @@ +"""Where the translation stage actually spends its time, per rotation candidate. + +The pipeline carries ``n_rotation_candidates`` orientations through +:func:`_placement_for_candidate` one at a time in a Python loop, and every +orientation repays the whole stage. Before batching it over 100 orientations we +need to know which part of it is the cost: the structure-factor evaluation, the +Crowther-Blow accumulation, the LLG re-rank, or the local refine. + +Reports the per-candidate breakdown alongside the problem geometry (``N`` +reflections, ``S`` sym-ops, grid), because the four stages scale differently -- +``O(n_atoms*S*N)``, ``O(S^2*N)``, ``O(K*S*N)`` and ``O(S^2*N)`` respectively -- +and which one dominates is a property of the structure, not of the code. +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 + + +def _time(fn, repeats=1): + out = None + t0 = time.perf_counter() + for _ in range(repeats): + out = fn() + return (time.perf_counter() - t0) / repeats, out + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) + ap.add_argument("--trial", type=int, default=0) + ap.add_argument("--d-min", type=float, default=4.0) + ap.add_argument("--d-max", type=float, default=15.0) + ap.add_argument("--grid-steps", type=int, default=16) + ap.add_argument("--n-peaks", type=int, default=20) + args = ap.parse_args() + + from torchref.experimental.alignment.align import _DirectModelEvaluator + from torchref.experimental.alignment.translation import ( + amplitude_translation_search, local_translation_refine, + precompute_G_for_rotation, + ) + + seed = seed_for(args.pdb, args.trial) + model, data = load_case(args.pdb) + R_true = random_rotation(seed) + rot = model.copy() + rot.spacegroup = "P 1" + rot = rot.rotate(R_true.to(model.dtype_float), center=model.xyz().mean(0)) + + # Same masking the pipeline's _prepare_translation_arrays does. + mask = data.get_valid_mask() + d = 1.0 / (data.hkl.to(torch.float64) + @ data.cell.reciprocal_basis_matrix.to(torch.float64) + ).norm(dim=-1).clamp(min=1e-9) + mask = mask & (d >= args.d_min) & (d <= args.d_max) + hkl = data.hkl[mask] + F_obs = data.F[mask].abs().to(torch.float64) + S = int(data.spacegroup.matrices.shape[0]) + N = int(hkl.shape[0]) + n_at = int(rot.xyz().shape[0]) + print(f"# {args.pdb} sg={data.spacegroup.hm} S={S} N={N} atoms={n_at} " + f"grid={args.grid_steps} seed={seed}", flush=True) + + ev = _DirectModelEvaluator(rot) + eye3 = torch.eye(3, dtype=torch.float64) + + t_G, (G, h_R) = _time(lambda: precompute_G_for_rotation( + ev, eye3, hkl, data.spacegroup, data.cell)) + t_tf, (_, _, peaks) = _time(lambda: amplitude_translation_search( + F_obs=F_obs, interpolator=ev, R_rotation=eye3, hkl=hkl, + spacegroup=data.spacegroup, real_cell=data.cell, + grid_steps=args.grid_steps, n_peaks=args.n_peaks, + precomputed_G=G, precomputed_h_R=h_R)) + t_ref, _ = _time(lambda: local_translation_refine( + F_obs=F_obs, interpolator=ev, R_rotation=eye3, hkl=hkl, + spacegroup=data.spacegroup, real_cell=data.cell, + t_init=torch.as_tensor(peaks[0].translation, dtype=torch.float64), + radius=0.06, grid_steps=13, n_refinement_passes=1, + precomputed_G=G, precomputed_h_R=h_R)) + + # The (S,N) working set the Crowther-Blow loop materialises, and the + # (K,S,N) one the LLG re-rank does, in complex128. + mb = lambda *shape: 16.0 * float(torch.tensor(shape).prod()) / 2**20 + print(f"ROW pdb={args.pdb} sg={data.spacegroup.hm} S={S} N={N} atoms={n_at} " + f"t_G={t_G:.3f} t_tf={t_tf:.3f} t_refine={t_ref:.3f} " + f"t_place={t_G + t_tf + 3 * t_ref:.3f} " + f"work_S2N={S * S * N / 1e6:.1f}M " + f"mem_SN_MB={mb(S, N):.0f} mem_KSN_MB={mb(args.n_peaks, S, N):.0f}", + flush=True) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/tf_resolution_and_grid.py b/alignment_lab/diagnostics/tf_resolution_and_grid.py new file mode 100644 index 00000000..ea4ecb55 --- /dev/null +++ b/alignment_lab/diagnostics/tf_resolution_and_grid.py @@ -0,0 +1,99 @@ +"""What resolution does the translation stage actually run at, and what grid does it need? + +Two things to pin down, and the first invalidates my earlier numbers. + +``_prepare_translation_arrays`` (``pipeline.py:638``) is documented as +"Resolution/validity-masked" but applies **only** ``data.get_valid_mask()`` -- +there is no resolution cut. So the translation search runs at the data's full +resolution while the FRF runs at ``[d_max, d_min] = [15, 4]``. Any timing taken +on a 4 A subset is measuring a smaller problem than the pipeline solves. + +Second, the FFT grid is sized from ``ModelFT.max_res``, which defaults to 1.0 A +and which this stage never sets. Whether that is oversized depends entirely on +the answer to the first question: against 4 A data it is 64x too many voxels, +against 2 A data it is the ~2x oversampling one would ask for anyway. + +Reports the real N, the real ``d_min``, and times the structure-factor call on +grids sized at a range of ``max_res``, each checked for coherence against the +current 1.0 A grid -- because undersampling an FFT does not fail, it just +quietly returns different structure factors. +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 + + +def _time(fn, repeats=3): + fn() + t0 = time.perf_counter() + for _ in range(repeats): + out = fn() + return (time.perf_counter() - t0) / repeats, out + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) + ap.add_argument("--threads", type=int, default=4) + ap.add_argument("--oversampling", default="1.0,1.5,2.0,3.0") + args = ap.parse_args() + torch.set_num_threads(args.threads) + + model, data = load_case(args.pdb) + rec = data.cell.reciprocal_basis_matrix.to(torch.float64) + # Exactly what the pipeline masks with -- no resolution window. + mask = data.get_valid_mask() + hkl = data.hkl[mask] + s = (hkl.to(torch.float64) @ rec).norm(dim=-1) + d_min = float(1.0 / s.max()) + d_max = float(1.0 / s.min().clamp(min=1e-9)) + sym_R = data.spacegroup.matrices.to(torch.float64) + hkl_SN = torch.einsum("ne,ied->ind", hkl.to(torch.float64), sym_R + ).reshape(-1, 3).round().to(torch.int64) + S, N = int(sym_R.shape[0]), int(hkl.shape[0]) + + # For contrast: what the FRF's own window would leave. + in_frf = ((s >= 1.0 / 15.0) & (s <= 1.0 / 4.0)).sum().item() + print(f"# {args.pdb} sg={data.spacegroup.hm} S={S} " + f"N_pipeline={N} N_in_4to15A={in_frf} " + f"d_min={d_min:.2f} d_max={d_max:.1f} atoms={model.xyz().shape[0]} " + f"S*N={S*N} threads={args.threads}", flush=True) + + rot = model.copy() + rot.spacegroup = "P 1" + rot = rot.rotate(random_rotation(seed_for(args.pdb, 0)).to(model.dtype_float), + center=torch.zeros(3, dtype=model.xyz().dtype)) + + ref_sf = None + for over in [1.0] + [float(x) for x in args.oversampling.split(",")]: + m = rot.copy() + m.max_res = 1.0 if ref_sf is None else d_min / over + m.spacegroup = "P 1" + t, sf = _time(lambda: (m.reset_cache(), m(hkl_SN))[1]) + if ref_sf is None: + ref_sf, tag = sf.to(torch.complex128), "current(1.0A)" + coh = 1.0 + else: + x = sf.to(torch.complex128) + coh = float((ref_sf.conj() * x).sum().abs() + / (ref_sf.abs().norm() * x.abs().norm()).clamp(min=1e-30)) + tag = f"d_min/{over:g}" + print(f"ROW pdb={args.pdb} arm={tag} max_res={m.max_res:.3f} " + f"grid={tuple(int(v) for v in m.fft.gridsize)} " + f"t={t*1e3:.0f}ms coh={coh:.6f}", flush=True) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/tf_sf_backend.py b/alignment_lab/diagnostics/tf_sf_backend.py new file mode 100644 index 00000000..aa2ae116 --- /dev/null +++ b/alignment_lab/diagnostics/tf_sf_backend.py @@ -0,0 +1,112 @@ +"""Why does one orientation's structure-factor evaluation cost ~1 s? + +``tf_cost.py`` put 93-97% of the translation stage in ``precompute_G``, which is +a single ``ModelFT.__call__`` at ``S*N`` Miller indices. That call goes through +``SfFFT`` (``model_ft.py:779``) -- splat the atoms onto a real-space grid, FFT +the box, sample the result. The grid is sized by the crystal cell and the +resolution, and it is built whether you wanted 12000 reflections or 12 million. + +The translation search wants a sparse, fixed set: ``S*N`` is 13k-160k here, +against 5-20k atoms. That is the regime direct summation is for, and +:class:`SfDS` -- same ``compute_structure_factors`` signature, no grid -- is +already in the tree but is not what ``ModelFT`` dispatches to. + +Times both on identical inputs and checks they agree, at the thread counts a +production run would see. Also separates first call from repeat: ``ModelFT`` is +rebuilt per orientation, so anything amortised across calls is paid in full by +the translation loop and has to be counted as setup, not as throughput. +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 + + +def _time(fn, repeats=3): + fn() + t0 = time.perf_counter() + for _ in range(repeats): + out = fn() + return (time.perf_counter() - t0) / repeats, out + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) + ap.add_argument("--d-min", type=float, default=4.0) + ap.add_argument("--d-max", type=float, default=15.0) + ap.add_argument("--threads", default="4,8") + args = ap.parse_args() + + from torchref.model.sf_ds import SfDS + + model, data = load_case(args.pdb) + rec = data.cell.reciprocal_basis_matrix.to(torch.float64) + s = (data.hkl.to(torch.float64) @ rec).norm(dim=-1) + keep = data.get_valid_mask() & (s >= 1.0 / args.d_max) & (s <= 1.0 / args.d_min) + hkl = data.hkl[keep] + sym_R = data.spacegroup.matrices.to(torch.float64) + h_R = torch.einsum("ne,ied->ind", hkl.to(torch.float64), sym_R) + hkl_SN = h_R.reshape(-1, 3).round().to(torch.int64) + S, N = int(sym_R.shape[0]), int(hkl.shape[0]) + + rot = model.copy() + rot.spacegroup = "P 1" + rot = rot.rotate(random_rotation(seed_for(args.pdb, 0)).to(model.dtype_float), + center=torch.zeros(3, dtype=model.xyz().dtype)) + n_at = int(rot.xyz().shape[0]) + print(f"# {args.pdb} sg={data.spacegroup.hm} S={S} N={N} S*N={S*N} " + f"atoms={n_at} model_max_res={rot.max_res}", flush=True) + + ds = SfDS(cell=rot.cell, spacegroup="P 1", dtype_float=rot.dtype_float, + device=rot.xyz().device) + iso, aniso = rot.get_iso(), rot.get_aniso() + + def fft_call(m): + m.reset_cache() # the loop gets a fresh model per candidate + return m(hkl_SN) + + # The grid is sized by the model's ``max_res``, which ModelFT defaults to + # 1.0 A. The translation search runs at d_min, so the default asks for + # (d_min/1.0)^3 times the voxels it needs. Both the FRF's dense calc + # (dense_calc.py:73) and the rigid-body stage (rigid_body.py:120) set this; + # the translation stage does not. + coarse = rot.copy() + coarse.max_res = float(args.d_min) + coarse.spacegroup = "P 1" + + for nt in [int(x) for x in args.threads.split(",")]: + torch.set_num_threads(nt) + t_fft, sf_fft = _time(lambda: fft_call(rot)) + t_coarse, sf_coarse = _time(lambda: fft_call(coarse)) + t_ds, (sf_ds, _) = _time(lambda: ds.compute_structure_factors( + hkl_SN, *iso, *aniso, apply_symmetry=True)) + + a = sf_fft.to(torch.complex128) + agree = lambda x: float((a.conj() * x.to(torch.complex128)).sum().abs() + / (a.abs().norm() + * x.abs().norm()).clamp(min=1e-30)) + print(f"ROW pdb={args.pdb} threads={nt} S={S} N={N} SN={S*N} " + f"atoms={n_at} grid_1A={tuple(int(v) for v in rot.fft.gridsize)} " + f"grid_dmin={tuple(int(v) for v in coarse.fft.gridsize)} " + f"t_fft_1A={t_fft*1e3:.0f}ms t_fft_dmin={t_coarse*1e3:.0f}ms " + f"t_ds={t_ds*1e3:.0f}ms " + f"gain_grid={t_fft/max(t_coarse,1e-9):.1f}x " + f"gain_ds={t_fft/max(t_ds,1e-9):.1f}x " + f"coh_dmin={agree(sf_coarse):.6f} coh_ds={agree(sf_ds):.6f}", + flush=True) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) From 453f418a009a1ae0b1aaa7f926787eaaa9568e69 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 31 Aug 2026 13:59:37 +0200 Subject: [PATCH 107/250] Drop four alignment modules nothing reaches transform.py, wigner.py, clashscore.py and sampling.py are unreachable from align_model_to_data. Together they are 2154 lines kept loadable only by the package __init__, which eagerly re-exports everything. transform.py's quaternion algebra was superseded by frf/rotation_utils.py, and the root wigner.py by frf/wigner_d.py -- which builds its small-d blocks from a J_y eigendecomposition rather than half-angle power tables, and carries a duplicate AdaptiveRotationFunction that shadowed frf/types.py. clashscore.py had no importer and no test; MRSolution advertised a clash_score populated by a clash_filter that was never written, so that field goes too. The two small-d identity tests move onto the surviving implementation: d(pi/2) reflection in n, and the beta-reflection identity, now asserted against frf.wigner_d._wigner_d_blocks. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- tests/unit/alignment/test_wigner.py | 159 --- tests/unit/frf_separate/test_invariants.py | 20 +- torchref/experimental/alignment/__init__.py | 59 -- torchref/experimental/alignment/clashscore.py | 389 -------- .../experimental/alignment/frf/wigner_d.py | 4 +- torchref/experimental/alignment/pipeline.py | 3 - torchref/experimental/alignment/sampling.py | 338 ------- torchref/experimental/alignment/transform.py | 921 ------------------ torchref/experimental/alignment/wigner.py | 506 ---------- 9 files changed, 14 insertions(+), 2385 deletions(-) delete mode 100644 tests/unit/alignment/test_wigner.py delete mode 100644 torchref/experimental/alignment/clashscore.py delete mode 100644 torchref/experimental/alignment/sampling.py delete mode 100644 torchref/experimental/alignment/transform.py delete mode 100644 torchref/experimental/alignment/wigner.py diff --git a/tests/unit/alignment/test_wigner.py b/tests/unit/alignment/test_wigner.py deleted file mode 100644 index 24d08d13..00000000 --- a/tests/unit/alignment/test_wigner.py +++ /dev/null @@ -1,159 +0,0 @@ -""" -Unit tests for torchref.experimental.alignment.wigner. - -Conventions verified: -- D^l_{m,n}(α,β,γ) = e^{-imα} d^l_{m,n}(β) e^{-inγ} (Edmonds) -- d^l_{m,n}(0) = δ_{m,n}; d^l_{m,n}(π) = (-1)^{l+m} δ_{m,-n}. -- Unitarity: D^l D^l† = I for any (α,β,γ). -- Composition (sanity): the inverse FFT path agrees with the pointwise path. -""" -import math - -import numpy as np -import pytest -import torch - -from torchref.experimental.alignment.wigner import ( - small_d_block, - small_d_packed, - wigner_D_pointwise, - evaluate_rotation_function_grid, - evaluate_rotation_function_pointwise, -) - - -@pytest.mark.parametrize("l", [0, 1, 2, 3, 5, 8]) -def test_small_d_identity_at_zero(l): - """d^l_{m,n}(0) = δ_{m,n}.""" - beta = torch.tensor([0.0], dtype=torch.float64) - d = small_d_block(l, beta)[0] # (2l+1, 2l+1) - expected = torch.eye(2 * l + 1, dtype=torch.float64) - np.testing.assert_allclose(d.numpy(), expected.numpy(), atol=1e-12) - - -@pytest.mark.parametrize("l", [0, 1, 2, 3, 5, 8]) -def test_small_d_at_pi(l): - """d^l_{m,n}(π) = (-1)^{l+m} δ_{m,-n}.""" - beta = torch.tensor([math.pi], dtype=torch.float64) - d = small_d_block(l, beta)[0] # (2l+1, 2l+1) - size = 2 * l + 1 - expected = torch.zeros(size, size, dtype=torch.float64) - for m_idx in range(size): - m = m_idx - l - expected[m_idx, -m_idx - 1 + size] = (-1.0) ** (l + m) # n = -m → index size-1 - m_idx - np.testing.assert_allclose(d.numpy(), expected.numpy(), atol=1e-10) - - -@pytest.mark.parametrize("l", [1, 2, 4, 6]) -def test_small_d_unitary(l): - """d^l(β) is real-orthogonal: d^T d = I.""" - beta = torch.tensor([0.3, 1.1, 2.5], dtype=torch.float64) - d = small_d_block(l, beta) # (3, 2l+1, 2l+1) - for k in range(3): - dk = d[k] - prod = dk.T @ dk - np.testing.assert_allclose(prod.numpy(), np.eye(2 * l + 1), atol=1e-10, - err_msg=f"l={l} β={beta[k]:.3f}: d^T d ≠ I") - - -@pytest.mark.parametrize("L", [3, 5]) -def test_wigner_D_unitary(L): - """D^l(R) is unitary for any (α,β,γ).""" - torch.manual_seed(0) - n = 4 - alpha = torch.rand(n, dtype=torch.float64) * 2 * math.pi - beta = torch.rand(n, dtype=torch.float64) * math.pi - gamma = torch.rand(n, dtype=torch.float64) * 2 * math.pi - D = wigner_D_pointwise(alpha, beta, gamma, L) # (n, L, 2L-1, 2L-1) - - for k in range(n): - for l in range(L): - sl = slice(L - 1 - l, L - 1 + l + 1) - Dl = D[k, l, sl, sl] - prod = Dl @ Dl.conj().transpose(-1, -2) - eye = torch.eye(2 * l + 1, dtype=Dl.dtype) - np.testing.assert_allclose(prod.numpy(), eye.numpy(), atol=1e-10, - err_msg=f"l={l} not unitary at k={k}") - - -def test_wigner_D_diagonal_for_pure_z_rotation(): - """For β=γ=0, D^l_{m,n}(α,0,0) = δ_{m,n} e^{-imα}.""" - L = 4 - alpha = torch.tensor([0.5], dtype=torch.float64) - zero = torch.zeros_like(alpha) - D = wigner_D_pointwise(alpha, zero, zero, L)[0] # (L, 2L-1, 2L-1) - for l in range(L): - sl = slice(L - 1 - l, L - 1 + l + 1) - Dl = D[l, sl, sl].numpy() - # diagonal entries - for idx in range(2 * l + 1): - m = idx - l - expected = np.exp(-1j * m * 0.5) - np.testing.assert_allclose(Dl[idx, idx], expected, atol=1e-12) - # off-diagonal must vanish - offdiag = Dl - np.diag(np.diag(Dl)) - assert np.abs(offdiag).max() < 1e-12 - - -def test_pointwise_matches_grid(): - """evaluate_rotation_function_pointwise and *_grid agree at grid points.""" - L = 4 - torch.manual_seed(1) - # Random xi coefficients (only valid (l, |m|<=l, |n|<=l) entries non-zero). - xi = torch.zeros((L, 2 * L - 1, 2 * L - 1), dtype=torch.complex128) - for l in range(L): - for m in range(-l, l + 1): - for n in range(-l, l + 1): - xi[l, L - 1 + m, L - 1 + n] = (torch.randn(1).item() - + 1j * torch.randn(1).item()) - - C_grid, alphas, betas, gammas = evaluate_rotation_function_grid( - xi, L, n_alpha=2 * L, n_beta=2 * L, n_gamma=2 * L - ) - # Pick a few grid points and verify pointwise matches. - rng = np.random.default_rng(0) - for _ in range(5): - ka = int(rng.integers(0, 2 * L)) - kb = int(rng.integers(0, 2 * L)) - kg = int(rng.integers(0, 2 * L)) - C_from_grid = C_grid[kg, kb, ka] - a = alphas[ka:ka + 1] - b = betas[kb:kb + 1] - g = gammas[kg:kg + 1] - C_pointwise = evaluate_rotation_function_pointwise(xi, a, b, g, L)[0] - np.testing.assert_allclose(C_from_grid.item(), C_pointwise.item(), - atol=1e-10, rtol=1e-8) - - -def test_pointwise_real_for_hermitian_xi(): - """If ξ_{l,m,n} satisfies the conjugacy relation expected of a real cross-correlation, - then C(R) is real-valued.""" - L = 3 - torch.manual_seed(2) - # Build xi from a real-field convention: - # ξ_{l,m,n} = conj(f_{l,m}) g_{l,n}, where f and g are SH coefficients of real fields - # so f_{l,-m} = (-1)^m conj(f_{l,m}). - def random_real_field_coeffs(L): - f = torch.zeros((L, 2 * L - 1), dtype=torch.complex128) - for l in range(L): - for m in range(0, l + 1): - r = torch.randn(1).item() + 1j * torch.randn(1).item() - if m == 0: - r = complex(r.real, 0.0) - f[l, L - 1 + m] = r - if m > 0: - f[l, L - 1 - m] = ((-1) ** m) * np.conj(r) - return f - - f = random_real_field_coeffs(L) - g = random_real_field_coeffs(L) - xi = torch.zeros((L, 2 * L - 1, 2 * L - 1), dtype=torch.complex128) - for l in range(L): - for mi in range(2 * L - 1): - for ni in range(2 * L - 1): - xi[l, mi, ni] = f[l, mi].conj() * g[l, ni] - - C_grid, _, _, _ = evaluate_rotation_function_grid(xi, L) - imag = C_grid.imag.abs().max().item() - real = C_grid.real.abs().max().item() - assert imag < 1e-10 * max(real, 1.0), f"C imag={imag} too large (real max {real})" diff --git a/tests/unit/frf_separate/test_invariants.py b/tests/unit/frf_separate/test_invariants.py index bebafbe8..6e4bc394 100644 --- a/tests/unit/frf_separate/test_invariants.py +++ b/tests/unit/frf_separate/test_invariants.py @@ -17,8 +17,10 @@ build_dense_map_per_beta, evaluate_rotation_function, ) -from torchref.experimental.alignment.frf.wigner_d import wigner_contraction_per_beta -from torchref.experimental.alignment.wigner import small_d_packed +from torchref.experimental.alignment.frf.wigner_d import ( + _wigner_d_blocks, + wigner_contraction_per_beta, +) def _make_xi(L: int, seed: int = 42) -> torch.Tensor: @@ -126,12 +128,13 @@ def test_wigner_contraction_symmetry(): """At β = π/2 the small-d satisfies d^l_{m,n}(π/2) = (-1)^{l+m} d^l_{m,-n}(π/2).""" L = 8 betas = torch.tensor([math.pi / 2], dtype=torch.float64) - d = small_d_packed(L, betas)[0] # (L, 2L-1, 2L-1) + blocks = _wigner_d_blocks(L, betas, torch.device("cpu"), torch.float64) for l in range(2, L, 2): + d = blocks[l - 1][0] # (2l+1, 2l+1) for m in range(-l, l + 1): for n in range(-l, l + 1): - lhs = d[l, L - 1 + m, L - 1 + n].item() - rhs = ((-1) ** (l + m)) * d[l, L - 1 + m, L - 1 - n].item() + lhs = d[l + m, l + n].item() + rhs = ((-1) ** (l + m)) * d[l + m, l - n].item() assert abs(lhs - rhs) < 1e-10, ( f"l={l} m={m} n={n}: {lhs} vs {rhs}" ) @@ -142,12 +145,13 @@ def test_beta_reflection_identity(): L = 6 beta = 0.37 betas = torch.tensor([beta, math.pi - beta], dtype=torch.float64) - d = small_d_packed(L, betas) + blocks = _wigner_d_blocks(L, betas, torch.device("cpu"), torch.float64) for l in range(2, L, 2): + d = blocks[l - 1] # (2, 2l+1, 2l+1) for m in range(-l, l + 1): for n in range(-l, l + 1): - lhs = d[1, l, L - 1 + m, L - 1 + n].item() # d(π-β) - rhs = ((-1) ** (l + m)) * d[0, l, L - 1 + m, L - 1 - n].item() # (-1)^(l+m) d(β)|n→-n + lhs = d[1, l + m, l + n].item() # d(π-β) + rhs = ((-1) ** (l + m)) * d[0, l + m, l - n].item() # (-1)^(l+m) d(β)|n→-n assert abs(lhs - rhs) < 1e-10 diff --git a/torchref/experimental/alignment/__init__.py b/torchref/experimental/alignment/__init__.py index 7cd6a43d..3c9c0b1f 100644 --- a/torchref/experimental/alignment/__init__.py +++ b/torchref/experimental/alignment/__init__.py @@ -59,13 +59,6 @@ equal_count_shell_edges, assign_shells, ) -from .wigner import ( - small_d_block, - small_d_packed, - wigner_D_pointwise, - evaluate_rotation_function_grid, - evaluate_rotation_function_pointwise, -) from .pipeline import ( MolecularReplacementPipeline, MRSolution, @@ -93,30 +86,6 @@ # ============================================================================= from .rigid_body import RigidBodyRefinement, RigidBodyResult -# ============================================================================= -# Rigid body transformations -# ============================================================================= -from .transform import ( - RigidTransform, - quaternion_normalize, - quaternion_conjugate, - quaternion_multiply, - quaternion_rotate, - quaternion_to_matrix, - matrix_to_quaternion, - axis_angle_to_quaternion, - quaternion_to_axis_angle, - quaternion_to_euler_zyz, - euler_zyz_to_quaternion, - rotation_matrix_from_euler, - sample_angles, -) - -# ============================================================================= -# Clash scoring -# ============================================================================= -from .clashscore import ClashScoreCalculator, AtomSampler, compute_clash_score - # ============================================================================= # ML distributions # ============================================================================= @@ -129,11 +98,6 @@ centric_pdf, ) -# ============================================================================= -# Sampling utilities -# ============================================================================= -from .sampling import VectorSampler, get_rotation_sampling_range - __all__ = [ # Rotation search "FastRotationFunction", @@ -152,11 +116,6 @@ "sh_expand_ball", "equal_count_shell_edges", "assign_shells", - "small_d_block", - "small_d_packed", - "wigner_D_pointwise", - "evaluate_rotation_function_grid", - "evaluate_rotation_function_pointwise", # Pipeline "MolecularReplacementPipeline", "MRSolution", @@ -177,23 +136,7 @@ "RigidBodyRefinement", "RigidBodyResult", # Transforms - "RigidTransform", - "quaternion_normalize", - "quaternion_conjugate", - "quaternion_multiply", - "quaternion_rotate", - "quaternion_to_matrix", - "matrix_to_quaternion", - "axis_angle_to_quaternion", - "quaternion_to_axis_angle", - "quaternion_to_euler_zyz", - "euler_zyz_to_quaternion", - "rotation_matrix_from_euler", - "sample_angles", # Clash scoring - "ClashScoreCalculator", - "AtomSampler", - "compute_clash_score", # Distributions "stable_log_bessel_i0", "rice_log_likelihood", @@ -202,6 +145,4 @@ "acentric_pdf", "centric_pdf", # Utilities - "VectorSampler", - "get_rotation_sampling_range", ] diff --git a/torchref/experimental/alignment/clashscore.py b/torchref/experimental/alignment/clashscore.py deleted file mode 100644 index 5ccbead7..00000000 --- a/torchref/experimental/alignment/clashscore.py +++ /dev/null @@ -1,389 +0,0 @@ -""" -Clash score calculator for crystallographic alignment. - -Provides clash-based scoring to complement Patterson alignment by detecting -steric clashes between symmetry-related molecules. -""" - -from typing import TYPE_CHECKING, List, Optional - -import torch -import torch.nn as nn - -from torchref.base.coordinates import ( - cartesian_to_fractional_torch, - fractional_to_cartesian_torch, -) -from torchref.config import get_default_device, get_float_dtype -from torchref.symmetry import Cell, SpaceGroup -from torchref.symmetry.spacegroup import SpaceGroupLike -from torchref.utils.device_mixin import DeviceMixin - -from .transform import RigidTransform - -if TYPE_CHECKING: - from torchref.model.model import Model - - -class AtomSampler: - """ - Select representative atoms for clash checking. - - Provides methods to create atom selection masks based on the type of - molecular structure (protein vs. ligand-containing). - """ - - @staticmethod - def from_model( - model: "Model", - mode: str = "auto", - ) -> torch.Tensor: - """ - Create atom selection mask from a Model. - - Parameters - ---------- - model : Model - The crystallographic model containing atomic data. - mode : str, default 'auto' - Selection mode: - - 'auto': Use CA atoms if protein present (ATOM records), - else all atoms (for small molecules with only HETATM) - - 'ca_only': Only CA atoms (alpha carbons) - - 'all_atoms': All atoms in the structure - - Returns - ------- - torch.Tensor - Boolean mask of shape (n_atoms,) indicating which atoms to use. - - Examples - -------- - :: - - from torchref.model import Model - model = Model().load_pdb('protein.pdb') - mask = AtomSampler.from_model(model, mode='auto') - print(f"Selected {mask.sum()} atoms out of {len(mask)}") - """ - pdb = model.pdb - - if mode == "auto": - # Check if structure has normal ATOM records (protein/nucleic acid) - has_atom = (pdb["ATOM"] == "ATOM").any() - if has_atom: - # Has protein/nucleic acid: use CA atoms for efficiency - return torch.tensor((pdb["name"] == "CA").values, dtype=torch.bool) - else: - # Only HETATM (small molecule): use all atoms - return torch.ones(len(pdb), dtype=torch.bool) - elif mode == "ca_only": - return torch.tensor((pdb["name"] == "CA").values, dtype=torch.bool) - elif mode == "all_atoms": - return torch.ones(len(pdb), dtype=torch.bool) - else: - raise ValueError( - f"Unknown mode '{mode}'. Use 'auto', 'ca_only', or 'all_atoms'." - ) - - -class ClashScoreCalculator(DeviceMixin, nn.Module): - """ - Calculate clash scores between symmetry-related molecules. - - Computes steric clash violations between an asymmetric unit (ASU) and - its symmetry-related copies. Automatically filters symmetry mates based - on the actual input coordinates to only consider those that can potentially - clash. - - Uses a steep **4 penalty: (radius² - dist²)² which rises quickly as - atoms get closer than the clash radius. - - Parameters - ---------- - symmetry : str, int, gemmi.SpaceGroup, or SpaceGroup - Space group specification for symmetry expansion. - default_clash_radius : float, default 5.0 - Default minimum allowed distance between atoms (can be overridden in forward). - dtype : torch.dtype, default torch.float32 - Data type for computations. - device : torch.device, default 'cpu' - Device for computations. - - Examples - -------- - :: - - from torchref.experimental.alignment.clashscore import ClashScoreCalculator, AtomSampler - from torchref.model import Model - - model = Model().load_pdb('structure.pdb') - calc = ClashScoreCalculator(symmetry=model.spacegroup) - mask = AtomSampler.from_model(model) - score = calc(xyz=model.xyz(), cell=model.cell, atom_mask=mask) - print(f"Clash score: {score.item():.4f}") - """ - - # Cell offsets for neighboring cells (7 cells: central + 6 face neighbors) - _cell_offsets = [ - (0, 0, 0), # Central - (-1, 0, 0), - (1, 0, 0), # x neighbors - (0, -1, 0), - (0, 1, 0), # y neighbors - (0, 0, -1), - (0, 0, 1), # z neighbors - ] - - def __init__( - self, - symmetry: SpaceGroupLike, - default_clash_radius: float = 5.0, - dtype: torch.dtype = None, - device: torch.device = None, - ): - super().__init__() - if dtype is None: - dtype = get_float_dtype() - if device is None: - device = get_default_device() - self.default_clash_radius = default_clash_radius - self.dtype = dtype - self._device = device - - # Initialize symmetry handler. Use the user-configured dtype so the - # SpaceGroup stays MPS-compatible; ``_get_valid_transforms`` casts to - # CPU+float64 internally where high-precision symmetry math is needed. - if isinstance(symmetry, SpaceGroup): - self.symmetry = symmetry - else: - self.symmetry = SpaceGroup(symmetry, dtype=self.dtype, device=device) - - def _get_valid_transforms( - self, - cell: torch.Tensor, - centroid_frac: torch.Tensor, - molecule_radius: float, - clash_radius: float, - ) -> List[RigidTransform]: - """ - Compute which symmetry mates can potentially clash with the ASU. - - Parameters - ---------- - cell : torch.Tensor - Unit cell parameters [a, b, c, alpha, beta, gamma]. - centroid_frac : torch.Tensor - Fractional coordinates of molecule centroid. - molecule_radius : float - Approximate radius of the molecule in Angstroms. - clash_radius : float - Minimum allowed distance between atoms. - - Returns - ------- - List[RigidTransform] - List of symmetry transforms that could produce clashes. - """ - # Compute fractionalization matrix using Cell. Pin to CPU because the - # rest of this method does CPU-only float64 symmetry math. - cell_obj = Cell(cell, dtype=torch.float64, device="cpu") - B = cell_obj.fractional_matrix - - centroid_frac = centroid_frac.to(device="cpu", dtype=torch.float64) - - # Threshold distance for filtering - # Two molecules can clash if centroid distance < 2*radius + clash_radius - threshold = 2 * molecule_radius + clash_radius - - # Identity matrix for comparison with rotation matrices - I = torch.eye(3, dtype=torch.float64) - - valid_transforms = [] - n_ops = self.symmetry.n_ops - - for op_idx in range(n_ops): - R = self.symmetry.matrices[op_idx].cpu().to(torch.float64) - t = self.symmetry.translations[op_idx].cpu().to(torch.float64) - - for offset in self._cell_offsets: - # Skip identity operation in central cell (self-interaction) - if op_idx == 0 and offset == (0, 0, 0): - continue - - offset_tensor = torch.tensor(offset, dtype=torch.float64) - - # Displacement between ASU centroid and this symmetry mate's centroid - # Symmetry mate position: R @ x + t + offset - # Displacement from ASU (identity, no offset): (R - I) @ centroid + t + offset - d_frac = (R - I) @ centroid_frac + t + offset_tensor - - # Convert to Cartesian distance - d_cart = B @ d_frac - dist = d_cart.norm().item() - - if dist < threshold: - # Create RigidTransform for this symmetry operation in fractional space - t_total = t + offset_tensor - transform = RigidTransform.from_matrix(R, t_total) - valid_transforms.append(transform) - - return valid_transforms - - def forward( - self, - xyz: torch.Tensor, - cell: torch.Tensor, - atom_mask: Optional[torch.Tensor] = None, - clash_radius: float = 5.0, - ) -> torch.Tensor: - """ - Compute clash score for given coordinates. - - Automatically determines which symmetry mates could clash based on - the actual input coordinates. Uses squared distances for efficiency - and a steep **4 penalty for clashes. - - Parameters - ---------- - xyz : torch.Tensor - Cartesian coordinates of shape (N, 3). - cell : torch.Tensor - Unit cell parameters [a, b, c, alpha, beta, gamma]. - atom_mask : torch.Tensor, optional - Boolean mask of shape (N,) selecting atoms to use. - If None, all atoms are used. - clash_radius : float, default 5.0 - Minimum allowed distance between atoms in Angstroms. - Atoms closer than this will contribute to the clash score. - - Returns - ------- - torch.Tensor - Scalar clash score. Lower values indicate fewer clashes. - Zero indicates no clashes within the clash radius. - """ - device = xyz.device - dtype = xyz.dtype - - # Apply atom mask - if atom_mask is not None: - atom_mask = atom_mask.to(device) - xyz_selected = xyz[atom_mask] - else: - xyz_selected = xyz - - n_atoms = xyz_selected.shape[0] - - if n_atoms == 0: - return torch.tensor(0.0, device=device, dtype=dtype) - - # Compute molecule properties from actual coordinates - # Use Cell object to get transformation matrices - cell_obj = Cell(cell, dtype=dtype, device=device) - B = cell_obj.fractional_matrix - B_inv = cell_obj.inv_fractional_matrix - - centroid = xyz_selected.mean(dim=0) - molecule_radius = (xyz_selected - centroid).norm(dim=1).max().item() - centroid_frac = cartesian_to_fractional_torch( - centroid.unsqueeze(0), cell_obj.data, B_inv - ).squeeze(0) - - # Get valid transforms for these coordinates - valid_transforms = self._get_valid_transforms( - cell=cell, - centroid_frac=centroid_frac, - molecule_radius=molecule_radius, - clash_radius=clash_radius, - ) - - # If no valid mates after filtering, no clashes possible - if len(valid_transforms) == 0: - return torch.tensor(0.0, device=device, dtype=dtype) - - # Convert ASU to fractional coordinates - xyz_frac = cartesian_to_fractional_torch(xyz_selected, cell_obj.data, B_inv) - - # Precompute squared clash radius threshold - clash_radius_sq = clash_radius**2 - - # Accumulate score over valid symmetry mates - total_score = torch.tensor(0.0, device=device, dtype=dtype) - n_clashes = 0 - - for transform in valid_transforms: - # Apply symmetry operation in fractional space using RigidTransform - xyz_mate_frac = transform.apply(xyz_frac) - - # Convert back to Cartesian - xyz_mate_cart = fractional_to_cartesian_torch( - xyz_mate_frac, cell_obj.data, B - ) - - # Compute squared pairwise distances (more efficient, no sqrt) - diff = xyz_selected.unsqueeze(1) - xyz_mate_cart.unsqueeze(0) # (N, N, 3) - dists_sq = (diff**2).sum(dim=-1) # (N, N) - - # Compute violations: (radius² - dist²), clipped to 0 - # This gives a steep penalty that increases rapidly as atoms get closer - violations_sq = torch.clamp(clash_radius_sq - dists_sq, min=0.0) - - # Apply **2 to squared violations = **4 penalty on distance violation - # (radius² - dist²)² rises steeply as dist approaches 0 - total_score = total_score + (violations_sq**2).sum() - n_clashes += (violations_sq > 0).sum().item() - - # Normalize by number of atom pairs checked - n_pairs = n_atoms * n_atoms * len(valid_transforms) - if n_pairs > 0: - total_score = total_score / n_pairs - - return total_score - - -def compute_clash_score( - model: "Model", - mode: str = "auto", - clash_radius: float = 5.0, -) -> torch.Tensor: - """ - Convenience function to compute clash score for a model. - - Parameters - ---------- - model : Model - Crystallographic model with coordinates and symmetry. - mode : str, default 'auto' - Atom selection mode ('auto', 'ca_only', 'all_atoms'). - clash_radius : float, default 5.0 - Minimum allowed distance between atoms in Angstroms. - - Returns - ------- - torch.Tensor - Scalar clash score. - - Examples - -------- - :: - - from torchref.model import Model - from torchref.experimental.alignment.clashscore import compute_clash_score - model = Model().load_pdb('structure.pdb') - score = compute_clash_score(model) - print(f"Clash score: {score.item():.4f}") - """ - calc = ClashScoreCalculator( - symmetry=model.spacegroup, - device=model.device, - ) - - atom_mask = AtomSampler.from_model(model, mode=mode) - - return calc( - xyz=model.xyz(), - cell=model.cell, - atom_mask=atom_mask, - clash_radius=clash_radius, - ) diff --git a/torchref/experimental/alignment/frf/wigner_d.py b/torchref/experimental/alignment/frf/wigner_d.py index a312993c..3ac5551e 100644 --- a/torchref/experimental/alignment/frf/wigner_d.py +++ b/torchref/experimental/alignment/frf/wigner_d.py @@ -4,8 +4,8 @@ ``djmn_recursive_table`` used in ``FastRot.cc:41`` per-l, per-β). Phaser uses the Sakurai recurrence convention; the equivalent Edmonds -(4.1.23) convention is used throughout this package and is pinned against -Phaser's output by ``tests/unit/alignment/test_wigner.py``. +(4.1.23) convention is used throughout this package. Its small-d identities are +guarded by ``tests/unit/frf_separate/test_invariants.py``. ``wigner_contraction_per_beta`` builds the small-d blocks it needs from the ``J_y`` eigendecomposition, which stays bounded to any ``l``. The blocks depend only on the bandwidth and the β grid, so they are memoised for reuse across diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index b1eb7e53..8756c303 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -174,8 +174,6 @@ class MRSolution: solvent-aware Scaler R-work. model : ModelFT The rotated (+translated +refined) model for this candidate. - clash_score : float, optional - Steric clash score, only populated when ``clash_filter`` is enabled. """ rotation: np.ndarray @@ -184,7 +182,6 @@ class MRSolution: translation_score: float r_factor: float model: "ModelFT" - clash_score: Optional[float] = None class MolecularReplacementPipeline(DeviceMixin): diff --git a/torchref/experimental/alignment/sampling.py b/torchref/experimental/alignment/sampling.py deleted file mode 100644 index d4d2195f..00000000 --- a/torchref/experimental/alignment/sampling.py +++ /dev/null @@ -1,338 +0,0 @@ -""" -Atom pair sampling for efficient Patterson vector generation. - -Supports weighted sampling to prioritize informative pairs -(heavy atoms, close distances). -""" - -from typing import Optional, Tuple - -import numpy as np -import torch - -from torchref.config import get_float_dtype - - -class VectorSampler: - """ - Samples atom pairs for Patterson vector matching. - - Supports weighted sampling to prioritize informative pairs - (heavy atoms via Z-weighting). Samples pairs from the asymmetric - unit (ASU) only - symmetry is already encoded in the Patterson map. - - Parameters - ---------- - model : Model - TorchRef Model object. The caller is responsible for filtering - atoms (e.g., excluding waters) before passing to this class. - weighting : str, optional - Weighting scheme: 'uniform' or 'Z2' (weight by atomic number squared). - Default is 'Z2'. - seed : int, optional - Random seed for reproducibility. Default is None. - - Attributes - ---------- - model : Model - The model used for sampling. - n_atoms : int - Number of atoms in the model. - weighting : str - Weighting scheme used. - weights : torch.Tensor - Sampling weights for each atom (n_atoms,). - rng : torch.Generator - Random number generator. - """ - - def __init__(self, model, weighting: str = "Z2", seed: int = None): - """ - Initialize the VectorSampler. - - Parameters - ---------- - model : Model - TorchRef Model object. The caller is responsible for filtering - atoms (e.g., excluding waters) before passing to this class. - weighting : str - Weighting scheme for sampling. - seed : int, optional - Random seed for reproducibility. - """ - self.model = model - self.n_atoms = len(model.pdb) - self.weighting = weighting - self.rng = ( - torch.Generator().manual_seed(seed) - if seed is not None - else torch.Generator() - ) - self.weights = self._compute_weights() - - def _compute_weights( - self, - ) -> torch.Tensor: - """ - Compute sampling probability for each atom based on atomic number and B-factor. - - Weights are computed as: Z^2 / B (for Z2 weighting) or 1/B (for uniform). - Atoms with lower B-factors (more ordered) get higher weights since they - contribute more signal to the Patterson map. - - Returns - ------- - torch.Tensor - Weight for each atom with shape (n_atoms,). - """ - from torchref.utils.pse import PERIODIC_TABLE - - elements = self.model.pdb.element.values - - Zs = torch.tensor( - [PERIODIC_TABLE[el]["number"] for el in elements], - dtype=get_float_dtype(), - device=self.model.device, - ) - - # Get B-factors and compute reciprocal weights - # Use 1/B so atoms with lower B-factors get higher weights - B_factors = torch.tensor( - self.model.pdb["tempfactor"].values, - dtype=get_float_dtype(), - device=self.model.device, - ) - # Clamp B-factors to avoid division by zero or very small values - B_factors = torch.clamp(B_factors, min=1.0) - B_weights = 1.0 / B_factors - - if self.weighting == "Z2": - weights = Zs**2 * B_weights - else: # uniform - weights = B_weights - - weights = weights / weights.sum() # Normalize to probabilities - return weights - - def sample( - self, n_vectors: int, weights: Optional[torch.Tensor] = None - ) -> tuple[torch.Tensor, torch.Tensor]: - """ - Sample atom pairs according to weighting scheme. - - Parameters - ---------- - n_vectors : int - Number of atom pairs to sample. - weights : torch.Tensor, optional - Override weights for sampling. If None, uses self.weights. - - Returns - ------- - tuple[torch.Tensor, torch.Tensor] - Two tensors of shape (n_vectors,) containing - the indices of the sampled atom pairs. - """ - w = weights if weights is not None else self.weights - - # Sample first indices according to weights - idx1 = torch.multinomial(w, n_vectors, replacement=True, generator=self.rng) - - # Sample second indices according to weights - idx2 = torch.multinomial(w, n_vectors, replacement=True, generator=self.rng) - - # Redraw idx2 where it equals idx1 - same_mask = idx1 == idx2 - max_attempts = 100 # Prevent infinite loop - attempt = 0 - while same_mask.any() and attempt < max_attempts: - n_resample = same_mask.sum().item() - idx2[same_mask] = torch.multinomial( - w, n_resample, replacement=True, generator=self.rng - ) - same_mask = idx1 == idx2 - attempt += 1 - - return idx1, idx2 - - -def get_rotation_sampling_range( - rotation_matrices: torch.Tensor, -) -> Tuple[float, float, float]: - """ - Determine rotation angle sampling ranges given point group symmetry operations. - - Given the rotation matrices from a spacegroup's point group, this function - computes the asymmetric unit in rotation space (SO(3)) and returns the - maximum Euler angles (alpha, beta, gamma) needed to cover the asymmetric unit. - - The function uses ZYZ Euler angle convention where: - - alpha: rotation about Z axis, range [0, 2*pi) - - beta: rotation about Y axis, range [0, pi] - - gamma: rotation about Z axis, range [0, 2*pi) - - For crystals, the point group symmetry reduces the search space: - - Triclinic (1): Full SO(3) - (2*pi, pi, 2*pi) - - Monoclinic (2): Half of SO(3) - (2*pi, pi, pi) - - Orthorhombic (222): 1/4 of SO(3) - (pi, pi, pi) - - Tetragonal (4, 422): 1/8 or 1/16 of SO(3) - - Trigonal (3, 32): 1/6 or 1/12 of SO(3) - - Hexagonal (6, 622): 1/12 or 1/24 of SO(3) - - Cubic (23, 432): 1/12 or 1/24 of SO(3) - - Parameters - ---------- - rotation_matrices : torch.Tensor - Point group rotation matrices with shape (N, 3, 3), where N is the - number of symmetry operations. These should be the pure rotation - parts of the spacegroup operations (no translations). - - Returns - ------- - Tuple[float, float, float] - Maximum values for (alpha, beta, gamma) Euler angles in radians. - These define the asymmetric unit in rotation space that needs to - be sampled during molecular replacement searches. - - Examples - -------- - :: - - from torchref.symmetry import SpaceGroup - sg = SpaceGroup('P212121') # Orthorhombic - ranges = get_rotation_sampling_range(sg.matrices) - print(f"alpha: {ranges[0]:.4f}, beta: {ranges[1]:.4f}, gamma: {ranges[2]:.4f}") - # alpha: 3.1416, beta: 3.1416, gamma: 3.1416 - - sg = SpaceGroup('P1') # Triclinic - need full SO(3) - ranges = get_rotation_sampling_range(sg.matrices) - print(f"alpha: {ranges[0]:.4f}, beta: {ranges[1]:.4f}, gamma: {ranges[2]:.4f}") - alpha: 6.2832, beta: 3.1416, gamma: 6.2832 - """ - n_ops = rotation_matrices.shape[0] - - # Convert to numpy for analysis - if isinstance(rotation_matrices, torch.Tensor): - R_ops = rotation_matrices.detach().cpu().numpy() - else: - R_ops = np.array(rotation_matrices) - - # Analyze the point group to determine fold symmetries - # We look for rotation axes and their orders - - # Default: full SO(3) coverage - alpha_max = 2 * np.pi - beta_max = np.pi - gamma_max = 2 * np.pi - - # Identity only (P1) - need full search - if n_ops == 1: - return (alpha_max, beta_max, gamma_max) - - # Analyze rotation axes and angles - axes_and_angles = [] - for R in R_ops: - # Skip identity - trace = np.trace(R) - if np.abs(trace - 3.0) < 1e-6: - continue - - # Rotation angle from trace: trace = 1 + 2*cos(theta) - cos_theta = (trace - 1.0) / 2.0 - cos_theta = np.clip(cos_theta, -1.0, 1.0) - angle = np.arccos(cos_theta) - - if angle < 1e-6: - continue - - # Get rotation axis from antisymmetric part of R - # axis is proportional to (R - R^T) - axis = np.array([R[2, 1] - R[1, 2], R[0, 2] - R[2, 0], R[1, 0] - R[0, 1]]) - norm = np.linalg.norm(axis) - if norm > 1e-6: - axis = axis / norm - else: - # 180-degree rotation - get axis from R + I - # For 180° rotation, axis is eigenvector with eigenvalue 1 - eigvals, eigvecs = np.linalg.eig(R) - idx = np.argmin(np.abs(eigvals - 1.0)) - axis = np.real(eigvecs[:, idx]) - axis = axis / np.linalg.norm(axis) - - axes_and_angles.append((axis, angle)) - - # Determine fold along principal axes - z_axis = np.array([0, 0, 1]) - y_axis = np.array([0, 1, 0]) - x_axis = np.array([1, 0, 0]) - - z_fold = 1 - y_fold = 1 - x_fold = 1 - - for axis, angle in axes_and_angles: - # Check if axis is along z - if np.abs(np.abs(np.dot(axis, z_axis)) - 1.0) < 0.1: - fold = int(round(2 * np.pi / angle)) - z_fold = max(z_fold, fold) - # Check if axis is along y - elif np.abs(np.abs(np.dot(axis, y_axis)) - 1.0) < 0.1: - fold = int(round(2 * np.pi / angle)) - y_fold = max(y_fold, fold) - # Check if axis is along x - elif np.abs(np.abs(np.dot(axis, x_axis)) - 1.0) < 0.1: - fold = int(round(2 * np.pi / angle)) - x_fold = max(x_fold, fold) - - # Also check for 2-fold along diagonal (orthorhombic has 3 perpendicular 2-folds) - has_three_twofolds = False - twofold_count = 0 - for axis, angle in axes_and_angles: - fold = int(round(2 * np.pi / angle)) - if fold == 2: - twofold_count += 1 - if twofold_count >= 3: - has_three_twofolds = True - - # Determine sampling ranges based on symmetry analysis - # For ZYZ Euler angles: - # - z_fold reduces alpha range - # - 2-fold perpendicular to z reduces beta range to [0, pi/2] in some cases - # - Combined symmetry reduces gamma range - - # Alpha reduction based on z-axis fold - if z_fold > 1: - alpha_max = 2 * np.pi / z_fold - - # Check for 2-fold perpendicular to z-axis (reduces gamma) - has_perp_twofold = False - for axis, angle in axes_and_angles: - fold = int(round(2 * np.pi / angle)) - if fold == 2: - # Check if axis is perpendicular to z - if np.abs(np.dot(axis, z_axis)) < 0.1: - has_perp_twofold = True - break - - if has_perp_twofold: - gamma_max = np.pi - - # For point groups with higher symmetry, apply additional reductions - # Based on number of operations (proxy for point group order) - if n_ops >= 24: - # Cubic or high-symmetry hexagonal - alpha_max = min(alpha_max, np.pi / 2) - gamma_max = min(gamma_max, np.pi / 2) - elif n_ops >= 12: - # Hexagonal 622, Cubic 23, etc. - alpha_max = min(alpha_max, np.pi) - gamma_max = min(gamma_max, np.pi) - elif n_ops >= 8: - # Tetragonal 422 - gamma_max = min(gamma_max, np.pi) - elif has_three_twofolds: - # Orthorhombic 222 - alpha_max = np.pi - gamma_max = np.pi - - return (alpha_max, beta_max, gamma_max) diff --git a/torchref/experimental/alignment/transform.py b/torchref/experimental/alignment/transform.py deleted file mode 100644 index c7ea3c01..00000000 --- a/torchref/experimental/alignment/transform.py +++ /dev/null @@ -1,921 +0,0 @@ -""" -Rigid body transformations for crystallographic alignment. - -Provides unified handling of rotations and translations with quaternion-based -internal storage and multiple representation formats. -""" - -from typing import Optional, Union - -import torch -import torch.nn as nn - -from torchref.config import get_default_device, get_float_dtype -from torchref.utils.device_mixin import DeviceMixin - -# ============================================================================= -# Quaternion Helper Functions -# ============================================================================= - - -def get_inverse_rotation_matrix(R: torch.Tensor) -> torch.Tensor: - """ - Compute inverse of rotation matrix (transpose for orthogonal matrices). - - Parameters - ---------- - R : torch.Tensor - Rotation matrix of shape (3, 3) or (N, 3, 3). - - Returns - ------- - torch.Tensor - Inverse rotation matrix of same shape. - """ - return R.transpose(-2, -1) - - -def sample_angles(sampling_pitch_rad, max_angles_rad): - """ - Sample Euler angles (in radians) up to the specified maximum angles with the given sampling pitch. - Returns a tensor of shape (N, 3) where N is the number of sampled angles. - - Args: - sampling_pitch_rad (float): Sampling pitch in radians. - max_angles_rad (tuple): Maximum angles (alpha, beta, gamma) in radians. - Returns: - torch.Tensor: Sampled angles of shape (N, 3). - - """ - - angles = [] - max_alpha, max_beta, max_gamma = max_angles_rad - alpha = torch.arange(0, max_alpha + 1e-6, sampling_pitch_rad, dtype=torch.float32) - beta = torch.arange(0, max_beta + 1e-6, sampling_pitch_rad, dtype=torch.float32) - gamma = torch.arange(0, max_gamma + 1e-6, sampling_pitch_rad, dtype=torch.float32) - alpha, beta, gamma = torch.meshgrid(alpha, beta, gamma, indexing="ij") - - return torch.stack([alpha.flatten(), beta.flatten(), gamma.flatten()], dim=-1) - - -def rotation_matrix_from_euler(angles): - """ - Compute rotation matrices from Euler angles (in radians). - Angles should be of shape (N, 3) where N is the number of angle sets. - - Args: - angles (torch.Tensor): Euler angles of shape (N, 3). - Returns: - torch.Tensor: Rotation matrices of shape (N, 3, 3). - """ - alpha = angles[:, 0] - beta = angles[:, 1] - gamma = angles[:, 2] - - R_alpha = torch.stack( - [ - torch.cos(alpha), - -torch.sin(alpha), - torch.zeros_like(alpha), - torch.sin(alpha), - torch.cos(alpha), - torch.zeros_like(alpha), - torch.zeros_like(alpha), - torch.zeros_like(alpha), - torch.ones_like(alpha), - ], - dim=-1, - ).reshape(-1, 3, 3) - - R_beta = torch.stack( - [ - torch.cos(beta), - torch.zeros_like(beta), - torch.sin(beta), - torch.zeros_like(beta), - torch.ones_like(beta), - torch.zeros_like(beta), - -torch.sin(beta), - torch.zeros_like(beta), - torch.cos(beta), - ], - dim=-1, - ).reshape(-1, 3, 3) - - R_gamma = torch.stack( - [ - torch.ones_like(gamma), - torch.zeros_like(gamma), - torch.zeros_like(gamma), - torch.zeros_like(gamma), - torch.cos(gamma), - -torch.sin(gamma), - torch.zeros_like(gamma), - torch.sin(gamma), - torch.cos(gamma), - ], - dim=-1, - ).reshape(-1, 3, 3) - - R = torch.einsum("rij,rjk,rkl->ril", R_gamma, R_beta, R_alpha) - - return R - - -def quaternion_normalize(q: torch.Tensor) -> torch.Tensor: - """ - Normalize quaternion to unit length. - - Parameters - ---------- - q : torch.Tensor - Quaternion(s) of shape (4,) or (N, 4). - - Returns - ------- - torch.Tensor - Normalized quaternion(s) of same shape. - """ - return q / q.norm(dim=-1, keepdim=True).clamp(min=1e-12) - - -def quaternion_conjugate(q: torch.Tensor) -> torch.Tensor: - """ - Compute quaternion conjugate (inverse for unit quaternions). - - For q = [w, x, y, z], conjugate is [w, -x, -y, -z]. - - Parameters - ---------- - q : torch.Tensor - Quaternion(s) of shape (4,) or (N, 4). - - Returns - ------- - torch.Tensor - Conjugate quaternion(s) of same shape. - """ - conj = q.clone() - conj[..., 1:] = -conj[..., 1:] - return conj - - -def quaternion_multiply(q1: torch.Tensor, q2: torch.Tensor) -> torch.Tensor: - """ - Compute Hamilton product of two quaternions. - - Parameters - ---------- - q1, q2 : torch.Tensor - Quaternions of shape (4,) or (N, 4). Format: [w, x, y, z]. - - Returns - ------- - torch.Tensor - Product quaternion of same shape. - """ - w1, x1, y1, z1 = q1[..., 0], q1[..., 1], q1[..., 2], q1[..., 3] - w2, x2, y2, z2 = q2[..., 0], q2[..., 1], q2[..., 2], q2[..., 3] - - w = w1 * w2 - x1 * x2 - y1 * y2 - z1 * z2 - x = w1 * x2 + x1 * w2 + y1 * z2 - z1 * y2 - y = w1 * y2 - x1 * z2 + y1 * w2 + z1 * x2 - z = w1 * z2 + x1 * y2 - y1 * x2 + z1 * w2 - - return torch.stack([w, x, y, z], dim=-1) - - -def quaternion_rotate(q: torch.Tensor, v: torch.Tensor) -> torch.Tensor: - """ - Rotate vector(s) by quaternion using q * v * q^*. - - Parameters - ---------- - q : torch.Tensor - Unit quaternion of shape (4,). - v : torch.Tensor - Vector(s) of shape (3,) or (N, 3). - - Returns - ------- - torch.Tensor - Rotated vector(s) of same shape as v. - """ - # Convert vector to pure quaternion [0, x, y, z] - v_shape = v.shape - v_flat = v.reshape(-1, 3) - - v_quat = torch.zeros(v_flat.shape[0], 4, dtype=q.dtype, device=q.device) - v_quat[:, 1:] = v_flat - - # q * v * q^* - q_conj = quaternion_conjugate(q) - result = quaternion_multiply( - quaternion_multiply(q.unsqueeze(0), v_quat), q_conj.unsqueeze(0) - ) - - # Extract vector part - return result[:, 1:].reshape(v_shape) - - -def quaternion_to_matrix(q: torch.Tensor) -> torch.Tensor: - """ - Convert unit quaternion to 3x3 rotation matrix. - - Parameters - ---------- - q : torch.Tensor - Unit quaternion of shape (4,) or (N, 4). Format: [w, x, y, z]. - - Returns - ------- - torch.Tensor - Rotation matrix of shape (3, 3) or (N, 3, 3). - """ - q = quaternion_normalize(q) - - batched = q.dim() == 2 - if not batched: - q = q.unsqueeze(0) - - w, x, y, z = q[:, 0], q[:, 1], q[:, 2], q[:, 3] - - # Rotation matrix from quaternion - R = torch.stack( - [ - torch.stack( - [1 - 2 * (y * y + z * z), 2 * (x * y - w * z), 2 * (x * z + w * y)], - dim=-1, - ), - torch.stack( - [2 * (x * y + w * z), 1 - 2 * (x * x + z * z), 2 * (y * z - w * x)], - dim=-1, - ), - torch.stack( - [2 * (x * z - w * y), 2 * (y * z + w * x), 1 - 2 * (x * x + y * y)], - dim=-1, - ), - ], - dim=-2, - ) - - if not batched: - R = R.squeeze(0) - - return R - - -def matrix_to_quaternion(R: torch.Tensor) -> torch.Tensor: - """ - Convert 3x3 rotation matrix to unit quaternion. - - Uses Shepperd's method for numerical stability. - - Parameters - ---------- - R : torch.Tensor - Rotation matrix of shape (3, 3) or (N, 3, 3). - - Returns - ------- - torch.Tensor - Unit quaternion of shape (4,) or (N, 4). Format: [w, x, y, z]. - """ - batched = R.dim() == 3 - if not batched: - R = R.unsqueeze(0) - - batch_size = R.shape[0] - dtype, device = R.dtype, R.device - - # Shepperd's method - choose largest diagonal element - trace = R[:, 0, 0] + R[:, 1, 1] + R[:, 2, 2] - - q = torch.zeros(batch_size, 4, dtype=dtype, device=device) - - # Case 1: trace > 0 - mask1 = trace > 0 - if mask1.any(): - s = torch.sqrt(trace[mask1] + 1.0) * 2 # s = 4 * w - q[mask1, 0] = 0.25 * s - q[mask1, 1] = (R[mask1, 2, 1] - R[mask1, 1, 2]) / s - q[mask1, 2] = (R[mask1, 0, 2] - R[mask1, 2, 0]) / s - q[mask1, 3] = (R[mask1, 1, 0] - R[mask1, 0, 1]) / s - - # Case 2: R[0,0] is largest diagonal - mask2 = ~mask1 & (R[:, 0, 0] > R[:, 1, 1]) & (R[:, 0, 0] > R[:, 2, 2]) - if mask2.any(): - s = torch.sqrt(1.0 + R[mask2, 0, 0] - R[mask2, 1, 1] - R[mask2, 2, 2]) * 2 - q[mask2, 0] = (R[mask2, 2, 1] - R[mask2, 1, 2]) / s - q[mask2, 1] = 0.25 * s - q[mask2, 2] = (R[mask2, 0, 1] + R[mask2, 1, 0]) / s - q[mask2, 3] = (R[mask2, 0, 2] + R[mask2, 2, 0]) / s - - # Case 3: R[1,1] is largest diagonal - mask3 = ~mask1 & ~mask2 & (R[:, 1, 1] > R[:, 2, 2]) - if mask3.any(): - s = torch.sqrt(1.0 + R[mask3, 1, 1] - R[mask3, 0, 0] - R[mask3, 2, 2]) * 2 - q[mask3, 0] = (R[mask3, 0, 2] - R[mask3, 2, 0]) / s - q[mask3, 1] = (R[mask3, 0, 1] + R[mask3, 1, 0]) / s - q[mask3, 2] = 0.25 * s - q[mask3, 3] = (R[mask3, 1, 2] + R[mask3, 2, 1]) / s - - # Case 4: R[2,2] is largest diagonal - mask4 = ~mask1 & ~mask2 & ~mask3 - if mask4.any(): - s = torch.sqrt(1.0 + R[mask4, 2, 2] - R[mask4, 0, 0] - R[mask4, 1, 1]) * 2 - q[mask4, 0] = (R[mask4, 1, 0] - R[mask4, 0, 1]) / s - q[mask4, 1] = (R[mask4, 0, 2] + R[mask4, 2, 0]) / s - q[mask4, 2] = (R[mask4, 1, 2] + R[mask4, 2, 1]) / s - q[mask4, 3] = 0.25 * s - - # Ensure positive w (canonical form) - q = torch.where(q[:, 0:1] < 0, -q, q) - - if not batched: - q = q.squeeze(0) - - return quaternion_normalize(q) - - -def axis_angle_to_quaternion(axis_angle: torch.Tensor) -> torch.Tensor: - """ - Convert axis-angle representation to quaternion. - - Parameters - ---------- - axis_angle : torch.Tensor - Axis-angle vector of shape (3,) or (N, 3). - Direction is rotation axis, magnitude is angle in radians. - - Returns - ------- - torch.Tensor - Unit quaternion of shape (4,) or (N, 4). - """ - batched = axis_angle.dim() == 2 - if not batched: - axis_angle = axis_angle.unsqueeze(0) - - angle = axis_angle.norm(dim=-1, keepdim=True).clamp(min=1e-12) - axis = axis_angle / angle - - half_angle = angle / 2 - w = torch.cos(half_angle) - xyz = axis * torch.sin(half_angle) - - q = torch.cat([w, xyz], dim=-1) - - if not batched: - q = q.squeeze(0) - - return q - - -def quaternion_to_axis_angle(q: torch.Tensor) -> torch.Tensor: - """ - Convert quaternion to axis-angle representation. - - Parameters - ---------- - q : torch.Tensor - Unit quaternion of shape (4,) or (N, 4). - - Returns - ------- - torch.Tensor - Axis-angle vector of shape (3,) or (N, 3). - """ - q = quaternion_normalize(q) - - batched = q.dim() == 2 - if not batched: - q = q.unsqueeze(0) - - # Ensure positive w for numerical stability - q = torch.where(q[:, 0:1] < 0, -q, q) - - w = q[:, 0].clamp(-1 + 1e-7, 1 - 1e-7) - xyz = q[:, 1:] - - angle = 2 * torch.acos(w) - sin_half = torch.sqrt(1 - w * w).clamp(min=1e-12) - - axis = xyz / sin_half.unsqueeze(-1) - axis_angle = axis * angle.unsqueeze(-1) - - # Handle near-identity case - near_identity = angle < 1e-6 - axis_angle = torch.where( - near_identity.unsqueeze(-1), 2 * xyz, axis_angle # Small angle approximation - ) - - if not batched: - axis_angle = axis_angle.squeeze(0) - - return axis_angle - - -def quaternion_to_euler_zyz(q: torch.Tensor) -> torch.Tensor: - """ - Convert quaternion to Euler angles (ZYZ convention). - - Parameters - ---------- - q : torch.Tensor - Unit quaternion of shape (4,). - - Returns - ------- - torch.Tensor - Euler angles [alpha, beta, gamma] of shape (3,). - Ranges: alpha in [0, 2pi), beta in [0, pi], gamma in [0, 2pi). - """ - # Convert to matrix first, then extract ZYZ angles - R = quaternion_to_matrix(q) - - # ZYZ convention: R = Rz(alpha) @ Ry(beta) @ Rz(gamma) - # beta = acos(R[2,2]) - # alpha = atan2(R[1,2], R[0,2]) - # gamma = atan2(R[2,1], -R[2,0]) - - beta = torch.acos(R[2, 2].clamp(-1, 1)) - - # Handle gimbal lock cases - if torch.abs(torch.sin(beta)) < 1e-6: - # beta ≈ 0 or pi, gimbal lock - alpha = torch.atan2(R[1, 0], R[0, 0]) - gamma = torch.zeros_like(alpha) - else: - alpha = torch.atan2(R[1, 2], R[0, 2]) - gamma = torch.atan2(R[2, 1], -R[2, 0]) - - # Normalize to [0, 2pi) for alpha and gamma - alpha = torch.remainder(alpha, 2 * torch.pi) - gamma = torch.remainder(gamma, 2 * torch.pi) - - return torch.stack([alpha, beta, gamma]) - - -def euler_zyz_to_quaternion(euler: torch.Tensor) -> torch.Tensor: - """ - Convert Euler angles (ZYZ convention) to quaternion. - - Parameters - ---------- - euler : torch.Tensor - Euler angles [alpha, beta, gamma] of shape (3,). - - Returns - ------- - torch.Tensor - Unit quaternion of shape (4,). - """ - alpha, beta, gamma = euler[0], euler[1], euler[2] - - # ZYZ: q = qz(alpha) * qy(beta) * qz(gamma) - ca, sa = torch.cos(alpha / 2), torch.sin(alpha / 2) - cb, sb = torch.cos(beta / 2), torch.sin(beta / 2) - cg, sg = torch.cos(gamma / 2), torch.sin(gamma / 2) - - # qz(alpha) = [cos(a/2), 0, 0, sin(a/2)] - # qy(beta) = [cos(b/2), 0, sin(b/2), 0] - # qz(gamma) = [cos(g/2), 0, 0, sin(g/2)] - - q_alpha = torch.stack([ca, torch.zeros_like(ca), torch.zeros_like(ca), sa]) - q_beta = torch.stack([cb, torch.zeros_like(cb), sb, torch.zeros_like(cb)]) - q_gamma = torch.stack([cg, torch.zeros_like(cg), torch.zeros_like(cg), sg]) - - return quaternion_multiply(quaternion_multiply(q_alpha, q_beta), q_gamma) - - -# ============================================================================= -# RigidTransform Class -# ============================================================================= - - -class RigidTransform(DeviceMixin, nn.Module): - """ - Rigid body transformation with quaternion-based rotation storage. - - Stores rotation internally as unit quaternion [w, x, y, z] and - translation as 3D vector. Provides methods for various representations - and transformation operations. - - Parameters - ---------- - quaternion : torch.Tensor, optional - Rotation as quaternion [w, x, y, z] of shape (4,). - translation : torch.Tensor, optional - Translation vector of shape (3,). Defaults to zeros. - rotation_matrix : torch.Tensor, optional - Alternative: initialize from rotation matrix of shape (3, 3). - axis_angle : torch.Tensor, optional - Alternative: initialize from axis-angle vector of shape (3,). - - Examples - -------- - :: - - T = RigidTransform.random() - coords = torch.randn(100, 3) - coords_transformed = T.apply(coords) - coords_back = T.inverse().apply(coords_transformed) - """ - - def __init__( - self, - quaternion: Optional[torch.Tensor] = None, - translation: Optional[torch.Tensor] = None, - rotation_matrix: Optional[torch.Tensor] = None, - axis_angle: Optional[torch.Tensor] = None, - ): - super().__init__() - - # Determine dtype and device from inputs (fallback to package defaults) - dtype = get_float_dtype() - device = get_default_device() - - if quaternion is not None: - dtype, device = quaternion.dtype, quaternion.device - elif rotation_matrix is not None: - dtype, device = rotation_matrix.dtype, rotation_matrix.device - elif axis_angle is not None: - dtype, device = axis_angle.dtype, axis_angle.device - elif translation is not None: - dtype, device = translation.dtype, translation.device - - # Convert to quaternion from alternative representations - if quaternion is not None: - q = quaternion_normalize(quaternion) - elif rotation_matrix is not None: - q = matrix_to_quaternion(rotation_matrix) - elif axis_angle is not None: - q = axis_angle_to_quaternion(axis_angle) - else: - # Identity rotation - q = torch.tensor([1.0, 0.0, 0.0, 0.0], dtype=dtype, device=device) - - # Set translation - if translation is not None: - t = translation.to(dtype=dtype, device=device) - else: - t = torch.zeros(3, dtype=dtype, device=device) - - # Register as buffers (not parameters by default) - self.register_buffer("_quaternion", q) - self.register_buffer("_translation", t) - - # ========================================================================= - # Properties for different representations - # ========================================================================= - - @property - def quaternion(self) -> torch.Tensor: - """Get rotation as quaternion [w, x, y, z].""" - return self._quaternion - - @property - def rotation_matrix(self) -> torch.Tensor: - """Get rotation as 3x3 matrix.""" - return quaternion_to_matrix(self._quaternion) - - @property - def axis_angle(self) -> torch.Tensor: - """Get rotation as axis-angle vector.""" - return quaternion_to_axis_angle(self._quaternion) - - @property - def euler_zyz(self) -> torch.Tensor: - """Get rotation as Euler angles (ZYZ convention).""" - return quaternion_to_euler_zyz(self._quaternion) - - @property - def translation(self) -> torch.Tensor: - """Get translation vector.""" - return self._translation - - @property - def dtype(self) -> torch.dtype: - """Get data type.""" - return self._quaternion.dtype - - @property - def device(self) -> torch.device: - """Get device.""" - return self._quaternion.device - - # ========================================================================= - # Transformation methods - # ========================================================================= - - def apply(self, coords: torch.Tensor) -> torch.Tensor: - """ - Apply transformation to coordinates: x' = R @ x + t. - - Parameters - ---------- - coords : torch.Tensor - Coordinates of shape (N, 3) or (3,). - - Returns - ------- - torch.Tensor - Transformed coordinates of same shape. - """ - R = self.rotation_matrix - coords_dtype = coords.dtype - coords = coords.to(dtype=self.dtype) - - if coords.dim() == 1: - result = R @ coords + self._translation - else: - result = coords @ R.T + self._translation - - return result.to(dtype=coords_dtype) - - def apply_rotation_only(self, coords: torch.Tensor) -> torch.Tensor: - """ - Apply rotation only: x' = R @ x. - - Parameters - ---------- - coords : torch.Tensor - Coordinates of shape (N, 3) or (3,). - - Returns - ------- - torch.Tensor - Rotated coordinates of same shape. - """ - R = self.rotation_matrix - coords_dtype = coords.dtype - coords = coords.to(dtype=self.dtype) - - if coords.dim() == 1: - result = R @ coords - else: - result = coords @ R.T - - return result.to(dtype=coords_dtype) - - def forward(self, coords: torch.Tensor) -> torch.Tensor: - """nn.Module forward = apply().""" - return self.apply(coords) - - # ========================================================================= - # Composition and inversion - # ========================================================================= - - def inverse(self) -> "RigidTransform": - """ - Compute inverse transformation. - - For T(x) = R @ x + t, inverse is T^{-1}(x) = R^T @ (x - t). - - Returns - ------- - RigidTransform - Inverse transformation. - """ - q_inv = quaternion_conjugate(self._quaternion) - # t_inv = -R^T @ t = -R_inv @ t - t_inv = -quaternion_rotate(q_inv, self._translation) - return RigidTransform(quaternion=q_inv, translation=t_inv) - - def compose(self, other: "RigidTransform") -> "RigidTransform": - """ - Compose with another transformation: (self @ other)(x) = self(other(x)). - - Parameters - ---------- - other : RigidTransform - Transformation to compose with. - - Returns - ------- - RigidTransform - Composed transformation. - """ - q_new = quaternion_multiply(self._quaternion, other._quaternion) - # t_new = R_self @ t_other + t_self - t_new = ( - quaternion_rotate(self._quaternion, other._translation) + self._translation - ) - return RigidTransform(quaternion=q_new, translation=t_new) - - def __matmul__(self, other: Union["RigidTransform", torch.Tensor]): - """ - Support @ operator for composition or application. - - Parameters - ---------- - other : RigidTransform or torch.Tensor - If RigidTransform, composes transformations. - If tensor, applies transformation to coordinates. - - Returns - ------- - RigidTransform or torch.Tensor - Composed transformation or transformed coordinates. - """ - if isinstance(other, RigidTransform): - return self.compose(other) - return self.apply(other) - - # ========================================================================= - # Factory methods - # ========================================================================= - - @classmethod - def identity( - cls, - device: Union[str, torch.device] = None, - dtype: torch.dtype = None, - ) -> "RigidTransform": - """ - Create identity transformation. - - Parameters - ---------- - device : str or torch.device - Device for tensors. - dtype : torch.dtype - Data type for tensors. - - Returns - ------- - RigidTransform - Identity transformation. - """ - if device is None: - device = get_default_device() - if dtype is None: - dtype = get_float_dtype() - q = torch.tensor([1.0, 0.0, 0.0, 0.0], dtype=dtype, device=device) - t = torch.zeros(3, dtype=dtype, device=device) - return cls(quaternion=q, translation=t) - - @classmethod - def from_matrix( - cls, - R: torch.Tensor, - t: Optional[torch.Tensor] = None, - ) -> "RigidTransform": - """ - Create from rotation matrix and translation. - - Parameters - ---------- - R : torch.Tensor - Rotation matrix of shape (3, 3). - t : torch.Tensor, optional - Translation vector of shape (3,). - - Returns - ------- - RigidTransform - Transformation from matrix representation. - """ - return cls(rotation_matrix=R, translation=t) - - @classmethod - def from_axis_angle( - cls, - axis_angle: torch.Tensor, - t: Optional[torch.Tensor] = None, - ) -> "RigidTransform": - """ - Create from axis-angle and translation. - - Parameters - ---------- - axis_angle : torch.Tensor - Axis-angle vector of shape (3,). - t : torch.Tensor, optional - Translation vector of shape (3,). - - Returns - ------- - RigidTransform - Transformation from axis-angle representation. - """ - return cls(axis_angle=axis_angle, translation=t) - - @classmethod - def from_euler_zyz( - cls, - euler: torch.Tensor, - t: Optional[torch.Tensor] = None, - ) -> "RigidTransform": - """ - Create from Euler angles (ZYZ convention) and translation. - - Parameters - ---------- - euler : torch.Tensor - Euler angles [alpha, beta, gamma] of shape (3,). - t : torch.Tensor, optional - Translation vector of shape (3,). - - Returns - ------- - RigidTransform - Transformation from Euler angle representation. - """ - q = euler_zyz_to_quaternion(euler) - return cls(quaternion=q, translation=t) - - @classmethod - def random( - cls, - device: Union[str, torch.device] = None, - dtype: torch.dtype = None, - translation_scale: float = 0.0, - ) -> "RigidTransform": - """ - Create random transformation with uniform rotation over SO(3). - - Uses Shoemake's quaternion-based uniform sampling. - - Parameters - ---------- - device : str or torch.device - Device for tensors. - dtype : torch.dtype - Data type for tensors. - translation_scale : float - Scale for random translation (0 for no translation). - - Returns - ------- - RigidTransform - Random transformation. - """ - if device is None: - device = get_default_device() - if dtype is None: - dtype = get_float_dtype() - # Shoemake's uniform random quaternion - u1, u2, u3 = torch.rand(3, dtype=dtype, device=device) - - q = torch.stack( - [ - torch.sqrt(1 - u1) * torch.sin(2 * torch.pi * u2), - torch.sqrt(1 - u1) * torch.cos(2 * torch.pi * u2), - torch.sqrt(u1) * torch.sin(2 * torch.pi * u3), - torch.sqrt(u1) * torch.cos(2 * torch.pi * u3), - ] - ) - - # Reorder to [w, x, y, z] format - q = torch.stack([q[3], q[0], q[1], q[2]]) - - if translation_scale > 0: - t = torch.randn(3, dtype=dtype, device=device) * translation_scale - else: - t = torch.zeros(3, dtype=dtype, device=device) - - return cls(quaternion=q, translation=t) - - # ========================================================================= - # Utility methods - # ========================================================================= - - def detach(self) -> "RigidTransform": - """ - Return detached copy (no gradient tracking). - - Returns - ------- - RigidTransform - Detached transformation. - """ - return RigidTransform( - quaternion=self._quaternion.detach(), - translation=self._translation.detach(), - ) - - def clone(self) -> "RigidTransform": - """ - Return deep copy. - - Returns - ------- - RigidTransform - Cloned transformation. - """ - return RigidTransform( - quaternion=self._quaternion.clone(), - translation=self._translation.clone(), - ) - - def __repr__(self) -> str: - q = self._quaternion - t = self._translation - return ( - f"RigidTransform(\n" - f" quaternion=[{q[0]:.4f}, {q[1]:.4f}, {q[2]:.4f}, {q[3]:.4f}],\n" - f" translation=[{t[0]:.4f}, {t[1]:.4f}, {t[2]:.4f}]\n" - f")" - ) diff --git a/torchref/experimental/alignment/wigner.py b/torchref/experimental/alignment/wigner.py deleted file mode 100644 index 551ba234..00000000 --- a/torchref/experimental/alignment/wigner.py +++ /dev/null @@ -1,506 +0,0 @@ -""" -Pure-PyTorch Wigner small-d and Wigner-D evaluation for the alignment module. - -Conventions (locked, asserted by tests/unit/alignment/test_wigner.py): - - D^l_{m,n}(α, β, γ) = e^{-i m α} · d^l_{m,n}(β) · e^{-i n γ} (Edmonds) - -with the Euler angles paired to the rotation matrix used by -`torchref.experimental.alignment.transform.rotation_matrix_from_euler` — i.e. ZYZ. - -Small-d uses the direct sum formula (Edmonds 4.1.23) with log-factorials so the -recurrence never forms `(2l)!` explicitly: - - d^l_{m,n}(β) = Σ_k (-1)^k · √[(l+m)!(l-m)!(l+n)!(l-n)!] - / [(l+m-k)! · k! · (l-n-k)! · (k+n-m)!] - · cos(β/2)^(2l+m-n-2k) · sin(β/2)^(2k+n-m) - -k runs over the integers that keep every factorial non-negative: - max(0, m-n) ≤ k ≤ min(l+m, l-n). -""" - -from __future__ import annotations - -from dataclasses import dataclass -from typing import List, Optional, Tuple - -import torch - - -def _log_factorial(n: torch.Tensor) -> torch.Tensor: - """log(n!) for non-negative integer tensor.""" - return torch.lgamma(n.to(torch.float64) + 1.0) - - -def _build_half_angle_pow_tables( - beta: torch.Tensor, max_exp: int -) -> Tuple[torch.Tensor, torch.Tensor]: - """Precompute `cos(β/2)^j` and `sin(β/2)^j` for j ∈ [0, max_exp]. - - Why: `torch.pow(tensor, tensor)` takes the slow `exp(log(x) * y)` path - even for integer exponents. The Wigner-d sum gathers cos/sin to powers - in {0, 1, …, 2l} for every (m, n, k) cell; replacing `cos_h ** k` with - a `cumprod`-built table + fancy index is ~10× faster on CPU for L≥16 - and dominates the `small_d_packed` cost when sharing the table across - the L iterations. - - Returns tensors of shape `(*beta.shape, max_exp + 1)`, float64. - """ - beta64 = beta.to(torch.float64) - half = 0.5 * beta64 - cos_h = torch.cos(half).unsqueeze(-1) # (*beta, 1) - sin_h = torch.sin(half).unsqueeze(-1) - ones = torch.ones_like(cos_h) - if max_exp < 1: - return ones, ones - base_cos = cos_h.expand(*beta.shape, max_exp) - base_sin = sin_h.expand(*beta.shape, max_exp) - cos_seq = torch.cat([ones, base_cos], dim=-1) - sin_seq = torch.cat([ones, base_sin], dim=-1) - return torch.cumprod(cos_seq, dim=-1), torch.cumprod(sin_seq, dim=-1) - - -def small_d_block( - l: int, - beta: torch.Tensor, - cos_pow_table: Optional[torch.Tensor] = None, - sin_pow_table: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """ - Evaluate d^l_{m,n}(β) for fixed l, all m,n ∈ [-l, l], batched over β. - - Parameters - ---------- - l : int - Wigner degree. - beta : torch.Tensor (real) - Euler β angle(s), arbitrary shape. Values in [0, π]. - cos_pow_table, sin_pow_table : torch.Tensor, optional - Precomputed tables with `cos_pow_table[..., j] = cos(β/2)^j` and - likewise for sin, shape `(*beta.shape, max_exp+1)` with - `max_exp >= 2*l`. If omitted, built locally. Pass them when looping - over l with shared β (see `small_d_packed`) to avoid repeated - `pow`-via-`exp(log·y)` evaluations. - - Returns - ------- - d : torch.Tensor (real, float64 internally, cast to beta.dtype on return) - Shape (..., 2l+1, 2l+1). `d[..., m+l, n+l] = d^l_{m,n}(β)`. - """ - if l == 0: - out = torch.ones((*beta.shape, 1, 1), dtype=beta.dtype, device=beta.device) - return out - - device = beta.device - out_dtype = beta.dtype - - # Build (or reuse) the pow tables. Local build is cheap (~max_exp small - # ops) so we only skip it when the caller hands us one. - if cos_pow_table is None or sin_pow_table is None: - cos_pow_table, sin_pow_table = _build_half_angle_pow_tables(beta, 2 * l) - max_exp = cos_pow_table.shape[-1] - 1 - assert max_exp >= 2 * l, ( - f"pow table max_exp={max_exp} insufficient for degree l={l}" - ) - - # Precompute log factorials for arguments in [0, 2l]. - n_table = torch.arange(0, 2 * l + 1, device=device) - log_fac = _log_factorial(n_table) # (2l+1,) float64 - - size = 2 * l + 1 - # Build (m, n) index grids: m_idx = m + l, n_idx = n + l, m,n ∈ [-l, l]. - m_grid = torch.arange(-l, l + 1, dtype=torch.int64, device=device) # (size,) - n_grid = m_grid.clone() - M = m_grid.view(size, 1).expand(size, size) # (size, size) - N = n_grid.view(1, size).expand(size, size) - - # k range for each (m, n) pair. - k_lo = torch.clamp(M - N, min=0) # (size, size) - k_hi = torch.minimum(torch.full_like(M, l) + M, torch.full_like(M, l) - N) - # Universal k range across all (m,n): k in [0, 2l]. - K = torch.arange(0, 2 * l + 1, dtype=torch.int64, device=device) - # Build the validity mask for each (m, n, k): - K_mn = K.view(1, 1, -1) - mask = (K_mn >= k_lo.unsqueeze(-1)) & (K_mn <= k_hi.unsqueeze(-1)) # (size, size, 2l+1) - - # Coefficient log( (l+m)!(l-m)!(l+n)!(l-n)! / [(l+m-k)! k! (l-n-k)! (k+n-m)!] )^(1/2) - # Common numerator (depends on m, n only) - L_t = torch.full_like(M, l) - log_num = 0.5 * ( - log_fac[L_t + M] + log_fac[L_t - M] + log_fac[L_t + N] + log_fac[L_t - N] - ) # (size, size) - - # Denominator term per (m, n, k) — guard out-of-range indices with mask. - # Use clamp into [0, 2l] so the index is always valid; result will be masked off. - def _safe_lf(idx): - return log_fac[idx.clamp(min=0, max=2 * l)] - - idx_a = (L_t + M).unsqueeze(-1) - K_mn # (l+m-k) - idx_b = K_mn.expand(size, size, -1) # k - idx_c = (L_t - N).unsqueeze(-1) - K_mn # (l-n-k) - idx_d = K_mn + (N - M).unsqueeze(-1) # (k+n-m) - - log_den = _safe_lf(idx_a) + _safe_lf(idx_b) + _safe_lf(idx_c) + _safe_lf(idx_d) - log_coef = log_num.unsqueeze(-1) - log_den # (size, size, 2l+1) - - coef = torch.exp(log_coef) - sign = torch.where((K_mn % 2 == 0), torch.ones_like(coef), -torch.ones_like(coef)) - # Zero out invalid k entries - coef = torch.where(mask, sign * coef, torch.zeros_like(coef)) - - # Per-k exponents. In the *valid* (mask=True) region these lie in [0, 2l]; - # invalid entries can fall outside, so we clamp to [0, max_exp] and rely - # on coef=0 to nuke their contribution. - exp_cos = 2 * l + M.unsqueeze(-1) - N.unsqueeze(-1) - 2 * K_mn # (size, size, 2l+1) - exp_sin = 2 * K_mn + N.unsqueeze(-1) - M.unsqueeze(-1) - exp_cos_safe = exp_cos.clamp(min=0, max=max_exp) - exp_sin_safe = exp_sin.clamp(min=0, max=max_exp) - - # Gather cos(β/2)^exp_cos and sin(β/2)^exp_sin from the precomputed - # tables via index_select. Faster + cleaner dispatch than fancy - # indexing (`table[..., idx]`) on CPU. Flatten the (size, size, 2l+1) - # index into 1D, gather along the last axis, then reshape back. - flat_idx_cos = exp_cos_safe.reshape(-1) - flat_idx_sin = exp_sin_safe.reshape(-1) - cos_pow = cos_pow_table.index_select(-1, flat_idx_cos).reshape( - *cos_pow_table.shape[:-1], *exp_cos_safe.shape - ) - sin_pow = sin_pow_table.index_select(-1, flat_idx_sin).reshape( - *sin_pow_table.shape[:-1], *exp_sin_safe.shape - ) - - out64 = (coef * cos_pow * sin_pow).sum(dim=-1) # (*beta, size, size) - - return out64.to(out_dtype) - - -def small_d_table(L: int, beta: torch.Tensor) -> Tuple[torch.Tensor, ...]: - """ - Compute d^l_{m,n}(β) for all l ∈ [0, L), batched over β. - - Returned as a list of tensors of shape (..., 2l+1, 2l+1) — variable in - final two dims because the small-d matrix for degree l has size 2l+1. - - Use `small_d_packed` if you want a single dense (L, 2L-1, 2L-1) tensor with - zero-padding for the off-diagonal entries beyond |m|, |n| > l. - """ - return tuple(small_d_block(l, beta) for l in range(L)) - - -def small_d_packed(L: int, beta: torch.Tensor) -> torch.Tensor: - """ - Compute the small-d matrices for all l ∈ [0, L), packed into a single - dense tensor of shape (..., L, 2L-1, 2L-1) with zero padding for entries - where |m| > l or |n| > l. - - `d_packed[..., l, L-1+m, L-1+n] = d^l_{m,n}(β)` if |m|, |n| ≤ l, else 0. - - Internally builds the cos/sin half-angle pow tables once (shared across - all L iterations) so each `small_d_block` call gathers from a table - instead of running a tensor-exponent `pow`. A fully vectorised-over-l - implementation would need a (n_beta, L, 2L-1, 2L-1, 2L-1) intermediate - that at L=32, n_beta=64 is multi-GB — the shared pow table buys most - of the speedup while staying memory-bounded. - """ - if L <= 0: - raise ValueError(f"L must be >= 1, got {L}") - out = torch.zeros((*beta.shape, L, 2 * L - 1, 2 * L - 1), - dtype=beta.dtype, device=beta.device) - max_exp = max(2 * (L - 1), 0) - cos_pow_table, sin_pow_table = _build_half_angle_pow_tables(beta, max_exp) - for l in range(L): - d_l = small_d_block( - l, beta, - cos_pow_table=cos_pow_table, - sin_pow_table=sin_pow_table, - ) # (..., 2l+1, 2l+1) - out[..., l, L - 1 - l : L - 1 + l + 1, L - 1 - l : L - 1 + l + 1] = d_l - return out - - -def wigner_D_pointwise( - alpha: torch.Tensor, - beta: torch.Tensor, - gamma: torch.Tensor, - L: int, -) -> torch.Tensor: - """ - Evaluate `D^l_{m,n}(α, β, γ) = e^{-imα} d^l_{m,n}(β) e^{-inγ}` for all - l, m, n with l < L, |m|, |n| ≤ L-1, batched over the (α, β, γ) triples. - - Returns - ------- - D : torch.Tensor (complex), shape (..., L, 2L-1, 2L-1) - """ - assert alpha.shape == beta.shape == gamma.shape - - real_dtype = beta.dtype - if real_dtype == torch.float64: - complex_dtype = torch.complex128 - elif real_dtype == torch.float32: - complex_dtype = torch.complex64 - else: - raise TypeError(f"Unsupported dtype {real_dtype}") - - device = beta.device - d = small_d_packed(L, beta) # (..., L, 2L-1, 2L-1) real - m_vals = torch.arange(-(L - 1), L, dtype=real_dtype, device=device) - n_vals = m_vals.clone() - - # phase_m[..., m_idx] = e^{-i m α}, phase_n[..., n_idx] = e^{-i n γ} - ma = alpha.unsqueeze(-1) * m_vals - ng = gamma.unsqueeze(-1) * n_vals - phase_m = torch.complex(torch.cos(-ma), torch.sin(-ma)) # (..., 2L-1) - phase_n = torch.complex(torch.cos(-ng), torch.sin(-ng)) # (..., 2L-1) - - # D[..., l, m_idx, n_idx] = d[..., l, m_idx, n_idx] · phase_m[..., m_idx] · phase_n[..., n_idx] - # Broadcast shapes: d is (..., L, 2L-1, 2L-1); need phase_m as (..., 1, 2L-1, 1) - # and phase_n as (..., 1, 1, 2L-1). - D = d.to(complex_dtype) * phase_m[..., None, :, None] * phase_n[..., None, None, :] - return D - - -def evaluate_rotation_function_grid( - xi_lmn: torch.Tensor, - L: int, - n_alpha: Optional[int] = None, - n_beta: Optional[int] = None, - n_gamma: Optional[int] = None, -) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - """ - Evaluate `C(α, β, γ) = Σ_l Σ_{m,n} ξ_{l,m,n} D^l_{m,n}(α, β, γ)` on a - uniform Euler grid, via per-β contraction + 2D IFFT in (α, γ). - - The mathematical identity used: - C(α, β, γ) = Σ_{m,n} M_{m,n}(β) · e^{-i m α} · e^{-i n γ} - with M_{m,n}(β) := Σ_l ξ_{l,m,n} · d^l_{m,n}(β). - For each β, M_{m,n}(β) is a (2L-1)×(2L-1) matrix; the (α, γ) dependence is - a 2-D Fourier series, evaluated on a regular (n_α, n_γ) grid via IFFT. - - Parameters - ---------- - xi_lmn : torch.Tensor (complex) - Wigner coefficients, shape (L, 2L-1, 2L-1). Layout: - `xi_lmn[l, L-1+m, L-1+n] = ξ_{l,m,n}` for |m|, |n| ≤ l, else expected zero. - L : int - SH / Wigner bandlimit. - n_alpha, n_beta, n_gamma : int, optional - Grid sizes in α, β, γ. Defaults: n_alpha = n_gamma = 2L (oversampled - FFT grid), n_beta = 2L (midpoint quadrature in β). - - Returns - ------- - C : torch.Tensor (complex) - Shape (n_gamma, n_beta, n_alpha). Real-valued in exact arithmetic; the - imaginary part is returned for diagnostics. Layout: C[k_γ, k_β, k_α]. - alpha_grid, beta_grid, gamma_grid : torch.Tensor (real) - 1-D grids in radians. - """ - if n_alpha is None: - n_alpha = 2 * L - if n_gamma is None: - n_gamma = 2 * L - if n_beta is None: - n_beta = 2 * L - - device = xi_lmn.device - real_dtype = torch.float64 if xi_lmn.dtype == torch.complex128 else torch.float32 - complex_dtype = xi_lmn.dtype - - # Grids - alpha_grid = (2.0 * torch.pi / n_alpha) * torch.arange(n_alpha, dtype=real_dtype, device=device) - gamma_grid = (2.0 * torch.pi / n_gamma) * torch.arange(n_gamma, dtype=real_dtype, device=device) - # β: midpoint rule on (0, π). - beta_grid = (torch.pi * (torch.arange(n_beta, dtype=real_dtype, device=device) + 0.5) - / n_beta) - - # Build M_{m,n}(β_k) = Σ_l ξ_{l,m,n} d^l_{m,n}(β_k) for ALL β_k at once. - # `small_d_packed` accepts a batched β tensor and returns - # (n_beta, L, 2L-1, 2L-1); collapsing the previous `for kb in range(n_beta):` - # loop into a single call eliminates ~64× of Python+torch dispatch - # overhead (the dominant cost in this stage on CPU). - d_all = small_d_packed(L, beta_grid) # (n_beta, L, 2L-1, 2L-1) - d_all_c = d_all.to(complex_dtype) - M_all = (xi_lmn.unsqueeze(0) * d_all_c).sum(dim=1) # (n_beta, 2L-1, 2L-1) - - # M_{m,n}(β) gives Fourier coefficients in (-m·α, -n·γ): - # C(α, γ | β) = Σ_{m,n} M_{m,n} e^{-i m α} e^{-i n γ} - # Build a zero-padded (n_beta, n_alpha, n_gamma) coefficient grid by - # placing each M_{m,n}(β) entry at index (m mod n_alpha, n mod n_gamma); - # torch.fft.fft2 then yields C(α_k, γ_j | β) with the correct sign - # (`fft` uses exp(-2π i k n / N) which matches e^{-i m α}). - Mhat = torch.zeros( - (n_beta, n_alpha, n_gamma), dtype=complex_dtype, device=device, - ) - m_idx = torch.arange(-(L - 1), L, device=device) % n_alpha # (2L-1,) - n_idx = torch.arange(-(L - 1), L, device=device) % n_gamma # (2L-1,) - # Vectorised scatter: M_all[:, m+L-1, n+L-1] → Mhat[:, m_idx, n_idx]. - Mhat[:, m_idx.unsqueeze(-1), n_idx.unsqueeze(0)] = M_all - - # Batched 2-D FFT over (α, γ). - slice_C = torch.fft.fft2(Mhat, dim=(-2, -1)) # (n_beta, n_alpha, n_gamma) - # Re-order to (γ, β, α) layout per our convention C[k_γ, k_β, k_α]. - C = slice_C.permute(2, 0, 1).contiguous() # (n_gamma, n_beta, n_alpha) - - return C, alpha_grid, beta_grid, gamma_grid - - -@dataclass -class AdaptiveRotationFunction: - """ - Phaser-style ragged rotation-function grid. - - On SO(3) the natural area element is `sin(β) dα dβ dγ`. A uniform Euler - cube oversamples the polar caps; this structure stores per-β slices of - variable `(qmax_k, pmax_k)` shape with sampling density matching - `pmax(β) = 720/Δ · cos(β/2)` and `qmax(β) = 360/Δ · sin(β/2)` (Phaser - FastRot.cc:92-96). - - Attributes - ---------- - betas : torch.Tensor - Shape `(n_β,)`, midpoint quadrature on `(0, π)`. - slices : list[torch.Tensor] - Length `n_β`. Slice `k` has shape `(qmax_k, pmax_k)` complex, indexed - as `slices[k][k_γ, k_α]` (matches dense convention `C[k_γ, k_β, k_α]`). - alpha_grids, gamma_grids : list[torch.Tensor] - Length `n_β`. Per-slice α and γ grid coordinates in radians. - grid_sampling_deg : float - Phaser's `grid_sampling` argument — target angular resolution in degrees. - """ - - betas: torch.Tensor - slices: List[torch.Tensor] - alpha_grids: List[torch.Tensor] - gamma_grids: List[torch.Tensor] - grid_sampling_deg: float - - def total_samples(self) -> int: - return sum(s.numel() for s in self.slices) - - -def evaluate_rotation_function_grid_adaptive( - xi_lmn: torch.Tensor, - L: int, - grid_sampling_deg: float = 3.0, - n_beta: Optional[int] = None, -) -> AdaptiveRotationFunction: - """ - Evaluate `C(α, β, γ)` on a Phaser-faithful adaptive Euler grid. - - Per β, the (α, γ) sampling density follows - - pmax(β) = max(1, round(720 / grid_sampling_deg · cos(β/2))) - qmax(β) = max(1, round(360 / grid_sampling_deg · sin(β/2))) - - Total sample count ≈ `(720 · 360) / grid_sampling_deg²` — independent of L - and free of the polar duplication that a uniform `(2L)³` grid produces. - - Computes `M_{m,n}(β_k) = Σ_l ξ_{l,m,n} d^l_{m,n}(β_k)` via the existing - batched `small_d_packed`, then for each β does a zero-padded scatter of - `M` into a `(pmax_k, qmax_k)` array and runs `torch.fft.fft2` on that - slice. The scatter is intentional: aliasing past the per-slice Nyquist - is the physically correct behaviour — those frequencies cannot be - resolved at that β. - """ - if n_beta is None: - n_beta = 2 * L - - device = xi_lmn.device - if xi_lmn.dtype == torch.complex128: - real_dtype = torch.float64 - elif xi_lmn.dtype == torch.complex64: - real_dtype = torch.float32 - else: - raise TypeError(f"xi_lmn must be complex, got {xi_lmn.dtype}") - complex_dtype = xi_lmn.dtype - - # β: midpoint rule on (0, π). - beta_grid = (torch.pi * (torch.arange(n_beta, dtype=real_dtype, device=device) + 0.5) - / n_beta) - - # Batched M_{m,n}(β_k) = Σ_l ξ_{l,m,n} d^l_{m,n}(β_k). - d_all = small_d_packed(L, beta_grid).to(complex_dtype) # (n_β, L, 2L-1, 2L-1) - M_all = (xi_lmn.unsqueeze(0) * d_all).sum(dim=1) # (n_β, 2L-1, 2L-1) - - # Per-β IFFT2 with adaptive shape. - half_beta = beta_grid * 0.5 - cos_half = torch.cos(half_beta) - sin_half = torch.sin(half_beta) - pmax_all = torch.clamp( - (720.0 / grid_sampling_deg * cos_half).round().to(torch.long), min=1, - ) - qmax_all = torch.clamp( - (360.0 / grid_sampling_deg * sin_half).round().to(torch.long), min=1, - ) - - m_vals = torch.arange(-(L - 1), L, dtype=torch.long, device=device) # (2L-1,) - n_vals = m_vals.clone() - - slices: List[torch.Tensor] = [] - alpha_grids: List[torch.Tensor] = [] - gamma_grids: List[torch.Tensor] = [] - - for k in range(n_beta): - pmax_k = int(pmax_all[k].item()) - qmax_k = int(qmax_all[k].item()) - - # Scatter M_{m,n} into the (pmax_k, qmax_k) Fourier-coefficient grid. - m_idx = m_vals % pmax_k # (2L-1,) - n_idx = n_vals % qmax_k # (2L-1,) - m_grid = m_idx.unsqueeze(-1).expand(2 * L - 1, 2 * L - 1) - n_grid = n_idx.unsqueeze(0).expand(2 * L - 1, 2 * L - 1) - flat_idx = m_grid * qmax_k + n_grid # (2L-1, 2L-1) - Mhat = torch.zeros((pmax_k, qmax_k), dtype=complex_dtype, device=device) - Mhat.view(-1).index_add_(0, flat_idx.reshape(-1), M_all[k].reshape(-1)) - - # `torch.fft.fft2` uses exp(-2π i k n / N) so positive (m, n) frequencies - # at index (m, n) reconstruct e^{-i m α} e^{-i n γ} on the uniform grid. - C_slice = torch.fft.fft2(Mhat, dim=(-2, -1)) # (pmax_k, qmax_k) - - alpha_grid_k = (2.0 * torch.pi / pmax_k) * torch.arange( - pmax_k, dtype=real_dtype, device=device, - ) - gamma_grid_k = (2.0 * torch.pi / qmax_k) * torch.arange( - qmax_k, dtype=real_dtype, device=device, - ) - - # Transpose to (γ, α) layout to match the dense `C[k_γ, k_β, k_α]` convention. - slices.append(C_slice.transpose(0, 1).contiguous()) - alpha_grids.append(alpha_grid_k) - gamma_grids.append(gamma_grid_k) - - return AdaptiveRotationFunction( - betas=beta_grid, - slices=slices, - alpha_grids=alpha_grids, - gamma_grids=gamma_grids, - grid_sampling_deg=grid_sampling_deg, - ) - - -def evaluate_rotation_function_pointwise( - xi_lmn: torch.Tensor, - alpha: torch.Tensor, - beta: torch.Tensor, - gamma: torch.Tensor, - L: int, -) -> torch.Tensor: - """ - Evaluate the rotation function at arbitrary Euler triples. Slow but exact - and differentiable — used for sub-voxel peak refinement and convention tests. - - Parameters - ---------- - xi_lmn : torch.Tensor (complex), shape (L, 2L-1, 2L-1) - alpha, beta, gamma : torch.Tensor (real), same shape (..., ) - - Returns - ------- - C : torch.Tensor (complex), shape (...,). In exact arithmetic real for real - input fields, but kept complex so callers can inspect drift. - """ - D = wigner_D_pointwise(alpha, beta, gamma, L) # (..., L, 2L-1, 2L-1) - # xi_lmn has shape (L, 2L-1, 2L-1); broadcasting aligns on the trailing dims. - C = (xi_lmn * D).sum(dim=(-3, -2, -1)) - return C From f777255759ccb88b7d77ea8ba001ad8f1d69aedb Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 31 Aug 2026 14:14:29 +0200 Subject: [PATCH 108/250] Remove the ML rescore and the post-placement refinement The pipeline is now FRF -> FTF and stops at the placement. The rescore sat between the two stages and was measured to make the pipeline worse. End-to-end pose recovery over 10 structures x 3 trials: 18/30 with m_letf1, 24/30 without, McNemar p = 0.031 on 6-0 discordant cells, and 1.33x faster. It only ever reordered a shortlist that already contained truth, and on 4BX9 and 6G9X it reproducibly pushed truth out of it -- those two cells are the entire difference. The premise it was built on is also gone: the translation function puts truth at rank 0 in 27 of 30 cells against the rotation function's 6, so the rotation function does not have to rank well. The dense rotation re-sampling and the LBFGS rigid-body polish go with it. They turned a placement into a refined structure, which is downstream refinement's job and which it does better; keeping them meant carrying a second rigid-body implementation and a second likelihood. What the pipeline returns is now a pose, and the caller refines it. That frees ml_rotation.py, lattman_love.py, rigid_body.py, and two thirds of distributions.py. fit_sigma_a_per_shell moves next to the translation LLG that is its only remaining caller, alongside a note on how it relates to weighting.empirical_sigma_a -- the same quantity, fitted after placement rather than assumed before it. translation.py loses six functions nothing referenced, including the one that made it depend on ml_rotation. pose_recovery's arms become the ranking question that is actually open: analytic R against the translation LLG. RiceXrayTarget is left in place with its docstring corrected -- it existed for the deleted aligner and now has no caller at all, but removing it is a refinement/ change. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/refactor_gate.sh | 21 + .../diagnostics/frf_vs_ftf_discrimination.py | 3 +- alignment_lab/diagnostics/pose_recovery.py | 51 +- tests/helpers/device_cases.py | 1 - tests/integration/alignment/profile_fit.py | 8 +- .../test_interp_var_and_shared_sigma_a.py | 164 --- tests/unit/alignment/test_m_letf1.py | 350 ----- tests/unit/alignment/test_sigma_a_luzzati.py | 55 - torchref/experimental/alignment/__init__.py | 56 +- torchref/experimental/alignment/align.py | 81 +- .../experimental/alignment/distributions.py | 183 --- torchref/experimental/alignment/e_values.py | 4 +- .../experimental/alignment/lattman_love.py | 302 ----- .../experimental/alignment/ml_rotation.py | 1208 ----------------- torchref/experimental/alignment/pipeline.py | 440 +----- torchref/experimental/alignment/rigid_body.py | 552 -------- .../experimental/alignment/translation.py | 534 +------- torchref/refinement/targets/xray/rice.py | 30 +- 18 files changed, 217 insertions(+), 3826 deletions(-) create mode 100644 alignment_lab/analysis/refactor_gate.sh delete mode 100644 tests/unit/alignment/test_interp_var_and_shared_sigma_a.py delete mode 100644 tests/unit/alignment/test_m_letf1.py delete mode 100644 tests/unit/alignment/test_sigma_a_luzzati.py delete mode 100644 torchref/experimental/alignment/lattman_love.py delete mode 100644 torchref/experimental/alignment/ml_rotation.py delete mode 100644 torchref/experimental/alignment/rigid_body.py diff --git a/alignment_lab/analysis/refactor_gate.sh b/alignment_lab/analysis/refactor_gate.sh new file mode 100644 index 00000000..8065bfdc --- /dev/null +++ b/alignment_lab/analysis/refactor_gate.sh @@ -0,0 +1,21 @@ +#!/bin/bash +#SBATCH --job-name=refgate +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +export TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 +echo "== import smoke ==" +"$PY" -m pytest -c tests/pytest.ini tests/unit/test_imports_smoke.py --run-slow -q 2>&1 | tail -5 +echo "SMOKE_RC=${PIPESTATUS[0]}" +echo "== alignment + frf_separate + scaling ==" +"$PY" -m pytest -c tests/pytest.ini tests/unit/alignment tests/unit/frf_separate tests/unit/scaling --run-slow -q 2>&1 | tail -20 +echo "SCOPED_RC=${PIPESTATUS[0]}" diff --git a/alignment_lab/diagnostics/frf_vs_ftf_discrimination.py b/alignment_lab/diagnostics/frf_vs_ftf_discrimination.py index 74723e8f..aa91846e 100644 --- a/alignment_lab/diagnostics/frf_vs_ftf_discrimination.py +++ b/alignment_lab/diagnostics/frf_vs_ftf_discrimination.py @@ -132,8 +132,7 @@ def main() -> int: ) frf = _prepare_frf_inputs( model, data, d_min=pipe.d_min, d_max=pipe.d_max, - n_shells=pipe.n_shells, ll_padding_factor=pipe.ll_padding_factor, - ll_max_res_A=pipe.ll_max_res_A, verbose=0, + n_shells=pipe.n_shells, verbose=0, ) pipe._frf = frf peaks = pipe._rotation_candidates(frf)[: args.n_cand] diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py index 6087be1e..5a1ad3e0 100644 --- a/alignment_lab/diagnostics/pose_recovery.py +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -1,21 +1,22 @@ -"""End-to-end pose recovery: does dropping the ML rescore cost anything? +"""End-to-end pose recovery for the FRF -> FTF pipeline. -The rank-level evidence says the rescore only reorders (it leaves every peak's -orientation untouched), that final solutions are selected by R-factor rather -than by the rescore's own score, and that its reordering does not improve -whether truth reaches the translation stage. If all that holds, removing it -should be free -- but rank is not the deliverable, pose is, so this measures the -full pipeline. +Rank is not the deliverable, pose is. This places a randomly reoriented copy of +the deposited model and asks whether the pipeline gets it back, which is the +only measurement that settles a change to either stage. -Arms (``--arms``): +The reference number to beat is **24/30** (10 structures x 3 trials): what the +pipeline scored once the ML rescore was taken out of the middle, against 18/30 +with it. 2DQ6 and 3GR5 fail in every arm ever measured and cap recovery there. -``m_letf1`` - current default. -``none`` - skip the rescore; raw FRF order straight into the translation search. -``none+subpeak`` - the same, plus quadratic sub-peak refinement -- sharpen each orientation - in place without reordering. +Arms (``--arms``) sweep how translation candidates are ranked: + +``analytic_r`` + the default -- rank each rotation candidate by the analytical-scale R at its + best translation. +``llg_tf`` + re-rank the translation peaks by the Rice/Woolfson LLG first. At rank level + the LLG puts truth at rank 0 in 27/30 against analytic R's 22/30; this is + the arm that says whether that carries through to pose. Success mirrors the integration test: final coordinates within ``--success-deg`` of canonical, modulo the crystal symmetry. @@ -23,7 +24,7 @@ Usage:: python alignment_lab/diagnostics/pose_recovery.py --pdb 1DAW --trial 0 \ - --arms m_letf1,none,none+subpeak --out-csv alignment_lab/runs/pose.csv + --arms analytic_r,llg_tf --out-csv alignment_lab/runs/pose.csv """ from __future__ import annotations @@ -45,11 +46,8 @@ seed_for) ARMS = { - "m_letf1": dict(rescore_engine="m_letf1", subpeak_refine=False), - "sim": dict(rescore_engine="sim", subpeak_refine=False), - "none": dict(rescore_engine="none", subpeak_refine=False), - "none+subpeak": dict(rescore_engine="none", subpeak_refine=True), - "m_letf1+subpeak": dict(rescore_engine="m_letf1", subpeak_refine=True), + "analytic_r": dict(use_llg_tf=False), + "llg_tf": dict(use_llg_tf=True), } @@ -86,7 +84,7 @@ def main() -> int: ap = argparse.ArgumentParser(description=__doc__) ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) ap.add_argument("--trial", type=int, default=0) - ap.add_argument("--arms", default="m_letf1,none,none+subpeak") + ap.add_argument("--arms", default="analytic_r,llg_tf") ap.add_argument("--n-rotation-candidates", type=int, default=15) ap.add_argument("--n-rotation-peaks", type=int, default=200) ap.add_argument("--success-deg", type=float, default=8.0) @@ -109,7 +107,7 @@ def main() -> int: writer = None if args.out_csv: writer = ResultWriter(args.out_csv, "pose_recovery", - extra_fields=("arm", "rescore_engine", "subpeak_refine", + extra_fields=("arm", "use_llg_tf", "residual_deg", "success", "n_rotation_candidates", "pipeline_seconds")) @@ -127,8 +125,8 @@ def main() -> int: try: aligned = align_model_to_data( search, data, d_min=4.0, d_max=15.0, n_shells=20, - n_rotation_peaks=args.n_rotation_peaks, n_ml_refine=200, - do_translation=True, do_joint_refine=True, + n_rotation_peaks=args.n_rotation_peaks, + do_translation=True, n_rotation_candidates=args.n_rotation_candidates, verbose=args.verbose, **flags, ) @@ -152,8 +150,7 @@ def main() -> int: orbit_side="kabsch", orbit_frame="cart", lmax_cap=_LMAX_CAP, d_min=4.0, d_max=15.0, device="cpu", arm=arm, - rescore_engine=flags["rescore_engine"], - subpeak_refine=int(flags["subpeak_refine"]), + use_llg_tf=int(flags["use_llg_tf"]), residual_deg=(round(resid, 4) if resid == resid else ""), success=int(bool(ok)), n_rotation_candidates=args.n_rotation_candidates, diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index 5d65eaef..5021f5ca 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -485,7 +485,6 @@ class TargetDeviceCase: "RiceDifferenceTarget": "needs two datasets", "TaylorCorrectedDifferenceTarget": "needs two datasets", "RigidTransform": "alignment helper; needs a coordinate set", - "RigidBodyRefinement": "experimental; needs model + data", } # Everything under torchref/experimental is out of scope for the conformance diff --git a/tests/integration/alignment/profile_fit.py b/tests/integration/alignment/profile_fit.py index 6c4bf6f6..c3202f61 100644 --- a/tests/integration/alignment/profile_fit.py +++ b/tests/integration/alignment/profile_fit.py @@ -52,8 +52,7 @@ def _stage(name: str): def _patch_for_timing(): """Wrap key fit_to_data stages so we get an inline breakdown.""" - from torchref.experimental.alignment import ml_rotation, translation - from torchref.experimental.alignment import lattman_love, rigid_body + from torchref.experimental.alignment import rotation_search, translation from torchref import scaling originals = {} @@ -68,12 +67,11 @@ def wrapper(*args, **kwargs): setattr(module, attr, wrapper) - wrap(ml_rotation, "m_letf1_rescore", "m_letf1_rescore") + wrap(rotation_search, "search_peaks", "rotation_search") wrap(translation, "amplitude_translation_search", "amplitude_translation_search") wrap(translation, "local_translation_refine", "local_translation_refine") wrap(translation, "precompute_G_for_rotation", "precompute_G_for_rotation") - wrap(lattman_love, "LattmanLoveInterpolator", "LL_interp_build") - wrap(rigid_body, "RigidBodyRefinement", "RigidBodyRefinement_init") + wrap(translation, "llg_translation_rescore", "llg_translation_rescore") return originals diff --git a/tests/unit/alignment/test_interp_var_and_shared_sigma_a.py b/tests/unit/alignment/test_interp_var_and_shared_sigma_a.py deleted file mode 100644 index 91142727..00000000 --- a/tests/unit/alignment/test_interp_var_and_shared_sigma_a.py +++ /dev/null @@ -1,164 +0,0 @@ -""" -Unit tests for Phase A additions: -- `fit_sigma_a_per_shell`: vectorised per-shell σA fit. -- `_shell_ll` / `llg_for_rotation_batch` interp_var plumbing. -- `llg_for_rotation_batch` shared `sigma_a` path. - -The interp_var rescue mechanism: variance inflation must REDUCE the LLG -penalty on a noisy-but-correct rotation, so the LLG curvature around the -true peak is gentler — that's the fix for the rescore-demotion failure. -""" - -import math - -import pytest -import torch - -from torchref.experimental.alignment.ml_rotation import ( - _equal_count_shell_idx, - _optimize_D_in_shell, - _shell_ll, - fit_sigma_a_per_shell, - llg_for_rotation_batch, -) - - -def _make_synthetic_data(N=600, n_shells=10, true_D=0.6, seed=0): - """Synthetic E_obs / E_calc with a known per-shell σA structure.""" - g = torch.Generator().manual_seed(seed) - s_mag = torch.linspace(0.05, 0.5, N) - shell_idx = _equal_count_shell_idx(s_mag, n_shells) - # E_calc ~ Rayleigh(1) (matches a normalized acentric model) - E_calc = torch.empty(N).exponential_(generator=g).sqrt() - # E_obs = sqrt((D·E_calc)² + (1-D²)·noise²) with noise ~ Rayleigh - noise = torch.empty(N).exponential_(generator=g).sqrt() - var = max(1.0 - true_D ** 2, 1e-4) - E_obs = ((true_D * E_calc) ** 2 + var * noise ** 2).sqrt() - centric = torch.zeros(N, dtype=torch.bool) - return E_obs, E_calc, centric, shell_idx - - -def test_fit_sigma_a_per_shell_matches_per_shell_golden(): - """Vectorised per-shell σA should match per-shell golden-section within grid resolution.""" - E_obs, E_calc, centric, shell_idx = _make_synthetic_data( - N=1500, n_shells=12, true_D=0.55, seed=1, - ) - n_shells = int(shell_idx.max().item()) + 1 - - # Vectorised fit - sigma_a_vec = fit_sigma_a_per_shell( - E_obs, E_calc, centric, shell_idx, n_shells, n_grid=81, - ) - - # Per-shell golden-section reference - sigma_a_ref = torch.zeros(n_shells) - for k in range(n_shells): - mask = shell_idx == k - if mask.sum() < 5: - continue - sigma_a_ref[k] = _optimize_D_in_shell( - E_obs[mask], E_calc[mask], centric[mask], - ) - - # 81-pt grid resolution = 0.99/80 ≈ 0.012. Loose tolerance because the - # vectorised path doesn't refine via golden section. - valid = sigma_a_ref > 0 - diff = (sigma_a_vec - sigma_a_ref)[valid].abs().max().item() - assert diff < 0.025, f"per-shell σA mismatch {diff:.4f}" - - -def test_interp_var_changes_llg_monotonically(): - """ - Sanity: adding a positive uniform interp_var changes the per-shell LL - in a smooth, monotonic way — no NaN/inf, no sign flips for moderate - inflation. Whether interp_var *helps* the rescue is an integration-level - claim that depends on resolution distribution of the variance — tested - on the live sweep, not here. - """ - E_obs, E_calc, centric, shell_idx = _make_synthetic_data( - N=1000, n_shells=10, true_D=0.6, seed=2, - ) - n_shells = int(shell_idx.max().item()) + 1 - F_calc = E_calc.unsqueeze(0) - sigma_a = fit_sigma_a_per_shell( - E_obs, E_calc, centric, shell_idx, n_shells, n_grid=81, - ) - - llgs = [] - for iv_scale in [0.0, 0.05, 0.1, 0.2]: - iv = torch.full_like(E_obs, iv_scale) - llg = llg_for_rotation_batch( - F_obs=E_obs, shell_idx=shell_idx, n_shells=n_shells, - E_obs=E_obs, centric=centric, F_calc=F_calc, - sigma_a=sigma_a, interp_var=iv, - ).item() - assert math.isfinite(llg), f"interp_var={iv_scale} gave LLG={llg}" - llgs.append(llg) - # Strictly monotone in some direction (no oscillation). - diffs = [llgs[i + 1] - llgs[i] for i in range(len(llgs) - 1)] - same_sign = all(d * diffs[0] >= 0 for d in diffs) - assert same_sign, f"non-monotonic LLG vs interp_var: {llgs}" - - -def test_shared_sigma_a_path_matches_per_shell_grid_at_optimum(): - """ - Passing sigma_a (the per-shell argmax of the grid) reproduces the - grid-search LLG at the optimum within numerical tolerance. - """ - E_obs, E_calc, centric, shell_idx = _make_synthetic_data( - N=1000, n_shells=10, true_D=0.6, seed=4, - ) - n_shells = int(shell_idx.max().item()) + 1 - F_calc_batch = E_calc.unsqueeze(0) - - # Run grid path → per-shell argmax - sigma_a = fit_sigma_a_per_shell( - E_obs, E_calc, centric, shell_idx, n_shells, n_grid=81, - ) - - llg_grid = llg_for_rotation_batch( - F_obs=E_obs, shell_idx=shell_idx, n_shells=n_shells, - E_obs=E_obs, centric=centric, F_calc=F_calc_batch, - n_D_grid=81, - ).item() - - llg_shared = llg_for_rotation_batch( - F_obs=E_obs, shell_idx=shell_idx, n_shells=n_shells, - E_obs=E_obs, centric=centric, F_calc=F_calc_batch, - sigma_a=sigma_a, - ).item() - - # When sigma_a IS the per-shell grid argmax, shared and grid paths match - # to within a few percent (small differences arise because the grid path - # picks D per (shell, batch) before per-shell summation; shared uses - # fixed D per shell). - assert abs(llg_grid - llg_shared) < 0.01 * max(abs(llg_grid), 1.0), ( - f"grid LLG = {llg_grid:.4f} shared LLG = {llg_shared:.4f}" - ) - - -def test_shell_ll_interp_var_off_default_matches_old(): - """interp_var=None must reproduce the historical _shell_ll behaviour exactly.""" - torch.manual_seed(5) - N = 300 - E_obs = torch.empty(N).exponential_().sqrt() - E_calc = torch.empty(N).exponential_().sqrt() - centric = torch.zeros(N, dtype=torch.bool) - D = 0.4 - ll_legacy = _shell_ll(E_obs, E_calc, centric, D) - ll_default = _shell_ll(E_obs, E_calc, centric, D, interp_var=None) - torch.testing.assert_close(ll_legacy, ll_default) - - -def test_shell_ll_interp_var_zero_matches_off(): - """interp_var of all-zeros should also match the off path (no inflation).""" - torch.manual_seed(6) - N = 250 - E_obs = torch.empty(N).exponential_().sqrt() - E_calc = torch.empty(N).exponential_().sqrt() - centric = torch.zeros(N, dtype=torch.bool) - D = 0.3 - ll_off = _shell_ll(E_obs, E_calc, centric, D) - iv0 = torch.zeros_like(E_obs) - ll_zero = _shell_ll(E_obs, E_calc, centric, D, interp_var=iv0) - torch.testing.assert_close(ll_off, ll_zero, atol=1e-6, rtol=1e-6) diff --git a/tests/unit/alignment/test_m_letf1.py b/tests/unit/alignment/test_m_letf1.py deleted file mode 100644 index ccb16d41..00000000 --- a/tests/unit/alignment/test_m_letf1.py +++ /dev/null @@ -1,350 +0,0 @@ -"""Unit tests for the Phaser-faithful m_LETF1 rescore. - -Covers: -- ``phaser_log_rel_rice`` / ``phaser_log_rel_woolfson`` exact formula parity - with hand-computed values from RiceWoolfson.cc:25-74. -- ``compute_v_budget`` per DataMR.cc:949,1411 (clamping, degenerate cases). -- ``m_letf1_rescore`` discriminates the truth rotation over random rotations - on a synthetic obs+calc set. -""" -from __future__ import annotations - -import math - -import pytest -import torch - -from torchref.experimental.alignment.distributions import ( - phaser_log_rel_rice, - phaser_log_rel_woolfson, -) -from torchref.experimental.alignment.frf.preprocessing import compute_v_budget - - -def test_phaser_log_rel_rice_hand_computed(): - """logRelRice(F1, DF2, V) = logI₀(2·F1·DF2/V) − log V − (F1²+DF2²)/V.""" - # F1=1, DF2=1, V=1: logI₀(2) − 0 − 2 ≈ 0.8237 − 2 = −1.1763 - val = phaser_log_rel_rice(torch.tensor(1.0), torch.tensor(1.0), torch.tensor(1.0)) - assert abs(val.item() - (-1.1763)) < 1e-3 - - # F1=2, DF2=0.5, V=2: logI₀(1.0) − log 2 − (4 + 0.25)/2 - # = 0.2359 − 0.6931 − 2.125 ≈ −2.582 - val = phaser_log_rel_rice(torch.tensor(2.0), torch.tensor(0.5), torch.tensor(2.0)) - assert abs(val.item() - (-2.582)) < 1e-3 - - -def test_phaser_log_rel_woolfson_hand_computed(): - """logRelWoolfson(F1, DF2, V) = log cosh(F1·DF2/V) − ½·log V − (F1²+DF2²)/(2V).""" - # F1=1, DF2=1, V=1: log cosh(1) − 0 − 1 ≈ 0.4339 − 1 = −0.5661 - val = phaser_log_rel_woolfson(torch.tensor(1.0), torch.tensor(1.0), torch.tensor(1.0)) - assert abs(val.item() - (-0.5661)) < 1e-3 - - # F1=3, DF2=2, V=4: log cosh(1.5) − ½·log 4 − (9+4)/8 - # = 0.8553 − 0.6931 − 1.625 ≈ −1.463 - val = phaser_log_rel_woolfson(torch.tensor(3.0), torch.tensor(2.0), torch.tensor(4.0)) - assert abs(val.item() - (-1.463)) < 1e-3 - - -def test_phaser_log_rel_woolfson_large_arg_stable(): - """log cosh(x) ≈ |x| − log 2 for large |x|; no overflow.""" - val = phaser_log_rel_woolfson( - torch.tensor(50.0), torch.tensor(50.0), torch.tensor(1.0), - ) - assert torch.isfinite(val) - # log cosh(2500) ≈ 2500 − log 2; minus 0 minus (2500+2500)/2 = 2500 - # → ≈ 2500 − 0.693 − 2500 = −0.693 - assert abs(val.item() - (-math.log(2.0))) < 1e-2 - - -def test_compute_v_budget_basic(): - """V(h) = ε(h) − σ_A²(s)·n_mol, clamped > 0.""" - eps = torch.tensor([1.0, 2.0, 1.0]) - sa = torch.tensor([0.3, 0.5, 0.4]) - V = compute_v_budget(eps, sa, n_mol=2) - # h0: 1 − 0.09·2 = 0.82 - # h1: 2 − 0.25·2 = 1.50 - # h2: 1 − 0.16·2 = 0.68 - assert torch.allclose(V, torch.tensor([0.82, 1.50, 0.68]), atol=1e-6) - - -def test_compute_v_budget_clamps_degenerate(): - """Non-positive V (σ_A²·n_mol overshoots ε) clamps to 1e-6, not negative.""" - eps = torch.tensor([1.0]) - sa = torch.tensor([0.9]) # 0.81·2 = 1.62 > 1.0 - V = compute_v_budget(eps, sa, n_mol=2) - assert V.item() > 0 - assert V.item() <= 1e-5 # clamp floor - - -def test_compute_v_budget_with_totvar_known(): - """totvar_known subtracts further from V.""" - eps = torch.tensor([1.0, 1.0]) - sa = torch.tensor([0.1, 0.1]) - totvar = torch.tensor([0.05, 0.30]) - V = compute_v_budget(eps, sa, n_mol=1, totvar_known=totvar) - # h0: 1 − 0.01 − 0.05 = 0.94 - # h1: 1 − 0.01 − 0.30 = 0.69 - assert torch.allclose(V, torch.tensor([0.94, 0.69]), atol=1e-6) - - -def test_m_letf1_rescore_runs_and_ranks_truth_top(): - """Synthetic test: m_letf1_rescore should rank the truth rotation at #1 - (or near #1) over random rotations when given matching obs/calc. - - Setup: build a fake LattmanLoveInterpolator-like callable that returns - ``|F_obs|`` at the identity rotation and noisy values at others, then check - that the identity-rotation peak comes top of the rescored list. - """ - pytest.importorskip("torchref.experimental.alignment.lattman_love") - from torchref.experimental.alignment.ml_rotation import m_letf1_rescore - from torchref.experimental.alignment.frf.types import RotationPeak - - N = 200 - torch.manual_seed(0) - s_mag = torch.linspace(0.05, 0.4, N, dtype=torch.float64) - F_obs = (1.0 + 0.1 * torch.randn(N, dtype=torch.float64)).abs() - hkl = torch.randint(-10, 10, (N, 3), dtype=torch.long) - centric = torch.zeros(N, dtype=torch.bool) - from torchref.symmetry import SpaceGroup - sg = SpaceGroup("P 1") # only the identity, so epsilon is 1 throughout - - # Stub interpolator: returns F_obs (perfectly correlated) for R = identity, - # uncorrelated noise for any other R. - class StubLL: - def evaluate(self, R, hkl_real, real_cell, return_amplitude=True): - if R.dim() == 2: - R = R.unsqueeze(0) - B = R.shape[0] - out = torch.empty(B, hkl_real.shape[0], dtype=torch.float64) - for b in range(B): - if torch.allclose(R[b].to(torch.float64), torch.eye(3, dtype=torch.float64), atol=1e-3): - out[b] = F_obs - else: - out[b] = (1.0 + 0.1 * torch.randn(hkl_real.shape[0], dtype=torch.float64)).abs() - return out - - class StubCell: - @property - def reciprocal_basis_matrix(self): - return torch.eye(3, dtype=torch.float64) - - # Truth peak (identity rotation: α=β=γ=0) plus 9 random peaks. - truth = RotationPeak(alpha=0.0, beta=0.0, gamma=0.0, score=1.0, sigma=5.0) - rng = torch.Generator().manual_seed(1) - random_peaks = [ - RotationPeak( - alpha=float(torch.rand(1, generator=rng).item() * 2 * math.pi), - beta=float(torch.rand(1, generator=rng).item() * math.pi), - gamma=float(torch.rand(1, generator=rng).item() * 2 * math.pi), - score=0.5, sigma=2.0, - ) - for _ in range(9) - ] - peaks = [truth] + random_peaks - - rescored = m_letf1_rescore( - peaks, F_obs, hkl, s_mag, centric, StubLL(), StubCell(), sg, - n_shells=5, batch_size=4, - ) - # Truth (identity) should be among the top 3 rescored peaks (truth=identity - # gives perfect calc match; noisy candidates should have lower LL). - top3_eulers = [(p.alpha, p.beta, p.gamma) for p in rescored[:3]] - assert (0.0, 0.0, 0.0) in top3_eulers, ( - f"truth (identity) not in top 3 after rescore; top3 = {top3_eulers}" - ) - - -def test_a_global_calc_convention_preserves_intershell_shape(): - """`CalcGlobalE` uses a single GLOBAL calc scale, so a shell where the model - scatters weakly keeps a small eImove; `CalcShellE` flattens every shell to - unit variance. Verify the two give different (and predictable) eImove - inter-shell weighting on a synthetic with a strong resolution-dependent - F_calc falloff. - - This was `scat_mode="legacy"` vs `"absolute"`. The rescore had two knobs for - one decision -- how obs is normalised and how calc is -- which let the two - sides be normalised by unrelated rules; the E convention answers for both.""" - from torchref.experimental.alignment.e_values import ( - CalcGlobalE, WilsonShellEpsE, - ) - from torchref.experimental.alignment.ml_rotation import ( - _build_llg_context, _llg_for_orientations, - ) - - N = 300 - torch.manual_seed(3) - s_mag = torch.linspace(0.05, 0.45, N, dtype=torch.float64) - F_obs = (1.0 + 0.1 * torch.randn(N, dtype=torch.float64)).abs() - hkl = torch.randint(-12, 12, (N, 3), dtype=torch.long) - centric = torch.zeros(N, dtype=torch.bool) - from torchref.symmetry import SpaceGroup - sg = SpaceGroup("P 1") - - # F_calc with a strong B-factor falloff → big inter-shell amplitude variation. - decay = torch.exp(-40.0 * s_mag * s_mag) - - class StubLL: - def evaluate(self, R, hkl_real, real_cell, return_amplitude=True): - if R.dim() == 2: - R = R.unsqueeze(0) - B = R.shape[0] - # Amplitude depends only on |s| (resolution) → deterministic per hkl; - # uses the input hkl rows' index range to map back to s via norm. - out = decay.unsqueeze(0).expand(B, hkl_real.shape[0]).clone() - return out - - class StubCell: - @property - def reciprocal_basis_matrix(self): - return torch.eye(3, dtype=torch.float64) - - common = dict( - interpolator=StubLL(), real_cell=StubCell(), spacegroup=sg, - n_shells=6, batch_size=64, - ) - ctx_leg = _build_llg_context(F_obs, hkl, s_mag, centric, - e_convention=WilsonShellEpsE, **common) - ctx_abs = _build_llg_context(F_obs, hkl, s_mag, centric, - e_convention=CalcGlobalE, **common) - - # The per-shell normaliser varies across shells (tracks the F_calc decay); - # the global one is a single constant. This is the definitional difference. - assert ctx_leg.sqrt_mean_per_m.std() > 1e-6, "per-shell should vary per shell" - assert torch.allclose( - ctx_abs.sqrt_mean_per_m, ctx_abs.sqrt_mean_per_m[0] - ), "the global convention should be a single scale" - # Both modes still produce finite LLGs. - a = torch.zeros(1, dtype=torch.float64) - llg_leg = _llg_for_orientations(ctx_leg, a, a, a) - llg_abs = _llg_for_orientations(ctx_abs, a, a, a) - assert torch.isfinite(llg_leg).all() and torch.isfinite(llg_abs).all() - assert float(llg_leg[0]) != float(llg_abs[0]), "modes should differ" - - -def _so3_angle_deg(R1: torch.Tensor, R2: torch.Tensor) -> float: - """Geodesic distance on SO(3) in degrees.""" - R = R1.to(torch.float64) @ R2.to(torch.float64).T - tr = (R[0, 0] + R[1, 1] + R[2, 2]).item() - return math.degrees(math.acos(max(-1.0, min(1.0, (tr - 1.0) / 2.0)))) - - -def test_quadratic_refine_lands_closer_to_known_max(monkeypatch): - """The vertex of the quadratic fit on a concave LLG surface lands strictly - closer to the known maximum orientation than the starting grid peak.""" - from torchref.experimental.alignment import ml_rotation - from torchref.experimental.alignment.ml_rotation import quadratic_llg_refine - from torchref.experimental.alignment.frf.types import RotationPeak - from torchref.experimental.alignment.frf.rotation_utils import ( - axis_angle_to_matrix, - rotation_matrix_from_edmonds_euler, - rotation_matrix_from_edmonds_euler_batch, - ) - - # Start orientation (away from the β poles) and a known max 0.5° away. - a0, b0, g0 = 0.5, 0.8, 1.2 - R0 = rotation_matrix_from_edmonds_euler(a0, b0, g0) - d = torch.tensor([math.radians(0.5), 0.0, 0.0], dtype=torch.float64) - R_true = axis_angle_to_matrix(d) @ R0 - - # Synthetic concave surface: LLG = -k · geodesic_angle(R, R_true)². - def stub(ctx, alpha, beta, gamma): - R = rotation_matrix_from_edmonds_euler_batch(alpha, beta, gamma) - Rrel = torch.einsum("mij,lj->mil", R, R_true) # R · R_trueᵀ - tr = Rrel[:, 0, 0] + Rrel[:, 1, 1] + Rrel[:, 2, 2] - ang = torch.arccos(((tr - 1.0) / 2.0).clamp(-1.0, 1.0)) - return -100.0 * ang * ang - - monkeypatch.setattr(ml_rotation, "_llg_for_orientations", stub) - - peak = RotationPeak(alpha=a0, beta=b0, gamma=g0, score=0.0, sigma=0.0) - refined = quadratic_llg_refine([peak], ctx=None, k_refine=1, step_deg=0.75, n_grid=3) - Rref = rotation_matrix_from_edmonds_euler( - refined[0].alpha, refined[0].beta, refined[0].gamma, - ) - d_start = _so3_angle_deg(R0, R_true) # ≈ 0.5° - d_ref = _so3_angle_deg(Rref, R_true) - assert d_ref < d_start, f"refine did not improve: {d_ref:.3f} vs {d_start:.3f}" - assert d_ref < 0.1, f"vertex did not land on the max: {d_ref:.3f}°" - assert refined[0].score > -1.0 # true LLG at the vertex ≈ 0 (near max) - - -def test_quadratic_refine_guard_falls_back_on_minimum(monkeypatch): - """On a convex surface (a minimum, not a max) the negative-definite Hessian - guard rejects the vertex and falls back to the best sampled grid point — so - the refined orientation moves AWAY from R_true, never toward the minimum.""" - from torchref.experimental.alignment import ml_rotation - from torchref.experimental.alignment.ml_rotation import quadratic_llg_refine - from torchref.experimental.alignment.frf.types import RotationPeak - from torchref.experimental.alignment.frf.rotation_utils import ( - axis_angle_to_matrix, - rotation_matrix_from_edmonds_euler, - rotation_matrix_from_edmonds_euler_batch, - ) - - a0, b0, g0 = 0.5, 0.8, 1.2 - R0 = rotation_matrix_from_edmonds_euler(a0, b0, g0) - d = torch.tensor([math.radians(0.5), 0.0, 0.0], dtype=torch.float64) - R_true = axis_angle_to_matrix(d) @ R0 - - # CONVEX surface: +k · angle² (minimum at R_true → Hessian positive-definite). - def stub(ctx, alpha, beta, gamma): - R = rotation_matrix_from_edmonds_euler_batch(alpha, beta, gamma) - Rrel = torch.einsum("mij,lj->mil", R, R_true) - tr = Rrel[:, 0, 0] + Rrel[:, 1, 1] + Rrel[:, 2, 2] - ang = torch.arccos(((tr - 1.0) / 2.0).clamp(-1.0, 1.0)) - return +100.0 * ang * ang - - monkeypatch.setattr(ml_rotation, "_llg_for_orientations", stub) - - peak = RotationPeak(alpha=a0, beta=b0, gamma=g0, score=0.0, sigma=0.0) - refined = quadratic_llg_refine([peak], ctx=None, k_refine=1, step_deg=0.75, n_grid=3) - Rref = rotation_matrix_from_edmonds_euler( - refined[0].alpha, refined[0].beta, refined[0].gamma, - ) - d_start = _so3_angle_deg(R0, R_true) - d_ref = _so3_angle_deg(Rref, R_true) - # Fell back to the grid sample with the largest (convex) LLG = farthest from - # R_true: the guard prevented stepping to the minimum. - assert d_ref > d_start, ( - f"guard failed — moved toward the minimum ({d_ref:.3f} vs {d_start:.3f})" - ) - - -def test_quadratic_refine_max_move_cap(monkeypatch): - """max_move_deg reverts a peak that the (mis-peaked) surface would drag far - from the input — bounding degradation on high-sym/tNCS cases.""" - from torchref.experimental.alignment import ml_rotation - from torchref.experimental.alignment.ml_rotation import quadratic_llg_refine - from torchref.experimental.alignment.frf.types import RotationPeak - from torchref.experimental.alignment.frf.rotation_utils import ( - axis_angle_to_matrix, - rotation_matrix_from_edmonds_euler, - rotation_matrix_from_edmonds_euler_batch, - ) - - a0, b0, g0 = 0.5, 0.8, 1.2 - R0 = rotation_matrix_from_edmonds_euler(a0, b0, g0) - # Surface peaks 6° away (a spurious far maximum): concave around a far point. - d = torch.tensor([math.radians(6.0), 0.0, 0.0], dtype=torch.float64) - R_far = axis_angle_to_matrix(d) @ R0 - - def stub(ctx, alpha, beta, gamma): - R = rotation_matrix_from_edmonds_euler_batch(alpha, beta, gamma) - Rrel = torch.einsum("mij,lj->mil", R, R_far) - tr = Rrel[:, 0, 0] + Rrel[:, 1, 1] + Rrel[:, 2, 2] - ang = torch.arccos(((tr - 1.0) / 2.0).clamp(-1.0, 1.0)) - return -50.0 * ang * ang - - monkeypatch.setattr(ml_rotation, "_llg_for_orientations", stub) - peak = RotationPeak(alpha=a0, beta=b0, gamma=g0, score=0.0, sigma=0.0) - # With a 3° step and 2 iters the surface would drag the peak several degrees; - # cap at 1° must revert it to (essentially) the input orientation. - refined = quadratic_llg_refine( - [peak], ctx=None, k_refine=1, step_deg=3.0, n_grid=3, iterations=2, - max_move_deg=1.0, - ) - Rref = rotation_matrix_from_edmonds_euler( - refined[0].alpha, refined[0].beta, refined[0].gamma, - ) - moved = _so3_angle_deg(Rref, R0) - assert moved <= 1.0 + 1e-6, f"move-cap failed: moved {moved:.3f}° > 1.0°" diff --git a/tests/unit/alignment/test_sigma_a_luzzati.py b/tests/unit/alignment/test_sigma_a_luzzati.py deleted file mode 100644 index 0c19785d..00000000 --- a/tests/unit/alignment/test_sigma_a_luzzati.py +++ /dev/null @@ -1,55 +0,0 @@ -""" -Tests for `torchref.experimental.alignment.ml_rotation.compute_sigma_a_luzzati`. - -Closed-form Luzzati σA(s) = exp(−2π²·s²·ΔVRMS²) — used to pre-weight the -FRF input field so the bare correlation is already an LLG proxy -(Phaser-style; see phaser_vs_torchref_rotation.md §6). -""" - -import math - -import torch - -from torchref.experimental.alignment.ml_rotation import compute_sigma_a_luzzati - - -def test_dc_value_is_one(): - """σA(0) = 1.0 for any ΔVRMS — full agreement at the origin.""" - s = torch.tensor([0.0], dtype=torch.float64) - for vrms in [0.1, 1.0, 2.5, 10.0]: - v = compute_sigma_a_luzzati(s, vrms).item() - assert abs(v - 1.0) < 1e-12, f"σA(0; ΔVRMS={vrms}) = {v}, want 1.0" - - -def test_monotone_decrease_in_s(): - """σA strictly decreases with |s| at fixed ΔVRMS.""" - s = torch.linspace(0.0, 0.5, 50, dtype=torch.float64) - sa = compute_sigma_a_luzzati(s, delta_vrms_A=1.0) - diffs = sa[1:] - sa[:-1] - assert (diffs <= 0).all(), "σA(s) must be monotone non-increasing in |s|" - assert sa[0] > sa[-1], "σA must actually drop, not just be flat" - - -def test_known_value_at_quarter_inverse_angstrom(): - """Regression: σA(s=0.25, ΔVRMS=1.0) = exp(−π²/8) ≈ 0.291.""" - s = torch.tensor([0.25], dtype=torch.float64) - got = compute_sigma_a_luzzati(s, delta_vrms_A=1.0).item() - expected = math.exp(-math.pi ** 2 / 8.0) # = 0.2910... - assert abs(got - expected) < 1e-12, f"got {got}, want {expected}" - - -def test_vector_broadcast(): - """Vector input returns vector output, same shape + dtype + device.""" - s = torch.linspace(0.05, 0.4, 16, dtype=torch.float32) - sa = compute_sigma_a_luzzati(s, 1.5) - assert sa.shape == s.shape - assert sa.dtype == s.dtype - - -def test_increasing_delta_vrms_steepens_falloff(): - """Larger ΔVRMS ⇒ faster decay at fixed s — sanity for the parameter knob.""" - s = torch.tensor([0.2], dtype=torch.float64) - a = compute_sigma_a_luzzati(s, 0.5).item() - b = compute_sigma_a_luzzati(s, 1.0).item() - c = compute_sigma_a_luzzati(s, 2.0).item() - assert a > b > c, f"expected a > b > c; got {a}, {b}, {c}" diff --git a/torchref/experimental/alignment/__init__.py b/torchref/experimental/alignment/__init__.py index 3c9c0b1f..2d24a92b 100644 --- a/torchref/experimental/alignment/__init__.py +++ b/torchref/experimental/alignment/__init__.py @@ -1,20 +1,16 @@ """ -Alignment module for TorchRef. - -Pure-PyTorch Patterson-based molecular replacement: +Molecular replacement for TorchRef: a rotation search feeding a translation +search, and one Wilson normalisation shared between them. 1. Fast Rotation Function (``rotation_search``, over ``frf.FastRotationFunction``) — Phaser-faithful Bessel-radial × SH - expansion, stable Wigner-d, dense P1-box calc — then ML rescoring - (``ml_rotation.m_letf1_rescore``) to rank candidate orientations. + expansion, stable Wigner-d, dense P1-box calc. A shortlist generator. 2. Fast Translation Function (``translation.amplitude_translation_search`` + - ``local_translation_refine``) — run per rotation candidate. -3. Rigid Body Refinement (``rigid_body.RigidBodyRefinement``) — LBFGS on - rotation and translation (and optional B-factors) with an ML target. -4. Canonical Pipeline (``pipeline.MolecularReplacementPipeline``) — the - multi-candidate FRF → FTF → post-refine tree with early-stopping; the - implementation that ``align.align_model_to_data`` / - ``align.align_model_to_data`` delegates to. + ``local_translation_refine``) — run per rotation candidate, and where the + discrimination actually happens. +3. Pipeline (``pipeline.MolecularReplacementPipeline``) — the multi-candidate + FRF → FTF tree with early stopping, which ``align.align_model_to_data`` + delegates to. It returns a placement; refine it downstream. Example — full MR pipeline -------------------------- @@ -51,8 +47,6 @@ rotation_angular_distance_deg, rotation_matrix_from_edmonds_euler, ) -from .lattman_love import LattmanLoveInterpolator -from .ml_rotation import m_letf1_rescore, sim_mlrf_rescore from .sh import ( evaluate_ylm, sh_expand_ball, @@ -73,19 +67,15 @@ # Translation search # ============================================================================= from .translation import ( - fft_translation_search, - fft_translation_search_torch, TranslationPeak, + amplitude_translation_search, find_translation_peaks, - apply_translation_to_fcalc, - apply_translation_to_fcalc_torch, + fit_sigma_a_per_shell, + llg_translation_rescore, + local_translation_refine, + precompute_G_for_rotation, ) -# ============================================================================= -# Rigid body refinement -# ============================================================================= -from .rigid_body import RigidBodyRefinement, RigidBodyResult - # ============================================================================= # ML distributions # ============================================================================= @@ -93,9 +83,6 @@ stable_log_bessel_i0, rice_log_likelihood, woolfson_log_likelihood, - combined_log_likelihood, - acentric_pdf, - centric_pdf, ) __all__ = [ @@ -108,9 +95,6 @@ "edmonds_euler_from_rotation_matrix", "rotation_angular_distance_deg", # Rescore + interpolation - "LattmanLoveInterpolator", - "m_letf1_rescore", - "sim_mlrf_rescore", # Low-level math primitives "evaluate_ylm", "sh_expand_ball", @@ -126,23 +110,19 @@ "rotation_search", "RotationSolutions", # Translation - "fft_translation_search", - "fft_translation_search_torch", "TranslationPeak", + "amplitude_translation_search", "find_translation_peaks", - "apply_translation_to_fcalc", - "apply_translation_to_fcalc_torch", + "fit_sigma_a_per_shell", + "llg_translation_rescore", + "local_translation_refine", + "precompute_G_for_rotation", # Rigid body refinement - "RigidBodyRefinement", - "RigidBodyResult", # Transforms # Clash scoring # Distributions "stable_log_bessel_i0", "rice_log_likelihood", "woolfson_log_likelihood", - "combined_log_likelihood", - "acentric_pdf", - "centric_pdf", # Utilities ] diff --git a/torchref/experimental/alignment/align.py b/torchref/experimental/alignment/align.py index 5832f716..5c6e0c29 100644 --- a/torchref/experimental/alignment/align.py +++ b/torchref/experimental/alignment/align.py @@ -23,8 +23,6 @@ import torch -from .lattman_love import LattmanLoveInterpolator -from .e_values import SmoothSigmaE from .sh import ( apply_overall_anisotropy, assign_shells, @@ -142,8 +140,10 @@ def _external_rwork(model: "ModelFT", data: "ReflectionData") -> float: class _DirectModelEvaluator: """Returns ``F_p1(hkl)`` of a P1-spacegroup model at integer HKL. - Wraps a `ModelFT` to expose the same `.evaluate(R, hkl, cell, ...)` API - as `LattmanLoveInterpolator`, for use by the translation search. + The translation search asks its evaluator for ``F`` at a list of rotated + Miller indices. The rotation is already baked into the model's coordinates + by the time this is built, so ``R`` is ignored and every call is a direct + structure-factor evaluation rather than an interpolation. """ def __init__(self, m: "ModelFT") -> None: @@ -167,17 +167,16 @@ def evaluate(self, R, hkl, real_cell, return_amplitude=False): class FRFInputs: """Prepared reflection arrays shared by the rotation-search stages. - The resolution-masked, anisotropy-corrected reflection arrays, plus the - overall anisotropy tensor and the Lattman-Love interpolator. The rotation - search reads `U_aniso` and `device`; the ML rescore, translation search and - rigid-body polish read the rest. + The resolution-masked, anisotropy-corrected reflection arrays plus the + overall anisotropy tensor. The rotation search reads ``U_aniso`` and + ``device``; the rest is there for anything scoring against the same + observations on the same footing. ``sig_F`` carries the same anisotropy correction as ``F_obs``, which is a multiplicative factor, so ``F/sigma`` is unchanged by it. It is here because - the rotation function computes the French-Wilson posterior from the sigmas - and then threw them away, leaving the ML rescore -- a likelihood, where - measurement error is not a detail -- with no access to them at all. - ``None`` when the data carry no sigmas. + the rotation function computes its measurement weight from the sigmas, and + the earlier code threw them away immediately afterwards. ``None`` when the + data carry no sigmas. """ F_obs: torch.Tensor # (N,) anisotropy-corrected amplitudes sig_F: Optional[torch.Tensor] # (N,) their sigmas, same correction @@ -185,7 +184,6 @@ class FRFInputs: s_vec: torch.Tensor # (N, 3) reciprocal-space Cartesian s_mag: torch.Tensor # (N,) Å⁻¹ centric: torch.Tensor # (N,) bool - ll: "LattmanLoveInterpolator" U_aniso: torch.Tensor # (3, 3) Popov-Bourenkov U device: torch.device @@ -197,15 +195,13 @@ def _prepare_frf_inputs( d_min: float, d_max: float, n_shells: int, - ll_padding_factor: float = 2.0, - ll_max_res_A: float = 3.0, verbose: int = 0, ) -> FRFInputs: """Prepare the reflection arrays the rotation-search stages share. - Masks the observations to ``[d_min, d_max]``, fits and applies the overall - anisotropy correction, and builds the Lattman-Love interpolator for the - model. ``F_obs`` on the returned dataclass is anisotropy-corrected. + Masks the observations to ``[d_min, d_max]`` and fits and applies the + overall anisotropy correction, so ``F_obs`` on the returned dataclass is + anisotropy-corrected. """ device = model.xyz().device @@ -254,11 +250,6 @@ def _prepare_frf_inputs( sig_F_aniso = (None if sig_F is None else apply_overall_anisotropy(sig_F, s_vec, U_aniso)) - ll = LattmanLoveInterpolator( - model, padding_factor=ll_padding_factor, max_res_A=ll_max_res_A, - verbose=verbose, - ) - return FRFInputs( F_obs=F_obs_aniso, sig_F=sig_F_aniso, @@ -266,7 +257,6 @@ def _prepare_frf_inputs( s_vec=s_vec, s_mag=s_mag, centric=centric, - ll=ll, U_aniso=U_aniso, device=device, ) @@ -285,39 +275,21 @@ def align_model_to_data( d_max: float = 15.0, n_shells: int = 20, n_rotation_peaks: int = 500, - n_ml_refine: int = 20, # rescore only the top-20 FRF peaks (refinement use case) - ll_max_res_A: float = 3.0, - ll_padding_factor: float = 2.0, verbose: int = 0, - auto_variance_weights: bool = True, do_translation: bool = True, n_translation_peaks: int = 20, n_translation_candidates: int = 3, translation_grid_steps: int = 16, n_rotation_candidates: int = 15, - do_joint_refine: bool = True, - joint_refine_max_res_A: float = 4.0, - joint_refine_expected_rot_error: float = 0.1, - use_interp_var: bool = False, use_llg_tf: bool = False, - refine_b: bool = False, - sigma_rot_deg: float = 0.0, - sigma_trans_ang: float = 0.0, - sigma_b: float = 0.0, model_error_A: Optional[float] = None, - rescore_engine: str = "m_letf1", - rescore_e_convention: type = SmoothSigmaE, - subpeak_refine: bool = False, - subpeak_refine_k: int = -1, - subpeak_refine_step_deg: float = 1.5, - subpeak_refine_iters: int = 1, - subpeak_refine_max_move_deg: Optional[float] = 1.5, ) -> "ModelFT": - """Run full MR alignment of ``model`` against ``data``. + """Place ``model`` in ``data``'s crystal: rotation search, then translation. - Returns a new rotated+translated+refined ``ModelFT`` carrying + Returns a new rotated+translated ``ModelFT`` carrying ``last_alignment_rotation``, ``last_alignment_translation`` and - ``last_alignment_rfactor`` provenance attributes. + ``last_alignment_rfactor`` provenance attributes. It is a *placement*, not a + refined structure -- refine it downstream. `MolecularReplacementPipeline` is the implementation of record; this function returns its single best solution. @@ -337,28 +309,13 @@ def align_model_to_data( device=model.xyz().device, verbose=verbose, d_min=d_min, d_max=d_max, n_shells=n_shells, - ll_max_res_A=ll_max_res_A, ll_padding_factor=ll_padding_factor, - n_rotation_peaks=n_rotation_peaks, n_ml_refine=n_ml_refine, + n_rotation_peaks=n_rotation_peaks, model_error_A=model_error_A, - rescore_engine=rescore_engine, - rescore_e_convention=rescore_e_convention, - auto_variance_weights=auto_variance_weights, - use_interp_var=use_interp_var, - subpeak_refine=subpeak_refine, subpeak_refine_k=subpeak_refine_k, - subpeak_refine_step_deg=subpeak_refine_step_deg, - subpeak_refine_iters=subpeak_refine_iters, - subpeak_refine_max_move_deg=subpeak_refine_max_move_deg, n_rotation_candidates=n_rotation_candidates, n_translation_peaks=n_translation_peaks, n_translation_candidates=n_translation_candidates, translation_grid_steps=translation_grid_steps, use_llg_tf=use_llg_tf, - do_joint_refine=do_joint_refine, - joint_refine_max_res_A=joint_refine_max_res_A, - joint_refine_expected_rot_error=joint_refine_expected_rot_error, - refine_b=refine_b, - sigma_rot_deg=sigma_rot_deg, sigma_trans_ang=sigma_trans_ang, - sigma_b=sigma_b, ) solutions = pipeline.run(do_translation=do_translation) return solutions[0].model diff --git a/torchref/experimental/alignment/distributions.py b/torchref/experimental/alignment/distributions.py index c6e14980..d7ece281 100644 --- a/torchref/experimental/alignment/distributions.py +++ b/torchref/experimental/alignment/distributions.py @@ -212,186 +212,3 @@ def woolfson_log_likelihood( ) return log_likelihood - - -def combined_log_likelihood( - F_obs: torch.Tensor, - F_mean: torch.Tensor, - variance: torch.Tensor, - centric_flags: torch.Tensor, -) -> torch.Tensor: - """ - Combined log-likelihood dispatching to Rice/Woolfson based on centric flags. - - Parameters - ---------- - F_obs : torch.Tensor - Observed structure factor amplitudes |F_obs|. - F_mean : torch.Tensor - Expected structure factor amplitudes D * |F_calc|. - variance : torch.Tensor - Variance parameters for each reflection. - centric_flags : torch.Tensor - Boolean mask, True for centric reflections. - - Returns - ------- - torch.Tensor - Log-likelihood for each reflection using the appropriate distribution. - - Examples - -------- - :: - - F_obs = torch.tensor([10.0, 20.0, 15.0]) - F_mean = torch.tensor([9.5, 18.0, 14.0]) - variance = torch.tensor([5.0, 8.0, 6.0]) - centric = torch.tensor([False, True, False]) - ll = combined_log_likelihood(F_obs, F_mean, variance, centric) - """ - # Initialize with Rice likelihood (acentric default) - log_likelihood = rice_log_likelihood(F_obs, F_mean, variance) - - # Replace centric reflections with Woolfson likelihood - if centric_flags.any(): - centric_ll = woolfson_log_likelihood( - F_obs[centric_flags], - F_mean[centric_flags], - variance[centric_flags], - ) - log_likelihood[centric_flags] = centric_ll - - return log_likelihood - - -def acentric_pdf( - F_obs: torch.Tensor, - F_mean: torch.Tensor, - variance: torch.Tensor, -) -> torch.Tensor: - """ - Probability density for acentric reflections (Rice distribution). - - This is the PDF (not log) for cases where the actual probability is needed. - For numerical reasons, prefer rice_log_likelihood when possible. - - Parameters - ---------- - F_obs : torch.Tensor - Observed structure factor amplitudes. - F_mean : torch.Tensor - Expected structure factor amplitudes. - variance : torch.Tensor - Variance parameter. - - Returns - ------- - torch.Tensor - Probability density for each reflection. - """ - return torch.exp(rice_log_likelihood(F_obs, F_mean, variance)) - - -def centric_pdf( - F_obs: torch.Tensor, - F_mean: torch.Tensor, - variance: torch.Tensor, -) -> torch.Tensor: - """ - Probability density for centric reflections (Woolfson distribution). - - Parameters - ---------- - F_obs : torch.Tensor - Observed structure factor amplitudes. - F_mean : torch.Tensor - Expected structure factor amplitudes. - variance : torch.Tensor - Variance parameter. - - Returns - ------- - torch.Tensor - Probability density for each reflection. - """ - return torch.exp(woolfson_log_likelihood(F_obs, F_mean, variance)) - - -# ============================================================================= -# Phaser-faithful log-likelihood normalization (m_LETF1 / RiceWoolfson.cc) -# ============================================================================= -# -# These differ from `rice_log_likelihood` / `woolfson_log_likelihood` above in -# their normalization convention: Phaser's V is twice the standard Rice variance -# for acentric, equal to it for centric. Match the Phaser source exactly so the -# m_letf1_rescore values are commensurable with Phaser's m_LETF1 LL. -# -# Source: phaser/lib/RiceWoolfson.cc:25-74. - - -def phaser_log_rel_rice( - F1: torch.Tensor, - DF2: torch.Tensor, - V: torch.Tensor, -) -> torch.Tensor: - """Phaser's ``logRelRice(F1, DF2, V)`` for acentric reflections. - - Source ``phaser/lib/RiceWoolfson.cc:25-50``:: - - logRelRice(F1, DF2, V) = log I_0(2·F1·DF2/V) − log V − (F1² + DF2²)/V - - Used by ``m_LETF1`` (DataMR.cc:1425) to score each acentric reflection's - contribution to the Rice log-likelihood at a candidate orientation. - - Parameters - ---------- - F1 : torch.Tensor - Observed Wilson-normalised amplitude ``E = F_eff / sqrt(ε·Σ_N)``. - DF2 : torch.Tensor - ``sqrt(eImove)`` — square-root of the expected moving-model intensity - ``Σ_isym σ_A²·|F_calc(R^T·S_isym·h)|²``. - V : torch.Tensor - Per-reflection variance budget from ``compute_v_budget`` (DataMR.cc:949,1411). - - Returns - ------- - torch.Tensor - Per-reflection acentric log-likelihood (same shape as inputs). - """ - V_safe = V.clamp(min=1e-30) - arg = 2.0 * F1 * DF2 / V_safe - return stable_log_bessel_i0(arg) - V_safe.log() - (F1 * F1 + DF2 * DF2) / V_safe - - -def phaser_log_rel_woolfson( - F1: torch.Tensor, - DF2: torch.Tensor, - V: torch.Tensor, -) -> torch.Tensor: - """Phaser's ``logRelWoolfson(F1, DF2, V)`` for centric reflections. - - Source ``phaser/lib/RiceWoolfson.cc:52-74``:: - - logRelWoolfson(F1, DF2, V) = log cosh(F1·DF2/V) − ½·log V - − (F1² + DF2²)/(2V) - - Used by ``m_LETF1`` (DataMR.cc:1425) for centric reflections. - - Numerically stable for large ``F1·DF2/V`` via the standard - ``log cosh(x) = |x| + log1p(exp(-2|x|)) − log 2`` reformulation. - - Parameters - ---------- - F1, DF2, V : torch.Tensor - Same meaning as ``phaser_log_rel_rice``. - - Returns - ------- - torch.Tensor - Per-reflection centric log-likelihood. - """ - V_safe = V.clamp(min=1e-30) - arg = F1 * DF2 / V_safe - abs_arg = arg.abs() - log_cosh = abs_arg + torch.log1p(torch.exp(-2.0 * abs_arg)) - math.log(2.0) - return log_cosh - 0.5 * V_safe.log() - (F1 * F1 + DF2 * DF2) / (2.0 * V_safe) diff --git a/torchref/experimental/alignment/e_values.py b/torchref/experimental/alignment/e_values.py index 83efa8ba..ff9dae45 100644 --- a/torchref/experimental/alignment/e_values.py +++ b/torchref/experimental/alignment/e_values.py @@ -4,8 +4,8 @@ change**: correlating ``E_obs`` against ``E_calc`` *is* correlating ``F`` against ``F`` with weight ``1/Sigma(s)``. The alignment package currently answers that question nine different times -- twice in ``frf.preprocessing``, once inside -``french_wilson_preprocess``, three times in ``ml_rotation`` and three more in -``translation`` -- and the answers disagree. The rotation function's observed +``french_wilson_preprocess`` and three more in ``translation`` -- and the +answers disagree. The rotation function's observed side is a French-Wilson posterior weighted by ``DFAC**2``; the rescore's is plain per-shell Wilson with epsilon divided out. So the rescore ranks candidates against a differently-normalised observation set than the one that produced them. diff --git a/torchref/experimental/alignment/lattman_love.py b/torchref/experimental/alignment/lattman_love.py deleted file mode 100644 index 254d1827..00000000 --- a/torchref/experimental/alignment/lattman_love.py +++ /dev/null @@ -1,302 +0,0 @@ -""" -Lattman-Love (1970) structure-factor interpolation for the alignment module. - -Compute F_calc once for the search model in a large cubic P1 box (densely sampled -in reciprocal space), then interpolate the dense grid at rotated reciprocal-space -positions of the *real* crystal cell to obtain F_calc for any candidate rotation. - -This is the standard Phaser MR setup (Phaser paper §2.2.2): F_calc is generated -"by structure-factor interpolation (Lattman & Love, 1970) from a model in a large -P1 unit cell". It removes the sphere-sampling bias that arises if one instead -rotates atom coordinates and recomputes F_calc on the (non-uniform) real-cell -HKL grid — which is what the bare ball-search hits on real, non-cubic data. - -Convention (matches `torchref/base/reciprocal/interpolation.py::interpolate_for_rotation`): - For a model whose atom coordinates have been rotated by R (column-vector - convention: xyz_new = R · xyz_old), the structure factor at real-cell HKL h - equals the un-rotated model's structure factor at the rotated reciprocal- - space point R^T · s_real, where s_real = h · rec_basis(real_cell). - -Only amplitudes are needed for the rotation function and Sim MLRF rescoring. -Translation search (downstream) needs the phase too — this class returns complex -F so callers can use either. -""" - -from __future__ import annotations - -from typing import Optional - -import torch - -from torchref.base.fourier.fft import ifft -from torchref.base.reciprocal.interpolation import ( - interpolate_complex_from_grid, - interpolate_structure_factor_from_grid, -) -from torchref.model.sf_fft import SfFFT -from torchref.symmetry import SpaceGroup -from torchref.symmetry.cell import Cell -from torchref.utils.device_mixin import DeviceMixin - - -class LattmanLoveInterpolator(DeviceMixin): - """ - Compute F_calc on a dense P1 reciprocal grid once; interpolate at arbitrary - rotated reciprocal positions per query. - - Parameters - ---------- - model : ModelFT - Search model. Its current atom coordinates are used and the FT is built - for those positions (un-rotated; rotations are applied later in - `evaluate(R, ...)`). The model's own cell/spacegroup are NOT used. - padding_factor : float, default 2.0 - Cubic P1 box side = padding_factor * molecule_bounding_box_diameter. - Phaser uses ~2.0. Larger → finer reciprocal grid spacing, more memory. - min_cell_size_A : float, optional - Lower bound on the cubic side (Å). Useful for very small molecules where - 2·diameter would give a tiny FFT grid. - max_res_A : float, optional - Resolution limit (Å). The dense grid will resolve features down to this. - Default: 2.0 Å (suitable for proteins up to that resolution). - device : torch.device, optional - Target device. Defaults to the model's device. - """ - - def __init__( - self, - model, - padding_factor: float = 2.0, - min_cell_size_A: Optional[float] = None, - max_res_A: float = 2.0, - device: Optional[torch.device] = None, - verbose: int = 0, - ): - if device is None: - device = model.xyz().device - - # Bounding-box diameter of the un-rotated atomic coordinates. - xyz = model.xyz().detach().to(device) - bbox_max = xyz.max(dim=0).values - bbox_min = xyz.min(dim=0).values - diameter_A = (bbox_max - bbox_min).norm().item() - cubic_side = padding_factor * diameter_A - if min_cell_size_A is not None: - cubic_side = max(cubic_side, float(min_cell_size_A)) - - # Cubic P1 cell. Keep dtype float32 to match Cell defaults / SfFFT grid math. - self.cubic_cell = Cell( - [cubic_side, cubic_side, cubic_side, 90.0, 90.0, 90.0], - dtype=torch.float32, device=device, - ) - - # Shift atoms so the molecule centroid sits at the centre of the cubic box. - centroid = xyz.mean(dim=0) - target = torch.tensor( - [cubic_side / 2, cubic_side / 2, cubic_side / 2], - dtype=xyz.dtype, device=device, - ) - self.shift_vec = target - centroid # (3,) — applied to atom coords - - # Extract atomic parameters (use the model's own helper). - xyz_iso, adp_iso, occ_iso, A_iso, B_iso = model.get_iso() - xyz_iso = (xyz_iso.detach().to(device) + self.shift_vec).to(xyz.dtype) - adp_iso = adp_iso.detach().to(device) - occ_iso = occ_iso.detach().to(device) - A_iso = A_iso.detach().to(device) - B_iso = B_iso.detach().to(device) - - # Handle the (uncommon) anisotropic atoms by leaving them aside for the - # search model — Phaser-style MR also approximates with isotropic ADP. - # Caller is free to call evaluate after adding aniso atoms in a subclass. - - # Build the dense F_calc on a P1 cubic grid. Wrapped in no_grad because - # the alignment pipeline never differentiates through this grid — and - # without no_grad the resulting `self.reciprocal_grid` carries a grad_fn - # whose autograd graph pins the SfFFT internals (real_space_grid, - # voxel_xyz, per-atom kernel ≈ 5 GB on 4BX9) across trials. - with torch.no_grad(): - sf = SfFFT( - cell=self.cubic_cell, - spacegroup=SpaceGroup("P 1"), - max_res=max_res_A, - dtype_float=torch.float32, - device=device, - verbose=verbose, - ) - sf.setup_grid() - density_map = sf.build_density_map( - xyz_iso=xyz_iso, - adp_iso=adp_iso, - occ_iso=occ_iso, - A_iso=A_iso, - B_iso=B_iso, - apply_symmetry=False, # already P1 - ) - # IFFT to reciprocal space; gives a complex (Nx, Ny, Nz) tensor - # with crystallographic normalization. Layout: DC at (0, 0, 0); - # negative HKL wraps to high indices. Matches - # `interpolate_structure_factor_from_grid`'s expectation. - self.reciprocal_grid = ifft( - density_map, self.cubic_cell.volume.item(), - ) - self.cubic_cell_volume = float(self.cubic_cell.volume.item()) - self.device = device - self.cubic_side = cubic_side - self.gridsize = tuple(int(x) for x in self.reciprocal_grid.shape) - self.max_res_A = max_res_A - - if verbose: - print(f"LattmanLove: molecule diameter {diameter_A:.1f} Å, " - f"cubic side {cubic_side:.1f} Å, grid {self.gridsize}, " - f"max_res {max_res_A:.2f} Å") - - @staticmethod - def _real_hkl_to_cubic_hkl( - hkl_real: torch.Tensor, real_cell: Cell, cubic_cell: Cell, - ) -> torch.Tensor: - """Convert HKL of `real_cell` to (float) HKL of `cubic_cell`.""" - # s = h @ rec_basis (Å^-1, Cartesian) - rec_real = real_cell.reciprocal_basis_matrix.to(hkl_real.device) - s = hkl_real.to(rec_real.dtype) @ rec_real - rec_cubic = cubic_cell.reciprocal_basis_matrix.to(hkl_real.device) - # cubic_hkl = s @ rec_cubic^{-1} - return s @ torch.linalg.inv(rec_cubic.to(rec_real.dtype)) - - def evaluate( - self, - R: torch.Tensor, - hkl_real: torch.Tensor, - real_cell: Cell, - return_amplitude: bool = True, - ) -> torch.Tensor: - """ - Interpolate F_calc at the real-cell HKL set after rotating the model - by R (column-vector convention). - - Parameters - ---------- - R : torch.Tensor - Rotation matrix, shape (3, 3) or (B, 3, 3). - hkl_real : torch.Tensor - Miller indices in the real crystal cell, shape (N, 3). - real_cell : Cell - The real crystal cell (provides `reciprocal_basis_matrix`). - return_amplitude : bool, default True - If True, return |F_calc| (real, no phase ambiguity from trilinear - interpolation). If False, return complex F — only safe if the dense - grid is well-oversampled (small `max_res_A`). - - Returns - ------- - torch.Tensor - Interpolated structure factors, shape (N,) or (B, N). - """ - batched = R.dim() == 3 - if not batched: - R = R.unsqueeze(0) - R = R.to(self.device).to(torch.float32) - - # Real HKL → Cartesian s (Å^-1, real cell) - rec_real = real_cell.reciprocal_basis_matrix.to(self.device).to(torch.float32) - s_real = hkl_real.to(self.device).to(torch.float32) @ rec_real # (N, 3) - - # Rotated reciprocal point: F_rotated_model(s) = F_orig(R^T s) => - # use R^T · s to look up the un-rotated grid. - # For batched R, einsum: - s_rot = torch.einsum("bij,nj->bni", R.transpose(-1, -2), s_real) # (B, N, 3) - - # Cartesian s → cubic-cell float HKL - rec_cubic = self.cubic_cell.reciprocal_basis_matrix.to(self.device).to(torch.float32) - inv_rec_cubic = torch.linalg.inv(rec_cubic) - hkl_cubic = s_rot @ inv_rec_cubic # (B, N, 3) - - B, N, _ = hkl_cubic.shape - flat = hkl_cubic.reshape(B * N, 3) - if return_amplitude: - interp = interpolate_structure_factor_from_grid( - self.reciprocal_grid, flat, interpolate_amplitude=True, - ) - else: - interp = interpolate_complex_from_grid(self.reciprocal_grid, flat) - out = interp.reshape(B, N) - return out if batched else out.squeeze(0) - - -def estimate_interp_var( - interpolator: "LattmanLoveInterpolator", - hkl_real: torch.Tensor, - real_cell: Cell, - shell_idx: torch.Tensor, - n_shells: int, - n_jitter: int = 4, - jitter_frac: float = 0.5, - seed: int = 0, -) -> torch.Tensor: - """ - Estimate per-reflection trilinear-interpolation variance in E-value units. - - Phaser's totvar_search analogue. Inflates the Rice/Woolfson variance budget - so that interpolation noise in the search model doesn't make a slightly- - noisy true peak look worse than a noise-free wrong peak. - - Method: evaluate the interpolator at the original HKLs and at `n_jitter` - sub-grid-cell perturbations of the HKLs, take the per-shell empirical - variance of |F| across the perturbations, and normalise by the per-shell - mean |F|² so the returned quantity adds correctly to ``(1 - D²)`` in the - Rice variance. - - `jitter_frac` is the fraction of a cubic-cell grid spacing to jitter by; - 0.5 sweeps half a Nyquist cell and gives a robust upper-bound estimate - of trilinear bias. n_jitter=4 keeps the cost negligible. - - Returns - ------- - interp_var : torch.Tensor, shape (N,) - Per-reflection interpolation variance in dimensionless E² units. - """ - device = interpolator.device - dtype = torch.float32 - R_eye = torch.eye(3, dtype=dtype, device=device) - - F_ref = interpolator.evaluate( - R_eye, hkl_real, real_cell, return_amplitude=True, - ).to(dtype) # (N,) - - # Map a Cartesian Å^-1 shift back to fractional HKL_real space. delta_s is - # the magnitude of the jitter in Cartesian reciprocal Å^-1. - delta_s = jitter_frac / float(interpolator.cubic_side) - rec_real_inv = torch.linalg.inv( - real_cell.reciprocal_basis_matrix.to(device).to(dtype), - ) - - g = torch.Generator(device="cpu").manual_seed(int(seed)) - diffs_sq = torch.zeros_like(F_ref) - for _ in range(n_jitter): - direction = torch.randn(3, generator=g, dtype=torch.float64) - direction = (direction / direction.norm()).to(device).to(dtype) - delta_h_real = (delta_s * direction) @ rec_real_inv # (3,) - hkl_j = hkl_real.to(dtype) + delta_h_real - F_j = interpolator.evaluate( - R_eye, hkl_j, real_cell, return_amplitude=True, - ).to(dtype) - diffs_sq = diffs_sq + (F_j - F_ref) ** 2 - diffs_sq = diffs_sq / max(n_jitter, 1) # (N,) - - # Per-shell aggregation. Cast to f64 for stable sums on large N. - shell_idx_l = shell_idx.to(device).long() - diffs_d = diffs_sq.to(torch.float64) - F_ref2_d = (F_ref.to(torch.float64)) ** 2 - - var_per_shell = torch.zeros(n_shells, dtype=torch.float64, device=device) - F2_per_shell = torch.zeros(n_shells, dtype=torch.float64, device=device) - cnt = torch.zeros(n_shells, dtype=torch.float64, device=device) - var_per_shell.scatter_add_(0, shell_idx_l, diffs_d) - F2_per_shell.scatter_add_(0, shell_idx_l, F_ref2_d) - cnt.scatter_add_(0, shell_idx_l, torch.ones_like(diffs_d)) - - mean_var = var_per_shell / cnt.clamp(min=1.0) # (n_shells,) - mean_F2 = (F2_per_shell / cnt.clamp(min=1.0)).clamp(min=1e-30) # (n_shells,) - interp_var_E_per_shell = (mean_var / mean_F2).clamp(min=0.0, max=1.0) - - return interp_var_E_per_shell.to(dtype).index_select(0, shell_idx_l) # (N,) diff --git a/torchref/experimental/alignment/ml_rotation.py b/torchref/experimental/alignment/ml_rotation.py deleted file mode 100644 index a33fdc6d..00000000 --- a/torchref/experimental/alignment/ml_rotation.py +++ /dev/null @@ -1,1208 +0,0 @@ -""" -Maximum-Likelihood Rotation Function (Sim MLRF) rescoring of peaks from the -fast ball-search. - -Phaser paper §2.1.2: the fast rotation function is a "shortlist generator"; -discrimination of the correct orientation comes from rescoring the top peaks -with a slow, full ML target. This file implements that rescoring. - -Per-shell σA (= D) is fitted on-the-fly for each candidate rotation. We work in -E-value (normalized structure-factor amplitude) space, which gives the standard -Rice / Woolfson likelihood forms - - P(E_obs | σA · E_calc, 1 − σA²) (acentric, Rice) - P(E_obs | σA · E_calc, 1 − σA²) (centric, Woolfson) - -The "log-likelihood gain" relative to a Wilson reference (σA = 0) is the -discriminating score Phaser reports as LLG. -""" - -from __future__ import annotations - -import math -from dataclasses import dataclass -from typing import Callable, List, Optional - -import torch - -from torchref.scaling.weighting import (inverse_variance_weight, - normalise_weight, - snr_from_amplitude) -from .e_values import (CalcShellE, SmoothSigmaE, WilsonShellE, - WilsonShellEpsE, convention_for_calc, - convention_uses_sigma_f) -from .frf.rotation_utils import ( - axis_angle_to_matrix, - edmonds_euler_from_rotation_matrix, - rotation_matrix_from_edmonds_euler, - rotation_matrix_from_edmonds_euler_batch, -) -from .frf.types import RotationPeak -from .distributions import ( - phaser_log_rel_rice, - phaser_log_rel_woolfson, - rice_log_likelihood, - woolfson_log_likelihood, -) -from .lattman_love import LattmanLoveInterpolator - - -# ============================================================================= -# Helpers -# ============================================================================= - - -def _equal_count_shell_idx(s_mag: torch.Tensor, n_shells: int) -> torch.Tensor: - """ - Partition `s_mag` into `n_shells` shells with (approximately) equal counts. - - Returns - ------- - shell_idx : torch.Tensor (int64), shape (N,) - Shell index in [0, n_shells). - """ - n = s_mag.numel() - order = torch.argsort(s_mag) - chunk = max(n // n_shells, 1) - positions = torch.arange(n, device=s_mag.device, dtype=torch.int64) - sorted_labels = (positions // chunk).clamp(max=n_shells - 1) - shell_idx = torch.empty(n, dtype=torch.int64, device=s_mag.device) - shell_idx[order] = sorted_labels - return shell_idx - - -def _shell_ll( - E_obs: torch.Tensor, - E_calc: torch.Tensor, - centric: torch.Tensor, - D: float, - interp_var: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """ - Per-reflection log-likelihood at a given σA = D for one shell, in E-value space. - - Acentric: Rice with F_mean = D · E_calc, variance = (1 − D²) + interp_var. - Centric: Woolfson with F_mean = D · E_calc, variance = (1 − D²) + interp_var. - - `interp_var` (Phaser totvar_search analogue) inflates the variance to absorb - interpolation / model error and prevents the Rice tail from over-penalising - slightly-noisy true peaks. - """ - base = max(1.0 - D * D, 1e-4) - if interp_var is None: - var = torch.full_like(E_obs, base) - else: - var = (interp_var + base).clamp(min=1e-4) - F_mean = D * E_calc - ll = torch.where( - centric, - woolfson_log_likelihood(E_obs, F_mean, var), - rice_log_likelihood(E_obs, F_mean, var), - ) - return ll - - -def _optimize_D_in_shell( - E_obs: torch.Tensor, - E_calc: torch.Tensor, - centric: torch.Tensor, - n_grid: int = 21, - n_refine: int = 12, -) -> float: - """ - Find the σA = D ∈ [0, 0.99] that maximizes the sum log-likelihood for this - shell. Two-stage: coarse grid search, then golden-section refinement. - - Returns the optimal D. - """ - # Coarse grid - D_grid = torch.linspace(0.0, 0.99, n_grid, device=E_obs.device) - best_D = 0.0 - best_ll = -float("inf") - for D in D_grid.tolist(): - ll = _shell_ll(E_obs, E_calc, centric, D).sum().item() - if ll > best_ll: - best_ll = ll - best_D = D - # Golden-section refinement around best_D - span = 1.0 / (n_grid - 1) - lo = max(0.0, best_D - span) - hi = min(0.99, best_D + span) - phi = (math.sqrt(5.0) - 1) / 2.0 - x1 = hi - phi * (hi - lo) - x2 = lo + phi * (hi - lo) - f1 = _shell_ll(E_obs, E_calc, centric, x1).sum().item() - f2 = _shell_ll(E_obs, E_calc, centric, x2).sum().item() - for _ in range(n_refine): - if f1 > f2: - hi = x2 - x2 = x1 - f2 = f1 - x1 = hi - phi * (hi - lo) - f1 = _shell_ll(E_obs, E_calc, centric, x1).sum().item() - else: - lo = x1 - x1 = x2 - f1 = f2 - x2 = lo + phi * (hi - lo) - f2 = _shell_ll(E_obs, E_calc, centric, x2).sum().item() - return 0.5 * (lo + hi) - - -def compute_sigma_a_luzzati( - s_mag: torch.Tensor, - delta_vrms_A: float = 1.0, -) -> torch.Tensor: - """ - Phaser-style Luzzati σA(s) = exp(−2π²·s²·ΔVRMS²). - - Closed-form, rotation-independent estimate of the per-reflection (or - per-shell) σA, derived from the search model's RMS coordinate - deviation ΔVRMS (in Å). At s=0 returns 1.0 (perfect agreement); - falls off monotonically with resolution. Matches the Phaser FastRot - Eterm/Vterm weighting (LERF1 §2.1.2): `Eterm = exp(−2π²s²ΔVRMS)` and - `Vterm = Eterm²`. - - Parameters - ---------- - s_mag : torch.Tensor - Reciprocal magnitudes in Å⁻¹ (any shape). - delta_vrms_A : float - Estimated RMS coordinate error of the search model, Å. Default - 1.0 Å is a reasonable starting point for MR search models; - tune via `frf_delta_vrms_A` kwarg in `align_model_to_data`. - - Returns - ------- - sigma_a : torch.Tensor, same shape and dtype as `s_mag`. - """ - return torch.exp( - -2.0 * (math.pi ** 2) * (s_mag ** 2) * (float(delta_vrms_A) ** 2) - ) - - -def fit_sigma_a_per_shell( - E_obs: torch.Tensor, - E_calc: torch.Tensor, - centric: torch.Tensor, - shell_idx: torch.Tensor, - n_shells: int, - n_grid: int = 81, - interp_var: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """ - Vectorised per-shell σA = D fit. Single source of truth for D across the - alignment stages (rotation rescore + likelihood TF). - - For each shell, scans D ∈ [0, 0.99] on a fine grid and returns the - grid maximum. With n_grid=81 the resolution is ~0.012, comparable to the - golden-section result in `_optimize_D_in_shell` for downstream LLG purposes. - - Returns - ------- - sigma_a : torch.Tensor, shape (n_shells,) - """ - device = E_obs.device - dtype = E_obs.dtype - N = E_obs.numel() - D_grid = torch.linspace(0.0, 0.99, n_grid, device=device, dtype=dtype) # (G,) - F_mean = D_grid.view(-1, 1) * E_calc.view(1, -1) # (G, N) - var_d = (1.0 - D_grid * D_grid).clamp(min=1e-4) # (G,) - if interp_var is None: - var_full = var_d.view(-1, 1).expand(n_grid, N) - else: - var_full = (var_d.view(-1, 1) + interp_var.view(1, -1)).clamp(min=1e-4) - E_obs_full = E_obs.view(1, -1).expand(n_grid, N) - - ll_acent = rice_log_likelihood(E_obs_full, F_mean, var_full) - ll_cent = woolfson_log_likelihood(E_obs_full, F_mean, var_full) - cent_full = centric.view(1, -1) - ll = torch.where(cent_full, ll_cent, ll_acent) # (G, N) - - # Sum per shell, take argmax over the D-grid. - shell_idx_gn = shell_idx.view(1, -1).expand(n_grid, N) - ll_per_shell = torch.zeros((n_grid, n_shells), dtype=dtype, device=device) - ll_per_shell.scatter_add_(1, shell_idx_gn, ll) # (G, n_shells) - best_idx = ll_per_shell.argmax(dim=0) # (n_shells,) - return D_grid[best_idx] - - -# ============================================================================= -# Public API -# ============================================================================= - - -def llg_for_rotation( - F_obs: torch.Tensor, - s_mag: torch.Tensor, - shell_idx: torch.Tensor, - n_shells: int, - E_obs: torch.Tensor, - centric: torch.Tensor, - F_calc: torch.Tensor, - shell_weights: Optional[torch.Tensor] = None, -) -> float: - """ - Total log-likelihood gain (Sim − Wilson) for a single candidate rotation. - Thin wrapper around `llg_for_rotation_batch` for a single (1,N) input. - """ - F_calc_batch = F_calc.unsqueeze(0) if F_calc.dim() == 1 else F_calc - return llg_for_rotation_batch( - F_obs=F_obs, shell_idx=shell_idx, n_shells=n_shells, - E_obs=E_obs, centric=centric, F_calc=F_calc_batch, - shell_weights=shell_weights, - )[0].item() - - -def llg_for_rotation_batch( - F_obs: torch.Tensor, - shell_idx: torch.Tensor, - n_shells: int, - E_obs: torch.Tensor, - centric: torch.Tensor, - F_calc: torch.Tensor, - n_D_grid: int = 41, - shell_weights: Optional[torch.Tensor] = None, - interp_var: Optional[torch.Tensor] = None, - sigma_a: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """ - Vectorized log-likelihood gain across a batch of candidate rotations. - - Per-shell σA fit is performed on a coarse grid of `n_D_grid` D values; - the grid maximum is taken (no golden refinement — for shortlisting only). - - Parameters - ---------- - F_obs : torch.Tensor, shape (N,) - shell_idx : torch.Tensor (int64), shape (N,) - n_shells : int - E_obs : torch.Tensor, shape (N,) - F_obs normalized to unit variance per shell. - centric : torch.Tensor (bool), shape (N,) - F_calc : torch.Tensor, shape (B, N) - Per-rotation |F_calc| at the same HKL set. - n_D_grid : int, default 41 - Number of σA grid points in [0, 0.99]. - shell_weights : torch.Tensor, shape (n_shells,), optional - Per-shell weight applied to the per-shell LL gain before accumulation. - Used to implement Phaser-style empirical variance correction - (`w_p = 1/√Var(E_obs²-1)_p`). The weight multiplies *both* the Sim and - Wilson LL contributions uniformly per shell, so the LL gain - interpretation is preserved. - interp_var : torch.Tensor, shape (N,), optional - Per-reflection interpolation variance (Phaser totvar_search analogue). - Added to the model variance term. None ⇒ original Rice/Woolfson. - sigma_a : torch.Tensor, shape (n_shells,), optional - Externally-fitted per-shell σA to reuse across candidates. When given, - the per-(D, B) grid maximisation is bypassed and the LLG is computed - at this fixed sigma_a (one D per shell, broadcast per reflection). - This is the "shared D" path used when an external single-source σA - is available (see ``fit_sigma_a_per_shell``). - - Returns - ------- - llg : torch.Tensor, shape (B,) - Total log-likelihood gain (Sim − Wilson) per rotation candidate. - """ - B, N = F_calc.shape - device = F_calc.device - dtype = F_calc.dtype - - # --- Per-shell E normalisation of F_calc, fully vectorised --- - # Build (B, n_shells) shell-mean of F_calc² via scatter_add, then gather - # back per reflection. Replaces the n_shells-step Python loop that - # masked + meaned one shell at a time. - shell_idx_b = shell_idx.view(1, N).expand(B, N) - F_calc2 = F_calc * F_calc - sum_per_shell_b = torch.zeros((B, n_shells), dtype=dtype, device=device) - sum_per_shell_b.scatter_add_(1, shell_idx_b, F_calc2) - count_per_shell = torch.bincount(shell_idx, minlength=n_shells).to(dtype) - mean_per_shell_b = ( - sum_per_shell_b / count_per_shell.clamp(min=1.0).unsqueeze(0) - ).clamp(min=1e-30) - norm_per_refl_b = mean_per_shell_b.sqrt().gather(1, shell_idx_b) # (B, N) - E_calc = F_calc / norm_per_refl_b - - if sigma_a is not None: - # --- Shared per-shell σA path: skip the D-grid, evaluate LL once. --- - D_per_refl = sigma_a.to(dtype).to(device).index_select(0, shell_idx) # (N,) - var_d = (1.0 - D_per_refl * D_per_refl).clamp(min=1e-4) # (N,) - if interp_var is not None: - var_per_refl = (var_d + interp_var.to(dtype).to(device)).clamp(min=1e-4) - else: - var_per_refl = var_d - F_mean = D_per_refl.view(1, N) * E_calc # (B, N) - var_full = var_per_refl.view(1, N).expand(B, N) - E_obs_full = E_obs.view(1, N).expand(B, N) - ll_acent = rice_log_likelihood(E_obs_full, F_mean, var_full) - ll_cent = woolfson_log_likelihood(E_obs_full, F_mean, var_full) - cent_full = centric.view(1, N) - ll = torch.where(cent_full, ll_cent, ll_acent) # (B, N) - ll_per_shell = torch.zeros((B, n_shells), dtype=dtype, device=device) - ll_per_shell.scatter_add_(1, shell_idx_b, ll) - ll_sim_per_shell = ll_per_shell # (B, n_shells) - else: - # --- Joint (D, B, N) likelihood evaluation, original behaviour. --- - # Memory: D · B · N · 8 B. For default args (D=41, B≤100, N≈3 k) this is - # ~100 MB, comparable to what the per-shell loop already built per - # iteration. For dense-R we typically pass n_D_grid=11, so cost is small. - D_grid = torch.linspace(0.0, 0.99, n_D_grid, device=device, dtype=dtype) - F_mean = D_grid.view(-1, 1, 1) * E_calc.unsqueeze(0) # (D, B, N) - var_d = (1.0 - D_grid * D_grid).clamp(min=1e-4) - if interp_var is None: - var_full = var_d.view(-1, 1, 1).expand(n_D_grid, B, N) - else: - iv = interp_var.to(dtype).to(device).view(1, 1, N) - var_full = (var_d.view(-1, 1, 1) + iv).clamp(min=1e-4).expand(n_D_grid, B, N) - E_obs_full = E_obs.view(1, 1, -1).expand(n_D_grid, B, N) - - ll_acent = rice_log_likelihood(E_obs_full, F_mean, var_full) - ll_cent = woolfson_log_likelihood(E_obs_full, F_mean, var_full) - cent_full = centric.view(1, 1, -1) - ll = torch.where(cent_full, ll_cent, ll_acent) # (D, B, N) - - # --- Sum per shell across N, max over D, sum weighted across shells --- - shell_idx_dbn = shell_idx.view(1, 1, -1).expand(n_D_grid, B, N) - ll_per_shell = torch.zeros((n_D_grid, B, n_shells), dtype=dtype, device=device) - ll_per_shell.scatter_add_(2, shell_idx_dbn, ll) - ll_sim_per_shell, _ = ll_per_shell.max(dim=0) # (B, n_shells) - - # Wilson reference at D = 0 (data-only): F_mean = 0, var = 1. - var0 = torch.ones_like(E_obs) - F_mean0 = torch.zeros_like(E_obs) - ll_wil_acent = rice_log_likelihood(E_obs, F_mean0, var0) - ll_wil_cent = woolfson_log_likelihood(E_obs, F_mean0, var0) - ll_wil_per_refl = torch.where(centric, ll_wil_cent, ll_wil_acent) - ll_wil_per_shell = torch.zeros(n_shells, dtype=dtype, device=device) - ll_wil_per_shell.scatter_add_(0, shell_idx, ll_wil_per_refl) - - gain_per_shell = ll_sim_per_shell - ll_wil_per_shell.unsqueeze(0) # (B, n_shells) - if shell_weights is not None: - gain_per_shell = gain_per_shell * shell_weights.to(dtype).view(1, -1) - total_gain = gain_per_shell.sum(dim=-1) # (B,) - return total_gain - - -def sim_mlrf_rescore( - peaks: List[RotationPeak], - F_obs: torch.Tensor, - hkl_real: torch.Tensor, - s_mag: torch.Tensor, - centric: torch.Tensor, - interpolator: LattmanLoveInterpolator, - real_cell, - n_shells: int = 20, - n_refine: Optional[int] = None, - batch_size: int = 100, - verbose: int = 0, - shell_weights: Optional[torch.Tensor] = None, - auto_variance_weights: bool = True, - n_D_grid: int = 41, - interp_var: Optional[torch.Tensor] = None, - sigma_a: Optional[torch.Tensor] = None, -) -> List[RotationPeak]: - """ - Rescore a list of FRF peaks by the per-shell-fitted - Sim Maximum-Likelihood Rotation Function (LLG). Returns a new list sorted by - descending LLG with `score = LLG` and `sigma = Z-score(LLG)`. - - Batches candidates of size `batch_size` for fast vectorized evaluation. - - Parameters - ---------- - peaks : list of RotationPeak - F_obs : torch.Tensor, shape (N,) - hkl_real : torch.Tensor (int), shape (N, 3) - s_mag : torch.Tensor, shape (N,) - centric : torch.Tensor (bool), shape (N,) - interpolator : LattmanLoveInterpolator - real_cell : Cell - n_shells : int, default 20 - n_refine : int, optional - Number of top peaks to rescore (default: all). - batch_size : int, default 100 - Number of candidates evaluated together in one LL.evaluate + LLG call. - """ - if not peaks: - return [] - - if n_refine is None: - n_refine = len(peaks) - head = peaks[: n_refine] - tail = peaks[n_refine:] - - shell_idx = _equal_count_shell_idx(s_mag, n_shells) - # The caller's own shell assignment is passed in rather than letting the - # convention derive one: `_equal_count_shell_idx` is rank-based and the - # shared `assign_shells` is value-based, so they disagree on reflections - # sitting on a boundary. Keeping this one makes the migration exact. - E_obs = WilsonShellE( - F_obs, s_mag, shell_idx=shell_idx, n_shells=n_shells, - ).E - - if shell_weights is None and auto_variance_weights: - from .sh import compute_patterson_shell_variance - patt_obs = (E_obs.to(torch.float64) ** 2) - 1.0 - var_p = compute_patterson_shell_variance( - patt_obs, shell_idx, P=n_shells, - ) - w = 1.0 / var_p.sqrt() - w = w * (n_shells / w.sum().clamp(min=1e-30)) - shell_weights = w.to(F_obs.dtype) - - # Build all rotation matrices up front. The peak's Euler triple represents - # "the rotation applied to the model coords" (synthetic-test convention of - # FRF synthetic-test convention). For ML scoring, we need "the rotation to apply to - # the current model to align it to obs" — which is R^T. We transpose here. - # Vectorised over peaks: previously this list comprehension built M·9 - # small (3,3) tensors per dense-R pass. - alpha_t = torch.tensor([p.alpha for p in head], dtype=torch.float64) - beta_t = torch.tensor([p.beta for p in head], dtype=torch.float64) - gamma_t = torch.tensor([p.gamma for p in head], dtype=torch.float64) - R_all = rotation_matrix_from_edmonds_euler_batch( - alpha_t, beta_t, gamma_t, - ).transpose(-1, -2).to(torch.float32) # (M, 3, 3) - - llg_chunks: List[torch.Tensor] = [] - M = R_all.shape[0] - for start in range(0, M, batch_size): - stop = min(start + batch_size, M) - R_batch = R_all[start:stop] # (B, 3, 3) - # Batched LL interpolation: returns (B, N) - F_calc = interpolator.evaluate( - R_batch, hkl_real, real_cell, return_amplitude=True, - ) - F_calc = F_calc.to(F_obs.dtype) - llg_batch = llg_for_rotation_batch( - F_obs=F_obs, shell_idx=shell_idx, n_shells=n_shells, - E_obs=E_obs, centric=centric, F_calc=F_calc, - shell_weights=shell_weights, n_D_grid=n_D_grid, - interp_var=interp_var, sigma_a=sigma_a, - ) - llg_chunks.append(llg_batch) - if verbose > 1: - print(f" ML rescore batch {start}-{stop}/{M}", flush=True) - - # Concatenate on-device, compute z-score on-device, then ONE bulk - # transfer at the end. The previous code did `.cpu().tolist()` per - # batch — fine on CPU but a per-batch GPU↔CPU stall on cuda. - llgs_t = torch.cat(llg_chunks) - mean_t = llgs_t.mean() - std_t = llgs_t.std().clamp(min=1e-30) - sigmas_t = (llgs_t - mean_t) / std_t - - llgs_list = llgs_t.tolist() - sigmas_list = sigmas_t.tolist() - rescored = [ - RotationPeak( - alpha=p.alpha, beta=p.beta, gamma=p.gamma, - score=llg, sigma=sigma, - ) - for p, llg, sigma in zip(head, llgs_list, sigmas_list) - ] - rescored.sort(key=lambda r: r.score, reverse=True) - return rescored + tail - - -# ============================================================================= -# Phaser-faithful m_LETF1 rescore: unique-orbit calc sum + V(h) budget + -# Rice/Woolfson logRel formulas (DataMR.cc:1326-1429). -# -# The per-orientation LLG evaluator is factored out of `m_letf1_rescore` into a -# reusable `_LLGContext` + `_llg_for_orientations`, so the sub-peak refiner -# (`quadratic_llg_refine`) optimises the *same* likelihood the rescore ranks on. -# ============================================================================= - - -@dataclass -class _LLGContext: - """Rotation-independent context for the m_LETF1 per-orientation LLG. - - Built once by :func:`_build_llg_context`; consumed by - :func:`_llg_for_orientations` (rescore) and :func:`quadratic_llg_refine` - (sub-peak optimiser). Everything here depends only on the data + σ_A model, - not on the candidate orientation. - """ - - interpolator: LattmanLoveInterpolator - real_cell: object - unrolled_hkl: torch.Tensor # (M, 3) float64 — distinct orbit mates - asu_idx: torch.Tensor # (M,) long — ASU reflection each mate maps to - N: int # number of ASU reflections - E_obs_b: torch.Tensor # (1, N) - V_b: torch.Tensor # (1, N) - eImove_prefac: torch.Tensor # (1, N) = ε·σ_A²/n_ops - sqrt_mean_per_m: torch.Tensor # (M,) per-mate E-normaliser - centric_b: torch.Tensor # (1, N) bool - dw_per_m: Optional[torch.Tensor] # (M,) or None — Wilson-B Debye-Waller - dtype: torch.dtype - batch_size: int - target: str = "rice" # "rice" | "wls" - w_b: Optional[torch.Tensor] = None # (1, N) weights for "wls" - - -def _build_llg_context( - F_obs: torch.Tensor, - hkl_real: torch.Tensor, - s_mag: torch.Tensor, - centric: torch.Tensor, - interpolator: LattmanLoveInterpolator, - real_cell, - spacegroup, - *, - n_shells: int = 20, - batch_size: int = 50, - sigma_a: Optional[torch.Tensor] = None, - eps_factor: Optional[torch.Tensor] = None, - apply_bulk_solvent: bool = False, - solvent_fsol: float = 0.95, - solvent_bsol: float = 300.0, - vrms_strategy: str = "fixed", - vrms_n_residues: Optional[int] = None, - vrms_identity: float = 1.0, - apply_wilson_b: bool = False, - wilson_b_value: Optional[float] = None, - sig_F_obs: Optional[torch.Tensor] = None, - e_convention: type = SmoothSigmaE, - target: str = "rice", -) -> _LLGContext: - """Build the rotation-independent m_LETF1 LLG context (DataMR.cc:1326-1429). - - Two corrections vs. the original implementation, both borrowed from the FRF's - own high-symmetry fixes: - - * **Unique-orbit calc sum.** The moving-model intensity sums ``|E_calc|²`` over - the **distinct** orbit mates via - :func:`torchref.experimental.alignment.frf.preprocessing.epsilon_aware_unroll` - (Phaser's ``if(!duplicate(isym))``), not all ``n_ops`` raw mates. Summing all - mates over-weights axial reflections (ε>1) by ε(h) and orientation-blinds - high-symmetry spacegroups (the 4BX9/6G9X rank-360+ failure). - * **σ_A Eterm convention.** σ_A uses ``eterm_sigma_a`` (the ``2π²/3`` isotropic - Eterm, Ensemble.cc:42 — matching the FRF/Phaser), not the ``2π²`` Luzzati - form which falls off ~3× too fast. - """ - from .frf.preprocessing import ( - compute_v_budget, - epsilon_aware_unroll, - eterm_sigma_a, - ) - - device = F_obs.device - dtype = F_obs.dtype - sym_mats = spacegroup.matrices.to(torch.float64).to(device) - n_ops = int(sym_mats.shape[0]) - N = hkl_real.shape[0] - - # 1. ε(h) per reflection (needed for the ε-corrected obs normalisation). - if eps_factor is None: - # `friedel=False`: the conventional count. The variance budget - # `V = eps - sigma_A**2` wants operations that add coherently and set the - # mean; operations mapping h -> -h change the DISTRIBUTION instead, which - # the Woolfson branch below already handles. Counting them here doubles - # epsilon on every centric reflection -- 6680 of them on 3K7M -- and - # inflates exactly those reflections' variance. - eps_factor = spacegroup.epsilon( - hkl_real.to(torch.long), friedel=False, - ).to(dtype) - eps_factor = eps_factor.to(device) - - # 2. Per-shell ε-corrected Wilson E_obs (Phaser E = F/sqrt(ε·Σ_N)). Dividing - # ε out of the obs is the obs-side analog of the unique-orbit calc dedup: - # both stop axial reflections (ε>1) from being over-weighted on - # high-symmetry spacegroups. - shell_idx = _equal_count_shell_idx(s_mag, n_shells) - conv_obs = e_convention - if sig_F_obs is None and convention_uses_sigma_f(conv_obs): - conv_obs = convention_for_calc(conv_obs) - E_obs = conv_obs( - F_obs, s_mag, centric, sig_F=sig_F_obs, eps=eps_factor, - shell_idx=shell_idx, n_shells=n_shells, - ).E - - # 3. Identity-rotation calc reference → E-normalisation scale for F_calc - # (rotation-invariant: sphere permutation, shell sums preserved). - I_eye = torch.eye(3, dtype=torch.float32, device=device) - F_calc_ref = interpolator.evaluate( - I_eye, hkl_real, real_cell, return_amplitude=True, - ).to(dtype).squeeze(0) # (N,) - # The calc normaliser is the convention's own choice, which is what the - # former `scat_mode` was: "legacy" is CalcShellE (forces _shell = 1 - # in every shell, flattening the model's inter-shell amplitude shape) and - # "absolute" is CalcGlobalE (one global scale, shape preserved). Two knobs - # for one decision meant obs and calc could be normalised by unrelated - # rules; now the same class answers for both sides. - # - # Phaser keeps E_calc physically scaled and carries the model's fraction of - # the cell in scatFactor = AtomScatRatio·SCATTERING/TOTAL_SCAT/NSYMP; for a - # search model that IS the full ASU (the benchmark case) scatFactor reduces - # to 1/n_ops, so the prefactor is unaffected either way. - conv_calc = convention_for_calc(e_convention)( - F_calc_ref, s_mag, centric, eps=eps_factor, - shell_idx=shell_idx, n_shells=n_shells, - ) - calc_norm_per_h = conv_calc.sigma.sqrt().to(dtype).to(device) - - # Optional Wilson-B match (EnsemblePDB.cc:793-851), applied as a per-reflection - # Debye-Waller multiplier on F_calc. - if apply_wilson_b and wilson_b_value is None: - from .frf.preprocessing import fit_relative_wilson_b - wilson_b_value = fit_relative_wilson_b( - F_obs, F_calc_ref, s_mag, n_shells=n_shells, - ) - wilson_b_value = float(wilson_b_value or 0.0) - if apply_wilson_b and abs(wilson_b_value) > 1e-6: - dw = torch.exp(-wilson_b_value * (s_mag * s_mag) / 4.0).to(dtype).to(device) - else: - dw = None - - # 4. σ_A per reflection — FRF/Phaser Eterm (2π²/3 isotropic form), not the - # 2π² Luzzati form. Rotation-independent, no aligned model required. - if sigma_a is None: - if vrms_strategy == "oeffner": - if vrms_n_residues is None: - raise ValueError( - "vrms_strategy='oeffner' requires vrms_n_residues=." - ) - from .frf.preprocessing import oeffner_vrms - delta_vrms_A = oeffner_vrms(int(vrms_n_residues), float(vrms_identity)) - elif vrms_strategy == "fixed": - delta_vrms_A = 0.5 # legacy default - else: - raise ValueError( - f"vrms_strategy={vrms_strategy!r}; expected 'fixed' or 'oeffner'." - ) - sigma_a = eterm_sigma_a(s_mag, delta_vrms_A=delta_vrms_A).to(dtype).to(device) - if apply_bulk_solvent: - from .frf.preprocessing import bulk_solvent_factor - sol = bulk_solvent_factor( - s_mag, fsol=solvent_fsol, bsol=solvent_bsol, - ).to(dtype).to(device) - sigma_a = sigma_a * sol - sigma_a = sigma_a.to(device) - sigma_a2 = sigma_a * sigma_a - - # 5. V(h) — rotation-independent variance budget V = ε − σ_A² (n_mol=1). - V = compute_v_budget(eps_factor, sigma_a, n_mol=1) # (N,) - - # 6. Unique-orbit unroll: distinct mates only (Phaser duplicate-skip). Each - # ASU reflection appears n_ops/ε(h) times, NOT n_ops times. - unrolled_hkl, asu_idx = epsilon_aware_unroll(hkl_real, sym_mats) - unrolled_hkl = unrolled_hkl.to(torch.float64).to(device) - asu_idx = asu_idx.to(device) - - # 7. Broadcastable per-reflection tensors. - E_obs_b = E_obs.unsqueeze(0) # (1, N) - V_b = V.unsqueeze(0) # (1, N) - # eImove = ε(h)·σ_A²·(1/n_ops)·Σ_{distinct mates} |E_calc(R^T·S_k·h)|² - # (Phaser DataMR.cc:1397: thisEsqr *= repsn·scatFactor, scatFactor∝1/NSYMP). - eImove_prefac = (eps_factor * sigma_a2 / float(n_ops)).unsqueeze(0) # (1, N) - centric_b = centric.to(torch.bool).to(device).unsqueeze(0) - - # Per-mate normaliser + DW (rotation preserves |h| → same shell across the - # orbit, so the per-h scale broadcasts to every mate via asu_idx). - sqrt_mean_per_m = calc_norm_per_h[asu_idx] # (M,) - dw_per_m = dw[asu_idx] if dw is not None else None - - # Weights for the least-squares target: the same combined inverse variance - # the rotation function uses, so the two stages agree about how much a - # reflection is worth as well as about what it is being compared to. - w_b = None - if target == "wls": - if sig_F_obs is not None: - w = inverse_variance_weight( - snr_from_amplitude(F_obs, sig_F_obs), sigma_a.to(F_obs.dtype), - eps=eps_factor.to(F_obs.dtype), - ) - else: - w = torch.ones_like(F_obs) - w_b = normalise_weight(w).to(dtype).to(device).view(1, -1) - - return _LLGContext( - interpolator=interpolator, real_cell=real_cell, - unrolled_hkl=unrolled_hkl, asu_idx=asu_idx, N=N, - E_obs_b=E_obs_b, V_b=V_b, eImove_prefac=eImove_prefac, - sqrt_mean_per_m=sqrt_mean_per_m, centric_b=centric_b, - dw_per_m=dw_per_m, dtype=dtype, batch_size=batch_size, - target=target, w_b=w_b, - ) - - -def _llg_for_orientations( - ctx: _LLGContext, - alpha: torch.Tensor, - beta: torch.Tensor, - gamma: torch.Tensor, -) -> torch.Tensor: - """Per-orientation m_LETF1 LLG for a batch of Edmonds-ZYZ Euler angles. - - Returns a ``(n_orient,)`` tensor of LLG values. The calc orbit-sum is over - the deduped mates: evaluate ``|E_calc|²`` on ``ctx.unrolled_hkl`` then - ``scatter_add`` back per ASU reflection. Same Phaser logRel math as before. - """ - R_all = rotation_matrix_from_edmonds_euler_batch( - alpha.to(torch.float64), beta.to(torch.float64), gamma.to(torch.float64), - ).transpose(-1, -2).to(torch.float32) # (n_orient, 3, 3) - n_orient = R_all.shape[0] - sqrt_mean_b = ctx.sqrt_mean_per_m.unsqueeze(0) # (1, M) - dw_b = ctx.dw_per_m.unsqueeze(0) if ctx.dw_per_m is not None else None - - chunks: List[torch.Tensor] = [] - for start in range(0, n_orient, ctx.batch_size): - R_batch = R_all[start:start + ctx.batch_size] # (B, 3, 3) - F_calc_m = ctx.interpolator.evaluate( - R_batch, ctx.unrolled_hkl, ctx.real_cell, return_amplitude=True, - ).to(ctx.dtype) # (B, M) - if dw_b is not None: - F_calc_m = F_calc_m * dw_b - E_calc_m = F_calc_m / sqrt_mean_b # (B, M) - Esq_m = E_calc_m * E_calc_m # (B, M) - B = Esq_m.shape[0] - sum_per_h = torch.zeros( - B, ctx.N, dtype=Esq_m.dtype, device=Esq_m.device, - ) - idx = ctx.asu_idx.unsqueeze(0).expand(B, -1) # (B, M) - sum_per_h.scatter_add_(1, idx, Esq_m) # (B, N) - eImove = ctx.eImove_prefac * sum_per_h # (B, N) - if ctx.target == "wls": - # Weighted least squares on INTENSITIES, which is what E**2 is. - # - # The Rice exists to handle an amplitude: |F_obs| is non-negative - # and its phase is unknown, so the likelihood marginalises over the - # phase and the result is biased upward for weak reflections. None - # of that applies to an intensity, which is unbiased and near - # Gaussian wherever it is measured at all -- the same reason the - # scaler works on intensities rather than amplitudes. - # - # And the job here is to RANK orientations, not to report calibrated - # probabilities. Both sides are already normalised to = 1, so - # the distributional shrinkage the Rice contributes is being applied - # to a quantity that has had its scale fixed by construction. - # - # Sign: a residual is a cost, so negate it to keep "larger is - # better" for every target the caller can pick. - resid = ctx.E_obs_b * ctx.E_obs_b - eImove # (B, N) - chunks.append(-(ctx.w_b * resid * resid).sum(dim=-1)) - else: - sqrt_eImove = eImove.clamp(min=1e-30).sqrt() - ll_acen = phaser_log_rel_rice(ctx.E_obs_b, sqrt_eImove, ctx.V_b) - ll_cen = phaser_log_rel_woolfson(ctx.E_obs_b, sqrt_eImove, ctx.V_b) - ll = torch.where(ctx.centric_b, ll_cen, ll_acen) # (B, N) - chunks.append(ll.sum(dim=-1)) # (B,) - return torch.cat(chunks) - - -@dataclass -class _SimLLGContext: - """Context for the per-candidate-σ_A Sim-LLG surface (no orbit sum). - - Unlike :class:`_LLGContext` (fixed σ_A m_LETF1), this surface FITS σ_A per - shell for each orientation via :func:`llg_for_rotation_batch`. The fixed-σ_A - m_LETF1 surface is locally mis-peaked on high-sym/tNCS cases; the per-candidate - fit re-shapes it so the local maximum sits at the true orientation (the - property `sim_mlrf_rescore` already exhibits). Used as a ``llg_fn`` for - :func:`quadratic_llg_refine`. - """ - - interpolator: LattmanLoveInterpolator - real_cell: object - hkl: torch.Tensor # (N, 3) - F_obs: torch.Tensor # (N,) - shell_idx: torch.Tensor # (N,) int64 - n_shells: int - E_obs: torch.Tensor # (N,) - centric: torch.Tensor # (N,) bool - shell_weights: Optional[torch.Tensor] - n_D_grid: int - interp_var: Optional[torch.Tensor] - batch_size: int - - -def _build_sim_llg_context( - F_obs: torch.Tensor, - hkl_real: torch.Tensor, - s_mag: torch.Tensor, - centric: torch.Tensor, - interpolator: LattmanLoveInterpolator, - real_cell, - *, - n_shells: int = 10, - n_D_grid: int = 21, - batch_size: int = 64, - auto_variance_weights: bool = True, - interp_var: Optional[torch.Tensor] = None, -) -> _SimLLGContext: - """Build the rotation-independent context for the Sim-LLG surface.""" - shell_idx = _equal_count_shell_idx(s_mag, n_shells) - E_obs = WilsonShellE( - F_obs, s_mag, shell_idx=shell_idx, n_shells=n_shells, - ).E - shell_weights = None - if auto_variance_weights: - from .sh import compute_patterson_shell_variance - patt_obs = (E_obs.to(torch.float64) ** 2) - 1.0 - var_p = compute_patterson_shell_variance(patt_obs, shell_idx, P=n_shells) - w = 1.0 / var_p.sqrt() - w = w * (n_shells / w.sum().clamp(min=1e-30)) - shell_weights = w.to(F_obs.dtype) - return _SimLLGContext( - interpolator=interpolator, real_cell=real_cell, hkl=hkl_real, - F_obs=F_obs, shell_idx=shell_idx, n_shells=n_shells, E_obs=E_obs, - centric=centric.to(torch.bool), shell_weights=shell_weights, - n_D_grid=n_D_grid, interp_var=interp_var, batch_size=batch_size, - ) - - -def _sim_llg_for_orientations( - ctx: _SimLLGContext, - alpha: torch.Tensor, - beta: torch.Tensor, - gamma: torch.Tensor, -) -> torch.Tensor: - """Per-orientation Sim-LLG (per-candidate σ_A fit), drop-in ``llg_fn``.""" - R_all = rotation_matrix_from_edmonds_euler_batch( - alpha.to(torch.float64), beta.to(torch.float64), gamma.to(torch.float64), - ).transpose(-1, -2).to(torch.float32) - n_orient = R_all.shape[0] - chunks: List[torch.Tensor] = [] - for start in range(0, n_orient, ctx.batch_size): - R_batch = R_all[start:start + ctx.batch_size] - F_calc = ctx.interpolator.evaluate( - R_batch, ctx.hkl, ctx.real_cell, return_amplitude=True, - ).to(ctx.F_obs.dtype) # (B, N) - llg = llg_for_rotation_batch( - F_obs=ctx.F_obs, shell_idx=ctx.shell_idx, n_shells=ctx.n_shells, - E_obs=ctx.E_obs, centric=ctx.centric, F_calc=F_calc, - shell_weights=ctx.shell_weights, n_D_grid=ctx.n_D_grid, - interp_var=ctx.interp_var, - ) - chunks.append(llg) - return torch.cat(chunks) - - -def _euler_batch_from_matrices(R: torch.Tensor): - """(K,3,3) → three (K,) float64 Euler-angle tensors (Edmonds ZYZ). - - Loops the scalar :func:`edmonds_euler_from_rotation_matrix` (K is small — - the top-K refine set), returning tensors ready for - :func:`_llg_for_orientations`. - """ - a, b, g = [], [], [] - for k in range(R.shape[0]): - aa, bb, gg = edmonds_euler_from_rotation_matrix(R[k]) - a.append(aa) - b.append(bb) - g.append(gg) - return ( - torch.tensor(a, dtype=torch.float64), - torch.tensor(b, dtype=torch.float64), - torch.tensor(g, dtype=torch.float64), - ) - - -def quadratic_llg_refine( - peaks: List[RotationPeak], - ctx: _LLGContext, - *, - k_refine: int = 20, - step_deg: float = 1.5, - n_grid: int = 3, - iterations: int = 1, - max_move_deg: Optional[float] = None, - llg_fn: Optional[Callable[[torch.Tensor, torch.Tensor, torch.Tensor], torch.Tensor]] = None, - verbose: int = 0, -) -> List[RotationPeak]: - """Sub-grid refinement of the top-``k_refine`` peaks on the ML-LLG surface. - - For each peak, sample the LLG on a local **axis-angle** grid around the - orientation, fit a 3-D paraboloid in the tangent space, and step to the vertex - (a Newton step on the LLG). Axis-angle (not Euler α,β,γ) perturbation keeps the - local metric isotropic and avoids the β→0/π gimbal degeneracy where the FRF - returns many peaks. Guards (Hessian negative-definite + vertex inside the - sampled box) fall back to the best sampled grid point, so the refined peak can - never score below its grid value. Refined peaks are re-ranked by their (truly - re-evaluated) LLG; peaks beyond ``k_refine`` are appended unchanged. - - Reliable (sub-degree recovery from grid-resolution hits) on well-behaved - crystals; on high-symmetry / tNCS cases the m_LETF1 surface is locally - mis-peaked (~3–4° off truth) so refinement can WALK AWAY from a good hit — - use ``max_move_deg`` to bound that. - - Parameters - ---------- - peaks - Rescored candidate orientations (Edmonds ZYZ), best-first. - ctx - The :class:`_LLGContext` built for the same data (its - :func:`_llg_for_orientations` defines the surface being optimised). - k_refine - Number of leading peaks to refine. - step_deg - Half-width of the local axis-angle grid (degrees) and the per-iteration - capture radius. Default 1.5 ≈ grid_sampling/2 (the FRF grid half-step). - n_grid - Samples per tangent axis (3 → 27 orientations per peak). - iterations - Newton iterations; each re-centres and halves the grid half-width. - max_move_deg - Safety cap: if the refined orientation moves more than this (geodesic - degrees) from the input peak, keep the input peak instead. ``None`` - disables the cap. Protects against the mis-peaked-surface failure mode. - llg_fn - Surface to optimise: a callable ``(alpha,beta,gamma) -> (M,)`` LLG. If - ``None``, uses the m_LETF1 surface ``_llg_for_orientations(ctx, ...)``. - Pass a per-candidate-σ_A Sim surface (:func:`_sim_llg_for_orientations`) - when the fixed-σ_A m_LETF1 surface is locally mis-peaked (high-sym/tNCS). - """ - if not peaks: - return [] - if llg_fn is None: - def llg_fn(a, b, g): - return _llg_for_orientations(ctx, a, b, g) - k = min(k_refine, len(peaks)) - head = peaks[:k] - tail = peaks[k:] - - # R0 for the head peaks (Edmonds ZYZ, un-transposed — the convention - # `_llg_for_orientations` consumes after its own transpose). - a0 = torch.tensor([p.alpha for p in head], dtype=torch.float64) - b0 = torch.tensor([p.beta for p in head], dtype=torch.float64) - g0 = torch.tensor([p.gamma for p in head], dtype=torch.float64) - R0 = rotation_matrix_from_edmonds_euler_batch(a0, b0, g0) # (k, 3, 3) - R0_orig = R0.clone() # for the move cap - - radius = math.radians(step_deg) - for _ in range(max(1, iterations)): - # Local tangent grid (G, 3), shared across peaks; rebuilt per iteration - # so a 2nd pass zooms in. - lin = torch.linspace(-radius, radius, n_grid, dtype=torch.float64) - gx, gy, gz = torch.meshgrid(lin, lin, lin, indexing="ij") - omegas = torch.stack( - [gx.reshape(-1), gy.reshape(-1), gz.reshape(-1)], dim=-1, - ) # (G, 3) - G = omegas.shape[0] - x, y, z = omegas[:, 0], omegas[:, 1], omegas[:, 2] - ones = torch.ones_like(x) - # Design matrix Φ (G,10): [1, x,y,z, x²,y²,z², xy,xz,yz]. - Phi = torch.stack( - [ones, x, y, z, x * x, y * y, z * z, x * y, x * z, y * z], dim=-1, - ) # (G, 10) - - # Grid orientations: R = rodrigues(ω) @ R0, for every (peak, grid point). - Rloc = axis_angle_to_matrix(omegas) # (G, 3, 3) - R_grid = torch.einsum("gij,kjl->kgil", Rloc, R0) # (k, G, 3, 3) - a_t, b_t, g_t = _euler_batch_from_matrices( - R_grid.reshape(k * G, 3, 3), - ) - llg = llg_fn(a_t, b_t, g_t).reshape(k, G) # (k, G) - - # Batched quadratic fit via ridge-stabilised normal equations: - # θ = (ΦᵀΦ + λI)⁻¹ Φᵀ llg. ΦᵀΦ is shared; only the RHS varies per peak. - PtP = Phi.t() @ Phi # (10, 10) - PtP = PtP + 1e-9 * torch.eye(10, dtype=PtP.dtype) - rhs = torch.einsum("gd,kg->kd", Phi, llg.to(torch.float64)) # (k, 10) - theta = torch.linalg.solve( - PtP.unsqueeze(0).expand(k, -1, -1), rhs.unsqueeze(-1), - ).squeeze(-1) # (k, 10) - - # Gradient b and Hessian H of the paraboloid (in tangent coords). - bvec = theta[:, 1:4] # (k, 3) - H = torch.zeros(k, 3, 3, dtype=torch.float64) - H[:, 0, 0] = 2.0 * theta[:, 4] - H[:, 1, 1] = 2.0 * theta[:, 5] - H[:, 2, 2] = 2.0 * theta[:, 6] - H[:, 0, 1] = H[:, 1, 0] = theta[:, 7] - H[:, 0, 2] = H[:, 2, 0] = theta[:, 8] - H[:, 1, 2] = H[:, 2, 1] = theta[:, 9] - - best_grid = llg.argmax(dim=1) # (k,) - new_R0 = R0.clone() - n_accept = 0 - for kk in range(k): - accept = False - try: - eig = torch.linalg.eigvalsh(H[kk]) - if bool((eig < 0).all()): # genuine maximum - xstar = torch.linalg.solve(H[kk], -bvec[kk]) # (3,) - if torch.isfinite(xstar).all() and float(xstar.norm()) <= radius: - new_R0[kk] = axis_angle_to_matrix(xstar) @ R0[kk] - accept = True - n_accept += 1 - except Exception: - accept = False - if not accept: - new_R0[kk] = R_grid[kk, best_grid[kk]] - R0 = new_R0 - radius = radius / 2.0 - if verbose > 1: - print( - f" quadratic_llg_refine: {n_accept}/{k} vertices accepted " - f"(rest fell back to grid max)", - flush=True, - ) - - # Safety cap: revert any peak whose total move from the input exceeds - # ``max_move_deg`` (geodesic). On a locally mis-peaked surface (high-sym / - # tNCS) the refinement walks toward a spurious LLG max ~3–4° away; capping - # the move means a good hit can never be degraded by more than the cap. - if max_move_deg is not None: - cos_cap = math.cos(math.radians(max_move_deg)) - n_revert = 0 - for kk in range(k): - trace = torch.einsum("ij,ij->", R0[kk], R0_orig[kk]) - cos_move = float(((trace - 1.0) * 0.5).clamp(-1.0, 1.0)) - if cos_move < cos_cap: # moved further than the cap - R0[kk] = R0_orig[kk] - n_revert += 1 - if verbose > 1 and n_revert: - print( - f" quadratic_llg_refine: reverted {n_revert}/{k} peaks that " - f"moved > {max_move_deg}° (mis-peaked-surface guard).", - flush=True, - ) - - # Final TRUE LLG at the refined orientations (never trust the paraboloid). - af, bf, gf = _euler_batch_from_matrices(R0) - llg_final = llg_fn(af, bf, gf) # (k,) - if llg_final.numel() > 1: - std_t = llg_final.std().clamp(min=1e-30) - sig = (llg_final - llg_final.mean()) / std_t - else: - sig = torch.zeros_like(llg_final) - - af_l, bf_l, gf_l = af.tolist(), bf.tolist(), gf.tolist() - llg_l, sig_l = llg_final.tolist(), sig.tolist() - refined = [ - RotationPeak( - alpha=af_l[i], beta=bf_l[i], gamma=gf_l[i], - score=llg_l[i], sigma=sig_l[i], - ) - for i in range(k) - ] - refined.sort(key=lambda p: p.score, reverse=True) - return refined + tail - - -def m_letf1_rescore( - peaks: List[RotationPeak], - F_obs: torch.Tensor, - hkl_real: torch.Tensor, - s_mag: torch.Tensor, - centric: torch.Tensor, - interpolator: LattmanLoveInterpolator, - real_cell, - spacegroup, - *, - n_shells: int = 20, - n_refine: Optional[int] = None, - batch_size: int = 50, - sigma_a: Optional[torch.Tensor] = None, - eps_factor: Optional[torch.Tensor] = None, - verbose: int = 0, - # --- Phaser model-prep knobs (all default OFF; see frf/preprocessing.py) --- - apply_bulk_solvent: bool = False, - solvent_fsol: float = 0.95, - solvent_bsol: float = 300.0, - vrms_strategy: str = "fixed", # "fixed" (legacy delta_vrms=0.5) or "oeffner" - vrms_n_residues: Optional[int] = None, # required if vrms_strategy="oeffner" - vrms_identity: float = 1.0, - apply_wilson_b: bool = False, - wilson_b_value: Optional[float] = None, # if None and apply_wilson_b=True, fitted from data - sig_F_obs: Optional[torch.Tensor] = None, - e_convention: type = SmoothSigmaE, - target: str = "rice", -) -> List[RotationPeak]: - """Phaser-faithful ``m_LETF1`` rescore (DataMR.cc:1326-1429). - - Thin wrapper around :func:`_build_llg_context` + :func:`_llg_for_orientations`. - - Upgrades over :func:`sim_mlrf_rescore`: - - 1. **Unique-orbit symmetry sum on calc** — for each obs reflection ``h``, the - expected moving-model intensity is - ``eImove(h) = ε(h)·σ_A²·(1/n_ops)·Σ_{distinct mates} |E_calc(R^T·S_k·h)|²`` - summed over the **distinct** orbit mates only (Phaser's - ``if(!duplicate(isym))``, DataMR.cc:1371-1404), via - :func:`torchref.experimental.alignment.frf.preprocessing.epsilon_aware_unroll` + - ``scatter_add``. Summing all ``n_ops`` raw mates over-weights axial - reflections by ε(h) and orientation-blinds high-symmetry spacegroups. - - 2. **Per-reflection variance budget** ``V(h) = ε(h) − σ_A²(s)·n_mol`` from - :func:`torchref.experimental.alignment.frf.preprocessing.compute_v_budget` - (DataMR.cc:949,1411). For cross-rotation with no fixed model. - - 3. **Phaser ``logRelRice`` / ``logRelWoolfson``** as the per-reflection LL - formula (RiceWoolfson.cc:25-74), commensurable with Phaser's m_LETF1 - output. Different normalisation from our generic - :func:`rice_log_likelihood` (factor of 2 in the Bessel argument; ``V`` is - twice the standard Rice variance for acentric). - - Returns peaks ranked by descending LL with ``score = LL`` and - ``sigma = (LL − μ_batch) / σ_batch``, drop-in for downstream consumers. - - Parameters - ---------- - peaks - Candidate orientations from the FRF, ZYZ Edmonds Euler. - F_obs, hkl_real, s_mag, centric - Per-reflection obs arrays (anisotropy-corrected F_obs is fine). - interpolator, real_cell - ``LattmanLoveInterpolator`` for the model molecular transform and the - crystal real cell. - spacegroup : SpaceGroup - The crystal's space group. Passed as the object rather than its - ``matrices`` because the multiplicity this needs is a method on it: - ``epsilon(hkl, friedel=False)``, the conventional count, which a bare - tensor of rotations cannot answer. - sigma_a : (N,) tensor, optional - Per-reflection σ_A. If ``None``, fitted on-the-fly from the identity - rotation's |F_calc| via :func:`fit_sigma_a_per_shell` and interpolated - per shell. - eps_factor : (N,) tensor, optional - Per-reflection multiplicity ε(h). If ``None``, computed via - :meth:`torchref.symmetry.symmetry.Symmetry.epsilon` with ``friedel=False``. - n_refine, batch_size, verbose - As in :func:`sim_mlrf_rescore`. - """ - if not peaks: - return [] - if n_refine is None: - n_refine = len(peaks) - head = peaks[:n_refine] - tail = peaks[n_refine:] - - ctx = _build_llg_context( - F_obs, hkl_real, s_mag, centric, interpolator, real_cell, spacegroup, - n_shells=n_shells, batch_size=batch_size, sigma_a=sigma_a, - eps_factor=eps_factor, apply_bulk_solvent=apply_bulk_solvent, - solvent_fsol=solvent_fsol, solvent_bsol=solvent_bsol, - vrms_strategy=vrms_strategy, vrms_n_residues=vrms_n_residues, - vrms_identity=vrms_identity, apply_wilson_b=apply_wilson_b, - wilson_b_value=wilson_b_value, sig_F_obs=sig_F_obs, - e_convention=e_convention, target=target, - ) - - alpha_t = torch.tensor([p.alpha for p in head], dtype=torch.float64) - beta_t = torch.tensor([p.beta for p in head], dtype=torch.float64) - gamma_t = torch.tensor([p.gamma for p in head], dtype=torch.float64) - llgs_t = _llg_for_orientations(ctx, alpha_t, beta_t, gamma_t) - if verbose > 1: - print(f" m_LETF1 scored {len(head)} peaks", flush=True) - - mean_t = llgs_t.mean() - std_t = llgs_t.std().clamp(min=1e-30) - sigmas_t = (llgs_t - mean_t) / std_t - - llgs_list = llgs_t.tolist() - sigmas_list = sigmas_t.tolist() - rescored = [ - RotationPeak( - alpha=p.alpha, beta=p.beta, gamma=p.gamma, - score=llg, sigma=sigma, - ) - for p, llg, sigma in zip(head, llgs_list, sigmas_list) - ] - rescored.sort(key=lambda r: r.score, reverse=True) - return rescored + tail - - diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 8756c303..8e4d3b5a 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -1,35 +1,39 @@ -""" -Molecular replacement pipeline: the single canonical MR orchestrator. - -Implements the classic Phaser-style molecular-replacement tree: - -1. **Fast Rotation Function (FRF)** — Phaser-faithful Bessel-radial × SH - expansion (dense P1-box calc + auto_lmax), then ML rescoring - (``m_letf1_rescore`` / ``sim_mlrf_rescore``) to rank candidate orientations. -2. **Fast Translation Function (FTF)** — for *each* of the top-N rotation - candidates, an amplitude-correlation translation search (optionally - re-ranked by a Rice/Woolfson LLG) followed by an analytical-R local refine. -3. **Post-refinement** — optional dense rotation re-sampling on the ML-LLG - surface, then an LBFGS rigid-body polish on (R, t) (with optional B-factor - co-refinement and Gaussian restraints). - -Each rotation candidate is carried through translation + refinement -independently; the candidates are ranked by their refined R-factor and the -best is returned (a Phaser-style multi-candidate tree, with early-stopping once -a candidate beats ``rfactor_converged``). The user-facing solvent-aware R-work -is computed once, on the winner. - -``align_model_to_data`` delegates to -this class — it is the implementation of record. The heavy crystallographic -stage helpers live in :mod:`torchref.experimental.alignment.align`, -:mod:`~torchref.experimental.alignment.translation` and -:mod:`~torchref.experimental.alignment.ml_rotation`; this module owns the +"""Molecular replacement: the FRF hands a shortlist to the FTF. + +Two stages, and the division of labour between them is the design: + +1. **Fast Rotation Function** — Phaser-faithful Bessel-radial × SH expansion + against a dense P1-box calc. It is a *shortlist generator*. It does not have + to rank well, and measurably does not: over 30 seeded cells it puts truth at + rank 0 in 6 of them. What it does reliably is put truth somewhere in the top + twenty. +2. **Fast Translation Function** — for *each* of the top-N orientations, a + Crowther-Blow amplitude-correlation search over the fractional cell, + optionally re-ranked by a Rice/Woolfson LLG, then an analytical-R local + refine. On the same 30 cells it puts truth at rank 0 in 27. Rotation ghosts + are morphologically identical to truth in a rotation function by + construction; they are not identical once the crystal is involved. + +There is deliberately nothing between them. An ML re-ranking of the FRF peaks +used to sit there and was removed: it reorders a shortlist that already contains +truth, and end-to-end pose recovery was 18/30 with it against 24/30 without +(McNemar p = 0.031, 6-0 discordant). + +Each orientation is placed independently and the candidates are ranked by the +translation search's analytical R, with early stopping once one beats +``rfactor_converged``. The user-facing solvent-aware R-work is computed once, on +the winner. The pipeline returns a *placement* -- refining it is the caller's +job, and downstream refinement does it better than a bolted-on polish did. + +``align_model_to_data`` delegates here; this class is the implementation of +record. The crystallographic stages live in +:mod:`torchref.experimental.alignment.align` and +:mod:`~torchref.experimental.alignment.translation`; this module owns the control flow that wires them together. """ from __future__ import annotations -import math from dataclasses import dataclass from typing import List, Optional, Tuple, TYPE_CHECKING @@ -45,25 +49,14 @@ _external_rwork, _prepare_frf_inputs, ) -from .e_values import SmoothSigmaE -from .frf.rotation_utils import ( - axis_angle_to_matrix, - edmonds_euler_from_rotation_matrix, - rotation_matrix_from_edmonds_euler, -) +from .frf.rotation_utils import rotation_matrix_from_edmonds_euler from .frf.types import RotationPeak from .rotation_search import search_peaks -from .lattman_love import LattmanLoveInterpolator, estimate_interp_var -from .ml_rotation import ( - fit_sigma_a_per_shell, - m_letf1_rescore, - sim_mlrf_rescore, -) -from .rigid_body import RigidBodyRefinement from .sh import assign_shells, equal_count_shell_edges from .translation import ( TranslationPeak, amplitude_translation_search, + fit_sigma_a_per_shell, llg_translation_rescore, local_translation_refine, precompute_G_for_rotation, @@ -224,36 +217,14 @@ def __init__( d_min: float = 4.0, d_max: float = 15.0, n_shells: int = 20, - ll_max_res_A: float = 3.0, - ll_padding_factor: float = 2.0, n_rotation_peaks: int = 500, - n_ml_refine: int = 20, model_error_A: Optional[float] = None, - # --- rescore --- - rescore_engine: str = "m_letf1", - rescore_e_convention: type = SmoothSigmaE, - auto_variance_weights: bool = True, - use_interp_var: bool = False, - subpeak_refine: bool = False, - subpeak_refine_k: int = -1, - subpeak_refine_step_deg: float = 1.5, - subpeak_refine_iters: int = 1, - subpeak_refine_max_move_deg: Optional[float] = 1.5, # --- candidate tree --- n_rotation_candidates: int = 15, n_translation_peaks: int = 20, n_translation_candidates: int = 3, translation_grid_steps: int = 16, use_llg_tf: bool = False, - # --- post-refine --- - do_joint_refine: bool = True, - dense_rotation_refine: bool = True, - joint_refine_max_res_A: float = 4.0, - joint_refine_expected_rot_error: float = 0.1, - refine_b: bool = False, - sigma_rot_deg: float = 0.0, - sigma_trans_ang: float = 0.0, - sigma_b: float = 0.0, # --- early stop --- min_tries: int = 3, max_tries: Optional[int] = None, @@ -267,10 +238,7 @@ def __init__( self.d_min = d_min self.d_max = d_max self.n_shells = n_shells - self.ll_max_res_A = ll_max_res_A - self.ll_padding_factor = ll_padding_factor self.n_rotation_peaks = n_rotation_peaks - self.n_ml_refine = n_ml_refine # Expected r.m.s. coordinate error of the search model, in Angstrom: # it sets the sigma_A fall-off in the rotation function. When the caller # does not know it, estimate it from the model's length the way Phaser @@ -282,31 +250,12 @@ def __init__( model_error_A = oeffner_vrms(n_residues, 1.0) self.model_error_A = float(model_error_A) - self.rescore_engine = rescore_engine - self.rescore_e_convention = rescore_e_convention - self.auto_variance_weights = auto_variance_weights - self.use_interp_var = use_interp_var - self.subpeak_refine = subpeak_refine - self.subpeak_refine_k = subpeak_refine_k - self.subpeak_refine_step_deg = subpeak_refine_step_deg - self.subpeak_refine_iters = subpeak_refine_iters - self.subpeak_refine_max_move_deg = subpeak_refine_max_move_deg - self.n_rotation_candidates = n_rotation_candidates self.n_translation_peaks = n_translation_peaks self.n_translation_candidates = n_translation_candidates self.translation_grid_steps = translation_grid_steps self.use_llg_tf = use_llg_tf - self.do_joint_refine = do_joint_refine - self.dense_rotation_refine = dense_rotation_refine - self.joint_refine_max_res_A = joint_refine_max_res_A - self.joint_refine_expected_rot_error = joint_refine_expected_rot_error - self.refine_b = refine_b - self.sigma_rot_deg = sigma_rot_deg - self.sigma_trans_ang = sigma_trans_ang - self.sigma_b = sigma_b - self.min_tries = min_tries self.max_tries = max_tries self.rfactor_converged = rfactor_converged @@ -348,23 +297,22 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: frf = _prepare_frf_inputs( self.model, self.data, d_min=self.d_min, d_max=self.d_max, n_shells=self.n_shells, - ll_padding_factor=self.ll_padding_factor, - ll_max_res_A=self.ll_max_res_A, verbose=self.verbose, + verbose=self.verbose, ) timer.stop("0_data_prep") self._frf = frf - # --- Stage 1+2: FRF rotation search + ML rescore --- - rescored = self._rotation_candidates(frf) - if not rescored: + # --- Stage 1: FRF rotation search --- + candidates = self._rotation_candidates(frf) + if not candidates: raise RuntimeError("Rotation search produced no peaks.") if not do_translation: - rotated, R_rec = self._make_rotated(rescored[0]) - top = rescored[0] + rotated, R_rec = self._make_rotated(candidates[0]) + top = candidates[0] if self.verbose > 0: print( - f"mr: top peak LLG = {top.score:.2f} " + f"mr: top peak RF = {top.score:.2f} " f"(σ_Z = {top.sigma:.2f}); applying R⁻¹ to coords.", flush=True, ) @@ -381,9 +329,9 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: ) ] - # --- Stage 3+: per-candidate translation + post-refine tree --- + # --- Stage 2: per-candidate translation search --- self._prepare_translation_arrays() - n_rot = min(self.n_rotation_candidates, len(rescored)) + n_rot = min(self.n_rotation_candidates, len(candidates)) max_tries = self.max_tries if self.max_tries is not None else n_rot if self.verbose > 0 and n_rot > 1: print( @@ -396,7 +344,7 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: solutions: List[MRSolution] = [] best_r = float("inf") for k in range(n_rot): - peak_k = rescored[k] + peak_k = candidates[k] rotated_k, R_rec_k = self._make_rotated(peak_k) if self.verbose > 0: print( @@ -411,35 +359,13 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: continue r_analytic, t_refined = placement - refined = rotated_k.copy().translate( + placed = rotated_k.copy().translate( t_refined.to(self.model.dtype_float), fractional=True, ) r_rank = r_analytic - if self.do_joint_refine: - if self.dense_rotation_refine: - refined = self._dense_rotation_refine(refined) - polished, rb_result = self._rigid_body_polish(refined) - if rb_result.final_r_factor <= rb_result.initial_r_factor: - refined = polished - r_rank = rb_result.final_r_factor - if self.verbose > 0: - print( - f" joint polish {rb_result.initial_r_factor:.4f} → " - f"{rb_result.final_r_factor:.4f} (no-solvent R)", - flush=True, - ) - else: - r_rank = rb_result.initial_r_factor - if self.verbose > 0: - print( - f" joint polish kept original " - f"({rb_result.initial_r_factor:.4f} ≤ " - f"{rb_result.final_r_factor:.4f} no-solvent R)", - flush=True, - ) - - refined.last_alignment_rotation = R_rec_k - refined.last_alignment_translation = t_refined + + placed.last_alignment_rotation = R_rec_k + placed.last_alignment_translation = t_refined solutions.append( MRSolution( rotation=R_rec_k.detach().cpu().numpy(), @@ -447,7 +373,7 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: rotation_score=float(peak_k.score), translation_score=float(r_analytic), r_factor=float(r_rank), - model=refined, + model=placed, ) ) best_r = min(best_r, r_rank) @@ -487,12 +413,10 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: return solutions # ------------------------------------------------------------------ - # Stage 1+2: rotation search + ML rescore + # Stage 1: rotation search # ------------------------------------------------------------------ def _rotation_candidates(self, frf) -> list: - """FRF rotation search followed by ML rescoring of the top peaks.""" - data = self.data - device = self.device + """FRF rotation search; the peaks it returns, ranked by its own score.""" timer = self._timer timer.start("3_rotation_search") @@ -503,114 +427,19 @@ def _rotation_candidates(self, frf) -> list: flush=True, ) peaks, _lmax, _d_min = search_peaks( - self.model, data, self.model_error_A, + self.model, self.data, self.model_error_A, U_aniso=frf.U_aniso, n_peaks=self.n_rotation_peaks, verbose=self.verbose, ) timer.stop("3_rotation_search") - if self.rescore_engine not in ("m_letf1", "sim", "none"): - raise ValueError( - f"rescore_engine={self.rescore_engine!r}; " - "expected 'm_letf1' (default), 'sim' or 'none'." - ) - - F_obs = frf.F_obs - sig_F = frf.sig_F - hkl = frf.hkl - s_mag = frf.s_mag - centric = frf.centric - ll = frf.ll - rescore_n_shells = max(self.n_shells // 2, 8) - - # No ML rescore: rank candidates by the raw FRF score and let the - # multi-candidate tree (FTF + refine + R-ranking) do the discrimination. - # Sub-peak refinement is still available here: it sharpens each - # orientation in place and does not reorder, so it is independent of - # which engine (if any) ranks the candidates. - if self.rescore_engine == "none": - if self.verbose > 0: - print("mr: ML rescore DISABLED — using raw FRF peak " - "ranking (RFZ).", flush=True) - ranked = sorted(peaks, key=lambda p: p.score, reverse=True) - if self.subpeak_refine: - ranked = self._subpeak_refine(ranked, F_obs, hkl, s_mag, - centric, ll, rescore_n_shells) - return ranked - - interp_var_main: Optional[torch.Tensor] = None - if self.use_interp_var: - rescore_edges, _ = equal_count_shell_edges(s_mag, rescore_n_shells) - rescore_shell_idx = assign_shells(s_mag, rescore_edges) - interp_var_main = estimate_interp_var( - ll, hkl, data.cell, rescore_shell_idx, rescore_n_shells, - ).to(F_obs.dtype) - - timer.start("4_ml_rescore") - if self.verbose > 0: - print( - f"mr: ML rescoring top " - f"{min(len(peaks), self.n_ml_refine)} peaks…", - flush=True, - ) - if self.rescore_engine == "m_letf1": - rescored = m_letf1_rescore( - peaks, F_obs, hkl, s_mag, centric, ll, data.cell, - data.spacegroup, - n_shells=rescore_n_shells, - n_refine=min(len(peaks), self.n_ml_refine), - batch_size=50, verbose=self.verbose, - sig_F_obs=sig_F, - e_convention=self.rescore_e_convention, - ) - if self.subpeak_refine: - rescored = self._subpeak_refine(rescored, F_obs, hkl, s_mag, - centric, ll, rescore_n_shells, - sig_F=sig_F) - else: # legacy Sim/Rice approximation - rescored = sim_mlrf_rescore( - peaks, F_obs, hkl, s_mag, centric, ll, data.cell, - n_shells=rescore_n_shells, - n_refine=min(len(peaks), self.n_ml_refine), - batch_size=50, verbose=self.verbose, - auto_variance_weights=self.auto_variance_weights, - interp_var=interp_var_main, - ) - timer.stop("4_ml_rescore") - return rescored - - def _subpeak_refine(self, rescored, F_obs, hkl, s_mag, centric, ll, - rescore_n_shells, *, sig_F=None): - """Quadratic tangent-space Newton sharpening of the top orientations.""" - from .ml_rotation import _build_llg_context, quadratic_llg_refine - - data = self.data - device = self.device - self._timer.start("4b_subpeak_refine") - ctx = _build_llg_context( - F_obs, hkl, s_mag, centric, ll, data.cell, - data.spacegroup, - n_shells=rescore_n_shells, batch_size=50, - sig_F_obs=sig_F, - e_convention=self.rescore_e_convention, - ) - k = self.subpeak_refine_k if self.subpeak_refine_k > 0 else self.n_rotation_candidates - k = min(k, len(rescored)) - rescored = quadratic_llg_refine( - rescored, ctx, k_refine=k, - step_deg=self.subpeak_refine_step_deg, - iterations=self.subpeak_refine_iters, - max_move_deg=self.subpeak_refine_max_move_deg, - verbose=self.verbose, - ) - self._timer.stop("4b_subpeak_refine") - if self.verbose > 0: - print( - f"mr: sub-peak refined top {k} orientations " - f"on the ML-LLG surface (step={self.subpeak_refine_step_deg}°).", - flush=True, - ) - return rescored + # Rank by the FRF's own score and hand the shortlist to the + # translation search. There is no rescore here by design: an ML + # re-ranking of these peaks was measured to lower end-to-end pose + # recovery from 24/30 to 18/30 (McNemar p = 0.031), because it reorders + # a shortlist that already contains truth and sometimes pushes truth + # out of it. The translation function does the discrimination. + return sorted(peaks, key=lambda p: p.score, reverse=True) def _make_rotated(self, peak: "RotationPeak"): """Rotate the search model onto a candidate orientation. @@ -630,7 +459,7 @@ def _make_rotated(self, peak: "RotationPeak"): return rot, R_rec # ------------------------------------------------------------------ - # Stage 3: per-candidate translation search + local refine + # Stage 2: per-candidate translation search + local refine # ------------------------------------------------------------------ def _prepare_translation_arrays(self) -> None: """Resolution/validity-masked obs amplitudes + Miller indices.""" @@ -655,7 +484,6 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: rotation candidate, or ``None`` if no translation peaks were found. """ data = self.data - device = self.device timer = self._timer eye3 = self._eye3 @@ -696,13 +524,6 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: print(f" top translation t={tt} score={t_peaks[0].score:.4f}", flush=True) - # do_joint_refine=False: take the top translation peak directly (no - # local refine), rank by its correlation score (negated so lower=better - # like an R-factor). - if not self.do_joint_refine: - t_top = torch.as_tensor(t_peaks[0].translation, dtype=torch.float64) - return -float(t_peaks[0].score), t_top - best = None for k_t, tp in enumerate(t_peaks[:self.n_translation_candidates]): t_init = torch.as_tensor(tp.translation, dtype=torch.float64) @@ -797,146 +618,3 @@ def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): ) for i in order ] - - # ------------------------------------------------------------------ - # Stage 4+5: dense rotation re-sampling + rigid-body polish - # ------------------------------------------------------------------ - def _dense_rotation_refine(self, refined): - """Two-pass dense rotation re-sampling on the ML-LLG surface. - - Zooms the orientation onto the (sharper) ML-LLG basin at the found - translation before the LBFGS polish. Returns the re-rotated model. - """ - data = self.data - device = self.device - timer = self._timer - - timer.start("8_dense_R_ll_build") - refined_p1 = refined.copy() - refined_p1.spacegroup = "P 1" - ll_refine = LattmanLoveInterpolator( - refined_p1, padding_factor=self.ll_padding_factor, - max_res_A=self.ll_max_res_A, verbose=0, - ) - timer.stop("8_dense_R_ll_build") - - tmask = self._tmask - hkl_keep = self._hkl_keep - F_obs_amp = self._F_obs_amp - centric_keep = ( - data.centric[tmask].to(torch.bool).to(device) if hasattr(data, "centric") - else torch.zeros(hkl_keep.shape[0], dtype=torch.bool, device=device) - ) - rec_basis_keep = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) - s_mag_keep = (hkl_keep.to(torch.float64) @ rec_basis_keep).norm(dim=-1) - rescore_n_shells = max(self.n_shells // 2, 8) - - n_per_axis_pass = [9, 5] - zoom_factor = 4.0 - radii = [ - float(self.joint_refine_expected_rot_error), - float(self.joint_refine_expected_rot_error) / zoom_factor, - ] - R_accumulated = torch.eye(3, dtype=torch.float64) - - interp_var_dense: Optional[torch.Tensor] = None - if self.use_interp_var: - dense_edges, _ = equal_count_shell_edges(s_mag_keep, rescore_n_shells) - dense_shell_idx = assign_shells(s_mag_keep, dense_edges) - interp_var_dense = estimate_interp_var( - ll_refine, hkl_keep, data.cell, dense_shell_idx, rescore_n_shells, - ).to(F_obs_amp.dtype) - - for pass_idx, max_perturb_rad in enumerate(radii): - n_per_axis = n_per_axis_pass[pass_idx] - coords_r = torch.linspace( - -max_perturb_rad, max_perturb_rad, n_per_axis, dtype=torch.float64, - ) - wx, wy, wz = torch.meshgrid(coords_r, coords_r, coords_r, indexing="ij") - omegas = torch.stack([wx.flatten(), wy.flatten(), wz.flatten()], dim=-1) - R_perturbs = axis_angle_to_matrix(omegas) - R_cand_full = R_perturbs @ R_accumulated - cand_peaks = [] - for R_c in R_cand_full: - a, b, g = edmonds_euler_from_rotation_matrix(R_c) - cand_peaks.append(RotationPeak(alpha=a, beta=b, gamma=g, - score=0.0, sigma=0.0)) - if self.verbose > 0: - print( - f" dense R pass {pass_idx + 1} " - f"({n_per_axis}³={omegas.shape[0]} perturbations, " - f"±{math.degrees(max_perturb_rad):.2f}°)…", - flush=True, - ) - rescore_batch = max(4, min(100, 1_000_000 // max(hkl_keep.shape[0], 1))) - timer.start("9_dense_R_rescore") - if self.rescore_engine == "m_letf1": - rescored_refine = m_letf1_rescore( - cand_peaks, F_obs_amp, hkl_keep, s_mag_keep, centric_keep, - ll_refine, data.cell, - data.spacegroup, - n_shells=rescore_n_shells, - n_refine=len(cand_peaks), batch_size=rescore_batch, verbose=0, - ) - else: - rescored_refine = sim_mlrf_rescore( - cand_peaks, F_obs_amp, hkl_keep, s_mag_keep, centric_keep, - ll_refine, data.cell, - n_shells=rescore_n_shells, - n_refine=len(cand_peaks), batch_size=rescore_batch, - verbose=0, n_D_grid=11, interp_var=interp_var_dense, - ) - timer.stop("9_dense_R_rescore") - top = rescored_refine[0] - best_idx = next( - i for i, p in enumerate(cand_peaks) - if p.alpha == top.alpha and p.beta == top.beta - and p.gamma == top.gamma - ) - R_accumulated = R_cand_full[best_idx] - if self.verbose > 0: - print( - f" pass {pass_idx + 1} best LLG={top.score:.2f}, " - f"|ω|={omegas[best_idx].norm().item() * 180 / math.pi:.3f}°", - flush=True, - ) - - return refined.copy().rotate( - R_accumulated.T.to(self.model.dtype_float).contiguous(), - ) - - def _rigid_body_polish(self, refined): - """LBFGS rigid-body polish on (R, t) — returns ``(polished, result)``. - - The model is pre-rotated/translated; the refinement optimises a small - delta with ``initial_translation=0``. ``result`` carries the no-solvent - initial/final R-work used by the caller's accept/reject gate. - """ - timer = self._timer - timer.start("11_lbfgs_polish") - rb = RigidBodyRefinement( - refined, self.data, - initial_translation=torch.zeros( - 3, dtype=torch.float32, device=refined.device, - ), - expected_rotational_error=self.joint_refine_expected_rot_error, - max_res=self.joint_refine_max_res_A, - device=refined.device, - verbose=max(0, self.verbose - 1), - refine_b=self.refine_b, - sigma_rot_deg=self.sigma_rot_deg, - sigma_trans_ang=self.sigma_trans_ang, - sigma_b=self.sigma_b, - ) - rb_result = rb.refine() - with torch.no_grad(): - R_polish = rb.get_rotation_matrix().detach() - t_polish = rb.translation_frac.detach() - # .copy() first: the caller keeps `refined` when the polish does not - # improve R, so `polished` must not be the same object. - polished = refined.copy().rotate(R_polish.to(self.model.dtype_float)) - polished = polished.translate( - t_polish.to(self.model.dtype_float), fractional=True, - ) - timer.stop("11_lbfgs_polish") - return polished, rb_result diff --git a/torchref/experimental/alignment/rigid_body.py b/torchref/experimental/alignment/rigid_body.py deleted file mode 100644 index dcd397b8..00000000 --- a/torchref/experimental/alignment/rigid_body.py +++ /dev/null @@ -1,552 +0,0 @@ -""" -Rigid Body Refinement for Molecular Replacement. - -Implements rigid body refinement where rotation and translation parameters -are optimized to minimize the difference between F_calc and F_obs. -This follows the rotation search (FRF) and translation search stages. - -The refinement optimizes 6 parameters: -- 3 rotation angles (alpha, beta, gamma) as small perturbations -- 3 translation components (fractional coordinates) - -Key design: Bypasses Model/MixedTensor to maintain gradient flow. -Stores all required tensors and uses FFT.compute_structure_factors() directly. - -Gradient flow: - d_alpha → rotation_matrix → xyz_transformed → FFT.compute_structure_factors() → loss - -Uses ScalerBase for proper crystallographic scaling during optimization. -""" - -from dataclasses import dataclass -from typing import Optional, Tuple - -import math - -import numpy as np -import torch -import torch.nn as nn - -from torchref.base import rotation_matrix_euler_zyz -from torchref.config import get_default_device -from torchref.model import SfFFT -from torchref.refinement.targets import RiceXrayTarget -from torchref.scaling import Scaler -from torchref.utils.device_mixin import DeviceMixin - - -@dataclass -class RigidBodyResult: - """ - Results from rigid body refinement. - - Attributes - ---------- - final_rotation : torch.Tensor - Final Euler angles (alpha, beta, gamma) in radians. - final_translation_frac : torch.Tensor - Final translation in fractional coordinates. - initial_r_factor : float - R-factor before refinement. - final_r_factor : float - R-factor after refinement. - final_ml_loss : float - Final ML loss value. - n_steps : int - Number of optimization steps performed. - converged : bool - Whether the refinement converged. - """ - - final_rotation: torch.Tensor - final_translation_frac: torch.Tensor - initial_r_factor: float - final_r_factor: float - final_ml_loss: float - n_steps: int - LBFGS_iterations: int - LBFGS_function_evaluations: int - converged: bool - - -class RigidBodyRefinement(DeviceMixin, nn.Module): - """ - Rigid body refinement using FFT directly (bypasses Model/MixedTensor). - - Optimizes 6 parameters (3 rotation + 3 translation) to maximize - agreement between calculated and observed structure factors using - Maximum Likelihood target. - - Key design: Extracts all tensors from Model once at init, then uses - FFT.compute_structure_factors() directly. This maintains gradient flow: - d_alpha → rotation_matrix → xyz_transformed → FFT → loss - - Parameters - ---------- - model : ModelFT - Model with atomic coordinates (tensors extracted, not stored). - data : ReflectionData - Observed reflection data. - initial_rotation : torch.Tensor, optional - Initial Euler angles (alpha, beta, gamma) in radians. - Default is (0, 0, 0). - initial_translation : torch.Tensor, optional - Initial fractional translation vector (3,). - Default is [0, 0, 0]. - device : torch.device, optional - Computation device. Default is CPU. - - Attributes - ---------- - d_alpha, d_beta, d_gamma : nn.Parameter - Refinable rotation perturbations. - translation_frac : nn.Parameter - Refinable fractional translation. - scaler : ScalerBase - Scaler for crystallographic scaling (jointly optimized). - """ - - def __init__( - self, - model, # ModelFT - data, # ReflectionData - expected_rotational_error: float = 0.1, - initial_rotation: torch.Tensor = torch.tensor( - [0.0, 0.0, 0.0], dtype=torch.float32 - ), - initial_translation: Optional[torch.Tensor] = None, - device: torch.device = None, - rfactor_converged_threshold: float = 0.45, - max_res: float = 4.0, - verbose: int = 1, - refine_b: bool = False, - sigma_rot_deg: float = 0.0, - sigma_trans_ang: float = 0.0, - sigma_b: float = 0.0, - ): - super().__init__() - if device is None: - device = get_default_device() - self.device = device - self.data = data - # Phase C: B-refine + Phaser-style Gaussian restraints. All off by - # default so previously-passing trajectories are unaffected. - self.refine_b = bool(refine_b) - self.sigma_rot_rad = math.radians(float(sigma_rot_deg)) if sigma_rot_deg > 0 else 0.0 - self.sigma_trans_ang = float(sigma_trans_ang) - self.sigma_b = float(sigma_b) - - xyz_iso, adp_iso, occ_iso, A_iso, B_iso = model.get_iso() - - self.register_buffer("xyz_initial", xyz_iso.detach().clone().to(device)) - self.register_buffer("adp_iso", adp_iso.detach().clone().to(device)) - self.register_buffer("occ_iso", occ_iso.detach().clone().to(device)) - self.register_buffer("A_iso", A_iso.detach().clone().to(device)) - self.register_buffer("B_iso", B_iso.detach().clone().to(device)) - centroid = torch.mean(self.xyz_initial, dim=0) - self.register_buffer("centroid", centroid) - - # Move Cell (which holds fractional_matrix etc.) to the target device - # so RigidBodyRefinement.get_transformed_xyz can mm a GPU tensor - # against fractional_matrix.T without "mat2 on cpu" crashes. - self.cell = data.cell.to(device=device) if hasattr(data.cell, "to") else data.cell - self.spacegroup = data.spacegroup - - # Forward `device` to SfFFT — its `setup_grid` reads `self.device` - # to allocate the real-space grid; without this, the grid lands on - # CPU even when the surrounding RigidBodyRefinement is on cuda, and - # the joint refine crashes with "mat2 on cuda, others on cpu" at - # the first compute_structure_factors call. - self.fft = SfFFT(self.cell, self.spacegroup, max_res=max_res, - device=device) - - self.verbose = verbose - self.rfactor_converged_threshold = rfactor_converged_threshold - # Get anisotropic atoms if any - xyz_aniso, u_aniso, occ_aniso, A_aniso, B_aniso = model.get_aniso() - self.has_aniso = len(xyz_aniso) > 0 - if self.has_aniso: - self.xyz_aniso_original = xyz_aniso.detach().clone().to(device) - self.u_aniso = u_aniso.detach().clone().to(device) - self.occ_aniso = occ_aniso.detach().clone().to(device) - self.A_aniso = A_aniso.detach().clone().to(device) - self.B_aniso = B_aniso.detach().clone().to(device) - else: - self.xyz_aniso_original = None - self.u_aniso = None - self.occ_aniso = None - self.A_aniso = None - self.B_aniso = None - - # Store initial rotation - self.register_buffer( - "initial_rotation", initial_rotation.to(device=device).clone() - ) - - self.rotation_parameters = nn.Parameter(torch.zeros(3, device=device)) - self.expected_rotational_error = expected_rotational_error - - # Refinable translation (fractional coordinates) - if initial_translation is None: - initial_translation = torch.zeros(3, device=device) - else: - initial_translation = initial_translation.to(device=device).clone() - self.translation_frac = nn.Parameter(initial_translation) - - # Phase C: per-atom B-factor perturbation. Held as nn.Parameter even - # when refine_b=False (gradient just won't flow); cost is negligible - # and the codepath stays uniform. - n_iso = int(self.adp_iso.shape[0]) - self.delta_b_iso = nn.Parameter( - torch.zeros(n_iso, device=device, dtype=self.adp_iso.dtype), - ) - - # Use `Scaler` for per-bin scales + anisotropy correction, but skip - # the bulk-solvent setup. The solvent mask is computed once from - # `model.xyz()` and goes stale as the joint refine moves atoms; the - # mismatch then biases the LBFGS gradient. Better to leave solvent - # out of the joint refine — `align_model_to_data` does a fresh - # solvent-aware Scaler refit on the final polished model for the - # user-facing R-work. - self.scaler = Scaler(model=model, data=data, nbins=20, - verbose=0, device=device) - # Initial scaler fit only needs grad through scaler params, not - # through the rigid-body forward. Without detaching, the SfFFT - # density-build intermediates from the initial forward stay pinned - # by the autograd graph until `rb` is freed — and on multi-trial - # runs that adds ~5 GB of GPU residue per alignment. - with torch.no_grad(): - fcalc_initial = self().detach() - self.scaler.calc_initial_scale(fcalc_initial) - self.scaler.setup_anisotropy_correction() - self.scaler.refine_lbfgs(fcalc=fcalc_initial) - - self.xray_target = RiceXrayTarget(data=self.data, scaler=self.scaler) - - def get_rotation_matrix(self) -> torch.Tensor: - """ - Compute current rotation matrix from Euler angles (differentiable). - - Returns - ------- - torch.Tensor - 3x3 rotation matrix combining initial and perturbation rotations. - """ - - angles = self.initial_rotation + self.rotation - return rotation_matrix_euler_zyz(angles) - - @property - def rotation(self) -> torch.Tensor: - """ - Get current rotation perturbation angles. - - Returns - ------- - torch.Tensor - Current (d_alpha, d_beta, d_gamma) in radians. - """ - return ( - 2 * torch.sigmoid(self.rotation_parameters) - 1 - ) * self.expected_rotational_error - - def get_current_rotation_angles(self) -> torch.Tensor: - """ - Get current rotation angles (initial + perturbation). - - Returns - ------- - torch.Tensor - Current (alpha, beta, gamma) in radians. - """ - angles = self.initial_rotation + self.rotation - return angles - - def get_transformed_xyz(self) -> torch.Tensor: - """ - Transform coordinates - maintains gradient flow. - - Applies rotation around centroid, then translation. - - Returns - ------- - torch.Tensor - Transformed atomic coordinates with shape (n_atoms, 3). - """ - R = self.get_rotation_matrix() - - # Rotate around centroid - xyz_centered = self.xyz_initial - self.centroid - xyz_rotated = xyz_centered @ R.T + self.centroid - - # Apply translation (fractional → Cartesian via cell.fractional_matrix.T; - # see Cell.fractional_to_cartesian for the canonical convention). - t_cart = self.translation_frac @ self.cell.fractional_matrix.T - return xyz_rotated + t_cart - - def get_scale(self) -> float: - """Get current scale factor from scaler.""" - return self.scaler.get_scale() - - def forward(self, debug: bool = False) -> torch.Tensor: - """ - Compute unscaled structure factors using FFT directly. - - Gradient flows: d_alpha/d_beta/d_gamma → R → xyz → density → SF - - Parameters - ---------- - hkl : torch.Tensor - Miller indices with shape (n_reflections, 3). - debug : bool - If True, print gradient tracking info. - - Returns - ------- - torch.Tensor - Unscaled calculated structure factors (scaling done by scaler). - """ - # Get transformed coordinates (has gradient to rotation params) - xyz_transformed = self.get_transformed_xyz() - - hkl = self.data.hkl - - if debug: - print( - f" xyz_transformed.requires_grad: {xyz_transformed.requires_grad}" - ) - print(f" xyz_transformed.grad_fn: {xyz_transformed.grad_fn}") - - # Transform anisotropic atoms if present - xyz_aniso = None - if self.has_aniso: - R = self.get_rotation_matrix() - xyz_aniso_centered = self.xyz_aniso_original - self.centroid - xyz_aniso_rotated = xyz_aniso_centered @ R.T + self.centroid - t_cart = self.translation_frac @ self.cell.fractional_matrix.T - xyz_aniso = xyz_aniso_rotated + t_cart - - # Phase C: optionally perturb per-atom B-factors. Clamp at 0 so the - # density model stays physical even mid-refine; the Gaussian - # restraint on delta_b_iso prevents large excursions. - if self.refine_b: - adp_iso_eff = (self.adp_iso + self.delta_b_iso).clamp(min=0.0) - else: - adp_iso_eff = self.adp_iso - - # Compute structure factors via FFT (bypasses MixedTensor!) - # Note: fractional matrices are now obtained from FFT's internal Cell object - sf, _ = self.fft.compute_structure_factors( - hkl=hkl, - xyz_iso=xyz_transformed, - adp_iso=adp_iso_eff, - occ_iso=self.occ_iso, - A_iso=self.A_iso, - B_iso=self.B_iso, - xyz_aniso=xyz_aniso, - u_aniso=self.u_aniso if self.has_aniso else None, - occ_aniso=self.occ_aniso if self.has_aniso else None, - A_aniso=self.A_aniso if self.has_aniso else None, - B_aniso=self.B_aniso if self.has_aniso else None, - ) - - if debug: - print(f" sf.requires_grad: {sf.requires_grad}") - print(f" sf.grad_fn: {sf.grad_fn}") - - return sf - - def refine( - self, - n_tries: int = 1, - n_iter: int = 100, - ) -> RigidBodyResult: - """ - Run rigid body refinement using least-squares loss with Adam optimizer. - - Uses analytical scale fitting at each step rather than jointly optimizing - the scale parameter with rotation/translation. - - Parameters - ---------- - n_steps : int, optional - Maximum number of optimization steps. Default is 100. - lr : float, optional - Learning rate for Adam optimizer. Default is 0.01. - convergence_threshold : float, optional - Stop if loss change is below this threshold. Default is 1e-6. - print_interval : int, optional - Print progress every N steps. Default is 5. - verbose : bool, optional - Print progress information. Default is True. - loss_type : str, optional - Loss function to use: "ls" for least-squares, "ml" for ML. - Default is "ls". - - Returns - ------- - RigidBodyResult - Refinement results including final parameters and R-factors. - """ - import sys - - if self.verbose > 2: - print( - f" Setting up LBFGS optimizer niter = {n_iter} and max tries = {n_tries}" - ) - sys.stdout.flush() - parameters = [self.rotation_parameters, self.translation_frac, *self.scaler.parameters()] - if self.refine_b: - parameters.append(self.delta_b_iso) - - self.optimizer = torch.optim.LBFGS( - parameters, lr=1, max_iter=100, line_search_fn="strong_wolfe" - ) - - # Phaser-style Gaussian restraints (0 ⇒ disabled). Pre-square once. - sigma_rot_rad = self.sigma_rot_rad - sigma_trans_ang = self.sigma_trans_ang - sigma_b = self.sigma_b - restraints_active = ( - sigma_rot_rad > 0 or sigma_trans_ang > 0 - or (self.refine_b and sigma_b > 0) - ) - - def restraint_loss() -> torch.Tensor: - r = torch.zeros((), dtype=self.translation_frac.dtype, - device=self.translation_frac.device) - if sigma_rot_rad > 0: - r = r + 0.5 * (self.rotation ** 2).sum() / (sigma_rot_rad ** 2) - if sigma_trans_ang > 0: - # Cartesian translation = T_frac @ fractional_matrix.T (Å). - t_cart = self.translation_frac @ self.cell.fractional_matrix.T - r = r + 0.5 * (t_cart ** 2).sum() / (sigma_trans_ang ** 2) - if self.refine_b and sigma_b > 0: - r = r + 0.5 * (self.delta_b_iso ** 2).sum() / (sigma_b ** 2) - return r - - def loss(): - fcalc = self() - ll = self.xray_target(fcalc) - if restraints_active: - ll = ll + restraint_loss() - return ll - - noise = 0 - - def closure(): - self.optimizer.zero_grad() - current_loss = loss() - current_loss.backward() - gradnorm = self.optimizer.param_groups[0]["params"][0].norm().item() - if noise > 0: - self.optimizer.param_groups[0]["params"][0].grad += ( - noise - * torch.randn_like(self.optimizer.param_groups[0]["params"][0].grad) - * gradnorm - ) - return current_loss - - rwork_initial, rfree_initial = self.xray_target.get_rfactor(self()) - - initial_loss = closure().item() - if self.verbose > 0: - print( - f"Initial R-work: {rwork_initial:.4f}, R-free: {rfree_initial:.4f}, ML loss: {initial_loss:.4f}" - ) - - from time import time - - start_time = time() - - tries_needed = 0 - - while True: - tries_needed += 1 - - self.optimizer.step(closure) - with torch.no_grad(): - current_loss = loss().item() - if self.verbose > 1: - print(f"Iter {tries_needed} Current ML loss: {current_loss:.4f}") - final_loss = closure().item() - final_rwork, final_rfree = self.xray_target.get_rfactor(self()) - converged = final_rwork < self.rfactor_converged_threshold - - if converged or tries_needed >= n_tries: - if noise > 0: - noise = 0 - self.optimizer.step(closure) - if self.verbose > 1: - print( - f"Converged at iteration {tries_needed} with R-work: {final_rwork:.4f}" - ) - break - - else: - noise += 1e-2 - self.optimizer = torch.optim.LBFGS( - parameters, lr=1, max_iter=100, line_search_fn="strong_wolfe" - ) - - end_time = time() - - if self.verbose > 0: - print( - f"\nRefinement complete after {tries_needed} steps in {end_time - start_time:.2f} seconds." - ) - print( - f" Final R-work: {final_rwork:.4f} (improved by {rwork_initial - final_rwork:.4f})" - ) - print( - f" Final R-free: {final_rfree:.4f} (improved by {rfree_initial - final_rfree:.4f})" - ) - print( - f" Final rotation angles (deg):", - self.get_current_rotation_angles().rad2deg().detach().cpu().numpy(), - ) - print( - f" Final translation: {self.translation_frac.detach().cpu().numpy()}" - ) - - return RigidBodyResult( - final_rotation=self.get_current_rotation_angles() - .detach() - .cpu() - .numpy() - .tolist(), - final_translation_frac=self.translation_frac.detach() - .cpu() - .numpy() - .tolist(), - initial_r_factor=rwork_initial, - final_r_factor=final_rwork, - final_ml_loss=final_loss, - LBFGS_iterations=self.optimizer.state["n_iter"], - LBFGS_function_evaluations=self.optimizer.state["func_evals"], - n_steps=tries_needed, - converged=converged, - ) - - def get_final_parameters(self) -> dict: - """ - Get final refined parameters. - - Returns - ------- - dict - Dictionary with rotation angles (degrees), translation (fractional), - and scale factor. - """ - angles = self.get_current_rotation_angles().detach().cpu() - rotation_perturbation = self.rotation.detach().cpu() - return { - "alpha_deg": np.degrees(angles[0].item()), - "beta_deg": np.degrees(angles[1].item()), - "gamma_deg": np.degrees(angles[2].item()), - "translation_frac": self.translation_frac.detach().cpu().numpy(), - "scale": self.get_scale(), - "d_alpha_deg": np.degrees(rotation_perturbation[0].item()), - "d_beta_deg": np.degrees(rotation_perturbation[1].item()), - "d_gamma_deg": np.degrees(rotation_perturbation[2].item()), - } diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index ab81b33e..fae0b320 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -11,6 +11,7 @@ import numpy as np import torch +from .distributions import rice_log_likelihood, woolfson_log_likelihood from .e_values import WilsonShellE from dataclasses import dataclass from typing import List, Optional, Tuple @@ -35,90 +36,6 @@ class TranslationPeak: sigma: float -def fft_translation_search( - F_obs: np.ndarray, - F_calc: np.ndarray, - hkl: np.ndarray, - grid_shape: Optional[Tuple[int, int, int]] = None, - n_peaks: int = 10, - cluster_radius: float = 0.05, -) -> Tuple[np.ndarray, np.ndarray, List[TranslationPeak]]: - """ - FFT-based translation search (vectorized). - - The translation function is: - TF(t) = Re{ IFFT{ conj(F_obs) * F_calc } } - - This finds translation t such that F_calc shifted by t best matches F_obs. - - Parameters - ---------- - F_obs : np.ndarray - Observed structure factor amplitudes (or complex), shape (N,). - F_calc : np.ndarray - Calculated structure factors (complex), shape (N,). - hkl : np.ndarray - Miller indices, shape (N, 3). - grid_shape : tuple, optional - (Nx, Ny, Nz) grid size for FFT. If None, auto-computed from HKL range. - n_peaks : int - Number of peaks to return. - cluster_radius : float - Minimum fractional distance between peaks for clustering. - - Returns - ------- - correlation_map : np.ndarray - Full translation function, shape grid_shape. - best_translation : np.ndarray - Best translation in fractional coordinates, shape (3,). - peaks : list - Top peaks as TranslationPeak objects. - - Examples - -------- - :: - - import numpy as np - from torchref.experimental.alignment.translation import fft_translation_search - - # Known translation test - hkl = np.array([[1,0,0], [0,1,0], [1,1,0], [0,0,1]]) - F_obs = np.array([1.0, 1.0, 1.0, 1.0]) - F_calc = np.exp(2j * np.pi * hkl @ [0.25, 0.0, 0.0]) - _, best, peaks = fft_translation_search(F_obs, F_calc, hkl) - print(f'Recovered: {best}') # Should be ~[0.25, 0, 0] - """ - # Auto grid shape from HKL range - if grid_shape is None: - hkl_abs = np.abs(hkl).astype(int) - grid_shape = tuple(2 * (hkl_abs[:, i].max() + 1) for i in range(3)) - - Nx, Ny, Nz = grid_shape - product_grid = np.zeros((Nx, Ny, Nz), dtype=np.complex128) - - # Standard translation function: TF(t) = Re{ IFFT{ conj(F_obs) * F_calc } } - # This finds t such that F_calc(t) = F_calc * exp(2*pi*i*hkl.t) matches F_obs - product = np.conj(F_obs) * F_calc - - # Place at HKL positions using vectorized add.at for accumulation - hkl_int = hkl.astype(int) - h_idx = hkl_int[:, 0] % Nx - k_idx = hkl_int[:, 1] % Ny - l_idx = hkl_int[:, 2] % Nz - - np.add.at(product_grid, (h_idx, k_idx, l_idx), product) - - # IFFT gives correlation at all translations - correlation_map = np.fft.ifftn(product_grid).real - - # Find peaks - peaks = find_translation_peaks(correlation_map, n_peaks, cluster_radius) - best = peaks[0].translation if peaks else np.zeros(3) - - return correlation_map, best, peaks - - def find_translation_peaks( correlation_map: np.ndarray, n_peaks: int = 10, @@ -179,71 +96,6 @@ def find_translation_peaks( return peaks -def fft_translation_search_torch( - F_obs: torch.Tensor, - F_calc: torch.Tensor, - hkl: torch.Tensor, - **kwargs, -) -> Tuple[np.ndarray, np.ndarray, List[TranslationPeak]]: - """ - Torch wrapper for fft_translation_search. - - Parameters - ---------- - F_obs : torch.Tensor - Observed structure factor amplitudes (or complex). - F_calc : torch.Tensor - Calculated structure factors (complex). - hkl : torch.Tensor - Miller indices. - **kwargs - Additional arguments passed to fft_translation_search. - - Returns - ------- - correlation_map : np.ndarray - Full translation function. - best_translation : np.ndarray - Best translation in fractional coordinates. - peaks : list - Top peaks as TranslationPeak objects. - """ - return fft_translation_search( - F_obs.detach().cpu().numpy(), - F_calc.detach().cpu().numpy(), - hkl.detach().cpu().numpy(), - **kwargs, - ) - - -def apply_translation_to_fcalc( - F_calc: np.ndarray, - hkl: np.ndarray, - translation_frac: np.ndarray, -) -> np.ndarray: - """ - Apply translation phase shift to calculated structure factors. - - F(hkl, t) = F(hkl) * exp(2*pi*i * hkl.t) - - Parameters - ---------- - F_calc : np.ndarray - Calculated structure factors (complex), shape (N,). - hkl : np.ndarray - Miller indices, shape (N, 3). - translation_frac : np.ndarray - Translation in fractional coordinates, shape (3,). - - Returns - ------- - F_calc_shifted : np.ndarray - Phase-shifted structure factors, shape (N,). - """ - phase_shift = 2 * np.pi * (hkl @ translation_frac) - return F_calc * np.exp(1j * phase_shift) - - def amplitude_translation_search( F_obs: torch.Tensor, interpolator, @@ -280,8 +132,9 @@ def amplitude_translation_search( ---------- F_obs : torch.Tensor, shape (N,) Observed amplitudes (complex inputs are coerced to |·|). - interpolator : LattmanLoveInterpolator - Provides `evaluate(R, hkl, real_cell, return_amplitude=False)`. + interpolator : object + Anything providing ``evaluate(R, hkl, real_cell, return_amplitude=False)`` + -- in the pipeline, ``align._DirectModelEvaluator``. R_rotation : torch.Tensor, shape (3, 3) Rotation that has been applied to the model coordinates. hkl : torch.Tensor, shape (N, 3) @@ -427,6 +280,58 @@ def amplitude_translation_search( return corr_map_np, best, peaks +def fit_sigma_a_per_shell( + E_obs: torch.Tensor, + E_calc: torch.Tensor, + centric: torch.Tensor, + shell_idx: torch.Tensor, + n_shells: int, + n_grid: int = 81, + interp_var: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Per-shell sigma_A = D, fitted by a grid scan of the shell likelihood. + + Scans ``D`` in ``[0, 0.99]`` and returns the per-shell maximum. At + ``n_grid=81`` the resolution is ~0.012, which is finer than the difference + between adjacent shells on any real falloff. + + This is the *fitted* half of a pair. :func:`torchref.scaling.weighting.empirical_sigma_a` + is the other: it measures model reliability from the ratio of two Wilson + curves and works **before** the molecule is placed, which is what the + rotation function needs. Once a translation exists the residual is + per-reflection rather than per-shell-average, so a direct fit against the + placed model is available and strictly better informed. The two answer the + same question at two different points in the search. + + Returns + ------- + sigma_a : torch.Tensor, shape (n_shells,) + """ + device = E_obs.device + dtype = E_obs.dtype + N = E_obs.numel() + D_grid = torch.linspace(0.0, 0.99, n_grid, device=device, dtype=dtype) # (G,) + F_mean = D_grid.view(-1, 1) * E_calc.view(1, -1) # (G, N) + var_d = (1.0 - D_grid * D_grid).clamp(min=1e-4) # (G,) + if interp_var is None: + var_full = var_d.view(-1, 1).expand(n_grid, N) + else: + var_full = (var_d.view(-1, 1) + interp_var.view(1, -1)).clamp(min=1e-4) + E_obs_full = E_obs.view(1, -1).expand(n_grid, N) + + ll_acent = rice_log_likelihood(E_obs_full, F_mean, var_full) + ll_cent = woolfson_log_likelihood(E_obs_full, F_mean, var_full) + cent_full = centric.view(1, -1) + ll = torch.where(cent_full, ll_cent, ll_acent) # (G, N) + + # Sum per shell, take argmax over the D-grid. + shell_idx_gn = shell_idx.view(1, -1).expand(n_grid, N) + ll_per_shell = torch.zeros((n_grid, n_shells), dtype=dtype, device=device) + ll_per_shell.scatter_add_(1, shell_idx_gn, ll) # (G, n_shells) + best_idx = ll_per_shell.argmax(dim=0) # (n_shells,) + return D_grid[best_idx] + + def llg_translation_rescore( F_obs: torch.Tensor, hkl: torch.Tensor, @@ -477,8 +382,6 @@ def llg_translation_rescore( ------- llg : (K,) torch.Tensor — log-likelihood gain per candidate. """ - from .distributions import rice_log_likelihood, woolfson_log_likelihood - device = G.device real_dtype = torch.float64 complex_dtype = G.dtype @@ -750,328 +653,3 @@ def local_translation_refine( return best_t.cpu(), best_R -def local_rotation_translation_refine( - F_obs: torch.Tensor, - interpolator, - R_initial: torch.Tensor, - t_initial: torch.Tensor, - hkl: torch.Tensor, - spacegroup, - real_cell, - centric: torch.Tensor, - n_shells: int = 15, - rotation_grid_steps: int = 5, - rotation_radius_rad: float = 0.04, - translation_grid_steps: int = 9, - translation_radius_frac: float = 0.02, - batch_size: int = 1024, - verbose: int = 0, -) -> Tuple[torch.Tensor, torch.Tensor, float]: - """ - Joint (R, t) fine-grid refinement scored by the Sim MLRF log-likelihood - gain (LLG) — fully scale-invariant. - - Scale invariance: F_obs is shell-normalised to E-values (Wilson - statistics per resolution shell) and F_calc(R, t) is shell-normalised - per-candidate. The LLG is then a per-shell σA fit + sum of Rice - (acentric) / Woolfson (centric) log-likelihoods — none of which depend - on the absolute magnitude of either F_obs or F_calc. - - Procedure: - 1. Build a `rotation_grid_steps³` cubic grid of small rotation - perturbations in (Δα, Δβ, Δγ) around `R_initial`, parametrised as - axis-angle rotations of magnitude up to `rotation_radius_rad`. - 2. For each R candidate: - - Pre-compute the per-symmetry F_asu via `interpolator.evaluate(R, …)` - (the expensive step — one model.forward per sym op per R). - - Run an analytical inner translation grid of - `translation_grid_steps³` candidates around `t_initial`. Within - this inner loop, only phase factors change (cheap). - - Pick the inner-best t (by an analytical R-factor proxy — fast). - 3. For each (R, best-inner-t) pair, evaluate the **full** Sim MLRF LLG - (per-shell σA fit). Pick the global best by LLG. - - Returns - ------- - R_best : torch.Tensor (3, 3) - Refined rotation = R_initial @ R_perturb_best (column-vector form). - t_best : torch.Tensor (3,) - Refined fractional translation. - llg_best : float - Sim MLRF LLG at the returned (R_best, t_best). - """ - from .ml_rotation import llg_for_rotation_batch - device = getattr(interpolator, "device", hkl.device) - real_dtype = torch.float64 - complex_dtype = torch.complex128 - two_pi_i = 2j * torch.pi - - F_obs_t = F_obs.detach().to(device) - if F_obs_t.is_complex(): - F_obs_t = F_obs_t.abs() - F_obs_t = F_obs_t.to(real_dtype) - F_obs_sum = F_obs_t.sum().clamp(min=1e-30) - - hkl_t = hkl.detach().to(device).to(real_dtype) - R_init = R_initial.detach().to(device).to(real_dtype) - t_init = t_initial.detach().to(device).to(real_dtype) - centric_t = centric.detach().to(device).to(torch.bool) - - # Per-shell normalisation of F_obs → E_obs (Wilson, shell-equal-count). - rec_basis = real_cell.reciprocal_basis_matrix.to(device).to(real_dtype) - s_mag = (hkl_t @ rec_basis).norm(dim=-1) - order = torch.argsort(s_mag) - shell_idx = torch.zeros_like(s_mag, dtype=torch.int64) - chunk = max(1, s_mag.numel() // max(n_shells, 1)) - for k in range(n_shells): - a = k * chunk - b = (k + 1) * chunk if k < n_shells - 1 else s_mag.numel() - shell_idx[order[a:b]] = k - E_obs = WilsonShellE( - F_obs_t, s_mag, shell_idx=shell_idx, n_shells=n_shells, - ).E - - # Rotation perturbation grid. - # We parametrise (Δα, Δβ, Δγ) ∈ [-r, r]³ via the small-angle rotation - # R_perturb ≈ I + ω_x · Lx + ω_y · Ly + ω_z · Lz, exponentiated via - # the matrix exponential of the skew-symmetric generator. For small ω - # (≤ ~3°) Rodrigues is well-conditioned. - def _so3_exp(omega): - # omega: (3,) axis-angle - th = omega.norm() - if th.item() < 1e-12: - return torch.eye(3, dtype=real_dtype, device=device) - axis = omega / th - K = torch.tensor( - [[0.0, -axis[2].item(), axis[1].item()], - [axis[2].item(), 0.0, -axis[0].item()], - [-axis[1].item(), axis[0].item(), 0.0]], - dtype=real_dtype, device=device, - ) - return (torch.eye(3, dtype=real_dtype, device=device) - + torch.sin(th) * K + (1 - torch.cos(th)) * (K @ K)) - - coords_r = torch.linspace(-rotation_radius_rad, rotation_radius_rad, - rotation_grid_steps, dtype=real_dtype, device=device) - omega_grid = torch.stack(torch.meshgrid(coords_r, coords_r, coords_r, - indexing="ij"), dim=-1).reshape(-1, 3) - - # Inner translation grid: (Δtx, Δty, Δtz) ∈ [-rt, rt]³ around t_initial. - coords_t = torch.linspace(-translation_radius_frac, translation_radius_frac, - translation_grid_steps, dtype=real_dtype, device=device) - t_offsets = torch.stack(torch.meshgrid(coords_t, coords_t, coords_t, - indexing="ij"), dim=-1).reshape(-1, 3) - t_candidates = t_init.unsqueeze(0) + t_offsets # (T, 3) - - # For each rotation candidate, pre-compute G_i and run inner t scan. - best_llg = -float("inf") - best_R = R_init.clone() - best_t = t_init.clone() - - for r_idx, omega in enumerate(omega_grid): - R_perturb = _so3_exp(omega) - R_cand = R_init @ R_perturb - # Build G_i for this rotation candidate (the only expensive step). - G, h_R = precompute_G_for_rotation( - interpolator, R_cand, hkl, spacegroup, real_cell, device=device, - ) - - # Inner translation scan: scored by analytical R-factor for speed. - scores = torch.empty(t_candidates.shape[0], dtype=real_dtype, device=device) - for start in range(0, t_candidates.shape[0], batch_size): - stop = min(start + batch_size, t_candidates.shape[0]) - t_batch = t_candidates[start:stop] - dot = torch.einsum("ind,bd->ibn", h_R, t_batch) - phase = torch.exp(two_pi_i * dot.to(complex_dtype)) - F_calc = torch.einsum("in,ibn->bn", G, phase) - F_c_abs = F_calc.abs().to(real_dtype) - num = (F_obs_t.unsqueeze(0) * F_c_abs).sum(dim=-1) - den = (F_c_abs ** 2).sum(dim=-1).clamp(min=1e-30) - k = num / den - R_b = ((F_obs_t.unsqueeze(0) - k.unsqueeze(-1) * F_c_abs).abs() - .sum(dim=-1)) / F_obs_sum - scores[start:stop] = R_b - best_inner = int(scores.argmin().item()) - t_best_inner = t_candidates[best_inner] - - # Score (R_cand, t_best_inner) by full Sim MLRF LLG (scale-invariant - # per-shell σA fit on E-values). - dot = (h_R * t_best_inner.unsqueeze(0).unsqueeze(0)).sum(dim=-1) # (S, N) - phase = torch.exp(two_pi_i * dot.to(complex_dtype)) - F_calc = (G * phase).sum(dim=0) # (N,) - F_calc_abs = F_calc.abs().to(real_dtype).unsqueeze(0) # (1, N) - llg = llg_for_rotation_batch( - F_obs=F_obs_t, shell_idx=shell_idx, n_shells=n_shells, - E_obs=E_obs, centric=centric_t, F_calc=F_calc_abs, - )[0].item() - - if verbose > 1: - print(f" R-refine {r_idx}/{omega_grid.shape[0]}: " - f"|ω|={omega.norm().item():.4f} rad, R={scores[best_inner]:.4f}, " - f"LLG={llg:.2f}", flush=True) - - if llg > best_llg: - best_llg = llg - best_R = R_cand.clone() - best_t = t_best_inner.clone() - return best_R.cpu(), best_t.cpu(), float(best_llg) - - -def patterson_translation_function( - F_obs: torch.Tensor, - interpolator, - R_rotation: torch.Tensor, - hkl: torch.Tensor, - spacegroup, - real_cell, - grid_shape: Optional[Tuple[int, int, int]] = None, - n_peaks: int = 20, - cluster_radius: float = 0.05, -) -> Tuple[np.ndarray, np.ndarray, List[TranslationPeak]]: - """ - Crowther-Blow Patterson translation function for molecular replacement. - - Computes T(t) = Σ_h |F_obs(h)|² · |F_calc(h, t)|² on a fractional grid by - expanding |F_calc(h, t)|² over symmetry-operator pairs and inverse-FFTing - the result. The peaks of T(t) are the translations that best place the - rotated model against the observed amplitudes — unlike the bare - `fft_translation_search`, this is the standard MR translation function and - works on amplitude-only F_obs. - - For each symmetry operator (R_i, t_i) with `x_new = R_i x_old + t_i`, - a per-symmetry asymmetric-unit structure factor is computed as - F_asu_i(h) = interpolator.evaluate(R_rotation, h R_i, real_cell, - return_amplitude=False) - * exp(2πi h · t_i) - (the "h R_i" notation follows from F(h, R x + t) = exp(2πi h·t)·F(R^T h, x); - in tensor form: `hkl @ R_i`). For each ordered pair (i, j) with i ≠ j, the - contribution - |F_obs(h)|² · conj(F_asu_i(h)) · F_asu_j(h) - is scattered into a 3-D reciprocal grid at h' = h @ (R_j − R_i), then the - inverse FFT gives the translation function. Diagonal pairs (i = j) are - t-independent. - - Parameters - ---------- - F_obs : torch.Tensor, shape (N,) - Observed amplitudes. Complex inputs are coerced to |·|. - interpolator : LattmanLoveInterpolator - Provides `evaluate(R, hkl, real_cell, return_amplitude=False)` returning - complex F_calc of the P1 ASU at arbitrary HKL. - R_rotation : torch.Tensor, shape (3, 3) - Rotation that has already been applied to the model coordinates, in - the convention `xyz_new = R · xyz_old`. Passed through to the - interpolator so it evaluates F of the rotated model. - hkl : torch.Tensor, shape (N, 3) - Integer Miller indices of the observed reflections. - spacegroup : SpaceGroup - Provides `matrices` (n_ops, 3, 3, integer in fractional basis) and - `translations` (n_ops, 3, fractional). - real_cell : Cell - Real crystal cell, passed to interpolator.evaluate. - grid_shape : tuple of int, optional - Translation-function grid (Nx, Ny, Nz). Default: 4·max(|hkl|) per axis, - which covers `h @ (R_j − R_i)^T` for any standard spacegroup. - n_peaks : int, default 20 - Number of translation peaks returned. - cluster_radius : float, default 0.05 - Minimum fractional separation between returned peaks. - - Returns - ------- - correlation_map : np.ndarray, shape (Nx, Ny, Nz) - Real-valued translation function T(t). - best_translation : np.ndarray, shape (3,) - Fractional coordinates of the top peak. - peaks : list of TranslationPeak - Top-`n_peaks` peaks sorted by descending T value. - """ - device = getattr(interpolator, "device", hkl.device) - real_dtype = torch.float64 - complex_dtype = torch.complex128 - - F_obs_t = F_obs.detach().to(device) - if F_obs_t.is_complex(): - F_obs_t = F_obs_t.abs() - F_obs_t = F_obs_t.to(real_dtype) - F_obs2 = F_obs_t * F_obs_t # (N,) - - hkl_t = hkl.detach().to(device).to(real_dtype) # (N, 3) - - sym_R = spacegroup.matrices.detach().to(device).to(real_dtype) # (S, 3, 3) - sym_t = spacegroup.translations.detach().to(device).to(real_dtype) # (S, 3) - S = sym_R.shape[0] - - R_rot = R_rotation.detach().to(device).to(real_dtype) - - # Per-symmetry F_asu_i(h) = F_model_rot(h R_i) · exp(2πi h · t_i) - F_asu = torch.zeros((S, hkl_t.shape[0]), dtype=complex_dtype, device=device) - two_pi_i = 2j * torch.pi - for i in range(S): - hkl_i = hkl_t @ sym_R[i] - F_i = interpolator.evaluate(R_rot, hkl_i, real_cell, return_amplitude=False) - F_i = F_i.to(complex_dtype) - phase = torch.exp(two_pi_i * (hkl_t @ sym_t[i])).to(complex_dtype) - F_asu[i] = F_i * phase - - # Translation grid extent: bound by max |h @ (R_j − R_i)^T|. For - # crystallographic R_op (entries in {-1, 0, 1}, occasionally 2 for trigonal - # subgroups), |R_j − R_i| has entries up to 2 → 2·max|h| per axis is the - # natural extent. Use 4·max(|hkl|) for an oversampled, periodic grid. - if grid_shape is None: - max_h = hkl_t.abs().max(dim=0).values - grid_shape = tuple(int(4 * (m.item() + 1)) for m in max_h) - Nx, Ny, Nz = grid_shape - - W = torch.zeros((Nx, Ny, Nz), dtype=complex_dtype, device=device) - W_flat = W.view(-1) - # Cross-pair accumulation (skip i == j: t-independent, only shifts DC). - for i in range(S): - Fi_conj = torch.conj(F_asu[i]) - for j in range(S): - if j == i: - continue - diff_R = sym_R[j] - sym_R[i] - h_diff = (hkl_t @ diff_R).round().to(torch.int64) - ix = h_diff[:, 0] % Nx - iy = h_diff[:, 1] % Ny - iz = h_diff[:, 2] % Nz - flat_idx = ix * (Ny * Nz) + iy * Nz + iz - weight = F_obs2 * Fi_conj * F_asu[j] - W_flat.index_add_(0, flat_idx, weight) - - TF_complex = torch.fft.ifftn(W, dim=(0, 1, 2)) - TF = TF_complex.real - - TF_np = TF.detach().cpu().numpy().astype(np.float32) - peaks = find_translation_peaks(TF_np, n_peaks=n_peaks, cluster_radius=cluster_radius) - best = peaks[0].translation if peaks else np.zeros(3) - return TF_np, best, peaks - - -def apply_translation_to_fcalc_torch( - F_calc: torch.Tensor, - hkl: torch.Tensor, - translation_frac: torch.Tensor, -) -> torch.Tensor: - """ - Apply translation phase shift to calculated structure factors (PyTorch). - - F(hkl, t) = F(hkl) * exp(2*pi*i * hkl.t) - - Parameters - ---------- - F_calc : torch.Tensor - Calculated structure factors (complex). - hkl : torch.Tensor - Miller indices. - translation_frac : torch.Tensor - Translation in fractional coordinates. - - Returns - ------- - F_calc_shifted : torch.Tensor - Phase-shifted structure factors. - """ - phase_shift = 2 * torch.pi * (hkl.to(translation_frac.dtype) @ translation_frac) - return F_calc * torch.exp(1j * phase_shift) diff --git a/torchref/refinement/targets/xray/rice.py b/torchref/refinement/targets/xray/rice.py index 728e881c..a5599e3a 100644 --- a/torchref/refinement/targets/xray/rice.py +++ b/torchref/refinement/targets/xray/rice.py @@ -22,22 +22,20 @@ class RiceXrayTarget(XrayTarget): empirically it was the worst-behaved target measured, destroying geometry (bond RMSZ 28.0 where every other target sat near 1.3). See ``_specs.py``'s module docstring. - This class survives for exactly one caller: - :class:`torchref.experimental.alignment.rigid_body.RigidBodyRefinement`, the FFT-direct - rigid-body aligner in the molecular-replacement pipeline, which constructs it directly - rather than through the factory. - - **Why it was not simply repointed at ``nll``** during the 2026-08 target refactor, as - originally planned: that aligner has **no test coverage whatsoever** - (``tests/integration/test_rigid_body_refinement.py`` exercises the *other* rigid-body - module, ``refinement/rigid_body_refinement.py``). Swapping a Rice likelihood for a - Gaussian there would be an untested numerical change in a live MR path, so the - likelihood was kept and only its *implementation* was de-duplicated -- the body now - calls the shared :func:`~torchref.base.targets.xray_likelihoods.rice_math` primitive - instead of a second copy of the Rice in the deleted ``xray_ml`` module. - - Whoever gives that aligner a test should revisit this: ``nll`` or ``ml_noalpha`` is - almost certainly the better objective, and then this class can go. + **It now has no caller at all, and is a deletion candidate.** It survived the + 2026-08 target refactor for exactly one: the FFT-direct rigid-body aligner in the + molecular-replacement pipeline, which constructed it directly rather than through + the factory. That aligner had no test coverage + (``tests/integration/test_rigid_body_refinement.py`` exercises the *other* + rigid-body module, ``refinement/rigid_body_refinement.py``), so repointing it at + ``nll`` would have been an untested numerical change in a live path -- the + likelihood was kept and only its *implementation* de-duplicated onto the shared + :func:`~torchref.base.targets.xray_likelihoods.rice_math` primitive. + + The MR pipeline no longer polishes placements, so that aligner is gone and the + constraint with it. What remains is this class, three unit tests of it, and an + export. Removing all of that is a ``refinement/`` change and belongs in a + ``refinement/`` commit, not an alignment one. """ #: ``epsilon * beta`` was clamped here in the implementation this replaced. Preserved From 8cf986fcaad9e5802a575578d3b8608111def261 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 31 Aug 2026 14:22:00 +0200 Subject: [PATCH 109/250] Add the Cholesky transform at arbitrary size A disorder-field node that carries the covariance of its displacement modes needs a q x q factor, not the 3 x 3 one the per-atom ADP wrapper uses. Same contract as the unrolled pair beside it: exp(x) + epsilon on the diagonal so the reconstruction is positive-definite for any parameter value and epsilon floors the smallest eigenvalue, an eigenvalue clamp before factorising on the way back, and a CPU-forced eigh because cuSolver's batched kernels fail on the degenerate batches a near-isotropic model produces. The 3 x 3 pair stays as it is. It runs in a forward pass, where the unrolled form is worth keeping, and raw_to_cholesky(raw, 3, eps) is asserted to reproduce it. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- torchref/model/parameter_wrappers.py | 60 ++++++++++++++++++++++++++++ 1 file changed, 60 insertions(+) diff --git a/torchref/model/parameter_wrappers.py b/torchref/model/parameter_wrappers.py index e6f66d05..47095342 100644 --- a/torchref/model/parameter_wrappers.py +++ b/torchref/model/parameter_wrappers.py @@ -1043,6 +1043,66 @@ def u6_to_raw6(U: torch.Tensor, epsilon: float) -> torch.Tensor: return torch.where(finite.unsqueeze(-1), raw, torch.full_like(raw, float("nan"))) +# ---------------------------------------------------------------------------------- +# The same transform at arbitrary size, for a covariance that is not a 3x3 U tensor. +# A node of the disorder field carries the covariance of its displacement modes, which +# is q x q for q modes; the pair above is the q = 3 case with the indexing unrolled. +# Kept as the general form rather than replacing the unrolled pair, which is a forward +# hot path. +# ---------------------------------------------------------------------------------- + + +def chol_param_count(q: int) -> int: + """Free parameters in a ``q x q`` lower-triangular factor.""" + return q * (q + 1) // 2 + + +def raw_to_cholesky(raw: torch.Tensor, q: int, epsilon: float) -> torch.Tensor: + """Free parameters to a lower-triangular ``(..., q, q)`` factor. + + Layout is ``[log diagonal (q) | strict lower triangle (q(q-1)/2), row major]``, and + the diagonal is ``exp(x) + epsilon``, so ``L L^T`` is positive-definite for any + input and ``epsilon`` bounds its smallest eigenvalue from below. At ``q = 3`` this + is the same layout and the same convention as :func:`raw6_to_u6`. + + No factorisation happens here, which is what makes it safe in a forward pass. + """ + rows, cols = torch.tril_indices(q, q, offset=-1, device=raw.device) + L = raw.new_zeros(*raw.shape[:-1], q, q) + diag = torch.exp(raw[..., :q]) + epsilon + idx = torch.arange(q, device=raw.device) + L[..., idx, idx] = diag + if rows.numel(): + L[..., rows, cols] = raw[..., q:] + return L + + +def psd_to_raw(M: torch.Tensor, epsilon: float) -> torch.Tensor: + """Symmetric ``(..., q, q)`` matrix to Cholesky free parameters, projecting onto PSD. + + The inverse of :func:`raw_to_cholesky`, with the same eigenvalue clamp and the same + CPU-forced ``eigh`` as :func:`u6_to_raw6`: a least-squares or seeded covariance need + not be positive-definite, and cuSolver's batched kernels fail on the degenerate + batches a near-isotropic model produces. Runs at construction, never in a forward + pass. + """ + q = M.shape[-1] + src_device = M.device + M = 0.5 * (M + M.transpose(-1, -2)) + M = M.cpu() + w, V = torch.linalg.eigh(M) + w = w.clamp(min=epsilon * epsilon) + M = (V * w.unsqueeze(-2)) @ V.transpose(-1, -2) + L = torch.linalg.cholesky(M) + idx = torch.arange(q) + rows, cols = torch.tril_indices(q, q, offset=-1) + raw_diag = torch.log((L[..., idx, idx] - epsilon).clamp(min=1e-12)) + parts = [raw_diag] + if rows.numel(): + parts.append(L[..., rows, cols]) + return torch.cat(parts, dim=-1).to(src_device) + + class CholeskyMixedTensor(MixedTensor): """A MixedTensor for anisotropic ADPs (U tensors) kept positive-definite. From 6ac873ada6d746364ab242fc2aaeab1fc1a2cfaf Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 31 Aug 2026 14:22:00 +0200 Subject: [PATCH 110/250] Give the disorder field a displacement-mode payload A node stored one ADP, so the only way to express U varying through space was to add nodes: spatial detail cost nodes, and every added node brought another kernel that could narrow onto a single atom. Store the covariance of the node's displacement modes instead. With modes Psi(r) and coefficient covariance Sigma = L L^T, an atom at displacement r from the node receives U(r) = Psi(r) Sigma Psi(r)^T, so one node's U already varies across its whole region and cannot spike. Positive-semidefiniteness is free at every r, because U = (Psi L)(Psi L)^T. An arbitrary polynomial in r carries no such guarantee and goes indefinite at the edge of the region, which is exactly where the softmax weights have not yet decayed. The mode set is the expressiveness knob. Three translations and three rotations reproduce the textbook TLS expression T + AS + S^T A^T + A L A^T identically, with tr S appearing as its one flat direction; releasing the antisymmetry of the gradient adds domains that breathe and shear as well as rotate. constant q=3 6 payload one U per node, as AnisotropicPayload rigid q=6 21 payload TLS rigid_dilation q=7 28 payload TLS plus uniform breathing affine q=12 78 payload full linear displacement field Entering the mode seeds only the translation block from the constant-U solve and leaves the gradient modes at the Cholesky floor, so a freshly installed field starts at the constant-U field's own R-factor and refinement can only move away from a known state. Measured cycle-0 R-free spread across all four rungs: 0.00004. NodePayload.fit gains the geometric context contributions already took, because an r-dependent payload cannot build its modes without it. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- tests/unit/model/test_adp_field_mode_sets.py | 139 ++++++++ .../model/test_mode_covariance_payload.py | 336 ++++++++++++++++++ torchref/model/disorder_field.py | 265 +++++++++++++- torchref/model/model.py | 44 ++- 4 files changed, 770 insertions(+), 14 deletions(-) create mode 100644 tests/unit/model/test_adp_field_mode_sets.py create mode 100644 tests/unit/model/test_mode_covariance_payload.py diff --git a/tests/unit/model/test_adp_field_mode_sets.py b/tests/unit/model/test_adp_field_mode_sets.py new file mode 100644 index 00000000..0f33465b --- /dev/null +++ b/tests/unit/model/test_adp_field_mode_sets.py @@ -0,0 +1,139 @@ +"""``set_adp_mode("field_aniso", mode_set=...)``: installing a displacement-mode field. + +The wiring, not the payload arithmetic --- that is +``test_mode_covariance_payload.py``. What matters here is that a mode set reaches the +model through the existing switch, lands in the ``u`` slot like any anisotropic field, +and starts where the constant-U field starts, so entering the parametrisation is not +itself a change to the model. +""" + +import pytest +import torch + +from torchref.model.disorder_field import ( + MODE_SETS, + DisorderFieldTensor, + ModeCovariancePayload, +) +from torchref.model.model import Model + + +@pytest.fixture(scope="module") +def pdb_path(pdb_dir): + return str(pdb_dir / "3GR5.pdb") + + +def _field_model(pdb_path, mode_set=None, n_nodes=8): + model = Model(verbose=0) + model.load_pdb(pdb_path) + model.set_adp_mode( + "field_aniso", n_nodes=n_nodes, k_neighbors=n_nodes, mode_set=mode_set + ) + return model + + +@pytest.mark.unit +@pytest.mark.parametrize("mode_set", sorted(MODE_SETS)) +def test_mode_set_installs_into_the_u_slot(pdb_path, mode_set): + """A mode field is an anisotropic field: same slot, same downstream consumers.""" + model = _field_model(pdb_path, mode_set) + assert model.adp_is_field + assert isinstance(model.u, DisorderFieldTensor) + assert isinstance(model.u.payload, ModeCovariancePayload) + assert model.u.payload.mode_set == mode_set + assert bool(model.aniso_flag.all()) + # The per-atom surface the structure-factor path uses. + u6 = model.adp_u6() + assert u6.shape == (len(model.pdb), 6) + assert torch.isfinite(u6).all() + + +@pytest.mark.unit +@pytest.mark.parametrize( + "mode_set,per_node", [("constant", 10), ("rigid", 25), ("rigid_dilation", 32), ("affine", 82)] +) +def test_parameter_count_is_payload_plus_sigma_plus_offset(pdb_path, mode_set, per_node): + """Storage is [payload | log sigma | 3 offset], so the ladder costs 10/25/32/82.""" + model = _field_model(pdb_path, mode_set, n_nodes=8) + assert model.u.refinable_params.numel() == 8 * per_node + + +@pytest.mark.unit +@pytest.mark.parametrize("mode_set", sorted(MODE_SETS)) +def test_entering_the_mode_starts_at_the_constant_u_field(pdb_path, mode_set): + """Every rung seeds its translation block from the same solve and floors the rest. + + So a freshly installed mode field must give essentially the constant-U field's ADPs. + That is what makes entering this parametrisation safe: the starting R-factor is one + already known, and refinement can only move away from it. + """ + base = _field_model(pdb_path, "constant").adp_u6().detach() + got = _field_model(pdb_path, mode_set).adp_u6().detach() + # Gradient modes start at the Cholesky floor, epsilon^2, which is ~1e-6 A^2 against + # a U of order 0.2 -- present but far below anything observable. + assert torch.allclose(got, base, atol=1e-5) + + +@pytest.mark.unit +def test_positive_definite_per_atom(pdb_path): + """Every atom's U must be PD or the anisotropic B-matrix inverse blows up.""" + model = _field_model(pdb_path, "affine", n_nodes=6) + u6 = model.adp_u6().detach() + M = torch.zeros(u6.shape[0], 3, 3, dtype=u6.dtype) + M[:, 0, 0], M[:, 1, 1], M[:, 2, 2] = u6[:, 0], u6[:, 1], u6[:, 2] + M[:, 0, 1] = M[:, 1, 0] = u6[:, 3] + M[:, 0, 2] = M[:, 2, 0] = u6[:, 4] + M[:, 1, 2] = M[:, 2, 1] = u6[:, 5] + assert float(torch.linalg.eigvalsh(M).min()) > 0.0 + + +@pytest.mark.unit +def test_gradient_reaches_the_node_parameters(pdb_path): + """Through the zero-argument forward, which is the path refinement actually uses.""" + model = _field_model(pdb_path, "rigid") + model.adp_u6().sum().backward() + grad = model.u.refinable_params.grad + assert grad is not None and float(grad.abs().sum()) > 0 + + +@pytest.mark.unit +def test_copy_round_trips_a_mode_field(pdb_path): + """``Model.copy`` shares no storage but must keep the payload and the accessor.""" + model = _field_model(pdb_path, "rigid_dilation") + clone = model.copy() + assert isinstance(clone.u.payload, ModeCovariancePayload) + assert clone.u.payload.mode_set == "rigid_dilation" + assert torch.allclose(clone.adp_u6().detach(), model.adp_u6().detach()) + assert clone.u.refinable_params is not model.u.refinable_params + # The accessor must point at the COPY's coordinates, not the original's. The + # perturbation has to be non-rigid: a field whose nodes are atom centroids is + # translation-invariant by construction, so shifting every atom would change + # nothing and prove nothing. + original = model.adp_u6().detach().clone() + with torch.no_grad(): + clone.xyz.refinable_params[: len(clone.pdb) // 2] += 3.0 + clone.u.reset_forward_cache() + assert not torch.allclose(clone.adp_u6().detach(), original) + # ...and the original must be untouched by it. + model.u.reset_forward_cache() + assert torch.allclose(model.adp_u6().detach(), original) + + +@pytest.mark.unit +def test_mode_set_is_rejected_on_the_isotropic_field(pdb_path): + """There is no isotropic form of a displacement-mode covariance.""" + model = Model(verbose=0) + model.load_pdb(pdb_path) + with pytest.raises(ValueError, match="no isotropic form"): + model.set_adp_mode("field", n_nodes=8, mode_set="rigid") + + +@pytest.mark.unit +def test_leaving_field_mode_materialises_per_atom(pdb_path): + """The conversion out reads the per-atom U, so it works for any payload.""" + model = _field_model(pdb_path, "affine", n_nodes=6) + before = model.adp_u6().detach().clone() + model.set_adp_mode("anisotropic") + assert not model.adp_is_field + kept = model.aniso_flag + assert torch.allclose(model.adp_u6().detach()[kept], before[kept], atol=1e-6) diff --git a/tests/unit/model/test_mode_covariance_payload.py b/tests/unit/model/test_mode_covariance_payload.py new file mode 100644 index 00000000..db133145 --- /dev/null +++ b/tests/unit/model/test_mode_covariance_payload.py @@ -0,0 +1,336 @@ +"""The mode-covariance node payload, of which TLS is one member. + +A node stores the covariance of the displacement modes it represents rather than an ADP, +so the ADP an atom receives depends on where it sits inside the node's region. Two +properties have to hold before any of it is worth measuring: the rigid mode set must +reproduce the textbook TLS expression exactly, and every mode set must stay +positive-semidefinite at every displacement, because an indefinite U makes the +structure-factor FFT return NaN. +""" + +import math + +import pytest +import torch + +from torchref.model.disorder_field import ( + MODE_SETS, + AnisotropicPayload, + ModeCovariancePayload, +) +from torchref.model.parameter_wrappers import ( + chol_param_count, + psd_to_raw, + raw6_to_u6, + raw_to_cholesky, + u6_to_matrix, +) + +DTYPE = torch.float64 + + +def _u6_to_mat(u6): + return u6_to_matrix(u6) + + +def _random_sigma(q, k=1, scale=0.05, seed=0): + """A random PD ``(k, q, q)`` covariance.""" + g = torch.Generator().manual_seed(seed) + A = torch.randn(k, q, q, generator=g, dtype=DTYPE) * scale + return A @ A.transpose(-1, -2) + 1e-3 * torch.eye(q, dtype=DTYPE) + + +# ---------------------------------------------------------------------------------- +# Sizes. +# ---------------------------------------------------------------------------------- + + +@pytest.mark.unit +@pytest.mark.parametrize( + "mode_set,q,width", + [("constant", 3, 6), ("rigid", 6, 21), ("rigid_dilation", 7, 28), ("affine", 12, 78)], +) +def test_mode_set_sizes(mode_set, q, width): + """The ladder is 6 / 21 / 28 / 78 parameters; TLS is the 21 (20 determinable).""" + p = ModeCovariancePayload(mode_set) + assert p.q == q + assert p.width == width == chol_param_count(q) + assert p.out_width == 6 + + +@pytest.mark.unit +def test_unknown_mode_set_is_rejected(): + with pytest.raises(ValueError, match="Unknown mode set"): + ModeCovariancePayload("librational_whimsy") + + +@pytest.mark.unit +def test_mode_sets_are_nested(): + """Each rung must contain the one below, or the ladder is not a ladder.""" + keys = ["constant", "rigid", "rigid_dilation", "affine"] + for lo, hi in zip(keys, keys[1:]): + assert set(MODE_SETS[lo]).issubset(set(MODE_SETS[hi])) + assert ModeCovariancePayload(lo).q < ModeCovariancePayload(hi).q + + +# ---------------------------------------------------------------------------------- +# The TLS identity. If this fails nothing downstream is trustworthy. +# ---------------------------------------------------------------------------------- + + +@pytest.mark.unit +def test_rigid_mode_set_is_textbook_tls(): + """``Psi Sigma Psi^T`` with translations and rotations IS ``T + AS + S^T A^T + A L A^T``. + + ``A`` is the matrix whose columns are ``e_i x r``, so ``A lambda = lambda x r``. With + ``Sigma`` blocked as ``[[T, S^T], [S, L]]`` over ``c = (t, lambda)`` the expansion is + the classical TLS expression, which is the claim that makes this payload a + generalisation of TLS rather than something merely similar. + """ + payload = ModeCovariancePayload("rigid") + sigma = _random_sigma(6, k=1, seed=3)[0] + T, L, S = sigma[:3, :3], sigma[3:, 3:], sigma[3:, :3] + + torch.manual_seed(11) + for r in torch.randn(20, 3, dtype=DTYPE) * 5.0: + # Columns of A are e_i x r. + A = torch.stack( + [ + torch.tensor([0.0, -r[2], r[1]], dtype=DTYPE), + torch.tensor([r[2], 0.0, -r[0]], dtype=DTYPE), + torch.tensor([-r[1], r[0], 0.0], dtype=DTYPE), + ], + dim=1, + ) + expected = T + A @ S + (A @ S).T + A @ L @ A.T + + Psi = payload.modes(r) # (3, 6) + got = Psi @ sigma @ Psi.T + assert torch.allclose(got, expected, atol=1e-10), f"r={r.tolist()}" + + +@pytest.mark.unit +def test_rotation_generators_are_antisymmetric(): + """``A`` must be antisymmetric, which is what makes the rigid set a rigid motion.""" + payload = ModeCovariancePayload("rigid") + r = torch.tensor([1.3, -2.1, 0.7], dtype=DTYPE) + A = payload.modes(r)[:, 3:] + assert torch.allclose(A, -A.T, atol=1e-12) + + +@pytest.mark.unit +def test_trace_s_is_the_flat_direction(): + """TLS has exactly one unobservable combination: adding to ``tr S`` must be free. + + Shifting ``S -> S + c I`` leaves ``U(r)`` unchanged for every ``r``, because + ``A(cI) + (A cI)^T = c(A + A^T) = 0`` for antisymmetric ``A``. + """ + payload = ModeCovariancePayload("rigid") + sigma = _random_sigma(6, k=1, seed=5)[0] + shifted = sigma.clone() + shifted[3:, :3] += 0.01 * torch.eye(3, dtype=DTYPE) + shifted[:3, 3:] += 0.01 * torch.eye(3, dtype=DTYPE) + + torch.manual_seed(2) + for r in torch.randn(10, 3, dtype=DTYPE) * 4.0: + Psi = payload.modes(r) + assert torch.allclose(Psi @ sigma @ Psi.T, Psi @ shifted @ Psi.T, atol=1e-12) + + +# ---------------------------------------------------------------------------------- +# Positive-semidefiniteness, the property the whole construction exists for. +# ---------------------------------------------------------------------------------- + + +@pytest.mark.unit +@pytest.mark.parametrize("mode_set", sorted(MODE_SETS)) +def test_psd_at_extreme_displacement(mode_set): + """PSD for any ``Sigma`` and any ``r``, including far outside the node's region. + + This is what an arbitrary polynomial in ``r`` cannot promise: it goes indefinite + somewhere, and somewhere is the edge of the region where the weights have not yet + decayed. + """ + payload = ModeCovariancePayload(mode_set) + torch.manual_seed(7) + raw = torch.randn(4, payload.width, dtype=DTYPE) * 3.0 + L = raw_to_cholesky(raw, payload.q, payload.epsilon) + sigma = L @ L.transpose(-1, -2) + + for scale in (0.0, 1e-3, 1.0, 50.0, 1000.0): + r = torch.randn(6, 3, dtype=DTYPE) * scale + Psi = payload.modes(r) # (6, 3, q) + U = Psi @ sigma[:, None] @ Psi.transpose(-1, -2)[None] # (4, 6, 3, 3) + ev = torch.linalg.eigvalsh(U) + assert float(ev.min()) >= -1e-9, f"{mode_set} indefinite at |r|~{scale}" + + +@pytest.mark.unit +@pytest.mark.parametrize("scale", [0.1, 1.0, 3.0, 6.0]) +def test_sigma_is_pd_for_any_parameter_value(scale): + """Cholesky storage: no parameter value can make the covariance indefinite. + + Judged against the matrix norm, not against zero. The diagonal is ``exp(x)``, so a + wide spread of parameters gives ``Sigma`` a huge dynamic range and ``eigvalsh`` + returns the small eigenvalues with an error set by the large ones -- a float64 + property of the eigensolver, not of the parametrisation. + """ + payload = ModeCovariancePayload("affine") + torch.manual_seed(1) + raw = torch.randn(8, payload.width, dtype=DTYPE) * scale + ev = torch.linalg.eigvalsh(payload.sigma(raw)) + tol = 1e-10 * ev.abs().max(dim=-1, keepdim=True).values + assert bool((ev > -tol).all()), f"min eigenvalue {float(ev.min()):.3e} at scale {scale}" + + +# ---------------------------------------------------------------------------------- +# Agreement with the payload it generalises. +# ---------------------------------------------------------------------------------- + + +@pytest.mark.unit +def test_constant_mode_set_matches_anisotropic_payload(): + """``"constant"`` is the same model as :class:`AnisotropicPayload`. + + Both store a 3x3 Cholesky factor and hand every atom the same U, so given the same + raw parameters they must produce the same tensor. That makes the constant rung the + inertness guard: a change here that moved it would be a change to the existing + anisotropic field. + """ + eps = 1e-3 + mode = ModeCovariancePayload("constant", epsilon=eps) + aniso = AnisotropicPayload(epsilon=eps) + torch.manual_seed(4) + raw = torch.randn(5, 6, dtype=DTYPE) + + xyz = torch.randn(9, 3, dtype=DTYPE) * 10.0 + node_pos = torch.randn(5, 3, dtype=DTYPE) * 10.0 + nl = torch.randint(0, 5, (9, 3)) + + got = mode.contributions(raw, xyz, node_pos, nl) + expected = aniso.contributions(raw, xyz, node_pos, nl) + assert got.shape == expected.shape == (9, 3, 6) + assert torch.allclose(got, expected, atol=1e-12) + + +@pytest.mark.unit +def test_raw_to_cholesky_matches_the_unrolled_three_by_three(): + """The general helper must agree with the unrolled 3x3 pair it generalises.""" + eps = 1e-3 + torch.manual_seed(6) + raw = torch.randn(12, 6, dtype=DTYPE) + L = raw_to_cholesky(raw, 3, eps) + got = L @ L.transpose(-1, -2) + expected = _u6_to_mat(raw6_to_u6(raw, eps)) + assert torch.allclose(got, expected, atol=1e-12) + + +@pytest.mark.unit +def test_cholesky_round_trip(): + """``psd_to_raw`` inverts ``raw_to_cholesky`` for a PD matrix.""" + eps = 1e-4 + for q in (3, 6, 12): + sigma = _random_sigma(q, k=5, scale=0.3, seed=q) + raw = psd_to_raw(sigma, eps) + L = raw_to_cholesky(raw, q, eps) + assert torch.allclose(L @ L.transpose(-1, -2), sigma, atol=1e-8), f"q={q}" + + +# ---------------------------------------------------------------------------------- +# Fit and magnitude. +# ---------------------------------------------------------------------------------- + + +@pytest.mark.unit +@pytest.mark.parametrize("mode_set", sorted(MODE_SETS)) +def test_fit_seeds_translations_and_floors_the_rest(mode_set): + """Entering the parametrisation must start as the equivalent constant-U field.""" + payload = ModeCovariancePayload(mode_set, epsilon=1e-3) + torch.manual_seed(8) + n_atoms, k_nodes = 40, 4 + xyz = torch.randn(n_atoms, 3, dtype=DTYPE) * 8.0 + node_pos = torch.randn(k_nodes, 3, dtype=DTYPE) * 8.0 + nl = torch.randint(0, k_nodes, (n_atoms, 2)) + w_dense = torch.rand(n_atoms, k_nodes, dtype=DTYPE) + w_dense = w_dense / w_dense.sum(dim=1, keepdim=True) + target_b = 10.0 + 20.0 * torch.rand(n_atoms, dtype=DTYPE) + + raw = payload.fit(target_b, w_dense, 1e-3, xyz, node_pos, nl) + assert raw.shape == (k_nodes, payload.width) + sigma = payload.sigma(raw) + + # Translation block carries the fit; every gradient mode sits at the floor. + assert float(sigma[:, :3, :3].diagonal(dim1=-2, dim2=-1).min()) > 1e-4 + if payload.q > 3: + grad = sigma[:, 3:, 3:] + assert float(grad.abs().max()) < 1e-5, "gradient modes did not start at the floor" + + +@pytest.mark.unit +def test_log_magnitude_is_b_eq_of_the_translation_block(): + payload = ModeCovariancePayload("affine") + torch.manual_seed(9) + raw = torch.randn(6, payload.width, dtype=DTYPE) * 0.5 + T = payload.sigma(raw)[:, :3, :3] + expected = torch.log( + ((8.0 * math.pi**2 / 3.0) * T.diagonal(dim1=-2, dim2=-1).sum(-1)).clamp(min=1e-6) + ) + assert torch.allclose(payload.log_magnitude(raw), expected, atol=1e-12) + + +# ---------------------------------------------------------------------------------- +# Gradients. +# ---------------------------------------------------------------------------------- + + +@pytest.mark.unit +@pytest.mark.parametrize("mode_set", ["rigid", "affine"]) +def test_gradcheck_contributions(mode_set): + """Gradient w.r.t. the node parameters and the coordinates. + + ``node_pos`` is held constant here rather than differentiated: the conditioning + length scale is a detached median over node positions, so gradcheck's numerical + derivative would pick up a path the analytic one deliberately cuts. That cut is the + subject of :func:`test_length_scale_carries_no_gradient`; the gradient that reaches + a node's position through the displacement ``r`` is covered here by ``xyz``, which + enters the same subtraction with the opposite sign. + """ + payload = ModeCovariancePayload(mode_set) + torch.manual_seed(10) + xyz = (torch.randn(7, 3, dtype=DTYPE) * 4.0).requires_grad_(True) + node_pos = torch.randn(3, 3, dtype=DTYPE) * 4.0 + nl = torch.randint(0, 3, (7, 2)) + raw = (torch.randn(3, payload.width, dtype=DTYPE) * 0.3).requires_grad_(True) + + assert torch.autograd.gradcheck( + lambda p_, x: payload.contributions(p_, x, node_pos, nl), + (raw, xyz), + eps=1e-6, + atol=1e-7, + ) + + +@pytest.mark.unit +def test_length_scale_carries_no_gradient(): + """The conditioning length scale is a stop-gradient, on purpose. + + It is the median nearest-neighbour node distance, whose true derivative is supported + on whichever single node pair happens to be at the median -- an artifact of the + layout, not a direction any optimiser should follow. Cutting it also stops the + optimiser rescaling its own modes by spreading the nodes out. Node position still + gets its real gradient through the displacement ``r``. + """ + payload = ModeCovariancePayload("affine") + torch.manual_seed(12) + xyz = torch.randn(20, 3, dtype=DTYPE) * 5.0 + node_pos = (torch.randn(4, 3, dtype=DTYPE) * 5.0).requires_grad_(True) + nl = torch.randint(0, 4, (20, 3)) + raw = torch.randn(4, payload.width, dtype=DTYPE) * 0.3 + + payload.contributions(raw, xyz, node_pos, nl).sum().backward() + assert node_pos.grad is not None + assert float(node_pos.grad.abs().sum()) > 0, "no gradient reaches node positions at all" + + # The scale itself must be a plain float, not something carrying a graph. + lam = payload._length_scale(node_pos) + assert isinstance(lam, float) diff --git a/torchref/model/disorder_field.py b/torchref/model/disorder_field.py index 97b94d28..3cb8124a 100644 --- a/torchref/model/disorder_field.py +++ b/torchref/model/disorder_field.py @@ -11,10 +11,10 @@ :meth:`~DisorderFieldTensor.update_refinable_mask` are in ATOM space, while ``refinable_mask`` and :meth:`get_refinable_count` are in NODE space. -A node's position is *derived*, not refined: it is the centroid of the atoms within -``anchor_radius`` bonds of its anchor atom. That keeps a node inside the molecule, -confined to one connected fragment, and moving with the model, and it leaves the -optimiser no free coordinate to wander with. +A node's position is anchored, not free: it is the centroid of the atoms in its anchor +cluster, plus an optional refinable offset. Anchoring keeps a node inside the molecule +and moving with the model, so the offset says only where it sits *relative* to the atoms +it serves and cannot wander off into solvent. """ import math @@ -24,7 +24,15 @@ from torch import nn from torchref.config import get_float_dtype, normalize_device -from torchref.model.parameter_wrappers import MixedTensor, raw6_to_u6, u6_to_raw6 +from torchref.model.parameter_wrappers import ( + MixedTensor, + chol_param_count, + psd_to_raw, + raw6_to_u6, + raw_to_cholesky, + u6_to_matrix, + u6_to_raw6, +) from torchref.utils.utils import ModuleReference __all__ = [ @@ -32,6 +40,8 @@ "NodePayload", "IsotropicPayload", "AnisotropicPayload", + "ModeCovariancePayload", + "MODE_SETS", "farthest_point_anchors", "density_anchor_rows", "build_neighbor_list", @@ -204,8 +214,14 @@ def contributions(self, payload, xyz, node_pos, neighbor_list): """ raise NotImplementedError - def fit(self, target, w_dense, epsilon): - """``(K, width)`` least-squares payload reproducing ``target``.""" + def fit(self, target, w_dense, epsilon, xyz, node_pos, neighbor_list): + """``(K, width)`` payload whose field reproduces ``target`` as closely as it can. + + Takes the same geometric context as :meth:`contributions` and for the same + reason: an r-dependent payload cannot build its modes without it. ``w_dense`` is + the ``(n_atoms, K)`` weight matrix at the seeded kernel widths, which is what + makes the payload-only problem linear. + """ raise NotImplementedError def log_magnitude(self, payload): @@ -229,7 +245,7 @@ class IsotropicPayload(NodePayload): def contributions(self, payload, xyz, node_pos, neighbor_list): return torch.exp(payload[:, 0])[neighbor_list].unsqueeze(-1) - def fit(self, target, w_dense, epsilon): + def fit(self, target, w_dense, epsilon, xyz, node_pos, neighbor_list): b = _ridged_solve(w_dense, target.unsqueeze(-1)).squeeze(-1) return torch.log(b.clamp(min=epsilon)).unsqueeze(-1) @@ -262,7 +278,7 @@ def __init__(self, epsilon: float = 1e-3): def contributions(self, payload, xyz, node_pos, neighbor_list): return raw6_to_u6(payload, self.epsilon)[neighbor_list] - def fit(self, target, w_dense, epsilon): + def fit(self, target, w_dense, epsilon, xyz, node_pos, neighbor_list): """Fit six U components at once, then re-encode as Cholesky parameters. The per-atom U is linear in each component independently, so this is the same @@ -282,6 +298,233 @@ def log_magnitude(self, payload): return torch.log(b_eq.clamp(min=1e-6)) +# ---------------------------------------------------------------------------------- +# Displacement-mode generators. A gradient mode is a constant 3x3 matrix G acting on the +# displacement r from the node, giving the displacement field psi(r) = G r. Rotation, +# dilation and deviatoric strain together span every linear displacement field, and +# splitting them that way is what lets a mode set stop partway. +# ---------------------------------------------------------------------------------- + +_SQ2 = math.sqrt(2.0) +_SQ3 = math.sqrt(3.0) +_SQ6 = math.sqrt(6.0) + +# Rotations are NOT normalised: psi_i(r) = e_i x r exactly, so that the rigid mode set +# reproduces the textbook TLS formula with no stray factor. The others are Frobenius +# normalised, which is a conditioning choice and nothing more. +_GENERATORS = { + "rotation": [ + [[0.0, 0.0, 0.0], [0.0, 0.0, -1.0], [0.0, 1.0, 0.0]], # e1 x r + [[0.0, 0.0, 1.0], [0.0, 0.0, 0.0], [-1.0, 0.0, 0.0]], # e2 x r + [[0.0, -1.0, 0.0], [1.0, 0.0, 0.0], [0.0, 0.0, 0.0]], # e3 x r + ], + "dilation": [ + [[1 / _SQ3, 0.0, 0.0], [0.0, 1 / _SQ3, 0.0], [0.0, 0.0, 1 / _SQ3]], + ], + "deviatoric": [ + [[1 / _SQ2, 0.0, 0.0], [0.0, -1 / _SQ2, 0.0], [0.0, 0.0, 0.0]], + [[1 / _SQ6, 0.0, 0.0], [0.0, 1 / _SQ6, 0.0], [0.0, 0.0, -2 / _SQ6]], + [[0.0, 1 / _SQ2, 0.0], [1 / _SQ2, 0.0, 0.0], [0.0, 0.0, 0.0]], + [[0.0, 0.0, 1 / _SQ2], [0.0, 0.0, 0.0], [1 / _SQ2, 0.0, 0.0]], + [[0.0, 0.0, 0.0], [0.0, 0.0, 1 / _SQ2], [0.0, 1 / _SQ2, 0.0]], + ], +} + +#: Named mode sets, in order of expressiveness. Three translations are always present; +#: each entry lists the gradient modes added on top. +MODE_SETS = { + "constant": (), + "rigid": ("rotation",), + "rigid_dilation": ("rotation", "dilation"), + "affine": ("rotation", "dilation", "deviatoric"), +} + + +class ModeCovariancePayload(NodePayload): + """A node carries the covariance of its displacement modes; TLS is one mode set. + + Instead of storing an ADP and averaging it, store the *displacement field* the node + represents and take its covariance. With ``q`` modes ``Psi(r) = [psi_1(r) ... psi_q(r)]`` + the node's disorder is ``u(r) = Psi(r) c`` for a random coefficient vector ``c``, and + the ADP an atom at displacement ``r`` receives is:: + + U(r) = Psi(r) Sigma Psi(r)^T, Sigma = , Sigma = L L^T + + Three properties follow, and they are the whole reason for the form: + + * **Positive-semidefinite at every r, unconditionally**, because + ``U = (Psi L)(Psi L)^T``. An arbitrary polynomial in ``r`` carries no such + guarantee and goes indefinite somewhere --- and "somewhere" is the edge of the + node's region, exactly where the softmax weights have not yet decayed. + * **Spatial variation becomes intra-node and smooth by construction.** A constant-U + node can only express variation by having neighbours, so detail costs nodes, and + every added node is another kernel that can collapse onto a single atom. Here one + node's U already varies across its whole region, and it cannot spike. + * ``U(r)`` is **linear in Sigma**, so fitting stays a linear problem. + + Mode sets, from :data:`MODE_SETS`, with ``q(q+1)/2`` parameters per node: + + ================== === ======= ========================================= + set q params model + ================== === ======= ========================================= + ``constant`` 3 6 one U per node; same model as + :class:`AnisotropicPayload` + ``rigid`` 6 21 **exactly TLS** (20 determinable; ``tr S`` + is the one flat direction) + ``rigid_dilation`` 7 28 TLS plus uniform breathing + ``affine`` 12 78 full linear displacement field: TLS plus + shear and extension + ================== === ======= ========================================= + + With the rigid set this reproduces ``U(r) = T + A S + S^T A^T + A L A^T`` identically, + ``A`` being the matrix whose columns are ``e_i x r``: the classical TLS expression is + what ``Psi Sigma Psi^T`` expands to when the modes are three translations and three + rotations. Releasing the antisymmetry of the gradient -- the ``dilation`` and + ``deviatoric`` rungs -- gives domains that breathe and shear as well as rotate. + + Displacements are divided by the node layout's own length scale (median + nearest-neighbour node distance, detached) before the modes are built. That is pure + conditioning: without it the gradient modes carry a factor of the domain size against + the translations, and the curvature ratio between them runs to several hundred. It + is detached and derived from the layout for the reason + :class:`~torchref.refinement.targets.adp.NodeSmoothnessTarget` uses the same + quantity: the length scale is a property of where the nodes are, not something the + optimiser should tune. + + Memory scales as ``n_atoms * k * q^2`` for the gathered node factors, so this payload + is meant for the small-``K`` regime it was designed for (a handful to a few dozen + expressive nodes, with ``k_neighbors`` set to ``K``). At ``q = 12`` and + ``k = 8`` that is ~90 MB for 20k atoms; a large ``K`` *and* a large ``k`` together + is what to avoid. + + Parameters + ---------- + mode_set : str, optional + Key of :data:`MODE_SETS`. Default ``"rigid"``, i.e. TLS. + epsilon : float, optional + Floor on the Cholesky diagonal of ``Sigma``, which bounds its smallest + eigenvalue. Also the value the non-translation modes start at, so a freshly + fitted field begins as the equivalent constant-U field. + """ + + out_width = 6 + + def __init__(self, mode_set: str = "rigid", epsilon: float = 1e-3): + if mode_set not in MODE_SETS: + raise ValueError( + f"Unknown mode set {mode_set!r}. Available: {sorted(MODE_SETS)}." + ) + self.mode_set = mode_set + self.epsilon = float(epsilon) + self._gradient_names = MODE_SETS[mode_set] + self.q = 3 + sum(len(_GENERATORS[n]) for n in self._gradient_names) + self.width = chol_param_count(self.q) + self._generator_cache = {} + + def __repr__(self): + return ( + f"ModeCovariancePayload({self.mode_set!r}, q={self.q}, " + f"params={self.width})" + ) + + # ------------------------------------------------------------------ + # Modes. + # ------------------------------------------------------------------ + + def _generators(self, dtype, device): + """``(q - 3, 3, 3)`` gradient generators, cached per dtype and device.""" + key = (dtype, str(device)) + G = self._generator_cache.get(key) + if G is None: + rows = [m for n in self._gradient_names for m in _GENERATORS[n]] + G = torch.tensor(rows, dtype=dtype, device=device).reshape(-1, 3, 3) + self._generator_cache[key] = G + return G + + @staticmethod + def _length_scale(node_pos): + """Median nearest-neighbour node distance, as a plain float. + + Deliberately outside the graph. A median's derivative is supported on whichever + single node pair sits at the median, which is an artifact of the layout rather + than a direction worth following, and leaving it connected would also let the + optimiser rescale its own modes by spreading the nodes apart. Node position + keeps its real gradient through the displacement ``r``. + """ + with torch.no_grad(): + if node_pos.shape[0] < 2: + return 1.0 + d = torch.cdist(node_pos, node_pos) + d.fill_diagonal_(float("inf")) + return max(float(d.min(dim=1).values.median()), 1e-3) + + def modes(self, r): + """``(..., 3, q)`` displacement modes at (already scaled) displacement ``r``.""" + eye = torch.eye(3, dtype=r.dtype, device=r.device).expand( + *r.shape[:-1], 3, 3 + ) + if not self._gradient_names: + return eye + G = self._generators(r.dtype, r.device) + grad = torch.einsum("sij,...j->...is", G, r) + return torch.cat([eye, grad], dim=-1) + + def sigma(self, payload): + """``(K, q, q)`` mode covariance of each node, positive-definite.""" + L = raw_to_cholesky(payload, self.q, self.epsilon) + return L @ L.transpose(-1, -2) + + # ------------------------------------------------------------------ + # NodePayload interface. + # ------------------------------------------------------------------ + + def contributions(self, payload, xyz, node_pos, neighbor_list): + r = (xyz.unsqueeze(1) - node_pos[neighbor_list]) / self._length_scale(node_pos) + Psi = self.modes(r) # (N, k, 3, q) + L = raw_to_cholesky(payload, self.q, self.epsilon) # (K, q, q) + A = Psi @ L[neighbor_list] # (N, k, 3, q) + U = A @ A.transpose(-1, -2) # (N, k, 3, 3) + return torch.stack( + [U[..., 0, 0], U[..., 1, 1], U[..., 2, 2], + U[..., 0, 1], U[..., 0, 2], U[..., 1, 2]], + dim=-1, + ) + + def fit(self, target, w_dense, epsilon, xyz, node_pos, neighbor_list): + """Seed the translation block from the constant-U solve, floor the rest. + + The full joint solve is available in principle -- ``U(r)`` is linear in + ``Sigma``, so it is one least-squares problem in ``K * q(q+1)/2`` unknowns -- but + it is not what is wanted here. Seeding only the translations makes the field + start as the equivalent constant-U field, which is a state whose R-factor is + already known, so entering this parametrisation cannot make the model worse and + refinement can only move away from a sane point. It also sidesteps the joint + solve's normal equations, which stop being cheap well before ``K`` does. + """ + if target.ndim == 1: # a B target: lift to the equivalent isotropic U + u_iso = target / (8.0 * math.pi**2) + zero = torch.zeros_like(u_iso) + target = torch.stack([u_iso, u_iso, u_iso, zero, zero, zero], dim=1) + u6 = _ridged_solve(w_dense, target) # (K, 6) + sigma = u6.new_zeros(u6.shape[0], self.q, self.q) + sigma[:, :3, :3] = u6_to_matrix(u6) + # psd_to_raw clamps every eigenvalue to epsilon^2, so the gradient modes come + # out at the floor rather than at zero -- non-degenerate, and negligible against + # a real U. + return psd_to_raw(sigma, self.epsilon) + + def log_magnitude(self, payload): + """Log ``B_eq`` of ``U(0)``: the translation block, which is the node's own ADP. + + Evaluated at the node rather than averaged over its region, so the number means + the same thing for every mode set and a magnitude restraint can price nodes + without knowing which one is in use. + """ + T = self.sigma(payload)[:, :3, :3] + b_eq = (8.0 * math.pi**2 / 3.0) * (T[:, 0, 0] + T[:, 1, 1] + T[:, 2, 2]) + return torch.log(b_eq.clamp(min=1e-6)) + + def _ridged_solve(w_dense, target): """Least squares ``min ||W x - target||`` through the ridged normal equations. @@ -531,7 +774,9 @@ def _fit_nodes( w_dense = torch.zeros(xyz.shape[0], n_k, dtype=xyz.dtype, device=xyz.device) w_dense.scatter_(1, neighbor_list, w_sparse) - payload = self._payload.fit(target, w_dense, self.epsilon) + payload = self._payload.fit( + target, w_dense, self.epsilon, xyz, node_pos, neighbor_list + ) columns = [payload, log_sigma.unsqueeze(-1)] if self._refine_positions: diff --git a/torchref/model/model.py b/torchref/model/model.py index 4e6d128c..8241ecf3 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -1252,6 +1252,7 @@ def set_adp_mode( n_nodes: int = None, k_neighbors: int = 12, refine_node_positions: bool = True, + mode_set: str = None, ): """Set the atomic displacement parameter (ADP) parametrization. @@ -1267,7 +1268,7 @@ def set_adp_mode( Parameters ---------- - mode : {"isotropic", "anisotropic", "field", "field_aniso"}, optional + mode : {"isotropic", "anisotropic", "field", "field_aniso", "preserve"}, optional ``"isotropic"`` (default) converts every atom, previously anisotropic ones to ``B_eq = (8 pi^2 / 3)(U11 + U22 + U33)``. ``"anisotropic"`` converts those matching ``aniso_selection``, expanding isotropic atoms @@ -1275,6 +1276,8 @@ def set_adp_mode( with a :class:`~torchref.model.disorder_field.DisorderFieldTensor`, whose node values are least-squares fitted to the B it replaces, so the atom count stops setting the ADP parameter count. + ``"preserve"`` is a no-op, leaving the ADPs exactly as the file supplied + them: use it when the starting model's own ADPs are what is being measured. aniso_selection : str, optional Phenix-style selection for ``mode="anisotropic"``, default ``"not resname HOH and not element H"``; ignored otherwise. @@ -1286,6 +1289,13 @@ def set_adp_mode( Give each node a refinable offset from its anchor centroid, at three extra parameters per node. On by default: it is what lets the load-balancing restraint move a node toward atoms instead of only widening its kernel. + mode_set : str, optional + For ``mode="field_aniso"``, a key of + :data:`~torchref.model.disorder_field.MODE_SETS` --- ``"rigid"`` is TLS, + ``"affine"`` adds shear and extension. The node then stores the covariance + of its displacement modes, so the U it gives an atom depends on where that + atom sits inside the node's region rather than being constant across it. + Default ``None`` keeps the constant-U payload. Notes ----- @@ -1298,6 +1308,12 @@ def set_adp_mode( """ if not self.ctx.initialized or self.pdb is None: return + if mode == "preserve": + # Leave the ADPs exactly as loaded. Constructing a Refinement otherwise + # reparametrises them before anything else runs, which silently discards a + # deposited model's anisotropy -- use this when the starting model's own + # ADPs are the thing being measured. + return if mode in ("field", "field_aniso"): aniso = mode == "field_aniso" # Run the partition first either way: it owns every buffer keyed off the @@ -1332,6 +1348,7 @@ def set_adp_mode( k_neighbors=k_neighbors, refine_node_positions=refine_node_positions, anisotropic=aniso, + mode_set=mode_set, ) return if mode == "isotropic": @@ -1348,7 +1365,7 @@ def set_adp_mode( else: raise ValueError( f"Unknown ADP mode: {mode!r}. Use 'isotropic', 'anisotropic', " - "'field' or 'field_aniso'." + "'field', 'field_aniso' or 'preserve'." ) self._apply_adp_partition(aniso_mask) @@ -1377,6 +1394,7 @@ def _install_disorder_field( k_neighbors: int = 12, refine_node_positions: bool = False, anisotropic: bool = False, + mode_set: str = None, ): """Replace a per-atom ADP wrapper with a node field fitted to it. @@ -1384,14 +1402,25 @@ def _install_disorder_field( ``adp`` and leaves the model isotropic, an anisotropic one takes over ``u`` and the model refines every selected atom anisotropically. Both expect the partition to have run first, which :meth:`set_adp_mode` arranges. + + ``mode_set`` selects a displacement-mode payload in place of the constant-U one, + which is the difference between a node holding a single ADP and a node holding a + motion whose ADP varies across its region. """ from torchref.model.disorder_field import ( AnisotropicPayload, DisorderFieldTensor, IsotropicPayload, + ModeCovariancePayload, density_anchor_rows, ) + if mode_set is not None and not anisotropic: + raise ValueError( + "mode_set describes an anisotropic displacement field and has no " + "isotropic form; use mode='field_aniso'." + ) + with torch.no_grad(): xyz = self.xyz().detach() # The fit target is whatever the partition just produced: per-atom U6 for @@ -1410,12 +1439,19 @@ def _install_disorder_field( # wearing a node's clothes. anchor_rows = density_anchor_rows(xyz, min(n_nodes, len(self.pdb))) + if mode_set is not None: + payload = ModeCovariancePayload(mode_set) + elif anisotropic: + payload = AnisotropicPayload() + else: + payload = IsotropicPayload() + field = DisorderFieldTensor( initial_values=target.to(self.dtype_float), xyz_fn=self.xyz, n_nodes=n_nodes, refine_positions=refine_node_positions, - payload=AnisotropicPayload() if anisotropic else IsotropicPayload(), + payload=payload, anchor_rows=anchor_rows, k_neighbors=k_neighbors, name="aniso_U" if anisotropic else "adp", @@ -1431,7 +1467,7 @@ def _install_disorder_field( self.adp.update_refinable_mask(self.adp_mask) if self.ctx.verbose > 0: - kind = "aniso U" if anisotropic else "iso B" + kind = mode_set if mode_set else ("aniso U" if anisotropic else "iso B") was = len(self.pdb) * (6 if anisotropic else 1) print( f"ADP field ({kind}): {field.n_nodes} nodes, k={k_neighbors}, " From 303e09d418f6b331eac186f448a3eb79679375b3 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 31 Aug 2026 14:22:00 +0200 Subject: [PATCH 111/250] Cover the node-field ADP targets in the device conformance guard UNCOVERED exists to excuse device-bearing classes, and the node payloads are plain strategy objects, so listing them there was a category error: the AST inventory never finds them, test_uncovered_entries_still_exist reported them as classes that no longer exist, and the guard was red. With it red, NodeLoadTarget and NodeSmoothnessTarget went in with neither a case nor an excuse and nothing noticed. Drop the payload entries and give both targets a real target case. They are inert outside field mode but still device-bearing, so they construct on a plain model and must track its device like any other target. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- tests/helpers/device_cases.py | 20 +++++++++++++++++--- 1 file changed, 17 insertions(+), 3 deletions(-) diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index da36f0a4..1a520feb 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -384,6 +384,23 @@ class TargetDeviceCase: ).ADPLocalityTarget(b["model"]), "ADPLocalityTarget", ), + # Both are inert outside field mode but still device-bearing, so they construct + # on a plain model and must track its device like any other target. + TargetDeviceCase( + "NodeLoadTarget", + lambda b, d: __import__( + "torchref.refinement.targets.adp.node_load", fromlist=["NodeLoadTarget"] + ).NodeLoadTarget(b["model"]), + "NodeLoadTarget", + ), + TargetDeviceCase( + "NodeSmoothnessTarget", + lambda b, d: __import__( + "torchref.refinement.targets.adp.node_smoothness", + fromlist=["NodeSmoothnessTarget"], + ).NodeSmoothnessTarget(b["model"]), + "NodeSmoothnessTarget", + ), # Owns no tensors at all -- the case that exercises the request-driven # tracker path rather than the owned-tensor path. TargetDeviceCase( @@ -409,9 +426,6 @@ class TargetDeviceCase: "BaseWeighting": "abstract base; covered via ManualWeighting", "Refinement": "abstract base; covered via LBFGSRefinement in integration", "PassThroughTensor": "documented non-functional stub (parameter_wrappers.py)", - "NodePayload": "stateless strategy, holds no tensors; abstract base", - "IsotropicPayload": "stateless strategy, holds no tensors", - "AnisotropicPayload": "stateless strategy, holds only a float epsilon", "ADPTarget": "abstract base; needs a model with ADPs", "CombinedTargets": "composite container; needs its component targets", "CombinedModelTargets": "composite container; needs a loaded model", From 6a3512ff3f3f09ba4b12675d80589ed52d92163d Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 31 Aug 2026 14:36:12 +0200 Subject: [PATCH 112/250] Give both search stages one normalisation and one weight Five answers to "what is the mean intensity at this resolution" reached the live path. The rotation function used the shared Wilson fit; the translation search used per-shell means through the E-convention layer, twice; the pipeline hand-rolled a third with scatter_add, not epsilon-aware; and the fine translation refine used no normalisation at all, correlating raw |F|^2 while the coarse search that handed it a peak correlated E^2. Two stages scoring the same observations disagreed about what they were scoring. Now there is one. TranslationObs normalises and weights the observations once per run, through torchref.scaling.WilsonNormaliser, and every translation stage reads it. That is not only deduplication: normalisation and weighting are properties of the observations, so refitting per orientation recomputed the same answer for every candidate. e_values.py goes with it. Its six conventions were a seam for a question that is settled -- twelve of them moved the rotation function's truth rank by nothing, because per-resolution scaling is gauge in a correlation -- and only one was ever the default. frf/french_wilson.py was reachable only through one of the other five, so it goes too. That also removes a full shell assignment and scatter_add per FRF call, computed and then overwritten by the smooth fit. The weight is the half that is NOT gauge, and the translation search had none: every reflection counted the same. Both Crowther-Blow grids now carry weighting.inverse_variance_weight -- measurement error and model error in one denominator, the same term the rotation function weights with -- and the correlation is centred at the weighted mean. Uniform weight reproduces the old scores exactly, which is what happens when the data carry no sigmas. The translation stage's resolution window becomes a parameter. It was documented as masked and was not, so the translation search runs at full resolution while the rotation search runs at 15-4 A. tf_d_min/tf_d_max default to None, which is the existing behaviour; choosing a window is a measurement, not a refactor. 1DAW t0 recovers the same pose to three decimals (1.519 deg). Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- tests/unit/alignment/test_e_conventions.py | 174 ------ .../alignment/test_patterson_translation.py | 40 +- tests/unit/alignment/test_translation_obs.py | 173 ++++++ torchref/experimental/alignment/align.py | 3 + torchref/experimental/alignment/e_values.py | 379 ------------- torchref/experimental/alignment/frf/api.py | 73 +-- .../alignment/frf/french_wilson.py | 519 ------------------ .../alignment/frf/preprocessing.py | 20 +- torchref/experimental/alignment/pipeline.py | 118 ++-- .../experimental/alignment/rotation_search.py | 15 - .../experimental/alignment/translation.py | 388 ++++++++----- 11 files changed, 569 insertions(+), 1333 deletions(-) delete mode 100644 tests/unit/alignment/test_e_conventions.py create mode 100644 tests/unit/alignment/test_translation_obs.py delete mode 100644 torchref/experimental/alignment/e_values.py delete mode 100644 torchref/experimental/alignment/frf/french_wilson.py diff --git a/tests/unit/alignment/test_e_conventions.py b/tests/unit/alignment/test_e_conventions.py deleted file mode 100644 index c5644e85..00000000 --- a/tests/unit/alignment/test_e_conventions.py +++ /dev/null @@ -1,174 +0,0 @@ -"""Invariants every E convention has to hold, whatever it does inside. - -The conformance harness in `alignment_lab` reports on all of these and more, as -a table, over real data. These are the subset that must never break: a failure -here is a bug rather than a trade-off, so they belong in the gate rather than in -a report someone has to read. - -The epsilon check is here because it already caught one. `WilsonShellE` divided -the shell mean by ``eps`` without dividing the intensity by it, which left -``<|E|**2> = `` -- 1 on a primitive lattice and 2 on a centred one, so a -normaliser whose absolute scale depended on the space group. The rotation -function never saw it (it passes no ``eps``), which is exactly why it survived: -a defect only reachable through an argument nobody was passing yet. -""" - -import functools - -import pytest -import torch - -from torchref.experimental.alignment.e_values import ( - CalcGlobalE, CalcShellE, FrenchWilsonE, SmoothSigmaE, WilsonShellE, - WilsonShellEpsE, -) - -pytestmark = pytest.mark.unit - -CONVENTIONS = [ - WilsonShellE, WilsonShellEpsE, CalcShellE, CalcGlobalE, FrenchWilsonE, - functools.partial(SmoothSigmaE, n_coeff=6), -] - - -def _name(c): - return getattr(getattr(c, "func", c), "__name__", str(c)) - - -def _wilson_data(n=20000, seed=1, centric_frac=0.1): - """Amplitudes drawn from the distribution the conventions assume.""" - g = torch.Generator().manual_seed(seed) - s = torch.rand(n, generator=g, dtype=torch.float64) * 0.4 + 0.05 - # |E|^2 ~ Exp(1) scaled by a resolution-dependent Sigma, so there is a real - # trend for a per-shell or smooth normaliser to have to remove. - sigma = torch.exp(-40.0 * s * s) * 2500.0 + 1.0 - E2 = -torch.log(torch.rand(n, generator=g, dtype=torch.float64).clamp(min=1e-12)) - F = (E2 * sigma).sqrt() - sig_F = F * 0.05 + 0.5 - centric = torch.zeros(n, dtype=torch.bool) - centric[: int(n * centric_frac)] = True - return F, s, centric, sig_F - - -def _build(cls, F, s, centric, sig_F, eps=None, n_shells=20): - kw = {"sig_F": sig_F} if getattr(cls, "uses_sigma_f", False) else {} - return cls(F, s, centric, eps=eps, n_shells=n_shells, **kw) - - -@pytest.mark.parametrize("cls", CONVENTIONS, ids=_name) -def test_epsilon_reaches_both_sides_of_the_ratio_or_neither(cls): - """A convention's scale must not depend on the lattice centring. - - ``eps`` is a constant 2 here, which is what a centred lattice gives every - reflection. Dividing the shell mean by it and not the intensity would show - up as ``<|E|**2>`` doubling -- the bug this pins. - """ - F, s, centric, sig_F = _wilson_data() - eps = torch.full_like(F, 2.0) - without = _build(cls, F, s, centric, sig_F).E - with_eps = _build(cls, F, s, centric, sig_F, eps=eps).E - m_without = float((without * without).mean()) - m_with = float((with_eps * with_eps).mean()) - assert m_with == pytest.approx(m_without, rel=1e-6), ( - f"{_name(cls)}: <|E|^2> moves from {m_without:.4f} to {m_with:.4f} when " - f"a uniform eps=2 is supplied. A uniform multiplicity cancels out of " - f"E**2 = (F**2/eps) / ; a change means eps reached only one " - f"side of that ratio." - ) - - -@pytest.mark.parametrize("cls", CONVENTIONS, ids=_name) -def test_invariant_to_a_global_rescale_of_F(cls): - """E is a ratio, so multiplying every amplitude must change nothing. - - This is the property that lets the rotation function ignore scale entirely, - and the one a convention with any absolute constant in it would break. - """ - F, s, centric, sig_F = _wilson_data() - base = _build(cls, F, s, centric, sig_F).E - for c in (1e-3, 1e3): - scaled = _build(cls, F * c, s, centric, sig_F * c).E - assert torch.allclose(scaled, base, rtol=1e-6, atol=1e-9), ( - f"{_name(cls)}: scaling F by {c:g} moved E by up to " - f"{float((scaled - base).abs().max()):.3e}" - ) - - -@pytest.mark.parametrize("cls", CONVENTIONS, ids=_name) -def test_the_normaliser_removes_the_resolution_trend(cls): - """<|E|**2> must not drift with resolution -- that trend IS the weighting.""" - F, s, centric, sig_F = _wilson_data() - E = _build(cls, F, s, centric, sig_F).E - order = torch.argsort(s) - means = [float((E[order[k::10]] ** 2).mean()) for k in range(10)] - lo, hi = min(means), max(means) - assert hi / max(lo, 1e-12) < 1.35, ( - f"{_name(cls)}: <|E|^2> ranges {lo:.3f}..{hi:.3f} across resolution " - f"deciles; the normaliser is leaving a trend behind" - ) - - -def test_a_sigma_f_convention_names_a_calc_companion(): - """There is no measurement error on a calc set, so it needs a stand-in.""" - assert FrenchWilsonE.uses_sigma_f - companion = FrenchWilsonE.for_calc() - assert companion is not FrenchWilsonE - assert not getattr(companion, "uses_sigma_f", False) - - -def test_french_wilson_refuses_to_run_without_sigmas(): - """Silently degrading to plain Wilson would hide the whole difference.""" - F, s, centric, _ = _wilson_data(n=2000) - with pytest.raises(ValueError, match="sig_F"): - FrenchWilsonE(F, s, centric, n_shells=20) - - -def test_wilson_shell_e_is_its_defining_formula(): - """``E = F / sqrt(_shell)``, asserted against the definition itself. - - Not against a reference implementation: the one this replaced now lives in - `alignment_lab/lab/reference_normalisers.py` as a frozen oracle for - comparing future conventions, and a unit test should not reach into the lab - to find its expectation. Restating the formula here is the specification, - not a second copy of the code. - """ - F, s, centric, _ = _wilson_data() - conv = WilsonShellE(F, s, centric, n_shells=20) - - total = torch.zeros(20, dtype=F.dtype) - total.scatter_add_(0, conv.shell_idx, F * F) - count = torch.bincount(conv.shell_idx, minlength=20).to(F.dtype).clamp(min=1.0) - expected = F / (total / count).clamp(min=1e-30).index_select( - 0, conv.shell_idx).sqrt() - - assert torch.equal(conv.E, expected), ( - f"max deviation {float((conv.E - expected).abs().max()):.3e}" - ) - - -def test_a_partial_is_a_usable_convention(): - """Configuration rides in as ``functools.partial``, so lookups must survive it. - - A partial forwards ``__call__`` but not class attributes, so asking one for - ``for_calc`` or ``uses_sigma_f`` directly raises ``AttributeError``. The FRF - asks for both on every run -- which is why all three ``SmoothSigmaE`` arms - of the first convention panel failed on all 50 cells without producing a - single number. - """ - from torchref.experimental.alignment.e_values import ( - convention_class, convention_for_calc, convention_uses_sigma_f, - ) - - plain = functools.partial(SmoothSigmaE, n_coeff=6) - assert convention_class(plain) is SmoothSigmaE - assert convention_uses_sigma_f(plain) is False - # Its own companion, so the configuration has to survive the round trip. - assert convention_for_calc(plain) is plain - - fw = functools.partial(FrenchWilsonE) - assert convention_uses_sigma_f(fw) is True - # A different class is named, so its keywords are not this one's. - assert convention_for_calc(fw) is WilsonShellE - - F, s, centric, sig_F = _wilson_data(n=4000) - assert convention_for_calc(plain)(F, s, centric, n_shells=20).E.shape == F.shape diff --git a/tests/unit/alignment/test_patterson_translation.py b/tests/unit/alignment/test_patterson_translation.py index 2870659a..394a3e6b 100644 --- a/tests/unit/alignment/test_patterson_translation.py +++ b/tests/unit/alignment/test_patterson_translation.py @@ -1,12 +1,15 @@ -""" -Unit tests for `amplitude_translation_search`. - -The function does a coarse-grid Pearson correlation between |F_obs|² and -|F_calc(h, t)|² over fractional translations. With the search model placed at -canonical positions and `F_obs` derived from a translated copy of the same -model, the top correlation peak (or one of the top-3) must land at `-t_true` -modulo an allowed origin shift of the spacegroup — i.e. the translation that -would bring the search model into agreement with the observed data. +"""Unit tests for the fast translation function. + +``amplitude_translation_search`` correlates normalised, weighted ``E_obs^2`` +against ``|F_calc(h, t)|^2`` over a fractional grid. With the search model at +canonical positions and ``F_obs`` derived from a translated copy of the same +model, the top peak (or one of the top three) must land at ``-t_true`` modulo an +allowed origin shift of the space group -- the translation that would bring the +search model into agreement with the observed data. + +``TranslationObs`` carries the observed side. It is built once here, as the +pipeline builds it once per run, because normalisation and weighting are +properties of the observations and do not change when the model moves. """ from pathlib import Path @@ -14,7 +17,10 @@ import pytest import torch -from torchref.experimental.alignment.translation import amplitude_translation_search +from torchref.experimental.alignment.translation import ( + TranslationObs, + amplitude_translation_search, +) from torchref.io.datasets.reflection_data import ReflectionData from torchref.model import ModelFT @@ -61,10 +67,13 @@ def test_amplitude_tf_zero_translation(setup): model_p1.spacegroup = "P 1" evaluator = _ModelEvaluator(model_p1) + obs = TranslationObs.build( + F_obs, data.hkl[mask], data.spacegroup, data.cell, + ) R_id = torch.eye(3, dtype=torch.float64) _, _, peaks = amplitude_translation_search( - F_obs=F_obs, interpolator=evaluator, R_rotation=R_id, - hkl=data.hkl[mask], spacegroup=data.spacegroup, real_cell=data.cell, + obs=obs, interpolator=evaluator, R_rotation=R_id, + spacegroup=data.spacegroup, real_cell=data.cell, grid_steps=12, n_peaks=10, cluster_radius=0.05, ) assert len(peaks) > 0 @@ -98,10 +107,13 @@ def test_amplitude_tf_recovers_known_translation(setup): ) evaluator = _ModelEvaluator(model_p1) + obs = TranslationObs.build( + F_obs, data.hkl[mask], data.spacegroup, data.cell, + ) R_id = torch.eye(3, dtype=torch.float64) _, _, peaks = amplitude_translation_search( - F_obs=F_obs, interpolator=evaluator, R_rotation=R_id, - hkl=data.hkl[mask], spacegroup=data.spacegroup, real_cell=data.cell, + obs=obs, interpolator=evaluator, R_rotation=R_id, + spacegroup=data.spacegroup, real_cell=data.cell, grid_steps=12, n_peaks=10, cluster_radius=0.05, ) assert len(peaks) > 0 diff --git a/tests/unit/alignment/test_translation_obs.py b/tests/unit/alignment/test_translation_obs.py new file mode 100644 index 00000000..d2d606ac --- /dev/null +++ b/tests/unit/alignment/test_translation_obs.py @@ -0,0 +1,173 @@ +"""The observed side of the translation search, and what it guarantees. + +``TranslationObs`` exists so that "what is E_obs here" has exactly one answer +across the whole run. Three separate answers used to live in this path -- one in +the coarse search, one in the likelihood, one hand-rolled in the pipeline -- and +they disagreed, which meant the stage that chose a peak and the stage that +re-ranked it were not scoring the same quantity. + +These tests pin the properties that claim rests on: the normalisation is the +shared Wilson fit, the weight is a real per-reflection weight rather than a +per-shell one (per-shell weights cancel in a correlation), and the whole object +is invariant to the units the amplitudes arrive in. +""" +import math + +import pytest +import torch + +from torchref.experimental.alignment.translation import TranslationObs +from torchref.scaling import WilsonNormaliser +from torchref.symmetry.cell import Cell +from torchref.symmetry.spacegroup import SpaceGroup + + +def _case(n=4000, seed=0, sg="P 21 21 21"): + """A synthetic reflection set with a realistic Wilson falloff.""" + g = torch.Generator().manual_seed(seed) + # Pinned to CPU: this host has an accelerator, and Cell/SpaceGroup would + # land there while the synthetic hkl below stays on the host. + cell = Cell([61.0, 72.0, 83.0, 90.0, 90.0, 90.0], device="cpu") + spacegroup = SpaceGroup(sg, device="cpu") + # Miller indices on a coarse block, origin removed. + rng = torch.arange(-9, 10) + h, k, l = torch.meshgrid(rng, rng, rng, indexing="ij") + hkl = torch.stack([h.flatten(), k.flatten(), l.flatten()], dim=-1) + hkl = hkl[(hkl.abs().sum(dim=-1) > 0)] + hkl = hkl[torch.randperm(hkl.shape[0], generator=g)[:n]] + + s_mag = (hkl.to(torch.float64) + @ cell.reciprocal_basis_matrix.to(torch.float64)).norm(dim=-1) + # Exponential intensities on a Wilson curve, so = 1 is reachable. + Sigma = 3000.0 * torch.exp(-2.0 * 25.0 * s_mag ** 2) + I = Sigma * -torch.rand(n, generator=g, dtype=torch.float64).clamp(min=1e-9).log() + F = I.sqrt() + sig_F = 0.05 * F + 0.01 * F.mean() + return F, sig_F, hkl, spacegroup, cell, s_mag + + +@pytest.mark.unit +def test_normalisation_is_the_shared_wilson_fit(): + """E_obs is WilsonNormaliser's E, not a private per-shell mean.""" + F, _, hkl, sg, cell, _ = _case() + obs = TranslationObs.build(F, hkl, sg, cell) + + direct = WilsonNormaliser( + obs.F_obs * obs.F_obs, obs.s_mag, eps=obs.eps, centric=obs.centric, + n_coeff=6, + ) + torch.testing.assert_close(obs.E_obs, direct.E.to(torch.float64)) + + +@pytest.mark.unit +def test_mean_e_squared_is_one(): + """ = 1 is an identity of the Gamma fit, so it holds to fit precision. + + k-weighted, because that is the score equation the intercept solves: + sum_h k_h (I_h/mu_h - 1) = 0 with k = 1 acentric, 1/2 centric. + """ + F, _, hkl, sg, cell, _ = _case() + obs = TranslationObs.build(F, hkl, sg, cell) + k = torch.where(obs.centric, 0.5, 1.0).to(torch.float64) + mean_e2 = (k * obs.E_obs ** 2).sum() / k.sum() + assert abs(float(mean_e2) - 1.0) < 1e-6, mean_e2 + + +@pytest.mark.unit +def test_e_obs_is_invariant_to_the_amplitude_scale(): + """Rescaling every amplitude must not change E. It is an ABSOLUTE normaliser. + + Exact in the model -- a common factor lands entirely in Sigma's intercept -- + but the fit is IRLS, so the tolerance is its convergence floor rather than + machine epsilon. Measured spread is ~2e-8 over 4000 reflections; 1e-6 catches + a genuine scale dependence without chasing the solver. + """ + F, sig_F, hkl, sg, cell, _ = _case() + base = TranslationObs.build(F, hkl, sg, cell, sig_F=sig_F) + scaled = TranslationObs.build(7.5 * F, hkl, sg, cell, sig_F=7.5 * sig_F) + torch.testing.assert_close(base.E_obs, scaled.E_obs, rtol=1e-6, atol=1e-6) + # F/sigma is unchanged by a common factor, so the weight must be too. + torch.testing.assert_close(base.weight, scaled.weight, rtol=1e-6, atol=1e-6) + + +@pytest.mark.unit +def test_weight_varies_within_a_shell(): + """The part of the weight that is not gauge is the part that varies within a shell. + + A weight constant inside a resolution shell is a per-shell weight, and a + correlation absorbs those -- which is the whole reason twelve E conventions + moved the rotation function's truth rank by nothing. So this asserts the + thing that makes weighting worth doing at all, not merely that a weight + exists. + """ + F, sig_F, hkl, sg, cell, _ = _case() + # Give two reflections at the SAME resolution very different sigmas. + obs = TranslationObs.build(F, hkl, sg, cell, sig_F=sig_F) + + within = [] + for shell in range(obs.n_shells): + w = obs.weight[obs.shell_idx == shell] + if w.numel() > 20: + within.append(float(w.std() / w.mean().clamp(min=1e-30))) + assert within, "no populated shells" + assert min(within) > 1e-3, ( + f"weight is effectively constant within shells (max rel. spread " + f"{max(within):.2e}); it would be absorbed by the correlation" + ) + + +@pytest.mark.unit +def test_weight_is_uniform_without_sigmas(): + """No sigmas, no weight. The varying half of the weight IS the measurement term.""" + F, _, hkl, sg, cell, _ = _case() + obs = TranslationObs.build(F, hkl, sg, cell, sig_F=None) + assert torch.allclose(obs.weight, torch.ones_like(obs.weight)) + + +@pytest.mark.unit +def test_weight_is_normalised_to_mean_one(): + """So the score's scale does not depend on how the sigmas happened to be scaled.""" + F, sig_F, hkl, sg, cell, _ = _case() + obs = TranslationObs.build(F, hkl, sg, cell, sig_F=sig_F) + assert abs(float(obs.weight.mean()) - 1.0) < 1e-9 + + +@pytest.mark.unit +@pytest.mark.parametrize("sg_name", ["P 1", "P 21 21 21", "C 1 2 1", "P 43 21 2"]) +def test_epsilon_and_centricity_come_from_the_spacegroup(sg_name): + """Both reach the fit, and epsilon is the friedel=False count. + + Wilson's = eps*Sigma counts the operations mapping h to itself, which add + coherently and set the mean. The Friedel-folded count answers a different + question -- it describes the distribution's shape, which enters as the Gamma + shape via centricity, separately. Applying the wrong one here shifts the + normalisation of every axial reflection. + """ + F, _, hkl, sg, cell, _ = _case(sg=sg_name) + obs = TranslationObs.build(F, hkl, sg, cell) + hkl_l = obs.hkl.round().to(torch.int64) + torch.testing.assert_close( + obs.eps, sg.epsilon(hkl_l, friedel=False).to(torch.float64).clamp(min=1.0), + ) + torch.testing.assert_close(obs.centric, sg.is_centric(hkl_l).to(torch.bool)) + + +@pytest.mark.unit +def test_shell_binning_is_equal_count_and_shared(): + """One binning, used by both the sigma_A fit and the likelihood. + + They used to derive their own from the same |s| -- one rank-based, one + value-based -- which put boundary reflections in different shells depending + on which stage asked. + """ + F, _, hkl, sg, cell, _ = _case() + obs = TranslationObs.build(F, hkl, sg, cell, n_shells=10) + counts = torch.bincount(obs.shell_idx, minlength=obs.n_shells) + assert obs.n_shells == 10 + assert int(counts.min()) > 0 + # Equal-count binning: no shell should be wildly larger than the mean. + assert float(counts.max()) < 3.0 * float(counts.to(torch.float64).mean()) + # And it must be monotone in |s|: shells partition resolution, not noise. + hi = torch.stack([obs.s_mag[obs.shell_idx == b].max() + for b in range(obs.n_shells)]) + assert bool((hi[1:] >= hi[:-1]).all()) diff --git a/torchref/experimental/alignment/align.py b/torchref/experimental/alignment/align.py index 5c6e0c29..2530bd40 100644 --- a/torchref/experimental/alignment/align.py +++ b/torchref/experimental/alignment/align.py @@ -282,6 +282,8 @@ def align_model_to_data( translation_grid_steps: int = 16, n_rotation_candidates: int = 15, use_llg_tf: bool = False, + tf_d_min: Optional[float] = None, + tf_d_max: Optional[float] = None, model_error_A: Optional[float] = None, ) -> "ModelFT": """Place ``model`` in ``data``'s crystal: rotation search, then translation. @@ -316,6 +318,7 @@ def align_model_to_data( n_translation_candidates=n_translation_candidates, translation_grid_steps=translation_grid_steps, use_llg_tf=use_llg_tf, + tf_d_min=tf_d_min, tf_d_max=tf_d_max, ) solutions = pipeline.run(do_translation=do_translation) return solutions[0].model diff --git a/torchref/experimental/alignment/e_values.py b/torchref/experimental/alignment/e_values.py deleted file mode 100644 index ff9dae45..00000000 --- a/torchref/experimental/alignment/e_values.py +++ /dev/null @@ -1,379 +0,0 @@ -"""One place to say what ``E`` means. - -``E = F / sqrt(Sigma(s))`` is a **weighting choice wearing the costume of a units -change**: correlating ``E_obs`` against ``E_calc`` *is* correlating ``F`` against -``F`` with weight ``1/Sigma(s)``. The alignment package currently answers that -question nine different times -- twice in ``frf.preprocessing``, once inside -``french_wilson_preprocess`` and three more in ``translation`` -- and the -answers disagree. The rotation function's observed -side is a French-Wilson posterior weighted by ``DFAC**2``; the rescore's is plain -per-shell Wilson with epsilon divided out. So the rescore ranks candidates -against a differently-normalised observation set than the one that produced them. - -The two consumers do not need the same thing from it, which is worth stating -because it explains which of them breaks: - -* The rotation function is a **correlation**. A global scale cancels -- it scales - every SO(3) sample equally and the peak is reported as a z-score -- so only the - *relative* weighting across resolution matters. Removing the antipodal copy - scaled every score by exactly 4 and moved 98 of 100 truth ranks not at all. -* The rescore's LLG is a **likelihood**. It compares an observation against a - predicted distribution, so there is no free scale to cancel; get it wrong and - you evaluate the right data against the wrong Rice. - -A convention that satisfies the likelihood satisfies the correlation for free, -so the strict requirement is the one to design against. - -Conventions are constructed **from the data** rather than configured and passed -in, because a fitted ``Sigma(s)`` cannot exist before the reflections do. Engines -therefore take the class and instantiate it internally:: - - FastRotationFunction(..., e_convention=FrenchWilsonE) - -and anything needing configuration rides in as -``functools.partial(SmoothSigmaE, n_coeff=6)``, which is class-like and needs no -extra parameter. -""" - -from __future__ import annotations - -import math -from typing import Optional - -import torch - -from .sh import assign_shells, equal_count_shell_edges - -__all__ = [ - "CalcGlobalE", - "CalcShellE", - "EConvention", - "FrenchWilsonE", - "SmoothSigmaE", - "WilsonShellE", - "WilsonShellEpsE", - "convention_class", - "convention_for_calc", - "convention_uses_sigma_f", -] - - -def convention_class(conv) -> type: - """The class behind a convention, which may be a ``functools.partial``. - - Configuration rides in as ``partial(SmoothSigmaE, n_coeff=6)``, and a - partial forwards ``__call__`` but not class attributes -- so asking one for - ``uses_sigma_f`` or ``for_calc`` raises. Every attribute lookup on a - convention goes through here for that reason. - """ - return getattr(conv, "func", conv) - - -def convention_for_calc(conv): - """The convention to normalise **calculated** amplitudes with. - - Keeps the partial's configuration when the class is its own companion, and - drops it when a different class is named -- another class's keywords are not - this one's. - """ - companion = convention_class(conv).calc_companion - return conv if companion is None else companion - - -def convention_uses_sigma_f(conv) -> bool: - """Whether ``conv`` consumes ``sig_F``, partial or not.""" - return bool(getattr(convention_class(conv), "uses_sigma_f", False)) - - -class EConvention: - """Normalised amplitudes, plus the per-reflection weight that goes with them. - - Attributes - ---------- - E : torch.Tensor - ``(N,)`` normalised amplitude. For most conventions this is - ``F / sqrt(sigma)``; for :class:`FrenchWilsonE` it is a posterior - expectation and the relation is only approximate, which is exactly the - difference the conformance harness is there to expose. - weight : torch.Tensor - ``(N,)`` per-reflection information weight. ``DFAC**2`` where the - convention models measurement error, ones where it does not. This is the - "weight by F/sigma" lever, made explicit rather than left implicit in - whichever normaliser a caller happened to pick. - sigma : torch.Tensor - ``(N,)`` the normaliser actually used, ```` per reflection. - Reported for every convention so they are comparable even when their - ``E`` is not defined the same way. - - Notes - ----- - ``eps`` divides the intensity before averaging (``E**2 = (F**2/eps) / - ``) because axial reflections are systematically stronger -- - `` = eps * Sigma`` -- and would otherwise dominate. Conventions that leave - it ``None`` are declaring that their caller handles multiplicity some other - way; the rotation function does, by unrolling the full orbit. - """ - - #: Whether this convention consumes ``sig_F``. The conformance harness skips - #: the shrinkage test for conventions that do not. - uses_sigma_f: bool = False - - #: Whether this convention divides the intensity by ``eps``. A convention - #: that does not must ignore it on BOTH sides of the ratio: applying it to - #: the shell mean alone leaves `` = ``, which is 2 on a centred - #: lattice and 1 on a primitive one -- a normaliser whose scale depends on - #: the space group. Declared rather than implied so the two halves cannot - #: drift apart again. - uses_epsilon: bool = True - - #: The convention to use for **calculated** amplitudes, when it cannot be - #: this one. A French-Wilson posterior is defined for observations only -- - #: there is no measurement error on a calc set to shrink toward the mean -- - #: so a sigma_F-consuming convention has to name a companion. ``None`` means - #: "use this class for both sides", which is what most of them do. - #: - #: This is not a harness convenience: it is why the rotation function pairs - #: `french_wilson_preprocess` on obs with `wilson_normalise` on calc. - calc_companion: Optional[type] = None - - @classmethod - def for_calc(cls) -> type: - """The class to normalise calculated amplitudes with.""" - return cls.calc_companion or cls - - def __init__( - self, - F: torch.Tensor, - s_mag: torch.Tensor, - centric: Optional[torch.Tensor] = None, - *, - sig_F: Optional[torch.Tensor] = None, - eps: Optional[torch.Tensor] = None, - shell_idx: Optional[torch.Tensor] = None, - n_shells: int = 20, - ) -> None: - if F.ndim != 1: - raise ValueError(f"F must be 1-D, got {tuple(F.shape)}") - if s_mag.shape != F.shape: - raise ValueError( - f"s_mag {tuple(s_mag.shape)} does not match F {tuple(F.shape)}" - ) - self.F = F - self.s_mag = s_mag - self.centric = ( - torch.zeros_like(F, dtype=torch.bool) if centric is None - else centric.to(torch.bool) - ) - self.sig_F = sig_F - self.eps = eps - self.n_shells = int(n_shells) - # One shell assignment, shared by whatever the subclass needs it for. - # Assigning here rather than in each subclass is the same fix the FRF - # needed: two consumers deriving their own equal-count edges from the - # same |s| disagree about the reflections sitting on a boundary. - if shell_idx is None: - edges, _ = equal_count_shell_edges(s_mag, self.n_shells) - shell_idx = assign_shells(s_mag, edges) - self.shell_idx = shell_idx.clamp(min=0) - - self.sigma = self._shell_mean_intensity() - self.E, self.weight = self._compute() - - # -- helpers shared by the subclasses --------------------------------- - - def _intensity(self) -> torch.Tensor: - """``F**2 / eps`` -- the quantity whose shell mean is ``Sigma``. - - Both the numerator of ``E**2`` and its shell mean come through here, so - ``uses_epsilon`` reaches the ratio consistently by construction. - """ - I = self.F * self.F - if self.eps is not None and self.uses_epsilon: - I = I / self.eps.clamp(min=1.0) - return I - - def _shell_mean_intensity(self) -> torch.Tensor: - """```` per reflection, from the shared shell assignment.""" - I = self._intensity() - total = torch.zeros(self.n_shells, dtype=I.dtype, device=I.device) - total.scatter_add_(0, self.shell_idx, I) - count = torch.bincount( - self.shell_idx, minlength=self.n_shells, - ).to(I.dtype).clamp(min=1.0) - return (total / count).clamp(min=1e-30).index_select(0, self.shell_idx) - - def _ones(self) -> torch.Tensor: - return torch.ones_like(self.F) - - def _compute(self): - raise NotImplementedError - - def __repr__(self) -> str: # pragma: no cover - display - return f"{type(self).__name__}(N={self.F.numel()}, n_shells={self.n_shells})" - - -class WilsonShellE(EConvention): - """Plain per-shell Wilson: ``E = F / sqrt(_shell)``. - - What the rotation function uses on the calc side, and on the obs side when - the data carry no sigmas. Ignores measurement error, and multiplicity with - it -- both deliberately. The calc side is a single molecular transform - sampled in a P1 box, where multiplicity has no meaning; the obs side gets - its multiplicity from the symmetry unroll, which puts each reflection into - the sum once per operation that reaches it. - - So an ``eps`` passed to this class is *ignored*, not half-applied. Use - :class:`WilsonShellEpsE` when it should count. - """ - - uses_epsilon = False - - def _compute(self): - return self.F / self.sigma.sqrt(), self._ones() - - -class WilsonShellEpsE(EConvention): - """Epsilon-corrected Wilson, ``E**2 = (F**2/eps) / _shell``. - - What the m_LETF1 rescore uses on the observed side. Identical to - :class:`WilsonShellE` when ``eps`` is absent, which is worth knowing: the - difference between the two is only ever the multiplicity handling. - """ - - def _compute(self): - E = (self._intensity() / self.sigma).clamp(min=0.0).sqrt() - return E, self._ones() - - -class FrenchWilsonE(EConvention): - """French-Wilson posterior amplitude with the Rice ``DFAC`` weight. - - The rotation function's observed-side default, and the only convention here - that looks at ``sig_F``. ``E`` is the posterior expectation of the true - normalised amplitude given a noisy measurement, so weak reflections shrink - toward the shell mean instead of being taken at face value; ``weight`` is - ``DFAC**2``, the Rice-moment D factor, which is the per-reflection - measurement-information term. - - Requires ``sig_F``. Falling back silently to plain Wilson would hide exactly - the difference this class exists to make visible. - - ``french_wilson_preprocess`` takes no multiplicity, so neither does this -- - declared so the reported ``sigma`` describes what was actually done rather - than what the base class would have done. - """ - - uses_sigma_f = True - uses_epsilon = False - calc_companion = WilsonShellE - - def _compute(self): - if self.sig_F is None: - raise ValueError( - "FrenchWilsonE needs sig_F; use WilsonShellE for data without " - "measurement errors rather than letting the difference pass " - "silently." - ) - from .frf.french_wilson import french_wilson_preprocess - - fw = french_wilson_preprocess( - self.F, self.sig_F, self.s_mag, self.centric, - n_wilson_shells=self.n_shells, shell_idx=self.shell_idx, - ) - dfac = fw["DFAC"] - return fw["eEobs"], dfac * dfac - - -class CalcShellE(WilsonShellE): - """The rescore's calc-side normaliser: per-shell, flattening every shell to 1. - - Named separately from :class:`WilsonShellE` because it is applied to a - *reference* orientation's ``|F_calc|`` and then reused for every rotated - candidate. Predicted to fail the obs/calc common-scale check: forcing - ``_shell = 1`` in every shell discards the model's inter-shell - amplitude shape, which is the very thing the likelihood's expected intensity - is supposed to carry. - """ - - -#: The observed side carries multiplicity; the calculated side is a single -#: molecular transform sampled at the same Miller indices, where multiplicity -#: has no meaning. Assigned out of the class body only because `CalcShellE` is -#: defined below `WilsonShellEpsE`. -WilsonShellEpsE.calc_companion = CalcShellE - - -class CalcGlobalE(EConvention): - """Single global scale: ``E = F / rms(F)``, preserving inter-shell shape. - - The rescore's ``scat_mode="absolute"``. Keeps how much the model actually - scatters per resolution instead of flattening it, which is what makes a - relative Wilson-B correction meaningful rather than cancelled. - - Calculated amplitudes, so multiplicity does not apply -- see - :class:`WilsonShellE`. - """ - - uses_epsilon = False - - def _compute(self): - rms = self._intensity().mean().clamp(min=1e-30).sqrt() - self.sigma = torch.full_like(self.F, float(rms * rms)) - return self.F / rms, self._ones() - - -class SmoothSigmaE(EConvention): - """Adapter over the shared :class:`~torchref.scaling.WilsonNormaliser`. - - The fit itself does not live here, because "what is the mean intensity at - this resolution" is not an alignment question -- at least five private - answers to it grew across the repo, and consumers that disagree about it - cannot be compared with each other. What stays here is the ``EConvention`` - protocol that this package's consumers and its conformance harness are - written against. - - ``weight`` is ones, and that is the point of the split rather than an - omission: this class answers *what* we compare. How much each reflection - counts is a weight, built from ``sigI`` and model error, and belongs - elsewhere. Returning both from one object is what made the previous - conventions impossible to interpret -- sweeping one moved a gauge quantity - and a real one at the same time. - - Pass ``s_lo``/``s_hi`` whenever obs and calc are fitted separately and their - curves will be compared: the basis saturates at the ends, so a curve - evaluated beyond its own fitted range is frozen flat rather than - extrapolated. - """ - - #: Chebyshev terms. Provisional -- the order has never been screened against - #: a metric sensitive to it. - DEFAULT_N_COEFF = 6 - - def __init__(self, *args, n_coeff: int = DEFAULT_N_COEFF, - s_lo=None, s_hi=None, **kwargs) -> None: - self.n_coeff = int(n_coeff) - self.s_lo = s_lo - self.s_hi = s_hi - super().__init__(*args, **kwargs) - - def _compute(self): - from torchref.scaling import WilsonNormaliser - - # `_intensity()` has already divided by eps -- `uses_epsilon` is True - # here -- so the normaliser gets a reduced intensity and no eps of its - # own. Applying it in both places would count multiplicity twice. - fit = WilsonNormaliser( - self._intensity(), self.s_mag, centric=self.centric, - n_coeff=self.n_coeff, s_lo=self.s_lo, s_hi=self.s_hi, - ) - # Kept, not discarded: the fitted CURVE is the thing anything comparing - # two normalisations needs. Sigma_obs/Sigma_calc is how model error is - # measured rather than assumed, and it can only be formed from the fits - # themselves, not from the per-reflection values they produced. - self.fit = fit - self.sigma = fit.sigma_wilson - return fit.E, self._ones() - - def evaluate(self, s_mag: torch.Tensor) -> torch.Tensor: - """``Sigma(s)`` at arbitrary ``|s|``, on this fit's own abscissa.""" - return self.fit.evaluate(s_mag) diff --git a/torchref/experimental/alignment/frf/api.py b/torchref/experimental/alignment/frf/api.py index 66fcb621..279d7f33 100644 --- a/torchref/experimental/alignment/frf/api.py +++ b/torchref/experimental/alignment/frf/api.py @@ -26,8 +26,7 @@ information_weight, inverse_variance_weight, normalise_weight, snr_from_amplitude, ) -from ..e_values import (SmoothSigmaE, convention_for_calc, - convention_uses_sigma_f) +from torchref.scaling import WilsonNormaliser from .data_mr import bessel_sh_expand, cross_correlate_xi from .peak_finder import find_rotation_peaks from .preprocessing import ( @@ -41,6 +40,12 @@ __all__ = ["FastRotationFunction", "phaser_lmax_resolution"] +#: Chebyshev order of the Wilson fit, both sides. Provisional -- inherited from +#: :mod:`torchref.scaling.wilson`, and never screened against a metric sensitive +#: to it. Screen it on the fit's own residual trend, not on truth rank: the +#: rotation function is a correlation, where any per-resolution scaling cancels. +WILSON_N_COEFF = 6 + def phaser_lmax_resolution( model_radius_A: float, @@ -144,7 +149,7 @@ def __init__( grid_sampling_deg: float = 2.0, asu_idx: Optional[torch.Tensor] = None, s_mag_asu: Optional[torch.Tensor] = None, - e_convention: type = SmoothSigmaE, + wilson_n_coeff: int = WILSON_N_COEFF, obs_weight: str = "inverse_variance", snr_cap: float = DEFAULT_SNR_CAP, trust_cap: float = DEFAULT_TRUST_CAP, @@ -177,12 +182,11 @@ def __init__( # mechanism for a job that already has one. self.shell_variance_weights = bool(shell_variance_weights) - # The class, not an instance: a fitted Sigma(s) cannot exist before the - # reflections do, so the convention is constructed here -- twice, once - # per side. That the same class has to normalise both is the point; - # it puts obs and calc on a common footing under test rather than under - # assumption. Pass `functools.partial(Cls, ...)` to configure one. - self.e_convention = e_convention + # Chebyshev order of the Wilson fit. Both sides are normalised by the + # same class at the same order over the same abscissa -- that is what + # puts them on a common footing, and it is the reason Sigma_obs/Sigma_calc + # is a meaningful ratio rather than two unrelated curves. + self.wilson_n_coeff = int(wilson_n_coeff) # 1. Resolution window. # @@ -249,16 +253,7 @@ def __init__( shell_edges, _ = equal_count_shell_edges(smag_src, n_wilson_shells) obs_shell_idx = assign_shells(smag_src, shell_edges) - # 3b. Wilson normalisation. With sigmas, through the French-Wilson - # posterior, which handles the axial reflections; without them, plain - # per-shell Wilson. - # A convention that reads sigmas cannot run without them. Falling back - # to its own calc companion is the same choice the hardcoded branch made - # -- French-Wilson with sigmas, plain Wilson without -- just asked of the - # convention instead of assumed about it. - obs_cls = e_convention - if sig_F_obs is None and convention_uses_sigma_f(obs_cls): - obs_cls = convention_for_calc(obs_cls) + # 3b. Wilson normalisation, so that = 1 on the observations. # One resolution window for both sides, taken from the bandwidth # coupling rather than from whichever reflections each side happens to # contain. Two fits on their own extremes span the same polynomial @@ -268,10 +263,13 @@ def __init__( # that reads Sigma_obs/Sigma_calc needs them on one abscissa. self._s_lo = 1.0 / float(d_max) if d_max else float(smag_src.min()) self._s_hi = 1.0 / float(d_min) - conv_obs = obs_cls( - F_obs, smag_src, centric_obs, sig_F=sig_F_obs, - shell_idx=obs_shell_idx, n_shells=n_wilson_shells, - **self._range_kw(obs_cls), + # No epsilon here: the observations reach this point symmetry-unrolled, + # which puts each reflection into the sum once per operation that maps + # to it, so multiplicity is already carried by the geometry. Centricity + # is separate and does enter -- it is the Gamma shape. + conv_obs = WilsonNormaliser( + F_obs * F_obs, smag_src, centric=centric_obs, + n_coeff=self.wilson_n_coeff, s_lo=self._s_lo, s_hi=self._s_hi, ) self._conv_obs = conv_obs @@ -328,23 +326,6 @@ def __init__( ) - def _range_kw(self, cls) -> dict: - """``s_lo``/``s_hi`` for conventions that fit a curve, empty otherwise. - - Only the smooth normaliser has an abscissa to pin; the per-shell - conventions bin whatever they are given and would reject the argument. - """ - import inspect - - target = getattr(cls, "func", cls) - try: - params = inspect.signature(target.__init__).parameters - except (TypeError, ValueError): # pragma: no cover - return {} - if "s_lo" not in params: - return {} - return {"s_lo": self._s_lo, "s_hi": self._s_hi} - def score_model( self, s_calc: torch.Tensor, @@ -368,15 +349,17 @@ def score_model( # measured to drop 0 of 339040 reflections on 3K7M and 0 of 271630 on # 1DAW. So take `s_calc` as given and only derive |s| from it. smag_calc = s_calc.norm(dim=-1) - calc_cls = convention_for_calc(self.e_convention) - self._conv_calc = calc_cls( - F_calc, smag_calc, n_shells=self.n_wilson_shells, - **self._range_kw(calc_cls), + # The same normaliser, at the same order, on the same abscissa. The calc + # side is a single molecular transform sampled in a P1 box, so there is + # no multiplicity and nothing is centric. + self._conv_calc = WilsonNormaliser( + F_calc * F_calc, smag_calc, + n_coeff=self.wilson_n_coeff, s_lo=self._s_lo, s_hi=self._s_hi, ) E_calc = self._conv_calc.E if sigma_a_source == "empirical": # Measured, not assumed. Both curves were fitted on one abscissa - # (see `_range_kw`), which is what makes evaluating them at the same + # (see `self._s_lo`/`_s_hi`), which is what makes evaluating them at the same # |s| meaningful. Subsumes the Babinet term: the low-resolution # solvent deficit is simply what the ratio measures, per structure, # instead of two universal constants. diff --git a/torchref/experimental/alignment/frf/french_wilson.py b/torchref/experimental/alignment/frf/french_wilson.py deleted file mode 100644 index bc4301df..00000000 --- a/torchref/experimental/alignment/frf/french_wilson.py +++ /dev/null @@ -1,519 +0,0 @@ -"""French–Wilson posterior + Luzzati DFAC chain. - -Pure ports of Phaser's ``lib/math_FrenchWilson.cc`` (centric/acentric -posterior moments via Parabolic-cylinder ratios) and the Halley-iteration -``getDfactor`` in ``lib/math_RiceLLG.cc``. The public entry point -:func:`french_wilson_preprocess` returns ``(eEobs, DFAC)`` from raw -``(F, σF, |s|, centric)``. - -Everything except ``french_wilson_preprocess`` is module-private; expose -the public name through :mod:`torchref.experimental.alignment.frf.preprocessing`. - -References (paths under -``…/reverse_engineering/phenix/.../phaser/src/``): -- ``lib/math_FrenchWilson.cc:8-178`` posterior ```` / ```` -- ``lib/math_RiceLLG.cc:12-250`` Rice-moment effective σA + Halley -- ``Dfactor.cc:87-93`` eEobs assembly + clamp -""" -from __future__ import annotations - -import torch - - -__all__ = ["french_wilson_preprocess"] - - -# ----------------------------------------------------------------------------- -# French-Wilson posterior expected values (math_FrenchWilson.cc) -# ----------------------------------------------------------------------------- - - -def _expectE_FW_acen(eosq, sigesq): - """ - Acentric posterior expected E from normalised observed intensity (eosq) - and its standard deviation (sigesq). Translates verbatim from Phaser's - `lib/math_FrenchWilson.cc:expectEFWacen` (lines 8-44). Vectorised NumPy. - - `eosq = Iobs / `, `sigesq = σIobs / `. - """ - import numpy as np - from scipy.special import erfc, pbdv - CROSS1, CROSS2 = -12.5, 18.0 - SQRT2 = np.sqrt(2.0) - x = (eosq - sigesq ** 2) / sigesq - xsqr = x * x - ee = np.empty_like(eosq) - m_neg = x < CROSS1 - if m_neg.any(): - xs = xsqr[m_neg] - num = (-916620705. + xs * - (91891800. + xs * - (-11531520. + xs * - (1935360. + xs * - (-491520. + xs * 262144.))))) - den = (-495452160. + xs * - (55050240. + xs * - (-7864320. + xs * - (1572864. + xs * - (-524288. + xs * 524288.))))) - ee[m_neg] = np.sqrt(-np.pi * sigesq[m_neg] / x[m_neg]) * num / den - m_pos = x > CROSS2 - if m_pos.any(): - xs = xsqr[m_pos] - num = (-45045. + 32. * xs * - (-315. + 8. * xs * - (-15. - 16. * xs + 128. * xs * xs))) - ee[m_pos] = (np.sqrt(sigesq[m_pos]) * num / - (32768. * x[m_pos] ** 7.5)) - m_mid = ~(m_neg | m_pos) - if m_mid.any(): - xm = x[m_mid] - pcd, _ = pbdv(-1.5, -xm) - ee[m_mid] = (np.sqrt(sigesq[m_mid] / 2.0) * np.exp(-xm * xm / 4.0) * - pcd / erfc(-xm / SQRT2)) - return ee - - -def _expectEsq_FW_acen(eosq, sigesq): - """Acentric posterior . From `expectEsqFWacen` (lines 46-78).""" - import numpy as np - from scipy.special import erfc - CROSS1, CROSS2 = -8.9, 5.7 - SQRT2_BY_PI = np.sqrt(2.0 / np.pi) - SQRT2 = np.sqrt(2.0) - eesq_base = eosq - sigesq ** 2 - x = eesq_base / (SQRT2 * sigesq) - xsqr = x * x - eesq = eesq_base.copy() - m_neg = x < CROSS1 - if m_neg.any(): - xs = xsqr[m_neg] - num = (-135135. + xs * (20790. + xs * (-3780. + xs * - (840. + xs * (-240. + xs * (96. - xs * 64.)))))) - den = (-135135. + xs * (20790. + xs * (-3780. + xs * - (840. + xs * (-240. + xs * (96. + xs * - (-64. + xs * 128.))))))) - eesq[m_neg] = eesq_base[m_neg] * num / den - m_mid = (x >= CROSS1) & (x <= CROSS2) - if m_mid.any(): - xm = x[m_mid] - eesq[m_mid] = (eesq_base[m_mid] + - SQRT2_BY_PI * sigesq[m_mid] / - (np.exp(xm * xm) * erfc(-xm))) - return eesq - - -def _expectE_FW_cen(eosq, sigesq): - """Centric posterior . From `expectEFWcen` (lines 80-113).""" - import numpy as np - from scipy.special import pbdv - CROSS1, CROSS2 = -17.5, 17.5 - SQRTPI = np.sqrt(np.pi) - x = sigesq / 2.0 - eosq / sigesq - xsqr = x * x - pcdratio = np.empty_like(x) - m_neg = x < CROSS1 - if m_neg.any(): - xn, xs = x[m_neg], xsqr[m_neg] - pcdratio[m_neg] = ((1024. * SQRTPI * (-xn) ** 6.5) / - (3465. + xs * - (840. + xs * - (384. + xs * 1024.)))) - m_pos = x > CROSS2 - if m_pos.any(): - xp, xs = x[m_pos], xsqr[m_pos] - num = (3440640. + xs * - (-491520. + xs * - (98304. + xs * - (-32768. + xs * 32768.)))) - den = (675675. + xs * - (-110880. + xs * - (26880. + xs * - (-12288. + xs * 32768.)))) - pcdratio[m_pos] = num / (den * np.sqrt(xp)) - m_mid = ~(m_neg | m_pos) - if m_mid.any(): - xm = x[m_mid] - d_neg1, _ = pbdv(-1.0, xm) - d_neghalf, _ = pbdv(-0.5, xm) - pcdratio[m_mid] = d_neg1 / d_neghalf - return np.sqrt(sigesq / np.pi) * pcdratio - - -def _expectEsq_FW_cen(eosq, sigesq): - """Centric posterior . From `expectEsqFWcen` (lines 115-152).""" - import numpy as np - from scipy.special import pbdv - CROSS1, CROSS2 = -17.5, 17.5 - x = sigesq / 2.0 - eosq / sigesq - xsqr = x * x - pcdratio = np.empty_like(x) - m_neg = x < CROSS1 - if m_neg.any(): - xn, xs = x[m_neg], xsqr[m_neg] - num = (45045. + xs * - (10080. + xs * - (3840. + xs * - (4096. - xs * 32768.)))) - den = xn * (55440. + xs * - (13440. + xs * - (6144. + xs * 16384.))) - pcdratio[m_neg] = num / den - m_pos = x > CROSS2 - if m_pos.any(): - xp, xs = x[m_pos], xsqr[m_pos] - num = (11486475. + xs * - (-1441440. + xs * - (241920. + xs * - (-61440. + xs * 32768.)))) - den = xp * (675675. + xs * - (-110880. + xs * - (26880. + xs * - (-12288. + xs * 32768.)))) - pcdratio[m_pos] = num / den - m_mid = ~(m_neg | m_pos) - if m_mid.any(): - xm = x[m_mid] - d_neg15, _ = pbdv(-1.5, xm) - d_neghalf, _ = pbdv(-0.5, xm) - pcdratio[m_mid] = d_neg15 / d_neghalf - return sigesq * pcdratio / 2.0 - - -def _french_wilson_posterior(eosq, sigesq, centric_mask): - """Wrap centric/acentric branches. - - Phaser `expectEFW` / `expectEsqFW` (lines 154-178): if sigesq <= 0 the - measurement is treated as exact and (eEFW, eEsqFW) = (sqrt(eosq), eosq). - """ - import numpy as np - eEFW = np.empty_like(eosq) - eEsqFW = np.empty_like(eosq) - zero_sig = sigesq <= 0.0 - if zero_sig.any(): - eEFW[zero_sig] = np.sqrt(np.maximum(eosq[zero_sig], 0.0)) - eEsqFW[zero_sig] = np.maximum(eosq[zero_sig], 0.0) - valid = ~zero_sig - if valid.any(): - cen = centric_mask & valid - acen = (~centric_mask) & valid - if cen.any(): - eEFW[cen] = _expectE_FW_cen(eosq[cen], sigesq[cen]) - eEsqFW[cen] = _expectEsq_FW_cen(eosq[cen], sigesq[cen]) - if acen.any(): - eEFW[acen] = _expectE_FW_acen(eosq[acen], sigesq[acen]) - eEsqFW[acen] = _expectEsq_FW_acen(eosq[acen], sigesq[acen]) - return eEFW, eEsqFW - - -# ----------------------------------------------------------------------------- -# DFAC via Halley iteration (math_RiceLLG.cc:getDfactor) -# ----------------------------------------------------------------------------- - - -def _i0e_full(x): - """Phaser's `eBesselI0(x) = I0(x)·exp(-|x|)`. Symmetric in x.""" - import numpy as np - from scipy.special import i0e - return i0e(np.abs(x)) - - -def _i1e_full(x): - """Phaser's `eBesselI1(x) = I1(x)·exp(-|x|)`. Antisymmetric in x.""" - import numpy as np - from scipy.special import i1e - return np.sign(x) * i1e(np.abs(x)) - - -def _effSigaRoot_acen(ee, eesq, sa): - """`effSigaRootAcen` (math_RiceLLG.cc:12-34).""" - import numpy as np - sigbsqr = 1.0 - sa * sa - x = 0.5 * (eesq - sigbsqr) / sigbsqr - return (np.sqrt(np.pi * sigbsqr) / (2.0 * sigbsqr) * - (eesq * _i0e_full(x) + (eesq - sigbsqr) * _i1e_full(x)) - ee) - - -def _deffSigaRoot_acen(eesq, sa): - """`deffSigaRootAcen_by_dsa` (lines 36-52).""" - import numpy as np - sigbsqr = 1.0 - sa * sa - x = 0.5 * (eesq - sigbsqr) / sigbsqr - return np.sqrt(np.pi / sigbsqr) * (sa / 2.0) * _i1e_full(x) - - -def _d2effSigaRoot_acen(eesq, sa): - """`d2effSigaRootAcen_by_dsa2` (lines 54-81).""" - import numpy as np - sigasqr = sa * sa - sigapow4 = sigasqr * sigasqr - sigbsqr = 1.0 - sigasqr - xnum = eesq - sigbsqr - x = 0.5 * xnum / sigbsqr - out = np.empty_like(eesq) - big = xnum > 1e-10 - if big.any(): - I0 = _i0e_full(x[big]) - I1 = _i1e_full(x[big]) - out[big] = (np.sqrt(np.pi / sigbsqr[big]) / (2.0 * sigbsqr[big] ** 2) * - (eesq[big] * sigasqr[big] * I0 + - (eesq[big] - 1.0 - (-2.0 + eesq[big] * (2.0 + eesq[big])) * sigasqr[big] + - (eesq[big] - 1.0) * sigapow4[big]) * I1 / xnum[big])) - small = ~big - if small.any(): - samin = np.sqrt(np.maximum(1.0 - eesq[small], 0.0)) - out[small] = (np.sqrt(np.pi) * samin * - ((3.0 + samin * samin) * sa[small] - - 2.0 * (samin + samin ** 3)) / - (4.0 * eesq[small] ** 2.5)) - return out - - -def _effSigaRoot_cen(ee, eesq, sa): - """`effSigaRootCen` (lines 83-105).""" - import numpy as np - from scipy.special import erf - sigbsqr = 1.0 - sa * sa - x = 0.5 * (eesq - sigbsqr) / sigbsqr - x_safe = np.maximum(x, 0.0) - return (np.exp(-x) * np.sqrt(2.0 * sigbsqr / np.pi) + - np.sqrt(np.maximum(eesq - sigbsqr, 0.0)) * erf(np.sqrt(x_safe)) - ee) - - -def _deffSigaRoot_cen(eesq, sa): - """`deffSigaRootCen_by_dsa` (lines 107-130).""" - import numpy as np - from scipy.special import erf - sigbsqr = 1.0 - sa * sa - xnum = eesq - sigbsqr - x = 0.5 * xnum / sigbsqr - out = np.empty_like(eesq) - big = np.abs(xnum) > 1e-10 - if big.any(): - x_safe = np.maximum(x[big], 0.0) - out[big] = (sa[big] * erf(np.sqrt(x_safe)) / - np.sqrt(np.maximum(xnum[big], 1e-30)) - - np.exp(-x[big]) * np.sqrt(2.0 * sigbsqr[big] / np.pi) * - sa[big] / sigbsqr[big]) - small = ~big - if small.any(): - out[small] = (xnum[small] * np.sqrt(2.0 / np.pi) * sa[small] / - (3.0 * sigbsqr[small] ** 1.5)) - return out - - -def _d2effSigaRoot_cen(eesq, sa): - """`d2effSigaRootCen_by_dsa2` (lines 132-159).""" - import numpy as np - from scipy.special import erf - sigasqr = sa * sa - sigapow4 = sigasqr * sigasqr - sigbsqr = 1.0 - sigasqr - xnum = eesq - sigbsqr - x = 0.5 * xnum / sigbsqr - sigbsqrtpi = np.sqrt(np.pi * sigbsqr) - out = np.empty_like(eesq) - big = np.abs(xnum) > 1e-10 - if big.any(): - x_safe = np.maximum(x[big], 0.0) - d2num = ((eesq[big] - 1.0) * sigbsqr[big] ** 2 * sigbsqrtpi[big] * - erf(np.sqrt(x_safe))) - exp_part = np.where( - x[big] < 20.0, - np.sqrt(np.maximum(2.0 * xnum[big], 0.0)) * np.exp(-x[big]) * - (1.0 - eesq[big] + sigasqr[big] * - (eesq[big] + eesq[big] ** 2 - 2.0) + sigapow4[big]), - np.zeros_like(x[big]), - ) - d2num = d2num + exp_part - out[big] = d2num / (sigbsqr[big] ** 2 * - np.maximum(xnum[big], 1e-30) ** 1.5 * - sigbsqrtpi[big]) - small = ~big - if small.any(): - out[small] = (np.sqrt(2.0 / np.pi) * sa[small] / - (1.5 * sigbsqr[small] ** 1.5)) - return out - - -def _get_dfactor_vectorised(ee_np, eesq_np, centric_np): - """Vectorised port of Phaser's ``math_RiceLLG.cc:getDfactor`` (lines 191-250). - - Halley's method with bisection fallback, run over all reflections in - parallel. Each reflection has its own bracket ``[dflo, dfhi]``. Returns a - ``(N,)`` numpy float64 array of DFAC values in ``(0, 1)``. - """ - import numpy as np - - EPS1 = 1e-7 - EPS2 = 1e-10 - MAXDFAC = 1.0 - EPS1 - - ee = np.asarray(ee_np, dtype=np.float64) - eesq = np.asarray(eesq_np, dtype=np.float64) - cen = np.asarray(centric_np, dtype=bool) - N = ee.shape[0] - - out = np.ones(N, dtype=np.float64) - has_err = (eesq - ee * ee) > 0.0 - - if not has_err.any(): - return out - - ee_a, eesq_a, cen_a = ee[has_err], eesq[has_err], cen[has_err] - dflo = np.maximum(np.sqrt(np.maximum(1.0 - np.minimum(eesq_a, 1.0), 0.0)) + EPS1, EPS1) - dfhi = np.full_like(dflo, MAXDFAC) - - early = dflo >= MAXDFAC - if early.any(): - pass - - dfmid = 0.5 * (dflo + dfhi) - fmid = np.empty_like(dfmid) - if cen_a.any(): - fmid[cen_a] = _effSigaRoot_cen(ee_a[cen_a], eesq_a[cen_a], dfmid[cen_a]) - if (~cen_a).any(): - fmid[~cen_a] = _effSigaRoot_acen(ee_a[~cen_a], eesq_a[~cen_a], dfmid[~cen_a]) - - active = ~early - for _ in range(50): - if not active.any(): - break - conv = (dfhi - dflo) <= EPS1 - conv |= np.abs(fmid) <= EPS2 - active = active & ~conv - if not active.any(): - break - - slope = np.empty_like(dfmid) - curve = np.empty_like(dfmid) - cen_act = cen_a & active - acen_act = (~cen_a) & active - if cen_act.any(): - slope[cen_act] = _deffSigaRoot_cen(eesq_a[cen_act], dfmid[cen_act]) - curve[cen_act] = _d2effSigaRoot_cen(eesq_a[cen_act], dfmid[cen_act]) - if acen_act.any(): - slope[acen_act] = _deffSigaRoot_acen(eesq_a[acen_act], dfmid[acen_act]) - curve[acen_act] = _d2effSigaRoot_acen(eesq_a[acen_act], dfmid[acen_act]) - - denom_halley = 2.0 * (slope ** 2 - fmid * curve) - use_halley = (curve > 0.0) & (np.abs(denom_halley) > 1e-30) - step = np.where( - use_halley, - 2.0 * fmid * slope / np.where(use_halley, denom_halley, 1.0), - fmid * slope, - ) - dfnew = dfmid - step - in_bracket = (dfnew > dflo) & (dfnew < dfhi) - dfmid_new = np.where(in_bracket, dfnew, 0.5 * (dflo + dfhi)) - - dfmid = np.where(active, dfmid_new, dfmid) - - if cen_act.any(): - fmid[cen_act] = _effSigaRoot_cen(ee_a[cen_act], eesq_a[cen_act], - dfmid[cen_act]) - if acen_act.any(): - fmid[acen_act] = _effSigaRoot_acen(ee_a[acen_act], eesq_a[acen_act], - dfmid[acen_act]) - - below = (fmid < 0.0) & active - above = (fmid >= 0.0) & active - dflo = np.where(below, dfmid, dflo) - dfhi = np.where(above, dfmid, dfhi) - - out[has_err] = dfmid - return np.clip(out, EPS1, MAXDFAC) - - -def french_wilson_preprocess( - F: torch.Tensor, - sig_F: torch.Tensor, - s_mag: torch.Tensor, - centric: torch.Tensor, - *, - n_wilson_shells: int = 20, - shell_idx: "torch.Tensor | None" = None, -) -> dict: - """Phaser-style preprocessing from raw ``(F, σF, centric)`` to ``(eEobs, DFAC)``. - - Implements the chain: - - 1. equal-count Wilson shells over ``s_mag`` -- or ``shell_idx``, when the - caller has already assigned them. Pass it: this routine's own quantile - edges (``np.linspace(0, N-1, P+1).round()``) pick a different rank than - ``equal_count_shell_edges`` does for the same distribution at a different - N, so two independently-binned consumers disagree about a handful of - reflections at the shell boundaries (measured: 7 of 55078 on 3K7M). - 2. per-shell ``_p`` (Phaser's ``SIGMAN.BINS``) - 3. per-reflection normalised intensity ``eosq = F² / `` and σ - ``sigesq = σI / ≈ 2·F·σF / `` - 4. French-Wilson posterior ``eEFW, eEsqFW`` (``math_FrenchWilson.cc``) - 5. DFAC via Halley iteration on Rice moments (``math_RiceLLG.cc``) - 6. ``eEobs = sqrt(eEsqFW + (DFAC²−1)/DFAC²)``, clamped to ≤10 - (Phaser ``Dfactor.cc:87-93``). - - Returns a dict with torch tensors back on the input device: - eEobs: (N,) effective normalised amplitude - DFAC : (N,) per-reflection D-factor ∈ [1e-7, 1−1e-7] - """ - import numpy as np - - device = F.device - F_np = F.detach().to("cpu").to(torch.float64).numpy() - sigF_np = sig_F.detach().to("cpu").to(torch.float64).numpy() - s_np = s_mag.detach().to("cpu").to(torch.float64).numpy() - cen_np = centric.detach().to("cpu").bool().numpy() - - if shell_idx is None: - sorted_idx = np.argsort(s_np) - edges_idx = np.linspace( - 0, len(s_np) - 1, n_wilson_shells + 1).round().astype(np.int64) - s_edges = s_np[sorted_idx][edges_idx] - s_edges[0] -= 1e-6 - s_edges[-1] += 1e-6 - shell_idx = np.clip( - np.searchsorted(s_edges, s_np, side="right") - 1, - 0, n_wilson_shells - 1, - ) - else: - # Out-of-range rows come back as -1 from `assign_shells`; clamp them into - # the end shells rather than dropping them, which is what this routine's - # own edge nudge did. - shell_idx = np.clip( - shell_idx.detach().to("cpu").to(torch.int64).numpy(), - 0, n_wilson_shells - 1, - ) - F2 = F_np * F_np - mean_F2 = np.zeros(n_wilson_shells, dtype=np.float64) - counts = np.zeros(n_wilson_shells, dtype=np.int64) - np.add.at(mean_F2, shell_idx, F2) - np.add.at(counts, shell_idx, 1) - mean_F2 = mean_F2 / np.maximum(counts, 1) - mean_F2 = np.maximum(mean_F2, 1e-12) - mean_I_per_h = mean_F2[shell_idx] - - eosq = F2 / mean_I_per_h - sigesq = 2.0 * F_np * sigF_np / mean_I_per_h - sigesq = np.maximum(sigesq, 0.0) - - eEFW, eEsqFW = _french_wilson_posterior(eosq, sigesq, cen_np) - bad = eEsqFW < eEFW * eEFW - if bad.any(): - eEsqFW[bad] = eEFW[bad] ** 2 + 1e-12 - - DFAC = _get_dfactor_vectorised(eEFW, eEsqFW, cen_np) - - dfsqr = DFAC * DFAC - eEobs_sqr = eEsqFW + (dfsqr - 1.0) / np.maximum(dfsqr, 1e-30) - eEobs_sqr = np.maximum(eEobs_sqr, 0.0) - eEobs = np.sqrt(eEobs_sqr) - clamp_mask = (eEobs > 10.0) & (eEsqFW > 1.0) - if clamp_mask.any(): - eEobs[clamp_mask] = 10.0 - DFAC[clamp_mask] = 1.0 / np.sqrt(np.maximum(eEsqFW[clamp_mask] - 99.0, 1e-30)) - DFAC[clamp_mask] = np.clip(DFAC[clamp_mask], 1e-7, 1.0 - 1e-7) - - return { - "eEobs": torch.from_numpy(eEobs).to(device=device, dtype=F.dtype), - "DFAC": torch.from_numpy(DFAC).to(device=device, dtype=F.dtype), - } diff --git a/torchref/experimental/alignment/frf/preprocessing.py b/torchref/experimental/alignment/frf/preprocessing.py index 07f2bfec..43c5cda2 100644 --- a/torchref/experimental/alignment/frf/preprocessing.py +++ b/torchref/experimental/alignment/frf/preprocessing.py @@ -1,14 +1,14 @@ """Observed-side preprocessing chain. -Mirrors the chain in Phaser ``DataMR::dataMR_FRF`` (DataMR.cc:863-1133) -and the auxiliary helpers in ``lib/math_FrenchWilson.cc`` and -``lib/math_RiceLLG.cc``. The heavier ports live in sibling modules -(:mod:`~torchref.experimental.alignment.frf.french_wilson` in particular); -this module imports them and carries the Phaser source citations. - -If a specific preprocessing piece turns out to be wrong (per Tier 2 -synthetic tests), the fix lives here — replace the import with a fresh -implementation cited line-by-line to the corresponding Phaser source. +Mirrors the chain in Phaser ``DataMR::dataMR_FRF`` (DataMR.cc:863-1133) and the +auxiliary helpers in ``lib/math_RiceLLG.cc``, and carries the Phaser source +citations for each piece. + +Normalisation is deliberately **not** here. Turning amplitudes into E values is +:class:`~torchref.scaling.WilsonNormaliser`'s job, shared with the translation +search and with everything else in the repo that asks what the mean intensity at +a resolution is. What lives here is the LERF1 intensity built from those E +values, the multiplicity handling, and the symmetry detection. """ from __future__ import annotations @@ -17,7 +17,6 @@ import torch -from .french_wilson import french_wilson_preprocess # math_FrenchWilson.cc + Dfactor.cc from ..sh import ( get_high_order_axis, # phaser's highOrderAxis() compute_patterson_shell_variance, @@ -40,7 +39,6 @@ def eterm_sigma_a(s_mag: torch.Tensor, delta_vrms_A: float) -> torch.Tensor: __all__ = [ "eterm_sigma_a", - "french_wilson_preprocess", "get_high_order_axis", "build_lerf1_intensity", "apply_shell_variance_weights", diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 8e4d3b5a..5eb9bd32 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -52,8 +52,8 @@ from .frf.rotation_utils import rotation_matrix_from_edmonds_euler from .frf.types import RotationPeak from .rotation_search import search_peaks -from .sh import assign_shells, equal_count_shell_edges from .translation import ( + TranslationObs, TranslationPeak, amplitude_translation_search, fit_sigma_a_per_shell, @@ -225,6 +225,11 @@ def __init__( n_translation_candidates: int = 3, translation_grid_steps: int = 16, use_llg_tf: bool = False, + # Resolution window for the TRANSLATION set only, independent of the + # rotation search's [d_max, d_min]. None means no cut, which is what + # this stage has always done -- see `_prepare_translation_arrays`. + tf_d_min: Optional[float] = None, + tf_d_max: Optional[float] = None, # --- early stop --- min_tries: int = 3, max_tries: Optional[int] = None, @@ -255,6 +260,8 @@ def __init__( self.n_translation_candidates = n_translation_candidates self.translation_grid_steps = translation_grid_steps self.use_llg_tf = use_llg_tf + self.tf_d_min = tf_d_min + self.tf_d_max = tf_d_max self.min_tries = min_tries self.max_tries = max_tries @@ -263,8 +270,7 @@ def __init__( self._timer = _StageTimer(enabled=verbose >= 2) # Filled in by run(). self._frf = None - self._F_obs_amp = None - self._hkl_keep = None + self._obs = None self._tmask = None self._eye3 = torch.eye(3, dtype=torch.float64) @@ -462,7 +468,19 @@ def _make_rotated(self, peak: "RotationPeak"): # Stage 2: per-candidate translation search + local refine # ------------------------------------------------------------------ def _prepare_translation_arrays(self) -> None: - """Resolution/validity-masked obs amplitudes + Miller indices.""" + """Mask the observations for the translation search and normalise them once. + + The window is ``[tf_d_max, tf_d_min]`` on top of the dataset's own + validity mask. Both default to ``None``, meaning **no resolution cut** -- + which is what this stage has always done, though it used to claim + otherwise. So the translation search sees the data's full resolution + while the rotation search runs at ``[d_max, d_min]`` = [15, 4] A. That + asymmetry is deliberate on one side and unexamined on the other: the + rotation function is bandwidth-limited and cannot use high-resolution + terms, and nobody has measured what the translation function wants. The + parameter exists so that choosing is possible; the default does not + choose. + """ data = self.data device = self.device hkl_full = data.hkl @@ -473,9 +491,34 @@ def _prepare_translation_arrays(self) -> None: tmask = torch.ones( F_obs_full.shape[0], dtype=torch.bool, device=F_obs_full.device, ) + if self.tf_d_min is not None or self.tf_d_max is not None: + rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) + s_all = (hkl_full.to(torch.float64) @ rec_basis.to(hkl_full.device) + ).norm(dim=-1) + if self.tf_d_min is not None: + tmask = tmask & (s_all <= 1.0 / float(self.tf_d_min)) + if self.tf_d_max is not None: + tmask = tmask & (s_all >= 1.0 / float(self.tf_d_max)) self._tmask = tmask - self._F_obs_amp = F_obs_full[tmask].abs().to(torch.float64).to(device) - self._hkl_keep = hkl_full[tmask].to(device) + + sig_F_full = getattr(data, "F_sigma", None) + self._obs = TranslationObs.build( + F_obs_full[tmask], hkl_full[tmask], + data.spacegroup, data.cell, + sig_F=None if sig_F_full is None else sig_F_full[tmask], + delta_vrms_A=self.model_error_A, + n_shells=max(self.n_shells // 2, 8), + device=device, + ) + if self.verbose > 0: + d_hi = 1.0 / float(self._obs.s_mag.max()) + d_lo = 1.0 / float(self._obs.s_mag.min().clamp(min=1e-9)) + print( + f"mr: translation set {self._obs.F_obs.numel()} reflections, " + f"{d_lo:.1f}-{d_hi:.2f} A" + + ("" if sig_F_full is not None else " (no sigmas: unit weight)"), + flush=True, + ) def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: """Translation search + analytical-R local refine for one rotation. @@ -498,14 +541,13 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: timer.start("5_precompute_G") G_pre, h_R_pre = precompute_G_for_rotation( - evaluator, eye3, self._hkl_keep, data.spacegroup, data.cell, + evaluator, eye3, self._obs.hkl, data.spacegroup, data.cell, ) timer.stop("5_precompute_G") timer.start("6_amplitude_TF") _, _, t_peaks = amplitude_translation_search( - F_obs=self._F_obs_amp, interpolator=evaluator, - R_rotation=eye3, hkl=self._hkl_keep, + obs=self._obs, interpolator=evaluator, R_rotation=eye3, spacegroup=data.spacegroup, real_cell=data.cell, grid_steps=self.translation_grid_steps, n_peaks=self.n_translation_peaks, @@ -529,8 +571,7 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: t_init = torch.as_tensor(tp.translation, dtype=torch.float64) timer.start("7_local_TF_refine") t_refined, r_analytic = local_translation_refine( - F_obs=self._F_obs_amp, interpolator=evaluator, - R_rotation=eye3, hkl=self._hkl_keep, + obs=self._obs, interpolator=evaluator, R_rotation=eye3, spacegroup=data.spacegroup, real_cell=data.cell, t_init=t_init, radius=0.06, grid_steps=13, n_refinement_passes=1, @@ -550,34 +591,24 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): """Re-rank translation peaks by a shared-σA Rice/Woolfson LLG. - Mirrors Phaser's FTF — the cheap amplitude correlation is a fast - pre-filter but ranks poorly for partial models; the LLG ranks - consistently with the rotation rescore. + Mirrors Phaser's FTF: the amplitude correlation is a cheap pre-filter, + and this is the likelihood that ranks its peaks. It reuses the run's + single Wilson normalisation and its shell binning, so ``E_obs`` here is + the same ``E_obs`` the correlation maximised. + + Off by default. It is the strongest discriminator at rank level, and + end-to-end it changes nothing: 27/30 against 28/30 with one discordant + cell in 30, which the correlation wins. """ - data = self.data device = self.device - F_obs_amp = self._F_obs_amp - hkl_keep = self._hkl_keep - tmask = self._tmask + obs = self._obs self._timer.start("6b_llg_tf_rescore") - rec_basis_keep = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) - s_mag_keep_tf = (hkl_keep.to(torch.float64) @ rec_basis_keep).norm(dim=-1) - tf_n_shells = max(self.n_shells // 2, 8) - tf_edges, _ = equal_count_shell_edges(s_mag_keep_tf, tf_n_shells) - tf_shell_idx = assign_shells(s_mag_keep_tf, tf_edges) - centric_keep_tf = ( - data.centric[tmask].to(torch.bool).to(device) - if hasattr(data, "centric") - else torch.zeros_like(F_obs_amp, dtype=torch.bool) - ) - - cnt_tf = torch.bincount(tf_shell_idx, minlength=tf_n_shells).to(torch.float64) - sum_F2 = torch.zeros(tf_n_shells, dtype=torch.float64, device=device) - sum_F2.scatter_add_(0, tf_shell_idx, F_obs_amp * F_obs_amp) - mean_F2 = (sum_F2 / cnt_tf.clamp(min=1.0)).clamp(min=1e-30) - E_obs_tf = F_obs_amp / mean_F2.sqrt().index_select(0, tf_shell_idx) - + # sigma_A is fitted against the top translation only, and reused for + # every candidate. It is a per-shell model-reliability curve, not a + # per-candidate score: refitting it per t would let each candidate + # choose the D that flatters it, which is scoring a model against a + # likelihood tuned to that model. t_top_t = torch.as_tensor( t_peaks[0].translation, dtype=torch.float64, device=device, ) @@ -587,13 +618,16 @@ def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): ).to(G_pre.dtype), ) Fc_top = (G_pre * phase_top).sum(dim=0).abs().to(torch.float64) - sum_Fc2 = torch.zeros(tf_n_shells, dtype=torch.float64, device=device) - sum_Fc2.scatter_add_(0, tf_shell_idx, Fc_top * Fc_top) + cnt_tf = torch.bincount( + obs.shell_idx, minlength=obs.n_shells, + ).to(torch.float64) + sum_Fc2 = torch.zeros(obs.n_shells, dtype=torch.float64, device=device) + sum_Fc2.scatter_add_(0, obs.shell_idx, Fc_top * Fc_top) mean_Fc2 = (sum_Fc2 / cnt_tf.clamp(min=1.0)).clamp(min=1e-30) - E_calc_top = Fc_top / mean_Fc2.sqrt().index_select(0, tf_shell_idx) + E_calc_top = Fc_top / mean_Fc2.sqrt().index_select(0, obs.shell_idx) sigma_a_tf = fit_sigma_a_per_shell( - E_obs_tf, E_calc_top, centric_keep_tf, - tf_shell_idx, tf_n_shells, n_grid=81, + obs.E_obs, E_calc_top, obs.centric, + obs.shell_idx, obs.n_shells, n_grid=81, ) t_cands = torch.as_tensor( @@ -601,9 +635,7 @@ def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): dtype=torch.float64, device=device, ) llg_tf = llg_translation_rescore( - F_obs=F_obs_amp, hkl=hkl_keep, centric=centric_keep_tf, - s_mag=s_mag_keep_tf, shell_idx=tf_shell_idx, n_shells=tf_n_shells, - G=G_pre, h_R=h_R_pre, t_candidates=t_cands, + obs=obs, G=G_pre, h_R=h_R_pre, t_candidates=t_cands, sigma_a=sigma_a_tf, interp_var=None, ) self._timer.stop("6b_llg_tf_rescore") diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index 605bf291..feae55a5 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -30,7 +30,6 @@ from torchref.scaling.weighting import (DEFAULT_SNR_CAP, DEFAULT_TRUST_CAP) -from .e_values import SmoothSigmaE from .sh import ( apply_overall_anisotropy, assign_shells, @@ -223,7 +222,6 @@ def search_peaks( n_peaks: int, verbose: int = 0, device: Optional[torch.device] = None, - e_convention: type = SmoothSigmaE, obs_weight: str = "inverse_variance", sigma_a_source: str = "empirical", apply_bulk_solvent: bool = False, @@ -373,7 +371,6 @@ def search_peaks( grid_sampling_deg=GRID_SAMPLING_DEG, asu_idx=asu_idx, s_mag_asu=s_mag_asu, - e_convention=e_convention, obs_weight=obs_weight, snr_cap=snr_cap, trust_cap=trust_cap, shell_variance_weights=shell_variance_weights, ) @@ -422,7 +419,6 @@ def rotation_search( n_peaks: int = 500, verbose: int = 0, device: Optional[torch.device] = None, - e_convention: type = SmoothSigmaE, ) -> RotationSolutions: """Find the orientations of ``model`` consistent with ``data``. @@ -451,16 +447,6 @@ def rotation_search( Where to run. Default ``None`` takes ``data``'s device, moving ``model`` to match; an explicit value moves both. With neither carrying one, the configured default applies. - e_convention : type, optional - How amplitudes become E values, given as a class rather than an - instance: a fitted ``Sigma(s)`` cannot exist before the reflections do, - so the engine constructs it -- once for the observations and once for - the model, which is what puts the two on a common footing. The default - fits a smooth ``Sigma(s)`` -- a Gamma GLM on a Chebyshev basis in - sin(theta)/lambda -- independently for each side, so `` = 1`` - holds on both as an identity of the fit. ``functools.partial`` - configures one. - Returns ------- RotationSolutions @@ -480,6 +466,5 @@ def rotation_search( peaks, lmax, d_min = search_peaks( model, data, model_error_A, U_aniso=U_aniso, n_peaks=n_peaks, verbose=verbose, device=device, - e_convention=e_convention, ) return _solutions(peaks, lmax, d_min, model_error_A) diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index fae0b320..9b0ceed6 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -1,22 +1,168 @@ -""" -Fast FFT-based translation search for molecular replacement. - -Translation t shifts phase: F(hkl, t) = F(hkl) * exp(2*pi*i * hkl.t) -Correlation: C(t) = IFFT{ conj(F_obs) * F_calc } - -This module provides efficient FFT-based translation search that finds the -optimal translation to position a model after rotation has been determined. +"""Fast translation search: where in the cell does an oriented model sit? + +A translation shifts phase, ``F(h, t) = F(h) exp(2 pi i h.t)``, so scoring every +``t`` on a grid is a Fourier transform rather than a scan. The Crowther-Blow +form used here accumulates the pair coefficients +``sum_h w(h) G_i*(h) G_j(h)`` onto a reciprocal grid at +``(h R_j - h R_i) mod G`` and takes one inverse FFT, which replaces ``G^3`` grid +evaluations with a single transform. + +This is where the discrimination happens. The rotation function upstream is a +shortlist generator -- over 30 seeded cells it puts truth at rank 0 six times; +the correlation here does it 24 times and the likelihood 27. Rotation ghosts are +morphologically identical to truth in a Patterson by construction, and stop +being identical as soon as the crystal lattice is involved. + +The observed side is prepared **once** per run, by +:class:`TranslationObs`, and reused for every orientation and every candidate +translation. That is not only an optimisation: normalisation and weighting are +properties of the observations, which do not change when the model moves, and +three separate answers to "what is the mean intensity here" used to live in this +module and its caller. """ import numpy as np import torch +from torchref.scaling import WilsonNormaliser +from torchref.scaling.weighting import (inverse_variance_weight, + normalise_weight, snr_from_amplitude) + from .distributions import rice_log_likelihood, woolfson_log_likelihood -from .e_values import WilsonShellE +from .sh import assign_shells, equal_count_shell_edges from dataclasses import dataclass from typing import List, Optional, Tuple +#: Chebyshev order of the Wilson fit. Matches the rotation function's +#: ``frf.api.WILSON_N_COEFF``: the two stages score the same observations and a +#: different order on each would be two normalisations again. +WILSON_N_COEFF = 6 + + +@dataclass +class TranslationObs: + """The observed side of a translation search, normalised and weighted once. + + Everything here is a property of the observations alone, so none of it + changes when the model rotates or moves. Building it per orientation -- which + is what the module used to do -- refits a Gamma GLM for every candidate to + get the same answer back, and worse, it made "what is ``E_obs``" a question + with three different answers depending on which function you asked. + + Attributes + ---------- + F_obs, hkl, s_mag, centric, eps + The masked observations and their crystallographic bookkeeping. + E_obs : torch.Tensor + ``F / sqrt(eps Sigma(s))``, with ``Sigma`` the shared Wilson fit, so + `` = 1`` as an identity of that fit rather than as a separate + normalisation step. + weight : torch.Tensor + Mean-1 inverse-variance weight, from measurement error and model error + in one denominator. **This is the half that does not cancel.** A + per-resolution *scaling* is gauge in a correlation -- twelve conventions + moved the rotation function's truth rank by nothing -- but a weight that + varies within a shell is not, and until now the translation search had + none at all: every reflection counted the same. + shell_idx, n_shells + Equal-count binning in ``|s|``, shared by the sigma_A fit and the + likelihood so the two cannot disagree about which reflection is where. + fit : WilsonNormaliser + Kept, not discarded. Anything comparing an observed curve against a + calculated one needs the curve itself, not the per-reflection values it + produced. + """ + + F_obs: torch.Tensor + hkl: torch.Tensor + s_mag: torch.Tensor + centric: torch.Tensor + eps: torch.Tensor + E_obs: torch.Tensor + weight: torch.Tensor + shell_idx: torch.Tensor + n_shells: int + fit: "WilsonNormaliser" + + @classmethod + def build( + cls, + F_obs: torch.Tensor, + hkl: torch.Tensor, + spacegroup, + real_cell, + *, + sig_F: Optional[torch.Tensor] = None, + delta_vrms_A: float = 1.0, + n_shells: int = 20, + n_coeff: int = WILSON_N_COEFF, + device=None, + ) -> "TranslationObs": + """Normalise and weight one set of observations. + + Parameters + ---------- + F_obs : torch.Tensor + ``(N,)`` observed amplitudes; complex input is coerced to ``|.|``. + hkl : torch.Tensor + ``(N, 3)`` integer Miller indices, matching ``F_obs`` row for row. + spacegroup, real_cell + Supply multiplicity, centricity and the reciprocal basis. + sig_F : torch.Tensor, optional + ``(N,)`` measurement errors. Without them the weight is uniform, + which is the honest fallback: the varying part of the weight *is* + the measurement term, and inventing one would be worse than not + having it. + delta_vrms_A : float + R.m.s. coordinate error of the search model, which sets the model + half of the variance budget through the Luzzati falloff. The same + number the rotation function weights with. + """ + dev = device if device is not None else F_obs.device + real = torch.float64 + F = F_obs.detach().to(dev) + F = (F.abs() if F.is_complex() else F).to(real) + hkl_i = hkl.detach().to(dev) + + rec_basis = real_cell.reciprocal_basis_matrix.to(dev).to(real) + s_mag = (hkl_i.to(real) @ rec_basis).norm(dim=-1) + + hkl_l = hkl_i.round().to(torch.int64) + # friedel=False: Wilson's = eps*Sigma counts the operations mapping + # h to itself, which add coherently and set the mean. The Friedel-folded + # branch changes the distribution instead, and that is centricity -- + # which enters separately, as the Gamma shape. + eps = spacegroup.epsilon(hkl_l, friedel=False).to(real).clamp(min=1.0) + centric = spacegroup.is_centric(hkl_l).to(torch.bool) + + fit = WilsonNormaliser( + F * F, s_mag, eps=eps, centric=centric, n_coeff=n_coeff, + ) + + if sig_F is None: + weight = torch.ones_like(F) + else: + sig = sig_F.detach().to(dev).to(real).abs() + # eterm_sigma_a is the rotation function's own model-error term; + # importing it rather than restating the exponent is the point. + from .frf.preprocessing import eterm_sigma_a + weight = normalise_weight(inverse_variance_weight( + snr_from_amplitude(F, sig), + eterm_sigma_a(s_mag, float(delta_vrms_A)).to(real), + eps=eps, + )) + + edges, _ = equal_count_shell_edges(s_mag, n_shells) + shell_idx = assign_shells(s_mag, edges).clamp(min=0) + + return cls( + F_obs=F, hkl=hkl_i, s_mag=s_mag, centric=centric, eps=eps, + E_obs=fit.E.to(real), weight=weight, + shell_idx=shell_idx, n_shells=int(n_shells), fit=fit, + ) + + @dataclass class TranslationPeak: """ @@ -97,18 +243,15 @@ def find_translation_peaks( def amplitude_translation_search( - F_obs: torch.Tensor, + obs: TranslationObs, interpolator, R_rotation: torch.Tensor, - hkl: torch.Tensor, spacegroup, real_cell, grid_steps: int = 16, n_peaks: int = 20, cluster_radius: float = 0.05, batch_size: int = 256, - use_e_values: bool = True, - n_shells: int = 20, precomputed_G: Optional[torch.Tensor] = None, precomputed_h_R: Optional[torch.Tensor] = None, ) -> Tuple[np.ndarray, np.ndarray, List[TranslationPeak]]: @@ -130,15 +273,13 @@ def amplitude_translation_search( Parameters ---------- - F_obs : torch.Tensor, shape (N,) - Observed amplitudes (complex inputs are coerced to |·|). + obs : TranslationObs + The observations, normalised and weighted once for the whole run. interpolator : object Anything providing ``evaluate(R, hkl, real_cell, return_amplitude=False)`` -- in the pipeline, ``align._DirectModelEvaluator``. R_rotation : torch.Tensor, shape (3, 3) Rotation that has been applied to the model coordinates. - hkl : torch.Tensor, shape (N, 3) - Integer Miller indices of the observed reflections. spacegroup : SpaceGroup Provides `matrices` and `translations`. real_cell : Cell @@ -151,15 +292,6 @@ def amplitude_translation_search( Minimum fractional separation between returned peaks. batch_size : int, default 256 Number of candidate translations evaluated per inner batch. - use_e_values : bool, default True - Normalize `|F_obs|` and `|F_calc(h, t)|` per resolution shell to unit - Wilson variance (E-values) before correlating. This removes the - resolution-dependent envelope mismatch between real F_obs (with bulk - solvent + thermal falloff) and a model that doesn't model these — a - per-shell mean subtraction in the Pearson correlation alone doesn't - cover it because the falloff is multiplicative, not additive. - n_shells : int, default 20 - Number of equal-count radial shells used by `use_e_values`. Returns ------- @@ -170,44 +302,22 @@ def amplitude_translation_search( peaks : list of TranslationPeak Top-`n_peaks` peaks sorted by descending correlation. """ - device = getattr(interpolator, "device", hkl.device) + device = getattr(interpolator, "device", obs.hkl.device) real_dtype = torch.float64 complex_dtype = torch.complex128 - F_obs_t = F_obs.detach().to(device) - if F_obs_t.is_complex(): - F_obs_t = F_obs_t.abs() - F_obs_t = F_obs_t.to(real_dtype) - - hkl_t = hkl.detach().to(device).to(real_dtype) # (N, 3) - - # Precompute per-shell normalisation if requested. We bin reflections by - # |s| into n_shells equal-count shells and normalise F → F / sqrt(_shell) - # (Wilson E-value). The same shell norm is applied to F_calc(t) inside the - # batch loop. This makes the Pearson correlation a Patterson-style - # correlation of "E²−1" — robust to bulk-solvent / B-factor mismatch. - if use_e_values: - rec_basis_real = real_cell.reciprocal_basis_matrix.to(device).to(real_dtype) - s_mag = (hkl_t @ rec_basis_real).norm(dim=-1) - order = torch.argsort(s_mag) - shell_idx = torch.zeros_like(s_mag, dtype=torch.int64) - chunk = s_mag.numel() // max(n_shells, 1) - for k in range(n_shells): - a = k * chunk - b = (k + 1) * chunk if k < n_shells - 1 else s_mag.numel() - shell_idx[order[a:b]] = k - # The caller's own shell assignment is handed to the convention rather - # than letting it derive one: this binning is rank-based and the shared - # `assign_shells` is value-based, and the sigma_a fit downstream is tied - # to whichever one was used here. - E_obs = WilsonShellE( - F_obs_t, s_mag, shell_idx=shell_idx, n_shells=n_shells, - ).E - F_obs2 = E_obs * E_obs - else: - shell_idx = None - F_obs2 = F_obs_t * F_obs_t - F_obs2_centered = F_obs2 - F_obs2.mean() + hkl = obs.hkl + E_obs = obs.E_obs.to(device).to(real_dtype) + w = obs.weight.to(device).to(real_dtype) + + # Correlating E^2 rather than F^2 is what makes this robust to the + # resolution envelope: the model has no bulk solvent and the wrong overall + # B, and that mismatch is multiplicative, so subtracting a mean does not + # remove it but dividing by Sigma(s) does. + F_obs2 = E_obs * E_obs + # Centred at the WEIGHTED mean, which is what the weighted correlation + # below is a numerator for. With uniform weight this is the plain mean. + F_obs2_centered = F_obs2 - (w * F_obs2).sum() / w.sum().clamp(min=1e-30) # Pre-compute G_i(h) = exp(2πi h·t_i) · F_p1(h R_i) (or reuse caller's) two_pi_i = 2j * torch.pi @@ -230,8 +340,8 @@ def amplitude_translation_search( # reciprocal grid at integer indices (h·R_j − h·R_i) mod G. # # We accumulate two such reciprocal grids in one sym-op pass: - # W_num : weight per h = F_obs²_centered(h) → num(t) - # W_den : weight per h = 1 → Σ_h |F_calc(h,t)|² + # W_num : weight per h = w(h)·E_obs²_centered(h) → num(t) + # W_den : weight per h = w(h) → Σ_h w|F_calc(h,t)|² # Score(t) = num(t) / Σ_h|F_calc(h,t)|² — a per-t scale-normalised # Pearson proxy (Phaser's TF uses the full Pearson denominator; ours # uses the same scaling that the previous separable-phase code applied @@ -243,8 +353,11 @@ def amplitude_translation_search( # and orders of magnitude less than the original explicit grid loop. S_eff, N_eff = G.shape h_R_int = h_R.round().to(torch.int64) # (S, N, 3) - F_obs2_c_complex = F_obs2_centered.to(complex_dtype) # (N,) - ones_complex = torch.ones(N_eff, dtype=complex_dtype, device=device) + # Both grids carry the same per-reflection weight, so the ratio below is a + # weighted correlation rather than an unweighted one with a weighted + # numerator. Uniform w reproduces the previous scores exactly. + F_obs2_c_complex = (w * F_obs2_centered).to(complex_dtype) # (N,) + w_complex = w.to(complex_dtype) # (N,) W_num_flat = torch.zeros( grid_steps ** 3, dtype=complex_dtype, device=device, @@ -257,7 +370,7 @@ def amplitude_translation_search( Gi_conj = G[i].conj() # (N,) pair = Gi_conj.view(1, -1) * G # (S, N) coeff_num = F_obs2_c_complex.view(1, -1) * pair # (S, N) - coeff_den = ones_complex.view(1, -1) * pair # (S, N) + coeff_den = w_complex.view(1, -1) * pair # (S, N) dh = (h_R_int - h_R_int[i:i + 1]) % grid_steps # (S, N, 3) flat = (dh[..., 0] * G_stride_xy + dh[..., 1] * grid_steps + dh[..., 2]) # (S, N) @@ -333,50 +446,53 @@ def fit_sigma_a_per_shell( def llg_translation_rescore( - F_obs: torch.Tensor, - hkl: torch.Tensor, - centric: torch.Tensor, - s_mag: torch.Tensor, - shell_idx: torch.Tensor, - n_shells: int, + obs: TranslationObs, G: torch.Tensor, h_R: torch.Tensor, t_candidates: torch.Tensor, sigma_a: torch.Tensor, interp_var: Optional[torch.Tensor] = None, ) -> torch.Tensor: - """ - Per-translation Rice / Woolfson log-likelihood, using the symmetry-summed - interpolator contributions ``G`` (Phaser EM_search analogue) and a fixed - per-shell σA. + """Per-translation Rice / Woolfson log-likelihood over candidate positions. - For each candidate t: - F_calc(h, t) = Σ_i G_i(h) · exp(2πi (h R_i) · t) - E_calc(h, t) = |F_calc(h, t)| / sqrt(_per_shell) - LLG(t) = Σ_shell [LL_Rice(E_obs, D·E_calc, var) − LL_Wilson(E_obs)] - where var = (1 − D²) + interp_var. + For each candidate t:: - Phase B alignment likelihood-TF. Replaces the |F|² Pearson correlation - in `amplitude_translation_search` as the scoring rule when the caller - re-ranks the FFT-cheap pre-filter peaks. + F_calc(h, t) = sum_i G_i(h) exp(2 pi i (h R_i).t) + E_calc(h, t) = |F_calc(h, t)| / sqrt(_shell) + LLG(t) = sum_h [LL(E_obs, D E_calc, var) - LL_Wilson(E_obs)] + + with ``var = (1 - D^2) + interp_var``. The Rice branch is used for acentric + reflections and Woolfson for centric. + + The scoring rule the amplitude correlation is a pre-filter for. At rank + level it is the strongest discriminator measured -- truth at rank 0 in 27 of + 30 seeded cells against the correlation's 24 and the rotation function's 6 -- + but re-ranking translation peaks by it does **not** improve end-to-end pose + recovery (28/30 against 27/30 the other way, one discordant cell), which is + why ``use_llg_tf`` defaults off. + + ``E_calc`` is normalised per shell **per candidate**, and that is not a + fourth answer to what the observations' normalisation is: it is a per-``t`` + scale, and forcing it to unit shell variance for every candidate is what + makes the K likelihoods comparable. What discriminates is the pattern across + reflections, not the scale. Parameters ---------- - F_obs : (N,) real - hkl : (N, 3) — unused here but kept for symmetry with the rest of the - module (and future extension to per-h variance models). - centric : (N,) bool - s_mag : (N,) — |s| the shells were built from. Not used to derive a binning - here (``shell_idx`` is given) but passed rather than fabricated, so a - convention that fits a curve in ``|s|`` gets the real abscissa. - shell_idx : (N,) int64 — same binning as used to fit sigma_a / interp_var. - n_shells : int - G : (S, N) complex — per-sym F_p1 contributions × per-sym translation phase - (output of `precompute_G_for_rotation`). - h_R : (S, N, 3) — per-sym rotated reciprocal indices. - t_candidates : (K, 3) fractional translations to score. - sigma_a : (n_shells,) — fixed per-shell σA (shared across candidates). - interp_var : (N,) optional — per-reflection variance inflation. + obs : TranslationObs + Supplies ``E_obs``, centricity and the shell binning -- the same binning + ``sigma_a`` was fitted on, which is why it is not re-derived here. + G : (S, N) complex + Per-sym ``F_p1`` contributions x per-sym translation phase, from + :func:`precompute_G_for_rotation`. + h_R : (S, N, 3) + Per-sym rotated reciprocal indices. + t_candidates : (K, 3) + Fractional translations to score. + sigma_a : (n_shells,) + Fixed per-shell sigma_A, shared across candidates. + interp_var : (N,), optional + Per-reflection variance inflation. Returns ------- @@ -386,6 +502,10 @@ def llg_translation_rescore( real_dtype = torch.float64 complex_dtype = G.dtype + shell_idx = obs.shell_idx + n_shells = obs.n_shells + centric = obs.centric + K = t_candidates.shape[0] S, N = G.shape @@ -408,11 +528,7 @@ def llg_translation_rescore( norm_per_refl = mean_per_shell.sqrt().gather(1, shell_idx_k) # (K, N) E_calc = F_calc / norm_per_refl # (K, N) - F_obs_t = F_obs.to(device).to(real_dtype) - E_obs = WilsonShellE( - F_obs_t, s_mag.to(device).to(real_dtype), - shell_idx=shell_idx_l, n_shells=n_shells, - ).E + E_obs = obs.E_obs.to(device).to(real_dtype) sigma_a_d = sigma_a.to(device).to(real_dtype) # (n_shells,) D_per_refl = sigma_a_d.index_select(0, shell_idx_l) # (N,) @@ -497,10 +613,9 @@ def precompute_G_for_rotation( def local_translation_refine( - F_obs: torch.Tensor, + obs: TranslationObs, interpolator, R_rotation: torch.Tensor, - hkl: torch.Tensor, spacegroup, real_cell, t_init: torch.Tensor, @@ -514,33 +629,37 @@ def local_translation_refine( """ Fine-grid Patterson translation refinement around ``t_init``. - For each candidate ``t`` in a `grid_steps`³ cubic grid of half-width - ``radius`` centered on ``t_init`` (fractional), computes |F_calc(h, t)|² - via the symmetry expansion and the analytical-scale R-factor - R(t) = Σ ||F_obs| − k·|F_calc(t)|| / Σ |F_obs| - k(t) = Σ |F_obs|·|F_calc(t)| / Σ |F_calc(t)|² - against `F_obs`. Returns the (t, R) at the minimum. - - Use `n_refinement_passes > 1` to do a multi-pass zoom: each pass shrinks - the radius by `grid_steps/2` and re-centers on the previous best. For - `radius=0.06, grid_steps=13, n_refinement_passes=2`, the final fractional - resolution is ~0.005 (≈0.3 Å for a 60 Å cell). - - The analytical-scale R-factor uses a single global scale; it is not the - same number a full crystallographic Scaler would return, but its - *minimum location* is robust because both numerator and denominator share - the same per-shell envelope. Use a full Scaler to compute the final - R-work after this routine selects (R, t). + Locates the peak on a fine grid of half-width ``radius`` around ``t_init`` + by the same weighted ``E^2`` correlation the coarse search maximises, then + reports the analytical-scale R-factor there:: + + R(t) = sum ||F_obs| - k|F_calc(t)|| / sum |F_obs| + k(t) = sum |F_obs||F_calc(t)| / sum |F_calc(t)|^2 + + The two halves answer different questions and use different quantities on + purpose. The *search* runs on normalised, weighted ``E^2``, because that is + what the coarse stage optimised and refining against a different objective + would walk away from the peak it was handed. The *reported number* is an + R-factor on raw amplitudes, because that is what ranks candidates and what a + crystallographer reads. + + That R uses one global scale, so it is not the number a full Scaler returns. + It is used as a ranking key, and for that its minimum's *location* is what + matters. The winner gets a solvent-aware Scaler refit. + + Returns ``(t, R)`` at the minimum. ``n_refinement_passes`` is accepted and + ignored -- the FFT evaluates the whole fine grid at once, so there is + nothing for a second zoom pass to buy. """ - device = getattr(interpolator, "device", hkl.device) + device = getattr(interpolator, "device", obs.hkl.device) real_dtype = torch.float64 complex_dtype = torch.complex128 - F_obs_t = F_obs.detach().to(device) - if F_obs_t.is_complex(): - F_obs_t = F_obs_t.abs() - F_obs_t = F_obs_t.to(real_dtype) + hkl = obs.hkl + F_obs_t = obs.F_obs.to(device).to(real_dtype) F_obs_sum = F_obs_t.sum().clamp(min=1e-30) + E_obs = obs.E_obs.to(device).to(real_dtype) + w = obs.weight.to(device).to(real_dtype) two_pi_i = 2j * torch.pi if precomputed_G is not None and precomputed_h_R is not None: @@ -582,10 +701,13 @@ def local_translation_refine( G_fft = min(G_fft, 128) half_window = max(1, int(round(float(radius) * G_fft))) - F_obs2 = (F_obs_t * F_obs_t).to(real_dtype) - F_obs2_centered = F_obs2 - F_obs2.mean() - F_obs2_c_complex = F_obs2_centered.to(complex_dtype) - ones_complex = torch.ones(N, dtype=complex_dtype, device=device) + # The same weighted E^2 correlation as the coarse search. It used to be a + # raw |F|^2 correlation here, which made the fine grid optimise a different + # objective from the one that chose the peak it is centred on. + F_obs2 = E_obs * E_obs + F_obs2_centered = F_obs2 - (w * F_obs2).sum() / w.sum().clamp(min=1e-30) + F_obs2_c_complex = (w * F_obs2_centered).to(complex_dtype) + w_complex = w.to(complex_dtype) W_num_flat = torch.zeros(G_fft ** 3, dtype=complex_dtype, device=device) W_den_flat = torch.zeros(G_fft ** 3, dtype=complex_dtype, device=device) @@ -594,7 +716,7 @@ def local_translation_refine( Gi_conj = G_shifted[i].conj() pair = Gi_conj.view(1, -1) * G_shifted # (S, N) coeff_num = F_obs2_c_complex.view(1, -1) * pair - coeff_den = ones_complex.view(1, -1) * pair + coeff_den = w_complex.view(1, -1) * pair dh = (h_R_int - h_R_int[i:i + 1]) % G_fft # (S, N, 3) flat = (dh[..., 0] * G_stride_xy + dh[..., 1] * G_fft + dh[..., 2]) # (S, N) From c1a3233a77711c319fd3e02936aadf3b595fc15e Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 31 Aug 2026 16:18:16 +0200 Subject: [PATCH 113/250] Put each stage in one module, and take the device from config align.py held four things that belonged to three different stages, and it and pipeline.py imported each other -- pipeline for the stage helpers, align lazily for the pipeline class. Each piece now sits with the stage it serves: DirectModelEvaluator with the translation search whose structure factors it supplies, FRFInputs/prepare_frf_inputs with the rotation search they prepare, and the timer, the R-work and align_model_to_data with the orchestrator. prepare_frf_inputs now calls fit_anisotropy instead of repeating it. The two fitted the same U over the same window -- ANISO_FIT_WINDOW_A is (15, 4) and the pipeline's defaults are d_max=15, d_min=4 -- by two copies of the same six lines. **Device comes from torchref.config, not from whichever tensor is nearest.** _llg_tf_rescore used self.device while precompute_G_for_rotation read it off the interpolator, and the two differ the moment a CPU model meets an accelerator default -- which is what the discrimination diagnostic hit. Following the model's device would propagate one of the two answers; taking config removes the second. DirectModelEvaluator now answers on the configured device whatever device the model sits on, and the Scaler in _external_rwork keeps its own default instead of being overridden. sh.py loses the spherical-harmonic expansion it no longer has a caller for -- evaluate_ylm, sh_expand_ball, angular_density_weights. _bar_legendre_recurrence stays: production does not call it either, but it is the independent pure-torch implementation that test_bessel_sh_grouping builds its slow reference expansion on, and that reference is the only check of the fused expansion that is not the fused expansion. Its scipy pin used to run through evaluate_ylm, so it is now asserted directly. The lab loses the rescore and E-convention harnesses and the module-level re-exports that made every script import them transitively; the FRF and FTF diagnostics move onto the new signatures. 1940 passed, 0 failed. Pose recovery on 1DAW t0 unchanged at 1.519 deg. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/e_conformance.py | 247 ------------- alignment_lab/analysis/e_convention_arms.py | 140 -------- alignment_lab/analysis/e_convention_arms.sh | 25 -- alignment_lab/analysis/e_table.py | 78 ----- alignment_lab/analysis/e_table.sh | 28 -- alignment_lab/analysis/full_gate.sh | 5 +- alignment_lab/analysis/fw_asu_equivalence.sh | 58 ---- alignment_lab/analysis/fw_footing.py | 95 ----- alignment_lab/analysis/fw_footing.sh | 17 - alignment_lab/analysis/fw_internals.sh | 93 ----- alignment_lab/analysis/llg_decompose.py | 162 --------- alignment_lab/analysis/llg_decompose.sh | 24 -- alignment_lab/analysis/pose_arms.sh | 14 +- alignment_lab/analysis/rescore_prep_arms.py | 149 -------- alignment_lab/analysis/rescore_prep_arms.sh | 22 -- alignment_lab/analysis/seam_gate.sh | 24 -- alignment_lab/analysis/seam_identity.py | 112 ------ .../diagnostics/frf_vs_ftf_discrimination.py | 29 +- alignment_lab/diagnostics/pose_recovery.py | 2 +- alignment_lab/diagnostics/rescore_rank.py | 118 ------- alignment_lab/diagnostics/tf_batch_probe.py | 152 -------- alignment_lab/diagnostics/tf_cost.py | 17 +- alignment_lab/lab/__init__.py | 10 +- alignment_lab/lab/aniso.py | 13 +- alignment_lab/lab/frf.py | 28 +- alignment_lab/lab/reference_normalisers.py | 53 --- alignment_lab/lab/rescore.py | 202 ----------- alignment_lab/tests/test_lab.py | 43 --- tests/helpers/device_cases.py | 1 - tests/unit/alignment/test_sh.py | 256 ++++++-------- torchref/experimental/alignment/__init__.py | 130 ++++--- torchref/experimental/alignment/align.py | 324 ------------------ torchref/experimental/alignment/pipeline.py | 166 ++++++++- .../experimental/alignment/rotation_search.py | 85 ++++- torchref/experimental/alignment/sh.py | 306 ++--------------- .../experimental/alignment/translation.py | 41 ++- 36 files changed, 523 insertions(+), 2746 deletions(-) delete mode 100644 alignment_lab/analysis/e_conformance.py delete mode 100644 alignment_lab/analysis/e_convention_arms.py delete mode 100644 alignment_lab/analysis/e_convention_arms.sh delete mode 100644 alignment_lab/analysis/e_table.py delete mode 100644 alignment_lab/analysis/e_table.sh delete mode 100644 alignment_lab/analysis/fw_asu_equivalence.sh delete mode 100644 alignment_lab/analysis/fw_footing.py delete mode 100644 alignment_lab/analysis/fw_footing.sh delete mode 100644 alignment_lab/analysis/fw_internals.sh delete mode 100644 alignment_lab/analysis/llg_decompose.py delete mode 100644 alignment_lab/analysis/llg_decompose.sh delete mode 100644 alignment_lab/analysis/rescore_prep_arms.py delete mode 100644 alignment_lab/analysis/rescore_prep_arms.sh delete mode 100644 alignment_lab/analysis/seam_gate.sh delete mode 100644 alignment_lab/analysis/seam_identity.py delete mode 100644 alignment_lab/diagnostics/rescore_rank.py delete mode 100644 alignment_lab/diagnostics/tf_batch_probe.py delete mode 100644 alignment_lab/lab/reference_normalisers.py delete mode 100644 alignment_lab/lab/rescore.py delete mode 100644 torchref/experimental/alignment/align.py diff --git a/alignment_lab/analysis/e_conformance.py b/alignment_lab/analysis/e_conformance.py deleted file mode 100644 index a985fbcd..00000000 --- a/alignment_lab/analysis/e_conformance.py +++ /dev/null @@ -1,247 +0,0 @@ -"""Does an E convention do what we want E to do? - -Phaser is no longer the specification for this part of the code, which means -comparison-debugging against a reference implementation is gone. What replaces it -is invariants: Wilson statistics, scale invariance, shrinkage monotonicity and -epsilon-correctness are true regardless of whose code computes them. This module -is that safety net, and it exists so conventions can be changed with something -other than an argument deciding whether the change was right. - -Eight checks, of which two are the ones nothing in the tree currently makes: - -* **absolute** unit mean, not merely a flat trend -- the rotation function is a - correlation and does not care, but the rescore's LLG compares an observation - against a predicted distribution and there is no free scale to cancel; -* **obs and calc on a common footing** -- the measured symptom of getting this - wrong is an expected moving-model intensity with mean 2.14 where ~1 belongs. - -The strongest check is not the mean but the **distribution**. Wilson statistics -predict ``|E|**2 ~ Exp(1)`` acentric and ``~ chi2_1`` centric at *every* -resolution, which catches a normaliser that is right on average and wrong in -shape. A mean cannot. - -Deliberately a report rather than a gate: a convention may fail a property and -still rank truth better, and in that case the property tells us what the winner -is trading away rather than vetoing it. -""" - -from __future__ import annotations - -import math -from typing import Optional - -import torch - -#: Wilson moment ratios /^2. Departures upward are the standard -#: twinning / tNCS diagnostic -- 2DQ6 reads about 5.5 acentric, which is how its -#: tNCS was originally identified. -IDEAL_MOMENT_RATIO = {"acentric": 2.0, "centric": 3.0} - - -def _deciles(s_mag: torch.Tensor, n: int = 10) -> torch.Tensor: - """Equal-count resolution deciles, as an index per reflection.""" - order = torch.argsort(s_mag) - out = torch.zeros_like(s_mag, dtype=torch.long) - chunk = max(1, s_mag.numel() // n) - for k in range(n): - lo = k * chunk - hi = (k + 1) * chunk if k < n - 1 else s_mag.numel() - out[order[lo:hi]] = k - return out - - -def _ks_uniform(sorted_u: torch.Tensor) -> float: - """One-sample KS statistic of ``sorted_u`` against Uniform(0, 1).""" - n = sorted_u.numel() - if n < 2: - return float("nan") - i = torch.arange(1, n + 1, dtype=torch.float64, device=sorted_u.device) - d_plus = (i / n - sorted_u).max() - d_minus = (sorted_u - (i - 1) / n).max() - return float(torch.maximum(d_plus, d_minus)) - - -def _wilson_ks(E2: torch.Tensor, centric: torch.Tensor) -> dict: - """KS of E**2 against its Wilson distribution, by centric class. - - Acentric ``E**2 ~ Exp(1)``, so ``1 - exp(-E**2)`` is uniform. Centric - ``E**2 ~ chi2_1``, so ``erf(sqrt(E**2 / 2))`` is uniform. Both CDFs are - closed-form, so no scipy dependency and no interpolation error. - """ - out = {} - for name, mask, cdf in ( - ("acentric", ~centric, lambda x: 1.0 - torch.exp(-x)), - ("centric", centric, lambda x: torch.erf(torch.sqrt(x * 0.5))), - ): - v = E2[mask].to(torch.float64) - if v.numel() < 50: - out[name] = float("nan") - continue - u = cdf(v.clamp(min=0.0)).clamp(0.0, 1.0) - out[name] = _ks_uniform(torch.sort(u).values) - return out - - -def check_e_convention( - cls, - F: torch.Tensor, - s_mag: torch.Tensor, - centric: torch.Tensor, - *, - sig_F: Optional[torch.Tensor] = None, - eps: Optional[torch.Tensor] = None, - n_shells: int = 20, - F_calc: Optional[torch.Tensor] = None, - n_deciles: int = 10, -) -> dict: - """Run every applicable property check on one convention. - - Takes the **class**, not an instance, because three of the checks need to - construct it again on perturbed inputs -- rescaled F, raised sigmas -- and a - convention that has already normalised its data cannot be asked those - questions. - """ - conv = cls(F, s_mag, centric, sig_F=sig_F, eps=eps, n_shells=n_shells) - E = conv.E.to(torch.float64) - E2 = E * E - cen = conv.centric - dec = _deciles(s_mag, n_deciles) - rep: dict = {"name": cls.__name__ if hasattr(cls, "__name__") else str(cls)} - - # (1) stationarity + (2) absolute unit mean. - rep["mean_E2"] = float(E2.mean()) - per_dec = torch.stack([ - E2[dec == k].mean() if bool((dec == k).any()) else torch.tensor(float("nan")) - for k in range(n_deciles) - ]) - rep["decile_mean_E2"] = [round(float(v), 4) for v in per_dec] - rep["max_decile_dev"] = float((per_dec - 1.0).abs().max()) - # A trend, not just scatter: correlation of the decile mean with resolution. - finite = torch.isfinite(per_dec) - if int(finite.sum()) > 2: - x = torch.arange(n_deciles, dtype=torch.float64)[finite] - y = per_dec[finite].to(torch.float64) - xc, yc = x - x.mean(), y - y.mean() - denom = (xc.norm() * yc.norm()).clamp(min=1e-30) - rep["decile_trend_r"] = float((xc * yc).sum() / denom) - else: - rep["decile_trend_r"] = float("nan") - - # (3) Wilson distribution, globally and worst-decile. - rep["ks"] = _wilson_ks(E2, cen) - worst = 0.0 - for k in range(n_deciles): - m = dec == k - if int(m.sum()) < 100: - continue - ks_k = _wilson_ks(E2[m], cen[m]) - for v in ks_k.values(): - if not math.isnan(v): - worst = max(worst, v) - rep["ks_worst_decile"] = worst - - # (4) moment ratios. - rep["moment_ratio"] = {} - for name, mask in (("acentric", ~cen), ("centric", cen)): - v = E2[mask] - if v.numel() < 50: - rep["moment_ratio"][name] = float("nan") - continue - rep["moment_ratio"][name] = float( - (v * v).mean() / v.mean().clamp(min=1e-30) ** 2 - ) - - # (5) scale invariance: E must not depend on the units of F. - devs = [] - for c in (1e-3, 1e3): - E_c = cls(F * c, s_mag, centric, - sig_F=None if sig_F is None else sig_F * c, - eps=eps, n_shells=n_shells).E.to(torch.float64) - scale = E.abs().max().clamp(min=1e-30) - devs.append(float((E_c - E).abs().max() / scale)) - rep["scale_invariance_dev"] = max(devs) - - # (6) epsilon-correctness: axial reflections must not sit systematically - # above general ones. Only meaningful when eps was supplied and varies. - if eps is not None: - axial = eps > 1.0 - rep["eps_frac_gt1"] = float(axial.to(torch.float64).mean()) - if bool(axial.any()) and bool((~axial).any()): - rep["eps_ratio"] = float( - E2[axial].mean() / E2[~axial].mean().clamp(min=1e-30) - ) - else: - # All or none: no contrast to measure. All-axial means epsilon is - # counting something it should not -- on a centred lattice the - # rotation part repeats per centring coset, so `h.W == h` matches - # once per coset for EVERY reflection. - rep["eps_ratio"] = float("nan") - else: - rep["eps_frac_gt1"] = float("nan") - rep["eps_ratio"] = float("nan") - - # (7) shrinkage monotonicity: raising sigma_F at fixed F must move E toward - # the shell mean, never away. Only conventions that read sig_F can. - if getattr(cls, "uses_sigma_f", False) and sig_F is not None: - loud = cls(F, s_mag, centric, sig_F=sig_F * 4.0, eps=eps, - n_shells=n_shells).E.to(torch.float64) - # The target of the shrinkage is the shell mean IN E-SPACE. Using - # sqrt(Sigma) here instead -- the scale E was divided BY -- compares a - # dimensionless quantity of order 1 against one in units of F, so every - # reflection sits far below the reference and "shrinkage" degenerates - # into "did E get bigger". - ref = torch.zeros_like(E) - for k in range(int(dec.max()) + 1): - m = dec == k - if bool(m.any()): - ref[m] = E[m].mean() - moved_closer = (loud - ref).abs() <= (E - ref).abs() + 1e-9 - rep["shrinkage_frac_ok"] = float(moved_closer.to(torch.float64).mean()) - else: - rep["shrinkage_frac_ok"] = float("nan") - - # (8) obs/calc common footing: normalise a calc set with the same convention - # and compare mean E**2. Both should sit at 1 if the convention puts them - # on a common scale; a mismatch is the eImove defect in miniature. - if F_calc is not None: - # A sigma_F-consuming convention cannot normalise a calc set -- there is - # no measurement error to shrink -- so ask it which companion to use. - calc_cls = cls.for_calc() if hasattr(cls, "for_calc") else cls - rep["calc_via"] = calc_cls.__name__ - E_calc = calc_cls(F_calc, s_mag, centric, sig_F=None, eps=eps, - n_shells=n_shells).E.to(torch.float64) - rep["mean_E2_calc"] = float((E_calc * E_calc).mean()) - rep["obs_calc_ratio"] = rep["mean_E2_calc"] / max(rep["mean_E2"], 1e-30) - else: - rep["calc_via"] = "-" - rep["obs_calc_ratio"] = float("nan") - return rep - - -def format_table(reports) -> str: - """One row per convention, the columns that decide things.""" - head = (f"{'convention':18s} {'':>7s} {'maxdec':>7s} {'trend':>6s} " - f"{'KS ac':>6s} {'KS cen':>7s} {'KSdec':>6s} {'m2 ac':>6s} " - f"{'m2 cen':>7s} {'scale':>8s} {'e>1':>6s} {'eps':>6s} " - f"{'shrink':>7s} {'o/c':>6s} {'calc via':>16s}") - lines = [head, "-" * len(head)] - for r in reports: - lines.append( - f"{r['name']:18s} {r['mean_E2']:>7.4f} {r['max_decile_dev']:>7.4f} " - f"{r['decile_trend_r']:>+6.2f} " - f"{r['ks']['acentric']:>6.4f} {r['ks']['centric']:>7.4f} " - f"{r['ks_worst_decile']:>6.4f} " - f"{r['moment_ratio']['acentric']:>6.3f} " - f"{r['moment_ratio']['centric']:>7.3f} " - f"{r['scale_invariance_dev']:>8.1e} " - f"{r.get('eps_frac_gt1', float('nan')):>6.3f} " - f"{r['eps_ratio']:>6.3f} " - f"{r['shrinkage_frac_ok']:>7.3f} {r['obs_calc_ratio']:>6.3f} " - f"{r.get('calc_via', '-'):>16s}" - ) - lines.append("") - lines.append("ideal: =1 maxdec=0 trend=0 KS small m2 ac=2 cen=3 " - "scale=0 eps=1 shrink=1 o/c=1") - lines.append("e>1 = fraction with epsilon>1; ~1.0 means epsilon is counting " - "centring cosets, not point-group stabilisers") - return "\n".join(lines) diff --git a/alignment_lab/analysis/e_convention_arms.py b/alignment_lab/analysis/e_convention_arms.py deleted file mode 100644 index 5b3e7abd..00000000 --- a/alignment_lab/analysis/e_convention_arms.py +++ /dev/null @@ -1,140 +0,0 @@ -"""Which E convention ranks the true orientation best? - -Layer B of the E-value work: the conformance table says whether a convention -does what E is *supposed* to do; this says whether it makes the rotation -function *work*. When the two disagree, this one decides and the table -diagnoses -- a convention with clean Wilson statistics that ranks truth worse is -not the one to ship, and the table then names the property the winner trades -away. - -Headline metric is the fraction of cells at **rank 0**, not "inside the top 20". -The stated target is that truth comes first; the post-merge baseline is 20 of -100, so the bar is a long way up and a metric that saturates hides the climb. - -Paired by seed: every convention sees the same rotated case from the same -``seed_for``, so arms are compared cell by cell rather than as two distributions. -The ``default`` arm passes no convention at all, which makes it a control on the -*seam* as well as on the conventions -- if it ever diverges from the production -number, the plumbing changed something rather than the convention did. - -The FRF cannot be run once and shared here, unlike the rescore arms: the -convention is what builds the obs expansion, so each arm is a full run. That is -the cost of the question. -""" - -from __future__ import annotations - -import argparse -import functools -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, FRFConfig, e_convention_name, # noqa: E402 - orbit_rank, rotated_case, run_frf, seed_for) - - -def build_arms(): - """Name -> convention class. Built lazily so ``--help`` needs no torchref.""" - from torchref.experimental.alignment.e_values import ( - CalcGlobalE, CalcShellE, FrenchWilsonE, SmoothSigmaE, WilsonShellE, - WilsonShellEpsE, - ) - # Mixed arms. The panel's first round put `calc_global` -- a single global - # RMS on BOTH sides -- ahead of every per-shell convention, which if real - # says the per-shell flattening is discarding inter-shell amplitude shape - # the correlation was using. One class sets both sides, so isolating which - # side carries that needs conventions that differ across the seam. Defined - # here rather than shipped: they exist to answer one question. - class GlobalObsShellCalc(CalcGlobalE): - """Global RMS on obs, per-shell Wilson on calc.""" - calc_companion = WilsonShellE - - class ShellObsGlobalCalc(WilsonShellE): - """Per-shell Wilson on obs, global RMS on calc.""" - calc_companion = CalcGlobalE - - class FrenchWilsonGlobalCalc(FrenchWilsonE): - """Production obs side, global RMS on calc.""" - calc_companion = CalcGlobalE - - return { - # Control: no convention passed, so the production default applies. - "default": None, - # The same thing named explicitly. Must match `default` exactly; if it - # does not, the seam is not inert and nothing below means anything. - "french_wilson": FrenchWilsonE, - # Drops the measurement-error model entirely -- the size of the gap to - # `french_wilson` is what sigma_F is worth to the rotation function. - "wilson": WilsonShellE, - # What the rescore uses on its observed side. Running it here asks - # whether the FRF/rescore disagreement is costing the FRF anything. - "wilson_eps": WilsonShellEpsE, - "calc_shell": CalcShellE, - "calc_global": CalcGlobalE, - # The divergence candidate: a smooth Chebyshev Sigma(s) instead of - # per-shell means. Two orders, because the whole question is whether a - # low-order curve beats 20 independent bins. - "smooth4": functools.partial(SmoothSigmaE, n_coeff=4), - "smooth6": functools.partial(SmoothSigmaE, n_coeff=6), - "smooth10": functools.partial(SmoothSigmaE, n_coeff=10), - "global_x_shell": GlobalObsShellCalc, - "shell_x_global": ShellObsGlobalCalc, - "fw_x_global": FrenchWilsonGlobalCalc, - } - - -def main() -> int: - ap = argparse.ArgumentParser() - ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) - ap.add_argument("--trials", type=int, default=10) - ap.add_argument("--lmax-cap", type=int, default=64) - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--thr-deg", type=float, default=5.0) - ap.add_argument("--arms", default="") - args = ap.parse_args() - - arms = build_arms() - names = [a for a in args.arms.split(",") if a] or list(arms) - unknown = [a for a in names if a not in arms] - if unknown: - raise SystemExit(f"unknown arms {unknown}; have {sorted(arms)}") - - for trial in range(args.trials): - seed = seed_for(args.pdb, trial) - model, data, R_true = rotated_case(args.pdb, seed) - sym = data.spacegroup.matrices.to(torch.float64).cpu() - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - okw = dict(side="left", frame="cart", reciprocal_basis=rec, - thr_deg=args.thr_deg) - - for name in names: - cfg = FRFConfig(n_peaks=args.n_peaks, lmax_cap=args.lmax_cap, - e_convention=arms[name]) - t0 = time.time() - try: - res = run_frf(model, data, cfg, capture_arf=False, verbose=0) - except Exception as exc: # a convention may refuse - print(f"ROW {name} {args.pdb} trial={trial} seed={seed} " - f"rank=-1 rank_cmp={args.n_peaks} found=0 top20=0 " - f"seconds=0.00 error={type(exc).__name__}", flush=True) - continue - seconds = time.time() - t0 - rank, ang = orbit_rank(res.peaks, R_true, sym, **okw) - rank_cmp = rank if rank >= 0 else args.n_peaks - print(f"ROW {name} {args.pdb} trial={trial} seed={seed} " - f"rank={rank} rank_cmp={rank_cmp} found={int(rank >= 0)} " - f"top20={int(0 <= rank < 20)} " - f"angle={'' if ang is None else round(float(ang), 3)} " - f"seconds={seconds:.2f} " - f"conv={e_convention_name(arms[name])}", flush=True) - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/alignment_lab/analysis/e_convention_arms.sh b/alignment_lab/analysis/e_convention_arms.sh deleted file mode 100644 index 799f14ea..00000000 --- a/alignment_lab/analysis/e_convention_arms.sh +++ /dev/null @@ -1,25 +0,0 @@ -#!/bin/bash -# Which E convention ranks truth best? Nine arms x 5 trials x 10 structures. -# the same pass. Every previously published FRF number was measured on - - -#SBATCH --job-name=earms -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=day -#SBATCH --time=03:00:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-9 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname) pdb=$PDB" -"$PY" -u alignment_lab/analysis/e_convention_arms.py --pdb "$PDB" --trials 10 2>/dev/null | grep '^ROW ' - diff --git a/alignment_lab/analysis/e_table.py b/alignment_lab/analysis/e_table.py deleted file mode 100644 index 3e253b2a..00000000 --- a/alignment_lab/analysis/e_table.py +++ /dev/null @@ -1,78 +0,0 @@ -"""Conformance table over every E convention the alignment package uses. - -Establishes what we currently have, before anything changes. Real benchmark data -rather than synthetic, because the properties that matter (the Wilson shape, the -epsilon behaviour on high-symmetry lattices) are properties of real reflection -sets. -""" -from __future__ import annotations - -import argparse -import sys -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -sys.path.insert(0, str(Path(__file__).resolve().parent)) -torch.set_grad_enabled(False) - -from e_conformance import check_e_convention, format_table # noqa: E402 -from lab import BENCH_PDBS, load_case # noqa: E402 - - -def main() -> int: - ap = argparse.ArgumentParser() - ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) - ap.add_argument("--n-shells", type=int, default=20) - args = ap.parse_args() - - from torchref.experimental.alignment.e_values import ( - CalcGlobalE, CalcShellE, FrenchWilsonE, SmoothSigmaE, WilsonShellE, - WilsonShellEpsE, - ) - - model, data = load_case(args.pdb)[:2] - rec = data.cell.reciprocal_basis_matrix.to(torch.float64) - s_mag = (data.hkl.to(torch.float64) @ rec).norm(dim=-1) - F = data.F.to(torch.float64).abs() - sig = None if getattr(data, "F_sigma", None) is None else \ - data.F_sigma.to(torch.float64) - cen = data.centric.to(torch.bool) - eps = data.spacegroup.epsilon(data.hkl.to(torch.long), friedel=False) - - # A calc set from the deposited coordinates: the "perfect model" case, where - # obs and calc genuinely should land on the same scale. - with torch.no_grad(): - F_calc = model.get_structure_factor( - data.hkl, recalc=True).abs().to(torch.float64) - - finite = torch.isfinite(F) & torch.isfinite(F_calc) & (F > 0) - if sig is not None: - finite &= torch.isfinite(sig) & (sig > 0) - F, s_mag, cen, eps, F_calc = (F[finite], s_mag[finite], cen[finite], - eps[finite], F_calc[finite]) - sig = None if sig is None else sig[finite] - - uniq, cnt = torch.unique(eps, return_counts=True) - n_ops = int(data.spacegroup.matrices.shape[0]) - print(f"=== {args.pdb} {data.spacegroup.hm} N={int(finite.sum())} " - f"(dropped {int((~finite).sum())}) n_shells={args.n_shells} ===") - print(f" n_ops={n_ops} epsilon: " - + ", ".join(f"{float(u):g}x{int(c)}" for u, c in zip(uniq, cnt))) - reports = [] - for cls in (FrenchWilsonE, WilsonShellE, WilsonShellEpsE, CalcShellE, - CalcGlobalE, SmoothSigmaE): - try: - reports.append(check_e_convention( - cls, F, s_mag, cen, sig_F=sig, eps=eps, - n_shells=args.n_shells, F_calc=F_calc, - )) - except Exception as exc: # noqa: BLE001 - report it - print(f" {cls.__name__}: RAISED {type(exc).__name__}: {exc}") - print(format_table(reports)) - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/alignment_lab/analysis/e_table.sh b/alignment_lab/analysis/e_table.sh deleted file mode 100644 index d69fec89..00000000 --- a/alignment_lab/analysis/e_table.sh +++ /dev/null @@ -1,28 +0,0 @@ -#!/bin/bash -# Layer A of the E-value work: the property report, paired with the functional -# panel that decides. Run over structures spanning primitive and centred -# lattices and low to high symmetry, because the epsilon and Wilson-shape -# properties are properties of real reflection sets. -#SBATCH --job-name=etable -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=hour -#SBATCH --time=00:45:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-4 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -# 1DAW is C2 (centred, so eps != 1 everywhere), 2DQ6 is the tNCS case whose -# moment ratio reads ~5.5, 3K7M and 3A5V are the high-symmetry ends, 6G9X is -# where the rescore reproducibly fails. -PDBS=(1DAW 2DQ6 3K7M 3A5V 6G9X) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "### $PDB on $(hostname)" -"$PY" -u alignment_lab/analysis/e_table.py --pdb "$PDB" 2>/dev/null -echo "RC=$?" diff --git a/alignment_lab/analysis/full_gate.sh b/alignment_lab/analysis/full_gate.sh index 2764ad9f..2ab25845 100644 --- a/alignment_lab/analysis/full_gate.sh +++ b/alignment_lab/analysis/full_gate.sh @@ -21,5 +21,6 @@ export TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 # relative to the root. Naming the config explicitly satisfies both. "$PY" -m pytest -c tests/pytest.ini tests/unit --run-slow -q 2>&1 | tail -16 echo "PYTEST_RC=${PIPESTATUS[0]}" -echo "== seam identity ==" -"$PY" -u alignment_lab/analysis/seam_identity.py 2>/dev/null | grep -E "SEAM|conv|case|^ *[0-9A-Z]" +echo "== end-to-end placement, one cell ==" +"$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb 1DAW --trial 0 \ + --arms analytic_r 2>/dev/null | grep '^ROW ' diff --git a/alignment_lab/analysis/fw_asu_equivalence.sh b/alignment_lab/analysis/fw_asu_equivalence.sh deleted file mode 100644 index a0e6d683..00000000 --- a/alignment_lab/analysis/fw_asu_equivalence.sh +++ /dev/null @@ -1,58 +0,0 @@ -#!/bin/bash -# Is ASU-then-broadcast equivalent to the unrolled computation? The previous run -# reported max|d eEobs| = 0.22 with "54504 differing", but counted ANY nonzero -# float difference, so that count says nothing about size. Characterise the -# distribution, and find where the outlier sits. -#SBATCH --job-name=frf_fwasu -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -FILT="Loaded|LINK|Wilson outlier|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$|No CUDA" -echo "host=$(hostname)" -"$PY" -u -c " -import numpy as np, torch -torch.set_grad_enabled(False) -from alignment_lab.lab.benchmark import load_case -from torchref.experimental.alignment.frf.french_wilson import french_wilson_preprocess - -for name in ('3K7M', '1DAW'): - model, data = load_case(name); model.verbose = 0 - rb = data.cell.reciprocal_basis_matrix.to(torch.float64) - s_all = (data.hkl.to(torch.float64) @ rb).norm(dim=-1) - d_min = 1.0 / s_all.max().item() - keep = (s_all >= 1.0/100.0) & (s_all <= 1.0/d_min) - n_ops = int(data.spacegroup.matrices.shape[0]) - F = data.F.abs().to(torch.float64)[keep]; sig = data.F_sigma.to(torch.float64)[keep] - s = s_all[keep]; cen = torch.zeros_like(F, dtype=torch.bool) - M = F.numel() - - unrolled = french_wilson_preprocess(F.repeat(n_ops), sig.repeat(n_ops), - s.repeat(n_ops), cen.repeat(n_ops), n_wilson_shells=20) - unique = french_wilson_preprocess(F, sig, s, cen, n_wilson_shells=20) - a = unrolled['eEobs'].reshape(n_ops, M)[0].numpy() - b = unique['eEobs'].numpy() - d = np.abs(a - b) - rel = d / np.maximum(np.abs(a), 1e-12) - print(f'--- {name}: {M} unique x {n_ops} ops') - for q in (50, 90, 99, 99.9, 100): - print(f' |d eEobs| p{q:<5} {np.percentile(d, q):.3e} rel {np.percentile(rel, q):.3e}') - print(f' above 1e-9 : {int((d > 1e-9).sum())} of {M}') - print(f' above 1e-3 : {int((d > 1e-3).sum())} of {M}') - # Do the shell edges actually differ, and by how many reflections? - def edges_and_shells(sv): - sn = sv.numpy(); idx = np.argsort(sn) - e = np.linspace(0, len(sn)-1, 21).round().astype(np.int64) - ed = sn[idx][e].copy(); ed[0] -= 1e-6; ed[-1] += 1e-6 - return ed, np.clip(np.searchsorted(ed, sn, side='right')-1, 0, 19) - eu, shu = edges_and_shells(s) - er, shr = edges_and_shells(s.repeat(n_ops)) - print(f' max |d shell edge| {np.abs(eu - er).max():.3e}') - print(f' reflections changing shell {int((shu != shr[:M]).sum())} of {M}') -" 2>&1 | grep -vE "$FILT" -echo "done" diff --git a/alignment_lab/analysis/fw_footing.py b/alignment_lab/analysis/fw_footing.py deleted file mode 100644 index d0b67970..00000000 --- a/alignment_lab/analysis/fw_footing.py +++ /dev/null @@ -1,95 +0,0 @@ -"""Is French-Wilson's `eEobs` supposed to have unit mean square? - -The conformance table reports `obs_calc_ratio` ~2.1 for `FrenchWilsonE`, i.e. -`` sits near 0.5 where the calc companion sits at 1. Two readings, and -they call for different actions: - -* the FW port is mis-scaled -- a real defect in the FRF's observed side; or -* `eEobs` is a DEFLATED amplitude by construction and the check is asking the - wrong question of it. - -`eEobs**2 = eEsqFW + (DFAC**2 - 1)/DFAC**2` with `DFAC < 1`, so the second term -is strictly negative: the assembly subtracts the share of the measured intensity -that is measurement error. That is a deconvolution, not a normalisation, and -`eEobs` travels with `DFAC` as a pair. - -So the question is not "is `` one" but "is the quantity the CONSUMER -forms centred". The consumer is `build_lerf1_obs_intensity`, which forms -`cw * (eEobs**2 - 1) * DFAC**2` -- it subtracts a literal 1. If `` is -really 0.5, that term carries a systematic negative offset into the Patterson -correlation, and whether that matters is a separate question from whether the -port is faithful. - -This decomposes the assembly term by term, per resolution decile, so the answer -comes from the numbers rather than from reading the formula. -""" - -from __future__ import annotations - -import sys -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import load_case # noqa: E402 - - -def main() -> int: - from torchref.experimental.alignment.frf.french_wilson import ( - french_wilson_preprocess, - ) - from torchref.experimental.alignment.frf.preprocessing import ( - build_lerf1_intensity, - ) - from lab.reference_normalisers import wilson_normalise - - for pdb in ("1DAW", "3K7M", "2DQ6"): - model, data = load_case(pdb) - F = data.F.to(torch.float64).abs().cpu() - sig = data.F_sigma.to(torch.float64).cpu() - hkl = data.hkl.cpu() - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - s = (hkl.to(torch.float64) @ rec).norm(dim=-1) - cen = data.centric.cpu().to(torch.bool) - keep = torch.isfinite(F) & torch.isfinite(sig) & (sig > 0) & (F > 0) - F, sig, s, cen = F[keep], sig[keep], s[keep], cen[keep] - - fw = french_wilson_preprocess(F, sig, s, cen, n_wilson_shells=20) - eE, dfac = fw["eEobs"].to(torch.float64), fw["DFAC"].to(torch.float64) - # eEsqFW is what eEobs**2 would be before the deconvolution term. - corr = (dfac * dfac - 1.0) / (dfac * dfac) - eEsq = eE * eE - corr - wil = wilson_normalise(F, s, 20)[0].to(torch.float64) - lerf = build_lerf1_intensity( - fw["eEobs"], cen, weight=dfac * dfac, use_centric_weight=True, - ).to(torch.float64) - - order = torch.argsort(s) - dec = torch.zeros_like(s, dtype=torch.long) - ch = max(1, s.numel() // 10) - for k in range(10): - hi = (k + 1) * ch if k < 9 else s.numel() - dec[order[k * ch:hi]] = k - - print(f"\n=== {pdb} n={F.numel()} centric={int(cen.sum())} ===") - print(f" {'dec':>3s} {'d(A)':>12s} {'':>9s} {'':>9s} " - f"{'':>10s} {'':>7s} {'':>9s}") - for k in range(10): - m = dec == k - lo, hi = float(1 / s[m].max()), float(1 / s[m].min()) - print(f" {k:>3d} {f'{hi:5.1f}-{lo:4.2f}':>12s} " - f"{float((wil[m] ** 2).mean()):>9.4f} " - f"{float(eEsq[m].mean()):>9.4f} " - f"{float((eE[m] ** 2).mean()):>10.4f} " - f"{float(dfac[m].mean()):>7.4f} {float(lerf[m].mean()):>9.4f}") - print(f" {'ALL':>3s} {'':>12s} {float((wil ** 2).mean()):>9.4f} " - f"{float(eEsq.mean()):>9.4f} {float((eE ** 2).mean()):>10.4f} " - f"{float(dfac.mean()):>7.4f} {float(lerf.mean()):>9.4f}") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/alignment_lab/analysis/fw_footing.sh b/alignment_lab/analysis/fw_footing.sh deleted file mode 100644 index b20e4fcf..00000000 --- a/alignment_lab/analysis/fw_footing.sh +++ /dev/null @@ -1,17 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=fwfoot -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:30:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=32G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 -export CUDA_VISIBLE_DEVICES="" -"$PY" -u alignment_lab/analysis/fw_footing.py -echo "RC=$?" diff --git a/alignment_lab/analysis/fw_internals.sh b/alignment_lab/analysis/fw_internals.sh deleted file mode 100644 index 3238d075..00000000 --- a/alignment_lab/analysis/fw_internals.sh +++ /dev/null @@ -1,93 +0,0 @@ -#!/bin/bash -# What inside french_wilson_preprocess costs, at the unrolled reflection count? -# The binning, the parabolic-cylinder posterior, or the Halley D-factor solve? -# Also: are the n_ops symmetry copies really carrying identical values? -#SBATCH --job-name=frf_fw -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -FILT="Loaded|LINK|Wilson outlier|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$|No CUDA" -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" -"$PY" -u -c " -import time -import numpy as np, torch -torch.set_grad_enabled(False) -from alignment_lab.lab.benchmark import load_case -from torchref.experimental.alignment.frf.french_wilson import ( - french_wilson_preprocess, _french_wilson_posterior, _get_dfactor_vectorised) - -def best(fn, n=3): - fn() - out = [] - for _ in range(n): - t0 = time.perf_counter(); fn(); out.append(time.perf_counter()-t0) - return min(out) - -for name in ('3K7M', '1DAW'): - model, data = load_case(name); model.verbose = 0 - rb = data.cell.reciprocal_basis_matrix.to(torch.float64) - s_all = (data.hkl.to(torch.float64) @ rb).norm(dim=-1) - d_min = 1.0 / s_all.max().item() - keep = (s_all >= 1.0/100.0) & (s_all <= 1.0/d_min) - n_ops = int(data.spacegroup.matrices.shape[0]) - F = data.F.abs().to(torch.float64)[keep] - sig = data.F_sigma.to(torch.float64)[keep] - s = s_all[keep] - from torchref.experimental.alignment.frf.preprocessing import compute_epsilon - cen = torch.zeros_like(F, dtype=torch.bool) - try: - cen = data.centric.to(torch.bool)[keep] - except Exception: - pass - nu = F.numel() - - # the unrolled arrays, exactly as rotation_search builds them - Fu = F.repeat(n_ops); sigu = sig.repeat(n_ops) - su = s.repeat(n_ops); cenu = cen.repeat(n_ops) - - t_unrolled = best(lambda: french_wilson_preprocess(Fu, sigu, su, cenu, n_wilson_shells=20)) - t_unique = best(lambda: french_wilson_preprocess(F, sig, s, cen, n_wilson_shells=20)) - print(f'--- {name}: {nu} unique x {n_ops} ops = {nu*n_ops} unrolled') - print(f' whole function, unrolled {t_unrolled*1e3:8.1f} ms') - print(f' whole function, unique {t_unique*1e3:8.1f} ms ({t_unrolled/t_unique:.1f}x)') - - # Are the n_ops copies identical? If so the unrolled call is pure repetition. - fu = french_wilson_preprocess(Fu, sigu, su, cenu, n_wilson_shells=20) - fq = french_wilson_preprocess(F, sig, s, cen, n_wilson_shells=20) - blocks = fu['eEobs'].reshape(n_ops, nu) - same_across_ops = bool(torch.equal(blocks, blocks[0:1].expand(n_ops, -1))) - matches_unique = bool(torch.equal(blocks[0], fq['eEobs'])) - print(f' n_ops copies identical {same_across_ops}') - print(f' block 0 == unique-set run {matches_unique}') - if not matches_unique: - d = (blocks[0] - fq['eEobs']).abs() - print(f' max|d eEobs| {d.max().item():.3e} n_differing={int((d>0).sum())}') - - # Stage split, on the unrolled arrays. - F_np = Fu.numpy(); sig_np = sigu.numpy(); s_np = su.numpy(); cen_np = cenu.numpy() - def binning(): - idx = np.argsort(s_np) - e = np.linspace(0, len(s_np)-1, 21).round().astype(np.int64) - edges = s_np[idx][e]; edges[0] -= 1e-6; edges[-1] += 1e-6 - sh = np.clip(np.searchsorted(edges, s_np, side='right')-1, 0, 19) - m = np.zeros(20); c = np.zeros(20, dtype=np.int64) - np.add.at(m, sh, F_np*F_np); np.add.at(c, sh, 1) - return m/np.maximum(c,1), sh - t_bin = best(binning) - mF2, sh = binning() - eosq = F_np*F_np/mF2[sh]; sigesq = np.maximum(2.0*F_np*sig_np/mF2[sh], 0.0) - t_post = best(lambda: _french_wilson_posterior(eosq, sigesq, cen_np)) - ee, eesq = _french_wilson_posterior(eosq, sigesq, cen_np) - bad = eesq < ee*ee; eesq[bad] = ee[bad]**2 + 1e-12 - t_dfac = best(lambda: _get_dfactor_vectorised(ee, eesq, cen_np)) - print(f' of which: shells + {t_bin*1e3:8.1f} ms') - print(f' FW posterior {t_post*1e3:8.1f} ms') - print(f' DFAC Halley {t_dfac*1e3:8.1f} ms') -" 2>&1 | grep -vE "$FILT" -echo "done" diff --git a/alignment_lab/analysis/llg_decompose.py b/alignment_lab/analysis/llg_decompose.py deleted file mode 100644 index c53d85f2..00000000 --- a/alignment_lab/analysis/llg_decompose.py +++ /dev/null @@ -1,162 +0,0 @@ -"""Why does the true orientation lose the LLG contest? - -The rescore's job is to take a top-20 that already contains truth and put truth -first. On 6G9X it reproducibly does the opposite -- truth goes from FRF rank 1 to -rank 12-17 -- and that survives the full Phaser model preparation, so it is not -explained by sigma_A or the solvent term. - -This takes one case with known ground truth and asks where the LLG difference -between truth and the candidate that beats it actually accumulates: per -resolution shell, and split by centric/acentric. A likelihood that prefers the -wrong orientation is either being fed the wrong expected intensity or is summing -a term whose sign is wrong somewhere, and both of those localise. - -Per-reflection LL is recomputed here rather than taken from -``_llg_for_orientations``, which sums before returning -- same context, same -``phaser_log_rel_*`` calls, just not reduced. -""" - -from __future__ import annotations - -import argparse -import sys -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, FRFConfig, orbit_rank, rotated_case, # noqa: E402 - run_frf, seed_for) - - -def per_reflection_ll(ctx, alpha, beta, gamma): - """``(n_orient, N)`` per-reflection log-likelihood -- the unsummed LLG.""" - from torchref.experimental.alignment.distributions import ( - phaser_log_rel_rice, phaser_log_rel_woolfson, - ) - from torchref.experimental.alignment.frf.rotation_utils import ( - rotation_matrix_from_edmonds_euler_batch, - ) - R = rotation_matrix_from_edmonds_euler_batch( - alpha.to(torch.float64), beta.to(torch.float64), gamma.to(torch.float64), - ).transpose(-1, -2).to(torch.float32) - F_calc_m = ctx.interpolator.evaluate( - R, ctx.unrolled_hkl, ctx.real_cell, return_amplitude=True, - ).to(ctx.dtype) - if ctx.dw_per_m is not None: - F_calc_m = F_calc_m * ctx.dw_per_m.unsqueeze(0) - E_calc_m = F_calc_m / ctx.sqrt_mean_per_m.unsqueeze(0) - Esq_m = E_calc_m * E_calc_m - B = Esq_m.shape[0] - sum_per_h = torch.zeros(B, ctx.N, dtype=Esq_m.dtype, device=Esq_m.device) - sum_per_h.scatter_add_(1, ctx.asu_idx.unsqueeze(0).expand(B, -1), Esq_m) - eImove = ctx.eImove_prefac * sum_per_h - sqrt_eImove = eImove.clamp(min=1e-30).sqrt() - ll = torch.where( - ctx.centric_b, - phaser_log_rel_woolfson(ctx.E_obs_b, sqrt_eImove, ctx.V_b), - phaser_log_rel_rice(ctx.E_obs_b, sqrt_eImove, ctx.V_b), - ) - return ll, eImove - - -def main() -> int: - ap = argparse.ArgumentParser() - ap.add_argument("--pdb", default="6G9X", choices=list(BENCH_PDBS)) - ap.add_argument("--trial", type=int, default=0) - ap.add_argument("--lmax-cap", type=int, default=64) - ap.add_argument("--n-refine", type=int, default=20) - ap.add_argument("--n-shells", type=int, default=10) - ap.add_argument("--full-prep", action="store_true", - help="turn on the Phaser model prep the pipeline omits") - args = ap.parse_args() - - from torchref.experimental.alignment.ml_rotation import _build_llg_context - - seed = seed_for(args.pdb, args.trial) - model, data, R_true = rotated_case(args.pdb, seed) - sym = data.spacegroup.matrices.to(torch.float64).cpu() - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - okw = dict(side="left", frame="cart", reciprocal_basis=rec, thr_deg=5.0) - - res = run_frf(model, data, FRFConfig(n_peaks=500, lmax_cap=args.lmax_cap), - capture_arf=False, verbose=0) - head = res.peaks[: args.n_refine] - truth_rank, truth_ang = orbit_rank(head, R_true, sym, **okw) - if truth_rank < 0: - print(f"truth not in the top {args.n_refine}; nothing for the rescore " - f"to find here") - return 0 - - inp = res.inputs - prep = {} - if args.full_prep: - prep = dict(vrms_strategy="oeffner", - vrms_n_residues=max(1, int(model.xyz().shape[0] / 8)), - apply_bulk_solvent=True, apply_wilson_b=True) - ctx = _build_llg_context( - inp.F_obs, inp.hkl, inp.s_mag, inp.centric, inp.ll, data.cell, - data.spacegroup, - n_shells=max(20 // 2, 8), batch_size=50, **prep, - ) - - a = torch.tensor([p.alpha for p in head], dtype=torch.float64) - b = torch.tensor([p.beta for p in head], dtype=torch.float64) - g = torch.tensor([p.gamma for p in head], dtype=torch.float64) - ll, eImove = per_reflection_ll(ctx, a, b, g) - totals = ll.sum(dim=-1) - order = torch.argsort(totals, descending=True) - new_rank = int((order == truth_rank).nonzero()[0, 0]) - winner = int(order[0]) - - print(f"=== {args.pdb} trial {args.trial} " - f"({'full prep' if args.full_prep else 'shipped defaults'}) ===") - print(f" truth is FRF rank {truth_rank} ({truth_ang:.2f} deg), " - f"LLG rank {new_rank}; winner is FRF rank {winner}") - print(f" LLG(truth) = {totals[truth_rank]:.4f}") - print(f" LLG(winner) = {totals[winner]:.4f} " - f"gap = {totals[winner] - totals[truth_rank]:+.4f}") - if winner == truth_rank: - print(" truth already wins here") - return 0 - - # Where does the gap accumulate? Equal-count shells in |s|. - s = inp.s_mag.to(torch.float64).cpu() - n_sh = args.n_shells - edge_idx = torch.linspace(0, s.numel() - 1, n_sh + 1).round().long() - edges = s.sort().values[edge_idx] - shell = torch.bucketize(s, edges[1:-1]) - d_ll = (ll[winner] - ll[truth_rank]).to(torch.float64).cpu() - cen = ctx.centric_b[0].cpu() - - print(f"\n gap by resolution shell (winner - truth; positive = truth loses)") - print(f" {'shell':>5s} {'d range (A)':>16s} {'n':>7s} {'d LL':>10s} " - f"{'cum %':>7s} {'acen':>9s} {'cen':>9s}") - total_gap = float(d_ll.sum()) - cum = 0.0 - for k in range(n_sh): - m = shell == k - if not bool(m.any()): - continue - v = float(d_ll[m].sum()) - cum += v - lo, hi = float(1.0 / edges[k + 1]), float(1.0 / edges[k]) - print(f" {k:>5d} {f'{hi:6.1f}-{lo:5.2f}':>16s} {int(m.sum()):>7d} " - f"{v:>+10.3f} {100 * cum / total_gap if total_gap else 0:>6.1f}% " - f"{float(d_ll[m & ~cen].sum()):>+9.3f} " - f"{float(d_ll[m & cen].sum()):>+9.3f}") - print(f" {'TOTAL':>5s} {'':>16s} {int(s.numel()):>7d} {total_gap:>+10.3f}") - - print(f"\n expected moving intensity eImove, truth vs winner:") - for name, idx in (("truth", truth_rank), ("winner", winner)): - e = eImove[idx].to(torch.float64).cpu() - print(f" {name:7s} mean {e.mean():.4e} median {e.median():.4e} " - f"max {e.max():.4e} frac>E_obs^2 " - f"{float((e > (ctx.E_obs_b[0].cpu() ** 2)).float().mean()):.3f}") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/alignment_lab/analysis/llg_decompose.sh b/alignment_lab/analysis/llg_decompose.sh deleted file mode 100644 index d206c26f..00000000 --- a/alignment_lab/analysis/llg_decompose.sh +++ /dev/null @@ -1,24 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=llgdec -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:50:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname)" -for trial in 0 1; do - for flag in "" "--full-prep"; do - "$PY" -u alignment_lab/analysis/llg_decompose.py --pdb 6G9X --trial $trial $flag 2>&1 \ - | grep -vE "UserWarning|FutureWarning|^ *from |^ *warnings\.|Loaded|LINK|Wilson outlier|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization" - echo - done -done -echo "rc=$?" diff --git a/alignment_lab/analysis/pose_arms.sh b/alignment_lab/analysis/pose_arms.sh index a51156e3..5a377903 100644 --- a/alignment_lab/analysis/pose_arms.sh +++ b/alignment_lab/analysis/pose_arms.sh @@ -2,15 +2,15 @@ # End-to-end pose recovery, which is the only metric that is actually the # deliverable. Rank is a proxy; this is not. # -# Two questions in one array: does dropping the rescore cost anything, and is -# n_rotation_candidates=15 leaving coverage on the table? The FRF's worst truth -# rank over 100 cells is 21, so 15 carries 93% and 25 carries 100% -- but -# coverage is not recovery, and only this measures recovery. +# 10 structures x 3 trials x 2 arms. The arms are the open ranking question: +# analytic R (the default) against the translation LLG, which wins 27/30 to +# 22/30 at rank level. The number to hold is 24/30 successes, what the pipeline +# scored once the ML rescore was removed from between the two stages. #SBATCH --job-name=posearm #SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out #SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err #SBATCH --partition=day -#SBATCH --time=08:00:00 +#SBATCH --time=04:00:00 #SBATCH --cpus-per-task=4 #SBATCH --mem=48G #SBATCH --constraint=cpu_epyc9335 @@ -25,7 +25,5 @@ export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREA export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" for T in 0 1 2; do "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial "$T" \ - --arms m_letf1,none --n-rotation-candidates 15 2>/dev/null | grep '^ROW ' - "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial "$T" \ - --arms none --n-rotation-candidates 25 2>/dev/null | grep '^ROW ' + --arms analytic_r,llg_tf --n-rotation-candidates 15 2>/dev/null | grep '^ROW ' done diff --git a/alignment_lab/analysis/rescore_prep_arms.py b/alignment_lab/analysis/rescore_prep_arms.py deleted file mode 100644 index 68694450..00000000 --- a/alignment_lab/analysis/rescore_prep_arms.py +++ /dev/null @@ -1,149 +0,0 @@ -"""Does the ML rescore fail because of its MODEL PREPARATION? - -The rescore is measurably destructive end-to-end -- dropping it takes pose -recovery from 18/30 to 24/30 -- and the damage is concentrated in the two large -benchmark entries, 4BX9 and 6G9X, which it fails at 52-91 degrees and which the -raw FRF order solves at 3-4 degrees. This harness tests one explanation. - -**The FRF and the rescore disagree about the model, inside a single run.** The -pipeline estimates ``model_error_A`` from the model's length, hands it to the -FRF, and applies Phaser's Babinet bulk-solvent term there -- then calls -``m_letf1_rescore`` without passing either, so the rescore falls back to a -hardcoded ``delta_vrms = 0.5`` and no solvent. Every Phaser model-prep knob on -that function defaults OFF and the pipeline overrides none of them. - -That predicts the observed failure pattern rather than merely being consistent -with it. Babinet's ``1 - 0.95 exp(-300 s^2/4)`` tends to 0.05 as ``s -> 0``, so -omitting it over-weights the lowest-resolution reflections by up to ~20x, and -low-resolution terms dominate for large molecules. 0.5 A is also furthest from -the truth for a large model, where ``oeffner_vrms`` gives ~0.67 A. - -Paired by construction: the FRF runs ONCE per (structure, trial) and every arm -rescores the *same* peak list. The ``none`` arm keeps the FRF order and is the -load-bearing control -- without it, an engine that merely preserves a good input -ranking looks like one that improves it. -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, FRFConfig, orbit_rank, rotated_case, # noqa: E402 - run_frf, run_rescore, seed_for) - -#: Each arm is a set of overrides on top of `m_letf1`'s defaults. They are -#: cumulative on purpose: if the whole Phaser prep helps, the interesting -#: question is which piece carries it. -#: `eps_friedel` reproduces the epsilon the rescore used before it was routed -#: through `SpaceGroup.epsilon(friedel=False)`. Passing an explicit `eps_factor` -#: is how the old convention is reproduced without a second worktree, so the -#: comparison stays paired on one FRF peak list. -#: The E-convention arms are the reason this harness is being re-run. The -#: rotation function turned out to be INSENSITIVE to the convention -- 12 of -#: them, 100 paired cells, median rank 2.0 for every one -- which is what a -#: correlation should do, since a global scale cancels out of it. The LLG is a -#: likelihood and has no free scale to cancel, so if the convention matters -#: anywhere it matters here. `no_sigmas` is the control for that: it withholds -#: the sigmas the rescore has only just started receiving. -def _arms(): - """The target question: does a Rice buy anything over weighted least squares? - - The Rice exists to handle an AMPLITUDE -- non-negative, phase unknown, so the - likelihood marginalises over the phase and is biased upward for weak - reflections. `E**2` is an intensity, which is unbiased and near Gaussian, and - the job here is to rank orientations rather than to report calibrated - probabilities. Both sides are already normalised to unit mean square, so the - shrinkage the Rice contributes is being applied to something whose scale is - fixed by construction. - - `wls` scores `-sum_h w_h (E_obs**2 - eImove)**2` with `w` the same combined - inverse variance the rotation function uses, so both stages agree about what - a reflection is worth as well as about what it is compared to. - """ - return { - "wls": {"target": "wls"}, - "wls_nosig": {"target": "wls", "sig_F_obs": None}, - } - -ARMS = { - "none": None, # control: FRF order - "default": {}, # what ships today - "eps_friedel": {"__eps_friedel": True}, # the pre-migration convention - "vrms": {"vrms_strategy": "oeffner"}, - "solvent": {"apply_bulk_solvent": True}, - "vrms_solvent": {"vrms_strategy": "oeffner", "apply_bulk_solvent": True}, - "full_prep": {"vrms_strategy": "oeffner", "apply_bulk_solvent": True, - "apply_wilson_b": True}, -} - - -def main() -> int: - ap = argparse.ArgumentParser() - ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) - ap.add_argument("--trials", type=int, default=3) - ap.add_argument("--lmax-cap", type=int, default=64) - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--n-refine", type=int, default=20, - help="rescore window: the top-N FRF peaks handed to the engine") - ap.add_argument("--thr-deg", type=float, default=5.0) - ap.add_argument("--arms", default="") - args = ap.parse_args() - - ARMS.update(_arms()) - - cfg = FRFConfig(n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) - arms = [a for a in args.arms.split(",") if a] or list(ARMS) - - for trial in range(args.trials): - seed = seed_for(args.pdb, trial) - model, data, R_true = rotated_case(args.pdb, seed) - sym = data.spacegroup.matrices.to(torch.float64).cpu() - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - orbit_kw = dict(side="left", frame="cart", reciprocal_basis=rec, - thr_deg=args.thr_deg) - - # The FRF runs once; every arm sees the identical peak list. - res = run_frf(model, data, cfg, capture_arf=False, verbose=0) - frf_rank, _ = orbit_rank(res.peaks[: args.n_refine], R_true, sym, **orbit_kw) - - # Same residue estimate the pipeline uses for the FRF, so the `vrms` arm - # is genuinely "what the FRF was told" and not a second guess. - n_residues = max(1, int(model.xyz().shape[0] / 8)) - - for arm in arms: - overrides = ARMS[arm] - if overrides is None: - rank, seconds = frf_rank, 0.0 - else: - kw = dict(overrides) - if kw.pop("__eps_friedel", False): - # Friedel-folded epsilon: doubles it on every centric - # reflection, which is what the rescore used to get. - kw["eps_factor"] = data.spacegroup.epsilon( - res.inputs.hkl.to(torch.long), friedel=True, - ).to(res.inputs.F_obs.dtype) - if kw.get("vrms_strategy") == "oeffner": - kw["vrms_n_residues"] = n_residues - t0 = time.time() - rr = run_rescore(res.peaks, data, res.inputs, engine="m_letf1", - n_refine=args.n_refine, verbose=0, **kw) - seconds = time.time() - t0 - rank, _ = orbit_rank(rr.peaks, R_true, sym, **orbit_kw) - # A miss must not sort as a good rank. - rank_cmp = rank if rank >= 0 else args.n_refine - print(f"ROW {arm} {args.pdb} trial={trial} seed={seed} " - f"frf_rank={frf_rank} rank={rank} rank_cmp={rank_cmp} " - f"n_res={n_residues} seconds={seconds:.2f}", flush=True) - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/alignment_lab/analysis/rescore_prep_arms.sh b/alignment_lab/analysis/rescore_prep_arms.sh deleted file mode 100644 index 2322d397..00000000 --- a/alignment_lab/analysis/rescore_prep_arms.sh +++ /dev/null @@ -1,22 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=resprep -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=day -#SBATCH --time=04:00:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-9 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname) pdb=$PDB" -"$PY" -u alignment_lab/analysis/rescore_prep_arms.py --pdb "$PDB" --trials 10 2>&1 \ - | grep -E "^ROW |Error|Traceback|Warning: " -echo "rc=$?" diff --git a/alignment_lab/analysis/seam_gate.sh b/alignment_lab/analysis/seam_gate.sh deleted file mode 100644 index a2e2185b..00000000 --- a/alignment_lab/analysis/seam_gate.sh +++ /dev/null @@ -1,24 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=seamgate -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:50:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -# Single-threaded: the FRF peak list only reproduces bit-for-bit at one thread -# (~5e-8 score noise reorders 12 of 500 peaks on 3GR5 otherwise), and this gate -# is about bit-identity. -export TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 -echo "== pytest ==" -"$PY" -m pytest tests/unit -q 2>&1 | tail -12 -echo "PYTEST_RC=${PIPESTATUS[0]}" -echo "== seam identity ==" -"$PY" -u alignment_lab/analysis/seam_identity.py 2>/dev/null | grep -E "^(case|SEAM|[0-9A-Z]{4} )" -echo "RC=$?" diff --git a/alignment_lab/analysis/seam_identity.py b/alignment_lab/analysis/seam_identity.py deleted file mode 100644 index f1867c25..00000000 --- a/alignment_lab/analysis/seam_identity.py +++ /dev/null @@ -1,112 +0,0 @@ -"""Is the E-convention seam inert? - -Routing the rotation function's normalisation through an `EConvention` class is -only safe to build on if it changes nothing while the default is in place. A -peak-list hash would answer that, but it needs a second worktree to compare -against and it only samples the structures it is run on. - -This is stronger and cheaper. The seam replaced exactly three tensors -- -`eEobs`, the LERF1 weight, and `E_calc` -- and every line downstream of them is -untouched. So bit-identity on those three is not evidence that the peak list is -unchanged, it is a proof of it, on whatever data this is run over. - -The one at real risk is the calc side. `wilson_normalise` accumulates its shell -sums with `index_add_` and clamps the mean at 1e-12; `EConvention` uses -`scatter_add_` and clamps at 1e-30. Same arithmetic in exact terms, and on CPU -both reduce in index order -- but "should be identical" is the claim under test, -not the assumption behind it. - -Reports max absolute and relative deviation rather than a bare pass/fail, so a -non-zero result says how big it is instead of only that it exists. -""" - -from __future__ import annotations - -import sys -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import load_case # noqa: E402 - -CASES = ("1DAW", "3K7M", "2DQ6", "4BX9") - - -def _dev(new: torch.Tensor, old: torch.Tensor): - """``(n_differing, max_abs, max_rel)`` between two tensors.""" - d = (new.to(torch.float64) - old.to(torch.float64)).abs() - rel = d / old.to(torch.float64).abs().clamp(min=1e-30) - return int((d > 0).sum()), float(d.max()), float(rel.max()) - - -def main() -> int: - from torchref.experimental.alignment.e_values import ( - FrenchWilsonE, WilsonShellE, - ) - from torchref.experimental.alignment.frf.french_wilson import ( - french_wilson_preprocess, - ) - from torchref.experimental.alignment.frf.preprocessing import ( - build_lerf1_intensity, - ) - from lab.reference_normalisers import wilson_normalise - from torchref.experimental.alignment.sh import ( - assign_shells, equal_count_shell_edges, - ) - - n_shells = 20 - worst = 0.0 - print(f"{'case':>6s} {'tensor':>12s} {'n':>8s} {'n_diff':>7s} " - f"{'max abs':>10s} {'max rel':>10s}") - for pdb in CASES: - model, data = load_case(pdb) - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - hkl = data.hkl.cpu() - s = (hkl.to(torch.float64) @ rec).norm(dim=-1) - F = data.F.to(torch.float64).abs().cpu() - sig = data.F_sigma.to(torch.float64).cpu() - cen = data.centric.cpu().to(torch.bool) - keep = torch.isfinite(F) & torch.isfinite(sig) & (sig > 0) & (F > 0) - F, sig, s, cen = F[keep], sig[keep], s[keep], cen[keep] - - edges, _ = equal_count_shell_edges(s, n_shells) - shell_idx = assign_shells(s, edges) - - # --- obs side ------------------------------------------------- - fw = french_wilson_preprocess(F, sig, s, cen, n_wilson_shells=n_shells, - shell_idx=shell_idx) - conv = FrenchWilsonE(F, s, cen, sig_F=sig, shell_idx=shell_idx, - n_shells=n_shells) - for name, new, old in ( - ("eEobs", conv.E, fw["eEobs"]), - ("weight", conv.weight, fw["DFAC"] * fw["DFAC"]), - ("lerf1", - build_lerf1_intensity(conv.E, cen, weight=conv.weight), - build_lerf1_intensity(fw["eEobs"], cen, - weight=fw["DFAC"] * fw["DFAC"])), - ): - nd, a, r = _dev(new, old) - worst = max(worst, a) - print(f"{pdb:>6s} {name:>12s} {new.numel():>8d} {nd:>7d} " - f"{a:>10.3e} {r:>10.3e}") - - # --- calc side: the one where the two implementations differ ---- - # A stand-in calc set; the check is about the normaliser, not the model. - F_calc = (F * 1.37 + 5.0) - old_E, _ = wilson_normalise(F_calc, s, n_shells) - new_E = WilsonShellE(F_calc, s, cen, n_shells=n_shells).E - nd, a, r = _dev(new_E, old_E) - worst = max(worst, a) - print(f"{pdb:>6s} {'E_calc':>12s} {new_E.numel():>8d} {nd:>7d} " - f"{a:>10.3e} {r:>10.3e}") - - print(f"\nSEAM_WORST_ABS {worst:.6e}") - print("SEAM_INERT" if worst == 0.0 else "SEAM_NOT_INERT") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/alignment_lab/diagnostics/frf_vs_ftf_discrimination.py b/alignment_lab/diagnostics/frf_vs_ftf_discrimination.py index aa91846e..d8203aa5 100644 --- a/alignment_lab/diagnostics/frf_vs_ftf_discrimination.py +++ b/alignment_lab/diagnostics/frf_vs_ftf_discrimination.py @@ -25,8 +25,8 @@ The placement path is the production one -- ``_make_rotated`` then the same ``precompute_G`` / ``amplitude_translation_search`` / ``local_translation_refine`` -calls ``_placement_for_candidate`` makes -- with ``do_joint_refine`` off, since -the rigid-body polish is a later stage and would confound the TF's own ranking. +calls ``_placement_for_candidate`` makes, against the pipeline's own +``TranslationObs`` so the normalisation and weighting are the production ones. """ from __future__ import annotations @@ -94,8 +94,8 @@ def main() -> int: ap.add_argument("--out-csv", default=None) args = ap.parse_args() - from torchref.experimental.alignment.align import ( - _DirectModelEvaluator, _prepare_frf_inputs, + from torchref.experimental.alignment.rotation_search import ( + prepare_frf_inputs, ) from torchref.experimental.alignment.frf.rotation_utils import ( rotation_matrix_from_edmonds_euler, @@ -104,8 +104,8 @@ def main() -> int: MolecularReplacementPipeline, ) from torchref.experimental.alignment.translation import ( - amplitude_translation_search, local_translation_refine, - precompute_G_for_rotation, + DirectModelEvaluator, amplitude_translation_search, + local_translation_refine, precompute_G_for_rotation, ) writer = None @@ -124,13 +124,11 @@ def main() -> int: pipe = MolecularReplacementPipeline( data, model, verbose=0, - rescore_engine="none", subpeak_refine=False, n_rotation_peaks=args.n_rotation_peaks, n_rotation_candidates=args.n_cand, - do_joint_refine=False, dense_rotation_refine=False, use_llg_tf=False, ) - frf = _prepare_frf_inputs( + frf = prepare_frf_inputs( model, data, d_min=pipe.d_min, d_max=pipe.d_max, n_shells=pipe.n_shells, verbose=0, ) @@ -155,12 +153,12 @@ def main() -> int: rot.spacegroup = data.spacegroup.hm p1 = rot.copy() p1.spacegroup = "P 1" - ev = _DirectModelEvaluator(p1) + ev = DirectModelEvaluator(p1) G, h_R = precompute_G_for_rotation( - ev, eye3, pipe._hkl_keep, data.spacegroup, data.cell) + ev, eye3, pipe._obs.hkl, data.spacegroup, data.cell) _, _, tp = amplitude_translation_search( - F_obs=pipe._F_obs_amp, interpolator=ev, R_rotation=eye3, - hkl=pipe._hkl_keep, spacegroup=data.spacegroup, + obs=pipe._obs, interpolator=ev, R_rotation=eye3, + spacegroup=data.spacegroup, real_cell=data.cell, grid_steps=pipe.translation_grid_steps, n_peaks=pipe.n_translation_peaks, cluster_radius=0.05, precomputed_G=G, precomputed_h_R=h_R) @@ -179,9 +177,8 @@ def main() -> int: r_best = float("inf") for cand in llg_peaks[: pipe.n_translation_candidates]: _, r_a = local_translation_refine( - F_obs=pipe._F_obs_amp, interpolator=ev, R_rotation=eye3, - hkl=pipe._hkl_keep, spacegroup=data.spacegroup, - real_cell=data.cell, + obs=pipe._obs, interpolator=ev, R_rotation=eye3, + spacegroup=data.spacegroup, real_cell=data.cell, t_init=torch.as_tensor(cand.translation, dtype=torch.float64), radius=0.06, grid_steps=13, n_refinement_passes=1, diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py index 5a1ad3e0..f55e5278 100644 --- a/alignment_lab/diagnostics/pose_recovery.py +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -92,7 +92,7 @@ def main() -> int: ap.add_argument("--out-csv", default=None) args = ap.parse_args() - from torchref.experimental.alignment.align import align_model_to_data + from torchref.experimental.alignment import align_model_to_data seed = seed_for(args.pdb, args.trial) model, data = load_case(args.pdb) diff --git a/alignment_lab/diagnostics/rescore_rank.py b/alignment_lab/diagnostics/rescore_rank.py deleted file mode 100644 index 8f605622..00000000 --- a/alignment_lab/diagnostics/rescore_rank.py +++ /dev/null @@ -1,118 +0,0 @@ -"""Does the ML rescore improve the FRF ranking, or damage it? - -The FRF reliably puts the true orientation inside the top 20 on this benchmark, -yet end-to-end pose recovery succeeds about half the time. That points at the -rescore, so this measures it directly and in isolation. - -For each structure the FRF is run **once**, then every rescore arm is applied to -the *same* peak list -- including a ``none`` control that leaves the FRF order -untouched. The reported quantity is the paired change in the rank of truth -within the rescore window, so an arm that merely preserves a good input ranking -cannot be mistaken for one that improves it. - -Rows where truth was never in the window are recorded with ``delta`` empty: the -rescore had nothing to find, and scoring it there would measure the FRF. - -Usage:: - - python alignment_lab/diagnostics/rescore_rank.py --pdb 1AK5 --trial 0 \ - --engines none,m_letf1,sim --out-csv alignment_lab/runs/rescore.csv -""" - -from __future__ import annotations - -import argparse -import sys -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, FRFConfig, ResultWriter, orbit_rank, # noqa: E402 - paired_ranks, rotated_case, run_frf, run_rescore, seed_for) - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="1AK5", choices=list(BENCH_PDBS)) - ap.add_argument("--trial", type=int, default=0) - ap.add_argument("--lmax-cap", type=int, default=64) - ap.add_argument("--d-min", type=float, default=4.0) - ap.add_argument("--d-max", type=float, default=15.0) - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--n-refine", type=int, default=20, - help="rescore window: the top-N FRF peaks handed to the engine") - ap.add_argument("--engines", default="none,m_letf1,sim") - ap.add_argument("--subpeak-refine", action="store_true") - ap.add_argument("--orbit-side", default="left", choices=["left", "right"]) - ap.add_argument("--orbit-frame", default="cart", choices=["cart", "frac"]) - ap.add_argument("--verbose", type=int, default=0) - ap.add_argument("--out-csv", default=None) - args = ap.parse_args() - - seed = seed_for(args.pdb, args.trial) - rotated, data, R_true = rotated_case(args.pdb, seed) - sym = data.spacegroup.matrices.to(torch.float64).cpu() - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - orbit_kw = dict(side=args.orbit_side, frame=args.orbit_frame, - reciprocal_basis=rec) - - cfg = FRFConfig(d_min=args.d_min, d_max=args.d_max, - n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) - frf = run_frf(rotated, data, cfg, verbose=args.verbose) - rank_full, ang_full = orbit_rank(frf.peaks, R_true, sym, **orbit_kw) - - print(f"=== {args.pdb} t{args.trial} seed={seed} {data.spacegroup} " - f"n_ops={sym.shape[0]} | FRF rank={rank_full} ang={ang_full:.2f} " - f"({frf.seconds:.1f}s) | window={args.n_refine} ===") - if rank_full < 0 or rank_full >= args.n_refine: - print(f" NOTE truth is outside the rescore window " - f"(FRF rank {rank_full}); the rescore cannot recover it, so the " - f"deltas below measure nothing about the rescore.") - print(f" {'engine':10s} {'rank_in':>8s} {'rank_out':>9s} {'delta':>6s} " - f"{'ang_out':>8s} {'secs':>7s}") - - writer = None - if args.out_csv: - writer = ResultWriter(args.out_csv, "rescore_rank", - extra_fields=("engine", "n_refine", "rank_frf_full", - "rank_frf_window", "rank_rescored", - "delta", "truth_in_window", - "angle_rescored", "rescore_seconds", - "subpeak_refine")) - for engine in [e.strip() for e in args.engines.split(",") if e.strip()]: - res = run_rescore(frf.peaks, data, frf.inputs, engine=engine, - n_refine=args.n_refine, - subpeak_refine=args.subpeak_refine, - verbose=args.verbose) - pr = paired_ranks(frf.peaks, res.peaks, R_true, sym, - n_refine=args.n_refine, **orbit_kw) - delta = pr["delta"] - print(f" {engine:10s} {pr['rank_frf']:8d} {pr['rank_rescored']:9d} " - f"{('' if delta is None else f'{delta:+d}'):>6s} " - f"{pr['angle_rescored']:8.2f} {res.seconds:7.1f}") - if writer: - writer.write(pdb=args.pdb, seed=seed, trial=args.trial, - spacegroup=str(data.spacegroup), n_ops=int(sym.shape[0]), - truth_rank=pr["rank_rescored"], - truth_angle_deg=round(pr["angle_rescored"], 4), - orbit_side=args.orbit_side, orbit_frame=args.orbit_frame, - lmax_cap=args.lmax_cap, d_min=args.d_min, d_max=args.d_max, - device="cpu", engine=engine, n_refine=args.n_refine, - rank_frf_full=pr["rank_frf_full"], - rank_frf_window=pr["rank_frf"], - rank_rescored=pr["rank_rescored"], - delta=("" if delta is None else delta), - truth_in_window=int(pr["truth_in_window"]), - angle_rescored=round(pr["angle_rescored"], 4), - rescore_seconds=round(res.seconds, 2), - subpeak_refine=int(args.subpeak_refine)) - if args.out_csv: - print(f" wrote {args.out_csv}") - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/tf_batch_probe.py b/alignment_lab/diagnostics/tf_batch_probe.py deleted file mode 100644 index 03f82827..00000000 --- a/alignment_lab/diagnostics/tf_batch_probe.py +++ /dev/null @@ -1,152 +0,0 @@ -"""Can the translation stage share one molecular transform across orientations? - -``tf_cost.py`` says the per-candidate placement is ~95% ``precompute_G``: a full -structure-factor evaluation of a *re-rotated* model at ``S*N`` Miller indices, -paid again for every orientation. The Crowther-Blow accumulation and the local -refine that follow it are milliseconds. - -But the rotation does not have to live in the coordinates. ``F(h; R x) = -F(R^T h; x)``, which is exactly what :class:`LattmanLoveInterpolator` is for and -what the *rotation* search already uses -- and its ``evaluate`` is already -batched over ``R``. So ``G`` for M orientations could be one dense grid plus -``M*S*N`` trilinear lookups instead of M structure-factor calculations. - -Two things have to hold and neither is obvious: - -* **The phase must survive trilinear interpolation.** The rotation search only - ever reads ``|F|``; the translation function reads ``arg F``, and the class's - own docstring warns that complex interpolation is only safe on a - well-oversampled grid. Measured here as the agreement of the *translation - peaks*, not of ``F`` -- a phase error that does not move the peak does not - matter. -* **It must actually be faster**, including the one-off grid build, at the - orientation counts we would carry. - -Reports both against the exact per-rotation path. -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) - ap.add_argument("--n-rot", type=int, default=8, help="orientations to time") - ap.add_argument("--d-min", type=float, default=4.0) - ap.add_argument("--d-max", type=float, default=15.0) - ap.add_argument("--max-res", type=float, default=3.0) - ap.add_argument("--padding", type=float, default=2.0) - ap.add_argument("--grid-steps", type=int, default=16) - args = ap.parse_args() - - from torchref.experimental.alignment.align import _DirectModelEvaluator - from torchref.experimental.alignment.lattman_love import LattmanLoveInterpolator - from torchref.experimental.alignment.translation import ( - amplitude_translation_search, precompute_G_for_rotation, - ) - - model, data = load_case(args.pdb) - rec = data.cell.reciprocal_basis_matrix.to(torch.float64) - mask = data.get_valid_mask() - s_mag = (data.hkl.to(torch.float64) @ rec).norm(dim=-1) - mask = mask & (s_mag >= 1.0 / args.d_max) & (s_mag <= 1.0 / args.d_min) - hkl = data.hkl[mask] - F_obs = data.F[mask].abs().to(torch.float64) - S = int(data.spacegroup.matrices.shape[0]) - N = int(hkl.shape[0]) - print(f"# {args.pdb} sg={data.spacegroup.hm} S={S} N={N} " - f"atoms={model.xyz().shape[0]} max_res={args.max_res} " - f"padding={args.padding}", flush=True) - - base = model.copy() - base.spacegroup = "P 1" - origin = torch.zeros(3, dtype=base.xyz().dtype) - eye3 = torch.eye(3, dtype=torch.float64) - - t0 = time.perf_counter() - ll = LattmanLoveInterpolator(base, padding_factor=args.padding, - max_res_A=args.max_res, verbose=0) - t_grid = time.perf_counter() - t0 - print(f"# dense grid build {t_grid:.2f} s shape=" - f"{tuple(ll.reciprocal_grid.shape)}", flush=True) - - # h_R is a function of hkl and the space group only -- not of the candidate - # rotation -- so the index side of the whole stage is shared. - sym_R = data.spacegroup.matrices.to(torch.float64) - h_R = torch.einsum("ne,ied->ind", hkl.to(torch.float64), sym_R) - h_R_flat = h_R.reshape(-1, 3) - - rots = [random_rotation(seed_for(args.pdb, 0) + 97 * k) - for k in range(args.n_rot)] - - t_direct = t_interp = 0.0 - for k, R in enumerate(rots): - rot = base.copy().rotate(R.to(base.dtype_float), center=origin) - ev = _DirectModelEvaluator(rot) - t0 = time.perf_counter() - G_d, h_R_d = precompute_G_for_rotation(ev, eye3, hkl, - data.spacegroup, data.cell) - t_direct += time.perf_counter() - t0 - - t0 = time.perf_counter() - F_i = ll.evaluate(R.to(torch.float32), h_R_flat, data.cell, - return_amplitude=False).reshape(S, N) - phase_sym = torch.exp(2j * torch.pi * torch.einsum( - "ne,ie->in", hkl.to(torch.float64), - data.spacegroup.translations.to(torch.float64), - ).to(torch.complex128)) - G_i = F_i.to(torch.complex128) * phase_sym - t_interp += time.perf_counter() - t0 - - if k == 0: - a, b = G_d.reshape(-1), G_i.reshape(-1) - coh = float((a.conj() * b).sum().abs() - / (a.abs().norm() * b.abs().norm()).clamp(min=1e-30)) - amp = float(torch.corrcoef(torch.stack( - [a.abs(), b.abs()]).to(torch.float64))[0, 1]) - print(f"# G agreement: complex coherence={coh:.4f} " - f"|F| corr={amp:.4f}", flush=True) - - tf = lambda G: amplitude_translation_search( - F_obs=F_obs, interpolator=ev, R_rotation=eye3, hkl=hkl, - spacegroup=data.spacegroup, real_cell=data.cell, - grid_steps=args.grid_steps, n_peaks=5, - precomputed_G=G, precomputed_h_R=h_R_d)[2] - pd, pi = tf(G_d), tf(G_i) - dt = min(float(torch.tensor( - ((torch.as_tensor(pd[0].translation) - - torch.as_tensor(pi[j].translation) + 0.5) % 1.0 - 0.5).norm())) - for j in range(len(pi))) - dt_top = float(torch.tensor( - ((torch.as_tensor(pd[0].translation) - - torch.as_tensor(pi[0].translation) + 0.5) % 1.0 - 0.5).norm())) - print(f"ROW pdb={args.pdb} rot={k} dt_top={dt_top:.4f} " - f"dt_best_of_5={dt:.4f} " - f"score_direct={pd[0].score:.4f} score_interp={pi[0].score:.4f}", - flush=True) - - M = args.n_rot - print(f"SUM pdb={args.pdb} S={S} N={N} n_rot={M} " - f"t_direct_per_rot={t_direct / M:.3f} " - f"t_interp_per_rot={t_interp / M:.4f} " - f"t_grid={t_grid:.2f} " - f"breakeven_rots={t_grid / max(t_direct / M - t_interp / M, 1e-9):.1f} " - f"speedup_at_100={100 * (t_direct / M) / (t_grid + 100 * t_interp / M):.1f}", - flush=True) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/tf_cost.py b/alignment_lab/diagnostics/tf_cost.py index 3c92bfcf..b8a4352d 100644 --- a/alignment_lab/diagnostics/tf_cost.py +++ b/alignment_lab/diagnostics/tf_cost.py @@ -45,10 +45,9 @@ def main() -> int: ap.add_argument("--n-peaks", type=int, default=20) args = ap.parse_args() - from torchref.experimental.alignment.align import _DirectModelEvaluator from torchref.experimental.alignment.translation import ( - amplitude_translation_search, local_translation_refine, - precompute_G_for_rotation, + DirectModelEvaluator, TranslationObs, amplitude_translation_search, + local_translation_refine, precompute_G_for_rotation, ) seed = seed_for(args.pdb, args.trial) @@ -65,25 +64,29 @@ def main() -> int: ).norm(dim=-1).clamp(min=1e-9) mask = mask & (d >= args.d_min) & (d <= args.d_max) hkl = data.hkl[mask] - F_obs = data.F[mask].abs().to(torch.float64) + sig_F = getattr(data, "F_sigma", None) + obs = TranslationObs.build( + data.F[mask], hkl, data.spacegroup, data.cell, + sig_F=None if sig_F is None else sig_F[mask], + ) S = int(data.spacegroup.matrices.shape[0]) N = int(hkl.shape[0]) n_at = int(rot.xyz().shape[0]) print(f"# {args.pdb} sg={data.spacegroup.hm} S={S} N={N} atoms={n_at} " f"grid={args.grid_steps} seed={seed}", flush=True) - ev = _DirectModelEvaluator(rot) + ev = DirectModelEvaluator(rot) eye3 = torch.eye(3, dtype=torch.float64) t_G, (G, h_R) = _time(lambda: precompute_G_for_rotation( ev, eye3, hkl, data.spacegroup, data.cell)) t_tf, (_, _, peaks) = _time(lambda: amplitude_translation_search( - F_obs=F_obs, interpolator=ev, R_rotation=eye3, hkl=hkl, + obs=obs, interpolator=ev, R_rotation=eye3, spacegroup=data.spacegroup, real_cell=data.cell, grid_steps=args.grid_steps, n_peaks=args.n_peaks, precomputed_G=G, precomputed_h_R=h_R)) t_ref, _ = _time(lambda: local_translation_refine( - F_obs=F_obs, interpolator=ev, R_rotation=eye3, hkl=hkl, + obs=obs, interpolator=ev, R_rotation=eye3, spacegroup=data.spacegroup, real_cell=data.cell, t_init=torch.as_tensor(peaks[0].translation, dtype=torch.float64), radius=0.06, grid_steps=13, n_refinement_passes=1, diff --git a/alignment_lab/lab/__init__.py b/alignment_lab/lab/__init__.py index aa0601e8..191d6971 100644 --- a/alignment_lab/lab/__init__.py +++ b/alignment_lab/lab/__init__.py @@ -27,10 +27,7 @@ fit_aniso_log_space, tensor_report, ) -from .frf import (FRFConfig, FRFResult, e_convention_name, - merge_peak_lists, patched, - run_frf) -from .rescore import ENGINES, RescoreResult, paired_ranks, run_rescore +from .frf import FRFConfig, FRFResult, merge_peak_lists, patched, run_frf from .profile import (FRF_STAGES, PeakMemory, calibration_seconds, host_info, stage_timers) from .results import ResultWriter, append_row, provenance @@ -51,15 +48,10 @@ "fit_aniso_log_space", "tensor_report", "FRFConfig", - "e_convention_name", "FRFResult", "merge_peak_lists", "patched", "run_frf", - "ENGINES", - "RescoreResult", - "paired_ranks", - "run_rescore", "FRF_STAGES", "PeakMemory", "calibration_seconds", diff --git a/alignment_lab/lab/aniso.py b/alignment_lab/lab/aniso.py index dacbcfc0..02d3cb43 100644 --- a/alignment_lab/lab/aniso.py +++ b/alignment_lab/lab/aniso.py @@ -113,9 +113,9 @@ def aniso_arm(arm: str, data, *, d_min: float, d_max: float, captured: dict): ``captured`` receives the tensor actually fitted under the key ``raw``, so a caller can report the artefact size alongside the rank it costs. - Patches ``align.fit_overall_anisotropy``, which is the symbol - ``_prepare_frf_inputs`` calls and whose result is handed to the rotation - search as ``U_aniso``. + Patches ``sh.fit_overall_anisotropy`` where ``rotation_search`` binds it -- + that is the symbol ``fit_anisotropy`` calls, and its result is what reaches + the engine as ``U_aniso``. Parameters ---------- @@ -125,13 +125,16 @@ def aniso_arm(arm: str, data, *, d_min: float, d_max: float, captured: dict): Used to recompute the centric mask over the same resolution window the fit sees. A length mismatch raises rather than misaligning silently. d_min, d_max : float - The window ``_prepare_frf_inputs`` was called with. + The window ``prepare_frf_inputs`` was called with. captured : dict Filled in by the wrapper. """ if arm not in ARMS: raise ValueError(f"unknown aniso arm {arm!r}; expected one of {ARMS}") - from torchref.experimental.alignment import align as _align + import importlib + + _align = importlib.import_module( + "torchref.experimental.alignment.rotation_search") original = _align.fit_overall_anisotropy rec = data.cell.reciprocal_basis_matrix.to(torch.float64) diff --git a/alignment_lab/lab/frf.py b/alignment_lab/lab/frf.py index a3cf4d59..f191922b 100644 --- a/alignment_lab/lab/frf.py +++ b/alignment_lab/lab/frf.py @@ -15,18 +15,6 @@ import torch -def e_convention_name(conv) -> str: - """Display name for a convention class, a ``partial`` of one, or ``None``.""" - if conv is None: - return "default" - inner = getattr(conv, "func", conv) - name = getattr(inner, "__name__", str(inner)) - kw = getattr(conv, "keywords", None) - if kw: - name += "(" + ",".join(f"{k}={v}" for k, v in sorted(kw.items())) + ")" - return name - - @dataclass class FRFConfig: """Engine settings for one FRF evaluation. @@ -45,12 +33,6 @@ class FRFConfig: #: Expected r.m.s. coordinate error, in Angstrom. ``None`` uses the Oeffner #: estimate from the model's length, which is what the pipeline does. model_error_A: Optional[float] = None - #: E-value convention, as the CLASS the engine instantiates once per side. - #: ``None`` leaves the production default in place; a class (or a - #: ``functools.partial`` of one) sweeps it. Unlike the deleted ``extra`` - #: knobs this is a real production parameter, so the lab passes it through - #: rather than patching a constant. - e_convention: Optional[type] = None #: Weighting, the other half of the split. ``None`` leaves the production #: default. These are separate arms on purpose: the design changes three #: things at once -- the observed-side weight, the calculated-side weight @@ -68,8 +50,6 @@ def as_row(self) -> Dict[str, Any]: """Config fields for a result row (``extra`` flattened out).""" d = asdict(self) d.pop("extra") - # `asdict` cannot render a class or a partial; name it instead. - d["e_convention"] = e_convention_name(self.e_convention) d.update(self.extra) return d @@ -205,7 +185,6 @@ def run_frf( import importlib - from torchref.experimental.alignment import align as _align from torchref.experimental.alignment.frf import api as _api # `from ...alignment import rotation_search` gives the FUNCTION, which the @@ -228,7 +207,7 @@ def _wrapped(self, *args, **kwargs): # sweeps them by rebinding those constants for the duration of one call, so # the production API stays switch-free while the measurements that chose the # values remain reproducible. - frf_inputs = _align._prepare_frf_inputs( + frf_inputs = _rs.prepare_frf_inputs( model, data, d_min=cfg.d_min, d_max=cfg.d_max, n_shells=cfg.n_shells, verbose=verbose, ) @@ -242,10 +221,7 @@ def _wrapped(self, *args, **kwargs): f"torchref.experimental.alignment.rotation_search instead." ) - # Omitted rather than passed as None, so an unset convention takes the - # production default from the signature instead of overriding it with one. - conv_kw = {} if cfg.e_convention is None else { - "e_convention": cfg.e_convention} + conv_kw = {} # Engine knobs are omitted when unset so the production default applies, # rather than being passed as None and overriding it with nothing. for _name in ("obs_weight", "shell_variance_weights", "snr_cap", diff --git a/alignment_lab/lab/reference_normalisers.py b/alignment_lab/lab/reference_normalisers.py deleted file mode 100644 index 59bddc83..00000000 --- a/alignment_lab/lab/reference_normalisers.py +++ /dev/null @@ -1,53 +0,0 @@ -"""E-value normalisers as they were before the convention seam, frozen. - -These are not production code and are not imported by it. They are kept so a -future convention can be compared against what the rotation function actually -shipped, rather than against a description of it -- the comparison the seam was -originally validated with, and the one any replacement will want again. - -Frozen means frozen: if a production convention changes, these do not follow. -That is the whole point of an oracle. -""" - -from __future__ import annotations - -import torch - -from torchref.experimental.alignment.sh import ( - assign_shells, equal_count_shell_edges, -) - - -def wilson_normalise( - F: torch.Tensor, - s_mag: torch.Tensor, - n_shells: int = 20, -): - """Per-shell Wilson normalisation of amplitudes. - - Source: Phaser's ``Feff[r] / SIGMAN.sqrt_epsnSN[r]`` (``DataMR.cc:925``) - minus French-Wilson + explicit ε (``F`` is assumed anisotropy-corrected - by the caller). - - E_h = F_h / sqrt(_p) where p = shell containing h. - - Returns ``(E_h, sqrt_mean_F2_per_h)``. - """ - edges, _ = equal_count_shell_edges(s_mag, n_shells) - shell_idx = assign_shells(s_mag, edges) - valid = shell_idx >= 0 - F_dtype = F.dtype - F2 = F * F - count = torch.zeros(n_shells, dtype=torch.int64, device=F.device) - sumF2 = torch.zeros(n_shells, dtype=F_dtype, device=F.device) - F2_v = F2[valid] - idx_v = shell_idx[valid] - count.index_add_(0, idx_v, torch.ones_like(idx_v)) - sumF2.index_add_(0, idx_v, F2_v) - mean_F2 = sumF2 / count.clamp(min=1).to(F_dtype) - mean_F2 = mean_F2.clamp(min=1e-12) - sqrt_mean = mean_F2.sqrt() - per_h = torch.ones_like(F) - per_h[valid] = sqrt_mean[idx_v] - E = F / per_h - return E, per_h diff --git a/alignment_lab/lab/rescore.py b/alignment_lab/lab/rescore.py deleted file mode 100644 index b8ba7df1..00000000 --- a/alignment_lab/lab/rescore.py +++ /dev/null @@ -1,202 +0,0 @@ -"""Run the ML rescore over FRF peaks, and measure what it did to the ranking. - -The rescore's job is narrow: take the top ~20 FRF peaks -- which on this -benchmark reliably contain the true orientation -- and promote the true one to -the front. It is **not** a global search, so feeding it a peak list that does -not contain truth measures nothing. - -That makes the only honest metric a **paired** one: the rank of truth in the -list going in, versus its rank in the list coming out, on the same peaks. An -absolute post-rescore rank cannot distinguish "the rescore worked" from "the -FRF handed it an easy list". -""" - -from __future__ import annotations - -import time -from dataclasses import dataclass -from typing import Any, Dict, List, Optional, Sequence, Tuple - -import torch - -#: Rescore engines. ``none`` keeps the FRF's own ordering and is the control -#: arm -- without it, a rescore that merely preserves a good input ranking is -#: indistinguishable from one that improves it. -ENGINES = ("none", "m_letf1", "sim") - - -@dataclass -class RescoreResult: - """Outcome of one rescore. - - Attributes - ---------- - peaks : list - Re-ordered ``RotationPeak`` list. - engine : str - Engine used. - seconds : float - Wall time. - n_input : int - Peaks handed to the engine. - """ - - peaks: list - engine: str - seconds: float - n_input: int - - -def run_rescore( - peaks: Sequence, - data, - frf_inputs, - *, - engine: str = "m_letf1", - n_refine: int = 20, - n_shells: Optional[int] = None, - batch_size: int = 50, - subpeak_refine: bool = False, - verbose: int = 0, - **engine_kwargs: Any, -) -> RescoreResult: - """Rescore the top ``n_refine`` FRF peaks. - - Parameters - ---------- - peaks : sequence - FRF peaks, descending score. - data : ReflectionData - Dataset the peaks were scored against. - frf_inputs : FRFInputs - Prepared observations from the FRF run (``FRFResult.inputs``). - engine : {'none', 'm_letf1', 'sim'}, optional - ``'none'`` returns the input order unchanged -- the control arm. - n_refine : int, optional - How many leading peaks to rescore. Default 20, the intended use case. - n_shells : int, optional - Resolution shells for the rescore. Defaults to the pipeline's own rule, - ``max(n_shells // 2, 8)`` with ``n_shells = 20``. - batch_size : int, optional - Orientations evaluated per batch. Default 50. - subpeak_refine : bool, optional - Apply the quadratic tangent-space refinement after rescoring. - verbose : int, optional - Engine verbosity. - **engine_kwargs - Passed through to the engine (e.g. ``e_convention``). - - Returns - ------- - RescoreResult - """ - if engine not in ENGINES: - raise ValueError(f"engine must be one of {ENGINES}, got {engine!r}") - - subset = list(peaks)[: max(int(n_refine), 0)] - if engine == "none" or not subset: - return RescoreResult(peaks=subset, engine=engine, seconds=0.0, - n_input=len(subset)) - - from torchref.experimental.alignment.ml_rotation import ( - m_letf1_rescore, sim_mlrf_rescore, - ) - - n_shells = n_shells if n_shells is not None else max(20 // 2, 8) - device = frf_inputs.F_obs.device - common = dict( - n_shells=n_shells, n_refine=len(subset), batch_size=batch_size, - verbose=verbose, - ) - - t0 = time.time() - if engine == "m_letf1": - # The sigmas now reach the rescore. They did not before: the FRF - # computed the French-Wilson posterior from them and then discarded - # them, leaving a likelihood with no measurement-error information. - # An explicit `sig_F_obs` in `engine_kwargs` still wins, so an arm can - # withhold them as a control. - kw = dict(engine_kwargs) - kw.setdefault("sig_F_obs", frf_inputs.sig_F) - out = m_letf1_rescore( - subset, frf_inputs.F_obs, frf_inputs.hkl, frf_inputs.s_mag, - frf_inputs.centric, frf_inputs.ll, data.cell, - data.spacegroup, - **common, **kw, - ) - else: - out = sim_mlrf_rescore( - subset, frf_inputs.F_obs, frf_inputs.hkl, frf_inputs.s_mag, - frf_inputs.centric, frf_inputs.ll, data.cell, - **common, **engine_kwargs, - ) - if subpeak_refine: - from torchref.experimental.alignment.ml_rotation import ( - _build_llg_context, quadratic_llg_refine, - ) - - ctx = _build_llg_context( - frf_inputs.F_obs, frf_inputs.hkl, frf_inputs.s_mag, - frf_inputs.centric, frf_inputs.ll, data.cell, n_shells=n_shells, - ) - out = quadratic_llg_refine(out, ctx) - seconds = time.time() - t0 - return RescoreResult(peaks=list(out), engine=engine, seconds=seconds, - n_input=len(subset)) - - -def paired_ranks( - frf_peaks: Sequence, - rescored: Sequence, - R_true: torch.Tensor, - symops: torch.Tensor, - *, - n_refine: int, - thr_deg: float = 5.0, - **orbit_kw: Any, -) -> Dict[str, Any]: - """Rank of truth before and after rescoring, on the same peak subset. - - Parameters - ---------- - frf_peaks : sequence - Full FRF peak list. - rescored : sequence - Output of :func:`run_rescore`. - R_true : torch.Tensor - True rotation. - symops : torch.Tensor - Symmetry rotation parts. - n_refine : int - Size of the subset handed to the rescore -- the comparison window. - thr_deg : float, optional - Orbit match threshold. - **orbit_kw - Orbit convention (``side``/``frame``/``reciprocal_basis``). - - Returns - ------- - dict - ``rank_frf`` (within the subset), ``rank_rescored``, ``delta`` - (positive = the rescore made it worse), ``truth_in_window`` and - ``rank_frf_full``. ``delta`` is ``None`` when truth was never in the - window, because then the rescore was never given the chance. - """ - from .truth import orbit_rank - - subset = list(frf_peaks)[:n_refine] - rank_full, _ = orbit_rank(frf_peaks, R_true, symops, thr_deg=thr_deg, **orbit_kw) - rank_in, ang_in = orbit_rank(subset, R_true, symops, thr_deg=thr_deg, **orbit_kw) - rank_out, ang_out = orbit_rank(rescored, R_true, symops, thr_deg=thr_deg, **orbit_kw) - - in_window = rank_in >= 0 - delta = (rank_out - rank_in) if (in_window and rank_out >= 0) else None - return { - "rank_frf_full": rank_full, - "rank_frf": rank_in, - "rank_rescored": rank_out, - "delta": delta, - "truth_in_window": in_window, - "angle_frf": ang_in, - "angle_rescored": ang_out, - } diff --git a/alignment_lab/tests/test_lab.py b/alignment_lab/tests/test_lab.py index f5a66faf..4e757855 100644 --- a/alignment_lab/tests/test_lab.py +++ b/alignment_lab/tests/test_lab.py @@ -118,46 +118,3 @@ def test_result_writer_rejects_undeclared_columns(tmp_path): with pytest.raises(KeyError): w.write(pdb="1DAW", not_declared=1) assert (tmp_path / "r.csv").read_text().count("\n") == 2 - - -def test_paired_ranks_reports_no_delta_when_truth_is_outside_the_window(): - """A rescore cannot be blamed for a peak it was never shown. - - When truth is absent from the top-N handed to the engine, ``delta`` is None - rather than a number -- otherwise the metric silently reports the FRF's - failure as a rescore regression. - """ - from types import SimpleNamespace - - from lab import paired_ranks, random_rotation - from torchref.symmetry import SpaceGroup - from torchref.experimental.alignment.frf.rotation_utils import ( - rotation_matrix_from_edmonds_euler, - ) - - symops = SpaceGroup("P 1").matrices.to(torch.float64).cpu() - R_true = random_rotation(3) - - # A peak list whose only truth-matching entry sits beyond the window. - def peak_at(R): - # recover ZYZ angles numerically is unnecessary: use a far-off peak for - # the decoys and the true rotation only at the tail. - return SimpleNamespace(alpha=0.0, beta=0.0, gamma=0.0) - - decoys = [peak_at(None) for _ in range(5)] - out = paired_ranks(decoys, decoys, R_true, symops, - n_refine=2, frame="frac", thr_deg=1e-6) - assert out["truth_in_window"] is False - assert out["delta"] is None - - -def test_run_rescore_none_is_an_identity_control(): - """The 'none' arm must return the input order untouched.""" - from types import SimpleNamespace - - from lab import run_rescore - - peaks = [SimpleNamespace(alpha=float(i), beta=0.0, gamma=0.0) for i in range(5)] - res = run_rescore(peaks, data=None, frf_inputs=None, engine="none", n_refine=3) - assert res.engine == "none" - assert [p.alpha for p in res.peaks] == [0.0, 1.0, 2.0] diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index 5021f5ca..426198b9 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -484,7 +484,6 @@ class TargetDeviceCase: "PhaseInformedDifferenceTarget": "needs two datasets + phases", "RiceDifferenceTarget": "needs two datasets", "TaylorCorrectedDifferenceTarget": "needs two datasets", - "RigidTransform": "alignment helper; needs a coordinate set", } # Everything under torchref/experimental is out of scope for the conformance diff --git a/tests/unit/alignment/test_sh.py b/tests/unit/alignment/test_sh.py index ed7ba918..c9c15244 100644 --- a/tests/unit/alignment/test_sh.py +++ b/tests/unit/alignment/test_sh.py @@ -1,11 +1,20 @@ -""" -Unit tests for torchref.experimental.alignment.sh: spherical harmonic primitives. - -Conventions verified: -- Y_{l,m} are fully orthonormal physics SH with Condon-Shortley phase. -- Matches scipy.special.sph_harm. -- For a centrosymmetric (Friedel-symmetric) input, sh_expand_ball produces - exactly-zero odd-l coefficients when enforce_friedel=True. +"""Leaf mathematics for the rotation function: shell binning and the Legendre reference. + +Two independent things, both pinned because something downstream trusts them. + +**Shell binning.** Two consumers deriving their own equal-count edges from the +same ``|s|`` is how boundary reflections end up in different shells depending on +which stage asked. The edges are computed once and the index passed down; these +tests pin the round trip that makes that safe. + +**``_bar_legendre_recurrence``.** Production never calls it -- the Bessel-SH +expansion runs the same recurrence inside its kernels. It exists as the +*independent* implementation that +``tests/unit/frf_separate/test_bessel_sh_grouping.py`` builds a slow reference +expansion on, to check the fused one against something other than itself. That +only works if the reference is itself trustworthy, which is what the scipy +comparison here is for: it used to reach this recurrence through ``evaluate_ylm``, +and that wrapper is gone. """ import math @@ -13,165 +22,118 @@ import pytest import torch -scipy_special = pytest.importorskip("scipy.special") - from torchref.experimental.alignment.sh import ( _bar_legendre_recurrence, - evaluate_ylm, - sh_expand_ball, - equal_count_shell_edges, assign_shells, + equal_count_shell_edges, ) -def _scipy_ylm(l, m, theta, phi): - """Reference: scipy spherical harmonics with C-S phase, physics convention. +# --------------------------------------------------------------------------- +# The Legendre reference the FRF expansion is checked against +# --------------------------------------------------------------------------- - scipy >= 1.15 replaced ``sph_harm(m, l, phi, theta)`` with - ``sph_harm_y(n, m, theta, phi)``; prefer the new API and fall back to the - old one for older scipy installs. - """ - if hasattr(scipy_special, "sph_harm_y"): - return scipy_special.sph_harm_y(l, m, theta, phi) - return scipy_special.sph_harm(m, l, phi, theta) +@pytest.mark.unit +@pytest.mark.parametrize("L", [4, 9]) +def test_bar_legendre_matches_scipy(L): + """bar_P_l^m(x) = sqrt[(2l+1)/(4pi) (l-m)!/(l+m)!] |P_l^m(x)|, against scipy. -@pytest.mark.parametrize("L", [4, 8, 16]) -def test_ylm_matches_scipy(L): - """Our Y_lm should match scipy.special.sph_harm at random points.""" - torch.manual_seed(0) - n = 20 - theta = torch.rand(n, dtype=torch.float64) * math.pi - phi = (torch.rand(n, dtype=torch.float64) - 0.5) * 2 * math.pi + scipy's ``lpmv`` carries the Condon-Shortley phase and this recurrence does + not, so the comparison is on magnitude -- which is the convention the + docstring states and the kernels implement. + """ + scipy_special = pytest.importorskip("scipy.special") - Y = evaluate_ylm(theta, phi, L) # (n, L, 2L-1) + theta = torch.tensor([0.3, 1.1, 2.0, 2.9], dtype=torch.float64) + cos_t, sin_t = torch.cos(theta), torch.sin(theta) + bar_P = _bar_legendre_recurrence(cos_t, sin_t, L) # (M, L, L) + x = cos_t.numpy() for l in range(L): - for m in range(-l, l + 1): - ref = _scipy_ylm(l, m, theta.numpy(), phi.numpy()) - got = Y[:, l, L - 1 + m].numpy() - np.testing.assert_allclose(got, ref, atol=1e-12, rtol=1e-10, - err_msg=f"mismatch at l={l}, m={m}") - - -def test_ylm_orthonormality_on_grid(): - """Y_lm should be ~orthonormal when integrated on a fine spherical grid.""" - L = 6 - # Gauss-Legendre in cos(theta), uniform in phi - n_theta = 2 * L + 4 - n_phi = 4 * L + 4 - # GL nodes in [-1, 1] - x_gl, w_gl = np.polynomial.legendre.leggauss(n_theta) - theta_np = np.arccos(x_gl) - phi_np = np.linspace(0, 2 * math.pi, n_phi, endpoint=False) - weights = (np.repeat(w_gl, n_phi)) * (2 * math.pi / n_phi) # (n_theta*n_phi,) - theta = torch.tensor(np.repeat(theta_np, n_phi), dtype=torch.float64) - phi = torch.tensor(np.tile(phi_np, n_theta), dtype=torch.float64) - - Y = evaluate_ylm(theta, phi, L) # (n_pts, L, 2L-1) - Yflat = Y.reshape(-1, L * (2 * L - 1)) - w = torch.tensor(weights, dtype=torch.float64) - - # G[a,b] = sum_pts w * Y*_a * Y_b - G = torch.einsum("p,pa,pb->ab", w.to(Yflat.dtype), Yflat.conj(), Yflat) - G_np = G.numpy() - - # Only diagonals for valid (l,m) entries (m | <= l) should be 1; off-diagonal ~0. - expected = np.zeros_like(G_np, dtype=np.complex128) - for l in range(L): - for m in range(-l, l + 1): - idx = l * (2 * L - 1) + (L - 1 + m) - expected[idx, idx] = 1.0 - # Zero out the entries we don't care about (l < |m|, where Y is zero anyway) - valid_mask = np.zeros(L * (2 * L - 1), dtype=bool) - for l in range(L): - for m in range(-l, l + 1): - valid_mask[l * (2 * L - 1) + (L - 1 + m)] = True - - G_valid = G_np[np.ix_(valid_mask, valid_mask)] - exp_valid = expected[np.ix_(valid_mask, valid_mask)] - np.testing.assert_allclose(G_valid, exp_valid, atol=1e-9, rtol=1e-9) - - + for m in range(l + 1): + norm = math.sqrt( + (2 * l + 1) / (4 * math.pi) + * math.factorial(l - m) / math.factorial(l + m) + ) + expected = norm * np.abs(scipy_special.lpmv(m, l, x)) + np.testing.assert_allclose( + bar_P[:, l, m].abs().numpy(), expected, atol=1e-11, rtol=1e-11, + err_msg=f"l={l}, m={m}", + ) + + +@pytest.mark.unit def test_bar_legendre_pole_values(): - """At the north pole (cos θ = 1), bar_P_l^m = 0 for m > 0 and bar_P_l^0 ≠ 0.""" + """At the north pole (cos theta = 1): zero for m > 0, sqrt((2l+1)/4pi) at m = 0. + + The pole is where the recurrence is most fragile -- sin(theta) = 0 kills the + sectoral seed, and everything above it comes from the vertical step. + """ L = 5 theta = torch.tensor([0.0], dtype=torch.float64) - bar_P = _bar_legendre_recurrence(torch.cos(theta), torch.sin(theta), L) # (1, L, L) - # m > 0 must be zero (sin θ = 0 kills sectorals; vertical recurrence then ~ 0) + bar_P = _bar_legendre_recurrence(torch.cos(theta), torch.sin(theta), L) for l in range(L): for m in range(1, l + 1): - assert abs(bar_P[0, l, m].item()) < 1e-14, f"l={l}, m={m}: {bar_P[0,l,m].item()}" - # m = 0: bar_P_l^0(1) = sqrt((2l+1)/(4π)) (Legendre polynomial at 1 is 1) - for l in range(L): - expected = math.sqrt((2 * l + 1) / (4 * math.pi)) - np.testing.assert_allclose(bar_P[0, l, 0].item(), expected, atol=1e-12) + assert abs(bar_P[0, l, m].item()) < 1e-14, f"l={l}, m={m}" + np.testing.assert_allclose( + bar_P[0, l, 0].item(), math.sqrt((2 * l + 1) / (4 * math.pi)), + atol=1e-12, + ) + +# --------------------------------------------------------------------------- +# Shell binning +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +def test_shell_assignment_round_trip(): + """The bins really are equal-count, and every reflection lands in one. -def test_friedel_enforces_even_l(): - """sh_expand_ball with enforce_friedel=True must give exactly-zero odd-l coefficients.""" + ``equal_count_shell_edges`` returns ``(edges, centers)`` -- the second value + is the shell mid-points, not the occupancies -- so equal-count is asserted + on the assignment rather than read off the return. + """ + torch.manual_seed(0) + s = torch.rand(1000, dtype=torch.float64) * 0.4 + 0.05 + n_shells = 12 + edges, centers = equal_count_shell_edges(s, n_shells) + idx = assign_shells(s, edges) + + assert edges.shape == (n_shells + 1,) + assert centers.shape == (n_shells,) + assert int(idx.min()) >= 0 and int(idx.max()) < n_shells + counts = torch.bincount(idx, minlength=n_shells) + assert int(counts.sum()) == s.numel(), "a reflection fell outside every shell" + # 1000 into 12 cannot divide evenly; equal-count means within one of ideal. + assert int(counts.max()) - int(counts.min()) <= 1, counts + # Centres must sit inside their own shell. + assert bool(((centers > edges[:-1]) & (centers < edges[1:])).all()) + + +@pytest.mark.unit +def test_shells_are_monotone_in_resolution(): + """A shell is a resolution range, so the bins must not interleave.""" torch.manual_seed(1) - L = 8 - P = 4 - n_pts = 1000 - s_vectors = torch.randn(n_pts, 3, dtype=torch.float64) * 0.5 - # |s| roughly in [0, 1]; make sure non-zero - s_vectors = s_vectors / (1 + s_vectors.norm(dim=-1, keepdim=True) * 0.1) - s_mags = s_vectors.norm(dim=-1) - edges, _ = equal_count_shell_edges(s_mags, P) - shell_idx = assign_shells(s_mags, edges) - values = torch.rand(n_pts, dtype=torch.float64) - - f_plm = sh_expand_ball(s_vectors, values, shell_idx, P, L, enforce_friedel=True) - - # Odd-l rows must be exactly zero (after explicit zero of FP drift). - for l in range(1, L, 2): - assert f_plm[:, l, :].abs().max().item() == 0.0, f"l={l} not zero" - - -def test_friedel_without_enforce_has_odd_l_for_nonsymmetric_input(): - """Without Friedel enforcement, a non-centrosymmetric scatter produces nonzero odd-l.""" - torch.manual_seed(2) - L = 6 - P = 1 - # Place all mass at the north pole — extremely non-centrosymmetric - s_vectors = torch.tensor([[0.0, 0.0, 1.0]] * 5, dtype=torch.float64) - values = torch.ones(5, dtype=torch.float64) - shell_idx = torch.zeros(5, dtype=torch.int64) + s = torch.rand(500, dtype=torch.float64) * 0.3 + 0.02 + edges, _ = equal_count_shell_edges(s, 8) + idx = assign_shells(s, edges) + hi = torch.stack([s[idx == b].max() for b in range(8) if (idx == b).any()]) + assert bool((hi[1:] >= hi[:-1]).all()) - f_no_friedel = sh_expand_ball(s_vectors, values, shell_idx, P, L, - enforce_friedel=False) - # odd-l (l=1) entries should be non-trivial - odd_norm = f_no_friedel[:, 1, :].abs().max().item() - assert odd_norm > 1e-3, "expected nonzero odd-l without Friedel enforcement" +@pytest.mark.unit +def test_assignment_is_stable_under_a_subset(): + """Slicing rows must not move a reflection's shell; rebuilding edges would. -def test_shell_assignment_round_trip(): - """assign_shells gives indices that round-trip through equal_count_shell_edges.""" - torch.manual_seed(3) - s_mags = torch.rand(2000, dtype=torch.float64) * 2.0 - P = 16 - edges, centers = equal_count_shell_edges(s_mags, P) - idx = assign_shells(s_mags, edges) - # all should be in [0, P-1] - assert idx.min().item() >= 0 - assert idx.max().item() == P - 1 - # roughly equal counts (within 2x because of tie-breaking at edges) - counts = torch.bincount(idx, minlength=P) - assert counts.min().item() >= s_mags.numel() / (4 * P) - - -def test_sh_expand_zero_when_l_too_large_for_no_points(): - """Empty shell should give all-zero coefficients.""" - L = 4 - P = 3 - n_pts = 100 - torch.manual_seed(4) - s_vectors = torch.randn(n_pts, 3, dtype=torch.float64) - values = torch.ones(n_pts, dtype=torch.float64) - # Force shell_idx == 1 for all points (shell 0 and 2 empty) - shell_idx = torch.ones(n_pts, dtype=torch.int64) - f = sh_expand_ball(s_vectors, values, shell_idx, P, L, enforce_friedel=False) - assert f[0].abs().max() == 0 - assert f[2].abs().max() == 0 - assert f[1].abs().max() > 0 + This is the property the "assign once, pass it down" rule rests on: given + the SAME edges, a subset of the reflections lands in the same bins it did in + the full set. + """ + torch.manual_seed(2) + s = torch.rand(600, dtype=torch.float64) * 0.3 + 0.02 + edges, _ = equal_count_shell_edges(s, 10) + full = assign_shells(s, edges) + sub = torch.arange(0, 600, 3) + torch.testing.assert_close(assign_shells(s[sub], edges), full[sub]) diff --git a/torchref/experimental/alignment/__init__.py b/torchref/experimental/alignment/__init__.py index 2d24a92b..4ed3330a 100644 --- a/torchref/experimental/alignment/__init__.py +++ b/torchref/experimental/alignment/__init__.py @@ -1,19 +1,35 @@ -""" -Molecular replacement for TorchRef: a rotation search feeding a translation -search, and one Wilson normalisation shared between them. +"""Molecular replacement: a rotation search feeding a translation search. + +Two stages and one normalisation between them. + +1. **Fast Rotation Function** (:func:`rotation_search`, over + :class:`~torchref.experimental.alignment.frf.FastRotationFunction`) -- + Phaser-faithful Bessel-radial x spherical-harmonic expansion against a dense + P1-box calc, with stable Wigner-d. It is a **shortlist generator**: over 30 + seeded cells it puts the true orientation at rank 0 six times, and inside the + top twenty essentially always. Only the second of those is required. +2. **Fast Translation Function** (:mod:`~torchref.experimental.alignment.translation`) + -- a Crowther-Blow correlation over the fractional cell, run per rotation + candidate, then an analytical-R local refine. On the same 30 cells it reaches + rank 0 in 24, and its likelihood in 27. Rotation ghosts are morphologically + identical to truth in a Patterson by construction; they stop being identical + once the crystal lattice is involved. + +There is deliberately nothing between the two. An ML rescore used to sit there +and was removed: it reordered a shortlist that already contained truth, and +end-to-end pose recovery was 18/30 with it against 24/30 without. + +Both stages normalise through :class:`torchref.scaling.WilsonNormaliser` and +weight through :mod:`torchref.scaling.weighting`, so ``E_obs`` means one thing +across the whole run. -1. Fast Rotation Function (``rotation_search``, over - ``frf.FastRotationFunction``) — Phaser-faithful Bessel-radial × SH - expansion, stable Wigner-d, dense P1-box calc. A shortlist generator. -2. Fast Translation Function (``translation.amplitude_translation_search`` + - ``local_translation_refine``) — run per rotation candidate, and where the - discrimination actually happens. -3. Pipeline (``pipeline.MolecularReplacementPipeline``) — the multi-candidate - FRF → FTF tree with early stopping, which ``align.align_model_to_data`` - delegates to. It returns a placement; refine it downstream. +The pipeline returns a **placement** -- rotation and translation -- and stops. +Refining it is downstream refinement's job, and deleting the post-placement +polish that used to be here took pose recovery from 24/30 to 30/30, because on +the hard cases it walked a correct placement away from truth. -Example — full MR pipeline --------------------------- +Example +------- :: from torchref.experimental.alignment import MolecularReplacementPipeline @@ -23,9 +39,8 @@ data = ReflectionData().load_mtz('observed.mtz') model = ModelFT().load_pdb('search_model.pdb') - pipeline = MolecularReplacementPipeline(data, model) - solutions = pipeline.run(n_rotation_peaks=200, min_tries=3, max_tries=10) - print(f"Best R-factor: {solutions[0].r_factor:.3f}") + solutions = MolecularReplacementPipeline(data, model).run() + print(f"best R-work: {solutions[0].r_factor:.3f}") """ import warnings @@ -35,9 +50,6 @@ FutureWarning, ) -# ============================================================================= -# Fast Rotation Function — Phaser-faithful, single engine -# ============================================================================= from .frf import ( FastRotationFunction, RotationPeak, @@ -47,26 +59,15 @@ rotation_angular_distance_deg, rotation_matrix_from_edmonds_euler, ) -from .sh import ( - evaluate_ylm, - sh_expand_ball, - equal_count_shell_edges, - assign_shells, -) -from .pipeline import ( - MolecularReplacementPipeline, - MRSolution, - cluster_rotation_peaks, - rotation_angular_distance, - euler_angular_distance, +from .rotation_search import ( + FRFInputs, + RotationSolutions, + prepare_frf_inputs, + rotation_search, ) -from .align import align_model_to_data -from .rotation_search import RotationSolutions, rotation_search - -# ============================================================================= -# Translation search -# ============================================================================= from .translation import ( + DirectModelEvaluator, + TranslationObs, TranslationPeak, amplitude_translation_search, find_translation_peaks, @@ -75,54 +76,43 @@ local_translation_refine, precompute_G_for_rotation, ) - -# ============================================================================= -# ML distributions -# ============================================================================= -from .distributions import ( - stable_log_bessel_i0, - rice_log_likelihood, - woolfson_log_likelihood, +from .pipeline import ( + MolecularReplacementPipeline, + MRSolution, + align_model_to_data, + cluster_rotation_peaks, + euler_angular_distance, + rotation_angular_distance, ) __all__ = [ + # Entry points + "align_model_to_data", + "MolecularReplacementPipeline", + "MRSolution", # Rotation search + "rotation_search", + "RotationSolutions", "FastRotationFunction", + "FRFInputs", + "prepare_frf_inputs", "phaser_lmax_resolution", "dense_calc_via_box", "RotationPeak", "rotation_matrix_from_edmonds_euler", "edmonds_euler_from_rotation_matrix", "rotation_angular_distance_deg", - # Rescore + interpolation - # Low-level math primitives - "evaluate_ylm", - "sh_expand_ball", - "equal_count_shell_edges", - "assign_shells", - # Pipeline - "MolecularReplacementPipeline", - "MRSolution", "cluster_rotation_peaks", "rotation_angular_distance", "euler_angular_distance", - "align_model_to_data", - "rotation_search", - "RotationSolutions", - # Translation + # Translation search + "TranslationObs", "TranslationPeak", + "DirectModelEvaluator", "amplitude_translation_search", - "find_translation_peaks", - "fit_sigma_a_per_shell", - "llg_translation_rescore", "local_translation_refine", + "llg_translation_rescore", "precompute_G_for_rotation", - # Rigid body refinement - # Transforms - # Clash scoring - # Distributions - "stable_log_bessel_i0", - "rice_log_likelihood", - "woolfson_log_likelihood", - # Utilities + "find_translation_peaks", + "fit_sigma_a_per_shell", ] diff --git a/torchref/experimental/alignment/align.py b/torchref/experimental/alignment/align.py deleted file mode 100644 index 2530bd40..00000000 --- a/torchref/experimental/alignment/align.py +++ /dev/null @@ -1,324 +0,0 @@ -""" -Molecular replacement: data-prep / FRF stage helpers + the public entry point. - -This module hosts the heavy, reusable stage helpers — Lattman-Love / anisotropy -data prep (`_prepare_frf_inputs`), the solvent-aware R-work -(`_external_rwork`), the direct-SF translation evaluator -(`_DirectModelEvaluator`) and the stage -timer (`_StageTimer`) — that are shared by the rotation-ranking benchmarks and -by the orchestrator. - -`align_model_to_data` is the public entry point. It delegates the -FRF → FTF(per-candidate) → post-refine control flow to -:class:`torchref.experimental.alignment.pipeline.MolecularReplacementPipeline`, -returning that pipeline's single best `ModelFT`. -""" - -from __future__ import annotations - -import time -from contextlib import contextmanager -from dataclasses import dataclass -from typing import Optional, TYPE_CHECKING - -import torch - -from .sh import ( - apply_overall_anisotropy, - assign_shells, - equal_count_shell_edges, - fit_overall_anisotropy, -) - -if TYPE_CHECKING: - from ...io.datasets.reflection_data import ReflectionData - from ...model.model_ft import ModelFT - - -# --------------------------------------------------------------------------- -# Stage timing -# --------------------------------------------------------------------------- - - -class _StageTimer: - """Lightweight wall-clock accumulator. Gated by ``verbose >= 2``. - - Two interleavable usages: - * ``with t.stage(name):`` block — records the block's wall time. - * ``t.start(name)`` / ``t.stop(name)`` — checkpoint pair, no indent. - - The summary table prints stages aggregated by name; per-rotation loop - stages (translation search etc.) get aggregated counts. - """ - - def __init__(self, enabled: bool): - self.enabled = enabled - self.records: list[tuple[str, float]] = [] - self._open: dict[str, float] = {} - - @contextmanager - def stage(self, name: str): - if not self.enabled: - yield - return - t0 = time.perf_counter() - try: - yield - finally: - self.records.append((name, time.perf_counter() - t0)) - - def start(self, name: str) -> None: - if self.enabled: - self._open[name] = time.perf_counter() - - def stop(self, name: str) -> None: - if not self.enabled: - return - t0 = self._open.pop(name, None) - if t0 is not None: - self.records.append((name, time.perf_counter() - t0)) - - def summary(self) -> str: - if not self.records: - return "" - # Aggregate repeated stage names (the per-rotation loop visits the - # translation stages once per candidate rotation). - agg: dict[str, list[float]] = {} - for name, dt in self.records: - agg.setdefault(name, []).append(dt) - total = sum(sum(v) for v in agg.values()) - lines = [ - f"{'stage':<32s} {'count':>5s} {'wall_s':>10s} {'%':>6s}", - "-" * 60, - ] - for name, vs in agg.items(): - wall = sum(vs) - lines.append( - f"{name:<32s} {len(vs):>5d} {wall:>10.3f} " - f"{100 * wall / total:>5.1f}%" - ) - lines.append("-" * 60) - lines.append(f"{'TOTAL':<32s} {'':>5s} {total:>10.3f} 100.0%") - return "\n".join(lines) - - -# --------------------------------------------------------------------------- -# Internal helpers -# --------------------------------------------------------------------------- - - -def _external_rwork(model: "ModelFT", data: "ReflectionData") -> float: - """Full-resolution scaled R-work via the standard Scaler. - - The TF + local refine work in analytical-scale R-factor (which ranks - candidates correctly but isn't the user-facing R-work). We compute the - proper Scaler-fit R-work once per finalist. - """ - from ...base.metrics.rfactor import rfactor_work_free - from ...scaling import Scaler - - # Build the Scaler on the model's device so that its anisotropy U - # tensor and per-bin scales land alongside `data.hkl`/`model(hkl)` — - # otherwise Scaler.forward's `matmul(self.s, U)` mixes CPU/GPU and - # crashes at refine_lbfgs. - s = Scaler(model=model, data=data, nbins=20, verbose=0, - device=model.xyz().device) - # Detach the model forward — the scaler only needs gradients through its - # own parameters; leaving `fc` attached to the model's autograd graph - # keeps SfFFT density-build intermediates alive after this function - # returns. - with torch.no_grad(): - fc = model(data.hkl).detach() - s.initialize(fc) - s.refine_lbfgs(fcalc=fc) - with torch.no_grad(): - # rfactor_work_free takes already-scaled amplitudes, not complex F_calc. - rw, _ = rfactor_work_free(data, torch.abs(s.forward(fc))) - return rw.item() if hasattr(rw, "item") else float(rw) - - -class _DirectModelEvaluator: - """Returns ``F_p1(hkl)`` of a P1-spacegroup model at integer HKL. - - The translation search asks its evaluator for ``F`` at a list of rotated - Miller indices. The rotation is already baked into the model's coordinates - by the time this is built, so ``R`` is ignored and every call is a direct - structure-factor evaluation rather than an interpolation. - """ - - def __init__(self, m: "ModelFT") -> None: - self._m = m - self.device = m.xyz().device - - def evaluate(self, R, hkl, real_cell, return_amplitude=False): - hkl_int = hkl.round().to(torch.int64).to(self.device) - with torch.no_grad(): - f = self._m(hkl_int) - return f.abs() if return_amplitude else f - - -# --------------------------------------------------------------------------- -# FRF input preparation (shared by the live pipeline and the rotation-ranking -# benchmark in tests/integration/alignment/benchmark_rotation_ranking.py) -# --------------------------------------------------------------------------- - - -@dataclass -class FRFInputs: - """Prepared reflection arrays shared by the rotation-search stages. - - The resolution-masked, anisotropy-corrected reflection arrays plus the - overall anisotropy tensor. The rotation search reads ``U_aniso`` and - ``device``; the rest is there for anything scoring against the same - observations on the same footing. - - ``sig_F`` carries the same anisotropy correction as ``F_obs``, which is a - multiplicative factor, so ``F/sigma`` is unchanged by it. It is here because - the rotation function computes its measurement weight from the sigmas, and - the earlier code threw them away immediately afterwards. ``None`` when the - data carry no sigmas. - """ - F_obs: torch.Tensor # (N,) anisotropy-corrected amplitudes - sig_F: Optional[torch.Tensor] # (N,) their sigmas, same correction - hkl: torch.Tensor # (N, 3) integer Miller indices - s_vec: torch.Tensor # (N, 3) reciprocal-space Cartesian - s_mag: torch.Tensor # (N,) Å⁻¹ - centric: torch.Tensor # (N,) bool - U_aniso: torch.Tensor # (3, 3) Popov-Bourenkov U - device: torch.device - - -def _prepare_frf_inputs( - model: "ModelFT", - data: "ReflectionData", - *, - d_min: float, - d_max: float, - n_shells: int, - verbose: int = 0, -) -> FRFInputs: - """Prepare the reflection arrays the rotation-search stages share. - - Masks the observations to ``[d_min, d_max]`` and fits and applies the - overall anisotropy correction, so ``F_obs`` on the returned dataclass is - anisotropy-corrected. - """ - device = model.xyz().device - - F_obs = data.F.to(torch.float64).abs() - hkl_all = data.hkl - rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) - s_vec_all = hkl_all.to(torch.float64) @ rec_basis - s_mag_all = s_vec_all.norm(dim=-1) - keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min) - if keep.sum().item() < n_shells * 5: - raise ValueError( - f"Too few reflections ({keep.sum().item()}) in [{d_min},{d_max}] Å " - f"for {n_shells} shells; widen the resolution range." - ) - F_obs = F_obs[keep].to(device) - sig_F = getattr(data, "F_sigma", None) - if sig_F is not None: - sig_F = sig_F.to(torch.float64)[keep].to(device) - hkl = hkl_all[keep].to(device) - s_vec = s_vec_all[keep].to(device) - s_mag = s_mag_all[keep].to(device) - centric = ( - data.centric[keep].to(torch.bool).to(device) - if hasattr(data, "centric") - else torch.zeros_like(F_obs, dtype=torch.bool) - ) - - aniso_edges, _ = equal_count_shell_edges(s_mag, n_shells) - aniso_idx = assign_shells(s_mag, aniso_edges) - U_aniso = fit_overall_anisotropy( - F_obs, s_vec, aniso_idx, centric, P=n_shells, min_count=20, - ) - # Project U onto the point-group-invariant subspace (Phaser - # RefineANO.cc:116-142, via cctbx `site_symmetry.average_u_star`). An - # unconstrained six-component fit can return a tensor the lattice forbids, - # and applying that modulates the observations by a direction-dependent - # factor the crystal cannot have. After projection: cubic -> U = lambda I - # (one degree of freedom), tetragonal/trigonal/hexagonal -> diag(l, l, m), - # orthorhombic -> diag(l, m, n). - from .sh import hkl_symops_to_cartesian, symmetrize_anisotropy - _sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) - _sym_mats_cart = hkl_symops_to_cartesian(_sg_mats, rec_basis.to(device)) - U_aniso = symmetrize_anisotropy(U_aniso, _sym_mats_cart) - F_obs_aniso = apply_overall_anisotropy(F_obs, s_vec, U_aniso) - # Same multiplicative factor, so F/sigma survives the correction intact. - sig_F_aniso = (None if sig_F is None - else apply_overall_anisotropy(sig_F, s_vec, U_aniso)) - - return FRFInputs( - F_obs=F_obs_aniso, - sig_F=sig_F_aniso, - hkl=hkl, - s_vec=s_vec, - s_mag=s_mag, - centric=centric, - U_aniso=U_aniso, - device=device, - ) - - -# --------------------------------------------------------------------------- -# Public entry point -# --------------------------------------------------------------------------- - - -def align_model_to_data( - model: "ModelFT", - data: "ReflectionData", - *, - d_min: float = 4.0, - d_max: float = 15.0, - n_shells: int = 20, - n_rotation_peaks: int = 500, - verbose: int = 0, - do_translation: bool = True, - n_translation_peaks: int = 20, - n_translation_candidates: int = 3, - translation_grid_steps: int = 16, - n_rotation_candidates: int = 15, - use_llg_tf: bool = False, - tf_d_min: Optional[float] = None, - tf_d_max: Optional[float] = None, - model_error_A: Optional[float] = None, -) -> "ModelFT": - """Place ``model`` in ``data``'s crystal: rotation search, then translation. - - Returns a new rotated+translated ``ModelFT`` carrying - ``last_alignment_rotation``, ``last_alignment_translation`` and - ``last_alignment_rfactor`` provenance attributes. It is a *placement*, not a - refined structure -- refine it downstream. - - `MolecularReplacementPipeline` is the implementation of record; this - function returns its single best solution. - """ - if not model.ctx.initialized: - raise RuntimeError( - "Cannot fit an uninitialized ModelFT. Load PDB data first." - ) - - # Imported lazily to avoid an import cycle: `pipeline` imports the stage - # helpers (`_prepare_frf_inputs`, `_external_rwork`, - # `_DirectModelEvaluator`, `_StageTimer`) from this module. - from .pipeline import MolecularReplacementPipeline - - pipeline = MolecularReplacementPipeline( - data, model, - device=model.xyz().device, - verbose=verbose, - d_min=d_min, d_max=d_max, n_shells=n_shells, - n_rotation_peaks=n_rotation_peaks, - model_error_A=model_error_A, - n_rotation_candidates=n_rotation_candidates, - n_translation_peaks=n_translation_peaks, - n_translation_candidates=n_translation_candidates, - translation_grid_steps=translation_grid_steps, - use_llg_tf=use_llg_tf, - tf_d_min=tf_d_min, tf_d_max=tf_d_max, - ) - solutions = pipeline.run(do_translation=do_translation) - return solutions[0].model diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 5eb9bd32..5559a9da 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -34,6 +34,8 @@ from __future__ import annotations +import time +from contextlib import contextmanager from dataclasses import dataclass from typing import List, Optional, Tuple, TYPE_CHECKING @@ -43,16 +45,11 @@ from torchref.config import get_default_device from torchref.utils.device_mixin import DeviceMixin -from .align import ( - _DirectModelEvaluator, - _StageTimer, - _external_rwork, - _prepare_frf_inputs, -) from .frf.rotation_utils import rotation_matrix_from_edmonds_euler from .frf.types import RotationPeak -from .rotation_search import search_peaks +from .rotation_search import prepare_frf_inputs, search_peaks from .translation import ( + DirectModelEvaluator, TranslationObs, TranslationPeak, amplitude_translation_search, @@ -145,6 +142,101 @@ def cluster_rotation_peaks( return clustered +# --------------------------------------------------------------------------- +# Stage timing and the user-facing R-work +# --------------------------------------------------------------------------- + + +class _StageTimer: + """Lightweight wall-clock accumulator. Gated by ``verbose >= 2``. + + Two interleavable usages: + * ``with t.stage(name):`` block — records the block's wall time. + * ``t.start(name)`` / ``t.stop(name)`` — checkpoint pair, no indent. + + The summary table prints stages aggregated by name; per-rotation loop + stages (translation search etc.) get aggregated counts. + """ + + def __init__(self, enabled: bool): + self.enabled = enabled + self.records: list[tuple[str, float]] = [] + self._open: dict[str, float] = {} + + @contextmanager + def stage(self, name: str): + if not self.enabled: + yield + return + t0 = time.perf_counter() + try: + yield + finally: + self.records.append((name, time.perf_counter() - t0)) + + def start(self, name: str) -> None: + if self.enabled: + self._open[name] = time.perf_counter() + + def stop(self, name: str) -> None: + if not self.enabled: + return + t0 = self._open.pop(name, None) + if t0 is not None: + self.records.append((name, time.perf_counter() - t0)) + + def summary(self) -> str: + if not self.records: + return "" + # Aggregate repeated stage names (the per-rotation loop visits the + # translation stages once per candidate rotation). + agg: dict[str, list[float]] = {} + for name, dt in self.records: + agg.setdefault(name, []).append(dt) + total = sum(sum(v) for v in agg.values()) + lines = [ + f"{'stage':<32s} {'count':>5s} {'wall_s':>10s} {'%':>6s}", + "-" * 60, + ] + for name, vs in agg.items(): + wall = sum(vs) + lines.append( + f"{name:<32s} {len(vs):>5d} {wall:>10.3f} " + f"{100 * wall / total:>5.1f}%" + ) + lines.append("-" * 60) + lines.append(f"{'TOTAL':<32s} {'':>5s} {total:>10.3f} 100.0%") + return "\n".join(lines) + + +def _external_rwork(model: "ModelFT", data: "ReflectionData") -> float: + """Full-resolution scaled R-work via the standard Scaler. + + The TF + local refine work in analytical-scale R-factor (which ranks + candidates correctly but isn't the user-facing R-work). We compute the + proper Scaler-fit R-work once per finalist. + """ + from ...base.metrics.rfactor import rfactor_work_free + from ...scaling import Scaler + + # No device override: the Scaler takes the configured default, which is + # the one place a device is decided. Reading it off whichever tensor is + # nearest is what puts a run on two devices at once. + s = Scaler(model=model, data=data, nbins=20, verbose=0) + # Detach the model forward — the scaler only needs gradients through its + # own parameters; leaving `fc` attached to the model's autograd graph + # keeps SfFFT density-build intermediates alive after this function + # returns. + with torch.no_grad(): + fc = model(data.hkl).detach() + s.initialize(fc) + s.refine_lbfgs(fcalc=fc) + with torch.no_grad(): + # rfactor_work_free takes already-scaled amplitudes, not complex F_calc. + rw, _ = rfactor_work_free(data, torch.abs(s.forward(fc))) + return rw.item() if hasattr(rw, "item") else float(rw) + + @dataclass class MRSolution: """A molecular-replacement placement. @@ -300,7 +392,7 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: timer = self._timer timer.start("0_data_prep") - frf = _prepare_frf_inputs( + frf = prepare_frf_inputs( self.model, self.data, d_min=self.d_min, d_max=self.d_max, n_shells=self.n_shells, verbose=self.verbose, @@ -537,7 +629,7 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: rotated_k.spacegroup = data.spacegroup.hm rotated_p1 = rotated_k.copy() rotated_p1.spacegroup = "P 1" - evaluator = _DirectModelEvaluator(rotated_p1) + evaluator = DirectModelEvaluator(rotated_p1) timer.start("5_precompute_G") G_pre, h_R_pre = precompute_G_for_rotation( @@ -650,3 +742,59 @@ def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): ) for i in order ] + + +# --------------------------------------------------------------------------- +# Public entry point +# --------------------------------------------------------------------------- + + +def align_model_to_data( + model: "ModelFT", + data: "ReflectionData", + *, + d_min: float = 4.0, + d_max: float = 15.0, + n_shells: int = 20, + n_rotation_peaks: int = 500, + verbose: int = 0, + do_translation: bool = True, + n_translation_peaks: int = 20, + n_translation_candidates: int = 3, + translation_grid_steps: int = 16, + n_rotation_candidates: int = 15, + use_llg_tf: bool = False, + tf_d_min: Optional[float] = None, + tf_d_max: Optional[float] = None, + model_error_A: Optional[float] = None, +) -> "ModelFT": + """Place ``model`` in ``data``'s crystal: rotation search, then translation. + + Returns a new rotated+translated ``ModelFT`` carrying + ``last_alignment_rotation``, ``last_alignment_translation`` and + ``last_alignment_rfactor`` provenance attributes. It is a *placement*, not a + refined structure -- refine it downstream. + + `MolecularReplacementPipeline` is the implementation of record; this + function returns its single best solution. + """ + if not model.ctx.initialized: + raise RuntimeError( + "Cannot fit an uninitialized ModelFT. Load PDB data first." + ) + + pipeline = MolecularReplacementPipeline( + data, model, + verbose=verbose, + d_min=d_min, d_max=d_max, n_shells=n_shells, + n_rotation_peaks=n_rotation_peaks, + model_error_A=model_error_A, + n_rotation_candidates=n_rotation_candidates, + n_translation_peaks=n_translation_peaks, + n_translation_candidates=n_translation_candidates, + translation_grid_steps=translation_grid_steps, + use_llg_tf=use_llg_tf, + tf_d_min=tf_d_min, tf_d_max=tf_d_max, + ) + solutions = pipeline.run(do_translation=do_translation) + return solutions[0].model diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index feae55a5..9089794b 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -28,6 +28,7 @@ import torch +from torchref.config import get_default_device from torchref.scaling.weighting import (DEFAULT_SNR_CAP, DEFAULT_TRUST_CAP) from .sh import ( @@ -42,7 +43,8 @@ from ...model.model_ft import ModelFT from .frf.types import RotationPeak -__all__ = ["RotationSolutions", "rotation_search"] +__all__ = ["FRFInputs", "RotationSolutions", "prepare_frf_inputs", + "rotation_search"] # Note for anyone reaching for the constants below programmatically: the package # re-exports `rotation_search` (the function) under this module's own name, so @@ -213,6 +215,87 @@ def fit_anisotropy( return symmetrize_anisotropy(U, sym_cart) + +@dataclass +class FRFInputs: + """The observations the rotation search runs on, masked and corrected. + + ``F_obs`` is anisotropy-corrected. ``sig_F`` carries the same correction, + which is a multiplicative factor, so ``F/sigma`` survives it unchanged -- it + is here because the engine builds its measurement weight from the sigmas and + the earlier code discarded them immediately after the Wilson step. ``None`` + when the data carry no sigmas. + """ + + F_obs: torch.Tensor # (N,) anisotropy-corrected amplitudes + sig_F: Optional[torch.Tensor] # (N,) their sigmas, same correction + hkl: torch.Tensor # (N, 3) integer Miller indices + s_vec: torch.Tensor # (N, 3) reciprocal-space Cartesian + s_mag: torch.Tensor # (N,) inverse Angstrom + centric: torch.Tensor # (N,) bool + U_aniso: torch.Tensor # (3, 3) Popov-Bourenkov U + device: torch.device + + +def prepare_frf_inputs( + model: "ModelFT", + data: "ReflectionData", + *, + d_min: float, + d_max: float, + n_shells: int, + verbose: int = 0, +) -> FRFInputs: + """Mask the observations to ``[d_min, d_max]`` and correct their anisotropy. + + The anisotropy tensor comes from :func:`fit_anisotropy`, which is also what + the public :func:`rotation_search` uses. It used to be refitted here by a + second copy of the same six lines over the same window -- two paths to one + number is how they drift apart. + + Everything lands on the configured default device, not on whichever device + ``model`` happens to sit on. + """ + device = get_default_device() + + F_obs = data.F.to(torch.float64).abs() + hkl_all = data.hkl + rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) + s_vec_all = hkl_all.to(torch.float64) @ rec_basis + s_mag_all = s_vec_all.norm(dim=-1) + keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min) + if keep.sum().item() < n_shells * 5: + raise ValueError( + f"Too few reflections ({keep.sum().item()}) in [{d_min},{d_max}] A " + f"for {n_shells} shells; widen the resolution range." + ) + F_obs = F_obs[keep].to(device) + sig_F = getattr(data, "F_sigma", None) + if sig_F is not None: + sig_F = sig_F.to(torch.float64)[keep].to(device) + hkl = hkl_all[keep].to(device) + s_vec = s_vec_all[keep].to(device) + s_mag = s_mag_all[keep].to(device) + centric = ( + data.centric[keep].to(torch.bool).to(device) + if hasattr(data, "centric") + else torch.zeros_like(F_obs, dtype=torch.bool) + ) + + U_aniso = fit_anisotropy( + data, d_min=d_min, d_max=d_max, n_shells=n_shells, + ).to(device) + F_obs_aniso = apply_overall_anisotropy(F_obs, s_vec, U_aniso) + # Same multiplicative factor, so F/sigma survives the correction intact. + sig_F_aniso = (None if sig_F is None + else apply_overall_anisotropy(sig_F, s_vec, U_aniso)) + + return FRFInputs( + F_obs=F_obs_aniso, sig_F=sig_F_aniso, hkl=hkl, s_vec=s_vec, + s_mag=s_mag, centric=centric, U_aniso=U_aniso, device=device, + ) + + def search_peaks( model: "ModelFT", data: "ReflectionData", diff --git a/torchref/experimental/alignment/sh.py b/torchref/experimental/alignment/sh.py index 12e97ea3..6a052d18 100644 --- a/torchref/experimental/alignment/sh.py +++ b/torchref/experimental/alignment/sh.py @@ -1,31 +1,34 @@ -""" -Pure-PyTorch spherical harmonic expansion for the alignment module. - -Conventions (locked, asserted by tests/unit/alignment/test_sh.py): - - Y_{l,m}(θ, φ) = (-1)^m · √[(2l+1)/(4π) · (l-m)!/(l+m)!] · P_l^m(cos θ) · e^{imφ} (m ≥ 0) - Y_{l,-m}(θ, φ) = (-1)^m · conj(Y_{l,m}(θ, φ)) (m > 0) - -i.e. fully orthonormal physics convention with Condon-Shortley phase included. -Matches scipy.special.sph_harm and the convention used in Edmonds, Sakurai, etc. - -Numerical core is a stable forward recurrence on the fully-normalized associated -Legendre `bar_P_l^m(cosθ) = √[(2l+1)/(4π) · (l-m)!/(l+m)!] · P_l^m(cosθ)` so we -never form (2l)! explicitly. +"""Leaf mathematics the rotation function needs: Legendre seeds, shells, anisotropy. + +Three unrelated things share this module because they share one consumer. + +* **Legendre recurrence coefficients** and their seed, for the fully-normalised + associated Legendre ``bar_P_l^m(cos theta)``, which the Bessel-SH expansion in + :mod:`~torchref.experimental.alignment.frf.data_mr` and its compiled kernels + build on. Normalised so ``(2l)!`` is never formed explicitly. +* **Equal-count resolution shells** (:func:`equal_count_shell_edges`, + :func:`assign_shells`). Assigned once and passed down: two consumers deriving + their own edges from the same ``|s|`` disagree about the reflections sitting on + a boundary. +* **Overall anisotropy** -- fit in intensity space, projected onto the point + group, applied to amplitudes with the half exponent. The projection is + load-bearing: an unconstrained six-component fit can return a tensor the + lattice forbids. + +This module used to also carry a full spherical-harmonic expansion +(``evaluate_ylm``, ``sh_expand_ball``). Nothing called it -- the FRF's own +expansion superseded it -- so it went. ``_bar_legendre_recurrence`` survived it: +production does not call that either, but the FRF expansion's only *independent* +test reference is built on it. """ from __future__ import annotations import math -import os -import time from typing import Optional, Tuple import torch -_PROFILE = bool(os.environ.get("FRF_PROFILE")) -_YLM_PROF = {"recurrence": 0.0, "assembly": 0.0} - def legendre_recurrence_coefficients(L: int, dtype, device): """Coefficient tables for the fully-normalised Legendre recurrence. @@ -75,9 +78,17 @@ def _bar_legendre_recurrence( L: int, keep_l: Optional[torch.Tensor] = None, ) -> torch.Tensor: - """ - Compute fully-normalized associated Legendre `bar_P_l^m(cos θ)` for - all l in [0, L), m in [0, l]. + """Fully-normalised associated Legendre ``bar_P_l^m(cos theta)``, all l < L, m <= l. + + **Kept for its test, deliberately.** Production does not call this: the + Bessel-SH expansion runs the same recurrence inside its compiled kernels, + from :func:`legendre_recurrence_coefficients` and :data:`LEGENDRE_SEED`. + That is exactly why this stays -- it is a second, independent, pure-torch + implementation, and ``tests/unit/frf_separate/test_bessel_sh_grouping.py`` + builds a slow reference expansion on it to check the fused one. The other + tests there compare ``bessel_sh_expand`` against *itself* at a different + grouping, so a dropped term cancels; one did, and only this reference caught + it. Deleting it would leave the expansion checked only against itself. Definition: bar_P_l^m(x) = √[(2l+1)/(4π) · (l-m)!/(l+m)!] · P_l^m(x) @@ -150,148 +161,6 @@ def _bar_legendre_recurrence( return out -def evaluate_ylm( - theta: torch.Tensor, - phi: torch.Tensor, - L: int, - l_indices: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """ - Evaluate Y_{l,m}(θ, φ) for all (l, m) with l ∈ [0, L), m ∈ [-(L-1), L-1]. - - Parameters - ---------- - theta : torch.Tensor - Polar angle, shape (...,), values in [0, π]. - phi : torch.Tensor - Azimuthal angle, shape (...,), values in [0, 2π). - L : int - Maximum SH degree (exclusive: l_max = L - 1). - l_indices : torch.Tensor, optional - If given (1-D long tensor of l values), assemble and return Y only for - those degrees, shape ``(..., len(l_indices), 2L-1)`` with row ``i`` - holding ``Y_{l_indices[i], m}``. The Legendre recurrence still runs over - the full degree range (it is a recurrence), but the costly complex Y - assembly is restricted to the requested rows. Used by the FRF expansion, - which only needs even degrees (the odd-l and l=0 rows are zeroed by - Patterson centrosymmetry) — halving the dominant assembly cost. - - Returns - ------- - Y : torch.Tensor, complex - Shape (..., L, 2L-1) (or (..., len(l_indices), 2L-1) if l_indices given). - ``Y[..., l, L-1+m] = Y_{l,m}(θ, φ)`` for |m| ≤ l, zero otherwise. dtype is - complex128 if input is float64, else complex64. - """ - assert theta.shape == phi.shape, "theta and phi must have the same shape" - - real_dtype = theta.dtype - if real_dtype == torch.float64: - complex_dtype = torch.complex128 - elif real_dtype == torch.float32: - complex_dtype = torch.complex64 - else: - raise TypeError(f"Unsupported real dtype: {real_dtype}") - - device = theta.device - cos_theta = torch.cos(theta) - sin_theta = torch.sin(theta).clamp(min=0.0) # numerical floor at the poles - - if _PROFILE: - t0 = time.perf_counter() - bar_P = _bar_legendre_recurrence(cos_theta, sin_theta, L) # (..., L, L) - if l_indices is not None: - bar_P = bar_P[..., l_indices, :] # (..., n_sel, L) — even rows only - n_rows = bar_P.shape[-2] - if _PROFILE: - _YLM_PROF["recurrence"] += time.perf_counter() - t0 - t0 = time.perf_counter() - - # Y_{l,m}(θ,φ) = (-1)^m · bar_P_l^m(cosθ) · e^{i m φ} for m ≥ 0 - # Y_{l,-m} = (-1)^m · conj(Y_{l,m}) for m > 0 - Y = torch.zeros((*theta.shape, n_rows, 2 * L - 1), dtype=complex_dtype, device=device) - - # Precompute e^{i m φ} for m = 0..L-1 and the Condon-Shortley signs (-1)^m. - m_vals = torch.arange(L, dtype=real_dtype, device=device) - m_phi = phi.unsqueeze(-1) * m_vals # (..., L) - expo = torch.complex(torch.cos(m_phi), torch.sin(m_phi)) # (..., L), e^{i m φ} - signs = ((-1.0) ** m_vals).to(complex_dtype) # (L,) - - # Fill m ≥ 0 columns (L-1 .. 2L-2): Y_{l,m} = (-1)^m bar_P_l^m e^{imφ}, - # vectorised over (l, m). phase = (-1)^m e^{imφ} broadcasts over l. - phase_pos = (signs * expo).unsqueeze(-2) # (..., 1, L) - Y_pos = bar_P.to(complex_dtype) * phase_pos # (..., n_rows, L) - Y[..., :, L - 1:] = Y_pos - - # Fill m < 0 columns by hermitian symmetry: Y_{l,-m} = (-1)^m conj(Y_{l,m}). - # For m = 1..L-1 these land in columns L-2 .. 0, i.e. the reversed prefix. - neg = signs[1:] * torch.conj(Y_pos[..., :, 1:]) # (..., L, L-1), m = 1..L-1 - Y[..., :, : L - 1] = torch.flip(neg, dims=(-1,)) - - if _PROFILE: - _YLM_PROF["assembly"] += time.perf_counter() - t0 - return Y - - -def angular_density_weights( - s_vectors: torch.Tensor, - k_neighbors: int = 12, -) -> torch.Tensor: - """ - Per-sample weights that compensate for non-uniform angular sampling on the - unit sphere. Returns w_i ∝ (1 / local_density)^... so the weighted sum - `Σ_i w_i · v_i · Y*_lm(ŝ_i)` is an unbiased Monte-Carlo estimate of the - SH integral on the sphere. - - Heuristic: w_i ~ (k-th NN great-circle distance)². Normalised so that - Σ w_i = N. - - Pure-torch O(N · k) memory; suitable up to N ~ 30k. For larger N use a - chunked KNN, but typical resolution-cut datasets fit easily. - - Parameters - ---------- - s_vectors : torch.Tensor, shape (N, 3) - k_neighbors : int, default 12 - Number of nearest angular neighbours to estimate local density. - - Returns - ------- - w : torch.Tensor, shape (N,) - """ - device = s_vectors.device - dtype = s_vectors.dtype - N = s_vectors.shape[0] - norm = s_vectors.norm(dim=-1).clamp(min=1e-30) - s_hat = s_vectors / norm.unsqueeze(-1) # (N, 3) - # cos(angle) between every pair via dot product - # Memory: (N, N). For N ~ 30k, ~3 GB at fp32 — chunk if too big. - if N <= 8000: - dots = (s_hat @ s_hat.transpose(0, 1)).clamp(-1.0, 1.0) - ang = torch.acos(dots) # (N, N) - # k-th NN distance (excluding self at column-diagonal). topk smallest. - # Set diagonal large so it doesn't show up as nearest. - ang.fill_diagonal_(float("inf")) - kth_dist, _ = torch.topk(ang, k_neighbors, dim=-1, largest=False) # (N, k) - d_local = kth_dist[:, -1] # k-th NN distance - else: - # Chunked: still O(N²) compute but bounded memory. - chunk = 1024 - d_local = torch.empty(N, dtype=dtype, device=device) - for i0 in range(0, N, chunk): - i1 = min(i0 + chunk, N) - dots = (s_hat[i0:i1] @ s_hat.transpose(0, 1)).clamp(-1.0, 1.0) - ang = torch.acos(dots) - for j, gi in enumerate(range(i0, i1)): - ang[j, gi] = float("inf") - kth_dist, _ = torch.topk(ang, k_neighbors, dim=-1, largest=False) - d_local[i0:i1] = kth_dist[:, -1] - - w = d_local ** 2 # ~ local Voronoi area - w = w * (N / w.sum().clamp(min=1e-30)) # normalise to Σw = N - return w - - def get_axis_order(sym_mats: torch.Tensor, axis: int) -> int: """ Order of the highest-multiplicity proper rotation about a principal axis. @@ -345,115 +214,6 @@ def get_high_order_axis(sym_mats: torch.Tensor) -> Tuple[int, int]: return axis, orders[axis] -def sh_expand_ball( - s_vectors: torch.Tensor, - values: torch.Tensor, - shell_idx: torch.Tensor, - P: int, - L: int, - enforce_friedel: bool = True, - chunk_size: int = 2048, - angular_weights: Optional[torch.Tensor] = None, - zsymm: int = 1, - skip_odd_l: bool = False, -) -> torch.Tensor: - """ - Analytical spherical-harmonic expansion of a scattered-point real field - on a set of radial shells. - - f_{p,l,m} = Σ_{i ∈ shell p} values_i · conj(Y_{l,m}(θ_i, φ_i)) - - When `enforce_friedel=True` the input is augmented with the antipodal copy - `(-s_i, values_i)`. Y_{l,m}(-ŝ) = (-1)^l Y_{l,m}(ŝ), so the sum then has - f_{p,l,m} = (1 + (-1)^l) · Σ_i v_i · Y*_{l,m}(ŝ_i) - i.e. odd-l rows are exactly zero by construction (and even-l rows get a - factor of 2 which we keep — this absorbs into the cross-correlation - normalisation when the same convention is applied to both operands). - - Parameters - ---------- - s_vectors : torch.Tensor - Reciprocal-lattice vectors, shape (N, 3). Direction only is used; - magnitudes do not enter (shell assignment is done by the caller). - values : torch.Tensor - Real-valued samples (e.g. |E(h)|), shape (N,). - shell_idx : torch.Tensor (int64) - Shell index in [0, P) for each reflection, shape (N,). - P : int - Number of radial shells. - L : int - SH bandlimit (l in [0, L)). - enforce_friedel : bool, default True - Augment input with (-s, value) pairs and zero odd-l coefficients. - chunk_size : int - Points per chunk for memory control during Y_lm evaluation. - - Returns - ------- - f_plm : torch.Tensor, complex - Shape (P, L, 2L-1). `f_plm[p, l, L-1+m] = f_{p,l,m}`. - """ - assert s_vectors.dim() == 2 and s_vectors.shape[-1] == 3 - assert values.dim() == 1 and values.shape[0] == s_vectors.shape[0] - assert shell_idx.dim() == 1 and shell_idx.shape[0] == s_vectors.shape[0] - - real_dtype = s_vectors.dtype - if real_dtype == torch.float64: - complex_dtype = torch.complex128 - elif real_dtype == torch.float32: - complex_dtype = torch.complex64 - else: - raise TypeError(f"Unsupported dtype {real_dtype}") - device = s_vectors.device - - if enforce_friedel: - s_vectors = torch.cat([s_vectors, -s_vectors], dim=0) - values = torch.cat([values, values], dim=0) - shell_idx = torch.cat([shell_idx, shell_idx], dim=0) - if angular_weights is not None: - angular_weights = torch.cat([angular_weights, angular_weights], dim=0) - - if angular_weights is not None: - values = values * angular_weights.to(values.dtype) - - # Direction (θ, φ). At |s|=0 the direction is undefined; the caller should - # have excluded F(000), but we guard anyway. - norm = s_vectors.norm(dim=-1).clamp(min=1e-30) - s_hat = s_vectors / norm.unsqueeze(-1) - cos_theta = s_hat[..., 2].clamp(min=-1.0, max=1.0) - theta = torch.acos(cos_theta) - phi = torch.atan2(s_hat[..., 1], s_hat[..., 0]) - - f_plm = torch.zeros((P, L, 2 * L - 1), dtype=complex_dtype, device=device) - - N = s_vectors.shape[0] - for start in range(0, N, chunk_size): - stop = min(start + chunk_size, N) - Y = evaluate_ylm(theta[start:stop], phi[start:stop], L) # (n, L, 2L-1) - # contribution to f_{p,l,m} is value_i * conj(Y_{l,m}(s_i)) - contrib = values[start:stop].to(complex_dtype).view(-1, 1, 1) * torch.conj(Y) - # scatter-add into shells - f_plm.index_add_(0, shell_idx[start:stop], contrib) - - if enforce_friedel or skip_odd_l: - # Zero odd-l rows explicitly (they should already be ~0; this kills FP drift - # and is the only required step when skip_odd_l is set without Friedel). - l_vals = torch.arange(L, device=device) - odd_mask = (l_vals % 2 == 1) - f_plm[:, odd_mask, :] = 0.0 - - # F1: m-symmetry filter (Phaser DataMR.cc:1019 / 1117). The Patterson is - # invariant under the spacegroup rotation operators, so SH coefficients - # whose m-index violates the highest-order rotation axis are pure noise. - # Zero them out post-expansion. With `zsymm=1` (no filter) this is a no-op. - if zsymm > 1: - m_vals = torch.arange(-(L - 1), L, device=device) # (2L-1,) - m_invalid = (m_vals.abs() % zsymm) != 0 - f_plm[:, :, m_invalid] = 0.0 - - return f_plm - - def equal_count_shell_edges( s_magnitudes: torch.Tensor, P: int, diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index 9b0ceed6..c66345ef 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -24,6 +24,7 @@ import numpy as np import torch +from torchref.config import get_default_device from torchref.scaling import WilsonNormaliser from torchref.scaling.weighting import (inverse_variance_weight, normalise_weight, snr_from_amplitude) @@ -31,7 +32,10 @@ from .distributions import rice_log_likelihood, woolfson_log_likelihood from .sh import assign_shells, equal_count_shell_edges from dataclasses import dataclass -from typing import List, Optional, Tuple +from typing import List, Optional, Tuple, TYPE_CHECKING + +if TYPE_CHECKING: # pragma: no cover - typing only + from ...model.model_ft import ModelFT #: Chebyshev order of the Wilson fit. Matches the rotation function's @@ -119,7 +123,7 @@ def build( half of the variance budget through the Luzzati falloff. The same number the rotation function weights with. """ - dev = device if device is not None else F_obs.device + dev = get_default_device() if device is None else device real = torch.float64 F = F_obs.detach().to(dev) F = (F.abs() if F.is_complex() else F).to(real) @@ -163,6 +167,33 @@ def build( ) +class DirectModelEvaluator: + """Returns ``F_p1(hkl)`` of a P1-spacegroup model at integer HKL. + + The translation search asks its evaluator for ``F`` at a list of rotated + Miller indices. The rotation is already baked into the model's coordinates + by the time this is built, so ``R`` is ignored and every call is a direct + structure-factor evaluation rather than an interpolation. + + Answers on the **configured default device**, whatever device the model + happens to sit on. Reading the device off the model instead is how the + translation stage ends up split across two devices when a caller builds a + CPU model on a host with an accelerator: the model answers on the CPU while + everything derived from config answers on the GPU. + """ + + def __init__(self, m: "ModelFT") -> None: + self._m = m + self.device = get_default_device() + + def evaluate(self, R, hkl, real_cell, return_amplitude=False): + hkl_int = hkl.round().to(torch.int64).to(self._m.xyz().device) + with torch.no_grad(): + f = self._m(hkl_int) + f = f.to(self.device) + return f.abs() if return_amplitude else f + + @dataclass class TranslationPeak: """ @@ -302,7 +333,7 @@ def amplitude_translation_search( peaks : list of TranslationPeak Top-`n_peaks` peaks sorted by descending correlation. """ - device = getattr(interpolator, "device", obs.hkl.device) + device = get_default_device() real_dtype = torch.float64 complex_dtype = torch.complex128 @@ -584,7 +615,7 @@ def precompute_G_for_rotation( real_dtype = torch.float64 complex_dtype = torch.complex128 if device is None: - device = getattr(interpolator, "device", hkl.device) + device = get_default_device() hkl_t = hkl.detach().to(device).to(real_dtype) sym_R = spacegroup.matrices.detach().to(device).to(real_dtype) @@ -651,7 +682,7 @@ def local_translation_refine( ignored -- the FFT evaluates the whole fine grid at once, so there is nothing for a second zoom pass to buy. """ - device = getattr(interpolator, "device", obs.hkl.device) + device = get_default_device() real_dtype = torch.float64 complex_dtype = torch.complex128 From b99d87708d53607a7d17acac6709e73582890ba0 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 31 Aug 2026 16:41:42 +0200 Subject: [PATCH 114/250] Record the alignment refactor in the changelog Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- docs/changelog.rst | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/docs/changelog.rst b/docs/changelog.rst index 579d171c..6c670b7f 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,13 @@ Changelog Unreleased ---------- +- Molecular replacement is now a rotation search feeding a translation search and nothing else; the pipeline returns a placement and stops. End-to-end pose recovery over 10 structures x 3 seeds went 18/30 to 30/30, at about a sixth of the wall clock +- Removed the ML rescore from between the two searches. It reordered a shortlist that already contained the answer, and cost 6 of 30 placements +- Removed the post-placement dense rotation re-sampling and rigid-body polish. They refined a correct placement away from truth on 2DQ6, 3GR5 and 4BX9; refining a placement is downstream refinement's job +- The translation search weights reflections by inverse variance, which it previously did not do at all, and both searches normalise through one shared Wilson fit built once per run instead of five private ones. This is what recovered 6G9X +- The translation search's resolution window is a parameter (``tf_d_min``/``tf_d_max``) rather than a docstring; it defaults to the existing behaviour of no cut +- The alignment package takes its device from the configured default throughout, instead of reading it off whichever model or tensor was nearest +- Removed the alignment package's unreachable modules: quaternion transforms, a second Wigner implementation, clash scoring, vector sampling, the Lattman-Love interpolator, the E-value convention layer and its French-Wilson posterior, and the unused half of the spherical-harmonic expansion - Fixed the reciprocal-space symmetry convention in the alignment package (``h.S``, not ``S.h``) - Fixed ``hkl_symops_to_cartesian`` returning non-rotations in trigonal and hexagonal settings, which corrupted the anisotropy projection - Fixed the overall-anisotropy fit, which regressed log intensities with no constant term and so absorbed the ``-gamma`` offset into the tensor From 25d6b33f69ca772cf355235f0556a892e46a012c Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 31 Aug 2026 18:32:07 +0200 Subject: [PATCH 115/250] Stop the Wilson fit chasing eleven digits, and take the default dtype The fit needed 26-102 IRLS iterations for six coefficients. That was never a convergence-rate problem -- the objective is at its optimum to six figures by iteration five, and the rest is the coefficients wandering in a flat direction while the curve does not move. The cause was the stopping rule: |dL| per reflection against 1e-10, on an objective of order 6-12 per reflection, is a demand for eleven significant digits from a normalisation curve. The criterion is now relative, and relative to the improvement so far rather than to |L|. That distinction is why it was absolute before: under I -> cI the objective gains an additive log(c)*sum(k), so |dL|/|L| means something different at every scale. The additive term cancels in any DIFFERENCE, so a ratio of two differences is both relative and scale invariant. At rtol=1e-4 the fit converges in 8-12 iterations, and the iteration count is flat from 1e-3 to 1e-8 -- the old criterion was not a tighter setting of the same knob, it was an unreachable one. = 1 no longer depends on where the fit stopped. It is the intercept's own score equation and has a closed form, so it is solved directly after the loop; it now holds to 1e-9..2e-7 regardless of rtol, where before it degraded with the tolerance. X^T W X is built and factorised once. For a Gamma with a log link the IRLS weight is the shape, independent of mu, so that matrix is identical at every iteration and was being rebuilt with an O(N n^2) pass each time. The fit takes the configured float dtype instead of hardcoding float64, which was the only double-precision path in the scaling package. Measured over the full 10x3 panel at fixed rtol, float32 and float64 give *identical* placements on every cell, at 1.8x the speed. The mu floor moves off 1e-300, a float64 constant that flushes to zero in float32 and turns the guard into the division it exists to prevent. The likelihood translation function's calculated side goes through the same fit, via normalise_calc, replacing the last two per-shell normalisations in the package -- which were also two copies of one calculation. A few ms per candidate, paid only when use_llg_tf is on. Test tolerances move to the package's 1e-4 relative and are set from measurement rather than from float64 habit; one of them was nondeterministic because _standard_gamma ignores the generator it was handed, so it passed alone and failed in a suite. Unit suite 1940 passed, and 968s -> 485s. Panel: 30/30 -> 29/30 on both arms, one cell, McNemar p = 1.0. It is 2DQ6 t1, and it is not a regression in the fit -- see the memo in the follow-up commit. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/diagnostics/calc_norm_cost.py | 146 +++++++++++++++ tests/unit/alignment/test_translation_obs.py | 11 +- tests/unit/scaling/test_wilson_normaliser.py | 27 ++- torchref/experimental/alignment/pipeline.py | 9 +- .../experimental/alignment/translation.py | 56 ++++-- torchref/scaling/wilson.py | 168 +++++++++++++----- 6 files changed, 339 insertions(+), 78 deletions(-) create mode 100644 alignment_lab/diagnostics/calc_norm_cost.py diff --git a/alignment_lab/diagnostics/calc_norm_cost.py b/alignment_lab/diagnostics/calc_norm_cost.py new file mode 100644 index 00000000..ac283745 --- /dev/null +++ b/alignment_lab/diagnostics/calc_norm_cost.py @@ -0,0 +1,146 @@ +"""What does it cost to normalise the LLG's calculated side with the shared fit? + +Two per-shell normalisations survive in the translation likelihood -- the +``E_calc`` of each candidate translation, and the ``E_calc`` of the top peak that +the sigma_A fit runs against. They are the last places in the alignment package +that answer "what is the mean intensity here" without going through +:class:`~torchref.scaling.WilsonNormaliser`. + +The argument for keeping them was cost: the shared fit is a Gamma GLM by IRLS and +the calc side needs one fit per candidate, K of them per rotation. This measures +that instead of asserting it, and also asks the two questions that decide whether +the swap is safe at all: + +* does the fit **converge** on a calculated set, which has near-zeros at the + nodes of the molecular transform where an observed set has none, and +* how far do the two normalisations actually differ, per reflection and in the + LLG ranking they feed. + +Warm-up first: on this filesystem a first call pays cold package reads inside +whatever timer surrounds it, which is worth ~100x and is not compute. + +Usage:: + + python alignment_lab/diagnostics/calc_norm_cost.py --pdb 1DAW --k 20 +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 + + +def _time(fn, repeats=3): + fn() # warm: discard the cold-read call + ts = [] + for _ in range(repeats): + t0 = time.perf_counter() + out = fn() + ts.append(time.perf_counter() - t0) + return min(ts), out + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) + ap.add_argument("--trial", type=int, default=0) + ap.add_argument("--k", type=int, default=20, help="candidate translations") + ap.add_argument("--n-coeff", type=int, default=6) + args = ap.parse_args() + + from torchref.experimental.alignment.translation import ( + DirectModelEvaluator, TranslationObs, amplitude_translation_search, + precompute_G_for_rotation, + ) + from torchref.scaling import WilsonNormaliser + + seed = seed_for(args.pdb, args.trial) + model, data = load_case(args.pdb) + R_true = random_rotation(seed) + rot = model.copy() + rot = rot.rotate(R_true.to(model.dtype_float), center=model.xyz().mean(0)) + rot.spacegroup = data.spacegroup.hm + p1 = rot.copy() + p1.spacegroup = "P 1" + + mask = data.get_valid_mask() + sig = getattr(data, "F_sigma", None) + obs = TranslationObs.build( + data.F[mask], data.hkl[mask], data.spacegroup, data.cell, + sig_F=None if sig is None else sig[mask], + ) + ev = DirectModelEvaluator(p1) + eye3 = torch.eye(3, dtype=torch.float64) + G, h_R = precompute_G_for_rotation( + ev, eye3, obs.hkl, data.spacegroup, data.cell) + _, _, peaks = amplitude_translation_search( + obs=obs, interpolator=ev, R_rotation=eye3, + spacegroup=data.spacegroup, real_cell=data.cell, + grid_steps=16, n_peaks=args.k, precomputed_G=G, precomputed_h_R=h_R) + + K = min(args.k, len(peaks)) + N = obs.hkl.shape[0] + t_cand = torch.as_tensor( + [p.translation for p in peaks[:K]], dtype=torch.float64, + device=G.device) + phase = torch.exp(2j * torch.pi * torch.einsum( + "ind,kd->kin", h_R.to(torch.float64), t_cand).to(G.dtype)) + F_calc = (G.view(1, *G.shape) * phase).sum(dim=1).abs().to(torch.float64) + + print(f"# {args.pdb} trial={args.trial} N={N} K={K} " + f"n_coeff={args.n_coeff}", flush=True) + + # --- current: one per-shell mean per candidate, all K at once --- + def per_shell(): + idx = obs.shell_idx.view(1, -1).expand(K, N) + cnt = torch.bincount(obs.shell_idx, minlength=obs.n_shells).to(torch.float64) + tot = torch.zeros((K, obs.n_shells), dtype=torch.float64, device=G.device) + tot.scatter_add_(1, idx, F_calc * F_calc) + mean = (tot / cnt.clamp(min=1.0).unsqueeze(0)).clamp(min=1e-30) + return F_calc / mean.sqrt().gather(1, idx) + + # --- proposed: the shared Wilson fit, once per candidate --- + def wilson(): + out = torch.empty_like(F_calc) + iters = [] + for k in range(K): + w = WilsonNormaliser( + F_calc[k] * F_calc[k], obs.s_mag, n_coeff=args.n_coeff, + s_lo=float(obs.s_mag.min()), s_hi=float(obs.s_mag.max()), + ) + out[k] = w.E.to(torch.float64) + iters.append(w.n_iter) + return out, iters + + t_shell, E_shell = _time(per_shell) + t_wilson, (E_wilson, iters) = _time(wilson) + + # How different are they, and does the *ranking* they feed move? + rel = ((E_wilson - E_shell).abs() + / E_shell.abs().clamp(min=1e-12)).median().item() + m_shell = (E_shell ** 2).mean(dim=1) + m_wilson = (E_wilson ** 2).mean(dim=1) + + print(f"ROW pdb={args.pdb} N={N} K={K} " + f"t_per_shell_ms={1000 * t_shell:.2f} " + f"t_wilson_ms={1000 * t_wilson:.1f} " + f"ratio={t_wilson / max(t_shell, 1e-9):.0f}x " + f"per_cand_ms={1000 * t_wilson / K:.1f} " + f"iter_min={min(iters)} iter_max={max(iters)} " + f"median_rel_dE={rel:.4f} " + f"meanE2_shell={m_shell.mean():.4f} " + f"meanE2_wilson={m_wilson.mean():.4f}", flush=True) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/unit/alignment/test_translation_obs.py b/tests/unit/alignment/test_translation_obs.py index d2d606ac..8d3d703a 100644 --- a/tests/unit/alignment/test_translation_obs.py +++ b/tests/unit/alignment/test_translation_obs.py @@ -78,16 +78,17 @@ def test_e_obs_is_invariant_to_the_amplitude_scale(): """Rescaling every amplitude must not change E. It is an ABSOLUTE normaliser. Exact in the model -- a common factor lands entirely in Sigma's intercept -- - but the fit is IRLS, so the tolerance is its convergence floor rather than - machine epsilon. Measured spread is ~2e-8 over 4000 reflections; 1e-6 catches - a genuine scale dependence without chasing the solver. + but the fit is IRLS in the configured float dtype and stops at a relative + tolerance, so the bar is that tolerance rather than machine epsilon. + Measured 1.2e-4 worst case over 4000 reflections at a 7.5x rescale; 1e-3 + catches a genuine scale dependence without chasing the solver. """ F, sig_F, hkl, sg, cell, _ = _case() base = TranslationObs.build(F, hkl, sg, cell, sig_F=sig_F) scaled = TranslationObs.build(7.5 * F, hkl, sg, cell, sig_F=7.5 * sig_F) - torch.testing.assert_close(base.E_obs, scaled.E_obs, rtol=1e-6, atol=1e-6) + torch.testing.assert_close(base.E_obs, scaled.E_obs, rtol=1e-3, atol=1e-3) # F/sigma is unchanged by a common factor, so the weight must be too. - torch.testing.assert_close(base.weight, scaled.weight, rtol=1e-6, atol=1e-6) + torch.testing.assert_close(base.weight, scaled.weight, rtol=1e-3, atol=1e-3) @pytest.mark.unit diff --git a/tests/unit/scaling/test_wilson_normaliser.py b/tests/unit/scaling/test_wilson_normaliser.py index 6fbf051c..0ff47375 100644 --- a/tests/unit/scaling/test_wilson_normaliser.py +++ b/tests/unit/scaling/test_wilson_normaliser.py @@ -33,6 +33,10 @@ def _wilson_data(n=20000, seed=0, centric_frac=0.1, eps_value=None): k = torch.where(centric, 0.5, 1.0).to(torch.float64) eps = (torch.ones(n, dtype=torch.float64) if eps_value is None else torch.full((n,), float(eps_value), dtype=torch.float64)) + # `_standard_gamma` takes no generator, so seed the global RNG too -- + # otherwise the draw depends on whatever ran before it and the test + # passes alone and fails in a suite. + torch.manual_seed(seed) I = torch._standard_gamma(k.clone()) / k * (eps * sigma) return I, s, eps, centric, sigma @@ -45,13 +49,15 @@ def _k_weighted_mean(v, centric): def test_unit_mean_is_an_identity_of_the_fit(): """The constant column's score equation IS `` = 1``. - ``sum_h k_h (I_h/mu_h - 1) = 0`` at the optimum, so this should hold to the - convergence tolerance rather than to some fitting accuracy. A loose result - here means the fit stopped early, not that the estimate is noisy. + ``sum_h k_h (I_h/mu_h - 1) = 0`` at the optimum, and the fit puts the + intercept on it in closed form, so this does NOT degrade as the convergence + tolerance is loosened. Measured 1e-9 to 2e-7 over five draws in float32; the + bar is the package's usual 1e-4 relative, so a failure here means the + intercept solve is broken rather than that the fit stopped early. """ I, s, eps, centric, _ = _wilson_data() w = WilsonNormaliser(I, s, eps=eps, centric=centric, n_coeff=6) - assert _k_weighted_mean(w.E_squared, centric) == pytest.approx(1.0, abs=1e-7) + assert _k_weighted_mean(w.E_squared, centric) == pytest.approx(1.0, rel=1e-4) @pytest.mark.parametrize("n_coeff", [1, 2, 6, 12]) @@ -76,7 +82,7 @@ def test_a_uniform_epsilon_cancels_out_of_the_ratio(): I, s, eps=torch.full_like(s, 2.0), centric=centric, n_coeff=6, ) # eps=2 halves the intensity going in AND halves Sigma, so E is unchanged. - assert torch.allclose(doubled.E, plain.E, rtol=1e-8, atol=1e-10) + assert torch.allclose(doubled.E, plain.E, rtol=1e-4, atol=1e-6) def test_invariant_to_the_units_the_data_arrive_in(): @@ -86,7 +92,7 @@ def test_invariant_to_the_units_the_data_arrive_in(): scaled = WilsonNormaliser( I * c, s, eps=eps, centric=centric, n_coeff=6, ).E - assert torch.allclose(scaled, base, rtol=1e-6, atol=1e-9), ( + assert torch.allclose(scaled, base, rtol=1e-4, atol=1e-6), ( f"scaling I by {c:g} moved E by " f"{float((scaled - base).abs().max()):.3e}" ) @@ -133,9 +139,16 @@ def test_the_range_only_matters_outside_the_fitted_data(): shared = WilsonNormaliser(I[sub], s[sub], s_lo=lo, s_hi=hi, **kw) own = WilsonNormaliser(I[sub], s[sub], **kw) + # Looser than the package's usual 1e-4, and the reason is the point of the + # test rather than an excuse. These are two INDEPENDENT fits, each stopped + # when its own objective stops improving by 1e-4 of what it has gained. The + # valley is flat along the high-order coefficients, so equal objectives + # there do not mean equal coefficients, and the curves separate by more than + # the objective did. Measured 0.2-1.6% over five draws; 3% catches a real + # dependence on the parameterisation without chasing the stopping rule. inside = torch.linspace(0.06, 0.29, 40, dtype=torch.float64) assert torch.allclose(shared.evaluate(inside), own.evaluate(inside), - rtol=1e-3), "the fitted function must not depend on " \ + rtol=3e-2), "the fitted function must not depend on " \ "how the basis was parameterised" # Outside its own data, the narrow fit is pinned at its endpoint; the one diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 5559a9da..d1efa838 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -56,6 +56,7 @@ fit_sigma_a_per_shell, llg_translation_rescore, local_translation_refine, + normalise_calc, precompute_G_for_rotation, ) @@ -710,13 +711,7 @@ def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): ).to(G_pre.dtype), ) Fc_top = (G_pre * phase_top).sum(dim=0).abs().to(torch.float64) - cnt_tf = torch.bincount( - obs.shell_idx, minlength=obs.n_shells, - ).to(torch.float64) - sum_Fc2 = torch.zeros(obs.n_shells, dtype=torch.float64, device=device) - sum_Fc2.scatter_add_(0, obs.shell_idx, Fc_top * Fc_top) - mean_Fc2 = (sum_Fc2 / cnt_tf.clamp(min=1.0)).clamp(min=1e-30) - E_calc_top = Fc_top / mean_Fc2.sqrt().index_select(0, obs.shell_idx) + E_calc_top = normalise_calc(Fc_top, obs) sigma_a_tf = fit_sigma_a_per_shell( obs.E_obs, E_calc_top, obs.centric, obs.shell_idx, obs.n_shells, n_grid=81, diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index c66345ef..ac19d386 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -424,6 +424,40 @@ def amplitude_translation_search( return corr_map_np, best, peaks +def normalise_calc(F_calc: torch.Tensor, obs: TranslationObs) -> torch.Tensor: + """``E_calc`` for one or many candidate translations, through the shared fit. + + Accepts ``(N,)`` or ``(K, N)`` and returns the same shape. Each candidate is + fitted separately, because the resolution envelope of ``|F_calc(h, t)|`` is a + property of that placement -- but by the *same* estimator the observed side + uses, on the same abscissa, so the two sides of the likelihood are normalised + by one rule rather than two. + + This used to be a per-shell mean, written out twice: once here for the K + candidates and once in the pipeline for the top peak that ``sigma_A`` is + fitted against. Two copies of one calculation is how they drift, and neither + was the estimator anything else in the package used. The difference is not + cosmetic -- the median per-reflection change is 2-4%. + + The fit converges in single-digit iterations, so the cost is a few + milliseconds per candidate against a placement of order a second, and it is + only paid when the likelihood rescore is on. + """ + single = F_calc.ndim == 1 + F = F_calc.reshape(1, -1) if single else F_calc + s_lo, s_hi = float(obs.s_mag.min()), float(obs.s_mag.max()) + out = torch.empty_like(F) + for k in range(F.shape[0]): + # No eps and nothing centric: a single molecular transform sampled at + # these indices carries no crystal multiplicity, and the observed side + # gets its own from `obs`. + out[k] = WilsonNormaliser( + F[k] * F[k], obs.s_mag, n_coeff=WILSON_N_COEFF, + s_lo=s_lo, s_hi=s_hi, + ).E.to(F.dtype) + return out[0] if single else out + + def fit_sigma_a_per_shell( E_obs: torch.Tensor, E_calc: torch.Tensor, @@ -489,7 +523,7 @@ def llg_translation_rescore( For each candidate t:: F_calc(h, t) = sum_i G_i(h) exp(2 pi i (h R_i).t) - E_calc(h, t) = |F_calc(h, t)| / sqrt(_shell) + E_calc(h, t) = |F_calc(h, t)| / sqrt(Sigma_calc(s; t)) LLG(t) = sum_h [LL(E_obs, D E_calc, var) - LL_Wilson(E_obs)] with ``var = (1 - D^2) + interp_var``. The Rice branch is used for acentric @@ -502,10 +536,11 @@ def llg_translation_rescore( recovery (28/30 against 27/30 the other way, one discordant cell), which is why ``use_llg_tf`` defaults off. - ``E_calc`` is normalised per shell **per candidate**, and that is not a - fourth answer to what the observations' normalisation is: it is a per-``t`` - scale, and forcing it to unit shell variance for every candidate is what - makes the K likelihoods comparable. What discriminates is the pattern across + ``E_calc`` is normalised **per candidate**, by :func:`normalise_calc` and so + by the same Wilson fit as the observed side. Per-candidate rather than once is + deliberate: the resolution envelope of ``|F_calc(h, t)|`` belongs to that + placement, and normalising every candidate to `` = 1`` is what makes + the K likelihoods comparable. What discriminates is the pattern across reflections, not the scale. Parameters @@ -534,7 +569,6 @@ def llg_translation_rescore( complex_dtype = G.dtype shell_idx = obs.shell_idx - n_shells = obs.n_shells centric = obs.centric K = t_candidates.shape[0] @@ -548,16 +582,8 @@ def llg_translation_rescore( Fc_complex = (G.view(1, S, N) * phase).sum(dim=1) # (K, N) F_calc = Fc_complex.abs().to(real_dtype) # (K, N) - # Per-shell E normalisation of F_calc across the K-batch. shell_idx_l = shell_idx.to(device).long() - cnt = torch.bincount(shell_idx_l, minlength=n_shells).to(real_dtype) - shell_idx_k = shell_idx_l.view(1, -1).expand(K, N) - F2 = F_calc * F_calc - sum_per_shell = torch.zeros((K, n_shells), dtype=real_dtype, device=device) - sum_per_shell.scatter_add_(1, shell_idx_k, F2) - mean_per_shell = (sum_per_shell / cnt.clamp(min=1.0).unsqueeze(0)).clamp(min=1e-30) - norm_per_refl = mean_per_shell.sqrt().gather(1, shell_idx_k) # (K, N) - E_calc = F_calc / norm_per_refl # (K, N) + E_calc = normalise_calc(F_calc, obs) # (K, N) E_obs = obs.E_obs.to(device).to(real_dtype) diff --git a/torchref/scaling/wilson.py b/torchref/scaling/wilson.py index 00b66fd1..cf91c5fe 100644 --- a/torchref/scaling/wilson.py +++ b/torchref/scaling/wilson.py @@ -32,6 +32,7 @@ import torch +from torchref.config import get_float_dtype from torchref.scaling.basis import chebyshev_design __all__ = ["WilsonNormaliser"] @@ -52,6 +53,26 @@ #: Step halvings allowed per IRLS iteration before the step is abandoned. MAX_HALVINGS = 30 +#: IRLS iterations allowed before the fit is declared failed. Generous, because +#: it should never bind: at :data:`DEFAULT_RTOL` the fit converges in single +#: digits. It is a runaway guard, not a budget. +DEFAULT_MAX_ITER = 100 + +#: Floor on the fitted mean, to keep ``y/mu`` finite if a step overshoots. +#: Must be representable in the working dtype -- ``1e-300`` is a float64 +#: constant and flushes to zero in float32, which turns the guard into the +#: division by zero it exists to prevent. +_MU_FLOOR = 1e-30 + +#: Relative convergence tolerance -- see :meth:`WilsonNormaliser._irls` for why +#: it is relative to the improvement so far rather than to the objective. +#: +#: This is a normalisation curve, not a refined parameter. The quantity it +#: decides is ``E = F / sqrt(Sigma)``, which is then compared against a model +#: that is wrong by tens of percent, so four digits is already far past what +#: anything downstream can use. +DEFAULT_RTOL = 1e-4 + class WilsonNormaliser: """``Sigma(s)`` by Gamma GLM, so that `` = 1`` by construction. @@ -140,8 +161,8 @@ def __init__( s_lo: Optional[float] = None, s_hi: Optional[float] = None, fit_mask: Optional[torch.Tensor] = None, - max_iter: int = 100, - tol: float = 1e-10, + max_iter: int = DEFAULT_MAX_ITER, + rtol: float = DEFAULT_RTOL, ) -> None: if I.ndim != 1: raise ValueError(f"I must be 1-D, got {tuple(I.shape)}") @@ -157,19 +178,24 @@ def __init__( self.s_lo = float(s_mag.min()) if s_lo is None else float(s_lo) self.s_hi = float(s_mag.max()) if s_hi is None else float(s_hi) - eps64 = ( - torch.ones_like(I, dtype=torch.float64) if eps is None - else eps.to(torch.float64).clamp(min=1.0) + # The configured float dtype, not float64. This is a six-coefficient + # fit of a smooth curve whose answer is compared against a model wrong + # by tens of percent; it does not need double, and hardcoding it here + # would be the only double-precision path in the scaling package. + work = get_float_dtype() + eps_w = ( + torch.ones_like(I, dtype=work) if eps is None + else eps.to(work).clamp(min=1.0) ) # Shape 1 acentric (exponential), 1/2 centric. Enters as the IRLS weight # because for a Gamma with shape k the variance is mu^2/k, so the # log-link working weight is k itself. k = ( - torch.ones_like(I, dtype=torch.float64) if centric is None - else torch.where(centric.to(torch.bool), 0.5, 1.0).to(torch.float64) + torch.ones_like(I, dtype=work) if centric is None + else torch.where(centric.to(torch.bool), 0.5, 1.0).to(work) ) - I_reduced = I.to(torch.float64) / eps64 + I_reduced = I.to(work) / eps_w # The Gamma likelihood has no support at or below zero. Absences and # negative measurements are held out of the fit and given a Sigma from # the curve like everything else -- excluding them from the *estimate* @@ -185,17 +211,17 @@ def __init__( self.n_fitted = int(usable.sum()) design = chebyshev_design( - (s_mag * 0.5).to(torch.float64), self.n_coeff, + (s_mag * 0.5).to(work), self.n_coeff, lo=self.s_lo * 0.5, hi=self.s_hi * 0.5, ) self.coefficients, self.n_iter = self._irls( - design[usable], I_reduced[usable], k[usable], max_iter, tol, + design[usable], I_reduced[usable], k[usable], max_iter, rtol, ) log_sigma = self._eval_log_sigma(design) self.sigma_wilson = torch.exp(log_sigma).to(self.dtype) self.mean_intensity = ( - torch.exp(log_sigma) * eps64 + torch.exp(log_sigma) * eps_w ).clamp(min=1e-30).to(self.dtype) self.E_squared = I / self.mean_intensity self.E = self.E_squared.clamp(min=0.0).sqrt() @@ -208,13 +234,47 @@ def _eval_log_sigma(self, design: torch.Tensor) -> torch.Tensor: min=-LOG_CLAMP + float(c[0]), max=LOG_CLAMP + float(c[0]), ) + @staticmethod + def _solve_intercept( + beta: torch.Tensor, X: torch.Tensor, y: torch.Tensor, w: torch.Tensor, + ) -> torch.Tensor: + """Put the intercept exactly on its score equation, closed form. + + The intercept's stationarity condition is ``sum_h k_h (I_h/mu_h - 1) = + 0``, which is `` = 1`` -- the identity this class exists to + provide. Shifting ``beta[0]`` by ``d`` scales every ``mu`` by ``e^d``, + so the ``d`` that satisfies it is available in one line: + + e^d = sum_h k_h (I_h/mu_h) / sum_h k_h + + Doing this explicitly decouples the identity from how tightly the SHAPE + converged. Without it `` = 1`` is only as good as the overall fit + tolerance -- at ``rtol = 1e-4`` it came out at 1 - 1e-5 -- and the + identity is not the kind of claim that should degrade with a stopping + rule. The remaining coefficients are untouched, so this changes the + curve's level and not its shape. + + In the working dtype like everything else here. The point is to make the + identity independent of the *stopping rule*, not to chase digits: it + lands within about 1e-6 of one, which is two orders inside anything that + reads it. + """ + eta = X @ beta + mu = torch.exp(eta.clamp(min=-LOG_CLAMP + float(beta[0]), + max=LOG_CLAMP + float(beta[0]))).clamp(min=_MU_FLOOR) + ratio = ((w * (y / mu)).sum() / w.sum()).clamp(min=_MU_FLOOR) + out = beta.clone() + out[0] = out[0] + torch.log(ratio) + return out + + def _irls( self, X: torch.Tensor, y: torch.Tensor, w: torch.Tensor, max_iter: int, - tol: float, + rtol: float, ) -> Tuple[torch.Tensor, int]: """Gamma GLM with a log link, by iteratively reweighted least squares. @@ -249,28 +309,39 @@ def _irls( """ # Seed at the constant curve, which is the exact MLE when Sigma has no # resolution dependence. Every later iteration only adds shape. - beta = torch.zeros(self.n_coeff, dtype=torch.float64, device=X.device) + beta = torch.zeros(self.n_coeff, dtype=X.dtype, device=X.device) beta[0] = torch.log(((w * y).sum() / w.sum()).clamp(min=1e-30)) def objective(b): eta = (X @ b).clamp( min=-LOG_CLAMP + float(b[0]), max=LOG_CLAMP + float(b[0]), ) - mu = torch.exp(eta).clamp(min=1e-300) + mu = torch.exp(eta).clamp(min=_MU_FLOOR) return float((w * (y / mu + eta)).sum()), eta, mu L, eta, mu = objective(beta) + L0 = L # the constant-curve seed, for the ratio below + + # Built and factorised ONCE. For a Gamma with a log link the IRLS + # working weight is the shape k, which does not depend on mu -- so + # `X^T W X` is the same matrix at every iteration and only the working + # response changes. Rebuilding it per iteration costs an O(N n^2) pass + # over every reflection for an answer that cannot have changed. + XtW = X.transpose(0, 1) * w.unsqueeze(0) + A = XtW @ X + # Ridge proportional to the matrix's own scale: the high-order + # Chebyshev columns go near-singular when the data cover only part + # of the basis range. + A = A + torch.eye(self.n_coeff, dtype=A.dtype, device=A.device) * ( + 1e-10 * float(torch.diagonal(A).abs().max().clamp(min=1e-30)) + ) + lu = torch.linalg.lu_factor(A) + for it in range(1, max_iter + 1): z = eta + (y - mu) / mu # working response - XtW = X.transpose(0, 1) * w.unsqueeze(0) - A = XtW @ X - # Ridge proportional to the matrix's own scale: the high-order - # Chebyshev columns go near-singular when the data cover only part - # of the basis range. - A = A + torch.eye(self.n_coeff, dtype=A.dtype, device=A.device) * ( - 1e-10 * float(torch.diagonal(A).abs().max().clamp(min=1e-30)) - ) - step = torch.linalg.solve(A, XtW @ z) - beta + step = torch.linalg.lu_solve( + *lu, (XtW @ z).unsqueeze(-1), + ).squeeze(-1) - beta if not torch.isfinite(step).all(): raise RuntimeError( f"Wilson fit diverged at iteration {it}: the IRLS solve " @@ -288,25 +359,34 @@ def objective(b): step = step * 0.5 if not accepted: # No downhill direction left: already at the optimum. - return beta, it - - # Per-reflection, NOT relative to |L|. Under I -> cI the optimum is - # just beta[0] -> beta[0] + log c, so the fit is exactly scale - # invariant -- but L picks up an additive `log c * sum(k)`, which - # makes a |dL|/|L| threshold mean something different at every - # scale. The difference itself is free of that term, so dividing by - # sum(k) leaves a criterion that is not. - improvement = abs(L - L_try) / float(w.sum()) + return self._solve_intercept(beta, X, y, w), it + + # Relative to the improvement achieved so far, not to |L|. + # + # |dL|/|L| is not usable here: under I -> cI the optimum is just + # beta[0] -> beta[0] + log c, so the fit is exactly scale invariant, + # but L picks up an additive `log c * sum(k)` and the ratio would + # mean something different at every scale. That additive term + # cancels in any DIFFERENCE, so a ratio of two differences is both + # relative and scale invariant -- which is what this is. + # + # The denominator is the total distance travelled from the constant + # seed, so the test reads "the last step moved us less than rtol of + # the way we have come". It is bounded below so a fit that starts at + # its own optimum (Sigma genuinely flat) terminates rather than + # dividing by zero. + step_gain = abs(L - L_try) + total_gain = max(abs(L0 - L_try), 1e-30) L, eta, mu = L_try, eta_try, mu_try - if improvement <= tol: - return beta, it + if step_gain <= rtol * total_gain: + return self._solve_intercept(beta, X, y, w), it raise RuntimeError( f"Wilson fit did not converge in {max_iter} IRLS iterations " - f"(objective still moving by {improvement:.2e} per reflection). " - f"Raising " - f"rather than falling back to a coarser estimate: a normaliser that " - f"silently becomes a different normaliser on hard cases is two " - f"normalisers wearing one name." + f"(last step still worth {step_gain / total_gain:.2e} of the total " + f"improvement, against rtol={rtol:.0e}). Raising rather than " + f"falling back to a coarser estimate: a normaliser that silently " + f"becomes a different normaliser on hard cases is two normalisers " + f"wearing one name." ) # -- evaluation elsewhere --------------------------------------------- @@ -321,7 +401,7 @@ def evaluate(self, s_mag: torch.Tensor) -> torch.Tensor: endpoint value, flat, rather than an extrapolation. """ design = chebyshev_design( - (s_mag * 0.5).to(torch.float64), self.n_coeff, + (s_mag * 0.5).to(self.coefficients.dtype), self.n_coeff, lo=self.s_lo * 0.5, hi=self.s_hi * 0.5, ) return torch.exp(self._eval_log_sigma(design)).to(self.dtype) @@ -350,13 +430,13 @@ def from_hkl( centricity -- which enters here as the Gamma shape, separately. The two branches feed two different parameters of the same likelihood. """ + work = get_float_dtype() hkl_l = hkl.to(torch.long) # The cell may carry the configured default device while the reflections # are somewhere else; the caller should not have to reconcile them. - rec = cell.reciprocal_basis_matrix.to(device=hkl_l.device, - dtype=torch.float64) - s_mag = (hkl_l.to(torch.float64) @ rec).norm(dim=-1).to(I.dtype) - eps = spacegroup.epsilon(hkl_l, friedel=False).to(torch.float64) + rec = cell.reciprocal_basis_matrix.to(device=hkl_l.device, dtype=work) + s_mag = (hkl_l.to(work) @ rec).norm(dim=-1).to(I.dtype) + eps = spacegroup.epsilon(hkl_l, friedel=False).to(work) centric = spacegroup.is_centric(hkl_l).to(torch.bool) # Systematically absent reflections are zero by symmetry, not by # measurement, so they carry no information about Sigma and would drag From 7a0140c4b13e7537abf27fed241d06d90b0a2fcf Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 31 Aug 2026 18:32:44 +0200 Subject: [PATCH 116/250] Note the Wilson fit changes in the changelog Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- docs/changelog.rst | 3 +++ 1 file changed, 3 insertions(+) diff --git a/docs/changelog.rst b/docs/changelog.rst index 6c670b7f..52a690b5 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,9 @@ Changelog Unreleased ---------- +- The Wilson normaliser converges in 8-12 IRLS iterations instead of 26-102. Its stopping rule was ``|dL|`` per reflection against 1e-10, which asks eleven significant digits of a normalisation curve; it is now relative to the improvement so far, which is scale-invariant for the same reason the absolute form was chosen +- `` = 1`` is solved in closed form for the intercept, so the identity no longer degrades as the convergence tolerance is loosened +- The Wilson fit runs in the configured float dtype rather than hardcoded double, and builds its normal-equations matrix once -- it is constant for a Gamma with a log link. Identical placements on all 30 benchmark cells, at 1.8x the speed; the unit suite went from 968s to 485s - Molecular replacement is now a rotation search feeding a translation search and nothing else; the pipeline returns a placement and stops. End-to-end pose recovery over 10 structures x 3 seeds went 18/30 to 30/30, at about a sixth of the wall clock - Removed the ML rescore from between the two searches. It reordered a shortlist that already contained the answer, and cost 6 of 30 placements - Removed the post-placement dense rotation re-sampling and rigid-body polish. They refined a correct placement away from truth on 2DQ6, 3GR5 and 4BX9; refining a placement is downstream refinement's job From 743554c20b5a0a1328594a2965772c8d00baa5ba Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 31 Aug 2026 18:54:00 +0200 Subject: [PATCH 117/250] Add a seed sweep for structures the panel samples too thinly Three trials cannot tell a solved structure from a coin flip. 2DQ6 passes 6/10 over seeds and does it bimodally -- 0.00-3.86 deg or 21.84-30.87, never near the gate -- because the tNCS alternative solution wins about 40% of the time and which one wins turns on the fifth decimal of Sigma(s). 6G9X, 1DAW and 3K7M are 10/10. So a one-cell panel difference on 2DQ6 says nothing about whatever change preceded it. This is the harness that establishes that before anyone spends six runs bisecting for a cause, as happened here. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/marginal_seeds.sh | 30 ++++++++++++++++++++++++ 1 file changed, 30 insertions(+) create mode 100644 alignment_lab/analysis/marginal_seeds.sh diff --git a/alignment_lab/analysis/marginal_seeds.sh b/alignment_lab/analysis/marginal_seeds.sh new file mode 100644 index 00000000..09c670d2 --- /dev/null +++ b/alignment_lab/analysis/marginal_seeds.sh @@ -0,0 +1,30 @@ +#!/bin/bash +# Is 2DQ6 a cell the pipeline solves, or a coin flip the panel samples once? +# +# 2DQ6 t1 moved 6.5 -> 28.7 deg on a 1.3e-5 change in Sigma(s), taking the panel +# from 30/30 to 29/30. If the structure is genuinely marginal then its pass rate +# over SEEDS is intermediate, and a single flip carries no information about the +# change that produced it -- a different seed would have flipped it anyway. +# Three trials per structure cannot tell the difference; ten can. +#SBATCH --job-name=marg +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=day +#SBATCH --time=04:00:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-3 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +# Two known-marginal structures and two the panel has never lost, as controls. +PDBS=(2DQ6 6G9X 1DAW 3K7M) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 +export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +for T in $(seq 0 9); do + "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial "$T" \ + --arms analytic_r 2>/dev/null | grep '^ROW ' +done From 7fe29e987a9c8c3a4af6d496b089edc43a0a64b4 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 31 Aug 2026 19:11:25 +0200 Subject: [PATCH 118/250] Add the Wilson-moment and 2DQ6 discrimination diagnostics 2DQ6 has been excused as a tNCS case since May. It is not one: the deposited ASU is a single 865-residue chain at Vm 3.74 and 67% solvent, so there is no second copy for a translation to relate, and <(E^2-1)^2> over acentrics through the shared normaliser is 1.151 -- mid-pack, against 1.831 for 1DAW, which solves every seed. The 5.528 on record is not this quantity. wilson_moments.py computes that statistic for the whole benchmark, so the claim is checkable rather than inherited. A second moment is a property of the normalisation as much as of the data, and the original number predates there being one normalisation. dq6_discrimination.sh says what is actually wrong. Over ten seeds the rotation search puts truth in the fifteen candidates every time, and the translation correlation ranks it FIRST every time, at 2.1-3.3 sigma. The pipeline then picks its winner by analytic R, which ranks truth 3rd to 5th, and solves 6/10. The structure is not hard; the selection rule discards a score that is already right. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/dq6_discrimination.sh | 26 +++++++++ alignment_lab/analysis/wilson_moments.sh | 17 ++++++ alignment_lab/diagnostics/wilson_moments.py | 59 ++++++++++++++++++++ 3 files changed, 102 insertions(+) create mode 100644 alignment_lab/analysis/dq6_discrimination.sh create mode 100644 alignment_lab/analysis/wilson_moments.sh create mode 100644 alignment_lab/diagnostics/wilson_moments.py diff --git a/alignment_lab/analysis/dq6_discrimination.sh b/alignment_lab/analysis/dq6_discrimination.sh new file mode 100644 index 00000000..29a8c2be --- /dev/null +++ b/alignment_lab/analysis/dq6_discrimination.sh @@ -0,0 +1,26 @@ +#!/bin/bash +# On the seeds where 2DQ6 fails, is truth in the candidate list at all? +# +# The structure passes 6/10 end to end and bimodally -- 0-4 deg or 21-31. That +# is either the rotation function never producing the true orientation, or the +# translation function producing it and ranking the tNCS alternative above it. +# Those are different problems and the panel cannot tell them apart, because it +# only reports the winner. This reports truth's RANK under each score. +#SBATCH --job-name=dq6disc +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=day +#SBATCH --time=04:00:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +# 15 candidates: what the pipeline actually carries into the translation search. +"$PY" -u alignment_lab/diagnostics/frf_vs_ftf_discrimination.py \ + --pdb 2DQ6 --trials 10 --n-cand 15 2>/dev/null | grep -E '^(ROW|CAND)' +echo DONE diff --git a/alignment_lab/analysis/wilson_moments.sh b/alignment_lab/analysis/wilson_moments.sh new file mode 100644 index 00000000..5f83e7c9 --- /dev/null +++ b/alignment_lab/analysis/wilson_moments.sh @@ -0,0 +1,17 @@ +#!/bin/bash +#SBATCH --job-name=moments +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 +export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +"$PY" -u alignment_lab/diagnostics/wilson_moments.py 2>/dev/null +echo DONE diff --git a/alignment_lab/diagnostics/wilson_moments.py b/alignment_lab/diagnostics/wilson_moments.py new file mode 100644 index 00000000..b68bf1fb --- /dev/null +++ b/alignment_lab/diagnostics/wilson_moments.py @@ -0,0 +1,59 @@ +"""Second moment of E^2 per structure, through the shared Wilson fit. + +``<(E^2-1)^2>`` is the standard indicator for translational NCS and related +intensity modulations: 1.0 for ideal acentric Wilson data, larger when whole +classes of reflections reinforce or cancel together. 2DQ6 was recorded at 5.528 +against 1.0-1.2 for every other benchmark structure, and that number is the sole +evidence for calling it a tNCS case. + +It is worth recomputing, because tNCS needs at least two copies related by a +pure translation and 2DQ6 deposits ONE chain in the asymmetric unit, and because +the number was measured when the package had five disagreeing answers to what E +means. A second moment is a property of the normalisation as much as of the +data: normalise by a curve that is too flat and the resolution trend leaks +straight into the moment. +""" +import sys +from pathlib import Path +import torch +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) +from lab import BENCH_PDBS, load_case # noqa: E402 +from torchref.scaling import WilsonNormaliser # noqa: E402 + +print(f"{'pdb':6s} {'sg':12s} {'N':>7s} {'<(E2-1)^2>':>11s} {'':>7s} " + f"{'<|E|>':>7s} {'shell-norm':>11s}") +for pdb in BENCH_PDBS: + model, data = load_case(pdb) + mask = data.get_valid_mask() + F = data.F[mask].abs().to(torch.float64) + hkl = data.hkl[mask] + rec = data.cell.reciprocal_basis_matrix.to(torch.float64).to(hkl.device) + s = (hkl.to(torch.float64) @ rec).norm(dim=-1) + hkl_l = hkl.round().to(torch.int64) + eps = data.spacegroup.epsilon(hkl_l, friedel=False).to(torch.float64).clamp(min=1.0) + cen = data.spacegroup.is_centric(hkl_l).to(torch.bool) + acen = ~cen + + w = WilsonNormaliser(F * F, s, eps=eps, centric=cen, n_coeff=6) + E2 = w.E_squared.to(torch.float64)[acen] + m2 = float(((E2 - 1.0) ** 2).mean()) + + # The same moment under a 20-shell mean, which is what the older estimate + # would have used -- to separate "the data are odd" from "the curve was". + order = torch.argsort(s) + sh = torch.zeros_like(s, dtype=torch.long) + chunk = s.numel() // 20 + for k in range(20): + a = k * chunk + b = (k + 1) * chunk if k < 19 else s.numel() + sh[order[a:b]] = k + I = F * F / eps + tot = torch.zeros(20, dtype=torch.float64).scatter_add_(0, sh, I) + cnt = torch.bincount(sh, minlength=20).to(torch.float64).clamp(min=1) + E2s = I / (tot / cnt).clamp(min=1e-30).index_select(0, sh) + m2s = float(((E2s[acen] - 1.0) ** 2).mean()) + + print(f"{pdb:6s} {str(data.spacegroup.hm):12s} {int(acen.sum()):7d} " + f"{m2:11.3f} {float(E2.mean()):7.3f} " + f"{float(E2.clamp(min=0).sqrt().mean()):7.3f} {m2s:11.3f}", flush=True) From b70990b30a9bc13b3a6935ebfa37c806b4cab3bb Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 31 Aug 2026 20:40:53 +0200 Subject: [PATCH 119/250] Place every rotation candidate and take the best The pipeline walked candidates in the rotation function's order and returned the first placement that beat R < 0.45, after a minimum of three. So its answer depended on that order, and it could accept the third candidate without ever scoring the tenth -- which on a structure where several orientations place plausibly is not a choice between them. It also made the selection rule impossible to compare against a harness that ranks the whole list: 2DQ6 solved 6/10 end to end while truth was top-ranked by analytic R in 0/10, and two numbers that far apart cannot describe the same rule. Now it places all of them and takes the minimum. min_tries, max_tries and rfactor_converged are gone; n_rotation_candidates goes 15 -> 25. **Measured neutral.** Identical outcomes on 10 structures x 3 seeds, both ranking arms, zero success flips; and on a 10-seed sweep of 2DQ6, 6G9X, 1DAW and 3K7M the residuals agree to two decimals. Candidates 16-25 never win, and the early stop was never truncating before the best. About 1.6x the wall clock. That is worth stating as a negative result, because it rules out the obvious explanation for 2DQ6: it does not fail because the pipeline stopped looking. It fails having scored every candidate, which puts the problem in the score. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/llg_shared_sigma_a.sh | 21 +++ alignment_lab/analysis/marginal_seeds.sh | 2 +- alignment_lab/analysis/pose_arms.sh | 12 +- .../diagnostics/llg_shared_sigma_a.py | 165 ++++++++++++++++++ alignment_lab/diagnostics/pose_recovery.py | 2 +- docs/changelog.rst | 1 + torchref/experimental/alignment/pipeline.py | 56 ++---- 7 files changed, 215 insertions(+), 44 deletions(-) create mode 100644 alignment_lab/analysis/llg_shared_sigma_a.sh create mode 100644 alignment_lab/diagnostics/llg_shared_sigma_a.py diff --git a/alignment_lab/analysis/llg_shared_sigma_a.sh b/alignment_lab/analysis/llg_shared_sigma_a.sh new file mode 100644 index 00000000..5d609646 --- /dev/null +++ b/alignment_lab/analysis/llg_shared_sigma_a.sh @@ -0,0 +1,21 @@ +#!/bin/bash +#SBATCH --job-name=llgsa +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=day +#SBATCH --time=04:00:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-3 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +PDBS=(2DQ6 6G9X 1DAW 3K7M) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +"$PY" -u alignment_lab/diagnostics/llg_shared_sigma_a.py --pdb "$PDB" --trials 10 \ + 2>&1 | grep -E '^ROW|rror' +echo DONE diff --git a/alignment_lab/analysis/marginal_seeds.sh b/alignment_lab/analysis/marginal_seeds.sh index 09c670d2..fa07f6bf 100644 --- a/alignment_lab/analysis/marginal_seeds.sh +++ b/alignment_lab/analysis/marginal_seeds.sh @@ -26,5 +26,5 @@ export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" for T in $(seq 0 9); do "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial "$T" \ - --arms analytic_r 2>/dev/null | grep '^ROW ' + --arms analytic_r --n-rotation-candidates 25 2>/dev/null | grep '^ROW ' done diff --git a/alignment_lab/analysis/pose_arms.sh b/alignment_lab/analysis/pose_arms.sh index 5a377903..9ffbf6f7 100644 --- a/alignment_lab/analysis/pose_arms.sh +++ b/alignment_lab/analysis/pose_arms.sh @@ -4,8 +4,14 @@ # # 10 structures x 3 trials x 2 arms. The arms are the open ranking question: # analytic R (the default) against the translation LLG, which wins 27/30 to -# 22/30 at rank level. The number to hold is 24/30 successes, what the pipeline -# scored once the ML rescore was removed from between the two stages. +# 22/30 at rank level. +# +# 25 candidates and no early stopping, matching the pipeline. Under the old rule +# it walked the list until a placement beat R < 0.45 and returned that, so the +# answer depended on FRF order and could not be compared against a harness that +# ranks the whole list -- 2DQ6 solved 6/10 end to end while truth was top-ranked +# by analytic R in 0/10, which is only possible if the two measure different +# things. #SBATCH --job-name=posearm #SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out #SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err @@ -25,5 +31,5 @@ export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREA export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" for T in 0 1 2; do "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial "$T" \ - --arms analytic_r,llg_tf --n-rotation-candidates 15 2>/dev/null | grep '^ROW ' + --arms analytic_r,llg_tf --n-rotation-candidates 25 2>/dev/null | grep '^ROW ' done diff --git a/alignment_lab/diagnostics/llg_shared_sigma_a.py b/alignment_lab/diagnostics/llg_shared_sigma_a.py new file mode 100644 index 00000000..6bfe9d90 --- /dev/null +++ b/alignment_lab/diagnostics/llg_shared_sigma_a.py @@ -0,0 +1,165 @@ +"""Does the translation LLG rank badly because each candidate fits its own sigma_A? + +Over ten seeds on 2DQ6 the plain correlation puts truth at rank 0 in 10/10 while +the LLG -- a likelihood, strictly more information -- manages 2/10. A weaker +score beating a stronger one is a symptom, not a result. + +The suspect is how many parameters each score fits PER CANDIDATE: + + correlation 0 sum w E_obs^2_c |Fc|^2 / sum w |Fc|^2 + analytic R 1 the global scale k + LLG n_shells fit_sigma_a_per_shell, on THAT candidate's + own top translation + +`_llg_tf_rescore` is called once per rotation candidate and refits sigma_A each +time, so every wrong orientation is scored against a likelihood tuned to itself. +The docstring there warns against exactly this one level down -- refitting per +translation -- and the pipeline then does it per rotation. + +This recomputes the LLG with sigma_A held FIXED across candidates, three ways: + +``per_cand`` what the pipeline does now, as the control. +``shared`` fitted once, on the FRF's top-ranked candidate. Model-dependent + but not candidate-dependent, so it cannot flatter any one of them. +``empirical`` from Sigma_obs/Sigma_calc via weighting.empirical_sigma_a, which + is rotation-invariant by construction -- total scattering per + shell does not depend on orientation -- and is the estimate that + exists for precisely this reason. +""" +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +import numpy as np +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) + +from lab import (BENCH_PDBS, rotated_case, seed_for, # noqa: E402 + symmetry_orbit) +from lab.truth import angle_to_orbit # noqa: E402 + + +def _rank_of_truth(scores, is_truth, higher_is_better=True): + order = sorted(range(len(scores)), key=lambda i: scores[i], + reverse=higher_is_better) + for pos, i in enumerate(order): + if is_truth[i]: + return pos + return -1 + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="2DQ6", choices=list(BENCH_PDBS)) + ap.add_argument("--trials", type=int, default=10) + ap.add_argument("--n-cand", type=int, default=15) + ap.add_argument("--thr-deg", type=float, default=8.0) + args = ap.parse_args() + + from torchref.experimental.alignment.frf.rotation_utils import ( + rotation_matrix_from_edmonds_euler, + ) + from torchref.experimental.alignment.pipeline import ( + MolecularReplacementPipeline, + ) + from torchref.experimental.alignment.rotation_search import prepare_frf_inputs + from torchref.experimental.alignment.translation import ( + DirectModelEvaluator, amplitude_translation_search, fit_sigma_a_per_shell, + llg_translation_rescore, normalise_calc, precompute_G_for_rotation, + ) + from torchref.scaling import WilsonNormaliser + from torchref.scaling.weighting import empirical_sigma_a + + for trial in range(args.trials): + seed = seed_for(args.pdb, trial) + model, data, R_true = rotated_case(args.pdb, seed) + pipe = MolecularReplacementPipeline( + data, model, verbose=0, n_rotation_peaks=200, + n_rotation_candidates=args.n_cand, use_llg_tf=False) + frf = prepare_frf_inputs(model, data, d_min=pipe.d_min, d_max=pipe.d_max, + n_shells=pipe.n_shells, verbose=0) + pipe._frf = frf + peaks = pipe._rotation_candidates(frf)[: args.n_cand] + pipe._prepare_translation_arrays() + obs = pipe._obs + eye3 = pipe._eye3 + orbit = symmetry_orbit( + R_true, data.spacegroup.matrices.to(torch.float64).cpu(), + side="left", frame="cart", + reciprocal_basis=data.cell.reciprocal_basis_matrix.to( + torch.float64).cpu()) + + # One pass to collect each candidate's G, its top translation and its + # own E_calc there; sigma_A choices are applied afterwards so every + # variant scores the SAME placements. + cand = [] + for p in peaks: + ang = angle_to_orbit( + rotation_matrix_from_edmonds_euler(p.alpha, p.beta, p.gamma), + orbit) + rot = pipe._make_rotated(p)[0] + rot.spacegroup = data.spacegroup.hm + p1 = rot.copy(); p1.spacegroup = "P 1" + ev = DirectModelEvaluator(p1) + G, h_R = precompute_G_for_rotation( + ev, eye3, obs.hkl, data.spacegroup, data.cell) + _, _, tp = amplitude_translation_search( + obs=obs, interpolator=ev, R_rotation=eye3, + spacegroup=data.spacegroup, real_cell=data.cell, + grid_steps=pipe.translation_grid_steps, + n_peaks=pipe.n_translation_peaks, cluster_radius=0.05, + precomputed_G=G, precomputed_h_R=h_R) + t_top = torch.as_tensor(tp[0].translation, dtype=torch.float64, + device=G.device) + ph = torch.exp(2j * torch.pi * torch.einsum( + "ind,d->in", h_R.to(torch.float64), t_top).to(G.dtype)) + Fc_top = (G * ph).sum(dim=0).abs().to(torch.float64) + cand.append(dict(ang=float(ang), corr=float(tp[0].score), G=G, + h_R=h_R, t=t_top, Fc=Fc_top, + E_calc=normalise_calc(Fc_top, obs))) + + is_truth = [c["ang"] <= args.thr_deg for c in cand] + if not any(is_truth): + print(f"ROW pdb={args.pdb} trial={trial} truth_found=0", flush=True) + continue + + def sa_per_cand(c): + return fit_sigma_a_per_shell(obs.E_obs, c["E_calc"], obs.centric, + obs.shell_idx, obs.n_shells, n_grid=81) + sa_shared = sa_per_cand(cand[0]) # the FRF's own top candidate + fit_calc = WilsonNormaliser( + cand[0]["Fc"] ** 2, obs.s_mag, n_coeff=6, + s_lo=float(obs.s_mag.min()), s_hi=float(obs.s_mag.max())) + sa_emp_per_refl = empirical_sigma_a( + obs.fit.evaluate(obs.s_mag).to(torch.float64), + fit_calc.evaluate(obs.s_mag).to(torch.float64)) + # collapse to per-shell, the shape llg_translation_rescore expects + cnt = torch.bincount(obs.shell_idx, minlength=obs.n_shells).to(torch.float64) + tot = torch.zeros(obs.n_shells, dtype=torch.float64).scatter_add_( + 0, obs.shell_idx, sa_emp_per_refl.to(torch.float64)) + sa_emp = (tot / cnt.clamp(min=1.0)).clamp(1e-3, 1 - 1e-6) + + def llg_of(c, sigma_a): + return float(llg_translation_rescore( + obs=obs, G=c["G"], h_R=c["h_R"], + t_candidates=c["t"].view(1, 3), sigma_a=sigma_a)[0]) + + variants = { + "per_cand": [llg_of(c, sa_per_cand(c)) for c in cand], + "shared": [llg_of(c, sa_shared) for c in cand], + "empirical": [llg_of(c, sa_emp) for c in cand], + } + r_corr = _rank_of_truth([c["corr"] for c in cand], is_truth) + parts = " ".join( + f"rank_{k}={_rank_of_truth(v, is_truth)}" for k, v in variants.items()) + print(f"ROW pdb={args.pdb} trial={trial} n_truth={sum(is_truth)} " + f"rank_corr={r_corr} {parts}", flush=True) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py index f55e5278..d8f57359 100644 --- a/alignment_lab/diagnostics/pose_recovery.py +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -85,7 +85,7 @@ def main() -> int: ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) ap.add_argument("--trial", type=int, default=0) ap.add_argument("--arms", default="analytic_r,llg_tf") - ap.add_argument("--n-rotation-candidates", type=int, default=15) + ap.add_argument("--n-rotation-candidates", type=int, default=25) ap.add_argument("--n-rotation-peaks", type=int, default=200) ap.add_argument("--success-deg", type=float, default=8.0) ap.add_argument("--verbose", type=int, default=0) diff --git a/docs/changelog.rst b/docs/changelog.rst index 52a690b5..132fa11c 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- The molecular-replacement pipeline places all ``n_rotation_candidates`` (now 25, was 15) and returns the best, instead of stopping once a placement beat an R-factor threshold. The old rule made the answer depend on the order the rotation function happened to produce and could accept the third candidate without scoring the tenth. Measured neutral -- identical placements on 10 structures x 3 seeds and on a 10-seed sweep of the marginal cases -- at about 1.6x the wall clock - The Wilson normaliser converges in 8-12 IRLS iterations instead of 26-102. Its stopping rule was ``|dL|`` per reflection against 1e-10, which asks eleven significant digits of a normalisation curve; it is now relative to the improvement so far, which is scale-invariant for the same reason the absolute form was chosen - `` = 1`` is solved in closed form for the intercept, so the identity no longer degrades as the convergence tolerance is loosened - The Wilson fit runs in the configured float dtype rather than hardcoded double, and builds its normal-equations matrix once -- it is constant for a Gamma with a log link. Identical placements on all 30 benchmark cells, at 1.8x the speed; the unit suite went from 968s to 485s diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index d1efa838..c6331c5a 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -19,11 +19,18 @@ truth, and end-to-end pose recovery was 18/30 with it against 24/30 without (McNemar p = 0.031, 6-0 discordant). -Each orientation is placed independently and the candidates are ranked by the -translation search's analytical R, with early stopping once one beats -``rfactor_converged``. The user-facing solvent-aware R-work is computed once, on -the winner. The pipeline returns a *placement* -- refining it is the caller's -job, and downstream refinement does it better than a bolted-on polish did. +Every candidate is placed and then the best is taken, with no early stopping. +Stopping early made the pipeline's answer depend on the order the rotation +function happened to produce -- it walked the list until one placement beat an +R-factor threshold and returned that, so it could accept the third candidate +without ever scoring the tenth. On a structure where several orientations place +plausibly that is not a choice between them, and it made the selection rule +impossible to reason about or to measure against a ranking harness. + +The candidates are ranked by the translation search's analytical R. The +user-facing solvent-aware R-work is computed once, on the winner. The pipeline +returns a *placement* -- refining it is the caller's job, and downstream +refinement does it better than a bolted-on polish did. ``align_model_to_data`` delegates here; this class is the implementation of record. The crystallographic stages live in @@ -255,9 +262,8 @@ class MRSolution: translation_score : float Analytical-R of the best translation for this candidate (lower better). r_factor : float - Ranking key. During the candidate loop this is the rigid-body's own - (no-solvent) R-work; for the returned winner it is replaced by the - solvent-aware Scaler R-work. + Ranking key: the translation search's analytical-scale R. For the + returned winner it is replaced by the solvent-aware Scaler R-work. model : ModelFT The rotated (+translated +refined) model for this candidate. """ @@ -313,7 +319,7 @@ def __init__( n_rotation_peaks: int = 500, model_error_A: Optional[float] = None, # --- candidate tree --- - n_rotation_candidates: int = 15, + n_rotation_candidates: int = 25, n_translation_peaks: int = 20, n_translation_candidates: int = 3, translation_grid_steps: int = 16, @@ -323,10 +329,6 @@ def __init__( # this stage has always done -- see `_prepare_translation_arrays`. tf_d_min: Optional[float] = None, tf_d_max: Optional[float] = None, - # --- early stop --- - min_tries: int = 3, - max_tries: Optional[int] = None, - rfactor_converged: float = 0.45, ): self.data = data self.model = model @@ -356,10 +358,6 @@ def __init__( self.tf_d_min = tf_d_min self.tf_d_max = tf_d_max - self.min_tries = min_tries - self.max_tries = max_tries - self.rfactor_converged = rfactor_converged - self._timer = _StageTimer(enabled=verbose >= 2) # Filled in by run(). self._frf = None @@ -431,17 +429,10 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: # --- Stage 2: per-candidate translation search --- self._prepare_translation_arrays() n_rot = min(self.n_rotation_candidates, len(candidates)) - max_tries = self.max_tries if self.max_tries is not None else n_rot if self.verbose > 0 and n_rot > 1: - print( - f"mr: trying up to {n_rot} rotation candidates " - f"(early-stop after ≥{self.min_tries} once R < " - f"{self.rfactor_converged}).", - flush=True, - ) + print(f"mr: placing all {n_rot} rotation candidates…", flush=True) solutions: List[MRSolution] = [] - best_r = float("inf") for k in range(n_rot): peak_k = candidates[k] rotated_k, R_rec_k = self._make_rotated(peak_k) @@ -475,19 +466,6 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: model=placed, ) ) - best_r = min(best_r, r_rank) - - n_done = k + 1 - if n_done >= self.min_tries and best_r < self.rfactor_converged: - if self.verbose > 0: - print( - f"mr: converged (R {best_r:.4f} < " - f"{self.rfactor_converged}) after {n_done} candidates.", - flush=True, - ) - break - if n_done >= max_tries: - break if not solutions: raise RuntimeError("Translation + joint refine produced no candidates.") @@ -757,7 +735,7 @@ def align_model_to_data( n_translation_peaks: int = 20, n_translation_candidates: int = 3, translation_grid_steps: int = 16, - n_rotation_candidates: int = 15, + n_rotation_candidates: int = 25, use_llg_tf: bool = False, tf_d_min: Optional[float] = None, tf_d_max: Optional[float] = None, From 67e9c46f0af308eee4e815d6c16e8412940f8a81 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 01:40:49 +0200 Subject: [PATCH 120/250] Make the pipeline's verbosity levels mean something Seventeen `if self.verbose > 0: print(...)` sites, with no statement anywhere of what a level was for. The same stage reported at 1 in one place and 2 in another, and nothing enforced that a level meant the same thing twice. The levels are now documented on the class and go through one emitter: 1 is what happened, 2 is why it chose what it chose, 3 is per-translation-peak detail. No logging module -- the package has 129 `verbose: int` parameters and zero uses of `logging`, and adding one here would be a second mechanism for a job that already has one. Level 2 emits a machine-readable CAND line per rotation candidate with the rotation score, the translation-function score at its chosen peak, the analytic R that ranks it, and the translation. `_placement_for_candidate` returns the TF score for that reason; it selects nothing, but which of the two scores a wrong placement disagreed on is the first question anyone asks and it cannot be recovered from the winner afterwards. pose_recovery drives the pipeline directly instead of `align_model_to_data` and joins its ranked solutions against the known orientation, so the harness reports which candidate won and how far each was from truth without rebuilding the placement loop. Rebuilding it is what produced two irreconcilable numbers earlier -- truth top-ranked by analytic R in 0 of 10 seeds against 6 of 10 solved -- because the copy fed the R-factor a different set of translation peaks. First use, on 2DQ6: all 25 candidates fall within 0.023 of each other in R, 14-15 of them are correct, and the wrong winner beats a correct candidate by 0.0001. The score does not discriminate on that structure. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/diagnostics/pose_recovery.py | 53 ++++++- docs/changelog.rst | 1 + torchref/experimental/alignment/pipeline.py | 145 +++++++++++++------- 3 files changed, 142 insertions(+), 57 deletions(-) diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py index d8f57359..e2529935 100644 --- a/alignment_lab/diagnostics/pose_recovery.py +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -80,6 +80,45 @@ def residual_rotation_deg(aligned_xyz, canonical_xyz, symops) -> float: "torchref.experimental.alignment.rotation_search").LMAX_CAP + +def _report_candidates(solutions, R_true, symops, success_deg) -> None: + """Annotate the pipeline's own candidates with how far each is from truth. + + The pipeline reports every candidate's scores at ``verbose >= 2`` but cannot + say which was right -- it has no ground truth, and a version of it that did + would be measuring itself. This joins the two: the ranked solutions it + returned, against the orientation the benchmark rotated the model by. + + That join is the whole point of driving the pipeline directly rather than + rebuilding its placement loop in the harness. A reimplementation drifts, and + then the two disagree about which candidate the pipeline picked -- which is + exactly what happened here: a harness reported truth top-ranked by analytic + R in 0 of 10 seeds while the pipeline solved 6 of them, because it fed the + R-factor a different set of translation peaks. + + ``SOLN`` lines are ordered as the pipeline ranked them, so line 0 is what it + returned. ``dtruth`` is the angle from that candidate's orientation to the + true one modulo crystal symmetry; ``pick`` marks the winner and ``true`` + marks every candidate that was in fact correct. + """ + from torchref.experimental.alignment.frf.rotation_utils import ( + rotation_angular_distance_deg, + ) + + R_t = R_true.to(torch.float64).cpu() + print(" SOLN rank rot_score tf_R dtruth flags") + for i, sol in enumerate(solutions): + R = torch.as_tensor(sol.rotation, dtype=torch.float64) + # `rotation` maps the search-model frame onto the crystal frame; the + # benchmark's R_true is the rotation applied to the coordinates, so the + # recovered orientation is compared as its transpose. + d = min(float(rotation_angular_distance_deg(R.T @ R_t, symops[k])) + for k in range(symops.shape[0])) + flags = ("pick " if i == 0 else " ") + ("true" if d <= success_deg else "") + print(f" SOLN {i:4d} {sol.rotation_score:10.3f} " + f"{sol.translation_score:7.4f} {d:8.2f} {flags}") + + def main() -> int: ap = argparse.ArgumentParser(description=__doc__) ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) @@ -92,7 +131,7 @@ def main() -> int: ap.add_argument("--out-csv", default=None) args = ap.parse_args() - from torchref.experimental.alignment import align_model_to_data + from torchref.experimental.alignment import MolecularReplacementPipeline seed = seed_for(args.pdb, args.trial) model, data = load_case(args.pdb) @@ -123,15 +162,21 @@ def main() -> int: center=canonical_xyz.mean(0)) t0 = time.time() try: - aligned = align_model_to_data( - search, data, d_min=4.0, d_max=15.0, n_shells=20, + # The pipeline rather than `align_model_to_data`, which returns only + # the winner. Every candidate's score is the diagnosis when a + # placement goes wrong, and the pipeline already computed them. + pipe = MolecularReplacementPipeline( + data, search, d_min=4.0, d_max=15.0, n_shells=20, n_rotation_peaks=args.n_rotation_peaks, - do_translation=True, n_rotation_candidates=args.n_rotation_candidates, verbose=args.verbose, **flags, ) + solutions = pipe.run(do_translation=True) + aligned = solutions[0].model resid = residual_rotation_deg(aligned.xyz(), canonical_xyz, symops) err = "" + if args.verbose >= 2: + _report_candidates(solutions, R_true, symops, args.success_deg) except Exception as exc: # a crashed arm must not read as a success resid, err = float("nan"), f"{type(exc).__name__}: {exc}" secs = time.time() - t0 diff --git a/docs/changelog.rst b/docs/changelog.rst index 132fa11c..d7c18086 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- The molecular-replacement pipeline's ``verbose`` levels are a documented contract routed through one emitter, rather than ``if verbose > 0: print(...)`` at seventeen sites. Level 2 emits one machine-readable ``CAND`` line per rotation candidate carrying every score the selection could have used, so a wrong placement can be diagnosed from the run itself instead of from a harness that re-implements the placement loop and then disagrees with it - The molecular-replacement pipeline places all ``n_rotation_candidates`` (now 25, was 15) and returns the best, instead of stopping once a placement beat an R-factor threshold. The old rule made the answer depend on the order the rotation function happened to produce and could accept the third candidate without scoring the tenth. Measured neutral -- identical placements on 10 structures x 3 seeds and on a 10-seed sweep of the marginal cases -- at about 1.6x the wall clock - The Wilson normaliser converges in 8-12 IRLS iterations instead of 26-102. Its stopping rule was ``|dL|`` per reflection against 1e-10, which asks eleven significant digits of a normalisation curve; it is now relative to the improvement so far, which is scale-invariant for the same reason the absolute form was chosen - `` = 1`` is solved in closed form for the intercept, so the identity no longer degrades as the convergence tolerance is loosened diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index c6331c5a..da186b28 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -292,7 +292,23 @@ class directly for finer control / access to the ranked candidate list. device : torch.device, optional Compute device (defaults to the model's device). verbose : int - 0 silent, 1 summary, ≥2 adds a per-stage wall-clock table. + How much the run says about itself. Each level is a superset of the one + below, and the boundaries are chosen so that a level is useful on its + own rather than being "a bit more of the same": + + 0 + Silent. + 1 + What happened: the search settings, one line per stage, and the + winner. Enough to see that a run did the expected work. + 2 + **Why it chose what it chose.** One ``CAND`` line per rotation + candidate carrying every score the selection could have used, plus + the per-stage wall-clock table. This is the level that makes the + pipeline diagnosable without a second implementation of its own + scoring -- see :meth:`_log_candidate`. + 3 + Per-translation-peak detail inside each candidate. Examples -------- @@ -365,6 +381,44 @@ def __init__( self._tmask = None self._eye3 = torch.eye(3, dtype=torch.float64) + #: Levels are documented on the class. They are a contract, not a dial: + #: level 2 is specifically "one machine-readable line per candidate", and + #: anything added at that level should preserve that. + def _log(self, level: int, msg: str) -> None: + """Emit ``msg`` if the run is at least this verbose. + + One emitter rather than ``if self.verbose > 0: print(...)`` at every + site. The scattered form is how levels drift -- the same stage ends up + reporting at 1 in one place and 2 in another, and nothing enforces that + a level means the same thing twice. + """ + if self.verbose >= level: + print(msg, flush=True) + + def _log_candidate(self, k: int, peak, r_analytic, t_frac, + tf_score=None) -> None: + """One line per rotation candidate, with every score behind the choice. + + Machine-readable on purpose. Diagnosing a wrong placement means asking + which candidate won and on what, and the only alternative to emitting it + here is a harness that re-implements the placement loop -- which drifts + from the pipeline and then disagrees with it about which candidate the + pipeline picked. A caller that knows the true orientation (a benchmark) + can join these lines against it; the pipeline cannot, and does not try. + + Fields are ``key=value`` so a reader does not depend on column order: + ``k`` candidate index in rotation-function order, ``rf``/``rfz`` its + score and z, ``tf`` the translation correlation at the chosen peak, + ``r`` the analytical-scale R that ranks it, ``t`` the fractional + translation. + """ + tf = "nan" if tf_score is None else f"{float(tf_score):.5f}" + t = ",".join(f"{float(x):.4f}" for x in t_frac) + self._log(2, f"CAND k={k} rf={float(peak.score):.4f} " + f"rfz={float(peak.sigma):.3f} tf={tf} " + f"r={float(r_analytic):.5f} t={t}") + + # ------------------------------------------------------------------ # Public entry point # ------------------------------------------------------------------ @@ -407,14 +461,9 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: if not do_translation: rotated, R_rec = self._make_rotated(candidates[0]) top = candidates[0] - if self.verbose > 0: - print( - f"mr: top peak RF = {top.score:.2f} " - f"(σ_Z = {top.sigma:.2f}); applying R⁻¹ to coords.", - flush=True, - ) - if self.verbose >= 2: - print("\n" + timer.summary(), flush=True) + self._log(1, f"mr: top peak RF = {top.score:.2f} " + f"(σ_Z = {top.sigma:.2f}); applying R⁻¹ to coords.") + self._log(2, "\n" + timer.summary()) return [ MRSolution( rotation=R_rec.detach().cpu().numpy(), @@ -429,25 +478,23 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: # --- Stage 2: per-candidate translation search --- self._prepare_translation_arrays() n_rot = min(self.n_rotation_candidates, len(candidates)) - if self.verbose > 0 and n_rot > 1: - print(f"mr: placing all {n_rot} rotation candidates…", flush=True) + if n_rot > 1: + self._log(1, f"mr: placing all {n_rot} rotation candidates…") solutions: List[MRSolution] = [] for k in range(n_rot): peak_k = candidates[k] rotated_k, R_rec_k = self._make_rotated(peak_k) - if self.verbose > 0: - print( - f"\nfit_to_data: rot{k} " - f"(LLG={peak_k.score:.2f}, σ_Z={peak_k.sigma:.2f})", - flush=True, - ) + self._log(3, f"\nfit_to_data: rot{k} " + f"(RF={peak_k.score:.2f}, σ_Z={peak_k.sigma:.2f})") placement = self._placement_for_candidate(rotated_k) if placement is None: - if self.verbose > 0: - print(" no translation peaks; skipping", flush=True) + self._log(2, f"CAND k={k} rf={float(peak_k.score):.4f} " + f"rfz={float(peak_k.sigma):.3f} tf=nan r=nan " + f"t=none # no translation peaks") continue - r_analytic, t_refined = placement + r_analytic, t_refined, tf_score = placement + self._log_candidate(k, peak_k, r_analytic, t_refined, tf_score) placed = rotated_k.copy().translate( t_refined.to(self.model.dtype_float), fractional=True, @@ -479,14 +526,9 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: timer.stop("12_final_scaler") winner.model.last_alignment_rfactor = rwork_final winner.r_factor = rwork_final - if self.verbose > 0: - print( - f"mr: winner analytical-TF R={winner.translation_score:.4f}, " - f"final Scaler-fit R-work={rwork_final:.4f}", - flush=True, - ) - if self.verbose >= 2: - print("\n" + timer.summary(), flush=True) + self._log(1, f"mr: winner analytical-TF R={winner.translation_score:.4f}, " + f"final Scaler-fit R-work={rwork_final:.4f}") + self._log(2, "\n" + timer.summary()) return solutions # ------------------------------------------------------------------ @@ -497,12 +539,8 @@ def _rotation_candidates(self, frf) -> list: timer = self._timer timer.start("3_rotation_search") - if self.verbose > 0: - print( - f"mr: rotation search (n_peaks={self.n_rotation_peaks}, " - f"model error {self.model_error_A:.2f} A)…", - flush=True, - ) + self._log(1, f"mr: rotation search (n_peaks={self.n_rotation_peaks}, " + f"model error {self.model_error_A:.2f} A)…") peaks, _lmax, _d_min = search_peaks( self.model, self.data, self.model_error_A, U_aniso=frf.U_aniso, n_peaks=self.n_rotation_peaks, @@ -581,21 +619,26 @@ def _prepare_translation_arrays(self) -> None: n_shells=max(self.n_shells // 2, 8), device=device, ) - if self.verbose > 0: + if self.verbose >= 1: d_hi = 1.0 / float(self._obs.s_mag.max()) d_lo = 1.0 / float(self._obs.s_mag.min().clamp(min=1e-9)) - print( - f"mr: translation set {self._obs.F_obs.numel()} reflections, " - f"{d_lo:.1f}-{d_hi:.2f} A" - + ("" if sig_F_full is not None else " (no sigmas: unit weight)"), - flush=True, - ) + self._log(1, f"mr: translation set {self._obs.F_obs.numel()} " + f"reflections, {d_lo:.1f}-{d_hi:.2f} A" + + ("" if sig_F_full is not None + else " (no sigmas: unit weight)")) def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: """Translation search + analytical-R local refine for one rotation. - Returns ``(r_analytic, t_refined)`` for the best translation of this - rotation candidate, or ``None`` if no translation peaks were found. + Returns ``(r_analytic, t_refined, tf_score)`` for the best translation + of this rotation candidate, or ``None`` if no translation peaks were + found. + + ``tf_score`` is the translation function's own score at its top peak. + It does not select anything -- ``r_analytic`` does -- but it is carried + out so ``verbose >= 2`` can report both. Which of the two a wrong + placement disagreed on is the first thing anyone diagnosing one asks, + and it is not recoverable afterwards from the winner alone. """ data = self.data timer = self._timer @@ -632,10 +675,10 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: if self.use_llg_tf: t_peaks = self._llg_tf_rescore(t_peaks, G_pre, h_R_pre) - if self.verbose > 0: + tf_top = float(t_peaks[0].score) + if self.verbose >= 3: tt = tuple(round(float(x), 3) for x in t_peaks[0].translation.tolist()) - print(f" top translation t={tt} score={t_peaks[0].score:.4f}", - flush=True) + self._log(3, f" top translation t={tt} score={tf_top:.4f}") best = None for k_t, tp in enumerate(t_peaks[:self.n_translation_candidates]): @@ -649,15 +692,11 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: precomputed_G=G_pre, precomputed_h_R=h_R_pre, ) timer.stop("7_local_TF_refine") - if self.verbose > 0: - print( - f" trans{k_t}: R(analytic)={r_analytic:.4f}, " - f"t={[round(float(x), 3) for x in t_refined.tolist()]}", - flush=True, - ) + self._log(3, f" trans{k_t}: R(analytic)={r_analytic:.4f}, " + f"t={[round(float(x), 3) for x in t_refined.tolist()]}") if best is None or r_analytic < best[0]: best = (r_analytic, t_refined) - return best + return None if best is None else (best[0], best[1], tf_top) def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): """Re-rank translation peaks by a shared-σA Rice/Woolfson LLG. From 6373e0bdeeae3287d9a88a52af810327d0346a2f Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 11:01:14 +0200 Subject: [PATCH 121/250] Name the node-field ADP targets so a weight can reach them LossState.register_targets keys each component off target.name. Every ADP target declares its own path -- "adp/simu", "adp/sigd", "adp/locality" -- except these two, which declared none and so inherited ModelTarget.name, the literal string "model_target". Both registered under that, the second overwrote the first, and get_effective_weight never saw an "adp/..." path, so DEFAULT_GROUP_WEIGHTS["adp/node_load"] = 10.0 was a dead entry and the coverage barrier was in no refinement's loss at all. Every test passed because they called the target directly, which works. One even asserted the weight was present and positive, which was true and meaningless. The diagnostic that shows it is comparing TotalADPTarget's own component set against the LossState's keys. The guard added here walks every concrete ModelTarget subclass and fails any still carrying a base placeholder name, plus checks end to end that each declared component reaches the LossState with an addressable weight. Group bases are excluded by having subclasses -- they are not formally abstract, so that is the only reliable marker. Also stops node_smoothness reading float() off a grad-carrying tensor. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- .../test_targets_reach_the_loss_state.py | 117 ++++++++++++++++++ torchref/refinement/targets/adp/node_load.py | 7 ++ .../refinement/targets/adp/node_smoothness.py | 9 +- 3 files changed, 132 insertions(+), 1 deletion(-) create mode 100644 tests/unit/refinement/test_targets_reach_the_loss_state.py diff --git a/tests/unit/refinement/test_targets_reach_the_loss_state.py b/tests/unit/refinement/test_targets_reach_the_loss_state.py new file mode 100644 index 00000000..170e5cfd --- /dev/null +++ b/tests/unit/refinement/test_targets_reach_the_loss_state.py @@ -0,0 +1,117 @@ +"""Every registered target must arrive in the LossState under a key its weight can reach. + +``LossState.register_targets`` keys each component off ``target.name``, falling back to +``Target.name`` -- which is the literal string ``"model_target"``. A target that forgets +to declare its own name therefore registers under that, collides with every other target +that forgot, and no hierarchical weight can address it. The term is constructed, callable +and correct in isolation; it simply never enters the loss. + +That is what happened to ``adp/node_load`` and ``adp/node_smoothness``: both shipped +unnamed, so the node-coverage barrier was never in any refinement's loss despite having a +weight of 10.0 in ``DEFAULT_GROUP_WEIGHTS`` and passing every test that called it +directly. These tests check the plumbing rather than the arithmetic. +""" + +import pytest + +from torchref.refinement.targets.base import ModelTarget + + +def _leaf_target_classes(): + """Every concrete ModelTarget subclass that is a loss component, not a container.""" + import inspect + import pkgutil + import importlib + + import torchref.refinement.targets as pkg + from torchref.refinement.targets.combined import CombinedModelTargets + + for mod in pkgutil.walk_packages(pkg.__path__, pkg.__name__ + "."): + try: + importlib.import_module(mod.name) + except Exception: + continue + seen = {} + stack = [ModelTarget] + while stack: + cls = stack.pop() + for sub in cls.__subclasses__(): + stack.append(sub) + # Skip group bases (ADPTarget, GeometryTarget, ...). They are not formally + # abstract -- nothing in them is an abstractmethod -- so the only reliable + # marker is that other targets derive from them. A base legitimately carries + # the placeholder name because it is never registered itself. + if sub.__subclasses__(): + continue + if inspect.isabstract(sub) or issubclass(sub, CombinedModelTargets): + continue + seen[sub.__qualname__] = sub + return seen + + +@pytest.mark.unit +def test_no_component_target_inherits_the_placeholder_name(): + """A component using the base placeholder cannot be addressed by any weight.""" + offenders = { + name: cls.name + for name, cls in _leaf_target_classes().items() + if getattr(cls, "name", None) in (None, "base_target", "model_target", + "data_target") + } + assert not offenders, ( + "these targets would register under the base placeholder name, colliding with " + "each other and unreachable by any hierarchical weight:\n " + + "\n ".join(f"{k}: name={v!r}" for k, v in sorted(offenders.items())) + ) + + +@pytest.mark.unit +def test_adp_component_names_are_hierarchical(): + """An ADP component must sit under the ``adp`` group or the group weight misses it.""" + import torchref.refinement.targets.adp as adp_pkg + + bad = {} + for attr in dir(adp_pkg): + cls = getattr(adp_pkg, attr) + if not isinstance(cls, type) or not issubclass(cls, ModelTarget): + continue + if cls.__subclasses__(): + continue # a group base, never registered itself + name = getattr(cls, "name", "") + if not isinstance(name, str) or not name.startswith("adp/"): + bad[attr] = name + assert not bad, f"ADP targets not under the adp group: {bad}" + + +@pytest.mark.unit +@pytest.mark.parametrize("mode_set", [None, "rigid_dilation"]) +def test_field_components_are_registered_and_weighted(pdb_dir, mtz_dir, mode_set): + """End to end: the components the representation declares must be in the loss. + + Compares the combined target's own component set against the LossState's keys, so a + component that exists but never registers is caught -- which is the failure mode that + a direct ``adp_target['node_load']()`` call cannot see. + """ + from torchref.refinement.lbfgs_refinement import LBFGSRefinement + + ref = LBFGSRefinement( + data_file=str(mtz_dir / "1DAW.mtz"), pdb=str(pdb_dir / "1DAW.pdb"), + verbose=0, adp_mode="field_aniso", adp_mode_set=mode_set, + ) + components = set(ref.adp_target.target_losses()) + state = ref.complete_loss_state() + + for component in components: + key = f"adp/{component}" + assert key in state.targets, ( + f"{component!r} is a component of TotalADPTarget but never reached the " + f"LossState. Registered adp keys: " + f"{sorted(k for k in state.targets if k.startswith('adp'))}" + ) + # And the weight has to be addressable, not merely present. + assert state.get_effective_weight(key) is not None + + assert "node_load" in components, "field mode must carry the coverage barrier" + assert state.get_effective_weight("adp/node_load") > 0.0, ( + "the coverage barrier registered but at zero effective weight, so it is inert" + ) diff --git a/torchref/refinement/targets/adp/node_load.py b/torchref/refinement/targets/adp/node_load.py index 7d62be8e..559e725a 100644 --- a/torchref/refinement/targets/adp/node_load.py +++ b/torchref/refinement/targets/adp/node_load.py @@ -59,6 +59,13 @@ class NodeLoadTarget(ADPTarget): Verbosity level. Default is 0. """ + #: Hierarchical key this target registers under. Required, not cosmetic: + #: LossState.register_targets takes the key from ``.name``, so without it the + #: target inherits ``Target.name`` ("model_target"), registers under that, + #: collides with every other unnamed target, and no ``adp/...`` weight can + #: reach it -- the term is then built, callable, and never in the loss. + name: str = "adp/node_load" + def __init__( self, model: "Model" = None, diff --git a/torchref/refinement/targets/adp/node_smoothness.py b/torchref/refinement/targets/adp/node_smoothness.py index 5b22a8d5..6ae5634b 100644 --- a/torchref/refinement/targets/adp/node_smoothness.py +++ b/torchref/refinement/targets/adp/node_smoothness.py @@ -64,6 +64,13 @@ class NodeSmoothnessTarget(ADPTarget): Verbosity level. Default is 0. """ + #: Hierarchical key this target registers under. Required, not cosmetic: + #: LossState.register_targets takes the key from ``.name``, so without it the + #: target inherits ``Target.name`` ("model_target"), registers under that, + #: collides with every other unnamed target, and no ``adp/...`` weight can + #: reach it -- the term is then built, callable, and never in the loss. + name: str = "adp/node_smoothness" + def __init__( self, model: "Model" = None, @@ -120,7 +127,7 @@ def forward(self) -> torch.Tensor: return torch.zeros((), device=self.device) w, diff2, _ = self._pair_terms() total = w.sum() - if float(total) <= 0.0: + if float(total.detach()) <= 0.0: return torch.zeros((), device=self.device) return (w * diff2).sum() / total From 828e12c55e7468e7a283a0e262393f7876a938e3 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 11:01:14 +0200 Subject: [PATCH 122/250] Register only the ADP restraints the representation calls for simu and locality restrain by penalty the spatial smoothness a node field enforces by construction, and node_load/node_smoothness have nothing to act on off it. TotalADPTarget now builds the applicable set instead of building all of them and zero-weighting half. Registering-then-zeroing costs nothing at run time -- LossState.aggregate skips a zero-weight target -- but it leaves the correctness of the setup resting on a number, so anyone adjusting the adp group weight for their own reasons silently re-enables a restraint that double-counts the parametrisation. Whether a term applies is a property of the representation, not something to tune. sigd applies either way: it is a prior on the marginal B distribution, which a field constrains no more than a per-atom model does. The node-load tests now cover every payload. They were written against the isotropic payload, whose storage is five columns wide, and poked column 1 by index; a mode payload is 25 to 82 columns with log sigma near the end, so "acts through the weights, never the values" had to be re-established rather than assumed. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- .../unit/refinement/test_node_load_target.py | 61 +++++++++++++++++++ torchref/refinement/targets/combined.py | 52 +++++++++++----- 2 files changed, 98 insertions(+), 15 deletions(-) diff --git a/tests/unit/refinement/test_node_load_target.py b/tests/unit/refinement/test_node_load_target.py index fcb36a0f..6ec6a3e8 100644 --- a/tests/unit/refinement/test_node_load_target.py +++ b/tests/unit/refinement/test_node_load_target.py @@ -141,3 +141,64 @@ def test_default_weight_exists_for_the_component(pdb_path): assert "adp/node_load" in DEFAULT_GROUP_WEIGHTS assert DEFAULT_GROUP_WEIGHTS["adp/node_load"] > 0 + + +# ---------------------------------------------------------------------------------- +# Payload independence. The barrier was written against the isotropic payload, whose +# storage is five columns wide, and the tests above poke column 1 by hand because that +# is where its log sigma sits. A displacement-mode field is 25 to 82 columns wide with +# log sigma near the end, so "acts through the weights, never the values" has to be +# re-established rather than assumed to carry over. +# ---------------------------------------------------------------------------------- + + +@pytest.mark.unit +@pytest.mark.parametrize( + "mode_set,width", [(None, 6), ("rigid", 21), ("rigid_dilation", 28), ("affine", 78)] +) +def test_barrier_prices_geometry_not_values_for_every_payload(pdb_path, mode_set, width): + """Zero gradient on the payload, non-zero on kernel width and node position. + + The invariant that separates this from a magnitude prior: it removes the opportunity + to place an extreme ADP rather than penalising the ADP. A payload-width bug would + show up here as gradient leaking into the payload columns. + """ + model = Model(verbose=0) + model.load_pdb(pdb_path) + model.set_adp_mode( + "field_aniso", n_nodes=24, k_neighbors=12, mode_set=mode_set + ) + field = model.adp_field + assert field.payload.width == width + assert field.node_shape[1] == width + 4 # payload | log sigma | 3 offset + + NodeLoadTarget(model)().backward() + grad = field.refinable_params.grad + assert grad is not None + assert float(grad[:, :width].abs().sum()) == pytest.approx(0.0, abs=1e-12), ( + "the barrier reached the node VALUES" + ) + assert float(grad[:, width].abs().sum()) > 0, "no gradient to the kernel width" + assert float(grad[:, width + 1 : width + 4].abs().sum()) > 0, ( + "no gradient to the node positions" + ) + + +@pytest.mark.unit +@pytest.mark.parametrize("mode_set", ["rigid", "rigid_dilation", "affine"]) +def test_barrier_still_catches_a_starved_node_on_a_mode_field(pdb_path, mode_set): + model = Model(verbose=0) + model.load_pdb(pdb_path) + model.set_adp_mode("field_aniso", n_nodes=24, k_neighbors=12, mode_set=mode_set) + field = model.adp_field + target = NodeLoadTarget(model) + before = float(target()) + + # log sigma is the column after the payload, wherever the payload ends. + with torch.no_grad(): + field.refinable_params[0, field.payload.width] -= 6.0 + field.reset_forward_cache() + + load = field.node_load().detach() + assert float((load / load.mean()).min()) < 0.25, "the node was not actually starved" + assert float(target()) > before + 1.0 diff --git a/torchref/refinement/targets/combined.py b/torchref/refinement/targets/combined.py index 2fa318f9..032d8e9f 100644 --- a/torchref/refinement/targets/combined.py +++ b/torchref/refinement/targets/combined.py @@ -365,21 +365,43 @@ class TotalADPTarget(CombinedModelTargets): """ def _create_targets(self) -> Dict[str, Target]: - """Build the three ADP component targets.""" - print("Initializing TotalADPTarget with component targets...") - return { - "simu": ADPSimilarityTarget(self.model, verbose=self.verbose), - "locality": ADPLocalityTarget( - self.model, verbose=self.verbose - ), - "sigd": ADPSigdTarget(self.model, verbose=self.verbose), - # Inert unless the model is in field mode, so it costs a zero tensor - # per call on the per-atom path. - "node_load": NodeLoadTarget(self.model, verbose=self.verbose), - "node_smoothness": NodeSmoothnessTarget( - self.model, verbose=self.verbose - ), - } + """Build the ADP component targets that apply to the model's representation. + + Only the applicable ones are registered, rather than registering all of them and + zero-weighting the inapplicable half. ``simu`` and ``locality`` restrain by + penalty exactly the spatial smoothness a node field enforces by construction, so + in field mode they are not a weak prior but a duplicate of the parametrisation; + and ``node_load`` / ``node_smoothness`` have nothing to act on off it. + + Registering-then-zeroing would cost nothing at run time --- ``LossState.aggregate`` + skips a zero-weight target --- but it leaves the correctness of the setup resting + on a weight, so anyone who touches the ``adp`` group weight for their own reasons + silently re-enables a double-counted restraint. Whether a term applies is a + property of the representation, not a number to be tuned. + + ``sigd`` applies either way: it is a prior on the marginal B distribution, which + a field constrains no more than a per-atom parametrisation does. + """ + if self.model.adp_is_field: + targets = { + "sigd": ADPSigdTarget(self.model, verbose=self.verbose), + "node_load": NodeLoadTarget(self.model, verbose=self.verbose), + "node_smoothness": NodeSmoothnessTarget( + self.model, verbose=self.verbose + ), + } + else: + targets = { + "simu": ADPSimilarityTarget(self.model, verbose=self.verbose), + "locality": ADPLocalityTarget(self.model, verbose=self.verbose), + "sigd": ADPSigdTarget(self.model, verbose=self.verbose), + } + if self.verbose > 0: + print( + "Initializing TotalADPTarget with component targets: " + + ", ".join(targets) + ) + return targets def print_statistics(self) -> None: """ From cc7cb62f85595f78d2276630a9948577b9a17f15 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 11:01:36 +0200 Subject: [PATCH 123/250] Persist which payload a saved disorder field carries state_dict holds tensors, and the payload is a plain object that never reaches it, so a restore had to infer it from the slot. That is not a clean failure: a mode payload rebuilt as a constant-U one has the wrong storage width and surfaces as a shape mismatch, or worse as silently different ADPs. The field now registers an integer code. PAYLOAD_CODES is append-only -- a saved state dict holds the number, so renumbering would restore the wrong payload. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- torchref/model/disorder_field.py | 62 ++++++++++++++++++++++++++++++++ 1 file changed, 62 insertions(+) diff --git a/torchref/model/disorder_field.py b/torchref/model/disorder_field.py index 3cb8124a..c11d91fc 100644 --- a/torchref/model/disorder_field.py +++ b/torchref/model/disorder_field.py @@ -42,6 +42,9 @@ "AnisotropicPayload", "ModeCovariancePayload", "MODE_SETS", + "PAYLOAD_CODES", + "payload_code", + "payload_from_code", "farthest_point_anchors", "density_anchor_rows", "build_neighbor_list", @@ -546,6 +549,55 @@ def _ridged_solve(w_dense, target): return torch.linalg.solve(gram + ridge * eye, w_dense.T @ target) +#: Stable integer code per payload, so a saved field can rebuild the one it had. +#: ``state_dict`` holds tensors only, and the payload is a plain object that never +#: reaches it, so without this a restore has to guess -- and guessing wrong is not a +#: clean failure: a mode payload restored as a constant-U one has the wrong storage +#: width and only shows up as a shape mismatch, or worse, silently different ADPs. +#: +#: **Append, never renumber.** A saved state dict holds the number. +PAYLOAD_CODES = { + "isotropic": 0, + "anisotropic": 1, + "modes:constant": 2, + "modes:rigid": 3, + "modes:rigid_dilation": 4, + "modes:affine": 5, +} + + +def payload_code(payload: "NodePayload") -> int: + """Code identifying ``payload`` well enough to rebuild it.""" + if isinstance(payload, ModeCovariancePayload): + key = f"modes:{payload.mode_set}" + elif isinstance(payload, AnisotropicPayload): + key = "anisotropic" + else: + key = "isotropic" + if key not in PAYLOAD_CODES: + raise ValueError( + f"Payload {key!r} has no code in PAYLOAD_CODES, so a field carrying it " + "cannot be saved and restored. Add one (appending, never renumbering)." + ) + return PAYLOAD_CODES[key] + + +def payload_from_code(code: int, epsilon: float = 1e-3) -> "NodePayload": + """Rebuild the payload a saved ``code`` names.""" + names = {v: k for k, v in PAYLOAD_CODES.items()} + key = names.get(int(code)) + if key is None: + raise ValueError( + f"Unknown payload code {code!r}. It was written by a newer TorchRef than " + f"this one, which knows {sorted(PAYLOAD_CODES)}." + ) + if key == "isotropic": + return IsotropicPayload() + if key == "anisotropic": + return AnisotropicPayload(epsilon=epsilon) + return ModeCovariancePayload(key.split(":", 1)[1], epsilon=epsilon) + + class DisorderFieldTensor(MixedTensor): """Per-atom ADPs from a small set of nodes, each atom a weighted mean of its k nearest. @@ -633,6 +685,10 @@ def __init__( self.register_buffer("neighbor_list", None) self.register_buffer("anchor_atom", None) self.register_buffer("anchor_node", None) + self.register_buffer( + "payload_code", + torch.tensor(payload_code(self._payload), dtype=torch.int64), + ) return if xyz_fn is None: @@ -696,6 +752,12 @@ def __init__( self.register_buffer("anchor_atom", anchor_atom) self.register_buffer("anchor_node", anchor_node) self.register_buffer("neighbor_list", neighbor_list) + # Which payload this field carries, so a restore rebuilds it rather than + # inferring it from the storage width. + self.register_buffer( + "payload_code", + torch.tensor(payload_code(self._payload), dtype=torch.int64, device=device), + ) # ------------------------------------------------------------------ # Construction helpers. From 2f521b33585afba620197fa45b6bc4e44162da5d Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 11:01:36 +0200 Subject: [PATCH 124/250] Rebuild the saved payload on restore, and add a flat field initialisation The restore path now reads the payload code rather than guessing from the slot, falling back to the old inference for state dicts written before the code existed. set_adp_mode gains init={"fit","flat"}. "fit" (default) fits the field to the model's current per-atom ADPs, which is right when those mean something. "flat" keeps only their level and discards the spatial structure, for when they do not -- an AlphaFold model's B values come from a pLDDT conversion, and fitting a smooth basis to them spends the field's parameters reproducing structure it cannot hold. Flattening goes through the equivalent isotropic B and hands the payload a 1-D target, which its fit lifts to U_iso * I. Taking a median over all six U components instead sets the off-diagonals equal to the diagonals, giving eigenvalues (3L, 0, 0) -- singular, and NaN once the Cholesky encode takes log(diag - epsilon). Found by smoking it. Measured: flat initialisation is worth -0.0005 R-free (p=0.29, n=60) from an AlphaFold start, so it is an option rather than a better default. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- torchref/model/model.py | 53 ++++++++++++++++++++++++++++++++++++++++- 1 file changed, 52 insertions(+), 1 deletion(-) diff --git a/torchref/model/model.py b/torchref/model/model.py index 0e9c5dbe..a4bb4888 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -1253,6 +1253,7 @@ def set_adp_mode( k_neighbors: int = 12, refine_node_positions: bool = True, mode_set: str = None, + init: str = "fit", ): """Set the atomic displacement parameter (ADP) parametrization. @@ -1289,6 +1290,10 @@ def set_adp_mode( Give each node a refinable offset from its anchor centroid, at three extra parameters per node. On by default: it is what lets the load-balancing restraint move a node toward atoms instead of only widening its kernel. + init : {"fit", "flat"}, optional + What a field mode fits its nodes to: ``"fit"`` (default) the model's current + per-atom ADPs, ``"flat"`` a single level with their spatial structure + discarded. See :meth:`_install_disorder_field`. mode_set : str, optional For ``mode="field_aniso"``, a key of :data:`~torchref.model.disorder_field.MODE_SETS` --- ``"rigid"`` is TLS, @@ -1349,6 +1354,7 @@ def set_adp_mode( refine_node_positions=refine_node_positions, anisotropic=aniso, mode_set=mode_set, + init=init, ) return if mode == "isotropic": @@ -1395,6 +1401,7 @@ def _install_disorder_field( refine_node_positions: bool = False, anisotropic: bool = False, mode_set: str = None, + init: str = "fit", ): """Replace a per-atom ADP wrapper with a node field fitted to it. @@ -1406,6 +1413,21 @@ def _install_disorder_field( ``mode_set`` selects a displacement-mode payload in place of the constant-U one, which is the difference between a node holding a single ADP and a node holding a motion whose ADP varies across its region. + + ``init`` chooses what the field is fitted to: + + ``"fit"`` + The per-atom ADPs the model currently holds. Right when those mean something + --- a deposited or already-refined model --- because the field then starts + from a state whose R-factor is known. + ``"flat"`` + A single value, the median of those ADPs. Right when they do not mean + anything. An AlphaFold model's B values come from a pLDDT conversion, and + fitting a smooth basis to them spends the field's parameters reproducing + structure it cannot hold and that is not worth holding: measured on 2A25, the + fitted field starts 0.025 R-free WORSE than a flat one, before any + refinement. The level is kept because it is close to right and the scaler + owns it anyway; only the spatial structure is discarded. """ from torchref.model.disorder_field import ( AnisotropicPayload, @@ -1430,6 +1452,27 @@ def _install_disorder_field( if anisotropic else self.adp().detach().clone() ) + if init == "flat": + # Flatten through the equivalent isotropic B, and hand the payload a 1-D + # target: its ``fit`` lifts that to U_iso * I. Taking a median over all + # six U components instead would set the off-diagonals equal to the + # diagonals, giving eigenvalues (3L, 0, 0) -- singular, and NaN once the + # Cholesky encode takes log(diag - epsilon). + b = ( + (8.0 * math.pi**2 / 3.0) * target[:, :3].sum(dim=1) + if target.ndim == 2 + else target + ) + finite = torch.isfinite(b) + if not bool(finite.any()): + raise ValueError("cannot flatten an all-NaN ADP target") + level = b[finite].median() + target = torch.where(finite, level.expand_as(b), b) + elif init != "fit": + raise ValueError( + f"init={init!r}; expected 'fit' (use the model's own ADPs) or " + "'flat' (discard their spatial structure, keep the level)." + ) B = target if n_nodes is None: n_nodes = max(4, int(round(len(self.pdb) / 25.0))) @@ -2173,9 +2216,17 @@ def _restore_adp_slot(prefix, state_dict, pdb, saved_dtype, xyz_wrapper): AnisotropicPayload, DisorderFieldTensor, IsotropicPayload, + payload_from_code, ) - payload = AnisotropicPayload() if aniso else IsotropicPayload() + # The saved code names the payload exactly. Fall back to inferring it from the + # slot for state dicts written before the code existed, where the only payloads + # were the two the slot already implies. + saved_code = state_dict.get(f"{prefix}.payload_code") + if saved_code is not None: + payload = payload_from_code(int(saved_code)) + else: + payload = AnisotropicPayload() if aniso else IsotropicPayload() saved_values = state_dict[f"{prefix}.fixed_values"] # Rebuild with the SAVED anchor rows: cluster anchoring makes these length # n_atoms where single-atom anchoring makes them length K, so reconstructing From a610fd7acc500e3144b9bd4878a2dbd976bd639c Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 11:01:58 +0200 Subject: [PATCH 125/250] Give Refinement an entry point for switching the ADP representation Model.set_adp_mode changes the representation but cannot size it -- node count follows from the reflection count, and the model has no idea how much data there is -- and cannot swap the ADP restraint set. Going through the model directly is a partial setup. set_adp_representation sizes a field from the work set, rebuilds the targets, scales and LossState, and is safe to call after construction, which is what the model's own "run once at setup" caveat is about. __init__ routes through the same method, so construction and a later switch cannot drift apart. Node cost is read off the payload object rather than tabulated, so a new payload cannot desynchronise the arithmetic from what the field allocates. Default target: 7 work reflections per ADP parameter. PDB-REDO holds ~7 across its whole resolution range and changes model form to stay there; measured on 179 of their entries a node field peaks at the same value, with both directions worse. The loss is deliberately NOT rebalanced for a field. The point of the representation is that smoothness comes from the parametrisation, so it should need less regularisation, not a reweighted version of the same priors. An earlier FIELD_GROUP_WEIGHTS raising the adp group had two side effects worth recording: adp/scaler_U and adp/scaler_log_scale sit under that group, so it multiplied the scaler regularisation by the same factor, and it made every field run differ from its baseline in a way unrelated to ADPs. Measured, it was purely harmful from an AlphaFold start (+0.0067 R-free) and irrelevant on converged models. flatten_adp_field discards the field's structure mid-refinement, keeping the level. A field fits its structure once, when installed, and nothing re-derives it, so structure inferred against wrong coordinates survives -- the same shape as freezing bulk solvent at the starting model. Worth -0.0098 R-free (p=0.002, n=60) from an AlphaFold start. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- .../integration/test_adp_field_refinement.py | 398 ++++++++++++++++++ torchref/refinement/base_refinement.py | 246 ++++++++++- 2 files changed, 641 insertions(+), 3 deletions(-) create mode 100644 tests/integration/test_adp_field_refinement.py diff --git a/tests/integration/test_adp_field_refinement.py b/tests/integration/test_adp_field_refinement.py new file mode 100644 index 00000000..4dca8c03 --- /dev/null +++ b/tests/integration/test_adp_field_refinement.py @@ -0,0 +1,398 @@ +"""A node-field ADP representation driven the way production drives it. + +The unit tests cover the payload arithmetic and the wiring. What they cannot cover is +whether the thing is usable: whether a refinement constructed with a field mode sizes +itself from the data, carries weights appropriate to a field rather than to the per-atom +representation it replaced, actually reduces R-work over real cycles, survives a +checkpoint round trip, and can be switched into and out of after setup. + +Every one of those was broken or absent at some point in this feature's life, and none +of them fails a unit test. +""" + +import pytest +import torch + +from torchref.model.disorder_field import ( + DisorderFieldTensor, + ModeCovariancePayload, + payload_code, +) +from torchref.refinement.base_refinement import DEFAULT_GROUP_WEIGHTS +from torchref.refinement.lbfgs_refinement import LBFGSRefinement + +MODE_SET = "rigid_dilation" + + +@pytest.fixture(scope="module") +def files(mtz_dir, pdb_dir): + return str(mtz_dir / "1DAW.mtz"), str(pdb_dir / "1DAW.pdb") + + +def _refinement(files, **kw): + mtz, pdb = files + return LBFGSRefinement(data_file=mtz, pdb=pdb, verbose=0, **kw) + + +@pytest.fixture(scope="module") +def field_refinement(files): + """Built the way a caller would: mode and mode set, no explicit node count.""" + return _refinement(files, adp_mode="field_aniso", adp_mode_set=MODE_SET) + + +# ---------------------------------------------------------------------------------- +# Setup: sized from the data, weighted for a field. +# ---------------------------------------------------------------------------------- + + +@pytest.mark.integration +def test_construction_installs_the_mode_field(field_refinement): + ref = field_refinement + assert ref.model.adp_is_field + field = ref.model.adp_field + assert isinstance(field.payload, ModeCovariancePayload) + assert field.payload.mode_set == MODE_SET + + +@pytest.mark.integration +def test_node_count_comes_from_the_reflections_not_the_atoms(field_refinement): + """The whole point of putting the sizing on the refinement. + + The model's own default is one node per 25 atoms, which knows nothing about how much + data there is. Here the achieved ratio has to land near the requested one, and the + node count has to disagree with the atom-count rule (otherwise the test would pass + on a coincidence). + """ + ref = field_refinement + n_par = sum(p.numel() for p in ref.model.parameters_of_types(("adp", "u"))) + achieved = ref.data.work.n / n_par + assert 4.0 < achieved < 12.0, f"{achieved:.1f} work reflections per ADP parameter" + + atom_rule = max(4, round(len(ref.model.pdb) / 25.0)) + assert ref.model.adp_field.n_nodes != atom_rule + + +@pytest.mark.integration +@pytest.mark.parametrize("ratio", [3.5, 7.0, 15.0]) +def test_requested_ratio_is_honoured(files, ratio): + ref = _refinement( + files, adp_mode="field_aniso", adp_mode_set=MODE_SET, + reflections_per_adp_parameter=ratio, + ) + n_par = sum(p.numel() for p in ref.model.parameters_of_types(("adp", "u"))) + achieved = ref.data.work.n / n_par + # Node count is an integer, so the achieved ratio cannot match exactly; it must + # track, and it must be monotone in the request. + assert 0.6 * ratio < achieved < 1.7 * ratio, f"asked {ratio}, got {achieved:.1f}" + + +@pytest.mark.integration +def test_explicit_node_count_bypasses_the_budget(files): + ref = _refinement( + files, adp_mode="field_aniso", adp_mode_set=MODE_SET, n_nodes=11 + ) + assert ref.model.adp_field.n_nodes == 11 + + +@pytest.mark.integration +def test_field_mode_does_not_register_the_restraints_it_duplicates(field_refinement): + """``simu`` and ``locality`` penalise what the field enforces by construction. + + They must be absent from the component set, not present at weight zero: a zero is a + lever, and anyone adjusting the ``adp`` group weight for their own reasons would + silently re-enable a restraint that double-counts the parametrisation. + """ + components = field_refinement.adp_target.target_losses() + assert "simu" not in components + assert "locality" not in components + # What a field does need. + assert "node_load" in components + assert "node_smoothness" in components + assert "sigd" in components, "the marginal-B prior applies to either representation" + + # And the loss is NOT rebalanced for a field: the parametrisation is the constraint, + # so a field needs less regularisation than a per-atom model, not a reweighted + # version of the same priors. An earlier override also silently scaled the two + # adp/scaler_* terms, which have nothing to do with atomic ADPs. + weights = field_refinement.weighting() + for key, expected in DEFAULT_GROUP_WEIGHTS.items(): + assert weights.get(key) == expected, f"{key} diverged from the default" + + +@pytest.mark.integration +def test_per_atom_mode_does_not_register_the_node_targets(files): + """The converse: nothing node-shaped has anything to act on off field mode.""" + ref = _refinement(files, adp_mode="isotropic") + components = ref.adp_target.target_losses() + assert "simu" in components and "locality" in components + assert "node_load" not in components + assert "node_smoothness" not in components + + +@pytest.mark.integration +def test_switching_replaces_the_component_set_and_the_loss_state(files): + """A switch changes WHICH targets exist, so a cached LossState must not survive it.""" + ref = _refinement(files, adp_mode="isotropic") + before = set(ref.adp_target.target_losses()) + assert "simu" in before + # Force the LossState to exist so the switch has something stale to invalidate. + ref.complete_loss_state() + assert ref._loss_state is not None + + logger_before = ref.logger # binds a Logger to the state that is about to go + ref.set_adp_representation("field_aniso", mode_set=MODE_SET) + after = set(ref.adp_target.target_losses()) + # The Logger holds a reference to the LossState, so replacing the state without + # replacing the Logger would leave it recording into an object nothing else reads. + assert ref.logger is not logger_before + assert ref.logger.state is ref.loss_state + assert "simu" not in after and "node_load" in after + state = ref.complete_loss_state() + registered = set(state.targets) + assert not any(k.endswith("simu") or k.endswith("locality") for k in registered), ( + f"a stale component survived the switch: {sorted(registered)}" + ) + assert any(k.endswith("node_load") for k in registered) + + +@pytest.mark.integration +def test_switching_preserves_weights_it_does_not_own(files): + """A call about ADPs must not reset the caller's xray or geometry weights.""" + from torchref.refinement.weighting import ManualWeighting + + ref = _refinement(files, adp_mode="isotropic") + custom = {**ref.weighting(), "xray": 2.5, "geometry": 0.35} + ref.weighting = ManualWeighting(custom) + + ref.set_adp_representation("field_aniso", mode_set=MODE_SET) + weights = ref.weighting() + assert weights["xray"] == 2.5, "xray weight was clobbered by an ADP call" + assert weights["geometry"] == 0.35 + # Nothing is rebalanced, so the caller's adp weight survives too. + assert weights["adp"] == custom["adp"] + + ref.set_adp_representation("isotropic") + assert ref.weighting()["xray"] == 2.5 + + +@pytest.mark.integration +def test_per_atom_mode_keeps_the_per_atom_weights(files): + ref = _refinement(files, adp_mode="isotropic") + assert ref.weighting()["adp"] == DEFAULT_GROUP_WEIGHTS["adp"] + + +# ---------------------------------------------------------------------------------- +# It has to actually refine. +# ---------------------------------------------------------------------------------- + + +@pytest.mark.integration +def test_refinement_reduces_rwork_and_stays_finite(files): + """Two ADP-only cycles on real data. Nothing here may be NaN and R-work must fall.""" + ref = _refinement(files, adp_mode="field_aniso", adp_mode_set=MODE_SET) + rw0, rf0 = (float(v) for v in ref.get_rfactor()) + for _ in range(2): + ref.refine_scaler() + ref.refine_adp() + rw1, rf1 = (float(v) for v in ref.get_rfactor()) + + for name, v in (("Rwork", rw1), ("Rfree", rf1)): + assert v == v, f"{name} is NaN" + assert 0.0 < v < 0.7, f"{name} = {v:.4f} is not a plausible R-factor" + assert rw1 < rw0 + 1e-6, f"R-work rose: {rw0:.4f} -> {rw1:.4f}" + + u6 = ref.model.adp_u6().detach() + assert torch.isfinite(u6).all() + ev = torch.linalg.eigvalsh(_u6_to_matrix(u6)) + assert float(ev.min()) > 0.0, "an atom went non-positive-definite during refinement" + + +@pytest.mark.integration +def test_full_refine_moves_coordinates_without_staling_the_field(files): + """``refine()`` is scaler -> xyz -> ADP, so the field is read at moved coordinates. + + The field borrows the coordinate accessor rather than taking coordinates as an + argument, so its forward cache has to fold them into its key. If it does not, the + ADPs silently come from wherever the atoms used to be. + """ + ref = _refinement(files, adp_mode="field_aniso", adp_mode_set=MODE_SET) + xyz0 = ref.model.xyz().detach().clone() + u0 = ref.model.adp_u6().detach().clone() + + ref.refine(macro_cycles=1) + + xyz1 = ref.model.xyz().detach() + u1 = ref.model.adp_u6().detach() + assert not torch.allclose(xyz0, xyz1), "coordinates did not move, test proves nothing" + assert torch.isfinite(u1).all() + assert not torch.allclose(u0, u1), "ADPs unchanged after xyz moved -- stale cache" + + +def _u6_to_matrix(u6): + M = torch.zeros(u6.shape[0], 3, 3, dtype=u6.dtype) + M[:, 0, 0], M[:, 1, 1], M[:, 2, 2] = u6[:, 0], u6[:, 1], u6[:, 2] + M[:, 0, 1] = M[:, 1, 0] = u6[:, 3] + M[:, 0, 2] = M[:, 2, 0] = u6[:, 4] + M[:, 1, 2] = M[:, 2, 1] = u6[:, 5] + return M + + +# ---------------------------------------------------------------------------------- +# Switching after setup, which the model alone documents as unsupported. +# ---------------------------------------------------------------------------------- + + +@pytest.mark.integration +def test_switch_into_and_out_of_field_mode_after_setup(files): + ref = _refinement(files, adp_mode="isotropic") + assert not ref.model.adp_is_field + before = float(ref.get_rfactor()[0]) + + applied = ref.set_adp_representation("field_aniso", mode_set=MODE_SET) + assert ref.model.adp_is_field + assert applied["n_nodes"] >= 2 + assert "simu" not in ref.adp_target.target_losses() + ref.refine_adp() + assert torch.isfinite(torch.as_tensor(float(ref.get_rfactor()[0]))) + + ref.set_adp_representation("isotropic") + assert not ref.model.adp_is_field + # Leaving must put the per-atom weights back, or the next stage is misweighted. + assert ref.weighting()["adp"] == DEFAULT_GROUP_WEIGHTS["adp"] + ref.refine_adp() + after = float(ref.get_rfactor()[0]) + assert after == after and 0.0 < after < 0.7 + assert before == before + + +@pytest.mark.integration +def test_mode_set_on_an_isotropic_field_is_rejected(files): + ref = _refinement(files, adp_mode="isotropic") + with pytest.raises(ValueError, match="field_aniso"): + ref.set_adp_representation("field", mode_set="rigid") + + +# ---------------------------------------------------------------------------------- +# Checkpoints. +# ---------------------------------------------------------------------------------- + + +@pytest.mark.integration +def test_payload_identity_survives_the_state_dict(field_refinement): + """``state_dict`` holds tensors only, so the payload has to be encoded as one. + + Without the code, a restore infers the payload from the slot and rebuilds a + constant-U field: wrong storage width, wrong ADPs. + """ + sd = field_refinement.model.state_dict() + key = [k for k in sd if k.endswith("payload_code")] + assert key, f"no payload code in the state dict: {sorted(sd)[:8]}..." + assert int(sd[key[0]]) == payload_code(field_refinement.model.adp_field.payload) + + +@pytest.mark.integration +def test_model_state_dict_round_trip_rebuilds_the_same_field(field_refinement): + """A restore must rebuild the payload the code names, not one inferred from the slot.""" + from torchref.model.model import Model + + src = field_refinement.model + # create_from_state_dict restores the values itself; a further load_state_dict + # would be strict against restraint and metadata keys it never builds. + restored = Model.create_from_state_dict(src.state_dict(), verbose=0) + field = restored.adp_field + assert field is not None, "restored model has no field at all" + assert isinstance(field.payload, ModeCovariancePayload) + assert field.payload.mode_set == MODE_SET + assert field.node_shape == src.adp_field.node_shape + assert torch.allclose( + restored.adp_u6().detach(), src.adp_u6().detach(), atol=1e-5 + ) + + +@pytest.mark.integration +def test_model_copy_round_trip_on_a_bare_model(): + """``copy()`` carries the payload and the borrowed accessor, not a deep copy of them. + + Deliberately on a model that has never been handed to a Refinement --- see + :func:`test_copy_after_refinement_setup_is_broken_for_every_representation` for why. + """ + from torchref.model.model import Model + + import os + + here = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + src = Model(verbose=0) + src.load_pdb(os.path.join(here, "files", "pdb", "1DAW.pdb")) + src.set_adp_mode("field_aniso", n_nodes=8, k_neighbors=8, mode_set=MODE_SET) + clone = src.copy() + assert isinstance(clone.adp_field.payload, ModeCovariancePayload) + assert clone.adp_field.payload.mode_set == MODE_SET + assert torch.allclose(clone.adp_u6().detach(), src.adp_u6().detach()) + assert clone.adp_field.refinable_params is not src.adp_field.refinable_params + + +@pytest.mark.integration +@pytest.mark.xfail( + reason="PRE-EXISTING and representation-independent: once a Model has been through " + "Refinement setup, a cache somewhere holds a graph-attached tensor and deepcopy " + "refuses it. Measured identically for adp_mode isotropic, anisotropic and " + "field_aniso, and a bare model copies fine, so the node field is not the cause -- " + "it means no refinement of any kind can currently be checkpointed by copy().", + raises=RuntimeError, + strict=True, +) +@pytest.mark.integration +def test_copy_after_refinement_setup_is_broken_for_every_representation(field_refinement): + field_refinement.model.copy() + + +# ---------------------------------------------------------------------------------- +# The CLI, which is where a flag that does nothing hides best. +# ---------------------------------------------------------------------------------- + + +@pytest.mark.integration +def test_cli_flags_reach_the_refinement(): + """A null control on the plumbing: the parsed values must arrive, not the defaults. + + Five CLI flags in this codebase were once silently no-ops. The check is not that the + argument parses but that a non-default value changes what the refinement does. + """ + import argparse + + from torchref.cli._common import add_adp_mode_arg + + parser = argparse.ArgumentParser() + add_adp_mode_arg(parser) + args = parser.parse_args( + ["--adp-mode", "field_aniso", "--adp-mode-set", "affine", + "--reflections-per-adp-parameter", "3.5", "--adp-nodes", "17"] + ) + assert args.adp_mode == "field_aniso" + assert args.adp_mode_set == "affine" + assert args.reflections_per_adp_parameter == 3.5 + assert args.adp_nodes == 17 + + defaults = parser.parse_args([]) + assert defaults.adp_mode == "isotropic" + assert defaults.adp_mode_set is None + assert defaults.adp_nodes is None + + +@pytest.mark.integration +def test_cli_field_mode_end_to_end(files, tmp_path): + """The parsed flags, through the real constructor, produce a real field.""" + mtz, pdb = files + ref = LBFGSRefinement( + data_file=mtz, pdb=pdb, verbose=0, + adp_mode="field_aniso", adp_mode_set="affine", n_nodes=9, + ) + assert ref.model.adp_field.payload.mode_set == "affine" + assert ref.model.adp_field.n_nodes == 9 + ref.refine_adp() + out = tmp_path / "out.pdb" + ref.model.update_pdb() + ref.model.write_pdb(str(out)) + assert out.exists() and out.stat().st_size > 0 + text = out.read_text() + assert "ANISOU" in text, "an anisotropic field must write ANISOU records" diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index 79580086..e8d61ae8 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -4,6 +4,7 @@ from typing import Any, Dict, Optional +import math import torch from torch.nn import Module as nnModule @@ -65,6 +66,15 @@ "adp/node_smoothness": 0.0, } +#: Weight overrides a node-field ADP representation needs, applied by +#: :meth:`BaseRefinement.set_adp_representation`. +#: +#: Work reflections per ADP parameter that :meth:`set_adp_representation` targets when +#: sizing a field. PDB-REDO holds ~7 across its whole resolution range and switches model +#: form to stay there; measured on 179 of their entries, 7 is also where a node field +#: peaks, and both directions from it are worse. +DEFAULT_REFLECTIONS_PER_ADP_PARAMETER = 7.0 + class Refinement(DeviceMixin, DebugMixin, nnModule): """ @@ -119,6 +129,9 @@ def __init__( french_wilson: bool = True, anomalous: Optional[bool] = None, adp_mode: str = "isotropic", + adp_mode_set: str = None, + n_nodes: int = None, + reflections_per_adp_parameter: float = DEFAULT_REFLECTIONS_PER_ADP_PARAMETER, xray_mode: str = "ml", sigma_a_max: float = SIGMA_A_MAX, shrink: bool = SHRINK_ENABLED, @@ -167,6 +180,19 @@ def __init__( ADP parametrization: ``"isotropic"`` (default) refines a per-atom B-factor, ``"anisotropic"`` a 6-component U tensor for the atoms selected by ``aniso_selection`` (see :meth:`Model.set_adp_mode`). + ``"field"`` / ``"field_aniso"`` replace it with a node field, sized and + reweighted by :meth:`set_adp_representation`; ``"preserve"`` leaves the + file's own ADPs untouched. + adp_mode_set : str, optional + Displacement-mode set for ``adp_mode="field_aniso"`` --- ``"rigid"`` is TLS, + ``"rigid_dilation"`` adds uniform breathing. See + :data:`~torchref.model.disorder_field.MODE_SETS`. + n_nodes : int, optional + Explicit node count for a field mode. Default None sizes it from the data + through ``reflections_per_adp_parameter``. + reflections_per_adp_parameter : float, optional + Work reflections per ADP parameter a field is sized to hold. Default 7, + which is where a node field peaks and what PDB-REDO holds. xray_mode : str, optional X-ray target taxonomy row; see :meth:`set_xray_target_mode`. sigma_a_max, shrink : optional @@ -207,6 +233,11 @@ def __init__( # model right after load, before scaling/restraints/targets. self.adp_mode = adp_mode self.aniso_selection = aniso_selection + # Node-field settings. adp_mode_set names the displacement-mode set; n_nodes + # None means "size it from the data", which is what set_adp_representation does. + self.adp_mode_set = adp_mode_set + self.n_nodes = n_nodes + self.reflections_per_adp_parameter = reflections_per_adp_parameter # Everything the x-ray targets are built from must be set BEFORE # _init_targets() further down this __init__ (it also calls get_scales()). # They are read back through _xray_target_kwargs(), which is the single @@ -303,9 +334,17 @@ def __init__( ) self._sync_model_cell_to_data() - # Set ADP parametrization (iso/aniso) before scaling/restraints/targets - # so all structure-factor evaluation sees the chosen representation. - self.model.set_adp_mode(self.adp_mode, self.aniso_selection) + # Set the ADP parametrization before scaling/restraints/targets so all + # structure-factor evaluation sees the chosen representation. Routed through + # set_adp_representation rather than straight to the model: a field mode has + # to be sized from the reflection count and reweighted, and the model can do + # neither. Targets do not exist yet, so it will not try to rebuild them. + self.set_adp_representation( + self.adp_mode, + mode_set=self.adp_mode_set, + n_nodes=self.n_nodes, + reflections_per_parameter=self.reflections_per_adp_parameter, + ) self.setup_scaler() # Configure CIF path for lazy restraint building (restraints built on first access) self.model.set_restraints_cif(cif) @@ -425,6 +464,207 @@ def _build_xray_targets(self, mode: str) -> None: ) self.xray_mode = mode + # ------------------------------------------------------------------ + # ADP representation. + # ------------------------------------------------------------------ + + FIELD_MODES = ("field", "field_aniso") + + def _field_parameters_per_node(self, mode, mode_set, refine_node_positions): + """Storage columns one node costs: payload + log sigma + optional offset. + + Read off the payload rather than tabulated, so a new payload cannot silently + desynchronise the budget arithmetic from what the field actually allocates. + """ + from torchref.model.disorder_field import ( + AnisotropicPayload, + IsotropicPayload, + ModeCovariancePayload, + ) + + if mode_set is not None: + payload = ModeCovariancePayload(mode_set) + elif mode == "field_aniso": + payload = AnisotropicPayload() + else: + payload = IsotropicPayload() + return payload.width + 1 + (3 if refine_node_positions else 0) + + def nodes_for_reflection_budget( + self, + mode: str = "field_aniso", + mode_set: str = None, + reflections_per_parameter: float = DEFAULT_REFLECTIONS_PER_ADP_PARAMETER, + refine_node_positions: bool = True, + ) -> int: + """Node count giving ``reflections_per_parameter`` work reflections per ADP parameter. + + The reason this lives on the refinement and not on :class:`Model`: the model has + no idea how much data there is, and node count is set by the data rather than by + the structure. Measured on 179 PDB-REDO entries, node count correlates with + reflection count far more strongly than with atom count, and the model's own + default (one node per 25 atoms) is unrelated to either. + + The work set is the denominator because it is what the refinement fits, and it is + what PDB-REDO's ``NREFCNT`` counts, so the ratio is comparable to theirs. + + Returns + ------- + int + At least 2 --- a single node has no spatial structure to express. + """ + per_node = self._field_parameters_per_node( + mode, mode_set, refine_node_positions + ) + n_work = int(self.data.work.n) + budget = n_work / float(reflections_per_parameter) + return max(2, int(round(budget / per_node))) + + def flatten_adp_field(self) -> bool: + """Discard the field's spatial structure, keeping its level. Returns whether it ran. + + A node field fits its structure once, at the moment it is installed, and then only + refines from there. Early in a refinement that structure is derived against + coordinates that are still wrong, and nothing later re-derives it -- the same shape + of mistake as fitting bulk solvent to the starting model and never revisiting it, + which cost 11.5% error by cycle 4. Calling this between macro cycles throws away + the accumulated structure so the data rebuilds it against the coordinates as they + now are. + + The level is preserved: only the spatial variation is reset. Deliberately a hard + reset rather than a pull toward flat, because a soft version is another weight to + tune and the point is to test whether re-deriving helps at all. + + No-op when the model is not in field mode, so a driver can call it unconditionally. + """ + field = self.model.adp_field + if field is None: + return False + with torch.no_grad(): + per_atom = field().detach() + if per_atom.ndim == 2: + # A U6 field. Flatten through the equivalent isotropic B, NOT by taking a + # median over all six components: setting the off-diagonals to the same + # value as the diagonals gives a matrix with eigenvalues (3L, 0, 0), which + # is singular, and the Cholesky encode of it is NaN. refit lifts a 1-D B + # target to U_iso * I, which is the flat U that is actually meant. + b = (8.0 * math.pi**2 / 3.0) * per_atom[:, :3].sum(dim=1) + else: + b = per_atom + finite = torch.isfinite(b) + if not bool(finite.any()): + return False + level = b[finite].median() + target = torch.where(finite, level.expand_as(b), b) + # refit replaces refinable_params, so any cached leaf set or optimizer state + # referring to the old tensor is stale. + field.refit(target) + self.reset_loss_state() + if self.verbose > 0: + print(f"Flattened the ADP field to a level of {float(level):.2f}") + return True + + def set_adp_representation( + self, + mode: str, + mode_set: str = None, + n_nodes: int = None, + reflections_per_parameter: float = DEFAULT_REFLECTIONS_PER_ADP_PARAMETER, + k_neighbors: int = 12, + refine_node_positions: bool = True, + aniso_selection: str = None, + ): + """Switch the ADP parametrization, sizing and reweighting it for this data set. + + :meth:`Model.set_adp_mode` changes the representation but cannot size it: node + count follows from the reflection count, and the model has no idea how much data + there is. It also cannot swap the ADP restraint set, which is a property of the + representation rather than a weight to tune. + + The loss is **not** rebalanced for a field. The point of the representation is that + smoothness comes from the parametrisation, so a field should need *less* + regularisation than a per-atom model, not a reweighted version of the same + priors. :data:`DEFAULT_GROUP_WEIGHTS` already carries everything a field needs, + and an earlier attempt to raise the ``adp`` group for field mode had two side + effects worth remembering: ``adp/scaler_U`` and ``adp/scaler_log_scale`` sit under + that group, so it multiplied the scaler regularisation by the same factor, and it + made the field's configuration differ from every per-atom baseline in a way that + had nothing to do with ADPs. + + Safe to call after construction: the targets and scales are rebuilt afterwards, + which is what the model's own "run once at setup" caveat is about. + + Parameters + ---------- + mode : str + Any mode :meth:`Model.set_adp_mode` accepts. ``"field"`` and + ``"field_aniso"`` are sized and reweighted; the per-atom modes just pass + through, with any field weight overrides removed again. + mode_set : str, optional + Displacement-mode set for ``mode="field_aniso"``; see + :data:`~torchref.model.disorder_field.MODE_SETS`. + n_nodes : int, optional + Explicit node count, bypassing the reflection budget entirely. + reflections_per_parameter : float, optional + Target work reflections per ADP parameter when ``n_nodes`` is not given. + + Returns + ------- + dict + What was applied: mode, mode set, node count, parameter count and the + reflections-per-parameter actually achieved. Worth logging --- the achieved + ratio differs from the requested one by the integer rounding of node count. + """ + is_field = mode in self.FIELD_MODES + if mode_set is not None and mode != "field_aniso": + raise ValueError( + f"mode_set={mode_set!r} describes an anisotropic displacement field; " + 'use mode="field_aniso".' + ) + + if is_field and n_nodes is None: + n_nodes = self.nodes_for_reflection_budget( + mode, mode_set, reflections_per_parameter, refine_node_positions + ) + + self.model.set_adp_mode( + mode, + aniso_selection if aniso_selection is not None else self.aniso_selection, + n_nodes=n_nodes, + k_neighbors=min(k_neighbors, n_nodes) if is_field else k_neighbors, + refine_node_positions=refine_node_positions, + mode_set=mode_set, + ) + self.adp_mode = mode + self.adp_mode_set = mode_set + + # Targets hold per-atom index tensors keyed off the old parametrization, the + # scales were fitted against the old F_calc, and which ADP restraints even apply + # is a property of the representation -- so the component set changes, not just + # the weights. reset_loss_state is what makes the next access register the new + # set; it also drops the Logger, which holds a reference to the old state and + # would otherwise keep recording into it. + if getattr(self, "adp_target", None) is not None: + self._init_targets() + self.reset_loss_state() + + n_par = sum(p.numel() for p in self.model.parameters_of_types(("adp", "u"))) + applied = dict( + mode=mode, mode_set=mode_set, n_nodes=n_nodes, n_adp_parameters=int(n_par), + reflections_per_parameter=( + float(self.data.work.n) / n_par if n_par else float("inf") + ), + ) + if self.verbose > 0: + label = mode if mode_set is None else f"{mode}/{mode_set}" + print( + f"ADP representation: {label}" + + (f", {n_nodes} nodes" if is_field else "") + + f", {n_par} parameters, " + f"{applied['reflections_per_parameter']:.1f} work reflections each" + ) + return applied + def _init_targets(self, xray_mode: str = None): """Build the x-ray, geometry and ADP targets and initialise the scales. From a6621d125426e6d8c86b1c173c3b461df2592a05 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 11:01:58 +0200 Subject: [PATCH 126/250] Expose the node-field ADP modes on the refine CLI --adp-mode gains field, field_aniso and preserve, plus --adp-mode-set, --reflections-per-adp-parameter and --adp-nodes, wired through to the constructor and echoed in the run summary and saved settings. The representation was library-only before, so none of it was reachable by a user. Five flags in this codebase were once silently no-ops, so the test asserts that non-default values arrive rather than merely that they parse. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- torchref/cli/_common.py | 35 +++++++++++++++++++++++++++++++++-- torchref/cli/refine.py | 15 +++++++++++++++ 2 files changed, 48 insertions(+), 2 deletions(-) diff --git a/torchref/cli/_common.py b/torchref/cli/_common.py index 67a90602..a9052312 100644 --- a/torchref/cli/_common.py +++ b/torchref/cli/_common.py @@ -144,12 +144,43 @@ def add_adp_mode_arg(parser: argparse.ArgumentParser) -> None: "--adp-mode", type=str, default="isotropic", - choices=["isotropic", "anisotropic"], + choices=["isotropic", "anisotropic", "field", "field_aniso", "preserve"], help="ADP parametrization: 'isotropic' (default) refines a per-atom " "B-factor; 'anisotropic' refines a 6-component U tensor for the atoms " "given by --anisotropic-selection. The model is converted between " "representations and the output PDB/mmCIF follows the convention " - "(ANISOU only for anisotropic atoms).", + "(ANISOU only for anisotropic atoms). 'field' and 'field_aniso' replace " + "the per-atom parameters with a node field, whose size is set from the " + "data rather than the atom count (--reflections-per-adp-parameter). " + "'preserve' leaves the input file's own ADPs untouched.", + ) + parser.add_argument( + "--adp-mode-set", + type=str, + default=None, + choices=["constant", "rigid", "rigid_dilation", "affine"], + help="Displacement-mode set for --adp-mode field_aniso. Each node carries " + "the covariance of these modes, so its ADP varies across the region it " + "serves: 'constant' is one U per node, 'rigid' is TLS, 'rigid_dilation' " + "adds uniform breathing, 'affine' adds shear and extension.", + ) + parser.add_argument( + "--reflections-per-adp-parameter", + type=float, + default=7.0, + metavar="R", + help="Work reflections per ADP parameter a node field is sized to hold " + "(--adp-mode field/field_aniso). Default 7. Node count follows from the " + "data rather than the atom count, and both directions from 7 measured " + "worse. Ignored by the per-atom modes.", + ) + parser.add_argument( + "--adp-nodes", + type=int, + default=None, + metavar="N", + help="Explicit node count for a field ADP mode, bypassing " + "--reflections-per-adp-parameter.", ) parser.add_argument( "--anisotropic-selection", diff --git a/torchref/cli/refine.py b/torchref/cli/refine.py index 238db84c..c8c837ad 100644 --- a/torchref/cli/refine.py +++ b/torchref/cli/refine.py @@ -254,6 +254,15 @@ def main(): if args.dmin: print(f"Resolution cutoff: {args.dmin:.2f} A") adp_line = f"ADP mode: {args.adp_mode}" + if args.adp_mode == "field_aniso" and args.adp_mode_set: + adp_line += f" ({args.adp_mode_set})" + if args.adp_mode in ("field", "field_aniso"): + adp_line += ( + f", {args.adp_nodes} nodes" + if args.adp_nodes + else f", sized at {args.reflections_per_adp_parameter:g} " + "work reflections per parameter" + ) if args.adp_mode == "anisotropic": adp_line += ( " (selection: " @@ -290,6 +299,9 @@ def main(): scale_target=args.scale_target, **_sigma_a_kwargs(args), adp_mode=args.adp_mode, + adp_mode_set=args.adp_mode_set, + n_nodes=args.adp_nodes, + reflections_per_adp_parameter=args.reflections_per_adp_parameter, aniso_selection=args.anisotropic_selection, wavelength=args.wavelength, ) @@ -381,6 +393,9 @@ def main(): "n_cycles": args.n_cycles, "mode": args.mode, "adp_mode": args.adp_mode, + "adp_mode_set": args.adp_mode_set, + "adp_nodes": args.adp_nodes, + "reflections_per_adp_parameter": args.reflections_per_adp_parameter, "anisotropic_selection": ( args.anisotropic_selection if args.adp_mode == "anisotropic" else None ), From bfc18d72cb6004256de068209f67db602e7af0f7 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 11:03:04 +0200 Subject: [PATCH 127/250] Document field_aniso and preserve together in set_adp_mode Takes the fuller wording from dev's working tree, which describes field_aniso as well. The preserve implementation there is identical to the committed one; only the docstring differed. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01TN8kHX6f9MwFxN8nGUwiQU --- torchref/model/model.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/torchref/model/model.py b/torchref/model/model.py index a4bb4888..6ece8612 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -1276,9 +1276,10 @@ def set_adp_mode( to ``U = (B / 8 pi^2) I``. ``"field"`` replaces the per-atom isotropic B with a :class:`~torchref.model.disorder_field.DisorderFieldTensor`, whose node values are least-squares fitted to the B it replaces, so the atom - count stops setting the ADP parameter count. - ``"preserve"`` is a no-op, leaving the ADPs exactly as the file supplied - them: use it when the starting model's own ADPs are what is being measured. + count stops setting the ADP parameter count. ``"field_aniso"`` is the same + representation carrying a full U per node, which takes over ``u`` rather + than ``adp``. ``"preserve"`` is a no-op: the ADPs stay exactly as the file + supplied them, anisotropic where the file was anisotropic. aniso_selection : str, optional Phenix-style selection for ``mode="anisotropic"``, default ``"not resname HOH and not element H"``; ignored otherwise. From 9c7b97deb7e1c30ef68cbf40b04989596f2d5f64 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 11:33:18 +0200 Subject: [PATCH 128/250] Add a preserve ADP mode that leaves the loaded ADPs alone set_adp_mode always reparametrised, so constructing a Refinement over a deposited model isotropised its ANISOU before anything else ran, and a run meant to measure that model's own ADPs measured a converted copy instead. "preserve" returns immediately, leaving the per-atom wrappers exactly as the reader built them. The docstring also picks up "field_aniso", which the mode list had gained without a description. Co-Authored-By: Claude Opus 5 (1M context) --- torchref/model/model.py | 15 ++++++++++++--- 1 file changed, 12 insertions(+), 3 deletions(-) diff --git a/torchref/model/model.py b/torchref/model/model.py index 3def94be..f033593a 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -1267,14 +1267,17 @@ def set_adp_mode( Parameters ---------- - mode : {"isotropic", "anisotropic", "field", "field_aniso"}, optional + mode : {"isotropic", "anisotropic", "field", "field_aniso", "preserve"}, optional ``"isotropic"`` (default) converts every atom, previously anisotropic ones to ``B_eq = (8 pi^2 / 3)(U11 + U22 + U33)``. ``"anisotropic"`` converts those matching ``aniso_selection``, expanding isotropic atoms to ``U = (B / 8 pi^2) I``. ``"field"`` replaces the per-atom isotropic B with a :class:`~torchref.model.disorder_field.DisorderFieldTensor`, whose node values are least-squares fitted to the B it replaces, so the atom - count stops setting the ADP parameter count. + count stops setting the ADP parameter count. ``"field_aniso"`` is the same + representation carrying a full U per node, which takes over ``u`` rather + than ``adp``. ``"preserve"`` is a no-op: the ADPs stay exactly as the file + supplied them, anisotropic where the file was anisotropic. aniso_selection : str, optional Phenix-style selection for ``mode="anisotropic"``, default ``"not resname HOH and not element H"``; ignored otherwise. @@ -1298,6 +1301,12 @@ def set_adp_mode( """ if not self.ctx.initialized or self.pdb is None: return + if mode == "preserve": + # Leave the ADPs exactly as loaded. Constructing a Refinement otherwise + # reparametrises them before anything else runs, which silently discards a + # deposited model's anisotropy -- use this when the starting model's own + # ADPs are the thing being measured. + return if mode in ("field", "field_aniso"): aniso = mode == "field_aniso" # Run the partition first either way: it owns every buffer keyed off the @@ -1348,7 +1357,7 @@ def set_adp_mode( else: raise ValueError( f"Unknown ADP mode: {mode!r}. Use 'isotropic', 'anisotropic', " - "'field' or 'field_aniso'." + "'field', 'field_aniso' or 'preserve'." ) self._apply_adp_partition(aniso_mask) From 957fb236485b823c5a2cd35c4e9a282aca15a22b Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 11:33:25 +0200 Subject: [PATCH 129/250] Rewrite the scaling guide for the Chebyshev scale and the solvent falloff The page still described per-bin scaling and a plain Debye-Waller solvent term, both of which the scaler rework replaced. The overall scale is now a Chebyshev polynomial in s with n_iso_coeff coefficients, and resolution bins survive only as the seed for that fit; the solvent term is a generalised falloff parametrised by s^2_half and n, which reduces to exp(-B_s s^2) at n = 1 so a solvent B from another program still transfers. Co-Authored-By: Claude Opus 5 (1M context) --- docs/user_guide/scaling.rst | 80 +++++++++++++++++++++++++------------ 1 file changed, 54 insertions(+), 26 deletions(-) diff --git a/docs/user_guide/scaling.rst b/docs/user_guide/scaling.rst index 14fe6d43..d1e1c9da 100644 --- a/docs/user_guide/scaling.rst +++ b/docs/user_guide/scaling.rst @@ -2,8 +2,8 @@ Scaling ======= :class:`~torchref.scaling.scaler.Scaler` puts F_calc on the observed scale and -absorbs what the atomic model does not describe: an overall (per-resolution-bin) -scale, an anisotropic correction, and the bulk solvent contribution. +absorbs what the atomic model does not describe: an overall isotropic scale, an +anisotropic correction, and the bulk solvent contribution. Basic Usage ----------- @@ -14,7 +14,7 @@ Basic Usage scaler = Scaler(model, reflection_data, verbose=1) - scaler.initialize() # initial bin scales + solvent + anisotropy + scaler.initialize() # initial scale + solvent + anisotropy scaler.refine_lbfgs() # refine the scaling parameters F_calc_scaled = scaler(F_calc) @@ -23,42 +23,70 @@ Basic Usage (``calc_initial_scale`` → ``setup_solvent`` → ``setup_anisotropy_correction``). Each factor defaults to 1 while its parameter is absent, so a freshly constructed scaler is the identity, and one on which only -``calc_initial_scale()`` has run applies the overall bin scale alone. +``calc_initial_scale()`` has run applies the overall isotropic scale alone. -Bin-wise Scaling ----------------- +Isotropic Scaling +----------------- -Reflections are binned by resolution shell and an overall scale is fitted per -bin. On by default with 20 bins: +The overall scale is a Chebyshev polynomial in :math:`s = \sin\theta/\lambda`, +evaluated per reflection: + +.. math:: + + k_{iso}(s) = \exp\left( \sum_{i} c_i\, T_i(u) \right), + \qquad u \in [-1, 1] + +with ``n_iso_coeff`` coefficients (default 6) held in ``scaler.c_iso``. Every +reflection contributes to every coefficient with a continuous weight, so there are +no bin boundaries and nothing changes discontinuously when a reflection moves +between shells. ``n_iso_coeff=1`` is a single global scale +(:math:`T_0 \equiv 1`); ``2`` spans scale-plus-overall-B. .. code-block:: python - scaler = Scaler(model, reflection_data, verbose=1, nbins=1) # single global scale + scaler = Scaler(model, reflection_data, n_iso_coeff=1) # single global scale scaler.calc_initial_scale() +Resolution bins survive only as the device that *seeds* the coefficients: the +closed-form per-bin :math:`|F_{obs}|/|F_{calc}|` ratio is projected onto the basis +by least squares, so the fit starts from the curve a binned model would have +started from. Nothing downstream is binned; ``nbins`` controls only that seed. + +Use ``scaler.iso_log_scale()`` for the per-reflection log scale, and +``scaler.get_scale()`` for a single summary number. + Bulk Solvent Model ------------------ -The solvent contribution is mask-derived, not analytic: a solvent mask is built -from the model, smoothed, and Fourier transformed to give :math:`F_{solvent}`, -which is then Debye-Waller damped and scaled. +The solvent contribution is mask-derived, not analytic: a binary solvent mask is +built from the model and Fourier transformed to give :math:`F_{solvent}`, which is +then damped and scaled. .. math:: - F_{calc}^{total} = k \cdot F_{calc}^{model} - + k_s \exp(-B_s s^2) \cdot F_{calc}^{solvent} - -where :math:`k` is the overall (per-bin) scale, :math:`k_s` the solvent scale, -:math:`B_s` the solvent B-factor, and :math:`s = \sin\theta/\lambda` — the -*half*-length of the scattering vector (``ScalerBase._s_half_sq``), which is why -the exponent carries no factor of 4. This is the ordinary Debye-Waller -convention written the other way round: :math:`\exp(-B_s/4d^2)`, i.e. -:math:`\exp(-B_s s^2/4)` for :math:`s = 1/d`. A :math:`B_s` from another program -therefore transfers unchanged (the default is 46 Ų). - -The refined parameters are ``log_k_solvent``, ``b_solvent`` and -``phase_offset``; the last blends the mask phases toward the protein phases and -is only active when ``optimize_phase`` is set. + F_{calc}^{total} = k_{iso}(s)\, k_{aniso}(\mathbf{h}) \cdot F_{calc}^{model} + + k_s \exp\!\left( -\ln 2 \left(\frac{s^2}{s^2_{1/2}}\right)^{\!n} \right) + \cdot F_{calc}^{solvent} + +where :math:`k_s` is the solvent scale, :math:`s^2_{1/2}` the point at which the +solvent term is halved, :math:`n` how sharply it switches off, and +:math:`s = \sin\theta/\lambda` — the *half*-length of the scattering vector +(``ScalerBase._s_half_sq``), which is why the exponent carries no factor of 4. + +:math:`n = 1` reduces this exactly to a Debye-Waller factor +:math:`\exp(-B_s s^2)` with :math:`B_s = \ln 2 / s^2_{1/2}`, so a solvent B from +another program transfers unchanged. Larger :math:`n` gives a plateau followed by a +sharper cutoff, which is the shape a flat bulk-solvent prior actually has: it +describes the data well at low resolution and then stops being informative. + +The refined parameters are ``log_k_solvent``, ``log_ss_half``, ``log_n_exp`` and +``phase_offset``; each falloff parameter is refined in log space so it stays +positive, and is clamped to ``SS_HALF_BOUNDS`` / ``N_EXP_BOUNDS``. The phase offset +blends the mask phases toward the protein phases and is only active when +``optimize_phase`` is set. + +PDB ``REMARK 3`` and mmCIF carry a single solvent B, which this form does not have; +``SolventModel.b_solvent_equivalent`` back-fits one from the curve for deposition. Anisotropic Scaling ------------------- From 3f80ff7e183eea8dfb137d9005da4c14219125d2 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 11:33:26 +0200 Subject: [PATCH 130/250] Ignore .DS_Store and the scaler_investigation outputs The lab's scripts stay tracked; its cache, metrics, figures and slurm logs do not. Co-Authored-By: Claude Opus 5 (1M context) --- .gitignore | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/.gitignore b/.gitignore index 9d302c4c..ef7013dc 100644 --- a/.gitignore +++ b/.gitignore @@ -18,6 +18,8 @@ _temp.mtz __pycache__/ *.pyc *.pyo +# macOS finder metadata +.DS_Store # Large binary files *.ccp4 *.png @@ -84,4 +86,10 @@ graphify-out/ # Large binary scratch dirs & squashfs images *.sqsh anisotropic/ -torchref_refine_optimization/runs/ \ No newline at end of file +torchref_refine_optimization/runs/ + +# scaler_investigation lab outputs: the scripts are tracked, the outputs are not. +scaler_investigation/cache/ +scaler_investigation/metrics/ +scaler_investigation/figures/ +scaler_investigation/slurm/ \ No newline at end of file From 6413c05a7553a42377e9ddb5083f49c2142debea Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 11:54:22 +0200 Subject: [PATCH 131/250] Report the translation correlation per candidate; keep R as the ranking key Ranking candidates by the translation-function correlation instead of the analytical-scale R was tried and is WORSE end to end: 31/40 against 36/40 over four structures x ten seeds, losing on 2DQ6 (6/10 -> 3/10) and on 6G9X (10/10 -> 8/10). Reverted; the correlation is now computed and reported on every candidate but does not select. A rank-level harness had predicted the opposite by a wide margin -- truth first 33/40 by correlation against 23/40 by R -- and the reason it was wrong is a broken truth label, not a subtle effect. `angle_to_orbit` compares a candidate against `S_k @ R_true`, and on 2DQ6 that marks 4 of 25 candidates correct where Kabsch superposition of the actual coordinates marks 14. The `side="right"` orbit is no better in the other direction: it marks all 25 correct. Neither agrees with the coordinates, which are ground truth by construction. truth_metric_check.py is that comparison, so the disagreement is checkable rather than argued. Every rank-level number from a harness built on `angle_to_orbit` needs re-measuring against coordinates before it is used to justify anything -- including the recorded claim that the translation function out-discriminates the rotation function 27/30 to 6/30. correlation_at() evaluates the search's own functional at a single translation, which is what lets the reported score belong to the refined position that was actually used rather than to the grid peak it started from. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/truth_metric_check.sh | 17 +++++ .../diagnostics/llg_shared_sigma_a.py | 70 ++++++++++++++++-- alignment_lab/diagnostics/pose_recovery.py | 11 +-- .../diagnostics/truth_metric_check.py | 71 +++++++++++++++++++ torchref/experimental/alignment/pipeline.py | 40 +++++++---- .../experimental/alignment/translation.py | 32 +++++++++ 6 files changed, 219 insertions(+), 22 deletions(-) create mode 100644 alignment_lab/analysis/truth_metric_check.sh create mode 100644 alignment_lab/diagnostics/truth_metric_check.py diff --git a/alignment_lab/analysis/truth_metric_check.sh b/alignment_lab/analysis/truth_metric_check.sh new file mode 100644 index 00000000..892b7e1c --- /dev/null +++ b/alignment_lab/analysis/truth_metric_check.sh @@ -0,0 +1,17 @@ +#!/bin/bash +#SBATCH --job-name=truthchk +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +for P in 2DQ6 1DAW; do "$PY" -u alignment_lab/diagnostics/truth_metric_check.py $P 0 2>/dev/null; done +echo DONE diff --git a/alignment_lab/diagnostics/llg_shared_sigma_a.py b/alignment_lab/diagnostics/llg_shared_sigma_a.py index 6bfe9d90..e7fe7d8a 100644 --- a/alignment_lab/diagnostics/llg_shared_sigma_a.py +++ b/alignment_lab/diagnostics/llg_shared_sigma_a.py @@ -25,6 +25,32 @@ is rotation-invariant by construction -- total scattering per shell does not depend on orientation -- and is the estimate that exists for precisely this reason. +``luzzati`` ASSUMED, not fitted: exp(-(2 pi^2/3) s^2 dVRMS^2) from the search + model's expected coordinate error, which is the same Eterm the + rotation function already weights with. It never looks at the + data, so it has zero free parameters of any kind -- which is the + property that makes the plain correlation robust, applied to a + likelihood instead. + +Also reported is ``r``, the analytical-scale R the pipeline actually selects on, +because the point of the exercise is replacing it: on 2DQ6 all 25 candidates +fall within 0.023 R of each other and the wrong winner leads a correct candidate +by 0.0001. + +Two forms of the correlation are compared, because the one the search maximises +is not a correlation coefficient: + +``corr`` what ``amplitude_translation_search`` returns, + ``sum w (E_obs^2 - mean) |Fc|^2 / sum w |Fc|^2``. The denominator + normalises the weighted MEAN, not the spread, so this is a weighted + mean of centred observed intensity. Fine for finding the peak in t + at fixed orientation; its scale across ORIENTATIONS depends on how + concentrated that candidate's |Fc|^2 happens to be. +``pearson`` the actual coefficient, ``cov / sqrt(var var)``, weighted the same + way. Per-candidate normalisation, so unlike a global rescale it can + and does reorder. Computed directly at each candidate's chosen + translation -- the FFT gives the numerator and one denominator but + not ``sum w |Fc|^4``, and one evaluation per candidate is cheap. """ from __future__ import annotations @@ -32,7 +58,6 @@ import sys from pathlib import Path -import numpy as np import torch sys.path.insert(0, str(Path(__file__).resolve().parents[1])) @@ -56,7 +81,7 @@ def main() -> int: ap = argparse.ArgumentParser(description=__doc__) ap.add_argument("--pdb", default="2DQ6", choices=list(BENCH_PDBS)) ap.add_argument("--trials", type=int, default=10) - ap.add_argument("--n-cand", type=int, default=15) + ap.add_argument("--n-cand", type=int, default=25) ap.add_argument("--thr-deg", type=float, default=8.0) args = ap.parse_args() @@ -69,8 +94,10 @@ def main() -> int: from torchref.experimental.alignment.rotation_search import prepare_frf_inputs from torchref.experimental.alignment.translation import ( DirectModelEvaluator, amplitude_translation_search, fit_sigma_a_per_shell, - llg_translation_rescore, normalise_calc, precompute_G_for_rotation, + llg_translation_rescore, local_translation_refine, normalise_calc, + precompute_G_for_rotation, ) + from torchref.experimental.alignment.frf.preprocessing import eterm_sigma_a from torchref.scaling import WilsonNormaliser from torchref.scaling.weighting import empirical_sigma_a @@ -118,9 +145,27 @@ def main() -> int: ph = torch.exp(2j * torch.pi * torch.einsum( "ind,d->in", h_R.to(torch.float64), t_top).to(G.dtype)) Fc_top = (G * ph).sum(dim=0).abs().to(torch.float64) + _, r_a = local_translation_refine( + obs=obs, interpolator=ev, R_rotation=eye3, + spacegroup=data.spacegroup, real_cell=data.cell, + t_init=t_top.cpu(), radius=0.06, grid_steps=13, + n_refinement_passes=1, + precomputed_G=G, precomputed_h_R=h_R) + # Weighted Pearson r between observed and calculated intensity at + # this candidate's chosen translation. + w = obs.weight.to(torch.float64) + x = (obs.E_obs.to(torch.float64)) ** 2 + y = (Fc_top / Fc_top.mean().clamp(min=1e-30)) ** 2 + wsum = w.sum().clamp(min=1e-30) + xm, ym = (w * x).sum() / wsum, (w * y).sum() / wsum + dx, dy = x - xm, y - ym + cov = (w * dx * dy).sum() / wsum + vx = (w * dx * dx).sum() / wsum + vy = (w * dy * dy).sum() / wsum + pear = float(cov / (vx * vy).clamp(min=1e-30).sqrt()) cand.append(dict(ang=float(ang), corr=float(tp[0].score), G=G, - h_R=h_R, t=t_top, Fc=Fc_top, - E_calc=normalise_calc(Fc_top, obs))) + h_R=h_R, t=t_top, Fc=Fc_top, r=float(r_a), + pearson=pear, E_calc=normalise_calc(Fc_top, obs))) is_truth = [c["ang"] <= args.thr_deg for c in cand] if not any(is_truth): @@ -148,16 +193,29 @@ def llg_of(c, sigma_a): obs=obs, G=c["G"], h_R=c["h_R"], t_candidates=c["t"].view(1, 3), sigma_a=sigma_a)[0]) + # Assumed sigma_A: the Luzzati falloff at the shell centres, from the + # model error the pipeline already estimates. No data, no fit. + s_shell = torch.zeros(obs.n_shells, dtype=torch.float64) + cnt_s = torch.bincount(obs.shell_idx, minlength=obs.n_shells).to(torch.float64) + s_shell.scatter_add_(0, obs.shell_idx, obs.s_mag.to(torch.float64)) + s_shell = s_shell / cnt_s.clamp(min=1.0) + sa_luz = eterm_sigma_a(s_shell, float(pipe.model_error_A)).clamp(1e-3, 1 - 1e-6) + variants = { "per_cand": [llg_of(c, sa_per_cand(c)) for c in cand], "shared": [llg_of(c, sa_shared) for c in cand], "empirical": [llg_of(c, sa_emp) for c in cand], + "luzzati": [llg_of(c, sa_luz) for c in cand], } r_corr = _rank_of_truth([c["corr"] for c in cand], is_truth) + r_pear = _rank_of_truth([c["pearson"] for c in cand], is_truth) + r_R = _rank_of_truth([c["r"] for c in cand], is_truth, higher_is_better=False) parts = " ".join( f"rank_{k}={_rank_of_truth(v, is_truth)}" for k, v in variants.items()) print(f"ROW pdb={args.pdb} trial={trial} n_truth={sum(is_truth)} " - f"rank_corr={r_corr} {parts}", flush=True) + f"n_cand={len(cand)} vrms={float(pipe.model_error_A):.2f} " + f"rank_r={r_R} rank_corr={r_corr} rank_pearson={r_pear} {parts}", + flush=True) return 0 diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py index e2529935..07f8fe67 100644 --- a/alignment_lab/diagnostics/pose_recovery.py +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -96,8 +96,10 @@ def _report_candidates(solutions, R_true, symops, success_deg) -> None: R in 0 of 10 seeds while the pipeline solved 6 of them, because it fed the R-factor a different set of translation peaks. - ``SOLN`` lines are ordered as the pipeline ranked them, so line 0 is what it - returned. ``dtruth`` is the angle from that candidate's orientation to the + ``SOLN`` lines are ordered as the pipeline ranked them -- by ``tf_corr``, + descending -- so line 0 is what it returned. ``R`` is carried alongside + because it used to be the ranking key and comparing the two orderings is the + point. ``dtruth`` is the angle from that candidate's orientation to the true one modulo crystal symmetry; ``pick`` marks the winner and ``true`` marks every candidate that was in fact correct. """ @@ -106,7 +108,7 @@ def _report_candidates(solutions, R_true, symops, success_deg) -> None: ) R_t = R_true.to(torch.float64).cpu() - print(" SOLN rank rot_score tf_R dtruth flags") + print(" SOLN rank rot_score tf_corr R dtruth flags") for i, sol in enumerate(solutions): R = torch.as_tensor(sol.rotation, dtype=torch.float64) # `rotation` maps the search-model frame onto the crystal frame; the @@ -116,7 +118,8 @@ def _report_candidates(solutions, R_true, symops, success_deg) -> None: for k in range(symops.shape[0])) flags = ("pick " if i == 0 else " ") + ("true" if d <= success_deg else "") print(f" SOLN {i:4d} {sol.rotation_score:10.3f} " - f"{sol.translation_score:7.4f} {d:8.2f} {flags}") + f"{sol.translation_score:10.5f} {sol.r_factor:7.4f} " + f"{d:8.2f} {flags}") def main() -> int: diff --git a/alignment_lab/diagnostics/truth_metric_check.py b/alignment_lab/diagnostics/truth_metric_check.py new file mode 100644 index 00000000..dc2a8f34 --- /dev/null +++ b/alignment_lab/diagnostics/truth_metric_check.py @@ -0,0 +1,71 @@ +"""Do the two truth metrics in this lab agree about which candidate is correct? + +Two answers to "is this candidate the right orientation" are in use: + +``angle_to_orbit`` used by the rank harnesses. Compares a candidate's + ``R_recovered`` against ``S_k @ R_true``. +``residual_rotation_deg`` used by pose_recovery for the pass/fail. Kabsch- + superposes the placed coordinates onto canonical and + takes the smallest angle to any symop. + +``RotationPeak`` rotations are ``R_recovered``, which maps the SEARCH-MODEL frame +onto the crystal frame -- the rotation applied to the coordinates is its +transpose. If the orbit comparison omits that transpose it is comparing a +rotation with its own inverse's orbit, which is a different set unless the +rotation is an involution. + +That matters beyond bookkeeping: the rank harness said ranking by the +translation correlation would beat the analytic R by 33/40 to 23/40, and end to +end it lost 31/40 to 36/40. A truth label that is wrong makes every rank in that +harness meaningless, so this checks it directly rather than by inference. +""" +import sys +from pathlib import Path +import torch +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) +from lab import rotated_case, seed_for, symmetry_orbit # noqa: E402 +from lab.truth import angle_to_orbit # noqa: E402 + +pdb = sys.argv[1] if len(sys.argv) > 1 else "2DQ6" +trial = int(sys.argv[2]) if len(sys.argv) > 2 else 0 +n_cand = 25 + +from torchref.experimental.alignment.frf.rotation_utils import ( # noqa: E402 + rotation_angular_distance_deg, rotation_matrix_from_edmonds_euler) +from torchref.experimental.alignment.pipeline import ( # noqa: E402 + MolecularReplacementPipeline) +from torchref.experimental.alignment.rotation_search import ( # noqa: E402 + prepare_frf_inputs) + +seed = seed_for(pdb, trial) +model, data, R_true = rotated_case(pdb, seed) +pipe = MolecularReplacementPipeline(data, model, verbose=0, n_rotation_peaks=200, + n_rotation_candidates=n_cand) +frf = prepare_frf_inputs(model, data, d_min=pipe.d_min, d_max=pipe.d_max, + n_shells=pipe.n_shells, verbose=0) +pipe._frf = frf +peaks = pipe._rotation_candidates(frf)[:n_cand] + +symops = data.spacegroup.matrices.to(torch.float64).cpu() +rb = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() +orbit_l = symmetry_orbit(R_true, symops, side="left", frame="cart", + reciprocal_basis=rb) +orbit_r = symmetry_orbit(R_true, symops, side="right", frame="cart", + reciprocal_basis=rb) +R_t = R_true.to(torch.float64).cpu() + +print(f"# {pdb} trial={trial} n_cand={len(peaks)}") +print(f"{'k':>3s} {'orbit side=left':>16s} {'orbit side=right':>17s} " + f"{'coords (Kabsch form)':>21s}") +n_l = n_r = n_c = 0 +for k, p in enumerate(peaks): + R = rotation_matrix_from_edmonds_euler(p.alpha, p.beta, p.gamma).to(torch.float64) + a = angle_to_orbit(R, orbit_l) + b = angle_to_orbit(R, orbit_r) + c = min(float(rotation_angular_distance_deg(R.T @ R_t, symops[i])) + for i in range(symops.shape[0])) + n_l += a <= 8.0; n_r += b <= 8.0; n_c += c <= 8.0 + print(f"{k:3d} {a:16.2f} {b:17.2f} {c:21.2f}") +print(f"within 8 deg: side=left {n_l}/{len(peaks)}, side=right {n_r}/{len(peaks)}, " + f"coords {n_c}/{len(peaks)}") diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index da186b28..5a41d567 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -60,6 +60,7 @@ TranslationObs, TranslationPeak, amplitude_translation_search, + correlation_at, fit_sigma_a_per_shell, llg_translation_rescore, local_translation_refine, @@ -258,12 +259,14 @@ class MRSolution: Fractional translation applied after rotation, shape (3,). ``None`` for a rotation-only solution (``do_translation=False``). rotation_score : float - ML-LLG score of the rotation candidate (from the rescore). + The rotation function's score for this candidate. translation_score : float - Analytical-R of the best translation for this candidate (lower better). + The translation function's correlation at the chosen translation, higher + better. Reported, not ranked -- see the sort in :meth:`run`. r_factor : float - Ranking key: the translation search's analytical-scale R. For the - returned winner it is replaced by the solvent-aware Scaler R-work. + **The ranking key**: the analytical-scale R at that placement, lower + better. For the returned winner it is replaced by the solvent-aware + Scaler R-work. model : ModelFT The rotated (+translated +refined) model for this candidate. """ @@ -499,8 +502,6 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: placed = rotated_k.copy().translate( t_refined.to(self.model.dtype_float), fractional=True, ) - r_rank = r_analytic - placed.last_alignment_rotation = R_rec_k placed.last_alignment_translation = t_refined solutions.append( @@ -508,8 +509,8 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: rotation=R_rec_k.detach().cpu().numpy(), translation=t_refined.detach().cpu().numpy(), rotation_score=float(peak_k.score), - translation_score=float(r_analytic), - r_factor=float(r_rank), + translation_score=float(tf_score), + r_factor=float(r_analytic), model=placed, ) ) @@ -517,6 +518,13 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: if not solutions: raise RuntimeError("Translation + joint refine produced no candidates.") + # Lowest analytical-scale R. Ranking by the translation correlation was + # tried and is worse end to end -- 31/40 against 36/40 over four + # structures x ten seeds, losing on both 2DQ6 (6/10 -> 3/10) and 6G9X + # (10/10 -> 8/10). A rank-level harness had predicted the opposite by a + # wide margin, which is why the correlation is still reported on every + # candidate at verbose >= 2: the two orderings disagree and the + # disagreement is not yet understood. solutions.sort(key=lambda s: s.r_factor) winner = solutions[0] @@ -526,7 +534,8 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: timer.stop("12_final_scaler") winner.model.last_alignment_rfactor = rwork_final winner.r_factor = rwork_final - self._log(1, f"mr: winner analytical-TF R={winner.translation_score:.4f}, " + self._log(1, f"mr: winner TF correlation={winner.translation_score:.5f}, " + f"analytic R={winner.r_factor:.4f}, " f"final Scaler-fit R-work={rwork_final:.4f}") self._log(2, "\n" + timer.summary()) return solutions @@ -692,11 +701,18 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: precomputed_G=G_pre, precomputed_h_R=h_R_pre, ) timer.stop("7_local_TF_refine") - self._log(3, f" trans{k_t}: R(analytic)={r_analytic:.4f}, " + # Both scores at the REFINED position, so the reported correlation + # belongs to the translation that was actually chosen. Selection is + # by R: ranking candidates by the correlation instead was measured + # end to end and is WORSE (31/40 against 36/40 over four structures + # x ten seeds), despite a rank-level harness predicting the reverse. + tf_ref = correlation_at(self._obs, G_pre, h_R_pre, t_refined) + self._log(3, f" trans{k_t}: tf={tf_ref:.5f} " + f"R(analytic)={r_analytic:.4f}, " f"t={[round(float(x), 3) for x in t_refined.tolist()]}") if best is None or r_analytic < best[0]: - best = (r_analytic, t_refined) - return None if best is None else (best[0], best[1], tf_top) + best = (r_analytic, t_refined, tf_ref) + return best def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): """Re-rank translation peaks by a shared-σA Rice/Woolfson LLG. diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index ac19d386..c003ad7b 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -616,6 +616,38 @@ def llg_translation_rescore( return ll.sum(dim=1) - ll_wil_total # (K,) +def correlation_at( + obs: TranslationObs, G: torch.Tensor, h_R: torch.Tensor, t: torch.Tensor, +) -> float: + """The translation search's own score at one translation. + + Same functional :func:`amplitude_translation_search` maximises -- + ``sum_h w (E_obs^2 - _w) |F_calc(h, t)|^2 / sum_h w |F_calc(h, t)|^2`` + -- evaluated at a single ``t`` rather than over a grid, which needs no FFT. + + The grid search returns its peaks' scores, but a peak is then refined and the + refined position is what gets used. Scoring the *used* translation is what + makes the number comparable with other candidates' used translations. It is + also what stops the ranking key and the returned placement coming from + different points, which is how a selection rule quietly stops meaning what + its name says. + """ + device = G.device + E = obs.E_obs.to(device).to(torch.float64) + w = obs.weight.to(device).to(torch.float64) + E2 = E * E + E2c = E2 - (w * E2).sum() / w.sum().clamp(min=1e-30) + + tt = t.detach().to(device).to(torch.float64).reshape(3) + phase = torch.exp(2j * torch.pi * torch.einsum( + "ind,d->in", h_R.to(torch.float64), tt).to(G.dtype)) + Fc2 = (G * phase).sum(dim=0).abs().to(torch.float64) ** 2 + + num = (w * E2c * Fc2).sum() + den = (w * Fc2).sum().clamp(min=1e-30) + return float(num / den) + + def precompute_G_for_rotation( interpolator, R_rotation: torch.Tensor, From 61b297ee463e73e7ac59b3891d86dcdbef329793 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 15:16:07 +0200 Subject: [PATCH 132/250] Make the candidate ranking rule a parameter and measure the three rank_by selects the winner among placed candidates by the analytical-scale R (default, unchanged), the translation function's correlation, or its Rice/Woolfson likelihood. It is not a tuning knob: the three disagree, a rank-level proxy got the ordering wrong, and the only way to compare them is end to end. Four structures x ten seeds, success = within 8 deg of canonical modulo crystal symmetry: llg 37/40 median residual 1.57 r 36/40 1.79 corr 32/40 2.77 The likelihood is at worst equal to R and picks visibly tighter placements -- 6 better, 2 worse, 27 tied on the cells both solve, and on 6G9X every residual falls under 2.3 deg against R's spread to 5.65. One cell of success difference over 40 is not significant, so the default does not move. The correlation is the one that matters here, and it answers the question that prompted this: no, the likelihood does NOT show what the correlation showed. The correlation's rank-level advantage (33/40 against R's 23/40) was an artefact of a truth label that disagrees with coordinate superposition; end to end it is the worst of the three. The likelihood's rank-level standing, which was middling, survives contact with coordinates. llg_at() evaluates the translation likelihood at a single translation, split from llg_translation_rescore because scoring a candidate and re-ranking a candidate's translations are different questions that happen to share a functional. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/marginal_seeds.sh | 2 +- alignment_lab/diagnostics/pose_recovery.py | 23 ++++-- docs/changelog.rst | 1 + torchref/experimental/alignment/pipeline.py | 76 ++++++++++++++----- .../experimental/alignment/translation.py | 21 +++++ 5 files changed, 100 insertions(+), 23 deletions(-) diff --git a/alignment_lab/analysis/marginal_seeds.sh b/alignment_lab/analysis/marginal_seeds.sh index fa07f6bf..7322bd07 100644 --- a/alignment_lab/analysis/marginal_seeds.sh +++ b/alignment_lab/analysis/marginal_seeds.sh @@ -26,5 +26,5 @@ export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" for T in $(seq 0 9); do "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial "$T" \ - --arms analytic_r --n-rotation-candidates 25 2>/dev/null | grep '^ROW ' + --arms analytic_r,corr,llg --n-rotation-candidates 25 2>/dev/null | grep '^ROW ' done diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py index 07f8fe67..0209d92b 100644 --- a/alignment_lab/diagnostics/pose_recovery.py +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -13,10 +13,18 @@ ``analytic_r`` the default -- rank each rotation candidate by the analytical-scale R at its best translation. +``corr`` + rank by the translation function's own correlation. Measured 31/40 against + ``analytic_r``'s 36/40 over four structures x ten seeds: worse, despite a + rank-level harness predicting the reverse on a truth label that disagreed + with coordinate superposition. +``llg`` + rank by the translation likelihood. The same rank-level harness rated it + between the other two, so it is here for the same reason: only the + end-to-end comparison is trustworthy. ``llg_tf`` - re-rank the translation peaks by the Rice/Woolfson LLG first. At rank level - the LLG puts truth at rank 0 in 27/30 against analytic R's 22/30; this is - the arm that says whether that carries through to pose. + a different question -- re-rank each candidate's TRANSLATIONS by the + likelihood, still selecting the candidate by R. Success mirrors the integration test: final coordinates within ``--success-deg`` of canonical, modulo the crystal symmetry. @@ -46,8 +54,13 @@ seed_for) ARMS = { - "analytic_r": dict(use_llg_tf=False), - "llg_tf": dict(use_llg_tf=True), + # How the winner is chosen among placed candidates. + "analytic_r": dict(use_llg_tf=False, rank_by="r"), + "corr": dict(use_llg_tf=False, rank_by="corr"), + "llg": dict(use_llg_tf=False, rank_by="llg"), + # Re-ranks each candidate's TRANSLATIONS by the likelihood, then still + # selects the candidate by R -- a different question from the three above. + "llg_tf": dict(use_llg_tf=True, rank_by="r"), } diff --git a/docs/changelog.rst b/docs/changelog.rst index d7c18086..6e97b50e 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- ``MolecularReplacementPipeline`` takes ``rank_by`` (``"r"``, ``"corr"``, ``"llg"``), because the three scores disagree about which placement is right and only an end-to-end comparison settles it. Over four structures x ten seeds: likelihood 37/40, analytical R 36/40 (the default), translation correlation 32/40. A rank-level harness had rated the correlation best by a wide margin, on a truth label that disagrees with coordinate superposition - The molecular-replacement pipeline's ``verbose`` levels are a documented contract routed through one emitter, rather than ``if verbose > 0: print(...)`` at seventeen sites. Level 2 emits one machine-readable ``CAND`` line per rotation candidate carrying every score the selection could have used, so a wrong placement can be diagnosed from the run itself instead of from a harness that re-implements the placement loop and then disagrees with it - The molecular-replacement pipeline places all ``n_rotation_candidates`` (now 25, was 15) and returns the best, instead of stopping once a placement beat an R-factor threshold. The old rule made the answer depend on the order the rotation function happened to produce and could accept the third candidate without scoring the tenth. Measured neutral -- identical placements on 10 structures x 3 seeds and on a 10-seed sweep of the marginal cases -- at about 1.6x the wall clock - The Wilson normaliser converges in 8-12 IRLS iterations instead of 26-102. Its stopping rule was ``|dL|`` per reflection against 1e-10, which asks eleven significant digits of a normalisation curve; it is now relative to the improvement so far, which is scale-invariant for the same reason the absolute form was chosen diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 5a41d567..6dd41c5c 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -61,6 +61,7 @@ TranslationPeak, amplitude_translation_search, correlation_at, + llg_at, fit_sigma_a_per_shell, llg_translation_rescore, local_translation_refine, @@ -264,9 +265,12 @@ class MRSolution: The translation function's correlation at the chosen translation, higher better. Reported, not ranked -- see the sort in :meth:`run`. r_factor : float - **The ranking key**: the analytical-scale R at that placement, lower - better. For the returned winner it is replaced by the solvent-aware - Scaler R-work. + The analytical-scale R at that placement, lower better. The default + ranking key; for the returned winner it is replaced by the + solvent-aware Scaler R-work. + llg_score : float + The translation likelihood at that placement, higher better. ``nan`` + unless ``rank_by="llg"``, since it costs a likelihood evaluation. model : ModelFT The rotated (+translated +refined) model for this candidate. """ @@ -277,6 +281,7 @@ class MRSolution: translation_score: float r_factor: float model: "ModelFT" + llg_score: float = float("nan") class MolecularReplacementPipeline(DeviceMixin): @@ -343,6 +348,12 @@ def __init__( n_translation_candidates: int = 3, translation_grid_steps: int = 16, use_llg_tf: bool = False, + # Which score picks the winner among placed candidates. "r" is the + # analytical-scale R-factor; "corr" the translation function's own + # correlation; "llg" its Rice/Woolfson likelihood. Not a tuning knob -- + # it exists because the three disagree and a rank-level proxy got the + # ordering wrong, so the comparison has to be made end to end. + rank_by: str = "r", # Resolution window for the TRANSLATION set only, independent of the # rotation search's [d_max, d_min]. None means no cut, which is what # this stage has always done -- see `_prepare_translation_arrays`. @@ -374,6 +385,10 @@ def __init__( self.n_translation_candidates = n_translation_candidates self.translation_grid_steps = translation_grid_steps self.use_llg_tf = use_llg_tf + if rank_by not in ("r", "corr", "llg"): + raise ValueError( + f"rank_by={rank_by!r}; expected 'r', 'corr' or 'llg'.") + self.rank_by = rank_by self.tf_d_min = tf_d_min self.tf_d_max = tf_d_max @@ -496,7 +511,7 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: f"rfz={float(peak_k.sigma):.3f} tf=nan r=nan " f"t=none # no translation peaks") continue - r_analytic, t_refined, tf_score = placement + r_analytic, t_refined, tf_score, llg_score = placement self._log_candidate(k, peak_k, r_analytic, t_refined, tf_score) placed = rotated_k.copy().translate( @@ -512,20 +527,25 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: translation_score=float(tf_score), r_factor=float(r_analytic), model=placed, + llg_score=float(llg_score), ) ) if not solutions: raise RuntimeError("Translation + joint refine produced no candidates.") - # Lowest analytical-scale R. Ranking by the translation correlation was - # tried and is worse end to end -- 31/40 against 36/40 over four - # structures x ten seeds, losing on both 2DQ6 (6/10 -> 3/10) and 6G9X - # (10/10 -> 8/10). A rank-level harness had predicted the opposite by a - # wide margin, which is why the correlation is still reported on every - # candidate at verbose >= 2: the two orderings disagree and the - # disagreement is not yet understood. - solutions.sort(key=lambda s: s.r_factor) + # Default is lowest analytical-scale R. Ranking by the correlation was + # measured end to end and is worse -- 31/40 against 36/40 over four + # structures x ten seeds -- despite a rank-level harness predicting the + # opposite by a wide margin, on a truth label that turned out to + # disagree with coordinate superposition. Hence `rank_by`: the scores + # disagree, and only the end-to-end comparison settles it. + if self.rank_by == "r": + solutions.sort(key=lambda s: s.r_factor) + elif self.rank_by == "corr": + solutions.sort(key=lambda s: -s.translation_score) + else: + solutions.sort(key=lambda s: -s.llg_score) winner = solutions[0] # Single solvent-aware Scaler refit on the winner for the user-facing R. @@ -639,9 +659,9 @@ def _prepare_translation_arrays(self) -> None: def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: """Translation search + analytical-R local refine for one rotation. - Returns ``(r_analytic, t_refined, tf_score)`` for the best translation - of this rotation candidate, or ``None`` if no translation peaks were - found. + Returns ``(r_analytic, t_refined, tf_score, llg_score)`` for the best + translation of this rotation candidate, or ``None`` if no translation + peaks were found. ``llg_score`` is ``nan`` unless it is the ranking key. ``tf_score`` is the translation function's own score at its top peak. It does not select anything -- ``r_analytic`` does -- but it is carried @@ -712,7 +732,28 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: f"t={[round(float(x), 3) for x in t_refined.tolist()]}") if best is None or r_analytic < best[0]: best = (r_analytic, t_refined, tf_ref) - return best + if best is None: + return None + llg = float("nan") + if self.rank_by == "llg": + # sigma_A at the chosen translation, per candidate. Fitting it once + # and sharing it across candidates was measured to give identical + # rankings, so the cheaper-to-reason-about form is used. + E_calc = normalise_calc( + self._fcalc_at(G_pre, h_R_pre, best[1]), self._obs) + sa = fit_sigma_a_per_shell( + self._obs.E_obs, E_calc, self._obs.centric, + self._obs.shell_idx, self._obs.n_shells, n_grid=81) + llg = llg_at(self._obs, G_pre, h_R_pre, best[1], sa) + return (best[0], best[1], best[2], llg) + + def _fcalc_at(self, G, h_R, t): + """``|F_calc(h, t)|`` from the precomputed per-symop contributions.""" + tt = t.detach().to(G.device).to(torch.float64).reshape(3) + phase = torch.exp(2j * torch.pi * torch.einsum( + "ind,d->in", h_R.to(torch.float64), tt).to(G.dtype)) + return (G * phase).sum(dim=0).abs().to(torch.float64) + def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): """Re-rank translation peaks by a shared-σA Rice/Woolfson LLG. @@ -792,6 +833,7 @@ def align_model_to_data( translation_grid_steps: int = 16, n_rotation_candidates: int = 25, use_llg_tf: bool = False, + rank_by: str = "r", tf_d_min: Optional[float] = None, tf_d_max: Optional[float] = None, model_error_A: Optional[float] = None, @@ -821,7 +863,7 @@ def align_model_to_data( n_translation_peaks=n_translation_peaks, n_translation_candidates=n_translation_candidates, translation_grid_steps=translation_grid_steps, - use_llg_tf=use_llg_tf, + use_llg_tf=use_llg_tf, rank_by=rank_by, tf_d_min=tf_d_min, tf_d_max=tf_d_max, ) solutions = pipeline.run(do_translation=do_translation) diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index c003ad7b..57b4bf54 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -616,6 +616,27 @@ def llg_translation_rescore( return ll.sum(dim=1) - ll_wil_total # (K,) +def llg_at( + obs: TranslationObs, + G: torch.Tensor, + h_R: torch.Tensor, + t: torch.Tensor, + sigma_a: torch.Tensor, +) -> float: + """The translation likelihood at one translation, as a candidate score. + + :func:`llg_translation_rescore` over a single ``t``. Split out because + scoring a *candidate* and re-ranking a candidate's *translations* are + different questions that happen to share a functional, and only the first + needs to be comparable across orientations. + """ + return float(llg_translation_rescore( + obs=obs, G=G, h_R=h_R, + t_candidates=t.detach().reshape(1, 3).to(G.device).to(torch.float64), + sigma_a=sigma_a, + )[0]) + + def correlation_at( obs: TranslationObs, G: torch.Tensor, h_R: torch.Tensor, t: torch.Tensor, ) -> float: From 5426c09f57dd969a213f1879ac27d9b4b2e43064 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 16:29:00 +0200 Subject: [PATCH 133/250] Rank molecular-replacement candidates by the translation likelihood Was the analytical-scale R-factor. On the ten-structure panel the likelihood gets 30/30 against 29/30, and over a ten-seed sweep of the marginal cases 37/40 against 36/40, with a median residual of 1.57 deg against 1.79. The success counts are one cell apart and that is not a result on its own. What carries the change is that the likelihood places better where both succeed -- 6 better against 2 worse, 27 tied -- and that it is the right object for the question. An R-factor compares a partial model against data at the resolution this stage runs at, and on 2DQ6 all 25 candidates land within 0.023 of each other in R while the wrong winner leads a correct candidate by 0.0001. It has almost nothing to distinguish with. The correlation stays available and is the cautionary case: a rank-level harness rated it best by a wide margin, 33/40 against R's 23/40, and end to end it is the worst of the three at 32/40. That harness's truth label disagrees with coordinate superposition. Costs about 2.5x the wall clock, from a sigma_A fit per candidate. Fitting it once and sharing it was measured to leave the ranking unchanged, so that saving is available and is not taken here. Every candidate now reports all three scores at verbose >= 2 whichever one ranks, because which score a wrong placement disagreed on is the question. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/marginal_seeds.sh | 2 +- alignment_lab/analysis/pose_arms.sh | 2 +- alignment_lab/diagnostics/pose_recovery.py | 13 ++-- docs/changelog.rst | 3 +- torchref/experimental/alignment/pipeline.py | 84 +++++++++++++-------- 5 files changed, 64 insertions(+), 40 deletions(-) diff --git a/alignment_lab/analysis/marginal_seeds.sh b/alignment_lab/analysis/marginal_seeds.sh index 7322bd07..e2bdf603 100644 --- a/alignment_lab/analysis/marginal_seeds.sh +++ b/alignment_lab/analysis/marginal_seeds.sh @@ -26,5 +26,5 @@ export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" for T in $(seq 0 9); do "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial "$T" \ - --arms analytic_r,corr,llg --n-rotation-candidates 25 2>/dev/null | grep '^ROW ' + --arms llg,analytic_r --n-rotation-candidates 25 2>/dev/null | grep '^ROW ' done diff --git a/alignment_lab/analysis/pose_arms.sh b/alignment_lab/analysis/pose_arms.sh index 9ffbf6f7..32641b11 100644 --- a/alignment_lab/analysis/pose_arms.sh +++ b/alignment_lab/analysis/pose_arms.sh @@ -31,5 +31,5 @@ export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREA export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" for T in 0 1 2; do "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial "$T" \ - --arms analytic_r,llg_tf --n-rotation-candidates 25 2>/dev/null | grep '^ROW ' + --arms llg,analytic_r --n-rotation-candidates 25 2>/dev/null | grep '^ROW ' done diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py index 0209d92b..acf5a601 100644 --- a/alignment_lab/diagnostics/pose_recovery.py +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -10,18 +10,17 @@ Arms (``--arms``) sweep how translation candidates are ranked: +``llg`` + the default -- rank each rotation candidate by the translation likelihood at + its best translation. 37/40 over four structures x ten seeds. ``analytic_r`` - the default -- rank each rotation candidate by the analytical-scale R at its - best translation. + rank by the analytical-scale R instead. 36/40, and places less well on the + cells both solve. ``corr`` rank by the translation function's own correlation. Measured 31/40 against ``analytic_r``'s 36/40 over four structures x ten seeds: worse, despite a rank-level harness predicting the reverse on a truth label that disagreed with coordinate superposition. -``llg`` - rank by the translation likelihood. The same rank-level harness rated it - between the other two, so it is here for the same reason: only the - end-to-end comparison is trustworthy. ``llg_tf`` a different question -- re-rank each candidate's TRANSLATIONS by the likelihood, still selecting the candidate by R. @@ -139,7 +138,7 @@ def main() -> int: ap = argparse.ArgumentParser(description=__doc__) ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) ap.add_argument("--trial", type=int, default=0) - ap.add_argument("--arms", default="analytic_r,llg_tf") + ap.add_argument("--arms", default="llg,analytic_r") ap.add_argument("--n-rotation-candidates", type=int, default=25) ap.add_argument("--n-rotation-peaks", type=int, default=200) ap.add_argument("--success-deg", type=float, default=8.0) diff --git a/docs/changelog.rst b/docs/changelog.rst index 6e97b50e..46a663dc 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,7 +4,8 @@ Changelog Unreleased ---------- -- ``MolecularReplacementPipeline`` takes ``rank_by`` (``"r"``, ``"corr"``, ``"llg"``), because the three scores disagree about which placement is right and only an end-to-end comparison settles it. Over four structures x ten seeds: likelihood 37/40, analytical R 36/40 (the default), translation correlation 32/40. A rank-level harness had rated the correlation best by a wide margin, on a truth label that disagrees with coordinate superposition +- Molecular-replacement candidates are ranked by the translation function's likelihood, not by an analytical-scale R-factor. 30/30 on the ten-structure panel against 29/30, and 37/40 against 36/40 over a ten-seed sweep, with tighter placements on the cells both solve. Selectable through ``rank_by``; the correlation is the third option and is the worst of the three at 32/40, despite a rank-level harness rating it best on a truth label that disagrees with coordinate superposition +- The likelihood ranking costs about 2.5x the wall clock, from a per-candidate sigma_A fit - The molecular-replacement pipeline's ``verbose`` levels are a documented contract routed through one emitter, rather than ``if verbose > 0: print(...)`` at seventeen sites. Level 2 emits one machine-readable ``CAND`` line per rotation candidate carrying every score the selection could have used, so a wrong placement can be diagnosed from the run itself instead of from a harness that re-implements the placement loop and then disagrees with it - The molecular-replacement pipeline places all ``n_rotation_candidates`` (now 25, was 15) and returns the best, instead of stopping once a placement beat an R-factor threshold. The old rule made the answer depend on the order the rotation function happened to produce and could accept the third candidate without scoring the tenth. Measured neutral -- identical placements on 10 structures x 3 seeds and on a 10-seed sweep of the marginal cases -- at about 1.6x the wall clock - The Wilson normaliser converges in 8-12 IRLS iterations instead of 26-102. Its stopping rule was ``|dL|`` per reflection against 1e-10, which asks eleven significant digits of a normalisation curve; it is now relative to the improvement so far, which is scale-invariant for the same reason the absolute form was chosen diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 6dd41c5c..160f04c9 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -27,10 +27,12 @@ plausibly that is not a choice between them, and it made the selection rule impossible to reason about or to measure against a ranking harness. -The candidates are ranked by the translation search's analytical R. The -user-facing solvent-aware R-work is computed once, on the winner. The pipeline -returns a *placement* -- refining it is the caller's job, and downstream -refinement does it better than a bolted-on polish did. +The candidates are ranked by the translation search's likelihood -- see +``rank_by``, and the sort in :meth:`MolecularReplacementPipeline.run` for what +the three available scores measured against each other. The user-facing +solvent-aware R-work is computed once, on the winner. The pipeline returns a +*placement* -- refining it is the caller's job, and downstream refinement does +it better than a bolted-on polish did. ``align_model_to_data`` delegates here; this class is the implementation of record. The crystallographic stages live in @@ -265,12 +267,13 @@ class MRSolution: The translation function's correlation at the chosen translation, higher better. Reported, not ranked -- see the sort in :meth:`run`. r_factor : float - The analytical-scale R at that placement, lower better. The default - ranking key; for the returned winner it is replaced by the - solvent-aware Scaler R-work. + The analytical-scale R at that placement, lower better. Reported, not + ranked; for the returned winner it is replaced by the solvent-aware + Scaler R-work, which is the number a caller reads. llg_score : float - The translation likelihood at that placement, higher better. ``nan`` - unless ``rank_by="llg"``, since it costs a likelihood evaluation. + **The ranking key**: the translation likelihood at that placement, + higher better. ``nan`` when ``rank_by`` is not ``"llg"``, since it costs + a sigma_A fit and a likelihood evaluation per candidate. model : ModelFT The rotated (+translated +refined) model for this candidate. """ @@ -348,12 +351,13 @@ def __init__( n_translation_candidates: int = 3, translation_grid_steps: int = 16, use_llg_tf: bool = False, - # Which score picks the winner among placed candidates. "r" is the - # analytical-scale R-factor; "corr" the translation function's own - # correlation; "llg" its Rice/Woolfson likelihood. Not a tuning knob -- - # it exists because the three disagree and a rank-level proxy got the - # ordering wrong, so the comparison has to be made end to end. - rank_by: str = "r", + # Which score picks the winner among placed candidates. "llg" is the + # translation function's Rice/Woolfson likelihood; "r" the + # analytical-scale R-factor; "corr" the translation correlation. Not a + # tuning knob -- it exists because the three disagree and a rank-level + # proxy got the ordering wrong, so the comparison has to be made end to + # end. See the sort in `run` for what that measured. + rank_by: str = "llg", # Resolution window for the TRANSLATION set only, independent of the # rotation search's [d_max, d_min]. None means no cut, which is what # this stage has always done -- see `_prepare_translation_arrays`. @@ -414,7 +418,7 @@ def _log(self, level: int, msg: str) -> None: print(msg, flush=True) def _log_candidate(self, k: int, peak, r_analytic, t_frac, - tf_score=None) -> None: + tf_score=None, llg_score=None) -> None: """One line per rotation candidate, with every score behind the choice. Machine-readable on purpose. Diagnosing a wrong placement means asking @@ -426,14 +430,17 @@ def _log_candidate(self, k: int, peak, r_analytic, t_frac, Fields are ``key=value`` so a reader does not depend on column order: ``k`` candidate index in rotation-function order, ``rf``/``rfz`` its - score and z, ``tf`` the translation correlation at the chosen peak, - ``r`` the analytical-scale R that ranks it, ``t`` the fractional - translation. + score and z, ``tf`` the translation correlation, ``llg`` the translation + likelihood, ``r`` the analytical-scale R, ``t`` the fractional + translation. All three placement scores are reported whichever one + ranks, because which of them a wrong placement disagreed on is the + question, and they do disagree. """ tf = "nan" if tf_score is None else f"{float(tf_score):.5f}" + llg = "nan" if llg_score is None else f"{float(llg_score):.1f}" t = ",".join(f"{float(x):.4f}" for x in t_frac) self._log(2, f"CAND k={k} rf={float(peak.score):.4f} " - f"rfz={float(peak.sigma):.3f} tf={tf} " + f"rfz={float(peak.sigma):.3f} tf={tf} llg={llg} " f"r={float(r_analytic):.5f} t={t}") @@ -441,7 +448,9 @@ def _log_candidate(self, k: int, peak, r_analytic, t_frac, # Public entry point # ------------------------------------------------------------------ def run(self, do_translation: bool = True) -> List[MRSolution]: - """Run the MR pipeline and return solutions ranked by R-factor. + """Run the MR pipeline and return solutions, best first. + + Ranked by ``rank_by`` -- the translation likelihood by default. Parameters ---------- @@ -512,7 +521,8 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: f"t=none # no translation peaks") continue r_analytic, t_refined, tf_score, llg_score = placement - self._log_candidate(k, peak_k, r_analytic, t_refined, tf_score) + self._log_candidate(k, peak_k, r_analytic, t_refined, tf_score, + llg_score) placed = rotated_k.copy().translate( t_refined.to(self.model.dtype_float), fractional=True, @@ -534,12 +544,24 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: if not solutions: raise RuntimeError("Translation + joint refine produced no candidates.") - # Default is lowest analytical-scale R. Ranking by the correlation was - # measured end to end and is worse -- 31/40 against 36/40 over four - # structures x ten seeds -- despite a rank-level harness predicting the - # opposite by a wide margin, on a truth label that turned out to - # disagree with coordinate superposition. Hence `rank_by`: the scores - # disagree, and only the end-to-end comparison settles it. + # Highest translation likelihood. Over four structures x ten seeds, + # success within 8 deg of canonical modulo crystal symmetry: + # + # llg 37/40 median residual 1.57 deg + # r 36/40 1.79 + # corr 32/40 2.77 + # + # The likelihood is chosen on the residuals rather than the success + # count -- one cell in 40 is not a result, but on the cells all three + # solve it places better 6 times against 2, and on 6G9X every residual + # falls under 2.3 deg where R spreads to 5.65. It is also the right + # object for the question: an R-factor on a partial model at the + # resolution this runs at has little to distinguish with. + # + # The correlation is here as a cautionary default-not-taken. A rank-level + # harness rated it best by a wide margin, 33/40 against R's 23/40, and + # end to end it is the worst of the three -- the harness's truth label + # disagreed with coordinate superposition. if self.rank_by == "r": solutions.sort(key=lambda s: s.r_factor) elif self.rank_by == "corr": @@ -554,7 +576,9 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: timer.stop("12_final_scaler") winner.model.last_alignment_rfactor = rwork_final winner.r_factor = rwork_final - self._log(1, f"mr: winner TF correlation={winner.translation_score:.5f}, " + self._log(1, f"mr: winner ({self.rank_by}) " + f"LLG={winner.llg_score:.1f} " + f"TF corr={winner.translation_score:.5f} " f"analytic R={winner.r_factor:.4f}, " f"final Scaler-fit R-work={rwork_final:.4f}") self._log(2, "\n" + timer.summary()) @@ -833,7 +857,7 @@ def align_model_to_data( translation_grid_steps: int = 16, n_rotation_candidates: int = 25, use_llg_tf: bool = False, - rank_by: str = "r", + rank_by: str = "llg", tf_d_min: Optional[float] = None, tf_d_max: Optional[float] = None, model_error_A: Optional[float] = None, From 46d20b31aae9e0c014949f3f406756ff4b0a9d0c Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Tue, 1 Sep 2026 16:35:44 +0200 Subject: [PATCH 134/250] Fixed some dtype inconsistenices and added test guarding inconsistent dtype handlign --- tests/helpers/dtype_inventory.py | 153 ++++++++++++++++++ tests/unit/test_dtype_conformance.py | 97 +++++++++++ torchref/base/direct_summation/_backends.py | 2 +- torchref/base/direct_summation/dispatch.py | 5 +- torchref/base/electron_density/_backends.py | 6 +- .../kernels/cpu/jit_reference.py | 16 +- .../kernels/cpu/sphere_splat.py | 2 +- .../kernels/cpu/variable_radius.py | 12 +- .../base/electron_density/map_building.py | 4 +- .../base/electron_density/solvent_mask.py | 2 +- torchref/base/electron_density/voxel_utils.py | 4 +- torchref/base/french_wilson.py | 18 +-- torchref/base/metrics/binwise_scale.py | 2 +- torchref/base/reciprocal/grid_operations.py | 16 +- torchref/base/reciprocal/interpolation.py | 10 +- torchref/base/reciprocal/symmetry.py | 2 +- torchref/base/scattering/scattering_table.py | 4 +- torchref/base/targets/_dispatch.py | 2 +- torchref/base/targets/xray_ml_full.py | 1 + torchref/base/wilson_outliers.py | 1 + torchref/cli/mtz2map.py | 7 +- torchref/cli/validate_ded.py | 6 +- torchref/experimental/alignment/clashscore.py | 12 +- torchref/experimental/alignment/pipeline.py | 4 + torchref/experimental/alignment/rigid_body.py | 20 ++- torchref/experimental/alignment/transform.py | 7 +- .../ensemble/ensemble_amber_kl.py | 2 +- .../experimental/ensemble/ensemble_model.py | 2 + torchref/experimental/ensemble/pca_model.py | 2 + .../ensemble/quasi_crystal_amber.py | 16 +- .../experimental/ensemble/rank_penalty.py | 2 + .../experimental/ensemble/wilson_prior.py | 2 +- torchref/experimental/kinetic/occupancies.py | 3 +- .../monolithic_refinement/density_scaler.py | 2 +- torchref/experimental/targets/amber_target.py | 12 +- .../experimental/targets/forcefield_target.py | 4 +- .../targets/sampled_ml_phase_target.py | 4 +- torchref/io/datasets/base.py | 9 +- torchref/io/datasets/fcalc_data.py | 2 +- torchref/io/datasets/reflection_data.py | 24 +-- torchref/model/disorder_field.py | 24 ++- torchref/model/model.py | 8 +- torchref/model/model_ft.py | 4 +- torchref/model/parameter_wrappers.py | 14 +- torchref/model/rigid_xyz.py | 2 +- torchref/refinement/base_refinement.py | 2 +- .../model_error_estimation/sigma_a.py | 6 +- torchref/refinement/optimizers/curvature.py | 2 +- torchref/refinement/targets/adp/rigid_bond.py | 2 +- torchref/refinement/targets/adp/sigd.py | 1 + torchref/refinement/targets/adp/similarity.py | 2 +- torchref/refinement/targets/difference.py | 4 +- .../refinement/targets/geometry/chiral.py | 2 +- .../refinement/targets/geometry/non_bonded.py | 4 +- torchref/refinement/targets/similarity.py | 12 +- torchref/scaling/collection_scaler.py | 4 +- torchref/scaling/scaler_base.py | 4 +- torchref/scaling/solvent.py | 6 +- torchref/symmetry/map_symmetry.py | 2 +- torchref/symmetry/reciprocal_symmetry.py | 24 +-- torchref/symmetry/symmetry.py | 10 +- torchref/topology/atom_graph.py | 12 +- torchref/topology/build.py | 4 +- torchref/topology/builders.py | 46 +++--- torchref/topology/edges.py | 4 +- torchref/topology/nonbonded.py | 34 ++-- torchref/topology/residue_graph.py | 2 +- torchref/topology/restraint_sets.py | 2 +- torchref/topology/restraints.py | 32 ++-- torchref/topology/riding.py | 44 ++--- torchref/topology/topology.py | 10 +- torchref/utils/device_mixin.py | 6 +- 72 files changed, 567 insertions(+), 269 deletions(-) create mode 100644 tests/helpers/dtype_inventory.py create mode 100644 tests/unit/test_dtype_conformance.py diff --git a/tests/helpers/dtype_inventory.py b/tests/helpers/dtype_inventory.py new file mode 100644 index 00000000..69848dae --- /dev/null +++ b/tests/helpers/dtype_inventory.py @@ -0,0 +1,153 @@ +"""Static inventory of hardcoded float/int/complex dtypes in the torchref source. + +The library resolves one dtype per category at import (``get_float_dtype`` / +``get_int_dtype`` / ``get_complex_dtype``) and every allocation on a live path is +expected to honour it. A literal ``torch.float32`` / ``torch.int64`` / +``torch.complex128`` baked into an allocation is a latent bug: MPS has no +float64, so a float64 config silently downcasts and a float32 config silently +upcasts -- neither raises, both corrupt results far from the cause. + +This module finds every guarded ``torch.`` reference by parsing the source +(AST, not regex, so dtypes named inside docstrings or comments do not count), +and reports whether each one carries an inline justification. The guard test in +``tests/unit/test_dtype_conformance.py`` turns that into a rule: outside a small +set of inherently-exempt modules, every hardcoded dtype must be justified with a +``# dtype-ok: `` marker on its own line or the comment block above it. + +``torch.bool`` is not guarded: a mask is categorical, not numeric precision. +""" + +from __future__ import annotations + +import ast +from pathlib import Path +from typing import List, NamedTuple + +__all__ = [ + "FLOAT_DTYPES", + "INT_DTYPES", + "COMPLEX_DTYPES", + "GUARDED_DTYPES", + "JUSTIFY_MARKER", + "EXEMPT_PREFIXES", + "DtypeUse", + "find_hardcoded_dtypes", + "is_exempt", +] + +# The dtypes that must not be hardcoded on a live path -- each has a config +# default (get_float_dtype / get_int_dtype / get_complex_dtype) that an +# allocation is meant to honour. ``torch.bool`` is deliberately excluded: a mask +# is categorical, not numeric precision, so pinning it is correct not a deviation. +FLOAT_DTYPES = frozenset( + {"float64", "float32", "float16", "double", "half", "bfloat16"} +) +INT_DTYPES = frozenset( + { + "int64", "int32", "int16", "int8", + "uint8", "uint16", "uint32", "uint64", + "long", "int", "short", "char", "byte", + } +) +COMPLEX_DTYPES = frozenset( + {"complex128", "complex64", "complex32", "cfloat", "cdouble", "chalf"} +) + +# Category lookup so a finding can say which config default it should use. +_CATEGORY = {name: "float" for name in FLOAT_DTYPES} +_CATEGORY.update({name: "int" for name in INT_DTYPES}) +_CATEGORY.update({name: "complex" for name in COMPLEX_DTYPES}) +GUARDED_DTYPES = frozenset(_CATEGORY) + +# A hardcoded float dtype is allowed when this marker appears on its line or the +# line immediately above it. The text after the colon is the required reason. +JUSTIFY_MARKER = "# dtype-ok:" + +# Module path prefixes (relative to the package parent, e.g. "torchref/...") +# where hardcoded float dtypes are inherent to the file's job and a per-line +# marker would be noise rather than signal: +# * triton kernels compile against explicit, static dtypes; +# * config.py *defines* the dtype maps the rest of the code reads; +# * scripts/ generate static on-disk tables offline, not model tensors. +EXEMPT_PREFIXES = ( + "torchref/base/targets/triton/", + "torchref/base/direct_summation/triton_ds.py", + "torchref/config.py", + "torchref/scripts/", +) + + +class DtypeUse(NamedTuple): + """One ``torch.`` reference (float, int, or complex) in the source.""" + + where: str # "relative/path.py:lineno" + rel_path: str # "relative/path.py" + lineno: int + dtype: str # e.g. "float64" + category: str # "float" | "int" | "complex" + line: str # the source line, stripped + justified: bool # carries JUSTIFY_MARKER on its line or the block above + + +def is_exempt(rel_path: str) -> bool: + """Whether ``rel_path`` is in an inherently-exempt module.""" + return any(rel_path.startswith(p) for p in EXEMPT_PREFIXES) + + +def _is_torch_dtype(node: ast.AST) -> str | None: + """Return the dtype name if ``node`` is a guarded ``torch.``, else None.""" + if ( + isinstance(node, ast.Attribute) + and node.attr in GUARDED_DTYPES + and isinstance(node.value, ast.Name) + and node.value.id == "torch" + ): + return node.attr + return None + + +def find_hardcoded_dtypes(package_root: Path) -> List[DtypeUse]: + """Every guarded ``torch.`` reference under ``package_root``. + + Uses the AST so references inside strings and comments are not counted, then + reads the raw source lines to decide whether each carries a justification. + """ + uses: List[DtypeUse] = [] + + for path in sorted(package_root.rglob("*.py")): + text = path.read_text(encoding="utf-8") + try: + tree = ast.parse(text, filename=str(path)) + except (SyntaxError, UnicodeDecodeError): # pragma: no cover + continue + lines = text.splitlines() + rel = str(path.relative_to(package_root.parent)) + + for node in ast.walk(tree): + dtype = _is_torch_dtype(node) + if dtype is None: + continue + lineno = node.lineno + this_line = lines[lineno - 1] if 0 < lineno <= len(lines) else "" + # A marker counts if it is on the reference's own line, or anywhere in + # the contiguous block of comment-only lines immediately above it -- so + # a multi-line justification works with the marker on any of its lines. + justified = JUSTIFY_MARKER in this_line + i = lineno - 2 + while not justified and i >= 0 and lines[i].strip().startswith("#"): + if JUSTIFY_MARKER in lines[i]: + justified = True + i -= 1 + uses.append( + DtypeUse( + where=f"{rel}:{lineno}", + rel_path=rel, + lineno=lineno, + dtype=dtype, + category=_CATEGORY[dtype], + line=this_line.strip(), + justified=justified, + ) + ) + + return uses diff --git a/tests/unit/test_dtype_conformance.py b/tests/unit/test_dtype_conformance.py new file mode 100644 index 00000000..8dbb8347 --- /dev/null +++ b/tests/unit/test_dtype_conformance.py @@ -0,0 +1,97 @@ +"""Dtype conformance: no unjustified hardcoded float dtype anywhere in the source. + +The package resolves one float dtype at import (``config.get_float_dtype()``) +and every allocation on a live path is meant to honour it. A literal +``torch.float32`` / ``torch.float64`` baked into an allocation is a latent bug: +MPS has no float64, so a float64 config silently downcasts and a float32 config +silently upcasts -- neither raises, and the corruption surfaces far from its +cause. + +This guard makes the rule enforceable. Outside a small set of inherently-exempt +modules (triton kernels, the config dtype maps, offline scripts), every +``torch.`` reference must carry a one-line justification:: + + x = torch.tensor(v, dtype=torch.float64) # dtype-ok: SVD needs f64 stability + +The marker may sit on the reference's own line or the line immediately above it. +The point is not to ban hardcoded dtypes -- some are correct (dtype validation, +deliberate high-precision accumulation, backend capability declarations) -- but +to force each one to say *why* it deviates, so a reviewer can tell a considered +choice from an oversight at a glance. +""" + +from pathlib import Path + +import pytest + +from tests.helpers.dtype_inventory import ( + EXEMPT_PREFIXES, + JUSTIFY_MARKER, + find_hardcoded_dtypes, + is_exempt, +) + +_PACKAGE_ROOT = Path(__file__).resolve().parents[2] / "torchref" + +# The config getter each category should defer to, quoted in the failure message. +_GETTER = { + "float": "config.get_float_dtype()", + "int": "config.get_int_dtype()", + "complex": "config.get_complex_dtype()", +} + + +@pytest.mark.unit +def test_no_unjustified_hardcoded_dtype(): + """Every hardcoded float/int/complex dtype on a live path is justified.""" + uses = find_hardcoded_dtypes(_PACKAGE_ROOT) + offenders = [u for u in uses if not is_exempt(u.rel_path) and not u.justified] + + def fmt(u): + return f"{u.where} {u.dtype} -> {_GETTER[u.category]} | {u.line}" + + assert not offenders, ( + f"{len(offenders)} hardcoded dtype(s) with no justification. Either switch " + f"to the config default for that category, or if the literal is deliberate " + f"add a '{JUSTIFY_MARKER} ' marker on the line or the block above:\n " + + "\n ".join(fmt(u) for u in offenders) + ) + + +@pytest.mark.unit +def test_justifications_carry_a_reason(): + """A ``# dtype-ok:`` marker must be followed by an actual reason, not left blank.""" + blank = [] + for path in sorted(_PACKAGE_ROOT.rglob("*.py")): + for i, line in enumerate(path.read_text(encoding="utf-8").splitlines(), 1): + if JUSTIFY_MARKER in line: + reason = line.split(JUSTIFY_MARKER, 1)[1].strip() + if not reason: + rel = path.relative_to(_PACKAGE_ROOT.parent) + blank.append(f"{rel}:{i}") + + assert not blank, ( + f"'{JUSTIFY_MARKER}' markers with no reason after the colon:\n " + + "\n ".join(blank) + ) + + +@pytest.mark.unit +def test_exempt_prefixes_still_match_something(): + """Stop ``EXEMPT_PREFIXES`` accumulating entries for paths that are long gone. + + An exemption that no longer matches any source is either a typo or a stale + excuse -- both hide the fact that the module it was meant to cover is now + unguarded (or renamed and silently re-included). + """ + uses = find_hardcoded_dtypes(_PACKAGE_ROOT) + seen_paths = {u.rel_path for u in uses} + stale = [ + p + for p in EXEMPT_PREFIXES + if not any(rp.startswith(p) for rp in seen_paths) + ] + assert not stale, ( + "EXEMPT_PREFIXES entries that match no source file with a hardcoded " + f"float dtype (stale or mistyped): {stale}" + ) diff --git a/torchref/base/direct_summation/_backends.py b/torchref/base/direct_summation/_backends.py index 9b581f32..7011f6bb 100644 --- a/torchref/base/direct_summation/_backends.py +++ b/torchref/base/direct_summation/_backends.py @@ -74,7 +74,7 @@ def _ds_aniso_triton(hkl, s_vec, xyz_frac, occ, U, A, B, max_memory_gb): name="ds_triton", kernel=(_THIS, "_ds_iso_triton", "_ds_aniso_triton"), device="cuda", - dtypes=(torch.float32,), + dtypes=(torch.float32,), # dtype-ok: backend capability declaration, not an allocation # Every argument except ``hkl`` (position 0), whose dtype provably costs # nothing -- see the module docstring. probes=(1, 2, 3, 4, 5, 6), diff --git a/torchref/base/direct_summation/dispatch.py b/torchref/base/direct_summation/dispatch.py index d3005af8..4fe3fd64 100644 --- a/torchref/base/direct_summation/dispatch.py +++ b/torchref/base/direct_summation/dispatch.py @@ -20,6 +20,7 @@ import torch +from torchref.config import get_complex_dtype from torchref.base.direct_summation.isotropic import ( _estimate_batch_size, iso_structure_factor_torched, @@ -241,7 +242,7 @@ def ds_iso(hkl, s, xyz_frac, occ, adp, A, B, *, force_portable=None, max_memory_ reference path regardless. """ if xyz_frac.shape[0] == 0: - return torch.zeros(hkl.shape[0], dtype=torch.complex64, device=hkl.device) + return torch.zeros(hkl.shape[0], dtype=get_complex_dtype(), device=hkl.device) return _dispatch( False, hkl, s, xyz_frac, occ, adp, A, B, force_portable, max_memory_gb ) @@ -253,7 +254,7 @@ def ds_aniso(hkl, s_vec, xyz_frac, occ, U, A, B, *, force_portable=None, max_mem See :func:`ds_iso` on ``force_portable=None``. """ if xyz_frac.shape[0] == 0: - return torch.zeros(hkl.shape[0], dtype=torch.complex64, device=hkl.device) + return torch.zeros(hkl.shape[0], dtype=get_complex_dtype(), device=hkl.device) return _dispatch( True, hkl, s_vec, xyz_frac, occ, U, A, B, force_portable, max_memory_gb ) diff --git a/torchref/base/electron_density/_backends.py b/torchref/base/electron_density/_backends.py index b43ca3aa..ef83ab1c 100644 --- a/torchref/base/electron_density/_backends.py +++ b/torchref/base/electron_density/_backends.py @@ -39,7 +39,7 @@ name="cuda_triton", kernel=(_CUDA, "add_isotropic_cuda_var", "add_anisotropic_cuda_var"), device="cuda", - dtypes=(torch.float32,), + dtypes=(torch.float32,), # dtype-ok: backend capability declaration, not an allocation probes=_ATOM_ARGS, probe=(_CUDA, "why_unavailable"), expect_available="cuda", @@ -53,7 +53,7 @@ name="mps_metal", kernel=(_MPS, "add_isotropic_mps_var", "add_anisotropic_mps_var"), device="mps", - dtypes=(torch.float32,), + dtypes=(torch.float32,), # dtype-ok: backend capability declaration, not an allocation probes=_ATOM_ARGS, probe=("torchref.base.electron_density.kernels.mps.compile", "why_unavailable"), @@ -66,7 +66,7 @@ kernel=(_SPHERE, "add_isotropic_cpu_sphere_var", "add_anisotropic_cpu_sphere_var"), device="cpu", - dtypes=(torch.float32, torch.float64), + dtypes=(torch.float32, torch.float64), # dtype-ok: backend capability declaration, not an allocation # Uniformity, not membership: the kernel picks one ``scalar_t`` from the output # map and then reads every other tensor through a raw pointer of that type, so a # float64 map beside float32 atoms would be a 2x out-of-bounds read. diff --git a/torchref/base/electron_density/kernels/cpu/jit_reference.py b/torchref/base/electron_density/kernels/cpu/jit_reference.py index 5394ee5d..257fa6d0 100644 --- a/torchref/base/electron_density/kernels/cpu/jit_reference.py +++ b/torchref/base/electron_density/kernels/cpu/jit_reference.py @@ -136,9 +136,9 @@ def forward( ny: int = density_map.shape[1] nz: int = density_map.shape[2] strides = torch.tensor( - [ny * nz, nz, 1], device=voxel_indices.device, dtype=torch.long + [ny * nz, nz, 1], device=voxel_indices.device, dtype=torch.long # dtype-ok: CPU-kernel strides for flat voxel index arithmetic; indexing requires long ) - index_flat = torch.sum(voxel_indices.to(torch.long) * strides, dim=-1).view(-1) + index_flat = torch.sum(voxel_indices.to(torch.long) * strides, dim=-1).view(-1) # dtype-ok: voxel indices flattened for scatter; indexing requires long density_map.view(-1).scatter_add_(0, index_flat, density.reshape(-1)) return density_map @@ -231,9 +231,9 @@ def forward( ny: int = density_map.shape[1] nz: int = density_map.shape[2] index_flat = ( - voxel_indices[:, :, 0].to(torch.int64) * (ny * nz) - + voxel_indices[:, :, 1].to(torch.int64) * nz - + voxel_indices[:, :, 2].to(torch.int64) + voxel_indices[:, :, 0].to(torch.int64) * (ny * nz) # dtype-ok: voxel-index flat-arithmetic term for scatter; requires int64 + + voxel_indices[:, :, 1].to(torch.int64) * nz # dtype-ok: voxel-index flat-arithmetic term for scatter; requires int64 + + voxel_indices[:, :, 2].to(torch.int64) # dtype-ok: voxel-index flat-arithmetic term for scatter; requires int64 ).flatten() density_map.view(-1).scatter_add_(0, index_flat, density.flatten()) @@ -307,9 +307,9 @@ def _add_to_map_gpu_simple( ny, nz = density_map.shape[1], density_map.shape[2] index_flat = ( - voxel_indices[:, :, 0].to(torch.int64) * (ny * nz) - + voxel_indices[:, :, 1].to(torch.int64) * nz - + voxel_indices[:, :, 2].to(torch.int64) + voxel_indices[:, :, 0].to(torch.int64) * (ny * nz) # dtype-ok: voxel-index flat-arithmetic term for scatter; requires int64 + + voxel_indices[:, :, 1].to(torch.int64) * nz # dtype-ok: voxel-index flat-arithmetic term for scatter; requires int64 + + voxel_indices[:, :, 2].to(torch.int64) # dtype-ok: voxel-index flat-arithmetic term for scatter; requires int64 ).flatten() density_map.view(-1).scatter_add_(0, index_flat, density.flatten()) diff --git a/torchref/base/electron_density/kernels/cpu/sphere_splat.py b/torchref/base/electron_density/kernels/cpu/sphere_splat.py index 9310769f..8d7dde39 100644 --- a/torchref/base/electron_density/kernels/cpu/sphere_splat.py +++ b/torchref/base/electron_density/kernels/cpu/sphere_splat.py @@ -608,7 +608,7 @@ def _prep(density_map, xyz, radius_per_atom, *tensors): f"sphere_splat is a CPU kernel; got device {density_map.device}" ) dtype = density_map.dtype - if dtype not in (torch.float32, torch.float64): + if dtype not in (torch.float32, torch.float64): # dtype-ok: validation guard, not an allocation raise ValueError(f"sphere_splat supports float32/float64, got {dtype}") for t in (xyz, radius_per_atom) + tensors: if t.dtype != dtype: diff --git a/torchref/base/electron_density/kernels/cpu/variable_radius.py b/torchref/base/electron_density/kernels/cpu/variable_radius.py index ba3393c0..bb19ae93 100644 --- a/torchref/base/electron_density/kernels/cpu/variable_radius.py +++ b/torchref/base/electron_density/kernels/cpu/variable_radius.py @@ -52,7 +52,7 @@ def _bucket_by_radius(radius: torch.Tensor, center_1d: torch.Tensor): spans.append((float(r), cursor, cursor + idx.numel())) cursor += idx.numel() order = (torch.cat(order_parts) if order_parts - else torch.zeros(0, dtype=torch.long, device=radius.device)) + else torch.zeros(0, dtype=torch.long, device=radius.device)) # dtype-ok: empty voxel-index fallback; must stay long for indexing return order, spans @@ -88,7 +88,7 @@ def _canonical_setup(xyz, inv_frac, frac, grid_dims, radius_per_atom, dtype): nx, ny, nz = grid_dims grid_f = torch.tensor(grid_dims, device=device, dtype=dtype) xyz_frac = (xyz @ inv_frac.T) % 1.0 - center_idx = torch.round(xyz_frac * grid_f).to(torch.long) + center_idx = torch.round(xyz_frac * grid_f).to(torch.long) # dtype-ok: rounded voxel center indices; torch indexing requires long # w0: atom position relative to its anchor node, in Cartesian. This is what # centres the sphere on the atom rather than on the node. w0 = (xyz_frac - center_idx.to(dtype) / grid_f) @ frac.T @@ -111,8 +111,8 @@ def add_isotropic_plain_var(density_map, xyz, adp, occ, A, B, device, dtype = xyz.device, density_map.dtype nx, ny, nz = (int(s) for s in density_map.shape) grid_dims = (nx, ny, nz) - strides = torch.tensor([ny * nz, nz, 1], device=device, dtype=torch.long) - grid_shape = torch.tensor(grid_dims, device=device, dtype=torch.long) + strides = torch.tensor([ny * nz, nz, 1], device=device, dtype=torch.long) # dtype-ok: strides for flat voxel-index arithmetic; indexing requires long + grid_shape = torch.tensor(grid_dims, device=device, dtype=torch.long) # dtype-ok: grid_shape for flat voxel-index arithmetic; indexing requires long order, spans, center_idx, w0 = _canonical_setup( xyz, inv_frac_matrix, frac_matrix, grid_dims, radius_per_atom, dtype) @@ -151,8 +151,8 @@ def add_anisotropic_plain_var(density_map, xyz, u, occ, A, B, device, dtype = xyz.device, density_map.dtype nx, ny, nz = (int(s) for s in density_map.shape) grid_dims = (nx, ny, nz) - strides = torch.tensor([ny * nz, nz, 1], device=device, dtype=torch.long) - grid_shape = torch.tensor(grid_dims, device=device, dtype=torch.long) + strides = torch.tensor([ny * nz, nz, 1], device=device, dtype=torch.long) # dtype-ok: strides for flat voxel-index arithmetic; indexing requires long + grid_shape = torch.tensor(grid_dims, device=device, dtype=torch.long) # dtype-ok: grid_shape for flat voxel-index arithmetic; indexing requires long order, spans, center_idx, w0 = _canonical_setup( xyz, inv_frac_matrix, frac_matrix, grid_dims, radius_per_atom, dtype) diff --git a/torchref/base/electron_density/map_building.py b/torchref/base/electron_density/map_building.py index 1d9b90fd..d00d7664 100644 --- a/torchref/base/electron_density/map_building.py +++ b/torchref/base/electron_density/map_building.py @@ -28,11 +28,11 @@ def scatter_add_nd(source, index, map): """Vectorized n-dimensional scatter-add: ``source`` ``(N,)`` into ``map`` ``(d1..dn)`` at ``index`` ``(N, ndim)``, returning the modified map. """ - map_shape = torch.tensor(map.shape, device=index.device, dtype=torch.int64) + map_shape = torch.tensor(map.shape, device=index.device, dtype=torch.int64) # dtype-ok: map_shape for stride/flat-index arithmetic feeding scatter_add; requires int64 # Convert n-dimensional indices to flat indices # For shape (d1, d2, d3, ..., dn), flat_index = i0 * (d1*d2*...*dn) + i1 * (d2*d3*...*dn) + ... + in - strides = torch.ones(len(map_shape), device=index.device, dtype=torch.int64) + strides = torch.ones(len(map_shape), device=index.device, dtype=torch.int64) # dtype-ok: strides for flat scatter_add index; requires int64 for i in range(len(map_shape) - 2, -1, -1): strides[i] = strides[i + 1] * map_shape[i + 1] diff --git a/torchref/base/electron_density/solvent_mask.py b/torchref/base/electron_density/solvent_mask.py index 08f1ceb7..2e36e709 100644 --- a/torchref/base/electron_density/solvent_mask.py +++ b/torchref/base/electron_density/solvent_mask.py @@ -113,7 +113,7 @@ def add_to_phenix_mask( ) # (N_atoms, N_voxels) # Flatten for scatter operations - voxel_indices_flat = voxel_indices.reshape(-1, 3).to(torch.long) + voxel_indices_flat = voxel_indices.reshape(-1, 3).to(torch.long) # dtype-ok: voxel indices for grid indexing; requires long # Create protein core mask using scatter_add int_dtype = dtypes.int diff --git a/torchref/base/electron_density/voxel_utils.py b/torchref/base/electron_density/voxel_utils.py index 88c2d037..3fc903d7 100644 --- a/torchref/base/electron_density/voxel_utils.py +++ b/torchref/base/electron_density/voxel_utils.py @@ -49,13 +49,13 @@ def find_relevant_voxels(real_space_grid, xyz, radius_angstrom=4, inv_frac_matri # This ensures atoms outside the unit cell are correctly wrapped xyz_frac = torch.matmul(inv_frac_matrix, xyz.T).T # (N, 3) xyz_frac = xyz_frac % 1.0 # Wrap to [0, 1] - center_idx = torch.round(xyz_frac * grid_shape.unsqueeze(0)).to(torch.int64) + center_idx = torch.round(xyz_frac * grid_shape.unsqueeze(0)).to(torch.int64) # dtype-ok: rounded voxel center indices; torch indexing requires int64 else: # Fallback for orthogonal cells (less accurate for non-orthogonal) voxelsize = real_space_grid[3, 3, 3] - real_space_grid[2, 2, 2] center_idx = torch.round( (xyz - grid_origin.unsqueeze(0)) / voxelsize.unsqueeze(0) - ).to(torch.int64) + ).to(torch.int64) # dtype-ok: voxel index cast; torch indexing requires int64 voxel_indices_wrapped = excise_angstrom_radius_around_coord( real_space_grid, center_idx, radius_angstrom diff --git a/torchref/base/french_wilson.py b/torchref/base/french_wilson.py index a48128db..a03c6df2 100644 --- a/torchref/base/french_wilson.py +++ b/torchref/base/french_wilson.py @@ -116,7 +116,7 @@ 2.906, 3.004, ], - dtype=torch.float32, + dtype=get_float_dtype(), ) AC_ZJ_SD = torch.tensor( @@ -193,7 +193,7 @@ 0.994, 0.996, ], - dtype=torch.float32, + dtype=get_float_dtype(), ) AC_ZF = torch.tensor( @@ -270,7 +270,7 @@ 1.676, 1.706, ], - dtype=torch.float32, + dtype=get_float_dtype(), ) AC_ZF_SD = torch.tensor( @@ -347,7 +347,7 @@ 0.310, 0.304, ], - dtype=torch.float32, + dtype=get_float_dtype(), ) # Centric lookup tables from French-Wilson supplement (1978) @@ -435,7 +435,7 @@ 3.753, 3.962, ], - dtype=torch.float32, + dtype=get_float_dtype(), ) C_ZJ_SD = torch.tensor( @@ -522,7 +522,7 @@ 1.029, 1.028, ], - dtype=torch.float32, + dtype=get_float_dtype(), ) C_ZF = torch.tensor( @@ -609,7 +609,7 @@ 1.917, 1.945, ], - dtype=torch.float32, + dtype=get_float_dtype(), ) C_ZF_SD = torch.tensor( @@ -696,7 +696,7 @@ 0.278, 0.272, ], - dtype=torch.float32, + dtype=get_float_dtype(), ) @@ -1194,7 +1194,7 @@ def estimate_mean_intensity_by_resolution( # Use scatter_add to compute sum of intensities per bin bin_sums = torch.zeros(actual_n_bins, dtype=I.dtype, device=I.device) - bin_counts = torch.zeros(actual_n_bins, dtype=torch.long, device=I.device) + bin_counts = torch.zeros(actual_n_bins, dtype=torch.long, device=I.device) # dtype-ok: count accumulator; scatter_add source is long ones, dtype must match bin_sums.scatter_add_(0, bin_indices, I_sorted) bin_counts.scatter_add_(0, bin_indices, torch.ones_like(bin_indices)) diff --git a/torchref/base/metrics/binwise_scale.py b/torchref/base/metrics/binwise_scale.py index 0ca6b213..2d803ef1 100644 --- a/torchref/base/metrics/binwise_scale.py +++ b/torchref/base/metrics/binwise_scale.py @@ -58,7 +58,7 @@ def binwise_scale( Fo = Fo.reshape(-1) device, dtype = Fc.device, Fc.dtype - bins = bins.reshape(-1).to(device=device, dtype=torch.int64) + bins = bins.reshape(-1).to(device=device, dtype=torch.int64) # dtype-ok: resolution-bin indices used as scatter_add index; requires int64 if nbins is None: nbins = int(bins.max().item()) + 1 if bins.numel() else 0 diff --git a/torchref/base/reciprocal/grid_operations.py b/torchref/base/reciprocal/grid_operations.py index 08fa138f..a631a640 100644 --- a/torchref/base/reciprocal/grid_operations.py +++ b/torchref/base/reciprocal/grid_operations.py @@ -47,14 +47,14 @@ def place_on_grid( dtype = structure_factor.dtype Nx, Ny, Nz = [int(x) for x in grid_size] hkls = hkls.to(device=device) - h = hkls[:, 0].to(torch.int64) - k = hkls[:, 1].to(torch.int64) - l = hkls[:, 2].to(torch.int64) + h = hkls[:, 0].to(torch.int64) # dtype-ok: hkl component cast to int64 for flat grid-index arithmetic; indexing requires long + k = hkls[:, 1].to(torch.int64) # dtype-ok: hkl component cast to int64 for flat grid-index arithmetic; indexing requires long + l = hkls[:, 2].to(torch.int64) # dtype-ok: hkl component cast to int64 for flat grid-index arithmetic; indexing requires long hi = torch.remainder(h, Nx) ki = torch.remainder(k, Ny) li = torch.remainder(l, Nz) - lin = (hi * (Ny * Nz) + ki * Nz + li).to(torch.int64) # (N,) + lin = (hi * (Ny * Nz) + ki * Nz + li).to(torch.int64) # (N,) # dtype-ok: flat grid index (lin) for scatter/gather; requires int64 grid = torch.zeros((B, Nx * Ny * Nz), dtype=dtype, device=device) grid = grid.index_add(1, lin, structure_factor) # (B, Nx*Ny*Nz) @@ -62,7 +62,7 @@ def place_on_grid( hi_sym = torch.remainder(-h, Nx) ki_sym = torch.remainder(-k, Ny) li_sym = torch.remainder(-l, Nz) - lin_sym = (hi_sym * (Ny * Nz) + ki_sym * Nz + li_sym).to(torch.int64) + lin_sym = (hi_sym * (Ny * Nz) + ki_sym * Nz + li_sym).to(torch.int64) # dtype-ok: symmetry flat grid index (lin_sym) for scatter/gather; requires int64 vals_conj = torch.conj(structure_factor) grid = grid.index_add(1, lin_sym, vals_conj) @@ -101,9 +101,9 @@ def extract_structure_factor_from_grid(reciprocal_grid, hkls) -> torch.Tensor: # Same wrapping convention as place_on_grid. hkls = hkls.to(device=device) - h = hkls[:, 0].to(torch.int64) - k = hkls[:, 1].to(torch.int64) - l = hkls[:, 2].to(torch.int64) + h = hkls[:, 0].to(torch.int64) # dtype-ok: hkl component cast to int64 for flat grid-index arithmetic; indexing requires long + k = hkls[:, 1].to(torch.int64) # dtype-ok: hkl component cast to int64 for flat grid-index arithmetic; indexing requires long + l = hkls[:, 2].to(torch.int64) # dtype-ok: hkl component cast to int64 for flat grid-index arithmetic; indexing requires long hi = torch.remainder(h, Nx) ki = torch.remainder(k, Ny) diff --git a/torchref/base/reciprocal/interpolation.py b/torchref/base/reciprocal/interpolation.py index e7db4f18..4b611634 100644 --- a/torchref/base/reciprocal/interpolation.py +++ b/torchref/base/reciprocal/interpolation.py @@ -39,6 +39,9 @@ def interpolate_structure_factor_from_grid( device = reciprocal_grid.device Nx, Ny, Nz = reciprocal_grid.shape + # dtype-ok: grid-index interpolation only. hkl are small integers, so floor + # and fractional weights are exact in float32; the weights are recast to the + # grid's dtype below, so nothing here mixes with config-dtype tensors. hkl_float = hkl_float.to(device=device, dtype=torch.float32) # Get the 8 corner indices for trilinear interpolation @@ -149,6 +152,9 @@ def interpolate_complex_from_grid( device = reciprocal_grid.device Nx, Ny, Nz = reciprocal_grid.shape + # dtype-ok: grid-index interpolation only. hkl are small integers, so floor + # and fractional weights are exact in float32; the weights are recast to the + # grid's dtype below, so nothing here mixes with config-dtype tensors. hkl_float = hkl_float.to(device=device, dtype=torch.float32) # Get the 8 corner indices for trilinear interpolation @@ -314,7 +320,9 @@ def interpolate_for_rotation(hkl, R, cell, reciprocal_space_grid): batched = False rotation_in_s = torch.einsum('ij, bjk -> bik', cell.reciprocal_basis_matrix, R.permute(0,2,1)) rotation_in_hkl = torch.einsum('bij, jk -> bik', rotation_in_s, cell.reciprocal_basis_matrix.inverse()) - reoriented_hkl = torch.einsum('aj, bji -> bai', hkl.to(torch.float32), rotation_in_hkl) + # Match the reciprocal-basis math's dtype (config float): a hardcoded float32 + # here mixes with a float64 rotation under a float64 config and raises. + reoriented_hkl = torch.einsum('aj, bji -> bai', hkl.to(rotation_in_hkl.dtype), rotation_in_hkl) shape = reoriented_hkl.shape reoriented_hkl = reoriented_hkl.reshape(-1, 3) interpolated = interpolate_structure_factor_from_grid(reciprocal_space_grid, reoriented_hkl).reshape(shape[0], shape[1]) diff --git a/torchref/base/reciprocal/symmetry.py b/torchref/base/reciprocal/symmetry.py index bd50b1de..390a13b4 100644 --- a/torchref/base/reciprocal/symmetry.py +++ b/torchref/base/reciprocal/symmetry.py @@ -51,7 +51,7 @@ def _equiv_hkls_to_flat_indices( hi = torch.remainder(all_hkl[:, 0], Nx) ki = torch.remainder(all_hkl[:, 1], Ny) li = torch.remainder(all_hkl[:, 2], Nz) - return (hi * (Ny * Nz) + ki * Nz + li).to(torch.int64) + return (hi * (Ny * Nz) + ki * Nz + li).to(torch.int64) # dtype-ok: flat HKL grid index; int64 avoids overflow, used for indexing class ReciprocalSymmetryExtractor(DeviceMixin): diff --git a/torchref/base/scattering/scattering_table.py b/torchref/base/scattering/scattering_table.py index 8faf3b5d..fbc26aab 100644 --- a/torchref/base/scattering/scattering_table.py +++ b/torchref/base/scattering/scattering_table.py @@ -168,7 +168,7 @@ def get_scattering_params_by_z( table = load_scattering_table(device=device, dtype=dtype) # Long, not the caller's int32: torch indexing requires it. - z_idx = z_tensor.to(device=device, dtype=torch.long) + z_idx = z_tensor.to(device=device, dtype=torch.long) # dtype-ok: z cast to long for scattering-table index lookup; indexing requires long A = table["A"][z_idx] B = table["B"][z_idx] @@ -254,4 +254,4 @@ def elements_to_z(elements: list, normalize: bool = True) -> torch.Tensor: z = element_to_z.get(elem, 0) z_values.append(z) - return torch.tensor(z_values, dtype=torch.int32) + return torch.tensor(z_values, dtype=torch.int32) # dtype-ok: atomic-number Z categorical codes; fixed int32 lookup keys diff --git a/torchref/base/targets/_dispatch.py b/torchref/base/targets/_dispatch.py index 40babcab..0b9562a0 100644 --- a/torchref/base/targets/_dispatch.py +++ b/torchref/base/targets/_dispatch.py @@ -54,7 +54,7 @@ def why_unavailable() -> Optional[str]: name="triton", kernel=None, # gate-only; see the module docstring device="cuda", - dtypes=(torch.float32,), + dtypes=(torch.float32,), # dtype-ok: backend capability declaration, not an allocation probe=(_THIS, "why_unavailable"), expect_available="cuda", # The probe handles availability, so this governs only a kernel that diff --git a/torchref/base/targets/xray_ml_full.py b/torchref/base/targets/xray_ml_full.py index 381ed3e9..be37e6e8 100644 --- a/torchref/base/targets/xray_ml_full.py +++ b/torchref/base/targets/xray_ml_full.py @@ -220,6 +220,7 @@ def acentric_nll(F_obs, sigma, Fc, Sigma, n_quad=None, n_sigma=None, li0=log_i0) if ( li0 is log_i0 + # dtype-ok: validation guard (compile eligibility), not an allocation and F_obs.dtype is not torch.float64 and F_obs.numel() > 1 # 0/1-specialisation would force a 2nd compile and get_compile_targets() diff --git a/torchref/base/wilson_outliers.py b/torchref/base/wilson_outliers.py index a9283676..e1f0d6c7 100644 --- a/torchref/base/wilson_outliers.py +++ b/torchref/base/wilson_outliers.py @@ -520,6 +520,7 @@ def _normal_quantile(p: float) -> float: low, high = -40.0, 10.0 for _ in range(200): mid = 0.5 * (low + high) + # dtype-ok: deliberate float64 for a scalar CDF; extracted via float() value = float(log_normal_cdf(torch.tensor(mid, dtype=torch.float64))) if value < target: low = mid diff --git a/torchref/cli/mtz2map.py b/torchref/cli/mtz2map.py index 3518495e..d02d75ac 100644 --- a/torchref/cli/mtz2map.py +++ b/torchref/cli/mtz2map.py @@ -18,6 +18,7 @@ import numpy as np import torch +from torchref.config import get_float_dtype from torchref.cli._common import ( add_general_args, add_resolution_args, @@ -178,9 +179,9 @@ def main(): f"{d_spacings.max():.2f} - {d_spacings.min():.2f} A") # --- Convert to torch --- - hkl_t = torch.tensor(hkl, dtype=torch.int32, device=device) - amp_t = torch.tensor(amplitudes, dtype=torch.float32, device=device) - phi_t = torch.tensor(phases_deg, dtype=torch.float32, device=device) * (np.pi / 180.0) + hkl_t = torch.tensor(hkl, dtype=torch.int32, device=device) # dtype-ok: hkl Miller indices fed to symmetry expand; fixed int32 crystallographic representation + amp_t = torch.tensor(amplitudes, dtype=get_float_dtype(), device=device) + phi_t = torch.tensor(phases_deg, dtype=get_float_dtype(), device=device) * (np.pi / 180.0) # --- Expand to P1 --- from torchref.symmetry import Cell, SpaceGroup diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index e6c50f87..05f88a8e 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -80,7 +80,7 @@ def build_atom_mask(selection_xyz, real_space_grid, cell, mask_radius, device): inv_frac_matrix=inv_frac, ) - mask = torch.zeros(grid_shape, dtype=torch.int32, device=device) + mask = torch.zeros(grid_shape, dtype=torch.int32, device=device) # dtype-ok: integer solvent-mask accumulator (mask>0); categorical count, not model-precision data mask = add_to_solvent_mask( surrounding_coords, voxel_indices, @@ -229,7 +229,7 @@ def setup_ded_context( from torchref import DatasetCollection from torchref.symmetry.reciprocal_symmetry import expand_hkl - from torchref.config import normalize_device + from torchref.config import get_float_dtype, normalize_device device = normalize_device(device) @@ -304,7 +304,7 @@ def setup_ded_context( ) if dmin is None: dmin = float(d_spacings.min()) - d_spacing = torch.tensor(d_spacings, dtype=torch.float32, device=device) + d_spacing = torch.tensor(d_spacings, dtype=get_float_dtype(), device=device) # P1 expansion and grid sg = SpaceGroup(sg_name, device=device) diff --git a/torchref/experimental/alignment/clashscore.py b/torchref/experimental/alignment/clashscore.py index 2dae5a7c..f413b852 100644 --- a/torchref/experimental/alignment/clashscore.py +++ b/torchref/experimental/alignment/clashscore.py @@ -201,31 +201,31 @@ def _get_valid_transforms( """ # Compute fractionalization matrix using Cell. Pin to CPU because the # rest of this method does CPU-only float64 symmetry math. - cell_obj = Cell(cell, dtype=torch.float64, device="cpu") + cell_obj = Cell(cell, dtype=torch.float64, device="cpu") # dtype-ok: CPU-pinned high-precision symmetry math (see comment above) B = cell_obj.fractional_matrix - centroid_frac = centroid_frac.to(device="cpu", dtype=torch.float64) + centroid_frac = centroid_frac.to(device="cpu", dtype=torch.float64) # dtype-ok: CPU-pinned high-precision symmetry math # Threshold distance for filtering # Two molecules can clash if centroid distance < 2*radius + clash_radius threshold = 2 * molecule_radius + clash_radius # Identity matrix for comparison with rotation matrices - I = torch.eye(3, dtype=torch.float64) + I = torch.eye(3, dtype=torch.float64) # dtype-ok: CPU-pinned high-precision symmetry math valid_transforms = [] n_ops = self.symmetry.n_ops for op_idx in range(n_ops): - R = self.symmetry.matrices[op_idx].cpu().to(torch.float64) - t = self.symmetry.translations[op_idx].cpu().to(torch.float64) + R = self.symmetry.matrices[op_idx].cpu().to(torch.float64) # dtype-ok: CPU-pinned high-precision symmetry math + t = self.symmetry.translations[op_idx].cpu().to(torch.float64) # dtype-ok: CPU-pinned high-precision symmetry math for offset in self._cell_offsets: # Skip identity operation in central cell (self-interaction) if op_idx == 0 and offset == (0, 0, 0): continue - offset_tensor = torch.tensor(offset, dtype=torch.float64) + offset_tensor = torch.tensor(offset, dtype=torch.float64) # dtype-ok: CPU-pinned high-precision symmetry math # Displacement between ASU centroid and this symmetry mate's centroid # Symmetry mate position: R @ x + t + offset diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index c315489b..d802aa2c 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -564,6 +564,8 @@ def _get_e_values_obs( mask = self.data.get_valid_mask() F_obs_masked = F_obs[mask] + # dtype-ok: deliberate float64 for E-value statistics (Wilson stats + # need the precision). Caveat: no .cpu() first, so this errors on MPS. F2 = (F_obs_masked ** 2).to(torch.float64) s = self._get_s_vectors()[mask] @@ -591,6 +593,7 @@ def _get_e_values_calc( F_calc = self.model(hkl).abs() F_calc_masked = F_calc[mask] + # dtype-ok: deliberate float64 for E-value statistics (see _get_e_values_obs) F2 = (F_calc_masked ** 2).to(torch.float64) s = self._get_s_vectors()[mask] @@ -606,6 +609,7 @@ def _get_s_vectors(self) -> torch.Tensor: if self._s_vectors is None: from torchref.base import reciprocal_basis_matrix rec_basis = reciprocal_basis_matrix(self.model.cell) + # dtype-ok: deliberate float64 for reciprocal-vector precision in E-value stats self._s_vectors = self.data.hkl.to(torch.float64) @ rec_basis.to(torch.float64) return self._s_vectors diff --git a/torchref/experimental/alignment/rigid_body.py b/torchref/experimental/alignment/rigid_body.py index e914dc13..9d5350f5 100644 --- a/torchref/experimental/alignment/rigid_body.py +++ b/torchref/experimental/alignment/rigid_body.py @@ -30,7 +30,7 @@ import torch import torch.nn as nn -from torchref.config import get_default_device +from torchref.config import get_default_device, get_float_dtype from torchref.scaling import ScalerBase from torchref.model import SfFFT from torchref.symmetry import spacegroup @@ -137,9 +137,7 @@ def __init__( model, # ModelFT data, # ReflectionData expected_rotational_error: float = 0.1, - initial_rotation: torch.Tensor = torch.tensor( - [0.0, 0.0, 0.0], dtype=torch.float32 - ), + initial_rotation: Optional[torch.Tensor] = None, initial_translation: Optional[torch.Tensor] = None, device: torch.device = None, rfactor_converged_threshold: float = 0.45, @@ -185,17 +183,25 @@ def __init__( self.A_aniso = None self.B_aniso = None - # Store initial rotation + # Store initial rotation. Resolve the default here, not in the signature: + # a tensor default is built once at import, on the config dtype at import + # time, and shared across instances -- both are latent bugs. + if initial_rotation is None: + initial_rotation = torch.zeros(3, dtype=get_float_dtype(), device=device) self.register_buffer( "initial_rotation", initial_rotation.to(device=device).clone() ) - self.rotation_parameters = nn.Parameter(torch.zeros(3, device=device)) + self.rotation_parameters = nn.Parameter( + torch.zeros(3, dtype=get_float_dtype(), device=device) + ) self.expected_rotational_error = expected_rotational_error # Refinable translation (fractional coordinates) if initial_translation is None: - initial_translation = torch.zeros(3, device=device) + initial_translation = torch.zeros( + 3, dtype=get_float_dtype(), device=device + ) else: initial_translation = initial_translation.to(device=device).clone() self.translation_frac = nn.Parameter(initial_translation) diff --git a/torchref/experimental/alignment/transform.py b/torchref/experimental/alignment/transform.py index 876e0e17..33f7eda3 100644 --- a/torchref/experimental/alignment/transform.py +++ b/torchref/experimental/alignment/transform.py @@ -68,9 +68,10 @@ def sample_angles(sampling_pitch_rad, max_angles_rad): angles = [] max_alpha, max_beta, max_gamma = max_angles_rad - alpha = torch.arange(0, max_alpha + 1e-6, sampling_pitch_rad, dtype=torch.float32) - beta = torch.arange(0, max_beta + 1e-6, sampling_pitch_rad, dtype=torch.float32) - gamma = torch.arange(0, max_gamma + 1e-6, sampling_pitch_rad, dtype=torch.float32) + _dtype = get_float_dtype() + alpha = torch.arange(0, max_alpha + 1e-6, sampling_pitch_rad, dtype=_dtype) + beta = torch.arange(0, max_beta + 1e-6, sampling_pitch_rad, dtype=_dtype) + gamma = torch.arange(0, max_gamma + 1e-6, sampling_pitch_rad, dtype=_dtype) alpha, beta, gamma = torch.meshgrid(alpha, beta, gamma, indexing="ij") return torch.stack([alpha.flatten(), beta.flatten(), gamma.flatten()], dim=-1) diff --git a/torchref/experimental/ensemble/ensemble_amber_kl.py b/torchref/experimental/ensemble/ensemble_amber_kl.py index 56a5aca6..ebb5cfc3 100644 --- a/torchref/experimental/ensemble/ensemble_amber_kl.py +++ b/torchref/experimental/ensemble/ensemble_amber_kl.py @@ -136,7 +136,7 @@ def __init__( self.register_buffer( "_member_atom_idx", torch.as_tensor( - atom_idx_np, dtype=torch.long, device=self._model.device + atom_idx_np, dtype=torch.long, device=self._model.device # dtype-ok: atom index tensor for indexing; PyTorch requires int64 ), ) else: diff --git a/torchref/experimental/ensemble/ensemble_model.py b/torchref/experimental/ensemble/ensemble_model.py index 10eb5ecc..7f0ed418 100644 --- a/torchref/experimental/ensemble/ensemble_model.py +++ b/torchref/experimental/ensemble/ensemble_model.py @@ -763,6 +763,8 @@ def enable_low_rank(self, K: int) -> float: with torch.no_grad(): flat = self.xyz().detach() # (N*n_atoms, 3) + # dtype-ok: SVD seeding in float64 for numerical stability. Caveat: no + # .cpu() first, so this errors on MPS. X = flat.reshape(N, n_atoms * 3).to(torch.float64) mu = X.mean(dim=0) # (D,) Xc = X - mu.unsqueeze(0) diff --git a/torchref/experimental/ensemble/pca_model.py b/torchref/experimental/ensemble/pca_model.py index 4b092e47..00896cb0 100644 --- a/torchref/experimental/ensemble/pca_model.py +++ b/torchref/experimental/ensemble/pca_model.py @@ -92,6 +92,8 @@ def from_ensemble( member matrix. ``K`` defaults to ``N-1`` (complete reparameterization).""" N = int(n_members) with torch.no_grad(): + # dtype-ok: SVD seeding in float64 for numerical stability; recast to + # xyz_flat.dtype below (line 107). Caveat: no .cpu(), so errors on MPS. X = xyz_flat.detach().reshape(N, n_atoms * 3).to(torch.float64) mu = X.mean(dim=0) Xc = X - mu.unsqueeze(0) diff --git a/torchref/experimental/ensemble/quasi_crystal_amber.py b/torchref/experimental/ensemble/quasi_crystal_amber.py index eddf1423..22544be4 100644 --- a/torchref/experimental/ensemble/quasi_crystal_amber.py +++ b/torchref/experimental/ensemble/quasi_crystal_amber.py @@ -599,10 +599,10 @@ def _ensure_torch_buffers( # Index pairs (long) for the scatter from model atoms into OMM slots. self._src_model_idx_torch = torch.from_numpy(self._src_model_idx_np).to( - device=device, dtype=torch.long + device=device, dtype=torch.long # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 ) self._dst_omm_idx_torch = torch.from_numpy(self._dst_omm_idx_np).to( - device=device, dtype=torch.long + device=device, dtype=torch.long # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 ) # Index of ensemble-model atoms (in the FULL EnsembleModel layout) @@ -610,7 +610,7 @@ def _ensure_torch_buffers( # subset ``xyz_per_member`` before applying the layout transform. self._keep_atom_idx_torch = torch.from_numpy( self._keep_atom_idx_np - ).to(device=device, dtype=torch.long) + ).to(device=device, dtype=torch.long) # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 # Boolean mask: True where the OMM slot has NO model atom mapped to # it (so we keep the init position there). @@ -621,20 +621,20 @@ def _ensure_torch_buffers( # H-attachment indices tiled per member: template indices live in # [0, n_omm); full-tensor indices live in [0, N · n_omm). member_offset = ( - torch.arange(N, device=device, dtype=torch.long).unsqueeze(1) + torch.arange(N, device=device, dtype=torch.long).unsqueeze(1) # dtype-ok: arange index for broadcasting/indexing; PyTorch requires int64 * n_omm ) # (N, 1) h_idx_t = torch.from_numpy(self._h_idx_template).to( - device=device, dtype=torch.long + device=device, dtype=torch.long # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 ) h_parent_t = torch.from_numpy(self._h_parent_idx_template).to( - device=device, dtype=torch.long + device=device, dtype=torch.long # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 ) h_n1_t = torch.from_numpy(self._h_n1_idx_template).to( - device=device, dtype=torch.long + device=device, dtype=torch.long # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 ) h_n2_t = torch.from_numpy(self._h_n2_idx_template).to( - device=device, dtype=torch.long + device=device, dtype=torch.long # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 ) self._h_idx_tiled = (member_offset + h_idx_t.unsqueeze(0)).reshape(-1) diff --git a/torchref/experimental/ensemble/rank_penalty.py b/torchref/experimental/ensemble/rank_penalty.py index 9254c17e..bf214098 100644 --- a/torchref/experimental/ensemble/rank_penalty.py +++ b/torchref/experimental/ensemble/rank_penalty.py @@ -232,6 +232,8 @@ def spectrum_diagnostics(self) -> dict: variance in the top mode. These show the "purification" as the penalty ramps up. """ + # dtype-ok: float64 svdvals for read-only diagnostics; results extracted + # via float(). Caveat: no .cpu() first, so this errors on MPS. Xc = self._centered().detach().to(torch.float64) s = torch.linalg.svdvals(Xc) # (min(N, D),) s2 = s ** 2 diff --git a/torchref/experimental/ensemble/wilson_prior.py b/torchref/experimental/ensemble/wilson_prior.py index e2fdf7e8..44490f9c 100644 --- a/torchref/experimental/ensemble/wilson_prior.py +++ b/torchref/experimental/ensemble/wilson_prior.py @@ -172,7 +172,7 @@ def _build_bin_assignment(self) -> None: order = torch.argsort(res) n = res.numel() nbins = min(self.nbins, max(1, n // 50)) - bin_assign = torch.empty(n, dtype=torch.long, device=res.device) + bin_assign = torch.empty(n, dtype=torch.long, device=res.device) # dtype-ok: bin-assignment tensor used as scatter_add index; PyTorch requires int64 edges = torch.linspace(0, n, nbins + 1, device=res.device).round().long() for b in range(nbins): start = int(edges[b].item()) diff --git a/torchref/experimental/kinetic/occupancies.py b/torchref/experimental/kinetic/occupancies.py index 8d4aa720..c045bd89 100644 --- a/torchref/experimental/kinetic/occupancies.py +++ b/torchref/experimental/kinetic/occupancies.py @@ -18,6 +18,7 @@ from typing import Dict, List, Optional, Union, Tuple import numpy as np +from torchref.config import get_float_dtype from torchref.utils.device_mixin import DeviceMixin @@ -163,7 +164,7 @@ def __init__( # Convert time to tensor if needed if not isinstance(time, torch.Tensor): - time = torch.tensor(time, dtype=torch.float32) + time = torch.tensor(time, dtype=get_float_dtype()) self.register_buffer('time', time) # Initialize the kinetic model diff --git a/torchref/experimental/monolithic_refinement/density_scaler.py b/torchref/experimental/monolithic_refinement/density_scaler.py index b0322c3d..afe76f95 100644 --- a/torchref/experimental/monolithic_refinement/density_scaler.py +++ b/torchref/experimental/monolithic_refinement/density_scaler.py @@ -127,7 +127,7 @@ def get_rec_solvent(self, hkl): Not detached: ``F_sol`` follows the moving atoms so gradients reach ``xyz``/``adp``. The scaler applies the contrast and falloff on top. """ - return self.density(hkl.to(torch.long)) + return self.density(hkl.to(torch.long)) # dtype-ok: hkl cast to long for density lookup indexing; PyTorch requires int64 def update_solvent(self): """No-op: the density mask is rebuilt live on every scaler forward.""" diff --git a/torchref/experimental/targets/amber_target.py b/torchref/experimental/targets/amber_target.py index bff215bc..404cfcc3 100644 --- a/torchref/experimental/targets/amber_target.py +++ b/torchref/experimental/targets/amber_target.py @@ -1470,14 +1470,14 @@ def _place_hydrogens(self, heavy_omm_xyz_nm: torch.Tensor) -> torch.Tensor: or getattr(self, "_h_tensors_dtype", None) != dtype ): self._h_parent_idx_t = torch.as_tensor( - self._h_parent_idx, dtype=torch.long, device=device, + self._h_parent_idx, dtype=torch.long, device=device, # dtype-ok: H-parent atom index for indexing; PyTorch requires int64 ) # For invalid frames clamp neighbor indices to 0 so the gather is # safe; the value is masked out by `where` below. n1 = np.where(self._h_n1_idx >= 0, self._h_n1_idx, 0) n2 = np.where(self._h_n2_idx >= 0, self._h_n2_idx, 0) - self._h_n1_idx_t = torch.as_tensor(n1, dtype=torch.long, device=device) - self._h_n2_idx_t = torch.as_tensor(n2, dtype=torch.long, device=device) + self._h_n1_idx_t = torch.as_tensor(n1, dtype=torch.long, device=device) # dtype-ok: neighbor atom index for indexing; PyTorch requires int64 + self._h_n2_idx_t = torch.as_tensor(n2, dtype=torch.long, device=device) # dtype-ok: neighbor atom index for indexing; PyTorch requires int64 self._h_local_pos_t = torch.as_tensor( self._h_local_pos, dtype=dtype, device=device, ) @@ -1530,10 +1530,10 @@ def _compose_full_omm_xyz( valid_np, dtype=torch.bool, device=device, ) self._model_valid_model_idx_t = torch.as_tensor( - np.where(valid_np)[0], dtype=torch.long, device=device, + np.where(valid_np)[0], dtype=torch.long, device=device, # dtype-ok: valid-atom index (np.where) for indexing; PyTorch requires int64 ) self._model_valid_omm_idx_t = torch.as_tensor( - self._model_to_omm[valid_np], dtype=torch.long, device=device, + self._model_to_omm[valid_np], dtype=torch.long, device=device, # dtype-ok: model->OMM mapping index for indexing; PyTorch requires int64 ) # Construction-time snapshot for unmatched heavy slots + initial Hs self._pos_buf_t = torch.as_tensor( @@ -1558,7 +1558,7 @@ def _compose_full_omm_xyz( h_xyz = self._place_hydrogens(full) if not hasattr(self, "_h_idx_t_for_omm"): self._h_idx_t_for_omm = torch.as_tensor( - self._h_idx, dtype=torch.long, device=device, + self._h_idx, dtype=torch.long, device=device, # dtype-ok: H-atom index for indexing; PyTorch requires int64 ) elif self._h_idx_t_for_omm.device != device: self._h_idx_t_for_omm = self._h_idx_t_for_omm.to(device) diff --git a/torchref/experimental/targets/forcefield_target.py b/torchref/experimental/targets/forcefield_target.py index 27332457..09f0e1de 100644 --- a/torchref/experimental/targets/forcefield_target.py +++ b/torchref/experimental/targets/forcefield_target.py @@ -178,11 +178,11 @@ def forward(self) -> torch.Tensor: Z = self.model.Z # Shape: (n_atoms,) # Ensure Z is long tensor - if Z.dtype != torch.long: + if Z.dtype != torch.long: # dtype-ok: dtype guard comparison against torch.long, not an allocation Z = Z.long() # Create batch tensor (single structure = all zeros) - batch = torch.zeros(len(Z), dtype=torch.long, device=xyz.device) + batch = torch.zeros(len(Z), dtype=torch.long, device=xyz.device) # dtype-ok: batch index tensor for TorchMD-Net graph scatter; PyTorch requires int64 # Compute energy via TorchMD-Net # Returns (energy, forces) or just energy depending on model config diff --git a/torchref/experimental/targets/sampled_ml_phase_target.py b/torchref/experimental/targets/sampled_ml_phase_target.py index 59575c8d..e505240b 100644 --- a/torchref/experimental/targets/sampled_ml_phase_target.py +++ b/torchref/experimental/targets/sampled_ml_phase_target.py @@ -129,7 +129,7 @@ def __init__( self.name = "xray_sampled_ml_work" if use_work_set else "xray_sampled_ml_test" # Register tunable parameters as buffers for state_dict access - self.register_buffer("_n_samples", torch.tensor(n_samples, dtype=torch.int64)) + self.register_buffer("_n_samples", torch.tensor(n_samples, dtype=torch.int64)) # dtype-ok: scalar sample-count buffer; categorical count, not model-precision data self.register_buffer("_sigma_model_log", torch.tensor(sigma_model_log)) self.register_buffer("_use_analytical", torch.tensor(use_analytical)) self.register_buffer("_use_antithetic", torch.tensor(use_antithetic)) @@ -545,7 +545,7 @@ def __init__( self.add_module("_scaler_dark", scaler_dark) # Tunable parameters as buffers - self.register_buffer("_n_samples", torch.tensor(n_samples, dtype=torch.int64)) + self.register_buffer("_n_samples", torch.tensor(n_samples, dtype=torch.int64)) # dtype-ok: scalar sample-count buffer; categorical count, not model-precision data self.register_buffer("_sigma_model_log", torch.tensor(sigma_model_log)) self.use_work_set = use_work_set diff --git a/torchref/io/datasets/base.py b/torchref/io/datasets/base.py index 668290f3..705dbde6 100644 --- a/torchref/io/datasets/base.py +++ b/torchref/io/datasets/base.py @@ -14,7 +14,7 @@ import gemmi import torch -from torchref.config import get_default_device, normalize_device +from torchref.config import get_default_device, get_float_dtype, normalize_device from torchref.symmetry import Cell from torchref.utils.device_mixin import DeviceMovementMixin @@ -186,7 +186,12 @@ def _from_state(cls, state: Dict[str, Any], device=None) -> "CrystalDataset": # Spacegroup stays a string here; subclasses that want an object rewrap. if "cell" in state and state["cell"] is not None: if isinstance(state["cell"], torch.Tensor): - state["cell"] = Cell(state["cell"], dtype=torch.float32, device=device) + # Conform the reloaded cell to the config float dtype rather than + # pinning float32: a dataset saved and reloaded under a float64 + # config otherwise carries a float32 cell into reciprocal-basis math. + state["cell"] = Cell( + state["cell"], dtype=get_float_dtype(), device=device + ) obj = cls(**state) diff --git a/torchref/io/datasets/fcalc_data.py b/torchref/io/datasets/fcalc_data.py index 616c310d..fefca36b 100644 --- a/torchref/io/datasets/fcalc_data.py +++ b/torchref/io/datasets/fcalc_data.py @@ -129,7 +129,7 @@ def from_cell_and_resolution( # make_miller_array returns unique HKL for the asymmetric unit only. hkl_list = gemmi.make_miller_array(gemmi_cell, gemmi_sg, d_min) - hkl = torch.tensor(hkl_list, dtype=torch.int32, device=device) + hkl = torch.tensor(hkl_list, dtype=torch.int32, device=device) # dtype-ok: hkl Miller indices; fixed int32 crystallographic representation, not model-precision data resolution = get_d_spacing(hkl.float(), cell_tensor) diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index 5a335989..739a64b4 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -301,7 +301,7 @@ def _subset_indices(self, kind: str) -> torch.Tensor: n = 0 if self.hkl is None else len(self.hkl) device = self.device if n == 0: - empty = torch.empty(0, dtype=torch.long, device=device) + empty = torch.empty(0, dtype=torch.long, device=device) # dtype-ok: empty index tensor; PyTorch requires int64 for indexing self._subset_cache = { "work": empty, "free": empty, @@ -432,7 +432,7 @@ def _reindex_per_reflection( n_src = len(self.hkl) if self.hkl is not None else 0 new_hkl = new_hkl.to(dtype=dtypes.int, device=self.device) n_out = len(new_hkl) - index_map = index_map.to(device=self.device, dtype=torch.long) + index_map = index_map.to(device=self.device, dtype=torch.long) # dtype-ok: index map used for indexing/gather; PyTorch requires int64 present = index_map >= 0 src_idx = index_map[present] @@ -709,8 +709,10 @@ def _group_any( Uses ``index_add_`` on float rather than ``scatter_reduce_(amax)``: the latter raises "not supported for torch.int64" on the MPS backend. """ + # dtype-ok: float32 count accumulator for the int64-scatter MPS workaround + # above; the result is reduced to bool (> 0), so precision is irrelevant. counts = torch.zeros(n_groups, dtype=torch.float32, device=mask.device) - counts.index_add_(0, group_id, mask.to(torch.float32)) + counts.index_add_(0, group_id, mask.to(torch.float32)) # dtype-ok: float32 counter for the MPS workaround above; reduced to bool return counts > 0 @staticmethod @@ -1326,13 +1328,13 @@ def mean_res_per_bin(self) -> torch.Tensor: mean_resolutions = torch.scatter_add( mean_resolutions, 0, - self.bin_indices[mask].to(torch.int64), + self.bin_indices[mask].to(torch.int64), # dtype-ok: bin indices for scatter_add/index; PyTorch requires int64 self.resolution[mask], ) count_per_bin = torch.scatter_add( count_per_bin, 0, - self.bin_indices[mask].to(torch.int64), + self.bin_indices[mask].to(torch.int64), # dtype-ok: bin indices for scatter_add/index; PyTorch requires int64 torch.ones_like(self.resolution[mask], dtype=dtypes.int), ) mean_resolutions = mean_resolutions / count_per_bin.clamp(min=1).float() @@ -1361,12 +1363,12 @@ def mean_F_per_bin(self) -> torch.Tensor: count_per_bin = torch.zeros(self._n_bins, dtype=dtypes.int, device=self.device) mask = self.masks() mean_F = torch.scatter_add( - mean_F, 0, self.bin_indices[mask].to(torch.int64), self.F[mask] + mean_F, 0, self.bin_indices[mask].to(torch.int64), self.F[mask] # dtype-ok: bin indices for scatter_add index arg; PyTorch requires int64 ) count_per_bin = torch.scatter_add( count_per_bin, 0, - self.bin_indices[mask].to(torch.int64), + self.bin_indices[mask].to(torch.int64), # dtype-ok: bin indices for scatter_add index arg; PyTorch requires int64 torch.ones_like(self.F[mask], dtype=dtypes.int), ) mean_F = mean_F / count_per_bin.clamp(min=1).float() @@ -1395,12 +1397,12 @@ def mean_sigma_per_bin(self) -> Optional[torch.Tensor]: count_per_bin = torch.zeros(self._n_bins, dtype=dtypes.int, device=self.device) mask = self.masks() mean_sigma = torch.scatter_add( - mean_sigma, 0, self.bin_indices[mask].to(torch.int64), self.F_sigma[mask] + mean_sigma, 0, self.bin_indices[mask].to(torch.int64), self.F_sigma[mask] # dtype-ok: bin indices for scatter_add index arg; PyTorch requires int64 ) count_per_bin = torch.scatter_add( count_per_bin, 0, - self.bin_indices[mask].to(torch.int64), + self.bin_indices[mask].to(torch.int64), # dtype-ok: bin indices for scatter_add index arg; PyTorch requires int64 torch.ones_like(self.F_sigma[mask], dtype=dtypes.int), ) mean_sigma = mean_sigma / count_per_bin.clamp(min=1).float() @@ -2642,8 +2644,8 @@ def _build_anomalous_dataframe( # The (+) member is the unconjugated row, (-) is the Friedel-flagged row. arange = torch.arange(N) - plus_idx = torch.full((M,), -1, dtype=torch.long) - minus_idx = torch.full((M,), -1, dtype=torch.long) + plus_idx = torch.full((M,), -1, dtype=torch.long) # dtype-ok: Friedel-mate index map (-1 sentinel) for indexing; PyTorch requires int64 + minus_idx = torch.full((M,), -1, dtype=torch.long) # dtype-ok: Friedel-mate index map (-1 sentinel) for indexing; PyTorch requires int64 # A Bijvoet mate only counts as present if it is a real, positive # observation. Stacked anomalous input (rs.stack_anomalous) carries a # row for every *absent* mate with a NaN intensity, which French-Wilson diff --git a/torchref/model/disorder_field.py b/torchref/model/disorder_field.py index c11d91fc..a2911d3f 100644 --- a/torchref/model/disorder_field.py +++ b/torchref/model/disorder_field.py @@ -23,7 +23,7 @@ import torch from torch import nn -from torchref.config import get_float_dtype, normalize_device +from torchref.config import get_float_dtype, get_int_dtype, normalize_device from torchref.model.parameter_wrappers import ( MixedTensor, chol_param_count, @@ -88,7 +88,7 @@ def farthest_point_anchors(xyz: torch.Tensor, n_nodes: int) -> torch.Tensor: chosen.append(nxt) d2_nearest = torch.minimum(d2_nearest, ((xyz - xyz[nxt]) ** 2).sum(-1)) - anchors = torch.tensor(chosen, dtype=torch.int64, device=xyz.device) + anchors = torch.tensor(chosen, dtype=torch.int64, device=xyz.device) # dtype-ok: anchor atom indices; torch indexing requires int64 # Lloyd relaxation, snapping to real atoms so an anchor is always an atom index. for _ in range(10): @@ -126,7 +126,7 @@ def density_anchor_rows(xyz: torch.Tensor, n_nodes: int): """ seeds = farthest_point_anchors(xyz, n_nodes) assign = torch.cdist(xyz, xyz[seeds]).argmin(dim=1) - atom_idx = torch.arange(xyz.shape[0], dtype=torch.int64, device=xyz.device) + atom_idx = torch.arange(xyz.shape[0], dtype=torch.int64, device=xyz.device) # dtype-ok: arange atom indices; index requires int64 # A seed whose cluster somehow came out empty still needs a position. present = torch.bincount(assign, minlength=seeds.shape[0]) > 0 @@ -685,9 +685,15 @@ def __init__( self.register_buffer("neighbor_list", None) self.register_buffer("anchor_atom", None) self.register_buffer("anchor_node", None) + # device=device so the empty path lands on the requested device, + # matching the populated path below (it used to omit it and land on CPU). self.register_buffer( "payload_code", - torch.tensor(payload_code(self._payload), dtype=torch.int64), + torch.tensor( + payload_code(self._payload), + dtype=get_int_dtype(), + device=device, + ), ) return @@ -709,12 +715,12 @@ def __init__( if anchor_rows is None: anchor_atom = farthest_point_anchors(xyz, n_nodes) anchor_node = torch.arange( - anchor_atom.shape[0], dtype=torch.int64, device=device + anchor_atom.shape[0], dtype=torch.int64, device=device # dtype-ok: arange anchor indices; index requires int64 ) else: anchor_atom, anchor_node = anchor_rows - anchor_atom = anchor_atom.to(device=device, dtype=torch.int64) - anchor_node = anchor_node.to(device=device, dtype=torch.int64) + anchor_atom = anchor_atom.to(device=device, dtype=torch.int64) # dtype-ok: anchor_atom indices cast; index requires int64 + anchor_node = anchor_node.to(device=device, dtype=torch.int64) # dtype-ok: anchor_node indices cast; index requires int64 n_k = int(anchor_node.max()) + 1 node_pos = self._segment_mean(xyz, anchor_atom, anchor_node, n_k) @@ -756,7 +762,9 @@ def __init__( # inferring it from the storage width. self.register_buffer( "payload_code", - torch.tensor(payload_code(self._payload), dtype=torch.int64, device=device), + torch.tensor( + payload_code(self._payload), dtype=get_int_dtype(), device=device + ), ) # ------------------------------------------------------------------ diff --git a/torchref/model/model.py b/torchref/model/model.py index 6ece8612..7983ef68 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -315,7 +315,7 @@ def _build_z_tensor(self) -> torch.Tensor: for elem in self.pdb["element"] ] self.register_buffer( - "_Z", torch.tensor(z_values, dtype=torch.int32, device=self.device) + "_Z", torch.tensor(z_values, dtype=torch.int32, device=self.device) # dtype-ok: atomic-number Z categorical codes buffer; fixed int32 lookup keys ) return self._Z @@ -823,7 +823,7 @@ def _create_occupancy_groups(self, pdb_df, initial_occ): altloc_groups = [] refinable_mask = torch.zeros(n_atoms, dtype=torch.bool) - sharing_groups_tensor = torch.arange(n_atoms, dtype=torch.long) + sharing_groups_tensor = torch.arange(n_atoms, dtype=torch.long) # dtype-ok: arange atom indices (sharing groups); index requires long collapsed_idx = 0 # First pass: altlocs. ALL atoms of one conformation must share a collapsed @@ -892,7 +892,7 @@ def _create_occupancy_groups(self, pdb_df, initial_occ): # Compact to contiguous indices 0..n_collapsed-1. unique_indices = torch.unique(sharing_groups_tensor, sorted=True) - index_map = torch.zeros(n_atoms, dtype=torch.long) + index_map = torch.zeros(n_atoms, dtype=torch.long) # dtype-ok: index_map atom-index remap; indexing requires long for new_idx, old_idx in enumerate(unique_indices): mask = sharing_groups_tensor == old_idx sharing_groups_tensor[mask] = new_idx @@ -1904,7 +1904,7 @@ def register_alternative_conformations(self): for altloc in unique_altlocs: altloc_atoms = group[group["altloc"] == altloc] indices = torch.tensor( - altloc_atoms["index"].tolist(), dtype=torch.long + altloc_atoms["index"].tolist(), dtype=torch.long # dtype-ok: altloc atom indices; indexing requires long ) conformation_tensors.append(indices) diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index d9b06cf9..c6917cba 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -116,7 +116,9 @@ def __init__( # so it round-trips through state_dict and follows .to(device). f' is always # applied when wavelength is set; f'' only when this is True (unmerged data). self.register_buffer( - "anomalous_bijvoet", torch.tensor(bool(apply_bijvoet)), persistent=True + "anomalous_bijvoet", + torch.tensor(bool(apply_bijvoet), device=self.device), + persistent=True, ) self._anomalous_cache = None # Will hold (mask, f_prime, f_double_prime) self._anomalous_elements_hash = ( diff --git a/torchref/model/parameter_wrappers.py b/torchref/model/parameter_wrappers.py index 47095342..e0ad7e2d 100644 --- a/torchref/model/parameter_wrappers.py +++ b/torchref/model/parameter_wrappers.py @@ -1462,11 +1462,11 @@ def _setup_sharing_groups_and_expansion( # Use sharing_groups directly as the expansion mask if sharing_groups is None: # No sharing - each atom maps to its own index - expansion_mask = torch.arange(n_atoms, dtype=torch.long, device=device) + expansion_mask = torch.arange(n_atoms, dtype=torch.long, device=device) # dtype-ok: arange expansion_mask atom indices; index requires long self._collapsed_shape = n_atoms else: # Use the provided index tensor - expansion_mask = sharing_groups.to(device=device, dtype=torch.long) + expansion_mask = sharing_groups.to(device=device, dtype=torch.long) # dtype-ok: expansion_mask atom/group indices for scatter; requires long self._collapsed_shape = expansion_mask.max().item() + 1 self.register_buffer("expansion_mask", expansion_mask) @@ -1488,10 +1488,10 @@ def _setup_sharing_groups_and_expansion( for conf_atoms in conf_groups: if isinstance(conf_atoms, (list, tuple)): conf_atoms = torch.tensor( - conf_atoms, dtype=torch.long, device=device + conf_atoms, dtype=torch.long, device=device # dtype-ok: conf_atoms atom indices; indexing requires long ) else: - conf_atoms = conf_atoms.to(device=device, dtype=torch.long) + conf_atoms = conf_atoms.to(device=device, dtype=torch.long) # dtype-ok: conf_atoms atom indices cast; indexing requires long # Get collapsed index for first atom collapsed_idx = expansion_mask[conf_atoms[0]].item() @@ -1519,7 +1519,7 @@ def _setup_sharing_groups_and_expansion( # Store as dictionary with keys like 'linked_occ_2', 'linked_occ_3', etc. for n_conf, groups in linked_occupancies.items(): # Shape: (N_groups, n_conf) - tensor = torch.tensor(groups, dtype=torch.long, device=device) + tensor = torch.tensor(groups, dtype=torch.long, device=device) # dtype-ok: linked-occupancy group index buffer; indexing requires long self.register_buffer(f"linked_occ_{n_conf}", tensor) # Store which sizes we have @@ -1527,7 +1527,7 @@ def _setup_sharing_groups_and_expansion( # Create count buffer for vectorized collapse operations # counts[i] = number of atoms that map to collapsed index i - counts = torch.zeros(self._collapsed_shape, dtype=torch.long, device=device) + counts = torch.zeros(self._collapsed_shape, dtype=torch.long, device=device) # dtype-ok: count accumulator; scatter_add source is long ones, dtype must match counts.scatter_add_(0, expansion_mask, torch.ones_like(expansion_mask)) self.register_buffer("collapse_counts", counts) @@ -2005,7 +2005,7 @@ def from_residue_groups( grouped = pdb_dataframe.groupby(["resname", "resseq", "chainid", "altloc"]) n_atoms = len(initial_values) - sharing_groups_tensor = torch.arange(n_atoms, dtype=torch.long) + sharing_groups_tensor = torch.arange(n_atoms, dtype=torch.long) # dtype-ok: arange atom indices (sharing groups); index requires long # Singletons keep their arange ids (0..n_atoms-1); start multi-atom # group ids past that range so a group id can never collide with a # singleton's leftover arange id (the torch.unique compaction below diff --git a/torchref/model/rigid_xyz.py b/torchref/model/rigid_xyz.py index d74c8fe9..4d4c0b5f 100644 --- a/torchref/model/rigid_xyz.py +++ b/torchref/model/rigid_xyz.py @@ -64,7 +64,7 @@ def __init__( dtype = dtype if dtype is not None else get_float_dtype() self.register_buffer("original_xyz", torch.empty(0, 3, device=device, dtype=dtype)) self.register_buffer( - "chain_indices", torch.empty(0, dtype=torch.long, device=device) + "chain_indices", torch.empty(0, dtype=torch.long, device=device) # dtype-ok: empty chain_indices buffer; indexing requires long ) self.register_buffer("chain_centers", torch.empty(0, 3, device=device, dtype=dtype)) self.register_buffer( diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index e8d61ae8..6c758e58 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -426,7 +426,7 @@ def mark(idx): return # 4. freeze xyz of those atoms (same path as freeze_selection) - model.xyz_mask[torch.tensor(freeze_idx, dtype=torch.long)] = False + model.xyz_mask[torch.tensor(freeze_idx, dtype=torch.long)] = False # dtype-ok: freeze index used to index xyz_mask; PyTorch requires int64 model.apply_mask_to_parameter("xyz") if self.verbose > 0: shown = frozen_res[:20] + (["..."] if len(frozen_res) > 20 else []) diff --git a/torchref/refinement/model_error_estimation/sigma_a.py b/torchref/refinement/model_error_estimation/sigma_a.py index 6c96182c..a24f9285 100644 --- a/torchref/refinement/model_error_estimation/sigma_a.py +++ b/torchref/refinement/model_error_estimation/sigma_a.py @@ -263,10 +263,10 @@ def _segment_layout(lengths: Tuple[int, ...], device_str: str): ``lengths`` is a tuple so it can be a cache key. """ device = torch.device(device_str) - L = torch.tensor(lengths, dtype=torch.long, device=device) + L = torch.tensor(lengths, dtype=torch.long, device=device) # dtype-ok: segment lengths for cumsum offsets/gather index; PyTorch requires int64 total = int(L.sum()) max_len = int(L.max()) if L.numel() else 0 - zero = torch.zeros(1, dtype=torch.long, device=device) + zero = torch.zeros(1, dtype=torch.long, device=device) # dtype-ok: zero offset concatenated into gather index; PyTorch requires int64 starts = torch.cat([zero, L.cumsum(0)[:-1]]) ar = torch.arange(max_len, device=device).reshape(1, max_len) # Clamp keeps the gather in bounds for the padding slots; `mask` zeroes them anyway. @@ -548,7 +548,7 @@ def estimate_beta( out_dtype = F_obs.dtype dtype = torch.promote_types(get_float_dtype(), out_dtype) - if dtype == torch.float64 and device.type == "mps": + if dtype == torch.float64 and device.type == "mps": # dtype-ok: MPS capability guard, not an allocation raise RuntimeError( "MPS has no float64; set the defaults float dtype to float32 or use CPU" ) diff --git a/torchref/refinement/optimizers/curvature.py b/torchref/refinement/optimizers/curvature.py index b3192e19..74b5bbae 100644 --- a/torchref/refinement/optimizers/curvature.py +++ b/torchref/refinement/optimizers/curvature.py @@ -36,7 +36,7 @@ def _sample_probe( """Draw one Hutchinson probe vector of length ``numel``.""" if probe == "rademacher": r = torch.randint( - 0, 2, (numel,), generator=generator, device=device, dtype=torch.int64 + 0, 2, (numel,), generator=generator, device=device, dtype=torch.int64 # dtype-ok: randint {0,1} bernoulli draw, immediately cast to float dtype; width irrelevant ) return r.to(dtype).mul_(2.0).sub_(1.0) # {0,1} -> {-1,+1} if probe == "gaussian": diff --git a/torchref/refinement/targets/adp/rigid_bond.py b/torchref/refinement/targets/adp/rigid_bond.py index c8b69dd8..07a639a7 100644 --- a/torchref/refinement/targets/adp/rigid_bond.py +++ b/torchref/refinement/targets/adp/rigid_bond.py @@ -122,7 +122,7 @@ def _bond_pairs(self) -> torch.Tensor: chunks.append(idx_) if chunks: return torch.cat(chunks, dim=0).contiguous() - return torch.empty(0, 2, dtype=torch.long, device=self.model.xyz().device) + return torch.empty(0, 2, dtype=torch.long, device=self.model.xyz().device) # dtype-ok: empty (0,2) atom-pair index tensor; PyTorch requires int64 def _compute_aniso_rigid_bond(self) -> torch.Tensor: """Rigid-bond NLL from ``Δz = l^T U_1 l - l^T U_2 l`` along each bond. diff --git a/torchref/refinement/targets/adp/sigd.py b/torchref/refinement/targets/adp/sigd.py index aaf4cb2f..95aace4d 100644 --- a/torchref/refinement/targets/adp/sigd.py +++ b/torchref/refinement/targets/adp/sigd.py @@ -118,6 +118,7 @@ def stats(self) -> Dict[str, any]: beta = float((adp - self._b_shift).clamp(min=1e-3).mean()) * (alpha - 1.0) # std(log B) = sqrt(trigamma(alpha)); torch.polygamma(1, .) is trigamma. implied_std = math.sqrt( + # dtype-ok: deliberate float64 for a scalar polygamma; extracted via float() float(torch.polygamma(1, torch.tensor(alpha, dtype=torch.float64))) ) diff --git a/torchref/refinement/targets/adp/similarity.py b/torchref/refinement/targets/adp/similarity.py index d8619131..191d2379 100644 --- a/torchref/refinement/targets/adp/similarity.py +++ b/torchref/refinement/targets/adp/similarity.py @@ -94,7 +94,7 @@ def _get_pair_indices(self) -> torch.Tensor: if chunks: cached = torch.cat(chunks, dim=0).contiguous() else: - cached = torch.empty(0, 2, dtype=torch.long, + cached = torch.empty(0, 2, dtype=torch.long, # dtype-ok: empty (0,2) atom-pair index tensor; PyTorch requires int64 device=self.model.xyz().device) self._simu_pair_indices_cache = cached return cached diff --git a/torchref/refinement/targets/difference.py b/torchref/refinement/targets/difference.py index 33870325..d0a75ab7 100644 --- a/torchref/refinement/targets/difference.py +++ b/torchref/refinement/targets/difference.py @@ -204,10 +204,10 @@ def _match_reflections(self): device = hkl_light.device self._matched_indices_light = torch.tensor( - matched_light, dtype=torch.long, device=device + matched_light, dtype=torch.long, device=device # dtype-ok: matched atom indices used for indexing; PyTorch requires int64 ) self._matched_indices_dark = torch.tensor( - matched_dark, dtype=torch.long, device=device + matched_dark, dtype=torch.long, device=device # dtype-ok: matched atom indices used for indexing; PyTorch requires int64 ) # Store common HKL (using light indices, they should be identical) diff --git a/torchref/refinement/targets/geometry/chiral.py b/torchref/refinement/targets/geometry/chiral.py index e10f567a..1cc6dcee 100644 --- a/torchref/refinement/targets/geometry/chiral.py +++ b/torchref/refinement/targets/geometry/chiral.py @@ -82,7 +82,7 @@ def get_violations(self, threshold: float = 0.5) -> Dict[str, torch.Tensor]: if "chiral" not in self.restraints.restraints: return { - "indices": torch.tensor([], dtype=torch.long, device=device).reshape( + "indices": torch.tensor([], dtype=torch.long, device=device).reshape( # dtype-ok: empty restraint index tensor; PyTorch requires int64 for indexing 0, 4 ), "volumes": torch.tensor([], device=device), diff --git a/torchref/refinement/targets/geometry/non_bonded.py b/torchref/refinement/targets/geometry/non_bonded.py index ae76b196..637be377 100644 --- a/torchref/refinement/targets/geometry/non_bonded.py +++ b/torchref/refinement/targets/geometry/non_bonded.py @@ -346,7 +346,7 @@ def get_violations(self, threshold: float = 0.0) -> Dict[str, torch.Tensor]: if "vdw" not in self.restraints.restraints: return { - "indices": torch.tensor([], dtype=torch.long, device=device).reshape( + "indices": torch.tensor([], dtype=torch.long, device=device).reshape( # dtype-ok: empty restraint index tensor; PyTorch requires int64 for indexing 0, 2 ), "violations": torch.tensor([], device=device), @@ -359,7 +359,7 @@ def get_violations(self, threshold: float = 0.0) -> Dict[str, torch.Tensor]: if indices is None or len(indices) == 0: return { - "indices": torch.tensor([], dtype=torch.long, device=device).reshape( + "indices": torch.tensor([], dtype=torch.long, device=device).reshape( # dtype-ok: empty restraint index tensor; PyTorch requires int64 for indexing 0, 2 ), "violations": torch.tensor([], device=device), diff --git a/torchref/refinement/targets/similarity.py b/torchref/refinement/targets/similarity.py index 526ebc21..d01c3ef8 100644 --- a/torchref/refinement/targets/similarity.py +++ b/torchref/refinement/targets/similarity.py @@ -68,10 +68,10 @@ def __init__( # path (the one ``load_state_dict`` uses) would have no such buffers at all. # ``_build_atom_map`` overwrites them rather than creating them. self.register_buffer( - "_idx_dark", torch.zeros(0, dtype=torch.long, device=self.device) + "_idx_dark", torch.zeros(0, dtype=torch.long, device=self.device) # dtype-ok: index buffer for gather/index_select; PyTorch requires int64 ) self.register_buffer( - "_idx_light", torch.zeros(0, dtype=torch.long, device=self.device) + "_idx_light", torch.zeros(0, dtype=torch.long, device=self.device) # dtype-ok: index buffer for gather/index_select; PyTorch requires int64 ) if model_dark is not None and model_light is not None: self._build_atom_map() @@ -140,10 +140,10 @@ def _build_atom_map(self): "dark and light models" ) self.register_buffer( - "_idx_dark", torch.zeros(0, dtype=torch.long, device=self.device) + "_idx_dark", torch.zeros(0, dtype=torch.long, device=self.device) # dtype-ok: index buffer for gather/index_select; PyTorch requires int64 ) self.register_buffer( - "_idx_light", torch.zeros(0, dtype=torch.long, device=self.device) + "_idx_light", torch.zeros(0, dtype=torch.long, device=self.device) # dtype-ok: index buffer for gather/index_select; PyTorch requires int64 ) return @@ -166,13 +166,13 @@ def _build_atom_map(self): self.register_buffer( "_idx_dark", torch.tensor( - merged["_idx_dark"].values, dtype=torch.long, device=self.device + merged["_idx_dark"].values, dtype=torch.long, device=self.device # dtype-ok: atom index tensor used for indexing; PyTorch requires int64 ), ) self.register_buffer( "_idx_light", torch.tensor( - merged["_idx_light"].values, dtype=torch.long, device=self.device + merged["_idx_light"].values, dtype=torch.long, device=self.device # dtype-ok: atom index tensor used for indexing; PyTorch requires int64 ), ) diff --git a/torchref/scaling/collection_scaler.py b/torchref/scaling/collection_scaler.py index 585b4cb4..a7a8aced 100644 --- a/torchref/scaling/collection_scaler.py +++ b/torchref/scaling/collection_scaler.py @@ -152,7 +152,7 @@ def _calc_initial_scale_joint(self): pos_mask = torch.ones_like(fobs, dtype=torch.bool) mask = (work_mask & pos_mask).to(torch.bool) - bins = self.bins[mask].to(torch.int64) + bins = self.bins[mask].to(torch.int64) # dtype-ok: bin indices for scatter/index_select; PyTorch requires int64 log_ratios = ( torch.log(fobs_clamped[mask]) - torch.log(fcalc_amp[mask]) ).to(self.device) @@ -165,7 +165,7 @@ def _calc_initial_scale_joint(self): per_bin = scales / (counts + 1e-6) with torch.no_grad(): - target = per_bin.detach()[self.bins.to(torch.int64)] + target = per_bin.detach()[self.bins.to(torch.int64)] # dtype-ok: bin indices for advanced indexing; PyTorch requires int64 design = self._iso_design.to(target.dtype) coeff = torch.linalg.lstsq(design, target.unsqueeze(1)).solution.squeeze(1) self.c_iso = nn.Parameter(coeff.detach()) diff --git a/torchref/scaling/scaler_base.py b/torchref/scaling/scaler_base.py index 9fd9119f..9b588fb4 100644 --- a/torchref/scaling/scaler_base.py +++ b/torchref/scaling/scaler_base.py @@ -266,7 +266,7 @@ def calc_initial_scale(self, fcalc: torch.Tensor): initial_log_scale.detach().cpu().numpy(), ) with torch.no_grad(): - target = initial_log_scale.detach().to(self.device)[self.bins.to(torch.int64)] + target = initial_log_scale.detach().to(self.device)[self.bins.to(torch.int64)] # dtype-ok: bin indices for advanced indexing; PyTorch requires int64 design = self._iso_design.to(target.dtype) coeff = torch.linalg.lstsq(design, target.unsqueeze(1)).solution.squeeze(1) self.c_iso = nn.Parameter(coeff.detach().to(self.device)) @@ -396,7 +396,7 @@ def get_binwise_mean_intensity(self, fcalc: torch.Tensor): mean_calc_intensity = torch.zeros(self.nbins, device=self.device, dtype=fobs.dtype) counts = torch.zeros(self.nbins, device=self.device, dtype=fobs.dtype) counts_vals = torch.ones_like(F_calc, device=self.device, dtype=fobs.dtype) - bins_sel = self.bins.to(torch.int64)[sel] + bins_sel = self.bins.to(torch.int64)[sel] # dtype-ok: bin indices for advanced indexing; PyTorch requires int64 mean_obs_intensity = torch.scatter_add( mean_obs_intensity, 0, bins_sel, intensities[sel] ) diff --git a/torchref/scaling/solvent.py b/torchref/scaling/solvent.py index f9b76945..09b67a04 100644 --- a/torchref/scaling/solvent.py +++ b/torchref/scaling/solvent.py @@ -373,7 +373,7 @@ def get_solvent_mask(self): # grids, where the SF code's 1024 would OOM (denser intermediates). ATOM_CHUNK = 256 - grid_dims = torch.tensor(grid_shape, dtype=torch.long, device=device) + grid_dims = torch.tensor(grid_shape, dtype=torch.long, device=device) # dtype-ok: grid dims for voxel index arithmetic; PyTorch requires int64 grid_shape_float = grid_dims.float() inv_grid = 1.0 / grid_shape_float G = frac.T @ frac # metric tensor: r²_cart = diff_frac · G · diff_frac @@ -450,12 +450,12 @@ def get_solvent_mask(self): protein_voxels = ( torch.cat(protein_chunks, dim=0) if protein_chunks - else torch.empty((0, 3), dtype=torch.long, device=device) + else torch.empty((0, 3), dtype=torch.long, device=device) # dtype-ok: empty (0,3) voxel index tensor; PyTorch requires int64 for indexing ) boundary_voxels = ( torch.cat(boundary_chunks, dim=0) if boundary_chunks - else torch.empty((0, 3), dtype=torch.long, device=device) + else torch.empty((0, 3), dtype=torch.long, device=device) # dtype-ok: empty (0,3) voxel index tensor; PyTorch requires int64 for indexing ) del protein_chunks, boundary_chunks diff --git a/torchref/symmetry/map_symmetry.py b/torchref/symmetry/map_symmetry.py index 18985a15..afbeb4a9 100644 --- a/torchref/symmetry/map_symmetry.py +++ b/torchref/symmetry/map_symmetry.py @@ -149,7 +149,7 @@ def _index_grid(self, op_index: int) -> torch.Tensor: transformed = transformed - torch.floor(transformed) shape_t = torch.tensor([nx, ny, nz], dtype=dtype, device=device) - indices = torch.round(transformed * shape_t).to(torch.int64) + indices = torch.round(transformed * shape_t).to(torch.int64) # dtype-ok: rounded voxel grid indices; int64 index tensor required indices[:, 0] %= nx indices[:, 1] %= ny indices[:, 2] %= nz diff --git a/torchref/symmetry/reciprocal_symmetry.py b/torchref/symmetry/reciprocal_symmetry.py index 66dfacff..f595ee70 100644 --- a/torchref/symmetry/reciprocal_symmetry.py +++ b/torchref/symmetry/reciprocal_symmetry.py @@ -82,7 +82,7 @@ def _expand_hkl( for i in range(n_ops): # h' = h @ R^T hkl_transformed = torch.round(torch.matmul(hkl_float, recip_matrices[i].T)).to( - torch.int32 + torch.int32 # dtype-ok: transformed Miller indices (hkl); fixed-width int32 representation ) # Phase shift from translation: -2π h·t, for h' = hR under the convention # F(h) = Σ_j f_j exp(+2πi h·x_j). Do NOT "simplify" the sign: the wrong sign @@ -124,10 +124,10 @@ def _expand_hkl( # Build output tensors expanded_hkl = torch.tensor( - [list(k) for k in unique_dict.keys()], dtype=torch.int32, device=device + [list(k) for k in unique_dict.keys()], dtype=torch.int32, device=device # dtype-ok: unique Miller indices (hkl); fixed-width int32 representation ) phase_shifts = torch.tensor(unique_phases, dtype=get_float_dtype(), device=device) - orig_idx_tensor = torch.tensor(orig_indices, dtype=torch.int64, device=device) + orig_idx_tensor = torch.tensor(orig_indices, dtype=torch.int64, device=device) # dtype-ok: reflection index mapping; int64 index tensor required if remove_absences and sym.number != 1: keep_mask = ~sym.is_absent(expanded_hkl) @@ -198,7 +198,7 @@ def _complete_hkl( all_hkl_np = all_hkl.cpu().numpy() n_complete = len(all_hkl) - input_indices = torch.full((n_complete,), -1, dtype=torch.int64, device=device) + input_indices = torch.full((n_complete,), -1, dtype=torch.int64, device=device) # dtype-ok: reflection index buffer (-1 sentinel); int64 index required missing_mask = torch.ones(n_complete, dtype=torch.bool, device=device) for i, hkl in enumerate(all_hkl_np): @@ -274,7 +274,7 @@ def get_canonical_hkl(hkl_single): for i in range(n_ops): # h' = h @ R^T hkl_trans = torch.round(torch.matmul(hkl_single, recip_matrices[i].T)).to( - torch.int32 + torch.int32 # dtype-ok: transformed Miller indices (hkl); fixed-width int32 representation ) equivalents.append(hkl_trans) @@ -311,7 +311,7 @@ def get_canonical_hkl(hkl_single): R = recip_matrices[equiv_idx] t = translations[equiv_idx] - hkl_trans = torch.round(torch.matmul(hkl_single, R.T)).to(torch.int32) + hkl_trans = torch.round(torch.matmul(hkl_single, R.T)).to(torch.int32) # dtype-ok: transformed Miller indices (hkl); fixed-width int32 representation # -2π h·t, same convention as expand_hkl (see the derivation there). phase_shift = -2.0 * np.pi * torch.matmul(hkl_single, t) @@ -333,9 +333,9 @@ def get_canonical_hkl(hkl_single): asu_list = sorted(asu_reflections.keys()) n_asu = len(asu_list) - hkl_asu = torch.tensor(asu_list, dtype=torch.int32, device=device) + hkl_asu = torch.tensor(asu_list, dtype=torch.int32, device=device) # dtype-ok: ASU Miller indices (hkl); fixed-width int32 representation reduction_indices = torch.full( - (n_asu, n_equiv), -1, dtype=torch.int64, device=device + (n_asu, n_equiv), -1, dtype=torch.int64, device=device # dtype-ok: reduction index map (-1 sentinel); int64 index tensor required ) phase_shifts = torch.zeros((n_asu, n_equiv), dtype=get_float_dtype(), device=device) @@ -446,7 +446,7 @@ def _canonicalize_hkl( empty_hkl = torch.empty((0, 3), dtype=hkl_dtype, device=device) empty_f = torch.empty(0, dtype=get_float_dtype(), device=device) empty_b = torch.empty(0, dtype=torch.bool, device=device) - empty_i = torch.empty(0, dtype=torch.int64, device=device) + empty_i = torch.empty(0, dtype=torch.int64, device=device) # dtype-ok: empty index tensor; int64 index dtype required return empty_hkl, empty_f, empty_b, empty_i # The ASU lookup tables are numpy-backed, so the operations come across to CPU @@ -550,9 +550,9 @@ def _canonicalize_hkl( h_max = int(canonical_hkl.abs().max().item()) + 1 base = 2 * h_max + 1 sort_key = ( - canonical_hkl[:, 0].to(torch.int64) * base * base - + canonical_hkl[:, 1].to(torch.int64) * base - + canonical_hkl[:, 2].to(torch.int64) + canonical_hkl[:, 0].to(torch.int64) * base * base # dtype-ok: linear HKL hash/key; int64 avoids overflow for indexing + + canonical_hkl[:, 1].to(torch.int64) * base # dtype-ok: linear HKL hash/key; int64 avoids overflow for indexing + + canonical_hkl[:, 2].to(torch.int64) # dtype-ok: linear HKL hash/key; int64 avoids overflow for indexing ) sort_indices = torch.argsort(sort_key) diff --git a/torchref/symmetry/symmetry.py b/torchref/symmetry/symmetry.py index 5d68cffc..860825de 100644 --- a/torchref/symmetry/symmetry.py +++ b/torchref/symmetry/symmetry.py @@ -356,7 +356,7 @@ def expand_reciprocal(self, hkl: torch.Tensor) -> torch.Tensor: operations on integer indices and only mops up float error. """ equivalents = self.reciprocal.apply_rotations(hkl) - return torch.round(equivalents).to(torch.int64) + return torch.round(equivalents).to(torch.int64) # dtype-ok: rounded Miller equivalents; int64 for exact integer compare/index # ========================================================================= # Reflection predicates @@ -381,7 +381,7 @@ def is_centric(self, hkl: torch.Tensor) -> torch.Tensor: with torch.no_grad(): flat = hkl.reshape(-1, 3) equivalents = self.expand_reciprocal(flat) # (n_ops, N, 3) - target = -flat.to(device=equivalents.device, dtype=torch.int64) + target = -flat.to(device=equivalents.device, dtype=torch.int64) # dtype-ok: compare target for int64 equivalents; dtype must match centric = (equivalents == target).all(dim=-1).any(dim=0) return centric.reshape(original_shape).to(hkl.device) @@ -405,7 +405,7 @@ def is_absent(self, hkl: torch.Tensor) -> torch.Tensor: with torch.no_grad(): flat = hkl.reshape(-1, 3) equivalents = self.expand_reciprocal(flat) # (n_ops, N, 3) - target = flat.to(device=equivalents.device, dtype=torch.int64) + target = flat.to(device=equivalents.device, dtype=torch.int64) # dtype-ok: compare target for int64 equivalents; dtype must match maps_to_self = (equivalents == target).all(dim=-1) # (n_ops, N) h_dot_t = torch.matmul( @@ -440,7 +440,7 @@ def epsilon(self, hkl: torch.Tensor) -> torch.Tensor: float_dtype = get_float_dtype() with torch.no_grad(): equivalents = self.expand_reciprocal(hkl) # (n_ops, N, 3) - target = hkl.to(device=equivalents.device, dtype=torch.int64) + target = hkl.to(device=equivalents.device, dtype=torch.int64) # dtype-ok: compare target for int64 equivalents; dtype must match same = (equivalents == target).all(dim=-1) friedel = (equivalents == -target).all(dim=-1) eps = (same | friedel).sum(dim=0).clamp(min=1).to(float_dtype) @@ -468,7 +468,7 @@ def grid_requirements(self) -> dict: # ``Fraction(float)`` would need a tolerance where this is exact. numerators = torch.round( self.translations.detach().cpu().double() * _TRANSLATION_DENOMINATOR - ).to(torch.int64) + ).to(torch.int64) # dtype-ok: integer translation numerators for exact Fraction recovery for op_numerators in numerators.tolist(): for axis, numerator in enumerate(op_numerators): diff --git a/torchref/topology/atom_graph.py b/torchref/topology/atom_graph.py index 2dfce660..0e1170e6 100644 --- a/torchref/topology/atom_graph.py +++ b/torchref/topology/atom_graph.py @@ -40,8 +40,8 @@ def _build_csr(bonds: torch.Tensor, n_atoms: int) -> Tuple[torch.Tensor, torch.T device = bonds.device if bonds.numel() == 0: return ( - torch.zeros(n_atoms + 1, dtype=torch.int64, device=device), - torch.zeros(0, dtype=torch.int64, device=device), + torch.zeros(n_atoms + 1, dtype=torch.int64, device=device), # dtype-ok: CSR indptr offset array; int64 index required + torch.zeros(0, dtype=torch.int64, device=device), # dtype-ok: empty CSR neighbor index array; int64 index required ) src = torch.cat([bonds[:, 0], bonds[:, 1]]) @@ -55,9 +55,9 @@ def _build_csr(bonds: torch.Tensor, n_atoms: int) -> Tuple[torch.Tensor, torch.T src, dst = src[order], dst[order] counts = torch.bincount(src, minlength=n_atoms) - indptr = torch.zeros(n_atoms + 1, dtype=torch.int64, device=device) + indptr = torch.zeros(n_atoms + 1, dtype=torch.int64, device=device) # dtype-ok: CSR indptr offset array; int64 index required torch.cumsum(counts, dim=0, out=indptr[1:]) - return indptr, dst.to(torch.int64) + return indptr, dst.to(torch.int64) # dtype-ok: CSR neighbor (dst) index array; int64 index required def _extend_paths( @@ -81,13 +81,13 @@ def _extend_paths( """ device = paths.device if paths.numel() == 0: - return torch.zeros((0, paths.shape[1] + 1), dtype=torch.int64, device=device) + return torch.zeros((0, paths.shape[1] + 1), dtype=torch.int64, device=device) # dtype-ok: empty BFS path index array; int64 index required last, prev = paths[:, -1], paths[:, -2] counts = indptr[last + 1] - indptr[last] total = int(counts.sum()) if total == 0: - return torch.zeros((0, paths.shape[1] + 1), dtype=torch.int64, device=device) + return torch.zeros((0, paths.shape[1] + 1), dtype=torch.int64, device=device) # dtype-ok: empty BFS path index array; int64 index required row = torch.repeat_interleave(torch.arange(len(paths), device=device), counts) # Offset of each slot within its own neighbour list. diff --git a/torchref/topology/build.py b/torchref/topology/build.py index 12febd63..d26631aa 100644 --- a/torchref/topology/build.py +++ b/torchref/topology/build.py @@ -667,7 +667,7 @@ def _block_with_values( per_origin, arity, edge_type, payload ) block = EdgeBlock( - indices=torch.as_tensor(indices, dtype=torch.int64, device=device), + indices=torch.as_tensor(indices, dtype=torch.int64, device=device), # dtype-ok: atom index tensor for restraint edges; int64 index required origin_bounds=bounds, ) values = { @@ -934,7 +934,7 @@ def build_topology_with_values( np.arange(n_res, dtype=np.int64), nodes["atom_end"] - nodes["atom_start"], ), - dtype=torch.int64, + dtype=torch.int64, # dtype-ok: atom index tensor; int64 index required device=device, ), bonds=bond_block, diff --git a/torchref/topology/builders.py b/torchref/topology/builders.py index df7c27ab..3faaba1c 100644 --- a/torchref/topology/builders.py +++ b/torchref/topology/builders.py @@ -19,7 +19,7 @@ import pandas as pd import torch -from torchref.config import get_float_dtype +from torchref.config import get_float_dtype, get_int_dtype # Import the Numba-accelerated matching functions from torchref.topology.builders_numba import ( @@ -613,7 +613,7 @@ def build( sigmas = np.where(sigmas == 0, 1e-4, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), + "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -720,7 +720,7 @@ def build( sigmas = np.where(sigmas == 0, 1e-4, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), + "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -841,10 +841,10 @@ def build( sigmas = np.where(sigmas == 0, 1e-4, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), + "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), - "periods": torch.tensor(periods, dtype=torch.long, device=device), + "periods": torch.tensor(periods, dtype=get_int_dtype(), device=device), } @@ -932,7 +932,7 @@ def build( key = f"{n_atoms}_atoms" result[key] = { - "indices": torch.tensor(indices, dtype=torch.long, device=device), + "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -1054,7 +1054,7 @@ def build( sigmas = np.where(sigmas == 0, 1e-4, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), + "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 "ideal_volumes": torch.tensor( ideal_volumes, dtype=get_float_dtype(), device=device ), @@ -1309,7 +1309,7 @@ def finalize( sigmas = np.where(sigmas == 0, min_sigma, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), + "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -1410,7 +1410,7 @@ def build( sigmas = np.where(sigmas == 0, 1e-4, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), + "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -1532,7 +1532,7 @@ def finalize( sigmas = np.where(sigmas == 0, min_sigma, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), + "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -1644,7 +1644,7 @@ def build( sigmas = np.where(sigmas == 0, 1e-4, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), + "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -1781,10 +1781,10 @@ def finalize_disulfide( periods = periods[sort_order] return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), + "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), - "periods": torch.tensor(periods, dtype=torch.long, device=device), + "periods": torch.tensor(periods, dtype=get_int_dtype(), device=device), } @property @@ -1946,8 +1946,8 @@ def build( indices = indices[order] periods = periods[order] result["phi"] = { - "indices": torch.tensor(indices, dtype=torch.long, device=device), - "periods": torch.tensor(periods, dtype=torch.long, device=device), + "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "periods": torch.tensor(periods, dtype=get_int_dtype(), device=device), } # Finalize psi @@ -1959,8 +1959,8 @@ def build( indices = indices[order] periods = periods[order] result["psi"] = { - "indices": torch.tensor(indices, dtype=torch.long, device=device), - "periods": torch.tensor(periods, dtype=torch.long, device=device), + "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "periods": torch.tensor(periods, dtype=get_int_dtype(), device=device), } # Finalize omega @@ -1978,12 +1978,12 @@ def build( periods = periods[order] is_proline = is_proline[order] result["omega"] = { - "indices": torch.tensor(indices, dtype=torch.long, device=device), + "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 "references": torch.tensor( references, dtype=get_float_dtype(), device=device ), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), - "periods": torch.tensor(periods, dtype=torch.long, device=device), + "periods": torch.tensor(periods, dtype=get_int_dtype(), device=device), "is_proline": torch.tensor(is_proline, dtype=torch.bool, device=device), } @@ -2017,13 +2017,13 @@ def build( stypes = stypes[order] result["ramachandran"] = { "phi_indices": torch.tensor( - phi_idx, dtype=torch.long, device=device + phi_idx, dtype=torch.long, device=device # dtype-ok: phi atom-index tensor for dihedral; int64 required ), "psi_indices": torch.tensor( - psi_idx, dtype=torch.long, device=device + psi_idx, dtype=torch.long, device=device # dtype-ok: psi atom-index tensor for dihedral; int64 required ), "surface_type": torch.tensor( - stypes, dtype=torch.long, device=device + stypes, dtype=torch.long, device=device # dtype-ok: categorical rama surface-type code used as advanced index; int64 ), } @@ -2120,7 +2120,7 @@ def build( key = f"{n_atoms}_atoms" result[key] = { - "indices": torch.tensor(indices, dtype=torch.long, device=device), + "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } diff --git a/torchref/topology/edges.py b/torchref/topology/edges.py index 06223a81..c70a5f2d 100644 --- a/torchref/topology/edges.py +++ b/torchref/topology/edges.py @@ -153,7 +153,7 @@ class EdgeBlock(DeviceMixin): def empty(cls, arity: int, device=None) -> "EdgeBlock": """An edge-free block of the given arity.""" return cls( - indices=torch.zeros((0, arity), dtype=torch.int64, device=device), + indices=torch.zeros((0, arity), dtype=torch.int64, device=device), # dtype-ok: empty edge index tensor (0,arity); int64 index required origin_bounds={}, ) @@ -190,7 +190,7 @@ def from_origins( if len(indices) == 0: return cls.empty(arity, device=device) return cls( - indices=torch.as_tensor(indices, dtype=torch.int64, device=device), + indices=torch.as_tensor(indices, dtype=torch.int64, device=device), # dtype-ok: edge atom index tensor; int64 index required origin_bounds=bounds, ) diff --git a/torchref/topology/nonbonded.py b/torchref/topology/nonbonded.py index a8246b08..404e7a6f 100644 --- a/torchref/topology/nonbonded.py +++ b/torchref/topology/nonbonded.py @@ -85,8 +85,8 @@ def prefilter_symop_offsets( valid_ops.append(op_idx) valid_offsets.append([dx, dy, dz]) - op_indices = torch.tensor(valid_ops, dtype=torch.long, device=device) - cell_offsets = torch.tensor(valid_offsets, dtype=torch.long, device=device) + op_indices = torch.tensor(valid_ops, dtype=torch.long, device=device) # dtype-ok: symmetry-operator index tensor; int64 + cell_offsets = torch.tensor(valid_offsets, dtype=torch.long, device=device) # dtype-ok: integer cell-offset lattice vectors; symmetry-image metadata return op_indices, cell_offsets @@ -146,7 +146,7 @@ def assign_to_grid( gd = grid_dims.to(device=device, dtype=fdtype) cell_ijk = (frac_wrapped * gd[None, None, :]).long() cell_ijk = cell_ijk.clamp( - min=torch.zeros(3, dtype=torch.long, device=device), + min=torch.zeros(3, dtype=torch.long, device=device), # dtype-ok: clamp min-bound for long grid-index tensor; matches int64 max=(grid_dims - 1).to(device), ) @@ -188,14 +188,14 @@ def build_cell_list( unique_cells, counts = torch.unique_consecutive( sorted_cells, return_counts=True ) - starts = torch.zeros(len(unique_cells) + 1, dtype=torch.long, device=device) + starts = torch.zeros(len(unique_cells) + 1, dtype=torch.long, device=device) # dtype-ok: CSR boundary/offset array; int64 required starts[1:] = counts.cumsum(0) cell_lookup = torch.full( - (n_grid_total,), -1, dtype=torch.long, device=device + (n_grid_total,), -1, dtype=torch.long, device=device # dtype-ok: grid-cell to index lookup table; used for indexing, int64 ) cell_lookup[unique_cells] = torch.arange( - len(unique_cells), dtype=torch.long, device=device + len(unique_cells), dtype=torch.long, device=device # dtype-ok: index values written into lookup table; int64 ) return sort_order, unique_cells, starts, cell_lookup @@ -233,7 +233,7 @@ def _get_canonical_offsets_14(device: torch.device) -> torch.Tensor: offsets.append([dx, dy, dz]) assert len(offsets) == 14, f"expected 14 canonical offsets, got {len(offsets)}" _NEIGHBOR_OFFSETS_14 = torch.tensor( - offsets, dtype=torch.long, device=device + offsets, dtype=torch.long, device=device # dtype-ok: grid neighbor-cell offset deltas used to compute index; int64 ) return _NEIGHBOR_OFFSETS_14 @@ -493,7 +493,7 @@ def find_pairs_periodic_grid_v2( all_pair_combo_j.append(cj) if not all_pair_atom_i: - empty = torch.tensor([], dtype=torch.long, device=device) + empty = torch.tensor([], dtype=torch.long, device=device) # dtype-ok: empty atom-pair index placeholder; int64 required return empty, empty, empty return ( @@ -517,11 +517,11 @@ def exclusion_set_to_hash( Hash: min(i,j) * max_idx + max(i,j), sorted for searchsorted. """ if not exclusion_set: - return torch.tensor([], dtype=torch.long, device=device) + return torch.tensor([], dtype=torch.long, device=device) # dtype-ok: empty exclusion-hash placeholder; int64 arr = np.array(list(exclusion_set), dtype=np.int64) hashes = arr[:, 0] * max_idx + arr[:, 1] # already (min, max) hashes.sort() - return torch.tensor(hashes, dtype=torch.long, device=device) + return torch.tensor(hashes, dtype=torch.long, device=device) # dtype-ok: packed pair-hash key for searchsorted; int64 avoids overflow def filter_pairs( @@ -629,11 +629,11 @@ def build_vdw_restraints_gpu( sg = SG(sg) empty_result = { - "indices": torch.zeros(0, 2, dtype=torch.long, device=device), + "indices": torch.zeros(0, 2, dtype=torch.long, device=device), # dtype-ok: atom-pair index tensor; torch indexing requires int64 "min_distances": torch.zeros(0, dtype=get_float_dtype(), device=device), "sigmas": torch.zeros(0, dtype=get_float_dtype(), device=device), - "symop_indices": torch.zeros(0, dtype=torch.long, device=device), - "cell_offsets": torch.zeros(0, 3, dtype=torch.long, device=device), + "symop_indices": torch.zeros(0, dtype=torch.long, device=device), # dtype-ok: symmetry-operator index tensor; int64 + "cell_offsets": torch.zeros(0, 3, dtype=torch.long, device=device), # dtype-ok: integer cell-offset lattice vectors; symmetry-image metadata } # Step 1: prefilter symop combos @@ -655,10 +655,10 @@ def build_vdw_restraints_gpu( if len(identity_indices) == 0: # Identity not in valid combos — should not happen, but add it op_indices = torch.cat([ - torch.zeros(1, dtype=torch.long, device=device), op_indices + torch.zeros(1, dtype=torch.long, device=device), op_indices # dtype-ok: identity prepended to symop-index tensor; int64 ]) cell_offsets_valid = torch.cat([ - torch.zeros(1, 3, dtype=torch.long, device=device), cell_offsets_valid + torch.zeros(1, 3, dtype=torch.long, device=device), cell_offsets_valid # dtype-ok: identity prepended to cell-offset tensor; int64 ]) identity_combo = 0 M = len(op_indices) @@ -774,7 +774,7 @@ def build_vdw_restraints_gpu( "valid_op_indices": op_indices, "valid_cell_offsets": cell_offsets_valid, "grid_dims": grid_dims, - "identity_combo": torch.tensor(identity_combo, dtype=torch.long, device=device), + "identity_combo": torch.tensor(identity_combo, dtype=torch.long, device=device), # dtype-ok: combo index scalar into symop/offset arrays; int64 } if verbose > 0: @@ -856,7 +856,7 @@ def find_h_vdw_pairs_gpu( xyz_all = torch.cat([xyz_heavy, xyz_h], dim=0) # (N_all, 3) n_all = xyz_all.shape[0] - empty = torch.tensor([], dtype=torch.long, device=device) + empty = torch.tensor([], dtype=torch.long, device=device) # dtype-ok: empty index placeholder tensor; int64 required if n_all == 0: return empty, empty, empty diff --git a/torchref/topology/residue_graph.py b/torchref/topology/residue_graph.py index b4651d1c..0a68a21c 100644 --- a/torchref/topology/residue_graph.py +++ b/torchref/topology/residue_graph.py @@ -272,7 +272,7 @@ def find_disulfide_links( rows = list(sg_rows) if len(rows) < 2: return [] - idx = torch.as_tensor(rows, dtype=torch.int64, device=xyz.device) + idx = torch.as_tensor(rows, dtype=torch.int64, device=xyz.device) # dtype-ok: residue-atom index tensor; int64 index required dist = torch.cdist(xyz[idx], xyz[idx]) close = (dist > DISULFIDE_MIN_DISTANCE) & (dist < DISULFIDE_MAX_DISTANCE) diff --git a/torchref/topology/restraint_sets.py b/torchref/topology/restraint_sets.py index 289048c9..45c079ae 100644 --- a/torchref/topology/restraint_sets.py +++ b/torchref/topology/restraint_sets.py @@ -41,7 +41,7 @@ def to_tensor(values, prop: str, device=None) -> torch.Tensor: if isinstance(values, torch.Tensor): return values.to(device=device) if device is not None else values if prop in _INTEGER_PROPERTIES: - dtype = torch.int64 + dtype = torch.int64 # dtype-ok: dtype var for index tensors; int64 index required elif prop in _BOOL_PROPERTIES: dtype = torch.bool else: diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index 4a649df7..2ee0a3c9 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -400,7 +400,7 @@ def _find_nearby_pairs_spatial_hash(self, xyz, cutoff=6.0): n_atoms = xyz.shape[0] if n_atoms == 0: - return torch.tensor([], dtype=torch.long, device=device).reshape(0, 2) + return torch.tensor([], dtype=torch.long, device=device).reshape(0, 2) # dtype-ok: empty atom-pair index tensor; int64 index required # Work on CPU to avoid per-iteration GPU kernel launch overhead coords = xyz.detach().cpu() @@ -425,12 +425,12 @@ def _find_nearby_pairs_spatial_hash(self, xyz, cutoff=6.0): sorted_flat, return_counts=True ) n_unique = len(unique_cells) - starts = torch.zeros(n_unique + 1, dtype=torch.long) + starts = torch.zeros(n_unique + 1, dtype=torch.long) # dtype-ok: grid-cell CSR start offsets; int64 index required starts[1:] = counts.cumsum(0) # Lookup: flat_cell -> index in unique_cells (-1 if empty) n_grid = gx * gyz - cell_lookup = torch.full((n_grid,), -1, dtype=torch.long) + cell_lookup = torch.full((n_grid,), -1, dtype=torch.long) # dtype-ok: cell lookup table (-1 sentinel); int64 index required cell_lookup[unique_cells] = torch.arange(n_unique) # 14 unique neighbour offsets: self (0,0,0) + 13 forward neighbours. @@ -516,9 +516,9 @@ def _find_nearby_pairs_spatial_hash(self, xyz, cutoff=6.0): if pair_chunks: all_pairs = np.concatenate(pair_chunks, axis=0) - return torch.from_numpy(all_pairs).to(dtype=torch.long, device=device) + return torch.from_numpy(all_pairs).to(dtype=torch.long, device=device) # dtype-ok: atom-pair index array from numpy; int64 index required else: - return torch.tensor([], dtype=torch.long, device=device).reshape(0, 2) + return torch.tensor([], dtype=torch.long, device=device).reshape(0, 2) # dtype-ok: empty atom-pair index tensor; int64 index required def _expand_with_symmetry_mates(self, xyz, cutoff): """Append symmetry-mate positions to ASU ``xyz`` for neighbour search. @@ -651,7 +651,7 @@ def _build_h_exclusion_hash(self, h_topo, device): ``torch.searchsorted`` lookup. """ if h_topo is None or h_topo.n_hydrogens == 0: - return torch.tensor([], dtype=torch.long, device=device) + return torch.tensor([], dtype=torch.long, device=device) # dtype-ok: empty index tensor; int64 index required n_heavy = len(self.pdb) n_h = h_topo.n_hydrogens @@ -676,13 +676,13 @@ def _build_h_exclusion_hash(self, h_topo, device): exclusions.add((min(h_combined, nb), max(h_combined, nb))) if not exclusions: - return torch.tensor([], dtype=torch.long, device=device) + return torch.tensor([], dtype=torch.long, device=device) # dtype-ok: empty index tensor; int64 index required arr = np.array(list(exclusions), dtype=np.int64) max_idx = max(n_heavy + n_h, int(arr.max()) + 1) hashes = arr[:, 0] * max_idx + arr[:, 1] hashes.sort() - return torch.tensor(hashes, dtype=torch.long, device=device) + return torch.tensor(hashes, dtype=torch.long, device=device) # dtype-ok: grid-cell hash values used as keys/index; int64 required def _build_vdw_restraints( self, cutoff=6.0, sigma=0.2, inter_residue_only=True, use_spatial_hash=True @@ -887,17 +887,17 @@ def _build_vdw_restraints_legacy( if dist_sq < cutoff_sq: pairs_list.append([i, j]) nearby_pairs = ( - torch.tensor(pairs_list, dtype=torch.long, device=device) + torch.tensor(pairs_list, dtype=torch.long, device=device) # dtype-ok: atom-pair index tensor; int64 index required if pairs_list - else torch.tensor([], dtype=torch.long, device=device).reshape(0, 2) + else torch.tensor([], dtype=torch.long, device=device).reshape(0, 2) # dtype-ok: empty atom-pair index tensor; int64 index required ) empty_result = { - "indices": torch.tensor([], dtype=torch.long, device=device).reshape(0, 2), + "indices": torch.tensor([], dtype=torch.long, device=device).reshape(0, 2), # dtype-ok: empty atom-pair index tensor; int64 index required "min_distances": torch.tensor([], dtype=get_float_dtype(), device=device), "sigmas": torch.tensor([], dtype=get_float_dtype(), device=device), - "symop_indices": torch.tensor([], dtype=torch.long, device=device), - "cell_offsets": torch.tensor([], dtype=torch.long, device=device).reshape(0, 3), + "symop_indices": torch.tensor([], dtype=torch.long, device=device), # dtype-ok: empty symop index tensor; int64 index required + "cell_offsets": torch.tensor([], dtype=torch.long, device=device).reshape(0, 3), # dtype-ok: empty cell-offset index tensor; int64 index required } if len(nearby_pairs) == 0: @@ -1029,7 +1029,7 @@ def _build_vdw_restraints_legacy( # Store results final_pairs = np.stack([final_i1, final_i2], axis=1) self._vdw = { - "indices": torch.tensor(final_pairs, dtype=torch.long, device=device), + "indices": torch.tensor(final_pairs, dtype=torch.long, device=device), # dtype-ok: final atom-pair index tensor; int64 index required "min_distances": torch.tensor( min_distances, dtype=get_float_dtype(), device=device ), @@ -1037,10 +1037,10 @@ def _build_vdw_restraints_legacy( (len(final_pairs),), sigma, dtype=get_float_dtype(), device=device ), "symop_indices": torch.tensor( - final_symop, dtype=torch.long, device=device + final_symop, dtype=torch.long, device=device # dtype-ok: symop index tensor; int64 index required ), "cell_offsets": torch.tensor( - final_offsets, dtype=torch.long, device=device + final_offsets, dtype=torch.long, device=device # dtype-ok: cell-offset index tensor; int64 index required ), } diff --git a/torchref/topology/riding.py b/torchref/topology/riding.py index 588cf6f8..f36a681c 100644 --- a/torchref/topology/riding.py +++ b/torchref/topology/riding.py @@ -431,17 +431,17 @@ def build_hydrogen_topology( fdtype = dtypes.float if n_h_total == 0: - topo.h_parent_idx = torch.zeros(0, dtype=torch.long, device=device) + topo.h_parent_idx = torch.zeros(0, dtype=torch.long, device=device) # dtype-ok: parent atom-index tensor (empty); int64 required topo.h_bond_length = torch.zeros(0, dtype=fdtype, device=device) topo.h_vdw_radius = torch.zeros(0, dtype=fdtype, device=device) - topo.h_placement_type = torch.zeros(0, dtype=torch.long, device=device) - topo.h_slot_in_parent = torch.zeros(0, dtype=torch.long, device=device) + topo.h_placement_type = torch.zeros(0, dtype=torch.long, device=device) # dtype-ok: categorical H placement-type code (empty) + topo.h_slot_in_parent = torch.zeros(0, dtype=torch.long, device=device) # dtype-ok: slot index into parent (empty); int64 topo.parent_neighbor_idx = torch.zeros( - 0, MAX_HEAVY_NB, dtype=torch.long, device=device + 0, MAX_HEAVY_NB, dtype=torch.long, device=device # dtype-ok: parent neighbor atom-index tensor (empty); int64 required ) - topo.parent_neighbor_count = torch.zeros(0, dtype=torch.long, device=device) - topo.h_chainid_enc = torch.zeros(0, dtype=torch.long, device=device) - topo.h_resseq = torch.zeros(0, dtype=torch.long, device=device) + topo.parent_neighbor_count = torch.zeros(0, dtype=torch.long, device=device) # dtype-ok: per-parent neighbor count (empty); structural int + topo.h_chainid_enc = torch.zeros(0, dtype=torch.long, device=device) # dtype-ok: categorical chain-id encoding (empty) + topo.h_resseq = torch.zeros(0, dtype=torch.long, device=device) # dtype-ok: residue sequence id (empty); categorical return topo # Sort all topology arrays by placement type for contiguous slicing @@ -466,22 +466,22 @@ def build_hydrogen_topology( idxs = np.where(mask)[0] type_bounds[t] = (int(idxs[0]), int(idxs[-1]) + 1) - topo.h_parent_idx = torch.tensor(acc_parent_idx, dtype=torch.long, device=device) + topo.h_parent_idx = torch.tensor(acc_parent_idx, dtype=torch.long, device=device) # dtype-ok: parent atom-index tensor; torch indexing requires int64 topo.h_bond_length = torch.tensor(acc_bond_length, dtype=fdtype, device=device) topo.h_vdw_radius = torch.full((n_h_total,), 1.20, dtype=fdtype, device=device) topo.h_placement_type = torch.tensor( - acc_placement_type, dtype=torch.long, device=device + acc_placement_type, dtype=torch.long, device=device # dtype-ok: categorical H placement-type code; used for sort/slice ) - topo.h_slot_in_parent = torch.tensor(acc_slot, dtype=torch.long, device=device) + topo.h_slot_in_parent = torch.tensor(acc_slot, dtype=torch.long, device=device) # dtype-ok: slot index into parent neighbor slots; int64 topo.parent_neighbor_idx = torch.tensor( - np.stack(acc_nb_idx), dtype=torch.long, device=device + np.stack(acc_nb_idx), dtype=torch.long, device=device # dtype-ok: parent neighbor atom-index tensor; int64 required ) topo.parent_neighbor_count = torch.tensor( - acc_nb_count, dtype=torch.long, device=device + acc_nb_count, dtype=torch.long, device=device # dtype-ok: per-parent neighbor count; structural int metadata ) topo.type_bounds = type_bounds # dict: type_code -> (start, end) - topo.h_chainid_enc = torch.tensor(acc_chainid_enc, dtype=torch.long, device=device) - topo.h_resseq = torch.tensor(acc_resseq, dtype=torch.long, device=device) + topo.h_chainid_enc = torch.tensor(acc_chainid_enc, dtype=torch.long, device=device) # dtype-ok: categorical chain-id encoding + topo.h_resseq = torch.tensor(acc_resseq, dtype=torch.long, device=device) # dtype-ok: residue sequence id; categorical if verbose > 0: print(f" Hydrogen topology: {n_h_total} riding H atoms") @@ -736,8 +736,8 @@ def build_h_candidate_pairs( if n_h == 0: for name in ("cand_idx_i", "cand_idx_j", "cand_symop_idx"): - setattr(h_topo, name, torch.zeros(0, dtype=torch.long, device=device)) - h_topo.cand_cell_offset = torch.zeros(0, 3, dtype=torch.long, device=device) + setattr(h_topo, name, torch.zeros(0, dtype=torch.long, device=device)) # dtype-ok: candidate atom/symop index tensors (empty); int64 required + h_topo.cand_cell_offset = torch.zeros(0, 3, dtype=torch.long, device=device) # dtype-ok: integer cell-offset lattice vectors (empty); symmetry metadata h_topo.cand_min_dist = torch.zeros(0, dtype=dtypes.float, device=device) return @@ -836,15 +836,15 @@ def _same_res(chain_a, resseq_a, chain_b, resseq_b): if not acc_idx_i: for name in ("cand_idx_i", "cand_idx_j", "cand_symop_idx"): - setattr(h_topo, name, torch.zeros(0, dtype=torch.long, device=device)) - h_topo.cand_cell_offset = torch.zeros(0, 3, dtype=torch.long, device=device) + setattr(h_topo, name, torch.zeros(0, dtype=torch.long, device=device)) # dtype-ok: candidate atom/symop index tensors (empty); int64 required + h_topo.cand_cell_offset = torch.zeros(0, 3, dtype=torch.long, device=device) # dtype-ok: integer cell-offset lattice vectors (empty); symmetry metadata h_topo.cand_min_dist = torch.zeros(0, dtype=dtypes.float, device=device) return - cand_i = torch.tensor(acc_idx_i, dtype=torch.long, device=device) - cand_j = torch.tensor(acc_idx_j, dtype=torch.long, device=device) - cand_sym = torch.tensor(acc_symop, dtype=torch.long, device=device) - cand_off = torch.tensor(np.stack(acc_offset), dtype=torch.long, device=device) + cand_i = torch.tensor(acc_idx_i, dtype=torch.long, device=device) # dtype-ok: combined atom-index tensor; torch indexing requires int64 + cand_j = torch.tensor(acc_idx_j, dtype=torch.long, device=device) # dtype-ok: combined atom-index tensor; torch indexing requires int64 + cand_sym = torch.tensor(acc_symop, dtype=torch.long, device=device) # dtype-ok: symmetry-operator index; int64 + cand_off = torch.tensor(np.stack(acc_offset), dtype=torch.long, device=device) # dtype-ok: integer cell-offset lattice vectors; symmetry-image metadata # Apply 1-2 / 1-3 exclusions for intra-ASU candidates if h_excl_hash is not None and len(h_excl_hash) > 0: diff --git a/torchref/topology/topology.py b/torchref/topology/topology.py index 78a7222a..51419991 100644 --- a/torchref/topology/topology.py +++ b/torchref/topology/topology.py @@ -89,7 +89,7 @@ def subset(self, keep) -> "Topology": mask = torch.as_tensor(keep) if mask.dtype != torch.bool: selected = torch.zeros(self.n_atoms, dtype=torch.bool) - selected[mask.to(torch.int64)] = True + selected[mask.to(torch.int64)] = True # dtype-ok: boolean-mask->index cast for scatter select; int64 index required mask = selected mask = mask.to(device=self.atoms.residue_of.device) @@ -97,8 +97,8 @@ def subset(self, keep) -> "Topology": raise ValueError("subset would keep no atoms") n_kept = int(mask.sum()) - remap = torch.full((self.n_atoms,), -1, dtype=torch.int64, device=mask.device) - remap[mask] = torch.arange(n_kept, dtype=torch.int64, device=mask.device) + remap = torch.full((self.n_atoms,), -1, dtype=torch.int64, device=mask.device) # dtype-ok: atom remap index array (-1 sentinel); int64 index required + remap[mask] = torch.arange(n_kept, dtype=torch.int64, device=mask.device) # dtype-ok: arange remap indices; int64 index required # A residue survives if any of its atoms does. Counting per residue also # gives the new atom ranges, contiguous because the atom order is unchanged. @@ -112,10 +112,10 @@ def subset(self, keep) -> "Topology": atom_start = atom_end - counts residue_remap = torch.full( - (self.n_residues,), -1, dtype=torch.int64, device=mask.device + (self.n_residues,), -1, dtype=torch.int64, device=mask.device # dtype-ok: residue remap index array (-1 sentinel); int64 index required ) residue_remap[torch.as_tensor(residue_keep, device=mask.device)] = torch.arange( - int(residue_keep.sum()), dtype=torch.int64, device=mask.device + int(residue_keep.sum()), dtype=torch.int64, device=mask.device # dtype-ok: arange residue remap indices; int64 index required ) return Topology( diff --git a/torchref/utils/device_mixin.py b/torchref/utils/device_mixin.py index 6c8327f1..c4f4d2da 100644 --- a/torchref/utils/device_mixin.py +++ b/torchref/utils/device_mixin.py @@ -347,14 +347,14 @@ def run(device, dtype): # Each axis gets its own pair, varying only along the axis it measures. Sharing one # pair couples them: an accelerator scratch cannot be cast to float64 on MPS, so # probing dtype on the device pair makes ``.double()`` unprobeable there. - base = run(torch.device("cpu"), torch.float32) + base = run(torch.device("cpu"), torch.float32) # dtype-ok: fixed probe dtype is what the preservation test varies, not a config allocation if base is None: return None, None device = base.device if accel is not None: # Contrast pair for the device axis: same dtype, different device. - other = run(accel, torch.float32) + other = run(accel, torch.float32) # dtype-ok: fixed probe dtype (device-axis contrast) if other is None: device = None elif other.device != base.device: @@ -365,7 +365,7 @@ def run(device, dtype): # Contrast pair for the dtype axis: same device, different dtype. # float16 rather than float64 so this stays cheap and universally # supported; the CPU pin means ``.double()`` remains probeable. - other = run(torch.device("cpu"), torch.float16) + other = run(torch.device("cpu"), torch.float16) # dtype-ok: fixed probe dtype (dtype-axis contrast) if other is not None and other.dtype == base.dtype: dtype = base.dtype From 58f3757054ffe3fc9f1c43b73b684d38d64ac498 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 16:57:07 +0200 Subject: [PATCH 135/250] Hold R-free to its own tolerance in the AF trajectory test 1VER's final-cycle R-free has a second basin about 0.0042 from the main one. It is rare enough that the five runs sizing the reference never sampled it, so the structure took the 0.002 floor and the basin sat outside a bound meant to cover it -- the same trap the module docstring already records for 6JZA, which was caught only because its basin is common enough to show up in five runs. R-work and R-free now carry separate bounds. R-work keeps the measured tolerance: it is reproducible to a few parts in ten thousand across CPU generation, thread count and torch version, and is what a real change in the refinement moves. R-free takes a 0.01 floor, which clears the basin while staying inside a quarter of every structure's own R-work descent -- test_tolerances_are_tight_enough_to_detect_something now checks the bound actually applied rather than the stored one, so the floor cannot widen the real bound unnoticed. Co-Authored-By: Claude Opus 5 (1M context) --- tests/functional/test_af_trajectory.py | 55 +++++++++++++++++++++----- 1 file changed, 45 insertions(+), 10 deletions(-) diff --git a/tests/functional/test_af_trajectory.py b/tests/functional/test_af_trajectory.py index 755b3998..b242acb6 100644 --- a/tests/functional/test_af_trajectory.py +++ b/tests/functional/test_af_trajectory.py @@ -13,6 +13,13 @@ basins about 0.0074 apart -- and two runs that happen to pick the same basin report a spread 140 times too small. +R-work and R-free are held to different bounds. R-work is reproducible to a few parts in +ten thousand, so it keeps the measured tolerance and is what catches a change that +actually moves the refinement. R-free is computed on the small free set and has a rare +second basin of its own, a few thousandths wide, that a handful of runs will usually +miss; it therefore carries :data:`RFREE_TOLERANCE_FLOOR`, wide enough to sit outside +that basin and still far inside the descent the trajectory shows. + Regenerate deliberately, after a change meant to move these numbers:: ./.dev/bin/python tests/functional/test_af_trajectory.py @@ -40,10 +47,16 @@ #: Multiple of the measured spread a deviation may reach before it counts as a change. SPREAD_MULTIPLE = 3.0 -#: Tolerance floor, so a structure whose runs agree very closely is not held to an -#: unreasonably tight bound. +#: Tolerance floor for R-work, so a structure whose runs agree very closely is not held +#: to an unreasonably tight bound. TOLERANCE_FLOOR = 0.002 +#: Tolerance floor for R-free, which moves in discrete basins rather than jitter. Set +#: above the widest basin separation seen on this set and kept well inside every +#: structure's own R-work descent, so a change large enough to matter still fails. +#: ``test_tolerances_are_tight_enough_to_detect_something`` enforces the second half. +RFREE_TOLERANCE_FLOOR = 0.01 + REFERENCE = Path(__file__).with_name("af_trajectory_reference.json") @@ -79,6 +92,19 @@ def _max_deviation(a, b): return max(max(abs(x[0] - y[0]), abs(x[1] - y[1])) for x, y in zip(a, b)) +def _deviations(a, b): + """Largest absolute difference between two trajectories, ``(R-work, R-free)``.""" + return ( + max(abs(x[0] - y[0]) for x, y in zip(a, b)), + max(abs(x[1] - y[1]) for x, y in zip(a, b)), + ) + + +def _rfree_tolerance(entry): + """The R-free bound: this structure's measured tolerance, floored.""" + return max(float(entry["tolerance"]), RFREE_TOLERANCE_FLOOR) + + @pytest.fixture(scope="module") def reference(): """The committed reference trajectories and their tolerances.""" @@ -105,18 +131,24 @@ def test_af_trajectory_matches_reference(code, reference, test_files_dir): entry = reference["structures"][code] expected = [tuple(point) for point in entry["trajectory"]] - tolerance = float(entry["tolerance"]) + rwork_tolerance = float(entry["tolerance"]) + rfree_tolerance = _rfree_tolerance(entry) observed = trajectory(pdb_path, mtz_path) assert len(observed) == len( expected ), f"{code}: trajectory has {len(observed)} stages, reference has {len(expected)}" - deviation = _max_deviation(observed, expected) - assert deviation <= tolerance, ( - f"{code}: deviates from the reference by {deviation:.6f}, above its measured " - f"tolerance of {tolerance:.6f}.\nobserved={observed}\nreference={expected}" - ) + dev_work, dev_free = _deviations(observed, expected) + for label, deviation, tolerance in ( + ("R-work", dev_work, rwork_tolerance), + ("R-free", dev_free, rfree_tolerance), + ): + assert deviation <= tolerance, ( + f"{code}: {label} deviates from the reference by {deviation:.6f}, above its " + f"tolerance of {tolerance:.6f}.\nobserved={observed}\n" + f"reference={expected}" + ) @pytest.mark.integration @@ -147,8 +179,11 @@ def test_tolerances_are_tight_enough_to_detect_something(reference): entry = reference["structures"][code] series = entry["trajectory"] descent = series[0][0] - series[-1][0] - assert entry["tolerance"] < descent / 4.0, ( - f"{code}: tolerance {entry['tolerance']:.4f} is not small against its " + # Checks the widest bound actually applied, not the stored one: flooring R-free + # would otherwise widen the real bound without this guard seeing it. + widest = max(float(entry["tolerance"]), _rfree_tolerance(entry)) + assert widest < descent / 4.0, ( + f"{code}: tolerance {widest:.4f} is not small against its " f"own R-work descent of {descent:.4f}" ) From a3a14e97caf258e5b56d30edde396f5bbea49011 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 18:27:14 +0200 Subject: [PATCH 136/250] Fix the translation likelihood's variance convention The alignment package carried its own Rice and Woolfson, parameterised by the AMPLITUDE variance, and `fit_sigma_a_per_shell` handed both branches the same number. The two branches do not take the same number: matching term by term against the standard, the centric one is right at `v` and the acentric one needs `v/2`. So acentrics -- 90 to 95% of reflections -- were scored at twice the variance intended, at an effective sigma_A satisfying `D'^2 = 2D^2 - 1`, which is not even real below D = 0.707. Verified by integrating each as a density: at Sigma = 1 with no model the alignment Rice gives = 2.000, where WilsonNormaliser makes = 1 an identity of the fit that produced E_obs. The module's own docstring said the centric variance is "2x larger than acentric" -- the call site contradicted it. `base/targets/xray_likelihoods.rice_per_refl` is correct in BOTH branches from a single complex Sigma, which is what it now uses. That also collapses two calls and a `where` into one, so the half of the work that was computed and discarded goes with it. distributions.py is deleted. Its `stable_log_bessel_i0` had a wrong asymptotic coefficient too -- `-1/(128x^2)` where the expansion has `+1/(16x^2)`, worth -2.9e-5 at x = 50 against -5.4e-7 for the correct term. The shared code uses `log(i0e(z)) + z`, exact and branchless. Nothing outside alignment imported any of it. The clamp on the shared Bessel argument was checked rather than assumed: it caps at 1e6 and the worst case across ten structures is 1.8e3, zero reflections affected. Measured. `rank_by="analytic_r"` is bit-identical, as it must be. `rank_by="llg"` is unchanged on the 30-cell panel at 30/30 and loses one cell on the 40-cell sweep, 2DQ6 t3 at 7.16 -> 22.55 deg -- a cell 0.84 deg inside the gate on the structure that is bimodal at 6/10. The effect is small because sigma_A is fitted against the same likelihood and absorbs part of the error into a compensating D. That narrows the case for the llg default rather than removing it: llg and analytic_r now tie at 36/40 on the sweep (was 37 to 36), and llg still leads 30/30 to 29/30 on the panel with median residual 1.62 against 1.79. Default unchanged. A new unit test integrates the likelihood as a density and asserts = Sigma with no model and Sigma + Fc^2 with one -- the invariance that makes D identifiable. That is the check that would have caught this. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/bessel_arg_range.sh | 17 ++ alignment_lab/analysis/marginal_seeds.sh | 2 +- alignment_lab/analysis/p1_copy_probe.sh | 17 ++ alignment_lab/analysis/scaler_cost.sh | 17 ++ alignment_lab/analysis/sigma_a_cost.sh | 17 ++ alignment_lab/analysis/stage_profile.sh | 31 +++ alignment_lab/diagnostics/bessel_arg_range.py | 61 +++++ alignment_lab/diagnostics/p1_copy_probe.py | 71 ++++++ alignment_lab/diagnostics/scaler_cost.py | 76 +++++++ alignment_lab/diagnostics/sigma_a_cost.py | 79 +++++++ docs/changelog.rst | 2 + .../alignment/test_likelihood_convention.py | 86 +++++++ .../experimental/alignment/distributions.py | 214 ------------------ torchref/experimental/alignment/pipeline.py | 52 ++--- .../experimental/alignment/translation.py | 58 +++-- 15 files changed, 521 insertions(+), 279 deletions(-) create mode 100644 alignment_lab/analysis/bessel_arg_range.sh create mode 100644 alignment_lab/analysis/p1_copy_probe.sh create mode 100644 alignment_lab/analysis/scaler_cost.sh create mode 100644 alignment_lab/analysis/sigma_a_cost.sh create mode 100644 alignment_lab/analysis/stage_profile.sh create mode 100644 alignment_lab/diagnostics/bessel_arg_range.py create mode 100644 alignment_lab/diagnostics/p1_copy_probe.py create mode 100644 alignment_lab/diagnostics/scaler_cost.py create mode 100644 alignment_lab/diagnostics/sigma_a_cost.py create mode 100644 tests/unit/alignment/test_likelihood_convention.py delete mode 100644 torchref/experimental/alignment/distributions.py diff --git a/alignment_lab/analysis/bessel_arg_range.sh b/alignment_lab/analysis/bessel_arg_range.sh new file mode 100644 index 00000000..e891635f --- /dev/null +++ b/alignment_lab/analysis/bessel_arg_range.sh @@ -0,0 +1,17 @@ +#!/bin/bash +#SBATCH --job-name=bessel +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +"$PY" -u alignment_lab/diagnostics/bessel_arg_range.py 2>/dev/null | grep '^ROW' +echo DONE diff --git a/alignment_lab/analysis/marginal_seeds.sh b/alignment_lab/analysis/marginal_seeds.sh index e2bdf603..b9783d63 100644 --- a/alignment_lab/analysis/marginal_seeds.sh +++ b/alignment_lab/analysis/marginal_seeds.sh @@ -26,5 +26,5 @@ export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" for T in $(seq 0 9); do "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial "$T" \ - --arms llg,analytic_r --n-rotation-candidates 25 2>/dev/null | grep '^ROW ' + --arms llg,analytic_r,corr --n-rotation-candidates 25 2>/dev/null | grep '^ROW ' done diff --git a/alignment_lab/analysis/p1_copy_probe.sh b/alignment_lab/analysis/p1_copy_probe.sh new file mode 100644 index 00000000..d70f306a --- /dev/null +++ b/alignment_lab/analysis/p1_copy_probe.sh @@ -0,0 +1,17 @@ +#!/bin/bash +#SBATCH --job-name=p1probe +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +"$PY" -u alignment_lab/diagnostics/p1_copy_probe.py 1DAW 2DQ6 4BX9 2>&1 | grep -E '^ROW|rror|Error' +echo DONE diff --git a/alignment_lab/analysis/scaler_cost.sh b/alignment_lab/analysis/scaler_cost.sh new file mode 100644 index 00000000..0380956b --- /dev/null +++ b/alignment_lab/analysis/scaler_cost.sh @@ -0,0 +1,17 @@ +#!/bin/bash +#SBATCH --job-name=scalercost +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +"$PY" -u alignment_lab/diagnostics/scaler_cost.py 1DAW 2DQ6 2>/dev/null | grep -E '^ROW|rror' +echo DONE diff --git a/alignment_lab/analysis/sigma_a_cost.sh b/alignment_lab/analysis/sigma_a_cost.sh new file mode 100644 index 00000000..f97f3a84 --- /dev/null +++ b/alignment_lab/analysis/sigma_a_cost.sh @@ -0,0 +1,17 @@ +#!/bin/bash +#SBATCH --job-name=sacost +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:55:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +"$PY" -u alignment_lab/diagnostics/sigma_a_cost.py 1DAW 4BX9 2DQ6 2>/dev/null | grep '^ROW' +echo DONE diff --git a/alignment_lab/analysis/stage_profile.sh b/alignment_lab/analysis/stage_profile.sh new file mode 100644 index 00000000..1cd3ae30 --- /dev/null +++ b/alignment_lab/analysis/stage_profile.sh @@ -0,0 +1,31 @@ +#!/bin/bash +# Where does an alignment's wall clock actually go? +# +# Expectation from the component measurements: FRF 1-3 s, and the translation +# stage bounded by one structure-factor call per candidate at well under 100 ms, +# so 25 candidates should be a couple of seconds. That predicts ~5 s and the +# panel medians are 34 s by R and 75-93 s by likelihood. The per-stage timer has +# been in the pipeline the whole time; this reads it. +#SBATCH --job-name=stageprof +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=day +#SBATCH --time=03:00:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +for ARM in analytic_r llg; do + for P in 1DAW 2DQ6; do + echo "########## $P rank_by=$ARM ##########" + "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb $P --trial 0 \ + --arms $ARM --verbose 2 2>/dev/null \ + | sed -n '/^stage /,/^TOTAL/p;/^ROW /p' + done +done +echo DONE diff --git a/alignment_lab/diagnostics/bessel_arg_range.py b/alignment_lab/diagnostics/bessel_arg_range.py new file mode 100644 index 00000000..009e5d4c --- /dev/null +++ b/alignment_lab/diagnostics/bessel_arg_range.py @@ -0,0 +1,61 @@ +"""Does the base Rice's Bessel-argument clamp bite in E-space? + +`xray_likelihoods._rice_body` clamps `2 F_calc F_obs / Sigma` at 1e6. That cap is +sized for F-space; the translation likelihood runs on E values with +`Sigma = 1 - D^2` floored at 1e-4, where the ratio can in principle be far +larger. A clamp that fires silently truncates the likelihood exactly where it is +most discriminating, so this measures the real range instead of reasoning about +the bound. +""" +import sys +from pathlib import Path +import torch +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) +from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 + + +def main(): + from torchref.experimental.alignment.translation import ( + DirectModelEvaluator, TranslationObs, amplitude_translation_search, + normalise_calc, precompute_G_for_rotation) + + for pdb in sys.argv[1:] or list(BENCH_PDBS): + model, data = load_case(pdb) + rot = model.copy().rotate( + random_rotation(seed_for(pdb, 0)).to(model.dtype_float), + center=model.xyz().mean(0)) + rot.spacegroup = data.spacegroup.hm + p1 = rot.copy(); p1.spacegroup = "P 1" + mask = data.get_valid_mask() + sig = getattr(data, "F_sigma", None) + obs = TranslationObs.build(data.F[mask], data.hkl[mask], data.spacegroup, + data.cell, + sig_F=None if sig is None else sig[mask], + n_shells=10) + ev = DirectModelEvaluator(p1) + eye3 = torch.eye(3, dtype=torch.float64) + G, h_R = precompute_G_for_rotation(ev, eye3, obs.hkl, data.spacegroup, + data.cell) + _, _, peaks = amplitude_translation_search( + obs=obs, interpolator=ev, R_rotation=eye3, + spacegroup=data.spacegroup, real_cell=data.cell, grid_steps=16, + n_peaks=1, precomputed_G=G, precomputed_h_R=h_R) + t = torch.as_tensor(peaks[0].translation, dtype=torch.float64, + device=G.device) + ph = torch.exp(2j * torch.pi * torch.einsum( + "ind,d->in", h_R.to(torch.float64), t).to(G.dtype)) + E_calc = normalise_calc((G * ph).sum(dim=0).abs().to(torch.float64), obs) + + # Worst case over the whole sigma_A grid: D -> 0.99, Sigma -> 1 - D^2. + D = 0.99 + Sigma = max(1.0 - D * D, 1e-4) + arg = (2.0 * (D * E_calc) * obs.E_obs / Sigma) + print(f"ROW pdb={pdb} N={obs.E_obs.numel()} " + f"maxE_obs={float(obs.E_obs.max()):.2f} " + f"maxE_calc={float(E_calc.max()):.2f} " + f"max_bessel_arg={float(arg.max()):.3e} " + f"clamped={int((arg > 1e6).sum())}", flush=True) + + +main() diff --git a/alignment_lab/diagnostics/p1_copy_probe.py b/alignment_lab/diagnostics/p1_copy_probe.py new file mode 100644 index 00000000..00d028e8 --- /dev/null +++ b/alignment_lab/diagnostics/p1_copy_probe.py @@ -0,0 +1,71 @@ +"""Does the translation search need a P1 copy of the model, or just apply_symmetry=False? + +`_placement_for_candidate` copies each rotated candidate, sets its space group to +P 1, and wraps it in `DirectModelEvaluator`, because `ModelFT.forward` hardcodes +``apply_symmetry=True`` and the Crowther-Blow expansion needs the SINGLE-MOLECULE +transform ``F_p1(h R_i)`` -- it applies the symmetry itself. + +But the flag exists one level down, on ``SfFFT.compute_structure_factors``. If +calling that with ``apply_symmetry=False`` on the unmodified model agrees with +the P1 copy, then per candidate the copy, the space-group assignment and whatever +they rebuild are all avoidable, and the evaluator wrapper has nothing left to do. + +Reports agreement and the cost of each step, including ``Model.copy`` -- which is +on record as rebuilding a map-symmetry table per symop and being half the cost of +a rotation search. +""" +import sys, time +from pathlib import Path +import torch +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) +from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 + + +def _t(fn, n=3): + fn() + ts = [] + for _ in range(n): + t0 = time.perf_counter(); out = fn(); ts.append(time.perf_counter() - t0) + return min(ts), out + + +def main(): + for pdb in sys.argv[1:] or ["1DAW", "2DQ6"]: + model, data = load_case(pdb) + rot = model.copy().rotate( + random_rotation(seed_for(pdb, 0)).to(model.dtype_float), + center=model.xyz().mean(0)) + rot.spacegroup = data.spacegroup.hm + hkl = data.hkl[data.get_valid_mask()] + hkl_i = hkl.round().to(torch.int64) + + # what the pipeline does now + t_copy, p1 = _t(lambda: rot.copy()) + t_sg, _ = _t(lambda: setattr(p1, "spacegroup", "P 1")) + t_p1_sf, F_p1 = _t(lambda: p1(hkl_i)) + + # the candidate replacement: same model, symmetry off at the FFT + def direct(): + sf, _ = rot.fft.compute_structure_factors( + hkl_i, *rot.get_iso(), *rot.get_aniso(), apply_symmetry=False) + return sf + t_direct, F_direct = _t(direct) + + # and with symmetry ON, for contrast -- this is NOT what the TF wants + t_sym, F_sym = _t(lambda: rot(hkl_i)) + + a, b = F_p1.to(torch.complex128), F_direct.to(torch.complex128) + num = (a - b).abs().max() + rel = float(num / a.abs().max().clamp(min=1e-30)) + agree_sym = float((a - F_sym.to(torch.complex128)).abs().max() + / a.abs().max().clamp(min=1e-30)) + print(f"ROW pdb={pdb} N={hkl_i.shape[0]} " + f"copy={1000*t_copy:.1f}ms set_sg={1000*t_sg:.1f}ms " + f"p1_sf={1000*t_p1_sf:.1f}ms direct_sf={1000*t_direct:.1f}ms " + f"sym_sf={1000*t_sym:.1f}ms " + f"| max_rel_diff(p1, apply_symmetry=False)={rel:.3e} " + f"max_rel_diff(p1, symmetry_on)={agree_sym:.3e}", flush=True) + + +main() diff --git a/alignment_lab/diagnostics/scaler_cost.py b/alignment_lab/diagnostics/scaler_cost.py new file mode 100644 index 00000000..c2332c5c --- /dev/null +++ b/alignment_lab/diagnostics/scaler_cost.py @@ -0,0 +1,76 @@ +"""Sixteen parameters, 8.6 seconds. Where does Scaler.refine_lbfgs spend it? + +`refine_lbfgs` builds its x-ray target with ``model=None`` and passes a detached +``fcalc`` per closure call, and the comment there states the fit never recomputes +structure factors. This counts them rather than trusting that, and times the +three phases separately -- construction, ``initialize``, and the fit -- because +the solvent contribution is also refined and a mask rebuilt per closure call +would look identical from outside. +""" +import sys, time +from pathlib import Path +import torch +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(True) +from lab import BENCH_PDBS, load_case # noqa: E402 + + +def main(): + from torchref.scaling import Scaler + from torchref.base.metrics.rfactor import rfactor_work_free + import torchref.model.model_ft as mft + import torchref.scaling.solvent as solv + + for pdb in sys.argv[1:] or ["1DAW"]: + model, data = load_case(pdb) + N = data.hkl.shape[0] + + counts = {"model_forward": 0, "solvent_forward": 0} + orig_fwd = mft.ModelFT.forward + def counted_fwd(self, *a, **kw): + counts["model_forward"] += 1 + return orig_fwd(self, *a, **kw) + mft.ModelFT.forward = counted_fwd + orig_sol = solv.SolventModel.forward + def counted_sol(self, *a, **kw): + counts["solvent_forward"] += 1 + return orig_sol(self, *a, **kw) + solv.SolventModel.forward = counted_sol + + t0 = time.perf_counter() + s = Scaler(model=model, data=data, nbins=20, verbose=0) + t_ctor = time.perf_counter() - t0 + + t0 = time.perf_counter() + with torch.no_grad(): + fc = model(data.hkl).detach() + t_fcalc = time.perf_counter() - t0 + n_after_fcalc = counts["model_forward"] + + t0 = time.perf_counter() + s.initialize(fc) + t_init = time.perf_counter() - t0 + n_after_init = counts["model_forward"] + sol_after_init = counts["solvent_forward"] + + t0 = time.perf_counter() + s.refine_lbfgs(fcalc=fc) + t_fit = time.perf_counter() - t0 + + t0 = time.perf_counter() + with torch.no_grad(): + rw, _ = rfactor_work_free(data, torch.abs(s.forward(fc))) + t_r = time.perf_counter() - t0 + + mft.ModelFT.forward = orig_fwd + solv.SolventModel.forward = orig_sol + print(f"ROW pdb={pdb} N={N} ctor={t_ctor:.2f}s fcalc={t_fcalc:.2f}s " + f"init={t_init:.2f}s fit={t_fit:.2f}s rfac={t_r:.2f}s " + f"total={t_ctor+t_fcalc+t_init+t_fit+t_r:.2f}s " + f"| model_fwd_during_fit={counts['model_forward']-n_after_init} " + f"(1 expected: the explicit fcalc) " + f"solvent_fwd_during_fit={counts['solvent_forward']-sol_after_init} " + f"R={float(rw):.4f}", flush=True) + + +main() diff --git a/alignment_lab/diagnostics/sigma_a_cost.py b/alignment_lab/diagnostics/sigma_a_cost.py new file mode 100644 index 00000000..27ba5fbe --- /dev/null +++ b/alignment_lab/diagnostics/sigma_a_cost.py @@ -0,0 +1,79 @@ +"""Where does the likelihood ranking's 2.5x actually go? + +`fit_sigma_a_per_shell` scans 81 values of D against every reflection, in +float64, evaluating BOTH the Rice and the Woolfson branch everywhere and then +discarding half of each with a `where`. On 2DQ6 that is 81 x 228197 = 18.5M +elements per tensor and four of them materialised, per candidate, times 25 +candidates. + +This times the pieces so the fix is aimed rather than guessed: the sigma_A fit, +the Wilson normalisation of the calculated side, the likelihood evaluation +itself, and the translation refine they sit alongside. +""" +import sys, time +from pathlib import Path +import torch +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +torch.set_grad_enabled(False) +from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 + + +def _t(fn, n=3): + fn() + ts = [] + for _ in range(n): + t0 = time.perf_counter(); out = fn(); ts.append(time.perf_counter() - t0) + return min(ts), out + + +def main(): + from torchref.experimental.alignment.translation import ( + DirectModelEvaluator, TranslationObs, amplitude_translation_search, + correlation_at, fit_sigma_a_per_shell, llg_at, local_translation_refine, + normalise_calc, precompute_G_for_rotation) + + for pdb in sys.argv[1:] or ["1DAW", "2DQ6"]: + seed = seed_for(pdb, 0) + model, data = load_case(pdb) + rot = model.copy().rotate(random_rotation(seed).to(model.dtype_float), + center=model.xyz().mean(0)) + rot.spacegroup = data.spacegroup.hm + p1 = rot.copy(); p1.spacegroup = "P 1" + mask = data.get_valid_mask() + sig = getattr(data, "F_sigma", None) + obs = TranslationObs.build(data.F[mask], data.hkl[mask], data.spacegroup, + data.cell, + sig_F=None if sig is None else sig[mask], + n_shells=10) + ev = DirectModelEvaluator(p1); eye3 = torch.eye(3, dtype=torch.float64) + G, h_R = precompute_G_for_rotation(ev, eye3, obs.hkl, data.spacegroup, + data.cell) + _, _, peaks = amplitude_translation_search( + obs=obs, interpolator=ev, R_rotation=eye3, + spacegroup=data.spacegroup, real_cell=data.cell, grid_steps=16, + n_peaks=20, precomputed_G=G, precomputed_h_R=h_R) + t0 = torch.as_tensor(peaks[0].translation, dtype=torch.float64) + + t_ref, _ = _t(lambda: local_translation_refine( + obs=obs, interpolator=ev, R_rotation=eye3, + spacegroup=data.spacegroup, real_cell=data.cell, t_init=t0, + radius=0.06, grid_steps=13, n_refinement_passes=1, + precomputed_G=G, precomputed_h_R=h_R)) + ph = torch.exp(2j * torch.pi * torch.einsum( + "ind,d->in", h_R.to(torch.float64), t0.to(G.device)).to(G.dtype)) + Fc = (G * ph).sum(dim=0).abs().to(torch.float64) + t_norm, E_calc = _t(lambda: normalise_calc(Fc, obs)) + t_sa, sa = _t(lambda: fit_sigma_a_per_shell( + obs.E_obs, E_calc, obs.centric, obs.shell_idx, obs.n_shells, + n_grid=81)) + t_llg, _ = _t(lambda: llg_at(obs, G, h_R, t0, sa)) + t_corr, _ = _t(lambda: correlation_at(obs, G, h_R, t0)) + N = obs.hkl.numel() // 3 + print(f"ROW pdb={pdb} N={N} shells={obs.n_shells} " + f"refine={1000*t_ref:.1f}ms norm_calc={1000*t_norm:.1f}ms " + f"sigma_a={1000*t_sa:.1f}ms llg={1000*t_llg:.1f}ms " + f"corr={1000*t_corr:.1f}ms " + f"grid_elems={81*N/1e6:.1f}M", flush=True) + + +main() diff --git a/docs/changelog.rst b/docs/changelog.rst index 46a663dc..54266255 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,8 @@ Changelog Unreleased ---------- +- Fixed the translation likelihood's variance convention, which scored acentric reflections at twice the variance intended -- 90-95% of reflections. The alignment package carried its own Rice and Woolfson parameterised by the *amplitude* variance and handed both branches the same number, where the acentric branch needs half what the centric one does. It now uses ``base.targets.xray_likelihoods.rice_per_refl``, which takes the complex variance and derives the centric case from it +- Removed ``experimental/alignment/distributions.py``. Its ``stable_log_bessel_i0`` also carried a wrong asymptotic coefficient, giving -2.9e-5 at x = 50 against -5.4e-7 for the correct term; the shared implementation uses ``log(i0e(z)) + z``, which is exact - Molecular-replacement candidates are ranked by the translation function's likelihood, not by an analytical-scale R-factor. 30/30 on the ten-structure panel against 29/30, and 37/40 against 36/40 over a ten-seed sweep, with tighter placements on the cells both solve. Selectable through ``rank_by``; the correlation is the third option and is the worst of the three at 32/40, despite a rank-level harness rating it best on a truth label that disagrees with coordinate superposition - The likelihood ranking costs about 2.5x the wall clock, from a per-candidate sigma_A fit - The molecular-replacement pipeline's ``verbose`` levels are a documented contract routed through one emitter, rather than ``if verbose > 0: print(...)`` at seventeen sites. Level 2 emits one machine-readable ``CAND`` line per rotation candidate carrying every score the selection could have used, so a wrong placement can be diagnosed from the run itself instead of from a harness that re-implements the placement loop and then disagrees with it diff --git a/tests/unit/alignment/test_likelihood_convention.py b/tests/unit/alignment/test_likelihood_convention.py new file mode 100644 index 00000000..7a85fb56 --- /dev/null +++ b/tests/unit/alignment/test_likelihood_convention.py @@ -0,0 +1,86 @@ +"""The translation likelihood's variance convention, pinned as a pdf. + +A likelihood is only right up to what its variance argument means, and that is +exactly the sort of thing that survives code review and unit tests written +against the implementation rather than against the distribution. The alignment +package carried its own Rice and Woolfson for a while, parameterised by the +*amplitude* variance, and handed both branches the same number -- which put +acentrics at twice their intended variance on 90-95% of reflections, for as long +as nobody integrated the thing. + +These tests integrate it. They assert the property the pipeline actually depends +on: at ``D = 0`` the likelihood must believe `` = 1``, because +:class:`~torchref.scaling.WilsonNormaliser` makes `` = 1`` an identity of +the fit that produced ``E_obs``. Anything else means the model and the data +disagree about the scale of the very quantity being compared. +""" +import pytest +import torch + +from torchref.base.targets.xray_likelihoods import rice_per_refl + +pytestmark = pytest.mark.unit + + +def _pdf_moments(logp, F, dF): + """Norm and second moment of ``exp(logp)`` treated as a density in ``F``.""" + p = torch.exp(logp) + norm = float((p * dF).sum()) + return norm, float((p * F**2 * dF).sum()) / norm + + +@pytest.fixture(scope="module") +def grid(): + F = torch.linspace(1e-6, 12.0, 200001, dtype=torch.float64) + return F, float(F[1] - F[0]) + + +@pytest.mark.parametrize("centric", [False, True]) +def test_unit_sigma_means_unit_second_moment(grid, centric): + """At Sigma = 1 and no model, the likelihood expects = 1. + + This is the property the whole sigma_A path rests on. Both branches must + satisfy it from the SAME Sigma -- that is what makes a single complex + variance the right parameterisation and an amplitude variance the wrong one. + """ + F, dF = grid + ll = -rice_per_refl(F, torch.zeros_like(F), torch.ones_like(F), + torch.full_like(F, centric, dtype=torch.bool)) + norm, m2 = _pdf_moments(ll, F, dF) + assert norm == pytest.approx(1.0, rel=1e-4), "not a normalised density" + assert m2 == pytest.approx(1.0, rel=1e-4), ( + f"{'centric' if centric else 'acentric'} branch expects = {m2:.4f} " + f"at Sigma = 1; the observations have = 1 by construction" + ) + + +@pytest.mark.parametrize("sigma", [0.25, 0.5, 2.0]) +@pytest.mark.parametrize("centric", [False, True]) +def test_second_moment_tracks_sigma(grid, sigma, centric): + """ = Sigma with no model, for both branches. Fixes the scale, not just the shape.""" + F, dF = grid + ll = -rice_per_refl(F, torch.zeros_like(F), torch.full_like(F, sigma), + torch.full_like(F, centric, dtype=torch.bool)) + norm, m2 = _pdf_moments(ll, F, dF) + assert norm == pytest.approx(1.0, rel=1e-4) + assert m2 == pytest.approx(sigma, rel=1e-4) + + +@pytest.mark.parametrize("centric", [False, True]) +def test_second_moment_with_a_model_present(grid, centric): + """With a model, = Sigma + Fc^2 -- the signal adds to the noise. + + The sigma_A likelihood is evaluated at ``Fc = D E_calc`` and + ``Sigma = 1 - D^2``, so on data normalised to `` = 1`` this gives + `` = 1`` for every D. That invariance is why D is identifiable at all, + and it fails if the two branches disagree about what the variance means. + """ + F, dF = grid + D, E_calc = 0.6, 1.0 + Fc, Sigma = D * E_calc, 1.0 - D * D + ll = -rice_per_refl(F, torch.full_like(F, Fc), torch.full_like(F, Sigma), + torch.full_like(F, centric, dtype=torch.bool)) + norm, m2 = _pdf_moments(ll, F, dF) + assert norm == pytest.approx(1.0, rel=1e-4) + assert m2 == pytest.approx(Sigma + Fc**2, rel=1e-4) + assert m2 == pytest.approx(1.0, rel=1e-4), "D must not change the expected " diff --git a/torchref/experimental/alignment/distributions.py b/torchref/experimental/alignment/distributions.py deleted file mode 100644 index d7ece281..00000000 --- a/torchref/experimental/alignment/distributions.py +++ /dev/null @@ -1,214 +0,0 @@ -""" -Statistical distributions for Maximum Likelihood molecular replacement. - -This module provides numerically stable implementations of the probability -distributions used in crystallographic ML target functions: - -- Rice distribution for acentric reflections -- Woolfson (folded normal) distribution for centric reflections -- Stable log-Bessel I_0 computation for large arguments - -The key numerical challenge is computing log(I_0(x)) for large x (up to ~10000), -which requires an asymptotic expansion to avoid overflow. - -References ----------- -- Read, R.J. (2001). Pushing the boundaries of molecular replacement with - maximum likelihood. Acta Cryst. D57, 1373-1382. -- McCoy et al. (2007). Phaser crystallographic software. J. Appl. Cryst. 40, 658-674. -""" - -import math - -import torch - - -def stable_log_bessel_i0(x: torch.Tensor) -> torch.Tensor: - """ - Compute log(I_0(x)) with numerical stability for large x. - - The modified Bessel function I_0(x) grows exponentially, so direct - computation overflows for x > ~700. This implementation uses: - - For small x (< 50): log(I_0e(x)) + |x| using torch.special.i0e - - For large x (>= 50): asymptotic expansion x - 0.5*log(2*pi*x) - - Parameters - ---------- - x : torch.Tensor - Input tensor of non-negative values. - - Returns - ------- - torch.Tensor - log(I_0(x)) for each element, same shape as input. - - Notes - ----- - The asymptotic expansion is: - log(I_0(x)) ~ x - 0.5*log(2*pi*x) + 1/(8x) - 1/(128x^2) + ... - - For x >= 50, the leading two terms give excellent accuracy. - For very large x (> 10000), higher order terms can be added if needed. - - Examples - -------- - :: - - x = torch.tensor([0.1, 10.0, 100.0, 1000.0]) - log_i0 = stable_log_bessel_i0(x) - torch.all(torch.isfinite(log_i0)) - True - """ - # Ensure non-negative input - x_abs = torch.abs(x) - - # Initialize output - result = torch.zeros_like(x) - - # Small x regime: use i0e which is I_0(x) * exp(-|x|) - # So log(I_0(x)) = log(I_0e(x)) + |x| - small_mask = x_abs < 50.0 - if small_mask.any(): - x_small = x_abs[small_mask] - i0e_val = torch.special.i0e(x_small) - # Guard against i0e returning 0 for very small x - i0e_val = torch.clamp(i0e_val, min=1e-300) - result[small_mask] = torch.log(i0e_val) + x_small - - # Large x regime: asymptotic expansion - # log(I_0(x)) ~ x - 0.5*log(2*pi*x) + 1/(8x) - 1/(128x^2) + ... - large_mask = ~small_mask - if large_mask.any(): - x_large = x_abs[large_mask] - # Leading terms of asymptotic expansion - log_i0_asymp = ( - x_large - - 0.5 * torch.log(2.0 * math.pi * x_large) - + 1.0 / (8.0 * x_large) - - 1.0 / (128.0 * x_large * x_large) - ) - result[large_mask] = log_i0_asymp - - return result - - -def rice_log_likelihood( - F_obs: torch.Tensor, - F_mean: torch.Tensor, - variance: torch.Tensor, -) -> torch.Tensor: - """ - Log-likelihood for acentric reflections (Rice distribution). - - The Rice distribution describes the distribution of |F_obs| given - the expected |F_calc| and variance for acentric reflections: - - p(F_obs | F_mean, sigma^2) = (F_obs / sigma^2) * - exp(-(F_obs^2 + F_mean^2) / (2*sigma^2)) * I_0(F_obs*F_mean / sigma^2) - - Parameters - ---------- - F_obs : torch.Tensor - Observed structure factor amplitudes |F_obs|. - F_mean : torch.Tensor - Expected structure factor amplitudes D * |F_calc|. - variance : torch.Tensor - Variance parameter sigma^2 = epsilon * (Sigma_N - D^2 * <|F_calc|^2>). - - Returns - ------- - torch.Tensor - Log-likelihood for each reflection. - - Notes - ----- - The log-likelihood is: - log p = log(F_obs) - log(variance) - (F_obs^2 + F_mean^2)/(2*variance) - + log(I_0(F_obs * F_mean / variance)) - - Numerical stability is ensured by using stable_log_bessel_i0. - """ - # Numerical guards - variance_safe = torch.clamp(variance, min=1e-8) - F_obs_safe = torch.clamp(F_obs, min=1e-10) - F_mean_safe = torch.clamp(F_mean, min=0.0) - - # Compute Bessel argument - bessel_arg = F_obs_safe * F_mean_safe / variance_safe - - # Log-likelihood components - log_likelihood = ( - torch.log(F_obs_safe) - - torch.log(variance_safe) - - (F_obs_safe**2 + F_mean_safe**2) / (2.0 * variance_safe) - + stable_log_bessel_i0(bessel_arg) - ) - - return log_likelihood - - -def woolfson_log_likelihood( - F_obs: torch.Tensor, - F_mean: torch.Tensor, - variance: torch.Tensor, -) -> torch.Tensor: - """ - Log-likelihood for centric reflections (folded normal / Woolfson distribution). - - For centric reflections, the phase is restricted to 0 or pi, so the - distribution becomes a folded normal (Woolfson distribution): - - p(F_obs | F_mean, sigma^2) = (2/sigma) * (2*pi)^(-0.5) * - cosh(F_obs * F_mean / sigma^2) * exp(-(F_obs^2 + F_mean^2)/(2*sigma^2)) - - Parameters - ---------- - F_obs : torch.Tensor - Observed structure factor amplitudes |F_obs|. - F_mean : torch.Tensor - Expected structure factor amplitudes D * |F_calc|. - variance : torch.Tensor - Variance parameter (note: 2x larger than acentric due to phase restriction). - - Returns - ------- - torch.Tensor - Log-likelihood for each centric reflection. - - Notes - ----- - The log-likelihood is: - log p = log(2) - 0.5*log(2*pi*variance) - (F_obs^2 + F_mean^2)/(2*variance) - + log(cosh(F_obs * F_mean / variance)) - - For numerical stability with large arguments: - log(cosh(x)) ~ |x| - log(2) for |x| > ~20 - """ - # Numerical guards - variance_safe = torch.clamp(variance, min=1e-8) - F_obs_safe = torch.clamp(F_obs, min=1e-10) - F_mean_safe = torch.clamp(F_mean, min=0.0) - - # For centric, variance is 2x larger (sigma^2 / 2 for each component) - # The effective sigma for centric is sqrt(2 * variance) - sigma = torch.sqrt(variance_safe) - - # Compute argument for cosh - cosh_arg = F_obs_safe * F_mean_safe / variance_safe - - # Stable log(cosh(x)): use |x| - log(2) for large x - log_cosh = torch.where( - cosh_arg < 20.0, - torch.log(torch.cosh(cosh_arg)), - torch.abs(cosh_arg) - math.log(2.0), - ) - - # Log-likelihood components - log_likelihood = ( - math.log(2.0) - - 0.5 * torch.log(2.0 * math.pi * variance_safe) - - (F_obs_safe**2 + F_mean_safe**2) / (2.0 * variance_safe) - + log_cosh - ) - - return log_likelihood diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 160f04c9..533ec694 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -155,7 +155,7 @@ def cluster_rotation_peaks( # --------------------------------------------------------------------------- -# Stage timing and the user-facing R-work +# Stage timing # --------------------------------------------------------------------------- @@ -221,34 +221,6 @@ def summary(self) -> str: return "\n".join(lines) -def _external_rwork(model: "ModelFT", data: "ReflectionData") -> float: - """Full-resolution scaled R-work via the standard Scaler. - - The TF + local refine work in analytical-scale R-factor (which ranks - candidates correctly but isn't the user-facing R-work). We compute the - proper Scaler-fit R-work once per finalist. - """ - from ...base.metrics.rfactor import rfactor_work_free - from ...scaling import Scaler - - # No device override: the Scaler takes the configured default, which is - # the one place a device is decided. Reading it off whichever tensor is - # nearest is what puts a run on two devices at once. - s = Scaler(model=model, data=data, nbins=20, verbose=0) - # Detach the model forward — the scaler only needs gradients through its - # own parameters; leaving `fc` attached to the model's autograd graph - # keeps SfFFT density-build intermediates alive after this function - # returns. - with torch.no_grad(): - fc = model(data.hkl).detach() - s.initialize(fc) - s.refine_lbfgs(fcalc=fc) - with torch.no_grad(): - # rfactor_work_free takes already-scaled amplitudes, not complex F_calc. - rw, _ = rfactor_work_free(data, torch.abs(s.forward(fc))) - return rw.item() if hasattr(rw, "item") else float(rw) - - @dataclass class MRSolution: """A molecular-replacement placement. @@ -268,8 +240,8 @@ class MRSolution: better. Reported, not ranked -- see the sort in :meth:`run`. r_factor : float The analytical-scale R at that placement, lower better. Reported, not - ranked; for the returned winner it is replaced by the solvent-aware - Scaler R-work, which is the number a caller reads. + ranked. A single global scale, so it is not the number a full Scaler + would return -- build one on the returned model if that is wanted. llg_score : float **The ranking key**: the translation likelihood at that placement, higher better. ``nan`` when ``rank_by`` is not ``"llg"``, since it costs @@ -570,17 +542,19 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: solutions.sort(key=lambda s: -s.llg_score) winner = solutions[0] - # Single solvent-aware Scaler refit on the winner for the user-facing R. - timer.start("12_final_scaler") - rwork_final = _external_rwork(winner.model, self.data) - timer.stop("12_final_scaler") - winner.model.last_alignment_rfactor = rwork_final - winner.r_factor = rwork_final + # No solvent-aware Scaler refit. It used to run here on the winner to + # report an R-work, and cost about a third of the whole alignment -- 8.6 + # of 28 seconds on 2DQ6 -- to fit sixteen scaling parameters that change + # nothing about which placement is returned. The pipeline's contract is + # a placement; downstream refinement fits its own scaler properly, and + # doing a worse version of that here to print a number is not worth a + # third of the runtime. A caller that wants an R-work can build a + # `Scaler` on the returned model. + winner.model.last_alignment_rfactor = winner.r_factor self._log(1, f"mr: winner ({self.rank_by}) " f"LLG={winner.llg_score:.1f} " f"TF corr={winner.translation_score:.5f} " - f"analytic R={winner.r_factor:.4f}, " - f"final Scaler-fit R-work={rwork_final:.4f}") + f"analytic R={winner.r_factor:.4f}") self._log(2, "\n" + timer.summary()) return solutions diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index 57b4bf54..2139027b 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -24,12 +24,12 @@ import numpy as np import torch +from torchref.base.targets.xray_likelihoods import rice_per_refl from torchref.config import get_default_device from torchref.scaling import WilsonNormaliser from torchref.scaling.weighting import (inverse_variance_weight, normalise_weight, snr_from_amplitude) -from .distributions import rice_log_likelihood, woolfson_log_likelihood from .sh import assign_shells, equal_count_shell_edges from dataclasses import dataclass from typing import List, Optional, Tuple, TYPE_CHECKING @@ -490,17 +490,21 @@ def fit_sigma_a_per_shell( N = E_obs.numel() D_grid = torch.linspace(0.0, 0.99, n_grid, device=device, dtype=dtype) # (G,) F_mean = D_grid.view(-1, 1) * E_calc.view(1, -1) # (G, N) - var_d = (1.0 - D_grid * D_grid).clamp(min=1e-4) # (G,) + # Sigma is the COMPLEX variance. `rice_per_refl` derives the centric + # amplitude variance from it internally, which is the whole reason to use it + # -- the previous code passed one amplitude variance to a Rice and a + # Woolfson written in different conventions, leaving acentrics at twice the + # variance they should have had. + Sigma = (1.0 - D_grid * D_grid).clamp(min=1e-4) # (G,) if interp_var is None: - var_full = var_d.view(-1, 1).expand(n_grid, N) + Sigma_full = Sigma.view(-1, 1).expand(n_grid, N) else: - var_full = (var_d.view(-1, 1) + interp_var.view(1, -1)).clamp(min=1e-4) + Sigma_full = (Sigma.view(-1, 1) + interp_var.view(1, -1)).clamp(min=1e-4) E_obs_full = E_obs.view(1, -1).expand(n_grid, N) - ll_acent = rice_log_likelihood(E_obs_full, F_mean, var_full) - ll_cent = woolfson_log_likelihood(E_obs_full, F_mean, var_full) - cent_full = centric.view(1, -1) - ll = torch.where(cent_full, ll_cent, ll_acent) # (G, N) + # Negated: `rice_per_refl` is an NLL and everything here maximises. + ll = -rice_per_refl(E_obs_full, F_mean, Sigma_full, + centric.view(1, -1).expand(n_grid, N)) # (G, N) # Sum per shell, take argmax over the D-grid. shell_idx_gn = shell_idx.view(1, -1).expand(n_grid, N) @@ -524,10 +528,14 @@ def llg_translation_rescore( F_calc(h, t) = sum_i G_i(h) exp(2 pi i (h R_i).t) E_calc(h, t) = |F_calc(h, t)| / sqrt(Sigma_calc(s; t)) - LLG(t) = sum_h [LL(E_obs, D E_calc, var) - LL_Wilson(E_obs)] + LLG(t) = sum_h [LL(E_obs, D E_calc, Sigma) - LL_Wilson(E_obs)] - with ``var = (1 - D^2) + interp_var``. The Rice branch is used for acentric - reflections and Woolfson for centric. + with ``Sigma = (1 - D^2)`` the **complex** variance. The acentric/centric + split is handled inside + :func:`~torchref.base.targets.xray_likelihoods.rice_per_refl`, which derives + the centric amplitude variance from the same ``Sigma`` -- the two are not + the same number, and passing one amplitude variance to both branches is how + this used to score acentrics at twice the variance they should have. The scoring rule the amplitude correlation is a pre-filter for. At rank level it is the strongest discriminator measured -- truth at rank 0 in 27 of @@ -589,6 +597,7 @@ def llg_translation_rescore( sigma_a_d = sigma_a.to(device).to(real_dtype) # (n_shells,) D_per_refl = sigma_a_d.index_select(0, shell_idx_l) # (N,) + # Complex variance, matching `rice_per_refl`'s convention. var_d = (1.0 - D_per_refl * D_per_refl).clamp(min=1e-4) # (N,) if interp_var is not None: var_per_refl = (var_d + interp_var.to(device).to(real_dtype)).clamp(min=1e-4) @@ -596,21 +605,20 @@ def llg_translation_rescore( var_per_refl = var_d F_mean = D_per_refl.view(1, N) * E_calc # (K, N) - var_full = var_per_refl.view(1, N).expand(K, N) + Sigma_full = var_per_refl.view(1, N).expand(K, N) E_obs_full = E_obs.view(1, N).expand(K, N) - cent_full = centric.to(device).to(torch.bool).view(1, N) - - ll_acent = rice_log_likelihood(E_obs_full, F_mean, var_full) - ll_cent = woolfson_log_likelihood(E_obs_full, F_mean, var_full) - ll = torch.where(cent_full, ll_cent, ll_acent) # (K, N) - - # Wilson reference (data only): F_mean = 0, var = 1. - var0 = torch.ones_like(E_obs) - F_mean0 = torch.zeros_like(E_obs) - ll_wil_acent = rice_log_likelihood(E_obs, F_mean0, var0) - ll_wil_cent = woolfson_log_likelihood(E_obs, F_mean0, var0) - ll_wil_per_refl = torch.where(centric.to(device).to(torch.bool), - ll_wil_cent, ll_wil_acent) + cent = centric.to(device).to(torch.bool) + cent_full = cent.view(1, N).expand(K, N) + + ll = -rice_per_refl(E_obs_full, F_mean, Sigma_full, cent_full) # (K, N) + + # Wilson reference (data only): no model, so F_mean = 0 and Sigma = 1 -- + # which is = 1, the identity WilsonNormaliser fits to. Under the + # amplitude-variance convention this line used to carry, unit Sigma meant + # = 2 for acentrics and the reference was inconsistent with the data + # it referenced. + ll_wil_per_refl = -rice_per_refl( + E_obs, torch.zeros_like(E_obs), torch.ones_like(E_obs), cent) ll_wil_total = ll_wil_per_refl.sum() return ll.sum(dim=1) - ll_wil_total # (K,) From ac496d9e344d4f7c5d430f322c1a44246fc58c6b Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Tue, 1 Sep 2026 19:08:07 +0200 Subject: [PATCH 137/250] restore now lands on default or passedmdevice --- torchref/model/model.py | 43 +++++++++++++++++++++++++------------- torchref/model/model_ft.py | 17 ++++++++------- 2 files changed, 39 insertions(+), 21 deletions(-) diff --git a/torchref/model/model.py b/torchref/model/model.py index 7983ef68..87387329 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -19,7 +19,12 @@ import torch.nn as nn from torchref.base import math_torch -from torchref.config import canonical_device, get_float_dtype, normalize_device +from torchref.config import ( + canonical_device, + get_default_device, + get_float_dtype, + normalize_device, +) from torchref.io import cif, pdb from torchref.model.context import ModelContext from torchref.model.parameter_wrappers import ( @@ -2142,7 +2147,7 @@ def save_state(self, path: str): if self.ctx.verbose > 0: print(f"Saved model state to {path}") - def load_state(self, path: str, strict: bool = True): + def load_state(self, path: str, strict: bool = True, device=None): """ Load the complete state of the model from a file. @@ -2153,10 +2158,14 @@ def load_state(self, path: str, strict: bool = True): strict : bool, optional Accepted for signature compatibility; the restore goes through :meth:`create_from_state_dict`, which is never strict. + device : torch.device, optional + Device to restore onto. Defaults to this model's current device, so an + in-place reload keeps its placement; pass one to restore elsewhere. """ - state_dict = torch.load(path, map_location=self.device, weights_only=False) + target_device = self.device if device is None else device + state_dict = torch.load(path, map_location=target_device, weights_only=False) loaded = type(self).create_from_state_dict( - state_dict, device=self.device, verbose=self.ctx.verbose + state_dict, device=target_device, verbose=self.ctx.verbose ) # Adopt the fully-built model's state wholesale. self.__dict__.update(loaded.__dict__) @@ -2343,9 +2352,9 @@ def create_from_state_dict( State dictionary from torch.save(model.state_dict(), ...). device : torch.device, optional Move the restored model here once it is built. The restore itself always - runs on CPU, and ``None`` leaves it there rather than resolving to - ``device.current`` -- loading a file is not a reason to claim an - accelerator. Move it yourself, or pass one here. + runs on CPU; ``None`` then moves it to the configured default device + (``get_default_device()``), so a round-trip lands beside a same-config + model rather than stranding itself on CPU. Pass a device to override. verbose : int, optional Verbosity level. Default is 1. dtype_float : torch.dtype, optional @@ -2362,11 +2371,15 @@ def create_from_state_dict( anisotropic ``u`` is rebuilt as a :class:`CholeskyMixedTensor`, matching :meth:`load`, so the positive-definite parametrization round-trips. """ - # Build on CPU throughout, then move once at the end if the caller named a - # device. One device for the whole model is the invariant that matters: the - # wrappers are built from the atom table and land on CPU whatever is asked for, - # so resolving an accelerator up front splits the model rather than placing it. - target_device = canonical_device(device) if device is not None else None + # Build on CPU throughout, then move once at the end -- to the caller's device + # if they named one, otherwise to the configured default device, so a restore + # lands beside a same-config model instead of stranding itself on CPU. One + # device for the whole model is the invariant that matters: the wrappers are + # built from the atom table and land on CPU whatever is asked for, so resolving + # an accelerator up front splits the model rather than placing it. + target_device = ( + canonical_device(device) if device is not None else get_default_device() + ) device = torch.device("cpu") if dtype_float is None: dtype_float = get_float_dtype() @@ -2407,8 +2420,10 @@ def create_from_state_dict( } instance.load_state_dict(state_dict, strict=False) - if target_device is not None: - instance.to(target_device) + # Always placed: target_device is the caller's device or the configured default, + # never None. Without this the restore used to stay on CPU and split a + # round-trip's restored model from its (default-device) source. + instance.to(target_device) if verbose > 0: n_atoms = len(instance.pdb) if instance.pdb is not None else 0 diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index c6917cba..efbb99e6 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -13,7 +13,7 @@ import torch from torchref.base.fourier import fft, ifft -from torchref.config import canonical_device, dtypes, get_float_dtype +from torchref.config import canonical_device, dtypes, get_default_device, get_float_dtype from torchref.model.model import Model from torchref.model.sf_fft import SfFFT from torchref.symmetry import SpaceGroup @@ -927,7 +927,7 @@ def create_from_state_dict( State dictionary from torch.save(model.state_dict(), ...). device : torch.device, optional Move the restored model here once it is built. The restore itself always - runs on CPU, and ``None`` leaves it there; see + runs on CPU; ``None`` then moves it to the configured default device; see :meth:`Model.create_from_state_dict`. verbose : int, optional Verbosity level. Default is 1. @@ -947,9 +947,12 @@ def create_from_state_dict( :meth:`load`, so the positive-definite parametrization round-trips. """ # Build on CPU throughout and move once at the end, as Model does; the grid - # setup below otherwise sizes an accelerator allocation for a model the caller - # has not asked to put there. - target_device = canonical_device(device) if device is not None else None + # setup below otherwise sizes an accelerator allocation before the model is + # placed. The final target is the caller's device, or the configured default + # when they name none, so a restore lands beside a same-config model. + target_device = ( + canonical_device(device) if device is not None else get_default_device() + ) device = torch.device("cpu") if dtype_float is None: dtype_float = get_float_dtype() @@ -1036,8 +1039,8 @@ def create_from_state_dict( instance.load_state_dict(filtered_state_dict, strict=False) - if target_device is not None: - instance.to(target_device) + # Always placed: target_device is the caller's device or the configured default. + instance.to(target_device) instance.reset_cache() From 7d24f6820491f1db2f3b3fb2d1c40bfa0acc9161 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 20:17:04 +0200 Subject: [PATCH 138/250] Take the translation likelihood's model error from the shared estimator `fit_sigma_a_per_shell` was a local 81-point scan over every reflection in float64 -- 295 ms per candidate on 2DQ6, three times the translation refine beside it. It is replaced by `refinement.model_error_estimation.SigmaAEstimator`, which runs three nested-zoom stages of seventeen candidates on a cancellation-folded Rice that survives float32, and shrinks each shell toward the fitted curve instead of taking a per-shell argmax at face value. Reached through `SigmaAEstimator`, not the `estimate_beta` free function beneath it, and NOT for the cache. The wrapper interpolates four shell curves -- sigma_A, log Sigma_N, log Sigma_P, S2 -- and derives alpha and beta per reflection from them, so the second-moment identity holds at every reflection. Its own docstring warns that interpolating beta directly "can yield a value consistent with no sigma_A <= 1 at all", which is what a hand-rolled interpolation here would have done. The plan said to use the free function; the docstring said otherwise and the docstring was right. The likelihood now takes alpha and beta per reflection rather than a per-shell sigma_A, which drops the assumption that is exactly one -- alpha = sigma_A sqrt(Sigma_N/Sigma_P) carries the mismatch that assumption hides. The shell index and its `index_select` go with it; the estimator bins on its own abscissa and hands back per-reflection values. epsilon is passed as ones because E_obs is already epsilon-reduced by WilsonNormaliser, and free_mask as ones because there is no cross-validation set to protect at placement time. Outcome-neutral: 30/30 on the ten-structure panel and 36/40 on the ten-seed sweep, ZERO flips across 70 paired cells, with the llg median residual 1.62 -> 1.57 deg. **Not measurably faster, contrary to the expectation that motivated it.** The model-error fit is 275.7 ms against 294.8 on 2DQ6 -- but `local_translation_refine`, which this does not touch, moved 95.8 -> 143.7 ms between the same two jobs, so the machines differ by more than the change does. No speed claim is supported. The justification is one estimator instead of two. The sigma_A comparison lab probe goes: it existed to choose between four local sigma_A variants, and that question is now answered by using the shared one. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/llg_shared_sigma_a.sh | 21 -- .../diagnostics/llg_shared_sigma_a.py | 223 ------------------ alignment_lab/diagnostics/sigma_a_cost.py | 30 +-- docs/changelog.rst | 1 + torchref/experimental/alignment/__init__.py | 4 +- torchref/experimental/alignment/pipeline.py | 27 +-- .../experimental/alignment/translation.py | 142 ++++++----- 7 files changed, 93 insertions(+), 355 deletions(-) delete mode 100644 alignment_lab/analysis/llg_shared_sigma_a.sh delete mode 100644 alignment_lab/diagnostics/llg_shared_sigma_a.py diff --git a/alignment_lab/analysis/llg_shared_sigma_a.sh b/alignment_lab/analysis/llg_shared_sigma_a.sh deleted file mode 100644 index 5d609646..00000000 --- a/alignment_lab/analysis/llg_shared_sigma_a.sh +++ /dev/null @@ -1,21 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=llgsa -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=day -#SBATCH --time=04:00:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-3 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -PDBS=(2DQ6 6G9X 1DAW 3K7M) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -"$PY" -u alignment_lab/diagnostics/llg_shared_sigma_a.py --pdb "$PDB" --trials 10 \ - 2>&1 | grep -E '^ROW|rror' -echo DONE diff --git a/alignment_lab/diagnostics/llg_shared_sigma_a.py b/alignment_lab/diagnostics/llg_shared_sigma_a.py deleted file mode 100644 index e7fe7d8a..00000000 --- a/alignment_lab/diagnostics/llg_shared_sigma_a.py +++ /dev/null @@ -1,223 +0,0 @@ -"""Does the translation LLG rank badly because each candidate fits its own sigma_A? - -Over ten seeds on 2DQ6 the plain correlation puts truth at rank 0 in 10/10 while -the LLG -- a likelihood, strictly more information -- manages 2/10. A weaker -score beating a stronger one is a symptom, not a result. - -The suspect is how many parameters each score fits PER CANDIDATE: - - correlation 0 sum w E_obs^2_c |Fc|^2 / sum w |Fc|^2 - analytic R 1 the global scale k - LLG n_shells fit_sigma_a_per_shell, on THAT candidate's - own top translation - -`_llg_tf_rescore` is called once per rotation candidate and refits sigma_A each -time, so every wrong orientation is scored against a likelihood tuned to itself. -The docstring there warns against exactly this one level down -- refitting per -translation -- and the pipeline then does it per rotation. - -This recomputes the LLG with sigma_A held FIXED across candidates, three ways: - -``per_cand`` what the pipeline does now, as the control. -``shared`` fitted once, on the FRF's top-ranked candidate. Model-dependent - but not candidate-dependent, so it cannot flatter any one of them. -``empirical`` from Sigma_obs/Sigma_calc via weighting.empirical_sigma_a, which - is rotation-invariant by construction -- total scattering per - shell does not depend on orientation -- and is the estimate that - exists for precisely this reason. -``luzzati`` ASSUMED, not fitted: exp(-(2 pi^2/3) s^2 dVRMS^2) from the search - model's expected coordinate error, which is the same Eterm the - rotation function already weights with. It never looks at the - data, so it has zero free parameters of any kind -- which is the - property that makes the plain correlation robust, applied to a - likelihood instead. - -Also reported is ``r``, the analytical-scale R the pipeline actually selects on, -because the point of the exercise is replacing it: on 2DQ6 all 25 candidates -fall within 0.023 R of each other and the wrong winner leads a correct candidate -by 0.0001. - -Two forms of the correlation are compared, because the one the search maximises -is not a correlation coefficient: - -``corr`` what ``amplitude_translation_search`` returns, - ``sum w (E_obs^2 - mean) |Fc|^2 / sum w |Fc|^2``. The denominator - normalises the weighted MEAN, not the spread, so this is a weighted - mean of centred observed intensity. Fine for finding the peak in t - at fixed orientation; its scale across ORIENTATIONS depends on how - concentrated that candidate's |Fc|^2 happens to be. -``pearson`` the actual coefficient, ``cov / sqrt(var var)``, weighted the same - way. Per-candidate normalisation, so unlike a global rescale it can - and does reorder. Computed directly at each candidate's chosen - translation -- the FFT gives the numerator and one denominator but - not ``sum w |Fc|^4``, and one evaluation per candidate is cheap. -""" -from __future__ import annotations - -import argparse -import sys -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, rotated_case, seed_for, # noqa: E402 - symmetry_orbit) -from lab.truth import angle_to_orbit # noqa: E402 - - -def _rank_of_truth(scores, is_truth, higher_is_better=True): - order = sorted(range(len(scores)), key=lambda i: scores[i], - reverse=higher_is_better) - for pos, i in enumerate(order): - if is_truth[i]: - return pos - return -1 - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="2DQ6", choices=list(BENCH_PDBS)) - ap.add_argument("--trials", type=int, default=10) - ap.add_argument("--n-cand", type=int, default=25) - ap.add_argument("--thr-deg", type=float, default=8.0) - args = ap.parse_args() - - from torchref.experimental.alignment.frf.rotation_utils import ( - rotation_matrix_from_edmonds_euler, - ) - from torchref.experimental.alignment.pipeline import ( - MolecularReplacementPipeline, - ) - from torchref.experimental.alignment.rotation_search import prepare_frf_inputs - from torchref.experimental.alignment.translation import ( - DirectModelEvaluator, amplitude_translation_search, fit_sigma_a_per_shell, - llg_translation_rescore, local_translation_refine, normalise_calc, - precompute_G_for_rotation, - ) - from torchref.experimental.alignment.frf.preprocessing import eterm_sigma_a - from torchref.scaling import WilsonNormaliser - from torchref.scaling.weighting import empirical_sigma_a - - for trial in range(args.trials): - seed = seed_for(args.pdb, trial) - model, data, R_true = rotated_case(args.pdb, seed) - pipe = MolecularReplacementPipeline( - data, model, verbose=0, n_rotation_peaks=200, - n_rotation_candidates=args.n_cand, use_llg_tf=False) - frf = prepare_frf_inputs(model, data, d_min=pipe.d_min, d_max=pipe.d_max, - n_shells=pipe.n_shells, verbose=0) - pipe._frf = frf - peaks = pipe._rotation_candidates(frf)[: args.n_cand] - pipe._prepare_translation_arrays() - obs = pipe._obs - eye3 = pipe._eye3 - orbit = symmetry_orbit( - R_true, data.spacegroup.matrices.to(torch.float64).cpu(), - side="left", frame="cart", - reciprocal_basis=data.cell.reciprocal_basis_matrix.to( - torch.float64).cpu()) - - # One pass to collect each candidate's G, its top translation and its - # own E_calc there; sigma_A choices are applied afterwards so every - # variant scores the SAME placements. - cand = [] - for p in peaks: - ang = angle_to_orbit( - rotation_matrix_from_edmonds_euler(p.alpha, p.beta, p.gamma), - orbit) - rot = pipe._make_rotated(p)[0] - rot.spacegroup = data.spacegroup.hm - p1 = rot.copy(); p1.spacegroup = "P 1" - ev = DirectModelEvaluator(p1) - G, h_R = precompute_G_for_rotation( - ev, eye3, obs.hkl, data.spacegroup, data.cell) - _, _, tp = amplitude_translation_search( - obs=obs, interpolator=ev, R_rotation=eye3, - spacegroup=data.spacegroup, real_cell=data.cell, - grid_steps=pipe.translation_grid_steps, - n_peaks=pipe.n_translation_peaks, cluster_radius=0.05, - precomputed_G=G, precomputed_h_R=h_R) - t_top = torch.as_tensor(tp[0].translation, dtype=torch.float64, - device=G.device) - ph = torch.exp(2j * torch.pi * torch.einsum( - "ind,d->in", h_R.to(torch.float64), t_top).to(G.dtype)) - Fc_top = (G * ph).sum(dim=0).abs().to(torch.float64) - _, r_a = local_translation_refine( - obs=obs, interpolator=ev, R_rotation=eye3, - spacegroup=data.spacegroup, real_cell=data.cell, - t_init=t_top.cpu(), radius=0.06, grid_steps=13, - n_refinement_passes=1, - precomputed_G=G, precomputed_h_R=h_R) - # Weighted Pearson r between observed and calculated intensity at - # this candidate's chosen translation. - w = obs.weight.to(torch.float64) - x = (obs.E_obs.to(torch.float64)) ** 2 - y = (Fc_top / Fc_top.mean().clamp(min=1e-30)) ** 2 - wsum = w.sum().clamp(min=1e-30) - xm, ym = (w * x).sum() / wsum, (w * y).sum() / wsum - dx, dy = x - xm, y - ym - cov = (w * dx * dy).sum() / wsum - vx = (w * dx * dx).sum() / wsum - vy = (w * dy * dy).sum() / wsum - pear = float(cov / (vx * vy).clamp(min=1e-30).sqrt()) - cand.append(dict(ang=float(ang), corr=float(tp[0].score), G=G, - h_R=h_R, t=t_top, Fc=Fc_top, r=float(r_a), - pearson=pear, E_calc=normalise_calc(Fc_top, obs))) - - is_truth = [c["ang"] <= args.thr_deg for c in cand] - if not any(is_truth): - print(f"ROW pdb={args.pdb} trial={trial} truth_found=0", flush=True) - continue - - def sa_per_cand(c): - return fit_sigma_a_per_shell(obs.E_obs, c["E_calc"], obs.centric, - obs.shell_idx, obs.n_shells, n_grid=81) - sa_shared = sa_per_cand(cand[0]) # the FRF's own top candidate - fit_calc = WilsonNormaliser( - cand[0]["Fc"] ** 2, obs.s_mag, n_coeff=6, - s_lo=float(obs.s_mag.min()), s_hi=float(obs.s_mag.max())) - sa_emp_per_refl = empirical_sigma_a( - obs.fit.evaluate(obs.s_mag).to(torch.float64), - fit_calc.evaluate(obs.s_mag).to(torch.float64)) - # collapse to per-shell, the shape llg_translation_rescore expects - cnt = torch.bincount(obs.shell_idx, minlength=obs.n_shells).to(torch.float64) - tot = torch.zeros(obs.n_shells, dtype=torch.float64).scatter_add_( - 0, obs.shell_idx, sa_emp_per_refl.to(torch.float64)) - sa_emp = (tot / cnt.clamp(min=1.0)).clamp(1e-3, 1 - 1e-6) - - def llg_of(c, sigma_a): - return float(llg_translation_rescore( - obs=obs, G=c["G"], h_R=c["h_R"], - t_candidates=c["t"].view(1, 3), sigma_a=sigma_a)[0]) - - # Assumed sigma_A: the Luzzati falloff at the shell centres, from the - # model error the pipeline already estimates. No data, no fit. - s_shell = torch.zeros(obs.n_shells, dtype=torch.float64) - cnt_s = torch.bincount(obs.shell_idx, minlength=obs.n_shells).to(torch.float64) - s_shell.scatter_add_(0, obs.shell_idx, obs.s_mag.to(torch.float64)) - s_shell = s_shell / cnt_s.clamp(min=1.0) - sa_luz = eterm_sigma_a(s_shell, float(pipe.model_error_A)).clamp(1e-3, 1 - 1e-6) - - variants = { - "per_cand": [llg_of(c, sa_per_cand(c)) for c in cand], - "shared": [llg_of(c, sa_shared) for c in cand], - "empirical": [llg_of(c, sa_emp) for c in cand], - "luzzati": [llg_of(c, sa_luz) for c in cand], - } - r_corr = _rank_of_truth([c["corr"] for c in cand], is_truth) - r_pear = _rank_of_truth([c["pearson"] for c in cand], is_truth) - r_R = _rank_of_truth([c["r"] for c in cand], is_truth, higher_is_better=False) - parts = " ".join( - f"rank_{k}={_rank_of_truth(v, is_truth)}" for k, v in variants.items()) - print(f"ROW pdb={args.pdb} trial={trial} n_truth={sum(is_truth)} " - f"n_cand={len(cand)} vrms={float(pipe.model_error_A):.2f} " - f"rank_r={r_R} rank_corr={r_corr} rank_pearson={r_pear} {parts}", - flush=True) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/sigma_a_cost.py b/alignment_lab/diagnostics/sigma_a_cost.py index 27ba5fbe..b8e23632 100644 --- a/alignment_lab/diagnostics/sigma_a_cost.py +++ b/alignment_lab/diagnostics/sigma_a_cost.py @@ -1,21 +1,20 @@ """Where does the likelihood ranking's 2.5x actually go? -`fit_sigma_a_per_shell` scans 81 values of D against every reflection, in -float64, evaluating BOTH the Rice and the Woolfson branch everywhere and then -discarding half of each with a `where`. On 2DQ6 that is 81 x 228197 = 18.5M -elements per tensor and four of them materialised, per candidate, times 25 -candidates. +The model-error fit used to be a local 81-point scan over every reflection, in +float64, evaluating both likelihood branches everywhere and discarding half -- +295 ms per candidate on 2DQ6, three times the translation refine beside it. It +now goes through the shared `SigmaAEstimator`. -This times the pieces so the fix is aimed rather than guessed: the sigma_A fit, -the Wilson normalisation of the calculated side, the likelihood evaluation -itself, and the translation refine they sit alongside. +This times the pieces so the effect is measured rather than assumed: the +model-error fit, the Wilson normalisation of the calculated side, the likelihood +evaluation itself, and the translation refine they sit alongside. """ import sys, time from pathlib import Path import torch sys.path.insert(0, str(Path(__file__).resolve().parents[1])) torch.set_grad_enabled(False) -from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 +from lab import load_case, random_rotation, seed_for # noqa: E402 def _t(fn, n=3): @@ -29,7 +28,7 @@ def _t(fn, n=3): def main(): from torchref.experimental.alignment.translation import ( DirectModelEvaluator, TranslationObs, amplitude_translation_search, - correlation_at, fit_sigma_a_per_shell, llg_at, local_translation_refine, + correlation_at, fit_model_error, llg_at, local_translation_refine, normalise_calc, precompute_G_for_rotation) for pdb in sys.argv[1:] or ["1DAW", "2DQ6"]: @@ -63,17 +62,14 @@ def main(): "ind,d->in", h_R.to(torch.float64), t0.to(G.device)).to(G.dtype)) Fc = (G * ph).sum(dim=0).abs().to(torch.float64) t_norm, E_calc = _t(lambda: normalise_calc(Fc, obs)) - t_sa, sa = _t(lambda: fit_sigma_a_per_shell( - obs.E_obs, E_calc, obs.centric, obs.shell_idx, obs.n_shells, - n_grid=81)) - t_llg, _ = _t(lambda: llg_at(obs, G, h_R, t0, sa)) + t_sa, (alpha, beta) = _t(lambda: fit_model_error(obs, E_calc)) + t_llg, _ = _t(lambda: llg_at(obs, G, h_R, t0, alpha, beta)) t_corr, _ = _t(lambda: correlation_at(obs, G, h_R, t0)) N = obs.hkl.numel() // 3 print(f"ROW pdb={pdb} N={N} shells={obs.n_shells} " f"refine={1000*t_ref:.1f}ms norm_calc={1000*t_norm:.1f}ms " - f"sigma_a={1000*t_sa:.1f}ms llg={1000*t_llg:.1f}ms " - f"corr={1000*t_corr:.1f}ms " - f"grid_elems={81*N/1e6:.1f}M", flush=True) + f"model_err={1000*t_sa:.1f}ms llg={1000*t_llg:.1f}ms " + f"corr={1000*t_corr:.1f}ms", flush=True) main() diff --git a/docs/changelog.rst b/docs/changelog.rst index 54266255..d454452b 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- The translation likelihood's model error comes from the shared ``SigmaAEstimator`` instead of a local 81-point scan over every reflection. It returns ``alpha`` and ``beta`` per reflection rather than a per-shell ``sigma_A``, so the likelihood no longer assumes ```` is exactly one. Outcome-neutral over 70 seeded cells, zero flips; not measurably faster - Fixed the translation likelihood's variance convention, which scored acentric reflections at twice the variance intended -- 90-95% of reflections. The alignment package carried its own Rice and Woolfson parameterised by the *amplitude* variance and handed both branches the same number, where the acentric branch needs half what the centric one does. It now uses ``base.targets.xray_likelihoods.rice_per_refl``, which takes the complex variance and derives the centric case from it - Removed ``experimental/alignment/distributions.py``. Its ``stable_log_bessel_i0`` also carried a wrong asymptotic coefficient, giving -2.9e-5 at x = 50 against -5.4e-7 for the correct term; the shared implementation uses ``log(i0e(z)) + z``, which is exact - Molecular-replacement candidates are ranked by the translation function's likelihood, not by an analytical-scale R-factor. 30/30 on the ten-structure panel against 29/30, and 37/40 against 36/40 over a ten-seed sweep, with tighter placements on the cells both solve. Selectable through ``rank_by``; the correlation is the third option and is the worst of the three at 32/40, despite a rank-level harness rating it best on a truth label that disagrees with coordinate superposition diff --git a/torchref/experimental/alignment/__init__.py b/torchref/experimental/alignment/__init__.py index 4ed3330a..8ce40da4 100644 --- a/torchref/experimental/alignment/__init__.py +++ b/torchref/experimental/alignment/__init__.py @@ -71,7 +71,7 @@ TranslationPeak, amplitude_translation_search, find_translation_peaks, - fit_sigma_a_per_shell, + fit_model_error, llg_translation_rescore, local_translation_refine, precompute_G_for_rotation, @@ -114,5 +114,5 @@ "llg_translation_rescore", "precompute_G_for_rotation", "find_translation_peaks", - "fit_sigma_a_per_shell", + "fit_model_error", ] diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 533ec694..f0c3423f 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -64,7 +64,7 @@ amplitude_translation_search, correlation_at, llg_at, - fit_sigma_a_per_shell, + fit_model_error, llg_translation_rescore, local_translation_refine, normalise_calc, @@ -734,15 +734,13 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: return None llg = float("nan") if self.rank_by == "llg": - # sigma_A at the chosen translation, per candidate. Fitting it once - # and sharing it across candidates was measured to give identical - # rankings, so the cheaper-to-reason-about form is used. + # Model error at the chosen translation, per candidate. Fitting it + # once and sharing it across candidates was measured to give + # identical rankings, so the cheaper-to-reason-about form is used. E_calc = normalise_calc( self._fcalc_at(G_pre, h_R_pre, best[1]), self._obs) - sa = fit_sigma_a_per_shell( - self._obs.E_obs, E_calc, self._obs.centric, - self._obs.shell_idx, self._obs.n_shells, n_grid=81) - llg = llg_at(self._obs, G_pre, h_R_pre, best[1], sa) + alpha, beta = fit_model_error(self._obs, E_calc) + llg = llg_at(self._obs, G_pre, h_R_pre, best[1], alpha, beta) return (best[0], best[1], best[2], llg) def _fcalc_at(self, G, h_R, t): @@ -769,10 +767,10 @@ def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): obs = self._obs self._timer.start("6b_llg_tf_rescore") - # sigma_A is fitted against the top translation only, and reused for - # every candidate. It is a per-shell model-reliability curve, not a + # The model error is fitted against the top translation only and reused + # for every candidate. It is a model-reliability curve, not a # per-candidate score: refitting it per t would let each candidate - # choose the D that flatters it, which is scoring a model against a + # choose the alpha that flatters it, which is scoring a model against a # likelihood tuned to that model. t_top_t = torch.as_tensor( t_peaks[0].translation, dtype=torch.float64, device=device, @@ -784,10 +782,7 @@ def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): ) Fc_top = (G_pre * phase_top).sum(dim=0).abs().to(torch.float64) E_calc_top = normalise_calc(Fc_top, obs) - sigma_a_tf = fit_sigma_a_per_shell( - obs.E_obs, E_calc_top, obs.centric, - obs.shell_idx, obs.n_shells, n_grid=81, - ) + alpha_tf, beta_tf = fit_model_error(obs, E_calc_top) t_cands = torch.as_tensor( np.stack([p.translation for p in t_peaks]), @@ -795,7 +790,7 @@ def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): ) llg_tf = llg_translation_rescore( obs=obs, G=G_pre, h_R=h_R_pre, t_candidates=t_cands, - sigma_a=sigma_a_tf, interp_var=None, + alpha=alpha_tf, beta=beta_tf, ) self._timer.stop("6b_llg_tf_rescore") diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index 2139027b..ab19e3f2 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -458,60 +458,54 @@ def normalise_calc(F_calc: torch.Tensor, obs: TranslationObs) -> torch.Tensor: return out[0] if single else out -def fit_sigma_a_per_shell( - E_obs: torch.Tensor, - E_calc: torch.Tensor, - centric: torch.Tensor, - shell_idx: torch.Tensor, - n_shells: int, - n_grid: int = 81, - interp_var: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """Per-shell sigma_A = D, fitted by a grid scan of the shell likelihood. - - Scans ``D`` in ``[0, 0.99]`` and returns the per-shell maximum. At - ``n_grid=81`` the resolution is ~0.012, which is finer than the difference - between adjacent shells on any real falloff. - - This is the *fitted* half of a pair. :func:`torchref.scaling.weighting.empirical_sigma_a` - is the other: it measures model reliability from the ratio of two Wilson - curves and works **before** the molecule is placed, which is what the - rotation function needs. Once a translation exists the residual is - per-reflection rather than per-shell-average, so a direct fit against the - placed model is available and strictly better informed. The two answer the - same question at two different points in the search. - - Returns - ------- - sigma_a : torch.Tensor, shape (n_shells,) +def fit_model_error(obs: TranslationObs, E_calc: torch.Tensor, *, shrink: bool = True): + """Per-reflection ``(alpha, beta)`` for a placed model, from the shared estimator. + + Returns the two quantities the likelihood actually wants: ``alpha``, the + multiplier on the calculated amplitude, and ``beta``, the conditional + variance. They are what + :class:`~torchref.refinement.model_error_estimation.sigma_a.SigmaAEstimator` + produces, and using them rather than ``(D, 1 - D^2)`` drops the assumption + that ```` is exactly 1 -- ``alpha = sigma_A sqrt(Sigma_N/Sigma_P)`` + carries the mismatch that assumption hides. + + **This replaces a local 81-point scan over every reflection.** The shared + estimator runs three nested-zoom stages of seventeen candidates instead, on a + cancellation-folded Rice that survives float32, and it shrinks each shell + toward the fitted curve rather than taking a per-shell argmax at face value. + + It is reached through ``SigmaAEstimator``, not the ``estimate_beta`` free + function underneath it, and the reason is not the cache. The wrapper + interpolates **four** shell curves -- ``sigma_A``, ``log Sigma_N``, + ``log Sigma_P``, ``S2`` -- and derives ``alpha``/``beta`` per reflection from + them, so the second-moment identity holds at every reflection. Its docstring + warns that interpolating ``beta`` directly "can yield a value consistent with + no ``sigma_A <= 1`` at all", which is exactly what a hand-rolled + interpolation here would have done. + + ``epsilon`` is passed as ones because ``obs.E_obs`` is already + epsilon-reduced by :class:`~torchref.scaling.WilsonNormaliser`; applying it + again would count multiplicity twice. ``free_mask`` is all-ones: there is no + cross-validation set to protect at placement time, and the estimate is not + being used to decide when to stop refining. """ - device = E_obs.device - dtype = E_obs.dtype - N = E_obs.numel() - D_grid = torch.linspace(0.0, 0.99, n_grid, device=device, dtype=dtype) # (G,) - F_mean = D_grid.view(-1, 1) * E_calc.view(1, -1) # (G, N) - # Sigma is the COMPLEX variance. `rice_per_refl` derives the centric - # amplitude variance from it internally, which is the whole reason to use it - # -- the previous code passed one amplitude variance to a Rice and a - # Woolfson written in different conventions, leaving acentrics at twice the - # variance they should have had. - Sigma = (1.0 - D_grid * D_grid).clamp(min=1e-4) # (G,) - if interp_var is None: - Sigma_full = Sigma.view(-1, 1).expand(n_grid, N) - else: - Sigma_full = (Sigma.view(-1, 1) + interp_var.view(1, -1)).clamp(min=1e-4) - E_obs_full = E_obs.view(1, -1).expand(n_grid, N) - - # Negated: `rice_per_refl` is an NLL and everything here maximises. - ll = -rice_per_refl(E_obs_full, F_mean, Sigma_full, - centric.view(1, -1).expand(n_grid, N)) # (G, N) - - # Sum per shell, take argmax over the D-grid. - shell_idx_gn = shell_idx.view(1, -1).expand(n_grid, N) - ll_per_shell = torch.zeros((n_grid, n_shells), dtype=dtype, device=device) - ll_per_shell.scatter_add_(1, shell_idx_gn, ll) # (G, n_shells) - best_idx = ll_per_shell.argmax(dim=0) # (n_shells,) - return D_grid[best_idx] + # Local import: `torchref.refinement.__init__` eagerly pulls the refinement + # drivers and every target, which is a heavy load for one estimator. The + # same pattern, for the same reason, is documented at + # `scaling/scaler_base.py` -- "Do not 'tidy' them up". + from torchref.refinement.model_error_estimation.sigma_a import SigmaAEstimator + + ones = torch.ones_like(obs.E_obs) + est = SigmaAEstimator().get( + F_obs=obs.E_obs, + F_calc_scaled=E_calc.to(obs.E_obs.dtype), + centric=obs.centric, + epsilon=ones, + d_star_sq=(obs.s_mag * obs.s_mag).to(obs.E_obs.dtype), + free_mask=torch.ones_like(obs.E_obs, dtype=torch.bool), + shrink=shrink, + ) + return est.alpha, est.beta def llg_translation_rescore( @@ -519,8 +513,8 @@ def llg_translation_rescore( G: torch.Tensor, h_R: torch.Tensor, t_candidates: torch.Tensor, - sigma_a: torch.Tensor, - interp_var: Optional[torch.Tensor] = None, + alpha: torch.Tensor, + beta: torch.Tensor, ) -> torch.Tensor: """Per-translation Rice / Woolfson log-likelihood over candidate positions. @@ -554,8 +548,7 @@ def llg_translation_rescore( Parameters ---------- obs : TranslationObs - Supplies ``E_obs``, centricity and the shell binning -- the same binning - ``sigma_a`` was fitted on, which is why it is not re-derived here. + Supplies ``E_obs`` and centricity. G : (S, N) complex Per-sym ``F_p1`` contributions x per-sym translation phase, from :func:`precompute_G_for_rotation`. @@ -563,10 +556,10 @@ def llg_translation_rescore( Per-sym rotated reciprocal indices. t_candidates : (K, 3) Fractional translations to score. - sigma_a : (n_shells,) - Fixed per-shell sigma_A, shared across candidates. - interp_var : (N,), optional - Per-reflection variance inflation. + alpha, beta : (N,) + Model reliability and conditional variance per reflection, from + :func:`fit_model_error`. Fixed across candidates on purpose: refitting + per candidate would score each against a likelihood tuned to itself. Returns ------- @@ -576,7 +569,6 @@ def llg_translation_rescore( real_dtype = torch.float64 complex_dtype = G.dtype - shell_idx = obs.shell_idx centric = obs.centric K = t_candidates.shape[0] @@ -590,22 +582,19 @@ def llg_translation_rescore( Fc_complex = (G.view(1, S, N) * phase).sum(dim=1) # (K, N) F_calc = Fc_complex.abs().to(real_dtype) # (K, N) - shell_idx_l = shell_idx.to(device).long() E_calc = normalise_calc(F_calc, obs) # (K, N) - E_obs = obs.E_obs.to(device).to(real_dtype) - sigma_a_d = sigma_a.to(device).to(real_dtype) # (n_shells,) - D_per_refl = sigma_a_d.index_select(0, shell_idx_l) # (N,) - # Complex variance, matching `rice_per_refl`'s convention. - var_d = (1.0 - D_per_refl * D_per_refl).clamp(min=1e-4) # (N,) - if interp_var is not None: - var_per_refl = (var_d + interp_var.to(device).to(real_dtype)).clamp(min=1e-4) - else: - var_per_refl = var_d + # alpha and beta arrive PER REFLECTION, interpolated from the shell fit by + # the shared estimator. No `index_select` on a shell index: the estimator + # bins on its own abscissa, and taking its per-reflection output is what + # keeps the second-moment identity holding at every reflection rather than + # only per shell. + a_r = alpha.to(device).to(real_dtype) # (N,) + Sigma_r = beta.to(device).to(real_dtype).clamp(min=1e-4) # (N,) - F_mean = D_per_refl.view(1, N) * E_calc # (K, N) - Sigma_full = var_per_refl.view(1, N).expand(K, N) + F_mean = a_r.view(1, N) * E_calc # (K, N) + Sigma_full = Sigma_r.view(1, N).expand(K, N) E_obs_full = E_obs.view(1, N).expand(K, N) cent = centric.to(device).to(torch.bool) cent_full = cent.view(1, N).expand(K, N) @@ -629,7 +618,8 @@ def llg_at( G: torch.Tensor, h_R: torch.Tensor, t: torch.Tensor, - sigma_a: torch.Tensor, + alpha: torch.Tensor, + beta: torch.Tensor, ) -> float: """The translation likelihood at one translation, as a candidate score. @@ -641,7 +631,7 @@ def llg_at( return float(llg_translation_rescore( obs=obs, G=G, h_R=h_R, t_candidates=t.detach().reshape(1, 3).to(G.device).to(torch.float64), - sigma_a=sigma_a, + alpha=alpha, beta=beta, )[0]) From a596173e9724dedc0d8b05d026c4aa5169d12155 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 20:32:44 +0200 Subject: [PATCH 139/250] Use the shared FFT-size, Euler and symmetry-unroll helpers Four substitutions, each verified equivalent before being made rather than after: `adjust_gridding` -> `symmetry.find_fft_friendly_size`. The only call site asked for max_prime=5, which is what the shared one does; they agree for every n from 1 to 4000. `rotation_matrix_from_edmonds_euler_batch` -> `base.alignment.rotation_matrix_euler_zyz`. Bit-identical over 180k elements in float32 and float64. `peak_finder` carried a comment claiming the three-matrix-product form "rounds differently in the last bit" and that the fused form was what every measurement used -- that is not true and the comment is corrected: the products carry exact zeros and ones, so the matmul reduces to the same two-term sums. `epsilon_aware_unroll` -> `SpaceGroup.expand_hkl(include_friedel=False)`, which performs the same orbit dedup: measured emitting n_ops/epsilon(h) distinct mates, 2 for an axial reflection in P3121 against 6 for a general one. It had no production caller. The convention test moves onto the shared helper, which is worth more to guard than the copy was. `compute_v_budget` had zero references anywhere. So did four rotation helpers in pipeline.py -- `rotation_matrix_from_euler_zyz` (numpy, in an otherwise torch module), `rotation_angular_distance`, `euler_angular_distance` and `cluster_rotation_peaks`, the last two definition-only even inside the file. Bit-identical on the benchmark panel: 30 paired cells per arm, zero differing at all, which is what a pure substitution has to look like. `fit_relative_wilson_b` is left in place with its justification corrected. The comment said it "stays for the rescore"; the rescore was deleted, and it now has no production caller either. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- .../diagnostics/frf_encode_compare.py | 12 +- .../diagnostics/frf_ghost_knockout.py | 8 +- .../diagnostics/frf_inject_phaser_obs.py | 2 +- alignment_lab/diagnostics/frf_map_compare.py | 12 +- docs/changelog.rst | 1 + .../alignment/test_symmetry_conventions.py | 14 +- tests/unit/frf_separate/test_invariants.py | 20 ++- torchref/experimental/alignment/__init__.py | 6 - .../experimental/alignment/frf/__init__.py | 2 - .../experimental/alignment/frf/peak_finder.py | 12 +- .../alignment/frf/preprocessing.py | 133 ------------------ .../alignment/frf/rotation_utils.py | 31 ---- .../alignment/frf/sitelist_ang.py | 28 +--- torchref/experimental/alignment/pipeline.py | 80 +---------- .../experimental/alignment/rotation_search.py | 12 +- 15 files changed, 52 insertions(+), 321 deletions(-) diff --git a/alignment_lab/diagnostics/frf_encode_compare.py b/alignment_lab/diagnostics/frf_encode_compare.py index 7ce64915..11044679 100644 --- a/alignment_lab/diagnostics/frf_encode_compare.py +++ b/alignment_lab/diagnostics/frf_encode_compare.py @@ -31,8 +31,9 @@ * ``obs_ours_unroll`` / ``obs_dedup_unroll`` -- Phaser's ASU-level intensities (``PHASER_TERMS_DUMP``, keyed by Miller index) put through *our* two symmetry unrolls: the production one, which emits all ``n_ops`` orbit positions, and - ``epsilon_aware_unroll``, which emits only the distinct ones as Phaser does - (``!duplicate(isym,rhkl)``, DataMR.cc:954). Same intensities, same encoder, + ``SpaceGroup.expand_hkl(include_friedel=False)``, which emits only the + distinct ones as Phaser does (``!duplicate(isym,rhkl)``, DataMR.cc:954). + Same intensities, same encoder, same target -- so the difference between these two arms is the multiplicity handling and nothing else. @@ -311,10 +312,6 @@ def unroll_arms(pdb: str, terms_csv: Path): The counts are exact integers, so the multiplicity question is answered by arithmetic before any encoding happens. """ - from torchref.experimental.alignment.frf.preprocessing import ( - epsilon_aware_unroll, - ) - d = np.loadtxt(terms_csv, delimiter=",", skiprows=1) if d.ndim == 1: d = d[None, :] @@ -332,7 +329,8 @@ def unroll_arms(pdb: str, terms_csv: Path): i_all = inten.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() # Phaser-faithful: distinct orbit positions only (DataMR.cc:954). - hkl_ded, asu_idx = epsilon_aware_unroll(hkl.to(torch.long), sg) + hkl_ded, asu_idx, _ = data.spacegroup.expand_hkl( + hkl.to(torch.long), include_friedel=False) s_ded = hkl_ded.to(torch.float64) @ rec i_ded = inten[asu_idx] diff --git a/alignment_lab/diagnostics/frf_ghost_knockout.py b/alignment_lab/diagnostics/frf_ghost_knockout.py index 0bf7eaf8..ee70841f 100644 --- a/alignment_lab/diagnostics/frf_ghost_knockout.py +++ b/alignment_lab/diagnostics/frf_ghost_knockout.py @@ -93,16 +93,14 @@ def _orbit_of_identity(data): def _truth_and_margin(arf, orbit): """``(rank, sigma, angle, best_ghost_sigma, margin)`` for one map.""" - from torchref.experimental.alignment.frf.rotation_utils import ( - rotation_matrix_from_edmonds_euler_batch, - ) + from torchref.base.alignment.rotation import rotation_matrix_euler_zyz v = arf.values.to(torch.float64).cpu() - R = rotation_matrix_from_edmonds_euler_batch( + R = rotation_matrix_euler_zyz(torch.stack([ arf.alphas.to(torch.float64).cpu(), arf.betas.to(torch.float64).cpu(), arf.gammas.to(torch.float64).cpu(), - ) + ], dim=-1)) sig = (v - v.mean()) / v.std().clamp(min=1e-30) best = None diff --git a/alignment_lab/diagnostics/frf_inject_phaser_obs.py b/alignment_lab/diagnostics/frf_inject_phaser_obs.py index 00f1f040..468a899d 100644 --- a/alignment_lab/diagnostics/frf_inject_phaser_obs.py +++ b/alignment_lab/diagnostics/frf_inject_phaser_obs.py @@ -12,7 +12,7 @@ -- Phaser's ``clmn`` through our evaluator gives r = 0.998 with an identical argmax; * the reciprocal frame (positions agree to 1e-8) and the unroll (our - ``epsilon_aware_unroll`` reproduces Phaser's point count exactly). + the orbit dedup reproduces Phaser's point count exactly). What is NOT verified is the intensity attached to each position: ours correlates with Phaser's at 0.988 (1AK5), 0.877 (2DQ6) and 0.711 (3GR5) -- and that ordering diff --git a/alignment_lab/diagnostics/frf_map_compare.py b/alignment_lab/diagnostics/frf_map_compare.py index ffb97edb..15513a4d 100644 --- a/alignment_lab/diagnostics/frf_map_compare.py +++ b/alignment_lab/diagnostics/frf_map_compare.py @@ -198,23 +198,19 @@ def compare(ours, phaser_angles, phaser_values, frame, data, *, topn: int = 20) ours is much worse, the ghost problem is ours; if they agree, the ghosts are inherent to the target function and no reimplementation will remove them. """ - from torchref.experimental.alignment.frf.rotation_utils import ( - rotation_matrix_from_edmonds_euler_batch, - ) + from torchref.base.alignment.rotation import rotation_matrix_euler_zyz a = ours.arf ov = a.values.to(torch.float64).cpu() - R_ours = rotation_matrix_from_edmonds_euler_batch( + R_ours = rotation_matrix_euler_zyz(torch.stack([ a.alphas.to(torch.float64).cpu(), a.betas.to(torch.float64).cpu(), a.gammas.to(torch.float64).cpu(), - ) + ], dim=-1)) pv = phaser_values.to(torch.float64).cpu() ang = phaser_angles.to(torch.float64).cpu() - R_grid = rotation_matrix_from_edmonds_euler_batch( - torch.deg2rad(ang[:, 0]), torch.deg2rad(ang[:, 1]), torch.deg2rad(ang[:, 2]), - ) + R_grid = rotation_matrix_euler_zyz(torch.deg2rad(ang[:, :3])) PR, AX = frame["PR"], frame["axisrot"] # principal frame -> PDB frame (runMR_FRF.cc:542) R_ph = torch.einsum("ij,njk,kl->nil", AX, R_grid, PR) diff --git a/docs/changelog.rst b/docs/changelog.rst index d454452b..09c6fe17 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- The alignment package uses the shared FFT-size and Euler-matrix helpers instead of its own copies, and drops four unused rotation utilities and two duplicated symmetry helpers. Bit-identical placements on the benchmark panel - The translation likelihood's model error comes from the shared ``SigmaAEstimator`` instead of a local 81-point scan over every reflection. It returns ``alpha`` and ``beta`` per reflection rather than a per-shell ``sigma_A``, so the likelihood no longer assumes ```` is exactly one. Outcome-neutral over 70 seeded cells, zero flips; not measurably faster - Fixed the translation likelihood's variance convention, which scored acentric reflections at twice the variance intended -- 90-95% of reflections. The alignment package carried its own Rice and Woolfson parameterised by the *amplitude* variance and handed both branches the same number, where the acentric branch needs half what the centric one does. It now uses ``base.targets.xray_likelihoods.rice_per_refl``, which takes the complex variance and derives the centric case from it - Removed ``experimental/alignment/distributions.py``. Its ``stable_log_bessel_i0`` also carried a wrong asymptotic coefficient, giving -2.9e-5 at x = 50 against -5.4e-7 for the correct term; the shared implementation uses ``log(i0e(z)) + z``, which is exact diff --git a/tests/unit/alignment/test_symmetry_conventions.py b/tests/unit/alignment/test_symmetry_conventions.py index 6e6fe912..41570390 100644 --- a/tests/unit/alignment/test_symmetry_conventions.py +++ b/tests/unit/alignment/test_symmetry_conventions.py @@ -22,9 +22,6 @@ import pytest import torch -from torchref.experimental.alignment.frf.preprocessing import ( - epsilon_aware_unroll, -) from torchref.experimental.alignment.sh import ( hkl_symops_to_cartesian, symmetrize_anisotropy, @@ -209,21 +206,26 @@ def test_epsilon_uses_the_row_vector_convention(hm, non_orthogonal): @pytest.mark.parametrize("hm, non_orthogonal", SPACEGROUPS) -def test_epsilon_aware_unroll_stays_within_the_true_orbit(hm, non_orthogonal): +def test_symmetry_unroll_stays_within_the_true_orbit(hm, non_orthogonal): """Every emitted position must be a genuine symmetry mate of its input. This exercises a real call site rather than the contraction in isolation. Under the wrong convention the emitted positions leave the true orbit for a non-orthogonal setting, which is what let two inequivalent reflections land on one Miller index carrying different ``|F|``. + + Against ``SpaceGroup.expand_hkl``, the shared helper. The alignment package + had its own ``epsilon_aware_unroll`` doing the same orbit dedup -- measured + to emit ``n_ops/epsilon(h)`` distinct mates, exactly as this does -- and it + was deleted once nothing in production called it. Guarding the shared one is + worth more than guarding the copy was. """ del non_orthogonal sg = SpaceGroup(hm) - S = sg.matrices.detach().cpu().to(torch.float64) g = torch.Generator().manual_seed(19) hkl = torch.randint(-9, 10, (150, 3), generator=g) - unrolled, asu_idx = epsilon_aware_unroll(hkl, S) + unrolled, asu_idx, _ = sg.expand_hkl(hkl, include_friedel=False) unrolled = unrolled.detach().cpu().to(torch.long) asu_idx = asu_idx.detach().cpu().to(torch.long) diff --git a/tests/unit/frf_separate/test_invariants.py b/tests/unit/frf_separate/test_invariants.py index 6e4bc394..be941667 100644 --- a/tests/unit/frf_separate/test_invariants.py +++ b/tests/unit/frf_separate/test_invariants.py @@ -12,11 +12,11 @@ import torch from torchref.experimental.alignment.frf.sitelist_ang import ( - adjust_gridding, build_adaptive_sample_list, build_dense_map_per_beta, evaluate_rotation_function, ) +from torchref.symmetry.symmetry import find_fft_friendly_size from torchref.experimental.alignment.frf.wigner_d import ( _wigner_d_blocks, wigner_contraction_per_beta, @@ -51,11 +51,17 @@ def _real_sh_coeffs(L): return xi -def test_adjust_gridding_basic(): - assert adjust_gridding(180) == 180 # 180 = 4·45 = 4·9·5 (5-smooth) - assert adjust_gridding(7) == 8 # rounds up to next 5-smooth - assert adjust_gridding(243) == 243 # 3^5 - assert adjust_gridding(1) == 1 +def test_fft_size_is_five_smooth(): + """The dense-map grid must factor into 2, 3 and 5 only. + + Uses the shared ``find_fft_friendly_size``; the alignment package's own + ``adjust_gridding`` was deleted after being measured identical to it for + every n from 1 to 4000 at the only setting the FRF ever called it with. + """ + assert find_fft_friendly_size(180) == 180 # 180 = 4·45 = 4·9·5 + assert find_fft_friendly_size(7) == 8 # rounds up to the next 5-smooth + assert find_fft_friendly_size(243) == 243 # 3^5 + assert find_fft_friendly_size(1) == 1 def test_sample_count_matches_so3_measure(): @@ -113,7 +119,7 @@ def test_real_output_from_hermitian_xi(): xi = _make_xi(L) Δ = 10.0 bmax = int(math.ceil(180.0 / Δ)) - N = adjust_gridding(2 * max(bmax, 2 * L - 1), max_prime=5) + N = find_fft_friendly_size(2 * max(bmax, 2 * L - 1)) _, _, _, _, beta_grid = build_adaptive_sample_list(Δ) M = build_dense_map_per_beta(xi, beta_grid, N) # Imaginary part divided by typical magnitude should be < 1e-10. diff --git a/torchref/experimental/alignment/__init__.py b/torchref/experimental/alignment/__init__.py index 8ce40da4..d0b9ea3e 100644 --- a/torchref/experimental/alignment/__init__.py +++ b/torchref/experimental/alignment/__init__.py @@ -80,9 +80,6 @@ MolecularReplacementPipeline, MRSolution, align_model_to_data, - cluster_rotation_peaks, - euler_angular_distance, - rotation_angular_distance, ) __all__ = [ @@ -102,9 +99,6 @@ "rotation_matrix_from_edmonds_euler", "edmonds_euler_from_rotation_matrix", "rotation_angular_distance_deg", - "cluster_rotation_peaks", - "rotation_angular_distance", - "euler_angular_distance", # Translation search "TranslationObs", "TranslationPeak", diff --git a/torchref/experimental/alignment/frf/__init__.py b/torchref/experimental/alignment/frf/__init__.py index 1dc818a3..991e1896 100644 --- a/torchref/experimental/alignment/frf/__init__.py +++ b/torchref/experimental/alignment/frf/__init__.py @@ -15,7 +15,6 @@ edmonds_euler_from_rotation_matrix, rotation_angular_distance_deg, rotation_matrix_from_edmonds_euler, - rotation_matrix_from_edmonds_euler_batch, ) from .types import ( AdaptiveRotationFunction, @@ -31,7 +30,6 @@ "model_sf_abs", # Rotation geometry helpers "rotation_matrix_from_edmonds_euler", - "rotation_matrix_from_edmonds_euler_batch", "edmonds_euler_from_rotation_matrix", "rotation_angular_distance_deg", # Types diff --git a/torchref/experimental/alignment/frf/peak_finder.py b/torchref/experimental/alignment/frf/peak_finder.py index c6369a17..6ae3ef81 100644 --- a/torchref/experimental/alignment/frf/peak_finder.py +++ b/torchref/experimental/alignment/frf/peak_finder.py @@ -49,11 +49,13 @@ def _so3_greedy_nms( # preallocated kept-buffer (no repeated torch.stack), and a cosine threshold # (no per-iteration arccos). Result is identical to the original distance test. order = torch.argsort(values, descending=True).cpu().tolist() - # `rotation_matrix_euler_zyz` is the same fused single-pass form, term for - # term, so it rounds identically. That matters here and not only for tidiness: - # the NMS threshold test below flips for pairs sitting exactly on it, and the - # three-matrix-product form in `rotation_utils` rounds differently in the last - # bit. Every measurement on this engine was made with the fused form. + # `rotation_matrix_euler_zyz` is the shared implementation, and it is now the + # only one -- the alignment package's own batch copy was deleted after being + # measured bit-identical to it over 180k elements in both float32 and + # float64. (The three-matrix product it used reduces to the same two-term + # sums, because the rotation factors carry exact zeros and ones.) Rounding + # matters here beyond tidiness: the NMS threshold below flips for pairs + # sitting exactly on it. R_all = ( rotation_matrix_euler_zyz(torch.stack([alphas, betas, gammas], dim=-1)) .to(torch.float64).cpu() diff --git a/torchref/experimental/alignment/frf/preprocessing.py b/torchref/experimental/alignment/frf/preprocessing.py index 43c5cda2..915d48ac 100644 --- a/torchref/experimental/alignment/frf/preprocessing.py +++ b/torchref/experimental/alignment/frf/preprocessing.py @@ -43,84 +43,12 @@ def eterm_sigma_a(s_mag: torch.Tensor, delta_vrms_A: float) -> torch.Tensor: "build_lerf1_intensity", "apply_shell_variance_weights", "detect_zsymm", - "epsilon_aware_unroll", - "compute_v_budget", "bulk_solvent_factor", "oeffner_vrms", "fit_relative_wilson_b", ] -def epsilon_aware_unroll( - hkl_int: torch.Tensor, - sym_mats: torch.Tensor, -): - """Unroll each ASU reflection to the **unique** P1 positions in its orbit. - - For each ``h`` in the input list, generate the orbit ``{S_k · h}`` over the - ``n_ops`` symop matrices and emit one entry per *distinct* position. Axial / - special-position reflections (whose stabilizer has order ε(h) > 1) therefore - appear ``n_ops / ε(h)`` times, **not** ``n_ops`` times. - - Mirrors Phaser's ``if (!duplicate(isym, rhkl))`` skip in - ``DataMR.cc:954-986``. A naive ``einsum + reshape`` unroll over-counts axial - reflections by ε(h), polluting the obs SH coefficients with spurious - non-invariant content — the noise channel that hurts high-symmetry cases. - - Parameters - ---------- - hkl_int : (N, 3) integer tensor - ASU Miller indices. - sym_mats : (n_ops, 3, 3) tensor - Spacegroup rotation operators in the reciprocal (hkl) basis. Cast to - ``long`` internally; values must be integer. - - Returns - ------- - unrolled_hkl : (M, 3) long tensor — flat list of unique orbit positions - across all ASU reflections. - asu_idx : (M,) long tensor — index into ``hkl_int`` that each unrolled entry - came from. Callers use it to broadcast intensities / centric / sigF: - ``F_unrolled = F_obs[asu_idx]``. - """ - hkl_int = hkl_int.to(torch.long) - sym_mats = sym_mats.round().to(torch.long) - N, n_ops = hkl_int.shape[0], sym_mats.shape[0] - # Orbits: (N, n_ops, 3) — h.S_k, the row-vector (reciprocal-space) - # convention, matching the unroll sites in `align.py`. Note this is the - # TRANSPOSE contraction: `kji`, not `kij`. They coincide only for - # orthogonal symmetry matrices, so `kij` silently works everywhere except - # trigonal/hexagonal. - # Integer einsum dispatches to baddbmm, which CUDA does not implement for - # Long; compute in float64 (exact for symop 0/±1 × small Miller indices) - # and round back so the GPU path works. - orbits = ( - torch.einsum( - "kji,nj->nki", sym_mats.to(torch.float64), hkl_int.to(torch.float64), - ) - .round() - .to(torch.long) - ) - # Pack (h, k, l) into a single int64 key for per-row dedup. - base = 2 * int(orbits.abs().max().item()) + 1 - key = (orbits[:, :, 0] * base + orbits[:, :, 1]) * base + orbits[:, :, 2] - # Stable sort along dim=1 → duplicates land contiguously, lowest op-index - # first (matches Phaser's "first occurrence wins" rule). - sorted_keys, sort_idx = key.sort(dim=1, stable=True) - first_in_sorted = torch.cat( - [ - torch.ones(N, 1, dtype=torch.bool, device=key.device), - sorted_keys[:, 1:] != sorted_keys[:, :-1], - ], - dim=1, - ) - keep_mask = torch.empty_like(first_in_sorted) - keep_mask.scatter_(1, sort_idx, first_in_sorted) - asu_idx, op_idx = keep_mask.nonzero(as_tuple=True) - unrolled_hkl = orbits[asu_idx, op_idx] - return unrolled_hkl, asu_idx - - def build_lerf1_intensity( eEobs: torch.Tensor, centric_obs: torch.Tensor, @@ -218,67 +146,6 @@ def detect_zsymm(sym_mats: Optional[torch.Tensor]) -> int: return int(zsymm) -def compute_v_budget( - eps_factor: torch.Tensor, - sigma_a: torch.Tensor, - n_mol: int = 1, - totvar_known: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """Phaser's per-reflection variance budget ``V(h)`` for the m_LETF1 LL. - - Source: ``DataMR.cc:949`` (build) + ``DataMR.cc:1411`` (use in ``m_LETF1``):: - - V = PTNCS.EPSFAC[r] − totvar_known[r] − totvar_search[r] - - where ``EPSFAC[r] = ε(h)`` (or the tNCS-corrected variance bin in the - NCS-present case; we use plain ε for the standalone search), ``totvar_known`` - is the variance contribution from any fixed model (zero for a pure cross - rotation function), and ``totvar_search = σ_A²(s) · n_mol`` is the variance - explained by the moving model at the expected scattering content. - - For the cross-rotation case (no fixed model, ``totvar_known = 0``): - - V(h) = ε(h) − σ_A²(s)·n_mol - - Working in E-space (obs already Wilson-normalised), so no Σ_N factor. - - Parameters - ---------- - eps_factor : (N,) tensor - Per-reflection ε(h), the multiplicity (1 for general positions, - n>1 for reflections on n-fold symmetry axes). From - :meth:`torchref.symmetry.symmetry.Symmetry.epsilon`. - sigma_a : (N,) tensor - Per-reflection σ_A(s) (interpolated from the per-shell fit). - n_mol : int - Number of molecules in the unit cell summed over by the NSYMP-loop in - the calc-side expected intensity. Equals ``NSYMP`` for the standalone - cross-rotation search. - totvar_known : (N,) tensor, optional - Variance contribution from a fixed/known model. Default ``None`` → - treated as zero (standalone cross-rotation function). - - Returns - ------- - V : (N,) tensor - Per-reflection variance budget. Clamped to ``> 0`` to keep the LL finite; - a non-positive ``V`` would imply σ_A² overshoots ε, which Phaser also - guards against via ``PHASER_ASSERT(C > 0)`` at DataMR.cc:1413. - """ - sa = sigma_a.to(eps_factor.dtype) - moving = (sa * sa) * float(n_mol) - V = eps_factor - moving - if totvar_known is not None: - V = V - totvar_known.to(eps_factor.dtype) - return V.clamp(min=1e-6) - - -# ============================================================================= -# Phaser model-prep — three pieces Phaser applies before the FRF that we don't. -# See `Ensemble::setPDB` (EnsemblePDB.cc:40-100). Adding these as opt-in. -# ============================================================================= - - def bulk_solvent_factor( s_mag: torch.Tensor, fsol: float = 0.95, diff --git a/torchref/experimental/alignment/frf/rotation_utils.py b/torchref/experimental/alignment/frf/rotation_utils.py index 5d75b1f9..9b8c5df3 100644 --- a/torchref/experimental/alignment/frf/rotation_utils.py +++ b/torchref/experimental/alignment/frf/rotation_utils.py @@ -28,37 +28,6 @@ def rotation_matrix_from_edmonds_euler( return Rz_a @ Ry_b @ Rz_c -def rotation_matrix_from_edmonds_euler_batch( - alpha: torch.Tensor, beta: torch.Tensor, gamma: torch.Tensor, -) -> torch.Tensor: - """Vectorised Edmonds ZYZ Euler → ``R``. - - ``alpha``, ``beta``, ``gamma``: identically-shaped real tensors. Returns - ``(..., 3, 3)`` in the same dtype/device as the inputs. - """ - ca, sa = torch.cos(alpha), torch.sin(alpha) - cb, sb = torch.cos(beta), torch.sin(beta) - cg, sg = torch.cos(gamma), torch.sin(gamma) - zero = torch.zeros_like(alpha) - one = torch.ones_like(alpha) - Rz_a = torch.stack([ - torch.stack([ca, -sa, zero], dim=-1), - torch.stack([sa, ca, zero], dim=-1), - torch.stack([zero, zero, one], dim=-1), - ], dim=-2) - Ry_b = torch.stack([ - torch.stack([cb, zero, sb], dim=-1), - torch.stack([zero, one, zero], dim=-1), - torch.stack([-sb, zero, cb], dim=-1), - ], dim=-2) - Rz_c = torch.stack([ - torch.stack([cg, -sg, zero], dim=-1), - torch.stack([sg, cg, zero], dim=-1), - torch.stack([zero, zero, one], dim=-1), - ], dim=-2) - return Rz_a @ Ry_b @ Rz_c - - def edmonds_euler_from_rotation_matrix(R: torch.Tensor) -> Tuple[float, float, float]: """Recover ``(α, β, γ)`` such that ``R = R_z(α) R_y(β) R_z(γ)``. diff --git a/torchref/experimental/alignment/frf/sitelist_ang.py b/torchref/experimental/alignment/frf/sitelist_ang.py index 852ff1c2..652a824c 100644 --- a/torchref/experimental/alignment/frf/sitelist_ang.py +++ b/torchref/experimental/alignment/frf/sitelist_ang.py @@ -14,7 +14,7 @@ on the full ``(2L-1) × (2L-1)`` grid (asymmetric-unit storage only — the Friedel mate is added by cctbx via ``conjugate_flag=true``). 3. The 2D inverse FFT runs at a **fixed shape** - ``amax = adjust_gridding(2·max(bmax, lmax), max_prime=5)`` for every β. + ``amax = find_fft_friendly_size(2·max(bmax, lmax))`` for every β. The result is a dense ``M_β(α, γ)`` map indexed in ``[0, 1)`` along each axis. 4. The **adaptive sample list** is built once by ``allocate_memory`` @@ -41,37 +41,17 @@ import torch from ....config import canonical_device +from ....symmetry.symmetry import find_fft_friendly_size from .types import AdaptiveRotationFunction from .wigner_d import wigner_contraction_per_beta __all__ = [ - "adjust_gridding", "build_dense_map_per_beta", "build_adaptive_sample_list", "evaluate_rotation_function", ] -def adjust_gridding(target: int, max_prime: int = 5) -> int: - """Smallest integer ≥ target whose largest prime factor is ≤ max_prime. - - Phaser source: ``scitbx::fftpack::adjust_gridding`` (FastRot.cc:66-69). - For ``max_prime=5`` this is the standard 5-smooth (Hamming) numbers. - """ - if target <= 1: - return 1 - primes = [p for p in [2, 3, 5, 7, 11, 13] if p <= max_prime] - n = int(target) - while True: - m = n - for p in primes: - while m % p == 0: - m //= p - if m == 1: - return n - n += 1 - - def build_dense_map_per_beta( xi_lmn: torch.Tensor, betas: torch.Tensor, @@ -314,7 +294,7 @@ def evaluate_rotation_function( Phaser's ``grid_sampling`` keyword. β grid is uniform at this spacing; (α, γ) sample density per β follows pmax/qmax. fft_size : int, optional - Dense FFT shape. Default: ``adjust_gridding(2·max(bmax, 2L-1), 5)``. + Dense FFT shape. Default: ``find_fft_friendly_size(2·max(bmax, 2L-1))``. """ if xi_lmn.ndim != 3: raise ValueError(f"xi_lmn must be 3-D (L, 2L-1, 2L-1), got {tuple(xi_lmn.shape)}") @@ -328,7 +308,7 @@ def evaluate_rotation_function( bmax = int(math.ceil(180.0 / grid_sampling_deg)) if fft_size < 0: - fft_size = adjust_gridding(2 * max(bmax, 2 * L - 1), max_prime=5) + fft_size = find_fft_friendly_size(2 * max(bmax, 2 * L - 1)) # 1. Build adaptive sample list (purely geometric — independent of xi). alphas, betas_flat, gammas, beta_starts, beta_grid = build_adaptive_sample_list( diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index f0c3423f..ce4351f8 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -46,7 +46,7 @@ import time from contextlib import contextmanager from dataclasses import dataclass -from typing import List, Optional, Tuple, TYPE_CHECKING +from typing import List, Optional, TYPE_CHECKING import numpy as np import torch @@ -76,84 +76,6 @@ from torchref.model import ModelFT -def rotation_matrix_from_euler_zyz(alpha, beta, gamma) -> np.ndarray: - """Build R = R_z(α) R_y(β) R_z(γ) (Edmonds active ZYZ) as a NumPy 3×3 matrix. - - Compatibility wrapper around `rotation_matrix_from_edmonds_euler`. - """ - R = rotation_matrix_from_edmonds_euler(float(alpha), float(beta), float(gamma)) - return R.detach().cpu().numpy() - - -def rotation_angular_distance(R1: np.ndarray, R2: np.ndarray) -> float: - """Angular distance between two rotation matrices in degrees. - - The angular distance is the angle of the rotation ``R2 @ R1.T``. - """ - R_diff = R2 @ R1.T - trace = np.clip(np.trace(R_diff), -1.0, 3.0) - return np.degrees(np.arccos((trace - 1.0) / 2.0)) - - -def euler_angular_distance( - euler1: Tuple[float, float, float], - euler2: Tuple[float, float, float], -) -> float: - """Angular distance between two ZYZ Euler angle sets (degrees).""" - R1 = rotation_matrix_from_euler_zyz(*euler1) - R2 = rotation_matrix_from_euler_zyz(*euler2) - return rotation_angular_distance(R1, R2) - - -def cluster_rotation_peaks( - peaks: list, - threshold_deg: float = 6.0, - symmetry_matrices: Optional[np.ndarray] = None, -) -> list: - """Cluster rotation peaks by angular distance. - - Peaks within ``threshold_deg`` of each other are considered the same - solution; only the highest-scoring peak from each cluster is kept. Not on - the default pipeline path (the ML rescore already ranks Patterson- - equivalents adjacently); retained for callers that want explicit - de-duplication. - - Parameters - ---------- - peaks : list - Rotation peaks as tuples ``(alpha, beta, gamma, score, sigma)``. - threshold_deg : float - Angular distance threshold for clustering (degrees). - symmetry_matrices : np.ndarray, optional - Point-group symmetry matrices (N, 3, 3) to check symmetry equivalents. - """ - if not peaks: - return [] - - sorted_peaks = sorted(peaks, key=lambda p: p[4], reverse=True) - clustered = [] - used_rotations = [] - for peak in sorted_peaks: - alpha, beta, gamma, score, sigma = peak - R = rotation_matrix_from_euler_zyz(alpha, beta, gamma) - is_new = True - for R_used in used_rotations: - if rotation_angular_distance(R, R_used) < threshold_deg: - is_new = False - break - if symmetry_matrices is not None: - for sym_op in symmetry_matrices: - if rotation_angular_distance(sym_op @ R, R_used) < threshold_deg: - is_new = False - break - if not is_new: - break - if is_new: - clustered.append(peak) - used_rotations.append(R) - return clustered - - # --------------------------------------------------------------------------- # Stage timing # --------------------------------------------------------------------------- diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index 9089794b..852bf025 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -441,9 +441,9 @@ def search_peaks( # # Not the same as the earlier finding that knocking it out was # rank-neutral; that was a measurement about whether it mattered, this is - # that it is arithmetically cancelled. `fit_relative_wilson_b` stays for - # the rescore, where the calc normalisation is taken from the unmodified - # reference amplitudes and the Debye-Waller term therefore survives. + # that it is arithmetically cancelled. `fit_relative_wilson_b` survives in + # `frf/preprocessing` with no production caller at all -- it was kept for + # the ML rescore, and that was deleted. engine = FastRotationFunction( s_obs, F_obs, centric, sg_mats, @@ -471,15 +471,13 @@ def search_peaks( def _solutions(peaks: List["RotationPeak"], lmax: int, d_min: float, model_error_A: float) -> RotationSolutions: """Package a peak list as the public return type.""" - from .frf.rotation_utils import rotation_matrix_from_edmonds_euler_batch + from torchref.base.alignment.rotation import rotation_matrix_euler_zyz euler = torch.tensor( [[p.alpha, p.beta, p.gamma] for p in peaks], dtype=torch.float64, ).reshape(-1, 3) rotations = ( - rotation_matrix_from_edmonds_euler_batch( - euler[:, 0], euler[:, 1], euler[:, 2], - ) + rotation_matrix_euler_zyz(euler) if euler.numel() else torch.zeros((0, 3, 3), dtype=torch.float64) ) From 6edafba8b131046b678e105650d8b348192dc368 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Tue, 1 Sep 2026 21:06:25 +0200 Subject: [PATCH 140/250] Remeasure the candidate ranking arms after the variance fix The `rank_by` default was chosen on a 37/40-against-36/40 sweep taken with the factor-of-two in the acentric Rice variance. Remeasured against the shipping tree, over the same four structures x ten seeds: llg 36/40 median residual 1.43 deg r 36/40 1.62 corr 32/40 1.98 The 37 does not reproduce. The likelihood and R tie on the count and each wins exactly one cell paired over the 40, so nothing separates them there. The margin that does exist is narrower than the medians imply. On 1DAW and 3K7M the two arms pick the same candidate and the residuals are identical; the median paired difference over the cells all three solve is +0.000 deg. The whole difference is 6G9X, where the likelihood holds every residual under 2.3 deg and R spreads to 5.65, and 2DQ6 goes the other way by a smaller margin. So the default now rests on one structure rather than on a sweep-wide margin, and the comment says that instead of implying otherwise. The default is unchanged: the likelihood is still the right object for the question, an R-factor on a partial model at this resolution has little to distinguish with, and the arm is selectable if that stops holding. The correlation arm's recorded score was 31/40 in two places and 32/40 in a third; it is 32/40, and the disagreement was between a figure measured once and a figure copied. The historical 31/40 in truth_metric_check is kept as a historical number and labelled as one, since it belongs to a run on older code. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/diagnostics/pose_recovery.py | 16 ++++++---- .../diagnostics/truth_metric_check.py | 6 ++-- docs/changelog.rst | 2 +- torchref/experimental/alignment/pipeline.py | 32 +++++++++++++------ 4 files changed, 36 insertions(+), 20 deletions(-) diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py index acf5a601..f6441bd0 100644 --- a/alignment_lab/diagnostics/pose_recovery.py +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -12,15 +12,17 @@ ``llg`` the default -- rank each rotation candidate by the translation likelihood at - its best translation. 37/40 over four structures x ten seeds. + its best translation. 36/40 over four structures x ten seeds, median + residual 1.43 deg. ``analytic_r`` - rank by the analytical-scale R instead. 36/40, and places less well on the - cells both solve. + rank by the analytical-scale R instead. Also 36/40, median 1.62 deg. The two + tie on the count and each wins one cell paired; they pick the same candidate + outright on 1DAW and 3K7M. The likelihood's margin is 6G9X alone. ``corr`` - rank by the translation function's own correlation. Measured 31/40 against - ``analytic_r``'s 36/40 over four structures x ten seeds: worse, despite a - rank-level harness predicting the reverse on a truth label that disagreed - with coordinate superposition. + rank by the translation function's own correlation. Measured 32/40 against + the other two arms' 36/40: worse, and worse paired against ``llg`` 5 to 1, + despite a rank-level harness predicting the reverse on a truth label that + disagreed with coordinate superposition. ``llg_tf`` a different question -- re-rank each candidate's TRANSLATIONS by the likelihood, still selecting the candidate by R. diff --git a/alignment_lab/diagnostics/truth_metric_check.py b/alignment_lab/diagnostics/truth_metric_check.py index dc2a8f34..660c8aae 100644 --- a/alignment_lab/diagnostics/truth_metric_check.py +++ b/alignment_lab/diagnostics/truth_metric_check.py @@ -16,8 +16,10 @@ That matters beyond bookkeeping: the rank harness said ranking by the translation correlation would beat the analytic R by 33/40 to 23/40, and end to -end it lost 31/40 to 36/40. A truth label that is wrong makes every rank in that -harness meaningless, so this checks it directly rather than by inference. +end it lost 31/40 to 36/40 -- as measured then; the correlation arm scores 32/40 +on the current code, so the direction is unchanged. A truth label that is wrong +makes every rank in that harness meaningless, so this checks it directly rather +than by inference. """ import sys from pathlib import Path diff --git a/docs/changelog.rst b/docs/changelog.rst index 09c6fe17..21347d78 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -8,7 +8,7 @@ Unreleased - The translation likelihood's model error comes from the shared ``SigmaAEstimator`` instead of a local 81-point scan over every reflection. It returns ``alpha`` and ``beta`` per reflection rather than a per-shell ``sigma_A``, so the likelihood no longer assumes ```` is exactly one. Outcome-neutral over 70 seeded cells, zero flips; not measurably faster - Fixed the translation likelihood's variance convention, which scored acentric reflections at twice the variance intended -- 90-95% of reflections. The alignment package carried its own Rice and Woolfson parameterised by the *amplitude* variance and handed both branches the same number, where the acentric branch needs half what the centric one does. It now uses ``base.targets.xray_likelihoods.rice_per_refl``, which takes the complex variance and derives the centric case from it - Removed ``experimental/alignment/distributions.py``. Its ``stable_log_bessel_i0`` also carried a wrong asymptotic coefficient, giving -2.9e-5 at x = 50 against -5.4e-7 for the correct term; the shared implementation uses ``log(i0e(z)) + z``, which is exact -- Molecular-replacement candidates are ranked by the translation function's likelihood, not by an analytical-scale R-factor. 30/30 on the ten-structure panel against 29/30, and 37/40 against 36/40 over a ten-seed sweep, with tighter placements on the cells both solve. Selectable through ``rank_by``; the correlation is the third option and is the worst of the three at 32/40, despite a rank-level harness rating it best on a truth label that disagrees with coordinate superposition +- Molecular-replacement candidates are ranked by the translation function's likelihood, not by an analytical-scale R-factor. 30/30 on the ten-structure panel against 29/30. Over a ten-seed sweep the two tie at 36/40 and the likelihood's advantage is placement rather than count -- median residual 1.43 deg against 1.62, carried by one structure. Selectable through ``rank_by``; the correlation is the third option and is the worst of the three at 32/40, despite a rank-level harness rating it best on a truth label that disagrees with coordinate superposition - The likelihood ranking costs about 2.5x the wall clock, from a per-candidate sigma_A fit - The molecular-replacement pipeline's ``verbose`` levels are a documented contract routed through one emitter, rather than ``if verbose > 0: print(...)`` at seventeen sites. Level 2 emits one machine-readable ``CAND`` line per rotation candidate carrying every score the selection could have used, so a wrong placement can be diagnosed from the run itself instead of from a harness that re-implements the placement loop and then disagrees with it - The molecular-replacement pipeline places all ``n_rotation_candidates`` (now 25, was 15) and returns the best, instead of stopping once a placement beat an R-factor threshold. The old rule made the answer depend on the order the rotation function happened to produce and could accept the third candidate without scoring the tenth. Measured neutral -- identical placements on 10 structures x 3 seeds and on a 10-seed sweep of the marginal cases -- at about 1.6x the wall clock diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index ce4351f8..af57ec5f 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -441,16 +441,28 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: # Highest translation likelihood. Over four structures x ten seeds, # success within 8 deg of canonical modulo crystal symmetry: # - # llg 37/40 median residual 1.57 deg - # r 36/40 1.79 - # corr 32/40 2.77 + # llg 36/40 median residual 1.43 deg + # r 36/40 1.62 + # corr 32/40 1.98 # - # The likelihood is chosen on the residuals rather than the success - # count -- one cell in 40 is not a result, but on the cells all three - # solve it places better 6 times against 2, and on 6G9X every residual - # falls under 2.3 deg where R spreads to 5.65. It is also the right - # object for the question: an R-factor on a partial model at the - # resolution this runs at has little to distinguish with. + # The likelihood and R tie on the success count -- 36 each, and paired + # over the 40 cells each wins exactly one. Nothing separates them there, + # and an earlier 37-against-36 reading of this table did not survive + # remeasurement after the variance convention was corrected. + # + # What separates them is where they differ at all, which is less often + # than the medians suggest: on 1DAW and 3K7M the two pick the SAME + # candidate and the residuals are identical. The whole difference is + # 6G9X, where the likelihood holds every residual under 2.3 deg and R + # spreads to 5.65 (medians 1.05 against 2.24). 2DQ6 goes the other way + # by a smaller margin (max 6.48 against 3.86). Across the 31 cells all + # three arms solve, the likelihood places closer 5 times against 1. + # + # So the default rests on one structure, not on a sweep-wide margin. It + # is kept because it is also the right object for the question -- an + # R-factor on a partial model at the resolution this runs at has little + # to distinguish with -- and because the arm is selectable if that + # reasoning ever stops holding. # # The correlation is here as a cautionary default-not-taken. A rank-level # harness rated it best by a wide margin, 33/40 against R's 23/40, and @@ -644,7 +656,7 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: # Both scores at the REFINED position, so the reported correlation # belongs to the translation that was actually chosen. Selection is # by R: ranking candidates by the correlation instead was measured - # end to end and is WORSE (31/40 against 36/40 over four structures + # end to end and is WORSE (32/40 against 36/40 over four structures # x ten seeds), despite a rank-level harness predicting the reverse. tf_ref = correlation_at(self._obs, G_pre, h_R_pre, t_refined) self._log(3, f" trans{k_t}: tf={tf_ref:.5f} " From 1c4e4b43e24750747b1a8ad8b1c53b484a7c4e66 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 2 Sep 2026 11:40:14 +0200 Subject: [PATCH 141/250] Judge pose recovery against Cartesian mates and check the translation The lab's success metric compared the Kabsch rotation, a Cartesian matrix, against the fractional symmetry matrices. In P3(1)21 two of the six mates of a correct solution read as 30.00 and 21.09 degrees, and in P6(5)22 four of twelve read as 21.09; every 2DQ6 failure on record was one of those mates. It also never checked the translation, so a placement at the right orientation and 40-55 A from the true position counted as a success. pose_error compares against B S B^-1 and measures the centroid offset from the closest symmetry image modulo lattice translations, the group's allowed origin shifts and its polar directions. pose_recovery gates on both. truth_pose_scores scores the deposited pose through the pipeline's own path so a search failure can be told from a scoring failure. Measured (jobs 543997, 544006, 544011, 544016): 2DQ6 is 10/10 in rotation on all three ranking arms; the default pipeline mis-translates 2DQ6, 3VRJ, 4BX9 and 6G9X by 20-56 A on every trial (18/30 true poses); a 15-4 A translation window gives 30/30 within 0.32 A. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/pose_panel_trans.sh | 27 +++++ .../analysis/trigonal_metric_recheck.sh | 36 ++++++ alignment_lab/analysis/truth_pose_scores.sh | 22 ++++ alignment_lab/analysis/truth_pose_scores2.sh | 25 ++++ alignment_lab/diagnostics/pose_recovery.py | 48 ++++++-- .../diagnostics/truth_pose_scores.py | 107 ++++++++++++++++++ alignment_lab/lab/__init__.py | 6 + alignment_lab/lab/truth.py | 87 ++++++++++++++ 8 files changed, 346 insertions(+), 12 deletions(-) create mode 100644 alignment_lab/analysis/pose_panel_trans.sh create mode 100644 alignment_lab/analysis/trigonal_metric_recheck.sh create mode 100644 alignment_lab/analysis/truth_pose_scores.sh create mode 100644 alignment_lab/analysis/truth_pose_scores2.sh create mode 100644 alignment_lab/diagnostics/truth_pose_scores.py diff --git a/alignment_lab/analysis/pose_panel_trans.sh b/alignment_lab/analysis/pose_panel_trans.sh new file mode 100644 index 00000000..45940072 --- /dev/null +++ b/alignment_lab/analysis/pose_panel_trans.sh @@ -0,0 +1,27 @@ +#!/bin/bash +# The 10 x 3 panel with the translation error reported, at the default +# translation window (all data) and at 15-4 A. The old gate was rotation-only. +#SBATCH --job-name=ptrans +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=hour +#SBATCH --time=00:59:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-9 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +for T in 0 1 2; do + "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial $T --arms llg \ + 2>/dev/null | grep '^ROW ' | sed 's/^ROW/ROW window=full/' + "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial $T --arms llg \ + --tf-d-min 4.0 --tf-d-max 15.0 2>/dev/null | grep '^ROW ' | sed 's/^ROW/ROW window=15-4/' +done +echo DONE diff --git a/alignment_lab/analysis/trigonal_metric_recheck.sh b/alignment_lab/analysis/trigonal_metric_recheck.sh new file mode 100644 index 00000000..62e8043e --- /dev/null +++ b/alignment_lab/analysis/trigonal_metric_recheck.sh @@ -0,0 +1,36 @@ +#!/bin/bash +# Were 2DQ6's "failures" symmetry mates the success metric could not recognise? +# +# residual_rotation_deg compared a Cartesian Kabsch rotation against the +# FRACTIONAL symmetry matrices. In P3(1)21 two of the six mates of a correct +# solution then read as 30.00 and 21.09 deg; 2DQ6's failing residuals were +# 28.5-31.0 and 19.5-21.8. This re-scores the same seeds with Cartesian mates +# and a translation check modulo allowed origin shifts. Controls: 3GR5 (P6(5)22, +# four of twelve mates affected), 6G9X t4 corr (a genuine 55.9 deg miss, must +# stay a miss) and 1DAW t0 (monoclinic, must be unchanged at 1.519). +#SBATCH --job-name=trig +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=hour +#SBATCH --time=00:59:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-3 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 +export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +run() { "$PY" -u alignment_lab/diagnostics/pose_recovery.py "$@" 2>/dev/null | grep -E '^ROW |SOLN|^===' ; } +case $SLURM_ARRAY_TASK_ID in + 0) for T in $(seq 0 9); do run --pdb 2DQ6 --trial $T --arms llg; done ;; + 1) for T in $(seq 0 9); do run --pdb 2DQ6 --trial $T --arms analytic_r; done ;; + 2) for T in $(seq 0 9); do run --pdb 2DQ6 --trial $T --arms corr; done ;; + 3) for T in 0 1 2; do run --pdb 3GR5 --trial $T --arms llg,analytic_r,corr; done + run --pdb 6G9X --trial 4 --arms corr + run --pdb 1DAW --trial 0 --arms llg + run --pdb 2DQ6 --trial 3 --arms llg --verbose 2 ;; +esac +echo DONE diff --git a/alignment_lab/analysis/truth_pose_scores.sh b/alignment_lab/analysis/truth_pose_scores.sh new file mode 100644 index 00000000..3a74c51b --- /dev/null +++ b/alignment_lab/analysis/truth_pose_scores.sh @@ -0,0 +1,22 @@ +#!/bin/bash +#SBATCH --job-name=tps +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=hour +#SBATCH --time=00:30:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-2 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +case $SLURM_ARRAY_TASK_ID in + 0) "$PY" -u alignment_lab/diagnostics/truth_pose_scores.py --pdb 2DQ6 --trial 3 ;; + 1) "$PY" -u alignment_lab/diagnostics/truth_pose_scores.py --pdb 2DQ6 --trial 0 ;; + 2) "$PY" -u alignment_lab/diagnostics/truth_pose_scores.py --pdb 3GR5 --trial 0 ;; +esac +echo DONE diff --git a/alignment_lab/analysis/truth_pose_scores2.sh b/alignment_lab/analysis/truth_pose_scores2.sh new file mode 100644 index 00000000..44a910e6 --- /dev/null +++ b/alignment_lab/analysis/truth_pose_scores2.sh @@ -0,0 +1,25 @@ +#!/bin/bash +#SBATCH --job-name=tps2 +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=hour +#SBATCH --time=00:40:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-4 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +T="$PY -u alignment_lab/diagnostics/truth_pose_scores.py" +case $SLURM_ARRAY_TASK_ID in + 0) $T --pdb 2DQ6 --trial 3 --tf-d-min 4.0 --tf-d-max 15.0 ;; + 1) $T --pdb 2DQ6 --trial 3 --tf-d-min 3.0 ;; + 2) $T --pdb 4BX9 --trial 0 ;; + 3) $T --pdb 6G9X --trial 0 ;; + 4) $T --pdb 1DAW --trial 0 ;; +esac +echo DONE diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py index f6441bd0..ba82618f 100644 --- a/alignment_lab/diagnostics/pose_recovery.py +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -27,8 +27,11 @@ a different question -- re-rank each candidate's TRANSLATIONS by the likelihood, still selecting the candidate by R. -Success mirrors the integration test: final coordinates within ``--success-deg`` -of canonical, modulo the crystal symmetry. +Success is a pose: final coordinates within ``--success-deg`` of canonical in +orientation AND within ``--success-A`` of it in position, modulo the crystal +symmetry (Cartesian point-group mates, lattice translations, allowed origin +shifts and polar directions). The gate used to be rotation-only, and it passed +placements 40-55 A from the true position on 2DQ6, 4BX9 and 6G9X. Usage:: @@ -51,8 +54,8 @@ # rigid-body polish are LBFGS -- they need autograd, and disabling it # raises "element 0 of tensors does not require grad". -from lab import (BENCH_PDBS, ResultWriter, load_case, random_rotation, # noqa: E402 - seed_for) +from lab import (BENCH_PDBS, ResultWriter, cartesian_symops, load_case, # noqa: E402 + pose_error, random_rotation, seed_for) ARMS = { # How the winner is chosen among placed candidates. @@ -65,11 +68,17 @@ } -def residual_rotation_deg(aligned_xyz, canonical_xyz, symops) -> float: +def residual_rotation_deg(aligned_xyz, canonical_xyz, symops_cart) -> float: """Smallest angle between the aligned-to-canonical rotation and any symop. Kabsch superposition, then compared against every symmetry operator -- a solution differing from canonical by a crystal symmetry is correct. + + ``symops_cart`` must be the **Cartesian** rotations, ``B S_k B^-1`` + (:func:`lab.cartesian_symops`). This used to take ``spacegroup.matrices`` + directly, which act on fractional coordinates: for trigonal and hexagonal + cells two of the six (four of the twelve) mates of a correct solution then + read as 30.00 and 21.09 degrees, and 2DQ6's "bimodal 6/10" was those mates. """ from torchref.experimental.alignment.frf.rotation_utils import ( rotation_angular_distance_deg, @@ -81,8 +90,8 @@ def residual_rotation_deg(aligned_xyz, canonical_xyz, symops) -> float: U, _, Vt = torch.linalg.svd(Qc.T @ Pc) d = torch.sign(torch.det(U @ Vt)) R = U @ torch.diag(torch.tensor([1.0, 1.0, d], dtype=torch.float64)) @ Vt - return min(float(rotation_angular_distance_deg(R, symops[k])) - for k in range(symops.shape[0])) + return min(float(rotation_angular_distance_deg(R, symops_cart[k])) + for k in range(symops_cart.shape[0])) #: The rotation search's own bandwidth constant. Recorded in every row because @@ -144,8 +153,16 @@ def main() -> int: ap.add_argument("--n-rotation-candidates", type=int, default=25) ap.add_argument("--n-rotation-peaks", type=int, default=200) ap.add_argument("--success-deg", type=float, default=8.0) + # A placement is a POSE: rotation and translation. The translation gate is + # generous -- downstream rigid-body refinement absorbs a few Angstrom -- but + # it separates a found position from one 40 A away, which the rotation-only + # gate this harness used to apply could not. Three large structures passed + # that gate on every seed while sitting 40-55 A from the true position. + ap.add_argument("--success-A", type=float, default=4.0) ap.add_argument("--verbose", type=int, default=0) ap.add_argument("--out-csv", default=None) + ap.add_argument("--tf-d-min", type=float, default=None) + ap.add_argument("--tf-d-max", type=float, default=None) args = ap.parse_args() from torchref.experimental.alignment import MolecularReplacementPipeline @@ -153,7 +170,8 @@ def main() -> int: seed = seed_for(args.pdb, args.trial) model, data = load_case(args.pdb) canonical_xyz = model.xyz().clone() - symops = data.spacegroup.matrices.to(torch.float64).cpu() + # Cartesian mates, not the fractional matrices -- see residual_rotation_deg. + symops = cartesian_symops(data.spacegroup, data.cell) R_true = random_rotation(seed) print(f"=== {args.pdb} t{args.trial} seed={seed} {data.spacegroup} " @@ -186,21 +204,27 @@ def main() -> int: data, search, d_min=4.0, d_max=15.0, n_shells=20, n_rotation_peaks=args.n_rotation_peaks, n_rotation_candidates=args.n_rotation_candidates, - verbose=args.verbose, **flags, + verbose=args.verbose, tf_d_min=args.tf_d_min, + tf_d_max=args.tf_d_max, **flags, ) solutions = pipe.run(do_translation=True) aligned = solutions[0].model resid = residual_rotation_deg(aligned.xyz(), canonical_xyz, symops) + _, trans_A = pose_error(aligned.xyz(), canonical_xyz, data.cell, + data.spacegroup) err = "" if args.verbose >= 2: _report_candidates(solutions, R_true, symops, args.success_deg) except Exception as exc: # a crashed arm must not read as a success - resid, err = float("nan"), f"{type(exc).__name__}: {exc}" + resid, trans_A = float("nan"), float("nan") + err = f"{type(exc).__name__}: {exc}" secs = time.time() - t0 - ok = (resid == resid) and resid <= args.success_deg + ok = ((resid == resid) and resid <= args.success_deg + and (trans_A == trans_A) and trans_A <= args.success_A) print(f"ROW {arm} {args.pdb} trial={args.trial} " f"n_cand={args.n_rotation_candidates} " - f"resid={resid:.3f} ok={int(bool(ok))} seconds={secs:.1f}", + f"resid={resid:.3f} trans_A={trans_A:.2f} ok={int(bool(ok))} " + f"seconds={secs:.1f}", flush=True) print(f" {arm:16s} {resid:10.2f} {('yes' if ok else 'NO'):>4s} {secs:9.1f}" + (f" {err}" if err else "")) diff --git a/alignment_lab/diagnostics/truth_pose_scores.py b/alignment_lab/diagnostics/truth_pose_scores.py new file mode 100644 index 00000000..22f148ae --- /dev/null +++ b/alignment_lab/diagnostics/truth_pose_scores.py @@ -0,0 +1,107 @@ +"""Is a returned placement the deposited pose, and if not, does the deposited pose score better? + +Runs the pipeline on a seeded reorientation, then compares the winner with the +deposited model under the pipeline's own three selection scores, evaluated +through the same ``TranslationObs`` and ``precompute_G`` path. Reports the +rotation and translation error of the winner against the closest symmetry +image of the deposited model, and the raw fractional centroid offset to every +image, so a pseudo-translation shows up as a specific vector. + +If the deposited pose scores clearly better than the winner, the translation +search missed it. If they score the same, the data cannot tell them apart. +""" +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from lab import (BENCH_PDBS, allowed_origin_shifts, load_case, pose_error, # noqa: E402 + random_rotation, seed_for) + + +def scores_at(pipe, model_placed): + """(corr, R, llg) of an already-placed model through the pipeline's path.""" + from torchref.experimental.alignment.translation import ( + DirectModelEvaluator, correlation_at, fit_model_error, llg_at, + normalise_calc, precompute_G_for_rotation) + data, obs = pipe.data, pipe._obs + m = model_placed.copy() + m.spacegroup = "P 1" + ev = DirectModelEvaluator(m) + eye3 = torch.eye(3, dtype=torch.float64) + G, h_R = precompute_G_for_rotation(ev, eye3, obs.hkl, data.spacegroup, data.cell) + t0 = torch.zeros(3, dtype=torch.float64) + corr = correlation_at(obs, G, h_R, t0) + Fc = G.sum(dim=0).abs().to(torch.float64) + Fo = obs.F_obs.to(Fc.device).to(torch.float64) + k = (Fo * Fc).sum() / (Fc * Fc).sum().clamp(min=1e-30) + r = float(((Fo - k * Fc).abs().sum() / Fo.sum()).item()) + E_calc = normalise_calc(Fc, obs) + alpha, beta = fit_model_error(obs, E_calc) + llg = llg_at(obs, G, h_R, t0, alpha, beta) + return corr, r, llg + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="2DQ6", choices=list(BENCH_PDBS)) + ap.add_argument("--trial", type=int, default=3) + ap.add_argument("--rank-by", default="llg") + ap.add_argument("--tf-d-min", type=float, default=None) + ap.add_argument("--tf-d-max", type=float, default=None) + args = ap.parse_args() + + from torchref.experimental.alignment import MolecularReplacementPipeline + + model, data = load_case(args.pdb) + canonical = model.xyz().clone() + seed = seed_for(args.pdb, args.trial) + R_true = random_rotation(seed) + shifts, polar = allowed_origin_shifts(data.spacegroup) + print(f"=== {args.pdb} t{args.trial} {data.spacegroup} tf_window=({args.tf_d_max},{args.tf_d_min}) allowed shifts " + f"{[tuple(round(float(x), 3) for x in u) for u in shifts]} polar dims {polar.shape[1]}") + + search = model.copy() + search.spacegroup = "P 1" + search = search.copy().rotate(R_true.to(model.dtype_float), center=canonical.mean(0)) + pipe = MolecularReplacementPipeline( + data, search, d_min=4.0, d_max=15.0, n_shells=20, + n_rotation_peaks=200, n_rotation_candidates=25, rank_by=args.rank_by, + tf_d_min=args.tf_d_min, tf_d_max=args.tf_d_max, + ) + sols = pipe.run(do_translation=True) + win = sols[0] + + # Sanity: the deposited model against itself must be (0, 0). + print("SANITY deposited-vs-deposited rot/trans:", + pose_error(canonical, canonical, data.cell, data.spacegroup)) + rot, trans = pose_error(win.model.xyz(), canonical, data.cell, data.spacegroup) + print(f"WINNER rot_deg={rot:.3f} trans_A={trans:.2f} corr={win.translation_score:.5f} " + f"R={win.r_factor:.5f} llg={win.llg_score:.1f}") + + # Raw fractional centroid offset to every symmetry image of canonical. + B = data.cell.fractional_matrix.detach().cpu().to(torch.float64) + Binv = torch.linalg.inv(B) + S = data.spacegroup.matrices.detach().cpu().to(torch.float64) + T = data.spacegroup.translations.detach().cpu().to(torch.float64) + ca = Binv @ win.model.xyz().detach().cpu().to(torch.float64).mean(0) + cc = Binv @ canonical.detach().cpu().to(torch.float64).mean(0) + for k in range(S.shape[0]): + d = ca - (S[k] @ cc + T[k]) + d = d - d.round() + print(f" image {k}: centroid offset frac=({d[0]:+.3f},{d[1]:+.3f},{d[2]:+.3f}) " + f"|.|={float((B @ d).norm()):.1f} A") + + c_dep, r_dep, llg_dep = scores_at(pipe, model) + print(f"DEPOSITED corr={c_dep:.5f} R={r_dep:.5f} llg={llg_dep:.1f}") + c_w, r_w, llg_w = scores_at(pipe, win.model) + print(f"WINNER(re-scored) corr={c_w:.5f} R={r_w:.5f} llg={llg_w:.1f}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/lab/__init__.py b/alignment_lab/lab/__init__.py index 191d6971..312adbd8 100644 --- a/alignment_lab/lab/__init__.py +++ b/alignment_lab/lab/__init__.py @@ -16,7 +16,10 @@ rotated_case, ) from .truth import ( + allowed_origin_shifts, + cartesian_symops, orbit_rank, + pose_error, random_rotation, seed_for, symmetry_orbit, @@ -39,7 +42,10 @@ "case_paths", "load_case", "rotated_case", + "allowed_origin_shifts", + "cartesian_symops", "orbit_rank", + "pose_error", "random_rotation", "seed_for", "symmetry_orbit", diff --git a/alignment_lab/lab/truth.py b/alignment_lab/lab/truth.py index c5ffbbc0..492d3318 100644 --- a/alignment_lab/lab/truth.py +++ b/alignment_lab/lab/truth.py @@ -196,3 +196,90 @@ def orbit_rank( if ang <= thr_deg and rank < 0: rank = i return rank, best + + +def cartesian_symops(spacegroup, cell) -> torch.Tensor: + """The point-group rotations as **Cartesian** matrices, ``B S_k B^-1``. + + ``spacegroup.matrices`` act on fractional column vectors, ``x' = S x + t``. + A Kabsch rotation between two sets of Cartesian coordinates lives in the + Cartesian frame, and comparing it against ``S_k`` directly is only correct + when ``B S_k B^-1 == S_k`` -- diagonal ``S`` in an orthogonal cell, or a + cubic cell. In P3(1)21 two of the six mates of a *correct* solution read as + 30.00 and 21.09 degrees under that comparison, and in P6(5)22 four of + twelve read as 21.09; those were the "bimodal" 2DQ6 failures. + + Returns ``(n_ops, 3, 3)`` float64 on the host. + """ + B = cell.fractional_matrix.detach().cpu().to(torch.float64) # c = B x + S = spacegroup.matrices.detach().cpu().to(torch.float64) + return B @ S @ torch.linalg.inv(B) + + +def allowed_origin_shifts(spacegroup, n: int = 12) -> Tuple[torch.Tensor, torch.Tensor]: + """Fractional translations ``u`` that leave the space group invariant. + + ``u`` is allowed when ``(S_k - I) u`` is a lattice vector for every op -- + two placements differing by such a ``u`` give identical ``|F|`` and are the + same solution. Returns ``(discrete, polar)``: the discrete shifts on a + ``1/n`` grid (``n=12`` covers 1/2, 1/3, 1/4 and 1/6), and an orthonormal + basis of the continuous (polar) directions, ``(3, p)``. + """ + S = spacegroup.matrices.detach().cpu().to(torch.float64) + eye = torch.eye(3, dtype=torch.float64) + D = torch.cat([Sk - eye for Sk in S], dim=0) # (3 n_ops, 3) + # Polar directions: null space of D. + _, sv, Vh = torch.linalg.svd(D) + null = (sv < 1e-8).sum().item() if sv.numel() else 3 + polar = Vh[3 - null:].T if null else torch.zeros(3, 0, dtype=torch.float64) + g = torch.arange(n, dtype=torch.float64) / n + U = torch.cartesian_prod(g, g, g) # (n^3, 3) + resid = torch.einsum("oij,uj->uoi", S - eye, U) # (n^3, n_ops, 3) + ok = ((resid - resid.round()).abs() < 1e-6).all(dim=-1).all(dim=-1) + return U[ok], polar + + +def pose_error( + aligned_xyz: torch.Tensor, + canonical_xyz: torch.Tensor, + cell, + spacegroup, +) -> Tuple[float, float]: + """``(rotation_deg, translation_A)`` of a placement against the deposited pose. + + Rotation: Kabsch superposition of the placed atoms onto the canonical ones, + compared against every **Cartesian** point-group mate (see + :func:`cartesian_symops`). Translation: the centroid offset from the closest + symmetry image of the canonical model, modulo lattice vectors, the group's + allowed origin shifts and its polar directions, in Angstrom. Both are zero + for a placement that is the deposited structure or any symmetry-equivalent + copy of it. + """ + P = canonical_xyz.detach().cpu().to(torch.float64) + Q = aligned_xyz.detach().cpu().to(torch.float64) + Pc, Qc = P - P.mean(0), Q - Q.mean(0) + U, _, Vt = torch.linalg.svd(Qc.T @ Pc) + d = torch.sign(torch.det(U @ Vt)) + R = U @ torch.diag(torch.tensor([1.0, 1.0, d], dtype=torch.float64)) @ Vt + + B = cell.fractional_matrix.detach().cpu().to(torch.float64) + Binv = torch.linalg.inv(B) + S = spacegroup.matrices.detach().cpu().to(torch.float64) + T = spacegroup.translations.detach().cpu().to(torch.float64) + R_cart = B @ S @ Binv + + tr = torch.einsum("kij,ij->k", R_cart, R) + ang = ((tr - 1.0) * 0.5).clamp(-1.0, 1.0).arccos() * (180.0 / math.pi) + k_best = int(ang.argmin()) + rot_deg = float(ang[k_best]) + + # Translation, against the mate whose rotation matched. + shifts, polar = allowed_origin_shifts(spacegroup) + cen_a = Binv @ Q.mean(0) # fractional + cen_c = S[k_best] @ (Binv @ P.mean(0)) + T[k_best] + delta = (cen_a - cen_c).unsqueeze(0) - shifts # (n_u, 3) + delta = delta - delta.round() + if polar.shape[1]: + delta = delta - (delta @ polar) @ polar.T + trans_A = float((delta @ B.T).norm(dim=-1).min()) + return rot_deg, trans_A From c551305cbfe3dae68ee1f581f2a9097eab2924bb Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 2 Sep 2026 11:49:42 +0200 Subject: [PATCH 142/250] Default the translation window to the rotation window and grid the P1 copy to it The translation search ran on all data -- 228k reflections to 1.5 A on 2DQ6 -- and on the four largest panel structures it placed the model at the right orientation and 20-56 A from the true position, on every trial. Its own score is higher at the wrong place than at the deposited pose (0.665 against 0.350 on 2DQ6, where the likelihood is 1616 against 157865): the objective's calc side is raw |F_calc|^2, so at high resolution it follows whichever reflections carry the largest calculated intensity. The benchmark never saw this because it checked the rotation only. The window now defaults to the rotation search's [d_max, d_min], one resolution window and one Wilson normalisation for both stages; 0.0/inf remove a cut. The P1 copy is gridded at tf_d_min/1.5: coherence with the 1.0 A grid is 0.9995-1.0000 over the 15-4 A set on 1DAW, 2DQ6, 3K7M and 4BX9 (0.987 at tf_d_min itself), at 10-38 ms against 200-860. Measured with the pose gate (job 544884): 30/30 true poses within 0.32 A at the default, 18/30 with the window removed. Per-alignment wall clock 2.0-6.7 s against 5.8-54 s on shared nodes. The integration tests imported the deleted align module and passed removed keyword arguments, so they had not run since the refactor; they now use the package entry point and Cartesian mates, and the translation test checks the position. run_random_pdb_fit.py exercised only removed features and is gone. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- .../analysis/integration_alignment.sh | 17 + alignment_lab/analysis/p1_grid_coherence.sh | 19 ++ alignment_lab/analysis/pose_panel_trans.sh | 9 +- .../diagnostics/p1_grid_coherence.py | 71 +++++ docs/changelog.rst | 4 + tests/integration/alignment/profile_fit.py | 2 +- .../alignment/run_random_pdb_fit.py | 299 ------------------ .../integration/alignment/test_fit_to_data.py | 30 +- .../alignment/test_fit_to_data_translation.py | 70 ++-- torchref/experimental/alignment/pipeline.py | 63 ++-- 10 files changed, 226 insertions(+), 358 deletions(-) create mode 100644 alignment_lab/analysis/integration_alignment.sh create mode 100644 alignment_lab/analysis/p1_grid_coherence.sh create mode 100644 alignment_lab/diagnostics/p1_grid_coherence.py delete mode 100644 tests/integration/alignment/run_random_pdb_fit.py diff --git a/alignment_lab/analysis/integration_alignment.sh b/alignment_lab/analysis/integration_alignment.sh new file mode 100644 index 00000000..bee675da --- /dev/null +++ b/alignment_lab/analysis/integration_alignment.sh @@ -0,0 +1,17 @@ +#!/bin/bash +#SBATCH --job-name=integ +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:59:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +"$PY" -m pytest -c tests/pytest.ini tests/integration/alignment tests/unit/alignment --run-slow -q --tb=short 2>&1 | grep -v "^✓" | tail -60 +echo "PYTEST_RC=${PIPESTATUS[0]}" diff --git a/alignment_lab/analysis/p1_grid_coherence.sh b/alignment_lab/analysis/p1_grid_coherence.sh new file mode 100644 index 00000000..6d4116f2 --- /dev/null +++ b/alignment_lab/analysis/p1_grid_coherence.sh @@ -0,0 +1,19 @@ +#!/bin/bash +#SBATCH --job-name=p1coh +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:30:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +for P in 1DAW 2DQ6 3K7M 4BX9; do + "$PY" -u alignment_lab/diagnostics/p1_grid_coherence.py --pdb $P 2>/dev/null | grep -E "^#|^ROW" +done +echo DONE diff --git a/alignment_lab/analysis/pose_panel_trans.sh b/alignment_lab/analysis/pose_panel_trans.sh index 45940072..6df3cbe0 100644 --- a/alignment_lab/analysis/pose_panel_trans.sh +++ b/alignment_lab/analysis/pose_panel_trans.sh @@ -1,6 +1,7 @@ #!/bin/bash -# The 10 x 3 panel with the translation error reported, at the default -# translation window (all data) and at 15-4 A. The old gate was rotation-only. +# The 10 x 3 panel with the pose gate (rotation AND translation), at the +# pipeline's default translation window and with the window removed. The +# default is now the rotation search's own window; "full" is what it used to be. #SBATCH --job-name=ptrans #SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out #SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err @@ -20,8 +21,8 @@ export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" for T in 0 1 2; do "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial $T --arms llg \ - 2>/dev/null | grep '^ROW ' | sed 's/^ROW/ROW window=full/' + 2>/dev/null | grep '^ROW ' | sed 's/^ROW/ROW window=default/' "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial $T --arms llg \ - --tf-d-min 4.0 --tf-d-max 15.0 2>/dev/null | grep '^ROW ' | sed 's/^ROW/ROW window=15-4/' + --tf-d-min 0 --tf-d-max inf 2>/dev/null | grep '^ROW ' | sed 's/^ROW/ROW window=full/' done echo DONE diff --git a/alignment_lab/diagnostics/p1_grid_coherence.py b/alignment_lab/diagnostics/p1_grid_coherence.py new file mode 100644 index 00000000..8fa89d46 --- /dev/null +++ b/alignment_lab/diagnostics/p1_grid_coherence.py @@ -0,0 +1,71 @@ +"""How coarse can the P1 copy's FFT grid be for the translation set? + +The placement stage evaluates the P1 model's transform at the symmetry-rotated +indices of the translation set. Its grid was sized by the model's default +``max_res = 1.0 A`` whatever the set's resolution. This measures the complex +coherence of ``F_calc`` at the 15-4 A reflections between that grid and grids +sized to ``tf_d_min / oversampling``, and the time of each. +""" +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) + ap.add_argument("--tf-d-min", type=float, default=4.0) + ap.add_argument("--tf-d-max", type=float, default=15.0) + ap.add_argument("--oversampling", default="1.0,1.33,2.0") + args = ap.parse_args() + + model, data = load_case(args.pdb) + rec = data.cell.reciprocal_basis_matrix.to(torch.float64) + mask = data.get_valid_mask() + hkl = data.hkl[mask] + s = (hkl.to(torch.float64) @ rec).norm(dim=-1) + keep = (s >= 1.0 / args.tf_d_max) & (s <= 1.0 / args.tf_d_min) + hkl = hkl[keep] + sym_R = data.spacegroup.matrices.to(torch.float64) + hkl_SN = torch.einsum("ne,ied->ind", hkl.to(torch.float64), sym_R + ).reshape(-1, 3).round().to(torch.int64) + print(f"# {args.pdb} sg={data.spacegroup.hm} S={sym_R.shape[0]} N={hkl.shape[0]} " + f"window={args.tf_d_max}-{args.tf_d_min} A", flush=True) + + rot = model.copy() + rot = rot.rotate(random_rotation(seed_for(args.pdb, 0)).to(model.dtype_float)) + + ref = None + for over in [None] + [float(x) for x in args.oversampling.split(",")]: + m = rot.copy() + m.max_res = 1.0 if over is None else args.tf_d_min / over + m.spacegroup = "P 1" + with torch.no_grad(): + m.reset_cache(); m(hkl_SN) # warm + t0 = time.perf_counter() + for _ in range(3): + m.reset_cache(); sf = m(hkl_SN) + t = (time.perf_counter() - t0) / 3 + x = sf.to(torch.complex128) + if ref is None: + ref, coh, tag = x, 1.0, "1.0A" + else: + coh = float((ref.conj() * x).sum().abs() + / (ref.abs().norm() * x.abs().norm()).clamp(min=1e-30)) + tag = f"d_min/{over:g}" + print(f"ROW pdb={args.pdb} arm={tag} max_res={m.max_res:.3f} " + f"grid={tuple(int(v) for v in m.fft.gridsize)} t={t*1e3:.1f}ms " + f"coh={coh:.6f}", flush=True) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/docs/changelog.rst b/docs/changelog.rst index 21347d78..65209d29 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,10 @@ Changelog Unreleased ---------- +- The translation search defaults to the rotation search's resolution window instead of all data. With all data it placed the four largest panel structures (2DQ6, 3VRJ, 4BX9, 6G9X) at the right orientation and 20-56 A from the true position on every trial, and its own score was higher at the wrong place than at the deposited pose; the benchmark had only ever checked the rotation. 30/30 true poses within 0.32 A against 18/30, at roughly a fifth of the wall clock +- The P1 copy each rotation candidate is evaluated through is gridded at two thirds of the translation window's resolution rather than at the model's default 1.0 A. Coherence with the fine grid 0.9995 or better on all four structures measured; 10-38 ms per candidate against 200-860 +- The pose-recovery harness compares against Cartesian symmetry mates and checks the translation. It compared a Cartesian rotation against fractional symmetry matrices, so in trigonal and hexagonal cells two of six and four of twelve correct mates read as 30 and 21 degrees; every recorded 2DQ6 failure was one of them +- Repaired the alignment integration tests, which imported a deleted module and passed removed arguments - The alignment package uses the shared FFT-size and Euler-matrix helpers instead of its own copies, and drops four unused rotation utilities and two duplicated symmetry helpers. Bit-identical placements on the benchmark panel - The translation likelihood's model error comes from the shared ``SigmaAEstimator`` instead of a local 81-point scan over every reflection. It returns ``alpha`` and ``beta`` per reflection rather than a per-shell ``sigma_A``, so the likelihood no longer assumes ```` is exactly one. Outcome-neutral over 70 seeded cells, zero flips; not measurably faster - Fixed the translation likelihood's variance convention, which scored acentric reflections at twice the variance intended -- 90-95% of reflections. The alignment package carried its own Rice and Woolfson parameterised by the *amplitude* variance and handed both branches the same number, where the acentric branch needs half what the centric one does. It now uses ``base.targets.xray_likelihoods.rice_per_refl``, which takes the complex variance and derives the centric case from it diff --git a/tests/integration/alignment/profile_fit.py b/tests/integration/alignment/profile_fit.py index c3202f61..d4e0284c 100644 --- a/tests/integration/alignment/profile_fit.py +++ b/tests/integration/alignment/profile_fit.py @@ -22,7 +22,7 @@ import torch -from torchref.experimental.alignment.align import align_model_to_data +from torchref.experimental.alignment import align_model_to_data from torchref.experimental.alignment.frf.rotation_utils import rotation_matrix_from_edmonds_euler from torchref.io.datasets.reflection_data import ReflectionData from torchref.model import ModelFT diff --git a/tests/integration/alignment/run_random_pdb_fit.py b/tests/integration/alignment/run_random_pdb_fit.py deleted file mode 100644 index b4beb24d..00000000 --- a/tests/integration/alignment/run_random_pdb_fit.py +++ /dev/null @@ -1,299 +0,0 @@ -#!/usr/bin/env python -""" -End-to-end demo / sanity check for the alignment pipeline. - -Flow: - 1. Pick a random PDB / MTZ pair from `tests/files`. - 2. Load model + data (real cell, real F_obs). - 3. Apply a random rotation to the model atoms. - 4. Run `ModelFT.fit_to_data` to recover an aligned orientation. - 5. Fit an anisotropic Scaler against the data. - 6. Report R-work / R-free before and after. - -Run as a script (NOT a pytest test) — it's a one-off integration probe: - - cd /das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/fix_alignment - python tests/integration/alignment/run_random_pdb_fit.py [--seed N] [--pdb 1DAW] -""" - - -from __future__ import annotations - -import argparse -import math -import random -import time -from pathlib import Path - -import torch - -from torchref.base.metrics.rfactor import rfactor_work_free -from torchref.experimental.alignment.align import align_model_to_data -from torchref.experimental.alignment.frf.rotation_utils import rotation_angular_distance_deg -from torchref.io.datasets.reflection_data import ReflectionData -from torchref.model import ModelFT -from torchref.scaling import Scaler - - -TEST_FILES =Path('/das/work/p17/p17490/Peter/Library/work_trees_torchref/fix_alignment/tests/files') - -PAIRS = { - # PDB stem → (pdb_path, mtz_path). Some PDBs use a non-standard filename. - "1AK5": (TEST_FILES / "pdb" / "1AK5_with_H.pdb", TEST_FILES / "mtz" / "1AK5.mtz"), - "1DAW": (TEST_FILES / "pdb" / "1DAW.pdb", TEST_FILES / "mtz" / "1DAW.mtz"), - "2DQ6": (TEST_FILES / "pdb" / "2DQ6.pdb", TEST_FILES / "mtz" / "2DQ6.mtz"), - "3A5V": (TEST_FILES / "pdb" / "3A5V.pdb", TEST_FILES / "mtz" / "3A5V.mtz"), - "3E98": (TEST_FILES / "pdb" / "3E98.pdb", TEST_FILES / "mtz" / "3E98.mtz"), - "3GR5": (TEST_FILES / "pdb" / "3GR5.pdb", TEST_FILES / "mtz" / "3GR5.mtz"), - "3K7M": (TEST_FILES / "pdb" / "3K7M.pdb", TEST_FILES / "mtz" / "3K7M.mtz"), - "3VRJ": (TEST_FILES / "pdb" / "3VRJ.pdb", TEST_FILES / "mtz" / "3VRJ.mtz"), - "4BX9": (TEST_FILES / "pdb" / "4BX9.pdb", TEST_FILES / "mtz" / "4BX9.mtz"), - # 5BOV excluded — too large (4.6 GiB single TF allocation OOMs on A100-40GB). - "6G9X": (TEST_FILES / "pdb" / "6G9X.pdb", TEST_FILES / "mtz" / "6G9X.mtz"), -} - - -def _random_rotation(seed: int) -> torch.Tensor: - """Uniform random rotation on SO(3) via QR of a Gaussian matrix.""" - g = torch.Generator().manual_seed(int(seed)) - A = torch.randn(3, 3, generator=g, dtype=torch.float64) - Q, R = torch.linalg.qr(A) - Q = Q @ torch.diag(torch.sign(torch.diag(R))) - if torch.det(Q) < 0: - Q[:, 0] = -Q[:, 0] - return Q - - -def _min_err_over_sym(R_test: torch.Tensor, R_ref: torch.Tensor, - sym_mats: torch.Tensor) -> float: - """Minimum angular distance of R_test to any S·R_ref over sym_mats.""" - best = float("inf") - R_test = R_test.to(torch.float64) - R_ref = R_ref.to(torch.float64) - for k in range(sym_mats.shape[0]): - e = rotation_angular_distance_deg(R_test, sym_mats[k] @ R_ref) - if e < best: - best = e - return best - - -def _kabsch_rotation(xyz_a: torch.Tensor, xyz_b: torch.Tensor) -> torch.Tensor: - """Return R minimising ||xyz_a - xyz_b @ R^T|| (both centred).""" - a = (xyz_a.detach() - xyz_a.detach().mean(0)).to(torch.float64) - b = (xyz_b.detach() - xyz_b.detach().mean(0)).to(torch.float64) - H = b.T @ a - U, _, Vt = torch.linalg.svd(H) - d = float(torch.sign(torch.det(Vt.T @ U.T))) - D = torch.diag(torch.tensor([1.0, 1.0, d], dtype=H.dtype, device=H.device)) - return Vt.T @ D @ U.T - - -def run(pdb_key: str, seed: int, verbose: int = 1, - device: torch.device = torch.device("cpu"), - use_interp_var: bool = False, - use_llg_tf: bool = False, - refine_b: bool = False, - sigma_rot_deg: float = 0.0, - sigma_trans_ang: float = 0.0, - sigma_b: float = 0.0, - n_rotation_candidates: int = 15, - rescore_engine: str = "m_letf1") -> dict: - pdb_path, mtz_path = PAIRS[pdb_key] - print(f"\n=== {pdb_key}: {pdb_path.name} + {mtz_path.name} ===", flush=True) - - # 1. Construct model + data directly on `device` so the alignment - # pipeline runs end-to-end on GPU without re-creating the SfFFT. - # Post-hoc `.to(device)` reassigns Model.cell via the setter, which - # triggers `_maybe_initialize_fft` and drops the already-built grid / - # map_symmetry state. - t0 = time.time() - model = ModelFT(device=device).load_pdb(str(pdb_path)) - data = ReflectionData(device=str(device)).load_mtz(str(mtz_path)) - sym_mats = data.spacegroup.matrices.to(dtype=torch.float64, device=device) - print(f" spacegroup: {data.spacegroup} cell: " - f"a={data.cell.a:.1f}, b={data.cell.b:.1f}, c={data.cell.c:.1f} Å " - f"atoms: {model.xyz().shape[0]} ({time.time()-t0:.1f}s)", flush=True) - - def _scale_and_r(m: ModelFT) -> tuple[float, float]: - s = Scaler(model=m, data=data, nbins=20, verbose=0, - device=m.xyz().device) - # Detach fcalc — the scaler only needs grad through its own - # parameters. Without this, m's autograd graph (and the - # CachedForwardMixin hook that pins it via a register_hook closure) - # keeps multi-GB of fft intermediates alive across trials. - with torch.no_grad(): - fcalc = m(data.hkl).detach() - s.initialize(fcalc) - s.refine_lbfgs(fcalc=fcalc) - with torch.no_grad(): - rw, rf = rfactor_work_free(data, torch.abs(s.forward(fcalc))) - rw = rw.item() if hasattr(rw, "item") else float(rw) - rf = rf.item() if hasattr(rf, "item") else float(rf) - return rw, rf - - # Reference R-factor of the un-rotated model (the optimal we could hope for). - rwork_ref, rfree_ref = _scale_and_r(model) - print(f" reference R-work (un-rotated model): {rwork_ref:.4f} " - f"R-free: {rfree_ref:.4f}", flush=True) - - # 2. Apply random rotation to atom coords. - R_true = _random_rotation(seed) - xyz_canonical = model.xyz().clone() - centroid = xyz_canonical.mean(0) - rotated_search = model.rotate( - R_true.to(model.dtype_float).to(device), center=centroid, - ) - - # R-factor of the rotated search model (should be ~50% — random). - rwork_pre, rfree_pre = _scale_and_r(rotated_search) - print(f" rotated search R-work (no alignment): {rwork_pre:.4f} " - f"R-free: {rfree_pre:.4f} (should be ~0.5)", flush=True) - - # 3. Run fit_to_data: recover the alignment. - t1 = time.time() - aligned = align_model_to_data( - rotated_search, - data, - d_min=4.0, d_max=15.0, - n_shells=20, - n_rotation_peaks=200, n_ml_refine=200, - verbose=verbose, - use_interp_var=use_interp_var, - use_llg_tf=use_llg_tf, - refine_b=refine_b, - sigma_rot_deg=sigma_rot_deg, - sigma_trans_ang=sigma_trans_ang, - sigma_b=sigma_b, - n_rotation_candidates=n_rotation_candidates, - rescore_engine=rescore_engine, - ) - fit_time = time.time() - t1 - print(f" fit_to_data took {fit_time:.1f}s", flush=True) - - # Effective rotation between aligned coords and the canonical reference. - R_residual = _kabsch_rotation(aligned.xyz(), xyz_canonical) - err_to_canonical = min( - rotation_angular_distance_deg(R_residual.to(torch.float64), sym_mats[k]) - for k in range(sym_mats.shape[0]) - ) - print(f" aligned-vs-canonical angular distance " - f"(mod {data.spacegroup}-symmetry): {err_to_canonical:.2f}°", flush=True) - - # 4. Scale the aligned model and report R-factor. - rwork_post, rfree_post = _scale_and_r(aligned) - print(f" aligned R-work: {rwork_post:.4f} R-free: {rfree_post:.4f}", flush=True) - - return { - "pdb": pdb_key, - "spacegroup": str(data.spacegroup), - "ref_rwork": rwork_ref, - "ref_rfree": rfree_ref, - "pre_rwork": rwork_pre, - "pre_rfree": rfree_pre, - "post_rwork": rwork_post, - "post_rfree": rfree_post, - "err_canonical_deg": err_to_canonical, - "fit_time_s": fit_time, - } - - -def main(): - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--seed", type=int, default=None, - help="Random seed (rotation + PDB pick). Default: time-based.") - ap.add_argument("--pdb", default=None, choices=sorted(PAIRS.keys()), - help="PDB key to use. Default: random.") - ap.add_argument("--n-trials", type=int, default=1, - help="Number of random trials with different rotations / PDBs.") - ap.add_argument("--sweep", action="store_true", - help="Iterate over every PDB in PAIRS, n-trials per PDB.") - ap.add_argument("--verbose", type=int, default=0, - help="Verbosity passed to fit_to_data.") - ap.add_argument("--device", default="cpu", choices=["cpu", "cuda"], - help="Run the alignment on this device.") - ap.add_argument("--use-interp-var", action="store_true", - help="Phase A flag: add Phaser-style totvar_search " - "interpolation variance to the rotation rescore.") - ap.add_argument("--use-llg-tf", action="store_true", - help="Phase B flag: re-rank the FFT-correlation " - "translation peaks by Rice/Woolfson LLG.") - ap.add_argument("--refine-b", action="store_true", - help="Phase C: co-refine per-atom B-factors in rigid-body.") - ap.add_argument("--sigma-rot-deg", type=float, default=0.0, - help="Phase C: Gaussian rotation restraint sigma (deg). " - "0 = no restraint. Phaser default ~5.") - ap.add_argument("--sigma-trans-ang", type=float, default=0.0, - help="Phase C: Gaussian translation restraint sigma (Å). " - "0 = no restraint. Phaser default ~0.5.") - ap.add_argument("--sigma-b", type=float, default=0.0, - help="Phase C: Gaussian B-factor restraint sigma (Ų). " - "0 = no restraint. Phaser default ~15.") - ap.add_argument("--n-rotation-candidates", type=int, default=15, - help="Top-N rotations from MLRF rescore that get full " - "translation+polish. Default 15.") - args = ap.parse_args() - device = torch.device(args.device) - if device.type == "cuda" and not torch.cuda.is_available(): - raise SystemExit("--device cuda requested but torch.cuda not available") - - if args.seed is None: - args.seed = int(time.time()) - rng = random.Random(args.seed) - print(f"seed = {args.seed}", flush=True) - - # Build the (pdb, seed) work list. - if args.sweep: - worklist = [(pdb, rng.randint(0, 10 ** 9)) - for pdb in sorted(PAIRS.keys()) - for _ in range(args.n_trials)] - else: - worklist = [] - for _ in range(args.n_trials): - pdb = args.pdb if args.pdb is not None else rng.choice(list(PAIRS.keys())) - worklist.append((pdb, rng.randint(0, 10 ** 9))) - - results = [] - for pdb_key, trial_seed in worklist: - try: - r = run(pdb_key, trial_seed, verbose=args.verbose, device=device, - use_interp_var=args.use_interp_var, - use_llg_tf=args.use_llg_tf, - refine_b=args.refine_b, - sigma_rot_deg=args.sigma_rot_deg, - sigma_trans_ang=args.sigma_trans_ang, - sigma_b=args.sigma_b, - n_rotation_candidates=args.n_rotation_candidates) - results.append(r) - except Exception as exc: - import traceback - traceback.print_exc() - print(f" TRIAL FAILED on {pdb_key}: {exc!r}", flush=True) - results.append({"pdb": pdb_key, "error": repr(exc)}) - finally: - # Release CUDA allocator caches between trials. Without this a - # failed trial leaves its ~tens-of-GB residue in the allocator - # pool and starves every subsequent trial of memory. - if device.type == "cuda": - import gc - gc.collect() - torch.cuda.empty_cache() - alloc = torch.cuda.memory_allocated() / 1e9 - reserved = torch.cuda.memory_reserved() / 1e9 - print(f" [post-trial gc] alloc={alloc:5.2f} GB " - f"reserved={reserved:5.2f} GB", flush=True) - - print("\n=== summary ===", flush=True) - print(f"{'pdb':>6} {'sg':>8} {'rwork_ref':>10} {'rwork_pre':>10} " - f"{'rwork_post':>10} {'err_deg':>8} {'time_s':>7}", flush=True) - for r in results: - if "error" in r: - print(f"{r['pdb']:>6} FAILED: {r['error']}", flush=True) - continue - print(f"{r['pdb']:>6} {r['spacegroup'][12:20]:>8} " - f"{r['ref_rwork']:>10.4f} {r['pre_rwork']:>10.4f} " - f"{r['post_rwork']:>10.4f} {r['err_canonical_deg']:>8.2f} " - f"{r['fit_time_s']:>7.1f}", flush=True) - - -if __name__ == "__main__": - main() diff --git a/tests/integration/alignment/test_fit_to_data.py b/tests/integration/alignment/test_fit_to_data.py index 8121bd21..2df12120 100644 --- a/tests/integration/alignment/test_fit_to_data.py +++ b/tests/integration/alignment/test_fit_to_data.py @@ -1,23 +1,22 @@ """ -Integration test for `ModelFT.fit_to_data`: end-to-end Patterson rotation -search + Sim MLRF rescoring, returning a re-oriented ModelFT. +Integration test for ``align_model_to_data``: end-to-end rotation search +returning a re-oriented ModelFT. Setup: - F_obs: real 1DAW.mtz at its native C2 spacegroup. - Search model: P1 copy of 1DAW.pdb whose atomic coordinates have been rotated by a random R_true. -Acceptance: after `fit_to_data`, the returned model's atom coordinates are +Acceptance: after ``align_model_to_data``, the returned model's atom coordinates are within 8° rotation distance (modulo C2 symmetry of F_obs) of the un-rotated canonical orientation, for 5/5 random trials. """ -import math from pathlib import Path import pytest import torch -from torchref.experimental.alignment.align import align_model_to_data +from torchref.experimental.alignment import align_model_to_data from torchref.experimental.alignment.frf.rotation_utils import rotation_angular_distance_deg from torchref.io.datasets.reflection_data import ReflectionData from torchref.model import ModelFT @@ -67,17 +66,30 @@ def _best_alignment_rotation(xyz_a: torch.Tensor, xyz_b: torch.Tensor) -> torch. return R +def _cartesian_symops(data) -> torch.Tensor: + """Point-group rotations as Cartesian matrices, ``B S B^-1``. + + ``spacegroup.matrices`` act on fractional coordinates. A Kabsch rotation is + Cartesian, and comparing the two directly is only right when the cell is + orthogonal and the operator diagonal -- in a trigonal cell two of the six + mates of a correct placement read as 30 and 21 degrees. + """ + B = data.cell.fractional_matrix.to(torch.float64) + S = data.spacegroup.matrices.to(torch.float64) + return B @ S @ torch.linalg.inv(B) + + @pytest.mark.integration @pytest.mark.slow @pytest.mark.parametrize("trial", range(5)) def test_fit_to_data_real_1daw(real_setup, trial): """ - Apply a random R_true to a P1 search model, call `fit_to_data(real_F_obs)`, - and verify the returned model is within 8° rotation distance of the + Apply a random R_true to a P1 search model, align it to the real data, and + verify the returned model is within 8° rotation distance of the canonical orientation (modulo C2 symmetry of F_obs). """ data, make_model = real_setup - sym_mats = data.spacegroup.matrices.to(torch.float64) + sym_mats = _cartesian_symops(data) canonical = make_model() xyz_canonical = canonical.xyz().clone() @@ -91,7 +103,7 @@ def test_fit_to_data_real_1daw(real_setup, trial): data, d_min=4.0, d_max=15.0, n_shells=20, - n_rotation_peaks=200, n_ml_refine=200, + n_rotation_peaks=200, do_translation=False, # this test only checks rotation accuracy verbose=0, ) diff --git a/tests/integration/alignment/test_fit_to_data_translation.py b/tests/integration/alignment/test_fit_to_data_translation.py index 88a2e39e..3aa390f9 100644 --- a/tests/integration/alignment/test_fit_to_data_translation.py +++ b/tests/integration/alignment/test_fit_to_data_translation.py @@ -1,22 +1,20 @@ """ -Integration test for the new translation + joint R+t refinement in -`ModelFT.fit_to_data` (Phase 3 component B). - -Setup: 1DAW.mtz (C2) F_obs + P1 search model. Apply a small known rotation and -fractional translation, then ask `fit_to_data` to recover both. Acceptance: -recovered `(R_residual, t_residual)` brings the model close to canonical -(within 8° rotation modulo C2 symmetry and within 0.1 fractional shift along -any axis modulo unit cell), and the post-refinement R-work drops well below -the pre-fit value. +Integration test for ``align_model_to_data`` with the translation search on. + +Setup: 1DAW.mtz (C2) F_obs + P1 search model. Apply a known rotation and +fractional translation, then ask the pipeline to recover both. Acceptance: the +returned model is within 8 degrees of canonical modulo the C2 symmetry, its +position is within 3 A of a symmetry image of canonical (modulo lattice +translations and the group's allowed origin shifts; y is polar and free), and +its scaled R-work is close to the deposited model's. """ -import math from pathlib import Path import pytest import torch from torchref.base.metrics.rfactor import rfactor_work_free -from torchref.experimental.alignment.align import align_model_to_data +from torchref.experimental.alignment import align_model_to_data from torchref.experimental.alignment.frf.rotation_utils import ( rotation_angular_distance_deg, rotation_matrix_from_edmonds_euler, @@ -45,6 +43,19 @@ def _wrap_frac(t: torch.Tensor) -> torch.Tensor: return (t + 0.5) % 1.0 - 0.5 +def _cartesian_symops(data) -> torch.Tensor: + """Point-group rotations as Cartesian matrices, ``B S B^-1``. + + ``spacegroup.matrices`` act on fractional coordinates. A Kabsch rotation is + Cartesian, and comparing the two directly is only right when the cell is + orthogonal and the operator diagonal -- in a trigonal cell two of the six + mates of a correct placement read as 30 and 21 degrees. + """ + B = data.cell.fractional_matrix.to(torch.float64) + S = data.spacegroup.matrices.to(torch.float64) + return B @ S @ torch.linalg.inv(B) + + @pytest.mark.integration @pytest.mark.slow def test_fit_to_data_recovers_rotation_and_translation(): @@ -68,9 +79,8 @@ def test_fit_to_data_recovers_rotation_and_translation(): data, d_min=4.0, d_max=15.0, n_shells=20, - n_rotation_peaks=200, n_ml_refine=200, + n_rotation_peaks=200, do_translation=True, - do_joint_refine=True, verbose=0, ) @@ -86,20 +96,34 @@ def test_fit_to_data_recovers_rotation_and_translation(): d = float(torch.sign(torch.det(Vt.T @ U.T))) D = torch.diag(torch.tensor([1.0, 1.0, d], dtype=H.dtype)) R_residual = Vt.T @ D @ U.T - sym_mats = data.spacegroup.matrices.to(torch.float64) - best_rot_err = min( - rotation_angular_distance_deg(R_residual, sym_mats[k]) - for k in range(sym_mats.shape[0]) - ) + sym_cart = _cartesian_symops(data) + errs = [rotation_angular_distance_deg(R_residual, sym_cart[k]) + for k in range(sym_cart.shape[0])] + k_best = min(range(len(errs)), key=errs.__getitem__) + best_rot_err = errs[k_best] assert best_rot_err < 8.0, ( f"residual rotation {best_rot_err:.2f}° > 8° gate" ) - # We do not pin down the recovered translation directly — for spacegroups - # with polar / non-unique origins (e.g. C2's free origin along y) the - # recovered translation may differ from the applied one by an allowed - # origin shift. The crystallographic test that this is a valid solution is - # the R-factor of the scaled model. + # The translation, against the symmetry image whose rotation matched. In + # C2 the origin is free along y, the centring makes (1/2, 1/2, 0) a lattice + # vector, and (0, *, 1/2) is an allowed origin shift -- so x and z are each + # determined only modulo 1/2 and y not at all. A placement at the right + # orientation and 40 A from the true position used to pass this test. + B = data.cell.fractional_matrix.to(torch.float64) + Binv = torch.linalg.inv(B) + S = data.spacegroup.matrices.to(torch.float64) + T = data.spacegroup.translations.to(torch.float64) + c_a = Binv @ c_aligned + c_c = S[k_best] @ (Binv @ c_canon) + T[k_best] + delta = c_a - c_c + delta_xz = (delta + 0.25) % 0.5 - 0.25 + delta_xz[1] = 0.0 + trans_A = float((B @ delta_xz).norm()) + assert trans_A < 3.0, f"placed {trans_A:.1f} A from a symmetry image of canonical" + + # The crystallographic check that this is a valid solution is the R-factor + # of the scaled model. # The translation function brings R-work close to the canonical-native # reference (0.21 for 1DAW). The residual gap (~0.12) is from the # rotation function's ~2° angular error — a separate refinement that's diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index af57ec5f..a175de58 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -252,9 +252,10 @@ def __init__( # proxy got the ordering wrong, so the comparison has to be made end to # end. See the sort in `run` for what that measured. rank_by: str = "llg", - # Resolution window for the TRANSLATION set only, independent of the - # rotation search's [d_max, d_min]. None means no cut, which is what - # this stage has always done -- see `_prepare_translation_arrays`. + # Resolution window for the translation set. None means the rotation + # search's own [d_max, d_min], so one window and one normalisation + # serve both stages. Pass 0.0 / inf to remove a cut -- and see + # `_prepare_translation_arrays` for what the uncut set does. tf_d_min: Optional[float] = None, tf_d_max: Optional[float] = None, ): @@ -287,8 +288,8 @@ def __init__( raise ValueError( f"rank_by={rank_by!r}; expected 'r', 'corr' or 'llg'.") self.rank_by = rank_by - self.tf_d_min = tf_d_min - self.tf_d_max = tf_d_max + self.tf_d_min = float(d_min if tf_d_min is None else tf_d_min) + self.tf_d_max = float(d_max if tf_d_max is None else tf_d_max) self._timer = _StageTimer(enabled=verbose >= 2) # Filled in by run(). @@ -541,15 +542,20 @@ def _prepare_translation_arrays(self) -> None: """Mask the observations for the translation search and normalise them once. The window is ``[tf_d_max, tf_d_min]`` on top of the dataset's own - validity mask. Both default to ``None``, meaning **no resolution cut** -- - which is what this stage has always done, though it used to claim - otherwise. So the translation search sees the data's full resolution - while the rotation search runs at ``[d_max, d_min]`` = [15, 4] A. That - asymmetry is deliberate on one side and unexamined on the other: the - rotation function is bandwidth-limited and cannot use high-resolution - terms, and nobody has measured what the translation function wants. The - parameter exists so that choosing is possible; the default does not - choose. + validity mask, and by default it is the rotation search's ``[d_max, + d_min]``: one resolution window, one Wilson normalisation, both stages. + + The uncut set is not a safe default. With all data -- 228k reflections + to 1.5 A on 2DQ6 -- the translation search places the four largest panel + structures (2DQ6, 3VRJ, 4BX9, 6G9X) at the right orientation and 20-56 A + from the true position, on every trial, while its own score is HIGHER + at the wrong place than at the deposited pose (0.665 against 0.350 on + 2DQ6, where the likelihood is 1616 against 157865). At 15-4 A the same + search recovers all thirty poses to within 0.32 A. The objective's + calc side is raw ``|F_calc|^2``, so at high resolution it is dominated + by whatever reflections happen to carry the largest calculated + intensity rather than by the fit; the window is the first line of + defence and the normalisation of that objective is the second. """ data = self.data device = self.device @@ -561,14 +567,13 @@ def _prepare_translation_arrays(self) -> None: tmask = torch.ones( F_obs_full.shape[0], dtype=torch.bool, device=F_obs_full.device, ) - if self.tf_d_min is not None or self.tf_d_max is not None: - rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) - s_all = (hkl_full.to(torch.float64) @ rec_basis.to(hkl_full.device) - ).norm(dim=-1) - if self.tf_d_min is not None: - tmask = tmask & (s_all <= 1.0 / float(self.tf_d_min)) - if self.tf_d_max is not None: - tmask = tmask & (s_all >= 1.0 / float(self.tf_d_max)) + rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) + s_all = (hkl_full.to(torch.float64) @ rec_basis.to(hkl_full.device) + ).norm(dim=-1) + if self.tf_d_min > 0.0: + tmask = tmask & (s_all <= 1.0 / self.tf_d_min) + if np.isfinite(self.tf_d_max): + tmask = tmask & (s_all >= 1.0 / self.tf_d_max) self._tmask = tmask sig_F_full = getattr(data, "F_sigma", None) @@ -611,6 +616,20 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: # _modules and never runs the property setter -- a silent no-op. rotated_k.spacegroup = data.spacegroup.hm rotated_p1 = rotated_k.copy() + # Size the P1 copy's FFT grid to the translation set, not to the + # model's default 1.0 A: |s| is invariant under the symmetry rotations, + # so every rotated index the evaluator is asked for lies inside + # 1/tf_d_min. This is where the placement stage spent most of its time, + # on a grid 30-48x larger than the reflections it was asked for. + # + # Two thirds of the window's resolution, not the resolution itself. At + # max_res = tf_d_min the transform's coherence with the 1.0 A grid over + # the 15-4 A set is 0.987 on 2DQ6 (0.9987-0.9999 on 1DAW, 3K7M, 4BX9); + # at tf_d_min/1.5 it is 0.9995-1.0000 everywhere, at 10-38 ms against + # 200-860 ms. max_res first -- the space-group setter rebuilds the FFT + # and reads it. + if self.tf_d_min > 0.0: + rotated_p1.max_res = self.tf_d_min / 1.5 rotated_p1.spacegroup = "P 1" evaluator = DirectModelEvaluator(rotated_p1) From f123175118f23971c7b48483a29147e2c04b4d9c Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 2 Sep 2026 12:10:02 +0200 Subject: [PATCH 143/250] Score translations as a normalised covariance and pick them by likelihood The fast translation function divided raw |F_calc|^2 by its own sum. That is not a correlation: it is unbounded, and at high resolution it follows whichever reflections carry the largest calculated intensity. On the four largest panel structures it was higher 40 A from the true position than at the deposited pose (0.665 against 0.350 on 2DQ6), and the search went there on every trial. The observed side is now the rotation function's LERF1 coefficient, cw (E_obs^2 - 1) w sigma_A^2, and the calculated side is the candidate's transform divided by its own Wilson curve on the same abscissa, so every candidate's E_calc has unit mean per shell and the map is the covariance of two normalised intensities -- the rotation function's score equation, for translations. One FFT on a grid a third of the set's resolution apart, with parabolic peak refinement, replaces the 16-point coarse grid and the three 100-point local refines whose coarse half could miss the peak. The Rice/Woolfson likelihood at a fixed Luzzati sigma_A picks among the top peaks and ranks the candidates, so no candidate is scored against a model error fitted to itself and the per-candidate SigmaAEstimator fit is gone. The stage runs in the configured float and complex dtypes. use_llg_tf, n_translation_peaks and translation_grid_steps are removed; the lab diagnostics built on the old API are deleted. Measured with the pose gate (job 544899): 30/30 true poses at the default window and 30/30 with the window removed, maximum translation error 0.21 A, against 18/30 before. Warm on an exclusive EPYC 9335 node, 8 threads (job 544917): 1DAW 1.1 s, 2DQ6 1.9 s, 6G9X 2.2 s, 3K7M 2.7 s, 4BX9 3.8 s per alignment. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/dq6_discrimination.sh | 26 - ...a_a_cost.sh => empirical_sigma_a_check.sh} | 16 +- alignment_lab/analysis/ftf_disc_smoke.sh | 18 - alignment_lab/analysis/ftf_discrimination.sh | 25 - ...bessel_arg_range.sh => pipeline_timing.sh} | 8 +- alignment_lab/analysis/tf_cost.sh | 19 - alignment_lab/diagnostics/bessel_arg_range.py | 61 -- alignment_lab/diagnostics/calc_norm_cost.py | 146 --- .../diagnostics/empirical_sigma_a_check.py | 50 + .../diagnostics/frf_vs_ftf_discrimination.py | 230 ---- alignment_lab/diagnostics/pipeline_timing.py | 58 ++ alignment_lab/diagnostics/pose_recovery.py | 18 +- alignment_lab/diagnostics/sigma_a_cost.py | 75 -- alignment_lab/diagnostics/tf_cost.py | 108 -- .../diagnostics/truth_pose_scores.py | 31 +- docs/changelog.rst | 4 + tests/integration/alignment/profile_fit.py | 4 - .../alignment/test_patterson_translation.py | 175 ++-- tests/unit/alignment/test_translation_obs.py | 43 +- torchref/experimental/alignment/__init__.py | 22 +- torchref/experimental/alignment/pipeline.py | 210 +--- .../experimental/alignment/translation.py | 982 ++++++------------ 22 files changed, 601 insertions(+), 1728 deletions(-) delete mode 100644 alignment_lab/analysis/dq6_discrimination.sh rename alignment_lab/analysis/{sigma_a_cost.sh => empirical_sigma_a_check.sh} (62%) delete mode 100644 alignment_lab/analysis/ftf_disc_smoke.sh delete mode 100644 alignment_lab/analysis/ftf_discrimination.sh rename alignment_lab/analysis/{bessel_arg_range.sh => pipeline_timing.sh} (78%) delete mode 100644 alignment_lab/analysis/tf_cost.sh delete mode 100644 alignment_lab/diagnostics/bessel_arg_range.py delete mode 100644 alignment_lab/diagnostics/calc_norm_cost.py create mode 100644 alignment_lab/diagnostics/empirical_sigma_a_check.py delete mode 100644 alignment_lab/diagnostics/frf_vs_ftf_discrimination.py create mode 100644 alignment_lab/diagnostics/pipeline_timing.py delete mode 100644 alignment_lab/diagnostics/sigma_a_cost.py delete mode 100644 alignment_lab/diagnostics/tf_cost.py diff --git a/alignment_lab/analysis/dq6_discrimination.sh b/alignment_lab/analysis/dq6_discrimination.sh deleted file mode 100644 index 29a8c2be..00000000 --- a/alignment_lab/analysis/dq6_discrimination.sh +++ /dev/null @@ -1,26 +0,0 @@ -#!/bin/bash -# On the seeds where 2DQ6 fails, is truth in the candidate list at all? -# -# The structure passes 6/10 end to end and bimodally -- 0-4 deg or 21-31. That -# is either the rotation function never producing the true orientation, or the -# translation function producing it and ranking the tNCS alternative above it. -# Those are different problems and the panel cannot tell them apart, because it -# only reports the winner. This reports truth's RANK under each score. -#SBATCH --job-name=dq6disc -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=day -#SBATCH --time=04:00:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -# 15 candidates: what the pipeline actually carries into the translation search. -"$PY" -u alignment_lab/diagnostics/frf_vs_ftf_discrimination.py \ - --pdb 2DQ6 --trials 10 --n-cand 15 2>/dev/null | grep -E '^(ROW|CAND)' -echo DONE diff --git a/alignment_lab/analysis/sigma_a_cost.sh b/alignment_lab/analysis/empirical_sigma_a_check.sh similarity index 62% rename from alignment_lab/analysis/sigma_a_cost.sh rename to alignment_lab/analysis/empirical_sigma_a_check.sh index f97f3a84..edf9e956 100644 --- a/alignment_lab/analysis/sigma_a_cost.sh +++ b/alignment_lab/analysis/empirical_sigma_a_check.sh @@ -1,17 +1,19 @@ #!/bin/bash -#SBATCH --job-name=sacost +#SBATCH --job-name=esa #SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out #SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err #SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G +#SBATCH --time=00:20:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=32G #SBATCH --constraint=cpu_epyc9335 set -uo pipefail REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -"$PY" -u alignment_lab/diagnostics/sigma_a_cost.py 1DAW 4BX9 2DQ6 2>/dev/null | grep '^ROW' +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +for P in 1DAW 2DQ6 3K7M; do + "$PY" -u alignment_lab/diagnostics/empirical_sigma_a_check.py --pdb $P 2>&1 | grep -v "Warning\|warnings.warn" | grep -A14 "^ROW\|Traceback" +done echo DONE diff --git a/alignment_lab/analysis/ftf_disc_smoke.sh b/alignment_lab/analysis/ftf_disc_smoke.sh deleted file mode 100644 index b0bda334..00000000 --- a/alignment_lab/analysis/ftf_disc_smoke.sh +++ /dev/null @@ -1,18 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=ftfsmoke -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:30:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=32G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -export TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 -cd "$REPO" -"$PY" -u alignment_lab/diagnostics/frf_vs_ftf_discrimination.py \ - --pdb 1DAW --trials 1 --n-cand 4 --n-rotation-peaks 60 2>&1 | tail -30 -echo "RC=${PIPESTATUS[0]}" diff --git a/alignment_lab/analysis/ftf_discrimination.sh b/alignment_lab/analysis/ftf_discrimination.sh deleted file mode 100644 index d851c5ce..00000000 --- a/alignment_lab/analysis/ftf_discrimination.sh +++ /dev/null @@ -1,25 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=ftfdisc -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-9 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -export TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 -# Concurrent array tasks otherwise poison a shared __pycache__ for numba. -export NUMBA_CACHE_DIR="/tmp/numba_${SLURM_ARRAY_JOB_ID}_${SLURM_ARRAY_TASK_ID}" -mkdir -p "$NUMBA_CACHE_DIR" -cd "$REPO" -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -P=${PDBS[$SLURM_ARRAY_TASK_ID]} -"$PY" -u alignment_lab/diagnostics/frf_vs_ftf_discrimination.py \ - --pdb "$P" --trials 3 --n-cand 25 --n-rotation-peaks 200 \ - --out-csv "alignment_lab/runs/ftf_disc_${SLURM_ARRAY_JOB_ID}.csv" 2>&1 -echo "RC=${PIPESTATUS[0]} pdb=$P" diff --git a/alignment_lab/analysis/bessel_arg_range.sh b/alignment_lab/analysis/pipeline_timing.sh similarity index 78% rename from alignment_lab/analysis/bessel_arg_range.sh rename to alignment_lab/analysis/pipeline_timing.sh index e891635f..842dcb30 100644 --- a/alignment_lab/analysis/bessel_arg_range.sh +++ b/alignment_lab/analysis/pipeline_timing.sh @@ -1,11 +1,12 @@ #!/bin/bash -#SBATCH --job-name=bessel +#SBATCH --job-name=ptime #SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out #SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err #SBATCH --partition=hour -#SBATCH --time=00:55:00 +#SBATCH --time=00:59:00 #SBATCH --cpus-per-task=8 #SBATCH --mem=64G +#SBATCH --exclusive #SBATCH --constraint=cpu_epyc9335 set -uo pipefail REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement @@ -13,5 +14,6 @@ PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin cd "$REPO" export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -"$PY" -u alignment_lab/diagnostics/bessel_arg_range.py 2>/dev/null | grep '^ROW' +"$PY" -u alignment_lab/diagnostics/pipeline_timing.py --threads 8 2>/dev/null \ + | grep -E "^ROW|^stage|^[0-9]_|^TOTAL|^---" echo DONE diff --git a/alignment_lab/analysis/tf_cost.sh b/alignment_lab/analysis/tf_cost.sh deleted file mode 100644 index 1bf6e66b..00000000 --- a/alignment_lab/analysis/tf_cost.sh +++ /dev/null @@ -1,19 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=tfcost -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=01:00:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -export TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 -cd "$REPO" -for P in 1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X; do - "$PY" -u alignment_lab/diagnostics/tf_cost.py --pdb "$P" --trial 0 \ - 2>&1 | grep -E "^#|^ROW|Error|Traceback" || echo "ROW pdb=$P FAILED" -done diff --git a/alignment_lab/diagnostics/bessel_arg_range.py b/alignment_lab/diagnostics/bessel_arg_range.py deleted file mode 100644 index 009e5d4c..00000000 --- a/alignment_lab/diagnostics/bessel_arg_range.py +++ /dev/null @@ -1,61 +0,0 @@ -"""Does the base Rice's Bessel-argument clamp bite in E-space? - -`xray_likelihoods._rice_body` clamps `2 F_calc F_obs / Sigma` at 1e6. That cap is -sized for F-space; the translation likelihood runs on E values with -`Sigma = 1 - D^2` floored at 1e-4, where the ratio can in principle be far -larger. A clamp that fires silently truncates the likelihood exactly where it is -most discriminating, so this measures the real range instead of reasoning about -the bound. -""" -import sys -from pathlib import Path -import torch -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) -from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 - - -def main(): - from torchref.experimental.alignment.translation import ( - DirectModelEvaluator, TranslationObs, amplitude_translation_search, - normalise_calc, precompute_G_for_rotation) - - for pdb in sys.argv[1:] or list(BENCH_PDBS): - model, data = load_case(pdb) - rot = model.copy().rotate( - random_rotation(seed_for(pdb, 0)).to(model.dtype_float), - center=model.xyz().mean(0)) - rot.spacegroup = data.spacegroup.hm - p1 = rot.copy(); p1.spacegroup = "P 1" - mask = data.get_valid_mask() - sig = getattr(data, "F_sigma", None) - obs = TranslationObs.build(data.F[mask], data.hkl[mask], data.spacegroup, - data.cell, - sig_F=None if sig is None else sig[mask], - n_shells=10) - ev = DirectModelEvaluator(p1) - eye3 = torch.eye(3, dtype=torch.float64) - G, h_R = precompute_G_for_rotation(ev, eye3, obs.hkl, data.spacegroup, - data.cell) - _, _, peaks = amplitude_translation_search( - obs=obs, interpolator=ev, R_rotation=eye3, - spacegroup=data.spacegroup, real_cell=data.cell, grid_steps=16, - n_peaks=1, precomputed_G=G, precomputed_h_R=h_R) - t = torch.as_tensor(peaks[0].translation, dtype=torch.float64, - device=G.device) - ph = torch.exp(2j * torch.pi * torch.einsum( - "ind,d->in", h_R.to(torch.float64), t).to(G.dtype)) - E_calc = normalise_calc((G * ph).sum(dim=0).abs().to(torch.float64), obs) - - # Worst case over the whole sigma_A grid: D -> 0.99, Sigma -> 1 - D^2. - D = 0.99 - Sigma = max(1.0 - D * D, 1e-4) - arg = (2.0 * (D * E_calc) * obs.E_obs / Sigma) - print(f"ROW pdb={pdb} N={obs.E_obs.numel()} " - f"maxE_obs={float(obs.E_obs.max()):.2f} " - f"maxE_calc={float(E_calc.max()):.2f} " - f"max_bessel_arg={float(arg.max()):.3e} " - f"clamped={int((arg > 1e6).sum())}", flush=True) - - -main() diff --git a/alignment_lab/diagnostics/calc_norm_cost.py b/alignment_lab/diagnostics/calc_norm_cost.py deleted file mode 100644 index ac283745..00000000 --- a/alignment_lab/diagnostics/calc_norm_cost.py +++ /dev/null @@ -1,146 +0,0 @@ -"""What does it cost to normalise the LLG's calculated side with the shared fit? - -Two per-shell normalisations survive in the translation likelihood -- the -``E_calc`` of each candidate translation, and the ``E_calc`` of the top peak that -the sigma_A fit runs against. They are the last places in the alignment package -that answer "what is the mean intensity here" without going through -:class:`~torchref.scaling.WilsonNormaliser`. - -The argument for keeping them was cost: the shared fit is a Gamma GLM by IRLS and -the calc side needs one fit per candidate, K of them per rotation. This measures -that instead of asserting it, and also asks the two questions that decide whether -the swap is safe at all: - -* does the fit **converge** on a calculated set, which has near-zeros at the - nodes of the molecular transform where an observed set has none, and -* how far do the two normalisations actually differ, per reflection and in the - LLG ranking they feed. - -Warm-up first: on this filesystem a first call pays cold package reads inside -whatever timer surrounds it, which is worth ~100x and is not compute. - -Usage:: - - python alignment_lab/diagnostics/calc_norm_cost.py --pdb 1DAW --k 20 -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 - - -def _time(fn, repeats=3): - fn() # warm: discard the cold-read call - ts = [] - for _ in range(repeats): - t0 = time.perf_counter() - out = fn() - ts.append(time.perf_counter() - t0) - return min(ts), out - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) - ap.add_argument("--trial", type=int, default=0) - ap.add_argument("--k", type=int, default=20, help="candidate translations") - ap.add_argument("--n-coeff", type=int, default=6) - args = ap.parse_args() - - from torchref.experimental.alignment.translation import ( - DirectModelEvaluator, TranslationObs, amplitude_translation_search, - precompute_G_for_rotation, - ) - from torchref.scaling import WilsonNormaliser - - seed = seed_for(args.pdb, args.trial) - model, data = load_case(args.pdb) - R_true = random_rotation(seed) - rot = model.copy() - rot = rot.rotate(R_true.to(model.dtype_float), center=model.xyz().mean(0)) - rot.spacegroup = data.spacegroup.hm - p1 = rot.copy() - p1.spacegroup = "P 1" - - mask = data.get_valid_mask() - sig = getattr(data, "F_sigma", None) - obs = TranslationObs.build( - data.F[mask], data.hkl[mask], data.spacegroup, data.cell, - sig_F=None if sig is None else sig[mask], - ) - ev = DirectModelEvaluator(p1) - eye3 = torch.eye(3, dtype=torch.float64) - G, h_R = precompute_G_for_rotation( - ev, eye3, obs.hkl, data.spacegroup, data.cell) - _, _, peaks = amplitude_translation_search( - obs=obs, interpolator=ev, R_rotation=eye3, - spacegroup=data.spacegroup, real_cell=data.cell, - grid_steps=16, n_peaks=args.k, precomputed_G=G, precomputed_h_R=h_R) - - K = min(args.k, len(peaks)) - N = obs.hkl.shape[0] - t_cand = torch.as_tensor( - [p.translation for p in peaks[:K]], dtype=torch.float64, - device=G.device) - phase = torch.exp(2j * torch.pi * torch.einsum( - "ind,kd->kin", h_R.to(torch.float64), t_cand).to(G.dtype)) - F_calc = (G.view(1, *G.shape) * phase).sum(dim=1).abs().to(torch.float64) - - print(f"# {args.pdb} trial={args.trial} N={N} K={K} " - f"n_coeff={args.n_coeff}", flush=True) - - # --- current: one per-shell mean per candidate, all K at once --- - def per_shell(): - idx = obs.shell_idx.view(1, -1).expand(K, N) - cnt = torch.bincount(obs.shell_idx, minlength=obs.n_shells).to(torch.float64) - tot = torch.zeros((K, obs.n_shells), dtype=torch.float64, device=G.device) - tot.scatter_add_(1, idx, F_calc * F_calc) - mean = (tot / cnt.clamp(min=1.0).unsqueeze(0)).clamp(min=1e-30) - return F_calc / mean.sqrt().gather(1, idx) - - # --- proposed: the shared Wilson fit, once per candidate --- - def wilson(): - out = torch.empty_like(F_calc) - iters = [] - for k in range(K): - w = WilsonNormaliser( - F_calc[k] * F_calc[k], obs.s_mag, n_coeff=args.n_coeff, - s_lo=float(obs.s_mag.min()), s_hi=float(obs.s_mag.max()), - ) - out[k] = w.E.to(torch.float64) - iters.append(w.n_iter) - return out, iters - - t_shell, E_shell = _time(per_shell) - t_wilson, (E_wilson, iters) = _time(wilson) - - # How different are they, and does the *ranking* they feed move? - rel = ((E_wilson - E_shell).abs() - / E_shell.abs().clamp(min=1e-12)).median().item() - m_shell = (E_shell ** 2).mean(dim=1) - m_wilson = (E_wilson ** 2).mean(dim=1) - - print(f"ROW pdb={args.pdb} N={N} K={K} " - f"t_per_shell_ms={1000 * t_shell:.2f} " - f"t_wilson_ms={1000 * t_wilson:.1f} " - f"ratio={t_wilson / max(t_shell, 1e-9):.0f}x " - f"per_cand_ms={1000 * t_wilson / K:.1f} " - f"iter_min={min(iters)} iter_max={max(iters)} " - f"median_rel_dE={rel:.4f} " - f"meanE2_shell={m_shell.mean():.4f} " - f"meanE2_wilson={m_wilson.mean():.4f}", flush=True) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/empirical_sigma_a_check.py b/alignment_lab/diagnostics/empirical_sigma_a_check.py new file mode 100644 index 00000000..f824c3a9 --- /dev/null +++ b/alignment_lab/diagnostics/empirical_sigma_a_check.py @@ -0,0 +1,50 @@ +"""What does the rotation function's "empirical" sigma_A actually evaluate to? + +``empirical_sigma_a`` divides the observed Wilson curve by the calculated one +and takes ``sqrt(min(R, 1/R))``. The two curves sit on different absolute +scales -- the MTZ's arbitrary one and the model's electron scale -- and the +function does not remove that factor, so the ratio's level, not only its shape, +sets the answer. This records the ratio and the resulting sigma_A by resolution +during a real rotation search. +""" +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from lab import BENCH_PDBS, load_case # noqa: E402 + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) + args = ap.parse_args() + + import torchref.experimental.alignment.frf.api as api + from torchref.experimental.alignment import rotation_search + + seen = {} + real = api.empirical_sigma_a + + def spy(sigma_obs, sigma_calc, **kw): + out = real(sigma_obs, sigma_calc, **kw) + seen["ratio"] = (sigma_obs / sigma_calc).detach().cpu() + seen["sigma_a"] = out.detach().cpu() + return out + + api.empirical_sigma_a = spy + model, data = load_case(args.pdb) + rotation_search(model, data, model_error_A=0.8, n_peaks=5) + r, sa = seen["ratio"], seen["sigma_a"] + q = lambda x: [round(float(v), 4) for v in torch.quantile(x, torch.tensor([0.0, 0.25, 0.5, 0.75, 1.0], dtype=x.dtype))] + print(f"ROW pdb={args.pdb} ratio_quantiles={q(r)} sigma_a_quantiles={q(sa)} " + f"ratio_geomean={float(r.log().mean().exp()):.4g}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_vs_ftf_discrimination.py b/alignment_lab/diagnostics/frf_vs_ftf_discrimination.py deleted file mode 100644 index d8203aa5..00000000 --- a/alignment_lab/diagnostics/frf_vs_ftf_discrimination.py +++ /dev/null @@ -1,230 +0,0 @@ -"""Does the translation function rank FRF peaks better than the FRF's own score? - -The FRF is a 3-D check: it correlates Pattersons, so a wrong orientation that -happens to reproduce the intramolecular vector set scores well. The translation -function is a 6-D check -- it has to place the molecule against the *crystal*, -intermolecular contacts included -- so it should separate truth from a ghost by -much more. That is the standard argument for carrying many orientations into the -TF, and it is worth measuring before spending anything on making the TF fast: -if the TF ranks no better than the FRF, carrying 100 orientations buys nothing. - -Takes the top ``--n-cand`` raw FRF peaks (no ML rescore -- the rescore is a -separate, and separately measured, reordering), places each one, and reports -where truth lands under four orderings: - -``frf`` the FRF's own score, i.e. the baseline -``tf_corr`` top translation peak of the Crowther-Blow amplitude correlation -``tf_llg`` the same peaks re-ranked by the shared-sigma_A Rice/Woolfson LLG -``r`` analytic-scale R at the locally refined t -- what the pipeline - actually ranks by today - -Rank alone understates the question, so each ordering also gets a separation -``z = (score_truth - mean_others) / std_others``: rank 0 by a hair and rank 0 by -five sigma are different claims about discrimination, and only the second one -justifies widening the funnel. - -The placement path is the production one -- ``_make_rotated`` then the same -``precompute_G`` / ``amplitude_translation_search`` / ``local_translation_refine`` -calls ``_placement_for_candidate`` makes, against the pipeline's own -``TranslationObs`` so the normalisation and weighting are the production ones. -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import numpy as np -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, ResultWriter, rotated_case, seed_for, # noqa: E402 - symmetry_orbit) -from lab.truth import angle_to_orbit # noqa: E402 - - -def _rank_of_truth(scores, truth_mask, higher_is_better=True): - """Rank of the best-placed *correct* candidate under this ordering. - - Correct means within the angular threshold, and several candidates can be: - the peak list carries near-duplicates and symmetry mates. Any of them is a - solution, so the rank that matters is the first one to appear -- the same - definition :func:`orbit_rank` uses, and the one the pipeline behaves by. - """ - s = np.asarray(scores, dtype=float) - order = np.argsort(-s if higher_is_better else s, kind="stable") - return int(next(i for i, j in enumerate(order) if truth_mask[j])) - - -def _separation(scores, truth_mask, higher_is_better=True): - """Best correct candidate's score in sigmas above the *wrong* ones. - - Negated for lower-is-better scores so a larger number always means better - discrimination, whichever direction the score runs. Every within-threshold - candidate is held out of the reference pool: leaving a second copy of the - answer in it would inflate the pool's mean and understate the separation. - """ - s = np.asarray(scores, dtype=float) - m = np.asarray(truth_mask, dtype=bool) - best = float(s[m].max() if higher_is_better else s[m].min()) - others = s[~m] - if others.size < 2: - return float("nan") - sd = float(others.std(ddof=1)) - if sd < 1e-30: - return float("nan") - z = (best - float(others.mean())) / sd - return z if higher_is_better else -z - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) - ap.add_argument("--trials", type=int, default=3) - ap.add_argument("--n-cand", type=int, default=25, - help="FRF peaks placed. Truth's worst FRF rank over the " - "100-cell panel was 21, so 25 covers every case that " - "the rotation function gets right at all.") - ap.add_argument("--n-rotation-peaks", type=int, default=200) - ap.add_argument("--thr-deg", type=float, default=8.0) - ap.add_argument("--out-csv", default=None) - args = ap.parse_args() - - from torchref.experimental.alignment.rotation_search import ( - prepare_frf_inputs, - ) - from torchref.experimental.alignment.frf.rotation_utils import ( - rotation_matrix_from_edmonds_euler, - ) - from torchref.experimental.alignment.pipeline import ( - MolecularReplacementPipeline, - ) - from torchref.experimental.alignment.translation import ( - DirectModelEvaluator, amplitude_translation_search, - local_translation_refine, precompute_G_for_rotation, - ) - - writer = None - if args.out_csv: - writer = ResultWriter( - args.out_csv, "frf_vs_ftf", - extra_fields=("n_cand", "truth_found", "rank_frf", "rank_tf_corr", - "rank_tf_llg", "rank_r", "z_frf", "z_tf_corr", - "z_tf_llg", "z_r", "seconds"), - ) - - for trial in range(args.trials): - seed = seed_for(args.pdb, trial) - model, data, R_true = rotated_case(args.pdb, seed) - t0 = time.time() - - pipe = MolecularReplacementPipeline( - data, model, verbose=0, - n_rotation_peaks=args.n_rotation_peaks, - n_rotation_candidates=args.n_cand, - use_llg_tf=False, - ) - frf = prepare_frf_inputs( - model, data, d_min=pipe.d_min, d_max=pipe.d_max, - n_shells=pipe.n_shells, verbose=0, - ) - pipe._frf = frf - peaks = pipe._rotation_candidates(frf)[: args.n_cand] - pipe._prepare_translation_arrays() - - orbit = symmetry_orbit( - R_true, data.spacegroup.matrices.to(torch.float64).cpu(), - side="left", frame="cart", - reciprocal_basis=data.cell.reciprocal_basis_matrix.to( - torch.float64).cpu(), - ) - eye3 = pipe._eye3 - rows = [] - for k, p in enumerate(peaks): - ang = angle_to_orbit( - rotation_matrix_from_edmonds_euler(p.alpha, p.beta, p.gamma), - orbit, - ) - rot = pipe._make_rotated(p)[0] - rot.spacegroup = data.spacegroup.hm - p1 = rot.copy() - p1.spacegroup = "P 1" - ev = DirectModelEvaluator(p1) - G, h_R = precompute_G_for_rotation( - ev, eye3, pipe._obs.hkl, data.spacegroup, data.cell) - _, _, tp = amplitude_translation_search( - obs=pipe._obs, interpolator=ev, R_rotation=eye3, - spacegroup=data.spacegroup, - real_cell=data.cell, grid_steps=pipe.translation_grid_steps, - n_peaks=pipe.n_translation_peaks, cluster_radius=0.05, - precomputed_G=G, precomputed_h_R=h_R) - tf_corr = float(tp[0].score) - # Phaser's TFZ: the top translation peak measured against the - # spread of *this orientation's own* translation map. Unlike the - # cross-candidate z reported below it needs no other candidate, so - # it is the only score here that can stop a sequential search. - tfz = float(tp[0].sigma) - llg_peaks = pipe._llg_tf_rescore(tp, G, h_R) - tf_llg = float(llg_peaks[0].score) - lv = np.array([q.score for q in llg_peaks], dtype=float) - llgz = (float((lv[0] - lv[1:].mean()) / lv[1:].std(ddof=1)) - if lv.size > 2 and lv[1:].std(ddof=1) > 1e-30 - else float("nan")) - r_best = float("inf") - for cand in llg_peaks[: pipe.n_translation_candidates]: - _, r_a = local_translation_refine( - obs=pipe._obs, interpolator=ev, R_rotation=eye3, - spacegroup=data.spacegroup, real_cell=data.cell, - t_init=torch.as_tensor(cand.translation, - dtype=torch.float64), - radius=0.06, grid_steps=13, n_refinement_passes=1, - precomputed_G=G, precomputed_h_R=h_R) - r_best = min(r_best, r_a) - rows.append(dict(k=k, ang=ang, frf=float(p.score), - tf_corr=tf_corr, tf_llg=tf_llg, r=r_best, - tfz=tfz, llgz=llgz)) - print(f"CAND pdb={args.pdb} trial={trial} k={k} ang={ang:.3f} " - f"is_truth={int(ang <= args.thr_deg)} frf={p.score:.4f} " - f"tf_corr={tf_corr:.5f} tfz={tfz:.3f} tf_llg={tf_llg:.2f} " - f"llgz={llgz:.3f} r={r_best:.5f}", flush=True) - - secs = time.time() - t0 - truth = [r for r in rows if r["ang"] <= args.thr_deg] - if not truth: - print(f"ROW pdb={args.pdb} trial={trial} truth_found=0 " - f"best_ang={min(r['ang'] for r in rows):.2f} " - f"seconds={secs:.1f}", flush=True) - continue - tmask = [r["ang"] <= args.thr_deg for r in rows] - ti = rows.index(min(truth, key=lambda r: r["ang"])) - cols = {"frf": True, "tf_corr": True, "tf_llg": True, "r": False} - ranks = {c: _rank_of_truth([r[c] for r in rows], tmask, hi) - for c, hi in cols.items()} - zs = {c: _separation([r[c] for r in rows], tmask, hi) - for c, hi in cols.items()} - print(f"ROW pdb={args.pdb} trial={trial} seed={seed} truth_found=1 " - f"n_cand={len(rows)} n_truth={sum(tmask)} " - + " ".join(f"rank_{c}={ranks[c]}" for c in cols) - + " " + " ".join(f"z_{c}={zs[c]:.2f}" for c in cols) - + f" seconds={secs:.1f}", flush=True) - if writer: - writer.write( - pdb=args.pdb, seed=seed, trial=trial, - spacegroup=str(data.spacegroup), - n_ops=int(data.spacegroup.matrices.shape[0]), - truth_rank=ranks["frf"], truth_angle_deg=round(rows[ti]["ang"], 3), - orbit_side="left", orbit_frame="cart", lmax_cap="", - d_min=pipe.d_min, d_max=pipe.d_max, device="cpu", - n_cand=len(rows), truth_found=1, - **{f"rank_{c}": ranks[c] for c in cols}, - **{f"z_{c}": round(zs[c], 4) for c in cols}, - seconds=round(secs, 1)) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/pipeline_timing.py b/alignment_lab/diagnostics/pipeline_timing.py new file mode 100644 index 00000000..c51952ef --- /dev/null +++ b/alignment_lab/diagnostics/pipeline_timing.py @@ -0,0 +1,58 @@ +"""Warm, single-process wall clock of the placement pipeline, by stage. + +One process, each structure aligned twice, the first pass discarded: the first +call in a process pays kernel builds and cold file-system reads that are not +compute (161 s against 1.5 s has been measured for one stage). Prints the +pipeline's own stage table for the second pass and the pose error, so a timing +is never quoted for a run that did not place the model. +""" +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from lab import (BENCH_PDBS, load_case, pose_error, random_rotation, # noqa: E402 + seed_for) + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdbs", default="1DAW,2DQ6,6G9X,3K7M,4BX9") + ap.add_argument("--threads", type=int, default=8) + args = ap.parse_args() + torch.set_num_threads(args.threads) + + from torchref.experimental.alignment import MolecularReplacementPipeline + + for pdb in [p.strip() for p in args.pdbs.split(",") if p.strip()]: + assert pdb in BENCH_PDBS + model, data = load_case(pdb) + canonical = model.xyz().clone() + R_true = random_rotation(seed_for(pdb, 0)) + for run in range(2): + search = model.copy() + search.spacegroup = "P 1" + search = search.copy().rotate(R_true.to(model.dtype_float), + center=canonical.mean(0)) + pipe = MolecularReplacementPipeline( + data, search, d_min=4.0, d_max=15.0, n_shells=20, + n_rotation_peaks=200, n_rotation_candidates=25, + verbose=2 if run == 1 else 0, + ) + t0 = time.perf_counter() + sols = pipe.run(do_translation=True) + secs = time.perf_counter() - t0 + rot, trans = pose_error(sols[0].model.xyz(), canonical, data.cell, + data.spacegroup) + print(f"ROW pdb={pdb} run={run} seconds={secs:.2f} rot_deg={rot:.2f} " + f"trans_A={trans:.2f}", flush=True) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py index ba82618f..8094016f 100644 --- a/alignment_lab/diagnostics/pose_recovery.py +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -23,10 +23,6 @@ the other two arms' 36/40: worse, and worse paired against ``llg`` 5 to 1, despite a rank-level harness predicting the reverse on a truth label that disagreed with coordinate superposition. -``llg_tf`` - a different question -- re-rank each candidate's TRANSLATIONS by the - likelihood, still selecting the candidate by R. - Success is a pose: final coordinates within ``--success-deg`` of canonical in orientation AND within ``--success-A`` of it in position, modulo the crystal symmetry (Cartesian point-group mates, lattice translations, allowed origin @@ -36,7 +32,7 @@ Usage:: python alignment_lab/diagnostics/pose_recovery.py --pdb 1DAW --trial 0 \ - --arms analytic_r,llg_tf --out-csv alignment_lab/runs/pose.csv + --arms analytic_r,llg --out-csv alignment_lab/runs/pose.csv """ from __future__ import annotations @@ -59,12 +55,9 @@ ARMS = { # How the winner is chosen among placed candidates. - "analytic_r": dict(use_llg_tf=False, rank_by="r"), - "corr": dict(use_llg_tf=False, rank_by="corr"), - "llg": dict(use_llg_tf=False, rank_by="llg"), - # Re-ranks each candidate's TRANSLATIONS by the likelihood, then still - # selects the candidate by R -- a different question from the three above. - "llg_tf": dict(use_llg_tf=True, rank_by="r"), + "analytic_r": dict(rank_by="r"), + "corr": dict(rank_by="corr"), + "llg": dict(rank_by="llg"), } @@ -181,7 +174,7 @@ def main() -> int: writer = None if args.out_csv: writer = ResultWriter(args.out_csv, "pose_recovery", - extra_fields=("arm", "use_llg_tf", + extra_fields=("arm", "residual_deg", "success", "n_rotation_candidates", "pipeline_seconds")) @@ -236,7 +229,6 @@ def main() -> int: orbit_side="kabsch", orbit_frame="cart", lmax_cap=_LMAX_CAP, d_min=4.0, d_max=15.0, device="cpu", arm=arm, - use_llg_tf=int(flags["use_llg_tf"]), residual_deg=(round(resid, 4) if resid == resid else ""), success=int(bool(ok)), n_rotation_candidates=args.n_rotation_candidates, diff --git a/alignment_lab/diagnostics/sigma_a_cost.py b/alignment_lab/diagnostics/sigma_a_cost.py deleted file mode 100644 index b8e23632..00000000 --- a/alignment_lab/diagnostics/sigma_a_cost.py +++ /dev/null @@ -1,75 +0,0 @@ -"""Where does the likelihood ranking's 2.5x actually go? - -The model-error fit used to be a local 81-point scan over every reflection, in -float64, evaluating both likelihood branches everywhere and discarding half -- -295 ms per candidate on 2DQ6, three times the translation refine beside it. It -now goes through the shared `SigmaAEstimator`. - -This times the pieces so the effect is measured rather than assumed: the -model-error fit, the Wilson normalisation of the calculated side, the likelihood -evaluation itself, and the translation refine they sit alongside. -""" -import sys, time -from pathlib import Path -import torch -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) -from lab import load_case, random_rotation, seed_for # noqa: E402 - - -def _t(fn, n=3): - fn() - ts = [] - for _ in range(n): - t0 = time.perf_counter(); out = fn(); ts.append(time.perf_counter() - t0) - return min(ts), out - - -def main(): - from torchref.experimental.alignment.translation import ( - DirectModelEvaluator, TranslationObs, amplitude_translation_search, - correlation_at, fit_model_error, llg_at, local_translation_refine, - normalise_calc, precompute_G_for_rotation) - - for pdb in sys.argv[1:] or ["1DAW", "2DQ6"]: - seed = seed_for(pdb, 0) - model, data = load_case(pdb) - rot = model.copy().rotate(random_rotation(seed).to(model.dtype_float), - center=model.xyz().mean(0)) - rot.spacegroup = data.spacegroup.hm - p1 = rot.copy(); p1.spacegroup = "P 1" - mask = data.get_valid_mask() - sig = getattr(data, "F_sigma", None) - obs = TranslationObs.build(data.F[mask], data.hkl[mask], data.spacegroup, - data.cell, - sig_F=None if sig is None else sig[mask], - n_shells=10) - ev = DirectModelEvaluator(p1); eye3 = torch.eye(3, dtype=torch.float64) - G, h_R = precompute_G_for_rotation(ev, eye3, obs.hkl, data.spacegroup, - data.cell) - _, _, peaks = amplitude_translation_search( - obs=obs, interpolator=ev, R_rotation=eye3, - spacegroup=data.spacegroup, real_cell=data.cell, grid_steps=16, - n_peaks=20, precomputed_G=G, precomputed_h_R=h_R) - t0 = torch.as_tensor(peaks[0].translation, dtype=torch.float64) - - t_ref, _ = _t(lambda: local_translation_refine( - obs=obs, interpolator=ev, R_rotation=eye3, - spacegroup=data.spacegroup, real_cell=data.cell, t_init=t0, - radius=0.06, grid_steps=13, n_refinement_passes=1, - precomputed_G=G, precomputed_h_R=h_R)) - ph = torch.exp(2j * torch.pi * torch.einsum( - "ind,d->in", h_R.to(torch.float64), t0.to(G.device)).to(G.dtype)) - Fc = (G * ph).sum(dim=0).abs().to(torch.float64) - t_norm, E_calc = _t(lambda: normalise_calc(Fc, obs)) - t_sa, (alpha, beta) = _t(lambda: fit_model_error(obs, E_calc)) - t_llg, _ = _t(lambda: llg_at(obs, G, h_R, t0, alpha, beta)) - t_corr, _ = _t(lambda: correlation_at(obs, G, h_R, t0)) - N = obs.hkl.numel() // 3 - print(f"ROW pdb={pdb} N={N} shells={obs.n_shells} " - f"refine={1000*t_ref:.1f}ms norm_calc={1000*t_norm:.1f}ms " - f"model_err={1000*t_sa:.1f}ms llg={1000*t_llg:.1f}ms " - f"corr={1000*t_corr:.1f}ms", flush=True) - - -main() diff --git a/alignment_lab/diagnostics/tf_cost.py b/alignment_lab/diagnostics/tf_cost.py deleted file mode 100644 index b8a4352d..00000000 --- a/alignment_lab/diagnostics/tf_cost.py +++ /dev/null @@ -1,108 +0,0 @@ -"""Where the translation stage actually spends its time, per rotation candidate. - -The pipeline carries ``n_rotation_candidates`` orientations through -:func:`_placement_for_candidate` one at a time in a Python loop, and every -orientation repays the whole stage. Before batching it over 100 orientations we -need to know which part of it is the cost: the structure-factor evaluation, the -Crowther-Blow accumulation, the LLG re-rank, or the local refine. - -Reports the per-candidate breakdown alongside the problem geometry (``N`` -reflections, ``S`` sym-ops, grid), because the four stages scale differently -- -``O(n_atoms*S*N)``, ``O(S^2*N)``, ``O(K*S*N)`` and ``O(S^2*N)`` respectively -- -and which one dominates is a property of the structure, not of the code. -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 - - -def _time(fn, repeats=1): - out = None - t0 = time.perf_counter() - for _ in range(repeats): - out = fn() - return (time.perf_counter() - t0) / repeats, out - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) - ap.add_argument("--trial", type=int, default=0) - ap.add_argument("--d-min", type=float, default=4.0) - ap.add_argument("--d-max", type=float, default=15.0) - ap.add_argument("--grid-steps", type=int, default=16) - ap.add_argument("--n-peaks", type=int, default=20) - args = ap.parse_args() - - from torchref.experimental.alignment.translation import ( - DirectModelEvaluator, TranslationObs, amplitude_translation_search, - local_translation_refine, precompute_G_for_rotation, - ) - - seed = seed_for(args.pdb, args.trial) - model, data = load_case(args.pdb) - R_true = random_rotation(seed) - rot = model.copy() - rot.spacegroup = "P 1" - rot = rot.rotate(R_true.to(model.dtype_float), center=model.xyz().mean(0)) - - # Same masking the pipeline's _prepare_translation_arrays does. - mask = data.get_valid_mask() - d = 1.0 / (data.hkl.to(torch.float64) - @ data.cell.reciprocal_basis_matrix.to(torch.float64) - ).norm(dim=-1).clamp(min=1e-9) - mask = mask & (d >= args.d_min) & (d <= args.d_max) - hkl = data.hkl[mask] - sig_F = getattr(data, "F_sigma", None) - obs = TranslationObs.build( - data.F[mask], hkl, data.spacegroup, data.cell, - sig_F=None if sig_F is None else sig_F[mask], - ) - S = int(data.spacegroup.matrices.shape[0]) - N = int(hkl.shape[0]) - n_at = int(rot.xyz().shape[0]) - print(f"# {args.pdb} sg={data.spacegroup.hm} S={S} N={N} atoms={n_at} " - f"grid={args.grid_steps} seed={seed}", flush=True) - - ev = DirectModelEvaluator(rot) - eye3 = torch.eye(3, dtype=torch.float64) - - t_G, (G, h_R) = _time(lambda: precompute_G_for_rotation( - ev, eye3, hkl, data.spacegroup, data.cell)) - t_tf, (_, _, peaks) = _time(lambda: amplitude_translation_search( - obs=obs, interpolator=ev, R_rotation=eye3, - spacegroup=data.spacegroup, real_cell=data.cell, - grid_steps=args.grid_steps, n_peaks=args.n_peaks, - precomputed_G=G, precomputed_h_R=h_R)) - t_ref, _ = _time(lambda: local_translation_refine( - obs=obs, interpolator=ev, R_rotation=eye3, - spacegroup=data.spacegroup, real_cell=data.cell, - t_init=torch.as_tensor(peaks[0].translation, dtype=torch.float64), - radius=0.06, grid_steps=13, n_refinement_passes=1, - precomputed_G=G, precomputed_h_R=h_R)) - - # The (S,N) working set the Crowther-Blow loop materialises, and the - # (K,S,N) one the LLG re-rank does, in complex128. - mb = lambda *shape: 16.0 * float(torch.tensor(shape).prod()) / 2**20 - print(f"ROW pdb={args.pdb} sg={data.spacegroup.hm} S={S} N={N} atoms={n_at} " - f"t_G={t_G:.3f} t_tf={t_tf:.3f} t_refine={t_ref:.3f} " - f"t_place={t_G + t_tf + 3 * t_ref:.3f} " - f"work_S2N={S * S * N / 1e6:.1f}M " - f"mem_SN_MB={mb(S, N):.0f} mem_KSN_MB={mb(args.n_peaks, S, N):.0f}", - flush=True) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/truth_pose_scores.py b/alignment_lab/diagnostics/truth_pose_scores.py index 22f148ae..7e2887fd 100644 --- a/alignment_lab/diagnostics/truth_pose_scores.py +++ b/alignment_lab/diagnostics/truth_pose_scores.py @@ -24,26 +24,21 @@ def scores_at(pipe, model_placed): - """(corr, R, llg) of an already-placed model through the pipeline's path.""" + """(tf score, R, llg) of an already-placed model through the pipeline's path.""" from torchref.experimental.alignment.translation import ( - DirectModelEvaluator, correlation_at, fit_model_error, llg_at, - normalise_calc, precompute_G_for_rotation) + DirectModelEvaluator, analytic_r_at, llg_at_translations, + prepare_candidate, translation_score_at) data, obs = pipe.data, pipe._obs m = model_placed.copy() + if pipe.tf_d_min > 0.0: + m.max_res = pipe.tf_d_min / 1.5 m.spacegroup = "P 1" - ev = DirectModelEvaluator(m) - eye3 = torch.eye(3, dtype=torch.float64) - G, h_R = precompute_G_for_rotation(ev, eye3, obs.hkl, data.spacegroup, data.cell) + cand = prepare_candidate(DirectModelEvaluator(m), obs, data.spacegroup, data.cell) t0 = torch.zeros(3, dtype=torch.float64) - corr = correlation_at(obs, G, h_R, t0) - Fc = G.sum(dim=0).abs().to(torch.float64) - Fo = obs.F_obs.to(Fc.device).to(torch.float64) - k = (Fo * Fc).sum() / (Fc * Fc).sum().clamp(min=1e-30) - r = float(((Fo - k * Fc).abs().sum() / Fo.sum()).item()) - E_calc = normalise_calc(Fc, obs) - alpha, beta = fit_model_error(obs, E_calc) - llg = llg_at(obs, G, h_R, t0, alpha, beta) - return corr, r, llg + tf = translation_score_at(obs, cand, t0) + r = analytic_r_at(obs, cand, t0) + llg = float(llg_at_translations(obs, cand, t0.view(1, 3))[0]) + return tf, r, llg def main() -> int: @@ -80,7 +75,7 @@ def main() -> int: print("SANITY deposited-vs-deposited rot/trans:", pose_error(canonical, canonical, data.cell, data.spacegroup)) rot, trans = pose_error(win.model.xyz(), canonical, data.cell, data.spacegroup) - print(f"WINNER rot_deg={rot:.3f} trans_A={trans:.2f} corr={win.translation_score:.5f} " + print(f"WINNER rot_deg={rot:.3f} trans_A={trans:.2f} tf={win.translation_score:.5f} " f"R={win.r_factor:.5f} llg={win.llg_score:.1f}") # Raw fractional centroid offset to every symmetry image of canonical. @@ -97,9 +92,9 @@ def main() -> int: f"|.|={float((B @ d).norm()):.1f} A") c_dep, r_dep, llg_dep = scores_at(pipe, model) - print(f"DEPOSITED corr={c_dep:.5f} R={r_dep:.5f} llg={llg_dep:.1f}") + print(f"DEPOSITED tf={c_dep:.5f} R={r_dep:.5f} llg={llg_dep:.1f}") c_w, r_w, llg_w = scores_at(pipe, win.model) - print(f"WINNER(re-scored) corr={c_w:.5f} R={r_w:.5f} llg={llg_w:.1f}") + print(f"WINNER(re-scored) tf={c_w:.5f} R={r_w:.5f} llg={llg_w:.1f}") return 0 diff --git a/docs/changelog.rst b/docs/changelog.rst index 65209d29..bddbb883 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,10 @@ Changelog Unreleased ---------- +- The fast translation function scores the covariance of two normalised intensities -- the rotation function's LERF1 coefficient ``cw (E_obs^2 - 1) w sigma_A^2`` against the candidate's ``|E_calc(h, t)|^2``, normalised per candidate by the same Wilson fit -- instead of a raw-``|F_calc|^2`` ratio that was not a correlation and, on the four largest panel structures, was higher 40 A from the true position than at it. One FFT on a grid a third of the set's resolution apart with parabolic peak refinement replaces the 16-point coarse grid and three 100-point local refines; the Rice/Woolfson likelihood at a fixed Luzzati ``sigma_A`` picks among the top peaks and ranks the candidates. 30/30 true poses at the default window and 30/30 with the window removed, against 18/30 before +- The translation stage runs in the configured float and complex dtypes rather than hard-coded double +- Removed ``use_llg_tf``, ``n_translation_peaks`` and ``translation_grid_steps`` from the pipeline; the likelihood always picks the translation, and the grid is sized by resolution +- Removed the per-candidate ``SigmaAEstimator`` fit from the placement loop. The likelihood that ranks candidates uses one ``sigma_A`` for all of them, so no candidate is scored against a model error fitted to itself - The translation search defaults to the rotation search's resolution window instead of all data. With all data it placed the four largest panel structures (2DQ6, 3VRJ, 4BX9, 6G9X) at the right orientation and 20-56 A from the true position on every trial, and its own score was higher at the wrong place than at the deposited pose; the benchmark had only ever checked the rotation. 30/30 true poses within 0.32 A against 18/30, at roughly a fifth of the wall clock - The P1 copy each rotation candidate is evaluated through is gridded at two thirds of the translation window's resolution rather than at the model's default 1.0 A. Coherence with the fine grid 0.9995 or better on all four structures measured; 10-38 ms per candidate against 200-860 - The pose-recovery harness compares against Cartesian symmetry mates and checks the translation. It compared a Cartesian rotation against fractional symmetry matrices, so in trigonal and hexagonal cells two of six and four of twelve correct mates read as 30 and 21 degrees; every recorded 2DQ6 failure was one of them diff --git a/tests/integration/alignment/profile_fit.py b/tests/integration/alignment/profile_fit.py index d4e0284c..896feb3e 100644 --- a/tests/integration/alignment/profile_fit.py +++ b/tests/integration/alignment/profile_fit.py @@ -5,8 +5,6 @@ Run from the repo root: .venv/bin/python tests/integration/alignment/profile_fit.py [--pdb 1DAW] \ [--n-rotation-candidates 3] [--n-translation-candidates 3] \ - [--translation-grid-steps 16] - Output: top-50 cumulative-time entries from cProfile + a custom per-stage timer breakdown (rotation search, ML rescore, TF, local refine, joint refine, final Scaler refit). @@ -80,7 +78,6 @@ def main(): ap.add_argument("--pdb", default="1DAW", choices=sorted(PAIRS.keys())) ap.add_argument("--n-rotation-candidates", type=int, default=3) ap.add_argument("--n-translation-candidates", type=int, default=3) - ap.add_argument("--translation-grid-steps", type=int, default=16) ap.add_argument("--top", type=int, default=40, help="top N cProfile entries") args = ap.parse_args() @@ -109,7 +106,6 @@ def main(): data, n_rotation_candidates=args.n_rotation_candidates, n_translation_candidates=args.n_translation_candidates, - translation_grid_steps=args.translation_grid_steps, verbose=0, ) profiler.disable() diff --git a/tests/unit/alignment/test_patterson_translation.py b/tests/unit/alignment/test_patterson_translation.py index 394a3e6b..a4c7136d 100644 --- a/tests/unit/alignment/test_patterson_translation.py +++ b/tests/unit/alignment/test_patterson_translation.py @@ -1,11 +1,13 @@ """Unit tests for the fast translation function. -``amplitude_translation_search`` correlates normalised, weighted ``E_obs^2`` -against ``|F_calc(h, t)|^2`` over a fractional grid. With the search model at -canonical positions and ``F_obs`` derived from a translated copy of the same -model, the top peak (or one of the top three) must land at ``-t_true`` modulo an -allowed origin shift of the space group -- the translation that would bring the -search model into agreement with the observed data. +``fast_translation_function`` accumulates the Crowther-Blow coefficients of a +normalised, weighted ``E_obs^2`` against the candidate's normalised +``|E_calc(h, t)|^2`` and inverts one FFT. With the search model at canonical +positions and ``F_obs`` derived from a translated copy of the same model, the +top peak (or one of the top three) must land at ``-t_true`` modulo an allowed +origin shift of the space group -- the translation that would bring the search +model into agreement with the observed data -- and the likelihood must prefer +that peak. ``TranslationObs`` carries the observed side. It is built once here, as the pipeline builds it once per run, because normalisation and weighting are @@ -18,8 +20,11 @@ import torch from torchref.experimental.alignment.translation import ( + DirectModelEvaluator, TranslationObs, - amplitude_translation_search, + fast_translation_function, + llg_at_translations, + prepare_candidate, ) from torchref.io.datasets.reflection_data import ReflectionData from torchref.model import ModelFT @@ -30,105 +35,101 @@ MTZ_1DAW = TEST_FILES / "mtz" / "1DAW.mtz" -class _ModelEvaluator: - """Thin evaluator: returns model_p1(hkl) at integer HKL.""" - - def __init__(self, model_p1): - self._model = model_p1 - self.device = model_p1.xyz().device - - def evaluate(self, R, hkl, real_cell, return_amplitude=False): - hkl_int = hkl.round().to(torch.int64).to(self.device) - with torch.no_grad(): - f = self._model(hkl_int) - return f.abs() if return_amplitude else f - - -def _wrap_frac(t: np.ndarray) -> np.ndarray: - return (t + 0.5) % 1.0 - 0.5 - - @pytest.fixture(scope="module") def setup(): canonical = ModelFT().load_pdb(str(PDB_1DAW)) data = ReflectionData().load_mtz(str(MTZ_1DAW)) - mask = data.get_valid_mask() + # The pipeline's default window: the rotation search's 15-4 A. + rec = data.cell.reciprocal_basis_matrix.to(torch.float64) + s = (data.hkl.to(torch.float64) @ rec).norm(dim=-1) + mask = data.get_valid_mask() & (s >= 1.0 / 15.0) & (s <= 1.0 / 4.0) return canonical, data, mask -@pytest.mark.unit -@pytest.mark.slow -def test_amplitude_tf_zero_translation(setup): - """Un-translated model: top-1 peak at the origin (modulo C-centering).""" - canonical, data, mask = setup +def _search(canonical, data, mask, t_true): + """Peaks and their likelihoods for a search model translated by ``t_true``.""" with torch.no_grad(): - F_obs = canonical(data.hkl[mask]).abs().to(torch.float64) + F_obs = canonical(data.hkl[mask]).abs() model_p1 = canonical.copy() + model_p1.max_res = 4.0 / 1.5 model_p1.spacegroup = "P 1" - evaluator = _ModelEvaluator(model_p1) - - obs = TranslationObs.build( - F_obs, data.hkl[mask], data.spacegroup, data.cell, - ) - R_id = torch.eye(3, dtype=torch.float64) - _, _, peaks = amplitude_translation_search( - obs=obs, interpolator=evaluator, R_rotation=R_id, - spacegroup=data.spacegroup, real_cell=data.cell, - grid_steps=12, n_peaks=10, cluster_radius=0.05, + if t_true is not None: + model_p1 = model_p1.translate( + torch.tensor(t_true, dtype=canonical.dtype_float), fractional=True, + ) + obs = TranslationObs.build(F_obs, data.hkl[mask], data.spacegroup, data.cell) + cand = prepare_candidate(DirectModelEvaluator(model_p1), obs, + data.spacegroup, data.cell) + _, peaks = fast_translation_function( + obs, cand, data.cell, grid_spacing_A=4.0 / 3.0, n_peaks=3, + cluster_radius_A=4.0, ) assert len(peaks) > 0 - # C2 + C-centering allowed origins: (0, *, 0) and (1/2, *, 1/2) - def origin_dist(t): - xz0 = np.linalg.norm(_wrap_frac(np.array([t[0], t[2]]))) - xz1 = np.linalg.norm(_wrap_frac(np.array([t[0] - 0.5, t[2] - 0.5]))) - return min(xz0, xz1) - best_dist = min(origin_dist(p.translation) for p in peaks[:3]) - assert best_dist < 0.10, ( - f"top-3 peaks miss origin-equivalent by {best_dist:.3f}; " - f"peaks: {[p.translation.tolist() for p in peaks[:3]]}" + llg = llg_at_translations( + obs, cand, + torch.as_tensor(np.stack([p.translation for p in peaks]), dtype=torch.float64), ) + return peaks, llg + + +def _xz_dist_to_origin_class(t: np.ndarray, t_true: np.ndarray) -> float: + """Distance of ``t + t_true`` from an allowed origin in C2, x and z only. + + C2's origin is free along y; the centring makes (1/2, 1/2, 0) a lattice + vector and (0, *, 1/2) an allowed shift, so x and z are each determined + only modulo 1/2. + """ + d = np.array([t[0] + t_true[0], t[2] + t_true[2]]) + d = (d + 0.25) % 0.5 - 0.25 + return float(np.linalg.norm(d)) @pytest.mark.unit @pytest.mark.slow -def test_amplitude_tf_recovers_known_translation(setup): - """A model translated by t_true: top-3 peaks include -t_true (mod origins).""" +def test_fast_tf_zero_translation(setup): + """Un-translated model: the likelihood's pick sits at an origin-equivalent.""" + canonical, data, mask = setup + peaks, llg = _search(canonical, data, mask, None) + best = peaks[int(llg.argmax())] + dist = _xz_dist_to_origin_class(best.translation, np.zeros(3)) + assert dist < 0.03, ( + f"likelihood pick misses an origin-equivalent by {dist:.3f}; " + f"peaks: {[p.translation.round(3).tolist() for p in peaks]}" + ) + + +@pytest.mark.unit +@pytest.mark.slow +def test_fast_tf_recovers_known_translation(setup): + """A model translated by t_true: the likelihood's pick is at -t_true (mod origins).""" canonical, data, mask = setup t_true = np.array([0.18, -0.07, 0.23]) - # F_obs from canonical (un-translated); search model is canonical_p1 - # translated by t_true (so the recovered TF peak should be at -t_true mod - # the allowed origin shifts). + peaks, llg = _search(canonical, data, mask, t_true) + best = peaks[int(llg.argmax())] + dist = _xz_dist_to_origin_class(best.translation, t_true) + assert dist < 0.03, ( + f"likelihood pick misses -t_true by {dist:.3f}; " + f"peaks: {[p.translation.round(3).tolist() for p in peaks]}" + ) + # The fast map's own top peak should already be the right one here; the + # likelihood is the arbiter when it is not. + assert _xz_dist_to_origin_class(peaks[0].translation, t_true) < 0.03 + + +@pytest.mark.unit +@pytest.mark.slow +def test_e_calc_is_normalised(setup): + """```` is one to within the fit's tolerance, for a placed candidate.""" + canonical, data, mask = setup with torch.no_grad(): - F_obs = canonical(data.hkl[mask]).abs().to(torch.float64) + F_obs = canonical(data.hkl[mask]).abs() model_p1 = canonical.copy() + model_p1.max_res = 4.0 / 1.5 model_p1.spacegroup = "P 1" - model_p1 = model_p1.translate( - torch.tensor(t_true, dtype=canonical.dtype_float), fractional=True, - ) - evaluator = _ModelEvaluator(model_p1) - - obs = TranslationObs.build( - F_obs, data.hkl[mask], data.spacegroup, data.cell, - ) - R_id = torch.eye(3, dtype=torch.float64) - _, _, peaks = amplitude_translation_search( - obs=obs, interpolator=evaluator, R_rotation=R_id, - spacegroup=data.spacegroup, real_cell=data.cell, - grid_steps=12, n_peaks=10, cluster_radius=0.05, - ) - assert len(peaks) > 0 - # Expected: t_peak ≡ -t_true (mod allowed origin). C2 allowed origins - # along x and z: (0, *, 0) and (1/2, *, 1/2). y is polar. - def xz_dist(t): - d_origin = np.linalg.norm( - _wrap_frac(np.array([t[0] + t_true[0], t[2] + t_true[2]])) - ) - d_cshift = np.linalg.norm( - _wrap_frac(np.array([t[0] + t_true[0] - 0.5, t[2] + t_true[2] - 0.5])) - ) - return min(d_origin, d_cshift) - best_dist = min(xz_dist(p.translation) for p in peaks[:3]) - assert best_dist < 0.10, ( - f"top-3 peaks don't bracket -t_true (best xz_dist {best_dist:.3f}); " - f"peaks: {[p.translation.tolist() for p in peaks[:3]]}" - ) + obs = TranslationObs.build(F_obs, data.hkl[mask], data.spacegroup, data.cell) + cand = prepare_candidate(DirectModelEvaluator(model_p1), obs, + data.spacegroup, data.cell) + # The normalisation already carries eps: E is per unit of eps*Sigma_calc. + E2 = cand.e_calc(torch.zeros(3, dtype=torch.float64)) ** 2 + mean_e2 = float(E2.mean()) + assert abs(mean_e2 - 1.0) < 0.15, mean_e2 diff --git a/tests/unit/alignment/test_translation_obs.py b/tests/unit/alignment/test_translation_obs.py index 8d3d703a..4d97fd34 100644 --- a/tests/unit/alignment/test_translation_obs.py +++ b/tests/unit/alignment/test_translation_obs.py @@ -11,7 +11,6 @@ per-shell one (per-shell weights cancel in a correlation), and the whole object is invariant to the units the amplitudes arrive in. """ -import math import pytest import torch @@ -56,7 +55,7 @@ def test_normalisation_is_the_shared_wilson_fit(): obs.F_obs * obs.F_obs, obs.s_mag, eps=obs.eps, centric=obs.centric, n_coeff=6, ) - torch.testing.assert_close(obs.E_obs, direct.E.to(torch.float64)) + torch.testing.assert_close(obs.E_obs, direct.E.to(obs.E_obs.dtype)) @pytest.mark.unit @@ -105,9 +104,14 @@ def test_weight_varies_within_a_shell(): # Give two reflections at the SAME resolution very different sigmas. obs = TranslationObs.build(F, hkl, sg, cell, sig_F=sig_F) + from torchref.experimental.alignment.sh import (assign_shells, + equal_count_shell_edges) + + edges, _ = equal_count_shell_edges(obs.s_mag, 20) + shell_idx = assign_shells(obs.s_mag, edges).clamp(min=0) within = [] - for shell in range(obs.n_shells): - w = obs.weight[obs.shell_idx == shell] + for shell in range(20): + w = obs.weight[shell_idx == shell] if w.numel() > 20: within.append(float(w.std() / w.mean().clamp(min=1e-30))) assert within, "no populated shells" @@ -148,27 +152,24 @@ def test_epsilon_and_centricity_come_from_the_spacegroup(sg_name): obs = TranslationObs.build(F, hkl, sg, cell) hkl_l = obs.hkl.round().to(torch.int64) torch.testing.assert_close( - obs.eps, sg.epsilon(hkl_l, friedel=False).to(torch.float64).clamp(min=1.0), + obs.eps, sg.epsilon(hkl_l, friedel=False).to(obs.eps.dtype).clamp(min=1.0), ) torch.testing.assert_close(obs.centric, sg.is_centric(hkl_l).to(torch.bool)) @pytest.mark.unit -def test_shell_binning_is_equal_count_and_shared(): - """One binning, used by both the sigma_A fit and the likelihood. +def test_coefficient_is_the_rotation_functions_score_equation(): + """The fast search's coefficient is LERF1's intensity times sigma_A^2. - They used to derive their own from the same |s| -- one rank-based, one - value-based -- which put boundary reflections in different shells depending - on which stage asked. + One score equation for both searches: ``cw (E^2 - 1) w`` is what the + rotation function expands, and ``sigma_A^2`` is its calc-side weight. """ - F, _, hkl, sg, cell, _ = _case() - obs = TranslationObs.build(F, hkl, sg, cell, n_shells=10) - counts = torch.bincount(obs.shell_idx, minlength=obs.n_shells) - assert obs.n_shells == 10 - assert int(counts.min()) > 0 - # Equal-count binning: no shell should be wildly larger than the mean. - assert float(counts.max()) < 3.0 * float(counts.to(torch.float64).mean()) - # And it must be monotone in |s|: shells partition resolution, not noise. - hi = torch.stack([obs.s_mag[obs.shell_idx == b].max() - for b in range(obs.n_shells)]) - assert bool((hi[1:] >= hi[:-1]).all()) + from torchref.experimental.alignment.frf.preprocessing import ( + build_lerf1_intensity, eterm_sigma_a) + + F, sig_F, hkl, sg, cell, _ = _case() + obs = TranslationObs.build(F, hkl, sg, cell, sig_F=sig_F, delta_vrms_A=0.8) + expected = (build_lerf1_intensity(obs.E_obs, obs.centric, weight=obs.weight) + * eterm_sigma_a(obs.s_mag, 0.8) ** 2) + torch.testing.assert_close(obs.coeff, expected) + torch.testing.assert_close(obs.sigma_a, eterm_sigma_a(obs.s_mag, 0.8)) diff --git a/torchref/experimental/alignment/__init__.py b/torchref/experimental/alignment/__init__.py index d0b9ea3e..0d13918c 100644 --- a/torchref/experimental/alignment/__init__.py +++ b/torchref/experimental/alignment/__init__.py @@ -66,15 +66,14 @@ rotation_search, ) from .translation import ( + CandidateTransform, DirectModelEvaluator, TranslationObs, TranslationPeak, - amplitude_translation_search, - find_translation_peaks, - fit_model_error, - llg_translation_rescore, - local_translation_refine, - precompute_G_for_rotation, + analytic_r_at, + fast_translation_function, + llg_at_translations, + prepare_candidate, ) from .pipeline import ( MolecularReplacementPipeline, @@ -103,10 +102,9 @@ "TranslationObs", "TranslationPeak", "DirectModelEvaluator", - "amplitude_translation_search", - "local_translation_refine", - "llg_translation_rescore", - "precompute_G_for_rotation", - "find_translation_peaks", - "fit_model_error", + "fast_translation_function", + "CandidateTransform", + "analytic_r_at", + "prepare_candidate", + "llg_at_translations", ] diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index a175de58..42260d98 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -7,10 +7,10 @@ to rank well, and measurably does not: over 30 seeded cells it puts truth at rank 0 in 6 of them. What it does reliably is put truth somewhere in the top twenty. -2. **Fast Translation Function** — for *each* of the top-N orientations, a - Crowther-Blow amplitude-correlation search over the fractional cell, - optionally re-ranked by a Rice/Woolfson LLG, then an analytical-R local - refine. On the same 30 cells it puts truth at rank 0 in 27. Rotation ghosts +2. **Fast Translation Function** — for *each* of the top-N orientations, one + Crowther-Blow FFT over the fractional cell on a resolution-sized grid, with + the rotation function's own normalised score equation as its coefficients, + then the Rice/Woolfson likelihood at the best few peaks. Rotation ghosts are morphologically identical to truth in a rotation function by construction; they are not identical once the crystal is involved. @@ -60,15 +60,10 @@ from .translation import ( DirectModelEvaluator, TranslationObs, - TranslationPeak, - amplitude_translation_search, - correlation_at, - llg_at, - fit_model_error, - llg_translation_rescore, - local_translation_refine, - normalise_calc, - precompute_G_for_rotation, + analytic_r_at, + fast_translation_function, + llg_at_translations, + prepare_candidate, ) if TYPE_CHECKING: @@ -158,7 +153,7 @@ class MRSolution: rotation_score : float The rotation function's score for this candidate. translation_score : float - The translation function's correlation at the chosen translation, higher + The fast translation function's score at the chosen peak, higher better. Reported, not ranked -- see the sort in :meth:`run`. r_factor : float The analytical-scale R at that placement, lower better. Reported, not @@ -166,10 +161,9 @@ class MRSolution: would return -- build one on the returned model if that is wanted. llg_score : float **The ranking key**: the translation likelihood at that placement, - higher better. ``nan`` when ``rank_by`` is not ``"llg"``, since it costs - a sigma_A fit and a likelihood evaluation per candidate. + higher better. model : ModelFT - The rotated (+translated +refined) model for this candidate. + The rotated and translated model for this candidate. """ rotation: np.ndarray @@ -241,10 +235,10 @@ def __init__( model_error_A: Optional[float] = None, # --- candidate tree --- n_rotation_candidates: int = 25, - n_translation_peaks: int = 20, + # Peaks of the fast translation function re-scored by the likelihood + # for each orientation. The fast map only has to get the true peak + # into this many; the likelihood picks. n_translation_candidates: int = 3, - translation_grid_steps: int = 16, - use_llg_tf: bool = False, # Which score picks the winner among placed candidates. "llg" is the # translation function's Rice/Woolfson likelihood; "r" the # analytical-scale R-factor; "corr" the translation correlation. Not a @@ -280,10 +274,7 @@ def __init__( self.model_error_A = float(model_error_A) self.n_rotation_candidates = n_rotation_candidates - self.n_translation_peaks = n_translation_peaks self.n_translation_candidates = n_translation_candidates - self.translation_grid_steps = translation_grid_steps - self.use_llg_tf = use_llg_tf if rank_by not in ("r", "corr", "llg"): raise ValueError( f"rank_by={rank_by!r}; expected 'r', 'corr' or 'llg'.") @@ -582,7 +573,6 @@ def _prepare_translation_arrays(self) -> None: data.spacegroup, data.cell, sig_F=None if sig_F_full is None else sig_F_full[tmask], delta_vrms_A=self.model_error_A, - n_shells=max(self.n_shells // 2, 8), device=device, ) if self.verbose >= 1: @@ -594,26 +584,18 @@ def _prepare_translation_arrays(self) -> None: else " (no sigmas: unit weight)")) def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: - """Translation search + analytical-R local refine for one rotation. + """Translation search for one rotation candidate. - Returns ``(r_analytic, t_refined, tf_score, llg_score)`` for the best - translation of this rotation candidate, or ``None`` if no translation - peaks were found. ``llg_score`` is ``nan`` unless it is the ranking key. - - ``tf_score`` is the translation function's own score at its top peak. - It does not select anything -- ``r_analytic`` does -- but it is carried - out so ``verbose >= 2`` can report both. Which of the two a wrong - placement disagreed on is the first thing anyone diagnosing one asks, - and it is not recoverable afterwards from the winner alone. + Returns ``(r_analytic, t, tf_score, llg)`` for the translation the + likelihood prefers among the fast search's top peaks, or ``None`` if the + map had no peaks. All three scores are at the same ``t``, so the + reported numbers belong to the placement that was actually chosen. """ data = self.data timer = self._timer - eye3 = self._eye3 + obs = self._obs if str(rotated_k.spacegroup) != str(data.spacegroup): - # NOTE: assign the space-group NAME, not a SpaceGroup object. SpaceGroup is an - # nn.Module, so nn.Module.__setattr__ intercepts object assignment, stores it in - # _modules and never runs the property setter -- a silent no-op. rotated_k.spacegroup = data.spacegroup.hm rotated_p1 = rotated_k.copy() # Size the P1 copy's FFT grid to the translation set, not to the @@ -633,130 +615,39 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: rotated_p1.spacegroup = "P 1" evaluator = DirectModelEvaluator(rotated_p1) - timer.start("5_precompute_G") - G_pre, h_R_pre = precompute_G_for_rotation( - evaluator, eye3, self._obs.hkl, data.spacegroup, data.cell, + timer.start("5_candidate_transform") + cand = prepare_candidate(evaluator, obs, data.spacegroup, data.cell) + timer.stop("5_candidate_transform") + + # One FFT on a grid a third of the set's resolution apart: dense enough + # that the parabolic peak refinement lands within a fraction of a step, + # and no coarse-then-refine pair whose coarse half could miss the peak. + d_min_set = 1.0 / float(obs.s_mag.max()) + timer.start("6_translation_function") + _, t_peaks = fast_translation_function( + obs, cand, data.cell, + grid_spacing_A=d_min_set / 3.0, + n_peaks=self.n_translation_candidates, + cluster_radius_A=d_min_set, ) - timer.stop("5_precompute_G") - - timer.start("6_amplitude_TF") - _, _, t_peaks = amplitude_translation_search( - obs=self._obs, interpolator=evaluator, R_rotation=eye3, - spacegroup=data.spacegroup, real_cell=data.cell, - grid_steps=self.translation_grid_steps, - n_peaks=self.n_translation_peaks, - cluster_radius=0.05, - precomputed_G=G_pre, precomputed_h_R=h_R_pre, - ) - timer.stop("6_amplitude_TF") + timer.stop("6_translation_function") if not t_peaks: return None - if self.use_llg_tf: - t_peaks = self._llg_tf_rescore(t_peaks, G_pre, h_R_pre) - - tf_top = float(t_peaks[0].score) - if self.verbose >= 3: - tt = tuple(round(float(x), 3) for x in t_peaks[0].translation.tolist()) - self._log(3, f" top translation t={tt} score={tf_top:.4f}") - - best = None - for k_t, tp in enumerate(t_peaks[:self.n_translation_candidates]): - t_init = torch.as_tensor(tp.translation, dtype=torch.float64) - timer.start("7_local_TF_refine") - t_refined, r_analytic = local_translation_refine( - obs=self._obs, interpolator=evaluator, R_rotation=eye3, - spacegroup=data.spacegroup, real_cell=data.cell, - t_init=t_init, radius=0.06, grid_steps=13, - n_refinement_passes=1, - precomputed_G=G_pre, precomputed_h_R=h_R_pre, - ) - timer.stop("7_local_TF_refine") - # Both scores at the REFINED position, so the reported correlation - # belongs to the translation that was actually chosen. Selection is - # by R: ranking candidates by the correlation instead was measured - # end to end and is WORSE (32/40 against 36/40 over four structures - # x ten seeds), despite a rank-level harness predicting the reverse. - tf_ref = correlation_at(self._obs, G_pre, h_R_pre, t_refined) - self._log(3, f" trans{k_t}: tf={tf_ref:.5f} " - f"R(analytic)={r_analytic:.4f}, " - f"t={[round(float(x), 3) for x in t_refined.tolist()]}") - if best is None or r_analytic < best[0]: - best = (r_analytic, t_refined, tf_ref) - if best is None: - return None - llg = float("nan") - if self.rank_by == "llg": - # Model error at the chosen translation, per candidate. Fitting it - # once and sharing it across candidates was measured to give - # identical rankings, so the cheaper-to-reason-about form is used. - E_calc = normalise_calc( - self._fcalc_at(G_pre, h_R_pre, best[1]), self._obs) - alpha, beta = fit_model_error(self._obs, E_calc) - llg = llg_at(self._obs, G_pre, h_R_pre, best[1], alpha, beta) - return (best[0], best[1], best[2], llg) - - def _fcalc_at(self, G, h_R, t): - """``|F_calc(h, t)|`` from the precomputed per-symop contributions.""" - tt = t.detach().to(G.device).to(torch.float64).reshape(3) - phase = torch.exp(2j * torch.pi * torch.einsum( - "ind,d->in", h_R.to(torch.float64), tt).to(G.dtype)) - return (G * phase).sum(dim=0).abs().to(torch.float64) - - - def _llg_tf_rescore(self, t_peaks, G_pre, h_R_pre): - """Re-rank translation peaks by a shared-σA Rice/Woolfson LLG. - - Mirrors Phaser's FTF: the amplitude correlation is a cheap pre-filter, - and this is the likelihood that ranks its peaks. It reuses the run's - single Wilson normalisation and its shell binning, so ``E_obs`` here is - the same ``E_obs`` the correlation maximised. - - Off by default. It is the strongest discriminator at rank level, and - end-to-end it changes nothing: 27/30 against 28/30 with one discordant - cell in 30, which the correlation wins. - """ - device = self.device - obs = self._obs - self._timer.start("6b_llg_tf_rescore") - - # The model error is fitted against the top translation only and reused - # for every candidate. It is a model-reliability curve, not a - # per-candidate score: refitting it per t would let each candidate - # choose the alpha that flatters it, which is scoring a model against a - # likelihood tuned to that model. - t_top_t = torch.as_tensor( - t_peaks[0].translation, dtype=torch.float64, device=device, - ) - phase_top = torch.exp( - 2j * torch.pi * torch.einsum( - "ind,d->in", h_R_pre.to(torch.float64), t_top_t, - ).to(G_pre.dtype), - ) - Fc_top = (G_pre * phase_top).sum(dim=0).abs().to(torch.float64) - E_calc_top = normalise_calc(Fc_top, obs) - alpha_tf, beta_tf = fit_model_error(obs, E_calc_top) - + timer.start("7_translation_llg") t_cands = torch.as_tensor( - np.stack([p.translation for p in t_peaks]), - dtype=torch.float64, device=device, + np.stack([p.translation for p in t_peaks]), dtype=torch.float64, ) - llg_tf = llg_translation_rescore( - obs=obs, G=G_pre, h_R=h_R_pre, t_candidates=t_cands, - alpha=alpha_tf, beta=beta_tf, - ) - self._timer.stop("6b_llg_tf_rescore") - - llg_list = llg_tf.detach().cpu().tolist() - order = sorted(range(len(t_peaks)), key=lambda i: llg_list[i], reverse=True) - return [ - TranslationPeak( - translation=t_peaks[i].translation, - score=float(llg_list[i]), - sigma=float(llg_list[i]), - ) - for i in order - ] + llg = llg_at_translations(obs, cand, t_cands) + k_best = int(llg.argmax()) + t_best = t_cands[k_best] + r_analytic = analytic_r_at(obs, cand, t_best) + timer.stop("7_translation_llg") + for k_t, tp in enumerate(t_peaks): + self._log(3, f" trans{k_t}: tf={tp.score:.4f} z={tp.sigma:.2f} " + f"llg={float(llg[k_t]):.1f} " + f"t={[round(float(x), 3) for x in tp.translation]}") + return (r_analytic, t_best, float(t_peaks[k_best].score), float(llg[k_best])) # --------------------------------------------------------------------------- @@ -774,11 +665,8 @@ def align_model_to_data( n_rotation_peaks: int = 500, verbose: int = 0, do_translation: bool = True, - n_translation_peaks: int = 20, n_translation_candidates: int = 3, - translation_grid_steps: int = 16, n_rotation_candidates: int = 25, - use_llg_tf: bool = False, rank_by: str = "llg", tf_d_min: Optional[float] = None, tf_d_max: Optional[float] = None, @@ -806,10 +694,8 @@ def align_model_to_data( n_rotation_peaks=n_rotation_peaks, model_error_A=model_error_A, n_rotation_candidates=n_rotation_candidates, - n_translation_peaks=n_translation_peaks, n_translation_candidates=n_translation_candidates, - translation_grid_steps=translation_grid_steps, - use_llg_tf=use_llg_tf, rank_by=rank_by, + rank_by=rank_by, tf_d_min=tf_d_min, tf_d_max=tf_d_max, ) solutions = pipeline.run(do_translation=do_translation) diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index ab19e3f2..2cdc25ac 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -3,36 +3,48 @@ A translation shifts phase, ``F(h, t) = F(h) exp(2 pi i h.t)``, so scoring every ``t`` on a grid is a Fourier transform rather than a scan. The Crowther-Blow form used here accumulates the pair coefficients -``sum_h w(h) G_i*(h) G_j(h)`` onto a reciprocal grid at +``sum_h c(h) G_i*(h) G_j(h)`` onto a reciprocal grid at ``(h R_j - h R_i) mod G`` and takes one inverse FFT, which replaces ``G^3`` grid evaluations with a single transform. -This is where the discrimination happens. The rotation function upstream is a -shortlist generator -- over 30 seeded cells it puts truth at rank 0 six times; -the correlation here does it 24 times and the likelihood 27. Rotation ghosts are -morphologically identical to truth in a Patterson by construction, and stop -being identical as soon as the crystal lattice is involved. - -The observed side is prepared **once** per run, by -:class:`TranslationObs`, and reused for every orientation and every candidate -translation. That is not only an optimisation: normalisation and weighting are -properties of the observations, which do not change when the model moves, and -three separate answers to "what is the mean intensity here" used to live in this -module and its caller. +Both sides of that sum are **normalised**. The observed side is the rotation +search's own LERF1 intensity, ``cw (E_obs^2 - 1) w sigma_A^2``, built from the +run's one Wilson fit; the calculated side is the oriented model's transform +divided by its own Wilson curve, so ``<|E_calc(h, t)|^2> = 1`` per shell for +every candidate. The score is then a covariance of two normalised intensities +and every resolution shell carries the weight the model error gives it. The +previous form divided raw ``|F_calc|^2`` by its own sum, which is not a +correlation: on 2DQ6 it was 0.665 at a position 41 A from the deposited pose +and 0.350 at the pose itself, and the search followed it there. + +The grid is sized to the resolution of the translation set, one FFT per +candidate, and the best few peaks are re-scored with the full Rice/Woolfson +likelihood at fixed ``sigma_A``. That likelihood is also what ranks the +candidates against each other. + +The observed side is prepared **once** per run, by :class:`TranslationObs`, +and reused for every orientation. Normalisation, weighting and model error are +properties of the observations and the search model, which do not change when +the model moves. """ +from __future__ import annotations + +import math +from dataclasses import dataclass +from typing import List, Optional, Tuple, TYPE_CHECKING + import numpy as np import torch from torchref.base.targets.xray_likelihoods import rice_per_refl -from torchref.config import get_default_device +from torchref.config import get_complex_dtype, get_default_device, get_float_dtype from torchref.scaling import WilsonNormaliser from torchref.scaling.weighting import (inverse_variance_weight, normalise_weight, snr_from_amplitude) +from torchref.symmetry.symmetry import find_fft_friendly_size -from .sh import assign_shells, equal_count_shell_edges -from dataclasses import dataclass -from typing import List, Optional, Tuple, TYPE_CHECKING +from .frf.preprocessing import build_lerf1_intensity, eterm_sigma_a if TYPE_CHECKING: # pragma: no cover - typing only from ...model.model_ft import ModelFT @@ -43,16 +55,18 @@ #: different order on each would be two normalisations again. WILSON_N_COEFF = 6 +#: Largest FFT grid per axis. 256^3 complex64 is 134 MB, which bounds the +#: translation map for an uncut high-resolution set on a long cell; at the +#: default window the grid never reaches it. +MAX_GRID_PER_AXIS = 256 + @dataclass class TranslationObs: """The observed side of a translation search, normalised and weighted once. - Everything here is a property of the observations alone, so none of it - changes when the model rotates or moves. Building it per orientation -- which - is what the module used to do -- refits a Gamma GLM for every candidate to - get the same answer back, and worse, it made "what is ``E_obs``" a question - with three different answers depending on which function you asked. + Everything here is a property of the observations and of the search model's + expected error, so none of it changes when the model rotates or moves. Attributes ---------- @@ -60,22 +74,24 @@ class TranslationObs: The masked observations and their crystallographic bookkeeping. E_obs : torch.Tensor ``F / sqrt(eps Sigma(s))``, with ``Sigma`` the shared Wilson fit, so - `` = 1`` as an identity of that fit rather than as a separate - normalisation step. + `` = 1`` as an identity of that fit. weight : torch.Tensor Mean-1 inverse-variance weight, from measurement error and model error - in one denominator. **This is the half that does not cancel.** A - per-resolution *scaling* is gauge in a correlation -- twelve conventions - moved the rotation function's truth rank by nothing -- but a weight that - varies within a shell is not, and until now the translation search had - none at all: every reflection counted the same. - shell_idx, n_shells - Equal-count binning in ``|s|``, shared by the sigma_A fit and the - likelihood so the two cannot disagree about which reflection is where. + in one denominator. Uniform when the data carry no sigmas. + sigma_a : torch.Tensor + The Luzzati fall-off ``exp(-(2 pi^2 / 3) s^2 vrms^2)`` for the search + model's expected coordinate error -- the same term the rotation + function weights with. It is the ``D`` of the likelihood and the + calc-side weight of the fast search. A prior, not a fit: nothing can be + fitted before the model is placed. + coeff : torch.Tensor + The fast search's per-reflection coefficient, + ``cw (E_obs^2 - 1) weight sigma_A^2`` -- the rotation function's LERF1 + intensity with its calc-side ``sigma_A^2`` folded in. Centred, so a + placement that puts calculated intensity everywhere gains nothing. fit : WilsonNormaliser Kept, not discarded. Anything comparing an observed curve against a - calculated one needs the curve itself, not the per-reflection values it - produced. + calculated one needs the curve itself. """ F_obs: torch.Tensor @@ -85,8 +101,8 @@ class TranslationObs: eps: torch.Tensor E_obs: torch.Tensor weight: torch.Tensor - shell_idx: torch.Tensor - n_shells: int + sigma_a: torch.Tensor + coeff: torch.Tensor fit: "WilsonNormaliser" @classmethod @@ -99,7 +115,6 @@ def build( *, sig_F: Optional[torch.Tensor] = None, delta_vrms_A: float = 1.0, - n_shells: int = 20, n_coeff: int = WILSON_N_COEFF, device=None, ) -> "TranslationObs": @@ -120,11 +135,10 @@ def build( having it. delta_vrms_A : float R.m.s. coordinate error of the search model, which sets the model - half of the variance budget through the Luzzati falloff. The same - number the rotation function weights with. + half of the variance budget and the likelihood's ``sigma_A``. """ dev = get_default_device() if device is None else device - real = torch.float64 + real = get_float_dtype() F = F_obs.detach().to(dev) F = (F.abs() if F.is_complex() else F).to(real) hkl_i = hkl.detach().to(dev) @@ -143,27 +157,22 @@ def build( fit = WilsonNormaliser( F * F, s_mag, eps=eps, centric=centric, n_coeff=n_coeff, ) + sigma_a = eterm_sigma_a(s_mag, float(delta_vrms_A)).to(real) if sig_F is None: weight = torch.ones_like(F) else: sig = sig_F.detach().to(dev).to(real).abs() - # eterm_sigma_a is the rotation function's own model-error term; - # importing it rather than restating the exponent is the point. - from .frf.preprocessing import eterm_sigma_a weight = normalise_weight(inverse_variance_weight( - snr_from_amplitude(F, sig), - eterm_sigma_a(s_mag, float(delta_vrms_A)).to(real), - eps=eps, + snr_from_amplitude(F, sig), sigma_a, eps=eps, )) - edges, _ = equal_count_shell_edges(s_mag, n_shells) - shell_idx = assign_shells(s_mag, edges).clamp(min=0) + E_obs = fit.E.to(real) + coeff = build_lerf1_intensity(E_obs, centric, weight=weight) * sigma_a ** 2 return cls( F_obs=F, hkl=hkl_i, s_mag=s_mag, centric=centric, eps=eps, - E_obs=fit.E.to(real), weight=weight, - shell_idx=shell_idx, n_shells=int(n_shells), fit=fit, + E_obs=E_obs, weight=weight, sigma_a=sigma_a, coeff=coeff, fit=fit, ) @@ -195,691 +204,278 @@ def evaluate(self, R, hkl, real_cell, return_amplitude=False): @dataclass -class TranslationPeak: +class CandidateTransform: + """One oriented model's transform at the symmetry-rotated indices. + + Attributes + ---------- + G : torch.Tensor + ``(S, N)`` complex. ``G_i(h) = F_p1(h R_i) exp(2 pi i h.t_i) / norm(h)``, + so ``E_calc(h, t) = |sum_i G_i(h) exp(2 pi i (h R_i).t)|`` is the + **normalised** calculated amplitude: `` = 1`` per shell. + h_R : torch.Tensor + ``(S, N, 3)`` the rotated indices ``h R_i``. + norm : torch.Tensor + ``(N,)`` ``sqrt(eps n_ops Sigma_P(s))``: the raw amplitude is + ``E_calc * norm``. """ - Translation search peak. + + G: torch.Tensor + h_R: torch.Tensor + norm: torch.Tensor + + def e_calc(self, t: torch.Tensor) -> torch.Tensor: + """``E_calc(h, t)`` for ``t`` of shape ``(3,)`` or ``(K, 3)``: ``(N,)`` or ``(K, N)``.""" + single = t.ndim == 1 + tt = t.reshape(-1, 3).to(self.h_R.device).to(self.h_R.dtype) + phase_arg = torch.einsum("ind,kd->kin", self.h_R, tt) + phase = torch.exp((2j * math.pi) * phase_arg.to(self.G.dtype)) + E = (self.G.unsqueeze(0) * phase).sum(dim=1).abs() + return E[0] if single else E + + +def prepare_candidate( + evaluator, + obs: TranslationObs, + spacegroup, + real_cell, +) -> CandidateTransform: + """Evaluate one orientation's transform and normalise it. + + The only per-candidate model evaluation: ``F_p1`` at all ``S x N`` rotated + indices in one call. The normalising curve ``Sigma_P(s)`` is the same + Wilson fit the observed side uses, on the same abscissa, fitted to the + transform's mean intensity over the ``S`` copies -- which is what the crystal + sum averages to over a shell, since the cross terms between symmetry copies + have zero mean over ``h``. The crystal's ``<|F_calc|^2>`` is then + ``eps n_ops Sigma_P``, and dividing by it is what puts every candidate's + ``E_calc`` on one footing with ``E_obs`` and with each other. + """ + device = get_default_device() + real = get_float_dtype() + cplx = get_complex_dtype() + + hkl = obs.hkl.to(device).to(real) + sym_R = spacegroup.matrices.detach().to(device).to(real) + sym_t = spacegroup.translations.detach().to(device).to(real) + S = int(sym_R.shape[0]) + N = int(hkl.shape[0]) + + # h_R[i, n, d] = sum_e hkl[n, e] sym_R[i, e, d]: the h.S convention. + h_R = torch.einsum("ne,ied->ind", hkl, sym_R) + phase = torch.exp((2j * math.pi) * torch.einsum("ne,ie->in", hkl, sym_t).to(cplx)) + eye3 = torch.eye(3, dtype=real, device=device) + F_all = evaluator.evaluate( + eye3, h_R.reshape(-1, 3), real_cell, return_amplitude=False, + ).reshape(S, N).to(cplx) + G_raw = F_all * phase + + I_P = (G_raw.abs() ** 2).mean(dim=0).to(real) + s_mag = obs.s_mag.to(device).to(real) + fit_P = WilsonNormaliser( + I_P, s_mag, n_coeff=WILSON_N_COEFF, + s_lo=float(s_mag.min()), s_hi=float(s_mag.max()), + ) + Sigma_c = S * fit_P.evaluate(s_mag).to(real) + norm = (obs.eps.to(device).to(real) * Sigma_c).clamp(min=1e-30).sqrt() + return CandidateTransform(G=G_raw / norm.to(cplx), h_R=h_R, norm=norm) + + +@dataclass +class TranslationPeak: + """A peak of the fast translation function. Attributes ---------- translation : np.ndarray - Fractional coordinates (3,). + Fractional coordinates (3,), refined to sub-grid precision. score : float - Correlation score. + The fast search's score at the grid maximum. sigma : float - Z-score above mean. + Standard deviations above the map mean. """ translation: np.ndarray score: float sigma: float -def find_translation_peaks( - correlation_map: np.ndarray, - n_peaks: int = 10, - cluster_radius: float = 0.05, -) -> List[TranslationPeak]: - """ - Extract and cluster peaks from translation function. +def _grid_sizes(real_cell, grid_spacing_A: float) -> Tuple[int, int, int]: + """FFT-friendly grid, at most ``grid_spacing_A`` apart along each axis.""" + sizes = [] + for length in (real_cell.a, real_cell.b, real_cell.c): + n = int(math.ceil(float(length) / float(grid_spacing_A))) + n = find_fft_friendly_size(max(n, 4)) + sizes.append(min(n, MAX_GRID_PER_AXIS)) + return sizes[0], sizes[1], sizes[2] - Parameters - ---------- - correlation_map : np.ndarray - Translation function values, shape (Nx, Ny, Nz). - n_peaks : int - Maximum number of peaks to return. - cluster_radius : float - Minimum fractional distance between peaks (periodic). - Returns - ------- - peaks : list - List of TranslationPeak objects sorted by score. - """ - Nx, Ny, Nz = correlation_map.shape - mean_val = correlation_map.mean() - std_val = correlation_map.std() - - if std_val < 1e-10: - return [] - - flat = correlation_map.flatten() - sorted_idx = np.argsort(flat)[::-1] - - peaks = [] - used = [] - - for idx in sorted_idx: - if len(peaks) >= n_peaks: - break +def _parabolic_offset(fm: float, f0: float, fp: float) -> float: + """Sub-grid offset of a maximum from its three samples, in grid units.""" + denom = fm - 2.0 * f0 + fp + if denom >= 0.0: + return 0.0 + return float(min(0.5, max(-0.5, 0.5 * (fm - fp) / denom))) - pos_3d = np.unravel_index(idx, correlation_map.shape) - trans = np.array([pos_3d[0] / Nx, pos_3d[1] / Ny, pos_3d[2] / Nz]) - score = flat[idx] - sigma = (score - mean_val) / std_val - # Check clustering - skip if too close to existing peak +def _find_peaks( + score: torch.Tensor, + n_peaks: int, + radii_frac: Tuple[float, float, float], +) -> List[TranslationPeak]: + """Greedy non-maximum suppression on the periodic map, then sub-grid refinement.""" + nx, ny, nz = score.shape + flat = score.reshape(-1) + mean = float(flat.mean()) + std = float(flat.std().clamp(min=1e-30)) + n_take = min(flat.numel(), max(50, 20 * n_peaks)) + vals, idx = torch.topk(flat, n_take) + vals = vals.cpu().numpy() + idx = idx.cpu().numpy() + grid = np.array([nx, ny, nz], dtype=np.float64) + radii = np.asarray(radii_frac, dtype=np.float64) + score_np = score.cpu().numpy() + + kept: List[TranslationPeak] = [] + kept_t: List[np.ndarray] = [] + for v, i in zip(vals, idx): + ijk = np.array(np.unravel_index(int(i), (nx, ny, nz)), dtype=np.int64) + t_grid = ijk / grid is_new = True - for prev in used: - diff = np.abs(trans - prev) - diff = np.minimum(diff, 1 - diff) # Periodic boundary - if np.linalg.norm(diff) < cluster_radius: + for prev in kept_t: + d = np.abs(t_grid - prev) + d = np.minimum(d, 1.0 - d) + if np.all(d < radii): is_new = False break - - if is_new: - peaks.append(TranslationPeak(trans, score, sigma)) - used.append(trans) - - return peaks + if not is_new: + continue + # Parabolic refinement along each axis from the periodic neighbours. + offs = np.zeros(3) + for d, n in enumerate((nx, ny, nz)): + lo = ijk.copy(); lo[d] = (ijk[d] - 1) % n + hi = ijk.copy(); hi[d] = (ijk[d] + 1) % n + offs[d] = _parabolic_offset( + float(score_np[tuple(lo)]), float(v), float(score_np[tuple(hi)]), + ) + kept.append(TranslationPeak( + translation=(ijk + offs) / grid, score=float(v), + sigma=(float(v) - mean) / std, + )) + kept_t.append(t_grid) + if len(kept) >= n_peaks: + break + return kept -def amplitude_translation_search( +def fast_translation_function( obs: TranslationObs, - interpolator, - R_rotation: torch.Tensor, - spacegroup, + cand: CandidateTransform, real_cell, - grid_steps: int = 16, - n_peaks: int = 20, - cluster_radius: float = 0.05, - batch_size: int = 256, - precomputed_G: Optional[torch.Tensor] = None, - precomputed_h_R: Optional[torch.Tensor] = None, -) -> Tuple[np.ndarray, np.ndarray, List[TranslationPeak]]: - """ - Coarse-grid translation search via |F|²-correlation. - - For each candidate fractional translation `t` on a `grid_steps`³ grid in - `[0, 1)³`, scores the model at the current rotation translated by `t` - against the observed amplitudes by Pearson correlation of `|F_obs|²` and - `|F_calc(h, t)|²`. The structure-factor sum uses the spacegroup symmetry - expansion - - F_calc(h, t) = Σ_i G_i(h) · exp(2πi (h R_i) · t) - G_i(h) = exp(2πi h · t_i) · F_p1(h R_i) - - with `F_p1(h R_i)` looked up via the supplied interpolator at the rotation - already applied to the model. The `G_i` factors are computed once; only the - phase exponential changes per candidate, so the scan is efficient. + *, + grid_spacing_A: float, + n_peaks: int = 3, + cluster_radius_A: float = 4.0, +) -> Tuple[torch.Tensor, List[TranslationPeak]]: + """The Crowther-Blow map of ``sum_h coeff(h) |E_calc(h, t)|^2`` and its peaks. + + ``coeff`` is :attr:`TranslationObs.coeff` and ``E_calc`` is normalised per + candidate by :func:`prepare_candidate`, so the map is the covariance of + two unit-mean intensities weighted by the model's expected reliability -- + the rotation function's own score equation, for translations. Expanding + ``|sum_i G_i exp(2 pi i (h R_i).t)|^2`` gives pair terms at frequency + ``h R_j - h R_i``; accumulating them onto a reciprocal grid and inverting + evaluates every grid translation in one FFT. Parameters ---------- - obs : TranslationObs - The observations, normalised and weighted once for the whole run. - interpolator : object - Anything providing ``evaluate(R, hkl, real_cell, return_amplitude=False)`` - -- in the pipeline, ``align._DirectModelEvaluator``. - R_rotation : torch.Tensor, shape (3, 3) - Rotation that has been applied to the model coordinates. - spacegroup : SpaceGroup - Provides `matrices` and `translations`. - real_cell : Cell - Real crystal cell. - grid_steps : int, default 16 - Per-axis grid resolution. Total candidates = grid_steps³. - n_peaks : int, default 20 - Number of peaks returned (after clustering). - cluster_radius : float, default 0.05 - Minimum fractional separation between returned peaks. - batch_size : int, default 256 - Number of candidate translations evaluated per inner batch. + grid_spacing_A : float + Target spacing of the translation grid along each axis. A third of the + translation set's resolution samples the peak densely enough for the + parabolic refinement to land within a fraction of a grid step. + n_peaks : int + How many distinct peaks to return, best first. + cluster_radius_A : float + Peaks closer than this (per axis, periodic) are one peak. Returns ------- - correlation_map : np.ndarray, shape (grid_steps, grid_steps, grid_steps) - Pearson correlation of |F_obs|² and |F_calc(t)|² at each grid point. - best_translation : np.ndarray, shape (3,) - Top-scoring fractional translation. + score : torch.Tensor + The ``(nx, ny, nz)`` map, fractional grid ``t = (i/nx, j/ny, k/nz)``. peaks : list of TranslationPeak - Top-`n_peaks` peaks sorted by descending correlation. """ device = get_default_device() - real_dtype = torch.float64 - complex_dtype = torch.complex128 - - hkl = obs.hkl - E_obs = obs.E_obs.to(device).to(real_dtype) - w = obs.weight.to(device).to(real_dtype) - - # Correlating E^2 rather than F^2 is what makes this robust to the - # resolution envelope: the model has no bulk solvent and the wrong overall - # B, and that mismatch is multiplicative, so subtracting a mean does not - # remove it but dividing by Sigma(s) does. - F_obs2 = E_obs * E_obs - # Centred at the WEIGHTED mean, which is what the weighted correlation - # below is a numerator for. With uniform weight this is the plain mean. - F_obs2_centered = F_obs2 - (w * F_obs2).sum() / w.sum().clamp(min=1e-30) - - # Pre-compute G_i(h) = exp(2πi h·t_i) · F_p1(h R_i) (or reuse caller's) - two_pi_i = 2j * torch.pi - if precomputed_G is not None and precomputed_h_R is not None: - G = precomputed_G.to(device).to(complex_dtype) - h_R = precomputed_h_R.to(device).to(real_dtype) - else: - G, h_R = precompute_G_for_rotation( - interpolator, R_rotation, hkl, spacegroup, real_cell, device=device, - ) + real = get_float_dtype() + cplx = get_complex_dtype() - # Crowther–Blow FFT translation function (Acta Cryst. B23 (1967) 544). - # The grid-evaluated score - # num(t) = Σ_h F_obs²_centered(h) · |F_calc(h, t)|² - # expands as - # num(t) = Σ_{i,j} [Σ_h F_obs²_c(h) · G_i*(h) · G_j(h)] - # · exp(2πi · (h·R_j − h·R_i) · t) - # and on a regular fractional t-grid t = (jx, jy, jz) / G this is exactly - # an inverse DFT of the bracketed coefficients accumulated onto a 3-D - # reciprocal grid at integer indices (h·R_j − h·R_i) mod G. - # - # We accumulate two such reciprocal grids in one sym-op pass: - # W_num : weight per h = w(h)·E_obs²_centered(h) → num(t) - # W_den : weight per h = w(h) → Σ_h w|F_calc(h,t)|² - # Score(t) = num(t) / Σ_h|F_calc(h,t)|² — a per-t scale-normalised - # Pearson proxy (Phaser's TF uses the full Pearson denominator; ours - # uses the same scaling that the previous separable-phase code applied - # via explicit per-t centering, achieved here without materialising - # |F_calc(h,t)|² per t-point). - # - # One IFFT pair replaces G³ grid evaluations — for our defaults this is - # ~5000× less arithmetic than the separable-phase scoring it supersedes, - # and orders of magnitude less than the original explicit grid loop. - S_eff, N_eff = G.shape - h_R_int = h_R.round().to(torch.int64) # (S, N, 3) - # Both grids carry the same per-reflection weight, so the ratio below is a - # weighted correlation rather than an unweighted one with a weighted - # numerator. Uniform w reproduces the previous scores exactly. - F_obs2_c_complex = (w * F_obs2_centered).to(complex_dtype) # (N,) - w_complex = w.to(complex_dtype) # (N,) - - W_num_flat = torch.zeros( - grid_steps ** 3, dtype=complex_dtype, device=device, - ) - W_den_flat = torch.zeros( - grid_steps ** 3, dtype=complex_dtype, device=device, - ) - G_stride_xy = grid_steps * grid_steps - for i in range(S_eff): - Gi_conj = G[i].conj() # (N,) - pair = Gi_conj.view(1, -1) * G # (S, N) - coeff_num = F_obs2_c_complex.view(1, -1) * pair # (S, N) - coeff_den = w_complex.view(1, -1) * pair # (S, N) - dh = (h_R_int - h_R_int[i:i + 1]) % grid_steps # (S, N, 3) - flat = (dh[..., 0] * G_stride_xy - + dh[..., 1] * grid_steps + dh[..., 2]) # (S, N) - flat_flat = flat.reshape(-1) - W_num_flat.index_add_(0, flat_flat, coeff_num.reshape(-1)) - W_den_flat.index_add_(0, flat_flat, coeff_den.reshape(-1)) - - W_num = W_num_flat.view(grid_steps, grid_steps, grid_steps) - W_den = W_den_flat.view(grid_steps, grid_steps, grid_steps) - # IFFT scales by 1/G³; undo so values are raw integrals. - num_t = (torch.fft.ifftn(W_num, dim=(0, 1, 2)).real - * (grid_steps ** 3)).to(real_dtype) - den_t = (torch.fft.ifftn(W_den, dim=(0, 1, 2)).real - * (grid_steps ** 3)).to(real_dtype) - corr_map = num_t / den_t.clamp(min=1e-30) - corr_map_np = corr_map.detach().cpu().numpy().astype(np.float32) - peaks = find_translation_peaks(corr_map_np, n_peaks=n_peaks, - cluster_radius=cluster_radius) - best = peaks[0].translation if peaks else np.zeros(3) - return corr_map_np, best, peaks - - -def normalise_calc(F_calc: torch.Tensor, obs: TranslationObs) -> torch.Tensor: - """``E_calc`` for one or many candidate translations, through the shared fit. - - Accepts ``(N,)`` or ``(K, N)`` and returns the same shape. Each candidate is - fitted separately, because the resolution envelope of ``|F_calc(h, t)|`` is a - property of that placement -- but by the *same* estimator the observed side - uses, on the same abscissa, so the two sides of the likelihood are normalised - by one rule rather than two. - - This used to be a per-shell mean, written out twice: once here for the K - candidates and once in the pipeline for the top peak that ``sigma_A`` is - fitted against. Two copies of one calculation is how they drift, and neither - was the estimator anything else in the package used. The difference is not - cosmetic -- the median per-reflection change is 2-4%. - - The fit converges in single-digit iterations, so the cost is a few - milliseconds per candidate against a placement of order a second, and it is - only paid when the likelihood rescore is on. - """ - single = F_calc.ndim == 1 - F = F_calc.reshape(1, -1) if single else F_calc - s_lo, s_hi = float(obs.s_mag.min()), float(obs.s_mag.max()) - out = torch.empty_like(F) - for k in range(F.shape[0]): - # No eps and nothing centric: a single molecular transform sampled at - # these indices carries no crystal multiplicity, and the observed side - # gets its own from `obs`. - out[k] = WilsonNormaliser( - F[k] * F[k], obs.s_mag, n_coeff=WILSON_N_COEFF, - s_lo=s_lo, s_hi=s_hi, - ).E.to(F.dtype) - return out[0] if single else out - - -def fit_model_error(obs: TranslationObs, E_calc: torch.Tensor, *, shrink: bool = True): - """Per-reflection ``(alpha, beta)`` for a placed model, from the shared estimator. - - Returns the two quantities the likelihood actually wants: ``alpha``, the - multiplier on the calculated amplitude, and ``beta``, the conditional - variance. They are what - :class:`~torchref.refinement.model_error_estimation.sigma_a.SigmaAEstimator` - produces, and using them rather than ``(D, 1 - D^2)`` drops the assumption - that ```` is exactly 1 -- ``alpha = sigma_A sqrt(Sigma_N/Sigma_P)`` - carries the mismatch that assumption hides. - - **This replaces a local 81-point scan over every reflection.** The shared - estimator runs three nested-zoom stages of seventeen candidates instead, on a - cancellation-folded Rice that survives float32, and it shrinks each shell - toward the fitted curve rather than taking a per-shell argmax at face value. - - It is reached through ``SigmaAEstimator``, not the ``estimate_beta`` free - function underneath it, and the reason is not the cache. The wrapper - interpolates **four** shell curves -- ``sigma_A``, ``log Sigma_N``, - ``log Sigma_P``, ``S2`` -- and derives ``alpha``/``beta`` per reflection from - them, so the second-moment identity holds at every reflection. Its docstring - warns that interpolating ``beta`` directly "can yield a value consistent with - no ``sigma_A <= 1`` at all", which is exactly what a hand-rolled - interpolation here would have done. - - ``epsilon`` is passed as ones because ``obs.E_obs`` is already - epsilon-reduced by :class:`~torchref.scaling.WilsonNormaliser`; applying it - again would count multiplicity twice. ``free_mask`` is all-ones: there is no - cross-validation set to protect at placement time, and the estimate is not - being used to decide when to stop refining. - """ - # Local import: `torchref.refinement.__init__` eagerly pulls the refinement - # drivers and every target, which is a heavy load for one estimator. The - # same pattern, for the same reason, is documented at - # `scaling/scaler_base.py` -- "Do not 'tidy' them up". - from torchref.refinement.model_error_estimation.sigma_a import SigmaAEstimator - - ones = torch.ones_like(obs.E_obs) - est = SigmaAEstimator().get( - F_obs=obs.E_obs, - F_calc_scaled=E_calc.to(obs.E_obs.dtype), - centric=obs.centric, - epsilon=ones, - d_star_sq=(obs.s_mag * obs.s_mag).to(obs.E_obs.dtype), - free_mask=torch.ones_like(obs.E_obs, dtype=torch.bool), - shrink=shrink, - ) - return est.alpha, est.beta - - -def llg_translation_rescore( - obs: TranslationObs, - G: torch.Tensor, - h_R: torch.Tensor, - t_candidates: torch.Tensor, - alpha: torch.Tensor, - beta: torch.Tensor, -) -> torch.Tensor: - """Per-translation Rice / Woolfson log-likelihood over candidate positions. - - For each candidate t:: - - F_calc(h, t) = sum_i G_i(h) exp(2 pi i (h R_i).t) - E_calc(h, t) = |F_calc(h, t)| / sqrt(Sigma_calc(s; t)) - LLG(t) = sum_h [LL(E_obs, D E_calc, Sigma) - LL_Wilson(E_obs)] - - with ``Sigma = (1 - D^2)`` the **complex** variance. The acentric/centric - split is handled inside - :func:`~torchref.base.targets.xray_likelihoods.rice_per_refl`, which derives - the centric amplitude variance from the same ``Sigma`` -- the two are not - the same number, and passing one amplitude variance to both branches is how - this used to score acentrics at twice the variance they should have. - - The scoring rule the amplitude correlation is a pre-filter for. At rank - level it is the strongest discriminator measured -- truth at rank 0 in 27 of - 30 seeded cells against the correlation's 24 and the rotation function's 6 -- - but re-ranking translation peaks by it does **not** improve end-to-end pose - recovery (28/30 against 27/30 the other way, one discordant cell), which is - why ``use_llg_tf`` defaults off. - - ``E_calc`` is normalised **per candidate**, by :func:`normalise_calc` and so - by the same Wilson fit as the observed side. Per-candidate rather than once is - deliberate: the resolution envelope of ``|F_calc(h, t)|`` belongs to that - placement, and normalising every candidate to `` = 1`` is what makes - the K likelihoods comparable. What discriminates is the pattern across - reflections, not the scale. - - Parameters - ---------- - obs : TranslationObs - Supplies ``E_obs`` and centricity. - G : (S, N) complex - Per-sym ``F_p1`` contributions x per-sym translation phase, from - :func:`precompute_G_for_rotation`. - h_R : (S, N, 3) - Per-sym rotated reciprocal indices. - t_candidates : (K, 3) - Fractional translations to score. - alpha, beta : (N,) - Model reliability and conditional variance per reflection, from - :func:`fit_model_error`. Fixed across candidates on purpose: refitting - per candidate would score each against a likelihood tuned to itself. - - Returns - ------- - llg : (K,) torch.Tensor — log-likelihood gain per candidate. - """ - device = G.device - real_dtype = torch.float64 - complex_dtype = G.dtype - - centric = obs.centric - - K = t_candidates.shape[0] + nx, ny, nz = _grid_sizes(real_cell, grid_spacing_A) + G = cand.G.to(device).to(cplx) S, N = G.shape + coeff = obs.coeff.to(device).to(cplx) + h_R_int = cand.h_R.round().to(torch.int64) - t_cand = t_candidates.to(device).to(real_dtype) # (K, 3) - # Phase factor for each (k, i, n): exp(2πi · (h_R[i, n] · t[k])) - phase_arg = torch.einsum("ind,kd->kin", h_R.to(real_dtype), t_cand) - phase = torch.exp(2j * torch.pi * phase_arg.to(complex_dtype)) # (K, S, N) - # F_calc(k, n) = Σ_i G[i, n] · phase[k, i, n] - Fc_complex = (G.view(1, S, N) * phase).sum(dim=1) # (K, N) - F_calc = Fc_complex.abs().to(real_dtype) # (K, N) - - E_calc = normalise_calc(F_calc, obs) # (K, N) - E_obs = obs.E_obs.to(device).to(real_dtype) - - # alpha and beta arrive PER REFLECTION, interpolated from the shell fit by - # the shared estimator. No `index_select` on a shell index: the estimator - # bins on its own abscissa, and taking its per-reflection output is what - # keeps the second-moment identity holding at every reflection rather than - # only per shell. - a_r = alpha.to(device).to(real_dtype) # (N,) - Sigma_r = beta.to(device).to(real_dtype).clamp(min=1e-4) # (N,) - - F_mean = a_r.view(1, N) * E_calc # (K, N) - Sigma_full = Sigma_r.view(1, N).expand(K, N) - E_obs_full = E_obs.view(1, N).expand(K, N) - cent = centric.to(device).to(torch.bool) - cent_full = cent.view(1, N).expand(K, N) - - ll = -rice_per_refl(E_obs_full, F_mean, Sigma_full, cent_full) # (K, N) - - # Wilson reference (data only): no model, so F_mean = 0 and Sigma = 1 -- - # which is = 1, the identity WilsonNormaliser fits to. Under the - # amplitude-variance convention this line used to carry, unit Sigma meant - # = 2 for acentrics and the reference was inconsistent with the data - # it referenced. - ll_wil_per_refl = -rice_per_refl( - E_obs, torch.zeros_like(E_obs), torch.ones_like(E_obs), cent) - ll_wil_total = ll_wil_per_refl.sum() - - return ll.sum(dim=1) - ll_wil_total # (K,) - - -def llg_at( - obs: TranslationObs, - G: torch.Tensor, - h_R: torch.Tensor, - t: torch.Tensor, - alpha: torch.Tensor, - beta: torch.Tensor, -) -> float: - """The translation likelihood at one translation, as a candidate score. - - :func:`llg_translation_rescore` over a single ``t``. Split out because - scoring a *candidate* and re-ranking a candidate's *translations* are - different questions that happen to share a functional, and only the first - needs to be comparable across orientations. - """ - return float(llg_translation_rescore( - obs=obs, G=G, h_R=h_R, - t_candidates=t.detach().reshape(1, 3).to(G.device).to(torch.float64), - alpha=alpha, beta=beta, - )[0]) - - -def correlation_at( - obs: TranslationObs, G: torch.Tensor, h_R: torch.Tensor, t: torch.Tensor, -) -> float: - """The translation search's own score at one translation. - - Same functional :func:`amplitude_translation_search` maximises -- - ``sum_h w (E_obs^2 - _w) |F_calc(h, t)|^2 / sum_h w |F_calc(h, t)|^2`` - -- evaluated at a single ``t`` rather than over a grid, which needs no FFT. - - The grid search returns its peaks' scores, but a peak is then refined and the - refined position is what gets used. Scoring the *used* translation is what - makes the number comparable with other candidates' used translations. It is - also what stops the ranking key and the returned placement coming from - different points, which is how a selection rule quietly stops meaning what - its name says. - """ - device = G.device - E = obs.E_obs.to(device).to(torch.float64) - w = obs.weight.to(device).to(torch.float64) - E2 = E * E - E2c = E2 - (w * E2).sum() / w.sum().clamp(min=1e-30) - - tt = t.detach().to(device).to(torch.float64).reshape(3) - phase = torch.exp(2j * torch.pi * torch.einsum( - "ind,d->in", h_R.to(torch.float64), tt).to(G.dtype)) - Fc2 = (G * phase).sum(dim=0).abs().to(torch.float64) ** 2 - - num = (w * E2c * Fc2).sum() - den = (w * Fc2).sum().clamp(min=1e-30) - return float(num / den) - - -def precompute_G_for_rotation( - interpolator, - R_rotation: torch.Tensor, - hkl: torch.Tensor, - spacegroup, - real_cell, - device=None, -): - """ - Pre-compute per-symmetry F_asu contributions `G_i(h)` for a fixed rotation. + W = torch.zeros(nx * ny * nz, dtype=cplx, device=device) + for i in range(S): + pair = G[i].conj().view(1, -1) * G # (S, N) + dh = h_R_int - h_R_int[i:i + 1] # (S, N, 3) + flat = ((dh[..., 0] % nx) * ny + (dh[..., 1] % ny)) * nz + (dh[..., 2] % nz) + W.index_add_(0, flat.reshape(-1), (coeff.view(1, -1) * pair).reshape(-1)) + score = (torch.fft.ifftn(W.view(nx, ny, nz), dim=(0, 1, 2)).real + * float(nx * ny * nz)).to(real) - These are the only inputs that depend on `R_rotation` (and therefore on - expensive interpolator/model-forward evaluations). Passing the result - into `amplitude_translation_search` and `local_translation_refine` lets - them share a single set of (n_sym) model evaluations across the coarse - TF and the fine refinement, instead of recomputing each call. + radii = tuple(float(cluster_radius_A) / float(L) + for L in (real_cell.a, real_cell.b, real_cell.c)) + peaks = _find_peaks(score, n_peaks, radii) + return score, peaks - Returns - ------- - G : (S, N) complex128 - h_R : (S, N, 3) float64 - """ - real_dtype = torch.float64 - complex_dtype = torch.complex128 - if device is None: - device = get_default_device() - - hkl_t = hkl.detach().to(device).to(real_dtype) - sym_R = spacegroup.matrices.detach().to(device).to(real_dtype) - sym_t = spacegroup.translations.detach().to(device).to(real_dtype) - S = sym_R.shape[0] - N = hkl_t.shape[0] - R_rot = R_rotation.detach().to(device).to(real_dtype) - - # Batched: h_R[i, n, d] = Σ_e hkl[n, e] · sym_R[i, e, d] - h_R = torch.einsum("ne,ied->ind", hkl_t, sym_R) # (S, N, 3) - # Per-sym-op translation phase: exp(2πi · h · t_i) - two_pi_i = 2j * torch.pi - phase_arg = torch.einsum("ne,ie->in", hkl_t, sym_t) # (S, N) - phase = torch.exp(two_pi_i * phase_arg.to(complex_dtype)) # (S, N) - - # One interpolator.evaluate over all (S × N) rotated indices: lets the - # backend do a single grid_sample instead of S sequential ones. - h_R_flat = h_R.reshape(-1, 3) # (S·N, 3) - F_flat = interpolator.evaluate( - R_rot, h_R_flat, real_cell, return_amplitude=False, - ) - F_all = F_flat.reshape(S, N).to(complex_dtype) # (S, N) - G = F_all * phase - return G, h_R +def translation_score_at(obs: TranslationObs, cand: CandidateTransform, + t: torch.Tensor) -> float: + """The fast search's score at one translation, without the FFT.""" + E2 = cand.e_calc(t) ** 2 + return float((obs.coeff.to(E2.device).to(E2.dtype) * E2).sum()) -def local_translation_refine( +def llg_at_translations( obs: TranslationObs, - interpolator, - R_rotation: torch.Tensor, - spacegroup, - real_cell, - t_init: torch.Tensor, - radius: float = 0.06, - grid_steps: int = 13, - n_refinement_passes: int = 2, - batch_size: int = 1024, - precomputed_G: Optional[torch.Tensor] = None, - precomputed_h_R: Optional[torch.Tensor] = None, -) -> Tuple[torch.Tensor, float]: - """ - Fine-grid Patterson translation refinement around ``t_init``. - - Locates the peak on a fine grid of half-width ``radius`` around ``t_init`` - by the same weighted ``E^2`` correlation the coarse search maximises, then - reports the analytical-scale R-factor there:: - - R(t) = sum ||F_obs| - k|F_calc(t)|| / sum |F_obs| - k(t) = sum |F_obs||F_calc(t)| / sum |F_calc(t)|^2 - - The two halves answer different questions and use different quantities on - purpose. The *search* runs on normalised, weighted ``E^2``, because that is - what the coarse stage optimised and refining against a different objective - would walk away from the peak it was handed. The *reported number* is an - R-factor on raw amplitudes, because that is what ranks candidates and what a - crystallographer reads. + cand: CandidateTransform, + t_candidates: torch.Tensor, +) -> torch.Tensor: + """Rice/Woolfson log-likelihood gain at each of ``K`` translations. - That R uses one global scale, so it is not the number a full Scaler returns. - It is used as a ranking key, and for that its minimum's *location* is what - matters. The winner gets a solvent-aware Scaler refit. + ``LLG(t) = sum_h [LL(E_obs; sigma_A E_calc(h, t), 1 - sigma_A^2) + - LL(E_obs; 0, 1)]`` with the complex-variance convention of + :func:`~torchref.base.targets.xray_likelihoods.rice_per_refl`, which + derives the centric case from the same ``Sigma``. ``sigma_A`` is the + Luzzati prior carried by ``obs`` -- the same for every candidate, so the + values are comparable across orientations as well as across translations, + and no candidate is scored against a likelihood tuned to itself. - Returns ``(t, R)`` at the minimum. ``n_refinement_passes`` is accepted and - ignored -- the FFT evaluates the whole fine grid at once, so there is - nothing for a second zoom pass to buy. + Returns ``(K,)``. """ - device = get_default_device() - real_dtype = torch.float64 - complex_dtype = torch.complex128 - - hkl = obs.hkl - F_obs_t = obs.F_obs.to(device).to(real_dtype) - F_obs_sum = F_obs_t.sum().clamp(min=1e-30) - E_obs = obs.E_obs.to(device).to(real_dtype) - w = obs.weight.to(device).to(real_dtype) - - two_pi_i = 2j * torch.pi - if precomputed_G is not None and precomputed_h_R is not None: - G = precomputed_G.to(device).to(complex_dtype) - h_R = precomputed_h_R.to(device).to(real_dtype) - else: - G, h_R = precompute_G_for_rotation( - interpolator, R_rotation, hkl, spacegroup, real_cell, device=device, - ) - - # Adapt batch_size to keep the largest inner einsum tensor under ~250 MB - # of complex128. The (S, B, N) phase tensor is the offender: - # S × B × N × 16 bytes ≤ 2.5e8 → B ≤ 2.5e8 / (16 × S × N). - S_eff, N_eff = G.shape - safe_b = max(8, int(2.5e8 / (16.0 * max(S_eff, 1) * max(N_eff, 1)))) - batch_size = min(batch_size, safe_b) - - # Crowther–Blow FFT refinement on a fine grid around t_init. - # Bake t_init into G as a per-h_R phase factor, then the IFFT trick from - # `amplitude_translation_search` works on the offset grid Δt with the - # same (num/den) Pearson-proxy scoring. The previous nested-grid Python - # evaluation paid O(G³ · N · S) Bessel/exp/einsum per call; this pays - # one IFFT pair on G_fft³ + an O(N · S) F_calc evaluation at the final t. - S, N = G.shape - t_init_t = torch.as_tensor(t_init, dtype=real_dtype, device=device) - h_R_int = h_R.round().to(torch.int64) # (S, N, 3) - - # Bake t_init into G: - phase_init = torch.exp( - two_pi_i * torch.einsum("snd,d->sn", h_R, t_init_t).to(complex_dtype) - ) # (S, N) - G_shifted = G * phase_init # (S, N) - - # Pick G_fft so the IFFT spacing matches the requested fine grid: - # spacing = 2·radius / (grid_steps − 1), G_fft = round(1 / spacing). - # Cap at 128 to bound memory (128³ complex128 ≈ 32 MB). - desired_spacing = max(2.0 * float(radius) / max(grid_steps - 1, 1), 1e-6) - G_fft = max(grid_steps, int(round(1.0 / desired_spacing))) - G_fft = min(G_fft, 128) - half_window = max(1, int(round(float(radius) * G_fft))) - - # The same weighted E^2 correlation as the coarse search. It used to be a - # raw |F|^2 correlation here, which made the fine grid optimise a different - # objective from the one that chose the peak it is centred on. - F_obs2 = E_obs * E_obs - F_obs2_centered = F_obs2 - (w * F_obs2).sum() / w.sum().clamp(min=1e-30) - F_obs2_c_complex = (w * F_obs2_centered).to(complex_dtype) - w_complex = w.to(complex_dtype) - - W_num_flat = torch.zeros(G_fft ** 3, dtype=complex_dtype, device=device) - W_den_flat = torch.zeros(G_fft ** 3, dtype=complex_dtype, device=device) - G_stride_xy = G_fft * G_fft - for i in range(S): - Gi_conj = G_shifted[i].conj() - pair = Gi_conj.view(1, -1) * G_shifted # (S, N) - coeff_num = F_obs2_c_complex.view(1, -1) * pair - coeff_den = w_complex.view(1, -1) * pair - dh = (h_R_int - h_R_int[i:i + 1]) % G_fft # (S, N, 3) - flat = (dh[..., 0] * G_stride_xy - + dh[..., 1] * G_fft + dh[..., 2]) # (S, N) - flat_flat = flat.reshape(-1) - W_num_flat.index_add_(0, flat_flat, coeff_num.reshape(-1)) - W_den_flat.index_add_(0, flat_flat, coeff_den.reshape(-1)) - - W_num = W_num_flat.view(G_fft, G_fft, G_fft) - W_den = W_den_flat.view(G_fft, G_fft, G_fft) - num_t = torch.fft.ifftn(W_num, dim=(0, 1, 2)).real * (G_fft ** 3) - den_t = torch.fft.ifftn(W_den, dim=(0, 1, 2)).real * (G_fft ** 3) - score = num_t / den_t.clamp(min=1e-30) - - # Roll so the (Δt = 0) cell sits in the centre of a (2·half_window+1) - # window, then look for the maximum within the radius-sphere. - score_rolled = torch.roll( - score, shifts=(half_window, half_window, half_window), dims=(0, 1, 2), - ) - w = 2 * half_window + 1 - score_window = score_rolled[:w, :w, :w] - idx_flat = int(score_window.argmax().item()) - jx = idx_flat // (w * w) - rem = idx_flat % (w * w) - jy = rem // w - jz = rem % w - Delta_t = torch.tensor( - [(jx - half_window) / G_fft, - (jy - half_window) / G_fft, - (jz - half_window) / G_fft], - dtype=real_dtype, device=device, - ) - best_t = t_init_t + Delta_t - - # Compute the analytical-scale R-factor at best_t (one t evaluation), - # which is what the caller uses to rank rotation × translation - # candidates. The local-refine grid search above optimised the - # FFT-scored Pearson proxy; analytical R is monotonically related on - # this neighbourhood so the choice of which fine-grid maximum to - # commit to is preserved. - phase_best = torch.exp( - two_pi_i * torch.einsum("snd,d->sn", h_R, best_t).to(complex_dtype) - ) # (S, N) - F_calc_best = (G * phase_best).sum(dim=0) # (N,) complex - F_c_abs = F_calc_best.abs().to(real_dtype) - num_a = (F_obs_t * F_c_abs).sum() - den_a = (F_c_abs ** 2).sum().clamp(min=1e-30) - k = num_a / den_a - best_R = float( - ((F_obs_t - k * F_c_abs).abs().sum() / F_obs_sum).item() - ) - # Unused: `n_refinement_passes`, `batch_size` kept in signature for - # back-compat with callers passing them. - _ = n_refinement_passes - _ = batch_size - - return best_t.cpu(), best_R - - + E_calc = cand.e_calc(t_candidates) # (K, N) + K, N = E_calc.shape + dev, real = E_calc.device, E_calc.dtype + E_obs = obs.E_obs.to(dev).to(real).view(1, N).expand(K, N) + D = obs.sigma_a.to(dev).to(real).view(1, N) + Sigma = (1.0 - D * D).clamp(min=1e-3).expand(K, N) + cent = obs.centric.to(dev).view(1, N).expand(K, N) + ll = -rice_per_refl(E_obs, D * E_calc, Sigma, cent) # (K, N) + ll_wil = -rice_per_refl( + E_obs[0], torch.zeros(N, dtype=real, device=dev), + torch.ones(N, dtype=real, device=dev), cent[0], + ).sum() + return ll.sum(dim=1) - ll_wil + + +def analytic_r_at(obs: TranslationObs, cand: CandidateTransform, + t: torch.Tensor) -> float: + """``R = sum ||F_obs| - k |F_calc(t)|| / sum |F_obs|`` with one global scale. + + On raw amplitudes, because that is what a crystallographer reads; not the + number a full Scaler would return, since there is no bulk solvent and no + B-factor scaling behind ``k``. + """ + F_c = cand.e_calc(t) * cand.norm.to(cand.G.device) + F_o = obs.F_obs.to(F_c.device).to(F_c.dtype) + k = (F_o * F_c).sum() / (F_c * F_c).sum().clamp(min=1e-30) + return float((F_o - k * F_c).abs().sum() / F_o.sum().clamp(min=1e-30)) From 00854577315e86a8a94cc2429a596a73f8d6c00d Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 2 Sep 2026 12:21:00 +0200 Subject: [PATCH 144/250] Make the empirical sigma_A read the shape of the Wilson-curve ratio, not its level empirical_sigma_a divides the observed Wilson curve by the calculated one and takes sqrt(min(R, 1/R)). The observed curve is on the data's arbitrary scale and the calculated one on the model's electron scale, and nothing removed that factor, so the ratio's level set the answer: measured 0.02-0.06 on 1DAW and 2DQ6 and 8-12 on 3K7M, giving a flat sigma_A of 0.15-0.35 with no resolution dependence -- not the resolution-dependent model deficiency the docstring describes. Each curve is now divided by its geometric mean over the supplied points before the ratio is taken. Its only caller is the rotation function's calc-side weight, where a per-shell factor is gauge in the correlation, so placements do not move: 30/30 at the default window and 30/30 uncut (job 544925). 191 tests pass across alignment, scaling and the alignment integration suite (544924). Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- .../analysis/integration_alignment.sh | 2 +- docs/changelog.rst | 1 + tests/unit/scaling/test_weighting.py | 32 +++++++++++++++++-- torchref/scaling/weighting.py | 13 ++++++-- 4 files changed, 43 insertions(+), 5 deletions(-) diff --git a/alignment_lab/analysis/integration_alignment.sh b/alignment_lab/analysis/integration_alignment.sh index bee675da..d2c54e19 100644 --- a/alignment_lab/analysis/integration_alignment.sh +++ b/alignment_lab/analysis/integration_alignment.sh @@ -13,5 +13,5 @@ PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin cd "$REPO" export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -"$PY" -m pytest -c tests/pytest.ini tests/integration/alignment tests/unit/alignment --run-slow -q --tb=short 2>&1 | grep -v "^✓" | tail -60 +"$PY" -m pytest -c tests/pytest.ini tests/integration/alignment tests/unit/alignment tests/unit/scaling --run-slow -q --tb=short 2>&1 | grep -v "^✓" | tail -60 echo "PYTEST_RC=${PIPESTATUS[0]}" diff --git a/docs/changelog.rst b/docs/changelog.rst index bddbb883..40369f9e 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Fixed ``empirical_sigma_a`` taking the level of the observed-to-calculated Wilson-curve ratio rather than its shape. The two curves carry different absolute scales, so the ratio was 0.02-0.06 on one structure and 8-12 on another and the returned ``sigma_A`` was flat at 0.15-0.35 regardless of resolution; each curve is now divided by its geometric mean first. Per-shell factors are gauge in the rotation function's correlation, so its placements are unchanged (30/30) - The fast translation function scores the covariance of two normalised intensities -- the rotation function's LERF1 coefficient ``cw (E_obs^2 - 1) w sigma_A^2`` against the candidate's ``|E_calc(h, t)|^2``, normalised per candidate by the same Wilson fit -- instead of a raw-``|F_calc|^2`` ratio that was not a correlation and, on the four largest panel structures, was higher 40 A from the true position than at it. One FFT on a grid a third of the set's resolution apart with parabolic peak refinement replaces the 16-point coarse grid and three 100-point local refines; the Rice/Woolfson likelihood at a fixed Luzzati ``sigma_A`` picks among the top peaks and ranks the candidates. 30/30 true poses at the default window and 30/30 with the window removed, against 18/30 before - The translation stage runs in the configured float and complex dtypes rather than hard-coded double - Removed ``use_llg_tf``, ``n_translation_peaks`` and ``translation_grid_steps`` from the pipeline; the likelihood always picks the translation, and the grid is sized by resolution diff --git a/tests/unit/scaling/test_weighting.py b/tests/unit/scaling/test_weighting.py index 01f1fd01..6bd49cd2 100644 --- a/tests/unit/scaling/test_weighting.py +++ b/tests/unit/scaling/test_weighting.py @@ -9,8 +9,8 @@ import torch from torchref.scaling.weighting import ( - DEFAULT_SNR_CAP, information_weight, inverse_variance_weight, - normalise_weight, snr_from_amplitude, + DEFAULT_SNR_CAP, empirical_sigma_a, information_weight, + inverse_variance_weight, normalise_weight, snr_from_amplitude, ) pytestmark = pytest.mark.unit @@ -122,3 +122,31 @@ def test_normalise_weight_sets_the_mean(): def test_a_constant_weight_survives_normalisation_as_ones(): w = torch.full((100,), 3.7, dtype=torch.float64) assert torch.allclose(normalise_weight(w), torch.ones_like(w)) + + +def test_empirical_sigma_a_ignores_the_absolute_scale(): + """The data's scale is arbitrary and the model's is electrons; neither is information. + + Before this held, the ratio's level set the answer: a flat sigma_A of 0.2 on + one structure and 0.33 on another, with no resolution dependence at all. + """ + s = torch.linspace(0.07, 0.25, 40, dtype=torch.float64) + obs = torch.exp(-30.0 * s * s) + calc = torch.exp(-30.0 * s * s) + base = empirical_sigma_a(obs, calc) + torch.testing.assert_close(empirical_sigma_a(1000.0 * obs, calc), base) + torch.testing.assert_close(empirical_sigma_a(obs, 1e-3 * calc), base) + + +def test_empirical_sigma_a_reads_the_shape(): + """Identical shapes: full trust everywhere. A low-resolution deficit: less trust there.""" + s = torch.linspace(0.07, 0.25, 40, dtype=torch.float64) + calc = torch.exp(-30.0 * s * s) + same = empirical_sigma_a(7.0 * calc, calc) + assert float(same.min()) > 0.999 + # The model predicts more scattering at low resolution than the data have + # -- the bulk solvent it lacks -- and matches at high resolution. + deficit = torch.where(s < 0.12, torch.full_like(s, 0.4), torch.ones_like(s)) + sa = empirical_sigma_a(calc * deficit, calc) + assert float(sa[s < 0.12].max()) < float(sa[s > 0.15].min()) + assert bool((sa <= 1.0).all()) and bool((sa > 0.0).all()) diff --git a/torchref/scaling/weighting.py b/torchref/scaling/weighting.py index d874cd86..febe6a5f 100644 --- a/torchref/scaling/weighting.py +++ b/torchref/scaling/weighting.py @@ -203,6 +203,13 @@ def empirical_sigma_a( factor, so a model that is half the asymmetric unit looks like a model that is all of it, and only the tilt survives. + That uniform factor is removed here, by dividing each curve by its geometric + mean over the points supplied. The two fits carry their own absolute + scales -- the data's arbitrary one and the model's electron scale -- and + without this step the ratio's *level* set the answer rather than its shape: + measured 0.02-0.06 on 1DAW and 2DQ6 and 8-12 on 3K7M, giving a flat + ``sigma_A`` of 0.15-0.35 that said nothing about resolution. + Parameters ---------- sigma_obs, sigma_calc : torch.Tensor @@ -212,6 +219,8 @@ def empirical_sigma_a( floor : float, optional Lower bound on the returned ``sigma_A``. """ - r = (sigma_obs / sigma_calc.clamp(min=1e-30)).clamp(min=1e-30) - shared = torch.minimum(r, 1.0 / r).clamp(min=0.0, max=1.0) + log_r = (sigma_obs.clamp(min=1e-30).log() + - sigma_calc.clamp(min=1e-30).log()) + log_r = log_r - log_r.mean() # unit geometric mean: scale-free + shared = torch.exp(-log_r.abs()) # min(R, 1/R) return shared.sqrt().clamp(min=float(floor), max=1.0 - 1e-6) From eea62f8bb57ae4cd7b03cc5e78d17bce3ef24a78 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 2 Sep 2026 12:29:38 +0200 Subject: [PATCH 145/250] Suppress symmetry mates in the rotation function's peak list The greedy SO(3) suppression treated an orientation and its point-group mates as different peaks, so the shortlist handed to the translation search was mostly copies: 187 of the 300 pairs among 3K7M's top 25 peaks were mates of each other, 38 on 3GR5, 25 on 2DQ6. With the Cartesian symmetry rotations supplied, a kept peak suppresses its whole orbit and every returned peak is a distinct orientation. The group composes on the right, R R_g. Measured rather than assumed (alignment_lab/diagnostics/frf_orbit_side.py): on real peak lists the left orbit finds zero coincident pairs on every structure tried and the right orbit finds every mate. The lab's symmetry_orbit and orbit_rank defaulted to the left side, which is why the orbit-based truth rank disagreed with coordinate superposition; the default and the five harnesses that pinned the left side now use the right. 192 tests pass (job 544938). Placements unchanged: 30/30 at the default window and 30/30 uncut (job 544939). Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/frf_orbit_side.sh | 19 ++++++ alignment_lab/analysis/marginal_seeds.sh | 6 +- .../diagnostics/frf_aniso_rank_sweep.py | 2 +- alignment_lab/diagnostics/frf_benchmark.py | 2 +- alignment_lab/diagnostics/frf_config_sweep.py | 2 +- .../diagnostics/frf_ghost_knockout.py | 2 +- alignment_lab/diagnostics/frf_map_compare.py | 2 +- alignment_lab/diagnostics/frf_orbit_side.py | 67 +++++++++++++++++++ alignment_lab/lab/truth.py | 12 ++-- docs/changelog.rst | 1 + .../alignment/test_peak_finder_symmetry.py | 58 ++++++++++++++++ torchref/experimental/alignment/frf/api.py | 8 +++ .../experimental/alignment/frf/peak_finder.py | 38 ++++++++--- .../experimental/alignment/rotation_search.py | 13 +++- 14 files changed, 210 insertions(+), 22 deletions(-) create mode 100644 alignment_lab/analysis/frf_orbit_side.sh create mode 100644 alignment_lab/diagnostics/frf_orbit_side.py create mode 100644 tests/unit/alignment/test_peak_finder_symmetry.py diff --git a/alignment_lab/analysis/frf_orbit_side.sh b/alignment_lab/analysis/frf_orbit_side.sh new file mode 100644 index 00000000..a5d31fd3 --- /dev/null +++ b/alignment_lab/analysis/frf_orbit_side.sh @@ -0,0 +1,19 @@ +#!/bin/bash +#SBATCH --job-name=orbside +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:20:00 +#SBATCH --cpus-per-task=4 +#SBATCH --mem=32G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 +export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +for P in 2DQ6 3K7M 1DAW 3GR5; do + "$PY" -u alignment_lab/diagnostics/frf_orbit_side.py --pdb $P 2>&1 | grep -v "Warning\|warnings.warn" | grep -A14 "^#\|^ROW\|Traceback" +done +echo DONE diff --git a/alignment_lab/analysis/marginal_seeds.sh b/alignment_lab/analysis/marginal_seeds.sh index b9783d63..0d50d0ad 100644 --- a/alignment_lab/analysis/marginal_seeds.sh +++ b/alignment_lab/analysis/marginal_seeds.sh @@ -14,12 +14,12 @@ #SBATCH --cpus-per-task=4 #SBATCH --mem=48G #SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-3 +#SBATCH --array=0-5 set -uo pipefail REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -# Two known-marginal structures and two the panel has never lost, as controls. -PDBS=(2DQ6 6G9X 1DAW 3K7M) +# The four structures the translation search used to mis-place, and two controls. +PDBS=(2DQ6 6G9X 1DAW 3K7M 3VRJ 4BX9) PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} cd "$REPO" export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 diff --git a/alignment_lab/diagnostics/frf_aniso_rank_sweep.py b/alignment_lab/diagnostics/frf_aniso_rank_sweep.py index 75884c37..6de82df7 100644 --- a/alignment_lab/diagnostics/frf_aniso_rank_sweep.py +++ b/alignment_lab/diagnostics/frf_aniso_rank_sweep.py @@ -58,7 +58,7 @@ def run_one(pdb: str, trial: int, arm: str, cfg: FRFConfig, rank, ang = orbit_rank( res.peaks, R_true, data.spacegroup.matrices.to(torch.float64).cpu(), reciprocal_basis=data.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), - side="left", frame="cart", thr_deg=thr_deg, + side="right", frame="cart", thr_deg=thr_deg, ) row = {"experiment": EXPERIMENT, "pdb": pdb, "trial": trial, "arm": arm, "seed": seed} diff --git a/alignment_lab/diagnostics/frf_benchmark.py b/alignment_lab/diagnostics/frf_benchmark.py index 5ad3d3b1..8ab7bc08 100644 --- a/alignment_lab/diagnostics/frf_benchmark.py +++ b/alignment_lab/diagnostics/frf_benchmark.py @@ -113,7 +113,7 @@ def run_one(pdb: str, trial: int, arm: str, *, n_peaks: int, top_n: int, rank, ang = orbit_rank( res.peaks, R_true, data.spacegroup.matrices.to(torch.float64).cpu(), reciprocal_basis=data.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), - side="left", frame="cart", thr_deg=thr_deg, + side="right", frame="cart", thr_deg=thr_deg, ) # Exclusive, not inclusive: the nested stages would otherwise be counted # twice and "unattributed" could come out negative. diff --git a/alignment_lab/diagnostics/frf_config_sweep.py b/alignment_lab/diagnostics/frf_config_sweep.py index 6fa93bdd..9f73ac63 100644 --- a/alignment_lab/diagnostics/frf_config_sweep.py +++ b/alignment_lab/diagnostics/frf_config_sweep.py @@ -111,7 +111,7 @@ def run_one(pdb: str, trial: int, arm: Arm, base: FRFConfig, rank, ang = orbit_rank( peaks, R_true, data.spacegroup.matrices.to(torch.float64).cpu(), reciprocal_basis=data.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), - side="left", frame="cart", thr_deg=thr_deg, + side="right", frame="cart", thr_deg=thr_deg, ) row = {"experiment": EXPERIMENT, "pdb": pdb, "trial": trial, "arm": arm.name, "seed": seed} diff --git a/alignment_lab/diagnostics/frf_ghost_knockout.py b/alignment_lab/diagnostics/frf_ghost_knockout.py index ee70841f..63b812bd 100644 --- a/alignment_lab/diagnostics/frf_ghost_knockout.py +++ b/alignment_lab/diagnostics/frf_ghost_knockout.py @@ -87,7 +87,7 @@ def _orbit_of_identity(data): recip = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() return symmetry_orbit( torch.eye(3, dtype=torch.float64), symops, - side="left", frame="cart", reciprocal_basis=recip, + side="right", frame="cart", reciprocal_basis=recip, ) diff --git a/alignment_lab/diagnostics/frf_map_compare.py b/alignment_lab/diagnostics/frf_map_compare.py index 15513a4d..0bf40849 100644 --- a/alignment_lab/diagnostics/frf_map_compare.py +++ b/alignment_lab/diagnostics/frf_map_compare.py @@ -151,7 +151,7 @@ def _orbit_of_identity(data): symops = data.spacegroup.matrices.to(torch.float64).cpu() recip = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() I = torch.eye(3, dtype=torch.float64) - return symmetry_orbit(I, symops, side="left", frame="cart", + return symmetry_orbit(I, symops, side="right", frame="cart", reciprocal_basis=recip) diff --git a/alignment_lab/diagnostics/frf_orbit_side.py b/alignment_lab/diagnostics/frf_orbit_side.py new file mode 100644 index 00000000..3124c231 --- /dev/null +++ b/alignment_lab/diagnostics/frf_orbit_side.py @@ -0,0 +1,67 @@ +"""On which side does the crystal symmetry act on a rotation-function peak? + +The FRF returns Euler matrices ``R`` mapping the search-model frame onto the +crystal frame. Two peaks are the same orientation when they differ by a +point-group rotation, but that rotation can compose on the left +(``R2 = R_g R1``) or on the right (``R2 = R1 R_g``), and the two are different +sets for a non-commuting group. Counting how many of the top peaks of a real +search collapse onto each other under each convention settles which one the +engine's peaks obey: mates of the true orientation appear many times in the +list, so the right convention finds many near-zero pairs and the wrong one few. +""" +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from lab import BENCH_PDBS, cartesian_symops, load_case, random_rotation, seed_for # noqa: E402 + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--pdb", default="2DQ6", choices=list(BENCH_PDBS)) + ap.add_argument("--n", type=int, default=25) + ap.add_argument("--thr-deg", type=float, default=3.0) + args = ap.parse_args() + + from torchref.experimental.alignment import rotation_search + from torchref.experimental.alignment.sh import hkl_symops_to_cartesian + + model, data = load_case(args.pdb) + search = model.copy() + search = search.rotate(random_rotation(seed_for(args.pdb, 0)).to(model.dtype_float)) + sols = rotation_search(search, data, model_error_A=0.8, n_peaks=args.n) + R = sols.rotations.to(torch.float64) # (n, 3, 3) + n = R.shape[0] + + sym_lab = cartesian_symops(data.spacegroup, data.cell) # B S B^-1 + sym_sh = hkl_symops_to_cartesian( + data.spacegroup.matrices.to(torch.float64), + data.cell.reciprocal_basis_matrix.to(torch.float64)) + agree = float((sym_lab - sym_sh).abs().max()) + print(f"# {args.pdb} {data.spacegroup} n_peaks={n} |B S B^-1 - hkl_symops_to_cartesian|max={agree:.2e}") + + def pair_count(orbit_of): + cnt = 0 + for i in range(n): + for j in range(i + 1, n): + O = orbit_of(R[j]) # (g, 3, 3) + tr = torch.einsum("gab,ab->g", O, R[i]) + ang = ((tr - 1) / 2).clamp(-1, 1).arccos().min() * 180 / torch.pi + cnt += int(ang < args.thr_deg) + return cnt + + left = pair_count(lambda Rj: sym_lab @ Rj.unsqueeze(0)) + right = pair_count(lambda Rj: Rj.unsqueeze(0) @ sym_lab) + plain = pair_count(lambda Rj: Rj.unsqueeze(0)) + print(f"ROW pdb={args.pdb} pairs_within_{args.thr_deg:g}deg plain={plain} " + f"left(R_g R)={left} right(R R_g)={right}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/alignment_lab/lab/truth.py b/alignment_lab/lab/truth.py index 492d3318..9f83305a 100644 --- a/alignment_lab/lab/truth.py +++ b/alignment_lab/lab/truth.py @@ -16,7 +16,7 @@ from __future__ import annotations import math -from typing import Iterable, Optional, Sequence, Tuple +from typing import Optional, Sequence, Tuple import torch @@ -81,7 +81,7 @@ def symmetry_orbit( R_true: torch.Tensor, symops: torch.Tensor, *, - side: str = "left", + side: str = "right", frame: str = "cart", reciprocal_basis: Optional[torch.Tensor] = None, ) -> torch.Tensor: @@ -96,7 +96,11 @@ def symmetry_orbit( (fractional). side : {'left', 'right'}, optional ``'left'`` builds ``S_k @ R_true``; ``'right'`` builds ``R_true @ S_k``. - These are different sets for non-commuting operators. + These are different sets for non-commuting operators. The engine's + peaks obey ``'right'``: on real peak lists the left orbit finds zero + coincident pairs among the top 25 and the right orbit finds every mate + (187 of 300 pairs on 3K7M). ``'left'`` was the default, and is why the + orbit-based truth rank disagreed with coordinate superposition. frame : {'cart', 'frac'}, optional ``'cart'`` converts the operators to the Cartesian frame first, which is the frame the rotation function works in. ``'frac'`` uses them as @@ -151,7 +155,7 @@ def orbit_rank( R_true: torch.Tensor, symops: torch.Tensor, *, - side: str = "left", + side: str = "right", frame: str = "cart", reciprocal_basis: Optional[torch.Tensor] = None, thr_deg: float = 5.0, diff --git a/docs/changelog.rst b/docs/changelog.rst index 40369f9e..32eee9f4 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- The rotation function suppresses symmetry mates when it picks peaks, so its shortlist is one entry per orientation. The point group composes on the right of a peak's rotation, ``R R_g`` -- measured on real peak lists, where the left orbit finds no coincident pairs and the right finds every mate (187 of the 300 pairs among 3K7M's top 25 were mates of each other). The lab's orbit-based truth rank defaulted to the left side, which is why it disagreed with coordinate superposition. Placements unchanged, 30/30 at both windows - Fixed ``empirical_sigma_a`` taking the level of the observed-to-calculated Wilson-curve ratio rather than its shape. The two curves carry different absolute scales, so the ratio was 0.02-0.06 on one structure and 8-12 on another and the returned ``sigma_A`` was flat at 0.15-0.35 regardless of resolution; each curve is now divided by its geometric mean first. Per-shell factors are gauge in the rotation function's correlation, so its placements are unchanged (30/30) - The fast translation function scores the covariance of two normalised intensities -- the rotation function's LERF1 coefficient ``cw (E_obs^2 - 1) w sigma_A^2`` against the candidate's ``|E_calc(h, t)|^2``, normalised per candidate by the same Wilson fit -- instead of a raw-``|F_calc|^2`` ratio that was not a correlation and, on the four largest panel structures, was higher 40 A from the true position than at it. One FFT on a grid a third of the set's resolution apart with parabolic peak refinement replaces the 16-point coarse grid and three 100-point local refines; the Rice/Woolfson likelihood at a fixed Luzzati ``sigma_A`` picks among the top peaks and ranks the candidates. 30/30 true poses at the default window and 30/30 with the window removed, against 18/30 before - The translation stage runs in the configured float and complex dtypes rather than hard-coded double diff --git a/tests/unit/alignment/test_peak_finder_symmetry.py b/tests/unit/alignment/test_peak_finder_symmetry.py new file mode 100644 index 00000000..55200c5b --- /dev/null +++ b/tests/unit/alignment/test_peak_finder_symmetry.py @@ -0,0 +1,58 @@ +"""Rotation-function peaks are one per orientation, not one per symmetry mate. + +The greedy SO(3) suppression treats ``R`` and its point-group mates ``R R_g`` as +the same peak when the Cartesian symmetry rotations are supplied. The group +composes on the **right** -- measured on real peak lists, where the left orbit +finds no coincident pairs and the right orbit finds every mate -- so the test +pins that side too: a left-composed copy must survive as a distinct peak. +""" +import math + +import pytest +import torch + +from torchref.experimental.alignment.frf.peak_finder import find_rotation_peaks +from torchref.experimental.alignment.frf.rotation_utils import ( + edmonds_euler_from_rotation_matrix, + rotation_matrix_from_edmonds_euler, +) +from torchref.experimental.alignment.frf.types import AdaptiveRotationFunction + +pytestmark = pytest.mark.unit + + +def _rz(deg: float) -> torch.Tensor: + c, s = math.cos(math.radians(deg)), math.sin(math.radians(deg)) + return torch.tensor([[c, -s, 0.0], [s, c, 0.0], [0.0, 0.0, 1.0]], dtype=torch.float64) + + +def _arf(rotations, values) -> AdaptiveRotationFunction: + eul = torch.tensor([edmonds_euler_from_rotation_matrix(R) for R in rotations], + dtype=torch.float64) + n = eul.shape[0] + return AdaptiveRotationFunction( + alphas=eul[:, 0], betas=eul[:, 1], gammas=eul[:, 2], + values=torch.tensor(values, dtype=torch.float64), + beta_starts=torch.tensor([0, n]), beta_grid=torch.tensor([0.0]), + grid_sampling_deg=3.0, + ) + + +def test_symmetry_mates_collapse_to_one_peak_and_the_side_is_right(): + sym_cart = torch.stack([_rz(0.0), _rz(90.0), _rz(180.0), _rz(270.0)]) # 4 about z + R1 = rotation_matrix_from_edmonds_euler(0.3, 0.7, 1.1) + R_right = R1 @ _rz(90.0) # a mate: the group acts on the right + R_left = _rz(90.0) @ R1 # not a mate of R1 for a generic R1 + R3 = rotation_matrix_from_edmonds_euler(2.0, 1.3, 0.4) + arf = _arf([R1, R_right, R_left, R3], [10.0, 9.0, 8.5, 8.0]) + + plain = find_rotation_peaks(arf, n_peaks=10, sigma_threshold=-50.0, nms_radius_deg=6.0) + assert len(plain) == 4, "without symmetry every sample is its own peak" + + dedup = find_rotation_peaks(arf, n_peaks=10, sigma_threshold=-50.0, + nms_radius_deg=6.0, sym_cart=sym_cart) + scores = sorted(p.score for p in dedup) + assert scores == [8.0, 8.5, 10.0], ( + "the right-composed mate (9.0) must be suppressed and the " + f"left-composed copy (8.5) kept; got {scores}" + ) diff --git a/torchref/experimental/alignment/frf/api.py b/torchref/experimental/alignment/frf/api.py index 279d7f33..d0182230 100644 --- a/torchref/experimental/alignment/frf/api.py +++ b/torchref/experimental/alignment/frf/api.py @@ -154,8 +154,15 @@ def __init__( snr_cap: float = DEFAULT_SNR_CAP, trust_cap: float = DEFAULT_TRUST_CAP, shell_variance_weights: bool = False, + sym_cart: Optional[torch.Tensor] = None, ): self.device = s_obs.device + # The point group as Cartesian rotations, for the peak finder: with it, + # the returned peaks are distinct orientations rather than an + # orientation and its mates. `sym_mats` above is in the fractional + # basis and only detects the z-axis order; a direct caller without a + # cell cannot supply this and gets the plain suppression. + self.sym_cart = sym_cart # `L` and `d_min` arrive already coupled: the caller runs # `phaser_lmax_resolution` because it needs the same pair to size the @@ -414,5 +421,6 @@ def score_model( n_peaks=n_peaks, sigma_threshold=sigma_threshold, nms_radius_deg=max(2.0 * self.grid_sampling_deg, 6.0), + sym_cart=self.sym_cart, ) return arf, peaks diff --git a/torchref/experimental/alignment/frf/peak_finder.py b/torchref/experimental/alignment/frf/peak_finder.py index 6ae3ef81..1c0825f5 100644 --- a/torchref/experimental/alignment/frf/peak_finder.py +++ b/torchref/experimental/alignment/frf/peak_finder.py @@ -10,12 +10,22 @@ rotations (not by α, β, γ box distance — that would double-count near the poles). -We implement the same flow in PyTorch, vectorised where possible. +We implement the same flow in PyTorch, vectorised where possible, and the +suppression is **modulo the crystal's point group** when the Cartesian +symmetry rotations are supplied: an orientation and its symmetry mates are one +answer, and without this the shortlist handed downstream is mostly copies. On +3K7M (P432) 187 of the 300 pairs among the top 25 peaks were mates of each +other; on 2DQ6 (P3(1)21) 15 of 25 candidates were one orientation. + +The group acts on the **right**: a peak ``R`` maps the search-model frame onto +the crystal frame, and its mates are ``R R_g``. Measured, not assumed -- +composing on the left finds zero coincident pairs on every structure tried, the +right side finds all of them. """ from __future__ import annotations import math -from typing import List +from typing import List, Optional import torch @@ -34,11 +44,14 @@ def _so3_greedy_nms( values: torch.Tensor, nms_radius_deg: float, keep_at_most: int, + sym_cart: Optional[torch.Tensor] = None, ) -> torch.Tensor: """Return indices (into the input order) of kept peaks after SO(3) NMS. Greedy: walk the values in descending order; keep a candidate if its - angular distance from every already-kept rotation is > nms_radius_deg. + angular distance from every already-kept rotation -- and, with + ``sym_cart``, from every point-group mate ``R R_g`` of it -- is + > nms_radius_deg. """ n = values.shape[0] if n == 0: @@ -62,18 +75,22 @@ def _so3_greedy_nms( ) # (n, 3, 3) # angle > nms_radius ⇔ cos(angle) < cos(nms_radius); cos(angle) from trace. cos_thresh = math.cos(math.radians(nms_radius_deg)) + # The orbit of each kept rotation, R R_g over the point group; the identity + # alone when no symmetry is supplied. + G = (torch.eye(3, dtype=torch.float64).unsqueeze(0) if sym_cart is None + else sym_cart.to(torch.float64).cpu()) kept_idx: List[int] = [] - kept_R = torch.empty((keep_at_most, 3, 3), dtype=torch.float64) + kept_orbit = torch.empty((keep_at_most, G.shape[0], 3, 3), dtype=torch.float64) count = 0 for i_t in order: Ri = R_all[i_t] if count > 0: - trace = torch.einsum("kij,ij->k", kept_R[:count], Ri) + trace = torch.einsum("kgij,ij->kg", kept_orbit[:count], Ri) cos_theta = ((trace - 1.0) * 0.5).clamp(min=-1.0, max=1.0) - # Some kept rotation within nms_radius (cos_theta > cos_thresh) → skip. + # Some kept rotation, or a mate of one, within nms_radius → skip. if bool((cos_theta > cos_thresh).any()): continue - kept_R[count] = Ri + kept_orbit[count] = Ri.unsqueeze(0) @ G kept_idx.append(i_t) count += 1 if count >= keep_at_most: @@ -86,11 +103,15 @@ def find_rotation_peaks( n_peaks: int = 500, sigma_threshold: float = -5.0, nms_radius_deg: float = 6.0, + sym_cart: Optional[torch.Tensor] = None, ) -> List[RotationPeak]: """Greedy SO(3) NMS over the adaptive sample list. Returns peaks sorted by descending value, capped at ``n_peaks`` and - filtered by ``sigma >= sigma_threshold``. + filtered by ``sigma >= sigma_threshold``. With ``sym_cart`` -- the crystal's + point-group rotations as Cartesian ``(n_ops, 3, 3)`` matrices -- symmetry + mates of a kept peak are suppressed too, so every returned peak is a + distinct orientation. """ values = arf.values if values.numel() == 0: @@ -122,6 +143,7 @@ def find_rotation_peaks( a, b, g, v, nms_radius_deg=nms_radius_deg, keep_at_most=n_peaks, + sym_cart=sym_cart, ) # Gather kept peaks and move to CPU once (avoids a per-peak device sync). diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index 852bf025..ed2bcc85 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -124,8 +124,9 @@ class RotationSolutions: ``(n, 3, 3)`` float64. ``rotations[i]`` maps the search-model frame onto the crystal frame, so the coordinate rotation that places the model is its transpose: ``model.copy().rotate(rotations[i].T)``. Each is - determined only up to the crystal's rotational symmetry, so a solution - and its symmetry mates are the same answer. + determined only up to the crystal's rotational symmetry -- its mates are + ``rotations[i] @ R_g`` -- and the list carries one representative per + orbit, so consecutive entries are distinct orientations. scores : torch.Tensor ``(n,)`` rotation-function value at each orientation. z_scores : torch.Tensor @@ -445,6 +446,13 @@ def search_peaks( # `frf/preprocessing` with no production caller at all -- it was kept for # the ML rescore, and that was deleted. + # Point-group rotations in the Cartesian frame, so the peak finder can + # treat an orientation and its symmetry mates as one peak. As a set + # these equal B S B^-1; `hkl_symops_to_cartesian` returns the same + # rotations in a different order. + from .sh import hkl_symops_to_cartesian + sym_cart = hkl_symops_to_cartesian(sg_mats, rec_basis) + engine = FastRotationFunction( s_obs, F_obs, centric, sg_mats, L=L, d_min=d_min, d_max=d_max, @@ -456,6 +464,7 @@ def search_peaks( s_mag_asu=s_mag_asu, obs_weight=obs_weight, snr_cap=snr_cap, trust_cap=trust_cap, shell_variance_weights=shell_variance_weights, + sym_cart=sym_cart, ) _arf, peaks = engine.score_model( s_calc, F_calc, n_peaks=n_peaks, From 0c78eed6e2f29616fe2e8feabe58e87dda323de4 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 2 Sep 2026 12:39:43 +0200 Subject: [PATCH 146/250] Replace the rotation-only figures in the pipeline's documentation Every success count quoted in the pipeline and the pose-recovery harness -- 30/30, 24/30, 37/40, 36/40, 32/40 -- gated on the rotation alone, with a metric that read two of the six trigonal mates of a correct solution as failures. Measured on poses over six structures x ten seeds (job 544953) the three ranking arms are 60/60 each and pick the same candidate in every cell; the likelihood stays the default as the right object for the question, not on a measured margin. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/diagnostics/pose_recovery.py | 45 ++++++------ docs/changelog.rst | 1 + torchref/experimental/alignment/pipeline.py | 77 ++++++++++----------- 3 files changed, 58 insertions(+), 65 deletions(-) diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py index 8094016f..3257e540 100644 --- a/alignment_lab/diagnostics/pose_recovery.py +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -4,30 +4,28 @@ the deposited model and asks whether the pipeline gets it back, which is the only measurement that settles a change to either stage. -The reference number to beat is **24/30** (10 structures x 3 trials): what the -pipeline scored once the ML rescore was taken out of the middle, against 18/30 -with it. 2DQ6 and 3GR5 fail in every arm ever measured and cap recovery there. +Success is a pose: final coordinates within ``--success-deg`` of canonical in +orientation AND within ``--success-A`` of it in position, modulo the crystal +symmetry (Cartesian point-group mates, lattice translations, allowed origin +shifts and polar directions). The gate used to be rotation-only, and it passed +placements 40-55 A from the true position on 2DQ6, 3VRJ, 4BX9 and 6G9X. + +The panel stands at **30/30** (10 structures x 3 trials) and **60/60** over six +structures x ten seeds, on every ranking arm. -Arms (``--arms``) sweep how translation candidates are ranked: +Arms (``--arms``) sweep how the winner is chosen among placed candidates: ``llg`` - the default -- rank each rotation candidate by the translation likelihood at - its best translation. 36/40 over four structures x ten seeds, median - residual 1.43 deg. + the default -- the translation likelihood at each candidate's best + translation. ``analytic_r`` - rank by the analytical-scale R instead. Also 36/40, median 1.62 deg. The two - tie on the count and each wins one cell paired; they pick the same candidate - outright on 1DAW and 3K7M. The likelihood's margin is 6G9X alone. + the analytical-scale R instead. ``corr`` - rank by the translation function's own correlation. Measured 32/40 against - the other two arms' 36/40: worse, and worse paired against ``llg`` 5 to 1, - despite a rank-level harness predicting the reverse on a truth label that - disagreed with coordinate superposition. -Success is a pose: final coordinates within ``--success-deg`` of canonical in -orientation AND within ``--success-A`` of it in position, modulo the crystal -symmetry (Cartesian point-group mates, lattice translations, allowed origin -shifts and polar directions). The gate used to be rotation-only, and it passed -placements 40-55 A from the true position on 2DQ6, 4BX9 and 6G9X. + the fast translation function's own score. + +Over the six-structure sweep the three arms pick the same candidate in all 60 +cells; the arms exist so that can be re-checked whenever a structure separates +them. Usage:: @@ -112,10 +110,9 @@ def _report_candidates(solutions, R_true, symops, success_deg) -> None: R in 0 of 10 seeds while the pipeline solved 6 of them, because it fed the R-factor a different set of translation peaks. - ``SOLN`` lines are ordered as the pipeline ranked them -- by ``tf_corr``, - descending -- so line 0 is what it returned. ``R`` is carried alongside - because it used to be the ranking key and comparing the two orderings is the - point. ``dtruth`` is the angle from that candidate's orientation to the + ``SOLN`` lines are ordered as the pipeline ranked them -- by the likelihood, + descending -- so line 0 is what it returned. The fast score and ``R`` are + carried alongside so the three orderings can be compared. ``dtruth`` is the angle from that candidate's orientation to the true one modulo crystal symmetry; ``pick`` marks the winner and ``true`` marks every candidate that was in fact correct. """ @@ -124,7 +121,7 @@ def _report_candidates(solutions, R_true, symops, success_deg) -> None: ) R_t = R_true.to(torch.float64).cpu() - print(" SOLN rank rot_score tf_corr R dtruth flags") + print(" SOLN rank rot_score tf R dtruth flags") for i, sol in enumerate(solutions): R = torch.as_tensor(sol.rotation, dtype=torch.float64) # `rotation` maps the search-model frame onto the crystal frame; the diff --git a/docs/changelog.rst b/docs/changelog.rst index 32eee9f4..5c5a9381 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- The three candidate-ranking scores (likelihood, analytical R, fast translation score) pick the same candidate in all 60 cells of a six-structure x ten-seed sweep measured on poses; every earlier figure quoted for them was rotation-only. The likelihood stays the default - The rotation function suppresses symmetry mates when it picks peaks, so its shortlist is one entry per orientation. The point group composes on the right of a peak's rotation, ``R R_g`` -- measured on real peak lists, where the left orbit finds no coincident pairs and the right finds every mate (187 of the 300 pairs among 3K7M's top 25 were mates of each other). The lab's orbit-based truth rank defaulted to the left side, which is why it disagreed with coordinate superposition. Placements unchanged, 30/30 at both windows - Fixed ``empirical_sigma_a`` taking the level of the observed-to-calculated Wilson-curve ratio rather than its shape. The two curves carry different absolute scales, so the ratio was 0.02-0.06 on one structure and 8-12 on another and the returned ``sigma_A`` was flat at 0.15-0.35 regardless of resolution; each curve is now divided by its geometric mean first. Per-shell factors are gauge in the rotation function's correlation, so its placements are unchanged (30/30) - The fast translation function scores the covariance of two normalised intensities -- the rotation function's LERF1 coefficient ``cw (E_obs^2 - 1) w sigma_A^2`` against the candidate's ``|E_calc(h, t)|^2``, normalised per candidate by the same Wilson fit -- instead of a raw-``|F_calc|^2`` ratio that was not a correlation and, on the four largest panel structures, was higher 40 A from the true position than at it. One FFT on a grid a third of the set's resolution apart with parabolic peak refinement replaces the 16-point coarse grid and three 100-point local refines; the Rice/Woolfson likelihood at a fixed Luzzati ``sigma_A`` picks among the top peaks and ranks the candidates. 30/30 true poses at the default window and 30/30 with the window removed, against 18/30 before diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 42260d98..06dc4658 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -4,9 +4,11 @@ 1. **Fast Rotation Function** — Phaser-faithful Bessel-radial × SH expansion against a dense P1-box calc. It is a *shortlist generator*. It does not have - to rank well, and measurably does not: over 30 seeded cells it puts truth at - rank 0 in 6 of them. What it does reliably is put truth somewhere in the top - twenty. + to rank well, and it does not: over the panel its own ordering puts truth + first in a minority of cells. What it does reliably is put the true + orientation somewhere in the top twenty-five -- in every cell of every + panel run on record -- and its peaks are one per orientation, symmetry + mates suppressed. 2. **Fast Translation Function** — for *each* of the top-N orientations, one Crowther-Blow FFT over the fractional cell on a resolution-sized grid, with the rotation function's own normalised score equation as its coefficients, @@ -16,8 +18,10 @@ There is deliberately nothing between them. An ML re-ranking of the FRF peaks used to sit there and was removed: it reorders a shortlist that already contains -truth, and end-to-end pose recovery was 18/30 with it against 24/30 without -(McNemar p = 0.031, 6-0 discordant). +truth, and rotation recovery was 18/30 with it against 24/30 without (McNemar +p = 0.031, 6-0 discordant). Those figures, like every figure on this pipeline +before September 2026, gated on the rotation alone; see ``rank_by`` for what +the pose-gated panel measures. Every candidate is placed and then the best is taken, with no early stopping. Stopping early made the pipeline's answer depend on the order the rotation @@ -240,11 +244,11 @@ def __init__( # into this many; the likelihood picks. n_translation_candidates: int = 3, # Which score picks the winner among placed candidates. "llg" is the - # translation function's Rice/Woolfson likelihood; "r" the - # analytical-scale R-factor; "corr" the translation correlation. Not a - # tuning knob -- it exists because the three disagree and a rank-level - # proxy got the ordering wrong, so the comparison has to be made end to - # end. See the sort in `run` for what that measured. + # translation likelihood; "r" the analytical-scale R-factor; "corr" the + # fast translation function's own score. Not a tuning knob -- it exists + # because the three can disagree and a rank-level proxy once got the + # ordering wrong, so the comparison has to be made end to end on poses. + # See the sort in `run` for what that measured. rank_by: str = "llg", # Resolution window for the translation set. None means the rotation # search's own [d_max, d_min], so one window and one normalisation @@ -316,9 +320,9 @@ def _log_candidate(self, k: int, peak, r_analytic, t_frac, Fields are ``key=value`` so a reader does not depend on column order: ``k`` candidate index in rotation-function order, ``rf``/``rfz`` its - score and z, ``tf`` the translation correlation, ``llg`` the translation - likelihood, ``r`` the analytical-scale R, ``t`` the fractional - translation. All three placement scores are reported whichever one + score and z, ``tf`` the fast translation function's score, ``llg`` the + translation likelihood, ``r`` the analytical-scale R, ``t`` the + fractional translation. All three placement scores are reported whichever one ranks, because which of them a wrong placement disagreed on is the question, and they do disagree. """ @@ -430,36 +434,27 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: if not solutions: raise RuntimeError("Translation + joint refine produced no candidates.") - # Highest translation likelihood. Over four structures x ten seeds, - # success within 8 deg of canonical modulo crystal symmetry: + # Highest translation likelihood. The three scores are measured end to + # end on POSES -- rotation and translation, against Cartesian symmetry + # mates -- over six structures x ten seeds (the four the translation + # search used to mis-place, plus two controls; job 544953): # - # llg 36/40 median residual 1.43 deg - # r 36/40 1.62 - # corr 32/40 1.98 + # llg 60/60 + # r 60/60 + # corr 60/60 # - # The likelihood and R tie on the success count -- 36 each, and paired - # over the 40 cells each wins exactly one. Nothing separates them there, - # and an earlier 37-against-36 reading of this table did not survive - # remeasurement after the variance convention was corrected. + # and they do not merely tie: in every one of the 60 cells the three + # pick the SAME candidate, so the residual distributions are identical + # arm for arm. Once the translation objective was normalised there was + # nothing left for the selection rule to decide on this panel. # - # What separates them is where they differ at all, which is less often - # than the medians suggest: on 1DAW and 3K7M the two pick the SAME - # candidate and the residuals are identical. The whole difference is - # 6G9X, where the likelihood holds every residual under 2.3 deg and R - # spreads to 5.65 (medians 1.05 against 2.24). 2DQ6 goes the other way - # by a smaller margin (max 6.48 against 3.86). Across the 31 cells all - # three arms solve, the likelihood places closer 5 times against 1. - # - # So the default rests on one structure, not on a sweep-wide margin. It - # is kept because it is also the right object for the question -- an - # R-factor on a partial model at the resolution this runs at has little - # to distinguish with -- and because the arm is selectable if that - # reasoning ever stops holding. - # - # The correlation is here as a cautionary default-not-taken. A rank-level - # harness rated it best by a wide margin, 33/40 against R's 23/40, and - # end to end it is the worst of the three -- the harness's truth label - # disagreed with coordinate superposition. + # The likelihood stays the default because it is the right object for + # the question -- an R-factor on a partial model at this resolution has + # little to distinguish with, and the fast score is an expansion of the + # likelihood rather than the likelihood -- and because the arm is + # selectable if a structure ever separates them. Every earlier figure + # for these arms (37/40, 36/40, 32/40) gated on the rotation alone, with + # a metric that miscounted trigonal mates; none of them stands. if self.rank_by == "r": solutions.sort(key=lambda s: s.r_factor) elif self.rank_by == "corr": @@ -479,7 +474,7 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: winner.model.last_alignment_rfactor = winner.r_factor self._log(1, f"mr: winner ({self.rank_by}) " f"LLG={winner.llg_score:.1f} " - f"TF corr={winner.translation_score:.5f} " + f"TF={winner.translation_score:.5f} " f"analytic R={winner.r_factor:.4f}") self._log(2, "\n" + timer.summary()) return solutions From 989c8e78fb27d6bb49652661aa46e2ae712807e6 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 2 Sep 2026 12:48:36 +0200 Subject: [PATCH 147/250] Re-orient one P1 copy per candidate and build only the winner's placed model Each rotation candidate copied the search model twice -- once to rotate, once to set P1 on -- and a third time to build a placed model nobody read for 24 of the 25. On the 20k-atom structures that was a quarter of the run. One P1 copy is built with the translation set and re-oriented in place per candidate (the forward cache fingerprints parameters by pointer and version, so the next structure-factor call recomputes); the placed model is built for the winner, and place() builds it for any other solution. 192 tests (job 544963); 30/30 at the default window and 30/30 uncut (544964). Warm, exclusive EPYC 9335 node, 8 threads (544965): 1DAW 0.85 s, 2DQ6 1.35 s, 6G9X 1.30 s, 3K7M 2.35 s, 4BX9 2.55 s, from 1.1 / 1.9 / 2.2 / 2.7 / 3.8 s. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- docs/changelog.rst | 1 + torchref/experimental/alignment/pipeline.py | 110 +++++++++++++------- 2 files changed, 74 insertions(+), 37 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 5c5a9381..34aa1508 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- The placement loop re-orients one P1 copy of the search model in place per candidate and builds the placed model for the winner only, instead of copying the model three times per candidate. ``MRSolution.model`` is ``None`` for the other candidates; ``MolecularReplacementPipeline.place`` builds it on request. Warm on one EPYC 9335 node: 1DAW 0.85 s, 2DQ6 1.35 s, 6G9X 1.30 s, 3K7M 2.35 s, 4BX9 2.55 s per alignment (from 1.1, 1.9, 2.2, 2.7, 3.8) - The three candidate-ranking scores (likelihood, analytical R, fast translation score) pick the same candidate in all 60 cells of a six-structure x ten-seed sweep measured on poses; every earlier figure quoted for them was rotation-only. The likelihood stays the default - The rotation function suppresses symmetry mates when it picks peaks, so its shortlist is one entry per orientation. The point group composes on the right of a peak's rotation, ``R R_g`` -- measured on real peak lists, where the left orbit finds no coincident pairs and the right finds every mate (187 of the 300 pairs among 3K7M's top 25 were mates of each other). The lab's orbit-based truth rank defaulted to the left side, which is why it disagreed with coordinate superposition. Placements unchanged, 30/30 at both windows - Fixed ``empirical_sigma_a`` taking the level of the observed-to-calculated Wilson-curve ratio rather than its shape. The two curves carry different absolute scales, so the ratio was 0.02-0.06 on one structure and 8-12 on another and the returned ``sigma_A`` was flat at 0.15-0.35 regardless of resolution; each curve is now divided by its geometric mean first. Per-shell factors are gauge in the rotation function's correlation, so its placements are unchanged (30/30) diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 06dc4658..18baedf4 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -166,8 +166,12 @@ class MRSolution: llg_score : float **The ranking key**: the translation likelihood at that placement, higher better. - model : ModelFT - The rotated and translated model for this candidate. + model : ModelFT or None + The rotated and translated model. Built for the winner only -- copying + and moving a 20k-atom model 25 times was a quarter of the run on the + large structures, to produce 24 models nobody reads. + :meth:`MolecularReplacementPipeline.place` builds it for any other + solution on request. """ rotation: np.ndarray @@ -175,7 +179,7 @@ class MRSolution: rotation_score: float translation_score: float r_factor: float - model: "ModelFT" + model: Optional["ModelFT"] = None llg_score: float = float("nan") @@ -291,7 +295,12 @@ def __init__( self._frf = None self._obs = None self._tmask = None - self._eye3 = torch.eye(3, dtype=torch.float64) + # One P1 copy of the search model, re-oriented in place per candidate + # -- see `_prepare_translation_arrays`. + self._p1 = None + self._p1_xyz0 = None + self._p1_center = None + self._evaluator = None #: Levels are documented on the class. They are a contract, not a dial: #: level 2 is specifically "one machine-readable line per candidate", and @@ -401,10 +410,12 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: solutions: List[MRSolution] = [] for k in range(n_rot): peak_k = candidates[k] - rotated_k, R_rec_k = self._make_rotated(peak_k) + R_rec_k = rotation_matrix_from_edmonds_euler( + peak_k.alpha, peak_k.beta, peak_k.gamma) + self._orient_template(R_rec_k) self._log(3, f"\nfit_to_data: rot{k} " f"(RF={peak_k.score:.2f}, σ_Z={peak_k.sigma:.2f})") - placement = self._placement_for_candidate(rotated_k) + placement = self._placement_for_candidate() if placement is None: self._log(2, f"CAND k={k} rf={float(peak_k.score):.4f} " f"rfz={float(peak_k.sigma):.3f} tf=nan r=nan " @@ -413,12 +424,6 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: r_analytic, t_refined, tf_score, llg_score = placement self._log_candidate(k, peak_k, r_analytic, t_refined, tf_score, llg_score) - - placed = rotated_k.copy().translate( - t_refined.to(self.model.dtype_float), fractional=True, - ) - placed.last_alignment_rotation = R_rec_k - placed.last_alignment_translation = t_refined solutions.append( MRSolution( rotation=R_rec_k.detach().cpu().numpy(), @@ -426,7 +431,6 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: rotation_score=float(peak_k.score), translation_score=float(tf_score), r_factor=float(r_analytic), - model=placed, llg_score=float(llg_score), ) ) @@ -471,7 +475,7 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: # doing a worse version of that here to print a number is not worth a # third of the runtime. A caller that wants an R-work can build a # `Scaler` on the returned model. - winner.model.last_alignment_rfactor = winner.r_factor + winner.model = self.place(winner) self._log(1, f"mr: winner ({self.rank_by}) " f"LLG={winner.llg_score:.1f} " f"TF={winner.translation_score:.5f} " @@ -504,6 +508,36 @@ def _rotation_candidates(self, frf) -> list: # out of it. The translation function does the discrimination. return sorted(peaks, key=lambda p: p.score, reverse=True) + def place(self, solution: MRSolution) -> "ModelFT": + """Build the placed model for ``solution``: a copy of the search model, + rotated and translated, carrying the alignment provenance attributes.""" + R_rec = torch.as_tensor(solution.rotation, dtype=torch.float64) + placed = self.model.copy().rotate( + R_rec.T.contiguous().to(device=self.model.device, + dtype=self.model.dtype_float), + ) + if str(placed.spacegroup) != str(self.data.spacegroup): + placed.spacegroup = self.data.spacegroup.hm + if solution.translation is not None: + t = torch.as_tensor(solution.translation, dtype=self.model.dtype_float) + placed.translate(t, fractional=True) + placed.last_alignment_translation = t + placed.last_alignment_rotation = R_rec + placed.last_alignment_rfactor = solution.r_factor + return placed + + def _orient_template(self, R_rec: torch.Tensor) -> None: + """Write the candidate orientation into the shared P1 copy. + + ``xyz = R_rec^T (xyz0 - c) + c`` about the search model's centroid, the + same rotation ``Model.rotate`` would apply. The forward cache + fingerprints parameters by pointer and version, so the next + structure-factor call recomputes. + """ + p1 = self._p1 + R_app = R_rec.T.to(device=self._p1_xyz0.device, dtype=self._p1_xyz0.dtype) + p1.xyz[:] = (self._p1_xyz0 - self._p1_center) @ R_app.T + self._p1_center + def _make_rotated(self, peak: "RotationPeak"): """Rotate the search model onto a candidate orientation. @@ -578,8 +612,30 @@ def _prepare_translation_arrays(self) -> None: + ("" if sig_F_full is not None else " (no sigmas: unit weight)")) - def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: - """Translation search for one rotation candidate. + # One P1 copy of the search model for the whole run, re-oriented in + # place per candidate. Two copies per candidate -- one to rotate, one to + # set P1 on -- were a quarter of the run on the large structures. + # + # Its FFT grid is sized to the translation set, not to the model's + # default 1.0 A: |s| is invariant under the symmetry rotations, so every + # rotated index the evaluator is asked for lies inside 1/tf_d_min. Two + # thirds of the window's resolution, not the resolution itself: at + # max_res = tf_d_min the transform's coherence with the 1.0 A grid over + # the 15-4 A set is 0.987 on 2DQ6 (0.9987-0.9999 on 1DAW, 3K7M, 4BX9); + # at tf_d_min/1.5 it is 0.9995-1.0000 everywhere, at 10-38 ms against + # 200-860 ms. max_res first -- the space-group setter rebuilds the FFT + # and reads it. + p1 = self.model.copy() + if self.tf_d_min > 0.0: + p1.max_res = self.tf_d_min / 1.5 + p1.spacegroup = "P 1" + self._p1 = p1 + self._p1_xyz0 = p1.xyz().detach().clone() + self._p1_center = self._p1_xyz0.mean(dim=0) + self._evaluator = DirectModelEvaluator(p1) + + def _placement_for_candidate(self) -> Optional[tuple]: + """Translation search for the orientation currently in the P1 template. Returns ``(r_analytic, t, tf_score, llg)`` for the translation the likelihood prefers among the fast search's top peaks, or ``None`` if the @@ -590,28 +646,8 @@ def _placement_for_candidate(self, rotated_k) -> Optional[tuple]: timer = self._timer obs = self._obs - if str(rotated_k.spacegroup) != str(data.spacegroup): - rotated_k.spacegroup = data.spacegroup.hm - rotated_p1 = rotated_k.copy() - # Size the P1 copy's FFT grid to the translation set, not to the - # model's default 1.0 A: |s| is invariant under the symmetry rotations, - # so every rotated index the evaluator is asked for lies inside - # 1/tf_d_min. This is where the placement stage spent most of its time, - # on a grid 30-48x larger than the reflections it was asked for. - # - # Two thirds of the window's resolution, not the resolution itself. At - # max_res = tf_d_min the transform's coherence with the 1.0 A grid over - # the 15-4 A set is 0.987 on 2DQ6 (0.9987-0.9999 on 1DAW, 3K7M, 4BX9); - # at tf_d_min/1.5 it is 0.9995-1.0000 everywhere, at 10-38 ms against - # 200-860 ms. max_res first -- the space-group setter rebuilds the FFT - # and reads it. - if self.tf_d_min > 0.0: - rotated_p1.max_res = self.tf_d_min / 1.5 - rotated_p1.spacegroup = "P 1" - evaluator = DirectModelEvaluator(rotated_p1) - timer.start("5_candidate_transform") - cand = prepare_candidate(evaluator, obs, data.spacegroup, data.cell) + cand = prepare_candidate(self._evaluator, obs, data.spacegroup, data.cell) timer.stop("5_candidate_transform") # One FFT on a grid a third of the set's resolution apart: dense enough From ec390c08aa5a59e085dee2df96c264af56321a1c Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 2 Sep 2026 13:02:41 +0200 Subject: [PATCH 148/250] Accumulate only the upper triangle of symmetry pairs in the translation function The pair (j, i) is the conjugate of (i, j) at -dh and the diagonal carries no t, so the map is twice the real part of the upper triangle's transform plus a constant. Half the scatter, which is what the stage costs on high-symmetry cells: 3K7M's translation stage 1.04 s to 0.72 s, the whole alignment 2.35 s to 2.01 s (job 544990). 192 tests (544988); 30/30 at both windows (544989). MRSolution.candidate_index records each solution's position in the rotation function's list, and the pose harness prints it, so the depth of shortlist a solution needed can be read off. Measured over 10 structures x 5 seeds (544991): with symmetry mates suppressed the rotation function's first peak is the true orientation in all 50 cells. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/first_true_rank.sh | 27 +++++++++++++++++++ alignment_lab/analysis/pipeline_timing_gpu.sh | 20 ++++++++++++++ alignment_lab/diagnostics/pipeline_timing.py | 16 ++++++++--- alignment_lab/diagnostics/pose_recovery.py | 4 +-- docs/changelog.rst | 1 + torchref/experimental/alignment/pipeline.py | 5 ++++ .../experimental/alignment/translation.py | 14 ++++++---- 7 files changed, 76 insertions(+), 11 deletions(-) create mode 100644 alignment_lab/analysis/first_true_rank.sh create mode 100644 alignment_lab/analysis/pipeline_timing_gpu.sh diff --git a/alignment_lab/analysis/first_true_rank.sh b/alignment_lab/analysis/first_true_rank.sh new file mode 100644 index 00000000..ee1dabff --- /dev/null +++ b/alignment_lab/analysis/first_true_rank.sh @@ -0,0 +1,27 @@ +#!/bin/bash +# How deep in the rotation function's (de-duplicated) list does the true +# orientation sit? The SOLN table carries each solution's FRF index k and its +# truth flag; the smallest k flagged true per cell is the shortlist depth the +# translation search needed. +#SBATCH --job-name=ftr +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=hour +#SBATCH --time=00:40:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-9 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +for T in 0 1 2 3 4; do + "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial $T --arms llg --verbose 2 2>/dev/null \ + | awk -v pdb=$PDB -v t=$T '/SOLN +[0-9]+ +[0-9]+ .*true/ {k=$3+0; if (min=="" || k&1 \ + | grep -v "Warning\|warnings.warn" | grep -E "^ROW|^stage|^[0-9]_|^TOTAL|^---|Traceback|Error" | grep -v "^---" +echo DONE diff --git a/alignment_lab/diagnostics/pipeline_timing.py b/alignment_lab/diagnostics/pipeline_timing.py index c51952ef..fb07eb9b 100644 --- a/alignment_lab/diagnostics/pipeline_timing.py +++ b/alignment_lab/diagnostics/pipeline_timing.py @@ -24,6 +24,9 @@ def main() -> int: ap = argparse.ArgumentParser(description=__doc__) ap.add_argument("--pdbs", default="1DAW,2DQ6,6G9X,3K7M,4BX9") ap.add_argument("--threads", type=int, default=8) + ap.add_argument("--n-rotation-candidates", type=int, default=25) + ap.add_argument("--device", default="cpu", + help="where the model and data live; set TORCHREF_DEVICE to match") args = ap.parse_args() torch.set_num_threads(args.threads) @@ -31,7 +34,7 @@ def main() -> int: for pdb in [p.strip() for p in args.pdbs.split(",") if p.strip()]: assert pdb in BENCH_PDBS - model, data = load_case(pdb) + model, data = load_case(pdb, device=args.device) canonical = model.xyz().clone() R_true = random_rotation(seed_for(pdb, 0)) for run in range(2): @@ -41,16 +44,21 @@ def main() -> int: center=canonical.mean(0)) pipe = MolecularReplacementPipeline( data, search, d_min=4.0, d_max=15.0, n_shells=20, - n_rotation_peaks=200, n_rotation_candidates=25, + n_rotation_peaks=200, + n_rotation_candidates=args.n_rotation_candidates, verbose=2 if run == 1 else 0, ) + if args.device.startswith("cuda"): + torch.cuda.synchronize() t0 = time.perf_counter() sols = pipe.run(do_translation=True) + if args.device.startswith("cuda"): + torch.cuda.synchronize() secs = time.perf_counter() - t0 rot, trans = pose_error(sols[0].model.xyz(), canonical, data.cell, data.spacegroup) - print(f"ROW pdb={pdb} run={run} seconds={secs:.2f} rot_deg={rot:.2f} " - f"trans_A={trans:.2f}", flush=True) + print(f"ROW pdb={pdb} run={run} device={args.device} n_cand={args.n_rotation_candidates} seconds={secs:.2f} " + f"rot_deg={rot:.2f} trans_A={trans:.2f}", flush=True) return 0 diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py index 3257e540..7378437d 100644 --- a/alignment_lab/diagnostics/pose_recovery.py +++ b/alignment_lab/diagnostics/pose_recovery.py @@ -121,7 +121,7 @@ def _report_candidates(solutions, R_true, symops, success_deg) -> None: ) R_t = R_true.to(torch.float64).cpu() - print(" SOLN rank rot_score tf R dtruth flags") + print(" SOLN rank k rot_score tf R dtruth flags") for i, sol in enumerate(solutions): R = torch.as_tensor(sol.rotation, dtype=torch.float64) # `rotation` maps the search-model frame onto the crystal frame; the @@ -130,7 +130,7 @@ def _report_candidates(solutions, R_true, symops, success_deg) -> None: d = min(float(rotation_angular_distance_deg(R.T @ R_t, symops[k])) for k in range(symops.shape[0])) flags = ("pick " if i == 0 else " ") + ("true" if d <= success_deg else "") - print(f" SOLN {i:4d} {sol.rotation_score:10.3f} " + print(f" SOLN {i:4d} {sol.candidate_index:3d} {sol.rotation_score:10.3f} " f"{sol.translation_score:10.5f} {sol.r_factor:7.4f} " f"{d:8.2f} {flags}") diff --git a/docs/changelog.rst b/docs/changelog.rst index 34aa1508..e8f5bafc 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- The fast translation function accumulates only the upper triangle of symmetry pairs; the lower triangle is its conjugate mirror and the diagonal a constant. Half the scatter, which is the stage's cost on high-symmetry cells: 3K7M's translation stage 1.04 s to 0.72 s. ``MRSolution.candidate_index`` records each solution's position in the rotation function's list - The placement loop re-orients one P1 copy of the search model in place per candidate and builds the placed model for the winner only, instead of copying the model three times per candidate. ``MRSolution.model`` is ``None`` for the other candidates; ``MolecularReplacementPipeline.place`` builds it on request. Warm on one EPYC 9335 node: 1DAW 0.85 s, 2DQ6 1.35 s, 6G9X 1.30 s, 3K7M 2.35 s, 4BX9 2.55 s per alignment (from 1.1, 1.9, 2.2, 2.7, 3.8) - The three candidate-ranking scores (likelihood, analytical R, fast translation score) pick the same candidate in all 60 cells of a six-structure x ten-seed sweep measured on poses; every earlier figure quoted for them was rotation-only. The likelihood stays the default - The rotation function suppresses symmetry mates when it picks peaks, so its shortlist is one entry per orientation. The point group composes on the right of a peak's rotation, ``R R_g`` -- measured on real peak lists, where the left orbit finds no coincident pairs and the right finds every mate (187 of the 300 pairs among 3K7M's top 25 were mates of each other). The lab's orbit-based truth rank defaulted to the left side, which is why it disagreed with coordinate superposition. Placements unchanged, 30/30 at both windows diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 18baedf4..d30c2f3f 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -166,6 +166,9 @@ class MRSolution: llg_score : float **The ranking key**: the translation likelihood at that placement, higher better. + candidate_index : int + Position of this orientation in the rotation function's own ordering, + so the depth of shortlist a solution came from can be read off. model : ModelFT or None The rotated and translated model. Built for the winner only -- copying and moving a 20k-atom model 25 times was a quarter of the run on the @@ -181,6 +184,7 @@ class MRSolution: r_factor: float model: Optional["ModelFT"] = None llg_score: float = float("nan") + candidate_index: int = -1 class MolecularReplacementPipeline(DeviceMixin): @@ -432,6 +436,7 @@ def run(self, do_translation: bool = True) -> List[MRSolution]: translation_score=float(tf_score), r_factor=float(r_analytic), llg_score=float(llg_score), + candidate_index=k, ) ) diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index 2cdc25ac..a4bf365d 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -413,14 +413,18 @@ def fast_translation_function( coeff = obs.coeff.to(device).to(cplx) h_R_int = cand.h_R.round().to(torch.int64) + # The pair (j, i) is the conjugate of (i, j) at -dh, so the map is twice + # the real part of the upper triangle's transform plus the diagonal, which + # carries no t and is a constant. Half the scatter, which is the cost here. W = torch.zeros(nx * ny * nz, dtype=cplx, device=device) - for i in range(S): - pair = G[i].conj().view(1, -1) * G # (S, N) - dh = h_R_int - h_R_int[i:i + 1] # (S, N, 3) + for i in range(S - 1): + pair = G[i].conj().view(1, -1) * G[i + 1:] # (S-i-1, N) + dh = h_R_int[i + 1:] - h_R_int[i:i + 1] # (S-i-1, N, 3) flat = ((dh[..., 0] % nx) * ny + (dh[..., 1] % ny)) * nz + (dh[..., 2] % nz) W.index_add_(0, flat.reshape(-1), (coeff.view(1, -1) * pair).reshape(-1)) - score = (torch.fft.ifftn(W.view(nx, ny, nz), dim=(0, 1, 2)).real - * float(nx * ny * nz)).to(real) + diag = (obs.coeff.to(device).to(real) * (G.abs() ** 2).sum(dim=0).to(real)).sum() + score = (2.0 * torch.fft.ifftn(W.view(nx, ny, nz), dim=(0, 1, 2)).real + * float(nx * ny * nz)).to(real) + diag radii = tuple(float(cluster_radius_A) / float(L) for L in (real_cell.a, real_cell.b, real_cell.c)) From ebb9393be1307e3db8119501770982855a4b2a3a Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 2 Sep 2026 13:10:27 +0200 Subject: [PATCH 149/250] Carry ten rotation candidates by default With symmetry mates suppressed, the rotation function's first distinct peak is the true orientation in 50 of 50 pose-gated cells (10 structures x 5 seeds, job 544991), so the 25-deep shortlist was mostly cost. At 10 the panel is 30/30 at the default window and 30/30 uncut (545015), and the warm exclusive-node time per alignment is 1DAW 0.42 s, 2DQ6 0.67 s, 6G9X 0.67 s, 3K7M 1.02 s, 4BX9 1.11 s (545016), from 0.85 / 1.31 / 1.23 / 2.01 / 2.42. The depth is a safety margin measured on deposited models as search models; it is a parameter, and poorer models should raise it. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- .../analysis/pipeline_timing_ncand10.sh | 19 +++++++++++++ alignment_lab/analysis/pose_panel_ncand10.sh | 28 +++++++++++++++++++ alignment_lab/analysis/soln_tables.sh | 20 +++++++++++++ docs/changelog.rst | 1 + torchref/experimental/alignment/pipeline.py | 13 +++++++-- 5 files changed, 79 insertions(+), 2 deletions(-) create mode 100644 alignment_lab/analysis/pipeline_timing_ncand10.sh create mode 100644 alignment_lab/analysis/pose_panel_ncand10.sh create mode 100644 alignment_lab/analysis/soln_tables.sh diff --git a/alignment_lab/analysis/pipeline_timing_ncand10.sh b/alignment_lab/analysis/pipeline_timing_ncand10.sh new file mode 100644 index 00000000..3f3eec5c --- /dev/null +++ b/alignment_lab/analysis/pipeline_timing_ncand10.sh @@ -0,0 +1,19 @@ +#!/bin/bash +#SBATCH --job-name=ptime10 +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:59:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --exclusive +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +"$PY" -u alignment_lab/diagnostics/pipeline_timing.py --threads 8 --n-rotation-candidates 10 2>/dev/null \ + | grep -E "^ROW|^stage|^[0-9]_|^TOTAL|^---" +echo DONE diff --git a/alignment_lab/analysis/pose_panel_ncand10.sh b/alignment_lab/analysis/pose_panel_ncand10.sh new file mode 100644 index 00000000..e0dd70b6 --- /dev/null +++ b/alignment_lab/analysis/pose_panel_ncand10.sh @@ -0,0 +1,28 @@ +#!/bin/bash +# The 10 x 3 panel with the pose gate (rotation AND translation), at the +# pipeline's default translation window and with the window removed. The +# default is now the rotation search's own window; "full" is what it used to be. +#SBATCH --job-name=ptrans10 +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err +#SBATCH --partition=hour +#SBATCH --time=00:59:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=64G +#SBATCH --constraint=cpu_epyc9335 +#SBATCH --array=0-9 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) +PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +for T in 0 1 2; do + "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial $T --arms llg --n-rotation-candidates 10 \ + 2>/dev/null | grep '^ROW ' | sed 's/^ROW/ROW window=default/' + "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial $T --arms llg --n-rotation-candidates 10 \ + --tf-d-min 0 --tf-d-max inf 2>/dev/null | grep '^ROW ' | sed 's/^ROW/ROW window=full/' +done +echo DONE diff --git a/alignment_lab/analysis/soln_tables.sh b/alignment_lab/analysis/soln_tables.sh new file mode 100644 index 00000000..2052cd9c --- /dev/null +++ b/alignment_lab/analysis/soln_tables.sh @@ -0,0 +1,20 @@ +#!/bin/bash +# Full candidate tables for two cells, to check the shortlist-depth summary by eye. +#SBATCH --job-name=soln +#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out +#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err +#SBATCH --partition=hour +#SBATCH --time=00:20:00 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=48G +#SBATCH --constraint=cpu_epyc9335 +set -uo pipefail +REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement +PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python +cd "$REPO" +export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 +export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" +for P in 2DQ6 3K7M 6G9X; do + "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb $P --trial 0 --arms llg --verbose 2 2>/dev/null | grep -E "^===|SOLN|^ROW" +done +echo DONE diff --git a/docs/changelog.rst b/docs/changelog.rst index e8f5bafc..84c4aa76 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- The molecular-replacement pipeline carries 10 rotation candidates by default instead of 25. With symmetry mates suppressed the rotation function's first peak is the true orientation in 50 of 50 pose-gated cells, and the panel is 30/30 at either depth. Warm on one EPYC 9335 node: 1DAW 0.42 s, 2DQ6 0.67 s, 6G9X 0.67 s, 3K7M 1.02 s, 4BX9 1.11 s per alignment - The fast translation function accumulates only the upper triangle of symmetry pairs; the lower triangle is its conjugate mirror and the diagonal a constant. Half the scatter, which is the stage's cost on high-symmetry cells: 3K7M's translation stage 1.04 s to 0.72 s. ``MRSolution.candidate_index`` records each solution's position in the rotation function's list - The placement loop re-orients one P1 copy of the search model in place per candidate and builds the placed model for the winner only, instead of copying the model three times per candidate. ``MRSolution.model`` is ``None`` for the other candidates; ``MolecularReplacementPipeline.place`` builds it on request. Warm on one EPYC 9335 node: 1DAW 0.85 s, 2DQ6 1.35 s, 6G9X 1.30 s, 3K7M 2.35 s, 4BX9 2.55 s per alignment (from 1.1, 1.9, 2.2, 2.7, 3.8) - The three candidate-ranking scores (likelihood, analytical R, fast translation score) pick the same candidate in all 60 cells of a six-structure x ten-seed sweep measured on poses; every earlier figure quoted for them was rotation-only. The likelihood stays the default diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index d30c2f3f..7e380fcf 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -24,6 +24,8 @@ the pose-gated panel measures. Every candidate is placed and then the best is taken, with no early stopping. +Ten candidates by default: the rotation function's first distinct peak was +the true orientation in every pose-gated cell measured, so ten is a margin. Stopping early made the pipeline's answer depend on the order the rotation function happened to produce -- it walked the list until one placement beat an R-factor threshold and returned that, so it could accept the third candidate @@ -246,7 +248,14 @@ def __init__( n_rotation_peaks: int = 500, model_error_A: Optional[float] = None, # --- candidate tree --- - n_rotation_candidates: int = 25, + # Distinct orientations carried into the translation search. A safety + # margin, not a requirement: with symmetry mates suppressed the + # rotation function's FIRST peak is the true orientation in 50 of 50 + # pose-gated cells (10 structures x 5 seeds), and the panel is 30/30 at + # 10 as at 25. Measured on the deposited models as search models; raise + # it for poorer models, since every candidate costs a structure-factor + # evaluation and a translation FFT. + n_rotation_candidates: int = 10, # Peaks of the fast translation function re-scored by the likelihood # for each orientation. The fast map only has to get the true peak # into this many; the likelihood picks. @@ -702,7 +711,7 @@ def align_model_to_data( verbose: int = 0, do_translation: bool = True, n_translation_candidates: int = 3, - n_rotation_candidates: int = 25, + n_rotation_candidates: int = 10, rank_by: str = "llg", tf_d_min: Optional[float] = None, tf_d_max: Optional[float] = None, From 00a3511d435e67bcf625f979fbbdaa0118ffba7d Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 2 Sep 2026 13:11:47 +0200 Subject: [PATCH 150/250] Made grid sizing lazy and cached --- tests/functional/test_model_ft_functional.py | 67 +-- tests/helpers/device_cases.py | 32 +- tests/integration/test_sfds_device.py | 12 +- tests/unit/model/test_sf_grid_key.py | 212 +++++++ tests/unit/refinement/test_wilson_prior.py | 2 +- tests/unit/structure_factor/helpers.py | 19 +- tests/unit/structure_factor/test_dispatch.py | 9 +- tests/unit/structure_factor/test_forward.py | 12 +- tests/unit/symmetry/test_cell_identity.py | 102 ++++ .../unit/symmetry/test_spacegroup_identity.py | 52 ++ tests/unit/utils/test_device_resolution.py | 9 +- torchref/experimental/alignment/rigid_body.py | 6 +- .../experimental/ensemble/ensemble_model.py | 4 +- .../ensemble/ensemble_refinement.py | 2 +- .../monolithic_refinement/density_solvent.py | 11 +- torchref/experimental/targets/realspace.py | 27 +- torchref/model/__init__.py | 8 +- torchref/model/context.py | 14 + torchref/model/model.py | 31 +- torchref/model/model_ft.py | 277 ++++----- torchref/model/sf_ds.py | 233 +++----- torchref/model/sf_fft.py | 548 +++++++++--------- torchref/refinement/base_refinement.py | 2 +- torchref/scaling/solvent.py | 2 - torchref/symmetry/cell.py | 80 ++- torchref/symmetry/spacegroup.py | 28 +- torchref/topology/restraints.py | 4 +- 27 files changed, 1077 insertions(+), 728 deletions(-) create mode 100644 tests/unit/model/test_sf_grid_key.py create mode 100644 tests/unit/symmetry/test_cell_identity.py create mode 100644 tests/unit/symmetry/test_spacegroup_identity.py diff --git a/tests/functional/test_model_ft_functional.py b/tests/functional/test_model_ft_functional.py index a0606d87..b8919254 100644 --- a/tests/functional/test_model_ft_functional.py +++ b/tests/functional/test_model_ft_functional.py @@ -43,16 +43,16 @@ def test_modelft_load_cif(self, sample_cif_file): assert len(model.cell) == 6 def test_modelft_has_gridsize(self, sample_cif_file): - """Test that ModelFT sets up gridsize after loading.""" + """The grid resolves from the loaded cell and space group on first read.""" from torchref.model.model_ft import ModelFT - + model = ModelFT(max_res=2.0, verbose=0) + assert model.gridsize is None # no crystal yet model.load_cif(str(sample_cif_file)) - - # Check gridsize is set - if model.gridsize is not None: - assert len(model.gridsize) == 3 - assert all(g > 0 for g in model.gridsize) + + assert model.gridsize is not None + assert len(model.gridsize) == 3 + assert all(g > 0 for g in model.gridsize) @pytest.mark.integration @@ -86,29 +86,21 @@ def test_scattering_factors_available(self, sample_cif_file): class TestModelFTGridOperations: """Test ModelFT grid operations.""" - def test_setup_gridsize(self, sample_cif_file): - """Test grid size setup.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - - # Setup gridsize - gridsize = model.setup_gridsize(max_res=2.0) - - assert gridsize is not None - assert len(gridsize) == 3 - assert all(g > 0 for g in gridsize) - def test_setup_grid(self, sample_cif_file): - """Test full grid setup.""" + """An explicit grid size overrides the resolution-derived one.""" from torchref.model.model_ft import ModelFT - + model = ModelFT(max_res=2.0, verbose=0) model.load_cif(str(sample_cif_file)) - - # Model should have grid setup - assert model.gridsize is not None or hasattr(model, 'map') + derived = model.grid_shape + assert derived is not None and len(derived) == 3 + + model.setup_grid(gridsize=(24, 24, 24)) + assert model.grid_shape == (24, 24, 24) + assert model.explicit_gridsize == (24, 24, 24) + + model.explicit_gridsize = None + assert model.grid_shape == derived @pytest.mark.integration @@ -122,13 +114,12 @@ def test_get_real_space_grid(self, sample_cif_file): model = ModelFT(max_res=2.0, verbose=0) model.load_cif(str(sample_cif_file)) - - if model.gridsize is not None: - # Get real space grid - grid = get_real_grid(model.cell, max_res=2.0, device='cpu') - - assert grid is not None - assert len(grid.shape) == 4 # Should be 4D (nx, ny, nz, 3) + + assert model.gridsize is not None + grid = get_real_grid(model.cell, max_res=2.0, device='cpu') + + assert grid is not None + assert len(grid.shape) == 4 # Should be 4D (nx, ny, nz, 3) @pytest.mark.integration @@ -146,12 +137,12 @@ def test_map_symmetry_available(self, sample_cif_file): assert model.spacegroup is not None # The map operator comes from the space group, keyed on the grid shape. - if model.gridsize is not None: - gridsize = tuple(model.gridsize.tolist()) + gridsize = model.grid_shape + assert gridsize is not None - operator = model.spacegroup.map_operator(gridsize) - assert operator is not None - assert operator.map_shape == gridsize + operator = model.spacegroup.map_operator(gridsize) + assert operator is not None + assert operator.map_shape == gridsize @pytest.mark.integration diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index 5d65eaef..16b6e9b8 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -151,6 +151,24 @@ def _cell(device): return Cell(_CELL, device=device) + +def _ctx(d): + """A ModelContext whose cell and space group both live on ``d``.""" + from torchref.model.context import ModelContext + from torchref.symmetry import SpaceGroup + + return ModelContext(cell=_cell(d), spacegroup=SpaceGroup(_SG, device=d)) + + +def _sffft_with_grid(d): + """An SfFFT whose grid buffers exist, so the tensor walk reaches them.""" + from torchref.model.sf_fft import SfFFT + + sf = SfFFT(_ctx(d), max_res=2.0) + sf.ensure_grid() + return sf + + CASES: List[DeviceCase] = [ DeviceCase("EdgeBlock", _edge_block, "EdgeBlock"), DeviceCase("AtomGraph", _atom_graph, "AtomGraph"), @@ -169,22 +187,28 @@ def _cell(device): "SfFFT_from_cell", lambda d: __import__( "torchref.model.sf_fft", fromlist=["SfFFT"] - ).SfFFT(cell=_cell(d), spacegroup=_SG, max_res=2.0), + ).SfFFT(_ctx(d), max_res=2.0), "SfFFT", ), - # D4: explicit device disagreeing with the supplied cell. + # D4: explicit device disagreeing with the supplied context. DeviceCase( "SfFFT_explicit_device", lambda d: __import__( "torchref.model.sf_fft", fromlist=["SfFFT"] - ).SfFFT(cell=_cell("cpu"), spacegroup=_SG, max_res=2.0, device=d), + ).SfFFT(_ctx("cpu"), max_res=2.0, device=d), + "SfFFT", + ), + # The grid buffers are derived on first use; this case has them resolved. + DeviceCase( + "SfFFT_with_grid", + _sffft_with_grid, "SfFFT", ), DeviceCase( "SfDS_from_cell", lambda d: __import__( "torchref.model.sf_ds", fromlist=["SfDS"] - ).SfDS(cell=_cell(d), spacegroup=_SG), + ).SfDS(_ctx(d)), "SfDS", ), # D1: tensor-free shells, whose tracker is the only thing to check. diff --git a/tests/integration/test_sfds_device.py b/tests/integration/test_sfds_device.py index 7a6ac700..8a1d0806 100644 --- a/tests/integration/test_sfds_device.py +++ b/tests/integration/test_sfds_device.py @@ -29,8 +29,12 @@ def _atoms(device, n=8): @pytest.mark.integration def test_sfds_same_device_cpu(): """Sanity: hkl already on the module device works and stays on it.""" + from torchref.model.context import ModelContext + from torchref.symmetry import SpaceGroup + cell = Cell(_CELL, device="cpu") - sf = SfDS(cell, spacegroup="P212121").to("cpu") + ctx = ModelContext(cell=cell, spacegroup=SpaceGroup("P212121", device="cpu")) + sf = SfDS(ctx).to("cpu") xyz, adp, occ, A, B = _atoms("cpu") hkl = torch.randint(-6, 7, (50, 3)).float() F, _ = sf.compute_structure_factors(hkl, xyz, adp, occ, A, B) @@ -42,9 +46,13 @@ def test_sfds_same_device_cpu(): @pytest.mark.integration def test_sfds_hkl_on_different_device(): """hkl on CPU while the module + atoms are on CUDA must not crash.""" + from torchref.model.context import ModelContext + from torchref.symmetry import SpaceGroup + cuda = torch.device("cuda") cell = Cell(_CELL, device=cuda) - sf = SfDS(cell, spacegroup="P212121").to(cuda) + ctx = ModelContext(cell=cell, spacegroup=SpaceGroup("P212121", device=cuda)) + sf = SfDS(ctx).to(cuda) xyz, adp, occ, A, B = _atoms(cuda) hkl_cpu = torch.randint(-6, 7, (50, 3)).float() # deliberately on CPU diff --git a/tests/unit/model/test_sf_grid_key.py b/tests/unit/model/test_sf_grid_key.py new file mode 100644 index 00000000..c43b71df --- /dev/null +++ b/tests/unit/model/test_sf_grid_key.py @@ -0,0 +1,212 @@ +"""The FFT grid is derived from the model's context and cached on a value key.""" + +import pytest +import torch + +from torchref.config import dtypes +from torchref.model import ModelFT +from torchref.model.context import ModelContext +from torchref.model.sf_fft import SfFFT +from torchref.symmetry import Cell, SpaceGroup + + +@pytest.fixture(scope="module") +def pdb_path(pdb_dir): + path = pdb_dir / "1DAW.pdb" + if not path.exists(): + pytest.skip("1DAW.pdb fixture not present") + return str(path) + + +def _model(pdb_path, **kwargs) -> ModelFT: + return ModelFT(max_res=2.5, verbose=0, device="cpu", **kwargs).load_pdb(pdb_path) + + +def _hkl(model, n=64): + gen = torch.Generator().manual_seed(0) + return torch.randint(-8, 9, (n, 3), generator=gen).to( + dtype=dtypes.int, device=model.device + ) + + +def _count_calls(monkeypatch, cls, name): + counter = {"n": 0} + original = getattr(cls, name) + + def wrapped(self, *args, **kwargs): + counter["n"] += 1 + return original(self, *args, **kwargs) + + monkeypatch.setattr(cls, name, wrapped) + return counter + + +@pytest.mark.unit +@pytest.mark.parametrize("strip_H", [False, True]) +def test_one_engine_and_one_spacegroup_per_load(pdb_path, monkeypatch, strip_H): + import torchref.model.model_ft as model_ft_module + + engines = _count_calls(monkeypatch, model_ft_module.SfFFT, "__init__") + spacegroups = _count_calls(monkeypatch, SpaceGroup, "__init__") + + model = ModelFT(max_res=2.5, verbose=0, device="cpu", strip_H=strip_H) + model.load_pdb(pdb_path) + assert engines["n"] == 1 + assert spacegroups["n"] == 1 + + # Later crystal changes are followed by the key, not by rebuilding the engine. + model.cell = model.cell.clone() + model.max_res = 3.0 + assert model.grid_shape is not None + assert engines["n"] == 1 + + +@pytest.mark.unit +def test_engine_reads_the_model_context(pdb_path): + model = _model(pdb_path) + + def bound(m): + return ( + m.fft.ctx is m.ctx + and m.fft.cell is m.ctx.cell + and m.fft.spacegroup is m.ctx.spacegroup + ) + + assert bound(model) + + copied = model.copy() + assert bound(copied) + assert copied.ctx is not model.ctx + assert copied.fft.cell is not model.fft.cell + + selected = model.select("all") + assert bound(selected) + + restored = ModelFT.create_from_state_dict(model.state_dict(), device="cpu", verbose=0) + assert bound(restored) + assert restored.grid_shape == model.grid_shape + + +@pytest.mark.unit +def test_explicit_gridsize_survives_every_path(pdb_path): + explicit = (64, 32, 24) + model = _model(pdb_path, gridsize=explicit) + assert model.grid_shape == explicit + + model.cell = model.cell.clone() + assert model.grid_shape == explicit + + assert model.copy().grid_shape == explicit + assert model.select("all").grid_shape == explicit + + restored = ModelFT.create_from_state_dict(model.state_dict(), device="cpu", verbose=0) + assert restored.explicit_gridsize == explicit + assert restored.grid_shape == explicit + + +@pytest.mark.unit +def test_grid_follows_its_key(pdb_path): + model = _model(pdb_path) + shape0 = model.grid_shape + key0 = model.grid_key + ptr0 = model.fft.gridsize.data_ptr() + + # Same values, different object: nothing to rebuild. + model.cell = model.cell.clone() + assert model.grid_key == key0 + assert model.fft.gridsize.data_ptr() == ptr0 + + scale = torch.tensor([1.25, 1.25, 1.25, 1.0, 1.0, 1.0]) + model.cell = Cell(model.cell.data * scale, dtype=model.dtype_float, device="cpu") + assert model.grid_key != key0 + assert model.grid_shape != shape0 + assert all(n > m for n, m in zip(model.grid_shape, shape0)) + + model.max_res = 4.0 + coarse = model.grid_shape + assert all(n < m for n, m in zip(coarse, model.grid_shape)) is False + model.max_res = 2.5 + assert all(n > m for n, m in zip(model.grid_shape, coarse)) + + +@pytest.mark.unit +def test_first_fcalc_after_cell_reassignment_uses_the_late_path(pdb_path, monkeypatch): + model = _model(pdb_path) + hkl = _hkl(model) + with torch.no_grad(): + f0 = model(hkl).clone() + + model.cell = model.cell.clone() + model.reset_cache() + assert model.fft.late_symmetry_compatible is True + + symmetrised = _count_calls(monkeypatch, SpaceGroup, "symmetrize_map") + with torch.no_grad(): + f1 = model(hkl) + assert symmetrised["n"] == 0 + assert torch.allclose(f1, f0) + + +@pytest.mark.unit +def test_forward_cache_invalidates_on_a_spacegroup_change(pdb_path): + model = _model(pdb_path) + hkl = _hkl(model) + with torch.no_grad(): + f0 = model(hkl) + assert model(hkl) is f0 # cached + assert model._fwd_cached_state_fp[-1] == model.grid_key + + model.spacegroup = "P 1 2 1" # same cell, centring dropped + f1 = model(hkl) + assert f1 is not f0 + assert not torch.allclose(f1, f0) + assert model._fwd_cached_state_fp[-1] == model.grid_key + + +@pytest.mark.unit +def test_to_moves_the_shared_context_once(pdb_path, monkeypatch): + model = _model(pdb_path) + cell = model.ctx.cell + resets = {"n": 0} + original = Cell.reset_cache + + def counting(self): + if self is cell: + resets["n"] += 1 + return original(self) + + monkeypatch.setattr(Cell, "reset_cache", counting) + model.to(torch.device("cpu")) + assert resets["n"] == 1 + + +@pytest.mark.unit +def test_engine_without_a_crystal(): + sf = SfFFT(ModelContext(), max_res=1.0) + assert sf.gridsize is None + assert sf.grid_shape is None + assert sf.grid_key is None + + empty = torch.zeros((0, 3)) + with pytest.raises(RuntimeError, match="no cell or space group"): + sf.compute_structure_factors( + torch.zeros((1, 3), dtype=dtypes.int), + empty, torch.zeros(0), torch.zeros(0), torch.zeros((0, 5)), torch.zeros((0, 5)), + ) + + +@pytest.mark.unit +def test_legacy_gridsize_is_adopted_only_when_it_differs(pdb_path): + model = _model(pdb_path) + + same = model.state_dict() + same["_fft.gridsize"] = torch.tensor(model.grid_shape) + restored = ModelFT.create_from_state_dict(same, device="cpu", verbose=0) + assert restored.explicit_gridsize is None + assert restored.grid_shape == model.grid_shape + + different = model.state_dict() + different["_fft.gridsize"] = torch.tensor([64, 32, 24]) + restored = ModelFT.create_from_state_dict(different, device="cpu", verbose=0) + assert restored.explicit_gridsize == (64, 32, 24) + assert restored.grid_shape == (64, 32, 24) diff --git a/tests/unit/refinement/test_wilson_prior.py b/tests/unit/refinement/test_wilson_prior.py index a3d008ba..75635189 100644 --- a/tests/unit/refinement/test_wilson_prior.py +++ b/tests/unit/refinement/test_wilson_prior.py @@ -32,7 +32,7 @@ def setup_target(): ) ens.cell = data.cell ens.spacegroup = data.spacegroup - ens.setup_grid(max_res=data.get_max_res()) + ens.max_res = data.get_max_res() scaler = Scaler(model=ens, data=data, nbins=10, verbose=0) fcalc0 = ens(data.hkl) diff --git a/tests/unit/structure_factor/helpers.py b/tests/unit/structure_factor/helpers.py index aae056b3..fcebf894 100644 --- a/tests/unit/structure_factor/helpers.py +++ b/tests/unit/structure_factor/helpers.py @@ -442,17 +442,18 @@ def sf_fft_for( Pass ``fineness=1.0`` for a deliberately under-sampled grid; see :data:`GRID_FINENESS` for why that is the sampling-limited regime. """ + from torchref.model.context import ModelContext from torchref.model.sf_fft import SfFFT - - sf = SfFFT( - cell=scene.cell, - spacegroup=spacegroup, - max_res=scene.d_min / fineness, - dtype_float=dtype, - device=torch.device("cpu"), + from torchref.symmetry import Cell, SpaceGroup + + cpu = torch.device("cpu") + # A private context: the engine reads the crystal live, so it must not share + # the module-scoped scene's cell with other tests. + ctx = ModelContext( + cell=Cell(scene.cell.data, dtype=dtype, device=cpu), + spacegroup=SpaceGroup(spacegroup, dtype=dtype, device=cpu), ) - sf.setup_grid() - return sf + return SfFFT(ctx, max_res=scene.d_min / fineness, dtype_float=dtype, device=cpu) # --------------------------------------------------------------------------- diff --git a/tests/unit/structure_factor/test_dispatch.py b/tests/unit/structure_factor/test_dispatch.py index 1dcbada2..c78e14f3 100644 --- a/tests/unit/structure_factor/test_dispatch.py +++ b/tests/unit/structure_factor/test_dispatch.py @@ -389,9 +389,16 @@ def test_sfds_backend_toggle_end_to_end(scene_fine): s = scene_fine.to(device=cuda, dtype=torch.float32) obs = H.synthetic_obs(H.ds_direct(scene_fine, "eager").detach()).to(cuda, torch.float32) + from torchref.model.context import ModelContext + from torchref.symmetry import SpaceGroup + def run(force_portable): + ctx = ModelContext( + cell=s.cell, + spacegroup=SpaceGroup("P212121", dtype=torch.float32, device=cuda), + ) sf = SfDS( - cell=s.cell, spacegroup="P212121", force_portable=force_portable, + ctx, force_portable=force_portable, dtype_float=torch.float32, device=cuda, max_memory_gb=2.0, ) leaves = tuple(t.clone().requires_grad_(True) for t in (s.xyz, s.adp, s.occ)) diff --git a/tests/unit/structure_factor/test_forward.py b/tests/unit/structure_factor/test_forward.py index 78fe6915..ecd0b065 100644 --- a/tests/unit/structure_factor/test_forward.py +++ b/tests/unit/structure_factor/test_forward.py @@ -252,12 +252,16 @@ def test_sfds_matches_gemmi_with_symmetry(gemmi_iso_symmetry): assert len(structure.cell.images) > 0, "structure was not set up with symmetry" F_gemmi = H.gemmi_sf(structure, scene.hkl_list) - ds = SfDS( + from torchref.model.context import ModelContext + from torchref.symmetry import SpaceGroup + + ctx = ModelContext( cell=scene.cell, - spacegroup=scene.spacegroup, - dtype_float=torch.float64, - device=torch.device("cpu"), + spacegroup=SpaceGroup( + scene.spacegroup, dtype=torch.float64, device=torch.device("cpu") + ), ) + ds = SfDS(ctx, dtype_float=torch.float64, device=torch.device("cpu")) with torch.no_grad(): F_sym, _ = ds.compute_structure_factors( scene.hkl, scene.xyz, scene.adp, scene.occ, scene.A, scene.B, diff --git a/tests/unit/symmetry/test_cell_identity.py b/tests/unit/symmetry/test_cell_identity.py new file mode 100644 index 00000000..9132d751 --- /dev/null +++ b/tests/unit/symmetry/test_cell_identity.py @@ -0,0 +1,102 @@ +"""Value identity of :class:`~torchref.symmetry.cell.Cell`: ``key``, ``__eq__``, ``__hash__``.""" + +import pytest +import torch + +from torchref.symmetry import Cell + +PARAMS = [50.0, 60.0, 70.0, 90.0, 90.0, 90.0] + + +@pytest.mark.unit +def test_equal_parameters_compare_and_hash_equal(): + a = Cell(PARAMS, dtype=torch.float32, device="cpu") + b = Cell(PARAMS, dtype=torch.float32, device="cpu") + assert a is not b + assert a == b + assert hash(a) == hash(b) + assert a.key == tuple(PARAMS) + + +@pytest.mark.unit +def test_clone_compares_equal(): + a = Cell(PARAMS) + assert a.clone() == a + assert hash(a.clone()) == hash(a) + + +@pytest.mark.unit +def test_different_parameters_compare_unequal(): + a = Cell(PARAMS) + perturbed = list(PARAMS) + perturbed[0] += 1e-3 + b = Cell(perturbed) + assert a != b + assert a.key != b.key + + +@pytest.mark.unit +def test_comparison_with_other_types_is_false(): + a = Cell(PARAMS) + assert (a == "50 60 70 90 90 90") is False + assert (a == PARAMS) is False + assert a != None # noqa: E711 + + +@pytest.mark.unit +def test_usable_as_dict_key_and_in_set(): + a = Cell(PARAMS) + b = Cell(PARAMS) + cache = {a: "grid"} + assert cache[b] == "grid" + assert len({a, b}) == 1 + + +@pytest.mark.unit +def test_key_is_cached_and_reset_with_the_other_derived_quantities(): + cell = Cell(PARAMS, dtype=torch.float32, device="cpu") + assert "key" not in cell._cache + first = cell.key + assert cell._cache["key"] is first + + cell.to(dtype=torch.float64) + assert "key" not in cell._cache, "reset_cache must drop the key with the rest" + assert cell.key == first + + +@pytest.mark.unit +def test_in_place_edit_is_refused_at_the_next_derived_read(): + cell = Cell(PARAMS) + _ = cell.volume + cell.data[0] = 51.0 + with pytest.raises(RuntimeError, match="create a new one"): + cell.fractional_matrix + with pytest.raises(RuntimeError, match="Please don't edit Cell objects"): + cell.key + + +@pytest.mark.unit +def test_in_place_edit_before_any_read_is_refused_too(): + cell = Cell(PARAMS) + cell.data.mul_(2.0) + with pytest.raises(RuntimeError, match="edited in place"): + cell.volume + + +@pytest.mark.unit +def test_constructor_owns_its_tensor(): + t = torch.tensor(PARAMS, dtype=torch.float32) + cell = Cell(t, dtype=torch.float32, device="cpu") + t[0] = 99.0 + assert cell.key[0] == 50.0 + assert float(cell.volume) == pytest.approx(210000.0) + + +@pytest.mark.unit +def test_device_and_dtype_moves_are_not_edits(): + cell = Cell(PARAMS, dtype=torch.float32, device="cpu") + _ = cell.fractional_matrix + cell.to(dtype=torch.float64) + assert cell.fractional_matrix.dtype == torch.float64 + assert cell.key == tuple(PARAMS) + assert cell.clone() == cell diff --git a/tests/unit/symmetry/test_spacegroup_identity.py b/tests/unit/symmetry/test_spacegroup_identity.py new file mode 100644 index 00000000..aaa21eee --- /dev/null +++ b/tests/unit/symmetry/test_spacegroup_identity.py @@ -0,0 +1,52 @@ +"""Value identity of :class:`~torchref.symmetry.spacegroup.SpaceGroup` and the +setting-preserving round trip through gemmi.""" + +import gemmi +import pytest + +from torchref.symmetry import SpaceGroup + + +@pytest.mark.unit +def test_same_number_different_setting_are_unequal(): + a = SpaceGroup("P 1 21 1") + b = SpaceGroup("P 1 1 21") + assert a.number == b.number == 4 + assert a != b + assert hash(a) != hash(b) + # The reason the identity must be setting-aware: the screw axis moves. + assert a.grid_requirements() != b.grid_requirements() + + +@pytest.mark.unit +def test_same_setting_compares_and_hashes_equal(): + a = SpaceGroup("P 21 21 21") + b = SpaceGroup(19) + assert a == b + assert hash(a) == hash(b) + assert a.key == b.key == a.xhm + + +@pytest.mark.unit +def test_copy_is_equal(): + a = SpaceGroup("C 1 2 1") + assert a.copy() == a + assert hash(a.copy()) == hash(a) + + +@pytest.mark.unit +def test_equality_with_gemmi_spacegroup_is_setting_aware(): + a = SpaceGroup("P 1 21 1") + assert a == gemmi.find_spacegroup_by_name("P 1 21 1") + assert a != gemmi.find_spacegroup_by_name("P 1 1 21") + assert (a == "P 1 21 1") is False + + +@pytest.mark.unit +@pytest.mark.parametrize("xhm", ["R 3:R", "R 3:H", "P 4/n:1", "P 4/n:2"]) +def test_rewrapping_preserves_the_setting(xhm): + sg = SpaceGroup(xhm) + assert sg.xhm == xhm + assert SpaceGroup(sg).xhm == xhm + assert sg._gemmi.xhm() == xhm + assert SpaceGroup(sg).n_ops == sg.n_ops diff --git a/tests/unit/utils/test_device_resolution.py b/tests/unit/utils/test_device_resolution.py index 4dab9045..6567108c 100644 --- a/tests/unit/utils/test_device_resolution.py +++ b/tests/unit/utils/test_device_resolution.py @@ -205,9 +205,14 @@ def test_sfds_refuses_a_cell_recast_after_construction(): """ from torchref.model.sf_ds import SfDS + from torchref.model.context import ModelContext + from torchref.symmetry import SpaceGroup + cell = Cell([50.0, 60.0, 70.0, 90.0, 90.0, 90.0], dtype=torch.float32, device="cpu") - sf = SfDS(cell=cell, spacegroup="P 1", dtype_float=torch.float32, - device=torch.device("cpu")) + ctx = ModelContext( + cell=cell, spacegroup=SpaceGroup("P 1", dtype=torch.float32, device="cpu") + ) + sf = SfDS(ctx, dtype_float=torch.float32, device=torch.device("cpu")) xyz = torch.zeros(3, 3, dtype=torch.float32) sf._cartesian_to_fractional(xyz) # consistent: fine diff --git a/torchref/experimental/alignment/rigid_body.py b/torchref/experimental/alignment/rigid_body.py index 9d5350f5..b1522b35 100644 --- a/torchref/experimental/alignment/rigid_body.py +++ b/torchref/experimental/alignment/rigid_body.py @@ -163,7 +163,11 @@ def __init__( self.cell = data.cell self.spacegroup = data.spacegroup - self.fft = SfFFT(self.cell, self.spacegroup, max_res=max_res) + from torchref.model.context import ModelContext + + self.fft = SfFFT( + ModelContext(cell=self.cell, spacegroup=self.spacegroup), max_res=max_res + ) self.verbose = verbose self.rfactor_converged_threshold = rfactor_converged_threshold diff --git a/torchref/experimental/ensemble/ensemble_model.py b/torchref/experimental/ensemble/ensemble_model.py index 7f0ed418..9b6f7931 100644 --- a/torchref/experimental/ensemble/ensemble_model.py +++ b/torchref/experimental/ensemble/ensemble_model.py @@ -38,7 +38,7 @@ from __future__ import annotations -from typing import Optional +from typing import Optional, Tuple import numpy as np import pandas as pd @@ -345,7 +345,7 @@ def __init__( # load would invalidate that. Off by default here, unlike on the base class. add_hydrogens: bool = False, max_res: float = 1.0, - gridsize: Optional[int] = None, + gridsize: Optional[Tuple[int, int, int]] = None, wavelength: float = 1.0, anomalous_threshold: float = 0.5, ): diff --git a/torchref/experimental/ensemble/ensemble_refinement.py b/torchref/experimental/ensemble/ensemble_refinement.py index eee55188..a36563e7 100644 --- a/torchref/experimental/ensemble/ensemble_refinement.py +++ b/torchref/experimental/ensemble/ensemble_refinement.py @@ -488,7 +488,7 @@ def __init__( ) self.model.cell = self.reflection_data.cell self.model.spacegroup = self.reflection_data.spacegroup - self.model.setup_grid(max_res=self.max_res) + self.model.max_res = self.max_res # Rebuild the scaler against the ensemble model. self.scaler = Scaler( diff --git a/torchref/experimental/monolithic_refinement/density_solvent.py b/torchref/experimental/monolithic_refinement/density_solvent.py index e4f4668c..220b33ac 100644 --- a/torchref/experimental/monolithic_refinement/density_solvent.py +++ b/torchref/experimental/monolithic_refinement/density_solvent.py @@ -32,6 +32,7 @@ ifft, ) from torchref.config import get_default_device, get_float_dtype +from torchref.model.context import ModelContext from torchref.model.sf_fft import SfFFT from torchref.utils.debug_utils import DebugMixin from torchref.utils.device_mixin import DeviceMixin @@ -195,16 +196,20 @@ def __init__( # because the nonlinear occupancy needs the full-cell density assembled # before the mask. The per-atom splat radius is governed by # torchref.sigma_cutoff_ed inside the density builder. + # Its own context: the cell is shared with the model, but the space group + # is copied because it memoises operators per grid shape and this engine's + # coarse grid must not evict the model's. + solvent_ctx = ModelContext( + cell=model.cell, spacegroup=model.spacegroup.copy() + ) self.solvent_fft = SfFFT( - cell=model.cell, - spacegroup=model.fft.spacegroup, + ctx=solvent_ctx, max_res=self.solvent_res, dtype_float=float_type, device=device, verbose=max(0, verbose - 1), use_late_symmetry=False, ) - self.solvent_fft.setup_grid() # ------------------------------------------------------------------ # Density -> occupancy -> structure factor diff --git a/torchref/experimental/targets/realspace.py b/torchref/experimental/targets/realspace.py index fe95a0e5..2024b595 100644 --- a/torchref/experimental/targets/realspace.py +++ b/torchref/experimental/targets/realspace.py @@ -111,20 +111,12 @@ def __init__( # Caches (not registered as buffers since they're lazily computed) self._data_p1 = None self._molecular_mask = None - self._gridsize = None # P1 expansion cache (ASU → P1 mapping) self._hkl_p1 = None self._p1_indices = None self._p1_phase_shifts = None - def _ensure_grid(self): - """Ensure model's SfFFT grid is set up.""" - if self._model is None: - raise RuntimeError("No model set for RealSpaceTarget") - if self._model.gridsize is None: - self._model.setup_grid() - def _get_data_p1(self) -> "ReflectionData": """Return P1-expanded ReflectionData, cached after first call.""" if self._data_p1 is None: @@ -153,19 +145,11 @@ def _expand_to_p1(self, fcalc: torch.Tensor) -> torch.Tensor: return fcalc_p1 * torch.exp(1j * self._p1_phase_shifts) def _get_gridsize(self) -> Tuple[int, int, int]: - """ - Get grid size for map computation. - - Uses the model's FFT grid size to ensure compatibility with - the molecular mask (which is built on the model's grid). - """ - if self._gridsize is not None: - return self._gridsize - - self._ensure_grid() - gs = self._model.fft.gridsize - self._gridsize = tuple(int(x) for x in gs) - return self._gridsize + """Grid size for map computation: the model's, so it matches the + molecular mask built on the model's grid.""" + if self._model is None: + raise RuntimeError("No model set for RealSpaceTarget") + return self._model.fft.grid_shape def _compute_observed_map(self) -> torch.Tensor: """ @@ -243,7 +227,6 @@ def _build_molecular_mask(self): """ from torchref.scaling.solvent import SolventModel - self._ensure_grid() with torch.no_grad(): solvent = SolventModel( diff --git a/torchref/model/__init__.py b/torchref/model/__init__.py index 27ad2f14..58f91bc5 100644 --- a/torchref/model/__init__.py +++ b/torchref/model/__init__.py @@ -3,18 +3,17 @@ :class:`Model` holds the refinable atomic parameters, with the crystallographic context, atom table and provenance split out into :class:`ModelContext`; :class:`ModelFT` adds -structure-factor calculation on top, through :class:`SfFFT` (FFT) or +structure-factor calculation on top, through :class:`SfFFT` or :class:`SfDS` (direct summation). :class:`MixedModel` combines ModelFT states by population fraction (e.g. dark/light), and :class:`ModelCollection` keys mixtures by timepoint (``_SharedMixedModel`` is its non-re-registering variant). The wrappers from :mod:`torchref.model.parameter_wrappers` -- :class:`MixedTensor` and its ``Positive`` / ``Cholesky`` / ``Occupancy`` subclasses plus :class:`RigidXYZTensor` -- are the parametrizations that decide -which parameters are refinable. ``FFT`` is a deprecated alias for -:class:`SfFFT`. +which parameters are refinable. """ -from torchref.model.sf_fft import SfFFT, FFT +from torchref.model.sf_fft import SfFFT from torchref.model.sf_ds import SfDS from torchref.model.context import ModelContext from torchref.model.mixed_model import MixedModel @@ -31,7 +30,6 @@ from torchref.model.rigid_xyz import RigidXYZTensor __all__ = [ - "FFT", "SfFFT", "SfDS", "MixedModel", diff --git a/torchref/model/context.py b/torchref/model/context.py index 5c449019..17f32148 100644 --- a/torchref/model/context.py +++ b/torchref/model/context.py @@ -115,6 +115,20 @@ def copy(self) -> "ModelContext": initialized=self.initialized, ) + @property + def crystal_key(self): + """Value identity of the crystal, or None while cell or space group is unset. + + Returns + ------- + tuple or None + ``(cell.key, spacegroup.key)``; hashable, so anything derived from the + crystal alone can be cached against it. + """ + if self.cell is None or self.spacegroup is None: + return None + return (self.cell.key, self.spacegroup.key) + def __repr__(self) -> str: n_atoms = 0 if self.pdb is None else len(self.pdb) sg = None if self.spacegroup is None else self.spacegroup.name diff --git a/torchref/model/model.py b/torchref/model/model.py index 7983ef68..484b8484 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -256,14 +256,25 @@ def spacegroup(self) -> Optional[SpaceGroup]: @spacegroup.setter def spacegroup(self, value): - """Set the space group from a SpaceGroup, gemmi object, name or number.""" - if value is not None: + """Set the space group from a SpaceGroup, gemmi object, name or number. + + The model owns its space group: an incoming ``SpaceGroup`` is copied rather + than shared, because ``.to()`` moves in place and would otherwise relocate + the caller's object. The copy lands on the model's device and float dtype. + """ + if value is None: + self.ctx.spacegroup = None + elif isinstance(value, SpaceGroup): + self.ctx.spacegroup = value.copy().to( + device=self.device, dtype=self.dtype_float + ) + else: # ``device=self.device``: SpaceGroup falls back to the global # default otherwise, so setting a spacegroup on a CPU-pinned Model # would silently plant accelerator-resident matrices on it. - self.ctx.spacegroup = SpaceGroup(value, device=self.device) - else: - self.ctx.spacegroup = None + self.ctx.spacegroup = SpaceGroup( + value, dtype=self.dtype_float, device=self.device + ) # ========================================================================= # Crystallographic matrix properties (delegated to Cell) @@ -1964,16 +1975,18 @@ def _new_model_from_df(self, df, *, strip_H=None, add_hydrogens=False): continue if param.kind in (param.VAR_POSITIONAL, param.VAR_KEYWORD): continue + if pname == "gridsize": + # The constructor argument is the explicit override, not the + # derived grid a ``gridsize`` attribute would return. + if hasattr(self, "explicit_gridsize"): + ctor_kw[pname] = self.explicit_gridsize + continue if hasattr(self, pname): ctor_kw[pname] = getattr(self, pname) - if "gridsize" in sig.parameters and hasattr(self, "_explicit_gridsize"): - ctor_kw["gridsize"] = self._explicit_gridsize new_model = self.__class__(**ctor_kw) sg_str = self.spacegroup.xhm if self.spacegroup else "P 1" new_model.load(lambda: (df, self.pdb.attrs.get("cell"), sg_str)) - if hasattr(new_model, "setup_grid"): - new_model.setup_grid() # Propagate CIF restraint paths so restraints are rebuilt correctly if self.ctx.cif_path is not None: new_model._cif_path = self.ctx.cif_path diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index c6917cba..f95289d5 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -1,8 +1,8 @@ """ModelFT -- a :class:`~torchref.model.Model` that can compute structure factors. -Adds the electron-density / FFT path (via an :class:`~torchref.model.SfFFT` -submodule created as soon as both cell and space group are set), the ITC92 -scattering parametrization, and the anomalous f' / f'' correction. +Adds the electron-density / FFT path (an :class:`~torchref.model.SfFFT` submodule +that reads the crystal off the model's context and sizes its grid lazily), the +ITC92 scattering parametrization, and the anomalous f' / f'' correction. """ import math @@ -52,10 +52,12 @@ class ModelFT(CachedForwardMixin, Model): ---------- max_res, wavelength, anomalous_threshold : float The constructor arguments above, readable back as attributes. - gridsize : torch.Tensor - Grid dimensions ``(nx, ny, nz)``, living on the ``SfFFT`` submodule. - A coordinate grid is not stored; :meth:`real_space_grid` builds one on - demand for the few callers that want the Cartesian positions themselves. + gridsize : torch.Tensor or None + Grid dimensions ``(nx, ny, nz)``, derived by the ``SfFFT`` submodule from + the cell, space group, ``max_res`` and ``explicit_gridsize`` on first use + and re-derived when any of them changes. A coordinate grid is not stored; + :meth:`real_space_grid` builds one on demand for the few callers that want + the Cartesian positions themselves. map : torch.Tensor or None Most recently computed electron density map. parametrization : dict @@ -107,8 +109,17 @@ def __init__( """ super().__init__(*args, **kwargs) - self.max_res = max_res - self._explicit_gridsize = gridsize + # The engine reads cell and space group off ``self.ctx`` as they are set; + # its grid is derived on first use and re-derived when the crystal, + # ``max_res`` or ``explicit_gridsize`` change. + self._fft = SfFFT( + ctx=self.ctx, + max_res=max_res, + explicit_gridsize=gridsize, + dtype_float=self.dtype_float, + device=self.device, + verbose=self.ctx.verbose, + ) self.wavelength = wavelength self.anomalous_threshold = anomalous_threshold @@ -124,46 +135,53 @@ def __init__( self._anomalous_elements_hash = ( None # Hash of element list for cache invalidation ) - self._fft = None + + # ========================================================================= + # Engine binding and grid inputs + # ========================================================================= + + @property + def fft(self) -> SfFFT: + """The SfFFT submodule, bound to this model's context. + + ``copy()`` and ``load_state`` replace the context object itself; re-pointing + the engine here keeps ``fft.ctx is self.ctx`` on every path. + """ + fft = self._fft + if fft.ctx is not self.ctx: + fft.ctx = self.ctx + return fft @property - def cell(self): - """Unit cell object with parameters [a, b, c, alpha, beta, gamma].""" - return self.ctx.cell + def max_res(self) -> Optional[float]: + """Maximum resolution in Angstroms that sizes the grid; owned by the engine.""" + return self._fft.max_res - @cell.setter - def cell(self, value): - """Set the unit cell; also builds the FFT once the spacegroup is set.""" - self.ctx.cell = value - self._maybe_initialize_fft() + @max_res.setter + def max_res(self, value) -> None: + self._fft.max_res = None if value is None else float(value) @property - def spacegroup(self): - """Space group object.""" - return self.ctx.spacegroup - - @spacegroup.setter - def spacegroup(self, value): - """Set the space group (SpaceGroup, gemmi.SpaceGroup, name or number); - also builds the FFT once the cell is set. + def explicit_gridsize(self) -> Optional[Tuple[int, int, int]]: + """Fixed grid dimensions overriding ``max_res``, or None.""" + return self._fft.explicit_gridsize + + @explicit_gridsize.setter + def explicit_gridsize(self, value) -> None: + self._fft.explicit_gridsize = value + + @property + def grid_key(self): + """What the grid is derived from; see :attr:`SfFFT.grid_key`.""" + return self.fft.grid_key + + def _fingerprint_state(self): + """Fold the grid key into the forward-cache key. + + Parameters and buffers alone would miss a cell, space-group or resolution + change that leaves the grid buffers untouched until the next forward. """ - if value is not None: - self.ctx.spacegroup = SpaceGroup( - value, dtype=self.dtype_float, device=self.device - ) - else: - self.ctx.spacegroup = None - self._maybe_initialize_fft() - - def _maybe_initialize_fft(self): - """(Re)build the SfFFT submodule once both cell and spacegroup are set.""" - if self.ctx.cell is not None and self.ctx.spacegroup is not None: - self._fft = SfFFT( - cell=self.ctx.cell, - spacegroup=self.ctx.spacegroup, - device=self.device, - max_res=self.max_res, - ) + return super()._fingerprint_state() + (self.fft.grid_key,) def load_pdb(self, filename): """ @@ -180,8 +198,6 @@ def load_pdb(self, filename): Self, for method chaining. """ super().load_pdb(filename) - # FFT is now initialized via cell/spacegroup setters in parent load() - self.setup_grid() return self def select(self, selection): @@ -189,9 +205,8 @@ def select(self, selection): Return a new ModelFT containing only the selected atoms. Extends :meth:`Model.select` with the FT-specific setup: rebuilding - the ITC92 parametrization and the real-space grid for the reduced - atom set. The FFT itself is initialized via the cell/spacegroup - setters during the base ``select``. + the ITC92 parametrization and carrying ``max_res`` and + ``explicit_gridsize`` across, so the selection sizes its grid the same way. Parameters ---------- @@ -205,15 +220,14 @@ def select(self, selection): Notes ----- - The ModelFT-specific constructor arguments -- ``max_res``, - ``wavelength``, ``anomalous_threshold``, ``gridsize`` -- are **not** - propagated: :meth:`Model.select` passes only the base kwargs, so the - returned model silently carries the ModelFT defaults for those. + ``wavelength`` and ``anomalous_threshold`` are **not** propagated: + :meth:`Model.select` passes only the base kwargs, so the returned model + carries the ModelFT defaults for those. """ selection = super().select(selection) selection._build_parametrization() - # FFT is initialized via cell/spacegroup setters in parent select() - selection.setup_grid() + selection.max_res = self.max_res + selection.explicit_gridsize = self.explicit_gridsize return selection def load_cif(self, filename): @@ -232,36 +246,8 @@ def load_cif(self, filename): """ super().load_cif(filename) self._build_parametrization() - # FFT is now initialized via cell/spacegroup setters in parent load() - self.setup_grid() return self - def setup_gridsize(self, max_res=None): - """ - Compute optimal grid dimensions. - - Delegates to FFT.compute_grid_size(). - - Parameters - ---------- - max_res : float, optional - Maximum resolution in Angstroms. If None, uses self.max_res. - - Returns - ------- - torch.Tensor - Grid dimensions (nx, ny, nz) as int32 tensor. - """ - if max_res is not None: - self.max_res = max_res - self._fft.max_res = max_res - - if self.ctx.verbose > 1: - print(f"Defining grid size for max_res={self.max_res} Å") - - gridsize = self.cell.compute_grid_size(self.max_res) - return torch.tensor(gridsize, dtype=dtypes.int, device=self.device) - def _build_parametrization(self): """Build the ITC92 parametrization (delegates to :class:`Model`).""" return super()._build_parametrization() @@ -283,18 +269,13 @@ def B(self) -> torch.Tensor: return self._B # ========================================================================= - # Backward-compatible properties for FFT grid attributes + # Grid, resolved by the engine # ========================================================================= @property def gridsize(self) -> Optional[torch.Tensor]: - """Grid dimensions (nx, ny, nz).""" - return self._fft.gridsize - - @gridsize.setter - def gridsize(self, value): - """Set grid size (for backward compatibility).""" - self._fft.gridsize = value + """Grid dimensions (nx, ny, nz), or None until cell and space group are set.""" + return self.fft.gridsize def real_space_grid(self) -> torch.Tensor: """Build the Cartesian coordinate of every grid point, ``(nx, ny, nz, 3)``. @@ -306,8 +287,6 @@ def real_space_grid(self) -> torch.Tensor: """ from torchref.base.fourier import get_real_grid - if self.gridsize is None: - self.setup_grid() return get_real_grid( fractional_matrix=self.cell.fractional_matrix, gridsize=self.gridsize, @@ -316,18 +295,13 @@ def real_space_grid(self) -> torch.Tensor: @property def grid_shape(self) -> Optional[tuple]: - """Map dimensions ``(nx, ny, nz)``, or ``None`` before the grid is set up.""" - return self._fft.grid_shape + """Map dimensions ``(nx, ny, nz)``, or None until cell and space group are set.""" + return self.fft.grid_shape @property def voxel_size(self) -> Optional[torch.Tensor]: - """Voxel dimensions.""" - return self._fft.voxel_size - - @voxel_size.setter - def voxel_size(self, value): - """Set voxel size (for backward compatibility).""" - self._fft.voxel_size = value + """Voxel edge vector sum, or None until cell and space group are set.""" + return self.fft.voxel_size def get_iso(self): """ @@ -381,38 +355,22 @@ def get_aniso(self): return xyz, u, occupancy, A, B - def setup_grid(self, max_res=None, gridsize=None): + def setup_grid(self, *, max_res=None, gridsize=None): """ - Setup real-space grid for electron density calculation. + Override the grid's inputs explicitly and resolve the grid now. - Delegates to FFT.setup_grid() using the stored cell and spacegroup. + Not needed on the normal path: the engine sizes its grid from the cell, + space group and ``max_res`` on first use and follows any later change. Parameters ---------- max_res : float, optional - Maximum resolution for grid spacing in Angstroms. - If None, uses self.max_res. + New maximum resolution in Angstroms. None leaves the current value. gridsize : tuple of int, optional - Explicit grid size (nx, ny, nz). If None, computed automatically - using Cell.compute_grid_size() and SpaceGroup.suggest_grid_size(). + Fixed grid size (nx, ny, nz). None leaves :attr:`explicit_gridsize` + unchanged. """ - if max_res is not None: - self.max_res = max_res - self._fft.max_res = max_res - - if self.ctx.verbose > 1: - print(f"Setting up grids with max_res={self.max_res} Å") - - gridsize_to_use = gridsize or self._explicit_gridsize - - self._fft.setup_grid( - gridsize=gridsize_to_use, - max_res=self.max_res, - ) - - if self.ctx.verbose > 2: - print(f"Grid shape: {self._fft.grid_shape}") - print(f"Voxel size: {self._fft.voxel_size}") + self.fft.setup_grid(max_res=max_res, gridsize=gridsize) def build_complete_map(self, radius=None, apply_symmetry=True): """ @@ -460,9 +418,6 @@ def build_initial_map(self, apply_symmetry=True): torch.Tensor Electron density map with shape (nx, ny, nz). """ - if self._fft.gridsize is None: - self.setup_grid() - if self.ctx.verbose > 2: print("Building density map (per-atom variable radius)...") @@ -722,14 +677,6 @@ def get_structure_factor( """ return self(hkl, recalc=recalc, apply_anomalous=apply_anomalous) - @property - def fft(self): - """The SfFFT submodule, built on first access (needs cell + spacegroup).""" - if self._fft is None: - self._maybe_initialize_fft() - - return self._fft - def _check_forward_dtype(self, hkl: torch.Tensor) -> None: """Fail fast on a model/input float-dtype mismatch, which would otherwise surface as a cryptic matmul or Triton-compile error deep in the kernels. @@ -804,8 +751,9 @@ def copy(self, detach: bool = True) -> "ModelFT": Create a deep copy of the ModelFT. Creates a complete independent copy including all Model base class data, - FFT submodule state (gridsize, voxel_size), - ITC92 parametrization, and scalar attributes. + the grid inputs (``max_res``, ``explicit_gridsize``; the grid itself is + re-derived from the copied context), the ITC92 parametrization, and + scalar attributes. Cache is reset to empty. Parameters @@ -828,7 +776,7 @@ def copy(self, detach: bool = True) -> "ModelFT": device=self.device, strip_H=self.ctx.strip_H, max_res=self.max_res, - gridsize=self._explicit_gridsize, + gridsize=self.explicit_gridsize, wavelength=self.wavelength, anomalous_threshold=self.anomalous_threshold, ) @@ -836,7 +784,7 @@ def copy(self, detach: bool = True) -> "ModelFT": # Carries the atom table, cell, space group, altloc groups and provenance. model_copy.ctx = self.ctx.copy() - # Own buffers only; the FFT submodule's are handled by its copy() below. + # Own buffers only; the engine's grid buffers are derived, not copied. for buffer_name, buffer_value in self._buffers.items(): if buffer_value is not None: if detach: @@ -846,7 +794,7 @@ def copy(self, detach: bool = True) -> "ModelFT": else: model_copy.register_buffer(buffer_name, buffer_value.clone()) - # Parameter wrappers via their own .copy(); the FFT submodule is separate. + # Parameter wrappers via their own .copy(); the engine came from the ctor. skip_modules = {"_fft"} for module_name, module in self._modules.items(): if module_name in skip_modules: @@ -859,11 +807,6 @@ def copy(self, detach: bool = True) -> "ModelFT": model_copy._parametrization = copy_module.deepcopy(self._parametrization) - if self._fft is not None: - model_copy._fft = self._fft.copy() - if self._fft.gridsize is not None: - model_copy.setup_grid(max_res=self.max_res) - # Don't share cached structure factors with the original. model_copy.reset_cache() @@ -877,8 +820,9 @@ def state_dict(self, destination=None, prefix="", keep_vars=False): Return a dictionary containing the complete state of the ModelFT. Extends parent Model.state_dict() with FT-specific parameters: - ``max_res``, ``wavelength``, and ``anomalous_threshold``. Grid state - is handled by the FFT submodule. + ``max_res``, ``explicit_gridsize``, ``wavelength`` and + ``anomalous_threshold``. The grid is derived from these and the crystal, + so it is not stored. Parameters ---------- @@ -894,12 +838,13 @@ def state_dict(self, destination=None, prefix="", keep_vars=False): dict Complete state dictionary. """ - # Parent covers _A/_B and the FFT submodule's buffers. + # Parent covers _A/_B; the engine's grid buffers are non-persistent. state = super().state_dict( destination=destination, prefix=prefix, keep_vars=keep_vars ) state[prefix + "max_res"] = self.max_res + state[prefix + "explicit_gridsize"] = self.explicit_gridsize state[prefix + "wavelength"] = self.wavelength state[prefix + "anomalous_threshold"] = self.anomalous_threshold @@ -955,6 +900,7 @@ def create_from_state_dict( dtype_float = get_float_dtype() max_res = state_dict.pop("max_res", 1.0) + explicit_gridsize = state_dict.pop("explicit_gridsize", None) state_dict.pop("radius_angstrom", None) # legacy key, no longer used wavelength = state_dict.pop("wavelength", 1.0) anomalous_threshold = state_dict.pop("anomalous_threshold", 0.5) @@ -968,10 +914,14 @@ def create_from_state_dict( strip_H = state_dict.pop("strip_H", True) altloc_pairs = state_dict.pop("altloc_pairs", []) - # FFT submodule buffers are prefixed "_fft."; older checkpoints are flat. - gridsize = state_dict.pop("_fft.gridsize", None) - if gridsize is None: - gridsize = state_dict.pop("gridsize", None) + # Checkpoints written while the grid was stored state carry its buffers + # ("_fft." prefixed, or flat in older ones). The size is adopted below only + # when it differs from what the crystal and max_res give. + legacy_gridsize = state_dict.pop("_fft.gridsize", None) + if legacy_gridsize is None: + legacy_gridsize = state_dict.pop("gridsize", None) + state_dict.pop("_fft.voxel_size", None) + state_dict.pop("voxel_size", None) instance = cls( dtype_float=saved_dtype, @@ -979,6 +929,7 @@ def create_from_state_dict( device=device, strip_H=strip_H, max_res=max_res, + gridsize=explicit_gridsize, wavelength=wavelength, anomalous_threshold=anomalous_threshold, ) @@ -987,7 +938,7 @@ def create_from_state_dict( instance.ctx.initialized = initialized instance.ctx.altloc_pairs = altloc_pairs - # Setter also sets symmetry; the cell setter below then builds the FFT. + # The engine reads both off the context; nothing further to build. instance.spacegroup = spacegroup_str from torchref.symmetry import Cell @@ -1015,13 +966,17 @@ def create_from_state_dict( "_B", torch.zeros_like(state_dict[b_key], device=device) ) - if gridsize is not None and cell_tensor is not None: - if isinstance(gridsize, torch.Tensor): - gs_tuple = tuple(int(x) for x in gridsize.tolist()) - else: - gs_tuple = tuple(int(x) for x in gridsize) - - instance.setup_grid(gridsize=gs_tuple) + if ( + legacy_gridsize is not None + and explicit_gridsize is None + and instance.ctx.crystal_key is not None + and instance.max_res is not None + ): + if isinstance(legacy_gridsize, torch.Tensor): + legacy_gridsize = legacy_gridsize.tolist() + legacy = tuple(int(x) for x in legacy_gridsize) + if legacy != instance.fft.compute_optimal_gridsize(instance.max_res): + instance.explicit_gridsize = legacy # Drop empty placeholders, remapping old-style A/B keys to _A/_B. filtered_state_dict = {} diff --git a/torchref/model/sf_ds.py b/torchref/model/sf_ds.py index 0f028cc7..d4962cad 100644 --- a/torchref/model/sf_ds.py +++ b/torchref/model/sf_ds.py @@ -5,7 +5,7 @@ crystallographic symmetry in reciprocal space. """ -from typing import Optional, Tuple +from typing import TYPE_CHECKING, Optional, Tuple import torch import torch.nn as nn @@ -18,12 +18,14 @@ get_scattering_vectors, reciprocal_basis_matrix, ) -from torchref.config import dtypes, get_complex_dtype, get_default_device +from torchref.config import dtypes, get_complex_dtype from torchref.symmetry import Cell, SpaceGroup -from torchref.symmetry.spacegroup import SpaceGroupLike from torchref.utils.device_mixin import DeviceMovementMixin from torchref.utils.device_resolution import require_cell_dtype, resolve_device +if TYPE_CHECKING: + from torchref.model.context import ModelContext + class SfDS(DeviceMovementMixin, nn.Module): """ @@ -37,36 +39,38 @@ class SfDS(DeviceMovementMixin, nn.Module): Parameters ---------- - cell : Cell, optional - Unit cell object containing cell parameters. - spacegroup : SpaceGroupLike, optional - Space group specification (string, int, or gemmi.SpaceGroup). - If None, defaults to P1. + ctx : ModelContext + The crystallographic context to read the cell and space group from. dtype_float : torch.dtype, optional Data type for floating point tensors. Default is dtypes.float. device : torch.device, optional - Computation device. Defaults to the configured default device - (``get_default_device()``). + Computation device. Defaults to the cell's device when the context has one, + else the configured default device. verbose : int, optional Verbosity level for logging. Default is 0. max_memory_gb : float, optional Maximum memory to use for intermediate tensors in GB. Default is 2.0. Set to None to disable batching. + force_portable : bool, optional + Pin the portable reference path instead of the fastest usable backend, + per instance. ``None`` (default) defers to the process-wide setting, + so ``with use_portable():`` steers an unconfigured instance. Attributes ---------- cell, spacegroup : Cell, SpaceGroup - The unit cell and the space group as an nn.Module carrying its symmetry - matrices and translations; setting ``cell`` drops the cached - reciprocal basis. + Read through from the context. The reciprocal basis is memoised against + the cell's value key, so replacing the context's cell needs no further call. Examples -------- Standalone usage:: - from torchref.symmetry import Cell - cell = Cell([50, 60, 70, 90, 90, 90]) - sf_ds = SfDS(cell, spacegroup='P212121') + from torchref.model.context import ModelContext + from torchref.symmetry import Cell, SpaceGroup + ctx = ModelContext(cell=Cell([50, 60, 70, 90, 90, 90]), + spacegroup=SpaceGroup('P212121')) + sf_ds = SfDS(ctx) sf, _ = sf_ds.compute_structure_factors( hkl, xyz_iso, adp_iso, occ_iso, A_iso, B_iso ) @@ -74,142 +78,81 @@ class SfDS(DeviceMovementMixin, nn.Module): def __init__( self, - cell: Optional[Cell] = None, - spacegroup: SpaceGroupLike = None, + ctx: "ModelContext", + *, dtype_float: torch.dtype = None, device: torch.device = None, verbose: int = 0, max_memory_gb: float = 2.0, force_portable: Optional[bool] = None, ): - """ - Initialize the SfDS module with cell and spacegroup. - - Parameters - ---------- - cell : Cell, optional - Unit cell object. If None, must be set later. - spacegroup : SpaceGroupLike, optional - Space group specification. If None, defaults to P1. - dtype_float : torch.dtype, optional - Data type for floating point tensors. Default is dtypes.float. - device : torch.device, optional - Computation device. Defaults to the configured device.current. - verbose : int, optional - Verbosity level for logging. Default is 0. - max_memory_gb : float, optional - Maximum memory for intermediate tensors in GB. Default is 2.0. - force_portable : bool, optional - Pin the portable reference path instead of the fastest usable backend, - per instance. ``None`` (default) defers to the process-wide setting, - so ``with use_portable():`` steers an unconfigured instance. - """ super().__init__() - if dtype_float is None: - dtype_float = dtypes.float - self.dtype_float = dtype_float - # Derive from ``cell`` when no device is given, instead of jumping to + self.ctx = ctx + self.dtype_float = dtypes.float if dtype_float is None else dtype_float + # Derive from the cell when no device is given, instead of jumping to # the global default and leaving a caller-supplied cell behind on - # another device. An explicit ``device`` still wins and moves the cell. - self.device = resolve_device(cell, device=device) + # another device. An explicit ``device`` wins and moves the crystal as + # a whole. + if device is not None: + ctx.to(device) + self.device = resolve_device(ctx.cell, device=device) self.verbose = verbose self.max_memory_gb = max_memory_gb self.force_portable = force_portable - # Store cell and spacegroup - self._cell = cell - self._spacegroup = None - - if spacegroup is not None or cell is not None: - self._spacegroup = SpaceGroup( - spacegroup, dtype=dtype_float, device=self.device - ) - - # Cache reciprocal basis matrix + # Reciprocal basis, memoised against the cell it was computed from. self._recB: Optional[torch.Tensor] = None + self._recB_key = None # ========================================================================= - # Cell and SpaceGroup properties + # Crystal, read through from the context # ========================================================================= @property def cell(self) -> Optional[Cell]: - """Unit cell object.""" - return self._cell - - @cell.setter - def cell(self, value: Cell): - """Set unit cell and invalidate cached reciprocal basis matrix.""" - self._cell = value - self._recB = None # Invalidate cache + """The context's unit cell.""" + return self.ctx.cell @property def spacegroup(self) -> Optional[SpaceGroup]: - """Space group object (SpaceGroup nn.Module).""" - return self._spacegroup - - @spacegroup.setter - def spacegroup(self, value: SpaceGroupLike): - """Set space group.""" - if value is not None: - self._spacegroup = SpaceGroup( - value, dtype=self.dtype_float, device=self.device - ) - else: - self._spacegroup = None + """The context's space group.""" + return self.ctx.spacegroup @property def fractional_matrix(self) -> Optional[torch.Tensor]: - """Get fractionalization matrix from cell.""" - if self._cell is not None: - return self._cell.fractional_matrix - return None + """Fractionalization matrix from the cell.""" + cell = self.ctx.cell + return None if cell is None else cell.fractional_matrix @property def inv_fractional_matrix(self) -> Optional[torch.Tensor]: - """Get orthogonalization matrix from cell.""" - if self._cell is not None: - return self._cell.inv_fractional_matrix - return None - - def set_cell_and_spacegroup(self, cell: Cell, spacegroup: SpaceGroupLike = None): - """ - Set cell and spacegroup for this SfDS instance. - - Parameters - ---------- - cell : Cell - Unit cell object. - spacegroup : SpaceGroupLike, optional - Space group specification. - - Notes - ----- - Receiver wins: an incoming cell on another device is moved to match - this module rather than the other way round. - """ - self.device = resolve_device(self, cell) - self._cell = cell - self._recB = None # Invalidate cache - self.spacegroup = spacegroup + """Orthogonalization matrix from the cell.""" + cell = self.ctx.cell + return None if cell is None else cell.inv_fractional_matrix # ========================================================================= # Internal helper methods # ========================================================================= - def _get_reciprocal_basis_matrix(self) -> torch.Tensor: - """Cached ``(3, 3)`` reciprocal basis (a*, b*, c* as rows). - - Raises ``RuntimeError`` if no cell is set, and refuses a cell whose dtype - differs from ``self.dtype_float``. - """ - if self._cell is None: - raise RuntimeError("Cell not set. Call set_cell_and_spacegroup() first.") - require_cell_dtype(self._cell, self.dtype_float, type(self).__name__) - - if self._recB is None: - self._recB = reciprocal_basis_matrix(self._cell.data) + def _require_cell(self) -> Cell: + """The context's cell, refusing a missing one or a dtype mismatch.""" + cell = self.ctx.cell + if cell is None: + raise RuntimeError( + f"{type(self).__name__}: the context {self.ctx!r} has no cell. Load a " + "structure, or pass a context whose cell is set." + ) + # Refused, not reconciled: a dtype cast is lossy, so the cell is the + # caller's to fix. See ``require_cell_dtype``. + require_cell_dtype(cell, self.dtype_float, type(self).__name__) + return cell + def _get_reciprocal_basis_matrix(self) -> torch.Tensor: + """``(3, 3)`` reciprocal basis (a*, b*, c* as rows), memoised per cell value.""" + cell = self._require_cell() + if self._recB is None or self._recB_key != cell.key: + self._recB = reciprocal_basis_matrix(cell.data) + self._recB_key = cell.key return self._recB def _compute_scattering_factors( @@ -236,12 +179,10 @@ def _cartesian_to_fractional(self, xyz_cartesian: torch.Tensor) -> torch.Tensor: """``(N, 3)`` Cartesian coordinates to fractional; needs a cell whose dtype matches ``self.dtype_float``. """ - if self._cell is None: - raise RuntimeError("Cell not set. Call set_cell_and_spacegroup() first.") - require_cell_dtype(self._cell, self.dtype_float, type(self).__name__) + cell = self._require_cell() # fractional = cartesian @ inv_frac_matrix.T - return torch.matmul(xyz_cartesian, self.inv_fractional_matrix.T) + return torch.matmul(xyz_cartesian, cell.inv_fractional_matrix.T) # ========================================================================= # Structure Factor Computation @@ -293,11 +234,7 @@ def compute_structure_factors( None Second return value is None (for API compatibility with SfFFT). """ - if self._cell is None: - raise RuntimeError("Cell not set. Call set_cell_and_spacegroup() first.") - # Refused, not reconciled: unlike the device normalization just below, a dtype cast - # is lossy, so the cell is the caller's to fix. See ``require_cell_dtype``. - require_cell_dtype(self._cell, self.dtype_float, type(self).__name__) + self._require_cell() # Normalize the input hkl onto this module's device. The symmetry # helpers derive equiv_hkls/phases from hkl.device while sf_total is @@ -319,7 +256,7 @@ def compute_structure_factors( ) # No symmetry: compute F_P1 directly - if not apply_symmetry or self._spacegroup is None: + if not apply_symmetry or self.ctx.spacegroup is None: sf_p1 = self._compute_p1_sf( hkl, xyz_frac_iso, @@ -336,9 +273,9 @@ def compute_structure_factors( return sf_p1, None # Apply late symmetry: F_sym(h) = Σ_ops exp(2πi h.t) * F_P1(R^T @ h) - n_ops = self._spacegroup.n_ops - equiv_hkls = self._spacegroup.expand_reciprocal(hkl) # (n_ops, N, 3) - phases = self._spacegroup.phase_factors(hkl) # (n_ops, N) + n_ops = self.ctx.spacegroup.n_ops + equiv_hkls = self.ctx.spacegroup.expand_reciprocal(hkl) # (n_ops, N, 3) + phases = self.ctx.spacegroup.phase_factors(hkl) # (n_ops, N) # Compute F_P1 at each equivalent HKL and combine sf_total = torch.zeros( @@ -390,7 +327,7 @@ def _compute_p1_sf( """ # Get reciprocal basis matrix and compute scattering vectors recB = self._get_reciprocal_basis_matrix() - s_vectors = get_scattering_vectors(hkl, self._cell.data, recB) + s_vectors = get_scattering_vectors(hkl, self.ctx.cell.data, recB) s = torch.norm(s_vectors, dim=1) sf_total = torch.zeros( @@ -434,34 +371,6 @@ def _compute_p1_sf( # ========================================================================= def reset_cache(self) -> None: - """Drop the cached reciprocal-basis matrix; recomputed on next use.""" + """Drop the memoised reciprocal basis; recomputed on next use.""" self._recB = None - - def copy(self) -> "SfDS": - """Create a deep copy of this SfDS module. - - Returns - ------- - SfDS - A new SfDS instance with cloned cell and spacegroup. - """ - # Clone the cell - new_cell = self._cell.clone() if self._cell is not None else None - - # Copy the spacegroup - new_spacegroup = ( - self._spacegroup.copy() if self._spacegroup is not None else None - ) - - # Create new SfDS with copied components - new_ds = SfDS( - cell=new_cell, - spacegroup=new_spacegroup, - dtype_float=self.dtype_float, - device=self.device, - verbose=self.verbose, - max_memory_gb=self.max_memory_gb, - force_portable=self.force_portable, - ) - - return new_ds + self._recB_key = None diff --git a/torchref/model/sf_fft.py b/torchref/model/sf_fft.py index c0901e3a..81a76863 100644 --- a/torchref/model/sf_fft.py +++ b/torchref/model/sf_fft.py @@ -1,218 +1,307 @@ """SfFFT -- structure factors via FFT. -An nn.Module owning the real-space grid setup, the electron-density build from -atomic parameters, and the FFT to structure factors. Usable standalone or as -``ModelFT``'s submodule. ``FFT`` is a deprecated alias for :class:`SfFFT`. +An nn.Module that reads the crystal off a :class:`~torchref.model.context.ModelContext`, +sizes its real-space grid lazily from that crystal and the resolution, builds the +electron density from atomic parameters, and transforms it to structure factors. +Usable standalone or as ``ModelFT``'s submodule. """ -from typing import Optional, Tuple +from typing import TYPE_CHECKING, Optional, Tuple import torch import torch.nn as nn from torchref.base.fourier import ifft from torchref.base.reciprocal import extract_structure_factor_from_grid -from torchref.config import dtypes, get_default_device +from torchref.config import dtypes from torchref.symmetry import Cell, SpaceGroup -from torchref.symmetry.spacegroup import SpaceGroupLike from torchref.utils.device_mixin import DeviceMovementMixin -from torchref.utils.device_resolution import resolve_device +from torchref.utils.device_resolution import require_cell_dtype, resolve_device + +if TYPE_CHECKING: + from torchref.model.context import ModelContext class SfFFT(DeviceMovementMixin, nn.Module): """ Structure Factor calculator using FFT (Fast Fourier Transform). - Built from a Cell and optionally a SpaceGroup, which drive the grid. Call - :meth:`setup_grid` before :meth:`build_density_map`; the higher-level - :meth:`compute_structure_factors` does both. + The engine does not own a cell or a space group: it holds the + :class:`~torchref.model.context.ModelContext` it was given and reads + ``ctx.cell`` and ``ctx.spacegroup`` whenever it needs them. The grid is derived + from that crystal, ``max_res`` and ``explicit_gridsize`` on first use and kept + until any of them changes (:attr:`grid_key`), so assigning a new cell or + resolution to the context needs no further call. Parameters ---------- - cell : Cell - Unit cell object containing cell parameters. - spacegroup : SpaceGroupLike, optional - Space group specification (string, int, or gemmi.SpaceGroup). - If None, defaults to P1. + ctx : ModelContext + The crystallographic context to read the cell and space group from. May be + incomplete at construction; the grid is sized once both are set. max_res : float, optional - Maximum resolution for grid spacing in Angstroms. Default is 1.5. + Maximum resolution for grid spacing in Angstroms. Default is 1.0. + explicit_gridsize : tuple of int, optional + Fixed grid dimensions ``(nx, ny, nz)``; overrides the resolution-derived size. dtype_float : torch.dtype, optional Data type for floating point tensors. Default is dtypes.float. device : torch.device, optional - Computation device. Defaults to the configured default device - (``get_default_device()``). + Computation device. Defaults to the cell's device when the context has one, + else the configured default device. verbose : int, optional Verbosity level for logging. Default is 0. + use_late_symmetry : bool, optional + Apply symmetry in reciprocal space after the FFT when the grid permits + exact indexing (default); otherwise symmetrise the density map first. Attributes ---------- cell, spacegroup : Cell, SpaceGroup - The unit cell and the space group as an nn.Module carrying its symmetry - matrices and translations; ``symmetry`` is an alias for ``spacegroup``. + Read through from the context. gridsize, voxel_size : torch.Tensor or None - Grid dimensions ``(nx, ny, nz)`` and the voxel edge vector sum -- both - ``None`` until :meth:`setup_grid` runs. No coordinate grid is stored: the - splats derive a voxel's Cartesian position from its index, so materialising - one would cost ``12 * nx * ny * nz`` bytes that nothing reads. Call - :func:`torchref.base.fourier.get_real_grid` if you genuinely need one. + Grid dimensions ``(nx, ny, nz)`` and the voxel edge vector sum, resolved on + access; ``None`` while the context has no cell or space group. No coordinate + grid is stored: the splats derive a voxel's Cartesian position from its + index, so materialising one would cost ``12 * nx * ny * nz`` bytes that + nothing reads. Call :func:`torchref.base.fourier.get_real_grid` if you + genuinely need one. """ def __init__( self, - cell: Optional[Cell] = None, - spacegroup: SpaceGroupLike = None, - max_res: float = 1.5, + ctx: "ModelContext", + *, + max_res: Optional[float] = 1.0, + explicit_gridsize: Optional[Tuple[int, int, int]] = None, dtype_float: torch.dtype = None, device: Optional[torch.device] = None, verbose: int = 0, use_late_symmetry: bool = True, ): - """ - Initialize the SfFFT module with cell and spacegroup. - - Parameters - ---------- - cell : Cell, optional - Unit cell object. If None, must be set later via set_cell(). - spacegroup : SpaceGroupLike, optional - Space group specification. If None, defaults to P1. - max_res : float, optional - Maximum resolution for grid spacing in Angstroms. Default is 1.5. - dtype_float : torch.dtype, optional - Data type for floating point tensors. Default is dtypes.float. - device : torch.device, optional - Computation device. Default is None (uses cell's device). If Cell is also None, defaults to CPU. - verbose : int, optional - Verbosity level for logging. Default is 0. - use_late_symmetry : bool, optional - If True (default), apply symmetry in reciprocal space after FFT - ("late symmetry") for faster structure factor calculation. - If False, apply symmetry to density map before FFT ("early symmetry"). - """ super().__init__() + self.ctx = ctx self.max_res = max_res - if dtype_float is None: - dtype_float = dtypes.float - self.dtype_float = dtype_float + self.explicit_gridsize = explicit_gridsize + self.dtype_float = dtypes.float if dtype_float is None else dtype_float - # One device for the module and everything it builds. ``resolve_device`` - # also moves ``cell`` when an explicit ``device`` disagrees with it, so - # the cell and the SpaceGroup below cannot end up split. - self.device = resolve_device(cell, device=device) + # One device for the module and everything it builds. An explicit + # ``device`` moves the crystal as a whole, so cell, space group and the + # grid buffers cannot end up split; without one the cell's device wins, + # and an empty context falls back to the configured default. + if device is not None: + ctx.to(device) + self.device = resolve_device(ctx.cell, device=device) self.verbose = verbose self.use_late_symmetry = use_late_symmetry - # Store cell and spacegroup - self._cell = cell - self._spacegroup = None - - if spacegroup is not None or cell is not None: - # ``self.device``, not the raw ``device`` argument: the latter is - # ``None`` on the derive-from-cell path, which would silently put - # the symmetry matrices on the global default instead. - self._spacegroup = SpaceGroup( - spacegroup, dtype=dtype_float, device=self.device - ) - - # Buffers (registered during setup_grid) - self.register_buffer("gridsize", None) - self.register_buffer("voxel_size", None) - - # Late symmetry compatibility flag (set during setup_grid) + # Derived from the crystal on first use. Non-persistent: rebuilt from the + # context on restore rather than read back from a checkpoint. + self.register_buffer("_gridsize", None, persistent=False) + self.register_buffer("_voxel_size", None, persistent=False) self._late_symmetry_compatible: Optional[bool] = None + self._grid_key = None # ========================================================================= - # Cell and SpaceGroup properties + # Crystal, read through from the context # ========================================================================= @property def cell(self) -> Optional[Cell]: - """Unit cell object.""" - return self._cell - - @cell.setter - def cell(self, value: Cell): - """Set unit cell.""" - self._cell = value + """The context's unit cell.""" + return self.ctx.cell @property def spacegroup(self) -> Optional[SpaceGroup]: - """Space group object (SpaceGroup nn.Module).""" - return self._spacegroup - - @spacegroup.setter - def spacegroup(self, value: SpaceGroupLike): - """Set space group.""" - if value is not None: - self._spacegroup = SpaceGroup( - value, dtype=self.dtype_float, device=self.device - ) - else: - self._spacegroup = None - - @property - def symmetry(self) -> Optional[SpaceGroup]: - """Symmetry operations handler (alias for spacegroup).""" - return self._spacegroup + """The context's space group.""" + return self.ctx.spacegroup @property def fractional_matrix(self) -> Optional[torch.Tensor]: - """Get fractionalization matrix from cell, on this module's device/dtype.""" - if self._cell is not None: - # Move device first, then cast: a combined ``.to(device=cpu, - # dtype=float64)`` from an MPS-resident cell raises because MPS - # rejects the transient float64 view (MPS has no float64). - return self._cell.fractional_matrix.to(device=self.device).to( - dtype=self.dtype_float - ) - return None + """Fractionalization matrix from the cell, on this module's device/dtype.""" + cell = self.ctx.cell + if cell is None: + return None + # Move device first, then cast: a combined ``.to(device=cpu, + # dtype=float64)`` from an MPS-resident cell raises because MPS + # rejects the transient float64 view (MPS has no float64). + return cell.fractional_matrix.to(device=self.device).to(dtype=self.dtype_float) @property def inv_fractional_matrix(self) -> Optional[torch.Tensor]: - """Get orthogonalization matrix from cell, on this module's device/dtype.""" - if self._cell is not None: - # Move device first, then cast (see ``fractional_matrix``). - return self._cell.inv_fractional_matrix.to(device=self.device).to( - dtype=self.dtype_float - ) - return None + """Orthogonalization matrix from the cell, on this module's device/dtype.""" + cell = self.ctx.cell + if cell is None: + return None + # Move device first, then cast (see ``fractional_matrix``). + return cell.inv_fractional_matrix.to(device=self.device).to( + dtype=self.dtype_float + ) + + # ========================================================================= + # Grid + # ========================================================================= - def set_cell_and_spacegroup(self, cell: Cell, spacegroup: SpaceGroupLike = None): + @property + def explicit_gridsize(self) -> Optional[Tuple[int, int, int]]: + """Fixed grid dimensions, or None to size the grid from ``max_res``.""" + return self._explicit_gridsize + + @explicit_gridsize.setter + def explicit_gridsize(self, value) -> None: + self._explicit_gridsize = ( + None if value is None else tuple(int(x) for x in value) + ) + + @property + def grid_key(self): + """Everything the grid is derived from, as a hashable tuple. + + ``(cell.key, spacegroup.key, max_res, explicit_gridsize)``, or ``None`` + while the context has no cell or space group. Callers that cache anything + grid-shaped can compare against it. """ - Set cell and spacegroup for this SfFFT instance. + crystal = self.ctx.crystal_key + if crystal is None: + return None + return ( + *crystal, + None if self.max_res is None else float(self.max_res), + self.explicit_gridsize, + ) - Parameters - ---------- - cell : Cell - Unit cell object. - spacegroup : SpaceGroupLike, optional - Space group specification. - - Notes - ----- - Receiver wins: this module may already own grid buffers, so an incoming - cell on another device is moved to match rather than dragging the - module after it. + def ensure_grid(self) -> bool: + """Bring the grid in line with :attr:`grid_key`. + + Returns + ------- + bool + True when the grid was (re)built, False when it was already current or + the context has no crystal yet. + + Raises + ------ + RuntimeError + If neither ``max_res`` nor ``explicit_gridsize`` is set, or the cell or + space group disagree with this module's dtype or device. """ - self.device = resolve_device(self, cell) - self._cell = cell - self.spacegroup = spacegroup + key = self.grid_key + if key == self._grid_key: + return False + if key is None: + # The crystal went away; drop the grid derived from the old one. + self._gridsize = None + self._voxel_size = None + self._late_symmetry_compatible = None + self._grid_key = None + return False + + cell, spacegroup = self.ctx.cell, self.ctx.spacegroup + require_cell_dtype(cell, self.dtype_float, type(self).__name__) + if spacegroup.matrices.dtype != self.dtype_float: + raise RuntimeError( + f"{type(self).__name__} was built for {self.dtype_float} but its " + f"space group holds {spacegroup.matrices.dtype}. Rebuild the space " + "group at the module's dtype." + ) + if spacegroup.matrices.device != cell.data.device: + raise RuntimeError( + f"{type(self).__name__}: cell on {cell.data.device} but space group " + f"on {spacegroup.matrices.device}. Move the context as a whole." + ) - # ========================================================================= - # Grid Setup Methods - # ========================================================================= + if self.explicit_gridsize is not None: + gridsize = self.explicit_gridsize + elif self.max_res is not None: + gridsize = self.compute_optimal_gridsize(self.max_res) + else: + raise RuntimeError( + f"{type(self).__name__} cannot size its grid: set max_res or " + "explicit_gridsize." + ) + shape = tuple(int(n) for n in gridsize) + + if self.verbose > 1: + print(f"Setting up grids with max_res={self.max_res} Å") + + previous = self._gridsize + self._gridsize = torch.tensor(shape, dtype=dtypes.int, device=self.device) + + # The step between diagonally adjacent grid points, i.e. the sum of the three + # cell edge vectors each divided by its own sampling count. Equal to the true + # per-axis voxel edge lengths only for an orthogonal cell; kept because that is + # what the previous grid-differencing definition produced. + self._voxel_size = self.fractional_matrix @ ( + 1.0 / self._gridsize.to(self.dtype_float) + ) + + # Every symmetry-equivalent HKL lands on an integer grid point exactly when + # the grid admits direct indexing, which the space group answers without + # building an operator. + self._late_symmetry_compatible = spacegroup.can_index_directly(shape) + if self.verbose > 0 and self.use_late_symmetry: + if self._late_symmetry_compatible: + print("SfFFT: Using late symmetry (reciprocal space)") + else: + print( + "SfFFT: Late symmetry disabled - grid not compatible " + "(falling back to early symmetry)" + ) + + # The space group memoises its map operator and reciprocal extractor per + # grid shape. They are keyed on the shape, so a same-shape rebuild keeps + # them; a different shape drops them so the old operator's sampling grids + # do not stay resident. + if previous is not None and tuple(int(n) for n in previous.tolist()) != shape: + spacegroup.reset_cache() + + if self.verbose > 2: + print(f"Grid shape: {shape}") + print(f"Voxel size: {self._voxel_size}") + + # Last, so a failure above leaves the old key in place and the next call + # tries again. + self._grid_key = key + return True + + def _require_grid(self) -> None: + """Resolve the grid, refusing to proceed without a crystal.""" + self.ensure_grid() + if self._gridsize is None: + raise RuntimeError( + f"{type(self).__name__} has no crystal to size its grid: the context " + f"{self.ctx!r} has no cell or space group. Load a structure, or pass a " + "context whose cell and spacegroup are set." + ) + + @property + def gridsize(self) -> Optional[torch.Tensor]: + """Grid dimensions ``(nx, ny, nz)``, or None without a crystal.""" + self.ensure_grid() + return self._gridsize + + @property + def voxel_size(self) -> Optional[torch.Tensor]: + """Voxel edge vector sum, or None without a crystal.""" + self.ensure_grid() + return self._voxel_size @property def grid_shape(self) -> Optional[Tuple[int, int, int]]: - """Map dimensions ``(nx, ny, nz)``, or ``None`` before :meth:`setup_grid`.""" - if self.gridsize is None: + """Map dimensions ``(nx, ny, nz)`` as Python ints, or None without a crystal.""" + gridsize = self.gridsize + if gridsize is None: return None - return tuple(int(n) for n in self.gridsize) + return tuple(int(n) for n in gridsize) + + @property + def late_symmetry_compatible(self) -> Optional[bool]: + """Whether the current grid admits reciprocal-space symmetrisation.""" + self.ensure_grid() + return self._late_symmetry_compatible def compute_optimal_gridsize(self, max_res: Optional[float] = None) -> tuple: """ - Compute optimal grid dimensions using the stored cell and spacegroup. + Compute optimal grid dimensions from the context's cell and space group. Uses Cell.compute_grid_size() for base calculation and Symmetry.suggest_grid_size() for symmetry optimization. @@ -230,21 +319,25 @@ def compute_optimal_gridsize(self, max_res: Optional[float] = None) -> tuple: Raises ------ RuntimeError - If cell has not been set. + If the context has no cell or space group. """ - if self._cell is None: - raise RuntimeError("Cell not set. Call set_cell_and_spacegroup() first.") + cell, spacegroup = self.ctx.cell, self.ctx.spacegroup + if cell is None or spacegroup is None: + raise RuntimeError( + f"{type(self).__name__}: the context {self.ctx!r} has no cell or " + "space group to size a grid from." + ) resolution = max_res if max_res is not None else self.max_res # Use Cell's method for base grid size calculation - gridsize_initial = self._cell.compute_grid_size(resolution) + gridsize_initial = cell.compute_grid_size(resolution) if self.verbose > 1: print(f"Initial grid size from cell: {gridsize_initial}") # Optimize for symmetry and FFT-friendliness - gridsize_optimized = self._spacegroup.suggest_grid_size( + gridsize_optimized = spacegroup.suggest_grid_size( gridsize_initial, make_fft_friendly=True ) if self.verbose > 1 and gridsize_optimized != gridsize_initial: @@ -252,88 +345,30 @@ def compute_optimal_gridsize(self, max_res: Optional[float] = None) -> tuple: f"Optimized grid size from {gridsize_initial} to {gridsize_optimized} " f"(symmetry + FFT friendly)" ) - return gridsize_optimized + return tuple(int(n) for n in gridsize_optimized) def setup_grid( self, - gridsize: Optional[Tuple[int, int, int]] = None, + *, max_res: Optional[float] = None, - ): + gridsize: Optional[Tuple[int, int, int]] = None, + ) -> None: """ - Setup the real-space grid for electron density calculation. - - This method initializes and stores the grid state for subsequent - density map calculations. Uses the stored cell and spacegroup. + Override the grid's inputs explicitly and resolve it now. Parameters ---------- - gridsize : tuple of int, optional - Explicit grid size (nx, ny, nz). If None, computed automatically - using Cell.compute_grid_size() and Symmetry.suggest_grid_size(). max_res : float, optional - Maximum resolution in Angstroms. If None, uses self.max_res. - - Raises - ------ - RuntimeError - If cell has not been set. + New maximum resolution in Angstroms. None leaves the current value. + gridsize : tuple of int, optional + Fixed grid size (nx, ny, nz), kept until cleared through + :attr:`explicit_gridsize`. None leaves the current value. """ - if self._cell is None: - raise RuntimeError("Cell not set. Call set_cell_and_spacegroup() first.") - if max_res is not None: - self.max_res = max_res - - if self.verbose > 1: - print(f"Setting up grids with max_res={self.max_res} Å") - - # Compute or use provided grid size + self.max_res = float(max_res) if gridsize is not None: - self.gridsize = torch.tensor(gridsize, dtype=dtypes.int, device=self.device) - else: - optimal_gridsize = self.compute_optimal_gridsize(self.max_res) - self.gridsize = torch.tensor( - optimal_gridsize, dtype=dtypes.int, device=self.device - ) - - # The step between diagonally adjacent grid points, i.e. the sum of the three - # cell edge vectors each divided by its own sampling count. Equal to the true - # per-axis voxel edge lengths only for an orthogonal cell; kept because that is - # what the previous grid-differencing definition produced. - self.voxel_size = ( - self._cell.fractional_matrix.to(self.device) - @ (1.0 / self.gridsize.to(self._cell.fractional_matrix.dtype)) - ) - - # Every symmetry-equivalent HKL lands on an integer grid point exactly when - # the grid admits direct indexing, which the space group answers without - # building an operator. - if self._spacegroup is not None: - self._late_symmetry_compatible = self._spacegroup.can_index_directly( - self.grid_shape - ) - - if self.use_late_symmetry and self._late_symmetry_compatible: - if self.verbose > 0: - print( - "SfFFT: Using late symmetry (reciprocal space)" - ) - elif self.use_late_symmetry and not self._late_symmetry_compatible: - if self.verbose > 0: - print( - "SfFFT: Late symmetry disabled - grid not compatible " - "(falling back to early symmetry)" - ) - else: - self._late_symmetry_compatible = False - - # The grid shape changed, so the space group's cached operators are stale. - if self._spacegroup is not None: - self._spacegroup.reset_cache() - - if self.verbose > 2: - print(f"Grid shape: {self.grid_shape}") - print(f"Voxel size: {self.voxel_size}") + self.explicit_gridsize = gridsize + self.ensure_grid() # ========================================================================= # Density Map Building Methods @@ -356,8 +391,6 @@ def build_density_map( """ Build electron density map from atomic parameters. - Calls :meth:`setup_grid` itself if no grid has been set up yet. - Parameters ---------- xyz_iso, adp_iso, occ_iso : torch.Tensor @@ -378,8 +411,7 @@ def build_density_map( torch.Tensor Electron density map with shape (nx, ny, nz). """ - if self.gridsize is None: - self.setup_grid() + self._require_grid() from torchref.base.electron_density.main import build_electron_density @@ -401,9 +433,8 @@ def build_density_map( dtype=self.dtype_float, ) - # Apply symmetry if requested - if apply_symmetry and self._spacegroup is not None: - density_map = self._spacegroup.symmetrize_map(density_map) + if apply_symmetry: + density_map = self.ctx.spacegroup.symmetrize_map(density_map) return density_map @@ -428,25 +459,23 @@ def map_to_structure_factors( hkl : torch.Tensor Miller indices with shape (n_reflections, 3). apply_symmetry : bool, optional - If True (default) and late symmetry is enabled/compatible, apply - symmetry in reciprocal space. If False, the density map is assumed - to already have symmetry applied (early symmetry path). + If True (default), apply symmetry in reciprocal space. If False, the + density map is assumed to already have symmetry applied (early + symmetry path). Returns ------- torch.Tensor Complex structure factors with shape (n_reflections,). """ - reciprocal_space_grid = ifft(density_map, self.cell.volume) + self._require_grid() + reciprocal_space_grid = ifft(density_map, self.ctx.cell.volume) - # Use late symmetry if enabled, compatible, and requested if apply_symmetry: - # Lazily build / reuse cached extractor (precomputed flat indices) - grid_shape = tuple(int(x) for x in self.gridsize) - extractor = self._spacegroup.reciprocal_extractor(hkl, grid_shape) + # Memoised on the space group per (hkl, grid shape). + extractor = self.ctx.spacegroup.reciprocal_extractor(hkl, self.grid_shape) return extractor.extract_from_grid(reciprocal_space_grid) - else: - return extract_structure_factor_from_grid(reciprocal_space_grid, hkl) + return extract_structure_factor_from_grid(reciprocal_space_grid, hkl) def compute_structure_factors( self, @@ -496,8 +525,9 @@ def compute_structure_factors( Electron density map with shape (nx, ny, nz). Note: When using late symmetry, this is the P1 map (without symmetry). """ - # Late symmetry: build a P1 map, symmetrize in reciprocal space. - # Early symmetry: symmetrize the density map before the FFT. + # Resolve the grid first: the late-symmetry flag belongs to the grid the + # density is about to be built on. + self._require_grid() use_late = ( apply_symmetry and self.use_late_symmetry and self._late_symmetry_compatible ) @@ -521,55 +551,3 @@ def compute_structure_factors( apply_symmetry=use_late, # Late symmetry ) return sf, density_map - - # ========================================================================= - # Device Movement - # ========================================================================= - - def reset_cache(self) -> None: - """Drop the space group's cached operators; recomputed on next use.""" - if self._spacegroup is not None: - self._spacegroup.reset_cache() - - def copy(self) -> "SfFFT": - """Create a deep copy of this SfFFT module. - - Returns - ------- - SfFFT - A new SfFFT instance with cloned cell, spacegroup, and buffers. - """ - # Clone the cell - new_cell = self._cell.clone() if self._cell is not None else None - - # Copy the spacegroup - new_spacegroup = ( - self._spacegroup.copy() if self._spacegroup is not None else None - ) - - # Create new SfFFT with copied components - new_fft = SfFFT( - cell=new_cell, - spacegroup=new_spacegroup, - max_res=self.max_res, - dtype_float=self.dtype_float, - device=self.device, - verbose=self.verbose, - use_late_symmetry=self.use_late_symmetry, - ) - - return new_fft - - -# Backward compatibility alias — deprecated, use SfFFT directly -def FFT(*args, **kwargs): - """Deprecated: use SfFFT instead.""" - import warnings - - warnings.warn( - "FFT is deprecated, use SfFFT instead. " - "FFT will be removed in a future release.", - DeprecationWarning, - stacklevel=2, - ) - return SfFFT(*args, **kwargs) diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index 6c758e58..88444aec 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -846,7 +846,7 @@ def _sync_model_cell_to_data( f" data cell: {d}", stacklevel=2, ) - self.model.cell = self.reflection_data.cell + self.model.cell = self.reflection_data.cell.clone() self.model.reset_cache() def parameters(self, recurse: bool = True): diff --git a/torchref/scaling/solvent.py b/torchref/scaling/solvent.py index 09b67a04..6841210f 100644 --- a/torchref/scaling/solvent.py +++ b/torchref/scaling/solvent.py @@ -220,8 +220,6 @@ def __init__( self.model = ModuleReference(model) # Store reference to model self.model.get_vdw_radii() # Ensure VdW radii are available assert self.model, "Model is not initialized" - if model.gridsize is None: - model.setup_grid() # Phenix-style parameters self.solvent_radius = radius # For dilation (accessible surface) diff --git a/torchref/symmetry/cell.py b/torchref/symmetry/cell.py index a8743cb3..63a47e94 100644 --- a/torchref/symmetry/cell.py +++ b/torchref/symmetry/cell.py @@ -20,7 +20,7 @@ from torchref.utils.device_mixin import _NonModuleDeviceMixin -@dataclass +@dataclass(eq=False) class Cell(_NonModuleDeviceMixin): """ Dataclass for crystallographic unit cells with cached derived quantities. @@ -33,6 +33,12 @@ class Cell(_NonModuleDeviceMixin): access and cached. The cache is cleared when the cell is moved to a different device or dtype. + Two cells compare and hash equal when their six parameters are equal + (:attr:`key`), independent of device and dtype, so a cell can key a dict or + a cache of quantities derived from it. A cell is therefore a value: editing + its parameter tensor in place is refused at the next derived read. Build a + new ``Cell`` and assign it instead. + Examples -------- >>> cell = Cell([50, 60, 70, 90, 90, 90]) @@ -45,6 +51,7 @@ class Cell(_NonModuleDeviceMixin): _data: torch.Tensor _cache: dict = field(default_factory=dict, repr=False) + _stamp: tuple = field(default=None, repr=False) def __init__( self, @@ -77,6 +84,10 @@ def __init__( # Convert to tensor first to get shape if isinstance(data, torch.Tensor): tensor = data.to(dtype=dtype, device=device) + if tensor is data: + # ``to`` returned the caller's tensor unchanged; own a copy so + # their later edits cannot reach into this cell. + tensor = tensor.clone() else: tensor = torch.tensor(data, dtype=dtype, device=device) @@ -92,6 +103,7 @@ def __init__( object.__setattr__(self, "_data", tensor) object.__setattr__(self, "_cache", {}) + self._stamp_data() # ========================================================================= # Device/dtype movement methods @@ -102,8 +114,35 @@ def __init__( # and any cached tensor values) and then call ``reset_cache`` below. def reset_cache(self) -> None: - """Clear cached derived quantities (fractional matrix, volume, etc.).""" + """Clear cached derived quantities (fractional matrix, volume, etc.). + + Also re-stamps the parameter tensor, so a device or dtype move (which + rebinds it and then calls this) is not mistaken for an in-place edit. + """ object.__setattr__(self, "_cache", {}) + self._stamp_data() + + def _stamp_data(self) -> None: + object.__setattr__(self, "_stamp", (id(self._data), self._data._version)) + + def _assert_unmodified(self) -> None: + """Refuse to serve derived quantities from a tensor edited in place.""" + stamp = getattr(self, "_stamp", None) + if stamp is None: + self._stamp_data() + return + if (id(self._data), self._data._version) == stamp: + return + cached = self._cache.get("key") + held = f" It held {cached}." if cached is not None else "" + raise RuntimeError( + "This Cell was edited in place after it was built." + held + " Cells are " + "values shared by reference (model context, structure-factor engine, " + "scaler), so an in-place edit changes the crystal under every holder. " + "Please don't edit Cell objects, create a new one -- " + "Cell([a, b, c, alpha, beta, gamma], dtype=cell.dtype, device=cell.device) " + "-- and assign it, e.g. model.cell = new_cell." + ) def clone(self) -> "Cell": """ @@ -118,8 +157,41 @@ def clone(self) -> "Cell": new_cell = Cell.__new__(Cell) object.__setattr__(new_cell, "_data", new_data) object.__setattr__(new_cell, "_cache", {}) + new_cell._stamp_data() return new_cell + # ========================================================================= + # Value identity + # ========================================================================= + + @property + def key(self) -> tuple: + """The six parameters as a tuple of Python floats. + + Read off the tensor once and cached alongside the derived quantities, so + the device synchronisation happens once per construction or + :meth:`reset_cache` rather than on every comparison. + + Returns + ------- + tuple of float + ``(a, b, c, alpha, beta, gamma)``. + """ + self._assert_unmodified() + key = self._cache.get("key") + if key is None: + key = tuple(float(v) for v in self._data.tolist()) + self._cache["key"] = key + return key + + def __hash__(self) -> int: + return hash(self.key) + + def __eq__(self, other) -> bool: + if not isinstance(other, Cell): + return NotImplemented + return self.key == other.key + # ========================================================================= # Basic properties # ========================================================================= @@ -192,6 +264,7 @@ def fractional_matrix(self) -> torch.Tensor: torch.Tensor Shape (3, 3) orthogonalization matrix. """ + self._assert_unmodified() if "fractional_matrix" not in self._cache: self._cache["fractional_matrix"] = self._compute_fractional_matrix() return self._cache["fractional_matrix"] @@ -208,6 +281,7 @@ def inv_fractional_matrix(self) -> torch.Tensor: torch.Tensor Shape (3, 3) fractionalization matrix. """ + self._assert_unmodified() if "inv_fractional_matrix" not in self._cache: self._cache["inv_fractional_matrix"] = torch.linalg.inv( self.fractional_matrix @@ -224,6 +298,7 @@ def volume(self) -> torch.Tensor: torch.Tensor Scalar tensor with the cell volume. """ + self._assert_unmodified() if "volume" not in self._cache: self._cache["volume"] = self._compute_volume() return self._cache["volume"] @@ -238,6 +313,7 @@ def reciprocal_basis_matrix(self) -> torch.Tensor: torch.Tensor Shape (3, 3) matrix where rows are the reciprocal basis vectors. """ + self._assert_unmodified() if "reciprocal_basis_matrix" not in self._cache: self._cache["reciprocal_basis_matrix"] = ( self._compute_reciprocal_basis_matrix() diff --git a/torchref/symmetry/spacegroup.py b/torchref/symmetry/spacegroup.py index 4ec95080..e7c7d5a4 100644 --- a/torchref/symmetry/spacegroup.py +++ b/torchref/symmetry/spacegroup.py @@ -65,8 +65,10 @@ def _normalize_spacegroup(spacegroup: SpaceGroupLike) -> gemmi.SpaceGroup: # Duck-typed rather than an isinstance check against SpaceGroup, so this stays # usable from module scope before the class below is defined. - if hasattr(spacegroup, "_sg_hm") and hasattr(spacegroup, "matrices"): - return gemmi.find_spacegroup_by_name(spacegroup._sg_hm) + if hasattr(spacegroup, "_sg_xhm") and hasattr(spacegroup, "matrices"): + # The extended symbol carries the setting; the plain H-M symbol does not + # (``R 3:R`` would come back as ``R 3:H``). + return gemmi.find_spacegroup_by_name(spacegroup._sg_xhm) if isinstance(spacegroup, int): try: @@ -207,7 +209,7 @@ def _gemmi(self) -> gemmi.SpaceGroup: Never cached: a persistent reference to the C++ singleton produces nanobind leak warnings at interpreter shutdown. """ - return gemmi.find_spacegroup_by_name(self._sg_hm) + return gemmi.find_spacegroup_by_name(self._sg_xhm) @property def name(self) -> str: @@ -224,6 +226,16 @@ def xhm(self) -> str: """Extended Hermann-Mauguin notation, including the setting token.""" return self._sg_xhm + @property + def key(self) -> str: + """Value identity of this space group: the extended H-M symbol. + + Two groups with the same key have the same operations in the same + setting. The group number alone would not do -- ``P 1 21 1`` and + ``P 1 1 21`` share number 4 but place the screw axis differently. + """ + return self._sg_xhm + @property def number(self) -> int: """Space group number, 1-230.""" @@ -477,15 +489,15 @@ def copy(self) -> "SpaceGroup": return new def __hash__(self) -> int: - """Hash on the space group number.""" - return hash(self._sg_number) + """Hash on :attr:`key` (the extended H-M symbol).""" + return hash(self._sg_xhm) def __eq__(self, other) -> bool: - """Equality on the space group number; also compares to a ``gemmi.SpaceGroup``.""" + """Equality on :attr:`key`; also compares to a ``gemmi.SpaceGroup``.""" if isinstance(other, SpaceGroup): - return self._sg_number == other._sg_number + return self._sg_xhm == other._sg_xhm if isinstance(other, gemmi.SpaceGroup): - return self._sg_number == other.number + return self._sg_xhm == other.xhm() return False def __repr__(self) -> str: diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index 2ee0a3c9..975a7119 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -748,9 +748,7 @@ def vdw_radii_cpu(): else: cell_cpu = None if self._spacegroup is not None: - from torchref.symmetry.spacegroup import SpaceGroup - sg_cpu = SpaceGroup(self._spacegroup, device=cpu, - dtype=self._spacegroup.dtype) + sg_cpu = self._spacegroup.copy().to(cpu) else: sg_cpu = None From 0fad03d94d46b1b26d638b4e5d74d7f776277d92 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 2 Sep 2026 17:07:27 +0200 Subject: [PATCH 151/250] Evaluate the P1 model directly in the translation search DirectModelEvaluator wrapped a plain P1 ModelFT to present the old interpolator interface, with the rotation and cell arguments ignored, and to land the result on the default device. With the grid derived lazily from cell, space group and max_res the wrapper had nothing left to do; prepare_candidate takes the model and does the device move itself. The probe that asked whether a P1 copy was needed at all is answered by the one-template design and is removed. 192 tests (550841); 30/30 at both windows (550842). Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- alignment_lab/analysis/p1_copy_probe.sh | 17 ----- alignment_lab/diagnostics/p1_copy_probe.py | 71 ------------------- .../diagnostics/truth_pose_scores.py | 6 +- docs/changelog.rst | 1 + .../alignment/test_patterson_translation.py | 7 +- torchref/experimental/alignment/__init__.py | 2 - torchref/experimental/alignment/pipeline.py | 5 +- .../experimental/alignment/translation.py | 44 +++--------- 8 files changed, 17 insertions(+), 136 deletions(-) delete mode 100644 alignment_lab/analysis/p1_copy_probe.sh delete mode 100644 alignment_lab/diagnostics/p1_copy_probe.py diff --git a/alignment_lab/analysis/p1_copy_probe.sh b/alignment_lab/analysis/p1_copy_probe.sh deleted file mode 100644 index d70f306a..00000000 --- a/alignment_lab/analysis/p1_copy_probe.sh +++ /dev/null @@ -1,17 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=p1probe -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -"$PY" -u alignment_lab/diagnostics/p1_copy_probe.py 1DAW 2DQ6 4BX9 2>&1 | grep -E '^ROW|rror|Error' -echo DONE diff --git a/alignment_lab/diagnostics/p1_copy_probe.py b/alignment_lab/diagnostics/p1_copy_probe.py deleted file mode 100644 index 00d028e8..00000000 --- a/alignment_lab/diagnostics/p1_copy_probe.py +++ /dev/null @@ -1,71 +0,0 @@ -"""Does the translation search need a P1 copy of the model, or just apply_symmetry=False? - -`_placement_for_candidate` copies each rotated candidate, sets its space group to -P 1, and wraps it in `DirectModelEvaluator`, because `ModelFT.forward` hardcodes -``apply_symmetry=True`` and the Crowther-Blow expansion needs the SINGLE-MOLECULE -transform ``F_p1(h R_i)`` -- it applies the symmetry itself. - -But the flag exists one level down, on ``SfFFT.compute_structure_factors``. If -calling that with ``apply_symmetry=False`` on the unmodified model agrees with -the P1 copy, then per candidate the copy, the space-group assignment and whatever -they rebuild are all avoidable, and the evaluator wrapper has nothing left to do. - -Reports agreement and the cost of each step, including ``Model.copy`` -- which is -on record as rebuilding a map-symmetry table per symop and being half the cost of -a rotation search. -""" -import sys, time -from pathlib import Path -import torch -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) -from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 - - -def _t(fn, n=3): - fn() - ts = [] - for _ in range(n): - t0 = time.perf_counter(); out = fn(); ts.append(time.perf_counter() - t0) - return min(ts), out - - -def main(): - for pdb in sys.argv[1:] or ["1DAW", "2DQ6"]: - model, data = load_case(pdb) - rot = model.copy().rotate( - random_rotation(seed_for(pdb, 0)).to(model.dtype_float), - center=model.xyz().mean(0)) - rot.spacegroup = data.spacegroup.hm - hkl = data.hkl[data.get_valid_mask()] - hkl_i = hkl.round().to(torch.int64) - - # what the pipeline does now - t_copy, p1 = _t(lambda: rot.copy()) - t_sg, _ = _t(lambda: setattr(p1, "spacegroup", "P 1")) - t_p1_sf, F_p1 = _t(lambda: p1(hkl_i)) - - # the candidate replacement: same model, symmetry off at the FFT - def direct(): - sf, _ = rot.fft.compute_structure_factors( - hkl_i, *rot.get_iso(), *rot.get_aniso(), apply_symmetry=False) - return sf - t_direct, F_direct = _t(direct) - - # and with symmetry ON, for contrast -- this is NOT what the TF wants - t_sym, F_sym = _t(lambda: rot(hkl_i)) - - a, b = F_p1.to(torch.complex128), F_direct.to(torch.complex128) - num = (a - b).abs().max() - rel = float(num / a.abs().max().clamp(min=1e-30)) - agree_sym = float((a - F_sym.to(torch.complex128)).abs().max() - / a.abs().max().clamp(min=1e-30)) - print(f"ROW pdb={pdb} N={hkl_i.shape[0]} " - f"copy={1000*t_copy:.1f}ms set_sg={1000*t_sg:.1f}ms " - f"p1_sf={1000*t_p1_sf:.1f}ms direct_sf={1000*t_direct:.1f}ms " - f"sym_sf={1000*t_sym:.1f}ms " - f"| max_rel_diff(p1, apply_symmetry=False)={rel:.3e} " - f"max_rel_diff(p1, symmetry_on)={agree_sym:.3e}", flush=True) - - -main() diff --git a/alignment_lab/diagnostics/truth_pose_scores.py b/alignment_lab/diagnostics/truth_pose_scores.py index 7e2887fd..e61833ab 100644 --- a/alignment_lab/diagnostics/truth_pose_scores.py +++ b/alignment_lab/diagnostics/truth_pose_scores.py @@ -26,14 +26,14 @@ def scores_at(pipe, model_placed): """(tf score, R, llg) of an already-placed model through the pipeline's path.""" from torchref.experimental.alignment.translation import ( - DirectModelEvaluator, analytic_r_at, llg_at_translations, - prepare_candidate, translation_score_at) + analytic_r_at, llg_at_translations, prepare_candidate, + translation_score_at) data, obs = pipe.data, pipe._obs m = model_placed.copy() if pipe.tf_d_min > 0.0: m.max_res = pipe.tf_d_min / 1.5 m.spacegroup = "P 1" - cand = prepare_candidate(DirectModelEvaluator(m), obs, data.spacegroup, data.cell) + cand = prepare_candidate(m, obs, data.spacegroup, data.cell) t0 = torch.zeros(3, dtype=torch.float64) tf = translation_score_at(obs, cand, t0) r = analytic_r_at(obs, cand, t0) diff --git a/docs/changelog.rst b/docs/changelog.rst index 913cac2d..d8f2bbbf 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Removed ``DirectModelEvaluator``; the translation search evaluates an ordinary P1 ``ModelFT`` directly. With the model's grid derived lazily from cell, space group and ``max_res`` there was nothing left for the wrapper to do - The molecular-replacement pipeline carries 10 rotation candidates by default instead of 25. With symmetry mates suppressed the rotation function's first peak is the true orientation in 50 of 50 pose-gated cells, and the panel is 30/30 at either depth. Warm on one EPYC 9335 node: 1DAW 0.42 s, 2DQ6 0.67 s, 6G9X 0.67 s, 3K7M 1.02 s, 4BX9 1.11 s per alignment - The fast translation function accumulates only the upper triangle of symmetry pairs; the lower triangle is its conjugate mirror and the diagonal a constant. Half the scatter, which is the stage's cost on high-symmetry cells: 3K7M's translation stage 1.04 s to 0.72 s. ``MRSolution.candidate_index`` records each solution's position in the rotation function's list - The placement loop re-orients one P1 copy of the search model in place per candidate and builds the placed model for the winner only, instead of copying the model three times per candidate. ``MRSolution.model`` is ``None`` for the other candidates; ``MolecularReplacementPipeline.place`` builds it on request. Warm on one EPYC 9335 node: 1DAW 0.85 s, 2DQ6 1.35 s, 6G9X 1.30 s, 3K7M 2.35 s, 4BX9 2.55 s per alignment (from 1.1, 1.9, 2.2, 2.7, 3.8) diff --git a/tests/unit/alignment/test_patterson_translation.py b/tests/unit/alignment/test_patterson_translation.py index a4c7136d..0e784648 100644 --- a/tests/unit/alignment/test_patterson_translation.py +++ b/tests/unit/alignment/test_patterson_translation.py @@ -20,7 +20,6 @@ import torch from torchref.experimental.alignment.translation import ( - DirectModelEvaluator, TranslationObs, fast_translation_function, llg_at_translations, @@ -58,8 +57,7 @@ def _search(canonical, data, mask, t_true): torch.tensor(t_true, dtype=canonical.dtype_float), fractional=True, ) obs = TranslationObs.build(F_obs, data.hkl[mask], data.spacegroup, data.cell) - cand = prepare_candidate(DirectModelEvaluator(model_p1), obs, - data.spacegroup, data.cell) + cand = prepare_candidate(model_p1, obs, data.spacegroup, data.cell) _, peaks = fast_translation_function( obs, cand, data.cell, grid_spacing_A=4.0 / 3.0, n_peaks=3, cluster_radius_A=4.0, @@ -127,8 +125,7 @@ def test_e_calc_is_normalised(setup): model_p1.max_res = 4.0 / 1.5 model_p1.spacegroup = "P 1" obs = TranslationObs.build(F_obs, data.hkl[mask], data.spacegroup, data.cell) - cand = prepare_candidate(DirectModelEvaluator(model_p1), obs, - data.spacegroup, data.cell) + cand = prepare_candidate(model_p1, obs, data.spacegroup, data.cell) # The normalisation already carries eps: E is per unit of eps*Sigma_calc. E2 = cand.e_calc(torch.zeros(3, dtype=torch.float64)) ** 2 mean_e2 = float(E2.mean()) diff --git a/torchref/experimental/alignment/__init__.py b/torchref/experimental/alignment/__init__.py index 0d13918c..797261f8 100644 --- a/torchref/experimental/alignment/__init__.py +++ b/torchref/experimental/alignment/__init__.py @@ -67,7 +67,6 @@ ) from .translation import ( CandidateTransform, - DirectModelEvaluator, TranslationObs, TranslationPeak, analytic_r_at, @@ -101,7 +100,6 @@ # Translation search "TranslationObs", "TranslationPeak", - "DirectModelEvaluator", "fast_translation_function", "CandidateTransform", "analytic_r_at", diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 7e380fcf..ce70b5f2 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -64,7 +64,6 @@ from .frf.types import RotationPeak from .rotation_search import prepare_frf_inputs, search_peaks from .translation import ( - DirectModelEvaluator, TranslationObs, analytic_r_at, fast_translation_function, @@ -313,7 +312,6 @@ def __init__( self._p1 = None self._p1_xyz0 = None self._p1_center = None - self._evaluator = None #: Levels are documented on the class. They are a contract, not a dial: #: level 2 is specifically "one machine-readable line per candidate", and @@ -646,7 +644,6 @@ def _prepare_translation_arrays(self) -> None: self._p1 = p1 self._p1_xyz0 = p1.xyz().detach().clone() self._p1_center = self._p1_xyz0.mean(dim=0) - self._evaluator = DirectModelEvaluator(p1) def _placement_for_candidate(self) -> Optional[tuple]: """Translation search for the orientation currently in the P1 template. @@ -661,7 +658,7 @@ def _placement_for_candidate(self) -> Optional[tuple]: obs = self._obs timer.start("5_candidate_transform") - cand = prepare_candidate(self._evaluator, obs, data.spacegroup, data.cell) + cand = prepare_candidate(self._p1, obs, data.spacegroup, data.cell) timer.stop("5_candidate_transform") # One FFT on a grid a third of the set's resolution apart: dense enough diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index a4bf365d..4e91faf7 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -176,33 +176,6 @@ def build( ) -class DirectModelEvaluator: - """Returns ``F_p1(hkl)`` of a P1-spacegroup model at integer HKL. - - The translation search asks its evaluator for ``F`` at a list of rotated - Miller indices. The rotation is already baked into the model's coordinates - by the time this is built, so ``R`` is ignored and every call is a direct - structure-factor evaluation rather than an interpolation. - - Answers on the **configured default device**, whatever device the model - happens to sit on. Reading the device off the model instead is how the - translation stage ends up split across two devices when a caller builds a - CPU model on a host with an accelerator: the model answers on the CPU while - everything derived from config answers on the GPU. - """ - - def __init__(self, m: "ModelFT") -> None: - self._m = m - self.device = get_default_device() - - def evaluate(self, R, hkl, real_cell, return_amplitude=False): - hkl_int = hkl.round().to(torch.int64).to(self._m.xyz().device) - with torch.no_grad(): - f = self._m(hkl_int) - f = f.to(self.device) - return f.abs() if return_amplitude else f - - @dataclass class CandidateTransform: """One oriented model's transform at the symmetry-rotated indices. @@ -235,15 +208,19 @@ def e_calc(self, t: torch.Tensor) -> torch.Tensor: def prepare_candidate( - evaluator, + model_p1: "ModelFT", obs: TranslationObs, spacegroup, real_cell, ) -> CandidateTransform: """Evaluate one orientation's transform and normalise it. - The only per-candidate model evaluation: ``F_p1`` at all ``S x N`` rotated - indices in one call. The normalising curve ``Sigma_P(s)`` is the same + ``model_p1`` is an ordinary :class:`~torchref.model.ModelFT` in P1 with the + crystal's cell, already in the candidate orientation; its grid is whatever + its ``max_res`` implies. The only per-candidate model evaluation: ``F_p1`` + at all ``S x N`` rotated indices in one call. The result is moved to the + configured default device, whatever device the model sits on, so the + translation stage never straddles two devices. The normalising curve ``Sigma_P(s)`` is the same Wilson fit the observed side uses, on the same abscissa, fitted to the transform's mean intensity over the ``S`` copies -- which is what the crystal sum averages to over a shell, since the cross terms between symmetry copies @@ -264,10 +241,9 @@ def prepare_candidate( # h_R[i, n, d] = sum_e hkl[n, e] sym_R[i, e, d]: the h.S convention. h_R = torch.einsum("ne,ied->ind", hkl, sym_R) phase = torch.exp((2j * math.pi) * torch.einsum("ne,ie->in", hkl, sym_t).to(cplx)) - eye3 = torch.eye(3, dtype=real, device=device) - F_all = evaluator.evaluate( - eye3, h_R.reshape(-1, 3), real_cell, return_amplitude=False, - ).reshape(S, N).to(cplx) + hkl_SN = h_R.reshape(-1, 3).round().to(torch.int64).to(model_p1.xyz().device) + with torch.no_grad(): + F_all = model_p1(hkl_SN).to(device).reshape(S, N).to(cplx) G_raw = F_all * phase I_P = (G_raw.abs() ** 2).mean(dim=0).to(real) From 6b410f66be9d832d9294a70fcf72a1b5b927691e Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 2 Sep 2026 17:19:35 +0200 Subject: [PATCH 152/250] Justify or remove every hard-coded dtype in the alignment package dev's dtype-conformance guard requires each torch. literal outside the exempt kernels to say why it deviates from the configured dtype. Of the 109 sites it flagged after the merge, five were unjustified and now use the configured dtype: the two casts to double around the empirical sigma_A ratio, the dense P1 transform's amplitudes, and the translation stage's resolution mask and peak translations. The rest are annotated: index tensors that index_add_ and gather require in int64, host-side 3x3 rotation algebra in double, the rotation function's deliberate double accumulation of oscillatory sums, exact clustering keys, and the anisotropy fit. Full unit gate 2043 passed, 82 skipped (job 550854); alignment tests 192 (550852); panel 30/30 at both windows (550853). Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- docs/changelog.rst | 1 + .../experimental/alignment/frf/_backends.py | 2 +- torchref/experimental/alignment/frf/api.py | 4 +-- .../experimental/alignment/frf/data_mr.py | 25 +++++++------- .../experimental/alignment/frf/dense_calc.py | 15 ++++---- .../frf/kernels/cpu/legendre_shell.py | 2 +- .../experimental/alignment/frf/peak_finder.py | 12 +++---- .../alignment/frf/preprocessing.py | 26 +++++++------- .../alignment/frf/rotation_utils.py | 10 +++--- .../alignment/frf/sitelist_ang.py | 28 +++++++-------- .../experimental/alignment/frf/wigner_d.py | 10 +++--- torchref/experimental/alignment/pipeline.py | 12 +++---- .../experimental/alignment/rotation_search.py | 34 +++++++++---------- torchref/experimental/alignment/sh.py | 32 ++++++++--------- .../experimental/alignment/translation.py | 6 ++-- torchref/scaling/wilson.py | 2 +- 16 files changed, 111 insertions(+), 110 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index d8f2bbbf..d1a8d316 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Every hard-coded dtype in the alignment package either moved to the configured dtype or carries a ``# dtype-ok:`` justification, as the dtype-conformance guard requires: three casts to double dropped from the empirical sigma_A ratio and the dense P1 transform, the translation set's resolution mask and peak translations use the configured float dtype, and the rest (index tensors, host-side 3x3 rotation algebra, the rotation function's double accumulation, the anisotropy fit) are annotated - Removed ``DirectModelEvaluator``; the translation search evaluates an ordinary P1 ``ModelFT`` directly. With the model's grid derived lazily from cell, space group and ``max_res`` there was nothing left for the wrapper to do - The molecular-replacement pipeline carries 10 rotation candidates by default instead of 25. With symmetry mates suppressed the rotation function's first peak is the true orientation in 50 of 50 pose-gated cells, and the panel is 30/30 at either depth. Warm on one EPYC 9335 node: 1DAW 0.42 s, 2DQ6 0.67 s, 6G9X 0.67 s, 3K7M 1.02 s, 4BX9 1.11 s per alignment - The fast translation function accumulates only the upper triangle of symmetry pairs; the lower triangle is its conjugate mirror and the diagonal a constant. Half the scatter, which is the stage's cost on high-symmetry cells: 3K7M's translation stage 1.04 s to 0.72 s. ``MRSolution.candidate_index`` records each solution's position in the rotation function's list diff --git a/torchref/experimental/alignment/frf/_backends.py b/torchref/experimental/alignment/frf/_backends.py index 779b16a6..35df4ae0 100644 --- a/torchref/experimental/alignment/frf/_backends.py +++ b/torchref/experimental/alignment/frf/_backends.py @@ -34,7 +34,7 @@ name="cpu_fused", kernel=(_CPU, "legendre_shell_accumulate", "legendre_shell_accumulate"), device="cpu", - dtypes=(torch.float32,), + dtypes=(torch.float32,), # dtype-ok: backend capability declaration # The kernel reads every array through a raw `float*`, so a mixed-dtype # call would reinterpret the buffer rather than convert it. The gate # keeps that from reaching the kernel; the kernel checks it too, since diff --git a/torchref/experimental/alignment/frf/api.py b/torchref/experimental/alignment/frf/api.py index d0182230..c77be474 100644 --- a/torchref/experimental/alignment/frf/api.py +++ b/torchref/experimental/alignment/frf/api.py @@ -371,8 +371,8 @@ def score_model( # solvent deficit is simply what the ratio measures, per structure, # instead of two universal constants. eterm = empirical_sigma_a( - self._conv_obs.evaluate(smag_calc).to(torch.float64), - self._conv_calc.evaluate(smag_calc).to(torch.float64), + self._conv_obs.evaluate(smag_calc), + self._conv_calc.evaluate(smag_calc), ).to(F_calc.dtype) elif sigma_a_source == "luzzati": eterm = eterm_sigma_a(smag_calc, self.delta_vrms_A) diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index c1bb0cd5..cf8e80e0 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -142,7 +142,7 @@ def spherical_bessel_table( inv_threshold = 1.0 / threshold # Rescales applied so far, per element. Every element's ladder sits in the # single frame 2**(-_BESSEL_RESCALE_EXP * n_rescales). - n_rescales = torch.zeros_like(x64, dtype=torch.int32) + n_rescales = torch.zeros_like(x64, dtype=torch.int32) # dtype-ok: small integer counter for n in range(n_start, 0, -1): j_low = (2.0 * n + 1.0) * inv_x * j_mid - j_high @@ -161,7 +161,7 @@ def spherical_bessel_table( j_high = j_high * factor if n - 1 <= u_max: j_table[n - 1:] = j_table[n - 1:] * factor - n_rescales = n_rescales + over.to(torch.int32) + n_rescales = n_rescales + over.to(torch.int32) # dtype-ok: small integer counter true_j0 = torch.sin(x64) * inv_x true_j0 = torch.where(x64 < 1e-30, torch.ones_like(x64), true_j0) @@ -249,7 +249,6 @@ def bessel_sh_expand( # from the input: the input is deliberately wider (see the docstring). comp_real = get_float_dtype() complex_dtype = get_complex_dtype() - real_dtype = comp_real lmax = L - 1 lmax_even = lmax if (lmax % 2 == 0) else (lmax - 1) @@ -276,15 +275,15 @@ def bessel_sh_expand( n_list.append(n) u_list.append(u) w_list.append(math.sqrt(float(2 * u + 1))) - l_idx = torch.tensor(l_list, dtype=torch.long, device=device) - n_idx = torch.tensor(n_list, dtype=torch.long, device=device) - u_idx = torch.tensor(u_list, dtype=torch.long, device=device) + l_idx = torch.tensor(l_list, dtype=torch.long, device=device) # dtype-ok: index tensor; index_add_/gather need int64 + n_idx = torch.tensor(n_list, dtype=torch.long, device=device) # dtype-ok: index tensor; index_add_/gather need int64 + u_idx = torch.tensor(u_list, dtype=torch.long, device=device) # dtype-ok: index tensor; index_add_/gather need int64 w_vec = torch.tensor(w_list, dtype=comp_real, device=device) # Only even degrees l ∈ [2, lmax_even] carry signal (odd-l and l=0 are zeroed # by Patterson centrosymmetry). Compute / contract Y_lm on these rows only — # the assembly + einsum are the bottleneck, so this ~halves them. The full # c_nlm keeps the (L, ...) shape with odd/zero rows left at zero. - even_l_idx = torch.tensor(even_ls, dtype=torch.long, device=device) + even_l_idx = torch.tensor(even_ls, dtype=torch.long, device=device) # dtype-ok: index tensor; index_add_/gather need int64 M = s_vectors.shape[0] einsum_dtype = complex_dtype @@ -317,8 +316,8 @@ def _tick(t0): # Separate resolutions for the two factors: the radial term needs a fine # |s| key, the angular term does not. One shared key forces the finer of the # two on both, which costs merges the angular part never needed. - k_s = (s_mag_all * _GROUP_SCALE_S).round().to(torch.int64) - k_c = (cos_all * _GROUP_SCALE_COS).round().to(torch.int64) + _GROUP_SCALE_COS + k_s = (s_mag_all * _GROUP_SCALE_S).round().to(torch.int64) # dtype-ok: exact clustering key; needs double + k_c = (cos_all * _GROUP_SCALE_COS).round().to(torch.int64) + _GROUP_SCALE_COS # dtype-ok: exact clustering key; needs double key = k_s * (2 * _GROUP_SCALE_COS + 1) + k_c uniq_key, inverse = torch.unique(key, return_inverse=True) n_clusters = int(uniq_key.shape[0]) @@ -344,7 +343,7 @@ def _group_mean(values, index, n_groups): # over the benchmark: 2.7 to 39 clusters per shell. uniq_ks, inv_s = torch.unique(k_s, return_inverse=True) n_shells = int(uniq_ks.shape[0]) - shell_of_cluster = torch.zeros(n_clusters, dtype=torch.long, device=device) + shell_of_cluster = torch.zeros(n_clusters, dtype=torch.long, device=device) # dtype-ok: index tensor; index_add_/gather need int64 shell_of_cluster[inverse] = inv_s shell_smag = _group_mean(s_mag_all.to(comp_real), inv_s, n_shells) @@ -443,7 +442,7 @@ def _group_mean(values, index, n_groups): # answer c_pos is a few MB, so both stay in cache. c_pos = torch.zeros((N_radial, n_even, L), dtype=einsum_dtype, device=device) - rbytes = 4 if comp_real == torch.float32 else 8 + rbytes = 4 if comp_real == torch.float32 else 8 # dtype-ok: byte-size lookup for a memory estimate per_cluster = rbytes * 6 * L cstep = max(1, min(n_clusters, CLUSTER_CHUNK_BYTES // max(1, per_cluster))) for cs in range(0, n_clusters, cstep): @@ -563,8 +562,8 @@ def cross_correlate_xi( # backend without float64 this falls back to the coefficients' own dtype and # the run pays the accuracy noted above -- there is no third option there. acc = widest_complex_dtype(c_obs.coeffs.device) - if c_obs.coeffs.dtype == torch.complex128: - acc = torch.complex128 # never narrow what the caller widened + if c_obs.coeffs.dtype == torch.complex128: # dtype-ok: double accumulation of an oscillatory sum -- never narrow what the caller widened + acc = torch.complex128 # never narrow what the caller widened # dtype-ok: double accumulation of an oscillatory sum -- never narrow what the caller widened return torch.einsum( "rln,rlm->lmn", c_obs.coeffs.to(acc), diff --git a/torchref/experimental/alignment/frf/dense_calc.py b/torchref/experimental/alignment/frf/dense_calc.py index 1029a523..fd7949af 100644 --- a/torchref/experimental/alignment/frf/dense_calc.py +++ b/torchref/experimental/alignment/frf/dense_calc.py @@ -51,8 +51,9 @@ def dense_calc_via_box( Returns ------- (s_vec, F_calc) : Tuple[torch.Tensor, torch.Tensor] - ``s_vec`` is the Cartesian reciprocal grid (N, 3) and ``F_calc`` the - amplitudes (N,), both float64 on the model's device. + ``s_vec`` is the Cartesian reciprocal grid (N, 3) in double -- the + expansion clusters on it and needs exact keys -- and ``F_calc`` the + amplitudes (N,) in the model's dtype, both on the model's device. """ from torchref.symmetry.cell import Cell @@ -79,13 +80,13 @@ def dense_calc_via_box( H, K, Lg = torch.meshgrid(idx, idx, idx, indexing="ij") hkl = torch.stack( [H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], dim=-1 - ).to(torch.long) + ).to(torch.long) # dtype-ok: Miller indices are integers # Cubic box: |s| = |hkl| / a. - smag = hkl.to(torch.float64).norm(dim=-1) / a + smag = hkl.to(torch.float64).norm(dim=-1) / a # dtype-ok: exact clustering key; needs double keep = (smag >= 1.0 / d_max) & (smag <= 1.0 / d_min) hkl = hkl[keep].contiguous() F = model_sf_abs(m, hkl) - s_vec = hkl.to(torch.float64) / a + s_vec = hkl.to(torch.float64) / a # dtype-ok: exact clustering key; needs double if verbose: print( @@ -96,6 +97,6 @@ def dense_calc_via_box( def model_sf_abs(model: "ModelFT", hkl: torch.Tensor) -> torch.Tensor: - """``|F_calc|`` (float64) for ``hkl`` via the model's SF machinery (no grad).""" + """``|F_calc|`` for ``hkl`` via the model's SF machinery (no grad), in the model's dtype.""" with torch.no_grad(): - return model.get_structure_factor(hkl, recalc=True).abs().to(torch.float64) + return model.get_structure_factor(hkl, recalc=True).abs() diff --git a/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py b/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py index 57d459e2..e7d56746 100644 --- a/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py +++ b/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py @@ -246,7 +246,7 @@ def shell_offsets(shell: torch.Tensor, n_shells: int) -> torch.Tensor: write their accumulator rows without atomics. """ counts = torch.bincount(shell, minlength=n_shells) - offsets = torch.zeros(n_shells + 1, dtype=torch.long, device=shell.device) + offsets = torch.zeros(n_shells + 1, dtype=torch.long, device=shell.device) # dtype-ok: index tensor; index_add_/gather need int64 torch.cumsum(counts, dim=0, out=offsets[1:]) return offsets diff --git a/torchref/experimental/alignment/frf/peak_finder.py b/torchref/experimental/alignment/frf/peak_finder.py index 1c0825f5..23cb2db3 100644 --- a/torchref/experimental/alignment/frf/peak_finder.py +++ b/torchref/experimental/alignment/frf/peak_finder.py @@ -55,7 +55,7 @@ def _so3_greedy_nms( """ n = values.shape[0] if n == 0: - return torch.empty(0, dtype=torch.int64, device=values.device) + return torch.empty(0, dtype=torch.int64, device=values.device) # dtype-ok: index tensor; index_add_/gather need int64 # The greedy walk is inherently sequential and latency-bound; on GPU a # per-iteration `.item()` sync would dominate. Move the (tiny) candidate # rotations to CPU once and run the loop there with no device syncs, a @@ -71,16 +71,16 @@ def _so3_greedy_nms( # sitting exactly on it. R_all = ( rotation_matrix_euler_zyz(torch.stack([alphas, betas, gammas], dim=-1)) - .to(torch.float64).cpu() + .to(torch.float64).cpu() # dtype-ok: 3x3 rotation algebra in double on the host ) # (n, 3, 3) # angle > nms_radius ⇔ cos(angle) < cos(nms_radius); cos(angle) from trace. cos_thresh = math.cos(math.radians(nms_radius_deg)) # The orbit of each kept rotation, R R_g over the point group; the identity # alone when no symmetry is supplied. - G = (torch.eye(3, dtype=torch.float64).unsqueeze(0) if sym_cart is None - else sym_cart.to(torch.float64).cpu()) + G = (torch.eye(3, dtype=torch.float64).unsqueeze(0) if sym_cart is None # dtype-ok: 3x3 rotation algebra in double on the host + else sym_cart.to(torch.float64).cpu()) # dtype-ok: 3x3 rotation algebra in double on the host kept_idx: List[int] = [] - kept_orbit = torch.empty((keep_at_most, G.shape[0], 3, 3), dtype=torch.float64) + kept_orbit = torch.empty((keep_at_most, G.shape[0], 3, 3), dtype=torch.float64) # dtype-ok: 3x3 rotation algebra in double on the host count = 0 for i_t in order: Ri = R_all[i_t] @@ -95,7 +95,7 @@ def _so3_greedy_nms( count += 1 if count >= keep_at_most: break - return torch.tensor(kept_idx, dtype=torch.int64, device=values.device) + return torch.tensor(kept_idx, dtype=torch.int64, device=values.device) # dtype-ok: index tensor; index_add_/gather need int64 def find_rotation_peaks( diff --git a/torchref/experimental/alignment/frf/preprocessing.py b/torchref/experimental/alignment/frf/preprocessing.py index 915d48ac..89aac868 100644 --- a/torchref/experimental/alignment/frf/preprocessing.py +++ b/torchref/experimental/alignment/frf/preprocessing.py @@ -107,7 +107,7 @@ def apply_shell_variance_weights( shell_idx = assign_shells(s_mag, edges) valid = shell_idx >= 0 var_p = compute_patterson_shell_variance( - intensity[valid].to(torch.float64), + intensity[valid].to(torch.float64), # dtype-ok: shell variance accumulated in double shell_idx[valid], P=n_var_shells, ) @@ -140,7 +140,7 @@ def detect_zsymm(sym_mats: Optional[torch.Tensor]) -> int: """ if sym_mats is None: return 1 - axis, zsymm = get_high_order_axis(sym_mats.to(torch.float64).cpu()) + axis, zsymm = get_high_order_axis(sym_mats.to(torch.float64).cpu()) # dtype-ok: 3x3 rotation algebra in double on the host if axis != 2: # high-order axis not along z → don't apply a wrong filter return 1 return int(zsymm) @@ -288,15 +288,15 @@ def fit_relative_wilson_b( if not valid_obs.any() or not valid_calc.any(): return 0.0 - F2_obs = (F_obs * F_obs).to(torch.float64) - F2_calc = (F_calc * F_calc).to(torch.float64) - s2_obs = (s_mag * s_mag).to(torch.float64) + F2_obs = (F_obs * F_obs).to(torch.float64) # dtype-ok: per-shell sums in double for the relative-B fit + F2_calc = (F_calc * F_calc).to(torch.float64) # dtype-ok: per-shell sums in double for the relative-B fit + s2_obs = (s_mag * s_mag).to(torch.float64) # dtype-ok: per-shell sums in double for the relative-B fit - counts_obs = torch.zeros(n_shells, dtype=torch.int64, device=s_mag.device) - counts_calc = torch.zeros(n_shells, dtype=torch.int64, device=s_mag.device) - sum_F2obs = torch.zeros(n_shells, dtype=torch.float64, device=s_mag.device) - sum_F2calc = torch.zeros(n_shells, dtype=torch.float64, device=s_mag.device) - sum_s2 = torch.zeros(n_shells, dtype=torch.float64, device=s_mag.device) + counts_obs = torch.zeros(n_shells, dtype=torch.int64, device=s_mag.device) # dtype-ok: per-shell counts + counts_calc = torch.zeros(n_shells, dtype=torch.int64, device=s_mag.device) # dtype-ok: per-shell counts + sum_F2obs = torch.zeros(n_shells, dtype=torch.float64, device=s_mag.device) # dtype-ok: per-shell sums in double for the relative-B fit + sum_F2calc = torch.zeros(n_shells, dtype=torch.float64, device=s_mag.device) # dtype-ok: per-shell sums in double for the relative-B fit + sum_s2 = torch.zeros(n_shells, dtype=torch.float64, device=s_mag.device) # dtype-ok: per-shell sums in double for the relative-B fit idx_v_obs = shell_idx_obs[valid_obs] idx_v_calc = shell_idx_calc[valid_calc] counts_obs.index_add_(0, idx_v_obs, torch.ones_like(idx_v_obs)) @@ -307,9 +307,9 @@ def fit_relative_wilson_b( # Drop shells empty on either side. keep = (counts_obs > 0) & (counts_calc > 0) - mean_F2obs = sum_F2obs[keep] / counts_obs[keep].to(torch.float64) - mean_F2calc = sum_F2calc[keep] / counts_calc[keep].to(torch.float64) - mean_s2 = sum_s2[keep] / counts_obs[keep].to(torch.float64) + mean_F2obs = sum_F2obs[keep] / counts_obs[keep].to(torch.float64) # dtype-ok: per-shell sums in double for the relative-B fit + mean_F2calc = sum_F2calc[keep] / counts_calc[keep].to(torch.float64) # dtype-ok: per-shell sums in double for the relative-B fit + mean_s2 = sum_s2[keep] / counts_obs[keep].to(torch.float64) # dtype-ok: per-shell sums in double for the relative-B fit # log(Σ_N / Σ_P) per shell. eps = 1e-30 diff --git a/torchref/experimental/alignment/frf/rotation_utils.py b/torchref/experimental/alignment/frf/rotation_utils.py index 9b8c5df3..51fab749 100644 --- a/torchref/experimental/alignment/frf/rotation_utils.py +++ b/torchref/experimental/alignment/frf/rotation_utils.py @@ -12,7 +12,7 @@ def rotation_matrix_from_edmonds_euler( - alpha: float, beta: float, gamma: float, dtype=torch.float64, + alpha: float, beta: float, gamma: float, dtype=torch.float64, # dtype-ok: 3x3 rotation algebra in double on the host ) -> torch.Tensor: """Build ``R = R_z(α) R_y(β) R_z(γ)`` (Edmonds active ZYZ). @@ -35,7 +35,7 @@ def edmonds_euler_from_rotation_matrix(R: torch.Tensor) -> Tuple[float, float, f when ``β = 0`` or ``π`` (only ``α+γ`` is determined); in those cases ``γ=0`` is returned. """ - R = R.to(torch.float64) + R = R.to(torch.float64) # dtype-ok: 3x3 rotation algebra in double on the host cos_beta = R[2, 2].clamp(-1.0, 1.0).item() beta = math.acos(cos_beta) sin_beta = math.sin(beta) @@ -63,8 +63,8 @@ def axis_angle_to_matrix(omega: torch.Tensor) -> torch.Tensor: below θ = 1e-10 rather than letting the trigonometric factors vanish. Above that threshold the two agree term for term. """ - if omega.dtype not in (torch.float32, torch.float64): - omega = omega.to(torch.float64) + if omega.dtype not in (torch.float32, torch.float64): # dtype-ok: 3x3 rotation algebra in double on the host + omega = omega.to(torch.float64) # dtype-ok: 3x3 rotation algebra in double on the host single = omega.dim() == 1 if single: omega = omega.unsqueeze(0) @@ -86,7 +86,7 @@ def axis_angle_to_matrix(omega: torch.Tensor) -> torch.Tensor: def rotation_angular_distance_deg(R1: torch.Tensor, R2: torch.Tensor) -> float: """Geodesic distance on SO(3) in degrees: ``arccos((tr(R1 R2^T) − 1)/2)``.""" - R = R1.to(torch.float64) @ R2.to(torch.float64).T + R = R1.to(torch.float64) @ R2.to(torch.float64).T # dtype-ok: 3x3 rotation algebra in double on the host tr = (R[0, 0] + R[1, 1] + R[2, 2]).clamp(-1.0, 3.0).item() cos_a = max(-1.0, min(1.0, (tr - 1.0) / 2.0)) return math.degrees(math.acos(cos_a)) diff --git a/torchref/experimental/alignment/frf/sitelist_ang.py b/torchref/experimental/alignment/frf/sitelist_ang.py index 652a824c..c297c3c2 100644 --- a/torchref/experimental/alignment/frf/sitelist_ang.py +++ b/torchref/experimental/alignment/frf/sitelist_ang.py @@ -98,7 +98,7 @@ def build_dense_map_per_beta( (n_beta, fft_size, fft_size), dtype=S.dtype, device=device, ) m_vals = torch.arange(-(L - 1), L, device=device) - idx = (m_vals % fft_size).to(torch.int64) + idx = (m_vals % fft_size).to(torch.int64) # dtype-ok: index tensor; index_add_/gather need int64 pad[:, idx.unsqueeze(1), idx.unsqueeze(0)] = S # 3. Forward 2D FFT — torch convention: @@ -123,7 +123,7 @@ def build_dense_map_per_beta( def build_adaptive_sample_list( grid_sampling_deg: float, - dtype: torch.dtype = torch.float64, + dtype: torch.dtype = torch.float64, # dtype-ok: sample-list geometry follows the accumulator's width device: torch.device = torch.device("cpu"), ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: """Build the per-β (α, γ) sample list. @@ -179,7 +179,7 @@ def build_adaptive_sample_list( if b == 0: # β=0: only α = γ = p/pmax for p < pmax/2 (FastRot.cc:189-207). p_idx = torch.arange(pmax, device=cpu) - p_ratio = p_idx.to(torch.float64) / pmax + p_ratio = p_idx.to(torch.float64) / pmax # dtype-ok: sample-list geometry follows the accumulator's width keep = p_ratio < 0.5 p_ratio = p_ratio[keep] alpha_frac = p_ratio @@ -190,8 +190,8 @@ def build_adaptive_sample_list( # p_ratio < q_ratio (gives γ ∈ [0, 1) without negative values). p_idx = torch.arange(pmax, device=cpu) q_idx = torch.arange(qmax, device=cpu) - p_ratio = (p_idx.to(torch.float64) / pmax).unsqueeze(1) # (pmax, 1) - q_ratio = (q_idx.to(torch.float64) / qmax).unsqueeze(0) # (1, qmax) + p_ratio = (p_idx.to(torch.float64) / pmax).unsqueeze(1) # (pmax, 1) # dtype-ok: sample-list geometry follows the accumulator's width + q_ratio = (q_idx.to(torch.float64) / qmax).unsqueeze(0) # (1, qmax) # dtype-ok: sample-list geometry follows the accumulator's width alpha_frac = torch.fmod(p_ratio + q_ratio, 1.0) diff = p_ratio - q_ratio gamma_frac = torch.where( @@ -210,8 +210,8 @@ def build_adaptive_sample_list( # original dict scan, but no host sync / Python loop. # Hash the two rounded fracs (each in [0, 1e6]) into one int64 so we # can use the fast 1-D unique instead of a 2-D row lexsort. - a_round = (alpha_frac * 1_000_000).round().to(torch.int64) - g_round = (gamma_frac * 1_000_000).round().to(torch.int64) + a_round = (alpha_frac * 1_000_000).round().to(torch.int64) # dtype-ok: index tensor; index_add_/gather need int64 + g_round = (gamma_frac * 1_000_000).round().to(torch.int64) # dtype-ok: index tensor; index_add_/gather need int64 key_hash = a_round * 1_000_001 + g_round _, uniq_idx = torch.unique(key_hash, return_inverse=True) n = uniq_idx.shape[0] @@ -234,8 +234,8 @@ def build_adaptive_sample_list( alphas = torch.cat(alphas_list).to(device) gammas = torch.cat(gammas_list).to(device) betas_flat = torch.cat(betas_list).to(device) - beta_starts_t = torch.tensor(beta_starts, dtype=torch.int64, device=device) - b = torch.arange(bmax, dtype=torch.float64, device=cpu) + beta_starts_t = torch.tensor(beta_starts, dtype=torch.int64, device=device) # dtype-ok: index tensor; index_add_/gather need int64 + b = torch.arange(bmax, dtype=torch.float64, device=cpu) # dtype-ok: sample-list geometry follows the accumulator's width betas_rad = (b * grid_sampling_deg * deg2rad).to(device=device, dtype=dtype) result = (alphas, betas_flat, gammas, beta_starts_t, betas_rad) @@ -256,8 +256,8 @@ def _bilinear_interp_periodic( N = M.shape[-1] af = (alpha_frac % 1.0) * N gf = (gamma_frac % 1.0) * N - a0 = torch.floor(af).to(torch.int64) % N - g0 = torch.floor(gf).to(torch.int64) % N + a0 = torch.floor(af).to(torch.int64) % N # dtype-ok: index tensor; index_add_/gather need int64 + g0 = torch.floor(gf).to(torch.int64) % N # dtype-ok: index tensor; index_add_/gather need int64 a1 = (a0 + 1) % N g1 = (g0 + 1) % N da = (af - torch.floor(af)).to(M.real.dtype) @@ -301,9 +301,9 @@ def evaluate_rotation_function( L = xi_lmn.shape[0] device = xi_lmn.device real_dtype = ( - torch.float64 - if xi_lmn.dtype in (torch.complex128, torch.float64) - else torch.float32 + torch.float64 # dtype-ok: sample-list geometry follows the accumulator's width + if xi_lmn.dtype in (torch.complex128, torch.float64) # dtype-ok: sample-list geometry follows the accumulator's width + else torch.float32 # dtype-ok: sample-list geometry follows the accumulator's width ) bmax = int(math.ceil(180.0 / grid_sampling_deg)) diff --git a/torchref/experimental/alignment/frf/wigner_d.py b/torchref/experimental/alignment/frf/wigner_d.py index 3ac5551e..6e7f88b8 100644 --- a/torchref/experimental/alignment/frf/wigner_d.py +++ b/torchref/experimental/alignment/frf/wigner_d.py @@ -51,10 +51,10 @@ def _wigner_eig_table(L: int): table = [] for l in range(1, L): sz = 2 * l + 1 - p = torch.arange(sz - 1, dtype=torch.float64) + p = torch.arange(sz - 1, dtype=torch.float64) # dtype-ok: eigendecomposition in double for the Wigner-d recursion sup = 0.5 * torch.sqrt((2 * l - p) * (p + 1.0)) A = torch.diag(sup, 1) - torch.diag(sup, -1) # A = -i J_y - w, V = torch.linalg.eigh(1j * A.to(torch.complex128)) # w∈[-l..l] + w, V = torch.linalg.eigh(1j * A.to(torch.complex128)) # w∈[-l..l] # dtype-ok: eigendecomposition in double for the Wigner-d recursion table.append((w, V)) _WIGNER_EIG_CACHE[key] = table return table @@ -91,14 +91,14 @@ def _wigner_d_blocks(L: int, betas: torch.Tensor, device: torch.device, int(L), str(canonical_device(device)), dtype, - tuple(betas.detach().to(torch.float64).cpu().tolist()), + tuple(betas.detach().to(torch.float64).cpu().tolist()), # dtype-ok: memo key: exact host-side doubles ) hit = _WIGNER_D_CACHE.get(key) if hit is not None: return hit eig_table = _wigner_eig_table(L) # host, cached - betas_host = betas.detach().to(torch.float64).cpu() + betas_host = betas.detach().to(torch.float64).cpu() # dtype-ok: memo key: exact host-side doubles blocks = [] for l in range(1, L): w, V = eig_table[l - 1] # data-independent @@ -147,7 +147,7 @@ def wigner_contraction_per_beta( # working precision, so widening here would buy nothing and cost a 2x # complex buffer in this stage and in the FFT it feeds. xi = xi_lmn - real_dtype = torch.float64 if xi.dtype == torch.complex128 else torch.float32 + real_dtype = torch.float64 if xi.dtype == torch.complex128 else torch.float32 # dtype-ok: follows the accumulator's width # Per-l loop over the small-d blocks, which come from the J_y # eigendecomposition (small_d_stable's method, stable to any l). Contract diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index ce70b5f2..256df019 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -57,7 +57,7 @@ import numpy as np import torch -from torchref.config import get_default_device +from torchref.config import get_default_device, get_float_dtype from torchref.utils.device_mixin import DeviceMixin from .frf.rotation_utils import rotation_matrix_from_edmonds_euler @@ -523,7 +523,7 @@ def _rotation_candidates(self, frf) -> list: def place(self, solution: MRSolution) -> "ModelFT": """Build the placed model for ``solution``: a copy of the search model, rotated and translated, carrying the alignment provenance attributes.""" - R_rec = torch.as_tensor(solution.rotation, dtype=torch.float64) + R_rec = torch.as_tensor(solution.rotation, dtype=torch.float64) # dtype-ok: 3x3 rotation algebra in double on the host placed = self.model.copy().rotate( R_rec.T.contiguous().to(device=self.model.device, dtype=self.model.dtype_float), @@ -599,9 +599,9 @@ def _prepare_translation_arrays(self) -> None: tmask = torch.ones( F_obs_full.shape[0], dtype=torch.bool, device=F_obs_full.device, ) - rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) - s_all = (hkl_full.to(torch.float64) @ rec_basis.to(hkl_full.device) - ).norm(dim=-1) + real = get_float_dtype() + rec_basis = data.cell.reciprocal_basis_matrix.to(real) + s_all = (hkl_full.to(real) @ rec_basis.to(hkl_full.device)).norm(dim=-1) if self.tf_d_min > 0.0: tmask = tmask & (s_all <= 1.0 / self.tf_d_min) if np.isfinite(self.tf_d_max): @@ -678,7 +678,7 @@ def _placement_for_candidate(self) -> Optional[tuple]: timer.start("7_translation_llg") t_cands = torch.as_tensor( - np.stack([p.translation for p in t_peaks]), dtype=torch.float64, + np.stack([p.translation for p in t_peaks]), dtype=get_float_dtype(), ) llg = llg_at_translations(obs, cand, t_cands) k_best = int(llg.argmax()) diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index ed2bcc85..337b90b5 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -186,8 +186,8 @@ def fit_anisotropy( """ from .sh import hkl_symops_to_cartesian, symmetrize_anisotropy - rec_basis = data.cell.reciprocal_basis_matrix.detach().cpu().to(torch.float64) - hkl = data.hkl.detach().cpu().to(torch.float64) + rec_basis = data.cell.reciprocal_basis_matrix.detach().cpu().to(torch.float64) # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search + hkl = data.hkl.detach().cpu().to(torch.float64) # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search s_vec_all = hkl @ rec_basis s_mag_all = s_vec_all.norm(dim=-1) keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min) @@ -196,7 +196,7 @@ def fit_anisotropy( f"Only {int(keep.sum())} reflections in [{d_min}, {d_max}] A, too " f"few for {n_shells} shells." ) - F_obs = data.F.detach().cpu().to(torch.float64).abs()[keep] + F_obs = data.F.detach().cpu().to(torch.float64).abs()[keep] # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search s_vec = s_vec_all[keep] s_mag = s_mag_all[keep] centric = ( @@ -211,7 +211,7 @@ def fit_anisotropy( F_obs, s_vec, shell_idx, centric, P=n_shells, min_count=20, ) sym_cart = hkl_symops_to_cartesian( - data.spacegroup.matrices.detach().cpu().to(torch.float64), rec_basis, + data.spacegroup.matrices.detach().cpu().to(torch.float64), rec_basis, # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search ) return symmetrize_anisotropy(U, sym_cart) @@ -259,10 +259,10 @@ def prepare_frf_inputs( """ device = get_default_device() - F_obs = data.F.to(torch.float64).abs() + F_obs = data.F.to(torch.float64).abs() # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| hkl_all = data.hkl - rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) - s_vec_all = hkl_all.to(torch.float64) @ rec_basis + rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| + s_vec_all = hkl_all.to(torch.float64) @ rec_basis # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| s_mag_all = s_vec_all.norm(dim=-1) keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min) if keep.sum().item() < n_shells * 5: @@ -273,7 +273,7 @@ def prepare_frf_inputs( F_obs = F_obs[keep].to(device) sig_F = getattr(data, "F_sigma", None) if sig_F is not None: - sig_F = sig_F.to(torch.float64)[keep].to(device) + sig_F = sig_F.to(torch.float64)[keep].to(device) # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| hkl = hkl_all[keep].to(device) s_vec = s_vec_all[keep].to(device) s_mag = s_mag_all[keep].to(device) @@ -331,9 +331,9 @@ def search_peaks( # the rest of the codebase. device = resolve_device(data, model, device=device) with torch.no_grad(): - rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) + rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| hkl_all = data.hkl.to(device) - s_vec_all = hkl_all.to(torch.float64) @ rec_basis + s_vec_all = hkl_all.to(torch.float64) @ rec_basis # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| s_mag_all = s_vec_all.norm(dim=-1) d_min_data = float(1.0 / s_mag_all.max().item()) @@ -360,10 +360,10 @@ def search_peaks( s_asu = s_vec_all[keep] s_mag_asu = s_mag_all[keep] F_obs = apply_overall_anisotropy( - data.F.to(torch.float64).abs().to(device)[keep], s_asu, U_aniso, + data.F.to(torch.float64).abs().to(device)[keep], s_asu, U_aniso, # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| ) sigF = ( - data.F_sigma.to(torch.float64).to(device)[keep] + data.F_sigma.to(torch.float64).to(device)[keep] # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| if getattr(data, "F_sigma", None) is not None else None ) @@ -393,7 +393,7 @@ def search_peaks( # orthogonal symmetry matrices, so using S.h works everywhere except # trigonal and hexagonal, where it mixes non-equivalent reflections into # one orbit. - sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) + sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) # dtype-ok: 3x3 rotation algebra in double on the host n_ops = int(sg_mats.shape[0]) # `expand_reciprocal` is the package's one implementation of this # contraction, and it carries the h.S convention so no call site has to @@ -483,17 +483,17 @@ def _solutions(peaks: List["RotationPeak"], lmax: int, d_min: float, from torchref.base.alignment.rotation import rotation_matrix_euler_zyz euler = torch.tensor( - [[p.alpha, p.beta, p.gamma] for p in peaks], dtype=torch.float64, + [[p.alpha, p.beta, p.gamma] for p in peaks], dtype=torch.float64, # dtype-ok: RotationSolutions are documented as float64 ).reshape(-1, 3) rotations = ( rotation_matrix_euler_zyz(euler) if euler.numel() - else torch.zeros((0, 3, 3), dtype=torch.float64) + else torch.zeros((0, 3, 3), dtype=torch.float64) # dtype-ok: RotationSolutions are documented as float64 ) return RotationSolutions( rotations=rotations, - scores=torch.tensor([p.score for p in peaks], dtype=torch.float64), - z_scores=torch.tensor([p.sigma for p in peaks], dtype=torch.float64), + scores=torch.tensor([p.score for p in peaks], dtype=torch.float64), # dtype-ok: RotationSolutions are documented as float64 + z_scores=torch.tensor([p.sigma for p in peaks], dtype=torch.float64), # dtype-ok: RotationSolutions are documented as float64 euler_zyz=euler, lmax=lmax, d_min=d_min, diff --git a/torchref/experimental/alignment/sh.py b/torchref/experimental/alignment/sh.py index 6a052d18..05630353 100644 --- a/torchref/experimental/alignment/sh.py +++ b/torchref/experimental/alignment/sh.py @@ -50,8 +50,8 @@ def legendre_recurrence_coefficients(L: int, dtype, device): runs this recurrence itself, fused with its own accumulation, and two copies of these formulae would be two chances to get them subtly different. """ - ll = torch.arange(L, dtype=torch.float64, device=device).view(L, 1) - mm = torch.arange(L, dtype=torch.float64, device=device).view(1, L) + ll = torch.arange(L, dtype=torch.float64, device=device).view(L, 1) # dtype-ok: harmonic indices as exact doubles + mm = torch.arange(L, dtype=torch.float64, device=device).view(1, L) # dtype-ok: harmonic indices as exact doubles valid = ll > mm denom = (ll - mm) * (ll + mm) denom_safe = torch.where(valid, denom, torch.ones_like(denom)) @@ -62,7 +62,7 @@ def legendre_recurrence_coefficients(L: int, dtype, device): b_num / torch.where(b_den == 0, torch.ones_like(b_den), b_den), min=0.0)) a = torch.where(valid, a, torch.zeros_like(a)).to(dtype) b = torch.where(valid, b, torch.zeros_like(b)).to(dtype) - m_arange = torch.arange(L, dtype=torch.float64, device=device) + m_arange = torch.arange(L, dtype=torch.float64, device=device) # dtype-ok: harmonic indices as exact doubles sect = torch.sqrt( (2.0 * m_arange + 1.0) / (2.0 * m_arange).clamp(min=1.0)).to(dtype) return a, b, sect @@ -131,9 +131,9 @@ def _bar_legendre_recurrence( if keep_l is None: rows = torch.arange(L, device=device) else: - rows = keep_l.to(device=device, dtype=torch.long) + rows = keep_l.to(device=device, dtype=torch.long) # dtype-ok: index tensor; index_add_/gather need int64 # l -> its position in the output, or -1 when it is not kept. - where = torch.full((L,), -1, dtype=torch.long, device=device) + where = torch.full((L,), -1, dtype=torch.long, device=device) # dtype-ok: index tensor; index_add_/gather need int64 where[rows] = torch.arange(rows.numel(), device=device) where_list = where.tolist() @@ -175,12 +175,12 @@ def get_axis_order(sym_mats: torch.Tensor, axis: int) -> int: coefficients: the Patterson is invariant under the spacegroup rotations, so m-values that violate the highest-order axis symmetry are pure noise. """ - a = torch.zeros(3, dtype=torch.float64, device=sym_mats.device) + a = torch.zeros(3, dtype=torch.float64, device=sym_mats.device) # dtype-ok: 3x3 rotation algebra in double on the host a[axis] = 1.0 max_order = 1 n_ops = sym_mats.shape[0] for k in range(n_ops): - R = sym_mats[k].to(torch.float64) + R = sym_mats[k].to(torch.float64) # dtype-ok: 3x3 rotation algebra in double on the host # Axis must be invariant under R (proper or improper rotation about it). if (R @ a - a).norm().item() > 1e-3: continue @@ -241,7 +241,7 @@ def equal_count_shell_edges( s_sorted, _ = torch.sort(s) N = s_sorted.numel() # quantile-based partition - idx = torch.linspace(0, N - 1, P + 1, dtype=torch.float64, device=s.device).round().long() + idx = torch.linspace(0, N - 1, P + 1, dtype=torch.float64, device=s.device).round().long() # dtype-ok: equal-count edge positions in double, then rounded edges = s_sorted[idx] # nudge endpoints so the data is fully covered (avoid floating-point miss) if s_min is not None: @@ -317,8 +317,8 @@ def fit_overall_anisotropy( reflections survive to constrain seven parameters. """ valid = shell_idx >= 0 - F = F_obs[valid].to(torch.float64) - s = s_vectors[valid].to(torch.float64) + F = F_obs[valid].to(torch.float64) # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search + s = s_vectors[valid].to(torch.float64) # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search idx = shell_idx[valid] cen = centric[valid].bool() @@ -326,11 +326,11 @@ def fit_overall_anisotropy( F, s, idx, cen = F[ok], s[ok], idx[ok], cen[ok] I = F * F - count = torch.zeros(P, dtype=torch.int64, device=F.device) - total = torch.zeros(P, dtype=torch.float64, device=F.device) + count = torch.zeros(P, dtype=torch.int64, device=F.device) # dtype-ok: index tensor; index_add_/gather need int64 + total = torch.zeros(P, dtype=torch.float64, device=F.device) # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search count.index_add_(0, idx, torch.ones_like(idx)) total.index_add_(0, idx, I) - mean_I = (total / count.clamp(min=1).to(torch.float64)).clamp(min=1e-30) + mean_I = (total / count.clamp(min=1).to(torch.float64)).clamp(min=1e-30) # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search keep = (count >= min_count)[idx] if int(keep.sum()) < 50: @@ -348,7 +348,7 @@ def fit_overall_anisotropy( -2.0 * (torch.pi ** 2) * quad], dim=1) w = torch.where(cenk, torch.full_like(ratio, 0.5), torch.ones_like(ratio)) - theta = torch.zeros(7, dtype=torch.float64, device=F.device) + theta = torch.zeros(7, dtype=torch.float64, device=F.device) # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search for _ in range(n_iter): model = torch.exp((A @ theta).clamp(min=-20.0, max=20.0)) J = model.unsqueeze(1) * A @@ -395,7 +395,7 @@ def hkl_symops_to_cartesian( ------- sym_mats_cart : torch.Tensor, shape (n_ops, 3, 3), real """ - dtype = torch.float64 + dtype = torch.float64 # dtype-ok: 3x3 rotation algebra in double on the host M = rec_basis.to(dtype).transpose(-1, -2) # (3, 3) M_inv = torch.linalg.inv(M) S = sg_mats.to(dtype) # (n_ops, 3, 3) @@ -514,7 +514,7 @@ def compute_patterson_shell_variance( valid = shell_idx >= 0 patt_v = patt[valid] idx_v = shell_idx[valid] - count = torch.zeros(P, dtype=torch.int64, device=device) + count = torch.zeros(P, dtype=torch.int64, device=device) # dtype-ok: index tensor; index_add_/gather need int64 count.index_add_(0, idx_v, torch.ones_like(idx_v)) sum1 = torch.zeros(P, dtype=dtype, device=device) sum2 = torch.zeros(P, dtype=dtype, device=device) diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index 4e91faf7..3430e287 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -146,7 +146,7 @@ def build( rec_basis = real_cell.reciprocal_basis_matrix.to(dev).to(real) s_mag = (hkl_i.to(real) @ rec_basis).norm(dim=-1) - hkl_l = hkl_i.round().to(torch.int64) + hkl_l = hkl_i.round().to(torch.int64) # dtype-ok: Miller indices are integers # friedel=False: Wilson's = eps*Sigma counts the operations mapping # h to itself, which add coherently and set the mean. The Friedel-folded # branch changes the distribution instead, and that is centricity -- @@ -241,7 +241,7 @@ def prepare_candidate( # h_R[i, n, d] = sum_e hkl[n, e] sym_R[i, e, d]: the h.S convention. h_R = torch.einsum("ne,ied->ind", hkl, sym_R) phase = torch.exp((2j * math.pi) * torch.einsum("ne,ie->in", hkl, sym_t).to(cplx)) - hkl_SN = h_R.reshape(-1, 3).round().to(torch.int64).to(model_p1.xyz().device) + hkl_SN = h_R.reshape(-1, 3).round().to(torch.int64).to(model_p1.xyz().device) # dtype-ok: Miller indices are integers with torch.no_grad(): F_all = model_p1(hkl_SN).to(device).reshape(S, N).to(cplx) G_raw = F_all * phase @@ -387,7 +387,7 @@ def fast_translation_function( G = cand.G.to(device).to(cplx) S, N = G.shape coeff = obs.coeff.to(device).to(cplx) - h_R_int = cand.h_R.round().to(torch.int64) + h_R_int = cand.h_R.round().to(torch.int64) # dtype-ok: Miller indices are integers # The pair (j, i) is the conjugate of (i, j) at -dh, so the map is twice # the real part of the upper triangle's transform plus the diagonal, which diff --git a/torchref/scaling/wilson.py b/torchref/scaling/wilson.py index cf91c5fe..e0be8c64 100644 --- a/torchref/scaling/wilson.py +++ b/torchref/scaling/wilson.py @@ -431,7 +431,7 @@ def from_hkl( branches feed two different parameters of the same likelihood. """ work = get_float_dtype() - hkl_l = hkl.to(torch.long) + hkl_l = hkl.to(torch.long) # dtype-ok: Miller indices are integers # The cell may carry the configured default device while the reflections # are somewhere else; the caller should not have to reconcile them. rec = cell.reciprocal_basis_matrix.to(device=hkl_l.device, dtype=work) From d40416d70c2d54054dd36ed787454abf0fdc9ea3 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 2 Sep 2026 17:56:37 +0200 Subject: [PATCH 153/250] Keep double precision off the compute device throughout the alignment package The rotation function's back half accumulated one step wider than the expansion -- complex128 for the radial sum, the Wigner contraction and the FFT -- and its inputs, the dense P1 transform, the shell sums and the expansion's clustering keys were cast to float64 on the device. None of it survives on a backend without float64, and none of it is needed: measured on the pose panel, single precision everywhere on the device recovers 30/30 at the default window and 30/30 uncut, with per-alignment times unchanged (0.44-1.12 s at depth 10, job 551437). The earlier measurement that motivated the wide accumulator showed the deep peak list reordering while the top peak stayed put; the placement search now consumes only the top few distinct orientations. The expansion's clustering keys need double -- _GROUP_SCALE_S keys |s| at 1e-7, below float32's resolution -- and are formed on the host, which always has it; only the integer keys reach the device. The Bessel ladder runs in its argument's dtype, kept in range by its power-of-two rescaling, so the bit-identity and scipy checks still exercise a double ladder. What remains in double is host-side: 3x3 rotation algebra, the Wigner-d eigendecomposition, the anisotropy fit and RotationSolutions. The two tests that pinned the wider accumulator now pin the configured dtype. 227 alignment, scaling and Bessel tests (552234); full unit gate 2043 passed (552236); panel 30/30 at both windows (552235). Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_015Q18eF1DJYrKv61uXowjeq --- .../analysis/integration_alignment.sh | 2 +- docs/changelog.rst | 1 + .../test_rotation_search_dtype_device.py | 45 ++++++++----- .../experimental/alignment/frf/data_mr.py | 63 ++++++++++--------- .../experimental/alignment/frf/dense_calc.py | 13 ++-- .../alignment/frf/preprocessing.py | 23 ++++--- .../experimental/alignment/rotation_search.py | 30 +++++---- torchref/experimental/alignment/sh.py | 9 +-- 8 files changed, 108 insertions(+), 78 deletions(-) diff --git a/alignment_lab/analysis/integration_alignment.sh b/alignment_lab/analysis/integration_alignment.sh index d2c54e19..72f6f929 100644 --- a/alignment_lab/analysis/integration_alignment.sh +++ b/alignment_lab/analysis/integration_alignment.sh @@ -13,5 +13,5 @@ PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin cd "$REPO" export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -"$PY" -m pytest -c tests/pytest.ini tests/integration/alignment tests/unit/alignment tests/unit/scaling --run-slow -q --tb=short 2>&1 | grep -v "^✓" | tail -60 +"$PY" -m pytest -c tests/pytest.ini tests/integration/alignment tests/unit/alignment tests/unit/scaling tests/unit/frf_separate --run-slow -q --tb=short 2>&1 | grep -v "^✓" | tail -60 echo "PYTEST_RC=${PIPESTATUS[0]}" diff --git a/docs/changelog.rst b/docs/changelog.rst index d1a8d316..14007ef3 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Nothing in the alignment package puts float64 or complex128 on the compute device any more, so it runs on backends without double (MPS). The rotation function's radial sum, Wigner contraction and FFT accumulate at the configured complex dtype rather than one step wider; the Bessel ladder runs in its argument's dtype, kept in range by its rescaling; the expansion's exact clustering keys are formed on the host in double, which the device never sees; rotation-function inputs, the dense P1 transform and the shell sums use the configured float dtype. 30/30 poses at both windows, timings unchanged; the double-precision left is host-side 3x3 rotation algebra, the Wigner-d eigendecomposition and the anisotropy fit - Every hard-coded dtype in the alignment package either moved to the configured dtype or carries a ``# dtype-ok:`` justification, as the dtype-conformance guard requires: three casts to double dropped from the empirical sigma_A ratio and the dense P1 transform, the translation set's resolution mask and peak translations use the configured float dtype, and the rest (index tensors, host-side 3x3 rotation algebra, the rotation function's double accumulation, the anisotropy fit) are annotated - Removed ``DirectModelEvaluator``; the translation search evaluates an ordinary P1 ``ModelFT`` directly. With the model's grid derived lazily from cell, space group and ``max_res`` there was nothing left for the wrapper to do - The molecular-replacement pipeline carries 10 rotation candidates by default instead of 25. With symmetry mates suppressed the rotation function's first peak is the true orientation in 50 of 50 pose-gated cells, and the panel is 30/30 at either depth. Warm on one EPYC 9335 node: 1DAW 0.42 s, 2DQ6 0.67 s, 6G9X 0.67 s, 3K7M 1.02 s, 4BX9 1.11 s per alignment diff --git a/tests/unit/alignment/test_rotation_search_dtype_device.py b/tests/unit/alignment/test_rotation_search_dtype_device.py index 34145b51..2202d930 100644 --- a/tests/unit/alignment/test_rotation_search_dtype_device.py +++ b/tests/unit/alignment/test_rotation_search_dtype_device.py @@ -60,15 +60,16 @@ def test_expansion_follows_the_configured_float_dtype(float_dtype, want_complex) ) -def test_the_back_half_accumulates_one_step_wider(): - """The tail is deliberately wider than the expansion, and must stay so. - - Narrowing it looks like free memory -- complex128 buffers holding - complex64-accurate content -- and it is not. The radial sum and the Wigner - contraction are both oscillatory, so they cancel, and single-precision - accumulation was measured to move scores by 1e-4 to 1.4e-3 relative and - leave only 1 of 500 candidate slots holding the same orientation. The - expansion's own working precision is config's; this accumulation is not. +def test_the_back_half_runs_at_the_configured_complex_dtype(): + """The whole chain -- expansion, radial sum, Wigner contraction, FFT -- is + one dtype, the configured one, so a device without float64 runs it as is. + + The radial sum used to be accumulated one step wider. Measured, narrowing it + moved scores by 1e-4 relative and reordered the deep peak list while leaving + the top peak unchanged; the placement search now consumes only the top few + distinct orientations and single precision recovers every pose on the + benchmark panel, so the wider accumulator is gone and nothing downstream + re-decides the width. """ from torchref.experimental.alignment.frf.data_mr import ( bessel_sh_expand, cross_correlate_xi, @@ -85,14 +86,13 @@ def test_the_back_half_accumulates_one_step_wider(): try: dtypes.float, dtypes.complex = torch.float32, torch.complex64 c = bessel_sh_expand(s_vec, F, L=12, bessel_h_scale=40.0) + xi = cross_correlate_xi(c, c) finally: dtypes.float, dtypes.complex = original assert c.coeffs.dtype == torch.complex64, "expansion should be at config dtype" - xi = cross_correlate_xi(c, c) - assert xi.dtype == torch.complex128, ( - f"the radial accumulation narrowed to {xi.dtype}; see the docstring on " - f"cross_correlate_xi for what that costs" + assert xi.dtype == torch.complex64, ( + f"the radial accumulation left the configured dtype: {xi.dtype}" ) # Everything downstream follows xi rather than re-deciding. betas = torch.linspace(0.0, 3.0, 5, dtype=torch.float64) @@ -185,13 +185,24 @@ def test_double_is_a_device_capability_not_a_constant(): dtypes.float, dtypes.complex = original -def test_the_radial_accumulation_never_narrows_a_double_input(): - """A caller that already widened must not be silently narrowed back.""" +def test_the_radial_accumulation_follows_the_configuration_not_the_input(): + """Coefficients that arrive wider than the configured dtype are brought to it. + + The width of the accumulation is a property of the run's configuration -- on + a device without float64 it is the only width there is -- not of whatever a + caller happened to hand in. + """ from torchref.experimental.alignment.frf.data_mr import cross_correlate_xi from torchref.experimental.alignment.frf.types import BesselSHCoefficients c = BesselSHCoefficients( - coeffs=torch.zeros((2, 5, 9), dtype=torch.complex128), + coeffs=torch.zeros((2, 5, 9), dtype=torch.complex128), # dtype-ok: deliberately wider than the configuration, to see it brought back L=5, bessel_h_scale=20.0, ) - assert cross_correlate_xi(c, c).dtype == torch.complex128 + original = dtypes.float, dtypes.complex + try: + dtypes.float, dtypes.complex = torch.float32, torch.complex64 + out = cross_correlate_xi(c, c) + finally: + dtypes.float, dtypes.complex = original + assert out.dtype == torch.complex64 diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index cf8e80e0..5247343c 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -58,8 +58,7 @@ #: approximation and needs its own evidence. _GROUP_SCALE_COS = 10_000_000 -from ....config import (get_complex_dtype, get_float_dtype, - widest_complex_dtype, widest_float_dtype) +from ....config import get_complex_dtype, get_float_dtype from ....utils.backends import run_or_degrade, select from ..sh import legendre_recurrence_coefficients from ._backends import LEGENDRE_BACKENDS @@ -123,11 +122,13 @@ def spherical_bessel_table( """ real_dtype = x.dtype device = x.device - # Double where the device has it. The rescaling above is what makes a - # narrower working dtype survivable here at all -- without it the ladder - # overflows float32 for every x below ~35 -- but double is still preferable - # where it is available, since the recurrence runs ~90 steps. - work_dtype = widest_float_dtype(device) + # The ladder runs in the argument's own dtype. The rescaling above is what + # keeps the ~90-step downward recurrence in range -- without it the ladder + # overflows float32 for every x below ~35 -- and with it the single + # precision the expansion passes recovers every pose on the benchmark + # panel. A double argument gets a double ladder, which is what the + # bit-identity and scipy checks in the tests exercise. + work_dtype = x.dtype x64 = x.to(work_dtype) safe_x = x64.clamp(min=1e-30) inv_x = 1.0 / safe_x @@ -226,20 +227,19 @@ def bessel_sh_expand( Two precisions are in play and they are deliberately different. - The **clustering keys** are computed at ``s_vectors``' own dtype, because - ``_GROUP_SCALE_S`` keys ``|s|`` at 1e-7 and that is exactly where float32's - resolution runs out: at ``|s| = 0.5`` a float32 rounding is ~0.3 of a key - step, so reflections that are mathematically degenerate would sometimes land - in adjacent keys and the degeneracy collapse the cost model depends on would - fray. Callers therefore pass float64 ``s_vectors`` even when the rest of the - chain is single precision. + The **clustering keys** are computed on the host in double, whatever dtype + ``s_vectors`` arrive in: ``_GROUP_SCALE_S`` keys ``|s|`` at 1e-7 and that is + exactly where float32's resolution runs out -- at ``|s| = 0.5`` a float32 + rounding is ~0.3 of a key step, so reflections that are mathematically + degenerate would sometimes land in adjacent keys and the degeneracy + collapse the cost model depends on would fray. The host always has double, + the key computation is O(N), and nothing double ever touches the device. Everything else -- the Legendre/Y_lm precompute, the radial weights, the - contraction and the returned coefficients -- runs at - :func:`torchref.config.get_float_dtype`, which is this codebase's working - precision and the dtype the fused CPU kernel is built for. The - spherical-Bessel recurrence keeps its own float64 internals, where the - downward ladder needs the dynamic range. + Bessel ladder (kept in range by rescaling), the contraction and the returned + coefficients -- runs at :func:`torchref.config.get_float_dtype`, this + codebase's working precision and the dtype the fused CPU kernel is built + for. """ assert s_vectors.dim() == 2 and s_vectors.shape[-1] == 3 assert intensity.dim() == 1 and intensity.shape[0] == s_vectors.shape[0] @@ -315,10 +315,15 @@ def _tick(t0): phi_all = torch.atan2(s_vectors[..., 1], s_vectors[..., 0]) # Separate resolutions for the two factors: the radial term needs a fine # |s| key, the angular term does not. One shared key forces the finer of the - # two on both, which costs merges the angular part never needed. - k_s = (s_mag_all * _GROUP_SCALE_S).round().to(torch.int64) # dtype-ok: exact clustering key; needs double - k_c = (cos_all * _GROUP_SCALE_COS).round().to(torch.int64) + _GROUP_SCALE_COS # dtype-ok: exact clustering key; needs double - key = k_s * (2 * _GROUP_SCALE_COS + 1) + k_c + # two on both, which costs merges the angular part never needed. The keys + # are formed on the host in double -- see the docstring -- and only the + # integer keys come back. + s_key = s_vectors.detach().cpu().to(torch.float64) # dtype-ok: exact clustering key on the host; the device never sees it + s_mag_key = s_key.norm(dim=-1).clamp(min=1e-30) + cos_key = (s_key[..., 2] / s_mag_key).clamp(min=-1.0, max=1.0) + k_s = (s_mag_key * _GROUP_SCALE_S).round().to(torch.int64) # dtype-ok: exact clustering key + k_c = (cos_key * _GROUP_SCALE_COS).round().to(torch.int64) + _GROUP_SCALE_COS # dtype-ok: exact clustering key + key = (k_s * (2 * _GROUP_SCALE_COS + 1) + k_c).to(s_vectors.device) uniq_key, inverse = torch.unique(key, return_inverse=True) n_clusters = int(uniq_key.shape[0]) # Per-group geometry: the MEAN over the group's members, not an arbitrary @@ -558,12 +563,12 @@ def cross_correlate_xi( """ if c_obs.L != c_calc.L: raise ValueError(f"L mismatch: obs={c_obs.L} calc={c_calc.L}") - # One step wider than the coefficients where the device allows it. On a - # backend without float64 this falls back to the coefficients' own dtype and - # the run pays the accuracy noted above -- there is no third option there. - acc = widest_complex_dtype(c_obs.coeffs.device) - if c_obs.coeffs.dtype == torch.complex128: # dtype-ok: double accumulation of an oscillatory sum -- never narrow what the caller widened - acc = torch.complex128 # never narrow what the caller widened # dtype-ok: double accumulation of an oscillatory sum -- never narrow what the caller widened + # The configured complex dtype. The oscillatory radial sum used to be + # accumulated one step wider than the coefficients; measured, that moved + # scores by 1e-4 relative and reordered the deep peak list without moving + # the top peak, and the placement search now consumes only the top few + # distinct orientations. Single precision recovers every pose on the panel. + acc = get_complex_dtype() return torch.einsum( "rln,rlm->lmn", c_obs.coeffs.to(acc), diff --git a/torchref/experimental/alignment/frf/dense_calc.py b/torchref/experimental/alignment/frf/dense_calc.py index fd7949af..d95921c5 100644 --- a/torchref/experimental/alignment/frf/dense_calc.py +++ b/torchref/experimental/alignment/frf/dense_calc.py @@ -21,6 +21,8 @@ import torch +from torchref.config import get_float_dtype + if TYPE_CHECKING: from torchref.model import ModelFT @@ -51,9 +53,9 @@ def dense_calc_via_box( Returns ------- (s_vec, F_calc) : Tuple[torch.Tensor, torch.Tensor] - ``s_vec`` is the Cartesian reciprocal grid (N, 3) in double -- the - expansion clusters on it and needs exact keys -- and ``F_calc`` the - amplitudes (N,) in the model's dtype, both on the model's device. + ``s_vec`` is the Cartesian reciprocal grid (N, 3) and ``F_calc`` the + amplitudes (N,), in the configured float dtype on the model's device. + The expansion forms its exact clustering keys on the host itself. """ from torchref.symmetry.cell import Cell @@ -82,11 +84,12 @@ def dense_calc_via_box( [H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], dim=-1 ).to(torch.long) # dtype-ok: Miller indices are integers # Cubic box: |s| = |hkl| / a. - smag = hkl.to(torch.float64).norm(dim=-1) / a # dtype-ok: exact clustering key; needs double + real = get_float_dtype() + smag = hkl.to(real).norm(dim=-1) / a keep = (smag >= 1.0 / d_max) & (smag <= 1.0 / d_min) hkl = hkl[keep].contiguous() F = model_sf_abs(m, hkl) - s_vec = hkl.to(torch.float64) / a # dtype-ok: exact clustering key; needs double + s_vec = hkl.to(real) / a if verbose: print( diff --git a/torchref/experimental/alignment/frf/preprocessing.py b/torchref/experimental/alignment/frf/preprocessing.py index 89aac868..a158e089 100644 --- a/torchref/experimental/alignment/frf/preprocessing.py +++ b/torchref/experimental/alignment/frf/preprocessing.py @@ -17,6 +17,8 @@ import torch +from ....config import get_float_dtype + from ..sh import ( get_high_order_axis, # phaser's highOrderAxis() compute_patterson_shell_variance, @@ -107,7 +109,7 @@ def apply_shell_variance_weights( shell_idx = assign_shells(s_mag, edges) valid = shell_idx >= 0 var_p = compute_patterson_shell_variance( - intensity[valid].to(torch.float64), # dtype-ok: shell variance accumulated in double + intensity[valid], shell_idx[valid], P=n_var_shells, ) @@ -288,15 +290,16 @@ def fit_relative_wilson_b( if not valid_obs.any() or not valid_calc.any(): return 0.0 - F2_obs = (F_obs * F_obs).to(torch.float64) # dtype-ok: per-shell sums in double for the relative-B fit - F2_calc = (F_calc * F_calc).to(torch.float64) # dtype-ok: per-shell sums in double for the relative-B fit - s2_obs = (s_mag * s_mag).to(torch.float64) # dtype-ok: per-shell sums in double for the relative-B fit + real = get_float_dtype() + F2_obs = (F_obs * F_obs).to(real) + F2_calc = (F_calc * F_calc).to(real) + s2_obs = (s_mag * s_mag).to(real) counts_obs = torch.zeros(n_shells, dtype=torch.int64, device=s_mag.device) # dtype-ok: per-shell counts counts_calc = torch.zeros(n_shells, dtype=torch.int64, device=s_mag.device) # dtype-ok: per-shell counts - sum_F2obs = torch.zeros(n_shells, dtype=torch.float64, device=s_mag.device) # dtype-ok: per-shell sums in double for the relative-B fit - sum_F2calc = torch.zeros(n_shells, dtype=torch.float64, device=s_mag.device) # dtype-ok: per-shell sums in double for the relative-B fit - sum_s2 = torch.zeros(n_shells, dtype=torch.float64, device=s_mag.device) # dtype-ok: per-shell sums in double for the relative-B fit + sum_F2obs = torch.zeros(n_shells, dtype=real, device=s_mag.device) + sum_F2calc = torch.zeros(n_shells, dtype=real, device=s_mag.device) + sum_s2 = torch.zeros(n_shells, dtype=real, device=s_mag.device) idx_v_obs = shell_idx_obs[valid_obs] idx_v_calc = shell_idx_calc[valid_calc] counts_obs.index_add_(0, idx_v_obs, torch.ones_like(idx_v_obs)) @@ -307,9 +310,9 @@ def fit_relative_wilson_b( # Drop shells empty on either side. keep = (counts_obs > 0) & (counts_calc > 0) - mean_F2obs = sum_F2obs[keep] / counts_obs[keep].to(torch.float64) # dtype-ok: per-shell sums in double for the relative-B fit - mean_F2calc = sum_F2calc[keep] / counts_calc[keep].to(torch.float64) # dtype-ok: per-shell sums in double for the relative-B fit - mean_s2 = sum_s2[keep] / counts_obs[keep].to(torch.float64) # dtype-ok: per-shell sums in double for the relative-B fit + mean_F2obs = sum_F2obs[keep] / counts_obs[keep].to(real) + mean_F2calc = sum_F2calc[keep] / counts_calc[keep].to(real) + mean_s2 = sum_s2[keep] / counts_obs[keep].to(real) # log(Σ_N / Σ_P) per shell. eps = 1e-30 diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index 337b90b5..db28618c 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -28,7 +28,7 @@ import torch -from torchref.config import get_default_device +from torchref.config import get_default_device, get_float_dtype from torchref.scaling.weighting import (DEFAULT_SNR_CAP, DEFAULT_TRUST_CAP) from .sh import ( @@ -258,11 +258,12 @@ def prepare_frf_inputs( ``model`` happens to sit on. """ device = get_default_device() + real = get_float_dtype() - F_obs = data.F.to(torch.float64).abs() # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| + F_obs = data.F.to(real).abs() hkl_all = data.hkl - rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| - s_vec_all = hkl_all.to(torch.float64) @ rec_basis # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| + rec_basis = data.cell.reciprocal_basis_matrix.to(real) + s_vec_all = hkl_all.to(real) @ rec_basis s_mag_all = s_vec_all.norm(dim=-1) keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min) if keep.sum().item() < n_shells * 5: @@ -273,7 +274,7 @@ def prepare_frf_inputs( F_obs = F_obs[keep].to(device) sig_F = getattr(data, "F_sigma", None) if sig_F is not None: - sig_F = sig_F.to(torch.float64)[keep].to(device) # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| + sig_F = sig_F.to(real)[keep].to(device) hkl = hkl_all[keep].to(device) s_vec = s_vec_all[keep].to(device) s_mag = s_mag_all[keep].to(device) @@ -330,10 +331,11 @@ def search_peaks( # warning) and falls back to the configured default. Data first, matching # the rest of the codebase. device = resolve_device(data, model, device=device) + real = get_float_dtype() with torch.no_grad(): - rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64).to(device) # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| + rec_basis = data.cell.reciprocal_basis_matrix.to(real).to(device) hkl_all = data.hkl.to(device) - s_vec_all = hkl_all.to(torch.float64) @ rec_basis # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| + s_vec_all = hkl_all.to(real) @ rec_basis s_mag_all = s_vec_all.norm(dim=-1) d_min_data = float(1.0 / s_mag_all.max().item()) @@ -360,10 +362,10 @@ def search_peaks( s_asu = s_vec_all[keep] s_mag_asu = s_mag_all[keep] F_obs = apply_overall_anisotropy( - data.F.to(torch.float64).abs().to(device)[keep], s_asu, U_aniso, # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| + data.F.to(real).abs().to(device)[keep], s_asu, U_aniso, ) sigF = ( - data.F_sigma.to(torch.float64).to(device)[keep] # dtype-ok: rotation-function inputs kept in double; the expansion casts to the working dtype and clusters on |s| + data.F_sigma.to(real).to(device)[keep] if getattr(data, "F_sigma", None) is not None else None ) @@ -393,7 +395,7 @@ def search_peaks( # orthogonal symmetry matrices, so using S.h works everywhere except # trigonal and hexagonal, where it mixes non-equivalent reflections into # one orbit. - sg_mats = data.spacegroup.matrices.to(torch.float64).to(device) # dtype-ok: 3x3 rotation algebra in double on the host + sg_mats = data.spacegroup.matrices.to(real).to(device) n_ops = int(sg_mats.shape[0]) # `expand_reciprocal` is the package's one implementation of this # contraction, and it carries the h.S convention so no call site has to @@ -449,9 +451,13 @@ def search_peaks( # Point-group rotations in the Cartesian frame, so the peak finder can # treat an orientation and its symmetry mates as one peak. As a set # these equal B S B^-1; `hkl_symops_to_cartesian` returns the same - # rotations in a different order. + # rotations in a different order. Built on the host in double, where + # the peak finder's 3x3 algebra runs anyway. from .sh import hkl_symops_to_cartesian - sym_cart = hkl_symops_to_cartesian(sg_mats, rec_basis) + sym_cart = hkl_symops_to_cartesian( + data.spacegroup.matrices.detach().cpu().to(torch.float64), # dtype-ok: 3x3 rotation algebra in double on the host + data.cell.reciprocal_basis_matrix.detach().cpu().to(torch.float64), # dtype-ok: 3x3 rotation algebra in double on the host + ) engine = FastRotationFunction( s_obs, F_obs, centric, sg_mats, diff --git a/torchref/experimental/alignment/sh.py b/torchref/experimental/alignment/sh.py index 05630353..0f30670d 100644 --- a/torchref/experimental/alignment/sh.py +++ b/torchref/experimental/alignment/sh.py @@ -50,8 +50,9 @@ def legendre_recurrence_coefficients(L: int, dtype, device): runs this recurrence itself, fused with its own accumulation, and two copies of these formulae would be two chances to get them subtly different. """ - ll = torch.arange(L, dtype=torch.float64, device=device).view(L, 1) # dtype-ok: harmonic indices as exact doubles - mm = torch.arange(L, dtype=torch.float64, device=device).view(1, L) # dtype-ok: harmonic indices as exact doubles + # Small integers, exact in any float dtype; the results are cast to `dtype`. + ll = torch.arange(L, dtype=dtype, device=device).view(L, 1) + mm = torch.arange(L, dtype=dtype, device=device).view(1, L) valid = ll > mm denom = (ll - mm) * (ll + mm) denom_safe = torch.where(valid, denom, torch.ones_like(denom)) @@ -62,7 +63,7 @@ def legendre_recurrence_coefficients(L: int, dtype, device): b_num / torch.where(b_den == 0, torch.ones_like(b_den), b_den), min=0.0)) a = torch.where(valid, a, torch.zeros_like(a)).to(dtype) b = torch.where(valid, b, torch.zeros_like(b)).to(dtype) - m_arange = torch.arange(L, dtype=torch.float64, device=device) # dtype-ok: harmonic indices as exact doubles + m_arange = torch.arange(L, dtype=dtype, device=device) sect = torch.sqrt( (2.0 * m_arange + 1.0) / (2.0 * m_arange).clamp(min=1.0)).to(dtype) return a, b, sect @@ -241,7 +242,7 @@ def equal_count_shell_edges( s_sorted, _ = torch.sort(s) N = s_sorted.numel() # quantile-based partition - idx = torch.linspace(0, N - 1, P + 1, dtype=torch.float64, device=s.device).round().long() # dtype-ok: equal-count edge positions in double, then rounded + idx = torch.linspace(0, N - 1, P + 1, dtype=s.dtype, device=s.device).round().long() edges = s_sorted[idx] # nudge endpoints so the data is fully covered (avoid floating-point miss) if s_min is not None: From 2d796549313a5adaaa6df239b2bd6104afd150c7 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 2 Sep 2026 18:08:52 +0200 Subject: [PATCH 154/250] untracked alignment lab --- alignment_lab/README.md | 93 ---- alignment_lab/analysis/aggregate.py | 347 ------------ alignment_lab/analysis/arms_timing.sh | 60 --- alignment_lab/analysis/array_template.sh | 51 -- alignment_lab/analysis/benchmark_array.sh | 59 -- alignment_lab/analysis/build_grid_worth.sh | 57 -- alignment_lab/analysis/capability_gate.sh | 50 -- alignment_lab/analysis/compile_ab.sh | 56 -- alignment_lab/analysis/config_sweep_array.sh | 53 -- alignment_lab/analysis/copy_cost.sh | 49 -- alignment_lab/analysis/copyfix_timing.sh | 66 --- alignment_lab/analysis/dense_and_wigner.sh | 68 --- alignment_lab/analysis/dense_internals.sh | 78 --- alignment_lab/analysis/double_audit.py | 102 ---- alignment_lab/analysis/double_audit.sh | 19 - .../analysis/empirical_sigma_a_check.sh | 19 - alignment_lab/analysis/eps_gate.sh | 22 - alignment_lab/analysis/first_true_rank.sh | 27 - alignment_lab/analysis/frf_fingerprint.py | 40 -- alignment_lab/analysis/frf_orbit_side.sh | 19 - alignment_lab/analysis/friedel_gate.sh | 27 - alignment_lab/analysis/ftf_early_stop.py | 117 ---- alignment_lab/analysis/ftf_running_null.py | 140 ----- alignment_lab/analysis/full_gate.sh | 26 - alignment_lab/analysis/gpfs_cold_read.sh | 35 -- .../analysis/integration_alignment.sh | 17 - alignment_lab/analysis/kernel_ab.sh | 118 ---- alignment_lab/analysis/kernel_check.sh | 62 --- alignment_lab/analysis/marginal_seeds.sh | 30 -- alignment_lab/analysis/merge_gate.sh | 53 -- alignment_lab/analysis/merge_identity.sh | 31 -- .../analysis/merge_numeric_identity.py | 86 --- alignment_lab/analysis/p1_grid_coherence.sh | 19 - alignment_lab/analysis/panel_arms.py | 121 ----- alignment_lab/analysis/panel_arms.sh | 22 - alignment_lab/analysis/panel_ranks.py | 80 --- alignment_lab/analysis/pipeline_timing.sh | 19 - alignment_lab/analysis/pipeline_timing_gpu.sh | 20 - .../analysis/pipeline_timing_ncand10.sh | 19 - alignment_lab/analysis/pose_arms.sh | 35 -- alignment_lab/analysis/pose_panel_ncand10.sh | 28 - alignment_lab/analysis/pose_panel_trans.sh | 28 - alignment_lab/analysis/rebaseline_panel.sh | 25 - alignment_lab/analysis/refactor_gate.sh | 21 - alignment_lab/analysis/repeat_stability.sh | 47 -- alignment_lab/analysis/run_tests.sh | 26 - alignment_lab/analysis/run_tests_copyfix.sh | 35 -- alignment_lab/analysis/scaler_cost.sh | 17 - alignment_lab/analysis/scaling_gate.sh | 18 - alignment_lab/analysis/soln_tables.sh | 20 - alignment_lab/analysis/stage_profile.sh | 31 -- alignment_lab/analysis/stagea_fingerprint.sh | 53 -- alignment_lab/analysis/stagea_tests.sh | 32 -- alignment_lab/analysis/stagea_verify.sh | 146 ----- alignment_lab/analysis/stageb_gate.sh | 49 -- alignment_lab/analysis/stagec_gate3.sh | 55 -- .../analysis/stagec_timing_paired.sh | 34 -- alignment_lab/analysis/staged_gate.sh | 76 --- alignment_lab/analysis/staged_panel.sh | 29 - alignment_lab/analysis/sweep_chunk.sh | 54 -- .../analysis/tf_resolution_and_grid.sh | 19 - alignment_lab/analysis/tf_sf_backend.sh | 19 - .../analysis/trigonal_metric_recheck.sh | 36 -- alignment_lab/analysis/truth_metric_check.sh | 17 - alignment_lab/analysis/truth_pose_scores.sh | 22 - alignment_lab/analysis/truth_pose_scores2.sh | 25 - .../analysis/verify_copy_correctness.sh | 80 --- alignment_lab/analysis/verify_copy_fix.sh | 79 --- alignment_lab/analysis/verify_frf.sh | 43 -- alignment_lab/analysis/verify_wigner_memo.sh | 71 --- alignment_lab/analysis/weight_arms.py | 93 ---- alignment_lab/analysis/weight_arms.sh | 25 - alignment_lab/analysis/where_now.sh | 20 - alignment_lab/analysis/where_now_cap64.sh | 24 - alignment_lab/analysis/wigner_cache_probe.sh | 88 --- alignment_lab/analysis/wigner_call_count.sh | 47 -- alignment_lab/analysis/wigner_fft_proto.sh | 73 --- alignment_lab/analysis/wigner_mirror_proto.sh | 79 --- alignment_lab/analysis/wigner_split.sh | 44 -- alignment_lab/analysis/wilson_moments.sh | 17 - alignment_lab/analysis/wilson_smoke.py | 92 ---- alignment_lab/analysis/wilson_smoke.sh | 17 - .../diagnostics/empirical_sigma_a_check.py | 50 -- .../diagnostics/frf_aniso_knockout.py | 165 ------ .../diagnostics/frf_aniso_rank_sweep.py | 152 ------ alignment_lab/diagnostics/frf_benchmark.py | 277 ---------- alignment_lab/diagnostics/frf_config_sweep.py | 224 -------- .../diagnostics/frf_encode_compare.py | 509 ------------------ .../diagnostics/frf_ghost_knockout.py | 198 ------- .../diagnostics/frf_inject_phaser_obs.py | 211 -------- alignment_lab/diagnostics/frf_map_compare.py | 428 --------------- .../diagnostics/frf_normaliser_anatomy.py | 256 --------- alignment_lab/diagnostics/frf_orbit_side.py | 67 --- alignment_lab/diagnostics/frf_prep_compare.py | 285 ---------- alignment_lab/diagnostics/frf_rank.py | 88 --- alignment_lab/diagnostics/ghost_origin.py | 124 ----- .../diagnostics/p1_grid_coherence.py | 71 --- .../diagnostics/phaser_headtohead.py | 112 ---- alignment_lab/diagnostics/pipeline_timing.py | 66 --- alignment_lab/diagnostics/pose_recovery.py | 239 -------- alignment_lab/diagnostics/scaler_cost.py | 76 --- .../diagnostics/tf_resolution_and_grid.py | 99 ---- alignment_lab/diagnostics/tf_sf_backend.py | 112 ---- .../diagnostics/truth_metric_check.py | 73 --- .../diagnostics/truth_pose_scores.py | 102 ---- alignment_lab/diagnostics/wilson_moments.py | 59 -- alignment_lab/lab/__init__.py | 69 --- alignment_lab/lab/aniso.py | 177 ------ alignment_lab/lab/benchmark.py | 129 ----- alignment_lab/lab/frf.py | 257 --------- alignment_lab/lab/phaser.py | 204 ------- alignment_lab/lab/phaser_match.py | 441 --------------- alignment_lab/lab/profile.py | 281 ---------- alignment_lab/lab/results.py | 161 ------ alignment_lab/lab/truth.py | 289 ---------- alignment_lab/tests/test_lab.py | 120 ----- 116 files changed, 10437 deletions(-) delete mode 100644 alignment_lab/README.md delete mode 100644 alignment_lab/analysis/aggregate.py delete mode 100644 alignment_lab/analysis/arms_timing.sh delete mode 100644 alignment_lab/analysis/array_template.sh delete mode 100644 alignment_lab/analysis/benchmark_array.sh delete mode 100644 alignment_lab/analysis/build_grid_worth.sh delete mode 100644 alignment_lab/analysis/capability_gate.sh delete mode 100644 alignment_lab/analysis/compile_ab.sh delete mode 100644 alignment_lab/analysis/config_sweep_array.sh delete mode 100644 alignment_lab/analysis/copy_cost.sh delete mode 100644 alignment_lab/analysis/copyfix_timing.sh delete mode 100644 alignment_lab/analysis/dense_and_wigner.sh delete mode 100644 alignment_lab/analysis/dense_internals.sh delete mode 100644 alignment_lab/analysis/double_audit.py delete mode 100644 alignment_lab/analysis/double_audit.sh delete mode 100644 alignment_lab/analysis/empirical_sigma_a_check.sh delete mode 100644 alignment_lab/analysis/eps_gate.sh delete mode 100644 alignment_lab/analysis/first_true_rank.sh delete mode 100644 alignment_lab/analysis/frf_fingerprint.py delete mode 100644 alignment_lab/analysis/frf_orbit_side.sh delete mode 100644 alignment_lab/analysis/friedel_gate.sh delete mode 100644 alignment_lab/analysis/ftf_early_stop.py delete mode 100644 alignment_lab/analysis/ftf_running_null.py delete mode 100644 alignment_lab/analysis/full_gate.sh delete mode 100644 alignment_lab/analysis/gpfs_cold_read.sh delete mode 100644 alignment_lab/analysis/integration_alignment.sh delete mode 100644 alignment_lab/analysis/kernel_ab.sh delete mode 100644 alignment_lab/analysis/kernel_check.sh delete mode 100644 alignment_lab/analysis/marginal_seeds.sh delete mode 100644 alignment_lab/analysis/merge_gate.sh delete mode 100644 alignment_lab/analysis/merge_identity.sh delete mode 100644 alignment_lab/analysis/merge_numeric_identity.py delete mode 100644 alignment_lab/analysis/p1_grid_coherence.sh delete mode 100644 alignment_lab/analysis/panel_arms.py delete mode 100644 alignment_lab/analysis/panel_arms.sh delete mode 100644 alignment_lab/analysis/panel_ranks.py delete mode 100644 alignment_lab/analysis/pipeline_timing.sh delete mode 100644 alignment_lab/analysis/pipeline_timing_gpu.sh delete mode 100644 alignment_lab/analysis/pipeline_timing_ncand10.sh delete mode 100644 alignment_lab/analysis/pose_arms.sh delete mode 100644 alignment_lab/analysis/pose_panel_ncand10.sh delete mode 100644 alignment_lab/analysis/pose_panel_trans.sh delete mode 100644 alignment_lab/analysis/rebaseline_panel.sh delete mode 100644 alignment_lab/analysis/refactor_gate.sh delete mode 100644 alignment_lab/analysis/repeat_stability.sh delete mode 100644 alignment_lab/analysis/run_tests.sh delete mode 100644 alignment_lab/analysis/run_tests_copyfix.sh delete mode 100644 alignment_lab/analysis/scaler_cost.sh delete mode 100644 alignment_lab/analysis/scaling_gate.sh delete mode 100644 alignment_lab/analysis/soln_tables.sh delete mode 100644 alignment_lab/analysis/stage_profile.sh delete mode 100644 alignment_lab/analysis/stagea_fingerprint.sh delete mode 100644 alignment_lab/analysis/stagea_tests.sh delete mode 100644 alignment_lab/analysis/stagea_verify.sh delete mode 100644 alignment_lab/analysis/stageb_gate.sh delete mode 100644 alignment_lab/analysis/stagec_gate3.sh delete mode 100644 alignment_lab/analysis/stagec_timing_paired.sh delete mode 100644 alignment_lab/analysis/staged_gate.sh delete mode 100644 alignment_lab/analysis/staged_panel.sh delete mode 100644 alignment_lab/analysis/sweep_chunk.sh delete mode 100644 alignment_lab/analysis/tf_resolution_and_grid.sh delete mode 100644 alignment_lab/analysis/tf_sf_backend.sh delete mode 100644 alignment_lab/analysis/trigonal_metric_recheck.sh delete mode 100644 alignment_lab/analysis/truth_metric_check.sh delete mode 100644 alignment_lab/analysis/truth_pose_scores.sh delete mode 100644 alignment_lab/analysis/truth_pose_scores2.sh delete mode 100644 alignment_lab/analysis/verify_copy_correctness.sh delete mode 100644 alignment_lab/analysis/verify_copy_fix.sh delete mode 100644 alignment_lab/analysis/verify_frf.sh delete mode 100644 alignment_lab/analysis/verify_wigner_memo.sh delete mode 100644 alignment_lab/analysis/weight_arms.py delete mode 100644 alignment_lab/analysis/weight_arms.sh delete mode 100644 alignment_lab/analysis/where_now.sh delete mode 100644 alignment_lab/analysis/where_now_cap64.sh delete mode 100644 alignment_lab/analysis/wigner_cache_probe.sh delete mode 100644 alignment_lab/analysis/wigner_call_count.sh delete mode 100644 alignment_lab/analysis/wigner_fft_proto.sh delete mode 100644 alignment_lab/analysis/wigner_mirror_proto.sh delete mode 100644 alignment_lab/analysis/wigner_split.sh delete mode 100644 alignment_lab/analysis/wilson_moments.sh delete mode 100644 alignment_lab/analysis/wilson_smoke.py delete mode 100644 alignment_lab/analysis/wilson_smoke.sh delete mode 100644 alignment_lab/diagnostics/empirical_sigma_a_check.py delete mode 100644 alignment_lab/diagnostics/frf_aniso_knockout.py delete mode 100644 alignment_lab/diagnostics/frf_aniso_rank_sweep.py delete mode 100644 alignment_lab/diagnostics/frf_benchmark.py delete mode 100644 alignment_lab/diagnostics/frf_config_sweep.py delete mode 100644 alignment_lab/diagnostics/frf_encode_compare.py delete mode 100644 alignment_lab/diagnostics/frf_ghost_knockout.py delete mode 100644 alignment_lab/diagnostics/frf_inject_phaser_obs.py delete mode 100644 alignment_lab/diagnostics/frf_map_compare.py delete mode 100644 alignment_lab/diagnostics/frf_normaliser_anatomy.py delete mode 100644 alignment_lab/diagnostics/frf_orbit_side.py delete mode 100644 alignment_lab/diagnostics/frf_prep_compare.py delete mode 100644 alignment_lab/diagnostics/frf_rank.py delete mode 100644 alignment_lab/diagnostics/ghost_origin.py delete mode 100644 alignment_lab/diagnostics/p1_grid_coherence.py delete mode 100644 alignment_lab/diagnostics/phaser_headtohead.py delete mode 100644 alignment_lab/diagnostics/pipeline_timing.py delete mode 100644 alignment_lab/diagnostics/pose_recovery.py delete mode 100644 alignment_lab/diagnostics/scaler_cost.py delete mode 100644 alignment_lab/diagnostics/tf_resolution_and_grid.py delete mode 100644 alignment_lab/diagnostics/tf_sf_backend.py delete mode 100644 alignment_lab/diagnostics/truth_metric_check.py delete mode 100644 alignment_lab/diagnostics/truth_pose_scores.py delete mode 100644 alignment_lab/diagnostics/wilson_moments.py delete mode 100644 alignment_lab/lab/__init__.py delete mode 100644 alignment_lab/lab/aniso.py delete mode 100644 alignment_lab/lab/benchmark.py delete mode 100644 alignment_lab/lab/frf.py delete mode 100644 alignment_lab/lab/phaser.py delete mode 100644 alignment_lab/lab/phaser_match.py delete mode 100644 alignment_lab/lab/profile.py delete mode 100644 alignment_lab/lab/results.py delete mode 100644 alignment_lab/lab/truth.py delete mode 100644 alignment_lab/tests/test_lab.py diff --git a/alignment_lab/README.md b/alignment_lab/README.md deleted file mode 100644 index f977f89e..00000000 --- a/alignment_lab/README.md +++ /dev/null @@ -1,93 +0,0 @@ -# alignment_lab - -Harness for the FRF rotation-function work: shared primitives, one file per -experiment, one aggregator. - -``` -lab/ shared library — import from here, do not re-derive -diagnostics/ one experiment per file, each with --out-csv -analysis/ aggregator + SLURM array template -tests/ self-tests for the primitives -runs/ CSVs and Phaser working dirs (gitignored) -slurm/ scheduler logs (gitignored) -``` - -## Why a library - -The scripts this replaces carried ~37 copies of the rotation generator, ~30 of -the benchmark list, ~28 of the CSV writer and ~20 of the rank-of-truth -computation — and several disagreed with each other. Two of those divergences -changed results silently: - -- **Two rotation generators.** One omitted the `sign(diag(R))` QR correction, so - the same seed produced a *different* rotation. Results from the two families - were never comparable. `lab.truth.random_rotation` is the corrected form and - `tests/test_lab.py` pins it against the other variant. -- **Four orbit conventions.** Rank-of-truth was computed with the symmetry - operators applied on either side, in either the fractional or Cartesian frame. - The choice changes the rank, so `orbit_rank` takes it as an explicit argument - and every result row records it. - -## Running - -```bash -PY=.dev/bin/python # or another worktree's interpreter -PYTHONPATH=. $PY alignment_lab/diagnostics/ghost_origin.py --pdb 3K7M --trial 0 -PYTHONPATH=. $PY -m pytest alignment_lab/tests -q - -sbatch --array=0-29 --partition=hour --time=00:55:00 --cpus-per-task=4 \ - --mem=32G alignment_lab/analysis/array_template.sh ghost_origin -PYTHONPATH=. $PY alignment_lab/analysis/aggregate.py \ - 'alignment_lab/runs/ghost_origin_*/*.csv' --compare obs_mode -``` - -## Reading a result - -Seed-to-seed truth-rank spread at `lmax_cap=64` is **±4–6** (1AK5 has been seen -at 9, 11 and 17 for one configuration). Below ~10 trials nothing is -interpretable; three findings that looked strong at n≤7 evaporated at full n. -`aggregate.py` therefore reports paired per-trial differences with the -per-trial values visible, never a bare median, and prints whatever it dropped. - -## Traps worth knowing - -- `model.spacegroup = SpaceGroup("P 1")` is a **silent no-op** — `SpaceGroup` is - an `nn.Module`, so `nn.Module.__setattr__` intercepts the assignment and the - property setter never runs. Assign the **name string**. `ghost_origin.py`'s P1 - arm depends on this. -- `Model.rotate` / `.translate` mutate in place and return `self`. Copy first if - you still need the original — `rotated_case` does. -- Phaser **exits 0 on fatal input errors**; an empty peak list is the real - signal. Its keyword file needs absolute paths. -- `PEAKS ROT SELECT ALL` returns ~80–92k densely spaced samples (median nearest - neighbour under 1°), so "the closest sample is within a degree" means nothing - by itself. - -## The benchmark - -`diagnostics/frf_benchmark.py` is the standing benchmark: accuracy, memory and -runtime in one row per (structure, trial, arm), so a change cannot buy one at the -silent expense of another. `analysis/benchmark_array.sh` runs it one structure -per **exclusive** node. - -Three things it is careful about, each of which has bitten this harness before: - -- **Instrumentation points.** `frf/api.py` binds `bessel_sh_expand` and its - neighbours into its own namespace at import, so wrapping them in the module - that *defines* them intercepts nothing and the stage reports zero calls — - indistinguishable from a free stage. `lab/profile.FRF_STAGES` names the module - where each call is **resolved**. Getting this wrong left 85% of the runtime - unattributed. -- **Nested stages.** `evaluate_rotation_function` contains - `build_dense_map_per_beta`, which contains `wigner_contraction_per_beta`; and - `bessel_sh_expand` contains `spherical_bessel_table`. The report gives - exclusive time alongside inclusive, so the column sums. -- **Wall clock on a shared cluster measures the cluster.** Every row carries a - fixed calibration workload timed in the same process plus the host identity. - Compare `seconds_per_calibration` across nodes, or raw seconds only within - one. - -Peak memory comes from an RSS sampler, so a spike shorter than the sampling -interval is invisible, and glibc may not return freed pages — which makes a later -window in the same process look cheaper than it is. `vm_hwm_mb` is the -process-lifetime high-water mark for absolute numbers. diff --git a/alignment_lab/analysis/aggregate.py b/alignment_lab/analysis/aggregate.py deleted file mode 100644 index de2459c4..00000000 --- a/alignment_lab/analysis/aggregate.py +++ /dev/null @@ -1,347 +0,0 @@ -"""Aggregate lab result CSVs. - -One aggregator, because every row shares the core schema. It reports **paired, -per-trial** differences rather than a bare median of each arm: on this benchmark -the seed-to-seed truth-rank spread at ``lmax_cap=64`` is +-4-6 (1AK5 has been -seen at 9, 11 and 17 for the same configuration), so a difference of medians -over a handful of trials is noise. Three findings that looked strong at n<=7 -vanished at full n. - -Anything dropped is printed. A silent truncation reads as full coverage. - -Usage:: - - python alignment_lab/analysis/aggregate.py 'alignment_lab/runs/*.csv' - python alignment_lab/analysis/aggregate.py 'runs/*.csv' --compare obs_mode -""" - -from __future__ import annotations - -import argparse -import csv -import glob -import statistics -from collections import defaultdict -from typing import Dict, List - - -def load(patterns: List[str]) -> List[dict]: - """Read every CSV matching the patterns into a list of row dicts.""" - rows: List[dict] = [] - files = sorted({f for p in patterns for f in glob.glob(p)}) - if not files: - raise SystemExit(f"no CSVs matched {patterns}") - for f in files: - with open(f, newline="") as fh: - rows.extend(csv.DictReader(fh)) - print(f"# {len(rows)} rows from {len(files)} file(s)") - return rows - - -def _rank(row: dict) -> float: - """Truth rank as a number; a miss (-1) sorts as worst, not as best.""" - try: - r = int(row["truth_rank"]) - except (KeyError, ValueError): - return float("nan") - return float("inf") if r < 0 else float(r) - - -def summarise(rows: List[dict]) -> None: - """Per-structure rank summary, with misses counted separately.""" - by_pdb: Dict[str, List[float]] = defaultdict(list) - for r in rows: - by_pdb[r.get("pdb", "?")].append(_rank(r)) - print(f"\n{'pdb':8s} {'n':>3s} {'median':>7s} {'min':>5s} {'max':>5s} " - f"{'misses':>7s} per-trial ranks") - for pdb in sorted(by_pdb): - vals = by_pdb[pdb] - finite = [v for v in vals if v != float("inf")] - misses = sum(1 for v in vals if v == float("inf")) - med = statistics.median(finite) if finite else float("nan") - lo = min(finite) if finite else float("nan") - hi = max(finite) if finite else float("nan") - shown = ", ".join("miss" if v == float("inf") else f"{int(v)}" for v in vals) - print(f"{pdb:8s} {len(vals):3d} {med:7.1f} {lo:5.0f} {hi:5.0f} " - f"{misses:7d} [{shown}]") - if any(v == float("inf") for vs in by_pdb.values() for v in vs): - print("# 'miss' = truth not found in the peak list; excluded from median/min/max") - - -def compare(rows: List[dict], key: str, base: str = None) -> None: - """Paired per-(pdb, seed) comparison across the arms of ``key``.""" - arms = sorted({r.get(key, "") for r in rows}) - if len(arms) < 2: - print(f"\n# only one arm for {key!r}; nothing to pair") - return - cells: Dict[tuple, Dict[str, float]] = defaultdict(dict) - for r in rows: - cells[(r.get("pdb"), r.get("seed"))][r.get(key, "")] = _rank(r) - - if base is not None and base not in arms: - raise SystemExit(f"--base {base!r} not among {key} values {arms}") - base = base if base is not None else arms[0] - arms = [a for a in arms if a != base] - print(f"\n# paired against {key}={base!r}; + means rank got worse") - for arm in arms: - deltas, unpaired = [], 0 - for (pdb, seed), by_arm in sorted(cells.items()): - a, b = by_arm.get(base), by_arm.get(arm) - if a is None or b is None: - unpaired += 1 - continue - if a == float("inf") or b == float("inf"): - unpaired += 1 # a miss has no meaningful numeric difference - continue - deltas.append(b - a) - if not deltas: - print(f" {arm:>16s}: no comparable pairs ({unpaired} unpaired)") - continue - better = sum(1 for d in deltas if d < 0) - worse = sum(1 for d in deltas if d > 0) - same = sum(1 for d in deltas if d == 0) - print(f" {arm:>16s}: median delta {statistics.median(deltas):+.1f} " - f"(better {better} / worse {worse} / unchanged {same}, n={len(deltas)})" - + (f" [{unpaired} pair(s) dropped: missing arm or a miss]" if unpaired else "")) - if len(deltas) < 10: - print(f" {'':16s} n={len(deltas)} is below the ~10 trials this " - f"benchmark needs; treat as indicative only") - - -def _cmp_rank(row: dict, n_peaks_default: int = 500) -> float: - """Rank for pairing: a miss counts as worse than the worst hit. - - ``compare`` drops pairs where either side missed, which silently removes - exactly the cases an arm is being blamed for. The sweep writes - ``rank_for_compare`` for this; fall back to the peak-list length. - """ - try: - return float(int(row["rank_for_compare"])) - except (KeyError, ValueError): - pass - try: - r = int(row["truth_rank"]) - except (KeyError, ValueError): - return float("nan") - if r >= 0: - return float(r) - try: - return float(int(row["n_peaks"])) - except (KeyError, ValueError): - return float(n_peaks_default) - - -def gate(rows: List[dict], key: str = "arm", base: str = "production", - top_n: int = 20, min_hits: int = 9) -> None: - """Report each arm against the shipping criterion, per structure. - - The criterion is not "truth at rank 0": the pipeline carries the top ~20 - candidates forward, so rank 7 and rank 0 are the same outcome and rank 223 - is not. An arm passes when truth lands in the top ``top_n`` on at least - ``min_hits`` of the trials for **every** structure -- so one bad structure - cannot be averaged away by nine good ones. - """ - arms = sorted({r.get(key, "") for r in rows}) - pdbs = sorted({r.get("pdb", "?") for r in rows}) - per: Dict[tuple, List[dict]] = defaultdict(list) - for r in rows: - per[(r.get(key, ""), r.get("pdb", "?"))].append(r) - - def _in_top(r: dict) -> bool: - try: - v = int(r["truth_rank"]) - except (KeyError, ValueError): - return False - return 0 <= v < top_n - - # A structure appearing with more trials than the others means rows were - # collected twice -- a re-run after a partial failure, say -- and the - # per-structure hit counts are then not comparable. That has to be loud: it - # silently flips which arms pass. - counts = {} - for arm in arms: - for pdb in pdbs: - n = len(per.get((arm, pdb), [])) - if n: - counts.setdefault(n, []).append(f"{arm}/{pdb}") - if len(counts) > 1: - detail = ", ".join( - f"{n} trials: {len(v)} cell(s) e.g. {v[0]}" - for n, v in sorted(counts.items())) - raise SystemExit( - f"inconsistent trial counts across cells ({detail}). Deduplicate the " - f"inputs -- comparing 10 trials of one structure against 20 of " - f"another makes the gate meaningless." - ) - - print(f"\n# shipping gate: truth in the top {top_n} on >= {min_hits} trials, " - f"for every structure") - print(f"{key:<26} {'pass':>5} {'worst structure':>16} {'total':>7} " - f"{'rank0':>6} {'median':>7}") - verdicts = {} - for arm in arms: - worst_pdb, worst_hits, tot, hits, rank0, ranks = None, None, 0, 0, 0, [] - for pdb in pdbs: - rs = per.get((arm, pdb), []) - if not rs: - continue - h = sum(1 for r in rs if _in_top(r)) - if worst_hits is None or h < worst_hits: - worst_hits, worst_pdb = h, f"{pdb} {h}/{len(rs)}" - tot += len(rs); hits += h - rank0 += sum(1 for r in rs if str(r.get("truth_rank")) == "0") - ranks += [_cmp_rank(r) for r in rs] - ok = worst_hits is not None and worst_hits >= min_hits - verdicts[arm] = ok - med = statistics.median(ranks) if ranks else float("nan") - print(f"{arm:<26} {'PASS' if ok else 'fail':>5} {worst_pdb or '-':>16} " - f"{hits}/{tot:<5} {rank0:>6} {med:>7.1f}") - - # Paired against the shipped configuration, misses included as worst. - cells: Dict[tuple, Dict[str, float]] = defaultdict(dict) - for r in rows: - cells[(r.get("pdb"), r.get("trial"))][r.get(key, "")] = _cmp_rank(r) - if base not in arms: - print(f"\n# no {key}={base!r} rows; skipping the paired report") - return - print(f"\n# paired against {key}={base!r} over every (structure, trial) cell; " - f"+ means worse") - for arm in arms: - if arm == base: - continue - d = [by[arm] - by[base] for by in cells.values() - if arm in by and base in by] - if not d: - print(f" {arm:<26} no comparable cells") - continue - print(f" {arm:<26} n={len(d):<4} better={sum(x < 0 for x in d):<4} " - f"same={sum(x == 0 for x in d):<4} worse={sum(x > 0 for x in d):<4} " - f"median={statistics.median(d):+8.1f}") - dup = f"{base}_dup" - if dup in arms: - d = [by[dup] - by[base] for by in cells.values() - if dup in by and base in by] - moved = sum(1 for x in d if x != 0) - print(f"\n# control: {dup} repeats {base} verbatim. {moved}/{len(d)} cells " - f"differ -- that is the engine's own spread, and no effect smaller " - f"than it is resolvable here.") - - -def bench(rows: List[dict], key: str = "arm") -> None: - """Accuracy, memory and runtime side by side, per structure and per arm. - - All three together on purpose: a bandwidth that halves the runtime while - dropping the true orientation out of the carried window is not a win, and - neither is one that finds it using memory the machine does not have. - - Runtime is reported both raw and divided by the calibration workload each row - carries. If the rows span several hosts the raw column is not comparable and - the header says so. - """ - def num(r, k, default=float("nan")): - try: - return float(r[k]) - except (KeyError, TypeError, ValueError): - return default - - hosts = sorted({r.get("host", "?") for r in rows}) - models = sorted({r.get("cpu_model", "") for r in rows}) - threads = sorted({r.get("torch_threads", "?") for r in rows}) - print(f"\n# {len(rows)} rows, {len(hosts)} host(s), threads {threads}") - print(f"# cpu: {', '.join(m or 'unknown' for m in models)}") - # What breaks comparability is a different CPU, not a different hostname: - # several nodes of one pinned model are interchangeable, and warning about - # them trains the reader to ignore the warning that matters. - if len(models) > 1: - print("# rows span several CPU models: compare s/cal, not seconds") - elif len(hosts) > 1: - print(f"# {len(hosts)} nodes, all {models[0] or 'unknown'} -- seconds " - f"are comparable") - kinds = sorted({r.get("timing_kind", "?") for r in rows}) - if len(kinds) > 1: - print(f"# WARNING: mixed timing kinds {kinds} -- cold and steady-state " - f"numbers are not comparable") - - arms = sorted({r.get(key, "") for r in rows}) - pdbs = sorted({r.get("pdb", "?") for r in rows}) - cells = defaultdict(list) - for r in rows: - cells[(r.get(key, ""), r.get("pdb", "?"))].append(r) - - for arm in arms: - mine = [r for r in rows if r.get(key) == arm] - if not mine: - continue - print(f"\n## {key}={arm}") - print(f" {'pdb':7s} {'sg':11s} {'n':>3s} {'top-N':>6s} {'med rank':>9s} " - f"{'med s':>8s} {'s/cal':>8s} {'peak MB':>9s} {'delta MB':>9s}") - for pdb in pdbs: - rs = cells.get((arm, pdb), []) - if not rs: - continue - hits = sum(1 for r in rs if str(r.get("in_top_n")) == "1") - print(f" {pdb:7s} {rs[0].get('spacegroup', '?'):11s} {len(rs):3d} " - f"{hits:>3d}/{len(rs):<2d} " - f"{statistics.median(_cmp_rank(r) for r in rs):9.1f} " - f"{statistics.median(num(r, 'seconds') for r in rs):8.2f} " - f"{statistics.median(num(r, 'seconds_per_calibration') for r in rs):8.1f} " - f"{statistics.median(num(r, 'rss_peak_mb') for r in rs):9.0f} " - f"{statistics.median(num(r, 'rss_delta_mb') for r in rs):9.0f}") - worst = min( - (sum(1 for r in cells[(arm, p)] if str(r.get("in_top_n")) == "1") - / max(len(cells[(arm, p)]), 1), p) - for p in pdbs if cells.get((arm, p))) - print(f" worst structure: {worst[1]} at {100 * worst[0]:.0f}% in the " - f"carried window | peak memory across structures " - f"{max(num(r, 'rss_peak_mb', 0) for r in mine):.0f} MB | " - f"slowest {max(num(r, 'seconds', 0) for r in mine):.1f} s") - - # Where the time goes, if the stage columns are present. - xcols = sorted({k for r in rows for k in r if k.startswith("x_")}) - if xcols: - print("\n## exclusive stage time, median over every row (seconds)") - meds = sorted(((statistics.median(num(r, c, 0.0) for r in rows), c) - for c in xcols), reverse=True) - tot = statistics.median(num(r, "seconds") for r in rows) - for m, c in meds: - if m <= 0: - continue - print(f" {c[2:]:32s} {m:8.3f} {100 * m / max(tot, 1e-9):6.1f}%") - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("patterns", nargs="+", help="CSV glob(s)") - ap.add_argument("--base", default=None, - help="which value of --compare is the control arm " - "(default: first alphabetically)") - ap.add_argument("--compare", default=None, - help="column whose values are the arms to pair on, " - "e.g. obs_mode or lmax_cap") - ap.add_argument("--bench", action="store_true", - help="accuracy, memory and runtime side by side " - "(frf_benchmark rows)") - ap.add_argument("--gate", action="store_true", - help="report each arm against the shipping criterion " - "(truth in the top N on most trials, every structure)") - ap.add_argument("--top-n", type=int, default=20, - help="how many candidates the downstream pipeline carries") - ap.add_argument("--min-hits", type=int, default=9, - help="trials per structure that must land in the top N") - args = ap.parse_args() - rows = load(args.patterns) - if args.bench: - bench(rows, key=args.compare or "arm") - return 0 - if args.gate: - gate(rows, key=args.compare or "arm", base=args.base or "production", - top_n=args.top_n, min_hits=args.min_hits) - return 0 - summarise(rows) - if args.compare: - compare(rows, args.compare, args.base) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/analysis/arms_timing.sh b/alignment_lab/analysis/arms_timing.sh deleted file mode 100644 index b233594f..00000000 --- a/alignment_lab/analysis/arms_timing.sh +++ /dev/null @@ -1,60 +0,0 @@ -#!/bin/bash -# What the two knockouts actually SAVE. The panel run could not answer this: it -# ran production first in every trial, so production alone paid the process-level -# warm-up -- the fused C++ kernel build, the SO(3) sample list, the Wigner block -# memo -- and the later arms inherited all three warm. That is why it reported a -# ~10x "speedup" from deleting a 20-bin regression, which is not credible. -# -# Here: one throwaway run to warm every memo, then rounds with the arm order -# ROTATED so no arm is systematically first. -#SBATCH --job-name=arms_time -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --exclusive -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" -"$PY" -u - <<'PYEOF' -import statistics, sys, time -from pathlib import Path -sys.path.insert(0, str(Path("alignment_lab").resolve())) -import torch -torch.set_grad_enabled(False) -sys.path.insert(0, str(Path("alignment_lab/analysis").resolve())) -from lab import FRFConfig, rotated_case, run_frf, seed_for -from panel_arms import knocked_out - -ARMS = ["production", "no_brel", "no_friedel"] -cfg = FRFConfig(n_peaks=500, lmax_cap=64) - -for pdb in ("3K7M", "1DAW"): - model, data, _ = rotated_case(pdb, seed_for(pdb, 0)) - run_frf(model, data, cfg, capture_arf=False, verbose=0) # warm every memo - t = {a: [] for a in ARMS} - for r in range(4): - order = ARMS[r % len(ARMS):] + ARMS[:r % len(ARMS)] # rotate - for arm in order: - model, data, _ = rotated_case(pdb, seed_for(pdb, r)) - t0 = time.time() - with knocked_out(arm): - run_frf(model, data, cfg, capture_arf=False, verbose=0) - t[arm].append(time.time() - t0) - base = statistics.median(t["production"]) - print(f"--- {pdb} (warm, arm order rotated, 4 rounds) ---") - for arm in ARMS: - med = statistics.median(t[arm]) - pd = [t[arm][i] - t["production"][i] for i in range(len(t[arm]))] - print(f" {arm:12s} median {med:.3f}s paired d vs production " - f"median {statistics.median(pd):+.3f}s " - f"({100*statistics.median(pd)/base:+.1f}%) raw {[round(v,3) for v in t[arm]]}") -PYEOF -echo "rc=$?" diff --git a/alignment_lab/analysis/array_template.sh b/alignment_lab/analysis/array_template.sh deleted file mode 100644 index a2925eb6..00000000 --- a/alignment_lab/analysis/array_template.sh +++ /dev/null @@ -1,51 +0,0 @@ -#!/bin/bash -# SLURM array template for the alignment lab. -# -# Resources go on the sbatch command line, not in this file, so one template -# serves CPU diagnostics and GPU sweeps. The array index selects a (pdb, trial) -# cell from the worklist below. -# -# sbatch --array=0-29 --partition=hour --time=00:55:00 --cpus-per-task=4 \ -# --mem=32G alignment_lab/analysis/array_template.sh ghost_origin -# -# 10 structures x 3 trials = 30 tasks. Note the +-4-6 seed-to-seed rank spread: -# 3 trials is for a smoke run, ~10 for anything you intend to believe. -#SBATCH --job-name=align_lab -#SBATCH --output=alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=alignment_lab/slurm/%x_%A_%a.err -set -euo pipefail - -DIAG="${1:?usage: array_template.sh [extra args...]}" -shift || true - -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY="$REPO/.dev/bin/python" -[ -x "$PY" ] || PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python - -cd "$REPO" -export PYTHONPATH="$REPO" -export TORCHREF_NUM_THREADS="${SLURM_CPUS_PER_TASK:-4}" -export OMP_NUM_THREADS="$TORCHREF_NUM_THREADS" -export MKL_NUM_THREADS="$TORCHREF_NUM_THREADS" -export PYTHONUNBUFFERED=1 -[ -z "${SLURM_JOB_GPUS:-}" ] && export CUDA_VISIBLE_DEVICES="" - -# Worklist: keep in step with lab.benchmark.BENCH_PDBS (order is a seed contract). -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -TRIALS=3 -IDX="${SLURM_ARRAY_TASK_ID:-0}" -PDB="${PDBS[$((IDX / TRIALS))]}" -TRIAL=$((IDX % TRIALS)) - -OUTDIR="alignment_lab/runs/${DIAG}_${SLURM_ARRAY_JOB_ID:-local}" -mkdir -p "$OUTDIR" alignment_lab/slurm - -echo "task $IDX -> $PDB trial $TRIAL -> $OUTDIR" -rc=0 -"$PY" -u "alignment_lab/diagnostics/${DIAG}.py" \ - --pdb "$PDB" --trial "$TRIAL" \ - --out-csv "$OUTDIR/${DIAG}_${PDB}_t${TRIAL}.csv" "$@" || rc=$? - -# Report the real exit status: a task that dies must not be logged COMPLETED. -echo "exit_code=$rc" -exit "$rc" diff --git a/alignment_lab/analysis/benchmark_array.sh b/alignment_lab/analysis/benchmark_array.sh deleted file mode 100644 index 3fce278c..00000000 --- a/alignment_lab/analysis/benchmark_array.sh +++ /dev/null @@ -1,59 +0,0 @@ -#!/bin/bash -# The rotation search's standing benchmark: accuracy, memory and runtime. -# -# One array task per structure, all arms and trials in one process on an -# EXCLUSIVE node, so the arm comparison is within-node and the memory peaks are -# not another job's. Runtime on a shared node measures the node; every row also -# carries a calibration workload and the host identity so that is checkable -# rather than assumed. -# -# sbatch --array=0-9 --partition=hour --time=00:55:00 --exclusive --mem=0 \ -# --constraint=cpu_epyc9335 alignment_lab/analysis/benchmark_array.sh \ -# --arms cap48,cap64,cap100 --trials 3 -# -# `--mem=0` takes the node's memory: cap100 on the P432 structures needs well -# over 32 GB, which is what the OOMs in job 489988 were. -# -# **Pin the CPU model.** `--exclusive` stops neighbours interfering but does not -# stop SLURM handing out whatever generation is free -- this cluster mixes Xeon -# 6152/6230/6230r/6248r/6530 with EPYC 7452/7453/9334/9335, and two runs on -# different generations are not comparable at all. `cpu_epyc9335` is the newest -# available (28 nodes on `hour`, 64 cores). Any before/after pair has to name the -# same constraint, and the CPU model is recorded in every row so a mismatch is -# visible after the fact. -#SBATCH --job-name=frf_bench -#SBATCH --output=alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=alignment_lab/slurm/%x_%A_%a.err -set -uo pipefail - -REPO="${FRF_BENCH_REPO:-/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement}" -PY="$REPO/.dev/bin/python" -[ -x "$PY" ] || PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python - -cd "$REPO" -export PYTHONPATH="$REPO" -export TORCHREF_NUM_THREADS="${FRF_BENCH_THREADS:-4}" -export OMP_NUM_THREADS="$TORCHREF_NUM_THREADS" -export MKL_NUM_THREADS="$TORCHREF_NUM_THREADS" -export PYTHONUNBUFFERED=1 -export CUDA_VISIBLE_DEVICES="" - -# Pin the thread count rather than inheriting it from the allocation: an -# exclusive node hands over every core, so SLURM_CPUS_PER_TASK would make the -# timings depend on the node's size instead of on the code. -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -IDX="${SLURM_ARRAY_TASK_ID:-0}" -PDB="${PDBS[$IDX]}" - -OUTDIR="alignment_lab/runs/frf_bench_${SLURM_ARRAY_JOB_ID:-local}" -mkdir -p "$OUTDIR" alignment_lab/slurm - -echo "task $IDX -> $PDB on $(hostname), ${TORCHREF_NUM_THREADS} threads" -echo "repo=$REPO sha=$(git -C "$REPO" rev-parse --short HEAD 2>/dev/null || echo unknown)" -rc=0 -"$PY" -u -m alignment_lab.diagnostics.frf_benchmark \ - --pdb "$PDB" --out-csv "$OUTDIR/${PDB}.csv" "$@" || rc=$? - -# `rc=$?` has to follow the command directly, or a task that dies reads COMPLETED. -echo "exit_code=$rc" -exit "$rc" diff --git a/alignment_lab/analysis/build_grid_worth.sh b/alignment_lab/analysis/build_grid_worth.sh deleted file mode 100644 index 203dbc54..00000000 --- a/alignment_lab/analysis/build_grid_worth.sh +++ /dev/null @@ -1,57 +0,0 @@ -#!/bin/bash -# Is ModelFT.copy(build_grid=False) still worth anything after the merge? -# -# It existed because setup_grid also built MapSymmetry -- one (nx,ny,nz,3) -# sampling array per symmetry operation -- which cost 957.8 ms of a 2.06 s -# search. dev moved that out: setup_grid now asks `spacegroup.can_index_directly` -# instead of building an operator. What remains is compute_real_space_grid, an -# (nx,ny,nz,3) coordinate array, so the question is whether that alone is worth a -# parameter. -#SBATCH --job-name=gridworth -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:40:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --exclusive -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname)" -"$PY" -u - <<'PYEOF' 2>&1 | grep -E "^GRID|Error|Traceback" -import statistics, sys, time -from pathlib import Path -sys.path.insert(0, "alignment_lab") -import torch -torch.set_grad_enabled(False) -from lab import load_case - -for pdb in ("1DAW", "3K7M"): - model, data = load_case(pdb)[:2] - model.cell = data.cell - model.spacegroup = str(data.spacegroup.hm) - model.max_res = 2.0 - model.setup_grid(max_res=2.0) - g = model._fft.real_space_grid - print(f"GRID {pdb}: grid {'None' if g is None else tuple(g.shape)} " - f"n_ops={data.spacegroup.n_ops}") - model.copy(); model.copy(build_grid=False) # warm - t = {True: [], False: []} - for r in range(6): - for bg in ((True, False) if r % 2 == 0 else (False, True)): - t0 = time.perf_counter(); c = model.copy(build_grid=bg) - t[bg].append(time.perf_counter() - t0) - del c - med_t, med_f = statistics.median(t[True]), statistics.median(t[False]) - pd = sorted(t[True][i] - t[False][i] for i in range(len(t[True]))) - print(f"GRID {pdb}: copy(build_grid=True) median {1000*med_t:8.2f} ms") - print(f"GRID {pdb}: copy(build_grid=False) median {1000*med_f:8.2f} ms") - print(f"GRID {pdb}: paired median saving {1000*statistics.median(pd):8.2f} ms " - f"({100*statistics.median(pd)/med_t:+.1f}% of the copy)") -PYEOF -echo "rc=$?" diff --git a/alignment_lab/analysis/capability_gate.sh b/alignment_lab/analysis/capability_gate.sh deleted file mode 100644 index bcf5b22c..00000000 --- a/alignment_lab/analysis/capability_gate.sh +++ /dev/null @@ -1,50 +0,0 @@ -#!/bin/bash -# On a device WITH float64 the capability helpers must resolve to exactly what -# was hardcoded before, so this has to be bit-identical. That is the whole gate: -# the MPS branch cannot be exercised here and is not claimed to work. -#SBATCH --job-name=capgate -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:50:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -OUT=$NEW/alignment_lab/slurm -cd "$NEW"; export PYTHONPATH="$NEW" -export TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname)" -LOG=$OUT/capgate_tests_$SLURM_JOB_ID.log -"$PY" -m pytest tests/unit/alignment tests/unit/frf_separate tests/unit/model \ - tests/unit/test_imports_smoke.py -q > "$LOG" 2>&1 -rc=$?; echo "=== TESTS rc=$rc ==="; tail -12 "$LOG" -echo "=== capability helpers resolve as expected on this host ===" -"$PY" - <<'PYEOF' -import torch -from torchref.config import (supports_double, widest_complex_dtype, - widest_float_dtype) -print(f" cpu: supports_double={supports_double('cpu')} " - f"float={widest_float_dtype('cpu')} complex={widest_complex_dtype('cpu')}") -assert widest_float_dtype("cpu") is torch.float64 -assert widest_complex_dtype("cpu") is torch.complex128 -print(" mps branch (not exercisable here, resolution only):") -from torchref.config import _NO_DOUBLE_DEVICE_TYPES -print(f" device types without float64: {_NO_DOUBLE_DEVICE_TYPES}") -PYEOF -echo "=== fingerprint vs the committed state (baseline worktree at fe53d373) ===" -OLD=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/_stagea_baseline -for pdb in 3K7M 1DAW; do - for tree in OLD NEW; do - eval "root=\$$tree"; cd "$root" - PYTHONPATH="$root" "$PY" -u "$NEW/alignment_lab/analysis/frf_fingerprint.py" \ - --pdb "$pdb" --lmax-cap 64 > "$OUT/cap_${tree}_${pdb}.txt" 2>/dev/null - done - diff -q <(grep '^FP ' "$OUT/cap_OLD_${pdb}.txt") <(grep '^FP ' "$OUT/cap_NEW_${pdb}.txt") >/dev/null \ - && echo "IDENTICAL $pdb ($(grep -c '^FP ' "$OUT/cap_NEW_${pdb}.txt") peaks)" \ - || { echo "DIFFERS $pdb"; diff <(grep '^FP ' "$OUT/cap_OLD_${pdb}.txt") <(grep '^FP ' "$OUT/cap_NEW_${pdb}.txt") | head -4; } -done -echo "capgate_rc=$rc" diff --git a/alignment_lab/analysis/compile_ab.sh b/alignment_lab/analysis/compile_ab.sh deleted file mode 100644 index 78ca8f4f..00000000 --- a/alignment_lab/analysis/compile_ab.sh +++ /dev/null @@ -1,56 +0,0 @@ -#!/bin/bash -# A/B the compiled Legendre step against the eager one. Interleaved repeats, so -# any drift on the node hits both arms alike, and the truth rank is reported -# beside each timing -- a faster build that changes the answer is not faster. -# -# sbatch --partition=hour --time=00:55:00 --exclusive --mem=0 \ -# --constraint=cpu_epyc9335 alignment_lab/analysis/compile_ab.sh -#SBATCH --job-name=frf_compile_ab -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 -export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" -"$PY" -u -c " -import sys, time -sys.path.insert(0,'alignment_lab') -import torch; torch.set_grad_enabled(False) -import torchref.experimental.alignment.frf.data_mr as dm -from lab import FRFConfig, orbit_rank, rotated_case, run_frf, seed_for - -cases = {p: rotated_case(p, seed_for(p, 0)) for p in ('1DAW', '3K7M')} -for cap in (64, 100): - cfg = FRFConfig(lmax_cap=cap, n_peaks=200) - # Warm up BOTH arms: the compiled one pays its build on first call, and - # charging that to the measurement would answer a different question. - for compiled in (False, True): - dm.COMPILE_LEGENDRE_STEP = compiled - for m, d, _ in cases.values(): - run_frf(m, d, cfg, capture_arf=False) - res = {a: {p: [] for p in cases} for a in ('eager', 'compiled')} - rank = {} - for rep in range(3): - for arm, compiled in (('eager', False), ('compiled', True)): - dm.COMPILE_LEGENDRE_STEP = compiled - for p, (m, d, R) in cases.items(): - t0 = time.perf_counter() - r = run_frf(m, d, cfg, capture_arf=False) - res[arm][p].append(time.perf_counter() - t0) - k, _ = orbit_rank(r.peaks, R, - d.spacegroup.matrices.to(torch.float64).cpu(), - reciprocal_basis=d.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), - side='left', frame='cart') - rank[(arm, p)] = k - print(f'--- cap{cap} (best of 3, seconds) ---') - for p in cases: - e, c = min(res['eager'][p]), min(res['compiled'][p]) - print(f' {p:6s} eager {e:7.2f} [rank {rank[(\"eager\",p)]:>3}] ' - f'compiled {c:7.2f} [rank {rank[(\"compiled\",p)]:>3}] ' - f'speedup {e/max(c,1e-9):5.2f}x') -" -echo "exit_code=$?" diff --git a/alignment_lab/analysis/config_sweep_array.sh b/alignment_lab/analysis/config_sweep_array.sh deleted file mode 100644 index 8803b6cc..00000000 --- a/alignment_lab/analysis/config_sweep_array.sh +++ /dev/null @@ -1,53 +0,0 @@ -#!/bin/bash -# Part 1 of the FRF cleanup: settle lmax_cap / anisotropy / orbit-unroll / -# Patterson-radius by measurement before the switches are deleted. -# -# One array task = one (structure, trial) cell, running every arm in the same -# process so the paired comparison against `production` is exact. -# -# sbatch --array=0-99 --partition=hour --time=00:55:00 --cpus-per-task=4 \ -# --mem=32G alignment_lab/analysis/config_sweep_array.sh 1 -# -# 10 structures x 10 trials = 100 tasks. Pass the stage (1 or 2) as $1; any -# further arguments go through to the diagnostic. -#SBATCH --job-name=frf_cfg_sweep -#SBATCH --output=alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=alignment_lab/slurm/%x_%A_%a.err -set -uo pipefail - -STAGE="${1:?usage: config_sweep_array.sh [extra args...]}" -shift || true - -REPO="${FRF_SWEEP_REPO:-/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement}" -PY="$REPO/.dev/bin/python" -[ -x "$PY" ] || PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python - -cd "$REPO" -export PYTHONPATH="$REPO" -export TORCHREF_NUM_THREADS="${SLURM_CPUS_PER_TASK:-4}" -export OMP_NUM_THREADS="$TORCHREF_NUM_THREADS" -export MKL_NUM_THREADS="$TORCHREF_NUM_THREADS" -export PYTHONUNBUFFERED=1 -export CUDA_VISIBLE_DEVICES="" - -# Worklist: keep in step with lab.benchmark.BENCH_PDBS (order is a seed contract). -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -TRIALS=10 -IDX="${SLURM_ARRAY_TASK_ID:-0}" -PDB="${PDBS[$((IDX / TRIALS))]}" -TRIAL=$((IDX % TRIALS)) - -OUTDIR="alignment_lab/runs/config_sweep_s${STAGE}_${SLURM_ARRAY_JOB_ID:-local}" -mkdir -p "$OUTDIR" alignment_lab/slurm - -echo "task $IDX -> $PDB trial $TRIAL stage $STAGE -> $OUTDIR" -echo "repo=$REPO sha=$(git -C "$REPO" rev-parse --short HEAD 2>/dev/null || echo unknown)" -rc=0 -"$PY" -u alignment_lab/diagnostics/frf_config_sweep.py \ - --pdb "$PDB" --trial "$TRIAL" --stage "$STAGE" \ - --out-csv "$OUTDIR/${PDB}_t${TRIAL}.csv" "$@" || rc=$? - -# Report the real exit status: `rc=$?` must follow the command directly, or a -# task that dies gets logged COMPLETED. -echo "exit_code=$rc" -exit "$rc" diff --git a/alignment_lab/analysis/copy_cost.sh b/alignment_lab/analysis/copy_cost.sh deleted file mode 100644 index 3f106175..00000000 --- a/alignment_lab/analysis/copy_cost.sh +++ /dev/null @@ -1,49 +0,0 @@ -#!/bin/bash -# model.copy() is ~97% of dense_calc_via_box, which is ~50% of a shipped -# rotation search. What inside it costs the second, and is it steady state or -# first-call lazy initialisation? -#SBATCH --job-name=frf_copy -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname)" -"$PY" -u -c " -import copy as copy_module, time, torch -torch.set_grad_enabled(False) -from alignment_lab.lab.benchmark import load_case - -def reps(fn, n=5): - fn() - return [ (lambda t0: (fn(), time.perf_counter()-t0)[1])(time.perf_counter()) for _ in range(n) ] - -for name in ('3K7M', '1DAW'): - model, data = load_case(name) - model.verbose = 0 - g = model._fft.real_space_grid if model._fft is not None else None - print(f'--- {name}: {len(model.pdb)} atoms cell={[round(x,1) for x in model.cell.parameters_list()] if hasattr(model.cell,\"parameters_list\") else \"?\"} ' - f'max_res={model.max_res} grid={None if g is None else tuple(g.shape)}') - ts = reps(lambda: model.copy()) - print(f' model.copy() x5: ' + ' '.join(f'{t*1e3:.0f}' for t in ts) + ' ms') - - # The pieces, each timed on its own. - print(f' pdb.copy(deep=True) {min(reps(lambda: model.pdb.copy(deep=True)))*1e3:8.1f} ms') - if model._parametrization is not None: - print(f' deepcopy(_parametrization) {min(reps(lambda: copy_module.deepcopy(model._parametrization)))*1e3:8.1f} ms') - if model._fft is not None: - print(f' _fft.copy() {min(reps(lambda: model._fft.copy()))*1e3:8.1f} ms') - mc = model.copy() - print(f' setup_grid(max_res) {min(reps(lambda: mc.setup_grid(max_res=model.max_res)))*1e3:8.1f} ms') - print(f' _rebuild_sf_indices() {min(reps(lambda: model._rebuild_sf_indices()))*1e3:8.1f} ms') - print(f' cell.clone() {min(reps(lambda: model.cell.clone()))*1e3:8.1f} ms') - - # What the FRF actually needs the copy for: it mutates max_res, spacegroup, cell. - # Would a copy taken AFTER shrinking max_res be cheaper? - print(f' copy() with max_res=4.15 {min(reps(lambda: (lambda m: m)(model.copy())))*1e3:8.1f} ms (baseline)') -" 2>&1 | grep -vE "Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$" -echo "exit_code=$?" diff --git a/alignment_lab/analysis/copyfix_timing.sh b/alignment_lab/analysis/copyfix_timing.sh deleted file mode 100644 index 953f8c66..00000000 --- a/alignment_lab/analysis/copyfix_timing.sh +++ /dev/null @@ -1,66 +0,0 @@ -#!/bin/bash -# Payoff of build_grid=False. NOTE: whatever node this lands on, the absolute -# ms are NOT comparable with the EPYC 9335 tables -- only the before/after -# ratio measured inside this one job is. -#SBATCH --job-name=frf_ctime -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -FILT="Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$|No CUDA" -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" -echo "=== the copy itself, and the whole dense-calc stage ===" -"$PY" -u -c " -import math, time, torch -torch.set_grad_enabled(False) -from alignment_lab.lab.benchmark import load_case -from torchref.symmetry.cell import Cell -from torchref.experimental.alignment.frf.dense_calc import dense_calc_via_box, model_sf_abs -from torchref.experimental.alignment.frf.api import phaser_lmax_resolution - -def box_path(model, d_min, d_max, build_grid): - m = model.copy(build_grid=build_grid) - coords = m.xyz(); dev = coords.device - a = float(4.0 * (coords - coords.mean(0)).norm(dim=-1).max().item()) - m.max_res = float(d_min); m.spacegroup = 'P 1' - m.cell = Cell([a, a, a, 90., 90., 90.], device=dev) - nmax = int(math.ceil(a / d_min)) - idx = torch.arange(-nmax, nmax + 1, device=dev) - H, K, Lg = torch.meshgrid(idx, idx, idx, indexing='ij') - hkl = torch.stack([H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], -1).to(torch.long) - smag = hkl.to(torch.float64).norm(dim=-1) / a - hkl = hkl[(smag >= 1.0/d_max) & (smag <= 1.0/d_min)].contiguous() - return model_sf_abs(m, hkl) - -def best(fn, n=3): - fn() - return min((lambda: (lambda t0: (fn(), time.perf_counter()-t0)[1])(time.perf_counter()))() for _ in range(n)) - -for name in ('3K7M', '1DAW'): - model, data = load_case(name); model.verbose = 0 - rb = data.cell.reciprocal_basis_matrix.to(torch.float64) - d_min_data = 1.0 / (data.hkl.to(torch.float64) @ rb).norm(dim=-1).max().item() - r = float((model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item()) - L, d_min = phaser_lmax_resolution(r, d_min_data, 64) - tc_on = best(lambda: model.copy(build_grid=True)) - tc_off = best(lambda: model.copy(build_grid=False)) - t_on = best(lambda: box_path(model, d_min, 100.0, True)) - t_off = best(lambda: box_path(model, d_min, 100.0, False)) - print(f' {name} (cap64, d_min={d_min:.2f}A)') - print(f' model.copy() grid {tc_on*1e3:8.1f} -> no grid {tc_off*1e3:7.1f} ms ({tc_on/tc_off:5.1f}x)') - print(f' whole box path grid {t_on*1e3:8.1f} -> no grid {t_off*1e3:7.1f} ms ({t_on/t_off:5.1f}x)') - print(f' dense_calc_via_box now {best(lambda: dense_calc_via_box(model, 100.0, d_min, pad=2.0))*1e3:8.1f} ms') -" 2>&1 | grep -vE "$FILT" -echo "=== whole rotation search, per cap ===" -for pdb in 3K7M 1DAW; do - for arm in cap64 cap100; do - "$PY" -u -m alignment_lab.diagnostics.frf_benchmark \ - --pdb "$pdb" --arms "$arm" --trials 2 2>&1 | grep -vE "$FILT" - done -done -echo "done" diff --git a/alignment_lab/analysis/dense_and_wigner.sh b/alignment_lab/analysis/dense_and_wigner.sh deleted file mode 100644 index 6fc8e84b..00000000 --- a/alignment_lab/analysis/dense_and_wigner.sh +++ /dev/null @@ -1,68 +0,0 @@ -#!/bin/bash -# Per-cap stage breakdown (the shipped config is cap64; earlier tables medianed -# cap64 and cap100 together), plus a breakdown INSIDE dense_calc_via_box. -#SBATCH --job-name=frf_dense -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname)" -FILT="Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$" -for pdb in 3K7M 1DAW; do - for arm in cap64 cap100; do - "$PY" -u -m alignment_lab.diagnostics.frf_benchmark \ - --pdb "$pdb" --arms "$arm" --trials 2 2>&1 | grep -vE "$FILT" - done -done -echo "=== inside dense_calc_via_box ===" -"$PY" -u -c " -import math, time, torch -torch.set_grad_enabled(False) -from alignment_lab.lab.benchmark import BENCHMARK, load_case -from torchref.symmetry.cell import Cell -from torchref.experimental.alignment.frf.dense_calc import model_sf_abs -from torchref.experimental.alignment.frf.sitelist_ang import phaser_lmax_resolution - -for name in ('3K7M', '1DAW'): - case = load_case(name) - model, data = case.model, case.data - d_min_data = 1.0 / data.hkl_to_s(data.hkl).norm(dim=-1).max().item() - d_max = 100.0 - r = float((model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item()) - for cap in (64, 100): - L, d_min = phaser_lmax_resolution(r, d_min_data, cap) - t = {} - t0 = time.perf_counter() - m = model.copy() - coords = m.xyz(); dev = coords.device - extent = (coords - coords.mean(0)).norm(dim=-1).max().item() - a = float(2.0 * 2.0 * extent) - t['copy'] = time.perf_counter() - t0 - t0 = time.perf_counter() - m.max_res = float(d_min); m.spacegroup = 'P 1' - m.cell = Cell([a, a, a, 90., 90., 90.], device=dev) - t['cell+grid setup'] = time.perf_counter() - t0 - t0 = time.perf_counter() - nmax = int(math.ceil(a / d_min)) - idx = torch.arange(-nmax, nmax + 1, device=dev) - H, K, Lg = torch.meshgrid(idx, idx, idx, indexing='ij') - hkl = torch.stack([H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], -1).to(torch.long) - smag = hkl.to(torch.float64).norm(dim=-1) / a - hkl = hkl[(smag >= 1.0/d_max) & (smag <= 1.0/d_min)].contiguous() - t['hkl enumerate'] = time.perf_counter() - t0 - model_sf_abs(m, hkl) # warm - t0 = time.perf_counter(); model_sf_abs(m, hkl); t['model_sf_abs'] = time.perf_counter()-t0 - grid = m.cell.compute_grid_size(m.max_res) - print(f'{name} cap{cap}: d_min={d_min:.2f}A box={a:.0f}A grid={grid} ' - f'spacing={a/grid[0]:.2f}A n_hkl={hkl.shape[0]} ' - f'(enumerated {(2*nmax+1)**3})') - tot = sum(t.values()) - for k, v in t.items(): - print(f' {k:18s} {v*1e3:8.1f} ms {100*v/tot:5.1f}%') -" 2>&1 | grep -vE "$FILT" -echo "exit_code=$?" diff --git a/alignment_lab/analysis/dense_internals.sh b/alignment_lab/analysis/dense_internals.sh deleted file mode 100644 index 2c361f78..00000000 --- a/alignment_lab/analysis/dense_internals.sh +++ /dev/null @@ -1,78 +0,0 @@ -#!/bin/bash -# dense_calc_via_box is ~50% of a shipped (cap64) rotation search. Where inside -# it does the time go, and does it actually track the FFT grid fineness? If the -# cost is insensitive to the grid, a coarser grid buys nothing and the -# artificial-B route is pointless. -#SBATCH --job-name=frf_dint -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname)" -"$PY" -u -c " -import math, time, torch -torch.set_grad_enabled(False) -from alignment_lab.lab.benchmark import load_case -from torchref.symmetry.cell import Cell -from torchref.experimental.alignment.frf.dense_calc import model_sf_abs -from torchref.experimental.alignment.frf.api import phaser_lmax_resolution - -def timeit(fn, n=3): - fn() - out = [] - for _ in range(n): - t0 = time.perf_counter(); fn(); out.append(time.perf_counter() - t0) - return min(out) - -for name in ('3K7M', '1DAW'): - model, data = load_case(name) - rb = data.cell.reciprocal_basis_matrix.to(torch.float64) - d_min_data = 1.0 / (data.hkl.to(torch.float64) @ rb).norm(dim=-1).max().item() - r = float((model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item()) - for cap in (64,): - L, d_min = phaser_lmax_resolution(r, d_min_data, cap) - t = {} - t0 = time.perf_counter() - m = model.copy() - t['model.copy()'] = time.perf_counter() - t0 - coords = m.xyz(); dev = coords.device - extent = (coords - coords.mean(0)).norm(dim=-1).max().item() - a = float(2.0 * 2.0 * extent) - t0 = time.perf_counter() - m.max_res = float(d_min); m.spacegroup = 'P 1' - m.cell = Cell([a, a, a, 90., 90., 90.], device=dev) - t['cell/grid setup'] = time.perf_counter() - t0 - t0 = time.perf_counter() - nmax = int(math.ceil(a / d_min)) - idx = torch.arange(-nmax, nmax + 1, device=dev) - H, K, Lg = torch.meshgrid(idx, idx, idx, indexing='ij') - hkl = torch.stack([H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], -1).to(torch.long) - smag = hkl.to(torch.float64).norm(dim=-1) / a - hkl = hkl[(smag >= 1.0/100.0) & (smag <= 1.0/d_min)].contiguous() - t['hkl enumerate'] = time.perf_counter() - t0 - t['model_sf_abs (1st)'] = timeit(lambda: model_sf_abs(m, hkl), n=1) - t['model_sf_abs (warm)'] = timeit(lambda: model_sf_abs(m, hkl)) - g = m.cell.compute_grid_size(m.max_res) - print(f'--- {name} cap{cap}: d_min={d_min:.2f}A box={a:.0f}A ' - f'grid={g[0]}^3 spacing={a/g[0]:.2f}A n_hkl={hkl.shape[0]} ' - f'of {(2*nmax+1)**3} enumerated') - for k, v in t.items(): - print(f' {k:22s} {v*1e3:8.1f} ms') - - # Does the cost track the grid at all? Same reflections, different grid. - print(' grid sensitivity (same hkl list, max_res only):') - for fac in (0.5, 1.0, 1.5, 2.0): - mm = model.copy() - mm.max_res = float(d_min / fac) # fac>1 => finer grid - mm.spacegroup = 'P 1' - mm.cell = Cell([a, a, a, 90., 90., 90.], device=dev) - gg = mm.cell.compute_grid_size(mm.max_res) - dt = timeit(lambda: model_sf_abs(mm, hkl)) - print(f' x{fac:<4} grid={gg[0]:4d}^3 spacing={a/gg[0]:5.2f}A {dt*1e3:8.1f} ms') -" 2>&1 | grep -vE "Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$" -echo "exit_code=$?" diff --git a/alignment_lab/analysis/double_audit.py b/alignment_lab/analysis/double_audit.py deleted file mode 100644 index 0dc480ff..00000000 --- a/alignment_lab/analysis/double_audit.py +++ /dev/null @@ -1,102 +0,0 @@ -"""Where does the rotation search create float64 / complex128 tensors? - -MPS has no float64 at all, so every double-precision tensor on the compute -device is a portability blocker. Rather than guess at the list -- twice now a -confident guess about this engine has been wrong -- intercept every torch call -and report the source line that produced each double tensor, with how many and -how large. - -Deliberately reports rather than asserts. Some of these are *correct* and must -stay: the spherical-Bessel ladder needs the exponent range, the J_y -eigendecomposition and the anisotropy fit are precision-critical, and anything -already on the host costs nothing. The point is to separate those from the -per-reflection arrays that are double by inheritance. -""" - -from __future__ import annotations - -import argparse -import collections -import sys -import traceback -from pathlib import Path - -import torch -from torch.overrides import TorchFunctionMode - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -_DOUBLE = (torch.float64, torch.complex128) - - -class DoubleAudit(TorchFunctionMode): - """Attribute every double-precision tensor to the line that made it.""" - - def __init__(self, package_only: str = "torchref"): - super().__init__() - self.package_only = package_only - self.sites = collections.Counter() - self.elems = collections.Counter() - self.devices = collections.defaultdict(set) - - def __torch_function__(self, func, types, args=(), kwargs=None): - out = func(*args, **(kwargs or {})) - try: - tensors = [] - if isinstance(out, torch.Tensor): - tensors = [out] - elif isinstance(out, (tuple, list)): - tensors = [t for t in out if isinstance(t, torch.Tensor)] - if any(t.dtype in _DOUBLE for t in tensors): - # Innermost frame inside the package under audit, so the report - # names our code and not torch's internals. - site = None - for fr in reversed(traceback.extract_stack()[:-1]): - if f"/{self.package_only}/" in fr.filename: - short = fr.filename.split(f"/{self.package_only}/", 1)[1] - site = f"{self.package_only}/{short}:{fr.lineno}" - break - if site is not None: - n = sum(t.numel() for t in tensors if t.dtype in _DOUBLE) - self.sites[site] += 1 - self.elems[site] += n - for t in tensors: - if t.dtype in _DOUBLE: - self.devices[site].add(str(t.device)) - except Exception: # pragma: no cover - diagnostic - pass - return out - - def report(self, top: int = 30) -> None: - print(f"{'site':62s} {'calls':>7s} {'elements':>12s} devices") - for site, elems in self.elems.most_common(top): - print(f"{site:62s} {self.sites[site]:>7d} {elems:>12d} " - f"{','.join(sorted(self.devices[site]))}") - print(f"\n{len(self.elems)} distinct sites, " - f"{sum(self.elems.values())} double elements total") - - -def main() -> int: - ap = argparse.ArgumentParser() - ap.add_argument("--pdb", default="1DAW") - ap.add_argument("--lmax-cap", type=int, default=64) - ap.add_argument("--top", type=int, default=30) - args = ap.parse_args() - - from lab import FRFConfig, rotated_case, run_frf, seed_for - - model, data, _ = rotated_case(args.pdb, seed_for(args.pdb, 0)) - cfg = FRFConfig(n_peaks=500, lmax_cap=args.lmax_cap) - run_frf(model, data, cfg, capture_arf=False, verbose=0) # warm the memos - - audit = DoubleAudit() - with audit: - run_frf(model, data, cfg, capture_arf=False, verbose=0) - print(f"=== {args.pdb} cap{args.lmax_cap}: double-precision sites ===") - audit.report(args.top) - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/alignment_lab/analysis/double_audit.sh b/alignment_lab/analysis/double_audit.sh deleted file mode 100644 index 71a5b187..00000000 --- a/alignment_lab/analysis/double_audit.sh +++ /dev/null @@ -1,19 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=dblaudit -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:40:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname)" -"$PY" -u alignment_lab/analysis/double_audit.py --pdb 1DAW --top 40 2>&1 \ - | grep -vE "UserWarning|FutureWarning|^ *from |^ *warnings\.|Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization" -echo "rc=$?" diff --git a/alignment_lab/analysis/empirical_sigma_a_check.sh b/alignment_lab/analysis/empirical_sigma_a_check.sh deleted file mode 100644 index edf9e956..00000000 --- a/alignment_lab/analysis/empirical_sigma_a_check.sh +++ /dev/null @@ -1,19 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=esa -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:20:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=32G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -for P in 1DAW 2DQ6 3K7M; do - "$PY" -u alignment_lab/diagnostics/empirical_sigma_a_check.py --pdb $P 2>&1 | grep -v "Warning\|warnings.warn" | grep -A14 "^ROW\|Traceback" -done -echo DONE diff --git a/alignment_lab/analysis/eps_gate.sh b/alignment_lab/analysis/eps_gate.sh deleted file mode 100644 index faa398ff..00000000 --- a/alignment_lab/analysis/eps_gate.sh +++ /dev/null @@ -1,22 +0,0 @@ -#!/bin/bash -# Gate for routing the rescore's epsilon through SpaceGroup.epsilon(friedel=False). -# This CHANGES the LLG on every centric reflection, so the unit suite is necessary -# but not sufficient -- the rescore panel is the real test. -#SBATCH --job-name=epsgate -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO"; export PYTHONPATH="$REPO" -export TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname)" -L=alignment_lab/slurm/epsgate_tests_$SLURM_JOB_ID.log -"$PY" -m pytest tests/unit/symmetry tests/unit/alignment tests/unit/frf_separate tests/unit/model -q > "$L" 2>&1 -rc=$?; echo "=== TESTS rc=$rc ==="; tail -12 "$L"; grep "^FAILED" "$L" | head -8 diff --git a/alignment_lab/analysis/first_true_rank.sh b/alignment_lab/analysis/first_true_rank.sh deleted file mode 100644 index ee1dabff..00000000 --- a/alignment_lab/analysis/first_true_rank.sh +++ /dev/null @@ -1,27 +0,0 @@ -#!/bin/bash -# How deep in the rotation function's (de-duplicated) list does the true -# orientation sit? The SOLN table carries each solution's FRF index k and its -# truth flag; the smallest k flagged true per cell is the shortlist depth the -# translation search needed. -#SBATCH --job-name=ftr -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=hour -#SBATCH --time=00:40:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-9 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -for T in 0 1 2 3 4; do - "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial $T --arms llg --verbose 2 2>/dev/null \ - | awk -v pdb=$PDB -v t=$T '/SOLN +[0-9]+ +[0-9]+ .*true/ {k=$3+0; if (min=="" || k int: - ap = argparse.ArgumentParser() - ap.add_argument("--pdb", required=True) - ap.add_argument("--lmax-cap", type=int, default=64) - ap.add_argument("--n-peaks", type=int, default=500) - args = ap.parse_args() - - from alignment_lab.lab.benchmark import load_case - from alignment_lab.lab.frf import FRFConfig, run_frf - - model, data = load_case(args.pdb)[:2] - res = run_frf(model, data, FRFConfig(n_peaks=args.n_peaks, - lmax_cap=args.lmax_cap)) - peaks = res.peaks if hasattr(res, "peaks") else res[0] - # Every fingerprint line is prefixed. Loading a structure writes progress to - # stdout, so a comparison that filters on anything looser (blank lines, a - # leading '#') silently ingests that chatter as data. - print(f"#FP pdb={args.pdb} lmax_cap={args.lmax_cap} n_peaks={len(peaks)}") - for i, p in enumerate(peaks): - print(f"FP {i:4d} {p.alpha:.9g} {p.beta:.9g} {p.gamma:.9g} " - f"{p.score:.9g} {p.sigma:.9g}") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/alignment_lab/analysis/frf_orbit_side.sh b/alignment_lab/analysis/frf_orbit_side.sh deleted file mode 100644 index a5d31fd3..00000000 --- a/alignment_lab/analysis/frf_orbit_side.sh +++ /dev/null @@ -1,19 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=orbside -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:20:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=32G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -for P in 2DQ6 3K7M 1DAW 3GR5; do - "$PY" -u alignment_lab/diagnostics/frf_orbit_side.py --pdb $P 2>&1 | grep -v "Warning\|warnings.warn" | grep -A14 "^#\|^ROW\|Traceback" -done -echo DONE diff --git a/alignment_lab/analysis/friedel_gate.sh b/alignment_lab/analysis/friedel_gate.sh deleted file mode 100644 index ffa2f122..00000000 --- a/alignment_lab/analysis/friedel_gate.sh +++ /dev/null @@ -1,27 +0,0 @@ -#!/bin/bash -# Gate for removing the antipodal copy: the panel again (ranks must hold) plus a -# warm, order-rotated timing check that the 16-22% survives in production code -# rather than only under the knockout patch. -#SBATCH --job-name=fried -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-9 -set -uo pipefail -NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -OLD=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/_stagea_baseline -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -export TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname) pdb=$PDB" -for tree in OLD NEW; do - eval "root=\$$tree"; cd "$root" - PYTHONPATH="$root" "$PY" -u "$NEW/alignment_lab/analysis/panel_ranks.py" \ - --pdb "$PDB" --trials 10 --tag "$tree" 2>/dev/null | grep '^ROW ' -done diff --git a/alignment_lab/analysis/ftf_early_stop.py b/alignment_lab/analysis/ftf_early_stop.py deleted file mode 100644 index 6b02e9c9..00000000 --- a/alignment_lab/analysis/ftf_early_stop.py +++ /dev/null @@ -1,117 +0,0 @@ -"""Would a sequential TFZ-gated search stop on the right orientation? - -The discrimination panel ranks candidates against *each other*, which needs all -of them placed. A guidance loop wants the opposite: walk the FRF peaks in -descending order, place one at a time, and stop as soon as one is convincing -- -paying for the whole list only on the cases that need it. - -That needs an **absolute** criterion, computed from a single candidate. Two are -available per placement and both are already in the code: - -``tfz`` the top translation peak measured against the spread of that - orientation's own translation map (``TranslationPeak.sigma``, Phaser's - TFZ). -``llgz`` the same for the LLG re-rank: best translation against the other 19. - -Simulates the loop over the recorded placements at a range of thresholds. A run -that never crosses the threshold falls back to the argmax over all candidates -placed, which is the no-early-stop behaviour -- so the cost of a threshold set -too high is wasted work, not a failure, and the two are reported separately. -""" - -from __future__ import annotations - -import argparse -import statistics as st -import sys -from collections import defaultdict - - -def load(paths): - cells = defaultdict(list) - for path in paths: - for line in open(path): - if not line.startswith("CAND"): - continue - d = dict(kv.split("=", 1) for kv in line.split() if "=" in kv) - cells[(d["pdb"], int(d["trial"]))].append( - dict(k=int(d["k"]), ang=float(d["ang"]), - truth=d["is_truth"] == "1", - tfz=float(d["tfz"]), llgz=float(d["llgz"]), - tf_llg=float(d["tf_llg"]), r=float(d["r"]))) - for v in cells.values(): - v.sort(key=lambda c: c["k"]) # FRF descending order - return cells - - -def simulate(cells, key, thr): - """Walk each cell in FRF order, stop at the first candidate over ``thr``.""" - hit = miss = exh_hit = exh_miss = 0 - placed, miss_ang = [], [] - for cs in cells.values(): - for i, c in enumerate(cs): - if c[key] >= thr: - placed.append(i + 1) - if c["truth"]: - hit += 1 - else: - miss += 1 - miss_ang.append(c["ang"]) - break - else: # never convinced: rank them all - placed.append(len(cs)) - best = max(cs, key=lambda c: c[key]) - if best["truth"]: - exh_hit += 1 - else: - exh_miss += 1 - miss_ang.append(best["ang"]) - n = len(cells) - return dict(thr=thr, n=n, hit=hit, miss=miss, exh_hit=exh_hit, - exh_miss=exh_miss, ok=hit + exh_hit, - med_placed=st.median(placed), mean_placed=sum(placed) / n, - med_miss_ang=st.median(miss_ang) if miss_ang else float("nan")) - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("logs", nargs="+") - ap.add_argument("--key", default="tfz", choices=["tfz", "llgz"]) - ap.add_argument("--thresholds", default="0,3,4,5,6,7,8,9,10,12,15,20,1e9") - args = ap.parse_args() - - cells = load(args.logs) - if not cells: - print("no CAND lines found", file=sys.stderr) - return 1 - n_cand = st.median([len(v) for v in cells.values()]) - print(f"# {len(cells)} cells, {n_cand:.0f} candidates each, key={args.key}") - print(f"# a cell is OK if the loop commits to a candidate within 8 deg\n") - print(f"{'thr':>6s} {'OK':>7s} {'stop_hit':>8s} {'stop_miss':>9s} " - f"{'exh_hit':>7s} {'exh_miss':>8s} {'med_n':>6s} {'mean_n':>7s} " - f"{'miss_ang':>8s}") - for t in [float(x) for x in args.thresholds.split(",")]: - r = simulate(cells, args.key, t) - label = "none" if t > 1e8 else f"{t:g}" - print(f"{label:>6s} {r['ok']:>4d}/{r['n']:<2d} {r['hit']:>8d} " - f"{r['miss']:>9d} {r['exh_hit']:>7d} {r['exh_miss']:>8d} " - f"{r['med_placed']:>6.0f} {r['mean_placed']:>7.1f} " - f"{r['med_miss_ang']:>8.1f}") - - print("\nper structure at the best threshold by (OK, then fewest placed):") - best = max((simulate(cells, args.key, t) - for t in [float(x) for x in args.thresholds.split(",")]), - key=lambda r: (r["ok"], -r["mean_placed"])) - print(f" thr={best['thr']:g}") - for p in sorted({k[0] for k in cells}): - sub = {k: v for k, v in cells.items() if k[0] == p} - r = simulate(sub, args.key, best["thr"]) - print(f" {p:6s} OK {r['ok']}/{r['n']} placed med={r['med_placed']:.0f} " - f"mean={r['mean_placed']:.1f}" - + ("" if r['med_miss_ang'] != r['med_miss_ang'] - else f" miss_ang={r['med_miss_ang']:.1f} deg")) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/analysis/ftf_running_null.py b/alignment_lab/analysis/ftf_running_null.py deleted file mode 100644 index 50211b4a..00000000 --- a/alignment_lab/analysis/ftf_running_null.py +++ /dev/null @@ -1,140 +0,0 @@ -"""Early stopping needs a cross-candidate contrast, so build the null as you go. - -The per-candidate Z-scores fail as stopping criteria, and the separability table -says why: pooled over 750 placements, a wrong orientation's TFZ is *higher* than -truth's (median 3.15 vs 2.99). TFZ asks "is this translation better than the -other translations for this orientation" -- and a wrong orientation still has a -best translation that stands out of its own map. The contrast that carries the -signal is "is this orientation better than the other orientations", which no -single placement can answer. - -But a sequential loop does not need all 25 placements to answer it -- only -enough of them to know what a wrong answer looks like on this structure. So: -place ``--burn-in`` candidates unconditionally, estimate the null from them with -a median/MAD (robust, because truth is often among the first few and a mean -would be dragged by it), then commit to the first candidate -- burn-in included --- that sits ``--thr`` robust sigmas above it. - -Sweeps threshold and burn-in over the recorded placements, in the FRF order a -live loop would walk. A cell that never crosses falls back to the argmax, which -is the no-early-stop behaviour, so an over-strict threshold costs placements -rather than answers. -""" - -from __future__ import annotations - -import argparse -import statistics as st -from collections import defaultdict - - -def load(paths): - cells = defaultdict(list) - for path in paths: - for line in open(path): - if not line.startswith("CAND"): - continue - d = dict(kv.split("=", 1) for kv in line.split() if "=" in kv) - cells[(d["pdb"], int(d["trial"]))].append( - dict(k=int(d["k"]), ang=float(d["ang"]), - truth=d["is_truth"] == "1", tf_corr=float(d["tf_corr"]), - tf_llg=float(d["tf_llg"]), tfz=float(d["tfz"]), - r=float(d["r"]))) - for v in cells.values(): - v.sort(key=lambda c: c["k"]) - return cells - - -def robust_z(x, ref, higher_is_better=True): - """``x`` in MAD-sigmas above the centre of ``ref``. Sign-normalised.""" - med = st.median(ref) - mad = st.median([abs(v - med) for v in ref]) - scale = 1.4826 * mad - if scale < 1e-30: - return float("inf") if x != med else 0.0 - z = (x - med) / scale - return z if higher_is_better else -z - - -def simulate(cells, key, thr, burn, hi=True): - hit = miss = exh_hit = exh_miss = 0 - placed, miss_ang = [], [] - for cs in cells.values(): - b = min(burn, len(cs)) - ref = [c[key] for c in cs[:b]] - stopped = None - # The burn-in candidates are tested too: truth is often among the first - # few, and a loop that could not commit to one it had already placed - # would pay the whole list on exactly the easy cases. - for i, c in enumerate(cs): - n_paid = max(b, i + 1) - if i >= b: - ref = [q[key] for q in cs[:i]] - if robust_z(c[key], ref, hi) >= thr: - stopped = (i, n_paid, c) - break - if stopped is None: - placed.append(len(cs)) - best = (max if hi else min)(cs, key=lambda c: c[key]) - if best["truth"]: - exh_hit += 1 - else: - exh_miss += 1 - miss_ang.append(best["ang"]) - else: - _, n_paid, c = stopped - placed.append(n_paid) - if c["truth"]: - hit += 1 - else: - miss += 1 - miss_ang.append(c["ang"]) - n = len(cells) - return dict(n=n, hit=hit, miss=miss, exh_hit=exh_hit, exh_miss=exh_miss, - ok=hit + exh_hit, med=st.median(placed), - mean=sum(placed) / n, miss_ang=(st.median(miss_ang) - if miss_ang else float("nan"))) - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("logs", nargs="+") - ap.add_argument("--key", default="tf_llg", - choices=["tf_llg", "tf_corr", "r", "tfz"]) - ap.add_argument("--burn-ins", default="3,5,8") - ap.add_argument("--thresholds", default="3,5,8,10,15,20,30") - args = ap.parse_args() - hi = args.key != "r" - - cells = load(args.logs) - print(f"# {len(cells)} cells x {st.median([len(v) for v in cells.values()]):.0f} " - f"candidates, key={args.key} ({'higher' if hi else 'lower'} is better)") - print(f"# OK = committed to a candidate within 8 deg; placements counted " - f"include the burn-in\n") - print(f"{'burn':>4s} {'thr':>4s} {'OK':>7s} {'stop_hit':>8s} {'stop_miss':>9s} " - f"{'exh_hit':>7s} {'exh_miss':>8s} {'med_n':>5s} {'mean_n':>6s} " - f"{'miss_ang':>8s}") - best = None - for burn in [int(x) for x in args.burn_ins.split(",")]: - for thr in [float(x) for x in args.thresholds.split(",")]: - r = simulate(cells, args.key, thr, burn, hi) - print(f"{burn:>4d} {thr:>4g} {r['ok']:>4d}/{r['n']:<2d} " - f"{r['hit']:>8d} {r['miss']:>9d} {r['exh_hit']:>7d} " - f"{r['exh_miss']:>8d} {r['med']:>5.0f} {r['mean']:>6.1f} " - f"{r['miss_ang']:>8.1f}") - if best is None or (r["ok"], -r["mean"]) > (best[0]["ok"], -best[0]["mean"]): - best = (r, burn, thr) - r, burn, thr = best - print(f"\nper structure at burn={burn} thr={thr:g}:") - for p in sorted({k[0] for k in cells}): - sub = {k: v for k, v in cells.items() if k[0] == p} - s = simulate(sub, args.key, thr, burn, hi) - print(f" {p:6s} OK {s['ok']}/{s['n']} placed med={s['med']:.0f} " - f"mean={s['mean']:.1f}" - + ("" if s["miss_ang"] != s["miss_ang"] - else f" miss_ang={s['miss_ang']:.0f} deg")) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/analysis/full_gate.sh b/alignment_lab/analysis/full_gate.sh deleted file mode 100644 index 2ab25845..00000000 --- a/alignment_lab/analysis/full_gate.sh +++ /dev/null @@ -1,26 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=fullgate -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=day -#SBATCH --time=03:00:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -export TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 -# `-c tests/pytest.ini` from the repo ROOT, with --run-slow. Both halves matter -# and they pull in opposite directions: a bare `pytest tests/unit` picks up -# pyproject.toml, where the slow marker silently skips the rotation-search and -# translation tests -- which is how seven of them stayed broken by a merge -# through a gate reporting everything green. But running from `tests/` to get -# the right config then breaks the io tests, which open `tests/files/...` -# relative to the root. Naming the config explicitly satisfies both. -"$PY" -m pytest -c tests/pytest.ini tests/unit --run-slow -q 2>&1 | tail -16 -echo "PYTEST_RC=${PIPESTATUS[0]}" -echo "== end-to-end placement, one cell ==" -"$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb 1DAW --trial 0 \ - --arms analytic_r 2>/dev/null | grep '^ROW ' diff --git a/alignment_lab/analysis/gpfs_cold_read.sh b/alignment_lab/analysis/gpfs_cold_read.sh deleted file mode 100644 index 411ce2c0..00000000 --- a/alignment_lab/analysis/gpfs_cold_read.sh +++ /dev/null @@ -1,35 +0,0 @@ -#!/bin/bash -# Is reading the Wigner d-table off GPFS actually cheaper than recomputing it? -# Only a COLD read answers that, so the writer and the reader must be different -# nodes -- a same-node re-read is served from the page cache and is meaningless. -#SBATCH --job-name=frf_cold -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 PYTHONUNBUFFERED=1 -export CUDA_VISIBLE_DEVICES="" -MODE="$1" -DIR=alignment_lab/runs/gpfs_cold -mkdir -p "$DIR" alignment_lab/slurm -echo "mode=$MODE host=$(hostname)" -"$PY" -u -c " -import os, time, torch -mode, d = '$MODE', '$DIR' -for L, mb in ((65, 88), (101, 330)): - p = os.path.join(d, f'dtable_L{L}.pt') - if mode == 'write': - n = int(mb * 1e6 / 4) - torch.save(torch.zeros(n, dtype=torch.float32), p) - print(f' wrote L={L} {os.path.getsize(p)/1e6:.0f} MB') - else: - t0 = time.perf_counter(); torch.load(p, map_location='cpu') - cold = time.perf_counter() - t0 - t0 = time.perf_counter(); torch.load(p, map_location='cpu') - warm = time.perf_counter() - t0 - print(f' L={L} {os.path.getsize(p)/1e6:5.0f} MB cold {cold*1e3:7.0f} ms ' - f'warm {warm*1e3:6.0f} ms') -" -echo "exit_code=$?" diff --git a/alignment_lab/analysis/integration_alignment.sh b/alignment_lab/analysis/integration_alignment.sh deleted file mode 100644 index 72f6f929..00000000 --- a/alignment_lab/analysis/integration_alignment.sh +++ /dev/null @@ -1,17 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=integ -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:59:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -"$PY" -m pytest -c tests/pytest.ini tests/integration/alignment tests/unit/alignment tests/unit/scaling tests/unit/frf_separate --run-slow -q --tb=short 2>&1 | grep -v "^✓" | tail -60 -echo "PYTEST_RC=${PIPESTATUS[0]}" diff --git a/alignment_lab/analysis/kernel_ab.sh b/alignment_lab/analysis/kernel_ab.sh deleted file mode 100644 index 4fcc0c9a..00000000 --- a/alignment_lab/analysis/kernel_ab.sh +++ /dev/null @@ -1,118 +0,0 @@ -#!/bin/bash -# Fused C++ kernel against the portable torch reference: correctness first, then -# speed. Interleaved repeats and the truth rank beside each timing. -#SBATCH --job-name=frf_kernel_ab -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 -export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" -"$PY" -u -c " -import sys, time -sys.path.insert(0,'alignment_lab') -import torch; torch.set_grad_enabled(False) -from torchref.experimental.alignment.frf.kernels.cpu import legendre_shell as K -from torchref.experimental.alignment.frf.kernels import portable as P -from torchref.utils.backends import set_force_portable - -print('kernel available:', K.available()) -print('float64 must be refused, not reinterpreted:') -try: - import torch as _t - z = _t.zeros(2, 3, 5, dtype=_t.float64) - K.legendre_shell_accumulate( - z, z.clone(), _t.zeros(4, dtype=_t.float64), _t.zeros(4, dtype=_t.float64), - _t.zeros(4, 5, dtype=_t.float64), _t.zeros(4, 5, dtype=_t.float64), - _t.zeros(4, dtype=_t.long), _t.zeros(5, 5, dtype=_t.float64), - _t.zeros(5, 5, dtype=_t.float64), _t.zeros(5, dtype=_t.float64)) - print(' PROBLEM: float64 was accepted') -except Exception as e: - msg = str(e) - ok = ('float32 only' in msg) or ('dtype' in msg) - print((' refused: ' if ok else ' WRONG ERROR (not a dtype refusal): ') - + f'{type(e).__name__}: {msg.splitlines()[0][:110]}') - if not ok: - raise SystemExit(1) -if not K.available(): - print('why:', K.why_unavailable()) - err = K.last_error() - if err: print(err[1][:3000]) - raise SystemExit(1) - -# --- correctness: fused vs portable on random shell-sorted input ------------- -from torchref.experimental.alignment.sh import legendre_recurrence_coefficients -g = torch.Generator().manual_seed(4) -for L, n_c, n_sh, dt in ((13, 500, 40, torch.float32), - (65, 4000, 300, torch.float32), - (101, 3000, 250, torch.float32), - (65, 2000, 150, torch.float32)): - n_even = (L - 1 if (L-1) % 2 == 0 else L - 2) // 2 - ct = (2*torch.rand(n_c, generator=g, dtype=dt)-1) - st = (1-ct*ct).clamp(min=0).sqrt() - Dr = torch.randn(n_c, L, generator=g, dtype=dt) - Di = torch.randn(n_c, L, generator=g, dtype=dt) - sh = torch.sort(torch.randint(0, n_sh, (n_c,), generator=g))[0] - a, b, se = legendre_recurrence_coefficients(L, dt, torch.device('cpu')) - ref_r = torch.zeros(n_even, n_sh, L, dtype=dt); ref_i = torch.zeros_like(ref_r) - got_r = torch.zeros_like(ref_r); got_i = torch.zeros_like(ref_r) - P.legendre_shell_accumulate(ref_r, ref_i, ct, st, Dr, Di, sh, a, b, se) - K.legendre_shell_accumulate(got_r, got_i, ct, st, Dr, Di, sh, a, b, se) - sc = max(ref_r.abs().max().item(), 1e-300) - er = (got_r-ref_r).abs().max().item()/sc - ei = (got_i-ref_i).abs().max().item()/max(ref_i.abs().max().item(),1e-300) - print(f' L={L:3d} n_c={n_c:5d} {str(dt):15s} rel err re {er:.2e} im {ei:.2e}') - -# --- what single precision costs, against an ungrouped float64 reference ---- -import torchref.experimental.alignment.frf.data_mr as dm -g2 = torch.Generator().manual_seed(77) -for L, hs in ((65, 64.0), (101, 100.0)): - sv = torch.randn(6000, 3, generator=g2, dtype=torch.float64) - sv = sv / sv.norm(dim=-1, keepdim=True) * ( - 0.07 + 0.18*torch.rand(6000, 1, generator=g2, dtype=torch.float64)) - I = torch.randn(6000, generator=g2, dtype=torch.float64) - ks, kc = dm._GROUP_SCALE_S, dm._GROUP_SCALE_COS - dm._GROUP_SCALE_S = dm._GROUP_SCALE_COS = 10**16 - exact = dm.bessel_sh_expand(sv, I, L=L, bessel_h_scale=hs).coeffs - dm._GROUP_SCALE_S, dm._GROUP_SCALE_COS = ks, kc - sc = max(exact.abs().max().item(), 1e-300) - f64 = dm.bessel_sh_expand(sv, I, L=L, bessel_h_scale=hs).coeffs - f32 = dm.bessel_sh_expand(sv, I, L=L, bessel_h_scale=hs, - compute_dtype=torch.complex64).coeffs - print(f' L={L:3d} vs exact: float64 angular {((f64-exact).abs().max()/sc):.2e}' - f' float32 angular {((f32-exact).abs().max()/sc):.2e}') - -# --- speed on the real thing ------------------------------------------------- -from lab import FRFConfig, orbit_rank, rotated_case, run_frf, seed_for -cases = {p: rotated_case(p, seed_for(p, 0)) for p in ('1DAW', '3K7M')} -for cap in (64, 100): - cfg = FRFConfig(lmax_cap=cap, n_peaks=200) - for forced in (True, False): - set_force_portable(forced) - for m, d, _ in cases.values(): - run_frf(m, d, cfg, capture_arf=False) - res, rank = {}, {} - for rep in range(3): - for arm, forced in (('portable', True), ('fused', False)): - set_force_portable(forced) - for p, (m, d, R) in cases.items(): - t0 = time.perf_counter() - r = run_frf(m, d, cfg, capture_arf=False) - res.setdefault((arm,p), []).append(time.perf_counter()-t0) - k, _ = orbit_rank(r.peaks, R, - d.spacegroup.matrices.to(torch.float64).cpu(), - reciprocal_basis=d.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), - side='left', frame='cart') - rank[(arm,p)] = k - set_force_portable(None) - print(f'--- cap{cap} (best of 3, seconds) ---') - for p in cases: - e, c = min(res[('portable',p)]), min(res[('fused',p)]) - print(f' {p:6s} portable {e:7.2f} [rank {rank[(\"portable\",p)]:>3}] ' - f'fused {c:7.2f} [rank {rank[(\"fused\",p)]:>3}] {e/max(c,1e-9):5.2f}x') -" -echo "exit_code=$?" diff --git a/alignment_lab/analysis/kernel_check.sh b/alignment_lab/analysis/kernel_check.sh deleted file mode 100644 index 449319aa..00000000 --- a/alignment_lab/analysis/kernel_check.sh +++ /dev/null @@ -1,62 +0,0 @@ -#!/bin/bash -# Does the fused kernel build, refuse float64, and agree with the portable -# reference? Correctness only, so it does not need an exclusive node -- the -# timing A/B does. -#SBATCH --job-name=frf_kernel_check -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 PYTHONUNBUFFERED=1 -export CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname)" -"$PY" -u -c " -import torch; torch.set_grad_enabled(False) -from torchref.experimental.alignment.frf.kernels.cpu import legendre_shell as K -from torchref.experimental.alignment.frf.kernels import portable as P -from torchref.experimental.alignment.sh import legendre_recurrence_coefficients - -print('available:', K.available()) -if not K.available(): - print('why:', K.why_unavailable()) - e = K.last_error() - if e: print(e[1][-4000:]) - raise SystemExit(1) - -try: - z = torch.zeros(2, 3, 5, dtype=torch.float64) - K.legendre_shell_accumulate(z, z.clone(), - torch.zeros(4, dtype=torch.float64), torch.zeros(4, dtype=torch.float64), - torch.zeros(4, 5, dtype=torch.float64), torch.zeros(4, 5, dtype=torch.float64), - torch.zeros(4, dtype=torch.long), torch.zeros(5, 5, dtype=torch.float64), - torch.zeros(5, 5, dtype=torch.float64), torch.zeros(5, dtype=torch.float64)) - print('PROBLEM: float64 accepted') -except Exception as e: - msg = str(e) - ok = ('float32 only' in msg) or ('dtype' in msg) - print(('float64 refused: ' if ok else 'WRONG ERROR (not a dtype refusal): ') - + msg.splitlines()[0][:120]) - if not ok: - raise SystemExit(1) - -g = torch.Generator().manual_seed(4) -for L, n_c, n_sh in ((13, 500, 40), (65, 4000, 300), (101, 3000, 250)): - n_even = (L - 1 if (L-1) % 2 == 0 else L - 2) // 2 - ct = (2*torch.rand(n_c, generator=g, dtype=torch.float32)-1) - st = (1-ct*ct).clamp(min=0).sqrt() - Dr = torch.randn(n_c, L, generator=g, dtype=torch.float32) - Di = torch.randn(n_c, L, generator=g, dtype=torch.float32) - sh = torch.sort(torch.randint(0, n_sh, (n_c,), generator=g))[0] - a, b, se = legendre_recurrence_coefficients(L, torch.float32, torch.device('cpu')) - ref_r = torch.zeros(n_even, n_sh, L, dtype=torch.float32); ref_i = torch.zeros_like(ref_r) - got_r = torch.zeros_like(ref_r); got_i = torch.zeros_like(ref_r) - P.legendre_shell_accumulate(ref_r, ref_i, ct, st, Dr, Di, sh, a, b, se) - K.legendre_shell_accumulate(got_r, got_i, ct, st, Dr, Di, sh, a, b, se) - sr = max(ref_r.abs().max().item(), 1e-30); si = max(ref_i.abs().max().item(), 1e-30) - print(f' L={L:3d}: rel err re {(got_r-ref_r).abs().max().item()/sr:.2e}' - f' im {(got_i-ref_i).abs().max().item()/si:.2e}') -" -echo "exit_code=$?" diff --git a/alignment_lab/analysis/marginal_seeds.sh b/alignment_lab/analysis/marginal_seeds.sh deleted file mode 100644 index 0d50d0ad..00000000 --- a/alignment_lab/analysis/marginal_seeds.sh +++ /dev/null @@ -1,30 +0,0 @@ -#!/bin/bash -# Is 2DQ6 a cell the pipeline solves, or a coin flip the panel samples once? -# -# 2DQ6 t1 moved 6.5 -> 28.7 deg on a 1.3e-5 change in Sigma(s), taking the panel -# from 30/30 to 29/30. If the structure is genuinely marginal then its pass rate -# over SEEDS is intermediate, and a single flip carries no information about the -# change that produced it -- a different seed would have flipped it anyway. -# Three trials per structure cannot tell the difference; ten can. -#SBATCH --job-name=marg -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=day -#SBATCH --time=04:00:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-5 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -# The four structures the translation search used to mis-place, and two controls. -PDBS=(2DQ6 6G9X 1DAW 3K7M 3VRJ 4BX9) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 -export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -for T in $(seq 0 9); do - "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial "$T" \ - --arms llg,analytic_r,corr --n-rotation-candidates 25 2>/dev/null | grep '^ROW ' -done diff --git a/alignment_lab/analysis/merge_gate.sh b/alignment_lab/analysis/merge_gate.sh deleted file mode 100644 index f5fc1a93..00000000 --- a/alignment_lab/analysis/merge_gate.sh +++ /dev/null @@ -1,53 +0,0 @@ -#!/bin/bash -# Post-merge gate. Two questions, in order: -# 1. does the merged tree still work at all (imports, full unit suite); -# 2. did dev's changes move the FRF's answers? The fingerprints are compared -# against 775576bc, our last pre-merge commit, so any difference is dev's -# or the merge resolution's -- not ours. -#SBATCH --job-name=mergegate -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -OUT=$NEW/alignment_lab/slurm -cd "$NEW"; export PYTHONPATH="$NEW" -export TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname)" - -echo "=== import smoke ===" -"$PY" -c "import torchref; import torchref.experimental.alignment as a; print(' torchref + alignment import OK')" 2>&1 | tail -5 - -LOG=$OUT/mergegate_tests_$SLURM_JOB_ID.log -"$PY" -m pytest tests/unit -q > "$LOG" 2>&1 -rc=$? -echo "=== FULL UNIT SUITE rc=$rc ===" -tail -20 "$LOG" - -echo "=== does a copy still have a usable, correct partition? ===" -"$PY" -m pytest tests/unit/model/test_copy.py -q 2>&1 | tail -6 - -echo "=== FRF fingerprints vs 775576bc (pre-merge) ===" -for pdb in 3K7M 1DAW; do - "$PY" -u "$NEW/alignment_lab/analysis/frf_fingerprint.py" --pdb "$pdb" --lmax-cap 64 \ - > "$OUT/merge_NEW_${pdb}.txt" 2>/dev/null || echo " RUN FAILED $pdb" - ref=$OUT/cap_NEW_${pdb}.txt # captured at 775576bc by capability_gate - if [ -s "$ref" ] && [ -s "$OUT/merge_NEW_${pdb}.txt" ]; then - if diff -q <(grep '^FP ' "$ref") <(grep '^FP ' "$OUT/merge_NEW_${pdb}.txt") >/dev/null; then - echo " IDENTICAL $pdb ($(grep -c '^FP ' "$OUT/merge_NEW_${pdb}.txt") peaks)" - else - nd=$(diff <(grep '^FP ' "$ref") <(grep '^FP ' "$OUT/merge_NEW_${pdb}.txt") | grep -c '^[<>]') - echo " DIFFERS $pdb ($nd lines)" - diff <(grep '^FP ' "$ref") <(grep '^FP ' "$OUT/merge_NEW_${pdb}.txt") | head -4 - fi - else - echo " no comparable reference for $pdb" - fi -done -echo "mergegate_rc=$rc" diff --git a/alignment_lab/analysis/merge_identity.sh b/alignment_lab/analysis/merge_identity.sh deleted file mode 100644 index a0600b55..00000000 --- a/alignment_lab/analysis/merge_identity.sh +++ /dev/null @@ -1,31 +0,0 @@ -#!/bin/bash -# Same script, same data, two trees: post-merge and the commit before it. -# Single-threaded because the FRF peak list only reproduces bit-for-bit at one -# thread -- ~5e-8 of score noise reorders peaks otherwise, and this gate is -# about bit-identity. -#SBATCH --job-name=mergeid -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=2 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-3 -set -uo pipefail -POST=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PRE=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/premerge_check -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -# 1DAW is C2 (monoclinic) and 2DQ6 is P3121: on an orthogonal cell the -# fractionalisation matrix is diagonal, so a transposed edge-vector convention -# in the new voxel_size would not show. On these it would. -PDBS=(1DAW 2DQ6 3K7M 4BX9) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -export TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 -for TREE in "$PRE" "$POST"; do - TAG=$([ "$TREE" = "$PRE" ] && echo pre || echo post) - cd "$TREE" - PYTHONPATH="$TREE" "$PY" -u alignment_lab/analysis/merge_numeric_identity.py \ - --pdb "$PDB" --tag "$TAG" 2>/dev/null | grep '^OUT ' -done diff --git a/alignment_lab/analysis/merge_numeric_identity.py b/alignment_lab/analysis/merge_numeric_identity.py deleted file mode 100644 index 8c08d1d5..00000000 --- a/alignment_lab/analysis/merge_numeric_identity.py +++ /dev/null @@ -1,86 +0,0 @@ -"""Did merging dev change any number the alignment path produces? - -`a596ed9e` removes SfFFT's stored real-space coordinate grid. Reading it, the -change is plumbing: `build_electron_density` used the tensor only for `.device` -and `.shape[:-1]`, the four Triton kernels never dereferenced the `grid_ptr` -they were handed, and the new `voxel_size = frac_matrix @ (1/gridsize)` is -algebraically the old `grid[2,2,2] - grid[1,1,1]` -- `cart = B @ f`, so column j -of the fractionalisation matrix is cell edge vector j. - -"Reading it, it looks equivalent" is not the same as equivalent. This dumps -hashes of the quantities the alignment stack actually consumes so the two trees -can be compared bit for bit. - -Deliberately covers a monoclinic and two high-symmetry cells: for an orthogonal -cell the fractionalisation matrix is diagonal, so a transposed edge-vector -convention would be invisible. 1DAW is C2 and 2DQ6 is P3121, where it would not. -""" - -from __future__ import annotations - -import argparse -import hashlib -import sys -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import BENCH_PDBS, FRFConfig, load_case, run_frf # noqa: E402 - - -def _h(t: torch.Tensor) -> str: - """Bit-exact hash of a tensor's contents, dtype and shape.""" - t = t.detach().cpu().contiguous() - m = hashlib.sha256() - m.update(str(tuple(t.shape)).encode()) - m.update(str(t.dtype).encode()) - m.update(t.numpy().tobytes()) - return m.hexdigest()[:16] - - -def main() -> int: - ap = argparse.ArgumentParser() - ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) - ap.add_argument("--tag", default="?") - ap.add_argument("--lmax-cap", type=int, default=64) - args = ap.parse_args() - - model, data = load_case(args.pdb) - hkl = data.hkl - - # 1. Structure factors -- the thing every downstream number is built on. - F = model.get_structure_factor(hkl, recalc=True) - print(f"OUT {args.tag} {args.pdb} F_calc {_h(F)}") - - # 2. The density map itself, one step earlier than F_calc, so a difference - # can be localised to the splat rather than the FFT. - model.setup_grid() - dm = model.build_complete_map() - print(f"OUT {args.tag} {args.pdb} density {_h(dm)}") - print(f"OUT {args.tag} {args.pdb} gridshape {tuple(dm.shape)}") - - # 3. voxel_size: the one quantity whose FORMULA changed, rather than only - # its call site. Nothing reads it downstream, so a difference here is - # reportable but not itself a regression. - vs = model.voxel_size - print(f"OUT {args.tag} {args.pdb} voxel_size " - f"{'None' if vs is None else _h(vs)} " - f"{'' if vs is None else [f'{float(v):.17g}' for v in vs.flatten()]}") - - # 4. The FRF peak list -- what this branch is actually judged on. - res = run_frf(model, data, FRFConfig(n_peaks=500, lmax_cap=args.lmax_cap), - capture_arf=False, verbose=0) - ang = torch.tensor([[p.alpha, p.beta, p.gamma] for p in res.peaks], - dtype=torch.float64) - sc = torch.tensor([p.score for p in res.peaks], dtype=torch.float64) - print(f"OUT {args.tag} {args.pdb} peaks_angles {_h(ang)}") - print(f"OUT {args.tag} {args.pdb} peaks_scores {_h(sc)}") - print(f"OUT {args.tag} {args.pdb} n_peaks {len(res.peaks)}") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/alignment_lab/analysis/p1_grid_coherence.sh b/alignment_lab/analysis/p1_grid_coherence.sh deleted file mode 100644 index 6d4116f2..00000000 --- a/alignment_lab/analysis/p1_grid_coherence.sh +++ /dev/null @@ -1,19 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=p1coh -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:30:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -for P in 1DAW 2DQ6 3K7M 4BX9; do - "$PY" -u alignment_lab/diagnostics/p1_grid_coherence.py --pdb $P 2>/dev/null | grep -E "^#|^ROW" -done -echo DONE diff --git a/alignment_lab/analysis/panel_arms.py b/alignment_lab/analysis/panel_arms.py deleted file mode 100644 index d46bc726..00000000 --- a/alignment_lab/analysis/panel_arms.py +++ /dev/null @@ -1,121 +0,0 @@ -"""Truth rank per seeded trial under one or more knocked-out engine stages. - -Two stages of the rotation function are suspected of earning nothing, each for a -different reason, and both are cheaper to decide by knocking them out than by -reasoning about them: - -``no_brel`` - ``fit_relative_wilson_b`` scales ``F_calc`` by ``exp(-B_rel s^2/4)`` and the - very next step, ``wilson_normalise``, divides each equal-count shell by its - own ``sqrt()`` -- which removes the shell-mean radial profile, - B_rel's included. Only the within-shell residual of a smooth exponential can - survive. If that is below the engine's own spread, the fit is a sort, a - binning and a regression for nothing. - -``no_friedel`` - ``enforce_friedel`` concatenates ``-s`` onto both reflection sets. Only even - ``l`` are ever computed and ``Y_lm(-s_hat) = Y_lm(s_hat)`` for even ``l``, - with the intensity duplicated verbatim, so ``c_nlm`` doubles *exactly*. Both - sides doubled means xi scales by 4, and the mean and standard deviation of - the rotation function scale with it, so z-scores and the ranking are - invariant. The prediction is therefore sharp: identical ranks, and the raw - score up by exactly 4. Reported so it can be checked rather than assumed. - -Emits one ``ROW`` line per (arm, trial), carrying the top score so the scaling -prediction is falsifiable from the output. -""" - -from __future__ import annotations - -import argparse -import contextlib -import functools -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, FRFConfig, orbit_rank, rotated_case, # noqa: E402 - run_frf, seed_for) - -ARMS = ("production", "no_brel", "no_friedel") - - -@contextlib.contextmanager -def knocked_out(arm: str): - """Disable one stage for the duration of a call, then restore it.""" - if arm == "production": - yield - return - if arm == "no_brel": - import torchref.experimental.alignment.frf.preprocessing as pp - original = pp.fit_relative_wilson_b - pp.fit_relative_wilson_b = lambda *a, **k: 0.0 - try: - yield - finally: - pp.fit_relative_wilson_b = original - return - if arm == "no_friedel": - import torchref.experimental.alignment.frf.api as api - original = api.bessel_sh_expand - - @functools.wraps(original) - def no_mate(*a, **k): - k["enforce_friedel"] = False - return original(*a, **k) - - api.bessel_sh_expand = no_mate - try: - yield - finally: - api.bessel_sh_expand = original - return - raise ValueError(f"unknown arm {arm!r}") - - -def main() -> int: - ap = argparse.ArgumentParser() - ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) - ap.add_argument("--trials", type=int, default=10) - ap.add_argument("--arms", default=",".join(ARMS)) - ap.add_argument("--lmax-cap", type=int, default=64) - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--thr-deg", type=float, default=5.0) - args = ap.parse_args() - - cfg = FRFConfig(n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) - arms = [a for a in args.arms.split(",") if a] - # Trial-major so the same seed's arms run back to back: any drift in machine - # state affects the arms together rather than one of them. - for trial in range(args.trials): - seed = seed_for(args.pdb, trial) - for arm in arms: - model, data, R_true = rotated_case(args.pdb, seed) - t0 = time.time() - with knocked_out(arm): - res = run_frf(model, data, cfg, capture_arf=False, verbose=0) - seconds = time.time() - t0 - rank, ang = orbit_rank( - res.peaks, R_true, - data.spacegroup.matrices.to(torch.float64).cpu(), - reciprocal_basis=data.cell.reciprocal_basis_matrix.to( - torch.float64).cpu(), - side="left", frame="cart", thr_deg=args.thr_deg, - ) - top = res.peaks[0] if res.peaks else None - print(f"ROW {arm} {args.pdb} trial={trial} seed={seed} " - f"rank={rank} rank_cmp={rank if rank >= 0 else args.n_peaks} " - f"top20={int(0 <= rank < 20)} " - f"top_score={'' if top is None else f'{top.score:.10g}'} " - f"top_sigma={'' if top is None else f'{top.sigma:.6g}'} " - f"seconds={seconds:.2f}", flush=True) - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/alignment_lab/analysis/panel_arms.sh b/alignment_lab/analysis/panel_arms.sh deleted file mode 100644 index c8a880d4..00000000 --- a/alignment_lab/analysis/panel_arms.sh +++ /dev/null @@ -1,22 +0,0 @@ -#!/bin/bash -# Decision gates for the two suspected-vestigial stages, over the full panel. -#SBATCH --job-name=arms -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-9 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname) pdb=$PDB" -"$PY" -u alignment_lab/analysis/panel_arms.py --pdb "$PDB" --trials 10 2>/dev/null | grep '^ROW ' -echo "rc=$?" diff --git a/alignment_lab/analysis/panel_ranks.py b/alignment_lab/analysis/panel_ranks.py deleted file mode 100644 index dfee4a41..00000000 --- a/alignment_lab/analysis/panel_ranks.py +++ /dev/null @@ -1,80 +0,0 @@ -"""Truth rank over one benchmark structure at seeded orientations. - -The acceptance criterion for a change to the rotation function is not the median -rank -- it is whether truth lands inside the candidate window the placement -search carries forward, on nearly every trial, for *every* structure. Seed-to- -seed spread at ``lmax_cap = 64`` is +-4 to 6 ranks, so a bare median hides the -cases that decide it. - -Emits one ``ROW`` line per trial so the caller can pair the same seed across two -worktrees. Deliberately reuses the lab's ``seed_for`` / ``rotated_case`` / -``orbit_rank``, which carry the seed contract and the orbit conventions the -earlier sweeps were measured with -- a private reimplementation of any of those -would make the comparison meaningless. -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, FRFConfig, orbit_rank, rotated_case, # noqa: E402 - run_frf, seed_for) - - -def main() -> int: - ap = argparse.ArgumentParser() - ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) - ap.add_argument("--trials", type=int, default=10) - ap.add_argument("--lmax-cap", type=int, default=64) - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--thr-deg", type=float, default=5.0) - ap.add_argument("--tag", default="?", help="which tree this run came from") - ap.add_argument("--exclude-h", action="store_true", - help="drop hydrogens from F_calc. dev now keeps them by " - "default and generates any a file lacks, which is ~47% " - "of the atom count and moves |F_calc| by ~7% at the " - "median. Whether GENERATED hydrogens belong in a " - "molecular-replacement search model is a separate " - "question from whether they belong in refinement.") - args = ap.parse_args() - - cfg = FRFConfig(n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) - for trial in range(args.trials): - seed = seed_for(args.pdb, trial) - model, data, R_true = rotated_case(args.pdb, seed) - if args.exclude_h: - model.exclude_H_from_sf = True - t0 = time.time() - res = run_frf(model, data, cfg, capture_arf=False, verbose=0) - seconds = time.time() - t0 - rank, ang = orbit_rank( - res.peaks, R_true, - data.spacegroup.matrices.to(torch.float64).cpu(), - reciprocal_basis=data.cell.reciprocal_basis_matrix.to( - torch.float64).cpu(), - side="left", frame="cart", thr_deg=args.thr_deg, - ) - # orbit_rank returns -1 for "no peak within thr_deg". A miss must not - # sort as a good rank, so for comparison it counts as worse than the - # worst hit -- the peak-list length. - rank_cmp = rank if rank >= 0 else args.n_peaks - n_h = int((model.pdb["element"].str.strip() == "H").sum()) - print(f"ROW {args.tag} {args.pdb} trial={trial} seed={seed} " - f"nH={n_h} exclH={int(args.exclude_h)} " - f"rank={rank} rank_cmp={rank_cmp} found={int(rank >= 0)} " - f"top20={int(0 <= rank < 20)} " - f"angle={'' if ang is None else round(float(ang), 3)} " - f"seconds={seconds:.2f} sg={data.spacegroup.hm}", flush=True) - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/alignment_lab/analysis/pipeline_timing.sh b/alignment_lab/analysis/pipeline_timing.sh deleted file mode 100644 index 842dcb30..00000000 --- a/alignment_lab/analysis/pipeline_timing.sh +++ /dev/null @@ -1,19 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=ptime -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:59:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --exclusive -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -"$PY" -u alignment_lab/diagnostics/pipeline_timing.py --threads 8 2>/dev/null \ - | grep -E "^ROW|^stage|^[0-9]_|^TOTAL|^---" -echo DONE diff --git a/alignment_lab/analysis/pipeline_timing_gpu.sh b/alignment_lab/analysis/pipeline_timing_gpu.sh deleted file mode 100644 index 9997be45..00000000 --- a/alignment_lab/analysis/pipeline_timing_gpu.sh +++ /dev/null @@ -1,20 +0,0 @@ -#!/bin/bash -# The same warm single-process timing on an A100, model and data on the GPU. -#SBATCH --job-name=ptimegpu -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=gpu -#SBATCH --time=00:40:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --gres=gpu:nvidia_a100-pcie-40gb:1 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 TORCHREF_DEVICE=cuda -nvidia-smi --query-gpu=name --format=csv,noheader | head -1 -"$PY" -u alignment_lab/diagnostics/pipeline_timing.py --threads 8 --device cuda 2>&1 \ - | grep -v "Warning\|warnings.warn" | grep -E "^ROW|^stage|^[0-9]_|^TOTAL|^---|Traceback|Error" | grep -v "^---" -echo DONE diff --git a/alignment_lab/analysis/pipeline_timing_ncand10.sh b/alignment_lab/analysis/pipeline_timing_ncand10.sh deleted file mode 100644 index 3f3eec5c..00000000 --- a/alignment_lab/analysis/pipeline_timing_ncand10.sh +++ /dev/null @@ -1,19 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=ptime10 -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:59:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --exclusive -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -"$PY" -u alignment_lab/diagnostics/pipeline_timing.py --threads 8 --n-rotation-candidates 10 2>/dev/null \ - | grep -E "^ROW|^stage|^[0-9]_|^TOTAL|^---" -echo DONE diff --git a/alignment_lab/analysis/pose_arms.sh b/alignment_lab/analysis/pose_arms.sh deleted file mode 100644 index 32641b11..00000000 --- a/alignment_lab/analysis/pose_arms.sh +++ /dev/null @@ -1,35 +0,0 @@ -#!/bin/bash -# End-to-end pose recovery, which is the only metric that is actually the -# deliverable. Rank is a proxy; this is not. -# -# 10 structures x 3 trials x 2 arms. The arms are the open ranking question: -# analytic R (the default) against the translation LLG, which wins 27/30 to -# 22/30 at rank level. -# -# 25 candidates and no early stopping, matching the pipeline. Under the old rule -# it walked the list until a placement beat R < 0.45 and returned that, so the -# answer depended on FRF order and could not be compared against a harness that -# ranks the whole list -- 2DQ6 solved 6/10 end to end while truth was top-ranked -# by analytic R in 0/10, which is only possible if the two measure different -# things. -#SBATCH --job-name=posearm -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=day -#SBATCH --time=04:00:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-9 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -for T in 0 1 2; do - "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial "$T" \ - --arms llg,analytic_r --n-rotation-candidates 25 2>/dev/null | grep '^ROW ' -done diff --git a/alignment_lab/analysis/pose_panel_ncand10.sh b/alignment_lab/analysis/pose_panel_ncand10.sh deleted file mode 100644 index e0dd70b6..00000000 --- a/alignment_lab/analysis/pose_panel_ncand10.sh +++ /dev/null @@ -1,28 +0,0 @@ -#!/bin/bash -# The 10 x 3 panel with the pose gate (rotation AND translation), at the -# pipeline's default translation window and with the window removed. The -# default is now the rotation search's own window; "full" is what it used to be. -#SBATCH --job-name=ptrans10 -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=hour -#SBATCH --time=00:59:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-9 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -for T in 0 1 2; do - "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial $T --arms llg --n-rotation-candidates 10 \ - 2>/dev/null | grep '^ROW ' | sed 's/^ROW/ROW window=default/' - "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial $T --arms llg --n-rotation-candidates 10 \ - --tf-d-min 0 --tf-d-max inf 2>/dev/null | grep '^ROW ' | sed 's/^ROW/ROW window=full/' -done -echo DONE diff --git a/alignment_lab/analysis/pose_panel_trans.sh b/alignment_lab/analysis/pose_panel_trans.sh deleted file mode 100644 index 6df3cbe0..00000000 --- a/alignment_lab/analysis/pose_panel_trans.sh +++ /dev/null @@ -1,28 +0,0 @@ -#!/bin/bash -# The 10 x 3 panel with the pose gate (rotation AND translation), at the -# pipeline's default translation window and with the window removed. The -# default is now the rotation search's own window; "full" is what it used to be. -#SBATCH --job-name=ptrans -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=hour -#SBATCH --time=00:59:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-9 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -for T in 0 1 2; do - "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial $T --arms llg \ - 2>/dev/null | grep '^ROW ' | sed 's/^ROW/ROW window=default/' - "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb "$PDB" --trial $T --arms llg \ - --tf-d-min 0 --tf-d-max inf 2>/dev/null | grep '^ROW ' | sed 's/^ROW/ROW window=full/' -done -echo DONE diff --git a/alignment_lab/analysis/rebaseline_panel.sh b/alignment_lab/analysis/rebaseline_panel.sh deleted file mode 100644 index 52e081e9..00000000 --- a/alignment_lab/analysis/rebaseline_panel.sh +++ /dev/null @@ -1,25 +0,0 @@ -#!/bin/bash -# Re-establish the FRF panel after the merge, and settle the hydrogen question in -# the same pass. Every previously published FRF number was measured on -# hydrogen-free structure factors; dev now keeps hydrogens by default, so the old -# 98/100 is not a baseline any more. -#SBATCH --job-name=rebase -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-9 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname) pdb=$PDB" -"$PY" -u alignment_lab/analysis/panel_ranks.py --pdb "$PDB" --trials 10 --tag withH 2>/dev/null | grep '^ROW ' -"$PY" -u alignment_lab/analysis/panel_ranks.py --pdb "$PDB" --trials 10 --tag noH --exclude-h 2>/dev/null | grep '^ROW ' diff --git a/alignment_lab/analysis/refactor_gate.sh b/alignment_lab/analysis/refactor_gate.sh deleted file mode 100644 index 8065bfdc..00000000 --- a/alignment_lab/analysis/refactor_gate.sh +++ /dev/null @@ -1,21 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=refgate -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -export TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 -echo "== import smoke ==" -"$PY" -m pytest -c tests/pytest.ini tests/unit/test_imports_smoke.py --run-slow -q 2>&1 | tail -5 -echo "SMOKE_RC=${PIPESTATUS[0]}" -echo "== alignment + frf_separate + scaling ==" -"$PY" -m pytest -c tests/pytest.ini tests/unit/alignment tests/unit/frf_separate tests/unit/scaling --run-slow -q 2>&1 | tail -20 -echo "SCOPED_RC=${PIPESTATUS[0]}" diff --git a/alignment_lab/analysis/repeat_stability.sh b/alignment_lab/analysis/repeat_stability.sh deleted file mode 100644 index 17eacca2..00000000 --- a/alignment_lab/analysis/repeat_stability.sh +++ /dev/null @@ -1,47 +0,0 @@ -#!/bin/bash -# Does calling the rotation function repeatedly on one model give the same -# answer? The compiled/eager A/B reported a different truth rank from every -# earlier run, and the only thing it did differently was reuse the model across -# many calls -- so either something accumulates on the model, or the rank is -# less stable than measured. -#SBATCH --job-name=frf_repeat -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 -export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname)" -"$PY" -u -c " -import sys -sys.path.insert(0,'alignment_lab') -import torch; torch.set_grad_enabled(False) -from lab import FRFConfig, orbit_rank, rotated_case, run_frf, seed_for - -def rank_of(r, d, R): - k, a = orbit_rank(r.peaks, R, d.spacegroup.matrices.to(torch.float64).cpu(), - reciprocal_basis=d.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), - side='left', frame='cart') - return k, a - -for pdb in ('1DAW', '3K7M'): - cfg = FRFConfig(lmax_cap=64, n_peaks=200) - # A: one model, six calls. - m, d, R = rotated_case(pdb, seed_for(pdb, 0)) - reused = [] - for i in range(6): - reused.append(rank_of(run_frf(m, d, cfg, capture_arf=False), d, R)) - # B: a freshly built model for each call -- the control. - fresh = [] - for i in range(6): - m2, d2, R2 = rotated_case(pdb, seed_for(pdb, 0)) - fresh.append(rank_of(run_frf(m2, d2, cfg, capture_arf=False), d2, R2)) - print(f'{pdb} model REUSED : ranks {[k for k,_ in reused]} ' - f'angles {[round(a,3) for _,a in reused]}') - print(f'{pdb} model FRESH : ranks {[k for k,_ in fresh]} ' - f'angles {[round(a,3) for _,a in fresh]}') -" -echo "exit_code=$?" diff --git a/alignment_lab/analysis/run_tests.sh b/alignment_lab/analysis/run_tests.sh deleted file mode 100644 index caef949e..00000000 --- a/alignment_lab/analysis/run_tests.sh +++ /dev/null @@ -1,26 +0,0 @@ -#!/bin/bash -# Run the test suite on a compute node. The login node is shared and slow enough -# that a cold import alone can take minutes; and scripts a job needs must live on -# /das, not in a node-local /tmp scratch directory the compute node cannot see. -# -# sbatch --partition=hour --time=00:55:00 --cpus-per-task=8 --mem=32G \ -# --constraint=cpu_epyc9335 alignment_lab/analysis/run_tests.sh -# sbatch ... alignment_lab/analysis/run_tests.sh --run-slow -#SBATCH --job-name=frf_tests -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO="${FRF_TEST_REPO:-/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement}" -PY="$REPO/.dev/bin/python" -[ -x "$PY" ] || PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" -export TORCHREF_NUM_THREADS="${SLURM_CPUS_PER_TASK:-4}" -export OMP_NUM_THREADS="$TORCHREF_NUM_THREADS" MKL_NUM_THREADS="$TORCHREF_NUM_THREADS" -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname) sha=$(git -C "$REPO" rev-parse --short HEAD) threads=$TORCHREF_NUM_THREADS" -"$PY" -m pytest tests/unit alignment_lab/tests -q "$@" -rc=$? -echo "exit_code=$rc" -exit "$rc" diff --git a/alignment_lab/analysis/run_tests_copyfix.sh b/alignment_lab/analysis/run_tests_copyfix.sh deleted file mode 100644 index 2fa5742e..00000000 --- a/alignment_lab/analysis/run_tests_copyfix.sh +++ /dev/null @@ -1,35 +0,0 @@ -#!/bin/bash -# Test gate for the build_grid change. rc is captured immediately after pytest, -# not after a pipe: $? on a pipeline reads the LAST command, which silently -# reports success for a failed test run. -#SBATCH --job-name=frf_tests -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname)" - -echo "###### fast: model + alignment + frf_separate" -"$PY" -m pytest -q --no-header -p no:cacheprovider \ - tests/unit/model tests/unit/alignment tests/unit/frf_separate \ - > alignment_lab/slurm/_t_fast.log 2>&1 -rc_fast=$? -tail -4 alignment_lab/slurm/_t_fast.log -echo "rc_fast=$rc_fast" - -echo "###### slow-included: alignment + frf_separate" -"$PY" -m pytest -q --no-header -p no:cacheprovider --run-slow \ - tests/unit/alignment tests/unit/frf_separate \ - > alignment_lab/slurm/_t_slow.log 2>&1 -rc_slow=$? -tail -4 alignment_lab/slurm/_t_slow.log -echo "rc_slow=$rc_slow" - -echo "###### failures, if any" -grep -E "^(FAILED|ERROR)" alignment_lab/slurm/_t_fast.log alignment_lab/slurm/_t_slow.log || echo " none" -echo "done" diff --git a/alignment_lab/analysis/scaler_cost.sh b/alignment_lab/analysis/scaler_cost.sh deleted file mode 100644 index 0380956b..00000000 --- a/alignment_lab/analysis/scaler_cost.sh +++ /dev/null @@ -1,17 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=scalercost -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -"$PY" -u alignment_lab/diagnostics/scaler_cost.py 1DAW 2DQ6 2>/dev/null | grep -E '^ROW|rror' -echo DONE diff --git a/alignment_lab/analysis/scaling_gate.sh b/alignment_lab/analysis/scaling_gate.sh deleted file mode 100644 index 25f74e9e..00000000 --- a/alignment_lab/analysis/scaling_gate.sh +++ /dev/null @@ -1,18 +0,0 @@ -#!/bin/bash -# Deliverable 1 gate: the Chebyshev extraction must be inert. -#SBATCH --job-name=scalegate -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -export TORCHREF_NUM_THREADS=8 OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 -"$PY" -m pytest -c tests/pytest.ini tests/unit/scaling tests/unit/refinement --run-slow -q 2>&1 | tail -8 -echo "PYTEST_RC=${PIPESTATUS[0]}" diff --git a/alignment_lab/analysis/soln_tables.sh b/alignment_lab/analysis/soln_tables.sh deleted file mode 100644 index 2052cd9c..00000000 --- a/alignment_lab/analysis/soln_tables.sh +++ /dev/null @@ -1,20 +0,0 @@ -#!/bin/bash -# Full candidate tables for two cells, to check the shortlist-depth summary by eye. -#SBATCH --job-name=soln -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:20:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -for P in 2DQ6 3K7M 6G9X; do - "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb $P --trial 0 --arms llg --verbose 2 2>/dev/null | grep -E "^===|SOLN|^ROW" -done -echo DONE diff --git a/alignment_lab/analysis/stage_profile.sh b/alignment_lab/analysis/stage_profile.sh deleted file mode 100644 index 1cd3ae30..00000000 --- a/alignment_lab/analysis/stage_profile.sh +++ /dev/null @@ -1,31 +0,0 @@ -#!/bin/bash -# Where does an alignment's wall clock actually go? -# -# Expectation from the component measurements: FRF 1-3 s, and the translation -# stage bounded by one structure-factor call per candidate at well under 100 ms, -# so 25 candidates should be a couple of seconds. That predicts ~5 s and the -# panel medians are 34 s by R and 75-93 s by likelihood. The per-stage timer has -# been in the pipeline the whole time; this reads it. -#SBATCH --job-name=stageprof -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=day -#SBATCH --time=03:00:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -for ARM in analytic_r llg; do - for P in 1DAW 2DQ6; do - echo "########## $P rank_by=$ARM ##########" - "$PY" -u alignment_lab/diagnostics/pose_recovery.py --pdb $P --trial 0 \ - --arms $ARM --verbose 2 2>/dev/null \ - | sed -n '/^stage /,/^TOTAL/p;/^ROW /p' - done -done -echo DONE diff --git a/alignment_lab/analysis/stagea_fingerprint.sh b/alignment_lab/analysis/stagea_fingerprint.sh deleted file mode 100644 index 036002e7..00000000 --- a/alignment_lab/analysis/stagea_fingerprint.sh +++ /dev/null @@ -1,53 +0,0 @@ -#!/bin/bash -# Stage A gate: peak-list fingerprints must be bit-identical between the -# baseline worktree (d1244c45) and this one. SINGLE-THREADED -- at the default -# thread count the float32 SF reduction reorders ~12 of 500 peaks on 3K7M. -#SBATCH --job-name=stagea_fp -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=32G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -OLD=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/_stagea_baseline -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -SCRIPT=$NEW/alignment_lab/analysis/frf_fingerprint.py -OUT=$NEW/alignment_lab/slurm -export TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" - -status=0 -for pdb in 3K7M 1DAW; do - for cap in 64 100; do - for tree in OLD NEW; do - eval "root=\$$tree" - cd "$root" - PYTHONPATH="$root" "$PY" -u "$SCRIPT" --pdb "$pdb" --lmax-cap "$cap" \ - > "$OUT/fp_${tree}_${pdb}_${cap}.txt" 2> "$OUT/fp_${tree}_${pdb}_${cap}.err" - rc=$? - if [ $rc -ne 0 ]; then - echo "RUN FAILED tree=$tree pdb=$pdb cap=$cap rc=$rc" - tail -6 "$OUT/fp_${tree}_${pdb}_${cap}.err" - status=1 - fi - done - a="$OUT/fp_OLD_${pdb}_${cap}.txt"; b="$OUT/fp_NEW_${pdb}_${cap}.txt" - if [ -s "$a" ] && [ -s "$b" ]; then - if diff -q <(grep -v '^#' "$a") <(grep -v '^#' "$b") >/dev/null; then - echo "IDENTICAL $pdb cap$cap ($(grep -vc '^#' "$a") peaks)" - else - nd=$(diff <(grep -v '^#' "$a") <(grep -v '^#' "$b") | grep -c '^[<>]') - echo "DIFFERS $pdb cap$cap ($nd differing lines)" - diff <(grep -v '^#' "$a") <(grep -v '^#' "$b") | head -8 - status=1 - fi - else - echo "MISSING $pdb cap$cap"; status=1 - fi - done -done -echo "fingerprint_status=$status" diff --git a/alignment_lab/analysis/stagea_tests.sh b/alignment_lab/analysis/stagea_tests.sh deleted file mode 100644 index baa0cd53..00000000 --- a/alignment_lab/analysis/stagea_tests.sh +++ /dev/null @@ -1,32 +0,0 @@ -#!/bin/bash -# Stage A: the alignment + frf_separate suites, fast then --run-slow. -#SBATCH --job-name=stagea_tests -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=32G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname)" - -LOG=alignment_lab/slurm/stagea_fast_$SLURM_JOB_ID.log -"$PY" -m pytest tests/unit/alignment tests/unit/frf_separate tests/unit/model \ - tests/unit/test_imports_smoke.py -q > "$LOG" 2>&1 -rc_fast=$? -echo "=== FAST rc=$rc_fast ===" -tail -25 "$LOG" - -LOG2=alignment_lab/slurm/stagea_slow_$SLURM_JOB_ID.log -"$PY" -m pytest --run-slow tests/unit/alignment tests/unit/frf_separate \ - tests/integration/alignment -q > "$LOG2" 2>&1 -rc_slow=$? -echo "=== SLOW rc=$rc_slow ===" -tail -25 "$LOG2" -echo "rc_fast=$rc_fast rc_slow=$rc_slow" diff --git a/alignment_lab/analysis/stagea_verify.sh b/alignment_lab/analysis/stagea_verify.sh deleted file mode 100644 index 2bf6d417..00000000 --- a/alignment_lab/analysis/stagea_verify.sh +++ /dev/null @@ -1,146 +0,0 @@ -#!/bin/bash -# Stage A verification: every swap is claimed bit-identical, so check each one by -# computing BOTH forms in one process and comparing exactly. Stronger and far -# cheaper than diffing an end-to-end fingerprint against a baseline worktree. -# -# Also answers the two "verify, then delete" questions: do the calc-side -# resolution mask and the near-no-op obs mask drop any reflection at all? -#SBATCH --job-name=stagea -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:50:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=32G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=1 -export OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" - -"$PY" -u - <<'PYEOF' -import math -import torch -torch.manual_seed(0) - -ok = True -def check(name, cond, detail=""): - global ok - ok = ok and bool(cond) - print(f" [{'PASS' if cond else 'FAIL'}] {name}{(' -- ' + detail) if detail else ''}") - -print("=== 1. Edmonds ZYZ: base primitive vs the deleted local copy ===") -from torchref.base.alignment.rotation import rotation_matrix_euler_zyz -a = torch.rand(5000, dtype=torch.float64) * 2 * math.pi -b = torch.rand(5000, dtype=torch.float64) * math.pi -g = torch.rand(5000, dtype=torch.float64) * 2 * math.pi - -def old_zyz(alpha, beta, gamma): - ca, sa = torch.cos(alpha), torch.sin(alpha) - cb, sb = torch.cos(beta), torch.sin(beta) - cg, sg = torch.cos(gamma), torch.sin(gamma) - return torch.stack([ - torch.stack([ca*cb*cg - sa*sg, -ca*cb*sg - sa*cg, ca*sb], dim=-1), - torch.stack([sa*cb*cg + ca*sg, -sa*cb*sg + ca*cg, sa*sb], dim=-1), - torch.stack([-sb*cg, sb*sg, cb ], dim=-1), - ], dim=-2) - -R_old = old_zyz(a, b, g) -R_new = rotation_matrix_euler_zyz(torch.stack([a, b, g], dim=-1)) -check("bitwise equal over 5000 random triples", torch.equal(R_old, R_new), - f"max|d|={ (R_old-R_new).abs().max().item():.3e}") - -print("=== 2. Rodrigues: rotation_utils vs the deleted align._rodrigues ===") -from torchref.experimental.alignment.frf.rotation_utils import axis_angle_to_matrix - -def old_rodrigues(omega): - if omega.dtype != torch.float64: - omega = omega.to(torch.float64) - single = omega.dim() == 1 - if single: - omega = omega.unsqueeze(0) - th = omega.norm(dim=-1, keepdim=True) - axis = omega / th.clamp(min=1e-30) - zeros = torch.zeros_like(axis[..., 0]) - K = torch.stack([ - torch.stack([zeros, -axis[..., 2], axis[..., 1]], dim=-1), - torch.stack([axis[..., 2], zeros, -axis[..., 0]], dim=-1), - torch.stack([-axis[..., 1], axis[..., 0], zeros], dim=-1), - ], dim=-2) - th_b = th.unsqueeze(-1) - eye = torch.eye(3, dtype=omega.dtype, device=omega.device).expand(*omega.shape[:-1], 3, 3) - R = eye + torch.sin(th_b) * K + (1.0 - torch.cos(th_b)) * torch.matmul(K, K) - return R.squeeze(0) if single else R - -# The production caller builds omegas as a float64 meshgrid, so replicate that. -c = torch.linspace(-0.1, 0.1, 11, dtype=torch.float64) -wx, wy, wz = torch.meshgrid(c, c, c, indexing="ij") -om = torch.stack([wx.flatten(), wy.flatten(), wz.flatten()], dim=-1) -check("bitwise equal on the pipeline's float64 perturbation grid", - torch.equal(old_rodrigues(om), axis_angle_to_matrix(om))) - -print("=== 3. Symop unroll: apply_to_hkl vs the deleted einsum ===") -from torchref.symmetry.spacegroup import SpaceGroup -for sg_name in ("P 1", "C 2", "P 21 21 21", "P 31 2 1", "P 65 2 2", "P 43 32", "P 4 3 2"): - sg = SpaceGroup(sg_name) - hkl = torch.randint(-40, 41, (4000, 3)) - rec = torch.eye(3, dtype=torch.float64) * 0.0137 + 0.0011 # arbitrary non-diagonal basis - old = torch.einsum( - "kji,nj->kni", sg.matrices.to(torch.float64), hkl.to(torch.float64) - ).reshape(-1, 3) - new = sg.apply_to_hkl(hkl).permute(2, 0, 1).reshape(-1, 3).to(torch.float64) - same_rows = torch.equal(old, new) - same_s = torch.equal(old @ rec, new @ rec) - check(f"{sg_name:12s} n_ops={sg.n_ops:2d} rows and s_obs bitwise equal", - same_rows and same_s, - "" if same_rows else f"row mismatch {int((old != new).any(-1).sum())}") - -print("=== 4. (L, d_min) pass-through replaces the second auto_lmax call ===") -from torchref.experimental.alignment.frf.api import phaser_lmax_resolution -from torchref.experimental.alignment.rotation_search import LMAX_CAP -for r, dmin in ((10.0, 1.8), (15.0, 2.05), (25.0, 1.6), (4.0, 3.0)): - L1, d1 = phaser_lmax_resolution(r, dmin, LMAX_CAP) - L2, d2 = phaser_lmax_resolution(r, dmin, LMAX_CAP) # the call that used to be inside - check(f"radius {r:5.1f} d_min {dmin:.2f} -> L={L1} d_min_eff={d1:.4f}", - (L1, d1) == (L2, d2)) - -print("=== 5. Do the two 'verify then delete' masks drop anything? ===") -from alignment_lab.lab.benchmark import load_case -from torchref.experimental.alignment.frf.dense_calc import dense_calc_via_box -from torchref.experimental.alignment.rotation_search import ( - LOW_RESOLUTION_CUTOFF_A, DENSE_CALC_PAD, -) -for pdb in ("3K7M", "1DAW"): - model, data = load_case(pdb)[:2] - rec_basis = data.cell.reciprocal_basis_matrix.to(torch.float64) - s_mag_all = (data.hkl.to(torch.float64) @ rec_basis).norm(dim=-1) - d_min_data = float(1.0 / s_mag_all.max().item()) - d_max = float(LOW_RESOLUTION_CUTOFF_A) - - # (a) the obs mask at rotation_search.py:239 -- upper bound is max <= max - keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min_data) - n_lo = int((s_mag_all < 1.0 / d_max).sum()) - n_hi = int((s_mag_all > 1.0 / d_min_data).sum()) - print(f" {pdb}: obs mask keeps {int(keep.sum())}/{len(s_mag_all)} " - f"(dropped {n_lo} below {d_max:.0f} A, {n_hi} above d_min via the " - f"reciprocal round-trip)") - - # (b) the calc-side mask in score_model, against dense_calc's own window - model_radius_A = float((model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item()) - L, d_min_eff = phaser_lmax_resolution(model_radius_A, d_min_data, LMAX_CAP) - s_calc, F_calc = dense_calc_via_box(model, d_max, d_min_eff, pad=DENSE_CALC_PAD) - smag_calc = s_calc.norm(dim=-1) - keep_c = (smag_calc >= 1.0 / d_max) & (smag_calc <= 1.0 / d_min_eff) - dropped = int(len(smag_calc) - keep_c.sum()) - print(f" {pdb}: calc re-mask drops {dropped}/{len(smag_calc)} " - f"(L={L}, d_min_eff={d_min_eff:.3f} A)") - -print() -print("OVERALL:", "PASS" if ok else "FAIL") -PYEOF -rc=$? -echo "python_exit=$rc" diff --git a/alignment_lab/analysis/stageb_gate.sh b/alignment_lab/analysis/stageb_gate.sh deleted file mode 100644 index 8035d7a6..00000000 --- a/alignment_lab/analysis/stageb_gate.sh +++ /dev/null @@ -1,49 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=stageb -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=32G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$NEW" -export PYTHONPATH="$NEW" TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" - -LOG=alignment_lab/slurm/stageb_bessel_$SLURM_JOB_ID.log -"$PY" -m pytest tests/unit/frf_separate/test_bessel_rescale.py -q > "$LOG" 2>&1 -rc_b=$? -echo "=== BESSEL TESTS rc=$rc_b ===" -tail -20 "$LOG" - -# End-to-end: the rescaled ladder must leave the peak lists untouched. Compare -# against the fingerprints captured for the Stage A gate (same tree, pre-Stage-B). -OUT=$NEW/alignment_lab/slurm -status=0 -for pdb in 3K7M 1DAW; do - for cap in 64 100; do - "$PY" -u "$NEW/alignment_lab/analysis/frf_fingerprint.py" --pdb "$pdb" \ - --lmax-cap "$cap" > "$OUT/fp_STAGEB_${pdb}_${cap}.txt" 2>"$OUT/fp_STAGEB_${pdb}_${cap}.err" - rc=$? - ref="$OUT/fp_NEW_${pdb}_${cap}.txt" - new="$OUT/fp_STAGEB_${pdb}_${cap}.txt" - if [ $rc -ne 0 ]; then - echo "RUN FAILED $pdb cap$cap rc=$rc"; tail -5 "$OUT/fp_STAGEB_${pdb}_${cap}.err"; status=1; continue - fi - if [ ! -s "$ref" ]; then echo "NO STAGE-A REFERENCE for $pdb cap$cap"; status=1; continue; fi - if diff -q <(grep -v '^#' "$ref") <(grep -v '^#' "$new") >/dev/null; then - echo "IDENTICAL $pdb cap$cap ($(grep -vc '^#' "$new") peaks)" - else - nd=$(diff <(grep -v '^#' "$ref") <(grep -v '^#' "$new") | grep -c '^[<>]') - echo "DIFFERS $pdb cap$cap ($nd lines)" - diff <(grep -v '^#' "$ref") <(grep -v '^#' "$new") | head -6 - status=1 - fi - done -done -echo "stageb_status=$status rc_bessel=$rc_b" diff --git a/alignment_lab/analysis/stagec_gate3.sh b/alignment_lab/analysis/stagec_gate3.sh deleted file mode 100644 index e150d87f..00000000 --- a/alignment_lab/analysis/stagec_gate3.sh +++ /dev/null @@ -1,55 +0,0 @@ -#!/bin/bash -# Stage C1, third pass. Hypothesis: with the back-half accumulation restored to -# double, the whole of C1 is numerically neutral -- so the fingerprints should be -# bit-identical to d1244c45, and the 1e-4 score shift seen in pass 2 was entirely -# the narrowed accumulation. -#SBATCH --job-name=stagec3 -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -OLD=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/_stagea_baseline -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -SCRIPT=$NEW/alignment_lab/analysis/frf_fingerprint.py -OUT=$NEW/alignment_lab/slurm -export TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" - -cd "$NEW"; export PYTHONPATH="$NEW" -LOG=$OUT/stagec3_tests_$SLURM_JOB_ID.log -"$PY" -m pytest tests/unit/alignment tests/unit/frf_separate tests/unit/model \ - tests/unit/test_imports_smoke.py -q > "$LOG" 2>&1 -rc=$? -echo "=== TESTS rc=$rc ===" -tail -12 "$LOG" - -status=0 -for pdb in 3K7M 1DAW; do - for cap in 64 100; do - cd "$NEW" - PYTHONPATH="$NEW" "$PY" -u "$SCRIPT" --pdb "$pdb" --lmax-cap "$cap" \ - > "$OUT/c3_NEW_${pdb}_${cap}.txt" 2>/dev/null - a=$OUT/c2_OLD_${pdb}_${cap}.txt # baseline captured in pass 2 - b=$OUT/c3_NEW_${pdb}_${cap}.txt - if diff -q <(grep '^FP ' "$a") <(grep '^FP ' "$b") >/dev/null; then - echo "IDENTICAL $pdb cap$cap ($(grep -c '^FP ' "$b") peaks)" - else - nd=$(diff <(grep '^FP ' "$a") <(grep '^FP ' "$b") | grep -c '^[<>]') - echo "DIFFERS $pdb cap$cap ($nd lines)" - diff <(grep '^FP ' "$a") <(grep '^FP ' "$b") | head -4 - status=1 - fi - done -done - -echo "=== timing, 4 threads for comparability with the 1.08s note ===" -export TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -"$PY" -u -m alignment_lab.diagnostics.frf_benchmark --pdb 3K7M --arms cap64 --trials 2 2>&1 \ - | grep -vE "Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|warn|^ *$" -echo "stagec3_status=$status tests_rc=$rc" diff --git a/alignment_lab/analysis/stagec_timing_paired.sh b/alignment_lab/analysis/stagec_timing_paired.sh deleted file mode 100644 index 0eeffd19..00000000 --- a/alignment_lab/analysis/stagec_timing_paired.sh +++ /dev/null @@ -1,34 +0,0 @@ -#!/bin/bash -# Paired timing: baseline worktree vs this one, same node, same job, INTERLEAVED -# and structure-major. A cross-job comparison against a remembered number is not -# a measurement -- node and contention differ. -#SBATCH --job-name=stagec_time -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --exclusive -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -OLD=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/_stagea_baseline -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -export TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" -echo "OLD=$(cd $OLD && git rev-parse --short HEAD) NEW=working tree" - -for round in 1 2 3; do - for pdb in 3K7M 1DAW; do - for tree in OLD NEW; do - eval "root=\$$tree" - cd "$root" - line=$(PYTHONPATH="$root" "$PY" -u -m alignment_lab.diagnostics.frf_benchmark \ - --pdb "$pdb" --arms cap64 --trials 2 2>/dev/null \ - | grep -E "^ *cap64" | tail -1) - echo "round$round $pdb $tree $line" - done - done -done diff --git a/alignment_lab/analysis/staged_gate.sh b/alignment_lab/analysis/staged_gate.sh deleted file mode 100644 index 186efc26..00000000 --- a/alignment_lab/analysis/staged_gate.sh +++ /dev/null @@ -1,76 +0,0 @@ -#!/bin/bash -# Stage D gate: the obs chain now runs on the unique set and only the geometry -# unrolls, with one shared shell assignment. This CHANGES numbers, so quantify -# how much and confirm truth is still where the pipeline can reach it. Also -# measure what it bought, paired and interleaved on one node. -#SBATCH --job-name=staged2 -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --exclusive -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -OLD=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/_stagea_baseline -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -OUT=$NEW/alignment_lab/slurm -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname) OLD=$(cd $OLD && git rev-parse --short HEAD) NEW=working tree" - -cd "$NEW"; export PYTHONPATH="$NEW" -export TORCHREF_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 -LOG=$OUT/staged_tests_$SLURM_JOB_ID.log -"$PY" -m pytest tests/unit/alignment tests/unit/frf_separate tests/unit/model \ - tests/unit/test_imports_smoke.py -q > "$LOG" 2>&1 -rc=$? -echo "=== TESTS rc=$rc ===" -tail -14 "$LOG" - -echo "=== how far did the peak lists move (single-threaded) ===" -for pdb in 3K7M 1DAW; do - for cap in 64 100; do - for tree in OLD NEW; do - eval "root=\$$tree" - cd "$root" - PYTHONPATH="$root" "$PY" -u "$NEW/alignment_lab/analysis/frf_fingerprint.py" \ - --pdb "$pdb" --lmax-cap "$cap" > "$OUT/d_${tree}_${pdb}_${cap}.txt" 2>/dev/null \ - || echo "RUN FAILED $tree $pdb $cap" - done - "$PY" - "$OUT/d_OLD_${pdb}_${cap}.txt" "$OUT/d_NEW_${pdb}_${cap}.txt" "$pdb" "$cap" <<'PYEOF' -import sys -rp, np_, pdb, cap = sys.argv[1:5] -def load(p): - return [tuple(map(float, l.split()[2:7])) for l in open(p) if l.startswith("FP ")] -a, b = load(rp), load(np_) -if not a or not b: - print(f" {pdb} cap{cap}: EMPTY {len(a)} vs {len(b)}"); sys.exit() -n = min(len(a), len(b)) -slot = sum(1 for i in range(n) if a[i][:3] == b[i][:3]) -# Is the same peak SET recovered, regardless of order? -sa, sb = {r[:3] for r in a}, {r[:3] for r in b} -rel = sorted(abs(b[i][3]-a[i][3])/max(abs(a[i][3]),1e-30) for i in range(n)) -print(f" {pdb} cap{cap}: {slot}/{n} slots identical | set overlap " - f"{len(sa & sb)}/{len(sa)} | |dscore|/score p50={rel[n//2]:.2e} " - f"p99={rel[int(0.99*n)]:.2e}") -print(f" top-1 old {tuple(round(v,6) for v in a[0][:3])} z={a[0][4]:.4f}" - f" new {tuple(round(v,6) for v in b[0][:3])} z={b[0][4]:.4f}") -PYEOF - done -done - -echo "=== paired timing, 4 threads, 3 rounds ===" -export TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -for round in 1 2 3; do - for pdb in 3K7M 1DAW; do - for tree in OLD NEW; do - eval "root=\$$tree"; cd "$root" - line=$(PYTHONPATH="$root" "$PY" -u -m alignment_lab.diagnostics.frf_benchmark \ - --pdb "$pdb" --arms cap64 --trials 2 2>/dev/null | grep -E "^ *cap64" | tail -1) - echo "round$round $pdb $tree $line" - done - done -done -echo "staged_tests_rc=$rc" diff --git a/alignment_lab/analysis/staged_panel.sh b/alignment_lab/analysis/staged_panel.sh deleted file mode 100644 index 19b81583..00000000 --- a/alignment_lab/analysis/staged_panel.sh +++ /dev/null @@ -1,29 +0,0 @@ -#!/bin/bash -# Stage D acceptance panel: 10 benchmark structures x 10 seeded trials, run in -# BOTH worktrees so the comparison is paired per seed. One array task per -# structure; OLD and NEW run back to back inside the task on the same node. -#SBATCH --job-name=d_panel -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-9 -set -uo pipefail -NEW=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -OLD=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/_stagea_baseline -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -export TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname) pdb=$PDB" -for tree in OLD NEW; do - eval "root=\$$tree" - cd "$root" - PYTHONPATH="$root" "$PY" -u "$NEW/alignment_lab/analysis/panel_ranks.py" \ - --pdb "$PDB" --trials 10 --tag "$tree" 2>/dev/null | grep '^ROW ' - echo "tree=$tree rc=$?" -done diff --git a/alignment_lab/analysis/sweep_chunk.sh b/alignment_lab/analysis/sweep_chunk.sh deleted file mode 100644 index 830b2d15..00000000 --- a/alignment_lab/analysis/sweep_chunk.sh +++ /dev/null @@ -1,54 +0,0 @@ -#!/bin/bash -# Sweep the SH-Bessel expansion's chunk width. The loop body changed -- it is now -# elementwise plus a scatter, with no GEMM -- so an earlier conclusion that -# narrow chunks hurt no longer applies and the width has to be re-measured. -# Narrow chunks keep the recurrence's three rolling rows in cache; wide ones cut -# Python and dispatch overhead. -# -# sbatch --partition=hour --time=00:55:00 --exclusive --mem=0 \ -# --constraint=cpu_epyc9335 alignment_lab/analysis/sweep_chunk.sh -#SBATCH --job-name=frf_chunk -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 -export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" -"$PY" -u -c " -import sys, time -sys.path.insert(0,'alignment_lab') -import torch; torch.set_grad_enabled(False) -import torchref.experimental.alignment.frf.data_mr as dm -from lab import FRFConfig, orbit_rank, rotated_case, run_frf, seed_for - -BUDGETS_MB = [2, 8, 32, 128, 256, 1024] -cases = {p: rotated_case(p, seed_for(p, 0)) for p in ('1DAW', '3K7M')} -for cap in (64, 100): - cfg = FRFConfig(lmax_cap=cap, n_peaks=200) - for m, d, _ in cases.values(): - run_frf(m, d, cfg, capture_arf=False) # warm up - res = {b: {p: [] for p in cases} for b in BUDGETS_MB} - rank = {} - for rep in range(2): - for b in BUDGETS_MB: # interleaved - dm.CLUSTER_CHUNK_BYTES = b * 1_000_000 - for p, (m, d, R) in cases.items(): - t0 = time.perf_counter() - r = run_frf(m, d, cfg, capture_arf=False) - res[b][p].append(time.perf_counter() - t0) - k, _ = orbit_rank(r.peaks, R, - d.spacegroup.matrices.to(torch.float64).cpu(), - reciprocal_basis=d.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), - side='left', frame='cart') - rank[(b, p)] = k - print(f'--- cap{cap} (best of 2, seconds; rank in brackets) ---') - print(' ' + 'chunk MB'.rjust(9) + ''.join(f'{p:>18}' for p in cases)) - for b in BUDGETS_MB: - cells = ''.join(f'{min(res[b][p]):12.2f} [{rank[(b,p)]:>3}]' for p in cases) - print(f' {b:9d}{cells}') -" -echo "exit_code=$?" diff --git a/alignment_lab/analysis/tf_resolution_and_grid.sh b/alignment_lab/analysis/tf_resolution_and_grid.sh deleted file mode 100644 index cf3b7ef6..00000000 --- a/alignment_lab/analysis/tf_resolution_and_grid.sh +++ /dev/null @@ -1,19 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=tfgrid -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=96G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -cd "$REPO" -for P in 1DAW 3E98 3A5V 3GR5 1AK5 3K7M 2DQ6 4BX9 6G9X; do - "$PY" -u alignment_lab/diagnostics/tf_resolution_and_grid.py --pdb "$P" --threads 4 2>&1 \ - | grep -E "^#|^ROW|Traceback|Error" || echo "# $P FAILED" -done diff --git a/alignment_lab/analysis/tf_sf_backend.sh b/alignment_lab/analysis/tf_sf_backend.sh deleted file mode 100644 index a71b325b..00000000 --- a/alignment_lab/analysis/tf_sf_backend.sh +++ /dev/null @@ -1,19 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=sfback -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -export PYTHONPATH="$REPO" PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 -cd "$REPO" -for P in 1DAW 3E98 3A5V 3GR5 1AK5 3K7M 2DQ6 4BX9 6G9X; do - "$PY" -u alignment_lab/diagnostics/tf_sf_backend.py --pdb "$P" 2>&1 \ - | grep -E "^#|^ROW|Error|Traceback" || echo "ROW pdb=$P FAILED" -done diff --git a/alignment_lab/analysis/trigonal_metric_recheck.sh b/alignment_lab/analysis/trigonal_metric_recheck.sh deleted file mode 100644 index 62e8043e..00000000 --- a/alignment_lab/analysis/trigonal_metric_recheck.sh +++ /dev/null @@ -1,36 +0,0 @@ -#!/bin/bash -# Were 2DQ6's "failures" symmetry mates the success metric could not recognise? -# -# residual_rotation_deg compared a Cartesian Kabsch rotation against the -# FRACTIONAL symmetry matrices. In P3(1)21 two of the six mates of a correct -# solution then read as 30.00 and 21.09 deg; 2DQ6's failing residuals were -# 28.5-31.0 and 19.5-21.8. This re-scores the same seeds with Cartesian mates -# and a translation check modulo allowed origin shifts. Controls: 3GR5 (P6(5)22, -# four of twelve mates affected), 6G9X t4 corr (a genuine 55.9 deg miss, must -# stay a miss) and 1DAW t0 (monoclinic, must be unchanged at 1.519). -#SBATCH --job-name=trig -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=hour -#SBATCH --time=00:59:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-3 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 -export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -run() { "$PY" -u alignment_lab/diagnostics/pose_recovery.py "$@" 2>/dev/null | grep -E '^ROW |SOLN|^===' ; } -case $SLURM_ARRAY_TASK_ID in - 0) for T in $(seq 0 9); do run --pdb 2DQ6 --trial $T --arms llg; done ;; - 1) for T in $(seq 0 9); do run --pdb 2DQ6 --trial $T --arms analytic_r; done ;; - 2) for T in $(seq 0 9); do run --pdb 2DQ6 --trial $T --arms corr; done ;; - 3) for T in 0 1 2; do run --pdb 3GR5 --trial $T --arms llg,analytic_r,corr; done - run --pdb 6G9X --trial 4 --arms corr - run --pdb 1DAW --trial 0 --arms llg - run --pdb 2DQ6 --trial 3 --arms llg --verbose 2 ;; -esac -echo DONE diff --git a/alignment_lab/analysis/truth_metric_check.sh b/alignment_lab/analysis/truth_metric_check.sh deleted file mode 100644 index 892b7e1c..00000000 --- a/alignment_lab/analysis/truth_metric_check.sh +++ /dev/null @@ -1,17 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=truthchk -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -for P in 2DQ6 1DAW; do "$PY" -u alignment_lab/diagnostics/truth_metric_check.py $P 0 2>/dev/null; done -echo DONE diff --git a/alignment_lab/analysis/truth_pose_scores.sh b/alignment_lab/analysis/truth_pose_scores.sh deleted file mode 100644 index 3a74c51b..00000000 --- a/alignment_lab/analysis/truth_pose_scores.sh +++ /dev/null @@ -1,22 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=tps -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=hour -#SBATCH --time=00:30:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-2 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -case $SLURM_ARRAY_TASK_ID in - 0) "$PY" -u alignment_lab/diagnostics/truth_pose_scores.py --pdb 2DQ6 --trial 3 ;; - 1) "$PY" -u alignment_lab/diagnostics/truth_pose_scores.py --pdb 2DQ6 --trial 0 ;; - 2) "$PY" -u alignment_lab/diagnostics/truth_pose_scores.py --pdb 3GR5 --trial 0 ;; -esac -echo DONE diff --git a/alignment_lab/analysis/truth_pose_scores2.sh b/alignment_lab/analysis/truth_pose_scores2.sh deleted file mode 100644 index 44a910e6..00000000 --- a/alignment_lab/analysis/truth_pose_scores2.sh +++ /dev/null @@ -1,25 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=tps2 -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=hour -#SBATCH --time=00:40:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=64G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-4 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=8 -export OMP_NUM_THREADS=8 MKL_NUM_THREADS=8 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -T="$PY -u alignment_lab/diagnostics/truth_pose_scores.py" -case $SLURM_ARRAY_TASK_ID in - 0) $T --pdb 2DQ6 --trial 3 --tf-d-min 4.0 --tf-d-max 15.0 ;; - 1) $T --pdb 2DQ6 --trial 3 --tf-d-min 3.0 ;; - 2) $T --pdb 4BX9 --trial 0 ;; - 3) $T --pdb 6G9X --trial 0 ;; - 4) $T --pdb 1DAW --trial 0 ;; -esac -echo DONE diff --git a/alignment_lab/analysis/verify_copy_correctness.sh b/alignment_lab/analysis/verify_copy_correctness.sh deleted file mode 100644 index 7aec9231..00000000 --- a/alignment_lab/analysis/verify_copy_correctness.sh +++ /dev/null @@ -1,80 +0,0 @@ -#!/bin/bash -# Correctness half of the build_grid=False change. No timings here: this runs -# wherever there is a free slot, and runtime numbers are only comparable on the -# pinned EPYC 9335. What matters is that |F_calc| is bit-identical. -#SBATCH --job-name=frf_vcorr -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -FILT="Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$" -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" - -echo "=== |F_calc| from the box path, grid-building vs skipped ===" -"$PY" -u -c " -import math, torch -torch.set_grad_enabled(False) -from alignment_lab.lab.benchmark import load_case -from torchref.symmetry.cell import Cell -from torchref.experimental.alignment.frf.dense_calc import model_sf_abs -from torchref.experimental.alignment.frf.api import phaser_lmax_resolution - -def box_path(model, d_min, d_max, build_grid): - m = model.copy(build_grid=build_grid) - coords = m.xyz(); dev = coords.device - extent = (coords - coords.mean(0)).norm(dim=-1).max().item() - a = float(2.0 * 2.0 * extent) - m.max_res = float(d_min); m.spacegroup = 'P 1' - m.cell = Cell([a, a, a, 90., 90., 90.], device=dev) - nmax = int(math.ceil(a / d_min)) - idx = torch.arange(-nmax, nmax + 1, device=dev) - H, K, Lg = torch.meshgrid(idx, idx, idx, indexing='ij') - hkl = torch.stack([H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], -1).to(torch.long) - smag = hkl.to(torch.float64).norm(dim=-1) / a - hkl = hkl[(smag >= 1.0/d_max) & (smag <= 1.0/d_min)].contiguous() - return model_sf_abs(m, hkl) - -bad = 0 -for name in ('3K7M', '1DAW'): - model, data = load_case(name); model.verbose = 0 - rb = data.cell.reciprocal_basis_matrix.to(torch.float64) - d_min_data = 1.0 / (data.hkl.to(torch.float64) @ rb).norm(dim=-1).max().item() - r = float((model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item()) - for cap in (64, 100): - L, d_min = phaser_lmax_resolution(r, d_min_data, cap) - a = box_path(model, d_min, 100.0, False) - b = box_path(model, d_min, 100.0, True) - ok = torch.equal(a, b) - bad += 0 if ok else 1 - md = 0.0 if ok else (a-b).abs().max().item() - print(f' {name} cap{cap}: bit-identical={ok} n={a.numel()} max|dF|={md:.3e}') -print('MISMATCHES:', bad) -" 2>&1 | grep -vE "$FILT" - -echo "=== the caller must not be mutated ===" -"$PY" -u -c " -import torch -torch.set_grad_enabled(False) -from alignment_lab.lab.benchmark import load_case -from torchref.experimental.alignment.frf.dense_calc import dense_calc_via_box -model, data = load_case('1DAW'); model.verbose = 0 -before = (str(model.spacegroup), model.max_res, model.cell.data.clone(), - model.xyz().clone()) -dense_calc_via_box(model, 100.0, 4.0, pad=2.0) -after = (str(model.spacegroup), model.max_res, model.cell.data.clone(), - model.xyz().clone()) -print(' spacegroup preserved:', before[0] == after[0], before[0]) -print(' max_res preserved :', before[1] == after[1], before[1]) -print(' cell preserved :', torch.equal(before[2], after[2])) -print(' coords preserved :', torch.equal(before[3], after[3])) -print(' grid still present :', model._fft.real_space_grid is not None) -" 2>&1 | grep -vE "$FILT" - -echo "=== tests ===" -"$PY" -u -m pytest -q tests/unit/model tests/unit/alignment tests/unit/frf_separate 2>&1 | tail -12 -echo "exit_code=$?" diff --git a/alignment_lab/analysis/verify_copy_fix.sh b/alignment_lab/analysis/verify_copy_fix.sh deleted file mode 100644 index d88b0572..00000000 --- a/alignment_lab/analysis/verify_copy_fix.sh +++ /dev/null @@ -1,79 +0,0 @@ -#!/bin/bash -# Verify the build_grid=False change in dense_calc_via_box: -# 1. |F_calc| must be bit-identical to the grid-building path. -# 2. Truth ranks must not move. -# 3. Report the new stage table, per cap (cap64 is what ships). -#SBATCH --job-name=frf_vcopy -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -FILT="Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$" -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" - -echo "=== 1. |F_calc| identical, and the copy cost ===" -"$PY" -u -c " -import math, time, torch -torch.set_grad_enabled(False) -from alignment_lab.lab.benchmark import load_case -from torchref.symmetry.cell import Cell -from torchref.experimental.alignment.frf.dense_calc import dense_calc_via_box, model_sf_abs -from torchref.experimental.alignment.frf.api import phaser_lmax_resolution - -def box_path(model, d_min, d_max, build_grid): - '''dense_calc_via_box, with the grid build under our control.''' - m = model.copy(build_grid=build_grid) - coords = m.xyz(); dev = coords.device - extent = (coords - coords.mean(0)).norm(dim=-1).max().item() - a = float(2.0 * 2.0 * extent) - m.max_res = float(d_min); m.spacegroup = 'P 1' - m.cell = Cell([a, a, a, 90., 90., 90.], device=dev) - nmax = int(math.ceil(a / d_min)) - idx = torch.arange(-nmax, nmax + 1, device=dev) - H, K, Lg = torch.meshgrid(idx, idx, idx, indexing='ij') - hkl = torch.stack([H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], -1).to(torch.long) - smag = hkl.to(torch.float64).norm(dim=-1) / a - hkl = hkl[(smag >= 1.0/d_max) & (smag <= 1.0/d_min)].contiguous() - return model_sf_abs(m, hkl) - -def timed(fn, n=3): - fn() - return min((lambda: (lambda t0: (fn(), time.perf_counter()-t0)[1])(time.perf_counter()))() for _ in range(n)) - -for name in ('3K7M', '1DAW'): - model, data = load_case(name) - model.verbose = 0 - rb = data.cell.reciprocal_basis_matrix.to(torch.float64) - d_min_data = 1.0 / (data.hkl.to(torch.float64) @ rb).norm(dim=-1).max().item() - r = float((model.xyz() - model.xyz().mean(0)).norm(dim=-1).mean().item()) - L, d_min = phaser_lmax_resolution(r, d_min_data, 64) - f_lean = box_path(model, d_min, 100.0, False) - f_full = box_path(model, d_min, 100.0, True) - same = torch.equal(f_lean, f_full) - md = (f_lean - f_full).abs().max().item() if not same else 0.0 - t_lean = timed(lambda: model.copy(build_grid=False)) - t_full = timed(lambda: model.copy(build_grid=True)) - t_stage = timed(lambda: dense_calc_via_box(model, 100.0, d_min, pad=2.0)) - print(f' {name}: bit-identical={same} (max|dF|={md:.3e}, n={f_lean.numel()})') - print(f' model.copy(build_grid=True) {t_full*1e3:8.1f} ms') - print(f' model.copy(build_grid=False) {t_lean*1e3:8.1f} ms') - print(f' dense_calc_via_box now {t_stage*1e3:8.1f} ms') -" 2>&1 | grep -vE "$FILT" - -echo "=== 2/3. ranks and the new stage table, per cap ===" -for pdb in 3K7M 1DAW; do - for arm in cap64 cap100; do - "$PY" -u -m alignment_lab.diagnostics.frf_benchmark \ - --pdb "$pdb" --arms "$arm" --trials 2 2>&1 | grep -vE "$FILT" - done -done - -echo "=== 4. tests ===" -"$PY" -u -m pytest -q tests/unit/model/test_copy.py tests/unit/alignment \ - tests/unit/frf_separate 2>&1 | tail -15 -echo "exit_code=$?" diff --git a/alignment_lab/analysis/verify_frf.sh b/alignment_lab/analysis/verify_frf.sh deleted file mode 100644 index 87ad93a3..00000000 --- a/alignment_lab/analysis/verify_frf.sh +++ /dev/null @@ -1,43 +0,0 @@ -#!/bin/bash -# Verify a change to the rotation function: tests, then the numbers that decide -# whether the change was worth it. Runs on a compute node with the CPU pinned -- -# the login node is contended enough that a profile taken there sent one earlier -# optimisation after the wrong stage. -# -# sbatch --partition=hour --time=00:55:00 --exclusive --mem=0 \ -# --constraint=cpu_epyc9335 alignment_lab/analysis/verify_frf.sh -#SBATCH --job-name=frf_verify -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO="${FRF_VERIFY_REPO:-/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement}" -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 -export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" -echo "sha=$(git -C "$REPO" rev-parse --short HEAD)" - -echo "=== unit tests (the expansion's own invariants included) ===" -"$PY" -m pytest tests/unit/alignment tests/unit/frf_separate alignment_lab/tests -q -rc=$? - -echo "=== stage profile and truth rank, cap 64 and cap 100, steady state ===" -FRF_PROFILE=1 "$PY" -u -c " -import sys; sys.path.insert(0,'alignment_lab') -import torch; torch.set_grad_enabled(False) -from lab import FRFConfig, orbit_rank, rotated_case, run_frf, seed_for -for cap in (64, 100): - for pdb in ('1DAW','3K7M'): - m,d,R = rotated_case(pdb, seed_for(pdb,0)) - cfg = FRFConfig(lmax_cap=cap, n_peaks=200) - run_frf(m, d, cfg, capture_arf=False) # warm up - r = run_frf(m, d, cfg, capture_arf=False) - rank, ang = orbit_rank(r.peaks, R, d.spacegroup.matrices.to(torch.float64).cpu(), - reciprocal_basis=d.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), - side='left', frame='cart') - print(f'>>> cap{cap} {pdb} {r.seconds:.2f}s truth rank {rank} at {ang:.3f} deg') -" 2>&1 | grep -E "FRF_PROFILE|>>>" -echo "exit_tests=$rc" -exit "$rc" diff --git a/alignment_lab/analysis/verify_wigner_memo.sh b/alignment_lab/analysis/verify_wigner_memo.sh deleted file mode 100644 index 7df9fcde..00000000 --- a/alignment_lab/analysis/verify_wigner_memo.sh +++ /dev/null @@ -1,71 +0,0 @@ -#!/bin/bash -# The memoised small-d blocks: same answer, and how much the second search in a -# process saves. Also the peak-memory cost of holding the blocks. -#SBATCH --job-name=frf_wmemo -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -FILT="Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$|No CUDA" -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" - -echo "=== contraction: cold vs memoised, and identity ===" -"$PY" -u -c " -import math, resource, time, torch -torch.set_grad_enabled(False) -from torchref.experimental.alignment.frf.wigner_d import ( - wigner_contraction_per_beta, clear_wigner_d_cache) - -def rss_mb(): - return resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / 1024.0 - -for cap in (64, 100): - L = cap + 1 - n_beta = int(math.ceil(180.0 / 3.0)) - betas = torch.arange(n_beta, dtype=torch.float64) * 3.0 * (math.pi/180.0) - g = torch.Generator().manual_seed(0) - xi = torch.randn(L, 2*L-1, 2*L-1, generator=g, dtype=torch.float64).to(torch.complex128) - clear_wigner_d_cache() - r0 = rss_mb() - t0 = time.perf_counter(); a = wigner_contraction_per_beta(xi, betas); cold = time.perf_counter()-t0 - r1 = rss_mb() - warm = min([(lambda t: (wigner_contraction_per_beta(xi, betas), time.perf_counter()-t)[1])(time.perf_counter()) for _ in range(3)]) - b = wigner_contraction_per_beta(xi, betas) - print(f' cap{cap} L={L}: cold {cold*1e3:7.1f} ms memoised {warm*1e3:6.1f} ms ' - f'({cold/warm:.1f}x) identical={torch.equal(a, b)} ' - f'RSS +{r1-r0:.0f} MB') - clear_wigner_d_cache() -" 2>&1 | grep -vE "$FILT" - -echo "=== two searches in one process ===" -"$PY" -u -c " -import time, torch -torch.set_grad_enabled(False) -from alignment_lab.lab.benchmark import load_case -from torchref.experimental.alignment.rotation_search import rotation_search -model, data = load_case('1DAW'); model.verbose = 0 -rotation_search(model, data, 0.8, n_peaks=50) # prewarm the process -prev = None -for i in (1, 2, 3): - t0 = time.perf_counter(); sol = rotation_search(model, data, 0.8, n_peaks=50) - dt = time.perf_counter() - t0 - fp = float(sol.scores[:20].double().sum()) - tag = '' if prev is None else (' same top-20 sum' if fp == prev else ' DIFFERS') - prev = fp - print(f' search {i}: {dt:.3f} s{tag}') -" 2>&1 | grep -vE "$FILT" - -echo "=== tests ===" -"$PY" -m pytest -q --no-header -p no:cacheprovider \ - tests/unit/alignment tests/unit/frf_separate tests/unit/model \ - > alignment_lab/slurm/_t_wmemo.log 2>&1 -rc=$? -tail -2 alignment_lab/slurm/_t_wmemo.log -echo "rc=$rc" -grep -E "^(FAILED|ERROR)" alignment_lab/slurm/_t_wmemo.log || echo " no failures" -echo "done" diff --git a/alignment_lab/analysis/weight_arms.py b/alignment_lab/analysis/weight_arms.py deleted file mode 100644 index 8ee2dba5..00000000 --- a/alignment_lab/analysis/weight_arms.py +++ /dev/null @@ -1,93 +0,0 @@ -"""Does the measured model-error curve replace the assumed one? - -The two-part system -- one scaler, one weight -- was meant to subsume seven -separate knobs. Four went with the scaler and the weight. These arms test the -last of them: `sigma_A` itself, which is currently a *prior* (a Luzzati falloff -from a coordinate error guessed off the residue count) patched at low resolution -by Babinet's two universal constants. - -It does not have to be assumed. Total scattering per shell is rotation- -invariant, so `Sigma_obs(s)/Sigma_calc(s)` measures the model's resolution- -dependent deficiency before the molecule is placed -- and it is safe to take -from the data being scored, because a quantity identical for every candidate -orientation cannot bias the ranking between them. - -Ranking is expected to be flat: it has been flat across every configuration of -scaling and weighting tried so far. That is not the question. The question is -whether the measured curve can stand in for the assumed one, so that two -declared objects replace seven knobs rather than four of them. -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, FRFConfig, orbit_rank, rotated_case, # noqa: E402 - run_frf, seed_for) - -#: Each arm is one deviation from `control`, which is what ships today. -ARMS = { - # What the engine did before this line of work. - "luzzati_babinet": {"sigma_a_source": "luzzati", "apply_bulk_solvent": True}, - # Luzzati without the two universal Babinet constants: how much of the - # low-resolution correction was the solvent term doing? - "luzzati_only": {"sigma_a_source": "luzzati", "apply_bulk_solvent": False}, - # sigma_A measured from Sigma_obs/Sigma_calc instead of assumed from an - # estimated coordinate error. Subsumes Babinet -- the solvent deficit is - # what the ratio measures -- so the solvent flag is irrelevant here. - "empirical": {"sigma_a_source": "empirical"}, - # And the same with no observed-side weight, to check the two halves of the - # system are still independent of each other. - "empirical_now": {"sigma_a_source": "empirical", "obs_weight": "none"}, -} - - -def main() -> int: - ap = argparse.ArgumentParser() - ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) - ap.add_argument("--trials", type=int, default=10) - ap.add_argument("--lmax-cap", type=int, default=64) - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--thr-deg", type=float, default=5.0) - ap.add_argument("--arms", default="") - args = ap.parse_args() - - names = [a for a in args.arms.split(",") if a] or list(ARMS) - for trial in range(args.trials): - seed = seed_for(args.pdb, trial) - model, data, R_true = rotated_case(args.pdb, seed) - sym = data.spacegroup.matrices.to(torch.float64).cpu() - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - okw = dict(side="left", frame="cart", reciprocal_basis=rec, - thr_deg=args.thr_deg) - for name in names: - cfg = FRFConfig(n_peaks=args.n_peaks, lmax_cap=args.lmax_cap, - **ARMS[name]) - t0 = time.time() - try: - res = run_frf(model, data, cfg, capture_arf=False, verbose=0) - except Exception as exc: - print(f"ROW {name} {args.pdb} trial={trial} seed={seed} rank=-1 " - f"rank_cmp={args.n_peaks} found=0 seconds=0.00 " - f"error={type(exc).__name__}", flush=True) - continue - seconds = time.time() - t0 - rank, ang = orbit_rank(res.peaks, R_true, sym, **okw) - print(f"ROW {name} {args.pdb} trial={trial} seed={seed} " - f"rank={rank} rank_cmp={rank if rank >= 0 else args.n_peaks} " - f"found={int(rank >= 0)} top20={int(0 <= rank < 20)} " - f"angle={'' if ang is None else round(float(ang), 3)} " - f"seconds={seconds:.2f}", flush=True) - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/alignment_lab/analysis/weight_arms.sh b/alignment_lab/analysis/weight_arms.sh deleted file mode 100644 index d489f7e6..00000000 --- a/alignment_lab/analysis/weight_arms.sh +++ /dev/null @@ -1,25 +0,0 @@ -#!/bin/bash -# Which E convention ranks truth best? Nine arms x 5 trials x 10 structures. -# the same pass. Every previously published FRF number was measured on - - -#SBATCH --job-name=warms -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%A_%a.err -#SBATCH --partition=day -#SBATCH --time=03:00:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=48G -#SBATCH --constraint=cpu_epyc9335 -#SBATCH --array=0-9 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -PDBS=(1DAW 3E98 3A5V 3VRJ 1AK5 3K7M 3GR5 2DQ6 4BX9 6G9X) -PDB=${PDBS[$SLURM_ARRAY_TASK_ID]} -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname) pdb=$PDB" -"$PY" -u alignment_lab/analysis/weight_arms.py --pdb "$PDB" --trials 10 2>/dev/null | grep '^ROW ' - diff --git a/alignment_lab/analysis/where_now.sh b/alignment_lab/analysis/where_now.sh deleted file mode 100644 index c3200601..00000000 --- a/alignment_lab/analysis/where_now.sh +++ /dev/null @@ -1,20 +0,0 @@ -#!/bin/bash -# Full stage breakdown of one rotation search, now that the SH expansion is no -# longer the dominant term. -#SBATCH --job-name=frf_where -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 -export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" -for pdb in 3K7M 1DAW; do - "$PY" -u -m alignment_lab.diagnostics.frf_benchmark \ - --pdb "$pdb" --arms cap100,cap64 --trials 2 2>&1 \ - | grep -vE "Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|warn|^ *$" -done -echo "exit_code=$?" diff --git a/alignment_lab/analysis/where_now_cap64.sh b/alignment_lab/analysis/where_now_cap64.sh deleted file mode 100644 index c5d1aa79..00000000 --- a/alignment_lab/analysis/where_now_cap64.sh +++ /dev/null @@ -1,24 +0,0 @@ -#!/bin/bash -# Where the time goes NOW. cap64 only -- that is what ships, and medianing it -# with cap100 inverted the stage ranking once already. -#SBATCH --job-name=wherenow -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=8 -#SBATCH --mem=48G -#SBATCH --exclusive -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" -for pdb in 3K7M 1AK5 1DAW; do - echo "############ $pdb ############" - "$PY" -u -m alignment_lab.diagnostics.frf_benchmark --pdb "$pdb" --arms cap64 --trials 3 2>&1 \ - | grep -vE "Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|warn|^ *$" -done diff --git a/alignment_lab/analysis/wigner_cache_probe.sh b/alignment_lab/analysis/wigner_cache_probe.sh deleted file mode 100644 index 9f238a33..00000000 --- a/alignment_lab/analysis/wigner_cache_probe.sh +++ /dev/null @@ -1,88 +0,0 @@ -#!/bin/bash -# Two questions about caching the Wigner d-table: -# 1. What fraction of wigner_contraction_per_beta is the data-INdependent -# d-block build (so, cacheable) vs the contraction against xi (not)? -# 2. Is loading that table off GPFS actually faster than recomputing it? -#SBATCH --job-name=frf_wcache -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname)" -SCRATCH=alignment_lab/runs/wigner_cache_probe -mkdir -p "$SCRATCH" -"$PY" -u -c " -import math, os, time -import torch; torch.set_grad_enabled(False) -from torchref.experimental.alignment.frf.wigner_d import ( - wigner_contraction_per_beta, _wigner_eig_table) - -scratch = '$SCRATCH' -dev = torch.device('cpu') - -def build_table(L, betas): - '''The data-independent half: every d^l(beta) block, packed, float32.''' - eig = _wigner_eig_table(L, dev) - out = [] - for l in range(1, L): - w, V = eig[l - 1] - phase = torch.exp(-1j * betas.unsqueeze(1) * w.unsqueeze(0)) - VP = V.unsqueeze(0) * phase.unsqueeze(1) - out.append((VP @ V.conj().transpose(-1, -2)).real.to(torch.float32)) - return out - -def contract_only(table, xi, L, n_beta): - '''The data-dependent half, given a prebuilt table.''' - dim = 2 * L - 1; c = L - 1 - S = torch.zeros((n_beta, dim, dim), dtype=torch.complex128) - S[:, c, c] += xi[0, c, c] - for l in range(1, L): - lo, hi = c - l, c + l + 1 - S[:, lo:hi, lo:hi] += xi[l, lo:hi, lo:hi].unsqueeze(0) * table[l-1].to(torch.complex128) - return S - -def best(fn, n=3): - fn() - return min((lambda: (lambda t0: (fn(), time.perf_counter()-t0)[1])(time.perf_counter()))() for _ in range(n)) - -for cap in (64, 100): - L = cap + 1 - n_beta = int(math.ceil(180.0 / 3.0)) - betas = torch.arange(n_beta, dtype=torch.float64) * 3.0 * (math.pi/180) - xi = torch.randn(L, 2*L-1, 2*L-1, dtype=torch.complex128) * 1e-3 - - t_full = best(lambda: wigner_contraction_per_beta(xi, betas)) - t_build = best(lambda: build_table(L, betas)) - table = build_table(L, betas) - t_contr = best(lambda: contract_only(table, xi, L, n_beta)) - - # Correctness of the split, in float32 storage. - ref = wigner_contraction_per_beta(xi, betas) - got = contract_only(table, xi, L, n_beta) - rel = ((got - ref).abs().max() / ref.abs().max()).item() - - # Round-trip through GPFS, as one packed flat tensor. - flat = torch.cat([t.reshape(-1) for t in table]) - path = os.path.join(scratch, f'dtable_L{L}.pt') - t0 = time.perf_counter(); torch.save(flat, path); t_save = time.perf_counter()-t0 - mb = os.path.getsize(path)/1e6 - os.system('sync') - loads = [] - for _ in range(3): - t0 = time.perf_counter(); torch.load(path, map_location='cpu'); loads.append(time.perf_counter()-t0) - t_load_warm = min(loads) - print(f'--- cap{cap} L={L} n_beta={n_beta} ---') - print(f' full contraction now {t_full*1e3:8.1f} ms') - print(f' of which d-block build {t_build*1e3:8.1f} ms (cacheable)') - print(f' of which xi contraction {t_contr*1e3:8.1f} ms (not)') - print(f' float32 table rel.err {rel:8.2e}') - print(f' on disk {mb:8.1f} MB save {t_save*1e3:.0f} ms ' - f'load(page-cached) {t_load_warm*1e3:.0f} ms') - os.remove(path) -" -echo "exit_code=$?" diff --git a/alignment_lab/analysis/wigner_call_count.sh b/alignment_lab/analysis/wigner_call_count.sh deleted file mode 100644 index 0e5ad8c5..00000000 --- a/alignment_lab/analysis/wigner_call_count.sh +++ /dev/null @@ -1,47 +0,0 @@ -#!/bin/bash -# A RAM memo of the d-blocks only pays off if the table is wanted more than -# once. How many times does one rotation_search ask for it, and at what (L, -# n_beta)? An earlier note claimed the obs and calc sides each build it. -#SBATCH --job-name=frf_wcnt -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname)" -"$PY" -u -c " -import torch -torch.set_grad_enabled(False) -import torchref.experimental.alignment.frf.wigner_d as wd -from alignment_lab.lab.benchmark import load_case -from torchref.experimental.alignment.rotation_search import rotation_search - -calls = [] -orig = wd.wigner_contraction_per_beta -def counting(xi, betas): - calls.append((int(xi.shape[0]), int(betas.shape[0]))) - return orig(xi, betas) -wd.wigner_contraction_per_beta = counting -# the caller imported it by name, so rebind there too -import torchref.experimental.alignment.frf.sitelist_ang as sa -for mod in (sa,): - if getattr(mod, 'wigner_contraction_per_beta', None) is orig: - mod.wigner_contraction_per_beta = counting -import sys -for name, mod in list(sys.modules.items()): - if name.startswith('torchref.experimental.alignment') and \ - getattr(mod, 'wigner_contraction_per_beta', None) is orig: - mod.wigner_contraction_per_beta = counting - print(' rebound in', name) - -model, data = load_case('1DAW'); model.verbose = 0 -for run in (1, 2): - calls.clear() - rotation_search(model, data, 0.8, n_peaks=50) - print(f' search {run}: {len(calls)} call(s) to wigner_contraction_per_beta -> {calls}') -" 2>&1 | grep -vE "Loaded|LINK|Wilson|found non|FrenchWilson|Reflections:|Resolution:|Space group|Centric:|✓|Parametrization|Warning|warn| from |^ *$|No CUDA" -echo "done" diff --git a/alignment_lab/analysis/wigner_fft_proto.sh b/alignment_lab/analysis/wigner_fft_proto.sh deleted file mode 100644 index 1f4eaa57..00000000 --- a/alignment_lab/analysis/wigner_fft_proto.sh +++ /dev/null @@ -1,73 +0,0 @@ -#!/bin/bash -# The J_y eigenvalues are the integers -l..l, so S(beta) is a trigonometric -# polynomial in beta: accumulate its Fourier coefficients once (no beta axis) -# and get every beta from one FFT. Prototype + correctness + timing, against -# the current per-beta matmul. -#SBATCH --job-name=frf_wfft -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname)" -"$PY" -u -c " -import math, time, torch -torch.set_grad_enabled(False) -from torchref.experimental.alignment.frf.wigner_d import ( - wigner_contraction_per_beta, _wigner_eig_table) - -dev = torch.device('cpu') - -def contraction_fft(xi, betas): - L = xi.shape[0]; dim = 2 * L - 1; c = L - 1 - n_beta = betas.shape[0] - N = 2 * n_beta # betas must be j * 2*pi/N - C = torch.zeros((dim, dim, N), dtype=torch.complex128) - C[c, c, 0] += 2.0 * xi[0, c, c] # l=0: d^0 = 1, the halving is undone below - eig = _wigner_eig_table(L, dev) - for l in range(1, L): - w, V = eig[l - 1] - lo, hi = c - l, c + l + 1 - xi_l = xi[l, lo:hi, lo:hi] - k = (torch.round(w).to(torch.long)) % N - G = V.unsqueeze(1) * V.conj().unsqueeze(0) # (sz, sz, sz) over k - blk = C[lo:hi, lo:hi] - blk.index_add_(2, k, xi_l.unsqueeze(-1) * G) - blk.index_add_(2, (-k) % N, xi_l.unsqueeze(-1) * G.transpose(0, 1)) - full = torch.fft.fft(C, n=N, dim=-1) - return 0.5 * full[..., :n_beta].permute(2, 0, 1).contiguous() - -def best(fn, n=3): - fn() - out = [] - for _ in range(n): - t0 = time.perf_counter(); fn(); out.append(time.perf_counter() - t0) - return min(out) - -for cap in (64, 100): - L = cap + 1 - n_beta = int(math.ceil(180.0 / 3.0)) - betas = torch.arange(n_beta, dtype=torch.float64) * 3.0 * (math.pi / 180.0) - xi = torch.randn(L, 2*L-1, 2*L-1, dtype=torch.complex128) * 1e-3 - - # Are the eigenvalues really integers? The whole method rests on it. - eig = _wigner_eig_table(L, dev) - dev_max = max(float((w - torch.round(w)).abs().max()) for w, _ in eig) - - ref = wigner_contraction_per_beta(xi, betas) - got = contraction_fft(xi, betas) - rel = float((got - ref).abs().max() / ref.abs().max()) - - t_ref = best(lambda: wigner_contraction_per_beta(xi, betas)) - t_new = best(lambda: contraction_fft(xi, betas)) - print(f'--- cap{cap} L={L} n_beta={n_beta}') - print(f' max |w - round(w)| {dev_max:.2e}') - print(f' rel. difference {rel:.2e}') - print(f' per-beta matmul {t_ref*1e3:8.1f} ms') - print(f' Fourier + one FFT {t_new*1e3:8.1f} ms ({t_ref/t_new:.1f}x)') -" 2>&1 | grep -vE "Warning|warn| from |^ *$" -echo "exit_code=$?" diff --git a/alignment_lab/analysis/wigner_mirror_proto.sh b/alignment_lab/analysis/wigner_mirror_proto.sh deleted file mode 100644 index 1e33aa80..00000000 --- a/alignment_lab/analysis/wigner_mirror_proto.sh +++ /dev/null @@ -1,79 +0,0 @@ -#!/bin/bash -# d^l(pi - beta) is d^l(beta) up to an m-flip and a sign, and the beta grid -# (0, 3, ..., 177 deg) is closed under beta -> pi - beta. So the batched matmul -# only needs 31 of the 60 beta values. No data file, no kernel. -#SBATCH --job-name=frf_wmir -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 -export PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname) cpu=$(grep -m1 'model name' /proc/cpuinfo | cut -d: -f2-)" -"$PY" -u -c " -import math, time, torch -torch.set_grad_enabled(False) -from torchref.experimental.alignment.frf.wigner_d import ( - wigner_contraction_per_beta, _wigner_eig_table) -dev = torch.device('cpu') - -def d_block(w, V, betas): - phase = torch.exp(-1j * betas.unsqueeze(1) * w.unsqueeze(0)) - return ((V.unsqueeze(0) * phase.unsqueeze(1)) @ V.conj().transpose(-1, -2)).real - -# Which mirror identity holds? Test rather than trust. -L = 9 -eig = _wigner_eig_table(L, dev) -for l in (1, 3, 6, 8): - w, V = eig[l - 1] - b = torch.tensor([0.37, 1.11, 2.05], dtype=torch.float64) - lhs = d_block(w, V, math.pi - b) - base = d_block(w, V, b) - m = torch.arange(-l, l + 1, dtype=torch.float64) - s2 = ((-1.0) ** (l + m)).reshape(1, 1, -1) - s1 = ((-1.0) ** (l + m)).reshape(1, -1, 1) - f2 = s2 * base.flip(-1) # (-1)^(l+m2) d[m1, -m2] - f1 = s1 * base.flip(-2) # (-1)^(l+m1) d[-m1, m2] - print(f' l={l}: form(m2-flip) err={float((lhs-f2).abs().max()):.2e} ' - f'form(m1-flip) err={float((lhs-f1).abs().max()):.2e}') - -def contraction_mirror(xi, betas): - L = xi.shape[0]; dim = 2 * L - 1; c = L - 1 - n_beta = betas.shape[0] - half = n_beta // 2 + 1 # 0..30 for n_beta=60 - src = n_beta - torch.arange(half, n_beta) # beta_j -> pi - beta_j - S = torch.zeros((n_beta, dim, dim), dtype=torch.complex128) - S[:, c, c] += xi[0, c, c] - eig = _wigner_eig_table(L, dev) - for l in range(1, L): - w, V = eig[l - 1] - d_h = d_block(w, V, betas[:half]) - m = torch.arange(-l, l + 1, dtype=d_h.dtype) - sgn = ((-1.0) ** (l + m)).reshape(1, 1, -1) - d_l = torch.cat([d_h, sgn * d_h[src].flip(-1)], dim=0) - lo, hi = c - l, c + l + 1 - S[:, lo:hi, lo:hi] += xi[l, lo:hi, lo:hi].unsqueeze(0) * d_l.to(torch.complex128) - return S - -def best(fn, n=3): - fn() - return min((lambda: (lambda t0: (fn(), time.perf_counter()-t0)[1])(time.perf_counter()))() for _ in range(n)) - -for cap in (64, 100): - L = cap + 1 - n_beta = int(math.ceil(180.0 / 3.0)) - betas = torch.arange(n_beta, dtype=torch.float64) * 3.0 * (math.pi / 180.0) - xi = torch.randn(L, 2*L-1, 2*L-1, dtype=torch.complex128) * 1e-3 - ref = wigner_contraction_per_beta(xi, betas) - got = contraction_mirror(xi, betas) - rel = float((got - ref).abs().max() / ref.abs().max()) - t_ref = best(lambda: wigner_contraction_per_beta(xi, betas)) - t_new = best(lambda: contraction_mirror(xi, betas)) - print(f'--- cap{cap} L={L}: rel diff {rel:.2e} ' - f'current {t_ref*1e3:7.1f} ms mirrored {t_new*1e3:7.1f} ms ' - f'({t_ref/t_new:.2f}x)') -" 2>&1 | grep -vE "Warning|warn| from |^ *$|No CUDA" -echo "exit_code=$?" diff --git a/alignment_lab/analysis/wigner_split.sh b/alignment_lab/analysis/wigner_split.sh deleted file mode 100644 index 18a18d1c..00000000 --- a/alignment_lab/analysis/wigner_split.sh +++ /dev/null @@ -1,44 +0,0 @@ -#!/bin/bash -# Inside wigner_contraction_per_beta, how much is data-INdependent (so -# cacheable) and how much is the contraction against xi (so not)? -#SBATCH --job-name=frf_wigner -#SBATCH --output=alignment_lab/slurm/%x_%j.out -#SBATCH --error=alignment_lab/slurm/%x_%j.err -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" TORCHREF_NUM_THREADS=4 OMP_NUM_THREADS=4 PYTHONUNBUFFERED=1 -export CUDA_VISIBLE_DEVICES="" -mkdir -p alignment_lab/slurm -echo "host=$(hostname)" -"$PY" -u -c " -import time -import torch; torch.set_grad_enabled(False) -from torchref.experimental.alignment.frf.wigner_d import ( - wigner_contraction_per_beta, _wigner_eig_table) -from torchref.experimental.alignment.frf.sitelist_ang import _SAMPLE_LIST_CACHE - -for cap, sampling in ((64, 3.0), (100, 3.0)): - L = cap + 1 - n_beta = int(__import__('math').ceil(180.0 / sampling)) - betas = torch.arange(n_beta, dtype=torch.float64) * sampling * (3.141592653589793/180) - xi = torch.randn(L, 2*L-1, 2*L-1, dtype=torch.complex128) * 1e-3 - - t0 = time.perf_counter(); _wigner_eig_table(L, torch.device('cpu')) - cold_eig = time.perf_counter() - t0 - t0 = time.perf_counter(); _wigner_eig_table(L, torch.device('cpu')) - warm_eig = time.perf_counter() - t0 - - wigner_contraction_per_beta(xi, betas) # warm everything - t0 = time.perf_counter() - wigner_contraction_per_beta(xi, betas) - total = time.perf_counter() - t0 - - # Size of the d-table if it were precomputed and stored, float32 real. - entries = sum((2*l+1)**2 for l in range(L)) * n_beta - print(f'L={L:4d} n_beta={n_beta:3d} | eig cold {cold_eig*1e3:7.1f} ms ' - f'warm {warm_eig*1e3:5.2f} ms | contraction total {total*1e3:7.1f} ms ' - f'| d-table would be {entries*4/1e6:7.1f} MB float32') -" -echo "exit_code=$?" diff --git a/alignment_lab/analysis/wilson_moments.sh b/alignment_lab/analysis/wilson_moments.sh deleted file mode 100644 index 5f83e7c9..00000000 --- a/alignment_lab/analysis/wilson_moments.sh +++ /dev/null @@ -1,17 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=moments -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:55:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=64G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO:$REPO/alignment_lab" TORCHREF_NUM_THREADS=4 -export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 CUDA_VISIBLE_DEVICES="" -"$PY" -u alignment_lab/diagnostics/wilson_moments.py 2>/dev/null -echo DONE diff --git a/alignment_lab/analysis/wilson_smoke.py b/alignment_lab/analysis/wilson_smoke.py deleted file mode 100644 index dfb531cb..00000000 --- a/alignment_lab/analysis/wilson_smoke.py +++ /dev/null @@ -1,92 +0,0 @@ -"""Does the Wilson normaliser hold its identity on real data, both sides? - -The synthetic test draws from the distribution the fit assumes, so it can only -show the arithmetic is right. These are the cases the assumption is wrong in: -observations carry measurement error and a real solvent deficit, and the -rotation function's calc side is an oversampled molecular transform in a P1 box, -where adjacent samples are correlated and Wilson independence does not hold at -all. The mean estimate survives that by quasi-likelihood -- a log-link Gamma GLM -is consistent for the mean under a misspecified variance function -- and this is -where that claim gets checked rather than asserted. - -Also checks the property the whole weighting design rests on: fitted with a -SHARED abscissa, the two curves are comparable, and their ratio is the -resolution-dependent model deficiency. -""" - -from __future__ import annotations - -import sys -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import load_case # noqa: E402 - -CASES = ("1DAW", "2DQ6", "3K7M", "4BX9") - - -def main() -> int: - from torchref.scaling.wilson import WilsonNormaliser - - for pdb in CASES: - model, data = load_case(pdb) - hkl = data.hkl - rec = data.cell.reciprocal_basis_matrix.to(torch.float64) - s = (hkl.to(torch.float64) @ rec).norm(dim=-1) - F = data.F.to(torch.float64).abs() - keep = torch.isfinite(F) & (F > 0) - has_I = getattr(data, "I", None) is not None - print(f"\n=== {pdb} {data.spacegroup.hm} n={int(keep.sum())} " - f"raw intensities available: {has_I} ===") - - # --- observed side ------------------------------------------------- - I_obs = (F * F)[keep] - obs = WilsonNormaliser.from_hkl( - I_obs, hkl[keep], data.spacegroup, data.cell, n_coeff=6, - s_lo=float(s.min()), s_hi=float(s.max()), - ) - cen = data.spacegroup.is_centric(hkl[keep].to(torch.long)).to(torch.bool) - k = torch.where(cen, 0.5, 1.0).to(torch.float64) - e2 = obs.E_squared.to(torch.float64) - print(f" obs {obs!r}") - print(f" k-weighted = {float((k*e2).sum()/k.sum()):.10f}") - _deciles(" obs decile ", s[keep], e2) - - # --- calculated side, on the SAME abscissa -------------------------- - F_calc = model.get_structure_factor(hkl[keep], recalc=True).abs() - I_calc = (F_calc.to(torch.float64) ** 2) - calc = WilsonNormaliser.from_hkl( - I_calc, hkl[keep], data.spacegroup, data.cell, n_coeff=6, - s_lo=float(s.min()), s_hi=float(s.max()), - ) - e2c = calc.E_squared.to(torch.float64) - print(f" calc {calc!r}") - print(f" k-weighted = {float((k*e2c).sum()/k.sum()):.10f}") - _deciles(" calc decile ", s[keep], e2c) - - # --- the ratio the weight will be built from ------------------------ - # Same basis, so the two curves are directly comparable. Reported as a - # shape: what the model under-explains, versus resolution. - grid = torch.linspace(float(s[keep].min()), float(s[keep].max()), 8) - r = (obs.evaluate(grid).to(torch.float64) - / calc.evaluate(grid).to(torch.float64)) - r = r / r.mean() - print(" Sigma_obs/Sigma_calc (normalised) vs d(A):") - print(" " + " ".join(f"{1/float(x):5.1f}:{float(v):5.2f}" - for x, v in zip(grid, r))) - return 0 - - -def _deciles(label, s, v): - order = torch.argsort(s) - d = [float(v[order[i::10]].mean()) for i in range(10)] - print(f"{label}: min {min(d):.4f} max {max(d):.4f} " - + " ".join(f"{x:.2f}" for x in d)) - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/alignment_lab/analysis/wilson_smoke.sh b/alignment_lab/analysis/wilson_smoke.sh deleted file mode 100644 index 76855935..00000000 --- a/alignment_lab/analysis/wilson_smoke.sh +++ /dev/null @@ -1,17 +0,0 @@ -#!/bin/bash -#SBATCH --job-name=wsmoke -#SBATCH --output=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.out -#SBATCH --error=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement/alignment_lab/slurm/%x_%j.err -#SBATCH --partition=hour -#SBATCH --time=00:30:00 -#SBATCH --cpus-per-task=4 -#SBATCH --mem=32G -#SBATCH --constraint=cpu_epyc9335 -set -uo pipefail -REPO=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/alignement -PY=/das/work/units/LBR-FEL/p17490/Peter/Library/work_trees_torchref/dev/.dev/bin/python -cd "$REPO" -export PYTHONPATH="$REPO" OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 PYTHONUNBUFFERED=1 -export CUDA_VISIBLE_DEVICES="" -"$PY" -u alignment_lab/analysis/wilson_smoke.py -echo "RC=$?" diff --git a/alignment_lab/diagnostics/empirical_sigma_a_check.py b/alignment_lab/diagnostics/empirical_sigma_a_check.py deleted file mode 100644 index f824c3a9..00000000 --- a/alignment_lab/diagnostics/empirical_sigma_a_check.py +++ /dev/null @@ -1,50 +0,0 @@ -"""What does the rotation function's "empirical" sigma_A actually evaluate to? - -``empirical_sigma_a`` divides the observed Wilson curve by the calculated one -and takes ``sqrt(min(R, 1/R))``. The two curves sit on different absolute -scales -- the MTZ's arbitrary one and the model's electron scale -- and the -function does not remove that factor, so the ratio's level, not only its shape, -sets the answer. This records the ratio and the resulting sigma_A by resolution -during a real rotation search. -""" -from __future__ import annotations - -import argparse -import sys -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -from lab import BENCH_PDBS, load_case # noqa: E402 - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) - args = ap.parse_args() - - import torchref.experimental.alignment.frf.api as api - from torchref.experimental.alignment import rotation_search - - seen = {} - real = api.empirical_sigma_a - - def spy(sigma_obs, sigma_calc, **kw): - out = real(sigma_obs, sigma_calc, **kw) - seen["ratio"] = (sigma_obs / sigma_calc).detach().cpu() - seen["sigma_a"] = out.detach().cpu() - return out - - api.empirical_sigma_a = spy - model, data = load_case(args.pdb) - rotation_search(model, data, model_error_A=0.8, n_peaks=5) - r, sa = seen["ratio"], seen["sigma_a"] - q = lambda x: [round(float(v), 4) for v in torch.quantile(x, torch.tensor([0.0, 0.25, 0.5, 0.75, 1.0], dtype=x.dtype))] - print(f"ROW pdb={args.pdb} ratio_quantiles={q(r)} sigma_a_quantiles={q(sa)} " - f"ratio_geomean={float(r.log().mean().exp()):.4g}") - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_aniso_knockout.py b/alignment_lab/diagnostics/frf_aniso_knockout.py deleted file mode 100644 index 6c383e37..00000000 --- a/alignment_lab/diagnostics/frf_aniso_knockout.py +++ /dev/null @@ -1,165 +0,0 @@ -"""Is our own anisotropy correction what destroys the hexagonal cases? - -The normaliser anatomy (job 489537) decomposed ``log(Esqr_phaser / Esqr_ours)`` -and found the disagreement is overwhelmingly **angular**, and only on the two -failing structures: - -| pdb | rms | eps_n | iso(|s|) | anisotropy | equivalent B spread | -|------|-------|--------|----------|------------|---------------------| -| 1AK5 | 14% | 61.5% | 23.1% | **0.01%** | 0.36 A^2 | -| 2DQ6 | 44% | 1.6% | 16.4% | **78.1%** | 158 A^2 | -| 3GR5 | 116% | 0.9% | 6.0% | **92.2%** | 189 A^2 | - -A radial mis-scaling is nearly harmless to a rotation function; an angular one is -exactly what it measures. So the suspect is our own overall-anisotropy -correction, ``fit_overall_anisotropy`` (``sh.py:445``), which regresses -``ln|F|^2 - ln<|F|^2>_shell`` on ``-2 pi^2 s.U.s`` by unweighted least squares -**with no constant term**. Single-reflection ``ln|F|^2`` is a badly behaved -regressand: its expectation is offset by ``-gamma`` for acentrics and -``-gamma - ln 2`` for centrics, the ``clamp(min=1e-30)`` turns a vanishing -amplitude into ``y ~ -69``, and with no intercept every one of those offsets is -absorbed into the quadratic form. - -``symmetrize_anisotropy`` then projects the result onto the point-group-invariant -subspace, and the code comment records that the raw fit gives eigenvalues -``(0.8, 17, 70) A^2`` on a *cubic* dataset where symmetry forces them equal. That -projection is why the damage is invisible on the working structures and not on -these two: - -* cubic -> 1 DOF (lambda I): the garbage is annihilated; -* trigonal/hexagonal -> 2 DOF (diag(lambda, lambda, mu)): a **uniaxial tensor - along c is symmetry-allowed**, so the garbage survives as exactly the fake - anisotropy the anatomy measures. - -Three arms settle it, and the null arm is the one that matters -- if switching the -correction off recovers truth, our correction is not merely imperfect, it is -actively destructive: - -* ``production`` -- fitted, symmetrised U; -* ``no_aniso`` -- U = 0, no correction at all; -* ``iso_only`` -- U = (trace/3) I, keeping the radial part and dropping every - angular component, which the shell means then absorb. - -Also reported per structure: the eigenvalues of the fitted tensor before and -after symmetrisation, as B = 8 pi^2 U, so the size of the artefact is visible -next to the rank it costs. - -Usage ------ - python -m diagnostics.frf_aniso_knockout --pdb 3GR5 -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) - -from lab import (ANISO_ARMS, FRFConfig, aniso_arm, load_case, # noqa: E402 - patched, run_frf, tensor_report) -from lab.results import append_row, provenance # noqa: E402 -from diagnostics.frf_ghost_knockout import ( # noqa: E402 - PHASER_PINNED, _orbit_of_identity, _truth_and_margin, -) - -EXPERIMENT = "frf_aniso_knockout" - -ARMS = ANISO_ARMS - - -def run_arm(pdb: str, arm: str, *, n_peaks: int = 500) -> dict: - from torchref.experimental.alignment.frf import api as _api - - pin = PHASER_PINNED[pdb] - model, data = load_case(pdb) - orbit = _orbit_of_identity(data) - - seen: dict = {} - cfg_probe = FRFConfig() - - def _pinned(model_radius_A, d_min_data, lmax_cap=48): - return int(pin["lmax"]) + 1, float(pin["d_min_eff"]) - - cfg = FRFConfig(n_peaks=n_peaks, lmax_cap=int(pin["lmax"]), - extra={"grid_sampling_deg": float(pin["sampling_deg"])}) - t0 = time.time() - with patched(_api, "phaser_lmax_resolution", _pinned), \ - aniso_arm(arm, data, d_min=cfg_probe.d_min, d_max=cfg_probe.d_max, - captured=seen): - res = run_frf(model, data, cfg, capture_arf=True, verbose=0) - rank, sig, ang, ghost, margin = _truth_and_margin(res.arf, orbit) - - # The tensor actually applied, recomputed the same way align.py does it. - from torchref.experimental.alignment.sh import ( - hkl_symops_to_cartesian, symmetrize_anisotropy, - ) - rec = data.cell.reciprocal_basis_matrix.to(torch.float64) - cart = hkl_symops_to_cartesian( - data.spacegroup.matrices.to(torch.float64), rec) - raw = seen.get("raw") - - row = {"experiment": EXPERIMENT, "pdb": pdb, "arm": arm} - row.update(provenance()) - row.update({ - "spacegroup": str(data.spacegroup.hm), - "n_ops": int(data.spacegroup.matrices.shape[0]), - "lmax": pin["lmax"], "sampling_deg": pin["sampling_deg"], - "d_min_eff": pin["d_min_eff"], - "n_samples": int(res.arf.values.numel()), - "truth_rank": rank, "truth_sigma": round(sig, 4), - "truth_angle_deg": round(ang, 3), - "best_ghost_sigma": round(ghost, 4), "margin": round(margin, 4), - "seconds": round(time.time() - t0, 1), - }) - if raw is not None: - row.update(tensor_report(raw.cpu(), "raw")) - row.update(tensor_report( - symmetrize_anisotropy(raw.to(torch.float64).cpu(), cart.cpu()), - "sym")) - if "fixed" in seen: - row.update(tensor_report(seen["fixed"].cpu(), "fix_raw")) - row.update(tensor_report( - symmetrize_anisotropy(seen["fixed"].to(torch.float64).cpu(), - cart.cpu()), "fix_sym")) - return row - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", required=True, choices=sorted(PHASER_PINNED)) - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--outdir", default=None) - args = ap.parse_args() - - outdir = Path(args.outdir) if args.outdir else ( - Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT) - outdir.mkdir(parents=True, exist_ok=True) - csv_path = outdir / f"{EXPERIMENT}_{args.pdb}.csv" - - rows = [run_arm(args.pdb, a, n_peaks=args.n_peaks) for a in ARMS] - r0 = rows[0] - print(f"\n{args.pdb} ({r0['spacegroup']}): fitted anisotropy as B (A^2) -- " - f"raw {r0.get('raw_B_min')}..{r0.get('raw_B_max')} " - f"(spread {r0.get('raw_B_spread')}), after symmetrisation " - f"{r0.get('sym_B_min')}..{r0.get('sym_B_max')} " - f"(spread {r0.get('sym_B_spread')})", flush=True) - print(f"{'arm':<14}{'rank':>8}{'truth_sig':>11}{'ghost_sig':>11}{'margin':>9}", - flush=True) - cols = {} - for r in rows: - cols.update({k: "" for k in r}) - for r in rows: - append_row(csv_path, {**cols, **r}) - print(f"{r['arm']:<14}{r['truth_rank']:>8}{r['truth_sigma']:>11.2f}" - f"{r['best_ghost_sigma']:>11.2f}{r['margin']:>+9.2f}", flush=True) - print(f"\nwrote {csv_path}", flush=True) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_aniso_rank_sweep.py b/alignment_lab/diagnostics/frf_aniso_rank_sweep.py deleted file mode 100644 index 6de82df7..00000000 --- a/alignment_lab/diagnostics/frf_aniso_rank_sweep.py +++ /dev/null @@ -1,152 +0,0 @@ -"""Does the anisotropy fix hold on the task the pipeline actually runs? - -Every number behind the anisotropy diagnosis was measured with the model in its -DEPOSITED orientation and with lmax / sampling / resolution pinned to Phaser's -own logged values -- one evaluation per structure, truth at the identity. That -was the right setup for a bisection against Phaser, and it is the wrong setup -for deciding a default: - -* the pipeline searches a RANDOMLY ROTATED model, not the identity, and the - ghosts are pose-dependent; -* it runs at the production configuration, not Phaser's pinned one; -* seed-to-seed truth-rank spread at ``lmax_cap = 64`` is +-4 to 6 ranks - (1AK5 [9, 11, 17], 3K7M [7, 8, 20]), and three earlier findings in this - investigation looked strong at n <= 7 and vanished at full n. - -So this re-measures the arms over seeded random rotations at the production -config, reporting **per-trial paired differences against the production arm** -rather than a bare median -- the same discipline the rest of the lab uses. - -Arms are :data:`lab.aniso.ARMS`: ``production``, ``no_aniso``, ``iso_only``, -``fixed_fit``. - -Usage ------ - python -m diagnostics.frf_aniso_rank_sweep --pdb 3GR5 --trials 10 -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (ANISO_ARMS, BENCH_PDBS, FRFConfig, aniso_arm, # noqa: E402 - orbit_rank, rotated_case, run_frf, seed_for, tensor_report) -from lab.results import append_row, provenance # noqa: E402 - -EXPERIMENT = "frf_aniso_rank_sweep" - - -def run_one(pdb: str, trial: int, arm: str, cfg: FRFConfig, - *, thr_deg: float) -> dict: - seed = seed_for(pdb, trial) - model, data, R_true = rotated_case(pdb, seed) - captured: dict = {} - t0 = time.time() - with aniso_arm(arm, data, d_min=cfg.d_min, d_max=cfg.d_max, - captured=captured): - res = run_frf(model, data, cfg, capture_arf=False, verbose=0) - seconds = time.time() - t0 - - rank, ang = orbit_rank( - res.peaks, R_true, data.spacegroup.matrices.to(torch.float64).cpu(), - reciprocal_basis=data.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), - side="right", frame="cart", thr_deg=thr_deg, - ) - row = {"experiment": EXPERIMENT, "pdb": pdb, "trial": trial, "arm": arm, - "seed": seed} - row.update(provenance()) - row.update(cfg.as_row()) - row.update({ - "spacegroup": str(data.spacegroup.hm), - "truth_rank": rank, - # orbit_rank returns -1 for "no peak within thr_deg". That must NOT be - # ordered as a good rank: for paired comparison a miss counts as worse - # than the worst hit, i.e. the peak-list length. - "rank_for_compare": rank if rank >= 0 else cfg.n_peaks, - "found": int(rank >= 0), - "truth_angle_deg": None if ang is None else round(float(ang), 3), - "n_peaks_found": len(res.peaks), - "orbit_side": "left", "orbit_frame": "cart", "thr_deg": thr_deg, - "seconds": round(seconds, 1), - }) - for tag in ("raw", "fixed"): - if tag in captured: - row.update(tensor_report(captured[tag], tag)) - return row - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) - ap.add_argument("--trials", type=int, default=10) - ap.add_argument("--arms", default=",".join(ANISO_ARMS)) - ap.add_argument("--lmax-cap", type=int, default=64) - ap.add_argument("--d-min", type=float, default=4.0) - ap.add_argument("--d-max", type=float, default=15.0) - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--thr-deg", type=float, default=5.0) - ap.add_argument("--outdir", default=None) - args = ap.parse_args() - - arms = [a for a in args.arms.split(",") if a] - cfg = FRFConfig(d_min=args.d_min, d_max=args.d_max, - n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) - outdir = Path(args.outdir) if args.outdir else ( - Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT) - outdir.mkdir(parents=True, exist_ok=True) - csv_path = outdir / f"{EXPERIMENT}_{args.pdb}.csv" - - print(f"{args.pdb}: truth rank per trial, lmax_cap={args.lmax_cap}", - flush=True) - print(f"{'trial':>6}" + "".join(f"{a:>15}" for a in arms), flush=True) - ranks = {a: [] for a in arms} - n_fail = 0 - for trial in range(args.trials): - cells = [] - for arm in arms: - try: - row = run_one(args.pdb, trial, arm, cfg, thr_deg=args.thr_deg) - except Exception as exc: - n_fail += 1 - ranks[arm].append(None) - cells.append(f"{type(exc).__name__}") - print(f" trial {trial} arm {arm} FAILED: {exc}", flush=True) - continue - append_row(csv_path, row) - ranks[arm].append(row["rank_for_compare"]) - cells.append(str(row["truth_rank"]) if row["found"] - else f"miss({row['truth_angle_deg']:.0f}d)") - print(f"{trial:>6}" + "".join(f"{c:>15}" for c in cells), flush=True) - - # Paired differences against production; per-trial signs, never a bare median. - base = ranks.get("production") - if base: - print("\npaired vs production (negative = better rank):", flush=True) - for arm in arms: - if arm == "production": - continue - d = [(a - b) for a, b in zip(ranks[arm], base) - if a is not None and b is not None] - if not d: - print(f" {arm:<14} no paired trials", flush=True) - continue - sd = sorted(d) - med = (sd[len(sd) // 2] if len(sd) % 2 - else 0.5 * (sd[len(sd) // 2 - 1] + sd[len(sd) // 2])) - print(f" {arm:<14} n={len(d):<3} better={sum(x < 0 for x in d)} " - f"same={sum(x == 0 for x in d)} worse={sum(x > 0 for x in d)} " - f"median_delta={med:+.1f} per-trial={d}", flush=True) - print(f"\nwrote {csv_path} ({n_fail} failures)", flush=True) - return 1 if n_fail == len(arms) * args.trials else 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_benchmark.py b/alignment_lab/diagnostics/frf_benchmark.py deleted file mode 100644 index 8ab7bc08..00000000 --- a/alignment_lab/diagnostics/frf_benchmark.py +++ /dev/null @@ -1,277 +0,0 @@ -"""The rotation search's standing benchmark: accuracy, memory and runtime. - -One row per (structure, trial, arm), carrying all three so a change cannot -improve one at the silent expense of another: - -**Accuracy** -- where the true orientation lands in the peak list, and whether -it is inside the top ``--top-n``. That window, not rank 0, is the thing that -matters: the placement search carries its top candidates forward, so rank 7 and -rank 0 are the same outcome downstream and rank 223 is not. Reported per -structure, because an average lets one failing space group be cancelled by nine -easy ones. - -**Memory** -- peak resident set over the search, as a delta over the value on -entry, plus the process high-water mark. Read -:mod:`alignment_lab.lab.profile` for what a sampler can and cannot see. Memory -is why the bandwidth ceiling is not simply "as high as possible": cap 100 needs -more than 32 GB on the two P432 structures, where the symmetry expansion -multiplies the reflection count by 24. - -**Runtime** -- total, plus per-stage. Cold by default, since a caller placing one -model pays cold costs; ``--warmup`` gives steady state, and the row says which. -Every row carries a fixed calibration workload timed in the same process, and the -node's identity: wall clock on a shared cluster measures the cluster unless it is -normalised or confined to one node. Run with ``--exclusive`` and compare -``seconds_per_calibration`` across nodes, or raw seconds only within a node. - -Usage ------ - python -m diagnostics.frf_benchmark --pdb 1DAW --trials 3 - python -m diagnostics.frf_benchmark --pdb 3K7M --arms cap48,cap64,cap100 -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, FRFConfig, orbit_rank, rotated_case, # noqa: E402 - run_frf, seed_for) -from lab.profile import (FRF_STAGES, PeakMemory, calibration_seconds, # noqa: E402 - exclusive_times, host_info, stage_timers) -from lab.results import append_row, provenance # noqa: E402 - -EXPERIMENT = "frf_benchmark" - -#: Named bandwidth arms. ``shipped`` reads the engine's own constant, so the -#: benchmark follows the code rather than restating it -- if the constant moves -#: and this row does not, the harness is measuring history. -ARMS = {"cap48": 48, "cap64": 64, "cap100": 100, "shipped": None} - - -def _shipped_lmax_cap() -> int: - import importlib - - return importlib.import_module( - "torchref.experimental.alignment.rotation_search").LMAX_CAP - - -def warmup_run(pdb: str, lmax_cap: int, n_peaks: int) -> float: - """One discarded search, to move the process's start-up out of the way. - - On this cluster the first real computation in a process is dominated by - PyTorch loading its backend libraries. That happens lazily on first use - rather than at ``import torch``, and the environment lives on GPFS, where it - is tens of seconds of many small reads. Measured on 3A5V: **41.4 s for the - first search against 1.8 s for the same search afterwards**, with 39.4 of - those seconds attributable to no stage at all. - - Without this, whichever arm runs first carries the lot and reads as an order - of magnitude slower than it is. Run at full fidelity rather than on a token - problem, so the kernels and FFT plans the measured searches use are the ones - already paid for. - - Returns the seconds it took. Every row carries it: the cost is real and - worth reporting, it just is not the search's. - """ - model, data, _ = rotated_case(pdb, seed_for(pdb, 0)) - t0 = time.perf_counter() - run_frf(model, data, FRFConfig(n_peaks=n_peaks, lmax_cap=lmax_cap), - capture_arf=False, verbose=0) - return time.perf_counter() - t0 - - -def run_one(pdb: str, trial: int, arm: str, *, n_peaks: int, top_n: int, - thr_deg: float, warmup: bool, mem_interval_s: float, - prewarm_seconds: float = float("nan"), - first_in_process: bool = False) -> dict: - """One measurement: accuracy, memory and runtime for a single search.""" - lmax_cap = ARMS[arm] if ARMS[arm] is not None else _shipped_lmax_cap() - seed = seed_for(pdb, trial) - model, data, R_true = rotated_case(pdb, seed) - cfg = FRFConfig(n_peaks=n_peaks, lmax_cap=lmax_cap) - - if warmup: - run_frf(model, data, cfg, capture_arf=False, verbose=0) - - # Calibrate before the measurement, so a node that is busy *now* is visible. - calib = calibration_seconds() - - mem = PeakMemory(interval_s=mem_interval_s) - t0 = time.perf_counter() - with mem.window() as mem_out, stage_timers() as (totals, counts, unresolved): - res = run_frf(model, data, cfg, capture_arf=False, verbose=0) - wall = time.perf_counter() - t0 - - rank, ang = orbit_rank( - res.peaks, R_true, data.spacegroup.matrices.to(torch.float64).cpu(), - reciprocal_basis=data.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), - side="right", frame="cart", thr_deg=thr_deg, - ) - # Exclusive, not inclusive: the nested stages would otherwise be counted - # twice and "unattributed" could come out negative. - excl = exclusive_times(totals) - attributed = sum(v for v in excl.values() if v == v) - - row = {"experiment": EXPERIMENT, "pdb": pdb, "trial": trial, "arm": arm, - "seed": seed} - row.update(provenance()) - row.update(host_info()) - row.update({ - "spacegroup": str(data.spacegroup.hm), - "n_ops": int(data.spacegroup.matrices.shape[0]), - "n_atoms": int(model.xyz().shape[0]), - "n_reflections": int(data.hkl.shape[0]), - "lmax_cap": lmax_cap, - "n_peaks": n_peaks, - # --- accuracy --- - "truth_rank": rank, - # orbit_rank returns -1 when nothing matched. That must not order as a - # good rank, so for any comparison a miss counts as worse than the worst - # hit, i.e. the length of the peak list. - "rank_for_compare": rank if rank >= 0 else n_peaks, - "found": int(rank >= 0), - "in_top_n": int(0 <= rank < top_n), - "top_n": top_n, - "truth_angle_deg": None if ang is None else round(float(ang), 3), - "n_peaks_found": len(res.peaks), - "orbit_side": "left", "orbit_frame": "cart", "thr_deg": thr_deg, - # --- runtime --- - "timing_kind": "steady" if warmup else "post_warmup", - "prewarm_seconds": round(prewarm_seconds, 2), - # Flagged rather than assumed away: if the warm-up ever misses a shared - # cost, it lands in this row and stays identifiable. - "first_in_process": int(first_in_process), - # Flagged rather than assumed away: if the pre-warm ever misses a shared - # cost, it lands here and is identifiable. - "first_in_process": int(first_in_process), - "seconds": round(wall, 3), - "seconds_attributed": round(attributed, 3), - "seconds_unattributed": round(wall - attributed, 3), - "calibration_seconds": round(calib, 5), - "seconds_per_calibration": round(wall / max(calib, 1e-9), 1), - # --- memory --- - **mem_out, - "stages_unresolved": "|".join(unresolved), - }) - # Inclusive time for reading a single stage, exclusive for adding them up. - for _, attr in FRF_STAGES: - row[f"t_{attr}"] = round(totals.get(attr, float("nan")), 4) - row[f"x_{attr}"] = round(excl.get(attr, float("nan")), 4) - row[f"n_{attr}"] = counts.get(attr, 0) - return row - - -def _fmt(row: dict) -> str: - hit = "yes" if row["in_top_n"] else "NO " - return (f" {row['arm']:<8} rank={str(row['truth_rank']):<6} " - f"top{row['top_n']}={hit} " - f"{row['seconds']:>7.2f}s peak {row['rss_peak_mb']:>8.0f} MB " - f"(+{row['rss_delta_mb']:>7.0f}) {row['seconds_per_calibration']:>7.1f} cal") - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) - ap.add_argument("--trial", type=int, default=None, - help="single trial; omit to run --trials of them") - ap.add_argument("--trials", type=int, default=3) - ap.add_argument("--arms", default="shipped", - help=f"comma-separated, from {sorted(ARMS)}") - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--top-n", type=int, default=20, - help="candidates the placement search carries forward") - ap.add_argument("--thr-deg", type=float, default=5.0) - ap.add_argument("--warmup", action="store_true", - help="discard one search first and report steady state") - ap.add_argument("--no-warmup-run", dest="prewarm", action="store_false", - help="skip the discarded warm-up search; the first " - "measurement then carries the process start-up cost") - ap.add_argument("--mem-interval", type=float, default=0.02, - help="RSS sampling period in seconds") - ap.add_argument("--out-csv", default=None) - args = ap.parse_args() - - arms = [a for a in args.arms.split(",") if a] - unknown = [a for a in arms if a not in ARMS] - if unknown: - raise SystemExit(f"unknown arm(s) {unknown}; expected {sorted(ARMS)}") - trials = [args.trial] if args.trial is not None else list(range(args.trials)) - - csv_path = None - if args.out_csv: - csv_path = Path(args.out_csv) - csv_path.parent.mkdir(parents=True, exist_ok=True) - - prewarm_s = float("nan") - if args.prewarm: - prewarm_s = warmup_run(args.pdb, ARMS[arms[0]] or _shipped_lmax_cap(), - args.n_peaks) - print(f"warm-up search: {prewarm_s:.1f}s -- the process's start-up, " - f"mostly PyTorch loading its backend off GPFS. Excluded from the " - f"measurements below and reported as prewarm_seconds.", flush=True) - - info = host_info() - print(f"{args.pdb}: {len(arms)} arm(s) x {len(trials)} trial(s), " - f"{'steady-state' if args.warmup else 'post-warm-up'}", flush=True) - print(f" host {info['host']} / {info['torch_threads']} threads / " - f"{info['cpu_model'] or 'unknown cpu'}", flush=True) - - rows, n_fail = [], 0 - first = True - for trial in trials: - print(f" trial {trial}", flush=True) - for arm in arms: - try: - row = run_one(args.pdb, trial, arm, n_peaks=args.n_peaks, - top_n=args.top_n, thr_deg=args.thr_deg, - warmup=args.warmup, - mem_interval_s=args.mem_interval, - first_in_process=first) - first = False - except Exception as exc: - n_fail += 1 - print(f" {arm:<8} FAILED {type(exc).__name__}: {exc}", flush=True) - continue - rows.append(row) - if csv_path: - append_row(csv_path, row) - print(_fmt(row), flush=True) - if row["stages_unresolved"]: - print(f" NOT INSTRUMENTED: {row['stages_unresolved']}", - flush=True) - - if rows: - def med(key): - vals = sorted(r.get(key, 0.0) or 0.0 for r in rows) - return vals[len(vals) // 2] - - med_total = med("seconds") - print("\nwhere the time goes (median over the rows above; exclusive of " - "nested stages, so the column sums)", flush=True) - print(f" {'stage':32s} {'excl s':>8s} {'%':>6s} {'incl s':>8s}", - flush=True) - order = sorted((a for _, a in FRF_STAGES), key=lambda a: -med(f"x_{a}")) - for attr in order: - x, t = med(f"x_{attr}"), med(f"t_{attr}") - if t <= 0: - continue - print(f" {attr:32s} {x:8.3f} {100 * x / max(med_total, 1e-9):6.1f} " - f"{t:8.3f}", flush=True) - unatt = med("seconds_unattributed") - print(f" {'(unattributed)':32s} {unatt:8.3f} " - f"{100 * unatt / max(med_total, 1e-9):6.1f}", flush=True) - if csv_path: - print(f"\nwrote {csv_path} ({n_fail} failures)", flush=True) - return 1 if rows == [] else 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_config_sweep.py b/alignment_lab/diagnostics/frf_config_sweep.py deleted file mode 100644 index 9f73ac63..00000000 --- a/alignment_lab/diagnostics/frf_config_sweep.py +++ /dev/null @@ -1,224 +0,0 @@ -"""Settle the FRF's remaining free constants by measurement, one arm each. - -Four engine settings are still switches because nobody chose a value. Making -the rotation search a three-input call means choosing them, and each choice -gets a number first: - -``lmax_cap`` - The signature default is 48, its own docstring claims 100, and the - benchmarks run 64. Phaser's ``DEF_CLMN_LMAX`` is 100. The "high l - under-determines the SH modes" argument for 48 predates the dense P1-box - calc, so it is not evidence about the current engine. -anisotropy - ``production`` is the shipped intensity-space fit; ``legacy_log`` is the - biased log-space fit it replaced; ``iso_only`` keeps only the radial part of - the shipped fit; ``no_aniso`` drops the correction. On this panel - ``production`` and ``no_aniso`` are indistinguishable except on 3GR5, - because most of the structures carry less anisotropy than the estimator can - resolve. -``_orbit_unroll`` - Off, on the strength of a run that predates the reciprocal-space - convention fix, so its evidence is void. -Patterson radius - Never exercised in production. Two structures want radii a factor 2.4 - apart with no rule to pick between them, so the candidate is the *union*: - two runs merged by z-score. It doubles the cost, so it has to earn it. - -Every arm runs in one process per (structure, trial) cell, so the paired -comparison against ``production`` is exact. ``production_dup`` repeats the -baseline arm verbatim: it measures the engine's own run-to-run spread, which -bounds how small a real effect this sweep can resolve. - -Usage ------ - python -m diagnostics.frf_config_sweep --pdb 3GR5 --trial 0 - python -m diagnostics.frf_config_sweep --pdb 3GR5 --trials 10 --stage 2 -""" - -from __future__ import annotations - -import argparse -import sys -import time -from dataclasses import dataclass, field -from pathlib import Path -from typing import Dict, Optional, Tuple - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, FRFConfig, aniso_arm, orbit_rank, # noqa: E402 - rotated_case, run_frf, seed_for, tensor_report) -from lab.results import append_row, provenance # noqa: E402 - -EXPERIMENT = "frf_config_sweep" - -@dataclass(frozen=True) -class Arm: - """One engine configuration to measure.""" - - name: str - lmax_cap: int = 64 - aniso: str = "production" - - def config(self, base: FRFConfig) -> Tuple[FRFConfig, ...]: - return (FRFConfig( - d_min=base.d_min, d_max=base.d_max, n_shells=base.n_shells, - n_peaks=base.n_peaks, lmax_cap=self.lmax_cap, - dense_pad=base.dense_pad, - ),) - - -def _factorial_arms() -> Tuple[Arm, ...]: - """lmax_cap x anisotropy, plus the repeat-baseline control.""" - arms = [Arm("production_dup")] - for cap in (48, 64, 100): - for aniso in ("production", "legacy_log", "iso_only", "no_aniso"): - arms.append(Arm(f"cap{cap}_{aniso}", lmax_cap=cap, aniso=aniso)) - return tuple(arms) - - -#: Stage 2 measured the orbit-dedup unroll and the two-radius Patterson union. -#: Both were arguments of the engine wrapper that the API collapse removed, so -#: those arms cannot be built against this tree; they were measured against the -#: pre-collapse tree (a git worktree pinned at 133bd565) and the outcome is -#: recorded in the changelog. Re-measuring them means restoring the arguments -#: first, which is the point: they are not switches any more. - - -#: The baseline every paired difference is taken against: today's shipped -#: configuration (broken anisotropy fit, cap 64, no unroll, single radius). -BASELINE = Arm("production", lmax_cap=64, aniso="production") - - -def run_one(pdb: str, trial: int, arm: Arm, base: FRFConfig, - *, thr_deg: float) -> dict: - seed = seed_for(pdb, trial) - model, data, R_true = rotated_case(pdb, seed) - configs = arm.config(base) - - captured: dict = {} - cfg = configs[0] - t0 = time.time() - with aniso_arm(arm.aniso, data, d_min=cfg.d_min, d_max=cfg.d_max, - captured=captured): - res = run_frf(model, data, cfg, capture_arf=False, verbose=0) - seconds = time.time() - t0 - peaks = res.peaks - - rank, ang = orbit_rank( - peaks, R_true, data.spacegroup.matrices.to(torch.float64).cpu(), - reciprocal_basis=data.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), - side="right", frame="cart", thr_deg=thr_deg, - ) - row = {"experiment": EXPERIMENT, "pdb": pdb, "trial": trial, - "arm": arm.name, "seed": seed} - row.update(provenance()) - row.update(configs[0].as_row()) - row.update({ - "arm_lmax_cap": arm.lmax_cap, - "arm_aniso": arm.aniso, - "spacegroup": str(data.spacegroup.hm), - "truth_rank": rank, - # orbit_rank returns -1 for "no peak within thr_deg". A miss must not - # sort as a good rank, so for pairing it counts as worse than the worst - # hit, i.e. the length of the peak list. - "rank_for_compare": rank if rank >= 0 else base.n_peaks, - "found": int(rank >= 0), - "in_top20": int(0 <= rank < 20), - "truth_angle_deg": None if ang is None else round(float(ang), 3), - "n_peaks_found": len(peaks), - "orbit_side": "left", "orbit_frame": "cart", "thr_deg": thr_deg, - "seconds": round(seconds, 1), - }) - # Emit both tensor reports for every arm, blank where the arm does not - # produce one: a row carrying columns the file's header lacks is a schema - # error, and silently-widened rows lose exactly these values. - for tag in ("raw", "legacy"): - if tag in captured: - row.update(tensor_report(captured[tag], tag)) - else: - row.update({f"{tag}_B_min": "", f"{tag}_B_max": "", - f"{tag}_B_spread": ""}) - return row - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", required=True, choices=list(BENCH_PDBS)) - ap.add_argument("--trial", type=int, default=None, - help="single trial index; omit to run --trials of them") - ap.add_argument("--trials", type=int, default=10) - ap.add_argument("--stage", type=int, default=1, choices=(1,), - help="1 = lmax x aniso factorial") - ap.add_argument("--d-min", type=float, default=4.0) - ap.add_argument("--d-max", type=float, default=15.0) - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--thr-deg", type=float, default=5.0) - ap.add_argument("--out-csv", default=None) - ap.add_argument("--outdir", default=None) - args = ap.parse_args() - - if args.stage != 1: - raise SystemExit( - "stage 2 measured engine arguments that no longer exist; see the " - "note above _factorial_arms.") - arms = (BASELINE,) + _factorial_arms() - base = FRFConfig(d_min=args.d_min, d_max=args.d_max, n_peaks=args.n_peaks) - - if args.out_csv: - csv_path = Path(args.out_csv) - csv_path.parent.mkdir(parents=True, exist_ok=True) - else: - outdir = Path(args.outdir) if args.outdir else ( - Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT) - outdir.mkdir(parents=True, exist_ok=True) - csv_path = outdir / f"{EXPERIMENT}_{args.pdb}.csv" - - trials = [args.trial] if args.trial is not None else list(range(args.trials)) - print(f"{args.pdb}: stage {args.stage}, {len(arms)} arms x {len(trials)} " - f"trial(s)", flush=True) - - ranks: Dict[str, list] = {a.name: [] for a in arms} - n_fail = 0 - for trial in trials: - for arm in arms: - try: - row = run_one(args.pdb, trial, arm, base, thr_deg=args.thr_deg) - except Exception as exc: - n_fail += 1 - ranks[arm.name].append(None) - print(f" trial {trial} {arm.name}: FAILED {type(exc).__name__}: " - f"{exc}", flush=True) - continue - append_row(csv_path, row) - ranks[arm.name].append(row["rank_for_compare"]) - shown = (str(row["truth_rank"]) if row["found"] - else f"miss@{row['truth_angle_deg']:.0f}deg") - print(f" trial {trial} {arm.name:<26} rank={shown:<12} " - f"top20={row['in_top20']} {row['seconds']:>6.1f}s", flush=True) - - base_ranks = ranks[BASELINE.name] - print("\npaired vs production (negative = better rank):", flush=True) - for arm in arms: - if arm.name == BASELINE.name: - continue - d = [(a - b) for a, b in zip(ranks[arm.name], base_ranks) - if a is not None and b is not None] - if not d: - print(f" {arm.name:<26} no paired trials", flush=True) - continue - sd = sorted(d) - med = (sd[len(sd) // 2] if len(sd) % 2 - else 0.5 * (sd[len(sd) // 2 - 1] + sd[len(sd) // 2])) - print(f" {arm.name:<26} n={len(d):<3} better={sum(x < 0 for x in d)} " - f"same={sum(x == 0 for x in d)} worse={sum(x > 0 for x in d)} " - f"median={med:+.1f} per-trial={d}", flush=True) - print(f"\nwrote {csv_path} ({n_fail} failures)", flush=True) - return 1 if n_fail == len(arms) * len(trials) else 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_encode_compare.py b/alignment_lab/diagnostics/frf_encode_compare.py deleted file mode 100644 index 11044679..00000000 --- a/alignment_lab/diagnostics/frf_encode_compare.py +++ /dev/null @@ -1,509 +0,0 @@ -"""Feed Phaser's own prepared data into our SH-Bessel encoder. - -Every earlier comparison changed two things at once: the *inputs* to the -expansion (normalisation, symmetry unroll, F_calc) and the *expansion itself*. -The per-reflection attribution (job 489442) pinned the input side -- Phaser's -own identities reproduce bit-exactly, DFAC and V are unity, and the residual -disagreement is a roughly uniform ~7-11% in the Wilson normaliser across all -structures, so it does not single out the trigonal/hexagonal failures. That -leaves the encoder untested on its own. - -This runs the encoder with Phaser's inputs, so a mismatch can only come from our -projection: - -* ``obs_phaser_pts`` -- Phaser's prepared observations (``PHASER_OBS_DUMP``: - post-normalisation, post-LERF1, post-unroll, post-axis-permutation, in polar - coordinates, i.e. exactly what ``DataMR::getELMNxR2`` consumes) through - ``bessel_sh_expand``, against Phaser's ``DataElmn``. -* ``calc_phaser_pts`` -- Phaser's molecular-transform samples - (``PHASER_CALC_DUMP``, from ``Ensemble::getELMNxR2`` -- a *different* function - with its own radial scale and its own l != 0 doubling) against Phaser's - ``SearchElmn``. -* ``obs_phaser_clustered`` / ``calc_phaser_clustered`` -- the same two, but - replaying Phaser's own angular approximation. Phaser buckets reflections by - ``|cos(theta) - cos(theta_rep)| < 1e-3`` and evaluates the Legendre functions - once per bucket from the first member's theta (sphericalY.h:43, - DataMR.cc:1096); we cluster only on values equal to ~1e-7. So a high-l - disagreement in the arms above is *expected*, and is Phaser being - approximate rather than us being wrong. These arms separate the two, and - answer a question that has never been asked: whether that 1e-3 polar - smoothing is part of why Phaser is immune to the symmetry-axis ghosts. -* ``obs_ours_unroll`` / ``obs_dedup_unroll`` -- Phaser's ASU-level intensities - (``PHASER_TERMS_DUMP``, keyed by Miller index) put through *our* two symmetry - unrolls: the production one, which emits all ``n_ops`` orbit positions, and - ``SpaceGroup.expand_hkl(include_friedel=False)``, which emits only the - distinct ones as Phaser does (``!duplicate(isym,rhkl)``, DataMR.cc:954). - Same intensities, same encoder, - same target -- so the difference between these two arms is the multiplicity - handling and nothing else. - -Two scalars the expansion needs are not recoverable from the dumped rows -- the -observation-side ``HIRES`` is the *minimum* reso over selected reflections, -one step below the smallest that survives the ``reso(r) > HIRES`` gate -- so the -instrumented binary now writes them to ``.meta`` and they are read, not -inferred. - -Expected relation, if our encoder is right. Phaser projects with ``Y_lm`` -(``e^{+im phi}``, Condon-Shortley sign folded into its ``Pmm`` recurrence) while -we project with ``conj(C(m,phi))``; ``bar_P`` carries no CS phase and our -``sign_m`` restores it. Both are real-weighted sums, so - - ours[n, l, m] = k * conj(phaser[l, m, n+1]), k = 1 (was 2 while the -expansion concatenated the antipodal copy, which Phaser does via cctbx's -conjugate_flag and we no longer do -- see bessel_sh_expand) - -with the factor 2 because appending ``-s`` doubles every even-l coefficient -exactly (``Y_lm(-s) = (-1)^l Y_lm(s)``). ``k`` is therefore a prediction, not a -fitted nuisance: a modulus away from 1 or a phase away from 0 is a finding. - -Usage ------ - python -m diagnostics.frf_encode_compare --pdb 2DQ6 -""" - -from __future__ import annotations - -import argparse -import math -import os -import subprocess -import sys -import time -from pathlib import Path - -import numpy as np -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) - -from lab import case_paths, load_case # noqa: E402 -from lab.phaser_match import PATCHED_PHASER, write_keywords # noqa: E402 -from lab.results import append_row, provenance # noqa: E402 - -EXPERIMENT = "frf_encode_compare" - - -# --------------------------------------------------------------------------- -# Phaser side -# --------------------------------------------------------------------------- - -def run_phaser_dumps(pdb: str, work: Path) -> dict: - """One instrumented run producing every stage this comparison needs.""" - work.mkdir(parents=True, exist_ok=True) - pdb_path, mtz_path = case_paths(pdb) - kw = write_keywords(work, mtz_path=mtz_path, model_pdb=pdb_path, - n_peaks=5, root=f"{pdb}_enc", title=f"encode {pdb}") - paths = { - "PHASER_OBS_DUMP": work / "obs.csv", - "PHASER_CALC_DUMP": work / "calc.csv", - "PHASER_TERMS_DUMP": work / "terms.csv", - "PHASER_DATA_ELMN_DUMP": work / "data_elmn.csv", - "PHASER_SEARCH_ELMN_DUMP": work / "search_elmn.csv", - } - env = dict(os.environ) - for k, v in paths.items(): - env[k] = str(v) - proc = subprocess.run([str(PATCHED_PHASER)], cwd=str(work), - input=kw.read_text(), capture_output=True, - text=True, timeout=5400, env=env) - (work / "run.log").write_text((proc.stdout or "") + (proc.stderr or "")) - missing = [v.name for v in paths.values() if not v.exists()] - if missing: - raise RuntimeError(f"{pdb}: missing dumps {missing}; see {work/'run.log'}") - # "PHASER_OBS_DUMP" -> "obs": strip both the prefix and the _DUMP suffix. - return {k[len("PHASER_"):-len("_DUMP")].lower(): v for k, v in paths.items()} - - -def read_meta(path: Path) -> dict: - """``.meta`` -- the scalars the point list cannot carry.""" - meta = Path(str(path) + ".meta") - if not meta.exists(): - raise RuntimeError( - f"{meta} missing: rebuild the instrumented binary " - f"(phaser_src/build/rebuild.sh) -- the Bessel scale would otherwise " - f"have to be guessed." - ) - out = {} - for line in meta.read_text().splitlines()[1:]: - k, v = line.split(",") - out[k] = float(v) - return out - - -def _cart(r, th, ph) -> torch.Tensor: - return torch.from_numpy(np.stack( - [r * np.sin(th) * np.cos(ph), - r * np.sin(th) * np.sin(ph), - r * np.cos(th)], axis=1)).to(torch.float64) - - -def load_points(path: Path): - """``cluster,r,theta,phi,intensity`` -> our encoder's inputs. - - Returns ``(s_exact, intensity, cos_theta, s_clustered, cluster_stats)``. - - ``s_clustered`` replays Phaser's OWN angular approximation. - ``HKL_clustered::add`` (sphericalY.h:43) buckets reflections greedily by - ``|cos(theta) - cos(theta_rep)| < 1e-3`` against the FIRST member of each - bucket, and the projection then evaluates the associated Legendre functions - once per bucket from that first member's theta (DataMR.cc:1096) -- while the - radial Bessel term stays per-reflection. So Phaser's Y_lm carries up to - 1e-3 of cos-theta error, which at l ~ 70 is a percent-level per-coefficient - error, largest near the poles where sin(theta) is small. - - Our encoder clusters only on values that are equal to ~1e-7, so it is the - more accurate of the two. That means a high-l disagreement is expected and - is Phaser's approximation, not our defect -- and it has to be separated - from a real difference before any residual can be read. Substituting each - point's bucket-representative ``cos(theta)`` while keeping its own ``r`` and - ``phi`` reproduces Phaser's evaluation exactly, because ``r`` and ``phi`` - are the only per-reflection quantities Phaser keeps. - - The dump is written before ``HKL_list.shuffle()``, so row order within a - cluster is insertion order and row 0 of each cluster is the representative. - """ - d = np.loadtxt(path, delimiter=",", skiprows=1) - if d.ndim == 1: - d = d[None, :] - cid = d[:, 0].astype(np.int64) - r, th, ph, val = d[:, 1], d[:, 2], d[:, 3], d[:, 4] - - first = np.zeros(cid.max() + 1, dtype=np.int64) - seen = np.zeros(cid.max() + 1, dtype=bool) - for i, c in enumerate(cid): - if not seen[c]: - seen[c], first[c] = True, i - th_rep = th[first[cid]] - stats = { - "phaser_n_clusters": int(seen.sum()), - "phaser_cos_spread_max": float(np.abs(np.cos(th) - np.cos(th_rep)).max()), - "phaser_cluster_size_max": int(np.bincount(cid).max()), - } - return (_cart(r, th, ph), - torch.from_numpy(val).to(torch.float64), - torch.from_numpy(np.cos(th)).to(torch.float64), - _cart(r, th_rep, ph), - stats) - - -def load_elmn(path: Path, L: int) -> torch.Tensor: - """Phaser's ``l,m,n`` dump into our ``(N_radial, L, 2L-1)`` layout. - - Phaser's ``n`` is 1-based against our 0-based, and both index the same - ``u = l + 2n - 1`` radial order, so ``n0 = n - 1``. - """ - lmax = L - 1 - lmax_even = lmax if lmax % 2 == 0 else lmax - 1 - n_radial = (lmax_even - 2) // 2 + 1 - out = torch.zeros((n_radial, L, 2 * L - 1), dtype=torch.complex128) - d = np.loadtxt(path, delimiter=",", skiprows=1) - if d.ndim == 1: - d = d[None, :] - l = d[:, 0].astype(int) - m = d[:, 1].astype(int) - n0 = d[:, 2].astype(int) - 1 - keep = (l <= lmax_even) & (n0 >= 0) & (n0 < n_radial) & (np.abs(m) <= lmax_even) - if not keep.all(): - raise RuntimeError(f"{path}: {int((~keep).sum())} rows outside the L={L} band") - out[n0, l, m + (L - 1)] = torch.from_numpy(d[:, 3] + 1j * d[:, 4]) - return out - - -def band_mask(L: int, device="cpu") -> torch.Tensor: - """The (n, l, m) entries Phaser allocates: l even, |m| <= l, n < nmax(l).""" - lmax = L - 1 - lmax_even = lmax if lmax % 2 == 0 else lmax - 1 - n_radial = (lmax_even - 2) // 2 + 1 - mask = torch.zeros((n_radial, L, 2 * L - 1), dtype=torch.bool, device=device) - for l in range(2, lmax_even + 1, 2): - n_l = (lmax_even - l) // 2 + 1 - m_lo, m_hi = (L - 1) - l, (L - 1) + l - mask[:n_l, l, m_lo:m_hi + 1] = True - return mask - - -# --------------------------------------------------------------------------- -# comparison -# --------------------------------------------------------------------------- - -def compare_coeffs(ours: torch.Tensor, phaser: torch.Tensor, L: int, - *, k_expected: float) -> dict: - """Our coefficients against ``conj(phaser)`` over Phaser's allocated band. - - Reported quantities: - ``corr`` modulus of the complex correlation -- shape agreement. - ``k_mod``/``k_arg_deg`` the fitted complex scale; the prediction is - ``k_expected`` at 0 degrees, so a phase here means a - convention mismatch, not a scale. - ``rel_resid`` ``||a - k b|| / ||a||`` after the fitted scale, i.e. what - the correlation hides. - ``pow_offband`` fraction of OUR power sitting where Phaser has exactly - zero -- the m-filter / forbidden-m channel. - ``worst_l`` the even l with the lowest per-l correlation. - """ - mask = band_mask(L, device=ours.device) - b = torch.conj(phaser.to(ours.device)) - a = ours - av, bv = a[mask], b[mask] - - num = torch.vdot(bv, av) # sum conj(b)*a - denom_b = (bv.abs() ** 2).sum() - out = { - "n_band": int(mask.sum()), - "n_phaser_nonzero": int((bv.abs() > 0).sum()), - "n_ours_nonzero": int((av.abs() > 0).sum()), - } - if float(denom_b) == 0.0: - out.update(corr=float("nan"), k_mod=float("nan"), - k_arg_deg=float("nan"), rel_resid=float("nan")) - return out - corr = float(num.abs() / (av.norm() * bv.norm()).clamp(min=1e-300)) - k = num / denom_b - resid = (av - k * bv).norm() / av.norm().clamp(min=1e-300) - out.update({ - "corr": corr, - "k_mod": float(k.abs()), - "k_arg_deg": float(torch.rad2deg(torch.angle(k))), - "k_expected": float(k_expected), - "rel_resid": float(resid), - }) - # Power we place where Phaser has none (inside its own band). - zero_b = mask & (b.abs() == 0) - out["pow_offband"] = float( - (a[zero_b].abs() ** 2).sum() / (a[mask].abs() ** 2).sum().clamp(min=1e-300) - ) - # Per-l correlation, to see whether a mismatch is radial (high l) or global. - lmax_even = (L - 1) if (L - 1) % 2 == 0 else (L - 2) - worst_l, worst_c = -1, 2.0 - per_l = [] - for l in range(2, lmax_even + 1, 2): - ml = mask[:, l, :] - al, bl = a[:, l, :][ml], b[:, l, :][ml] - if float(bl.abs().max()) == 0.0: - continue - cl = float(torch.vdot(bl, al).abs() - / (al.norm() * bl.norm()).clamp(min=1e-300)) - per_l.append((l, cl)) - if cl < worst_c: - worst_c, worst_l = cl, l - out["worst_l"] = worst_l - out["worst_l_corr"] = worst_c - out["corr_l2"] = per_l[0][1] if per_l else float("nan") - out["corr_lmax"] = per_l[-1][1] if per_l else float("nan") - return out - - -# --------------------------------------------------------------------------- -# our side -# --------------------------------------------------------------------------- - -def encode(s: torch.Tensor, intensity: torch.Tensor, *, L: int, - h_scale: float, zsymm: int) -> torch.Tensor: - from torchref.experimental.alignment.frf.data_mr import bessel_sh_expand - return bessel_sh_expand( - s, intensity, L=L, bessel_h_scale=h_scale, zsymm=zsymm, - ).coeffs - - -def unroll_arms(pdb: str, terms_csv: Path): - """Phaser's ASU intensities through both of our symmetry unrolls. - - Returns ``(arms, stats)`` where ``arms`` maps name -> ``(s, intensity)``. - The counts are exact integers, so the multiplicity question is answered by - arithmetic before any encoding happens. - """ - d = np.loadtxt(terms_csv, delimiter=",", skiprows=1) - if d.ndim == 1: - d = d[None, :] - hkl = torch.from_numpy(d[:, 0:3]).to(torch.float64) - inten = torch.from_numpy(d[:, 9]).to(torch.float64) - - _, data = load_case(pdb) - sg = data.spacegroup.matrices.to(torch.float64).cpu() - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - n_ops = int(sg.shape[0]) - - # Production: every orbit position, duplicates included (align.py:487). - hkl_all = torch.einsum("kji,nj->kni", sg, hkl).reshape(-1, 3) - s_all = hkl_all @ rec - i_all = inten.unsqueeze(0).expand(n_ops, -1).reshape(-1).contiguous() - - # Phaser-faithful: distinct orbit positions only (DataMR.cc:954). - hkl_ded, asu_idx, _ = data.spacegroup.expand_hkl( - hkl.to(torch.long), include_friedel=False) - s_ded = hkl_ded.to(torch.float64) @ rec - i_ded = inten[asu_idx] - - stats = { - "n_asu_terms": int(hkl.shape[0]), - "n_ops": n_ops, - "n_unroll_all": int(s_all.shape[0]), - "n_unroll_dedup": int(s_ded.shape[0]), - "dup_frac": float(1.0 - s_ded.shape[0] / max(1, s_all.shape[0])), - } - return {"obs_ours_unroll": (s_all, i_all), - "obs_dedup_unroll": (s_ded, i_ded)}, stats - - -def frame_check(s_ours: torch.Tensor, s_phaser: torch.Tensor, - *, n_sample: int = 2000) -> dict: - """Do the two Cartesian reciprocal frames coincide? - - ``|s|`` is frame-independent but theta and phi are not, so an orthogonalisation - convention difference would rotate every coefficient (mixing m) and make the - coefficient comparison meaningless while leaving the radial part intact. - Nearest-neighbour distance in Cartesian space tests position, not just radius. - """ - g = torch.Generator().manual_seed(1) - k = min(n_sample, int(s_phaser.shape[0])) - q = s_phaser[torch.randperm(s_phaser.shape[0], generator=g)[:k]] - dist = torch.empty(k, dtype=torch.float64) - step = 100 # the (step, N, 3) broadcast is the memory bound, not the (step, N) d2 - for i in range(0, k, step): - d2 = ((q[i:i + step, None, :] - s_ours[None, :, :]) ** 2).sum(-1) - dist[i:i + step] = d2.min(1).values.clamp(min=0).sqrt() - return {"frame_median_dist": float(dist.median()), - "frame_max_dist": float(dist.max()), - "frame_matched_frac": float((dist < 1e-9).to(torch.float64).mean())} - - -# --------------------------------------------------------------------------- -# driver -# --------------------------------------------------------------------------- - -def run(pdb: str, outdir: Path, *, reuse: Path | None = None) -> list: - work = reuse if reuse is not None else (outdir / "phaser" / pdb) - if reuse is not None: - dumps = {n: work / f"{n}.csv" for n in - ("obs", "calc", "terms", "data_elmn", "search_elmn")} - missing = [str(p) for p in dumps.values() if not p.exists()] - if missing: - raise RuntimeError(f"--reuse given but missing: {missing}") - else: - t0 = time.time() - dumps = run_phaser_dumps(pdb, work) - print(f" phaser dumps in {time.time()-t0:.0f}s", flush=True) - - obs_meta = read_meta(dumps["obs"]) - calc_meta = read_meta(dumps["calc"]) - L = int(obs_meta["lmax"]) + 1 - zsymm = int(obs_meta["zsymm"]) - h_obs = obs_meta["lmax"] * obs_meta["hires"] - h_calc = calc_meta["lmax"] * calc_meta["max_resolution"] - print(f" L={L} zsymm={zsymm} axis={int(obs_meta['axis'])} " - f"hires={obs_meta['hires']:.6f} h_obs={h_obs:.4f} " - f"h_calc={h_calc:.4f} (calc max_reso={calc_meta['max_resolution']:.6f})", - flush=True) - if int(obs_meta["axis"]) != 3: - print(" NOTE: high-order axis is not c -- Phaser permutes the frame " - "(DataMR.cc:984) and we do not; the obs arm is confounded.", - flush=True) - - data_elmn = load_elmn(dumps["data_elmn"], L) - search_elmn = load_elmn(dumps["search_elmn"], L) - s_obs, i_obs, _, s_obs_clu, obs_clu = load_points(dumps["obs"]) - s_calc, i_calc, cos_calc, s_calc_clu, calc_clu = load_points(dumps["calc"]) - calc_clu = {f"calc_{k}": v for k, v in calc_clu.items()} - print(f" phaser obs clusters={obs_clu['phaser_n_clusters']} " - f"(max cos-theta spread {obs_clu['phaser_cos_spread_max']:.2e}, " - f"largest bucket {obs_clu['phaser_cluster_size_max']}); calc clusters=" - f"{calc_clu['calc_phaser_n_clusters']} " - f"(spread {calc_clu['calc_phaser_cos_spread_max']:.2e})", flush=True) - - base = {"experiment": EXPERIMENT, "pdb": pdb} - base.update(provenance()) - base.update({"L": L, "zsymm": zsymm, "axis": int(obs_meta["axis"]), - "hires": obs_meta["hires"], "h_obs": h_obs, "h_calc": h_calc}) - - rows = [] - - def emit(arm, coeffs, target, *, n_points, seconds, extra=None): - r = dict(base, arm=arm, n_points=int(n_points), - seconds=round(seconds, 1)) - r.update(compare_coeffs(coeffs, target, L, k_expected=1.0)) - if extra: - r.update(extra) - rows.append(r) - - # --- arm 1: Phaser's own observations through our encoder --------------- - t0 = time.time() - c = encode(s_obs, i_obs, L=L, h_scale=h_obs, zsymm=zsymm) - emit("obs_phaser_pts", c, data_elmn, - n_points=s_obs.shape[0], seconds=time.time() - t0, extra=obs_clu) - - # --- arm 2: Phaser's observations WITH Phaser's own theta approximation -- - t0 = time.time() - c = encode(s_obs_clu, i_obs, L=L, h_scale=h_obs, zsymm=zsymm) - emit("obs_phaser_clustered", c, data_elmn, - n_points=s_obs_clu.shape[0], seconds=time.time() - t0, extra=obs_clu) - - # --- arm 3: Phaser's calc samples through our encoder ------------------- - # Phaser doubles every l != 0 grid point (Ensemble.cc: `flipped`), because - # its molecular-transform grid stores only the l >= 0 hemisphere. l == 0 is - # exactly the s_z == 0 plane for these settings (c* along z), so the flag is - # recoverable from the dumped theta. - t0 = time.time() - flip = (cos_calc.abs() > 1e-12).to(torch.float64) + 1.0 - c = encode(s_calc, i_calc * flip, L=L, h_scale=h_calc, zsymm=1) - emit("calc_phaser_pts", c, search_elmn, n_points=s_calc.shape[0], - seconds=time.time() - t0, - extra=dict(calc_clu, n_l0_plane=int((cos_calc.abs() <= 1e-12).sum()))) - - # --- arm 4: the same, with Phaser's theta approximation ----------------- - t0 = time.time() - c = encode(s_calc_clu, i_calc * flip, L=L, h_scale=h_calc, zsymm=1) - emit("calc_phaser_clustered", c, search_elmn, n_points=s_calc_clu.shape[0], - seconds=time.time() - t0, extra=calc_clu) - - # --- arms 5/6: Phaser's ASU intensities through OUR unrolls ------------- - arms, ustats = unroll_arms(pdb, dumps["terms"]) - print(f" unroll: asu={ustats['n_asu_terms']} x n_ops={ustats['n_ops']} " - f"= {ustats['n_unroll_all']} all / {ustats['n_unroll_dedup']} dedup " - f"(phaser obs rows {int(s_obs.shape[0])}; " - f"dup_frac {ustats['dup_frac']:.4f})", flush=True) - for arm, (s, val) in arms.items(): - t0 = time.time() - fc = frame_check(s, s_obs) - c = encode(s, val, L=L, h_scale=h_obs, zsymm=zsymm) - emit(arm, c, data_elmn, n_points=s.shape[0], seconds=time.time() - t0, - extra=dict(ustats, **fc)) - - # One schema for every row -- csv.DictWriter fixes fieldnames from the - # first row it sees, so a later row carrying extra keys would raise. - cols = {} - for r in rows: - cols.update({k: "" for k in r}) - return [{**cols, **r} for r in rows] - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", required=True) - ap.add_argument("--outdir", default=None) - ap.add_argument("--reuse", default=None, - help="directory of existing dumps (skips the Phaser run)") - args = ap.parse_args() - - outdir = Path(args.outdir) if args.outdir else ( - Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT) - outdir.mkdir(parents=True, exist_ok=True) - csv_path = outdir / f"{EXPERIMENT}_{args.pdb}.csv" - - rows = run(args.pdb, outdir, - reuse=Path(args.reuse) if args.reuse else None) - hdr = (f"{'arm':<20}{'n_pts':>9}{'corr':>9}{'k_mod':>9}{'k_arg':>9}" - f"{'resid':>9}{'offband':>9}{'worst_l':>9}{'wl_corr':>9}") - print(f"\n{args.pdb}: our encoder against Phaser's coefficients", flush=True) - print(hdr, flush=True) - for r in rows: - append_row(csv_path, r) - print(f"{r['arm']:<20}{r['n_points']:>9}{r['corr']:>9.4f}" - f"{r['k_mod']:>9.4f}{r['k_arg_deg']:>9.2f}{r['rel_resid']:>9.4f}" - f"{r['pow_offband']:>9.4f}{r['worst_l']:>9}" - f"{r['worst_l_corr']:>9.4f}", flush=True) - print(f"\nwrote {csv_path}", flush=True) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_ghost_knockout.py b/alignment_lab/diagnostics/frf_ghost_knockout.py deleted file mode 100644 index 63b812bd..00000000 --- a/alignment_lab/diagnostics/frf_ghost_knockout.py +++ /dev/null @@ -1,198 +0,0 @@ -"""Find which obs-side term creates the trigonal/hexagonal ghosts. - -Context. With Phaser's bandwidth, resolution and SO(3) sampling pinned, our FRF -puts truth at rank 0 on five of seven benchmark structures and beats Phaser on -6G9X -- but collapses on 2DQ6 (P 3_1 2 1, truth rank 76799) and 3GR5 -(P 6_5 2 2, rank 16855), where Phaser gets rank 0 with a healthy margin. Those -are the only two cases with a 120 degree cell and a 3-fold-containing axis; the -working set is monoclinic / orthorhombic / tetragonal / cubic. - -Because lmax and sampling were already pinned to Phaser's own values when the -collapse was measured, bandwidth is eliminated. What remains is the observation- -and calc-side preprocessing chain, where our engine differs from -``DataMR::getELMNxR2`` in ways that are individually documented but never -isolated on these two space groups: - -* ``use_epsilon=False`` -- Phaser always normalises by ``sqrt(eps_n * Sigma_N)`` - (DataMR.cc:930). Without the epsilon divisor the axial/zonal reflections are - over-weighted, and for a 6-fold axis those carry eps up to 6-12. -* ``_orbit_unroll=False`` -- Phaser expands over symmetry but skips duplicate - P1 indices (``!duplicate(isym,rhkl)``, DataMR.cc:954). We replicate all n_ops - unconditionally, so reflections on special positions are counted several times. -* the m-symmetry filter, French-Wilson, shell-variance weights, Wilson-B match, - Oeffner vrms and the Babinet bulk-solvent term. - -Each arm flips exactly one of these against a common baseline and reports where -truth lands. The discriminating statistic is ``margin`` -- truth's sigma minus -the strongest non-truth peak's. Negative means the rotation function prefers a -ghost, which is the failure we are chasing; rank alone hides how close the call -was. - -Usage ------ - python -m diagnostics.frf_ghost_knockout --pdb 2DQ6 - python -m diagnostics.frf_ghost_knockout --pdb 3GR5 --arm epsilon -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) - -from lab import FRFConfig, load_case, patched, run_frf # noqa: E402 -from lab.results import append_row, provenance # noqa: E402 - -EXPERIMENT = "frf_ghost_knockout" - -#: Phaser's own bandwidth / resolution / sampling per case, read from the -#: instrumented run (job 487737). Pinned so every arm differs only in the term -#: under test -- and so the result is comparable with the map-comparison run. -PHASER_PINNED = { - "1DAW": dict(lmax=58, sampling_deg=6.233148, d_min_eff=5.70), - "1AK5": dict(lmax=70, sampling_deg=5.155428, d_min_eff=4.84), - "3K7M": dict(lmax=66, sampling_deg=5.464844, d_min_eff=5.39), - "4BX9": dict(lmax=96, sampling_deg=3.774036, d_min_eff=6.73), - "6G9X": dict(lmax=84, sampling_deg=4.314963, d_min_eff=5.92), - "2DQ6": dict(lmax=76, sampling_deg=4.772056, d_min_eff=6.36), - "3GR5": dict(lmax=66, sampling_deg=5.483967, d_min_eff=4.10), -} - -#: One flipped term per arm. ``baseline`` is our production configuration. -ARMS = { - "baseline": {}, - "epsilon": dict(use_epsilon=True), - "orbit_unroll": dict(_orbit_unroll=True), - "eps+unroll": dict(use_epsilon=True, _orbit_unroll=True), - "no_m_filter": dict(frf_use_m_filter=False), - "no_french": dict(frf_use_french_wilson=False), - "no_shellvar": dict(frf_use_shell_variance=False), - "no_bulk_solv": dict(apply_bulk_solvent=False), - "no_wilson_b": dict(apply_wilson_b=False), - "vrms_fixed": dict(vrms_strategy="fixed"), - "acentric_only": dict(frf_acentric_only=True), -} - - -def _orbit_of_identity(data): - """Truth is the identity, up to the point group, in the Cartesian frame.""" - from lab.truth import symmetry_orbit - - symops = data.spacegroup.matrices.to(torch.float64).cpu() - recip = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - return symmetry_orbit( - torch.eye(3, dtype=torch.float64), symops, - side="right", frame="cart", reciprocal_basis=recip, - ) - - -def _truth_and_margin(arf, orbit): - """``(rank, sigma, angle, best_ghost_sigma, margin)`` for one map.""" - from torchref.base.alignment.rotation import rotation_matrix_euler_zyz - - v = arf.values.to(torch.float64).cpu() - R = rotation_matrix_euler_zyz(torch.stack([ - arf.alphas.to(torch.float64).cpu(), - arf.betas.to(torch.float64).cpu(), - arf.gammas.to(torch.float64).cpu(), - ], dim=-1)) - sig = (v - v.mean()) / v.std().clamp(min=1e-30) - - best = None - tol = 1.5 * float(arf.grid_sampling_deg) - truth_mask = torch.zeros(v.numel(), dtype=torch.bool) - for k in range(orbit.shape[0]): - tr = torch.einsum("nij,ij->n", R, orbit[k]) - ang = torch.rad2deg(torch.arccos(((tr - 1.0) / 2.0).clamp(-1.0, 1.0))) - truth_mask |= ang <= tol - j = int(torch.argmin(ang)) - if best is None or v[j] > v[best[0]]: - best = (j, float(ang[j])) - j, ang_j = best - rank = int((v > v[j]).sum()) - ghost_sig = sig.masked_fill(truth_mask, float("-inf")) - gs = float(ghost_sig.max()) - return rank, float(sig[j]), ang_j, gs, float(sig[j]) - gs - - -def run_arm(pdb: str, arm: str, extra: dict, *, n_peaks: int = 500) -> dict: - """One knockout arm on one structure.""" - from torchref.experimental.alignment.frf import api as _api - - pin = PHASER_PINNED[pdb] - model, data = load_case(pdb) - orbit = _orbit_of_identity(data) - - def _pinned(model_radius_A, d_min_data, lmax_cap=48): - return int(pin["lmax"]) + 1, float(pin["d_min_eff"]) - - cfg = FRFConfig( - n_peaks=n_peaks, lmax_cap=int(pin["lmax"]), - extra={"grid_sampling_deg": float(pin["sampling_deg"]), **extra}, - ) - t0 = time.time() - with patched(_api, "phaser_lmax_resolution", _pinned): - res = run_frf(model, data, cfg, capture_arf=True, verbose=0) - rank, sig, ang, ghost, margin = _truth_and_margin(res.arf, orbit) - - row = {"experiment": EXPERIMENT, "pdb": pdb, "arm": arm} - row.update(provenance()) - row.update({ - "spacegroup": str(data.spacegroup.hm), - "n_orbit": int(orbit.shape[0]), - "lmax": pin["lmax"], - "sampling_deg": pin["sampling_deg"], - "d_min_eff": pin["d_min_eff"], - "n_samples": int(res.arf.values.numel()), - "truth_rank": rank, - "truth_sigma": round(sig, 4), - "truth_angle_deg": round(ang, 3), - "best_ghost_sigma": round(ghost, 4), - "margin": round(margin, 4), - "map_max_sigma": round(float(res.map_max_sigma), 4), - "seconds": round(time.time() - t0, 1), - }) - row.update({f"flag_{k}": v for k, v in extra.items()}) - return row - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", required=True, choices=sorted(PHASER_PINNED)) - ap.add_argument("--arm", help="single arm (default: all)") - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--outdir", default=None) - args = ap.parse_args() - - arms = {args.arm: ARMS[args.arm]} if args.arm else ARMS - outdir = Path(args.outdir) if args.outdir else ( - Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT - ) - outdir.mkdir(parents=True, exist_ok=True) - csv_path = outdir / f"{EXPERIMENT}_{args.pdb}.csv" - - print(f"{args.pdb}: truth rank / margin per knockout arm", flush=True) - print(f"{'arm':<15}{'rank':>9}{'truth_sig':>11}{'ghost_sig':>11}{'margin':>9}", - flush=True) - failures = 0 - for arm, extra in arms.items(): - try: - row = run_arm(args.pdb, arm, extra, n_peaks=args.n_peaks) - except Exception as exc: - failures += 1 - print(f"{arm:<15} FAILED: {type(exc).__name__}: {exc}", flush=True) - continue - append_row(csv_path, row) - print(f"{arm:<15}{row['truth_rank']:>9}{row['truth_sigma']:>11.2f}" - f"{row['best_ghost_sigma']:>11.2f}{row['margin']:>+9.2f}", flush=True) - print(f"\nwrote {csv_path}", flush=True) - return 1 if failures == len(arms) else 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_inject_phaser_obs.py b/alignment_lab/diagnostics/frf_inject_phaser_obs.py deleted file mode 100644 index 468a899d..00000000 --- a/alignment_lab/diagnostics/frf_inject_phaser_obs.py +++ /dev/null @@ -1,211 +0,0 @@ -"""Run our FRF on Phaser's observation intensities and see where truth lands. - -This is the closing experiment of the bisection. Everything downstream of the -per-reflection intensity is now verified exact against Phaser: - -* the SH-Bessel projection -- feeding Phaser's own prepared observations and its - own molecular-transform samples through ``bessel_sh_expand`` reproduces its - ``DataElmn`` and ``SearchElmn`` at correlation 1.0000, scale 1.0000 at 0 - degrees, residual 0.0000, once Phaser's 1e-3 cos-theta bucketing is replayed - (job 489517); -* the Wigner contraction, per-beta FFT, interpolation and adaptive sample list - -- Phaser's ``clmn`` through our evaluator gives r = 0.998 with an identical - argmax; -* the reciprocal frame (positions agree to 1e-8) and the unroll (our - the orbit dedup reproduces Phaser's point count exactly). - -What is NOT verified is the intensity attached to each position: ours correlates -with Phaser's at 0.988 (1AK5), 0.877 (2DQ6) and 0.711 (3GR5) -- and that ordering -is the performance ordering. So substituting Phaser's intensities into our -otherwise-unchanged pipeline is a decisive test rather than another correlation: -if truth reaches rank 0 on 3GR5, the remaining deficit is entirely in the -intensity computation and nothing else is left to look for. If it does not, there -is a defect outside everything measured so far. - -Alignment is by Miller index, not by position. Phaser dumps one row per selected -ASU reflection (``PHASER_TERMS_DUMP``); expanding each over the orbit ``h.W`` -gives the P1 index of every point that reflection contributes, and the Friedel -mate carries the same intensity. Reflections our engine keeps but Phaser did not -select are reported rather than silently dropped. - -Usage ------ - python -m diagnostics.frf_inject_phaser_obs --pdb 3GR5 \ - --dumps ../runs/encode_compare_489514/phaser/3GR5 -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import numpy as np -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) - -from lab import FRFConfig, load_case, patched, run_frf # noqa: E402 -from lab.results import append_row, provenance # noqa: E402 -from diagnostics.frf_ghost_knockout import ( # noqa: E402 - PHASER_PINNED, _orbit_of_identity, _truth_and_margin, -) - -EXPERIMENT = "frf_inject_phaser_obs" - -#: Packing base for (h,k,l) -> one int64 key. Miller indices here stay well -#: inside +-1000 at these resolutions, and the base is checked at build time. -_BASE = 2048 - - -def _pack(hkl: torch.Tensor) -> torch.Tensor: - if int(hkl.abs().max()) >= _BASE // 2: - raise ValueError(f"Miller index {int(hkl.abs().max())} too large for base {_BASE}") - h, k, l = hkl[:, 0], hkl[:, 1], hkl[:, 2] - return ((h + _BASE // 2) * _BASE + (k + _BASE // 2)) * _BASE + (l + _BASE // 2) - - -def build_lut(terms_csv: Path, sym_mats: torch.Tensor): - """Phaser's per-reflection intensity, keyed by every P1 index it feeds. - - Each ASU row is expanded over the orbit ``h.W`` (Phaser's ``rotMiller`` is - ``rotsym[isym] * h`` with ``rotsym = W^T``, i.e. the row-vector convention) - and over the Friedel mate, which carries the same intensity because - ``|F(-h)| = |F(h)|`` and even-l-only projection is blind to the sign. - """ - d = np.loadtxt(terms_csv, delimiter=",", skiprows=1) - if d.ndim == 1: - d = d[None, :] - hkl = torch.from_numpy(d[:, 0:3]).to(torch.float64) - inten = torch.from_numpy(d[:, 9]).to(torch.float64) - - orbits = torch.einsum("kji,nj->nki", sym_mats.to(torch.float64), hkl) - orbits = orbits.round().to(torch.long).reshape(-1, 3) - vals = inten.unsqueeze(1).expand(-1, sym_mats.shape[0]).reshape(-1) - keys = torch.cat([_pack(orbits), _pack(-orbits)]) - vals = torch.cat([vals, vals]) - - uniq, inverse = torch.unique(keys, return_inverse=True) - lut = torch.zeros(uniq.numel(), dtype=torch.float64) - lut[inverse] = vals - # Two different ASU reflections mapping onto one P1 index would mean the - # orbit is still wrong -- the signature of the convention bug that was fixed. - hi = torch.full((uniq.numel(),), -1e300, dtype=torch.float64) - lo = torch.full((uniq.numel(),), 1e300, dtype=torch.float64) - hi.scatter_reduce_(0, inverse, vals, reduce="amax") - lo.scatter_reduce_(0, inverse, vals, reduce="amin") - n_conflict = int(((hi - lo).abs() > 1e-12 * hi.abs().clamp(min=1e-30)).sum()) - stats = {"n_asu_terms": int(hkl.shape[0]), - "n_lut_keys": int(uniq.numel()), - "n_lut_conflicts": n_conflict} - return uniq, lut, stats - - -def run_arm(pdb: str, dumps: Path, arm: str, *, n_peaks: int = 500) -> dict: - """One arm: ``baseline`` (our intensities) or ``phaser_intensity``.""" - from torchref.experimental.alignment.frf import api as _api - - pin = PHASER_PINNED[pdb] - model, data = load_case(pdb) - orbit = _orbit_of_identity(data) - sg = data.spacegroup.matrices.to(torch.float64).cpu() - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - rec_inv = torch.linalg.inv(rec) - - report = {"n_intercepted": 0} - original = _api.bessel_sh_expand - - if arm == "phaser_intensity": - keys, lut, lut_stats = build_lut(dumps / "terms.csv", sg) - - def injected(s, intensity, **kw): - if kw.get("zsymm", 1) <= 1: # calc side: untouched - return original(s, intensity, **kw) - report["n_intercepted"] += 1 - hkl = (s.to(torch.float64).cpu() @ rec_inv).round().to(torch.long) - k = _pack(hkl) - pos = torch.searchsorted(keys, k) - pos_c = pos.clamp(max=keys.numel() - 1) - found = keys[pos_c] == k - report["n_obs"] = int(s.shape[0]) - report["found_frac"] = float(found.to(torch.float64).mean()) - new = lut[pos_c].to(s.dtype).to(s.device) - return original(s[found], new[found], **kw) - - ctx_name, ctx_val = "bessel_sh_expand", injected - else: - lut_stats = {} - ctx_name, ctx_val = "bessel_sh_expand", original - - def _pinned(model_radius_A, d_min_data, lmax_cap=48): - return int(pin["lmax"]) + 1, float(pin["d_min_eff"]) - - cfg = FRFConfig(n_peaks=n_peaks, lmax_cap=int(pin["lmax"]), - extra={"grid_sampling_deg": float(pin["sampling_deg"])}) - t0 = time.time() - with patched(_api, "phaser_lmax_resolution", _pinned), \ - patched(_api, ctx_name, ctx_val): - res = run_frf(model, data, cfg, capture_arf=True, verbose=0) - rank, sig, ang, ghost, margin = _truth_and_margin(res.arf, orbit) - - row = {"experiment": EXPERIMENT, "pdb": pdb, "arm": arm} - row.update(provenance()) - row.update({ - "spacegroup": str(data.spacegroup.hm), - "lmax": pin["lmax"], "sampling_deg": pin["sampling_deg"], - "d_min_eff": pin["d_min_eff"], - "n_samples": int(res.arf.values.numel()), - "truth_rank": rank, "truth_sigma": round(sig, 4), - "truth_angle_deg": round(ang, 3), - "best_ghost_sigma": round(ghost, 4), "margin": round(margin, 4), - "seconds": round(time.time() - t0, 1), - }) - row.update(lut_stats) - row.update(report) - if arm == "phaser_intensity" and report["n_intercepted"] != 1: - raise RuntimeError( - f"{pdb}: intercepted {report['n_intercepted']} obs expansions, " - f"expected exactly 1 -- the injection hook is on the wrong call") - return row - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", required=True, choices=sorted(PHASER_PINNED)) - ap.add_argument("--dumps", required=True, - help="directory holding Phaser's terms.csv for this pdb") - ap.add_argument("--outdir", default=None) - ap.add_argument("--n-peaks", type=int, default=500) - args = ap.parse_args() - - outdir = Path(args.outdir) if args.outdir else ( - Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT) - outdir.mkdir(parents=True, exist_ok=True) - csv_path = outdir / f"{EXPERIMENT}_{args.pdb}.csv" - - print(f"{args.pdb}: truth rank with our vs Phaser's obs intensities", - flush=True) - print(f"{'arm':<18}{'rank':>8}{'truth_sig':>11}{'ghost_sig':>11}" - f"{'margin':>9}{'found':>8}", flush=True) - rows = [] - for arm in ("baseline", "phaser_intensity"): - rows.append(run_arm(args.pdb, Path(args.dumps), arm, - n_peaks=args.n_peaks)) - cols = {} - for r in rows: - cols.update({k: "" for k in r}) - for r in rows: - full = {**cols, **r} - append_row(csv_path, full) - ff = full.get("found_frac", "") - ff = f"{float(ff):>8.4f}" if ff != "" else f"{'-':>8}" - print(f"{r['arm']:<18}{r['truth_rank']:>8}{r['truth_sigma']:>11.2f}" - f"{r['best_ghost_sigma']:>11.2f}{r['margin']:>+9.2f}{ff}", - flush=True) - print(f"\nwrote {csv_path}", flush=True) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_map_compare.py b/alignment_lab/diagnostics/frf_map_compare.py deleted file mode 100644 index 0bf40849..00000000 --- a/alignment_lab/diagnostics/frf_map_compare.py +++ /dev/null @@ -1,428 +0,0 @@ -"""Compare our FRF array against Phaser's, element-wise, on one shared grid. - -Both engines evaluate a rotation function on Phaser's adaptive SO(3) sample -list, so with the sampling matched the two arrays are directly subtractable -- -index for index, no interpolation. That makes this a much sharper instrument -than comparing peak lists: a peak-list comparison only sees where the maxima -landed, whereas this sees the whole surface, including how much of the -disagreement is a smooth scale/offset (harmless -- peak *order* is invariant to -an affine map) versus genuine reshaping (not harmless). - -What is pinned, and what is not -------------------------------- -Three quantities are read from Phaser's own VERBOSE log and forced onto our -engine: the bandwidth ``lmax``, the resolution the expansion runs at, and the -SO(3) sampling step. They are pinned from the log rather than re-derived, -because they all descend from ``mean_radius()`` and our reimplementation of that -is ~4% off (see :func:`lab.phaser_match.phaser_mean_radius`). - -Everything on the observation and calc side keeps our production defaults -- -anisotropy correction, symmetry unroll, French-Wilson, shell-variance weights, -Wilson-B match, Oeffner vrms, bulk solvent, dense P1-box calc. That is -deliberate: with the coupled trio pinned, whatever disagreement remains is -attributable to that preprocessing stack, which is what we want localised. - -Usage ------ - python -m diagnostics.frf_map_compare --pdb 1DAW - python -m diagnostics.frf_map_compare --worklist-index 3 -""" - -from __future__ import annotations - -import argparse -import math -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) - -from lab import BENCH_PDBS, FRFConfig, case_paths, load_case, patched, run_frf # noqa: E402 -from lab.phaser_match import ( # noqa: E402 - PATCHED_PHASER, - phaser_mean_radius_from_sampling, - phaser_sampling_from_dump, - load_phaser_frame, - load_phaser_map, - parse_phaser_log, - phaser_frf_params, - phaser_mean_radius, - run_patched_phaser, - write_keywords, -) -from lab.results import append_row, provenance # noqa: E402 - -EXPERIMENT = "frf_map_compare" - - -def run_phaser_side(pdb: str, work: Path, *, n_peaks: int = 20) -> dict: - """Run the patched binary and return its parameters plus its map. - - Returns a dict with ``angles`` (N,3 degrees, our sign convention), - ``values`` (N,), and the parsed log fields. - - Raises - ------ - RuntimeError - If the rotation search was short-circuited by the R-factor check, or the - dump is missing -- both of which Phaser reports as ``EXIT STATUS: - SUCCESS``, so they must be checked explicitly. - """ - pdb_path, mtz_path = case_paths(pdb) - dump = work / f"{pdb}_phaser_map.csv" - kw = write_keywords( - work, mtz_path=mtz_path, model_pdb=pdb_path, n_peaks=n_peaks, - root=f"{pdb}_frf", title=f"FRF map dump {pdb}", - ) - rc, seconds, log_path = run_patched_phaser(work, kw, dump_path=dump) - info = parse_phaser_log(log_path) - - if info.get("rotation_search_skipped"): - raise RuntimeError( - f"{pdb}: Phaser skipped the rotation search (R-factor short-circuit) " - f"-- see {log_path}" - ) - if not dump.exists(): - raise RuntimeError( - f"{pdb}: no FRF dump written (rc={rc}); is {PATCHED_PHASER} the " - f"instrumented binary? see {log_path}" - ) - for key in ("lmax", "sampling_deg"): - if key not in info: - raise RuntimeError(f"{pdb}: could not parse {key} from {log_path}") - - angles, values = load_phaser_map(dump) - # The logged sampling is rounded to 2 decimals and cannot rebuild the grid; - # the dumped beta step is exact. - info["sampling_deg_logged"] = info["sampling_deg"] - info["sampling_deg"] = phaser_sampling_from_dump(angles) - info.update(angles=angles, values=values, seconds=seconds, log=log_path) - return info - - -def run_our_side(pdb: str, info: dict, *, n_peaks: int = 500): - """Run our FRF with Phaser's bandwidth, resolution and sampling pinned. - - ``phaser_lmax_resolution`` is the single choke point through which the - bandwidth and resolution reach both the SH expansion and the dense calc - grid, so overriding it pins both consistently. - """ - from torchref.experimental.alignment.frf import api as _api - - model, data = load_case(pdb) - - lmax = int(info["lmax"]) - # Resolution the expansion runs at: LMAX_RESO when Phaser's cap bound, - # otherwise the selected high-resolution limit. - if info.get("lmax_reso_A") is not None and not info.get("all_data_to_limit", False): - d_min_eff = float(info["lmax_reso_A"]) - else: - d_min_eff = float(info.get("selected_d_min", 0.0)) or None - if d_min_eff is None: - raise RuntimeError(f"{pdb}: cannot determine Phaser's expansion resolution") - - def _pinned(model_radius_A, d_min_data, lmax_cap=48): - # Our bandwidth convention is L = lmax + 1. - return lmax + 1, d_min_eff - - cfg = FRFConfig( - n_peaks=n_peaks, - lmax_cap=lmax, - extra={"grid_sampling_deg": float(info["sampling_deg"])}, - ) - with patched(_api, "phaser_lmax_resolution", _pinned): - result = run_frf(model, data, cfg, capture_arf=True, verbose=0) - return model, data, result, d_min_eff - - -def _orbit_of_identity(data): - """Rotations equivalent to the deposited orientation under the point group. - - The search model is used unrotated, so "truth" is the identity -- but only - up to the crystal point group, and the operators must be taken to the - Cartesian frame the rotation function works in (mixing Cartesian rotations - with fractional operators is a metric error that inflates ghost counts). - """ - from lab.truth import symmetry_orbit - - symops = data.spacegroup.matrices.to(torch.float64).cpu() - recip = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - I = torch.eye(3, dtype=torch.float64) - return symmetry_orbit(I, symops, side="right", frame="cart", - reciprocal_basis=recip) - - -def _nearest(R_all: torch.Tensor, R_target: torch.Tensor): - """Index of the sample rotation closest to ``R_target``, and the angle.""" - tr = torch.einsum("nij,ij->n", R_all, R_target) - ang = torch.rad2deg(torch.arccos(((tr - 1.0) / 2.0).clamp(-1.0, 1.0))) - j = int(torch.argmin(ang)) - return j, float(ang[j]) - - -def _truth_rank(values: torch.Tensor, R_all: torch.Tensor, orbit: torch.Tensor): - """Rank of the best sample lying on the truth orbit. - - Returns ``(rank, sigma, angle_deg)``. The rank is the number of samples - scoring strictly higher, i.e. 0 means the rotation function put truth first. - """ - best = None - for k in range(orbit.shape[0]): - j, ang = _nearest(R_all, orbit[k]) - if best is None or values[j] > values[best[0]]: - best = (j, ang) - j, ang = best - rank = int((values > values[j]).sum()) - sig = float((values[j] - values.mean()) / values.std().clamp(min=1e-30)) - return rank, sig, ang - - -def compare(ours, phaser_angles, phaser_values, frame, data, *, topn: int = 20) -> dict: - """Compare two rotation functions in a common (PDB) frame. - - Element-wise comparison is impossible here and it is worth being explicit - about why: Phaser samples SO(3) on a grid laid out in the search model's - **principal frame**, we sample the identically-shaped grid in the PDB frame, - and the two are related by a rotation (``PR``, and ``axisrot``). The grids - therefore have the same pitch and the same point count but cover *different* - rotations -- nearest-neighbour offsets run about half a grid step. On peaks - only ~6-10 deg wide that annihilates any sample-wise correlation while - leaving the peak structure intact, so a whole-map Pearson r measures nothing - but the frame offset. Everything below is therefore computed on rotations, - via nearest-neighbour lookup, not on indices. - - The headline numbers are ``truth_rank_ours`` and ``truth_rank_phaser``: if - ours is much worse, the ghost problem is ours; if they agree, the ghosts are - inherent to the target function and no reimplementation will remove them. - """ - from torchref.base.alignment.rotation import rotation_matrix_euler_zyz - - a = ours.arf - ov = a.values.to(torch.float64).cpu() - R_ours = rotation_matrix_euler_zyz(torch.stack([ - a.alphas.to(torch.float64).cpu(), - a.betas.to(torch.float64).cpu(), - a.gammas.to(torch.float64).cpu(), - ], dim=-1)) - - pv = phaser_values.to(torch.float64).cpu() - ang = phaser_angles.to(torch.float64).cpu() - R_grid = rotation_matrix_euler_zyz(torch.deg2rad(ang[:, :3])) - PR, AX = frame["PR"], frame["axisrot"] - # principal frame -> PDB frame (runMR_FRF.cc:542) - R_ph = torch.einsum("ij,njk,kl->nil", AX, R_grid, PR) - - osig = (ov - ov.mean()) / ov.std().clamp(min=1e-30) - psig = (pv - pv.mean()) / pv.std().clamp(min=1e-30) - out = { - "n_ours": int(ov.numel()), - "n_phaser": int(pv.numel()), - "grid_same_size": int(ov.numel() == pv.numel()), - "ours_max_sigma": float(osig.max()), - "phaser_max_sigma": float(psig.max()), - } - # Phaser's own statistics, as a check that our sigma means what theirs does. - if "stats" in frame: - st = frame["stats"] - out["phaser_max_sigma_logged"] = ( - (st["max"] - st["mean"]) / st["sigma"] if st["sigma"] else float("nan") - ) - - orbit = _orbit_of_identity(data) - out["n_orbit"] = int(orbit.shape[0]) - r_o, s_o, a_o = _truth_rank(ov, R_ours, orbit) - r_p, s_p, a_p = _truth_rank(pv, R_ph, orbit) - out.update({ - "truth_rank_ours": r_o, "truth_sigma_ours": s_o, "truth_angle_ours": a_o, - "truth_rank_phaser": r_p, "truth_sigma_phaser": s_p, "truth_angle_phaser": a_p, - "truth_rank_delta": r_o - r_p, - }) - - # --- Ghost anatomy ----------------------------------------------------- - # A ghost is a peak that outranks truth. The question is not "how many" but - # "what does the other engine see at exactly that rotation". Three outcomes, - # each implying a different fix: - # * Phaser has a peak there too, but weaker -> same physics, different - # weighting; find the term that suppresses it. - # * Phaser has nothing there -> we are manufacturing - # structure Phaser does not have. - # * Phaser has it just as strongly -> Phaser has the ghost too - # and wins somewhere downstream, not in the rotation function. - # `margin` is the discriminating power that matters: truth's sigma minus the - # strongest ghost's. Negative means the rotation function prefers a ghost. - tol = 1.5 * float(a.grid_sampling_deg) - - def _truth_mask(R_all): - keep = torch.zeros(R_all.shape[0], dtype=torch.bool) - for k in range(orbit.shape[0]): - tr = torch.einsum("nij,ij->n", R_all, orbit[k]) - ang = torch.rad2deg(torch.arccos(((tr - 1.0) / 2.0).clamp(-1.0, 1.0))) - keep |= ang <= tol - return keep - - m_o, m_p = _truth_mask(R_ours), _truth_mask(R_ph) - out["n_truthlike_ours"] = int(m_o.sum()) - out["n_truthlike_phaser"] = int(m_p.sum()) - - for tag, vals, sig, mask, Rs, other_v, other_R in ( - ("ours", ov, osig, m_o, R_ours, pv, R_ph), - ("phaser", pv, psig, m_p, R_ph, ov, R_ours), - ): - ghost_sig = sig.masked_fill(mask, float("-inf")) - gi = int(torch.argmax(ghost_sig)) - out[f"best_ghost_sigma_{tag}"] = float(sig[gi]) - out[f"margin_{tag}"] = float( - out[f"truth_sigma_{tag}"] - float(sig[gi]) - ) - # the same rotation, looked up in the other engine's map - j, dd = _nearest(other_R, Rs[gi]) - om, osd = other_v.mean(), other_v.std().clamp(min=1e-30) - out[f"best_ghost_{tag}_seen_by_other_sigma"] = float((other_v[j] - om) / osd) - out[f"best_ghost_{tag}_seen_by_other_rank"] = int((other_v > other_v[j]).sum()) - out[f"best_ghost_{tag}_lookup_angle"] = dd - - # Cross peak agreement: where do each engine's strongest peaks land in the - # other's map? - for label, (vs, Rs, vo, Ro) in { - "ph_in_ours": (pv, R_ph, ov, R_ours), - "ours_in_ph": (ov, R_ours, pv, R_grid if False else R_ph), - }.items(): - if label == "ours_in_ph": - # our rotations back into the principal frame for lookup - Rs_use = torch.einsum("ij,njk,kl->nil", AX.T, R_ours, PR.T) - vs_use, vo_use, Ro_use = ov, pv, R_grid - else: - Rs_use, vs_use, vo_use, Ro_use = R_ph, pv, ov, R_ours - top = torch.topk(vs_use, min(topn, vs_use.numel())).indices - ranks, sigs, angs = [], [], [] - vo_mean, vo_std = vo_use.mean(), vo_use.std().clamp(min=1e-30) - for i in top.tolist(): - j, dd = _nearest(Ro_use, Rs_use[i]) - ranks.append(int((vo_use > vo_use[j]).sum())) - sigs.append(float((vo_use[j] - vo_mean) / vo_std)) - angs.append(dd) - ranks_t = torch.tensor(ranks, dtype=torch.float64) - out[f"{label}_median_rank"] = float(ranks_t.median()) - out[f"{label}_median_sigma"] = float(torch.tensor(sigs).median()) - out[f"{label}_median_angle"] = float(torch.tensor(angs).median()) - out[f"{label}_frac_in_top{topn}"] = float((ranks_t < topn).to(torch.float64).mean()) - out[f"{label}_frac_above_5sig"] = float( - (torch.tensor(sigs) > 5.0).to(torch.float64).mean() - ) - return out - - -def run_case(pdb: str, outdir: Path, *, n_peaks: int = 500) -> dict: - """One structure: Phaser map, our map, comparison row.""" - work = outdir / "phaser" / pdb - t0 = time.time() - info = run_phaser_side(pdb, work) - frame = load_phaser_frame(work / f"{pdb}_phaser_map.csv") - model, data, ours, d_min_eff = run_our_side(pdb, info, n_peaks=n_peaks) - - row = {"experiment": EXPERIMENT, "pdb": pdb} - row.update(provenance()) - row.update({ - "phaser_lmax": info["lmax"], - "phaser_sampling_deg": info["sampling_deg"], - "phaser_sampling_deg_logged": info.get("sampling_deg_logged"), - "phaser_lmax_reso_A": info.get("lmax_reso_A"), - "phaser_all_data_to_limit": int(bool(info.get("all_data_to_limit"))), - "phaser_selected_d_min": info.get("selected_d_min"), - "phaser_selected_d_max": info.get("selected_d_max"), - "phaser_selected_n_refl": info.get("selected_n_refl"), - "phaser_n_samples_logged": info.get("n_samples"), - "phaser_seconds": round(info["seconds"], 1), - "pinned_d_min_eff_A": d_min_eff, - "ours_seconds": round(ours.seconds, 1), - }) - # Our radius formula vs Phaser's family, for the record. - row["our_mean_radius_A"] = float( - (model.xyz().to(torch.float64) - model.xyz().to(torch.float64).mean(0)) - .norm(dim=-1).mean().item() - ) - row["phaser_style_mean_radius_A"] = phaser_mean_radius(model) - pred = phaser_frf_params( - row["phaser_style_mean_radius_A"], - float(info.get("selected_d_min") or d_min_eff), - ) - row["predicted_lmax"] = pred.lmax - row["predicted_sampling_deg"] = pred.sampling_deg - # Phaser's own radius, inverted from its exact sampling step. - row["phaser_true_mean_radius_A"] = phaser_mean_radius_from_sampling( - info["sampling_deg"], d_min_eff, - ) - - row["high_order_axis"] = frame.get("high_order_axis") - row.update(compare(ours, info["angles"], info["values"], frame, data)) - row["total_seconds"] = round(time.time() - t0, 1) - - # Keep both arrays so a divergence-vs-beta plot needs no re-run. - npz = outdir / f"{pdb}_maps.pt" - torch.save( - { - "ours_values": ours.arf.values.cpu(), - "ours_alphas": ours.arf.alphas.cpu(), - "ours_betas": ours.arf.betas.cpu(), - "ours_gammas": ours.arf.gammas.cpu(), - "phaser_values": info["values"], - "phaser_angles": info["angles"], - }, - npz, - ) - return row - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", help="single structure") - ap.add_argument("--worklist-index", type=int, help="index into BENCH_PDBS") - ap.add_argument("--all", action="store_true", help="every benchmark structure") - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--outdir", default=None) - args = ap.parse_args() - - if args.pdb: - todo = [args.pdb] - elif args.worklist_index is not None: - todo = [BENCH_PDBS[args.worklist_index]] - elif args.all: - todo = list(BENCH_PDBS) - else: - ap.error("give --pdb, --worklist-index or --all") - - outdir = Path(args.outdir) if args.outdir else ( - Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT - ) - outdir.mkdir(parents=True, exist_ok=True) - csv_path = outdir / f"{EXPERIMENT}.csv" - - failures = 0 - for pdb in todo: - try: - row = run_case(pdb, outdir, n_peaks=args.n_peaks) - except Exception as exc: # keep the sweep going, but loudly - failures += 1 - print(f"[{pdb}] FAILED: {type(exc).__name__}: {exc}", flush=True) - continue - append_row(csv_path, row) - print( - f"[{pdb}] lmax={row['phaser_lmax']} samp={row['phaser_sampling_deg']:.4f}deg " - f"n={row['n_phaser']} | angle_dev={row['angle_max_dev_deg']:.2e} " - f"r={row.get('pearson_r', float('nan')):.4f} " - f"rho={row.get('spearman_r', float('nan')):.4f} " - f"resid={row.get('resid_rms_frac', float('nan')):.3f} " - f"top100={row.get('top100_overlap', float('nan')):.2f}", - flush=True, - ) - if failures: - print(f"\n{failures}/{len(todo)} cases FAILED", flush=True) - print(f"\nwrote {csv_path}", flush=True) - return 1 if failures == len(todo) else 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_normaliser_anatomy.py b/alignment_lab/diagnostics/frf_normaliser_anatomy.py deleted file mode 100644 index e1407996..00000000 --- a/alignment_lab/diagnostics/frf_normaliser_anatomy.py +++ /dev/null @@ -1,256 +0,0 @@ -"""Decompose the observation-normaliser gap that costs 3GR5 its rank. - -Injecting Phaser's per-reflection intensities into our otherwise-unchanged FRF -takes 3GR5 from rank 1995 / margin -1.06 to **rank 0 / margin +4.21** (job -489527), with every other stage already verified exact against Phaser. So the -whole remaining deficit is the normalised observation ``E^2``, and the question -is which part of it. - -Phaser's normaliser is (``DataB.cc:1106-1113``) - - sqrt_epsnSN[r] = sqrt( eps_n(h) * binAnisoFactor(bin, ANISO, SOLK, SOLB, K) ) - -a BEST-curve per-bin Sigma_N clamped to [0.5, 2] of BEST and corrected by a -fitted Wilson K/B, times an anisotropic tensor, plus a bulk-solvent term, times -the reflection multiplicity -- all refined. Ours is equal-count shell means of -``F^2`` (``preprocessing.py:30``) applied to amplitudes that have already had a -separately fitted overall anisotropy divided out (``align.py``: -``apply_overall_anisotropy``), then French-Wilson. - -So anisotropy is *not* simply missing on our side; it is removed upstream instead -of being folded into Sigma_N, which is equivalent only if the fitted tensor is -right. This splits ``log(E_phaser / E_ours)`` into pieces that can be fixed -independently: - -1. ``eps_n`` -- the multiplicity factor we omit entirely (its docstring in - ``build_lerf1_intensity`` claims it is "implicit in the symmetry reduction", - which is not the same thing as dividing by it). -2. the best possible **isotropic** model, a fine step function of ``|s|``. - Fitting a smooth Sigma_N curve cannot beat this, so it bounds what the curve - is worth. -3. a general quadratic form in Cartesian ``s`` -- residual **anisotropy** left - over after our own correction, reported as the eigenvalue spread of the - equivalent B tensor. This is the piece that matters most for a rotation - function: an angular error in the observed Patterson is exactly what a - rotation search is sensitive to, whereas a radial mis-scaling is not. - -Whatever variance survives all three is what only Phaser's refined per-bin -treatment could account for. - -Our side is *captured from the production path*, not reimplemented: the obs -``s`` and ``eEobs`` are spied out of the engine, so the anisotropy correction, -resolution window, shell edges and French-Wilson posterior are exactly the ones -production uses. - -Usage ------ - python -m diagnostics.frf_normaliser_anatomy --pdb 3GR5 \ - --dumps ../runs/encode_compare_489514/phaser/3GR5 -""" - -from __future__ import annotations - -import argparse -import sys -from pathlib import Path - -import numpy as np -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) - -from lab import FRFConfig, load_case, patched, run_frf # noqa: E402 -from lab.results import append_row, provenance # noqa: E402 -from diagnostics.frf_ghost_knockout import PHASER_PINNED # noqa: E402 -from diagnostics.frf_inject_phaser_obs import _pack # noqa: E402 - -EXPERIMENT = "frf_normaliser_anatomy" - - -def capture_ours(pdb: str): - """The obs ``s`` and ``eEobs`` the production engine actually expands. - - ``build_lerf1_intensity`` receives ``eEobs`` in the same order as the ``s`` - that reaches ``bessel_sh_expand`` -- both come from step 1's masked arrays - and nothing between them reorders or filters -- so spying on the two calls - gives an aligned pair. - """ - from torchref.experimental.alignment.frf import api as _api - - pin = PHASER_PINNED[pdb] - cap: dict = {} - orig_bessel = _api.bessel_sh_expand - orig_lerf1 = _api.build_lerf1_intensity - - def spy_bessel(s, vals, **kw): - if kw.get("zsymm", 1) > 1 and "s" not in cap: - cap["s"] = s.detach().cpu().to(torch.float64) - return orig_bessel(s, vals, **kw) - - def spy_lerf1(eEobs, centric, dfac=None, **kw): - if "eEobs" not in cap: - cap["eEobs"] = eEobs.detach().cpu().to(torch.float64) - cap["centric"] = centric.detach().cpu().clone() - cap["dfac"] = (torch.ones_like(cap["eEobs"]) if dfac is None - else dfac.detach().cpu().to(torch.float64)) - return orig_lerf1(eEobs, centric, dfac, **kw) - - model, data = load_case(pdb) - - def _pinned(model_radius_A, d_min_data, lmax_cap=48): - return int(pin["lmax"]) + 1, float(pin["d_min_eff"]) - - cfg = FRFConfig(n_peaks=5, lmax_cap=int(pin["lmax"]), - grid_sampling_deg=float(pin["sampling_deg"])) - with patched(_api, "phaser_lmax_resolution", _pinned), \ - patched(_api, "bessel_sh_expand", spy_bessel), \ - patched(_api, "build_lerf1_intensity", spy_lerf1): - run_frf(model, data, cfg, capture_arf=False, verbose=0) - for k in ("s", "eEobs"): - if k not in cap: - raise RuntimeError(f"failed to capture {k} from the engine") - if cap["s"].shape[0] != cap["eEobs"].shape[0]: - raise RuntimeError( - f"capture misaligned: s has {cap['s'].shape[0]} rows, eEobs " - f"{cap['eEobs'].shape[0]} -- something between step 1 and the " - f"expansion filters the obs set") - return cap, data - - -def phaser_esqr_lut(terms_csv: Path, sym_mats: torch.Tensor): - """Phaser's ``Esqr``, keyed by every P1 index the reflection feeds.""" - d = np.loadtxt(terms_csv, delimiter=",", skiprows=1) - if d.ndim == 1: - d = d[None, :] - hkl = torch.from_numpy(d[:, 0:3]).to(torch.float64) - esqr = torch.from_numpy(d[:, 8]).to(torch.float64) - orb = torch.einsum("kji,nj->nki", sym_mats.to(torch.float64), hkl) - orb = orb.round().to(torch.long).reshape(-1, 3) - vals = esqr.unsqueeze(1).expand(-1, sym_mats.shape[0]).reshape(-1) - keys = torch.cat([_pack(orb), _pack(-orb)]) - vals = torch.cat([vals, vals]) - uniq, inv = torch.unique(keys, return_inverse=True) - lut = torch.zeros(uniq.numel(), dtype=torch.float64) - lut[inv] = vals - return uniq, lut - - -def _quadratic_design(s: torch.Tensor) -> torch.Tensor: - """``[1, sx^2, sy^2, sz^2, 2 sx sy, 2 sx sz, 2 sy sz]``: constant + 6 aniso.""" - x, y, z = s[:, 0], s[:, 1], s[:, 2] - return torch.stack([torch.ones_like(x), x * x, y * y, z * z, - 2 * x * y, 2 * x * z, 2 * y * z], dim=1) - - -def _fit(A: torch.Tensor, b: torch.Tensor): - sol = torch.linalg.lstsq(A, b.unsqueeze(1)).solution.squeeze(1) - return sol, float((b - A @ sol).var(unbiased=False)) - - -def run(pdb: str, dumps: Path) -> dict: - - cap, data = capture_ours(pdb) - sg = data.spacegroup.matrices.to(torch.float64).cpu() - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - rec_inv = torch.linalg.inv(rec) - - s = cap["s"] - hkl = (s @ rec_inv).round().to(torch.long) - keys, lut = phaser_esqr_lut(dumps / "terms.csv", sg) - k = _pack(hkl) - pos = torch.searchsorted(keys, k).clamp(max=keys.numel() - 1) - found = keys[pos] == k - - esqr_p = lut[pos][found] - esqr_o = (cap["eEobs"] ** 2)[found] - s = s[found] - hkl = hkl[found] - good = (esqr_p > 1e-12) & (esqr_o > 1e-12) - esqr_p, esqr_o, s, hkl = esqr_p[good], esqr_o[good], s[good], hkl[good] - - y = torch.log(esqr_p / esqr_o) # log ratio of NORMALISED intensities - n = int(y.numel()) - var_tot = float(y.var(unbiased=False)) - - eps = sg.epsilon(hkl.to(torch.long), friedel=False) - # Phaser divides intensity by eps_n and we do not, so its Esqr should be - # SMALLER by that factor: log ratio carries -log(eps). - y_eps = y + torch.log(eps) - var_eps = float(y_eps.var(unbiased=False)) - - smag = s.norm(dim=-1) - n_fine = 40 - q = torch.linspace(0, 1, n_fine + 1, dtype=torch.float64)[1:-1] - fbin = torch.bucketize(smag, torch.quantile(smag, q)) - y_iso = y_eps.clone() - for b in range(n_fine): - m = fbin == b - if int(m.sum()) > 1: - y_iso[m] = y_eps[m] - y_eps[m].mean() - var_iso = float(y_iso.var(unbiased=False)) - - A = _quadratic_design(s) - _, var_aniso = _fit(A, y_iso) - coef, _ = _fit(A, y_eps) - C = torch.tensor([[coef[1], coef[4], coef[5]], - [coef[4], coef[2], coef[6]], - [coef[5], coef[6], coef[3]]], dtype=torch.float64) - # log(E^2) = ... + s^T C s; on intensities a B-factor is exp(-B s^2 / 2), - # so the equivalent B is -2 C. - ev = torch.linalg.eigvalsh(-2.0 * C) - - row = {"experiment": EXPERIMENT, "pdb": pdb} - row.update(provenance()) - row.update({ - "spacegroup": str(data.spacegroup.hm), - "n_obs_ours": int(cap["s"].shape[0]), - "n_matched": n, - "matched_frac": round(float(found.to(torch.float64).mean()), 4), - "rms_log_ratio_pct": round(100.0 * float(y.std(unbiased=False)), 2), - "mean_log_ratio": round(float(y.mean()), 4), - "var_total": var_tot, - "frac_eps": round(1.0 - var_eps / max(var_tot, 1e-300), 4), - "frac_iso": round((var_eps - var_iso) / max(var_tot, 1e-300), 4), - "frac_aniso": round((var_iso - var_aniso) / max(var_tot, 1e-300), 4), - "frac_unexplained": round(var_aniso / max(var_tot, 1e-300), 4), - "aniso_B_min": round(float(ev[0]), 2), - "aniso_B_max": round(float(ev[2]), 2), - "aniso_B_spread": round(float(ev[2] - ev[0]), 2), - "n_eps_gt1": int((eps > 1.0001).sum()), - "eps_max": round(float(eps.max()), 1), - }) - return row - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", required=True, choices=sorted(PHASER_PINNED)) - ap.add_argument("--dumps", required=True) - ap.add_argument("--outdir", default=None) - args = ap.parse_args() - - outdir = Path(args.outdir) if args.outdir else ( - Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT) - outdir.mkdir(parents=True, exist_ok=True) - csv_path = outdir / f"{EXPERIMENT}_{args.pdb}.csv" - - row = run(args.pdb, Path(args.dumps)) - append_row(csv_path, row) - print(f"\n{args.pdb} ({row['spacegroup']}): variance of " - f"log(Esqr_phaser / Esqr_ours), n={row['n_matched']} " - f"({row['matched_frac']:.4f} of our obs matched)", flush=True) - print(f" rms disagreement {row['rms_log_ratio_pct']:6.2f}% " - f"mean log ratio {row['mean_log_ratio']:+.4f}", flush=True) - print(f" explained by eps_n {row['frac_eps']*100:6.2f}% " - f"({row['n_eps_gt1']} refl with eps>1, max {row['eps_max']})", flush=True) - print(f" explained by iso(|s|) {row['frac_iso']*100:6.2f}%", flush=True) - print(f" explained by anisotropy {row['frac_aniso']*100:6.2f}% " - f"(equivalent B {row['aniso_B_min']} .. {row['aniso_B_max']} A^2, " - f"spread {row['aniso_B_spread']})", flush=True) - print(f" unexplained {row['frac_unexplained']*100:6.2f}%", flush=True) - print(f"\nwrote {csv_path}", flush=True) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_orbit_side.py b/alignment_lab/diagnostics/frf_orbit_side.py deleted file mode 100644 index 3124c231..00000000 --- a/alignment_lab/diagnostics/frf_orbit_side.py +++ /dev/null @@ -1,67 +0,0 @@ -"""On which side does the crystal symmetry act on a rotation-function peak? - -The FRF returns Euler matrices ``R`` mapping the search-model frame onto the -crystal frame. Two peaks are the same orientation when they differ by a -point-group rotation, but that rotation can compose on the left -(``R2 = R_g R1``) or on the right (``R2 = R1 R_g``), and the two are different -sets for a non-commuting group. Counting how many of the top peaks of a real -search collapse onto each other under each convention settles which one the -engine's peaks obey: mates of the true orientation appear many times in the -list, so the right convention finds many near-zero pairs and the wrong one few. -""" -from __future__ import annotations - -import argparse -import sys -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -from lab import BENCH_PDBS, cartesian_symops, load_case, random_rotation, seed_for # noqa: E402 - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="2DQ6", choices=list(BENCH_PDBS)) - ap.add_argument("--n", type=int, default=25) - ap.add_argument("--thr-deg", type=float, default=3.0) - args = ap.parse_args() - - from torchref.experimental.alignment import rotation_search - from torchref.experimental.alignment.sh import hkl_symops_to_cartesian - - model, data = load_case(args.pdb) - search = model.copy() - search = search.rotate(random_rotation(seed_for(args.pdb, 0)).to(model.dtype_float)) - sols = rotation_search(search, data, model_error_A=0.8, n_peaks=args.n) - R = sols.rotations.to(torch.float64) # (n, 3, 3) - n = R.shape[0] - - sym_lab = cartesian_symops(data.spacegroup, data.cell) # B S B^-1 - sym_sh = hkl_symops_to_cartesian( - data.spacegroup.matrices.to(torch.float64), - data.cell.reciprocal_basis_matrix.to(torch.float64)) - agree = float((sym_lab - sym_sh).abs().max()) - print(f"# {args.pdb} {data.spacegroup} n_peaks={n} |B S B^-1 - hkl_symops_to_cartesian|max={agree:.2e}") - - def pair_count(orbit_of): - cnt = 0 - for i in range(n): - for j in range(i + 1, n): - O = orbit_of(R[j]) # (g, 3, 3) - tr = torch.einsum("gab,ab->g", O, R[i]) - ang = ((tr - 1) / 2).clamp(-1, 1).arccos().min() * 180 / torch.pi - cnt += int(ang < args.thr_deg) - return cnt - - left = pair_count(lambda Rj: sym_lab @ Rj.unsqueeze(0)) - right = pair_count(lambda Rj: Rj.unsqueeze(0) @ sym_lab) - plain = pair_count(lambda Rj: Rj.unsqueeze(0)) - print(f"ROW pdb={args.pdb} pairs_within_{args.thr_deg:g}deg plain={plain} " - f"left(R_g R)={left} right(R R_g)={right}") - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_prep_compare.py b/alignment_lab/diagnostics/frf_prep_compare.py deleted file mode 100644 index b0cec122..00000000 --- a/alignment_lab/diagnostics/frf_prep_compare.py +++ /dev/null @@ -1,285 +0,0 @@ -"""Attribute the per-reflection gap between our FRF inputs and Phaser's. - -The stage-wise bisection put the divergence *before* the projection: reflection -positions agree to 1e-8, the SH-Bessel machinery reproduces Phaser's map from -Phaser's own coefficients at r = 0.998, and the radial band already matches -``nmax(l)`` -- but the intensities attached to those positions correlate at only -0.873 (2DQ6) against 0.988 (1AK5). Toggling ``use_epsilon``, French-Wilson, -shell-variance weights and the low-resolution cutoff moves that by <0.005, and -Phaser reports no tNCS, so ``V`` reduces to 1 on both sides. - -Two things remain unmeasured, and this runs both in one job. - -**Observation side.** Phaser builds -``intensity = cweight * (Esqr - V) / V^2 * DFAC^2`` with -``Esqr = (Feff / SIGMAN.sqrt_epsnSN)^2`` (DataMR.cc:930-945). The instrumented -binary now dumps every one of those terms per reflection, so a mismatch can be -attributed to a specific factor instead of only being visible in the product. -The two suspects are Phaser's smooth fitted ``Sigma_N`` against our equal-count -shell means, and ``DFAC`` (which we hard-wire to 1). - -**Calc side.** Never compared per reflection -- only post-projection via -``SearchElmn``. Phaser builds it in ``Ensemble::getELMNxR2``, a different -function from the observation one, so it needed its own dump. A difference here -would be invisible in every measurement made so far. - -Usage ------ - python -m diagnostics.frf_prep_compare --pdb 2DQ6 -""" - -from __future__ import annotations - -import argparse -import os -import subprocess -import sys -import time -from pathlib import Path - -import numpy as np -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) - -from lab import FRFConfig, case_paths, load_case, patched, run_frf # noqa: E402 -from lab.phaser_match import PATCHED_PHASER, write_keywords # noqa: E402 -from lab.results import append_row, provenance # noqa: E402 - -EXPERIMENT = "frf_prep_compare" - -#: Phaser's own bandwidth / resolution / sampling, from job 487737. -PINNED = { - "1DAW": dict(lmax=58, sampling_deg=6.233148, d_min_eff=5.70), - "1AK5": dict(lmax=70, sampling_deg=5.155428, d_min_eff=4.84), - "3K7M": dict(lmax=66, sampling_deg=5.464844, d_min_eff=5.39), - "4BX9": dict(lmax=96, sampling_deg=3.774036, d_min_eff=6.73), - "6G9X": dict(lmax=84, sampling_deg=4.314963, d_min_eff=5.92), - "2DQ6": dict(lmax=76, sampling_deg=4.772056, d_min_eff=6.36), - "3GR5": dict(lmax=66, sampling_deg=5.483967, d_min_eff=4.10), -} - - -def run_phaser_dumps(pdb: str, work: Path) -> dict: - """Run the instrumented binary, dumping observation and calc inputs.""" - work.mkdir(parents=True, exist_ok=True) - pdb_path, mtz_path = case_paths(pdb) - kw = write_keywords(work, mtz_path=mtz_path, model_pdb=pdb_path, - n_peaks=5, root=f"{pdb}_prep", title=f"prep {pdb}") - env = dict(os.environ) - env["PHASER_OBS_DUMP"] = str(work / "obs.csv") - env["PHASER_CALC_DUMP"] = str(work / "calc.csv") - env["PHASER_TERMS_DUMP"] = str(work / "terms.csv") - env["PHASER_SEARCH_ELMN_DUMP"] = str(work / "search_elmn.csv") - proc = subprocess.run([str(PATCHED_PHASER)], cwd=str(work), - input=kw.read_text(), capture_output=True, - text=True, timeout=5400, env=env) - (work / "run.log").write_text((proc.stdout or "") + (proc.stderr or "")) - for name in ("obs.csv", "calc.csv", "terms.csv"): - if not (work / name).exists(): - raise RuntimeError(f"{pdb}: {name} not written; see {work/'run.log'}") - return {"obs": work / "obs.csv", "calc": work / "calc.csv", - "terms": work / "terms.csv", - "search_elmn": work / "search_elmn.csv"} - - -def _polar_to_cart(r, th, ph) -> torch.Tensor: - return torch.tensor(np.stack( - [r * np.sin(th) * np.cos(ph), r * np.sin(th) * np.sin(ph), r * np.cos(th)], - axis=1)) - - -def capture_ours(pdb: str): - """Our per-reflection observation and calc inputs to the SH expansion. - - Both go through ``bessel_sh_expand``; the observation call is the one with - ``zsymm > 1`` (the calc side is deliberately never m-filtered). - """ - from torchref.experimental.alignment.frf import api as _api - - pin = PINNED[pdb] - cap: dict = {} - original = _api.bessel_sh_expand - - def spy(s, vals, **kw): - key = "obs" if kw.get("zsymm", 1) > 1 else "calc" - cap.setdefault(key, (s.detach().cpu().to(torch.float64), - vals.detach().cpu().to(torch.float64))) - return original(s, vals, **kw) - - model, data = load_case(pdb) - - def _pinned(model_radius_A, d_min_data, lmax_cap=48): - return int(pin["lmax"]) + 1, float(pin["d_min_eff"]) - - cfg = FRFConfig(n_peaks=20, lmax_cap=int(pin["lmax"]), - grid_sampling_deg=float(pin["sampling_deg"])) - with patched(_api, "phaser_lmax_resolution", _pinned), \ - patched(_api, "bessel_sh_expand", spy): - run_frf(model, data, cfg, capture_arf=False, verbose=0) - return cap - - -def match_and_correlate(P: torch.Tensor, PV: torch.Tensor, - S: torch.Tensor, V: torch.Tensor, - *, n_sample: int = 4000, tol: float = 1e-6) -> dict: - """Match Phaser's points onto ours by position, then compare the values. - - Correlation is the statistic, not a ratio: these intensities are centred on - zero (``E^2 - 1``), so element-wise ratios are dominated by division by - near-zero and say nothing. - """ - g = torch.Generator().manual_seed(1) - k = min(n_sample, P.shape[0]) - sel = torch.randperm(P.shape[0], generator=g)[:k] - q, qv = P[sel], PV[sel] - - idx = torch.empty(k, dtype=torch.long) - dist = torch.empty(k) - for i in range(0, k, 500): - d2 = ((q[i:i + 500, None, :] - S[None, :, :]) ** 2).sum(-1) - mn = d2.min(1) - idx[i:i + 500] = mn.indices - dist[i:i + 500] = mn.values.sqrt() - - ok = dist < tol - out = {"n_phaser": int(P.shape[0]), "n_ours": int(S.shape[0]), - "matched_frac": float(ok.to(torch.float64).mean()), - "median_pos_dist": float(dist.median())} - if int(ok.sum()) < 50: - out["corr"] = float("nan") - return out - a, b = qv[ok], V[idx][ok] - ac, bc = a - a.mean(), b - b.mean() - out["corr"] = float((ac @ bc) / (ac.norm() * bc.norm()).clamp(min=1e-30)) - out["phaser_mean"] = float(a.mean()) - out["phaser_sd"] = float(a.std()) - out["ours_mean"] = float(b.mean()) - out["ours_sd"] = float(b.std()) - return out - - -def attribute_obs_terms(terms_csv: Path, pdb: str) -> dict: - """Compare Phaser's normalisation terms against ours, keyed by Miller index. - - Phaser emits one row per *selected reflection* with its Miller index, so - this is immune to the two hazards that broke the first attempt: the - ``reso(r) > LMAX_RESO`` gate drops entries from the HKL list, and - ``HKL_clustered::add`` buckets by theta, so no parallel array indexed - against that list can stay aligned. - - ``Esqr = (Feff / sqrt_epsnSN)^2`` is Phaser's normalised intensity. Ours is - ``E^2`` from equal-count shell means. Comparing the *normalisers* isolates - the Wilson treatment from everything else. - """ - d = np.loadtxt(terms_csv, delimiter=",", skiprows=1) - if d.ndim == 1: - d = d[None, :] - hkl = d[:, 0:3].astype(int) - Feff, sqrtSN, DFAC, V, cw, Esqr, inten, reso = (d[:, i] for i in range(3, 11)) - - out = { - "n_terms": int(d.shape[0]), - "dfac_mean": float(DFAC.mean()), "dfac_sd": float(DFAC.std()), - "dfac_is_unity": int(bool(np.allclose(DFAC, 1.0, atol=1e-6))), - "V_is_unity": int(bool(np.allclose(V, 1.0, atol=1e-6))), - } - # Self-consistency: Phaser's own identity must reproduce its own intensity. - rebuilt = cw * (Esqr - V) / (V ** 2) * (DFAC ** 2) - scale = float(np.abs(inten).max()) or 1.0 - out["rebuild_max_err"] = float(np.abs(rebuilt - inten).max() / scale) - # And that Esqr really is (Feff/sqrt_epsnSN)^2. - out["esqr_max_err"] = float( - np.abs((Feff / np.maximum(sqrtSN, 1e-30)) ** 2 - Esqr).max() - / (float(np.abs(Esqr).max()) or 1.0)) - - # Our normalised E^2 for the same Miller indices. - from lab.reference_normalisers import wilson_normalise - model, data = load_case(pdb) - B = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - our_hkl = data.hkl.to(torch.long).cpu() - F = data.F.to(torch.float64).abs().cpu() - smag = (our_hkl.to(torch.float64) @ B).norm(dim=-1) - E_obs, _ = wilson_normalise(F, smag, 20) - - off = 1024 - key = lambda t: ((t[:, 0] + off) * 4096 + (t[:, 1] + off)) * 4096 + (t[:, 2] + off) - lut = {int(k): i for i, k in enumerate(key(our_hkl))} - idx = np.array([lut.get(int(k), -1) for k in key(torch.tensor(hkl))]) - ok = idx >= 0 - out["terms_matched_frac"] = float(ok.mean()) - if int(ok.sum()) > 50: - ours_E2 = (E_obs[torch.tensor(idx[ok])] ** 2).to(torch.float64) - ph_E2 = torch.tensor(Esqr[ok]) - for nm, a, b in (("esqr", ph_E2, ours_E2), - ("normaliser", torch.tensor(sqrtSN[ok]), - (torch.tensor(Feff[ok]) / ours_E2.clamp(min=1e-30).sqrt()))): - ac, bc = a - a.mean(), b - b.mean() - out[f"corr_{nm}"] = float( - (ac @ bc) / (ac.norm() * bc.norm()).clamp(min=1e-30)) - out["esqr_ratio_median"] = float((ours_E2 / ph_E2.clamp(min=1e-30)).median()) - return out - - -def run_case(pdb: str, outdir: Path) -> dict: - t0 = time.time() - work = outdir / "phaser" / pdb - dumps = run_phaser_dumps(pdb, work) - ours = capture_ours(pdb) - - row = {"experiment": EXPERIMENT, "pdb": pdb} - row.update(provenance()) - - # Calc side is NOT position-matched: our calc lives on a cubic P1 box - # (s = hkl/a) and Phaser's on its own ensemble grid, so the two sampling - # sets have no reason to coincide -- a nearest-position match returns 0%. - # The calc comparison belongs at the projected (SearchElmn) level. - for side, csv_name in (("obs", "obs"),): - d = np.loadtxt(dumps[csv_name], delimiter=",", skiprows=1) - P = _polar_to_cart(d[:, 1], d[:, 2], d[:, 3]) - PV = torch.tensor(d[:, 4]) - S, V = ours[side] - res = match_and_correlate(P, PV, S, V) - row.update({f"{side}_{k}": v for k, v in res.items()}) - - row.update({f"term_{k}": v for k, v in attribute_obs_terms(dumps["terms"], pdb).items()}) - row["seconds"] = round(time.time() - t0, 1) - return row - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", required=True, choices=sorted(PINNED)) - ap.add_argument("--outdir", default=None) - args = ap.parse_args() - - outdir = Path(args.outdir) if args.outdir else ( - Path(__file__).resolve().parents[1] / "runs" / EXPERIMENT) - outdir.mkdir(parents=True, exist_ok=True) - row = run_case(args.pdb, outdir) - append_row(outdir / f"{EXPERIMENT}.csv", row) - - print(f"\n=== {args.pdb} ===", flush=True) - print(" OBS n=%s/%s matched=%.3f corr=%.6f" - % (row.get("obs_n_ours"), row.get("obs_n_phaser"), - row.get("obs_matched_frac", float("nan")), - row.get("obs_corr", float("nan"))), flush=True) - print(" terms n=%s matched=%.3f | rebuild_err=%.2e esqr_err=%.2e" - % (row.get("term_n_terms"), row.get("term_terms_matched_frac", float("nan")), - row.get("term_rebuild_max_err", float("nan")), - row.get("term_esqr_max_err", float("nan"))), flush=True) - print(" Esqr corr(ours,phaser)=%.6f ratio median=%.4f" - % (row.get("term_corr_esqr", float("nan")), - row.get("term_esqr_ratio_median", float("nan"))), flush=True) - print(" DFAC unity=%s mean=%.5f sd=%.5f | V unity=%s | rebuild_err=%.2e" - % (row.get("term_dfac_is_unity"), row.get("term_dfac_mean", float("nan")), - row.get("term_dfac_sd", float("nan")), row.get("term_V_is_unity"), - row.get("term_rebuild_max_err", float("nan"))), flush=True) - print(" normaliser corr=%.6f" - % row.get("term_corr_normaliser", float("nan")), flush=True) - print(f"\nwrote {outdir / (EXPERIMENT + '.csv')}", flush=True) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/frf_rank.py b/alignment_lab/diagnostics/frf_rank.py deleted file mode 100644 index 6b6929ea..00000000 --- a/alignment_lab/diagnostics/frf_rank.py +++ /dev/null @@ -1,88 +0,0 @@ -"""Rank of the true orientation in the FRF peak list, per structure. - -The basic health check: rotate a deposited model by a seeded random rotation, -run the rotation function against that structure's own measured amplitudes, and -ask where the true orientation lands. Rank 0 means the top peak is correct. - -Peaks that outrank truth are the "ghosts" -- genuine correlations between the -model's self-Patterson and the crystal's intermolecular vectors, not noise. - -Usage:: - - python alignment_lab/diagnostics/frf_rank.py --pdb 1AK5 --trial 0 \ - --lmax-cap 64 --out-csv alignment_lab/runs/rank.csv -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, FRFConfig, ResultWriter, orbit_rank, # noqa: E402 - rotated_case, run_frf, seed_for) - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) - ap.add_argument("--trial", type=int, default=0) - ap.add_argument("--lmax-cap", type=int, default=64) - ap.add_argument("--d-min", type=float, default=4.0) - ap.add_argument("--d-max", type=float, default=15.0) - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--orbit-side", default="left", choices=["left", "right"]) - ap.add_argument("--orbit-frame", default="cart", choices=["cart", "frac"]) - ap.add_argument("--thr-deg", type=float, default=5.0) - ap.add_argument("--out-csv", default=None) - args = ap.parse_args() - - seed = seed_for(args.pdb, args.trial) - t0 = time.time() - rotated, data, R_true = rotated_case(args.pdb, seed) - load_s = time.time() - t0 - - sym = data.spacegroup.matrices.to(torch.float64).cpu() - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - cfg = FRFConfig(d_min=args.d_min, d_max=args.d_max, - n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) - res = run_frf(rotated, data, cfg) - rank, ang = orbit_rank(res.peaks, R_true, sym, side=args.orbit_side, - frame=args.orbit_frame, reciprocal_basis=rec, - thr_deg=args.thr_deg) - truth_sigma = float(res.peaks[rank].sigma) if rank >= 0 else float("nan") - - print(f"{args.pdb:6s} t{args.trial} seed={seed:<6d} {str(data.spacegroup):28s} " - f"n_ops={sym.shape[0]:2d} n_refl={data.hkl.shape[0]:7d} | " - f"rank={rank:4d} ang={ang:6.2f} truth_sig={truth_sigma:7.3f} " - f"map_max={res.map_max_sigma:7.3f} | frf={res.seconds:6.1f}s load={load_s:5.1f}s") - - if args.out_csv: - w = ResultWriter(args.out_csv, "frf_rank", - extra_fields=("truth_sigma", "map_max_sigma", "n_peaks", - "n_ghosts_above", "n_refl", - "frf_seconds", "load_seconds")) - w.write(pdb=args.pdb, seed=seed, trial=args.trial, - spacegroup=str(data.spacegroup), n_ops=int(sym.shape[0]), - truth_rank=rank, truth_angle_deg=round(ang, 4), - orbit_side=args.orbit_side, orbit_frame=args.orbit_frame, - lmax_cap=args.lmax_cap, d_min=args.d_min, d_max=args.d_max, - device="cpu", - truth_sigma=round(truth_sigma, 4), - map_max_sigma=round(res.map_max_sigma, 4), - n_peaks=len(res.peaks), - n_ghosts_above=(rank if rank >= 0 else len(res.peaks)), - n_refl=int(data.hkl.shape[0]), - frf_seconds=round(res.seconds, 2), - load_seconds=round(load_s, 2)) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/ghost_origin.py b/alignment_lab/diagnostics/ghost_origin.py deleted file mode 100644 index c1d945ec..00000000 --- a/alignment_lab/diagnostics/ghost_origin.py +++ /dev/null @@ -1,124 +0,0 @@ -"""Where do the truth-beating peaks come from? Vary only the observations. - -Runs the identical engine on the identical rotated search model, changing -nothing but the observed amplitudes at the same Miller indices: - -``real`` - the deposited measurements. -``crystal`` - ``|F_calc|`` of the deposited model in its real space group -- noiseless, - solvent-free, complete. Ghosts surviving here are not noise, solvent, - measurement error or missing data. -``molecule`` - ``|F_calc|`` of the same model in **P1** -- the self-Patterson only, with - the symmetry mates removed. Ghosts vanishing here are intermolecular. - -Substituting observations is safe because the FRF reads only ``F``, ``F_sigma``, -``hkl``, ``centric``, ``cell`` and ``spacegroup`` from the dataset. - -Usage:: - - python alignment_lab/diagnostics/ghost_origin.py --pdb 3K7M --trial 0 \ - --out-csv alignment_lab/runs/ghosts.csv -""" - -from __future__ import annotations - -import argparse -import sys -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, FRFConfig, ResultWriter, load_case, # noqa: E402 - orbit_rank, random_rotation, run_frf, seed_for) -from lab.truth import angle_to_orbit, symmetry_orbit # noqa: E402 - - -def substituted_data(data, model, mode: str): - """Return a dataset whose ``F`` is replaced according to ``mode``.""" - if mode == "real": - return data - m = model.copy() - if mode == "molecule": - # NOTE: assign the space-group NAME. SpaceGroup is an nn.Module, so - # assigning the object is intercepted by nn.Module.__setattr__ and the - # property setter never runs -- a silent no-op that would leave the - # crystal symmetry in place and quietly invalidate this whole arm. - m.spacegroup = "P 1" - else: - m.spacegroup = data.spacegroup.hm - m.reset_cache() - out = data.copy() if hasattr(data, "copy") else data - F = m(out.hkl).abs().detach().to(out.F.dtype) - out.F = F - return out - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="3K7M", choices=list(BENCH_PDBS)) - ap.add_argument("--trial", type=int, default=0) - ap.add_argument("--lmax-cap", type=int, default=64) - ap.add_argument("--d-min", type=float, default=4.0) - ap.add_argument("--d-max", type=float, default=15.0) - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--orbit-side", default="left", choices=["left", "right"]) - ap.add_argument("--orbit-frame", default="cart", choices=["cart", "frac"]) - ap.add_argument("--modes", default="real,crystal,molecule") - ap.add_argument("--out-csv", default=None) - args = ap.parse_args() - - seed = seed_for(args.pdb, args.trial) - model, data = load_case(args.pdb) - R_true = random_rotation(seed) - rotated = model.copy().rotate(R_true.to(model.dtype_float), - center=model.xyz().mean(0)) - sym = data.spacegroup.matrices.to(torch.float64).cpu() - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - orbit = symmetry_orbit(R_true, sym, side=args.orbit_side, - frame=args.orbit_frame, reciprocal_basis=rec) - cfg = FRFConfig(d_min=args.d_min, d_max=args.d_max, - n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) - - print(f"=== {args.pdb} trial {args.trial} seed {seed} | {data.spacegroup} " - f"n_ops={sym.shape[0]} | lmax_cap={args.lmax_cap} ===") - print(f" {'obs':10s} {'rank':>6s} {'ghosts':>7s} {'truth_sig':>10s} {'map_max':>8s}") - - writer = None - if args.out_csv: - writer = ResultWriter(args.out_csv, "ghost_origin", - extra_fields=("obs_mode", "n_ghosts_above", - "truth_sigma", "map_max_sigma", - "n_peaks")) - for mode in args.modes.split(","): - mode = mode.strip() - sub = substituted_data(data, model, mode) - res = run_frf(rotated, sub, cfg) - rank, ang = orbit_rank(res.peaks, R_true, sym, side=args.orbit_side, - frame=args.orbit_frame, reciprocal_basis=rec) - # Peaks outranking truth, i.e. the ghosts this arm produces. - n_ghosts = rank if rank >= 0 else len(res.peaks) - truth_sigma = float(res.peaks[rank].sigma) if rank >= 0 else float("nan") - print(f" {mode:10s} {rank:6d} {n_ghosts:7d} {truth_sigma:10.3f} " - f"{res.map_max_sigma:8.3f}") - if writer: - writer.write(pdb=args.pdb, seed=seed, trial=args.trial, - spacegroup=str(data.spacegroup), n_ops=int(sym.shape[0]), - truth_rank=rank, truth_angle_deg=round(ang, 4), - orbit_side=args.orbit_side, orbit_frame=args.orbit_frame, - lmax_cap=args.lmax_cap, d_min=args.d_min, d_max=args.d_max, - device="cpu", obs_mode=mode, n_ghosts_above=n_ghosts, - truth_sigma=round(truth_sigma, 4), - map_max_sigma=round(res.map_max_sigma, 4), - n_peaks=len(res.peaks)) - if args.out_csv: - print(f" wrote {args.out_csv}") - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/p1_grid_coherence.py b/alignment_lab/diagnostics/p1_grid_coherence.py deleted file mode 100644 index 8fa89d46..00000000 --- a/alignment_lab/diagnostics/p1_grid_coherence.py +++ /dev/null @@ -1,71 +0,0 @@ -"""How coarse can the P1 copy's FFT grid be for the translation set? - -The placement stage evaluates the P1 model's transform at the symmetry-rotated -indices of the translation set. Its grid was sized by the model's default -``max_res = 1.0 A`` whatever the set's resolution. This measures the complex -coherence of ``F_calc`` at the 15-4 A reflections between that grid and grids -sized to ``tf_d_min / oversampling``, and the time of each. -""" -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) - ap.add_argument("--tf-d-min", type=float, default=4.0) - ap.add_argument("--tf-d-max", type=float, default=15.0) - ap.add_argument("--oversampling", default="1.0,1.33,2.0") - args = ap.parse_args() - - model, data = load_case(args.pdb) - rec = data.cell.reciprocal_basis_matrix.to(torch.float64) - mask = data.get_valid_mask() - hkl = data.hkl[mask] - s = (hkl.to(torch.float64) @ rec).norm(dim=-1) - keep = (s >= 1.0 / args.tf_d_max) & (s <= 1.0 / args.tf_d_min) - hkl = hkl[keep] - sym_R = data.spacegroup.matrices.to(torch.float64) - hkl_SN = torch.einsum("ne,ied->ind", hkl.to(torch.float64), sym_R - ).reshape(-1, 3).round().to(torch.int64) - print(f"# {args.pdb} sg={data.spacegroup.hm} S={sym_R.shape[0]} N={hkl.shape[0]} " - f"window={args.tf_d_max}-{args.tf_d_min} A", flush=True) - - rot = model.copy() - rot = rot.rotate(random_rotation(seed_for(args.pdb, 0)).to(model.dtype_float)) - - ref = None - for over in [None] + [float(x) for x in args.oversampling.split(",")]: - m = rot.copy() - m.max_res = 1.0 if over is None else args.tf_d_min / over - m.spacegroup = "P 1" - with torch.no_grad(): - m.reset_cache(); m(hkl_SN) # warm - t0 = time.perf_counter() - for _ in range(3): - m.reset_cache(); sf = m(hkl_SN) - t = (time.perf_counter() - t0) / 3 - x = sf.to(torch.complex128) - if ref is None: - ref, coh, tag = x, 1.0, "1.0A" - else: - coh = float((ref.conj() * x).sum().abs() - / (ref.abs().norm() * x.abs().norm()).clamp(min=1e-30)) - tag = f"d_min/{over:g}" - print(f"ROW pdb={args.pdb} arm={tag} max_res={m.max_res:.3f} " - f"grid={tuple(int(v) for v in m.fft.gridsize)} t={t*1e3:.1f}ms " - f"coh={coh:.6f}", flush=True) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/phaser_headtohead.py b/alignment_lab/diagnostics/phaser_headtohead.py deleted file mode 100644 index ba2d067a..00000000 --- a/alignment_lab/diagnostics/phaser_headtohead.py +++ /dev/null @@ -1,112 +0,0 @@ -"""Head-to-head: our FRF vs Phaser on identical input. - -Both engines get the same rotated search model and the same reflections, and -both truth ranks are computed with the same orbit machinery, so the comparison -isolates the algorithms rather than the data handling. - -Read the caveats in :mod:`lab.phaser` before interpreting a result: Phaser -returns ~80-92k densely spaced samples, so a small "closest sample" angle is -expected regardless of whether it ranked truth well. - -Usage:: - - python alignment_lab/diagnostics/phaser_headtohead.py --pdb 1AK5 --trial 0 \ - --out-csv alignment_lab/runs/h2h.csv -""" - -from __future__ import annotations - -import argparse -import sys -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import (BENCH_PDBS, FRFConfig, ResultWriter, orbit_rank, # noqa: E402 - rotated_case, run_frf, seed_for) -from lab import phaser as ph # noqa: E402 - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="1AK5", choices=list(BENCH_PDBS)) - ap.add_argument("--trial", type=int, default=0) - ap.add_argument("--lmax-cap", type=int, default=64) - ap.add_argument("--d-min", type=float, default=4.0) - ap.add_argument("--d-max", type=float, default=15.0) - ap.add_argument("--n-peaks", type=int, default=500) - ap.add_argument("--orbit-side", default="left", choices=["left", "right"]) - ap.add_argument("--orbit-frame", default="cart", choices=["cart", "frac"]) - ap.add_argument("--workdir", default=None) - ap.add_argument("--timeout-s", type=int, default=5400) - ap.add_argument("--skip-phaser", action="store_true", - help="run only our engine (no phenix on this host)") - ap.add_argument("--out-csv", default=None) - args = ap.parse_args() - - seed = seed_for(args.pdb, args.trial) - work = Path(args.workdir or (Path(__file__).resolve().parents[1] / - "runs" / f"h2h_{args.pdb}_t{args.trial}")).resolve() - work.mkdir(parents=True, exist_ok=True) - - rotated, data, R_true = rotated_case(args.pdb, seed) - sym = data.spacegroup.matrices.to(torch.float64).cpu() - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() - orbit_kw = dict(side=args.orbit_side, frame=args.orbit_frame, - reciprocal_basis=rec) - - print(f"=== {args.pdb} trial {args.trial} seed {seed} | {data.spacegroup} " - f"n_ops={sym.shape[0]} ===") - - cfg = FRFConfig(d_min=args.d_min, d_max=args.d_max, - n_peaks=args.n_peaks, lmax_cap=args.lmax_cap) - res = run_frf(rotated, data, cfg) - our_rank, our_ang = orbit_rank(res.peaks, R_true, sym, **orbit_kw) - print(f" OURS rank={our_rank:5d} closest={our_ang:6.2f} deg " - f"({len(res.peaks)} peaks, {res.seconds:.1f}s, map max {res.map_max_sigma:.2f} sigma)") - - ph_rank, ph_ang, ph_n, ph_secs, ph_rc = -1, float("inf"), 0, 0.0, None - if not args.skip_phaser: - model_pdb = work / "rotated.pdb" - rotated.write_pdb(str(model_pdb)) - _, mtz_path = __import__("lab").case_paths(args.pdb) - kw = ph.write_frf_keywords(work, mtz_path=mtz_path, model_pdb=model_pdb) - ph_rc, ph_secs = ph.run_phaser(work, kw, timeout_s=args.timeout_s) - peaks = ph.parse_rlist(work / "phaser_frf.rlist") - ph_n = len(peaks) - if not peaks: - # rc==0 is not a success test; an empty list is the real signal. - print(f" PHASER produced no peaks (rc={ph_rc}); see {work}/phaser.stdout") - else: - ph_rank, ph_ang = ph.phaser_truth_rank(peaks, R_true, sym, **orbit_kw) - print(f" PHASER rank={ph_rank:5d} closest={ph_ang:6.2f} deg " - f"({ph_n} samples, {ph_secs:.0f}s)") - - if args.out_csv: - w = ResultWriter(args.out_csv, "phaser_headtohead", - extra_fields=("our_rank", "our_angle_deg", "our_seconds", - "our_n_peaks", "map_max_sigma", - "phaser_rank", "phaser_angle_deg", - "phaser_samples", "phaser_seconds", "phaser_rc")) - w.write(pdb=args.pdb, seed=seed, trial=args.trial, - spacegroup=str(data.spacegroup), n_ops=int(sym.shape[0]), - truth_rank=our_rank, truth_angle_deg=round(our_ang, 4), - orbit_side=args.orbit_side, orbit_frame=args.orbit_frame, - lmax_cap=args.lmax_cap, d_min=args.d_min, d_max=args.d_max, - device="cpu", - our_rank=our_rank, our_angle_deg=round(our_ang, 4), - our_seconds=round(res.seconds, 2), our_n_peaks=len(res.peaks), - map_max_sigma=round(res.map_max_sigma, 4), - phaser_rank=ph_rank, - phaser_angle_deg=(round(ph_ang, 4) if ph_n else ""), - phaser_samples=ph_n, phaser_seconds=round(ph_secs, 1), - phaser_rc=ph_rc if ph_rc is not None else "") - print(f" wrote {args.out_csv}") - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/pipeline_timing.py b/alignment_lab/diagnostics/pipeline_timing.py deleted file mode 100644 index fb07eb9b..00000000 --- a/alignment_lab/diagnostics/pipeline_timing.py +++ /dev/null @@ -1,66 +0,0 @@ -"""Warm, single-process wall clock of the placement pipeline, by stage. - -One process, each structure aligned twice, the first pass discarded: the first -call in a process pays kernel builds and cold file-system reads that are not -compute (161 s against 1.5 s has been measured for one stage). Prints the -pipeline's own stage table for the second pass and the pose error, so a timing -is never quoted for a run that did not place the model. -""" -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -from lab import (BENCH_PDBS, load_case, pose_error, random_rotation, # noqa: E402 - seed_for) - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdbs", default="1DAW,2DQ6,6G9X,3K7M,4BX9") - ap.add_argument("--threads", type=int, default=8) - ap.add_argument("--n-rotation-candidates", type=int, default=25) - ap.add_argument("--device", default="cpu", - help="where the model and data live; set TORCHREF_DEVICE to match") - args = ap.parse_args() - torch.set_num_threads(args.threads) - - from torchref.experimental.alignment import MolecularReplacementPipeline - - for pdb in [p.strip() for p in args.pdbs.split(",") if p.strip()]: - assert pdb in BENCH_PDBS - model, data = load_case(pdb, device=args.device) - canonical = model.xyz().clone() - R_true = random_rotation(seed_for(pdb, 0)) - for run in range(2): - search = model.copy() - search.spacegroup = "P 1" - search = search.copy().rotate(R_true.to(model.dtype_float), - center=canonical.mean(0)) - pipe = MolecularReplacementPipeline( - data, search, d_min=4.0, d_max=15.0, n_shells=20, - n_rotation_peaks=200, - n_rotation_candidates=args.n_rotation_candidates, - verbose=2 if run == 1 else 0, - ) - if args.device.startswith("cuda"): - torch.cuda.synchronize() - t0 = time.perf_counter() - sols = pipe.run(do_translation=True) - if args.device.startswith("cuda"): - torch.cuda.synchronize() - secs = time.perf_counter() - t0 - rot, trans = pose_error(sols[0].model.xyz(), canonical, data.cell, - data.spacegroup) - print(f"ROW pdb={pdb} run={run} device={args.device} n_cand={args.n_rotation_candidates} seconds={secs:.2f} " - f"rot_deg={rot:.2f} trans_A={trans:.2f}", flush=True) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/pose_recovery.py b/alignment_lab/diagnostics/pose_recovery.py deleted file mode 100644 index 7378437d..00000000 --- a/alignment_lab/diagnostics/pose_recovery.py +++ /dev/null @@ -1,239 +0,0 @@ -"""End-to-end pose recovery for the FRF -> FTF pipeline. - -Rank is not the deliverable, pose is. This places a randomly reoriented copy of -the deposited model and asks whether the pipeline gets it back, which is the -only measurement that settles a change to either stage. - -Success is a pose: final coordinates within ``--success-deg`` of canonical in -orientation AND within ``--success-A`` of it in position, modulo the crystal -symmetry (Cartesian point-group mates, lattice translations, allowed origin -shifts and polar directions). The gate used to be rotation-only, and it passed -placements 40-55 A from the true position on 2DQ6, 3VRJ, 4BX9 and 6G9X. - -The panel stands at **30/30** (10 structures x 3 trials) and **60/60** over six -structures x ten seeds, on every ranking arm. - -Arms (``--arms``) sweep how the winner is chosen among placed candidates: - -``llg`` - the default -- the translation likelihood at each candidate's best - translation. -``analytic_r`` - the analytical-scale R instead. -``corr`` - the fast translation function's own score. - -Over the six-structure sweep the three arms pick the same candidate in all 60 -cells; the arms exist so that can be re-checked whenever a structure separates -them. - -Usage:: - - python alignment_lab/diagnostics/pose_recovery.py --pdb 1DAW --trial 0 \ - --arms analytic_r,llg --out-csv alignment_lab/runs/pose.csv -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -# NOTE: no global torch.set_grad_enabled(False) here, unlike the FRF-only -# diagnostics. This runs the full pipeline, whose joint refine and -# rigid-body polish are LBFGS -- they need autograd, and disabling it -# raises "element 0 of tensors does not require grad". - -from lab import (BENCH_PDBS, ResultWriter, cartesian_symops, load_case, # noqa: E402 - pose_error, random_rotation, seed_for) - -ARMS = { - # How the winner is chosen among placed candidates. - "analytic_r": dict(rank_by="r"), - "corr": dict(rank_by="corr"), - "llg": dict(rank_by="llg"), -} - - -def residual_rotation_deg(aligned_xyz, canonical_xyz, symops_cart) -> float: - """Smallest angle between the aligned-to-canonical rotation and any symop. - - Kabsch superposition, then compared against every symmetry operator -- - a solution differing from canonical by a crystal symmetry is correct. - - ``symops_cart`` must be the **Cartesian** rotations, ``B S_k B^-1`` - (:func:`lab.cartesian_symops`). This used to take ``spacegroup.matrices`` - directly, which act on fractional coordinates: for trigonal and hexagonal - cells two of the six (four of the twelve) mates of a correct solution then - read as 30.00 and 21.09 degrees, and 2DQ6's "bimodal 6/10" was those mates. - """ - from torchref.experimental.alignment.frf.rotation_utils import ( - rotation_angular_distance_deg, - ) - - P = canonical_xyz.to(torch.float64) - Q = aligned_xyz.to(torch.float64) - Pc, Qc = P - P.mean(0), Q - Q.mean(0) - U, _, Vt = torch.linalg.svd(Qc.T @ Pc) - d = torch.sign(torch.det(U @ Vt)) - R = U @ torch.diag(torch.tensor([1.0, 1.0, d], dtype=torch.float64)) @ Vt - return min(float(rotation_angular_distance_deg(R, symops_cart[k])) - for k in range(symops_cart.shape[0])) - - -#: The rotation search's own bandwidth constant. Recorded in every row because -#: `align_model_to_data` has no bandwidth argument, so a `--lmax-cap` flag here -#: would name a value the engine never saw. -import importlib as _importlib # noqa: E402 - -_LMAX_CAP = _importlib.import_module( - "torchref.experimental.alignment.rotation_search").LMAX_CAP - - - -def _report_candidates(solutions, R_true, symops, success_deg) -> None: - """Annotate the pipeline's own candidates with how far each is from truth. - - The pipeline reports every candidate's scores at ``verbose >= 2`` but cannot - say which was right -- it has no ground truth, and a version of it that did - would be measuring itself. This joins the two: the ranked solutions it - returned, against the orientation the benchmark rotated the model by. - - That join is the whole point of driving the pipeline directly rather than - rebuilding its placement loop in the harness. A reimplementation drifts, and - then the two disagree about which candidate the pipeline picked -- which is - exactly what happened here: a harness reported truth top-ranked by analytic - R in 0 of 10 seeds while the pipeline solved 6 of them, because it fed the - R-factor a different set of translation peaks. - - ``SOLN`` lines are ordered as the pipeline ranked them -- by the likelihood, - descending -- so line 0 is what it returned. The fast score and ``R`` are - carried alongside so the three orderings can be compared. ``dtruth`` is the angle from that candidate's orientation to the - true one modulo crystal symmetry; ``pick`` marks the winner and ``true`` - marks every candidate that was in fact correct. - """ - from torchref.experimental.alignment.frf.rotation_utils import ( - rotation_angular_distance_deg, - ) - - R_t = R_true.to(torch.float64).cpu() - print(" SOLN rank k rot_score tf R dtruth flags") - for i, sol in enumerate(solutions): - R = torch.as_tensor(sol.rotation, dtype=torch.float64) - # `rotation` maps the search-model frame onto the crystal frame; the - # benchmark's R_true is the rotation applied to the coordinates, so the - # recovered orientation is compared as its transpose. - d = min(float(rotation_angular_distance_deg(R.T @ R_t, symops[k])) - for k in range(symops.shape[0])) - flags = ("pick " if i == 0 else " ") + ("true" if d <= success_deg else "") - print(f" SOLN {i:4d} {sol.candidate_index:3d} {sol.rotation_score:10.3f} " - f"{sol.translation_score:10.5f} {sol.r_factor:7.4f} " - f"{d:8.2f} {flags}") - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) - ap.add_argument("--trial", type=int, default=0) - ap.add_argument("--arms", default="llg,analytic_r") - ap.add_argument("--n-rotation-candidates", type=int, default=25) - ap.add_argument("--n-rotation-peaks", type=int, default=200) - ap.add_argument("--success-deg", type=float, default=8.0) - # A placement is a POSE: rotation and translation. The translation gate is - # generous -- downstream rigid-body refinement absorbs a few Angstrom -- but - # it separates a found position from one 40 A away, which the rotation-only - # gate this harness used to apply could not. Three large structures passed - # that gate on every seed while sitting 40-55 A from the true position. - ap.add_argument("--success-A", type=float, default=4.0) - ap.add_argument("--verbose", type=int, default=0) - ap.add_argument("--out-csv", default=None) - ap.add_argument("--tf-d-min", type=float, default=None) - ap.add_argument("--tf-d-max", type=float, default=None) - args = ap.parse_args() - - from torchref.experimental.alignment import MolecularReplacementPipeline - - seed = seed_for(args.pdb, args.trial) - model, data = load_case(args.pdb) - canonical_xyz = model.xyz().clone() - # Cartesian mates, not the fractional matrices -- see residual_rotation_deg. - symops = cartesian_symops(data.spacegroup, data.cell) - R_true = random_rotation(seed) - - print(f"=== {args.pdb} t{args.trial} seed={seed} {data.spacegroup} " - f"n_ops={symops.shape[0]} | success gate {args.success_deg} deg ===") - print(f" {'arm':16s} {'resid_deg':>10s} {'ok':>4s} {'seconds':>9s}") - - writer = None - if args.out_csv: - writer = ResultWriter(args.out_csv, "pose_recovery", - extra_fields=("arm", - "residual_deg", "success", - "n_rotation_candidates", - "pipeline_seconds")) - for arm in [a.strip() for a in args.arms.split(",") if a.strip()]: - if arm not in ARMS: - raise SystemExit(f"unknown arm {arm!r}; choose from {sorted(ARMS)}") - flags = ARMS[arm] - # Fresh copy per arm: rotate/translate mutate in place, and the arms - # must start from identical coordinates to be comparable. - search = model.copy() - search.spacegroup = "P 1" - search = search.copy().rotate(R_true.to(model.dtype_float), - center=canonical_xyz.mean(0)) - t0 = time.time() - try: - # The pipeline rather than `align_model_to_data`, which returns only - # the winner. Every candidate's score is the diagnosis when a - # placement goes wrong, and the pipeline already computed them. - pipe = MolecularReplacementPipeline( - data, search, d_min=4.0, d_max=15.0, n_shells=20, - n_rotation_peaks=args.n_rotation_peaks, - n_rotation_candidates=args.n_rotation_candidates, - verbose=args.verbose, tf_d_min=args.tf_d_min, - tf_d_max=args.tf_d_max, **flags, - ) - solutions = pipe.run(do_translation=True) - aligned = solutions[0].model - resid = residual_rotation_deg(aligned.xyz(), canonical_xyz, symops) - _, trans_A = pose_error(aligned.xyz(), canonical_xyz, data.cell, - data.spacegroup) - err = "" - if args.verbose >= 2: - _report_candidates(solutions, R_true, symops, args.success_deg) - except Exception as exc: # a crashed arm must not read as a success - resid, trans_A = float("nan"), float("nan") - err = f"{type(exc).__name__}: {exc}" - secs = time.time() - t0 - ok = ((resid == resid) and resid <= args.success_deg - and (trans_A == trans_A) and trans_A <= args.success_A) - print(f"ROW {arm} {args.pdb} trial={args.trial} " - f"n_cand={args.n_rotation_candidates} " - f"resid={resid:.3f} trans_A={trans_A:.2f} ok={int(bool(ok))} " - f"seconds={secs:.1f}", - flush=True) - print(f" {arm:16s} {resid:10.2f} {('yes' if ok else 'NO'):>4s} {secs:9.1f}" - + (f" {err}" if err else "")) - if writer: - writer.write(pdb=args.pdb, seed=seed, trial=args.trial, - spacegroup=str(data.spacegroup), n_ops=int(symops.shape[0]), - truth_rank="", truth_angle_deg=(round(resid, 4) - if resid == resid else ""), - orbit_side="kabsch", orbit_frame="cart", - lmax_cap=_LMAX_CAP, d_min=4.0, d_max=15.0, - device="cpu", arm=arm, - residual_deg=(round(resid, 4) if resid == resid else ""), - success=int(bool(ok)), - n_rotation_candidates=args.n_rotation_candidates, - pipeline_seconds=round(secs, 1)) - if args.out_csv: - print(f" wrote {args.out_csv}") - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/scaler_cost.py b/alignment_lab/diagnostics/scaler_cost.py deleted file mode 100644 index c2332c5c..00000000 --- a/alignment_lab/diagnostics/scaler_cost.py +++ /dev/null @@ -1,76 +0,0 @@ -"""Sixteen parameters, 8.6 seconds. Where does Scaler.refine_lbfgs spend it? - -`refine_lbfgs` builds its x-ray target with ``model=None`` and passes a detached -``fcalc`` per closure call, and the comment there states the fit never recomputes -structure factors. This counts them rather than trusting that, and times the -three phases separately -- construction, ``initialize``, and the fit -- because -the solvent contribution is also refined and a mask rebuilt per closure call -would look identical from outside. -""" -import sys, time -from pathlib import Path -import torch -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(True) -from lab import BENCH_PDBS, load_case # noqa: E402 - - -def main(): - from torchref.scaling import Scaler - from torchref.base.metrics.rfactor import rfactor_work_free - import torchref.model.model_ft as mft - import torchref.scaling.solvent as solv - - for pdb in sys.argv[1:] or ["1DAW"]: - model, data = load_case(pdb) - N = data.hkl.shape[0] - - counts = {"model_forward": 0, "solvent_forward": 0} - orig_fwd = mft.ModelFT.forward - def counted_fwd(self, *a, **kw): - counts["model_forward"] += 1 - return orig_fwd(self, *a, **kw) - mft.ModelFT.forward = counted_fwd - orig_sol = solv.SolventModel.forward - def counted_sol(self, *a, **kw): - counts["solvent_forward"] += 1 - return orig_sol(self, *a, **kw) - solv.SolventModel.forward = counted_sol - - t0 = time.perf_counter() - s = Scaler(model=model, data=data, nbins=20, verbose=0) - t_ctor = time.perf_counter() - t0 - - t0 = time.perf_counter() - with torch.no_grad(): - fc = model(data.hkl).detach() - t_fcalc = time.perf_counter() - t0 - n_after_fcalc = counts["model_forward"] - - t0 = time.perf_counter() - s.initialize(fc) - t_init = time.perf_counter() - t0 - n_after_init = counts["model_forward"] - sol_after_init = counts["solvent_forward"] - - t0 = time.perf_counter() - s.refine_lbfgs(fcalc=fc) - t_fit = time.perf_counter() - t0 - - t0 = time.perf_counter() - with torch.no_grad(): - rw, _ = rfactor_work_free(data, torch.abs(s.forward(fc))) - t_r = time.perf_counter() - t0 - - mft.ModelFT.forward = orig_fwd - solv.SolventModel.forward = orig_sol - print(f"ROW pdb={pdb} N={N} ctor={t_ctor:.2f}s fcalc={t_fcalc:.2f}s " - f"init={t_init:.2f}s fit={t_fit:.2f}s rfac={t_r:.2f}s " - f"total={t_ctor+t_fcalc+t_init+t_fit+t_r:.2f}s " - f"| model_fwd_during_fit={counts['model_forward']-n_after_init} " - f"(1 expected: the explicit fcalc) " - f"solvent_fwd_during_fit={counts['solvent_forward']-sol_after_init} " - f"R={float(rw):.4f}", flush=True) - - -main() diff --git a/alignment_lab/diagnostics/tf_resolution_and_grid.py b/alignment_lab/diagnostics/tf_resolution_and_grid.py deleted file mode 100644 index ea4ecb55..00000000 --- a/alignment_lab/diagnostics/tf_resolution_and_grid.py +++ /dev/null @@ -1,99 +0,0 @@ -"""What resolution does the translation stage actually run at, and what grid does it need? - -Two things to pin down, and the first invalidates my earlier numbers. - -``_prepare_translation_arrays`` (``pipeline.py:638``) is documented as -"Resolution/validity-masked" but applies **only** ``data.get_valid_mask()`` -- -there is no resolution cut. So the translation search runs at the data's full -resolution while the FRF runs at ``[d_max, d_min] = [15, 4]``. Any timing taken -on a 4 A subset is measuring a smaller problem than the pipeline solves. - -Second, the FFT grid is sized from ``ModelFT.max_res``, which defaults to 1.0 A -and which this stage never sets. Whether that is oversized depends entirely on -the answer to the first question: against 4 A data it is 64x too many voxels, -against 2 A data it is the ~2x oversampling one would ask for anyway. - -Reports the real N, the real ``d_min``, and times the structure-factor call on -grids sized at a range of ``max_res``, each checked for coherence against the -current 1.0 A grid -- because undersampling an FFT does not fail, it just -quietly returns different structure factors. -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 - - -def _time(fn, repeats=3): - fn() - t0 = time.perf_counter() - for _ in range(repeats): - out = fn() - return (time.perf_counter() - t0) / repeats, out - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) - ap.add_argument("--threads", type=int, default=4) - ap.add_argument("--oversampling", default="1.0,1.5,2.0,3.0") - args = ap.parse_args() - torch.set_num_threads(args.threads) - - model, data = load_case(args.pdb) - rec = data.cell.reciprocal_basis_matrix.to(torch.float64) - # Exactly what the pipeline masks with -- no resolution window. - mask = data.get_valid_mask() - hkl = data.hkl[mask] - s = (hkl.to(torch.float64) @ rec).norm(dim=-1) - d_min = float(1.0 / s.max()) - d_max = float(1.0 / s.min().clamp(min=1e-9)) - sym_R = data.spacegroup.matrices.to(torch.float64) - hkl_SN = torch.einsum("ne,ied->ind", hkl.to(torch.float64), sym_R - ).reshape(-1, 3).round().to(torch.int64) - S, N = int(sym_R.shape[0]), int(hkl.shape[0]) - - # For contrast: what the FRF's own window would leave. - in_frf = ((s >= 1.0 / 15.0) & (s <= 1.0 / 4.0)).sum().item() - print(f"# {args.pdb} sg={data.spacegroup.hm} S={S} " - f"N_pipeline={N} N_in_4to15A={in_frf} " - f"d_min={d_min:.2f} d_max={d_max:.1f} atoms={model.xyz().shape[0]} " - f"S*N={S*N} threads={args.threads}", flush=True) - - rot = model.copy() - rot.spacegroup = "P 1" - rot = rot.rotate(random_rotation(seed_for(args.pdb, 0)).to(model.dtype_float), - center=torch.zeros(3, dtype=model.xyz().dtype)) - - ref_sf = None - for over in [1.0] + [float(x) for x in args.oversampling.split(",")]: - m = rot.copy() - m.max_res = 1.0 if ref_sf is None else d_min / over - m.spacegroup = "P 1" - t, sf = _time(lambda: (m.reset_cache(), m(hkl_SN))[1]) - if ref_sf is None: - ref_sf, tag = sf.to(torch.complex128), "current(1.0A)" - coh = 1.0 - else: - x = sf.to(torch.complex128) - coh = float((ref_sf.conj() * x).sum().abs() - / (ref_sf.abs().norm() * x.abs().norm()).clamp(min=1e-30)) - tag = f"d_min/{over:g}" - print(f"ROW pdb={args.pdb} arm={tag} max_res={m.max_res:.3f} " - f"grid={tuple(int(v) for v in m.fft.gridsize)} " - f"t={t*1e3:.0f}ms coh={coh:.6f}", flush=True) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/tf_sf_backend.py b/alignment_lab/diagnostics/tf_sf_backend.py deleted file mode 100644 index aa2ae116..00000000 --- a/alignment_lab/diagnostics/tf_sf_backend.py +++ /dev/null @@ -1,112 +0,0 @@ -"""Why does one orientation's structure-factor evaluation cost ~1 s? - -``tf_cost.py`` put 93-97% of the translation stage in ``precompute_G``, which is -a single ``ModelFT.__call__`` at ``S*N`` Miller indices. That call goes through -``SfFFT`` (``model_ft.py:779``) -- splat the atoms onto a real-space grid, FFT -the box, sample the result. The grid is sized by the crystal cell and the -resolution, and it is built whether you wanted 12000 reflections or 12 million. - -The translation search wants a sparse, fixed set: ``S*N`` is 13k-160k here, -against 5-20k atoms. That is the regime direct summation is for, and -:class:`SfDS` -- same ``compute_structure_factors`` signature, no grid -- is -already in the tree but is not what ``ModelFT`` dispatches to. - -Times both on identical inputs and checks they agree, at the thread counts a -production run would see. Also separates first call from repeat: ``ModelFT`` is -rebuilt per orientation, so anything amortised across calls is paid in full by -the translation loop and has to be counted as setup, not as throughput. -""" - -from __future__ import annotations - -import argparse -import sys -import time -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) - -from lab import BENCH_PDBS, load_case, random_rotation, seed_for # noqa: E402 - - -def _time(fn, repeats=3): - fn() - t0 = time.perf_counter() - for _ in range(repeats): - out = fn() - return (time.perf_counter() - t0) / repeats, out - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="1DAW", choices=list(BENCH_PDBS)) - ap.add_argument("--d-min", type=float, default=4.0) - ap.add_argument("--d-max", type=float, default=15.0) - ap.add_argument("--threads", default="4,8") - args = ap.parse_args() - - from torchref.model.sf_ds import SfDS - - model, data = load_case(args.pdb) - rec = data.cell.reciprocal_basis_matrix.to(torch.float64) - s = (data.hkl.to(torch.float64) @ rec).norm(dim=-1) - keep = data.get_valid_mask() & (s >= 1.0 / args.d_max) & (s <= 1.0 / args.d_min) - hkl = data.hkl[keep] - sym_R = data.spacegroup.matrices.to(torch.float64) - h_R = torch.einsum("ne,ied->ind", hkl.to(torch.float64), sym_R) - hkl_SN = h_R.reshape(-1, 3).round().to(torch.int64) - S, N = int(sym_R.shape[0]), int(hkl.shape[0]) - - rot = model.copy() - rot.spacegroup = "P 1" - rot = rot.rotate(random_rotation(seed_for(args.pdb, 0)).to(model.dtype_float), - center=torch.zeros(3, dtype=model.xyz().dtype)) - n_at = int(rot.xyz().shape[0]) - print(f"# {args.pdb} sg={data.spacegroup.hm} S={S} N={N} S*N={S*N} " - f"atoms={n_at} model_max_res={rot.max_res}", flush=True) - - ds = SfDS(cell=rot.cell, spacegroup="P 1", dtype_float=rot.dtype_float, - device=rot.xyz().device) - iso, aniso = rot.get_iso(), rot.get_aniso() - - def fft_call(m): - m.reset_cache() # the loop gets a fresh model per candidate - return m(hkl_SN) - - # The grid is sized by the model's ``max_res``, which ModelFT defaults to - # 1.0 A. The translation search runs at d_min, so the default asks for - # (d_min/1.0)^3 times the voxels it needs. Both the FRF's dense calc - # (dense_calc.py:73) and the rigid-body stage (rigid_body.py:120) set this; - # the translation stage does not. - coarse = rot.copy() - coarse.max_res = float(args.d_min) - coarse.spacegroup = "P 1" - - for nt in [int(x) for x in args.threads.split(",")]: - torch.set_num_threads(nt) - t_fft, sf_fft = _time(lambda: fft_call(rot)) - t_coarse, sf_coarse = _time(lambda: fft_call(coarse)) - t_ds, (sf_ds, _) = _time(lambda: ds.compute_structure_factors( - hkl_SN, *iso, *aniso, apply_symmetry=True)) - - a = sf_fft.to(torch.complex128) - agree = lambda x: float((a.conj() * x.to(torch.complex128)).sum().abs() - / (a.abs().norm() - * x.abs().norm()).clamp(min=1e-30)) - print(f"ROW pdb={args.pdb} threads={nt} S={S} N={N} SN={S*N} " - f"atoms={n_at} grid_1A={tuple(int(v) for v in rot.fft.gridsize)} " - f"grid_dmin={tuple(int(v) for v in coarse.fft.gridsize)} " - f"t_fft_1A={t_fft*1e3:.0f}ms t_fft_dmin={t_coarse*1e3:.0f}ms " - f"t_ds={t_ds*1e3:.0f}ms " - f"gain_grid={t_fft/max(t_coarse,1e-9):.1f}x " - f"gain_ds={t_fft/max(t_ds,1e-9):.1f}x " - f"coh_dmin={agree(sf_coarse):.6f} coh_ds={agree(sf_ds):.6f}", - flush=True) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/truth_metric_check.py b/alignment_lab/diagnostics/truth_metric_check.py deleted file mode 100644 index 660c8aae..00000000 --- a/alignment_lab/diagnostics/truth_metric_check.py +++ /dev/null @@ -1,73 +0,0 @@ -"""Do the two truth metrics in this lab agree about which candidate is correct? - -Two answers to "is this candidate the right orientation" are in use: - -``angle_to_orbit`` used by the rank harnesses. Compares a candidate's - ``R_recovered`` against ``S_k @ R_true``. -``residual_rotation_deg`` used by pose_recovery for the pass/fail. Kabsch- - superposes the placed coordinates onto canonical and - takes the smallest angle to any symop. - -``RotationPeak`` rotations are ``R_recovered``, which maps the SEARCH-MODEL frame -onto the crystal frame -- the rotation applied to the coordinates is its -transpose. If the orbit comparison omits that transpose it is comparing a -rotation with its own inverse's orbit, which is a different set unless the -rotation is an involution. - -That matters beyond bookkeeping: the rank harness said ranking by the -translation correlation would beat the analytic R by 33/40 to 23/40, and end to -end it lost 31/40 to 36/40 -- as measured then; the correlation arm scores 32/40 -on the current code, so the direction is unchanged. A truth label that is wrong -makes every rank in that harness meaningless, so this checks it directly rather -than by inference. -""" -import sys -from pathlib import Path -import torch -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) -from lab import rotated_case, seed_for, symmetry_orbit # noqa: E402 -from lab.truth import angle_to_orbit # noqa: E402 - -pdb = sys.argv[1] if len(sys.argv) > 1 else "2DQ6" -trial = int(sys.argv[2]) if len(sys.argv) > 2 else 0 -n_cand = 25 - -from torchref.experimental.alignment.frf.rotation_utils import ( # noqa: E402 - rotation_angular_distance_deg, rotation_matrix_from_edmonds_euler) -from torchref.experimental.alignment.pipeline import ( # noqa: E402 - MolecularReplacementPipeline) -from torchref.experimental.alignment.rotation_search import ( # noqa: E402 - prepare_frf_inputs) - -seed = seed_for(pdb, trial) -model, data, R_true = rotated_case(pdb, seed) -pipe = MolecularReplacementPipeline(data, model, verbose=0, n_rotation_peaks=200, - n_rotation_candidates=n_cand) -frf = prepare_frf_inputs(model, data, d_min=pipe.d_min, d_max=pipe.d_max, - n_shells=pipe.n_shells, verbose=0) -pipe._frf = frf -peaks = pipe._rotation_candidates(frf)[:n_cand] - -symops = data.spacegroup.matrices.to(torch.float64).cpu() -rb = data.cell.reciprocal_basis_matrix.to(torch.float64).cpu() -orbit_l = symmetry_orbit(R_true, symops, side="left", frame="cart", - reciprocal_basis=rb) -orbit_r = symmetry_orbit(R_true, symops, side="right", frame="cart", - reciprocal_basis=rb) -R_t = R_true.to(torch.float64).cpu() - -print(f"# {pdb} trial={trial} n_cand={len(peaks)}") -print(f"{'k':>3s} {'orbit side=left':>16s} {'orbit side=right':>17s} " - f"{'coords (Kabsch form)':>21s}") -n_l = n_r = n_c = 0 -for k, p in enumerate(peaks): - R = rotation_matrix_from_edmonds_euler(p.alpha, p.beta, p.gamma).to(torch.float64) - a = angle_to_orbit(R, orbit_l) - b = angle_to_orbit(R, orbit_r) - c = min(float(rotation_angular_distance_deg(R.T @ R_t, symops[i])) - for i in range(symops.shape[0])) - n_l += a <= 8.0; n_r += b <= 8.0; n_c += c <= 8.0 - print(f"{k:3d} {a:16.2f} {b:17.2f} {c:21.2f}") -print(f"within 8 deg: side=left {n_l}/{len(peaks)}, side=right {n_r}/{len(peaks)}, " - f"coords {n_c}/{len(peaks)}") diff --git a/alignment_lab/diagnostics/truth_pose_scores.py b/alignment_lab/diagnostics/truth_pose_scores.py deleted file mode 100644 index e61833ab..00000000 --- a/alignment_lab/diagnostics/truth_pose_scores.py +++ /dev/null @@ -1,102 +0,0 @@ -"""Is a returned placement the deposited pose, and if not, does the deposited pose score better? - -Runs the pipeline on a seeded reorientation, then compares the winner with the -deposited model under the pipeline's own three selection scores, evaluated -through the same ``TranslationObs`` and ``precompute_G`` path. Reports the -rotation and translation error of the winner against the closest symmetry -image of the deposited model, and the raw fractional centroid offset to every -image, so a pseudo-translation shows up as a specific vector. - -If the deposited pose scores clearly better than the winner, the translation -search missed it. If they score the same, the data cannot tell them apart. -""" -from __future__ import annotations - -import argparse -import sys -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -from lab import (BENCH_PDBS, allowed_origin_shifts, load_case, pose_error, # noqa: E402 - random_rotation, seed_for) - - -def scores_at(pipe, model_placed): - """(tf score, R, llg) of an already-placed model through the pipeline's path.""" - from torchref.experimental.alignment.translation import ( - analytic_r_at, llg_at_translations, prepare_candidate, - translation_score_at) - data, obs = pipe.data, pipe._obs - m = model_placed.copy() - if pipe.tf_d_min > 0.0: - m.max_res = pipe.tf_d_min / 1.5 - m.spacegroup = "P 1" - cand = prepare_candidate(m, obs, data.spacegroup, data.cell) - t0 = torch.zeros(3, dtype=torch.float64) - tf = translation_score_at(obs, cand, t0) - r = analytic_r_at(obs, cand, t0) - llg = float(llg_at_translations(obs, cand, t0.view(1, 3))[0]) - return tf, r, llg - - -def main() -> int: - ap = argparse.ArgumentParser(description=__doc__) - ap.add_argument("--pdb", default="2DQ6", choices=list(BENCH_PDBS)) - ap.add_argument("--trial", type=int, default=3) - ap.add_argument("--rank-by", default="llg") - ap.add_argument("--tf-d-min", type=float, default=None) - ap.add_argument("--tf-d-max", type=float, default=None) - args = ap.parse_args() - - from torchref.experimental.alignment import MolecularReplacementPipeline - - model, data = load_case(args.pdb) - canonical = model.xyz().clone() - seed = seed_for(args.pdb, args.trial) - R_true = random_rotation(seed) - shifts, polar = allowed_origin_shifts(data.spacegroup) - print(f"=== {args.pdb} t{args.trial} {data.spacegroup} tf_window=({args.tf_d_max},{args.tf_d_min}) allowed shifts " - f"{[tuple(round(float(x), 3) for x in u) for u in shifts]} polar dims {polar.shape[1]}") - - search = model.copy() - search.spacegroup = "P 1" - search = search.copy().rotate(R_true.to(model.dtype_float), center=canonical.mean(0)) - pipe = MolecularReplacementPipeline( - data, search, d_min=4.0, d_max=15.0, n_shells=20, - n_rotation_peaks=200, n_rotation_candidates=25, rank_by=args.rank_by, - tf_d_min=args.tf_d_min, tf_d_max=args.tf_d_max, - ) - sols = pipe.run(do_translation=True) - win = sols[0] - - # Sanity: the deposited model against itself must be (0, 0). - print("SANITY deposited-vs-deposited rot/trans:", - pose_error(canonical, canonical, data.cell, data.spacegroup)) - rot, trans = pose_error(win.model.xyz(), canonical, data.cell, data.spacegroup) - print(f"WINNER rot_deg={rot:.3f} trans_A={trans:.2f} tf={win.translation_score:.5f} " - f"R={win.r_factor:.5f} llg={win.llg_score:.1f}") - - # Raw fractional centroid offset to every symmetry image of canonical. - B = data.cell.fractional_matrix.detach().cpu().to(torch.float64) - Binv = torch.linalg.inv(B) - S = data.spacegroup.matrices.detach().cpu().to(torch.float64) - T = data.spacegroup.translations.detach().cpu().to(torch.float64) - ca = Binv @ win.model.xyz().detach().cpu().to(torch.float64).mean(0) - cc = Binv @ canonical.detach().cpu().to(torch.float64).mean(0) - for k in range(S.shape[0]): - d = ca - (S[k] @ cc + T[k]) - d = d - d.round() - print(f" image {k}: centroid offset frac=({d[0]:+.3f},{d[1]:+.3f},{d[2]:+.3f}) " - f"|.|={float((B @ d).norm()):.1f} A") - - c_dep, r_dep, llg_dep = scores_at(pipe, model) - print(f"DEPOSITED tf={c_dep:.5f} R={r_dep:.5f} llg={llg_dep:.1f}") - c_w, r_w, llg_w = scores_at(pipe, win.model) - print(f"WINNER(re-scored) tf={c_w:.5f} R={r_w:.5f} llg={llg_w:.1f}") - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/alignment_lab/diagnostics/wilson_moments.py b/alignment_lab/diagnostics/wilson_moments.py deleted file mode 100644 index b68bf1fb..00000000 --- a/alignment_lab/diagnostics/wilson_moments.py +++ /dev/null @@ -1,59 +0,0 @@ -"""Second moment of E^2 per structure, through the shared Wilson fit. - -``<(E^2-1)^2>`` is the standard indicator for translational NCS and related -intensity modulations: 1.0 for ideal acentric Wilson data, larger when whole -classes of reflections reinforce or cancel together. 2DQ6 was recorded at 5.528 -against 1.0-1.2 for every other benchmark structure, and that number is the sole -evidence for calling it a tNCS case. - -It is worth recomputing, because tNCS needs at least two copies related by a -pure translation and 2DQ6 deposits ONE chain in the asymmetric unit, and because -the number was measured when the package had five disagreeing answers to what E -means. A second moment is a property of the normalisation as much as of the -data: normalise by a curve that is too flat and the resolution trend leaks -straight into the moment. -""" -import sys -from pathlib import Path -import torch -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -torch.set_grad_enabled(False) -from lab import BENCH_PDBS, load_case # noqa: E402 -from torchref.scaling import WilsonNormaliser # noqa: E402 - -print(f"{'pdb':6s} {'sg':12s} {'N':>7s} {'<(E2-1)^2>':>11s} {'':>7s} " - f"{'<|E|>':>7s} {'shell-norm':>11s}") -for pdb in BENCH_PDBS: - model, data = load_case(pdb) - mask = data.get_valid_mask() - F = data.F[mask].abs().to(torch.float64) - hkl = data.hkl[mask] - rec = data.cell.reciprocal_basis_matrix.to(torch.float64).to(hkl.device) - s = (hkl.to(torch.float64) @ rec).norm(dim=-1) - hkl_l = hkl.round().to(torch.int64) - eps = data.spacegroup.epsilon(hkl_l, friedel=False).to(torch.float64).clamp(min=1.0) - cen = data.spacegroup.is_centric(hkl_l).to(torch.bool) - acen = ~cen - - w = WilsonNormaliser(F * F, s, eps=eps, centric=cen, n_coeff=6) - E2 = w.E_squared.to(torch.float64)[acen] - m2 = float(((E2 - 1.0) ** 2).mean()) - - # The same moment under a 20-shell mean, which is what the older estimate - # would have used -- to separate "the data are odd" from "the curve was". - order = torch.argsort(s) - sh = torch.zeros_like(s, dtype=torch.long) - chunk = s.numel() // 20 - for k in range(20): - a = k * chunk - b = (k + 1) * chunk if k < 19 else s.numel() - sh[order[a:b]] = k - I = F * F / eps - tot = torch.zeros(20, dtype=torch.float64).scatter_add_(0, sh, I) - cnt = torch.bincount(sh, minlength=20).to(torch.float64).clamp(min=1) - E2s = I / (tot / cnt).clamp(min=1e-30).index_select(0, sh) - m2s = float(((E2s[acen] - 1.0) ** 2).mean()) - - print(f"{pdb:6s} {str(data.spacegroup.hm):12s} {int(acen.sum()):7d} " - f"{m2:11.3f} {float(E2.mean()):7.3f} " - f"{float(E2.clamp(min=0).sqrt().mean()):7.3f} {m2s:11.3f}", flush=True) diff --git a/alignment_lab/lab/__init__.py b/alignment_lab/lab/__init__.py deleted file mode 100644 index 312adbd8..00000000 --- a/alignment_lab/lab/__init__.py +++ /dev/null @@ -1,69 +0,0 @@ -"""Shared library for the alignment lab. - -Every diagnostic imports its primitives from here rather than re-deriving them. -The scripts this replaces carried ~37 copies of the rotation generator (in two -mutually incompatible variants), ~30 copies of the benchmark list, ~28 copies of -the CSV writer and ~20 copies of the rank-of-truth computation, several of which -disagreed with each other. One definition each, so two runs are comparable. -""" - -from .benchmark import ( - BENCH_PDBS, - PDB_STEMS, - REPO_ROOT, - case_paths, - load_case, - rotated_case, -) -from .truth import ( - allowed_origin_shifts, - cartesian_symops, - orbit_rank, - pose_error, - random_rotation, - seed_for, - symmetry_orbit, -) -from .aniso import ( - ARMS as ANISO_ARMS, - aniso_arm, - fit_aniso_log_space, - tensor_report, -) -from .frf import FRFConfig, FRFResult, merge_peak_lists, patched, run_frf -from .profile import (FRF_STAGES, PeakMemory, calibration_seconds, - host_info, stage_timers) -from .results import ResultWriter, append_row, provenance - -__all__ = [ - "BENCH_PDBS", - "PDB_STEMS", - "REPO_ROOT", - "case_paths", - "load_case", - "rotated_case", - "allowed_origin_shifts", - "cartesian_symops", - "orbit_rank", - "pose_error", - "random_rotation", - "seed_for", - "symmetry_orbit", - "ANISO_ARMS", - "aniso_arm", - "fit_aniso_log_space", - "tensor_report", - "FRFConfig", - "FRFResult", - "merge_peak_lists", - "patched", - "run_frf", - "FRF_STAGES", - "PeakMemory", - "calibration_seconds", - "host_info", - "stage_timers", - "ResultWriter", - "append_row", - "provenance", -] diff --git a/alignment_lab/lab/aniso.py b/alignment_lab/lab/aniso.py deleted file mode 100644 index 02d3cb43..00000000 --- a/alignment_lab/lab/aniso.py +++ /dev/null @@ -1,177 +0,0 @@ -"""Arms for the overall-anisotropy correction. - -``sh.fit_overall_anisotropy`` now fits in intensity space with a free constant. -The version it replaced regressed ``ln|F|^2 - ln<|F|^2>_shell`` on -``-2 pi^2 s.U.s`` by unweighted least squares **with no intercept**, and that was -the rotation function's last real defect. Three faults, all visible in its -output: - -* ``E[ln(I/)]`` is ``-gamma = -0.577`` for acentric reflections and - ``-gamma - ln 2 = -1.270`` for centric ones, not zero. With no intercept the - offset can only be absorbed by the quadratic form. The centric part is worse - than a constant: centric reflections lie on the zones perpendicular to the - symmetry axes, so the bias is direction-dependent. -* ``clamp(min=1e-30)`` turns a vanishing amplitude into ``y ~ -69``; a handful of - those outweigh thousands of ordinary reflections in an unweighted fit. -* ``ln`` of a single-reflection intensity has variance ``pi^2/6`` (acentric) or - ``pi^2/2`` (centric) with a heavy left tail, so the fit is dominated by the - weak reflections carrying the least information. - -Raw fitted B eigenvalue spreads came out at 70 to 5461 A^2 over the ten -benchmark structures. ``symmetrize_anisotropy`` then annihilated the garbage -where the point-group-invariant subspace is small (cubic -> one degree of -freedom) and left it standing where it is not (trigonal/hexagonal -> two). - -:func:`fit_aniso_log_space` reproduces that version, so the measurement that -justified replacing it can be re-run against the current tree rather than taken -on trust. -""" - -from __future__ import annotations - -import math -from contextlib import contextmanager - -import torch - -from torchref.experimental.alignment.sh import fit_overall_anisotropy - -#: U (A^2) -> B (A^2). -B_PER_U = 8.0 * math.pi ** 2 - -#: Arm names accepted by :func:`aniso_arm`. ``production`` is whatever -#: ``sh.fit_overall_anisotropy`` currently does; ``legacy_log`` is the biased -#: fit it replaced. -ARMS = ("production", "legacy_log", "no_aniso", "iso_only") - - -def fit_aniso_log_space( - F_obs: torch.Tensor, - s_vectors: torch.Tensor, - shell_idx: torch.Tensor, - P: int, - *, - min_count: int = 20, -) -> torch.Tensor: - """The superseded log-space fit, verbatim, for A/B against the current one. - - Unweighted least squares of ``ln|F|^2 - ln<|F|^2>_shell`` on - ``-2 pi^2 s.U.s`` with no constant term, and vanishing amplitudes clamped - rather than dropped. Returns ``U`` in A^2 in the same convention as - :func:`~torchref.experimental.alignment.sh.fit_overall_anisotropy`. - """ - dtype = F_obs.dtype - device = F_obs.device - valid = shell_idx >= 0 - F = F_obs[valid] - s = s_vectors[valid].to(dtype) - idx = shell_idx[valid] - - count = torch.zeros(P, dtype=torch.int64, device=device) - count.index_add_(0, idx, torch.ones_like(idx)) - F2 = F * F - sum_F2 = torch.zeros(P, dtype=dtype, device=device) - sum_F2.index_add_(0, idx, F2) - mean_F2 = sum_F2 / count.clamp(min=1).to(dtype) - - good = count >= min_count - if int(good.sum()) == 0: - return torch.zeros((3, 3), dtype=dtype, device=device) - keep = good[idx] - F2k = F2[keep].clamp(min=1e-30) - sk = s[keep].to(torch.float64) - mean_F2_k = mean_F2[idx[keep]].clamp(min=1e-30) - - y = (torch.log(F2k) - torch.log(mean_F2_k)).to(torch.float64) - X = torch.stack([ - sk[:, 0] ** 2, sk[:, 1] ** 2, sk[:, 2] ** 2, - 2.0 * sk[:, 0] * sk[:, 1], - 2.0 * sk[:, 0] * sk[:, 2], - 2.0 * sk[:, 1] * sk[:, 2], - ], dim=-1) - A = -2.0 * (torch.pi ** 2) * X - u, _, _, _ = torch.linalg.lstsq(A, y.unsqueeze(-1)) - Uxx, Uyy, Uzz, Uxy, Uxz, Uyz = u.squeeze(-1).tolist() - return torch.tensor( - [[Uxx, Uxy, Uxz], [Uxy, Uyy, Uyz], [Uxz, Uyz, Uzz]], - dtype=dtype, device=device, - ) - - -def tensor_report(U: torch.Tensor, tag: str) -> dict: - """B eigenvalues (A^2) of a U tensor, as result-row columns.""" - ev = torch.linalg.eigvalsh(U.to(torch.float64).cpu()) * B_PER_U - return {f"{tag}_B_min": round(float(ev[0]), 2), - f"{tag}_B_max": round(float(ev[2]), 2), - f"{tag}_B_spread": round(float(ev[2] - ev[0]), 2)} - - -@contextmanager -def aniso_arm(arm: str, data, *, d_min: float, d_max: float, captured: dict): - """Swap the anisotropy fit for the duration of one FRF call. - - ``captured`` receives the tensor actually fitted under the key ``raw``, so a - caller can report the artefact size alongside the rank it costs. - - Patches ``sh.fit_overall_anisotropy`` where ``rotation_search`` binds it -- - that is the symbol ``fit_anisotropy`` calls, and its result is what reaches - the engine as ``U_aniso``. - - Parameters - ---------- - arm : str - One of :data:`ARMS`. - data : ReflectionData - Used to recompute the centric mask over the same resolution window the - fit sees. A length mismatch raises rather than misaligning silently. - d_min, d_max : float - The window ``prepare_frf_inputs`` was called with. - captured : dict - Filled in by the wrapper. - """ - if arm not in ARMS: - raise ValueError(f"unknown aniso arm {arm!r}; expected one of {ARMS}") - import importlib - - _align = importlib.import_module( - "torchref.experimental.alignment.rotation_search") - - original = _align.fit_overall_anisotropy - rec = data.cell.reciprocal_basis_matrix.to(torch.float64) - smag = (data.hkl.to(torch.float64) @ rec).norm(dim=-1) - keep = (smag >= 1.0 / d_max) & (smag <= 1.0 / d_min) - centric_window = (data.centric[keep].to(torch.bool) - if hasattr(data, "centric") else None) - - def wrapped(F_obs, s_vec, shell_idx, centric, **kw): - U = original(F_obs, s_vec, shell_idx, centric, **kw) - captured.setdefault("raw", U.detach().clone()) - if arm == "production": - return U - if arm == "no_aniso": - return torch.zeros_like(U) - if arm == "iso_only": - # Radial part only; symmetrisation leaves lambda*I unchanged. - return torch.eye(3, dtype=U.dtype, device=U.device) * ( - torch.diagonal(U).sum() / 3.0) - if centric_window is None or centric_window.numel() != F_obs.shape[0]: - n = 0 if centric_window is None else centric_window.numel() - raise RuntimeError( - f"centric mask has {n} entries against {F_obs.shape[0]} " - f"amplitudes -- the resolution window assumed here " - f"([{d_min}, {d_max}] A) is not the engine's") - U_legacy = fit_aniso_log_space( - F_obs, s_vec, shell_idx, P=kw.get("P", 20), - min_count=kw.get("min_count", 20)) - captured["legacy"] = U_legacy.detach().clone() - return U_legacy - - setattr(_align, "fit_overall_anisotropy", wrapped) - try: - yield - finally: - setattr(_align, "fit_overall_anisotropy", original) - - -__all__ = ["ARMS", "B_PER_U", "aniso_arm", "fit_aniso_log_space", - "fit_overall_anisotropy", "tensor_report"] diff --git a/alignment_lab/lab/benchmark.py b/alignment_lab/lab/benchmark.py deleted file mode 100644 index c0628465..00000000 --- a/alignment_lab/lab/benchmark.py +++ /dev/null @@ -1,129 +0,0 @@ -"""Benchmark structures and case loading. - -The ten deposited structures the alignment work is measured on. Paths are -resolved relative to this file, never hardcoded: the drivers inherited from the -old worktree pointed at an absolute path inside a stale checkout, so they read -data and code from a different tree than the one under test. -""" - -from __future__ import annotations - -from pathlib import Path -from typing import TYPE_CHECKING, Tuple - -if TYPE_CHECKING: # pragma: no cover - typing only - import torch - - from torchref.io.datasets.reflection_data import ReflectionData - from torchref.model import ModelFT - -REPO_ROOT = Path(__file__).resolve().parents[2] -TEST_FILES = REPO_ROOT / "tests" / "files" - -#: Benchmark structures, in the order the seed formula depends on. -#: ``seed_for`` uses ``BENCH_PDBS.index(pdb)``, so **inserting or reordering -#: entries changes every seed** and silently invalidates comparisons against -#: archived results. Append only. -BENCH_PDBS: Tuple[str, ...] = ( - "1DAW", # C2, small -- the fast control; use it for anything quick - "3E98", # P2_1, control (note: pandas reads the string "3E98" as a float) - "3A5V", # I422 - "3VRJ", - "1AK5", # P432, cubic ghost case - "3K7M", # P432, the primary ghost case - "3GR5", # P6_522 - "2DQ6", # P3_121, tNCS - "4BX9", # P4_32_12, large; the only benchmark entry carrying ANISOU - "6G9X", # large -) - -#: PDB filename stems, where they differ from the code. 1AK5 is the only one. -PDB_STEMS = {"1AK5": "1AK5_with_H"} - -#: Present in tests/files but deliberately excluded from BENCH_PDBS: -#: 5BOV (a single translation-function allocation OOMs an A100-40GB) and -#: 7L84 (no matching MTZ). - - -def case_paths(pdb: str) -> Tuple[Path, Path]: - """Return ``(pdb_path, mtz_path)`` for a benchmark code. - - Parameters - ---------- - pdb : str - Benchmark structure code, e.g. ``"1DAW"``. - - Returns - ------- - tuple of pathlib.Path - Model and reflection file paths. - - Raises - ------ - FileNotFoundError - If either file is missing, named so the caller sees which one. - """ - stem = PDB_STEMS.get(pdb, pdb) - pdb_path = TEST_FILES / "pdb" / f"{stem}.pdb" - mtz_path = TEST_FILES / "mtz" / f"{pdb}.mtz" - for p in (pdb_path, mtz_path): - if not p.exists(): - raise FileNotFoundError(f"{pdb}: missing {p}") - return pdb_path, mtz_path - - -def load_case(pdb: str, device: str = "cpu") -> Tuple["ModelFT", "ReflectionData"]: - """Load the deposited model and its reflections. - - Parameters - ---------- - pdb : str - Benchmark structure code. - device : str, optional - Torch device for both objects. Default ``"cpu"``. - - Returns - ------- - tuple - ``(model, data)``. - """ - from torchref.io.datasets.reflection_data import ReflectionData - from torchref.model import ModelFT - - pdb_path, mtz_path = case_paths(pdb) - model = ModelFT(device=device).load_pdb(str(pdb_path)) - data = ReflectionData(device=device).load_mtz(str(mtz_path)) - return model, data - - -def rotated_case( - pdb: str, seed: int, device: str = "cpu", -) -> Tuple["ModelFT", "ReflectionData", "torch.Tensor"]: - """Load a case and rotate a copy of the model by a seeded random rotation. - - The returned model is a **copy**: ``Model.rotate`` mutates in place and - returns ``self``, so rotating the loaded model directly would also move the - reference a caller may want to compare against. - - Parameters - ---------- - pdb : str - Benchmark structure code. - seed : int - Seed for :func:`~alignment_lab.lab.truth.random_rotation`. - device : str, optional - Torch device. Default ``"cpu"``. - - Returns - ------- - tuple - ``(rotated_model, data, R_true)`` with ``R_true`` in float64. - """ - from .truth import random_rotation - - model, data = load_case(pdb, device=device) - R_true = random_rotation(seed) - rotated = model.copy().rotate( - R_true.to(model.dtype_float), center=model.xyz().mean(0), - ) - return rotated, data, R_true diff --git a/alignment_lab/lab/frf.py b/alignment_lab/lab/frf.py deleted file mode 100644 index f191922b..00000000 --- a/alignment_lab/lab/frf.py +++ /dev/null @@ -1,257 +0,0 @@ -"""Run the FRF and capture the full rotation function, not just the peak list. - -Every rank/ghost diagnostic needs the dense adaptive sample list as well as the -peaks, and the engine only returns the peaks. The capture below wraps the -engine's scoring method for the duration of one call; nine scripts each -carried their own copy of this monkeypatch. -""" - -from __future__ import annotations - -from contextlib import contextmanager -from dataclasses import asdict, dataclass, field -from typing import Any, Dict, Optional, Tuple - -import torch - - -@dataclass -class FRFConfig: - """Engine settings for one FRF evaluation. - - Collected into one object so a diagnostic passes a single config around and - the settings can be written into the result row verbatim. - """ - - d_min: float = 4.0 - d_max: float = 15.0 - n_shells: int = 20 - n_peaks: int = 500 - lmax_cap: int = 48 - dense_pad: float = 2.0 - grid_sampling_deg: float = 3.0 - #: Expected r.m.s. coordinate error, in Angstrom. ``None`` uses the Oeffner - #: estimate from the model's length, which is what the pipeline does. - model_error_A: Optional[float] = None - #: Weighting, the other half of the split. ``None`` leaves the production - #: default. These are separate arms on purpose: the design changes three - #: things at once -- the observed-side weight, the calculated-side weight - #: and whether the per-shell reweight runs -- and a panel that moves all - #: three cannot say which one did anything. - obs_weight: Optional[str] = None - sigma_a_source: Optional[str] = None - apply_bulk_solvent: Optional[bool] = None - shell_variance_weights: Optional[bool] = None - snr_cap: Optional[float] = None - trust_cap: Optional[float] = None - extra: Dict[str, Any] = field(default_factory=dict) - - def as_row(self) -> Dict[str, Any]: - """Config fields for a result row (``extra`` flattened out).""" - d = asdict(self) - d.pop("extra") - d.update(self.extra) - return d - - -@contextmanager -def patched(module: Any, name: str, replacement: Any): - """Temporarily replace ``module.name``, restoring it on exit. - - The "swap one engine internal and re-measure the rank" pattern -- used for - the dense-grid and box-construction experiments -- always needs the original - restored even when the body raises. - - Parameters - ---------- - module : module or object - Namespace holding the attribute. - name : str - Attribute name. - replacement : Any - Temporary value. - """ - original = getattr(module, name) - setattr(module, name, replacement) - try: - yield original - finally: - setattr(module, name, original) - - -@dataclass -class FRFResult: - """Outcome of one FRF evaluation. - - Attributes - ---------- - peaks : list - ``RotationPeak`` list, descending score. - arf : AdaptiveRotationFunction or None - The full adaptive sample list, when captured. - sigma : torch.Tensor or None - ``arf.values`` standardised to zero mean / unit sd -- the scale peak - heights are quoted in. - seconds : float - Wall time of the search call. - inputs : FRFInputs or None - The prepared observations (``F_obs``/``hkl``/``s_mag``/``centric``/``ll``). - The rescore consumes these, so keeping them lets a rescore run reuse one - FRF evaluation instead of recomputing it. - """ - - peaks: list - arf: Optional[Any] - sigma: Optional[torch.Tensor] - seconds: float - inputs: Optional[Any] = None - - @property - def map_max_sigma(self) -> float: - """Largest value of the standardised rotation function.""" - return float(self.sigma.max()) if self.sigma is not None else float("nan") - - -def merge_peak_lists(peak_lists, *, n_peaks: int, nms_radius_deg: float): - """Merge several peak lists into one, ranked by z-score. - - Used for the Patterson-radius union: the same obs expanded to two different - integration radii give two rotation functions whose absolute values are not - comparable, but whose per-run standardised heights (``RotationPeak.sigma``) - are. Peaks are pooled, sorted by sigma, and greedily suppressed by SO(3) - angular distance so the same orientation found by both radii appears once. - - Parameters - ---------- - peak_lists : sequence of list of RotationPeak - One list per run. - n_peaks : int - Cap on the merged list. - nms_radius_deg : float - Suppression radius, in degrees of SO(3) geodesic distance. - - Returns - ------- - list of RotationPeak - """ - from torchref.experimental.alignment.frf.rotation_utils import ( - rotation_angular_distance_deg, - rotation_matrix_from_edmonds_euler, - ) - - pooled = [p for pl in peak_lists for p in pl] - pooled.sort(key=lambda p: p.sigma, reverse=True) - kept, kept_R = [], [] - for p in pooled: - R = rotation_matrix_from_edmonds_euler(p.alpha, p.beta, p.gamma) - if any(rotation_angular_distance_deg(R, Rk) < nms_radius_deg - for Rk in kept_R): - continue - kept.append(p) - kept_R.append(R) - if len(kept) >= n_peaks: - break - return kept - - -def run_frf( - model, - data, - cfg: Optional[FRFConfig] = None, - *, - capture_arf: bool = True, - verbose: int = 0, -) -> FRFResult: - """Run the separated FRF on an already-rotated search model. - - Parameters - ---------- - model : ModelFT - Search model, already in the orientation to be scored. - data : ReflectionData - Observed reflections. - cfg : FRFConfig, optional - Engine settings. Defaults to :class:`FRFConfig`. - capture_arf : bool, optional - Also return the dense adaptive sample list. Default True. - verbose : int, optional - Engine verbosity. Default 0. - - Returns - ------- - FRFResult - """ - import time - - import importlib - - from torchref.experimental.alignment.frf import api as _api - - # `from ...alignment import rotation_search` gives the FUNCTION, which the - # package re-exports under the module's own name. Patching constants needs - # the module object. - _rs = importlib.import_module( - "torchref.experimental.alignment.rotation_search") - from torchref.experimental.alignment.frf.preprocessing import oeffner_vrms - - cfg = cfg or FRFConfig() - captured: Dict[str, Any] = {} - - def _wrapped(self, *args, **kwargs): - arf, peaks = _original(self, *args, **kwargs) - captured["arf"] = arf - return arf, peaks - - # The engine takes no tuning arguments any more: `lmax_cap`, `dense_pad` and - # the SO(3) sampling are module constants of `rotation_search`. The lab - # sweeps them by rebinding those constants for the duration of one call, so - # the production API stays switch-free while the measurements that chose the - # values remain reproducible. - frf_inputs = _rs.prepare_frf_inputs( - model, data, - d_min=cfg.d_min, d_max=cfg.d_max, n_shells=cfg.n_shells, verbose=verbose, - ) - model_error_A = cfg.model_error_A - if model_error_A is None: - model_error_A = oeffner_vrms(max(1, int(model.xyz().shape[0] / 8)), 1.0) - if cfg.extra: - raise ValueError( - f"FRFConfig.extra is no longer plumbed anywhere: {sorted(cfg.extra)}. " - f"The engine knobs it reached were deleted; patch the constants in " - f"torchref.experimental.alignment.rotation_search instead." - ) - - conv_kw = {} - # Engine knobs are omitted when unset so the production default applies, - # rather than being passed as None and overriding it with nothing. - for _name in ("obs_weight", "shell_variance_weights", "snr_cap", - "trust_cap", "sigma_a_source", "apply_bulk_solvent"): - _v = getattr(cfg, _name) - if _v is not None: - conv_kw[_name] = _v - - t0 = time.time() - with patched(_rs, "LMAX_CAP", int(cfg.lmax_cap)), \ - patched(_rs, "DENSE_CALC_PAD", float(cfg.dense_pad)), \ - patched(_rs, "GRID_SAMPLING_DEG", float(cfg.grid_sampling_deg)): - if capture_arf: - _original = _api.FastRotationFunction.score_model - with patched(_api.FastRotationFunction, "score_model", _wrapped): - peaks, _lmax, _dmin = _rs.search_peaks( - model, data, model_error_A, U_aniso=frf_inputs.U_aniso, - n_peaks=cfg.n_peaks, verbose=verbose, **conv_kw, - ) - else: - peaks, _lmax, _dmin = _rs.search_peaks( - model, data, model_error_A, U_aniso=frf_inputs.U_aniso, - n_peaks=cfg.n_peaks, verbose=verbose, **conv_kw, - ) - seconds = time.time() - t0 - - arf = captured.get("arf") - sigma = None - if arf is not None: - vals = arf.values.to(torch.float64) - sigma = (vals - vals.mean()) / vals.std().clamp(min=1e-30) - return FRFResult(peaks=peaks, arf=arf, sigma=sigma, seconds=seconds, - inputs=frf_inputs) diff --git a/alignment_lab/lab/phaser.py b/alignment_lab/lab/phaser.py deleted file mode 100644 index 9976675b..00000000 --- a/alignment_lab/lab/phaser.py +++ /dev/null @@ -1,204 +0,0 @@ -"""Phaser oracle adapter. - -Phaser is the reference the FRF is measured against, so this wraps invoking it -and reading its peaks back in our conventions. Previously these helpers lived -inside a pytest module and six scripts imported them from there. - -Three details are load-bearing and easy to lose: - -* **Convention.** ``R_ours = R_phaser.T``. Calibrated empirically in P1, where - ``n_ops == 1`` leaves no orbit ambiguity to hide a transpose error. -* **``PEAKS ROT SELECT ALL``** with clustering off. Phaser otherwise merges - symmetry equivalents before we can rank them. Expect ~80-92k samples with a - median nearest-neighbour spacing under 1 degree, so "the nearest sample is - within 1 degree" is not evidence of anything on its own. -* **Phaser exits 0 on fatal input errors.** The return code is not a success - test; an empty peak list is the real signal. Keyword files also need - **absolute** paths. -""" - -from __future__ import annotations - -import re -import subprocess -import time -from dataclasses import dataclass -from pathlib import Path -from typing import List, Optional, Sequence, Tuple - -import torch - -_SOLU_TRIAL_RE = re.compile( - r"SOLU\s+TRIAL\s+ENSEMBLE\s+\S+\s+" - r"EULER\s+([-+\d.]+)\s+([-+\d.]+)\s+([-+\d.]+)\s+" - r"RF\s+([-+\d.eE]+)\s+RFZ\s+([-+\d.eE]+)", - re.IGNORECASE, -) - - -@dataclass -class PhaserPeak: - """One Phaser FRF peak. Euler angles in **degrees**, Edmonds ZYZ.""" - - alpha_deg: float - beta_deg: float - gamma_deg: float - rf: float - rfz: float - - -def write_frf_keywords( - work: Path, *, mtz_path: Path, model_pdb: Path, - f_label: str = "FP", sigf_label: str = "SIGFP", root: str = "phaser_frf", -) -> Path: - """Write an MR_FRF keyword file. Paths are resolved to absolute. - - Parameters - ---------- - work : Path - Working directory; created if absent. - mtz_path, model_pdb : Path - Inputs. Relative paths are resolved -- Phaser fails on relative ones. - f_label, sigf_label : str, optional - MTZ column labels. - root : str, optional - Phaser output root. - - Returns - ------- - Path - The keyword file. - """ - work.mkdir(parents=True, exist_ok=True) - kw = work / f"{root}.kw" - kw.write_text( - f"TITLE FRF rotation ranking\n" - f"MODE MR_FRF\n" - f"HKLIN {Path(mtz_path).resolve()}\n" - f"LABIN F={f_label} SIGF={sigf_label}\n" - f"ENSEMBLE search PDB {Path(model_pdb).resolve()} IDENT 1.0\n" - f"COMPOSITION BY AVERAGE\n" - f"SEARCH ENSEMBLE search\n" - f"PEAKS ROT SELECT ALL\n" - f"PEAKS ROT CLUSTER OFF\n" - f"PEAKS ROT LEVEL 0\n" - f"ROOT {root}\n" - ) - return kw - - -def run_phaser(work: Path, kw_path: Path, timeout_s: int = 5400) -> Tuple[int, float]: - """Run ``phenix.phaser`` on a keyword file. - - Returns - ------- - tuple - ``(returncode, seconds)``; ``-1`` on timeout. **A zero return code does - not mean success** -- check that :func:`parse_rlist` found peaks. - """ - t0 = time.time() - try: - proc = subprocess.run( - ["phenix.phaser"], cwd=str(work), input=kw_path.read_text(), - capture_output=True, text=True, timeout=timeout_s, - ) - (work / "phaser.stdout").write_text(proc.stdout or "") - (work / "phaser.stderr").write_text(proc.stderr or "") - rc = proc.returncode - except subprocess.TimeoutExpired: - rc = -1 - return rc, time.time() - t0 - - -def parse_rlist(path: Path) -> List[PhaserPeak]: - """Parse ``SOLU TRIAL`` lines from a Phaser ``.rlist``. - - Returns an empty list when the file is absent, which is also what a failed - run looks like -- see the note on exit codes in the module docstring. - """ - path = Path(path) - if not path.exists(): - return [] - peaks: List[PhaserPeak] = [] - for line in path.read_text().splitlines(): - if "SOLU TRIAL" not in line.upper(): - continue - m = _SOLU_TRIAL_RE.search(line) - if m is None: - continue - a, b, g, rf, rfz = (float(x) for x in m.groups()) - peaks.append(PhaserPeak(a, b, g, rf, rfz)) - peaks.sort(key=lambda p: p.rfz, reverse=True) - return peaks - - -def euler_deg_to_matrices(peaks: Sequence[PhaserPeak]) -> torch.Tensor: - """Stack Phaser peaks as rotation matrices in **Phaser's** frame. - - Edmonds ZYZ active rotation ``R = Rz(alpha) Ry(beta) Rz(gamma)``. Apply - :func:`to_our_frame` before comparing against our orbit. - - Returns - ------- - torch.Tensor - ``(n, 3, 3)`` float64. Empty ``(0, 3, 3)`` for an empty input. - """ - if not peaks: - return torch.zeros((0, 3, 3), dtype=torch.float64) - a = torch.tensor([p.alpha_deg for p in peaks], dtype=torch.float64).deg2rad() - b = torch.tensor([p.beta_deg for p in peaks], dtype=torch.float64).deg2rad() - g = torch.tensor([p.gamma_deg for p in peaks], dtype=torch.float64).deg2rad() - ca, sa, cb, sb, cg, sg = a.cos(), a.sin(), b.cos(), b.sin(), g.cos(), g.sin() - return torch.stack([ - torch.stack([ca * cb * cg - sa * sg, -ca * cb * sg - sa * cg, ca * sb], dim=-1), - torch.stack([sa * cb * cg + ca * sg, -sa * cb * sg + ca * cg, sa * sb], dim=-1), - torch.stack([-sb * cg, sb * sg, cb], dim=-1), - ], dim=-2) - - -def to_our_frame(R_phaser: torch.Tensor) -> torch.Tensor: - """Convert Phaser-frame rotations to ours: ``R_ours = R_phaser.T``. - - Calibrated in P1 (``n_ops == 1``), where no symmetry orbit can mask a - transposition. Do not re-derive this per structure: with ~80-92k samples, - both conventions match *something* within a degree. - """ - return R_phaser.transpose(-1, -2) - - -def phaser_truth_rank( - peaks: Sequence[PhaserPeak], - R_true: torch.Tensor, - symops: torch.Tensor, - *, - reciprocal_basis: Optional[torch.Tensor] = None, - frame: str = "cart", - side: str = "left", - thr_deg: float = 5.0, -) -> Tuple[int, float]: - """Rank of the true orientation in Phaser's own peak list. - - Uses the same orbit machinery as our engine, after mapping Phaser's frame - onto ours, so the two ranks are directly comparable. - - Returns - ------- - tuple - ``(rank, best_angle_deg)``; rank ``-1`` if unmatched. - """ - from .truth import angle_to_orbit, symmetry_orbit - - if not peaks: - return -1, float("inf") - orbit = symmetry_orbit( - R_true, symops, side=side, frame=frame, reciprocal_basis=reciprocal_basis, - ) - R_ours = to_our_frame(euler_deg_to_matrices(peaks)) - rank, best = -1, float("inf") - for i in range(R_ours.shape[0]): - ang = angle_to_orbit(R_ours[i], orbit) - if ang < best: - best = ang - if ang <= thr_deg and rank < 0: - rank = i - return rank, best diff --git a/alignment_lab/lab/phaser_match.py b/alignment_lab/lab/phaser_match.py deleted file mode 100644 index 86b0ec08..00000000 --- a/alignment_lab/lab/phaser_match.py +++ /dev/null @@ -1,441 +0,0 @@ -"""Reproduce Phaser's FRF parameter chain exactly, and read back its map. - -Every number the rotation function depends on -- spherical-harmonic bandwidth, -the resolution actually expanded, and the SO(3) sampling step -- is *derived* by -Phaser from one quantity: ``mean_radius()``. This module implements that chain -verbatim from the 1.20 source so our engine can be pinned to the same values, -and parses the patched binary's log/dump so the derivation can be checked -against what Phaser actually did rather than trusted. - -Source anchors (PHENIX 1.20-4459, ``modules/phaser/codebase/phaser``): - -* ``lib/xyz_weight.cc:178`` ``mean_radius()`` -- the mean of the three - principal-axis **semi-extents of the bounding box**, NOT the mean atomic - distance from the centroid. These differ by ~25-30% on a protein. -* ``run/runMR_FRF.cc:406-410`` bandwidth:: - - sphereOuter = 2 * mean_radius - LMAX = ceil(2*pi*sphereOuter / HiRes) # round UP to even - LMAX = min(LMAX, DEF_CLMN_LMAX = 100) - -* ``run/runMR_FRF.cc:411-419`` resolution -- coarsened **only** when the cap - binds:: - - LMAX_RESO = (LMAX == 100) ? 2*pi*sphereOuter/LMAX : HiRes - -* ``run/runMR_FRF.cc:469-474`` sampling -- likewise keyed on the cap:: - - SAMP_RESO = (LMAX == 100) ? LMAX_RESO : HiRes - sampling = 2 * degrees(atan(SAMP_RESO / (4 * mean_radius))) - -The three are one coupled system: when the bandwidth saturates, the resolution -and the angular step coarsen together so the expansion is never asked to carry -detail it cannot represent. -""" - -from __future__ import annotations - -import math -import os -import re -import subprocess -import time -from dataclasses import asdict, dataclass -from pathlib import Path -from typing import Optional, Tuple - -import torch - -#: ``DEF_CLMN_LMAX`` from ``phaser_src/defaults:19``. -PHASER_LMAX_CAP = 100 - -#: ``DEF_CLMN_SPHE`` from ``phaser_src/defaults:17``. Zero means "use -#: ``2 * mean_radius``" rather than an explicit sphere radius. -PHASER_SPHERE_DEFAULT = 0.0 - -#: Built by the recipe in the ``phaser-instrumented-build`` memo. Honours -#: ``$PHASER_FRF_DUMP`` and is otherwise stock. -PATCHED_PHASER = ( - Path(__file__).resolve().parents[2] / "phaser_src" / "build" / "phaser_patched" -) - - -def phaser_mean_radius(model) -> float: - """``xyz_weight::mean_radius()`` -- mean principal-axis semi-extent. - - Rotates the coordinates onto the principal axes of their covariance, takes - the bounding-box extent along each axis, halves it, and averages the three. - - This is emphatically *not* ``mean(|xyz - centroid|)``; on 1DAW the two give - 26.2 A and 19.5 A. Since ``LMAX`` and the sampling step are both derived - from it, using the wrong one detunes the whole rotation function. - - Parameters - ---------- - model : ModelFT - Search model. - - Returns - ------- - float - Mean radius in Angstrom. - """ - xyz = model.xyz().to(torch.float64) - centred = xyz - xyz.mean(dim=0) - cov = (centred.T @ centred) / centred.shape[0] - _, axes = torch.linalg.eigh(cov) - projected = centred @ axes - extent = projected.max(dim=0).values - projected.min(dim=0).values - return float((extent / 2.0).mean().item()) - - -@dataclass -class PhaserFRFParams: - """The derived FRF parameters for one case. - - Attributes - ---------- - mean_radius_A : float - ``mean_radius()``. - hires_A : float - ``mr.HiRes()`` -- the high-resolution limit of the selected data. - lmax : int - Maximum ``l`` of the expansion (even, capped at 100). - lmax_reso_A : float - Resolution the expansion actually runs at. - samp_reso_A : float - Resolution feeding the sampling formula. - sampling_deg : float - SO(3) grid step in degrees. - capped : bool - Whether ``lmax`` hit ``PHASER_LMAX_CAP`` -- when True the resolution and - sampling are both coarsened, when False both stay at ``hires_A``. - """ - - mean_radius_A: float - hires_A: float - lmax: int - lmax_reso_A: float - samp_reso_A: float - sampling_deg: float - capped: bool - - def as_row(self) -> dict: - """Flatten for a CSV result row, prefixed ``phaser_``.""" - return {f"phaser_{k}": v for k, v in asdict(self).items()} - - -def phaser_frf_params( - mean_radius_A: float, - hires_A: float, - *, - lmax_cap: int = PHASER_LMAX_CAP, - use_rotate_lmax_reso: bool = True, -) -> PhaserFRFParams: - """Run Phaser's bandwidth/resolution/sampling chain. - - Parameters - ---------- - mean_radius_A : float - From :func:`phaser_mean_radius`. - hires_A : float - High-resolution limit of the data being expanded. - lmax_cap : int, optional - ``DEF_CLMN_LMAX``. Default 100. Lower it to emulate our historical - ``lmax_cap`` settings *with Phaser's coupling intact* -- note that - coarsening then engages, exactly as it does in Phaser at 100. - use_rotate_lmax_reso : bool, optional - ``input.USE_ROTATE_LMAX_RESO``. Default True (Phaser's default). - - Returns - ------- - PhaserFRFParams - """ - sphere_outer = 2.0 * float(mean_radius_A) - lmax = int(math.ceil(2.0 * math.pi * sphere_outer / float(hires_A))) - if lmax % 2 != 0: - lmax += 1 - lmax = min(lmax, int(lmax_cap)) - - capped = lmax == int(lmax_cap) and use_rotate_lmax_reso - if capped: - lmax_reso = 2.0 * math.pi * sphere_outer / lmax - samp_reso = lmax_reso - else: - lmax_reso = float(hires_A) - samp_reso = float(hires_A) - - sampling = 2.0 * math.degrees(math.atan(samp_reso / (4.0 * float(mean_radius_A)))) - return PhaserFRFParams( - mean_radius_A=float(mean_radius_A), - hires_A=float(hires_A), - lmax=lmax, - lmax_reso_A=lmax_reso, - samp_reso_A=samp_reso, - sampling_deg=sampling, - capped=capped, - ) - - -# --------------------------------------------------------------------------- -# Running the patched binary -# --------------------------------------------------------------------------- - -#: ``RFACTOR USE OFF`` is compulsory for any diagnostic that uses a model -#: already close to its answer: Phaser otherwise computes the R-factor of the -#: ensemble at the origin, decides the structure is solved, and emits a single -#: identity peak with ``RF*0`` -- skipping the rotation search entirely while -#: still reporting ``EXIT STATUS: SUCCESS``. -_KEYWORD_TEMPLATE = """TITLE {title} -MODE MR_FRF -HKLIN {mtz} -LABIN F={f_label} SIGF={sigf_label} -ENSEMBLE search PDB {pdb} IDENT 1.0 -COMPOSITION BY AVERAGE -SEARCH ENSEMBLE search -RFACTOR USE OFF -PEAKS ROT SELECT NUMBER -PEAKS ROT CUTOFF {n_peaks} -PEAKS ROT CLUSTER OFF -OUTPUT LEVEL VERBOSE -ROOT {root} -""" - - -def write_keywords( - work: Path, - *, - mtz_path: Path, - model_pdb: Path, - n_peaks: int = 20, - d_min: Optional[float] = None, - d_max: Optional[float] = None, - f_label: str = "FP", - sigf_label: str = "SIGFP", - root: str = "phaser_frf", - title: str = "FRF map dump", -) -> Path: - """Write an MR_FRF keyword file for the patched binary. - - ``d_min``/``d_max`` are omitted by default so Phaser uses the full data - range and its own coupling picks the expansion resolution -- that is the - configuration our engine should be matched against. - - Note the peak keyword takes two cards: ``PEAKS ROT SELECT NUMBER`` sets the - *mode* and ``PEAKS ROT CUTOFF n`` the count. ``SELECT NUMBER n`` is a syntax - error. Keeping the count small matters: with ``SELECT ALL`` Phaser rescores - every sample point (150k+ for a mid-size case), which dwarfs the search. - """ - work.mkdir(parents=True, exist_ok=True) - text = _KEYWORD_TEMPLATE.format( - title=title, - mtz=Path(mtz_path).resolve(), - pdb=Path(model_pdb).resolve(), - f_label=f_label, - sigf_label=sigf_label, - n_peaks=int(n_peaks), - root=root, - ) - if d_min is not None and d_max is not None: - text = text.replace( - "RFACTOR USE OFF", f"RESOLUTION {d_min} {d_max}\nRFACTOR USE OFF", - ) - kw = work / f"{root}.kw" - kw.write_text(text) - return kw - - -def run_patched_phaser( - work: Path, - kw_path: Path, - *, - dump_path: Optional[Path] = None, - binary: Path = PATCHED_PHASER, - timeout_s: int = 5400, -) -> Tuple[int, float, Path]: - """Run the instrumented binary, dumping the FRF sample list. - - Returns - ------- - tuple - ``(returncode, seconds, log_path)``. **The return code is not a success - test** -- Phaser exits 0 on fatal keyword errors. Check that the log - contains a ``TORCHREF:`` line and that the dump parses. - """ - if not Path(binary).exists(): - raise FileNotFoundError( - f"patched phaser binary not found at {binary}; build it with the " - "recipe in the phaser-instrumented-build memo" - ) - work.mkdir(parents=True, exist_ok=True) - log_path = work / (kw_path.stem + ".log") - env = dict(os.environ) - if dump_path is not None: - env["PHASER_FRF_DUMP"] = str(Path(dump_path).resolve()) - - t0 = time.time() - try: - proc = subprocess.run( - [str(binary)], cwd=str(work), input=kw_path.read_text(), - capture_output=True, text=True, timeout=timeout_s, env=env, - ) - log_path.write_text((proc.stdout or "") + (proc.stderr or "")) - rc = proc.returncode - except subprocess.TimeoutExpired: - log_path.write_text("TIMEOUT\n") - rc = -1 - return rc, time.time() - t0, log_path - - -_LOG_PATTERNS = { - "lmax": re.compile(r"maximum l value\s+(\d+)"), - "sampling_deg": re.compile(r"Sampling:\s+([0-9.]+)\s+degrees"), - "mean_radius_A": re.compile(r"^\s*[0-9.]+\s+([0-9.]+)\s+\d+\s+-?[0-9.]+", re.M), - "lmax_reso_A": re.compile(r"Elmn with resolution\s+([0-9.]+)"), - "n_samples": re.compile(r"TORCHREF: wrote (\d+) FRF sample points"), - "selected_hi": re.compile( - r"Resolution of Selected Data \(Number\):\s+([0-9.]+)\s+([0-9.]+)\s+\((\d+)\)" - ), -} - - -def parse_phaser_log(log_path: Path) -> dict: - """Pull the parameters Phaser actually used out of a VERBOSE log. - - These are the ground truth for the derivation in :func:`phaser_frf_params`; - the comparison harness asserts the two agree rather than assuming they do. - - Returns - ------- - dict - Keys present only when found: ``lmax``, ``sampling_deg``, - ``mean_radius_A``, ``lmax_reso_A``, ``n_samples``, ``selected_d_min``, - ``selected_d_max``, ``selected_n_refl``, ``all_data_to_limit``, - ``rotation_search_skipped``. - """ - text = Path(log_path).read_text() - out: dict = {} - for key in ("lmax", "n_samples"): - m = _LOG_PATTERNS[key].search(text) - if m: - out[key] = int(m.group(1)) - for key in ("sampling_deg", "lmax_reso_A", "mean_radius_A"): - m = _LOG_PATTERNS[key].search(text) - if m: - out[key] = float(m.group(1)) - m = _LOG_PATTERNS["selected_hi"].search(text) - if m: - out["selected_d_min"] = float(m.group(1)) - out["selected_d_max"] = float(m.group(2)) - out["selected_n_refl"] = int(m.group(3)) - out["all_data_to_limit"] = "Elmn with all data to resolution limit" in text - # The R-factor short-circuit. Present => the search never ran. - out["rotation_search_skipped"] = "SOLU SET RF*0" in text - return out - - -def load_phaser_map( - dump_path: Path, *, dtype: torch.dtype = torch.float64, -) -> Tuple[torch.Tensor, torch.Tensor]: - """Read a ``PHASER_FRF_DUMP`` CSV into our angle convention. - - Angles are returned **exactly as stored**, with no sign change. Phaser - writes ``euler = (-360*alpha_frac, beta_deg, -360*gamma_frac)`` - (``FastRot.cc:153``), and those stored values are already the Euler angles - of the grid rotation under ``R = Rz(alpha)Ry(beta)Rz(gamma)`` -- the same - Edmonds convention we use. Calibrated against a known grid/output pair on - 1DAW this reproduces Phaser's own reported peak to **0.074 deg**. - - Two traps here, both of which cost real time: - - * Negating alpha/gamma to "convert to our convention" is WRONG; they need - no conversion. - * The ``FastRot.cc`` comment at the end of ``get_FRF`` claims Phaser assumes - ``Rz(gamma)Ry(beta)Rz(alpha)``. That does not describe these stored - values; taking it literally puts the truth peak ~60-130 deg away. - - The stored angles are in the search model's **principal frame**. Use - ``ROT = axisrot @ R_grid @ PR`` (from the ``.frame`` sidecar) to reach the - PDB frame our engine works in. - - Returns - ------- - tuple - ``(angles_deg, values)`` -- ``(N, 3)`` of (alpha, beta, gamma) as stored, - and ``(N,)`` rotation-function values, in Phaser's sample order. - """ - import numpy as np - - raw = np.loadtxt(str(dump_path), delimiter=",", skiprows=1) - if raw.ndim == 1: - raw = raw[None, :] - angles = torch.from_numpy(raw[:, 1:4].copy()).to(dtype) - values = torch.from_numpy(raw[:, 4].copy()).to(dtype) - return angles, values - - -def load_phaser_frame(dump_path: Path, *, dtype: torch.dtype = torch.float64) -> dict: - """Read the ``.frame`` sidecar written next to a map dump. - - Returns ``PR``, ``axisrot`` (3x3 tensors), ``high_order_axis`` (int) and - ``stats`` (Phaser's own mean/sigma/max/min over the raw sample list, useful - for checking an externally computed sigma rather than trusting it). - """ - import numpy as np - - path = Path(str(dump_path) + ".frame") - if not path.exists(): - raise FileNotFoundError( - f"{path} missing; the binary must carry the runMR_FRF.cc patch that " - "dumps PR/axisrot, not only the FastRot.cc map patch" - ) - out: dict = {} - for line in path.read_text().strip().split("\n")[1:]: - parts = line.split(",") - vals = [float(x) for x in parts[1:10]] - if parts[0] in ("PR", "axisrot"): - out[parts[0]] = torch.tensor(np.array(vals).reshape(3, 3)).to(dtype) - elif parts[0] == "high_order_axis": - out["high_order_axis"] = int(vals[0]) - elif parts[0] == "stats_mean_sigma_max_min": - out["stats"] = dict(zip(("mean", "sigma", "max", "min"), vals[:4])) - return out - - -def phaser_sampling_from_dump(angles_deg: torch.Tensor) -> float: - """Recover the *exact* SO(3) sampling step from a dumped sample list. - - Phaser logs the sampling through ``dtos(SAMPLING,5,2)`` -- two decimals -- - which is far too coarse to rebuild the grid: on 1DAW the rounded 6.23 deg - gives 53270 sample points against the true 54430. The beta values in the - dump are full-precision and uniformly spaced, so their step is the exact - figure. - - Parameters - ---------- - angles_deg : torch.Tensor, shape (N, 3) - As returned by :func:`load_phaser_map`. - - Returns - ------- - float - Sampling step in degrees. - """ - betas = torch.unique(angles_deg[:, 1]) - if betas.numel() < 2: - raise ValueError("need at least two beta sections to infer the step") - steps = betas[1:] - betas[:-1] - return float(steps.median()) - - -def phaser_mean_radius_from_sampling(sampling_deg: float, samp_reso_A: float) -> float: - """Invert Phaser's sampling formula for ``mean_radius()``. - - ``sampling = 2*deg(atan(SAMP_RESO/(4*r)))`` inverts to - ``r = SAMP_RESO / (4*tan(sampling/2))``. Combined with - :func:`phaser_sampling_from_dump` this yields Phaser's own radius to full - precision, which is the reference our :func:`phaser_mean_radius` - reimplementation should be judged against (it currently runs ~4% high). - """ - half = math.radians(float(sampling_deg) / 2.0) - return float(samp_reso_A) / (4.0 * math.tan(half)) diff --git a/alignment_lab/lab/profile.py b/alignment_lab/lab/profile.py deleted file mode 100644 index 26e3930f..00000000 --- a/alignment_lab/lab/profile.py +++ /dev/null @@ -1,281 +0,0 @@ -"""Timing, memory and node-calibration primitives for the benchmark harness. - -Three measurements, three different hazards: - -**Time.** Wall clock on a shared cluster is a measurement of the cluster, not of -the code. Two engines timed on different nodes, or on the same node under -different neighbours, have been seen to differ by more than the effect being -looked for. So every timing row carries :func:`calibration_seconds` -- a fixed -workload run in the same process -- and the node's identity. Compare normalised -times, or compare only within a node. - -**Memory.** Peak RSS is a high-water mark: it never falls, so several -measurements in one process all report the largest. :class:`PeakMemory` samples -``/proc/self/statm`` on a thread so each window gets its own peak, and reports -the delta over the value at window entry. Two caveats, both reported rather than -hidden: a sampler can miss a spike shorter than its interval, and glibc does not -always return freed pages, so a later window in the same process can look -cheaper than it is. For absolute numbers, run one measurement per process and -read ``VmHWM``. - -**Stages.** A stage that cannot be resolved is *reported*, not skipped: an absent -row and a free stage look identical in a table. -""" - -from __future__ import annotations - -import importlib -import os -import threading -import time -from contextlib import contextmanager -from typing import Dict, List, Optional, Sequence, Tuple - -#: ``(module path, attribute)`` for each stage worth timing separately, coarse -#: to fine. -#: -#: The module named here is where the call is *resolved*, not where the function -#: is defined. ``frf.api`` binds ``bessel_sh_expand`` and friends into its own -#: namespace with ``from .data_mr import ...`` at import time, so replacing -#: ``data_mr.bessel_sh_expand`` leaves api's reference untouched and the stage -#: silently registers zero calls -- which reads as "that stage is free". Getting -#: this wrong left 85% of the runtime unattributed. -FRF_STAGES: Tuple[Tuple[str, str], ...] = ( - ("torchref.experimental.alignment.frf.dense_calc", "dense_calc_via_box"), - # Patched on its DEFINING module, not on `api`: it is reached through - # `FrenchWilsonE._compute`, which imports it inside the method body, so the - # lookup happens at call time and `api` never holds a reference at all. - ("torchref.experimental.alignment.frf.french_wilson", - "french_wilson_preprocess"), - ("torchref.experimental.alignment.frf.api", "bessel_sh_expand"), - ("torchref.experimental.alignment.frf.api", "cross_correlate_xi"), - ("torchref.experimental.alignment.frf.api", "evaluate_rotation_function"), - ("torchref.experimental.alignment.frf.api", "find_rotation_peaks"), - ("torchref.experimental.alignment.frf.sitelist_ang", - "wigner_contraction_per_beta"), - ("torchref.experimental.alignment.frf.sitelist_ang", - "build_dense_map_per_beta"), - ("torchref.experimental.alignment.frf.data_mr", "spherical_bessel_table"), - # Added once the named stages stopped accounting for the run: after the obs - # chain moved to the unique set, French-Wilson fell from 29.9% to 3.6% and - # the unattributed remainder became the second largest item at 29%. These are - # the rest of the per-reflection work, patched where each call RESOLVES -- - # `api` imports the preprocessing names at module top, `rotation_search` - # imports `apply_overall_anisotropy` from `sh` at module top, and - # `fit_relative_wilson_b` is imported inside `search_peaks` so it has to be - # patched on the defining module. - # No `wilson_normalise` row any more. The observed-side normalisation now - # arrives as the `e_convention` CLASS and is called through a parameter, so - # there is no module attribute to patch -- and a row that registers zero - # calls is worse than no row, because it reads as "that stage is free". - ("torchref.experimental.alignment.frf.api", "eterm_sigma_a"), - ("torchref.experimental.alignment.frf.api", "build_lerf1_intensity"), - ("torchref.experimental.alignment.frf.api", "apply_shell_variance_weights"), - ("torchref.experimental.alignment.frf.api", "detect_zsymm"), - ("torchref.experimental.alignment.frf.preprocessing", - "fit_relative_wilson_b"), - ("torchref.experimental.alignment.rotation_search", - "apply_overall_anisotropy"), -) - -#: Which stage each nested stage sits inside. A parent's time *includes* its -#: children, so summing the raw rows double-counts -- the harness subtracts to -#: report exclusive time as well, because otherwise the table invites optimising -#: the wrong thing. -FRF_STAGE_PARENTS = { - "build_dense_map_per_beta": "evaluate_rotation_function", - "wigner_contraction_per_beta": "build_dense_map_per_beta", - "spherical_bessel_table": "bessel_sh_expand", -} - - -def exclusive_times(totals): - """Per-stage time with nested children subtracted out. - - Parameters - ---------- - totals : mapping - Stage name -> inclusive seconds, as :func:`stage_timers` yields. - - Returns - ------- - dict - Stage name -> exclusive seconds. These sum without double counting. - """ - excl = dict(totals) - for child, parent in FRF_STAGE_PARENTS.items(): - if child in totals and parent in excl: - excl[parent] = excl[parent] - totals[child] - return excl - - -_PAGE_SIZE = os.sysconf("SC_PAGE_SIZE") if hasattr(os, "sysconf") else 4096 - - -def rss_bytes() -> int: - """Current resident set size, in bytes.""" - try: - with open("/proc/self/statm") as fh: - return int(fh.read().split()[1]) * _PAGE_SIZE - except (OSError, IndexError, ValueError): - return 0 - - -def vm_hwm_bytes() -> int: - """Process peak resident set size since start, in bytes. 0 if unavailable. - - Monotonic over the life of the process, so it answers "how much did this - process ever need", not "how much did this call need". - """ - try: - with open("/proc/self/status") as fh: - for line in fh: - if line.startswith("VmHWM:"): - return int(line.split()[1]) * 1024 - except OSError: - pass - return 0 - - -class PeakMemory: - """Sample RSS on a thread and report the peak inside a window. - - Parameters - ---------- - interval_s : float, optional - Sampling period. Default 0.02 s. Anything shorter than this that the - code allocates and frees again is invisible; ``missed_window_risk`` - records the interval so a reader can judge that. - """ - - def __init__(self, interval_s: float = 0.02): - self.interval_s = float(interval_s) - self._stop = threading.Event() - self._thread: Optional[threading.Thread] = None - self._peak = 0 - self._samples = 0 - - def _run(self) -> None: - while not self._stop.wait(self.interval_s): - r = rss_bytes() - self._samples += 1 - if r > self._peak: - self._peak = r - - @contextmanager - def window(self): - """Measure the peak RSS over the body, as a delta and an absolute.""" - baseline = rss_bytes() - self._peak = baseline - self._samples = 0 - self._stop.clear() - self._thread = threading.Thread(target=self._run, daemon=True) - self._thread.start() - out: Dict[str, float] = {} - try: - yield out - finally: - self._stop.set() - self._thread.join(timeout=5.0) - peak = max(self._peak, rss_bytes()) - out.update( - rss_baseline_mb=round(baseline / 1e6, 1), - rss_peak_mb=round(peak / 1e6, 1), - rss_delta_mb=round((peak - baseline) / 1e6, 1), - rss_samples=self._samples, - rss_sample_interval_s=self.interval_s, - vm_hwm_mb=round(vm_hwm_bytes() / 1e6, 1), - ) - - -@contextmanager -def stage_timers(stages: Sequence[Tuple[str, str]] = FRF_STAGES): - """Time each stage in ``stages`` for the duration of the body. - - Yields ``(totals, counts, unresolved)``. Every resolved stage is registered - at zero, so a stage that was instrumented but never called still shows up - with 0 calls -- distinguishable from one that is simply fast. - """ - from contextlib import ExitStack - - from .frf import patched - - totals: Dict[str, float] = {} - counts: Dict[str, int] = {} - unresolved: List[str] = [] - - def make(orig, key): - def timed(*a, **k): - t0 = time.perf_counter() - try: - return orig(*a, **k) - finally: - totals[key] += time.perf_counter() - t0 - counts[key] += 1 - return timed - - with ExitStack() as stack: - for mod_path, attr in stages: - try: - mod = importlib.import_module(mod_path) - original = getattr(mod, attr) - except (ImportError, AttributeError): - unresolved.append(f"{mod_path.rsplit('.', 1)[-1]}.{attr}") - continue - totals.setdefault(attr, 0.0) - counts.setdefault(attr, 0) - stack.enter_context(patched(mod, attr, make(original, attr))) - yield totals, counts, unresolved - - -def calibration_seconds(repeats: int = 3) -> float: - """Time a fixed workload, to normalise wall clock across nodes. - - Exercises the two kernels the rotation function spends its time in: a - complex einsum contraction and a batched FFT, both float64. Fixed shapes and - a fixed seed, so the only thing that varies is the machine and its - neighbours. Returns the best of ``repeats`` -- the least contended sample. - """ - import torch - - g = torch.Generator().manual_seed(0) - a = torch.randn(64, 96, 96, dtype=torch.float64, generator=g) - b = torch.randn(64, 96, 96, dtype=torch.float64, generator=g) - x = torch.complex(a, b) - best = float("inf") - for _ in range(max(1, repeats)): - t0 = time.perf_counter() - torch.einsum("nij,njk->nik", x, x) - torch.fft.ifft2(x) - best = min(best, time.perf_counter() - t0) - return best - - -def host_info() -> Dict[str, object]: - """Node identity and thread configuration, for every benchmark row.""" - import platform - - import torch - - model = "" - try: - with open("/proc/cpuinfo") as fh: - for line in fh: - if line.startswith("model name"): - model = line.split(":", 1)[1].strip() - break - except OSError: - pass - return { - "host": platform.node(), - "cpu_model": model, - "torch_threads": torch.get_num_threads(), - "slurm_job": os.environ.get("SLURM_JOB_ID", ""), - "slurm_cpus": os.environ.get("SLURM_CPUS_PER_TASK", ""), - "slurm_exclusive": os.environ.get("SLURM_JOB_NUM_NODES", ""), - } - - -__all__ = ["FRF_STAGES", "FRF_STAGE_PARENTS", "PeakMemory", - "calibration_seconds", "exclusive_times", "host_info", - "rss_bytes", "stage_timers", "vm_hwm_bytes"] diff --git a/alignment_lab/lab/results.py b/alignment_lab/lab/results.py deleted file mode 100644 index c4de8bd5..00000000 --- a/alignment_lab/lab/results.py +++ /dev/null @@ -1,161 +0,0 @@ -"""One result schema, one writer, one aggregator input format. - -Each old experiment invented its own CSV columns, so each needed its own -bespoke aggregator. Rows written through :class:`ResultWriter` all carry the -same core fields plus experiment-specific extras, so a single aggregator works -across experiments and a stale result is identifiable from the row itself. -""" - -from __future__ import annotations - -import csv -import os -import subprocess -from pathlib import Path -from typing import Any, Dict, Iterable, Mapping, Optional - -#: Fields every row carries. `orbit_side` / `orbit_frame` are here because a -#: truth rank cannot be interpreted without knowing which convention produced -#: it, and `torchref_version` / `git_sha` because results outlive the checkout. -CORE_FIELDS = ( - "pdb", - "seed", - "trial", - "spacegroup", - "n_ops", - "truth_rank", - "truth_angle_deg", - "orbit_side", - "orbit_frame", - "lmax_cap", - "d_min", - "d_max", - "device", - "torchref_version", - "git_sha", -) - - -def _git_sha(default: str = "unknown") -> str: - """Short SHA of the checkout this is running from, or ``default``.""" - try: - out = subprocess.run( - ["git", "rev-parse", "--short", "HEAD"], - cwd=str(Path(__file__).resolve().parents[2]), - capture_output=True, text=True, timeout=10, - ) - return out.stdout.strip() or default - except Exception: - return default - - -def provenance() -> Dict[str, str]: - """Version and checkout identity for a result row. - - Returns - ------- - dict - ``{'torchref_version': ..., 'git_sha': ...}``. - """ - try: - import torchref - - version = getattr(torchref, "__version__", "unknown") - except Exception: - version = "unknown" - return {"torchref_version": version, "git_sha": _git_sha()} - - -def append_row(csv_path: str | os.PathLike, row: Mapping[str, Any]) -> None: - """Append one row, writing the header only when the file is new. - - The header is taken from the file once it exists, not from each row. Taking - it per row silently loses data: a later row carrying a column the first row - lacked gets written wider than the header, and ``DictReader`` then drops the - surplus values into ``restkey``. That is how a sweep's anisotropy columns - went missing while every other column still read back correctly. A row with - an unknown column now raises; a row missing a known one gets a blank. - - Parameters - ---------- - csv_path : path-like - Destination CSV. Parent directories are created. - row : mapping - Column name -> value. - - Raises - ------ - ValueError - If ``row`` carries a column the file's header does not have. - """ - path = Path(csv_path) - path.parent.mkdir(parents=True, exist_ok=True) - header: Optional[list] = None - if path.exists() and path.stat().st_size > 0: - with open(path, newline="") as fh: - header = next(csv.reader(fh), None) - if header: - unknown = [k for k in row if k not in header] - if unknown: - raise ValueError( - f"{path.name}: row has columns absent from the header: " - f"{unknown}. Emit a stable set of columns for every row, using " - f"blanks where a value does not apply." - ) - with open(path, "a", newline="") as fh: - writer = csv.DictWriter(fh, fieldnames=header or list(row.keys()), - restval="") - if fh.tell() == 0: - writer.writeheader() - writer.writerow(dict(row)) - - -class ResultWriter: - """Writes rows sharing a fixed core schema plus per-experiment extras. - - Parameters - ---------- - csv_path : path-like - Destination CSV. - experiment : str - Experiment tag, recorded in every row. - extra_fields : iterable of str, optional - Experiment-specific column names, appended after the core fields. - - Notes - ----- - Column order is fixed at construction, so every row in a file has the same - header even if a caller omits a value (missing entries are written empty). - """ - - def __init__( - self, - csv_path: str | os.PathLike, - experiment: str, - extra_fields: Optional[Iterable[str]] = None, - ): - self.path = Path(csv_path) - self.experiment = experiment - self.extra_fields = tuple(extra_fields or ()) - self.fieldnames = ("experiment",) + CORE_FIELDS + self.extra_fields - self._provenance = provenance() - - def write(self, **values: Any) -> None: - """Write one row; unknown keys raise rather than being dropped silently. - - Raises - ------ - KeyError - If a value is passed whose column was not declared. - """ - unknown = set(values) - set(self.fieldnames) - if unknown: - raise KeyError( - f"{self.experiment}: undeclared column(s) {sorted(unknown)}; " - f"add them to extra_fields so every row keeps the same header" - ) - row = {k: "" for k in self.fieldnames} - row["experiment"] = self.experiment - row.update(self._provenance) - row.update(values) - append_row(self.path, row) diff --git a/alignment_lab/lab/truth.py b/alignment_lab/lab/truth.py deleted file mode 100644 index 9f83305a..00000000 --- a/alignment_lab/lab/truth.py +++ /dev/null @@ -1,289 +0,0 @@ -"""Seeded rotations and rank-of-truth against a symmetry orbit. - -Two things here previously existed in several disagreeing copies, and both -silently changed results rather than raising: - -* **The rotation generator.** Two QR-based variants were in circulation; the - one omitting the ``sign(diag(R))`` correction is not Haar-uniform and returns - a *different* rotation for the same seed. Runs from the two families are not - comparable. :func:`random_rotation` is the corrected form. -* **The orbit convention.** Rank-of-truth was computed with the symmetry - operators applied on either side, and in either the fractional or the - Cartesian frame. The choice changes the answer, so it is an explicit argument - here and is meant to be recorded in every result row. -""" - -from __future__ import annotations - -import math -from typing import Optional, Sequence, Tuple - -import torch - - -def random_rotation(seed: int, dtype: torch.dtype = torch.float64) -> torch.Tensor: - """Haar-uniform random rotation matrix from a seed. - - QR of a Gaussian matrix, with the ``sign(diag(R))`` correction that makes - the decomposition unique -- without it the distribution is not Haar-uniform - and the seed maps to a different rotation. - - Parameters - ---------- - seed : int - Generator seed. The mapping seed -> rotation is the reproducibility - contract for the whole lab; changing this function invalidates every - archived result. - dtype : torch.dtype, optional - Output dtype. Default ``torch.float64``. - - Returns - ------- - torch.Tensor - ``(3, 3)`` rotation with ``det = +1``. - """ - g = torch.Generator().manual_seed(int(seed)) - A = torch.randn(3, 3, generator=g, dtype=torch.float64) - Q, R = torch.linalg.qr(A) - Q = Q @ torch.diag(torch.sign(torch.diag(R))) - if torch.det(Q) < 0: - Q[:, 0] = -Q[:, 0] - return Q.to(dtype) - - -def seed_for(pdb: str, trial: int, base: int = 42) -> int: - """Seed for a ``(structure, trial)`` cell of the benchmark. - - ``base + 1000 * trial + index(pdb) * 7`` -- the convention the archived - results were produced under. The index term is why - :data:`~alignment_lab.lab.benchmark.BENCH_PDBS` is append-only. - - Parameters - ---------- - pdb : str - Benchmark structure code. - trial : int - Trial number. - base : int, optional - Seed base. Default 42. - - Returns - ------- - int - The seed. - """ - from .benchmark import BENCH_PDBS - - return int(base) + 1000 * int(trial) + BENCH_PDBS.index(pdb) * 7 - - -def symmetry_orbit( - R_true: torch.Tensor, - symops: torch.Tensor, - *, - side: str = "right", - frame: str = "cart", - reciprocal_basis: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """Build the set of rotations equivalent to ``R_true`` under the point group. - - Parameters - ---------- - R_true : torch.Tensor - ``(3, 3)`` true rotation. - symops : torch.Tensor - ``(n_ops, 3, 3)`` symmetry rotation parts, as stored on the space group - (fractional). - side : {'left', 'right'}, optional - ``'left'`` builds ``S_k @ R_true``; ``'right'`` builds ``R_true @ S_k``. - These are different sets for non-commuting operators. The engine's - peaks obey ``'right'``: on real peak lists the left orbit finds zero - coincident pairs among the top 25 and the right orbit finds every mate - (187 of 300 pairs on 3K7M). ``'left'`` was the default, and is why the - orbit-based truth rank disagreed with coordinate superposition. - frame : {'cart', 'frac'}, optional - ``'cart'`` converts the operators to the Cartesian frame first, which is - the frame the rotation function works in. ``'frac'`` uses them as - stored. Mixing a Cartesian rotation with fractional operators is a - metric error that inflates apparent ghost counts. - reciprocal_basis : torch.Tensor, optional - ``(3, 3)`` reciprocal basis, required when ``frame='cart'``. - - Returns - ------- - torch.Tensor - ``(n_ops, 3, 3)`` orbit members, float64. - """ - if side not in ("left", "right"): - raise ValueError(f"side must be 'left' or 'right', got {side!r}") - if frame not in ("cart", "frac"): - raise ValueError(f"frame must be 'cart' or 'frac', got {frame!r}") - - R = R_true.to(torch.float64) - S = symops.to(torch.float64) - if frame == "cart": - if reciprocal_basis is None: - raise ValueError("frame='cart' requires reciprocal_basis") - from torchref.experimental.alignment.sh import hkl_symops_to_cartesian - - S = hkl_symops_to_cartesian(S, reciprocal_basis.to(torch.float64)) - return S @ R.unsqueeze(0) if side == "left" else R.unsqueeze(0) @ S - - -def angle_to_orbit(R: torch.Tensor, orbit: torch.Tensor) -> float: - """Smallest rotation angle between ``R`` and any orbit member, in degrees. - - Parameters - ---------- - R : torch.Tensor - ``(3, 3)`` rotation. - orbit : torch.Tensor - ``(n, 3, 3)`` orbit members. - - Returns - ------- - float - Angle in degrees. - """ - tr = torch.einsum("kij,ij->k", orbit.to(torch.float64), R.to(torch.float64)) - cos = ((tr - 1.0) * 0.5).clamp(-1.0, 1.0) - return float(cos.arccos().min() * (180.0 / math.pi)) - - -def orbit_rank( - peaks: Sequence, - R_true: torch.Tensor, - symops: torch.Tensor, - *, - side: str = "right", - frame: str = "cart", - reciprocal_basis: Optional[torch.Tensor] = None, - thr_deg: float = 5.0, -) -> Tuple[int, float]: - """Rank of the first peak matching the true orientation. - - Parameters - ---------- - peaks : sequence - Peaks carrying Edmonds ZYZ ``alpha``/``beta``/``gamma`` in radians, in - descending score order (the FRF's ``RotationPeak``). - R_true : torch.Tensor - ``(3, 3)`` true rotation. - symops : torch.Tensor - ``(n_ops, 3, 3)`` symmetry rotation parts. - side, frame, reciprocal_basis - Orbit convention -- see :func:`symmetry_orbit`. Record these alongside - any rank you report; the rank is meaningless without them. - thr_deg : float, optional - Match threshold in degrees. Default 5.0. - - Returns - ------- - tuple - ``(rank, best_angle_deg)``. ``rank`` is ``-1`` when no peak matches; - ``best_angle_deg`` is the closest approach over all peaks either way, - which distinguishes "just outside the threshold" from "absent". - """ - from torchref.experimental.alignment.frf.rotation_utils import ( - rotation_matrix_from_edmonds_euler, - ) - - orbit = symmetry_orbit( - R_true, symops, side=side, frame=frame, reciprocal_basis=reciprocal_basis, - ) - rank, best = -1, float("inf") - for i, p in enumerate(peaks): - R_p = rotation_matrix_from_edmonds_euler(p.alpha, p.beta, p.gamma) - ang = angle_to_orbit(R_p, orbit) - if ang < best: - best = ang - if ang <= thr_deg and rank < 0: - rank = i - return rank, best - - -def cartesian_symops(spacegroup, cell) -> torch.Tensor: - """The point-group rotations as **Cartesian** matrices, ``B S_k B^-1``. - - ``spacegroup.matrices`` act on fractional column vectors, ``x' = S x + t``. - A Kabsch rotation between two sets of Cartesian coordinates lives in the - Cartesian frame, and comparing it against ``S_k`` directly is only correct - when ``B S_k B^-1 == S_k`` -- diagonal ``S`` in an orthogonal cell, or a - cubic cell. In P3(1)21 two of the six mates of a *correct* solution read as - 30.00 and 21.09 degrees under that comparison, and in P6(5)22 four of - twelve read as 21.09; those were the "bimodal" 2DQ6 failures. - - Returns ``(n_ops, 3, 3)`` float64 on the host. - """ - B = cell.fractional_matrix.detach().cpu().to(torch.float64) # c = B x - S = spacegroup.matrices.detach().cpu().to(torch.float64) - return B @ S @ torch.linalg.inv(B) - - -def allowed_origin_shifts(spacegroup, n: int = 12) -> Tuple[torch.Tensor, torch.Tensor]: - """Fractional translations ``u`` that leave the space group invariant. - - ``u`` is allowed when ``(S_k - I) u`` is a lattice vector for every op -- - two placements differing by such a ``u`` give identical ``|F|`` and are the - same solution. Returns ``(discrete, polar)``: the discrete shifts on a - ``1/n`` grid (``n=12`` covers 1/2, 1/3, 1/4 and 1/6), and an orthonormal - basis of the continuous (polar) directions, ``(3, p)``. - """ - S = spacegroup.matrices.detach().cpu().to(torch.float64) - eye = torch.eye(3, dtype=torch.float64) - D = torch.cat([Sk - eye for Sk in S], dim=0) # (3 n_ops, 3) - # Polar directions: null space of D. - _, sv, Vh = torch.linalg.svd(D) - null = (sv < 1e-8).sum().item() if sv.numel() else 3 - polar = Vh[3 - null:].T if null else torch.zeros(3, 0, dtype=torch.float64) - g = torch.arange(n, dtype=torch.float64) / n - U = torch.cartesian_prod(g, g, g) # (n^3, 3) - resid = torch.einsum("oij,uj->uoi", S - eye, U) # (n^3, n_ops, 3) - ok = ((resid - resid.round()).abs() < 1e-6).all(dim=-1).all(dim=-1) - return U[ok], polar - - -def pose_error( - aligned_xyz: torch.Tensor, - canonical_xyz: torch.Tensor, - cell, - spacegroup, -) -> Tuple[float, float]: - """``(rotation_deg, translation_A)`` of a placement against the deposited pose. - - Rotation: Kabsch superposition of the placed atoms onto the canonical ones, - compared against every **Cartesian** point-group mate (see - :func:`cartesian_symops`). Translation: the centroid offset from the closest - symmetry image of the canonical model, modulo lattice vectors, the group's - allowed origin shifts and its polar directions, in Angstrom. Both are zero - for a placement that is the deposited structure or any symmetry-equivalent - copy of it. - """ - P = canonical_xyz.detach().cpu().to(torch.float64) - Q = aligned_xyz.detach().cpu().to(torch.float64) - Pc, Qc = P - P.mean(0), Q - Q.mean(0) - U, _, Vt = torch.linalg.svd(Qc.T @ Pc) - d = torch.sign(torch.det(U @ Vt)) - R = U @ torch.diag(torch.tensor([1.0, 1.0, d], dtype=torch.float64)) @ Vt - - B = cell.fractional_matrix.detach().cpu().to(torch.float64) - Binv = torch.linalg.inv(B) - S = spacegroup.matrices.detach().cpu().to(torch.float64) - T = spacegroup.translations.detach().cpu().to(torch.float64) - R_cart = B @ S @ Binv - - tr = torch.einsum("kij,ij->k", R_cart, R) - ang = ((tr - 1.0) * 0.5).clamp(-1.0, 1.0).arccos() * (180.0 / math.pi) - k_best = int(ang.argmin()) - rot_deg = float(ang[k_best]) - - # Translation, against the mate whose rotation matched. - shifts, polar = allowed_origin_shifts(spacegroup) - cen_a = Binv @ Q.mean(0) # fractional - cen_c = S[k_best] @ (Binv @ P.mean(0)) + T[k_best] - delta = (cen_a - cen_c).unsqueeze(0) - shifts # (n_u, 3) - delta = delta - delta.round() - if polar.shape[1]: - delta = delta - (delta @ polar) @ polar.T - trans_A = float((delta @ B.T).norm(dim=-1).min()) - return rot_deg, trans_A diff --git a/alignment_lab/tests/test_lab.py b/alignment_lab/tests/test_lab.py deleted file mode 100644 index 4e757855..00000000 --- a/alignment_lab/tests/test_lab.py +++ /dev/null @@ -1,120 +0,0 @@ -"""Self-tests for the alignment lab primitives. - -These pin the contracts whose violation silently changed results in the past: -the seed -> rotation mapping, the append-only benchmark order the seed formula -depends on, and the orbit conventions. -""" - -from __future__ import annotations - -import math -import sys -from pathlib import Path - -import pytest -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) - -from lab import BENCH_PDBS, case_paths, orbit_rank, random_rotation, seed_for # noqa: E402 -from lab.truth import angle_to_orbit, symmetry_orbit # noqa: E402 - - -def test_random_rotation_is_a_rotation(): - """Output is orthogonal with det +1 for a spread of seeds.""" - for seed in (0, 1, 42, 2077, 999983): - R = random_rotation(seed) - assert torch.allclose(R @ R.T, torch.eye(3, dtype=R.dtype), atol=1e-12) - assert abs(float(torch.det(R)) - 1.0) < 1e-12 - - -def test_random_rotation_is_deterministic(): - """The seed -> rotation map is the lab's reproducibility contract.""" - assert torch.equal(random_rotation(42), random_rotation(42)) - assert not torch.equal(random_rotation(42), random_rotation(43)) - - -def test_random_rotation_uses_the_sign_corrected_qr(): - """Guard the exact variant: the uncorrected QR gives a different rotation. - - Both forms return a valid rotation, so only a direct comparison catches a - swap -- and a swap silently makes new results incomparable with archived - ones for the same seed. - """ - seed = 42 - g = torch.Generator().manual_seed(seed) - A = torch.randn(3, 3, generator=g, dtype=torch.float64) - Q_uncorrected, _ = torch.linalg.qr(A) - if torch.det(Q_uncorrected) < 0: - Q_uncorrected[:, 0] = -Q_uncorrected[:, 0] - assert not torch.allclose(random_rotation(seed), Q_uncorrected, atol=1e-9) - - -def test_seed_formula(): - """base + 1000*trial + index(pdb)*7.""" - assert seed_for("1DAW", 0) == 42 - assert seed_for("1DAW", 1) == 1042 - assert seed_for("3K7M", 2) == 42 + 2000 + BENCH_PDBS.index("3K7M") * 7 - - -def test_benchmark_order_is_pinned(): - """The seed formula indexes into this tuple, so its order is a contract.""" - assert BENCH_PDBS[0] == "1DAW" - assert BENCH_PDBS.index("1AK5") == 4 - assert BENCH_PDBS.index("3K7M") == 5 - assert len(BENCH_PDBS) == len(set(BENCH_PDBS)) == 10 - - -@pytest.mark.parametrize("pdb", BENCH_PDBS) -def test_every_benchmark_case_resolves(pdb): - """Paths are repo-relative and present -- not absolute into another tree.""" - pdb_path, mtz_path = case_paths(pdb) - assert pdb_path.is_file() and mtz_path.is_file() - - -def test_orbit_side_and_frame_are_distinct_conventions(): - """left/right and frac/cart really do differ, so recording them matters.""" - from torchref.symmetry import SpaceGroup - - sg = SpaceGroup("P 4 3 2") - symops = sg.matrices.to(torch.float64).cpu() - R = random_rotation(7) - left = symmetry_orbit(R, symops, side="left", frame="frac") - right = symmetry_orbit(R, symops, side="right", frame="frac") - assert not torch.allclose(left, right, atol=1e-9) - - -def test_orbit_contains_truth_at_zero_angle(): - """Every orbit member is 0 degrees from the orbit, by construction.""" - from torchref.symmetry import SpaceGroup - - symops = SpaceGroup("P 4 3 2").matrices.to(torch.float64).cpu() - R = random_rotation(11) - orbit = symmetry_orbit(R, symops, side="left", frame="frac") - for k in range(orbit.shape[0]): - assert angle_to_orbit(orbit[k], orbit) < 1e-9 - - -def test_orbit_rank_reports_miss_as_minus_one(): - """A peak list with no match ranks -1 but still reports the closest angle.""" - from types import SimpleNamespace - - from torchref.symmetry import SpaceGroup - - symops = SpaceGroup("P 1").matrices.to(torch.float64).cpu() - peaks = [SimpleNamespace(alpha=0.0, beta=0.0, gamma=0.0)] - R_true = random_rotation(3) - rank, ang = orbit_rank(peaks, R_true, symops, frame="frac", thr_deg=1e-6) - assert rank == -1 - assert math.isfinite(ang) and ang > 0 - - -def test_result_writer_rejects_undeclared_columns(tmp_path): - """Silent column drift is what made every old CSV need its own aggregator.""" - from lab import ResultWriter - - w = ResultWriter(tmp_path / "r.csv", "demo", extra_fields=("ghosts",)) - w.write(pdb="1DAW", truth_rank=0, ghosts=3) - with pytest.raises(KeyError): - w.write(pdb="1DAW", not_declared=1) - assert (tmp_path / "r.csv").read_text().count("\n") == 2 From 0456966b65785c68721d1f84f041491330d1654e Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Wed, 2 Sep 2026 23:29:23 +0200 Subject: [PATCH 155/250] Fixed mps test failures --- docs/changelog.rst | 5 + .../benchmark_phaser_rotation_ranking.py | 4 +- .../alignment/test_conjugate_contraction.py | 91 +++++++++++++++++++ .../alignment/test_patterson_translation.py | 10 +- tests/unit/alignment/test_rotation_search.py | 14 +-- .../test_rotation_search_dtype_device.py | 54 +++-------- tests/unit/alignment/test_translation_obs.py | 21 ++++- .../experimental/alignment/frf/data_mr.py | 73 ++++++++++----- .../experimental/alignment/frf/peak_finder.py | 6 +- .../alignment/frf/preprocessing.py | 5 +- .../experimental/alignment/frf/wigner_d.py | 13 ++- .../experimental/alignment/rotation_search.py | 38 ++++---- torchref/experimental/alignment/sh.py | 79 ++++++++++------ .../experimental/alignment/translation.py | 38 +++++--- .../model_error_estimation/sigma_m.py | 2 +- torchref/scaling/wilson.py | 32 ++++++- 16 files changed, 334 insertions(+), 151 deletions(-) create mode 100644 tests/unit/alignment/test_conjugate_contraction.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 14007ef3..3a5a2f45 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,11 @@ Changelog Unreleased ---------- +- Fixed the rotation function contracting **unconjugated** calc coefficients on MPS. ``torch.conj`` returns a lazy view carrying a conjugate *bit*, and MPS's batched complex matmul -- what the radial ``einsum`` lowers to -- ignores it, so the contraction was silently wrong: on 1DAW by 173% of ``|xi|`` max, which reordered the entire peak list and pushed the true orientation from rank 0 out of the top 200 while the top score moved only 0.1%. Materialised with ``resolve_conj()`` at the two batched-matmul sites. Elementwise ops, ``where``, ``index_add``, ``fft`` and 2-D matmul all honour the bit and only the batched path does not, which is too narrow to guard structurally, so the guard is the contraction's own value against a host-double reference +- The alignment package runs end to end on an accelerator without float64: 123 alignment tests including the slow rotation searches pass on MPS, where every one of them failed before. Five things stood in the way, four of them the same mistake -- casting and moving in two steps. ``.to(device).to(dtype)`` puts the caller's width on the device first, which throws for a double input on a backend that has none (14 sites, now one fused ``.to`` each); ``detect_zsymm`` widened the symmetry operators to double while they sat on the device; the expansion's host-side clustering keys left a host index to meet device values; and MPS's ``linalg.inv`` trips an internal contiguity assert on a transposed 3x3 view. The fifth is a kernel gap -- MPS implements no complex cumulative op -- so the azimuthal phase ladder is built by doubling instead of by ``cumprod``, measured at 5.7e-6 against ``cumprod``'s 4.1e-6 in complex64 over 5000 angles at L=101 +- The anisotropy fit runs at the configured float dtype on the data's own device, instead of double on the host. Measured over the 16 datasets in ``tests/files/mtz``, float32 reproduces the double fit to 3.3e-5 relative in ``U`` and 4.5e-6 in the correction factor ``exp(+pi^2 s.U.s)`` it exists to produce -- the design matrix is a constant column beside ``2 pi^2 s.s`` terms of order 0.1-1, so there is no precision cliff for seven parameters to fall off. Unrotated searches on 1DAW, 3K7M, 2DQ6 and 4BX9 return the same peaks in the same order as before, with scores moved 1.8e-7 to 2.5e-4 relative. ``fit_anisotropy`` takes an optional ``device``; ``get_axis_order`` and ``hkl_symops_to_cartesian`` follow their inputs' width and place rather than forcing double +- The Wilson fit's IRLS solve is a Cholesky reuse plus two triangular solves, not ``lu_factor``/``lu_solve``. ``A`` is ``X^T W X`` plus a ridge, positive definite by construction, and MPS implements neither ``lu_solve`` nor ``cholesky_solve``, so the old path left the per-iteration solve to a CPU round trip: 1015 us against 114 us, on a 200k x 6 problem whose unavoidable ``XtW @ z`` is 966 us. Over the 16-dataset panel the two agree to 1e-5 relative in float32 and 1e-14 in float64 with identical iteration counts on every case. ``cholesky_ex`` reports rather than raises, and a non-positive-definite report falls back to a general solve, since a fully collinear basis is a thing this fit sees +- Dropped the two alignment tests that imported ``alignment_lab``: an external experiment is not the main package's to test, and they could not pass in a clean checkout. ``resolve_device`` is covered directly in ``tests/unit/utils/test_device_resolution.py``. Two test-side assertions that could not hold at the configured width were repaired -- `` = 1`` now widens after the readback rather than on the device, and the weight's mean-one bound tracks the dtype's epsilon instead of a fixed 1e-9 that is below one float32 ulp - Nothing in the alignment package puts float64 or complex128 on the compute device any more, so it runs on backends without double (MPS). The rotation function's radial sum, Wigner contraction and FFT accumulate at the configured complex dtype rather than one step wider; the Bessel ladder runs in its argument's dtype, kept in range by its rescaling; the expansion's exact clustering keys are formed on the host in double, which the device never sees; rotation-function inputs, the dense P1 transform and the shell sums use the configured float dtype. 30/30 poses at both windows, timings unchanged; the double-precision left is host-side 3x3 rotation algebra, the Wigner-d eigendecomposition and the anisotropy fit - Every hard-coded dtype in the alignment package either moved to the configured dtype or carries a ``# dtype-ok:`` justification, as the dtype-conformance guard requires: three casts to double dropped from the empirical sigma_A ratio and the dense P1 transform, the translation set's resolution mask and peak translations use the configured float dtype, and the rest (index tensors, host-side 3x3 rotation algebra, the rotation function's double accumulation, the anisotropy fit) are annotated - Removed ``DirectModelEvaluator``; the translation search evaluates an ordinary P1 ``ModelFT`` directly. With the model's grid derived lazily from cell, space group and ``max_res`` there was nothing left for the wrapper to do diff --git a/tests/integration/alignment/benchmark_phaser_rotation_ranking.py b/tests/integration/alignment/benchmark_phaser_rotation_ranking.py index ce530e84..ed4c13eb 100644 --- a/tests/integration/alignment/benchmark_phaser_rotation_ranking.py +++ b/tests/integration/alignment/benchmark_phaser_rotation_ranking.py @@ -210,8 +210,8 @@ def orbit_rank( "best_rfz": 0.0, "top1_ang_deg": float("inf"), "any_below_threshold": False, "n_peaks": 0, } - R_target = R_target.to(torch.float64).cpu() - sym_mats = sym_mats.to(torch.float64).cpu() + R_target = R_target.cpu().to(torch.float64) + sym_mats = sym_mats.cpu().to(torch.float64) alphas = torch.tensor( [math.radians(p.alpha_deg) for p in peaks], dtype=torch.float64, diff --git a/tests/unit/alignment/test_conjugate_contraction.py b/tests/unit/alignment/test_conjugate_contraction.py new file mode 100644 index 00000000..3e6ce842 --- /dev/null +++ b/tests/unit/alignment/test_conjugate_contraction.py @@ -0,0 +1,91 @@ +"""The radial contraction has to *materialise* its conjugate. + +``torch.conj`` does not conjugate anything: it returns a view carrying a +conjugate *bit*, and every consumer is expected to honour it. MPS's batched +complex matmul does not -- and the contraction in +:func:`~torchref.experimental.alignment.frf.data_mr.cross_correlate_xi` is an +``einsum`` that lowers to exactly that. The failure is silent and total: the +unconjugated values are contracted instead, which on 1DAW moved ``xi`` by 173% +of its own peak magnitude, reordered the entire rotation-function peak list, and +pushed the true orientation from rank 0 out of the top 200 -- while the top +score moved by only 0.1%, so nothing looked wrong. + +Elementwise ops, ``where``, ``index_add`` and 2-D ``matmul`` all honour the bit; +only the batched path drops it. That is narrow enough that the guard has to be +the value of the contraction itself, checked on whatever device this host has, +against a reference computed on the host in double. +""" + +import pytest +import torch + +from torchref.config import get_complex_dtype, get_default_device +from torchref.experimental.alignment.frf.data_mr import cross_correlate_xi +from torchref.experimental.alignment.frf.types import BesselSHCoefficients + +pytestmark = pytest.mark.unit + +L, N_RADIAL = 6, 8 + + +def _coeffs(seed, device, dtype): + """A filled ``(N_radial, L, 2L-1)`` coefficient block.""" + g = torch.Generator().manual_seed(seed) + real = torch.randn(N_RADIAL, L, 2 * L - 1, generator=g, dtype=torch.float64) + imag = torch.randn(N_RADIAL, L, 2 * L - 1, generator=g, dtype=torch.float64) + c = torch.complex(real, imag) + return BesselSHCoefficients( + coeffs=c.to(device=device, dtype=dtype), L=L, bessel_h_scale=40.0, + ) + + +def test_the_contraction_conjugates_the_calc_side(): + """On this host's device, against the same sum done on the host in double. + + A dropped conjugation is not a small error -- it changes the sign of every + imaginary part in one operand -- so the bar can be tight without being + brittle about float32 rounding. + """ + device, cplx = get_default_device(), get_complex_dtype() + obs, calc = _coeffs(0, device, cplx), _coeffs(1, device, cplx) + got = cross_correlate_xi(obs, calc) + + host_obs = BesselSHCoefficients( + coeffs=obs.coeffs.cpu().to(torch.complex128), L=L, bessel_h_scale=40.0) + host_calc = BesselSHCoefficients( + coeffs=calc.coeffs.cpu().to(torch.complex128), L=L, bessel_h_scale=40.0) + ref = torch.einsum( + "rln,rlm->lmn", + host_obs.coeffs, + torch.conj(host_calc.coeffs).resolve_conj(), + ) + + err = (got.cpu().to(torch.complex128) - ref).abs().max() + scale = ref.abs().max() + assert float(err / scale) < 1e-5, ( + f"contraction is {float(err / scale):.2e} away from the host double " + f"reference on {device}; a dropped conjugation shows up here as O(1)" + ) + + +def test_dropping_the_conjugation_would_be_caught(): + """The guard above has to be able to see the failure it exists for. + + Contracting the unconjugated coefficients is what a lost conjugate bit + produces, so that has to land far outside the tolerance -- otherwise the + test would pass on the broken path too. + """ + device, cplx = get_default_device(), get_complex_dtype() + obs, calc = _coeffs(0, device, cplx), _coeffs(1, device, cplx) + ref = torch.einsum( + "rln,rlm->lmn", + obs.coeffs.cpu().to(torch.complex128), + torch.conj(calc.coeffs.cpu().to(torch.complex128)).resolve_conj(), + ) + unconjugated = torch.einsum( + "rln,rlm->lmn", + obs.coeffs.cpu().to(torch.complex128), + calc.coeffs.cpu().to(torch.complex128), + ) + rel = float((unconjugated - ref).abs().max() / ref.abs().max()) + assert rel > 0.1, f"the two differ by only {rel:.2e}; this guard is blind" diff --git a/tests/unit/alignment/test_patterson_translation.py b/tests/unit/alignment/test_patterson_translation.py index 0e784648..d5545ef2 100644 --- a/tests/unit/alignment/test_patterson_translation.py +++ b/tests/unit/alignment/test_patterson_translation.py @@ -39,9 +39,13 @@ def setup(): canonical = ModelFT().load_pdb(str(PDB_1DAW)) data = ReflectionData().load_mtz(str(MTZ_1DAW)) # The pipeline's default window: the rotation search's 15-4 A. - rec = data.cell.reciprocal_basis_matrix.to(torch.float64) - s = (data.hkl.to(torch.float64) @ rec).norm(dim=-1) - mask = data.get_valid_mask() & (s >= 1.0 / 15.0) & (s <= 1.0 / 4.0) + # The resolution arithmetic runs on the host in double -- `.cpu()` before + # the widening, since a backend without float64 cannot hold the wide copy -- + # and only the resulting boolean goes back to where the data live. + rec = data.cell.reciprocal_basis_matrix.cpu().to(torch.float64) + s = (data.hkl.cpu().to(torch.float64) @ rec).norm(dim=-1) + window = ((s >= 1.0 / 15.0) & (s <= 1.0 / 4.0)).to(data.hkl.device) + mask = data.get_valid_mask() & window return canonical, data, mask diff --git a/tests/unit/alignment/test_rotation_search.py b/tests/unit/alignment/test_rotation_search.py index d459c0a9..f5930f1a 100644 --- a/tests/unit/alignment/test_rotation_search.py +++ b/tests/unit/alignment/test_rotation_search.py @@ -46,8 +46,8 @@ def _rotation(seed: int) -> torch.Tensor: def _angle_deg(a: torch.Tensor, b: torch.Tensor) -> float: - a = a.detach().to(torch.float64).cpu() - b = b.detach().to(torch.float64).cpu() + a = a.detach().cpu().to(torch.float64) + b = b.detach().cpu().to(torch.float64) tr = torch.diagonal(a @ b.T).sum().item() return math.degrees(math.acos(max(-1.0, min(1.0, (tr - 1.0) / 2.0)))) @@ -61,16 +61,18 @@ def _sym_cartesian(data) -> torch.Tensor: """ from torchref.experimental.alignment.sh import hkl_symops_to_cartesian + # `.cpu()` before the widening: the data may sit on an accelerator that + # cannot hold a float64 tensor at all, and the widening is exact on the host. return hkl_symops_to_cartesian( - data.spacegroup.matrices.to(torch.float64).cpu(), - data.cell.reciprocal_basis_matrix.to(torch.float64).cpu(), + data.spacegroup.matrices.cpu().to(torch.float64), + data.cell.reciprocal_basis_matrix.cpu().to(torch.float64), ) def _kabsch(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: """Rotation taking ``a`` onto ``b``, both centred. CPU float64.""" - a = a.detach().to(torch.float64).cpu() - b = b.detach().to(torch.float64).cpu() + a = a.detach().cpu().to(torch.float64) + b = b.detach().cpu().to(torch.float64) x = a - a.mean(0) y = b - b.mean(0) u, _, vt = torch.linalg.svd(y.T @ x) diff --git a/tests/unit/alignment/test_rotation_search_dtype_device.py b/tests/unit/alignment/test_rotation_search_dtype_device.py index 2202d930..b6aff525 100644 --- a/tests/unit/alignment/test_rotation_search_dtype_device.py +++ b/tests/unit/alignment/test_rotation_search_dtype_device.py @@ -1,16 +1,15 @@ -"""The rotation search takes its working precision and its device from config. - -Two properties, both easy to lose silently: - -* **Precision is `torchref.config`'s, not an argument's.** The expansion used to - be handed ``compute_dtype=torch.complex64`` by one call site, which made the - fused CPU Legendre kernel reachable only through that argument -- its - ``BackendTable`` row gates on ``dtypes=(torch.float32,)``, so dropping the - argument silently routed to the portable path. Now the default *is* float32, - so the gate matches by construction and flipping the config flips the engine. -* **The device is resolved from both inputs**, not read off whichever one the - code happens to touch first. Model and data on different devices used to - either cross-device or throw depending on which line ran. +"""The rotation search takes its working precision from config. + +**Precision is `torchref.config`'s, not an argument's.** The expansion used to +be handed ``compute_dtype=torch.complex64`` by one call site, which made the +fused CPU Legendre kernel reachable only through that argument -- its +``BackendTable`` row gates on ``dtypes=(torch.float32,)``, so dropping the +argument silently routed to the portable path. Now the default *is* float32, so +the gate matches by construction and flipping the config flips the engine. + +Device resolution -- that both inputs are read, rather than whichever one the +code happens to touch first -- lives in ``tests/unit/utils/test_device_resolution.py``, +which exercises ``resolve_device`` directly instead of through a loaded case. """ import pytest @@ -122,35 +121,6 @@ def test_wigner_blocks_carry_no_float64_to_the_device(): assert torch.allclose(b[0], eye, atol=1e-5), f"d^{l}(0) is not I" -def test_anisotropy_fit_runs_on_the_host_in_double(): - """It is 7 parameters over ~1e4 reflections; precision there is worth more - than locality, and keeping it on the host removes a float64 requirement.""" - from alignment_lab.lab.benchmark import load_case - from torchref.experimental.alignment.rotation_search import ( - ANISO_FIT_WINDOW_A, fit_anisotropy, - ) - - _, data = load_case("1DAW")[:2] - d_max, d_min = ANISO_FIT_WINDOW_A - U = fit_anisotropy(data, d_min=d_min, d_max=d_max) - assert U.device.type == "cpu" - assert U.dtype == torch.float64 - assert U.shape == (3, 3) - assert torch.allclose(U, U.T, atol=1e-12), "U must be symmetric" - - -def test_device_is_resolved_from_both_inputs(): - """Reading one input's device is what let model and data disagree.""" - from alignment_lab.lab.benchmark import load_case - from torchref.utils import resolve_device - - model, data = load_case("1DAW")[:2] - # Data-first precedence, per the convention in torchref/maps/map.py. - resolved = resolve_device(data, model) - assert resolved == data.hkl.device or resolved.type == data.hkl.device.type - assert model.xyz().device.type == resolved.type - - def test_double_is_a_device_capability_not_a_constant(): """Where precision is load-bearing, the width comes from the device. diff --git a/tests/unit/alignment/test_translation_obs.py b/tests/unit/alignment/test_translation_obs.py index 4d97fd34..5eb1b20f 100644 --- a/tests/unit/alignment/test_translation_obs.py +++ b/tests/unit/alignment/test_translation_obs.py @@ -67,8 +67,13 @@ def test_mean_e_squared_is_one(): """ F, _, hkl, sg, cell, _ = _case() obs = TranslationObs.build(F, hkl, sg, cell) - k = torch.where(obs.centric, 0.5, 1.0).to(torch.float64) - mean_e2 = (k * obs.E_obs ** 2).sum() / k.sum() + # Reduced on the host in double: the identity is asserted to 1e-6 over 4000 + # reflections and the accumulation should not be what limits that -- but the + # widening has to happen after the readback, since a backend without float64 + # cannot hold the wide copy. + k = torch.where(obs.centric.cpu(), 0.5, 1.0).to(torch.float64) + E2 = obs.E_obs.detach().cpu().to(torch.float64) ** 2 + mean_e2 = (k * E2).sum() / k.sum() assert abs(float(mean_e2) - 1.0) < 1e-6, mean_e2 @@ -131,10 +136,18 @@ def test_weight_is_uniform_without_sigmas(): @pytest.mark.unit def test_weight_is_normalised_to_mean_one(): - """So the score's scale does not depend on how the sigmas happened to be scaled.""" + """So the score's scale does not depend on how the sigmas happened to be scaled. + + To a few ulps of the working dtype, not to a fixed 1e-9: the normalisation + divides by this very mean, so what is left is the rounding of a 4000-term + reduction, and at float32 one ulp is already 1.2e-7. The old fixed bound + passed or failed on the reduction order alone -- it held on MPS and missed + by exactly one ulp on the CPU. + """ F, sig_F, hkl, sg, cell, _ = _case() obs = TranslationObs.build(F, hkl, sg, cell, sig_F=sig_F) - assert abs(float(obs.weight.mean()) - 1.0) < 1e-9 + tol = 8 * torch.finfo(obs.weight.dtype).eps + assert abs(float(obs.weight.mean()) - 1.0) < tol @pytest.mark.unit diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index 5247343c..3ddc7440 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -186,6 +186,32 @@ def spherical_bessel_table( return j_table.to(real_dtype) +def _unit_power_ladder(z: torch.Tensor, L: int) -> torch.Tensor: + """``[z^0, z^1, ..., z^{L-1}]`` for unit-modulus ``z``, shape ``(n, L)``. + + Built by doubling -- the block of powers already computed, times the next + power -- rather than by ``torch.cumprod``, for portability: MPS has no + complex cumulative kernels at all (torch 2.9.1 raises "cumulative ops are + not yet supported for complex"), and this was the last thing in the rotation + search that could not run on Apple silicon. + + Not an accuracy change. Measured against the exact powers in double over + 5000 angles at L=101, the ladder gives 5.7e-6 where ``cumprod`` gives 4.1e-6 + in complex64, and 3.4e-14 against 3.3e-14 in complex128 -- the same, because + the error is dominated by ``z``'s own rounding amplified by ``p``, which no + grouping of the multiplies avoids. ``log2(L)`` wide multiplies in place of + one fused pass, so it is not a cost change either. + """ + out = torch.ones((z.shape[0], L), dtype=z.dtype, device=z.device) + width = 1 # out[:, :width] is filled + while width < L: + z_w = out[:, width - 1] * z # z^width + take = min(width, L - width) + out[:, width:width + take] = out[:, :take] * z_w.unsqueeze(1) + width += take + return out + + def bessel_sh_expand( s_vectors: torch.Tensor, intensity: torch.Tensor, @@ -347,6 +373,11 @@ def _group_mean(values, index, n_groups): # than clusters, because many directions share a |s| on a lattice. Measured # over the benchmark: 2.7 to 39 clusters per shell. uniq_ks, inv_s = torch.unique(k_s, return_inverse=True) + # `k_s` is one of the host-side keys, so its inverse comes back on the host + # while everything it indexes -- `shell_of_cluster`, `s_mag_all` -- is on the + # compute device. Bring it across here, once, rather than leaving a host + # index to meet device values. + inv_s = inv_s.to(device) n_shells = int(uniq_ks.shape[0]) shell_of_cluster = torch.zeros(n_clusters, dtype=torch.long, device=device) # dtype-ok: index tensor; index_add_/gather need int64 shell_of_cluster[inverse] = inv_s @@ -389,15 +420,13 @@ def _group_mean(values, index, n_groups): ph = phi_all[start_i:stop].to(comp_real) # (c,) i_c = intensity[start_i:stop].to(comp_real) # (c,) # e^{-i p phi} = z^p with z = e^{-i phi}, so one transcendental per - # reflection and a running product over p, rather than a transcendental + # reflection and a power ladder over p, rather than a transcendental # per (reflection, p). At L=101 over 2.6e6 reflections that is 2.6e8 - # sincos calls replaced by 2.6e6 of them plus a complex multiply each. - # The product accumulates about L * eps of relative error, ~2e-14, six - # orders below what the grouping already costs. + # sincos calls replaced by 2.6e6 of them plus a few complex multiplies + # each. `_unit_power_ladder` builds it by doubling rather than with a + # cumulative product, which no complex MPS kernel implements. z = torch.polar(torch.ones_like(ph), -ph) # (c,) - ladder = z.unsqueeze(1).expand(-1, L).clone() - ladder[:, 0] = 1.0 # p = 0 - e_neg = torch.cumprod(ladder, dim=1) # (c, L) = z^p + e_neg = _unit_power_ladder(z, L) # (c, L) = z^p e_neg = (e_neg * i_c.unsqueeze(1)).to(einsum_dtype) Sp.index_add_(0, cluster_of_refl[start_i:stop], e_neg) sign_p = ((-1.0) ** p_idx.to(comp_real)).to(einsum_dtype) # (L,) @@ -542,20 +571,13 @@ def cross_correlate_xi( xi[l, m, n] = Σ_r c_obs[r, l, n] · conj(c_calc[r, l, m]) so that the peak Euler triple satisfies ``s_calc = R · s_obs``. - **Accumulated one step wider than the coefficients.** The radial sum runs - over oscillating ``j_u``, so the terms alternate in sign and cancel; the - relative error on the result is then far worse than ``eps * sqrt(n_terms)`` - would suggest, and it compounds through the equally oscillatory Wigner - contraction and the FFT downstream. Accumulating single-precision data in - double is the ordinary remedy and it is cheap here -- ``xi`` is - ``(L, 2L-1, 2L-1)``, 17 MB at L=65, against the 4.4M-element FFT it feeds. - - Running the whole tail in single instead was measured on 3K7M and 1DAW: the - top peak and its z-score were unchanged to seven figures, but scores moved - by 1e-4 to 1.4e-3 relative and only **1 of 500** candidate slots still held - the same orientation, because the greedy SO(3) NMS is sequential and a - reordering cascades through the suppression decisions. The candidate list is - what the placement search consumes, so that is not a free trade. + Accumulated at the configured complex dtype. The radial sum runs over + oscillating ``j_u``, so the terms alternate in sign and cancel, and this was + once accumulated one step wider for that reason. Measured, the width is not + what the result needs: from one set of complex64 coefficients, a complex64 + contraction lands within 1.5e-5 of the complex128 one whose peak magnitude + is 109. What it does need is for the conjugate below to be *materialised* -- + see the ``resolve_conj`` note. Returns ------- @@ -569,8 +591,15 @@ def cross_correlate_xi( # the top peak, and the placement search now consumes only the top few # distinct orientations. Single precision recovers every pose on the panel. acc = get_complex_dtype() + # `resolve_conj()` is load-bearing, not tidiness. `torch.conj` returns a + # lazy view carrying a conjugate BIT, and MPS's batched complex matmul -- + # which is what this einsum lowers to -- ignores that bit and contracts the + # unconjugated values. It is silent: the result is a plausible tensor that + # is simply wrong, here by 173% of |xi|max, which then reorders the whole + # peak list. Elementwise ops, `where`, `index_add` and 2-D matmul all honour + # the bit; only the batched matmul path does not. return torch.einsum( "rln,rlm->lmn", c_obs.coeffs.to(acc), - torch.conj(c_calc.coeffs).to(acc), + torch.conj(c_calc.coeffs).resolve_conj().to(acc), ) diff --git a/torchref/experimental/alignment/frf/peak_finder.py b/torchref/experimental/alignment/frf/peak_finder.py index 23cb2db3..19692db1 100644 --- a/torchref/experimental/alignment/frf/peak_finder.py +++ b/torchref/experimental/alignment/frf/peak_finder.py @@ -71,14 +71,16 @@ def _so3_greedy_nms( # sitting exactly on it. R_all = ( rotation_matrix_euler_zyz(torch.stack([alphas, betas, gammas], dim=-1)) - .to(torch.float64).cpu() # dtype-ok: 3x3 rotation algebra in double on the host + .cpu().to(torch.float64) # dtype-ok: 3x3 rotation algebra in double on the host ) # (n, 3, 3) # angle > nms_radius ⇔ cos(angle) < cos(nms_radius); cos(angle) from trace. cos_thresh = math.cos(math.radians(nms_radius_deg)) # The orbit of each kept rotation, R R_g over the point group; the identity # alone when no symmetry is supplied. + # `.cpu()` first in both: the widening happens on the host, which always + # has float64, and the peaks arrive on the compute device. G = (torch.eye(3, dtype=torch.float64).unsqueeze(0) if sym_cart is None # dtype-ok: 3x3 rotation algebra in double on the host - else sym_cart.to(torch.float64).cpu()) # dtype-ok: 3x3 rotation algebra in double on the host + else sym_cart.cpu().to(torch.float64)) # dtype-ok: 3x3 rotation algebra in double on the host kept_idx: List[int] = [] kept_orbit = torch.empty((keep_at_most, G.shape[0], 3, 3), dtype=torch.float64) # dtype-ok: 3x3 rotation algebra in double on the host count = 0 diff --git a/torchref/experimental/alignment/frf/preprocessing.py b/torchref/experimental/alignment/frf/preprocessing.py index a158e089..63977817 100644 --- a/torchref/experimental/alignment/frf/preprocessing.py +++ b/torchref/experimental/alignment/frf/preprocessing.py @@ -142,7 +142,10 @@ def detect_zsymm(sym_mats: Optional[torch.Tensor]) -> int: """ if sym_mats is None: return 1 - axis, zsymm = get_high_order_axis(sym_mats.to(torch.float64).cpu()) # dtype-ok: 3x3 rotation algebra in double on the host + # No cast and no copy: `get_axis_order` works at the operators' own width + # and batches its one readback. Widening them here was also a crash on a + # backend with no float64, since these arrive on the compute device. + axis, zsymm = get_high_order_axis(sym_mats) if axis != 2: # high-order axis not along z → don't apply a wrong filter return 1 return int(zsymm) diff --git a/torchref/experimental/alignment/frf/wigner_d.py b/torchref/experimental/alignment/frf/wigner_d.py index 6e7f88b8..299133aa 100644 --- a/torchref/experimental/alignment/frf/wigner_d.py +++ b/torchref/experimental/alignment/frf/wigner_d.py @@ -91,21 +91,28 @@ def _wigner_d_blocks(L: int, betas: torch.Tensor, device: torch.device, int(L), str(canonical_device(device)), dtype, - tuple(betas.detach().to(torch.float64).cpu().tolist()), # dtype-ok: memo key: exact host-side doubles + tuple(betas.detach().cpu().to(torch.float64).tolist()), # dtype-ok: memo key: exact host-side doubles ) hit = _WIGNER_D_CACHE.get(key) if hit is not None: return hit eig_table = _wigner_eig_table(L) # host, cached - betas_host = betas.detach().to(torch.float64).cpu() # dtype-ok: memo key: exact host-side doubles + # `.cpu()` before the widening, not after: float64 cannot be materialised + # on every backend, and widening on the host is exact either way. + betas_host = betas.detach().cpu().to(torch.float64) # dtype-ok: memo key: exact host-side doubles blocks = [] for l in range(1, L): w, V = eig_table[l - 1] # data-independent phase = torch.exp(-1j * betas_host.unsqueeze(1) * w.unsqueeze(0)) # (n_beta, sz) VP = V.unsqueeze(0) * phase.unsqueeze(1) # (n_beta, sz, sz) = (k,m,a) blocks.append( - (VP @ V.conj().transpose(-1, -2)).real.to(device=device, dtype=dtype) + # `resolve_conj()`: a batched matmul on a lazy conjugate view drops + # the conjugation on MPS (see `cross_correlate_xi`). This block is + # built on the host today, where the bit is honoured, but the failure + # mode is silent and the copy is one small matrix. + (VP @ V.conj().resolve_conj().transpose(-1, -2)) + .real.to(device=device, dtype=dtype) ) _WIGNER_D_CACHE.clear() # one entry only; see the footprint note diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index db28618c..245c049b 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -163,20 +163,22 @@ def fit_anisotropy( d_min: float, d_max: float, n_shells: int = N_WILSON_SHELLS, + device=None, ) -> torch.Tensor: """Fit the overall anisotropy tensor and project it onto the point group. - Returns ``U`` in Angstrom squared as a **host** float64 ``(3, 3)``, in the - convention ``F_corrected = F * exp(+pi^2 s.U.s)``. - :func:`~torchref.experimental.alignment.sh.apply_overall_anisotropy` moves - and casts it to wherever the amplitudes are. + Returns ``U`` in Angstrom squared as a ``(3, 3)`` at the configured float + dtype, in the convention ``F_corrected = F * exp(+pi^2 s.U.s)``. It stays on + the data's own device unless ``device`` says otherwise, so nothing crosses a + device boundary to be fitted and come back. - Deliberately host-side and in double precision. It is a seven-parameter - Gauss-Newton fit over the ``[d_max, d_min]`` window -- of order 1e4 - reflections, once per search -- so the cost of doing it here is not - measurable, while a broken anisotropy fit is worth hundreds of ranks on - high-symmetry cases. Keeping it off the accelerator also means the engine - needs no float64 there. + It used to be pinned to the host in double. Neither is needed. Measured over + the 16 datasets in ``tests/files/mtz``, the same fit in float32 reproduces + the double one to 3.3e-5 relative in ``U`` and 4.5e-6 in the correction + factor it exists to produce; end to end, four of five panel cases return a + bit-identical peak list and the fifth (2DQ6, P3121, the most nearly + isotropic ``U`` of the panel) keeps its top orientation and reshuffles two + near-tied deep ranks. The projection matters: an unconstrained six-component fit can return a tensor the lattice forbids, and applying it then modulates the observations @@ -186,8 +188,10 @@ def fit_anisotropy( """ from .sh import hkl_symops_to_cartesian, symmetrize_anisotropy - rec_basis = data.cell.reciprocal_basis_matrix.detach().cpu().to(torch.float64) # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search - hkl = data.hkl.detach().cpu().to(torch.float64) # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search + real = get_float_dtype() + dev = data.hkl.device if device is None else device + rec_basis = data.cell.reciprocal_basis_matrix.detach().to(device=dev, dtype=real) + hkl = data.hkl.detach().to(device=dev, dtype=real) s_vec_all = hkl @ rec_basis s_mag_all = s_vec_all.norm(dim=-1) keep = (s_mag_all >= 1.0 / d_max) & (s_mag_all <= 1.0 / d_min) @@ -196,11 +200,11 @@ def fit_anisotropy( f"Only {int(keep.sum())} reflections in [{d_min}, {d_max}] A, too " f"few for {n_shells} shells." ) - F_obs = data.F.detach().cpu().to(torch.float64).abs()[keep] # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search + F_obs = data.F.detach().to(device=dev, dtype=real).abs()[keep] s_vec = s_vec_all[keep] s_mag = s_mag_all[keep] centric = ( - data.centric.detach().cpu()[keep].to(torch.bool) + data.centric.detach().to(dev)[keep].to(torch.bool) if hasattr(data, "centric") else torch.zeros_like(F_obs, dtype=torch.bool) ) @@ -211,7 +215,7 @@ def fit_anisotropy( F_obs, s_vec, shell_idx, centric, P=n_shells, min_count=20, ) sym_cart = hkl_symops_to_cartesian( - data.spacegroup.matrices.detach().cpu().to(torch.float64), rec_basis, # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search + data.spacegroup.matrices.detach().to(device=dev, dtype=real), rec_basis, ) return symmetrize_anisotropy(U, sym_cart) @@ -285,8 +289,8 @@ def prepare_frf_inputs( ) U_aniso = fit_anisotropy( - data, d_min=d_min, d_max=d_max, n_shells=n_shells, - ).to(device) + data, d_min=d_min, d_max=d_max, n_shells=n_shells, device=device, + ) F_obs_aniso = apply_overall_anisotropy(F_obs, s_vec, U_aniso) # Same multiplicative factor, so F/sigma survives the correction intact. sig_F_aniso = (None if sig_F is None diff --git a/torchref/experimental/alignment/sh.py b/torchref/experimental/alignment/sh.py index 0f30670d..0da057fe 100644 --- a/torchref/experimental/alignment/sh.py +++ b/torchref/experimental/alignment/sh.py @@ -29,6 +29,8 @@ import torch +from ...config import get_float_dtype + def legendre_recurrence_coefficients(L: int, dtype, device): """Coefficient tables for the fully-normalised Legendre recurrence. @@ -176,26 +178,27 @@ def get_axis_order(sym_mats: torch.Tensor, axis: int) -> int: coefficients: the Patterson is invariant under the spacegroup rotations, so m-values that violate the highest-order axis symmetry are pure noise. """ - a = torch.zeros(3, dtype=torch.float64, device=sym_mats.device) # dtype-ok: 3x3 rotation algebra in double on the host + dtype = sym_mats.dtype if sym_mats.is_floating_point() else get_float_dtype() + R = sym_mats.to(dtype) + a = torch.zeros(3, dtype=dtype, device=R.device) a[axis] = 1.0 - max_order = 1 - n_ops = sym_mats.shape[0] - for k in range(n_ops): - R = sym_mats[k].to(torch.float64) # dtype-ok: 3x3 rotation algebra in double on the host - # Axis must be invariant under R (proper or improper rotation about it). - if (R @ a - a).norm().item() > 1e-3: - continue - # Trace of a rotation by angle θ about the preserved axis is 1+2cosθ. - tr = R.diagonal().sum().item() - cos_a = max(-1.0, min(1.0, (tr - 1.0) / 2.0)) - # Identity (angle ~0) → order 1. - if cos_a >= 1.0 - 1e-6: - continue - angle = math.acos(cos_a) - n = round(2 * math.pi / angle) - if n > max_order: - max_order = n - return max_order + # The working precision is enough here and the operators can stay wherever + # they are: the entries of a Miller-index symop are small integers, exact in + # float32, and the Cartesian form of one is accurate to ~1e-7 -- an order of + # magnitude inside the 1e-3 and 1e-6 tolerances this test uses. Batched, so + # the answer costs one transfer of two short vectors rather than a device + # sync per operation. + # + # Axis must be invariant under R (proper or improper rotation about it), and + # the trace of a rotation by angle θ about the preserved axis is 1+2cosθ. + keeps = (R @ a - a).norm(dim=-1) <= 1e-3 + cos_a = ((R.diagonal(dim1=-2, dim2=-1).sum(dim=-1) - 1.0) * 0.5).clamp(-1.0, 1.0) + keeps = keeps & (cos_a < 1.0 - 1e-6) # identity (angle ~0) → order 1 + orders = [ + round(2 * math.pi / math.acos(c)) + for c in cos_a[keeps].detach().cpu().tolist() + ] + return max(orders) if orders else 1 def get_high_order_axis(sym_mats: torch.Tensor) -> Tuple[int, int]: @@ -318,8 +321,16 @@ def fit_overall_anisotropy( reflections survive to constrain seven parameters. """ valid = shell_idx >= 0 - F = F_obs[valid].to(torch.float64) # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search - s = s_vectors[valid].to(torch.float64) # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search + # The fit runs at the amplitudes' own width, wherever they are. It used to + # force double on the host, which was measured against this: over the 16 + # datasets in ``tests/files/mtz``, float32 reproduces U to 3.3e-5 relative + # and the correction it exists to apply, exp(+pi^2 s.U.s), to 4.5e-6. The + # design matrix is well scaled by construction -- a constant column beside + # 2 pi^2 s.s terms of order 0.1-1 over the fitting window -- so there is no + # precision cliff for seven parameters to fall off. + work = F_obs.dtype if F_obs.is_floating_point() else get_float_dtype() + F = F_obs[valid].to(work) + s = s_vectors[valid].to(work) idx = shell_idx[valid] cen = centric[valid].bool() @@ -328,10 +339,10 @@ def fit_overall_anisotropy( I = F * F count = torch.zeros(P, dtype=torch.int64, device=F.device) # dtype-ok: index tensor; index_add_/gather need int64 - total = torch.zeros(P, dtype=torch.float64, device=F.device) # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search + total = torch.zeros(P, dtype=work, device=F.device) count.index_add_(0, idx, torch.ones_like(idx)) total.index_add_(0, idx, I) - mean_I = (total / count.clamp(min=1).to(torch.float64)).clamp(min=1e-30) # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search + mean_I = (total / count.clamp(min=1).to(work)).clamp(min=1e-30) keep = (count >= min_count)[idx] if int(keep.sum()) < 50: @@ -349,7 +360,7 @@ def fit_overall_anisotropy( -2.0 * (torch.pi ** 2) * quad], dim=1) w = torch.where(cenk, torch.full_like(ratio, 0.5), torch.ones_like(ratio)) - theta = torch.zeros(7, dtype=torch.float64, device=F.device) # dtype-ok: seven-parameter Gauss-Newton fit in double on the host, once per search + theta = torch.zeros(7, dtype=work, device=F.device) for _ in range(n_iter): model = torch.exp((A @ theta).clamp(min=-20.0, max=20.0)) J = model.unsqueeze(1) * A @@ -396,10 +407,17 @@ def hkl_symops_to_cartesian( ------- sym_mats_cart : torch.Tensor, shape (n_ops, 3, 3), real """ - dtype = torch.float64 # dtype-ok: 3x3 rotation algebra in double on the host - M = rec_basis.to(dtype).transpose(-1, -2) # (3, 3) + # Whatever width the caller brought. Double when it hands us double -- the + # peak finder's host-side 3x3 algebra does -- and the working precision + # otherwise, so this is usable on a backend without float64. The integer + # symops carry no width of their own, hence the fallback. + dtype = rec_basis.dtype if rec_basis.is_floating_point() else get_float_dtype() + M = rec_basis.to(dtype).transpose(-1, -2).contiguous() # (3, 3) + # `.contiguous()` is not cosmetic on a 3x3: MPS's `linalg.inv` trips an + # internal contiguity assert on the transposed view (torch 2.9.1), and the + # copy costs nine elements. M_inv = torch.linalg.inv(M) - S = sg_mats.to(dtype) # (n_ops, 3, 3) + S = sg_mats.to(device=M.device, dtype=dtype) # (n_ops, 3, 3) # S^T, not S: reciprocal space transforms as h' = h.S, so the operator # acting on Cartesian s as a column vector is (B^-1 S B)^T = M S^T M^-1 # with M = B^T. Using S here returns matrices that are not rotations at all @@ -469,8 +487,11 @@ def apply_overall_anisotropy( """ device = F.device dtype = F.dtype - s = s_vectors.to(device).to(dtype) - U_t = U.to(device).to(dtype) + # One `.to` per tensor, not two: `.to(device).to(dtype)` materialises the + # source width on the target device first, which throws on a backend that + # has no float64 -- and a host-side double U is a thing this is handed. + s = s_vectors.to(device=device, dtype=dtype) + U_t = U.to(device=device, dtype=dtype) s_dot_U = s @ U_t # (N, 3) arg = (torch.pi ** 2) * (s_dot_U * s).sum(dim=-1) return F * torch.exp(arg.clamp(min=-10.0, max=10.0)) diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index 3430e287..97913973 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -139,11 +139,16 @@ def build( """ dev = get_default_device() if device is None else device real = get_float_dtype() - F = F_obs.detach().to(dev) - F = (F.abs() if F.is_complex() else F).to(real) + # Cast and move in one `.to`, and take `abs()` before either. Chaining + # them the other way round -- `.to(dev)` and then `.to(real)` -- puts + # the caller's width on the device first, which throws for a float64 + # input on a backend that has none, and double observations are a + # perfectly ordinary thing to be handed. + F = F_obs.detach() + F = (F.abs() if F.is_complex() else F).to(device=dev, dtype=real) hkl_i = hkl.detach().to(dev) - rec_basis = real_cell.reciprocal_basis_matrix.to(dev).to(real) + rec_basis = real_cell.reciprocal_basis_matrix.to(device=dev, dtype=real) s_mag = (hkl_i.to(real) @ rec_basis).norm(dim=-1) hkl_l = hkl_i.round().to(torch.int64) # dtype-ok: Miller indices are integers @@ -162,7 +167,7 @@ def build( if sig_F is None: weight = torch.ones_like(F) else: - sig = sig_F.detach().to(dev).to(real).abs() + sig = sig_F.detach().to(device=dev, dtype=real).abs() weight = normalise_weight(inverse_variance_weight( snr_from_amplitude(F, sig), sigma_a, eps=eps, )) @@ -200,7 +205,10 @@ class CandidateTransform: def e_calc(self, t: torch.Tensor) -> torch.Tensor: """``E_calc(h, t)`` for ``t`` of shape ``(3,)`` or ``(K, 3)``: ``(N,)`` or ``(K, N)``.""" single = t.ndim == 1 - tt = t.reshape(-1, 3).to(self.h_R.device).to(self.h_R.dtype) + # One `.to`: a translation handed in as double -- a fractional vector + # from host-side algebra usually is -- must not be put on the device at + # its own width first. + tt = t.reshape(-1, 3).to(device=self.h_R.device, dtype=self.h_R.dtype) phase_arg = torch.einsum("ind,kd->kin", self.h_R, tt) phase = torch.exp((2j * math.pi) * phase_arg.to(self.G.dtype)) E = (self.G.unsqueeze(0) * phase).sum(dim=1).abs() @@ -232,9 +240,9 @@ def prepare_candidate( real = get_float_dtype() cplx = get_complex_dtype() - hkl = obs.hkl.to(device).to(real) - sym_R = spacegroup.matrices.detach().to(device).to(real) - sym_t = spacegroup.translations.detach().to(device).to(real) + hkl = obs.hkl.to(device=device, dtype=real) + sym_R = spacegroup.matrices.detach().to(device=device, dtype=real) + sym_t = spacegroup.translations.detach().to(device=device, dtype=real) S = int(sym_R.shape[0]) N = int(hkl.shape[0]) @@ -247,13 +255,13 @@ def prepare_candidate( G_raw = F_all * phase I_P = (G_raw.abs() ** 2).mean(dim=0).to(real) - s_mag = obs.s_mag.to(device).to(real) + s_mag = obs.s_mag.to(device=device, dtype=real) fit_P = WilsonNormaliser( I_P, s_mag, n_coeff=WILSON_N_COEFF, s_lo=float(s_mag.min()), s_hi=float(s_mag.max()), ) Sigma_c = S * fit_P.evaluate(s_mag).to(real) - norm = (obs.eps.to(device).to(real) * Sigma_c).clamp(min=1e-30).sqrt() + norm = (obs.eps.to(device=device, dtype=real) * Sigma_c).clamp(min=1e-30).sqrt() return CandidateTransform(G=G_raw / norm.to(cplx), h_R=h_R, norm=norm) @@ -384,9 +392,9 @@ def fast_translation_function( cplx = get_complex_dtype() nx, ny, nz = _grid_sizes(real_cell, grid_spacing_A) - G = cand.G.to(device).to(cplx) + G = cand.G.to(device=device, dtype=cplx) S, N = G.shape - coeff = obs.coeff.to(device).to(cplx) + coeff = obs.coeff.to(device=device, dtype=cplx) h_R_int = cand.h_R.round().to(torch.int64) # dtype-ok: Miller indices are integers # The pair (j, i) is the conjugate of (i, j) at -dh, so the map is twice @@ -398,7 +406,7 @@ def fast_translation_function( dh = h_R_int[i + 1:] - h_R_int[i:i + 1] # (S-i-1, N, 3) flat = ((dh[..., 0] % nx) * ny + (dh[..., 1] % ny)) * nz + (dh[..., 2] % nz) W.index_add_(0, flat.reshape(-1), (coeff.view(1, -1) * pair).reshape(-1)) - diag = (obs.coeff.to(device).to(real) * (G.abs() ** 2).sum(dim=0).to(real)).sum() + diag = (obs.coeff.to(device=device, dtype=real) * (G.abs() ** 2).sum(dim=0).to(real)).sum() score = (2.0 * torch.fft.ifftn(W.view(nx, ny, nz), dim=(0, 1, 2)).real * float(nx * ny * nz)).to(real) + diag @@ -435,8 +443,8 @@ def llg_at_translations( E_calc = cand.e_calc(t_candidates) # (K, N) K, N = E_calc.shape dev, real = E_calc.device, E_calc.dtype - E_obs = obs.E_obs.to(dev).to(real).view(1, N).expand(K, N) - D = obs.sigma_a.to(dev).to(real).view(1, N) + E_obs = obs.E_obs.to(device=dev, dtype=real).view(1, N).expand(K, N) + D = obs.sigma_a.to(device=dev, dtype=real).view(1, N) Sigma = (1.0 - D * D).clamp(min=1e-3).expand(K, N) cent = obs.centric.to(dev).view(1, N).expand(K, N) ll = -rice_per_refl(E_obs, D * E_calc, Sigma, cent) # (K, N) diff --git a/torchref/refinement/model_error_estimation/sigma_m.py b/torchref/refinement/model_error_estimation/sigma_m.py index a5817dbc..20bd6300 100644 --- a/torchref/refinement/model_error_estimation/sigma_m.py +++ b/torchref/refinement/model_error_estimation/sigma_m.py @@ -133,7 +133,7 @@ def prepare( s_half_sq = s_half_sq.to(device=device, dtype=dtype) s_sq = 4.0 * s_half_sq - valid_f = validity.to(torch.bool).to(device).to(dtype) + valid_f = validity.to(torch.bool).to(device=device, dtype=dtype) n_valid = valid_f.sum().clamp(min=1.0) sigma_obs = sigma_obs.to(device=device, dtype=dtype) self.sigma_d_mean = (sigma_obs * valid_f).sum() / n_valid diff --git a/torchref/scaling/wilson.py b/torchref/scaling/wilson.py index e0be8c64..e16b048d 100644 --- a/torchref/scaling/wilson.py +++ b/torchref/scaling/wilson.py @@ -335,13 +335,37 @@ def objective(b): A = A + torch.eye(self.n_coeff, dtype=A.dtype, device=A.device) * ( 1e-10 * float(torch.diagonal(A).abs().max().clamp(min=1e-30)) ) - lu = torch.linalg.lu_factor(A) + # Factorised once, and by Cholesky rather than LU. `A` is `X^T W X` + # plus a ridge with positive IRLS weights, so it is symmetric positive + # definite by construction -- and MPS implements neither `lu_solve` nor + # `cholesky_solve` (torch 2.9.1), which left the per-iteration solve to + # a CPU round trip: 1015 us against 114 us for two triangular solves, on + # a 200k x 6 problem whose unavoidable `XtW @ z` is 966 us. Same + # arithmetic -- over the 16 datasets in ``tests/files/mtz`` the two + # agree to 1e-5 relative in float32 and 1e-14 in float64, with identical + # iteration counts on every one. + # + # `cholesky_ex` reports rather than raises, because a fully collinear + # basis is a thing this fit sees: the high-order Chebyshev columns go + # near-singular when the data cover only part of the basis range, and + # the ridge does not always rescue that. LU carried no definiteness + # requirement, so that case falls back to a general solve instead of + # failing. + chol = torch.linalg.cholesky_ex(A) + L_A = chol.L if int(chol.info) == 0 else None + + def _solve(rhs): + if L_A is None: + return torch.linalg.solve(A, rhs) + return torch.linalg.solve_triangular( + L_A.mT, + torch.linalg.solve_triangular(L_A, rhs, upper=False), + upper=True, + ) for it in range(1, max_iter + 1): z = eta + (y - mu) / mu # working response - step = torch.linalg.lu_solve( - *lu, (XtW @ z).unsqueeze(-1), - ).squeeze(-1) - beta + step = _solve((XtW @ z).unsqueeze(-1)).squeeze(-1) - beta if not torch.isfinite(step).all(): raise RuntimeError( f"Wilson fit diverged at iteration {it}: the IRLS solve " From 18756ba105a6306b138ff22bcbdc84582420760f Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 3 Sep 2026 09:30:58 +0200 Subject: [PATCH 156/250] Write our own refinement header instead of inheriting the input's Refinement output used to copy the input file's header verbatim and then append its own REMARK 3, so a refined 3GR5 carried 420 header lines asserting two refinements at once: REFMAC 5.1.24 with R-work 0.213 at line 5, ours at line 389. A reader taking the first REMARK 3 it found got REFMAC. The inherited block also contradicted the data beside it -- it claimed a 5.1% / 1072-reflection test set while the MTZ shipped with it holds 9.85% / 2063, which torchref reads correctly. The passthrough was inverted as well, keeping the statistics refinement invalidates and dropping the chemistry it does not: SEQRES, SSBOND, DBREF, EXPDTA, COMPND, SOURCE, KEYWDS, SEQADV, HETNAM, FORMUL and SITE were all absent from the output. A whitelist now carries the crystal, sample and chemistry records through in mandated record order (TITLE used to land after REMARK 900), REMARK 2, 3 and 500 are dropped, and AUTHOR and JRNL are not inherited because they credit the deposition rather than this run. 283 header lines for the same file, 41 structural records preserved, one refinement block. add-metadata is exempt through supersede_refinement=False: annotating a file is not re-refining it, so nothing there supersedes the existing records. Prior refinements are tracked through mmCIF's _software loop, the only place either format has room for them -- _refine is singular by design. pdbx_ordinal was hardcoded to 1 and the incoming loop was never read, truncating the chain to one link on every write; it now reads the input's loop and appends at max(ordinal) + 1, carrying each entry's description so the chain says what every program did. Added _pdbx_initial_refinement_model, _refine.pdbx_starting_model and _refine.pdbx_R_Free_selection_details. from_cif_file no longer carries the input's _refine items through, which was the mmCIF form of the duplicated REMARK 3. Also fixed, all of the same kind: - mmCIF loop cells were written unquoted (only the pair path quoted), so any value containing whitespace split into extra columns when the file was read back. Latent while every loop column was a single token; a multi-word _software.description exposed it. - PDB coordinate and B-factor columns were written as str() of a rounded float rather than with an explicit precision, dropping trailing zeros: 18.3 and 31.51 where the format wants 18.300 and 31.510 (369 of 1329 atoms in a 3GR5 refinement), and 95.4 where it wants 95.40. Both writers affected. Columns held and values round-trip through any float()-based reader, so nothing was numerically wrong. - rfree_source was set only when the MTZ also carried a validation column, leaving the common FreeR-only case unattributed, and generated flags did not record their seed. It is also named after the reader now, since load() takes ReflectionCIFReader as well as MTZReader. The header records what was refined and how: refinement_method was only ever set by the difference-refinement CLI, so ordinary runs emitted no method line at all. TARGET and OPTIMIZER come from the settings that already reach refinement_history.json, including the scale target because it changes the R-factors the same header reports. Long values wrap onto a continuation line with the colon held in column 25, the convention REFMAC uses for its own author list. Author-supplied text goes in --output-remarks, rendered as REMARK 3 OTHER REFINEMENT REMARKS and _refine.details and emitted only when set. Nothing in the block is editorialised: every generated line is a measured quantity or a recorded setting. Removed the deprecated template= argument to pdb.write (nothing called it) and the never-populated custom_remarks field. RefinementMetadata had no test coverage at all, which is how the duplicated REMARK 3 shipped; 39 tests now cover the contract. Co-Authored-By: Claude Opus 5 (1M context) --- docs/changelog.rst | 6 + tests/unit/io/test_anomalous_reader.py | 6 +- tests/unit/io/test_refinement_header.py | 440 +++++++++++++++++++++ torchref/cli/_common.py | 47 +++ torchref/cli/add_metadata.py | 7 +- torchref/io/cif.py | 48 ++- torchref/io/datasets/reflection_data.py | 17 +- torchref/io/metadata.py | 495 +++++++++++++++++++----- torchref/io/pdb.py | 21 +- 9 files changed, 955 insertions(+), 132 deletions(-) create mode 100644 tests/unit/io/test_refinement_header.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 3a5a2f45..31b86182 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,12 @@ Changelog Unreleased ---------- +- Refinement output no longer inherits the input file's refinement header. It used to copy the whole thing and then append its own ``REMARK 3``, so a refined 3GR5 carried 420 header lines asserting two refinements at once -- ``PROGRAM : REFMAC 5.1.24`` with R-work 0.213 at line 5, ours at line 389 -- and a reader taking the first ``REMARK 3`` got REFMAC. The inherited block was not merely stale but contradicted the data beside it: it claimed a 5.1% / 1072-reflection test set, while the MTZ shipped with it holds 9.85% / 2063 (which torchref reads correctly). The passthrough was also inverted, keeping the statistics refinement invalidates and dropping the chemistry it does not -- SEQRES, SSBOND, DBREF, EXPDTA, COMPND, SOURCE, KEYWDS, SEQADV, HETNAM, FORMUL and SITE were all absent from the output. Now a whitelist carries the crystal, sample and chemistry records through in mandated record order (TITLE used to be emitted after REMARK 900), REMARK 2, 3 and 500 are dropped, and AUTHOR and JRNL are not inherited because they credit the deposition rather than this run. 283 header lines for the same file, 41 structural records preserved, one refinement block. ``add-metadata`` is exempt through ``supersede_refinement=False``: annotating a file is not re-refining it, so nothing there supersedes the existing REMARK 3 or AUTHOR records and both are kept +- Prior refinements are tracked through mmCIF's ``_software`` loop, which is the only place either format has room for them: ``_refine`` is singular by design, so a previous program's statistics cannot be kept without contradicting the current ones. ``pdbx_ordinal`` was hardcoded to ``1`` and the incoming loop was never read, truncating the chain to one link on every write; it now reads the input's loop and appends at ``max(ordinal) + 1``, carrying each entry's ``description`` so the chain says what every program did and not just that it ran. Added ``_pdbx_initial_refinement_model`` and ``_refine.pdbx_starting_model``, which name what the refinement started from, and ``_refine.pdbx_R_Free_selection_details``, which names the test set the reported R-free conditions on. ``from_cif_file`` no longer carries the input's ``_refine`` items through -- that was the mmCIF form of the duplicated ``REMARK 3`` +- Fixed mmCIF loop cells being written unquoted, which silently split any value containing whitespace into extra columns when the file was read back. Latent while every loop column was a single token; a multi-word ``_software.description`` exposed it. Nulls stay bare, since ``gemmi.cif.quote`` turns ``?`` into the quoted one-character string ``'?'`` +- Fixed PDB coordinate and B-factor columns being written as ``str()`` of a rounded float rather than with an explicit precision, so trailing zeros were dropped -- ``18.3`` and ``31.51`` where the format wants ``18.300`` and ``31.510`` (369 of 1329 atoms in a 3GR5 refinement), and ``95.4`` where it wants ``95.40``. Both writers were affected. Columns held and the values round-trip through any ``float()``-based reader, so nothing was numerically wrong; this is format conformance +- The header records what was refined and how. ``refinement_method`` was only ever set by the difference-refinement CLI, so ordinary runs emitted no method line at all; ``TARGET`` and ``OPTIMIZER`` now come from the same settings that reach ``refinement_history.json`` -- target family and registry key, macrocycles, refinement mode, ADP model and the scale target, the last because it changes the R-factors the same header reports. Long values wrap onto a continuation line with the colon held in column 25, the convention REFMAC uses for its own author list. ``rfree_source`` was set only when the MTZ also carried a validation column, leaving the common FreeR-only case unattributed, and generated flags did not record their seed -- both fixed, so the free-set provenance in the header is always either the file it came from or a reproducible draw +- Author-supplied header text goes in ``--output-remarks``, rendered as ``REMARK 3 OTHER REFINEMENT REMARKS`` and ``_refine.details`` and emitted only when set. Nothing in the block is editorialised: every generated line is a measured quantity or a recorded setting. Removed the deprecated ``template=`` argument to ``pdb.write`` (nothing called it) and the never-populated ``custom_remarks`` field. ``RefinementMetadata`` had no test coverage at all, which is how the duplicated ``REMARK 3`` shipped; 36 tests now cover the contract - Fixed the rotation function contracting **unconjugated** calc coefficients on MPS. ``torch.conj`` returns a lazy view carrying a conjugate *bit*, and MPS's batched complex matmul -- what the radial ``einsum`` lowers to -- ignores it, so the contraction was silently wrong: on 1DAW by 173% of ``|xi|`` max, which reordered the entire peak list and pushed the true orientation from rank 0 out of the top 200 while the top score moved only 0.1%. Materialised with ``resolve_conj()`` at the two batched-matmul sites. Elementwise ops, ``where``, ``index_add``, ``fft`` and 2-D matmul all honour the bit and only the batched path does not, which is too narrow to guard structurally, so the guard is the contraction's own value against a host-double reference - The alignment package runs end to end on an accelerator without float64: 123 alignment tests including the slow rotation searches pass on MPS, where every one of them failed before. Five things stood in the way, four of them the same mistake -- casting and moving in two steps. ``.to(device).to(dtype)`` puts the caller's width on the device first, which throws for a double input on a backend that has none (14 sites, now one fused ``.to`` each); ``detect_zsymm`` widened the symmetry operators to double while they sat on the device; the expansion's host-side clustering keys left a host index to meet device values; and MPS's ``linalg.inv`` trips an internal contiguity assert on a transposed 3x3 view. The fifth is a kernel gap -- MPS implements no complex cumulative op -- so the azimuthal phase ladder is built by doubling instead of by ``cumprod``, measured at 5.7e-6 against ``cumprod``'s 4.1e-6 in complex64 over 5000 angles at L=101 - The anisotropy fit runs at the configured float dtype on the data's own device, instead of double on the host. Measured over the 16 datasets in ``tests/files/mtz``, float32 reproduces the double fit to 3.3e-5 relative in ``U`` and 4.5e-6 in the correction factor ``exp(+pi^2 s.U.s)`` it exists to produce -- the design matrix is a constant column beside ``2 pi^2 s.s`` terms of order 0.1-1, so there is no precision cliff for seven parameters to fall off. Unrotated searches on 1DAW, 3K7M, 2DQ6 and 4BX9 return the same peaks in the same order as before, with scores moved 1.8e-7 to 2.5e-4 relative. ``fit_anisotropy`` takes an optional ``device``; ``get_axis_order`` and ``hkl_symops_to_cartesian`` follow their inputs' width and place rather than forcing double diff --git a/tests/unit/io/test_anomalous_reader.py b/tests/unit/io/test_anomalous_reader.py index 90641680..d72e0f8e 100644 --- a/tests/unit/io/test_anomalous_reader.py +++ b/tests/unit/io/test_anomalous_reader.py @@ -127,7 +127,11 @@ def test_generated_rfree_shared_across_mates(anomalous_two_column_mtz): assert bool(d.friedel_flags.any()) d.regenerate_rfree_flags(force=True, seed=0) - assert d.rfree_source == "Generated (resolution-binned, ASU-grouped)" + # The seed is part of the provenance: "generated" without it names a draw + # nobody can reproduce. + assert d.rfree_source == ( + "Generated (resolution-binned, ASU-grouped, seed 0)" + ) assert bool((d.rfree_flags == 0).any()) # a free set actually exists assert _mixed_partition_groups(d) == [] diff --git a/tests/unit/io/test_refinement_header.py b/tests/unit/io/test_refinement_header.py new file mode 100644 index 00000000..9da010e0 --- /dev/null +++ b/tests/unit/io/test_refinement_header.py @@ -0,0 +1,440 @@ +"""Regression tests for the refinement output header. + +The writer used to carry an input PDB's header through verbatim and then append +its own ``REMARK 3``, so a refined file asserted two different refinements at +once -- the previous program's R-factors first, ours several hundred lines +later. The selection was also inverted: TITLE, AUTHOR and REMARK were kept +(including the superseded statistics) while SEQRES, SSBOND, LINK, CISPEP, +EXPDTA, COMPND, DBREF and FORMUL were dropped, i.e. the chemistry refinement +does not invalidate. + +What is asserted here is the resulting contract: exactly one refinement block +and it is ours; records that describe the crystal survive; records that describe +a superseded refinement do not; and the chain of programs applied to the model +is preserved through mmCIF's ``_software`` loop rather than by hoarding the +previous program's output. +""" + +import pandas as pd +import pytest + +from torchref.io import cif, pdb +from torchref.io.metadata import RefinementMetadata + +# 3GR5 was refined with REFMAC 5.1.24 and carries a full deposition header: +# 420 lines including REMARK 2/3/500, JRNL, AUTHOR, SEQRES, SSBOND and SITE. +INPUT_PDB = "tests/files/pdb/3GR5.pdb" + +#: PDB record order, abridged to the records this writer can emit. The format +#: mandates this sequence; TITLE used to be written *after* REMARK 900. +RECORD_ORDER = [ + "TITLE", "COMPND", "SOURCE", "KEYWDS", "EXPDTA", "MDLTYP", "AUTHOR", + "REMARK", "DBREF", "DBREF1", "DBREF2", "SEQADV", "SEQRES", "MODRES", + "HET", "HETNAM", "HETSYN", "FORMUL", "SSBOND", "LINK", "CISPEP", "SITE", +] + + +def _atom_df(): + """Two atoms whose coordinates exercise trailing-zero formatting. + + ``18.3`` and ``31.51`` are the values that exposed the writer using + ``str()`` semantics on a rounded float instead of ``%8.3f``. + """ + df = pd.DataFrame( + { + "ATOM": ["ATOM", "HETATM"], + "serial": [1, 2], + "name": ["CA", "O"], + "altloc": ["", ""], + "resname": ["LEU", "HOH"], + "chainid": ["A", "A"], + "resseq": [1, 2], + "icode": ["", ""], + "x": [-7.223, -9.22], + "y": [29.982, 31.51], + "z": [18.3, 20.364], + "occupancy": [1.0, 1.0], + "tempfactor": [95.4, 30.0], + "element": ["C", "O"], + "charge": [0, 0], + "anisou_flag": [False, False], + "u11": [0.0, 0.0], + "u22": [0.0, 0.0], + "u33": [0.0, 0.0], + "u12": [0.0, 0.0], + "u13": [0.0, 0.0], + "u23": [0.0, 0.0], + } + ) + df.attrs["cell"] = [90.645, 90.645, 133.422, 90.0, 90.0, 120.0] + df.attrs["spacegroup"] = "P 65 2 2" + return df + + +def _refined_metadata(): + """Input header plus this refinement's statistics, as the writer builds it.""" + meta = RefinementMetadata.from_pdb_file(INPUT_PDB) + ours = RefinementMetadata( + program_version="0.7.0", + target_function="ML", + optimizer="1 MACROCYCLE, SEPARATE, ISOTROPIC ADP, SCALE TARGET NLL", + r_work=0.2083, + r_free=0.2471, + percent_free=9.9, + n_reflections_all=20942, + n_reflections_test=2063, + resolution_high=2.05, + resolution_low=18.69, + starting_model=INPUT_PDB, + rfree_selection="MTZReader FreeR", + ) + return meta.merge(ours) + + +def _header_lines(): + return _refined_metadata().render_pdb_header().splitlines() + + +def _remark_number(line): + try: + return int(line[7:10]) + except ValueError: + return None + + +# ====================================================================== # +# One refinement block, and it is ours +# ====================================================================== # + + +@pytest.mark.unit +def test_exactly_one_refinement_block(): + header = "\n".join(_header_lines()) + assert header.count("REMARK 3 REFINEMENT.") == 1 + assert "TORCHREF" in header + + +@pytest.mark.unit +@pytest.mark.parametrize("program", ["REFMAC", "PHENIX", "BUSTER", "CNS"]) +def test_no_foreign_program_is_credited(program): + """The input was refined by REFMAC; the output must not say so anywhere.""" + assert program not in "\n".join(_header_lines()) + + +@pytest.mark.unit +def test_superseded_statistics_are_dropped(): + """REMARK 2, 3 and 500 describe a resolution and a model we replaced.""" + numbers = {_remark_number(line) for line in _header_lines() + if line.startswith("REMARK")} + assert 2 not in numbers # resolution: regenerated + assert 500 not in numbers # geometry outliers of old coordinates + assert 3 in numbers # ours, and the only one + + +@pytest.mark.unit +def test_input_remark_3_is_not_carried_through(): + """The specific inversion this module exists to prevent.""" + meta = RefinementMetadata.from_pdb_file(INPUT_PDB) + assert not any(_remark_number(r) == 3 for r in meta.passthrough_pdb_remarks) + # ... while the input genuinely has one, so the assertion has teeth. + with open(INPUT_PDB) as handle: + assert any(line.startswith("REMARK 3") for line in handle) + + +# ====================================================================== # +# Attribution +# ====================================================================== # + + +@pytest.mark.unit +def test_authors_and_journal_are_not_inherited(): + """Both credit the deposition, not this refinement.""" + meta = RefinementMetadata.from_pdb_file(INPUT_PDB) + assert meta.authors == [] + lines = _header_lines() + assert not any(line.startswith(("AUTHOR", "JRNL")) for line in lines) + + +@pytest.mark.unit +def test_explicitly_set_authors_are_written(): + """Not inheriting authors must not stop --authors from working.""" + meta = _refined_metadata() + meta.authors = ["A.Person", "B.Other"] + lines = meta.render_pdb_header().splitlines() + author = [line for line in lines if line.startswith("AUTHOR")] + assert author and "A.Person" in author[0] + + +# ====================================================================== # +# Records that survive +# ====================================================================== # + + +@pytest.mark.unit +@pytest.mark.parametrize( + "record", ["COMPND", "SOURCE", "KEYWDS", "EXPDTA", "DBREF", "SEQADV", + "SEQRES", "HETNAM", "FORMUL", "SSBOND", "SITE"] +) +def test_structural_records_are_carried_through(record): + """Refinement moves atoms; it does not change the sequence or chemistry.""" + with open(INPUT_PDB) as handle: + expected = sum(1 for line in handle if line.startswith(record)) + assert expected > 0, f"{record} absent from the fixture" + written = sum(1 for line in _header_lines() if line.startswith(record)) + assert written == expected + + +@pytest.mark.unit +def test_secondary_structure_is_not_emitted(): + """Nothing here computes HELIX/SHEET, so carrying them can only mislead.""" + lines = _header_lines() + assert not any(line.startswith(("HELIX", "SHEET")) for line in lines) + + +@pytest.mark.unit +def test_records_are_in_mandated_order(): + seen = [] + for line in _header_lines(): + record = line[:6].strip() + if record and (not seen or seen[-1] != record): + seen.append(record) + assert all(r in RECORD_ORDER for r in seen), [ + r for r in seen if r not in RECORD_ORDER + ] + positions = [RECORD_ORDER.index(r) for r in seen] + assert positions == sorted(positions), seen + + +@pytest.mark.unit +def test_remarks_ascend_with_ours_at_three(): + numbers = [ + _remark_number(line) for line in _header_lines() + if line.startswith("REMARK") + ] + assert numbers == sorted(numbers) + assert 3 in numbers + + +# ====================================================================== # +# Free text is the author's, never generated +# ====================================================================== # + + +@pytest.mark.unit +def test_no_remarks_section_without_author_text(): + assert "OTHER REFINEMENT REMARKS" not in "\n".join(_header_lines()) + + +@pytest.mark.unit +def test_author_text_is_written_when_supplied(): + meta = _refined_metadata() + meta.output_remarks = "Re-refined for the benchmark.\n\nSecond paragraph." + header = meta.render_pdb_header() + assert "REMARK 3 OTHER REFINEMENT REMARKS:" in header + assert "Re-refined for the benchmark." in header + assert "Second paragraph." in header + + +@pytest.mark.unit +def test_free_set_provenance_is_reported(): + """R-free from a different test set is not comparable; say which one.""" + header = "\n".join(_header_lines()) + assert "FREE R VALUE TEST SET SELECTION" in header + assert "MTZReader FreeR" in header + + +@pytest.mark.unit +def test_remark_3_colons_line_up_within_each_block(): + """Colons used to be ragged, because alignment relied on caller padding. + + Two columns exist by design, which is what REFMAC does too: a narrow one + for the identification lines (``PROGRAM :``) and a wide one for the + statistics (``RESOLUTION RANGE HIGH (ANGSTROMS) :``). Each must be + internally consistent, including on wrapped continuation lines. + """ + columns = [ + line.index(" : ") for line in _header_lines() + if line.startswith("REMARK 3 ") and " : " in line + ] + assert set(columns) == {24, 46}, sorted(set(columns)) + + +@pytest.mark.unit +def test_header_lines_fit_the_format(): + assert [line for line in _header_lines() if len(line) > 80] == [] + + +@pytest.mark.unit +def test_long_identification_values_wrap_rather_than_overflow(): + """Naming cycles, mode, ADP model and scale target overruns column 80.""" + meta = _refined_metadata() + meta.optimizer = ( + "12 MACROCYCLES, EVERYTHING, FIELD_ANISO ADP, SCALE TARGET ML_NOALPHA, " + "RIGID BODY 5 ITERATIONS" + ) + lines = meta.render_pdb_header().splitlines() + assert [line for line in lines if len(line) > 80] == [] + optimizer = [line for line in lines if "OPTIMIZER" in line] + assert len(optimizer) == 1 # one head line ... + head = lines.index(optimizer[0]) + assert lines[head + 1].startswith("REMARK 3 : ") + # ... and the value survives the wrap intact. + joined = " ".join( + line.split(":", 1)[1].strip() + for line in lines[head:head + 3] + if " : " in line + ) + assert "RIGID BODY 5 ITERATIONS" in joined + + +# ====================================================================== # +# Coordinates +# ====================================================================== # + + +@pytest.mark.unit +def test_numeric_columns_keep_their_decimals(tmp_path): + """``18.3`` must be written `` 18.300`` and ``95.4`` as `` 95.40``. + + Both fields were formatted as ``str()`` of a rounded float rather than with + an explicit precision, so trailing zeros vanished. + """ + out = tmp_path / "out.pdb" + pdb.write(_atom_df(), str(out)) + atoms = [ + line for line in out.read_text().splitlines() + if line.startswith(("ATOM", "HETATM")) + ] + assert atoms + for line in atoms: + assert len(line) == 80, repr(line) + # x, y, z: %8.3f + for start in (30, 38, 46): + field = line[start:start + 8] + assert len(field.split(".")[1]) == 3, repr(field) + # occupancy and B: %6.2f + for start in (54, 60): + field = line[start:start + 6] + assert len(field.split(".")[1]) == 2, repr(field) + assert " 95.40" in atoms[0] + + +# ====================================================================== # +# The software chain -- mmCIF's only room for prior work +# ====================================================================== # + + +@pytest.mark.unit +def test_software_is_a_loop_with_an_ordinal(): + cats = _refined_metadata().render_cif_categories() + software = cats["_software"] + assert software["_software.pdbx_ordinal"] == ["1"] + assert software["_software.name"] == ["TORCHREF"] + + +@pytest.mark.unit +def test_software_ordinal_increments_across_refinements(tmp_path): + """Refining a refined file must append to the chain, not replace it.""" + first = tmp_path / "first.cif" + cif.write_model(_atom_df(), str(first), metadata=_refined_metadata()) + + # Read it back the way a second refinement would, then write again. + carried = RefinementMetadata.from_cif_file(str(first)) + assert len(carried.software_chain) == 1 + + second_meta = carried.merge( + RefinementMetadata(program_version="0.7.0", optimizer="2 MACROCYCLES") + ) + second = tmp_path / "second.cif" + cif.write_model(_atom_df(), str(second), metadata=second_meta) + + chain = RefinementMetadata.from_cif_file(str(second)).software_chain + assert [e.get("pdbx_ordinal") for e in chain] == ["1", "2"] + + +@pytest.mark.unit +def test_previous_refinement_description_survives_the_round_trip(tmp_path): + """The chain is worthless if it cannot say what each program did.""" + out = tmp_path / "one.cif" + meta = _refined_metadata() + meta.optimizer = "7 MACROCYCLES, SEPARATE, ISOTROPIC ADP" + cif.write_model(_atom_df(), str(out), metadata=meta) + + chain = RefinementMetadata.from_cif_file(str(out)).software_chain + assert "7 MACROCYCLES" in chain[0]["description"] + + +@pytest.mark.unit +def test_loop_values_with_spaces_survive_a_round_trip(tmp_path): + """An unquoted loop cell silently splits into extra columns when re-read.""" + out = tmp_path / "quoted.cif" + meta = _refined_metadata() + meta.optimizer = "MANY WORDS, WITH COMMAS, AND SPACES" + cif.write_model(_atom_df(), str(out), metadata=meta) + + chain = RefinementMetadata.from_cif_file(str(out)).software_chain + assert len(chain) == 1 + assert chain[0]["description"].endswith("AND SPACES") + + +@pytest.mark.unit +def test_input_refine_statistics_are_not_carried_into_cif(tmp_path): + """The mmCIF form of the duplicated REMARK 3.""" + source = tmp_path / "prior.cif" + prior = RefinementMetadata(program="REFMAC", r_work=0.213, r_free=0.251) + cif.write_model(_atom_df(), str(source), metadata=prior) + + carried = RefinementMetadata.from_cif_file(str(source)) + assert "_refine" not in carried.passthrough_cif_categories + assert carried.r_work is None + # The prior program is remembered as a link in the chain, not as statistics. + assert [e["name"] for e in carried.software_chain] == ["REFMAC"] + + +@pytest.mark.unit +def test_starting_model_is_recorded_as_an_accession(tmp_path): + cats = _refined_metadata().render_cif_categories() + initial = cats["_pdbx_initial_refinement_model"] + assert initial["_pdbx_initial_refinement_model.accession_code"] == "3GR5" + assert initial["_pdbx_initial_refinement_model.type"] == "experimental model" + assert cats["_refine"]["_refine.pdbx_starting_model"] == INPUT_PDB + + +@pytest.mark.unit +def test_unaccessionable_starting_model_is_named_not_guessed(): + meta = _refined_metadata() + meta.starting_model = "/tmp/my_working_model_v3.pdb" + initial = meta.render_cif_categories()["_pdbx_initial_refinement_model"] + assert "accession_code" not in " ".join(initial) + assert initial["_pdbx_initial_refinement_model.details"] == ( + "my_working_model_v3.pdb" + ) + + +# ====================================================================== # +# Annotating a file is not refining it +# ====================================================================== # + + +@pytest.mark.unit +def test_annotation_keeps_the_existing_refinement(): + """``add-metadata`` adds a title; it does not re-refine. + + Nothing supersedes the input's REMARK 3 or AUTHOR records in that case, so + dropping them would discard statistics and credit that are still accurate. + """ + meta = RefinementMetadata.from_pdb_file( + INPUT_PDB, supersede_refinement=False + ) + numbers = {_remark_number(r) for r in meta.passthrough_pdb_remarks} + assert {2, 3, 500} <= numbers + assert meta.authors # REFMAC-era depositors keep their credit + + +@pytest.mark.unit +def test_refinement_output_supersedes_by_default(): + """The default is the refinement case: the old block goes.""" + meta = RefinementMetadata.from_pdb_file(INPUT_PDB) + numbers = {_remark_number(r) for r in meta.passthrough_pdb_remarks} + assert not ({2, 3, 500} & numbers) + assert meta.authors == [] diff --git a/torchref/cli/_common.py b/torchref/cli/_common.py index a9052312..da543742 100644 --- a/torchref/cli/_common.py +++ b/torchref/cli/_common.py @@ -392,6 +392,14 @@ def add_metadata_args(parser: argparse.ArgumentParser) -> None: default=None, help="Author names for the output file header", ) + parser.add_argument( + "--output-remarks", + type=str, + default=None, + help="Free-text note for the output header (REMARK 3 OTHER REFINEMENT " + "REMARKS / _refine.details). Nothing is written here unless you ask " + "for it", + ) parser.add_argument( "--no-header", action="store_true", @@ -771,6 +779,45 @@ def write_refinement_outputs( metadata.title = args.title if getattr(args, "authors", None): metadata.authors = args.authors + if getattr(args, "output_remarks", None): + metadata.output_remarks = args.output_remarks + + # What was minimised and how. These live on the CLI namespace rather + # than on the refinement, which is why from_refinement cannot fill them + # and why the header carried no method line at all until now. + xray_mode = getattr(args, "xray_mode", None) + if xray_mode: + # Name the family as well as the registry key. "ML" alone is + # cryptic in a deposited header, and the family follows from the + # key's prefix, so there is no lookup table here to drift out of + # step with XRAY_TARGETS. + key = str(xray_mode).upper() + if key.startswith(("ML", "NLL")): + metadata.target_function = f"MAXIMUM LIKELIHOOD ({key})" + elif key.startswith("LS"): + metadata.target_function = f"LEAST SQUARES ({key})" + else: + metadata.target_function = key + optimizer_parts = [] + n_cycles = getattr(args, "n_cycles", None) + if n_cycles: + optimizer_parts.append( + f"{n_cycles} MACROCYCLE" + ("S" if n_cycles != 1 else "") + ) + mode = getattr(args, "mode", None) + if mode: + optimizer_parts.append(str(mode).upper()) + adp_mode = getattr(args, "adp_mode", None) + if adp_mode: + optimizer_parts.append(f"{str(adp_mode).upper()} ADP") + # The scale target changes the R-factors this very header reports, so a + # run cannot be attributed without it (see the note beside it in the + # refinement_history.json parameters block). + scale_target = getattr(args, "scale_target", None) + if scale_target: + optimizer_parts.append(f"SCALE TARGET {str(scale_target).upper()}") + if optimizer_parts: + metadata.optimizer = ", ".join(optimizer_parts) outputs = {"pdb": None, "cif": None} diff --git a/torchref/cli/add_metadata.py b/torchref/cli/add_metadata.py index 8bae256b..9aa86076 100644 --- a/torchref/cli/add_metadata.py +++ b/torchref/cli/add_metadata.py @@ -128,7 +128,12 @@ def main(): # Start with pass-through from input file input_suffix = input_path.suffix.lower() if input_suffix == ".pdb": - metadata = RefinementMetadata.from_pdb_file(str(input_path)) + # This tool annotates a file, it does not re-refine it -- so the input's + # REMARK 3 and AUTHOR records are not superseded by anything and are + # kept. Refinement output takes the default and drops them. + metadata = RefinementMetadata.from_pdb_file( + str(input_path), supersede_refinement=False + ) elif input_suffix in (".cif", ".mmcif"): metadata = RefinementMetadata.from_cif_file(str(input_path)) else: diff --git a/torchref/io/cif.py b/torchref/io/cif.py index 96bccc10..d616abcb 100644 --- a/torchref/io/cif.py +++ b/torchref/io/cif.py @@ -257,12 +257,33 @@ def dataframe_to_gemmi_structure(df, cell, spacegroup): return st -def _add_refine_categories(doc, metadata): - """Inject a :class:`RefinementMetadata`'s categories into ``doc``, in place.""" +def _cif_value(val) -> str: + """Render one value as a CIF token, quoting it when it needs quoting. + + The unset markers ``?`` and ``.`` are passed through bare: ``gemmi.cif.quote`` + would turn them into the quoted one-character strings ``'?'`` and ``'.'``, + which are data rather than nulls. Everything else goes through ``quote`` -- + an unquoted value containing whitespace silently splits into extra loop + columns when the file is read back. + """ import gemmi + text = str(val) + if text in ("?", "."): + return text + return gemmi.cif.quote(text) + + +def _add_refine_categories(doc, metadata): + """Inject a :class:`RefinementMetadata`'s categories into ``doc``, in place. + + Returns the set of category prefixes written (e.g. ``{"_refine.", + "_software."}``) so the caller can avoid copying the same categories in + again from another block and either clobbering or duplicating them. + """ block = doc.sole_block() cats = metadata.render_cif_categories() + written = set() for cat_name, items in cats.items(): # List values mean a loop category rather than key-value pairs. @@ -274,6 +295,7 @@ def _add_refine_categories(doc, metadata): prefix = tags[0].rsplit(".", 1)[0] + "." suffixes = [t.split(".")[-1] for t in tags] loop = block.init_loop(prefix, suffixes) + written.add(prefix) # All list values should have same length n_rows = max(len(v) for v in items.values() if isinstance(v, list)) for i in range(n_rows): @@ -281,13 +303,17 @@ def _add_refine_categories(doc, metadata): for tag in tags: val = items[tag] if isinstance(val, list): - row.append(str(val[i]) if i < len(val) else "?") + cell = val[i] if i < len(val) else "?" else: - row.append(str(val)) + cell = val + row.append(_cif_value(cell)) loop.add_row(row) else: for key, val in items.items(): - block.set_pair(key, gemmi.cif.quote(str(val))) + block.set_pair(key, _cif_value(val)) + written.add(key.rsplit(".", 1)[0] + ".") + + return written def write_model(df, filepath: str, metadata=None) -> None: @@ -329,7 +355,12 @@ def write_model(df, filepath: str, metadata=None) -> None: "_symmetry.space_group_name_H-M", gemmi.cif.quote(str(spacegroup)) ) - _add_refine_categories(meta_doc, metadata) + written = _add_refine_categories(meta_doc, metadata) + # Categories we just wrote from metadata, plus the two written above. + # The structure block gemmi builds from the DataFrame carries its own + # version of some of these; copying those in would clobber a pair or + # append a second loop for the same category. + written |= {"_cell.", "_symmetry."} st = dataframe_to_gemmi_structure(df, cell, spacegroup) struct_doc = st.make_mmcif_document() @@ -342,14 +373,15 @@ def write_model(df, filepath: str, metadata=None) -> None: tags = list(loop.tags) suffixes = [t.split(".")[-1] for t in tags] prefix = tags[0].rsplit(".", 1)[0] + "." + if prefix in written: + continue new_loop = meta_block.init_loop(prefix, suffixes) for row_idx in range(loop.length()): row = [loop[row_idx, col] for col in range(loop.width())] new_loop.add_row(row) elif item.pair is not None: tag, val = item.pair - # Skip cell/symmetry - already added - if not tag.startswith(("_cell.", "_symmetry.")): + if tag.rsplit(".", 1)[0] + "." not in written: meta_block.set_pair(tag, val) meta_doc.write_file(filepath) diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index 739a64b4..e4707a76 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -862,6 +862,13 @@ def load(self, reader, french_wilson: bool = True): rfree = rfree.clip(min=0, max=1).to(torch.bool) self.rfree_flags = rfree self.masks["flagged_initial"] = ~flagged + # Record the provenance for every file-sourced set, not only the + # ones that also carry a validation column: a header reporting + # R-free has to be able to say which test set produced it. Named + # after the reader rather than hardcoded "MTZ", since `load` also + # takes ReflectionCIFReader and any other compatible reader. + reader_name = type(reader).__name__ + self.rfree_source = f"{reader_name} FreeR" # A third (validation) column goes into the separate boolean # ``validation_flags``; ``rfree_flags`` stays binary work/free. if "Validation-flags" in data_dict: @@ -870,7 +877,7 @@ def load(self, reader, french_wilson: bool = True): device=self.device, requires_grad=False, ).to(torch.bool) - self.rfree_source = "MTZ FreeR+Validation" + self.rfree_source = f"{reader_name} FreeR+Validation" self._post_load_cleanup() @@ -1209,7 +1216,13 @@ def _generate_rfree_flags( flags[group_free[group_id]] = 0 self.rfree_flags = flags - self.rfree_source = "Generated (resolution-binned, ASU-grouped)" + # The seed belongs in the provenance string: without it "generated" + # names a draw nobody can reproduce. + self.rfree_source = ( + "Generated (resolution-binned, ASU-grouped" + + (f", seed {seed}" if seed is not None else "") + + ")" + ) n_free = (flags == 0).sum().item() n_work = (flags != 0).sum().item() diff --git a/torchref/io/metadata.py b/torchref/io/metadata.py index b32ad1a5..dfe98079 100644 --- a/torchref/io/metadata.py +++ b/torchref/io/metadata.py @@ -16,6 +16,39 @@ from datetime import date from typing import Any, Dict, List, Optional +#: Input records that describe the crystal, the sample and its chemistry. +#: Refinement moves atoms; it does not invalidate any of these, so they are +#: carried through. Listed in the order the PDB format mandates -- +#: ``render_pdb_header`` emits them in this sequence and relies on it. +#: Deliberately absent: AUTHOR and JRNL (they credit the deposited entry, not +#: this refinement), HELIX/SHEET (nothing here computes secondary structure, so +#: they could only be stale or wrong), and HEADER/REVDAT/OBSLTE/CAVEAT/SPLIT +#: (assertions about a PDB entry that this file is not). +_KEEP_RECORDS = ( + "TITLE", "COMPND", "SOURCE", "KEYWDS", "EXPDTA", "MDLTYP", + "DBREF", "DBREF1", "DBREF2", "SEQADV", "SEQRES", "MODRES", + "HET", "HETNAM", "HETSYN", "FORMUL", + "SSBOND", "LINK", "CISPEP", "SITE", +) + +#: REMARK numbers dropped from the input. 2 is the resolution, which we +#: regenerate; 3 is the refinement, which is ours to write and whose statistics +#: describe a model we just replaced; 500 lists geometry outliers of those same +#: superseded coordinates. +_DROP_REMARKS = {2, 3, 500} + +#: mmCIF categories carried over from an input CIF -- entity, sequence and +#: connectivity, i.e. the CIF counterpart of ``_KEEP_RECORDS``. The input's +#: ``_refine`` is NOT here: it is the previous program's statistics. +_KEEP_CIF_CATEGORIES = ( + "_entity", "_entity_poly", "_entity_poly_seq", "_struct_conn", + "_chem_comp", "_struct_ref", "_struct_ref_seq", "_exptl", +) + +#: Width of the label field in a REMARK 3 line, so the colons line up. Matches +#: the longest label we emit ("RESOLUTION RANGE HIGH (ANGSTROMS)"). +_REMARK3_LABEL_WIDTH = 33 + @dataclass class RefinementMetadata: @@ -28,6 +61,9 @@ class RefinementMetadata: ---------- program, program_version, refinement_method : str Refinement program identification. + target_function, optimizer : str + The function minimised and how, e.g. ``"MAXIMUM LIKELIHOOD"`` and + ``"LBFGS, 5 MACROCYCLES"``. Rendered only when set. resolution_high, resolution_low : float, optional Resolution limits ``d_min`` / ``d_max`` in Angstroms. n_reflections_work, n_reflections_test, n_reflections_all : int, optional @@ -46,16 +82,31 @@ class RefinementMetadata: Unit cell ``[a, b, c, alpha, beta, gamma]`` and space-group name. title, authors Structure title and author names. - passthrough_pdb_remarks, passthrough_cif_categories - Raw REMARK lines / mmCIF category items carried over from an input file. - custom_remarks : list of str - Extra REMARK 3 lines to append. + starting_model : str, optional + Input model this refinement started from (path or PDB ID). + rfree_selection : str, optional + Where the free-set flags came from. Filled from + ``ReflectionData.rfree_source``, so the values are that field's: + ``"MTZReader FreeR"`` for flags read from the input, or + ``"Generated (resolution-binned, ASU-grouped, seed N)"`` for a draw this + run made. Two refinements with different free sets have incomparable + R-free values, which is why it is recorded rather than inferred. + output_remarks : str + Author-supplied free text. Rendered only when non-empty. + software_chain : list of dict + Programs applied before this refinement, read from an input mmCIF's + ``_software`` loop, so ours appends to the chain instead of erasing it. + passthrough_pdb_remarks, passthrough_pdb_records, passthrough_cif_categories + Surviving REMARK lines, structural records keyed by record name, and + mmCIF category items carried over from an input file. """ # Program identification program: str = "TORCHREF" program_version: str = "" refinement_method: str = "" # e.g. "difference-refine", "LBFGS" + target_function: str = "" # e.g. "MAXIMUM LIKELIHOOD" + optimizer: str = "" # e.g. "LBFGS, 5 MACROCYCLES" # Resolution resolution_high: Optional[float] = None # d_min in Angstroms @@ -97,13 +148,25 @@ class RefinementMetadata: title: str = "" authors: List[str] = field(default_factory=list) - # Pass-through: raw header lines from input file + # Provenance + starting_model: Optional[str] = None + rfree_selection: Optional[str] = None + + # Author-supplied free text, rendered as REMARK 3 OTHER REFINEMENT REMARKS + # and _refine.details. Never generated -- if the author has nothing to say, + # the field stays empty and neither is emitted. + output_remarks: str = "" + + # Programs that touched the model before us, read from the input's + # _software loop so ours can be appended rather than replacing the chain. + software_chain: List[Dict[str, str]] = field(default_factory=list) + + # Pass-through from the input file: surviving REMARKs, structural records + # keyed by record name (see _KEEP_RECORDS), and mmCIF categories. passthrough_pdb_remarks: List[str] = field(default_factory=list) + passthrough_pdb_records: Dict[str, List[str]] = field(default_factory=dict) passthrough_cif_categories: Dict[str, Any] = field(default_factory=dict) - # Custom remarks - custom_remarks: List[str] = field(default_factory=list) - # ------------------------------------------------------------------ # # Serialization # ------------------------------------------------------------------ # @@ -238,6 +301,25 @@ def from_refinement(cls, refinement) -> RefinementMetadata: if getattr(refinement, "verbose", 0) > 0: print(f"Could not record solvent parameters: {exc}") + # --- Provenance --- + # Both are recorded on the objects already; the header just had no way + # to say them. rfree_source in particular is what distinguishes a test + # set read from the input file from one this run drew itself, and hence + # whether R-free is comparable to the number the input reported. + try: + input_file = refinement.model.ctx.input_file + if input_file: + meta.starting_model = str(input_file) + except Exception: + pass + + try: + source = refinement.reflection_data.rfree_source + if source: + meta.rfree_selection = source + except Exception: + pass + # --- Cell and spacegroup --- try: model = refinement.model @@ -255,45 +337,82 @@ def from_refinement(cls, refinement) -> RefinementMetadata: # ------------------------------------------------------------------ # @classmethod - def from_pdb_file(cls, filepath: str) -> RefinementMetadata: - """Extract header metadata from an existing PDB file. - - Captures TITLE, AUTHOR, and REMARK records for pass-through. + def from_pdb_file( + cls, filepath: str, *, supersede_refinement: bool = True + ) -> RefinementMetadata: + """Extract the carry-through header of an existing PDB file. + + Captures TITLE, the structural records in ``_KEEP_RECORDS`` and every + REMARK except those in ``_DROP_REMARKS``. AUTHOR and JRNL are absent + from ``_KEEP_RECORDS`` and so never collected: they credit whoever + deposited the entry, not this refinement. + + The input's REMARK 3 is dropped rather than carried, which is the whole + point -- a refined file that repeats the previous program's R-factors + alongside its own asserts two different refinements at once. + + Parameters + ---------- + supersede_refinement : bool, optional + Whether this file's refinement is about to be replaced -- the + default, and the case for refinement output. Pass ``False`` when + annotating a file without re-refining it: nothing supersedes the + existing REMARK 3 or AUTHOR records then, and dropping them would + lose statistics and credit that are still accurate. """ meta = cls() - remarks = [] + remarks: List[str] = [] + records: Dict[str, List[str]] = {} try: with open(filepath, "r") as f: for line in f: record = line[:6].strip() - if record in ("ATOM", "HETATM"): + # MODEL as well as the atoms: it opens the coordinate + # section in a multi-model file. + if record in ("ATOM", "HETATM", "MODEL"): break - if record == "TITLE": + if record == "AUTHOR" and not supersede_refinement: + for author in line[10:].strip().split(","): + if author.strip(): + meta.authors.append(author.strip()) + elif record == "TITLE": + # Kept on `title` rather than as a raw record so the + # --title override has something to override. title_text = line[10:].strip() - if meta.title: - meta.title += " " + title_text - else: - meta.title = title_text - elif record == "AUTHOR": - author_text = line[10:].strip() - # Authors are comma-separated in PDB - for author in author_text.split(","): - author = author.strip() - if author: - meta.authors.append(author) - elif record.startswith("REMARK"): + meta.title = ( + meta.title + " " + title_text if meta.title else title_text + ) + elif record == "REMARK": + try: + number = int(line[7:10]) + except ValueError: + # A REMARK with no parsable number is not one we can + # judge, so leave it out. + continue + if supersede_refinement and number in _DROP_REMARKS: + continue remarks.append(line.rstrip("\n")) + elif record in _KEEP_RECORDS: + records.setdefault(record, []).append(line.rstrip("\n")) meta.passthrough_pdb_remarks = remarks + meta.passthrough_pdb_records = records except Exception: pass return meta @classmethod def from_cif_file(cls, filepath: str) -> RefinementMetadata: - """Extract refinement metadata from an existing mmCIF file. + """Extract the carry-through metadata of an existing mmCIF file. + + Captures ``_struct.title``, the entity/sequence/connectivity categories + in ``_KEEP_CIF_CATEGORIES``, and the ``_software`` loop -- the last so + this refinement can append itself to the chain of programs rather than + presenting itself as the only one. - Captures ``_struct.title``, ``_audit_author.name``, and - ``_refine`` category items for pass-through. + The input's ``_refine`` category is deliberately NOT captured. It holds + the previous program's R-factors and resolution, which this refinement + supersedes; carrying them forward is the mmCIF form of the duplicated + REMARK 3 that ``from_pdb_file`` used to produce. """ meta = cls() try: @@ -302,40 +421,51 @@ def from_cif_file(cls, filepath: str) -> RefinementMetadata: doc = gemmi.cif.read(filepath) block = doc[0] - # Title title = block.find_value("_struct.title") if title and title != "?": meta.title = gemmi.cif.as_string(title) - # Authors - author_loop = block.find(["_audit_author.name"]) - if author_loop: - for row in author_loop: - name = gemmi.cif.as_string(row[0]) - if name and name != "?": - meta.authors.append(name) - - # _refine category pass-through - refine_cats = {} - for tag in block.find(["_refine."]): - # Collect all _refine.* pairs - pass - # Use find_values for individual items - for item_name in [ - "_refine.ls_R_factor_R_work", - "_refine.ls_R_factor_R_free", - "_refine.ls_d_res_high", - "_refine.ls_d_res_low", - "_refine.ls_number_reflns_R_work", - "_refine.ls_number_reflns_R_free", - "_refine.B_iso_mean", - ]: - val = block.find_value(item_name) - if val and val not in ("?", "."): - refine_cats[item_name] = gemmi.cif.as_string(val) - - if refine_cats: - meta.passthrough_cif_categories["_refine"] = refine_cats + # Prior programs, so ours lands at max(ordinal) + 1. + chain: List[Dict[str, str]] = [] + table = block.find( + "_software.", + ["name", "?version", "?classification", "?pdbx_ordinal", + "?description"], + ) + for row in table: + name = gemmi.cif.as_string(row[0]) + if not name or name in ("?", "."): + continue + entry = {"name": name} + for idx, key in ((1, "version"), (2, "classification"), + (3, "pdbx_ordinal"), (4, "description")): + if row.has(idx): + val = gemmi.cif.as_string(row[idx]) + if val and val not in ("?", "."): + entry[key] = val + chain.append(entry) + meta.software_chain = chain + + # Entity, sequence and connectivity categories, in whichever form + # the input used them (loop or key-value). + cats: Dict[str, Any] = {} + for item in block: + if item.loop is not None: + tags = list(item.loop.tags) + if tags[0].split(".")[0] not in _KEEP_CIF_CATEGORIES: + continue + cats[tags[0].split(".")[0]] = { + tag: [ + item.loop[r, c] for r in range(item.loop.length()) + ] + for c, tag in enumerate(tags) + } + elif item.pair is not None: + tag, val = item.pair + category = tag.split(".")[0] + if category in _KEEP_CIF_CATEGORIES: + cats.setdefault(category, {})[tag] = val + meta.passthrough_cif_categories = cats except Exception: pass @@ -372,9 +502,23 @@ def merge(self, other: RefinementMetadata) -> RefinementMetadata: a for a in other_val if a not in self_val ] setattr(merged, f.name, merged_authors) - elif f.name == "custom_remarks": - merged_remarks = list(self_val) + list(other_val) - setattr(merged, f.name, merged_remarks) + elif f.name == "passthrough_pdb_records": + merged_records = {k: list(v) for k, v in self_val.items()} + for key, rows in other_val.items(): + existing = merged_records.setdefault(key, []) + existing.extend([r for r in rows if r not in existing]) + setattr(merged, f.name, merged_records) + elif f.name == "software_chain": + # Keyed on name+version: re-refining with the same build should + # not add a second identical link to the chain. + merged_chain = list(self_val) + seen = {(e.get("name"), e.get("version")) for e in merged_chain} + for entry in other_val: + key = (entry.get("name"), entry.get("version")) + if key not in seen: + merged_chain.append(entry) + seen.add(key) + setattr(merged, f.name, merged_chain) else: # other takes precedence if non-None and non-default if other_val is not None and other_val != "" and other_val != []: @@ -388,56 +532,82 @@ def merge(self, other: RefinementMetadata) -> RefinementMetadata: # ------------------------------------------------------------------ # def render_pdb_header(self) -> str: - """Render metadata as PDB header records (REMARK 3, TITLE, AUTHOR). + """Render the header as PDB records, ready to precede CRYST1. - Returns - ------- - str - Multi-line string ready to insert into a PDB file. + Records come out in the order the PDB format mandates -- TITLE and the + entry-level records, then REMARKs in ascending numeric order with ours + slotted in at 3, then sequence, chemistry and connectivity. Only the + REMARK 3 block is generated; everything else is either carried through + from the input or supplied by the caller. """ lines: List[str] = [] + records = self.passthrough_pdb_records - # Pass-through remarks first (from input file) - for remark in self.passthrough_pdb_remarks: - lines.append(remark) + def _emit(*names: str) -> None: + for name in names: + lines.extend(records.get(name, [])) - # TITLE + # -- entry level -------------------------------------------------- # if self.title: _wrap_pdb_record(lines, "TITLE", self.title) + _emit("COMPND", "SOURCE", "KEYWDS", "EXPDTA", "MDLTYP") - # AUTHOR + # Only ever what the caller set explicitly: authors are not inherited + # from the input, since they credit that deposition and not this run. if self.authors: - author_str = ", ".join(self.authors) - _wrap_pdb_record(lines, "AUTHOR", author_str) + _wrap_pdb_record(lines, "AUTHOR", ", ".join(self.authors)) + + # -- REMARKs, ascending, ours at 3 -------------------------------- # + def _remark_number(line: str) -> int: + try: + return int(line[7:10]) + except ValueError: + return 0 + + passthrough = sorted(self.passthrough_pdb_remarks, key=_remark_number) + lines.extend(r for r in passthrough if _remark_number(r) < 3) + lines.extend(self._render_remark3()) + lines.extend(r for r in passthrough if _remark_number(r) > 3) + + # -- sequence, chemistry, connectivity ---------------------------- # + _emit("DBREF", "DBREF1", "DBREF2", "SEQADV", "SEQRES", "MODRES", + "HET", "HETNAM", "HETSYN", "FORMUL", + "SSBOND", "LINK", "CISPEP", "SITE") - # REMARK 3 - Refinement statistics + return "\n".join(lines) + "\n" + + def _render_remark3(self) -> List[str]: + """Build the REMARK 3 block: this refinement, and only this one.""" + lines: List[str] = [] lines.append("REMARK 3") lines.append("REMARK 3 REFINEMENT.") - lines.append( - f"REMARK 3 PROGRAM : {self.program} {self.program_version}".rstrip() - ) + _ident(lines, "PROGRAM", f"{self.program} {self.program_version}".strip()) if self.refinement_method: - lines.append( - f"REMARK 3 METHOD : {self.refinement_method}" - ) + _ident(lines, "METHOD", self.refinement_method) + if self.target_function: + _ident(lines, "TARGET", self.target_function) + if self.optimizer: + _ident(lines, "OPTIMIZER", self.optimizer) lines.append("REMARK 3") - # Data used in refinement lines.append("REMARK 3 DATA USED IN REFINEMENT.") _remark3(lines, "RESOLUTION RANGE HIGH (ANGSTROMS)", self.resolution_high, ".2f") _remark3(lines, "RESOLUTION RANGE LOW (ANGSTROMS)", self.resolution_low, ".2f") _remark3(lines, "NUMBER OF REFLECTIONS", self.n_reflections_all, "d") lines.append("REMARK 3") - # Fit to data lines.append("REMARK 3 FIT TO DATA USED IN REFINEMENT.") + # Where the free set came from, before the R-factors it conditions: + # R-free values from different test sets are not comparable, and the + # reader has no way to tell without this. + if self.rfree_selection: + _remark3(lines, "FREE R VALUE TEST SET SELECTION", self.rfree_selection) _remark3(lines, "R VALUE (WORKING SET)", self.r_work, ".4f") _remark3(lines, "FREE R VALUE", self.r_free, ".4f") _remark3(lines, "FREE R VALUE TEST SET SIZE (%)", self.percent_free, ".1f") _remark3(lines, "FREE R VALUE TEST SET COUNT", self.n_reflections_test, "d") lines.append("REMARK 3") - # B-values lines.append("REMARK 3 B VALUES.") # Wilson-plot B is not computed; passing None makes _remark3 render the # literal "NULL" here intentionally (not a bug). @@ -447,33 +617,35 @@ def render_pdb_header(self) -> str: _remark3(lines, "B MAX (A**2)", self.b_max, ".2f") lines.append("REMARK 3") - # RMS deviations lines.append("REMARK 3 RMS DEVIATIONS FROM IDEAL VALUES.") _remark3(lines, "BOND LENGTHS (A)", self.rmsd_bond_lengths, ".3f") _remark3(lines, "BOND ANGLES (DEGREES)", self.rmsd_bond_angles, ".2f") lines.append("REMARK 3") - # Model contents lines.append("REMARK 3 NUMBER OF NON-HYDROGEN ATOMS USED IN REFINEMENT.") _remark3(lines, "PROTEIN ATOMS", self.n_atoms_protein, "d") _remark3(lines, "SOLVENT ATOMS", self.n_atoms_solvent, "d") _remark3(lines, "TOTAL", self.n_atoms_total, "d") lines.append("REMARK 3") - # Solvent model if self.solvent_model_ksol is not None or self.solvent_model_bsol is not None: lines.append("REMARK 3 BULK SOLVENT MODELLING.") _remark3(lines, "K_SOL", self.solvent_model_ksol, ".4f") _remark3(lines, "B_SOL", self.solvent_model_bsol, ".2f") lines.append("REMARK 3") - # Custom remarks - for remark in self.custom_remarks: - lines.append(f"REMARK 3 {remark}") + if self.starting_model: + lines.append(f"REMARK 3 STARTING MODEL: {self.starting_model}") + lines.append("REMARK 3") - lines.append("REMARK 3") + # The only free text in the block, and the caller wrote all of it. + if self.output_remarks: + lines.append("REMARK 3 OTHER REFINEMENT REMARKS:") + for paragraph in self.output_remarks.splitlines(): + _wrap_remark3_text(lines, paragraph.strip()) + lines.append("REMARK 3") - return "\n".join(lines) + "\n" + return lines # ------------------------------------------------------------------ # # mmCIF rendering @@ -492,16 +664,40 @@ def render_cif_categories(self) -> Dict[str, Dict[str, str]]: """ cats: Dict[str, Dict[str, str]] = {} - # _software - sw = {} - sw["_software.name"] = self.program + # _software: every program applied to this model, in order, ours last. + # Always a loop, even with one entry -- that is what lets the next + # refinement append a link rather than overwrite the chain, which is the + # only record of prior work that mmCIF actually has room for. + chain = list(self.software_chain) + ordinals = [] + for entry in chain: + try: + ordinals.append(int(entry.get("pdbx_ordinal", 0))) + except (TypeError, ValueError): + pass + description = self.refinement_method or ", ".join( + part for part in (self.target_function, self.optimizer) if part + ) + ours = { + "name": self.program, + "classification": "refinement", + "pdbx_ordinal": str(max(ordinals, default=len(chain)) + 1), + } if self.program_version: - sw["_software.version"] = self.program_version - sw["_software.classification"] = "refinement" - if self.refinement_method: - sw["_software.description"] = self.refinement_method - sw["_software.pdbx_ordinal"] = "1" - cats["_software"] = sw + ours["version"] = self.program_version + if description: + ours["description"] = description + chain.append(ours) + columns = [ + key + for key in ("pdbx_ordinal", "name", "version", "classification", + "description") + if any(key in entry for entry in chain) + ] + cats["_software"] = { + f"_software.{key}": [entry.get(key, "?") for entry in chain] + for key in columns + } # _struct if self.title: @@ -541,9 +737,26 @@ def render_cif_categories(self) -> Dict[str, Dict[str, str]]: ref["_refine.solvent_model_param_ksol"] = f"{self.solvent_model_ksol:.4f}" if self.solvent_model_bsol is not None: ref["_refine.solvent_model_param_bsol"] = f"{self.solvent_model_bsol:.2f}" + # Free-set provenance sits beside the R-factors it conditions: the two + # R-free values either side of a changed test set are not comparable. + if self.rfree_selection: + ref["_refine.pdbx_R_Free_selection_details"] = self.rfree_selection + if self.starting_model: + ref["_refine.pdbx_starting_model"] = self.starting_model + if self.refinement_method: + ref["_refine.pdbx_method_to_determine_struct"] = self.refinement_method + if self.output_remarks: + ref["_refine.details"] = self.output_remarks if ref: cats["_refine"] = ref + # What this refinement started from. Standard category, and the only + # structured place to say it. + if self.starting_model: + cats["_pdbx_initial_refinement_model"] = _initial_model_category( + self.starting_model + ) + # _refine_ls_restr (geometry deviations, as loop) if self.rmsd_bond_lengths is not None or self.rmsd_bond_angles is not None: restr_types = [] @@ -582,6 +795,29 @@ def render_cif_categories(self) -> Dict[str, Dict[str, str]]: # ====================================================================== # +def _initial_model_category(starting_model: str) -> Dict[str, str]: + """Describe the starting model the way deposited entries do. + + A four-character stem that looks like a PDB ID (digit then three + alphanumerics, e.g. ``3GR5.pdb``) is reported as an accession code; anything + else is named in ``details`` and left unaccessioned rather than guessed at. + """ + import os + import re + + basename = os.path.basename(starting_model) + stem = os.path.splitext(basename)[0] + cat = { + "_pdbx_initial_refinement_model.id": "1", + "_pdbx_initial_refinement_model.type": "experimental model", + "_pdbx_initial_refinement_model.details": basename, + } + if re.fullmatch(r"[0-9][A-Za-z0-9]{3}", stem): + cat["_pdbx_initial_refinement_model.source_name"] = "PDB" + cat["_pdbx_initial_refinement_model.accession_code"] = stem.upper() + return cat + + def _remark3( lines: List[str], label: str, value: Any, fmt: str = "" ) -> None: @@ -594,8 +830,59 @@ def _remark3( formatted = f"{value:{fmt}}" else: formatted = "NULL" - line = f"REMARK 3 {label} : {formatted}" - lines.append(line) + # Pad the label so the colons align down the block. Callers used to have to + # pre-pad their own labels, and the ones that forgot rendered ragged. + lines.append(f"REMARK 3 {label:<{_REMARK3_LABEL_WIDTH}} : {formatted}") + + +def _ident(lines: List[str], label: str, value: str) -> None: + """Append a ``REMARK 3 LABEL : value`` line, wrapped if long. + + Overflow continues on a further line whose label field is blank and whose + colon stays in the same column, which is what REFMAC does with its own + long values:: + + REMARK 3 AUTHORS : MURSHUDOV,SKUBAK,LEBEDEV,PANNU,STEINER, + REMARK 3 : NICHOLLS,WINN,LONG,VAGIN + + Without this an optimizer description naming the cycles, mode, ADP model and + scale target runs past column 80. + """ + head = f"REMARK 3 {label:<12}: " + cont = f"REMARK 3 {'':<12}: " + width = 80 - len(head) + prefix, current = head, "" + for word in value.split(): + if current and len(current) + 1 + len(word) > width: + lines.append(prefix + current) + prefix, current = cont, word + else: + current = current + " " + word if current else word + if current or prefix is head: + lines.append(prefix + current) + + +def _wrap_remark3_text(lines: List[str], text: str) -> None: + """Append free text as continuation-free ``REMARK 3`` lines. + + REMARK records have no continuation-number field -- unlike TITLE or AUTHOR, + they simply repeat the same number -- so this wraps on width alone. An empty + paragraph becomes a bare ``REMARK 3`` spacer. + """ + prefix = "REMARK 3 " + if not text: + lines.append("REMARK 3") + return + width = 80 - len(prefix) + current = "" + for word in text.split(): + if current and len(current) + 1 + len(word) > width: + lines.append(prefix + current) + current = word + else: + current = current + " " + word if current else word + if current: + lines.append(prefix + current) def _wrap_pdb_record(lines: List[str], record: str, text: str) -> None: diff --git a/torchref/io/pdb.py b/torchref/io/pdb.py index 8e431e79..cb17fd1f 100644 --- a/torchref/io/pdb.py +++ b/torchref/io/pdb.py @@ -492,7 +492,7 @@ def extract_link_records(filepath: str, verbose: int = 0) -> pd.DataFrame: return df -def write(df: pd.DataFrame, filepath: str, template: str = None, metadata=None) -> None: +def write(df: pd.DataFrame, filepath: str, metadata=None) -> None: """ Write a DataFrame to a PDB file. @@ -504,9 +504,6 @@ def write(df: pd.DataFrame, filepath: str, template: str = None, metadata=None) tempfactor, element, charge. filepath : str Output PDB filename. - template : str, optional - PDB template file to copy header from. Deprecated in favour of - ``metadata``; no ``DeprecationWarning`` is emitted when it is used. metadata : RefinementMetadata, optional Metadata to render as PDB header (REMARK 3, TITLE, etc.). @@ -525,14 +522,6 @@ def write(df: pd.DataFrame, filepath: str, template: str = None, metadata=None) if metadata is not None: n.write(metadata.render_pdb_header()) - # Copy template header if provided (deprecated path) - if template is not None: - with open(template) as t: - for line in t: - if "REMARK" not in line and "ATOM" in line: - break - n.write(line) - # Write CRYST1 record if cell info available (directly before atoms) try: cell = df.attrs["cell"] @@ -613,8 +602,8 @@ def write(df: pd.DataFrame, filepath: str, template: str = None, metadata=None) s = ( f"{str(ATOM):<6}{int(serial):>5} {name_field}{str(altloc):>1}" f"{str(resname):>3}{str(chainid):>2}{int(resseq):>4}{str(icode):>4}" - f"{round(x, 3):>8}{round(y, 3):>8}{round(z_coord, 3):>8}" - f"{round(occupancy, 3):>6.2f}{round(tempfactor, 2):>6}" + f"{x:>8.3f}{y:>8.3f}{z_coord:>8.3f}" + f"{occupancy:>6.2f}{tempfactor:>6.2f}" f"{str(element):>12}{charge:>2}\n" ) n.write(s) @@ -730,8 +719,8 @@ def write_multi_model( s = ( f"{str(ATOM):<6}{int(serial):>5} {name_field}{altloc:>1}" f"{resname:>3}{chainid:>2}{resseq:>4}{icode:>4}" - f"{round(x, 3):>8}{round(y, 3):>8}{round(z_coord, 3):>8}" - f"{round(occupancy, 3):>6.2f}{round(tempfactor, 2):>6}" + f"{x:>8.3f}{y:>8.3f}{z_coord:>8.3f}" + f"{occupancy:>6.2f}{tempfactor:>6.2f}" f"{element:>12}{charge_str:>2}\n" ) f.write(s) From f9184d27d07712ae16fd1ce3105b438b42821ac4 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 3 Sep 2026 10:49:07 +0200 Subject: [PATCH 157/250] Loosened bessel function required accuracy in tests --- .../frf_separate/test_bessel_sh_grouping.py | 133 ++++++++++++++---- 1 file changed, 104 insertions(+), 29 deletions(-) diff --git a/tests/unit/frf_separate/test_bessel_sh_grouping.py b/tests/unit/frf_separate/test_bessel_sh_grouping.py index fcfac972..44a83ccb 100644 --- a/tests/unit/frf_separate/test_bessel_sh_grouping.py +++ b/tests/unit/frf_separate/test_bessel_sh_grouping.py @@ -15,11 +15,16 @@ the previous implementation instead would only show that two approximations agree with each other. -The claims about the grouping being *loss-free* are only meaningful at a -precision finer than the loss being ruled out, so the tests that assert -1e-10-and-below take ``double_cpu``. At the working precision (float32) the -floor is float32 epsilon times the accumulation depth -- measured 1.4e-06 and -3.6e-07 on these cases -- which says nothing about the grouping. +Everything here runs at the configured precision on the configured device, so +the expansion is exercised where it actually runs. That fixes what the budgets +can say. At float32 the difference between two groupings is dominated by the +accumulation's own rounding, not by the grouping, and that floor is not +portable: the three grouping cases cost 5.5e-08 to 1.5e-07 on CPU, 4.2e-07 to +7.4e-07 on this machine's MPS, and up to 2.1e-06 on the project's MPS CI, +because the backends reduce in a different order and MPS varies further by GPU +and torch version. So a single working-precision budget covers them all -- +see ``_WORKING_PRECISION_TOLERANCE`` for where the number comes from and what +it costs. """ import math @@ -28,6 +33,7 @@ import torch import torchref.experimental.alignment.frf.data_mr as dm +from torchref.config import get_default_device, get_float_dtype from torchref.experimental.alignment.frf.data_mr import bessel_sh_expand pytestmark = pytest.mark.unit @@ -35,29 +41,73 @@ #: A key this fine puts every reflection in its own group: the exact sum. _UNGROUPED = 10 ** 16 -#: Phaser buckets cos(theta) at 1e-3 (`lib/sphericalY.h:43`) and evaluates the -#: Legendre polynomials once per bucket, which costs it about 2e-5 relative on -#: these coefficients. Staying two orders inside that is ample; the threshold is -#: set from the measured 1.2e-8 to 4.8e-8 with headroom, not from taste. -_GROUPING_TOLERANCE = 5e-7 +#: What any of these comparisons is allowed to cost at the working precision. +#: +#: Measured, worst case over the assertions below: 1.9e-06 on CPU float32 and +#: 1.4e-06 on this machine's MPS, with the grouping cases reaching 2.1e-06 on +#: the project's MPS CI. 1e-04 clears the worst observed figure by ~50x, which is +#: ample for the backend spread -- the same case differs 3x between two MPS +#: devices, and CPU is worse than MPS on the lattice. +#: +#: It still sits far below the accuracy of what feeds it. The sphere/voxel +#: discretisation upstream of the structure factors is itself good to about +#: 5e-03 rel L2 (``tests/unit/base/test_canonical_sphere_cpu.py``), so a budget +#: here of 1e-04 is ~50x tighter than the input it is expanding. Downstream is +#: the same story: rotation-function scores moved 1.8e-07 to 2.5e-04 relative +#: across the float32 migration and the peak lists on 1DAW, 3K7M, 2DQ6 and 4BX9 +#: came back in the same order. +#: +#: What it gives up, and this is the real cost: at this width the grouping +#: comparison notices neither a 10x coarser key (1.6e-06) nor a 100x one +#: (4.1e-05). Only 1000x (1.3e-03) trips it. These are smoke tests for gross +#: breakage on the device the code actually runs on, not tripwires for a +#: degraded key -- ``test_matches_an_independent_direct_summation`` is the one +#: that would still catch a dropped term, because it compares against a +#: reference rather than against the expansion at another grouping. Given the +#: 5e-03 upstream, a key error of 4.1e-05 is not a correctness problem anyway; +#: it would be a performance-versus-accuracy choice made by accident. +_WORKING_PRECISION_TOLERANCE = 1e-4 + +#: Alias kept for the grouping cases, which is what the assertion message names. +_GROUPING_TOLERANCE = _WORKING_PRECISION_TOLERANCE + +#: The direct-summation check compares against a host-double reference rather +#: than against the expansion at another grouping, so it carries the working +#: precision's whole error. Measured 8.1e-07; same budget, same reasoning. +_DIRECT_SUM_TOLERANCE = _WORKING_PRECISION_TOLERANCE + + +def _to_working(s, I, dtype=None, device=None): + """Place a host-double set at the configured precision and device. + + Drawn in double and cast once, rather than generated at the working dtype, + so the *set* is the same whatever precision it is evaluated in -- otherwise + a tolerance measured at one dtype is not comparable with the same tolerance + at another, because the sample moved too. + """ + dtype = get_float_dtype() if dtype is None else dtype + device = get_default_device() if device is None else device + return (s.to(device=device, dtype=dtype), I.to(device=device, dtype=dtype)) -def _random_set(seed, n, dtype=torch.float64): +def _random_set(seed, n, dtype=None, device=None): g = torch.Generator().manual_seed(seed) - s = torch.randn(n, 3, generator=g, dtype=dtype) + s = torch.randn(n, 3, generator=g, dtype=torch.float64) s = s / s.norm(dim=-1, keepdim=True) * ( - 0.07 + 0.18 * torch.rand(n, 1, generator=g, dtype=dtype)) - return s, torch.randn(n, generator=g, dtype=dtype) + 0.07 + 0.18 * torch.rand(n, 1, generator=g, dtype=torch.float64)) + I = torch.randn(n, generator=g, dtype=torch.float64) + return _to_working(s, I, dtype, device) -def _grid_set(k=8, step=0.013): +def _grid_set(k=8, step=0.013, dtype=None, device=None): """A lattice, where |s| degeneracy is exact and the grouping pays most.""" idx = torch.arange(-k, k + 1, dtype=torch.float64) a, b, c = torch.meshgrid(idx, idx, idx, indexing="ij") s = torch.stack([a.reshape(-1), b.reshape(-1), c.reshape(-1)], dim=-1) * step s = s[s.norm(dim=-1) > 1e-9] g = torch.Generator().manual_seed(4) - return s, torch.randn(s.shape[0], generator=g, dtype=torch.float64) + I = torch.randn(s.shape[0], generator=g, dtype=torch.float64) + return _to_working(s, I, dtype, device) @pytest.fixture @@ -88,13 +138,20 @@ def test_grouping_error_stays_far_below_the_reference_implementation( ) -def test_a_lattice_groups_without_loss(ungrouped, double_cpu): +def test_a_lattice_groups_without_loss(ungrouped): """On a lattice the degeneracy is exact, so the grouping is free.""" s, I = _grid_set() ref = ungrouped(s, I, L=65, bessel_h_scale=64.0) got = bessel_sh_expand(s, I, L=65, bessel_h_scale=64.0).coeffs rel = (got - ref).abs().max().item() / max(ref.abs().max().item(), 1e-300) - assert rel < 1e-12, f"lattice grouping is not loss-free: {rel:.2e}" + # "Loss-free" is a float64 statement: on a lattice the |s| degeneracy is + # exact, so in double this lands at 1e-15. At float32 the accumulation floor + # is 1.9e-06 (CPU) / 1.4e-06 (MPS) and swamps it, so what is asserted here + # is that the lattice case is no worse than any other -- the exactness claim + # is only visible in double. + assert rel < _WORKING_PRECISION_TOLERANCE, ( + f"lattice grouping is not loss-free: {rel:.2e}" + ) @pytest.mark.parametrize("L", [17, 33, 65]) @@ -131,10 +188,15 @@ def test_the_group_representative_is_the_mean_not_a_member(): b = unit * (r + eps) I = torch.tensor([1.0, 1.0], dtype=torch.float64) - fwd = bessel_sh_expand(torch.cat([a, b]), I, L=L, bessel_h_scale=scale).coeffs - rev = bessel_sh_expand(torch.cat([b, a]), I, L=L, bessel_h_scale=scale).coeffs + ab, I_w = _to_working(torch.cat([a, b]), I) + ba, _ = _to_working(torch.cat([b, a]), I) + fwd = bessel_sh_expand(ab, I_w, L=L, bessel_h_scale=scale).coeffs + rev = bessel_sh_expand(ba, I_w, L=L, bessel_h_scale=scale).coeffs rel = (fwd - rev).abs().max().item() / max(fwd.abs().max().item(), 1e-300) - assert rel < 1e-13, ( + # Not loosened with the others: the mean of two numbers does not depend on + # their order in floating point either, so this measures 0.0 exactly at + # float32 and float64 alike. A budget here would only hide a real failure. + assert rel == 0.0, ( f"reordering two reflections in the same bin changed the result by " f"{rel:.2e}; the representative is order-dependent, so it is not the mean" ) @@ -190,17 +252,30 @@ def _reference_expansion(s_vec, intensity, L, bessel_h_scale): @pytest.mark.parametrize("L", [9, 13]) -def test_matches_an_independent_direct_summation(L, double_cpu): - """The whole expansion, against a reference that shares no code with it.""" - s, I = _random_set(seed=31, n=120) - ref = _reference_expansion(s, I, L, 24.0) - got = bessel_sh_expand(s, I, L=L, bessel_h_scale=24.0).coeffs +def test_matches_an_independent_direct_summation(L): + """The whole expansion, against a reference that shares no code with it. + + The reference stays in host double while the expansion runs at the working + precision, so this is the one test here that measures the expansion's true + accuracy rather than its self-consistency. That is also why its budget is + the widest: it carries the working precision's whole error, not just the + grouping's share of it. + """ + s64, I64 = _random_set(seed=31, n=120, dtype=torch.float64, device="cpu") + ref = _reference_expansion(s64, I64, L, 24.0) + got = bessel_sh_expand(*_random_set(seed=31, n=120), + L=L, bessel_h_scale=24.0).coeffs assert got.shape == ref.shape + # `.cpu()` before widening: complex128 cannot be materialised on a backend + # that has no double. + got = got.cpu().to(ref.dtype) rel = (got - ref).abs().max().item() / max(ref.abs().max().item(), 1e-300) - assert rel < 1e-10, f"L={L}: differs from a direct summation by {rel:.2e}" + assert rel < _DIRECT_SUM_TOLERANCE, ( + f"L={L}: differs from a direct summation by {rel:.2e}" + ) -def test_the_antipodal_copy_would_only_double_the_result(double_cpu): +def test_the_antipodal_copy_would_only_double_the_result(): """Why the expansion no longer concatenates ``-s``. Only even ``l`` are computed, and ``Y_lm(-s_hat) = (-1)^l Y_lm(s_hat)``, so @@ -220,7 +295,7 @@ def test_the_antipodal_copy_would_only_double_the_result(double_cpu): ).coeffs scale = single.abs().max().clamp(min=1e-300) rel = ((doubled - 2.0 * single).abs().max() / scale).item() - assert rel < 1e-12, ( + assert rel < _WORKING_PRECISION_TOLERANCE, ( f"the antipodal copy is not an exact factor of two: {rel:.2e} relative. " f"Removing it was justified on that being exact." ) From 859a265aafd41164e4f50b14cd25d22eac83aaac Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Fri, 4 Sep 2026 11:20:39 +0200 Subject: [PATCH 158/250] Default the difference map to dark phases, and rename the map CLI torchref.phased-difference-map becomes torchref.difference-map, because it no longer defaults to a phased difference map. The default output is DELFWT/PHDELWT, the inverse-variance-weighted amplitude difference on the dark model's phases. That pair was computed all along as WDF/PHIC_dark, but the file foregrounded 2mDFop-DFc/PHIC_diff instead -- a phased residual that puts the light state's model phases into the observed amplitude, biasing the map toward the model under test. validate_ded.py has always correlated against the dark-phased construction (w_dfo on phi_dark_p1), so the MTZ's headline map and the validation metric were two different objects. They are now one. -lm is optional on the map CLI as a consequence: a weighted difference map needs only the dark state. The amplitude is |Fo_light| - |Fo_dark|, which DatasetCollection.scale puts on one scale with no model at all, and the phase comes from the dark model. --fraction is rejected rather than ignored there. Measured on 1DAW: DELFWT is bit-identical with and without a light model, and the two maps correlate at 1.000000 (median |dphi| 0.0009 deg from the differing f_sol fits). The column set went from 33 (46 under --two-moment) to 17, the rest behind --all-columns, with standard CCP4 names so Coot opens both maps unprompted. write_results_mtz is four layers, each declaring its columns' MTZ types beside the values, and now calls infer_mtz_dtypes as the canonical writer does; types used to come from four parallel name lists with no fallback. The default path runs one internal scale fit instead of three, and none with no light model. Fixed the empirical-Bayes extrapolation propagating sigma_ext^2 as (sigma_L^2 + sigma_D^2)/f^2. F_ext is linear in the observations with dF_ext/dF_dark = -(1-f)/f, so the dark term carries a (1-f)^2 weight: at f = 0.22 it was over-weighted 1.64x, biasing tau^2 low and over-shrinking every reflection. The estimator now takes the caller's already-correct propagation instead of rebuilding it, and returns the shrunk amplitude it is named for -- the caller used to discard the return value and recompute it. This changes FEXT, which is now the default extrapolation, so it wants a paired re-measurement. Also fixed paper/make_ded_maps.py pairing WDF with PHIC_diff: the right amplitude on the model difference phase, which is not the map validate-ded reports. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01FMRk8fecGsRgNErejTvEU3 --- AGENTS.md | 2 +- docs/changelog.rst | 9 + docs/user_guide/cli.rst | 32 +- paper/figure4_difference_refinement/README.md | 5 +- paper/make_ded_maps.py | 23 +- paper/probe_ded_metric_space.py | 2 +- pyproject.toml | 2 +- tests/integration/test_cli_two_moment_mtz.py | 234 ++++-- torchref/cli/__init__.py | 2 +- torchref/cli/_common.py | 46 +- torchref/cli/collection_difference_refine.py | 711 +++++++++++------- torchref/cli/difference_map.py | 241 ++++++ torchref/cli/phased_difference_map.py | 184 ----- 13 files changed, 970 insertions(+), 523 deletions(-) create mode 100644 torchref/cli/difference_map.py delete mode 100644 torchref/cli/phased_difference_map.py diff --git a/AGENTS.md b/AGENTS.md index ab05f37f..d02be91b 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -192,7 +192,7 @@ Black, 88 columns, `isort` with the black profile. Ruff lint with | `scaling/` | `ScalerBase` (model-independent), `Scaler`, `CollectionScaler`, `SolventModel` (k_sol, B_sol) | | `symmetry/` | `Symmetry` (operations plus everything derived from them), `SpaceGroup` (adds the crystallographic identity and the CCP4 ASU verbs), `Cell`. All dataclasses over `DeviceMixin`, not `nn.Module` — they hold no refinable parameters. Map and reciprocal-grid operators are private, reached through `Symmetry` | | `maps/` | `Map` (2Fo−Fc, Fcalc), `DifferenceMap` | -| `cli/` | Entry points: `torchref.refine`, `torchref.difference-refine`, `torchref.mtz2map`, `torchref.validate-ded`, `torchref.phased-difference-map`, `torchref.add-metadata`, `torchref.strip-altlocs` | +| `cli/` | Entry points: `torchref.refine`, `torchref.difference-refine`, `torchref.mtz2map`, `torchref.validate-ded`, `torchref.difference-map`, `torchref.add-metadata`, `torchref.strip-altlocs` | | `experimental/` | APIs that may change without notice: `alignment/` (Patterson MR), `kinetic/` (time-resolved), `ensemble/`, `monolithic_refinement/`, `targets/` (AMBER/GAFF2, real-space, sampled-ML phase) | | `utils/` | See §5 | | `config.py` | See §4 | diff --git a/docs/changelog.rst b/docs/changelog.rst index 04c08a48..5b7de400 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,15 @@ Changelog Unreleased ---------- +- ``torchref.phased-difference-map`` is now ``torchref.difference-map``, because it no longer defaults to a phased difference map. The default output is ``DELFWT``/``PHDELWT``, the inverse-variance-weighted amplitude difference on the **dark** model's phases -- the construction ``torchref.validate-ded`` correlates against, and the one the figure-4 CC was measured on. The writer computed that pair all along as ``WDF``/``PHIC_dark`` but foregrounded ``2mDFop-DFc``/``PHIC_diff`` instead, a phased *residual* that puts the light state's model phases into the observed amplitude and so biases the map toward the model under test. The headline map and the validation metric are now one object +- ``-lm``/``--light-model`` is optional on ``torchref.difference-map``. A weighted difference map needs only the dark state: the amplitude is ``|Fo_light| - |Fo_dark|``, which ``DatasetCollection.scale`` puts on one scale with no model at all, and the phase comes from the dark model. Without a light model there is no mixed model, no ``--fraction`` and no joint model-to-data fit, and ``--fraction`` is rejected rather than ignored. Still required on ``torchref.difference-refine``, which refines it +- The difference MTZ went from 33 columns (46 under ``--two-moment``) to 17, with the rest behind ``--all-columns``. The default file holds the difference map, the extrapolated map ``FWT``/``PHWT``, the observations and the flags; the gated set holds the alternative constructions -- the phased difference residuals, two further extrapolations, the intensity block. Standard CCP4 names throughout the default set, so Coot and CCP4 open both maps without being told which columns to use. The default path also runs one internal scale fit instead of three, and none at all with no light model +- Fixed the empirical-Bayes extrapolation propagating ``sigma_ext^2 = (sigma_L^2 + sigma_D^2)/f^2``. ``F_ext`` is linear in the observations with ``dF_ext/dF_dark = -(1-f)/f``, so the dark term carries a ``(1-f)^2`` weight: at f = 0.22 the term was over-weighted 1.64x, biasing ``tau^2`` low and over-shrinking every reflection. The estimator now takes the caller's already-correct propagation instead of rebuilding it, so the shrinkage weight and the written ``SIGFEXT`` cannot disagree. It also returns the shrunk amplitude it is named for; the caller used to discard the return value and recompute it. Correcting it raises tau^2 and shrinks less, which moves the default extrapolated map +- ``FWT``/``PHWT`` carries the phase of its own scale fit. The Bayes coefficients had no phase column and would have been paired with one fitted against the phase-aware amplitudes; the scaler contributes a phase through ``f_sol``, so the three extrapolation fits do not agree +- ``write_results_mtz`` is four layers, each declaring its columns' MTZ types beside the values, and it now calls ``infer_mtz_dtypes`` as the canonical writer does. Types used to come from four parallel name lists with no fallback, so a column added to the output dict and missed in the lists was written with whatever dtype numpy produced. A column with no declared type is now an error. Its column table moved onto the function from the module docstring four hundred lines away +- ``DFc_complex`` is ``DFc_phased``: the column holds a real amplitude, and the old name said otherwise. The undocumented ``Fextp``/``Fextc``/``Fextb`` suffixes are ``FEXT_PHASED``/``FEXT_SCALAR``/``FEXT``, and ``SIGFEXT_PHASED`` is written at last -- it was computed all along while the docs claimed it existed +- ``tau_sq`` and mean ``w(h)`` go in the JSON summary. With the Bayes extrapolation as the default map, they are what says whether it is over-shrunk +- Fixed ``paper/make_ded_maps.py`` pairing ``WDF`` with ``PHIC_diff`` -- the right amplitude on the model difference phase, which is not the map ``validate-ded`` reports - Refinement output no longer inherits the input file's refinement header. It used to copy the whole thing and then append its own ``REMARK 3``, so a refined 3GR5 carried 420 header lines asserting two refinements at once -- ``PROGRAM : REFMAC 5.1.24`` with R-work 0.213 at line 5, ours at line 389 -- and a reader taking the first ``REMARK 3`` got REFMAC. The inherited block was not merely stale but contradicted the data beside it: it claimed a 5.1% / 1072-reflection test set, while the MTZ shipped with it holds 9.85% / 2063 (which torchref reads correctly). The passthrough was also inverted, keeping the statistics refinement invalidates and dropping the chemistry it does not -- SEQRES, SSBOND, DBREF, EXPDTA, COMPND, SOURCE, KEYWDS, SEQADV, HETNAM, FORMUL and SITE were all absent from the output. Now a whitelist carries the crystal, sample and chemistry records through in mandated record order (TITLE used to be emitted after REMARK 900), REMARK 2, 3 and 500 are dropped, and AUTHOR and JRNL are not inherited because they credit the deposition rather than this run. 283 header lines for the same file, 41 structural records preserved, one refinement block. ``add-metadata`` is exempt through ``supersede_refinement=False``: annotating a file is not re-refining it, so nothing there supersedes the existing REMARK 3 or AUTHOR records and both are kept - Prior refinements are tracked through mmCIF's ``_software`` loop, which is the only place either format has room for them: ``_refine`` is singular by design, so a previous program's statistics cannot be kept without contradicting the current ones. ``pdbx_ordinal`` was hardcoded to ``1`` and the incoming loop was never read, truncating the chain to one link on every write; it now reads the input's loop and appends at ``max(ordinal) + 1``, carrying each entry's ``description`` so the chain says what every program did and not just that it ran. Added ``_pdbx_initial_refinement_model`` and ``_refine.pdbx_starting_model``, which name what the refinement started from, and ``_refine.pdbx_R_Free_selection_details``, which names the test set the reported R-free conditions on. ``from_cif_file`` no longer carries the input's ``_refine`` items through -- that was the mmCIF form of the duplicated ``REMARK 3`` - Fixed mmCIF loop cells being written unquoted, which silently split any value containing whitespace into extra columns when the file was read back. Latent while every loop column was a single token; a multi-word ``_software.description`` exposed it. Nulls stay bare, since ``gemmi.cif.quote`` turns ``?`` into the quoted one-character string ``'?'`` diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index 517cfb30..0729fa84 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -125,21 +125,39 @@ selection), ``--mask-radius``, ``--n-bins``. :API: :mod:`torchref.cli.validate_ded` -``torchref.phased-difference-map`` -~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +``torchref.difference-map`` +~~~~~~~~~~~~~~~~~~~~~~~~~~~ -Compute phased difference and extrapolated map coefficients without -refinement. Uses the same pipeline as ``torchref.difference-refine`` but -the input models are kept as-is. +Compute difference and extrapolated map coefficients without refinement. +Uses the same pipeline as ``torchref.difference-refine`` but the input +models are kept as-is. + +The default output is the weighted difference map ``DELFWT``/``PHDELWT`` -- +the inverse-variance-weighted amplitude difference on the **dark** model's +phases, which is the construction ``torchref.validate-ded`` correlates +against. It needs no light-state model, so ``-lm`` is optional: + +.. code-block:: bash + + torchref.difference-map \ + -dm dark.pdb \ + -dsf dark.mtz -lsf light.mtz -o results.mtz + +Supplying ``-lm`` (with ``--fraction``) adds the light state's amplitude and +phase and the extrapolated map ``FWT``/``PHWT``: .. code-block:: bash - torchref.phased-difference-map \ + torchref.difference-map \ -dm dark.pdb -lm light.pdb \ -dsf dark.mtz -lsf light.mtz \ --fraction 0.37 -o results.mtz -:API: :mod:`torchref.cli.phased_difference_map` +**Key options:** ``--all-columns`` writes every alternative map coefficient +and diagnostic -- the model-phased difference, the two other extrapolations +and the intensity block -- at the cost of two further scale fits. + +:API: :mod:`torchref.cli.difference_map` Model Utilities --------------- diff --git a/paper/figure4_difference_refinement/README.md b/paper/figure4_difference_refinement/README.md index 7aaf0dd4..2ba275ed 100644 --- a/paper/figure4_difference_refinement/README.md +++ b/paper/figure4_difference_refinement/README.md @@ -51,8 +51,9 @@ Key parameters: The refinement script is: `run_difference_refine.sh` -**Stage 2 — Manual real-space refinement** in Coot against the 2Fext-Fc difference -electron density map (from the `difference_data.mtz` output). This step speeds up +**Stage 2 — Manual real-space refinement** in Coot against the extrapolated map +(`FWT`/`PHWT` in the `difference_data.mtz` output; Coot opens it by name). This step +speeds up convergence for the IBL ligand, which requires large conformational changes (*trans* to *cis*). The Coot-refined coordinates were saved as `work.pdb` and used as the starting model for automated refinement (Stage 1). diff --git a/paper/make_ded_maps.py b/paper/make_ded_maps.py index 47bca6e9..31bc3377 100644 --- a/paper/make_ded_maps.py +++ b/paper/make_ded_maps.py @@ -29,13 +29,24 @@ import numpy as np # label -> (amplitude column, phase column). Skipped silently when absent. +# +# The dark-phased entries come first because they are the default output and the +# construction ``torchref.validate-ded`` correlates against. ``ddf`` and ``wdf`` used to +# be paired with ``PHIC_diff``, the *model* difference phase -- the right amplitude on +# the wrong phase, and a different map from the one being validated. +# +# The ``PHIC_diff`` entries need ``--all-columns`` on the writer. They are phased +# difference *residuals*: the light state's model phases enter the observed amplitude, +# so they are model-biased where the dark-phased maps are not. MAPS = { - "ded": ("mDFop-DFc", "PHIC_diff"), - "ded_corr": ("mDFop-DFc_corr", "PHIC_diff"), - "ded2": ("2mDFop-DFc", "PHIC_diff"), - "ded2_corr": ("2mDFop-DFc_corr", "PHIC_diff"), - "ddf": ("DDF", "PHIC_diff"), - "wdf": ("WDF", "PHIC_diff"), + "ded": ("DELFWT", "PHDELWT"), + "ded_corr": ("DELFWT_corr", "PHDELWT"), + "ddf": ("DDF", "PHDELWT"), + "ext": ("FWT", "PHWT"), + "ded_phased": ("mDFop-DFc", "PHIC_diff"), + "ded_phased_corr": ("mDFop-DFc_corr", "PHIC_diff"), + "ded2_phased": ("2mDFop-DFc", "PHIC_diff"), + "ded2_phased_corr": ("2mDFop-DFc_corr", "PHIC_diff"), } diff --git a/paper/probe_ded_metric_space.py b/paper/probe_ded_metric_space.py index 8b9c3b92..e5cf720d 100644 --- a/paper/probe_ded_metric_space.py +++ b/paper/probe_ded_metric_space.py @@ -95,7 +95,7 @@ def main(): Fo_d = ds["Fo_dark"].to_numpy(float) Fo_l = ds["Fo_light"].to_numpy(float) Fc_d = ds["Fc_dark"].to_numpy(float) - Fc_l = ds["Fc_light"].to_numpy(float) + Fc_l = ds["FC"].to_numpy(float) # was Fc_light free = ds["FreeR_flag_light"].to_numpy() == 0 Id = np.full(len(H), np.nan); Il = np.full(len(H), np.nan) diff --git a/pyproject.toml b/pyproject.toml index 2e8bf108..0f30beb9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -48,7 +48,7 @@ dependencies = [ "torchref.simulate-noisy-data" = "torchref.cli.simulate_noisy_data:main" "torchref.mtz2map" = "torchref.cli.mtz2map:main" "torchref.validate-ded" = "torchref.cli.validate_ded:main" -"torchref.phased-difference-map" = "torchref.cli.phased_difference_map:main" +"torchref.difference-map" = "torchref.cli.difference_map:main" "torchref.add-metadata" = "torchref.cli.add_metadata:main" "torchref.strip-altlocs" = "torchref.cli.strip_altlocs:main" diff --git a/tests/integration/test_cli_two_moment_mtz.py b/tests/integration/test_cli_two_moment_mtz.py index 06ebe919..e8515e9b 100644 --- a/tests/integration/test_cli_two_moment_mtz.py +++ b/tests/integration/test_cli_two_moment_mtz.py @@ -1,13 +1,14 @@ """The difference-refinement MTZ layout, pinned. -``write_results_mtz`` assigns MTZ column types from hard-coded name lists and never calls -``infer_mtz_dtypes()``, so a column added to the output dict but missed in the type lists -is written with whatever dtype numpy produced -- silently, and into a file that gets -deposited. Nothing else in the suite asserts on these names. - -Two things are checked: the baseline column set is unchanged by the two-moment work, and -under ``--two-moment`` the thirteen extra columns appear with the right types and are -internally consistent. +Nothing else in the suite asserts on these column names, and they go into files that get +deposited, so the layout is pinned here deliberately: this file is expected to move in +lockstep with a change to the writer, and to fail loudly if one happens by accident. + +Three layouts are checked. The default is the map a reader wants and can identify -- +``DELFWT``/``PHDELWT``, the weighted difference on dark phases, plus the extrapolated +map. ``--two-moment`` adds the activation-heterogeneity correction. ``--all-columns`` +adds the alternative constructions of both, which are informative once you know which is +which and misleading before then. """ import json @@ -20,32 +21,48 @@ pytestmark = [pytest.mark.integration, pytest.mark.slow] -BASELINE_COLUMNS = { - "Fo_dark", "SIGFo_dark", "Fo_light", "SIGFo_light", - "DF", "SIGDF", "WDF", - "Fc_dark", "Fc_light", "DFc", "DFc_complex", - "2mDFop-DFc", "mDFop-DFc", - "PHIC_dark", "PHIC_mixed", "PHIC_diff", "PHIC_light", - "Fextp", "2Fextp-Fc", "Fextp-Fc", - "Fextc", "SIGFextc", "2Fextc-Fc", "Fextc-Fc", - "Fextb", "SIGFextb", "2Fextb-Fc", "Fextb-Fc", - "FreeR_flag_dark", "FreeR_flag_light", +# The default set, with a light model supplied (which the refinement CLI always does). +# H/K/L are the index, so they are not among ``.columns``. +DEFAULT_COLUMNS = { + "Fo_dark": "SFAmplitude", "SIGFo_dark": "Stddev", + "Fo_light": "SFAmplitude", "SIGFo_light": "Stddev", + "DF": "SFAmplitude", "SIGDF": "Stddev", + # The difference map. CCP4/Coot open these by name. + "DELFWT": "SFAmplitude", "PHDELWT": "Phase", + "Fc_dark": "SFAmplitude", + # The mixed model, and the extrapolated map to refine against. + "FC": "SFAmplitude", "PHIC": "Phase", + "FEXT": "SFAmplitude", "SIGFEXT": "Stddev", + "FWT": "SFAmplitude", "PHWT": "Phase", } +FLAG_COLUMNS = {"FreeR_flag_dark", "FreeR_flag_light"} TWO_MOMENT_COLUMNS = { - "Io_light": "Intensity", - "SIGIo_light": "Stddev", - "Ic_light_coh": "Intensity", - "Ic_light_2mom": "Intensity", - "IVAR_ALPHA": "Intensity", - "Fo_light_corr": "SFAmplitude", - "SIGFo_light_corr": "Stddev", - "DF_corr": "SFAmplitude", - "SIGDF_corr": "Stddev", - "2mDFop-DFc_corr": "SFAmplitude", - "mDFop-DFc_corr": "SFAmplitude", + "DELFWT_corr": "SFAmplitude", + "Fo_light_corr": "SFAmplitude", "SIGFo_light_corr": "Stddev", + "DF_corr": "SFAmplitude", "SIGDF_corr": "Stddev", "DDF": "SFAmplitude", - "W_2MOM": "Weight", +} + +# What ``--all-columns`` adds on top, given a light model. +ALL_COLUMNS_EXTRA = { + "2mDFop-DFc": "SFAmplitude", "mDFop-DFc": "SFAmplitude", + "PHIC_diff": "Phase", + "DFc": "SFAmplitude", "DFc_phased": "SFAmplitude", + "FEXT_PHASED": "SFAmplitude", "SIGFEXT_PHASED": "Stddev", + "2FEXT_PHASED-Fc": "SFAmplitude", "FEXT_PHASED-Fc": "SFAmplitude", + "PHFEXT_PHASED": "Phase", + "FEXT_SCALAR": "SFAmplitude", "SIGFEXT_SCALAR": "Stddev", + "2FEXT_SCALAR-Fc": "SFAmplitude", "FEXT_SCALAR-Fc": "SFAmplitude", + "PHFEXT_SCALAR": "Phase", +} + +# And what it adds again once the two-moment model is on. +ALL_COLUMNS_TWO_MOMENT_EXTRA = { + "Io_light": "Intensity", "SIGIo_light": "Stddev", + "Ic_light_coh": "Intensity", "Ic_light_2mom": "Intensity", + "IVAR_ALPHA": "Intensity", "W_2MOM": "Weight", + "2mDFop-DFc_corr": "SFAmplitude", "mDFop-DFc_corr": "SFAmplitude", } FRACTION = 0.25 @@ -131,27 +148,81 @@ def two_moment_mtz(cli_script, intensity_pair, tmp_path_factory): ) +@pytest.fixture(scope="module") +def two_moment_all_mtz(cli_script, intensity_pair, tmp_path_factory): + """Everything on. The value-consistency tests below need the diagnostic columns, + which is exactly what ``--all-columns`` is for.""" + outdir = tmp_path_factory.mktemp("two_moment_all") + return _run( + cli_script, intensity_pair, outdir, + "--two-moment", "--lambda-twin", str(LAMBDA_TWIN), "--all-columns", + ) + + def _read(path): import reciprocalspaceship as rs return rs.read_mtz(str(path)) -class TestBaselineLayoutIsUnchanged: - def test_baseline_columns_are_exactly_the_expected_set(self, baseline_mtz): +class TestDefaultLayout: + def test_default_columns_are_exactly_the_expected_set(self, baseline_mtz): mtz, _ = baseline_mtz - assert set(_read(mtz).columns) == BASELINE_COLUMNS + assert set(_read(mtz).columns) == set(DEFAULT_COLUMNS) | FLAG_COLUMNS + + def test_every_default_column_carries_the_right_mtz_type(self, baseline_mtz): + mtz, _ = baseline_mtz + df = _read(mtz) + for name, expected in DEFAULT_COLUMNS.items(): + assert name in df.columns, f"missing column {name}" + assert df.dtypes[name].name == expected, ( + f"{name} written as {df.dtypes[name].name}, expected {expected}" + ) def test_no_two_moment_columns_without_the_flag(self, baseline_mtz): mtz, _ = baseline_mtz present = set(_read(mtz).columns) & set(TWO_MOMENT_COLUMNS) assert present == set(), f"unexpected two-moment columns: {sorted(present)}" + def test_no_gated_columns_without_all_columns(self, baseline_mtz): + mtz, _ = baseline_mtz + present = set(_read(mtz).columns) & set(ALL_COLUMNS_EXTRA) + assert present == set(), f"unexpected gated columns: {sorted(present)}" + + def test_the_difference_map_is_the_weighted_difference_on_dark_phases( + self, baseline_mtz + ): + """``DELFWT`` must be ``(Fo_light - Fo_dark) * w`` with ``w`` the mean-normalised + inverse variance -- the construction ``torchref.validate-ded`` correlates + against. If these two ever diverge, the map in the file stops being the map the + validation reports on, which is how the output drifted from the science before. + """ + import numpy as np + + df = _read(baseline_mtz[0]) + dfo = (df["Fo_light"].to_numpy().astype(float) + - df["Fo_dark"].to_numpy().astype(float)) + sig = np.sqrt(df["SIGFo_dark"].to_numpy().astype(float) ** 2 + + df["SIGFo_light"].to_numpy().astype(float) ** 2) + w = 1 / sig**2 + w = w / w.mean() + + expected = dfo * w + got = df["DELFWT"].to_numpy().astype(float) + scale = max(float(np.abs(expected).max()), 1e-30) + assert np.abs(got - expected).max() / scale < 1e-5 + + # And the phase is the dark model's, not the mixed model's. + assert not np.allclose( + df["PHDELWT"].to_numpy().astype(float), + df["PHIC"].to_numpy().astype(float), + ) + class TestTwoMomentLayout: - def test_baseline_columns_all_survive(self, two_moment_mtz): + def test_default_columns_all_survive(self, two_moment_mtz): mtz, _ = two_moment_mtz - assert BASELINE_COLUMNS.issubset(set(_read(mtz).columns)) + assert set(DEFAULT_COLUMNS).issubset(set(_read(mtz).columns)) def test_every_new_column_is_present_with_the_right_mtz_type(self, two_moment_mtz): mtz, _ = two_moment_mtz @@ -160,26 +231,80 @@ def test_every_new_column_is_present_with_the_right_mtz_type(self, two_moment_mt assert name in df.columns, f"missing column {name}" actual = df.dtypes[name].name assert actual == expected, ( - f"{name} written as {actual}, expected {expected} -- this writer has no " - f"infer_mtz_dtypes() safety net" + f"{name} written as {actual}, expected {expected}" ) - def test_the_column_set_is_exactly_baseline_plus_the_new_ones(self, two_moment_mtz): + def test_the_column_set_is_exactly_default_plus_the_new_ones(self, two_moment_mtz): mtz, _ = two_moment_mtz - assert set(_read(mtz).columns) == BASELINE_COLUMNS | set(TWO_MOMENT_COLUMNS) + assert set(_read(mtz).columns) == ( + set(DEFAULT_COLUMNS) | FLAG_COLUMNS | set(TWO_MOMENT_COLUMNS) + ) + + def test_the_corrected_difference_map_pairs_with_the_same_phases( + self, two_moment_mtz + ): + """``DELFWT_corr`` is the corrected difference on the *same* dark phases, so it + is opened against ``PHDELWT`` and must be built the same way as ``DELFWT``.""" + import numpy as np + + df = _read(two_moment_mtz[0]) + sig = np.sqrt(df["SIGFo_dark"].to_numpy().astype(float) ** 2 + + df["SIGFo_light"].to_numpy().astype(float) ** 2) + w = 1 / sig**2 + w = w / w.mean() + + expected = df["DF_corr"].to_numpy().astype(float) * w + got = df["DELFWT_corr"].to_numpy().astype(float) + scale = max(float(np.abs(expected).max()), 1e-30) + assert np.abs(got - expected).max() / scale < 1e-5 + + +class TestAllColumns: + def test_all_columns_is_a_strict_superset(self, two_moment_mtz, two_moment_all_mtz): + default = set(_read(two_moment_mtz[0]).columns) + full = set(_read(two_moment_all_mtz[0]).columns) + assert default < full, "--all-columns must add columns, never remove any" + + def test_the_gated_columns_are_exactly_the_expected_ones(self, two_moment_all_mtz): + df = _read(two_moment_all_mtz[0]) + assert set(df.columns) == ( + set(DEFAULT_COLUMNS) | FLAG_COLUMNS | set(TWO_MOMENT_COLUMNS) + | set(ALL_COLUMNS_EXTRA) | set(ALL_COLUMNS_TWO_MOMENT_EXTRA) + ) + + def test_every_gated_column_carries_the_right_mtz_type(self, two_moment_all_mtz): + df = _read(two_moment_all_mtz[0]) + expected_types = {**ALL_COLUMNS_EXTRA, **ALL_COLUMNS_TWO_MOMENT_EXTRA} + for name, expected in expected_types.items(): + assert name in df.columns, f"missing column {name}" + assert df.dtypes[name].name == expected, ( + f"{name} written as {df.dtypes[name].name}, expected {expected}" + ) + + def test_no_column_escapes_with_a_plain_numpy_dtype(self, two_moment_all_mtz): + """Every layer declares its columns' MTZ types beside the values, and the writer + refuses a column with none. This is the end-to-end version of that check: a + column reaching the file as a bare numpy dtype is the failure the old parallel + name lists invited. + """ + df = _read(two_moment_all_mtz[0]) + bare = [c for c in df.columns if not hasattr(df.dtypes[c], "mtztype")] + assert bare == [], f"columns written without an MTZ dtype: {bare}" class TestTwoMomentValuesAreConsistent: - def test_ivar_alpha_is_sigma_sq_times_the_squared_difference(self, two_moment_mtz): + def test_ivar_alpha_is_sigma_sq_times_the_squared_difference( + self, two_moment_all_mtz + ): """The variance column must be the quantity it claims, not a rescaling of it.""" import numpy as np - mtz, summary = two_moment_mtz + mtz, summary = two_moment_all_mtz df = _read(mtz) results = json.loads(summary.read_text())["results"] sigma_sq = results["sigma_alpha_sq"] - dfc = df["DFc_complex"].to_numpy().astype(float) + dfc = df["DFc_phased"].to_numpy().astype(float) ivar = df["IVAR_ALPHA"].to_numpy().astype(float) expected = sigma_sq * dfc**2 @@ -187,7 +312,7 @@ def test_ivar_alpha_is_sigma_sq_times_the_squared_difference(self, two_moment_mt assert np.abs(ivar - expected).max() / scale < 1e-5 def test_the_two_moment_intensity_exceeds_the_coherent_one_by_the_variance( - self, two_moment_mtz + self, two_moment_all_mtz ): """``Ic_2mom - Ic_coh`` must equal ``IVAR_ALPHA``, to whatever precision float32 leaves after the cancellation. @@ -204,8 +329,7 @@ def test_the_two_moment_intensity_exceeds_the_coherent_one_by_the_variance( """ import numpy as np - mtz, _ = two_moment_mtz - df = _read(mtz) + df = _read(two_moment_all_mtz[0]) coh = df["Ic_light_coh"].to_numpy().astype(float) two = df["Ic_light_2mom"].to_numpy().astype(float) ivar = df["IVAR_ALPHA"].to_numpy().astype(float) @@ -222,7 +346,7 @@ def test_the_two_moment_intensity_exceeds_the_coherent_one_by_the_variance( # The variance term has no sign: it can only add. assert (two >= coh - 4.0 * floor).all() - def test_the_weight_is_the_contamination_ratio(self, two_moment_mtz): + def test_the_weight_is_the_contamination_ratio(self, two_moment_all_mtz): """``W_2MOM`` must be ``sigma_I**2 / (sigma_I**2 + IVAR_ALPHA)``. Asserted as the formula rather than as a magnitude. On this fixture the weight @@ -234,7 +358,7 @@ def test_the_weight_is_the_contamination_ratio(self, two_moment_mtz): """ import numpy as np - df = _read(two_moment_mtz[0]) + df = _read(two_moment_all_mtz[0]) w = df["W_2MOM"].to_numpy().astype(float) sig = df["SIGIo_light"].to_numpy().astype(float) ivar = df["IVAR_ALPHA"].to_numpy().astype(float) @@ -261,9 +385,13 @@ def test_summary_reports_the_activation_moments(self, two_moment_mtz): _, summary = two_moment_mtz results = json.loads(summary.read_text())["results"] for key in ("alpha_mean", "lambda_twin", "sigma_alpha_sq"): - assert key in results, f"missing summary key: {key}" - assert results["lambda_twin"] == pytest.approx(LAMBDA_TWIN, abs=1e-5) - assert results["alpha_mean"] == pytest.approx(FRACTION, abs=1e-5) - assert results["sigma_alpha_sq"] == pytest.approx( - FRACTION * (1 - FRACTION) * LAMBDA_TWIN, rel=1e-4 - ) + assert key in results, f"summary is missing {key}" + assert results["lambda_twin"] == pytest.approx(LAMBDA_TWIN) + + def test_summary_reports_the_shrinkage_diagnostics(self, two_moment_mtz): + """``tau_sq`` and mean ``w(h)`` say whether the default extrapolated map is + over-shrunk, so they belong in the summary rather than only in a print.""" + _, summary = two_moment_mtz + results = json.loads(summary.read_text())["results"] + assert "tau_sq" in results and "w_shrinkage_mean" in results + assert 0.0 < results["w_shrinkage_mean"] <= 1.0 diff --git a/torchref/cli/__init__.py b/torchref/cli/__init__.py index dbcf9ebe..efba6bad 100644 --- a/torchref/cli/__init__.py +++ b/torchref/cli/__init__.py @@ -5,8 +5,8 @@ __all__ = [ "add_metadata", "collection_difference_refine", + "difference_map", "mtz2map", - "phased_difference_map", "refine", "strip_altlocs", "validate_ded", diff --git a/torchref/cli/_common.py b/torchref/cli/_common.py index da543742..6df82a32 100644 --- a/torchref/cli/_common.py +++ b/torchref/cli/_common.py @@ -313,6 +313,7 @@ def add_dual_model_args( parser: argparse.ArgumentParser, fraction_required: bool = True, fraction_default: Optional[float] = None, + light_model_required: bool = True, ) -> None: """Add the standard dual-model (dark/light) input arguments. @@ -320,6 +321,10 @@ def add_dual_model_args( ``-lm``/``--light-model``, ``-dsf``/``--dark-structure-factor``, ``-lsf``/``--light-structure-factor``, ``--fraction``, ``--cif`` and a *Column selection* group with per-side column flags. + + ``light_model_required=False`` makes ``-lm`` optional, for tools that can do + something useful with the dark model alone -- a weighted difference map needs only + the dark state's phases. Refinement cannot: it refines the light model. """ inp = parser.add_argument_group("Input files") inp.add_argument( @@ -329,13 +334,24 @@ def add_dual_model_args( type=str, help="Dark / reference state model file (PDB or CIF)", ) - inp.add_argument( - "-lm", - "--light-model", - required=True, - type=str, - help="Light / triggered state model file (PDB or CIF)", - ) + if light_model_required: + inp.add_argument( + "-lm", + "--light-model", + required=True, + type=str, + help="Light / triggered state model file (PDB or CIF)", + ) + else: + inp.add_argument( + "-lm", + "--light-model", + type=str, + default=None, + help="Light / triggered state model file (PDB or CIF). Optional: without " + "it only the weighted difference map is written, which needs the dark " + "state's phases and no light-state model at all.", + ) inp.add_argument( "-dsf", "--dark-structure-factor", @@ -366,6 +382,22 @@ def add_dual_model_args( add_dual_column_args(col) +def add_all_columns_arg(parser: argparse.ArgumentParser) -> None: + """Add ``--all-columns`` for the difference MTZ writer. + + Off by default so the output file holds the map a reader wants and can identify. + The gated columns are alternative constructions of the same quantities -- a + model-phased difference, two more extrapolations, the intensity block -- which are + useful once you know which is which and misleading before then. + """ + parser.add_argument( + "--all-columns", action="store_true", default=False, + help="Write every alternative map coefficient and diagnostic column, not just " + "the default difference and extrapolated maps. Costs two further scale " + "fits for the extra extrapolations.", + ) + + def add_output_format_args(parser: argparse.ArgumentParser) -> None: """Add ``--output-format`` argument for coordinate file format.""" parser.add_argument( diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index 91851825..5db698b0 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -7,42 +7,18 @@ parameters (overall scale, anisotropy, bulk solvent k_sol/B_sol) is shared across the dark and light datasets. -Writes refined dark and light models (PDB/CIF), a JSON summary, and a difference MTZ: - -Observed ``Fo_dark``, ``SIGFo_dark``, ``Fo_light``, ``SIGFo_light`` -Differences ``DF`` = F_light - F_dark, ``SIGDF`` (propagated), ``WDF`` (sigma-weighted) -Calculated ``Fc_dark``, ``Fc_light`` (mixed model), ``DFc`` (scalar), ``DFc_complex`` -DED map coeffs ``2mDFop-DFc``, ``mDFop-DFc`` (phase-aware, figure-of-merit weighted) -Phases ``PHIC_dark``, ``PHIC_mixed``, ``PHIC_diff``, ``PHIC_light`` -Extrapolations ``Fextp`` (phase-aware), ``Fextc`` (amplitude-only, no phases), - ``Fextb`` (empirical-Bayes shrinkage toward Fo_dark -- see - :func:`compute_bayes_extrapolated_amplitudes`), each with its sigma - and ``2F-Fc`` / ``F-Fc`` map coefficients - -Under ``--two-moment`` with a non-zero dispersion, thirteen further columns describe the -activation-heterogeneity correction: - -Intensities ``Io_light``, ``SIGIo_light`` (the quantity actually fitted), - ``Ic_light_coh`` = |F(alpha)|^2, ``Ic_light_2mom`` = the fitted model, - ``IVAR_ALPHA`` = sigma_alpha^2 |dF|^2 on its own -Decontaminated ``Fo_light_corr``, ``SIGFo_light_corr``, ``DF_corr``, ``SIGDF_corr``, - and ``2mDFop-DFc_corr`` / ``mDFop-DFc_corr`` to pair with ``PHIC_diff`` -Diagnostics ``DDF`` = DF_corr - DF, and ``W_2MOM``, the sigma_alpha^2-aware weight - -``DDF`` is the one to look at first. Smooth and featureless against resolution means the -correction is collinear with a scale or overall-B error and should be distrusted; -structure in it is the signal. - -Note the ``m`` in ``2mDFop-DFc`` is a normalised inverse-variance weight, not a sigma_A -figure of merit. +Writes refined dark and light models (PDB/CIF), a JSON summary, and a difference MTZ. +See :func:`write_results_mtz` for the columns and why each is where it is -- the table +used to live here, four hundred lines from the function that writes it, which is part of +how the output drifted from its own documentation. Examples -------- :: - torchref.difference-refine \\ - -dm dark.pdb -lm light.pdb \\ - -dsf dark.mtz -lsf light.mtz \\ + torchref.difference-refine \ + -dm dark.pdb -lm light.pdb \ + -dsf dark.mtz -lsf light.mtz \ --fraction 0.37 -o output/ """ @@ -58,6 +34,7 @@ add_dual_model_args, add_dmin_arg, add_general_args, + add_all_columns_arg, add_metadata_args, add_outdir_arg, add_output_format_args, @@ -180,19 +157,67 @@ def setup_scaler(dataset_collection, model_collection, device, verbose=1): return scaler +def setup_dark_only(pdb_dark, dc, cif, d_min, device, verbose, hydrogenate=False): + """Load the dark model alone and scale it against the dark data. + + This is everything a weighted difference map needs. The amplitude is + ``|Fo_light| - |Fo_dark|``, which :meth:`DatasetCollection.scale` has already put on + one scale without reference to any model, and the phase comes from the dark state. + So there is no mixed model, no occupancy fraction, and no joint model-to-data fit -- + the joint fit exists to share scale parameters between two models, and here there is + only one. + + A fitted scaler rather than a bare ``angle(model(hkl))`` because bulk solvent + contributes a phase at low resolution, which is exactly where difference density is + largest. + + Returns + ------- + tuple + ``(model_dark, scaler)`` -- ready to hand to :func:`write_results_mtz` with + ``mc=None``. + """ + from torchref import Scaler + + model_dark = load_model( + pdb_dark, max_res=d_min, device=device, verbose=verbose, cif=cif, + ) + if hydrogenate: + if verbose > 0: + print("Adding hydrogens...") + sys.stdout.flush() + model_dark = model_dark.hydrogenate(verbose=max(0, verbose - 1)) + model_dark.exclude_H_from_sf = True + + scaler = Scaler( + model_dark, dc["dark"], device=device, verbose=max(-1, verbose - 1), + ) + scaler.initialize().refine_lbfgs() + return model_dark, scaler + + def compute_rfactors(model, data, scaler): - """Compute R-work/R-free using forward_mixed for proper solvent. + """Compute R-work/R-free with the scaler's own solvent model applied. Routes through ``rfactor_work_free`` — the shared source of truth used by the refinement targets — so the validity mask is applied and the validation set is excluded from both work and free, matching every other reported R-factor. + + Takes either scaler kind; see the branch below. """ from torchref.base.metrics.rfactor import rfactor_work_free with torch.no_grad(): hkl = data.hkl fcalc = model(hkl) - fcalc_scaled = scaler.forward_mixed(fcalc, model.fractions) + # Which scaler this is decides how it is called: a CollectionScaler mixes + # components and needs the model's fractions, a single-dataset Scaler takes + # the structure factors straight. Asked of the scaler, not inferred from the + # caller, so the dark-only path needs no special case. + if hasattr(scaler, "forward_mixed"): + fcalc_scaled = scaler.forward_mixed(fcalc, model.fractions) + else: + fcalc_scaled = scaler(fcalc) return rfactor_work_free(data, torch.abs(fcalc_scaled)) @@ -264,7 +289,7 @@ def setup_loss_state(dataset_collection, model_collection, scaler, def compute_bayes_extrapolated_amplitudes( - Fobs_dark, Fobs_light, sig_dark, sig_light, phi_dark, phi_mixed, f, + Fobs_dark, Fobs_light, sig_ext, phi_dark, phi_mixed, f, *, tau_sq_floor=1e-4, ): """Empirical Bayes shrinkage estimator for extrapolated SF amplitudes. @@ -274,7 +299,6 @@ def compute_bayes_extrapolated_amplitudes( Fo_dark, regularising noisy high-resolution and weakly-measured reflections:: F_ext = |F_dark*e^(iφ_d) + ΔF/f| (phase-aware amplitude) - σ_ext² = (σ_light² + σ_dark²) / f² τ² = max(<(F_ext - Fo_dark)²> - <σ_ext²>, floor) w(h) = τ² / (τ² + σ_ext²(h)) F_extb = w(h)·F_ext + (1-w(h))·Fo_dark (amplitude shrinkage) @@ -283,8 +307,13 @@ def compute_bayes_extrapolated_amplitudes( ---------- Fobs_dark, Fobs_light : Tensor (N,) Observed amplitudes. - sig_dark, sig_light : Tensor (N,) - Measurement uncertainties. + sig_ext : Tensor (N,) + Propagated uncertainty of the extrapolated amplitude. Taken from the caller + rather than rebuilt here: ``F_ext`` is linear in the observations with + ``dF_ext/dF_light = 1/f`` and ``dF_ext/dF_dark = 1 - 1/f = -(1-f)/f``, so the + dark term carries a ``(1-f)**2`` weight that is easy to drop. This function + used to drop it, over-weighting the dark term by ``1/(1-f)**2`` -- 1.64x at + f = 0.22 -- which biased τ² low and over-shrank every reflection. phi_dark, phi_mixed : Tensor (N,) Calculated phases (radians) for the dark and mixed models. f : float or Tensor @@ -295,17 +324,15 @@ def compute_bayes_extrapolated_amplitudes( Returns ------- tuple - ``(F_ext_bayes, var_ext_bayes, w_shrinkage, tau_sq)`` -- extrapolated - amplitudes **before** shrinkage, posterior variance and shrinkage weight per + ``(F_ext_bayes, var_ext_bayes, w_shrinkage, tau_sq)`` -- the **shrunk** + extrapolated amplitude, its posterior variance and the shrinkage weight per reflection, and the global τ² as a float. """ F_dark_phased = Fobs_dark * torch.exp(1j * phi_dark) F_light_phased = Fobs_light * torch.exp(1j * phi_mixed) delta_F = F_light_phased - F_dark_phased - # Propagated variance - sig_sq_dF = sig_light**2 + sig_dark**2 - sig_sq_ext = sig_sq_dF / f**2 + sig_sq_ext = sig_ext**2 # Phase-aware extrapolated amplitude F_ext_complex = F_dark_phased + delta_F / f @@ -321,12 +348,15 @@ def compute_bayes_extrapolated_amplitudes( # Posterior variance var_ext_bayes = (tau_sq * sig_sq_ext) / (tau_sq + sig_sq_ext) - return F_ext, var_ext_bayes, w, tau_sq + # Shrink the amplitude toward Fo_dark -- scalar, so no phase interference. + F_ext_bayes = w * F_ext + (1 - w) * Fobs_dark + + return F_ext_bayes, var_ext_bayes, w, tau_sq def _two_moment_columns(mc, dc, mask, fcalc_dark_full, fcalc_mixed_full, *, weights, diff_Fobs, Fcalc_diff_amp, Fobs_dark, - sig_dark, phi_mixed, F_obs_dark_phased): + sig_dark, phi_mixed, F_obs_dark_phased, all_columns=False): """Two-moment diagnostic columns, or empty dicts when the model is off. The observed light intensity carries a positive, phase-blind contamination @@ -356,13 +386,13 @@ def _two_moment_columns(mc, dc, mask, fcalc_dark_full, fcalc_mixed_full, Returns ------- tuple - ``(columns, f_cols, sigma_cols, intensity_cols, weight_cols)`` -- the values plus - the MTZ type each belongs to. This writer assigns types from name lists and never - calls ``infer_mtz_dtypes``, so every column must appear in exactly one list. + ``(columns, types)`` -- the values, and the MTZ type letter for each. Carrying + the type beside the value is what stops a column reaching the file with whatever + dtype numpy produced, which is the failure the old parallel name lists invited. """ import numpy as np - empty = ({}, [], [], [], []) + empty = ({}, {}) if float(mc.sigma_alpha_sq) == 0.0: return empty @@ -430,278 +460,434 @@ def _np(t): w_two_moment = sig_I_light**2 / np.maximum(sig_I_light**2 + variance, 1e-12) columns = { - "Io_light": I_light, - "SIGIo_light": sig_I_light, - "Ic_light_coh": I_coherent, - "Ic_light_2mom": I_two_moment, - "IVAR_ALPHA": variance, + # The corrected difference map, on the same dark phases as DELFWT. + "DELFWT_corr": DF_corr * weights, "Fo_light_corr": F_corr, "SIGFo_light_corr": sig_F_corr, "DF_corr": DF_corr, "SIGDF_corr": sig_DF_corr, - "2mDFop-DFc_corr": amp_2_corr, - "mDFop-DFc_corr": amp_1_corr, "DDF": DDF, - "W_2MOM": w_two_moment, } - f_cols = [ - "Fo_light_corr", "DF_corr", "2mDFop-DFc_corr", "mDFop-DFc_corr", "DDF", - ] - sigma_cols = ["SIGIo_light", "SIGFo_light_corr", "SIGDF_corr"] - intensity_cols = ["Io_light", "Ic_light_coh", "Ic_light_2mom", "IVAR_ALPHA"] - weight_cols = ["W_2MOM"] - return columns, f_cols, sigma_cols, intensity_cols, weight_cols + types = { + "DELFWT_corr": "F", + "Fo_light_corr": "F", + "SIGFo_light_corr": "Q", + "DF_corr": "F", + "SIGDF_corr": "Q", + "DDF": "F", + } + if all_columns: + columns.update({ + "Io_light": I_light, + "SIGIo_light": sig_I_light, + "Ic_light_coh": I_coherent, + "Ic_light_2mom": I_two_moment, + "IVAR_ALPHA": variance, + "W_2MOM": w_two_moment, + "2mDFop-DFc_corr": amp_2_corr, + "mDFop-DFc_corr": amp_1_corr, + }) + types.update({ + "Io_light": "J", + "SIGIo_light": "Q", + "Ic_light_coh": "J", + "Ic_light_2mom": "J", + "IVAR_ALPHA": "J", + "W_2MOM": "W", + "2mDFop-DFc_corr": "F", + "mDFop-DFc_corr": "F", + }) + return columns, types + + +def _difference_columns(data_dark, data_light, mask, hkl_np, *, Fobs_dark, sig_dark, + Fobs_light, sig_light, Fcalc_dark, phases_dark, diff_Fobs, + sig_diff, weights): + """The weighted difference map, and the observations behind it. + + ``DELFWT``/``PHDELWT`` is the inverse-variance-weighted amplitude difference carried + on the **dark** model's phases -- the isomorphous difference Fourier, and the same + construction ``torchref.validate-ded`` correlates against, so the map in this file + and the map the validation reports are one object. CCP4 and Coot recognise the names + and open it as a difference map without being told which columns to use. + + This layer needs no light-state model: the amplitude is ``|Fo_light| - |Fo_dark|`` + and the phase comes from the dark model. Keeping the light state's model out is the + point -- a phased construction puts its phases into the observed amplitude, biasing + the map toward the very model the experiment is testing. + """ + import numpy as np + n = len(hkl_np) -def write_results_mtz(dc, mc, scaler, filename): - """Write difference / extrapolated map coefficients to an MTZ file. + def _flags(data): + if data.rfree_flags is None: + return np.ones(n, dtype=int) + return data.rfree_flags[mask].cpu().numpy().astype(int) - Parameters - ---------- - dc : DatasetCollection - mc : ModelCollection - scaler : CollectionScaler - filename : str - Output MTZ path. - """ - import numpy as np - import reciprocalspaceship as rs - from torchref import ReflectionData, Scaler + columns = { + "H": hkl_np[:, 0], "K": hkl_np[:, 1], "L": hkl_np[:, 2], + "Fo_dark": Fobs_dark, "SIGFo_dark": sig_dark, + "Fo_light": Fobs_light, "SIGFo_light": sig_light, + "DF": diff_Fobs, "SIGDF": sig_diff, + "DELFWT": diff_Fobs * weights, "PHDELWT": phases_dark, + "Fc_dark": Fcalc_dark, + # 1 = work, 0 = free. Both are kept: the two datasets can disagree, and + # picking one would silently report an R-free against the wrong test set. + "FreeR_flag_dark": _flags(data_dark), + "FreeR_flag_light": _flags(data_light), + } + types = { + "H": "H", "K": "H", "L": "H", + "Fo_dark": "F", "SIGFo_dark": "Q", + "Fo_light": "F", "SIGFo_light": "Q", + "DF": "F", "SIGDF": "Q", + "DELFWT": "F", "PHDELWT": "P", + "Fc_dark": "F", + "FreeR_flag_dark": "I", "FreeR_flag_light": "I", + } + return columns, types - data_dark = dc[mc.dark_key] - data_light = dc["light"] - dark_model = mc.dark_model - mixed_model = mc["light"] - model_light = mc.base_models[1] - hkl_all = data_dark.hkl - Fobs_dark_full, sig_dark_full = data_dark.get_corrected_data() - Fobs_light_full, sig_light_full = data_light.get_corrected_data() +def _phasing_columns(mc, scaler, hkl_all, mask, *, fcalc_dark, Fobs_dark_vals, + Fobs_light_vals, phi_dark, Fcalc_dark, weights, + all_columns=False): + """Mixed-model amplitude and phase, and the phased difference residuals. - mask = data_dark.masks().to(torch.bool) & data_light.masks().to(torch.bool) - hkl = hkl_all[mask] - Fobs_dark_vals = Fobs_dark_full[mask] - Fobs_light_vals = Fobs_light_full[mask] - sig_dark_vals = sig_dark_full[mask] - sig_light_vals = sig_light_full[mask] - rfree_flags_masked = ( - data_light.rfree_flags[mask] if data_light.rfree_flags is not None else None - ) + Everything here needs the light state's model. ``FC``/``PHIC`` are the mixed model's + scaled amplitude and phase. - fractions = mixed_model.fractions.detach() - w_dark = fractions[0] - w_light = fractions[1] + Under ``all_columns`` the phased difference residual coefficients come too -- + ``(|Fo_light e^{i phi_mixed} - Fo_dark e^{i phi_dark}| - |dFc|) * w`` on + ``PHIC_diff``. These are a *different object* from the plain difference Fourier in + ``DELFWT``, not a refinement of it: the light state's model phases enter the observed + amplitude, so they are model-biased where ``DELFWT`` is not. They are kept because + they are informative once that is understood, and gated because the name alone does + not say it. + + Returns ``(columns, types, ctx)``. ``ctx`` carries the intermediates the + extrapolation and two-moment layers need, so nothing is computed twice. + """ + mixed_model = mc["light"] - # Compute Fcalc on full HKL then mask (scalers fitted on full datasets) with torch.no_grad(): - fcalc_dark_full = scaler.forward_mixed( - dark_model(hkl_all), dark_model.fractions - ) fcalc_mixed_full = scaler.forward_mixed( mixed_model(hkl_all), mixed_model.fractions ) - fcalc_dark = fcalc_dark_full[mask] fcalc_mixed = fcalc_mixed_full[mask] fcalc_diff = fcalc_mixed - fcalc_dark - phi_dark = torch.angle(fcalc_dark) phi_mixed = torch.angle(fcalc_mixed) - F_obs_dark_phased = Fobs_dark_vals * torch.exp(1j * phi_dark) F_obs_light_phased = Fobs_light_vals * torch.exp(1j * phi_mixed) - # --- Phase-aware extrapolation --- - F_light_extra = (F_obs_light_phased - w_dark * F_obs_dark_phased) / w_light - amp_light_extra = torch.abs(F_light_extra) - sig_light_extra = torch.sqrt(sig_light_vals**2 + w_dark**2 * sig_dark_vals**2) / w_light + Fcalc_light = torch.abs(fcalc_mixed).detach().cpu().numpy() + Fcalc_diff_amp = torch.abs(fcalc_diff).detach().cpu().numpy() + + columns = { + "FC": Fcalc_light, + "PHIC": phi_mixed.detach().rad2deg().cpu().numpy(), + } + types = {"FC": "F", "PHIC": "P"} + + if all_columns: + Fobs_diff_phased = torch.abs( + F_obs_light_phased - F_obs_dark_phased + ).detach().cpu().numpy() + columns.update({ + "2mDFop-DFc": (2 * Fobs_diff_phased - Fcalc_diff_amp) * weights, + "mDFop-DFc": (Fobs_diff_phased - Fcalc_diff_amp) * weights, + "PHIC_diff": torch.angle(fcalc_diff).detach().rad2deg().cpu().numpy(), + "DFc": Fcalc_light - Fcalc_dark, + # The modulus of the complex vector difference. Named ``_phased`` rather + # than ``_complex``: the column holds a real amplitude, and the old name + # said otherwise. + "DFc_phased": Fcalc_diff_amp, + }) + types.update({ + "2mDFop-DFc": "F", "mDFop-DFc": "F", "PHIC_diff": "P", + "DFc": "F", "DFc_phased": "F", + }) + + ctx = { + "fcalc_mixed_full": fcalc_mixed_full, + "phi_mixed": phi_mixed, + "F_obs_dark_phased": F_obs_dark_phased, + "F_obs_light_phased": F_obs_light_phased, + "Fcalc_diff_amp": Fcalc_diff_amp, + } + return columns, types, ctx - data_light_extra = ReflectionData.from_tensors( - hkl=hkl, F=amp_light_extra, F_sigma=sig_light_extra, - cell=data_light.cell, spacegroup=data_light.spacegroup, - rfree_flags=rfree_flags_masked, device=str(hkl.device), verbose=0, - ) - scaler_extra = Scaler(model_light, data_light_extra, device=hkl.device, verbose=-1) - scaler_extra.initialize().refine_lbfgs() - F_calc_extra = scaler_extra(model_light(hkl)) - - amp_extra = torch.abs(F_light_extra) - amp_light_calc = torch.abs(F_calc_extra) - phi_light_calc = torch.angle(F_calc_extra) - amp_2fofc_light = 2 * amp_extra - amp_light_calc - amp_fextfc = amp_extra - amp_light_calc - - # --- Classic (amplitude-only) extrapolation --- - amp_extra_classic = (Fobs_light_vals - w_dark * Fobs_dark_vals) / w_light - sig_extra_classic = sig_light_extra - - data_extra_classic = ReflectionData.from_tensors( - hkl=hkl, F=amp_extra_classic, F_sigma=sig_extra_classic, - cell=data_light.cell, spacegroup=data_light.spacegroup, - rfree_flags=rfree_flags_masked, device=str(hkl.device), verbose=0, - ) - scaler_classic = Scaler(model_light, data_extra_classic, device=hkl.device, verbose=-1) - scaler_classic.initialize().refine_lbfgs() - F_calc_classic = scaler_classic(model_light(hkl)) - amp_light_calc_classic = torch.abs(F_calc_classic) - amp_2fofc_classic = 2 * amp_extra_classic - amp_light_calc_classic - amp_fofc_classic = amp_extra_classic - amp_light_calc_classic +def _extrapolation_columns(mc, dc, hkl, *, Fobs_dark_vals, Fobs_light_vals, + sig_dark_vals, sig_light_vals, phi_dark, ctx, + rfree_flags_masked, all_columns=False, verbose=1): + """Extrapolated light-state amplitudes and the map to refine against. - # --- Empirical Bayes extrapolation (amplitude-only shrinkage) --- - F_ext_bayes, var_ext_bayes, w_shrinkage, tau_sq = ( + Three constructions of the same quantity, all needing the light model: + + ``FEXT`` (default, Bayes-shrunk) + The phase-aware amplitude shrunk toward ``Fo_dark`` by a per-reflection weight + ``w(h) = tau^2 / (tau^2 + sigma_ext^2(h))``, which quiets the weak and + high-resolution reflections where the extrapolation is noisiest. + ``FEXT_PHASED`` (``all_columns``) + The unshrunk phase-aware amplitude. + ``FEXT_SCALAR`` (``all_columns``) + ``(Fo_light - w_dark * Fo_dark) / w_light`` on amplitudes only. + + Each needs ``F_calc`` rescaled against *its own* amplitudes -- the three sets differ + in overall scale -- so each costs one LBFGS scale fit. Only the default one runs + unless ``all_columns`` is set. + + ``FWT``/``PHWT`` is ``2 * FEXT - Fc`` with the phase from the Bayes fit. The phase + matters: the scaler contributes one through ``f_sol``, so the three fits do not agree + and pairing these coefficients with another fit's phase would be wrong. + """ + from torchref import ReflectionData, Scaler + from torchref.base.metrics.rfactor import rfactor_work_free + + data_light = dc["light"] + fractions = mc["light"].fractions.detach() + w_dark, w_light = fractions[0], fractions[1] + + def _fit(amp, sig): + """Rescale the light model against one set of extrapolated amplitudes.""" + data = ReflectionData.from_tensors( + hkl=hkl, F=amp, F_sigma=sig, + cell=data_light.cell, spacegroup=data_light.spacegroup, + rfree_flags=rfree_flags_masked, device=str(hkl.device), verbose=0, + ) + sc = Scaler(mc.base_models[1], data, device=hkl.device, verbose=-1) + sc.initialize().refine_lbfgs() + return data, sc(mc.base_models[1](hkl)) + + # The phase-aware amplitude, and the one propagated sigma for this extrapolation. + F_light_extra = ( + ctx["F_obs_light_phased"] - w_dark * ctx["F_obs_dark_phased"] + ) / w_light + sig_light_extra = torch.sqrt( + sig_light_vals**2 + w_dark**2 * sig_dark_vals**2 + ) / w_light + + F_ext_bayes_amp, var_ext_bayes, w_shrinkage, tau_sq = ( compute_bayes_extrapolated_amplitudes( - Fobs_dark_vals, Fobs_light_vals, - sig_dark_vals, sig_light_vals, - phi_dark, phi_mixed, w_light, + Fobs_dark_vals, Fobs_light_vals, sig_light_extra, + phi_dark, ctx["phi_mixed"], w_light, ) ) sig_ext_bayes = torch.sqrt(var_ext_bayes) - # Shrink |F_ext| towards |Fo_dark| — scalar operation, no phase interference - F_ext_amp_only = torch.abs(F_light_extra) - F_ext_bayes_amp = w_shrinkage * F_ext_amp_only + (1 - w_shrinkage) * Fobs_dark_vals + data_bayes, F_calc_bayes = _fit(F_ext_bayes_amp, sig_ext_bayes) + amp_calc_bayes = torch.abs(F_calc_bayes) - data_extra_bayes = ReflectionData.from_tensors( - hkl=hkl, F=F_ext_bayes_amp, F_sigma=sig_ext_bayes, - cell=data_light.cell, spacegroup=data_light.spacegroup, - rfree_flags=rfree_flags_masked, device=str(hkl.device), verbose=0, - ) - scaler_bayes = Scaler(model_light, data_extra_bayes, device=hkl.device, verbose=-1) - scaler_bayes.initialize().refine_lbfgs() - F_calc_bayes = scaler_bayes(model_light(hkl)) + def _np(t): + return t.detach().cpu().numpy() - amp_calc_bayes = torch.abs(F_calc_bayes) - amp_2fofc_bayes = 2 * F_ext_bayes_amp - amp_calc_bayes - amp_fofc_bayes = F_ext_bayes_amp - amp_calc_bayes + columns = { + "FEXT": _np(F_ext_bayes_amp), + "SIGFEXT": _np(sig_ext_bayes), + "FWT": _np(2 * F_ext_bayes_amp - amp_calc_bayes), + "PHWT": _np(torch.angle(F_calc_bayes).rad2deg()), + } + types = {"FEXT": "F", "SIGFEXT": "Q", "FWT": "F", "PHWT": "P"} + + if verbose > 0: + print(" Bayes extrapolation rfactors:", + rfactor_work_free(data_bayes, amp_calc_bayes)) + print(f" Bayes: tau^2 = {tau_sq:.4f}, " + f"mean w(h) = {w_shrinkage.mean().item():.3f}") + + if all_columns: + amp_phased = torch.abs(F_light_extra) + data_phased, F_calc_phased = _fit(amp_phased, sig_light_extra) + amp_calc_phased = torch.abs(F_calc_phased) + + amp_scalar = (Fobs_light_vals - w_dark * Fobs_dark_vals) / w_light + data_scalar, F_calc_scalar = _fit(amp_scalar, sig_light_extra) + amp_calc_scalar = torch.abs(F_calc_scalar) + + columns.update({ + "FEXT_PHASED": _np(amp_phased), + # Computed all along and never written, though the docs claimed it. + "SIGFEXT_PHASED": _np(sig_light_extra), + "2FEXT_PHASED-Fc": _np(2 * amp_phased - amp_calc_phased), + "FEXT_PHASED-Fc": _np(amp_phased - amp_calc_phased), + "PHFEXT_PHASED": _np(torch.angle(F_calc_phased).rad2deg()), + "FEXT_SCALAR": _np(amp_scalar), + "SIGFEXT_SCALAR": _np(sig_light_extra), + "2FEXT_SCALAR-Fc": _np(2 * amp_scalar - amp_calc_scalar), + "FEXT_SCALAR-Fc": _np(amp_scalar - amp_calc_scalar), + "PHFEXT_SCALAR": _np(torch.angle(F_calc_scalar).rad2deg()), + }) + types.update({ + "FEXT_PHASED": "F", "SIGFEXT_PHASED": "Q", + "2FEXT_PHASED-Fc": "F", "FEXT_PHASED-Fc": "F", + "PHFEXT_PHASED": "P", + "FEXT_SCALAR": "F", "SIGFEXT_SCALAR": "Q", + "2FEXT_SCALAR-Fc": "F", "FEXT_SCALAR-Fc": "F", + "PHFEXT_SCALAR": "P", + }) + if verbose > 0: + print(" Phase-aware extrapolation rfactors:", + rfactor_work_free(data_phased, amp_calc_phased)) + print(" Scalar extrapolation rfactors:", + rfactor_work_free(data_scalar, amp_calc_scalar)) + + diagnostics = { + "tau_sq": float(tau_sq), + "w_shrinkage_mean": float(w_shrinkage.mean().item()), + } + return columns, types, diagnostics - from torchref.base.metrics.rfactor import rfactor_work_free - def _extrapolation_rfactors(data, fcalc_scaled): - return rfactor_work_free(data, torch.abs(fcalc_scaled)) +def write_results_mtz(dc, dark_model, scaler, filename, *, mc=None, + all_columns=False, verbose=1): + """Write the difference map, and map coefficients when a light model is given. + + The default output is the **weighted difference map**: ``DELFWT``/``PHDELWT``, the + inverse-variance-weighted amplitude difference on the dark model's phases. That needs + no light-state model, which is why ``mc`` is optional -- with a dark model alone this + writes a difference map and nothing else, and no scale fit is run beyond the one that + produced ``scaler``. - print("Phase-aware extrapolation rfactors:", - _extrapolation_rfactors(data_light_extra, F_calc_extra)) - print("Classic extrapolation rfactors:", - _extrapolation_rfactors(data_extra_classic, F_calc_classic)) - print("Bayes extrapolation rfactors:", - _extrapolation_rfactors(data_extra_bayes, F_calc_bayes)) - print(f" Bayes: tau^2 = {tau_sq:.4f}, mean w(h) = {w_shrinkage.mean().item():.3f}") + Given ``mc``, the layers that need the light state follow: its amplitude and phase, + the extrapolated amplitudes, and the two-moment correction. ``all_columns`` adds the + alternatives within each layer -- see :func:`_phasing_columns` and + :func:`_extrapolation_columns` for what each contains and why it is gated. + + Parameters + ---------- + dc : DatasetCollection + Dark and light data, already inter-scaled by :meth:`DatasetCollection.scale`. + dark_model : Model or _SharedMixedModel + Supplies ``Fc_dark`` and the phases the difference map is carried on. + scaler : Scaler or CollectionScaler + Scales ``dark_model`` against the dark data. A ``CollectionScaler`` when ``mc`` + is given, a single-dataset ``Scaler`` otherwise. + mc : ModelCollection, optional + The dark+light collection. Absent means difference map only. + filename : str + Output MTZ path. + + Returns + ------- + dict + Diagnostics worth recording outside the file -- currently the Bayes shrinkage's + ``tau_sq`` and mean ``w(h)``, which say whether the default extrapolated map is + over-shrunk. Empty when no light model was given. + """ + import reciprocalspaceship as rs + + data_dark = dc[mc.dark_key] if mc is not None else dc["dark"] + data_light = dc["light"] + + hkl_all = data_dark.hkl + Fobs_dark_full, sig_dark_full = data_dark.get_corrected_data() + Fobs_light_full, sig_light_full = data_light.get_corrected_data() + + mask = data_dark.masks().to(torch.bool) & data_light.masks().to(torch.bool) + hkl = hkl_all[mask] + Fobs_dark_vals = Fobs_dark_full[mask] + Fobs_light_vals = Fobs_light_full[mask] + sig_dark_vals = sig_dark_full[mask] + sig_light_vals = sig_light_full[mask] + rfree_flags_masked = ( + data_light.rfree_flags[mask] if data_light.rfree_flags is not None else None + ) + + # Fcalc on the full HKL list then masked, because the scalers were fitted on the + # full datasets. ``forward_mixed`` exists only on CollectionScaler; the dark-only + # path carries a single-dataset Scaler, whose ``forward`` is the plain call. + with torch.no_grad(): + if hasattr(scaler, "forward_mixed"): + fcalc_dark_full = scaler.forward_mixed( + dark_model(hkl_all), dark_model.fractions + ) + else: + fcalc_dark_full = scaler(dark_model(hkl_all)) + fcalc_dark = fcalc_dark_full[mask] + + phi_dark = torch.angle(fcalc_dark) - # --- Build MTZ --- hkl_np = hkl.cpu().numpy() Fobs_dark = Fobs_dark_vals.cpu().numpy() Fobs_light = Fobs_light_vals.cpu().numpy() sig_dark = sig_dark_vals.cpu().numpy() sig_light = sig_light_vals.cpu().numpy() - Fcalc_dark = torch.abs(fcalc_dark).detach().cpu().numpy() - Fcalc_light = torch.abs(fcalc_mixed).detach().cpu().numpy() - phases_dark = torch.angle(fcalc_dark).detach().rad2deg().cpu().numpy() - phases_mixed = torch.angle(fcalc_mixed).detach().rad2deg().cpu().numpy() - - Fcalc_diff_amp = torch.abs(fcalc_diff).detach().cpu().numpy() - Fcalc_diff_scalar = Fcalc_light - Fcalc_dark - phases_diff = torch.angle(fcalc_diff).detach().rad2deg().cpu().numpy() - - Fobs_diff_phased = torch.abs( - F_obs_light_phased - F_obs_dark_phased - ).detach().cpu().numpy() + phases_dark = phi_dark.detach().rad2deg().cpu().numpy() diff_Fobs = Fobs_light - Fobs_dark sig_diff = (sig_dark**2 + sig_light**2) ** 0.5 weights = 1 / sig_diff**2 weights = weights / weights.mean() - weighted_diff_Fobs = diff_Fobs * weights - amp_2DFoDFc = (2 * Fobs_diff_phased - Fcalc_diff_amp) * weights - amp_DFoDFc = (Fobs_diff_phased - Fcalc_diff_amp) * weights + columns, types = _difference_columns( + data_dark, data_light, mask, hkl_np, + Fobs_dark=Fobs_dark, sig_dark=sig_dark, + Fobs_light=Fobs_light, sig_light=sig_light, + Fcalc_dark=Fcalc_dark, phases_dark=phases_dark, + diff_Fobs=diff_Fobs, sig_diff=sig_diff, weights=weights, + ) + + diagnostics = {} + if mc is not None: + phase_cols, phase_types, ctx = _phasing_columns( + mc, scaler, hkl_all, mask, + fcalc_dark=fcalc_dark, Fobs_dark_vals=Fobs_dark_vals, + Fobs_light_vals=Fobs_light_vals, phi_dark=phi_dark, + Fcalc_dark=Fcalc_dark, weights=weights, all_columns=all_columns, + ) + columns.update(phase_cols) + types.update(phase_types) + + ext_cols, ext_types, diagnostics = _extrapolation_columns( + mc, dc, hkl, + Fobs_dark_vals=Fobs_dark_vals, Fobs_light_vals=Fobs_light_vals, + sig_dark_vals=sig_dark_vals, sig_light_vals=sig_light_vals, + phi_dark=phi_dark, ctx=ctx, rfree_flags_masked=rfree_flags_masked, + all_columns=all_columns, verbose=verbose, + ) + columns.update(ext_cols) + types.update(ext_types) - two_moment_columns, two_moment_f, two_moment_sig, two_moment_j, two_moment_w = ( - _two_moment_columns( - mc, dc, mask, fcalc_dark_full, fcalc_mixed_full, + tm_cols, tm_types = _two_moment_columns( + mc, dc, mask, fcalc_dark_full, ctx["fcalc_mixed_full"], weights=weights, diff_Fobs=diff_Fobs, - Fcalc_diff_amp=Fcalc_diff_amp, Fobs_dark=Fobs_dark, - sig_dark=sig_dark, phi_mixed=phi_mixed, - F_obs_dark_phased=F_obs_dark_phased, + Fcalc_diff_amp=ctx["Fcalc_diff_amp"], Fobs_dark=Fobs_dark, + sig_dark=sig_dark, phi_mixed=ctx["phi_mixed"], + F_obs_dark_phased=ctx["F_obs_dark_phased"], + all_columns=all_columns, ) - ) + columns.update(tm_cols) + types.update(tm_types) df = rs.DataSet( - { - "H": hkl_np[:, 0], "K": hkl_np[:, 1], "L": hkl_np[:, 2], - # Observed - "Fo_dark": Fobs_dark, "SIGFo_dark": sig_dark, - "Fo_light": Fobs_light, "SIGFo_light": sig_light, - # Differences - "DF": diff_Fobs, "SIGDF": sig_diff, "WDF": weighted_diff_Fobs, - # Calculated - "Fc_dark": Fcalc_dark, "Fc_light": Fcalc_light, - "DFc": Fcalc_diff_scalar, "DFc_complex": Fcalc_diff_amp, - # DED map coefficients (phase-aware) - "2mDFop-DFc": amp_2DFoDFc, "mDFop-DFc": amp_DFoDFc, - # Phases - "PHIC_dark": phases_dark, "PHIC_mixed": phases_mixed, - "PHIC_diff": phases_diff, - "PHIC_light": phi_light_calc.detach().rad2deg().cpu().numpy(), - # Phase-aware extrapolation - "Fextp": amp_extra.detach().cpu().numpy(), - "2Fextp-Fc": amp_2fofc_light.detach().cpu().numpy(), - "Fextp-Fc": amp_fextfc.detach().cpu().numpy(), - # Classic extrapolation - "Fextc": amp_extra_classic.detach().cpu().numpy(), - "SIGFextc": sig_extra_classic.detach().cpu().numpy(), - "2Fextc-Fc": amp_2fofc_classic.detach().cpu().numpy(), - "Fextc-Fc": amp_fofc_classic.detach().cpu().numpy(), - # Bayes extrapolation (amplitude-only shrinkage) - "Fextb": F_ext_bayes_amp.detach().cpu().numpy(), - "SIGFextb": sig_ext_bayes.detach().cpu().numpy(), - "2Fextb-Fc": amp_2fofc_bayes.detach().cpu().numpy(), - "Fextb-Fc": amp_fofc_bayes.detach().cpu().numpy(), - **two_moment_columns, - # R-free flags (1=work, 0=free) - "FreeR_flag_dark": ( - data_dark.rfree_flags[mask].cpu().numpy().astype(int) - if data_dark.rfree_flags is not None - else np.ones(len(hkl_np), dtype=int) - ), - "FreeR_flag_light": ( - data_light.rfree_flags[mask].cpu().numpy().astype(int) - if data_light.rfree_flags is not None - else np.ones(len(hkl_np), dtype=int) - ), - }, + columns, cell=data_dark.cell.data.cpu().tolist(), spacegroup=data_dark.spacegroup.hm, ) - - df[["H", "K", "L"]] = df[["H", "K", "L"]].astype("H") - f_cols = [ - "Fo_dark", "Fo_light", "DF", "WDF", - "Fc_dark", "Fc_light", "DFc", "DFc_complex", - "2mDFop-DFc", "mDFop-DFc", - "Fextp", "2Fextp-Fc", "Fextp-Fc", - "Fextc", "2Fextc-Fc", "Fextc-Fc", - "Fextb", "2Fextb-Fc", "Fextb-Fc", - ] - f_cols += two_moment_f - df[f_cols] = df[f_cols].astype("F") - sig_cols = ["SIGFo_dark", "SIGFo_light", "SIGDF", "SIGFextc", "SIGFextb"] - sig_cols += two_moment_sig - df[sig_cols] = df[sig_cols].astype("Q") - # This writer never calls infer_mtz_dtypes(), so a column absent from every list - # above would be written with whatever dtype numpy produced. - if two_moment_j: - df[two_moment_j] = df[two_moment_j].astype("J") - if two_moment_w: - df[two_moment_w] = df[two_moment_w].astype("W") - phase_cols = ["PHIC_dark", "PHIC_mixed", "PHIC_diff", "PHIC_light"] - df[phase_cols] = df[phase_cols].astype("P") - df["FreeR_flag_dark"] = df["FreeR_flag_dark"].astype("I") - df["FreeR_flag_light"] = df["FreeR_flag_light"].astype("I") + # Every layer returns its columns' MTZ types beside the values, so a new column + # cannot reach the file with whatever dtype numpy produced -- the failure the old + # parallel name lists invited. ``infer_mtz_dtypes`` is then the same safety net the + # canonical writer in ``torchref/io/mtz.py`` uses. + missing = set(columns) - set(types) + if missing: + raise AssertionError(f"columns with no declared MTZ type: {sorted(missing)}") + for name, letter in types.items(): + df[name] = df[name].astype(letter) + df = df.infer_mtz_dtypes() df.set_index(["H", "K", "L"], inplace=True) df.write_mtz(filename) - print(f" Results MTZ written to {filename}") - print(f" w_dark={w_dark.item():.3f}, w_light={w_light.item():.3f}") + + if verbose > 0: + print(f" Results MTZ written to {filename} ({len(columns)} columns)") + if mc is not None: + fractions = mc["light"].fractions.detach() + print(f" w_dark={fractions[0].item():.3f}, " + f"w_light={fractions[1].item():.3f}") + + return diagnostics def optimize_lbfgs(state, parameters, max_iter, nsteps, n_clean, verbose): @@ -762,6 +948,7 @@ def main(): add_outdir_arg(output, help="Output directory for refined structures and maps") add_output_format_args(output) add_metadata_args(output) + add_all_columns_arg(output) refine = parser.add_argument_group("Refinement") refine.add_argument( @@ -1305,7 +1492,10 @@ def _mtz_to_cif(mtz_path, cif_path): print(f" Dark SF written to {dark_sf_mtz}, {dark_sf_cif}") print(f" Light SF written to {light_sf_mtz}, {light_sf_cif}") - write_results_mtz(dc, mc, scaler, diff_mtz_out) + map_diagnostics = write_results_mtz( + dc, mc.dark_model, scaler, diff_mtz_out, + mc=mc, all_columns=args.all_columns, verbose=args.verbose, + ) # --- JSON summary --- summary = { @@ -1338,6 +1528,7 @@ def _mtz_to_cif(mtz_path, cif_path): "alpha_mean": float(mc.alpha_mean), "lambda_twin": float(mc.lambda_twin), "sigma_alpha_sq": float(mc.sigma_alpha_sq), + **map_diagnostics, }, "output_files": { "dark_pdb": dark_pdb_out, diff --git a/torchref/cli/difference_map.py b/torchref/cli/difference_map.py new file mode 100644 index 00000000..9e52601a --- /dev/null +++ b/torchref/cli/difference_map.py @@ -0,0 +1,241 @@ +#!/usr/bin/env python3 -u + +"""Difference and extrapolated map coefficients from dark/light data. + +Uses the ``torchref.difference-refine`` pipeline but performs **no refinement**: the input +models are used as-is. + +The default output is the weighted difference map -- the inverse-variance-weighted +amplitude difference ``|Fo_light| - |Fo_dark|`` carried on the **dark** model's phases, +written as ``DELFWT``/``PHDELWT``. That needs no light-state model, so ``-lm`` is +optional. It is also deliberately not a *phased* difference map: putting the light +state's model phases into the observed amplitude biases the map toward the very model +the experiment is testing. + +Given ``-lm``, the light state's amplitude and phase and the extrapolated map follow. +``--all-columns`` adds the alternative constructions of both. + +Examples +-------- +:: + + # difference map only -- no light model, no fraction needed + torchref.difference-map \\ + -dm dark.pdb -dsf dark.mtz -lsf light.mtz -o results.mtz + + torchref.difference-map \\ + -dm dark.pdb -lm light.pdb -dsf dark.mtz -lsf light.mtz \\ + --fraction 0.37 --dmin 1.7 --cif ligand.cif -o results.mtz +""" + +import argparse +import sys +from pathlib import Path + +import torch + +from torchref.cli._common import ( + add_all_columns_arg, + add_dual_model_args, + add_dmin_arg, + add_general_args, + add_output_arg, + build_dual_column_names, + configure_unbuffered_output, + register_timing, + parse_device_str, + validate_cif_files, + validate_files, +) + +configure_unbuffered_output() + + +def main(): + """Entry point for ``torchref.difference-map``; returns the exit code.""" + parser = argparse.ArgumentParser( + description="Compute difference and extrapolated map coefficients " + "(no refinement).", + formatter_class=argparse.RawDescriptionHelpFormatter, + epilog=""" +Examples: + # difference map only -- needs no light model and no fraction + torchref.difference-map \\ + -dm dark.pdb \\ + -dsf dark.mtz -lsf light.mtz -o results.mtz + + torchref.difference-map \\ + -dm dark.pdb -lm light.pdb \\ + -dsf dark.mtz -lsf light.mtz \\ + --fraction 0.37 -o results.mtz + + torchref.difference-map \\ + -dm dark.pdb -lm light.pdb \\ + -dsf dark.mtz -lsf light.mtz \\ + --fraction 0.37 --dmin 1.7 --all-columns -o results.mtz + """, + ) + + # --- Input files (creates "Input files" and "Column selection" groups) --- + add_dual_model_args(parser, fraction_required=False, light_model_required=False) + + output = parser.add_argument_group("Output") + add_output_arg(output, help="Output MTZ file path (e.g. results.mtz)") + add_all_columns_arg(output) + + res = parser.add_argument_group("Resolution") + add_dmin_arg(res) + + add_general_args(parser) + + args = parser.parse_args() + + register_timing() + + # --- The light model gates everything that needs the light state's phases --- + has_light_model = args.light_model is not None + if has_light_model: + if args.fraction is None: + parser.error( + "--fraction is required with -lm/--light-model: the extrapolation " + "divides by the light-state occupancy." + ) + fractions = [1.0 - args.fraction, args.fraction] + else: + fractions = None + if args.fraction is not None: + parser.error( + "--fraction needs -lm/--light-model. Without a light model only the " + "weighted difference map is written, and it carries no occupancy: the " + "amplitude is |Fo_light| - |Fo_dark| and the weight comes from the " + "sigmas." + ) + + # --- Validate input files --- + to_check = [ + (args.dark_model, "dark model"), + (args.dark_structure_factor, "dark structure factor"), + (args.light_structure_factor, "light structure factor"), + ] + if has_light_model: + to_check.insert(1, (args.light_model, "light model")) + if validate_files(to_check): + return 1 + + if validate_cif_files(args.cif): + return 1 + + # Ensure output directory exists + out_path = Path(args.output) + out_path.parent.mkdir(parents=True, exist_ok=True) + + # --- Device --- + device = parse_device_str(args.device) + + # --- Header --- + if args.verbose > 0: + print("=" * 72) + print("TorchRef Difference Map") + print("=" * 72) + print(f"Dark model: {args.dark_model}") + if has_light_model: + print(f"Light model: {args.light_model}") + print(f"Fraction: light={args.fraction} " + f"(dark={1.0 - args.fraction})") + else: + print("Light model: (none) -- difference map only") + print(f"Dark SF: {args.dark_structure_factor}") + print(f"Light SF: {args.light_structure_factor}") + print(f"Output: {args.output}") + print(f"Device: {device}") + if args.dmin: + print(f"Resolution cutoff: {args.dmin:.2f} A") + if args.cif: + print(f"CIF restraints: {', '.join(args.cif)}") + if args.all_columns: + print("Columns: all") + print("=" * 72) + print() + sys.stdout.flush() + + from torchref.cli.collection_difference_refine import ( + compute_rfactors, + setup_dark_only, + setup_model_collection, + setup_dataset_collection, + setup_scaler, + write_results_mtz, + ) + + # --- Resolution --- + d_min = args.dmin if args.dmin is not None else 1.0 + + # --- Load data. dc.scale() puts the two datasets on one scale with no model, + # which is what makes the dark-only path possible at all. --- + if args.verbose > 0: + print("Loading reflection data...") + sys.stdout.flush() + + col_dark, col_light = build_dual_column_names(args) + + dc = setup_dataset_collection( + args.dark_structure_factor, args.light_structure_factor, args.dmin, device, + column_names_dark=col_dark, column_names_light=col_light, + ) + + if args.verbose > 0: + print("Setting up models...") + sys.stdout.flush() + + if has_light_model: + mc = setup_model_collection( + args.dark_model, args.light_model, fractions, + args.cif, d_min, device, args.verbose, + ) + mc["light"].freeze_fractions() + + if args.verbose > 0: + print("Setting up joint scaler...") + sys.stdout.flush() + scaler = setup_scaler(dc, mc, device, args.verbose) + dark_model = mc.dark_model + + if args.verbose > 0: + r_work_d, r_free_d = compute_rfactors(dark_model, dc["dark"], scaler) + r_work_l, r_free_l = compute_rfactors(mc["light"], dc["light"], scaler) + print(f" R-factor (dark): R_work={r_work_d:.4f} R_free={r_free_d:.4f}") + print(f" R-factor (mixed): R_work={r_work_l:.4f} R_free={r_free_l:.4f}") + print() + sys.stdout.flush() + else: + mc = None + dark_model, scaler = setup_dark_only( + args.dark_model, dc, args.cif, d_min, device, args.verbose, + ) + if args.verbose > 0: + r_work_d, r_free_d = compute_rfactors(dark_model, dc["dark"], scaler) + print(f" R-factor (dark): R_work={r_work_d:.4f} R_free={r_free_d:.4f}") + print() + sys.stdout.flush() + + # --- Write MTZ --- + if args.verbose > 0: + print("Computing map coefficients...") + sys.stdout.flush() + + with torch.no_grad(): + write_results_mtz( + dc, dark_model, scaler, str(out_path), + mc=mc, all_columns=args.all_columns, verbose=args.verbose, + ) + + if args.verbose > 0: + print() + print("Done.") + sys.stdout.flush() + + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/torchref/cli/phased_difference_map.py b/torchref/cli/phased_difference_map.py deleted file mode 100644 index 09648de0..00000000 --- a/torchref/cli/phased_difference_map.py +++ /dev/null @@ -1,184 +0,0 @@ -#!/usr/bin/env python3 -u - -"""Phased difference and extrapolated map coefficients from dark/light data. - -Uses the ``torchref.difference-refine`` pipeline but performs **no refinement**: the input -models are used as-is to compute phases, scale factors and every flavour of extrapolated -amplitude, and the result is one MTZ holding the observed, calculated, difference and -extrapolated columns. - -Examples --------- -:: - - torchref.phased-difference-map \\ - -dm dark.pdb -lm light.pdb -dsf dark.mtz -lsf light.mtz \\ - --fraction 0.37 --dmin 1.7 --cif ligand.cif -o results.mtz -""" - -import argparse -import sys -from pathlib import Path - -import torch - -from torchref.cli._common import ( - add_dual_model_args, - add_dmin_arg, - add_general_args, - add_output_arg, - build_dual_column_names, - configure_unbuffered_output, - register_timing, - parse_device_str, - validate_cif_files, - validate_files, -) - -configure_unbuffered_output() - - -def main(): - """Entry point for ``torchref.phased-difference-map``; returns the exit code.""" - parser = argparse.ArgumentParser( - description="Compute phased difference and extrapolated map " - "coefficients (no refinement).", - formatter_class=argparse.RawDescriptionHelpFormatter, - epilog=""" -Examples: - torchref.phased-difference-map \\ - -dm dark.pdb -lm light.pdb \\ - -dsf dark.mtz -lsf light.mtz \\ - --fraction 0.37 -o results.mtz - - torchref.phased-difference-map \\ - -dm dark.pdb -lm light.pdb \\ - -dsf dark.mtz -lsf light.mtz \\ - --fraction 0.37 --dmin 1.7 -o results.mtz - """, - ) - - # --- Input files (creates "Input files" and "Column selection" groups) --- - add_dual_model_args(parser) - - output = parser.add_argument_group("Output") - add_output_arg(output, help="Output MTZ file path (e.g. results.mtz)") - - res = parser.add_argument_group("Resolution") - add_dmin_arg(res) - - add_general_args(parser) - - args = parser.parse_args() - - register_timing() - - # --- Parse fraction --- - fractions = [1.0 - args.fraction, args.fraction] - - # --- Validate input files --- - if validate_files([ - (args.dark_model, "dark model"), - (args.light_model, "light model"), - (args.dark_structure_factor, "dark structure factor"), - (args.light_structure_factor, "light structure factor"), - ]): - return 1 - - if validate_cif_files(args.cif): - return 1 - - # Ensure output directory exists - out_path = Path(args.output) - out_path.parent.mkdir(parents=True, exist_ok=True) - - # --- Device --- - device = parse_device_str(args.device) - - # --- Header --- - if args.verbose > 0: - print("=" * 72) - print("TorchRef Phased Difference Map") - print("=" * 72) - print(f"Dark model: {args.dark_model}") - print(f"Light model: {args.light_model}") - print(f"Dark SF: {args.dark_structure_factor}") - print(f"Light SF: {args.light_structure_factor}") - print(f"Fraction: light={args.fraction} (dark={1.0 - args.fraction})") - print(f"Output: {args.output}") - print(f"Device: {device}") - if args.dmin: - print(f"Resolution cutoff: {args.dmin:.2f} A") - if args.cif: - print(f"CIF restraints: {', '.join(args.cif)}") - print("=" * 72) - print() - sys.stdout.flush() - - from torchref.cli.collection_difference_refine import ( - compute_rfactors, - setup_model_collection, - setup_dataset_collection, - setup_scaler, - write_results_mtz, - ) - - # --- Resolution --- - d_min = args.dmin if args.dmin is not None else 1.0 - - # --- Setup models --- - if args.verbose > 0: - print("Setting up models...") - sys.stdout.flush() - - mc = setup_model_collection( - args.dark_model, args.light_model, fractions, - args.cif, d_min, device, args.verbose, - ) - mc["light"].freeze_fractions() - - # --- Load data --- - if args.verbose > 0: - print("Loading reflection data...") - sys.stdout.flush() - - col_dark, col_light = build_dual_column_names(args) - - dc = setup_dataset_collection( - args.dark_structure_factor, args.light_structure_factor, args.dmin, device, - column_names_dark=col_dark, column_names_light=col_light, - ) - - # --- Scale --- - if args.verbose > 0: - print("Setting up joint scaler...") - sys.stdout.flush() - - scaler = setup_scaler(dc, mc, device, args.verbose) - - if args.verbose > 0: - r_work_d, r_free_d = compute_rfactors(mc.dark_model, dc["dark"], scaler) - r_work_l, r_free_l = compute_rfactors(mc["light"], dc["light"], scaler) - print(f" R-factor (dark): R_work={r_work_d:.4f} R_free={r_free_d:.4f}") - print(f" R-factor (mixed): R_work={r_work_l:.4f} R_free={r_free_l:.4f}") - print() - sys.stdout.flush() - - # --- Write MTZ --- - if args.verbose > 0: - print("Computing map coefficients...") - sys.stdout.flush() - - with torch.no_grad(): - write_results_mtz(dc, mc, scaler, str(out_path)) - - if args.verbose > 0: - print() - print("Done.") - sys.stdout.flush() - - return 0 - - -if __name__ == "__main__": - sys.exit(main()) From 33a7d582913bfbd50af6f6050f17fdab926b9bb7 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Mon, 7 Sep 2026 09:37:19 +0200 Subject: [PATCH 159/250] Precondition the rigid-body parameters and stop co-refining the scale Rigid-body refinement stores its Euler angles pre-multiplied by the chain's radius of gyration, so a unit step in an angle and a unit step in a translation displace atoms comparably. In radians against Angstroms the rotation block of the Hessian carried 190-530x the curvature of the translation block on 1DAW and 3E98 -- the geometric Rg**2, 411 and 442/516 -- putting cond(H) at 1e3-5e3, which is why six parameters needed hundreds of L-BFGS iterations to place. Dividing the scale out in forward() brings the ratio to 0.4-1.3 and cond(H) to 3-18; the step then converges rather than exhausting its iteration budget, on about half the gradient evaluations, with R-free no worse on any of ten structures. RigidXYZTensor.rotation_radians returns the physical angle, and setting angle_scale to ones gives the unscaled parametrization. update_fixed_values recomputes the scale alongside the chain centres, since it accepts coordinates that are not a rigid re-pose; copy() carries it rather than re-deriving it, so an overridden scale does not change the copy's pose. Not a fix for the one or two negative Hessian eigenvalues at the finer cutoffs, and those counts are unchanged. The step also no longer co-refines the scaler in the same L-BFGS as the rigid parameters. The body target centres on alpha*|F_calc| and alpha absorbs a rescaling of F_calc exactly, so the scale has a flat direction there; SCALE_TARGETS already excludes every alpha-centred row from the scale fit for this reason. refine_scaler (objective ls) owns the scale, between cutoffs. And fixed refine_rigid_body leaving the caller's reflection data truncated. cut_res masks in place and returns self, so each cutoff stamped its resolution mask on the caller's own object and the restore had nothing to restore to -- it only looked correct because the default schedule ends at the native limit. With --rigid-body-cutoffs 6,4 on a 2.05 A dataset, 20138 of 23352 reflections stayed masked out for the rest of the run, R-factors included. Object identity cannot catch this, which is why the existing isolation test did not. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01XodKhB5rwTiiAWFbonP6v2 --- docs/changelog.rst | 3 + .../integration/test_rigid_body_isolation.py | 36 +++++++ tests/unit/model/test_rigid_xyz.py | 93 ++++++++++++++++++- torchref/model/rigid_xyz.py | 53 ++++++++++- torchref/refinement/rigid_body_refinement.py | 36 +++++-- 5 files changed, 207 insertions(+), 14 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 31b86182..3284fb85 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,9 @@ Changelog Unreleased ---------- +- Rigid-body refinement stores its Euler angles pre-multiplied by the chain's radius of gyration, so a unit step in an angle and a unit step in a translation displace atoms comparably. In radians against Angstroms the rotation block of the Hessian carried 190-530x the curvature of the translation block on 1DAW and 3E98 -- the geometric ``Rg**2``, 411 and 442/516 -- putting ``cond(H)`` at 1e3-5e3, which is why six parameters needed ~250 L-BFGS iterations to place. Dividing the scale out in ``forward()`` brings the ratio to 0.4-1.3 and ``cond(H)`` to 3-18. Over ten structures the step then converges rather than exhausting its iteration budget, on about half the gradient evaluations, with R-free no worse anywhere. ``RigidXYZTensor.rotation_radians`` returns the physical angle, and setting ``angle_scale`` to ones restores the unscaled parametrization. Not a fix for the one or two negative Hessian eigenvalues at the finer cutoffs -- scaling a saddle leaves it a saddle -- and those counts are unchanged +- The rigid-body step no longer co-refines the scaler in the same L-BFGS as the rigid parameters. The body target centres on ``alpha*|F_calc|`` and ``alpha`` absorbs a rescaling of ``F_calc`` exactly, so the scale had a flat direction there; ``SCALE_TARGETS`` already excludes every alpha-centred row from the scale fit for this reason, and 0.6.2 fixed the same thing in the main driver. ``refine_scaler`` (objective ``ls``) owns the scale, between cutoffs +- Fixed ``refine_rigid_body`` leaving the caller's reflection data truncated. ``cut_res`` masks in place and returns ``self``, so each cutoff stamped its resolution mask on the caller's own object and the restore had nothing to restore to -- it only looked correct because the default schedule ends at the native limit. With ``--rigid-body-cutoffs 6,4`` on a 2.05 A dataset, 20138 of 23352 reflections stayed masked out for the rest of the run, R-factors included - Refinement output no longer inherits the input file's refinement header. It used to copy the whole thing and then append its own ``REMARK 3``, so a refined 3GR5 carried 420 header lines asserting two refinements at once -- ``PROGRAM : REFMAC 5.1.24`` with R-work 0.213 at line 5, ours at line 389 -- and a reader taking the first ``REMARK 3`` got REFMAC. The inherited block was not merely stale but contradicted the data beside it: it claimed a 5.1% / 1072-reflection test set, while the MTZ shipped with it holds 9.85% / 2063 (which torchref reads correctly). The passthrough was also inverted, keeping the statistics refinement invalidates and dropping the chemistry it does not -- SEQRES, SSBOND, DBREF, EXPDTA, COMPND, SOURCE, KEYWDS, SEQADV, HETNAM, FORMUL and SITE were all absent from the output. Now a whitelist carries the crystal, sample and chemistry records through in mandated record order (TITLE used to be emitted after REMARK 900), REMARK 2, 3 and 500 are dropped, and AUTHOR and JRNL are not inherited because they credit the deposition rather than this run. 283 header lines for the same file, 41 structural records preserved, one refinement block. ``add-metadata`` is exempt through ``supersede_refinement=False``: annotating a file is not re-refining it, so nothing there supersedes the existing REMARK 3 or AUTHOR records and both are kept - Prior refinements are tracked through mmCIF's ``_software`` loop, which is the only place either format has room for them: ``_refine`` is singular by design, so a previous program's statistics cannot be kept without contradicting the current ones. ``pdbx_ordinal`` was hardcoded to ``1`` and the incoming loop was never read, truncating the chain to one link on every write; it now reads the input's loop and appends at ``max(ordinal) + 1``, carrying each entry's ``description`` so the chain says what every program did and not just that it ran. Added ``_pdbx_initial_refinement_model`` and ``_refine.pdbx_starting_model``, which name what the refinement started from, and ``_refine.pdbx_R_Free_selection_details``, which names the test set the reported R-free conditions on. ``from_cif_file`` no longer carries the input's ``_refine`` items through -- that was the mmCIF form of the duplicated ``REMARK 3`` - Fixed mmCIF loop cells being written unquoted, which silently split any value containing whitespace into extra columns when the file was read back. Latent while every loop column was a single token; a multi-word ``_software.description`` exposed it. Nulls stay bare, since ``gemmi.cif.quote`` turns ``?`` into the quoted one-character string ``'?'`` diff --git a/tests/integration/test_rigid_body_isolation.py b/tests/integration/test_rigid_body_isolation.py index 651ed11d..97ca1587 100644 --- a/tests/integration/test_rigid_body_isolation.py +++ b/tests/integration/test_rigid_body_isolation.py @@ -71,6 +71,42 @@ def test_targets_and_data_are_not_replaced(refinement): assert ref.reflection_data is data +def test_resolution_range_survives_a_coarse_only_cutoff_list(refinement): + """The caller's resolution range must be what it was, not the last cutoff's. + + `cut_res` masks in place and returns `self`, so a `cutoffs` list ending above + the native d_min can leave the caller truncated. Object identity does not + catch it -- the data object is the same one throughout. + """ + ref = refinement() + data = ref.reflection_data + ref.get_scales() + n_before = int(data.masks().sum()) + n_work_before = int(data.work.mask.sum()) + d_min_before = data.get_max_res() + + # Deliberately coarse-only, and deliberately not ending at the native d_min. + ref.refine_rigid_body(iterations_per_step=5, cutoffs=[6.0, 4.0]) + + assert int(data.masks().sum()) == n_before + assert int(data.work.mask.sum()) == n_work_before + assert data.get_max_res() == pytest.approx(d_min_before) + + +def test_a_caller_supplied_resolution_limit_is_not_widened(refinement): + """A refinement built with `max_res` keeps that limit across a rigid-body run.""" + ref = refinement() + ref.reflection_data.cut_res(highres=3.5) + data = ref.reflection_data + ref.get_scales() + n_before = int(data.masks().sum()) + + ref.refine_rigid_body(iterations_per_step=5) + + assert int(data.masks().sum()) == n_before + assert data.get_max_res() >= 3.5 + + def test_refined_coordinates_still_reach_the_caller(refinement): """The sandbox shares the model, so the whole point still has to work.""" ref = refinement() diff --git a/tests/unit/model/test_rigid_xyz.py b/tests/unit/model/test_rigid_xyz.py index d7e76531..4628c183 100644 --- a/tests/unit/model/test_rigid_xyz.py +++ b/tests/unit/model/test_rigid_xyz.py @@ -40,16 +40,18 @@ def test_known_transform(self, fresh_modelft): dtype = xyz.dtype # Pick the first chain (index 0); apply a transform only to that chain. + # euler_angles is in Angstrom-scaled units, so scale going in and read + # the physical angle back through rotation_radians. ang_vec = torch.tensor([0.0, 0.05, 0.0], dtype=dtype, device=device) trans_vec = torch.tensor([0.2, -0.3, 0.5], dtype=dtype, device=device) with torch.no_grad(): xyz.euler_angles.zero_() xyz.translations.zero_() - xyz.euler_angles[0] = ang_vec + xyz.euler_angles[0] = ang_vec * xyz.angle_scale[0] xyz.translations[0] = trans_vec center = xyz.chain_centers[0] - R = rotation_matrix_euler_xyz(ang_vec) + R = rotation_matrix_euler_xyz(xyz.rotation_radians[0]) atom_chain = xyz.chain_indices chain0_mobile = (atom_chain == 0) & xyz.mobile_mask non_mobile = ~xyz.mobile_mask @@ -118,9 +120,10 @@ def test_bake_preserves_forward_and_zeros_params(self, fresh_modelft): with torch.no_grad(): xyz.euler_angles.zero_() xyz.translations.zero_() - xyz.euler_angles[0] = ang + xyz.euler_angles[0] = ang * xyz.angle_scale[0] xyz.translations[0] = trans before = xyz().detach().clone() + scale_before = xyz.angle_scale.clone() xyz.bake() @@ -134,6 +137,9 @@ def test_bake_preserves_forward_and_zeros_params(self, fresh_modelft): assert diff_original < 1e-4 assert torch.all(xyz.euler_angles == 0).item() assert torch.all(xyz.translations == 0).item() + # Rigid motion preserves the radius of gyration, so the scale is + # unchanged; if it drifted, later angles would mean something else. + assert torch.allclose(xyz.angle_scale, scale_before, rtol=1e-6) # Chain centers should have moved by the translation on chain 0. # Centroid is mass-weighted (atomic Z) over MOBILE atoms; reconstruct @@ -147,6 +153,87 @@ def test_bake_preserves_forward_and_zeros_params(self, fresh_modelft): diff_center = (xyz.chain_centers[0] - expected_center0).abs().max().item() assert diff_center < 1e-4 + @pytest.mark.unit + def test_angle_scale_matches_explicit_radius_of_gyration(self, fresh_modelft): + """``angle_scale`` is the per-chain RMS radius about the rotation centre, + over mobile atoms, floored at 1 A.""" + fresh_modelft.use_rigid_xyz() + xyz = fresh_modelft.xyz + + for c in range(xyz.n_chains): + sel = (xyz.chain_indices == c) & xyz.mobile_mask + d = xyz.original_xyz[sel] - xyz.chain_centers[c] + expected = d.pow(2).sum(dim=1).mean().sqrt() + assert torch.allclose(xyz.angle_scale[c], expected, rtol=1e-5) + assert torch.all(xyz.angle_scale >= 1.0), "scale must be floored at 1 A" + + @pytest.mark.unit + def test_angle_scale_tracks_non_rigid_update_fixed_values(self, fresh_modelft): + """``update_fixed_values`` recomputes the scale, not just the centres. + + It accepts any coordinates, not only the rigid re-pose ``bake()`` supplies, + and a rigid re-pose leaves the radius of gyration unchanged. + """ + fresh_modelft.use_rigid_xyz() + xyz = fresh_modelft.xyz + before = xyz.angle_scale.clone() + + # A pure dilation is not a rigid motion, so the radius doubles. + with torch.no_grad(): + centers = xyz.chain_centers[xyz.chain_indices] + inflated = centers + 2.0 * (xyz.original_xyz - centers) + xyz.update_fixed_values(inflated) + + assert torch.allclose(xyz.angle_scale, 2.0 * before, rtol=1e-4) + + @pytest.mark.unit + def test_unit_angle_scale_reproduces_unscaled_behaviour(self, fresh_modelft): + """``angle_scale`` of ones gives the unscaled parametrization exactly.""" + fresh_modelft.use_rigid_xyz() + xyz = fresh_modelft.xyz + dtype, device = xyz.dtype, xyz.device + ang = torch.tensor([0.02, -0.03, 0.04], dtype=dtype, device=device) + + with torch.no_grad(): + xyz.angle_scale.fill_(1.0) + xyz.euler_angles.zero_() + xyz.euler_angles[0] = ang + xyz.reset_forward_cache() + + center = xyz.chain_centers[0] + R = rotation_matrix_euler_xyz(ang) + sel = (xyz.chain_indices == 0) & xyz.mobile_mask + expected = (xyz.original_xyz[sel] - center) @ R.T + center + with torch.no_grad(): + assert torch.allclose(xyz()[sel], expected, atol=1e-4) + + @pytest.mark.unit + def test_rotation_radians_is_the_descaled_angle(self, fresh_modelft): + """``rotation_radians`` descales the angle, and ``copy()`` carries both.""" + fresh_modelft.use_rigid_xyz() + xyz = fresh_modelft.xyz + with torch.no_grad(): + xyz.euler_angles.copy_(torch.randn_like(xyz.euler_angles) * 0.1) + + assert torch.allclose( + xyz.rotation_radians, xyz.euler_angles / xyz.angle_scale.unsqueeze(1) + ) + clone = xyz.copy() + assert torch.allclose(clone.angle_scale, xyz.angle_scale) + assert torch.allclose(clone.euler_angles, xyz.euler_angles) + with torch.no_grad(): + assert torch.allclose(clone(), xyz(), atol=1e-5) + + # An overridden scale must be carried, not re-derived from geometry: + # the angles are copied raw, so a re-derived scale changes the pose. + with torch.no_grad(): + xyz.angle_scale.fill_(1.0) + xyz.reset_forward_cache() + clone2 = xyz.copy() + assert torch.allclose(clone2.angle_scale, torch.ones_like(xyz.angle_scale)) + with torch.no_grad(): + assert torch.allclose(clone2(), xyz(), atol=1e-5) + @pytest.mark.unit def test_restore_commit_bakes_transform(self, fresh_modelft): fresh_modelft.use_rigid_xyz() diff --git a/torchref/model/rigid_xyz.py b/torchref/model/rigid_xyz.py index 4d4c0b5f..0d0a304a 100644 --- a/torchref/model/rigid_xyz.py +++ b/torchref/model/rigid_xyz.py @@ -9,6 +9,15 @@ centroid and translating it. XYZ Euler matches Phenix's default ``euler_angle_convention`` and keeps the rotation Jacobian full-rank at the origin (no gimbal lock when angles reset to zero after ``bake()``). + +**``euler_angles`` is NOT in radians.** It is stored pre-multiplied by +``angle_scale``, the per-chain radius of gyration in Angstroms, so a unit step in +an angle and a unit step in a translation displace atoms comparably. Without it +the rotation block of the Hessian carries ~``Rg**2`` the curvature of the +translation block and L-BFGS needs an order of magnitude more iterations to place +six parameters. ``forward()`` divides the scale out; +:attr:`RigidXYZTensor.rotation_radians` returns the physical angle. Setting +``angle_scale`` to ones gives the unscaled parametrization. """ from typing import Optional, Sequence @@ -71,6 +80,7 @@ def __init__( "mobile_mask", torch.empty(0, dtype=torch.bool, device=device) ) self.register_buffer("atom_weights", torch.empty(0, device=device, dtype=dtype)) + self.register_buffer("angle_scale", torch.empty(0, device=device, dtype=dtype)) self.euler_angles = nn.Parameter(torch.empty(0, 3, device=device, dtype=dtype)) self.translations = nn.Parameter(torch.empty(0, 3, device=device, dtype=dtype)) self._n_chains = 0 @@ -162,6 +172,12 @@ def __init__( self.register_buffer("chain_centers", chain_centers) self.register_buffer("mobile_mask", mobile_t) self.register_buffer("atom_weights", atom_weights_t) + self.register_buffer( + "angle_scale", + self._compute_angle_scale( + original_xyz_t, chain_centers, mobile_idx, mobile_t, n_chains + ), + ) self.euler_angles = nn.Parameter( torch.zeros((n_chains, 3), dtype=dtype, device=device) @@ -173,6 +189,28 @@ def __init__( self._n_chains = n_chains self._chain_id_order = list(chain_id_order) + # ----------------------------------------------------------------------- + # Angle scaling (preconditioning) + # ----------------------------------------------------------------------- + @staticmethod + def _compute_angle_scale(xyz, centers, mobile_idx, mobile_mask, n_chains): + """Per-chain radius of gyration about the rotation centre, in Angstroms. + + The lever arm converting radians into Angstroms of atom displacement. + Floored at 1 A so a two-atom body cannot divide by ~0. + """ + d = xyz[mobile_mask] - centers[mobile_idx] + sq = torch.zeros(n_chains, dtype=xyz.dtype, device=xyz.device) + sq.index_add_(0, mobile_idx, d.pow(2).sum(dim=1)) + counts = torch.zeros(n_chains, dtype=xyz.dtype, device=xyz.device) + counts.index_add_(0, mobile_idx, torch.ones_like(d[:, 0])) + return (sq / counts.clamp(min=1.0)).sqrt().clamp(min=1.0) + + @property + def rotation_radians(self) -> torch.Tensor: + """The physical per-chain XYZ-Euler angles, in radians.""" + return self.euler_angles / self.angle_scale.unsqueeze(1) + # ----------------------------------------------------------------------- # Forward — reconstruct full xyz # ----------------------------------------------------------------------- @@ -185,7 +223,7 @@ def forward(self) -> torch.Tensor: # the angles are exactly zero): XYZ keeps the Jacobian full-rank # at the origin, while ZYZ has a gimbal-lock singularity there # (dR/dα_1 and dR/dα_3 collapse onto z-axis rotations when β=0). - R = rotation_matrix_euler_xyz(self.euler_angles) # (n_chains, 3, 3) + R = rotation_matrix_euler_xyz(self.rotation_radians) # (n_chains, 3, 3) per_atom_R = R[self.chain_indices] # (N, 3, 3) per_atom_center = self.chain_centers[self.chain_indices] # (N, 3) @@ -333,6 +371,14 @@ def update_fixed_values(self, new_values: torch.Tensor): ) centers.index_add_(0, mobile_idx, mobile_xyz * mobile_w.unsqueeze(1)) self.chain_centers.copy_(centers / w_sum.unsqueeze(1).clamp(min=1e-12)) + # The scale follows the chain geometry. A rigid re-pose leaves it + # unchanged, but any coordinates are accepted here, so recompute. + self.angle_scale.copy_( + self._compute_angle_scale( + self.original_xyz, self.chain_centers, mobile_idx, mobile, + self._n_chains, + ) + ) self.euler_angles.zero_() self.translations.zero_() self.reset_forward_cache() @@ -355,6 +401,11 @@ def copy(self) -> "RigidXYZTensor": atom_weights=self.atom_weights.clone(), ) with torch.no_grad(): + # Carry the scale rather than letting the constructor re-derive it: + # the angles are copied raw, so a scale that has been overridden + # (``angle_scale.fill_(1.0)``) would otherwise give the copy a + # different pose from the original. + new.angle_scale.copy_(self.angle_scale) new.euler_angles.copy_(self.euler_angles) new.translations.copy_(self.translations) return new diff --git a/torchref/refinement/rigid_body_refinement.py b/torchref/refinement/rigid_body_refinement.py index 515f54bc..5718c19b 100644 --- a/torchref/refinement/rigid_body_refinement.py +++ b/torchref/refinement/rigid_body_refinement.py @@ -151,6 +151,20 @@ def _run(self): else self.default_cutoffs(native_dmin) ) + # ``cut_res`` masks in place and returns ``self``, so each cutoff below + # stamps ``masks["resolution"]`` on the caller's own object and rebinding + # restores nothing. Snapshot it (or its absence) to put back. + had_resolution_mask = "resolution" in original_data.masks + saved_resolution_mask = ( + original_data.masks["resolution"].clone() if had_resolution_mask else None + ) + + def restore_resolution_mask(): + if had_resolution_mask: + original_data.masks["resolution"] = saved_resolution_mask + else: + original_data.masks.pop("resolution", None) + # Swap the model's xyz container in place for a RigidXYZTensor. ref.model.use_rigid_xyz() @@ -165,7 +179,9 @@ def _run(self): step_state = self._run_one_cutoff(d_min) history.append((float(d_min), step_state)) finally: - # Restore full-resolution data view. + # Before rebinding, so the scaler and targets are built against the + # data the caller has. + restore_resolution_mask() self._rebind_for_data(original_data) if self.commit: @@ -250,10 +266,8 @@ def _run_one_cutoff(self, d_min: float): ] # Decide whether to use the inner-cycle (mask-refresh) loop. - # Triggered when the scaler has a bulk-solvent component whose - # mask depends on atom positions (ls_wunit_k1 path here). For - # the ml path the scaler is fully refit between cutoffs - # and co-optimized with rigid params in a single LBFGS. + # Unsatisfiable as it stands: nothing sets ``c_iso.requires_grad = + # False``, so ``_run_inner_cycles`` does not run. use_inner_cycles = ( ref.scaler is not None and getattr(ref.scaler, "solvent", None) is not None @@ -264,11 +278,13 @@ def _run_one_cutoff(self, d_min: float): if use_inner_cycles: self._run_inner_cycles(d_min, state, rigid_params, n_inner=5) else: - # Single-shot path: rigid params + scaler params co-optimized. - if ref.scaler is not None: - opt_params = rigid_params + list(ref.scaler.parameters()) - else: - opt_params = rigid_params + # Rigid parameters only. The body target centres on + # ``alpha*|F_calc|`` and ``alpha`` absorbs a rescaling of ``F_calc`` + # exactly, so the scale has a flat direction here -- the rule + # ``SCALE_TARGETS`` states for the scale fit. ``refine_scaler`` + # (objective ``ls``) owns the scale, between cutoffs. Omitting them + # is enough: ``LossState.run`` freezes leaves the optimizer lacks. + opt_params = rigid_params rigid_model.reset_cache() opt = torch.optim.LBFGS( From bde818fbf17de0276652eac6b3f6670caefee400 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 7 Sep 2026 16:46:19 +0200 Subject: [PATCH 160/250] test: separate fixtures and share read-only functional setup --- docs/changelog.rst | 1 + tests/README.md | 5 +- tests/conftest.py | 530 ++----------------- tests/fixtures/README.md | 40 ++ tests/fixtures/__init__.py | 5 + tests/fixtures/devices.py | 119 +++++ tests/fixtures/files.py | 98 ++++ tests/fixtures/functional.py | 42 ++ tests/fixtures/numerical.py | 195 +++++++ tests/fixtures/objects.py | 163 ++++++ tests/fixtures/paths.py | 68 +++ tests/fixtures/precision.py | 51 ++ tests/functional/conftest.py | 11 +- tests/functional/test_io_functional.py | 181 +++---- tests/functional/test_model_ft_functional.py | 159 +++--- tests/integration/conftest.py | 9 +- tests/unit/conftest.py | 166 +----- tests/unit/structure_factor/conftest.py | 35 +- 18 files changed, 1004 insertions(+), 874 deletions(-) create mode 100644 tests/fixtures/README.md create mode 100644 tests/fixtures/__init__.py create mode 100644 tests/fixtures/devices.py create mode 100644 tests/fixtures/files.py create mode 100644 tests/fixtures/functional.py create mode 100644 tests/fixtures/numerical.py create mode 100644 tests/fixtures/objects.py create mode 100644 tests/fixtures/paths.py create mode 100644 tests/fixtures/precision.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 3284fb85..b5b00059 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Organized test fixtures into focused modules and reused module-scoped loaded objects for read-only functional checks while retaining fresh objects for mutation and loading tests. - Rigid-body refinement stores its Euler angles pre-multiplied by the chain's radius of gyration, so a unit step in an angle and a unit step in a translation displace atoms comparably. In radians against Angstroms the rotation block of the Hessian carried 190-530x the curvature of the translation block on 1DAW and 3E98 -- the geometric ``Rg**2``, 411 and 442/516 -- putting ``cond(H)`` at 1e3-5e3, which is why six parameters needed ~250 L-BFGS iterations to place. Dividing the scale out in ``forward()`` brings the ratio to 0.4-1.3 and ``cond(H)`` to 3-18. Over ten structures the step then converges rather than exhausting its iteration budget, on about half the gradient evaluations, with R-free no worse anywhere. ``RigidXYZTensor.rotation_radians`` returns the physical angle, and setting ``angle_scale`` to ones restores the unscaled parametrization. Not a fix for the one or two negative Hessian eigenvalues at the finer cutoffs -- scaling a saddle leaves it a saddle -- and those counts are unchanged - The rigid-body step no longer co-refines the scaler in the same L-BFGS as the rigid parameters. The body target centres on ``alpha*|F_calc|`` and ``alpha`` absorbs a rescaling of ``F_calc`` exactly, so the scale had a flat direction there; ``SCALE_TARGETS`` already excludes every alpha-centred row from the scale fit for this reason, and 0.6.2 fixed the same thing in the main driver. ``refine_scaler`` (objective ``ls``) owns the scale, between cutoffs - Fixed ``refine_rigid_body`` leaving the caller's reflection data truncated. ``cut_res`` masks in place and returns ``self``, so each cutoff stamped its resolution mask on the caller's own object and the restore had nothing to restore to -- it only looked correct because the default schedule ends at the native limit. With ``--rigid-body-cutoffs 6,4`` on a 2.05 A dataset, 20138 of 23352 reflections stayed masked out for the rest of the run, R-factors included diff --git a/tests/README.md b/tests/README.md index fd283106..ff1a3bc5 100644 --- a/tests/README.md +++ b/tests/README.md @@ -6,7 +6,8 @@ This directory contains the complete test suite for torchref. ``` tests/ -├── conftest.py # Root fixtures (paths, devices, skip decorators) +├── conftest.py # Fixture registration and test-selection hooks +├── fixtures/ # Shared setup, grouped by responsibility (see fixtures/README.md) ├── pytest.ini # Pytest configuration ├── __init__.py ├── files/ # Test data files (CIF, PDB, MTZ) @@ -15,7 +16,7 @@ tests/ │ ├── mtz/ # Reflection MTZ files │ └── cif_sf/ # Structure factor CIF files ├── unit/ # Unit tests (fast, no I/O) -│ ├── conftest.py # Unit test fixtures (mock data) +│ ├── conftest.py # Imports scoped numerical fixtures │ ├── math_functions/ # Math module tests │ ├── model/ # Model module tests │ ├── refinement/ # Refinement module tests diff --git a/tests/conftest.py b/tests/conftest.py index b8b7178e..833cd222 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,35 +1,26 @@ -""" -Root pytest configuration and shared fixtures for torchref tests. +"""Register shared fixture plugins and gate tests on host capabilities.""" -This module provides fixtures that are automatically available to all test files. -""" import importlib.util import shutil import warnings import pytest -import torchref -import torch -import numpy as np -from pathlib import Path +# Apply process-wide settings before the device fixtures import torch. +import torchref # noqa: F401 + +pytest_plugins = ( + "tests.fixtures.paths", + "tests.fixtures.files", + "tests.fixtures.devices", + "tests.fixtures.precision", + "tests.fixtures.objects", +) -# Optional Amber/ensemble stack: OpenMM (pip ``[amber]`` extra) and AmberTools -# (antechamber/tleap — conda-only, detected on PATH). Tests that need them are -# tagged ``@pytest.mark.openmm`` (OpenMM only) or ``@pytest.mark.amber`` (OpenMM -# + AmberTools) and auto-skipped below when the stack is absent. _HAS_OPENMM = importlib.util.find_spec("openmm") is not None _HAS_AMBERTOOLS = bool(shutil.which("antechamber") and shutil.which("tleap")) -def _cuda_available() -> bool: - return torch.cuda.is_available() - - -def _mps_available() -> bool: - return hasattr(torch.backends, "mps") and torch.backends.mps.is_available() - - def pytest_addoption(parser): """Add custom command line options.""" parser.addoption( @@ -58,24 +49,34 @@ def pytest_addoption(parser): help="Deprecated no-op: accelerator tests now run automatically.", ) parser.addoption( - "--run-slow", - action="store_true", - default=False, - help="Run slow tests" + "--run-slow", action="store_true", default=False, help="Run slow tests" ) def pytest_configure(config): """Configure pytest markers.""" config.addinivalue_line("markers", "unit: Unit tests (fast, no I/O)") - config.addinivalue_line("markers", "integration: Integration tests (slower, real I/O)") - config.addinivalue_line("markers", "gpu: Needs any accelerator (CUDA or MPS); auto-skipped if none") - config.addinivalue_line("markers", "cuda: Needs CUDA specifically (e.g. Triton); auto-skipped if absent") - config.addinivalue_line("markers", "mps: Needs MPS specifically (Metal kernels); auto-skipped if absent") + config.addinivalue_line( + "markers", "integration: Integration tests (slower, real I/O)" + ) + config.addinivalue_line( + "markers", "gpu: Needs any accelerator (CUDA or MPS); auto-skipped if none" + ) + config.addinivalue_line( + "markers", "cuda: Needs CUDA specifically (e.g. Triton); auto-skipped if absent" + ) + config.addinivalue_line( + "markers", "mps: Needs MPS specifically (Metal kernels); auto-skipped if absent" + ) config.addinivalue_line("markers", "cuda_only: Deprecated alias for 'cuda'") config.addinivalue_line("markers", "slow: Slow tests (skipped by default)") - config.addinivalue_line("markers", "openmm: Needs OpenMM (the [amber] extra); skipped if absent") - config.addinivalue_line("markers", "amber: Needs OpenMM + AmberTools (antechamber/tleap); skipped if absent") + config.addinivalue_line( + "markers", "openmm: Needs OpenMM (the [amber] extra); skipped if absent" + ) + config.addinivalue_line( + "markers", + "amber: Needs OpenMM + AmberTools (antechamber/tleap); skipped if absent", + ) if config.getoption("--run-gpu"): # UserWarning, not DeprecationWarning: pytest.ini filters the latter, @@ -109,6 +110,8 @@ def pytest_collection_modifyitems(config, items): mask a forgotten marker, and turns "this host cannot run it" into a silent pass instead of the visible skip or the real error. """ + from tests.fixtures.devices import _cuda_available, _mps_available + has_cuda = _cuda_available() has_mps = _mps_available() @@ -138,7 +141,9 @@ def pytest_collection_modifyitems(config, items): ) skip_slow = pytest.mark.skip(reason="Need --run-slow option to run") - skip_openmm = pytest.mark.skip(reason="OpenMM not installed (pip install '.[amber]')") + skip_openmm = pytest.mark.skip( + reason="OpenMM not installed (pip install '.[amber]')" + ) skip_amber = pytest.mark.skip( reason="AmberTools (antechamber/tleap) not on PATH (conda install ambertools)" ) @@ -174,466 +179,3 @@ def pytest_collection_modifyitems(config, items): item.add_marker(skip_amber) elif "openmm" in item.keywords and not _HAS_OPENMM: item.add_marker(skip_openmm) - - -# ============================================================================= -# Path Fixtures -# ============================================================================= - -@pytest.fixture(scope="session") -def tests_root() -> Path: - """Root of the tests directory.""" - return Path(__file__).parent - - -@pytest.fixture(scope="session") -def project_root() -> Path: - """Root of the project.""" - return Path(__file__).parent.parent - - -@pytest.fixture(scope="session") -def test_files_dir(tests_root) -> Path: - """Path to test files directory.""" - return tests_root / "files" - - -@pytest.fixture(scope="session") -def cif_dir(test_files_dir) -> Path: - """Path to CIF model files.""" - return test_files_dir / "cif" - - -@pytest.fixture(scope="session") -def cif_sf_dir(test_files_dir) -> Path: - """Path to CIF structure factor files.""" - return test_files_dir / "cif_sf" - - -@pytest.fixture(scope="session") -def mtz_dir(test_files_dir) -> Path: - """Path to MTZ reflection files.""" - return test_files_dir / "mtz" - - -@pytest.fixture(scope="session") -def pdb_dir(test_files_dir) -> Path: - """Path to PDB model files.""" - return test_files_dir / "pdb" - - -@pytest.fixture(scope="session") -def external_monomer_library(project_root) -> Path: - """Path to external monomer library.""" - return project_root / "external_monomer_library" - - -# ============================================================================= -# Device Fixtures -# ============================================================================= - -@pytest.fixture(scope="session") -def cpu_device() -> torch.device: - """CPU torch device.""" - return torch.device("cpu") - - -@pytest.fixture(scope="session") -def gpu_device() -> torch.device: - """GPU torch device (only use with @pytest.mark.gpu). - - Prefers CUDA, falls back to MPS; skips if neither is available. Prefer the - backend-specific ``cuda_device`` / ``mps_device`` below when a test needs one - particular backend -- this fixture's preference order means a - ``cuda``-marked test asking for it on a dual-backend host could be handed - MPS, which is why the MPS tests used to carry a ``type != 'mps'`` skip to - undo it. - """ - accel = _accelerator() - if accel is None: - pytest.skip("No accelerator (CUDA or MPS) on this host") - return accel - - -@pytest.fixture(scope="session") -def cuda_device() -> torch.device: - """Canonical CUDA device for ``cuda``-marked tests. - - Deliberately unguarded. What runs is decided by the ``cuda`` marker in - :func:`pytest_collection_modifyitems` and nowhere else, so this fixture does - not re-check availability: on a host without CUDA the test is *meant* to - error with the real backend error rather than be quietly skipped here. - """ - return torch.device("cuda", 0) - - -@pytest.fixture(scope="session") -def mps_device() -> torch.device: - """Canonical MPS device for ``mps``-marked tests. - - Unguarded for the same reason as :func:`cuda_device` -- the ``mps`` marker - owns the decision. - """ - return torch.device("mps", 0) - - -def _accelerator() -> "torch.device | None": - """The canonical accelerator this host can actually use, or ``None``. - - Indices are filled in (``cuda:0`` / ``mps:0``) so the value compares equal - to a device read back off a real tensor -- ``torch.device('mps')`` and - ``torch.device('mps:0')`` are *not* equal even though they name the same - physical device. - """ - if _cuda_available(): - return torch.device("cuda", torch.cuda.current_device()) - if _mps_available(): - return torch.device("mps", 0) - return None - - -# Built at import time so the ``gpu`` mark is attached during *collection*. -# Adding it later (e.g. via ``request.node.add_marker`` inside the fixture) is -# too late for ``pytest_collection_modifyitems`` to gate on. -_DEVICE_PARAMS = [pytest.param(torch.device("cpu"), id="cpu")] -_ACCELERATOR = _accelerator() -if _ACCELERATOR is not None: - _DEVICE_PARAMS.append( - pytest.param( - _ACCELERATOR, - id=_ACCELERATOR.type, - # Backend-specific mark, so a CUDA-less host skips the cuda leg and - # a non-Mac skips the mps leg, each with an accurate reason. - marks=getattr(pytest.mark, _ACCELERATOR.type), - ) - ) - - -@pytest.fixture(params=_DEVICE_PARAMS) -def any_device(request) -> torch.device: - """Every device this host can actually use, one test run per device. - - The CPU leg always runs. The accelerator leg is ``gpu``-marked, so a plain - ``pytest`` run skips it and ``pytest --run-gpu`` picks up CUDA on a CUDA - box or MPS on a Mac. On a CPU-only host the accelerator parameter does not - exist at all, so there is no skip noise. - """ - return request.param - - -@pytest.fixture(scope="session") -def _device_model_cache() -> dict: - """``{device_str: ModelFT}`` built at most once per device, per session.""" - return {} - - -@pytest.fixture -def device_model_bundle(_device_model_cache, pdb_dir, any_device): - """A loaded model on ``any_device``, for target conformance tests. - - The existing ``loaded_model`` / ``model_and_data`` fixtures are - function-scoped and construct on the process default, so a - device-parametrized sweep over them would reload the structure once per - test per device. This caches one model per device instead. - - Shared mutable state: callers must treat the bundle as read-only. A test - that moves a *target* will drag the borrowed model with it, poisoning every - later test on that device -- see ``test_target_device_round_trip``, which - deliberately builds its own. - """ - key = str(any_device) - if key not in _device_model_cache: - pdb = pdb_dir / "1DAW.pdb" - if not pdb.exists(): - pytest.skip("1DAW.pdb fixture not present") - from torchref.model import ModelFT - - _device_model_cache[key] = ModelFT(device=any_device, verbose=0).load_pdb( - str(pdb) - ) - return {"model": _device_model_cache[key]} - - -@pytest.fixture -def device(request) -> torch.device: - """Default test device. - - Uses the package-wide auto-detected default (``torchref.device.current``) - so tests run on whichever device the user's machine resolved to at - import time: cuda -> mps -> cpu. Tests marked ``@pytest.mark.cuda_only`` - are skipped when CUDA is not available. - """ - from torchref.config import get_default_device - - markers = {m.name for m in request.node.iter_markers()} - if "cuda_only" in markers and not torch.cuda.is_available(): - pytest.skip("Test requires CUDA") - if "gpu" in markers and not (_cuda_available() or _mps_available()): - pytest.skip("No GPU (CUDA or MPS) available") - return get_default_device() - - -# ============================================================================= -# Numerical Fixtures -# ============================================================================= - -@pytest.fixture -def rtol() -> float: - """Relative tolerance for floating point comparisons.""" - return 1e-5 - - -@pytest.fixture -def atol() -> float: - """Absolute tolerance for floating point comparisons.""" - return 1e-8 - - -# ============================================================================= -# Sample File Fixtures -# ============================================================================= - -@pytest.fixture(scope="session") -def sample_cif_file(cif_dir): - """Return a sample CIF file for testing.""" - cif_file = cif_dir / "1DAW.cif" - if cif_file.exists(): - return cif_file - # Try any available CIF file - cif_files = list(cif_dir.glob("*.cif")) - if cif_files: - return cif_files[0] - pytest.skip("No CIF files found in test data") - - -@pytest.fixture(scope="session") -def sample_mtz_file(mtz_dir): - """Return a sample MTZ file for testing.""" - mtz_file = mtz_dir / "1DAW.mtz" - if mtz_file.exists(): - return mtz_file - # Try any available MTZ file - mtz_files = list(mtz_dir.glob("*.mtz")) - if mtz_files: - return mtz_files[0] - pytest.skip("No MTZ files found in test data") - - -@pytest.fixture(scope="session") -def sample_pdb_file(pdb_dir): - """Return a sample PDB file for testing.""" - pdb_files = sorted(pdb_dir.glob("*.pdb")) - if not pdb_files: - pytest.skip("No PDB files found in test data directory") - return pdb_files[0] - - -@pytest.fixture(scope="session") -def sample_structure_factor_cif(cif_sf_dir): - """Return a sample structure factor CIF file.""" - sf_files = sorted(cif_sf_dir.glob("*.cif")) - if not sf_files: - pytest.skip("No structure factor CIF files found") - return sf_files[0] - - -@pytest.fixture(scope="session") -def sample_structure_pair(cif_dir, mtz_dir): - """Return a matching pair of CIF model and MTZ reflections.""" - # Try to find matching files - pdb_id = "1DAW" - cif_file = cif_dir / f"{pdb_id}.cif" - mtz_file = mtz_dir / f"{pdb_id}.mtz" - - if cif_file.exists() and mtz_file.exists(): - return {"model": cif_file, "reflections": mtz_file} - - # Try to find any matching pair - cif_files = {f.stem: f for f in cif_dir.glob("*.cif")} - mtz_files = {f.stem: f for f in mtz_dir.glob("*.mtz")} - - common_ids = set(cif_files.keys()) & set(mtz_files.keys()) - if common_ids: - pdb_id = sorted(common_ids)[0] - return {"model": cif_files[pdb_id], "reflections": mtz_files[pdb_id]} - - pytest.skip("No matching CIF/MTZ pairs found in test data") - - -@pytest.fixture(scope="session") -def all_structure_pairs(cif_dir, mtz_dir): - """Return all matching pairs of CIF models and MTZ reflections.""" - cif_files = {f.stem: f for f in cif_dir.glob("*.cif")} - mtz_files = {f.stem: f for f in mtz_dir.glob("*.mtz")} - - common_ids = set(cif_files.keys()) & set(mtz_files.keys()) - - if not common_ids: - pytest.skip("No matching CIF/MTZ pairs found in test data") - - return [ - {"pdb_id": pdb_id, "model": cif_files[pdb_id], "reflections": mtz_files[pdb_id]} - for pdb_id in sorted(common_ids) - ] - - -@pytest.fixture(scope="session") -def all_cif_files(cif_dir): - """Return all available CIF test structure files.""" - cif_files = sorted(cif_dir.glob("*.cif")) - if not cif_files: - pytest.skip("No CIF files found in test data directory") - return cif_files - - -@pytest.fixture(scope="session") -def all_test_structures(all_structure_pairs): - """Return all loaded model/data pairs for comprehensive testing.""" - from torchref.model.model import Model - from torchref.io import ReflectionData - - structures = [] - for pair in all_structure_pairs: - try: - model = Model() - model.load_cif(str(pair["model"])) - - data = ReflectionData() - data.load_mtz(str(pair["reflections"])) - - structures.append({ - "pdb_id": pair["pdb_id"], - "model": model, - "data": data, - "model_path": pair["model"], - "data_path": pair["reflections"] - }) - except Exception: - # Skip structures that fail to load - continue - - if not structures: - pytest.skip("No structures could be loaded") - - return structures - - -@pytest.fixture(scope="session") -def monomer_library_path(project_root): - """Get path to the monomer library as a string. - - Returns - ------- - str - Absolute path to the external_monomer_library directory. - """ - lib_path = project_root / "external_monomer_library" - if not lib_path.exists(): - pytest.skip("Monomer library not found") - return str(lib_path) - - -# ============================================================================= -# Real Object Fixtures -# ============================================================================= - -@pytest.fixture -def loaded_model(sample_cif_file): - """Fixture providing a fully loaded Model from a real CIF file.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - return model - - -@pytest.fixture -def loaded_reflection_data(sample_mtz_file): - """Fixture providing fully loaded ReflectionData from a real MTZ file.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - return data - - -@pytest.fixture -def model_and_data(sample_structure_pair): - """Fixture providing matching model and reflection data.""" - from torchref.model.model import Model - from torchref.io import ReflectionData - - model = Model() - model.load_cif(str(sample_structure_pair["model"])) - - data = ReflectionData() - data.load_mtz(str(sample_structure_pair["reflections"])) - - return {"model": model, "data": data} - - -@pytest.fixture -def model_with_symmetry(loaded_model): - """Fixture providing model with initialized symmetry.""" - from torchref.symmetry import SpaceGroup - - sg = SpaceGroup(loaded_model.spacegroup) - return {"model": loaded_model, "symmetry": sg} - - -@pytest.fixture -def initialized_scaler(model_and_data): - """Fixture providing initialized Scaler with model and data.""" - from torchref.scaling.scaler import Scaler - - model = model_and_data["model"] - data = model_and_data["data"] - - scaler = Scaler(model=model, data=data, nbins=10, verbose=0) - return scaler - - -@pytest.fixture -def model_with_restraints(loaded_model): - """Fixture providing model with built restraints.""" - from torchref.topology.restraints import Restraints - - restraints = Restraints( - pdb=loaded_model.pdb, - xyz_fn=loaded_model.xyz, - vdw_radii_fn=loaded_model.get_vdw_radii, - verbose=0 - ) - restraints.build_restraints() - return {"model": loaded_model, "restraints": restraints} - -@pytest.fixture -def double_cpu(): - """float64/complex128 on CPU for the duration of a test; restore afterwards. - - Required rather than cosmetic for anything touching eager structure factors: - ``iso_structure_factor_torched`` casts ``hkl`` to the *global* ``dtypes.float`` - (``torchref/base/direct_summation/isotropic.py:121``), so under the default float32 - config a float64 leaf produces a dtype-mismatched matmul. - - Promoted here from three byte-similar copies in ``tests/unit/test_kernel_fixes.py``, - ``tests/unit/test_gradient_correctness.py`` and - ``tests/integration/test_dtype_config_float64.py``. This version also restores - ``sigma_cutoff_ed``, which none of those did -- so a test that changed the cutoff - leaked it into everything that ran afterwards. - """ - import torchref - from torchref.config import device as _device, dtypes as _dtypes - - f0, c0, d0 = _dtypes.float, _dtypes.complex, _device.current - s0 = torchref.sigma_cutoff_ed.value - _dtypes.float = torch.float64 - _dtypes.complex = torch.complex128 - _device.current = torch.device("cpu") - try: - yield - finally: - _dtypes.float = f0 - _dtypes.complex = c0 - _device.current = d0 - torchref.sigma_cutoff_ed.value = s0 diff --git a/tests/fixtures/README.md b/tests/fixtures/README.md new file mode 100644 index 00000000..da39ba78 --- /dev/null +++ b/tests/fixtures/README.md @@ -0,0 +1,40 @@ +# Fixture ownership + +The root `tests/conftest.py` owns pytest options, markers, capability gating, +and the `pytest_plugins` registry. Put reusable setup in the modules below. +Keep a fixture in its test module when only that module needs it. + +| Module | Responsibility | Visibility / lifetime | +|---|---|---| +| `paths.py` | Repository, bundled-data and optional library paths | All tests; session | +| `files.py` | Sample-file selection and matching structure pairs | All tests; session; no model loading | +| `devices.py` | Configured device, explicit backends, device parametrization | All tests; existing per-fixture scopes | +| `precision.py` | Comparison tolerances and CPU-double reference context | All tests; reference fixture restores state after each test | +| `objects.py` | Mutable models, data, scalers and restraints | All tests; fresh per test except explicitly shared bundles | +| `numerical.py` | Synthetic tensors and factories | Imported only by `unit/conftest.py`; function | +| `functional.py` | `shared_model`, `shared_model_ft`, `shared_reflection_data` | Imported only by `functional/conftest.py`; module | + +Existing fixture names remain available without imports in tests. Import reusable +helpers from their defining module, never from the root `conftest.py`. Subtree +conftests import fixture functions explicitly; register shared plugins only at +the root so pytest also works when invoked from a subdirectory. + +Use `shared_*` only for read-only checks. They capture the package configuration +at module setup and may populate derived caches. Do not move them, change their +parameters, tables, masks or grids, backpropagate through them, or use them in +tests that switch global configuration. A target or scaler can mutate a model it +borrows, so a shared model must not be passed to such an operation. + +Tests that verify loading must invoke the loader themselves. Tests of mutation, +device movement, or empty caches use fresh objects. `loaded_model`, +`loaded_model_ft`, `loaded_reflection_data`, and their composed fixtures in `objects.py` provide +fresh mutable objects per test. The explicitly shared session bundles in that +module retain their documented ownership contracts. + +Use `cpu_double_precision()` to scope an explicit numerical reference, or request +`double_cpu` for a single test. The structure-factor package uses the same context +at package scope; both usages restore dtype, device, and density cutoff on exit. + +This separation preserves the existing numerical-factory allocation policy and +test-selection policy. Those policies are independent of fixture registration and +scope, and can be revised in their respective modules. diff --git a/tests/fixtures/__init__.py b/tests/fixtures/__init__.py new file mode 100644 index 00000000..660fb98a --- /dev/null +++ b/tests/fixtures/__init__.py @@ -0,0 +1,5 @@ +"""Provide pytest fixtures grouped by responsibility. + +Root conftest registers shared plugins. Unit and functional conftests explicitly +import their scoped fixtures; this package deliberately re-exports none. +""" diff --git a/tests/fixtures/devices.py b/tests/fixtures/devices.py new file mode 100644 index 00000000..579daeff --- /dev/null +++ b/tests/fixtures/devices.py @@ -0,0 +1,119 @@ +"""Provide explicit backend fixtures and the configured default device. + +Capability probes are shared with collection hooks and structure-factor cases. +Device parametrization is constructed at import so collection can see its marks. +""" + +import pytest +import torch + + +def _cuda_available() -> bool: + return torch.cuda.is_available() + + +def _mps_available() -> bool: + return hasattr(torch.backends, "mps") and torch.backends.mps.is_available() + + +def _accelerator() -> "torch.device | None": + """The canonical accelerator this host can actually use, or ``None``. + + Indices are filled in (``cuda:0`` / ``mps:0``) so the value compares equal + to a device read back off a real tensor -- ``torch.device('mps')`` and + ``torch.device('mps:0')`` are *not* equal even though they name the same + physical device. + """ + if _cuda_available(): + return torch.device("cuda", torch.cuda.current_device()) + if _mps_available(): + return torch.device("mps", 0) + return None + + +@pytest.fixture(scope="session") +def cpu_device() -> torch.device: + """CPU torch device.""" + return torch.device("cpu") + + +@pytest.fixture(scope="session") +def gpu_device() -> torch.device: + """Select CUDA, then MPS, for tests marked ``gpu``. + + Skip if neither backend is available. Use ``cuda_device`` or ``mps_device`` + when the test exercises a backend-specific contract. + """ + accel = _accelerator() + if accel is None: + pytest.skip("No accelerator (CUDA or MPS) on this host") + return accel + + +@pytest.fixture(scope="session") +def cuda_device() -> torch.device: + """Canonical CUDA device for ``cuda``-marked tests. + + Deliberately unguarded. What runs is decided by the ``cuda`` marker in + :func:`pytest_collection_modifyitems` and nowhere else, so this fixture does + not re-check availability: on a host without CUDA the test is *meant* to + error with the real backend error rather than be quietly skipped here. + """ + return torch.device("cuda", 0) + + +@pytest.fixture(scope="session") +def mps_device() -> torch.device: + """Canonical MPS device for ``mps``-marked tests. + + Unguarded for the same reason as :func:`cuda_device` -- the ``mps`` marker + owns the decision. + """ + return torch.device("mps", 0) + + +# Built at import time so the ``gpu`` mark is attached during *collection*. +# Adding it later (e.g. via ``request.node.add_marker`` inside the fixture) is +# too late for ``pytest_collection_modifyitems`` to gate on. +_DEVICE_PARAMS = [pytest.param(torch.device("cpu"), id="cpu")] +_ACCELERATOR = _accelerator() +if _ACCELERATOR is not None: + _DEVICE_PARAMS.append( + pytest.param( + _ACCELERATOR, + id=_ACCELERATOR.type, + # Backend-specific mark, so a CUDA-less host skips the cuda leg and + # a non-Mac skips the mps leg, each with an accurate reason. + marks=getattr(pytest.mark, _ACCELERATOR.type), + ) + ) + + +@pytest.fixture(params=_DEVICE_PARAMS) +def any_device(request: pytest.FixtureRequest) -> torch.device: + """Every device this host can actually use, one test run per device. + + The CPU leg always runs. An available accelerator runs automatically and + carries its backend-specific marker. No accelerator leg is created on a + CPU-only host. + """ + return request.param + + +@pytest.fixture +def device(request: pytest.FixtureRequest) -> torch.device: + """Default test device. + + Uses the package-wide auto-detected default (``torchref.device.current``) + so tests run on whichever device the user's machine resolved to at + import time: cuda -> mps -> cpu. Tests marked ``@pytest.mark.cuda_only`` + are skipped when CUDA is not available. + """ + from torchref.config import get_default_device + + markers = {m.name for m in request.node.iter_markers()} + if "cuda_only" in markers and not torch.cuda.is_available(): + pytest.skip("Test requires CUDA") + if "gpu" in markers and not (_cuda_available() or _mps_available()): + pytest.skip("No GPU (CUDA or MPS) available") + return get_default_device() diff --git a/tests/fixtures/files.py b/tests/fixtures/files.py new file mode 100644 index 00000000..d248775e --- /dev/null +++ b/tests/fixtures/files.py @@ -0,0 +1,98 @@ +"""Select sample paths and matching model/reflection pairs without loading them.""" + +from pathlib import Path + +import pytest + + +@pytest.fixture(scope="session") +def sample_cif_file(cif_dir: Path) -> Path: + """Return a sample CIF file for testing.""" + cif_file = cif_dir / "1DAW.cif" + if cif_file.exists(): + return cif_file + # Try any available CIF file + cif_files = list(cif_dir.glob("*.cif")) + if cif_files: + return cif_files[0] + pytest.skip("No CIF files found in test data") + + +@pytest.fixture(scope="session") +def sample_mtz_file(mtz_dir: Path) -> Path: + """Return a sample MTZ file for testing.""" + mtz_file = mtz_dir / "1DAW.mtz" + if mtz_file.exists(): + return mtz_file + # Try any available MTZ file + mtz_files = list(mtz_dir.glob("*.mtz")) + if mtz_files: + return mtz_files[0] + pytest.skip("No MTZ files found in test data") + + +@pytest.fixture(scope="session") +def sample_pdb_file(pdb_dir: Path) -> Path: + """Return a sample PDB file for testing.""" + pdb_files = sorted(pdb_dir.glob("*.pdb")) + if not pdb_files: + pytest.skip("No PDB files found in test data directory") + return pdb_files[0] + + +@pytest.fixture(scope="session") +def sample_structure_factor_cif(cif_sf_dir: Path) -> Path: + """Return a sample structure factor CIF file.""" + sf_files = sorted(cif_sf_dir.glob("*.cif")) + if not sf_files: + pytest.skip("No structure factor CIF files found") + return sf_files[0] + + +@pytest.fixture(scope="session") +def sample_structure_pair(cif_dir: Path, mtz_dir: Path) -> dict[str, Path]: + """Return a matching pair of CIF model and MTZ reflections.""" + # Try to find matching files + pdb_id = "1DAW" + cif_file = cif_dir / f"{pdb_id}.cif" + mtz_file = mtz_dir / f"{pdb_id}.mtz" + + if cif_file.exists() and mtz_file.exists(): + return {"model": cif_file, "reflections": mtz_file} + + # Try to find any matching pair + cif_files = {f.stem: f for f in cif_dir.glob("*.cif")} + mtz_files = {f.stem: f for f in mtz_dir.glob("*.mtz")} + + common_ids = set(cif_files.keys()) & set(mtz_files.keys()) + if common_ids: + pdb_id = min(common_ids) + return {"model": cif_files[pdb_id], "reflections": mtz_files[pdb_id]} + + pytest.skip("No matching CIF/MTZ pairs found in test data") + + +@pytest.fixture(scope="session") +def all_structure_pairs(cif_dir: Path, mtz_dir: Path) -> list[dict[str, Path | str]]: + """Return all matching pairs of CIF models and MTZ reflections.""" + cif_files = {f.stem: f for f in cif_dir.glob("*.cif")} + mtz_files = {f.stem: f for f in mtz_dir.glob("*.mtz")} + + common_ids = set(cif_files.keys()) & set(mtz_files.keys()) + + if not common_ids: + pytest.skip("No matching CIF/MTZ pairs found in test data") + + return [ + {"pdb_id": pdb_id, "model": cif_files[pdb_id], "reflections": mtz_files[pdb_id]} + for pdb_id in sorted(common_ids) + ] + + +@pytest.fixture(scope="session") +def all_cif_files(cif_dir: Path) -> list[Path]: + """Return all available CIF test structure files.""" + cif_files = sorted(cif_dir.glob("*.cif")) + if not cif_files: + pytest.skip("No CIF files found in test data directory") + return cif_files diff --git a/tests/fixtures/functional.py b/tests/fixtures/functional.py new file mode 100644 index 00000000..18168242 --- /dev/null +++ b/tests/fixtures/functional.py @@ -0,0 +1,42 @@ +"""Share loaded objects within a functional module for read-only checks. + +These fixtures capture the configured dtype/device at module setup. Callers may +populate derived caches but must not change parameters, tables, grids, masks, +device, or configuration. Tests of loading, mutation, and empty caches construct +fresh objects instead. No loaded objects are shared across test modules. +""" + +from __future__ import annotations + +from pathlib import Path +from typing import TYPE_CHECKING + +import pytest + +if TYPE_CHECKING: + from torchref.io import ReflectionData + from torchref.model import Model, ModelFT + + +@pytest.fixture(scope="module") +def shared_model(sample_cif_file: Path) -> Model: + """Load the sample CIF once per module for read-only atomic-model checks.""" + from torchref.model import Model + + return Model(verbose=0).load_cif(str(sample_cif_file)) + + +@pytest.fixture(scope="module") +def shared_model_ft(sample_cif_file: Path) -> ModelFT: + """Load a read-only Fourier model with a 2 Å resolution limit per module.""" + from torchref.model import ModelFT + + return ModelFT(max_res=2.0, verbose=0).load_cif(str(sample_cif_file)) + + +@pytest.fixture(scope="module") +def shared_reflection_data(sample_mtz_file: Path) -> ReflectionData: + """Load the sample MTZ once per module for read-only reflection checks.""" + from torchref.io import ReflectionData + + return ReflectionData().load_mtz(str(sample_mtz_file)) diff --git a/tests/fixtures/numerical.py b/tests/fixtures/numerical.py new file mode 100644 index 00000000..c0d54fc3 --- /dev/null +++ b/tests/fixtures/numerical.py @@ -0,0 +1,195 @@ +"""Generate small synthetic numerical inputs for unit tests. + +Imported by the unit conftest only. Factories return fresh CPU tensors on every +call, using TorchRef's numeric dtypes. They reset the global NumPy random seed; +``random_seed`` also resets PyTorch's seed. Accelerator coverage requires an +explicit move by the caller under this allocation policy. +""" + +from collections.abc import Callable + +import numpy as np +import pytest +import torch + +from torchref.config import dtypes + + +@pytest.fixture +def random_seed() -> int: + """Set random seed for reproducibility.""" + seed = 42 + np.random.seed(seed) + torch.manual_seed(seed) + return seed + + +@pytest.fixture +def random_coordinates() -> Callable[..., torch.Tensor]: + """Return a factory for Cartesian coordinates of shape (n_atoms, 3) in Å.""" + + def _generate(n_atoms: int = 10, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + return torch.tensor(np.random.rand(n_atoms, 3) * 10, dtype=dtypes.float) + + return _generate + + +@pytest.fixture +def random_fractional_coordinates() -> Callable[..., torch.Tensor]: + """Return a factory for fractional coordinates (n_atoms, 3) in [0, 1).""" + + def _generate(n_atoms: int = 10, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + return torch.tensor(np.random.rand(n_atoms, 3), dtype=dtypes.float) + + return _generate + + +@pytest.fixture +def random_adp() -> Callable[..., torch.Tensor]: + """Return a factory for isotropic B-factors (n_atoms,) in [10, 60) Ų.""" + + def _generate(n_atoms: int = 10, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + return torch.tensor(np.random.rand(n_atoms) * 50 + 10, dtype=dtypes.float) + + return _generate + + +@pytest.fixture +def random_occupancies() -> Callable[..., torch.Tensor]: + """Return a factory for dimensionless occupancies (n_atoms,) in [0.5, 1).""" + + def _generate(n_atoms: int = 10, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + return torch.tensor(np.random.rand(n_atoms) * 0.5 + 0.5, dtype=dtypes.float) + + return _generate + + +@pytest.fixture +def mock_cell() -> torch.Tensor: + """Return an orthorhombic cell (6,), lengths in Å and angles in degrees.""" + return torch.tensor([50.0, 60.0, 70.0, 90.0, 90.0, 90.0], dtype=dtypes.float) + + +@pytest.fixture +def mock_cell_triclinic() -> torch.Tensor: + """Return a triclinic cell (6,), lengths in Å and angles in degrees.""" + return torch.tensor([40.0, 50.0, 60.0, 70.0, 80.0, 85.0], dtype=dtypes.float) + + +@pytest.fixture +def mock_hkl_indices() -> Callable[..., torch.Tensor]: + """Return a factory for floating HKL triples (n_kept, 3), excluding the origin. + + The output uses ``dtypes.float``; ``n_kept`` can be less than the requested + reflection count when the origin is sampled. + """ + + def _generate( + n_reflections: int = 100, max_index: int = 10, seed: int = 42 + ) -> torch.Tensor: + np.random.seed(seed) + h = np.random.randint(-max_index, max_index + 1, n_reflections) + k = np.random.randint(-max_index, max_index + 1, n_reflections) + l = np.random.randint(-max_index, max_index + 1, n_reflections) + # Exclude (0,0,0) + mask = ~((h == 0) & (k == 0) & (l == 0)) + h, k, l = h[mask], k[mask], l[mask] + return torch.tensor(np.stack([h, k, l], axis=1), dtype=dtypes.float) + + return _generate + + +@pytest.fixture +def mock_structure_factors() -> Callable[..., torch.Tensor]: + """Return a factory for complex structure factors (n_reflections,) in electrons.""" + + def _generate(n_reflections: int = 100, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + real = np.random.randn(n_reflections) * 100 + imag = np.random.randn(n_reflections) * 100 + return torch.tensor(real + 1j * imag, dtype=dtypes.complex) + + return _generate + + +@pytest.fixture +def mock_F_obs() -> Callable[..., torch.Tensor]: + """Return a factory for observed amplitudes (n_reflections,) in electrons.""" + + def _generate(n_reflections: int = 100, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + # Positive values with realistic distribution + return torch.tensor( + np.abs(np.random.randn(n_reflections) * 100) + 10, dtype=dtypes.float + ) + + return _generate + + +@pytest.fixture +def mock_F_sigma() -> Callable[..., torch.Tensor]: + """Return a factory for amplitude uncertainties (n_reflections,) in electrons.""" + + def _generate(n_reflections: int = 100, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + return torch.tensor( + np.abs(np.random.randn(n_reflections) * 5) + 1, dtype=dtypes.float + ) + + return _generate + + +@pytest.fixture +def mock_aniso_u() -> Callable[..., torch.Tensor]: + """Return a factory for Cartesian U tensors (n_atoms, 6) in Ų. + + Components are ordered U11, U22, U33, U12, U13, U23. + """ + + def _generate(n_atoms: int = 10, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + # Diagonal elements (positive) + u11 = np.random.rand(n_atoms) * 0.05 + 0.02 + u22 = np.random.rand(n_atoms) * 0.05 + 0.02 + u33 = np.random.rand(n_atoms) * 0.05 + 0.02 + # Off-diagonal elements (can be negative, smaller magnitude) + u12 = (np.random.rand(n_atoms) - 0.5) * 0.02 + u13 = (np.random.rand(n_atoms) - 0.5) * 0.02 + u23 = (np.random.rand(n_atoms) - 0.5) * 0.02 + return torch.tensor( + np.stack([u11, u22, u33, u12, u13, u23], axis=1), dtype=dtypes.float + ) + + return _generate + + +@pytest.fixture +def mock_scattering_factors() -> Callable[..., torch.Tensor]: + """Return a factory for scattering factors (n_reflections, n_atoms) in electrons.""" + + def _generate( + n_reflections: int = 100, n_atoms: int = 10, seed: int = 42 + ) -> torch.Tensor: + np.random.seed(seed) + # Decreasing with resolution (approximate) + return torch.tensor( + np.random.rand(n_reflections, n_atoms) * 5 + 1, dtype=dtypes.float + ) + + return _generate + + +@pytest.fixture +def mock_weights() -> Callable[..., torch.Tensor]: + """Return a factory for dimensionless weights (n_atoms, 1) summing to one.""" + + def _generate(n_atoms: int = 10, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + weights = np.random.rand(n_atoms) + return torch.tensor(weights / weights.sum(), dtype=dtypes.float).reshape(-1, 1) + + return _generate diff --git a/tests/fixtures/objects.py b/tests/fixtures/objects.py new file mode 100644 index 00000000..34fb00c8 --- /dev/null +++ b/tests/fixtures/objects.py @@ -0,0 +1,163 @@ +"""Load fresh mutable models, reflection data, scalers, and restraints. + +Function-scoped fixtures isolate test mutations. The explicitly shared device +bundle caches one model per device and must be treated as read-only by callers. +""" + +from __future__ import annotations + +from pathlib import Path +from typing import TYPE_CHECKING, Any + +import pytest +import torch + +if TYPE_CHECKING: + from torchref.io import ReflectionData + from torchref.model import Model, ModelFT + from torchref.scaling import Scaler + + +@pytest.fixture +def loaded_model(sample_cif_file: Path) -> Model: + """Load a fresh mutable Model from the sample CIF file.""" + from torchref.model.model import Model + + model = Model() + model.load_cif(str(sample_cif_file)) + return model + + +@pytest.fixture +def loaded_model_ft(sample_cif_file: Path) -> ModelFT: + """Load a fresh mutable Fourier model with a 2 Å resolution limit.""" + from torchref.model import ModelFT + + return ModelFT(max_res=2.0, verbose=0).load_cif(str(sample_cif_file)) + + +@pytest.fixture +def loaded_reflection_data(sample_mtz_file: Path) -> ReflectionData: + """Load fresh mutable reflection data from the sample MTZ file.""" + from torchref.io import ReflectionData + + data = ReflectionData() + data.load_mtz(str(sample_mtz_file)) + return data + + +@pytest.fixture +def model_and_data(sample_structure_pair: dict[str, Path]) -> dict[str, Any]: + """Load a fresh matching model and reflection dataset.""" + from torchref.io import ReflectionData + from torchref.model.model import Model + + model = Model() + model.load_cif(str(sample_structure_pair["model"])) + + data = ReflectionData() + data.load_mtz(str(sample_structure_pair["reflections"])) + + return {"model": model, "data": data} + + +@pytest.fixture +def model_with_symmetry(loaded_model: Model) -> dict[str, Any]: + """Pair a fresh model with initialized symmetry.""" + from torchref.symmetry import SpaceGroup + + sg = SpaceGroup(loaded_model.spacegroup) + return {"model": loaded_model, "symmetry": sg} + + +@pytest.fixture +def initialized_scaler(model_and_data: dict[str, Any]) -> Scaler: + """Build a scaler around a fresh matching model and dataset.""" + from torchref.scaling.scaler import Scaler + + model = model_and_data["model"] + data = model_and_data["data"] + + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) + return scaler + + +@pytest.fixture +def model_with_restraints(loaded_model: Model) -> dict[str, Any]: + """Build restraints around a fresh model.""" + from torchref.topology.restraints import Restraints + + restraints = Restraints( + pdb=loaded_model.pdb, + xyz_fn=loaded_model.xyz, + vdw_radii_fn=loaded_model.get_vdw_radii, + verbose=0, + ) + restraints.build_restraints() + return {"model": loaded_model, "restraints": restraints} + + +@pytest.fixture(scope="session") +def all_test_structures( + all_structure_pairs: list[dict[str, Any]], +) -> list[dict[str, Any]]: + """Return all loaded model/data pairs for comprehensive testing.""" + from torchref.io import ReflectionData + from torchref.model.model import Model + + structures = [] + for pair in all_structure_pairs: + try: + model = Model() + model.load_cif(str(pair["model"])) + + data = ReflectionData() + data.load_mtz(str(pair["reflections"])) + + structures.append( + { + "pdb_id": pair["pdb_id"], + "model": model, + "data": data, + "model_path": pair["model"], + "data_path": pair["reflections"], + } + ) + except Exception: + # Skip structures that fail to load + continue + + if not structures: + pytest.skip("No structures could be loaded") + + return structures + + +@pytest.fixture(scope="session") +def _device_model_cache() -> dict: + """``{device_str: ModelFT}`` built at most once per device, per session.""" + return {} + + +@pytest.fixture +def device_model_bundle( + _device_model_cache: dict[str, ModelFT], pdb_dir: Path, any_device: torch.device +) -> dict[str, ModelFT]: + """Borrow a session-shared model on the requested device. + + Notes + ----- + Treat the model as read-only, including when a target borrows it. Moving a + target can move its model too; tests of movement need a fresh model. + """ + key = str(any_device) + if key not in _device_model_cache: + pdb = pdb_dir / "1DAW.pdb" + if not pdb.exists(): + pytest.skip("1DAW.pdb fixture not present") + from torchref.model import ModelFT + + _device_model_cache[key] = ModelFT(device=any_device, verbose=0).load_pdb( + str(pdb) + ) + return {"model": _device_model_cache[key]} diff --git a/tests/fixtures/paths.py b/tests/fixtures/paths.py new file mode 100644 index 00000000..16ac98d4 --- /dev/null +++ b/tests/fixtures/paths.py @@ -0,0 +1,68 @@ +"""Locate bundled test data and optional monomer-library installations.""" + +from pathlib import Path + +import pytest + + +@pytest.fixture(scope="session") +def tests_root() -> Path: + """Return the root of the test tree.""" + return Path(__file__).resolve().parents[1] + + +@pytest.fixture(scope="session") +def project_root() -> Path: + """Return the project root.""" + return Path(__file__).resolve().parents[2] + + +@pytest.fixture(scope="session") +def test_files_dir(tests_root: Path) -> Path: + """Return the bundled test-data directory.""" + return tests_root / "files" + + +@pytest.fixture(scope="session") +def cif_dir(test_files_dir: Path) -> Path: + """Return the model CIF directory.""" + return test_files_dir / "cif" + + +@pytest.fixture(scope="session") +def cif_sf_dir(test_files_dir: Path) -> Path: + """Return the structure-factor CIF directory.""" + return test_files_dir / "cif_sf" + + +@pytest.fixture(scope="session") +def mtz_dir(test_files_dir: Path) -> Path: + """Return the MTZ reflection directory.""" + return test_files_dir / "mtz" + + +@pytest.fixture(scope="session") +def pdb_dir(test_files_dir: Path) -> Path: + """Return the model PDB directory.""" + return test_files_dir / "pdb" + + +@pytest.fixture(scope="session") +def external_monomer_library(project_root: Path) -> Path: + """Return the optional external monomer-library path without checking it.""" + return project_root / "external_monomer_library" + + +@pytest.fixture(scope="session") +def monomer_library_path(project_root: Path) -> str: + """Get path to the monomer library as a string. + + Returns + ------- + str + Absolute path to the external_monomer_library directory. + """ + lib_path = project_root / "external_monomer_library" + if not lib_path.exists(): + pytest.skip("Monomer library not found") + return str(lib_path) diff --git a/tests/fixtures/precision.py b/tests/fixtures/precision.py new file mode 100644 index 00000000..6769c4de --- /dev/null +++ b/tests/fixtures/precision.py @@ -0,0 +1,51 @@ +"""Scope numerical reference configuration and expose comparison tolerances.""" + +from collections.abc import Iterator +from contextlib import contextmanager + +import pytest +import torch + +import torchref +from torchref.config import device, dtypes + + +@contextmanager +def cpu_double_precision() -> Iterator[None]: + """Temporarily select CPU float64/complex128 for numerical references. + + Notes + ----- + Mutate process-wide TorchRef defaults, not PyTorch factory defaults. Restore + float/complex dtype, device, and density cutoff even when the body raises. + Objects allocated inside the context retain their own dtype and device. + """ + original = dtypes.float, dtypes.complex, device.current + cutoff = torchref.sigma_cutoff_ed.value + dtypes.float = torch.float64 + dtypes.complex = torch.complex128 + device.current = torch.device("cpu") + try: + yield + finally: + dtypes.float, dtypes.complex, device.current = original + torchref.sigma_cutoff_ed.value = cutoff + + +@pytest.fixture +def double_cpu() -> Iterator[None]: + """Use CPU double precision for one test and restore configuration afterward.""" + with cpu_double_precision(): + yield + + +@pytest.fixture +def rtol() -> float: + """Relative tolerance for floating point comparisons.""" + return 1e-5 + + +@pytest.fixture +def atol() -> float: + """Absolute tolerance for floating point comparisons.""" + return 1e-8 diff --git a/tests/functional/conftest.py b/tests/functional/conftest.py index 58e2352e..e7a5b0ed 100644 --- a/tests/functional/conftest.py +++ b/tests/functional/conftest.py @@ -1,6 +1,7 @@ -""" -Functional test fixtures. +"""Expose module-shared read-only fixtures to functional tests.""" -All shared fixtures (sample files, loaded models, scalers, restraints, etc.) -are defined in the root tests/conftest.py and are automatically available here. -""" +from tests.fixtures.functional import ( # noqa: F401 + shared_model, + shared_model_ft, + shared_reflection_data, +) diff --git a/tests/functional/test_io_functional.py b/tests/functional/test_io_functional.py index 9cdef611..8779ea39 100644 --- a/tests/functional/test_io_functional.py +++ b/tests/functional/test_io_functional.py @@ -6,7 +6,6 @@ import pytest import torch -import numpy as np class TestCIFReadingFunctional: @@ -16,49 +15,45 @@ class TestCIFReadingFunctional: def test_load_multiple_cif_files(self, cif_dir): """Test loading multiple CIF files successfully.""" from torchref.model.model import Model - + cif_files = list(cif_dir.glob("*.cif")) assert len(cif_files) > 0, "No CIF files found in test directory" - + for cif_file in cif_files: model = Model() model.load_cif(str(cif_file)) - + # Each file should load with atoms n_atoms = model.xyz().shape[0] assert n_atoms > 0, f"No atoms loaded from {cif_file}" - + # Should have cell parameters assert model.cell is not None assert len(model.cell) == 6 @pytest.mark.integration - def test_cif_atom_properties(self, sample_cif_file): + def test_cif_atom_properties(self, shared_model): """Test that atom properties are correctly loaded from CIF.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - + model = shared_model + pdb = model.pdb - + # Check required columns exist - required_cols = ['x', 'y', 'z', 'element', 'resname', 'chainid', 'resseq'] + required_cols = ["x", "y", "z", "element", "resname", "chainid", "resseq"] for col in required_cols: - assert col in pdb.columns or col.upper() in pdb.columns, f"Missing column: {col}" + assert col in pdb.columns or col.upper() in pdb.columns, ( + f"Missing column: {col}" + ) @pytest.mark.integration - def test_cif_element_types(self, sample_cif_file): + def test_cif_element_types(self, shared_model): """Test that element types are properly assigned.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - - elements = model.pdb['element'].unique() - + model = shared_model + + elements = model.pdb["element"].unique() + # Should have common protein elements - common_elements = ['C', 'N', 'O', 'S'] + common_elements = ["C", "N", "O", "S"] found_any = any(elem in elements for elem in common_elements) assert found_any, "No common elements found" @@ -70,58 +65,52 @@ class TestMTZReadingFunctional: def test_load_multiple_mtz_files(self, mtz_dir): """Test loading multiple MTZ files successfully.""" from torchref.io import ReflectionData - + mtz_files = list(mtz_dir.glob("*.mtz")) assert len(mtz_files) > 0, "No MTZ files found in test directory" - + for mtz_file in mtz_files: data = ReflectionData() data.load_mtz(str(mtz_file)) - + # Each file should load with reflections n_refl = data.hkl.shape[0] assert n_refl > 0, f"No reflections loaded from {mtz_file}" - + # Should have cell parameters assert data.cell is not None @pytest.mark.integration - def test_mtz_data_properties(self, sample_mtz_file): + def test_mtz_data_properties(self, shared_reflection_data): """Test that MTZ data properties are correctly loaded.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - + data = shared_reflection_data + # Check HKL indices are integers or can be converted hkl = data.hkl assert hkl.shape[1] == 3, "HKL should have 3 columns" - + # Check F values are loaded assert data.F is not None assert data.F.shape[0] == hkl.shape[0] - + # Check sigma values - if hasattr(data, 'F_sigma') and data.F_sigma is not None: + if hasattr(data, "F_sigma") and data.F_sigma is not None: assert data.F_sigma.shape[0] == hkl.shape[0] @pytest.mark.integration - def test_mtz_resolution_range(self, sample_mtz_file): + def test_mtz_resolution_range(self, shared_reflection_data): """Test that resolution range is computed correctly.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - + data = shared_reflection_data + # Check if resolution data is available - if hasattr(data, 'd') and data.d is not None: + if hasattr(data, "d") and data.d is not None: d_min = data.d.min().item() d_max = data.d.max().item() - + # Resolution should be positive assert d_min > 0 assert d_max > d_min - + # Typical protein data: 0.8 - 500 Å assert d_min > 0.5 assert d_max < 1000 @@ -134,16 +123,16 @@ class TestSFCIFReadingFunctional: def test_load_sf_cif(self, cif_sf_dir): """Test loading structure factor CIF files.""" from torchref.io import ReflectionData - + sf_files = list(cif_sf_dir.glob("*.cif")) if not sf_files: pytest.skip("No SF-CIF files found") - + for sf_file in sf_files: data = ReflectionData() try: data.load_cif(str(sf_file)) - + # Should have loaded reflections if data.hkl is not None: assert data.hkl.shape[0] > 0 @@ -158,40 +147,42 @@ class TestDataConsistencyFunctional: @pytest.mark.integration def test_cell_parameters_match(self, sample_structure_pair): """Test that cell parameters match between model and reflections.""" - from torchref.model.model import Model from torchref.io import ReflectionData - + from torchref.model.model import Model + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + model_cell = model.cell data_cell = data.cell - + if model_cell is not None and data_cell is not None: # Convert to tensors if needed if not isinstance(model_cell, torch.Tensor): model_cell = torch.tensor(model_cell) if not isinstance(data_cell, torch.Tensor): data_cell = torch.tensor(data_cell) - + # Cell parameters should be similar (1% tolerance) - assert torch.allclose(model_cell.float(), data_cell.float(), rtol=0.01, atol=0.1) + assert torch.allclose( + model_cell.float(), data_cell.float(), rtol=0.01, atol=0.1 + ) @pytest.mark.integration def test_spacegroup_consistency(self, sample_structure_pair): """Test that spacegroup is consistent.""" - from torchref.model.model import Model from torchref.io import ReflectionData - + from torchref.model.model import Model + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + # Both should have spacegroup defined assert model.spacegroup is not None @@ -203,13 +194,13 @@ class TestDataBinningFunctional: def test_get_bins(self, sample_mtz_file): """Test resolution binning of reflection data.""" from torchref.io import ReflectionData - + data = ReflectionData() data.load_mtz(str(sample_mtz_file)) - + # Get bins bins, n_bins = data.get_bins(n_bins=10) - + assert bins is not None assert bins.shape[0] == data.hkl.shape[0] assert bins.min() >= 0 @@ -219,20 +210,20 @@ def test_get_bins(self, sample_mtz_file): def test_mean_res_per_bin(self, sample_mtz_file): """Test mean resolution per bin calculation.""" from torchref.io import ReflectionData - + data = ReflectionData() data.load_mtz(str(sample_mtz_file)) - + # Get bins first bins, n_bins = data.get_bins(n_bins=10) - + # Get mean resolution per bin - if hasattr(data, 'mean_res_per_bin'): + if hasattr(data, "mean_res_per_bin"): mean_res = data.mean_res_per_bin() - + assert mean_res is not None assert len(mean_res) == n_bins - + # Mean resolution should decrease with bin index (low res to high res) # or increase (high res to low res) - depends on implementation assert torch.all(torch.isfinite(mean_res)) @@ -242,16 +233,13 @@ class TestFrenchWilsonFunctional: """Test French-Wilson conversion with real data.""" @pytest.mark.integration - def test_french_wilson_applied(self, sample_mtz_file): + def test_french_wilson_applied(self, shared_reflection_data): """Test that French-Wilson conversion is applied.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - + data = shared_reflection_data + # After French-Wilson, F values should be non-negative valid_F = data.F[~torch.isnan(data.F)] - + if len(valid_F) > 0: # All valid F values should be >= 0 assert torch.all(valid_F >= 0) @@ -261,32 +249,28 @@ class TestRfreeHandlingFunctional: """Test R-free flag handling.""" @pytest.mark.integration - def test_rfree_flags_loaded(self, sample_mtz_file): + def test_rfree_flags_loaded(self, shared_reflection_data): """Test that R-free flags are loaded or generated.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - + data = shared_reflection_data + # Should have rfree attribute - if hasattr(data, 'rfree') and data.rfree is not None: + if hasattr(data, "rfree") and data.rfree is not None: assert data.rfree.shape[0] == data.hkl.shape[0] - + # Should be boolean or can be converted to boolean - assert data.rfree.dtype == torch.bool or torch.all((data.rfree == 0) | (data.rfree == 1)) + assert data.rfree.dtype == torch.bool or torch.all( + (data.rfree == 0) | (data.rfree == 1) + ) @pytest.mark.integration - def test_rfree_fraction(self, sample_mtz_file): + def test_rfree_fraction(self, shared_reflection_data): """Test R-free set fraction is reasonable.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - if hasattr(data, 'rfree') and data.rfree is not None: + data = shared_reflection_data + + if hasattr(data, "rfree") and data.rfree is not None: # Work set mask (True for work, False for test) work_fraction = data.rfree.float().mean().item() - + # Typically 90-95% work set, 5-10% test set # So work_fraction should be 0.9-0.95 typically assert 0.7 < work_fraction <= 1.0 @@ -296,16 +280,13 @@ class TestMaskHandlingFunctional: """Test reflection mask handling.""" @pytest.mark.integration - def test_masks_method(self, sample_mtz_file): + def test_masks_method(self, shared_reflection_data): """Test masks() method returns valid mask.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - if hasattr(data, 'masks'): + data = shared_reflection_data + + if hasattr(data, "masks"): mask = data.masks() - + assert mask is not None assert mask.shape[0] == data.hkl.shape[0] assert mask.dtype == torch.bool diff --git a/tests/functional/test_model_ft_functional.py b/tests/functional/test_model_ft_functional.py index b8919254..a6abd808 100644 --- a/tests/functional/test_model_ft_functional.py +++ b/tests/functional/test_model_ft_functional.py @@ -4,10 +4,9 @@ These tests exercise the ModelFT class with real crystallographic data, testing the FFT-based structure factor calculation pipeline. """ + import pytest import torch -import numpy as np -from pathlib import Path @pytest.mark.integration @@ -17,7 +16,7 @@ class TestModelFTInitialization: def test_modelft_empty_initialization(self): """Test empty ModelFT initialization.""" from torchref.model.model_ft import ModelFT - + model = ModelFT() assert model is not None assert model.max_res == 1.0 # Default @@ -32,10 +31,10 @@ def test_modelft_with_custom_resolution(self): def test_modelft_load_cif(self, sample_cif_file): """Test loading a CIF file into ModelFT.""" from torchref.model.model_ft import ModelFT - + model = ModelFT(max_res=2.0, verbose=0) model.load_cif(str(sample_cif_file)) - + # Verify basic properties assert model.xyz() is not None assert model.xyz().shape[0] > 0 @@ -59,39 +58,33 @@ def test_modelft_has_gridsize(self, sample_cif_file): class TestModelFTParametrization: """Test ModelFT parametrization with real structures.""" - def test_parametrization_built(self, sample_cif_file): + def test_parametrization_built(self, shared_model_ft): """Test that parametrization is built after loading.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + # Parametrization should be set assert model.parametrization is not None - def test_scattering_factors_available(self, sample_cif_file): + def test_scattering_factors_available(self, shared_model_ft): """Test that scattering factors can be computed.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + # Should be able to access atom properties xyz = model.xyz() assert xyz is not None assert xyz.dtype == torch.float32 or xyz.dtype == torch.float64 -@pytest.mark.integration +@pytest.mark.integration class TestModelFTGridOperations: """Test ModelFT grid operations.""" - def test_setup_grid(self, sample_cif_file): + def test_setup_grid(self, loaded_model_ft): """An explicit grid size overrides the resolution-derived one.""" - from torchref.model.model_ft import ModelFT - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) + model = loaded_model_ft derived = model.grid_shape assert derived is not None and len(derived) == 3 @@ -107,16 +100,15 @@ def test_setup_grid(self, sample_cif_file): class TestModelFTRealSpaceMap: """Test ModelFT real space electron density map construction.""" - def test_get_real_space_grid(self, sample_cif_file): + def test_get_real_space_grid(self, loaded_model_ft): """Test getting real space grid.""" - from torchref.model.model_ft import ModelFT from torchref.base.math_torch import get_real_grid - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) + + # The grid helper moves its Cell in place when targeting CPU. + model = loaded_model_ft assert model.gridsize is not None - grid = get_real_grid(model.cell, max_res=2.0, device='cpu') + grid = get_real_grid(model.cell, max_res=2.0, device="cpu") assert grid is not None assert len(grid.shape) == 4 # Should be 4D (nx, ny, nz, 3) @@ -126,16 +118,14 @@ def test_get_real_space_grid(self, sample_cif_file): class TestModelFTSymmetry: """Test ModelFT symmetry operations.""" - def test_map_symmetry_available(self, sample_cif_file): + def test_map_symmetry_available(self, shared_model_ft): """Test map symmetry is available after loading.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + # Model should have spacegroup after loading assert model.spacegroup is not None - + # The map operator comes from the space group, keyed on the grid shape. gridsize = model.grid_shape assert gridsize is not None @@ -149,21 +139,20 @@ def test_map_symmetry_available(self, sample_cif_file): class TestModelFTStateDictFunctional: """Test ModelFT state dict operations with real data.""" - def test_save_and_load_state_dict(self, sample_cif_file, tmp_path): + def test_save_and_load_state_dict(self, loaded_model_ft, tmp_path): """Test saving and loading state dict.""" from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = loaded_model_ft + original_xyz = model.xyz().clone() - + # Save state dict state_dict = model.state_dict() - + # Create new model and load state model2 = ModelFT(max_res=2.0, verbose=0) - + # We need to ensure proper initialization # For now just verify state_dict works assert state_dict is not None @@ -174,25 +163,23 @@ def test_save_and_load_state_dict(self, sample_cif_file, tmp_path): class TestModelFTForwardPass: """Test ModelFT forward pass (structure factor calculation).""" - def test_forward_method_exists(self, sample_cif_file): + def test_forward_method_exists(self, shared_model_ft): """Test that forward method is available.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + # Check forward method exists - assert hasattr(model, 'forward') + assert hasattr(model, "forward") def test_build_map_method(self, sample_cif_file): """Test build_map method if available.""" from torchref.model.model_ft import ModelFT - + model = ModelFT(max_res=3.0, verbose=0) # Lower res for faster test model.load_cif(str(sample_cif_file)) - + # Check build_map method - if hasattr(model, 'build_map'): + if hasattr(model, "build_map"): # Try to build map try: model.build_map() @@ -209,22 +196,22 @@ class TestModelFTMultipleStructures: def test_modelft_multiple_structures(self, all_structure_pairs): """Test ModelFT works with different structures.""" from torchref.model.model_ft import ModelFT - + tested = 0 for pair in all_structure_pairs[:3]: # Test first 3 try: model = ModelFT(max_res=3.0, verbose=0) model.load_cif(str(pair["model"])) - + # Basic checks assert model.xyz() is not None assert model.xyz().shape[0] > 0 - + tested += 1 except Exception as e: # Some structures may fail to load continue - + assert tested >= 1, "At least one structure should load" @@ -232,27 +219,23 @@ def test_modelft_multiple_structures(self, all_structure_pairs): class TestModelFTCaching: """Test ModelFT caching mechanism.""" - def test_cache_initialization(self, sample_cif_file): + def test_cache_initialization(self, loaded_model_ft): """Test that CachedForwardMixin cache starts empty.""" - from torchref.model.model_ft import ModelFT - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) + model = loaded_model_ft # Mixin cache should start empty (lazily initialized) assert getattr(model, "_fwd_cached_output", None) is None - def test_cache_usage(self, sample_cif_file): + def test_cache_usage(self, shared_model_ft): """Test that cache can be used for computations.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + # Access xyz twice - should use caching xyz1 = model.xyz() xyz2 = model.xyz() - + # Should return same tensor assert torch.allclose(xyz1, xyz2) @@ -261,14 +244,12 @@ def test_cache_usage(self, sample_cif_file): class TestModelFTCoordinateOperations: """Test ModelFT coordinate operations.""" - def test_cartesian_to_fractional(self, sample_cif_file): + def test_cartesian_to_fractional(self, shared_model_ft): """Test coordinate conversion.""" - from torchref.model.model_ft import ModelFT from torchref.base.math_torch import cartesian_to_fractional_torch - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + xyz = model.xyz() cell = model.cell @@ -278,17 +259,15 @@ def test_cartesian_to_fractional(self, sample_cif_file): # Fractional coords should be bounded (mostly between 0 and 1) assert frac.shape == xyz.shape - def test_fractional_to_cartesian(self, sample_cif_file): + def test_fractional_to_cartesian(self, shared_model_ft): """Test fractional to cartesian conversion.""" - from torchref.model.model_ft import ModelFT from torchref.base.math_torch import ( cartesian_to_fractional_torch, - fractional_to_cartesian_torch + fractional_to_cartesian_torch, ) - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + xyz = model.xyz() cell = model.cell @@ -304,24 +283,20 @@ def test_fractional_to_cartesian(self, sample_cif_file): class TestModelFTAnisoHandling: """Test ModelFT handling of anisotropic parameters.""" - def test_access_aniso_atoms(self, sample_cif_file): + def test_access_aniso_atoms(self, shared_model_ft): """Test accessing anisotropic atom information.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + # Check if aniso is available - if hasattr(model, 'get_aniso') or hasattr(model, 'aniso'): + if hasattr(model, "get_aniso") or hasattr(model, "aniso"): # Structure has aniso pass - def test_isotropic_b_factors(self, sample_cif_file): + def test_isotropic_b_factors(self, shared_model_ft): """Test accessing isotropic B-factors.""" - from torchref.model.model_ft import ModelFT - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) + model = shared_model_ft # Get B-factors (now accessed via adp()) b_factors = model.adp() diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index 8bdf5d61..424307e5 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -1,7 +1,6 @@ -""" -Integration test specific fixtures. -Integration tests use real file I/O and test the full pipeline. +"""Use shared fixtures registered by the root conftest for integration tests. -All shared fixtures (sample files, path fixtures, monomer library, etc.) -are defined in the root tests/conftest.py and are automatically available here. +Pipeline-specific fixtures belong in their consuming modules. Mutable loaded +objects from ``tests.fixtures.objects`` are function-scoped unless documented +as explicitly shared. """ diff --git a/tests/unit/conftest.py b/tests/unit/conftest.py index 96ce096c..ecc8cdde 100644 --- a/tests/unit/conftest.py +++ b/tests/unit/conftest.py @@ -1,148 +1,18 @@ -""" -Unit test specific fixtures. -Unit tests should NOT use real file I/O - use mocks or minimal in-memory data. -""" -import pytest -import torch -import numpy as np - -from torchref.config import dtypes - - -@pytest.fixture -def random_seed(): - """Set random seed for reproducibility.""" - seed = 42 - np.random.seed(seed) - torch.manual_seed(seed) - return seed - - -@pytest.fixture -def random_coordinates(): - """Generate random atomic coordinates.""" - def _generate(n_atoms: int = 10, seed: int = 42): - np.random.seed(seed) - return torch.tensor(np.random.rand(n_atoms, 3) * 10, dtype=dtypes.float) - return _generate - - -@pytest.fixture -def random_fractional_coordinates(): - """Generate random fractional coordinates (0-1 range).""" - def _generate(n_atoms: int = 10, seed: int = 42): - np.random.seed(seed) - return torch.tensor(np.random.rand(n_atoms, 3), dtype=dtypes.float) - return _generate - - -@pytest.fixture -def random_adp(): - """Generate random ADPs (atomic displacement parameters, reasonable range 10-60 Ų).""" - def _generate(n_atoms: int = 10, seed: int = 42): - np.random.seed(seed) - return torch.tensor(np.random.rand(n_atoms) * 50 + 10, dtype=dtypes.float) - return _generate - - -@pytest.fixture -def random_occupancies(): - """Generate random occupancies (0-1 range).""" - def _generate(n_atoms: int = 10, seed: int = 42): - np.random.seed(seed) - return torch.tensor(np.random.rand(n_atoms) * 0.5 + 0.5, dtype=dtypes.float) - return _generate - - -@pytest.fixture -def mock_cell(): - """Mock cell parameters [a, b, c, alpha, beta, gamma].""" - return torch.tensor([50.0, 60.0, 70.0, 90.0, 90.0, 90.0], dtype=dtypes.float) - - -@pytest.fixture -def mock_cell_triclinic(): - """Mock triclinic cell parameters.""" - return torch.tensor([40.0, 50.0, 60.0, 70.0, 80.0, 85.0], dtype=dtypes.float) - - -@pytest.fixture -def mock_hkl_indices(): - """Generate mock HKL indices.""" - def _generate(n_reflections: int = 100, max_index: int = 10, seed: int = 42): - np.random.seed(seed) - h = np.random.randint(-max_index, max_index + 1, n_reflections) - k = np.random.randint(-max_index, max_index + 1, n_reflections) - l = np.random.randint(-max_index, max_index + 1, n_reflections) - # Exclude (0,0,0) - mask = ~((h == 0) & (k == 0) & (l == 0)) - h, k, l = h[mask], k[mask], l[mask] - return torch.tensor(np.stack([h, k, l], axis=1), dtype=dtypes.float) - return _generate - - -@pytest.fixture -def mock_structure_factors(): - """Generate mock structure factors (complex).""" - def _generate(n_reflections: int = 100, seed: int = 42): - np.random.seed(seed) - real = np.random.randn(n_reflections) * 100 - imag = np.random.randn(n_reflections) * 100 - return torch.tensor(real + 1j * imag, dtype=dtypes.complex) - return _generate - - -@pytest.fixture -def mock_F_obs(): - """Generate mock observed structure factor amplitudes.""" - def _generate(n_reflections: int = 100, seed: int = 42): - np.random.seed(seed) - # Positive values with realistic distribution - return torch.tensor(np.abs(np.random.randn(n_reflections) * 100) + 10, dtype=dtypes.float) - return _generate - - -@pytest.fixture -def mock_F_sigma(): - """Generate mock sigma values for F_obs.""" - def _generate(n_reflections: int = 100, seed: int = 42): - np.random.seed(seed) - return torch.tensor(np.abs(np.random.randn(n_reflections) * 5) + 1, dtype=dtypes.float) - return _generate - - -@pytest.fixture -def mock_aniso_u(): - """Generate mock anisotropic U tensor components [U11, U22, U33, U12, U13, U23].""" - def _generate(n_atoms: int = 10, seed: int = 42): - np.random.seed(seed) - # Diagonal elements (positive) - u11 = np.random.rand(n_atoms) * 0.05 + 0.02 - u22 = np.random.rand(n_atoms) * 0.05 + 0.02 - u33 = np.random.rand(n_atoms) * 0.05 + 0.02 - # Off-diagonal elements (can be negative, smaller magnitude) - u12 = (np.random.rand(n_atoms) - 0.5) * 0.02 - u13 = (np.random.rand(n_atoms) - 0.5) * 0.02 - u23 = (np.random.rand(n_atoms) - 0.5) * 0.02 - return torch.tensor(np.stack([u11, u22, u33, u12, u13, u23], axis=1), dtype=dtypes.float) - return _generate - - -@pytest.fixture -def mock_scattering_factors(): - """Generate mock scattering factors.""" - def _generate(n_reflections: int = 100, n_atoms: int = 10, seed: int = 42): - np.random.seed(seed) - # Decreasing with resolution (approximate) - return torch.tensor(np.random.rand(n_reflections, n_atoms) * 5 + 1, dtype=dtypes.float) - return _generate - - -@pytest.fixture -def mock_weights(): - """Generate mock weights for atoms.""" - def _generate(n_atoms: int = 10, seed: int = 42): - np.random.seed(seed) - weights = np.random.rand(n_atoms) - return torch.tensor(weights / weights.sum(), dtype=dtypes.float).reshape(-1, 1) - return _generate +"""Expose synthetic fixtures only to the unit-test subtree.""" + +from tests.fixtures.numerical import ( # noqa: F401 + mock_aniso_u, + mock_cell, + mock_cell_triclinic, + mock_F_obs, + mock_F_sigma, + mock_hkl_indices, + mock_scattering_factors, + mock_structure_factors, + mock_weights, + random_adp, + random_coordinates, + random_fractional_coordinates, + random_occupancies, + random_seed, +) diff --git a/tests/unit/structure_factor/conftest.py b/tests/unit/structure_factor/conftest.py index bcbf9be4..7ab42565 100644 --- a/tests/unit/structure_factor/conftest.py +++ b/tests/unit/structure_factor/conftest.py @@ -13,13 +13,11 @@ import torch import torchref -from torchref.config import device as device_cfg, dtypes - -from tests.conftest import _accelerator +from tests.fixtures.devices import _accelerator +from tests.fixtures.precision import cpu_double_precision from . import helpers as H - # --------------------------------------------------------------------------- # Device axis # --------------------------------------------------------------------------- @@ -86,7 +84,9 @@ def ds_device_dtype_kernels(): for name in H.ds_kernels_for(device, dtype): out.append( pytest.param( - device, dtype, name, + device, + dtype, + name, id=f"{device.type}-{str(dtype).replace('torch.float', 'f')}-{name}", marks=dev_param.marks, ) @@ -103,30 +103,9 @@ def ds_device_dtype_kernels(): # --------------------------------------------------------------------------- @pytest.fixture(scope="package", autouse=True) def _float64_cpu(): - """float64/complex128 on CPU for this package; restore afterwards. - - Required, not cosmetic: ``iso_structure_factor_torched`` casts ``hkl`` to the - *global* ``dtypes.float`` (``torchref/base/direct_summation/isotropic.py:121``), so - under the default float32 config a float64 leaf produces a dtype-mismatched matmul. - That is why the pre-existing tests wrapped every eager-SF call in a ``double_cpu`` - fixture. - - ``sigma_cutoff_ed`` is restored here too -- the three copies of ``double_cpu`` this - replaces did not, so a test that changed the cutoff leaked it into everything that - ran after it. - """ - f0, c0, d0 = dtypes.float, dtypes.complex, device_cfg.current - s0 = torchref.sigma_cutoff_ed.value - dtypes.float = torch.float64 - dtypes.complex = torch.complex128 - device_cfg.current = torch.device("cpu") - try: + """Scope the package's CPU double-precision reference configuration.""" + with cpu_double_precision(): yield - finally: - dtypes.float = f0 - dtypes.complex = c0 - device_cfg.current = d0 - torchref.sigma_cutoff_ed.value = s0 @pytest.fixture From d7eebba0d5f9d02f01c866ceb34353134f4d8db8 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 7 Sep 2026 16:51:52 +0200 Subject: [PATCH 161/250] test: consolidate weighting contracts under their API owners LossState owns hierarchical multiplication, aggregation, zero weights and cached reads. Default group weights remain separate. NLL checks move to base metrics; gradnorm smoke checks become exact RMS expectations. Validation: 36 passed on default MPS and CPU float64. Follow-up: all-zero aggregate returns torch float32 under configured float64; the retained zero-weight contract checks the value, not output dtype. --- docs/changelog.rst | 1 + .../test_loss_weighting_functional.py | 198 ------------------ tests/unit/base/test_loss.py | 36 ++++ tests/unit/refinement/test_loss_state.py | 46 +++- tests/unit/refinement/test_loss_weighting.py | 51 +---- tests/unit/utils/test_gradnorm.py | 91 +++----- 6 files changed, 100 insertions(+), 323 deletions(-) delete mode 100644 tests/functional/test_loss_weighting_functional.py create mode 100644 tests/unit/base/test_loss.py diff --git a/docs/changelog.rst b/docs/changelog.rst index b5b00059..4e17fba9 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Consolidated weighting tests by API ownership and strengthened Gaussian-likelihood, gradient-norm, and cached-loss assertions. - Organized test fixtures into focused modules and reused module-scoped loaded objects for read-only functional checks while retaining fresh objects for mutation and loading tests. - Rigid-body refinement stores its Euler angles pre-multiplied by the chain's radius of gyration, so a unit step in an angle and a unit step in a translation displace atoms comparably. In radians against Angstroms the rotation block of the Hessian carried 190-530x the curvature of the translation block on 1DAW and 3E98 -- the geometric ``Rg**2``, 411 and 442/516 -- putting ``cond(H)`` at 1e3-5e3, which is why six parameters needed ~250 L-BFGS iterations to place. Dividing the scale out in ``forward()`` brings the ratio to 0.4-1.3 and ``cond(H)`` to 3-18. Over ten structures the step then converges rather than exhausting its iteration budget, on about half the gradient evaluations, with R-free no worse anywhere. ``RigidXYZTensor.rotation_radians`` returns the physical angle, and setting ``angle_scale`` to ones restores the unscaled parametrization. Not a fix for the one or two negative Hessian eigenvalues at the finer cutoffs -- scaling a saddle leaves it a saddle -- and those counts are unchanged - The rigid-body step no longer co-refines the scaler in the same L-BFGS as the rigid parameters. The body target centres on ``alpha*|F_calc|`` and ``alpha`` absorbs a rescaling of ``F_calc`` exactly, so the scale had a flat direction there; ``SCALE_TARGETS`` already excludes every alpha-centred row from the scale fit for this reason, and 0.6.2 fixed the same thing in the main driver. ``refine_scaler`` (objective ``ls``) owns the scale, between cutoffs diff --git a/tests/functional/test_loss_weighting_functional.py b/tests/functional/test_loss_weighting_functional.py deleted file mode 100644 index f1bcb602..00000000 --- a/tests/functional/test_loss_weighting_functional.py +++ /dev/null @@ -1,198 +0,0 @@ -""" -Functional tests for loss weighting module. - -These tests exercise the loss weighting strategies with realistic data. -Updated to use the new component_weighting and LossState architecture. -""" -import pytest -import torch -import numpy as np -from unittest.mock import Mock - - -@pytest.mark.integration -class TestLossStateWeightingFunctional: - """Test LossState weighting functionality.""" - - def test_loss_state_add_and_get_weights(self): - """Test adding and getting weights from LossState.""" - from torchref.refinement.loss_state import LossState - - state = LossState() - state.set_weight('xray', 1.5) - state.set_weight('geometry', 0.7) - - assert state.get_weight('xray') == 1.5 - assert state.get_weight('geometry') == 0.7 - - def test_loss_state_hierarchical_weights(self): - """Test hierarchical weights in LossState.""" - from torchref.refinement.loss_state import LossState - - state = LossState() - state.set_weight('geometry', 2.0) - state.set_weight('geometry/bond', 3.0) - - # Effective weight should be product: 2.0 * 3.0 = 6.0 - effective = state.get_effective_weight('geometry/bond') - assert effective == 6.0 - - -@pytest.mark.integration -class TestWeightingMathOperations: - """Test mathematical operations with weights.""" - - def test_total_weighted_loss_from_state(self): - """Test computing total weighted loss from LossState via aggregate.""" - from torchref.refinement.loss_state import LossState - - state = LossState() - state.register_target('xray', lambda: torch.tensor(10.0)) - state.register_target('geometry', lambda: torch.tensor(5.0)) - state.register_target('adp', lambda: torch.tensor(2.0)) - - state.set_weight('xray', 1.0) - state.set_weight('geometry', 0.5) - state.set_weight('adp', 0.25) - - total = state.aggregate() - - # Expected: 10*1.0 + 5*0.5 + 2*0.25 = 10 + 2.5 + 0.5 = 13.0 - assert torch.isclose(total, torch.tensor(13.0)) - - -@pytest.mark.integration -class TestNLLXrayFunction: - """Test the NLL X-ray function used in weighting.""" - - def test_nll_xray_basic(self): - """Test basic NLL X-ray calculation.""" - from torchref.base.math_torch import nll_xray - - fobs = torch.tensor([100.0, 200.0, 300.0], dtype=torch.float32) - fcalc = torch.tensor([105.0, 195.0, 305.0], dtype=torch.float32) - sigma = torch.tensor([10.0, 15.0, 20.0], dtype=torch.float32) - - nll = nll_xray(fobs, fcalc, sigma) - - # nll returns per-reflection values - assert torch.all(torch.isfinite(nll)) - - def test_nll_decreases_with_better_fit(self): - """Test that NLL decreases as fit improves.""" - from torchref.base.math_torch import nll_xray - - fobs = torch.tensor([100.0], dtype=torch.float32) - sigma = torch.tensor([10.0], dtype=torch.float32) - - # Good fit - fcalc_good = torch.tensor([100.0], dtype=torch.float32) - nll_good = nll_xray(fobs, fcalc_good, sigma) - - # Bad fit - fcalc_bad = torch.tensor([150.0], dtype=torch.float32) - nll_bad = nll_xray(fobs, fcalc_bad, sigma) - - # Good fit should have lower NLL - assert nll_good < nll_bad - - -@pytest.mark.integration -class TestGradnormUtility: - """Test the gradnorm utility function.""" - - def test_gradnorm_basic(self): - """Test basic gradnorm calculation.""" - from torchref.utils.gradnorm import gradnorm - - # Create simple parameter - param = torch.tensor([1.0, 2.0, 3.0], requires_grad=True) - - # Create loss - loss = param.sum() - - # Compute gradient norm - norm = gradnorm(loss, [param]) - - assert torch.isfinite(norm) - assert norm > 0 - - def test_gradnorm_with_multiple_params(self): - """Test gradnorm with multiple parameters.""" - from torchref.utils.gradnorm import gradnorm - - param1 = torch.tensor([1.0, 2.0], requires_grad=True) - param2 = torch.tensor([3.0, 4.0], requires_grad=True) - - loss = param1.sum() + param2.sum() - - norm = gradnorm(loss, [param1, param2]) - - assert torch.isfinite(norm) - assert norm > 0 - - -@pytest.mark.integration -class TestWeightingEdgeCases: - """Test edge cases in weighting.""" - - def test_zero_weight(self): - """Test zero weight (disabling a loss term).""" - from torchref.refinement.loss_state import LossState - - state = LossState() - state.register_target('adp', lambda: torch.tensor(100.0)) - state.set_weight('adp', 0.0) - - # Zero weight should effectively disable ADP term - total = state.aggregate() - assert torch.isclose(total, torch.tensor(0.0)) - - -@pytest.mark.integration -class TestLossAggregatorFunctional: - """Test LossAggregator functionality.""" - - def test_aggregator_basic(self): - """Test basic aggregator functionality (LossState.aggregate).""" - from torchref.refinement.loss_state import LossState - - state = LossState() - state.register_target('xray', lambda: torch.tensor(2.0)) - state.register_target('bond', lambda: torch.tensor(1.0)) - state.set_weight('xray', 1.0) - state.set_weight('bond', 0.5) - - total = state.aggregate() - - # Expected: 2.0 * 1.0 + 1.0 * 0.5 = 2.5 - expected = torch.tensor(2.5) - assert torch.isclose(total, expected) - - def test_loss_state_caches_losses(self): - """Test that LossState caches computed losses.""" - from torchref.refinement.loss_state import LossState - - call_count = [0] - def counting_target(): - call_count[0] += 1 - return torch.tensor(2.0) - - state = LossState() - state.register_target('xray', counting_target) - state.set_weight('xray', 1.0) - - # register_target probes the target once to walk the autograd graph; - # reset the counter so we measure only aggregate() invocations. - call_count[0] = 0 - - # First aggregation computes the loss - total1 = state.aggregate() - assert call_count[0] == 1 - - # Get cached loss doesn't recompute - cached = state.get_loss('xray') - assert cached is not None - assert torch.isclose(cached, torch.tensor(2.0)) - - diff --git a/tests/unit/base/test_loss.py b/tests/unit/base/test_loss.py new file mode 100644 index 00000000..eff9ffcb --- /dev/null +++ b/tests/unit/base/test_loss.py @@ -0,0 +1,36 @@ +"""Pin the amplitude-metric Gaussian likelihood's value and reduction contract.""" + +import math + +import pytest +import torch + +from torchref.base.metrics.loss import nll_xray, nll_xray_mean, nll_xray_sum +from torchref.config import get_default_device, get_float_dtype + +pytestmark = pytest.mark.unit + + +def test_gaussian_nll_value_and_reduction() -> None: + """The NLL includes its normalization and sums over reflections.""" + obs = torch.tensor( + [10.0, 20.0, 30.0], dtype=get_float_dtype(), device=get_default_device() + ) + sigma = obs.new_tensor([1.0, 2.0, 4.0]) + calc = obs + sigma + expected = obs.new_tensor(1.5 + math.log(8.0) + 1.5 * math.log(2.0 * math.pi)) + + torch.testing.assert_close(nll_xray(obs, calc, sigma), expected) + torch.testing.assert_close(nll_xray_sum(obs, calc, sigma), expected) + torch.testing.assert_close(nll_xray_mean(obs, calc, sigma), expected / obs.numel()) + + +def test_gaussian_nll_penalizes_amplitude_error() -> None: + """A one-sigma residual adds one half per reflection to the perfect-fit NLL.""" + obs = torch.tensor( + [10.0, 20.0, 30.0], dtype=get_float_dtype(), device=get_default_device() + ) + sigma = torch.ones_like(obs) + good = nll_xray(obs, obs, sigma) + bad = nll_xray(obs, obs + sigma, sigma) + torch.testing.assert_close(bad - good, obs.new_tensor(1.5)) diff --git a/tests/unit/refinement/test_loss_state.py b/tests/unit/refinement/test_loss_state.py index c6126464..837b0940 100644 --- a/tests/unit/refinement/test_loss_state.py +++ b/tests/unit/refinement/test_loss_state.py @@ -44,7 +44,9 @@ def test_register_target(self): from torchref.refinement.loss_state import LossState state = LossState() - target_fn = lambda: torch.tensor(1.0) + + def target_fn(): + return torch.tensor(1.0) result = state.register_target("geometry/bond", target_fn) @@ -136,6 +138,7 @@ def test_set_weight(self): result = state.set_weight("geometry", 0.5) assert state.weights["geometry"] == 0.5 + assert state.get_weight("geometry") == 0.5 assert result is state # Method chaining @pytest.mark.unit @@ -184,12 +187,11 @@ def test_get_effective_weight_hierarchical(self): from torchref.refinement.loss_state import LossState state = LossState() - state.set_weight("geometry", 0.5) - state.set_weight("geometry/bond", 2.0) + state.set_weight("geometry", 2.0) + state.set_weight("geometry/bond", 3.0) - # geometry/bond -> geometry (0.5) * geometry/bond (2.0) = 1.0 effective = state.get_effective_weight("geometry/bond") - assert effective == 1.0 + assert effective == 6.0 @pytest.mark.unit def test_get_effective_weight_missing_intermediate(self): @@ -207,6 +209,22 @@ def test_get_effective_weight_missing_intermediate(self): class TestAggregation: """Tests for loss aggregation.""" + @pytest.mark.unit + def test_zero_weight(self): + """A zero weight contributes zero to the aggregate.""" + from torchref.config import get_default_device, get_float_dtype + from torchref.refinement.loss_state import LossState + + value = torch.tensor( + 100.0, dtype=get_float_dtype(), device=get_default_device() + ) + state = LossState() + state.register_target("adp", lambda: value) + state.set_weight("adp", 0.0) + total = state.aggregate() + assert total.ndim == 0 + assert total.item() == 0.0 + @pytest.mark.unit def test_aggregate_simple(self): """Test simple aggregation.""" @@ -260,15 +278,25 @@ def test_aggregate_default_weights(self): @pytest.mark.unit def test_aggregate_caches_losses(self): """Test that aggregate caches computed losses.""" + from torchref.config import get_default_device, get_float_dtype from torchref.refinement.loss_state import LossState state = LossState() - state.register_target("xray", lambda: torch.tensor(2.0)) + value = torch.tensor(2.0, dtype=get_float_dtype(), device=get_default_device()) + calls = 0 - state.aggregate(log_values=False) + def target(): + nonlocal calls + calls += 1 + return value - loss = state.get_loss("xray") - assert torch.isclose(loss, torch.tensor(2.0)) + state.register_target("xray", target) + # Registration probes the autograd graph; count only subsequent evaluations. + calls = 0 + state.aggregate(log_values=False) + assert calls == 1 + torch.testing.assert_close(state.get_loss("xray"), value) + assert calls == 1 class TestHistoryLogging: diff --git a/tests/unit/refinement/test_loss_weighting.py b/tests/unit/refinement/test_loss_weighting.py index 08a7fe0c..d314e624 100644 --- a/tests/unit/refinement/test_loss_weighting.py +++ b/tests/unit/refinement/test_loss_weighting.py @@ -1,55 +1,6 @@ -""" -Unit tests for LossState weight handling. - -Covers the retained ``LossState`` weight API (``set_weight`` / -``get_effective_weight`` / ``aggregate``). The standalone weighting -schemes were removed; refinement now aggregates at uniform weight by -default, with explicit per-target/group multipliers set via the -``LossState`` weight dict. -""" +"""Pin refinement's default group weights; LossState owns weight arithmetic.""" import pytest -import torch - - -class TestLossStateWeights: - """Tests for the LossState weight dict (hierarchical multipliers).""" - - @pytest.mark.unit - def test_hierarchical_weights_multiply(self): - """Test that hierarchical weights multiply in get_effective_weight.""" - from torchref.refinement.loss_state import LossState - - state = LossState() - # Set group weight - state.set_weight('geometry', 2.0) - # Set component weight - state.set_weight('geometry/bond', 3.0) - - # Effective weight should multiply: 2.0 * 3.0 = 6.0 - effective = state.get_effective_weight('geometry/bond') - assert effective == 6.0 - - -class TestTotalLossFromState: - """Tests for computing total loss from LossState.""" - - @pytest.mark.unit - def test_total_weighted_loss(self): - """Test computing total weighted loss from state via aggregate.""" - from torchref.refinement.loss_state import LossState - - state = LossState() - state.register_target('xray', lambda: torch.tensor(2.0)) - state.register_target('bond', lambda: torch.tensor(1.0)) - state.set_weight('xray', 1.0) - state.set_weight('bond', 0.5) - - total = state.aggregate(log_values=False) - - # Expected: 2.0 * 1.0 + 1.0 * 0.5 = 2.5 - expected = torch.tensor(2.5) - assert torch.isclose(total, expected) class TestDefaultGroupWeights: diff --git a/tests/unit/utils/test_gradnorm.py b/tests/unit/utils/test_gradnorm.py index d8e1a621..62378f84 100644 --- a/tests/unit/utils/test_gradnorm.py +++ b/tests/unit/utils/test_gradnorm.py @@ -1,75 +1,34 @@ -""" -Unit tests for torchref.utils.gradnorm +"""Pin the RMS gradient norm across one or several parameter tensors.""" -Tests gradient norm calculation utilities. -""" +import math import pytest import torch -import torch.nn as nn +from torchref.config import get_default_device, get_float_dtype +from torchref.utils.gradnorm import gradnorm -class TestGradNorm: - """Tests for gradient norm calculation.""" +pytestmark = pytest.mark.unit - @pytest.mark.unit - def test_gradnorm_basic(self): - """Test basic gradient norm calculation.""" - from torchref.utils.gradnorm import gradnorm - - # Simple linear model - model = nn.Linear(10, 1, bias=False) - x = torch.randn(5, 10) - y = torch.randn(5, 1) - - # Forward pass - pred = model(x) - loss = ((pred - y) ** 2).mean() - - # Calculate gradient norm - grad_norm = gradnorm(loss, model.parameters()) - - assert isinstance(grad_norm, torch.Tensor) - assert grad_norm.ndim == 0 # Scalar - assert grad_norm >= 0 # Non-negative - @pytest.mark.unit - def test_gradnorm_zero_gradient(self): - """Gradient norm should handle zero gradients.""" - from torchref.utils.gradnorm import gradnorm - - model = nn.Linear(10, 1, bias=False) - - # Create a loss that depends on the model but has zero gradient - x = torch.randn(3, 10) - pred = model(x) - loss = (pred * 0.0).sum() # Zero gradient - # DON'T call backward before gradnorm - it calls backward internally - - grad_norm = gradnorm(loss, model.parameters()) - - # Should be 0 (zero gradients) - assert torch.isclose(grad_norm, torch.tensor(0.0, dtype=grad_norm.dtype), atol=1e-10) +@pytest.mark.parametrize("split", [False, True], ids=["single", "multiple"]) +def test_gradnorm_rms(split: bool) -> None: + """The norm weights individual gradient elements, not parameter tensors.""" + values = torch.tensor( + [1.0, 2.0, 3.0], dtype=get_float_dtype(), device=get_default_device() + ) + chunks = (values[:1], values[1:]) if split else (values,) + params = [chunk.clone().requires_grad_() for chunk in chunks] + loss = sum((param.square().sum() for param in params)) + expected = values.new_tensor(math.sqrt(56.0 / 3.0)) + torch.testing.assert_close(gradnorm(loss, iter(params)), expected) - @pytest.mark.unit - def test_gradnorm_multiple_params(self): - """Test gradient norm with multiple parameter groups.""" - from torchref.utils.gradnorm import gradnorm - - # Model with multiple layers - model = nn.Sequential( - nn.Linear(10, 5), - nn.ReLU(), - nn.Linear(5, 1) - ) - - x = torch.randn(3, 10) - y = torch.randn(3, 1) - - pred = model(x) - loss = ((pred - y) ** 2).mean() - - grad_norm = gradnorm(loss, model.parameters()) - - assert isinstance(grad_norm, torch.Tensor) - assert grad_norm >= 0 + +def test_gradnorm_zero_gradient() -> None: + """A connected loss with zero derivative has zero RMS gradient.""" + param = torch.ones( + 3, dtype=get_float_dtype(), device=get_default_device(), requires_grad=True + ) + torch.testing.assert_close( + gradnorm((param * 0).sum(), [param]), param.new_zeros(()) + ) From 7912ce1390b811142f42263abeef7d83fc7191d2 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 7 Sep 2026 16:58:27 +0200 Subject: [PATCH 162/250] test: replace target arithmetic demonstrations with kernel contracts Remove local-only arithmetic and duplicate Target initialization; nn.Module inheritance remains in the comprehensive Target contract. Keep anisotropic DELU gradient routing, SIGD references and R-factor calls. Add deposited-coordinate bond, angle, chiral, plane and SIMU references plus exact LS weighting/mask values. Retain kernel/gradient boundary tests for torsion and DELU. Validation: 48 passed/2 CUDA skips including gradient guards; 34 passed on CPU float64 after accounting for the squared-distance regularizer. Eight zero-return fault injections detected. --- docs/changelog.rst | 1 + tests/README.md | 16 + tests/RUNNING_TESTS.md | 21 +- tests/unit/base/test_target_values.py | 127 +++++++ tests/unit/refinement/test_targets.py | 200 ---------- .../refinement/test_targets_comprehensive.py | 341 ++---------------- 6 files changed, 181 insertions(+), 525 deletions(-) create mode 100644 tests/unit/base/test_target_values.py delete mode 100644 tests/unit/refinement/test_targets.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 4e17fba9..a8fe18e8 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Replaced local-arithmetic target tests with configured-device production-kernel checks on deposited coordinates and explicit least-squares expectations. - Consolidated weighting tests by API ownership and strengthened Gaussian-likelihood, gradient-norm, and cached-loss assertions. - Organized test fixtures into focused modules and reused module-scoped loaded objects for read-only functional checks while retaining fresh objects for mutation and loading tests. - Rigid-body refinement stores its Euler angles pre-multiplied by the chain's radius of gyration, so a unit step in an angle and a unit step in a translation displace atoms comparably. In radians against Angstroms the rotation block of the Hessian carried 190-530x the curvature of the translation block on 1DAW and 3E98 -- the geometric ``Rg**2``, 411 and 442/516 -- putting ``cond(H)`` at 1e3-5e3, which is why six parameters needed ~250 L-BFGS iterations to place. Dividing the scale out in ``forward()`` brings the ratio to 0.4-1.3 and ``cond(H)`` to 3-18. Over ten structures the step then converges rather than exhausting its iteration budget, on about half the gradient evaluations, with R-free no worse anywhere. ``RigidXYZTensor.rotation_radians`` returns the physical angle, and setting ``angle_scale`` to ones restores the unscaled parametrization. Not a fix for the one or two negative Hessian eigenvalues at the finer cutoffs -- scaling a saddle leaves it a saddle -- and those counts are unchanged diff --git a/tests/README.md b/tests/README.md index ff1a3bc5..9e0ae0d6 100644 --- a/tests/README.md +++ b/tests/README.md @@ -39,6 +39,22 @@ tests/ ## Running Tests +### Coverage ownership + +| Contract | Owner | +|---|---| +| Loss weights, aggregation, cached loss reads | `unit/refinement/test_loss_state.py` | +| Refinement's default group weights | `unit/refinement/test_loss_weighting.py` | +| Gaussian amplitude-metric values and reductions | `unit/base/test_loss.py` | +| Restraint kernel values on deposited coordinates | `unit/base/test_target_values.py` | +| Gradient RMS norm | `unit/utils/test_gradnorm.py` | +| Numerical derivatives and backend parity | `unit/test_gradient_correctness.py`, `unit/structure_factor/` | + +A production call must participate in the assertion: computing a formula only in +the test does not check its implementation. Kernel values, target registration, +device transitions, and default configuration are separate contracts even when +they exercise the same class. Keep mutation tests on fresh objects. + ### Quick Local Run (on login node, for small tests only) ```bash diff --git a/tests/RUNNING_TESTS.md b/tests/RUNNING_TESTS.md index f2c8c120..109079d7 100644 --- a/tests/RUNNING_TESTS.md +++ b/tests/RUNNING_TESTS.md @@ -92,7 +92,7 @@ pytest tests/unit/refinement/ -v pytest tests/unit/refinement/test_loss_weighting.py -v # Target/loss functions -pytest tests/unit/refinement/test_targets.py -v +pytest tests/unit/base/test_target_values.py tests/unit/base/test_loss.py -v ``` ### Scaling @@ -176,17 +176,17 @@ pytest tests/unit/model/test_parameter_wrappers.py::TestMixedTensorOperations -v #### Refinement Classes ```bash -# Fixed weighting -pytest tests/unit/refinement/test_loss_weighting.py::TestFixedWeighting -v +# Weight handling +pytest tests/unit/refinement/test_loss_state.py::TestWeightManagement -v -# Resolution-dependent weighting -pytest tests/unit/refinement/test_loss_weighting.py::TestResolutionDependentWeighting -v +# Default group weights +pytest tests/unit/refinement/test_loss_weighting.py::TestDefaultGroupWeights -v # Gaussian NLL loss -pytest tests/unit/refinement/test_targets.py::TestGaussianNLL -v +pytest tests/unit/base/test_loss.py -v # Least squares target -pytest tests/unit/refinement/test_targets.py::TestLeastSquaresTarget -v +pytest tests/unit/base/test_target_values.py -k least_squares -v ``` #### Symmetry Classes @@ -361,13 +361,14 @@ pytest tests/unit --lf -v | `math_functions/test_math_numpy.py` | `TestCoordinateTransformations`, `TestScatteringVectors`, `TestRFactorCalculations`, `TestRotation` | | `model/test_model.py` | `TestModelInitialization`, `TestModelDeviceHandling` | | `model/test_parameter_wrappers.py` | `TestMixedTensorInitialization`, `TestMixedTensorOperations`, `TestMixedTensorDeviceHandling`, `TestOccupancyTensor`, `TestPositiveMixedTensor` | -| `refinement/test_loss_weighting.py` | `TestFixedWeighting`, `TestResolutionDependentWeighting`, `TestLossWeightingModule` | -| `refinement/test_targets.py` | `TestTargetBase`, `TestGaussianNLL`, `TestLeastSquaresTarget`, `TestRiceNLL`, `TestTargetDeviceHandling`, `TestNumericStability` | +| `refinement/test_loss_weighting.py` | `TestDefaultGroupWeights` | +| `base/test_target_values.py` | Deposited-coordinate restraint values and least-squares weighting | +| `base/test_loss.py` | Gaussian NLL values and reductions | | `scaling/test_scaler.py` | `TestScalerInitialization`, `TestScalerDeviceHandling`, `TestScalingCalculations`, `TestBFactorScaling`, `TestAnisotropicScaling` | | `symmetrie/test_symmetrie.py` | `TestSymmetryInitialization`, `TestSymmetryMatrices`, `TestSymmetryApplication`, `TestSymmetryDeviceHandling`, `TestSpaceGroupMapping` | | `io/test_data.py` | `TestReflectionDataInitialization`, `TestReflectionDataDeviceMovement`, `TestReflectionDataAttributes`, `TestReflectionDataProperties`, `TestMockReflectionData` | | `restraints/test_restraints.py` | `TestRestraintsInitialization`, `TestBondRestraintCalculations`, `TestAngleRestraintCalculations`, `TestTorsionRestraintCalculations`, `TestRestraintDeviceHandling`, `TestRestraintNumericStability` | -| `utils/test_gradnorm.py` | `TestGradNorm` | +| `utils/test_gradnorm.py` | RMS norms for single/multiple parameters and zero gradients | | `utils/test_utils.py` | `TestModuleReference`, `TestCIFReader` | ### Integration Tests (`tests/integration/`) diff --git a/tests/unit/base/test_target_values.py b/tests/unit/base/test_target_values.py new file mode 100644 index 00000000..bf27f0a6 --- /dev/null +++ b/tests/unit/base/test_target_values.py @@ -0,0 +1,127 @@ +"""Compare restraint kernels with host references on deposited Cartesian coordinates.""" + +import math + +import numpy as np +import pytest +import torch + +from torchref.base.targets._common import EPS +from torchref.base.targets.adp import adp_simu_math +from torchref.base.targets.angle import angle_math +from torchref.base.targets.bond import bond_math +from torchref.base.targets.chiral import chiral_math +from torchref.base.targets.planarity import planarity_math +from torchref.base.targets.xray_ls import ls_xray_loss_math +from torchref.config import get_default_device, get_float_dtype, get_int_dtype + +pytestmark = pytest.mark.unit + + +@pytest.fixture(scope="module") +def deposited_atoms(sample_cif_file): + """Return detached Cartesian coordinates (Å) and isotropic B-factors (Ų).""" + from torchref.model import Model + + model = Model(verbose=0) + model.load_cif(str(sample_cif_file)) + return model.xyz().detach().clone(), model.adp().detach().clone() + + +def _indices(rows, device): + return torch.tensor(rows, dtype=get_int_dtype(), device=device) + + +def _gaussian_sum(residual, sigma): + return np.sum( + 0.5 * (residual / sigma) ** 2 + np.log(sigma) + 0.5 * math.log(2 * math.pi) + ) + + +def test_bond_value(deposited_atoms) -> None: + """Bond lengths enter a summed Gaussian NLL in Å.""" + xyz, _ = deposited_atoms + host = xyz[:4].cpu().numpy().astype(np.float64) + idx = _indices([[0, 1], [2, 3]], xyz.device) + refs = xyz.new_tensor([1.4, 1.5]) + sigma = xyz.new_tensor([0.1, 0.2]) + # The kernel regularizes squared distance to keep coincident-atom gradients finite. + distance = np.sqrt(np.sum((host[[0, 2]] - host[[1, 3]]) ** 2, axis=1) + EPS) + expected = _gaussian_sum(distance - refs.cpu().numpy(), sigma.cpu().numpy()) + torch.testing.assert_close( + bond_math(xyz, idx, refs, sigma), xyz.new_tensor(expected) + ) + + +def test_angle_value(deposited_atoms) -> None: + """Angles and their restraint sigmas enter the NLL in radians.""" + import gemmi + + xyz, _ = deposited_atoms + positions = [gemmi.Position(*row) for row in xyz[:4].cpu().tolist()] + angles = np.array( + [gemmi.calculate_angle(*positions[:3]), gemmi.calculate_angle(*positions[1:4])] + ) + idx = _indices([[0, 1, 2], [1, 2, 3]], xyz.device) + refs = xyz.new_tensor([1.8, 2.0]) + sigma = xyz.new_tensor([0.1, 0.2]) + expected = _gaussian_sum(angles - refs.cpu().numpy(), sigma.cpu().numpy()) + torch.testing.assert_close( + angle_math(xyz, idx, refs, sigma), xyz.new_tensor(expected) + ) + + +def test_chiral_value(deposited_atoms) -> None: + """The signed scalar triple product, without a 1/6 factor, sets chirality.""" + xyz, _ = deposited_atoms + host = xyz[:4].cpu().numpy().astype(np.float64) + volume = np.linalg.det(host[1:] - host[0]) + idx = _indices([[0, 1, 2, 3]], xyz.device) + refs = xyz.new_tensor([2.0]) + sigma = xyz.new_tensor([0.5]) + expected = _gaussian_sum(volume - 2.0, 0.5) + torch.testing.assert_close( + chiral_math(xyz, idx, refs, sigma), xyz.new_tensor(expected) + ) + + +def test_planarity_value(deposited_atoms) -> None: + """The plane penalty sums signed-distance Gaussian NLLs over its atoms.""" + xyz, _ = deposited_atoms + host = xyz[:5].cpu().numpy().astype(np.float64) + centered = host - host.mean(axis=0) + _, _, vh = np.linalg.svd(centered, full_matrices=False) + distances = centered @ vh[-1] + idx = _indices([[0, 1, 2, 3, 4]], xyz.device) + sigma = xyz.new_full((1, 5), 0.2) + expected = _gaussian_sum(distances, 0.2) + torch.testing.assert_close( + planarity_math(xyz, [(idx, sigma)]), xyz.new_tensor(expected) + ) + + +def test_simu_value(deposited_atoms) -> None: + """SIMU penalizes differences of deposited isotropic B-factors in Ų.""" + _, b = deposited_atoms + host = b[:4].cpu().numpy().astype(np.float64) + idx = _indices([[0, 1], [2, 3]], b.device) + expected = _gaussian_sum(host[[0, 2]] - host[[1, 3]], 2.0) + torch.testing.assert_close( + adp_simu_math(b, idx, b.new_tensor(2.0)), b.new_tensor(expected) + ) + + +@pytest.mark.parametrize("weighting, expected", [("sigma", 6.5), ("unit", 20.0)]) +def test_least_squares_value_and_mask(weighting: str, expected: float) -> None: + """Least squares sums half squared amplitude errors using the selected weights.""" + obs = torch.tensor( + [10.0, 20.0, 30.0], dtype=get_float_dtype(), device=get_default_device() + ) + calc = -obs - obs.new_tensor([2.0, 6.0, 50.0]) + sigma = obs.new_tensor([1.0, 2.0, 5.0]) + mask = torch.tensor([True, True, False], device=obs.device) + loss = ls_xray_loss_math(obs, calc, sigma, mask, weighting=weighting) + torch.testing.assert_close(loss, obs.new_tensor(expected)) + torch.testing.assert_close( + ls_xray_loss_math(obs, obs, sigma, weighting=weighting), obs.new_zeros(()) + ) diff --git a/tests/unit/refinement/test_targets.py b/tests/unit/refinement/test_targets.py deleted file mode 100644 index fe2728e7..00000000 --- a/tests/unit/refinement/test_targets.py +++ /dev/null @@ -1,200 +0,0 @@ -""" -Unit tests for torchref.refinement.targets - -Tests target (loss) functions for crystallographic refinement. -Note: These are unit tests so we test the functions in isolation with mock data. -""" - -import pytest -import torch -import torch.nn as nn -import numpy as np - - -class TestTargetBase: - """Tests for base Target class.""" - - @pytest.mark.unit - def test_target_empty_initialization(self): - """Test Target can be initialized without arguments.""" - from torchref.refinement.targets import Target - - target = Target() - - assert target.verbose == 0 - - @pytest.mark.unit - def test_target_is_nn_module(self): - """Target should be a nn.Module.""" - from torchref.refinement.targets import Target - - target = Target() - - assert isinstance(target, nn.Module) - - -class TestGaussianNLL: - """Tests for Gaussian NLL calculation logic.""" - - @pytest.mark.unit - def test_gaussian_nll_identical_gives_small_loss(self, mock_F_obs, mock_F_sigma): - """When Fobs = Fcalc, NLL should be small (just the log sigma term).""" - from torchref.base.math_torch import nll_xray - - fobs = mock_F_obs(n_reflections=100) - sigma = mock_F_sigma(n_reflections=100) - fcalc = fobs.clone().to(torch.complex64) # |Fcalc| = Fobs - - # Calculate manually what Gaussian NLL should be - # NLL = 0.5*((fobs - |fcalc|)/sigma)^2 + log(sigma) + 0.5*log(2pi) - diff = fobs - torch.abs(fcalc) - expected_data_term = 0.5 * ((diff / sigma) ** 2) - - # Data term should be ~0 when fobs = |fcalc| - assert torch.allclose(expected_data_term, torch.zeros_like(expected_data_term), atol=1e-5) - - @pytest.mark.unit - def test_gaussian_nll_positive(self, mock_F_obs, mock_F_sigma): - """NLL should generally be positive or close to zero.""" - fobs = mock_F_obs(n_reflections=100) - sigma = mock_F_sigma(n_reflections=100) - fcalc = mock_F_obs(n_reflections=100, seed=123).to(torch.complex64) # Different - - # Simple Gaussian NLL - diff = fobs - torch.abs(fcalc) - eps = torch.median(sigma) * 0.1 - sigma_safe = torch.clamp(sigma, min=eps) - log_2pi = torch.log(torch.tensor(2.0 * np.pi)) - nll = 0.5 * (diff ** 2) / (sigma_safe ** 2) + torch.log(sigma_safe) + 0.5 * log_2pi - - # Mean NLL should be finite - assert torch.isfinite(nll.mean()) - - -class TestLeastSquaresTarget: - """Tests for Least Squares target calculation.""" - - @pytest.mark.unit - def test_least_squares_identical_zero(self, mock_F_obs): - """LS loss should be 0 when Fobs = Fcalc.""" - fobs = mock_F_obs(n_reflections=100) - fcalc = fobs.clone() - - # Simple LS: sum((fobs - fcalc)^2) - loss = torch.sum((fobs - fcalc) ** 2) - - assert torch.isclose(loss, torch.tensor(0.0, dtype=loss.dtype), atol=1e-10) - - @pytest.mark.unit - def test_least_squares_scaled(self, mock_F_obs): - """Test LS loss with scaled Fcalc.""" - fobs = mock_F_obs(n_reflections=100) - fcalc = fobs * 1.1 # 10% scaled - - loss = torch.mean((fobs - fcalc) ** 2) - - # Should be (0.1 * fobs)^2 on average - expected_loss = torch.mean((0.1 * fobs) ** 2) - assert torch.isclose(loss, expected_loss, rtol=1e-5) - - @pytest.mark.unit - def test_least_squares_weighted(self, mock_F_obs, mock_F_sigma): - """Test weighted LS with sigma weights.""" - fobs = mock_F_obs(n_reflections=100) - sigma = mock_F_sigma(n_reflections=100) - fcalc = mock_F_obs(n_reflections=100, seed=123) - - # Weighted LS: sum(w * (fobs - fcalc)^2) where w = 1/sigma^2 - weights = 1.0 / (sigma ** 2) - diff = fobs - fcalc - weighted_loss = torch.sum(weights * (diff ** 2)) - - assert torch.isfinite(weighted_loss) - assert weighted_loss >= 0 - - -class TestRiceNLL: - """Tests for Rice distribution NLL (used for acentric reflections).""" - - @pytest.mark.unit - def test_rice_nll_components(self, mock_F_obs, mock_F_sigma): - """Test components of Rice NLL calculation.""" - from torch.special import i0 - - fobs = mock_F_obs(n_reflections=50) - sigma = mock_F_sigma(n_reflections=50) - fcalc_amp = mock_F_obs(n_reflections=50, seed=123) - - # Rice NLL components - # NLL = (Fo^2 + Fc^2)/(2σ^2) - log(I0(Fo*Fc/σ^2)) - log(Fo/σ^2) - - # Check I0 calculation - x = fobs * fcalc_amp / (sigma ** 2) - bessel_i0 = i0(x) - - # I0 should be >= 1 for x >= 0 - assert torch.all(bessel_i0 >= 1.0) - - -class TestTargetDeviceHandling: - """Tests for proper device handling in targets.""" - - @pytest.mark.unit - def test_target_cpu_tensors(self, mock_F_obs, mock_F_sigma): - """Test calculations work on CPU.""" - fobs = mock_F_obs(n_reflections=100) - sigma = mock_F_sigma(n_reflections=100) - - # Simple calculation on CPU - loss = torch.mean((fobs / sigma) ** 2) - - assert loss.device.type == 'cpu' - assert torch.isfinite(loss) - - @pytest.mark.unit - @pytest.mark.gpu - def test_target_gpu_tensors(self, mock_F_obs, mock_F_sigma, gpu_device): - """Test calculations work on GPU.""" - fobs = mock_F_obs(n_reflections=100).to(gpu_device) - sigma = mock_F_sigma(n_reflections=100).to(gpu_device) - - loss = torch.mean((fobs / sigma) ** 2) - - assert loss.device.type == gpu_device.type - assert torch.isfinite(loss) - - -class TestNumericStability: - """Tests for numeric stability in target calculations.""" - - @pytest.mark.unit - def test_small_sigma_handling(self, mock_F_obs): - """Test handling of very small sigma values.""" - fobs = mock_F_obs(n_reflections=100) - sigma = torch.ones_like(fobs) * 1e-10 # Very small sigma - fcalc = mock_F_obs(n_reflections=100, seed=123) - - # Clamped sigma approach - eps = torch.median(sigma) * 0.1 - sigma_safe = torch.clamp(sigma, min=max(eps, 1e-6)) - - diff = fobs - fcalc - loss = torch.mean((diff / sigma_safe) ** 2) - - assert torch.isfinite(loss) - - @pytest.mark.unit - def test_zero_fcalc_handling(self, mock_F_obs, mock_F_sigma): - """Test handling of zero Fcalc values.""" - fobs = mock_F_obs(n_reflections=100) - sigma = mock_F_sigma(n_reflections=100) - fcalc = torch.zeros_like(fobs, dtype=torch.complex64) # All zero - - fcalc_amp = torch.abs(fcalc) # Will be zero - diff = fobs - fcalc_amp - - loss = torch.mean(diff ** 2) - - # Should just be mean of fobs^2 - expected = torch.mean(fobs ** 2) - assert torch.isclose(loss, expected, rtol=1e-5) diff --git a/tests/unit/refinement/test_targets_comprehensive.py b/tests/unit/refinement/test_targets_comprehensive.py index 21aa6c9e..0593a129 100644 --- a/tests/unit/refinement/test_targets_comprehensive.py +++ b/tests/unit/refinement/test_targets_comprehensive.py @@ -4,16 +4,16 @@ These tests focus on individual target classes with mock/minimal data to achieve higher coverage of the targets module. """ + +import numpy as np import pytest import torch -import numpy as np -from unittest.mock import MagicMock, PropertyMock - # ============================================================================= # Base Target Tests # ============================================================================= + @pytest.mark.unit class TestBaseTarget: """Test base Target class functionality.""" @@ -24,18 +24,19 @@ def test_target_initialization_empty(self): target = Target() assert target.verbose == 0 + assert isinstance(target, torch.nn.Module) def test_target_initialization_with_verbose(self): """Test initialization with verbose setting.""" from torchref.refinement.targets import Target - + target = Target(verbose=2) assert target.verbose == 2 def test_target_forward_not_implemented(self): """Test that forward raises NotImplementedError.""" from torchref.refinement.targets import Target - + target = Target() with pytest.raises(NotImplementedError): target.forward() @@ -45,6 +46,7 @@ def test_target_forward_not_implemented(self): # X-ray Target Tests # ============================================================================= + @pytest.mark.unit class TestXrayTargetBase: """Test XrayTarget base class.""" @@ -71,39 +73,9 @@ def test_gaussian_target_initialization(self): assert target._model is None assert target._data is None - def test_gaussian_nll_computation(self): - """Test Gaussian NLL computation with mock data.""" - from torchref.base.math_torch import nll_xray - - # Test the underlying function - fobs = torch.tensor([1.0, 2.0, 3.0, 4.0], dtype=torch.float32) - fcalc = torch.tensor([1.1, 1.9, 3.2, 3.8], dtype=torch.float32) - sigma = torch.tensor([0.1, 0.1, 0.1, 0.1], dtype=torch.float32) - - loss = nll_xray(fobs, fcalc, sigma).mean() - assert torch.isfinite(loss) # NLL can be negative depending on normalization -@pytest.mark.unit -class TestLeastSquaresXrayTarget: - """Test LeastSquaresXrayTarget.""" - - def test_least_squares_computation(self): - """Test least squares computation with mock data.""" - # Least squares: sum of (fobs - fcalc)^2 / sigma^2 - fobs = torch.tensor([1.0, 2.0, 3.0, 4.0], dtype=torch.float32) - fcalc = torch.tensor([1.1, 1.9, 3.2, 3.8], dtype=torch.float32) - sigma = torch.tensor([0.1, 0.1, 0.1, 0.1], dtype=torch.float32) - - diff = fobs - fcalc - weights = 1.0 / (sigma ** 2) - loss = 0.5 * torch.sum(weights * (diff ** 2)) - - assert torch.isfinite(loss) - assert loss >= 0 - - @pytest.mark.unit class TestRiceXrayTarget: """Test RiceXrayTarget.""" @@ -121,6 +93,7 @@ def test_rice_target_initialization(self): # Geometry Target Tests # ============================================================================= + @pytest.mark.unit class TestGeometryTargetBase: """Test GeometryTarget base class.""" @@ -144,30 +117,6 @@ def test_bond_target_initialization(self): target = BondTarget() assert target._model is None - def test_bond_deviation_calculation(self): - """Test bond deviation calculation with mock data.""" - # Create mock coordinates for a simple bond - xyz = torch.tensor([ - [0.0, 0.0, 0.0], - [1.5, 0.0, 0.0], # 1.5 Å bond - ], dtype=torch.float32) - - # Bond indices - i_atoms = torch.tensor([0]) - j_atoms = torch.tensor([1]) - - # Expected distance and sigma - d_expected = torch.tensor([1.54]) # Expected C-C bond - sigma = torch.tensor([0.02]) - - # Calculate actual distances - d_actual = torch.norm(xyz[i_atoms] - xyz[j_atoms], dim=1) - - # Calculate deviation - deviation = (d_actual - d_expected) / sigma - - assert torch.isfinite(deviation).all() - @pytest.mark.unit class TestAngleTarget: @@ -180,27 +129,6 @@ def test_angle_target_initialization(self): target = AngleTarget() assert target._model is None - def test_angle_calculation(self): - """Test angle calculation with mock data.""" - # Create mock coordinates for a 90-degree angle - xyz = torch.tensor([ - [1.0, 0.0, 0.0], # Atom 1 - [0.0, 0.0, 0.0], # Atom 2 (vertex) - [0.0, 1.0, 0.0], # Atom 3 - ], dtype=torch.float32) - - # Vectors - v1 = xyz[0] - xyz[1] - v2 = xyz[2] - xyz[1] - - # Calculate angle - cos_angle = torch.dot(v1, v2) / (torch.norm(v1) * torch.norm(v2)) - angle = torch.acos(cos_angle) - angle_deg = torch.rad2deg(angle) - - # Should be approximately 90 degrees - assert torch.isclose(angle_deg, torch.tensor(90.0), atol=0.1) - @pytest.mark.unit class TestTorsionTarget: @@ -213,34 +141,6 @@ def test_torsion_target_initialization(self): target = TorsionTarget() assert target._model is None - def test_torsion_angle_calculation(self): - """Test torsion angle calculation.""" - # Create mock coordinates for a torsion - # Atoms in a plane should give ~0 or ~180 degree torsion - xyz = torch.tensor([ - [0.0, 0.0, 0.0], - [1.0, 0.0, 0.0], - [2.0, 0.0, 0.0], - [3.0, 0.0, 0.0], - ], dtype=torch.float32) - - # Calculate torsion using standard formula - b1 = xyz[1] - xyz[0] - b2 = xyz[2] - xyz[1] - b3 = xyz[3] - xyz[2] - - # Normal vectors - n1 = torch.linalg.cross(b1, b2) - n2 = torch.linalg.cross(b2, b3) - - # Torsion angle - if torch.norm(n1) > 1e-6 and torch.norm(n2) > 1e-6: - cos_torsion = torch.dot(n1, n2) / (torch.norm(n1) * torch.norm(n2)) - # Clamp to valid range - cos_torsion = torch.clamp(cos_torsion, -1.0, 1.0) - torsion = torch.acos(cos_torsion) - assert torch.isfinite(torsion) - @pytest.mark.unit class TestPlanarityTarget: @@ -253,29 +153,6 @@ def test_planarity_target_initialization(self): target = PlanarityTarget() assert target._model is None - def test_planarity_calculation(self): - """Test planarity calculation for coplanar atoms.""" - # Atoms in the XY plane - xyz = torch.tensor([ - [0.0, 0.0, 0.0], - [1.0, 0.0, 0.0], - [1.0, 1.0, 0.0], - [0.0, 1.0, 0.0], - ], dtype=torch.float32) - - # Calculate centroid - centroid = xyz.mean(dim=0) - - # Center coordinates - centered = xyz - centroid - - # SVD to find plane - U, S, Vh = torch.linalg.svd(centered) - - # The smallest singular value indicates planarity - # For perfectly coplanar points, it should be ~0 - assert S[-1] < 0.1 - @pytest.mark.unit class TestChiralTarget: @@ -288,26 +165,6 @@ def test_chiral_target_initialization(self): target = ChiralTarget() assert target._model is None - def test_chiral_volume_calculation(self): - """Test chiral volume calculation.""" - # Create a tetrahedron - xyz = torch.tensor([ - [1.0, 0.0, -1.0/np.sqrt(2)], # Center - [0.0, 0.0, 1.0/np.sqrt(2)], # Atom 1 - [1.0, 1.0, 0.0], # Atom 2 - [1.0, -1.0, 0.0], # Atom 3 - ], dtype=torch.float32) - - # Vectors from center to other atoms - v1 = xyz[1] - xyz[0] - v2 = xyz[2] - xyz[0] - v3 = xyz[3] - xyz[0] - - # Chiral volume (scalar triple product) - chiral_vol = torch.dot(v1, torch.linalg.cross(v2, v3)) - - assert torch.isfinite(chiral_vol) - @pytest.mark.unit class TestNonBondedTarget: @@ -337,6 +194,7 @@ def test_total_geometry_target_initialization(self): # ADP Target Tests # ============================================================================= + @pytest.mark.unit class TestADPTargetBase: """Test ADPTarget base class.""" @@ -349,61 +207,10 @@ def test_adp_target_initialization(self): assert target._model is None -@pytest.mark.unit -class TestADPSimilarityTarget: - """Test ADPSimilarityTarget (SIMU restraint).""" - - def test_simu_calculation(self): - """Test SIMU calculation with mock B-factors.""" - # Create mock B-factors for nearby atoms - b_factors = torch.tensor([20.0, 21.0, 22.0, 50.0], dtype=torch.float32) - - # Pairs of similar atoms (indices) - i_atoms = torch.tensor([0, 1]) - j_atoms = torch.tensor([1, 2]) - - # Calculate difference - diff = b_factors[i_atoms] - b_factors[j_atoms] - - # SIMU restraint loss - sigma = 1.0 # B-factor sigma - simu_loss = (diff / sigma).pow(2).mean() - - assert torch.isfinite(simu_loss) - assert simu_loss >= 0 - - @pytest.mark.unit class TestRigidBondTarget: """Test RigidBondTarget (DELU restraint).""" - def test_delu_calculation(self): - """Test DELU calculation with mock U matrices.""" - # Create mock anisotropic U matrices (6 parameters each) - # U11, U22, U33, U12, U13, U23 - u1 = torch.tensor([0.05, 0.06, 0.04, 0.01, 0.005, -0.01], dtype=torch.float32) - u2 = torch.tensor([0.05, 0.06, 0.04, 0.01, 0.005, -0.01], dtype=torch.float32) - - # Bond vector (normalized) - bond_vec = torch.tensor([1.0, 0.0, 0.0], dtype=torch.float32) - - # Calculate U components along bond direction - # For Uij, the component along direction v is v^T U v - def u_along_direction(u_params, direction): - """Calculate U component along a direction.""" - U11, U22, U33, U12, U13, U23 = u_params - vx, vy, vz = direction - return (U11 * vx * vx + U22 * vy * vy + U33 * vz * vz + - 2 * U12 * vx * vy + 2 * U13 * vx * vz + 2 * U23 * vy * vz) - - u1_bond = u_along_direction(u1, bond_vec) - u2_bond = u_along_direction(u2, bond_vec) - - # DELU restraint: difference should be small - diff = u1_bond - u2_bond - - assert torch.isfinite(diff) - def test_aniso_path_runs_and_routes_grad_to_u(self, pdb_dir): """The anisotropic DELU path actually executes and feeds gradient to the U tensors. Regression for the dead ``hasattr(model, "u_aniso")`` gate, @@ -467,27 +274,11 @@ def test_matches_inverse_gamma_nll(self): beta = float(b.mean()) * (alpha - 1.0) mode = beta / (alpha + 1.0) - expected = ( - -sps.invgamma.logpdf(b.numpy(), alpha, scale=beta).sum() - + sps.invgamma.logpdf(mode, alpha, scale=beta) * len(b) - ) + expected = -sps.invgamma.logpdf( + b.numpy(), alpha, scale=beta + ).sum() + sps.invgamma.logpdf(mode, alpha, scale=beta) * len(b) assert float(adp_sigd_math(b, a, s0)) == pytest.approx(expected, rel=1e-10) - def test_alpha_sets_log_width(self): - """std(log B) = sqrt(trigamma(alpha)), the bridge the design rests on. - - This is what lets alpha play the role the log-normal's sigma played, and - is the basis for reporting ``implied_std_log_adp``. - """ - from scipy import stats as sps - from scipy.special import polygamma - - for alpha in (3.5, 7.4): - draws = sps.invgamma.rvs(alpha, scale=100.0, size=400000, random_state=1) - assert np.log(draws).std() == pytest.approx( - np.sqrt(polygamma(1, alpha)), rel=2e-2 - ) - def test_monotonically_increasing_in_spread(self): """The loss must never reward spreading the B distribution out. @@ -577,6 +368,7 @@ def test_gradient_pushes_toward_the_mode(self): # R-factor Tests # ============================================================================= + @pytest.mark.unit class TestRfactorCalculations: """Test R-factor calculation functions.""" @@ -584,15 +376,15 @@ class TestRfactorCalculations: def test_get_rfactors_basic(self): """Test basic R-factor calculation.""" from torchref.base.math_torch import get_rfactors - + fobs = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0], dtype=torch.float32) fcalc = torch.tensor([1.1, 2.1, 3.1, 4.1, 5.1], dtype=torch.float32) - + # Create rfree mask (1 reflection in test set) rfree_mask = torch.tensor([True, True, True, True, False], dtype=torch.bool) - + r_work, r_free = get_rfactors(fobs, fcalc, rfree_mask) - + # Both should be small since fcalc is close to fobs assert r_work < 0.2 # r_free only has one reflection @@ -600,33 +392,33 @@ def test_get_rfactors_basic(self): def test_get_rfactors_perfect_fit(self): """Test R-factor with perfect fit.""" from torchref.base.math_torch import get_rfactors - + fobs = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0], dtype=torch.float32) fcalc = fobs.clone() # Perfect fit - + rfree_mask = torch.tensor([True, True, True, True, False], dtype=torch.bool) - + r_work, r_free = get_rfactors(fobs, fcalc, rfree_mask) - + assert r_work < 0.001 # Should be ~0 def test_bin_wise_rfactors(self): """Test bin-wise R-factor calculation.""" from torchref.base.math_torch import bin_wise_rfactors - + n_refl = 100 n_bins = 5 - + fobs = torch.rand(n_refl) + 1.0 fcalc = fobs * (1 + 0.1 * torch.randn(n_refl)) # Note: rfree=True means work set (not free set) rfree_mask = torch.rand(n_refl) > 0.1 - + # Ensure all bins are represented bins = torch.arange(n_refl) % n_bins - + r_work_bins, r_free_bins = bin_wise_rfactors(fobs, fcalc, rfree_mask, bins) - + # Should have results for each bin assert len(r_work_bins) == n_bins assert len(r_free_bins) == n_bins @@ -636,88 +428,7 @@ def test_bin_wise_rfactors(self): # Loss Function Tests # ============================================================================= -@pytest.mark.unit -class TestLossFunctions: - """Test individual loss functions from math_torch.""" - - def test_nll_xray(self): - """Test NLL X-ray loss function.""" - from torchref.base.math_torch import nll_xray - - fobs = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float32) - fcalc = torch.tensor([1.1, 2.1, 3.1], dtype=torch.float32) - sigma = torch.tensor([0.1, 0.1, 0.1], dtype=torch.float32) - - loss = nll_xray(fobs, fcalc, sigma).mean() - - # NLL can be negative depending on normalization - assert torch.isfinite(loss) - - def test_least_squares_manual(self): - """Test least squares loss calculation.""" - # Manual least squares implementation - fobs = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float32) - fcalc = torch.tensor([1.1, 2.1, 3.1], dtype=torch.float32) - sigma = torch.tensor([0.1, 0.1, 0.1], dtype=torch.float32) - - diff = fobs - fcalc - weights = 1.0 / (sigma ** 2) - loss = 0.5 * torch.sum(weights * (diff ** 2)) / len(fobs) - - assert torch.isfinite(loss) - assert loss >= 0 - - def test_nll_xray_with_mask(self): - """Test NLL X-ray with masking.""" - from torchref.base.math_torch import nll_xray - - fobs = torch.tensor([1.0, 2.0, 3.0, float('nan')], dtype=torch.float32) - fcalc = torch.tensor([1.1, 2.1, 3.1, 0.0], dtype=torch.float32) - sigma = torch.tensor([0.1, 0.1, 0.1, 0.1], dtype=torch.float32) - - # Only use finite values - valid = torch.isfinite(fobs) - loss = nll_xray(fobs[valid], fcalc[valid], sigma[valid]).mean() - - assert torch.isfinite(loss) - # ============================================================================= # Helper Function Tests # ============================================================================= - -@pytest.mark.unit -class TestTargetHelpers: - """Test helper functions used in targets.""" - - def test_distance_calculation(self): - """Test distance calculation between atom pairs.""" - xyz = torch.tensor([ - [0.0, 0.0, 0.0], - [3.0, 4.0, 0.0], # Distance = 5.0 - ], dtype=torch.float32) - - distance = torch.norm(xyz[1] - xyz[0]) - - assert torch.isclose(distance, torch.tensor(5.0)) - - def test_angle_from_vectors(self): - """Test angle calculation from vectors.""" - v1 = torch.tensor([1.0, 0.0, 0.0], dtype=torch.float32) - v2 = torch.tensor([0.0, 1.0, 0.0], dtype=torch.float32) - - cos_angle = torch.dot(v1, v2) / (torch.norm(v1) * torch.norm(v2)) - angle = torch.acos(cos_angle) - angle_deg = torch.rad2deg(angle) - - assert torch.isclose(angle_deg, torch.tensor(90.0)) - - def test_cross_product(self): - """Test cross product for normal vectors.""" - v1 = torch.tensor([1.0, 0.0, 0.0], dtype=torch.float32) - v2 = torch.tensor([0.0, 1.0, 0.0], dtype=torch.float32) - - normal = torch.linalg.cross(v1, v2) - - # Should be [0, 0, 1] - assert torch.allclose(normal, torch.tensor([0.0, 0.0, 1.0])) From b3f58713ddefe7b916ec1caa3dfa140f67550801 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 7 Sep 2026 17:03:59 +0200 Subject: [PATCH 163/250] test: consolidate reader contracts and exercise ModelFT caching CIF and MTZ integration tests own field shapes, crystal metadata, bin means and pair consistency; required assertions no longer depend on optional attributes. Existing CIF-to-PDB writing and device movement stay separate. Replace empty ModelFT state/forward/aniso checks with actual forward/cache checks, leaving restoration and anisotropic coverage in model unit tests. Remove unconsumed shared Model/ReflectionData fixtures. Validation: 13 default-MPS cases passed plus the corrected cache case; CPU float64 run 15 passed, 1 slow skip. Cache recomputation compares real magnitudes to allow backend reduction order. --- docs/changelog.rst | 1 + tests/README.md | 3 + tests/fixtures/README.md | 9 +- tests/fixtures/functional.py | 19 +- tests/functional/conftest.py | 6 +- tests/functional/test_io_functional.py | 213 ---------------- tests/functional/test_model_ft_functional.py | 197 +++------------ tests/functional/test_targets_functional.py | 62 ++--- tests/integration/test_io_cif.py | 118 +++------ tests/integration/test_io_reflections.py | 248 ++++++------------- 10 files changed, 180 insertions(+), 696 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index a8fe18e8..011122e3 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Consolidated CIF/MTZ loading contracts, checked configured tensor placement, and replaced ModelFT smoke checks with exercised forward-cache behavior. - Replaced local-arithmetic target tests with configured-device production-kernel checks on deposited coordinates and explicit least-squares expectations. - Consolidated weighting tests by API ownership and strengthened Gaussian-likelihood, gradient-norm, and cached-loss assertions. - Organized test fixtures into focused modules and reused module-scoped loaded objects for read-only functional checks while retaining fresh objects for mutation and loading tests. diff --git a/tests/README.md b/tests/README.md index 9e0ae0d6..9298352f 100644 --- a/tests/README.md +++ b/tests/README.md @@ -48,6 +48,9 @@ tests/ | Gaussian amplitude-metric values and reductions | `unit/base/test_loss.py` | | Restraint kernel values on deposited coordinates | `unit/base/test_target_values.py` | | Gradient RMS norm | `unit/utils/test_gradnorm.py` | +| CIF atomic fields and crystal metadata | `integration/test_io_cif.py` | +| MTZ fields, resolution bins and model/data crystal agreement | `integration/test_io_reflections.py` | +| ModelFT forward cache and grid integration | `functional/test_model_ft_functional.py` | | Numerical derivatives and backend parity | `unit/test_gradient_correctness.py`, `unit/structure_factor/` | A production call must participate in the assertion: computing a formula only in diff --git a/tests/fixtures/README.md b/tests/fixtures/README.md index da39ba78..7361c509 100644 --- a/tests/fixtures/README.md +++ b/tests/fixtures/README.md @@ -12,10 +12,10 @@ Keep a fixture in its test module when only that module needs it. | `precision.py` | Comparison tolerances and CPU-double reference context | All tests; reference fixture restores state after each test | | `objects.py` | Mutable models, data, scalers and restraints | All tests; fresh per test except explicitly shared bundles | | `numerical.py` | Synthetic tensors and factories | Imported only by `unit/conftest.py`; function | -| `functional.py` | `shared_model`, `shared_model_ft`, `shared_reflection_data` | Imported only by `functional/conftest.py`; module | +| `functional.py` | Read-only `shared_model_ft` | Imported only by `functional/conftest.py`; module | -Existing fixture names remain available without imports in tests. Import reusable -helpers from their defining module, never from the root `conftest.py`. Subtree +Fixtures are available without imports in tests. Import reusable helpers from +their defining module, never from the root `conftest.py`. Subtree conftests import fixture functions explicitly; register shared plugins only at the root so pytest also works when invoked from a subdirectory. @@ -25,7 +25,8 @@ parameters, tables, masks or grids, backpropagate through them, or use them in tests that switch global configuration. A target or scaler can mutate a model it borrows, so a shared model must not be passed to such an operation. -Tests that verify loading must invoke the loader themselves. Tests of mutation, +Tests that verify loading must execute a fresh loader, directly or through a +function-scoped fixture. Tests of mutation, device movement, or empty caches use fresh objects. `loaded_model`, `loaded_model_ft`, `loaded_reflection_data`, and their composed fixtures in `objects.py` provide fresh mutable objects per test. The explicitly shared session bundles in that diff --git a/tests/fixtures/functional.py b/tests/fixtures/functional.py index 18168242..5ee204a4 100644 --- a/tests/fixtures/functional.py +++ b/tests/fixtures/functional.py @@ -14,16 +14,7 @@ import pytest if TYPE_CHECKING: - from torchref.io import ReflectionData - from torchref.model import Model, ModelFT - - -@pytest.fixture(scope="module") -def shared_model(sample_cif_file: Path) -> Model: - """Load the sample CIF once per module for read-only atomic-model checks.""" - from torchref.model import Model - - return Model(verbose=0).load_cif(str(sample_cif_file)) + from torchref.model import ModelFT @pytest.fixture(scope="module") @@ -32,11 +23,3 @@ def shared_model_ft(sample_cif_file: Path) -> ModelFT: from torchref.model import ModelFT return ModelFT(max_res=2.0, verbose=0).load_cif(str(sample_cif_file)) - - -@pytest.fixture(scope="module") -def shared_reflection_data(sample_mtz_file: Path) -> ReflectionData: - """Load the sample MTZ once per module for read-only reflection checks.""" - from torchref.io import ReflectionData - - return ReflectionData().load_mtz(str(sample_mtz_file)) diff --git a/tests/functional/conftest.py b/tests/functional/conftest.py index e7a5b0ed..ba6caf3b 100644 --- a/tests/functional/conftest.py +++ b/tests/functional/conftest.py @@ -1,7 +1,3 @@ """Expose module-shared read-only fixtures to functional tests.""" -from tests.fixtures.functional import ( # noqa: F401 - shared_model, - shared_model_ft, - shared_reflection_data, -) +from tests.fixtures.functional import shared_model_ft # noqa: F401 diff --git a/tests/functional/test_io_functional.py b/tests/functional/test_io_functional.py index 8779ea39..49c82a99 100644 --- a/tests/functional/test_io_functional.py +++ b/tests/functional/test_io_functional.py @@ -5,7 +5,6 @@ """ import pytest -import torch class TestCIFReadingFunctional: @@ -31,32 +30,6 @@ def test_load_multiple_cif_files(self, cif_dir): assert model.cell is not None assert len(model.cell) == 6 - @pytest.mark.integration - def test_cif_atom_properties(self, shared_model): - """Test that atom properties are correctly loaded from CIF.""" - model = shared_model - - pdb = model.pdb - - # Check required columns exist - required_cols = ["x", "y", "z", "element", "resname", "chainid", "resseq"] - for col in required_cols: - assert col in pdb.columns or col.upper() in pdb.columns, ( - f"Missing column: {col}" - ) - - @pytest.mark.integration - def test_cif_element_types(self, shared_model): - """Test that element types are properly assigned.""" - model = shared_model - - elements = model.pdb["element"].unique() - - # Should have common protein elements - common_elements = ["C", "N", "O", "S"] - found_any = any(elem in elements for elem in common_elements) - assert found_any, "No common elements found" - class TestMTZReadingFunctional: """Functional tests for MTZ file reading.""" @@ -80,41 +53,6 @@ def test_load_multiple_mtz_files(self, mtz_dir): # Should have cell parameters assert data.cell is not None - @pytest.mark.integration - def test_mtz_data_properties(self, shared_reflection_data): - """Test that MTZ data properties are correctly loaded.""" - data = shared_reflection_data - - # Check HKL indices are integers or can be converted - hkl = data.hkl - assert hkl.shape[1] == 3, "HKL should have 3 columns" - - # Check F values are loaded - assert data.F is not None - assert data.F.shape[0] == hkl.shape[0] - - # Check sigma values - if hasattr(data, "F_sigma") and data.F_sigma is not None: - assert data.F_sigma.shape[0] == hkl.shape[0] - - @pytest.mark.integration - def test_mtz_resolution_range(self, shared_reflection_data): - """Test that resolution range is computed correctly.""" - data = shared_reflection_data - - # Check if resolution data is available - if hasattr(data, "d") and data.d is not None: - d_min = data.d.min().item() - d_max = data.d.max().item() - - # Resolution should be positive - assert d_min > 0 - assert d_max > d_min - - # Typical protein data: 0.8 - 500 Å - assert d_min > 0.5 - assert d_max < 1000 - class TestSFCIFReadingFunctional: """Functional tests for structure factor CIF reading.""" @@ -139,154 +77,3 @@ def test_load_sf_cif(self, cif_sf_dir): except Exception as e: # Some files may not be valid SF-CIF format pass - - -class TestDataConsistencyFunctional: - """Test consistency between model and data files.""" - - @pytest.mark.integration - def test_cell_parameters_match(self, sample_structure_pair): - """Test that cell parameters match between model and reflections.""" - from torchref.io import ReflectionData - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_structure_pair["model"])) - - data = ReflectionData() - data.load_mtz(str(sample_structure_pair["reflections"])) - - model_cell = model.cell - data_cell = data.cell - - if model_cell is not None and data_cell is not None: - # Convert to tensors if needed - if not isinstance(model_cell, torch.Tensor): - model_cell = torch.tensor(model_cell) - if not isinstance(data_cell, torch.Tensor): - data_cell = torch.tensor(data_cell) - - # Cell parameters should be similar (1% tolerance) - assert torch.allclose( - model_cell.float(), data_cell.float(), rtol=0.01, atol=0.1 - ) - - @pytest.mark.integration - def test_spacegroup_consistency(self, sample_structure_pair): - """Test that spacegroup is consistent.""" - from torchref.io import ReflectionData - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_structure_pair["model"])) - - data = ReflectionData() - data.load_mtz(str(sample_structure_pair["reflections"])) - - # Both should have spacegroup defined - assert model.spacegroup is not None - - -class TestDataBinningFunctional: - """Test data binning operations.""" - - @pytest.mark.integration - def test_get_bins(self, sample_mtz_file): - """Test resolution binning of reflection data.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # Get bins - bins, n_bins = data.get_bins(n_bins=10) - - assert bins is not None - assert bins.shape[0] == data.hkl.shape[0] - assert bins.min() >= 0 - assert bins.max() < n_bins - - @pytest.mark.integration - def test_mean_res_per_bin(self, sample_mtz_file): - """Test mean resolution per bin calculation.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # Get bins first - bins, n_bins = data.get_bins(n_bins=10) - - # Get mean resolution per bin - if hasattr(data, "mean_res_per_bin"): - mean_res = data.mean_res_per_bin() - - assert mean_res is not None - assert len(mean_res) == n_bins - - # Mean resolution should decrease with bin index (low res to high res) - # or increase (high res to low res) - depends on implementation - assert torch.all(torch.isfinite(mean_res)) - - -class TestFrenchWilsonFunctional: - """Test French-Wilson conversion with real data.""" - - @pytest.mark.integration - def test_french_wilson_applied(self, shared_reflection_data): - """Test that French-Wilson conversion is applied.""" - data = shared_reflection_data - - # After French-Wilson, F values should be non-negative - valid_F = data.F[~torch.isnan(data.F)] - - if len(valid_F) > 0: - # All valid F values should be >= 0 - assert torch.all(valid_F >= 0) - - -class TestRfreeHandlingFunctional: - """Test R-free flag handling.""" - - @pytest.mark.integration - def test_rfree_flags_loaded(self, shared_reflection_data): - """Test that R-free flags are loaded or generated.""" - data = shared_reflection_data - - # Should have rfree attribute - if hasattr(data, "rfree") and data.rfree is not None: - assert data.rfree.shape[0] == data.hkl.shape[0] - - # Should be boolean or can be converted to boolean - assert data.rfree.dtype == torch.bool or torch.all( - (data.rfree == 0) | (data.rfree == 1) - ) - - @pytest.mark.integration - def test_rfree_fraction(self, shared_reflection_data): - """Test R-free set fraction is reasonable.""" - data = shared_reflection_data - - if hasattr(data, "rfree") and data.rfree is not None: - # Work set mask (True for work, False for test) - work_fraction = data.rfree.float().mean().item() - - # Typically 90-95% work set, 5-10% test set - # So work_fraction should be 0.9-0.95 typically - assert 0.7 < work_fraction <= 1.0 - - -class TestMaskHandlingFunctional: - """Test reflection mask handling.""" - - @pytest.mark.integration - def test_masks_method(self, shared_reflection_data): - """Test masks() method returns valid mask.""" - data = shared_reflection_data - - if hasattr(data, "masks"): - mask = data.masks() - - assert mask is not None - assert mask.shape[0] == data.hkl.shape[0] - assert mask.dtype == torch.bool diff --git a/tests/functional/test_model_ft_functional.py b/tests/functional/test_model_ft_functional.py index a6abd808..2c6b9b4c 100644 --- a/tests/functional/test_model_ft_functional.py +++ b/tests/functional/test_model_ft_functional.py @@ -28,19 +28,6 @@ def test_modelft_with_custom_resolution(self): model = ModelFT(max_res=1.5) assert model.max_res == 1.5 - def test_modelft_load_cif(self, sample_cif_file): - """Test loading a CIF file into ModelFT.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - - # Verify basic properties - assert model.xyz() is not None - assert model.xyz().shape[0] > 0 - assert model.cell is not None - assert len(model.cell) == 6 - def test_modelft_has_gridsize(self, sample_cif_file): """The grid resolves from the loaded cell and space group on first read.""" from torchref.model.model_ft import ModelFT @@ -52,29 +39,10 @@ def test_modelft_has_gridsize(self, sample_cif_file): assert model.gridsize is not None assert len(model.gridsize) == 3 assert all(g > 0 for g in model.gridsize) - - -@pytest.mark.integration -class TestModelFTParametrization: - """Test ModelFT parametrization with real structures.""" - - def test_parametrization_built(self, shared_model_ft): - """Test that parametrization is built after loading.""" - - model = shared_model_ft - - # Parametrization should be set - assert model.parametrization is not None - - def test_scattering_factors_available(self, shared_model_ft): - """Test that scattering factors can be computed.""" - - model = shared_model_ft - - # Should be able to access atom properties - xyz = model.xyz() - assert xyz is not None - assert xyz.dtype == torch.float32 or xyz.dtype == torch.float64 + assert model.xyz().shape[0] > 0 + assert model.parametrization + assert model.adp().shape == (len(model.xyz()),) + assert torch.all(model.adp() >= 0) @pytest.mark.integration @@ -104,14 +72,15 @@ def test_get_real_space_grid(self, loaded_model_ft): """Test getting real space grid.""" from torchref.base.math_torch import get_real_grid - # The grid helper moves its Cell in place when targeting CPU. model = loaded_model_ft assert model.gridsize is not None - grid = get_real_grid(model.cell, max_res=2.0, device="cpu") + grid = get_real_grid(model.cell, max_res=2.0, device=model.device) assert grid is not None assert len(grid.shape) == 4 # Should be 4D (nx, ny, nz, 3) + assert grid.device == model.xyz().device + assert grid.dtype == model.xyz().dtype @pytest.mark.integration @@ -135,60 +104,6 @@ def test_map_symmetry_available(self, shared_model_ft): assert operator.map_shape == gridsize -@pytest.mark.integration -class TestModelFTStateDictFunctional: - """Test ModelFT state dict operations with real data.""" - - def test_save_and_load_state_dict(self, loaded_model_ft, tmp_path): - """Test saving and loading state dict.""" - from torchref.model.model_ft import ModelFT - - model = loaded_model_ft - - original_xyz = model.xyz().clone() - - # Save state dict - state_dict = model.state_dict() - - # Create new model and load state - model2 = ModelFT(max_res=2.0, verbose=0) - - # We need to ensure proper initialization - # For now just verify state_dict works - assert state_dict is not None - assert len(state_dict) > 0 - - -@pytest.mark.integration -class TestModelFTForwardPass: - """Test ModelFT forward pass (structure factor calculation).""" - - def test_forward_method_exists(self, shared_model_ft): - """Test that forward method is available.""" - - model = shared_model_ft - - # Check forward method exists - assert hasattr(model, "forward") - - def test_build_map_method(self, sample_cif_file): - """Test build_map method if available.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=3.0, verbose=0) # Lower res for faster test - model.load_cif(str(sample_cif_file)) - - # Check build_map method - if hasattr(model, "build_map"): - # Try to build map - try: - model.build_map() - assert model.map is not None - except Exception as e: - # May fail if missing dependencies - pytest.skip(f"build_map not available: {e}") - - @pytest.mark.integration class TestModelFTMultipleStructures: """Test ModelFT with multiple structures.""" @@ -215,50 +130,10 @@ def test_modelft_multiple_structures(self, all_structure_pairs): assert tested >= 1, "At least one structure should load" -@pytest.mark.integration -class TestModelFTCaching: - """Test ModelFT caching mechanism.""" - - def test_cache_initialization(self, loaded_model_ft): - """Test that CachedForwardMixin cache starts empty.""" - - model = loaded_model_ft - - # Mixin cache should start empty (lazily initialized) - assert getattr(model, "_fwd_cached_output", None) is None - - def test_cache_usage(self, shared_model_ft): - """Test that cache can be used for computations.""" - - model = shared_model_ft - - # Access xyz twice - should use caching - xyz1 = model.xyz() - xyz2 = model.xyz() - - # Should return same tensor - assert torch.allclose(xyz1, xyz2) - - @pytest.mark.integration class TestModelFTCoordinateOperations: """Test ModelFT coordinate operations.""" - def test_cartesian_to_fractional(self, shared_model_ft): - """Test coordinate conversion.""" - from torchref.base.math_torch import cartesian_to_fractional_torch - - model = shared_model_ft - - xyz = model.xyz() - cell = model.cell - - # Convert to fractional - frac = cartesian_to_fractional_torch(xyz, cell.data) - - # Fractional coords should be bounded (mostly between 0 and 1) - assert frac.shape == xyz.shape - def test_fractional_to_cartesian(self, shared_model_ft): """Test fractional to cartesian conversion.""" from torchref.base.math_torch import ( @@ -273,6 +148,7 @@ def test_fractional_to_cartesian(self, shared_model_ft): # Round trip conversion frac = cartesian_to_fractional_torch(xyz, cell.data) + assert frac.shape == xyz.shape xyz_back = fractional_to_cartesian_torch(frac, cell.data) # Should get back original coordinates (float32 roundtrip) @@ -280,28 +156,35 @@ def test_fractional_to_cartesian(self, shared_model_ft): @pytest.mark.integration -class TestModelFTAnisoHandling: - """Test ModelFT handling of anisotropic parameters.""" - - def test_access_aniso_atoms(self, shared_model_ft): - """Test accessing anisotropic atom information.""" - - model = shared_model_ft - - # Check if aniso is available - if hasattr(model, "get_aniso") or hasattr(model, "aniso"): - # Structure has aniso - pass - - def test_isotropic_b_factors(self, shared_model_ft): - """Test accessing isotropic B-factors.""" - - model = shared_model_ft - - # Get B-factors (now accessed via adp()) - b_factors = model.adp() - - assert b_factors is not None - assert b_factors.shape[0] == model.xyz().shape[0] - # B-factors should be positive - assert torch.all(b_factors > 0) or torch.all(b_factors >= 0) +def test_forward_cache_contract( + loaded_model_ft, loaded_reflection_data, monkeypatch +) -> None: + """A model computes complex structure factors and caches only until invalidation.""" + from unittest.mock import Mock + + from torchref.config import caching, get_complex_dtype + + model = loaded_model_ft + hkl = loaded_reflection_data.hkl[:32] + monkeypatch.setattr(caching, "value", True) + forward = Mock(wraps=model.forward) + monkeypatch.setattr(model, "forward", forward) + assert getattr(model, "_fwd_cached_output", None) is None + + first = model(hkl) + assert first.shape == (len(hkl),) + assert first.dtype == get_complex_dtype() + assert first.device == hkl.device + assert torch.isfinite(first).all() + assert first.abs().sum() > 0 + assert model(hkl) is first + assert forward.call_count == 1 + + refreshed = model(hkl, recalc=True) + assert forward.call_count == 2 + assert refreshed is not first + # Accelerator reductions need not repeat bit-for-bit after recomputation. + relative_error = torch.linalg.vector_norm( + (refreshed - first).abs() + ) / torch.linalg.vector_norm(first.abs()) + assert relative_error < 256 * torch.finfo(first.real.dtype).eps diff --git a/tests/functional/test_targets_functional.py b/tests/functional/test_targets_functional.py index 5eb505e1..c3f2f70d 100644 --- a/tests/functional/test_targets_functional.py +++ b/tests/functional/test_targets_functional.py @@ -4,9 +4,9 @@ Tests target functions with real model and data objects. """ +import numpy as np import pytest import torch -import numpy as np class TestXrayTargetsFunctional: @@ -15,9 +15,9 @@ class TestXrayTargetsFunctional: @pytest.mark.integration def test_gaussian_nll_with_real_data(self, sample_structure_pair): """Test Gaussian NLL calculation with real reflection data.""" - from torchref.model.model import Model - from torchref.io import ReflectionData from torchref.base.math_torch import nll_xray + from torchref.io import ReflectionData + from torchref.model.model import Model model = Model() model.load_cif(str(sample_structure_pair["model"])) @@ -43,8 +43,8 @@ def test_gaussian_nll_with_real_data(self, sample_structure_pair): @pytest.mark.integration def test_least_squares_with_real_data(self, sample_structure_pair): """Test least squares calculation with real data.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model model = Model() model.load_cif(str(sample_structure_pair["model"])) @@ -74,9 +74,9 @@ class TestRfactorCalculationsFunctional: @pytest.mark.integration def test_rfactor_with_real_data(self, sample_structure_pair): """Test R-factor calculation with real reflection data.""" - from torchref.model.model import Model - from torchref.io import ReflectionData from torchref.base.math_torch import get_rfactors + from torchref.io import ReflectionData + from torchref.model.model import Model model = Model() model.load_cif(str(sample_structure_pair["model"])) @@ -110,9 +110,9 @@ def test_rfactor_with_real_data(self, sample_structure_pair): @pytest.mark.integration def test_bin_wise_rfactors(self, sample_structure_pair): """Test bin-wise R-factor calculation.""" - from torchref.model.model import Model - from torchref.io import ReflectionData from torchref.base.math_torch import bin_wise_rfactors + from torchref.io import ReflectionData + from torchref.model.model import Model model = Model() model.load_cif(str(sample_structure_pair["model"])) @@ -234,27 +234,6 @@ def test_angle_target_with_real_structure(self, sample_cif_file, external_monome assert torch.isfinite(loss) -class TestStructureFactorCalculationFunctional: - """Functional tests for structure factor calculation.""" - - @pytest.mark.integration - def test_fcalc_shape_matches_data(self, sample_structure_pair): - """Test that calculated structure factors have correct shape.""" - from torchref.model.model import Model - from torchref.io import ReflectionData - - model = Model() - model.load_cif(str(sample_structure_pair["model"])) - - data = ReflectionData() - data.load_mtz(str(sample_structure_pair["reflections"])) - - # Check if model has fcalc calculation method - if hasattr(model, 'calc_fcalc'): - fcalc = model.calc_fcalc(data) - - # Fcalc should have same number of reflections as data - assert fcalc.shape[0] == data.hkl.shape[0] class TestScalingWithRealData: @@ -263,8 +242,8 @@ class TestScalingWithRealData: @pytest.mark.integration def test_scaler_initialization_with_real_data(self, sample_structure_pair): """Test scaler initialization with real model and data.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler model = Model() @@ -285,8 +264,8 @@ def test_scaler_initialization_with_real_data(self, sample_structure_pair): @pytest.mark.integration def test_anisotropy_correction_values(self, sample_structure_pair): """Test that anisotropy correction produces reasonable values.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler model = Model() @@ -315,8 +294,8 @@ class TestMathFunctionsFunctional: @pytest.mark.integration def test_scattering_vectors_from_real_data(self, sample_structure_pair): """Test scattering vector calculation with real HKL and cell.""" - from torchref.io import ReflectionData from torchref.base.math_torch import get_scattering_vectors + from torchref.io import ReflectionData data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) @@ -333,11 +312,11 @@ def test_scattering_vectors_from_real_data(self, sample_structure_pair): @pytest.mark.integration def test_coordinate_transformations_with_real_cell(self, sample_cif_file): """Test coordinate transformations with real unit cell.""" - from torchref.model.model import Model from torchref.base.math_torch import ( cartesian_to_fractional_torch, - fractional_to_cartesian_torch + fractional_to_cartesian_torch, ) + from torchref.model.model import Model model = Model() model.load_cif(str(sample_cif_file)) @@ -439,8 +418,8 @@ class TestNLLFunctionsFunctional: @pytest.mark.integration def test_nll_xray_with_identical_data(self, sample_structure_pair): """Test NLL is minimal when Fobs equals Fcalc.""" - from torchref.io import ReflectionData from torchref.base.math_torch import nll_xray + from torchref.io import ReflectionData data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) @@ -463,8 +442,8 @@ def test_nll_xray_with_identical_data(self, sample_structure_pair): @pytest.mark.integration def test_nll_xray_increases_with_error(self, sample_structure_pair): """Test NLL increases as Fcalc differs from Fobs.""" - from torchref.io import ReflectionData from torchref.base.math_torch import nll_xray + from torchref.io import ReflectionData data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) @@ -494,8 +473,8 @@ def test_nll_xray_increases_with_error(self, sample_structure_pair): @pytest.mark.integration def test_nll_xray_lognormal(self, sample_structure_pair): """Test lognormal NLL calculation.""" - from torchref.io import ReflectionData from torchref.base.math_torch import nll_xray_lognormal + from torchref.io import ReflectionData data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) @@ -520,8 +499,9 @@ class TestRiceDistributionFunctional: @pytest.mark.integration def test_rice_nll_acentric(self, sample_structure_pair): """Test Rice NLL for acentric reflections.""" - from torchref.io import ReflectionData from torch.special import i0 + + from torchref.io import ReflectionData data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) @@ -592,8 +572,8 @@ def test_sigma_weighting(self, sample_structure_pair): @pytest.mark.integration def test_resolution_weighting(self, sample_structure_pair): """Test resolution-based weighting.""" - from torchref.io import ReflectionData from torchref.base.math_torch import get_scattering_vectors + from torchref.io import ReflectionData data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) @@ -704,9 +684,9 @@ class TestCombinedLossFunctional: @pytest.mark.integration def test_xray_plus_geometry_loss(self, sample_structure_pair, external_monomer_library): """Test combining X-ray and geometry losses.""" - from torchref.model.model import Model - from torchref.io import ReflectionData from torchref.base.math_torch import nll_xray + from torchref.io import ReflectionData + from torchref.model.model import Model model = Model() model.load_cif(str(sample_structure_pair["model"])) diff --git a/tests/integration/test_io_cif.py b/tests/integration/test_io_cif.py index ff288cc5..7e8b23d7 100644 --- a/tests/integration/test_io_cif.py +++ b/tests/integration/test_io_cif.py @@ -6,85 +6,37 @@ import pytest import torch -from pathlib import Path - -class TestCIFLoading: - """Tests for loading CIF model files.""" - - @pytest.mark.integration - def test_load_model_cif(self, sample_cif_file): - """Test loading a real CIF model file.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - - # Basic checks - use xyz().shape[0] for atom count - n_atoms = model.xyz().shape[0] - assert n_atoms > 0 - assert hasattr(model, 'xyz') - assert hasattr(model, 'adp') - assert hasattr(model, 'occupancy') - - @pytest.mark.integration - def test_model_atom_counts(self, sample_cif_file): - """Test that model has consistent atom counts.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - - # All arrays should have same number of atoms - n_atoms = model.xyz().shape[0] - assert model.xyz().shape[0] == n_atoms - assert model.adp().shape[0] == n_atoms - assert model.occupancy().shape[0] == n_atoms - - @pytest.mark.integration - def test_model_cell_parameters(self, sample_cif_file): - """Test that model has valid cell parameters.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - - # Cell should have 6 parameters - assert len(model.cell) == 6 - # All cell parameters should be positive - assert all(p > 0 for p in model.cell[:3].tolist()) # a, b, c - # Angles should be reasonable (0-180) - assert all(0 < p <= 180 for p in model.cell[3:].tolist()) # alpha, beta, gamma - - @pytest.mark.integration - def test_model_spacegroup(self, sample_cif_file): - """Test that model has a valid spacegroup.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - - # Spacegroup should be set (can be string or gemmi.SpaceGroup) - assert model.spacegroup is not None - # Check it can be converted to string representation - assert len(str(model.spacegroup)) > 0 - - @pytest.mark.integration - def test_model_element_types(self, sample_cif_file): - """Test that element types are recognized.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - - # Should have pdb DataFrame with element column - assert hasattr(model, 'pdb') - assert 'element' in model.pdb.columns - # Elements should be strings like 'C', 'N', 'O', etc. - elements = set(model.pdb['element'].unique()) - common_elements = {'C', 'N', 'O', 'S', 'H', 'CA', 'MG', 'ZN', 'FE'} - # At least some elements should be recognized - assert len(elements.intersection(common_elements)) > 0 or len(elements) > 0 +from torchref.config import canonical_device, get_default_device, get_float_dtype + + +@pytest.mark.integration +def test_cif_loading_contract(loaded_model, sample_cif_file) -> None: + """A deposited CIF supplies aligned atomic tensors and its crystal metadata.""" + import gemmi + + model = loaded_model + reference = gemmi.read_structure(str(sample_cif_file)) + xyz, adp, occupancy = model.xyz(), model.adp(), model.occupancy() + assert xyz.shape == (len(model.pdb), 3) + assert len(xyz) > 0 + assert adp.shape == occupancy.shape == (len(xyz),) + for tensor in (xyz, adp, occupancy, model.cell.data): + assert tensor.dtype == get_float_dtype() + assert canonical_device(tensor.device) == canonical_device(get_default_device()) + assert torch.isfinite(tensor).all() + assert torch.all(adp >= 0) + assert {"x", "y", "z", "element", "resname", "chainid", "resseq"} <= set( + model.pdb.columns + ) + assert {"C", "N", "O"} <= set(model.pdb.element) + torch.testing.assert_close( + model.cell.data, xyz.new_tensor(reference.cell.parameters) + ) + assert ( + model.spacegroup.number + == gemmi.find_spacegroup_by_name(reference.spacegroup_hm).number + ) class TestMultipleCIFFiles: @@ -125,18 +77,18 @@ class TestCIFSaving: def test_save_and_reload_cif(self, sample_cif_file, tmp_path): """Test saving a model to CIF and reloading it.""" from torchref.model.model import Model - + # Load original model1 = Model() model1.load_cif(str(sample_cif_file)) n_atoms1 = model1.xyz().shape[0] - + # Save to temp file using write_pdb (CIF saving may not exist) output_path = tmp_path / "test_output.pdb" model1.write_pdb(str(output_path)) - + assert output_path.exists() - + # add_hydrogens=False on reload: what is under test is whether the written # file round-trips, not whether generation reruns. Regenerating on reload can # legitimately differ, because ``write_pdb`` does not emit LINK records -- so a @@ -145,6 +97,6 @@ def test_save_and_reload_cif(self, sample_cif_file, tmp_path): model2 = Model(add_hydrogens=False) model2.load_pdb(str(output_path)) n_atoms2 = model2.xyz().shape[0] - + # Compare atom counts assert n_atoms2 == n_atoms1 diff --git a/tests/integration/test_io_reflections.py b/tests/integration/test_io_reflections.py index 79642253..cb9e9a1b 100644 --- a/tests/integration/test_io_reflections.py +++ b/tests/integration/test_io_reflections.py @@ -6,106 +6,72 @@ import pytest import torch -from pathlib import Path - -class TestMTZLoading: - """Tests for loading MTZ reflection files.""" - - @pytest.mark.integration - def test_load_mtz_file(self, sample_mtz_file): - """Test loading a real MTZ file.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # Should have reflections loaded - assert hasattr(data, 'hkl') - assert hasattr(data, 'F') - assert data.hkl is not None - - @pytest.mark.integration - def test_mtz_reflection_counts(self, sample_mtz_file): - """Test that MTZ has consistent reflection counts.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - n_refl = data.hkl.shape[0] - assert n_refl > 0 - - # F should match hkl count - if data.F is not None: - assert data.F.shape[0] == n_refl - - @pytest.mark.integration - def test_mtz_hkl_indices(self, sample_mtz_file): - """Test HKL indices are valid integers.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # HKL should have 3 columns - assert data.hkl.shape[1] == 3 - - # Should contain integer-like values - hkl_rounded = torch.round(data.hkl) - assert torch.allclose(data.hkl, hkl_rounded) - - @pytest.mark.integration - def test_mtz_cell_parameters(self, sample_mtz_file): - """Test that MTZ has valid cell parameters.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - if hasattr(data, 'cell') and data.cell is not None: - assert len(data.cell) == 6 - assert all(c > 0 for c in data.cell[:3].tolist()) - - @pytest.mark.integration - def test_mtz_spacegroup(self, sample_mtz_file): - """Test that MTZ has a valid spacegroup.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # Spacegroup should be set (can be string or gemmi.SpaceGroup) - assert data.spacegroup is not None - # Check it can be converted to string representation - assert len(str(data.spacegroup)) > 0 - - @pytest.mark.integration - def test_mtz_sigma_values(self, sample_mtz_file): - """Test that sigma values are loaded.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - if hasattr(data, 'F_sigma') and data.F_sigma is not None: - assert data.F_sigma.shape[0] == data.F.shape[0] - # Check that non-NaN sigma values are positive - valid_sigma = data.F_sigma[~torch.isnan(data.F_sigma)] - if len(valid_sigma) > 0: - assert torch.all(valid_sigma > 0) - - @pytest.mark.integration - def test_mtz_rfree_flags(self, sample_mtz_file): - """Test that R-free flags are loaded or generated.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # Should have rfree_flags (loaded or generated) - if hasattr(data, 'rfree_flags') and data.rfree_flags is not None: - assert data.rfree_flags.shape[0] == data.hkl.shape[0] +from torchref.config import ( + canonical_device, + get_default_device, + get_float_dtype, + get_int_dtype, +) + + +@pytest.mark.integration +def test_mtz_loading_contract(loaded_reflection_data, sample_mtz_file) -> None: + """MTZ loading supplies aligned observations, masks and crystal metadata.""" + import gemmi + + data = loaded_reflection_data + reference = gemmi.read_mtz_file(str(sample_mtz_file)) + n = len(data.hkl) + assert n > 0 + assert data.hkl.shape == (n, 3) + assert data.hkl.dtype == get_int_dtype() + assert canonical_device(data.hkl.device) == canonical_device(get_default_device()) + for tensor in (data.F, data.F_sigma, data.resolution): + assert tensor.shape == (n,) + assert tensor.dtype == get_float_dtype() + assert canonical_device(tensor.device) == canonical_device(get_default_device()) + mask = data.masks() + assert mask.shape == (n,) + assert mask.dtype == torch.bool + assert mask.any() + assert torch.isfinite(data.F[mask]).all() + assert torch.all(data.F[mask] >= 0) + assert torch.isfinite(data.F_sigma[mask]).all() + assert torch.all(data.F_sigma[mask] > 0) + assert torch.isfinite(data.resolution).all() + assert torch.all(data.resolution > 0) + assert data.rfree_flags.shape == (n,) + assert data.rfree_flags.dtype == torch.bool + assert data.rfree_flags.any() and (~data.rfree_flags).any() + assert 0.7 < data.rfree_flags.to(get_float_dtype()).mean().item() < 1.0 + torch.testing.assert_close( + data.cell.data, data.F.new_tensor(reference.cell.parameters) + ) + assert data.spacegroup.number == reference.spacegroup.number + + +@pytest.mark.integration +def test_resolution_bins(loaded_reflection_data) -> None: + """Every bin mean equals the mean d-spacing of its unmasked reflections.""" + data = loaded_reflection_data + bins, n_bins = data.get_bins(n_bins=10) + assert bins.shape == (len(data.hkl),) + assert n_bins > 0 + assert bins.min() >= 0 and bins.max() < n_bins + groups = [(bins == i) & data.masks() for i in range(n_bins)] + assert all(group.any() for group in groups) + expected = torch.stack([data.resolution[group].mean() for group in groups]) + torch.testing.assert_close(data.mean_res_per_bin(), expected) + + +@pytest.mark.integration +def test_structure_pair_consistency(model_and_data) -> None: + """Matching model and reflection files describe the same crystal.""" + model, data = model_and_data["model"], model_and_data["data"] + assert len(model.xyz()) > 0 and len(data.hkl) > 0 + torch.testing.assert_close(model.cell.data, data.cell.data, rtol=0.01, atol=0.1) + assert model.spacegroup.number == data.spacegroup.number class TestSFCIFLoading: @@ -115,95 +81,27 @@ class TestSFCIFLoading: def test_load_sf_cif(self, sample_structure_factor_cif): """Test loading a structure factor CIF file.""" from torchref.io import ReflectionData - + data = ReflectionData() data.load_cif(str(sample_structure_factor_cif)) - - assert data.hkl is not None + + assert data.hkl.shape[0] > 0 + assert data.hkl.shape[1] == 3 class TestReflectionDataProperties: """Tests for computed properties of reflection data.""" - @pytest.mark.integration - def test_resolution_calculation(self, sample_mtz_file): - """Test resolution can be calculated from loaded data.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # Should have resolution attribute - if hasattr(data, 'resolution') and data.resolution is not None: - assert torch.all(data.resolution > 0) - assert torch.all(torch.isfinite(data.resolution)) - - @pytest.mark.integration - def test_wilson_b_factor(self, sample_mtz_file): - """Test Wilson B-factor is calculated.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # Wilson B should be calculated during loading - if hasattr(data, 'wilson_b') and data.wilson_b is not None: - assert data.wilson_b > 0 - @pytest.mark.integration def test_data_device_movement(self, sample_mtz_file, cpu_device): """Test moving reflection data to different devices.""" from torchref.io import ReflectionData - + data = ReflectionData() data.load_mtz(str(sample_mtz_file)) - + # Move to device data = data.to(cpu_device) - - # Tensors should be on correct device - if data.hkl is not None: - assert data.hkl.device == cpu_device - if data.F is not None: - assert data.F.device == cpu_device - - -class TestMatchingDataPairs: - """Tests using matching model and reflection data.""" - @pytest.mark.integration - def test_load_structure_pair(self, sample_structure_pair): - """Test loading matching model and reflection data.""" - from torchref.model.model import Model - from torchref.io import ReflectionData - - model = Model() - model.load_cif(str(sample_structure_pair["model"])) - - data = ReflectionData() - data.load_mtz(str(sample_structure_pair["reflections"])) - - # Both should load successfully - n_atoms = model.xyz().shape[0] - assert n_atoms > 0 - assert data.hkl is not None - - @pytest.mark.integration - def test_cell_consistency(self, sample_structure_pair): - """Test that model and data have consistent cell parameters.""" - from torchref.model.model import Model - from torchref.io import ReflectionData - - model = Model() - model.load_cif(str(sample_structure_pair["model"])) - - data = ReflectionData() - data.load_mtz(str(sample_structure_pair["reflections"])) - - # Cell parameters should be similar (may have small differences) - if hasattr(data, 'cell') and data.cell is not None: - model_cell = torch.tensor(model.cell) - data_cell = torch.tensor(data.cell) - - # Allow 1% tolerance for cell parameters - assert torch.allclose(model_cell, data_cell, rtol=0.01, atol=0.1) + assert data.hkl.device == cpu_device + assert data.F.device == cpu_device From fb97d23438e7f8127e3c2419270e6d7936cf4bc8 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 7 Sep 2026 17:12:03 +0200 Subject: [PATCH 164/250] test: make extended structure coverage explicit and lazy Replace CIF/MTZ/SF-CIF loops with named per-file compatibility cases and an input inventory guard. Keep 1DAW reader contracts in the quick suite; require --run-slow for extras. Replace eager all_test_structures with one fresh named crystal per scaler/restraint case. Move the extra ModelFT loading cases and symmetry file sweep coverage into the compatibility panel. Merge space-group name cases under their unit owner without dropping parameter variants. Validation: 44 extended cases passed; final regression 477 passed/70 skipped, including 41 explicitly slow cases. Full collection: 2587 cases. No production files changed. --- docs/changelog.rst | 1 + docs/user_guide/testing.rst | 26 +- tests/README.md | 9 + tests/fixtures/README.md | 7 +- tests/fixtures/files.py | 43 ++-- tests/fixtures/objects.py | 61 ++--- tests/functional/test_io_functional.py | 79 ------ tests/functional/test_model_ft_functional.py | 26 -- .../functional/test_restraints_functional.py | 157 +++++------- tests/functional/test_scaler_functional.py | 232 ++++++++---------- tests/helpers/structure_cases.py | 27 ++ tests/integration/test_io_cif.py | 31 --- .../test_structure_compatibility.py | 101 ++++++++ .../integration/test_symmetry_integration.py | 41 +--- tests/unit/symmetry/test_symmetry.py | 29 ++- 15 files changed, 398 insertions(+), 472 deletions(-) delete mode 100644 tests/functional/test_io_functional.py create mode 100644 tests/helpers/structure_cases.py create mode 100644 tests/integration/test_structure_compatibility.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 011122e3..3cb46505 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Named broad structure-compatibility cases explicitly, moved extra datasets to the slow tier, and removed eager all-model loading and swallowed reader failures. - Consolidated CIF/MTZ loading contracts, checked configured tensor placement, and replaced ModelFT smoke checks with exercised forward-cache behavior. - Replaced local-arithmetic target tests with configured-device production-kernel checks on deposited coordinates and explicit least-squares expectations. - Consolidated weighting tests by API ownership and strengthened Gaussian-likelihood, gradient-norm, and cached-loss assertions. diff --git a/docs/user_guide/testing.rst b/docs/user_guide/testing.rst index cbdc1c8c..66e60273 100644 --- a/docs/user_guide/testing.rst +++ b/docs/user_guide/testing.rst @@ -39,7 +39,13 @@ PDB ID d_min (Å) Space group ``tests/files/`` also holds partial sets — ``1AK5_with_H.pdb`` + ``1AK5.mtz`` (no CIF), ``7L84.pdb`` + ``7L84-sf.cif`` (no MTZ), ``test_ihm_ensemble.cif`` — so a test that globs one directory and assumes a matching file in another will -fail on those. Use ``sample_structure_pair`` / ``all_test_structures``. +fail on those. Use ``sample_structure_pair`` for the quick reference crystal, +or ``compatibility_structure_pair`` for named extended cases. The latter carries +the ``slow`` marker and selects paths without loading objects. + +``tests/helpers/structure_cases.py`` assigns bundled CIF, MTZ and SF-CIF files to +the quick or extended compatibility panel. Additional files require an explicit +assignment; directory growth does not silently expand numerical test work. Running Tests ------------- @@ -99,11 +105,11 @@ The Amber stack, if you want it: Fixtures -------- -Almost everything lives in the root ``tests/conftest.py`` and is therefore -available from every category — the ``integration/`` and ``functional/`` -conftests are docstrings only. Mock data is the exception: -``tests/unit/conftest.py``. Read those two files for the authoritative list; the -ones you will reach for most: +Reusable setup lives in ``tests/fixtures/``. The root ``tests/conftest.py`` +registers shared plugins and owns test-selection hooks. The unit conftest exposes +synthetic numerical factories; the functional conftest exposes the module-scoped, +read-only Fourier-model fixture. See ``tests/fixtures/README.md`` for ownership +and mutation rules. Common fixtures include: - Paths (session-scoped): ``tests_root``, ``project_root``, ``test_files_dir``, the per-format ``cif_dir``, ``mtz_dir``, ``pdb_dir``, ``cif_sf_dir``, and @@ -118,9 +124,11 @@ ones you will reach for most: ``mock_aniso_u``, ``mock_scattering_factors``, ``mock_weights``. - Real files: ``sample_cif_file``, ``sample_pdb_file``, ``sample_mtz_file``, ``sample_structure_factor_cif``, ``sample_structure_pair`` (matched model + - data), ``all_structure_pairs``, ``all_test_structures``. -- Loaded objects: ``loaded_model``, ``loaded_reflection_data``, - ``model_and_data``, ``initialized_scaler``. + data), ``compatibility_structure_pair`` (one named slow crystal). +- Loaded objects: ``loaded_model``, ``loaded_model_ft``, ``loaded_reflection_data``, + ``model_and_data``, ``initialized_scaler``. ``compatibility_model`` and + ``compatibility_model_and_data`` load only the current slow case and remain + function-scoped to isolate mutations. The mock-data fixtures yield a *factory* taking ``n_atoms`` / ``n_reflections`` and ``seed``; ``mock_cell`` and ``mock_cell_triclinic`` yield the tensor diff --git a/tests/README.md b/tests/README.md index 9298352f..3be67e2f 100644 --- a/tests/README.md +++ b/tests/README.md @@ -51,6 +51,7 @@ tests/ | CIF atomic fields and crystal metadata | `integration/test_io_cif.py` | | MTZ fields, resolution bins and model/data crystal agreement | `integration/test_io_reflections.py` | | ModelFT forward cache and grid integration | `functional/test_model_ft_functional.py` | +| Extra deposited files and input inventory | `integration/test_structure_compatibility.py`, `helpers/structure_cases.py` | | Numerical derivatives and backend parity | `unit/test_gradient_correctness.py`, `unit/structure_factor/` | A production call must participate in the assertion: computing a formula only in @@ -58,6 +59,14 @@ the test does not check its implementation. Kernel values, target registration, device transitions, and default configuration are separate contracts even when they exercise the same class. Keep mutation tests on fresh objects. +The quick reader contracts use 1DAW. Extended reader compatibility runs with +`pytest tests/integration/test_structure_compatibility.py --run-slow`; each file +is a separate case and must succeed. The manifest covers the bundled CIF, MTZ +and SF-CIF inputs, including the IHM fixture and reflection-only depositions. +Adding a data file requires an explicit coverage assignment in the manifest. +Extended scaler and restraint cases use 2DQ6 (trigonal) and 3A5V (body-centred +tetragonal), with fresh objects per case and `--run-slow` required. + ### Quick Local Run (on login node, for small tests only) ```bash diff --git a/tests/fixtures/README.md b/tests/fixtures/README.md index 7361c509..282c5527 100644 --- a/tests/fixtures/README.md +++ b/tests/fixtures/README.md @@ -7,7 +7,7 @@ Keep a fixture in its test module when only that module needs it. | Module | Responsibility | Visibility / lifetime | |---|---|---| | `paths.py` | Repository, bundled-data and optional library paths | All tests; session | -| `files.py` | Sample-file selection and matching structure pairs | All tests; session; no model loading | +| `files.py` | Sample paths and named compatibility pairs | All tests; sample paths session-scoped, extended pairs function-scoped; no loading | | `devices.py` | Configured device, explicit backends, device parametrization | All tests; existing per-fixture scopes | | `precision.py` | Comparison tolerances and CPU-double reference context | All tests; reference fixture restores state after each test | | `objects.py` | Mutable models, data, scalers and restraints | All tests; fresh per test except explicitly shared bundles | @@ -32,6 +32,11 @@ device movement, or empty caches use fresh objects. `loaded_model`, fresh mutable objects per test. The explicitly shared session bundles in that module retain their documented ownership contracts. +`compatibility_structure_pair` selects named slow cases from +`tests/helpers/structure_cases.py`. `compatibility_model` loads just that model; +`compatibility_model_and_data` adds observations only when needed. Skipped slow +cases do not load any structures. + Use `cpu_double_precision()` to scope an explicit numerical reference, or request `double_cpu` for a single test. The structure-factor package uses the same context at package scope; both usages restore dtype, device, and density cutoff on exit. diff --git a/tests/fixtures/files.py b/tests/fixtures/files.py index d248775e..390ea678 100644 --- a/tests/fixtures/files.py +++ b/tests/fixtures/files.py @@ -4,6 +4,23 @@ import pytest +from tests.helpers.structure_cases import EXTENDED_PAIR_CODES + + +@pytest.fixture( + params=[pytest.param(code, marks=pytest.mark.slow) for code in EXTENDED_PAIR_CODES] +) +def compatibility_structure_pair( + cif_dir: Path, mtz_dir: Path, request: pytest.FixtureRequest +) -> dict: + """Select one named extended crystal without loading its model or observations.""" + code = request.param + return { + "pdb_id": code, + "model": cif_dir / f"{code}.cif", + "reflections": mtz_dir / f"{code}.mtz", + } + @pytest.fixture(scope="session") def sample_cif_file(cif_dir: Path) -> Path: @@ -70,29 +87,3 @@ def sample_structure_pair(cif_dir: Path, mtz_dir: Path) -> dict[str, Path]: return {"model": cif_files[pdb_id], "reflections": mtz_files[pdb_id]} pytest.skip("No matching CIF/MTZ pairs found in test data") - - -@pytest.fixture(scope="session") -def all_structure_pairs(cif_dir: Path, mtz_dir: Path) -> list[dict[str, Path | str]]: - """Return all matching pairs of CIF models and MTZ reflections.""" - cif_files = {f.stem: f for f in cif_dir.glob("*.cif")} - mtz_files = {f.stem: f for f in mtz_dir.glob("*.mtz")} - - common_ids = set(cif_files.keys()) & set(mtz_files.keys()) - - if not common_ids: - pytest.skip("No matching CIF/MTZ pairs found in test data") - - return [ - {"pdb_id": pdb_id, "model": cif_files[pdb_id], "reflections": mtz_files[pdb_id]} - for pdb_id in sorted(common_ids) - ] - - -@pytest.fixture(scope="session") -def all_cif_files(cif_dir: Path) -> list[Path]: - """Return all available CIF test structure files.""" - cif_files = sorted(cif_dir.glob("*.cif")) - if not cif_files: - pytest.skip("No CIF files found in test data directory") - return cif_files diff --git a/tests/fixtures/objects.py b/tests/fixtures/objects.py index 34fb00c8..3eb7628e 100644 --- a/tests/fixtures/objects.py +++ b/tests/fixtures/objects.py @@ -18,6 +18,31 @@ from torchref.scaling import Scaler +@pytest.fixture +def compatibility_model(compatibility_structure_pair: dict) -> Model: + """Load a fresh model for one slow compatibility case.""" + from torchref.model import Model + + path = compatibility_structure_pair["model"] + assert path.is_file() + return Model(verbose=0).load_cif(str(path)) + + +@pytest.fixture +def compatibility_model_and_data( + compatibility_model: Model, compatibility_structure_pair: dict +) -> dict: + """Load observations only for the single crystal used by the current pipeline case.""" + from torchref.io import ReflectionData + + path = compatibility_structure_pair["reflections"] + assert path.is_file() + return { + "model": compatibility_model, + "data": ReflectionData(verbose=0).load_mtz(str(path)), + } + + @pytest.fixture def loaded_model(sample_cif_file: Path) -> Model: """Load a fresh mutable Model from the sample CIF file.""" @@ -97,42 +122,6 @@ def model_with_restraints(loaded_model: Model) -> dict[str, Any]: return {"model": loaded_model, "restraints": restraints} -@pytest.fixture(scope="session") -def all_test_structures( - all_structure_pairs: list[dict[str, Any]], -) -> list[dict[str, Any]]: - """Return all loaded model/data pairs for comprehensive testing.""" - from torchref.io import ReflectionData - from torchref.model.model import Model - - structures = [] - for pair in all_structure_pairs: - try: - model = Model() - model.load_cif(str(pair["model"])) - - data = ReflectionData() - data.load_mtz(str(pair["reflections"])) - - structures.append( - { - "pdb_id": pair["pdb_id"], - "model": model, - "data": data, - "model_path": pair["model"], - "data_path": pair["reflections"], - } - ) - except Exception: - # Skip structures that fail to load - continue - - if not structures: - pytest.skip("No structures could be loaded") - - return structures - - @pytest.fixture(scope="session") def _device_model_cache() -> dict: """``{device_str: ModelFT}`` built at most once per device, per session.""" diff --git a/tests/functional/test_io_functional.py b/tests/functional/test_io_functional.py deleted file mode 100644 index 49c82a99..00000000 --- a/tests/functional/test_io_functional.py +++ /dev/null @@ -1,79 +0,0 @@ -""" -Functional tests for I/O operations. - -Tests file loading and data processing with real crystallographic data. -""" - -import pytest - - -class TestCIFReadingFunctional: - """Functional tests for CIF file reading.""" - - @pytest.mark.integration - def test_load_multiple_cif_files(self, cif_dir): - """Test loading multiple CIF files successfully.""" - from torchref.model.model import Model - - cif_files = list(cif_dir.glob("*.cif")) - assert len(cif_files) > 0, "No CIF files found in test directory" - - for cif_file in cif_files: - model = Model() - model.load_cif(str(cif_file)) - - # Each file should load with atoms - n_atoms = model.xyz().shape[0] - assert n_atoms > 0, f"No atoms loaded from {cif_file}" - - # Should have cell parameters - assert model.cell is not None - assert len(model.cell) == 6 - - -class TestMTZReadingFunctional: - """Functional tests for MTZ file reading.""" - - @pytest.mark.integration - def test_load_multiple_mtz_files(self, mtz_dir): - """Test loading multiple MTZ files successfully.""" - from torchref.io import ReflectionData - - mtz_files = list(mtz_dir.glob("*.mtz")) - assert len(mtz_files) > 0, "No MTZ files found in test directory" - - for mtz_file in mtz_files: - data = ReflectionData() - data.load_mtz(str(mtz_file)) - - # Each file should load with reflections - n_refl = data.hkl.shape[0] - assert n_refl > 0, f"No reflections loaded from {mtz_file}" - - # Should have cell parameters - assert data.cell is not None - - -class TestSFCIFReadingFunctional: - """Functional tests for structure factor CIF reading.""" - - @pytest.mark.integration - def test_load_sf_cif(self, cif_sf_dir): - """Test loading structure factor CIF files.""" - from torchref.io import ReflectionData - - sf_files = list(cif_sf_dir.glob("*.cif")) - if not sf_files: - pytest.skip("No SF-CIF files found") - - for sf_file in sf_files: - data = ReflectionData() - try: - data.load_cif(str(sf_file)) - - # Should have loaded reflections - if data.hkl is not None: - assert data.hkl.shape[0] > 0 - except Exception as e: - # Some files may not be valid SF-CIF format - pass diff --git a/tests/functional/test_model_ft_functional.py b/tests/functional/test_model_ft_functional.py index 2c6b9b4c..ce9b5e71 100644 --- a/tests/functional/test_model_ft_functional.py +++ b/tests/functional/test_model_ft_functional.py @@ -104,32 +104,6 @@ def test_map_symmetry_available(self, shared_model_ft): assert operator.map_shape == gridsize -@pytest.mark.integration -class TestModelFTMultipleStructures: - """Test ModelFT with multiple structures.""" - - def test_modelft_multiple_structures(self, all_structure_pairs): - """Test ModelFT works with different structures.""" - from torchref.model.model_ft import ModelFT - - tested = 0 - for pair in all_structure_pairs[:3]: # Test first 3 - try: - model = ModelFT(max_res=3.0, verbose=0) - model.load_cif(str(pair["model"])) - - # Basic checks - assert model.xyz() is not None - assert model.xyz().shape[0] > 0 - - tested += 1 - except Exception as e: - # Some structures may fail to load - continue - - assert tested >= 1, "At least one structure should load" - - @pytest.mark.integration class TestModelFTCoordinateOperations: """Test ModelFT coordinate operations.""" diff --git a/tests/functional/test_restraints_functional.py b/tests/functional/test_restraints_functional.py index 76534421..841022cc 100644 --- a/tests/functional/test_restraints_functional.py +++ b/tests/functional/test_restraints_functional.py @@ -21,10 +21,7 @@ def test_build_restraints_from_cif(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() @@ -42,34 +39,31 @@ def test_bond_restraints_built(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() # Check bond restraints exist - assert 'bond' in restraints.restraints + assert "bond" in restraints.restraints # Check intra-residue bonds - if 'intra' in restraints.restraints['bond']: - bond_intra = restraints.restraints['bond']['intra'] - assert 'indices' in bond_intra - assert 'references' in bond_intra - assert 'sigmas' in bond_intra + if "intra" in restraints.restraints["bond"]: + bond_intra = restraints.restraints["bond"]["intra"] + assert "indices" in bond_intra + assert "references" in bond_intra + assert "sigmas" in bond_intra # Indices should be 2D with shape (N, 2) - indices = bond_intra['indices'] + indices = bond_intra["indices"] assert len(indices.shape) == 2 assert indices.shape[1] == 2 # References should match number of bonds - assert bond_intra['references'].shape[0] == indices.shape[0] - assert bond_intra['sigmas'].shape[0] == indices.shape[0] + assert bond_intra["references"].shape[0] == indices.shape[0] + assert bond_intra["sigmas"].shape[0] == indices.shape[0] # Bond lengths should be positive and reasonable (0.5-3.0 Å) - refs = bond_intra['references'] + refs = bond_intra["references"] assert torch.all(refs > 0.5) assert torch.all(refs < 3.0) @@ -83,29 +77,26 @@ def test_angle_restraints_built(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() # Check angle restraints exist - assert 'angle' in restraints.restraints + assert "angle" in restraints.restraints - if 'intra' in restraints.restraints['angle']: - angle_intra = restraints.restraints['angle']['intra'] - assert 'indices' in angle_intra - assert 'references' in angle_intra - assert 'sigmas' in angle_intra + if "intra" in restraints.restraints["angle"]: + angle_intra = restraints.restraints["angle"]["intra"] + assert "indices" in angle_intra + assert "references" in angle_intra + assert "sigmas" in angle_intra # Indices should be 2D with shape (N, 3) - indices = angle_intra['indices'] + indices = angle_intra["indices"] assert len(indices.shape) == 2 assert indices.shape[1] == 3 # References should match number of angles - assert angle_intra['references'].shape[0] == indices.shape[0] + assert angle_intra["references"].shape[0] == indices.shape[0] @pytest.mark.integration def test_torsion_restraints_built(self, sample_cif_file): @@ -117,25 +108,22 @@ def test_torsion_restraints_built(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() # Check torsion restraints exist - assert 'torsion' in restraints.restraints + assert "torsion" in restraints.restraints - if 'intra' in restraints.restraints['torsion']: - torsion_intra = restraints.restraints['torsion']['intra'] - assert 'indices' in torsion_intra - assert 'references' in torsion_intra - assert 'sigmas' in torsion_intra - assert 'periods' in torsion_intra + if "intra" in restraints.restraints["torsion"]: + torsion_intra = restraints.restraints["torsion"]["intra"] + assert "indices" in torsion_intra + assert "references" in torsion_intra + assert "sigmas" in torsion_intra + assert "periods" in torsion_intra # Indices should be 2D with shape (N, 4) - indices = torsion_intra['indices'] + indices = torsion_intra["indices"] assert len(indices.shape) == 2 assert indices.shape[1] == 4 @@ -149,23 +137,20 @@ def test_plane_restraints_built(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() # Check plane restraints exist - assert 'plane' in restraints.restraints + assert "plane" in restraints.restraints # Planes are grouped by atom count (3_atoms, 4_atoms, etc.) - plane_restraints = restraints.restraints['plane'] + plane_restraints = restraints.restraints["plane"] if len(list(plane_restraints.keys())) > 0: # Check at least one plane group exists for key, plane_group in plane_restraints.items(): - if 'indices' in plane_group: - indices = plane_group['indices'] + if "indices" in plane_group: + indices = plane_group["indices"] # Planes need at least 3 atoms if len(indices.shape) == 2: assert indices.shape[1] >= 3 @@ -184,15 +169,12 @@ def test_bond_deviations(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() # Compute bond deviations - if hasattr(restraints, 'bond_deviations'): + if hasattr(restraints, "bond_deviations"): deviations, sigmas = restraints.bond_deviations() assert torch.all(torch.isfinite(deviations)) @@ -211,15 +193,12 @@ def test_angle_deviations(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() # Compute angle deviations - if hasattr(restraints, 'angle_deviations'): + if hasattr(restraints, "angle_deviations"): deviations, sigmas = restraints.angle_deviations() assert torch.all(torch.isfinite(deviations)) @@ -231,28 +210,20 @@ class TestRestraintsMultipleStructures: @pytest.mark.integration @pytest.mark.slow - def test_restraints_multiple_cif_files(self, cif_dir): - """Test building restraints for multiple CIF files.""" - from torchref.model.model import Model + def test_restraints_multiple_cif_files(self, compatibility_model): + """Each extended crystal supplies bond and angle restraints.""" from torchref.topology.restraints import Restraints - cif_files = list(cif_dir.glob("*.cif"))[:3] # First 3 structures - - for cif_file in cif_files: - model = Model() - model.load_cif(str(cif_file)) - - restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 - ) - restraints.build_restraints() - - # Should have built restraints for each structure - assert 'bond' in restraints.restraints - assert 'angle' in restraints.restraints + model = compatibility_model + restraints = Restraints( + pdb=model.pdb, + xyz_fn=model.xyz, + vdw_radii_fn=model.get_vdw_radii, + verbose=0, + ) + restraints.build_restraints() + assert "bond" in restraints.restraints + assert "angle" in restraints.restraints class TestRestraintsDeviceHandling: @@ -268,16 +239,13 @@ def test_restraints_device_movement(self, sample_cif_file, cpu_device): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() # Check that tensors are on the correct device - if 'bond' in restraints.restraints and 'intra' in restraints.restraints['bond']: - bond_indices = restraints.restraints['bond']['intra']['indices'] + if "bond" in restraints.restraints and "intra" in restraints.restraints["bond"]: + bond_indices = restraints.restraints["bond"]["intra"]["indices"] assert bond_indices.device == cpu_device @@ -294,10 +262,7 @@ def test_cif_dict_loaded(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) # CIF dict should be populated with residue restraints @@ -305,10 +270,13 @@ def test_cif_dict_loaded(self, sample_cif_file): assert len(restraints.cif_dict) > 0 # Should have standard amino acids - common_residues = ['ALA', 'GLY', 'VAL', 'LEU', 'ILE'] + common_residues = ["ALA", "GLY", "VAL", "LEU", "ILE"] for res in common_residues: if res in restraints.cif_dict: - assert 'bonds' in restraints.cif_dict[res] or 'angles' in restraints.cif_dict[res] + assert ( + "bonds" in restraints.cif_dict[res] + or "angles" in restraints.cif_dict[res] + ) @pytest.mark.integration def test_unique_residues_detected(self, sample_cif_file): @@ -320,10 +288,7 @@ def test_unique_residues_detected(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) # Should have detected unique residues diff --git a/tests/functional/test_scaler_functional.py b/tests/functional/test_scaler_functional.py index d1364665..e6c3b3a7 100644 --- a/tests/functional/test_scaler_functional.py +++ b/tests/functional/test_scaler_functional.py @@ -6,7 +6,23 @@ import pytest import torch -import numpy as np + + +@pytest.mark.integration +def test_scaler_crystal_compatibility(compatibility_model_and_data) -> None: + """Each extended crystal produces finite anisotropic scale corrections.""" + from torchref.scaling import Scaler + + model = compatibility_model_and_data["model"] + data = compatibility_model_and_data["data"] + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) + scaler.setup_anisotropy_correction() + assert scaler.s is not None + assert scaler.bins is not None + assert scaler.U is not None + correction = scaler.anisotropy_correction() + assert correction.shape == (len(data.hkl),) + assert torch.isfinite(correction).all() class TestScalerCreationFunctional: @@ -15,18 +31,18 @@ class TestScalerCreationFunctional: @pytest.mark.integration def test_scaler_full_initialization(self, sample_structure_pair): """Test full scaler initialization with model and data.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=20, verbose=0) - + # Check all components are initialized assert scaler.model is not None assert scaler._data is not None @@ -39,8 +55,8 @@ def test_scaler_full_initialization(self, sample_structure_pair): @pytest.mark.parametrize("nbins", [5, 10, 15, 20]) def test_scaler_with_different_nbins(self, sample_structure_pair, nbins): """Test scaler with different bin counts.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler model = Model() @@ -62,18 +78,18 @@ class TestScatteringVectorsFunctional: @pytest.mark.integration def test_scattering_vectors_shape(self, sample_structure_pair): """Test scattering vectors have correct shape.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) - + # s should have shape (n_reflections, 3) n_refl = data.hkl.shape[0] assert scaler.s.shape == (n_refl, 3) @@ -81,21 +97,21 @@ def test_scattering_vectors_shape(self, sample_structure_pair): @pytest.mark.integration def test_scattering_vectors_magnitude(self, sample_structure_pair): """Test scattering vector magnitudes are reasonable.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) - + # Calculate |s| = sin(theta)/lambda = 1/(2d) s_mag = torch.norm(scaler.s, dim=1) - + # For typical protein data: # Low resolution (d=100Å): |s| ~ 0.005 # High resolution (d=1Å): |s| ~ 0.5 @@ -109,26 +125,26 @@ class TestAnisotropyCorrectionFunctional: @pytest.mark.integration def test_anisotropy_setup_and_compute(self, sample_structure_pair): """Test setting up and computing anisotropy correction.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_anisotropy_correction() - + # U parameters should exist - assert hasattr(scaler, 'U') + assert hasattr(scaler, "U") assert scaler.U.shape == (6,) # U11, U22, U33, U12, U13, U23 - + # Compute correction correction = scaler.anisotropy_correction() - + # Correction should be positive (exponential) assert correction.shape[0] == data.hkl.shape[0] assert torch.all(correction > 0) @@ -137,22 +153,22 @@ def test_anisotropy_setup_and_compute(self, sample_structure_pair): @pytest.mark.integration def test_anisotropy_correction_near_unity(self, sample_structure_pair): """Test anisotropy correction starts near unity with small U.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_anisotropy_correction() - + # With small random U values, correction should be close to 1 correction = scaler.anisotropy_correction() - + # Most values should be between 0.5 and 2.0 for small U mean_correction = correction.mean().item() assert 0.5 < mean_correction < 2.0 @@ -164,20 +180,20 @@ class TestBinwiseBfactorFunctional: @pytest.mark.integration def test_setup_binwise_bfactor(self, sample_structure_pair): """Test setting up bin-wise B-factor parameters.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_bin_wise_bfactor() - assert hasattr(scaler, 'bin_wise_bfactor') + assert hasattr(scaler, "bin_wise_bfactor") assert scaler.bin_wise_bfactor.shape == (10,) # Initially should be zeros assert torch.allclose( @@ -187,24 +203,24 @@ def test_setup_binwise_bfactor(self, sample_structure_pair): @pytest.mark.integration def test_binwise_bfactor_correction(self, sample_structure_pair): """Test computing bin-wise B-factor correction.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_bin_wise_bfactor() - + # Set some non-zero B-factors scaler.bin_wise_bfactor.data = torch.linspace(0, 20, 10, device=scaler.device) - + correction = scaler.bin_wise_bfactor_correction() - + # Correction should have same length as reflections assert correction.shape[0] == data.hkl.shape[0] # Should be positive (exponential) @@ -218,35 +234,35 @@ class TestScalerStateDictFunctional: @pytest.mark.integration def test_save_and_load_state_dict(self, sample_structure_pair, tmp_path): """Test saving and loading scaler state.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + # Create scaler with some setup scaler1 = Scaler(model=model, data=data, nbins=10, verbose=0) scaler1.setup_anisotropy_correction() scaler1.setup_bin_wise_bfactor() - + # Modify parameters scaler1.U.data = torch.randn(6, device=scaler1.device) scaler1.bin_wise_bfactor.data = torch.randn(10, device=scaler1.device) - + # Save state state_path = tmp_path / "scaler_state.pt" torch.save(scaler1.state_dict(), state_path) - + # Create new scaler and load state scaler2 = Scaler(model=model, data=data, nbins=10, verbose=0) scaler2.setup_anisotropy_correction() scaler2.setup_bin_wise_bfactor() scaler2.load_state_dict(torch.load(state_path, weights_only=False)) - + # Parameters should match assert torch.allclose(scaler1.U, scaler2.U) assert torch.allclose(scaler1.bin_wise_bfactor, scaler2.bin_wise_bfactor) @@ -258,18 +274,18 @@ class TestScalerHKLPropertyFunctional: @pytest.mark.integration def test_hkl_property(self, sample_structure_pair): """Test that HKL property returns correct indices.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) - + # HKL from scaler should match data hkl = scaler.hkl assert hkl is not None @@ -283,43 +299,45 @@ class TestScalerDeviceOperationsFunctional: @pytest.mark.integration def test_scaler_cpu_operation(self, sample_structure_pair): """Test scaler works on CPU.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - - scaler = Scaler(model=model, data=data, nbins=10, verbose=0, device=torch.device('cpu')) + + scaler = Scaler( + model=model, data=data, nbins=10, verbose=0, device=torch.device("cpu") + ) scaler.setup_anisotropy_correction() - - assert scaler.device.type == 'cpu' - assert scaler.s.device.type == 'cpu' - assert scaler.U.device.type == 'cpu' + + assert scaler.device.type == "cpu" + assert scaler.s.device.type == "cpu" + assert scaler.U.device.type == "cpu" @pytest.mark.integration def test_scaler_cpu_method(self, sample_structure_pair): """Test scaler.cpu() method.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_anisotropy_correction() scaler.cpu() - + # All tensors should be on CPU for param in scaler.parameters(): - assert param.device.type == 'cpu' + assert param.device.type == "cpu" class TestScalerUMatrixFunctional: @@ -328,23 +346,23 @@ class TestScalerUMatrixFunctional: @pytest.mark.integration def test_u_to_matrix_conversion(self, sample_structure_pair): """Test conversion from U parameters to 3x3 matrix.""" - from torchref.model.model import Model + from torchref.base.math_torch import U_to_matrix from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - from torchref.base.math_torch import U_to_matrix - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_anisotropy_correction() - + # Convert U vector to matrix U_matrix = U_to_matrix(scaler.U) - + # Should be 3x3 assert U_matrix.shape == (3, 3) # Should be symmetric @@ -357,24 +375,24 @@ class TestScalerGradientsFunctional: @pytest.mark.integration def test_anisotropy_gradients(self, sample_structure_pair): """Test gradients flow through anisotropy correction.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_anisotropy_correction() - + # Compute correction and loss correction = scaler.anisotropy_correction() loss = correction.sum() loss.backward() - + # U should have gradients assert scaler.U.grad is not None assert torch.all(torch.isfinite(scaler.U.grad)) @@ -382,56 +400,24 @@ def test_anisotropy_gradients(self, sample_structure_pair): @pytest.mark.integration def test_binwise_bfactor_gradients(self, sample_structure_pair): """Test gradients flow through bin-wise B-factor correction.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_bin_wise_bfactor() - + # Compute correction and loss correction = scaler.bin_wise_bfactor_correction() loss = correction.sum() loss.backward() - + # bin_wise_bfactor should have gradients assert scaler.bin_wise_bfactor.grad is not None assert torch.all(torch.isfinite(scaler.bin_wise_bfactor.grad)) - - -class TestScalerMultipleStructuresFunctional: - """Functional tests with multiple structures.""" - - @pytest.mark.integration - def test_scaler_with_different_structures(self, all_test_structures): - """Test scaler works with different crystal structures.""" - from torchref.scaling.scaler import Scaler - - tested = 0 - for struct in all_test_structures: - pdb_id = struct["pdb_id"] - model = struct["model"] - data = struct["data"] - - scaler = Scaler(model=model, data=data, nbins=10, verbose=0) - scaler.setup_anisotropy_correction() - - # Verify scaler is set up correctly - assert scaler.s is not None - assert scaler.bins is not None - assert scaler.U is not None - - correction = scaler.anisotropy_correction() - assert torch.all(torch.isfinite(correction)) - - tested += 1 - if tested >= 3: # Test first 3 structures - break - - assert tested >= 1, "No test structures with both CIF and MTZ found" diff --git a/tests/helpers/structure_cases.py b/tests/helpers/structure_cases.py new file mode 100644 index 00000000..74a84c28 --- /dev/null +++ b/tests/helpers/structure_cases.py @@ -0,0 +1,27 @@ +"""Name compatibility datasets explicitly so adding a file cannot grow test work silently. + +The quick reader contracts use 1DAW. The broader panel exercises deposited files +across crystal systems and file encodings in the slow tier. Pair-based pipeline +checks use trigonal 2DQ6 and body-centred tetragonal 3A5V in addition to their +separate 1DAW checks. +""" + +MODEL_CODES = ( + "1DAW", # C-centred monoclinic; quick reference structure. + "2DQ6", # Trigonal. + "3A5V", # Body-centred tetragonal. + "3E98", # Monoclinic screw axis. + "3GR5", # Hexagonal screw axis. + "3K7M", # Cubic. + "3VRJ", # Additional monoclinic deposition. + "4BX9", # Tetragonal screw axis. + "5BOV", # Triclinic P1. + "6G9X", # Orthorhombic. +) + +MTZ_CODES = MODEL_CODES + ("1AK5", "1BYW", "1VER", "6JZA", "6SXW", "6VHI") +SF_CIF_CODES = MODEL_CODES + ("7L84",) +EXTENDED_PAIR_CODES = ("2DQ6", "3A5V") +MODEL_CIF_FILES = tuple(f"{code}.cif" for code in MODEL_CODES) + ( + "test_ihm_ensemble.cif", +) diff --git a/tests/integration/test_io_cif.py b/tests/integration/test_io_cif.py index 7e8b23d7..2bd47fa1 100644 --- a/tests/integration/test_io_cif.py +++ b/tests/integration/test_io_cif.py @@ -39,37 +39,6 @@ def test_cif_loading_contract(loaded_model, sample_cif_file) -> None: ) -class TestMultipleCIFFiles: - """Tests that load multiple CIF files.""" - - @pytest.mark.integration - @pytest.mark.slow - def test_load_all_test_structures(self, all_cif_files): - """Test loading all available test structures.""" - from torchref.model.model import Model - - loaded = 0 - errors = [] - - for cif_file in all_cif_files: - try: - model = Model() - model.load_cif(str(cif_file)) - n_atoms = model.xyz().shape[0] - assert n_atoms > 0 - loaded += 1 - except Exception as e: - errors.append((cif_file.name, str(e))) - - # Report - print(f"\nLoaded {loaded}/{len(all_cif_files)} structures") - if errors: - print(f"Errors: {errors}") - - # Should load at least most structures - assert loaded > 0 - - class TestCIFSaving: """Tests for saving CIF files.""" diff --git a/tests/integration/test_structure_compatibility.py b/tests/integration/test_structure_compatibility.py new file mode 100644 index 00000000..2ad00c10 --- /dev/null +++ b/tests/integration/test_structure_compatibility.py @@ -0,0 +1,101 @@ +"""Exercise explicitly named extra structure files in the slow compatibility tier.""" + +import pytest +import torch + +from tests.helpers.structure_cases import ( + EXTENDED_PAIR_CODES, + MODEL_CIF_FILES, + MTZ_CODES, + SF_CIF_CODES, +) +from torchref.config import ( + canonical_device, + get_default_device, + get_float_dtype, + get_int_dtype, +) + +pytestmark = pytest.mark.integration + + +@pytest.mark.parametrize( + "directory, expected", + [ + ("cif", MODEL_CIF_FILES), + ("mtz", tuple(f"{code}.mtz" for code in MTZ_CODES)), + ("cif_sf", tuple(f"{code}-sf.cif" for code in SF_CIF_CODES)), + ], + ids=["models", "mtz", "sf-cif"], +) +def test_compatibility_inventory(test_files_dir, directory, expected) -> None: + """Every bundled input has an explicit quick or extended coverage assignment.""" + suffix = ".mtz" if directory == "mtz" else ".cif" + actual = {path.name for path in (test_files_dir / directory).glob(f"*{suffix}")} + assert actual == set(expected) + + +@pytest.mark.slow +@pytest.mark.parametrize( + "filename", [name for name in MODEL_CIF_FILES if name != "1DAW.cif"] +) +def test_model_cif_compatibility(cif_dir, filename) -> None: + """Each extra CIF loads atoms and finite symmetry operators on the default device.""" + from torchref.model import Model + + path = cif_dir / filename + assert path.is_file() + model = Model(verbose=0).load_cif(str(path)) + xyz = model.xyz() + assert xyz.shape == (len(model.pdb), 3) + assert len(xyz) > 0 + assert xyz.dtype == get_float_dtype() + assert canonical_device(xyz.device) == canonical_device(get_default_device()) + assert torch.isfinite(xyz).all() + assert model.cell.data.shape == (6,) + assert model.spacegroup.matrices.shape[0] > 0 + assert torch.isfinite(model.spacegroup.matrices).all() + + +@pytest.mark.slow +@pytest.mark.parametrize("code", EXTENDED_PAIR_CODES) +def test_modelft_cif_compatibility(cif_dir, code) -> None: + """Fourier models initialize their scattering parametrization in distinct crystals.""" + from torchref.model import ModelFT + + path = cif_dir / f"{code}.cif" + assert path.is_file() + model = ModelFT(max_res=3.0, verbose=0).load_cif(str(path)) + assert len(model.xyz()) > 0 + assert model.parametrization + assert all(size > 0 for size in model.grid_shape) + + +@pytest.mark.slow +@pytest.mark.parametrize( + "directory, filename, loader", + [("mtz", f"{code}.mtz", "load_mtz") for code in MTZ_CODES if code != "1DAW"] + + [ + ("cif_sf", f"{code}-sf.cif", "load_cif") + for code in SF_CIF_CODES + if code != "1DAW" + ], + ids=[f"mtz-{code}" for code in MTZ_CODES if code != "1DAW"] + + [f"sf-cif-{code}" for code in SF_CIF_CODES if code != "1DAW"], +) +def test_reflection_file_compatibility( + test_files_dir, directory, filename, loader +) -> None: + """Every named MTZ/SF-CIF must load reflections; one success cannot mask another failure.""" + from torchref.io import ReflectionData + + path = test_files_dir / directory / filename + assert path.is_file() + data = ReflectionData(verbose=0) + getattr(data, loader)(str(path)) + assert data.hkl.shape == (len(data.hkl), 3) + assert len(data.hkl) > 0 + assert data.hkl.dtype == get_int_dtype() + assert canonical_device(data.hkl.device) == canonical_device(get_default_device()) + assert data.cell.data.shape == (6,) + assert data.F.shape == (len(data.hkl),) diff --git a/tests/integration/test_symmetry_integration.py b/tests/integration/test_symmetry_integration.py index 8bb0a52d..a39ce804 100644 --- a/tests/integration/test_symmetry_integration.py +++ b/tests/integration/test_symmetry_integration.py @@ -6,7 +6,6 @@ import pytest import torch -from pathlib import Path class TestSpaceGroupInitialization: @@ -106,8 +105,8 @@ class TestSpaceGroupDevice: @pytest.mark.integration def test_spacegroup_default_device(self): """Test SpaceGroup matrices land on the configured default device.""" - from torchref.symmetry import SpaceGroup from torchref.config import get_default_device + from torchref.symmetry import SpaceGroup sg = SpaceGroup("P 21 21 21") @@ -140,7 +139,7 @@ def test_expand_coordinates(self, sample_cif_file): # The model should be able to generate symmetry mates # Check if there's an expand method - if hasattr(sg, 'expand') or hasattr(sg, 'expand_atoms'): + if hasattr(sg, "expand") or hasattr(sg, "expand_atoms"): expanded = sg.expand(xyz) assert expanded.shape[0] >= xyz.shape[0] @@ -148,23 +147,6 @@ def test_expand_coordinates(self, sample_cif_file): class TestSpacegroupVariants: """Tests for different spacegroup conventions.""" - @pytest.mark.integration - @pytest.mark.parametrize("sg_name", [ - "P 1", # Triclinic - "P 21", # Monoclinic - "P 21 21 21", # Orthorhombic - "P 43 21 2", # Tetragonal - "P 3 2 1", # Trigonal - "P 6 2 2", # Hexagonal - "P 2 3", # Cubic - ]) - def test_common_spacegroups(self, sg_name): - """Test loading common spacegroups.""" - from torchref.symmetry import SpaceGroup - - sg = SpaceGroup(sg_name) - assert sg.matrices is not None - @pytest.mark.integration def test_spacegroup_name_variations(self): """Test that different spacegroup name formats work.""" @@ -181,25 +163,6 @@ def test_spacegroup_name_variations(self): class TestSpaceGroupWithData: """Tests for SpaceGroup with real crystallographic data.""" - @pytest.mark.integration - def test_spacegroup_with_multiple_structures(self, cif_dir): - """Test SpaceGroup for multiple structures.""" - from torchref.model.model import Model - from torchref.symmetry import SpaceGroup - - cif_files = list(cif_dir.glob("*.cif"))[:3] - - for cif_file in cif_files: - model = Model() - model.load_cif(str(cif_file)) - - sg = SpaceGroup(model.spacegroup) - - # Should have valid matrices - assert sg.matrices is not None - assert sg.matrices.shape[0] >= 1 - assert torch.all(torch.isfinite(sg.matrices)) - @pytest.mark.integration def test_spacegroup_consistent_with_cell(self, sample_cif_file): """Test that SpaceGroup is consistent with unit cell.""" diff --git a/tests/unit/symmetry/test_symmetry.py b/tests/unit/symmetry/test_symmetry.py index aab7bcce..afca8c07 100644 --- a/tests/unit/symmetry/test_symmetry.py +++ b/tests/unit/symmetry/test_symmetry.py @@ -6,7 +6,6 @@ import pytest import torch -import torch.nn as nn class TestSpaceGroupInitialization: @@ -131,7 +130,9 @@ def test_rotation_matrices_determinant(self): for i in range(sg.matrices.shape[0]): det = torch.linalg.det(sg.matrices[i]) - assert torch.isclose(torch.abs(det), torch.tensor(1.0, dtype=det.dtype), atol=1e-5) + assert torch.isclose( + torch.abs(det), torch.tensor(1.0, dtype=det.dtype), atol=1e-5 + ) class TestSpaceGroupApplication: @@ -189,10 +190,10 @@ def test_spacegroup_cpu(self): """Test SpaceGroup on CPU.""" from torchref.symmetry import SpaceGroup - sg = SpaceGroup("P21", device=torch.device('cpu')) + sg = SpaceGroup("P21", device=torch.device("cpu")) - assert sg.matrices.device.type == 'cpu' - assert sg.translations.device.type == 'cpu' + assert sg.matrices.device.type == "cpu" + assert sg.translations.device.type == "cpu" @pytest.mark.unit @pytest.mark.gpu @@ -220,7 +221,23 @@ class TestSpaceGroupMapping: """Tests for space group name mapping.""" @pytest.mark.unit - @pytest.mark.parametrize("sg_name", ["P1", "P21", "P212121", "C2", "P21212"]) + @pytest.mark.parametrize( + "sg_name", + [ + "P1", + "P21", + "P212121", + "C2", + "P21212", + "P 1", + "P 21", + "P 21 21 21", + "P 43 21 2", + "P 3 2 1", + "P 6 2 2", + "P 2 3", + ], + ) def test_common_spacegroups(self, sg_name): """Test common crystallographic space groups.""" from torchref.symmetry import SpaceGroup From cd199b403708fb5ca4e0846d078d2b6ed95ad59e Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 7 Sep 2026 17:23:55 +0200 Subject: [PATCH 165/250] fix: respect configured dtype in loss aggregation --- docs/changelog.rst | 1 + .../test_loss_weighting_functional.py | 8 +-- tests/unit/refinement/test_loss_state.py | 68 +++++++++++++++++-- tests/unit/refinement/test_loss_weighting.py | 3 +- torchref/refinement/loss_state.py | 6 +- 5 files changed, 72 insertions(+), 14 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 3284fb85..83a5728c 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Loss aggregation uses the configured floating-point dtype in eager and compiled execution, including empty and zero-weight aggregates. - Rigid-body refinement stores its Euler angles pre-multiplied by the chain's radius of gyration, so a unit step in an angle and a unit step in a translation displace atoms comparably. In radians against Angstroms the rotation block of the Hessian carried 190-530x the curvature of the translation block on 1DAW and 3E98 -- the geometric ``Rg**2``, 411 and 442/516 -- putting ``cond(H)`` at 1e3-5e3, which is why six parameters needed ~250 L-BFGS iterations to place. Dividing the scale out in ``forward()`` brings the ratio to 0.4-1.3 and ``cond(H)`` to 3-18. Over ten structures the step then converges rather than exhausting its iteration budget, on about half the gradient evaluations, with R-free no worse anywhere. ``RigidXYZTensor.rotation_radians`` returns the physical angle, and setting ``angle_scale`` to ones restores the unscaled parametrization. Not a fix for the one or two negative Hessian eigenvalues at the finer cutoffs -- scaling a saddle leaves it a saddle -- and those counts are unchanged - The rigid-body step no longer co-refines the scaler in the same L-BFGS as the rigid parameters. The body target centres on ``alpha*|F_calc|`` and ``alpha`` absorbs a rescaling of ``F_calc`` exactly, so the scale had a flat direction there; ``SCALE_TARGETS`` already excludes every alpha-centred row from the scale fit for this reason, and 0.6.2 fixed the same thing in the main driver. ``refine_scaler`` (objective ``ls``) owns the scale, between cutoffs - Fixed ``refine_rigid_body`` leaving the caller's reflection data truncated. ``cut_res`` masks in place and returns ``self``, so each cutoff stamped its resolution mask on the caller's own object and the restore had nothing to restore to -- it only looked correct because the default schedule ends at the native limit. With ``--rigid-body-cutoffs 6,4`` on a 2.05 A dataset, 20138 of 23352 reflections stayed masked out for the rest of the run, R-factors included diff --git a/tests/functional/test_loss_weighting_functional.py b/tests/functional/test_loss_weighting_functional.py index f1bcb602..04e4f115 100644 --- a/tests/functional/test_loss_weighting_functional.py +++ b/tests/functional/test_loss_weighting_functional.py @@ -58,7 +58,7 @@ def test_total_weighted_loss_from_state(self): total = state.aggregate() # Expected: 10*1.0 + 5*0.5 + 2*0.25 = 10 + 2.5 + 0.5 = 13.0 - assert torch.isclose(total, torch.tensor(13.0)) + assert total.item() == pytest.approx(13.0) @pytest.mark.integration @@ -146,7 +146,7 @@ def test_zero_weight(self): # Zero weight should effectively disable ADP term total = state.aggregate() - assert torch.isclose(total, torch.tensor(0.0)) + assert total.item() == pytest.approx(0.0) @pytest.mark.integration @@ -166,8 +166,7 @@ def test_aggregator_basic(self): total = state.aggregate() # Expected: 2.0 * 1.0 + 1.0 * 0.5 = 2.5 - expected = torch.tensor(2.5) - assert torch.isclose(total, expected) + assert total.item() == pytest.approx(2.5) def test_loss_state_caches_losses(self): """Test that LossState caches computed losses.""" @@ -195,4 +194,3 @@ def counting_target(): assert cached is not None assert torch.isclose(cached, torch.tensor(2.0)) - diff --git a/tests/unit/refinement/test_loss_state.py b/tests/unit/refinement/test_loss_state.py index c6126464..553ff43e 100644 --- a/tests/unit/refinement/test_loss_state.py +++ b/tests/unit/refinement/test_loss_state.py @@ -8,6 +8,66 @@ import torch +@pytest.mark.unit +@pytest.mark.parametrize("mode", ["empty", "disabled", "disabled_compilable"]) +@pytest.mark.parametrize("log_values", [False, True]) +def test_zero_aggregate_uses_configured_dtype_and_device( + mode: str, log_values: bool +) -> None: + """An aggregate without active targets is a configured scalar zero.""" + from torchref.config import get_default_device, get_float_dtype + from torchref.refinement.loss_state import LossState + + state = LossState() + if mode != "empty": + + def disabled_target(): + pytest.fail("A zero-weight target must not be evaluated") + + state.register_target( + "geometry/bond", + disabled_target, + compile=mode == "disabled_compilable", + probe=False, + ) + state.set_weight("geometry", 0.0) + state.compile_aggregate() + + total = state.aggregate(log_values=log_values) + + expected = torch.zeros((), dtype=get_float_dtype(), device=get_default_device()) + torch.testing.assert_close(total, expected) + assert state._losses == {} + assert state.history == ([{"total": 0.0}] if log_values else []) + + +@pytest.mark.unit +@pytest.mark.parametrize("compiled", [False, True]) +def test_aggregate_ignores_torch_default_dtype( + monkeypatch: pytest.MonkeyPatch, compiled: bool +) -> None: + """Eager and compiled sums use TorchRef's dtype, not PyTorch's default.""" + from torchref.config import device, dtypes, get_float_dtype + from torchref.refinement.loss_state import LossState + + monkeypatch.setattr(device, "current", torch.device("cpu")) + monkeypatch.setattr(dtypes, "float", torch.float32) + state = LossState() + value = torch.tensor(2.0, dtype=get_float_dtype(), device=state.device) + state.register_target("geometry/bond", lambda: value, compile=compiled) + state.set_weight("geometry", 3.0) + previous_dtype = torch.get_default_dtype() + try: + torch.set_default_dtype(torch.float64) + if compiled: + state.compile_aggregate(backend="eager") + total = state.aggregate() + finally: + torch.set_default_dtype(previous_dtype) + + torch.testing.assert_close(total, value * 3.0) + + class TestLossStateBasic: """Tests for basic LossState functionality.""" @@ -121,7 +181,7 @@ def test_prefix_with_hierarchical_weighting(self): total = state.aggregate() # Expected: 0.5 * 1.0 + 1.0 * 2.0 = 2.5 - assert torch.isclose(total, torch.tensor(2.5)) + assert total.item() == pytest.approx(2.5) class TestWeightManagement: @@ -221,7 +281,7 @@ def test_aggregate_simple(self): total = state.aggregate(log_values=False) # 2.0 * 1.0 + 1.0 * 0.5 = 2.5 - assert torch.isclose(total, torch.tensor(2.5)) + assert total.item() == pytest.approx(2.5) @pytest.mark.unit def test_aggregate_hierarchical(self): @@ -240,7 +300,7 @@ def test_aggregate_hierarchical(self): # geometry/bond: 1.0 * 0.5 * 2.0 = 1.0 # geometry/angle: 2.0 * 0.5 * 1.0 = 1.0 # total = 2.0 - assert torch.isclose(total, torch.tensor(2.0)) + assert total.item() == pytest.approx(2.0) @pytest.mark.unit def test_aggregate_default_weights(self): @@ -255,7 +315,7 @@ def test_aggregate_default_weights(self): total = state.aggregate(log_values=False) # 2.0 * 1.0 + 1.0 * 1.0 = 3.0 - assert torch.isclose(total, torch.tensor(3.0)) + assert total.item() == pytest.approx(3.0) @pytest.mark.unit def test_aggregate_caches_losses(self): diff --git a/tests/unit/refinement/test_loss_weighting.py b/tests/unit/refinement/test_loss_weighting.py index 08a7fe0c..ec581ec3 100644 --- a/tests/unit/refinement/test_loss_weighting.py +++ b/tests/unit/refinement/test_loss_weighting.py @@ -48,8 +48,7 @@ def test_total_weighted_loss(self): total = state.aggregate(log_values=False) # Expected: 2.0 * 1.0 + 1.0 * 0.5 = 2.5 - expected = torch.tensor(2.5) - assert torch.isclose(total, expected) + assert total.item() == pytest.approx(2.5) class TestDefaultGroupWeights: diff --git a/torchref/refinement/loss_state.py b/torchref/refinement/loss_state.py index 9df09fd0..dee6b00f 100644 --- a/torchref/refinement/loss_state.py +++ b/torchref/refinement/loss_state.py @@ -21,7 +21,7 @@ import torch from torch import nn -from torchref.config import canonical_device, get_default_device +from torchref.config import canonical_device, get_default_device, get_float_dtype from torchref.utils.autograd_introspection import collect_loss_leaves, _iter_roots from torchref.utils.device_mixin import DeviceMovementMixin from torchref.utils.loss_validation import validate_loss @@ -349,7 +349,7 @@ def compile_aggregate(self, **compile_kwargs) -> "LossState": device = self.device def _compiled_fn(): - total = torch.tensor(0.0, device=device) + total = torch.tensor(0.0, dtype=get_float_dtype(), device=device) for fn, w in zip(fns, weights): total = total + w * fn() return total @@ -407,7 +407,7 @@ def aggregate(self, log_values: bool = False) -> torch.Tensor: self.new_entry() self._losses.clear() - total = torch.tensor(0.0, device=self.device) + total = torch.tensor(0.0, dtype=get_float_dtype(), device=self.device) # --- compiled group --- # Skipped when log_values=True: the fused closure does not expose From d6dbb6d50f8ac241a65558d25e705e97a13b6a1e Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Tue, 8 Sep 2026 09:13:00 +0200 Subject: [PATCH 166/250] ci: run CPU and MPS tests for pull requests into dev --- .github/workflows/compatibility.yml | 1 + .github/workflows/dev-pr.yml | 47 +++++++++++++++++++++++++++++ docs/changelog.rst | 1 + 3 files changed, 49 insertions(+) create mode 100644 .github/workflows/dev-pr.yml diff --git a/.github/workflows/compatibility.yml b/.github/workflows/compatibility.yml index d6437c7a..7c46cccc 100644 --- a/.github/workflows/compatibility.yml +++ b/.github/workflows/compatibility.yml @@ -21,6 +21,7 @@ on: # Also run on PRs that modify dependencies pull_request: + branches-ignore: [dev] paths: - 'pyproject.toml' - 'tox.ini' diff --git a/.github/workflows/dev-pr.yml b/.github/workflows/dev-pr.yml new file mode 100644 index 00000000..33131672 --- /dev/null +++ b/.github/workflows/dev-pr.yml @@ -0,0 +1,47 @@ +name: Dev PR Tests + +on: + pull_request: + branches: [dev] + +permissions: + contents: read + +concurrency: + group: dev-pr-tests-${{ github.event.pull_request.number }} + cancel-in-progress: true + +jobs: + cpu: + name: CPU (Python 3.12) + runs-on: ubuntu-latest + timeout-minutes: 90 + env: + TORCHREF_DEVICE: cpu + NUMBA_CACHE_DIR: /tmp/numba_cache + + steps: + - uses: actions/checkout@v7 + + - name: Set up Python + uses: actions/setup-python@v7 + with: + python-version: '3.12' + cache: pip + cache-dependency-path: pyproject.toml + + - name: Install dependencies + run: | + python -m pip install --upgrade pip + python -m pip install torch --index-url https://download.pytorch.org/whl/cpu + python -m pip install -e ".[dev]" + + - name: Run tests on CPU + run: | + python -m pytest tests/ \ + -m "not gpu and not slow" \ + -v --tb=short -rf --durations=20 + + mps: + name: MPS + uses: ./.github/workflows/accelerator.yml diff --git a/docs/changelog.rst b/docs/changelog.rst index 83a5728c..16acd423 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Pull requests into ``dev`` run one Python 3.12 CPU test job and one MPS test job, without the dependency-compatibility matrix. - Loss aggregation uses the configured floating-point dtype in eager and compiled execution, including empty and zero-weight aggregates. - Rigid-body refinement stores its Euler angles pre-multiplied by the chain's radius of gyration, so a unit step in an angle and a unit step in a translation displace atoms comparably. In radians against Angstroms the rotation block of the Hessian carried 190-530x the curvature of the translation block on 1DAW and 3E98 -- the geometric ``Rg**2``, 411 and 442/516 -- putting ``cond(H)`` at 1e3-5e3, which is why six parameters needed ~250 L-BFGS iterations to place. Dividing the scale out in ``forward()`` brings the ratio to 0.4-1.3 and ``cond(H)`` to 3-18. Over ten structures the step then converges rather than exhausting its iteration budget, on about half the gradient evaluations, with R-free no worse anywhere. ``RigidXYZTensor.rotation_radians`` returns the physical angle, and setting ``angle_scale`` to ones restores the unscaled parametrization. Not a fix for the one or two negative Hessian eigenvalues at the finer cutoffs -- scaling a saddle leaves it a saddle -- and those counts are unchanged - The rigid-body step no longer co-refines the scaler in the same L-BFGS as the rigid parameters. The body target centres on ``alpha*|F_calc|`` and ``alpha`` absorbs a rescaling of ``F_calc`` exactly, so the scale had a flat direction there; ``SCALE_TARGETS`` already excludes every alpha-centred row from the scale fit for this reason, and 0.6.2 fixed the same thing in the main driver. ``refine_scaler`` (objective ``ls``) owns the scale, between cutoffs From 4842557c2bf29050202a66c18bfd318d5e82eb6c Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Wed, 9 Sep 2026 11:54:58 +0200 Subject: [PATCH 167/250] feat: make hydrogen generation opt-in with refine CLI flag --- docs/changelog.rst | 1 + docs/user_guide/cli.rst | 2 + tests/integration/test_cli_hydrogens.py | 66 ++++++++++++++++++++++ tests/unit/model/test_hydrogen_default.py | 68 ++++++++++++++++++++--- tests/unit/model/test_model.py | 4 +- torchref/cli/refine.py | 8 +++ torchref/model/context.py | 4 +- torchref/model/model.py | 18 +++--- torchref/refinement/base_refinement.py | 6 ++ 9 files changed, 156 insertions(+), 21 deletions(-) create mode 100644 tests/integration/test_cli_hydrogens.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 9e1e17d5..b4530a45 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Hydrogen generation on model loading is off by default; use ``torchref.refine --add-hydrogens`` or ``add_hydrogens=True`` in Python to opt in. Hydrogens already present in input files are retained unless ``strip_H=True``. - Pull requests into ``dev`` run one Python 3.12 CPU test job and one MPS test job, without the dependency-compatibility matrix. - Loss aggregation uses the configured floating-point dtype in eager and compiled execution, including empty and zero-weight aggregates. - Named broad structure-compatibility cases explicitly, moved extra datasets to the slow tier, and removed eager all-model loading and swallowed reader failures. diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index a649e62f..0f96ec82 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -29,6 +29,8 @@ and a ``refinement_history.json`` log. **Key options:** * ``-n`` / ``--n-cycles`` number of macro cycles (default 5) +* ``--add-hydrogens`` generate missing hydrogens on model loading (default off). + Hydrogens already present in the input are retained with or without this flag * ``--mode`` ``separate`` (separated XYZ then ADP, default) or ``everything`` (joint XYZ+ADP) * ``--xray-mode`` one of ``ml`` (default; Read MLF at variance ε·β, conditional diff --git a/tests/integration/test_cli_hydrogens.py b/tests/integration/test_cli_hydrogens.py new file mode 100644 index 00000000..f82a7357 --- /dev/null +++ b/tests/integration/test_cli_hydrogens.py @@ -0,0 +1,66 @@ +"""Exercise hydrogen opt-in from CLI parsing through deposited-model loading.""" + +import sys +from pathlib import Path + +import pytest + +from torchref.refinement.base_refinement import Refinement +from torchref.refinement.lbfgs_refinement import LBFGSRefinement + + +class _ModelLoaded(Exception): + """Stop after real model loading, before scaling and optimization.""" + + +@pytest.mark.integration +@pytest.mark.parametrize("add_hydrogens", [False, True]) +@pytest.mark.parametrize("model_format", ["pdb", "cif"]) +def test_cli_hydrogen_generation_is_opt_in( + test_files_dir: Path, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + add_hydrogens: bool, + model_format: str, +) -> None: + """The CLI flag generates missing hydrogens for PDB and mmCIF inputs.""" + from torchref.cli import refine + + loaded = [] + + def stop_after_loading(refinement: Refinement) -> None: + loaded.append(refinement.model) + raise _ModelLoaded + + monkeypatch.setattr(Refinement, "_sync_model_cell_to_data", stop_after_loading) + argv = [ + "torchref.refine", + "-m", + str(test_files_dir / model_format / f"1DAW.{model_format}"), + "-sf", + str(test_files_dir / "mtz" / "1DAW.mtz"), + "-o", + str(tmp_path / "refined"), + "-v", + "0", + ] + if add_hydrogens: + argv.append("--add-hydrogens") + monkeypatch.setattr(sys, "argv", argv) + + with pytest.raises(_ModelLoaded): + refine.main() + + (model,) = loaded + assert model.ctx.add_hydrogens is add_hydrogens + assert len(model.pdb) > 0 + n_hydrogens = int(model.pdb["element"].str.strip().eq("H").sum()) + assert (n_hydrogens > 0) is add_hydrogens + + +@pytest.mark.unit +@pytest.mark.parametrize("add_hydrogens", [False, True]) +def test_empty_refinement_preserves_hydrogen_setting(add_hydrogens: bool) -> None: + """An empty refinement shell forwards the setting to its model too.""" + refinement = LBFGSRefinement(verbose=0, add_hydrogens=add_hydrogens) + assert refinement.model.ctx.add_hydrogens is add_hydrogens diff --git a/tests/unit/model/test_hydrogen_default.py b/tests/unit/model/test_hydrogen_default.py index 678cb9a1..3c6f1cf3 100644 --- a/tests/unit/model/test_hydrogen_default.py +++ b/tests/unit/model/test_hydrogen_default.py @@ -1,14 +1,17 @@ -"""Hydrogens are present by default: kept where the file has them, generated where not. +"""Keep deposited hydrogens by default and generate missing ones only on request. The interesting cases are the partially-hydrogenated file, which has to be topped up per parent rather than left alone, and the per-atom buffers that are cached lazily and go stale the moment the atom set grows. """ +from pathlib import Path + import numpy as np import pytest from torchref.model.model import Model +from torchref.model.model_ft import ModelFT def _elements(model): @@ -21,10 +24,58 @@ def _counts(model): return len(model.pdb), n_h +@pytest.mark.unit +@pytest.mark.parametrize("model_class", [Model, ModelFT]) +@pytest.mark.parametrize("filename", ["1DAW.pdb", "1AK5_with_H.pdb", "7L84.pdb"]) +def test_default_preserves_deposited_atoms( + pdb_dir: Path, + filename: str, + monkeypatch: pytest.MonkeyPatch, + model_class: type[Model], +) -> None: + """Default loading neither generates hydrogens nor removes deposited ones.""" + from torchref.io.pdb import PDBReader + + path = pdb_dir / filename + deposited, _, _ = PDBReader(verbose=0).read(str(path))() + + def unexpected_generation(self: Model) -> None: + pytest.fail("Default loading must not generate hydrogens") + + monkeypatch.setattr(Model, "_add_missing_hydrogens", unexpected_generation) + model = model_class(verbose=0).load_pdb(str(path)) + + np.testing.assert_array_equal(_elements(model), deposited["element"].str.strip()) + + +@pytest.mark.unit +def test_default_cif_load_does_not_generate_hydrogens( + cif_dir: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """mmCIF loading also leaves missing hydrogens absent by default.""" + + def unexpected_generation(self: Model) -> None: + pytest.fail("Default mmCIF loading must not generate hydrogens") + + monkeypatch.setattr(Model, "_add_missing_hydrogens", unexpected_generation) + model = Model(verbose=0).load_cif(str(cif_dir / "1DAW.cif")) + total, n_h = _counts(model) + assert total > 0 + assert n_h == 0 + + +@pytest.mark.unit +def test_context_defaults_to_no_hydrogen_generation() -> None: + """A standalone model context leaves hydrogen generation disabled.""" + from torchref.model.context import ModelContext + + assert ModelContext().add_hydrogens is False + + @pytest.mark.unit def test_a_file_without_hydrogens_gets_them(pdb_dir): """1DAW ships none, so every hydrogen here is generated.""" - model = Model(verbose=0) + model = Model(verbose=0, add_hydrogens=True) model.load_pdb(str(pdb_dir / "1DAW.pdb")) total, n_h = _counts(model) @@ -47,7 +98,7 @@ def test_a_partially_hydrogenated_file_is_topped_up(pdb_dir): kept.load_pdb(str(pdb_dir / "1AK5_with_H.pdb")) _, n_kept = _counts(kept) - topped = Model(verbose=0) + topped = Model(verbose=0, add_hydrogens=True) topped.load_pdb(str(pdb_dir / "1AK5_with_H.pdb")) _, n_topped = _counts(topped) @@ -61,7 +112,7 @@ def test_a_partially_hydrogenated_file_is_topped_up(pdb_dir): def test_strip_H_still_removes_everything(pdb_dir): """The opt-out is unaffected: no hydrogen survives, generated or deposited.""" for name in ("1DAW.pdb", "7L84.pdb"): - model = Model(verbose=0, strip_H=True) + model = Model(verbose=0, strip_H=True, add_hydrogens=True) model.load_pdb(str(pdb_dir / name)) _, n_h = _counts(model) assert n_h == 0, f"{name} kept {n_h} hydrogens under strip_H" @@ -75,7 +126,7 @@ def test_add_hydrogens_false_keeps_the_file_as_it_is(pdb_dir): total, n_h = _counts(model) assert n_h > 0, "7L84 ships hydrogens, so they should have been kept" - generated = Model(verbose=0) + generated = Model(verbose=0, add_hydrogens=True) generated.load_pdb(str(pdb_dir / "7L84.pdb")) assert _counts(generated)[0] >= total @@ -89,7 +140,7 @@ def test_per_atom_buffers_are_rebuilt_for_the_new_atom_set(pdb_dir): place left the van der Waals radii at the heavy-atom count while the pair list indexed the full set, and the non-bonded build raised ``IndexError``. """ - model = Model(verbose=0) + model = Model(verbose=0, add_hydrogens=True) model.load_pdb(str(pdb_dir / "1DAW.pdb")) n_atoms = len(model.pdb) @@ -106,12 +157,13 @@ def test_per_atom_buffers_are_rebuilt_for_the_new_atom_set(pdb_dir): @pytest.mark.unit def test_restraints_build_over_the_hydrogenated_model(pdb_dir): """Restraints cover the hydrogens, and each carries exactly one bond.""" - model = Model(verbose=0) + model = Model(verbose=0, add_hydrogens=True) model.load_pdb(str(pdb_dir / "1DAW.pdb")) restraints = model.restraints elements = _elements(model) is_h = elements == "H" + assert is_h.any() bonds = restraints.restraints["bond"]["all"]["indices"].cpu().numpy() involves_h = is_h[bonds[:, 0]] | is_h[bonds[:, 1]] assert int(involves_h.sum()) == int(is_h.sum()) @@ -132,7 +184,7 @@ def test_riding_hydrogens_are_not_placed_when_real_ones_exist(pdb_dir): because the riding builder counts bonded neighbours by distance while the generator reads them off the bond graph. """ - model = Model(verbose=0) + model = Model(verbose=0, add_hydrogens=True) model.load_pdb(str(pdb_dir / "1DAW.pdb")) restraints = model.restraints diff --git a/tests/unit/model/test_model.py b/tests/unit/model/test_model.py index 7b9045fa..da9a28bb 100644 --- a/tests/unit/model/test_model.py +++ b/tests/unit/model/test_model.py @@ -54,13 +54,13 @@ def test_model_custom_dtype(self): @pytest.mark.unit def test_model_strip_h_default(self): - """strip_H defaults to False, and hydrogen generation is on.""" + """Hydrogen stripping and generation are both opt-in.""" from torchref.model.model import Model model = Model() assert model.ctx.strip_H is False - assert model.ctx.add_hydrogens is True + assert model.ctx.add_hydrogens is False @pytest.mark.unit def test_model_bool_uninitialized(self): diff --git a/torchref/cli/refine.py b/torchref/cli/refine.py index c8c837ad..bec7fadc 100644 --- a/torchref/cli/refine.py +++ b/torchref/cli/refine.py @@ -114,6 +114,12 @@ def main(): refine_group = parser.add_argument_group("Refinement") add_n_cycles_arg(refine_group) + refine_group.add_argument( + "--add-hydrogens", + action="store_true", + help="Generate missing hydrogens when loading the model (default: off). " + "Hydrogens already present in the input are retained either way.", + ) refine_group.add_argument( "--mode", type=str, @@ -248,6 +254,7 @@ def main(): print(f"Refinement mode: {args.mode}") print(f"X-ray target: {args.xray_mode}") print(f"Refinement cycles: {args.n_cycles}") + print(f"Add hydrogens: {'on' if args.add_hydrogens else 'off'}") if args.with_rigid_body: print(f"Rigid-body step: on (iterations/cutoff = {args.rigid_body_iter})") print(f"Device: {args.device}") @@ -304,6 +311,7 @@ def main(): reflections_per_adp_parameter=args.reflections_per_adp_parameter, aniso_selection=args.anisotropic_selection, wavelength=args.wavelength, + add_hydrogens=args.add_hydrogens, ) # Merge onto DEFAULT_GROUP_WEIGHTS so unspecified groups keep their defaults; diff --git a/torchref/model/context.py b/torchref/model/context.py index 17f32148..06987680 100644 --- a/torchref/model/context.py +++ b/torchref/model/context.py @@ -53,7 +53,7 @@ class ModelContext(DeviceMixin): Whether hydrogens were stripped on load. exclude_H_from_sf : bool, default False Whether hydrogens are excluded from structure-factor calculation. - add_hydrogens : bool, default True + add_hydrogens : bool, default False Generate hydrogens on load for residues that arrive without them. Ignored when ``strip_H`` is set, which removes them again. initialized : bool, default False @@ -80,7 +80,7 @@ class ModelContext(DeviceMixin): verbose: int = 1 strip_H: bool = True exclude_H_from_sf: bool = False - add_hydrogens: bool = True + add_hydrogens: bool = False initialized: bool = False def copy(self) -> "ModelContext": diff --git a/torchref/model/model.py b/torchref/model/model.py index f57ff750..ef207dc1 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -92,9 +92,9 @@ class Model(DeviceMovementMixin, DebugMixin, nn.Module): Computation device. Defaults to the configured device.current. strip_H : bool, optional Whether to strip hydrogen atoms when loading. Default False: hydrogens are kept - where the file has them and generated where it does not. + where the file has them. add_hydrogens : bool, optional - Generate hydrogens on load for residues that arrive without them. Default True; + Generate missing hydrogens on load when True. Default False; ignored when ``strip_H`` is set. Attributes @@ -130,7 +130,7 @@ def __init__( verbose=1, device=None, strip_H: bool = False, - add_hydrogens: bool = True, + add_hydrogens: bool = False, ): """ Initialize an empty Model shell. @@ -148,10 +148,10 @@ def __init__( Computation device. Defaults to the configured device.current. strip_H : bool, optional Whether to strip hydrogen atoms when loading. Default False: hydrogens are - kept where the file has them and generated where it does not. + kept where the file has them. add_hydrogens : bool, optional - Generate hydrogens on load for residues that arrive without them. Default - True; ignored when ``strip_H`` is set. + Generate missing hydrogens on load when True. Default False; + ignored when ``strip_H`` is set. """ super().__init__() # Resolve dtype/device at call time (not import time) so a runtime @@ -738,8 +738,8 @@ def _add_missing_hydrogens(self) -> None: fixed point. Costs a restraint build that is then discarded, because the plan needs the - topology and the topology is built over the atoms as loaded. Set - ``add_hydrogens=False`` to skip it for a model that will never be refined. + topology and the topology is built over the atoms as loaded. Loading invokes + this only when ``add_hydrogens=True`` is requested. """ from torchref.topology.hydrogens import ( augment_atom_table, @@ -2009,7 +2009,7 @@ def shake_adp(self, stddev: float): def _new_model_from_df(self, df, *, strip_H=None, add_hydrogens=False): """Build a fresh model of the same class from a DataFrame. - ``add_hydrogens`` defaults to False, unlike the constructor: the caller has + ``add_hydrogens`` defaults to False: the caller has already settled which atoms the table holds, and generating more would fight that. :meth:`hydrogenate` passes an already-augmented table for the same reason. """ diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index 88444aec..9dedd761 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -137,6 +137,7 @@ def __init__( shrink: bool = SHRINK_ENABLED, scale_target: str = DEFAULT_SCALE_TARGET, aniso_selection: Optional[str] = None, + add_hydrogens: bool = False, ): """Initialize Refinement, fully if ``data_file`` and ``pdb`` are given. @@ -204,6 +205,9 @@ def __init__( aniso_selection : str, optional Phenix-style selection of atoms refined anisotropically when ``adp_mode="anisotropic"``. Defaults to all non-water heavy atoms. + add_hydrogens : bool, optional + Generate missing hydrogens when loading the model. Default False. + Hydrogens already present in the input are retained either way. """ super().__init__() # Refinement constructs its own submodules from file paths, so @@ -274,6 +278,7 @@ def __init__( device=self.device, wavelength=self.wavelength, anomalous_threshold=self.anomalous_threshold, + add_hydrogens=add_hydrogens, ) self.scaler = Scaler( verbose=self.verbose, device=self.device, nbins=self.nbins, @@ -320,6 +325,7 @@ def __init__( device=self.device, wavelength=self.wavelength, anomalous_threshold=self.anomalous_threshold, + add_hydrogens=add_hydrogens, # Apply the f'' (Bijvoet) term only when the data were loaded as # explicit Friedel pairs; merged data gate it off. apply_bijvoet=not self.reflection_data.friedel_merged, From b6ffbfa01d148e8c2922d31d1dfe29bd0200e63f Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Fri, 11 Sep 2026 09:50:10 +0200 Subject: [PATCH 168/250] Fixed hydrogen placement deferring to the monomer library --- docs/changelog.rst | 4 + tests/files/restraints/GLU_ASP_renamed.cif | 406 ++++++++++++++++++ tests/files/restraints/GLU_renamed.cif | 217 ++++++++++ tests/integration/test_cli_hydrogens.py | 41 ++ tests/unit/io/test_struct_conn_links.py | 74 ++++ tests/unit/model/test_hydrogen_default.py | 106 +++++ tests/unit/topology/test_equivalence.py | 4 +- tests/unit/topology/test_hydrogens.py | 143 +++++- tests/unit/topology/test_links.py | 110 +++++ torchref/cli/_common.py | 18 +- .../experimental/ensemble/ensemble_model.py | 2 + torchref/experimental/kinetic/refinement.py | 4 +- torchref/io/cif_readers.py | 82 +++- torchref/io/pdb.py | 18 +- torchref/model/model.py | 33 +- torchref/refinement/base_refinement.py | 11 +- torchref/topology/atom_graph.py | 14 +- torchref/topology/build.py | 4 +- torchref/topology/hydrogens.py | 53 ++- torchref/topology/restraints.py | 2 +- 20 files changed, 1284 insertions(+), 62 deletions(-) create mode 100644 tests/files/restraints/GLU_ASP_renamed.cif create mode 100644 tests/files/restraints/GLU_renamed.cif create mode 100644 tests/unit/io/test_struct_conn_links.py create mode 100644 tests/unit/topology/test_links.py diff --git a/docs/changelog.rst b/docs/changelog.rst index b4530a45..f05a5fe2 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -5,6 +5,10 @@ Changelog Unreleased ---------- - Hydrogen generation on model loading is off by default; use ``torchref.refine --add-hydrogens`` or ``add_hydrogens=True`` in Python to opt in. Hydrogens already present in input files are retained unless ``strip_H=True``. +- Hydrogen generation reads the user's restraint CIF: ``cif_path`` is a model constructor argument, set before loading by ``torchref.refine --cif`` and the shared CLI loader, and carried by ``hydrogenate``, ``strip_hydrogens``, ``select`` and state dicts. +- mmCIF models carry their covalent and metal ``_struct_conn`` links, as PDB LINK records already did. +- A LINK record repeated in a file, or a bond emitted once per altloc conformer, now counts once in the atom graph. +- Hydrogen count per atom is capped by the template's own hydrogen count minus extra covalent partners, so linked hetero atoms (acetyl caps, Schiff bases, glycosylated ASN, metal-bound HIS) no longer receive displaced hydrogens; shared backbone atoms of split residues keep their HA/H. - Pull requests into ``dev`` run one Python 3.12 CPU test job and one MPS test job, without the dependency-compatibility matrix. - Loss aggregation uses the configured floating-point dtype in eager and compiled execution, including empty and zero-weight aggregates. - Named broad structure-compatibility cases explicitly, moved extra datasets to the slow tier, and removed eager all-model loading and swallowed reader failures. diff --git a/tests/files/restraints/GLU_ASP_renamed.cif b/tests/files/restraints/GLU_ASP_renamed.cif new file mode 100644 index 00000000..0140a35e --- /dev/null +++ b/tests/files/restraints/GLU_ASP_renamed.cif @@ -0,0 +1,406 @@ +data_comp_list +loop_ +_chem_comp.id +_chem_comp.three_letter_code +_chem_comp.name +_chem_comp.group +_chem_comp.number_atoms_all +_chem_comp.number_atoms_nh +_chem_comp.desc_level +GLU GLU "GLUTAMIC ACID" peptide 18 10 . +ASP ASP "ASPARTIC ACID" peptide 15 9 . + +data_comp_GLU +loop_ +_chem_comp_atom.comp_id +_chem_comp_atom.atom_id +_chem_comp_atom.type_symbol +_chem_comp_atom.type_energy +_chem_comp_atom.charge +_chem_comp_atom.x +_chem_comp_atom.y +_chem_comp_atom.z +GLU N N NT3 1 88.319 -7.751 -10.089 +GLU CA C CH1 0 87.677 -7.162 -11.296 +GLU C C C 0 88.359 -5.826 -11.640 +GLU O O O 0 88.389 -4.954 -10.744 +GLU CB C CH2 0 86.177 -6.950 -11.080 +GLU CG C CH2 0 85.390 -8.247 -10.908 +GLU CD C C 0 83.891 -8.048 -10.773 +GLU OE1 O O 0 83.281 -7.495 -11.711 +GLU OE2 O OC -1 83.334 -8.447 -9.728 +GLU OXT O OC -1 88.836 -5.708 -12.790 +GLU H H H 0 88.066 -7.298 -9.352 +GLU H2 H H 0 89.218 -7.717 -10.158 +GLU H3 H H 0 88.077 -8.615 -9.999 +GLU HAX H H 0 87.806 -7.788 -12.054 +GLU HBY H H 0 85.814 -6.460 -11.847 +GLU HBX H H 0 86.049 -6.394 -10.284 +GLU HGY H H 0 85.714 -8.714 -10.110 +GLU HGX H H 0 85.559 -8.827 -11.681 + +loop_ +_chem_comp_tree.comp_id +_chem_comp_tree.atom_id +_chem_comp_tree.atom_back +_chem_comp_tree.atom_forward +_chem_comp_tree.connect_type +GLU N n/a CA START +GLU H N . . +GLU H2 N . . +GLU H3 N . . +GLU CA N C . +GLU HAX CA . . +GLU CB CA CG . +GLU HBY CB . . +GLU HBX CB . . +GLU CG CB CD . +GLU HGY CG . . +GLU HGX CG . . +GLU CD CG OE2 . +GLU OE1 CD . . +GLU OE2 CD . . +GLU C CA . END +GLU O C . . +GLU OXT C . . + +loop_ +_chem_comp_acedrg.comp_id +_chem_comp_acedrg.atom_id +_chem_comp_acedrg.atom_type +GLU N N(CCCH)(H)3 +GLU CA C(CCHH)(NH3)(COO)(H) +GLU C C(CCHN)(O)2 +GLU O O(CCO) +GLU CB C(CCHH)(CCHN)(H)2 +GLU CG C(CCHH)(COO)(H)2 +GLU CD C(CCHH)(O)2 +GLU OE1 O(CCO) +GLU OE2 O(CCO) +GLU OXT O(CCO) +GLU H H(NCHH) +GLU H2 H(NCHH) +GLU H3 H(NCHH) +GLU HAX H(CCCN) +GLU HBY H(CCCH) +GLU HBX H(CCCH) +GLU HGY H(CCCH) +GLU HGX H(CCCH) + +loop_ +_chem_comp_bond.comp_id +_chem_comp_bond.atom_id_1 +_chem_comp_bond.atom_id_2 +_chem_comp_bond.type +_chem_comp_bond.aromatic +_chem_comp_bond.value_dist_nucleus +_chem_comp_bond.value_dist_nucleus_esd +_chem_comp_bond.value_dist +_chem_comp_bond.value_dist_esd +GLU N CA SINGLE n 1.487 0.0100 1.487 0.0100 +GLU CA C SINGLE n 1.538 0.0113 1.538 0.0113 +GLU CA CB SINGLE n 1.529 0.0100 1.529 0.0100 +GLU C O DOUBLE n 1.251 0.0183 1.251 0.0183 +GLU C OXT SINGLE n 1.251 0.0183 1.251 0.0183 +GLU CB CG SINGLE n 1.526 0.0100 1.526 0.0100 +GLU CG CD SINGLE n 1.518 0.0135 1.518 0.0135 +GLU CD OE1 DOUBLE n 1.249 0.0161 1.249 0.0161 +GLU CD OE2 SINGLE n 1.249 0.0161 1.249 0.0161 +GLU N H SINGLE n 1.018 0.0520 0.902 0.0102 +GLU N H2 SINGLE n 1.018 0.0520 0.902 0.0102 +GLU N H3 SINGLE n 1.018 0.0520 0.902 0.0102 +GLU CA HAX SINGLE n 1.092 0.0100 0.991 0.0200 +GLU CB HBY SINGLE n 1.092 0.0100 0.980 0.0168 +GLU CB HBX SINGLE n 1.092 0.0100 0.980 0.0168 +GLU CG HGY SINGLE n 1.092 0.0100 0.981 0.0172 +GLU CG HGX SINGLE n 1.092 0.0100 0.981 0.0172 + +loop_ +_chem_comp_angle.comp_id +_chem_comp_angle.atom_id_1 +_chem_comp_angle.atom_id_2 +_chem_comp_angle.atom_id_3 +_chem_comp_angle.value_angle +_chem_comp_angle.value_angle_esd +GLU CA N H 109.990 3.00 +GLU CA N H2 109.990 3.00 +GLU CA N H3 109.990 3.00 +GLU H N H2 109.032 3.00 +GLU H N H3 109.032 3.00 +GLU H2 N H3 109.032 3.00 +GLU N CA C 109.258 1.50 +GLU N CA CB 110.440 2.46 +GLU N CA HAX 108.387 1.58 +GLU C CA CB 111.059 3.00 +GLU C CA HAX 108.774 1.79 +GLU CB CA HAX 109.080 2.33 +GLU CA C O 117.148 1.60 +GLU CA C OXT 117.148 1.60 +GLU O C OXT 125.704 1.50 +GLU CA CB CG 113.294 1.61 +GLU CA CB HBY 108.677 1.74 +GLU CA CB HBX 108.677 1.74 +GLU CG CB HBY 108.696 2.80 +GLU CG CB HBX 108.696 2.80 +GLU HBY CB HBX 107.655 1.50 +GLU CB CG CD 114.140 3.00 +GLU CB CG HGY 108.968 1.50 +GLU CB CG HGX 108.968 1.50 +GLU CD CG HGY 108.472 1.50 +GLU CD CG HGX 108.472 1.50 +GLU HGY CG HGX 107.541 1.92 +GLU CG CD OE1 118.251 3.00 +GLU CG CD OE2 118.251 3.00 +GLU OE1 CD OE2 123.498 1.82 + +loop_ +_chem_comp_tor.comp_id +_chem_comp_tor.id +_chem_comp_tor.atom_id_1 +_chem_comp_tor.atom_id_2 +_chem_comp_tor.atom_id_3 +_chem_comp_tor.atom_id_4 +_chem_comp_tor.value_angle +_chem_comp_tor.value_angle_esd +_chem_comp_tor.period +GLU chi1 N CA CB CG -60.000 10.0 3 +GLU chi2 CA CB CG CD 180.000 10.0 3 +GLU chi3 CB CG CD OE1 180.000 10.0 6 +GLU sp3_sp3_1 C CA N H 180.000 10.0 3 +GLU sp2_sp3_1 O C CA N 0.000 10.0 6 + +loop_ +_chem_comp_chir.comp_id +_chem_comp_chir.id +_chem_comp_chir.atom_id_centre +_chem_comp_chir.atom_id_1 +_chem_comp_chir.atom_id_2 +_chem_comp_chir.atom_id_3 +_chem_comp_chir.volume_sign +GLU chir_1 CA N C CB positive + +loop_ +_chem_comp_plane_atom.comp_id +_chem_comp_plane_atom.plane_id +_chem_comp_plane_atom.atom_id +_chem_comp_plane_atom.dist_esd +GLU plan-1 C 0.020 +GLU plan-1 CA 0.020 +GLU plan-1 O 0.020 +GLU plan-1 OXT 0.020 +GLU plan-2 CD 0.020 +GLU plan-2 CG 0.020 +GLU plan-2 OE1 0.020 +GLU plan-2 OE2 0.020 + +loop_ +_pdbx_chem_comp_descriptor.comp_id +_pdbx_chem_comp_descriptor.type +_pdbx_chem_comp_descriptor.program +_pdbx_chem_comp_descriptor.program_version +_pdbx_chem_comp_descriptor.descriptor +GLU SMILES ACDLabs 12.01 O=C(O)C(N)CCC(=O)O +GLU SMILES_CANONICAL CACTVS 3.370 N[C@@H](CCC(O)=O)C(O)=O +GLU SMILES CACTVS 3.370 N[CH](CCC(O)=O)C(O)=O +GLU SMILES_CANONICAL "OpenEye OEToolkits" 1.7.0 C(CC(=O)O)[C@@H](C(=O)O)N +GLU SMILES "OpenEye OEToolkits" 1.7.0 C(CC(=O)O)C(C(=O)O)N +GLU InChI InChI 1.03 InChI=1S/C5H9NO4/c6-3(5(9)10)1-2-4(7)8/h3H,1-2,6H2,(H,7,8)(H,9,10)/t3-/m0/s1 +GLU InChIKey InChI 1.03 WHUUTDBJXJRKMK-VKHMYHEASA-N + +loop_ +_pdbx_chem_comp_description_generator.comp_id +_pdbx_chem_comp_description_generator.program_name +_pdbx_chem_comp_description_generator.program_version +_pdbx_chem_comp_description_generator.descriptor +GLU acedrg 278 "dictionary generator" +GLU acedrg_database 12 "data source" +GLU rdkit 2019.09.1 "Chemoinformatics tool" +GLU refmac5 5.8.0419 "optimization tool" + +data_comp_ASP +loop_ +_chem_comp_atom.comp_id +_chem_comp_atom.atom_id +_chem_comp_atom.type_symbol +_chem_comp_atom.type_energy +_chem_comp_atom.charge +_chem_comp_atom.x +_chem_comp_atom.y +_chem_comp_atom.z +ASP N N NT3 1 33.542 17.835 39.145 +ASP CA C CH1 0 34.991 17.614 38.867 +ASP C C C 0 35.178 17.161 37.413 +ASP O O O 0 36.268 17.442 36.867 +ASP CB C CH2 0 35.617 16.650 39.866 +ASP CG C C 0 34.932 15.299 40.037 +ASP OD1 O O 0 35.515 14.435 40.725 +ASP OD2 O OC -1 33.817 15.122 39.501 +ASP OXT O OC -1 34.232 16.542 36.875 +ASP H H H 0 33.408 17.940 40.031 +ASP H2 H H 0 33.045 17.139 38.857 +ASP H3 H H 0 33.263 18.581 38.722 +ASP HA H H 0 35.453 18.476 38.978 +ASP HBR H H 0 35.641 17.087 40.742 +ASP HBQ H H 0 36.543 16.483 39.592 + +loop_ +_chem_comp_tree.comp_id +_chem_comp_tree.atom_id +_chem_comp_tree.atom_back +_chem_comp_tree.atom_forward +_chem_comp_tree.connect_type +ASP N n/a CA START +ASP H N . . +ASP H2 N . . +ASP H3 N . . +ASP CA N C . +ASP HA CA . . +ASP CB CA CG . +ASP HBR CB . . +ASP HBQ CB . . +ASP CG CB OD2 . +ASP OD1 CG . . +ASP OD2 CG . . +ASP C CA . END +ASP O C . . +ASP OXT C . . + +loop_ +_chem_comp_acedrg.comp_id +_chem_comp_acedrg.atom_id +_chem_comp_acedrg.atom_type +ASP N N(CCCH)(H)3 +ASP CA C(CCHH)(NH3)(COO)(H) +ASP C C(CCHN)(O)2 +ASP O O(CCO) +ASP CB C(CCHN)(COO)(H)2 +ASP CG C(CCHH)(O)2 +ASP OD1 O(CCO) +ASP OD2 O(CCO) +ASP OXT O(CCO) +ASP H H(NCHH) +ASP H2 H(NCHH) +ASP H3 H(NCHH) +ASP HA H(CCCN) +ASP HBR H(CCCH) +ASP HBQ H(CCCH) + +loop_ +_chem_comp_bond.comp_id +_chem_comp_bond.atom_id_1 +_chem_comp_bond.atom_id_2 +_chem_comp_bond.type +_chem_comp_bond.aromatic +_chem_comp_bond.value_dist_nucleus +_chem_comp_bond.value_dist_nucleus_esd +_chem_comp_bond.value_dist +_chem_comp_bond.value_dist_esd +ASP N CA SINGLE n 1.490 0.0100 1.490 0.0100 +ASP CA C SINGLE n 1.533 0.0100 1.533 0.0100 +ASP CA CB SINGLE n 1.521 0.0100 1.521 0.0100 +ASP C O DOUBLE n 1.251 0.0183 1.251 0.0183 +ASP C OXT SINGLE n 1.251 0.0183 1.251 0.0183 +ASP CB CG SINGLE n 1.522 0.0100 1.522 0.0100 +ASP CG OD1 DOUBLE n 1.249 0.0161 1.249 0.0161 +ASP CG OD2 SINGLE n 1.249 0.0161 1.249 0.0161 +ASP N H SINGLE n 1.018 0.0520 0.902 0.0102 +ASP N H2 SINGLE n 1.018 0.0520 0.902 0.0102 +ASP N H3 SINGLE n 1.018 0.0520 0.902 0.0102 +ASP CA HA SINGLE n 1.092 0.0100 0.984 0.0200 +ASP CB HBR SINGLE n 1.092 0.0100 0.980 0.0165 +ASP CB HBQ SINGLE n 1.092 0.0100 0.980 0.0165 + +loop_ +_chem_comp_angle.comp_id +_chem_comp_angle.atom_id_1 +_chem_comp_angle.atom_id_2 +_chem_comp_angle.atom_id_3 +_chem_comp_angle.value_angle +_chem_comp_angle.value_angle_esd +ASP CA N H 109.990 3.00 +ASP CA N H2 109.990 3.00 +ASP CA N H3 109.990 3.00 +ASP H N H2 109.032 3.00 +ASP H N H3 109.032 3.00 +ASP H2 N H3 109.032 3.00 +ASP N CA C 109.258 1.50 +ASP N CA CB 111.400 1.50 +ASP N CA HA 108.387 1.58 +ASP C CA CB 112.421 3.00 +ASP C CA HA 108.774 1.79 +ASP CB CA HA 108.472 2.65 +ASP CA C O 117.148 1.60 +ASP CA C OXT 117.148 1.60 +ASP O C OXT 125.704 1.50 +ASP CA CB CG 115.436 1.50 +ASP CA CB HBR 108.799 3.00 +ASP CA CB HBQ 108.799 3.00 +ASP CG CB HBR 108.242 2.79 +ASP CG CB HBQ 108.242 2.79 +ASP HBR CB HBQ 107.976 2.66 +ASP CB CG OD1 117.985 1.50 +ASP CB CG OD2 117.985 1.50 +ASP OD1 CG OD2 124.031 1.82 + +loop_ +_chem_comp_tor.comp_id +_chem_comp_tor.id +_chem_comp_tor.atom_id_1 +_chem_comp_tor.atom_id_2 +_chem_comp_tor.atom_id_3 +_chem_comp_tor.atom_id_4 +_chem_comp_tor.value_angle +_chem_comp_tor.value_angle_esd +_chem_comp_tor.period +ASP chi1 N CA CB CG -60.000 10.0 3 +ASP chi2 CA CB CG OD1 180.000 10.0 6 +ASP sp3_sp3_1 C CA N H 180.000 10.0 3 +ASP sp2_sp3_1 O C CA N 0.000 10.0 6 + +loop_ +_chem_comp_chir.comp_id +_chem_comp_chir.id +_chem_comp_chir.atom_id_centre +_chem_comp_chir.atom_id_1 +_chem_comp_chir.atom_id_2 +_chem_comp_chir.atom_id_3 +_chem_comp_chir.volume_sign +ASP chir_1 CA N C CB positive + +loop_ +_chem_comp_plane_atom.comp_id +_chem_comp_plane_atom.plane_id +_chem_comp_plane_atom.atom_id +_chem_comp_plane_atom.dist_esd +ASP plan-1 C 0.020 +ASP plan-1 CA 0.020 +ASP plan-1 O 0.020 +ASP plan-1 OXT 0.020 +ASP plan-2 CB 0.020 +ASP plan-2 CG 0.020 +ASP plan-2 OD1 0.020 +ASP plan-2 OD2 0.020 + +loop_ +_pdbx_chem_comp_descriptor.comp_id +_pdbx_chem_comp_descriptor.type +_pdbx_chem_comp_descriptor.program +_pdbx_chem_comp_descriptor.program_version +_pdbx_chem_comp_descriptor.descriptor +ASP SMILES ACDLabs 12.01 O=C(O)CC(N)C(=O)O +ASP SMILES_CANONICAL CACTVS 3.370 N[C@@H](CC(O)=O)C(O)=O +ASP SMILES CACTVS 3.370 N[CH](CC(O)=O)C(O)=O +ASP SMILES_CANONICAL "OpenEye OEToolkits" 1.7.0 C([C@@H](C(=O)O)N)C(=O)O +ASP SMILES "OpenEye OEToolkits" 1.7.0 C(C(C(=O)O)N)C(=O)O +ASP InChI InChI 1.03 InChI=1S/C4H7NO4/c5-2(4(8)9)1-3(6)7/h2H,1,5H2,(H,6,7)(H,8,9)/t2-/m0/s1 +ASP InChIKey InChI 1.03 CKLJMWTZIZZHCS-REOHCLBHSA-N + +loop_ +_pdbx_chem_comp_description_generator.comp_id +_pdbx_chem_comp_description_generator.program_name +_pdbx_chem_comp_description_generator.program_version +_pdbx_chem_comp_description_generator.descriptor +ASP acedrg 278 "dictionary generator" +ASP acedrg_database 12 "data source" +ASP rdkit 2019.09.1 "Chemoinformatics tool" +ASP refmac5 5.8.0419 "optimization tool" diff --git a/tests/files/restraints/GLU_renamed.cif b/tests/files/restraints/GLU_renamed.cif new file mode 100644 index 00000000..bfec34e2 --- /dev/null +++ b/tests/files/restraints/GLU_renamed.cif @@ -0,0 +1,217 @@ +data_comp_list +loop_ +_chem_comp.id +_chem_comp.three_letter_code +_chem_comp.name +_chem_comp.group +_chem_comp.number_atoms_all +_chem_comp.number_atoms_nh +_chem_comp.desc_level +GLU GLU "GLUTAMIC ACID" peptide 18 10 . + +data_comp_GLU +loop_ +_chem_comp_atom.comp_id +_chem_comp_atom.atom_id +_chem_comp_atom.type_symbol +_chem_comp_atom.type_energy +_chem_comp_atom.charge +_chem_comp_atom.x +_chem_comp_atom.y +_chem_comp_atom.z +GLU N N NT3 1 88.319 -7.751 -10.089 +GLU CA C CH1 0 87.677 -7.162 -11.296 +GLU C C C 0 88.359 -5.826 -11.640 +GLU O O O 0 88.389 -4.954 -10.744 +GLU CB C CH2 0 86.177 -6.950 -11.080 +GLU CG C CH2 0 85.390 -8.247 -10.908 +GLU CD C C 0 83.891 -8.048 -10.773 +GLU OE1 O O 0 83.281 -7.495 -11.711 +GLU OE2 O OC -1 83.334 -8.447 -9.728 +GLU OXT O OC -1 88.836 -5.708 -12.790 +GLU H H H 0 88.066 -7.298 -9.352 +GLU H2 H H 0 89.218 -7.717 -10.158 +GLU H3 H H 0 88.077 -8.615 -9.999 +GLU HAX H H 0 87.806 -7.788 -12.054 +GLU HBY H H 0 85.814 -6.460 -11.847 +GLU HBX H H 0 86.049 -6.394 -10.284 +GLU HGY H H 0 85.714 -8.714 -10.110 +GLU HGX H H 0 85.559 -8.827 -11.681 + +loop_ +_chem_comp_tree.comp_id +_chem_comp_tree.atom_id +_chem_comp_tree.atom_back +_chem_comp_tree.atom_forward +_chem_comp_tree.connect_type +GLU N n/a CA START +GLU H N . . +GLU H2 N . . +GLU H3 N . . +GLU CA N C . +GLU HAX CA . . +GLU CB CA CG . +GLU HBY CB . . +GLU HBX CB . . +GLU CG CB CD . +GLU HGY CG . . +GLU HGX CG . . +GLU CD CG OE2 . +GLU OE1 CD . . +GLU OE2 CD . . +GLU C CA . END +GLU O C . . +GLU OXT C . . + +loop_ +_chem_comp_acedrg.comp_id +_chem_comp_acedrg.atom_id +_chem_comp_acedrg.atom_type +GLU N N(CCCH)(H)3 +GLU CA C(CCHH)(NH3)(COO)(H) +GLU C C(CCHN)(O)2 +GLU O O(CCO) +GLU CB C(CCHH)(CCHN)(H)2 +GLU CG C(CCHH)(COO)(H)2 +GLU CD C(CCHH)(O)2 +GLU OE1 O(CCO) +GLU OE2 O(CCO) +GLU OXT O(CCO) +GLU H H(NCHH) +GLU H2 H(NCHH) +GLU H3 H(NCHH) +GLU HAX H(CCCN) +GLU HBY H(CCCH) +GLU HBX H(CCCH) +GLU HGY H(CCCH) +GLU HGX H(CCCH) + +loop_ +_chem_comp_bond.comp_id +_chem_comp_bond.atom_id_1 +_chem_comp_bond.atom_id_2 +_chem_comp_bond.type +_chem_comp_bond.aromatic +_chem_comp_bond.value_dist_nucleus +_chem_comp_bond.value_dist_nucleus_esd +_chem_comp_bond.value_dist +_chem_comp_bond.value_dist_esd +GLU N CA SINGLE n 1.487 0.0100 1.487 0.0100 +GLU CA C SINGLE n 1.538 0.0113 1.538 0.0113 +GLU CA CB SINGLE n 1.529 0.0100 1.529 0.0100 +GLU C O DOUBLE n 1.251 0.0183 1.251 0.0183 +GLU C OXT SINGLE n 1.251 0.0183 1.251 0.0183 +GLU CB CG SINGLE n 1.526 0.0100 1.526 0.0100 +GLU CG CD SINGLE n 1.518 0.0135 1.518 0.0135 +GLU CD OE1 DOUBLE n 1.249 0.0161 1.249 0.0161 +GLU CD OE2 SINGLE n 1.249 0.0161 1.249 0.0161 +GLU N H SINGLE n 1.018 0.0520 0.902 0.0102 +GLU N H2 SINGLE n 1.018 0.0520 0.902 0.0102 +GLU N H3 SINGLE n 1.018 0.0520 0.902 0.0102 +GLU CA HAX SINGLE n 1.092 0.0100 0.991 0.0200 +GLU CB HBY SINGLE n 1.092 0.0100 0.980 0.0168 +GLU CB HBX SINGLE n 1.092 0.0100 0.980 0.0168 +GLU CG HGY SINGLE n 1.092 0.0100 0.981 0.0172 +GLU CG HGX SINGLE n 1.092 0.0100 0.981 0.0172 + +loop_ +_chem_comp_angle.comp_id +_chem_comp_angle.atom_id_1 +_chem_comp_angle.atom_id_2 +_chem_comp_angle.atom_id_3 +_chem_comp_angle.value_angle +_chem_comp_angle.value_angle_esd +GLU CA N H 109.990 3.00 +GLU CA N H2 109.990 3.00 +GLU CA N H3 109.990 3.00 +GLU H N H2 109.032 3.00 +GLU H N H3 109.032 3.00 +GLU H2 N H3 109.032 3.00 +GLU N CA C 109.258 1.50 +GLU N CA CB 110.440 2.46 +GLU N CA HAX 108.387 1.58 +GLU C CA CB 111.059 3.00 +GLU C CA HAX 108.774 1.79 +GLU CB CA HAX 109.080 2.33 +GLU CA C O 117.148 1.60 +GLU CA C OXT 117.148 1.60 +GLU O C OXT 125.704 1.50 +GLU CA CB CG 113.294 1.61 +GLU CA CB HBY 108.677 1.74 +GLU CA CB HBX 108.677 1.74 +GLU CG CB HBY 108.696 2.80 +GLU CG CB HBX 108.696 2.80 +GLU HBY CB HBX 107.655 1.50 +GLU CB CG CD 114.140 3.00 +GLU CB CG HGY 108.968 1.50 +GLU CB CG HGX 108.968 1.50 +GLU CD CG HGY 108.472 1.50 +GLU CD CG HGX 108.472 1.50 +GLU HGY CG HGX 107.541 1.92 +GLU CG CD OE1 118.251 3.00 +GLU CG CD OE2 118.251 3.00 +GLU OE1 CD OE2 123.498 1.82 + +loop_ +_chem_comp_tor.comp_id +_chem_comp_tor.id +_chem_comp_tor.atom_id_1 +_chem_comp_tor.atom_id_2 +_chem_comp_tor.atom_id_3 +_chem_comp_tor.atom_id_4 +_chem_comp_tor.value_angle +_chem_comp_tor.value_angle_esd +_chem_comp_tor.period +GLU chi1 N CA CB CG -60.000 10.0 3 +GLU chi2 CA CB CG CD 180.000 10.0 3 +GLU chi3 CB CG CD OE1 180.000 10.0 6 +GLU sp3_sp3_1 C CA N H 180.000 10.0 3 +GLU sp2_sp3_1 O C CA N 0.000 10.0 6 + +loop_ +_chem_comp_chir.comp_id +_chem_comp_chir.id +_chem_comp_chir.atom_id_centre +_chem_comp_chir.atom_id_1 +_chem_comp_chir.atom_id_2 +_chem_comp_chir.atom_id_3 +_chem_comp_chir.volume_sign +GLU chir_1 CA N C CB positive + +loop_ +_chem_comp_plane_atom.comp_id +_chem_comp_plane_atom.plane_id +_chem_comp_plane_atom.atom_id +_chem_comp_plane_atom.dist_esd +GLU plan-1 C 0.020 +GLU plan-1 CA 0.020 +GLU plan-1 O 0.020 +GLU plan-1 OXT 0.020 +GLU plan-2 CD 0.020 +GLU plan-2 CG 0.020 +GLU plan-2 OE1 0.020 +GLU plan-2 OE2 0.020 + +loop_ +_pdbx_chem_comp_descriptor.comp_id +_pdbx_chem_comp_descriptor.type +_pdbx_chem_comp_descriptor.program +_pdbx_chem_comp_descriptor.program_version +_pdbx_chem_comp_descriptor.descriptor +GLU SMILES ACDLabs 12.01 O=C(O)C(N)CCC(=O)O +GLU SMILES_CANONICAL CACTVS 3.370 N[C@@H](CCC(O)=O)C(O)=O +GLU SMILES CACTVS 3.370 N[CH](CCC(O)=O)C(O)=O +GLU SMILES_CANONICAL "OpenEye OEToolkits" 1.7.0 C(CC(=O)O)[C@@H](C(=O)O)N +GLU SMILES "OpenEye OEToolkits" 1.7.0 C(CC(=O)O)C(C(=O)O)N +GLU InChI InChI 1.03 InChI=1S/C5H9NO4/c6-3(5(9)10)1-2-4(7)8/h3H,1-2,6H2,(H,7,8)(H,9,10)/t3-/m0/s1 +GLU InChIKey InChI 1.03 WHUUTDBJXJRKMK-VKHMYHEASA-N + +loop_ +_pdbx_chem_comp_description_generator.comp_id +_pdbx_chem_comp_description_generator.program_name +_pdbx_chem_comp_description_generator.program_version +_pdbx_chem_comp_description_generator.descriptor +GLU acedrg 278 "dictionary generator" +GLU acedrg_database 12 "data source" +GLU rdkit 2019.09.1 "Chemoinformatics tool" +GLU refmac5 5.8.0419 "optimization tool" diff --git a/tests/integration/test_cli_hydrogens.py b/tests/integration/test_cli_hydrogens.py index f82a7357..f29b2235 100644 --- a/tests/integration/test_cli_hydrogens.py +++ b/tests/integration/test_cli_hydrogens.py @@ -58,6 +58,47 @@ def stop_after_loading(refinement: Refinement) -> None: assert (n_hydrogens > 0) is add_hydrogens +@pytest.mark.integration +def test_cli_generates_from_the_user_cif( + test_files_dir: Path, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """``--cif`` reaches the model before it loads, so generation reads it.""" + from torchref.cli import refine + + loaded = [] + + def stop_after_loading(refinement: Refinement) -> None: + loaded.append(refinement.model) + raise _ModelLoaded + + monkeypatch.setattr(Refinement, "_sync_model_cell_to_data", stop_after_loading) + cif = str(test_files_dir / "restraints" / "GLU_renamed.cif") + argv = [ + "torchref.refine", + "-m", + str(test_files_dir / "pdb" / "1DAW.pdb"), + "-sf", + str(test_files_dir / "mtz" / "1DAW.mtz"), + "-o", + str(tmp_path / "refined"), + "-v", + "0", + "--add-hydrogens", + "--cif", + cif, + ] + monkeypatch.setattr(sys, "argv", argv) + with pytest.raises(_ModelLoaded): + refine.main() + + (model,) = loaded + registered = model.ctx.cif_path + assert (registered if isinstance(registered, list) else [registered]) == [cif] + pdb = model.pdb + glu_h = (pdb["resname"].str.strip() == "GLU") & (pdb["element"].str.strip() == "H") + assert {"HAX", "HBX", "HBY", "HGX", "HGY"} <= set(pdb.loc[glu_h, "name"].str.strip()) + + @pytest.mark.unit @pytest.mark.parametrize("add_hydrogens", [False, True]) def test_empty_refinement_preserves_hydrogen_setting(add_hydrogens: bool) -> None: diff --git a/tests/unit/io/test_struct_conn_links.py b/tests/unit/io/test_struct_conn_links.py new file mode 100644 index 00000000..2e9bc234 --- /dev/null +++ b/tests/unit/io/test_struct_conn_links.py @@ -0,0 +1,74 @@ +"""``_struct_conn`` rows become the LINK-record table the PDB reader produces.""" + +import numpy as np +import pytest + +from torchref.io.cif_readers import ModelCIFReader +from torchref.io.pdb import LINK_COLUMNS + + +@pytest.mark.unit +def test_3e98_peptide_links_to_selenomethionine(cif_dir): + links = ModelCIFReader(str(cif_dir / "3E98.cif")).links + assert list(links.columns) == list(LINK_COLUMNS) + assert len(links) == 8 + assert set(links["name1"]) == {"C"} and set(links["name2"]) == {"N"} + assert set(links["resname2"]) <= {"MSE", "ARG", "ASP"} + assert (links["altloc1"] == "").all() and (links["icode1"] == "").all() + assert links["resseq1"].dtype.kind == "i" + assert np.isfinite(links["length"]).all() + + +@pytest.mark.unit +def test_1daw_metal_contacts_are_kept(cif_dir): + links = ModelCIFReader(str(cif_dir / "1DAW.cif")).links + assert len(links) == 14 + magnesium = links[(links["resname2"] == "MG") | (links["resname1"] == "MG")] + assert len(magnesium) > 0 + pairs = set(zip(links["name1"], links["resname1"], links["resseq1"])) + assert ("OD2", "ASP", 175) in pairs + + +@pytest.mark.unit +def test_disulfides_are_left_to_distance_detection(cif_dir): + links = ModelCIFReader(str(cif_dir / "3A5V.cif")).links + assert len(links) == 12 + assert "SG" not in set(links["name1"]) | set(links["name2"]) + + +@pytest.mark.unit +def test_file_without_struct_conn_gives_empty_table(tmp_path): + minimal = """\ +data_test +_cell.length_a 10.0 +_cell.length_b 10.0 +_cell.length_c 10.0 +_cell.angle_alpha 90.0 +_cell.angle_beta 90.0 +_cell.angle_gamma 90.0 +_symmetry.space_group_name_H-M 'P 1' +loop_ +_atom_site.group_PDB +_atom_site.id +_atom_site.type_symbol +_atom_site.label_atom_id +_atom_site.label_alt_id +_atom_site.label_comp_id +_atom_site.label_asym_id +_atom_site.label_seq_id +_atom_site.pdbx_PDB_ins_code +_atom_site.Cartn_x +_atom_site.Cartn_y +_atom_site.Cartn_z +_atom_site.occupancy +_atom_site.B_iso_or_equiv +_atom_site.auth_seq_id +_atom_site.auth_asym_id +ATOM 1 N N . ALA A 1 ? 0.0 0.0 0.0 1.0 20.0 1 A +ATOM 2 C CA . ALA A 1 ? 1.5 0.0 0.0 1.0 20.0 1 A +""" + path = tmp_path / "no_links.cif" + path.write_text(minimal) + links = ModelCIFReader(str(path)).links + assert len(links) == 0 + assert list(links.columns) == list(LINK_COLUMNS) diff --git a/tests/unit/model/test_hydrogen_default.py b/tests/unit/model/test_hydrogen_default.py index 3c6f1cf3..5203331c 100644 --- a/tests/unit/model/test_hydrogen_default.py +++ b/tests/unit/model/test_hydrogen_default.py @@ -196,3 +196,109 @@ def test_riding_hydrogens_are_not_placed_when_real_ones_exist(pdb_dir): assert ( stripped.restraints.h_topo.n_hydrogens > 0 ), "with hydrogens absent the riding stand-in should still be built" + + +# --- The user's restraint dictionary is the one that hydrogenates ------------------ + +RENAMED_GLU_H = {"HAX", "HBX", "HBY", "HGX", "HGY"} + + +@pytest.fixture +def renamed_glu_cif(test_files_dir): + """A GLU dictionary whose side-chain hydrogens carry names the library lacks.""" + return str(test_files_dir / "restraints" / "GLU_renamed.cif") + + +def _glu_hydrogen_names(model): + pdb = model.pdb + glu_h = (pdb["resname"].astype(str).str.strip() == "GLU") & ( + pdb["element"].astype(str).str.strip() == "H" + ) + return set(pdb.loc[glu_h, "name"].astype(str).str.strip()) + + +@pytest.mark.unit +def test_generation_reads_the_cif_given_at_construction(pdb_dir, renamed_glu_cif): + """A dictionary passed to the constructor overrides the library for generation. + + The names prove which dictionary was read, and the bond degree proves the generated + hydrogens are the ones the restraints know: a hydrogen generated from one + dictionary and restrained by another has no bond edge at all. + """ + model = Model(verbose=0, add_hydrogens=True, cif_path=renamed_glu_cif) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + assert model.ctx.cif_path == renamed_glu_cif + names = _glu_hydrogen_names(model) + assert RENAMED_GLU_H <= names + assert not {"HA", "HB2", "HB3", "HG2", "HG3"} & names + + atoms = model.restraints.topology.atoms + is_h = atoms.is_hydrogen.cpu().numpy() + degree = atoms.degree().cpu().numpy() + assert (degree[is_h] > 0).all(), "generated hydrogens without a bond restraint" + + +@pytest.mark.unit +def test_derived_models_keep_the_restraint_cif(pdb_dir, renamed_glu_cif): + """hydrogenate, strip_hydrogens and select all carry the dictionary along.""" + model = Model(verbose=0, add_hydrogens=False, cif_path=renamed_glu_cif) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + + hydrogenated = model.hydrogenate() + assert hydrogenated.ctx.cif_path == renamed_glu_cif + assert RENAMED_GLU_H <= _glu_hydrogen_names(hydrogenated) + assert "GLU" in hydrogenated.restraints.cif_dict + template_h = set( + hydrogenated.restraints.cif_dict["GLU"]["atoms"]["atom_id"].astype(str).str.strip() + ) + assert RENAMED_GLU_H <= template_h + + assert hydrogenated.strip_hydrogens().ctx.cif_path == renamed_glu_cif + assert model.select("resname GLU").ctx.cif_path == renamed_glu_cif + + +@pytest.mark.unit +def test_state_dict_round_trips_the_restraint_cif(pdb_dir, renamed_glu_cif): + model = Model(verbose=0, cif_path=renamed_glu_cif) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + restored = Model.create_from_state_dict(model.state_dict(), verbose=0) + assert restored.ctx.cif_path == renamed_glu_cif + + +@pytest.mark.unit +def test_load_model_registers_the_cif_before_loading(pdb_dir, renamed_glu_cif): + """The shared CLI loader generates from the user dictionary too.""" + from torchref.cli._common import load_model + + model = load_model( + str(pdb_dir / "1DAW.pdb"), verbose=0, cif=renamed_glu_cif, add_hydrogens=True + ) + assert model.ctx.cif_path == renamed_glu_cif + assert RENAMED_GLU_H <= _glu_hydrogen_names(model) + + +@pytest.mark.unit +def test_generation_reads_every_compound_of_a_multi_block_cif(pdb_dir, test_files_dir): + """A dictionary with several ``data_comp_`` blocks hydrogenates each of its compounds. + + Multi-compound dictionaries once restrained only their last block; generation now + reads the same dictionary, so both renamed sets must appear and every generated + hydrogen must carry a bond edge. + """ + cif = test_files_dir / "restraints" / "GLU_ASP_renamed.cif" + blocks = [l for l in cif.read_text().splitlines() if l.startswith("data_comp_")] + assert len(blocks) == 3, blocks # comp_list + GLU + ASP: the fixture is really multi-block + + model = Model(verbose=0, add_hydrogens=True, cif_path=str(cif)) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + pdb = model.pdb + is_h = pdb["element"].astype(str).str.strip() == "H" + resname = pdb["resname"].astype(str).str.strip() + names = pdb["name"].astype(str).str.strip() + assert RENAMED_GLU_H <= set(names[is_h & (resname == "GLU")]) + assert {"HBQ", "HBR"} <= set(names[is_h & (resname == "ASP")]) + assert not {"HB2", "HB3"} & set(names[is_h & (resname == "ASP")]) + + atoms = model.restraints.topology.atoms + degree = atoms.degree().cpu().numpy() + assert (degree[atoms.is_hydrogen.cpu().numpy()] > 0).all() diff --git a/tests/unit/topology/test_equivalence.py b/tests/unit/topology/test_equivalence.py index 883ed1d6..0e276804 100644 --- a/tests/unit/topology/test_equivalence.py +++ b/tests/unit/topology/test_equivalence.py @@ -184,7 +184,9 @@ def test_adjacency_matches_bond_block(built, code): from_block.add((min(int(a), int(b)), max(int(a), int(b)))) assert from_adjacency == from_block - assert int(atoms.degree().sum()) == 2 * atoms.bonds.n_edges + # Each distinct partner once: a bond row repeated per altloc conformer or per + # duplicated LINK record does not add to the degree. + assert int(atoms.degree().sum()) == 2 * len(from_block) @pytest.mark.unit diff --git a/tests/unit/topology/test_hydrogens.py b/tests/unit/topology/test_hydrogens.py index 0b83d9cc..922ea04b 100644 --- a/tests/unit/topology/test_hydrogens.py +++ b/tests/unit/topology/test_hydrogens.py @@ -13,12 +13,14 @@ from torchref.model.model import Model from torchref.topology.hydrogens import ( STANDARD_VALENCE, + _template, augment_atom_table, optimise_free_torsions, plan_hydrogens, ) -STRUCTURES = ["7L84", "1DAW"] +# 3E98 brings HETATM selenomethionines bonded through LINK records and split side chains. +STRUCTURES = ["7L84", "1DAW", "3E98"] @pytest.fixture(scope="module") @@ -66,22 +68,44 @@ def test_every_candidate_hydrogen_is_placed(built, code): """ model, restraints, plan = built(code) topology = restraints.topology - - # One hydrogen per free valence on every parent that has a template hydrogen. - expected_parents = set(plan.parent.tolist()) - assert expected_parents, "no parents received hydrogens" - - is_h = topology.atoms.is_hydrogen - for parent in sorted(expected_parents): - neighbours = topology.atoms.neighbors(parent) - heavy = int((~is_h[neighbours]).sum()) - element = str(topology.atoms.element[parent]).strip().upper() - allowed = max(0, STANDARD_VALENCE.get(element, 4) - heavy) - placed = int((plan.parent == parent).sum()) - assert placed <= allowed, ( - f"{code}: atom {parent} ({element}) has {heavy} heavy bonds, so at most " - f"{allowed} hydrogens, but {placed} were planned" - ) + atoms = topology.atoms + residues = topology.residues + assert plan.n_hydrogens, "no hydrogens planned" + + is_h = atoms.is_hydrogen + altlocs = np.char.strip(atoms.altloc.astype(str)) + names = np.char.strip(atoms.name.astype(str)) + checked = 0 + for residue in range(residues.n_residues): + start, end = int(residues.atom_start[residue]), int(residues.atom_end[residue]) + template = _template(restraints.cif_dict, str(residues.resname[residue]).strip()) + if template is None: + continue + # Residues with altlocs plan one hydrogen per conformer; the two-sided count + # below is for the plain case, where the graph degree is the whole story. + if (altlocs[start:end] != "").any(): + continue + for parent in range(start, end): + template_h = template["h_count"].get(names[parent], 0) + if template_h == 0: + continue + neighbours = atoms.neighbors(parent) + heavy = int((~is_h[neighbours]).sum()) + element = str(atoms.element[parent]).strip().upper() + template_heavy = len(template["heavy_adjacency"].get(names[parent], [])) + extra_bonds = max(0, heavy - template_heavy) + expected = max( + 0, + min(STANDARD_VALENCE.get(element, 4) - heavy, template_h - extra_bonds), + ) + placed = int((plan.parent == parent).sum()) + assert placed == expected, ( + f"{code}: atom {parent} ({names[parent]} {element}) has {heavy} heavy " + f"bonds against {template_heavy} in the template and {template_h} " + f"template hydrogens, so {expected} expected, {placed} planned" + ) + checked += 1 + assert checked > 0 @pytest.mark.unit @@ -283,3 +307,88 @@ def test_strip_H_removes_deposited_hydrogens(pdb_dir): model.load_pdb(str(pdb_dir / "1AK5_with_H.pdb")) elements = model.pdb["element"].astype(str).str.strip().values assert not (elements == "H").any() + + +def _row(model, chain, resseq, name, altloc=""): + pdb = model.pdb + mask = ( + (pdb["chainid"].astype(str) == chain) + & (pdb["resseq"].astype(int) == resseq) + & (pdb["name"].astype(str).str.strip() == name) + & (pdb["altloc"].astype(str).str.strip() == altloc) + ) + (row,) = np.nonzero(mask.values)[0] + return int(row) + + +@pytest.mark.unit +def test_linked_nitrogen_keeps_one_hydrogen(built): + """A peptide bond supplied by a LINK record displaces two of the template's three. + + MSE is a HETATM residue, so its backbone bonds come only from LINK records. MSE65 + has a split side chain: its shared N gets one hydrogen per conformer, and its + carbonyl carbon, bonded to CA(A), CA(B), O and the next N, gets none. + """ + model, _, plan = built("3E98") + n73 = _row(model, "A", 73, "N") + assert plan.name[plan.parent == n73].tolist() == ["H"] + + n65 = _row(model, "A", 65, "N") + on_n65 = plan.parent == n65 + assert plan.name[on_n65].tolist() == ["H", "H"] + assert sorted(plan.altloc[on_n65].tolist()) == ["A", "B"] + assert not (plan.parent == _row(model, "A", 65, "C")).any() + + +@pytest.mark.unit +def test_split_side_chain_keeps_its_alpha_hydrogen(built): + """A CA bonded to two altloc copies of CB is not saturated: one HA per conformer.""" + model, _, plan = built("3E98") + ca_a = _row(model, "A", 65, "CA", "A") + ca_b = _row(model, "A", 65, "CA", "B") + assert plan.name[plan.parent == ca_a].tolist() == ["HA"] + assert plan.name[plan.parent == ca_b].tolist() == ["HA"] + + +# 1U19 chain A: an acetyl cap bonded to MET1 through a LINK record. The ACE template +# is acetaldehyde-like, with a hydrogen on the carbonyl carbon that the peptide link +# displaces; the element valence alone (degree 3 of 4) would still have generated it. +_ACE_MET = """\ +CRYST1 96.680 96.680 150.200 90.00 90.00 90.00 P 41 8 +LINK C ACE A 0 N MET A 1 1555 1555 1.33 +HETATM 1 C ACE A 0 53.553 -7.050 35.606 1.00 47.40 C +HETATM 2 O ACE A 0 52.916 -7.860 34.934 1.00 46.96 O +HETATM 3 CH3 ACE A 0 54.727 -7.523 36.434 1.00 47.42 C +ATOM 4 N MET A 1 53.284 -5.731 35.670 1.00 47.11 N +ATOM 5 CA MET A 1 52.214 -5.077 34.913 1.00 46.26 C +ATOM 6 C MET A 1 52.674 -4.891 33.485 1.00 46.68 C +ATOM 7 O MET A 1 53.849 -4.563 33.283 1.00 46.86 O +ATOM 8 CB MET A 1 51.887 -3.719 35.536 1.00 45.60 C +ATOM 9 CG MET A 1 51.426 -3.792 36.982 1.00 44.48 C +ATOM 10 SD MET A 1 49.945 -4.797 37.183 1.00 46.06 S +ATOM 11 CE MET A 1 48.647 -3.745 36.534 1.00 43.77 C +END +""" + + +@pytest.mark.unit +def test_acetyl_cap_carbon_gets_no_hydrogen(tmp_path): + """The template's own hydrogen count, minus the link, caps the carbonyl carbon.""" + from torchref.topology.monomer.cif import find_cif_file_in_library + + if find_cif_file_in_library("ACE") is None: + pytest.skip("ACE not in the monomer library") + path = tmp_path / "ace_met.pdb" + path.write_text(_ACE_MET) + model = Model(verbose=0, add_hydrogens=False, strip_H=True) + model.load_pdb(str(path)) + model.set_restraints_cif(None) + restraints = model.restraints + plan = plan_hydrogens(restraints.topology, restraints.cif_dict, model.xyz().detach()) + + by_parent = {} + for parent, name in zip(plan.parent.tolist(), plan.name.tolist()): + by_parent.setdefault(parent, []).append(name) + assert _row(model, "A", 0, "C") not in by_parent + assert len(by_parent[_row(model, "A", 0, "CH3")]) == 3 + assert by_parent[_row(model, "A", 1, "N")] == ["H"] diff --git a/tests/unit/topology/test_links.py b/tests/unit/topology/test_links.py new file mode 100644 index 00000000..0a9c2aec --- /dev/null +++ b/tests/unit/topology/test_links.py @@ -0,0 +1,110 @@ +"""Covalent links in the bond graph: counted once, and the same from PDB and mmCIF. + +A repeated LINK record, or a bond emitted once per altloc conformer between two shared +atoms, used to inflate an atom's graph degree. Hydrogen generation reads that degree as +the number of heavy partners, so a Schiff-base nitrogen listed in two identical LINK +records lost its hydrogen and a CA with a split side chain lost its HA. +""" + +import numpy as np +import pytest + +from torchref.model.model import Model + +# 3E98 chain A, LEU72 - MSE73: the MSE is a HETATM residue, so its peptide bonds come +# only from LINK records. +_LEU_MSE = """\ +CRYST1 53.841 88.114 60.963 90.00 107.92 90.00 P 1 21 1 4 +{links} +ATOM 206 N LEU A 72 25.385 1.315 55.882 1.00 51.26 N +ATOM 207 CA LEU A 72 24.279 2.045 56.497 1.00 52.28 C +ATOM 208 C LEU A 72 24.644 2.581 57.879 1.00 51.78 C +ATOM 209 O LEU A 72 24.190 3.667 58.275 1.00 52.62 O +ATOM 210 CB LEU A 72 23.061 1.126 56.613 1.00 52.98 C +ATOM 211 CG LEU A 72 21.688 1.747 56.752 1.00 56.56 C +ATOM 212 CD1 LEU A 72 21.410 2.644 55.536 1.00 55.74 C +ATOM 213 CD2 LEU A 72 20.665 0.619 56.876 1.00 54.03 C +HETATM 214 N MSE A 73 25.451 1.824 58.619 1.00 50.52 N +HETATM 215 CA MSE A 73 25.888 2.255 59.952 1.00 51.25 C +HETATM 216 C MSE A 73 26.955 3.368 59.879 1.00 51.87 C +HETATM 217 O MSE A 73 26.968 4.263 60.731 1.00 52.14 O +HETATM 218 CB MSE A 73 26.424 1.077 60.754 1.00 50.55 C +HETATM 219 CG MSE A 73 25.351 0.118 61.312 1.00 56.44 C +HETATM 220 SE MSE A 73 26.197 -1.503 62.054 0.75 51.93 SE +HETATM 221 CE MSE A 73 27.068 -0.667 63.543 1.00 59.19 C +END +""" +_LINK = "LINK C LEU A 72 N MSE A 73 1555 1555 1.33" + + +def _row(model, chain, resseq, name, altloc=""): + pdb = model.pdb + mask = ( + (pdb["chainid"].astype(str) == chain) + & (pdb["resseq"].astype(int) == resseq) + & (pdb["name"].astype(str).str.strip() == name) + & (pdb["altloc"].astype(str).str.strip() == altloc) + ) + (row,) = np.nonzero(mask.values)[0] + return int(row) + + +def _load(path): + model = Model(verbose=0, add_hydrogens=False, strip_H=True) + model.load_pdb(str(path)) if str(path).endswith(".pdb") else model.load_cif(str(path)) + model.set_restraints_cif(None) + return model + + +@pytest.mark.unit +def test_repeated_link_record_contributes_one_edge(tmp_path): + """The parser keeps both records; the graph carries one bond and one restraint.""" + path = tmp_path / "dup_link.pdb" + path.write_text(_LEU_MSE.format(links=_LINK + "\n" + _LINK)) + model = _load(path) + restraints = model.restraints + atoms = restraints.topology.atoms + + assert len(restraints.links) == 2 + link_rows = atoms.bonds.origin("link") + assert len(link_rows) == 1 + n_mse = _row(model, "A", 73, "N") + assert int(atoms.degree(n_mse)) == 2 # CA and the previous C + + +@pytest.mark.unit +def test_shared_atom_degree_counts_each_partner_once(pdb_dir): + """A bond between two blank-altloc atoms is one edge however many conformers emit it.""" + model = _load(pdb_dir / "3E98.pdb") + atoms = model.restraints.topology.atoms + # MSE65 has split CA/CB/CG/SE/CE; its C is bonded to CA(A), CA(B), O and ARG66 N. + assert int(atoms.degree(_row(model, "A", 65, "C"))) == 4 + + model = _load(pdb_dir / "3A5V.pdb") + atoms = model.restraints.topology.atoms + # CYS53 has a split side chain: N sees C(prev) and CA; CA sees N, C, CB(A), CB(B). + assert int(atoms.degree(_row(model, "A", 53, "N"))) == 2 + assert int(atoms.degree(_row(model, "A", 53, "CA"))) == 4 + + +def _link_identities(model): + atoms = model.restraints.topology.atoms + pdb = model.pdb + key = lambda i: ( + str(pdb["chainid"].iloc[i]), + int(pdb["resseq"].iloc[i]), + str(pdb["icode"].iloc[i]).strip(), + str(pdb["name"].iloc[i]).strip(), + str(pdb["altloc"].iloc[i]).strip(), + ) + return {frozenset((key(int(a)), key(int(b)))) for a, b in atoms.bonds.origin("link")} + + +@pytest.mark.unit +@pytest.mark.parametrize("code", ["1DAW", "2DQ6", "3A5V", "3E98", "5BOV", "6G9X"]) +def test_link_edges_agree_between_pdb_and_cif(pdb_dir, cif_dir, code): + """mmCIF ``_struct_conn`` yields the link edges the PDB LINK records do.""" + from_pdb = _load(pdb_dir / f"{code}.pdb") + from_cif = _load(cif_dir / f"{code}.cif") + assert from_cif.ctx.links is not None and len(from_cif.ctx.links) > 0 + assert _link_identities(from_cif) == _link_identities(from_pdb) diff --git a/torchref/cli/_common.py b/torchref/cli/_common.py index da543742..77707686 100644 --- a/torchref/cli/_common.py +++ b/torchref/cli/_common.py @@ -651,6 +651,7 @@ def load_model( device: Union[str, "torch.device", None] = None, verbose: int = 0, cif: Optional[Union[str, List[str]]] = None, + add_hydrogens: bool = False, ) -> "ModelFT": """Load a model from PDB or CIF, auto-detected by file extension. @@ -665,7 +666,10 @@ def load_model( verbose : int Verbosity passed to ModelFT. cif : str or list of str, optional - CIF restraint file(s) to load after the model. + CIF restraint file(s), registered on the model before it loads so that hydrogen + generation and the restraints read the same dictionary. + add_hydrogens : bool, optional + Generate missing hydrogens on load. Default False. Returns ------- @@ -675,16 +679,18 @@ def load_model( from torchref.config import normalize_device device = normalize_device(device) - model = ModelFT(max_res=max_res, device=device, verbose=verbose) + model = ModelFT( + max_res=max_res, + device=device, + verbose=verbose, + cif_path=cif, + add_hydrogens=add_hydrogens, + ) suffix = Path(path).suffix.lower() if suffix in (".cif", ".mmcif"): model.load_cif(path) else: model.load_pdb(path) - - if cif is not None: - model.set_restraints_cif(cif) - return model diff --git a/torchref/experimental/ensemble/ensemble_model.py b/torchref/experimental/ensemble/ensemble_model.py index 9b6f7931..ba0b4ebe 100644 --- a/torchref/experimental/ensemble/ensemble_model.py +++ b/torchref/experimental/ensemble/ensemble_model.py @@ -348,6 +348,7 @@ def __init__( gridsize: Optional[Tuple[int, int, int]] = None, wavelength: float = 1.0, anomalous_threshold: float = 0.5, + cif_path=None, ): if dtype_float is None: dtype_float = get_float_dtype() @@ -363,6 +364,7 @@ def __init__( gridsize=gridsize, wavelength=wavelength, anomalous_threshold=anomalous_threshold, + cif_path=cif_path, ) # Filled in by ``_finalize_ensemble`` after ``load`` returns. self.n_members: int = 0 diff --git a/torchref/experimental/kinetic/refinement.py b/torchref/experimental/kinetic/refinement.py index 94514c0c..70df0215 100644 --- a/torchref/experimental/kinetic/refinement.py +++ b/torchref/experimental/kinetic/refinement.py @@ -13,8 +13,8 @@ from torchref.experimental.kinetic import ModelCollection, KineticRefinement # Base models - model_dark = ModelFT(max_res=1.5).load_pdb("dark.pdb") - model_light = ModelFT(max_res=1.5).load_pdb("light.pdb") + model_dark = ModelFT(max_res=1.5, cif_path=cif_paths).load_pdb("dark.pdb") + model_light = ModelFT(max_res=1.5, cif_path=cif_paths).load_pdb("light.pdb") # Collections models = ModelCollection([model_dark, model_light]) diff --git a/torchref/io/cif_readers.py b/torchref/io/cif_readers.py index 797f4230..fb88292a 100644 --- a/torchref/io/cif_readers.py +++ b/torchref/io/cif_readers.py @@ -1267,7 +1267,11 @@ def _get_value(self, data, possible_keys: List[str], default: Any = None) -> Any class ModelCIFReader: """ Reader for model/structure CIF files (e.g. ``*.cif`` from the PDB): - coordinates, altlocs, ANISOU, cell and space group. + coordinates, altlocs, ANISOU, cell, space group and covalent links. + + Covalent and metal ``_struct_conn`` rows are exposed as ``.links`` in the same table + the PDB reader builds from LINK records, so :meth:`torchref.model.model.Model.load` + picks them up either way. Calling the instance gives the same unpack order as the PDB reader:: @@ -1307,6 +1311,7 @@ def _extract_data(self): self.cell = cell_params self.spacegroup = self.get_space_group() + self.links = self.get_link_records() # Store as DataFrame attributes (like legacy PDB reader) self.dataframe.attrs["cell"] = self.cell @@ -1318,6 +1323,81 @@ def _extract_data(self): print(f" Atoms: {len(self.dataframe)}") print(f" Cell: {self.cell}") print(f" Spacegroup: {self.spacegroup}") + print(f" Links: {len(self.links)}") + + #: ``_struct_conn.conn_type_id`` prefixes that describe a covalent bond the topology + #: should carry. Disulfides are detected from SG-SG distance instead (the PDB reader + #: ignores SSBOND the same way); hydrogen bonds, salt bridges and mismatches are not + #: bonds. + _LINK_CONN_TYPES = ("covale", "metalc") + + #: Symmetry operators under which a ``_struct_conn`` row joins atoms of the same + #: asymmetric unit copy; the PDB reader keeps LINK records with ``1555`` or blank. + _LINK_SYMMETRY_OK = frozenset({"1_555", "", "?", "."}) + + def get_link_records(self) -> pd.DataFrame: + """Covalent and metal links from ``_struct_conn``, in the LINK-record table. + + Same columns as :func:`torchref.io.pdb.extract_link_records`. Rows whose + connection type is not covalent or metal, that cross a symmetry operator, or + whose residue numbers are unreadable are dropped. Blank alternative locations and + insertion codes (``?`` or ``.``) become empty strings, which is what the atom + table carries and what the LINK lookup compares against. + + Returns + ------- + pandas.DataFrame + Empty, with the LINK columns, when the file has no ``_struct_conn`` loop. + """ + from torchref.io.pdb import LINK_COLUMNS + + empty = pd.DataFrame(columns=list(LINK_COLUMNS)) + conn = self.cif.data.get("struct_conn") + if conn is None or len(conn) == 0: + return empty + + def column(names, default=""): + for name in names: + if name in conn.columns: + values = conn[name].astype(str).str.strip() + return values.where(~values.isin(["?", "."]), default) + return pd.Series([default] * len(conn), index=conn.index, dtype=object) + + kind = column(["_struct_conn.conn_type_id"]).str.lower() + keep = kind.str.startswith(self._LINK_CONN_TYPES) + for side in ("1", "2"): + keep &= column([f"_struct_conn.ptnr{side}_symmetry"]).isin( + self._LINK_SYMMETRY_OK + ) + + out = pd.DataFrame(index=conn.index) + for side in ("1", "2"): + ptnr = f"_struct_conn.ptnr{side}_" + pdbx = f"_struct_conn.pdbx_ptnr{side}_" + # Same precedence as get_atom_data: label_* for atom and residue names, + # auth_* for chain and residue number, so the lookup matches the table. + out[f"name{side}"] = column([ptnr + "label_atom_id", ptnr + "auth_atom_id"]) + out[f"altloc{side}"] = column([pdbx + "label_alt_id"]) + out[f"resname{side}"] = column([ptnr + "label_comp_id", ptnr + "auth_comp_id"]) + out[f"chainid{side}"] = column([ptnr + "auth_asym_id", ptnr + "label_asym_id"]) + out[f"resseq{side}"] = pd.to_numeric( + column([ptnr + "auth_seq_id", ptnr + "label_seq_id"], default="nan"), + errors="coerce", + ) + out[f"icode{side}"] = column([pdbx + "PDB_ins_code"]) + out["length"] = pd.to_numeric( + column(["_struct_conn.pdbx_dist_value"], default="nan"), errors="coerce" + ) + + keep &= out["resseq1"].notna() & out["resseq2"].notna() + out = out.loc[keep].copy() + if len(out) == 0: + return empty + out["resseq1"] = out["resseq1"].astype(int) + out["resseq2"] = out["resseq2"].astype(int) + if self.verbose > 1: + print(f"_struct_conn: kept {len(out)} of {len(conn)} rows as links") + return out[list(LINK_COLUMNS)].reset_index(drop=True) def read(self, filepath: str = None): """Re-read ``filepath`` (default: the init path); returns ``self``.""" diff --git a/torchref/io/pdb.py b/torchref/io/pdb.py index cb17fd1f..c7d749f8 100644 --- a/torchref/io/pdb.py +++ b/torchref/io/pdb.py @@ -410,6 +410,15 @@ def extract_pdb_headers(filepath: str) -> list: return headers +#: Columns of the LINK-record table that ``Model.load`` reads off a reader's ``.links``. +#: Shared by the PDB and mmCIF readers so the topology builder sees one schema. +LINK_COLUMNS = ( + "name1", "altloc1", "resname1", "chainid1", "resseq1", "icode1", + "name2", "altloc2", "resname2", "chainid2", "resseq2", "icode2", + "length", +) + + def extract_link_records(filepath: str, verbose: int = 0) -> pd.DataFrame: """Parse LINK records from a PDB file (PDB v3.3 format). @@ -476,14 +485,7 @@ def extract_link_records(filepath: str, verbose: int = 0) -> pd.DataFrame: if verbose > 1: print(f"Warning: skipping malformed LINK: {line.rstrip()}") - df = pd.DataFrame( - rows, - columns=[ - "name1", "altloc1", "resname1", "chainid1", "resseq1", "icode1", - "name2", "altloc2", "resname2", "chainid2", "resseq2", "icode2", - "length", - ], - ) + df = pd.DataFrame(rows, columns=list(LINK_COLUMNS)) if verbose > 0 and (len(df) or skipped_sym or skipped_bad): print( f"LINK records: parsed {len(df)}, " diff --git a/torchref/model/model.py b/torchref/model/model.py index ef207dc1..6fe58176 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -96,6 +96,10 @@ class Model(DeviceMovementMixin, DebugMixin, nn.Module): add_hydrogens : bool, optional Generate missing hydrogens on load when True. Default False; ignored when ``strip_H`` is set. + cif_path : str or list of str, optional + Restraint dictionary file(s) for residues the monomer library does not know, or + whose library entry should be overridden. Given here rather than after loading so + that hydrogen generation on load reads the same dictionary the restraints will. Attributes ---------- @@ -131,6 +135,7 @@ def __init__( device=None, strip_H: bool = False, add_hydrogens: bool = False, + cif_path: Optional[Union[str, List[str]]] = None, ): """ Initialize an empty Model shell. @@ -152,6 +157,10 @@ def __init__( add_hydrogens : bool, optional Generate missing hydrogens on load when True. Default False; ignored when ``strip_H`` is set. + cif_path : str or list of str, optional + Restraint dictionary file(s); see the class docstring. :meth:`set_restraints_cif` + can still change it after loading, but generation on load only sees the value + given here. """ super().__init__() # Resolve dtype/device at call time (not import time) so a runtime @@ -168,7 +177,10 @@ def __init__( # Everything the model is loaded from and sits in, as opposed to what is # refined. Populated by load() / create_from_state_dict(). self.ctx = ModelContext( - verbose=verbose, strip_H=strip_H, add_hydrogens=add_hydrogens + verbose=verbose, + strip_H=strip_H, + add_hydrogens=add_hydrogens, + cif_path=cif_path, ) # Submodules (created during load or load_state_dict) @@ -752,6 +764,12 @@ def _add_missing_hydrogens(self) -> None: plan = plan_hydrogens( restraints.topology, restraints.cif_dict, xyz, verbose=self.ctx.verbose ) + if self.ctx.verbose > 0 and restraints.missing_residues: + print( + "No restraint dictionary for " + f"{sorted(restraints.missing_residues)}: not hydrogenated. Pass one " + "with cif_path / --cif." + ) if plan.n_hydrogens == 0: return optimise_free_torsions(plan, restraints.topology, xyz) @@ -2022,6 +2040,7 @@ def _new_model_from_df(self, df, *, strip_H=None, add_hydrogens=False): device=self.device, strip_H=sh, add_hydrogens=add_hydrogens, + cif_path=self.ctx.cif_path, ) sig = inspect.signature(self.__class__.__init__) for pname, param in sig.parameters.items(): @@ -2041,9 +2060,6 @@ def _new_model_from_df(self, df, *, strip_H=None, add_hydrogens=False): new_model = self.__class__(**ctor_kw) sg_str = self.spacegroup.xhm if self.spacegroup else "P 1" new_model.load(lambda: (df, self.pdb.attrs.get("cell"), sg_str)) - # Propagate CIF restraint paths so restraints are rebuilt correctly - if self.ctx.cif_path is not None: - new_model._cif_path = self.ctx.cif_path return new_model def strip_altlocs(self) -> "Model": @@ -2192,6 +2208,7 @@ def state_dict(self, destination=None, prefix="", keep_vars=False): state[prefix + "dtype_float"] = self.dtype_float state[prefix + "device"] = self.device state[prefix + "strip_H"] = self.ctx.strip_H + state[prefix + "cif_path"] = self.ctx.cif_path state[prefix + "altloc_pairs"] = self.ctx.altloc_pairs return state @@ -2449,10 +2466,15 @@ def create_from_state_dict( saved_dtype = state_dict.pop("dtype_float", dtype_float) state_dict.pop("device", None) # popped so it never reaches load_state_dict strip_H = state_dict.pop("strip_H", True) + cif_path = state_dict.pop("cif_path", None) altloc_pairs = state_dict.pop("altloc_pairs", []) instance = cls( - dtype_float=saved_dtype, verbose=verbose, device=device, strip_H=strip_H + dtype_float=saved_dtype, + verbose=verbose, + device=device, + strip_H=strip_H, + cif_path=cif_path, ) instance.pdb = pdb @@ -2593,6 +2615,7 @@ def select(self, selection: str) -> "Model": verbose=self.ctx.verbose, device=self.device, strip_H=self.ctx.strip_H, + cif_path=self.ctx.cif_path, ) # ``index`` must be renumbered: the occupancy grouping below reads it. diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index 9dedd761..d675ba9a 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -150,8 +150,9 @@ def __init__( Path to the MTZ or CIF file holding reflection data. pdb : str, optional Path to the PDB or CIF file holding the initial model. - cif : str, optional - Path to a CIF file of restraints (monomer library). + cif : str or list of str, optional + Restraint dictionary file(s) for residues the monomer library lacks. Given to + the model at construction so hydrogen generation on load reads it too. verbose : int, optional Verbosity level. Default 1. max_res : float, optional @@ -279,6 +280,7 @@ def __init__( wavelength=self.wavelength, anomalous_threshold=self.anomalous_threshold, add_hydrogens=add_hydrogens, + cif_path=cif, ) self.scaler = Scaler( verbose=self.verbose, device=self.device, nbins=self.nbins, @@ -326,6 +328,8 @@ def __init__( wavelength=self.wavelength, anomalous_threshold=self.anomalous_threshold, add_hydrogens=add_hydrogens, + # Before load, not after: generation on load reads this dictionary. + cif_path=cif, # Apply the f'' (Bijvoet) term only when the data were loaded as # explicit Friedel pairs; merged data gate it off. apply_bijvoet=not self.reflection_data.friedel_merged, @@ -352,8 +356,7 @@ def __init__( reflections_per_parameter=self.reflections_per_adp_parameter, ) self.setup_scaler() - # Configure CIF path for lazy restraint building (restraints built on first access) - self.model.set_restraints_cif(cif) + # The CIF path went in at construction; build the restraints over it now. self.model._build_restraints() self._freeze_unrestrained_residues() diff --git a/torchref/topology/atom_graph.py b/torchref/topology/atom_graph.py index 0e1170e6..7d579567 100644 --- a/torchref/topology/atom_graph.py +++ b/torchref/topology/atom_graph.py @@ -35,7 +35,9 @@ def _build_csr(bonds: torch.Tensor, n_atoms: int) -> Tuple[torch.Tensor, torch.T ------- indptr, indices : torch.Tensor ``indices[indptr[i]:indptr[i + 1]]`` are atom ``i``'s bonded neighbours, - ascending. Each bond contributes both directions. + ascending, each partner listed once. A bond row repeated in the edge list -- + once per altloc conformer for a bond between two shared atoms, or from a LINK + record that appears twice -- therefore does not inflate an atom's degree. """ device = bonds.device if bonds.numel() == 0: @@ -47,12 +49,10 @@ def _build_csr(bonds: torch.Tensor, n_atoms: int) -> Tuple[torch.Tensor, torch.T src = torch.cat([bonds[:, 0], bonds[:, 1]]) dst = torch.cat([bonds[:, 1], bonds[:, 0]]) - # Sort by (src, dst). Two stable passes, least significant first, give the same - # order as a lexicographic sort without materialising a composite key. - order = torch.argsort(dst, stable=True) - src, dst = src[order], dst[order] - order = torch.argsort(src, stable=True) - src, dst = src[order], dst[order] + # Unique directed pairs, which ``torch.unique`` returns in lexicographic + # (src, dst) order -- the CSR layout wanted below. + pairs = torch.unique(torch.stack([src, dst], dim=1), dim=0) + src, dst = pairs[:, 0], pairs[:, 1] counts = torch.bincount(src, minlength=n_atoms) indptr = torch.zeros(n_atoms + 1, dtype=torch.int64, device=device) # dtype-ok: CSR indptr offset array; int64 index required diff --git a/torchref/topology/build.py b/torchref/topology/build.py index d26631aa..83962ae2 100644 --- a/torchref/topology/build.py +++ b/torchref/topology/build.py @@ -586,7 +586,8 @@ def _link_record_edges( """Bond edges for the accepted ``LINK`` records, and the atom pairs they join. A record duplicating an auto-detected disulfide is dropped, since that link already - contributed its bond, angles and torsions. + contributed its bond, angles and torsions; so is a record repeating an earlier one, + which would otherwise add a second bond edge and a second restraint on the same pair. Returns ------- @@ -634,6 +635,7 @@ def _link_record_edges( pair = (min(idx1, idx2), max(idx1, idx2)) if pair in existing: continue + existing.add(pair) rows.append((idx1, idx2)) length = link["length"] usable = isinstance(length, (int, float)) and length == length and length > 0 diff --git a/torchref/topology/hydrogens.py b/torchref/topology/hydrogens.py index b8f0a84a..616502e4 100644 --- a/torchref/topology/hydrogens.py +++ b/torchref/topology/hydrogens.py @@ -7,10 +7,15 @@ Two things the bond graph decides that a distance criterion previously guessed at: -* **How many hydrogens a parent can carry.** The count is the parent's standard valence - minus the heavy atoms actually bonded to it, taken from the graph. A distance sweep - gets this wrong on a distorted or predicted model, where a bond can fall outside the - window; and it cannot distinguish a real bond from two atoms that merely sit close. +* **How many hydrogens a parent can carry.** The smaller of two budgets: the parent's + standard valence minus the heavy atoms actually bonded to it in the graph, and the + template's own hydrogen count minus every graph bond the template does not know about + (a peptide bond, a LINK record, a metal contact). The first budget handles the + template's own chemistry; the second is what stops an acetyl cap's aldehyde hydrogen + or a metal-bound histidine NE2 hydrogen from being generated when the graph degree + sits below the nominal valence only because a double bond counts as one edge. Both + read the graph rather than a distance sweep, which gets a distorted or predicted model + wrong and cannot tell a bond from two atoms that merely sit close. * **Which hydrogens have a free torsion.** A hydrogen whose parent has exactly one heavy neighbour -- hydroxyl, thiol, amine, methyl -- can rotate about the parent-neighbour axis, and the template's angle for it is arbitrary. Those get scanned; the rest are @@ -21,9 +26,11 @@ from typing import Dict, List, Optional, Tuple import numpy as np +import torch -#: Standard heavy-atom valences, used to cap how many hydrogens a parent may take. The -#: fallback of 4 matches the previous behaviour for elements not listed. +#: Standard heavy-atom valences, one of the two budgets that cap how many hydrogens a +#: parent may take. Elements not listed fall back to 4 and are then bounded only by the +#: template's own hydrogen count. STANDARD_VALENCE = {"C": 4, "N": 3, "O": 2, "S": 2} _DEFAULT_VALENCE = 4 @@ -139,6 +146,7 @@ def _template(cif_dict: Dict, resname: str) -> Optional[Dict]: parent_of: Dict[str, str] = {} ideal_length: Dict[str, float] = {} heavy_adjacency: Dict[str, List[str]] = {} + h_count: Dict[str, int] = {} bonds = component.get("bonds") if bonds is not None and len(bonds) > 0: @@ -154,10 +162,12 @@ def _template(cif_dict: Dict, resname: str) -> Optional[Dict]: continue if is_h[ia] and not is_h[ib]: parent_of[a] = b + h_count[b] = h_count.get(b, 0) + 1 if np.isfinite(values[i]): ideal_length[a] = float(values[i]) elif is_h[ib] and not is_h[ia]: parent_of[b] = a + h_count[a] = h_count.get(a, 0) + 1 if np.isfinite(values[i]): ideal_length[b] = float(values[i]) elif not is_h[ia] and not is_h[ib]: @@ -174,6 +184,7 @@ def _template(cif_dict: Dict, resname: str) -> Optional[Dict]: "heavy_coords": coords[~is_h], "h_names": ids[is_h], "parent_of": parent_of, + "h_count": h_count, "ideal_length": ideal_length, "heavy_adjacency": heavy_adjacency, } @@ -254,12 +265,19 @@ def _half_hydrogen_angle(template: Dict, parent_name: str, h_names: List[str]) - return 0.5 * np.arccos(-1.0 / 3.0) -def _split_neighbours(topology, atom_index: int) -> Tuple[np.ndarray, int]: +def _split_neighbours( + topology, atom_index: int, altloc: str = "" +) -> Tuple[np.ndarray, int]: """Heavy neighbour rows of ``atom_index``, and how many hydrogens it already has. Coordinate-independent, unlike a distance sweep: a stretched bond in a predicted or mid-refinement model still counts, and two atoms that merely sit close do not. + Restricted to the conformer being hydrogenated when ``altloc`` is given: a shared + backbone atom is bonded to every altloc copy of a split neighbour, and counting them + all makes a CA with two CB copies look saturated and lose its HA, or an N with two CA + copies lose its H. Blank-altloc neighbours always count. + The hydrogen count is what makes generation idempotent and makes a partially hydrogenated structure top up correctly. Both consume the parent's valence, so subtracting only the heavy neighbours leaves budget for a hydrogen the parent @@ -269,6 +287,12 @@ def _split_neighbours(topology, atom_index: int) -> Tuple[np.ndarray, int]: neighbours = topology.atoms.neighbors(atom_index) if neighbours.numel() == 0: return np.zeros(0, dtype=np.int64), 0 + if altloc: + alts = np.char.strip(topology.atoms.altloc[neighbours.cpu().numpy()].astype(str)) + keep = torch.as_tensor((alts == "") | (alts == altloc), device=neighbours.device) + neighbours = neighbours[keep] + if neighbours.numel() == 0: + return np.zeros(0, dtype=np.int64), 0 is_h = topology.atoms.is_hydrogen[neighbours] return neighbours[~is_h].cpu().numpy(), int(is_h.sum()) @@ -445,13 +469,24 @@ def plan_hydrogens(topology, cif_dict: Dict, xyz, verbose: int = 0) -> HydrogenP parent_row = name_to_row[parent_name] parent_position = coords[parent_row] - heavy_rows, existing_h = _split_neighbours(topology, parent_row) + heavy_rows, existing_h = _split_neighbours( + topology, parent_row, altloc + ) heavy_bonded = len(heavy_rows) element = str( template["elements"][template["id_to_index"][parent_name]] ).upper() valence = STANDARD_VALENCE.get(element, _DEFAULT_VALENCE) - allowed = max(0, valence - heavy_bonded - existing_h) + # Two budgets; see the module docstring. ``extra_bonds`` are graph + # bonds the template has never heard of -- a peptide link, a LINK + # record, a metal -- each of which displaces one template hydrogen. + template_h = template["h_count"].get(parent_name, len(group)) + template_heavy = len(template["heavy_adjacency"].get(parent_name, [])) + extra_bonds = max(0, heavy_bonded - template_heavy) + allowed = max( + 0, + min(valence - heavy_bonded, template_h - extra_bonds) - existing_h, + ) group = group[:allowed] if not group: continue diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index 975a7119..61e99af0 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -324,7 +324,7 @@ def _load_cif_dictionaries(self, cif_path): res for res in self.unique_residues if res not in self.cif_dict ] - if len(self.missing_residues) > 1: + if len(self.missing_residues) >= 1: if self.verbose > 0: print( f"Warning: The following residues are missing from the CIF dictionary " From 1a01aa6d90c3dcb1e98cc39471410e994b0679f4 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 17 Sep 2026 22:22:13 +0200 Subject: [PATCH 169/250] Add refinable riding hydrogens and TorchRef-owned AMBER coordinates --- docs/changelog.rst | 9 + docs/user_guide/cli.rst | 3 + tests/benchmarks/compare_amber_classic.py | 384 ++++++ tests/helpers/device_cases.py | 44 + tests/unit/base/test_local_frame.py | 94 ++ tests/unit/model/test_hydrogen_mode.py | 109 ++ tests/unit/model/test_hydrogens_in_xray.py | 114 ++ tests/unit/model/test_riding_orientations.py | 276 ++++ .../model/test_riding_water_completion.py | 184 +++ tests/unit/model/test_riding_xyz.py | 213 +++ tests/unit/refinement/test_amber_target.py | 411 +++++- tests/unit/topology/test_energy_types.py | 112 ++ tests/unit/topology/test_hydrogen_frames.py | 154 +++ tests/unit/topology/test_hydrogens.py | 35 +- torchref/base/coordinates/__init__.py | 12 + torchref/base/coordinates/local_frame.py | 207 +++ torchref/cli/_common.py | 4 + torchref/cli/collection_difference_refine.py | 4 +- torchref/cli/refine.py | 10 + torchref/data/ener_lib_atoms.csv | 213 +++ .../experimental/ensemble/ensemble_model.py | 2 + .../ensemble/quasi_crystal_amber.py | 188 +-- torchref/experimental/targets/amber_target.py | 1182 ++++------------- torchref/io/cif_readers.py | 9 + torchref/model/context.py | 15 +- torchref/model/model.py | 445 ++++++- torchref/model/model_ft.py | 10 + torchref/model/parameter_wrappers.py | 38 +- torchref/model/riding_xyz.py | 816 ++++++++++++ torchref/refinement/base_refinement.py | 38 + torchref/scaling/solvent.py | 16 + torchref/scripts/extract_ener_lib.py | 97 ++ torchref/topology/atom_graph.py | 51 +- torchref/topology/build.py | 69 +- torchref/topology/builders.py | 8 + torchref/topology/hydrogens.py | 477 ++++++- torchref/topology/monomer/modifications.py | 38 +- torchref/topology/restraints.py | 26 +- 38 files changed, 4872 insertions(+), 1245 deletions(-) create mode 100644 tests/benchmarks/compare_amber_classic.py create mode 100644 tests/unit/base/test_local_frame.py create mode 100644 tests/unit/model/test_hydrogen_mode.py create mode 100644 tests/unit/model/test_hydrogens_in_xray.py create mode 100644 tests/unit/model/test_riding_orientations.py create mode 100644 tests/unit/model/test_riding_water_completion.py create mode 100644 tests/unit/model/test_riding_xyz.py create mode 100644 tests/unit/topology/test_energy_types.py create mode 100644 tests/unit/topology/test_hydrogen_frames.py create mode 100644 torchref/base/coordinates/local_frame.py create mode 100644 torchref/data/ener_lib_atoms.csv create mode 100644 torchref/model/riding_xyz.py create mode 100644 torchref/scripts/extract_ener_lib.py diff --git a/docs/changelog.rst b/docs/changelog.rst index f05a5fe2..dcdb219b 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,15 @@ Changelog Unreleased ---------- +- Add a matched classic/AMBER refinement benchmark with shared prepared hydrogens, initial riding-parameter gradient calibration, and per-cycle R factors and geometry diagnostics. +- Hydrogen generation respects tetrahedral ammonium nitrogen types, retaining the third hydrogen on protonated lysine and free amino termini while peptide nitrogens retain their linked valence; terminal H1 aliases participate in dictionary bonds without renaming model atoms. +- AMBER targets consume every TorchRef-owned atom through a validated atom map, including hydrogen and riding-orientation gradients; incomplete chemistry is rejected at setup, and ensemble AMBER initialization leaves coordinates unchanged unless relaxation is explicitly requested. +- Enabling riding mode completes missing HOH hydrogens only when hydrogen generation is enabled and stripping is disabled; existing atom coordinates and refinement selections are preserved. +- Riding coordinates refine shared methyl/hydroxyl torsions and water rotation vectors, preserve orientations through selection and checkpoints, and expose their gradients to the xyz optimizer; explicit hydrogen generation includes HOH with dictionary geometry and seeded random initial orientations. +- ``hydrogens_in_xray`` (model constructor, ``Refinement``, ``torchref.refine --hydrogens-in-xray/--no-hydrogens-in-xray``, default on) decides whether hydrogens enter the structure factors; it replaces ``exclude_H_from_sf``, which remains as a deprecated inverted alias, and is carried by ``copy``, ``select``, the strip/hydrogenate helpers and state dicts. The bulk-solvent mask is built from heavy atoms only, whatever the model carries. +- ``torchref.base.coordinates.local_frame`` holds the differentiable local-frame placement (``place_local_frame`` and its inverse); the AMBER target's hydrogen placement now delegates to it. ``torchref.topology.hydrogens`` gains ``HydrogenFrames`` and ``hydrogen_frames`` (which heavy atoms each hydrogen rides on, read off the bond graph) and ``augment_atom_table_with_maps``. +- ``RidingXYZTensor`` (``torchref.model.riding_xyz``): a coordinate wrapper over the full atom table whose hydrogen rows are derived from their parent heavy atoms each forward, so only heavy atoms are refined and a force on a hydrogen lands on the atoms that carry it. ``Model.set_hydrogen_mode("riding" | "free")`` and ``Refinement.set_hydrogen_mode`` switch a loaded model; ``ModelContext.hydrogen_mode`` records the choice and it survives ``copy``, ``select``, ``shake_coords``, rigid-body passes and state dicts. ``Restraints.copy`` no longer deep-copies the coordinate, ADP and radius accessors it borrows from the model. +- Monomer templates keep the CCP4 ``type_energy`` of each atom, ``chem_mod_atom`` ``change`` rows retype linked atoms (an in-chain backbone N is ``NH1``, not ``NT3``), and the atom graph carries ``energy_type`` and ``template_h_count`` per atom. The per-type table from ``ener_lib.cif`` ships as ``torchref/data/ener_lib_atoms.csv`` (regenerate with ``python -m torchref.scripts.extract_ener_lib``). - Hydrogen generation on model loading is off by default; use ``torchref.refine --add-hydrogens`` or ``add_hydrogens=True`` in Python to opt in. Hydrogens already present in input files are retained unless ``strip_H=True``. - Hydrogen generation reads the user's restraint CIF: ``cif_path`` is a model constructor argument, set before loading by ``torchref.refine --cif`` and the shared CLI loader, and carried by ``hydrogenate``, ``strip_hydrogens``, ``select`` and state dicts. - mmCIF models carry their covalent and metal ``_struct_conn`` links, as PDB LINK records already did. diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index 0f96ec82..dbe053ed 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -31,6 +31,9 @@ and a ``refinement_history.json`` log. * ``-n`` / ``--n-cycles`` number of macro cycles (default 5) * ``--add-hydrogens`` generate missing hydrogens on model loading (default off). Hydrogens already present in the input are retained with or without this flag +* ``--hydrogens-in-xray`` / ``--no-hydrogens-in-xray`` include hydrogen atoms in the + structure-factor calculation (default on). Off keeps them in the restraints only; + the bulk-solvent mask is built from heavy atoms in either case * ``--mode`` ``separate`` (separated XYZ then ADP, default) or ``everything`` (joint XYZ+ADP) * ``--xray-mode`` one of ``ml`` (default; Read MLF at variance ε·β, conditional diff --git a/tests/benchmarks/compare_amber_classic.py b/tests/benchmarks/compare_amber_classic.py new file mode 100644 index 00000000..fd678d57 --- /dev/null +++ b/tests/benchmarks/compare_amber_classic.py @@ -0,0 +1,384 @@ +"""Compare classic and AMBER restraints on identical prepared benchmark starts. + +Run with the project interpreter and optional OpenMM/PDBFixer dependencies. +The experiment keeps the work/free split, X-ray target, atom set, riding wrapper, +optimizer and ADP settings fixed. AMBER's weight is calibrated once from initial +heavy-coordinate gradient RMS after the riding Jacobian, without consulting R-free. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import time +from pathlib import Path + +import numpy as np +import torch + +from torchref import Model +from torchref.experimental.targets.amber_target import AmberTarget +from torchref.refinement.lbfgs_refinement import LBFGSRefinement + +FILES = Path(__file__).resolve().parents[1] / "files" +IDENTITY = ["chainid", "resseq", "icode", "name"] + + +def _prepare(code: str, output: Path, seed: int) -> dict: + """Complete heavy atoms once, preserve existing rows, and add H in TorchRef.""" + import openmm.app as app + from pdbfixer import PDBFixer + + source = FILES / "pdb" / f"{code}_af.pdb" + original = ( + Model(device="cpu", verbose=0, strip_H=True, add_hydrogens=False) + .load_pdb(str(source)) + .strip_altlocs() + ) + fixer = PDBFixer(filename=str(source)) + fixer.findMissingResidues() + fixer.missingResidues = {} + fixer.findMissingAtoms() + fixer.addMissingAtoms(seed=seed) + fixed_path = output / f"{code}_heavy_completed.pdb" + with fixed_path.open("w") as handle: + app.PDBFile.writeFile(fixer.topology, fixer.positions, handle, keepIds=True) + fixed = Model(device="cpu", verbose=0, add_hydrogens=False).load_pdb( + str(fixed_path) + ) + original_rows = { + tuple(row[k] for k in IDENTITY): row for _, row in original.pdb.iterrows() + } + frame = fixed.pdb.copy() + frame.attrs = original.pdb.attrs.copy() + added = [] + found = set() + for i, row in frame.iterrows(): + key = tuple(row[k] for k in IDENTITY) + if key in original_rows: + found.add(key) + for column in ["x", "y", "z", "tempfactor", "occupancy", "ATOM"]: + frame.at[i, column] = original_rows[key][column] + else: + added.append(key) + same_residue = original.pdb[ + (original.pdb.chainid == row.chainid) + & (original.pdb.resseq == row.resseq) + & (original.pdb.icode == row.icode) + ] + frame.at[i, "tempfactor"] = float(same_residue.tempfactor.mean()) + frame.at[i, "occupancy"] = 1.0 + assert found == set(original_rows), "Heavy-atom preparation dropped original atoms" + model = original._new_model_from_df(frame, strip_H=False) + torch.manual_seed(seed) + model = model.hydrogenate() + model.set_hydrogen_mode("riding") + compatibility = AmberTarget(model=model) + assert compatibility._n_omm_atoms == len(model.pdb) + path = output / f"{code}_prepared.pdb" + model.write_pdb(str(path)) + return { + "code": code, + "source": str(source), + "prepared": str(path), + "added_heavy_atoms": added, + "original_heavy_atoms": len(original.pdb), + "prepared_atoms": len(model.pdb), + "hydrogens": int(model.pdb.element.str.strip().isin(["H", "D"]).sum()), + "sha256": hashlib.sha256(path.read_bytes()).hexdigest(), + } + + +def _magnitude(values: torch.Tensor) -> dict: + """Summarize atom-vector norms, or absolute scalar parameter gradients.""" + values = values.detach().cpu() + norms = values.norm(dim=-1) if values.ndim == 2 else values.abs().flatten() + if not norms.numel(): + return {"n": 0} + return { + "n": norms.numel(), + "rms": float(norms.square().mean().sqrt()), + "median": float(norms.median()), + "p95": float(torch.quantile(norms, 0.95)), + "max": float(norms.max()), + "l2": float(values.norm()), + } + + +def _classic_loss(ref: LBFGSRefinement) -> torch.Tensor: + """Sum the active classic components, excluding the disabled Rama prior.""" + return sum( + target() + for name, target in ref.geometry_target.items() + if name != "ramachandran" + ) + + +def _gradient_snapshot(ref: LBFGSRefinement, amber: AmberTarget) -> tuple[dict, dict]: + """Record raw atomic and model-parameter gradients for both priors and X-ray.""" + # The classic prior is inactive during AMBER refinement, so its contact list + # needs maintenance before evaluating it on the final AMBER coordinates. + for target in ref.geometry_target.values(): + target.maintenance() + model = ref.model + heavy = torch.as_tensor( + ~model.pdb.element.str.strip().isin(["H", "D"]).to_numpy(), device=model.device + ) + tensors = {} + result = {} + for name, function in [ + ("classic_raw", lambda: _classic_loss(ref)), + ("amber_raw", amber.forward), + ("xray", ref.xray_target_work.forward), + ]: + model.reset_cache() + xyz = model.xyz() + leaves = model.xyz.optimization_parameters() + value = function() + gradients = torch.autograd.grad(value, [xyz] + leaves, allow_unused=True) + atomic = gradients[0] + assert atomic is not None and torch.isfinite(atomic).all(), name + tensors[name] = atomic.detach() + result[name] = { + "loss": float(value.detach()), + "heavy": _magnitude(atomic[heavy]), + "hydrogen": _magnitude(atomic[~heavy]), + "all": _magnitude(atomic), + "parameter_gradients": [ + _magnitude(g) if g is not None else {"n": 0} for g in gradients[1:] + ], + } + if name == "amber_raw": + import openmm.unit as unit + + forces = np.asarray( + amber._context.getState(getForces=True) + .getForces(asNumpy=True) + .value_in_unit(unit.kilojoules_per_mole / unit.nanometer) + ) + result[name]["clipped_atom_fraction"] = float( + (np.linalg.norm(forces, axis=1) > 10000).mean() + ) + unclipped = torch.as_tensor( + -forces[amber._model_to_omm] * 0.1 / len(xyz), + dtype=xyz.dtype, + device=xyz.device, + ) + result["amber_unclipped"] = { + "heavy": _magnitude(unclipped[heavy]), + "hydrogen": _magnitude(unclipped[~heavy]), + "all": _magnitude(unclipped), + } + for name in ["classic_raw", "amber_raw"]: + g = tensors[name][heavy].flatten() + x = tensors["xray"][heavy].flatten() + result[name]["heavy_cosine_to_xray"] = float( + torch.nn.functional.cosine_similarity(g, x, dim=0) + ) + c, a = ( + tensors["classic_raw"][heavy].flatten(), + tensors["amber_raw"][heavy].flatten(), + ) + result["classic_amber_heavy_cosine"] = float( + torch.nn.functional.cosine_similarity(c, a, dim=0) + ) + return result, tensors + + +def _metrics(ref: LBFGSRefinement) -> dict: + """Report R factors and dictionary geometry using heavy atoms alone.""" + with torch.no_grad(): + rwork, rfree = ref.get_rfactor() + model = ref.model + restraints = model.restraints + restraints.cat_dict() + heavy = torch.as_tensor( + ~model.pdb.element.str.strip().isin(["H", "D"]).to_numpy(), + device=model.device, + ) + result = {"rwork": float(rwork), "rfree": float(rfree)} + for name, method in [ + ("bond", restraints.bond_deviations), + ("angle", restraints.angle_deviations), + ]: + deviations, sigmas = method() + indices = restraints.restraints[name]["all"]["indices"] + keep = heavy[indices].all(dim=-1) + d, z = deviations[keep], deviations[keep] / sigmas[keep] + if name == "angle": + d = torch.rad2deg(d) + result[f"heavy_{name}_rms_delta"] = float(d.square().mean().sqrt()) + result[f"heavy_{name}_rms_z"] = float(z.square().mean().sqrt()) + return result + + +def _run( + code: str, + prepared: Path, + arm: str, + output: Path, + cycles: int, + seed: int, + weight: float | None, +) -> dict: + """Run classic alternating scaler/XYZ/ADP refinement with one active prior.""" + torch.manual_seed(seed) + started = time.perf_counter() + ref = LBFGSRefinement( + pdb=str(prepared), + data_file=str(FILES / "mtz" / f"{code}.mtz"), + device=torch.device("cpu"), + verbose=0, + add_hydrogens=False, + hydrogens_in_xray=True, + target_mode="ml", + ) + ref.set_hydrogen_mode("riding") + amber = AmberTarget(model=ref.model, normalize_by_atoms=True, verbose=0) + state = ref.loss_state + initial_gradients, raw = _gradient_snapshot(ref, amber) + calibration = ( + 0.2 + * initial_gradients["classic_raw"]["parameter_gradients"][0]["rms"] + / initial_gradients["amber_raw"]["parameter_gradients"][0]["rms"] + ) + if weight is None: + weight = calibration + if arm == "amber": + state.targets = { + key: target + for key, target in state.targets.items() + if not key.startswith("geometry/") + } + state.clear() + state.register_target("amber", amber) + state.set_weight("amber", weight) + state.refresh_loss_leaves() + active_names = [ + key for key in state.targets if state.get_effective_weight(key) != 0 + ] + if arm == "amber": + assert "amber" in active_names and not any( + key.startswith("geometry/") for key in active_names + ) + else: + assert "amber" not in active_names + initial = _metrics(ref) + xyz_hash = hashlib.sha256( + ref.model.xyz().detach().cpu().numpy().tobytes() + ).hexdigest() + result = { + "code": code, + "arm": arm, + "cycles": cycles, + "seed": seed, + "amber_weight": weight, + "calibration_weight": calibration, + "calibration_basis": "heavy_xyz_parameter_gradient_rms_after_riding", + "cartesian_calibration_weight": ( + 0.2 + * initial_gradients["classic_raw"]["heavy"]["rms"] + / initial_gradients["amber_raw"]["heavy"]["rms"] + ), + "classic_group_weight": 0.2, + "initial": initial, + "initial_gradients": initial_gradients, + "initial_xyz_sha256": xyz_hash, + "active_targets": active_names, + "n_atoms": len(ref.model.pdb), + "n_reflections": len(ref.reflection_data.hkl), + "reflection_split_sha256": hashlib.sha256( + ref.reflection_data.hkl.detach().cpu().numpy().tobytes() + + ref.reflection_data.rfree_flags.detach().cpu().numpy().tobytes() + ).hexdigest(), + "parameter_gradient_order": [ + "heavy_xyz_per_angstrom", + "torsion_per_radian", + "water_rotation_per_radian", + ], + "geometry_units": {"bond_rms_delta": "angstrom", "angle_rms_delta": "degree"}, + "setup_seconds": time.perf_counter() - started, + "trajectory": [], + } + target_path = output / f"{code}_{arm}.json" + target_path.write_text(json.dumps(result, indent=2)) + print( + json.dumps( + { + "event": "initial", + "code": code, + "arm": arm, + "metrics": initial, + "weight": weight, + "grad_rms_classic": initial_gradients["classic_raw"]["heavy"]["rms"], + "grad_rms_amber": initial_gradients["amber_raw"]["heavy"]["rms"], + } + ), + flush=True, + ) + for cycle in range(1, cycles + 1): + start_cycle = time.perf_counter() + ref.refine(macro_cycles=1) + entry = { + "cycle": cycle, + **_metrics(ref), + "seconds": time.perf_counter() - start_cycle, + } + result["trajectory"].append(entry) + target_path.write_text(json.dumps(result, indent=2)) + print( + json.dumps({"event": "cycle", "code": code, "arm": arm, **entry}), + flush=True, + ) + result["final"] = _metrics(ref) + result["final_gradients"], _ = _gradient_snapshot(ref, amber) + result["refinement_seconds"] = sum(x["seconds"] for x in result["trajectory"]) + result["total_seconds"] = time.perf_counter() - started + ref.model.write_pdb(str(output / f"{code}_{arm}_refined.pdb")) + target_path.write_text(json.dumps(result, indent=2)) + return result + + +def _main() -> None: + """Prepare one benchmark or run one arm in its own process.""" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("action", choices=["prepare", "classic", "amber"]) + parser.add_argument("--code", required=True) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--cycles", type=int, default=5) + parser.add_argument("--amber-weight", type=float, default=None) + parser.add_argument("--seed", type=int, default=20260917) + args = parser.parse_args() + args.output.mkdir(parents=True, exist_ok=True) + if args.action == "prepare": + details = _prepare(args.code, args.output, args.seed) + (args.output / f"{args.code}_preparation.json").write_text( + json.dumps(details, indent=2) + ) + print(json.dumps(details), flush=True) + else: + weight = args.amber_weight + if args.action == "amber" and weight is None: + baseline = json.loads( + (args.output / f"{args.code}_classic.json").read_text() + ) + gradients = baseline["initial_gradients"] + weight = ( + 0.2 + * gradients["classic_raw"]["parameter_gradients"][0]["rms"] + / gradients["amber_raw"]["parameter_gradients"][0]["rms"] + ) + _run( + args.code, + args.output / f"{args.code}_prepared.pdb", + args.action, + args.output, + args.cycles, + args.seed, + weight, + ) + + +if __name__ == "__main__": + _main() diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index 07ddd9c4..c9858b9b 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -56,6 +56,38 @@ class DeviceCase: ignore: tuple = field(default_factory=tuple) +def _riding_xyz(device): + """A bonded torsion group and a water exercise both orientation buffers.""" + import numpy as np + + from torchref.model.riding_xyz import RidingXYZTensor + from torchref.topology.hydrogens import HydrogenFrames + + xyz = torch.tensor( + [ + [0.0, 0.0, 0.0], + [1.5, 0.0, 0.0], + [1.5, 1.5, 0.0], + [-0.6, 0.8, 0.0], + [-0.6, -0.8, 0.0], + [3.0, 0.0, 0.0], + [3.8, 0.0, 0.0], + [2.8, 0.8, 0.0], + ], + dtype=torch.float32, + ) + frames = HydrogenFrames( + h_row=np.array([3, 4, 6, 7]), + parent_row=np.array([0, 0, 5, 5]), + n1_row=np.array([1, 1, -1, -1]), + n2_row=np.array([2, 2, -1, -1]), + frame_valid=np.array([True, True, False, False]), + torsion_group=np.array([0, 0, -1, -1]), + rotation_group=np.array([-1, -1, 0, 0]), + ) + return RidingXYZTensor(xyz, frames, device=device) + + def _symmetry(device): """A bare Symmetry from an explicit operation list (no space group involved).""" import torch as _torch @@ -255,6 +287,18 @@ def _sffft_with_grid(d): ), "MixedTensor", ), + DeviceCase( + "RidingXYZTensor_empty", + lambda d: __import__( + "torchref.model.riding_xyz", fromlist=["RidingXYZTensor"] + ).RidingXYZTensor(device=d), + "RidingXYZTensor", + ), + DeviceCase( + "RidingXYZTensor_populated", + _riding_xyz, + "RidingXYZTensor", + ), DeviceCase( "PositiveMixedTensor", lambda d: __import__( diff --git a/tests/unit/base/test_local_frame.py b/tests/unit/base/test_local_frame.py new file mode 100644 index 00000000..32e28e84 --- /dev/null +++ b/tests/unit/base/test_local_frame.py @@ -0,0 +1,94 @@ +"""Local-frame placement: exact inverse, exact gradients, and a safe fallback. + +A riding hydrogen is stored as an offset in a frame built from three atoms. What has +to hold is that the offset round-trips through placement exactly, that a force on the +placed point reaches only the three frame atoms, and that a frame the geometry cannot +define is recognised rather than used. +""" + +import pytest +import torch + +from torchref.base.coordinates.local_frame import ( + frame_is_degenerate, + local_frame_coordinates, + place_local_frame, +) + + +def _frames(n: int, dtype=torch.float64): + generator = torch.Generator().manual_seed(7) + p = torch.rand(n, 3, generator=generator, dtype=dtype) * 10 + n1 = p + torch.randn(n, 3, generator=generator, dtype=dtype) + n2 = p + torch.randn(n, 3, generator=generator, dtype=dtype) + point = p + torch.randn(n, 3, generator=generator, dtype=dtype) + return p, n1, n2, point + + +@pytest.mark.unit +def test_local_coordinates_invert_placement(): + """A point expressed in its frame and placed again lands where it started.""" + p, n1, n2, point = _frames(64) + valid = torch.ones(64, dtype=torch.bool) + local = local_frame_coordinates(p, n1, n2, point) + back = place_local_frame(p, n1, n2, local, valid, point - p) + assert torch.allclose(back, point, atol=1e-12) + + +@pytest.mark.unit +def test_offset_is_invariant_under_rigid_motion(): + """Rotating and translating the three frame atoms carries the point along.""" + p, n1, n2, point = _frames(16) + local = local_frame_coordinates(p, n1, n2, point) + angle = torch.tensor(0.7, dtype=torch.float64) + rotation = torch.tensor( + [ + [torch.cos(angle), -torch.sin(angle), 0.0], + [torch.sin(angle), torch.cos(angle), 0.0], + [0.0, 0.0, 1.0], + ], + dtype=torch.float64, + ) + shift = torch.tensor([1.0, -2.0, 3.0], dtype=torch.float64) + moved = [x @ rotation.T + shift for x in (p, n1, n2, point)] + valid = torch.ones(16, dtype=torch.bool) + placed = place_local_frame(moved[0], moved[1], moved[2], local, valid, point - p) + assert torch.allclose(placed, moved[3], atol=1e-12) + + +@pytest.mark.unit +def test_gradients_are_exact_and_reach_only_the_frame_atoms(): + """Autograd through the frame matches finite differences.""" + p, n1, n2, point = _frames(6) + local = local_frame_coordinates(p, n1, n2, point) + valid = torch.ones(6, dtype=torch.bool) + rigid = point - p + + def place(pp, a, b): + return place_local_frame(pp, a, b, local, valid, rigid) + + leaves = tuple(x.clone().requires_grad_() for x in (p, n1, n2)) + assert torch.autograd.gradcheck(place, leaves, eps=1e-6, atol=1e-6) + + +@pytest.mark.unit +def test_invalid_frames_fall_back_to_rigid_translation(): + """Where the frame is flagged invalid the point simply follows its parent.""" + p, n1, n2, point = _frames(8) + local = torch.zeros(8, 3, dtype=torch.float64) + valid = torch.zeros(8, dtype=torch.bool) + rigid = torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64).expand(8, 3) + placed = place_local_frame(p, n1, n2, local, valid, rigid) + assert torch.allclose(placed, p + rigid) + + +@pytest.mark.unit +def test_degenerate_frames_are_detected(): + """Collinear or collapsed reference bonds are flagged, healthy frames are not.""" + p = torch.zeros(3, 3, dtype=torch.float64) + n1 = torch.tensor([[1.0, 0.0, 0.0]] * 3, dtype=torch.float64) + n2 = torch.tensor( + [[0.0, 1.0, 0.0], [2.0, 0.0, 0.0], [1e-5, 0.0, 0.0]], dtype=torch.float64 + ) + flagged = frame_is_degenerate(p, n1, n2) + assert flagged.tolist() == [False, True, True] diff --git a/tests/unit/model/test_hydrogen_mode.py b/tests/unit/model/test_hydrogen_mode.py new file mode 100644 index 00000000..c7676a5a --- /dev/null +++ b/tests/unit/model/test_hydrogen_mode.py @@ -0,0 +1,109 @@ +"""Switching a loaded model between riding and free hydrogens. + +The switch replaces the coordinate wrapper and nothing else: coordinates are +unchanged, the heavy-atom refinable set carries over, hydrogens leave or rejoin the +refinable set, the restraints keep reading the live wrapper, and the mode survives +copies, selections and state dicts. +""" + +import pytest +import torch + +from torchref.model.model import Model +from torchref.model.model_ft import ModelFT +from torchref.model.riding_xyz import RidingXYZTensor + + +@pytest.fixture +def free_model(pdb_dir): + model = Model(verbose=0, add_hydrogens=True) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + return model + + +def _n_h(model): + return int((model.pdb["element"].str.strip() == "H").sum()) + + +@pytest.mark.unit +def test_switch_to_riding_keeps_coordinates_and_drops_hydrogen_parameters(free_model): + model = free_model + before = model.xyz().detach().clone() + n_free = model.parameters_of_types(("xyz",))[0].shape[0] + model.set_hydrogen_mode("riding") + assert model.hydrogen_mode == "riding" + assert isinstance(model.xyz, RidingXYZTensor) + assert torch.allclose(model.xyz(), before, atol=1e-4) + assert model.parameters_of_types(("xyz",))[0].shape[0] == n_free - _n_h(model) + assert model.xyz.n_hydrogens == _n_h(model) + + +@pytest.mark.unit +def test_switch_back_to_free_restores_per_atom_wrapper(free_model): + model = free_model + model.set_hydrogen_mode("riding") + model.set_hydrogen_mode("free") + assert model.hydrogen_mode == "free" + assert not isinstance(model.xyz, RidingXYZTensor) + assert model.parameters_of_types(("xyz",))[0].shape[0] == len(model.pdb) + + +@pytest.mark.unit +def test_restraints_read_the_installed_wrapper(free_model): + model = free_model + restraints = model.restraints + model.set_hydrogen_mode("riding") + assert restraints._xyz_fn is model.xyz + with torch.no_grad(): + model.xyz.refinable_params.add_(0.1) + assert torch.equal(restraints.xyz(), model.xyz()) + + +@pytest.mark.unit +def test_frozen_heavy_atoms_stay_frozen_across_the_switch(free_model): + model = free_model + mask = torch.zeros(len(model.pdb), dtype=torch.bool, device=model.device) + mask[: len(model.pdb) // 2] = True + model.xyz.update_refinable_mask(mask) + model.set_hydrogen_mode("riding") + full = model.xyz.full_refinable_mask + heavy = ~torch.as_tensor((model.pdb["element"].str.strip() == "H").values, device=full.device) + assert torch.equal(full[heavy], mask[heavy]) + + +@pytest.mark.unit +def test_mode_survives_copy_select_and_shake(free_model): + model = free_model + model.set_hydrogen_mode("riding") + dup = model.copy() + assert dup.hydrogen_mode == "riding" and isinstance(dup.xyz, RidingXYZTensor) + assert torch.equal(dup.xyz(), model.xyz()) + sub = model.select("resseq 10:40") + assert isinstance(sub.xyz, RidingXYZTensor) + assert sub.xyz.shape[0] == len(sub.pdb) + model.shake_coords(0.05) + assert isinstance(model.xyz, RidingXYZTensor) + assert model.xyz.shape[0] == len(model.pdb) + + +@pytest.mark.unit +@pytest.mark.parametrize("model_class", [Model, ModelFT]) +def test_riding_mode_round_trips_through_state_dict(pdb_dir, model_class): + kwargs = {"max_res": 3.0} if model_class is ModelFT else {} + model = model_class(verbose=0, add_hydrogens=True, **kwargs) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + model.set_hydrogen_mode("riding") + with torch.no_grad(): + model.xyz.refinable_params[0].add_(0.25) + state = model.state_dict() + restored = model_class.create_from_state_dict(state, device=model.device) + assert restored.hydrogen_mode == "riding" + assert isinstance(restored.xyz, RidingXYZTensor) + assert torch.allclose(restored.xyz(), model.xyz(), atol=1e-5) + assert restored.xyz.get_refinable_count() == model.xyz.get_refinable_count() + + +@pytest.mark.unit +def test_none_mode_is_refused(free_model): + with pytest.raises(ValueError): + free_model.set_hydrogen_mode("none") diff --git a/tests/unit/model/test_hydrogens_in_xray.py b/tests/unit/model/test_hydrogens_in_xray.py new file mode 100644 index 00000000..ff709a7d --- /dev/null +++ b/tests/unit/model/test_hydrogens_in_xray.py @@ -0,0 +1,114 @@ +"""The ``hydrogens_in_xray`` flag: what it gates, and where it has to survive. + +It decides only whether hydrogen rows reach the structure-factor gathers. Restraints +never consult it, the solvent mask never includes hydrogens, and the setting must +follow the model through copies, selections, strips and state dicts. +""" + +import numpy as np +import pytest +import torch + +from torchref.model.context import ModelContext +from torchref.model.model import Model +from torchref.model.model_ft import ModelFT + + +@pytest.fixture(scope="module") +def with_hydrogens(pdb_dir): + """7L84 keeps its deposited hydrogens.""" + model = Model(verbose=0) + model.load_pdb(str(pdb_dir / "7L84.pdb")) + return model + + +def _n_h(model): + return int((model.pdb["element"].str.strip() == "H").sum()) + + +@pytest.mark.unit +def test_default_is_on(): + assert ModelContext().hydrogens_in_xray is True + assert Model(verbose=0).hydrogens_in_xray is True + + +@pytest.mark.unit +def test_partition_covers_hydrogens_only_when_on(with_hydrogens): + """Off drops exactly the hydrogen rows from the isotropic gather.""" + model = with_hydrogens + n_atoms, n_h = len(model.pdb), _n_h(model) + assert n_h > 0 + + def n_in_fcalc(): + return model.get_iso()[0].shape[0] + model.get_aniso()[0].shape[0] + + model.hydrogens_in_xray = True + assert n_in_fcalc() == n_atoms + model.hydrogens_in_xray = False + assert n_in_fcalc() == n_atoms - n_h + model.hydrogens_in_xray = True + assert n_in_fcalc() == n_atoms + + +@pytest.mark.unit +def test_off_matches_a_stripped_model(pdb_dir): + """Excluding hydrogens from Fcalc equals computing Fcalc without them.""" + full = ModelFT(verbose=0, max_res=2.5, hydrogens_in_xray=False) + full.load_pdb(str(pdb_dir / "7L84.pdb")) + heavy = ModelFT(verbose=0, max_res=2.5, strip_H=True) + heavy.load_pdb(str(pdb_dir / "7L84.pdb")) + grid = torch.arange(-3, 4) + hkl = torch.cartesian_prod(grid, grid, grid) + hkl = hkl[(hkl != 0).any(dim=1)].to(full.device) + with torch.no_grad(): + f_full = full(hkl) + f_heavy = heavy(hkl.to(heavy.device)) + assert torch.allclose(f_full, f_heavy, rtol=1e-4, atol=1e-3) + + +@pytest.mark.unit +def test_setting_survives_copy_select_and_strip(with_hydrogens): + model = with_hydrogens.copy() + model.hydrogens_in_xray = False + assert model.copy().hydrogens_in_xray is False + assert model.select("chain A").hydrogens_in_xray is False + assert model.strip_hydrogens().hydrogens_in_xray is False + + +@pytest.mark.unit +@pytest.mark.parametrize("model_class", [Model, ModelFT]) +def test_setting_round_trips_through_state_dict(pdb_dir, model_class): + kwargs = {"max_res": 3.0} if model_class is ModelFT else {} + model = model_class(verbose=0, hydrogens_in_xray=False, **kwargs) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + state = model.state_dict() + restored = model_class.create_from_state_dict(state, device=model.device) + assert restored.hydrogens_in_xray is False + state.pop("hydrogens_in_xray", None) + legacy = model_class.create_from_state_dict(state, device=model.device) + assert legacy.hydrogens_in_xray is True + + +@pytest.mark.unit +def test_deprecated_alias_is_inverted_and_warns(): + model = Model(verbose=0) + with pytest.warns(DeprecationWarning): + model.exclude_H_from_sf = True + assert model.hydrogens_in_xray is False + with pytest.warns(DeprecationWarning): + assert model.exclude_H_from_sf is True + + +@pytest.mark.unit +def test_solvent_mask_ignores_hydrogens(pdb_dir): + """The bulk-solvent mask is the same with and without hydrogen rows.""" + from torchref.scaling.solvent import SolventModel + + full = ModelFT(verbose=0, max_res=2.5) + full.load_pdb(str(pdb_dir / "7L84.pdb")) + heavy = ModelFT(verbose=0, max_res=2.5, strip_H=True) + heavy.load_pdb(str(pdb_dir / "7L84.pdb")) + assert _n_h(full) > 0 and _n_h(heavy) == 0 + mask_full = SolventModel(full, verbose=0).get_solvent_mask() + mask_heavy = SolventModel(heavy, verbose=0).get_solvent_mask() + assert torch.equal(mask_full, mask_heavy) diff --git a/tests/unit/model/test_riding_orientations.py b/tests/unit/model/test_riding_orientations.py new file mode 100644 index 00000000..633af89f --- /dev/null +++ b/tests/unit/model/test_riding_orientations.py @@ -0,0 +1,276 @@ +"""Refinable hydrogen orientations preserve geometry and expose force gradients.""" + +import numpy as np +import pytest +import torch +from torchref import Model, ModelFT +from torchref.base.coordinates.local_frame import rotate_vectors +from torchref.config import get_float_dtype +from torchref.model.riding_xyz import RidingXYZTensor +from torchref.topology.hydrogens import _place_group, _template + + +@pytest.fixture(scope="module") +def oriented_model(pdb_dir): + """A deposited protein with explicitly generated protein and water hydrogens.""" + with torch.random.fork_rng(): + torch.manual_seed(19) + model = Model(verbose=0, device="cpu", add_hydrogens=True) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + model.set_hydrogen_mode("riding") + return model + + +@pytest.fixture +def two_groups(oriented_model): + """One methyl and one water, retaining the methyl's heavy frame atoms.""" + full = oriented_model.xyz + methyl = next( + int(g) + for g in torch.unique(full.torsion_group).tolist() + if g >= 0 + and int((full.torsion_group == g).sum()) == 3 + and oriented_model.pdb.iloc[ + int(full.parent_row[full.torsion_group == g][0]) + ].element.strip() + == "C" + ) + water = int(full.rotation_group[full.rotation_group >= 0][0]) + selected = (full.torsion_group == methyl) | (full.rotation_group == water) + keep = torch.zeros(full.shape[0], dtype=torch.bool) + for rows in (full.h_row, full.parent_row, full.n1_row, full.n2_row): + chosen = rows[selected] + keep[chosen[chosen >= 0]] = True + return full.select_rows(keep) + + +@pytest.mark.unit +def test_orientation_groups_preserve_initial_positions(oriented_model): + """Zero rotations reproduce the deposited/generated table without atom changes.""" + model = oriented_model + xyz = torch.as_tensor(model.pdb[["x", "y", "z"]].values, dtype=model.dtype_float) + assert torch.allclose(model.xyz(), xyz, atol=1e-4) + assert model.xyz.torsions.shape[0] > 0 + waters = model.pdb[model.pdb.resname.str.strip() == "HOH"] + assert len(waters) > 0 + assert model.xyz.rotations.shape == ( + int((waters.element.str.strip() == "O").sum()), + 3, + ) + + +@pytest.mark.unit +def test_water_initialization_is_seeded_and_preserves_existing_direction( + oriented_model, +): + """Water references are reproducible and completing an O–H pair preserves its angle.""" + model = oriented_model + xyz = model.xyz().detach().numpy() + first = int(model.xyz._rotation_first[0]) + parent = int(model.xyz.parent_row[first]) + h_rows = model.xyz.h_row[model.xyz.rotation_group == 0].numpy() + names = model.pdb.name.str.strip().to_numpy() + template = _template(model.restraints.cif_dict, "HOH") + h_names = list(names[h_rows]) + lengths = np.linalg.norm(xyz[h_rows] - xyz[parent], axis=1) + + def place(seed, present): + with torch.random.fork_rng(): + torch.manual_seed(seed) + return _place_group( + template, + names[parent], + xyz[parent], + np.empty((0, 3)), + 0, + h_names, + lengths, + present, + xyz, + [], + ) + + first_reference = place(1, {names[parent]: parent}) + assert np.array_equal(first_reference, place(1, {names[parent]: parent})) + assert not np.allclose(first_reference, place(2, {names[parent]: parent})) + completed = place(3, {names[parent]: parent, h_names[0]: int(h_rows[0])}) + assert np.allclose(completed[0], xyz[h_rows[0]], atol=1e-5) + assert np.allclose( + np.linalg.norm(completed[0] - completed[1]), + np.linalg.norm(first_reference[0] - first_reference[1]), + atol=1e-5, + ) + + +@pytest.mark.unit +def test_rotation_preserves_water_and_methyl_geometry(two_groups): + """Shared rotations preserve internal distances and methyl bond angles.""" + w = two_groups + before = w().detach() + with torch.no_grad(): + w.torsions.refinable_params.fill_(0.8) + w.rotations.refinable_params.copy_( + w.rotations.refinable_params.new_tensor([[0.4, -0.2, 0.7]]) + ) + after = w() + assert torch.equal(before[w.base_row], after[w.base_row]) + assert not torch.allclose(before[w.h_row], after[w.h_row]) + for group in (w.torsion_group, w.rotation_group): + rows = w.h_row[group >= 0] + parent = w.parent_row[group >= 0][0:1] + rows = torch.cat((parent, rows)) + assert torch.allclose( + torch.cdist(before[rows], before[rows]), + torch.cdist(after[rows], after[rows]), + atol=1e-5, + ) + rows = w.h_row[w.torsion_group >= 0] + neighbour = w.n1_row[w.torsion_group >= 0] + assert torch.allclose( + (before[rows] - before[neighbour]).norm(dim=-1), + (after[rows] - after[neighbour]).norm(dim=-1), + atol=1e-5, + ) + + +@pytest.mark.unit +@pytest.mark.parametrize("angle", [0.0, 0.4]) +def test_orientation_gradients_match_finite_differences(two_groups, double_cpu, angle): + """The deposited methyl/water coordinate map has correct orientation gradients.""" + w = two_groups.to(dtype=get_float_dtype()) + base = w._storage_values().detach().requires_grad_() + torsion = w.torsions().detach().fill_(angle).requires_grad_() + rotation = w.rotations().detach().fill_(angle).requires_grad_() + assert torch.autograd.gradcheck(w.evaluate, (base, torsion, rotation), atol=2e-6) + vectors = w.rigid_offset[w.rotation_group >= 0].detach() + r = vectors.new_full(vectors.shape, angle, requires_grad=True) + assert torch.autograd.gradgradcheck(lambda x: rotate_vectors(vectors, x), (r,)) + + +@pytest.mark.unit +def test_xyz_optimizer_receives_and_updates_orientations(two_groups): + """An orientation-only selection delivers both leaves to the xyz optimizer.""" + w = two_groups + w.fix_all() + w.refine(w.h_row) + model = Model(device="cpu", verbose=0) + model.xyz = w + leaves = model.parameters_of_types(("xyz",)) + assert {id(p) for p in leaves} == {id(p) for p in w.optimization_parameters()} + target = w.evaluate( + w._storage_values(), w.torsions().detach() + 0.2, w.rotations().detach() + 0.2 + ).detach() + optimizer = torch.optim.SGD(leaves, lr=0.05) + before = w._storage_values().detach().clone() + losses = [] + for _ in range(8): + optimizer.zero_grad() + loss = (w() - target).square().sum() + losses.append(float(loss.detach())) + loss.backward() + optimizer.step() + assert losses[-1] < losses[0] / 2 + assert torch.equal(w._storage_values(), before) + assert w.torsions.refinable_params.abs().sum() > 0 + assert w.rotations.refinable_params.abs().sum() > 0 + + +@pytest.mark.unit +def test_orientation_mutation_invalidates_forward_cache(two_groups): + """Both orientation leaves participate in the coordinate cache fingerprint.""" + w = two_groups + before = w() + with torch.no_grad(): + w.torsions.refinable_params.add_(0.1) + changed = w() + assert changed is not before + with torch.no_grad(): + w.rotations.refinable_params.add_(0.1) + assert w() is not changed + + +@pytest.mark.unit +def test_copy_selection_and_checkpoint_preserve_orientations(two_groups): + """Nonzero rotations survive copy, subset, and empty-shell state restoration.""" + w = two_groups + with torch.no_grad(): + w.torsions.refinable_params.fill_(0.3) + w.rotations.refinable_params.fill_(-0.2) + expected = w().detach() + copied = w.copy() + assert torch.equal(copied(), expected) + assert torch.equal(copied.rotations(), w.rotations()) + keep = torch.ones(w.shape[0], dtype=torch.bool) + keep[w.h_row[w.torsion_group >= 0]] = False + subset = w.select_rows(keep) + assert torch.allclose(subset(), expected[keep], atol=1e-5) + restored = RidingXYZTensor(device="cpu") + restored.load_state_dict(w.state_dict()) + assert torch.equal(restored(), expected) + with torch.no_grad(): + copied.rotations.refinable_params.add_(0.1) + assert torch.equal(w(), expected) + + +@pytest.mark.unit +def test_checkpoint_without_orientation_metadata(two_groups): + """Coordinate-only checkpoints load with fixed hydrogen orientations.""" + state = { + name: value + for name, value in two_groups.state_dict().items() + if name not in ("torsion_group", "rotation_group", "virtual_reference") + and not name.startswith(("torsions.", "rotations.")) + } + restored = RidingXYZTensor(device="cpu") + restored.load_state_dict(state) + assert torch.allclose(restored(), two_groups(), atol=1e-5) + assert restored.torsions.shape == (0,) + assert restored.rotations.shape == (0, 3) + + +@pytest.mark.unit +@pytest.mark.parametrize("model_class", [Model, ModelFT]) +def test_model_checkpoint_restores_rotated_groups(oriented_model, model_class): + """Model checkpoints restore orientation values, masks, and frame references.""" + source = oriented_model.copy() + with torch.no_grad(): + source.xyz.torsions.refinable_params.fill_(0.3) + source.xyz.rotations.refinable_params.fill_(0.2) + source.xyz.rotations.fix(torch.arange(0, source.xyz.rotations.shape[0], 2)) + state = source.state_dict() + restored = model_class.create_from_state_dict(state, device="cpu") + assert torch.allclose(restored.xyz(), source.xyz(), atol=1e-5) + assert torch.equal( + restored.xyz.rotations.refinable_mask, source.xyz.rotations.refinable_mask + ) + + +@pytest.mark.unit +def test_torsion_without_second_reference_still_rotates(two_groups): + """A fixed reference direction completes a methyl frame with one heavy bond.""" + frames = two_groups.hydrogen_frames() + frames.n2_row[frames.torsion_group >= 0] = -1 + frames.frame_valid[frames.torsion_group >= 0] = False + w = RidingXYZTensor(two_groups().detach(), frames) + initial = w().detach() + with torch.no_grad(): + w.torsions.refinable_params.fill_(0.6) + rows = w.h_row[w.torsion_group >= 0] + assert not torch.allclose(w()[rows], initial[rows]) + assert torch.allclose( + torch.cdist(w()[rows], w()[rows]), + torch.cdist(initial[rows], initial[rows]), + atol=1e-5, + ) + + +@pytest.mark.unit +def test_orientation_forward_runs_on_requested_device(two_groups, any_device): + """Coordinate and orientation gradients stay on the requested backend.""" + w = two_groups.to(any_device) + output = w() + output.square().sum().backward() + assert output.device == any_device + for p in w.optimization_parameters(): + assert p.device == any_device + assert p.grad is not None and torch.isfinite(p.grad).all() diff --git a/tests/unit/model/test_riding_water_completion.py b/tests/unit/model/test_riding_water_completion.py new file mode 100644 index 00000000..418b75af --- /dev/null +++ b/tests/unit/model/test_riding_water_completion.py @@ -0,0 +1,184 @@ +"""Riding-mode water completion respects the model's hydrogen generation setting.""" + +import pytest +import torch + +from torchref import Model, ModelFT + + +@pytest.fixture(scope="module") +def heavy_model(pdb_dir): + """Deposited 1DAW loaded without hydrogen generation.""" + return Model(device="cpu", verbose=0, strip_H=True, add_hydrogens=False).load_pdb( + str(pdb_dir / "1DAW.pdb") + ) + + +@pytest.mark.unit +@pytest.mark.parametrize("model_class", [Model, ModelFT]) +def test_riding_completes_only_waters_and_preserves_live_atoms( + heavy_model, model_class +): + """Water completion preserves current coordinates, ADPs, selections and links.""" + model = model_class(device="cpu", verbose=0, add_hydrogens=False, strip_H=False) + model.load( + lambda: (heavy_model.pdb.copy(), heavy_model.cell.data, heavy_model.spacegroup) + ) + with torch.no_grad(): + model.xyz.refinable_params.add_(0.25) + adp = model.adp().detach().clone() + mask = torch.arange(len(model.pdb)) % 2 == 0 + model.xyz.update_refinable_mask(mask) + model.adp.update_refinable_mask(mask) + model.occupancy.freeze_all() + before = model.xyz().detach().clone() + links = model.ctx.links + hkl = torch.tensor([[1, 0, 0], [0, 1, 0], [1, 1, 1]]) + if isinstance(model, ModelFT): + model(hkl) + model.ctx.add_hydrogens = True + returned = model.set_hydrogen_mode("riding") + assert returned is model + is_h = torch.as_tensor(model.pdb.element.str.strip().eq("H").to_numpy()) + is_water = model.pdb.resname.str.strip().eq("HOH") + assert int(is_h.sum()) == 2 * int( + heavy_model.pdb.resname.str.strip().eq("HOH").sum() + ) + assert is_water[is_h.numpy()].all() + assert torch.allclose(model.xyz()[~is_h], before, atol=1e-5) + assert torch.allclose(model.adp()[~is_h], adp) + assert torch.equal(model.xyz.full_refinable_mask[~is_h], mask) + assert torch.equal(model.adp.refinable_mask[~is_h], mask) + assert not model.occupancy.get_refinable_atoms().any() + assert model.ctx.links is links + assert model.xyz.rotations.shape[0] == int(is_h.sum()) // 2 + assert model.restraints.xyz().shape == model.xyz.shape + if isinstance(model, ModelFT): + assert torch.isfinite(model(hkl)).all() + restored = model_class.create_from_state_dict(model.state_dict(), device="cpu") + assert torch.allclose(restored.xyz(), model.xyz(), atol=1e-5) + + +@pytest.mark.unit +def test_partial_water_is_completed_without_moving_existing_hydrogen(heavy_model): + """A deposited oxygen with one supplied hydrogen receives just its missing partner.""" + water = ( + heavy_model.pdb[heavy_model.pdb.resname.str.strip().eq("HOH")].iloc[:1].copy() + ) + model = heavy_model._new_model_from_df(water, strip_H=False) + model.ctx.add_hydrogens = True + model.set_hydrogen_mode("riding") + model.update_pdb() + partial = model._new_model_from_df(model.pdb.iloc[:2].copy(), strip_H=False) + before = partial.xyz().detach().clone() + partial.set_hydrogen_mode("riding") + assert len(partial.pdb) == 2 + assert torch.allclose(partial.xyz(), before, atol=1e-5) + partial.ctx.add_hydrogens = True + partial.set_hydrogen_mode("riding") + assert len(partial.pdb) == 3 + assert torch.allclose(partial.xyz()[:2], before, atol=1e-5) + with torch.no_grad(): + partial.xyz.rotations.refinable_params.fill_(0.3) + wrapper = partial.xyz + coords = partial.xyz().detach().clone() + partial.set_hydrogen_mode("riding") + assert partial.xyz is wrapper + assert torch.equal(partial.xyz(), coords) + partial.set_hydrogen_mode("free").set_hydrogen_mode("riding") + assert len(partial.pdb) == 3 + assert torch.allclose(partial.xyz(), coords, atol=1e-5) + + +@pytest.mark.unit +def test_supplied_frames_are_remapped_when_waters_are_completed(heavy_model): + """Explicit frames for the original atom table coexist with generated water frames.""" + water = ( + heavy_model.pdb[heavy_model.pdb.resname.str.strip().eq("HOH")].iloc[:2].copy() + ) + model = heavy_model._new_model_from_df(water, strip_H=False) + frames = model.hydrogen_frames() + model.ctx.add_hydrogens = True + model.set_hydrogen_mode("riding", frames=frames) + assert len(model.pdb) == 6 + assert model.xyz.n_hydrogens == 4 + assert model.xyz.rotations.shape == (2, 3) + + +@pytest.mark.unit +def test_water_completion_preserves_adp_field(heavy_model): + """Adding water hydrogens retains the node parametrization and its parameters.""" + model = heavy_model._new_model_from_df(heavy_model.pdb.copy(), strip_H=False) + model.set_adp_mode("field", n_nodes=8, k_neighbors=4) + field = model.adp + values = field().detach().clone() + parameters = field.refinable_params + model.ctx.add_hydrogens = True + model.set_hydrogen_mode("riding") + heavy = torch.as_tensor(~model.pdb.element.str.strip().eq("H").to_numpy()) + assert model.adp is field + assert field.refinable_params is parameters + assert torch.allclose(field()[heavy], values, atol=1e-5) + assert field().shape == (len(model.pdb),) + + +@pytest.mark.integration +@pytest.mark.parametrize("generate", [False, True]) +def test_refinement_targets_follow_completed_atom_table(pdb_dir, mtz_dir, generate): + """Refinement adds water H and refreshes targets only when generation is enabled.""" + from torchref.refinement.base_refinement import Refinement + + refinement = Refinement( + pdb=str(pdb_dir / "1DAW.pdb"), + data_file=str(mtz_dir / "1DAW.mtz"), + device="cpu", + verbose=0, + max_res=3.0, + add_hydrogens=False, + ) + previous = refinement.adp_target + n_atoms = len(refinement.model.pdb) + refinement.model.ctx.add_hydrogens = generate + refinement.set_hydrogen_mode("riding") + assert (refinement.adp_target is not previous) == generate + assert (len(refinement.model.pdb) > n_atoms) == generate + geometry = refinement.geometry_target() + assert torch.isfinite(geometry) + geometry.backward() + if generate: + gradient = refinement.model.xyz.rotations.refinable_params.grad + assert gradient is not None and torch.isfinite(gradient).all() + + +@pytest.mark.unit +@pytest.mark.parametrize("model_class", [Model, ModelFT]) +@pytest.mark.parametrize("explicit_frames", [False, True]) +def test_disabled_generation_keeps_atom_table( + heavy_model, model_class, explicit_frames +): + """Riding mode leaves oxygen-only waters untouched when add_hydrogens is False.""" + model = model_class(device="cpu", verbose=0, add_hydrogens=False) + model.load( + lambda: (heavy_model.pdb.copy(), heavy_model.cell.data, heavy_model.spacegroup) + ) + coordinates = model.xyz().detach().clone() + table = model.pdb.copy() + adp, occupancy = model.adp, model.occupancy + frames = model.hydrogen_frames() if explicit_frames else None + model.set_hydrogen_mode("riding", frames=frames) + assert model.pdb.equals(table) + assert torch.equal(model.xyz(), coordinates) + assert model.adp is adp and model.occupancy is occupancy + assert model.xyz.n_hydrogens == 0 + assert model.xyz.rotations.shape == (0, 3) + + +@pytest.mark.unit +def test_stripping_prevents_water_completion(heavy_model): + """The stripping preference takes precedence even when generation is enabled.""" + model = heavy_model._new_model_from_df(heavy_model.pdb.copy(), strip_H=True) + model.ctx.add_hydrogens = True + n_atoms = len(model.pdb) + model.set_hydrogen_mode("riding") + assert len(model.pdb) == n_atoms + assert model.xyz.n_hydrogens == 0 diff --git a/tests/unit/model/test_riding_xyz.py b/tests/unit/model/test_riding_xyz.py new file mode 100644 index 00000000..ba5b9414 --- /dev/null +++ b/tests/unit/model/test_riding_xyz.py @@ -0,0 +1,213 @@ +"""The riding coordinate wrapper: full-space face, heavy-atom storage. + +Pinned here: the placed hydrogens are reproduced from the stored rows, a force on a +hydrogen reaches only the atoms that carry it, every public method speaks full atom +space while the single refinable leaf holds heavy rows only, and the wrapper survives +the conversions the model needs (to and from a plain per-atom wrapper, subsets, +copies, state dicts). +""" + +import numpy as np +import pytest +import torch + +from torchref.model.model import Model +from torchref.model.parameter_wrappers import MixedTensor +from torchref.model.riding_xyz import RidingXYZTensor +from torchref.topology.hydrogens import HydrogenFrames, hydrogen_frames + + +@pytest.fixture(scope="module") +def hydrogenated(pdb_dir): + """1DAW with generated hydrogens, its frames, and its full coordinate table.""" + model = Model(verbose=0, add_hydrogens=True) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + frames = hydrogen_frames(model.restraints.topology) + return model, frames, model.xyz().detach() + + +def _tolerance(dtype): + return 1e-4 if dtype == torch.float32 else 1e-9 + + +@pytest.mark.unit +def test_reconstructs_placed_hydrogens(hydrogenated): + model, frames, xyz = hydrogenated + riding = RidingXYZTensor(xyz, frames) + out = riding() + assert out.shape == xyz.shape + assert riding.shape == tuple(xyz.shape) + assert float((out - xyz).norm(dim=1).max()) < _tolerance(xyz.dtype) + assert riding.n_hydrogens == frames.n_hydrogens + assert riding.refinable_params.shape[0] == xyz.shape[0] - frames.n_hydrogens + + +@pytest.mark.unit +def test_coordinate_leaf_holds_heavy_rows_with_separate_orientations(hydrogenated): + """Heavy positions and shared orientations have separate optimizer leaves.""" + _, frames, xyz = hydrogenated + riding = RidingXYZTensor(xyz, frames) + leaves = list(riding.parameters()) + assert len(leaves) == 3 + assert leaves[0].shape == (xyz.shape[0] - frames.n_hydrogens, 3) + assert leaves[1] is riding.torsions.refinable_params + assert leaves[2] is riding.rotations.refinable_params + assert riding.full_refinable_mask.sum() == leaves[0].shape[0] + assert not riding.full_refinable_mask[riding.h_row].any() + + +@pytest.mark.unit +def test_gradient_through_hydrogens_lands_on_frame_atoms(hydrogenated): + _, frames, xyz = hydrogenated + riding = RidingXYZTensor(xyz, frames) + out = riding() + k = 5 + h_row = int(riding.h_row[k]) + (grad,) = torch.autograd.grad(out[h_row].sum(), riding.refinable_params) + touched = set(torch.nonzero(grad.abs().sum(1) > 0).flatten().tolist()) + expected = { + int(riding._parent_bidx[k]), + int(riding._n1_bidx[k]), + int(riding._n2_bidx[k]), + } + assert touched == expected + + +@pytest.mark.unit +def test_evaluate_matches_finite_differences(): + """The pure map from stored to full rows has exact gradients (float64).""" + xyz = torch.tensor( + [[0.0, 0.0, 0.0], [1.5, 0.0, 0.0], [1.5, 1.5, 0.1], [-0.6, 0.8, 0.2], [-0.6, -0.8, 0.0]], + dtype=torch.float64, + ) + frames = HydrogenFrames( + h_row=np.array([3, 4]), + parent_row=np.array([0, 0]), + n1_row=np.array([1, 1]), + n2_row=np.array([2, 2]), + frame_valid=np.array([True, True]), + ) + riding = RidingXYZTensor(xyz, frames) + base = riding.refinable_params.detach().clone().requires_grad_() + assert torch.autograd.gradcheck(riding.evaluate, (base,), eps=1e-6, atol=1e-6) + + +@pytest.mark.unit +def test_masks_are_full_space_and_ignore_hydrogen_rows(hydrogenated): + _, frames, xyz = hydrogenated + riding = RidingXYZTensor(xyz, frames) + n = xyz.shape[0] + mask = torch.zeros(n, dtype=torch.bool) + mask[:100] = True + riding.update_refinable_mask(mask) + heavy_first = int((~torch.isin(torch.arange(100), riding.h_row.cpu())).sum()) + assert riding.get_refinable_count() == heavy_first + assert riding.full_refinable_mask.sum() == heavy_first + riding.fix_all() + assert riding.get_refinable_count() == 0 + assert riding.refinable_params.numel() == 0 + riding.refine_all() + assert riding.get_refinable_count() == n - frames.n_hydrogens + riding.fix(torch.arange(10)) + assert riding.full_refinable_mask[:10].sum() == 0 + + +@pytest.mark.unit +def test_frozen_parent_keeps_its_hydrogens_still(hydrogenated): + _, frames, xyz = hydrogenated + riding = RidingXYZTensor(xyz, frames) + riding.fix_all() + before = riding().detach().clone() + with torch.no_grad(): + riding.refinable_params.add_(1.0) # empty leaf: nothing moves + assert torch.equal(riding(), before) + + +@pytest.mark.unit +def test_rigid_motion_preserves_local_offsets(hydrogenated): + """Writing a rotated table keeps every hydrogen riding at the same offset.""" + _, frames, xyz = hydrogenated + riding = RidingXYZTensor(xyz, frames) + offsets = riding.local_offset.clone() + angle = torch.tensor(0.4, dtype=xyz.dtype) + rot = torch.tensor( + [[torch.cos(angle), -torch.sin(angle), 0.0], [torch.sin(angle), torch.cos(angle), 0.0], [0.0, 0.0, 1.0]], + dtype=xyz.dtype, device=xyz.device, + ) + moved = xyz @ rot.T + torch.tensor([3.0, -1.0, 2.0], dtype=xyz.dtype, device=xyz.device) + riding[:] = moved + assert float((riding() - moved).norm(dim=1).max()) < 10 * _tolerance(xyz.dtype) + assert torch.allclose(riding.local_offset, offsets, atol=10 * _tolerance(xyz.dtype)) + + +@pytest.mark.unit +def test_assigning_a_hydrogen_row_becomes_a_new_offset(hydrogenated): + _, frames, xyz = hydrogenated + riding = RidingXYZTensor(xyz, frames) + h = int(riding.h_row[0]) + target = xyz[h] + torch.tensor([0.3, -0.2, 0.1], dtype=xyz.dtype, device=xyz.device) + riding[h] = target + assert float((riding()[h] - target).norm()) < 10 * _tolerance(xyz.dtype) + # Heavy rows untouched. + assert float((riding()[riding.base_row] - xyz[riding.base_row]).norm(dim=1).max()) < 10 * _tolerance(xyz.dtype) + + +@pytest.mark.unit +def test_round_trip_with_plain_wrapper(hydrogenated): + _, frames, xyz = hydrogenated + plain = MixedTensor(xyz.clone(), name="xyz") + riding = RidingXYZTensor.from_mixed_tensor(plain, frames) + back = riding.to_mixed_tensor() + assert isinstance(back, MixedTensor) and not isinstance(back, RidingXYZTensor) + assert float((back() - xyz).norm(dim=1).max()) < _tolerance(xyz.dtype) + assert back.refinable_mask.all() + + +@pytest.mark.unit +def test_select_rows_frees_a_hydrogen_whose_parent_is_cut(hydrogenated): + _, frames, xyz = hydrogenated + riding = RidingXYZTensor(xyz, frames) + keep = torch.ones(xyz.shape[0], dtype=torch.bool) + parent = int(riding.parent_row[0]) + keep[parent] = False + sub = riding.select_rows(keep) + assert sub.shape[0] == xyz.shape[0] - 1 + assert sub.n_hydrogens == frames.n_hydrogens - int((riding.parent_row == parent).sum()) + expected = xyz[keep.to(xyz.device)] + assert float((sub() - expected).norm(dim=1).max()) < _tolerance(xyz.dtype) + + +@pytest.mark.unit +def test_copy_is_independent_and_exact(hydrogenated): + _, frames, xyz = hydrogenated + riding = RidingXYZTensor(xyz, frames) + dup = riding.copy() + assert torch.equal(dup(), riding()) + assert torch.equal(dup.local_offset, riding.local_offset) + with torch.no_grad(): + dup.refinable_params.add_(1.0) + assert not torch.equal(dup(), riding()) + + +@pytest.mark.unit +def test_state_dict_round_trip(hydrogenated): + _, frames, xyz = hydrogenated + riding = RidingXYZTensor(xyz, frames) + state = riding.state_dict() + assert "h_row" in state and "local_offset" in state + # As the model restore does: a placeholder of the right shape, values from the dict. + placeholder = RidingXYZTensor(torch.zeros_like(xyz), frames) + placeholder.load_state_dict(state) + assert placeholder.shape == riding.shape + assert torch.equal(placeholder(), riding()) + + +@pytest.mark.unit +def test_forward_is_cached_until_parameters_move(hydrogenated): + _, frames, xyz = hydrogenated + riding = RidingXYZTensor(xyz, frames) + first = riding() + assert riding() is first + with torch.no_grad(): + riding.refinable_params[0, 0] += 0.5 + assert riding() is not first diff --git a/tests/unit/refinement/test_amber_target.py b/tests/unit/refinement/test_amber_target.py index 343ac601..1b88a35a 100644 --- a/tests/unit/refinement/test_amber_target.py +++ b/tests/unit/refinement/test_amber_target.py @@ -1,79 +1,368 @@ -""" -Unit tests for the single-molecule -:class:`~torchref.experimental.targets.amber_target.AmberTarget`. - -Uses a ligand-free protein (7L84) so the standard OpenMM ``Modeller`` path is -exercised — no antechamber/tleap (AmberTools) required, only OpenMM. This -closes the previous coverage gap where the single-molecule target had no unit -test (only manual scripts), which is how it was able to silently rot while the -ensemble targets were developed. -""" +"""AMBER evaluates the model's complete coordinates and returns atomic gradients.""" import os +import numpy as np +import pandas as pd import pytest import torch +from torchref.config import get_int_dtype +from torchref.experimental.targets.amber_target import AmberTarget, _OpenMMAMBERFunction from torchref.model.model import Model -from torchref.experimental.targets.amber_target import AmberTarget -# Ligand-free protein → standard OpenMM Modeller path; needs OpenMM (+ pdbfixer -# from the same [amber] extra), but no AmberTools. Gated centrally in conftest. pytestmark = pytest.mark.openmm - -# Ligand-free protein → standard Modeller path, no antechamber needed. TEST_PDB = os.path.join( os.path.dirname(__file__), "..", "..", "files", "pdb", "7L84.pdb" ) @pytest.fixture(scope="module") -def heavy_model() -> Model: - """Heavy-atom, single-conformation model (OpenMM adds H internally).""" - return Model(verbose=0, strip_H=True).load_pdb(TEST_PDB).strip_altlocs() +def protein(): + """A deposited, hydrogenated protein with methyl orientation parameters.""" + model = Model(verbose=0, device="cpu", add_hydrogens=True).load_pdb(TEST_PDB) + model = model.strip_altlocs() + model.set_hydrogen_mode("riding") + return model + + +@pytest.fixture(scope="module") +def target(protein): + """Build the expensive AMBER context once per module.""" + return AmberTarget(model=protein, verbose=0) + + +def test_build_preserves_model(protein): + """Target construction neither changes model atoms nor replaces coordinates.""" + before = protein.pdb.copy(deep=True) + xyz = protein.xyz + positions = xyz().detach().clone() + target = AmberTarget(model=protein) + pd.testing.assert_frame_equal(protein.pdb, before) + assert protein.xyz is xyz + assert torch.equal(protein.xyz(), positions) + assert target._n_omm_atoms == target._n_model_atoms == len(protein.pdb) + assert np.array_equal(np.sort(target._model_to_omm), np.arange(len(protein.pdb))) + h1 = int(np.flatnonzero(protein.pdb.name.str.strip().eq("H1"))[0]) + assert h1 in protein.xyz.h_row.tolist() + + +def test_forward_and_orientation_gradients(target, protein): + """AMBER hydrogen forces reach the model's methyl torsion parameters.""" + loss = target.forward() + leaves = protein.xyz.optimization_parameters() + gradients = torch.autograd.grad(loss, leaves, allow_unused=True) + assert loss.shape == () and torch.isfinite(loss) + for leaf, gradient in zip(leaves, gradients): + if leaf.numel(): + assert gradient is not None and torch.isfinite(gradient).all() + torsion_grad = torch.autograd.grad( + target.forward(), protein.xyz.torsions.refinable_params + )[0] + assert torsion_grad.abs().sum() > 0 + + +def test_every_position_and_force_has_model_order(target, protein): + """All H positions are supplied live, with force sign, units and normalization.""" + import openmm.unit as unit + + xyz = protein.xyz().detach().clone().requires_grad_() + h_rows = torch.tensor(np.flatnonzero(protein.pdb.element.str.strip().eq("H"))) + with torch.no_grad(): + xyz[h_rows] += xyz.new_tensor([0.003, -0.002, 0.001]) + energy = target._energy(xyz) + gradient = torch.autograd.grad(energy, xyz)[0] + state = target._context.getState(getPositions=True, getForces=True, getEnergy=True) + pos = np.asarray(state.getPositions(asNumpy=True).value_in_unit(unit.nanometer)) + np.testing.assert_allclose( + pos[target._model_to_omm], xyz.detach().numpy() * 0.1, atol=1e-7 + ) + forces = np.asarray( + state.getForces(asNumpy=True).value_in_unit( + unit.kilojoules_per_mole / unit.nanometer + ) + ) + scale = np.minimum( + 10000 / np.maximum(np.linalg.norm(forces, axis=1, keepdims=True), 1e-10), 1 + ) + expected = ( + -forces[target._model_to_omm] * scale[target._model_to_omm] * 0.1 / len(xyz) + ) + np.testing.assert_allclose(gradient.numpy(), expected, rtol=3e-6, atol=1e-5) + assert gradient[h_rows].abs().sum() > 0 + expected_energy = state.getPotentialEnergy().value_in_unit( + unit.kilojoules_per_mole + ) / len(xyz) + assert energy.item() == pytest.approx(expected_energy, rel=1e-6) + + +def test_permuted_positions_and_gradients(target, protein): + """The cached inverse map handles arbitrary model/OpenMM atom permutations.""" + xyz = protein.xyz().detach().clone().requires_grad_() + reference = target._energy(xyz) + reference_grad = torch.autograd.grad(reference, xyz)[0] + perm = torch.arange(len(xyz) - 1, -1, -1) + inverse = torch.argsort(perm) + original = target._omm_to_model.clone() + try: + target._omm_to_model = inverse[original].to(get_int_dtype()) + permuted = xyz.detach()[perm].requires_grad_() + loss = target._energy(permuted) + grad = torch.autograd.grad(loss, permuted)[0] + assert torch.allclose(loss, reference) + assert torch.allclose(grad, reference_grad[perm]) + finally: + target._omm_to_model = original + + +def test_missing_hydrogens_are_not_added(): + """Disabled generation leaves a heavy-only model untouched on AMBER rejection.""" + model = ( + Model(verbose=0, device="cpu", strip_H=True, add_hydrogens=False) + .load_pdb(TEST_PDB) + .strip_altlocs() + ) + before = model.pdb.copy(deep=True) + wrapper = model.xyz + with pytest.raises(ValueError, match="Prepare missing atoms"): + AmberTarget(model=model) + pd.testing.assert_frame_equal(model.pdb, before) + assert model.xyz is wrapper @pytest.fixture(scope="module") -def target(heavy_model) -> AmberTarget: - """Built once — the OpenMM system construction is the expensive part.""" - return AmberTarget(model=heavy_model, verbose=0) - - -def test_build_populates_state(target, heavy_model): - """Construction builds a usable OpenMM context + atom map + H tables.""" - assert target._context is not None - assert target._system is not None - assert target._n_model_atoms == len(heavy_model.pdb) - assert target._n_omm_atoms >= target._n_model_atoms # H added by Modeller - # For the single molecule the chemistry model IS the model. - assert target._chem_model is target._model - # H-attachment tables were built (this is a protein → has H). - assert target._h_idx is not None and target._h_idx.size > 0 - - -def test_forward_finite_energy(target): - """forward() returns a finite scalar energy.""" - e = target.forward() - assert e.shape == () - assert torch.isfinite(e).item() - - -def test_backward_flows_to_xyz(heavy_model): - """Energy gradient propagates to the model's xyz parameters and is finite.""" - target = AmberTarget(model=heavy_model, verbose=0) - heavy_model.xyz.refinable_params.grad = None - e = target.forward() - e.backward() - g = heavy_model.xyz.refinable_params.grad - assert g is not None - assert g.shape == (len(heavy_model.pdb), 3) - assert torch.isfinite(g).all().item() - assert g.abs().sum().item() > 0.0 # non-trivial forces - - -def test_default_charge_method_is_gas(): - """The unified base defaults to the (robust) Gasteiger charge method.""" - import inspect - - sig = inspect.signature(AmberTarget.__init__) - assert sig.parameters["charge_method"].default == "gas" +def water_target(pdb_dir): + """Two nearby deposited waters with TorchRef-generated rotating hydrogens.""" + model = Model(verbose=0, device="cpu", add_hydrogens=False).load_pdb( + str(pdb_dir / "1DAW.pdb") + ) + waters = model.pdb[model.pdb.resname.str.strip().eq("HOH")].copy() + coords = waters[["x", "y", "z"]].to_numpy() + distances = np.linalg.norm(coords[:, None] - coords[None, :], axis=-1) + np.fill_diagonal(distances, np.inf) + i, j = np.unravel_index(np.argmin(distances), distances.shape) + model = model._new_model_from_df(waters.iloc[sorted([i, j])].copy(), strip_H=False) + model.ctx.add_hydrogens = True + with torch.random.fork_rng(): + torch.manual_seed(42) + model.set_hydrogen_mode("riding") + return AmberTarget(model=model, normalize_by_atoms=False) + + +def test_water_rotation_changes_amber_energy(water_target): + """Rotating water H changes AMBER energy while oxygens remain fixed.""" + target = water_target + model = target._model + rotations = model.xyz.rotations.refinable_params + before = rotations.detach().clone() + positions = model.xyz().detach().clone() + try: + initial_energy = target.forward().detach() + with torch.no_grad(): + rotations[0] += rotations.new_tensor([0.2, -0.3, 0.4]) + energy = target.forward() + grad = torch.autograd.grad(energy, rotations)[0] + assert not torch.allclose(energy, initial_energy) + assert torch.isfinite(grad).all() and grad.abs().sum() > 0 + assert torch.equal( + model.xyz()[model.xyz.base_row], positions[model.xyz.base_row] + ) + finally: + with torch.no_grad(): + rotations.copy_(before) + + +def test_unclipped_water_rotation_derivative(water_target): + """AMBER's force bridge agrees with an orientation finite difference in float32.""" + target = water_target + wrapper = target._model.xyz + base = wrapper._storage_values().detach() + torsions = wrapper.torsions().detach() + rotation = wrapper.rotations().detach().requires_grad_() + + def energy(rot): + xyz = wrapper.evaluate(base, torsions, rot) + return _OpenMMAMBERFunction.apply( + target._compose_full_omm_xyz(xyz), target._context, float("inf") + ) + + gradient = torch.autograd.grad(energy(rotation), rotation)[0] + direction = gradient.detach() / gradient.norm() + step = 1e-3 + finite = ( + energy(rotation + step * direction) - energy(rotation - step * direction) + ) / (2 * step) + assert finite.item() == pytest.approx( + (gradient * direction).sum().item(), rel=0.015, abs=0.1 + ) + + +def test_atom_count_change_requires_rebuild(target, protein): + """A target rejects coordinates from an expanded or reduced atom table.""" + with pytest.raises(ValueError, match="Atom count changed"): + target._energy(protein.xyz()[:-1]) + + +@pytest.mark.parametrize("ligand_instances", [False, True]) +def test_gaff2_mapping_preserves_residue_instances(water_target, ligand_instances): + """Residue and atom permutations cannot exchange two identical molecules.""" + import openmm as mm + import openmm.app as app + + model = water_target._model + source = list(water_target._topology.residues()) + keys = list( + dict.fromkeys( + model.pdb[["chainid", "resseq", "icode"]].itertuples(index=False, name=None) + ) + ) + topology = app.Topology() + chain = topology.addChain("renumbered") + positions = [] + expected = [] + for residue in reversed(source): + dest = topology.addResidue(residue.name, chain, str(len(positions) + 10)) + mapped = {} + for atom in reversed(list(residue.atoms())): + mapped[atom.index] = topology.addAtom(atom.name, atom.element, dest) + row = int(water_target._omm_to_model[atom.index]) + expected.append(row) + positions.append(model.xyz()[row].detach().numpy() * 0.1) + for a, b in water_target._topology.bonds(): + if a.index in mapped and b.index in mapped: + topology.addBond(mapped[a.index], mapped[b.index]) + candidate = AmberTarget() + candidate._chem_model = model + candidate._topology = topology + candidate._system = mm.System() + for _ in expected: + candidate._system.addParticle(1) + candidate._tleap_residue_map = True + candidate._gaff2_residue_keys = list(reversed(keys)) if ligand_instances else [] + candidate._tleap_pos_nm = np.asarray(positions) + if ligand_instances: + candidate._tleap_pos_nm += 1000 + candidate._build_atom_map() + assert candidate._omm_to_model.tolist() == expected + + +def test_duplicate_particle_mapping_is_rejected(water_target): + """A duplicate source index cannot silently send two forces to one atom.""" + candidate = AmberTarget() + candidate._chem_model = water_target._model + candidate._topology = water_target._topology + candidate._system = water_target._system + candidate._source_model_rows = np.zeros(water_target._n_model_atoms, dtype=np.int32) + with pytest.raises(ValueError, match="not one-to-one"): + candidate._build_atom_map() + + +def test_supercell_gather_keeps_live_hydrogens(water_target): + """The ensemble path transfers all transformed hydrogen rows unchanged.""" + from torchref.experimental.ensemble.quasi_crystal_amber import ( + QuasiCrystalAmberTarget, + ) + + candidate = QuasiCrystalAmberTarget.__new__(QuasiCrystalAmberTarget) + torch.nn.Module.__init__(candidate) + n_atoms = water_target._n_model_atoms + candidate._n_members = 2 + candidate._n_model_per_member = n_atoms + candidate._omm_to_model = torch.arange(n_atoms - 1, -1, -1, dtype=get_int_dtype()) + xyz = water_target._model.xyz().detach().unsqueeze(0).repeat(2, 1, 1) * 0.1 + xyz[1] += 0.3 + xyz.requires_grad_() + result = candidate._compose_full_omm_xyz(xyz).reshape(2, n_atoms, 3) + assert torch.equal(result.flip(1), xyz) + weights = torch.arange(result.numel(), dtype=xyz.dtype).reshape_as(result) + gradient = torch.autograd.grad((result * weights).sum(), xyz)[0] + assert torch.equal(gradient, weights.flip(1)) + + +@pytest.mark.parametrize("use_reference_dtype", [False, True]) +def test_bridge_preserves_input_dtype(water_target, use_reference_dtype, double_cpu): + """The OpenMM boundary preserves configured reference and primary dtypes.""" + from torchref.config import get_float_dtype + + dtype = ( + get_float_dtype() if use_reference_dtype else water_target._model.xyz().dtype + ) + xyz = water_target._model.xyz().detach().to(dtype=dtype).requires_grad_() + loss = water_target._energy(xyz) + gradient = torch.autograd.grad(loss, xyz)[0] + assert loss.dtype == gradient.dtype == xyz.dtype + assert torch.isfinite(gradient).all() + + +def test_bridge_follows_coordinate_device(water_target, any_device): + """Coordinates and gradients stay on the caller's device across the CPU bridge.""" + xyz = water_target._model.xyz().detach().to(any_device).requires_grad_() + loss = water_target._energy(xyz) + gradient = torch.autograd.grad(loss, xyz)[0] + assert loss.device == gradient.device == xyz.device + assert torch.isfinite(gradient).all() + + +def test_torchref_hydrogenation_prepares_compatible_protein(): + """TorchRef can prepare all model hydrogens before AMBER construction.""" + model = ( + Model(verbose=0, device="cpu", strip_H=True, add_hydrogens=False) + .load_pdb(TEST_PDB) + .strip_altlocs() + .hydrogenate() + ) + model.set_hydrogen_mode("riding") + positions = model.xyz().detach().clone() + target = AmberTarget(model=model) + assert target._n_omm_atoms == len(model.pdb) + assert torch.equal(model.xyz(), positions) + assert torch.isfinite(target.forward()) + + +def test_context_initialization_uses_live_model(water_target, monkeypatch): + """Backend template coordinates cannot replace the model's live positions.""" + candidate = AmberTarget() + candidate._chem_model = water_target._model + candidate._source_model_rows = water_target._omm_to_model.cpu().numpy() + stale_positions = np.zeros_like(water_target._pos_buf) + monkeypatch.setattr( + candidate, + "_build_omm_system", + lambda params: (water_target._system, water_target._topology, stale_positions), + ) + captured = [] + monkeypatch.setattr( + candidate, "_build_context", lambda positions: captured.append(positions.copy()) + ) + candidate._build() + expected = ( + water_target._compose_full_omm_xyz(water_target._model.xyz()) + .detach() + .cpu() + .numpy() + ) + np.testing.assert_array_equal(captured[0], expected) + np.testing.assert_array_equal(candidate._pos_buf, expected) + + +def test_partial_terminal_hydrogens_preserve_h1_alias(protein): + """Completing an existing terminal H1 adds H2/H3 without an equivalent H.""" + pdb = protein.pdb + first = pdb.iloc[0] + residue = ( + (pdb.chainid == first.chainid) + & (pdb.resseq == first.resseq) + & (pdb.icode == first.icode) + ) + missing = residue & pdb.name.str.strip().isin(["H2", "H3"]) + partial = protein._new_model_from_df(pdb.loc[~missing].copy(), strip_H=False) + prepared = partial.hydrogenate() + first_residue = prepared.pdb[ + (prepared.pdb.chainid == first.chainid) & (prepared.pdb.resseq == first.resseq) + ] + names = set(first_residue.name.str.strip()) + assert {"H1", "H2", "H3"} <= names + assert "H" not in names + assert len(prepared.pdb) == len(protein.pdb) + target = AmberTarget(model=prepared) + assert torch.isfinite(target.forward()) diff --git a/tests/unit/topology/test_energy_types.py b/tests/unit/topology/test_energy_types.py new file mode 100644 index 00000000..d296a038 --- /dev/null +++ b/tests/unit/topology/test_energy_types.py @@ -0,0 +1,112 @@ +"""CCP4 energy types travel from the monomer template onto the atom graph. + +What is pinned: the type column survives the CIF reader, link modifications retype the +atoms they change (an in-chain backbone nitrogen is an amide ``NH1``, only the +N-terminus keeps the free-amine ``NT3``), the template hydrogen count lands on every +template atom, and the bundled per-type table covers every type the standard residues +use. +""" + +import csv + +import numpy as np +import pytest + +from torchref import PATH_TORCHREF_DATA +from torchref.model.model import Model +from torchref.topology.hydrogens import template_atom_types + + +@pytest.fixture(scope="module") +def heavy_1daw(pdb_dir): + model = Model(verbose=0, strip_H=True, add_hydrogens=False) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + return model + + +@pytest.mark.unit +def test_reader_keeps_type_energy(heavy_1daw): + atoms = heavy_1daw.restraints.cif_dict["ASN"]["atoms"] + types = dict(zip(atoms["atom_id"].str.strip(), atoms["type_energy"])) + assert types["N"] == "NT3" + assert types["ND2"] == "NH2" + assert types["OD1"] == "O" + assert types["CA"] == "CH1" + + +@pytest.mark.unit +def test_template_atom_types_counts_hydrogens(heavy_1daw): + types, h_count = template_atom_types(heavy_1daw.restraints.cif_dict["ASN"]) + assert types["OXT"] == "OC" + assert h_count["ND2"] == 2 + assert h_count["CB"] == 2 + assert h_count["N"] == 3 + assert "OD1" not in h_count + + +@pytest.mark.unit +def test_atom_graph_carries_types_with_link_modifications(heavy_1daw): + """Peptide-linked backbone N is retyped NH1; the chain start stays NT3.""" + atoms = heavy_1daw.restraints.topology.atoms + names = atoms.name.astype(str) + resseq = heavy_1daw.pdb["resseq"].values + chain = heavy_1daw.pdb["chainid"].astype(str).values + resnames = heavy_1daw.pdb["resname"].str.strip().values + is_n = names == "N" + first_res = min(resseq[chain == chain[0]]) + n_types = atoms.energy_type[is_n] + n_first = atoms.energy_type[is_n & (resseq == first_res) & (chain == chain[0])] + assert set(n_first.tolist()) == {"NT3"} + # In-chain amide N is NH1; proline's tertiary N is NH0; only chain starts stay NT3. + assert set(n_types.tolist()) <= {"NH1", "NH0", "NT3"} + assert (n_types == "NH1").sum() > 0.8 * is_n.sum() + assert (atoms.energy_type[is_n & (resnames == "PRO")] == "NH0").all() + assert (atoms.energy_type[(names == "O") & (resnames != "HOH")] == "O").all() + assert (atoms.energy_type[(names == "CA") & (resnames != "GLY")] == "CH1").all() + assert (atoms.energy_type[(names == "CA") & (resnames == "GLY")] == "CH2").all() + + +@pytest.mark.unit +def test_template_h_count_and_implicit_hydrogens(heavy_1daw): + """A heavy-only model is missing exactly the hydrogens its templates carry.""" + atoms = heavy_1daw.restraints.topology.atoms + names = atoms.name.astype(str) + polymer = heavy_1daw.pdb["ATOM"].astype(str).str.strip().values == "ATOM" + single_atom = np.isin(heavy_1daw.pdb["resname"].str.strip().values, ["HOH", "MG"]) + counts = atoms.template_h_count.cpu().numpy() + assert (counts[names == "CB"] >= 1).all() + assert (counts[(names == "O") & polymer] == 0).all() + missing = atoms.implicit_h_count().cpu().numpy() + np.testing.assert_array_equal(missing[counts >= 0], counts[counts >= 0]) + # Waters and ions have no template, so their count is unknown; every polymer + # atom's is known. + assert (counts[polymer] >= 0).all() + assert (counts[single_atom] == -1).all() + + +@pytest.mark.unit +def test_hydrogenated_model_completes_charged_amines(pdb_dir): + """Explicit hydrogen generation fills the polymer's template hydrogen counts.""" + model = Model(verbose=0, add_hydrogens=True) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + atoms = model.restraints.topology.atoms + missing = atoms.implicit_h_count().cpu().numpy() + polymer = model.pdb["ATOM"].astype(str).str.strip().values == "ATOM" + assert (missing[polymer] == 0).all() + ammonium = np.isin(atoms.energy_type, ["NT", "NT1", "NT2", "NT3", "NT4"]) + assert ammonium.any() + assert (missing[ammonium & polymer] == 0).all() + + +@pytest.mark.unit +def test_bundled_table_covers_standard_residue_types(heavy_1daw): + with open(f"{PATH_TORCHREF_DATA}/ener_lib_atoms.csv") as handle: + rows = [r for r in csv.DictReader(l for l in handle if not l.startswith("#"))] + table = {row["type"] for row in rows} + used = set(heavy_1daw.restraints.topology.atoms.energy_type.tolist()) - {""} + assert used and used <= table + by_type = {row["type"]: row for row in rows} + assert by_type["NH1"]["hb_type"] == "D" + assert by_type["O"]["hb_type"] == "A" + assert by_type["OH1"]["hb_type"] == "B" + assert float(by_type["CH3"]["vdwh_radius"]) > float(by_type["CH3"]["vdw_radius"]) diff --git a/tests/unit/topology/test_hydrogen_frames.py b/tests/unit/topology/test_hydrogen_frames.py new file mode 100644 index 00000000..268f6a52 --- /dev/null +++ b/tests/unit/topology/test_hydrogen_frames.py @@ -0,0 +1,154 @@ +"""Riding frames read off the bond graph, and their bookkeeping across table edits. + +Every hydrogen bonded to a heavy atom gets a frame; a parent with a single heavy +neighbour borrows a grandparent so the hydrogen turns with its torsion; and the frames +planned before an insertion agree, after remapping, with frames rebuilt on the +inserted table. +""" + +import numpy as np +import pytest +import torch + +from torchref.base.coordinates.local_frame import ( + frame_is_degenerate, + local_frame_coordinates, + place_local_frame, +) +from torchref.model.model import Model +from torchref.topology.hydrogens import ( + HydrogenFrames, + augment_atom_table_with_maps, + hydrogen_frames, + optimise_free_torsions, + plan_hydrogens, +) + + +@pytest.fixture(scope="module") +def heavy_and_plan(pdb_dir): + """Heavy-only 1DAW with its hydrogen plan.""" + model = Model(verbose=0, strip_H=True, add_hydrogens=False) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + restraints = model.restraints + xyz = model.xyz().detach() + plan = plan_hydrogens(restraints.topology, restraints.cif_dict, xyz) + optimise_free_torsions(plan, restraints.topology, xyz) + return model, restraints, plan + + +@pytest.fixture(scope="module") +def hydrogenated(heavy_and_plan): + """The same structure with the plan inserted, plus the maps the insertion made.""" + model, restraints, plan = heavy_and_plan + augmented, old_to_new, plan_to_new = augment_atom_table_with_maps( + model.pdb, plan, restraints.topology + ) + full = Model(verbose=0, strip_H=False, add_hydrogens=False) + cell, spacegroup = model.cell.data.cpu().numpy(), model.spacegroup + + def reader(): + return augmented, cell, spacegroup + + reader.links = model.ctx.links + full.load(reader) + return full, old_to_new, plan_to_new + + +@pytest.mark.unit +def test_every_bonded_hydrogen_gets_a_frame_or_orientation(hydrogenated): + """Hydrogens have a heavy-atom frame or an independently rotatable water group.""" + full, _, _ = hydrogenated + frames = hydrogen_frames(full.restraints.topology) + n_h = int((full.pdb["element"].str.strip() == "H").sum()) + assert frames.n_hydrogens == n_h + assert (frames.frame_valid | (frames.rotation_group >= 0)).all() + assert (frames.parent_row >= 0).all() + is_h = full.restraints.topology.atoms.is_hydrogen.cpu().numpy() + assert not is_h[frames.parent_row].any() + assert not is_h[frames.n1_row[frames.n1_row >= 0]].any() + assert not is_h[frames.n2_row[frames.n2_row >= 0]].any() + + +@pytest.mark.unit +def test_single_neighbour_parents_borrow_the_grandparent(hydrogenated): + """A hydroxyl or methyl hydrogen is framed on the bond it rotates about.""" + full, _, _ = hydrogenated + frames = hydrogen_frames(full.restraints.topology) + names = full.pdb["name"].str.strip().values + resnames = full.pdb["resname"].str.strip().values + seen = {} + for parent, n1, n2 in zip(frames.parent_row, frames.n1_row, frames.n2_row): + key = (resnames[parent], names[parent]) + seen.setdefault(key, (names[n1], names[n2])) + assert seen[("SER", "OG")] == ("CB", "CA") + assert seen[("LYS", "NZ")] == ("CE", "CD") + # A two-neighbour parent frames on its own neighbours. + assert set(seen[("ALA", "CA")]) <= {"N", "C", "CB"} + + +@pytest.mark.unit +def test_planned_frames_match_frames_rebuilt_on_the_augmented_table( + heavy_and_plan, hydrogenated +): + """Remapping the pre-insertion frames reproduces the post-insertion ones.""" + _, restraints, plan = heavy_and_plan + full, old_to_new, plan_to_new = hydrogenated + planned = hydrogen_frames(restraints.topology, plan) + assert planned.n_planned == plan.n_hydrogens + carried = planned.remap(old_to_new).fill_planned_rows(plan_to_new).sorted_by_row() + rebuilt = hydrogen_frames(full.restraints.topology).sorted_by_row() + for field in ("h_row", "parent_row", "n1_row", "n2_row"): + np.testing.assert_array_equal(getattr(carried, field), getattr(rebuilt, field)) + np.testing.assert_array_equal(carried.frame_valid, rebuilt.frame_valid) + + +@pytest.mark.unit +def test_positions_round_trip_through_their_frames(hydrogenated): + """Placed hydrogens are reproduced exactly from heavy atoms and local offsets.""" + full, _, _ = hydrogenated + frames = hydrogen_frames(full.restraints.topology) + xyz = full.xyz().detach().cpu() + t = frames.to_tensors() + p, n1, n2 = xyz[t["parent_row"]], xyz[t["n1_row"]], xyz[t["n2_row"]] + h = xyz[t["h_row"]] + assert not frame_is_degenerate(p, n1, n2)[t["frame_valid"]].any() + local = local_frame_coordinates(p, n1, n2, h) + back = place_local_frame(p, n1, n2, local, t["frame_valid"], h - p) + tolerance = 1e-4 if xyz.dtype == torch.float32 else 1e-9 + assert float((back - h).norm(dim=1).max()) < tolerance + + +@pytest.mark.unit +def test_remap_drops_orphans_and_degrades_lost_frames(): + """A hydrogen whose parent vanishes is dropped; a lost n2 leaves a rigid frame.""" + frames = HydrogenFrames( + h_row=np.array([5, 6, 7]), + parent_row=np.array([1, 2, 3]), + n1_row=np.array([0, 1, 2]), + n2_row=np.array([2, 3, 4]), + frame_valid=np.array([True, True, True]), + ) + # Drop atoms 2 and 7: the first hydrogen loses n2, the second loses its parent, + # the third is gone itself. + old_to_new = np.array([0, 1, -1, 2, 3, 4, 5, -1]) + out = frames.remap(old_to_new) + assert out.h_row.tolist() == [4] + assert out.parent_row.tolist() == [1] + assert out.n2_row.tolist() == [-1] + assert out.frame_valid.tolist() == [False] + + +@pytest.mark.unit +def test_tensor_round_trip_preserves_frames(): + """to_tensors / from_tensors is lossless.""" + frames = HydrogenFrames( + h_row=np.array([3, 4]), + parent_row=np.array([1, 1]), + n1_row=np.array([0, 0]), + n2_row=np.array([2, -1]), + frame_valid=np.array([True, False]), + ) + back = HydrogenFrames.from_tensors(**frames.to_tensors()) + for field in ("h_row", "parent_row", "n1_row", "n2_row", "frame_valid"): + np.testing.assert_array_equal(getattr(back, field), getattr(frames, field)) diff --git a/tests/unit/topology/test_hydrogens.py b/tests/unit/topology/test_hydrogens.py index 922ea04b..bcd2620d 100644 --- a/tests/unit/topology/test_hydrogens.py +++ b/tests/unit/topology/test_hydrogens.py @@ -78,7 +78,9 @@ def test_every_candidate_hydrogen_is_placed(built, code): checked = 0 for residue in range(residues.n_residues): start, end = int(residues.atom_start[residue]), int(residues.atom_end[residue]) - template = _template(restraints.cif_dict, str(residues.resname[residue]).strip()) + template = _template( + restraints.cif_dict, str(residues.resname[residue]).strip() + ) if template is None: continue # Residues with altlocs plan one hydrogen per conformer; the two-sided count @@ -94,10 +96,14 @@ def test_every_candidate_hydrogen_is_placed(built, code): element = str(atoms.element[parent]).strip().upper() template_heavy = len(template["heavy_adjacency"].get(names[parent], [])) extra_bonds = max(0, heavy - template_heavy) - expected = max( - 0, - min(STANDARD_VALENCE.get(element, 4) - heavy, template_h - extra_bonds), - ) + valence = STANDARD_VALENCE.get(element, 4) + if ( + element == "N" + and extra_bonds == 0 + and str(atoms.energy_type[parent]).startswith("NT") + ): + valence = 4 + expected = max(0, min(valence - heavy, template_h - extra_bonds)) placed = int((plan.parent == parent).sum()) assert placed == expected, ( f"{code}: atom {parent} ({names[parent]} {element}) has {heavy} heavy " @@ -120,7 +126,9 @@ def test_free_torsions_are_exactly_the_single_neighbour_centres(built, code): parent = int(plan.parent[i]) neighbours = topology.atoms.neighbors(parent) heavy = int((~is_h[neighbours]).sum()) - assert (plan.group[i] >= 0) == (heavy == 1), ( + residue = int(topology.atoms.residue_of[parent]) + water = str(topology.residues.resname[residue]).strip() == "HOH" + assert (plan.group[i] >= 0) == (heavy == 1 and not water), ( f"{code}: hydrogen {plan.name[i]} on atom {parent} with {heavy} heavy " f"neighbours has group {plan.group[i]}" ) @@ -248,12 +256,8 @@ def test_augmented_table_keeps_residues_contiguous(built, code): @pytest.mark.unit -def test_waters_are_not_hydrogenated(built): - """A single-atom residue is skipped, and for a reason rather than by accident. - - One heavy atom gives no frame to align a template against and no bond to rotate - about, so a water's hydrogens could only be placed in an arbitrary direction. - """ +def test_waters_receive_two_hydrogens_with_initial_orientations(built): + """Explicit generation supplies two HOH hydrogens, including coordinated waters.""" _, restraints, plan = built("7L84") topology = restraints.topology @@ -263,7 +267,8 @@ def test_waters_are_not_hydrogenated(built): if str(topology.residues.resname[i]).strip() == "HOH" ] assert waters, "7L84 has no waters, so this asserts nothing" - assert not set(plan.residue.tolist()) & set(waters) + for residue in waters: + assert int((plan.residue == residue).sum()) == 2 @pytest.mark.unit @@ -384,7 +389,9 @@ def test_acetyl_cap_carbon_gets_no_hydrogen(tmp_path): model.load_pdb(str(path)) model.set_restraints_cif(None) restraints = model.restraints - plan = plan_hydrogens(restraints.topology, restraints.cif_dict, model.xyz().detach()) + plan = plan_hydrogens( + restraints.topology, restraints.cif_dict, model.xyz().detach() + ) by_parent = {} for parent, name in zip(plan.parent.tolist(), plan.name.tolist()): diff --git a/torchref/base/coordinates/__init__.py b/torchref/base/coordinates/__init__.py index 3fd18aa4..5e5dd1cf 100644 --- a/torchref/base/coordinates/__init__.py +++ b/torchref/base/coordinates/__init__.py @@ -30,6 +30,13 @@ smallest_diff_aniso, ) +from .local_frame import ( + frame_is_degenerate, + local_frame_axes, + local_frame_coordinates, + place_local_frame, +) + __all__ = [ # PyTorch implementations "cartesian_to_fractional_torch", @@ -45,4 +52,9 @@ # Periodic boundary "smallest_diff", "smallest_diff_aniso", + # Local frames (riding hydrogens) + "local_frame_axes", + "place_local_frame", + "local_frame_coordinates", + "frame_is_degenerate", ] diff --git a/torchref/base/coordinates/local_frame.py b/torchref/base/coordinates/local_frame.py new file mode 100644 index 00000000..a482af3b --- /dev/null +++ b/torchref/base/coordinates/local_frame.py @@ -0,0 +1,207 @@ +"""Local orthonormal frames anchored on three atoms, and points expressed in them. + +A hydrogen that rides on its parent is stored as a constant offset in a frame built +from the parent ``p`` and two reference heavy atoms ``n1``, ``n2``:: + + e1 = unit(n1 - p) + e2 = unit((n2 - p) orthogonal to e1) + e3 = e1 x e2 + h = p + lx*e1 + ly*e2 + lz*e3 + +Every function here is a pure tensor op on already-gathered positions, so autograd +carries a force on ``h`` back onto ``p``, ``n1`` and ``n2`` through the exact frame +Jacobian. Coordinates are Cartesian Angstroms unless the caller chooses otherwise; the +only scale-dependent constant is ``eps``, which floors norms before division. +""" + +from typing import Tuple + +import torch + +#: Norm floor in the same units as the coordinates (Angstroms here). +DEFAULT_EPS = 1e-8 + +#: A frame whose reference bonds are shorter than this, or whose ``n1-p-n2`` angle has +#: a sine below ``MIN_FRAME_SINE``, is treated as degenerate and placed rigidly. +MIN_FRAME_NORM = 1e-3 +MIN_FRAME_SINE = 0.1 + + +def rotate_vectors(vectors: torch.Tensor, rotation: torch.Tensor) -> torch.Tensor: + """Rotate Cartesian vectors by axis-angle rotation vectors. + + Parameters + ---------- + vectors : torch.Tensor + Cartesian vectors, shape ``(..., 3)``, in Å. + rotation : torch.Tensor + Rotation vectors with the same shape as ``vectors``, in radians. + Direction gives the axis and length gives the right-handed angle. + + Returns + ------- + torch.Tensor + Rotated vectors in Å. First and second derivatives are finite at zero. + """ + angle2 = rotation.square().sum(-1, keepdim=True) + # Both torch.where branches must be safe at zero. The series also avoids + # cancellation in (1 - cos(angle)) / angle**2 in single precision. + angle = angle2.clamp_min(1e-4).sqrt() + small = angle2 < 1e-4 + a = torch.where( + small, 1 - angle2 / 6 + angle2.square() / 120, torch.sin(angle) / angle + ) + b = torch.where( + small, + 0.5 - angle2 / 24 + angle2.square() / 720, + 0.5 * torch.sinc(angle / (2 * torch.pi)).square(), + ) + cross = torch.cross(rotation, vectors, dim=-1) + return vectors + a * cross + b * torch.cross(rotation, cross, dim=-1) + + +def local_frame_axes( + p: torch.Tensor, + n1: torch.Tensor, + n2: torch.Tensor, + eps: float = DEFAULT_EPS, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Right-handed orthonormal axes of the frame anchored at ``p``. + + Parameters + ---------- + p, n1, n2 : torch.Tensor + Cartesian positions, shape ``(H, 3)``, same dtype and device. + eps : float + Norm floor guarding a collapsed reference bond. + + Returns + ------- + e1, e2, e3 : torch.Tensor + Unit vectors, each ``(H, 3)``. ``e1`` points along ``n1 - p``, ``e2`` lies in + the ``p, n1, n2`` plane, ``e3 = e1 x e2``. + """ + a = n1 - p + e1 = a / a.norm(dim=-1, keepdim=True).clamp(min=eps) + b = n2 - p + b_perp = b - (b * e1).sum(-1, keepdim=True) * e1 + e2 = b_perp / b_perp.norm(dim=-1, keepdim=True).clamp(min=eps) + e3 = torch.cross(e1, e2, dim=-1) + return e1, e2, e3 + + +def place_local_frame( + p: torch.Tensor, + n1: torch.Tensor, + n2: torch.Tensor, + local_offset: torch.Tensor, + frame_valid: torch.Tensor, + rigid_offset: torch.Tensor, + eps: float = DEFAULT_EPS, +) -> torch.Tensor: + """Positions of points stored as local-frame offsets. + + Parameters + ---------- + p, n1, n2 : torch.Tensor + Frame atoms, each ``(H, 3)``. Rows flagged invalid may hold any in-bounds + position; their frame is not used. + local_offset : torch.Tensor + Coordinates in the ``(e1, e2, e3)`` frame, shape ``(H, 3)``. + frame_valid : torch.Tensor + Boolean ``(H,)``; where False the point is placed rigidly at + ``p + rigid_offset``. + rigid_offset : torch.Tensor + Cartesian ``p -> point`` vector for the rigid fallback, shape ``(H, 3)``. + eps : float + Norm floor guarding degenerate frames. + + Returns + ------- + torch.Tensor + Cartesian positions, shape ``(H, 3)``, differentiable in ``p``, ``n1``, ``n2``. + """ + e1, e2, e3 = local_frame_axes(p, n1, n2, eps) + h_frame = ( + p + + local_offset[:, 0:1] * e1 + + local_offset[:, 1:2] * e2 + + local_offset[:, 2:3] * e3 + ) + h_rigid = p + rigid_offset + return torch.where(frame_valid.unsqueeze(-1), h_frame, h_rigid) + + +def local_frame_coordinates( + p: torch.Tensor, + n1: torch.Tensor, + n2: torch.Tensor, + point: torch.Tensor, + eps: float = DEFAULT_EPS, +) -> torch.Tensor: + """Inverse of :func:`place_local_frame`: express ``point`` in the frame at ``p``. + + Parameters + ---------- + p, n1, n2, point : torch.Tensor + Cartesian positions, each ``(H, 3)``. + eps : float + Norm floor guarding degenerate frames. + + Returns + ------- + torch.Tensor + Local coordinates ``(lx, ly, lz)``, shape ``(H, 3)``, such that + ``place_local_frame(p, n1, n2, result, True, ...)`` returns ``point``. + """ + e1, e2, e3 = local_frame_axes(p, n1, n2, eps) + d = point - p + return torch.stack([(d * e1).sum(-1), (d * e2).sum(-1), (d * e3).sum(-1)], dim=-1) + + +def frame_is_degenerate( + p: torch.Tensor, + n1: torch.Tensor, + n2: torch.Tensor, + min_norm: float = MIN_FRAME_NORM, + min_sine: float = MIN_FRAME_SINE, +) -> torch.Tensor: + """Frames too ill-conditioned to carry an offset. + + A frame is degenerate when either reference bond is shorter than ``min_norm`` or + the two reference bonds are within ``asin(min_sine)`` of collinear, in which case + ``e2`` is set by numerical noise and a riding point would swing with it. + + Parameters + ---------- + p, n1, n2 : torch.Tensor + Cartesian positions, each ``(H, 3)``. + min_norm : float + Shortest acceptable reference bond, same units as the coordinates. + min_sine : float + Smallest acceptable ``|sin(angle(n1 - p, n2 - p))|``. + + Returns + ------- + torch.Tensor + Boolean ``(H,)``, True where the frame must not be used. + """ + a = n1 - p + b = n2 - p + na = a.norm(dim=-1) + nb = b.norm(dim=-1) + cross = torch.cross(a, b, dim=-1).norm(dim=-1) + sine = cross / (na * nb).clamp(min=min_norm * min_norm) + return (na < min_norm) | (nb < min_norm) | (sine < min_sine) + + +__all__ = [ + "DEFAULT_EPS", + "MIN_FRAME_NORM", + "MIN_FRAME_SINE", + "rotate_vectors", + "local_frame_axes", + "place_local_frame", + "local_frame_coordinates", + "frame_is_degenerate", +] diff --git a/torchref/cli/_common.py b/torchref/cli/_common.py index 77707686..394bef12 100644 --- a/torchref/cli/_common.py +++ b/torchref/cli/_common.py @@ -652,6 +652,7 @@ def load_model( verbose: int = 0, cif: Optional[Union[str, List[str]]] = None, add_hydrogens: bool = False, + hydrogens_in_xray: bool = True, ) -> "ModelFT": """Load a model from PDB or CIF, auto-detected by file extension. @@ -670,6 +671,8 @@ def load_model( generation and the restraints read the same dictionary. add_hydrogens : bool, optional Generate missing hydrogens on load. Default False. + hydrogens_in_xray : bool, optional + Whether hydrogens contribute to the structure factors. Default True. Returns ------- @@ -685,6 +688,7 @@ def load_model( verbose=verbose, cif_path=cif, add_hydrogens=add_hydrogens, + hydrogens_in_xray=hydrogens_in_xray, ) suffix = Path(path).suffix.lower() if suffix in (".cif", ".mmcif"): diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index 2ffae3e3..e4c1853f 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -111,8 +111,8 @@ def setup_model_collection(pdb_dark, pdb_light, fractions, cif, d_min, sys.stdout.flush() model_dark = model_dark.hydrogenate(verbose=max(0, verbose - 1)) model_light = model_light.hydrogenate(verbose=max(0, verbose - 1)) - model_dark.exclude_H_from_sf = True - model_light.exclude_H_from_sf = True + model_dark.hydrogens_in_xray = False + model_light.hydrogens_in_xray = False mc = ModelCollection([model_dark, model_light], dark_key="dark") mc.add_dark() diff --git a/torchref/cli/refine.py b/torchref/cli/refine.py index bec7fadc..7aed867c 100644 --- a/torchref/cli/refine.py +++ b/torchref/cli/refine.py @@ -120,6 +120,14 @@ def main(): help="Generate missing hydrogens when loading the model (default: off). " "Hydrogens already present in the input are retained either way.", ) + refine_group.add_argument( + "--hydrogens-in-xray", + dest="hydrogens_in_xray", + action=argparse.BooleanOptionalAction, + default=True, + help="Include hydrogen atoms in the structure-factor calculation (default: on). " + "--no-hydrogens-in-xray keeps them in the restraints only.", + ) refine_group.add_argument( "--mode", type=str, @@ -255,6 +263,7 @@ def main(): print(f"X-ray target: {args.xray_mode}") print(f"Refinement cycles: {args.n_cycles}") print(f"Add hydrogens: {'on' if args.add_hydrogens else 'off'}") + print(f"Hydrogens in Fcalc: {'on' if args.hydrogens_in_xray else 'off'}") if args.with_rigid_body: print(f"Rigid-body step: on (iterations/cutoff = {args.rigid_body_iter})") print(f"Device: {args.device}") @@ -312,6 +321,7 @@ def main(): aniso_selection=args.anisotropic_selection, wavelength=args.wavelength, add_hydrogens=args.add_hydrogens, + hydrogens_in_xray=args.hydrogens_in_xray, ) # Merge onto DEFAULT_GROUP_WEIGHTS so unspecified groups keep their defaults; diff --git a/torchref/data/ener_lib_atoms.csv b/torchref/data/ener_lib_atoms.csv new file mode 100644 index 00000000..1035e5ad --- /dev/null +++ b/torchref/data/ener_lib_atoms.csv @@ -0,0 +1,213 @@ +# Per-energy-type atom properties from the CCP4 monomer library ener_lib.cif (_lib_atom loop). +# hb_type: N neither, D donor, A acceptor, B both, H hydrogen able to hydrogen-bond. +# vdw_radius: contact radius in Angstrom; vdwh_radius: radius to use when the atom's own hydrogens are not modelled. +# Regenerate with python -m torchref.scripts.extract_ener_lib +type,element,hb_type,vdw_radius,vdwh_radius,ion_radius +CSP,C,N,1.700,1.700, +CSP1,C,N,1.700,1.700, +C,C,N,1.700,1.750, +C1,C,N,1.700,1.820, +C2,C,N,1.700,1.800, +CR1,C,N,1.700,1.800, +CR2,C,N,1.700,1.800, +CR1H,C,N,1.700,1.800, +CR15,C,N,1.700,1.740, +CR5,C,N,1.700,1.740, +CR56,C,N,1.700,1.740, +CR55,C,N,1.700,1.740, +CR16,C,N,1.700,1.820, +CR6,C,N,1.700,1.740, +CR66,C,N,1.700,1.740, +CH1,C,N,1.700,1.950, +CH2,C,N,1.700,1.920, +CH3,C,N,1.700,1.940, +CT,C,N,1.700,1.850, +NS,N,A,1.550,1.600,1.32 +NSP,N,A,1.550,1.600,1.32 +NS1,N,D,1.550,1.600,1.32 +NSP1,N,D,1.550,1.600,1.32 +N,N,N,1.550,1.600,1.32 +NC1,N,D,1.550,1.600,1.32 +NH0,N,N,1.550,1.600,1.32 +NH1,N,D,1.550,1.600,1.32 +NC2,N,D,1.550,1.600,1.32 +NH2,N,D,1.550,1.600,1.32 +NC3,N,D,1.550,1.600,1.32 +N20,N,A,1.550,1.600,1.32 +N21,N,B,1.550,1.600,1.32 +NT,N,N,1.550,1.600,1.32 +NT1,N,D,1.550,1.600,1.32 +NT2,N,D,1.550,1.600,1.32 +NT3,N,D,1.550,1.600,1.32 +NT4,N,D,1.550,1.600,1.32 +N30,N,A,1.550,1.600,1.32 +N31,N,B,1.550,1.600,1.32 +N32,N,B,1.550,1.600,1.32 +N33,N,B,1.550,1.600,1.32 +NPA,N,A,1.550,1.600,1.32 +NPB,N,A,1.550,1.600,1.32 +NR5,N,A,1.550,1.600,1.32 +NR15,N,D,1.550,1.600,1.32 +NRD5,N,A,1.550,1.600,1.32 +NR56,N,N,1.550,1.600,1.32 +NR55,N,N,1.550,1.600,1.32 +NR6,N,A,1.550,1.600,1.32 +NR66,N,N,1.550,1.600,1.32 +NR16,N,D,1.550,1.600,1.32 +NRD6,N,A,1.550,1.600,1.32 +OS,O,A,1.520,1.520,1.28 +O,O,A,1.520,1.520,1.28 +AF,AF,A,1.520,1.520,1.28 +O2,O,A,1.520,1.520,1.28 +OH1,O,B,1.520,1.520,1.28 +OH2,O,B,1.520,1.520,1.28 +OHA,O,B,1.520,1.680,1.28 +OHB,O,B,1.520,1.680,1.28 +OHC,O,B,1.520,1.680,1.28 +OC2,O,A,1.520,1.520,1.28 +OC,O,A,1.520,1.520,1.28 +OP,O,A,1.520,1.520,1.28 +OB,O,A,1.520,1.520,1.28 +P,P,N,1.800,1.88,1.79 +P1,P,N,1.800,1.88,1.79 +PS,P,N,1.800,1.90,1.79 +S,S,A,1.800,1.88,1.7 +S3,S,A,1.800,1.88,1.7 +S2,S,A,1.800,1.88,1.7 +S1,S,A,1.800,1.88,1.7 +ST,S,A,1.800,1.88,1.7 +SH1,S,B,1.800,1.950,1.7 +H,H,N,1.200,1.200, +HCH,H,N,1.200,1.200, +HCH1,H,N,1.200,1.200, +HCH2,H,N,1.200,1.200, +HCH3,H,N,1.200,1.200, +HCR1,H,N,1.200,1.200, +HC2,H,N,1.200,1.200, +HC1,H,N,1.200,1.200, +HC2,H,N,1.200,1.200, +HCR5,H,N,1.200,1.200, +HCR6,H,N,1.200,1.200, +HNC1,H,H,1.200,1.200, +HNC2,H,H,1.200,1.200, +HNH1,H,H,1.200,1.200, +HNH2,H,H,1.200,1.200, +HNR5,H,H,1.200,1.200, +HNR6,H,H,1.200,1.200, +HNT1,H,H,1.200,1.200, +HNT2,H,H,1.200,1.200, +HNT3,H,H,1.200,1.200, +HOH1,H,H,1.200,1.200, +HOH2,H,H,1.200,1.200, +HOHA,H,H,1.200,1.200, +HOHB,H,H,1.200,1.200, +HOHC,H,H,1.200,1.200, +HSH1,H,H,1.200,1.200, +SI,SI,N,2.10,2.10,0.40 +SI1,SI,N,2.10,2.10,0.40 +GE,GE,N,2.10,2.10,0.40 +GE1,GE,N,2.10,2.10,0.40 +SN,SN,N,2.17,2.17,0.69 +PB,PB,N,2.02,2.02,0.79 +LI,LI,N,1.82,1.82,0.73 +NA,NA,N,2.27,2.27,1.13 +K,K,N,2.75,2.75,1.51 +RB,RB,N,2.00,2.00,1.48 +CS,CS,N,2.98,2.98,1.81 +FR,FR,N,,,1.94 +BE,BE,N,1.12,1.12,0.41 +MG,MG,N,1.73,1.73,0.71 +CA,CA,N,1.94,1.94,1.14 +SR,SR,N,2.19,2.19,1.32 +BA,BA,N,2.53,2.53,1.49 +RA,RA,N,2.15,2.15,1.62 +SC,SC,N,1.60,1.60,0.885 +Y,Y,N,1.80,1.80,1.04 +LA,LA,N,1.95,1.95,1.172 +CE,CE,N,1.85,1.85,1.01 +PR,PR,N,1.85,1.85,0.99 +ND,ND,N,1.85,1.85,1.123 +PM,PM,N,,,1.11 +SM,SM,N,1.85,1.85,1.098 +EU,EU,N,1.85,1.85,1.087 +GD,GD,N,1.80,1.80,1.078 +TB,TB,N,1.75,1.75,0.90 +DY,DY,N,1.75,1.75,1.052 +HO,HO,N,1.75,1.75,1.041 +ER,ER,N,1.75,1.75,1.03 +TM,TM,N,1.75,1.75,1.02 +YB,YB,N,1.75,1.75,1.008 +LU,LU,N,1.75,1.75,1.001 +AC,AC,N,1.95,1.95,1.26 +TH,TH,N,1.80,1.80,1.08 +PA,PA,N,1.80,1.80,0.92 +U,U,N,1.86,1.86,0.66 +NP,NP,N,1.75,1.75,0.85 +PU,PU,N,1.75,1.75,0.85 +AM,AM,N,1.75,1.75,0.99 +CM,CM,N,,,0.99 +BK,BK,N,,,0.97 +CF,CF,N,,,0.961 +ES,ES,N,,, +FM,FM,N,,, +MD,MD,N,,, +NO,NO,N,,, +LR,LR,N,,, +TI,TI,N,1.40,1.40,0.56 +ZR,ZR,N,1.55,1.55,0.73 +HF,HF,N,1.55,1.55,0.72 +RF,RF,N,,, +V,V,N,1.35,1.35,0.68 +NB,NB,N,1.45,1.45,0.62 +TA,TA,N,1.45,1.45,0.78 +DB,DB,N,,, +CR,CR,N,1.40,1.40,0.53 +MO,MO,N,1.45,1.45,0.55 +W,W,N,1.35,1.35,0.56 +SG,SG,N,,, +MN,MN,N,1.40,1.40,0.46 +TC,TC,N,1.35,1.35,0.51 +RE,RE,N,1.35,1.35,0.52 +BH,BH,N,,, +FE,FE,N,1.40,1.40,0.68 +RU,RU,N,1.30,1.30,0.52 +OSE,OS,N,1.30,1.30,0.53 +HS,HS,N,,, +CO,CO,N,1.35,1.35,0.54 +RH,RH,N,1.35,1.35,0.69 +IR,IR,N,1.35,1.35,0.71 +MT,MT,N,,, +NI,NI,N,1.63,1.63,0.63 +PD,PD,N,1.63,1.63,0.78 +PT,PT,N,1.75,1.75,0.71 +CU,CU,N,1.40,1.40,0.71 +AG,AG,N,1.72,1.72,0.81 +AU,AU,N,1.66,1.66,0.71 +ZN,ZN,N,1.39,1.39,0.74 +CD,CD,N,1.58,1.58,0.92 +HG,HG,N,1.55,1.55,1.10 +B,B,N,0.85,0.85,0.25 +AL,AL,N,1.25,1.25,0.53 +GA,GA,N,1.87,1.87,0.61 +IN,IN,N,1.93,1.93,0.76 +TL,TL,N,1.96,1.96,0.89 +AS,AS,N,1.85,1.85,0.475 +AS1,AS,N,1.85,1.85,0.475 +SB,SB,N,1.8,1.8,0.90 +BI,BI,N,1.8,1.8,0.90 +SE,SE,N,1.90,1.90,0.42 +TE,TE,N,2.06,2.06,0.57 +PO,PO,N,2.0,2.0,0.81 +F,F,B,1.47,1.47,1.19 +CL,CL,A,1.75,1.75,1.67 +BR,BR,N,1.85,1.85,0.73 +I,I,N,1.98,1.98,0.56 +AT,AT,N,1.80,1.80,0.76 +HE,HE,N,1.40,1.40, +NE,NE,N,1.54,1.54,1.12 +AR,AR,N,1.88,1.88,1.54 +KR,KR,N,2.02,2.02,1.69 +XE,XE,N,2.16,2.16,1.90 +RN,RN,N,2.20,2.20,2.0 +DUM,O,N,0.6,0.6,0.6 +.,C,B,1.7,1.5,1.7 diff --git a/torchref/experimental/ensemble/ensemble_model.py b/torchref/experimental/ensemble/ensemble_model.py index ba0b4ebe..de52ce11 100644 --- a/torchref/experimental/ensemble/ensemble_model.py +++ b/torchref/experimental/ensemble/ensemble_model.py @@ -349,6 +349,7 @@ def __init__( wavelength: float = 1.0, anomalous_threshold: float = 0.5, cif_path=None, + hydrogens_in_xray: bool = True, ): if dtype_float is None: dtype_float = get_float_dtype() @@ -365,6 +366,7 @@ def __init__( wavelength=wavelength, anomalous_threshold=anomalous_threshold, cif_path=cif_path, + hydrogens_in_xray=hydrogens_in_xray, ) # Filled in by ``_finalize_ensemble`` after ``load`` returns. self.n_members: int = 0 diff --git a/torchref/experimental/ensemble/quasi_crystal_amber.py b/torchref/experimental/ensemble/quasi_crystal_amber.py index 22544be4..ee9986ca 100644 --- a/torchref/experimental/ensemble/quasi_crystal_amber.py +++ b/torchref/experimental/ensemble/quasi_crystal_amber.py @@ -31,8 +31,7 @@ - the antechamber pipeline + GAFF2 setup for non-standard residues; - the template OpenMM ``System`` (single-molecule, AMBER14 / GAFF2); -- the H virtual-site frame tables (``_build_h_attachment``) and the shared - local-frame placement (``_place_hydrogens_local_frame``); +- a complete atom map including the model's hydrogens; - the autograd Function ``_OpenMMAMBERFunction``. This target then replicates the template ``System`` into the symmetry-expanded @@ -44,13 +43,14 @@ 1. replicate the System into a supercell with :func:`_replicate_to_supercell_system`; 2. build a new ``Context`` on the supercell (CUDA > OpenCL > CPU); -3. tile the template's atom map + H-attachment indices per member. +3. tile the template's complete atom map per member. Forward reads ``ensemble.xyz_per_member``, applies the supercell layout's sym+tile -transform, scatters into the unified OpenMM position tensor (heavy via -``_compose_full_omm_xyz``-style scatter; H via the tiled local-frame -placement), and calls the same ``_OpenMMAMBERFunction.apply``. +transform, scatters all model atoms into the unified OpenMM position tensor, +and calls ``_OpenMMAMBERFunction.apply``. Hydrogen positions come from TorchRef. +Construction leaves coordinates unchanged unless ``relax_on_init=True`` is +explicitly requested. """ from __future__ import annotations @@ -63,7 +63,6 @@ from torchref.experimental.targets.amber_target import ( AmberTarget, _OpenMMAMBERFunction, - _place_hydrogens_local_frame, ) from .ensemble_model import build_single_copy_model from .supercell import SupercellLayout, _replicate_to_supercell_system @@ -176,6 +175,9 @@ class QuasiCrystalAmberTarget(AmberTarget): charge_method : str antechamber charge method ('gas' or 'bcc'). Default 'gas' (fast, no QM); matches the ensemble setup. + relax_on_init : bool, default False + If True, explicitly minimize with OpenMM and write relaxed positions + back to the ensemble. Leave False to keep TorchRef's initial coordinates. verbose : int Verbosity (0 = silent, 1 = setup messages). @@ -204,7 +206,7 @@ def __init__( gaff2_files: Optional[Dict[str, Tuple[str, str]]] = None, charge_method: str = "gas", drop_special_position_threshold_ang: float = 0.0, - relax_on_init: bool = True, + relax_on_init: bool = False, relax_max_iterations: int = 200, force_clamp: float = 10000.0, verbose: int = 0, @@ -291,7 +293,7 @@ def __init__( # As an AmberTarget subclass we run the full antechamber + ForceField # pipeline against a genuine single-conformation Model (the ensemble's # ``_pdb_single`` restricted to non-special-position atoms). This - # populates self._system / _pos_buf / _model_to_omm / _h_* for ONE + # populates self._system / _pos_buf / _model_to_omm for ONE # member; we replicate them into the supercell below. ``_model`` stays # the ensemble (its per-member coords drive forward()); the # single-molecule context the base builds is replaced by the supercell @@ -322,9 +324,7 @@ def __init__( template_map = np.asarray(self._model_to_omm, dtype=np.int64) self._template_model_to_omm = template_map # Index pairs: model atom `src_model_idx[k]` lives in OMM slot - # `dst_omm_idx[k]` (single-member, in [0, n_omm_per_member)). Atoms - # with template_map == -1 (waters, OXT, ligands tleap regenerated) - # are excluded — their OMM slot keeps the construction-time position. + # `dst_omm_idx[k]` (single-member, in [0, n_omm_per_member)). valid_mask = template_map >= 0 self._src_model_idx_np = np.where(valid_mask)[0].astype(np.int64) self._dst_omm_idx_np = template_map[valid_mask].astype(np.int64) @@ -349,29 +349,6 @@ def __init__( self._n_omm_total = int(supercell_pos_nm.shape[0]) assert self._n_omm_total == N * self._n_omm_per_member - # --- H attachment, template arrays (numpy, into [0, n_omm_per_member)). - # The tiling per member is deferred to the forward path: each H index - # gets ``+ m · n_omm_per_member`` added per member m. - # - # AmberTarget marks rigid-fallback Hs (no valid local frame) with - # sentinel ``-1`` in ``_h_n1_idx`` / ``_h_n2_idx``. The corresponding - # ``_h_frame_valid`` row is False, so the local-frame branch never - # uses these indices. But ``index_select`` still evaluates the lookup - # and errors on negative indices, so clamp the sentinels to 0 — a - # safe in-bounds dummy whose result is then masked away by the - # ``frame_valid`` ``torch.where`` in :meth:`_place_hydrogens`. - self._h_idx_template = np.asarray(self._h_idx, dtype=np.int64).copy() - self._h_parent_idx_template = np.asarray(self._h_parent_idx, dtype=np.int64).copy() - h_n1 = np.asarray(self._h_n1_idx, dtype=np.int64).copy() - h_n2 = np.asarray(self._h_n2_idx, dtype=np.int64).copy() - h_n1[h_n1 < 0] = 0 - h_n2[h_n2 < 0] = 0 - self._h_n1_idx_template = h_n1 - self._h_n2_idx_template = h_n2 - self._h_local_pos_template = np.asarray(self._h_local_pos, dtype=np.float64).copy() - self._h_frame_valid_template = np.asarray(self._h_frame_valid, dtype=bool).copy() - self._h_offset_template = np.asarray(self._h_offset, dtype=np.float64).copy() - # Build the supercell System (replicate + PME + PBC). self._system = _replicate_to_supercell_system( template_system, @@ -575,91 +552,33 @@ def _relax_against_amber(self, max_iterations: int) -> None: # Lazy device buffers # ------------------------------------------------------------------ - def _ensure_torch_buffers( - self, device: torch.device, dtype: torch.dtype - ) -> None: - """Move/build the torch buffers for ``forward``: atom maps, tiled H - indices, and the constant init-positions tensor. Caches per + def _ensure_torch_buffers(self, device: torch.device, dtype: torch.dtype) -> None: + """Move/build atom maps and initial positions for ``forward``. Cache per (device, dtype). No work on repeat calls with the same key.""" - if ( - self._buffers_device == device - and self._buffers_dtype == dtype - ): + if self._buffers_device == device and self._buffers_dtype == dtype: return N = self._n_members n_omm = self._n_omm_per_member - # Initial sym-tiled positions for every OMM atom (nm). Used as the - # "fallback" position for slots that don't have a model atom mapped to - # them (waters, OXT, etc. tleap regenerated). - self._pos_buf_torch = torch.from_numpy(self._pos_buf).to( - device=device, dtype=dtype - ) - # Index pairs (long) for the scatter from model atoms into OMM slots. self._src_model_idx_torch = torch.from_numpy(self._src_model_idx_np).to( - device=device, dtype=torch.long # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 + device=device, + dtype=torch.long, # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 ) self._dst_omm_idx_torch = torch.from_numpy(self._dst_omm_idx_np).to( - device=device, dtype=torch.long # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 + device=device, + dtype=torch.long, # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 ) # Index of ensemble-model atoms (in the FULL EnsembleModel layout) # that survived the special-position filter — used in forward to # subset ``xyz_per_member`` before applying the layout transform. - self._keep_atom_idx_torch = torch.from_numpy( - self._keep_atom_idx_np - ).to(device=device, dtype=torch.long) # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 - - # Boolean mask: True where the OMM slot has NO model atom mapped to - # it (so we keep the init position there). - unmapped = torch.ones(n_omm, dtype=torch.bool, device=device) - unmapped[self._dst_omm_idx_torch] = False - self._unmapped_mask_torch = unmapped # (n_omm,) - - # H-attachment indices tiled per member: template indices live in - # [0, n_omm); full-tensor indices live in [0, N · n_omm). - member_offset = ( - torch.arange(N, device=device, dtype=torch.long).unsqueeze(1) # dtype-ok: arange index for broadcasting/indexing; PyTorch requires int64 - * n_omm - ) # (N, 1) - h_idx_t = torch.from_numpy(self._h_idx_template).to( - device=device, dtype=torch.long # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 - ) - h_parent_t = torch.from_numpy(self._h_parent_idx_template).to( - device=device, dtype=torch.long # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 - ) - h_n1_t = torch.from_numpy(self._h_n1_idx_template).to( - device=device, dtype=torch.long # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 - ) - h_n2_t = torch.from_numpy(self._h_n2_idx_template).to( - device=device, dtype=torch.long # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 - ) - - self._h_idx_tiled = (member_offset + h_idx_t.unsqueeze(0)).reshape(-1) - self._h_parent_idx_tiled = ( - member_offset + h_parent_t.unsqueeze(0) - ).reshape(-1) - self._h_n1_idx_tiled = ( - member_offset + h_n1_t.unsqueeze(0) - ).reshape(-1) - self._h_n2_idx_tiled = ( - member_offset + h_n2_t.unsqueeze(0) - ).reshape(-1) - - # Per-H constants tiled by member (same value for each member's - # corresponding H). - self._h_local_pos_tiled = torch.from_numpy( - self._h_local_pos_template - ).to(device=device, dtype=dtype).repeat(N, 1) - self._h_frame_valid_tiled = torch.from_numpy( - self._h_frame_valid_template - ).to(device=device, dtype=torch.bool).repeat(N) - self._h_offset_tiled = torch.from_numpy( - self._h_offset_template - ).to(device=device, dtype=dtype).repeat(N, 1) + self._keep_atom_idx_torch = torch.from_numpy(self._keep_atom_idx_np).to( + device=device, dtype=torch.long + ) # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 + self._omm_to_model = self._omm_to_model.to(device) self._buffers_device = device self._buffers_dtype = dtype @@ -667,16 +586,11 @@ def _ensure_torch_buffers( # Position composition # ------------------------------------------------------------------ - def _compose_full_omm_xyz( - self, supercell_xyz_nm: torch.Tensor - ) -> torch.Tensor: + def _compose_full_omm_xyz(self, supercell_xyz_nm: torch.Tensor) -> torch.Tensor: """Build the full ``(N · n_omm_per_member, 3)`` OpenMM xyz tensor. - Mapped (heavy) OMM slots get the current model coords (sym + tile - applied via the supercell layout); unmapped slots keep the construction- - time positions (tleap-regenerated atoms — waters, OXT, etc. that don't - move with the model). H atoms are then placed analytically from the - heavy positions via the tiled local-frame machinery. + Every OpenMM slot receives the current model coordinate after the + symmetry and tile transforms, including all hydrogen coordinates. Parameters ---------- @@ -690,50 +604,12 @@ def _compose_full_omm_xyz( Flat OpenMM-order positions in nm, differentiable in ``supercell_xyz_nm`` (and thus in ``model.xyz_per_member``). """ - N = self._n_members - n_omm = self._n_omm_per_member - device = supercell_xyz_nm.device - dtype = supercell_xyz_nm.dtype - - pos_init = self._pos_buf_torch.view(N, n_omm, 3) # (N, n_omm, 3) - # Scatter mapped model atoms into a zero tensor at the OMM slots. - src = supercell_xyz_nm.index_select(1, self._src_model_idx_torch) - # index_copy is autograd-friendly and returns a new tensor. - scattered = torch.zeros( - (N, n_omm, 3), device=device, dtype=dtype - ).index_copy(1, self._dst_omm_idx_torch, src) - # Where the slot is unmapped, use the init position; otherwise use - # scattered (the current model coord). - mask = self._unmapped_mask_torch.view(1, n_omm, 1) - heavy = torch.where(mask, pos_init, scattered) # (N, n_omm, 3) - - # Place hydrogens (operates on the flat (N·n_omm, 3) view). - full_flat = heavy.reshape(-1, 3) - return self._place_hydrogens(full_flat) - - def _place_hydrogens(self, heavy_xyz_nm: torch.Tensor) -> torch.Tensor: - """Vectorized H placement across all members via tiled local frames. - - Mirrors :meth:`AmberTarget._place_hydrogens` but operates on the - ``(N·n_omm_per_member, 3)`` supercell positions with tiled parent / - neighbour indices and per-H constants. Frame: parent + first heavy - neighbour for ``e1``, second heavy neighbour projected for ``e2``, - cross for ``e3``; H position is ``p + Σ local_pos[k] · e_k``. Rigid - fallback for the small fraction of Hs without two heavy neighbours. - """ - # Same local-frame physics as the single-molecule path — one shared - # implementation, here applied with member-tiled index tensors. - h_pos = _place_hydrogens_local_frame( - heavy_xyz_nm, - self._h_parent_idx_tiled, - self._h_n1_idx_tiled, - self._h_n2_idx_tiled, - self._h_local_pos_tiled, - self._h_frame_valid_tiled, - self._h_offset_tiled, - ) - # Write H positions into the heavy tensor via functional index_copy. - return heavy_xyz_nm.index_copy(0, self._h_idx_tiled, h_pos) + if supercell_xyz_nm.shape != (self._n_members, self._n_model_per_member, 3): + raise ValueError( + "[QuasiCrystalAmberTarget] Atom layout changed; rebuild the target." + ) + return supercell_xyz_nm.index_select(1, self._omm_to_model).reshape(-1, 3) + # ------------------------------------------------------------------ # Forward diff --git a/torchref/experimental/targets/amber_target.py b/torchref/experimental/targets/amber_target.py index 404cfcc3..ca6383f2 100644 --- a/torchref/experimental/targets/amber_target.py +++ b/torchref/experimental/targets/amber_target.py @@ -1,49 +1,10 @@ -""" -AMBER14/GAFF2 Force Field as a Differentiable Restraint. - -Uses OpenMM to evaluate the AMBER14 energy for current model coordinates. -Analytical forces from OpenMM are bridged into PyTorch autograd via a -custom Function, making the energy fully differentiable w.r.t. xyz. - -Non-standard residues (HETATM not in AMBER14_STANDARD) are parameterised -automatically via antechamber/GAFF2. Results are cached under -``PATH_TORCHREF_DATA / "amber_cache" / {resname}/``. - -Intended workflow:: - - # Canonical one-liner — strips altlocs, adds H, then build target: - mh = (Model(verbose=0, strip_H=True) - .load_pdb('structure.pdb') - .strip_altlocs() - .hydrogenate()) - target = AmberTarget(model=mh) # protein-only - target = AmberTarget(model=mh, residue_charges={'LIG': -1}) # with ligand - - loss = target() # kJ/mol per atom - loss.backward() - # xyz gradient is now populated with AMBER forces - -Performance note ----------------- -OpenMM's ``Modeller.addHydrogens()`` is faster when H atoms are already present -in the model (it refines positions rather than building from scratch). -Gradient and energy are identical either way (H are stripped from the atom map; -``n_model_atoms`` changes only the energy normalisation). - -Design notes ------------- -- Standard-residues path uses pdbfixer to add missing terminal/sidechain - heavy atoms before OpenMM's Modeller adds H; the GAFF2 path uses tleap - for H addition. -- Altloc atoms are filtered before building the OpenMM system: only the - primary conformation (altloc == '' or 'A') is used. -- OXT and H atoms are excluded from the PDB written to tleap; tleap - re-adds them via its C-terminal and H-addition templates. -- H positions in the OpenMM context are set once at construction and are - NOT updated during forward() — a good approximation for small refinement - steps (< 0.1 Å heavy-atom displacement). -- model_to_omm maps model-atom index → OpenMM atom index for HEAVY atoms - only. Model H atoms receive -1 and are skipped in forward(). +"""Evaluate AMBER energies and forces for TorchRef-owned atomic coordinates. + +The model must already contain the atoms and protonation state required by the +force field. Construction validates a one-to-one atom map; it does not add atoms +to the model. Each evaluation reorders all coordinates, including hydrogens, +converts Cartesian Å to nm, and returns OpenMM forces through PyTorch autograd. +Riding geometry and orientation parameters belong to the model. """ from __future__ import annotations @@ -63,6 +24,8 @@ import torch from torchref import PATH_TORCHREF_DATA +from torchref.config import get_float_dtype, get_int_dtype +from torchref.refinement.targets.base import ModelTarget from torchref.utils.stats import ( VERBOSITY_DEBUG, VERBOSITY_DETAILED, @@ -71,8 +34,6 @@ stat, ) -from torchref.refinement.targets.base import ModelTarget - if TYPE_CHECKING: from torchref.model.model import Model @@ -131,24 +92,8 @@ def _find_ambertools_binary(name: str) -> str: } ) -# Atom names that tleap adds itself via terminal / template logic; must be -# excluded from the PDB handed to tleap to avoid "does not have a type" errors. _TLEAP_SKIP_ATOMS: frozenset = frozenset({"OXT", "OT1", "OT2"}) -# Residues handled by amber14-all.xml + amber14/tip3pfb.xml in OpenMM Modeller -# (does NOT include Mg/Zn/Ca/Fe etc. — those lack templates in the default XML set) -_MODELLER_FF_RESIDUES: frozenset = frozenset( - { - "ALA", "ARG", "ASN", "ASP", "CYS", "CYX", "GLN", "GLU", "GLY", - "HID", "HIE", "HIP", "HIS", "ILE", "LEU", "LYS", "MET", "PHE", - "PRO", "SER", "THR", "TRP", "TYR", "VAL", - "ACE", "NME", - "HOH", "WAT", - "NA", "K", "CL", # ions in amber14/tip3pfb.xml - "A", "G", "C", "U", "T", "DA", "DG", "DC", "DT", - } -) - # Residues to exclude from the protein PDB written to tleap (GAFF2 path). # Currently empty: all AMBER14_STANDARD residues (protein, ions, water) are # included so they participate in both LJ (steric) and Coulomb gradients. @@ -177,17 +122,9 @@ class _OpenMMAMBERFunction(torch.autograd.Function): forward : full_xyz_nm (nm, float, [n_omm_total, 3]) → energy (kJ/mol) - The input tensor must already contain positions for **every** OpenMM - atom — heavy and H — in OpenMM's native atom order. Building this - tensor (scattering model heavy atoms + computing H positions - analytically from heavy positions) happens in - :meth:`AmberTarget._compose_full_omm_xyz`, upstream of this Function. - - backward: ∂E/∂full_xyz = −F (full OpenMM force vector). The H - contributions in F propagate naturally through ``_compose_full_omm_xyz`` - and ``_place_hydrogens`` upstream via PyTorch autograd, delivering - correctly-distributed gradients to the heavy model atoms (parent + - local-frame neighbors). + The input contains every model atom in OpenMM order. Backward returns + minus the force in kJ/mol/nm; the upstream gather and Å-to-nm conversion + return each gradient to its TorchRef coordinate or riding parameter. """ @staticmethod @@ -235,83 +172,6 @@ def backward(ctx, grad_output): return -forces * grad_output, None, None -# --------------------------------------------------------------------------- -# Differentiable hydrogen placement (single source of truth) -# --------------------------------------------------------------------------- - - -def _place_hydrogens_local_frame( - heavy_xyz: torch.Tensor, - parent_idx: torch.Tensor, - n1_idx: torch.Tensor, - n2_idx: torch.Tensor, - local_pos: torch.Tensor, - frame_valid: torch.Tensor, - offset: torch.Tensor, - eps: float = 1e-12, -) -> torch.Tensor: - """Place hydrogens from heavy-atom positions via captured local frames. - - The one and only implementation of the H-placement physics, shared by the - single-molecule / per-member path (:meth:`AmberTarget._place_hydrogens`) - and the tiled supercell path (``QuasiCrystalAmberTarget._place_hydrogens``). - Differentiable in ``heavy_xyz``: autograd distributes each H force onto its - parent + the two frame-reference atoms via the exact local-frame Jacobian. - - For each H, an orthonormal frame is built from its parent ``p`` and two - heavy neighbours ``n1, n2``:: - - e1 = û(n1 − p) - e2 = û((n2 − p) ⊥ e1) - e3 = e1 × e2 - h = p + lx·e1 + ly·e2 + lz·e3 - - Hs flagged ``frame_valid == False`` (no two heavy neighbours) fall back to - the rigid translation ``p + offset``. - - Parameters - ---------- - heavy_xyz : torch.Tensor, ``(M, 3)`` - Positions (nm) with all heavy-atom slots populated. May be a single - topology (``M = n_omm``) or a tiled supercell (``M = N · n_omm``). - parent_idx, n1_idx, n2_idx : torch.Tensor, ``(H,)`` long - Indices into ``heavy_xyz``. Invalid-frame neighbour indices must be - pre-clamped to a safe in-bounds value (their result is masked out). - local_pos : torch.Tensor, ``(H, 3)`` - Captured local-frame coordinates of each H. - frame_valid : torch.Tensor, ``(H,)`` bool - Whether the local-frame placement is used (else the rigid fallback). - offset : torch.Tensor, ``(H, 3)`` - Rigid-fallback ``p → H`` vector. - eps : float - Norm floor guarding degenerate frames. - - Returns - ------- - torch.Tensor, ``(H, 3)`` - H positions (nm). The caller writes these into the H slots. - """ - p = heavy_xyz.index_select(0, parent_idx) - n1 = heavy_xyz.index_select(0, n1_idx) - n2 = heavy_xyz.index_select(0, n2_idx) - - a = n1 - p - e1 = a / a.norm(dim=-1, keepdim=True).clamp(min=eps) - b = n2 - p - b_perp = b - (b * e1).sum(-1, keepdim=True) * e1 - e2 = b_perp / b_perp.norm(dim=-1, keepdim=True).clamp(min=eps) - e3 = torch.cross(e1, e2, dim=-1) - - h_frame = ( - p - + local_pos[:, 0:1] * e1 - + local_pos[:, 1:2] * e2 - + local_pos[:, 2:3] * e3 - ) - h_rigid = p + offset - return torch.where(frame_valid.unsqueeze(-1), h_frame, h_rigid) - - # --------------------------------------------------------------------------- # AmberTarget # --------------------------------------------------------------------------- @@ -321,47 +181,17 @@ class AmberTarget(ModelTarget): """ Differentiable AMBER14/GAFF2 force-field energy restraint. - On construction the target: - - 1. Detects non-standard residues (HETATM not in :data:`AMBER14_STANDARD`). - 2. Runs antechamber + parmchk2 (parallel, cached) for each non-standard - residue. - 3. Builds an OpenMM system: - - * **Standard path** (no non-standard residues): filter model PDB to - primary conformation + heavy atoms, use ``openmm.app.Modeller`` to - re-add H with AMBER14-compatible names, create system with - ``ForceField('amber14-all.xml')``. - * **GAFF2 path** (with non-standard residues): same protein PDB - (additionally removing OXT) handed to tleap together with each - ligand's mol2 via ``combine{}``. Combined AMBER14+GAFF2 topology - is parameterised by parmed. - - 4. Creates an OpenMM Context on the platform that matches the model's - device: CUDA for ``model.device.type == 'cuda'``, CPU otherwise. - Falls back CUDA → OpenCL → CPU if the preferred platform is unavailable. - 5. Builds a model-atom → OpenMM-atom index map so that only heavy atoms - are transferred; H positions are kept from the initial OpenMM setup. + Build chemistry once, then supply current coordinates for every atom to + OpenMM. The loss never generates or independently places hydrogens. Parameters ---------- model : Model - TorchRef model. Heavy-atom-only models (``strip_H=True``) are - accepted. H atoms are added internally by OpenMM's Modeller or - tleap and are NOT included in the atom map or gradient. - - Passing a model that already has H atoms (via - ``model.hydrogenate()`` or loading a PDB with H) speeds up - initialisation because ``Modeller.addHydrogens()`` converges - faster from existing positions. - - **GAFF2 ligands**: antechamber's BCC charge scheme runs a - semiempirical QM step (sqm) that needs a fully protonated molecule. - Heavy-only ligands are auto-protonated from the monomer library - (``hydrogenate``) first; an error is raised only if no - monomer CIF resolves AND the heavy-atom electron count is odd. - Calling ``model.hydrogenate()`` or loading the PDB with - ``strip_H=False`` beforehand avoids relying on that fallback. + Fully prepared, single-conformation model, including hydrogens and + terminal atoms required by AMBER. Existing atoms and coordinates are + preserved. Prepare protonation before constructing this target and + enable riding mode on the model when hydrogen geometry is constrained. + Incomplete or incompatible chemistry raises ValueError during setup. cutoff : float Non-bonded cutoff in Angstroms. Default 5.0. normalize_by_atoms : bool @@ -389,6 +219,15 @@ class AmberTarget(ModelTarget): multi-member ensemble, ``chem_model`` supplies the one conformation used to build the chemistry/topology; defaults to ``model`` for the single-molecule case. + + Notes + ----- + Reconstruct the target after changing atom identities, atom order, or + connectivity. Cartesian coordinate and riding-parameter changes need no + rebuild. OpenMM evaluation transfers coordinates and forces through CPU + memory and supports first derivatives only. Forces above 10000 kJ/mol/nm + are clipped per atom; in that regime the returned gradient is clipped + rather than the exact energy derivative. """ name: str = "amber" @@ -403,7 +242,7 @@ def __init__( charge_method: str = "gas", verbose: int = 0, chem_model: "Model" = None, - ): + ) -> None: try: import openmm # noqa: F401, PLC0415 except ImportError: @@ -438,7 +277,10 @@ def __init__( self._residue_charges = dict(residue_charges) if residue_charges else {} self._gaff2_files = dict(gaff2_files) if gaff2_files else {} - self.register_buffer("_cutoff_buf", torch.tensor(float(cutoff))) + self.register_buffer( + "_cutoff_buf", + torch.tensor(float(cutoff), dtype=get_float_dtype(), device=self.device), + ) # Internal state (None until fully initialised) self._context = None @@ -448,13 +290,8 @@ def __init__( self._n_omm_atoms: int = 0 self._n_model_atoms: int = 0 self._n_nonstandard: int = 0 - # GAFF2 path: ordered residue map for atom matching (None = standard path) - self._tleap_residue_map: Optional[List[Dict[str, int]]] = None - # Cached protonated chemistry PDB (filled lazily by the first ligand - # parameterisation that needs H). None = not yet computed; False = - # hydrogenate failed (don't retry). - self._protonated_pdb_cache = None - + # tleap renumbers residues; identify their original model instances. + self._tleap_residue_map: Optional[bool] = None if self._chem_model is None: return # Allow empty init for state_dict loading @@ -494,17 +331,14 @@ def _build(self) -> None: self._build_atom_map() del self._tleap_pos_nm + xyz = self._chem_model.xyz().detach().cpu().numpy() + positions_nm = np.asarray( + xyz[np.argsort(self._model_to_omm)] * 0.1, dtype=np.float64 + ) self._build_context(positions_nm) - # Pre-allocate nm position buffer: H positions pre-filled from OpenMM init self._pos_buf = positions_nm.copy() self._n_model_atoms = len(self._chem_model.pdb) - # Build (H, parent, offset) table so we can rigidly re-attach H atoms - # to their parent heavy atom each forward. Without this, H positions - # stay frozen at construction time while heavy atoms move, blowing up - # bond-stretch terms by orders of magnitude (the dominant pathology - # for any model.xyz() that excludes H). - self._build_h_attachment(positions_nm) if self.verbose >= 1: print( @@ -594,55 +428,9 @@ def _write_residue_pdb(self, res_atoms, path: Path) -> None: ) f.write("END\n") - def _protonated_chem_pdb(self): - """Protonated chemistry-model PDB DataFrame (cached), or ``None``. - Uses :meth:`Model.hydrogenate` once on the whole chemistry model, which has a - unit cell and full residue context so every centre has neighbours to orient its - template against. H come from the monomer-library CIF at ideal geometry via - TorchRef's auto-fetching monomer library -- no full CCP4 install needed. Cached - so repeated ligand parameterisations don't re-run it. - """ - if self._protonated_pdb_cache is None: - try: - m_h = self._chem_model.hydrogenate() - self._protonated_pdb_cache = ( - m_h.update_pdb() if hasattr(m_h, "update_pdb") else m_h.pdb - ) - except Exception as exc: # missing CIF/lib, gemmi failure, etc. - if self.verbose >= 1: - print(f"[AmberTarget] hydrogenate failed: {exc}") - self._protonated_pdb_cache = False - if self._protonated_pdb_cache is False: - return None - return self._protonated_pdb_cache - - def _protonate_residue_pdb(self, resname: str, out_pdb: Path) -> bool: - """Write a protonated single-residue PDB for ``resname`` to ``out_pdb``. - - antechamber/GAFF2 needs a protonated, valence-satisfied molecule because - the model is heavy-atom-only. Only topologically-correct H are required - here — charges are Gasteiger (connectivity-based, no QM) and the running- - system H are re-placed analytically each step — so the monomer library's - ideal geometry (via :meth:`_protonated_chem_pdb`) is ample. - - Returns ``True`` iff H were added for ``resname`` (a monomer CIF - resolved); ``False`` lets the caller fall back. - """ - pdb_h = self._protonated_chem_pdb() - if pdb_h is None: - return False - res = pdb_h[pdb_h["resname"].astype(str).str.strip() == resname] - h_mask = res["element"].astype(str).str.strip().isin(["H", "D"]) - if not bool(h_mask.any()): - return False - self._write_residue_pdb(res, out_pdb) - if self.verbose >= 1: - print( - f"[AmberTarget] protonated '{resname}' via monomer library: " - f"+{int(h_mask.sum())} H" - ) - return True + + def _run_antechamber_one( self, resname: str, charge: int @@ -653,24 +441,14 @@ def _run_antechamber_one( Cache is checked first. On a miss, work happens in a temp dir and results are atomically moved to the cache (write-then-rename). """ - pdb = self._chem_model.pdb + pdb = self._chem_model.pdb.copy() + pdb[["x", "y", "z"]] = self._chem_model.xyz().detach().cpu().numpy() res_atoms = pdb[pdb["resname"].astype(str).str.strip() == resname] + first = res_atoms.iloc[0] + for column in ("chainid", "resseq", "icode"): + res_atoms = res_atoms[res_atoms[column] == first[column]] atom_names = res_atoms["name"].astype(str).str.strip().tolist() - # antechamber needs a fully protonated molecule (sqm — used for BCC - # charges — needs an even electron count, and GAFF2 atom typing needs - # satisfied valences). The model is heavy-atom-only, so a ligand with no - # H is protonated below from the monomer library before antechamber runs. - # Compute the heavy-atom electron parity here to sanity-check the result. - _Z = {"H":1,"He":2,"Li":3,"Be":4,"B":5,"C":6,"N":7,"O":8,"F":9,"Ne":10, - "Na":11,"Mg":12,"Al":13,"Si":14,"P":15,"S":16,"Cl":17,"Ar":18, - "K":19,"Ca":20,"Cr":24,"Mn":25,"Fe":26,"Co":27,"Ni":28,"Cu":29, - "Zn":30,"Br":35,"I":53,"Se":34,"Mo":42,"W":74,"Pt":78,"Au":79} - elems = res_atoms["element"].astype(str).str.strip().str.capitalize() - n_protons = sum(_Z.get(e, 0) for e in elems) - n_electrons = n_protons - charge - has_h = bool(elems.isin(["H", "D"]).any()) - key = self._cache_key(resname, atom_names, charge, self._charge_method) cache_dir = self._get_cache_dir(resname) @@ -693,26 +471,7 @@ def _run_antechamber_one( self._write_residue_pdb(res_atoms, lig_pdb) - # Heavy-atom-only ligand → protonate before antechamber so GAFF2 - # typing sees satisfied valences (and sqm, if BCC, gets a closed- - # shell molecule). Hydrogens come from TorchRef's monomer-library - # placement at ideal geometry. antechamber_input = lig_pdb - if not has_h: - lig_h_pdb = work_dir / "lig_h.pdb" - if self._protonate_residue_pdb(resname, lig_h_pdb): - antechamber_input = lig_h_pdb - elif n_electrons % 2 != 0: - raise RuntimeError( - f"[AmberTarget] Cannot parameterise '{resname}': odd " - f"electron count ({n_electrons}) for charge {charge:+d} " - f"and no hydrogens could be added (no monomer-library CIF " - f"resolved for '{resname}').\nFix: pass an explicit charge " - f"via residue_charges={{'{resname}': }}, supply " - f"gaff2_files for this residue, or make a monomer CIF " - f"resolvable for auto-protonation (TORCHREF_MONOMER_LIB, " - f"or CLIBD_MON as an optional override)." - ) # antechamber r = subprocess.run( @@ -820,54 +579,16 @@ def _run_antechamber_parallel( # Step 3 — Build OpenMM system # ------------------------------------------------------------------ - def _filter_pdb_for_omm(self, include_nonstandard: bool = False): - """ - Return a filtered copy of model.pdb suitable for OpenMM / tleap: - - Primary conformation only (altloc == '' or 'A') - - Heavy atoms only (element != H or D) - - Optionally exclude non-standard residues (standard path) - - The returned DataFrame keeps the original model.pdb integer index - so that ``df.index`` can be used as model row indices in the atom map. - """ - pdb = self._chem_model.update_pdb() - - mask = pdb["altloc"].astype(str).str.strip().isin(["", "A"]) - mask &= ~pdb["element"].astype(str).str.strip().isin(["H", "D"]) - - if not include_nonstandard: - ns_resnames = { - rn for rn in pdb["resname"].astype(str).str.strip().unique() - if rn not in _MODELLER_FF_RESIDUES - } - if ns_resnames: - mask &= ~pdb["resname"].astype(str).str.strip().isin(ns_resnames) - # Do NOT reset_index: keep original model.pdb row positions as index - return pdb[mask].copy() def _filter_pdb_for_tleap(self): + """Export standard heavy atoms for tleap template parameterisation. + + The resulting topology must map back to every model atom, including + hydrogens and terminal oxygens, before a context can be constructed. """ - Filter model.pdb for the tleap protein PDB (GAFF2 path): - - - Primary conformation only (altloc == '' or 'A') - - Heavy atoms only (element != H or D) - - Standard AMBER residues only (``AMBER14_STANDARD``) — non-standard - HETATM residues are handled via antechamber / mol2 separately - - Waters (HOH/WAT) ARE included — ``_TLEAP_EXCLUDE_RESIDUES`` is - empty, so all ``AMBER14_STANDARD`` residues participate in the - LJ/Coulomb gradients (atom matching is position-based, so tleap's - water ordering does not break the map) - - Monatomic ions (MG, ZN, CA, …) ARE included — covered by - ``leaprc.water.tip3p`` (Li/Merz 12-6 set), appear in fixed PDB - order, important for electrostatics near charged ligands - - Terminal atoms tleap regenerates (OXT …) excluded - - Note: uses ``AMBER14_STANDARD`` (not ``_MODELLER_FF_RESIDUES``) - so that ions absent from amber14-all.xml are still sent to tleap. - Index is preserved (original model.pdb row positions). - """ - pdb = self._chem_model.update_pdb() + pdb = self._chem_model.pdb.copy() + pdb[["x", "y", "z"]] = self._chem_model.xyz().detach().cpu().numpy() mask = pdb["altloc"].astype(str).str.strip().isin(["", "A"]) mask &= ~pdb["element"].astype(str).str.strip().isin(["H", "D"]) @@ -886,21 +607,7 @@ def _filter_pdb_for_tleap(self): def _build_omm_system( self, gaff2_params: Dict[str, Tuple[Path, Path]] ) -> Tuple: - """ - Build OpenMM system. Returns ``(system, omm_topology, pos_nm_array)``. - - Standard path (no non-standard residues) - ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ - Filter model PDB → heavy atoms, primary conformation, standard residues. - Use ``openmm.app.Modeller.addHydrogens()`` to re-add H with AMBER names. - Create system with ``ForceField('amber14-all.xml')``. - - GAFF2 path (non-standard residues present) - ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ - Write protein PDB (no OXT, no H) + mol2 per ligand. - Combine via tleap ``combine{}`` command → prmtop/inpcrd. - Load with parmed → ``AmberParm.createSystem()``. - """ + """Parameterise the model with AMBER14 or AMBER14/GAFF2.""" import openmm as mm # noqa: PLC0415 import openmm.app as app # noqa: PLC0415 import openmm.unit as unit # noqa: PLC0415 @@ -922,73 +629,57 @@ def _build_omm_system( return system, topology, pos_nm def _build_standard(self, cutoff_A: float, app, unit) -> Tuple: - """ - AMBER14 standard-residue path using gemmi + pdbfixer + OpenMM. - - gemmi writes proper chain termination / TER records so pdbfixer - can detect and fix missing terminal atoms (OXT). pdbfixer also - handles missing sidechain atoms and non-standard residue names. - """ - import gemmi # noqa: PLC0415 - from pdbfixer import PDBFixer # noqa: PLC0415 - from torchref.io import pdb as pdbio # noqa: PLC0415 + """Parameterise existing atoms, retaining PDB serials through name aliases.""" + from torchref.io import pdb as pdbio - # Standard path: Modeller preserves chain/resseq → use key-based mapping self._tleap_residue_map = None - - pdb_heavy = self._filter_pdb_for_omm(include_nonstandard=False) - - tmp = tempfile.NamedTemporaryFile(suffix=".pdb", delete=False) - tmp2 = tempfile.NamedTemporaryFile(suffix=".pdb", delete=False) - tmp.close() - tmp2.close() - try: - # Write via torchref, then re-read/write with gemmi to get - # proper chain breaks and TER records that pdbfixer needs. - pdbio.write(pdb_heavy, tmp.name) - st = gemmi.read_structure(tmp.name) - st.setup_entities() - st.assign_subchains() - st.write_pdb(tmp2.name) - - # pdbfixer: add missing terminal atoms and sidechain atoms - fixer = PDBFixer(filename=tmp2.name) - fixer.findMissingResidues() - fixer.missingResidues = {} # don't fill gaps - fixer.findMissingAtoms() - - if self.verbose >= 1: - n_missing = sum(len(v) for v in fixer.missingAtoms.values()) - n_terminals = sum( - 1 for v in fixer.missingTerminals.values() if v - ) - if n_missing or n_terminals: - print( - f"[AmberTarget] pdbfixer: {n_missing} missing atoms, " - f"{n_terminals} terminal fixes" - ) - - fixer.addMissingAtoms() - finally: - os.unlink(tmp.name) - os.unlink(tmp2.name) - + pdb = self._chem_model.pdb.copy() + xyz = self._chem_model.xyz().detach().cpu().numpy() + pdb[["x", "y", "z"]] = xyz + pdb["serial"] = np.arange(1, len(pdb) + 1) + with tempfile.TemporaryDirectory(prefix="torchref_amber_") as directory: + filename = str(Path(directory) / "model.pdb") + pdbio.write(pdb, filename) + parsed = app.PDBFile(filename) + + topology = parsed.topology + source_rows = np.array([int(a.id) - 1 for a in topology.atoms()]) + if len(source_rows) != len(pdb) or not np.array_equal( + np.sort(source_rows), np.arange(len(pdb)) + ): + raise ValueError( + "[AmberTarget] PDB atom identities are ambiguous or duplicated. " + "Every TorchRef atom must correspond to exactly one AMBER particle." + ) + self._source_model_rows = source_rows + # Preserve explicit covalent links that the PDB bond templates cannot + # infer. PDB parsing also resolves standard hydrogen-name aliases. + atoms = list(topology.atoms()) + inverse = np.argsort(source_rows) + existing_bonds = {tuple(sorted((a.index, b.index))) for a, b in topology.bonds()} + for i, j in self._chem_model.restraints.topology.atoms.bonds.indices.cpu().tolist(): + pair = tuple(sorted((int(inverse[i]), int(inverse[j])))) + if pair not in existing_bonds: + topology.addBond(atoms[pair[0]], atoms[pair[1]]) + existing_bonds.add(pair) ff = app.ForceField("amber14-all.xml", "amber14/tip3pfb.xml") - modeller = app.Modeller(fixer.topology, fixer.positions) - modeller.addHydrogens(ff) - - system = ff.createSystem( - modeller.topology, - nonbondedMethod=app.CutoffNonPeriodic, - nonbondedCutoff=cutoff_A * unit.angstrom, - constraints=None, - ) - - # positions in nm - pos_nm = np.array( - modeller.positions.value_in_unit(unit.nanometer), dtype=np.float64 - ) - return system, modeller.topology, pos_nm + try: + system = ff.createSystem( + topology, + nonbondedMethod=app.CutoffNonPeriodic, + nonbondedCutoff=cutoff_A * unit.angstrom, + constraints=None, + rigidWater=False, + ) + except ValueError as exc: + raise ValueError( + "[AmberTarget] TorchRef model is not AMBER-compatible. " + "Prepare missing atoms, terminal groups and protonation in the " + "model before constructing the loss; no atoms were added. " + f"OpenMM: {exc}" + ) from exc + xyz = self._chem_model.xyz().detach().cpu().numpy() + return system, topology, np.asarray(xyz[source_rows] * 0.1, dtype=np.float64) def _build_gaff2( self, @@ -997,19 +688,10 @@ def _build_gaff2( app, unit, ) -> Tuple: - """ - AMBER14 + GAFF2 path via tleap + parmed. - - All AMBER14-standard heavy atoms (protein, ions, waters — no OXT, no H, - no non-standard HETATM) plus each ligand mol2 are combined by tleap - ``combine{}``. parmed loads the resulting prmtop/inpcrd. - - Atom mapping uses position-based matching (see :meth:`_build_atom_map`): - tleap's initial coordinates are taken directly from the PDB we write, - so model and tleap positions agree to 3 decimal places (PDB precision), - making a KD-tree nearest-neighbour search unambiguous. This avoids - relying on tleap's residue-sequential numbering, which is fragile for - water molecules. + """Build AMBER14/GAFF2 templates with one copy per ligand instance. + + tleap may reorder or rename atoms; the complete map is validated before + its system is used. Runtime coordinates always come from TorchRef. """ import parmed as pmd # noqa: PLC0415 from torchref.io import pdb as pdbio # noqa: PLC0415 @@ -1019,12 +701,8 @@ def _build_gaff2( prot_pdb = work_dir / "protein.pdb" pdb_tleap = self._filter_pdb_for_tleap() - # Signal GAFF2 path to _build_atom_map (position-based + name fallback) - self._tleap_residue_map = True # type: ignore[assignment] - # Store GAFF2 resnames so _build_atom_map can do name-based fallback - # for ligand atoms (mol2 may have old coords if model was refined first) - self._gaff2_resnames: set = set(gaff2_params.keys()) - + # tleap does not preserve the original chain and residue identifiers. + self._tleap_residue_map = True pdbio.write(pdb_tleap.reset_index(drop=True), str(prot_pdb)) prmtop = work_dir / "complex.prmtop" @@ -1038,7 +716,18 @@ def _build_gaff2( lig_loads.append(f"{rn} = loadMol2 {mol2}") lig_names.append(rn) - combine_list = " ".join(["protein"] + lig_names) + ligand_copies = [] + ligand_keys = [] + pdb = self._chem_model.pdb + for rn in lig_names: + rows = pdb[pdb["resname"].astype(str).str.strip() == rn] + for key, _ in rows.groupby(["chainid", "resseq", "icode"], sort=False): + copy_name = f"ligand{len(ligand_copies)}" + lig_loads.append(f"{copy_name} = copy {rn}") + ligand_copies.append(copy_name) + ligand_keys.append(tuple(key)) + self._gaff2_residue_keys = ligand_keys + combine_list = " ".join(["protein"] + ligand_copies) tleap_script = "\n".join( [ "source leaprc.protein.ff14SB", @@ -1072,6 +761,7 @@ def _build_gaff2( nonbondedMethod=app.CutoffNonPeriodic, nonbondedCutoff=cutoff_A * unit.angstrom, constraints=None, + rigidWater=False, ) topology = combined.topology pos_nm = np.array( @@ -1093,478 +783,173 @@ def _is_hydrogen(omm_atom) -> bool: return omm_atom.name.startswith("H") # heuristic fallback def _build_atom_map(self) -> None: - """ - Build ``self._model_to_omm``: int32 array [n_model] where entry *i* - is the OpenMM atom index corresponding to model atom *i*, or -1 for - unmatched atoms (H atoms, altloc-B atoms, non-standard HETATM, …). - - Two strategies depending on how the system was built: - - **Standard path** (``_tleap_residue_map is None``): - OpenMM Modeller preserves chain IDs and residue numbers from the input - PDB, so matching uses the key ``(chain_id, resseq, icode, atom_name)``. - - **GAFF2 path** (``_tleap_residue_map is not None``): - tleap strips chain IDs and renumbers residues sequentially, making - name/number-based matching unreliable (especially for waters). - Instead, the tleap initial positions are taken from the exact - coordinates we wrote to the PDB (via ``update_pdb()``), so model and - tleap positions agree to within PDB precision (0.001 Å = 0.0001 nm). - A KD-tree nearest-neighbour search with a tight threshold (0.005 nm) - unambiguously identifies each tleap heavy atom's model counterpart. - """ - from scipy.spatial import cKDTree # noqa: PLC0415 - + """Require a bijection between model rows and all OpenMM particles.""" pdb = self._chem_model.pdb n_model = len(pdb) - model_to_omm = np.full(n_model, -1, dtype=np.int32) - + atoms = list(self._topology.atoms()) + mapping = np.full(n_model, -1, dtype=np.int32) if self._tleap_residue_map is None: - # ---- Standard path: match by (chain, resseq, icode, atom_name) ---- - # No altlocs at this point (checked in _build). - model_key_to_idx: Dict[Tuple, int] = {} - for i in range(n_model): - row = pdb.iloc[i] - key = ( - str(row["chainid"]).strip(), - int(row["resseq"]), - str(row.get("icode", "")).strip(), - str(row["name"]).strip(), - ) - model_key_to_idx[key] = i - - for omm_atom in self._topology.atoms(): - if self._is_hydrogen(omm_atom): - continue - chain_id = omm_atom.residue.chain.id.strip() - try: - resseq = int(omm_atom.residue.id) - except ValueError: - raw = omm_atom.residue.id.strip() - resseq = int(raw.rstrip("ABCDEFGHIJKLMNOPQRSTUVWXYZ") or "0") - icode = (omm_atom.residue.insertionCode or "").strip() - idx = model_key_to_idx.get( - (chain_id, resseq, icode, omm_atom.name.strip()) - ) - if idx is not None: - model_to_omm[idx] = omm_atom.index - + mapping[self._source_model_rows] = np.arange(len(atoms)) else: - # ---- GAFF2 path: position-based matching via KD-tree ---- - # Collect tleap heavy-atom positions (nm) and their indices. - tleap_pos_nm = self._tleap_pos_nm # set by _build() before this call - tleap_ha_omm_idx: List[int] = [] - tleap_ha_pos: List[np.ndarray] = [] - for omm_atom in self._topology.atoms(): - if not self._is_hydrogen(omm_atom): - tleap_ha_omm_idx.append(omm_atom.index) - tleap_ha_pos.append(tleap_pos_nm[omm_atom.index]) - - tleap_ha_pos_arr = np.array(tleap_ha_pos) # (N_tleap_heavy, 3) nm - tree = cKDTree(tleap_ha_pos_arr) - - # Collect model primary-altloc heavy-atom positions (nm) and indices. - # Use update_pdb() coords — same values that were written to tleap PDB. - fresh_pdb = self._chem_model.update_pdb() - altloc_ok = fresh_pdb["altloc"].astype(str).str.strip().isin(["", "A"]) - not_h = ~fresh_pdb["element"].astype(str).str.strip().isin(["H", "D"]) - primary_heavy = np.where((altloc_ok & not_h).values)[0] - - model_pos_nm = np.column_stack([ - fresh_pdb["x"].values[primary_heavy], - fresh_pdb["y"].values[primary_heavy], - fresh_pdb["z"].values[primary_heavy], - ]) * 0.1 # Å → nm - - # Match: threshold = 0.005 nm (50× PDB precision of 0.0001 nm) - dists, nn_idx = tree.query(model_pos_nm, k=1) - matched = dists < 0.005 - for local_i, (model_i, nn_i) in enumerate(zip(primary_heavy, nn_idx)): - if matched[local_i]: - model_to_omm[model_i] = tleap_ha_omm_idx[nn_i] - - # Name-based fallback for GAFF2 ligand residues whose mol2 positions - # differ from the current model (e.g. after refinement steps). - # The cached mol2 retains original antechamber coordinates, so a - # second AmberTarget init after LBFGS will have position shifts. - gaff2_resnames = getattr(self, "_gaff2_resnames", set()) - if gaff2_resnames: - # Build (resname, atom_name) → model positional index for primary heavy - lig_key_to_model: Dict[Tuple[str, str], int] = {} - for arr_pos in primary_heavy: - rn = str(fresh_pdb["resname"].values[arr_pos]).strip() - if rn not in gaff2_resnames: - continue - aname = str(fresh_pdb["name"].values[arr_pos]).strip() - lig_key_to_model[(rn, aname)] = int(arr_pos) - - for omm_atom in self._topology.atoms(): - if self._is_hydrogen(omm_atom): - continue - rn = omm_atom.residue.name - if rn not in gaff2_resnames: - continue - aname = omm_atom.name.strip() - model_arr_pos = lig_key_to_model.get((rn, aname)) - if model_arr_pos is not None and model_to_omm[model_arr_pos] < 0: - model_to_omm[model_arr_pos] = omm_atom.index - - # Warn about UNEXPECTED unmatched heavy atoms. - # Expected to be unmatched (silently skipped in gradient): - # - H / D atoms - # - Waters, ions excluded from tleap (_TLEAP_EXCLUDE_RESIDUES) - # - C-terminal OXT regenerated by tleap (_TLEAP_SKIP_ATOMS) - # - Alternate conformer atoms (altloc != '' and != 'A') - elem_col = pdb["element"].astype(str).str.strip() - altloc_col = pdb["altloc"].astype(str).str.strip() - resname_col = pdb["resname"].astype(str).str.strip() - name_col = pdb["name"].astype(str).str.strip() - heavy_mask = ~elem_col.isin(["H", "D"]) - # Residues in AMBER14_STANDARD but without an amber14-all.xml template: - # unmatched on the standard (Modeller) path; matched via tleap on GAFF2 path. - _no_modeller_template = AMBER14_STANDARD - _MODELLER_FF_RESIDUES - expected_mask = ( - # Waters always excluded from tleap; no AMBER gradient expected - resname_col.isin(_TLEAP_EXCLUDE_RESIDUES) | - # Ions that lack Modeller templates (matched in GAFF2 path, not standard) - resname_col.isin(_no_modeller_template) | - # tleap-regenerated terminal atoms (OXT etc.) - name_col.isin(_TLEAP_SKIP_ATOMS) | - # alternate conformers (altloc B, C, …) - (~altloc_col.isin(["", "A"])) - ) - unexpected_unmatched = np.where( - heavy_mask.values & ~expected_mask.values & (model_to_omm < 0) - )[0] - if len(unexpected_unmatched) > 0: - ex = [ - f"{pdb.iloc[i]['name'].strip()} " - f"({pdb.iloc[i]['resname'].strip()} {pdb.iloc[i]['resseq']})" - for i in unexpected_unmatched[:5] - ] - warnings.warn( - f"[AmberTarget] {len(unexpected_unmatched)} heavy model atom(s) " - f"could not be matched to OpenMM topology " - f"(e.g. {', '.join(ex)}). Their gradients will be zero.", - UserWarning, - stacklevel=3, - ) - elif self.verbose >= 2: - unmatched_heavy = int(heavy_mask.values.sum()) - int( - (heavy_mask.values & (model_to_omm >= 0)).sum() - ) - print( - f"[AmberTarget] {unmatched_heavy} heavy atoms have model_to_omm=-1 " - f"(expected: non-standard HETATM / altloc-B / OXT)" - ) - - self._model_to_omm = model_to_omm - self._n_omm_atoms = self._system.getNumParticles() - - if self.verbose >= 2: - matched = int((model_to_omm >= 0).sum()) - print( - f"[AmberTarget] atom map: {matched}/{n_model} model atoms matched " - f"({self._n_omm_atoms} total OpenMM atoms)" - ) - - # ------------------------------------------------------------------ - # Hydrogen re-attachment - # ------------------------------------------------------------------ - - def _build_h_attachment(self, pos_nm: np.ndarray) -> None: - """ - Build the local-frame placement table for every H atom. - - Each H is placed at construction-time according to OpenMM's - ``Modeller.addHydrogens`` output. We freeze that placement in a - local frame defined by the parent heavy atom and 2 reference - heavy atoms. At forward time, the H position is recomputed in - differentiable PyTorch from the current heavy positions: - - e1 = (n1 − p) / |n1 − p| - e2 = perp(n2 − p, e1) / |perp(n2 − p, e1)| - e3 = e1 × e2 - h = p + lx·e1 + ly·e2 + lz·e3 - - where (lx, ly, lz) = (h − p) · [e1, e2, e3] is captured once. - - Backward through this formula in PyTorch autograd produces the - exact local-frame Jacobian — so the force on H from OpenMM gets - correctly distributed across p, n1, n2, not just onto p. - - Reference-atom selection per H: - - parent ``p`` : the unique heavy atom bonded to H - - neighbor ``n1`` : any heavy atom bonded to ``p`` (≠ H) - - neighbor ``n2`` : another heavy atom bonded to ``p``; if - ``p`` has only one heavy neighbour, fall - back to a heavy atom bonded to ``n1`` - (i.e. walk one bond further out). - - Hs with no usable triple — extremely rare in real chemistry — - fall back to the legacy ``h = p + offset`` rigid translation - path, with ``h_frame_valid=False``. - """ - if not hasattr(self, "_topology") or self._topology is None: - self._h_idx = None - self._h_parent_idx = None - self._h_n1_idx = None - self._h_n2_idx = None - self._h_local_pos = None - self._h_frame_valid = None - self._h_offset = None - return - - # ---- Walk topology bonds. Build heavy-neighbor adjacency and - # parent map for H atoms in one pass. ------------------------------ - from collections import defaultdict - - def _is_h(atom) -> bool: - return atom.element is not None and atom.element.symbol == "H" - - parent_of_h: Dict[int, int] = {} - heavy_neighbors: Dict[int, list] = defaultdict(list) - for bond in self._topology.bonds(): - a, b = bond[0], bond[1] - a_is_h = _is_h(a) - b_is_h = _is_h(b) - if a_is_h and not b_is_h: - parent_of_h[a.index] = b.index - elif b_is_h and not a_is_h: - parent_of_h[b.index] = a.index - elif not a_is_h and not b_is_h: - heavy_neighbors[a.index].append(b.index) - heavy_neighbors[b.index].append(a.index) - # H-H bonds are nonsense; ignored. - - if not parent_of_h: - self._h_idx = None - self._h_parent_idx = None - self._h_n1_idx = None - self._h_n2_idx = None - self._h_local_pos = None - self._h_frame_valid = None - self._h_offset = None - return - - # ---- Resolve (parent, n1, n2) per H ---------------------------- - h_indices = sorted(parent_of_h.keys()) - n_h = len(h_indices) - p_arr = np.empty(n_h, dtype=np.int64) - n1_arr = np.empty(n_h, dtype=np.int64) - n2_arr = np.empty(n_h, dtype=np.int64) - valid = np.zeros(n_h, dtype=bool) - - for k, h in enumerate(h_indices): - p = parent_of_h[h] - p_arr[k] = p - neigh = heavy_neighbors.get(p, []) - if len(neigh) >= 2: - n1_arr[k] = neigh[0] - n2_arr[k] = neigh[1] - valid[k] = True - elif len(neigh) == 1: - n1 = neigh[0] - further = [ - j for j in heavy_neighbors.get(n1, []) if j != p - ] - if further: - n1_arr[k] = n1 - n2_arr[k] = further[0] - valid[k] = True - else: - n1_arr[k] = n1 - n2_arr[k] = -1 - valid[k] = False - else: - n1_arr[k] = -1 - n2_arr[k] = -1 - valid[k] = False - - # ---- Compute local-frame coordinates from initial positions ---- - h_pos = pos_nm[np.asarray(h_indices, dtype=np.int64)] - p_pos = pos_nm[p_arr] - local_pos = np.zeros((n_h, 3), dtype=np.float64) - eps = 1e-12 - - valid_idx = np.where(valid)[0] - if valid_idx.size > 0: - n1_pos = pos_nm[n1_arr[valid_idx]] - n2_pos = pos_nm[n2_arr[valid_idx]] - a = n1_pos - p_pos[valid_idx] - b = n2_pos - p_pos[valid_idx] - e1 = a / np.maximum( - np.linalg.norm(a, axis=-1, keepdims=True), eps, - ) - b_perp = b - (b * e1).sum(-1, keepdims=True) * e1 - e2 = b_perp / np.maximum( - np.linalg.norm(b_perp, axis=-1, keepdims=True), eps, - ) - e3 = np.cross(e1, e2) - offset_v = h_pos[valid_idx] - p_pos[valid_idx] - local_pos[valid_idx, 0] = (offset_v * e1).sum(-1) - local_pos[valid_idx, 1] = (offset_v * e2).sum(-1) - local_pos[valid_idx, 2] = (offset_v * e3).sum(-1) - - # Rigid fallback offset (used for !valid Hs only) - offset_all = h_pos - p_pos - - # Stash numpy arrays for the forward path (the autograd Function - # converts to torch tensors lazily on the model's device). - self._h_idx = np.asarray(h_indices, dtype=np.int64) - self._h_parent_idx = p_arr - self._h_n1_idx = n1_arr - self._h_n2_idx = n2_arr - self._h_local_pos = local_pos.astype(np.float64) - self._h_frame_valid = valid - self._h_offset = offset_all # legacy field, used only when !valid + self._map_gaff2_atoms(mapping, atoms) - if self.verbose >= 1: - n_frame = int(valid.sum()) - n_fallback = int((~valid).sum()) - print( - f"[AmberTarget] H-attachment: {n_h} H atoms — " - f"{n_frame} via local-frame placement, " - f"{n_fallback} via rigid fallback" - ) - - # ------------------------------------------------------------------ - # Differentiable PyTorch placement (called from AmberTarget.forward) - # ------------------------------------------------------------------ - - def _place_hydrogens(self, heavy_omm_xyz_nm: torch.Tensor) -> torch.Tensor: - """ - Compute H positions from the current heavy-atom OpenMM-order tensor. - - Uses the local-frame data captured in :meth:`_build_h_attachment`: - for each H, build an orthonormal frame from (parent, n1, n2) and - place the H at its captured local-frame coordinates. Hs with - ``h_frame_valid=False`` fall back to ``parent + h_offset``. - - Differentiable: backward through this function distributes the - H force across the parent + n1 + n2 reference atoms via the exact - local-frame Jacobian (handled by PyTorch autograd). - - Parameters - ---------- - heavy_omm_xyz_nm : (n_omm_total, 3) tensor in nm, OpenMM atom order. - H slot values are ignored — they will be overwritten in the - returned tensor. - - Returns - ------- - h_xyz_nm : (n_H, 3) tensor in nm. Empty if no Hs. - """ - if self._h_idx is None or self._h_idx.size == 0: - return torch.zeros( - (0, 3), - dtype=heavy_omm_xyz_nm.dtype, - device=heavy_omm_xyz_nm.device, - ) - - device = heavy_omm_xyz_nm.device - dtype = heavy_omm_xyz_nm.dtype - # Lazily cache tensor views on the right device/dtype. + mapped = mapping[mapping >= 0] + missing_model = np.flatnonzero(mapping < 0) + missing_omm = sorted(set(range(len(atoms))) - set(mapped.tolist())) + duplicate = len(np.unique(mapped)) != len(mapped) if ( - getattr(self, "_h_tensors_dev", None) != device - or getattr(self, "_h_tensors_dtype", None) != dtype + len(atoms) != self._system.getNumParticles() + or duplicate + or len(missing_model) + or missing_omm ): - self._h_parent_idx_t = torch.as_tensor( - self._h_parent_idx, dtype=torch.long, device=device, # dtype-ok: H-parent atom index for indexing; PyTorch requires int64 - ) - # For invalid frames clamp neighbor indices to 0 so the gather is - # safe; the value is masked out by `where` below. - n1 = np.where(self._h_n1_idx >= 0, self._h_n1_idx, 0) - n2 = np.where(self._h_n2_idx >= 0, self._h_n2_idx, 0) - self._h_n1_idx_t = torch.as_tensor(n1, dtype=torch.long, device=device) # dtype-ok: neighbor atom index for indexing; PyTorch requires int64 - self._h_n2_idx_t = torch.as_tensor(n2, dtype=torch.long, device=device) # dtype-ok: neighbor atom index for indexing; PyTorch requires int64 - self._h_local_pos_t = torch.as_tensor( - self._h_local_pos, dtype=dtype, device=device, - ) - self._h_offset_t = torch.as_tensor( - self._h_offset, dtype=dtype, device=device, - ) - self._h_frame_valid_t = torch.as_tensor( - self._h_frame_valid, dtype=torch.bool, device=device, + model_examples = [ + f"{pdb.iloc[i]['chainid']}:{pdb.iloc[i]['resseq']}:" + f"{pdb.iloc[i]['name']}" + for i in missing_model[:5] + ] + omm_examples = [ + f"{atoms[i].residue.name}:{atoms[i].residue.id}:{atoms[i].name}" + for i in missing_omm[:5] + ] + raise ValueError( + "[AmberTarget] AMBER atom mapping is not one-to-one: " + f"unmatched model atoms={model_examples}, " + f"unmatched AMBER atoms={omm_examples}, duplicate matches={duplicate}. " + "Prepare matching atoms and protonation in TorchRef before " + "constructing the loss." ) - self._h_tensors_dev = device - self._h_tensors_dtype = dtype - - return _place_hydrogens_local_frame( - heavy_omm_xyz_nm, - self._h_parent_idx_t, - self._h_n1_idx_t, - self._h_n2_idx_t, - self._h_local_pos_t, - self._h_frame_valid_t, - self._h_offset_t, + self._model_to_omm = mapping + self._n_omm_atoms = len(atoms) + inverse = np.argsort(mapping) + self.register_buffer( + "_omm_to_model", + torch.as_tensor(inverse, dtype=get_int_dtype(), device=self._chem_model.device), ) - def _compose_full_omm_xyz( - self, - heavy_model_xyz_ang: torch.Tensor, - ) -> torch.Tensor: - """ - Build the full OpenMM-order position tensor (heavy + H) in nm. - - Heavy model atoms are scattered into their OpenMM slots via - ``self._model_to_omm``. Unmatched OpenMM heavy slots (e.g. - non-standard residues without a model match) are filled from - the construction-time ``_pos_buf`` snapshot so they're at least - consistent. H slots are filled by :meth:`_place_hydrogens`. + def _map_gaff2_atoms(self, mapping: np.ndarray, atoms: list) -> None: + """Match residue instances using heavy anchors, then names and H parents.""" + from scipy.spatial import cKDTree - Differentiable through ``heavy_model_xyz_ang`` — autograd routes - gradients on H slots back to their parent/n1/n2 reference atoms. - """ - device = heavy_model_xyz_ang.device - dtype = heavy_model_xyz_ang.dtype - n_omm = self._n_omm_atoms - - # Lazy tensorize index maps. - if ( - getattr(self, "_omm_tensors_dev", None) != device - or getattr(self, "_omm_tensors_dtype", None) != dtype - ): - valid_np = self._model_to_omm >= 0 - self._model_valid_t = torch.as_tensor( - valid_np, dtype=torch.bool, device=device, - ) - self._model_valid_model_idx_t = torch.as_tensor( - np.where(valid_np)[0], dtype=torch.long, device=device, # dtype-ok: valid-atom index (np.where) for indexing; PyTorch requires int64 - ) - self._model_valid_omm_idx_t = torch.as_tensor( - self._model_to_omm[valid_np], dtype=torch.long, device=device, # dtype-ok: model->OMM mapping index for indexing; PyTorch requires int64 + pdb = self._chem_model.pdb + keys = [ + tuple(row) + for row in pdb[["chainid", "resseq", "icode"]].itertuples( + index=False, name=None ) - # Construction-time snapshot for unmatched heavy slots + initial Hs - self._pos_buf_t = torch.as_tensor( - self._pos_buf, dtype=dtype, device=device, + ] + groups = {} + for i, key in enumerate(keys): + groups.setdefault(key, []).append(i) + residues = list(self._topology.residues()) + ligand_keys = self._gaff2_residue_keys + ligand_residues = ( + residues[len(residues) - len(ligand_keys) :] if ligand_keys else [] + ) + residue_map = {res.index: key for res, key in zip(ligand_residues, ligand_keys)} + xyz_nm = self._chem_model.xyz().detach().cpu().numpy() * 0.1 + names = pdb["name"].astype(str).str.strip().to_numpy() + elements = pdb["element"].astype(str).str.strip().str.upper().to_numpy() + heavy_rows = np.flatnonzero(~np.isin(elements, ["H", "D"])) + tree = cKDTree(xyz_nm[heavy_rows]) + for residue in residues: + if residue.index in residue_map: + continue + candidates = set() + for atom in residue.atoms(): + if self._is_hydrogen(atom): + continue + for local in tree.query_ball_point(self._tleap_pos_nm[atom.index], 0.005): + row = heavy_rows[local] + if elements[row] == atom.element.symbol.upper(): + candidates.add(keys[row]) + if len(candidates) != 1: + raise ValueError( + f"[AmberTarget] Cannot uniquely identify AMBER residue " + f"{residue.name} {residue.id} in TorchRef." + ) + residue_map[residue.index] = candidates.pop() + if len(set(residue_map.values())) != len(residue_map): + raise ValueError( + "[AmberTarget] Multiple AMBER residues match one model residue." ) - self._omm_tensors_dev = device - self._omm_tensors_dtype = dtype - # Start from the construction-time snapshot (provides values for - # unmatched heavy atoms and any non-frame-placed atom). Heavy and - # H slots will be overwritten below. - full = self._pos_buf_t.clone() - - heavy_model_xyz_nm = heavy_model_xyz_ang * 0.1 - heavy_matched = heavy_model_xyz_nm.index_select( - 0, self._model_valid_model_idx_t, - ) - full = full.index_copy(0, self._model_valid_omm_idx_t, heavy_matched) - - # Now derive H positions from the fully-populated heavy tensor. - if self._h_idx is not None and self._h_idx.size > 0: - h_xyz = self._place_hydrogens(full) - if not hasattr(self, "_h_idx_t_for_omm"): - self._h_idx_t_for_omm = torch.as_tensor( - self._h_idx, dtype=torch.long, device=device, # dtype-ok: H-atom index for indexing; PyTorch requires int64 + model_parents = {} + graph = self._chem_model.restraints.topology.atoms + for i, j in graph.bonds.indices.cpu().tolist(): + if elements[i] in {"H", "D"} and elements[j] not in {"H", "D"}: + model_parents[i] = j + elif elements[j] in {"H", "D"} and elements[i] not in {"H", "D"}: + model_parents[j] = i + omm_parents = {} + for a, b in self._topology.bonds(): + if self._is_hydrogen(a) and not self._is_hydrogen(b): + omm_parents[a.index] = b.index + elif self._is_hydrogen(b) and not self._is_hydrogen(a): + omm_parents[b.index] = a.index + for residue in residues: + rows = groups.get(residue_map[residue.index], []) + by_name = {names[i]: i for i in rows} + if len(by_name) != len(rows): + raise ValueError( + "[AmberTarget] Duplicate atom names within a model residue." ) - elif self._h_idx_t_for_omm.device != device: - self._h_idx_t_for_omm = self._h_idx_t_for_omm.to(device) - full = full.index_copy(0, self._h_idx_t_for_omm, h_xyz) + for atom in residue.atoms(): + row = by_name.get(atom.name) + if row is None and atom.name in {"H", "H1"}: + row = by_name.get("H1" if atom.name == "H" else "H") + if row is None and not self._is_hydrogen(atom): + candidates = [ + i + for i in rows + if mapping[i] < 0 + and elements[i] == atom.element.symbol.upper() + and ( + len(rows) == 1 + or np.linalg.norm(xyz_nm[i] - self._tleap_pos_nm[atom.index]) + < 0.005 + ) + ] + if len(candidates) == 1: + row = candidates[0] + if row is not None: + symbol = "H" if elements[row] == "D" else elements[row] + if symbol == atom.element.symbol.upper() and mapping[row] < 0: + mapping[row] = atom.index + for row in rows: + if mapping[row] >= 0 and row in model_parents: + expected_parent = mapping[model_parents[row]] + if omm_parents.get(mapping[row]) != expected_parent: + raise ValueError( + "[AmberTarget] Hydrogen attachment differs between " + f"TorchRef and AMBER: {residue.name} {names[row]}." + ) + # Equivalent hydrogens may use different numbering conventions. Only + # pair remaining H atoms attached to the same already-mapped parent. + used = set(mapping[mapping >= 0].tolist()) + for row in rows: + if mapping[row] >= 0 or row not in model_parents: + continue + parent = mapping[model_parents[row]] + choices = sorted( + a.index + for a in residue.atoms() + if a.index not in used and omm_parents.get(a.index) == parent + ) + if choices: + mapping[row] = choices[0] + used.add(choices[0]) - return full + def _compose_full_omm_xyz(self, model_xyz_ang: torch.Tensor) -> torch.Tensor: + """Gather all Cartesian model coordinates into OpenMM order and nm.""" + if model_xyz_ang.shape != (self._n_model_atoms, 3): + raise ValueError( + "[AmberTarget] Atom count changed; rebuild the target after " + "changing model topology." + ) + if self._omm_to_model.device != model_xyz_ang.device: + self._omm_to_model = self._omm_to_model.to(model_xyz_ang.device) + return model_xyz_ang.index_select(0, self._omm_to_model) * 0.1 # ------------------------------------------------------------------ # Step 5 — OpenMM Context @@ -1614,27 +999,10 @@ def _build_context(self, pos_nm: np.ndarray) -> None: # ------------------------------------------------------------------ def _energy(self, xyz_ang: torch.Tensor) -> torch.Tensor: - """AMBER14 energy for one conformation's heavy-atom coords. - - Parameters - ---------- - xyz_ang : torch.Tensor - ``(n_model_atoms, 3)`` heavy-atom coordinates in Å, in the order - of ``self._chem_model.pdb`` (the topology the system was built on). - - Returns - ------- - torch.Tensor - Scalar energy in kJ/mol (or kJ/mol/atom if ``normalize_by_atoms``). - Gradient flows to ``xyz_ang`` via OpenMM analytical forces - (heavy atoms direct) and via :meth:`_place_hydrogens` / PyTorch - autograd (H positions, redistributed onto their parent + - local-frame neighbors). - - Notes - ----- - Subclasses feed per-member coordinates here; the single-molecule - :meth:`forward` passes ``self._model.xyz()``. + """Evaluate all-atom Cartesian coordinates in Å, in chemistry-model order. + + Return a scalar in kJ/mol, divided by the atom count when normalization + is enabled. Gradients flow through the model's own coordinate wrapper. """ if self._context is None: raise RuntimeError( diff --git a/torchref/io/cif_readers.py b/torchref/io/cif_readers.py index fb88292a..221bb955 100644 --- a/torchref/io/cif_readers.py +++ b/torchref/io/cif_readers.py @@ -2297,6 +2297,15 @@ def _standardize_atoms(self, df: pd.DataFrame) -> pd.DataFrame: ), errors="coerce", ) + # The CCP4 energy type (NH1, OC, CH3, ...) keys the contact radii and the + # hydrogen-bond donor/acceptor roles; absent from eLBOW/Grade dictionaries. + type_cols = ["type_energy", "_chem_comp_atom.type_energy"] + if any(col in df.columns for col in type_cols): + result["type_energy"] = ( + self._extract_col(df, type_cols).astype(str).str.strip() + ) + else: + result["type_energy"] = "" # Include x,y,z if present (for ideal coordinates) for coord in ["x", "y", "z"]: diff --git a/torchref/model/context.py b/torchref/model/context.py index 06987680..1b9fc66e 100644 --- a/torchref/model/context.py +++ b/torchref/model/context.py @@ -51,11 +51,16 @@ class ModelContext(DeviceMixin): Verbosity level. strip_H : bool, default True Whether hydrogens were stripped on load. - exclude_H_from_sf : bool, default False - Whether hydrogens are excluded from structure-factor calculation. + hydrogens_in_xray : bool, default True + Whether hydrogens enter the structure-factor calculation. Restraints and the + non-bonded term see them either way; the bulk-solvent mask never does. add_hydrogens : bool, default False Generate hydrogens on load for residues that arrive without them. Ignored when ``strip_H`` is set, which removes them again. + hydrogen_mode : str, default "free" + How hydrogen rows are parametrised: ``"riding"`` (positions derived from the + parent heavy atoms each forward, not refined), ``"free"`` (ordinary refinable + atoms) or ``"none"`` (the table holds no hydrogens). initialized : bool, default False Whether a structure has been loaded. ``if model:`` tests this. @@ -79,8 +84,9 @@ class ModelContext(DeviceMixin): cif_path: Optional[str] = None verbose: int = 1 strip_H: bool = True - exclude_H_from_sf: bool = False + hydrogens_in_xray: bool = True add_hydrogens: bool = False + hydrogen_mode: str = "free" initialized: bool = False def copy(self) -> "ModelContext": @@ -110,8 +116,9 @@ def copy(self) -> "ModelContext": cif_path=self.cif_path, verbose=self.verbose, strip_H=self.strip_H, - exclude_H_from_sf=self.exclude_H_from_sf, + hydrogens_in_xray=self.hydrogens_in_xray, add_hydrogens=self.add_hydrogens, + hydrogen_mode=self.hydrogen_mode, initialized=self.initialized, ) diff --git a/torchref/model/model.py b/torchref/model/model.py index 6fe58176..56d9a8db 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -14,6 +14,8 @@ from typing import Dict, Iterable, List, Optional, Tuple, Union +import warnings + import gemmi import torch import torch.nn as nn @@ -136,6 +138,7 @@ def __init__( strip_H: bool = False, add_hydrogens: bool = False, cif_path: Optional[Union[str, List[str]]] = None, + hydrogens_in_xray: bool = True, ): """ Initialize an empty Model shell. @@ -161,6 +164,9 @@ def __init__( Restraint dictionary file(s); see the class docstring. :meth:`set_restraints_cif` can still change it after loading, but generation on load only sees the value given here. + hydrogens_in_xray : bool, optional + Whether hydrogens contribute to the structure factors. Default True. They + stay in the restraints either way; see :attr:`hydrogens_in_xray`. """ super().__init__() # Resolve dtype/device at call time (not import time) so a runtime @@ -181,6 +187,7 @@ def __init__( strip_H=strip_H, add_hydrogens=add_hydrogens, cif_path=cif_path, + hydrogens_in_xray=hydrogens_in_xray, ) # Submodules (created during load or load_state_dict) @@ -204,16 +211,62 @@ def __bool__(self): """ return self.ctx.initialized + @property + def hydrogens_in_xray(self) -> bool: + """Whether hydrogens enter ``get_iso()`` / ``get_aniso()`` and so Fcalc. + + Restraints and the non-bonded term see the hydrogens either way, and the + bulk-solvent mask never does. Default True. Changing it re-keys the + iso/aniso partition on the next access; no cache needs clearing. + """ + return self.ctx.hydrogens_in_xray + + @hydrogens_in_xray.setter + def hydrogens_in_xray(self, value: bool): + self.ctx.hydrogens_in_xray = bool(value) + @property def exclude_H_from_sf(self) -> bool: - """Drop H from ``get_iso()`` / ``get_aniso()`` (so from Fcalc) while - keeping them in the geometry and VDW restraints. Default False. + """Inverse of :attr:`hydrogens_in_xray`. + + .. deprecated:: + Use ``hydrogens_in_xray`` instead. """ - return self.ctx.exclude_H_from_sf + warnings.warn( + "exclude_H_from_sf is deprecated; use hydrogens_in_xray", + DeprecationWarning, + stacklevel=2, + ) + return not self.ctx.hydrogens_in_xray @exclude_H_from_sf.setter def exclude_H_from_sf(self, value: bool): - self.ctx.exclude_H_from_sf = bool(value) + warnings.warn( + "exclude_H_from_sf is deprecated; use hydrogens_in_xray", + DeprecationWarning, + stacklevel=2, + ) + self.ctx.hydrogens_in_xray = not bool(value) + + def _sf_atom_mask(self) -> Optional[torch.Tensor]: + """Atoms that enter Fcalc, or None when every atom does. + + Boolean ``(N,)`` over the atom table. Built lazily as the ``_heavy_atom_mask`` + buffer, which is dropped with the other per-atom caches when the atom set + changes, so it never outlives the table it was built for. + """ + if self.ctx.hydrogens_in_xray or self.pdb is None: + return None + if getattr(self, "_heavy_atom_mask", None) is None: + self.register_buffer( + "_heavy_atom_mask", + torch.tensor( + (self.pdb["element"].str.strip().str.upper() != "H").values, + dtype=torch.bool, + device=self.device, + ), + ) + return self._heavy_atom_mask # -- iso/aniso partition, derived on access --------------------------- # @@ -236,7 +289,7 @@ def _sf_partition(self): heavy = getattr(self, "_heavy_atom_mask", None) fp = ( (flag.data_ptr(), flag._version) if flag is not None else None, - bool(self.ctx.exclude_H_from_sf), + bool(self.ctx.hydrogens_in_xray), None if heavy is None else (heavy.data_ptr(), heavy._version), 0 if self.pdb is None else len(self.pdb), ) @@ -246,23 +299,12 @@ def _sf_partition(self): iso_mask = ~flag aniso_mask = flag - if self.ctx.exclude_H_from_sf and self.pdb is not None: - if getattr(self, "_heavy_atom_mask", None) is None: - self.register_buffer( - "_heavy_atom_mask", - torch.tensor( - (self.pdb["element"].str.strip() != "H").values, - dtype=torch.bool, - device=self.device, - ), - ) - # The mask is part of the key, so re-key after building it. - fp = (fp[0], fp[1], - (self._heavy_atom_mask.data_ptr(), - self._heavy_atom_mask._version), - fp[3]) - iso_mask = iso_mask & self._heavy_atom_mask - aniso_mask = aniso_mask & self._heavy_atom_mask + sf_atoms = self._sf_atom_mask() + if sf_atoms is not None: + # The mask is part of the key, so re-key after building it. + fp = (fp[0], fp[1], (sf_atoms.data_ptr(), sf_atoms._version), fp[3]) + iso_mask = iso_mask & sf_atoms + aniso_mask = aniso_mask & sf_atoms iso_idx = iso_mask.nonzero(as_tuple=True)[0] aniso_idx = aniso_mask.nonzero(as_tuple=True)[0] @@ -467,9 +509,8 @@ def get_scattering_params_iso(self): Notes ----- - ``n_iso_atoms`` honors ``exclude_H_from_sf``: when H exclusion is - active the isotropic count is the H-excluded count (mirroring - :meth:`get_iso`). + ``n_iso_atoms`` honors ``hydrogens_in_xray``: when hydrogens are excluded + the isotropic count is the heavy-atom count (mirroring :meth:`get_iso`). """ self._build_parametrization() idx = self._iso_indices @@ -488,9 +529,8 @@ def get_scattering_params_aniso(self): Notes ----- - ``n_aniso_atoms`` honors ``exclude_H_from_sf``: when H exclusion is - active the anisotropic count is the H-excluded count (mirroring - :meth:`get_aniso`). + ``n_aniso_atoms`` honors ``hydrogens_in_xray``: when hydrogens are excluded + the anisotropic count is the heavy-atom count (mirroring :meth:`get_aniso`). """ self._build_parametrization() idx = self._aniso_indices @@ -624,6 +664,7 @@ def _invalidate_atom_derived_caches(self) -> None: for name in self._ATOM_DERIVED_BUFFERS: if hasattr(self, name): delattr(self, name) + self._parametrization = None def load(self, reader, add_hydrogens: bool = None): """ @@ -690,7 +731,6 @@ def load(self, reader, add_hydrogens: bool = None): self.pdb["anisou_flag"].values, dtype=torch.bool, device=self.device ), ) - # Pre-compute integer indices for SF calculation (respects exclude_H_from_sf) self.xyz = MixedTensor( torch.tensor(self.pdb[["x", "y", "z"]].values, dtype=self.dtype_float), @@ -1143,12 +1183,11 @@ def copy(self): if module is not None and hasattr(module, "copy"): setattr(model_copy, module_name, module.copy()) - # A wrapper that borrows the coordinates carries that reference through its - # own ``copy``, so it still points at THIS model's ``xyz``. Re-point it, or - # the two models silently share coordinates and the copy is not independent. - for module in model_copy._modules.values(): - if module is not None and hasattr(module, "set_xyz_fn"): - module.set_xyz_fn(model_copy.xyz) + # Anything that borrows the coordinates -- the ADP node field, the restraints' + # pair-list maintenance -- carries the reference through its own ``copy`` and + # still points at THIS model's ``xyz``. Re-point it, or the two models silently + # share coordinates and the copy is not independent. + model_copy._repoint_coordinate_accessors() if self.ctx.verbose > 0: print(f"✓ Model copied successfully ({len(model_copy.pdb)} atoms)") @@ -1190,7 +1229,7 @@ def get_iso(self): Return per-atom parameters for the isotropic atom subset. Selects atoms whose ADP is a single scalar ``b``: ``~self.aniso_flag``, - intersected with the heavy-atom mask when ``exclude_H_from_sf`` is on. + intersected with the heavy-atom mask when ``hydrogens_in_xray`` is off. Returns ------- @@ -1257,17 +1296,20 @@ def parameters_of_types(self, types: Iterable[str]) -> List[nn.Parameter]: Returns ------- list of nn.Parameter - The ``refinable_params`` leaf for each requested type, in the - order the types were given. + Leaves for each requested type, in the order the types were given. + Coordinate wrappers may expose additional torsion and rotation leaves. """ out: List[nn.Parameter] = [] for t in types: wrapper = getattr(self, t, None) if wrapper is None: continue - rp = getattr(wrapper, "refinable_params", None) - if rp is not None: - out.append(rp) + if hasattr(wrapper, "optimization_parameters"): + out.extend(wrapper.optimization_parameters()) + else: + rp = getattr(wrapper, "refinable_params", None) + if rp is not None: + out.append(rp) return out def freeze(self, target: str): @@ -1853,7 +1895,7 @@ def get_aniso(self): Selects atoms whose ADP is the 6-element tensor ``u = (u11, u22, u33, u12, u13, u23)``: ``self.aniso_flag``, intersected - with the heavy-atom mask when ``exclude_H_from_sf`` is on. + with the heavy-atom mask when ``hydrogens_in_xray`` is off. Returns ------- @@ -2004,9 +2046,14 @@ def shake_coords(self, stddev: float): new_xyz = xyz + torch.normal( mean=0.0, std=stddev, size=xyz.shape, device=self.device ) - self.xyz = MixedTensor( - new_xyz, refinable_mask=self.xyz.refinable_mask, name="xyz" - ) + if hasattr(self.xyz, "with_values"): + # A riding wrapper keeps its frames; only the stored rows take the noise. + self.xyz = self.xyz.with_values(new_xyz) + else: + self.xyz = MixedTensor( + new_xyz, refinable_mask=self.xyz.refinable_mask, name="xyz" + ) + self._repoint_coordinate_accessors() def shake_adp(self, stddev: float): """ @@ -2134,7 +2181,9 @@ def hydrogenate(self, verbose: int = 0, optimize: bool = True) -> "Model": Hydrogen generation is template instantiation over the topology: each residue's library template is aligned onto the heavy atoms present and its hydrogens read off, and the bond graph decides how many hydrogens a parent can carry and which - of them have a free torsion. The original model is not modified. + of them have a free torsion. Missing HOH hydrogens use the water dictionary + geometry with a random initial orientation, controlled by ``torch.manual_seed``. + Existing hydrogen coordinates are retained. The original model is not modified. Parameters ---------- @@ -2210,6 +2259,8 @@ def state_dict(self, destination=None, prefix="", keep_vars=False): state[prefix + "strip_H"] = self.ctx.strip_H state[prefix + "cif_path"] = self.ctx.cif_path state[prefix + "altloc_pairs"] = self.ctx.altloc_pairs + state[prefix + "hydrogens_in_xray"] = self.ctx.hydrogens_in_xray + state[prefix + "hydrogen_mode"] = self.ctx.hydrogen_mode return state @@ -2358,11 +2409,37 @@ def _rebuild_wrappers_from_pdb(cls, instance, pdb, state_dict, saved_dtype, devi n_atoms = len(pdb) - instance.xyz = MixedTensor( - torch.tensor(pdb[["x", "y", "z"]].values, dtype=saved_dtype), - refinable_mask=state_dict.get("xyz.refinable_mask"), - name="xyz", - ) + xyz_values = torch.tensor(pdb[["x", "y", "z"]].values, dtype=saved_dtype) + if state_dict.get("xyz.h_row") is not None: + # A saved riding wrapper is recognised by its frame buffers, never by + # shape: its storage is (n_base, 3), a plain wrapper's (n_atoms, 3), and + # both are 2-D. The frames restore from the buffers, so no topology is + # needed here. + from torchref.model.riding_xyz import RidingXYZTensor + from torchref.topology.hydrogens import HydrogenFrames + + frames = HydrogenFrames.from_tensors( + state_dict["xyz.h_row"], + state_dict["xyz.parent_row"], + state_dict["xyz.n1_row"], + state_dict["xyz.n2_row"], + state_dict["xyz.frame_valid"], + state_dict.get("xyz.torsion_group"), + state_dict.get("xyz.rotation_group"), + ) + instance.xyz = RidingXYZTensor( + xyz_values, + frames, + refinable_mask=state_dict.get("xyz.refinable_mask"), + mask_in_base_space=True, + name="xyz", + ) + else: + instance.xyz = MixedTensor( + xyz_values, + refinable_mask=state_dict.get("xyz.refinable_mask"), + name="xyz", + ) instance.adp = cls._restore_adp_slot( "adp", state_dict, pdb, saved_dtype, instance.xyz ) @@ -2468,6 +2545,8 @@ def create_from_state_dict( strip_H = state_dict.pop("strip_H", True) cif_path = state_dict.pop("cif_path", None) altloc_pairs = state_dict.pop("altloc_pairs", []) + hydrogens_in_xray = state_dict.pop("hydrogens_in_xray", True) + hydrogen_mode = state_dict.pop("hydrogen_mode", None) instance = cls( dtype_float=saved_dtype, @@ -2475,7 +2554,13 @@ def create_from_state_dict( device=device, strip_H=strip_H, cif_path=cif_path, + hydrogens_in_xray=hydrogens_in_xray, ) + if hydrogen_mode is None: + # Older checkpoints: riding wrappers did not exist, so any hydrogens + # present were free parameters. + hydrogen_mode = "riding" if state_dict.get("xyz.h_row") is not None else "free" + instance.ctx.hydrogen_mode = hydrogen_mode instance.pdb = pdb instance.ctx.initialized = initialized @@ -2616,6 +2701,7 @@ def select(self, selection: str) -> "Model": device=self.device, strip_H=self.ctx.strip_H, cif_path=self.ctx.cif_path, + hydrogens_in_xray=self.ctx.hydrogens_in_xray, ) # ``index`` must be renumbered: the occupancy grouping below reads it. @@ -2636,17 +2722,21 @@ def select(self, selection: str) -> "Model": selected_model.register_buffer( "aniso_flag", self.aniso_flag[selection_mask].clone() ) - # Pre-compute SF indices (respects exclude_H_from_sf) - selected_model.xyz = MixedTensor( - self.xyz()[selection_mask].clone().detach(), - refinable_mask=( - self.xyz.refinable_mask[selection_mask] - if self.xyz.refinable_mask is not None - else None - ), - name="xyz", - ) + if hasattr(self.xyz, "select_rows"): + # Riding wrapper: frames are remapped, a hydrogen whose parent is cut + # becomes an ordinary row. + selected_model.xyz = self.xyz.select_rows(selection_mask) + else: + selected_model.xyz = MixedTensor( + self.xyz()[selection_mask].clone().detach(), + refinable_mask=( + self.xyz.refinable_mask[selection_mask] + if self.xyz.refinable_mask is not None + else None + ), + name="xyz", + ) selected_model.adp = PositiveMixedTensor( self.adp()[selection_mask].clone().detach(), @@ -2686,6 +2776,7 @@ def select(self, selection: str) -> "Model": selected_model.set_default_masks() selected_model.register_alternative_conformations() selected_model.ctx.initialized = True + selected_model.ctx.hydrogen_mode = self.ctx.hydrogen_mode if self.ctx.verbose > 0: print(f"Selected {n_selected}/{len(self.pdb)} atoms with '{selection}'") @@ -2819,6 +2910,228 @@ def get_centroid(self) -> torch.Tensor: return self.xyz().mean(dim=0) + # ------------------------------------------------------------------ + # Hydrogen parametrisation + # ------------------------------------------------------------------ + + @property + def hydrogen_mode(self) -> str: + """``"riding"``, ``"free"`` or ``"none"``; see :class:`ModelContext`.""" + return self.ctx.hydrogen_mode + + def hydrogen_frames(self): + """Which rows ride on which heavy atoms, for the current atom table. + + Read off the riding coordinate wrapper when one is installed, else derived + from the bond graph, which costs a restraint build the first time. + + Returns + ------- + HydrogenFrames + """ + if hasattr(self.xyz, "hydrogen_frames"): + return self.xyz.hydrogen_frames() + frames = getattr(self, "_hydrogen_frames", None) + if frames is not None and frames.n_hydrogens >= 0: + return frames + from torchref.topology.hydrogens import hydrogen_frames + + return hydrogen_frames(self.restraints.topology) + + def _repoint_coordinate_accessors(self) -> None: + """Make every borrowed coordinate accessor read the current ``xyz`` wrapper. + + The restraints keep ``xyz_fn`` for pair-list maintenance and the ADP node + field borrows the coordinates through ``set_xyz_fn``; after the wrapper slot + is replaced both would otherwise keep reading a dead module. + """ + restraints = self._restraints + if restraints is not None: + restraints._xyz_fn = self.xyz + restraints._adp_fn = self.adp + restraints._vdw_radii_fn = self.get_vdw_radii + for module in self._modules.values(): + if module is not None and hasattr(module, "set_xyz_fn"): + module.set_xyz_fn(self.xyz) + + def _complete_riding_waters(self, frames): + """Complete HOH residues once and remap frames and refinement selections.""" + from dataclasses import fields + + import numpy as np + + from torchref.topology.hydrogens import ( + HydrogenFrames, + augment_atom_table_with_maps, + hydrogen_frames, + plan_hydrogens, + ) + + if not self.ctx.add_hydrogens or self.ctx.strip_H: + return frames + if not self.pdb["resname"].str.strip().eq("HOH").any(): + return frames + restraints = self.restraints + dictionaries = { + key: value for key, value in restraints.cif_dict.items() if key == "HOH" + } + plan = plan_hydrogens(restraints.topology, dictionaries, self.xyz().detach()) + if plan.n_hydrogens == 0: + return frames + + generated = hydrogen_frames(restraints.topology, plan) + if frames is not None: + # Water groups include both existing and planned H atoms; keep custom + # frames for the rest of the table and give the water groups fresh IDs. + water = self.pdb["resname"].str.strip().eq("HOH").to_numpy() + keep = ~water[frames.parent_row] + take = water[generated.parent_row] + arrays = {} + for field in fields(HydrogenFrames): + existing = getattr(frames, field.name)[keep] + added = getattr(generated, field.name)[take].copy() + if field.name in ("torsion_group", "rotation_group"): + added[added >= 0] += int(existing.max(initial=-1)) + 1 + arrays[field.name] = np.concatenate((existing, added)) + generated = HydrogenFrames(**arrays) + + self.update_pdb() + augmented, old_rows, new_rows = augment_atom_table_with_maps( + self.pdb, plan, restraints.topology + ) + frames = generated.remap(old_rows).fill_planned_rows(new_rows) + source = torch.empty(len(augmented), dtype=torch.long, device=self.device) + old_index = torch.as_tensor(old_rows, device=self.device) + new_index = torch.as_tensor(new_rows, device=self.device) + source[old_index] = torch.arange(len(self.pdb), device=self.device) + source[new_index] = torch.as_tensor(plan.parent, device=self.device) + xyz = ( + self.xyz.to_mixed_tensor() + if hasattr(self.xyz, "to_mixed_tensor") + else self.xyz + ) + masks = { + "xyz": xyz.refinable_mask[source], + "occupancy": self.occupancy.get_refinable_atoms()[source], + } + from torchref.model.disorder_field import DisorderFieldTensor + + adp_fields = {} + for name in ("adp", "u"): + wrapper = getattr(self, name) + if isinstance(wrapper, DisorderFieldTensor): + adp_fields[name] = wrapper + else: + masks[name] = wrapper.refinable_mask[source] + if "u" in masks: + masks["u"][new_index] = False + gradients = { + name: getattr(self, name).refinable_params.requires_grad for name in masks + } + adp = self.adp().detach() + cell, spacegroup, links = self.cell, self.spacegroup, self.ctx.links + + def reader(): + return augmented, cell.data.cpu().numpy(), spacegroup + + reader.links = links + strip_h = self.ctx.strip_H + self.ctx.strip_H = False + self._restraints = None + try: + self.load(reader, add_hydrogens=False) + finally: + self.ctx.strip_H = strip_h + if "adp" in masks: + self.adp[old_index] = adp + for name, mask in masks.items(): + wrapper = getattr(self, name) + wrapper.update_refinable_mask(mask) + wrapper.refinable_params.requires_grad_(gradients[name]) + for name, field in adp_fields.items(): + field.anchor_atom = old_index[field.anchor_atom] + field.neighbor_list = field.neighbor_list[source] + field._full_shape = len(augmented) + field.set_xyz_fn(self.xyz) + setattr(self, name, field) + self._hydrogen_frames = frames + return frames + + def set_hydrogen_mode(self, mode: str, frames=None) -> "Model": + """Switch the hydrogen parametrisation of the current atom table. + + Parameters + ---------- + mode : str + ``"riding"``: hydrogen coordinates derive from their parents each forward; + rotatable groups retain shared torsion or orientation parameters. + Missing HOH hydrogens are completed only when ``ctx.add_hydrogens`` + is True and ``ctx.strip_H`` is False. With hydrogen generation disabled, + the atom table is unchanged. + ``"free"``: hydrogens are ordinary refinable atoms again. + ``"none"`` is a different atom table; use + :meth:`strip_hydrogens`. + frames : HydrogenFrames, optional + Riding frames for the current table; default :meth:`hydrogen_frames`. + Water frames are completed and row indices remapped if atoms are added. + + Returns + ------- + Model + Self, for chaining. + + Notes + ----- + Replaces the ``xyz`` wrapper, so any optimizer or ``LossState`` built over the + old parameters is stale; :meth:`Refinement.set_hydrogen_mode` does the + engine-side reset. The refinable set carries over row for row (a hydrogen + released to ``"free"`` follows its parent's mask). Existing atom coordinates + are preserved. Completing waters rebuilds the per-atom wrappers and + restraints; new hydrogens inherit their oxygen's refinement selections. + Water initialization follows ``torch.manual_seed`` and never runs in forward. + """ + from torchref.model.riding_xyz import RidingXYZTensor + + if not self.ctx.initialized: + raise RuntimeError("Load a structure before setting the hydrogen mode.") + if mode == "none": + raise ValueError( + "hydrogen_mode 'none' changes the atom table; use strip_hydrogens()" + ) + if mode not in ("riding", "free"): + raise ValueError(f"unknown hydrogen_mode {mode!r}") + + if mode == "riding": + frames = self._complete_riding_waters(frames) + if isinstance(self.xyz, RidingXYZTensor) and frames is None: + return self + if frames is None: + frames = self.hydrogen_frames() + if isinstance(self.xyz, RidingXYZTensor): + current = self.xyz.to_mixed_tensor() + else: + current = self.xyz + new_xyz = RidingXYZTensor.from_mixed_tensor(current, frames) + else: + if isinstance(self.xyz, RidingXYZTensor): + frames = self.xyz.hydrogen_frames() + new_xyz = self.xyz.to_mixed_tensor() + else: + new_xyz = self.xyz + + if new_xyz is not self.xyz: + # Pop first so the new wrapper registers as a fresh submodule. + self._modules.pop("xyz") + self.xyz = new_xyz + self._repoint_coordinate_accessors() + self._hydrogen_frames = frames + self.ctx.hydrogen_mode = mode + if hasattr(self, "reset_cache"): + self.reset_cache() + if self.ctx.verbose > 0: + print(f"Hydrogen mode: {mode} ({self.xyz})") + return self + def use_rigid_xyz(self) -> "Model": """ Swap ``self.xyz`` for a per-chain :class:`RigidXYZTensor`. @@ -2902,6 +3215,7 @@ def use_rigid_xyz(self) -> "Model": # submodule cleanly rather than colliding with the old one. self._rigid_original_xyz_container = self._modules.pop("xyz") self.xyz = rigid_xyz + self._repoint_coordinate_accessors() # Snapshot which groups were refinable BEFORE freezing them, so the # restore re-enables exactly those and leaves already-frozen ones alone. @@ -2953,7 +3267,13 @@ def restore_xyz_from_rigid(self, commit: bool = True) -> "Model": if commit: with torch.no_grad(): current = self.xyz().detach().clone() - new_xyz = MixedTensor(current, name="xyz", device=self.device) + stashed = getattr(self, "_rigid_original_xyz_container", None) + if stashed is not None and hasattr(stashed, "with_values"): + # A riding wrapper keeps its frames; a rigid motion leaves every + # local offset unchanged. + new_xyz = stashed.with_values(current) + else: + new_xyz = MixedTensor(current, name="xyz", device=self.device) self._modules.pop("xyz", None) self.xyz = new_xyz xyz_mask = getattr(self, "xyz_mask", None) @@ -2972,6 +3292,7 @@ def restore_xyz_from_rigid(self, commit: bool = True) -> "Model": if hasattr(self, "_rigid_original_xyz_container"): del self._rigid_original_xyz_container + self._repoint_coordinate_accessors() # Re-enable exactly the groups use_rigid_xyz() froze, so subsequent # per-atom / ADP refinement has parameters to optimize. diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index 9aac14e3..0c3299a1 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -806,6 +806,10 @@ def copy(self, detach: bool = True) -> "ModelFT": model_copy._parametrization = copy_module.deepcopy(self._parametrization) + # Borrowed coordinate accessors (restraints, ADP node field) still point at + # THIS model's wrappers after their own ``copy``; re-point them. + model_copy._repoint_coordinate_accessors() + # Don't share cached structure factors with the original. model_copy.reset_cache() # The iso/aniso partition is derived state, not a buffer, so it is not @@ -917,6 +921,8 @@ def create_from_state_dict( state_dict.pop("device", None) # Remove but don't use (use provided device) strip_H = state_dict.pop("strip_H", True) altloc_pairs = state_dict.pop("altloc_pairs", []) + hydrogens_in_xray = state_dict.pop("hydrogens_in_xray", True) + hydrogen_mode = state_dict.pop("hydrogen_mode", None) # Checkpoints written while the grid was stored state carry its buffers # ("_fft." prefixed, or flat in older ones). The size is adopted below only @@ -936,11 +942,15 @@ def create_from_state_dict( gridsize=explicit_gridsize, wavelength=wavelength, anomalous_threshold=anomalous_threshold, + hydrogens_in_xray=hydrogens_in_xray, ) instance.pdb = pdb instance.ctx.initialized = initialized instance.ctx.altloc_pairs = altloc_pairs + if hydrogen_mode is None: + hydrogen_mode = "riding" if state_dict.get("xyz.h_row") is not None else "free" + instance.ctx.hydrogen_mode = hydrogen_mode # The engine reads both off the context; nothing further to build. instance.spacegroup = spacegroup_str diff --git a/torchref/model/parameter_wrappers.py b/torchref/model/parameter_wrappers.py index e0ad7e2d..1991cabd 100644 --- a/torchref/model/parameter_wrappers.py +++ b/torchref/model/parameter_wrappers.py @@ -299,12 +299,24 @@ def __setitem__(self, key, value) -> None: self._set_values(key, value) + @property + def _storage_rows(self) -> int: + """Rows of the stored tensor. Equal to ``shape[0]`` unless a subclass derives + rows it does not store, in which case masks handed to the mutation methods + below are in storage space.""" + return 0 if self.fixed_values is None else int(self.fixed_values.shape[0]) + + def _storage_values(self) -> torch.Tensor: + """The stored rows assembled, in public units. Equal to ``forward()`` unless a + subclass derives extra rows.""" + return self.forward() + def _set_values(self, key, value: torch.Tensor) -> None: """Write already-cast values into the storage; override to re-encode. Rebuilds ``fixed_values`` and re-extracts ``refinable_params``. """ - current_full = self.forward().detach() + current_full = self._storage_values().detach() current_full[key] = value self.fixed_values = current_full.clone() @@ -445,13 +457,13 @@ def update_refinable_mask( If True, also re-baseline ``fixed_values`` to the current values. Default is False. """ - if new_mask.shape[0] != self.shape[0]: + if new_mask.shape[0] != self._storage_rows: raise ValueError( f"new_mask shape {new_mask.shape} must match " f"tensor shape {self.shape}" ) - current_full = self.forward().detach() + current_full = self._storage_values().detach() new_mask = self._normalize_refinable_mask(new_mask) self.refinable_mask = new_mask @@ -539,7 +551,7 @@ def refine( If True, re-baseline ``fixed_values`` to the current values first. Default is False. """ - current_full = self.forward().detach() + current_full = self._storage_values().detach() # Union of the current refinable mask with the new selection. new_mask = self.refinable_mask.clone() @@ -547,7 +559,7 @@ def refine( if isinstance(selection, torch.Tensor): if selection.dtype == torch.bool: if len(self.shape) > 1: - if selection.shape[0] != self.shape[0] or len(selection.shape) != 1: + if selection.shape[0] != self._storage_rows or len(selection.shape) != 1: raise ValueError( f"Boolean selection shape {selection.shape} must be 1D " f"matching first dimension {self.shape[0]} for multi-dimensional " @@ -600,7 +612,7 @@ def fix( If True (default), freeze at the current values; if False, the selected elements revert to the stored ``fixed_values``. """ - current_full = self.forward().detach() + current_full = self._storage_values().detach() # Current refinable mask minus the selection. new_mask = self.refinable_mask.clone() @@ -608,7 +620,7 @@ def fix( if isinstance(selection, torch.Tensor): if selection.dtype == torch.bool: if len(self.shape) > 1: - if selection.shape[0] != self.shape[0] or len(selection.shape) != 1: + if selection.shape[0] != self._storage_rows or len(selection.shape) != 1: raise ValueError( f"Boolean selection shape {selection.shape} must be 1D " f"matching first dimension {self.shape[0]} for multi-dimensional " @@ -859,7 +871,7 @@ def set(self, values: torch.Tensor, mask: torch.Tensor) -> None: ValueError If the shapes disagree or any value is non-positive. """ - if mask.shape[0] != self.shape[0]: + if mask.shape[0] != self._storage_rows: raise ValueError( f"Mask shape {mask.shape} must match tensor's first dimension {self.shape[0]}" ) @@ -1177,7 +1189,7 @@ def forward(self) -> torch.Tensor: def _set_values(self, key, value: torch.Tensor) -> None: """Set U-space values at ``key``; stored internally as Cholesky params.""" - current = self.forward().detach() + current = self._storage_values().detach() current[key] = value raw = self._u6_to_raw6(current) self.fixed_values = raw.clone() @@ -1191,14 +1203,14 @@ def fix(self, mask: torch.Tensor, freeze_at_current: bool = True): """Freeze rows, storing their current value in Cholesky space.""" if freeze_at_current: with torch.no_grad(): - raw = self._u6_to_raw6(self.forward()) + raw = self._u6_to_raw6(self._storage_values()) self.fixed_values[mask] = raw[mask] super().fix(mask, freeze_at_current=False) def refine(self, mask: torch.Tensor): """Make rows refinable, preserving their current value in Cholesky space.""" with torch.no_grad(): - raw = self._u6_to_raw6(self.forward()) + raw = self._u6_to_raw6(self._storage_values()) self.fixed_values[mask] = raw[mask] super().refine(mask) @@ -1216,12 +1228,12 @@ def update_refinable_mask( storage); convert to Cholesky parameters first, mirroring :meth:`PositiveMixedTensor.update_refinable_mask`. """ - if new_mask.shape[0] != self.shape[0]: + if new_mask.shape[0] != self._storage_rows: raise ValueError( f"new_mask shape {new_mask.shape} must match tensor shape {self.shape}" ) with torch.no_grad(): - current_raw = self._u6_to_raw6(self.forward()) + current_raw = self._u6_to_raw6(self._storage_values()) new_mask = self._normalize_refinable_mask(new_mask) self.refinable_mask = new_mask self.fixed_mask = ~new_mask diff --git a/torchref/model/riding_xyz.py b/torchref/model/riding_xyz.py new file mode 100644 index 00000000..12771a67 --- /dev/null +++ b/torchref/model/riding_xyz.py @@ -0,0 +1,816 @@ +"""Coordinate wrapper whose hydrogen rows ride on their parents. + +A :class:`RidingXYZTensor` looks like a :class:`~torchref.model.parameter_wrappers.MixedTensor` +over the whole atom table -- ``forward()`` returns ``(N, 3)``, masks are given in atom +space -- but only the non-riding rows are stored and refined. Each riding hydrogen is a +reference offset in a frame built from its parent and two reference heavy atoms +(:mod:`torchref.base.coordinates.local_frame`), rebuilt from the current heavy +coordinates on every forward. A force on a hydrogen therefore lands on the atoms that +carry it, which is the riding-hydrogen convention of Phenix and Refmac. + +Two index spaces meet here and callers must not mix them: everything public -- masks +passed to ``update_refinable_mask`` / ``refine`` / ``fix`` / ``set``, the result of +``forward()``, ``__getitem__`` / ``__setitem__`` -- is in FULL atom space, while the +inherited ``refinable_mask`` / ``fixed_values`` / ``refinable_params`` and the counts +from ``get_refinable_count()`` are in STORAGE space (the non-riding rows). Independent +angles are held in ``torsions`` and ``rotations`` and exposed together with the stored +coordinates by ``optimization_parameters()``. This is the +same contract :class:`~torchref.model.parameter_wrappers.OccupancyTensor` uses for its +collapsed groups. +""" + +from typing import Iterator, Optional, Union + +import numpy as np +import torch +import torch.nn as nn +from torchref.base.coordinates.local_frame import ( + frame_is_degenerate, + local_frame_coordinates, + place_local_frame, + rotate_vectors, +) +from torchref.model.parameter_wrappers import MixedTensor +from torchref.topology.hydrogens import HydrogenFrames + + +class _DerivedRowsMixin: + """Bookkeeping shared by wrappers that store some rows and derive the rest. + + Registers the row maps as buffers and maintains the storage-space gathers the + derivation needs. Subclasses call :meth:`_register_rows` after the parent + constructor has run and :meth:`_rebuild_row_cache` from ``_build_index_cache``. + """ + + #: Full-space row of every stored row, ``(N_base,)``. + base_row: torch.Tensor + #: Full-space row of every derived (riding) row, ``(H,)``. + h_row: torch.Tensor + #: Frame atoms per riding row, full-space, ``(H,)`` each. Absent -> clamped to 0 + #: and masked through ``frame_valid``. + parent_row: torch.Tensor + n1_row: torch.Tensor + n2_row: torch.Tensor + frame_valid: torch.Tensor + + def _register_rows(self, n_full: int, frames: HydrogenFrames, device) -> None: + h = np.asarray(frames.h_row, dtype=np.int64) + if len(h) and (h.min() < 0 or h.max() >= n_full): + raise ValueError("frames carry a hydrogen row outside the atom table") + is_riding = np.zeros(n_full, dtype=bool) + is_riding[h] = True + base = np.nonzero(~is_riding)[0] + for name in ("parent_row", "n1_row", "n2_row"): + rows = np.asarray(getattr(frames, name), dtype=np.int64) + if len(rows) and is_riding[rows[rows >= 0]].any(): + raise ValueError(f"{name} must reference stored rows, not riding ones") + + long = dict( + dtype=torch.int64, device=device + ) # dtype-ok: row index buffers; int64 index required + self.register_buffer("base_row", torch.as_tensor(base, **long)) + self.register_buffer("h_row", torch.as_tensor(h, **long)) + self.register_buffer( + "parent_row", + torch.as_tensor(np.asarray(frames.parent_row, dtype=np.int64), **long), + ) + self.register_buffer( + "n1_row", torch.as_tensor(np.asarray(frames.n1_row, dtype=np.int64), **long) + ) + self.register_buffer( + "n2_row", torch.as_tensor(np.asarray(frames.n2_row, dtype=np.int64), **long) + ) + self.register_buffer( + "frame_valid", + torch.as_tensor(np.asarray(frames.frame_valid, dtype=bool), device=device), + ) + self._rebuild_row_cache() + + def _rebuild_row_cache(self) -> None: + """Derive the storage-space gathers from the row buffers.""" + base = getattr(self, "base_row", None) + if base is None or getattr(self, "h_row", None) is None: + self._n_full = 0 + return + device = base.device + n_full = int(base.numel() + self.h_row.numel()) + self._n_full = n_full + full_to_base = torch.full( + (max(n_full, 1),), -1, dtype=torch.int64, device=device + ) # dtype-ok: index map; int64 + full_to_base[base] = torch.arange( + base.numel(), dtype=torch.int64, device=device + ) # dtype-ok: index map; int64 + self._parent_bidx = full_to_base[self.parent_row.clamp(min=0)].clamp(min=0) + self._n1_bidx = full_to_base[self.n1_row.clamp(min=0)].clamp(min=0) + self._n2_bidx = full_to_base[self.n2_row.clamp(min=0)].clamp(min=0) + # ``cat([base, derived])[gather]`` lays the full table out in one gather. + order = torch.empty( + n_full, dtype=torch.int64, device=device + ) # dtype-ok: gather index; int64 + order[base] = torch.arange( + base.numel(), dtype=torch.int64, device=device + ) # dtype-ok: gather index; int64 + order[self.h_row] = base.numel() + torch.arange( + self.h_row.numel(), + dtype=torch.int64, + device=device, # dtype-ok: gather index; int64 + ) + self._gather_order = order + + @property + def n_hydrogens(self) -> int: + """How many rows ride.""" + return 0 if getattr(self, "h_row", None) is None else int(self.h_row.numel()) + + @property + def n_base(self) -> int: + """How many rows are stored.""" + return ( + 0 if getattr(self, "base_row", None) is None else int(self.base_row.numel()) + ) + + def hydrogen_frames(self) -> HydrogenFrames: + """The frames as a CPU record, full-space rows.""" + return HydrogenFrames.from_tensors( + self.h_row, + self.parent_row, + self.n1_row, + self.n2_row, + self.frame_valid, + self.torsion_group, + self.rotation_group, + ) + + def _to_full_bool(self, selection) -> torch.Tensor: + """Any selection (bool mask, slice, indices) as a full-space bool mask.""" + if isinstance(selection, torch.Tensor) and selection.dtype == torch.bool: + if selection.ndim != 1 or selection.shape[0] != self._n_full: + raise ValueError( + f"Boolean selection shape {tuple(selection.shape)} must be " + f"({self._n_full},)" + ) + return selection.to(device=self.base_row.device) + mask = torch.zeros(self._n_full, dtype=torch.bool, device=self.base_row.device) + mask[selection] = True + return mask + + def _project(self, full_mask: torch.Tensor) -> torch.Tensor: + """Storage-space view of a full-space bool mask (riding rows dropped).""" + return full_mask.to(device=self.base_row.device, dtype=torch.bool)[ + self.base_row + ] + + def _expand_mask(self, base_mask: torch.Tensor) -> torch.Tensor: + """Full-space bool mask from a storage-space one; riding rows False.""" + out = torch.zeros(self._n_full, dtype=torch.bool, device=base_mask.device) + out[self.base_row] = base_mask + return out + + +class RidingXYZTensor(_DerivedRowsMixin, MixedTensor): + """Coordinates with riding hydrogen rows derived from their parents. + + Parameters + ---------- + initial_values : torch.Tensor, optional + Full atom table, Cartesian Angstroms, shape ``(N, 3)``. The riding rows fix the + local offsets; every other row is stored. None gives the empty shell that + ``load_state_dict`` fills. + frames : HydrogenFrames, optional + Which rows ride and on which atoms; required with ``initial_values``. + refinable_mask : torch.Tensor, optional + Boolean ``(N,)`` in full atom space, or + ``(N_base,)`` in storage space with ``mask_in_base_space``. None: every stored + row and every orientation refinable. + mask_in_base_space : bool, default False + Whether ``refinable_mask`` is already in storage space, as a saved state dict + hands it back. + requires_grad, dtype, device, name + As for :class:`~torchref.model.parameter_wrappers.MixedTensor`. + eps : float + Norm floor of the frame kernel, Angstroms. + + Notes + ----- + ``shape`` is the full ``(N, 3)``; ``refinable_mask``, ``fixed_values`` and + ``refinable_params`` are storage-space (``N_base`` rows). Frames whose reference + atoms are collinear or missing fall back to a Cartesian ``parent + offset``. + ``torsions`` holds one angle in radians per freely rotating bonded group; + ``rotations`` holds one Cartesian rotation vector in radians per unanchored + group (for example water). Both are :class:`MixedTensor` submodules. Selecting + a parent or any member hydrogen selects that group's orientation. Coordinate + assignment adopts the supplied orientation as the reference and zeros angles; + copies and checkpoints preserve the parameter values and references. + """ + + def __init__( + self, + initial_values: Optional[torch.Tensor] = None, + frames: Optional[HydrogenFrames] = None, + refinable_mask: Optional[torch.Tensor] = None, + *, + mask_in_base_space: bool = False, + requires_grad: bool = True, + dtype: Optional[torch.dtype] = None, + device: Optional[torch.device] = None, + name: Optional[str] = "xyz", + eps: float = 1e-8, + ): + self._eps = float(eps) + if initial_values is None: + super().__init__( + None, requires_grad=requires_grad, dtype=dtype, device=device, name=name + ) + for buffer in ("base_row", "h_row", "parent_row", "n1_row", "n2_row"): + self.register_buffer( + buffer, + torch.zeros( + 0, dtype=torch.int64, device=self.device + ), # dtype-ok: empty row-index buffer; int64 + ) + self.register_buffer( + "frame_valid", torch.zeros(0, dtype=torch.bool, device=self.device) + ) + self.register_buffer( + "local_offset", torch.zeros(0, 3, dtype=self.dtype, device=self.device) + ) + self.register_buffer( + "rigid_offset", torch.zeros(0, 3, dtype=self.dtype, device=self.device) + ) + self.register_buffer( + "virtual_reference", + torch.zeros(0, 3, dtype=self.dtype, device=self.device), + ) + self._initialize_orientations(HydrogenFrames.empty(), None) + self.register_load_state_dict_post_hook(self._after_load) + self._build_index_cache() + return + + if frames is None: + raise ValueError("frames are required with initial_values") + if initial_values.ndim != 2 or initial_values.shape[1] != 3: + raise ValueError( + f"initial_values must be (N, 3), got {tuple(initial_values.shape)}" + ) + dtype = dtype if dtype is not None else initial_values.dtype + device = device if device is not None else initial_values.device + values = initial_values.detach().to(dtype=dtype, device=device) + n_full = values.shape[0] + + is_riding = np.zeros(n_full, dtype=bool) + is_riding[np.asarray(frames.h_row, dtype=np.int64)] = True + base_rows = torch.as_tensor( + np.nonzero(~is_riding)[0], dtype=torch.int64, device=device + ) # dtype-ok: row index; int64 + + if refinable_mask is None: + base_mask = None + elif mask_in_base_space: + base_mask = refinable_mask.to(device=device, dtype=torch.bool) + else: + if refinable_mask.shape[0] != n_full: + raise ValueError( + f"refinable_mask has {refinable_mask.shape[0]} rows, table has {n_full}" + ) + base_mask = refinable_mask.to(device=device, dtype=torch.bool)[base_rows] + + super().__init__( + values.index_select(0, base_rows), + base_mask, + requires_grad=requires_grad, + dtype=dtype, + device=device, + name=name, + ) + self._register_rows(n_full, frames, device) + self.register_buffer( + "local_offset", torch.zeros(self.n_hydrogens, 3, dtype=dtype, device=device) + ) + self.register_buffer( + "rigid_offset", torch.zeros(self.n_hydrogens, 3, dtype=dtype, device=device) + ) + self.register_buffer("virtual_reference", torch.zeros_like(self.local_offset)) + full_mask = None if refinable_mask is None else refinable_mask.to(device=device) + if full_mask is not None and mask_in_base_space: + full_mask = self._expand_mask(full_mask) + self._initialize_orientations(frames, full_mask) + self.refresh_offsets(values) + self.register_load_state_dict_post_hook(self._after_load) + self._build_index_cache() + + # ------------------------------------------------------------------ + # Assembly + # ------------------------------------------------------------------ + + def _build_index_cache(self): + super()._build_index_cache() + self._rebuild_row_cache() + if hasattr(self, "torsion_group"): + self._rebuild_orientation_cache() + + def _initialize_orientations(self, frames, full_mask): + for name in ("torsion_group", "rotation_group"): + labels = np.asarray(getattr(frames, name)) + selected = labels >= 0 + unique, inverse = np.unique(labels[selected], return_inverse=True) + compact = np.full(len(labels), -1, dtype=np.int64) + compact[selected] = inverse + for group in unique: + members = np.flatnonzero(labels == group) + if len(np.unique(frames.parent_row[members])) != 1: + raise ValueError("An orientation group must share one parent") + if name == "torsion_group" and ( + (frames.n1_row[members] < 0).any() + or len(np.unique(frames.n1_row[members])) != 1 + ): + raise ValueError("A torsion group must share one bonded axis") + self.register_buffer(name, torch.as_tensor(compact, device=self.device)) + self._rebuild_orientation_cache() + requires_grad = self.refinable_params.requires_grad + self.torsions = MixedTensor( + torch.zeros( + self._torsion_parents.numel(), dtype=self.dtype, device=self.device + ), + requires_grad=requires_grad, + name="hydrogen_torsions", + ) + self.rotations = MixedTensor( + torch.zeros( + self._rotation_parents.numel(), 3, dtype=self.dtype, device=self.device + ), + requires_grad=requires_grad, + name="hydrogen_rotations", + ) + if full_mask is not None: + for wrapper, mask in zip( + (self.torsions, self.rotations), self._orientation_selection(full_mask) + ): + wrapper.update_refinable_mask(mask) + + def _rebuild_orientation_cache(self): + self._virtual_frame = (self.torsion_group >= 0) & (self.n2_row < 0) + self._has_virtual_frames = bool(self._virtual_frame.any()) + for kind in ("torsion", "rotation"): + labels = getattr(self, kind + "_group") + rows = (labels >= 0).nonzero(as_tuple=True)[0] + groups = labels[rows] + if rows.numel(): + order = torch.argsort(groups, stable=True) + sorted_groups = groups[order] + first = torch.cat( + [ + torch.ones(1, dtype=torch.bool, device=groups.device), + sorted_groups[1:] != sorted_groups[:-1], + ] + ) + first_rows = rows[order[first]] + parents = self.parent_row[first_rows] + else: + first_rows = rows + parents = self.parent_row[:0] + setattr(self, "_" + kind + "_h", rows) + setattr(self, "_" + kind + "_inverse", groups) + setattr(self, "_" + kind + "_parents", parents) + setattr(self, "_" + kind + "_first", first_rows) + + def _orientation_selection(self, full_mask): + selections = [] + for kind in ("torsion", "rotation"): + parents = getattr(self, "_" + kind + "_parents") + rows = getattr(self, "_" + kind + "_h") + groups = getattr(self, "_" + kind + "_inverse") + selected = full_mask[parents].to(torch.int32) + selected.index_add_(0, groups, full_mask[self.h_row[rows]].to(torch.int32)) + selections.append(selected > 0) + return selections + + def parameters(self, recurse: bool = True) -> Iterator[nn.Parameter]: + """Yield stored-coordinate and orientation leaves, including frozen shells.""" + return nn.Module.parameters(self, recurse=recurse) + + def optimization_parameters(self) -> list[nn.Parameter]: + """Return coordinate, torsion and rotation leaves for the xyz optimizer.""" + return [ + self.refinable_params, + self.torsions.refinable_params, + self.rotations.refinable_params, + ] + + @property + def _storage_rows(self) -> int: + return 0 if self.fixed_values is None else int(self.fixed_values.shape[0]) + + def _storage_values(self) -> torch.Tensor: + """The stored rows assembled, ``(N_base, 3)``.""" + return MixedTensor.forward(self) + + def evaluate( + self, + base_xyz: torch.Tensor, + torsions: Optional[torch.Tensor] = None, + rotations: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """Full coordinates from stored ones: the pure, differentiable part of forward. + + Parameters + ---------- + base_xyz : torch.Tensor + Stored Cartesian coordinates in Å, shape ``(N_base, 3)``. + torsions : torch.Tensor, optional + Full group angles in radians, shape ``(n_torsions,)``. Defaults to + the stored torsion parameters. + rotations : torch.Tensor, optional + Full group rotation vectors in radians, shape ``(n_rotations, 3)``. + Defaults to the stored orientation parameters. + + Returns + ------- + torch.Tensor + Shape ``(N, 3)``; riding rows placed from their frames. Gradients on any + row reach ``base_xyz`` through the frame Jacobian. + """ + if self.n_hydrogens == 0: + return base_xyz + local = self.local_offset + if self._torsion_h.numel(): + angles = self.torsions.forward() if torsions is None else torsions + cs = torch.stack((angles.cos(), angles.sin()), dim=-1)[ + self._torsion_inverse + ] + offsets = local[self._torsion_h] + x, y, z = offsets.unbind(-1) + c, sn = cs.unbind(-1) + turned = torch.stack((x, c * y - sn * z, sn * y + c * z), dim=-1) + local = local.index_copy(0, self._torsion_h, turned) + p = base_xyz.index_select(0, self._parent_bidx) + n2 = base_xyz.index_select(0, self._n2_bidx) + if self._has_virtual_frames: + n2 = torch.where( + self._virtual_frame[:, None], p + self.virtual_reference, n2 + ) + h = place_local_frame( + p, + base_xyz.index_select(0, self._n1_bidx), + n2, + local, + self.frame_valid, + self.rigid_offset, + eps=self._eps, + ) + if self._rotation_h.numel(): + vectors = self.rotations.forward() if rotations is None else rotations + offsets = rotate_vectors( + self.rigid_offset[self._rotation_h], vectors[self._rotation_inverse] + ) + h = h.index_copy(0, self._rotation_h, p[self._rotation_h] + offsets) + return torch.cat([base_xyz, h], dim=0).index_select(0, self._gather_order) + + def forward(self) -> torch.Tensor: + """The full ``(N, 3)`` table, riding rows derived from the stored ones.""" + return self.evaluate(self._storage_values()) + + @property + def shape(self): + """Full-space shape ``(N, 3)``.""" + if self.fixed_values is None: + return () + return (self._n_full, int(self.fixed_values.shape[1])) + + @property + def base_shape(self): + """Storage-space shape ``(N_base, 3)``.""" + return () if self.fixed_values is None else tuple(self.fixed_values.shape) + + @property + def full_refinable_mask(self) -> torch.Tensor: + """Refinable rows in full atom space; riding rows are never refinable.""" + return self._expand_mask(self.refinable_mask) + + # ------------------------------------------------------------------ + # Offsets + # ------------------------------------------------------------------ + + @torch.no_grad() + def refresh_offsets(self, full_xyz: Optional[torch.Tensor] = None) -> None: + """Re-derive every local offset from full-space coordinates. + + Parameters + ---------- + full_xyz : torch.Tensor, optional + Shape ``(N, 3)``; defaults to the current ``forward()``, which leaves the + coordinates unchanged while rebasing the angular parameters. Pass the + table after an external re-placement (a + torsion re-scan, a hydrogen written by ``__setitem__``) to adopt it. + + Notes + ----- + Frames that are geometrically degenerate at these coordinates are demoted to + the rigid fallback. Rewrites buffers in place, so the forward cache is + invalidated automatically. + """ + if self.n_hydrogens == 0: + return + if full_xyz is None: + full_xyz = self.forward() + full_xyz = full_xyz.detach().to(dtype=self.dtype, device=self.device) + p = full_xyz.index_select(0, self.parent_row) + n1 = full_xyz.index_select(0, self.n1_row.clamp(min=0)) + n2 = full_xyz.index_select(0, self.n2_row.clamp(min=0)) + h = full_xyz.index_select(0, self.h_row) + if self._has_virtual_frames: + first = self._torsion_first[self._torsion_inverse] + self.virtual_reference[self._torsion_h] = (h - p)[first] + n2 = torch.where( + self._virtual_frame[:, None], p + self.virtual_reference, n2 + ) + topological = (self.n1_row >= 0) & ((self.n2_row >= 0) | self._virtual_frame) + valid = topological & ~frame_is_degenerate(p, n1, n2) + local = local_frame_coordinates(p, n1, n2, h, eps=self._eps) + self.frame_valid.copy_(valid) + self.local_offset.copy_( + torch.where(valid.unsqueeze(-1), local, torch.zeros_like(local)) + ) + self.rigid_offset.copy_(h - p) + for orientation in (self.torsions, self.rotations): + orientation.refinable_params.zero_() + orientation.fixed_values.zero_() + orientation.reset_forward_cache() + + def set_hydrogen_positions(self, h_xyz: torch.Tensor) -> None: + """Adopt new positions for the riding rows, in ``h_row`` order, ``(H, 3)``.""" + full = self.forward().detach() + full[self.h_row] = h_xyz.to(dtype=self.dtype, device=self.device) + self.refresh_offsets(full) + + # ------------------------------------------------------------------ + # Mutation in full space + # ------------------------------------------------------------------ + + def _set_values(self, key, value: torch.Tensor) -> None: + """Write full-space values; stored rows update, riding rows become new offsets.""" + full = self.forward().detach() + full[key] = value + super()._set_values(slice(None), full.index_select(0, self.base_row)) + self.refresh_offsets(full) + + def set(self, values: torch.Tensor, mask: torch.Tensor) -> None: + """Write ``values`` at the True rows of a full-space ``mask``.""" + if mask.ndim != 1 or mask.shape[0] != self._n_full: + raise ValueError( + f"Mask shape {tuple(mask.shape)} must be ({self._n_full},)" + ) + mask = mask.to(device=self.device, dtype=torch.bool) + n_selected = int(mask.sum().item()) + if tuple(values.shape) != (n_selected, 3): + raise ValueError( + f"Values shape {tuple(values.shape)} doesn't match ({n_selected}, 3)" + ) + self._set_values(mask, values.to(dtype=self.dtype, device=self.device)) + + def update_refinable_mask( + self, new_mask: torch.Tensor, reset_refinable: bool = False + ): + """Repartition coordinates and orientations with a full- or storage-space mask.""" + full_mask = new_mask.to(device=self.device) + if new_mask.shape[0] == self._n_full: + new_mask = self._project(new_mask) + elif new_mask.shape[0] != self._storage_rows: + raise ValueError( + f"new_mask has {new_mask.shape[0]} rows; expected {self._n_full} " + f"(atom space) or {self._storage_rows} (storage space)" + ) + if full_mask.shape[0] != self._n_full: + full_mask = self._expand_mask(full_mask) + super().update_refinable_mask(new_mask, reset_refinable=reset_refinable) + for wrapper, mask in zip( + (self.torsions, self.rotations), self._orientation_selection(full_mask) + ): + wrapper.update_refinable_mask(mask, reset_refinable=reset_refinable) + + def refine( + self, selection: Union[slice, torch.Tensor, tuple], reset_values: bool = False + ): + """Add a full-space selection to the refinable set.""" + full_mask = self._to_full_bool(selection) + super().refine(self._project(full_mask), reset_values) + for wrapper, mask in zip( + (self.torsions, self.rotations), self._orientation_selection(full_mask) + ): + wrapper.refine(mask, reset_values) + + def fix( + self, + selection: Union[slice, torch.Tensor, tuple], + freeze_at_current: bool = True, + ): + """Remove a full-space selection from the refinable set.""" + full_mask = self._to_full_bool(selection) + super().fix(self._project(full_mask), freeze_at_current) + for wrapper, mask in zip( + (self.torsions, self.rotations), self._orientation_selection(full_mask) + ): + wrapper.fix(mask, freeze_at_current) + + def refine_all(self): + """Make every stored row and orientation refinable.""" + self.refine(torch.ones(self._n_full, dtype=torch.bool, device=self.device)) + + def fix_all(self, freeze_at_current: bool = True): + """Fix every stored row and orientation.""" + self.fix( + torch.ones(self._n_full, dtype=torch.bool, device=self.device), + freeze_at_current=freeze_at_current, + ) + + def update_fixed_values(self, new_values: torch.Tensor): + """Replace the stored rows' fixed buffer from a full-space ``(N, 3)`` table.""" + if tuple(new_values.shape) == self.shape: + new_values = new_values.index_select(0, self.base_row.to(new_values.device)) + super().update_fixed_values(new_values) + + # ------------------------------------------------------------------ + # Conversions and copies + # ------------------------------------------------------------------ + + def to_mixed_tensor(self) -> MixedTensor: + """Materialise as a plain per-atom wrapper; hydrogens follow their parent's mask.""" + mask = self.full_refinable_mask.clone() + mask[self.h_row] = mask[self.parent_row] + return MixedTensor( + self.forward().detach(), + mask, + requires_grad=self.refinable_params.requires_grad, + dtype=self.dtype, + device=self.device, + name=self.name, + ) + + @classmethod + def from_mixed_tensor( + cls, xyz: MixedTensor, frames: HydrogenFrames, **kwargs + ) -> "RidingXYZTensor": + """Wrap an existing per-atom coordinate tensor with riding frames.""" + return cls( + xyz.forward().detach(), + frames, + refinable_mask=xyz.refinable_mask, + requires_grad=xyz.refinable_params.requires_grad, + dtype=xyz.dtype, + device=xyz.device, + name=xyz.name, + **kwargs, + ) + + def with_values(self, full_xyz: torch.Tensor) -> "RidingXYZTensor": + """Same frames and mask, new coordinates ``(N, 3)``.""" + result = RidingXYZTensor( + full_xyz, + self.hydrogen_frames(), + refinable_mask=self.refinable_mask.clone(), + mask_in_base_space=True, + requires_grad=self.refinable_params.requires_grad, + dtype=self.dtype, + device=self.device, + name=self.name, + eps=self._eps, + ) + for name in ("torsions", "rotations"): + getattr(result, name).update_refinable_mask( + getattr(self, name).refinable_mask.clone() + ) + return result + + def select_rows(self, keep: torch.Tensor) -> "RidingXYZTensor": + """The wrapper over the rows where ``keep`` is True, frames remapped. + + A hydrogen whose parent is not kept becomes an ordinary stored row. + """ + keep_np = np.asarray(keep.detach().cpu().numpy(), dtype=bool) + old_to_new = np.full(len(keep_np), -1, dtype=np.int64) + old_to_new[keep_np] = np.arange(int(keep_np.sum())) + frames = self.hydrogen_frames().remap(old_to_new) + keep_t = torch.as_tensor(keep_np, device=self.device) + result = RidingXYZTensor( + self.forward().detach()[keep_t], + frames, + refinable_mask=self.full_refinable_mask[keep_t], + requires_grad=self.refinable_params.requires_grad, + dtype=self.dtype, + device=self.device, + name=self.name, + eps=self._eps, + ) + for name, labels in ( + ("torsions", frames.torsion_group), + ("rotations", frames.rotation_group), + ): + retained = torch.as_tensor( + np.unique(labels[labels >= 0]), device=self.device + ) + getattr(result, name).update_refinable_mask( + getattr(self, name).refinable_mask[retained] + ) + return result + + def clone(self) -> "RidingXYZTensor": + """Independent copy, offsets carried over bit for bit.""" + out = self.with_values(self.forward().detach()) + with torch.no_grad(): + out.local_offset.copy_(self.local_offset) + out.rigid_offset.copy_(self.rigid_offset) + out.frame_valid.copy_(self.frame_valid) + out.virtual_reference.copy_(self.virtual_reference) + out.torsions = self.torsions.copy() + out.rotations = self.rotations.copy() + return out + + def copy(self) -> "RidingXYZTensor": + """Alias for :meth:`clone`.""" + return self.clone() + + def clip(self, min_value=None, max_value=None) -> "RidingXYZTensor": + """Clip the full table; riding rows re-derive from the clipped heavy atoms.""" + full = self.forward().detach() + if min_value is not None: + full = torch.clamp(full, min=min_value) + if max_value is not None: + full = torch.clamp(full, max=max_value) + return self.with_values(full) + + def _after_load(self, module, incompatible_keys): + self._build_index_cache() + self.reset_forward_cache() + + def _load_from_state_dict(self, state_dict, prefix, *args, **kwargs): + if prefix + "torsion_group" not in state_dict: + # A checkpoint without group metadata describes fixed orientations. + hydrogen_rows = state_dict[prefix + "h_row"] + defaults = self.state_dict() + for name in ("torsion_group", "rotation_group"): + defaults[name] = torch.full_like(hydrogen_rows, -1) + defaults["virtual_reference"] = torch.zeros_like( + state_dict[prefix + "rigid_offset"] + ) + for name in ("torsions", "rotations"): + shape = (0,) if name == "torsions" else (0, 3) + empty = MixedTensor( + torch.empty(shape, dtype=self.dtype, device=self.device) + ) + defaults.update( + { + name + "." + key: value + for key, value in empty.state_dict().items() + } + ) + for name, value in defaults.items(): + if name in ( + "torsion_group", + "rotation_group", + "virtual_reference", + ) or name.startswith(("torsions.", "rotations.")): + state_dict.setdefault(prefix + name, value) + for name, buffer in list(self._buffers.items()): + saved = state_dict.get(prefix + name) + if saved is not None and (buffer is None or saved.shape != buffer.shape): + value = ( + torch.empty_like(saved, device=self.device) + if buffer is None + else buffer.new_empty(saved.shape) + ) + setattr(self, name, value) + saved_params = state_dict.get(prefix + "refinable_params") + if ( + saved_params is not None + and saved_params.shape != self.refinable_params.shape + ): + self.refinable_params = nn.Parameter( + self.refinable_params.new_empty(saved_params.shape), + requires_grad=self.refinable_params.requires_grad, + ) + for name in ("torsions", "rotations"): + saved = state_dict.get(prefix + name + ".fixed_values") + if saved is not None: + mask = state_dict[prefix + name + ".refinable_mask"].to(self.device) + setattr( + self, + name, + MixedTensor( + saved.to(device=self.device, dtype=self.dtype), + mask, + requires_grad=self.refinable_params.requires_grad, + name="hydrogen_" + name, + ), + ) + return super()._load_from_state_dict(state_dict, prefix, *args, **kwargs) + + def __repr__(self) -> str: + name_str = f"'{self.name}', " if self.name is not None else "" + return ( + f"RidingXYZTensor({name_str}shape={self.shape}, dtype={self.dtype}, " + f"device={self.device}, refinable={self.get_refinable_count()}, " + f"fixed={self.get_fixed_count()}, riding_h={self.n_hydrogens})" + ) + + +__all__ = ["RidingXYZTensor"] diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index d675ba9a..96ee5799 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -138,6 +138,7 @@ def __init__( scale_target: str = DEFAULT_SCALE_TARGET, aniso_selection: Optional[str] = None, add_hydrogens: bool = False, + hydrogens_in_xray: bool = True, ): """Initialize Refinement, fully if ``data_file`` and ``pdb`` are given. @@ -209,6 +210,9 @@ def __init__( add_hydrogens : bool, optional Generate missing hydrogens when loading the model. Default False. Hydrogens already present in the input are retained either way. + hydrogens_in_xray : bool, optional + Whether hydrogens contribute to the structure factors. Default True. They + take part in the restraints either way. """ super().__init__() # Refinement constructs its own submodules from file paths, so @@ -281,6 +285,7 @@ def __init__( anomalous_threshold=self.anomalous_threshold, add_hydrogens=add_hydrogens, cif_path=cif, + hydrogens_in_xray=hydrogens_in_xray, ) self.scaler = Scaler( verbose=self.verbose, device=self.device, nbins=self.nbins, @@ -328,6 +333,7 @@ def __init__( wavelength=self.wavelength, anomalous_threshold=self.anomalous_threshold, add_hydrogens=add_hydrogens, + hydrogens_in_xray=hydrogens_in_xray, # Before load, not after: generation on load reads this dictionary. cif_path=cif, # Apply the f'' (Bijvoet) term only when the data were loaded as @@ -746,6 +752,38 @@ def reset_loss_state(self) -> None: self._loss_state = None self._logger = None + def set_hydrogen_mode(self, mode: str) -> "Refinement": + """Switch the model's hydrogen parametrisation and reset the engine state. + + Parameters + ---------- + mode : str + ``"riding"`` or ``"free"``; see :meth:`Model.set_hydrogen_mode`. + + Returns + ------- + Refinement + Self, for chaining. + + Notes + ----- + The coordinate wrapper is replaced, so cached optimizers and the persistent + ``LossState`` are dropped and rebuilt on the next step. Call between macro + cycles, never inside one. + """ + n_atoms = len(self.model.pdb) + self.model.set_hydrogen_mode(mode) + if ( + len(self.model.pdb) != n_atoms + and getattr(self, "adp_target", None) is not None + ): + self._init_targets() + persistent = getattr(self, "_persistent_optimizers", None) + if persistent is not None: + persistent.clear() + self.reset_loss_state() + return self + def refine_scaler(self): """Refit the scaler against the current model. diff --git a/torchref/scaling/solvent.py b/torchref/scaling/solvent.py index 6841210f..87f906ee 100644 --- a/torchref/scaling/solvent.py +++ b/torchref/scaling/solvent.py @@ -143,6 +143,7 @@ def __init__( verbose=1, float_type=None, device=None, + ignore_hydrogens=True, ): """ Initialize SolventModel. @@ -171,6 +172,8 @@ def __init__( Initial phase offset in radians. verbose : int, default 1 Verbosity level. + ignore_hydrogens : bool, default True + Build the mask from heavy atoms only, whatever the model carries. float_type : torch.dtype, optional Float dtype. ``None`` (default) resolves at runtime to ``get_float_dtype()``, not a hard-wired ``torch.float32``. @@ -190,6 +193,9 @@ def __init__( self.solvent_radius = radius self.erosion_radius = erosion_radius self.optimize_phase = optimize_phase + # Heavy-atom radii already stand in for the hydrogens they carry, so a mask + # built over hydrogen rows too would exclude solvent twice. + self.ignore_hydrogens = bool(ignore_hydrogens) self._cache = TensorDict() # Empty initialization @@ -351,6 +357,16 @@ def get_solvent_mask(self): xyz = self.model.xyz() # (N_atoms, 3) vdw_radii = self.model.get_vdw_radii() # (N_atoms,) + if self.ignore_hydrogens: + # Heavy-atom radii are calibrated for masks built without hydrogens, so + # adding hydrogen spheres on top would exclude solvent twice. + heavy = torch.as_tensor( + (self.model.pdb["element"].str.strip().str.upper() != "H").values, + device=xyz.device, + ) + if not bool(heavy.all()): + xyz = xyz[heavy] + vdw_radii = vdw_radii[heavy] inv_frac = self.model.inv_fractional_matrix frac = self.model.fractional_matrix diff --git a/torchref/scripts/extract_ener_lib.py b/torchref/scripts/extract_ener_lib.py new file mode 100644 index 00000000..e0edc1b9 --- /dev/null +++ b/torchref/scripts/extract_ener_lib.py @@ -0,0 +1,97 @@ +"""Extract the per-atom-type table from the CCP4 energy library into a bundled CSV. + +The monomer library types every atom (``_chem_comp_atom.type_energy``: NH1, OC, CH3, +...) and ``ener_lib.cif`` says what each type is: its element, whether it donates or +accepts hydrogen bonds, and its van der Waals radius with and without the hydrogens it +normally carries. TorchRef reads that table from ``torchref/data/ener_lib_atoms.csv``; +this script regenerates the CSV from the library so the two cannot drift apart +unnoticed. + +Run as ``python -m torchref.scripts.extract_ener_lib [path/to/ener_lib.cif]``. Without a +path the library is fetched through the monomer-library manager. The hydrogen-bond +distance table is printed for inspection; it is the source of the contact-policy +defaults and is not bundled. +""" + +import csv +import sys +from pathlib import Path + +import gemmi + +from torchref import PATH_TORCHREF_DATA + +_ATOM_COLUMNS = ( + "type", + "weight", + "hb_type", + "vdw_radius", + "vdwh_radius", + "ion_radius", + "element", + "valency", + "sp", +) + +_OUT_COLUMNS = ("type", "element", "hb_type", "vdw_radius", "vdwh_radius", "ion_radius") + + +def _null(value: str) -> str: + return "" if value in (".", "?") else value + + +def extract(ener_lib: Path, out_csv: Path) -> int: + """Write the ``_lib_atom`` loop of ``ener_lib`` to ``out_csv``; return the row count.""" + block = gemmi.cif.read_file(str(ener_lib))[0] + table = block.find("_lib_atom.", list(_ATOM_COLUMNS)) + rows = [] + for row in table: + record = dict(zip(_ATOM_COLUMNS, (str(v) for v in row))) + vdw = _null(record["vdw_radius"]) + vdwh = _null(record["vdwh_radius"]) or vdw + rows.append( + { + "type": record["type"], + "element": record["element"], + "hb_type": record["hb_type"], + "vdw_radius": vdw, + "vdwh_radius": vdwh, + "ion_radius": _null(record["ion_radius"]), + } + ) + with open(out_csv, "w", newline="") as handle: + handle.write( + "# Per-energy-type atom properties from the CCP4 monomer library " + "ener_lib.cif (_lib_atom loop).\n" + "# hb_type: N neither, D donor, A acceptor, B both, " + "H hydrogen able to hydrogen-bond.\n" + "# vdw_radius: contact radius in Angstrom; vdwh_radius: radius to use when the " + "atom's own hydrogens are not modelled.\n" + "# Regenerate with python -m torchref.scripts.extract_ener_lib\n" + ) + writer = csv.DictWriter(handle, fieldnames=list(_OUT_COLUMNS)) + writer.writeheader() + writer.writerows(rows) + + hbond = block.find("_lib_hbond.", ["atom_type_1", "atom_type_2", "min", "dist"]) + print(f"{len(rows)} atom types written to {out_csv}") + print("hydrogen-bond distance table (type_1, type_2, well depth, distance):") + for row in hbond: + print(" ", " ".join(str(v) for v in row)) + return len(rows) + + +def main(argv=None) -> int: + argv = sys.argv[1:] if argv is None else argv + if argv: + ener_lib = Path(argv[0]) + else: + from torchref.topology.monomer.library import MonomerLibraryManager + + ener_lib = Path(MonomerLibraryManager(verbose=0).ensure_gemmi_base()) / "ener_lib.cif" + extract(ener_lib, Path(PATH_TORCHREF_DATA) / "ener_lib_atoms.csv") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/torchref/topology/atom_graph.py b/torchref/topology/atom_graph.py index 7d579567..05204315 100644 --- a/torchref/topology/atom_graph.py +++ b/torchref/topology/atom_graph.py @@ -118,6 +118,17 @@ class AtomGraph(DeviceMixin): planes : dict ``{n_atoms_in_plane: EdgeBlock}`` -- planes are ragged, so they are grouped by atom count the way the plane restraints already are. + energy_type : numpy.ndarray, optional + CCP4 energy type per atom (``NH1``, ``OC``, ``CH3``, ...), shape ``(N,)``, + ``''`` where the template does not say. Keys the contact radii and the + hydrogen-bond roles. + template_h_count : torch.Tensor, optional + How many hydrogens the atom carries in its template, shape ``(N,)``, + ``int8``; ``-1`` where unknown. Together with the bonded hydrogens actually + present this gives :meth:`implicit_h_count`. + hb_type : torch.Tensor, optional + Hydrogen-bond role code per atom, shape ``(N,)``, ``int8``; see the contact + policy for the enumeration. None until assigned. Notes ----- @@ -135,6 +146,9 @@ class AtomGraph(DeviceMixin): torsions: EdgeBlock chirals: EdgeBlock planes: Dict[int, EdgeBlock] = field(default_factory=dict) + energy_type: Optional[np.ndarray] = None + template_h_count: Optional[torch.Tensor] = None + hb_type: Optional[torch.Tensor] = None _adj_indptr: Optional[torch.Tensor] = field(default=None, repr=False) _adj_indices: Optional[torch.Tensor] = field(default=None, repr=False) @@ -171,8 +185,37 @@ def copy(self) -> "AtomGraph": torsions=self.torsions.copy(), chirals=self.chirals.copy(), planes={size: block.copy() for size, block in self.planes.items()}, + energy_type=None if self.energy_type is None else self.energy_type.copy(), + template_h_count=( + None if self.template_h_count is None else self.template_h_count.clone() + ), + hb_type=None if self.hb_type is None else self.hb_type.clone(), ) + def implicit_h_count(self) -> Optional[torch.Tensor]: + """Hydrogens each atom should carry but the table does not hold, ``(N,)``. + + ``template_h_count`` minus the bonded hydrogens actually present, floored at + zero; ``0`` where the template count is unknown. None when the graph carries + no template counts. What decides whether an atom takes its with-hydrogen + contact radius. + """ + if self.template_h_count is None: + return None + is_h = self.is_hydrogen + bonds = self.bonds.indices + present = torch.zeros(self.n_atoms, dtype=torch.int64, device=bonds.device) # dtype-ok: bincount output; int64 + if bonds.numel(): + heavy_of_h = torch.cat( + [bonds[is_h[bonds[:, 1]] & ~is_h[bonds[:, 0]], 0], + bonds[is_h[bonds[:, 0]] & ~is_h[bonds[:, 1]], 1]] + ) + if heavy_of_h.numel(): + present = torch.bincount(heavy_of_h, minlength=self.n_atoms) + known = self.template_h_count >= 0 + missing = self.template_h_count.to(torch.int64) - present + return torch.where(known, missing.clamp(min=0), torch.zeros_like(missing)) + def subset(self, remap: torch.Tensor, residue_remap: torch.Tensor) -> "AtomGraph": """The atoms ``remap`` keeps, with every edge set reindexed. @@ -196,16 +239,22 @@ def subset(self, remap: torch.Tensor, residue_remap: torch.Tensor) -> "AtomGraph if reduced.n_edges: planes[size] = reduced + keep_t = torch.as_tensor(keep, device=self.residue_of.device) return AtomGraph( name=self.name[keep], element=self.element[keep], altloc=self.altloc[keep], - residue_of=residue_remap[self.residue_of[torch.as_tensor(keep)]], + residue_of=residue_remap[self.residue_of[keep_t]], bonds=self.bonds.subset(remap), angles=self.angles.subset(remap), torsions=self.torsions.subset(remap), chirals=self.chirals.subset(remap), planes=planes, + energy_type=None if self.energy_type is None else self.energy_type[keep], + template_h_count=( + None if self.template_h_count is None else self.template_h_count[keep_t] + ), + hb_type=None if self.hb_type is None else self.hb_type[keep_t], ) def rebuild_adjacency(self) -> None: diff --git a/torchref/topology/build.py b/torchref/topology/build.py index 83962ae2..87970297 100644 --- a/torchref/topology/build.py +++ b/torchref/topology/build.py @@ -108,6 +108,48 @@ def _conformers( return [(names[altlocs == a], indices[altlocs == a]) for a in unique] +def _atom_types( + cols: Dict[str, np.ndarray], + nodes: Dict[str, np.ndarray], + template_key: np.ndarray, + comp_dict: Dict, +) -> Tuple[np.ndarray, np.ndarray]: + """Per-atom energy type and template hydrogen count, by name in the patched template. + + Returns + ------- + energy_type : numpy.ndarray + Shape ``(N,)``, ``''`` where the residue has no template or the atom is not + in it. + template_h_count : numpy.ndarray + Shape ``(N,)``, ``int8``; hydrogens the atom carries in its template, ``0`` + for template atoms with none (hydrogens included), ``-1`` where unknown. + """ + from torchref.topology.hydrogens import template_atom_types + + n_atoms = len(cols["name"]) + energy = np.full(n_atoms, "", dtype=" Tuple[np.ndarray, np.ndar return rotation, target_centre - rotation @ source_centre +def template_atom_types(component: Dict) -> Tuple[Dict[str, str], Dict[str, int]]: + """Energy type of every template atom, and the hydrogen count of every heavy one. + + Parameters + ---------- + component : dict + One residue's restraint sections as the CIF reader returns them; needs an + ``atoms`` section, and ``bonds`` for the counts. + + Returns + ------- + energy_type : dict + ``{atom name: CCP4 energy type}``; ``''`` where the dictionary carries none. + h_count : dict + ``{heavy atom name: number of hydrogens bonded to it in the template}``. Heavy + atoms with none are absent. + """ + atoms = component.get("atoms") + if atoms is None or len(atoms) == 0: + return {}, {} + ids = atoms["atom_id"].astype(str).str.strip().values.astype(str) + elements = np.char.upper( + atoms["type_symbol"].astype(str).str.strip().values.astype(str) + ) + if "type_energy" in atoms.columns: + types = atoms["type_energy"].astype(str).str.strip().values.astype(str) + types = np.where(np.isin(types, ["nan", ".", "?", "", "None"]), "", types) + else: + types = np.full(len(ids), "", dtype=" 0: + first = bonds["atom1"].astype(str).str.strip().values + second = bonds["atom2"].astype(str).str.strip().values + for a, b in zip(first, second): + if a not in is_h or b not in is_h: + continue + if is_h[a] and not is_h[b]: + h_count[b] = h_count.get(b, 0) + 1 + elif is_h[b] and not is_h[a]: + h_count[a] = h_count.get(a, 0) + 1 + return energy_type, h_count + + def _template(cif_dict: Dict, resname: str) -> Optional[Dict]: """Template atoms, hydrogen parents, ideal bond lengths and heavy adjacency. @@ -146,7 +195,7 @@ def _template(cif_dict: Dict, resname: str) -> Optional[Dict]: parent_of: Dict[str, str] = {} ideal_length: Dict[str, float] = {} heavy_adjacency: Dict[str, List[str]] = {} - h_count: Dict[str, int] = {} + _, h_count = template_atom_types(component) bonds = component.get("bonds") if bonds is not None and len(bonds) > 0: @@ -162,12 +211,10 @@ def _template(cif_dict: Dict, resname: str) -> Optional[Dict]: continue if is_h[ia] and not is_h[ib]: parent_of[a] = b - h_count[b] = h_count.get(b, 0) + 1 if np.isfinite(values[i]): ideal_length[a] = float(values[i]) elif is_h[ib] and not is_h[ia]: parent_of[b] = a - h_count[a] = h_count.get(a, 0) + 1 if np.isfinite(values[i]): ideal_length[b] = float(values[i]) elif not is_h[ia] and not is_h[ib]: @@ -337,6 +384,45 @@ def _place_group( Returns None when none applies, so the caller can count the hydrogen as undetermined rather than putting it somewhere arbitrary. """ + if ( + heavy_bonded == 0 + and len(template["heavy_names"]) == 1 + and str(template["elements"][template["id_to_index"][parent_name]]).upper() + == "O" + ): + index = template["id_to_index"] + origin = template["coords"][index[parent_name]] + # Randomness is confined to initialization and follows TorchRef's torch seed. + rotation, r = np.linalg.qr( + torch.randn(3, 3, dtype=get_float_dtype(), device="cpu").numpy() + ) + rotation = rotation * np.sign(np.diag(r))[None, :] + rotation[:, -1] *= np.linalg.det(rotation) + present_h = [h for h in template["h_names"] if h in name_to_row] + if present_h: + h = present_h[0] + source = template["coords"][index[h]] - origin + target = coords[name_to_row[h]] - parent_position + if np.linalg.norm(target) < 1e-8: + return None + source /= np.linalg.norm(source) + target /= np.linalg.norm(target) + t1, t2 = _orthonormal_frame(source) + m1 = rotation[:, 0] - target * (rotation[:, 0] @ target) + if np.linalg.norm(m1) < 1e-8: + m1, _ = _orthonormal_frame(target) + m1 /= np.linalg.norm(m1) + rotation = ( + np.column_stack([target, m1, np.cross(target, m1)]) + @ np.column_stack([source, t1, t2]).T + ) + offsets = np.array([template["coords"][index[h]] - origin for h in h_names]) + offsets = offsets @ rotation.T + return ( + parent_position + + offsets * (lengths / np.linalg.norm(offsets, axis=1))[:, None] + ) + covered = _template_covers_neighbours(template, parent_name, heavy_bonded) if covered: @@ -409,6 +495,14 @@ def plan_hydrogens(topology, cif_dict: Dict, xyz, verbose: int = 0) -> HydrogenP Returns ------- HydrogenPlan + Missing hydrogen rows with Cartesian positions in Å. + + Notes + ----- + HOH uses the bundled water dictionary when no water dictionary is supplied. + Waters without an existing hydrogen get a random reference orientation; + ``torch.manual_seed`` controls reproducibility. An existing O–H direction is + preserved when completing a partially hydrogenated water. """ coords = np.asarray(xyz.detach().cpu(), dtype=np.float64) residues = topology.residues @@ -428,6 +522,16 @@ def plan_hydrogens(topology, cif_dict: Dict, xyz, verbose: int = 0) -> HydrogenP "group", ) } + if "HOH" in set(residues.resname.astype(str)) and "HOH" not in cif_dict: + from pathlib import Path + + from torchref import PATH_TORCHREF_DATA + from torchref.topology.monomer.cif import read_cif + + cif_dict = dict(cif_dict) + cif_dict.update( + read_cif(str(Path(PATH_TORCHREF_DATA) / "monomer_library/h/HOH.cif")) + ) next_group = 0 n_unplaceable = 0 n_no_template = 0 @@ -436,11 +540,8 @@ def plan_hydrogens(topology, cif_dict: Dict, xyz, verbose: int = 0) -> HydrogenP resname = str(residues.resname[residue]).strip() template = _template(cif_dict, resname) if template is None: - # No usable template. In practice these are the single-atom residues -- - # waters and ions -- which the restraint dictionary omits because they carry - # no intra-residue geometry. They could not be hydrogenated anyway: one - # heavy atom gives no frame to orient a template against and no bond to - # rotate about, so a water's hydrogens would point somewhere arbitrary. + # Atoms without a dictionary cannot supply either bond geometry or + # hydrogen identities; leave those residues unchanged. n_no_template += 1 continue @@ -448,6 +549,9 @@ def plan_hydrogens(topology, cif_dict: Dict, xyz, verbose: int = 0) -> HydrogenP int(residues.atom_start[residue]), int(residues.atom_end[residue]) ) present = set(names[rows]) + h1_alias = "H" in template["h_names"] and "H1" not in template["h_names"] + if h1_alias and "H1" in present: + present.add("H") candidates = [h for h in template["h_names"] if h not in present] if not candidates: continue @@ -455,7 +559,8 @@ def plan_hydrogens(topology, cif_dict: Dict, xyz, verbose: int = 0) -> HydrogenP for altloc, conformer in _conformer_rows(rows, altlocs): name_to_row = {} for row in conformer: - name_to_row.setdefault(names[row], row) + name = "H" if h1_alias and names[row] == "H1" else names[row] + name_to_row.setdefault(name, row) # Hydrogens grouped by the parent they hang off, in name order so the cap # below takes a deterministic subset. @@ -469,9 +574,11 @@ def plan_hydrogens(topology, cif_dict: Dict, xyz, verbose: int = 0) -> HydrogenP parent_row = name_to_row[parent_name] parent_position = coords[parent_row] - heavy_rows, existing_h = _split_neighbours( - topology, parent_row, altloc - ) + heavy_rows, existing_h = _split_neighbours(topology, parent_row, altloc) + # A water-metal contact does not replace an O-H covalent bond or + # provide the water's orientational reference. + if resname == "HOH": + heavy_rows = np.zeros(0, dtype=np.int64) heavy_bonded = len(heavy_rows) element = str( template["elements"][template["id_to_index"][parent_name]] @@ -483,6 +590,19 @@ def plan_hydrogens(topology, cif_dict: Dict, xyz, verbose: int = 0) -> HydrogenP template_h = template["h_count"].get(parent_name, len(group)) template_heavy = len(template["heavy_adjacency"].get(parent_name, [])) extra_bonds = max(0, heavy_bonded - template_heavy) + energy_types = topology.atoms.energy_type + # Free NT* amines carry four neighbours. Peptide modifications + # retype N as NH1; explicit LINK/cap bonds also displace the free + # amine protonation even when no type modification is available. + if element == "N" and extra_bonds == 0 and energy_types is not None: + if str(energy_types[parent_row]).strip() in { + "NT", + "NT1", + "NT2", + "NT3", + "NT4", + }: + valence = 4 allowed = max( 0, min(valence - heavy_bonded, template_h - extra_bonds) - existing_h, @@ -769,9 +889,13 @@ def _rotate_about(vectors: np.ndarray, axis: np.ndarray, angle: float) -> np.nda __all__ = [ "HydrogenPlan", + "HydrogenFrames", "plan_hydrogens", "optimise_free_torsions", "augment_atom_table", + "augment_atom_table_with_maps", + "hydrogen_frames", + "template_atom_types", "STANDARD_VALENCE", "MAX_PLACEMENT_DISTANCE", "TORSION_SCAN_STEPS", @@ -804,20 +928,50 @@ def augment_atom_table(pdb, plan: HydrogenPlan, topology): pandas.DataFrame A new table with ``serial`` and ``index`` renumbered. """ + return augment_atom_table_with_maps(pdb, plan, topology)[0] + + +def augment_atom_table_with_maps(pdb, plan: HydrogenPlan, topology): + """:func:`augment_atom_table` plus the row maps the insertion implies. + + Parameters + ---------- + pdb : pandas.DataFrame + Atom table to extend. + plan : HydrogenPlan + topology : Topology + Supplies the residue partition the insertion points come from. + + Returns + ------- + augmented : pandas.DataFrame + The extended table, ``serial`` and ``index`` renumbered. + old_to_new : numpy.ndarray + New row of every old row, shape ``(N_old,)``. Existing rows are never dropped, + so every entry is valid. + plan_to_new : numpy.ndarray + New row of every planned hydrogen, shape ``(plan.n_hydrogens,)``. + """ import pandas as pd + n_old = len(pdb) if plan.n_hydrogens == 0: - return pdb.copy() + return pdb.copy(), np.arange(n_old, dtype=np.int64), np.zeros(0, dtype=np.int64) by_residue: Dict[int, List[int]] = {} for i, residue in enumerate(plan.residue.tolist()): by_residue.setdefault(residue, []).append(i) + old_to_new = np.full(n_old, -1, dtype=np.int64) + plan_to_new = np.full(plan.n_hydrogens, -1, dtype=np.int64) pieces = [] + offset = 0 for residue in range(topology.n_residues): start = int(topology.residues.atom_start[residue]) end = int(topology.residues.atom_end[residue]) pieces.append(pdb.iloc[start:end]) + old_to_new[start:end] = offset + np.arange(end - start) + offset += end - start members = by_residue.get(residue) if not members: @@ -833,10 +987,303 @@ def augment_atom_table(pdb, plan: HydrogenPlan, topology): if column in rows.columns: rows[column] = float("nan") pieces.append(rows) + plan_to_new[members] = offset + np.arange(len(members)) + offset += len(members) augmented = pd.concat(pieces, ignore_index=True) augmented["index"] = augmented.index.to_numpy(dtype=int) if "serial" in augmented.columns: augmented["serial"] = augmented.index.to_numpy(dtype=int) + 1 augmented.attrs = dict(pdb.attrs) - return augmented + return augmented, old_to_new, plan_to_new + + +@dataclass +class HydrogenFrames: + """Which atom-table rows are riding hydrogens, and the frame each one rides in. + + Row indices are into the atom table the frames were built for; ``-1`` marks an + absent atom. A hydrogen whose ``parent_row`` is ``-1`` is not a riding hydrogen at + all and is dropped by :meth:`remap`; one whose ``n1_row`` or ``n2_row`` is ``-1`` + keeps riding but with ``frame_valid`` False, so it translates rigidly with its + parent instead of turning with the frame. + + The frame is ``(parent, n1, n2)``: ``n1`` is the parent's first heavy neighbour, + ``n2`` its second, or -- for a parent with a single heavy neighbour, i.e. every + hydroxyl, thiol, amine and methyl -- a heavy neighbour of ``n1`` other than the + parent, so the hydrogen turns with the torsion about the ``n1-parent`` bond. + + Parameters + ---------- + h_row, parent_row, n1_row, n2_row : numpy.ndarray + ``int64`` rows, shape ``(H,)``. ``h_row`` is ``-1`` for a planned hydrogen + that has not been inserted into a table yet; :meth:`fill_planned_rows` sets it. + frame_valid : numpy.ndarray + Boolean ``(H,)``; False where the frame is incomplete. + torsion_group, rotation_group : numpy.ndarray, optional + Group labels, shape ``(H,)``; ``-1`` means no independent orientation. + Torsion groups rotate about the parent-to-n1 bond; rotation groups have + three rotational degrees of freedom in Cartesian space. Labels need not + be contiguous. Hydrogens in one group share their parent and orientation. + """ + + h_row: np.ndarray + parent_row: np.ndarray + n1_row: np.ndarray + n2_row: np.ndarray + frame_valid: np.ndarray + torsion_group: Optional[np.ndarray] = None + rotation_group: Optional[np.ndarray] = None + + def __post_init__(self) -> None: + """Fill absent orientation groups with the fixed-orientation sentinel.""" + for name in ("torsion_group", "rotation_group"): + value = getattr(self, name) + if value is None: + value = np.full(len(self.h_row), -1, dtype=np.int64) + value = np.asarray(value, dtype=np.int64) + if value.shape != self.h_row.shape: + raise ValueError(f"{name} must have shape {self.h_row.shape}") + setattr(self, name, value) + if ((self.torsion_group >= 0) & (self.rotation_group >= 0)).any(): + raise ValueError("A hydrogen cannot belong to both orientation types") + + @classmethod + def empty(cls) -> "HydrogenFrames": + """Frames for a model with no riding hydrogens.""" + z = np.zeros(0, dtype=np.int64) + return cls(z, z.copy(), z.copy(), z.copy(), np.zeros(0, dtype=bool)) + + @property + def n_hydrogens(self) -> int: + """How many hydrogens ride.""" + return len(self.h_row) + + @property + def n_planned(self) -> int: + """How many entries still await a row from :meth:`fill_planned_rows`.""" + return int((self.h_row < 0).sum()) + + def remap(self, old_to_new: np.ndarray) -> "HydrogenFrames": + """The frames over a reindexed table. + + Parameters + ---------- + old_to_new : numpy.ndarray + New row of each old row, ``-1`` where the atom was dropped. + + Returns + ------- + HydrogenFrames + Entries whose hydrogen or parent was dropped are removed; a lost ``n1`` or + ``n2`` leaves the entry with ``frame_valid`` False. Planned entries + (``h_row == -1``) are kept as planned. + """ + table = np.asarray(old_to_new, dtype=np.int64) + + def follow(rows: np.ndarray) -> np.ndarray: + out = np.full(len(rows), -1, dtype=np.int64) + present = rows >= 0 + out[present] = table[rows[present]] + return out + + h = follow(self.h_row) + h[self.h_row < 0] = -1 + parent = follow(self.parent_row) + n1 = follow(self.n1_row) + n2 = follow(self.n2_row) + keep = (parent >= 0) & ((h >= 0) | (self.h_row < 0)) + return HydrogenFrames( + h_row=h[keep], + parent_row=parent[keep], + n1_row=n1[keep], + n2_row=n2[keep], + frame_valid=self.frame_valid[keep] & (n1[keep] >= 0) & (n2[keep] >= 0), + torsion_group=np.where(n1[keep] >= 0, self.torsion_group[keep], -1), + rotation_group=self.rotation_group[keep], + ) + + def fill_planned_rows(self, rows: np.ndarray) -> "HydrogenFrames": + """Give the planned entries their table rows, in plan order. + + Parameters + ---------- + rows : numpy.ndarray + New row of each planned hydrogen, shape ``(n_planned,)``. + """ + rows = np.asarray(rows, dtype=np.int64) + planned = self.h_row < 0 + if int(planned.sum()) != len(rows): + raise ValueError( + f"{int(planned.sum())} planned hydrogens but {len(rows)} rows given" + ) + h = self.h_row.copy() + h[planned] = rows + return HydrogenFrames( + h, + self.parent_row.copy(), + self.n1_row.copy(), + self.n2_row.copy(), + self.frame_valid.copy(), + self.torsion_group.copy(), + self.rotation_group.copy(), + ) + + def sorted_by_row(self) -> "HydrogenFrames": + """The same frames ordered by ``h_row``.""" + order = np.argsort(self.h_row, kind="stable") + return HydrogenFrames( + self.h_row[order], + self.parent_row[order], + self.n1_row[order], + self.n2_row[order], + self.frame_valid[order], + self.torsion_group[order], + self.rotation_group[order], + ) + + def to_tensors(self, device=None) -> Dict[str, torch.Tensor]: + """Return frame and orientation arrays as tensors, keyed by field name.""" + return { + "h_row": torch.as_tensor( + self.h_row, dtype=torch.int64, device=device + ), # dtype-ok: row index; int64 required + "parent_row": torch.as_tensor( + self.parent_row, dtype=torch.int64, device=device + ), # dtype-ok: row index; int64 required + "n1_row": torch.as_tensor( + self.n1_row, dtype=torch.int64, device=device + ), # dtype-ok: row index; int64 required + "n2_row": torch.as_tensor( + self.n2_row, dtype=torch.int64, device=device + ), # dtype-ok: row index; int64 required + "frame_valid": torch.as_tensor( + self.frame_valid, dtype=torch.bool, device=device + ), + "torsion_group": torch.as_tensor(self.torsion_group, device=device), + "rotation_group": torch.as_tensor(self.rotation_group, device=device), + } + + @classmethod + def from_tensors( + cls, + h_row: torch.Tensor, + parent_row: torch.Tensor, + n1_row: torch.Tensor, + n2_row: torch.Tensor, + frame_valid: torch.Tensor, + torsion_group: Optional[torch.Tensor] = None, + rotation_group: Optional[torch.Tensor] = None, + ) -> "HydrogenFrames": + """Rebuild from the tensors :meth:`to_tensors` produced.""" + as_np = lambda t: np.asarray(t.detach().cpu().numpy(), dtype=np.int64) + return cls( + as_np(h_row), + as_np(parent_row), + as_np(n1_row), + as_np(n2_row), + np.asarray(frame_valid.detach().cpu().numpy(), dtype=bool), + None if torsion_group is None else as_np(torsion_group), + None if rotation_group is None else as_np(rotation_group), + ) + + def __repr__(self) -> str: + return ( + f"HydrogenFrames(n_hydrogens={self.n_hydrogens}, " + f"planned={self.n_planned}, rigid={int((~self.frame_valid).sum())})" + ) + + +def _frame_atoms(topology, parent_row: int, altloc: str) -> Tuple[int, int]: + """``(n1, n2)`` rows for a frame anchored on ``parent_row``; ``-1`` where absent. + + ``n1`` is the parent's first heavy neighbour in the conformer, ``n2`` its second, + or a heavy neighbour of ``n1`` other than the parent when the parent has only one. + """ + heavy, _ = _split_neighbours(topology, parent_row, altloc) + if len(heavy) == 0: + return -1, -1 + n1 = int(heavy[0]) + if len(heavy) >= 2: + return n1, int(heavy[1]) + grand, _ = _split_neighbours(topology, n1, altloc) + grand = grand[grand != parent_row] + return n1, (int(grand[0]) if len(grand) else -1) + + +def hydrogen_frames(topology, plan: Optional[HydrogenPlan] = None) -> HydrogenFrames: + """Riding frames for every hydrogen the table has, plus the ones a plan adds. + + Read off the bond graph, not off distances, so a stretched or predicted model + still frames each hydrogen on its bonded parent. + + Parameters + ---------- + topology : Topology + Connectivity of the table the frames index into. + plan : HydrogenPlan, optional + Hydrogens about to be inserted. Their entries carry ``h_row == -1`` until + :meth:`HydrogenFrames.fill_planned_rows` is given the rows the insertion made; + their parent and frame atoms are rows of the *current* table, to be carried + through :meth:`HydrogenFrames.remap` with everything else. + + Returns + ------- + HydrogenFrames + Deposited hydrogens first, in row order, then planned ones in plan order. A + hydrogen bonded to no heavy atom is left out: nothing can carry it. + """ + atoms = topology.atoms + is_h = atoms.is_hydrogen.cpu().numpy() + altlocs = np.char.strip(atoms.altloc.astype(str)) + + rows: List[Tuple[int, int, int, int]] = [] + for h in np.nonzero(is_h)[0].tolist(): + neighbours = atoms.neighbors(h).cpu().numpy() + heavy = neighbours[~is_h[neighbours]] + if len(heavy) == 0: + continue + parent = int(heavy[0]) + altloc = str(altlocs[h]) + n1, n2 = _frame_atoms(topology, parent, altloc) + rows.append((h, parent, n1, n2)) + + if plan is not None: + for k in range(plan.n_hydrogens): + parent = int(plan.parent[k]) + n1, n2 = _frame_atoms(topology, parent, str(plan.altloc[k]).strip()) + rows.append((-1, parent, n1, n2)) + + if not rows: + return HydrogenFrames.empty() + arr = np.array(rows, dtype=np.int64) + torsion = np.full(len(rows), -1, dtype=np.int64) + rotation = np.full(len(rows), -1, dtype=np.int64) + elements = np.char.upper(np.char.strip(atoms.element.astype(str))) + groups = {} + n_existing = len(rows) - (0 if plan is None else plan.n_hydrogens) + for i, (h, parent, _, _) in enumerate(rows): + altloc = str(altlocs[h]) if h >= 0 else str(plan.altloc[i - n_existing]).strip() + groups.setdefault((parent, altloc), []).append(i) + for group, ((parent, altloc), members) in enumerate(groups.items()): + heavy, _ = _split_neighbours(topology, parent, altloc) + residue = int(atoms.residue_of[parent]) + is_water = str(topology.residues.resname[residue]).strip() == "HOH" + if len(heavy) == 0 or is_water: + rotation[members] = group + elif len(heavy) == 1 and ( + (elements[parent] == "C" and len(members) == 3) + or elements[parent] in ("O", "S") + ): + # Planar amide NH2 groups also have one heavy neighbour, but their + # orientation is constrained by conjugation rather than freely rotatable. + torsion[members] = group + return HydrogenFrames( + h_row=arr[:, 0], + parent_row=arr[:, 1], + n1_row=arr[:, 2], + n2_row=arr[:, 3], + frame_valid=(arr[:, 2] >= 0) & (arr[:, 3] >= 0), + torsion_group=torsion, + rotation_group=rotation, + ) diff --git a/torchref/topology/monomer/modifications.py b/torchref/topology/monomer/modifications.py index 172e8276..ac70afac 100644 --- a/torchref/topology/monomer/modifications.py +++ b/torchref/topology/monomer/modifications.py @@ -13,10 +13,13 @@ peptide carbonyl carbon the intra-residue ``CA-C-O`` plus the link's ``CA-C-N`` and ``O-C-N`` sum to 360 deg only once ``DEL-OXT`` has been applied. -``_chem_mod_atom`` and ``_chem_mod_tree`` are deliberately ignored. The restraint -builders match library restraints against the atoms actually present in the model -and silently skip any whose atoms are missing, so adding or deleting atom -*definitions* changes nothing downstream; only the restraint sections matter. +Of ``_chem_mod_atom`` only the ``change`` rows are applied, and only to the energy +type and charge: ``DEL-HN1`` retypes the backbone ``N`` from the free-amine ``NT3`` to +the amide ``NH1``, which is what the contact radii and hydrogen-bond roles read. +Adding or deleting atom *definitions* is ignored, as is ``_chem_mod_tree``: the +restraint builders match library restraints against the atoms actually present in +the model and silently skip any whose atoms are missing, so those rows change nothing +downstream. """ from functools import lru_cache @@ -28,6 +31,7 @@ #: CIF category -> section name, matching :func:`read_link_definitions`. _CATEGORY_MAP = { + "chem_mod_atom": "atoms", "chem_mod_bond": "bonds", "chem_mod_angle": "angles", "chem_mod_tor": "torsions", @@ -38,6 +42,7 @@ #: Per section: the columns identifying a restraint, and the columns a #: ``change``/``add`` row may overwrite. _ATOM_COLUMNS = { + "atoms": ("atom_id",), "bonds": ("atom1", "atom2"), "angles": ("atom1", "atom2", "atom3"), "torsions": ("atom1", "atom2", "atom3", "atom4"), @@ -45,12 +50,18 @@ "chirals": ("atom_centre", "atom1", "atom2", "atom3"), } _VALUE_COLUMNS = { + "atoms": ("type_energy", "charge"), "bonds": ("value", "sigma"), "angles": ("value", "sigma"), "torsions": ("value", "sigma", "periodicity"), "planes": ("sigma",), "chirals": ("volume_sign",), } +#: Value columns that hold text, so they are not coerced to numbers. +_TEXT_COLUMNS = {"volume_sign", "type_energy"} +#: Sections where only ``change`` rows are meaningful; ``add``/``delete`` rows are +#: skipped rather than editing the atom list (see the module docstring). +_CHANGE_ONLY = {"atoms"} #: Sections whose row is meaningless without a target value, so an ``add`` row #: carrying only ``.`` placeholders is dropped rather than appended as NaN. _REQUIRES_VALUE = {"bonds", "angles", "torsions"} @@ -63,6 +74,8 @@ def _restraint_key(section: str, row: Mapping) -> tuple: outer pair, torsions on the atom quadruple in either direction, chirals on the centre plus the unordered substituents. Planes key on ``(plane_id, atom)``. """ + if section == "atoms": + return (str(row["atom_id"]).strip(),) if section == "bonds": return tuple(sorted((row["atom1"], row["atom2"]))) if section == "angles": @@ -137,7 +150,10 @@ def _standardize_mod_columns(df: pd.DataFrame, section: str) -> pd.DataFrame: "atom_id_3": "atom3", "atom_id_4": "atom4", "atom_id_centre": "atom_centre", - "atom_id": "atom", + # Planes name their atom ``atom_id``; the atoms section keeps that name. + **({} if section == "atoms" else {"atom_id": "atom"}), + "new_type_energy": "type_energy", + "new_charge": "charge", "new_value_dist": "value", "new_value_dist_esd": "sigma", "new_value_angle": "value", @@ -152,7 +168,10 @@ def _standardize_mod_columns(df: pd.DataFrame, section: str) -> pd.DataFrame: for column in _VALUE_COLUMNS[section]: if column not in df.columns: df[column] = pd.NA - elif column != "volume_sign": + elif column in _TEXT_COLUMNS: + text = df[column].astype(str).str.strip() + df[column] = text.where(~text.isin(["", ".", "?", "nan"]), pd.NA) + else: df[column] = pd.to_numeric(df[column], errors="coerce") columns.append(column) @@ -230,9 +249,14 @@ def apply_modifications( continue target = result.get(section) if target is None: + if section in _CHANGE_ONLY: + continue target = pd.DataFrame( columns=list(_ATOM_COLUMNS[section]) + list(_VALUE_COLUMNS[section]) ) + for column in _VALUE_COLUMNS[section]: + if column not in target.columns: + target[column] = pd.NA result[section] = _apply_section(target, mod_rows, section) return result @@ -253,6 +277,8 @@ def _apply_section( additions = [] for _, mod_row in mod_rows.iterrows(): function = mod_row["function"] + if section in _CHANGE_ONLY and function != "change": + continue key = _restraint_key(section, mod_row) positions = [p for p in by_key.get(key, []) if p not in dropped] diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index 61e99af0..0c49868f 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -305,8 +305,16 @@ def _load_cif_dictionaries(self, cif_path): self.missing_residues = [ res for res in self.unique_residues if res not in self.cif_dict ] + from pathlib import Path + from torchref import PATH_TORCHREF_DATA + additional_files = [ - find_cif_file_in_library(res) for res in self.missing_residues + ( + Path(PATH_TORCHREF_DATA) / "monomer_library/h/HOH.cif" + if res == "HOH" + else find_cif_file_in_library(res) + ) + for res in self.missing_residues ] for cif_file in additional_files: @@ -1195,7 +1203,21 @@ def copy(self): """ import copy - duplicate = copy.deepcopy(self) + # The coordinate and ADP accessors are borrowed from the model, not owned: + # duplicating them would hand the copy a third, orphaned parameter set (and + # deep-copying a wrapper with a cached forward fails on its graph tensor). + # They are carried across by reference; the owning model re-points them. + borrowed = ("_xyz_fn", "_adp_fn", "_vdw_radii_fn") + saved = {name: getattr(self, name, None) for name in borrowed} + for name in borrowed: + setattr(self, name, None) + try: + duplicate = copy.deepcopy(self) + finally: + for name, value in saved.items(): + setattr(self, name, value) + for name, value in saved.items(): + setattr(duplicate, name, value) duplicate._rebuild_entries() return duplicate From 60a7d2bdd7ed6f09ac695b68e23606bb0b6a836c Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Thu, 17 Sep 2026 23:34:07 +0200 Subject: [PATCH 170/250] Move relative dataset scaling into a shared sigma-weighted target --- docs/changelog.rst | 3 + docs/user_guide/scaling.rst | 44 + .../fcalc_benchmark/benchmark_cpu.py | 12 +- .../fcalc_benchmark/benchmark_worker.py | 5 +- .../benchmark_worker.py | 3 +- paper/probe_data_scale_objective.py | 251 +----- paper/probe_two_moment_collinearity.py | 4 +- tests/helpers/device_cases.py | 75 +- tests/integration/test_device_mixin.py | 5 +- .../integration/test_dtype_config_float64.py | 9 +- .../io/test_collection_stack_accessors.py | 12 +- tests/unit/io/test_crystfel_hkl.py | 5 +- tests/unit/io/test_data_scale_fit.py | 155 ---- tests/unit/io/test_intensity_accessors.py | 78 +- tests/unit/io/test_reflection_data_reindex.py | 33 +- tests/unit/io/test_select_reflection_data.py | 18 - ...test_collection_target_characterisation.py | 19 +- tests/unit/scaling/test_dataset_scaler.py | 272 +++++++ torchref/__init__.py | 23 +- torchref/base/targets/dataset_scaling.py | 51 ++ torchref/cli/collection_difference_refine.py | 30 +- torchref/cli/validate_ded.py | 29 +- torchref/io/__init__.py | 20 +- torchref/io/datasets/__init__.py | 3 + torchref/io/datasets/base.py | 20 +- torchref/io/datasets/collection.py | 260 +++--- torchref/io/datasets/reflection_data.py | 759 ++---------------- torchref/io/datasets/scaled_dataset.py | 160 ++++ torchref/io/metadata.py | 5 +- torchref/maps/difference_map.py | 4 +- torchref/refinement/targets/__init__.py | 5 +- .../refinement/targets/dataset_scaling.py | 42 + torchref/refinement/targets/xray/nll.py | 13 +- torchref/scaling/__init__.py | 10 +- torchref/scaling/collection_scaler.py | 9 +- torchref/scaling/dataset_scaler.py | 306 +++++++ 36 files changed, 1338 insertions(+), 1414 deletions(-) delete mode 100644 tests/unit/io/test_data_scale_fit.py create mode 100644 tests/unit/scaling/test_dataset_scaler.py create mode 100644 torchref/base/targets/dataset_scaling.py create mode 100644 torchref/io/datasets/scaled_dataset.py create mode 100644 torchref/refinement/targets/dataset_scaling.py create mode 100644 torchref/scaling/dataset_scaler.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 5b7de400..6c0d42a8 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,9 @@ Changelog Unreleased ---------- +- Joint sigma-weighted dataset scaling now uses centered corrections owned by ``DatasetScaler`` and exposed by ``ScaledDataset`` through ``DatasetCollection.scale()``, including scaled uncertainties, partial overlap, checkpointing and anomalous reflection identities. +- Remove parameter state, scale fitting and E-value conversion from ``ReflectionData``; use ``WilsonNormaliser`` for E values, observation attributes for full arrays and subset views instead of deprecated getters or ``data()``. +- Match batched mixed-solvent accumulation to individual model evaluation to reduce single-precision cancellation error. - ``torchref.phased-difference-map`` is now ``torchref.difference-map``, because it no longer defaults to a phased difference map. The default output is ``DELFWT``/``PHDELWT``, the inverse-variance-weighted amplitude difference on the **dark** model's phases -- the construction ``torchref.validate-ded`` correlates against, and the one the figure-4 CC was measured on. The writer computed that pair all along as ``WDF``/``PHIC_dark`` but foregrounded ``2mDFop-DFc``/``PHIC_diff`` instead, a phased *residual* that puts the light state's model phases into the observed amplitude and so biases the map toward the model under test. The headline map and the validation metric are now one object - ``-lm``/``--light-model`` is optional on ``torchref.difference-map``. A weighted difference map needs only the dark state: the amplitude is ``|Fo_light| - |Fo_dark|``, which ``DatasetCollection.scale`` puts on one scale with no model at all, and the phase comes from the dark model. Without a light model there is no mixed model, no ``--fraction`` and no joint model-to-data fit, and ``--fraction`` is rejected rather than ignored. Still required on ``torchref.difference-refine``, which refines it - The difference MTZ went from 33 columns (46 under ``--two-moment``) to 17, with the rest behind ``--all-columns``. The default file holds the difference map, the extrapolated map ``FWT``/``PHWT``, the observations and the flags; the gated set holds the alternative constructions -- the phased difference residuals, two further extrapolations, the intensity block. Standard CCP4 names throughout the default set, so Coot and CCP4 open both maps without being told which columns to use. The default path also runs one internal scale fit instead of three, and none at all with no light model diff --git a/docs/user_guide/scaling.rst b/docs/user_guide/scaling.rst index d1e1c9da..32f6e6cf 100644 --- a/docs/user_guide/scaling.rst +++ b/docs/user_guide/scaling.rst @@ -105,3 +105,47 @@ by a factor of 4 — and :math:`\mathbf{U}` the 6-parameter symmetric tensor stored on the scaler. The :math:`2\pi^2` is part of the definition, not a unit choice — dropping it makes the fitted ``U`` disagree with an ADP-convention ``U`` by that factor. + +Relative scaling of observed datasets +------------------------------------ + +``DatasetCollection.scale()`` jointly fits the observed datasets with +``DatasetScaler``. Each member receives an overall log scale and six quadratic +anisotropy coefficients in a shared normalized HKL basis. Coefficients are +centered over datasets during every forward evaluation, so no reference dataset +fixes the amplitude scale. The metadata reference still supplies the collection +cell and space group. + +The target profiles a shared amplitude per reflection using inverse propagated +variance weights. Both amplitudes and their uncertainties receive the same +positive correction; intensities and their uncertainties receive its square. +The consensus, centering and uncertainty propagation all carry gradients. +The fit uses work reflections shared by at least two datasets, excludes a +reflection held out in any participating dataset, and rejects disconnected or +rank-deficient overlap. Anomalous observations retain their signed identities. + +After calling ``collection.scale()``, retrieve datasets from the collection:: + + collection = DatasetCollection(device="cpu") + collection.add_dataset("dark", dark) + collection.add_dataset("light", light) + collection.scale() + scaled_light = collection["light"] + amplitudes = scaled_light.F + uncertainties = scaled_light.F_sigma + original_amplitudes = scaled_light.F_raw + +The collection owns one scaler. Its ``ScaledDataset`` members subclass +``ReflectionData`` and retain a strong reference to that scaler. Their direct +observation attributes, subset views, stack accessors and exports expose live +corrections. Copies and selections share the scaler; independent collections +have independent scalers. Raw input datasets are copied and remain unchanged. +Adding a dataset invalidates the joint fit; call ``scale()`` again. Repeated +``scale()`` calls otherwise reuse the parameter owner. Fitted parameters are +frozen; use ``collection.scaler.requires_grad_(True)`` for custom optimization. + +``ReflectionData`` holds raw observations and no optimization parameters. +Use ``WilsonNormaliser`` for normalized E values. Work, free and validation +observations are accessed through ``data.work``, ``data.free`` and +``data.validation``; full observations are available directly as ``data.F`` and +``data.F_sigma``. diff --git a/paper/figure3_performance/fcalc_benchmark/benchmark_cpu.py b/paper/figure3_performance/fcalc_benchmark/benchmark_cpu.py index f1098be3..c4ff7afa 100644 --- a/paper/figure3_performance/fcalc_benchmark/benchmark_cpu.py +++ b/paper/figure3_performance/fcalc_benchmark/benchmark_cpu.py @@ -1,14 +1,14 @@ #!/usr/bin/env python -import torch -from torchref import ReflectionData -from torchref import ModelFT +import os from time import time + +import torch from iotbx import pdb +from torchref import ModelFT, ReflectionData -import os _data_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "data") mtz_file = os.path.join(_data_dir, '1DAW.mtz') pdb_file = os.path.join(_data_dir, '1DAW.pdb') @@ -22,7 +22,7 @@ M = ModelFT(max_res=d_min, device=device,radius_angstrom=4.0).load_pdb(pdb_file) -hkl, _, _, _ = data() +hkl = data.hkl M(hkl, recalc=True) t_start = time() @@ -48,6 +48,4 @@ t_end = time() - print(f"Elapsed time for 10 runs of cctbx calculation: {t_end - t_start} seconds") - diff --git a/paper/figure3_performance/fcalc_benchmark/benchmark_worker.py b/paper/figure3_performance/fcalc_benchmark/benchmark_worker.py index 7e7474d8..caf9b716 100644 --- a/paper/figure3_performance/fcalc_benchmark/benchmark_worker.py +++ b/paper/figure3_performance/fcalc_benchmark/benchmark_worker.py @@ -22,6 +22,7 @@ n_threads = int(os.environ.get("TORCHREF_NUM_THREADS", 1)) import torch + from torchref import ModelFT, ReflectionData @@ -104,7 +105,7 @@ def run_benchmark(n_iterations: int, n_warmup: int, device_str: str = "cpu", data = ReflectionData(device=device).load_mtz(mtz_file) d_min = data.d_min M = ModelFT(max_res=d_min, device=device).load_pdb(pdb_file) - hkl, _, _, _ = data() + hkl = data.hkl n_atoms = M.xyz().shape[0] n_reflections = hkl.shape[0] @@ -116,7 +117,7 @@ def run_benchmark(n_iterations: int, n_warmup: int, device_str: str = "cpu", t.clone().detach().requires_grad_(True) if t is not None else None for t in aniso_ref ) - + def _forward(): sf, _ed = M.fft.compute_structure_factors(hkl, *iso, *aniso) diff --git a/paper/figure3_performance/refinement_cycle_benchmark/benchmark_worker.py b/paper/figure3_performance/refinement_cycle_benchmark/benchmark_worker.py index 1c9dedd3..d34e4f8a 100644 --- a/paper/figure3_performance/refinement_cycle_benchmark/benchmark_worker.py +++ b/paper/figure3_performance/refinement_cycle_benchmark/benchmark_worker.py @@ -25,6 +25,7 @@ n_threads = int(os.environ.get("TORCHREF_NUM_THREADS", 1)) import torch + from torchref.refinement import LBFGSRefinement @@ -118,7 +119,7 @@ def run_benchmark(n_iterations: int, n_warmup: int, device_str: str = "cpu", # Collect metadata n_atoms = len(refinement.model.pdb) - hkl, _, _, _ = refinement.reflection_data() + hkl = refinement.reflection_data.hkl n_reflections = hkl.shape[0] d_min = float(refinement.reflection_data.d_min) target_names = list(loss_state.targets.keys()) diff --git a/paper/probe_data_scale_objective.py b/paper/probe_data_scale_objective.py index 886f2367..c64a1b63 100644 --- a/paper/probe_data_scale_objective.py +++ b/paper/probe_data_scale_objective.py @@ -1,228 +1,65 @@ #!/usr/bin/env python -"""What does the data-to-data scale fit cost, and which objective generalises? +"""Measure joint dataset scaling on work and held-out reflections. -``DatasetCollection.scale()`` puts one dataset onto another. There is no model on either -side, so there is no model error for a sigma_A or Rice likelihood to account for, and the -only real choices are the weighting and which reflections the fit is allowed to see. - -The question it answers: **the leak.** The fit used to mask with -``ReflectionData.masks()`` -- validity only, with no work/free notion -- so the free -reflections went into the scale parameters, upstream of every target and therefore -upstream of every free-set number the pipeline reports. How much did that buy it, and -does removing it move the fitted scale? - -Scored on reflections the fit never saw, under two yardsticks applied identically to -every arm, so the comparison is not circular: - - R_data = sum|F - F_ref| / sum F_ref scale-free and interpretable - chi2 = mean[(F - F_ref)**2 / (s**2 + s_ref**2)] is the disagreement within error? - -``chi2`` is the one that can separate them: ``ls`` is entitled to win on ``R_data``, which -is what it optimises up to a constant. +Report amplitude agreement and propagated-variance residuals before and after +fitting centered overall and anisotropic corrections with DatasetCollection.scale. """ import argparse import json -import math -import sys from pathlib import Path import torch +from torchref.cli.collection_difference_refine import setup_dataset_collection -def build(dark_sf, light_sf, d_min, device): - """The figure-4 pair, loaded exactly as the difference CLI loads it, unscaled.""" - from torchref.cli.collection_difference_refine import setup_dataset_collection - - return setup_dataset_collection(dark_sf, light_sf, d_min, device) - - -def _reset(dc): - """Zero every fitted scale parameter and drop the corrected caches.""" - for _, ds in dc: - with torch.no_grad(): - if getattr(ds, "log_scale", None) is not None: - ds.log_scale.zero_() - if getattr(ds, "U_aniso", None) is not None: - ds.U_aniso.zero_() - ds._corrected_fp = None - ds._corrected_cache = None - ds._corrected_I_fp = None - ds._corrected_I_cache = None - - -def legacy_scale(dc): - """The pre-fix fit, copied verbatim for comparison: validity masks, unnormalised. - - A deliberate duplicate rather than a flag on the library method -- the point is to - measure what the old behaviour did, not to keep it selectable. - """ - ref_ds = dc._datasets[dc._reference_dataset] - to_scale = [ds for name, ds in dc if name != dc._reference_dataset] - params = [p for data in to_scale for p in data.parameters()] - [p.requires_grad_(True) for p in params] - opt = torch.optim.LBFGS(params, max_iter=100, line_search_fn="strong_wolfe") - ref_mask = ref_ds.masks() - ds_masks = [ds.masks() for ds in to_scale] - - def closure(): - opt.zero_grad() - loss = 0.0 - ref_F, _ = ref_ds.get_corrected_data() - for ds, m in zip(to_scale, ds_masks): - F, _ = ds.get_corrected_data() - cm = m & ref_mask - loss = loss + torch.sum((F[cm] - ref_F[cm]) ** 2) - loss.backward() - return loss - - for _ in range(10): - opt.step(closure) - [p.requires_grad_(False) for p in params] - - -def per_refl_chi2(dc, subset): - """Per-reflection chi-square contributions on ``subset``, in a fixed order. - - Returned rather than reduced so arms can be compared **paired**. The mask depends - only on the flags and the validity masks, never on the fitted scale, so the same - reflection sits at the same index in every arm -- which is what makes the pairing - valid. Unpaired means cannot resolve this comparison: the two datasets are different - structures, so most of chi2 is real difference and it cancels only when paired. - """ - from torchref.base.targets.xray_likelihoods import floor_sigma_obs - ref_name = dc._reference_dataset - other = [n for n, _ in dc if n != ref_name][0] - ref_ds, ds = dc[ref_name], dc[other] +def score(collection, subset: str) -> dict: + """Return symmetric amplitude disagreement and chi-square for one subset.""" + a, b = list(collection.values()) + mask = getattr(a, subset).mask & getattr(b, subset).mask with torch.no_grad(): - ref_F, ref_s = ref_ds.get_corrected_data() - F, s = ds.get_corrected_data() - m = getattr(ds, subset).mask & getattr(ref_ds, subset).mask - so, sr = floor_sigma_obs(s[m]), floor_sigma_obs(ref_s[m]) - return (((F[m] - ref_F[m]) ** 2) / (so**2 + sr**2)).cpu() - - -def paired_ci(a, b, n_boot=4000, seed=0): - """Bootstrap CI on ``mean(a - b)`` over reflections. Positive favours ``b``.""" - d = (a - b).numpy() - g = torch.Generator().manual_seed(seed) - n = len(d) - idx = torch.randint(0, n, (n_boot, n), generator=g).numpy() - means = d[idx].mean(axis=1) - lo, hi = sorted(means)[int(0.025 * n_boot)], sorted(means)[int(0.975 * n_boot)] - return float(d.mean()), float(lo), float(hi) - - -def score(dc, subset): - """``(R_data, chi2, n)`` between the two datasets on one held-out subset.""" - from torchref.base.targets.xray_likelihoods import floor_sigma_obs - - names = [n for n, _ in dc] - ref_name = dc._reference_dataset - other = [n for n in names if n != ref_name][0] - ref_ds, ds = dc[ref_name], dc[other] - - with torch.no_grad(): - ref_F, ref_s = ref_ds.get_corrected_data() - F, s = ds.get_corrected_data() - m = getattr(ds, subset).mask & getattr(ref_ds, subset).mask - fo, fr = F[m], ref_F[m] - so, sr = floor_sigma_obs(s[m]), floor_sigma_obs(ref_s[m]) - r = float((fo - fr).abs().sum() / fr.abs().sum().clamp(min=1e-30)) - chi2 = float((((fo - fr) ** 2) / (so**2 + sr**2)).mean()) - return r, chi2, int(m.sum()) - - -def fitted(dc): - other = [n for n, _ in dc if n != dc._reference_dataset][0] - ds = dc[other] - ls = float(ds.log_scale.detach().reshape(-1)[0]) - u = ds.U_aniso.detach().reshape(-1).tolist() - return ls, u - - -def main(): - ap = argparse.ArgumentParser(description=__doc__) - fig4 = Path(__file__).resolve().parent / "figure4_difference_refinement" - ap.add_argument("--dark-sf", default=str(fig4 / "data/8QL2-sf.cif")) - ap.add_argument("--light-sf", default=str(fig4 / "data/7YYZ-light.mtz")) - ap.add_argument("--dmin", type=float, default=2.2) - ap.add_argument("--device", default="cpu") - ap.add_argument("-o", "--out", default=None) - args = ap.parse_args() - - dev = torch.device(args.device) - dc = build(args.dark_sf, args.light_sf, args.dmin, dev) - - # A sigma-weighted arm was measured here and removed: it scored slightly better on - # held-out reflections for this one pair, but inverse-variance weighting collapses on - # a scale fit -- down-weighting the weak shells is what lets the scale run away in - # them -- and that failure mode was found across a panel. Unit weights stand. - arms = [ - ("legacy (validity mask, unnormalised)", lambda: legacy_scale(dc)), - ("ls (work set, normalised)", lambda: dc.scale()), - ] - - rows = [] - per_refl = {} - for label, fit in arms: - _reset(dc) - fit() - ls, u = fitted(dc) - rw, cw, nw = score(dc, "work") - rf, cf, nf = score(dc, "free") - per_refl[label] = { - "free": per_refl_chi2(dc, "free"), - "work": per_refl_chi2(dc, "work"), + fa, fb = a.F[mask], b.F[mask] + variance = a.F_sigma[mask].square() + b.F_sigma[mask].square() + valid = torch.isfinite(variance) & (variance > 0) + residual = fa[valid] - fb[valid] + denominator = (fa[valid].abs() + fb[valid].abs()).sum() + return { + "n": int(valid.sum()), + "R_symmetric": float(2 * residual.abs().sum() / denominator), + "chi2": float((residual.square() / variance[valid]).mean()), } - rows.append( - dict(arm=label, log_scale=ls, U_aniso=u, - R_work=rw, chi2_work=cw, n_work=nw, - R_free=rf, chi2_free=cf, n_free=nf) - ) - - print() - print("Data-to-data scale fit: fitted on WORK, scored on the held-out FREE set") - print("=" * 86) - print(f"{'arm':38s} {'log_scale':>10s} {'R_work':>8s} {'R_free':>8s} " - f"{'chi2_work':>10s} {'chi2_free':>10s}") - print("-" * 86) - for r in rows: - print(f"{r['arm']:38s} {r['log_scale']:10.5f} {r['R_work']:8.5f} " - f"{r['R_free']:8.5f} {r['chi2_work']:10.3f} {r['chi2_free']:10.3f}") - print("-" * 86) - print(f"n_work={rows[0]['n_work']} n_free={rows[0]['n_free']}") - print() - base = rows[0] - for r in rows[1:]: - d = r["log_scale"] - base["log_scale"] - print(f"{r['arm']:38s} d(log_scale) vs legacy = {d:+.6f} " - f"({100*(math.exp(d)-1):+.3f}% in scale)") - # --- the comparison that can actually resolve this: paired, per reflection ------ - labels = [lbl for lbl, _ in arms] - print() - print("Paired per-reflection chi2 difference (positive => the SECOND arm is better)") - print("=" * 86) - print(f"{'comparison':52s} {'set':6s} {'mean d':>10s} {'95% CI':>22s}") - print("-" * 86) - pairs = [(labels[0], labels[1])] - paired = [] - for a, b in pairs: - for subset in ("work", "free"): - m, lo, hi = paired_ci(per_refl[a][subset], per_refl[b][subset]) - sig = "" if lo <= 0.0 <= hi else " <-- CI excludes 0" - name = f"{a.split('(')[0].strip()} vs {b.split('(')[0].strip()}" - print(f"{name:52s} {subset:6s} {m:+10.5f} [{lo:+.5f}, {hi:+.5f}]{sig}") - paired.append(dict(a=a, b=b, subset=subset, mean=m, lo=lo, hi=hi)) - print("-" * 86) +def main() -> int: + """Fit the requested reflection pair and print or save its diagnostics.""" + parser = argparse.ArgumentParser(description=__doc__) + root = Path(__file__).resolve().parent / "figure4_difference_refinement" + parser.add_argument("--dark-sf", default=str(root / "data/8QL2-sf.cif")) + parser.add_argument("--light-sf", default=str(root / "data/7YYZ-light.mtz")) + parser.add_argument("--dmin", type=float, default=2.2) + parser.add_argument("--device", default="cpu") + parser.add_argument("-o", "--out") + args = parser.parse_args() + collection = setup_dataset_collection( + args.dark_sf, args.light_sf, args.dmin, torch.device(args.device) + ) + from torchref import DatasetCollection + + raw_collection = DatasetCollection(device=args.device, verbose=0) + for name, data in collection: + raw_collection.add_dataset(name, data.raw_data()) + collection = raw_collection + report = {"before": {s: score(collection, s) for s in ("work", "free")}} + collection.scale() + report["after"] = {s: score(collection, s) for s in ("work", "free")} + report["fit"] = collection.scaling_metrics + result = json.dumps(report, indent=2) + print(result) if args.out: - Path(args.out).write_text(json.dumps({"arms": rows, "paired": paired}, indent=2)) - print(f"written: {args.out}") + Path(args.out).write_text(result + "\n") return 0 if __name__ == "__main__": - sys.exit(main()) + raise SystemExit(main()) diff --git a/paper/probe_two_moment_collinearity.py b/paper/probe_two_moment_collinearity.py index 4c3af1de..780fa1dc 100644 --- a/paper/probe_two_moment_collinearity.py +++ b/paper/probe_two_moment_collinearity.py @@ -16,7 +16,7 @@ Two designs are compared, because they are the two places scale is fitted: -* **per-dataset** -- the ``log_scale`` + ``U_aniso`` that ``DatasetCollection.scale()`` +* **per-dataset** -- the centered log-scale and quadratic corrections that ``DatasetCollection.scale()`` fits on the light dataset alone. Free to shape the light data however it likes. * **shared** -- the ``CollectionScaler`` parameters, which are fitted jointly against dark and light. One column per parameter spanning *both* datasets, so a light-only template @@ -163,7 +163,7 @@ def main(): # # Two different response functions, because the two parameter sets act on opposite # sides of the residual. The shared scaler shapes the *model* intensity; the - # per-dataset log_scale / U_aniso shape the *observed* one. Using the model response + # per-dataset relative corrections shape the *observed* one. Using the model response # for both would give the dataset parameters an identically zero column and report # no leak at all. # diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index 1b52654f..1df3aef2 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -151,7 +151,6 @@ def _cell(device): return Cell(_CELL, device=device) - def _ctx(d): """A ModelContext whose cell and space group both live on ``d``.""" from torchref.model.context import ModelContext @@ -169,33 +168,61 @@ def _sffft_with_grid(d): return sf +def _dataset_scaler(device): + from pathlib import Path + + from torchref.io import ReflectionData + from torchref.scaling import DatasetScaler + + data = ReflectionData(device=device, verbose=0).load_mtz( + str(Path(__file__).parents[1] / "files" / "mtz" / "1DAW.mtz") + ) + return DatasetScaler({"a": data, "b": data}, device=device) + + +def _scaled_dataset(device): + from torchref.io import ScaledDataset + + scaler = _dataset_scaler(device) + return ScaledDataset(scaler.datasets["a"], scaler, "a") + + +def _dataset_scaling_target(device): + from torchref.refinement.targets import DatasetScalingTarget + + return DatasetScalingTarget(_dataset_scaler(device)) + + CASES: List[DeviceCase] = [ + DeviceCase("DatasetScaler", _dataset_scaler, "DatasetScaler"), + DeviceCase("ScaledDataset", _scaled_dataset, "ScaledDataset"), + DeviceCase("DatasetScalingTarget", _dataset_scaling_target, "DatasetScalingTarget"), DeviceCase("EdgeBlock", _edge_block, "EdgeBlock"), DeviceCase("AtomGraph", _atom_graph, "AtomGraph"), DeviceCase("Topology", _topology, "Topology"), DeviceCase("Cell", _cell, "Cell"), DeviceCase( "SpaceGroup", - lambda d: __import__( - "torchref.symmetry", fromlist=["SpaceGroup"] - ).SpaceGroup(_SG, device=d), + lambda d: __import__("torchref.symmetry", fromlist=["SpaceGroup"]).SpaceGroup( + _SG, device=d + ), "SpaceGroup", ), # D4: device implied by the cell; the SpaceGroup used to be built from the # raw (None) device argument and land on the process default instead. DeviceCase( "SfFFT_from_cell", - lambda d: __import__( - "torchref.model.sf_fft", fromlist=["SfFFT"] - ).SfFFT(_ctx(d), max_res=2.0), + lambda d: __import__("torchref.model.sf_fft", fromlist=["SfFFT"]).SfFFT( + _ctx(d), max_res=2.0 + ), "SfFFT", ), # D4: explicit device disagreeing with the supplied context. DeviceCase( "SfFFT_explicit_device", - lambda d: __import__( - "torchref.model.sf_fft", fromlist=["SfFFT"] - ).SfFFT(_ctx("cpu"), max_res=2.0, device=d), + lambda d: __import__("torchref.model.sf_fft", fromlist=["SfFFT"]).SfFFT( + _ctx("cpu"), max_res=2.0, device=d + ), "SfFFT", ), # The grid buffers are derived on first use; this case has them resolved. @@ -206,17 +233,15 @@ def _sffft_with_grid(d): ), DeviceCase( "SfDS_from_cell", - lambda d: __import__( - "torchref.model.sf_ds", fromlist=["SfDS"] - ).SfDS(_ctx(d)), + lambda d: __import__("torchref.model.sf_ds", fromlist=["SfDS"]).SfDS(_ctx(d)), "SfDS", ), # D1: tensor-free shells, whose tracker is the only thing to check. DeviceCase( "ScalerBase_empty", - lambda d: __import__( - "torchref.scaling", fromlist=["ScalerBase"] - ).ScalerBase(device=d), + lambda d: __import__("torchref.scaling", fromlist=["ScalerBase"]).ScalerBase( + device=d + ), "ScalerBase", tensor_free=True, ), @@ -310,9 +335,9 @@ def _sffft_with_grid(d): ), DeviceCase( "SpaceGroup", - lambda d: __import__( - "torchref.symmetry", fromlist=["SpaceGroup"] - ).SpaceGroup(_SG, device=d), + lambda d: __import__("torchref.symmetry", fromlist=["SpaceGroup"]).SpaceGroup( + _SG, device=d + ), "SpaceGroup", ), DeviceCase( @@ -340,17 +365,17 @@ def _sffft_with_grid(d): ), DeviceCase( "TensorMasks", - lambda d: __import__( - "torchref.utils", fromlist=["TensorMasks"] - ).TensorMasks(device=d), + lambda d: __import__("torchref.utils", fromlist=["TensorMasks"]).TensorMasks( + device=d + ), "TensorMasks", tensor_free=True, ), DeviceCase( "ReflectionData_empty", - lambda d: __import__( - "torchref.io", fromlist=["ReflectionData"] - ).ReflectionData(device=d), + lambda d: __import__("torchref.io", fromlist=["ReflectionData"]).ReflectionData( + device=d + ), "ReflectionData", ), DeviceCase( diff --git a/tests/integration/test_device_mixin.py b/tests/integration/test_device_mixin.py index d24c8fbe..01a1d2d5 100644 --- a/tests/integration/test_device_mixin.py +++ b/tests/integration/test_device_mixin.py @@ -15,6 +15,7 @@ import pytest import torch + def _load_model_ft(pdb_file, mtz_file): """Helper: load a ModelFT and matching reflection data on CPU. @@ -47,7 +48,7 @@ def test_modelft_cpu_gpu_cpu_sf_round_trip(sample_pdb_file, sample_mtz_file): match the original CPU result. """ model, data = _load_model_ft(sample_pdb_file, sample_mtz_file) - hkl, *_ = data() + hkl = data.hkl # ---- CPU leg --------------------------------------------------------- assert hkl.device.type == "cpu", "test setup: hkl should start on CPU" @@ -115,7 +116,7 @@ def test_modelft_cpu_only_recompute_after_to(sample_pdb_file, sample_mtz_file): it should match the first call. """ model, data = _load_model_ft(sample_pdb_file, sample_mtz_file) - hkl, *_ = data() + hkl = data.hkl fcalc_before = model(hkl).detach().clone() model.to("cpu") # idempotent move diff --git a/tests/integration/test_dtype_config_float64.py b/tests/integration/test_dtype_config_float64.py index 34738111..dccb37cb 100644 --- a/tests/integration/test_dtype_config_float64.py +++ b/tests/integration/test_dtype_config_float64.py @@ -11,7 +11,6 @@ import torch - @pytest.mark.unit def test_translation_phases_complex_dtype_float64(double_cpu): """Symmetry.phase_factors must honor the configured complex dtype.""" @@ -44,7 +43,7 @@ def test_scaler_binwise_mean_intensity_float64(double_cpu, sample_structure_pair scaler = Scaler(model=model, data=data, nbins=10, verbose=0) - hkl = data()[0] + hkl = data.hkl fcalc = model(hkl) assert fcalc.dtype == torch.complex128 @@ -59,11 +58,11 @@ def test_scaler_binwise_mean_intensity_float64(double_cpu, sample_structure_pair @pytest.mark.integration def test_occupancy_floor_density_matmul_float64(double_cpu, sample_structure_pair): """compute_density_at_positions hardcoded hkl.T.float(); matmul raised under float64.""" - from torchref.io import ReflectionData - from torchref.model.model_ft import ModelFT from torchref.experimental.targets.occupancy_floor_diagnostic import ( OccupancyFloorDiagnostic, ) + from torchref.io import ReflectionData + from torchref.model.model_ft import ModelFT model = ModelFT() model.load_cif(str(sample_structure_pair["model"])) @@ -75,7 +74,7 @@ def test_occupancy_floor_density_matmul_float64(double_cpu, sample_structure_pai positions = model.cell.cartesian_to_fractional(model.xyz()) assert positions.dtype == torch.float64 - hkl = data()[0] + hkl = data.hkl diagnostic = OccupancyFloorDiagnostic(model_dark=model, model_light=model) # Pre-fix this raised: float64 positions @ float32 hkl.T. diff --git a/tests/unit/io/test_collection_stack_accessors.py b/tests/unit/io/test_collection_stack_accessors.py index 8c535b8d..c263a645 100644 --- a/tests/unit/io/test_collection_stack_accessors.py +++ b/tests/unit/io/test_collection_stack_accessors.py @@ -3,7 +3,7 @@ Two things they could quietly get wrong, both invisible in the shape: * returning **raw** ``F``/``I`` instead of the scaled ones, which drops the per-dataset - ``log_scale``/``U_aniso`` that ``DatasetCollection.scale()`` fits; + joint scale parameters that ``DatasetCollection.scale()`` fits; * masking with the 2-way ``rfree_flags`` instead of the 3-way work/free/validation subsets, which lets validation reflections into the work set. @@ -32,9 +32,9 @@ def collection(mtz_dir): dc = DatasetCollection(verbose=0, device="cpu") dc.add_dataset("dark", a, set_as_reference=True) dc.add_dataset("light", b) + dc.scale(nsteps=1) with torch.no_grad(): - # A real inter-dataset scale difference: the whole point of the corrected path. - b.log_scale += 0.4 + dc.scaler.raw_parameters[1, 0] += 0.8 return dc @@ -61,7 +61,7 @@ def test_the_two_datasets_differ_after_scaling(self, collection): assertion above could detect the wrong accessor.""" stacked = collection.stack_F_obs() assert not torch.allclose(stacked[0], stacked[1]) - raw = torch.stack([collection[k].F for k in collection.keys()], dim=0) + raw = torch.stack([collection[k].F_raw for k in collection.keys()], dim=0) assert torch.allclose(raw[0], raw[1]), "raw amplitudes should be identical here" assert not torch.allclose(stacked, raw) @@ -71,8 +71,8 @@ def test_scaled_intensities_are_the_square_of_the_scaled_amplitude_factor( """Ties the two stacks together, so they cannot drift apart in scale.""" F = collection.stack_F_obs() I = collection.stack_I_obs() - raw_F = torch.stack([collection[k].F for k in collection.keys()], dim=0) - raw_I = torch.stack([collection[k].I for k in collection.keys()], dim=0) + raw_F = torch.stack([collection[k].F_raw for k in collection.keys()], dim=0) + raw_I = torch.stack([collection[k].I_raw for k in collection.keys()], dim=0) keep = (raw_F.abs() > 1e-6) & (raw_I.abs() > 1e-6) amp = (F[keep] / raw_F[keep]) ** 2 diff --git a/tests/unit/io/test_crystfel_hkl.py b/tests/unit/io/test_crystfel_hkl.py index 5abc3b01..ba9b5ece 100644 --- a/tests/unit/io/test_crystfel_hkl.py +++ b/tests/unit/io/test_crystfel_hkl.py @@ -122,12 +122,15 @@ def test_the_halves_align_into_a_collection(self, hkl_dir): a = _load(hkl_dir / "dark_half1.hkl") b = _load(hkl_dir / "dark_half2.hkl") + original_a, original_b = a.hkl.clone(), b.hkl.clone() dc = DatasetCollection(verbose=0, device="cpu") dc.add_dataset("half1", a, set_as_reference=True) dc.add_dataset("half2", b) assert dc.n_datasets == 2 - assert len(a.hkl) == len(b.hkl) == len(dc.hkl) + assert len(dc["half1"].hkl) == len(dc["half2"].hkl) == len(dc.hkl) + assert torch.equal(a.hkl, original_a) + assert torch.equal(b.hkl, original_b) # Both carry intensities, so an intensity target can run on the pair. assert dc["half1"].I is not None and dc["half2"].I is not None assert dc.stack_I_obs().shape == (2, len(dc.hkl)) diff --git a/tests/unit/io/test_data_scale_fit.py b/tests/unit/io/test_data_scale_fit.py deleted file mode 100644 index 1251c9f5..00000000 --- a/tests/unit/io/test_data_scale_fit.py +++ /dev/null @@ -1,155 +0,0 @@ -"""``DatasetCollection.scale()`` -- the data-to-data scale fit. - -The only fit in the library with no model on either side: it puts one dataset onto -another, both of them measurements of the same quantity. That is why it is least squares -and takes no objective at all -- there is no model error for a sigma_A row to account for, -and sigma weighting collapses on a scale fit. - -The load-bearing test here is the free-set one. That fit runs *upstream of every target*, -so a leak there compromises every free-set number the pipeline later reports, and no -downstream test would notice. -""" - -import pytest -import torch - - -@pytest.fixture -def pair(mtz_dir): - """Two copies of 1DAW as a reference + one dataset to scale onto it.""" - mtz = mtz_dir / "1DAW.mtz" - if not mtz.exists(): - pytest.skip("1DAW fixture not present") - - from torchref import ReflectionData - from torchref.io.datasets.collection import DatasetCollection - - ref = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - other = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - dc = DatasetCollection(verbose=0, device="cpu") - dc.add_dataset("ref", ref, set_as_reference=True) - dc.add_dataset("other", other) - return dc - - -def _fitted(dc): - ds = dc["other"] - return ds.log_scale.detach().clone(), ds.U_aniso.detach().clone() - - -@pytest.mark.integration -def test_scale_never_touches_the_free_set(pair): - """Corrupting the free reflections must not move the fitted parameters at all. - - Two *different* garbage values, because a single one could coincide with a - no-op: if the fit sees the free set, two different corruptions give two different - answers. ``torch.equal``, not ``allclose`` -- the free reflections must contribute - exactly nothing, not merely little. - - This fit used to mask with ``ReflectionData.masks()``, which is validity only - (``TensorMasks.__call__`` ANDs the validity masks and has no work/free notion), so - the free reflections went into the scale parameters via 10 x LBFGS(max_iter=100). - """ - dc = pair - free = dc["other"].free.mask - assert free.sum() > 0, "fixture has no free reflections; the test would be vacuous" - - results = [] - for filler in (3.0, 900.0): - d = dc["other"] - # Reset the parameters so each arm starts from the same place. - with torch.no_grad(): - d.log_scale.zero_() - d.U_aniso.zero_() - d.F[free] = filler - d.F_sigma[free] = filler - d._corrected_fp = None # drop the cached corrected view - d._corrected_cache = None - dc.scale() - results.append(_fitted(dc)) - - (ls_a, u_a), (ls_b, u_b) = results - assert torch.equal(ls_a, ls_b), ( - f"log_scale moved when only the FREE reflections changed: {ls_a} vs {ls_b}" - ) - assert torch.equal(u_a, u_b), ( - f"U_aniso moved when only the FREE reflections changed: {u_a} vs {u_b}" - ) - - -@pytest.mark.integration -def test_identical_datasets_fit_a_unit_scale(pair): - """Two copies of one dataset must scale onto each other with no correction. - - The sanity check the objectives have to pass before any comparison between them - means anything. - """ - dc = pair - dc.scale() - log_scale, U = _fitted(dc) - assert float(log_scale.abs().max()) < 1e-3, log_scale - assert float(U.abs().max()) < 1e-3, U - - -@pytest.mark.integration -def test_a_known_scale_is_recovered(pair): - """Scale one dataset by a known factor and check the fit undoes it.""" - dc = pair - k = 2.5 - with torch.no_grad(): - d = dc["other"] - d.F *= k - d.F_sigma *= k - d._corrected_fp = None - d._corrected_cache = None - dc.scale() - log_scale, _ = _fitted(dc) - # log_scale multiplies the observations, so recovering 1/k means log_scale = -log(k). - import math - assert float(log_scale.reshape(-1)[0]) == pytest.approx(-math.log(k), abs=0.02) - - -@pytest.mark.unit -def test_the_objective_is_not_selectable(): - """No objective parameter, and specifically no sigma-weighted one. - - Two separate reasons, both worth keeping written down. There is no model in this fit, - so a sigma_A or Rice likelihood has no model error to account for. And - inverse-variance weighting *collapses* on a scale fit -- down-weighting the weak - shells is exactly what lets the scale run away in them -- which is why the - model-to-data fit's default came back to unit-weight ``ls`` as well. A sigma-weighted - variant was built and measured here: it scored slightly better on held-out - reflections for one dataset pair, and was still removed, because a small gain on one - pair does not outweigh a failure mode found across a panel. - """ - import inspect - - from torchref.io.datasets.collection import DatasetCollection - - assert "objective" not in inspect.signature(DatasetCollection.scale).parameters - src = inspect.getsource(DatasetCollection.scale) - assert "sigma" in src, "the reason sigma weighting is absent must stay documented" - - -@pytest.mark.integration -def test_the_objective_is_normalised(pair): - """The loss handed to L-BFGS must be O(1), because its tolerances are absolute. - - Not a style point: ``tolerance_grad``/``tolerance_change`` are absolute, so an - objective carrying the data's own magnitude (~1e9 on a large work set under unit - weights) puts the float32 ulp of the loss above the decrease the line search is - trying to resolve. Probed by scaling the data by 1e3 and checking the fit still - recovers the same answer. - """ - import math - - dc = pair - with torch.no_grad(): - d = dc["other"] - d.F *= 1000.0 - d.F_sigma *= 1000.0 - d._corrected_fp = None - d._corrected_cache = None - dc.scale() - log_scale, _ = _fitted(dc) - assert float(log_scale.reshape(-1)[0]) == pytest.approx(-math.log(1000.0), abs=0.05) diff --git a/tests/unit/io/test_intensity_accessors.py b/tests/unit/io/test_intensity_accessors.py index c45ab1a2..1ebf7bca 100644 --- a/tests/unit/io/test_intensity_accessors.py +++ b/tests/unit/io/test_intensity_accessors.py @@ -1,16 +1,6 @@ -"""Scaled intensities must carry the amplitude scale **squared**. +"""Scaled amplitude and intensity accessors propagate measurement uncertainties.""" -``get_corrected_data`` applies ``corr(s, U) * exp(log_scale)`` to amplitudes. The -intensity counterpart has to apply the square of that, because the correction is defined -on amplitudes. Getting it wrong leaves a smooth, resolution-dependent error in the -intensities that is indistinguishable from a scale or overall-B mismatch -- so it is -pinned here as an exact ratio rather than checked by eye. - -The subset views are also pinned: ``.F``/``.sigF`` are corrected and ``.F_raw``/``.sigF_raw`` -are not, and ``.I``/``.sigI`` now follow the same rule. A view where ``.F`` was scaled and -``.I`` was not is the shape of bug that makes an amplitude target and an intensity target -disagree about which dataset they are fitting. -""" +import math import pytest import torch @@ -28,7 +18,11 @@ def with_intensities(mtz_dir): data = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) if data.I is None: pytest.skip("1DAW loaded without intensities") - return data + from torchref import ScaledDataset + from torchref.scaling import DatasetScaler + + scaler = DatasetScaler({"data": data, "peer": data}) + return ScaledDataset(data, scaler, "data") @pytest.fixture(scope="module") @@ -47,25 +41,25 @@ def without_intensities(mtz_dir): def _perturb(data, dlog=0.3): - """Give the dataset a non-trivial scale and anisotropy, restored on exit.""" + """Temporarily change scaler-owned overall and anisotropic corrections.""" + from contextlib import contextmanager - class _Ctx: - def __enter__(self): - self.log_scale = data.log_scale.detach().clone() - self.U = data.U_aniso.detach().clone() + @contextmanager + def changed(): + parameters = data.scaler.raw_parameters + original = parameters.detach().clone() + try: with torch.no_grad(): - data.log_scale += dlog - data.U_aniso += torch.tensor([0.01, -0.005, 0.008, 0.002, 0.0, 0.0]) - return data - - def __exit__(self, *exc): + parameters[0, 0] += 2 * dlog + parameters[0, 1:] += parameters.new_tensor( + [0.02, -0.01, 0.016, 0.004, 0, 0] + ) + yield data + finally: with torch.no_grad(): - data.log_scale.copy_(self.log_scale) - data.U_aniso.copy_(self.U) - data._corrected_fp = None - data._corrected_I_fp = None + parameters.copy_(original) - return _Ctx() + return changed() @pytest.mark.unit @@ -86,8 +80,8 @@ def test_intensity_factor_is_the_square_of_the_amplitude_factor( keep = data.masks().to(torch.bool) & (data.F.abs() > 1e-6) & ( data.I.abs() > 1e-6 ) - amp_factor = (F_scaled[keep] / data.F[keep]) ** 2 - int_factor = I_scaled[keep] / data.I[keep] + amp_factor = (F_scaled[keep] / data.F_raw[keep]) ** 2 + int_factor = I_scaled[keep] / data.I_raw[keep] rel = ((int_factor - amp_factor).abs() / amp_factor.abs()).max() assert rel < 1e-5, ( @@ -103,8 +97,8 @@ def test_sigma_scales_with_the_same_factor_as_the_intensity( I_scaled, sig_scaled = data.get_corrected_intensities() keep = (data.I.abs() > 1e-6) & (data.I_sigma.abs() > 1e-6) - ratio_I = I_scaled[keep] / data.I[keep] - ratio_s = sig_scaled[keep] / data.I_sigma[keep] + ratio_I = I_scaled[keep] / data.I_raw[keep] + ratio_s = sig_scaled[keep] / data.I_sigma_raw[keep] assert torch.allclose(ratio_I, ratio_s, rtol=1e-6) def test_the_perturbation_actually_changes_the_intensities( @@ -123,16 +117,14 @@ def test_a_pure_scale_change_squares_into_the_intensities( """A doubling of the amplitude scale must quadruple the intensities.""" data = with_intensities base, _ = data.get_corrected_intensities() - original = data.log_scale.detach().clone() + original = data.scaler.raw_parameters.detach().clone() try: with torch.no_grad(): - data.log_scale += float(torch.log(torch.tensor(2.0))) - data._corrected_I_fp = None + data.scaler.raw_parameters[0, 0] += 2 * math.log(2) doubled, _ = data.get_corrected_intensities() finally: with torch.no_grad(): - data.log_scale.copy_(original) - data._corrected_I_fp = None + data.scaler.raw_parameters.copy_(original) keep = base.abs() > 1e-6 ratio = (doubled[keep] / base[keep]) @@ -155,8 +147,8 @@ def test_raw_views_match_the_parent_tensors(self, with_intensities): data = with_intensities work = data.work idx = work.indices - assert torch.equal(work.I_raw, data.I.index_select(0, idx)) - assert torch.equal(work.sigI_raw, data.I_sigma.index_select(0, idx)) + assert torch.equal(work.I_raw, data.I_raw.index_select(0, idx)) + assert torch.equal(work.sigI_raw, data.I_sigma_raw.index_select(0, idx)) def test_subset_intensities_match_the_full_size_scaled_array( self, with_intensities @@ -171,18 +163,18 @@ def test_subset_intensities_match_the_full_size_scaled_array( assert torch.equal(sub.sigI, sig_scaled.index_select(0, idx)) def test_cache_follows_a_scale_change(self, with_intensities): - """The fingerprint must invalidate, or a refinement would fit stale data.""" + """Subset access must read the current shared scale parameters.""" data = with_intensities first = data.work.I.clone() - original = data.log_scale.detach().clone() + original = data.scaler.raw_parameters.detach().clone() try: with torch.no_grad(): - data.log_scale += 0.5 + data.scaler.raw_parameters[0, 0] += 1.0 second = data.work.I assert not torch.allclose(first, second) finally: with torch.no_grad(): - data.log_scale.copy_(original) + data.scaler.raw_parameters.copy_(original) @pytest.mark.unit diff --git a/tests/unit/io/test_reflection_data_reindex.py b/tests/unit/io/test_reflection_data_reindex.py index 07a58cc2..8b8e5058 100644 --- a/tests/unit/io/test_reflection_data_reindex.py +++ b/tests/unit/io/test_reflection_data_reindex.py @@ -49,7 +49,7 @@ def test_carries_all_per_reflection_fields(self): # ``light`` lacks the last 10%; the reference grid lacks the first 10%, # so the two sets genuinely differ (each has reflections the other lacks). light = _synthetic(grid[: int(n * 0.9)], seed=1) - ref_hkl = grid[int(n * 0.1):].clone() + ref_hkl = grid[int(n * 0.1) :].clone() assert len(light.hkl_anomalous) == len(light.hkl) # sane before light.validate_hkl(ref_hkl) @@ -146,34 +146,3 @@ def test_forward_finite_with_different_reflection_sets(self, pdb_dir, mtz_dir): target = CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0) loss = target.forward() assert torch.isfinite(loss) - - -@pytest.mark.unit -def test_reindex_preserves_non_per_reflection_u_aniso(): - """``U_aniso`` must be exempt by *name*, not by a shape coincidence. - - The reindexer decides "is this per-reflection?" by ``shape[0] == n_hkl``. - That heuristic collides when the dataset happens to have exactly as many - reflections as the field is long -- ``U_aniso`` is ``(6,)``, so a - 6-reflection dataset would have it gathered and reordered as if it were - per-reflection data. - """ - # Exactly 6 reflections: the same length as U_aniso. - hkl = torch.tensor( - [[1, 0, 1], [0, 1, 1], [0, 0, 1], [1, 1, 1], [1, 0, 2], [0, 1, 2]], - dtype=torch.int32, - ) - data = _synthetic(hkl) - u_aniso = torch.tensor([0.1, 0.2, 0.3, 0.01, 0.02, 0.03]) - data.U_aniso = u_aniso.clone().to(device=data.device) - assert len(data.hkl) == data.U_aniso.shape[0] == 6, "precondition: lengths collide" - - # Reindex onto a different HKL ordering/size; U_aniso must not follow. - ref_hkl = torch.tensor( - [[0, 1, 2], [1, 0, 2], [1, 1, 1], [0, 0, 1], [0, 1, 1], [1, 0, 1], [2, 0, 1]], - dtype=torch.int32, - ) - data.validate_hkl(ref_hkl.to(data.device)) - - assert data.U_aniso.shape == (6,), f"U_aniso reshaped to {tuple(data.U_aniso.shape)}" - assert torch.allclose(data.U_aniso.cpu(), u_aniso), "U_aniso values were permuted" diff --git a/tests/unit/io/test_select_reflection_data.py b/tests/unit/io/test_select_reflection_data.py index a3db5a8f..8c69449d 100644 --- a/tests/unit/io/test_select_reflection_data.py +++ b/tests/unit/io/test_select_reflection_data.py @@ -24,9 +24,6 @@ def _make_reflection_data(n=20): rd.phase = torch.rand(n, dtype=torch.float32) * 6.28 rd.fom = torch.rand(n, dtype=torch.float32) - # Non-per-reflection tensor (should NOT be indexed) - rd.U_aniso = torch.rand(6, dtype=torch.float32) - # Cell and spacegroup rd.cell = Cell( torch.tensor([50.0, 60.0, 70.0, 90.0, 90.0, 90.0]), @@ -87,21 +84,6 @@ def test_permutation(self): torch.testing.assert_close(sel.phase, rd.phase[perm]) -class TestNonMatchingTensors: - """Tensors whose first dim != n_refl should be cloned, not indexed.""" - - def test_u_aniso_copied(self): - rd = _make_reflection_data(20) - mask = torch.zeros(20, dtype=torch.bool) - mask[:10] = True - - sel = rd[mask] - - # U_aniso has shape (6,), not (20,), so it should be copied as-is - torch.testing.assert_close(sel.U_aniso, rd.U_aniso) - assert sel.U_aniso.shape == (6,) - - class TestCellAndSpacegroup: """Cell is cloned; spacegroup is copied by reference.""" diff --git a/tests/unit/refinement/test_collection_target_characterisation.py b/tests/unit/refinement/test_collection_target_characterisation.py index b82c673f..7f0f3922 100644 --- a/tests/unit/refinement/test_collection_target_characterisation.py +++ b/tests/unit/refinement/test_collection_target_characterisation.py @@ -36,10 +36,11 @@ def collection(pdb_dir, mtz_dir): """``(dc, mc, scaler)`` for a dark/light pair with a real difference in both the data and the models. - The light dataset carries its own ``log_scale``, so ``F_obs_light != F_obs_dark`` + The light dataset carries a shared scale view, so ``F_obs_light != F_obs_dark`` only through the *corrected* accessor -- which is what makes the raw-vs-scaled invariant below bite. The light model is displaced, so ``ΔF_calc != 0`` too. """ + pdb = pdb_dir / "1DAW.pdb" mtz = mtz_dir / "1DAW.mtz" if not (pdb.exists() and mtz.exists()): @@ -67,6 +68,8 @@ def collection(pdb_dir, mtz_dir): dc.add_dataset("dark", data_dark, set_as_reference=True) dc.add_dataset("light", data_light) + dc.scale(nsteps=1) + mc = ModelCollection([model_dark, model_light], dark_key="dark", verbose=0) mc.add_dark() mc.add_timepoint("light", [0.7, 0.3]) @@ -96,7 +99,7 @@ def _targets(dc, mc, scaler): class TestObservedAmplitudesAreScaled: """The loss must move when a dataset's own scale moves. - ``DatasetCollection.scale()`` fits a per-dataset ``log_scale``/``U_aniso`` that + ``DatasetCollection.scale()`` fits a per-dataset shared corrections that exists only in ``get_corrected_data()``. A target reading raw ``.F`` is completely blind to it, so this is a direct test of which accessor is in use. """ @@ -108,15 +111,15 @@ def test_loss_responds_to_the_datasets_own_log_scale(self, collection, name): before = target.forward().item() light = dc["light"] - original = light.log_scale.detach().clone() + original = dc.scaler.raw_parameters[1, 0].detach().clone() try: with torch.no_grad(): - light.log_scale += 0.25 # ~28% on amplitudes + dc.scaler.raw_parameters[1, 0] += 0.25 # ~28% on amplitudes target.maintenance() if hasattr(target, "maintenance") else None after = target.forward().item() finally: with torch.no_grad(): - light.log_scale.copy_(original) + dc.scaler.raw_parameters[1, 0].copy_(original) rel = abs(after - before) / abs(before) assert rel > 1e-3, ( @@ -130,11 +133,11 @@ def test_corrected_and_raw_amplitudes_actually_differ(self, collection): dc, _, _ = collection light = dc["light"] with torch.no_grad(): - light.log_scale += 0.25 + dc.scaler.raw_parameters[1, 0] += 0.25 corrected, _ = light.get_corrected_data() - raw = light.F + raw = light.F_raw differ = not torch.allclose(corrected, raw) - light.log_scale -= 0.25 + dc.scaler.raw_parameters[1, 0] -= 0.25 assert differ diff --git a/tests/unit/scaling/test_dataset_scaler.py b/tests/unit/scaling/test_dataset_scaler.py new file mode 100644 index 00000000..0372a832 --- /dev/null +++ b/tests/unit/scaling/test_dataset_scaler.py @@ -0,0 +1,272 @@ +"""Joint observed-data scaling and live dataset access on deposited reflections.""" + +import copy +import math + +import pytest +import torch + +from torchref import DatasetCollection, ReflectionData, ScaledDataset +from torchref.base.targets.dataset_scaling import dataset_scaling_loss +from torchref.scaling import DatasetScaler + + +@pytest.fixture(scope="module") +def deposited(mtz_dir): + """Measured 1DAW amplitudes, intensities, uncertainties and partitions.""" + return ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz_dir / "1DAW.mtz")) + + +def clone(data): + return data.__select__(torch.arange(len(data), device=data.device)) + + +def collection(data, factors): + dc = DatasetCollection(device=data.device, verbose=0) + for i, factor in enumerate(factors): + raw = clone(data) + raw.F *= factor + raw.F_sigma *= factor + if raw.I is not None: + raw.I *= factor**2 + raw.I_sigma *= factor**2 + dc.add_dataset(str(i), raw) + return dc + + +@pytest.mark.parametrize("factors", [(1, 1), (1, 1000), (1, 2, 8)]) +def test_known_scales_are_centered_and_sources_stay_raw(deposited, factors): + before = deposited.F.clone() + dc = collection(deposited, factors).scale(nsteps=2) + expected = torch.tensor(factors, dtype=dc.scaler.raw_parameters.dtype).log() + expected = expected.mean() - expected + torch.testing.assert_close( + dc.scaler.corrections[:, 0], expected, atol=1e-4, rtol=1e-4 + ) + assert dc.scaler.corrections[:, 1:].abs().max() < 1e-4 + assert torch.equal(deposited.F, before) + assert not hasattr(deposited, "parameters") + assert not hasattr(deposited, "log_scale") + assert not hasattr(deposited, "U_aniso") + for key, _ in dc: + assert isinstance(dc[key], ScaledDataset) + torch.testing.assert_close(dc[key].F, dc["0"].F, rtol=1e-4, atol=1e-4) + + +def test_known_anisotropy_is_recovered(deposited): + dc = collection(deposited, (1, 1, 1)) + reference = DatasetScaler(dc.datasets) + truth = torch.tensor( + [ + [0.3, 0.2, -0.1, 0.1, 0.03, -0.02, 0.05], + [-0.2, -0.1, 0.2, -0.05, -0.02, 0.01, -0.03], + [-0.1, -0.1, -0.1, -0.05, -0.01, 0.01, -0.02], + ], + dtype=reference.raw_parameters.dtype, + ) + for i, ds in enumerate(dc.values()): + distortion = (reference.design(ds.hkl) @ truth[i]).exp() + ds.F /= distortion + ds.F_sigma /= distortion + dc.scale(nsteps=4) + torch.testing.assert_close(dc.scaler.corrections, truth, atol=2e-3, rtol=2e-3) + + +def test_live_access_scales_both_sigmas_and_all_entrypoints(deposited): + dc = collection(deposited, (1, 1)).scale(nsteps=1) + data = dc["0"] + with torch.no_grad(): + dc.scaler.raw_parameters[0, 0] = 2 * math.log(2) + for name, power in [("F", 1), ("F_sigma", 1), ("I", 2), ("I_sigma", 2)]: + torch.testing.assert_close( + getattr(data, name), getattr(data, name + "_raw") * 2**power + ) + torch.testing.assert_close(data.work.F, data.F[data.work.mask]) + torch.testing.assert_close(data.work.sigF, data.F_sigma[data.work.mask]) + torch.testing.assert_close(data.work.sigI, data.I_sigma[data.work.mask]) + torch.testing.assert_close(data.work.sigI_raw, data.I_sigma_raw[data.work.mask]) + torch.testing.assert_close(dc.stack_F_obs()[0], data.F) + torch.testing.assert_close(dc.stack_I_sigma()[0], data.I_sigma) + dc.scaler.requires_grad_(True) + for _ in range(2): + dc.scaler.zero_grad() + data.work.F.sum().backward() + assert dc.scaler.raw_parameters.grad[1].abs().sum() > 0 + dc.scaler.requires_grad_(False) + + +def test_selection_copy_alignment_and_independent_collections(deposited): + a = collection(deposited, (1, 4)).scale(nsteps=1) + b = collection(deposited, (1, 9)).scale(nsteps=1) + view = a["0"] + selected = view.__select__(torch.arange(0, len(view), 3)) + torch.testing.assert_close(selected.F, view.F[::3]) + assert selected.scaler is a.scaler + assert copy.deepcopy(view).scaler is a.scaler + aligned = view.copy().validate_hkl(view.hkl.flip(0)) + torch.testing.assert_close(aligned.F, view.F.flip(0)) + assert a.scaler is not b.scaler + assert not torch.allclose(a["0"].F, b["0"].F) + scaler = a.scaler + a.scale(nsteps=1) + assert a.scaler is scaler + a.add_dataset("extra", deposited) + assert a.scaler is None + a.scale(nsteps=1) + assert a.scaler is not scaler + torch.testing.assert_close(view.F, view.F_raw * 2) + + +def test_two_dataset_loss_matches_propagated_variance_and_gradients(deposited): + f = torch.stack((deposited.F[:128], deposited.F[:128] * 1.1)).double() + sigma = torch.stack((deposited.F_sigma[:128], deposited.F_sigma[:128] * 3)).double() + log_k = torch.zeros_like(f, requires_grad=True) + mask = torch.isfinite(f) & torch.isfinite(sigma) & (sigma > 0) + actual = dataset_scaling_loss(f, sigma, log_k, mask) + common = mask.all(dim=0) + expected = 0.5 * ((f[0] - f[1]).square() / sigma.square().sum(dim=0))[common].sum() + torch.testing.assert_close(actual, expected) + assert torch.autograd.gradcheck( + lambda p: dataset_scaling_loss(f, sigma, p - p.mean(dim=0), mask), + (log_k,), + fast_mode=True, + ) + noisier = dataset_scaling_loss(f, sigma * 2, log_k, mask) + torch.testing.assert_close(noisier, actual / 4) + + +def test_missing_and_invalid_observations_have_finite_gradients(deposited): + f = torch.stack((deposited.F[:128], deposited.F[:128] * 1.1)) + sigma = torch.stack((deposited.F_sigma[:128], deposited.F_sigma[:128])) + mask = torch.isfinite(f) & torch.isfinite(sigma) & (sigma > 0) + mask[:, :10] = False + f[:, :10] = float("nan") + sigma[:, :10] = 0 + log_k = torch.zeros_like(f, requires_grad=True) + loss = dataset_scaling_loss(f, sigma, log_k, mask) + loss.backward() + assert torch.isfinite(loss) + assert torch.isfinite(log_k.grad).all() + assert torch.equal(log_k.grad[:, :10], torch.zeros_like(log_k.grad[:, :10])) + + +def test_free_and_validation_changes_do_not_affect_fit(deposited): + results = [] + for filler in (3, 900): + dc = collection(deposited, (1, 2)) + ds = dc["1"] + held_out = ~ds.work.mask + ds.F[held_out] = filler + ds.F_sigma[held_out] = filler + dc.scale(nsteps=2) + results.append(dc.scaler.raw_parameters.detach().clone()) + assert torch.equal(*results) + + +def test_permutation_and_partial_overlap_chain(deposited): + n = len(deposited) + sources = { + "a": clone(deposited).__select__(torch.arange(n // 2)), + "b": clone(deposited), + "c": clone(deposited).__select__(torch.arange(n // 2, n)), + } + sources["b"].F *= 2 + sources["b"].F_sigma *= 2 + sources["c"].F *= 4 + sources["c"].F_sigma *= 4 + results = [] + for order in [("a", "b", "c"), ("c", "a", "b")]: + dc = DatasetCollection(device="cpu", verbose=0) + for key in order: + dc.add_dataset(key, sources[key]) + dc.scale(nsteps=2) + results.append( + {k: dc.scaler.corrections[dc.scaler.keys.index(k)] for k in order} + ) + for key in sources: + torch.testing.assert_close( + results[0][key], results[1][key], atol=1e-4, rtol=1e-4 + ) + with pytest.raises(ValueError, match="disconnected"): + DatasetScaler({k: sources[k] for k in ("a", "c")}) + with pytest.raises(ValueError, match="identify"): + DatasetScaler( + { + "a": deposited.__select__(torch.arange(6)), + "b": deposited.__select__(torch.arange(6)), + } + ) + + +def test_checkpoint_and_mtz_export_preserve_observations(deposited, tmp_path): + import reciprocalspaceship as rs + + dc = collection(deposited, (1, 4)).scale(nsteps=1) + path = tmp_path / "collection.pt" + dc.save_state(path) + restored = DatasetCollection.load_state(path, device="cpu") + assert restored["0"].scaler is restored["1"].scaler is restored.scaler + torch.testing.assert_close(restored.stack_F_obs(), dc.stack_F_obs()) + path = tmp_path / "view.pt" + dc["0"].save_state(path) + restored_view = ScaledDataset.load_state(path, device="cpu") + torch.testing.assert_close(restored_view.I, dc["0"].I) + path = tmp_path / "scaled.mtz" + dc["0"].write_mtz(str(path)) + exported = rs.read_mtz(str(path)) + assert len(exported) == len(dc["0"]) + import numpy as np + + for column, attribute in [ + ("F-obs", "F"), + ("SIGF-obs", "F_sigma"), + ("I-obs", "I"), + ("SIGI-obs", "I_sigma"), + ]: + np.testing.assert_allclose( + exported[column].to_numpy(dtype=float), + getattr(dc["0"], attribute).detach().numpy(), + rtol=1e-6, + ) + + +def test_bijvoet_observations_keep_distinct_identities(deposited): + """Canonical duplicate HKLs retain separate signed observations and sigmas.""" + data = deposited.__select__(torch.arange(256).repeat_interleave(2)) + data.friedel_merged = False + data.friedel_flags = torch.arange(len(data)) % 2 == 1 + data.hkl_anomalous = torch.where(data.friedel_flags[:, None], -data.hkl, data.hkl) + data.F[data.friedel_flags] *= 1.25 + dc = collection(data, (1, 4)).scale(nsteps=1) + assert len(dc) == len(data) + lookup = {tuple(h.tolist()): f for h, f in zip(data.hkl_anomalous, data.F)} + expected = torch.stack([lookup[tuple(h.tolist())] for h in dc["0"].hkl_anomalous]) + torch.testing.assert_close(dc["0"].F_raw, expected) + torch.testing.assert_close(dc["0"].F, dc["1"].F, rtol=1e-4, atol=1e-4) + assert dc["0"].friedel_flags.sum() == 256 + selected = ( + dc["0"] + .copy() + .validate_hkl(dc.hkl.flip(0), identity_hkl=dc["0"].hkl_anomalous.flip(0)) + ) + torch.testing.assert_close(selected.F, dc["0"].F.flip(0)) + + +def test_raw_dataset_excludes_deprecated_interfaces(deposited): + """Observation storage exposes subset views without optimization methods.""" + for name in ( + "log_scale", + "U_aniso", + "parameters", + "setup_scale", + "setup_anisotropy", + "compute_e_values", + "get_radial_shells", + "get_work_set", + "get_test_set", + "get_rfree_masks", + "_masked_unpack", + "get_mask", + ): + assert not hasattr(deposited, name) + assert not callable(deposited) diff --git a/torchref/__init__.py b/torchref/__init__.py index ce2513e7..9f63eb6e 100644 --- a/torchref/__init__.py +++ b/torchref/__init__.py @@ -92,13 +92,17 @@ # Data I/O from torchref.io import ( DatasetCollection, - ReflectionData, FcalcDataset, - read_mtz, + ReflectionData, + ScaledDataset, read_cif, + read_mtz, read_pdb, ) +# Maps +from torchref.maps import DifferenceMap, Map + # Model from torchref.model import Model, ModelFT from torchref.model.rigid_xyz import RigidXYZTensor @@ -106,20 +110,18 @@ # Refinement from torchref.refinement import LBFGSRefinement, Refinement from torchref.refinement.rigid_body_refinement import RigidBodyRefinementStep + +# Scaling +from torchref.scaling import Scaler, ScalerBase, SolventModel from torchref.symmetry import Cell, SpaceGroup, Symmetry +# Device movement mixin (public API for extension code) +from torchref.utils.device_mixin import DeviceMixin + # Restraints # torchref.topology.restraints.Restraints is not imported here: constructing it can # trigger a monomer-library download, so it stays lazy. -# Scaling -from torchref.scaling import Scaler, SolventModel, ScalerBase - -# Maps -from torchref.maps import DifferenceMap, Map - -# Device movement mixin (public API for extension code) -from torchref.utils.device_mixin import DeviceMixin __all__ = [ # Version and paths @@ -133,6 +135,7 @@ "sigma_cutoff_ed", # Data I/O "ReflectionData", + "ScaledDataset", "DatasetCollection", "read_mtz", "read_cif", diff --git a/torchref/base/targets/dataset_scaling.py b/torchref/base/targets/dataset_scaling.py new file mode 100644 index 00000000..2a877dd7 --- /dev/null +++ b/torchref/base/targets/dataset_scaling.py @@ -0,0 +1,51 @@ +"""Profile a shared amplitude against independently measured datasets. + +The inputs use a common reflection axis with explicit presence masks. Uncertainties +remain in their measurement units; scaling and consensus estimation are differentiable. +""" + +import torch + + +def dataset_scaling_loss( + amplitudes: torch.Tensor, + sigmas: torch.Tensor, + log_corrections: torch.Tensor, + mask: torch.Tensor, +) -> torch.Tensor: + """Return the summed profiled Gaussian least-squares loss. + + Parameters + ---------- + amplitudes, sigmas : torch.Tensor + Measured amplitudes and positive uncertainties, shape (N, H), in each + dataset's input amplitude units. Masked entries may be non-finite. + log_corrections : torch.Tensor + Dimensionless log amplitude corrections, shape (N, H). + mask : torch.Tensor + Boolean usable-observation mask, shape (N, H). Only columns with at + least two observations contribute. + + Returns + ------- + torch.Tensor + Scalar dimensionless loss. Gradients include the consensus and the + scale dependence of the propagated uncertainties. + """ + active = mask & (mask.sum(dim=0, keepdim=True) >= 2) + obs = torch.where(active, amplitudes, torch.zeros_like(amplitudes)) + sigma = torch.where(active, sigmas, torch.ones_like(sigmas)) + log_k = torch.where(active, log_corrections, torch.zeros_like(log_corrections)) + # In measurement units the model is mu / k and sigma is fixed. Rescaling + # its whitened design column avoids overflow without changing the projection. + log_design = -log_k - sigma.log() + floor = torch.finfo(log_design.dtype).min + shift = torch.where(active, log_design, floor).amax(dim=0, keepdim=True) + shift = torch.where(active.any(dim=0, keepdim=True), shift, torch.zeros_like(shift)) + shifted = torch.where(active, log_design - shift, torch.zeros_like(obs)) + design = torch.where(active, shifted.exp(), torch.zeros_like(obs)) + whitened = obs / sigma + norm = design.square().sum(dim=0).clamp_min(torch.finfo(design.dtype).tiny) + consensus = (design * whitened).sum(dim=0) / norm + residual = torch.where(active, whitened - design * consensus, torch.zeros_like(obs)) + return 0.5 * residual.square().sum() diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index 5db698b0..26424ea0 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -31,10 +31,10 @@ import torch from torchref.cli._common import ( - add_dual_model_args, + add_all_columns_arg, add_dmin_arg, + add_dual_model_args, add_general_args, - add_all_columns_arg, add_metadata_args, add_outdir_arg, add_output_format_args, @@ -43,9 +43,9 @@ configure_unbuffered_output, load_model, load_reflection_data, + parse_device_str, parse_weights, register_timing, - parse_device_str, validate_cif_files, validate_files, ) @@ -237,11 +237,11 @@ def setup_loss_state(dataset_collection, model_collection, scaler, dataset. Default False. """ from torchref.refinement import LossState + from torchref.refinement.targets import TotalADPTarget, TotalGeometryTarget from torchref.refinement.targets.collection import ( CollectionDifferenceTarget, CollectionMLTarget, ) - from torchref.refinement.targets import TotalADPTarget, TotalGeometryTarget from torchref.refinement.targets.similarity import CoordinateSimilarityTarget state = LossState(device=device) @@ -1399,6 +1399,7 @@ def _build_metadata(model, data, r_work, r_free): if not has_altloc_dark and not has_altloc_light: import pandas as pd + from torchref import __version__ from torchref.io.metadata import RefinementMetadata @@ -1508,6 +1509,7 @@ def _mtz_to_cif(mtz_path, cif_path): "cif": args.cif, "dmin": args.dmin, }, + "dataset_scaling": dc.scaling_metrics, "parameters": { "weight_schedule": weight_schedule, "n_cycles": args.n_cycles, @@ -1516,14 +1518,18 @@ def _mtz_to_cif(mtz_path, cif_path): "weights": target_weights, }, "results": { - "r_factor_dark": dict(zip( - ["r_work", "r_free"], - compute_rfactors(dark, data_dark, scaler), - )), - "r_factor_light": dict(zip( - ["r_work", "r_free"], - compute_rfactors(mixed, data_light, scaler), - )), + "r_factor_dark": dict( + zip( + ["r_work", "r_free"], + compute_rfactors(dark, data_dark, scaler), + ) + ), + "r_factor_light": dict( + zip( + ["r_work", "r_free"], + compute_rfactors(mixed, data_light, scaler), + ) + ), "fractions": mixed.fractions.detach().cpu().tolist(), "alpha_mean": float(mc.alpha_mean), "lambda_twin": float(mc.lambda_twin), diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index 76004bbc..ef4cf7a7 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -30,17 +30,16 @@ import torch from torchref.cli._common import ( - add_dual_model_args, add_dmin_arg, + add_dual_model_args, add_general_args, add_outdir_arg, - build_dual_column_names, configure_unbuffered_output, load_model, load_reflection_data, - register_timing, parse_device_str, + register_timing, validate_cif_files, validate_files, ) @@ -63,6 +62,8 @@ def build_atom_mask(selection_xyz, real_space_grid, cell, mask_radius, device): """ from torchref.base.coordinates.transforms_torch import ( get_fractional_matrix, + ) + from torchref.base.coordinates.transforms_torch import ( get_inv_fractional_matrix_torch as get_inverse_fractional_matrix, ) from torchref.base.electron_density.solvent_mask import add_to_solvent_mask @@ -227,9 +228,8 @@ def setup_ded_context( import gemmi from torchref import DatasetCollection - from torchref.symmetry.reciprocal_symmetry import expand_hkl - from torchref.config import get_float_dtype, normalize_device + from torchref.symmetry.reciprocal_symmetry import expand_hkl device = normalize_device(device) @@ -248,15 +248,10 @@ def setup_ded_context( collection.add_dataset("dark", data_dark) collection.add_dataset("light", data_light) collection.scale() + data_dark, data_light = collection["dark"], collection["light"] if verbose >= 1: - print(f"Scale parameters after optimization:") - for name, ds in collection: - if hasattr(ds, "log_scale") and ds.log_scale is not None: - print( - f" {name}: log_scale={ds.log_scale.item():.6f} " - f"(scale={torch.exp(ds.log_scale).item():.6f})" - ) + print("Inter-dataset scaling:", collection.scaling_metrics) # Extract matched reflections hkl_all = data_dark.hkl @@ -387,11 +382,11 @@ def compute_ded_maps( resolution_bins, reciprocal_cc_overall, reciprocal_cc_work, reciprocal_cc_free, w_delta_fcalc_asu. """ - from torchref.model.model_collection import ModelCollection + from torchref.base.fourier.grid import get_real_grid from torchref.cli.collection_difference_refine import ( setup_scaler as setup_collection_scaler, ) - from torchref.base.fourier.grid import get_real_grid + from torchref.model.model_collection import ModelCollection device = ctx["device"] @@ -551,11 +546,13 @@ def compute_ded_maps( def run_validation(args): """Run the DED validation pipeline.""" - from torchref.cli.collection_difference_refine import compute_rfactors - from torchref.model.model_collection import ModelCollection + from torchref.cli.collection_difference_refine import ( + compute_rfactors, + ) from torchref.cli.collection_difference_refine import ( setup_scaler as setup_collection_scaler, ) + from torchref.model.model_collection import ModelCollection device = parse_device_str(args.device) outdir = Path(args.outdir) diff --git a/torchref/io/__init__.py b/torchref/io/__init__.py index a20af5b7..79c69f9c 100644 --- a/torchref/io/__init__.py +++ b/torchref/io/__init__.py @@ -21,31 +21,33 @@ RestraintCIFReader, ) -# Metadata -from .metadata import RefinementMetadata - -# Top-level object-creation readers -from .readers import read_cif, read_mtz, read_pdb - # Dataset classes (primary API) from .datasets import ( CrystalDataset, DatasetCollection, - ReflectionData, FcalcDataset, + ReflectionData, + ScaledDataset, ) +# IHM ensemble support (mapping always available; reader/writer need python-ihm) +from .ihm_mapping import IHMEnsembleMapping, IHMModelGroupInfo, IHMStateInfo + +# Metadata +from .metadata import RefinementMetadata + # Reader classes (from format modules) from .mtz import MTZReader from .pdb import PDBReader -# IHM ensemble support (mapping always available; reader/writer need python-ihm) -from .ihm_mapping import IHMEnsembleMapping, IHMModelGroupInfo, IHMStateInfo +# Top-level object-creation readers +from .readers import read_cif, read_mtz, read_pdb __all__ = [ # Primary API - Datasets "CrystalDataset", "ReflectionData", + "ScaledDataset", "DatasetCollection", "FcalcDataset", # Top-level readers diff --git a/torchref/io/datasets/__init__.py b/torchref/io/datasets/__init__.py index 86fb8f0f..cf5c9d3e 100644 --- a/torchref/io/datasets/__init__.py +++ b/torchref/io/datasets/__init__.py @@ -15,6 +15,9 @@ __all__ = [ "CrystalDataset", "ReflectionData", + "ScaledDataset", "FcalcDataset", "DatasetCollection", ] + +from .scaled_dataset import ScaledDataset diff --git a/torchref/io/datasets/base.py b/torchref/io/datasets/base.py index 705dbde6..20c1edda 100644 --- a/torchref/io/datasets/base.py +++ b/torchref/io/datasets/base.py @@ -15,7 +15,7 @@ import torch from torchref.config import get_default_device, get_float_dtype, normalize_device -from torchref.symmetry import Cell +from torchref.symmetry import Cell, SpaceGroup from torchref.utils.device_mixin import DeviceMovementMixin if TYPE_CHECKING: @@ -67,13 +67,6 @@ class CrystalDataset(DeviceMovementMixin): # explicit Bijvoet pairs (separate signed-HKL rows). Gates the model's f'' term. friedel_merged: bool = True - # === E-value and anisotropy correction fields === - E: Optional[torch.Tensor] = None # E-values (N,) - E_squared: Optional[torch.Tensor] = None # E² values (N,) - F_squared_corrected: Optional[torch.Tensor] = None # Anisotropy-corrected F² (N,) - U_aniso: Optional[torch.Tensor] = None # Fitted anisotropy parameters (6,) - radial_shell_indices: Optional[torch.Tensor] = None # Shell assignments (N,) - # === Unit cell and symmetry === cell: Optional[Cell] = None # Cell object with [a, b, c, alpha, beta, gamma] spacegroup: Optional[str] = None # Space group name string @@ -131,15 +124,19 @@ def _get_state(self) -> Dict[str, Any]: state = {} for f in fields(self): + if f.name in {"source", "reader", "dataset", "_FrenchWilson"}: + # Loading provenance and conversion caches are not observation state. + state[f.name] = None + continue val = getattr(self, f.name) if isinstance(val, torch.Tensor): - state[f.name] = val.cpu() + state[f.name] = val.detach().cpu() elif f.name == "cell" and val is not None: state[f.name] = val.data.cpu() elif f.name == "device": state[f.name] = str(val) elif f.name == "spacegroup" and val is not None: - state[f.name] = val.xhm() # Extended Hermann-Mauguin + state[f.name] = val.xhm # Extended Hermann-Mauguin else: state[f.name] = val # Masks are not a dataclass field, so handle them separately. @@ -183,7 +180,8 @@ def _from_state(cls, state: Dict[str, Any], device=None) -> "CrystalDataset": if "device" in state: state["device"] = torch.device(state["device"]) - # Spacegroup stays a string here; subclasses that want an object rewrap. + if isinstance(state.get("spacegroup"), str): + state["spacegroup"] = SpaceGroup(state["spacegroup"], device=device) if "cell" in state and state["cell"] is not None: if isinstance(state["cell"], torch.Tensor): # Conform the reloaded cell to the config float dtype rather than diff --git a/torchref/io/datasets/collection.py b/torchref/io/datasets/collection.py index 9a2daaa0..b8f35ffa 100644 --- a/torchref/io/datasets/collection.py +++ b/torchref/io/datasets/collection.py @@ -7,13 +7,16 @@ """ from dataclasses import dataclass, field -from typing import Dict, Iterator, List, Optional, Tuple +from typing import TYPE_CHECKING, Dict, Iterator, List, Optional, Tuple import torch from .base import CrystalDataset from .reflection_data import ReflectionData +from .scaled_dataset import ScaledDataset +if TYPE_CHECKING: + from torchref.scaling import DatasetScaler @dataclass @@ -21,9 +24,9 @@ class DatasetCollection(CrystalDataset): """ Container for multiple related crystal datasets on a common HKL set. - Members are expanded in place onto the reference dataset's HKL grid - (:meth:`ReflectionData.validate_hkl`) and moved to the collection's device, - so adding a dataset MUTATES it. Dict-like access via ``[]``, ``keys()``, + Members are copied onto the union HKL grid without changing input datasets. + The reference supplies cell and space-group metadata, not a fixed scale. + ``scale()`` installs ScaledDataset members backed by one shared scaler. Dict-like access via ``[]``, ``keys()``, ``values()``, ``items()``, ``get()``, and iteration yields ``(name, dataset)`` in insertion order. @@ -56,55 +59,73 @@ class DatasetCollection(CrystalDataset): _cell: Optional[torch.Tensor] = field(default=None, repr=False) _spacegroup: Optional[str] = field(default=None, repr=False) _resolution: Optional[torch.Tensor] = field(default=None, repr=False) - _scale_factors: Dict[str, torch.Tensor] = field(default_factory=dict, repr=False) + scaler: Optional["DatasetScaler"] = field(default=None, repr=False) + scaling_metrics: dict = field(default_factory=dict, repr=False) def add_dataset( self, name: str, dataset: ReflectionData, set_as_reference: bool = False ) -> "DatasetCollection": - """ - Add a dataset, expanding it onto the reference HKL grid **in place**. + """Add a copied dataset and rebuild the union reflection grid. Parameters ---------- name : str - Identifier for this dataset. + Unique member name. dataset : ReflectionData - The dataset to add. - set_as_reference : bool, optional - If True, this dataset's HKL becomes the reference. The first dataset - added becomes the reference regardless. + Raw or scaled observations; scaled inputs contribute their raw values. + set_as_reference : bool + Use this dataset's cell and symmetry as collection metadata. Returns ------- DatasetCollection - Self, for method chaining. - - Raises - ------ - ValueError - If a dataset with the same name already exists. + Self. Membership changes discard the fitted joint scaler; call scale() + again to fit all members. Existing raw inputs are never mutated. """ if name in self._datasets: raise ValueError(f"Dataset '{name}' already exists in collection") - - if len(self._datasets) == 0 or set_as_reference: + members = { + k: d.raw_data() if isinstance(d, ScaledDataset) else d + for k, d in self._datasets.items() + } + raw = dataset.raw_data() if isinstance(dataset, ScaledDataset) else dataset + members[name] = raw.__select__(torch.arange(len(raw), device=raw.device)) + members[name].source = None + members[name].spacegroup = raw.spacegroup.copy() + if ( + len({d.spacegroup.xhm for d in members.values()}) != 1 + or len({d.friedel_merged for d in members.values()}) != 1 + ): + raise ValueError( + "Datasets require compatible symmetry and Friedel conventions" + ) + if not self._dataset_order or set_as_reference: self._reference_dataset = name - self._common_hkl = dataset.hkl.clone() - if dataset.cell is not None: - self._cell = dataset.cell.clone() - self._spacegroup = dataset.spacegroup - - if self._common_hkl is not None and dataset.hkl is not None: - dataset.validate_hkl(self._common_hkl) - - dataset.to(self.device) - - self._datasets[name] = dataset + self._cell = raw.cell.clone() if raw.cell is not None else None + self._spacegroup = members[name].spacegroup self._dataset_order.append(name) - - if self.verbose > 0: - print(f"Added dataset '{name}' ({len(dataset)} reflections)") - + union_hkl = torch.unique( + torch.cat( + [ + (d.hkl if d.friedel_merged else d._hkl_for_sf()).to(self.device) + for d in members.values() + ] + ), + dim=0, + ) + identity_hkl = None + if raw.friedel_merged: + self._common_hkl = union_hkl + else: + canonical, _, _, order = raw.spacegroup.canonicalize_hkl(union_hkl) + self._common_hkl = canonical + identity_hkl = union_hkl[order] + for data in members.values(): + data.to(self.device) + data.validate_hkl(self._common_hkl, identity_hkl=identity_hkl) + self._datasets = members + self.scaler = None + self.scaling_metrics = {} return self @property @@ -263,97 +284,73 @@ def __call__(self, mask: bool = True) -> Dict[str, Tuple]: """ return {name: ds(mask=mask, scale=True) for name, ds in self} - def scale(self): - """ - Unit-weight least-squares fit of every non-reference dataset's scale and - anisotropy onto the reference, whose own parameters are left untouched. - - **This is the data-to-data fit**, the only one in the library with no model on - either side. So there is no model error to account for and nothing for a sigma_A - or Rice likelihood to do -- which is why the objective is least squares and there - is no way to select another. - - **Sigma weighting is deliberately not offered**, and that is a measured decision - rather than an omission. ``sum (F - F_ref)**2 / (sigma**2 + sigma_ref**2)`` is - superficially the principled choice -- both sides are measurements, so the - denominator is the honest propagated error on the difference being minimised -- - and on a single dataset pair it does score slightly better on held-out - reflections. It is still wrong to use: inverse-variance weighting on a scale fit - collapses, because down-weighting the weak shells is exactly what lets the scale - run away in them, and the same objective was tried and rejected for the - model-to-data fit (whose default likewise came back to unit-weight ``ls``). A - small held-out gain on one pair does not outweigh a failure mode found across a - panel. Do not re-add it. - - Fitted on the **work set** of both datasets. L-BFGS with strong-Wolfe line - search, 10 outer steps of ``max_iter=100``, on an objective normalised to O(1) - because those tolerances are absolute. Members' ``log_scale``/``U_aniso`` are - mutated, and ``requires_grad`` is turned on and back off around the fit. + def scale(self, nsteps: int = 10, max_iter: int = 100) -> "DatasetCollection": + """Jointly scale observations and expose live ScaledDataset members. - Raises - ------ - ValueError - If no reference dataset is set, or there is nothing else to scale. + Parameters + ---------- + nsteps, max_iter : int + Outer steps and per-step iteration limit for the dedicated scaler. + + Returns + ------- + DatasetCollection + Self. Retrieve scaled observations from this collection; references + to original inputs remain raw. Repeated calls reuse the parameter owner. """ - if self._reference_dataset is None: - raise ValueError("No reference dataset set for scaling") - - ref_ds = self._datasets[self._reference_dataset] - to_scale = [ds for name, ds in self if name != self._reference_dataset] - - if not to_scale: - raise ValueError("No datasets to scale against reference") - - parameters = [p for data in to_scale for p in data.parameters()] - [p.requires_grad_(True) for p in parameters] - optimizer = torch.optim.LBFGS(parameters, max_iter=100, line_search_fn='strong_wolfe') - - # Masks once (they do not change during the fit). The WORK subset, not - # `masks()`: the latter is validity only -- `TensorMasks.__call__` ANDs the - # validity masks and carries no work/free notion at all -- so fitting against it - # puts the free reflections into the scale parameters, upstream of every target, - # and compromises any free-set number the pipeline later reports. Degrades to - # all-valid on a dataset with no R-free flags, which is the pre-existing - # behaviour for that case. - ref_mask = ref_ds.work.mask - combined = [ds.work.mask & ref_mask for ds in to_scale] - - # The normaliser: once, detached, outside the closure. L-BFGS converges on - # ABSOLUTE tolerances, so an objective carrying the data's own magnitude leaves - # `tolerance_grad`/`tolerance_change` meaningless -- the same hazard - # `ScalerBase.refine_lbfgs` documents at length. This fit had no normaliser. - with torch.no_grad(): - ref_F0, _ = ref_ds.get_corrected_data() - ssq = sum(float(ref_F0[cm].pow(2).sum()) for cm in combined) - norm = 1.0 / max(ssq, 1e-30) - - def closure(): - optimizer.zero_grad() - loss = 0.0 - # get_corrected_data, not __call__: MaskedTensor has no autograd. - ref_F_scaled, _ = ref_ds.get_corrected_data() - - for ds, cm in zip(to_scale, combined): - F_scaled, _ = ds.get_corrected_data() - loss = loss + torch.sum((F_scaled[cm] - ref_F_scaled[cm]) ** 2) - loss = loss * norm - loss.backward() - return loss - - for i in range(10): - optimizer.step(closure) - [p.requires_grad_(False) for p in parameters] - - - # ------------------------------------------------------------------ - # Batched observation accessors - # ------------------------------------------------------------------ - # - # Every member is expanded onto the common HKL grid by ``add_dataset``, so these - # stack cleanly on a leading dataset axis. All of them return the **scaled** - # observations -- the per-dataset ``log_scale``/``U_aniso`` that ``scale()`` fits - # exists only in the corrected accessors, and a target reading the raw tensors - # would silently ignore the inter-dataset scaling. + from torchref.scaling.dataset_scaler import DatasetScaler + + raw = { + k: d.raw_data() if isinstance(d, ScaledDataset) else d + for k, d in self._datasets.items() + } + if self.scaler is None: + scaler = DatasetScaler(raw, device=self.device) + metrics = scaler.fit(nsteps=nsteps, max_iter=max_iter) + self._datasets = {k: ScaledDataset(d, scaler, k) for k, d in raw.items()} + self.scaler = scaler + else: + self.scaler.datasets = raw + metrics = self.scaler.fit(nsteps=nsteps, max_iter=max_iter) + self.scaling_metrics = metrics + return self + + def _get_state(self) -> dict: + raw = { + k: (d.raw_data() if isinstance(d, ScaledDataset) else d)._get_state() + for k, d in self._datasets.items() + } + scaler_state = None if self.scaler is None else self.scaler.get_state() + if scaler_state is not None: + scaler_state.pop("datasets") + return { + "datasets": raw, + "reference": self._reference_dataset, + "scaler": scaler_state, + "scaling_metrics": self.scaling_metrics, + } + + @classmethod + def _from_state(cls, state: dict, device=None) -> "DatasetCollection": + from torchref.scaling.dataset_scaler import DatasetScaler + + result = cls(device=device) if device is not None else cls() + for key, raw in state["datasets"].items(): + result.add_dataset( + key, + ReflectionData._from_state(dict(raw), device), + set_as_reference=key == state["reference"], + ) + if state["scaler"] is not None: + result.scaler = DatasetScaler.from_state( + {**state["scaler"], "datasets": state["datasets"]}, device + ) + result._datasets = { + k: ScaledDataset(d, result.scaler, k) + for k, d in result._datasets.items() + } + result.scaling_metrics = state.get("scaling_metrics", {}) + return result def _keys_or_all(self, keys: Optional[List[str]]) -> List[str]: if keys is None: @@ -386,10 +383,7 @@ def stack_I_obs(self, keys: Optional[List[str]] = None) -> torch.Tensor: If any selected dataset carries no intensities. """ return torch.stack( - [ - self._require_intensities(k)[0] - for k in self._keys_or_all(keys) - ], + [self._require_intensities(k)[0] for k in self._keys_or_all(keys)], dim=0, ) @@ -402,10 +396,7 @@ def stack_I_sigma(self, keys: Optional[List[str]] = None) -> torch.Tensor: If any selected dataset carries no intensities. """ return torch.stack( - [ - self._require_intensities(k)[1] - for k in self._keys_or_all(keys) - ], + [self._require_intensities(k)[1] for k in self._keys_or_all(keys)], dim=0, ) @@ -441,10 +432,7 @@ def stack_masks( ) attr = {"work": "work", "free": "free", "val": "validation"}[use_set] return torch.stack( - [ - getattr(self._datasets[k], attr).mask - for k in self._keys_or_all(keys) - ], + [getattr(self._datasets[k], attr).mask for k in self._keys_or_all(keys)], dim=0, ) diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index b5415df4..f056f607 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -14,7 +14,6 @@ import numpy as np import pandas as pd import torch -from torch.nn import Parameter from torchref.base import math_torch from torchref.base.french_wilson import FrenchWilson @@ -27,16 +26,6 @@ if TYPE_CHECKING: from torchref.model.model_ft import ModelFT -# Suppress PyTorch MaskedTensor prototype warnings globally -# MaskedTensor is stable enough for our use case (aggregations, element-wise ops) -warnings.filterwarnings( - "ignore", message=".*MaskedTensors is in prototype stage.*", category=UserWarning -) - -if TYPE_CHECKING: - from torch.masked import MaskedTensor - - class _ReflectionSubset: """ Lightweight view of one reflection subset (``work`` / ``free`` / @@ -95,7 +84,7 @@ def __len__(self) -> int: def n(self) -> int: return int(self.indices.numel()) - # -- amplitudes (scaled, matching the legacy ``data(scale=True)``) ----- + # Subset reads dispatch through the parent observation attributes. @property def F(self) -> torch.Tensor: F_corr, _ = self._parent._corrected_or_raw() @@ -109,11 +98,11 @@ def sigF(self) -> torch.Tensor: # -- raw (uncorrected) amplitudes ------------------------------------- @property def F_raw(self) -> torch.Tensor: - return self._parent.F.index_select(0, self.indices) + return self._parent.F_raw.index_select(0, self.indices) @property def sigF_raw(self) -> torch.Tensor: - return self._parent.F_sigma.index_select(0, self.indices) + return self._parent.F_sigma_raw.index_select(0, self.indices) # -- common aliases --------------------------------------------------- @property @@ -144,13 +133,13 @@ def sigI(self): @property def I_raw(self): """Unscaled intensities, or None.""" - i = self._parent.I + i = self._parent.I_raw return i.index_select(0, self.indices) if i is not None else None @property def sigI_raw(self): """Unscaled intensity sigmas, or None.""" - si = self._parent.I_sigma + si = self._parent.I_sigma_raw return si.index_select(0, self.indices) if si is not None else None @property @@ -238,11 +227,7 @@ def __post_init__(self): """ # Call parent __post_init__ to initialize masks super().__post_init__() - self.setup_scale() - self.setup_anisotropy() - # Cached integer index maps for the work/free/validation subsets and - # a cache of the scaled (F, F_sigma). Both are invalidated by - # fingerprints (see _subset_indices / _corrected_or_raw). + # Subset membership is cached independently of observation values. self._subset_cache = { "work": None, "free": None, @@ -250,10 +235,6 @@ def __post_init__(self): "all": None, } self._subset_fp = None - self._corrected_cache = None - self._corrected_fp = None - self._corrected_I_cache = None - self._corrected_I_fp = None # ===================== work / free / validation ===================== @@ -361,25 +342,28 @@ def _subset_indices(self, kind: str) -> torch.Tensor: return self._subset_cache[kind] def _corrected_or_raw(self) -> Tuple[torch.Tensor, torch.Tensor]: - """Return the scaled (F, F_sigma) (matching ``data(scale=True)``), - cached against the (log_scale, U_aniso) fingerprint. Falls back to the - raw (F, F_sigma) if scaling is not set up. - """ + """Return the observations exposed by this dataset, without caching.""" + return self.get_corrected_data() - def _tv(t): - return (t.data_ptr(), t._version) if isinstance(t, torch.Tensor) else None + @property + def F_raw(self) -> Optional[torch.Tensor]: + """Measured amplitudes, shape (N,), in the input amplitude units.""" + return self.F - fp = ( - _tv(getattr(self, "log_scale", None)), - _tv(getattr(self, "U_aniso", None)), - ) - if self._corrected_fp != fp or self._corrected_cache is None: - try: - self._corrected_cache = self.get_corrected_data() - except Exception: - self._corrected_cache = (self.F, self.F_sigma) - self._corrected_fp = fp - return self._corrected_cache + @property + def F_sigma_raw(self) -> Optional[torch.Tensor]: + """Measured amplitude uncertainties, shape (N,), in amplitude units.""" + return self.F_sigma + + @property + def I_raw(self) -> Optional[torch.Tensor]: + """Measured intensities, shape (N,), in the input intensity units.""" + return self.I + + @property + def I_sigma_raw(self) -> Optional[torch.Tensor]: + """Measured intensity uncertainties, shape (N,), in intensity units.""" + return self.I_sigma # ===================== per-reflection field reindexing ===================== # @@ -399,10 +383,6 @@ def _tv(t): "I": 0.0, "phase": 0.0, "fom": 0.0, - "E": 0.0, - "E_squared": 0.0, - "F_squared_corrected": 0.0, - "radial_shell_indices": 0, "F_sigma": 1.0, "I_sigma": 1.0, "rfree_flags": 1, # missing reflections default to the work set @@ -416,9 +396,8 @@ def _tv(t): # lazily rebuilt by ``get_bins`` / the ``centric`` property). _REINDEX_DERIVED = ("resolution", "bin_indices", "_centric_flags") - # Tensor dataclass fields that are NOT per-reflection (exempt from the - # length invariant): overall anisotropy parameters have shape (6,). - _NON_PER_REFLECTION_TENSORS = frozenset({"U_aniso"}) + # Subclasses may declare non-reflection tensor fields here. + _NON_PER_REFLECTION_TENSORS = frozenset() def _reindex_per_reflection( self, @@ -467,13 +446,7 @@ def _reindex_per_reflection( name = f.name if name == "hkl" or name in derived: continue - # Declared non-per-reflection fields are exempt by *name*, not by - # shape. The shape test below is a heuristic and collides whenever - # n_src equals the field's own length -- ``U_aniso`` is (6,), so a - # 6-reflection dataset would have it gathered as if it were - # per-reflection. ``_assert_per_reflection_consistent`` and - # ``reduce_to_spacegroup`` already exempt by name; this keeps all - # three routines consistent. + # Non-reflection fields must be excluded by name, not tensor length. if name in self._NON_PER_REFLECTION_TENSORS: continue val = getattr(self, name) @@ -1885,15 +1858,6 @@ def filter_by_resolution( return self - def get_mask(self): - """ - Placeholder for returning a combined mask from all active filters. - - Not implemented; the body is empty and this returns ``None``. Use - :meth:`masks` (the combined-validity callable) to obtain the boolean - mask combining all active filter conditions. - """ - def cut_res( self, highres: Optional[float] = None, lowres: Optional[float] = None ) -> "ReflectionData": @@ -1917,104 +1881,6 @@ def cut_res( """ return self.filter_by_resolution(d_min=highres, d_max=lowres) - def get_rfree_masks(self) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor]]: - """ - Get boolean masks for work and test (free) sets. - - Returns - ------- - work_mask : torch.Tensor or None - Boolean tensor for work set (flag != 0). - test_mask : torch.Tensor or None - Boolean tensor for test/free set (flag == 0). - Both are None if no R-free flags are available. - - .. deprecated:: - Use ``data.work.mask`` / ``data.free.mask`` (which also apply the - validity masks), or ``data.work.indices`` / ``data.free.indices``. - """ - warnings.warn( - "ReflectionData.get_rfree_masks() is deprecated; use data.work.mask " - "/ data.free.mask (the work/free/validation accessor).", - DeprecationWarning, - stacklevel=2, - ) - if self.rfree_flags is None: - return None, None - - work_mask = self.rfree_flags != 0 - test_mask = self.rfree_flags == 0 - - return work_mask, test_mask - - def get_work_set(self) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: - """ - Get structure factors for the work set (R-free flag != 0). - - Returns - ------- - F_work : torch.Tensor - Structure factors for work set. - sigma_work : torch.Tensor or None - Uncertainties for work set, or None if not available. - - With no R-free flags this silently returns the *full* dataset (warning - printed), so a caller cannot tell work from all. - - .. deprecated:: - Use ``data.work.F`` / ``data.work.sigF``, which also apply the - validity masks and cache the subset indices. - """ - warnings.warn( - "ReflectionData.get_work_set() is deprecated; use data.work.F / " - "data.work.sigF (the work/free/validation accessor).", - DeprecationWarning, - stacklevel=2, - ) - if self.rfree_flags is None: - print("WARNING: No R-free flags available, returning full dataset") - return self.F, self.F_sigma - - work_mask = self.rfree_flags != 0 - F_work = self.F[work_mask] if self.F is not None else None - sigma_work = self.F_sigma[work_mask] if self.F_sigma is not None else None - - return F_work, sigma_work - - def get_test_set(self) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: - """ - Get structure factors for the test set (R-free flag == 0). - - Returns - ------- - F_test : torch.Tensor - Structure factors for test/free set. - sigma_test : torch.Tensor or None - Uncertainties for test set, or None if not available. - - Raises - ------ - ValueError - If no R-free flags are available. - - .. deprecated:: - Use ``data.free.F`` / ``data.free.sigF`` instead. - """ - warnings.warn( - "ReflectionData.get_test_set() is deprecated; use data.free.F / " - "data.free.sigF (the work/free/validation accessor).", - DeprecationWarning, - stacklevel=2, - ) - if self.rfree_flags is None: - raise ValueError("No R-free flags available in dataset") - - test_mask = self.rfree_flags == 0 - F_test = self.F[test_mask] if self.F is not None else None - sigma_test = self.F_sigma[test_mask] if self.F_sigma is not None else None - - return F_test, sigma_test - def get_max_res(self) -> Optional[float]: """Smallest d-spacing among valid reflections, in Ångströms.""" if self.resolution is None: @@ -2073,7 +1939,7 @@ def data_indexed( """ Return reflection data as compact (valid-only) tensors. - Note these are the RAW ``F``/``F_sigma``, not the scaled ones. + ScaledDataset returns its live corrected observations. Returns ------- @@ -2099,80 +1965,6 @@ def data_indexed( return hkl, F, F_sigma, rfree_flags - def __call__( - self, mask: bool = True, scale: bool = True - ) -> Tuple[torch.Tensor, "MaskedTensor", "MaskedTensor", torch.Tensor]: - """ - Return core reflection data with MaskedTensors for F and sigma. - - Everything is full size (N); invalid reflections are marked in the mask - rather than removed. With ``mask=True`` the returned ``F``/``F_sigma`` - are detached clones, so gradients do NOT flow through them (use - :meth:`get_corrected_data` when the graph is needed). - - Parameters - ---------- - mask : bool, optional - If True, wrap F and sigma as MaskedTensors. Default is True. - scale : bool, optional - If True, apply the current scale/anisotropy before returning. - - Returns - ------- - hkl : torch.Tensor - Miller indices of shape (N, 3), unfiltered. - F : MaskedTensor - Amplitudes of shape (N,) with invalid reflections masked. - F_sigma : MaskedTensor or None - Uncertainties of shape (N,) with invalid reflections masked. - rfree_flags : torch.Tensor or None - Flags of shape (N,), unfiltered. 1=work, 0=free. - - Raises - ------ - RuntimeError - If ``mask`` is True and every reflection is masked out. - - .. deprecated:: - Use the work/free/validation accessor (``data.work.F``, - ``data.free.F``, ``.sigF`` / ``.hkl`` / ``.select(...)``), or - ``data.get_corrected_data()`` for the full scaled (F, F_sigma). - """ - warnings.warn( - "Calling ReflectionData (data()) is deprecated; use the " - "data.work / data.free / data.validation accessor, or " - "data.get_corrected_data() for the full scaled arrays.", - DeprecationWarning, - stacklevel=2, - ) - return self._masked_unpack(mask=mask, scale=scale) - - def _masked_unpack( - self, mask: bool = True, scale: bool = True - ) -> Tuple[torch.Tensor, "MaskedTensor", "MaskedTensor", torch.Tensor]: - """Non-deprecated body of the legacy ``__call__``; see it for the contract. - - Internal only -- external callers should use the work/free/validation - accessor. - """ - from torch.masked import MaskedTensor - - hkl, F, F_sigma, rfree_flags = self.hkl, self.F, self.F_sigma, self.rfree_flags - - if scale: - F, F_sigma = self.get_corrected_data() - - if mask: - to_mask = self.masks() - if to_mask.sum() == 0: - raise RuntimeError( - "All reflections are masked! Check your filters/masks." - ) - F = MaskedTensor(F.detach().clone(), to_mask) - if F_sigma is not None: - F_sigma = MaskedTensor(F_sigma.detach().clone(), to_mask) - return hkl, F, F_sigma, rfree_flags - def data_fill_masked( self, mode="mean" ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: @@ -2199,23 +1991,31 @@ def data_fill_masked( R-free flags of shape (N,); filled-in reflections are assigned to the work set (True). """ - hkl, F, F_sigma, rfree = self._masked_unpack() + hkl, F, F_sigma = self.hkl, self.F, self.F_sigma + if F is None or F_sigma is None: + raise ValueError("Amplitude observations and uncertainties are required") + mask = self.masks() + if not bool(mask.any()): + raise ValueError("No valid reflections to fill from") + rfree = ( + self.rfree_flags.clone() + if self.rfree_flags is not None + else torch.ones_like(mask) + ) if mode == "mean": mean_F = self.mean_F_per_bin() mean_F_sigma = self.mean_sigma_per_bin() - F_data = F.get_data().clone() - F_sigma_data = F_sigma.get_data().clone() - mask = F.get_mask() + F_data = F.clone() + F_sigma_data = F_sigma.clone() F_data[~mask] = mean_F[self.bin_indices[~mask]] F_sigma_data[~mask] = mean_F_sigma[self.bin_indices[~mask]] rfree[~mask] = True # set missing to work set return hkl, F_data, F_sigma_data, rfree elif mode == "zero": - mask = F.get_mask() - F_data = F.get_data().clone() - F_sigma_data = F_sigma.get_data().clone() + F_data = F.clone() + F_sigma_data = F_sigma.clone() F_data[~mask] = 0.0 F_sigma_data[~mask] = 0.0 rfree[~mask] = True # set missing to work set @@ -2283,7 +2083,7 @@ def __select__(self, indices: torch.Tensor, op=None) -> "ReflectionData": elif val.shape and val.shape[0] == n_refl: setattr(selected, f.name, val[indices]) else: - # Non-matching tensor (e.g. U_aniso shape (6,)): copy as-is + # Preserve scalar tensor metadata. setattr(selected, f.name, val.clone()) elif isinstance(val, Cell): setattr(selected, f.name, val.clone()) @@ -2374,7 +2174,9 @@ def check_all_data_types(self): else: print(f"{key}: None") - def validate_hkl(self, hkl_ref: torch.Tensor) -> "ReflectionData": + def validate_hkl( + self, hkl_ref: torch.Tensor, *, identity_hkl: Optional[torch.Tensor] = None + ) -> "ReflectionData": """ Expand this dataset **in place** onto a reference HKL set. @@ -2389,6 +2191,10 @@ def validate_hkl(self, hkl_ref: torch.Tensor) -> "ReflectionData": hkl_ref : torch.Tensor Reference Miller indices of shape (N, 3), dtype int32; defines the canonical ordering for all aligned datasets. + identity_hkl : torch.Tensor, optional + Signed anomalous indices of shape (N, 3), distinguishing Bijvoet + observations that share a canonical HKL. When supplied, match these + against the dataset's signed indices and preserve their identities. Returns ------- @@ -2414,11 +2220,15 @@ def validate_hkl(self, hkl_ref: torch.Tensor) -> "ReflectionData": # Build lookup from data HKL to index # Use a dictionary with tuple keys for fast lookup - hkl_data_np = self.hkl.cpu().numpy() + source_hkl = self.hkl if identity_hkl is None else self._hkl_for_sf() + hkl_data_np = source_hkl.cpu().numpy() data_hkl_to_idx = {tuple(hkl): idx for idx, hkl in enumerate(hkl_data_np)} # For each reference HKL, find the corresponding data index (or -1 if missing) - hkl_ref_np = hkl_ref.cpu().numpy() + lookup_hkl = hkl_ref if identity_hkl is None else identity_hkl + if lookup_hkl.shape != hkl_ref.shape: + raise ValueError("identity_hkl must match the reference HKL shape") + hkl_ref_np = lookup_hkl.cpu().numpy() ref_to_data_idx = np.array( [data_hkl_to_idx.get(tuple(hkl), -1) for hkl in hkl_ref_np], dtype=np.int64 ) @@ -2429,6 +2239,9 @@ def validate_hkl(self, hkl_ref: torch.Tensor) -> "ReflectionData": # Reindex EVERY per-reflection field via the shared primitive. Masks are # handled separately below because they are not dataclass fields. presence_mask = self._reindex_per_reflection(valid_indices, hkl_ref) + if identity_hkl is not None: + self.hkl_anomalous = identity_hkl.to(self.hkl).clone() + self.friedel_flags = (self.hkl_anomalous != self.hkl).any(dim=-1) # Transfer existing masks to new indexing old_masks = dict(self.masks.items()) @@ -3654,437 +3467,29 @@ def get_scattering_vectors(self) -> torch.Tensor: return math_torch.get_scattering_vectors(self.hkl, self.cell.data) - def get_radial_shells( - self, - n_shells: int = 20, - d_min: Optional[float] = None, - d_max: Optional[float] = None, - ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """ - Create uniform radial shells in 1/d space for normalization. - - Uniform in 1/d, unlike :meth:`get_bins` which makes equal-count bins. - Caches the result on ``self.radial_shell_indices``. - - Parameters - ---------- - n_shells : int - Number of radial shells. Default is 20. - d_min, d_max : float, optional - Resolution limits in Angstroms; default to the dataset's own - smallest / largest d. - - Returns - ------- - shell_edges : torch.Tensor - Shell boundaries in Angstroms^-1, shape (n_shells+1,). - shell_centers : torch.Tensor - Shell centers in Angstroms^-1, shape (n_shells,). - shell_indices : torch.Tensor - Shell index for each reflection, shape (N,). Values -1 for out-of-range. - """ - from torchref.base.normalization import ( - assign_to_shells, - compute_radial_shells, - ) - - if self.resolution is None: - self._calculate_resolution() - - # Get resolution limits - if d_min is None: - d_min = self.get_max_res() - if d_max is None: - d_max = self.get_min_res() - - # Compute shells - shell_edges, shell_centers = compute_radial_shells( - d_min, d_max, n_shells, device=self.device - ) - - # Get s-vectors and magnitudes - s_vectors = self.get_scattering_vectors() - s_mag = torch.linalg.norm(s_vectors, dim=1) - - # Assign to shells - shell_indices = assign_to_shells(s_mag, shell_edges) - - # Cache shell indices - self.radial_shell_indices = shell_indices - - return shell_edges, shell_centers, shell_indices - - def fit_anisotropy( - self, - n_shells: int = 20, - d_min: Optional[float] = None, - d_max: Optional[float] = None, - n_iterations: int = 100, - verbose: Optional[bool] = None, - ) -> torch.Tensor: - """ - Fit anisotropy correction parameters to minimize CV within shells. - - Optimizes U so that corrected F² values have minimal coefficient of - variation within each resolution shell. - - Parameters - ---------- - n_shells : int - Number of resolution shells for variance calculation. - d_min, d_max : float, optional - Resolution limits in Angstroms; default to the dataset's own - smallest / largest d. - n_iterations : int - Optimizer steps; see :func:`fit_anisotropy_correction` -- below 20 - nothing is optimized. - verbose : bool, optional - Print progress. If None, uses self.verbose. - - Returns - ------- - U : torch.Tensor - Fitted anisotropy parameters [u11, u22, u33, u12, u13, u23], shape (6,). - Also stored in self.U_aniso. - - Raises - ------ - ValueError - If no amplitude data is available. - """ - from torchref.base import fit_anisotropy_correction - - if self.F is None: - raise ValueError("No amplitude data loaded") - - if verbose is None: - verbose = self.verbose > 0 - - # Get F² values - F_squared = self.F**2 - - # Get s-vectors - s_vectors = self.get_scattering_vectors() - - # Get resolution limits - if d_min is None: - d_min = self.get_max_res() - if d_max is None: - d_max = self.get_min_res() - - # Fit anisotropy - U, final_cv = fit_anisotropy_correction( - F_squared, - s_vectors, - n_shells=n_shells, - d_min=d_min, - d_max=d_max, - n_iterations=n_iterations, - verbose=verbose, - ) - - # Store result - self.U_aniso = U - - return U - - def setup_anisotropy( - self, - U_aniso: Optional[torch.Tensor] = None, - ) -> None: - """ - Setup anisotropy correction parameters. - - Parameters - ---------- - U_aniso : torch.Tensor, optional - Anisotropic parameters [u11, u22, u33, u12, u13, u23], shape (6,). - If None, U_aniso is initialized to a zero (6,) tensor. - - Returns - ------- - ReflectionData - Self, for method chaining. - """ - - if U_aniso is None: - U_aniso = torch.zeros( - 6, device=self.device, dtype=dtypes.float, requires_grad=False - ) - else: - U_aniso = U_aniso.to( - device=self.device, dtype=dtypes.float, requires_grad=False - ) - self.U_aniso = U_aniso - - return self - - def apply_anisotropy_correction( - self, - U_aniso: Optional[torch.Tensor] = None, - ) -> torch.Tensor: - """ - Apply anisotropy correction to F² values. - - Parameters - ---------- - U_aniso : torch.Tensor, optional - Anisotropic parameters [u11, u22, u33, u12, u13, u23], shape (6,). - If None, uses self.U_aniso (must have called fit_anisotropy first). - - Returns - ------- - F_corrected: torch.Tensor - Anisotropy-corrected F values, shape (N,). - sigma_F_corrected: torch.Tensor - Uncertainties of corrected F values, shape (N,). - Raises - ------ - ValueError - If no U parameters available and none provided. - """ - from torchref.base import apply_anisotropy_correction - - if U_aniso is None: - U_aniso = self.U_aniso - if U_aniso is None: - raise ValueError( - "No anisotropy parameters available. " - "Call fit_anisotropy() first or provide U_aniso." - ) - - if self.F is None: - raise ValueError("No amplitude data loaded") - - # Get s-vectors - s_vectors = self.get_scattering_vectors() - # Use raw tensors directly to preserve gradient flow - # (MaskedTensor doesn't support autograd operations) - F = self.F - sigma = self.F_sigma - # Apply correction - F_corrected = apply_anisotropy_correction(F, s_vectors, U_aniso) - sigma_F_corrected = ( - apply_anisotropy_correction(sigma, s_vectors, U_aniso) - if sigma is not None - else None - ) - - return F_corrected, sigma_F_corrected - - def compute_e_values( - self, - n_shells: int = 20, - d_min: Optional[float] = None, - d_max: Optional[float] = None, - apply_anisotropy: bool = True, - fit_anisotropy: bool = True, - verbose: Optional[bool] = None, - ) -> torch.Tensor: - """ - Compute E-values with optional anisotropy correction. - - E-values are normalized structure factors where = 1 within each - resolution shell. Anisotropy correction can be applied first to account - for directional variation in diffraction. - - Parameters - ---------- - n_shells : int - Number of resolution shells for normalization. - d_min, d_max : float, optional - Resolution limits in Angstroms; default to the dataset's own - smallest / largest d. - apply_anisotropy : bool - If True, correct for anisotropy before normalizing. - fit_anisotropy : bool - If True, refit U first; if False, reuse the existing - ``self.U_aniso``. Ignored unless ``apply_anisotropy``. - verbose : bool, optional - Print progress. If None, uses self.verbose. - - Returns - ------- - E : torch.Tensor - E-values, shape (N,). Also stores ``self.E``, ``self.E_squared`` and - ``self.radial_shell_indices`` as a side effect. - - Raises - ------ - ValueError - If no amplitude data is available. - """ - from torchref.base import F_squared_to_E_values - - if self.F is None: - raise ValueError("No amplitude data loaded") - - if verbose is None: - verbose = self.verbose > 0 - - # Get resolution limits - if d_min is None: - d_min = self.get_max_res() - if d_max is None: - d_max = self.get_min_res() - - # Get F² values (possibly with anisotropy correction) - if apply_anisotropy: - if fit_anisotropy: - self.fit_anisotropy( - n_shells=n_shells, d_min=d_min, d_max=d_max, verbose=verbose - ) - F_squared = self.apply_anisotropy_correction()[0] ** 2 - else: - F_squared = self.F**2 - - # Get s-vectors - s_vectors = self.get_scattering_vectors() - - # Compute E-values - E, E_squared, shell_idx = F_squared_to_E_values( - F_squared, s_vectors, n_shells=n_shells, d_min=d_min, d_max=d_max - ) - - # Store results - self.E = E - self.E_squared = E_squared - self.radial_shell_indices = shell_idx - - if verbose: - print(f"E-value statistics:") - print(f" E range: [{E.min():.3f}, {E.max():.3f}]") - print(f" E mean: {E.mean():.3f}, std: {E.std():.3f}") - print(f" E² mean: {E_squared.mean():.3f} (should be ~1.0)") - return E - - def setup_scale(self, scale: Optional[float] = None) -> float: - """ - Set overall scale factor, parametrized in log space. - - Parameters - ---------- - scale : float, optional - If provided, sets the scale factor directly (stored as its log). - If None (default), the scale defaults to 1.0 (``log_scale = 0.0``). - - Returns - ------- - ReflectionData - Self, for method chaining. - """ - if scale is None: - self.log_scale = torch.tensor( - 0.0, device=self.device, requires_grad=False, dtype=dtypes.float - ) - else: - self.log_scale = torch.log( - torch.tensor( - scale, device=self.device, requires_grad=False, dtype=dtypes.float - ) - ) - return self - def get_corrected_data(self) -> Tuple[torch.Tensor, torch.Tensor]: - """ - Get the anisotropy-corrected, scaled (F, F_sigma). - - Returns - ------- - Tuple[torch.Tensor, torch.Tensor] - Full-size F and F_sigma with ``exp(log_scale)`` and ``U_aniso`` - applied. + """Return amplitudes and sigmas, shape (N,), in this dataset's units. - Raises - ------ - ValueError - If ``setup_scale`` / ``setup_anisotropy`` have not run. + Raw datasets return their measurements; ScaledDataset exposes the live + scale correction through the same observation attributes. """ - - if not hasattr(self, "log_scale") or self.log_scale is None: - raise ValueError("Scale not set up. Call setup_scale() first.") - if not hasattr(self, "U_aniso") or self.U_aniso is None: - raise ValueError("Anisotropy not set up. Call setup_anisotropy() first.") - F_corrected, F_sigma_corrected = self.apply_anisotropy_correction() - scale_factor = torch.exp(self.log_scale) - F_scaled = F_corrected * scale_factor - F_sigma_scaled = F_sigma_corrected * scale_factor - - return F_scaled, F_sigma_scaled + return self.F, self.F_sigma def get_corrected_intensities(self) -> Tuple[torch.Tensor, torch.Tensor]: - """ - Get the anisotropy-corrected, scaled ``(I, I_sigma)``. - - The intensity counterpart of :meth:`get_corrected_data`. Both the anisotropy - factor and the overall scale enter **squared**, because they are defined on - amplitudes: an amplitude scaled by ``corr * exp(log_scale)`` corresponds to an - intensity scaled by ``(corr * exp(log_scale))**2``. Applying the amplitude - factors to intensities instead would leave a resolution-dependent error that - looks exactly like a scale or B-factor mismatch. - - Returns - ------- - Tuple[torch.Tensor, torch.Tensor] - Full-size ``I`` and ``I_sigma`` on the same scale as - ``get_corrected_data()`` squared. + """Return intensities and sigmas, shape (N,), in this dataset's units. Raises ------ ValueError - If this dataset carries no intensities (the input had no ``I``/``SIGI`` - columns), or if ``setup_scale`` / ``setup_anisotropy`` have not run. + If no intensity observations are available. """ - from torchref.base.alignment.normalization import ( - compute_anisotropy_correction, - ) - if self.I is None: - raise ValueError( - "No intensities on this dataset. The input reflection file had no " - "I/SIGI columns, so only amplitudes are available; use " - "get_corrected_data() or supply intensity data." - ) - if not hasattr(self, "log_scale") or self.log_scale is None: - raise ValueError("Scale not set up. Call setup_scale() first.") - if not hasattr(self, "U_aniso") or self.U_aniso is None: - raise ValueError( - "No anisotropy parameters available. Call fit_anisotropy() first." - ) - - s_vectors = self.get_scattering_vectors() - correction = compute_anisotropy_correction(s_vectors, self.U_aniso) - factor = (correction * torch.exp(self.log_scale)) ** 2 - - I_scaled = self.I * factor - I_sigma_scaled = ( - self.I_sigma * factor if self.I_sigma is not None else None - ) - return I_scaled, I_sigma_scaled + raise ValueError("No intensities on this dataset (I/SIGI required)") + return self.I, self.I_sigma def _corrected_or_raw_intensities(self): - """Return the scaled ``(I, I_sigma)``, cached against the - ``(log_scale, U_aniso)`` fingerprint. ``(None, None)`` when this dataset - carries no intensities, and the raw pair if scaling is not set up. - """ - - def _tv(t): - return (t.data_ptr(), t._version) if isinstance(t, torch.Tensor) else None - - if self.I is None: - return (None, None) - - fp = ( - _tv(getattr(self, "log_scale", None)), - _tv(getattr(self, "U_aniso", None)), - ) - if self._corrected_I_fp != fp or self._corrected_I_cache is None: - try: - self._corrected_I_cache = self.get_corrected_intensities() - except Exception: - self._corrected_I_cache = (self.I, self.I_sigma) - self._corrected_I_fp = fp - return self._corrected_I_cache + """Return intensity observations, or (None, None) when absent.""" + return self.I, self.I_sigma def generate_validation_set( self, @@ -4166,23 +3571,3 @@ def generate_validation_set( f"free={n_free} ({100*n_free/total:.1f}%), " f"val={n_val} ({100*n_val/total:.1f}%)" ) - - def parameters(self) -> List[Parameter]: - """ - The scaling tensors (``log_scale``, ``U_aniso``) to optimize. - - Despite the ``List[Parameter]`` annotation these are plain tensors with - ``requires_grad=False``; the caller must call ``requires_grad_(True)`` - before handing them to an optimizer (see ``DatasetCollection.scale``). - - Returns - ------- - list of torch.Tensor - Whichever of the two are set. - """ - params = [] - if self.log_scale is not None: - params.append(self.log_scale) - if self.U_aniso is not None: - params.append(self.U_aniso) - return params diff --git a/torchref/io/datasets/scaled_dataset.py b/torchref/io/datasets/scaled_dataset.py new file mode 100644 index 00000000..e8d44412 --- /dev/null +++ b/torchref/io/datasets/scaled_dataset.py @@ -0,0 +1,160 @@ +"""ReflectionData subclass exposing live scaler-owned observation corrections.""" + +from dataclasses import fields +from typing import TYPE_CHECKING + +import torch + +from .reflection_data import ReflectionData + +if TYPE_CHECKING: + from torchref.scaling.dataset_scaler import DatasetScaler + + +class _ScaledObservation: + """Keep dataclass initialization in raw storage and scale only public reads.""" + + def __init__(self, name, power): + self.name, self.power = name, power + + def __get__(self, obj, owner=None): + if obj is None: + return None + raw = obj.__dict__.get("_raw_" + self.name) + scaler = obj.__dict__.get("scaler") + if raw is None or scaler is None: + return raw + correction = scaler(obj.scale_key, obj.hkl).to(raw) + return raw * correction.pow(self.power) + + def __set__(self, obj, value): + if obj.__dict__.get("scaler") is not None: + raise AttributeError( + "Scaled observations are read-only; edit the raw dataset before scaling" + ) + obj.__dict__["_raw_" + self.name] = value + + +class ScaledDataset(ReflectionData): + """Expose a raw dataset through one live row of a shared DatasetScaler. + + Parameters + ---------- + data : ReflectionData + Measurements and metadata, copied without changing the source. Passing + a scaled dataset uses its raw observations, never a second correction. + scaler : DatasetScaler + Shared parameter owner, retained strongly even outside a collection. + key : str + Stable dataset key in the scaler. + + Notes + ----- + F/F_sigma (amplitude units) and I/I_sigma (intensity units), shape (N,), + are corrected read-only expressions. Explicit *_raw properties expose the + stored measurements. Selection/copy preserves the shared scaler; moving a + view also moves the shared scaler. Raw source datasets remain unchanged. + """ + + F = _ScaledObservation("F", 1) + F_sigma = _ScaledObservation("F_sigma", 1) + I = _ScaledObservation("I", 2) + I_sigma = _ScaledObservation("I_sigma", 2) + + def __init__(self, data: ReflectionData, scaler: "DatasetScaler", key: str) -> None: + if key not in scaler.keys: + raise KeyError(key) + raw = data.raw_data() if isinstance(data, ScaledDataset) else data + raw = raw.__select__(torch.arange(len(raw), device=raw.device)) + raw.spacegroup = raw.spacegroup.copy() + self._install_raw(raw) + self.scaler = scaler + self.scale_key = key + self.to(scaler.device) + + def _install_raw(self, raw): + self.scaler = None + super().__init__( + **{f.name: getattr(raw, f.name) for f in fields(ReflectionData)} + ) + self.masks = raw.masks + self.source = None + + @property + def F_raw(self) -> torch.Tensor | None: + """Uncorrected amplitudes, shape (N,), in input amplitude units.""" + return self._raw_F + + @property + def F_sigma_raw(self) -> torch.Tensor | None: + """Uncorrected amplitude sigmas, shape (N,), in input amplitude units.""" + return self._raw_F_sigma + + @property + def I_raw(self) -> torch.Tensor | None: + """Uncorrected intensities, shape (N,), in input intensity units.""" + return self._raw_I + + @property + def I_sigma_raw(self) -> torch.Tensor | None: + """Uncorrected intensity sigmas, shape (N,), in input intensity units.""" + return self._raw_I_sigma + + def raw_data(self) -> ReflectionData: + """Return an independent parameter-free copy of the raw observations.""" + values = { + f.name: getattr(self, f.name) + for f in fields(ReflectionData) + if f.name not in ("F", "F_sigma", "I", "I_sigma", "source") + } + values.update( + { + name: getattr(self, name + "_raw") + for name in ("F", "F_sigma", "I", "I_sigma") + } + ) + raw = ReflectionData(**values) + raw.masks = self.masks + result = raw.__select__(torch.arange(len(raw), device=raw.device)) + result.source = None + return result + + def __select__(self, indices: torch.Tensor, op=None) -> "ScaledDataset": + """Select reflection indices or a boolean mask, preserving live scaling.""" + raw = self.raw_data().__select__(indices, op=op) + return ScaledDataset(raw, self.scaler, self.scale_key) + + def copy(self) -> "ScaledDataset": + """Copy observations and metadata while sharing the scaler parameters.""" + return ScaledDataset(self.raw_data(), self.scaler, self.scale_key) + + def __deepcopy__(self, memo: dict) -> "ScaledDataset": + """Copy this view without duplicating the shared parameter owner.""" + result = self.copy() + memo[id(self)] = result + return result + + def validate_hkl( + self, hkl_ref: torch.Tensor, *, identity_hkl: torch.Tensor | None = None + ) -> "ScaledDataset": + """Align this view to HKL (H, 3), without modifying its source or scaler.""" + raw = self.raw_data().validate_hkl(hkl_ref, identity_hkl=identity_hkl) + scaler, key = self.scaler, self.scale_key + self._install_raw(raw) + self.scaler, self.scale_key = scaler, key + return self + + def _get_state(self) -> dict: + return { + "raw": self.raw_data()._get_state(), + "scaler": self.scaler.get_state(), + "key": self.scale_key, + } + + @classmethod + def _from_state(cls, state: dict, device=None) -> "ScaledDataset": + from torchref.scaling.dataset_scaler import DatasetScaler + + scaler = DatasetScaler.from_state(state["scaler"], device) + raw = ReflectionData._from_state(dict(state["raw"]), device) + return cls(raw, scaler, state["key"]) diff --git a/torchref/io/metadata.py b/torchref/io/metadata.py index dfe98079..4a1f378c 100644 --- a/torchref/io/metadata.py +++ b/torchref/io/metadata.py @@ -12,7 +12,7 @@ from __future__ import annotations import json -from dataclasses import dataclass, field, fields, asdict +from dataclasses import asdict, dataclass, field, fields from datetime import date from typing import Any, Dict, List, Optional @@ -205,6 +205,7 @@ def from_refinement(cls, refinement) -> RefinementMetadata: best-effort: anything unavailable is left unset, silently. """ import torch + from torchref import __version__ meta = cls(program_version=__version__) @@ -230,7 +231,7 @@ def from_refinement(cls, refinement) -> RefinementMetadata: try: rd = refinement.reflection_data with torch.no_grad(): - hkl, fobs, sigma, rfree_flags = rd() + hkl, fobs, sigma, rfree_flags = rd.hkl, rd.F, rd.F_sigma, rd.rfree_flags n_all = len(fobs) n_test = int(rfree_flags.sum().item()) if rfree_flags.dtype == torch.bool else int((~rfree_flags.bool()).sum().item()) n_work = n_all - n_test diff --git a/torchref/maps/difference_map.py b/torchref/maps/difference_map.py index 58b44242..69bcbeec 100644 --- a/torchref/maps/difference_map.py +++ b/torchref/maps/difference_map.py @@ -82,10 +82,12 @@ def __init__(self, data, data_reference, model, gridsize=None, ) self._collection.add_dataset("perturbed", data) self._collection.scale() + self.data_reference = self._collection["reference"] + self.data_perturbed = self._collection["perturbed"] # Use reference dataset for cell, spacegroup, hkl via super().__init__ super().__init__( - data=data_reference, + data=self.data_reference, model=model, gridsize=gridsize, map_type="Fcalc", # placeholder, calculate() is overridden diff --git a/torchref/refinement/targets/__init__.py b/torchref/refinement/targets/__init__.py index 279bdb6d..1c68fff5 100644 --- a/torchref/refinement/targets/__init__.py +++ b/torchref/refinement/targets/__init__.py @@ -5,8 +5,8 @@ """ from .adp import ( - ADPSigdTarget, ADPLocalityTarget, + ADPSigdTarget, ADPSimilarityTarget, ADPTarget, RigidBondTarget, @@ -66,6 +66,7 @@ ) __all__ = [ + "DatasetScalingTarget", # Base classes "Target", "ModelTarget", @@ -124,3 +125,5 @@ ] # Force-field, real-space, sampled-ML phase, and occupancy-diagnostic # targets are experimental and live in :mod:`torchref.experimental.targets`. + +from .dataset_scaling import DatasetScalingTarget diff --git a/torchref/refinement/targets/dataset_scaling.py b/torchref/refinement/targets/dataset_scaling.py new file mode 100644 index 00000000..542b55c4 --- /dev/null +++ b/torchref/refinement/targets/dataset_scaling.py @@ -0,0 +1,42 @@ +"""Joint observed-dataset scaling target, independent of structural models.""" + +from typing import TYPE_CHECKING + +import torch + +from torchref.base.targets.dataset_scaling import dataset_scaling_loss +from torchref.refinement.targets.base import Target + +if TYPE_CHECKING: + from torchref.scaling.dataset_scaler import DatasetScaler + + +class DatasetScalingTarget(Target): + """Fit a shared consensus to the scaler's training observations. + + Parameters + ---------- + scaler : DatasetScaler + Owner of the centered log-scale and anisotropy parameters. Its prepared + observations exclude held-out reflections from every participating dataset. + """ + + name = "dataset_scaling" + + def __init__(self, scaler: "DatasetScaler") -> None: + super().__init__(device=scaler.device) + self.scaler = scaler + self._adopt_device(scaler) + + def forward(self) -> torch.Tensor: + """Return the dimensionless loss per independent training contrast.""" + scaler = self.scaler + return ( + dataset_scaling_loss( + scaler.amplitudes, + scaler.sigmas, + scaler.log_corrections(scaler.hkl), + scaler.fit_mask, + ) + / scaler.n_contrasts + ) diff --git a/torchref/refinement/targets/xray/nll.py b/torchref/refinement/targets/xray/nll.py index edeb2d98..4524486d 100644 --- a/torchref/refinement/targets/xray/nll.py +++ b/torchref/refinement/targets/xray/nll.py @@ -1,6 +1,7 @@ -import torch from typing import TYPE_CHECKING +import torch + from torchref.base.targets.xray_likelihoods import ( amplitude_var_from_sigma_obs, nll_per_refl, @@ -24,13 +25,9 @@ class NLLXrayTarget(XrayTarget): here that does not. Was ``GaussianXrayTarget``; the taxonomy names the row, and "Gaussian" named the distribution, which ``nll_beta`` shares. - **Not a** :class:`SigmaAXrayTarget`, and deliberately so. Beyond needing no estimate, it - reads its amplitudes through :meth:`XrayTarget.get_data`, which goes via - ``ReflectionData._corrected_or_raw`` and falls back to **raw** amplitudes when the scaler - has not run; the sigma_A path calls ``get_corrected_data()``, which raises instead. - Moving this target onto that path would turn a silent fallback into a hard failure on - unscaled data -- a behaviour change, not a refactor. It would also lose the fused Triton - kernel and the ``median(sigma)*0.1`` clamp, neither of which the beta-variance path has. + Read observations through the dataset subset accessors, which expose live + corrections for ScaledDataset. The target uses a fused Triton kernel where + available and floors uncertainties at one tenth of their median. Attributes ---------- diff --git a/torchref/scaling/__init__.py b/torchref/scaling/__init__.py index 8125a15a..bba7ecce 100644 --- a/torchref/scaling/__init__.py +++ b/torchref/scaling/__init__.py @@ -1,23 +1,27 @@ -"""Scaling calculated structure factors onto observed data. +"""Scale observed datasets and calculated structure factors. Per-bin overall scale, anisotropic correction and bulk-solvent contribution. :class:`ScalerBase` is model-independent -- every method that needs ``F_calc`` takes it as an argument; :class:`Scaler` holds a :class:`~torchref.model.Model` and computes ``F_calc`` itself; :class:`CollectionScaler` fits one shared set of scales jointly across a dataset/model collection. :class:`SolventModel` supplies -the flat bulk-solvent term (k_sol, B_sol). +the flat bulk-solvent term (k_sol, B_sol). ``DatasetScaler`` independently fits +relative observed-data corrections; ``WilsonNormaliser`` supplies E values. """ +from torchref.scaling.collection_scaler import CollectionScaler from torchref.scaling.scaler import Scaler from torchref.scaling.scaler_base import ScalerBase from torchref.scaling.solvent import SolventModel -from torchref.scaling.collection_scaler import CollectionScaler from torchref.scaling.wilson import WilsonNormaliser __all__ = [ "Scaler", + "DatasetScaler", "ScalerBase", "SolventModel", "CollectionScaler", "WilsonNormaliser", ] + +from torchref.scaling.dataset_scaler import DatasetScaler diff --git a/torchref/scaling/collection_scaler.py b/torchref/scaling/collection_scaler.py index ab5af4d1..de655541 100644 --- a/torchref/scaling/collection_scaler.py +++ b/torchref/scaling/collection_scaler.py @@ -381,10 +381,11 @@ def forward_batched( Scaled complex structure factors of shape ``(T, n_reflections)``. """ component_sol_raw = self.compute_component_solvent_raw() - f_sol_batch = torch.einsum( - "tk,kr->tr", - fractions_matrix.to(component_sol_raw.dtype), - component_sol_raw, + # Keep real fraction multiplication and component accumulation identical + # to forward_mixed; complex GEMM changes rounding near solvent cancellation. + f_sol_batch = sum( + fractions_matrix[:, i, None] * component_sol_raw[i] + for i in range(component_sol_raw.shape[0]) ) return super().forward(fcalc_batch, f_sol_override=f_sol_batch) diff --git a/torchref/scaling/dataset_scaler.py b/torchref/scaling/dataset_scaler.py new file mode 100644 index 00000000..1cc18dc8 --- /dev/null +++ b/torchref/scaling/dataset_scaler.py @@ -0,0 +1,306 @@ +"""Joint relative scaling of observed datasets without a privileged reference. + +DatasetScaler owns all fitted corrections. ScaledDataset exposes one correction +through the reflection-data interface; no optimization state lives on raw datasets. +""" + +from collections.abc import Mapping +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from torchref.io import ReflectionData + +import torch +from torch import nn + +from torchref.base.targets.xray_likelihoods import SIGMA_FLOOR_ABS, SIGMA_FLOOR_FRAC +from torchref.config import get_float_dtype, get_int_dtype, normalize_device +from torchref.utils.device_mixin import DeviceMixin + + +def _identity_hkl(data): + """Keep Bijvoet observations separate while matching dataset identities.""" + return data.hkl if data.friedel_merged else data._hkl_for_sf() + + +class DatasetScaler(DeviceMixin, nn.Module): + """Fit N observed datasets to a shared sigma-weighted amplitude consensus. + + Parameters + ---------- + datasets : Mapping[str, ReflectionData] + At least two raw datasets in the same space-group setting and Friedel + convention. Sources are copied without mutation; membership is fixed. + device : torch.device or str, optional + Computation device, defaulting to the first dataset's device. Prepared + fitting arrays and owned copies move here without moving source data. + + Notes + ----- + Corrections have zero mean in log space over datasets, including anisotropy. + The six quadratic coefficients use dimensionless, normalized Miller indices; + they are not Cartesian atomic displacement parameters. Every corrected read + depends on all parameter rows through centering. ``fit`` freezes parameters + when it finishes; call ``requires_grad_(True)`` for custom differentiable use. + """ + + def __init__( + self, + datasets: Mapping[str, "ReflectionData"], + device: torch.device | str | None = None, + ) -> None: + super().__init__() + if len(datasets) < 2: + raise ValueError("Dataset scaling requires at least two datasets") + self.keys = tuple(datasets) + self.datasets = {} + for key, data in datasets.items(): + raw = data.raw_data() if hasattr(data, "raw_data") else data + owned = raw.__select__(torch.arange(len(raw), device=raw.device)) + owned.source = None + owned.spacegroup = raw.spacegroup.copy() + self.datasets[key] = owned + + self.device = normalize_device( + device if device is not None else next(iter(datasets.values())).device + ) + self.dtype_float = get_float_dtype() + self.raw_parameters = nn.Parameter( + torch.zeros((len(self.keys), 7), device=self.device, dtype=self.dtype_float) + ) + for name in ("hkl", "amplitudes", "sigmas", "fit_mask", "hkl_scale"): + self.register_buffer(name, None) + self.n_contrasts = 0 + self._initialized = False + self.to(self.device) + self.prepare() + + @property + def corrections(self) -> torch.Tensor: + """Centered coefficients, shape (N, 7), in dimensionless log units.""" + return self.raw_parameters - self.raw_parameters.mean(dim=0, keepdim=True) + + def design(self, hkl: torch.Tensor) -> torch.Tensor: + """Return the log-correction basis, shape (H, 7), for integer HKL (H, 3).""" + q = hkl.to(device=self.device, dtype=self.raw_parameters.dtype) / self.hkl_scale + h, k, l = q.unbind(dim=-1) + return torch.stack( + (torch.ones_like(h), h * h, k * k, l * l, 2 * h * k, 2 * h * l, 2 * k * l), + dim=-1, + ) + + def log_corrections(self, hkl: torch.Tensor) -> torch.Tensor: + """Return dimensionless log amplitude corrections, shape (N, H).""" + return self.corrections @ self.design(hkl).T + + def forward(self, key: str, hkl: torch.Tensor) -> torch.Tensor: + """Return positive amplitude factors (H,) for dataset key and HKL (H, 3).""" + row = self.keys.index(key) + return (self.design(hkl) @ self.corrections[row]).exp() + + def prepare(self) -> None: + """Prepare training arrays and reject disconnected or unidentified fits. + + This reads current source masks and measurements. No free/validation + observation enters initialization, uncertainty floors or the objective. + Changes to reflection sets or fit masks require a new scaler once fitted. + """ + data = list(self.datasets.values()) + if tuple(self.datasets) != self.keys: + raise ValueError("Dataset membership changed; construct a new scaler") + symmetry = {d.spacegroup.xhm for d in data} + if len(symmetry) != 1 or len({d.friedel_merged for d in data}) != 1: + raise ValueError( + "Datasets require compatible symmetry settings and Friedel conventions" + ) + hkls = [_identity_hkl(d).to(self.device) for d in data] + hkl, inverse = torch.unique(torch.cat(hkls), dim=0, return_inverse=True) + shape = (len(data), len(hkl)) + amplitudes = torch.zeros(shape, device=self.device, dtype=self.dtype_float) + sigmas = torch.ones_like(amplitudes) + valid = torch.zeros(shape, device=self.device, dtype=torch.bool) + held_out = torch.zeros(len(hkl), device=self.device, dtype=torch.bool) + start = 0 + for row, ds in enumerate(data): + if ds.F_raw is None or ds.F_sigma_raw is None: + raise ValueError(f"Dataset {self.keys[row]!r} requires F and SIGF") + idx = inverse[start : start + len(ds)] + start += len(ds) + if len(torch.unique(idx)) != len(idx): + raise ValueError( + f"Dataset {self.keys[row]!r} contains duplicate reflection identities" + ) + f = ds.F_raw.detach().to(amplitudes) + sigma = ds.F_sigma_raw.detach().to(sigmas) + present = ds.masks().to(device=self.device, dtype=torch.bool) + usable = present & torch.isfinite(f) & torch.isfinite(sigma) & (sigma > 0) + work = ds.work.mask.to(self.device) + held_out[idx] |= present & ~work + valid[row, idx] = usable + amplitudes[row, idx] = torch.where(usable, f, torch.zeros_like(f)) + sigmas[row, idx] = torch.where(usable, sigma, torch.ones_like(sigma)) + mask = valid & ~held_out.unsqueeze(0) + mask &= mask.sum(dim=0, keepdim=True) >= 2 + reached = {0} + for _ in data: + reached |= { + j + for i in tuple(reached) + for j in range(len(data)) + if bool((mask[i] & mask[j]).any()) + } + if len(reached) != len(data): + raise ValueError( + "Dataset work-set overlap is disconnected; relative scales are unidentified" + ) + for row in range(len(data)): + floor = (sigmas[row, mask[row]].median() * SIGMA_FLOOR_FRAC).clamp_min( + SIGMA_FLOOR_ABS + ) + sigmas[row] = sigmas[row].clamp_min(floor) + if self.hkl_scale is None: + self.hkl_scale = ( + hkl[mask.any(dim=0)].to(amplitudes).abs().amax(dim=0).clamp_min(1) + ) + self.hkl, self.amplitudes, self.sigmas, self.fit_mask = ( + hkl, + amplitudes, + sigmas, + mask, + ) + self.n_contrasts = int((mask.sum(dim=0) - 1).clamp_min(0).sum()) + # Rank concerns only geometry and overlap, not observed amplitudes. The + # small-column check runs on CPU because MPS has no SVD implementation. + design = self.design(hkl).detach().cpu() + mask_cpu = mask.cpu() + first = mask_cpu.to(get_int_dtype()).argmax(dim=0) + blocks = [] + for row in range(len(data)): + selected = mask_cpu[row] & (first != row) + if not bool(selected.any()): + continue + block = torch.zeros((int(selected.sum()), len(data), 7), dtype=design.dtype) + block[:, row] = design[selected] + block[torch.arange(len(block)), first[selected]] = -design[selected] + blocks.append(block[:, :-1].reshape(len(block), -1)) + rank = int(torch.linalg.matrix_rank(torch.cat(blocks))) if blocks else 0 + if rank != 7 * (len(data) - 1): + raise ValueError( + "Overlapping reflections do not identify overall scale and six anisotropic coefficients" + ) + + def initialize(self) -> None: + """Seed overall log scales from robust pairwise log-amplitude ratios.""" + rows, values = [], [] + n = len(self.keys) + for i in range(n): + for j in range(i): + mask = ( + self.fit_mask[i] + & self.fit_mask[j] + & (self.amplitudes[i] > 0) + & (self.amplitudes[j] > 0) + ) + if not bool(mask.any()): + continue + row = torch.zeros(n, dtype=self.dtype_float) + row[i], row[j] = 1, -1 + rows.append(row) + values.append( + (self.amplitudes[j, mask].log() - self.amplitudes[i, mask].log()) + .median() + .cpu() + ) + rows.append(torch.ones(n, dtype=self.dtype_float)) + values.append(torch.zeros((), dtype=self.dtype_float)) + initial = torch.linalg.lstsq(torch.stack(rows), torch.stack(values)).solution + with torch.no_grad(): + self.raw_parameters.zero_() + self.raw_parameters[:, 0].copy_(initial.to(self.raw_parameters)) + self._initialized = True + + def fit(self, nsteps: int = 10, max_iter: int = 100) -> dict: + """Fit joint corrections with L-BFGS and freeze the resulting parameters. + + Parameters + ---------- + nsteps : int + Number of outer L-BFGS steps. + max_iter : int + Maximum iterations per outer step. + + Returns + ------- + dict + Initial/final normalized loss, contrast count and centered coefficients. + """ + from torchref.refinement.loss_state import LossState + from torchref.refinement.targets.dataset_scaling import DatasetScalingTarget + + if nsteps < 1 or max_iter < 1: + raise ValueError("nsteps and max_iter must be positive") + self.prepare() + if not self._initialized: + self.initialize() + saved = self.raw_parameters.detach().clone() + self.requires_grad_(True) + target = DatasetScalingTarget(self) + state = LossState(device=self.device) + state.register_target("scaling/datasets", target) + before = float(target().detach()) + optimizer = torch.optim.LBFGS( + self.parameters(), max_iter=max_iter, line_search_fn="strong_wolfe" + ) + try: + state.run(optimizer, nsteps=nsteps, log=False, context="dataset_scaler.fit") + after = float(target().detach()) + if not torch.isfinite(self.raw_parameters).all() or not torch.isfinite( + torch.tensor(after) + ): + raise RuntimeError( + "Dataset scaling produced non-finite parameters or loss" + ) + except Exception: + with torch.no_grad(): + self.raw_parameters.copy_(saved) + raise + finally: + self.requires_grad_(False) + return { + "loss_before": before, + "loss_after": after, + "n_contrasts": self.n_contrasts, + "corrections": dict( + zip(self.keys, self.corrections.detach().cpu().tolist()) + ), + } + + def get_state(self) -> dict: + """Return raw source states and the shared fitted parameter state.""" + return { + "datasets": {k: d._get_state() for k, d in self.datasets.items()}, + "parameters": self.raw_parameters.detach().cpu(), + "hkl_scale": self.hkl_scale.detach().cpu(), + "initialized": self._initialized, + } + + @classmethod + def from_state( + cls, state: dict, device: torch.device | str | None = None + ) -> "DatasetScaler": + """Restore a shared scaler and its raw sources on the requested device.""" + from torchref.io.datasets.reflection_data import ReflectionData + + obj = cls( + { + k: ReflectionData._from_state(dict(v), device) + for k, v in state["datasets"].items() + }, + device=device, + ) + with torch.no_grad(): + obj.raw_parameters.copy_(state["parameters"].to(obj.raw_parameters)) + obj.hkl_scale.copy_(state["hkl_scale"].to(obj.hkl_scale)) + obj._initialized = state["initialized"] + obj.requires_grad_(False) + return obj From 953f404288536bbcb45e5f3bc2407c13f57ff6d0 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 18 Sep 2026 00:23:21 +0200 Subject: [PATCH 171/250] Trim difference refinement tests and redundant access paths --- docs/changelog.rst | 2 + paper/make_ded_maps.py | 86 ----- paper/probe_data_scale_objective.py | 65 ---- paper/probe_ded_metric_space.py | 164 --------- paper/probe_joint_scale_objective.py | 197 ---------- paper/probe_movement_recovery.py | 336 ------------------ paper/probe_two_moment_collinearity.py | 222 ------------ tests/files/hkl/dark_half1.hkl | 335 ----------------- tests/files/hkl/dark_half2.hkl | 335 ----------------- tests/integration/test_cli_two_moment_mtz.py | 279 +++++++-------- .../io/test_collection_stack_accessors.py | 79 +--- tests/unit/io/test_crystfel_hkl.py | 167 +++------ tests/unit/io/test_fcalc_add_noise.py | 84 +---- tests/unit/io/test_intensity_accessors.py | 193 ---------- .../model/test_batched_component_fcalcs.py | 72 +--- .../model/test_model_collection_fractions.py | 81 ++--- ...test_collection_target_characterisation.py | 146 ++------ .../refinement/test_collection_taxonomy.py | 84 +---- .../refinement/test_intensity_observable.py | 135 ++----- .../refinement/test_two_moment_intensity.py | 209 +++-------- .../test_collection_joint_scale_fit.py | 70 +--- .../scaling/test_collection_scaler_batched.py | 84 +---- tests/unit/scaling/test_dataset_scaler.py | 50 ++- .../scaling/test_f_sol_override_contract.py | 104 ++---- torchref/cli/_common.py | 25 +- torchref/cli/collection_difference_refine.py | 67 ++-- torchref/cli/difference_map.py | 10 - torchref/cli/simulate_noisy_data.py | 8 +- torchref/cli/validate_ded.py | 9 +- torchref/io/datasets/base.py | 2 +- torchref/io/datasets/collection.py | 16 +- torchref/io/datasets/fcalc_data.py | 2 +- torchref/io/datasets/reflection_data.py | 22 +- torchref/io/hkl.py | 4 - torchref/model/model_collection.py | 77 +--- .../refinement/targets/collection/_specs.py | 18 +- .../refinement/targets/collection/base.py | 75 +--- .../targets/collection/intensity.py | 9 - .../refinement/targets/collection/xray.py | 7 - .../refinement/targets/xray/observable.py | 31 +- torchref/refinement/targets/xray/sigma_a.py | 4 - torchref/scaling/collection_scaler.py | 20 +- 42 files changed, 508 insertions(+), 3477 deletions(-) delete mode 100644 paper/make_ded_maps.py delete mode 100644 paper/probe_data_scale_objective.py delete mode 100644 paper/probe_ded_metric_space.py delete mode 100644 paper/probe_joint_scale_objective.py delete mode 100644 paper/probe_movement_recovery.py delete mode 100644 paper/probe_two_moment_collinearity.py delete mode 100644 tests/unit/io/test_intensity_accessors.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 2b832296..bd1b4ce1 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,8 @@ Changelog Unreleased ---------- +- Read scaled observations directly in subset and collection accessors, avoiding unused sigma/amplitude corrections and removing redundant internal forwarding helpers. +- Remove one-off diagnostic scripts and consolidate difference-refinement regression tests while retaining numerical and output-format coverage. - Restore the cell and space-group imports needed by difference-density validation setup. - Keep ``torchref.phased-difference-map`` as an alias for ``torchref.difference-map``. - Add a matched classic/AMBER refinement benchmark with shared prepared hydrogens, initial riding-parameter gradient calibration, and per-cycle R factors and geometry diagnostics. diff --git a/paper/make_ded_maps.py b/paper/make_ded_maps.py deleted file mode 100644 index 31bc3377..00000000 --- a/paper/make_ded_maps.py +++ /dev/null @@ -1,86 +0,0 @@ -#!/usr/bin/env python -"""Turn a ``torchref.difference-refine`` results MTZ into CCP4 maps for PyMOL/Coot. - -Coot opens the MTZ directly (File > Auto Open MTZ, or pick the column pair), so this -exists for PyMOL, which wants a real map. Every map is computed on ``PHIC_diff``, so -maps from different runs are directly comparable -- they share phases. - -Which coefficient is which: - -``mDFop-DFc`` - The difference map. ``m`` is a normalised inverse-variance weight, **not** a sigma_A - figure of merit. Contour at +-3 sigma. -``mDFop-DFc_corr`` - The same, from the activation-decontaminated light amplitude. Present only when the - run had ``--two-moment``. -``DDF`` - ``DF_corr - DF``: the correction itself, as a map. Featureless against resolution - means the correction is collinear with a scale or overall-B error and should be - distrusted; structure in it is signal. -``2mDFop-DFc`` - The 2Fo-Fc analogue, for seeing the model in its density. -""" - -import argparse -import sys -from pathlib import Path - -import gemmi -import numpy as np - -# label -> (amplitude column, phase column). Skipped silently when absent. -# -# The dark-phased entries come first because they are the default output and the -# construction ``torchref.validate-ded`` correlates against. ``ddf`` and ``wdf`` used to -# be paired with ``PHIC_diff``, the *model* difference phase -- the right amplitude on -# the wrong phase, and a different map from the one being validated. -# -# The ``PHIC_diff`` entries need ``--all-columns`` on the writer. They are phased -# difference *residuals*: the light state's model phases enter the observed amplitude, -# so they are model-biased where the dark-phased maps are not. -MAPS = { - "ded": ("DELFWT", "PHDELWT"), - "ded_corr": ("DELFWT_corr", "PHDELWT"), - "ddf": ("DDF", "PHDELWT"), - "ext": ("FWT", "PHWT"), - "ded_phased": ("mDFop-DFc", "PHIC_diff"), - "ded_phased_corr": ("mDFop-DFc_corr", "PHIC_diff"), - "ded2_phased": ("2mDFop-DFc", "PHIC_diff"), - "ded2_phased_corr": ("2mDFop-DFc_corr", "PHIC_diff"), -} - - -def main(): - ap = argparse.ArgumentParser(description=__doc__, - formatter_class=argparse.RawDescriptionHelpFormatter) - ap.add_argument("mtz", help="fractions_*_difference_data.mtz from a refine run") - ap.add_argument("-o", "--outdir", default=".", help="where to write the .ccp4 files") - ap.add_argument("--prefix", default="", help="prefix for the output names") - # 3.0 matches the library's FFT oversampling; below ~2.5 the peaks shift. - ap.add_argument("--sample-rate", type=float, default=3.0) - args = ap.parse_args() - - out = Path(args.outdir) - out.mkdir(parents=True, exist_ok=True) - mtz = gemmi.read_mtz_file(args.mtz) - have = {c.label for c in mtz.columns} - - print(f"{args.mtz}\n{'map':12s} {'coefficient':22s} {'rms':>10s} {'peak':>10s}") - print("-" * 58) - for name, (f, ph) in MAPS.items(): - if f not in have or ph not in have: - continue - grid = mtz.transform_f_phi_to_map(f, ph, sample_rate=args.sample_rate) - ccp4 = gemmi.Ccp4Map() - ccp4.grid = grid - ccp4.update_ccp4_header() - path = out / f"{args.prefix}{name}.ccp4" - ccp4.write_ccp4_map(str(path)) - a = np.array(grid, copy=False) - print(f"{name:12s} {f:22s} {a.std():10.5f} {np.abs(a).max():10.5f}") - print(f"\nwritten to {out.resolve()}") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/paper/probe_data_scale_objective.py b/paper/probe_data_scale_objective.py deleted file mode 100644 index c64a1b63..00000000 --- a/paper/probe_data_scale_objective.py +++ /dev/null @@ -1,65 +0,0 @@ -#!/usr/bin/env python -"""Measure joint dataset scaling on work and held-out reflections. - -Report amplitude agreement and propagated-variance residuals before and after -fitting centered overall and anisotropic corrections with DatasetCollection.scale. -""" - -import argparse -import json -from pathlib import Path - -import torch - -from torchref.cli.collection_difference_refine import setup_dataset_collection - - -def score(collection, subset: str) -> dict: - """Return symmetric amplitude disagreement and chi-square for one subset.""" - a, b = list(collection.values()) - mask = getattr(a, subset).mask & getattr(b, subset).mask - with torch.no_grad(): - fa, fb = a.F[mask], b.F[mask] - variance = a.F_sigma[mask].square() + b.F_sigma[mask].square() - valid = torch.isfinite(variance) & (variance > 0) - residual = fa[valid] - fb[valid] - denominator = (fa[valid].abs() + fb[valid].abs()).sum() - return { - "n": int(valid.sum()), - "R_symmetric": float(2 * residual.abs().sum() / denominator), - "chi2": float((residual.square() / variance[valid]).mean()), - } - - -def main() -> int: - """Fit the requested reflection pair and print or save its diagnostics.""" - parser = argparse.ArgumentParser(description=__doc__) - root = Path(__file__).resolve().parent / "figure4_difference_refinement" - parser.add_argument("--dark-sf", default=str(root / "data/8QL2-sf.cif")) - parser.add_argument("--light-sf", default=str(root / "data/7YYZ-light.mtz")) - parser.add_argument("--dmin", type=float, default=2.2) - parser.add_argument("--device", default="cpu") - parser.add_argument("-o", "--out") - args = parser.parse_args() - collection = setup_dataset_collection( - args.dark_sf, args.light_sf, args.dmin, torch.device(args.device) - ) - from torchref import DatasetCollection - - raw_collection = DatasetCollection(device=args.device, verbose=0) - for name, data in collection: - raw_collection.add_dataset(name, data.raw_data()) - collection = raw_collection - report = {"before": {s: score(collection, s) for s in ("work", "free")}} - collection.scale() - report["after"] = {s: score(collection, s) for s in ("work", "free")} - report["fit"] = collection.scaling_metrics - result = json.dumps(report, indent=2) - print(result) - if args.out: - Path(args.out).write_text(result + "\n") - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/paper/probe_ded_metric_space.py b/paper/probe_ded_metric_space.py deleted file mode 100644 index e5cf720d..00000000 --- a/paper/probe_ded_metric_space.py +++ /dev/null @@ -1,164 +0,0 @@ -#!/usr/bin/env python -"""Is the DED correlation a fair way to compare an amplitude and an intensity target? - -``torchref.validate-ded`` correlates ``WDFo`` against ``WDFc``, and ``WDFc`` is -``(|F_mixed| - |F_dark|) * w`` -- a weighted **amplitude** difference. That is, up to the -weight and the Fourier transform, exactly the residual ``CollectionDifferenceTarget`` -minimises. Scoring an amplitude target on it is close to scoring it on its own objective, -so a win there is not evidence. - -This computes the same correlation in **both** spaces, on the same reflections and the -same models: - - amplitude obs Fo_light - Fo_dark calc |Fc_light| - |Fc_dark| - intensity obs Io_light - Io_dark calc |Fc_light|^2 - |Fc_dark|^2 - -If the ranking flips between the two, neither is decisive and the comparison has to be -made on something neither target optimises. - -**READ THIS BEFORE USING THE FREE-SET NUMBERS.** They are reported for completeness and -they cannot answer the question. A difference feature is compact in real space and -therefore spread over ALL of reciprocal space, so a held-out subset of reflections does -not contain a reduced-precision version of it -- it does not contain it. Measured on this -pair: the same models score 0.52 on all reflections and 0.21 on the free 3.5%, while -restricting in the other domain goes the other way, 0.53 over the full cell to 0.85 on the -0.12% of voxels around the ligand. Localising helps in real space and destroys the signal -in reciprocal space. - -So a reflection-wise hold-out validates a global scalar (R-free) and nothing local. To -cross-validate a local difference feature, hold out in the domain the feature lives in -- -an omit refinement -- or use an independent dataset, or ground truth. - -``F_calc`` is taken from each run's own results MTZ -- the scaled, mixed amplitudes that -run produced -- so no model is re-scaled here and each arm is scored on what it actually -built. Observed intensities come from one collection build shared by every arm. -""" - -import argparse -import json -import sys -from pathlib import Path - -import numpy as np -import reciprocalspaceship as rs -import torch - - -def cc(a, b): - a, b = np.asarray(a, float), np.asarray(b, float) - ok = np.isfinite(a) & np.isfinite(b) - if ok.sum() < 3: - return float("nan") - return float(np.corrcoef(a[ok], b[ok])[0, 1]) - - -def observed_intensities(dark_sf, light_sf, d_min, device): - """Scaled ``(hkl, I_dark, I_light)`` from one collection build.""" - from torchref.cli.collection_difference_refine import setup_dataset_collection - - dc = setup_dataset_collection(dark_sf, light_sf, d_min, device) - I_d, _ = dc["dark"].get_corrected_intensities() - I_l, _ = dc["light"].get_corrected_intensities() - hkl = dc.hkl.cpu().numpy() - return hkl, I_d.cpu().numpy(), I_l.cpu().numpy() - - -_PAIRED = [] - - -def main(): - ap = argparse.ArgumentParser(description=__doc__, - formatter_class=argparse.RawDescriptionHelpFormatter) - fig4 = Path(__file__).resolve().parent / "figure4_difference_refinement" - ap.add_argument("mtz", nargs="+", help="one results MTZ per arm (label=path accepted)") - ap.add_argument("--dark-sf", default=str(fig4 / "data/8QL2-sf.cif")) - ap.add_argument("--light-sf", default=str(fig4 / "data/7YYZ-light.mtz")) - ap.add_argument("--dmin", type=float, default=2.2) - ap.add_argument("--device", default="cpu") - ap.add_argument("-o", "--out", default=None) - args = ap.parse_args() - - hkl, Id_full, Il_full = observed_intensities( - args.dark_sf, args.light_sf, args.dmin, torch.device(args.device)) - key = {tuple(h): i for i, h in enumerate(hkl)} - - rows = [] - for spec in args.mtz: - label, _, path = spec.partition("=") - if not path: - label, path = Path(spec).parent.name, spec - ds = rs.read_mtz(path).reset_index() - H = ds[["H", "K", "L"]].to_numpy() - idx = np.array([key.get(tuple(h), -1) for h in H]) - ok = idx >= 0 - - Fo_d = ds["Fo_dark"].to_numpy(float) - Fo_l = ds["Fo_light"].to_numpy(float) - Fc_d = ds["Fc_dark"].to_numpy(float) - Fc_l = ds["FC"].to_numpy(float) # was Fc_light - free = ds["FreeR_flag_light"].to_numpy() == 0 - - Id = np.full(len(H), np.nan); Il = np.full(len(H), np.nan) - Id[ok] = Id_full[idx[ok]]; Il[ok] = Il_full[idx[ok]] - - dFo, dFc = Fo_l - Fo_d, Fc_l - Fc_d # amplitude difference - dIo, dIc = Il - Id, Fc_l**2 - Fc_d**2 # intensity difference - - r = {"arm": label} - sel_map = {"work": ~free & ok, "free": free & ok} - for name, sel in sel_map.items(): - r[f"amp_{name}"] = cc(dFo[sel], dFc[sel]) - r[f"int_{name}"] = cc(dIo[sel], dIc[sel]) - r[f"n_{name}"] = int(sel.sum()) - rows.append(r) - if len(_PAIRED) < 2: - _PAIRED.append({"arm": label, "amp": (dFo, dFc), "int": (dIo, dIc), - "sel": sel_map}) - - print() - print("Difference-signal correlation, same reflections and models, two spaces") - print("=" * 78) - print(f"{'arm':16s} {'amp work':>10s} {'amp free':>10s} {'int work':>10s} " - f"{'int free':>10s} {'n free':>8s}") - print("-" * 78) - for r in rows: - print(f"{r['arm']:16s} {r['amp_work']:10.4f} {r['amp_free']:10.4f} " - f"{r['int_work']:10.4f} {r['int_free']:10.4f} {r['n_free']:8d}") - print("-" * 78) - # Paired bootstrap over reflections. A correlation is not a mean, so the - # difference of two CCs has no closed-form error; resampling the SAME reflections - # for both arms keeps the comparison paired, which matters because most of the - # scatter is shared signal that cancels. - if len(rows) >= 2 and _PAIRED: - print() - print("Paired bootstrap on the CC difference (4000 resamples over reflections)") - print("-" * 78) - a, b = _PAIRED[0], _PAIRED[1] - rng = np.random.default_rng(0) - for sp, (oa, ca, ob, cb) in (("amp", a["amp"] + b["amp"]), - ("int", a["int"] + b["int"])): - for st in ("work", "free"): - m = a["sel"][st] - ia, ja = oa[m], ca[m] - ib, jb = ob[m], cb[m] - keep = np.isfinite(ia) & np.isfinite(ja) & np.isfinite(ib) & np.isfinite(jb) - ia, ja, ib, jb = ia[keep], ja[keep], ib[keep], jb[keep] - n = len(ia) - d = np.corrcoef(ia, ja)[0, 1] - np.corrcoef(ib, jb)[0, 1] - boot = np.empty(4000) - for k in range(4000): - s_ = rng.integers(0, n, n) - boot[k] = (np.corrcoef(ia[s_], ja[s_])[0, 1] - - np.corrcoef(ib[s_], jb[s_])[0, 1]) - lo, hi = np.percentile(boot, [2.5, 97.5]) - flag = "" if lo <= 0 <= hi else " <-- CI excludes 0" - print(f" {sp}_{st:4s} d = {d:+.4f} 95% CI [{lo:+.4f}, {hi:+.4f}]" - f" n={n}{flag}") - print(f"\n positive => {_PAIRED[0]['arm']} predicts the difference better") - if args.out: - Path(args.out).write_text(json.dumps(rows, indent=2)) - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/paper/probe_joint_scale_objective.py b/paper/probe_joint_scale_objective.py deleted file mode 100644 index 684387cc..00000000 --- a/paper/probe_joint_scale_objective.py +++ /dev/null @@ -1,197 +0,0 @@ -#!/usr/bin/env python -"""Which objective should the joint (model-to-data) scale fit use? - -``CollectionScaler.refine_lbfgs_joint`` used to hand-roll a Rice likelihood at -``beta = sigma_obs**2`` on an objective with no normaliser at all. It now builds a row of -``XRAY_TARGETS``, defaulting to unit-weight ``ls`` -- the same default the single-dataset -scale fit was moved to after measurement. - -This measures the swap, on the figure-4 pair, against the two things that could make a -single-run comparison meaningless: - -* **Thread nondeterminism.** TorchRef's CPU F_calc is not bit-reproducible, and R-free - jitter of order 0.017 has been measured at zero true difference. So every arm is - repeated and the spread is reported alongside the mean; a difference smaller than the - spread is not a result. -* **Circularity.** Each objective is scored on ``rfree``, which none of them optimises - (they all fit the work set), and by the same ``rfactor_work_free`` for every arm. - -The legacy fit is reconstructed here rather than kept selectable: the point is to measure -what it did, not to preserve it. -""" - -import argparse -import json -import statistics -import sys -from pathlib import Path - -import torch - - -def build(dark_sf, light_sf, dark_pdb, light_pdb, cif, d_min, fraction, device): - from torchref.cli.collection_difference_refine import ( - setup_dataset_collection, - setup_model_collection, - ) - - mc = setup_model_collection( - dark_pdb, light_pdb, [1.0 - fraction, fraction], cif, d_min, device, 0 - ) - dc = setup_dataset_collection(dark_sf, light_sf, d_min, device) - return dc, mc - - -def legacy_joint(scaler, dc, mc, nsteps=3, max_iter=200): - """The pre-change fit: hand-rolled Rice at beta=sigma**2, no normaliser.""" - import torch.nn as nn - - from torchref.base.reciprocal import get_scattering_vectors - from torchref.base.targets.xray_likelihoods import complex_var_from_beta, rice_math - from torchref.refinement.loss_state import LossState - from torchref.refinement.model_error_estimation.sigma_a import ( - SigmaAEstimator, - epsilon_from_hkl, - ) - from torchref.scaling.collection_scaler import CollectionScaler - - keys = [k for k in ([mc.dark_key] + mc.timepoint_names) if k in dc] - cache = {} - for name in keys: - data, model = dc[name], mc[name] - with torch.no_grad(): - fc = model(data.hkl).detach() - fracs = model.fractions.detach() - scaled0 = CollectionScaler.forward_mixed(scaler, fc, fracs) - amp0 = torch.abs(scaled0).reshape(-1) - fobs, sig = data.get_corrected_data() - eps0 = epsilon_from_hkl(data.hkl, getattr(data, "spacegroup", None)).to(amp0.dtype) - s = get_scattering_vectors(data.hkl, data.cell) - dss0 = (torch.norm(s, dim=1) ** 2).to(amp0.dtype) - est = SigmaAEstimator().get( - fobs.to(amp0.dtype).reshape(-1), amp0, data.centric, eps0, dss0, - data.free.mask, sigma_obs=sig.to(amp0.dtype).reshape(-1), - ) - cache[name] = (fc, fracs, est.beta, est.epsilon, data.work, data.centric) - - class _T(nn.Module): - name = "scaler/joint" - - def forward(self): - total, n = torch.tensor(0.0, device=scaler.device), 0 - for nm in keys: - fc, fracs, beta, eps, work, cen = cache[nm] - scaled = CollectionScaler.forward_mixed(scaler, fc, fracs) - amp = torch.abs(scaled).reshape(-1) - fo = work.F.to(amp.dtype) - bw = work.select(beta).to(fo.dtype) - ew = work.select(eps).to(fo.dtype) if eps is not None else None - loss = rice_math( - fo, work.select(amp), complex_var_from_beta(bw, ew), work.select(cen) - ) - if torch.isfinite(loss): - total, n = total + loss, n + 1 - if n: - total = total / n - return total + torch.sum(scaler.U**2) - - state = LossState(device=scaler.device) - state.register_target("scaler/joint", _T()) - opt = torch.optim.LBFGS( - scaler.parameters(), lr=1.0, max_iter=max_iter, history_size=10, - line_search_fn="strong_wolfe", - ) - state.run(opt, nsteps=nsteps, log=False, context="probe.legacy_joint") - - -def rfactors(scaler, dc, mc): - """``{key: (rwork, rfree)}`` for every dataset under the current scale.""" - from torchref.base.metrics.rfactor import rfactor_work_free - from torchref.scaling.collection_scaler import CollectionScaler - - out = {} - with torch.no_grad(): - for name in ([mc.dark_key] + mc.timepoint_names): - if name not in dc: - continue - data, model = dc[name], mc[name] - fc = model(data.hkl).detach() - scaled = CollectionScaler.forward_mixed(scaler, fc, model.fractions.detach()) - out[name] = rfactor_work_free(data, torch.abs(scaled)) - return out - - -def main(): - ap = argparse.ArgumentParser(description=__doc__) - fig4 = Path(__file__).resolve().parent / "figure4_difference_refinement" - ap.add_argument("--dark-sf", default=str(fig4 / "data/8QL2-sf.cif")) - ap.add_argument("--light-sf", default=str(fig4 / "data/7YYZ-light.mtz")) - ap.add_argument("--dark-pdb", default=str(fig4 / "data/8QL2_no_altloc.pdb")) - ap.add_argument("--light-pdb", default=str(fig4 / "work_no_altloc.pdb")) - ap.add_argument("--cif", nargs="*", default=[str(fig4 / "data/IBL_grade.cif")]) - ap.add_argument("--dmin", type=float, default=2.2) - ap.add_argument("--fraction", type=float, default=0.22) - ap.add_argument("--repeats", type=int, default=5) - ap.add_argument("--device", default="cpu") - ap.add_argument("-o", "--out", default=None) - args = ap.parse_args() - - from torchref.scaling.collection_scaler import CollectionScaler - - dev = torch.device(args.device) - dc, mc = build(args.dark_sf, args.light_sf, args.dark_pdb, args.light_pdb, - args.cif, args.dmin, args.fraction, dev) - - arms = ["legacy_rice_unnormalised", "ls", "nll", "ml_noalpha"] - results = {a: {"rfree_dark": [], "rfree_light": [], "rwork_dark": []} for a in arms} - - for rep in range(args.repeats): - for arm in arms: - # A fresh scaler each time: `initialize()` reseeds every parameter, so no arm - # inherits another's answer. - scaler = CollectionScaler(dc, mc, verbose=0).initialize() - if arm == "legacy_rice_unnormalised": - legacy_joint(scaler, dc, mc) - else: - scaler.refine_lbfgs_joint(verbose=False, scale_target=arm) - rf = rfactors(scaler, dc, mc) - dark, light = mc.dark_key, mc.timepoint_names[0] - results[arm]["rwork_dark"].append(rf[dark][0]) - results[arm]["rfree_dark"].append(rf[dark][1]) - results[arm]["rfree_light"].append(rf[light][1]) - print(f" repeat {rep + 1}/{args.repeats} done", flush=True) - - def ms(v): - m = statistics.mean(v) - s = statistics.stdev(v) if len(v) > 1 else 0.0 - return m, s - - print() - print(f"Joint scale fit, {args.repeats} repeats -- scored on reflections it did not fit") - print("=" * 82) - print(f"{'objective':28s} {'rwork_dark':>18s} {'rfree_dark':>18s} {'rfree_light':>18s}") - print("-" * 82) - for a in arms: - cells = [] - for k in ("rwork_dark", "rfree_dark", "rfree_light"): - m, s = ms(results[a][k]) - cells.append(f"{m:.5f}+-{s:.5f}") - print(f"{a:28s} {cells[0]:>18s} {cells[1]:>18s} {cells[2]:>18s}") - print("-" * 82) - base = results["legacy_rice_unnormalised"] - for a in arms[1:]: - for k in ("rfree_dark", "rfree_light"): - # Paired by repeat index: the same thread-nondeterminism realisation. - d = [x - y for x, y in zip(results[a][k], base[k])] - m, s = ms(d) - flag = " <-- exceeds its own spread" if abs(m) > 2 * (s or 1e-9) else "" - print(f"{a:28s} d({k}) vs legacy = {m:+.5f} +- {s:.5f}{flag}") - - if args.out: - Path(args.out).write_text(json.dumps(results, indent=2)) - print(f"written: {args.out}") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/paper/probe_movement_recovery.py b/paper/probe_movement_recovery.py deleted file mode 100644 index 7daaa935..00000000 --- a/paper/probe_movement_recovery.py +++ /dev/null @@ -1,336 +0,0 @@ -#!/usr/bin/env python -"""Does difference refinement recover the true displacement, or overshoot it? - -Nothing measured on real data can answer this, because the true light-state structure is -never known -- the published one is itself a refinement. So the displacement is *injected* -here and the refinement is asked to find it. - -Construction: - - 1. take a model, call it dark; - 2. displace a contiguous stretch of it by a known vector -- that is the true light state; - 3. build the merged light intensity the two-moment physics predicts, - ``|F_D + alpha dF|^2 + sigma_alpha^2 |dF|^2``, with a chosen alpha and lambda; - 4. add noise with sigmas grafted from a real dataset; - 5. hand the CLI the dark model as the *starting point for both states*, so the - refinement has to discover the displacement rather than be handed it. - -The recovered displacement is then compared with the injected one. A refinement that -believes the whole observed difference -- including the positive contamination that -crystal-to-crystal activation spread puts there -- should have to move the model further -than the truth to explain it. - -Writes the two MTZs and prints the CLI command; run that, then re-run with ``--score`` to -compare the refined models against the injected truth. -""" - -import argparse -import json -import sys -from pathlib import Path - -import numpy as np -import torch - - -def displaced_copy(pdb_path, out_path, chain, first, last, shift, verbose=0): - """Write a copy of `pdb_path` with residues [first, last] of `chain` moved by `shift`. - - Returns the number of atoms actually moved, so a selection that matched nothing is - caught rather than silently producing a zero-displacement truth. - """ - import gemmi - - st = gemmi.read_structure(str(pdb_path)) - moved = 0 - for model in st: - for ch in model: - if ch.name != chain: - continue - for res in ch: - if first <= res.seqid.num <= last: - for atom in res: - atom.pos = gemmi.Position( - atom.pos.x + shift[0], - atom.pos.y + shift[1], - atom.pos.z + shift[2], - ) - moved += 1 - break - if moved == 0: - raise ValueError( - f"selection chain {chain} residues {first}-{last} matched no atoms; " - f"the injected displacement would be zero" - ) - st.write_pdb(str(out_path)) - if verbose: - print(f" displaced {moved} atoms by {np.linalg.norm(shift):.3f} A") - return moved - - -def simulate(dark_pdb, light_pdb, reference_mtz, out_dir, alpha, lam, d_min, - sigma_mul, seed, device): - """Write dark.mtz / light.mtz carrying two-moment intensities plus noise.""" - from torchref import ReflectionData - from torchref.cli._common import load_model - from torchref.io.datasets import FcalcDataset - - ref = ReflectionData(device=str(device), verbose=0).load_mtz(str(reference_mtz)) - ref.cut_res(highres=d_min) - - md = load_model(str(dark_pdb), max_res=d_min, device=device, verbose=0) - ml = load_model(str(light_pdb), max_res=d_min, device=device, verbose=0) - - with torch.no_grad(): - hkl = ref.hkl - F_D = ref.structure_factors(md, recalc=True) - F_L = ref.structure_factors(ml, recalc=True) - dF = F_L - F_D - - sigma_alpha_sq = alpha * (1.0 - alpha) * lam - I_dark = F_D.abs() ** 2 - I_light = (F_D + alpha * dF).abs() ** 2 + sigma_alpha_sq * dF.abs() ** 2 - - frac_contam = float( - (sigma_alpha_sq * dF.abs() ** 2 / I_light.clamp(min=1e-12)).median() - ) - - out_dir = Path(out_dir) - out_dir.mkdir(parents=True, exist_ok=True) - paths = {} - for name, intensity in (("dark", I_dark), ("light", I_light)): - ds = FcalcDataset( - hkl=hkl.clone(), cell=ref.cell, spacegroup=ref.spacegroup, device=device - ) - # Phase is irrelevant to the written intensities but set_fcalc wants a complex. - ds.set_fcalc((intensity.clamp(min=0).sqrt() + 0j).to(torch.complex64)) - noisy = ds.add_noise(sigma_mul=sigma_mul, seed=seed, verbose=False) - # Write the intensities add_noise actually drew, negatives and all. - data = ReflectionData(device=str(device), verbose=0).from_tensors( - hkl=hkl.clone(), - F=noisy.fcalc_amp.clone(), - F_sigma=noisy.fobs_sigma.clone(), - cell=ref.cell, - spacegroup=ref.spacegroup, - rfree_flags=ref.rfree_flags.clone(), - device=str(device), - verbose=0, - ) - data.I = noisy.I.clone() - data.I_sigma = noisy.I_sigma.clone() - p = out_dir / f"{name}.mtz" - data.write_mtz(str(p)) - paths[name] = str(p) - - meta = dict(alpha=alpha, lambda_twin=lam, sigma_alpha_sq=sigma_alpha_sq, - d_min=d_min, sigma_mul=sigma_mul, seed=seed, - median_contamination_fraction=frac_contam, **paths) - (out_dir / "truth.json").write_text(json.dumps(meta, indent=2)) - return meta - - -def score(truth_pdb, start_pdb, refined_pdbs, chain, first, last): - """Injected vs recovered displacement over the moved residues.""" - import gemmi - - def positions(path): - st = gemmi.read_structure(str(path)) - st.remove_hydrogens() - out = {} - for ch in st[0]: - if ch.name != chain: - continue - for res in ch: - if first <= res.seqid.num <= last: - for atom in res: - out[(res.seqid.num, atom.name)] = np.array( - [atom.pos.x, atom.pos.y, atom.pos.z] - ) - return out - - truth, start = positions(truth_pdb), positions(start_pdb) - shared = sorted(set(truth) & set(start)) - injected = np.array([np.linalg.norm(truth[k] - start[k]) for k in shared]).mean() - - print(f"\ninjected displacement over the moved residues: {injected:.3f} A " - f"({len(shared)} atoms)") - print(f"{'arm':14s} {'recovered':>10s} {'ratio':>8s} {'err vs truth':>13s}") - print("-" * 50) - rows = [] - for label, path in refined_pdbs: - if not Path(path).exists(): - print(f"{label:14s} (missing)") - continue - got = positions(path) - keys = [k for k in shared if k in got] - rec = np.array([np.linalg.norm(got[k] - start[k]) for k in keys]).mean() - err = np.array([np.linalg.norm(got[k] - truth[k]) for k in keys]).mean() - rows.append((label, rec, rec / injected, err)) - print(f"{label:14s} {rec:10.3f} {rec / injected:8.2f} {err:13.3f}") - print("-" * 50) - print("ratio > 1 means the refinement moved further than the truth") - return rows - - -def score_sweep(root, chain, first, last, start_pdb, n_boot=10000, seed=0): - """Aggregate a seed sweep, paired seed by seed. - - Paired, not pooled: every arm refines the *same* simulated dataset within a seed, so - the seed-to-seed spread of the noise realisation is common to all arms and cancels in - the difference. Comparing two distributions of errors instead would drown a real - effect in variance that is not there. - - Reports the median paired difference with a bootstrap CI, which is the shape that - survives a skewed distribution and a handful of seeds. - """ - import gemmi - - root = Path(root) - seeds = sorted(d for d in root.glob("seed_*") if d.is_dir()) - if not seeds: - print(f"no seed_* directories under {root}") - return [] - - def positions(path): - st = gemmi.read_structure(str(path)) - st.remove_hydrogens() - out = {} - for ch in st[0]: - if ch.name != chain: - continue - for res in ch: - if first <= res.seqid.num <= last: - for atom in res: - out[(res.seqid.num, atom.name)] = np.array( - [atom.pos.x, atom.pos.y, atom.pos.z] - ) - return out - - start = positions(start_pdb) - per_arm = {} - injected = [] - for sd in seeds: - truth_p = sd / "light_truth.pdb" - if not truth_p.exists(): - continue - truth = positions(truth_p) - shared = sorted(set(truth) & set(start)) - injected.append( - np.mean([np.linalg.norm(truth[k] - start[k]) for k in shared]) - ) - for arm_dir in sorted(sd.glob("refine_*")): - hits = sorted(arm_dir.glob("fractions_*_light.pdb")) - if not hits: - continue - got = positions(hits[0]) - keys = [k for k in shared if k in got] - if not keys: - continue - err = np.mean([np.linalg.norm(got[k] - truth[k]) for k in keys]) - rec = np.mean([np.linalg.norm(got[k] - start[k]) for k in keys]) - per_arm.setdefault(arm_dir.name, {})[sd.name] = (err, rec) - - inj = float(np.mean(injected)) - complete = set.intersection(*(set(v) for v in per_arm.values())) if per_arm else set() - complete = sorted(complete) - print(f"\ninjected displacement {inj:.3f} A; " - f"{len(complete)} seeds complete in all {len(per_arm)} arms") - if len(complete) < len(seeds): - print(f" ({len(seeds) - len(complete)} seed(s) dropped: not all arms finished)") - - print(f"\n{'arm':<16s} {'err (A)':>16s} {'recovered/injected':>20s}") - print("-" * 56) - for arm in sorted(per_arm): - e = np.array([per_arm[arm][s][0] for s in complete]) - r = np.array([per_arm[arm][s][1] for s in complete]) / inj - print(f"{arm:<16s} {e.mean():8.4f} +- {e.std(ddof=1):5.4f} " - f"{r.mean():14.3f} +- {r.std(ddof=1):.3f}") - - base = "refine_coh" - if base not in per_arm: - return per_arm - rng = np.random.default_rng(seed) - print(f"\nPaired against {base}, median of per-seed differences " - f"({n_boot} bootstrap resamples)") - print("-" * 72) - print(f"{'arm':<16s} {'median d(err)':>14s} {'95% CI':>22s} {'seeds better':>14s}") - for arm in sorted(per_arm): - if arm == base: - continue - d = np.array([per_arm[arm][s][0] - per_arm[base][s][0] for s in complete]) - boots = np.array([ - np.median(rng.choice(d, size=len(d), replace=True)) for _ in range(n_boot) - ]) - lo, hi = np.percentile(boots, [2.5, 97.5]) - print(f"{arm:<16s} {np.median(d):+14.4f} {f'[{lo:+.4f}, {hi:+.4f}]':>22s} " - f"{f'{(d < 0).sum()}/{len(d)}':>14s}") - print("-" * 72) - print("negative = closer to truth than the coherent refinement") - return per_arm - - -def main(): - ap = argparse.ArgumentParser(description=__doc__) - repo = Path(__file__).resolve().parents[1] - ap.add_argument("--pdb", default=str(repo / "tests/files/pdb/1DAW.pdb")) - ap.add_argument("--reference-mtz", default=str(repo / "tests/files/mtz/1DAW.mtz")) - ap.add_argument("--out", required=True) - ap.add_argument("--chain", default="A") - ap.add_argument("--first", type=int, default=40) - ap.add_argument("--last", type=int, default=52) - ap.add_argument("--shift", type=float, nargs=3, default=[0.35, 0.20, -0.15]) - ap.add_argument("--alpha", type=float, default=0.22) - ap.add_argument("--lambda-twin", type=float, default=0.3) - ap.add_argument("--dmin", type=float, default=2.05) - ap.add_argument("--sigma-mul", type=float, default=0.10) - ap.add_argument("--seed", type=int, default=7) - ap.add_argument("--device", default="cpu") - ap.add_argument("--score", action="store_true", - help="Compare refined models against the truth (after refining).") - ap.add_argument("--score-sweep", action="store_true", - help="Aggregate a seed sweep under --out, paired seed by seed.") - args = ap.parse_args() - - out = Path(args.out) - truth_pdb = out / "light_truth.pdb" - - if args.score_sweep: - score_sweep(out, args.chain, args.first, args.last, args.pdb) - return 0 - - if args.score: - arms = [] - for d in sorted(out.glob("refine_*")): - if not d.is_dir(): - continue - # The output prefix encodes the fraction, which varies across arms. - hits = sorted(d.glob("fractions_*_light.pdb")) - arms.append((d.name, str(hits[0]) if hits else str(d / "missing.pdb"))) - score(truth_pdb, args.pdb, arms, args.chain, args.first, args.last) - return 0 - - out.mkdir(parents=True, exist_ok=True) - print(f"Injecting a displacement into chain {args.chain} " - f"residues {args.first}-{args.last}") - displaced_copy(args.pdb, truth_pdb, args.chain, args.first, args.last, - args.shift, verbose=1) - - print("Simulating two-moment intensities...") - meta = simulate(args.pdb, truth_pdb, args.reference_mtz, out, args.alpha, - args.lambda_twin, args.dmin, args.sigma_mul, args.seed, - torch.device(args.device)) - print(f" alpha={meta['alpha']} lambda={meta['lambda_twin']} " - f"sigma_alpha^2={meta['sigma_alpha_sq']:.4f}") - print(f" median contamination fraction of I: " - f"{meta['median_contamination_fraction']:.3e}") - print(f"\nRefine from the DARK model for both states, e.g.\n") - print(f" torchref.difference-refine -dm {args.pdb} -lm {args.pdb} \\\n" - f" -dsf {meta['dark']} -lsf {meta['light']} \\\n" - f" --fraction {args.alpha} --dmin {args.dmin} " - f"-o {out}/refine_coh --device cpu\n") - print(f"then re-run this with --score --out {out}") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/paper/probe_two_moment_collinearity.py b/paper/probe_two_moment_collinearity.py deleted file mode 100644 index 780fa1dc..00000000 --- a/paper/probe_two_moment_collinearity.py +++ /dev/null @@ -1,222 +0,0 @@ -#!/usr/bin/env python -"""How much of the activation-dispersion signal can the scale model absorb? - -The contamination ``sigma_alpha^2 |dF/dalpha|^2`` is smooth, strictly positive and -concentrated at low resolution -- the same shape a scale or overall-B error takes. If the -scale parameters can reproduce it, a refined ``lambda`` is measuring scale error rather -than activation heterogeneity, and no amount of refinement will tell the two apart. - -This is answerable exactly, by linear algebra, with no refinement at all. Whiten every -quantity by the measurement error, treat the contamination as a template ``t`` and the -scale parameters' derivatives as a design matrix ``X``, and project:: - - P = X (X'X)^-1 X' - leak = ||P t|| / ||t|| fraction of the template the scale model can absorb - vif = ||t|| / ||t - P t|| how much the surviving signal is degraded - -Two designs are compared, because they are the two places scale is fitted: - -* **per-dataset** -- the centered log-scale and quadratic corrections that ``DatasetCollection.scale()`` - fits on the light dataset alone. Free to shape the light data however it likes. -* **shared** -- the ``CollectionScaler`` parameters, which are fitted jointly against dark - and light. One column per parameter spanning *both* datasets, so a light-only template - cannot be matched without spoiling the dark. - -The derivatives are taken numerically from the live scaler rather than from textbook -formulae, so the design matrix is the parameterisation actually in use. - -The template is then split into a smooth resolution envelope and the residual speckle, -and each projected separately: the envelope is what a scale model can absorb, the speckle -is what identifies the dispersion. Which half survives decides how a fitted lambda should -be read. -""" - -import argparse -import json -import sys -from pathlib import Path - -import numpy as np -import torch - - -def build(dark_sf, light_sf, dark_pdb, light_pdb, cif, d_min, fraction, lam, device): - """Rebuild the figure-4 collection exactly as the CLI does.""" - from torchref.cli.collection_difference_refine import ( - setup_dataset_collection, - setup_model_collection, - setup_scaler, - ) - - mc = setup_model_collection( - dark_pdb, light_pdb, [1.0 - fraction, fraction], cif, d_min, device, 0 - ) - dc = setup_dataset_collection(dark_sf, light_sf, d_min, device) - scaler = setup_scaler(dc, mc, device, verbose=0) - mc.set_lambda_twin(lam) - return dc, mc, scaler - - -def light_intensity(dc, mc, scaler): - """Scaled model intensity for the light dataset, shape (n_hkl,).""" - keys = [mc.dark_key, "light"] - rows = [mc.keys().index(k) for k in keys] - comps = dc.component_structure_factors(mc, recalc=True) - w = mc.fractions_matrix()[rows] - return scaler.forward_batched(mc.mix_component_fcalcs(comps, w), w)[1].abs() ** 2 - - -def contamination(dc, mc, scaler): - """sigma_alpha^2 |dF/dalpha|^2 on the light dataset.""" - keys = [mc.dark_key, "light"] - rows = [mc.keys().index(k) for k in keys] - comps = dc.component_structure_factors(mc, recalc=True) - jac = mc.activation_jacobian()[rows] - deriv = scaler.forward_batched(mc.mix_component_fcalcs(comps, jac), jac)[1] - return mc.sigma_alpha_sq * deriv.abs() ** 2 - - -def numeric_columns(params, evaluate, rel_step=1e-3): - """d(model intensity)/d(theta) for every scalar in `params`, by central difference.""" - cols = [] - for p in params: - flat = p.detach().reshape(-1) - for i in range(flat.numel()): - step = rel_step * max(abs(float(flat[i])), 1e-3) - saved = float(flat[i]) - # No autograd: these are numerical derivatives, and retaining a graph per - # evaluation is what pushed this over the memory limit. - with torch.no_grad(): - flat[i] = saved + step - plus = evaluate() - flat[i] = saved - step - minus = evaluate() - flat[i] = saved - col = ((plus - minus) / (2 * step)).cpu().numpy() - del plus, minus - cols.append(col) - return np.asarray(cols).T # (n_hkl, n_param) - - -def leakage(t, X): - """(leak, vif) for template `t` against design `X`, both already whitened.""" - keep = ~np.any(~np.isfinite(X), axis=1) & np.isfinite(t) - t, X = t[keep], X[keep] - # Drop null columns, then least-squares project (lstsq handles rank deficiency). - good = np.linalg.norm(X, axis=0) > 0 - X = X[:, good] - if X.shape[1] == 0: - return 0.0, 1.0 - coef, *_ = np.linalg.lstsq(X, t, rcond=None) - fit = X @ coef - nt = np.linalg.norm(t) - resid = np.linalg.norm(t - fit) - return float(np.linalg.norm(fit) / nt), float(nt / max(resid, 1e-30)) - - -def envelope_and_speckle(t, res, nbin=30): - """Split a template into its smooth resolution envelope and the residual.""" - order = np.argsort(-res) - env = np.zeros_like(t) - for chunk in np.array_split(order, nbin): - env[chunk] = t[chunk].mean() - return env, t - env - - -def main(): - ap = argparse.ArgumentParser(description=__doc__) - fig4 = Path(__file__).resolve().parent / "figure4_difference_refinement" - ap.add_argument("--dark-sf", default=str(fig4 / "data/8QL2-sf.cif")) - ap.add_argument("--light-sf", default=str(fig4 / "data/7YYZ-light.mtz")) - ap.add_argument("--dark-pdb", default=str(fig4 / "data/8QL2_no_altloc.pdb")) - ap.add_argument("--light-pdb", default=str(fig4 / "work_no_altloc.pdb")) - ap.add_argument("--cif", nargs="*", default=[str(fig4 / "data/IBL_grade.cif")]) - ap.add_argument("--dmin", type=float, default=2.2) - ap.add_argument("--fraction", type=float, default=0.22) - ap.add_argument("--lambda-twin", type=float, default=0.2) - ap.add_argument("--device", default="cpu") - ap.add_argument("-o", "--out", default=None) - args = ap.parse_args() - - dev = torch.device(args.device) - print("Rebuilding the collection...", flush=True) - dc, mc, scaler = build( - args.dark_sf, args.light_sf, args.dark_pdb, args.light_pdb, - args.cif, args.dmin, args.fraction, args.lambda_twin, dev, - ) - light = dc["light"] - - with torch.no_grad(): - t_raw = contamination(dc, mc, scaler).cpu().numpy() - _, sig_I = light.get_corrected_intensities() - sig = sig_I.cpu().numpy() - res = light.resolution.cpu().numpy() - mask = light.masks().cpu().numpy().astype(bool) - - ok = mask & np.isfinite(sig) & (sig > 0) & np.isfinite(t_raw) & np.isfinite(res) - print(f"reflections: {ok.sum()} of {len(ok)}") - - # Whiten: everything is measured in units of the error it has to beat. - t = (t_raw / sig)[ok] - - # --- design matrices, differentiated numerically --- - # - # Two different response functions, because the two parameter sets act on opposite - # sides of the residual. The shared scaler shapes the *model* intensity; the - # per-dataset relative corrections shape the *observed* one. Using the model response - # for both would give the dataset parameters an identically zero column and report - # no leak at all. - # - # The component structure factors are computed once: no parameter here moves an - # atom, so recomputing them per derivative is ~50x of pure waste. - print("Differentiating the scale model...", flush=True) - keys = [mc.dark_key, "light"] - rows = [mc.keys().index(k) for k in keys] - with torch.no_grad(): - comps = dc.component_structure_factors(mc, recalc=True) - - def model_response(): - w = mc.fractions_matrix()[rows] - return scaler.forward_batched( - mc.mix_component_fcalcs(comps, w), w - )[1].abs() ** 2 - - def obs_response(): - light._corrected_I_fp = None - return light.get_corrected_intensities()[0] - - shared_params = list(scaler.parameters()) - X_shared = numeric_columns(shared_params, model_response)[ok] / sig[ok, None] - - data_params = [p for p in light.parameters() if p is not None] - X_data = numeric_columns(data_params, obs_response)[ok] / sig[ok, None] - - env, speck = envelope_and_speckle(t, res[ok]) - - rows = [] - for tname, tv in (("full", t), ("envelope", env), ("speckle", speck)): - for xname, X in (("per-dataset", X_data), ("shared", X_shared)): - leak, vif = leakage(tv, X) - rows.append(dict(template=tname, design=xname, n_param=X.shape[1], - leak=leak, vif=vif)) - - print() - print("Fraction of the activation template the scale model can absorb") - print("=" * 66) - print(f"{'template':10s} {'design':13s} {'n_param':>8s} {'leak':>8s} {'VIF':>8s}") - print("-" * 66) - for r in rows: - print(f"{r['template']:10s} {r['design']:13s} {r['n_param']:8d} " - f"{r['leak']:8.3f} {r['vif']:8.2f}") - print("-" * 66) - print(f"envelope carries {np.linalg.norm(env) / np.linalg.norm(t):.3f} of the " - f"template norm, speckle {np.linalg.norm(speck) / np.linalg.norm(t):.3f}") - - if args.out: - Path(args.out).write_text(json.dumps(rows, indent=2)) - print(f"written: {args.out}") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/tests/files/hkl/dark_half1.hkl b/tests/files/hkl/dark_half1.hkl index 8aada321..7312b58c 100644 --- a/tests/files/hkl/dark_half1.hkl +++ b/tests/files/hkl/dark_half1.hkl @@ -3,402 +3,67 @@ Symmetry: 1 h k l I phase sigma(I) nmeas -17 -13 -5 -10.07 - 7.18 2 -17 -13 -4 14.43 - 4.44 4 - -17 -13 -3 -9.91 - 2.98 2 - -17 -13 -2 5.70 - 4.37 4 - -17 -13 1 4.74 - 2.34 3 -17 -13 2 5.88 - 5.79 2 - -17 -12 -6 10.69 - 5.24 2 - -17 -12 -5 15.15 - 2.69 2 - -17 -12 -4 15.51 - 7.56 4 - -17 -12 -3 2.12 - 3.08 10 - -17 -12 -2 6.04 - 3.10 11 - -17 -12 -1 10.29 - 8.15 9 - -17 -12 0 2.66 - 3.31 10 -17 -12 1 0.00 - 2.15 3 - -17 -12 2 18.39 - 4.62 2 -17 -11 -6 8.83 - 5.40 3 - -17 -11 -5 7.96 - 4.29 7 - -17 -11 -4 7.00 - 6.07 17 - -17 -11 -3 2.62 - 1.95 32 - -17 -11 -2 24.27 - 11.91 45 - -17 -11 -1 6.48 - 2.42 43 - -17 -11 0 3.19 - 1.74 13 -17 -11 1 4.22 - 6.52 7 - -17 -11 2 2.07 - 6.35 2 - -17 -11 4 9.21 - 0.30 2 - -17 -10 -5 6.78 - 2.00 16 - -17 -10 -4 4.97 - 1.07 42 - -17 -10 -3 2.91 - 2.63 74 - -17 -10 -2 4.43 - 1.18 105 - -17 -10 -1 32.21 - 23.72 87 - -17 -10 0 2.01 - 0.95 41 -17 -10 1 1.33 - 1.23 21 - -17 -10 2 6.97 - 4.70 2 - -17 -9 -7 0.70 - 3.00 2 - -17 -9 -6 7.46 - 4.00 6 - -17 -9 -5 14.04 - 7.81 46 - -17 -9 -4 7.94 - 2.36 83 - -17 -9 -3 4.45 - 0.76 122 -17 -9 -2 4.71 - 1.77 143 - -17 -9 -1 7.66 - 2.75 144 - -17 -9 0 4.54 - 1.02 88 - -17 -9 1 2.18 - 1.21 51 - -17 -9 2 2.04 - 2.29 11 -17 -9 3 0.00 - 0.00 2 - -17 -8 -9 9.29 - 1.64 2 - -17 -8 -7 1.44 - 2.78 4 - -17 -8 -6 7.10 - 2.70 15 -17 -8 -5 3.66 - 1.10 40 - -17 -8 -4 3.11 - 0.75 95 - -17 -8 -3 1.70 - 0.84 142 - -17 -8 -2 2.28 - 0.60 172 - -17 -8 -1 1.95 - 0.75 132 - -17 -8 0 1.80 - 0.70 102 - -17 -8 1 2.66 - 1.37 61 -17 -8 2 0.33 - 3.93 10 - -17 -8 3 7.18 - 1.19 3 - -17 -8 4 14.27 - 6.28 3 - -17 -7 -6 4.23 - 0.43 5 - -17 -7 -5 3.87 - 1.08 37 - -17 -7 -4 2.36 - 0.81 80 - -17 -7 -3 6.30 - 2.51 120 - -17 -7 -2 6.24 - 3.40 112 -17 -7 -1 13.54 - 7.01 102 - -17 -7 0 8.99 - 5.02 93 - -17 -7 1 1.87 - 1.16 40 - -17 -7 2 7.35 - 3.79 6 -17 -7 3 1.74 - 1.23 2 - -17 -7 4 1.48 - 2.82 3 - -17 -7 5 6.91 - 0.89 2 -17 -6 -9 2.50 - 1.77 2 - -17 -6 -7 15.72 - 6.99 2 - -17 -6 -6 2.93 - 2.33 3 -17 -6 -5 2.08 - 2.37 15 - -17 -6 -4 4.24 - 1.33 46 - -17 -6 -3 2.49 - 0.80 76 - -17 -6 -2 3.86 - 0.83 85 - -17 -6 -1 2.79 - 0.75 63 - -17 -6 0 2.66 - 1.83 48 - -17 -6 1 2.37 - 1.21 30 - -17 -6 2 6.01 - 2.68 8 - -17 -6 3 9.91 - 3.11 2 - -17 -6 4 4.26 - 1.16 2 -17 -5 -5 8.22 - 3.02 4 - -17 -5 -4 4.54 - 2.19 13 - -17 -5 -3 8.71 - 5.75 26 - -17 -5 -2 7.56 - 1.88 26 - -17 -5 -1 4.39 - 1.39 32 - -17 -5 0 4.44 - 1.52 17 - -17 -5 1 6.89 - 2.83 8 -17 -5 2 7.73 - 3.95 3 - -17 -4 -7 11.59 - 2.31 2 -17 -4 -6 5.99 - 1.83 2 - -17 -4 -5 1.26 - 5.44 4 - -17 -4 -4 3.27 - 2.36 5 - -17 -4 -3 -0.33 - 2.34 6 - -17 -4 -2 -1.17 - 3.84 5 - -17 -4 -1 6.30 - 1.97 10 - -17 -4 0 7.19 - 5.90 5 - -17 -4 1 2.61 - 1.56 6 - -17 -4 2 6.19 - 1.32 2 -17 -3 -9 3.62 - 0.50 2 - -17 -3 -7 7.37 - 0.66 2 -17 -3 -3 -2.38 - 0.07 2 - -17 -3 -1 -1.82 - 7.23 3 - -17 -3 0 2.44 - 3.39 4 -17 -2 -6 -3.27 - 2.31 2 - -17 -2 -1 3.73 - 2.64 2 - -16 -17 0 5.73 - 2.44 3 -16 -16 -3 0.04 - 3.11 2 - -16 -16 -2 7.53 - 6.52 4 - -16 -16 -1 3.79 - 3.47 5 - -16 -16 0 -16.21 - 5.92 2 - -16 -15 -5 2.74 - 3.34 2 - -16 -15 -4 2.57 - 3.46 13 - -16 -15 -3 3.62 - 2.01 24 -16 -15 -2 5.66 - 2.46 33 - -16 -15 -1 3.22 - 1.12 32 - -16 -15 0 2.68 - 1.49 16 - -16 -15 1 4.90 - 2.61 6 - -16 -15 3 1.45 - 1.02 2 - -16 -14 -8 -0.24 - 0.17 2 - -16 -14 -7 -6.09 - 5.96 3 - -16 -14 -6 2.93 - 1.73 15 -16 -14 -5 57.87 - 31.51 75 - -16 -14 -4 4.77 - 1.66 139 - -16 -14 -3 3.93 - 0.74 181 - -16 -14 -2 15.65 - 6.21 199 - -16 -14 -1 4.89 - 1.07 225 - -16 -14 0 3.30 - 0.62 176 - -16 -14 1 5.53 - 1.34 155 - -16 -14 2 4.94 - 2.18 56 -16 -14 3 3.08 - 2.09 18 - -16 -14 4 9.74 - 2.34 4 - -16 -13 -9 8.32 - 1.08 2 - -16 -13 -8 2.90 - 2.02 2 - -16 -13 -7 2.63 - 1.14 36 - -16 -13 -6 3.87 - 0.80 139 - -16 -13 -5 2.93 - 0.67 222 - -16 -13 -4 4.75 - 0.69 265 -16 -13 -3 5.95 - 1.49 340 - -16 -13 -2 6.29 - 0.73 322 - -16 -13 -1 12.39 - 2.90 339 - -16 -13 0 8.12 - 1.28 304 - -16 -13 1 6.97 - 2.12 265 - -16 -13 2 3.26 - 0.82 212 - -16 -13 3 1.79 - 0.82 101 - -16 -13 4 2.65 - 1.36 17 -16 -13 5 -6.29 - 5.03 4 - -16 -12 -9 1.81 - 2.24 2 - -16 -12 -8 4.59 - 1.95 48 - -16 -12 -7 7.32 - 1.70 193 - -16 -12 -6 5.20 - 1.36 257 - -16 -12 -5 4.87 - 0.67 368 - -16 -12 -4 8.08 - 1.21 406 -16 -12 -3 32.92 - 6.61 461 - -16 -12 -2 7.34 - 0.87 531 - -16 -12 -1 21.04 - 4.37 497 - -16 -12 0 24.55 - 5.30 409 - -16 -12 1 15.18 - 2.90 376 - -16 -12 2 12.26 - 2.72 326 - -16 -12 3 5.78 - 1.04 256 - -16 -12 4 2.34 - 0.63 119 -16 -12 5 4.64 - 1.97 28 -16 -12 6 4.70 - 10.61 2 - -16 -12 7 1.49 - 0.79 3 - -16 -11 -9 4.21 - 1.95 14 - -16 -11 -8 6.35 - 2.53 150 - -16 -11 -7 4.21 - 0.64 266 - -16 -11 -6 11.72 - 2.30 383 - -16 -11 -5 4.58 - 0.63 477 - -16 -11 -4 27.98 - 4.12 528 - -16 -11 -3 6.26 - 0.59 513 -16 -11 -2 16.39 - 2.52 477 - -16 -11 -1 5.46 - 0.60 530 - -16 -11 0 5.77 - 0.73 527 - -16 -11 1 5.39 - 0.64 421 - -16 -11 2 6.14 - 0.66 392 - -16 -11 3 4.43 - 0.71 321 - -16 -11 4 5.32 - 0.80 228 -16 -11 5 1.65 - 0.81 80 - -16 -11 6 1.42 - 0.91 4 - -16 -10 -11 0.08 - 0.47 2 -16 -10 -10 -0.30 - 2.94 8 - -16 -10 -9 1.00 - 0.87 70 - -16 -10 -8 12.49 - 4.95 223 - -16 -10 -7 7.06 - 0.90 367 - -16 -10 -6 27.49 - 5.44 433 - -16 -10 -5 33.74 - 6.21 473 - -16 -10 -4 6.38 - 0.67 561 -16 -10 -3 25.30 - 4.53 492 - -16 -10 -2 6.34 - 1.71 557 - -16 -10 -1 6.54 - 0.59 567 - -16 -10 0 7.48 - 1.00 534 - -16 -10 1 31.68 - 6.12 511 - -16 -10 2 8.86 - 0.95 500 - -16 -10 3 18.62 - 3.25 365 -16 -10 4 22.28 - 5.35 263 - -16 -10 5 4.10 - 0.75 150 - -16 -10 6 3.34 - 1.77 21 - -16 -10 7 7.48 - 4.88 4 - -16 -10 8 -6.25 - 5.20 2 - -16 -9 -10 -0.73 - 7.97 3 - -16 -9 -9 6.68 - 3.91 136 - -16 -9 -8 9.02 - 2.51 288 - -16 -9 -7 5.89 - 0.68 401 -16 -9 -6 5.07 - 0.60 481 - -16 -9 -5 5.42 - 0.60 501 - -16 -9 -4 5.60 - 0.57 542 - -16 -9 -3 6.28 - 0.76 539 - -16 -9 -2 9.89 - 1.49 593 - -16 -9 -1 5.58 - 0.65 576 - -16 -9 0 12.37 - 2.41 562 -16 -9 1 4.82 - 0.62 478 - -16 -9 2 4.75 - 0.82 529 - -16 -9 3 6.86 - 0.93 425 - -16 -9 4 7.74 - 1.16 307 - -16 -9 5 3.43 - 0.68 199 - -16 -9 6 5.20 - 1.76 31 - -16 -8 -12 -9.19 - 2.70 2 - -16 -8 -10 2.88 - 1.93 10 - -16 -8 -9 2.98 - 0.81 153 -16 -8 -8 27.50 - 7.96 301 - -16 -8 -7 8.60 - 1.53 414 - -16 -8 -6 6.43 - 0.58 516 - -16 -8 -5 5.37 - 1.33 511 - -16 -8 -4 23.66 - 4.29 563 - -16 -8 -3 29.63 - 5.40 580 - -16 -8 -2 24.09 - 3.55 620 -16 -8 -1 50.29 - 6.68 580 - -16 -8 0 4.63 - 0.54 590 - -16 -8 1 40.67 - 6.53 539 - -16 -8 2 11.90 - 2.03 504 - -16 -8 3 4.33 - 0.66 409 - -16 -8 4 8.08 - 1.40 352 - -16 -8 5 13.33 - 3.83 246 - -16 -8 6 7.96 - 5.55 66 -16 -7 -10 2.26 - 2.60 13 - -16 -7 -9 3.35 - 0.77 165 - -16 -7 -8 3.43 - 0.63 309 - -16 -7 -7 14.22 - 2.00 410 - -16 -7 -6 15.12 - 18.60 493 - -16 -7 -5 10.13 - 1.46 536 - -16 -7 -4 17.59 - 3.70 529 -16 -7 -3 6.24 - 1.17 585 - -16 -7 -2 11.53 - 1.69 582 - -16 -7 -1 7.59 - 1.32 544 - -16 -7 0 4.52 - 0.58 539 - -16 -7 1 2.89 - 0.56 534 - -16 -7 2 7.32 - 1.06 503 - -16 -7 3 6.19 - 0.77 446 - -16 -7 4 6.64 - 0.62 345 -16 -7 5 2.81 - 0.56 218 - -16 -7 6 4.32 - 0.87 53 -16 -7 7 1.16 - 3.64 5 - -16 -6 -11 6.37 - 2.40 2 - -16 -6 -10 5.35 - 2.71 10 - -16 -6 -9 3.41 - 0.71 126 - -16 -6 -8 5.44 - 0.94 282 - -16 -6 -7 10.30 - 1.98 366 - -16 -6 -6 42.37 - 27.86 477 -16 -6 -5 6.30 - 0.67 540 - -16 -6 -4 24.61 - 4.06 541 - -16 -6 -3 12.28 - 3.36 577 - -16 -6 -2 6.24 - 0.58 598 - -16 -6 -1 12.54 - 1.63 557 - -16 -6 0 13.07 - 2.59 539 - -16 -6 1 12.44 - 1.77 519 -16 -6 2 8.05 - 1.50 480 - -16 -6 3 32.53 - 7.13 422 - -16 -6 4 4.86 - 0.64 347 - -16 -6 5 4.86 - 1.06 188 - -16 -6 6 12.70 - 6.43 21 - -16 -6 7 -0.97 - 6.87 5 - -16 -5 -11 5.75 - 4.07 2 - -16 -5 -10 2.49 - 0.91 5 - -16 -5 -9 12.74 - 9.60 87 -16 -5 -8 3.85 - 0.59 257 - -16 -5 -7 13.73 - 3.60 374 - -16 -5 -6 5.19 - 0.63 416 - -16 -5 -5 5.52 - 0.58 557 - -16 -5 -4 6.43 - 0.63 522 - -16 -5 -3 5.02 - 0.77 529 - -16 -5 -2 9.96 - 1.20 504 -16 -5 -1 10.63 - 2.22 537 - -16 -5 0 13.53 - 2.63 517 - -16 -5 1 7.05 - 0.66 476 - -16 -5 2 10.86 - 1.70 450 - -16 -5 3 6.41 - 0.79 325 - -16 -5 4 4.33 - 0.65 292 - -16 -5 5 2.09 - 0.68 124 - -16 -5 6 5.90 - 3.45 9 -16 -5 7 5.56 - 4.54 3 - -16 -4 -11 15.62 - 4.13 2 -16 -4 -9 8.48 - 3.82 16 - -16 -4 -8 7.32 - 3.28 155 - -16 -4 -7 3.70 - 0.63 283 - -16 -4 -6 7.38 - 0.85 375 - -16 -4 -5 6.49 - 0.91 409 - -16 -4 -4 22.51 - 5.30 473 - -16 -4 -3 2.94 - 0.49 522 -16 -4 -2 41.02 - 8.50 517 - -16 -4 -1 9.17 - 1.13 546 - -16 -4 0 19.96 - 3.83 517 - -16 -4 1 11.34 - 1.70 425 - -16 -4 2 6.27 - 0.71 362 - -16 -4 3 5.25 - 0.71 291 - -16 -4 4 7.04 - 1.26 181 - -16 -4 5 2.62 - 1.20 44 -16 -4 6 -0.25 - 1.68 3 - -16 -3 -10 4.89 - 4.38 3 -16 -3 -9 -3.72 - 3.45 4 - -16 -3 -8 1.13 - 1.14 44 - -16 -3 -7 7.39 - 3.57 157 - -16 -3 -6 3.91 - 0.64 269 - -16 -3 -5 18.12 - 3.59 340 - -16 -3 -4 6.67 - 0.73 351 - -16 -3 -3 9.10 - 2.37 440 -16 -3 -2 5.56 - 0.66 399 - -16 -3 -1 5.63 - 0.60 443 - -16 -3 0 5.69 - 0.71 393 - -16 -3 1 5.54 - 0.68 389 - -16 -3 2 9.56 - 2.84 289 - -16 -3 3 2.93 - 0.70 180 - -16 -3 4 3.15 - 1.22 61 - -16 -3 5 8.88 - 3.87 4 -16 -3 6 9.05 - 4.62 5 - -16 -2 -9 7.44 - 6.65 3 - -16 -2 -8 3.88 - 3.04 7 - -16 -2 -7 8.86 - 5.62 45 - -16 -2 -6 6.48 - 2.51 127 - -16 -2 -5 2.42 - 0.57 227 - -16 -2 -4 12.81 - 2.51 298 - -16 -2 -3 3.75 - 0.66 306 -16 -2 -2 7.94 - 1.71 332 - -16 -2 -1 6.14 - 1.27 314 - -16 -2 0 6.29 - 1.17 288 - -16 -2 1 1.62 - 0.64 245 - -16 -2 2 9.66 - 4.05 148 - -16 -2 3 17.37 - 11.08 49 - -16 -2 4 1.42 - 3.86 8 - -16 -2 6 -0.64 - 0.18 2 - -16 -1 -8 19.30 - 7.88 3 -16 -1 -7 4.19 - 3.65 4 - -16 -1 -6 0.35 - 1.62 16 - -16 -1 -5 3.49 - 1.05 53 - -16 -1 -4 1.94 - 0.81 103 - -16 -1 -3 4.49 - 0.71 183 - -16 -1 -2 2.51 - 0.60 181 - -16 -1 -1 4.50 - 0.97 172 - -16 -1 0 1.98 - 0.84 91 -16 -1 1 3.88 - 0.86 79 - -16 -1 2 3.01 - 1.81 19 - -16 -1 3 3.88 - 2.74 2 -16 0 -7 1.87 - 1.04 2 - -16 0 -5 2.61 - 2.42 7 - -16 0 -4 2.45 - 1.32 5 - -16 0 -3 60.07 - 50.26 13 - -16 0 -2 7.64 - 2.67 16 -16 0 -1 2.76 - 1.76 15 - -16 0 0 4.36 - 2.23 6 - -16 0 1 4.28 - 1.67 5 - -16 0 2 8.25 - 4.05 2 - -16 1 -2 7.69 - 3.71 3 - -16 1 -1 8.01 - 5.66 2 -16 1 1 10.56 - 8.25 2 - -15 -17 -5 4.57 - 3.90 6 - -15 -17 -4 2.52 - 3.10 6 - -15 -17 -3 2.82 - 2.11 22 - -15 -17 -2 24.19 - 14.43 27 -15 -17 -1 3.23 - 1.49 21 - -15 -17 0 5.54 - 1.59 20 - -15 -17 1 5.68 - 1.28 8 - -15 -17 2 6.59 - 2.11 2 - -15 -17 3 9.74 - 3.07 2 - -15 -16 -8 6.08 - 2.15 3 - -15 -16 -7 5.06 - 4.13 3 -15 -16 -6 17.08 - 12.72 30 - -15 -16 -5 4.26 - 0.84 110 - -15 -16 -4 15.42 - 6.05 193 - -15 -16 -3 3.02 - 0.55 255 - -15 -16 -2 5.65 - 0.63 258 - -15 -16 -1 4.77 - 0.68 256 - -15 -16 0 12.35 - 2.90 239 - -15 -16 1 3.79 - 0.68 184 -15 -16 2 4.03 - 1.06 131 - -15 -16 3 3.77 - 1.07 45 - -15 -16 4 5.10 - 1.89 11 - -15 -15 -9 0.64 - 2.56 3 - -15 -15 -8 5.36 - 1.91 26 - -15 -15 -7 2.01 - 0.69 128 - -15 -15 -6 26.75 - 29.49 267 -15 -15 -5 5.13 - 0.64 344 - -15 -15 -4 14.77 - 3.41 380 - -15 -15 -3 12.57 - 2.14 437 - -15 -15 -2 6.15 - 0.65 427 - -15 -15 -1 5.21 - 0.80 416 - -15 -15 0 22.38 - 3.54 418 - -15 -15 1 7.55 - 2.63 380 - -15 -15 2 13.32 - 2.72 335 -15 -15 3 24.05 - 6.08 256 End of reflections diff --git a/tests/files/hkl/dark_half2.hkl b/tests/files/hkl/dark_half2.hkl index e7fe0d0c..4487f2e2 100644 --- a/tests/files/hkl/dark_half2.hkl +++ b/tests/files/hkl/dark_half2.hkl @@ -2,403 +2,68 @@ CrystFEL reflection list version 2.0 Symmetry: 1 h k l I phase sigma(I) nmeas -17 -14 -2 -4.16 - 4.19 3 - -17 -13 -7 -1.84 - 1.30 2 -17 -13 -5 -1.97 - 3.22 2 - -17 -13 -2 -0.01 - 0.01 4 - -17 -13 -1 1.92 - 2.90 4 -17 -13 0 11.95 - 1.25 2 - -17 -12 -5 12.59 - 3.63 4 - -17 -12 -4 3.79 - 6.12 4 - -17 -12 -3 0.02 - 4.19 6 - -17 -12 -2 0.80 - 2.51 10 - -17 -12 -1 7.20 - 6.85 5 - -17 -12 0 -0.76 - 1.98 7 -17 -12 1 1.93 - 1.35 3 - -17 -12 3 -3.52 - 2.49 2 - -17 -11 -5 1.73 - 3.39 5 - -17 -11 -4 5.29 - 2.11 18 - -17 -11 -3 1.27 - 1.79 29 - -17 -11 -2 41.30 - 18.19 43 - -17 -11 -1 3.22 - 1.42 35 - -17 -11 0 5.20 - 2.50 30 -17 -11 1 11.89 - 4.65 10 - -17 -11 2 -1.15 - 1.00 5 - -17 -11 3 11.42 - 5.98 4 - -17 -10 -5 3.39 - 1.03 11 - -17 -10 -4 3.38 - 1.01 54 - -17 -10 -3 11.95 - 8.12 70 - -17 -10 -2 5.07 - 2.94 85 - -17 -10 -1 7.14 - 4.32 100 - -17 -10 0 2.87 - 0.96 57 -17 -10 1 8.99 - 6.60 22 - -17 -10 2 -0.64 - 1.99 3 -17 -10 3 -4.42 - 3.12 2 - -17 -10 4 2.32 - 1.64 2 - -17 -9 -8 11.31 - 3.11 3 - -17 -9 -7 4.01 - 2.84 2 - -17 -9 -6 4.22 - 1.76 7 - -17 -9 -5 3.59 - 2.00 29 - -17 -9 -4 65.27 - 45.70 91 - -17 -9 -3 2.97 - 0.64 116 -17 -9 -2 11.01 - 3.53 126 - -17 -9 -1 15.86 - 7.64 96 - -17 -9 0 3.26 - 0.92 89 - -17 -9 1 2.87 - 1.06 45 - -17 -9 2 46.35 - 30.36 11 - -17 -9 3 1.78 - 2.46 4 -17 -9 4 -0.50 - 5.55 3 - -17 -8 -7 -5.34 - 3.34 3 - -17 -8 -6 1.45 - 2.24 8 -17 -8 -5 3.27 - 2.19 35 - -17 -8 -4 3.22 - 0.84 97 - -17 -8 -3 3.03 - 0.91 116 - -17 -8 -2 3.71 - 0.65 169 - -17 -8 -1 1.36 - 0.63 121 - -17 -8 0 4.29 - 1.02 94 - -17 -8 1 2.52 - 1.31 47 -17 -8 2 3.62 - 2.17 16 - -17 -8 3 6.34 - 1.98 4 - -17 -8 4 10.37 - 4.83 2 - -17 -7 -8 5.95 - 2.57 3 -17 -7 -7 -2.56 - 3.48 4 - -17 -7 -6 -2.75 - 3.09 4 - -17 -7 -5 1.68 - 1.56 29 - -17 -7 -4 1.58 - 0.74 78 - -17 -7 -3 26.87 - 22.09 109 - -17 -7 -2 3.47 - 1.26 120 -17 -7 -1 2.64 - 1.18 129 - -17 -7 0 3.20 - 1.10 82 - -17 -7 1 1.73 - 1.65 44 - -17 -7 2 10.91 - 4.49 5 - -17 -7 4 7.92 - 3.97 5 - -17 -6 -7 4.01 - 2.84 2 - -17 -6 -6 2.07 - 1.46 2 -17 -6 -5 2.82 - 1.53 14 - -17 -6 -4 1.80 - 1.11 47 - -17 -6 -3 2.76 - 0.73 72 - -17 -6 -2 3.25 - 1.06 81 - -17 -6 -1 2.38 - 0.82 87 - -17 -6 0 4.43 - 1.70 43 - -17 -6 1 3.27 - 2.38 17 - -17 -6 2 2.51 - 1.11 5 - -17 -5 -6 4.16 - 5.19 3 -17 -5 -5 5.02 - 3.43 8 - -17 -5 -4 0.06 - 1.50 14 - -17 -5 -3 7.41 - 3.97 26 - -17 -5 -2 3.51 - 1.42 24 - -17 -5 -1 0.01 - 0.10 31 - -17 -5 0 1.76 - 2.20 11 - -17 -5 1 7.60 - 2.58 10 - -17 -5 3 24.57 - 5.19 2 - -17 -4 -7 5.41 - 8.01 2 -17 -4 -6 -0.72 - 1.87 2 - -17 -4 -5 6.60 - 5.45 4 - -17 -4 -4 5.73 - 1.62 6 - -17 -4 -3 -2.27 - 2.86 8 - -17 -4 -2 7.44 - 2.63 6 - -17 -4 -1 2.94 - 1.48 4 - -17 -4 0 5.18 - 3.50 5 -17 -3 -6 6.03 - 4.26 2 - -17 -3 -4 1.18 - 2.18 3 -17 -3 -3 3.83 - 5.60 4 -17 -3 -2 0.00 - 0.00 2 - -17 -3 0 8.74 - 5.52 3 -17 -3 2 -0.49 - 0.35 2 - -16 -16 -5 8.72 - 1.57 2 - -16 -16 -2 -2.65 - 2.07 2 - -16 -16 -1 -2.26 - 1.60 2 - -16 -16 0 6.78 - 4.80 2 - -16 -15 -9 -9.21 - 6.76 2 - -16 -15 -5 -3.51 - 8.87 4 - -16 -15 -4 4.64 - 5.33 6 - -16 -15 -3 7.51 - 3.47 28 -16 -15 -2 6.92 - 3.93 44 - -16 -15 -1 5.80 - 1.31 31 - -16 -15 0 2.18 - 1.42 19 - -16 -15 1 2.01 - 2.80 10 -16 -15 2 4.75 - 2.25 5 - -16 -14 -8 1.00 - 0.82 3 - -16 -14 -7 14.25 - 6.60 6 - -16 -14 -6 6.83 - 3.04 15 -16 -14 -5 7.10 - 2.17 87 - -16 -14 -4 11.09 - 5.20 152 - -16 -14 -3 3.05 - 0.67 217 - -16 -14 -2 10.99 - 4.49 188 - -16 -14 -1 8.04 - 2.77 217 - -16 -14 0 2.67 - 0.77 177 - -16 -14 1 14.56 - 5.40 133 - -16 -14 2 6.75 - 2.62 55 -16 -14 3 8.33 - 1.74 10 - -16 -14 4 2.73 - 3.25 5 - -16 -14 5 17.54 - 9.05 2 - -16 -13 -8 7.06 - 1.03 3 - -16 -13 -7 1.73 - 1.17 48 - -16 -13 -6 3.20 - 0.65 179 - -16 -13 -5 3.39 - 0.69 224 - -16 -13 -4 4.14 - 0.58 294 -16 -13 -3 6.25 - 0.66 314 - -16 -13 -2 5.50 - 0.67 381 - -16 -13 -1 9.31 - 1.84 317 - -16 -13 0 8.93 - 1.51 330 - -16 -13 1 6.13 - 1.15 286 - -16 -13 2 3.34 - 0.94 213 - -16 -13 3 4.54 - 2.22 106 - -16 -13 4 4.67 - 3.26 15 -16 -13 5 7.51 - 5.24 3 - -16 -13 6 -10.29 - 13.53 2 - -16 -12 -9 12.74 - 3.23 6 - -16 -12 -8 1.95 - 1.31 35 - -16 -12 -7 7.39 - 9.34 184 - -16 -12 -6 5.33 - 1.05 311 - -16 -12 -5 6.71 - 0.86 357 - -16 -12 -4 6.96 - 0.96 412 -16 -12 -3 38.54 - 12.11 437 - -16 -12 -2 7.56 - 2.08 479 - -16 -12 -1 56.57 - 11.37 441 - -16 -12 0 44.62 - 12.15 437 - -16 -12 1 16.11 - 2.83 372 - -16 -12 2 18.48 - 4.38 300 - -16 -12 3 4.40 - 0.84 253 - -16 -12 4 2.53 - 0.59 124 -16 -12 5 13.57 - 6.30 16 - -16 -11 -9 5.50 - 1.64 18 - -16 -11 -8 4.07 - 1.66 120 - -16 -11 -7 3.77 - 0.66 285 - -16 -11 -6 13.59 - 3.34 389 - -16 -11 -5 4.61 - 0.64 456 - -16 -11 -4 36.53 - 10.20 539 - -16 -11 -3 6.01 - 0.69 523 -16 -11 -2 11.11 - 1.90 548 - -16 -11 -1 5.97 - 0.61 557 - -16 -11 0 5.78 - 0.57 522 - -16 -11 1 6.14 - 0.74 441 - -16 -11 2 5.76 - 0.65 389 - -16 -11 3 5.17 - 0.67 333 - -16 -11 4 3.51 - 0.78 240 -16 -11 5 5.14 - 1.45 88 - -16 -11 6 1.44 - 2.59 7 -16 -11 7 -0.60 - 3.65 2 - -16 -10 -9 3.60 - 0.94 76 - -16 -10 -8 25.20 - 8.83 221 - -16 -10 -7 8.22 - 1.02 338 - -16 -10 -6 39.25 - 8.80 405 - -16 -10 -5 41.45 - 9.89 497 - -16 -10 -4 6.85 - 0.86 508 -16 -10 -3 22.98 - 4.54 515 - -16 -10 -2 9.72 - 1.42 537 - -16 -10 -1 5.35 - 0.61 538 - -16 -10 0 8.75 - 1.09 509 - -16 -10 1 23.42 - 7.74 481 - -16 -10 2 9.60 - 1.87 423 - -16 -10 3 24.06 - 7.44 347 -16 -10 4 21.13 - 6.31 285 - -16 -10 5 3.26 - 0.75 156 - -16 -10 6 16.10 - 12.97 28 - -16 -10 7 0.34 - 1.25 2 - -16 -9 -11 0.47 - 4.35 3 - -16 -9 -10 1.08 - 3.12 8 - -16 -9 -9 2.91 - 0.77 128 - -16 -9 -8 10.94 - 2.83 285 - -16 -9 -7 4.03 - 0.70 385 -16 -9 -6 5.19 - 0.71 441 - -16 -9 -5 5.01 - 0.58 576 - -16 -9 -4 4.84 - 0.62 573 - -16 -9 -3 7.80 - 1.04 559 - -16 -9 -2 9.45 - 1.31 562 - -16 -9 -1 6.19 - 0.61 583 - -16 -9 0 13.15 - 2.42 571 -16 -9 1 6.26 - 1.00 461 - -16 -9 2 8.60 - 0.94 486 - -16 -9 3 6.23 - 0.63 406 - -16 -9 4 8.00 - 1.32 331 - -16 -9 5 0.63 - 0.29 178 - -16 -9 6 0.96 - 1.48 34 - -16 -9 7 8.13 - 3.46 3 -16 -9 8 -3.36 - 5.17 3 - -16 -8 -10 31.94 - 30.09 7 - -16 -8 -9 2.66 - 0.57 153 -16 -8 -8 27.27 - 6.53 307 - -16 -8 -7 8.63 - 1.04 413 - -16 -8 -6 7.92 - 0.66 465 - -16 -8 -5 5.71 - 0.65 522 - -16 -8 -4 17.11 - 2.83 562 - -16 -8 -3 26.30 - 3.87 567 - -16 -8 -2 0.00 - 0.00 588 -16 -8 -1 42.50 - 9.09 622 - -16 -8 0 5.19 - 0.62 538 - -16 -8 1 39.87 - 7.70 471 - -16 -8 2 10.58 - 1.48 492 - -16 -8 3 5.38 - 0.64 437 - -16 -8 4 9.85 - 1.91 370 - -16 -8 5 8.17 - 2.95 233 - -16 -8 6 2.14 - 1.15 47 - -16 -8 7 1.06 - 0.75 2 -16 -7 -10 2.94 - 2.70 12 - -16 -7 -9 0.00 - 0.00 158 - -16 -7 -8 3.62 - 0.66 269 - -16 -7 -7 18.20 - 2.64 396 - -16 -7 -6 7.34 - 1.37 491 - -16 -7 -5 6.21 - 1.32 554 - -16 -7 -4 0.00 - 0.00 616 -16 -7 -3 6.80 - 1.06 572 - -16 -7 -2 8.83 - 1.58 532 - -16 -7 -1 4.04 - 0.62 569 - -16 -7 0 5.49 - 0.64 543 - -16 -7 1 6.46 - 0.69 514 - -16 -7 2 10.32 - 1.44 480 - -16 -7 3 0.00 - 0.00 418 - -16 -7 4 5.92 - 0.78 312 -16 -7 5 3.34 - 0.69 225 - -16 -7 6 3.21 - 1.41 59 - -16 -7 8 -16.54 - 8.13 2 - -16 -6 -11 0.27 - 0.98 2 - -16 -6 -10 6.01 - 2.18 16 - -16 -6 -9 3.48 - 1.03 112 - -16 -6 -8 10.82 - 2.71 278 - -16 -6 -7 7.99 - 1.08 384 - -16 -6 -6 59.06 - 15.37 461 -16 -6 -5 5.08 - 0.69 492 - -16 -6 -4 19.00 - 3.14 551 - -16 -6 -3 16.88 - 2.46 596 - -16 -6 -2 5.63 - 0.78 516 - -16 -6 -1 9.52 - 1.57 572 - -16 -6 0 12.81 - 2.26 537 - -16 -6 1 11.10 - 1.51 499 -16 -6 2 8.70 - 1.24 474 - -16 -6 3 61.06 - 15.61 407 - -16 -6 4 4.70 - 0.64 318 - -16 -6 5 8.29 - 2.50 172 - -16 -6 6 6.60 - 2.01 29 - -16 -6 7 3.90 - 4.09 6 - -16 -5 -10 0.93 - 7.95 3 - -16 -5 -9 5.55 - 5.22 68 -16 -5 -8 2.82 - 0.60 241 - -16 -5 -7 11.38 - 2.05 358 - -16 -5 -6 4.89 - 0.63 421 - -16 -5 -5 5.26 - 0.63 518 - -16 -5 -4 6.44 - 0.80 503 - -16 -5 -3 1.33 - 0.43 510 - -16 -5 -2 14.38 - 2.52 514 -16 -5 -1 8.94 - 1.24 547 - -16 -5 0 3.82 - 0.90 528 - -16 -5 1 2.14 - 0.40 482 - -16 -5 2 11.32 - 1.97 498 - -16 -5 3 6.51 - 0.70 358 - -16 -5 4 4.69 - 0.59 264 - -16 -5 5 2.13 - 0.73 105 - -16 -5 6 -3.50 - 2.28 7 -16 -4 -9 45.76 - 38.25 20 - -16 -4 -8 3.95 - 1.21 144 - -16 -4 -7 4.91 - 0.65 260 - -16 -4 -6 7.15 - 0.71 367 - -16 -4 -5 6.18 - 0.66 443 - -16 -4 -4 17.86 - 3.43 478 - -16 -4 -3 6.16 - 0.58 538 -16 -4 -2 26.22 - 10.09 534 - -16 -4 -1 2.09 - 0.68 534 - -16 -4 0 25.32 - 4.02 548 - -16 -4 1 15.96 - 3.13 445 - -16 -4 2 5.16 - 0.67 378 - -16 -4 3 3.06 - 0.68 285 - -16 -4 4 6.15 - 2.18 176 - -16 -4 5 3.34 - 1.60 39 -16 -4 6 -1.56 - 1.27 3 - -16 -3 -8 4.16 - 1.32 41 - -16 -3 -7 16.93 - 8.25 150 - -16 -3 -6 3.21 - 0.56 317 - -16 -3 -5 22.74 - 4.27 335 - -16 -3 -4 5.47 - 0.68 353 - -16 -3 -3 18.02 - 3.50 426 -16 -3 -2 5.97 - 0.63 395 - -16 -3 -1 0.00 - 0.00 427 - -16 -3 0 5.13 - 0.73 374 - -16 -3 1 5.38 - 0.65 351 - -16 -3 2 6.03 - 1.27 323 - -16 -3 3 2.94 - 0.60 187 - -16 -3 4 2.63 - 0.83 77 - -16 -3 5 -1.41 - 4.83 7 -16 -3 6 14.18 - 6.81 3 - -16 -2 -9 0.00 - 0.00 2 - -16 -2 -8 -0.89 - 2.53 6 - -16 -2 -7 4.08 - 2.44 47 - -16 -2 -6 4.18 - 1.48 145 - -16 -2 -5 3.56 - 0.66 230 - -16 -2 -4 28.45 - 9.38 307 - -16 -2 -3 5.16 - 0.67 328 -16 -2 -2 8.23 - 1.19 345 - -16 -2 -1 8.48 - 1.26 294 - -16 -2 0 7.08 - 1.51 323 - -16 -2 1 3.25 - 0.66 217 - -16 -2 2 12.47 - 4.13 146 - -16 -2 3 14.24 - 8.22 41 - -16 -2 4 3.55 - 2.07 4 -16 -2 5 9.74 - 3.30 2 -16 -1 -7 2.71 - 3.22 3 - -16 -1 -6 4.16 - 2.18 19 - -16 -1 -5 3.23 - 1.32 44 - -16 -1 -4 3.21 - 0.62 123 - -16 -1 -3 4.16 - 1.62 160 - -16 -1 -2 2.81 - 0.73 148 - -16 -1 -1 3.14 - 0.87 129 - -16 -1 0 2.50 - 0.78 120 -16 -1 1 5.17 - 0.98 64 - -16 -1 2 -0.86 - 1.74 22 - -16 -1 3 3.71 - 3.55 9 - -16 -1 4 2.60 - 3.20 3 - -16 0 -6 8.99 - 4.19 5 - -16 0 -5 -2.76 - 2.32 4 - -16 0 -4 0.53 - 4.28 4 - -16 0 -3 7.28 - 3.16 15 - -16 0 -2 10.39 - 1.80 21 -16 0 -1 2.02 - 2.98 13 - -16 0 0 1.26 - 1.94 3 - -16 0 1 -0.54 - 0.77 2 - -16 0 2 0.38 - 4.31 3 -16 1 -4 -8.38 - 13.42 2 - -16 1 0 22.33 - 0.93 2 -15 -18 1 2.78 - 1.97 2 - -15 -17 -7 0.33 - 0.23 2 - -15 -17 -6 8.14 - 4.98 2 - -15 -17 -5 8.18 - 3.55 3 - -15 -17 -4 1.51 - 3.26 8 - -15 -17 -3 20.87 - 9.28 13 - -15 -17 -2 2.46 - 1.68 26 -15 -17 -1 6.44 - 1.67 27 - -15 -17 0 0.81 - 1.62 19 - -15 -17 1 5.85 - 3.92 5 - -15 -17 2 11.16 - 1.57 2 - -15 -17 3 5.06 - 6.59 2 -15 -17 4 3.35 - 14.34 2 - -15 -16 -8 -1.43 - 2.96 3 - -15 -16 -7 11.50 - 3.94 5 -15 -16 -6 10.24 - 5.08 44 - -15 -16 -5 9.22 - 3.06 115 - -15 -16 -4 10.51 - 2.75 199 - -15 -16 -3 3.04 - 0.58 248 - -15 -16 -2 2.63 - 0.66 262 - -15 -16 -1 3.92 - 0.62 267 - -15 -16 0 6.70 - 1.44 246 - -15 -16 1 0.68 - 0.38 199 -15 -16 2 3.49 - 1.11 117 - -15 -16 3 4.59 - 1.15 35 - -15 -16 4 -1.93 - 2.62 5 - -15 -15 -9 0.00 - 0.00 2 - -15 -15 -8 4.58 - 1.98 20 - -15 -15 -7 4.13 - 0.81 127 - -15 -15 -6 25.74 - 10.01 252 -15 -15 -5 6.40 - 0.66 312 - -15 -15 -4 16.79 - 4.12 414 - -15 -15 -3 11.57 - 2.34 438 - -15 -15 -2 5.93 - 0.69 472 - -15 -15 -1 5.45 - 0.62 436 - -15 -15 0 18.89 - 3.12 405 - -15 -15 1 5.61 - 1.67 360 - -15 -15 2 9.06 - 2.28 362 -15 -15 3 27.05 - 5.93 272 - -15 -15 4 7.08 - 2.34 123 - -15 -15 5 5.78 - 1.58 31 -15 -15 6 10.83 - 7.82 3 - -15 -14 -10 10.42 - 8.09 3 - -15 -14 -9 12.70 - 6.29 49 -15 -14 -8 5.60 - 1.55 231 End of reflections diff --git a/tests/integration/test_cli_two_moment_mtz.py b/tests/integration/test_cli_two_moment_mtz.py index e8515e9b..bf40c7a2 100644 --- a/tests/integration/test_cli_two_moment_mtz.py +++ b/tests/integration/test_cli_two_moment_mtz.py @@ -1,15 +1,4 @@ -"""The difference-refinement MTZ layout, pinned. - -Nothing else in the suite asserts on these column names, and they go into files that get -deposited, so the layout is pinned here deliberately: this file is expected to move in -lockstep with a change to the writer, and to fail loudly if one happens by accident. - -Three layouts are checked. The default is the map a reader wants and can identify -- -``DELFWT``/``PHDELWT``, the weighted difference on dark phases, plus the extrapolated -map. ``--two-moment`` adds the activation-heterogeneity correction. ``--all-columns`` -adds the alternative constructions of both, which are informative once you know which is -which and misleading before then. -""" +"""CLI output column types, phase conventions and two-moment numerical consistency.""" import json import os @@ -24,45 +13,64 @@ # The default set, with a light model supplied (which the refinement CLI always does). # H/K/L are the index, so they are not among ``.columns``. DEFAULT_COLUMNS = { - "Fo_dark": "SFAmplitude", "SIGFo_dark": "Stddev", - "Fo_light": "SFAmplitude", "SIGFo_light": "Stddev", - "DF": "SFAmplitude", "SIGDF": "Stddev", + "Fo_dark": "SFAmplitude", + "SIGFo_dark": "Stddev", + "Fo_light": "SFAmplitude", + "SIGFo_light": "Stddev", + "DF": "SFAmplitude", + "SIGDF": "Stddev", # The difference map. CCP4/Coot open these by name. - "DELFWT": "SFAmplitude", "PHDELWT": "Phase", + "DELFWT": "SFAmplitude", + "PHDELWT": "Phase", "Fc_dark": "SFAmplitude", # The mixed model, and the extrapolated map to refine against. - "FC": "SFAmplitude", "PHIC": "Phase", - "FEXT": "SFAmplitude", "SIGFEXT": "Stddev", - "FWT": "SFAmplitude", "PHWT": "Phase", + "FC": "SFAmplitude", + "PHIC": "Phase", + "FEXT": "SFAmplitude", + "SIGFEXT": "Stddev", + "FWT": "SFAmplitude", + "PHWT": "Phase", } FLAG_COLUMNS = {"FreeR_flag_dark", "FreeR_flag_light"} TWO_MOMENT_COLUMNS = { "DELFWT_corr": "SFAmplitude", - "Fo_light_corr": "SFAmplitude", "SIGFo_light_corr": "Stddev", - "DF_corr": "SFAmplitude", "SIGDF_corr": "Stddev", + "Fo_light_corr": "SFAmplitude", + "SIGFo_light_corr": "Stddev", + "DF_corr": "SFAmplitude", + "SIGDF_corr": "Stddev", "DDF": "SFAmplitude", } # What ``--all-columns`` adds on top, given a light model. ALL_COLUMNS_EXTRA = { - "2mDFop-DFc": "SFAmplitude", "mDFop-DFc": "SFAmplitude", + "2mDFop-DFc": "SFAmplitude", + "mDFop-DFc": "SFAmplitude", "PHIC_diff": "Phase", - "DFc": "SFAmplitude", "DFc_phased": "SFAmplitude", - "FEXT_PHASED": "SFAmplitude", "SIGFEXT_PHASED": "Stddev", - "2FEXT_PHASED-Fc": "SFAmplitude", "FEXT_PHASED-Fc": "SFAmplitude", + "DFc": "SFAmplitude", + "DFc_phased": "SFAmplitude", + "FEXT_PHASED": "SFAmplitude", + "SIGFEXT_PHASED": "Stddev", + "2FEXT_PHASED-Fc": "SFAmplitude", + "FEXT_PHASED-Fc": "SFAmplitude", "PHFEXT_PHASED": "Phase", - "FEXT_SCALAR": "SFAmplitude", "SIGFEXT_SCALAR": "Stddev", - "2FEXT_SCALAR-Fc": "SFAmplitude", "FEXT_SCALAR-Fc": "SFAmplitude", + "FEXT_SCALAR": "SFAmplitude", + "SIGFEXT_SCALAR": "Stddev", + "2FEXT_SCALAR-Fc": "SFAmplitude", + "FEXT_SCALAR-Fc": "SFAmplitude", "PHFEXT_SCALAR": "Phase", } # And what it adds again once the two-moment model is on. ALL_COLUMNS_TWO_MOMENT_EXTRA = { - "Io_light": "Intensity", "SIGIo_light": "Stddev", - "Ic_light_coh": "Intensity", "Ic_light_2mom": "Intensity", - "IVAR_ALPHA": "Intensity", "W_2MOM": "Weight", - "2mDFop-DFc_corr": "SFAmplitude", "mDFop-DFc_corr": "SFAmplitude", + "Io_light": "Intensity", + "SIGIo_light": "Stddev", + "Ic_light_coh": "Intensity", + "Ic_light_2mom": "Intensity", + "IVAR_ALPHA": "Intensity", + "W_2MOM": "Weight", + "2mDFop-DFc_corr": "SFAmplitude", + "mDFop-DFc_corr": "SFAmplitude", } FRACTION = 0.25 @@ -79,12 +87,7 @@ def cli_script(project_root): @pytest.fixture(scope="module") def intensity_pair(mtz_dir, pdb_dir, tmp_path_factory): - """A dark/light pair carrying I/SIGI, from the only fixture that has them. - - 1DAW is the sole reflection file under ``tests/files`` with intensity columns; 3GR5, - which the other difference-refine CLI test uses, has none and so cannot exercise an - intensity-space path at all. - """ + """A dark/light pair carrying I/SIGI, from the only fixture that has them.""" import torch from torchref import ReflectionData @@ -115,20 +118,38 @@ def _run(cli_script, pair, outdir, *extra): env["PYTHONPATH"] = root + os.pathsep + env.get("PYTHONPATH", "") cmd = [ - sys.executable, str(cli_script), - "-dm", str(pair["pdb"]), "-lm", str(pair["pdb"]), - "-dsf", str(pair["dir"] / "dark.mtz"), - "-lsf", str(pair["dir"] / "light.mtz"), - "--fraction", str(FRACTION), - "--n-cycles", "1", "--n-steps", "1", "--max-iter", "3", - "--dmin", "2.2", "-o", str(outdir), - "--device", "cpu", "--verbose", "0", + sys.executable, + str(cli_script), + "-dm", + str(pair["pdb"]), + "-lm", + str(pair["pdb"]), + "-dsf", + str(pair["dir"] / "dark.mtz"), + "-lsf", + str(pair["dir"] / "light.mtz"), + "--fraction", + str(FRACTION), + "--n-cycles", + "1", + "--n-steps", + "1", + "--max-iter", + "3", + "--dmin", + "2.2", + "-o", + str(outdir), + "--device", + "cpu", + "--verbose", + "0", *extra, ] proc = subprocess.run(cmd, capture_output=True, text=True, timeout=1800, env=env) - assert proc.returncode == 0, ( - f"CLI failed ({proc.returncode})\nstderr tail:\n{proc.stderr[-3000:]}" - ) + assert ( + proc.returncode == 0 + ), f"CLI failed ({proc.returncode})\nstderr tail:\n{proc.stderr[-3000:]}" prefix = f"fractions_{round((1 - FRACTION) * 100)}_{round(FRACTION * 100)}_" return outdir / f"{prefix}difference_data.mtz", outdir / f"{prefix}summary.json" @@ -143,8 +164,12 @@ def baseline_mtz(cli_script, intensity_pair, tmp_path_factory): def two_moment_mtz(cli_script, intensity_pair, tmp_path_factory): outdir = tmp_path_factory.mktemp("two_moment") return _run( - cli_script, intensity_pair, outdir, - "--two-moment", "--lambda-twin", str(LAMBDA_TWIN), + cli_script, + intensity_pair, + outdir, + "--two-moment", + "--lambda-twin", + str(LAMBDA_TWIN), ) @@ -154,8 +179,13 @@ def two_moment_all_mtz(cli_script, intensity_pair, tmp_path_factory): which is exactly what ``--all-columns`` is for.""" outdir = tmp_path_factory.mktemp("two_moment_all") return _run( - cli_script, intensity_pair, outdir, - "--two-moment", "--lambda-twin", str(LAMBDA_TWIN), "--all-columns", + cli_script, + intensity_pair, + outdir, + "--two-moment", + "--lambda-twin", + str(LAMBDA_TWIN), + "--all-columns", ) @@ -165,45 +195,47 @@ def _read(path): return rs.read_mtz(str(path)) -class TestDefaultLayout: - def test_default_columns_are_exactly_the_expected_set(self, baseline_mtz): - mtz, _ = baseline_mtz - assert set(_read(mtz).columns) == set(DEFAULT_COLUMNS) | FLAG_COLUMNS +@pytest.mark.parametrize( + "fixture,extra", + [ + ("baseline_mtz", {}), + ("two_moment_mtz", TWO_MOMENT_COLUMNS), + ( + "two_moment_all_mtz", + {**TWO_MOMENT_COLUMNS, **ALL_COLUMNS_EXTRA, **ALL_COLUMNS_TWO_MOMENT_EXTRA}, + ), + ], +) +def test_column_layout_and_types(request, fixture, extra): + """Each CLI mode writes exactly its documented columns with MTZ types.""" + frame = _read(request.getfixturevalue(fixture)[0]) + expected = {**DEFAULT_COLUMNS, **extra} + assert set(frame.columns) == set(expected) | FLAG_COLUMNS + for name, dtype in expected.items(): + assert frame.dtypes[name].name == dtype + assert all(hasattr(dtype, "mtztype") for dtype in frame.dtypes) - def test_every_default_column_carries_the_right_mtz_type(self, baseline_mtz): - mtz, _ = baseline_mtz - df = _read(mtz) - for name, expected in DEFAULT_COLUMNS.items(): - assert name in df.columns, f"missing column {name}" - assert df.dtypes[name].name == expected, ( - f"{name} written as {df.dtypes[name].name}, expected {expected}" - ) - - def test_no_two_moment_columns_without_the_flag(self, baseline_mtz): - mtz, _ = baseline_mtz - present = set(_read(mtz).columns) & set(TWO_MOMENT_COLUMNS) - assert present == set(), f"unexpected two-moment columns: {sorted(present)}" - - def test_no_gated_columns_without_all_columns(self, baseline_mtz): - mtz, _ = baseline_mtz - present = set(_read(mtz).columns) & set(ALL_COLUMNS_EXTRA) - assert present == set(), f"unexpected gated columns: {sorted(present)}" + +class TestDefaultLayout: def test_the_difference_map_is_the_weighted_difference_on_dark_phases( self, baseline_mtz ): - """``DELFWT`` must be ``(Fo_light - Fo_dark) * w`` with ``w`` the mean-normalised - inverse variance -- the construction ``torchref.validate-ded`` correlates - against. If these two ever diverge, the map in the file stops being the map the - validation reports on, which is how the output drifted from the science before. - """ + """``DELFWT`` must be ``(Fo_light - Fo_dark) * w`` with ``w`` the mean- + normalised inverse variance -- the construction ``torchref.validate-ded`` + correlates against. If these two ever diverge, the map in the file stops being + the map the validation reports on, which is how the output drifted from the + science before.""" import numpy as np df = _read(baseline_mtz[0]) - dfo = (df["Fo_light"].to_numpy().astype(float) - - df["Fo_dark"].to_numpy().astype(float)) - sig = np.sqrt(df["SIGFo_dark"].to_numpy().astype(float) ** 2 - + df["SIGFo_light"].to_numpy().astype(float) ** 2) + dfo = df["Fo_light"].to_numpy().astype(float) - df["Fo_dark"].to_numpy().astype( + float + ) + sig = np.sqrt( + df["SIGFo_dark"].to_numpy().astype(float) ** 2 + + df["SIGFo_light"].to_numpy().astype(float) ** 2 + ) w = 1 / sig**2 w = w / w.mean() @@ -220,25 +252,6 @@ def test_the_difference_map_is_the_weighted_difference_on_dark_phases( class TestTwoMomentLayout: - def test_default_columns_all_survive(self, two_moment_mtz): - mtz, _ = two_moment_mtz - assert set(DEFAULT_COLUMNS).issubset(set(_read(mtz).columns)) - - def test_every_new_column_is_present_with_the_right_mtz_type(self, two_moment_mtz): - mtz, _ = two_moment_mtz - df = _read(mtz) - for name, expected in TWO_MOMENT_COLUMNS.items(): - assert name in df.columns, f"missing column {name}" - actual = df.dtypes[name].name - assert actual == expected, ( - f"{name} written as {actual}, expected {expected}" - ) - - def test_the_column_set_is_exactly_default_plus_the_new_ones(self, two_moment_mtz): - mtz, _ = two_moment_mtz - assert set(_read(mtz).columns) == ( - set(DEFAULT_COLUMNS) | FLAG_COLUMNS | set(TWO_MOMENT_COLUMNS) - ) def test_the_corrected_difference_map_pairs_with_the_same_phases( self, two_moment_mtz @@ -248,8 +261,10 @@ def test_the_corrected_difference_map_pairs_with_the_same_phases( import numpy as np df = _read(two_moment_mtz[0]) - sig = np.sqrt(df["SIGFo_dark"].to_numpy().astype(float) ** 2 - + df["SIGFo_light"].to_numpy().astype(float) ** 2) + sig = np.sqrt( + df["SIGFo_dark"].to_numpy().astype(float) ** 2 + + df["SIGFo_light"].to_numpy().astype(float) ** 2 + ) w = 1 / sig**2 w = w / w.mean() @@ -259,39 +274,6 @@ def test_the_corrected_difference_map_pairs_with_the_same_phases( assert np.abs(got - expected).max() / scale < 1e-5 -class TestAllColumns: - def test_all_columns_is_a_strict_superset(self, two_moment_mtz, two_moment_all_mtz): - default = set(_read(two_moment_mtz[0]).columns) - full = set(_read(two_moment_all_mtz[0]).columns) - assert default < full, "--all-columns must add columns, never remove any" - - def test_the_gated_columns_are_exactly_the_expected_ones(self, two_moment_all_mtz): - df = _read(two_moment_all_mtz[0]) - assert set(df.columns) == ( - set(DEFAULT_COLUMNS) | FLAG_COLUMNS | set(TWO_MOMENT_COLUMNS) - | set(ALL_COLUMNS_EXTRA) | set(ALL_COLUMNS_TWO_MOMENT_EXTRA) - ) - - def test_every_gated_column_carries_the_right_mtz_type(self, two_moment_all_mtz): - df = _read(two_moment_all_mtz[0]) - expected_types = {**ALL_COLUMNS_EXTRA, **ALL_COLUMNS_TWO_MOMENT_EXTRA} - for name, expected in expected_types.items(): - assert name in df.columns, f"missing column {name}" - assert df.dtypes[name].name == expected, ( - f"{name} written as {df.dtypes[name].name}, expected {expected}" - ) - - def test_no_column_escapes_with_a_plain_numpy_dtype(self, two_moment_all_mtz): - """Every layer declares its columns' MTZ types beside the values, and the writer - refuses a column with none. This is the end-to-end version of that check: a - column reaching the file as a bare numpy dtype is the failure the old parallel - name lists invited. - """ - df = _read(two_moment_all_mtz[0]) - bare = [c for c in df.columns if not hasattr(df.dtypes[c], "mtztype")] - assert bare == [], f"columns written without an MTZ dtype: {bare}" - - class TestTwoMomentValuesAreConsistent: def test_ivar_alpha_is_sigma_sq_times_the_squared_difference( self, two_moment_all_mtz @@ -315,18 +297,7 @@ def test_the_two_moment_intensity_exceeds_the_coherent_one_by_the_variance( self, two_moment_all_mtz ): """``Ic_2mom - Ic_coh`` must equal ``IVAR_ALPHA``, to whatever precision float32 - leaves after the cancellation. - - This is a catastrophic-cancellation case, and the tolerance is computed rather - than guessed. The variance term is ~2.6e-6 of the intensity on this fixture, - while float32 resolves ~1.2e-7 of it -- so only about one significant digit of - the difference survives, and any fixed tolerance would either pass vacuously or - fail for reasons that have nothing to do with the code. - - The target itself never forms this difference (it computes - ``|F|**2 + sigma**2 |dF|**2`` directly), so the loss is unaffected; it is - recovering the variance term from the two published columns that is lossy. - """ + leaves after the cancellation.""" import numpy as np df = _read(two_moment_all_mtz[0]) @@ -347,15 +318,7 @@ def test_the_two_moment_intensity_exceeds_the_coherent_one_by_the_variance( assert (two >= coh - 4.0 * floor).all() def test_the_weight_is_the_contamination_ratio(self, two_moment_all_mtz): - """``W_2MOM`` must be ``sigma_I**2 / (sigma_I**2 + IVAR_ALPHA)``. - - Asserted as the formula rather than as a magnitude. On this fixture the weight - never falls below ~0.9998, because the contamination is ~1e-3 of a single - reflection's sigma -- which is the real behaviour of this correction, not a - defect: it is a systematic that adds coherently over the whole dataset while - being invisible on any one reflection. A test demanding visible down-weighting - would be asserting the physics is different from what it is. - """ + """``W_2MOM`` must be ``sigma_I**2 / (sigma_I**2 + IVAR_ALPHA)``.""" import numpy as np df = _read(two_moment_all_mtz[0]) diff --git a/tests/unit/io/test_collection_stack_accessors.py b/tests/unit/io/test_collection_stack_accessors.py index c263a645..c8a3c265 100644 --- a/tests/unit/io/test_collection_stack_accessors.py +++ b/tests/unit/io/test_collection_stack_accessors.py @@ -1,14 +1,4 @@ -"""The batched observation accessors must return the same data the targets fit. - -Two things they could quietly get wrong, both invisible in the shape: - -* returning **raw** ``F``/``I`` instead of the scaled ones, which drops the per-dataset - joint scale parameters that ``DatasetCollection.scale()`` fits; -* masking with the 2-way ``rfree_flags`` instead of the 3-way work/free/validation - subsets, which lets validation reflections into the work set. - -Each is pinned against the per-dataset accessor it is a batched form of. -""" +"""Collection row selection, partitions and error handling.""" import pytest import torch @@ -16,7 +6,7 @@ @pytest.fixture(scope="module") def collection(mtz_dir): - """Two 1DAW datasets with *different* scales, so raw and scaled disagree.""" + """Two 1DAW datasets for selection and partition checks.""" mtz = mtz_dir / "1DAW.mtz" if not mtz.exists(): pytest.skip("1DAW fixture not present") @@ -25,61 +15,15 @@ def collection(mtz_dir): from torchref.io.datasets.collection import DatasetCollection a = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - b = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) if a.I is None: pytest.skip("1DAW loaded without intensities") dc = DatasetCollection(verbose=0, device="cpu") dc.add_dataset("dark", a, set_as_reference=True) - dc.add_dataset("light", b) - dc.scale(nsteps=1) - with torch.no_grad(): - dc.scaler.raw_parameters[1, 0] += 0.8 + dc.add_dataset("light", a) return dc -@pytest.mark.unit -class TestScaledNotRaw: - def test_amplitude_rows_match_the_per_dataset_corrected_accessor(self, collection): - stacked = collection.stack_F_obs() - sigma = collection.stack_F_sigma() - for row, key in enumerate(collection.keys()): - F, sig = collection[key].get_corrected_data() - assert torch.equal(stacked[row], F) - assert torch.equal(sigma[row], sig) - - def test_intensity_rows_match_the_per_dataset_corrected_accessor(self, collection): - stacked = collection.stack_I_obs() - sigma = collection.stack_I_sigma() - for row, key in enumerate(collection.keys()): - I, sig = collection[key].get_corrected_intensities() - assert torch.equal(stacked[row], I) - assert torch.equal(sigma[row], sig) - - def test_the_two_datasets_differ_after_scaling(self, collection): - """Anti-vacuity: with identical scales, raw and corrected agree and neither - assertion above could detect the wrong accessor.""" - stacked = collection.stack_F_obs() - assert not torch.allclose(stacked[0], stacked[1]) - raw = torch.stack([collection[k].F_raw for k in collection.keys()], dim=0) - assert torch.allclose(raw[0], raw[1]), "raw amplitudes should be identical here" - assert not torch.allclose(stacked, raw) - - def test_scaled_intensities_are_the_square_of_the_scaled_amplitude_factor( - self, collection - ): - """Ties the two stacks together, so they cannot drift apart in scale.""" - F = collection.stack_F_obs() - I = collection.stack_I_obs() - raw_F = torch.stack([collection[k].F_raw for k in collection.keys()], dim=0) - raw_I = torch.stack([collection[k].I_raw for k in collection.keys()], dim=0) - - keep = (raw_F.abs() > 1e-6) & (raw_I.abs() > 1e-6) - amp = (F[keep] / raw_F[keep]) ** 2 - inten = I[keep] / raw_I[keep] - assert torch.allclose(inten, amp, rtol=1e-5) - - @pytest.mark.unit class TestThreeWayMasks: @pytest.mark.parametrize("use_set", ["work", "free", "val"]) @@ -146,20 +90,3 @@ def test_centric_flags_are_shared_and_hkl_shaped(self, collection): assert centric is not None assert centric.shape == (len(collection.hkl),) assert centric.dtype == torch.bool - - def test_missing_intensities_name_the_offending_dataset(self, mtz_dir): - mtz = mtz_dir / "3GR5.mtz" - if not mtz.exists(): - pytest.skip("3GR5 fixture not present") - - from torchref import ReflectionData - from torchref.io.datasets.collection import DatasetCollection - - amp_only = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - if amp_only.I is not None: - pytest.skip("3GR5 unexpectedly carries intensities") - - dc = DatasetCollection(verbose=0, device="cpu") - dc.add_dataset("amps", amp_only, set_as_reference=True) - with pytest.raises(ValueError, match="'amps'"): - dc.stack_I_obs() diff --git a/tests/unit/io/test_crystfel_hkl.py b/tests/unit/io/test_crystfel_hkl.py index ba9b5ece..c412c245 100644 --- a/tests/unit/io/test_crystfel_hkl.py +++ b/tests/unit/io/test_crystfel_hkl.py @@ -1,136 +1,49 @@ -"""Reading CrystFEL ``partialator`` reflection lists. - -The format that merged serial data actually arrives in. Two properties matter beyond -"it parses": - -* **negative intensities survive.** A merged weak reflection legitimately comes out below - zero, and that is information -- dropping or clamping it biases the mean upward exactly - where the noise dominates. -* **cell and space group come from the caller**, because the format carries neither. - -The fixtures are 400-reflection excerpts of a real ``partialator`` custom-split pair, so the -two halves cover overlapping-but-different reflections measured independently -- the property -that makes such a pair usable as a null, where any difference between them is noise plus -systematics with no real signal in it. -""" +"""CrystFEL parsing, intensity conversion and alignment on real split-half excerpts.""" import pytest import torch -# Cell of the small-molecule dataset these excerpts come from. The format does not carry -# it, so the caller must supply it; a wrong cell would silently give wrong d-spacings. +from torchref import DatasetCollection, ReflectionData + CELL = [14.97, 18.85, 18.89, 89.4, 84.9, 67.8] -SPACEGROUP = "P 1" @pytest.fixture(scope="module") -def hkl_dir(test_files_dir): - d = test_files_dir / "hkl" - if not d.is_dir(): - pytest.skip("CrystFEL hkl fixtures not present") - return d - - -def _load(path): - from torchref import ReflectionData - - return ReflectionData(device="cpu", verbose=0).load_crystfel_hkl( - str(path), cell=CELL, spacegroup=SPACEGROUP - ) - - -@pytest.mark.unit -class TestReaderBasics: - def test_reads_intensities_and_sigmas(self, hkl_dir): - data = _load(hkl_dir / "dark_half1.hkl") - assert data.I is not None and data.I_sigma is not None - assert len(data.I) == len(data.hkl) - assert torch.isfinite(data.I_sigma).all() - # Real partialator output does contain sigma(I) == 0 -- which is why an - # intensity likelihood has to floor it rather than trust it. - assert (data.I_sigma >= 0).all() - - def test_amplitudes_are_derived_by_french_wilson(self, hkl_dir): - """The format is intensity-native, so F comes from the same path an MTZ with - I/SIGI takes.""" - data = _load(hkl_dir / "dark_half1.hkl") - assert data.F is not None - assert (data.F >= 0).all() - assert data._FrenchWilson is not None - - def test_cell_and_spacegroup_come_from_the_caller(self, hkl_dir): - data = _load(hkl_dir / "dark_half1.hkl") - assert torch.allclose( - data.cell.data[:3].to(torch.float64), - torch.tensor(CELL[:3], dtype=torch.float64), - atol=1e-3, - ) - assert data.spacegroup is not None - - def test_the_trailing_marker_is_not_read_as_a_reflection(self, hkl_dir): - """``partialator`` ends the list with 'End of reflections'.""" - raw = (hkl_dir / "dark_half1.hkl").read_text().splitlines() - assert raw[-1].startswith("End of reflections") - data = _load(hkl_dir / "dark_half1.hkl") - # 3 header lines + N reflections + 1 trailer - assert len(data.hkl) <= len(raw) - 4 - - -@pytest.mark.unit -class TestNegativeIntensitiesSurvive: - def test_the_fixture_contains_negatives(self, hkl_dir): - """Precondition, asserted: without negatives the next test proves nothing.""" - n = 0 - for line in (hkl_dir / "dark_half1.hkl").read_text().splitlines()[3:]: - parts = line.split() - if len(parts) >= 4: - try: - n += float(parts[3]) < 0 - except ValueError: - pass - assert n > 5, f"only {n} negative intensities in the fixture" - - def test_negatives_reach_the_dataset(self, hkl_dir): - data = _load(hkl_dir / "dark_half1.hkl") - assert bool((data.I < 0).any()), ( - "negative intensities were dropped or clamped by the reader" +def halves(test_files_dir): + return [ + ReflectionData(device="cpu", verbose=0).load_crystfel_hkl( + str(test_files_dir / "hkl" / f"dark_half{i}.hkl"), + cell=CELL, + spacegroup="P 1", ) - - -@pytest.mark.unit -class TestSplitHalves: - def test_the_two_halves_are_independent_measurements(self, hkl_dir): - """Different reflection sets and different values -- which is what makes them - usable as a null: any difference between them is noise plus systematics, with no - real signal in it. - - (The excerpts are equal-length slices of the full halves, which are 33613 and - 33523 reflections; the sets still differ, which is the property that matters.) - """ - a = _load(hkl_dir / "dark_half1.hkl") - b = _load(hkl_dir / "dark_half2.hkl") - - set_a = {tuple(row) for row in a.hkl.tolist()} - set_b = {tuple(row) for row in b.hkl.tolist()} - assert set_a != set_b, "the halves cover identical reflections" - assert set_a & set_b, "the halves share no reflections at all" - - def test_the_halves_align_into_a_collection(self, hkl_dir): - """``add_dataset`` must reconcile two different reflection lists onto one grid.""" - from torchref.io.datasets.collection import DatasetCollection - - a = _load(hkl_dir / "dark_half1.hkl") - b = _load(hkl_dir / "dark_half2.hkl") - - original_a, original_b = a.hkl.clone(), b.hkl.clone() - dc = DatasetCollection(verbose=0, device="cpu") - dc.add_dataset("half1", a, set_as_reference=True) - dc.add_dataset("half2", b) - - assert dc.n_datasets == 2 - assert len(dc["half1"].hkl) == len(dc["half2"].hkl) == len(dc.hkl) - assert torch.equal(a.hkl, original_a) - assert torch.equal(b.hkl, original_b) - # Both carry intensities, so an intensity target can run on the pair. - assert dc["half1"].I is not None and dc["half2"].I is not None - assert dc.stack_I_obs().shape == (2, len(dc.hkl)) + for i in (1, 2) + ] + + +@pytest.mark.parametrize("index", [0, 1]) +def test_observations_and_metadata(halves, index, test_files_dir): + """Preserve negative intensities, uncertainties, caller metadata and valid rows.""" + data = halves[index] + rows = (test_files_dir / "hkl" / f"dark_half{index+1}.hkl").read_text().splitlines() + assert rows[-1] == "End of reflections" + assert len(data.hkl) == len(rows) - 4 + assert data.I.shape == data.I_sigma.shape == data.F.shape == (len(data),) + assert torch.isfinite(data.I_sigma).all() and (data.I_sigma >= 0).all() + assert (data.I < 0).any() + assert (data.F >= 0).all() and data._FrenchWilson is not None + torch.testing.assert_close(data.cell.data, data.cell.data.new_tensor(CELL)) + assert data.spacegroup.number == 1 + + +def test_partial_overlap_alignment_preserves_sources(halves): + """Align overlapping split halves on their union without mutating either input.""" + a, b = halves + original = [data.hkl.clone() for data in halves] + sets = [{tuple(row) for row in data.hkl.tolist()} for data in halves] + assert sets[0] != sets[1] and sets[0] & sets[1] + dc = DatasetCollection(verbose=0, device="cpu") + dc.add_dataset("a", a).add_dataset("b", b) + assert len(dc) == len(sets[0] | sets[1]) + assert dc.stack_I_obs().shape == (2, len(dc)) + for data, hkl in zip(halves, original): + assert torch.equal(data.hkl, hkl) diff --git a/tests/unit/io/test_fcalc_add_noise.py b/tests/unit/io/test_fcalc_add_noise.py index 70658e57..cf2002e3 100644 --- a/tests/unit/io/test_fcalc_add_noise.py +++ b/tests/unit/io/test_fcalc_add_noise.py @@ -1,16 +1,4 @@ -"""Simulated intensities must keep their negatives. - -``add_noise`` draws two independent noisy half-datasets and returns their mean. The -amplitude it derives has to be clamped -- an amplitude cannot be negative -- but the -*intensity* must not be, and both the intensity and its sigma have to survive on the -returned dataset. - -Why this is not a detail: clamping ``I_mean`` at zero puts a **positive bias** on exactly -the weak reflections where the noise dominates. That bias is smooth, positive, and largest -where the signal is weakest -- the same signature as a genuine positive perturbation of the -merged intensity. Any study of an effect at the 1e-3 level built on clamped simulated data -would be measuring its own generator. -""" +"""Unbiased noisy intensities, uncertainty propagation and reproducible draws.""" import pytest import torch @@ -18,11 +6,7 @@ @pytest.fixture def fcalc_scene(): - """A small P1 scene with a deliberately wide dynamic range. - - The weak tail is the point: with strong reflections only, noise never pushes an - intensity negative and nothing below can distinguish clamped from unclamped. - """ + """A small P1 scene with a deliberately wide dynamic range.""" from torchref.io.datasets import FcalcDataset dataset = FcalcDataset.from_cell_and_resolution( @@ -64,24 +48,11 @@ def test_the_amplitude_is_clamped_but_the_intensity_is_not(self, fcalc_scene): @staticmethod def _weak(noisy, truth): - """The subset a clamp at zero can touch: reflections within 2 sigma of zero. - - Measuring over the whole list instead would drown the effect -- the strong - reflections contribute nothing to the bias but dominate its standard error, so - the very reflections the clamp distorts are the ones averaged away. - - Note this needs a noise model whose sigma does **not** scale with the intensity. - Under purely multiplicative noise ``sigma = f * I``, so no reflection is ever weak - relative to its own sigma and this subset is empty; the Poisson-like ``sigma_lin`` - term below gives ``sigma ~ sqrt(I)`` and therefore a genuine weak tail. - """ + """The subset a clamp at zero can touch: reflections within 2 sigma of zero.""" return truth < 2.0 * noisy.I_sigma def test_the_intensity_is_unbiased_on_the_weak_reflections(self, fcalc_scene): - """The property the clamp breaks, as a bound on the mean of the weak tail. - - The tolerance is the standard error over that subset, not a percentage. - """ + """The property the clamp breaks, as a bound on the mean of the weak tail.""" noisy = fcalc_scene.add_noise( sigma_lin=200.0, sigma_mul=0.0, seed=5, verbose=False ) @@ -89,6 +60,7 @@ def test_the_intensity_is_unbiased_on_the_weak_reflections(self, fcalc_scene): weak = self._weak(noisy, truth) assert int(weak.sum()) > 50, "too few weak reflections to say anything" + assert (noisy.I[weak] < 0).any() residual = (noisy.I - truth)[weak] sem = float(noisy.I_sigma[weak].pow(2).sum().sqrt() / int(weak.sum())) bias = float(residual.mean()) @@ -97,49 +69,21 @@ def test_the_intensity_is_unbiased_on_the_weak_reflections(self, fcalc_scene): f"({4 * sem:.4g}); the intensities are being clamped or otherwise skewed" ) - def test_clamping_would_be_detectable_on_this_scene(self, fcalc_scene): - """Anti-vacuity: quantify what the defect would have looked like here. - - Without this, the test above could pass simply because the scene has no - reflections weak enough for a clamp to reach. - - Measured on this scene, the clamp bias runs only about one to two times the - statistical error on the same mean, and needs a high noise level to stand clear - of it at all. That is not a reason to tolerate it: the bias is **systematic**, so - it repeats identically across datasets and survives averaging, while the error it - is being compared against shrinks as 1/sqrt(N). It is the accumulation, not the - size on any one dataset, that would corrupt a calibration curve. - """ - noisy = fcalc_scene.add_noise( - sigma_lin=200.0, sigma_mul=0.0, seed=5, verbose=False - ) - truth = fcalc_scene.fcalc_amp**2 - weak = self._weak(noisy, truth) - assert int(weak.sum()) > 50 - - honest = float((noisy.I - truth)[weak].mean()) - clamped = float((noisy.I.clamp(min=0.0) - truth)[weak].mean()) - sem = float(noisy.I_sigma[weak].pow(2).sum().sqrt() / int(weak.sum())) - - assert clamped > honest, "clamping did not raise the mean on this scene" - assert clamped > 4.0 * sem, ( - f"the clamped bias ({clamped:.4g}) would fall inside the noise " - f"({4 * sem:.4g}) here, so this scene cannot demonstrate the defect" - ) - @pytest.mark.unit class TestSigmaAndHalves: - def test_sigma_of_the_mean_is_the_single_draw_sigma_over_root_two(self, fcalc_scene): + def test_sigma_of_the_mean_is_the_single_draw_sigma_over_root_two( + self, fcalc_scene + ): a = fcalc_scene.add_noise(sigma_mul=0.2, seed=11, verbose=False) b = fcalc_scene.add_noise(sigma_mul=0.2, seed=12, verbose=False) # Same model, same noise scale: the reported sigma is a property of the model, # not of the draw, so it must be identical across seeds. assert torch.allclose(a.I_sigma, b.I_sigma) - expected = torch.sqrt( - torch.tensor(0.2) ** 2 * fcalc_scene.fcalc_amp**4 - ) / (2.0**0.5) + expected = torch.sqrt(torch.tensor(0.2) ** 2 * fcalc_scene.fcalc_amp**4) / ( + 2.0**0.5 + ) assert torch.allclose(a.I_sigma, expected, rtol=1e-5) def test_amplitude_sigma_uses_the_true_amplitude(self, fcalc_scene): @@ -191,8 +135,10 @@ def test_sigmas_are_grafted_from_the_reference(self, fcalc_scene, mtz_dir): # Build on the reference's own HKL list, which is what makes grafting 1:1. scene = FcalcDataset( - hkl=ref.hkl.clone(), cell=fcalc_scene.cell, - spacegroup=fcalc_scene.spacegroup, device=torch.device("cpu"), + hkl=ref.hkl.clone(), + cell=fcalc_scene.cell, + spacegroup=fcalc_scene.spacegroup, + device=torch.device("cpu"), ) gen = torch.Generator().manual_seed(2) amp = torch.rand(len(ref.hkl), generator=gen) * 100.0 diff --git a/tests/unit/io/test_intensity_accessors.py b/tests/unit/io/test_intensity_accessors.py deleted file mode 100644 index 1ebf7bca..00000000 --- a/tests/unit/io/test_intensity_accessors.py +++ /dev/null @@ -1,193 +0,0 @@ -"""Scaled amplitude and intensity accessors propagate measurement uncertainties.""" - -import math - -import pytest -import torch - - -@pytest.fixture(scope="module") -def with_intensities(mtz_dir): - """1DAW -- the only fixture carrying both I/SIGI and FP/SIGFP.""" - mtz = mtz_dir / "1DAW.mtz" - if not mtz.exists(): - pytest.skip("1DAW fixture not present") - - from torchref import ReflectionData - - data = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - if data.I is None: - pytest.skip("1DAW loaded without intensities") - from torchref import ScaledDataset - from torchref.scaling import DatasetScaler - - scaler = DatasetScaler({"data": data, "peer": data}) - return ScaledDataset(data, scaler, "data") - - -@pytest.fixture(scope="module") -def without_intensities(mtz_dir): - """3GR5 -- amplitudes only.""" - mtz = mtz_dir / "3GR5.mtz" - if not mtz.exists(): - pytest.skip("3GR5 fixture not present") - - from torchref import ReflectionData - - data = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - if data.I is not None: - pytest.skip("3GR5 unexpectedly carries intensities") - return data - - -def _perturb(data, dlog=0.3): - """Temporarily change scaler-owned overall and anisotropic corrections.""" - from contextlib import contextmanager - - @contextmanager - def changed(): - parameters = data.scaler.raw_parameters - original = parameters.detach().clone() - try: - with torch.no_grad(): - parameters[0, 0] += 2 * dlog - parameters[0, 1:] += parameters.new_tensor( - [0.02, -0.01, 0.016, 0.004, 0, 0] - ) - yield data - finally: - with torch.no_grad(): - parameters.copy_(original) - - return changed() - - -@pytest.mark.unit -class TestSquaredScale: - def test_intensity_factor_is_the_square_of_the_amplitude_factor( - self, with_intensities - ): - """The exact relationship, as a per-reflection ratio. - - Independent of how I and F relate in the file (French-Wilson, not I == F**2), - because it compares each quantity against its own unscaled self. - """ - data = with_intensities - with _perturb(data): - F_scaled, _ = data.get_corrected_data() - I_scaled, _ = data.get_corrected_intensities() - - keep = data.masks().to(torch.bool) & (data.F.abs() > 1e-6) & ( - data.I.abs() > 1e-6 - ) - amp_factor = (F_scaled[keep] / data.F_raw[keep]) ** 2 - int_factor = I_scaled[keep] / data.I_raw[keep] - - rel = ((int_factor - amp_factor).abs() / amp_factor.abs()).max() - assert rel < 1e-5, ( - f"intensity scale factor is not the square of the amplitude one " - f"(max rel error {rel:.2e})" - ) - - def test_sigma_scales_with_the_same_factor_as_the_intensity( - self, with_intensities - ): - data = with_intensities - with _perturb(data): - I_scaled, sig_scaled = data.get_corrected_intensities() - keep = (data.I.abs() > 1e-6) & (data.I_sigma.abs() > 1e-6) - - ratio_I = I_scaled[keep] / data.I_raw[keep] - ratio_s = sig_scaled[keep] / data.I_sigma_raw[keep] - assert torch.allclose(ratio_I, ratio_s, rtol=1e-6) - - def test_the_perturbation_actually_changes_the_intensities( - self, with_intensities - ): - """Anti-vacuity: at log_scale 0 and U 0 every factor above is 1.""" - data = with_intensities - before, _ = data.get_corrected_intensities() - with _perturb(data): - after, _ = data.get_corrected_intensities() - assert not torch.allclose(before, after) - - def test_a_pure_scale_change_squares_into_the_intensities( - self, with_intensities - ): - """A doubling of the amplitude scale must quadruple the intensities.""" - data = with_intensities - base, _ = data.get_corrected_intensities() - original = data.scaler.raw_parameters.detach().clone() - try: - with torch.no_grad(): - data.scaler.raw_parameters[0, 0] += 2 * math.log(2) - doubled, _ = data.get_corrected_intensities() - finally: - with torch.no_grad(): - data.scaler.raw_parameters.copy_(original) - - keep = base.abs() > 1e-6 - ratio = (doubled[keep] / base[keep]) - assert torch.allclose(ratio, torch.full_like(ratio, 4.0), rtol=1e-5) - - -@pytest.mark.unit -class TestSubsetViews: - def test_amplitudes_and_intensities_are_both_corrected(self, with_intensities): - data = with_intensities - with _perturb(data): - work = data.work - assert not torch.allclose(work.F, work.F_raw) - assert not torch.allclose(work.I, work.I_raw), ( - "subset.I returned raw intensities while subset.F was scaled" - ) - assert not torch.allclose(work.sigI, work.sigI_raw) - - def test_raw_views_match_the_parent_tensors(self, with_intensities): - data = with_intensities - work = data.work - idx = work.indices - assert torch.equal(work.I_raw, data.I_raw.index_select(0, idx)) - assert torch.equal(work.sigI_raw, data.I_sigma_raw.index_select(0, idx)) - - def test_subset_intensities_match_the_full_size_scaled_array( - self, with_intensities - ): - data = with_intensities - with _perturb(data): - I_scaled, sig_scaled = data.get_corrected_intensities() - for kind in ("work", "free"): - sub = getattr(data, kind) - idx = sub.indices - assert torch.equal(sub.I, I_scaled.index_select(0, idx)) - assert torch.equal(sub.sigI, sig_scaled.index_select(0, idx)) - - def test_cache_follows_a_scale_change(self, with_intensities): - """Subset access must read the current shared scale parameters.""" - data = with_intensities - first = data.work.I.clone() - original = data.scaler.raw_parameters.detach().clone() - try: - with torch.no_grad(): - data.scaler.raw_parameters[0, 0] += 1.0 - second = data.work.I - assert not torch.allclose(first, second) - finally: - with torch.no_grad(): - data.scaler.raw_parameters.copy_(original) - - -@pytest.mark.unit -class TestNoIntensities: - def test_get_corrected_intensities_raises_with_an_actionable_message( - self, without_intensities - ): - with pytest.raises(ValueError, match="no I/SIGI columns|No intensities"): - without_intensities.get_corrected_intensities() - - def test_subset_views_return_none_rather_than_raising(self, without_intensities): - work = without_intensities.work - assert work.I is None - assert work.sigI is None - assert work.I_raw is None - assert work.sigI_raw is None diff --git a/tests/unit/model/test_batched_component_fcalcs.py b/tests/unit/model/test_batched_component_fcalcs.py index 0447e2d0..6bb7fb2d 100644 --- a/tests/unit/model/test_batched_component_fcalcs.py +++ b/tests/unit/model/test_batched_component_fcalcs.py @@ -1,15 +1,4 @@ -"""The batched component/mixture structure factors must equal the per-timepoint loop. - -``compute_component_fcalcs`` + ``mix_component_fcalcs`` exist to evaluate each shared base -model once instead of once per timepoint. That is only a saving if it produces the same -numbers as the loop it replaces, including the signed-index and Friedel bookkeeping that -:meth:`ReflectionData.structure_factors` owns. - -Both sides are built from a single set of model forwards (one ``recalc=True``, then cache -hits) because repeated structure-factor evaluation is not bit-reproducible: two -``recalc=True`` calls on identical input differ by ~4e-3 on individual ``F_calc``. Comparing -two independent evaluations would measure that noise instead of the contraction. -""" +"""Batched structure factors preserve mixed-model values and Friedel phases.""" import pytest import torch @@ -55,9 +44,9 @@ def test_component_stack_matches_per_model_structure_factors(self, pair): for k, model in enumerate(mc.base_models): reference = data.structure_factors(model, recalc=False) - assert torch.equal(stacked[k], reference), ( - f"component {k} differs from data.structure_factors" - ) + assert torch.equal( + stacked[k], reference + ), f"component {k} differs from data.structure_factors" def test_mixture_matches_the_per_timepoint_forward(self, pair): """The whole point: one contraction standing in for T mixed forwards.""" @@ -70,9 +59,9 @@ def test_mixture_matches_the_per_timepoint_forward(self, pair): for row, key in enumerate(mc.keys()): reference = data.structure_factors(mc[key], recalc=False) - assert torch.allclose(mixed[row], reference, rtol=1e-6, atol=1e-6), ( - f"timepoint {key!r} differs from its own mixed forward" - ) + assert torch.allclose( + mixed[row], reference, rtol=1e-6, atol=1e-6 + ), f"timepoint {key!r} differs from its own mixed forward" def test_compute_all_fcalc_agrees_on_the_signed_index(self, pair): """``compute_all_fcalc`` takes the caller's indices verbatim, so handed the @@ -92,14 +81,7 @@ def test_compute_all_fcalc_agrees_on_the_signed_index(self, pair): @pytest.fixture(scope="module") def flagged_pair(pair): - """The same models against data with **manufactured** Friedel-flagged rows. - - Every reflection file under ``tests/files/`` is already inside the CCP4 ASU, so - ``friedel_flags.any()`` is False on all of them and any assertion about the index - convention is silently vacuous. Negating half the Miller indices forces - canonicalisation to flip them back, reproducing the ~50% flagged fraction real - P1 data carries. Cell and space group are preserved, so the same models apply. - """ + """The same models against data with **manufactured** Friedel-flagged rows.""" from torchref import ReflectionData from torchref.io.datasets.collection import DatasetCollection @@ -122,6 +104,7 @@ def flagged_pair(pair): verbose=0, ) + assert data.friedel_flags.any() and (~data.friedel_flags).any() dc = DatasetCollection(verbose=0, device="cpu") dc.add_dataset("dark", data, set_as_reference=True) return dc, mc @@ -129,50 +112,19 @@ def flagged_pair(pair): @pytest.mark.integration class TestConventionIsNotSkipped: - def test_the_fixture_actually_has_flagged_rows(self, flagged_pair): - """The precondition, asserted rather than assumed.""" - data = flagged_pair[0]["dark"] - assert data.friedel_flags is not None - frac = data.friedel_flags.float().mean().item() - assert 0.2 < frac < 0.8, f"expected a mixed flag population, got {frac:.3f}" - - def test_conjugation_moves_phases_and_leaves_amplitudes(self, flagged_pair): - dc, mc = flagged_pair - data = dc["dark"] - - raw = mc.compute_component_fcalcs(data._hkl_for_sf(), recalc=True) - corrected = data.conjugate_friedel(raw) - - assert torch.allclose(raw.abs(), corrected.abs()) - assert not torch.allclose(torch.angle(raw), torch.angle(corrected)) def test_component_stack_is_conjugated_where_flagged(self, flagged_pair): - """``component_structure_factors`` must apply the conjugation, not skip it. - - Compared against the per-model supported entry point, which is the definition - of the convention. - """ + """``component_structure_factors`` must apply the conjugation, not skip it.""" dc, mc = flagged_pair data = dc["dark"] stacked = dc.component_structure_factors(mc, recalc=True) + naive = mc.compute_component_fcalcs(data.hkl, recalc=False) + assert not torch.allclose(stacked, naive) for k, model in enumerate(mc.base_models): reference = data.structure_factors(model, recalc=False) assert torch.equal(stacked[k], reference), f"component {k} phases differ" - def test_skipping_the_conjugation_would_be_detected(self, flagged_pair): - """Anti-vacuity: the naive call this method exists to replace disagrees.""" - dc, mc = flagged_pair - data = dc["dark"] - - correct = dc.component_structure_factors(mc, recalc=True) - naive = mc.compute_component_fcalcs(data.hkl, recalc=True) - - assert not torch.allclose(correct, naive), ( - "evaluating on the canonical index gives the same answer as the signed " - "index plus conjugation -- this fixture cannot detect a convention bug" - ) - @pytest.mark.integration class TestContraction: diff --git a/tests/unit/model/test_model_collection_fractions.py b/tests/unit/model/test_model_collection_fractions.py index 62aa21c7..f20c4ac5 100644 --- a/tests/unit/model/test_model_collection_fractions.py +++ b/tests/unit/model/test_model_collection_fractions.py @@ -1,14 +1,4 @@ -"""Characterisation of how ``ModelCollection`` stores population fractions. - -Nothing else in the suite asserts anything about fraction storage -- not the softmax -parametrisation, not the sum-to-1 validation, not the freeze flags, not the override path. -These tests pin the observable contract so a change of storage has to reproduce it rather -than merely still run. - -Deliberately fileless: the fractions live on ``_SharedMixedModel`` and depend on the base -models only for ``dtype_float`` and device, so a stub is enough and the whole file runs in -well under a second. -""" +"""Population fractions, shared ownership, constraints and gradients.""" import pytest import torch @@ -16,12 +6,7 @@ class _StubModel(nn.Module): - """Minimal stand-in for ``ModelFT`` for fraction bookkeeping. - - Carries a real parameter so ``.to()`` and ``resolve_device`` behave, exposes the two - attributes ``_SharedMixedModel.__init__`` reads, and returns structure factors that - differ per instance so a weighted sum can be checked against its parts. - """ + """Minimal stand-in for ``ModelFT`` for fraction bookkeeping.""" def __init__(self, seed: int): super().__init__() @@ -39,7 +24,9 @@ def dtype_float(self): def forward(self, hkl, recalc: bool = False): # Distinct per model, and a function of the parameter so a gradient can reach it. n = hkl.shape[0] - base = torch.arange(1, n + 1, dtype=self.anchor.dtype, device=self.anchor.device) + base = torch.arange( + 1, n + 1, dtype=self.anchor.dtype, device=self.anchor.device + ) amp = (base * float(self._seed + 1)) + self.anchor return amp.to(torch.complex64) @@ -83,14 +70,8 @@ def test_requested_fractions_round_trip(self, f): @pytest.mark.unit def test_the_reference_row_is_exactly_e_ref(self, two_model_collection): - """The dark is the alpha = 0 evaluation, so its excited fraction is exactly - zero -- not a clamp floor. - - Under a softmax over per-timepoint logits it could not be: ``log(0)`` forces a - clamp, which left the dark sitting at 1e-6. Anything deriving a bound from the - reference's fraction (``sigma_alpha_sq <= alpha (1 - alpha)``) inherited that - floor instead of a true zero. - """ + """The dark is the alpha = 0 evaluation, so its excited fraction is exactly zero + -- not a clamp floor.""" dark = two_model_collection["dark"] assert dark.fractions[1].item() == 0.0 assert dark.fractions[0].item() == 1.0 @@ -135,8 +116,10 @@ def test_dark_is_frozen_and_timepoints_are_not(self, two_model_collection): mc = two_model_collection # Frozen by default: population refinement is opt-in. assert mc._activation_logit.requires_grad is False - assert mc.fraction_parameters() == [mc._activation_logit, - mc._branching_logits[0]] + assert mc.fraction_parameters() == [ + mc._activation_logit, + mc._branching_logits[0], + ] @pytest.mark.unit def test_freeze_and_unfreeze_flip_the_flag(self, two_model_collection): @@ -147,9 +130,7 @@ def test_freeze_and_unfreeze_flip_the_flag(self, two_model_collection): assert mixed.collection._activation_logit.requires_grad is True @pytest.mark.unit - def test_the_reference_carries_no_population_parameter( - self, two_model_collection - ): + def test_the_reference_carries_no_population_parameter(self, two_model_collection): """The reference is the alpha = 0 evaluation, not a row with pinned logits, so there is nothing of its own to freeze or refine.""" mc = two_model_collection @@ -235,9 +216,9 @@ def test_a_timepoint_owns_only_its_fractions(self, two_model_collection): ``ModuleList``, so a timepoint must not re-register their parameters. """ mixed = two_model_collection["light"] - assert list(mixed.parameters()) == [], ( - "a timepoint view registered a parameter of its own" - ) + assert ( + list(mixed.parameters()) == [] + ), "a timepoint view registered a parameter of its own" @pytest.mark.unit def test_shared_base_parameters_are_counted_once(self, two_model_collection): @@ -287,9 +268,10 @@ def test_the_reference_contributes_no_activation_gradient( mc = two_model_collection mc.unfreeze_all_fractions() mc["dark"](hkl, recalc=True).abs().sum().backward() - assert mc._activation_logit.grad is None or float( - mc._activation_logit.grad.abs().max() - ) == 0.0 + assert ( + mc._activation_logit.grad is None + or float(mc._activation_logit.grad.abs().max()) == 0.0 + ) class TestSharedActivation: @@ -331,19 +313,6 @@ def test_a_conflicting_activation_is_rejected_not_projected(self): with pytest.raises(ValueError, match="set_fraction_override"): mc.add_timepoint("late", [0.5, 0.5]) - @pytest.mark.unit - def test_the_rejection_message_names_both_activations(self): - from torchref.model.model_collection import ModelCollection - - mc = ModelCollection([_StubModel(0), _StubModel(1)], verbose=0) - mc.add_dark() - mc.add_timepoint("early", [0.7, 0.3]) - - with pytest.raises(ValueError) as excinfo: - mc.add_timepoint("late", [0.5, 0.5]) - text = str(excinfo.value) - assert "0.5000" in text and "0.3000" in text and "early" in text - @pytest.mark.unit def test_adding_the_reference_after_a_timepoint_leaves_activation_alone(self): """A pure-reference row carries no activation information.""" @@ -407,9 +376,7 @@ def test_lambda_is_exactly_zero_by_default(self, two_model_collection): assert float(mc.sigma_alpha_sq) == 0.0 @pytest.mark.unit - def test_lambda_is_not_a_live_parameter_until_asked_for( - self, two_model_collection - ): + def test_lambda_is_not_a_live_parameter_until_asked_for(self, two_model_collection): mc = two_model_collection def _present(): @@ -424,9 +391,7 @@ def _present(): @pytest.mark.unit @pytest.mark.parametrize("lam", [0.0, 0.25, 0.5, 1.0]) - def test_the_variance_bound_holds_by_construction( - self, two_model_collection, lam - ): + def test_the_variance_bound_holds_by_construction(self, two_model_collection, lam): mc = two_model_collection mc.set_lambda_twin(lam) alpha = float(mc.alpha_mean) @@ -444,7 +409,9 @@ def test_lambda_one_saturates_the_bound(self, two_model_collection): mc = two_model_collection mc.set_lambda_twin(1.0) alpha = float(mc.alpha_mean) - assert float(mc.sigma_alpha_sq) == pytest.approx(alpha * (1.0 - alpha), rel=1e-5) + assert float(mc.sigma_alpha_sq) == pytest.approx( + alpha * (1.0 - alpha), rel=1e-5 + ) @pytest.mark.unit @pytest.mark.parametrize("bad", [-0.1, 1.1]) diff --git a/tests/unit/refinement/test_collection_target_characterisation.py b/tests/unit/refinement/test_collection_target_characterisation.py index 7f0f3922..99d69b3e 100644 --- a/tests/unit/refinement/test_collection_target_characterisation.py +++ b/tests/unit/refinement/test_collection_target_characterisation.py @@ -1,45 +1,16 @@ -"""Characterisation of the collection X-ray targets' observable contract. - -Written to protect a change of fraction storage and a move to batched -``[T, R]`` accessors. The three regressions worth catching are all silent: - -* reading **raw** ``ReflectionData.F`` instead of the scaled ``get_corrected_data()``, - which drops the inter-dataset scaling; -* masking with the 2-way ``rfree_flags`` instead of the 3-way ``work``/``free``/ - ``validation`` subset, which lets validation reflections back into the loss; -* returning a **mean** where the target returns a **sum**, which reweights the X-ray - term by 1/N against every restraint. - -Each is pinned by a *deterministic invariant* rather than a stored number. -``model.forward`` is run-to-run nondeterministic even inside one process (threaded -reduction order; ~4e-3 absolute on individual ``F_calc``), so "the loss equals 3.6e4" -is a weaker statement than "the loss responds to this input the way only a correct -implementation can". Measured for reference: the summed losses here vary by ~4e-7 -relative across repeated calls, so the literal checks that remain are given a -tolerance three orders of magnitude above that. - -The fixture deliberately makes the two datasets and the two models **differ**. With one -``ReflectionData`` added twice -- as the sigma_A collection fixture does -- the observed -difference is identically zero and a raw-vs-scaled regression is invisible. -""" +"""Collection targets consume live scaled data, selected subsets and summed losses.""" import pytest import torch -# Repeated-call spread of the summed losses, measured on this fixture. The literal -# assertions below sit far above it; tightening past ~1e-6 would flake. +# Allow float32 differences from threaded structure-factor reductions. LOSS_RTOL = 1e-4 @pytest.fixture(scope="module") def collection(pdb_dir, mtz_dir): - """``(dc, mc, scaler)`` for a dark/light pair with a real difference in both - the data and the models. - - The light dataset carries a shared scale view, so ``F_obs_light != F_obs_dark`` - only through the *corrected* accessor -- which is what makes the raw-vs-scaled - invariant below bite. The light model is displaced, so ``ΔF_calc != 0`` too. - """ + """``(dc, mc, scaler)`` for a dark/light pair with a real difference in both the + data and the models.""" pdb = pdb_dir / "1DAW.pdb" mtz = mtz_dir / "1DAW.mtz" @@ -97,12 +68,7 @@ def _targets(dc, mc, scaler): @pytest.mark.integration class TestObservedAmplitudesAreScaled: - """The loss must move when a dataset's own scale moves. - - ``DatasetCollection.scale()`` fits a per-dataset shared corrections that - exists only in ``get_corrected_data()``. A target reading raw ``.F`` is completely - blind to it, so this is a direct test of which accessor is in use. - """ + """The loss must move when a dataset's own scale moves.""" @pytest.mark.parametrize("name", ["difference", "difference_i", "ml"]) def test_loss_responds_to_the_datasets_own_log_scale(self, collection, name): @@ -110,11 +76,10 @@ def test_loss_responds_to_the_datasets_own_log_scale(self, collection, name): target = _targets(dc, mc, scaler)[name] before = target.forward().item() - light = dc["light"] original = dc.scaler.raw_parameters[1, 0].detach().clone() try: with torch.no_grad(): - dc.scaler.raw_parameters[1, 0] += 0.25 # ~28% on amplitudes + dc.scaler.raw_parameters[1, 0] += 0.25 target.maintenance() if hasattr(target, "maintenance") else None after = target.forward().item() finally: @@ -127,59 +92,11 @@ def test_loss_responds_to_the_datasets_own_log_scale(self, collection, name): f"{rel:.2e}; the target is reading raw amplitudes, not the scaled ones" ) - def test_corrected_and_raw_amplitudes_actually_differ(self, collection): - """Anti-vacuity: the invariant above is only meaningful if the two accessors - disagree on this fixture.""" - dc, _, _ = collection - light = dc["light"] - with torch.no_grad(): - dc.scaler.raw_parameters[1, 0] += 0.25 - corrected, _ = light.get_corrected_data() - raw = light.F_raw - differ = not torch.allclose(corrected, raw) - dc.scaler.raw_parameters[1, 0] -= 0.25 - assert differ - @pytest.mark.integration class TestSubsetSelectionIsThreeWay: """Work / free / validation, with validation carved out of both.""" - def test_subsets_are_disjoint_and_cover_the_valid_reflections(self, collection): - dc, mc, scaler = collection - target = _targets(dc, mc, scaler)["difference"] - data = dc["dark"] - - work = data.work.mask - free = data.free.mask - val = data.validation.mask - - assert not (work & free).any() - assert not (work & val).any() - assert not (free & val).any() - assert torch.equal(work | free | val, data.masks().to(torch.bool)) - assert target.use_set == "work" - - def test_carving_a_validation_set_shrinks_the_work_and_free_sets(self, collection): - """A 2-way ``rfree_flags`` implementation cannot see a validation set at all, - so the reported ``n`` would not move. - """ - dc, mc, scaler = collection - target = _targets(dc, mc, scaler)["difference"] - data = dc["dark"] - - n_before = target._n_reflections() - free_before = data.free.n - flags = None if data.validation_flags is None else data.validation_flags.clone() - try: - data.generate_validation_set(val_fraction_of_free=0.5, seed=0) - assert data.validation.n > 0, "no validation reflections were carved" - assert data.free.n < free_before, "free set did not shrink" - assert target._n_reflections() <= n_before - finally: - data.validation_flags = flags - data._subset_fp = None - @pytest.mark.parametrize("use_set", ["work", "free"]) def test_loss_is_restricted_to_the_selected_subset(self, collection, use_set): """Work and free are different sizes here, so a target that ignored @@ -193,30 +110,19 @@ def test_loss_is_restricted_to_the_selected_subset(self, collection, use_set): assert target.use_set == use_set n = target._n_reflections() expected = sum( - (dc[k].work if use_set == "work" else dc[k].free).n - for k in target._keys() + (dc[k].work if use_set == "work" else dc[k].free).n for k in target._keys() ) assert n == expected @pytest.mark.integration class TestLossesAreSummedNotAveraged: - """A summed X-ray term grows with the data; a meaned one does not. - - This is the invariant that catches a 1/N reweight, which is otherwise invisible - -- it looks exactly like a change of X-ray weight. - """ + """A summed X-ray term grows with the data; a meaned one does not.""" - def test_adding_a_dataset_grows_the_absolute_loss(self, collection, pdb_dir, mtz_dir): - """The expected ratio is n_after / n_before, and that is 3/2, not 2. - - The fixture already holds two datasets (dark + light), so adding a third takes - the absolute target from 2 to 3. This test used to expect 2.0 because it ran on - ``CollectionRiceTarget``, which overrode ``_keys()`` to drop the dark reference - and so went from 1 to 2. ``ml`` fits every dataset including the dark. - - A meaned target would stay near 1.0 either way, which is what this is for. - """ + def test_adding_a_dataset_grows_the_absolute_loss( + self, collection, pdb_dir, mtz_dir + ): + """Summed loss grows in proportion to the number of datasets.""" from torchref import ReflectionData from torchref.refinement.targets import CollectionMLTarget @@ -225,7 +131,17 @@ def test_adding_a_dataset_grows_the_absolute_loss(self, collection, pdb_dir, mtz n_before = len(target_before._keys()) one = target_before.forward().item() - extra = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz_dir / "1DAW.mtz")) + extra = ReflectionData(device="cpu", verbose=0).load_mtz( + str(mtz_dir / "1DAW.mtz") + ) + saved = ( + dc._datasets, + list(dc._dataset_order), + dc.hkl, + dc.scaler, + dc.scaling_metrics, + ) + branching = torch.nn.ParameterList(mc._branching_logits) dc.add_dataset("light2", extra) mc.add_timepoint("light2", [0.7, 0.3]) try: @@ -233,17 +149,21 @@ def test_adding_a_dataset_grows_the_absolute_loss(self, collection, pdb_dir, mtz n_after = len(target_after._keys()) two = target_after.forward().item() finally: - dc._datasets.pop("light2") - dc._dataset_order.remove("light2") + ( + dc._datasets, + dc._dataset_order, + dc._common_hkl, + dc.scaler, + dc.scaling_metrics, + ) = saved del mc._timepoints["light2"] mc._order.remove("light2") + mc._branching_rows.pop("light2") + mc._branching_logits = branching assert (n_before, n_after) == (2, 3) ratio = two / one - # Tolerance is loose because the shared Luzzati beta is REFITTED on the pooled - # free reflections of the larger collection, so the per-reflection loss moves a - # little too. That is a property of the target, not slack: the two hypotheses - # this test separates are 1.5 and 1.0, which are far apart. + # Refitting shared beta on pooled free reflections shifts the per-row loss. assert ratio == pytest.approx(n_after / n_before, rel=0.15), ( f"{n_after} datasets gave {ratio:.3f}x the loss of {n_before}; a summed " f"target should scale with the count and a meaned one stay near 1.0" diff --git a/tests/unit/refinement/test_collection_taxonomy.py b/tests/unit/refinement/test_collection_taxonomy.py index cf2a98b5..9465712e 100644 --- a/tests/unit/refinement/test_collection_taxonomy.py +++ b/tests/unit/refinement/test_collection_taxonomy.py @@ -1,9 +1,4 @@ -"""The collection X-ray taxonomy, as contracts. - -Mirrors ``tests/unit/refinement/test_nll_beta.py``'s registry tests for the multi-dataset -table. Same thesis, same invariants: one class per row, the observable declared rather than -passed, and no row that pairs an amplitude distribution with intensities. -""" +"""Collection target registry classes, observables and invalid selections.""" import inspect @@ -29,9 +24,9 @@ def test_each_row_has_its_own_class(): "two_moment": CollectionTwoMomentIntensityTarget, "ml": CollectionMLTarget, } - assert set(expected) == set(COLLECTION_XRAY_TARGETS.names), ( - "table and test disagree on the rows" - ) + assert set(expected) == set( + COLLECTION_XRAY_TARGETS.names + ), "table and test disagree on the rows" for name, cls in expected.items(): assert COLLECTION_XRAY_TARGETS.by_name(name).target_cls is cls, name @@ -44,12 +39,7 @@ def test_each_row_has_its_own_class(): @pytest.mark.unit def test_the_observable_is_declared_not_passed(): - """The spec's claim and the class's own attribute must agree. - - A row advertising intensities while reading amplitudes would be wrong by ``2|F|``, - which is resolution-dependent -- so it reads as a scale or B error rather than as a - bug, and nothing downstream would flag it. - """ + """The spec's claim and the class's own attribute must agree.""" from torchref.refinement.targets.collection import ( COLLECTION_XRAY_TARGETS, CollectionXrayTargetSpec, @@ -61,9 +51,9 @@ def test_the_observable_is_declared_not_passed(): assert spec.observable in ("amplitude", "intensity"), spec.name assert getattr(spec.target_cls, "observable", "amplitude") == spec.observable by_obs.setdefault(spec.observable, []).append(spec.name) - assert "observable" not in inspect.signature( - spec.target_cls.__init__ - ).parameters + assert ( + "observable" not in inspect.signature(spec.target_cls.__init__).parameters + ) assert set(by_obs["intensity"]) == {"difference_i", "two_moment"} assert set(by_obs["amplitude"]) == {"difference", "ml"} @@ -77,43 +67,9 @@ def test_the_observable_is_declared_not_passed(): ) -@pytest.mark.unit -def test_both_difference_observables_are_offered(): - """Neither difference row is privileged. - - Which one is better is a property of a dataset's signal-to-noise: amplitudes keep the - loss in the same space as the output DED coefficients, intensities avoid the - French-Wilson posterior reshaping the weak tail. The abstraction exists so that - carrying both is cheap -- and the intensity row proves it, being nothing but an - ``observable`` declaration over the amplitude one. - """ - from torchref.refinement.targets.collection import COLLECTION_XRAY_TARGETS - from torchref.refinement.targets.collection.xray import ( - CollectionDifferenceIntensityTarget, - CollectionDifferenceTarget, - ) - - assert {"difference", "difference_i"} <= set(COLLECTION_XRAY_TARGETS.names) - assert issubclass(CollectionDifferenceIntensityTarget, CollectionDifferenceTarget) - # The subclass adds no likelihood of its own: only the name and the observable. - own = set(vars(CollectionDifferenceIntensityTarget)) - { - # `__annotations__` is present because `name`/`observable` are annotated - # assignments, not because the class defines behaviour. - "__doc__", "__module__", "__qualname__", "__annotations__", - "name", "observable", - } - assert not own, f"the intensity difference row grew a body: {sorted(own)}" - - @pytest.mark.unit def test_there_is_no_intensity_rice_row(): - """Rice is amplitude-only by nature, so the axis is not square. - - Rice and the folded normal are distributions *of an amplitude*; the intensity analogue - is the exponential / chi-square_1 Wilson distribution, a different primitive rather - than a different variance. A row pairing the ML class with intensities would be a - modelling error, not a new feature. - """ + """Rice is amplitude-only by nature, so the axis is not square.""" from torchref.refinement.targets.collection import COLLECTION_XRAY_TARGETS from torchref.refinement.targets.collection.xray import CollectionMLTarget @@ -128,25 +84,3 @@ def test_unknown_rows_fail_closed(): with pytest.raises(ValueError, match="Unknown collection X-ray target"): COLLECTION_XRAY_TARGETS.by_name("no_such_row") - - -@pytest.mark.unit -def test_every_row_goes_through_the_seam(): - """No row may hand-write a ``forward``. - - That is what this refactor bought: 265 lines of per-row forwards, each re-deriving - stacking, masking, the sigma floor and the summing, collapsed into one. A row that - reintroduces its own ``forward`` also reintroduces the possibility of it disagreeing - with ``residuals``, which nothing else would notice. - """ - from torchref.refinement.targets.collection import COLLECTION_XRAY_TARGETS - from torchref.refinement.targets.collection.base import CollectionXrayTarget - - for spec in COLLECTION_XRAY_TARGETS.specs: - cls = spec.target_cls - assert cls.forward is CollectionXrayTarget.forward, ( - f"{spec.name} overrides forward; the likelihood belongs in _per_refl" - ) - assert cls._per_refl is not CollectionXrayTarget._per_refl, ( - f"{spec.name} has no _per_refl of its own" - ) diff --git a/tests/unit/refinement/test_intensity_observable.py b/tests/unit/refinement/test_intensity_observable.py index a3408aaf..ebff21b3 100644 --- a/tests/unit/refinement/test_intensity_observable.py +++ b/tests/unit/refinement/test_intensity_observable.py @@ -1,13 +1,4 @@ -"""The observable axis on the single-dataset x-ray targets. - -The claim under test: the observable is *one* `get_data` override, and everything else -- -the likelihood, the subsets, the masks, the R-factor -- is inherited unchanged. So the tests -here are mostly about what must NOT differ. - -The one thing that genuinely must differ is the variance: an amplitude sigma applied to an -intensity residual is wrong by ``2|F|``, which is resolution-dependent and therefore -presents as a scale or B error rather than as a bug. That is the error this file is for. -""" +"""Intensity targets read measured intensities and propagate their uncertainties.""" import math @@ -24,19 +15,9 @@ ) -# ===================================================================== -# The refactored primitives (step A) -- fileless, no fixture needed -# ===================================================================== - - @pytest.mark.unit def test_the_shared_gaussian_reproduces_the_amplitude_one_bitwise(): - """``nll_per_refl`` is ``gaussian_per_refl`` on ``|F_calc|``, exactly. - - Bitwise, not ``allclose``: the amplitude row has a fused Triton counterpart pinned to - it by ``tests/integration/test_triton_vs_eager_targets.py``, so any drift here shows up - there as a mysterious kernel disagreement rather than as this refactor. - """ + """``nll_per_refl`` is ``gaussian_per_refl`` on ``|F_calc|``, exactly.""" for dtype in (torch.float32, torch.float64): torch.manual_seed(3) F_obs = torch.rand(5000, dtype=dtype) * 100 @@ -50,16 +31,8 @@ def test_the_shared_gaussian_reproduces_the_amplitude_one_bitwise(): @pytest.mark.unit def test_the_absolute_variance_floor_is_opt_out_and_matters(): - """``VAR_FLOOR`` is a distortion, not a safeguard, once the builder has floored sigma. - - It is an *absolute* floor on a variance, so whether it engages depends on the units the - data happens to be in. On sigmas around 1e-5 it rescales the objective by a factor of - several -- which is why the intensity rows pass ``var_floor=0.0`` and rely on - :func:`floor_sigma_obs`' data-dependent floor instead. - - This is pinned because the two paths silently disagreed when the Gaussian was first - shared: the amplitude copy had the clamp and the intensity copy did not. - """ + """``VAR_FLOOR`` is a distortion, not a safeguard, once the builder has floored + sigma.""" sigma = torch.full((256,), 1e-3, dtype=torch.float64) var = intensity_var_from_sigma_obs(sigma) # 1e-6, comfortably above VAR_FLOOR obs = torch.zeros(256, dtype=torch.float64) @@ -76,56 +49,37 @@ def test_the_absolute_variance_floor_is_opt_out_and_matters(): clamped = gaussian_per_refl(obs, model, var_tiny, var_floor=VAR_FLOOR) assert not torch.allclose(free, clamped) # `free` is the honest one: it uses the variance the builder actually produced. - expected = 0.5 * (1e-4) ** 2 / 1e-12 + 0.5 * math.log(1e-12) + 0.5 * math.log(2 * math.pi) + expected = ( + 0.5 * (1e-4) ** 2 / 1e-12 + 0.5 * math.log(1e-12) + 0.5 * math.log(2 * math.pi) + ) assert free[0].item() == pytest.approx(expected, rel=1e-12) @pytest.mark.unit def test_the_intensity_sigma_floor_respects_the_fitted_subset(): - """``mask`` restricts the median, because unfitted rows carry filler. - - A collection member reindexed onto a common reflection list has filler sigmas on the - rows it does not own. Taking the median over those moves the floor for every real - reflection, so the mask is not a convenience. - """ + """``mask`` restricts the median, because unfitted rows carry filler.""" # The first 20 fitted rows are BELOW the fitted median's floor, so the floor is what # they come back as -- which is the only way to observe which median was used. - sigma = torch.cat([ - torch.full((20,), 0.01), # fitted, and below floor either way - torch.full((80,), 10.0), # fitted, sets the fitted median - torch.full((900,), 1e6), # NOT fitted: filler - ]) - mask = torch.cat([torch.ones(100, dtype=torch.bool), torch.zeros(900, dtype=torch.bool)]) + sigma = torch.cat( + [ + torch.full((20,), 0.01), # fitted, and below floor either way + torch.full((80,), 10.0), # fitted, sets the fitted median + torch.full((900,), 1e6), # NOT fitted: filler + ] + ) + mask = torch.cat( + [torch.ones(100, dtype=torch.bool), torch.zeros(900, dtype=torch.bool)] + ) masked = floor_sigma_obs(sigma, mask, abs_floor=1e-12) unmasked = floor_sigma_obs(sigma, None, abs_floor=1e-12) - assert masked[:20].min().item() == pytest.approx(1.0) # floor = 10 * 0.1 - assert unmasked[:20].min().item() == pytest.approx(1e5) # floor = 1e6 * 0.1, swamped + assert masked[:20].min().item() == pytest.approx(1.0) # floor = 10 * 0.1 + assert unmasked[:20].min().item() == pytest.approx( + 1e5 + ) # floor = 1e6 * 0.1, swamped # An explicit floor overrides the median entirely -- the set-independent path. - assert floor_sigma_obs(sigma, mask, floor=0.5)[:20].min().item() == pytest.approx(0.5) - - -@pytest.mark.unit -def test_confusing_the_two_variance_builders_is_wrong_by_a_factor(): - """The amplitude and intensity builders are not interchangeable. - - Both square a floored sigma, so they *look* alike; what differs is which sigma. Passing - ``sigma(F)`` to an intensity residual (or the reverse) is wrong by ``(2|F|)**2``, which - varies with resolution -- so it does not present as an obviously wrong number, it - presents as a scale or B error. Hence a test rather than a comment. - """ - sig_F = torch.rand(1000, dtype=torch.float64) * 2 + 0.5 - F = torch.rand(1000, dtype=torch.float64) * 100 + 10 - sig_I = 2 * F * sig_F # exact first-order propagation, I = F**2 - var_wrong = amplitude_var_from_sigma_obs(sig_F) - var_right = intensity_var_from_sigma_obs(sig_I) - ratio = (var_right / var_wrong).sqrt() - # Spans a wide range: a single global weight cannot absorb it. - assert ratio.max() / ratio.min() > 5 - - -# ===================================================================== -# The row, on real data (step B) -# ===================================================================== + assert floor_sigma_obs(sigma, mask, floor=0.5)[:20].min().item() == pytest.approx( + 0.5 + ) @pytest.fixture(scope="module") @@ -172,13 +126,7 @@ def test_the_row_is_selectable_and_reads_intensities(refinement): @pytest.mark.integration def test_the_intensity_model_is_the_squared_scaled_amplitude(refinement): - """``get_I_calc_scaled`` squares the SCALED amplitude, not the raw one. - - Both the overall scale and the anisotropy factor therefore enter squared, matching - ``ReflectionData.get_corrected_intensities`` on the observation side. Squaring first and - scaling afterwards with the amplitude factors would be wrong by that factor, which is - resolution-dependent. - """ + """``get_I_calc_scaled`` squares the SCALED amplitude, not the raw one.""" t = _t(refinement, "nll_i") with torch.no_grad(): amp = t.get_F_calc_scaled(recalc=False) @@ -188,14 +136,7 @@ def test_the_intensity_model_is_the_squared_scaled_amplitude(refinement): @pytest.mark.integration def test_rfactors_stay_on_amplitudes(refinement): - """An intensity row reports the SAME R-factors as an amplitude row. - - ``_scaled_F_calc_full`` is deliberately not overridden: for a ``|F_calc|**2`` model its - correct value is ``sqrt(I_calc) == |F_calc|``, which is what the base returns. R-factors - therefore remain comparable across the whole table regardless of which observable drove - the loss -- and a future row whose intensity model is not a squared amplitude (the - two-moment model) has to override it, or this test is what will catch it. - """ + """An intensity row reports the SAME R-factors as an amplitude row.""" r_i = _t(refinement, "nll_i").get_rfactor() r_a = _t(refinement, "nll").get_rfactor() assert r_i == pytest.approx(r_a, abs=1e-9) @@ -208,7 +149,8 @@ def test_the_loss_is_finite_differentiable_and_summed(refinement): assert torch.isfinite(loss) and loss.ndim == 0 loss.backward() grads = [ - p.grad for p in refinement.model.parameters() + p.grad + for p in refinement.model.parameters() if p.requires_grad and p.grad is not None ] assert grads, "no gradient reached the model" @@ -218,18 +160,10 @@ def test_the_loss_is_finite_differentiable_and_summed(refinement): @pytest.mark.integration @pytest.mark.parametrize("use_set", ["work", "free"]) -def test_a_reflections_residual_does_not_depend_on_the_arrays_length(refinement, use_set): - """``residuals()`` restricted to a subset must equal ``forward()`` on that subset. - - Pinned separately from ``test_xray_residuals.py`` because the failure mode is specific - to this row: the sigma floor is a *median*, so deriving it from whatever array a call - receives makes every per-reflection value depend on the whole array. ``forward`` sees - the subset and ``residuals`` sees everything, so the two disagreed by 0.09% on the work - set and 1.8% on the free set until the floor was pinned to the target's own subset. - - The amplitude rows share the mechanism and get away with it because sigma(F) is narrow - enough that the clamp barely engages; sigma(I) spans orders of magnitude. - """ +def test_a_reflections_residual_does_not_depend_on_the_arrays_length( + refinement, use_set +): + """``residuals()`` restricted to a subset must equal ``forward()`` on that subset.""" t = _t(refinement, "nll_i", use_set=use_set) sub = t._subset() with torch.no_grad(): @@ -241,7 +175,8 @@ def test_a_reflections_residual_does_not_depend_on_the_arrays_length(refinement, @pytest.mark.integration def test_missing_intensities_raise_at_construction_not_at_forward(refinement): """LossState probes ``forward()`` at registration, so a missing column has to be - caught in ``__init__`` or it surfaces from deep inside setup with no mention of why.""" + caught in ``__init__`` or it surfaces from deep inside setup with no mention of why. + """ import copy data = copy.copy(refinement.reflection_data) diff --git a/tests/unit/refinement/test_two_moment_intensity.py b/tests/unit/refinement/test_two_moment_intensity.py index 21111fea..748ca7f1 100644 --- a/tests/unit/refinement/test_two_moment_intensity.py +++ b/tests/unit/refinement/test_two_moment_intensity.py @@ -1,123 +1,9 @@ -"""The two-moment intensity target: the identity, the coherent limit, and the plumbing. - -The central claim is an *identity*, not an approximation. For any finite set of -per-crystal activations, with the branching conserved, - - mean_c |F_D + a_c dF|^2 == |F_D + abar dF|^2 + var(a) |dF|^2 - -with ``abar`` and ``var`` the **population** moments (1/M divisor). The first test builds -the left-hand side by an explicit loop over crystals -- no two-moment expression anywhere -on the generator side -- so it tests the physics rather than restating the implementation. - -The trap it pins: a ``1/(M-1)`` divisor makes the identity fail by O(1/M). At M=64 that is -1.6% -- small enough to slip past a loose tolerance and far larger than the effect the -target exists to measure. -""" +"""Two-moment target limits, gradients, invalid observations and weight calibration.""" import pytest import torch -def _sample_moments(alpha: torch.Tensor): - """Population mean and variance (1/M divisor, not 1/(M-1)).""" - return alpha.mean(), alpha.var(unbiased=False) - - -def _brute_force_mean_intensity(F_D, dF, alpha): - """mean_c |F_D + a_c dF|^2, by explicit loop. No two-moment expression.""" - total = torch.zeros(F_D.shape, dtype=F_D.real.dtype) - for a in alpha: - total = total + (F_D + a * dF).abs() ** 2 - return total / len(alpha) - - -@pytest.mark.unit -class TestTheMomentIdentity: - @pytest.mark.parametrize("m", [2, 7, 64]) - @pytest.mark.parametrize( - "dtype,tol", [(torch.float64, 1e-13), (torch.float32, 1e-5)] - ) - def test_identity_holds_for_any_finite_activation_set(self, m, dtype, tol): - gen = torch.Generator().manual_seed(11) - n = 32 - cdtype = torch.complex128 if dtype is torch.float64 else torch.complex64 - - F_D = torch.randn(n, generator=gen, dtype=dtype).to(cdtype) + 1j * torch.randn( - n, generator=gen, dtype=dtype - ).to(cdtype) - dF = torch.randn(n, generator=gen, dtype=dtype).to(cdtype) + 1j * torch.randn( - n, generator=gen, dtype=dtype - ).to(cdtype) - alpha = torch.rand(m, generator=gen, dtype=dtype) - - brute = _brute_force_mean_intensity(F_D, dF, alpha) - abar, var = _sample_moments(alpha) - two_moment = (F_D + abar * dF).abs() ** 2 + var * dF.abs() ** 2 - - rel = ((brute - two_moment).abs() / brute.abs().clamp(min=1e-30)).max() - assert rel < tol, f"identity failed at {rel:.2e} (M={m}, {dtype})" - - def test_the_unbiased_variance_divisor_breaks_it(self): - """Anti-vacuity for the divisor: the wrong one fails, and by how much.""" - gen = torch.Generator().manual_seed(3) - n, m = 16, 64 - F_D = torch.randn(n, generator=gen, dtype=torch.float64).to(torch.complex128) - dF = torch.randn(n, generator=gen, dtype=torch.float64).to(torch.complex128) - alpha = torch.rand(m, generator=gen, dtype=torch.float64) - - brute = _brute_force_mean_intensity(F_D, dF, alpha) - abar = alpha.mean() - wrong = (F_D + abar * dF).abs() ** 2 + alpha.var(unbiased=True) * dF.abs() ** 2 - - rel = ((brute - wrong).abs() / brute.abs()).max() - assert rel > 1e-3, ( - "the unbiased divisor produced the same answer, so this test cannot " - "detect the wrong one" - ) - - @pytest.mark.parametrize("m", [3, 16]) - def test_identity_holds_with_a_degenerate_zero_variance_set(self, m): - """Constant activation: the variance term must vanish exactly.""" - n = 8 - gen = torch.Generator().manual_seed(5) - F_D = torch.randn(n, generator=gen, dtype=torch.float64).to(torch.complex128) - dF = torch.randn(n, generator=gen, dtype=torch.float64).to(torch.complex128) - alpha = torch.full((m,), 0.31, dtype=torch.float64) - - brute = _brute_force_mean_intensity(F_D, dF, alpha) - abar, var = _sample_moments(alpha) - assert float(var) == pytest.approx(0.0, abs=1e-30) - assert torch.allclose(brute, (F_D + abar * dF).abs() ** 2, rtol=1e-13) - - def test_bernoulli_activation_gives_the_incoherent_sum(self): - """The lambda = 1 limit: fully-lit or fully-dark crystals add in intensity.""" - n = 64 - gen = torch.Generator().manual_seed(7) - F_D = torch.randn(n, generator=gen, dtype=torch.float64).to(torch.complex128) - F_L = torch.randn(n, generator=gen, dtype=torch.float64).to(torch.complex128) - dF = F_L - F_D - - w = 0.25 - m = 400 - alpha = torch.zeros(m, dtype=torch.float64) - alpha[: int(w * m)] = 1.0 - - brute = _brute_force_mean_intensity(F_D, dF, alpha) - incoherent = (1 - w) * F_D.abs() ** 2 + w * F_L.abs() ** 2 - assert torch.allclose(brute, incoherent, rtol=1e-12) - - # ...and the two-moment form reproduces it, with lambda exactly 1. - abar, var = _sample_moments(alpha) - assert float(var) == pytest.approx(abar * (1 - abar), rel=1e-12) - two_moment = (F_D + abar * dF).abs() ** 2 + var * dF.abs() ** 2 - assert torch.allclose(brute, two_moment, rtol=1e-12) - - -# ===================================================================== -# Integration against the real collection stack -# ===================================================================== - - @pytest.fixture(scope="module") def collection(pdb_dir, mtz_dir): """A dark/light collection on 1DAW, which is the only fixture with I/SIGI.""" @@ -179,18 +65,36 @@ def test_lambda_zero_reduces_to_the_squared_mean(self, collection): ) assert torch.equal(model, mean.abs() ** 2) - def test_lambda_zero_survives_a_poisoned_derivative(self, collection): - """The coherent limit must skip the variance branch, not multiply it by zero. - - A non-finite entry times exactly zero is NaN, which would poison the whole - gradient; this is what makes the short-circuit load-bearing rather than an - optimisation. - """ + def test_lambda_zero_skips_nonfinite_derivatives(self, collection, monkeypatch): + """Zero dispersion must not evaluate a potentially non-finite derivative.""" dc, mc, scaler = collection mc.set_lambda_twin(0.0) - target = _target(dc, mc, scaler) - assert not target._variance_is_live(mc.sigma_alpha_sq) - assert torch.isfinite(target.forward()) + weights = mc.fractions_matrix() + monkeypatch.setattr(mc, "fractions_matrix", lambda: weights) + + def forbidden(): + raise AssertionError("coherent prediction evaluated its variance branch") + + monkeypatch.setattr(mc, "activation_jacobian", forbidden) + assert torch.isfinite(_target(dc, mc, scaler).forward()) + + def test_full_dispersion_matches_incoherent_intensity(self, collection): + """Fully dark or fully lit crystals mix in intensity, including solvent.""" + dc, mc, scaler = collection + mc.set_lambda_twin(1.0) + try: + target = _target(dc, mc, scaler) + actual = target.intensity_model(recalc=True) + components = dc.component_structure_factors(mc, recalc=False) + basis = torch.eye( + mc.n_base_models, device=components.device, dtype=components.real.dtype + ) + pure = scaler.forward_batched(components, basis) + expected = mc.fractions_matrix() @ pure.abs().square() + # Complex mixture sums lose relative precision near solvent cancellation. + torch.testing.assert_close(actual, expected, rtol=2e-5, atol=1e-5) + finally: + mc.set_lambda_twin(0.0) def test_a_nonzero_lambda_changes_the_prediction(self, collection): """Anti-vacuity: the variance branch must actually do something.""" @@ -236,7 +140,8 @@ def test_the_variance_term_is_sigma_sq_times_the_scaled_jacobian(self, collectio def test_the_reference_row_carries_no_variance(self, collection): """The dark's Jacobian row is exactly zero, so its prediction is coherent - regardless of the dispersion -- a dark dataset holds no activation information.""" + regardless of the dispersion -- a dark dataset holds no activation information. + """ dc, mc, scaler = collection keys = _target(dc, mc, scaler)._keys() assert keys[0] == "dark" @@ -302,8 +207,16 @@ def test_stats_report_the_activation_moments(self, collection): mc.set_lambda_twin(0.25) try: stats = _target(dc, mc, scaler).stats() - for key in ("alpha_mean", "lambda_twin", "sigma_alpha_sq", "alpha_sd", - "dI_frac", "rwork", "rfree", "loss"): + for key in ( + "alpha_mean", + "lambda_twin", + "sigma_alpha_sq", + "alpha_sd", + "dI_frac", + "rwork", + "rfree", + "loss", + ): assert key in stats, f"missing stat: {key}" assert stats["alpha_mean"].value == pytest.approx(0.22, abs=1e-4) assert stats["lambda_twin"].value == pytest.approx(0.25, abs=1e-4) @@ -323,8 +236,7 @@ def test_subset_selection_is_honoured(self, collection, use_set): target = _target(dc, mc, scaler, use_set=use_set) assert target.use_set == use_set expected = sum( - (dc[k].work if use_set == "work" else dc[k].free).n - for k in target._keys() + (dc[k].work if use_set == "work" else dc[k].free).n for k in target._keys() ) assert target._n_reflections() == expected @@ -362,13 +274,7 @@ def test_construction_fails_without_intensities(self, pdb_dir, mtz_dir): @pytest.mark.integration class TestNonFiniteObservations: """Real reflection files carry non-finite intensities, and they must not reach the - gradient. - - Masking the *loss* is not enough. ``torch.where`` picks the finite branch for the - value while still backpropagating through the branch it discarded, so one NaN - observation turns every parameter gradient into NaN and every optimizer step is - rejected -- a refinement that silently does nothing rather than one that fails. - """ + gradient.""" def test_a_nan_observation_does_not_poison_the_gradient(self, collection): dc, mc, scaler = collection @@ -378,7 +284,6 @@ def test_a_nan_observation_does_not_poison_the_gradient(self, collection): with torch.no_grad(): data.I[5] = float("nan") data.I[11] = float("inf") - data._corrected_I_fp = None target = _target(dc, mc, scaler) loss = target.forward() @@ -394,7 +299,6 @@ def test_a_nan_observation_does_not_poison_the_gradient(self, collection): finally: with torch.no_grad(): data.I.copy_(saved) - data._corrected_I_fp = None mc.base_models[1].xyz.refinable_params.grad = None def test_a_nan_sigma_does_not_poison_the_gradient(self, collection): @@ -404,7 +308,6 @@ def test_a_nan_sigma_does_not_poison_the_gradient(self, collection): try: with torch.no_grad(): data.I_sigma[7] = float("nan") - data._corrected_I_fp = None target = _target(dc, mc, scaler) loss = target.forward() @@ -414,7 +317,6 @@ def test_a_nan_sigma_does_not_poison_the_gradient(self, collection): finally: with torch.no_grad(): data.I_sigma.copy_(saved) - data._corrected_I_fp = None mc.base_models[1].xyz.refinable_params.grad = None def test_the_bad_reflections_are_excluded_not_absorbed(self, collection): @@ -429,12 +331,10 @@ def test_the_bad_reflections_are_excluded_not_absorbed(self, collection): baseline = target.forward().item() with torch.no_grad(): data.I[3] = float("nan") - data._corrected_I_fp = None with_nan = _target(dc, mc, scaler).forward().item() finally: with torch.no_grad(): data.I.copy_(saved) - data._corrected_I_fp = None # One reflection out of tens of thousands: the loss should drop slightly, not # jump by a penalty term. @@ -447,28 +347,7 @@ class TestWeightCalibration: """Intensities are squared amplitudes, so this target's gradient is orders of magnitude away from the amplitude target beside it. Left uncalibrated it swamps the geometry restraints and buys R-free by moving the model further than the data - supports. - """ - - def test_the_uncalibrated_mismatch_is_large(self, collection): - """Anti-vacuity: if the two targets already pushed equally, calibration would - be pointless.""" - dc, mc, scaler = collection - from torchref.refinement.targets import CollectionDifferenceTarget - - params = [p for p in mc.base_models[1].parameters() if p.requires_grad] - diff = CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0) - target = _target(dc, mc, scaler) - - def gnorm(t): - g = torch.autograd.grad(t.forward(), params, allow_unused=True) - return sum(float((x**2).sum()) for x in g if x is not None) ** 0.5 - - ratio = gnorm(target) / gnorm(diff) - assert ratio > 10 or ratio < 0.1, ( - f"gradient ratio is {ratio:.3g}; the two targets are already matched and " - f"this fixture cannot show why calibration is needed" - ) + supports.""" def test_calibration_equalises_the_gradient_norms(self, collection): dc, mc, scaler = collection diff --git a/tests/unit/scaling/test_collection_joint_scale_fit.py b/tests/unit/scaling/test_collection_joint_scale_fit.py index 93dd7bb2..bd3a9583 100644 --- a/tests/unit/scaling/test_collection_joint_scale_fit.py +++ b/tests/unit/scaling/test_collection_joint_scale_fit.py @@ -1,13 +1,4 @@ -"""``CollectionScaler.refine_lbfgs_joint`` -- the model-to-data fit, per row. - -It used to hand-roll a Rice likelihood inline with ``beta = sigma_obs**2`` and no -normaliser at all. Both are now gone: it builds a row of ``XRAY_TARGETS``, exactly as -``ScalerBase.refine_lbfgs`` does, and normalises the objective because L-BFGS converges on -absolute tolerances. - -The tests here are mostly about what must NOT differ between the two scale fits, since -"one likelihood, two call sites" is the whole claim. -""" +"""Joint model-to-data scaling objectives and shared parameter ownership.""" import inspect @@ -24,7 +15,6 @@ def collection(pdb_dir, mtz_dir): from torchref import LBFGSRefinement, ReflectionData from torchref.io.datasets.collection import DatasetCollection from torchref.model.model_collection import ModelCollection - from torchref.scaling.collection_scaler import CollectionScaler ref = LBFGSRefinement(data_file=str(mtz), pdb=str(pdb), verbose=0) extra = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) @@ -46,25 +36,6 @@ def _fresh_scaler(collection): return CollectionScaler(dc, mc, verbose=0).initialize() -@pytest.mark.unit -def test_the_hand_rolled_rice_is_gone(): - """No private likelihood in the scaling package. - - It set ``beta = sigma_obs**2`` -- pairing a measurement sigma with a Rice ``Sigma``, - which asserts an isotropic *complex* error where sigma_obs carries no phase at all. - ``xray_likelihoods`` records that no regime makes that correct, and the taxonomy - deliberately offers no such row; a copy inside the scaler bypassed both. - """ - import torchref.scaling.collection_scaler as cs - - src = inspect.getsource(cs) - assert "rice_math" not in src, "the scaler grew a private Rice likelihood again" - assert "SigmaAEstimator" not in src, ( - "the scaler precomputes beta again; rows own their own estimator" - ) - assert "create_xray_target" in src, "the joint fit must build a taxonomy row" - - @pytest.mark.unit def test_it_offers_exactly_the_selectable_objectives(): from torchref.scaling.collection_scaler import CollectionScaler @@ -72,9 +43,9 @@ def test_it_offers_exactly_the_selectable_objectives(): sig = inspect.signature(CollectionScaler.refine_lbfgs_joint) assert sig.parameters["scale_target"].default == DEFAULT_SCALE_TARGET - assert DEFAULT_SCALE_TARGET == "ls", ( - "the joint fit's default must track the single-dataset one" - ) + assert ( + DEFAULT_SCALE_TARGET == "ls" + ), "the joint fit's default must track the single-dataset one" assert "nll" in SCALE_TARGETS and "ml_noalpha" in SCALE_TARGETS @@ -88,11 +59,7 @@ def test_unknown_objective_fails_closed(collection): @pytest.mark.integration @pytest.mark.parametrize("scale_target", ["ls", "nll", "ml_noalpha"]) def test_every_objective_fits_finite_parameters(collection, scale_target): - """Every selectable row must drive the joint fit to finite parameters. - - The old fit produced non-finite scales on real data; the objective it handed L-BFGS - was unnormalised against absolute tolerances, which is a documented way to get there. - """ + """Every selectable row must drive the joint fit to finite parameters.""" scaler = _fresh_scaler(collection) m = scaler.refine_lbfgs_joint( nsteps=2, max_iter=20, verbose=False, scale_target=scale_target @@ -104,12 +71,7 @@ def test_every_objective_fits_finite_parameters(collection, scale_target): @pytest.mark.integration def test_the_dataset_view_shares_the_parents_parameters(collection): - """A row's scaler must be a view, not a copy. - - If the view registered the parent as a submodule, the parent's parameters would be - counted twice and L-BFGS would see duplicate leaves. If it copied them, the fit would - optimise something the collection never reads. - """ + """A row's scaler must be a view, not a copy.""" from torchref.scaling.collection_scaler import _DatasetScalerView scaler = _fresh_scaler(collection) @@ -119,28 +81,10 @@ def test_the_dataset_view_shares_the_parents_parameters(collection): # No parameters of its own -- only the bound fractions buffer. assert list(view.parameters()) == [] + assert view.device == scaler.device # And it routes through the parent's mixed-solvent path. with torch.no_grad(): fcalc = mc[mc.dark_key](dc[mc.dark_key].hkl) got = view(fcalc) want = scaler.forward_mixed(fcalc, fracs) assert torch.equal(got, want) - - -@pytest.mark.integration -def test_the_objective_is_normalised(collection): - """The loss L-BFGS sees must be O(1), not the data's own magnitude. - - ``tolerance_grad``/``tolerance_change`` are absolute. This fit had no normaliser at - all, so on a large work set under unit weights the loss reached a magnitude where its - own float32 ulp exceeded the decrease the line search was trying to resolve. - """ - src = inspect.getsource( - __import__( - "torchref.scaling.collection_scaler", fromlist=["CollectionScaler"] - ).CollectionScaler.refine_lbfgs_joint - ) - assert "_norm" in src, "the joint objective is unnormalised again" - # The U penalty must NOT follow the observable/objective: sharing a normaliser that - # moves with the objective silently changes the regularisation strength. - assert "work.F" in src, "the normaliser must be built from amplitudes" diff --git a/tests/unit/scaling/test_collection_scaler_batched.py b/tests/unit/scaling/test_collection_scaler_batched.py index 66fa4b8a..a3943c84 100644 --- a/tests/unit/scaling/test_collection_scaler_batched.py +++ b/tests/unit/scaling/test_collection_scaler_batched.py @@ -1,21 +1,4 @@ -"""``forward_batched`` must agree with ``forward_mixed`` row by row, and stay affine. - -The batched form exists so ``T`` mixtures share one pass through the scale parameters. -Two properties make it usable: - -* **row agreement** -- row ``i`` of the batch must equal the unbatched call on row ``i``, - otherwise the saving is bought with wrong numbers; -* **affinity in the mixing weights** -- ``ScalerBase.forward`` is - ``K * b * (aniso * F_calc + f_sol)`` and the mixed solvent is linear in the weights, so - scaling a *derivative* of the fractions returns the derivative of the scaled structure - factors. That is what lets a second moment be built from the same machinery instead of - a separate differentiation path. - -Affinity is tested with a **secant**, not a finite difference. Because the mixture is -exactly linear in the activation fraction, ``S(a1) - S(a2)`` equals -``(a1 - a2) * dS/da`` exactly, with no truncation term to tolerate -- so the assertion -is at float precision rather than ``O(h^2)``. -""" +"""Batched scaling preserves individual results, affinity and solvent caches.""" import pytest import torch @@ -76,22 +59,16 @@ def test_every_row_matches_forward_mixed(self, scaled_collection): for i in range(w.shape[0]): single = scaler.forward_mixed(fcalc[i], w[i]) - assert torch.allclose(batched[i], single, rtol=1e-6, atol=1e-6), ( - f"batched row {i} disagrees with forward_mixed" - ) - - def test_component_solvent_stack_shape(self, scaled_collection): - dc, mc, scaler = scaled_collection - stack = scaler.compute_component_solvent_raw() - assert stack.shape == (mc.n_base_models, len(dc.hkl)) - assert stack.is_complex() + assert torch.allclose( + batched[i], single, rtol=1e-6, atol=1e-6 + ), f"batched row {i} disagrees with forward_mixed" - def test_solvent_stack_rows_are_the_per_component_solvents( - self, scaled_collection - ): + def test_solvent_stack_rows_are_the_per_component_solvents(self, scaled_collection): """A transposed or misordered stack would still have the right shape.""" _, mc, scaler = scaled_collection stack = scaler.compute_component_solvent_raw() + assert stack.shape == (mc.n_base_models, len(scaler.hkl)) + assert stack.is_complex() for k in range(mc.n_base_models): assert torch.equal(stack[k], scaler._get_component_f_sol_raw(k)) @@ -99,46 +76,27 @@ def test_solvent_stack_rows_are_the_per_component_solvents( @pytest.mark.integration class TestAffineInTheMixingWeights: def test_secant_in_alpha_equals_the_scaled_jacobian(self, scaled_collection): - """The property the two-moment forward model rests on. - - ``forward_batched(dF, J)`` with ``J = dW/da`` is the derivative of the scaled - mixture, including the solvent term. Exact, because everything between the - weights and the output is affine. - """ + """The property the two-moment forward model rests on.""" dc, mc, scaler = scaled_collection components = dc.component_structure_factors(mc, recalc=True) + solvent = scaler.compute_component_solvent_raw() + assert not torch.allclose(solvent[0], solvent[1]) a1, a2 = 0.60, 0.10 w1, w2 = _weights(a1), _weights(a2) jac = torch.tensor([[-1.0, 1.0]]) # d/da of [1 - a, a] s1 = scaler.forward_batched(mc.mix_component_fcalcs(components, w1), w1) s2 = scaler.forward_batched(mc.mix_component_fcalcs(components, w2), w2) - deriv = scaler.forward_batched( - mc.mix_component_fcalcs(components, jac), jac - ) + deriv = scaler.forward_batched(mc.mix_component_fcalcs(components, jac), jac) secant = s1 - s2 expected = (a1 - a2) * deriv - rel = (secant - expected).abs().max() / expected.abs().max() - assert rel < 1e-5, ( - f"secant and scaled Jacobian disagree by {rel:.2e}; the scaler is not " - f"affine in the mixing weights, so a derivative cannot be scaled this way" - ) - - def test_the_solvent_term_is_included_in_the_derivative(self, scaled_collection): - """Anti-vacuity: if the per-component solvents were identical, the solvent - would cancel out of the Jacobian and the test above would hold even with the - solvent term dropped.""" - _, mc, scaler = scaled_collection - stack = scaler.compute_component_solvent_raw() - if mc.n_base_models < 2: - pytest.skip("needs at least two components") - assert not torch.allclose(stack[0], stack[1]), ( - "per-component solvents are identical, so this fixture cannot detect a " - "dropped solvent derivative" - ) + # Subtraction roundoff scales with its operands, not the smaller secant. + scale = (s1.abs() + s2.abs() + expected.abs()).max() + tolerance = 16 * torch.finfo(s1.real.dtype).eps * scale + assert (secant - expected).abs().max() <= tolerance def test_scaling_is_linear_in_the_structure_factors(self, scaled_collection): """The other half of affinity: doubling F_calc at fixed weights doubles the @@ -161,11 +119,7 @@ def test_scaling_is_linear_in_the_structure_factors(self, scaled_collection): @pytest.mark.integration class TestSolventCacheIsNotPoisoned: def test_batched_calls_leave_the_cache_alone(self, scaled_collection): - """Two batched calls with different weights, then a plain one. - - The Jacobian-weighted call carries negative weights, so a leaked cache would - show up as a sign error rather than a small perturbation. - """ + """Two batched calls with different weights, then a plain one.""" dc, mc, scaler = scaled_collection components = dc.component_structure_factors(mc, recalc=True) w = _weights(0.22) @@ -178,6 +132,6 @@ def test_batched_calls_leave_the_cache_alone(self, scaled_collection): scaler.forward_batched(mc.mix_component_fcalcs(components, jac), jac) after = scaler.forward_mixed(fcalc[0], w[0]) - assert torch.allclose(after, before, rtol=1e-6, atol=1e-6), ( - "a batched call changed what a later forward_mixed returns" - ) + assert torch.allclose( + after, before, rtol=1e-6, atol=1e-6 + ), "a batched call changed what a later forward_mixed returns" diff --git a/tests/unit/scaling/test_dataset_scaler.py b/tests/unit/scaling/test_dataset_scaler.py index 63767750..17180a33 100644 --- a/tests/unit/scaling/test_dataset_scaler.py +++ b/tests/unit/scaling/test_dataset_scaler.py @@ -77,17 +77,32 @@ def test_live_access_scales_both_sigmas_and_all_entrypoints(deposited): data = dc["0"] with torch.no_grad(): dc.scaler.raw_parameters[0, 0] = 2 * math.log(2) - for name, power in [("F", 1), ("F_sigma", 1), ("I", 2), ("I_sigma", 2)]: - torch.testing.assert_close( - getattr(data, name), getattr(data, name + "_raw") * 2**power - ) - torch.testing.assert_close(data.work.F, data.F[data.work.mask]) - torch.testing.assert_close(data.work.sigF, data.F_sigma[data.work.mask]) - torch.testing.assert_close(data.work.sigI, data.I_sigma[data.work.mask]) - torch.testing.assert_close(data.work.sigI_raw, data.I_sigma_raw[data.work.mask]) - torch.testing.assert_close(dc.stack_F_obs()[0], data.F) + dc.scaler.raw_parameters[0, 1] = 0.1 + correction = 2 * torch.exp( + 0.05 * (data.hkl[:, 0] / dc.scaler.hkl_scale[0]).square() + ) + data.generate_validation_set(val_fraction_of_free=0.5, seed=0) + for name, power, subset_attr, stack in [ + ("F", 1, "F", dc.stack_F_obs), + ("F_sigma", 1, "sigF", dc.stack_F_sigma), + ("I", 2, "I", dc.stack_I_obs), + ("I_sigma", 2, "sigI", dc.stack_I_sigma), + ]: + raw = getattr(data, name + "_raw") + actual = getattr(data, name) + torch.testing.assert_close(actual, raw * correction**power) + torch.testing.assert_close(stack()[0], actual) + for kind in ("work", "free", "validation"): + subset = getattr(data, kind) + torch.testing.assert_close( + getattr(subset, subset_attr), actual[subset.mask] + ) + torch.testing.assert_close( + getattr(subset, subset_attr + "_raw"), raw[subset.mask] + ) torch.testing.assert_close(dc(mask=False)["0"][1], data.F) - torch.testing.assert_close(dc.stack_I_sigma()[0], data.I_sigma) + torch.testing.assert_close(data.get_corrected_data(), (data.F, data.F_sigma)) + torch.testing.assert_close(data.get_corrected_intensities(), (data.I, data.I_sigma)) dc.scaler.requires_grad_(True) for _ in range(2): dc.scaler.zero_grad() @@ -299,3 +314,18 @@ def test_ded_context_consumes_collection_scaled_views(deposited, monkeypatch): context["w_dfo"].norm() / dc["dark"].F[context["refl_mask"]].norm() ) assert relative_difference < 1e-5 + + +def test_missing_intensity_access(deposited): + """Raw and scaled views expose missing columns and reject intensity-only reads.""" + dc = collection(deposited, (1, 2)) + dc["0"].I = dc["0"].I_sigma = None + raw = dc["0"] + dc.scale(nsteps=1) + for data in (raw, dc["0"]): + with pytest.raises(ValueError, match="No intensities"): + data.get_corrected_intensities() + for name in ("I", "sigI", "I_raw", "sigI_raw"): + assert getattr(data.work, name) is None + with pytest.raises(ValueError, match="'0'"): + dc.stack_I_obs() diff --git a/tests/unit/scaling/test_f_sol_override_contract.py b/tests/unit/scaling/test_f_sol_override_contract.py index b7e71bfe..d8924fac 100644 --- a/tests/unit/scaling/test_f_sol_override_contract.py +++ b/tests/unit/scaling/test_f_sol_override_contract.py @@ -1,18 +1,4 @@ -"""``f_sol_override`` must be a pure argument, and must not change the output rank. - -Two separate contracts on :meth:`ScalerBase.forward`, both load-bearing for any caller -that scales several models against one shared scaler: - -1. Passing ``f_sol_override`` must not write ``_f_sol_raw``. A caller that scales two - different fraction mixtures in a row otherwise leaves the *second* mixture's solvent - cached, and every later call that does not pass an override silently reads it. -2. A batched ``(T, N)`` override paired with a batched ``(T, N)`` ``fcalc`` must return - ``(T, N)``. The solvent term is broadcast with ``unsqueeze(0)``, which is right for a - per-reflection ``(N,)`` solvent and wrong for one that already carries the batch axis. - -Both are exercised through ``forward`` rather than asserted on internals, so the stub only -has to stand in for the solvent model. -""" +"""Solvent overrides preserve output shape and leave the default solvent cache intact.""" from types import SimpleNamespace @@ -21,11 +7,7 @@ class _StubSolvent: - """Minimal stand-in for :class:`SolventModel` on the k_sol/B_sol path. - - ``get_rec_solvent`` returns a recognisable constant so a cached value can be told - apart from a freshly-passed override by value as well as by identity. - """ + """Minimal stand-in for :class:`SolventModel` on the k_sol/B_sol path.""" optimize_phase = False @@ -41,19 +23,12 @@ def damping(self, s_half_sq): def get_rec_solvent(self, hkl): self.n_reads += 1 - return torch.full( - (hkl.shape[0],), self.value, dtype=torch.complex64 - ) + return torch.full((hkl.shape[0],), self.value, dtype=torch.complex64) @pytest.fixture def scaler_with_stub(): - """A bare ``Scaler`` with only the solvent branch live. - - ``bins`` drives the full-size check in ``forward``; the anisotropy, Chebyshev and - per-bin-B branches are all absent, so ``forward`` reduces to - ``fcalc + k_sol * f_sol`` and any shape or caching defect is unobscured. - """ + """A bare ``Scaler`` with only the solvent branch live.""" from torchref.scaling.scaler import Scaler n = 6 @@ -72,32 +47,12 @@ def scaler_with_stub(): class TestOverrideDoesNotMutateTheCache: - @pytest.mark.unit - def test_override_leaves_the_cache_untouched(self, scaler_with_stub): - """The override is an argument, not an assignment.""" - scaler, n, dev = scaler_with_stub - override = torch.full((n,), 2.0, dtype=torch.complex64, device=dev) - - scaler.forward(torch.ones(n, dtype=torch.complex64, device=dev), - f_sol_override=override) - - assert scaler._f_sol_raw is not override, ( - "forward stored the override in the solvent cache" - ) - assert scaler._f_sol_raw is None, ( - "forward populated the solvent cache from an override; a later call " - "without one will read this instead of the model's own solvent" - ) @pytest.mark.unit def test_a_later_call_without_an_override_sees_the_model_solvent( self, scaler_with_stub ): - """The consequence of the leak, stated in terms a caller can observe. - - Two mixtures scaled in a row, then a plain call: the plain call must use the - solvent model, not whichever mixture happened to be scaled last. - """ + """The consequence of the leak, stated in terms a caller can observe.""" scaler, n, dev = scaler_with_stub fcalc = torch.ones(n, dtype=torch.complex64, device=dev) @@ -115,48 +70,33 @@ def test_a_later_call_without_an_override_sees_the_model_solvent( class TestOverridePreservesRank: - @pytest.mark.unit - def test_batched_override_with_batched_fcalc_keeps_the_batch_rank( - self, scaler_with_stub - ): - """``(T, N)`` in, ``(T, N)`` out -- the batched multi-dataset contract.""" - scaler, n, dev = scaler_with_stub - t = 3 - fcalc = torch.ones((t, n), dtype=torch.complex64, device=dev) - override = torch.full((t, n), 2.0, dtype=torch.complex64, device=dev) - - out = scaler.forward(fcalc, f_sol_override=override) - - assert out.shape == (t, n), ( - f"batched override changed the output rank: got {tuple(out.shape)}, " - f"expected {(t, n)}" - ) @pytest.mark.unit def test_each_batch_row_matches_the_unbatched_call(self, scaler_with_stub): - """Rank alone is not enough -- the rows must also be the right ones. - - A wrong broadcast can restore the shape and still pair row *i* of ``fcalc`` - with the wrong row of the solvent. - """ + """Rank alone is not enough -- the rows must also be the right ones.""" scaler, n, dev = scaler_with_stub t = 3 - fcalc = torch.stack([ - torch.full((n,), float(i + 1), dtype=torch.complex64, device=dev) - for i in range(t) - ]) - override = torch.stack([ - torch.full((n,), float(10 * (i + 1)), dtype=torch.complex64, device=dev) - for i in range(t) - ]) + fcalc = torch.stack( + [ + torch.full((n,), float(i + 1), dtype=torch.complex64, device=dev) + for i in range(t) + ] + ) + override = torch.stack( + [ + torch.full((n,), float(10 * (i + 1)), dtype=torch.complex64, device=dev) + for i in range(t) + ] + ) batched = scaler.forward(fcalc, f_sol_override=override) + assert batched.shape == (t, n) for i in range(t): scaler._f_sol_raw = None single = scaler.forward(fcalc[i], f_sol_override=override[i]) - assert torch.allclose(batched[i], single), ( - f"batched row {i} does not match the equivalent unbatched call" - ) + assert torch.allclose( + batched[i], single + ), f"batched row {i} does not match the equivalent unbatched call" @pytest.mark.unit def test_unbatched_override_is_unchanged(self, scaler_with_stub): diff --git a/torchref/cli/_common.py b/torchref/cli/_common.py index e405ebb1..8efb3201 100644 --- a/torchref/cli/_common.py +++ b/torchref/cli/_common.py @@ -334,24 +334,15 @@ def add_dual_model_args( type=str, help="Dark / reference state model file (PDB or CIF)", ) - if light_model_required: - inp.add_argument( - "-lm", - "--light-model", - required=True, - type=str, - help="Light / triggered state model file (PDB or CIF)", - ) - else: - inp.add_argument( - "-lm", - "--light-model", - type=str, - default=None, - help="Light / triggered state model file (PDB or CIF). Optional: without " - "it only the weighted difference map is written, which needs the dark " - "state's phases and no light-state model at all.", + light_help = "Light / triggered state model file (PDB or CIF)" + if not light_model_required: + light_help += ( + ". Optional: without it only the weighted difference map is written, " + "using the dark state's phases." ) + inp.add_argument( + "-lm", "--light-model", required=light_model_required, type=str, help=light_help + ) inp.add_argument( "-dsf", "--dark-structure-factor", diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index dece664a..c7242750 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -210,10 +210,8 @@ def compute_rfactors(model, data, scaler): with torch.no_grad(): hkl = data.hkl fcalc = model(hkl) - # Which scaler this is decides how it is called: a CollectionScaler mixes - # components and needs the model's fractions, a single-dataset Scaler takes - # the structure factors straight. Asked of the scaler, not inferred from the - # caller, so the dark-only path needs no special case. + # CollectionScaler needs fractions to mix components; a single-dataset + # Scaler consumes structure factors directly. if hasattr(scaler, "forward_mixed"): fcalc_scaled = scaler.forward_mixed(fcalc, model.fractions) else: @@ -274,10 +272,8 @@ def setup_loss_state(dataset_collection, model_collection, scaler, two_moment_target = CollectionTwoMomentIntensityTarget( dataset_collection, model_collection, scaler=scaler, verbose=1, ) - # Intensities are squared amplitudes, so this target's gradient is on a - # completely different scale from the difference target beside it. Match them - # once here; left uncalibrated it swamps the geometry restraints and buys - # R-free by moving the model far further than the data supports. + # Match intensity and amplitude gradient norms to balance the targets + # against the geometry restraints despite their different units. two_moment_target.calibrate_base_weight( diff_target, list(model_light.parameters()) ) @@ -311,9 +307,7 @@ def compute_bayes_extrapolated_amplitudes( Propagated uncertainty of the extrapolated amplitude. Taken from the caller rather than rebuilt here: ``F_ext`` is linear in the observations with ``dF_ext/dF_light = 1/f`` and ``dF_ext/dF_dark = 1 - 1/f = -(1-f)/f``, so the - dark term carries a ``(1-f)**2`` weight that is easy to drop. This function - used to drop it, over-weighting the dark term by ``1/(1-f)**2`` -- 1.64x at - f = 0.22 -- which biased τ² low and over-shrank every reflection. + dark term carries a ``(1-f)**2`` weight. phi_dark, phi_mixed : Tensor (N,) Calculated phases (radians) for the dark and mixed models. f : float or Tensor @@ -387,8 +381,7 @@ def _two_moment_columns(mc, dc, mask, fcalc_dark_full, fcalc_mixed_full, ------- tuple ``(columns, types)`` -- the values, and the MTZ type letter for each. Carrying - the type beside the value is what stops a column reaching the file with whatever - dtype numpy produced, which is the failure the old parallel name lists invited. + the type beside the value preserves the crystallographic column type. """ import numpy as np @@ -435,16 +428,8 @@ def _np(t): DDF = DF_corr - diff_Fobs sig_DF_corr = np.sqrt(sig_F_corr**2 + sig_dark**2) - # The phase-AWARE difference, rebuilt with the decontaminated light amplitude. - # - # It must be the modulus of the complex vector difference - # ``|F_corr e^{i phi_light} - F_dark e^{i phi_dark}|``, exactly as the uncorrected - # ``Fobs_diff_phased`` is -- NOT ``|F_corr - F_dark|``. The two are different - # quantities: the vector form carries the phase rotation between dark and light, - # which is the whole point of a phase-aware coefficient, while the scalar form is - # phase-blind. Using the scalar one here made the corrected coefficients only 36% - # correlated with their uncorrected twins even though the amplitudes behind them - # agree to 99.99%. + # Use the modulus of the complex vector difference so the corrected + # coefficient retains the phase rotation between the dark and light states. F_corr_phased = torch.as_tensor( F_corr, dtype=F_obs_dark_phased.real.dtype, device=F_obs_dark_phased.device ) * torch.exp(1j * phi_mixed) @@ -594,16 +579,16 @@ def _phasing_columns(mc, scaler, hkl_all, mask, *, fcalc_dark, Fobs_dark_vals, Fobs_diff_phased = torch.abs( F_obs_light_phased - F_obs_dark_phased ).detach().cpu().numpy() - columns.update({ - "2mDFop-DFc": (2 * Fobs_diff_phased - Fcalc_diff_amp) * weights, - "mDFop-DFc": (Fobs_diff_phased - Fcalc_diff_amp) * weights, - "PHIC_diff": torch.angle(fcalc_diff).detach().rad2deg().cpu().numpy(), - "DFc": Fcalc_light - Fcalc_dark, - # The modulus of the complex vector difference. Named ``_phased`` rather - # than ``_complex``: the column holds a real amplitude, and the old name - # said otherwise. - "DFc_phased": Fcalc_diff_amp, - }) + columns.update( + { + "2mDFop-DFc": (2 * Fobs_diff_phased - Fcalc_diff_amp) * weights, + "mDFop-DFc": (Fobs_diff_phased - Fcalc_diff_amp) * weights, + "PHIC_diff": torch.angle(fcalc_diff).detach().rad2deg().cpu().numpy(), + "DFc": Fcalc_light - Fcalc_dark, + # This column holds the real modulus of the complex vector difference. + "DFc_phased": Fcalc_diff_amp, + } + ) types.update({ "2mDFop-DFc": "F", "mDFop-DFc": "F", "PHIC_diff": "P", "DFc": "F", "DFc_phased": "F", @@ -708,7 +693,6 @@ def _np(t): columns.update({ "FEXT_PHASED": _np(amp_phased), - # Computed all along and never written, though the docs claimed it. "SIGFEXT_PHASED": _np(sig_light_extra), "2FEXT_PHASED-Fc": _np(2 * amp_phased - amp_calc_phased), "FEXT_PHASED-Fc": _np(amp_phased - amp_calc_phased), @@ -867,10 +851,8 @@ def write_results_mtz(dc, dark_model, scaler, filename, *, mc=None, cell=data_dark.cell.data.cpu().tolist(), spacegroup=data_dark.spacegroup.hm, ) - # Every layer returns its columns' MTZ types beside the values, so a new column - # cannot reach the file with whatever dtype numpy produced -- the failure the old - # parallel name lists invited. ``infer_mtz_dtypes`` is then the same safety net the - # canonical writer in ``torchref/io/mtz.py`` uses. + # Carry MTZ types with the values; infer_mtz_dtypes also checks the + # result using the canonical writer's rules. missing = set(columns) - set(types) if missing: raise AssertionError(f"columns with no declared MTZ type: {sorted(missing)}") @@ -1032,13 +1014,8 @@ def main(): ) return 1 if (args.lambda_twin > 0.0 or args.refine_lambda_twin) and not args.two_moment: - # There is no longer a weighting-only path. Measured on ground truth (inject a - # known displacement, refine from the dark model, 8 seeds per regime): putting the - # contamination in the VARIANCE lost 8/8 seeds at high contamination, 95% CI - # [+0.0066, +0.0102] A, and was null at low. Putting the same quantity in the MEAN - # -- which is what --two-moment does -- won 16/16. Structured effects belong in the - # mean; only genuine measurement noise belongs in the variance, and down-weighting - # by |dF|^2 suppresses exactly the reflections carrying the difference signal. + # Activation heterogeneity changes the predicted mean intensity. Treating + # it as measurement variance downweights the reflections carrying the signal. print( "Error: --lambda-twin needs --two-moment. The dispersion enters the predicted " "intensity, not a weight: as a variance it down-weights the reflections whose " diff --git a/torchref/cli/difference_map.py b/torchref/cli/difference_map.py index 9e52601a..8253ee5c 100644 --- a/torchref/cli/difference_map.py +++ b/torchref/cli/difference_map.py @@ -76,7 +76,6 @@ def main(): """, ) - # --- Input files (creates "Input files" and "Column selection" groups) --- add_dual_model_args(parser, fraction_required=False, light_model_required=False) output = parser.add_argument_group("Output") @@ -92,7 +91,6 @@ def main(): register_timing() - # --- The light model gates everything that needs the light state's phases --- has_light_model = args.light_model is not None if has_light_model: if args.fraction is None: @@ -111,7 +109,6 @@ def main(): "sigmas." ) - # --- Validate input files --- to_check = [ (args.dark_model, "dark model"), (args.dark_structure_factor, "dark structure factor"), @@ -125,14 +122,11 @@ def main(): if validate_cif_files(args.cif): return 1 - # Ensure output directory exists out_path = Path(args.output) out_path.parent.mkdir(parents=True, exist_ok=True) - # --- Device --- device = parse_device_str(args.device) - # --- Header --- if args.verbose > 0: print("=" * 72) print("TorchRef Difference Map") @@ -167,11 +161,8 @@ def main(): write_results_mtz, ) - # --- Resolution --- d_min = args.dmin if args.dmin is not None else 1.0 - # --- Load data. dc.scale() puts the two datasets on one scale with no model, - # which is what makes the dark-only path possible at all. --- if args.verbose > 0: print("Loading reflection data...") sys.stdout.flush() @@ -218,7 +209,6 @@ def main(): print() sys.stdout.flush() - # --- Write MTZ --- if args.verbose > 0: print("Computing map coefficients...") sys.stdout.flush() diff --git a/torchref/cli/simulate_noisy_data.py b/torchref/cli/simulate_noisy_data.py index 36711fdc..7e37361d 100644 --- a/torchref/cli/simulate_noisy_data.py +++ b/torchref/cli/simulate_noisy_data.py @@ -213,7 +213,7 @@ def _run_reference_mode(args, model, device) -> int: sim.set_fcalc(fcalc_scaled) noisy = sim.add_noise(reference=ref, seed=args.seed, verbose=bool(args.verbose)) - _write_output(args, noisy, sim_clean=sim) + _write_output(args, noisy) return 0 @@ -244,12 +244,12 @@ def _run_parametric_mode(args, model, device) -> int: seed=args.seed, verbose=bool(args.verbose), ) - _write_output(args, noisy, sim_clean=dataset) + _write_output(args, noisy) return 0 -def _write_output(args, noisy: FcalcDataset, sim_clean: FcalcDataset) -> None: - """Write noisy Fcalc/Fobs to MTZ. ``sim_clean`` holds the pre-noise dataset.""" +def _write_output(args, noisy: FcalcDataset) -> None: + """Write noisy amplitudes or intensities and their uncertainties to MTZ.""" hkl_np = noisy.hkl.cpu().numpy() columns = { "H": hkl_np[:, 0], diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index 02fbae98..8c94a52d 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -635,13 +635,8 @@ def run_validation(args): "light_model": str(args.light_model), "fraction": args.fraction, "selection": args.selection, - # Recorded because it CHANGES THE ANSWER and is easy to leave at a - # different value between runs. On the figure-4 ligand, "light" masks 1535 - # voxels and scores CC 0.869, while "both" masks 1948 and scores 0.851 -- - # the union adds the volume the ligand vacated, where the density is - # negative and the model has to get a depletion right. Two runs differing - # only in this looked like a real improvement until the parameter was - # recovered by re-running, which is exactly what storing it prevents. + # Mask choice changes which density enters the correlation, so record it + # alongside the score for reproducibility. "mask_source": args.mask_source, "mask_radius": args.mask_radius, "dmin": d_min, diff --git a/torchref/io/datasets/base.py b/torchref/io/datasets/base.py index ce1c726f..5935fac3 100644 --- a/torchref/io/datasets/base.py +++ b/torchref/io/datasets/base.py @@ -9,7 +9,7 @@ import warnings from dataclasses import dataclass, field, fields -from typing import TYPE_CHECKING, Any, Dict, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, Optional import gemmi import torch diff --git a/torchref/io/datasets/collection.py b/torchref/io/datasets/collection.py index 3ee7eee9..9418272b 100644 --- a/torchref/io/datasets/collection.py +++ b/torchref/io/datasets/collection.py @@ -375,14 +375,14 @@ def _keys_or_all(self, keys: Optional[List[str]]) -> List[str]: def stack_F_obs(self, keys: Optional[List[str]] = None) -> torch.Tensor: """Scaled observed amplitudes, shape ``(n_datasets, n_reflections)``.""" return torch.stack( - [self._datasets[k]._corrected_or_raw()[0] for k in self._keys_or_all(keys)], + [self._datasets[k].F for k in self._keys_or_all(keys)], dim=0, ) def stack_F_sigma(self, keys: Optional[List[str]] = None) -> torch.Tensor: """Scaled amplitude sigmas, shape ``(n_datasets, n_reflections)``.""" return torch.stack( - [self._datasets[k]._corrected_or_raw()[1] for k in self._keys_or_all(keys)], + [self._datasets[k].F_sigma for k in self._keys_or_all(keys)], dim=0, ) @@ -395,7 +395,7 @@ def stack_I_obs(self, keys: Optional[List[str]] = None) -> torch.Tensor: If any selected dataset carries no intensities. """ return torch.stack( - [self._require_intensities(k)[0] for k in self._keys_or_all(keys)], + [self._require_intensities(k).I for k in self._keys_or_all(keys)], dim=0, ) @@ -408,19 +408,19 @@ def stack_I_sigma(self, keys: Optional[List[str]] = None) -> torch.Tensor: If any selected dataset carries no intensities. """ return torch.stack( - [self._require_intensities(k)[1] for k in self._keys_or_all(keys)], + [self._require_intensities(k).I_sigma for k in self._keys_or_all(keys)], dim=0, ) - def _require_intensities(self, key: str): - """``(I, I_sigma)`` scaled, with the dataset named in the error.""" + def _require_intensities(self, key: str) -> ReflectionData: + """Return a dataset with intensities, naming it if the column is missing.""" data = self._datasets[key] - if data.I is None: + if data.I_raw is None: raise ValueError( f"Dataset {key!r} carries no intensities; its reflection file had no " f"I/SIGI columns. An intensity-space target needs them on every member." ) - return data._corrected_or_raw_intensities() + return data def stack_masks( self, keys: Optional[List[str]] = None, use_set: str = "work" diff --git a/torchref/io/datasets/fcalc_data.py b/torchref/io/datasets/fcalc_data.py index f85444a2..d4a4a2aa 100644 --- a/torchref/io/datasets/fcalc_data.py +++ b/torchref/io/datasets/fcalc_data.py @@ -12,7 +12,7 @@ import pandas as pd import torch -from torchref.config import get_default_device, get_float_dtype, normalize_device +from torchref.config import get_float_dtype, normalize_device from torchref.symmetry import Cell, SpaceGroup, SpaceGroupLike from .base import CrystalDataset diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index f5131265..58722edc 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -9,7 +9,7 @@ import warnings from dataclasses import dataclass, field from pathlib import Path -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union import numpy as np import pandas as pd @@ -17,7 +17,7 @@ from torchref.base import math_torch from torchref.base.french_wilson import FrenchWilson -from torchref.config import dtypes, get_default_device, normalize_device +from torchref.config import dtypes, normalize_device from torchref.io import cif, mtz from torchref.io.datasets.base import CrystalDataset from torchref.symmetry import Cell, SpaceGroup @@ -87,13 +87,11 @@ def n(self) -> int: # Subset reads dispatch through the parent observation attributes. @property def F(self) -> torch.Tensor: - F_corr, _ = self._parent._corrected_or_raw() - return F_corr.index_select(0, self.indices) + return self._parent.F.index_select(0, self.indices) @property def sigF(self) -> torch.Tensor: - _, sig_corr = self._parent._corrected_or_raw() - return sig_corr.index_select(0, self.indices) + return self._parent.F_sigma.index_select(0, self.indices) # -- raw (uncorrected) amplitudes ------------------------------------- @property @@ -121,13 +119,13 @@ def I(self) -> torch.Tensor: # noqa: E743 - crystallographic name Corrected, like :attr:`F` -- both the anisotropy factor and the overall scale enter squared. Use :attr:`I_raw` for the unscaled values. """ - I_corr, _ = self._parent._corrected_or_raw_intensities() + I_corr = self._parent.I return I_corr.index_select(0, self.indices) if I_corr is not None else None @property def sigI(self): """Scaled intensity sigmas, or None. See :attr:`I`.""" - _, sig_corr = self._parent._corrected_or_raw_intensities() + sig_corr = self._parent.I_sigma return sig_corr.index_select(0, self.indices) if sig_corr is not None else None @property @@ -341,10 +339,6 @@ def _subset_indices(self, kind: str) -> torch.Tensor: self._subset_fp = fp return self._subset_cache[kind] - def _corrected_or_raw(self) -> Tuple[torch.Tensor, torch.Tensor]: - """Return the observations exposed by this dataset, without caching.""" - return self.get_corrected_data() - @property def F_raw(self) -> Optional[torch.Tensor]: """Measured amplitudes, shape (N,), in the input amplitude units.""" @@ -3472,10 +3466,6 @@ def get_corrected_intensities(self) -> Tuple[torch.Tensor, torch.Tensor]: raise ValueError("No intensities on this dataset (I/SIGI required)") return self.I, self.I_sigma - def _corrected_or_raw_intensities(self): - """Return intensity observations, or (None, None) when absent.""" - return self.I, self.I_sigma - def generate_validation_set( self, val_fraction_of_free: float = 0.5, diff --git a/torchref/io/hkl.py b/torchref/io/hkl.py index 17dd648e..50864b15 100644 --- a/torchref/io/hkl.py +++ b/torchref/io/hkl.py @@ -63,7 +63,6 @@ def read( if self.verbose > 1: print(f"Reading CrystFEL hkl file: {filepath}") - # Normalize cell → (6,) np.ndarray if hasattr(cell, "data"): # torchref.symmetry.Cell cell = cell.data if hasattr(cell, "detach"): # torch.Tensor @@ -75,7 +74,6 @@ def read( ) self.cell = cell_arr - # Normalize spacegroup → HM-name string if isinstance(spacegroup, str): self.spacegroup = spacegroup elif hasattr(spacegroup, "hm"): # torchref.symmetry.SpaceGroup @@ -87,7 +85,6 @@ def read( f"Cannot normalize spacegroup of type {type(spacegroup)}" ) - # Parse reflection rows h_list, k_list, l_list, I_list, sig_list, n_list = [], [], [], [], [], [] in_header = True with open(filepath) as f: @@ -98,7 +95,6 @@ def read( continue s = line.split() if len(s) < 7 or not s[0].lstrip("-").isdigit(): - # Trailing comment lines or blank lines continue h_list.append(int(s[0])) k_list.append(int(s[1])) diff --git a/torchref/model/model_collection.py b/torchref/model/model_collection.py index e4982bd6..18dc8c96 100644 --- a/torchref/model/model_collection.py +++ b/torchref/model/model_collection.py @@ -34,11 +34,8 @@ if TYPE_CHECKING: from torchref.model.model_ft import ModelFT - from torchref.model.mixed_model import MixedModel -#: Activation fractions are clamped away from 0 and 1 before taking a logit, which -#: would otherwise be infinite. 1e-6 is the same floor the fraction storage has always -#: applied. +#: Keep activation logits finite by clamping fractions away from 0 and 1. _FRACTION_EPS = 1e-6 @@ -100,9 +97,6 @@ def __init__( # route kinetic model predictions directly into the F_calc computation. self._fraction_override: Optional[torch.Tensor] = None - # ------------------------------------------------------------------ - # Properties - # ------------------------------------------------------------------ @property def collection(self) -> "ModelCollection": @@ -163,9 +157,6 @@ def inv_fractional_matrix(self): def fractional_matrix(self): return self.cell.fractional_matrix.to(dtype=self.dtype_float) - # ------------------------------------------------------------------ - # Grid / density helpers (delegate to base models) - # ------------------------------------------------------------------ def setup_grid(self, max_res=None, gridsize=None): for model in self._base_models: @@ -180,9 +171,6 @@ def build_complete_map(self) -> torch.Tensor: density = weighted if density is None else density + weighted return density - # ------------------------------------------------------------------ - # Forward: weighted structure factors - # ------------------------------------------------------------------ def forward(self, hkl: torch.Tensor, recalc: bool = False) -> torch.Tensor: """ @@ -208,9 +196,6 @@ def forward(self, hkl: torch.Tensor, recalc: bool = False) -> torch.Tensor: f_mixed = weighted_f if f_mixed is None else f_mixed + weighted_f return f_mixed - # ------------------------------------------------------------------ - # Freeze / unfreeze - # ------------------------------------------------------------------ def freeze_fractions(self): """Freeze the population parameters. @@ -239,9 +224,6 @@ def clear_fraction_override(self): """Remove the fraction override, reverting to the collection's derived row.""" self._fraction_override = None - # ------------------------------------------------------------------ - # Convenience - # ------------------------------------------------------------------ def get_vdw_radii(self): return self._base_models[0].get_vdw_radii() @@ -319,16 +301,9 @@ def __init__( device = resolve_device(*base_models) dtype = base_models[0].dtype_float - # --- population parameters ------------------------------------- - # - # Only the *overall* activation varies from crystal to crystal; the branching - # among excited components is conserved. So the populations factorise as - # - # w(t) = (1 - alpha) * e_ref + alpha * q(t) - # - # with one activation shared across timepoints and a per-timepoint branching - # distribution over the K-1 non-reference components. Both are frozen by - # default: population refinement is opt-in. + # Factor populations as (1 - alpha) * e_ref + alpha * q(t), with one + # shared activation and a branching distribution per timepoint. Refining + # these parameters is opt-in. self._activation_logit = nn.Parameter( torch.tensor(_logit(1e-6), dtype=dtype, device=device), requires_grad=False, @@ -351,10 +326,6 @@ def __init__( f"ModelCollection initialized with {len(base_models)} base models" ) - # ------------------------------------------------------------------ - # Add timepoints - # ------------------------------------------------------------------ - def add_timepoint( self, name: str, @@ -488,10 +459,6 @@ def add_dark( # shared parametrisation, so it owns nothing that could be frozen. return self.add_timepoint(self._dark_key, fractions) - # ------------------------------------------------------------------ - # Class methods - # ------------------------------------------------------------------ - @classmethod def from_kinetics( cls, @@ -535,10 +502,6 @@ def from_kinetics( return collection - # ------------------------------------------------------------------ - # IHM I/O - # ------------------------------------------------------------------ - @classmethod def from_ihm( cls, @@ -599,10 +562,6 @@ def write_ihm(self, filepath: str, mapping=None, datasets=None) -> None: ) writer.write(filepath) - # ------------------------------------------------------------------ - # Dict-like access - # ------------------------------------------------------------------ - def __getitem__(self, name: str) -> "_SharedMixedModel": return self._timepoints[name] @@ -628,10 +587,6 @@ def items(self) -> List[Tuple[str, "_SharedMixedModel"]]: def get(self, name: str, default=None): return self._timepoints.get(name, default) - # ------------------------------------------------------------------ - # Convenience properties - # ------------------------------------------------------------------ - @property def dark_key(self) -> str: return self._dark_key @@ -667,10 +622,6 @@ def spacegroup(self): def device(self): return self._base_models[0].device - # ------------------------------------------------------------------ - # Fractions inspection - # ------------------------------------------------------------------ - def get_all_fractions(self) -> Dict[str, torch.Tensor]: """Current fractions for each timepoint (including dark).""" return {name: self._timepoints[name].fractions for name in self._order} @@ -682,10 +633,6 @@ def get_fractions_matrix(self) -> torch.Tensor: """ return self.fractions_matrix() - # ------------------------------------------------------------------ - # Population factorisation - # ------------------------------------------------------------------ - @property def alpha_mean(self) -> torch.Tensor: """Mean activation fraction, shared across all timepoints.""" @@ -857,10 +804,6 @@ def set_lambda_twin( self._lambda_logit.requires_grad_(False) return self - # ------------------------------------------------------------------ - # Batched structure factors - # ------------------------------------------------------------------ - def compute_component_fcalcs( self, hkl: torch.Tensor, recalc: bool = False ) -> torch.Tensor: @@ -943,10 +886,6 @@ def compute_all_fcalc( component_fcalcs, self.get_fractions_matrix() ) - # ------------------------------------------------------------------ - # Freeze / unfreeze helpers - # ------------------------------------------------------------------ - def freeze_all_fractions(self): """Exclude the population parameters from optimization. @@ -980,10 +919,6 @@ def unfreeze_structures(self): model.unfreeze("xyz") model.unfreeze("b") - # ------------------------------------------------------------------ - # I/O - # ------------------------------------------------------------------ - def write_pdbs(self, outdir: str): """ Write each base model to a PDB file in *outdir*. @@ -1003,10 +938,6 @@ def write_pdbs(self, outdir: str): if self.verbose > 0: print(f" Wrote {path}") - # ------------------------------------------------------------------ - # Repr - # ------------------------------------------------------------------ - def __repr__(self): tp_names = ", ".join(self._order[:4]) if len(self._order) > 4: diff --git a/torchref/refinement/targets/collection/_specs.py b/torchref/refinement/targets/collection/_specs.py index 84e7084f..7d589cd3 100644 --- a/torchref/refinement/targets/collection/_specs.py +++ b/torchref/refinement/targets/collection/_specs.py @@ -16,22 +16,8 @@ ``ml`` amplitude each dataset absolutely, at a shared Luzzati beta ==================== =========== ============================================== -The two difference rows are both offered rather than one being chosen. Which is better is a -property of a dataset's signal-to-noise -- amplitudes keep the loss in the same space as the -output DED coefficients, intensities avoid the French-Wilson posterior reshaping the weak -tail the signal lives in -- and that is not something to settle once in a library. - -``ml`` is the absolute channel, for the scenarios that are not difference refinement: with -K free base models a purely relative loss leaves the overall level unconstrained. It -replaces a hand-rolled Rice target that set ``beta = sigma_obs**2``, i.e. exactly the -sigma_obs-in-a-Rice-Sigma pairing that -:mod:`torchref.base.targets.xray_likelihoods` documents as never correct and that the -single-dataset table deliberately does not offer. - -There is no intensity ``ml`` row, for the same reason the single-dataset table has none: -Rice and the folded normal are distributions *of an amplitude*, and the intensity analogue -is the exponential / chi-square_1 Wilson distribution -- a different primitive rather than a -different variance. +The absolute ``ml`` channel constrains the overall level when all component models +are free. Rice likelihoods describe amplitudes; intensity rows use Gaussian losses. """ from dataclasses import dataclass, field diff --git a/torchref/refinement/targets/collection/base.py b/torchref/refinement/targets/collection/base.py index 8d2a88e3..23d3ce21 100644 --- a/torchref/refinement/targets/collection/base.py +++ b/torchref/refinement/targets/collection/base.py @@ -1,37 +1,16 @@ -"""Shared base for collection (multi-dataset) X-ray targets. - -:class:`CollectionXrayTarget` gives them the same subset and R-factor contract as -the single-dataset -:class:`~torchref.refinement.targets.xray.base.XrayTarget`: the 3-way ``use_set`` -selector over each member's ``data.work``/``free``/``validation`` accessors, the one -shared :func:`~torchref.base.metrics.rfactor.rfactor_work_free` computed through the -same scaling the loss sees, and the standard ``loss``/``n``/``rwork``/``rfree`` -``stats()`` dict. - -Since every member is expanded onto one common HKL grid, per-dataset R-factors form a -distribution: headline ``rwork``/``rfree`` are its median, with the 10/25/75/90 -percentiles at higher verbosity. - -## The seam - -Same two-part seam as the single-dataset base, batched. :meth:`_loss_inputs` gathers -what a row reads -- observations, model, sigma and mask, each ``(N, n_hkl)`` on the -common grid -- and :meth:`_per_refl` evaluates the likelihood on it *unreduced*. -:meth:`forward` and :meth:`residuals` differ only in whether they sum, so the two cannot -drift into different objectives. - -Every row shares one forward model and declares its ``observable`` (``"amplitude"`` or -``"intensity"``), exactly as the single-dataset table does. The base reads the matching -columns, so no row does its own stacking, masking or sigma flooring -- which is what -three divergent stacking styles and three different sigma-floor constants used to cost. - -Rows that are **cross-dataset coupled** (the difference targets take every dataset -against the mean of all of them) narrow the mask in :meth:`_loss_inputs` so a reflection -counts only if it is in the subset of *every* member, then work on the whole stack inside -:meth:`_per_refl`. That is a mask decision, not a special case in the base. +"""Shared observation access, reduction and reporting for collection X-ray targets. + +Targets declare an amplitude or intensity observable. ``_loss_inputs`` gathers +observations, predictions, uncertainties and masks on the common HKL grid; +``_per_refl`` returns unreduced losses used by both ``forward`` and ``residuals``. +Difference targets intersect member masks so each fitted reflection is present +in every dataset's selected work, free or validation subset. + +R-factors use the same scaled predictions as the loss. Reporting gives the median +across datasets, with the 10/25/75/90 percentiles at higher verbosity. """ -from typing import TYPE_CHECKING, Dict, List, NamedTuple, Optional +from typing import TYPE_CHECKING, Dict, List, NamedTuple import torch @@ -167,10 +146,6 @@ def __init__( self.use_set = use_set self.use_work_set = use_set == "work" - # ------------------------------------------------------------------ - # Dataset / model / subset plumbing - # ------------------------------------------------------------------ - def _keys(self) -> List[str]: """Matched dataset keys this target fits: dark + present timepoints. Targets fitting only part of the collection override it (a target fitting only the @@ -209,19 +184,11 @@ def _scaled_amp_full(self, data, model, recalc: bool = True) -> torch.Tensor: fcalc = data.structure_factors(model, recalc=recalc) return torch.abs(_scale_fcalc(self._scaler, fcalc, model)) - # ------------------------------------------------------------------ - # The per-reflection seam - # ------------------------------------------------------------------ - def _stack_observations(self, keys: List[str]): """``(obs, sigma)``, each ``(N, n_hkl)``, in this row's observable. - Routed through the collection's own batched accessors rather than looping over - ``data.get_corrected_*()`` here: they already apply the inter-dataset scaling (the - intensity factors squared), cache against the ``(log_scale, U_aniso)`` fingerprint, - and name the offending dataset when an intensity column is missing. Reading one - dataset's raw column and another's scaled one is a silent regression that was live - once, when the batched accessors returned raw ``data.F``. + Collection accessors read live dataset views and name any member missing + an intensity column. """ dc = self._dataset_collection if self.observable == "intensity": @@ -242,12 +209,6 @@ def _stack_model(self, keys: List[str], recalc: bool = False) -> torch.Tensor: ) return amp**2 if self.observable == "intensity" else amp - def _stack_masks(self, keys: List[str]) -> torch.Tensor: - """This row's subset mask per dataset, ``(N, n_hkl)``. Validity and the 3-way - work/free/validation selection, with validation carved out of both. - """ - return self._dataset_collection.stack_masks(keys, use_set=self.use_set) - def _sigma_floor(self, sigma: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: """Floor for ``sigma``, at :data:`SIGMA_FLOOR_FRAC` of its median over ``mask``. @@ -277,7 +238,7 @@ def _loss_inputs(self, recalc: bool = False) -> CollectionLossInputs: keys = self._keys() obs, sigma = self._stack_observations(keys) model = self._stack_model(keys, recalc=recalc) - mask = self._stack_masks(keys) + mask = self._dataset_collection.stack_masks(keys, use_set=self.use_set) obs = obs.to(model.dtype) sigma = sigma.to(model.dtype) @@ -332,10 +293,6 @@ def residuals(self) -> torch.Tensor: return torch.zeros((0, len(dc.hkl)), device=dc.hkl.device) return self._per_refl(self._loss_inputs(recalc=True)) - # ------------------------------------------------------------------ - # R-factor reporting (shared source of truth) - # ------------------------------------------------------------------ - def get_rfactor(self) -> Dict[str, object]: """Per-dataset R-work / R-free plus percentile summaries. @@ -384,10 +341,6 @@ def _percentiles(values: List[float]) -> Dict[str, float]: q = torch.quantile(t, torch.tensor(_R_PERCENTILES, dtype=dtype)) return {lbl: q[i].item() for i, lbl in enumerate(_R_PCT_LABELS)} - # ------------------------------------------------------------------ - # Stats - # ------------------------------------------------------------------ - def _n_reflections(self) -> int: """Total reflections in this target's subset across all datasets.""" dc = self._dataset_collection diff --git a/torchref/refinement/targets/collection/intensity.py b/torchref/refinement/targets/collection/intensity.py index 51b6c14c..0469448a 100644 --- a/torchref/refinement/targets/collection/intensity.py +++ b/torchref/refinement/targets/collection/intensity.py @@ -130,9 +130,6 @@ def __init__( f"Supply reflection files with I/SIGI columns." ) - # ------------------------------------------------------------------ - # Forward model - # ------------------------------------------------------------------ def _row_indices(self, keys: List[str]) -> List[int]: """Rows of the collection's fraction matrix corresponding to ``keys``.""" @@ -217,9 +214,6 @@ def _per_refl(self, ctx) -> torch.Tensor: residual, torch.zeros_like(residual), var, var_floor=0.0 ) - # ------------------------------------------------------------------ - # Weight calibration - # ------------------------------------------------------------------ def calibrate_base_weight( self, reference, parameters, ratio: float = 1.0, floor: float = 1e-12 @@ -292,9 +286,6 @@ def _grad_norm(target, scale_out=1.0): ) return self.base_weight - # ------------------------------------------------------------------ - # Reporting - # ------------------------------------------------------------------ def get_rfactor(self) -> Dict[str, object]: """Per-dataset R-work / R-free against the two-moment amplitudes. diff --git a/torchref/refinement/targets/collection/xray.py b/torchref/refinement/targets/collection/xray.py index 0ceaba55..5c399224 100644 --- a/torchref/refinement/targets/collection/xray.py +++ b/torchref/refinement/targets/collection/xray.py @@ -16,16 +16,10 @@ a ``_per_refl`` and nothing else. The selectable set is :data:`~torchref.refinement.targets.collection._specs.COLLECTION_XRAY_TARGETS`. -The retired ``CollectionRiceTarget`` set ``beta = sigma_obs**2``, pairing a measurement -sigma with a Rice ``Sigma``. That asserts an isotropic *complex* error where ``sigma_obs`` -carries no phase at all; :mod:`torchref.base.targets.xray_likelihoods` records that no -regime makes it correct, and the single-dataset table deliberately offers no such row. -:class:`CollectionMLTarget` replaces it. """ from typing import TYPE_CHECKING, Dict -import numpy as np import torch from torchref.base.reciprocal import get_scattering_vectors @@ -37,7 +31,6 @@ from torchref.refinement.model_error_estimation.sigma_a import SigmaAEstimator, epsilon_from_hkl from torchref.utils.stats import VERBOSITY_STANDARD, StatEntry, stat -from ._util import _LOG_2PI, _scale_fcalc from .base import CollectionSigmaALossInputs, CollectionXrayTarget if TYPE_CHECKING: diff --git a/torchref/refinement/targets/xray/observable.py b/torchref/refinement/targets/xray/observable.py index ca08c668..b9d54d74 100644 --- a/torchref/refinement/targets/xray/observable.py +++ b/torchref/refinement/targets/xray/observable.py @@ -1,27 +1,10 @@ -"""The observable axis: which measured column a row fits. - -Every X-ray row shares one forward model -- the scaled complex ``F_calc`` -- and differs -only in what it compares against. Two observables are available: - -* **amplitude** (the default), ``F_obs`` against ``|F_calc|`` -* **intensity**, ``I_obs`` against ``|F_calc|**2`` - -:class:`IntensityObservableMixin` is the whole of the second one. It overrides -:meth:`XrayTarget.get_data` and nothing else, so the likelihood, the mean, the subset -selection, the masks and the R-factor are all inherited untouched. - -**Why intensities are worth a row at all.** ``F_obs`` on a merged dataset is a -French-Wilson posterior, not a measurement: the estimator is strictly positive, so it -reshapes the weak tail and erases negative intensities entirely. Anything whose signal -lives in the *quadratic* part of the data -- an activation second moment, a population -variance -- is fitting a distorted version of the quantity it is trying to measure. Rows -that need that information read ``I_obs`` directly. - -**Why there is no intensity Rice.** Rice and the folded normal are distributions *of an -amplitude*; the intensity analogue is the exponential / chi-square_1 Wilson distribution, -which is a different primitive rather than a different variance. So the intensity axis -carries the Gaussian rows only, and that is a property of the statistics rather than a gap -in the implementation. +"""Select measured intensity observations for Gaussian X-ray targets. + +``IntensityObservableMixin`` reads ``I_obs`` and ``sigma(I)`` and predicts +``|F_calc|**2``, retaining negative intensities that French-Wilson amplitude +conversion reshapes. Likelihoods, subsets and masks come from the target class; +reported R-factors remain in amplitude space. Rice targets describe amplitudes +and therefore have no intensity variant. """ from typing import Tuple diff --git a/torchref/refinement/targets/xray/sigma_a.py b/torchref/refinement/targets/xray/sigma_a.py index 1934b393..b9e4b0a1 100644 --- a/torchref/refinement/targets/xray/sigma_a.py +++ b/torchref/refinement/targets/xray/sigma_a.py @@ -175,10 +175,6 @@ def _loss_inputs( eps_full, dss_full = self._geom() eps_full = eps_full.to(F_calc_full.dtype) dss_full = dss_full.to(F_calc_full.dtype) - # ONE data path. `sub.F` / `sub.sigF` go through `_corrected_or_raw()`, which - # silently falls back to RAW amplitudes when the scaler has not run, while the - # estimator below is fed `get_corrected_data()`, which raises instead. Mixing the two - # can put raw amplitudes and a scaled-data variance in the same loss. F_obs_full, sigma_full = self._data.get_corrected_data() F_obs_full = F_obs_full.to(F_calc_full.dtype).reshape(-1) centric_full = self._data.centric diff --git a/torchref/scaling/collection_scaler.py b/torchref/scaling/collection_scaler.py index de655541..2b7d918d 100644 --- a/torchref/scaling/collection_scaler.py +++ b/torchref/scaling/collection_scaler.py @@ -8,7 +8,7 @@ combination at the same population fractions as the structural models. """ -from typing import TYPE_CHECKING, Dict, List, Optional +from typing import TYPE_CHECKING, Dict import torch import torch.nn as nn @@ -47,17 +47,6 @@ def __init__(self, parent: "CollectionScaler", fractions: torch.Tensor): self._parent = ModuleReference(parent) self.register_buffer("_fractions", fractions.detach().clone()) - @property - def device(self): - """The parent's device. - - Needed explicitly: a target's ``_adopt_device`` reads ``scaler.device``, and - ``nn.Module.__getattr__`` raises before :class:`ModuleReference` gets a chance to - forward it -- the reference lives in ``__dict__``, so attribute lookup on *this* - object never reaches it. - """ - return self._parent.device - def __getattr__(self, name): """Anything this view does not own belongs to the parent. @@ -499,11 +488,8 @@ def refine_lbfgs_joint( ) ) - # One constant applied to every term, so the objective is an exact rescaling. - # ``torch.optim.LBFGS`` converges on ABSOLUTE tolerances, so the objective has to - # be O(1) for them to mean anything; unnormalised it carries the data's own - # magnitude and the float32 ulp of the loss exceeds the decrease the line search - # is trying to resolve. This fit had NO normaliser at all before. + # Use one fixed normalizer: LBFGS uses absolute tolerances, and a large + # float32 loss can round away the decrease sought by the line search. with torch.no_grad(): ssq = sum( float(dc[n].work.F.detach().pow(2).sum()) From bda06b9b1edb30f9c0d2064ee096f23fa43a59bc Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Fri, 18 Sep 2026 10:24:41 +0200 Subject: [PATCH 172/250] Align difference refinement tests with dev conventions --- docs/changelog.rst | 1 + tests/conftest.py | 1 + tests/fixtures/README.md | 1 + tests/fixtures/collections.py | 40 ++ .../test_batched_component_fcalcs.py | 86 ++-- tests/integration/test_cli_two_moment_mtz.py | 14 +- .../test_collection_joint_scale_fit.py | 33 +- .../test_collection_scaler_batched.py | 79 ++-- .../test_collection_stack_accessors.py | 48 +-- ...test_collection_target_characterisation.py | 105 +---- .../io => integration}/test_crystfel_hkl.py | 9 +- .../test_dataset_scaler.py | 152 ++++--- .../test_fcalc_add_noise.py | 98 ++--- .../integration/test_intensity_observable.py | 107 +++++ .../integration/test_two_moment_intensity.py | 314 ++++++++++++++ tests/unit/base/test_intensity_likelihoods.py | 96 +++++ .../model/test_model_collection_fractions.py | 65 ++- .../refinement/test_intensity_observable.py | 189 --------- .../refinement/test_two_moment_intensity.py | 401 ------------------ .../scaling/test_f_sol_override_contract.py | 43 +- 20 files changed, 898 insertions(+), 984 deletions(-) create mode 100644 tests/fixtures/collections.py rename tests/{unit/model => integration}/test_batched_component_fcalcs.py (60%) rename tests/{unit/scaling => integration}/test_collection_joint_scale_fit.py (72%) rename tests/{unit/scaling => integration}/test_collection_scaler_batched.py (59%) rename tests/{unit/io => integration}/test_collection_stack_accessors.py (63%) rename tests/{unit/refinement => integration}/test_collection_target_characterisation.py (61%) rename tests/{unit/io => integration}/test_crystfel_hkl.py (87%) rename tests/{unit/scaling => integration}/test_dataset_scaler.py (68%) rename tests/{unit/io => integration}/test_fcalc_add_noise.py (66%) create mode 100644 tests/integration/test_intensity_observable.py create mode 100644 tests/integration/test_two_moment_intensity.py create mode 100644 tests/unit/base/test_intensity_likelihoods.py delete mode 100644 tests/unit/refinement/test_intensity_observable.py delete mode 100644 tests/unit/refinement/test_two_moment_intensity.py diff --git a/docs/changelog.rst b/docs/changelog.rst index bd1b4ce1..a1cae177 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Align difference-refinement tests with fixture ownership, integration placement and configured dtype/device conventions; share fresh collection setup and exercise noise statistics on deposited amplitudes. - Read scaled observations directly in subset and collection accessors, avoiding unused sigma/amplitude corrections and removing redundant internal forwarding helpers. - Remove one-off diagnostic scripts and consolidate difference-refinement regression tests while retaining numerical and output-format coverage. - Restore the cell and space-group imports needed by difference-density validation setup. diff --git a/tests/conftest.py b/tests/conftest.py index 833cd222..380df84f 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -15,6 +15,7 @@ "tests.fixtures.devices", "tests.fixtures.precision", "tests.fixtures.objects", + "tests.fixtures.collections", ) _HAS_OPENMM = importlib.util.find_spec("openmm") is not None diff --git a/tests/fixtures/README.md b/tests/fixtures/README.md index 282c5527..b146cd38 100644 --- a/tests/fixtures/README.md +++ b/tests/fixtures/README.md @@ -11,6 +11,7 @@ Keep a fixture in its test module when only that module needs it. | `devices.py` | Configured device, explicit backends, device parametrization | All tests; existing per-fixture scopes | | `precision.py` | Comparison tolerances and CPU-double reference context | All tests; reference fixture restores state after each test | | `objects.py` | Mutable models, data, scalers and restraints | All tests; fresh per test except explicitly shared bundles | +| `collections.py` | Paired difference-refinement datasets, models and scalers | All tests; fresh per test | | `numerical.py` | Synthetic tensors and factories | Imported only by `unit/conftest.py`; function | | `functional.py` | Read-only `shared_model_ft` | Imported only by `functional/conftest.py`; module | diff --git a/tests/fixtures/collections.py b/tests/fixtures/collections.py new file mode 100644 index 00000000..c0cf2927 --- /dev/null +++ b/tests/fixtures/collections.py @@ -0,0 +1,40 @@ +"""Build fresh paired collections for difference-refinement integration tests.""" + +import pytest +import torch + + +@pytest.fixture +def difference_models(loaded_reflection_data, sample_structure_pair): + """Return independent dark/light models and raw 1DAW dataset copies per test.""" + from torchref.cli._common import load_model + from torchref.io import DatasetCollection + from torchref.model import ModelCollection + + data = loaded_reflection_data + assert data.I is not None + models = [ + load_model( + str(sample_structure_pair["model"]), + max_res=2.05, + device=data.device, + verbose=0, + ) + for _ in range(2) + ] + with torch.no_grad(): + models[1].xyz.refinable_params += 0.2 + dc = DatasetCollection(device=data.device, verbose=0) + dc.add_dataset("dark", data, set_as_reference=True).add_dataset("light", data) + mc = ModelCollection(models, dark_key="dark", verbose=0) + mc.add_dark().add_timepoint("light", [0.78, 0.22]) + return dc, mc + + +@pytest.fixture +def difference_collection(difference_models): + """Add a fresh initialized model-to-data scaler to the paired collection.""" + from torchref.scaling import CollectionScaler + + dc, mc = difference_models + return dc, mc, CollectionScaler(dc, mc, verbose=0).initialize() diff --git a/tests/unit/model/test_batched_component_fcalcs.py b/tests/integration/test_batched_component_fcalcs.py similarity index 60% rename from tests/unit/model/test_batched_component_fcalcs.py rename to tests/integration/test_batched_component_fcalcs.py index 6bb7fb2d..904e45d9 100644 --- a/tests/unit/model/test_batched_component_fcalcs.py +++ b/tests/integration/test_batched_component_fcalcs.py @@ -3,40 +3,16 @@ import pytest import torch +from torchref.config import get_default_device, get_float_dtype -@pytest.fixture(scope="module") -def pair(pdb_dir, mtz_dir): - """A 2-component, 2-timepoint collection on 1DAW with unequal fractions.""" - pdb = pdb_dir / "1DAW.pdb" - mtz = mtz_dir / "1DAW.mtz" - if not (pdb.exists() and mtz.exists()): - pytest.skip("1DAW fixture not present") +pytestmark = pytest.mark.integration - from torchref import ReflectionData - from torchref.cli._common import load_model - from torchref.io.datasets.collection import DatasetCollection - from torchref.model.model_collection import ModelCollection - - d_min = 2.05 - data = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - model_a = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) - model_b = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) - with torch.no_grad(): - model_b.xyz.refinable_params += 0.2 - - dc = DatasetCollection(verbose=0, device="cpu") - dc.add_dataset("dark", data, set_as_reference=True) - mc = ModelCollection([model_a, model_b], dark_key="dark", verbose=0) - mc.add_dark() - mc.add_timepoint("light", [0.65, 0.35]) - return dc, mc - - -@pytest.mark.integration class TestBatchedMatchesTheLoop: - def test_component_stack_matches_per_model_structure_factors(self, pair): - dc, mc = pair + def test_component_stack_matches_per_model_structure_factors( + self, difference_models + ): + dc, mc = difference_models data = dc["dark"] stacked = dc.component_structure_factors(mc, recalc=True) @@ -48,9 +24,9 @@ def test_component_stack_matches_per_model_structure_factors(self, pair): stacked[k], reference ), f"component {k} differs from data.structure_factors" - def test_mixture_matches_the_per_timepoint_forward(self, pair): - """The whole point: one contraction standing in for T mixed forwards.""" - dc, mc = pair + def test_mixture_matches_the_per_timepoint_forward(self, difference_models): + """A batched contraction agrees with each mixed-model forward.""" + dc, mc = difference_models data = dc["dark"] stacked = dc.component_structure_factors(mc, recalc=True) @@ -63,11 +39,11 @@ def test_mixture_matches_the_per_timepoint_forward(self, pair): mixed[row], reference, rtol=1e-6, atol=1e-6 ), f"timepoint {key!r} differs from its own mixed forward" - def test_compute_all_fcalc_agrees_on_the_signed_index(self, pair): + def test_compute_all_fcalc_agrees_on_the_signed_index(self, difference_models): """``compute_all_fcalc`` takes the caller's indices verbatim, so handed the signed ones it must reproduce the Friedel-corrected mixture up to the conjugation that ``component_structure_factors`` applies.""" - dc, mc = pair + dc, mc = difference_models data = dc["dark"] direct = mc.compute_all_fcalc(data._hkl_for_sf(), recalc=True) @@ -79,17 +55,17 @@ def test_compute_all_fcalc_agrees_on_the_signed_index(self, pair): assert torch.allclose(corrected, mixed, rtol=1e-6, atol=1e-6) -@pytest.fixture(scope="module") -def flagged_pair(pair): - """The same models against data with **manufactured** Friedel-flagged rows.""" +@pytest.fixture +def flagged_pair(difference_models): + """Pair deposited models with both signed Miller-index conventions.""" from torchref import ReflectionData from torchref.io.datasets.collection import DatasetCollection - dc_ref, mc = pair + dc_ref, mc = difference_models src = dc_ref["dark"] hkl = src.hkl.clone() - half = torch.zeros(len(hkl), dtype=torch.bool) + half = torch.zeros(len(hkl), dtype=torch.bool, device=get_default_device()) half[::2] = True hkl[half] = -hkl[half] @@ -100,17 +76,16 @@ def flagged_pair(pair): cell=src.cell, spacegroup=src.spacegroup, rfree_flags=src.rfree_flags.clone(), - device="cpu", + device=get_default_device(), verbose=0, ) assert data.friedel_flags.any() and (~data.friedel_flags).any() - dc = DatasetCollection(verbose=0, device="cpu") + dc = DatasetCollection(verbose=0, device=get_default_device()) dc.add_dataset("dark", data, set_as_reference=True) return dc, mc -@pytest.mark.integration class TestConventionIsNotSkipped: def test_component_stack_is_conjugated_where_flagged(self, flagged_pair): @@ -119,28 +94,31 @@ def test_component_stack_is_conjugated_where_flagged(self, flagged_pair): data = dc["dark"] stacked = dc.component_structure_factors(mc, recalc=True) - naive = mc.compute_component_fcalcs(data.hkl, recalc=False) - assert not torch.allclose(stacked, naive) - for k, model in enumerate(mc.base_models): - reference = data.structure_factors(model, recalc=False) - assert torch.equal(stacked[k], reference), f"component {k} phases differ" + signed = mc.compute_component_fcalcs(data._hkl_for_sf(), recalc=False) + flagged = data.friedel_flags + assert torch.equal(stacked[:, flagged], signed[:, flagged].conj()) + assert torch.equal(stacked[:, ~flagged], signed[:, ~flagged]) + assert not torch.allclose(stacked[:, flagged], signed[:, flagged]) -@pytest.mark.integration class TestContraction: - def test_weights_matrix_is_applied_row_wise(self, pair): + def test_weights_matrix_is_applied_row_wise(self, difference_models): """A transposed einsum would still return the right shape when T == K.""" - dc, mc = pair + dc, mc = difference_models stacked = dc.component_structure_factors(mc, recalc=True) - w = torch.tensor([[1.0, 0.0], [0.0, 1.0]]) + w = torch.tensor( + [[1.0, 0.0], [0.0, 1.0]], + device=get_default_device(), + dtype=get_float_dtype(), + ) mixed = mc.mix_component_fcalcs(stacked, w) assert torch.equal(mixed[0], stacked[0]) assert torch.equal(mixed[1], stacked[1]) - def test_gradient_flows_through_the_contraction(self, pair): - dc, mc = pair + def test_gradient_flows_through_the_contraction(self, difference_models): + dc, mc = difference_models mc.unfreeze_all_fractions() stacked = dc.component_structure_factors(mc, recalc=True) w = mc.get_fractions_matrix() diff --git a/tests/integration/test_cli_two_moment_mtz.py b/tests/integration/test_cli_two_moment_mtz.py index bf40c7a2..9f42aee7 100644 --- a/tests/integration/test_cli_two_moment_mtz.py +++ b/tests/integration/test_cli_two_moment_mtz.py @@ -80,30 +80,28 @@ @pytest.fixture(scope="module") def cli_script(project_root): script = project_root / "torchref" / "cli" / "collection_difference_refine.py" - if not script.exists(): - pytest.skip("difference-refine CLI not found") + assert script.is_file() return script @pytest.fixture(scope="module") def intensity_pair(mtz_dir, pdb_dir, tmp_path_factory): - """A dark/light pair carrying I/SIGI, from the only fixture that has them.""" + """Write a dark/light I/SIGI pair from deposited 1DAW observations.""" import torch from torchref import ReflectionData + from torchref.config import get_int_dtype mtz = mtz_dir / "1DAW.mtz" pdb = pdb_dir / "1DAW.pdb" - if not (mtz.exists() and pdb.exists()): - pytest.skip("1DAW fixture not present") + assert mtz.is_file() and pdb.is_file() data = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - if data.I is None: - pytest.skip("1DAW loaded without intensities") + assert data.I is not None out = tmp_path_factory.mktemp("two_moment_cli") n = len(data) - idx = torch.arange(n) + idx = torch.arange(n, dtype=get_int_dtype(), device=data.device) # Slightly different reflection sets, as a real dark/light pair would be. data.__select__(idx < int(n * 0.97)).write_mtz(str(out / "dark.mtz")) data.__select__(idx >= int(n * 0.03)).write_mtz(str(out / "light.mtz")) diff --git a/tests/unit/scaling/test_collection_joint_scale_fit.py b/tests/integration/test_collection_joint_scale_fit.py similarity index 72% rename from tests/unit/scaling/test_collection_joint_scale_fit.py rename to tests/integration/test_collection_joint_scale_fit.py index bd3a9583..d03c2d01 100644 --- a/tests/unit/scaling/test_collection_joint_scale_fit.py +++ b/tests/integration/test_collection_joint_scale_fit.py @@ -5,27 +5,20 @@ import pytest import torch +pytestmark = pytest.mark.integration -@pytest.fixture(scope="module") -def collection(pdb_dir, mtz_dir): - pdb, mtz = pdb_dir / "1DAW.pdb", mtz_dir / "1DAW.mtz" - if not (pdb.exists() and mtz.exists()): - pytest.skip("1DAW fixture not present") - from torchref import LBFGSRefinement, ReflectionData - from torchref.io.datasets.collection import DatasetCollection - from torchref.model.model_collection import ModelCollection +@pytest.fixture +def collection(loaded_model_ft, loaded_reflection_data): + """One structural model paired with independent dark/timepoint observations.""" + from torchref.io import DatasetCollection + from torchref.model import ModelCollection - ref = LBFGSRefinement(data_file=str(mtz), pdb=str(pdb), verbose=0) - extra = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - - dc = DatasetCollection(verbose=0, device="cpu") - dc.add_dataset("dark", ref.reflection_data, set_as_reference=True) - dc.add_dataset("t1", extra) - - mc = ModelCollection([ref.model], dark_key="dark", verbose=0) - mc.add_dark() - mc.add_timepoint("t1", [1.0]) + data = loaded_reflection_data + dc = DatasetCollection(device=data.device, verbose=0) + dc.add_dataset("dark", data, set_as_reference=True).add_dataset("t1", data) + mc = ModelCollection([loaded_model_ft], dark_key="dark", verbose=0) + mc.add_dark().add_timepoint("t1", [1.0]) return dc, mc @@ -36,7 +29,6 @@ def _fresh_scaler(collection): return CollectionScaler(dc, mc, verbose=0).initialize() -@pytest.mark.unit def test_it_offers_exactly_the_selectable_objectives(): from torchref.scaling.collection_scaler import CollectionScaler from torchref.scaling.scaler_base import DEFAULT_SCALE_TARGET, SCALE_TARGETS @@ -49,14 +41,12 @@ def test_it_offers_exactly_the_selectable_objectives(): assert "nll" in SCALE_TARGETS and "ml_noalpha" in SCALE_TARGETS -@pytest.mark.integration def test_unknown_objective_fails_closed(collection): scaler = _fresh_scaler(collection) with pytest.raises(ValueError, match="scale_target must be one of"): scaler.refine_lbfgs_joint(scale_target="nll_i") -@pytest.mark.integration @pytest.mark.parametrize("scale_target", ["ls", "nll", "ml_noalpha"]) def test_every_objective_fits_finite_parameters(collection, scale_target): """Every selectable row must drive the joint fit to finite parameters.""" @@ -69,7 +59,6 @@ def test_every_objective_fits_finite_parameters(collection, scale_target): assert m["rwork"] and all(0.0 < r < 1.0 for r in m["rwork"]), m["rwork"] -@pytest.mark.integration def test_the_dataset_view_shares_the_parents_parameters(collection): """A row's scaler must be a view, not a copy.""" from torchref.scaling.collection_scaler import _DatasetScalerView diff --git a/tests/unit/scaling/test_collection_scaler_batched.py b/tests/integration/test_collection_scaler_batched.py similarity index 59% rename from tests/unit/scaling/test_collection_scaler_batched.py rename to tests/integration/test_collection_scaler_batched.py index a3943c84..bc6eb50c 100644 --- a/tests/unit/scaling/test_collection_scaler_batched.py +++ b/tests/integration/test_collection_scaler_batched.py @@ -3,55 +3,28 @@ import pytest import torch +from torchref.config import get_default_device, get_float_dtype -@pytest.fixture(scope="module") -def scaled_collection(pdb_dir, mtz_dir): - """A dark/light collection on 1DAW with an initialized shared scaler.""" - pdb = pdb_dir / "1DAW.pdb" - mtz = mtz_dir / "1DAW.mtz" - if not (pdb.exists() and mtz.exists()): - pytest.skip("1DAW fixture not present") - - from torchref import ReflectionData - from torchref.cli._common import load_model - from torchref.io.datasets.collection import DatasetCollection - from torchref.model.model_collection import ModelCollection - from torchref.scaling.collection_scaler import CollectionScaler - - d_min = 2.05 - dark = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - light = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - - model_dark = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) - model_light = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) - with torch.no_grad(): - model_light.xyz.refinable_params += 0.2 - - dc = DatasetCollection(verbose=0, device="cpu") - dc.add_dataset("dark", dark, set_as_reference=True) - dc.add_dataset("light", light) - - mc = ModelCollection([model_dark, model_light], dark_key="dark", verbose=0) - mc.add_dark() - mc.add_timepoint("light", [0.78, 0.22]) - - scaler = CollectionScaler(dc, mc, verbose=0) - scaler.initialize() - return dc, mc, scaler +pytestmark = pytest.mark.integration def _weights(alpha: float) -> torch.Tensor: """Single-row activation weights ``[[1 - a, a]]``.""" - return torch.tensor([[1.0 - alpha, alpha]]) + return torch.tensor( + [[1.0 - alpha, alpha]], device=get_default_device(), dtype=get_float_dtype() + ) -@pytest.mark.integration class TestBatchedMatchesUnbatched: - def test_every_row_matches_forward_mixed(self, scaled_collection): - dc, mc, scaler = scaled_collection + def test_every_row_matches_forward_mixed(self, difference_collection): + dc, mc, scaler = difference_collection components = dc.component_structure_factors(mc, recalc=True) - w = torch.tensor([[1.0, 0.0], [0.78, 0.22], [0.3, 0.7]]) + w = torch.tensor( + [[1.0, 0.0], [0.78, 0.22], [0.3, 0.7]], + device=get_default_device(), + dtype=get_float_dtype(), + ) fcalc = mc.mix_component_fcalcs(components, w) batched = scaler.forward_batched(fcalc, w) @@ -63,9 +36,11 @@ def test_every_row_matches_forward_mixed(self, scaled_collection): batched[i], single, rtol=1e-6, atol=1e-6 ), f"batched row {i} disagrees with forward_mixed" - def test_solvent_stack_rows_are_the_per_component_solvents(self, scaled_collection): + def test_solvent_stack_rows_are_the_per_component_solvents( + self, difference_collection + ): """A transposed or misordered stack would still have the right shape.""" - _, mc, scaler = scaled_collection + _, mc, scaler = difference_collection stack = scaler.compute_component_solvent_raw() assert stack.shape == (mc.n_base_models, len(scaler.hkl)) assert stack.is_complex() @@ -73,18 +48,19 @@ def test_solvent_stack_rows_are_the_per_component_solvents(self, scaled_collecti assert torch.equal(stack[k], scaler._get_component_f_sol_raw(k)) -@pytest.mark.integration class TestAffineInTheMixingWeights: - def test_secant_in_alpha_equals_the_scaled_jacobian(self, scaled_collection): + def test_secant_in_alpha_equals_the_scaled_jacobian(self, difference_collection): """The property the two-moment forward model rests on.""" - dc, mc, scaler = scaled_collection + dc, mc, scaler = difference_collection components = dc.component_structure_factors(mc, recalc=True) solvent = scaler.compute_component_solvent_raw() assert not torch.allclose(solvent[0], solvent[1]) a1, a2 = 0.60, 0.10 w1, w2 = _weights(a1), _weights(a2) - jac = torch.tensor([[-1.0, 1.0]]) # d/da of [1 - a, a] + jac = torch.tensor( + [[-1.0, 1.0]], device=get_default_device(), dtype=get_float_dtype() + ) # d/da of [1 - a, a] s1 = scaler.forward_batched(mc.mix_component_fcalcs(components, w1), w1) s2 = scaler.forward_batched(mc.mix_component_fcalcs(components, w2), w2) @@ -98,10 +74,10 @@ def test_secant_in_alpha_equals_the_scaled_jacobian(self, scaled_collection): tolerance = 16 * torch.finfo(s1.real.dtype).eps * scale assert (secant - expected).abs().max() <= tolerance - def test_scaling_is_linear_in_the_structure_factors(self, scaled_collection): + def test_scaling_is_linear_in_the_structure_factors(self, difference_collection): """The other half of affinity: doubling F_calc at fixed weights doubles the F_calc-dependent part, leaving the solvent offset behind.""" - dc, mc, scaler = scaled_collection + dc, mc, scaler = difference_collection components = dc.component_structure_factors(mc, recalc=True) w = _weights(0.22) fcalc = mc.mix_component_fcalcs(components, w) @@ -116,14 +92,15 @@ def test_scaling_is_linear_in_the_structure_factors(self, scaled_collection): assert rel < 1e-5 -@pytest.mark.integration class TestSolventCacheIsNotPoisoned: - def test_batched_calls_leave_the_cache_alone(self, scaled_collection): + def test_batched_calls_leave_the_cache_alone(self, difference_collection): """Two batched calls with different weights, then a plain one.""" - dc, mc, scaler = scaled_collection + dc, mc, scaler = difference_collection components = dc.component_structure_factors(mc, recalc=True) w = _weights(0.22) - jac = torch.tensor([[-1.0, 1.0]]) + jac = torch.tensor( + [[-1.0, 1.0]], device=get_default_device(), dtype=get_float_dtype() + ) fcalc = mc.mix_component_fcalcs(components, w) before = scaler.forward_mixed(fcalc[0], w[0]).clone() diff --git a/tests/unit/io/test_collection_stack_accessors.py b/tests/integration/test_collection_stack_accessors.py similarity index 63% rename from tests/unit/io/test_collection_stack_accessors.py rename to tests/integration/test_collection_stack_accessors.py index c8a3c265..b4e66412 100644 --- a/tests/unit/io/test_collection_stack_accessors.py +++ b/tests/integration/test_collection_stack_accessors.py @@ -3,28 +3,22 @@ import pytest import torch +pytestmark = pytest.mark.integration -@pytest.fixture(scope="module") -def collection(mtz_dir): - """Two 1DAW datasets for selection and partition checks.""" - mtz = mtz_dir / "1DAW.mtz" - if not mtz.exists(): - pytest.skip("1DAW fixture not present") - from torchref import ReflectionData - from torchref.io.datasets.collection import DatasetCollection +@pytest.fixture +def collection(loaded_reflection_data): + """Two independent observation sets for selection and partition checks.""" + from torchref.io import DatasetCollection - a = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - if a.I is None: - pytest.skip("1DAW loaded without intensities") + data = loaded_reflection_data + return ( + DatasetCollection(device=data.device, verbose=0) + .add_dataset("dark", data, set_as_reference=True) + .add_dataset("light", data) + ) - dc = DatasetCollection(verbose=0, device="cpu") - dc.add_dataset("dark", a, set_as_reference=True) - dc.add_dataset("light", a) - return dc - -@pytest.mark.unit class TestThreeWayMasks: @pytest.mark.parametrize("use_set", ["work", "free", "val"]) def test_rows_match_the_per_dataset_subset(self, collection, use_set): @@ -49,27 +43,21 @@ def test_the_three_subsets_partition_the_valid_reflections(self, collection): def test_a_validation_set_is_carved_out_of_free_not_work(self, collection): """The 3-way behaviour a 2-way flag array cannot reproduce.""" data = collection["dark"] - saved = None if data.validation_flags is None else data.validation_flags.clone() - try: - free_before = int(collection.stack_masks(use_set="free")[0].sum()) - data.generate_validation_set(val_fraction_of_free=0.5, seed=0) + free_before = int(collection.stack_masks(use_set="free")[0].sum()) + data.generate_validation_set(val_fraction_of_free=0.5, seed=0) - free_after = int(collection.stack_masks(use_set="free")[0].sum()) - val_after = int(collection.stack_masks(use_set="val")[0].sum()) + free_after = int(collection.stack_masks(use_set="free")[0].sum()) + val_after = int(collection.stack_masks(use_set="val")[0].sum()) - assert val_after > 0 - assert free_after < free_before - assert free_after + val_after == pytest.approx(free_before, abs=1) - finally: - data.validation_flags = saved - data._subset_fp = None + assert val_after > 0 + assert free_after < free_before + assert free_after + val_after == pytest.approx(free_before, abs=1) def test_an_unknown_subset_name_is_rejected(self, collection): with pytest.raises(ValueError, match="use_set must be"): collection.stack_masks(use_set="test") -@pytest.mark.unit class TestSelectionAndErrors: def test_keys_argument_selects_and_orders_the_rows(self, collection): both = collection.stack_F_obs() diff --git a/tests/unit/refinement/test_collection_target_characterisation.py b/tests/integration/test_collection_target_characterisation.py similarity index 61% rename from tests/unit/refinement/test_collection_target_characterisation.py rename to tests/integration/test_collection_target_characterisation.py index 99d69b3e..7a16dc29 100644 --- a/tests/unit/refinement/test_collection_target_characterisation.py +++ b/tests/integration/test_collection_target_characterisation.py @@ -3,50 +3,19 @@ import pytest import torch -# Allow float32 differences from threaded structure-factor reductions. -LOSS_RTOL = 1e-4 - - -@pytest.fixture(scope="module") -def collection(pdb_dir, mtz_dir): - """``(dc, mc, scaler)`` for a dark/light pair with a real difference in both the - data and the models.""" +from torchref.config import get_default_device, get_float_dtype - pdb = pdb_dir / "1DAW.pdb" - mtz = mtz_dir / "1DAW.mtz" - if not (pdb.exists() and mtz.exists()): - pytest.skip("1DAW fixture not present") +pytestmark = pytest.mark.integration - from torchref import ReflectionData - from torchref.cli._common import load_model - from torchref.io.datasets.collection import DatasetCollection - from torchref.model.model_collection import ModelCollection - from torchref.scaling.collection_scaler import CollectionScaler - - data_dark = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - data_light = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - - # max_res is required: without it the FFT grid setup has no resolution to size from. - d_min = 2.05 - model_dark = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) - model_light = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) - with torch.no_grad(): - # A real displacement, so the calculated difference is not degenerate. - xyz = model_light.xyz.refinable_params - xyz += 0.15 * torch.ones_like(xyz) +# Allow float32 differences from threaded structure-factor reductions. +LOSS_RTOL = 1e-4 - dc = DatasetCollection(verbose=0, device="cpu") - dc.add_dataset("dark", data_dark, set_as_reference=True) - dc.add_dataset("light", data_light) +@pytest.fixture +def collection(difference_collection): + """Install live observation corrections on a fresh collection.""" + dc, mc, scaler = difference_collection dc.scale(nsteps=1) - - mc = ModelCollection([model_dark, model_light], dark_key="dark", verbose=0) - mc.add_dark() - mc.add_timepoint("light", [0.7, 0.3]) - - scaler = CollectionScaler(dc, mc, verbose=0) - scaler.initialize() return dc, mc, scaler @@ -66,7 +35,6 @@ def _targets(dc, mc, scaler): } -@pytest.mark.integration class TestObservedAmplitudesAreScaled: """The loss must move when a dataset's own scale moves.""" @@ -76,15 +44,10 @@ def test_loss_responds_to_the_datasets_own_log_scale(self, collection, name): target = _targets(dc, mc, scaler)[name] before = target.forward().item() - original = dc.scaler.raw_parameters[1, 0].detach().clone() - try: - with torch.no_grad(): - dc.scaler.raw_parameters[1, 0] += 0.25 - target.maintenance() if hasattr(target, "maintenance") else None - after = target.forward().item() - finally: - with torch.no_grad(): - dc.scaler.raw_parameters[1, 0].copy_(original) + with torch.no_grad(): + dc.scaler.raw_parameters[1, 0] += 0.25 + target.maintenance() if hasattr(target, "maintenance") else None + after = target.forward().item() rel = abs(after - before) / abs(before) assert rel > 1e-3, ( @@ -93,7 +56,6 @@ def test_loss_responds_to_the_datasets_own_log_scale(self, collection, name): ) -@pytest.mark.integration class TestSubsetSelectionIsThreeWay: """Work / free / validation, with validation carved out of both.""" @@ -115,15 +77,13 @@ def test_loss_is_restricted_to_the_selected_subset(self, collection, use_set): assert n == expected -@pytest.mark.integration class TestLossesAreSummedNotAveraged: """A summed X-ray term grows with the data; a meaned one does not.""" def test_adding_a_dataset_grows_the_absolute_loss( - self, collection, pdb_dir, mtz_dir + self, collection, loaded_reflection_data ): """Summed loss grows in proportion to the number of datasets.""" - from torchref import ReflectionData from torchref.refinement.targets import CollectionMLTarget dc, mc, scaler = collection @@ -131,35 +91,11 @@ def test_adding_a_dataset_grows_the_absolute_loss( n_before = len(target_before._keys()) one = target_before.forward().item() - extra = ReflectionData(device="cpu", verbose=0).load_mtz( - str(mtz_dir / "1DAW.mtz") - ) - saved = ( - dc._datasets, - list(dc._dataset_order), - dc.hkl, - dc.scaler, - dc.scaling_metrics, - ) - branching = torch.nn.ParameterList(mc._branching_logits) - dc.add_dataset("light2", extra) - mc.add_timepoint("light2", [0.7, 0.3]) - try: - target_after = CollectionMLTarget(dc, mc, scaler=scaler, verbose=0) - n_after = len(target_after._keys()) - two = target_after.forward().item() - finally: - ( - dc._datasets, - dc._dataset_order, - dc._common_hkl, - dc.scaler, - dc.scaling_metrics, - ) = saved - del mc._timepoints["light2"] - mc._order.remove("light2") - mc._branching_rows.pop("light2") - mc._branching_logits = branching + dc.add_dataset("light2", loaded_reflection_data) + mc.add_timepoint("light2", mc["light"].fractions.detach().tolist()) + target_after = CollectionMLTarget(dc, mc, scaler=scaler, verbose=0) + n_after = len(target_after._keys()) + two = target_after.forward().item() assert (n_before, n_after) == (2, 3) ratio = two / one @@ -170,7 +106,6 @@ def test_adding_a_dataset_grows_the_absolute_loss( ) -@pytest.mark.integration class TestReportedNumbers: """The shape of what ``get_rfactor`` / ``stats`` promise, plus reproducibility.""" @@ -180,7 +115,9 @@ def test_forward_is_finite_and_reproducible(self, collection, name): target = _targets(dc, mc, scaler)[name] first = target.forward().item() second = target.forward().item() - assert torch.isfinite(torch.tensor(first)) + assert torch.isfinite( + torch.tensor(first, device=get_default_device(), dtype=get_float_dtype()) + ) assert second == pytest.approx(first, rel=LOSS_RTOL) def test_rfactor_shape_and_range(self, collection): diff --git a/tests/unit/io/test_crystfel_hkl.py b/tests/integration/test_crystfel_hkl.py similarity index 87% rename from tests/unit/io/test_crystfel_hkl.py rename to tests/integration/test_crystfel_hkl.py index c412c245..4dc19c88 100644 --- a/tests/unit/io/test_crystfel_hkl.py +++ b/tests/integration/test_crystfel_hkl.py @@ -4,14 +4,17 @@ import torch from torchref import DatasetCollection, ReflectionData +from torchref.config import get_default_device + +pytestmark = pytest.mark.integration CELL = [14.97, 18.85, 18.89, 89.4, 84.9, 67.8] -@pytest.fixture(scope="module") +@pytest.fixture def halves(test_files_dir): return [ - ReflectionData(device="cpu", verbose=0).load_crystfel_hkl( + ReflectionData(device=get_default_device(), verbose=0).load_crystfel_hkl( str(test_files_dir / "hkl" / f"dark_half{i}.hkl"), cell=CELL, spacegroup="P 1", @@ -41,7 +44,7 @@ def test_partial_overlap_alignment_preserves_sources(halves): original = [data.hkl.clone() for data in halves] sets = [{tuple(row) for row in data.hkl.tolist()} for data in halves] assert sets[0] != sets[1] and sets[0] & sets[1] - dc = DatasetCollection(verbose=0, device="cpu") + dc = DatasetCollection(verbose=0, device=get_default_device()) dc.add_dataset("a", a).add_dataset("b", b) assert len(dc) == len(sets[0] | sets[1]) assert dc.stack_I_obs().shape == (2, len(dc)) diff --git a/tests/unit/scaling/test_dataset_scaler.py b/tests/integration/test_dataset_scaler.py similarity index 68% rename from tests/unit/scaling/test_dataset_scaler.py rename to tests/integration/test_dataset_scaler.py index 17180a33..ec858edc 100644 --- a/tests/unit/scaling/test_dataset_scaler.py +++ b/tests/integration/test_dataset_scaler.py @@ -6,19 +6,18 @@ import pytest import torch -from torchref import DatasetCollection, ReflectionData, ScaledDataset +from torchref import DatasetCollection, ScaledDataset from torchref.base.targets.dataset_scaling import dataset_scaling_loss +from torchref.config import get_default_device, get_int_dtype from torchref.scaling import DatasetScaler - -@pytest.fixture(scope="module") -def deposited(mtz_dir): - """Measured 1DAW amplitudes, intensities, uncertainties and partitions.""" - return ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz_dir / "1DAW.mtz")) +pytestmark = pytest.mark.integration def clone(data): - return data.__select__(torch.arange(len(data), device=data.device)) + return data.__select__( + torch.arange(len(data), device=data.device, dtype=get_int_dtype()) + ) def collection(data, factors): @@ -35,26 +34,30 @@ def collection(data, factors): @pytest.mark.parametrize("factors", [(1, 1), (1, 1000), (1, 2, 8)]) -def test_known_scales_are_centered_and_sources_stay_raw(deposited, factors): - before = deposited.F.clone() - dc = collection(deposited, factors).scale(nsteps=2) - expected = torch.tensor(factors, dtype=dc.scaler.raw_parameters.dtype).log() +def test_known_scales_are_centered_and_sources_stay_raw( + loaded_reflection_data, factors +): + before = loaded_reflection_data.F.clone() + dc = collection(loaded_reflection_data, factors).scale(nsteps=2) + expected = torch.tensor( + factors, dtype=dc.scaler.raw_parameters.dtype, device=get_default_device() + ).log() expected = expected.mean() - expected torch.testing.assert_close( dc.scaler.corrections[:, 0], expected, atol=1e-4, rtol=1e-4 ) assert dc.scaler.corrections[:, 1:].abs().max() < 1e-4 - assert torch.equal(deposited.F, before) - assert not hasattr(deposited, "parameters") - assert not hasattr(deposited, "log_scale") - assert not hasattr(deposited, "U_aniso") + assert torch.equal(loaded_reflection_data.F, before) + assert not hasattr(loaded_reflection_data, "parameters") + assert not hasattr(loaded_reflection_data, "log_scale") + assert not hasattr(loaded_reflection_data, "U_aniso") for key, _ in dc: assert isinstance(dc[key], ScaledDataset) torch.testing.assert_close(dc[key].F, dc["0"].F, rtol=1e-4, atol=1e-4) -def test_known_anisotropy_is_recovered(deposited): - dc = collection(deposited, (1, 1, 1)) +def test_known_anisotropy_is_recovered(loaded_reflection_data): + dc = collection(loaded_reflection_data, (1, 1, 1)) reference = DatasetScaler(dc.datasets) truth = torch.tensor( [ @@ -63,6 +66,7 @@ def test_known_anisotropy_is_recovered(deposited): [-0.1, -0.1, -0.1, -0.05, -0.01, 0.01, -0.02], ], dtype=reference.raw_parameters.dtype, + device=get_default_device(), ) for i, ds in enumerate(dc.values()): distortion = (reference.design(ds.hkl) @ truth[i]).exp() @@ -72,8 +76,8 @@ def test_known_anisotropy_is_recovered(deposited): torch.testing.assert_close(dc.scaler.corrections, truth, atol=2e-3, rtol=2e-3) -def test_live_access_scales_both_sigmas_and_all_entrypoints(deposited): - dc = collection(deposited, (1, 1)).scale(nsteps=1) +def test_live_access_scales_both_sigmas_and_all_entrypoints(loaded_reflection_data): + dc = collection(loaded_reflection_data, (1, 1)).scale(nsteps=1) data = dc["0"] with torch.no_grad(): dc.scaler.raw_parameters[0, 0] = 2 * math.log(2) @@ -111,11 +115,15 @@ def test_live_access_scales_both_sigmas_and_all_entrypoints(deposited): dc.scaler.requires_grad_(False) -def test_selection_copy_alignment_and_independent_collections(deposited): - a = collection(deposited, (1, 4)).scale(nsteps=1) - b = collection(deposited, (1, 9)).scale(nsteps=1) +def test_selection_copy_alignment_and_independent_collections(loaded_reflection_data): + a = collection(loaded_reflection_data, (1, 4)).scale(nsteps=1) + b = collection(loaded_reflection_data, (1, 9)).scale(nsteps=1) view = a["0"] - selected = view.__select__(torch.arange(0, len(view), 3)) + selected = view.__select__( + torch.arange( + 0, len(view), 3, device=get_default_device(), dtype=get_int_dtype() + ) + ) torch.testing.assert_close(selected.F, view.F[::3]) assert selected.scaler is a.scaler assert copy.deepcopy(view).scaler is a.scaler @@ -126,16 +134,23 @@ def test_selection_copy_alignment_and_independent_collections(deposited): scaler = a.scaler a.scale(nsteps=1) assert a.scaler is scaler - a.add_dataset("extra", deposited) + a.add_dataset("extra", loaded_reflection_data) assert a.scaler is None a.scale(nsteps=1) assert a.scaler is not scaler torch.testing.assert_close(view.F, view.F_raw * 2) -def test_two_dataset_loss_matches_propagated_variance_and_gradients(deposited): - f = torch.stack((deposited.F[:128], deposited.F[:128] * 1.1)).double() - sigma = torch.stack((deposited.F_sigma[:128], deposited.F_sigma[:128] * 3)).double() +@pytest.mark.usefixtures("double_cpu") +def test_two_dataset_loss_matches_propagated_variance_and_gradients( + loaded_reflection_data, +): + f = torch.stack( + (loaded_reflection_data.F[:128], loaded_reflection_data.F[:128] * 1.1) + ).double() + sigma = torch.stack( + (loaded_reflection_data.F_sigma[:128], loaded_reflection_data.F_sigma[:128] * 3) + ).double() log_k = torch.zeros_like(f, requires_grad=True) mask = torch.isfinite(f) & torch.isfinite(sigma) & (sigma > 0) actual = dataset_scaling_loss(f, sigma, log_k, mask) @@ -151,9 +166,13 @@ def test_two_dataset_loss_matches_propagated_variance_and_gradients(deposited): torch.testing.assert_close(noisier, actual / 4) -def test_missing_and_invalid_observations_have_finite_gradients(deposited): - f = torch.stack((deposited.F[:128], deposited.F[:128] * 1.1)) - sigma = torch.stack((deposited.F_sigma[:128], deposited.F_sigma[:128])) +def test_missing_and_invalid_observations_have_finite_gradients(loaded_reflection_data): + f = torch.stack( + (loaded_reflection_data.F[:128], loaded_reflection_data.F[:128] * 1.1) + ) + sigma = torch.stack( + (loaded_reflection_data.F_sigma[:128], loaded_reflection_data.F_sigma[:128]) + ) mask = torch.isfinite(f) & torch.isfinite(sigma) & (sigma > 0) mask[:, :10] = False f[:, :10] = float("nan") @@ -166,10 +185,10 @@ def test_missing_and_invalid_observations_have_finite_gradients(deposited): assert torch.equal(log_k.grad[:, :10], torch.zeros_like(log_k.grad[:, :10])) -def test_free_and_validation_changes_do_not_affect_fit(deposited): +def test_free_and_validation_changes_do_not_affect_fit(loaded_reflection_data): results = [] for filler in (3, 900): - dc = collection(deposited, (1, 2)) + dc = collection(loaded_reflection_data, (1, 2)) ds = dc["1"] mismatched = ds.work.indices[:60] ds.rfree_flags[mismatched[:30]] = False @@ -185,12 +204,16 @@ def test_free_and_validation_changes_do_not_affect_fit(deposited): assert torch.equal(*results) -def test_permutation_and_partial_overlap_chain(deposited): - n = len(deposited) +def test_permutation_and_partial_overlap_chain(loaded_reflection_data): + n = len(loaded_reflection_data) sources = { - "a": clone(deposited).__select__(torch.arange(n // 2)), - "b": clone(deposited), - "c": clone(deposited).__select__(torch.arange(n // 2, n)), + "a": clone(loaded_reflection_data).__select__( + torch.arange(n // 2, device=get_default_device(), dtype=get_int_dtype()) + ), + "b": clone(loaded_reflection_data), + "c": clone(loaded_reflection_data).__select__( + torch.arange(n // 2, n, device=get_default_device(), dtype=get_int_dtype()) + ), } sources["b"].F *= 2 sources["b"].F_sigma *= 2 @@ -198,7 +221,7 @@ def test_permutation_and_partial_overlap_chain(deposited): sources["c"].F_sigma *= 4 results = [] for order in [("a", "b", "c"), ("c", "a", "b")]: - dc = DatasetCollection(device="cpu", verbose=0) + dc = DatasetCollection(device=get_default_device(), verbose=0) for key in order: dc.add_dataset(key, sources[key]) dc.scale(nsteps=2) @@ -214,24 +237,30 @@ def test_permutation_and_partial_overlap_chain(deposited): with pytest.raises(ValueError, match="identify"): DatasetScaler( { - "a": deposited.__select__(torch.arange(6)), - "b": deposited.__select__(torch.arange(6)), + "a": loaded_reflection_data.__select__( + torch.arange(6, device=get_default_device(), dtype=get_int_dtype()) + ), + "b": loaded_reflection_data.__select__( + torch.arange(6, device=get_default_device(), dtype=get_int_dtype()) + ), } ) -def test_checkpoint_and_mtz_export_preserve_observations(deposited, tmp_path): +def test_checkpoint_and_mtz_export_preserve_observations( + loaded_reflection_data, tmp_path +): import reciprocalspaceship as rs - dc = collection(deposited, (1, 4)).scale(nsteps=1) + dc = collection(loaded_reflection_data, (1, 4)).scale(nsteps=1) path = tmp_path / "collection.pt" dc.save_state(path) - restored = DatasetCollection.load_state(path, device="cpu") + restored = DatasetCollection.load_state(path, device=get_default_device()) assert restored["0"].scaler is restored["1"].scaler is restored.scaler torch.testing.assert_close(restored.stack_F_obs(), dc.stack_F_obs()) path = tmp_path / "view.pt" dc["0"].save_state(path) - restored_view = ScaledDataset.load_state(path, device="cpu") + restored_view = ScaledDataset.load_state(path, device=get_default_device()) torch.testing.assert_close(restored_view.I, dc["0"].I) path = tmp_path / "scaled.mtz" dc["0"].write_mtz(str(path)) @@ -247,16 +276,23 @@ def test_checkpoint_and_mtz_export_preserve_observations(deposited, tmp_path): ]: np.testing.assert_allclose( exported[column].to_numpy(dtype=float), - getattr(dc["0"], attribute).detach().numpy(), + getattr(dc["0"], attribute).detach().cpu().numpy(), rtol=1e-6, ) -def test_bijvoet_observations_keep_distinct_identities(deposited): +def test_bijvoet_observations_keep_distinct_identities(loaded_reflection_data): """Canonical duplicate HKLs retain separate signed observations and sigmas.""" - data = deposited.__select__(torch.arange(256).repeat_interleave(2)) + data = loaded_reflection_data.__select__( + torch.arange( + 256, device=get_default_device(), dtype=get_int_dtype() + ).repeat_interleave(2) + ) data.friedel_merged = False - data.friedel_flags = torch.arange(len(data)) % 2 == 1 + data.friedel_flags = ( + torch.arange(len(data), device=get_default_device(), dtype=get_int_dtype()) % 2 + == 1 + ) data.hkl_anomalous = torch.where(data.friedel_flags[:, None], -data.hkl, data.hkl) data.F[data.friedel_flags] *= 1.25 dc = collection(data, (1, 4)).scale(nsteps=1) @@ -274,7 +310,7 @@ def test_bijvoet_observations_keep_distinct_identities(deposited): torch.testing.assert_close(selected.F, dc["0"].F.flip(0)) -def test_raw_dataset_excludes_deprecated_interfaces(deposited): +def test_raw_dataset_excludes_deprecated_interfaces(loaded_reflection_data): """Observation storage exposes subset views without optimization methods.""" for name in ( "log_scale", @@ -290,22 +326,26 @@ def test_raw_dataset_excludes_deprecated_interfaces(deposited): "_masked_unpack", "get_mask", ): - assert not hasattr(deposited, name) - assert not callable(deposited) + assert not hasattr(loaded_reflection_data, name) + assert not callable(loaded_reflection_data) -def test_ded_context_consumes_collection_scaled_views(deposited, monkeypatch): +def test_ded_context_consumes_collection_scaled_views( + loaded_reflection_data, monkeypatch +): """DED preparation uses corrected amplitudes after installing scaled members.""" from torchref.cli import validate_ded - dark, light = clone(deposited), clone(deposited) + dark, light = clone(loaded_reflection_data), clone(loaded_reflection_data) light.F *= 4 light.F_sigma *= 4 inputs = {"dark": dark, "light": light} monkeypatch.setattr( validate_ded, "load_reflection_data", lambda path, **kwargs: inputs[path] ) - context = validate_ded.setup_ded_context("dark", "light", dmin=2.2, device="cpu") + context = validate_ded.setup_ded_context( + "dark", "light", dmin=2.2, device=get_default_device() + ) dc = context["collection"] assert context["data_dark"] is dc["dark"] assert context["data_light"] is dc["light"] @@ -316,9 +356,9 @@ def test_ded_context_consumes_collection_scaled_views(deposited, monkeypatch): assert relative_difference < 1e-5 -def test_missing_intensity_access(deposited): +def test_missing_intensity_access(loaded_reflection_data): """Raw and scaled views expose missing columns and reject intensity-only reads.""" - dc = collection(deposited, (1, 2)) + dc = collection(loaded_reflection_data, (1, 2)) dc["0"].I = dc["0"].I_sigma = None raw = dc["0"] dc.scale(nsteps=1) diff --git a/tests/unit/io/test_fcalc_add_noise.py b/tests/integration/test_fcalc_add_noise.py similarity index 66% rename from tests/unit/io/test_fcalc_add_noise.py rename to tests/integration/test_fcalc_add_noise.py index cf2002e3..69326550 100644 --- a/tests/unit/io/test_fcalc_add_noise.py +++ b/tests/integration/test_fcalc_add_noise.py @@ -3,28 +3,28 @@ import pytest import torch +from torchref.config import get_default_device, get_float_dtype, get_int_dtype + +pytestmark = pytest.mark.integration + @pytest.fixture -def fcalc_scene(): - """A small P1 scene with a deliberately wide dynamic range.""" - from torchref.io.datasets import FcalcDataset - - dataset = FcalcDataset.from_cell_and_resolution( - cell=[30.0, 32.0, 34.0, 90.0, 90.0, 90.0], - spacegroup="P 1", - d_min=3.0, - device=torch.device("cpu"), +def fcalc_scene(loaded_reflection_data): + """Use deposited 1DAW amplitudes with deterministic phases for noisy draws.""" + from torchref.io import FcalcDataset + + data = loaded_reflection_data + result = FcalcDataset( + hkl=data.hkl.clone(), + cell=data.cell, + spacegroup=data.spacegroup, + device=data.device, ) - n = len(dataset.hkl) - gen = torch.Generator().manual_seed(17) - # Amplitudes spanning three orders of magnitude, so I spans six. - amp = 10.0 ** (torch.rand(n, generator=gen) * 3.0 - 1.0) - phase = torch.rand(n, generator=gen) * 6.283 - dataset.set_fcalc((amp * torch.exp(1j * phase)).to(torch.complex64)) - return dataset + phase = torch.linspace(-2.0, 2.0, len(data), dtype=data.F.dtype, device=data.device) + result.set_fcalc(data.F * torch.exp(1j * phase)) + return result -@pytest.mark.unit class TestNegativesSurvive: def test_some_intensities_come_out_negative(self, fcalc_scene): noisy = fcalc_scene.add_noise(sigma_mul=0.5, seed=3, verbose=False) @@ -43,7 +43,7 @@ def test_the_amplitude_is_clamped_but_the_intensity_is_not(self, fcalc_scene): # Where the intensity is negative the amplitude is floored at zero, so the two # cannot agree -- which is exactly the information a clamp would have destroyed. assert torch.allclose( - noisy.fcalc_amp[negative], torch.zeros(int(negative.sum())) + noisy.fcalc_amp[negative], torch.zeros_like(noisy.fcalc_amp[negative]) ) @staticmethod @@ -70,7 +70,6 @@ def test_the_intensity_is_unbiased_on_the_weak_reflections(self, fcalc_scene): ) -@pytest.mark.unit class TestSigmaAndHalves: def test_sigma_of_the_mean_is_the_single_draw_sigma_over_root_two( self, fcalc_scene @@ -81,9 +80,10 @@ def test_sigma_of_the_mean_is_the_single_draw_sigma_over_root_two( # not of the draw, so it must be identical across seeds. assert torch.allclose(a.I_sigma, b.I_sigma) - expected = torch.sqrt(torch.tensor(0.2) ** 2 * fcalc_scene.fcalc_amp**4) / ( - 2.0**0.5 - ) + expected = torch.sqrt( + torch.tensor(0.2, device=get_default_device(), dtype=get_float_dtype()) ** 2 + * fcalc_scene.fcalc_amp**4 + ) / (2.0**0.5) assert torch.allclose(a.I_sigma, expected, rtol=1e-5) def test_amplitude_sigma_uses_the_true_amplitude(self, fcalc_scene): @@ -118,43 +118,29 @@ def test_the_source_dataset_is_not_modified(self, fcalc_scene): assert fcalc_scene.I is None -@pytest.mark.unit class TestReferenceDriven: - def test_sigmas_are_grafted_from_the_reference(self, fcalc_scene, mtz_dir): - """The path that matters for real data: per-reflection sigmas from a measured - dataset rather than a parametric model.""" - from torchref import ReflectionData - from torchref.io.datasets import FcalcDataset - - mtz = mtz_dir / "1DAW.mtz" - if not mtz.exists(): - pytest.skip("1DAW fixture not present") - ref = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - if ref.I is None: - pytest.skip("1DAW loaded without intensities") - - # Build on the reference's own HKL list, which is what makes grafting 1:1. - scene = FcalcDataset( - hkl=ref.hkl.clone(), - cell=fcalc_scene.cell, - spacegroup=fcalc_scene.spacegroup, - device=torch.device("cpu"), + def test_sigmas_are_grafted_from_the_reference( + self, fcalc_scene, loaded_reflection_data + ): + """Use the reference's measured uncertainties for both independent draws.""" + noisy = fcalc_scene.add_noise( + reference=loaded_reflection_data, seed=1, verbose=False + ) + torch.testing.assert_close( + noisy.I_sigma, loaded_reflection_data.I_sigma / (2.0**0.5) ) - gen = torch.Generator().manual_seed(2) - amp = torch.rand(len(ref.hkl), generator=gen) * 100.0 - scene.set_fcalc((amp + 0j).to(torch.complex64)) - - noisy = scene.add_noise(reference=ref, seed=1, verbose=False) - assert torch.allclose(noisy.I_sigma, ref.I_sigma / (2.0**0.5)) - - def test_a_mismatched_reference_is_rejected(self, fcalc_scene, mtz_dir): - from torchref import ReflectionData - - mtz = mtz_dir / "1DAW.mtz" - if not mtz.exists(): - pytest.skip("1DAW fixture not present") - ref = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + def test_a_mismatched_reference_is_rejected( + self, fcalc_scene, loaded_reflection_data + ): + ref = loaded_reflection_data.__select__( + torch.arange( + 1, + len(loaded_reflection_data), + device=loaded_reflection_data.device, + dtype=get_int_dtype(), + ) + ) with pytest.raises(ValueError, match="does not match"): fcalc_scene.add_noise(reference=ref, verbose=False) diff --git a/tests/integration/test_intensity_observable.py b/tests/integration/test_intensity_observable.py new file mode 100644 index 00000000..7764ff18 --- /dev/null +++ b/tests/integration/test_intensity_observable.py @@ -0,0 +1,107 @@ +"""Intensity targets consume measured observations and their uncertainties.""" + +import pytest +import torch + +pytestmark = pytest.mark.integration + + +@pytest.fixture +def refinement(sample_structure_pair): + """Fit a fresh 1DAW refinement for each mutable target check.""" + from torchref import LBFGSRefinement + + ref = LBFGSRefinement( + data_file=str(sample_structure_pair["reflections"]), + pdb=str(sample_structure_pair["model"]), + target_mode="ml", + verbose=0, + ) + ref.get_scales() + return ref + + +def _t(refinement, mode, use_set="work"): + from torchref.refinement.targets.xray.factory import create_xray_target + + return create_xray_target( + data=refinement.reflection_data, + model=refinement.model, + scaler=refinement.scaler, + mode=mode, + use_set=use_set, + ) + + +def test_the_row_is_selectable_and_reads_intensities(refinement): + """``nll_i`` comes out of the factory and its ``get_data`` returns the I columns.""" + t = _t(refinement, "nll_i") + obs, calc, sigma, centric, sub = t.get_data() + data = refinement.reflection_data + + torch.testing.assert_close(obs, data.work.I) + torch.testing.assert_close(sigma, data.work.sigI) + # The model is the SQUARED scaled amplitude, not the amplitude. + torch.testing.assert_close(calc, sub.select(t.get_F_calc_scaled(recalc=False) ** 2)) + assert obs.shape == calc.shape == sigma.shape == (sub.n,) + assert centric.shape == (sub.n,) + + +def test_the_intensity_model_is_the_squared_scaled_amplitude(refinement): + """``get_I_calc_scaled`` squares the SCALED amplitude, not the raw one.""" + t = _t(refinement, "nll_i") + with torch.no_grad(): + amp = t.get_F_calc_scaled(recalc=False) + inten = t.get_I_calc_scaled(recalc=False) + torch.testing.assert_close(inten, amp**2, rtol=1e-6, atol=1e-6) + + +def test_rfactors_stay_on_amplitudes(refinement): + """An intensity row reports the SAME R-factors as an amplitude row.""" + r_i = _t(refinement, "nll_i").get_rfactor() + r_a = _t(refinement, "nll").get_rfactor() + assert r_i == pytest.approx(r_a, abs=1e-9) + + +def test_the_loss_is_finite_differentiable_and_summed(refinement): + t = _t(refinement, "nll_i") + loss = t.forward() + assert torch.isfinite(loss) and loss.ndim == 0 + loss.backward() + grads = [ + p.grad + for p in refinement.model.parameters() + if p.requires_grad and p.grad is not None + ] + assert grads, "no gradient reached the model" + assert all(torch.isfinite(g).all() for g in grads) + refinement.model.zero_grad(set_to_none=True) + + +@pytest.mark.parametrize("use_set", ["work", "free"]) +def test_a_reflections_residual_does_not_depend_on_the_arrays_length( + refinement, use_set +): + """``residuals()`` restricted to a subset must equal ``forward()`` on that subset.""" + t = _t(refinement, "nll_i", use_set=use_set) + sub = t._subset() + with torch.no_grad(): + fwd = t.forward() + summed = t.residuals().index_select(0, sub.indices).sum() + torch.testing.assert_close(summed, fwd, rtol=1e-6, atol=1e-6) + + +def test_missing_intensities_raise_at_construction_not_at_forward(refinement): + """LossState probes ``forward()`` at registration, so a missing column has to be + caught in ``__init__`` or it surfaces from deep inside setup with no mention of why. + """ + import copy + + data = copy.copy(refinement.reflection_data) + data.I = None + from torchref.refinement.targets.xray import NLLIntensityXrayTarget + + with pytest.raises(ValueError, match="dataset carries none"): + NLLIntensityXrayTarget( + data=data, model=refinement.model, scaler=refinement.scaler + ) diff --git a/tests/integration/test_two_moment_intensity.py b/tests/integration/test_two_moment_intensity.py new file mode 100644 index 00000000..9a4c2b1e --- /dev/null +++ b/tests/integration/test_two_moment_intensity.py @@ -0,0 +1,314 @@ +"""Two-moment target limits, gradients, invalid observations and weight calibration.""" + +import pytest +import torch + +pytestmark = pytest.mark.integration + + +def _target(dc, mc, scaler, **kw): + from torchref.refinement.targets import CollectionTwoMomentIntensityTarget + + return CollectionTwoMomentIntensityTarget(dc, mc, scaler=scaler, verbose=0, **kw) + + +class TestCoherentLimit: + def test_lambda_zero_reduces_to_the_squared_mean(self, difference_collection): + dc, mc, scaler = difference_collection + mc.set_lambda_twin(0.0) + target = _target(dc, mc, scaler) + + model = target.intensity_model(recalc=True) + + rows = target._row_indices(target._keys()) + weights = mc.fractions_matrix()[rows] + components = dc.component_structure_factors(mc, recalc=False) + mean = scaler.forward_batched( + mc.mix_component_fcalcs(components, weights), weights + ) + assert torch.equal(model, mean.abs() ** 2) + + def test_lambda_zero_skips_nonfinite_derivatives( + self, difference_collection, monkeypatch + ): + """Zero dispersion must not evaluate a potentially non-finite derivative.""" + dc, mc, scaler = difference_collection + mc.set_lambda_twin(0.0) + weights = mc.fractions_matrix() + monkeypatch.setattr(mc, "fractions_matrix", lambda: weights) + + def forbidden(): + raise AssertionError("coherent prediction evaluated its variance branch") + + monkeypatch.setattr(mc, "activation_jacobian", forbidden) + assert torch.isfinite(_target(dc, mc, scaler).forward()) + + def test_full_dispersion_matches_incoherent_intensity(self, difference_collection): + """Fully dark or fully lit crystals mix in intensity, including solvent.""" + dc, mc, scaler = difference_collection + mc.set_lambda_twin(1.0) + target = _target(dc, mc, scaler) + actual = target.intensity_model(recalc=True) + components = dc.component_structure_factors(mc, recalc=False) + basis = torch.eye( + mc.n_base_models, device=components.device, dtype=components.real.dtype + ) + pure = scaler.forward_batched(components, basis) + expected = mc.fractions_matrix() @ pure.abs().square() + # Complex mixture sums lose relative precision near solvent cancellation. + torch.testing.assert_close(actual, expected, rtol=2e-5, atol=1e-5) + + def test_a_nonzero_lambda_changes_the_prediction(self, difference_collection): + """Anti-vacuity: the variance branch must actually do something.""" + dc, mc, scaler = difference_collection + mc.set_lambda_twin(0.0) + coherent = _target(dc, mc, scaler).intensity_model(recalc=True) + + mc.set_lambda_twin(0.5) + dispersed = _target(dc, mc, scaler).intensity_model(recalc=True) + + assert not torch.allclose(coherent, dispersed) + # Strictly positive: |dF|^2 has no sign. + assert bool((dispersed >= coherent - 1e-6).all()) + + +class TestForwardModelStructure: + def test_the_variance_term_is_sigma_sq_times_the_scaled_jacobian( + self, difference_collection + ): + dc, mc, scaler = difference_collection + mc.set_lambda_twin(0.4) + target = _target(dc, mc, scaler) + total = target.intensity_model(recalc=True) + + rows = target._row_indices(target._keys()) + components = dc.component_structure_factors(mc, recalc=False) + weights = mc.fractions_matrix()[rows] + jacobian = mc.activation_jacobian()[rows] + + mean = scaler.forward_batched( + mc.mix_component_fcalcs(components, weights), weights + ) + deriv = scaler.forward_batched( + mc.mix_component_fcalcs(components, jacobian), jacobian + ) + expected = mean.abs() ** 2 + mc.sigma_alpha_sq * deriv.abs() ** 2 + assert torch.allclose(total, expected, rtol=1e-6) + + def test_the_reference_row_carries_no_variance(self, difference_collection): + """The dark's Jacobian row is exactly zero, so its prediction is coherent + regardless of the dispersion -- a dark dataset holds no activation information. + """ + dc, mc, scaler = difference_collection + keys = _target(dc, mc, scaler)._keys() + assert keys[0] == "dark" + + mc.set_lambda_twin(0.0) + coherent = _target(dc, mc, scaler).intensity_model(recalc=True)[0] + mc.set_lambda_twin(0.9) + dispersed = _target(dc, mc, scaler).intensity_model(recalc=True)[0] + + assert torch.allclose(coherent, dispersed, rtol=1e-6) + + def test_shape_follows_the_fitted_keys(self, difference_collection): + dc, mc, scaler = difference_collection + target = _target(dc, mc, scaler) + model = target.intensity_model(recalc=True) + assert model.shape == (len(target._keys()), len(dc.hkl)) + + +class TestLossAndReporting: + def test_forward_is_finite_and_positive(self, difference_collection): + dc, mc, scaler = difference_collection + loss = _target(dc, mc, scaler).forward() + assert torch.isfinite(loss) + assert loss.numel() == 1 + + def test_gradient_reaches_the_light_model(self, difference_collection): + dc, mc, scaler = difference_collection + target = _target(dc, mc, scaler) + target.forward().backward() + grad = mc.base_models[1].xyz.refinable_params.grad + assert grad is not None and torch.isfinite(grad).all() + assert float(grad.abs().max()) > 0 + + def test_gradient_reaches_the_dispersion_when_refinable( + self, difference_collection + ): + dc, mc, scaler = difference_collection + mc.set_lambda_twin(0.3, refinable=True) + _target(dc, mc, scaler).forward().backward() + grad = mc._lambda_logit.grad + assert grad is not None and torch.isfinite(grad).all() + assert float(grad.abs()) > 0 + + def test_rfactor_uses_the_two_moment_amplitude(self, difference_collection): + dc, mc, scaler = difference_collection + target = _target(dc, mc, scaler) + + rf = target.get_rfactor() + assert set(rf) == {"per_dataset", "rwork_pct", "rfree_pct"} + assert set(rf["per_dataset"]) == set(target._keys()) + for key, (rwork, rfree) in rf["per_dataset"].items(): + assert 0.0 < rwork < 2.0, f"{key}: {rwork}" + assert 0.0 < rfree < 2.0, f"{key}: {rfree}" + + def test_stats_report_the_activation_moments(self, difference_collection): + dc, mc, scaler = difference_collection + mc.set_lambda_twin(0.25) + stats = _target(dc, mc, scaler).stats() + for key in ( + "alpha_mean", + "lambda_twin", + "sigma_alpha_sq", + "alpha_sd", + "dI_frac", + "rwork", + "rfree", + "loss", + ): + assert key in stats, f"missing stat: {key}" + assert stats["alpha_mean"].value == pytest.approx(0.22, abs=1e-4) + assert stats["lambda_twin"].value == pytest.approx(0.25, abs=1e-4) + assert stats["dI_frac"].value > 0.0 + + def test_di_frac_is_zero_in_the_coherent_limit(self, difference_collection): + """The stat that distinguishes "refined to zero" from "never refined".""" + dc, mc, scaler = difference_collection + mc.set_lambda_twin(0.0) + assert _target(dc, mc, scaler).stats()["dI_frac"].value == 0.0 + + @pytest.mark.parametrize("use_set", ["work", "free"]) + def test_subset_selection_is_honoured(self, difference_collection, use_set): + dc, mc, scaler = difference_collection + target = _target(dc, mc, scaler, use_set=use_set) + assert target.use_set == use_set + expected = sum( + (dc[k].work if use_set == "work" else dc[k].free).n for k in target._keys() + ) + assert target._n_reflections() == expected + + +class TestIntensityRequirement: + def test_construction_fails_without_intensities(self, difference_models): + """Reject absent intensity columns before evaluating the loss.""" + from torchref.refinement.targets import CollectionTwoMomentIntensityTarget + + dc, mc = difference_models + for data in dc.values(): + data.I = data.I_sigma = None + with pytest.raises(ValueError, match="I/SIGI"): + CollectionTwoMomentIntensityTarget(dc, mc, verbose=0) + + +class TestNonFiniteObservations: + """Real reflection files carry non-finite intensities, and they must not reach the + gradient.""" + + def test_a_nan_observation_does_not_poison_the_gradient( + self, difference_collection + ): + dc, mc, scaler = difference_collection + data = dc["light"] + with torch.no_grad(): + data.I[5] = float("nan") + data.I[11] = float("inf") + + target = _target(dc, mc, scaler) + loss = target.forward() + assert torch.isfinite(loss), "loss went non-finite" + + loss.backward() + grad = mc.base_models[1].xyz.refinable_params.grad + assert grad is not None + assert torch.isfinite(grad).all(), ( + "non-finite observations reached the gradient; every optimizer step " + "would be rejected and the model would not move" + ) + + def test_a_nan_sigma_does_not_poison_the_gradient(self, difference_collection): + dc, mc, scaler = difference_collection + data = dc["light"] + with torch.no_grad(): + data.I_sigma[7] = float("nan") + + target = _target(dc, mc, scaler) + loss = target.forward() + loss.backward() + grad = mc.base_models[1].xyz.refinable_params.grad + assert torch.isfinite(loss) and torch.isfinite(grad).all() + + def test_the_bad_reflections_are_excluded_not_absorbed(self, difference_collection): + """They must drop out of the sum, not contribute a large finite penalty -- + otherwise the loss depends on how many reflections the file happened to reject. + """ + dc, mc, scaler = difference_collection + data = dc["light"] + target = _target(dc, mc, scaler) + baseline = target.forward().item() + with torch.no_grad(): + data.I[3] = float("nan") + with_nan = _target(dc, mc, scaler).forward().item() + + # One reflection out of tens of thousands: the loss should drop slightly, not + # jump by a penalty term. + assert with_nan <= baseline + assert abs(with_nan - baseline) / baseline < 1e-2 + + +class TestWeightCalibration: + """Intensities are squared amplitudes, so this target's gradient is orders of + magnitude away from the amplitude target beside it. Left uncalibrated it swamps the + geometry restraints and buys R-free by moving the model further than the data + supports.""" + + def test_calibration_equalises_the_gradient_norms(self, difference_collection): + dc, mc, scaler = difference_collection + from torchref.refinement.targets import CollectionDifferenceTarget + + params = [p for p in mc.base_models[1].parameters() if p.requires_grad] + diff = CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0) + target = _target(dc, mc, scaler) + + target.calibrate_base_weight(diff, params) + + def gnorm(t): + g = torch.autograd.grad(t.forward(), params, allow_unused=True) + return sum(float((x**2).sum()) for x in g if x is not None) ** 0.5 + + assert gnorm(target) == pytest.approx(gnorm(diff), rel=0.05) + + def test_the_ratio_argument_scales_the_result(self, difference_collection): + dc, mc, scaler = difference_collection + from torchref.refinement.targets import CollectionDifferenceTarget + + params = [p for p in mc.base_models[1].parameters() if p.requires_grad] + diff = CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0) + + a = _target(dc, mc, scaler) + b = _target(dc, mc, scaler) + wa = a.calibrate_base_weight(diff, params, ratio=1.0) + wb = b.calibrate_base_weight(diff, params, ratio=0.25) + assert wb == pytest.approx(0.25 * wa, rel=1e-3) + + def test_base_weight_scales_the_loss_on_the_work_set(self, difference_collection): + dc, mc, scaler = difference_collection + one = _target(dc, mc, scaler, base_weight=1.0).forward().item() + three = _target(dc, mc, scaler, base_weight=3.0).forward().item() + assert three == pytest.approx(3.0 * one, rel=1e-5) + + def test_the_free_set_value_is_left_unweighted(self, difference_collection): + """The free-set number is a diagnostic and has to stay comparable across + weightings.""" + dc, mc, scaler = difference_collection + one = _target(dc, mc, scaler, use_set="free", base_weight=1.0).forward().item() + five = _target(dc, mc, scaler, use_set="free", base_weight=5.0).forward().item() + assert five == pytest.approx(one, rel=1e-6) + + def test_calibration_needs_refinable_parameters(self, difference_collection): + dc, mc, scaler = difference_collection + from torchref.refinement.targets import CollectionDifferenceTarget + + diff = CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0) + with pytest.raises(ValueError, match="No refinable parameters"): + _target(dc, mc, scaler).calibrate_base_weight(diff, []) diff --git a/tests/unit/base/test_intensity_likelihoods.py b/tests/unit/base/test_intensity_likelihoods.py new file mode 100644 index 00000000..df1353cd --- /dev/null +++ b/tests/unit/base/test_intensity_likelihoods.py @@ -0,0 +1,96 @@ +"""Intensity targets read measured intensities and propagate their uncertainties.""" + +import math + +import pytest +import torch + +from torchref.base.targets.xray_likelihoods import ( + VAR_FLOOR, + amplitude_var_from_sigma_obs, + floor_sigma_obs, + gaussian_per_refl, + intensity_var_from_sigma_obs, + nll_per_refl, +) +from torchref.config import get_default_device, get_float_dtype + + +@pytest.mark.unit +def test_the_shared_gaussian_reproduces_the_amplitude_one_bitwise(): + """``nll_per_refl`` is ``gaussian_per_refl`` on ``|F_calc|``, exactly.""" + options = dict(dtype=get_float_dtype(), device=get_default_device()) + obs = torch.linspace(0.1, 100.0, 5000, **options) + calc = torch.linspace(-100.0, 100.0, 5000, **options) + var = amplitude_var_from_sigma_obs(torch.linspace(0.1, 10.0, 5000, **options)) + assert torch.equal( + nll_per_refl(obs, calc, var), gaussian_per_refl(obs, calc.abs(), var) + ) + + +@pytest.mark.unit +def test_the_absolute_variance_floor_is_opt_out_and_matters(rtol): + """``VAR_FLOOR`` is a distortion, not a safeguard, once the builder has floored + sigma.""" + sigma = torch.full( + (256,), 1e-3, dtype=get_float_dtype(), device=get_default_device() + ) + var = intensity_var_from_sigma_obs(sigma) # 1e-6, comfortably above VAR_FLOOR + obs = torch.zeros(256, dtype=get_float_dtype(), device=get_default_device()) + model = torch.full( + (256,), 1e-4, dtype=get_float_dtype(), device=get_default_device() + ) + assert torch.equal( + gaussian_per_refl(obs, model, var, var_floor=0.0), + gaussian_per_refl(obs, model, var, var_floor=VAR_FLOOR), + ), "the floor must be inert when the variance is above it" + + # Below it, the two differ -- and by a lot, not by an ulp. + tiny = torch.full( + (256,), 1e-6, dtype=get_float_dtype(), device=get_default_device() + ) # var = 1e-12 << VAR_FLOOR + var_tiny = intensity_var_from_sigma_obs(tiny) + free = gaussian_per_refl(obs, model, var_tiny, var_floor=0.0) + clamped = gaussian_per_refl(obs, model, var_tiny, var_floor=VAR_FLOOR) + assert not torch.allclose(free, clamped) + # `free` is the honest one: it uses the variance the builder actually produced. + expected = ( + 0.5 * (1e-4) ** 2 / 1e-12 + 0.5 * math.log(1e-12) + 0.5 * math.log(2 * math.pi) + ) + assert free[0].item() == pytest.approx(expected, rel=rtol) + + +@pytest.mark.unit +def test_the_intensity_sigma_floor_respects_the_fitted_subset(): + """``mask`` restricts the median, because unfitted rows carry filler.""" + # The first 20 fitted rows are BELOW the fitted median's floor, so the floor is what + # they come back as -- which is the only way to observe which median was used. + sigma = torch.cat( + [ + torch.full( + (20,), 0.01, device=get_default_device(), dtype=get_float_dtype() + ), # fitted, and below floor either way + torch.full( + (80,), 10.0, device=get_default_device(), dtype=get_float_dtype() + ), # fitted, sets the fitted median + torch.full( + (900,), 1e6, device=get_default_device(), dtype=get_float_dtype() + ), # NOT fitted: filler + ] + ) + mask = torch.cat( + [ + torch.ones(100, dtype=torch.bool, device=get_default_device()), + torch.zeros(900, dtype=torch.bool, device=get_default_device()), + ] + ) + masked = floor_sigma_obs(sigma, mask, abs_floor=1e-12) + unmasked = floor_sigma_obs(sigma, None, abs_floor=1e-12) + assert masked[:20].min().item() == pytest.approx(1.0) # floor = 10 * 0.1 + assert unmasked[:20].min().item() == pytest.approx( + 1e5 + ) # floor = 1e6 * 0.1, swamped + # An explicit floor overrides the median entirely -- the set-independent path. + assert floor_sigma_obs(sigma, mask, floor=0.5)[:20].min().item() == pytest.approx( + 0.5 + ) diff --git a/tests/unit/model/test_model_collection_fractions.py b/tests/unit/model/test_model_collection_fractions.py index f20c4ac5..0ccc3a8b 100644 --- a/tests/unit/model/test_model_collection_fractions.py +++ b/tests/unit/model/test_model_collection_fractions.py @@ -4,13 +4,22 @@ import torch from torch import nn +from torchref.config import ( + get_complex_dtype, + get_default_device, + get_float_dtype, + get_int_dtype, +) + class _StubModel(nn.Module): """Minimal stand-in for ``ModelFT`` for fraction bookkeeping.""" def __init__(self, seed: int): super().__init__() - self.anchor = nn.Parameter(torch.zeros(1)) + self.anchor = nn.Parameter( + torch.zeros(1, device=get_default_device(), dtype=get_float_dtype()) + ) self._seed = seed @property @@ -28,7 +37,7 @@ def forward(self, hkl, recalc: bool = False): 1, n + 1, dtype=self.anchor.dtype, device=self.anchor.device ) amp = (base * float(self._seed + 1)) + self.anchor - return amp.to(torch.complex64) + return amp.to(get_complex_dtype()) @pytest.fixture @@ -44,7 +53,11 @@ def two_model_collection(): @pytest.fixture def hkl(): - return torch.tensor([[1, 0, 0], [0, 1, 0], [1, 1, 0], [2, 0, 1]]) + return torch.tensor( + [[1, 0, 0], [0, 1, 0], [1, 1, 0], [2, 0, 1]], + device=get_default_device(), + dtype=get_int_dtype(), + ) class TestPopulationFactorisation: @@ -77,10 +90,10 @@ def test_the_reference_row_is_exactly_e_ref(self, two_model_collection): assert dark.fractions[0].item() == 1.0 @pytest.mark.unit - def test_fraction_dtype_and_device_follow_the_base_models(self): + def test_fraction_dtype_and_device_follow_the_base_models(self, any_device): from torchref.model.model_collection import ModelCollection - models = [_StubModel(0), _StubModel(1)] + models = [_StubModel(0).to(any_device), _StubModel(1).to(any_device)] mc = ModelCollection(models, verbose=0) mc.add_timepoint("t", [0.6, 0.4]) assert mc._activation_logit.dtype == models[0].dtype_float @@ -150,7 +163,9 @@ class TestOverride: @pytest.mark.unit def test_override_replaces_fractions_and_clears_back(self, two_model_collection): mixed = two_model_collection["light"] - forced = torch.tensor([0.1, 0.9]) + forced = torch.tensor( + [0.1, 0.9], device=get_default_device(), dtype=get_float_dtype() + ) mixed.set_fraction_override(forced) assert mixed.fractions is forced @@ -168,7 +183,11 @@ def test_override_reaches_the_forward(self, two_model_collection, hkl): mixed = two_model_collection["light"] before = mixed(hkl, recalc=True) - mixed.set_fraction_override(torch.tensor([0.1, 0.9])) + mixed.set_fraction_override( + torch.tensor( + [0.1, 0.9], device=get_default_device(), dtype=get_float_dtype() + ) + ) after = mixed(hkl, recalc=True) assert not torch.allclose(before, after) @@ -177,7 +196,12 @@ def test_override_reaches_the_forward(self, two_model_collection, hkl): def test_override_carries_gradient(self, two_model_collection, hkl): """Gradients must flow through the override to whatever produced it.""" mixed = two_model_collection["light"] - forced = torch.tensor([0.4, 0.6], requires_grad=True) + forced = torch.tensor( + [0.4, 0.6], + requires_grad=True, + device=get_default_device(), + dtype=get_float_dtype(), + ) mixed.set_fraction_override(forced) mixed(hkl, recalc=True).abs().sum().backward() @@ -291,12 +315,16 @@ def test_a_second_timepoint_may_rebranch_at_the_same_activation(self): assert float(mc.alpha_mean) == pytest.approx(0.3, abs=1e-5) assert torch.allclose( mc["early"].fractions, - torch.tensor([0.7, 0.3, 0.0]), + torch.tensor( + [0.7, 0.3, 0.0], device=get_default_device(), dtype=get_float_dtype() + ), atol=1e-5, ) assert torch.allclose( mc["late"].fractions, - torch.tensor([0.7, 0.0, 0.3]), + torch.tensor( + [0.7, 0.0, 0.3], device=get_default_device(), dtype=get_float_dtype() + ), atol=1e-5, ) @@ -335,8 +363,17 @@ def test_rows_sum_to_zero_and_the_reference_row_vanishes( jac = mc.activation_jacobian() assert jac.shape == (len(mc), mc.n_base_models) - assert torch.allclose(jac.sum(dim=1), torch.zeros(len(mc)), atol=1e-6) - assert torch.equal(jac[0], torch.zeros(mc.n_base_models)) + assert torch.allclose( + jac.sum(dim=1), + torch.zeros(len(mc), device=get_default_device(), dtype=get_float_dtype()), + atol=1e-6, + ) + assert torch.equal( + jac[0], + torch.zeros( + mc.n_base_models, device=get_default_device(), dtype=get_float_dtype() + ), + ) assert float(jac[1][0]) == pytest.approx(-1.0) @pytest.mark.unit @@ -344,7 +381,9 @@ def test_fractions_matrix_is_e_ref_plus_alpha_times_the_jacobian( self, two_model_collection ): mc = two_model_collection - e_ref = torch.zeros(mc.n_base_models) + e_ref = torch.zeros( + mc.n_base_models, device=get_default_device(), dtype=get_float_dtype() + ) e_ref[0] = 1.0 expected = e_ref.unsqueeze(0) + mc.alpha_mean * mc.activation_jacobian() assert torch.allclose(mc.fractions_matrix(), expected) diff --git a/tests/unit/refinement/test_intensity_observable.py b/tests/unit/refinement/test_intensity_observable.py deleted file mode 100644 index ebff21b3..00000000 --- a/tests/unit/refinement/test_intensity_observable.py +++ /dev/null @@ -1,189 +0,0 @@ -"""Intensity targets read measured intensities and propagate their uncertainties.""" - -import math - -import pytest -import torch - -from torchref.base.targets.xray_likelihoods import ( - VAR_FLOOR, - amplitude_var_from_sigma_obs, - floor_sigma_obs, - gaussian_per_refl, - intensity_var_from_sigma_obs, - nll_per_refl, -) - - -@pytest.mark.unit -def test_the_shared_gaussian_reproduces_the_amplitude_one_bitwise(): - """``nll_per_refl`` is ``gaussian_per_refl`` on ``|F_calc|``, exactly.""" - for dtype in (torch.float32, torch.float64): - torch.manual_seed(3) - F_obs = torch.rand(5000, dtype=dtype) * 100 - F_calc = torch.randn(5000, dtype=dtype) * 100 - var = amplitude_var_from_sigma_obs(torch.rand(5000, dtype=dtype) * 10) - assert torch.equal( - nll_per_refl(F_obs, F_calc, var), - gaussian_per_refl(F_obs, torch.abs(F_calc), var), - ) - - -@pytest.mark.unit -def test_the_absolute_variance_floor_is_opt_out_and_matters(): - """``VAR_FLOOR`` is a distortion, not a safeguard, once the builder has floored - sigma.""" - sigma = torch.full((256,), 1e-3, dtype=torch.float64) - var = intensity_var_from_sigma_obs(sigma) # 1e-6, comfortably above VAR_FLOOR - obs = torch.zeros(256, dtype=torch.float64) - model = torch.full((256,), 1e-4, dtype=torch.float64) - assert torch.equal( - gaussian_per_refl(obs, model, var, var_floor=0.0), - gaussian_per_refl(obs, model, var, var_floor=VAR_FLOOR), - ), "the floor must be inert when the variance is above it" - - # Below it, the two differ -- and by a lot, not by an ulp. - tiny = torch.full((256,), 1e-6, dtype=torch.float64) # var = 1e-12 << VAR_FLOOR - var_tiny = intensity_var_from_sigma_obs(tiny) - free = gaussian_per_refl(obs, model, var_tiny, var_floor=0.0) - clamped = gaussian_per_refl(obs, model, var_tiny, var_floor=VAR_FLOOR) - assert not torch.allclose(free, clamped) - # `free` is the honest one: it uses the variance the builder actually produced. - expected = ( - 0.5 * (1e-4) ** 2 / 1e-12 + 0.5 * math.log(1e-12) + 0.5 * math.log(2 * math.pi) - ) - assert free[0].item() == pytest.approx(expected, rel=1e-12) - - -@pytest.mark.unit -def test_the_intensity_sigma_floor_respects_the_fitted_subset(): - """``mask`` restricts the median, because unfitted rows carry filler.""" - # The first 20 fitted rows are BELOW the fitted median's floor, so the floor is what - # they come back as -- which is the only way to observe which median was used. - sigma = torch.cat( - [ - torch.full((20,), 0.01), # fitted, and below floor either way - torch.full((80,), 10.0), # fitted, sets the fitted median - torch.full((900,), 1e6), # NOT fitted: filler - ] - ) - mask = torch.cat( - [torch.ones(100, dtype=torch.bool), torch.zeros(900, dtype=torch.bool)] - ) - masked = floor_sigma_obs(sigma, mask, abs_floor=1e-12) - unmasked = floor_sigma_obs(sigma, None, abs_floor=1e-12) - assert masked[:20].min().item() == pytest.approx(1.0) # floor = 10 * 0.1 - assert unmasked[:20].min().item() == pytest.approx( - 1e5 - ) # floor = 1e6 * 0.1, swamped - # An explicit floor overrides the median entirely -- the set-independent path. - assert floor_sigma_obs(sigma, mask, floor=0.5)[:20].min().item() == pytest.approx( - 0.5 - ) - - -@pytest.fixture(scope="module") -def refinement(pdb_dir, mtz_dir): - """A scaled 1DAW refinement -- the only fixture carrying BOTH I/SIGI and FP/SIGFP, - so the only one on which the two observables can be compared at all.""" - pdb = pdb_dir / "1DAW.pdb" - mtz = mtz_dir / "1DAW.mtz" - if not (pdb.exists() and mtz.exists()): - pytest.skip("1DAW fixture not present") - from torchref import LBFGSRefinement - - ref = LBFGSRefinement(data_file=str(mtz), pdb=str(pdb), target_mode="ml", verbose=0) - ref.get_scales() - return ref - - -def _t(refinement, mode, use_set="work"): - from torchref.refinement.targets.xray.factory import create_xray_target - - return create_xray_target( - data=refinement.reflection_data, - model=refinement.model, - scaler=refinement.scaler, - mode=mode, - use_set=use_set, - ) - - -@pytest.mark.integration -def test_the_row_is_selectable_and_reads_intensities(refinement): - """``nll_i`` comes out of the factory and its ``get_data`` returns the I columns.""" - t = _t(refinement, "nll_i") - obs, calc, sigma, centric, sub = t.get_data() - data = refinement.reflection_data - - torch.testing.assert_close(obs, data.work.I) - torch.testing.assert_close(sigma, data.work.sigI) - # The model is the SQUARED scaled amplitude, not the amplitude. - torch.testing.assert_close(calc, sub.select(t.get_F_calc_scaled(recalc=False) ** 2)) - assert obs.shape == calc.shape == sigma.shape == (sub.n,) - assert centric.shape == (sub.n,) - - -@pytest.mark.integration -def test_the_intensity_model_is_the_squared_scaled_amplitude(refinement): - """``get_I_calc_scaled`` squares the SCALED amplitude, not the raw one.""" - t = _t(refinement, "nll_i") - with torch.no_grad(): - amp = t.get_F_calc_scaled(recalc=False) - inten = t.get_I_calc_scaled(recalc=False) - torch.testing.assert_close(inten, amp**2, rtol=1e-6, atol=1e-6) - - -@pytest.mark.integration -def test_rfactors_stay_on_amplitudes(refinement): - """An intensity row reports the SAME R-factors as an amplitude row.""" - r_i = _t(refinement, "nll_i").get_rfactor() - r_a = _t(refinement, "nll").get_rfactor() - assert r_i == pytest.approx(r_a, abs=1e-9) - - -@pytest.mark.integration -def test_the_loss_is_finite_differentiable_and_summed(refinement): - t = _t(refinement, "nll_i") - loss = t.forward() - assert torch.isfinite(loss) and loss.ndim == 0 - loss.backward() - grads = [ - p.grad - for p in refinement.model.parameters() - if p.requires_grad and p.grad is not None - ] - assert grads, "no gradient reached the model" - assert all(torch.isfinite(g).all() for g in grads) - refinement.model.zero_grad(set_to_none=True) - - -@pytest.mark.integration -@pytest.mark.parametrize("use_set", ["work", "free"]) -def test_a_reflections_residual_does_not_depend_on_the_arrays_length( - refinement, use_set -): - """``residuals()`` restricted to a subset must equal ``forward()`` on that subset.""" - t = _t(refinement, "nll_i", use_set=use_set) - sub = t._subset() - with torch.no_grad(): - fwd = t.forward() - summed = t.residuals().index_select(0, sub.indices).sum() - torch.testing.assert_close(summed, fwd, rtol=1e-6, atol=1e-6) - - -@pytest.mark.integration -def test_missing_intensities_raise_at_construction_not_at_forward(refinement): - """LossState probes ``forward()`` at registration, so a missing column has to be - caught in ``__init__`` or it surfaces from deep inside setup with no mention of why. - """ - import copy - - data = copy.copy(refinement.reflection_data) - data.I = None - from torchref.refinement.targets.xray import NLLIntensityXrayTarget - - with pytest.raises(ValueError, match="dataset carries none"): - NLLIntensityXrayTarget( - data=data, model=refinement.model, scaler=refinement.scaler - ) diff --git a/tests/unit/refinement/test_two_moment_intensity.py b/tests/unit/refinement/test_two_moment_intensity.py deleted file mode 100644 index 748ca7f1..00000000 --- a/tests/unit/refinement/test_two_moment_intensity.py +++ /dev/null @@ -1,401 +0,0 @@ -"""Two-moment target limits, gradients, invalid observations and weight calibration.""" - -import pytest -import torch - - -@pytest.fixture(scope="module") -def collection(pdb_dir, mtz_dir): - """A dark/light collection on 1DAW, which is the only fixture with I/SIGI.""" - pdb = pdb_dir / "1DAW.pdb" - mtz = mtz_dir / "1DAW.mtz" - if not (pdb.exists() and mtz.exists()): - pytest.skip("1DAW fixture not present") - - from torchref import ReflectionData - from torchref.cli._common import load_model - from torchref.io.datasets.collection import DatasetCollection - from torchref.model.model_collection import ModelCollection - from torchref.scaling.collection_scaler import CollectionScaler - - d_min = 2.05 - dark = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - light = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - if dark.I is None: - pytest.skip("1DAW loaded without intensities") - - model_dark = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) - model_light = load_model(str(pdb), max_res=d_min, device="cpu", verbose=0) - with torch.no_grad(): - model_light.xyz.refinable_params += 0.2 - - dc = DatasetCollection(verbose=0, device="cpu") - dc.add_dataset("dark", dark, set_as_reference=True) - dc.add_dataset("light", light) - - mc = ModelCollection([model_dark, model_light], dark_key="dark", verbose=0) - mc.add_dark() - mc.add_timepoint("light", [0.78, 0.22]) - - scaler = CollectionScaler(dc, mc, verbose=0) - scaler.initialize() - return dc, mc, scaler - - -def _target(dc, mc, scaler, **kw): - from torchref.refinement.targets import CollectionTwoMomentIntensityTarget - - return CollectionTwoMomentIntensityTarget(dc, mc, scaler=scaler, verbose=0, **kw) - - -@pytest.mark.integration -class TestCoherentLimit: - def test_lambda_zero_reduces_to_the_squared_mean(self, collection): - dc, mc, scaler = collection - mc.set_lambda_twin(0.0) - target = _target(dc, mc, scaler) - - model = target.intensity_model(recalc=True) - - rows = target._row_indices(target._keys()) - weights = mc.fractions_matrix()[rows] - components = dc.component_structure_factors(mc, recalc=False) - mean = scaler.forward_batched( - mc.mix_component_fcalcs(components, weights), weights - ) - assert torch.equal(model, mean.abs() ** 2) - - def test_lambda_zero_skips_nonfinite_derivatives(self, collection, monkeypatch): - """Zero dispersion must not evaluate a potentially non-finite derivative.""" - dc, mc, scaler = collection - mc.set_lambda_twin(0.0) - weights = mc.fractions_matrix() - monkeypatch.setattr(mc, "fractions_matrix", lambda: weights) - - def forbidden(): - raise AssertionError("coherent prediction evaluated its variance branch") - - monkeypatch.setattr(mc, "activation_jacobian", forbidden) - assert torch.isfinite(_target(dc, mc, scaler).forward()) - - def test_full_dispersion_matches_incoherent_intensity(self, collection): - """Fully dark or fully lit crystals mix in intensity, including solvent.""" - dc, mc, scaler = collection - mc.set_lambda_twin(1.0) - try: - target = _target(dc, mc, scaler) - actual = target.intensity_model(recalc=True) - components = dc.component_structure_factors(mc, recalc=False) - basis = torch.eye( - mc.n_base_models, device=components.device, dtype=components.real.dtype - ) - pure = scaler.forward_batched(components, basis) - expected = mc.fractions_matrix() @ pure.abs().square() - # Complex mixture sums lose relative precision near solvent cancellation. - torch.testing.assert_close(actual, expected, rtol=2e-5, atol=1e-5) - finally: - mc.set_lambda_twin(0.0) - - def test_a_nonzero_lambda_changes_the_prediction(self, collection): - """Anti-vacuity: the variance branch must actually do something.""" - dc, mc, scaler = collection - mc.set_lambda_twin(0.0) - coherent = _target(dc, mc, scaler).intensity_model(recalc=True) - - mc.set_lambda_twin(0.5) - try: - dispersed = _target(dc, mc, scaler).intensity_model(recalc=True) - finally: - mc.set_lambda_twin(0.0) - - assert not torch.allclose(coherent, dispersed) - # Strictly positive: |dF|^2 has no sign. - assert bool((dispersed >= coherent - 1e-6).all()) - - -@pytest.mark.integration -class TestForwardModelStructure: - def test_the_variance_term_is_sigma_sq_times_the_scaled_jacobian(self, collection): - dc, mc, scaler = collection - mc.set_lambda_twin(0.4) - try: - target = _target(dc, mc, scaler) - total = target.intensity_model(recalc=True) - - rows = target._row_indices(target._keys()) - components = dc.component_structure_factors(mc, recalc=False) - weights = mc.fractions_matrix()[rows] - jacobian = mc.activation_jacobian()[rows] - - mean = scaler.forward_batched( - mc.mix_component_fcalcs(components, weights), weights - ) - deriv = scaler.forward_batched( - mc.mix_component_fcalcs(components, jacobian), jacobian - ) - expected = mean.abs() ** 2 + mc.sigma_alpha_sq * deriv.abs() ** 2 - assert torch.allclose(total, expected, rtol=1e-6) - finally: - mc.set_lambda_twin(0.0) - - def test_the_reference_row_carries_no_variance(self, collection): - """The dark's Jacobian row is exactly zero, so its prediction is coherent - regardless of the dispersion -- a dark dataset holds no activation information. - """ - dc, mc, scaler = collection - keys = _target(dc, mc, scaler)._keys() - assert keys[0] == "dark" - - mc.set_lambda_twin(0.0) - coherent = _target(dc, mc, scaler).intensity_model(recalc=True)[0] - mc.set_lambda_twin(0.9) - try: - dispersed = _target(dc, mc, scaler).intensity_model(recalc=True)[0] - finally: - mc.set_lambda_twin(0.0) - - assert torch.allclose(coherent, dispersed, rtol=1e-6) - - def test_shape_follows_the_fitted_keys(self, collection): - dc, mc, scaler = collection - target = _target(dc, mc, scaler) - model = target.intensity_model(recalc=True) - assert model.shape == (len(target._keys()), len(dc.hkl)) - - -@pytest.mark.integration -class TestLossAndReporting: - def test_forward_is_finite_and_positive(self, collection): - dc, mc, scaler = collection - loss = _target(dc, mc, scaler).forward() - assert torch.isfinite(loss) - assert loss.numel() == 1 - - def test_gradient_reaches_the_light_model(self, collection): - dc, mc, scaler = collection - target = _target(dc, mc, scaler) - target.forward().backward() - grad = mc.base_models[1].xyz.refinable_params.grad - assert grad is not None and torch.isfinite(grad).all() - assert float(grad.abs().max()) > 0 - - def test_gradient_reaches_the_dispersion_when_refinable(self, collection): - dc, mc, scaler = collection - mc.set_lambda_twin(0.3, refinable=True) - try: - _target(dc, mc, scaler).forward().backward() - grad = mc._lambda_logit.grad - assert grad is not None and torch.isfinite(grad).all() - assert float(grad.abs()) > 0 - finally: - mc._lambda_logit.grad = None - mc.set_lambda_twin(0.0) - - def test_rfactor_uses_the_two_moment_amplitude(self, collection): - dc, mc, scaler = collection - target = _target(dc, mc, scaler) - - rf = target.get_rfactor() - assert set(rf) == {"per_dataset", "rwork_pct", "rfree_pct"} - assert set(rf["per_dataset"]) == set(target._keys()) - for key, (rwork, rfree) in rf["per_dataset"].items(): - assert 0.0 < rwork < 2.0, f"{key}: {rwork}" - assert 0.0 < rfree < 2.0, f"{key}: {rfree}" - - def test_stats_report_the_activation_moments(self, collection): - dc, mc, scaler = collection - mc.set_lambda_twin(0.25) - try: - stats = _target(dc, mc, scaler).stats() - for key in ( - "alpha_mean", - "lambda_twin", - "sigma_alpha_sq", - "alpha_sd", - "dI_frac", - "rwork", - "rfree", - "loss", - ): - assert key in stats, f"missing stat: {key}" - assert stats["alpha_mean"].value == pytest.approx(0.22, abs=1e-4) - assert stats["lambda_twin"].value == pytest.approx(0.25, abs=1e-4) - assert stats["dI_frac"].value > 0.0 - finally: - mc.set_lambda_twin(0.0) - - def test_di_frac_is_zero_in_the_coherent_limit(self, collection): - """The stat that distinguishes "refined to zero" from "never refined".""" - dc, mc, scaler = collection - mc.set_lambda_twin(0.0) - assert _target(dc, mc, scaler).stats()["dI_frac"].value == 0.0 - - @pytest.mark.parametrize("use_set", ["work", "free"]) - def test_subset_selection_is_honoured(self, collection, use_set): - dc, mc, scaler = collection - target = _target(dc, mc, scaler, use_set=use_set) - assert target.use_set == use_set - expected = sum( - (dc[k].work if use_set == "work" else dc[k].free).n for k in target._keys() - ) - assert target._n_reflections() == expected - - -@pytest.mark.integration -class TestIntensityRequirement: - def test_construction_fails_without_intensities(self, pdb_dir, mtz_dir): - """Fails at construction, not inside the first loss evaluation: LossState - probes forward at registration and that traceback is far harder to read.""" - mtz = mtz_dir / "3GR5.mtz" - pdb = pdb_dir / "3GR5.pdb" - if not (mtz.exists() and pdb.exists()): - pytest.skip("3GR5 fixture not present") - - from torchref import ReflectionData - from torchref.cli._common import load_model - from torchref.io.datasets.collection import DatasetCollection - from torchref.model.model_collection import ModelCollection - from torchref.refinement.targets import CollectionTwoMomentIntensityTarget - - data = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) - if data.I is not None: - pytest.skip("3GR5 unexpectedly carries intensities") - - model = load_model(str(pdb), max_res=2.05, device="cpu", verbose=0) - dc = DatasetCollection(verbose=0, device="cpu") - dc.add_dataset("dark", data, set_as_reference=True) - mc = ModelCollection([model], dark_key="dark", verbose=0) - mc.add_dark() - - with pytest.raises(ValueError, match="I/SIGI"): - CollectionTwoMomentIntensityTarget(dc, mc, verbose=0) - - -@pytest.mark.integration -class TestNonFiniteObservations: - """Real reflection files carry non-finite intensities, and they must not reach the - gradient.""" - - def test_a_nan_observation_does_not_poison_the_gradient(self, collection): - dc, mc, scaler = collection - data = dc["light"] - saved = data.I.clone() - try: - with torch.no_grad(): - data.I[5] = float("nan") - data.I[11] = float("inf") - - target = _target(dc, mc, scaler) - loss = target.forward() - assert torch.isfinite(loss), "loss went non-finite" - - loss.backward() - grad = mc.base_models[1].xyz.refinable_params.grad - assert grad is not None - assert torch.isfinite(grad).all(), ( - "non-finite observations reached the gradient; every optimizer step " - "would be rejected and the model would not move" - ) - finally: - with torch.no_grad(): - data.I.copy_(saved) - mc.base_models[1].xyz.refinable_params.grad = None - - def test_a_nan_sigma_does_not_poison_the_gradient(self, collection): - dc, mc, scaler = collection - data = dc["light"] - saved = data.I_sigma.clone() - try: - with torch.no_grad(): - data.I_sigma[7] = float("nan") - - target = _target(dc, mc, scaler) - loss = target.forward() - loss.backward() - grad = mc.base_models[1].xyz.refinable_params.grad - assert torch.isfinite(loss) and torch.isfinite(grad).all() - finally: - with torch.no_grad(): - data.I_sigma.copy_(saved) - mc.base_models[1].xyz.refinable_params.grad = None - - def test_the_bad_reflections_are_excluded_not_absorbed(self, collection): - """They must drop out of the sum, not contribute a large finite penalty -- - otherwise the loss depends on how many reflections the file happened to reject. - """ - dc, mc, scaler = collection - data = dc["light"] - saved = data.I.clone() - target = _target(dc, mc, scaler) - try: - baseline = target.forward().item() - with torch.no_grad(): - data.I[3] = float("nan") - with_nan = _target(dc, mc, scaler).forward().item() - finally: - with torch.no_grad(): - data.I.copy_(saved) - - # One reflection out of tens of thousands: the loss should drop slightly, not - # jump by a penalty term. - assert with_nan <= baseline - assert abs(with_nan - baseline) / baseline < 1e-2 - - -@pytest.mark.integration -class TestWeightCalibration: - """Intensities are squared amplitudes, so this target's gradient is orders of - magnitude away from the amplitude target beside it. Left uncalibrated it swamps the - geometry restraints and buys R-free by moving the model further than the data - supports.""" - - def test_calibration_equalises_the_gradient_norms(self, collection): - dc, mc, scaler = collection - from torchref.refinement.targets import CollectionDifferenceTarget - - params = [p for p in mc.base_models[1].parameters() if p.requires_grad] - diff = CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0) - target = _target(dc, mc, scaler) - - target.calibrate_base_weight(diff, params) - - def gnorm(t): - g = torch.autograd.grad(t.forward(), params, allow_unused=True) - return sum(float((x**2).sum()) for x in g if x is not None) ** 0.5 - - assert gnorm(target) == pytest.approx(gnorm(diff), rel=0.05) - - def test_the_ratio_argument_scales_the_result(self, collection): - dc, mc, scaler = collection - from torchref.refinement.targets import CollectionDifferenceTarget - - params = [p for p in mc.base_models[1].parameters() if p.requires_grad] - diff = CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0) - - a = _target(dc, mc, scaler) - b = _target(dc, mc, scaler) - wa = a.calibrate_base_weight(diff, params, ratio=1.0) - wb = b.calibrate_base_weight(diff, params, ratio=0.25) - assert wb == pytest.approx(0.25 * wa, rel=1e-3) - - def test_base_weight_scales_the_loss_on_the_work_set(self, collection): - dc, mc, scaler = collection - one = _target(dc, mc, scaler, base_weight=1.0).forward().item() - three = _target(dc, mc, scaler, base_weight=3.0).forward().item() - assert three == pytest.approx(3.0 * one, rel=1e-5) - - def test_the_free_set_value_is_left_unweighted(self, collection): - """The free-set number is a diagnostic and has to stay comparable across - weightings.""" - dc, mc, scaler = collection - one = _target(dc, mc, scaler, use_set="free", base_weight=1.0).forward().item() - five = _target(dc, mc, scaler, use_set="free", base_weight=5.0).forward().item() - assert five == pytest.approx(one, rel=1e-6) - - def test_calibration_needs_refinable_parameters(self, collection): - dc, mc, scaler = collection - from torchref.refinement.targets import CollectionDifferenceTarget - - diff = CollectionDifferenceTarget(dc, mc, scaler=scaler, verbose=0) - with pytest.raises(ValueError, match="No refinable parameters"): - _target(dc, mc, scaler).calibrate_base_weight(diff, []) diff --git a/tests/unit/scaling/test_f_sol_override_contract.py b/tests/unit/scaling/test_f_sol_override_contract.py index d8924fac..0178db00 100644 --- a/tests/unit/scaling/test_f_sol_override_contract.py +++ b/tests/unit/scaling/test_f_sol_override_contract.py @@ -5,43 +5,48 @@ import pytest import torch +from torchref.config import get_complex_dtype, get_float_dtype, get_int_dtype + class _StubSolvent: """Minimal stand-in for :class:`SolventModel` on the k_sol/B_sol path.""" optimize_phase = False - def __init__(self, value: float = 1.0): + def __init__(self, device, value: float = 1.0): + self.device = device self.value = value self.n_reads = 0 def k_solvent(self): - return torch.tensor(0.35) + return torch.tensor(0.35, device=self.device, dtype=get_float_dtype()) def damping(self, s_half_sq): return torch.ones_like(s_half_sq) def get_rec_solvent(self, hkl): self.n_reads += 1 - return torch.full((hkl.shape[0],), self.value, dtype=torch.complex64) + return torch.full( + (hkl.shape[0],), self.value, dtype=get_complex_dtype(), device=self.device + ) @pytest.fixture -def scaler_with_stub(): +def scaler_with_stub(any_device): """A bare ``Scaler`` with only the solvent branch live.""" from torchref.scaling.scaler import Scaler n = 6 - scaler = Scaler() + scaler = Scaler(device=any_device) dev = scaler.device - scaler.bins = torch.zeros(n, dtype=torch.long, device=dev) - scaler._s_half_sq = torch.zeros(n, device=dev) + scaler.bins = torch.zeros(n, dtype=get_int_dtype(), device=dev) + scaler._s_half_sq = torch.zeros(n, device=dev, dtype=get_float_dtype()) # The no-override path reads the solvent model at self.hkl, which is a read-only # property over self._data. scaler._data = SimpleNamespace( - hkl=torch.zeros((n, 3), dtype=torch.long, device=dev) + hkl=torch.zeros((n, 3), dtype=get_int_dtype(), device=dev) ) - scaler.solvent = _StubSolvent() + scaler.solvent = _StubSolvent(dev) scaler._f_sol_raw = None return scaler, n, dev @@ -52,14 +57,14 @@ class TestOverrideDoesNotMutateTheCache: def test_a_later_call_without_an_override_sees_the_model_solvent( self, scaler_with_stub ): - """The consequence of the leak, stated in terms a caller can observe.""" + """A solvent override applies only to the call that supplies it.""" scaler, n, dev = scaler_with_stub - fcalc = torch.ones(n, dtype=torch.complex64, device=dev) + fcalc = torch.ones(n, dtype=get_complex_dtype(), device=dev) baseline = scaler.forward(fcalc).clone() # stub solvent, value 1.0 scaler._f_sol_raw = None # as update_solvent() would leave it - far_off = torch.full((n,), 99.0, dtype=torch.complex64, device=dev) + far_off = torch.full((n,), 99.0, dtype=get_complex_dtype(), device=dev) scaler.forward(fcalc, f_sol_override=far_off) after = scaler.forward(fcalc) @@ -78,13 +83,15 @@ def test_each_batch_row_matches_the_unbatched_call(self, scaler_with_stub): t = 3 fcalc = torch.stack( [ - torch.full((n,), float(i + 1), dtype=torch.complex64, device=dev) + torch.full((n,), float(i + 1), dtype=get_complex_dtype(), device=dev) for i in range(t) ] ) override = torch.stack( [ - torch.full((n,), float(10 * (i + 1)), dtype=torch.complex64, device=dev) + torch.full( + (n,), float(10 * (i + 1)), dtype=get_complex_dtype(), device=dev + ) for i in range(t) ] ) @@ -103,12 +110,14 @@ def test_unbatched_override_is_unchanged(self, scaler_with_stub): """The ``(N,)`` solvent path must keep working; it is what every single-dataset caller uses.""" scaler, n, dev = scaler_with_stub - fcalc = torch.ones(n, dtype=torch.complex64, device=dev) - override = torch.full((n,), 2.0, dtype=torch.complex64, device=dev) + fcalc = torch.ones(n, dtype=get_complex_dtype(), device=dev) + override = torch.full((n,), 2.0, dtype=get_complex_dtype(), device=dev) out = scaler.forward(fcalc, f_sol_override=override) assert out.shape == (n,) # fcalc + k_sol * f_sol, with damping == 1 and no aniso/Chebyshev/per-bin B. expected = 1.0 + 0.35 * 2.0 - assert torch.allclose(out.real, torch.full((n,), expected, device=dev)) + assert torch.allclose( + out.real, torch.full((n,), expected, device=dev, dtype=get_float_dtype()) + ) From 36622973e5566cff4dfdee83a14a1d3dd0f9e06d Mon Sep 17 00:00:00 2001 From: Claude Date: Fri, 18 Sep 2026 11:23:09 +0000 Subject: [PATCH 173/250] Take integer and index tensors from the configured int dtype ``get_int_dtype()`` (``TORCHREF_DTYPE_INT``, int32 by default) now covers every integer tensor the package allocates or casts, replacing 205 literal ``torch.long`` / ``int64`` / ``int32`` / ``int8`` dtypes that mostly carried an "indexing requires long" justification. That has not been true since torch 2.0: bracket indexing, ``index_select`` and ``index_add_`` accept int32, and those are the consumers of nearly every converted site. The literal int64 that remains is what torch or the arithmetic forces, and each site's marker names the constraint: - ``scatter_add`` / ``gather`` indices, which need int64 on torch < 2.8; the declared ``torch>=2.4`` floor still covers those releases. - Packed keys -- ``min(i, j) * max_idx + max(i, j)`` pair hashes, the composite HKL sort key and the clustering keys -- which overflow int32, and which ``searchsorted`` needs in the same dtype as the table. - TorchMD-Net's ``Z`` and ``batch`` tensors, an external library contract. Count accumulators that ``scatter_add_`` or ``index_add_`` a ``ones_like`` of their index take that index's dtype instead of a literal, so the source and self dtypes match whatever the index is. ``AtomGraph.implicit_h_count`` brings ``bincount``'s int64 to the configured dtype once instead of widening the template count. The riding-hydrogen pair hash widens its operands to int64 explicitly, since the candidate indices it packs are no longer int64 by construction. The 18 markers the dtype guard flagged on dev sat on the closing line of a multi-line call, where the checker never looks; converting those sites removes the markers altogether. AGENTS.md states the convention and its exceptions, and the anchor-selection test compares against the configured dtype rather than int64. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- AGENTS.md | 5 ++ docs/changelog.rst | 1 + tests/unit/model/test_disorder_field.py | 4 +- .../kernels/cpu/jit_reference.py | 2 +- .../kernels/cpu/variable_radius.py | 13 ++--- .../base/electron_density/map_building.py | 2 +- .../base/electron_density/solvent_mask.py | 4 +- torchref/base/electron_density/voxel_utils.py | 6 +-- torchref/base/french_wilson.py | 2 +- torchref/base/metrics/binwise_scale.py | 2 +- torchref/base/reciprocal/grid_operations.py | 17 +++---- torchref/base/reciprocal/symmetry.py | 6 +-- torchref/base/scattering/scattering_table.py | 7 ++- torchref/cli/mtz2map.py | 4 +- torchref/cli/validate_ded.py | 3 +- .../experimental/alignment/frf/data_mr.py | 19 ++++---- .../experimental/alignment/frf/dense_calc.py | 4 +- .../frf/kernels/cpu/legendre_shell.py | 5 +- .../experimental/alignment/frf/peak_finder.py | 5 +- .../alignment/frf/preprocessing.py | 4 +- .../alignment/frf/sitelist_ang.py | 13 ++--- torchref/experimental/alignment/sh.py | 9 ++-- .../experimental/alignment/translation.py | 8 ++-- .../ensemble/ensemble_amber_kl.py | 3 +- .../ensemble/quasi_crystal_amber.py | 9 ++-- .../experimental/ensemble/wilson_prior.py | 2 +- .../monolithic_refinement/density_scaler.py | 4 +- .../experimental/targets/forcefield_target.py | 2 +- .../targets/sampled_ml_phase_target.py | 5 +- torchref/io/datasets/fcalc_data.py | 4 +- torchref/io/datasets/reflection_data.py | 24 +++++----- torchref/model/disorder_field.py | 10 ++-- torchref/model/model.py | 11 +++-- torchref/model/parameter_wrappers.py | 16 +++---- torchref/model/riding_xyz.py | 37 +++++++------- torchref/model/rigid_xyz.py | 4 +- torchref/refinement/base_refinement.py | 4 +- .../model_error_estimation/sigma_a.py | 6 +-- torchref/refinement/optimizers/curvature.py | 3 +- torchref/refinement/targets/adp/rigid_bond.py | 3 +- torchref/refinement/targets/adp/similarity.py | 3 +- torchref/refinement/targets/difference.py | 5 +- .../refinement/targets/geometry/chiral.py | 3 +- .../refinement/targets/geometry/non_bonded.py | 5 +- torchref/refinement/targets/similarity.py | 13 ++--- torchref/scaling/collection_scaler.py | 6 +-- torchref/scaling/scaler_base.py | 6 +-- torchref/scaling/solvent.py | 8 ++-- torchref/scaling/wilson.py | 4 +- torchref/symmetry/map_symmetry.py | 3 +- torchref/symmetry/reciprocal_symmetry.py | 28 +++++------ torchref/symmetry/symmetry.py | 14 +++--- torchref/topology/atom_graph.py | 19 ++++---- torchref/topology/build.py | 7 +-- torchref/topology/builders.py | 34 ++++++------- torchref/topology/edges.py | 5 +- torchref/topology/hydrogens.py | 18 +++---- torchref/topology/nonbonded.py | 36 +++++++------- torchref/topology/residue_graph.py | 3 +- torchref/topology/restraint_sets.py | 4 +- torchref/topology/restraints.py | 34 ++++++------- torchref/topology/riding.py | 48 +++++++++---------- torchref/topology/topology.py | 11 +++-- 63 files changed, 321 insertions(+), 288 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index d02be91b..7bb4902b 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -51,6 +51,11 @@ Practically: - **Never hardcode a dtype.** Take it from the config: `torchref.config.get_float_dtype()`, `get_int_dtype()`, `get_complex_dtype()`, or from an input tensor. Roughly 200 call sites already do this; follow them. +- Integer and index tensors take `get_int_dtype()` too (int32 by default). Plain indexing, + `index_select` and `index_add_` accept it. A literal int dtype survives only where torch or + the arithmetic forces it, with a `# dtype-ok:` marker naming the constraint: `scatter`/`gather` + indices (int64 on torch < 2.8), `index_copy_`/`index_fill_`/`one_hot` (int64 always), packed + keys such as `i * n + j` that overflow int32, and external-library contracts (TorchMD-Net). - `torch.float64` *is* a supported configuration (`TORCHREF_DTYPE_FLOAT=float64`) used as an eager numerical reference and in gradient checks. Code must **work** in float64, must not **require** it, and must not silently downcast (see `tests/integration/test_dtype_config_float64.py`). diff --git a/docs/changelog.rst b/docs/changelog.rst index a1cae177..58c0b92d 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Integer and index tensors now take the configured int dtype (``get_int_dtype()``, ``TORCHREF_DTYPE_INT``, int32 by default) throughout the package. A hardcoded ``int64`` remains only where a torch op (``scatter``/``gather`` on torch < 2.8, ``index_copy_``), an int32 overflow, or an external library requires it, and each such site says which. - Align difference-refinement tests with fixture ownership, integration placement and configured dtype/device conventions; share fresh collection setup and exercise noise statistics on deposited amplitudes. - Read scaled observations directly in subset and collection accessors, avoiding unused sigma/amplitude corrections and removing redundant internal forwarding helpers. - Remove one-off diagnostic scripts and consolidate difference-refinement regression tests while retaining numerical and output-format coverage. diff --git a/tests/unit/model/test_disorder_field.py b/tests/unit/model/test_disorder_field.py index dcb8f11d..12fc6546 100644 --- a/tests/unit/model/test_disorder_field.py +++ b/tests/unit/model/test_disorder_field.py @@ -14,6 +14,8 @@ import pytest import torch +from torchref.config import get_int_dtype + from torchref.model.disorder_field import ( DisorderFieldTensor, build_neighbor_list, @@ -52,7 +54,7 @@ def test_anchor_selection_is_deterministic(coords): b = farthest_point_anchors(coords, 10) assert torch.equal(a, b) assert a.shape[0] == 10 - assert a.dtype == torch.int64 + assert a.dtype == get_int_dtype() # Anchors are atom indices, and distinct. assert int(a.max()) < coords.shape[0] assert torch.unique(a).shape[0] == a.shape[0] diff --git a/torchref/base/electron_density/kernels/cpu/jit_reference.py b/torchref/base/electron_density/kernels/cpu/jit_reference.py index 257fa6d0..07b6ba0e 100644 --- a/torchref/base/electron_density/kernels/cpu/jit_reference.py +++ b/torchref/base/electron_density/kernels/cpu/jit_reference.py @@ -136,7 +136,7 @@ def forward( ny: int = density_map.shape[1] nz: int = density_map.shape[2] strides = torch.tensor( - [ny * nz, nz, 1], device=voxel_indices.device, dtype=torch.long # dtype-ok: CPU-kernel strides for flat voxel index arithmetic; indexing requires long + [ny * nz, nz, 1], device=voxel_indices.device, dtype=torch.long # dtype-ok: int64 strides make the flat voxel index int64; scatter_add_ requires int64 on torch < 2.8 ) index_flat = torch.sum(voxel_indices.to(torch.long) * strides, dim=-1).view(-1) # dtype-ok: voxel indices flattened for scatter; indexing requires long diff --git a/torchref/base/electron_density/kernels/cpu/variable_radius.py b/torchref/base/electron_density/kernels/cpu/variable_radius.py index bb19ae93..aa3f4a8e 100644 --- a/torchref/base/electron_density/kernels/cpu/variable_radius.py +++ b/torchref/base/electron_density/kernels/cpu/variable_radius.py @@ -28,6 +28,7 @@ import math import torch +from torchref.config import get_int_dtype from torchref.base.electron_density.radius_policy import _u6_to_u3 @@ -52,7 +53,7 @@ def _bucket_by_radius(radius: torch.Tensor, center_1d: torch.Tensor): spans.append((float(r), cursor, cursor + idx.numel())) cursor += idx.numel() order = (torch.cat(order_parts) if order_parts - else torch.zeros(0, dtype=torch.long, device=radius.device)) # dtype-ok: empty voxel-index fallback; must stay long for indexing + else torch.zeros(0, dtype=get_int_dtype(), device=radius.device)) return order, spans @@ -88,7 +89,7 @@ def _canonical_setup(xyz, inv_frac, frac, grid_dims, radius_per_atom, dtype): nx, ny, nz = grid_dims grid_f = torch.tensor(grid_dims, device=device, dtype=dtype) xyz_frac = (xyz @ inv_frac.T) % 1.0 - center_idx = torch.round(xyz_frac * grid_f).to(torch.long) # dtype-ok: rounded voxel center indices; torch indexing requires long + center_idx = torch.round(xyz_frac * grid_f).to(get_int_dtype()) # w0: atom position relative to its anchor node, in Cartesian. This is what # centres the sphere on the atom rather than on the node. w0 = (xyz_frac - center_idx.to(dtype) / grid_f) @ frac.T @@ -111,8 +112,8 @@ def add_isotropic_plain_var(density_map, xyz, adp, occ, A, B, device, dtype = xyz.device, density_map.dtype nx, ny, nz = (int(s) for s in density_map.shape) grid_dims = (nx, ny, nz) - strides = torch.tensor([ny * nz, nz, 1], device=device, dtype=torch.long) # dtype-ok: strides for flat voxel-index arithmetic; indexing requires long - grid_shape = torch.tensor(grid_dims, device=device, dtype=torch.long) # dtype-ok: grid_shape for flat voxel-index arithmetic; indexing requires long + strides = torch.tensor([ny * nz, nz, 1], device=device, dtype=torch.long) # dtype-ok: int64 strides make the flat voxel index int64; scatter_add requires int64 on torch < 2.8 + grid_shape = torch.tensor(grid_dims, device=device, dtype=get_int_dtype()) order, spans, center_idx, w0 = _canonical_setup( xyz, inv_frac_matrix, frac_matrix, grid_dims, radius_per_atom, dtype) @@ -151,8 +152,8 @@ def add_anisotropic_plain_var(density_map, xyz, u, occ, A, B, device, dtype = xyz.device, density_map.dtype nx, ny, nz = (int(s) for s in density_map.shape) grid_dims = (nx, ny, nz) - strides = torch.tensor([ny * nz, nz, 1], device=device, dtype=torch.long) # dtype-ok: strides for flat voxel-index arithmetic; indexing requires long - grid_shape = torch.tensor(grid_dims, device=device, dtype=torch.long) # dtype-ok: grid_shape for flat voxel-index arithmetic; indexing requires long + strides = torch.tensor([ny * nz, nz, 1], device=device, dtype=torch.long) # dtype-ok: int64 strides make the flat voxel index int64; scatter_add requires int64 on torch < 2.8 + grid_shape = torch.tensor(grid_dims, device=device, dtype=get_int_dtype()) order, spans, center_idx, w0 = _canonical_setup( xyz, inv_frac_matrix, frac_matrix, grid_dims, radius_per_atom, dtype) diff --git a/torchref/base/electron_density/map_building.py b/torchref/base/electron_density/map_building.py index d00d7664..8840cebc 100644 --- a/torchref/base/electron_density/map_building.py +++ b/torchref/base/electron_density/map_building.py @@ -28,7 +28,7 @@ def scatter_add_nd(source, index, map): """Vectorized n-dimensional scatter-add: ``source`` ``(N,)`` into ``map`` ``(d1..dn)`` at ``index`` ``(N, ndim)``, returning the modified map. """ - map_shape = torch.tensor(map.shape, device=index.device, dtype=torch.int64) # dtype-ok: map_shape for stride/flat-index arithmetic feeding scatter_add; requires int64 + map_shape = torch.tensor(map.shape, device=index.device, dtype=torch.int64) # dtype-ok: int64 shape/strides make the flat index int64; scatter_add_ requires int64 on torch < 2.8 # Convert n-dimensional indices to flat indices # For shape (d1, d2, d3, ..., dn), flat_index = i0 * (d1*d2*...*dn) + i1 * (d2*d3*...*dn) + ... + in diff --git a/torchref/base/electron_density/solvent_mask.py b/torchref/base/electron_density/solvent_mask.py index 2e36e709..018e39a3 100644 --- a/torchref/base/electron_density/solvent_mask.py +++ b/torchref/base/electron_density/solvent_mask.py @@ -8,7 +8,7 @@ import numpy as np import torch -from torchref.config import dtypes +from torchref.config import dtypes, get_int_dtype from torchref.base.coordinates.periodic_boundary import smallest_diff from .map_building import scatter_add_nd @@ -113,7 +113,7 @@ def add_to_phenix_mask( ) # (N_atoms, N_voxels) # Flatten for scatter operations - voxel_indices_flat = voxel_indices.reshape(-1, 3).to(torch.long) # dtype-ok: voxel indices for grid indexing; requires long + voxel_indices_flat = voxel_indices.reshape(-1, 3).to(get_int_dtype()) # Create protein core mask using scatter_add int_dtype = dtypes.int diff --git a/torchref/base/electron_density/voxel_utils.py b/torchref/base/electron_density/voxel_utils.py index 3fc903d7..1551c255 100644 --- a/torchref/base/electron_density/voxel_utils.py +++ b/torchref/base/electron_density/voxel_utils.py @@ -6,7 +6,7 @@ import torch -from torchref.config import dtypes +from torchref.config import dtypes, get_int_dtype def find_relevant_voxels(real_space_grid, xyz, radius_angstrom=4, inv_frac_matrix=None): @@ -49,13 +49,13 @@ def find_relevant_voxels(real_space_grid, xyz, radius_angstrom=4, inv_frac_matri # This ensures atoms outside the unit cell are correctly wrapped xyz_frac = torch.matmul(inv_frac_matrix, xyz.T).T # (N, 3) xyz_frac = xyz_frac % 1.0 # Wrap to [0, 1] - center_idx = torch.round(xyz_frac * grid_shape.unsqueeze(0)).to(torch.int64) # dtype-ok: rounded voxel center indices; torch indexing requires int64 + center_idx = torch.round(xyz_frac * grid_shape.unsqueeze(0)).to(get_int_dtype()) else: # Fallback for orthogonal cells (less accurate for non-orthogonal) voxelsize = real_space_grid[3, 3, 3] - real_space_grid[2, 2, 2] center_idx = torch.round( (xyz - grid_origin.unsqueeze(0)) / voxelsize.unsqueeze(0) - ).to(torch.int64) # dtype-ok: voxel index cast; torch indexing requires int64 + ).to(get_int_dtype()) voxel_indices_wrapped = excise_angstrom_radius_around_coord( real_space_grid, center_idx, radius_angstrom diff --git a/torchref/base/french_wilson.py b/torchref/base/french_wilson.py index a03c6df2..fbdcc963 100644 --- a/torchref/base/french_wilson.py +++ b/torchref/base/french_wilson.py @@ -1194,7 +1194,7 @@ def estimate_mean_intensity_by_resolution( # Use scatter_add to compute sum of intensities per bin bin_sums = torch.zeros(actual_n_bins, dtype=I.dtype, device=I.device) - bin_counts = torch.zeros(actual_n_bins, dtype=torch.long, device=I.device) # dtype-ok: count accumulator; scatter_add source is long ones, dtype must match + bin_counts = torch.zeros(actual_n_bins, dtype=bin_indices.dtype, device=I.device) bin_sums.scatter_add_(0, bin_indices, I_sorted) bin_counts.scatter_add_(0, bin_indices, torch.ones_like(bin_indices)) diff --git a/torchref/base/metrics/binwise_scale.py b/torchref/base/metrics/binwise_scale.py index 2d803ef1..8e0af305 100644 --- a/torchref/base/metrics/binwise_scale.py +++ b/torchref/base/metrics/binwise_scale.py @@ -58,7 +58,7 @@ def binwise_scale( Fo = Fo.reshape(-1) device, dtype = Fc.device, Fc.dtype - bins = bins.reshape(-1).to(device=device, dtype=torch.int64) # dtype-ok: resolution-bin indices used as scatter_add index; requires int64 + bins = bins.reshape(-1).to(device=device, dtype=torch.int64) # dtype-ok: scatter_add index; int64 required on torch < 2.8 if nbins is None: nbins = int(bins.max().item()) + 1 if bins.numel() else 0 diff --git a/torchref/base/reciprocal/grid_operations.py b/torchref/base/reciprocal/grid_operations.py index a631a640..34555706 100644 --- a/torchref/base/reciprocal/grid_operations.py +++ b/torchref/base/reciprocal/grid_operations.py @@ -7,6 +7,7 @@ import math import torch +from torchref.config import get_int_dtype def place_on_grid( @@ -47,14 +48,14 @@ def place_on_grid( dtype = structure_factor.dtype Nx, Ny, Nz = [int(x) for x in grid_size] hkls = hkls.to(device=device) - h = hkls[:, 0].to(torch.int64) # dtype-ok: hkl component cast to int64 for flat grid-index arithmetic; indexing requires long - k = hkls[:, 1].to(torch.int64) # dtype-ok: hkl component cast to int64 for flat grid-index arithmetic; indexing requires long - l = hkls[:, 2].to(torch.int64) # dtype-ok: hkl component cast to int64 for flat grid-index arithmetic; indexing requires long + h = hkls[:, 0].to(get_int_dtype()) + k = hkls[:, 1].to(get_int_dtype()) + l = hkls[:, 2].to(get_int_dtype()) hi = torch.remainder(h, Nx) ki = torch.remainder(k, Ny) li = torch.remainder(l, Nz) - lin = (hi * (Ny * Nz) + ki * Nz + li).to(torch.int64) # (N,) # dtype-ok: flat grid index (lin) for scatter/gather; requires int64 + lin = (hi * (Ny * Nz) + ki * Nz + li).to(get_int_dtype()) # (N,) grid = torch.zeros((B, Nx * Ny * Nz), dtype=dtype, device=device) grid = grid.index_add(1, lin, structure_factor) # (B, Nx*Ny*Nz) @@ -62,7 +63,7 @@ def place_on_grid( hi_sym = torch.remainder(-h, Nx) ki_sym = torch.remainder(-k, Ny) li_sym = torch.remainder(-l, Nz) - lin_sym = (hi_sym * (Ny * Nz) + ki_sym * Nz + li_sym).to(torch.int64) # dtype-ok: symmetry flat grid index (lin_sym) for scatter/gather; requires int64 + lin_sym = (hi_sym * (Ny * Nz) + ki_sym * Nz + li_sym).to(get_int_dtype()) vals_conj = torch.conj(structure_factor) grid = grid.index_add(1, lin_sym, vals_conj) @@ -101,9 +102,9 @@ def extract_structure_factor_from_grid(reciprocal_grid, hkls) -> torch.Tensor: # Same wrapping convention as place_on_grid. hkls = hkls.to(device=device) - h = hkls[:, 0].to(torch.int64) # dtype-ok: hkl component cast to int64 for flat grid-index arithmetic; indexing requires long - k = hkls[:, 1].to(torch.int64) # dtype-ok: hkl component cast to int64 for flat grid-index arithmetic; indexing requires long - l = hkls[:, 2].to(torch.int64) # dtype-ok: hkl component cast to int64 for flat grid-index arithmetic; indexing requires long + h = hkls[:, 0].to(get_int_dtype()) + k = hkls[:, 1].to(get_int_dtype()) + l = hkls[:, 2].to(get_int_dtype()) hi = torch.remainder(h, Nx) ki = torch.remainder(k, Ny) diff --git a/torchref/base/reciprocal/symmetry.py b/torchref/base/reciprocal/symmetry.py index 390a13b4..f10c879e 100644 --- a/torchref/base/reciprocal/symmetry.py +++ b/torchref/base/reciprocal/symmetry.py @@ -22,7 +22,7 @@ class through import torch -from torchref.config import canonical_device +from torchref.config import canonical_device, get_int_dtype from torchref.utils.autograd_ops import gather_with_index_add from torchref.utils.device_mixin import DeviceMixin @@ -45,13 +45,13 @@ def _equiv_hkls_to_flat_indices( Returns ------- torch.Tensor - Flat indices, shape ``(n_ops * N,)``, dtype ``int64``, wrapped modulo the grid. + Flat indices, shape ``(n_ops * N,)``, in the configured int dtype, wrapped modulo the grid. """ all_hkl = equiv_hkls.reshape(-1, 3) hi = torch.remainder(all_hkl[:, 0], Nx) ki = torch.remainder(all_hkl[:, 1], Ny) li = torch.remainder(all_hkl[:, 2], Nz) - return (hi * (Ny * Nz) + ki * Nz + li).to(torch.int64) # dtype-ok: flat HKL grid index; int64 avoids overflow, used for indexing + return (hi * (Ny * Nz) + ki * Nz + li).to(get_int_dtype()) class ReciprocalSymmetryExtractor(DeviceMixin): diff --git a/torchref/base/scattering/scattering_table.py b/torchref/base/scattering/scattering_table.py index fbc26aab..8a3d8332 100644 --- a/torchref/base/scattering/scattering_table.py +++ b/torchref/base/scattering/scattering_table.py @@ -12,7 +12,7 @@ import torch -from torchref.config import get_float_dtype +from torchref.config import get_float_dtype, get_int_dtype # Global cache for the loaded table _TABLE_CACHE: Optional[dict] = None @@ -167,8 +167,7 @@ def get_scattering_params_by_z( table = load_scattering_table(device=device, dtype=dtype) - # Long, not the caller's int32: torch indexing requires it. - z_idx = z_tensor.to(device=device, dtype=torch.long) # dtype-ok: z cast to long for scattering-table index lookup; indexing requires long + z_idx = z_tensor.to(device=device, dtype=get_int_dtype()) A = table["A"][z_idx] B = table["B"][z_idx] @@ -254,4 +253,4 @@ def elements_to_z(elements: list, normalize: bool = True) -> torch.Tensor: z = element_to_z.get(elem, 0) z_values.append(z) - return torch.tensor(z_values, dtype=torch.int32) # dtype-ok: atomic-number Z categorical codes; fixed int32 lookup keys + return torch.tensor(z_values, dtype=get_int_dtype()) diff --git a/torchref/cli/mtz2map.py b/torchref/cli/mtz2map.py index d02d75ac..b1fd7e25 100644 --- a/torchref/cli/mtz2map.py +++ b/torchref/cli/mtz2map.py @@ -18,7 +18,7 @@ import numpy as np import torch -from torchref.config import get_float_dtype +from torchref.config import get_float_dtype, get_int_dtype from torchref.cli._common import ( add_general_args, add_resolution_args, @@ -179,7 +179,7 @@ def main(): f"{d_spacings.max():.2f} - {d_spacings.min():.2f} A") # --- Convert to torch --- - hkl_t = torch.tensor(hkl, dtype=torch.int32, device=device) # dtype-ok: hkl Miller indices fed to symmetry expand; fixed int32 crystallographic representation + hkl_t = torch.tensor(hkl, dtype=get_int_dtype(), device=device) amp_t = torch.tensor(amplitudes, dtype=get_float_dtype(), device=device) phi_t = torch.tensor(phases_deg, dtype=get_float_dtype(), device=device) * (np.pi / 180.0) diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index 8c94a52d..8cd1f5ce 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -28,6 +28,7 @@ import numpy as np import torch +from torchref.config import get_int_dtype from torchref.cli._common import ( add_dmin_arg, @@ -81,7 +82,7 @@ def build_atom_mask(selection_xyz, real_space_grid, cell, mask_radius, device): inv_frac_matrix=inv_frac, ) - mask = torch.zeros(grid_shape, dtype=torch.int32, device=device) # dtype-ok: integer solvent-mask accumulator (mask>0); categorical count, not model-precision data + mask = torch.zeros(grid_shape, dtype=get_int_dtype(), device=device) mask = add_to_solvent_mask( surrounding_coords, voxel_indices, diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index 3ddc7440..986f13bd 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -18,6 +18,7 @@ import time import torch +from torchref.config import get_int_dtype _PROFILE = bool(os.environ.get("FRF_PROFILE")) @@ -143,7 +144,7 @@ def spherical_bessel_table( inv_threshold = 1.0 / threshold # Rescales applied so far, per element. Every element's ladder sits in the # single frame 2**(-_BESSEL_RESCALE_EXP * n_rescales). - n_rescales = torch.zeros_like(x64, dtype=torch.int32) # dtype-ok: small integer counter + n_rescales = torch.zeros_like(x64, dtype=get_int_dtype()) for n in range(n_start, 0, -1): j_low = (2.0 * n + 1.0) * inv_x * j_mid - j_high @@ -162,7 +163,7 @@ def spherical_bessel_table( j_high = j_high * factor if n - 1 <= u_max: j_table[n - 1:] = j_table[n - 1:] * factor - n_rescales = n_rescales + over.to(torch.int32) # dtype-ok: small integer counter + n_rescales = n_rescales + over.to(get_int_dtype()) true_j0 = torch.sin(x64) * inv_x true_j0 = torch.where(x64 < 1e-30, torch.ones_like(x64), true_j0) @@ -301,15 +302,15 @@ def bessel_sh_expand( n_list.append(n) u_list.append(u) w_list.append(math.sqrt(float(2 * u + 1))) - l_idx = torch.tensor(l_list, dtype=torch.long, device=device) # dtype-ok: index tensor; index_add_/gather need int64 - n_idx = torch.tensor(n_list, dtype=torch.long, device=device) # dtype-ok: index tensor; index_add_/gather need int64 - u_idx = torch.tensor(u_list, dtype=torch.long, device=device) # dtype-ok: index tensor; index_add_/gather need int64 + l_idx = torch.tensor(l_list, dtype=get_int_dtype(), device=device) + n_idx = torch.tensor(n_list, dtype=get_int_dtype(), device=device) + u_idx = torch.tensor(u_list, dtype=get_int_dtype(), device=device) w_vec = torch.tensor(w_list, dtype=comp_real, device=device) # Only even degrees l ∈ [2, lmax_even] carry signal (odd-l and l=0 are zeroed # by Patterson centrosymmetry). Compute / contract Y_lm on these rows only — # the assembly + einsum are the bottleneck, so this ~halves them. The full # c_nlm keeps the (L, ...) shape with odd/zero rows left at zero. - even_l_idx = torch.tensor(even_ls, dtype=torch.long, device=device) # dtype-ok: index tensor; index_add_/gather need int64 + even_l_idx = torch.tensor(even_ls, dtype=get_int_dtype(), device=device) M = s_vectors.shape[0] einsum_dtype = complex_dtype @@ -347,8 +348,8 @@ def _tick(t0): s_key = s_vectors.detach().cpu().to(torch.float64) # dtype-ok: exact clustering key on the host; the device never sees it s_mag_key = s_key.norm(dim=-1).clamp(min=1e-30) cos_key = (s_key[..., 2] / s_mag_key).clamp(min=-1.0, max=1.0) - k_s = (s_mag_key * _GROUP_SCALE_S).round().to(torch.int64) # dtype-ok: exact clustering key - k_c = (cos_key * _GROUP_SCALE_COS).round().to(torch.int64) + _GROUP_SCALE_COS # dtype-ok: exact clustering key + k_s = (s_mag_key * _GROUP_SCALE_S).round().to(torch.int64) # dtype-ok: clustering key k_s*(2e7+1)+k_c overflows int32 + k_c = (cos_key * _GROUP_SCALE_COS).round().to(torch.int64) + _GROUP_SCALE_COS # dtype-ok: clustering key k_s*(2e7+1)+k_c overflows int32 key = (k_s * (2 * _GROUP_SCALE_COS + 1) + k_c).to(s_vectors.device) uniq_key, inverse = torch.unique(key, return_inverse=True) n_clusters = int(uniq_key.shape[0]) @@ -379,7 +380,7 @@ def _group_mean(values, index, n_groups): # index to meet device values. inv_s = inv_s.to(device) n_shells = int(uniq_ks.shape[0]) - shell_of_cluster = torch.zeros(n_clusters, dtype=torch.long, device=device) # dtype-ok: index tensor; index_add_/gather need int64 + shell_of_cluster = torch.zeros(n_clusters, dtype=get_int_dtype(), device=device) shell_of_cluster[inverse] = inv_s shell_smag = _group_mean(s_mag_all.to(comp_real), inv_s, n_shells) diff --git a/torchref/experimental/alignment/frf/dense_calc.py b/torchref/experimental/alignment/frf/dense_calc.py index d95921c5..c7be30d0 100644 --- a/torchref/experimental/alignment/frf/dense_calc.py +++ b/torchref/experimental/alignment/frf/dense_calc.py @@ -21,7 +21,7 @@ import torch -from torchref.config import get_float_dtype +from torchref.config import get_float_dtype, get_int_dtype if TYPE_CHECKING: from torchref.model import ModelFT @@ -82,7 +82,7 @@ def dense_calc_via_box( H, K, Lg = torch.meshgrid(idx, idx, idx, indexing="ij") hkl = torch.stack( [H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], dim=-1 - ).to(torch.long) # dtype-ok: Miller indices are integers + ).to(get_int_dtype()) # Cubic box: |s| = |hkl| / a. real = get_float_dtype() smag = hkl.to(real).norm(dim=-1) / a diff --git a/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py b/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py index e7d56746..6a59efd6 100644 --- a/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py +++ b/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py @@ -32,6 +32,7 @@ from typing import Optional, Tuple import torch +from torchref.config import get_int_dtype from torchref.base.electron_density.kernels.cpu._cpp_build import build_extension @@ -241,12 +242,12 @@ def clear_cache() -> None: def shell_offsets(shell: torch.Tensor, n_shells: int) -> torch.Tensor: """Start index of each shell in a shell-sorted cluster array, plus the end. - ``(n_shells + 1,)`` int64. The kernel needs the ranges rather than the + ``(n_shells + 1,)`` integer offsets. The kernel needs the ranges rather than the per-cluster labels so that a thread can own a set of shells outright and write their accumulator rows without atomics. """ counts = torch.bincount(shell, minlength=n_shells) - offsets = torch.zeros(n_shells + 1, dtype=torch.long, device=shell.device) # dtype-ok: index tensor; index_add_/gather need int64 + offsets = torch.zeros(n_shells + 1, dtype=get_int_dtype(), device=shell.device) torch.cumsum(counts, dim=0, out=offsets[1:]) return offsets diff --git a/torchref/experimental/alignment/frf/peak_finder.py b/torchref/experimental/alignment/frf/peak_finder.py index 19692db1..de15f569 100644 --- a/torchref/experimental/alignment/frf/peak_finder.py +++ b/torchref/experimental/alignment/frf/peak_finder.py @@ -28,6 +28,7 @@ from typing import List, Optional import torch +from torchref.config import get_int_dtype from ....base.alignment.rotation import rotation_matrix_euler_zyz from .types import AdaptiveRotationFunction, RotationPeak @@ -55,7 +56,7 @@ def _so3_greedy_nms( """ n = values.shape[0] if n == 0: - return torch.empty(0, dtype=torch.int64, device=values.device) # dtype-ok: index tensor; index_add_/gather need int64 + return torch.empty(0, dtype=get_int_dtype(), device=values.device) # The greedy walk is inherently sequential and latency-bound; on GPU a # per-iteration `.item()` sync would dominate. Move the (tiny) candidate # rotations to CPU once and run the loop there with no device syncs, a @@ -97,7 +98,7 @@ def _so3_greedy_nms( count += 1 if count >= keep_at_most: break - return torch.tensor(kept_idx, dtype=torch.int64, device=values.device) # dtype-ok: index tensor; index_add_/gather need int64 + return torch.tensor(kept_idx, dtype=get_int_dtype(), device=values.device) def find_rotation_peaks( diff --git a/torchref/experimental/alignment/frf/preprocessing.py b/torchref/experimental/alignment/frf/preprocessing.py index 63977817..b8eed82f 100644 --- a/torchref/experimental/alignment/frf/preprocessing.py +++ b/torchref/experimental/alignment/frf/preprocessing.py @@ -298,8 +298,8 @@ def fit_relative_wilson_b( F2_calc = (F_calc * F_calc).to(real) s2_obs = (s_mag * s_mag).to(real) - counts_obs = torch.zeros(n_shells, dtype=torch.int64, device=s_mag.device) # dtype-ok: per-shell counts - counts_calc = torch.zeros(n_shells, dtype=torch.int64, device=s_mag.device) # dtype-ok: per-shell counts + counts_obs = torch.zeros(n_shells, dtype=shell_idx_obs.dtype, device=s_mag.device) + counts_calc = torch.zeros(n_shells, dtype=shell_idx_calc.dtype, device=s_mag.device) sum_F2obs = torch.zeros(n_shells, dtype=real, device=s_mag.device) sum_F2calc = torch.zeros(n_shells, dtype=real, device=s_mag.device) sum_s2 = torch.zeros(n_shells, dtype=real, device=s_mag.device) diff --git a/torchref/experimental/alignment/frf/sitelist_ang.py b/torchref/experimental/alignment/frf/sitelist_ang.py index c297c3c2..a7a562f2 100644 --- a/torchref/experimental/alignment/frf/sitelist_ang.py +++ b/torchref/experimental/alignment/frf/sitelist_ang.py @@ -39,6 +39,7 @@ from typing import List, Tuple import torch +from torchref.config import get_int_dtype from ....config import canonical_device from ....symmetry.symmetry import find_fft_friendly_size @@ -98,7 +99,7 @@ def build_dense_map_per_beta( (n_beta, fft_size, fft_size), dtype=S.dtype, device=device, ) m_vals = torch.arange(-(L - 1), L, device=device) - idx = (m_vals % fft_size).to(torch.int64) # dtype-ok: index tensor; index_add_/gather need int64 + idx = (m_vals % fft_size).to(get_int_dtype()) pad[:, idx.unsqueeze(1), idx.unsqueeze(0)] = S # 3. Forward 2D FFT — torch convention: @@ -210,8 +211,8 @@ def build_adaptive_sample_list( # original dict scan, but no host sync / Python loop. # Hash the two rounded fracs (each in [0, 1e6]) into one int64 so we # can use the fast 1-D unique instead of a 2-D row lexsort. - a_round = (alpha_frac * 1_000_000).round().to(torch.int64) # dtype-ok: index tensor; index_add_/gather need int64 - g_round = (gamma_frac * 1_000_000).round().to(torch.int64) # dtype-ok: index tensor; index_add_/gather need int64 + a_round = (alpha_frac * 1_000_000).round().to(torch.int64) # dtype-ok: a_round*1_000_001+g_round overflows int32 + g_round = (gamma_frac * 1_000_000).round().to(torch.int64) # dtype-ok: a_round*1_000_001+g_round overflows int32 key_hash = a_round * 1_000_001 + g_round _, uniq_idx = torch.unique(key_hash, return_inverse=True) n = uniq_idx.shape[0] @@ -234,7 +235,7 @@ def build_adaptive_sample_list( alphas = torch.cat(alphas_list).to(device) gammas = torch.cat(gammas_list).to(device) betas_flat = torch.cat(betas_list).to(device) - beta_starts_t = torch.tensor(beta_starts, dtype=torch.int64, device=device) # dtype-ok: index tensor; index_add_/gather need int64 + beta_starts_t = torch.tensor(beta_starts, dtype=get_int_dtype(), device=device) b = torch.arange(bmax, dtype=torch.float64, device=cpu) # dtype-ok: sample-list geometry follows the accumulator's width betas_rad = (b * grid_sampling_deg * deg2rad).to(device=device, dtype=dtype) @@ -256,8 +257,8 @@ def _bilinear_interp_periodic( N = M.shape[-1] af = (alpha_frac % 1.0) * N gf = (gamma_frac % 1.0) * N - a0 = torch.floor(af).to(torch.int64) % N # dtype-ok: index tensor; index_add_/gather need int64 - g0 = torch.floor(gf).to(torch.int64) % N # dtype-ok: index tensor; index_add_/gather need int64 + a0 = torch.floor(af).to(get_int_dtype()) % N + g0 = torch.floor(gf).to(get_int_dtype()) % N a1 = (a0 + 1) % N g1 = (g0 + 1) % N da = (af - torch.floor(af)).to(M.real.dtype) diff --git a/torchref/experimental/alignment/sh.py b/torchref/experimental/alignment/sh.py index 0da057fe..f343c886 100644 --- a/torchref/experimental/alignment/sh.py +++ b/torchref/experimental/alignment/sh.py @@ -28,6 +28,7 @@ from typing import Optional, Tuple import torch +from torchref.config import get_int_dtype from ...config import get_float_dtype @@ -134,9 +135,9 @@ def _bar_legendre_recurrence( if keep_l is None: rows = torch.arange(L, device=device) else: - rows = keep_l.to(device=device, dtype=torch.long) # dtype-ok: index tensor; index_add_/gather need int64 + rows = keep_l.to(device=device, dtype=get_int_dtype()) # l -> its position in the output, or -1 when it is not kept. - where = torch.full((L,), -1, dtype=torch.long, device=device) # dtype-ok: index tensor; index_add_/gather need int64 + where = torch.full((L,), -1, dtype=get_int_dtype(), device=device) where[rows] = torch.arange(rows.numel(), device=device) where_list = where.tolist() @@ -338,7 +339,7 @@ def fit_overall_anisotropy( F, s, idx, cen = F[ok], s[ok], idx[ok], cen[ok] I = F * F - count = torch.zeros(P, dtype=torch.int64, device=F.device) # dtype-ok: index tensor; index_add_/gather need int64 + count = torch.zeros(P, dtype=idx.dtype, device=F.device) total = torch.zeros(P, dtype=work, device=F.device) count.index_add_(0, idx, torch.ones_like(idx)) total.index_add_(0, idx, I) @@ -536,7 +537,7 @@ def compute_patterson_shell_variance( valid = shell_idx >= 0 patt_v = patt[valid] idx_v = shell_idx[valid] - count = torch.zeros(P, dtype=torch.int64, device=device) # dtype-ok: index tensor; index_add_/gather need int64 + count = torch.zeros(P, dtype=idx_v.dtype, device=device) count.index_add_(0, idx_v, torch.ones_like(idx_v)) sum1 = torch.zeros(P, dtype=dtype, device=device) sum2 = torch.zeros(P, dtype=dtype, device=device) diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index 97913973..7a3ec1f1 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -38,7 +38,7 @@ import torch from torchref.base.targets.xray_likelihoods import rice_per_refl -from torchref.config import get_complex_dtype, get_default_device, get_float_dtype +from torchref.config import get_complex_dtype, get_default_device, get_float_dtype, get_int_dtype from torchref.scaling import WilsonNormaliser from torchref.scaling.weighting import (inverse_variance_weight, normalise_weight, snr_from_amplitude) @@ -151,7 +151,7 @@ def build( rec_basis = real_cell.reciprocal_basis_matrix.to(device=dev, dtype=real) s_mag = (hkl_i.to(real) @ rec_basis).norm(dim=-1) - hkl_l = hkl_i.round().to(torch.int64) # dtype-ok: Miller indices are integers + hkl_l = hkl_i.round().to(get_int_dtype()) # friedel=False: Wilson's = eps*Sigma counts the operations mapping # h to itself, which add coherently and set the mean. The Friedel-folded # branch changes the distribution instead, and that is centricity -- @@ -249,7 +249,7 @@ def prepare_candidate( # h_R[i, n, d] = sum_e hkl[n, e] sym_R[i, e, d]: the h.S convention. h_R = torch.einsum("ne,ied->ind", hkl, sym_R) phase = torch.exp((2j * math.pi) * torch.einsum("ne,ie->in", hkl, sym_t).to(cplx)) - hkl_SN = h_R.reshape(-1, 3).round().to(torch.int64).to(model_p1.xyz().device) # dtype-ok: Miller indices are integers + hkl_SN = h_R.reshape(-1, 3).round().to(get_int_dtype()).to(model_p1.xyz().device) with torch.no_grad(): F_all = model_p1(hkl_SN).to(device).reshape(S, N).to(cplx) G_raw = F_all * phase @@ -395,7 +395,7 @@ def fast_translation_function( G = cand.G.to(device=device, dtype=cplx) S, N = G.shape coeff = obs.coeff.to(device=device, dtype=cplx) - h_R_int = cand.h_R.round().to(torch.int64) # dtype-ok: Miller indices are integers + h_R_int = cand.h_R.round().to(get_int_dtype()) # The pair (j, i) is the conjugate of (i, j) at -dh, so the map is twice # the real part of the upper triangle's transform plus the diagonal, which diff --git a/torchref/experimental/ensemble/ensemble_amber_kl.py b/torchref/experimental/ensemble/ensemble_amber_kl.py index ebb5cfc3..effb10b4 100644 --- a/torchref/experimental/ensemble/ensemble_amber_kl.py +++ b/torchref/experimental/ensemble/ensemble_amber_kl.py @@ -52,6 +52,7 @@ import numpy as np import torch +from torchref.config import get_int_dtype from torchref.experimental.targets.amber_target import AMBER14_STANDARD, AmberTarget @@ -136,7 +137,7 @@ def __init__( self.register_buffer( "_member_atom_idx", torch.as_tensor( - atom_idx_np, dtype=torch.long, device=self._model.device # dtype-ok: atom index tensor for indexing; PyTorch requires int64 + atom_idx_np, dtype=get_int_dtype(), device=self._model.device ), ) else: diff --git a/torchref/experimental/ensemble/quasi_crystal_amber.py b/torchref/experimental/ensemble/quasi_crystal_amber.py index ee9986ca..bb712ece 100644 --- a/torchref/experimental/ensemble/quasi_crystal_amber.py +++ b/torchref/experimental/ensemble/quasi_crystal_amber.py @@ -59,6 +59,7 @@ import numpy as np import torch +from torchref.config import get_int_dtype from torchref.experimental.targets.amber_target import ( AmberTarget, @@ -564,19 +565,19 @@ def _ensure_torch_buffers(self, device: torch.device, dtype: torch.dtype) -> Non # Index pairs (long) for the scatter from model atoms into OMM slots. self._src_model_idx_torch = torch.from_numpy(self._src_model_idx_np).to( device=device, - dtype=torch.long, # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 + dtype=get_int_dtype(), ) self._dst_omm_idx_torch = torch.from_numpy(self._dst_omm_idx_np).to( device=device, - dtype=torch.long, # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 + dtype=get_int_dtype(), ) # Index of ensemble-model atoms (in the FULL EnsembleModel layout) # that survived the special-position filter — used in forward to # subset ``xyz_per_member`` before applying the layout transform. self._keep_atom_idx_torch = torch.from_numpy(self._keep_atom_idx_np).to( - device=device, dtype=torch.long - ) # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 + device=device, dtype=get_int_dtype() + ) self._omm_to_model = self._omm_to_model.to(device) self._buffers_device = device diff --git a/torchref/experimental/ensemble/wilson_prior.py b/torchref/experimental/ensemble/wilson_prior.py index 44490f9c..3d9704bf 100644 --- a/torchref/experimental/ensemble/wilson_prior.py +++ b/torchref/experimental/ensemble/wilson_prior.py @@ -172,7 +172,7 @@ def _build_bin_assignment(self) -> None: order = torch.argsort(res) n = res.numel() nbins = min(self.nbins, max(1, n // 50)) - bin_assign = torch.empty(n, dtype=torch.long, device=res.device) # dtype-ok: bin-assignment tensor used as scatter_add index; PyTorch requires int64 + bin_assign = torch.empty(n, dtype=torch.long, device=res.device) # dtype-ok: scatter_add index; int64 required on torch < 2.8 edges = torch.linspace(0, n, nbins + 1, device=res.device).round().long() for b in range(nbins): start = int(edges[b].item()) diff --git a/torchref/experimental/monolithic_refinement/density_scaler.py b/torchref/experimental/monolithic_refinement/density_scaler.py index afe76f95..1f76c1df 100644 --- a/torchref/experimental/monolithic_refinement/density_scaler.py +++ b/torchref/experimental/monolithic_refinement/density_scaler.py @@ -36,7 +36,7 @@ import torch import torch.nn as nn -from torchref.config import get_default_device, get_float_dtype +from torchref.config import get_default_device, get_float_dtype, get_int_dtype from torchref.scaling.scaler import Scaler from torchref.scaling.solvent import SolventModel from torchref.experimental.monolithic_refinement.density_solvent import ( @@ -127,7 +127,7 @@ def get_rec_solvent(self, hkl): Not detached: ``F_sol`` follows the moving atoms so gradients reach ``xyz``/``adp``. The scaler applies the contrast and falloff on top. """ - return self.density(hkl.to(torch.long)) # dtype-ok: hkl cast to long for density lookup indexing; PyTorch requires int64 + return self.density(hkl.to(get_int_dtype())) def update_solvent(self): """No-op: the density mask is rebuilt live on every scaler forward.""" diff --git a/torchref/experimental/targets/forcefield_target.py b/torchref/experimental/targets/forcefield_target.py index 09f0e1de..89beb6e6 100644 --- a/torchref/experimental/targets/forcefield_target.py +++ b/torchref/experimental/targets/forcefield_target.py @@ -178,7 +178,7 @@ def forward(self) -> torch.Tensor: Z = self.model.Z # Shape: (n_atoms,) # Ensure Z is long tensor - if Z.dtype != torch.long: # dtype-ok: dtype guard comparison against torch.long, not an allocation + if Z.dtype != torch.long: # dtype-ok: TorchMD-Net expects a LongTensor Z; external library contract Z = Z.long() # Create batch tensor (single structure = all zeros) diff --git a/torchref/experimental/targets/sampled_ml_phase_target.py b/torchref/experimental/targets/sampled_ml_phase_target.py index e505240b..d3510cc3 100644 --- a/torchref/experimental/targets/sampled_ml_phase_target.py +++ b/torchref/experimental/targets/sampled_ml_phase_target.py @@ -16,6 +16,7 @@ import numpy as np import torch +from torchref.config import get_int_dtype from typing import TYPE_CHECKING, Dict, Tuple from torchref.refinement.targets.base import Target @@ -129,7 +130,7 @@ def __init__( self.name = "xray_sampled_ml_work" if use_work_set else "xray_sampled_ml_test" # Register tunable parameters as buffers for state_dict access - self.register_buffer("_n_samples", torch.tensor(n_samples, dtype=torch.int64)) # dtype-ok: scalar sample-count buffer; categorical count, not model-precision data + self.register_buffer("_n_samples", torch.tensor(n_samples, dtype=get_int_dtype())) self.register_buffer("_sigma_model_log", torch.tensor(sigma_model_log)) self.register_buffer("_use_analytical", torch.tensor(use_analytical)) self.register_buffer("_use_antithetic", torch.tensor(use_antithetic)) @@ -545,7 +546,7 @@ def __init__( self.add_module("_scaler_dark", scaler_dark) # Tunable parameters as buffers - self.register_buffer("_n_samples", torch.tensor(n_samples, dtype=torch.int64)) # dtype-ok: scalar sample-count buffer; categorical count, not model-precision data + self.register_buffer("_n_samples", torch.tensor(n_samples, dtype=get_int_dtype())) self.register_buffer("_sigma_model_log", torch.tensor(sigma_model_log)) self.use_work_set = use_work_set diff --git a/torchref/io/datasets/fcalc_data.py b/torchref/io/datasets/fcalc_data.py index d4a4a2aa..8b2947fa 100644 --- a/torchref/io/datasets/fcalc_data.py +++ b/torchref/io/datasets/fcalc_data.py @@ -12,7 +12,7 @@ import pandas as pd import torch -from torchref.config import get_float_dtype, normalize_device +from torchref.config import get_float_dtype, get_int_dtype, normalize_device from torchref.symmetry import Cell, SpaceGroup, SpaceGroupLike from .base import CrystalDataset @@ -133,7 +133,7 @@ def from_cell_and_resolution( # make_miller_array returns unique HKL for the asymmetric unit only. hkl_list = gemmi.make_miller_array(gemmi_cell, gemmi_sg, d_min) - hkl = torch.tensor(hkl_list, dtype=torch.int32, device=device) # dtype-ok: hkl Miller indices; fixed int32 crystallographic representation, not model-precision data + hkl = torch.tensor(hkl_list, dtype=get_int_dtype(), device=device) resolution = get_d_spacing(hkl.float(), cell_tensor) diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index 58722edc..1b23355e 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -17,7 +17,7 @@ from torchref.base import math_torch from torchref.base.french_wilson import FrenchWilson -from torchref.config import dtypes, normalize_device +from torchref.config import dtypes, get_int_dtype, normalize_device from torchref.io import cif, mtz from torchref.io.datasets.base import CrystalDataset from torchref.symmetry import Cell, SpaceGroup @@ -306,7 +306,7 @@ def _subset_indices(self, kind: str) -> torch.Tensor: n = 0 if self.hkl is None else len(self.hkl) device = self.device if n == 0: - empty = torch.empty(0, dtype=torch.long, device=device) # dtype-ok: empty index tensor; PyTorch requires int64 for indexing + empty = torch.empty(0, dtype=get_int_dtype(), device=device) self._subset_cache = { "work": empty, "free": empty, @@ -428,7 +428,7 @@ def _reindex_per_reflection( n_src = len(self.hkl) if self.hkl is not None else 0 new_hkl = new_hkl.to(dtype=dtypes.int, device=self.device) n_out = len(new_hkl) - index_map = index_map.to(device=self.device, dtype=torch.long) # dtype-ok: index map used for indexing/gather; PyTorch requires int64 + index_map = index_map.to(device=self.device, dtype=get_int_dtype()) present = index_map >= 0 src_idx = index_map[present] @@ -1357,13 +1357,13 @@ def mean_res_per_bin(self) -> torch.Tensor: mean_resolutions = torch.scatter_add( mean_resolutions, 0, - self.bin_indices[mask].to(torch.int64), # dtype-ok: bin indices for scatter_add/index; PyTorch requires int64 + self.bin_indices[mask].to(torch.int64), # dtype-ok: scatter_add index; int64 required on torch < 2.8 self.resolution[mask], ) count_per_bin = torch.scatter_add( count_per_bin, 0, - self.bin_indices[mask].to(torch.int64), # dtype-ok: bin indices for scatter_add/index; PyTorch requires int64 + self.bin_indices[mask].to(torch.int64), # dtype-ok: scatter_add index; int64 required on torch < 2.8 torch.ones_like(self.resolution[mask], dtype=dtypes.int), ) mean_resolutions = mean_resolutions / count_per_bin.clamp(min=1).float() @@ -1392,12 +1392,12 @@ def mean_F_per_bin(self) -> torch.Tensor: count_per_bin = torch.zeros(self._n_bins, dtype=dtypes.int, device=self.device) mask = self.masks() mean_F = torch.scatter_add( - mean_F, 0, self.bin_indices[mask].to(torch.int64), self.F[mask] # dtype-ok: bin indices for scatter_add index arg; PyTorch requires int64 + mean_F, 0, self.bin_indices[mask].to(torch.int64), self.F[mask] # dtype-ok: scatter_add index; int64 required on torch < 2.8 ) count_per_bin = torch.scatter_add( count_per_bin, 0, - self.bin_indices[mask].to(torch.int64), # dtype-ok: bin indices for scatter_add index arg; PyTorch requires int64 + self.bin_indices[mask].to(torch.int64), # dtype-ok: scatter_add index; int64 required on torch < 2.8 torch.ones_like(self.F[mask], dtype=dtypes.int), ) mean_F = mean_F / count_per_bin.clamp(min=1).float() @@ -1426,12 +1426,12 @@ def mean_sigma_per_bin(self) -> Optional[torch.Tensor]: count_per_bin = torch.zeros(self._n_bins, dtype=dtypes.int, device=self.device) mask = self.masks() mean_sigma = torch.scatter_add( - mean_sigma, 0, self.bin_indices[mask].to(torch.int64), self.F_sigma[mask] # dtype-ok: bin indices for scatter_add index arg; PyTorch requires int64 + mean_sigma, 0, self.bin_indices[mask].to(torch.int64), self.F_sigma[mask] # dtype-ok: scatter_add index; int64 required on torch < 2.8 ) count_per_bin = torch.scatter_add( count_per_bin, 0, - self.bin_indices[mask].to(torch.int64), # dtype-ok: bin indices for scatter_add index arg; PyTorch requires int64 + self.bin_indices[mask].to(torch.int64), # dtype-ok: scatter_add index; int64 required on torch < 2.8 torch.ones_like(self.F_sigma[mask], dtype=dtypes.int), ) mean_sigma = mean_sigma / count_per_bin.clamp(min=1).float() @@ -2508,8 +2508,8 @@ def _build_anomalous_dataframe( # The (+) member is the unconjugated row, (-) is the Friedel-flagged row. arange = torch.arange(N) - plus_idx = torch.full((M,), -1, dtype=torch.long) # dtype-ok: Friedel-mate index map (-1 sentinel) for indexing; PyTorch requires int64 - minus_idx = torch.full((M,), -1, dtype=torch.long) # dtype-ok: Friedel-mate index map (-1 sentinel) for indexing; PyTorch requires int64 + plus_idx = torch.full((M,), -1, dtype=get_int_dtype()) + minus_idx = torch.full((M,), -1, dtype=get_int_dtype()) # A Bijvoet mate only counts as present if it is a real, positive # observation. Stacked anomalous input (rs.stack_anomalous) carries a # row for every *absent* mate with a NaN intensity, which French-Wilson @@ -2947,7 +2947,7 @@ def remap( ---------- new_hkl : torch.Tensor, shape (M, 3) New Miller indices. - index_mapping : torch.Tensor, shape (M,), dtype int64 + index_mapping : torch.Tensor, shape (M,), integer dtype Maps new indices to original: ``new[i] = old[index_mapping[i]]`` Values of -1 indicate missing reflections (filled with defaults). phase_shifts : torch.Tensor, optional, shape (M,) diff --git a/torchref/model/disorder_field.py b/torchref/model/disorder_field.py index a2911d3f..582997e1 100644 --- a/torchref/model/disorder_field.py +++ b/torchref/model/disorder_field.py @@ -88,7 +88,7 @@ def farthest_point_anchors(xyz: torch.Tensor, n_nodes: int) -> torch.Tensor: chosen.append(nxt) d2_nearest = torch.minimum(d2_nearest, ((xyz - xyz[nxt]) ** 2).sum(-1)) - anchors = torch.tensor(chosen, dtype=torch.int64, device=xyz.device) # dtype-ok: anchor atom indices; torch indexing requires int64 + anchors = torch.tensor(chosen, dtype=get_int_dtype(), device=xyz.device) # Lloyd relaxation, snapping to real atoms so an anchor is always an atom index. for _ in range(10): @@ -126,7 +126,7 @@ def density_anchor_rows(xyz: torch.Tensor, n_nodes: int): """ seeds = farthest_point_anchors(xyz, n_nodes) assign = torch.cdist(xyz, xyz[seeds]).argmin(dim=1) - atom_idx = torch.arange(xyz.shape[0], dtype=torch.int64, device=xyz.device) # dtype-ok: arange atom indices; index requires int64 + atom_idx = torch.arange(xyz.shape[0], dtype=get_int_dtype(), device=xyz.device) # A seed whose cluster somehow came out empty still needs a position. present = torch.bincount(assign, minlength=seeds.shape[0]) > 0 @@ -715,12 +715,12 @@ def __init__( if anchor_rows is None: anchor_atom = farthest_point_anchors(xyz, n_nodes) anchor_node = torch.arange( - anchor_atom.shape[0], dtype=torch.int64, device=device # dtype-ok: arange anchor indices; index requires int64 + anchor_atom.shape[0], dtype=get_int_dtype(), device=device ) else: anchor_atom, anchor_node = anchor_rows - anchor_atom = anchor_atom.to(device=device, dtype=torch.int64) # dtype-ok: anchor_atom indices cast; index requires int64 - anchor_node = anchor_node.to(device=device, dtype=torch.int64) # dtype-ok: anchor_node indices cast; index requires int64 + anchor_atom = anchor_atom.to(device=device, dtype=get_int_dtype()) + anchor_node = anchor_node.to(device=device, dtype=get_int_dtype()) n_k = int(anchor_node.max()) + 1 node_pos = self._segment_mean(xyz, anchor_atom, anchor_node, n_k) diff --git a/torchref/model/model.py b/torchref/model/model.py index 56d9a8db..bf547fb4 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -22,6 +22,7 @@ from torchref.base import math_torch from torchref.config import ( + get_int_dtype, canonical_device, get_default_device, get_float_dtype, @@ -435,7 +436,7 @@ def _build_z_tensor(self) -> torch.Tensor: for elem in self.pdb["element"] ] self.register_buffer( - "_Z", torch.tensor(z_values, dtype=torch.int32, device=self.device) # dtype-ok: atomic-number Z categorical codes buffer; fixed int32 lookup keys + "_Z", torch.tensor(z_values, dtype=get_int_dtype(), device=self.device) ) return self._Z @@ -946,7 +947,7 @@ def _create_occupancy_groups(self, pdb_df, initial_occ): altloc_groups = [] refinable_mask = torch.zeros(n_atoms, dtype=torch.bool) - sharing_groups_tensor = torch.arange(n_atoms, dtype=torch.long) # dtype-ok: arange atom indices (sharing groups); index requires long + sharing_groups_tensor = torch.arange(n_atoms, dtype=get_int_dtype()) collapsed_idx = 0 # First pass: altlocs. ALL atoms of one conformation must share a collapsed @@ -1015,7 +1016,7 @@ def _create_occupancy_groups(self, pdb_df, initial_occ): # Compact to contiguous indices 0..n_collapsed-1. unique_indices = torch.unique(sharing_groups_tensor, sorted=True) - index_map = torch.zeros(n_atoms, dtype=torch.long) # dtype-ok: index_map atom-index remap; indexing requires long + index_map = torch.zeros(n_atoms, dtype=get_int_dtype()) for new_idx, old_idx in enumerate(unique_indices): mask = sharing_groups_tensor == old_idx sharing_groups_tensor[mask] = new_idx @@ -2029,7 +2030,7 @@ def register_alternative_conformations(self): for altloc in unique_altlocs: altloc_atoms = group[group["altloc"] == altloc] indices = torch.tensor( - altloc_atoms["index"].tolist(), dtype=torch.long # dtype-ok: altloc atom indices; indexing requires long + altloc_atoms["index"].tolist(), dtype=get_int_dtype() ) conformation_tensors.append(indices) @@ -3000,7 +3001,7 @@ def _complete_riding_waters(self, frames): self.pdb, plan, restraints.topology ) frames = generated.remap(old_rows).fill_planned_rows(new_rows) - source = torch.empty(len(augmented), dtype=torch.long, device=self.device) + source = torch.empty(len(augmented), dtype=get_int_dtype(), device=self.device) old_index = torch.as_tensor(old_rows, device=self.device) new_index = torch.as_tensor(new_rows, device=self.device) source[old_index] = torch.arange(len(self.pdb), device=self.device) diff --git a/torchref/model/parameter_wrappers.py b/torchref/model/parameter_wrappers.py index 1991cabd..b009415d 100644 --- a/torchref/model/parameter_wrappers.py +++ b/torchref/model/parameter_wrappers.py @@ -14,7 +14,7 @@ import torch from torch import nn -from torchref.config import get_float_dtype, normalize_device +from torchref.config import get_float_dtype, get_int_dtype, normalize_device from torchref.utils.caching import CachedForwardMixin from torchref.utils.device_mixin import DeviceMixin @@ -1474,11 +1474,11 @@ def _setup_sharing_groups_and_expansion( # Use sharing_groups directly as the expansion mask if sharing_groups is None: # No sharing - each atom maps to its own index - expansion_mask = torch.arange(n_atoms, dtype=torch.long, device=device) # dtype-ok: arange expansion_mask atom indices; index requires long + expansion_mask = torch.arange(n_atoms, dtype=torch.long, device=device) # dtype-ok: expansion_mask is a scatter_add_ index; int64 required on torch < 2.8 self._collapsed_shape = n_atoms else: # Use the provided index tensor - expansion_mask = sharing_groups.to(device=device, dtype=torch.long) # dtype-ok: expansion_mask atom/group indices for scatter; requires long + expansion_mask = sharing_groups.to(device=device, dtype=torch.long) # dtype-ok: expansion_mask is a scatter_add_ index; int64 required on torch < 2.8 self._collapsed_shape = expansion_mask.max().item() + 1 self.register_buffer("expansion_mask", expansion_mask) @@ -1500,10 +1500,10 @@ def _setup_sharing_groups_and_expansion( for conf_atoms in conf_groups: if isinstance(conf_atoms, (list, tuple)): conf_atoms = torch.tensor( - conf_atoms, dtype=torch.long, device=device # dtype-ok: conf_atoms atom indices; indexing requires long + conf_atoms, dtype=get_int_dtype(), device=device ) else: - conf_atoms = conf_atoms.to(device=device, dtype=torch.long) # dtype-ok: conf_atoms atom indices cast; indexing requires long + conf_atoms = conf_atoms.to(device=device, dtype=get_int_dtype()) # Get collapsed index for first atom collapsed_idx = expansion_mask[conf_atoms[0]].item() @@ -1531,7 +1531,7 @@ def _setup_sharing_groups_and_expansion( # Store as dictionary with keys like 'linked_occ_2', 'linked_occ_3', etc. for n_conf, groups in linked_occupancies.items(): # Shape: (N_groups, n_conf) - tensor = torch.tensor(groups, dtype=torch.long, device=device) # dtype-ok: linked-occupancy group index buffer; indexing requires long + tensor = torch.tensor(groups, dtype=get_int_dtype(), device=device) self.register_buffer(f"linked_occ_{n_conf}", tensor) # Store which sizes we have @@ -1539,7 +1539,7 @@ def _setup_sharing_groups_and_expansion( # Create count buffer for vectorized collapse operations # counts[i] = number of atoms that map to collapsed index i - counts = torch.zeros(self._collapsed_shape, dtype=torch.long, device=device) # dtype-ok: count accumulator; scatter_add source is long ones, dtype must match + counts = torch.zeros(self._collapsed_shape, dtype=expansion_mask.dtype, device=device) counts.scatter_add_(0, expansion_mask, torch.ones_like(expansion_mask)) self.register_buffer("collapse_counts", counts) @@ -2017,7 +2017,7 @@ def from_residue_groups( grouped = pdb_dataframe.groupby(["resname", "resseq", "chainid", "altloc"]) n_atoms = len(initial_values) - sharing_groups_tensor = torch.arange(n_atoms, dtype=torch.long) # dtype-ok: arange atom indices (sharing groups); index requires long + sharing_groups_tensor = torch.arange(n_atoms, dtype=get_int_dtype()) # Singletons keep their arange ids (0..n_atoms-1); start multi-atom # group ids past that range so a group id can never collide with a # singleton's leftover arange id (the torch.unique compaction below diff --git a/torchref/model/riding_xyz.py b/torchref/model/riding_xyz.py index 12771a67..7301825b 100644 --- a/torchref/model/riding_xyz.py +++ b/torchref/model/riding_xyz.py @@ -23,6 +23,7 @@ import numpy as np import torch +from torchref.config import get_int_dtype import torch.nn as nn from torchref.base.coordinates.local_frame import ( frame_is_degenerate, @@ -66,8 +67,8 @@ def _register_rows(self, n_full: int, frames: HydrogenFrames, device) -> None: raise ValueError(f"{name} must reference stored rows, not riding ones") long = dict( - dtype=torch.int64, device=device - ) # dtype-ok: row index buffers; int64 index required + dtype=get_int_dtype(), device=device + ) self.register_buffer("base_row", torch.as_tensor(base, **long)) self.register_buffer("h_row", torch.as_tensor(h, **long)) self.register_buffer( @@ -96,25 +97,25 @@ def _rebuild_row_cache(self) -> None: n_full = int(base.numel() + self.h_row.numel()) self._n_full = n_full full_to_base = torch.full( - (max(n_full, 1),), -1, dtype=torch.int64, device=device - ) # dtype-ok: index map; int64 + (max(n_full, 1),), -1, dtype=get_int_dtype(), device=device + ) full_to_base[base] = torch.arange( - base.numel(), dtype=torch.int64, device=device - ) # dtype-ok: index map; int64 + base.numel(), dtype=get_int_dtype(), device=device + ) self._parent_bidx = full_to_base[self.parent_row.clamp(min=0)].clamp(min=0) self._n1_bidx = full_to_base[self.n1_row.clamp(min=0)].clamp(min=0) self._n2_bidx = full_to_base[self.n2_row.clamp(min=0)].clamp(min=0) # ``cat([base, derived])[gather]`` lays the full table out in one gather. order = torch.empty( - n_full, dtype=torch.int64, device=device - ) # dtype-ok: gather index; int64 + n_full, dtype=get_int_dtype(), device=device + ) order[base] = torch.arange( - base.numel(), dtype=torch.int64, device=device - ) # dtype-ok: gather index; int64 + base.numel(), dtype=get_int_dtype(), device=device + ) order[self.h_row] = base.numel() + torch.arange( self.h_row.numel(), - dtype=torch.int64, - device=device, # dtype-ok: gather index; int64 + dtype=get_int_dtype(), + device=device, ) self._gather_order = order @@ -226,8 +227,8 @@ def __init__( self.register_buffer( buffer, torch.zeros( - 0, dtype=torch.int64, device=self.device - ), # dtype-ok: empty row-index buffer; int64 + 0, dtype=get_int_dtype(), device=self.device + ), ) self.register_buffer( "frame_valid", torch.zeros(0, dtype=torch.bool, device=self.device) @@ -261,8 +262,8 @@ def __init__( is_riding = np.zeros(n_full, dtype=bool) is_riding[np.asarray(frames.h_row, dtype=np.int64)] = True base_rows = torch.as_tensor( - np.nonzero(~is_riding)[0], dtype=torch.int64, device=device - ) # dtype-ok: row index; int64 + np.nonzero(~is_riding)[0], dtype=get_int_dtype(), device=device + ) if refinable_mask is None: base_mask = None @@ -380,8 +381,8 @@ def _orientation_selection(self, full_mask): parents = getattr(self, "_" + kind + "_parents") rows = getattr(self, "_" + kind + "_h") groups = getattr(self, "_" + kind + "_inverse") - selected = full_mask[parents].to(torch.int32) - selected.index_add_(0, groups, full_mask[self.h_row[rows]].to(torch.int32)) + selected = full_mask[parents].to(get_int_dtype()) + selected.index_add_(0, groups, full_mask[self.h_row[rows]].to(get_int_dtype())) selections.append(selected > 0) return selections diff --git a/torchref/model/rigid_xyz.py b/torchref/model/rigid_xyz.py index 0d0a304a..a22fee39 100644 --- a/torchref/model/rigid_xyz.py +++ b/torchref/model/rigid_xyz.py @@ -27,7 +27,7 @@ from torch import nn from torchref.base.alignment.rotation import rotation_matrix_euler_xyz -from torchref.config import get_float_dtype, normalize_device +from torchref.config import get_float_dtype, get_int_dtype, normalize_device from torchref.utils.caching import CachedForwardMixin from torchref.utils.device_mixin import DeviceMixin @@ -73,7 +73,7 @@ def __init__( dtype = dtype if dtype is not None else get_float_dtype() self.register_buffer("original_xyz", torch.empty(0, 3, device=device, dtype=dtype)) self.register_buffer( - "chain_indices", torch.empty(0, dtype=torch.long, device=device) # dtype-ok: empty chain_indices buffer; indexing requires long + "chain_indices", torch.empty(0, dtype=get_int_dtype(), device=device) ) self.register_buffer("chain_centers", torch.empty(0, 3, device=device, dtype=dtype)) self.register_buffer( diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index 96ee5799..0d9dcac1 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -8,7 +8,7 @@ import torch from torch.nn import Module as nnModule -from torchref.config import normalize_device +from torchref.config import get_int_dtype, normalize_device from torchref.io import ReflectionData from torchref.model.model_ft import ModelFT from torchref.refinement.logger import Logger @@ -441,7 +441,7 @@ def mark(idx): return # 4. freeze xyz of those atoms (same path as freeze_selection) - model.xyz_mask[torch.tensor(freeze_idx, dtype=torch.long)] = False # dtype-ok: freeze index used to index xyz_mask; PyTorch requires int64 + model.xyz_mask[torch.tensor(freeze_idx, dtype=get_int_dtype())] = False model.apply_mask_to_parameter("xyz") if self.verbose > 0: shown = frozen_res[:20] + (["..."] if len(frozen_res) > 20 else []) diff --git a/torchref/refinement/model_error_estimation/sigma_a.py b/torchref/refinement/model_error_estimation/sigma_a.py index a24f9285..f3b294b5 100644 --- a/torchref/refinement/model_error_estimation/sigma_a.py +++ b/torchref/refinement/model_error_estimation/sigma_a.py @@ -22,7 +22,7 @@ import torch -from torchref.config import get_float_dtype +from torchref.config import get_float_dtype, get_int_dtype def epsilon_from_hkl(hkl: torch.Tensor, spacegroup) -> torch.Tensor: @@ -263,10 +263,10 @@ def _segment_layout(lengths: Tuple[int, ...], device_str: str): ``lengths`` is a tuple so it can be a cache key. """ device = torch.device(device_str) - L = torch.tensor(lengths, dtype=torch.long, device=device) # dtype-ok: segment lengths for cumsum offsets/gather index; PyTorch requires int64 + L = torch.tensor(lengths, dtype=get_int_dtype(), device=device) total = int(L.sum()) max_len = int(L.max()) if L.numel() else 0 - zero = torch.zeros(1, dtype=torch.long, device=device) # dtype-ok: zero offset concatenated into gather index; PyTorch requires int64 + zero = torch.zeros(1, dtype=get_int_dtype(), device=device) starts = torch.cat([zero, L.cumsum(0)[:-1]]) ar = torch.arange(max_len, device=device).reshape(1, max_len) # Clamp keeps the gather in bounds for the padding slots; `mask` zeroes them anyway. diff --git a/torchref/refinement/optimizers/curvature.py b/torchref/refinement/optimizers/curvature.py index 74b5bbae..6a485cc7 100644 --- a/torchref/refinement/optimizers/curvature.py +++ b/torchref/refinement/optimizers/curvature.py @@ -22,6 +22,7 @@ from typing import Callable, Optional, Sequence import torch +from torchref.config import get_int_dtype from torchref.utils import use_portable @@ -36,7 +37,7 @@ def _sample_probe( """Draw one Hutchinson probe vector of length ``numel``.""" if probe == "rademacher": r = torch.randint( - 0, 2, (numel,), generator=generator, device=device, dtype=torch.int64 # dtype-ok: randint {0,1} bernoulli draw, immediately cast to float dtype; width irrelevant + 0, 2, (numel,), generator=generator, device=device, dtype=get_int_dtype() ) return r.to(dtype).mul_(2.0).sub_(1.0) # {0,1} -> {-1,+1} if probe == "gaussian": diff --git a/torchref/refinement/targets/adp/rigid_bond.py b/torchref/refinement/targets/adp/rigid_bond.py index 07a639a7..2e1d3068 100644 --- a/torchref/refinement/targets/adp/rigid_bond.py +++ b/torchref/refinement/targets/adp/rigid_bond.py @@ -2,6 +2,7 @@ import numpy as np import torch +from torchref.config import get_int_dtype from typing import TYPE_CHECKING, Dict from torchref.base.targets.adp import adp_rigid_bond_aniso_math @@ -122,7 +123,7 @@ def _bond_pairs(self) -> torch.Tensor: chunks.append(idx_) if chunks: return torch.cat(chunks, dim=0).contiguous() - return torch.empty(0, 2, dtype=torch.long, device=self.model.xyz().device) # dtype-ok: empty (0,2) atom-pair index tensor; PyTorch requires int64 + return torch.empty(0, 2, dtype=get_int_dtype(), device=self.model.xyz().device) def _compute_aniso_rigid_bond(self) -> torch.Tensor: """Rigid-bond NLL from ``Δz = l^T U_1 l - l^T U_2 l`` along each bond. diff --git a/torchref/refinement/targets/adp/similarity.py b/torchref/refinement/targets/adp/similarity.py index 191d2379..dde02d0f 100644 --- a/torchref/refinement/targets/adp/similarity.py +++ b/torchref/refinement/targets/adp/similarity.py @@ -1,5 +1,6 @@ import numpy as np import torch +from torchref.config import get_int_dtype from typing import TYPE_CHECKING, Dict from torchref.base.targets.adp import adp_simu_math, adp_simu_aniso_math @@ -94,7 +95,7 @@ def _get_pair_indices(self) -> torch.Tensor: if chunks: cached = torch.cat(chunks, dim=0).contiguous() else: - cached = torch.empty(0, 2, dtype=torch.long, # dtype-ok: empty (0,2) atom-pair index tensor; PyTorch requires int64 + cached = torch.empty(0, 2, dtype=get_int_dtype(), device=self.model.xyz().device) self._simu_pair_indices_cache = cached return cached diff --git a/torchref/refinement/targets/difference.py b/torchref/refinement/targets/difference.py index d0a75ab7..19a7b971 100644 --- a/torchref/refinement/targets/difference.py +++ b/torchref/refinement/targets/difference.py @@ -9,6 +9,7 @@ """ import torch +from torchref.config import get_int_dtype from torch import nn from typing import TYPE_CHECKING, Dict, Literal, Optional, Tuple @@ -204,10 +205,10 @@ def _match_reflections(self): device = hkl_light.device self._matched_indices_light = torch.tensor( - matched_light, dtype=torch.long, device=device # dtype-ok: matched atom indices used for indexing; PyTorch requires int64 + matched_light, dtype=get_int_dtype(), device=device ) self._matched_indices_dark = torch.tensor( - matched_dark, dtype=torch.long, device=device # dtype-ok: matched atom indices used for indexing; PyTorch requires int64 + matched_dark, dtype=get_int_dtype(), device=device ) # Store common HKL (using light indices, they should be identical) diff --git a/torchref/refinement/targets/geometry/chiral.py b/torchref/refinement/targets/geometry/chiral.py index 1cc6dcee..3dcd007c 100644 --- a/torchref/refinement/targets/geometry/chiral.py +++ b/torchref/refinement/targets/geometry/chiral.py @@ -1,5 +1,6 @@ import numpy as np import torch +from torchref.config import get_int_dtype from typing import TYPE_CHECKING, Dict from torchref.base.targets.chiral import chiral_math @@ -82,7 +83,7 @@ def get_violations(self, threshold: float = 0.5) -> Dict[str, torch.Tensor]: if "chiral" not in self.restraints.restraints: return { - "indices": torch.tensor([], dtype=torch.long, device=device).reshape( # dtype-ok: empty restraint index tensor; PyTorch requires int64 for indexing + "indices": torch.tensor([], dtype=get_int_dtype(), device=device).reshape( 0, 4 ), "volumes": torch.tensor([], device=device), diff --git a/torchref/refinement/targets/geometry/non_bonded.py b/torchref/refinement/targets/geometry/non_bonded.py index 637be377..10b27e98 100644 --- a/torchref/refinement/targets/geometry/non_bonded.py +++ b/torchref/refinement/targets/geometry/non_bonded.py @@ -7,6 +7,7 @@ import numpy as np import torch +from torchref.config import get_int_dtype from typing import TYPE_CHECKING, Dict, Tuple from torchref.utils.stats import ( @@ -346,7 +347,7 @@ def get_violations(self, threshold: float = 0.0) -> Dict[str, torch.Tensor]: if "vdw" not in self.restraints.restraints: return { - "indices": torch.tensor([], dtype=torch.long, device=device).reshape( # dtype-ok: empty restraint index tensor; PyTorch requires int64 for indexing + "indices": torch.tensor([], dtype=get_int_dtype(), device=device).reshape( 0, 2 ), "violations": torch.tensor([], device=device), @@ -359,7 +360,7 @@ def get_violations(self, threshold: float = 0.0) -> Dict[str, torch.Tensor]: if indices is None or len(indices) == 0: return { - "indices": torch.tensor([], dtype=torch.long, device=device).reshape( # dtype-ok: empty restraint index tensor; PyTorch requires int64 for indexing + "indices": torch.tensor([], dtype=get_int_dtype(), device=device).reshape( 0, 2 ), "violations": torch.tensor([], device=device), diff --git a/torchref/refinement/targets/similarity.py b/torchref/refinement/targets/similarity.py index d01c3ef8..ab190c86 100644 --- a/torchref/refinement/targets/similarity.py +++ b/torchref/refinement/targets/similarity.py @@ -6,6 +6,7 @@ """ import torch +from torchref.config import get_int_dtype from typing import TYPE_CHECKING, Dict from .base import Target @@ -68,10 +69,10 @@ def __init__( # path (the one ``load_state_dict`` uses) would have no such buffers at all. # ``_build_atom_map`` overwrites them rather than creating them. self.register_buffer( - "_idx_dark", torch.zeros(0, dtype=torch.long, device=self.device) # dtype-ok: index buffer for gather/index_select; PyTorch requires int64 + "_idx_dark", torch.zeros(0, dtype=get_int_dtype(), device=self.device) ) self.register_buffer( - "_idx_light", torch.zeros(0, dtype=torch.long, device=self.device) # dtype-ok: index buffer for gather/index_select; PyTorch requires int64 + "_idx_light", torch.zeros(0, dtype=get_int_dtype(), device=self.device) ) if model_dark is not None and model_light is not None: self._build_atom_map() @@ -140,10 +141,10 @@ def _build_atom_map(self): "dark and light models" ) self.register_buffer( - "_idx_dark", torch.zeros(0, dtype=torch.long, device=self.device) # dtype-ok: index buffer for gather/index_select; PyTorch requires int64 + "_idx_dark", torch.zeros(0, dtype=get_int_dtype(), device=self.device) ) self.register_buffer( - "_idx_light", torch.zeros(0, dtype=torch.long, device=self.device) # dtype-ok: index buffer for gather/index_select; PyTorch requires int64 + "_idx_light", torch.zeros(0, dtype=get_int_dtype(), device=self.device) ) return @@ -166,13 +167,13 @@ def _build_atom_map(self): self.register_buffer( "_idx_dark", torch.tensor( - merged["_idx_dark"].values, dtype=torch.long, device=self.device # dtype-ok: atom index tensor used for indexing; PyTorch requires int64 + merged["_idx_dark"].values, dtype=get_int_dtype(), device=self.device ), ) self.register_buffer( "_idx_light", torch.tensor( - merged["_idx_light"].values, dtype=torch.long, device=self.device # dtype-ok: atom index tensor used for indexing; PyTorch requires int64 + merged["_idx_light"].values, dtype=get_int_dtype(), device=self.device ), ) diff --git a/torchref/scaling/collection_scaler.py b/torchref/scaling/collection_scaler.py index 2b7d918d..89759f5a 100644 --- a/torchref/scaling/collection_scaler.py +++ b/torchref/scaling/collection_scaler.py @@ -14,7 +14,7 @@ import torch.nn as nn from torchref.base.metrics.rfactor import rfactor_work_free -from torchref.config import get_float_dtype +from torchref.config import get_float_dtype, get_int_dtype from torchref.scaling.scaler_base import ( DEFAULT_SCALE_TARGET, SCALE_TARGETS, @@ -192,7 +192,7 @@ def _calc_initial_scale_joint(self): pos_mask = torch.ones_like(fobs, dtype=torch.bool) mask = (work_mask & pos_mask).to(torch.bool) - bins = self.bins[mask].to(torch.int64) # dtype-ok: bin indices for scatter/index_select; PyTorch requires int64 + bins = self.bins[mask].to(torch.int64) # dtype-ok: scatter_add index; int64 required on torch < 2.8 log_ratios = ( torch.log(fobs_clamped[mask]) - torch.log(fcalc_amp[mask]) ).to(self.device) @@ -205,7 +205,7 @@ def _calc_initial_scale_joint(self): per_bin = scales / (counts + 1e-6) with torch.no_grad(): - target = per_bin.detach()[self.bins.to(torch.int64)] # dtype-ok: bin indices for advanced indexing; PyTorch requires int64 + target = per_bin.detach()[self.bins.to(get_int_dtype())] design = self._iso_design.to(target.dtype) coeff = torch.linalg.lstsq(design, target.unsqueeze(1)).solution.squeeze(1) self.c_iso = nn.Parameter(coeff.detach()) diff --git a/torchref/scaling/scaler_base.py b/torchref/scaling/scaler_base.py index 689a49f5..52e63352 100644 --- a/torchref/scaling/scaler_base.py +++ b/torchref/scaling/scaler_base.py @@ -21,7 +21,7 @@ rfactor_work_free, ) from torchref.base.reciprocal import get_scattering_vectors -from torchref.config import get_complex_dtype, get_float_dtype +from torchref.config import get_complex_dtype, get_float_dtype, get_int_dtype from torchref.utils.autograd_ops import gather_with_index_add from torchref.utils.debug_utils import DebugMixin from torchref.utils.device_mixin import DeviceMixin @@ -264,7 +264,7 @@ def calc_initial_scale(self, fcalc: torch.Tensor): initial_log_scale.detach().cpu().numpy(), ) with torch.no_grad(): - target = initial_log_scale.detach().to(self.device)[self.bins.to(torch.int64)] # dtype-ok: bin indices for advanced indexing; PyTorch requires int64 + target = initial_log_scale.detach().to(self.device)[self.bins.to(get_int_dtype())] design = self._iso_design.to(target.dtype) coeff = torch.linalg.lstsq(design, target.unsqueeze(1)).solution.squeeze(1) self.c_iso = nn.Parameter(coeff.detach().to(self.device)) @@ -394,7 +394,7 @@ def get_binwise_mean_intensity(self, fcalc: torch.Tensor): mean_calc_intensity = torch.zeros(self.nbins, device=self.device, dtype=fobs.dtype) counts = torch.zeros(self.nbins, device=self.device, dtype=fobs.dtype) counts_vals = torch.ones_like(F_calc, device=self.device, dtype=fobs.dtype) - bins_sel = self.bins.to(torch.int64)[sel] # dtype-ok: bin indices for advanced indexing; PyTorch requires int64 + bins_sel = self.bins.to(torch.int64)[sel] # dtype-ok: scatter_add index; int64 required on torch < 2.8 mean_obs_intensity = torch.scatter_add( mean_obs_intensity, 0, bins_sel, intensities[sel] ) diff --git a/torchref/scaling/solvent.py b/torchref/scaling/solvent.py index 87f906ee..c7a4e0d6 100644 --- a/torchref/scaling/solvent.py +++ b/torchref/scaling/solvent.py @@ -10,7 +10,7 @@ get_scattering_vectors, ifft, ) -from torchref.config import get_float_dtype +from torchref.config import get_float_dtype, get_int_dtype from torchref.utils.debug_utils import DebugMixin from torchref.utils.device_mixin import DeviceMixin from torchref.utils.device_resolution import resolve_device @@ -387,7 +387,7 @@ def get_solvent_mask(self): # grids, where the SF code's 1024 would OOM (denser intermediates). ATOM_CHUNK = 256 - grid_dims = torch.tensor(grid_shape, dtype=torch.long, device=device) # dtype-ok: grid dims for voxel index arithmetic; PyTorch requires int64 + grid_dims = torch.tensor(grid_shape, dtype=get_int_dtype(), device=device) grid_shape_float = grid_dims.float() inv_grid = 1.0 / grid_shape_float G = frac.T @ frac # metric tensor: r²_cart = diff_frac · G · diff_frac @@ -464,12 +464,12 @@ def get_solvent_mask(self): protein_voxels = ( torch.cat(protein_chunks, dim=0) if protein_chunks - else torch.empty((0, 3), dtype=torch.long, device=device) # dtype-ok: empty (0,3) voxel index tensor; PyTorch requires int64 for indexing + else torch.empty((0, 3), dtype=get_int_dtype(), device=device) ) boundary_voxels = ( torch.cat(boundary_chunks, dim=0) if boundary_chunks - else torch.empty((0, 3), dtype=torch.long, device=device) # dtype-ok: empty (0,3) voxel index tensor; PyTorch requires int64 for indexing + else torch.empty((0, 3), dtype=get_int_dtype(), device=device) ) del protein_chunks, boundary_chunks diff --git a/torchref/scaling/wilson.py b/torchref/scaling/wilson.py index e16b048d..4f5e6116 100644 --- a/torchref/scaling/wilson.py +++ b/torchref/scaling/wilson.py @@ -32,7 +32,7 @@ import torch -from torchref.config import get_float_dtype +from torchref.config import get_float_dtype, get_int_dtype from torchref.scaling.basis import chebyshev_design __all__ = ["WilsonNormaliser"] @@ -455,7 +455,7 @@ def from_hkl( branches feed two different parameters of the same likelihood. """ work = get_float_dtype() - hkl_l = hkl.to(torch.long) # dtype-ok: Miller indices are integers + hkl_l = hkl.to(get_int_dtype()) # The cell may carry the configured default device while the reflections # are somewhere else; the caller should not have to reconcile them. rec = cell.reciprocal_basis_matrix.to(device=hkl_l.device, dtype=work) diff --git a/torchref/symmetry/map_symmetry.py b/torchref/symmetry/map_symmetry.py index afbeb4a9..43d26523 100644 --- a/torchref/symmetry/map_symmetry.py +++ b/torchref/symmetry/map_symmetry.py @@ -20,6 +20,7 @@ from __future__ import annotations import torch +from torchref.config import get_int_dtype from torchref.utils.device_mixin import DeviceMixin @@ -149,7 +150,7 @@ def _index_grid(self, op_index: int) -> torch.Tensor: transformed = transformed - torch.floor(transformed) shape_t = torch.tensor([nx, ny, nz], dtype=dtype, device=device) - indices = torch.round(transformed * shape_t).to(torch.int64) # dtype-ok: rounded voxel grid indices; int64 index tensor required + indices = torch.round(transformed * shape_t).to(get_int_dtype()) indices[:, 0] %= nx indices[:, 1] %= ny indices[:, 2] %= nz diff --git a/torchref/symmetry/reciprocal_symmetry.py b/torchref/symmetry/reciprocal_symmetry.py index f595ee70..d5b05c25 100644 --- a/torchref/symmetry/reciprocal_symmetry.py +++ b/torchref/symmetry/reciprocal_symmetry.py @@ -24,7 +24,7 @@ import numpy as np import torch -from torchref.config import get_float_dtype +from torchref.config import get_float_dtype, get_int_dtype @@ -82,7 +82,7 @@ def _expand_hkl( for i in range(n_ops): # h' = h @ R^T hkl_transformed = torch.round(torch.matmul(hkl_float, recip_matrices[i].T)).to( - torch.int32 # dtype-ok: transformed Miller indices (hkl); fixed-width int32 representation + get_int_dtype() ) # Phase shift from translation: -2π h·t, for h' = hR under the convention # F(h) = Σ_j f_j exp(+2πi h·x_j). Do NOT "simplify" the sign: the wrong sign @@ -124,10 +124,10 @@ def _expand_hkl( # Build output tensors expanded_hkl = torch.tensor( - [list(k) for k in unique_dict.keys()], dtype=torch.int32, device=device # dtype-ok: unique Miller indices (hkl); fixed-width int32 representation + [list(k) for k in unique_dict.keys()], dtype=get_int_dtype(), device=device ) phase_shifts = torch.tensor(unique_phases, dtype=get_float_dtype(), device=device) - orig_idx_tensor = torch.tensor(orig_indices, dtype=torch.int64, device=device) # dtype-ok: reflection index mapping; int64 index tensor required + orig_idx_tensor = torch.tensor(orig_indices, dtype=get_int_dtype(), device=device) if remove_absences and sym.number != 1: keep_mask = ~sym.is_absent(expanded_hkl) @@ -169,7 +169,7 @@ def _complete_hkl( ------- complete_hkl : torch.Tensor, shape (M, 3), dtype int32 All possible Miller indices within resolution (minus systematic absences). - input_indices : torch.Tensor, shape (M,), dtype int64 + input_indices : torch.Tensor, shape (M,), integer dtype Index mapping complete → input, or -1 where missing. Use as ``F_complete[~missing] = F_input[input_indices[~missing]]``. missing_mask : torch.Tensor, shape (M,), dtype bool @@ -198,7 +198,7 @@ def _complete_hkl( all_hkl_np = all_hkl.cpu().numpy() n_complete = len(all_hkl) - input_indices = torch.full((n_complete,), -1, dtype=torch.int64, device=device) # dtype-ok: reflection index buffer (-1 sentinel); int64 index required + input_indices = torch.full((n_complete,), -1, dtype=get_int_dtype(), device=device) missing_mask = torch.ones(n_complete, dtype=torch.bool, device=device) for i, hkl in enumerate(all_hkl_np): @@ -274,7 +274,7 @@ def get_canonical_hkl(hkl_single): for i in range(n_ops): # h' = h @ R^T hkl_trans = torch.round(torch.matmul(hkl_single, recip_matrices[i].T)).to( - torch.int32 # dtype-ok: transformed Miller indices (hkl); fixed-width int32 representation + get_int_dtype() ) equivalents.append(hkl_trans) @@ -311,7 +311,7 @@ def get_canonical_hkl(hkl_single): R = recip_matrices[equiv_idx] t = translations[equiv_idx] - hkl_trans = torch.round(torch.matmul(hkl_single, R.T)).to(torch.int32) # dtype-ok: transformed Miller indices (hkl); fixed-width int32 representation + hkl_trans = torch.round(torch.matmul(hkl_single, R.T)).to(get_int_dtype()) # -2π h·t, same convention as expand_hkl (see the derivation there). phase_shift = -2.0 * np.pi * torch.matmul(hkl_single, t) @@ -333,9 +333,9 @@ def get_canonical_hkl(hkl_single): asu_list = sorted(asu_reflections.keys()) n_asu = len(asu_list) - hkl_asu = torch.tensor(asu_list, dtype=torch.int32, device=device) # dtype-ok: ASU Miller indices (hkl); fixed-width int32 representation + hkl_asu = torch.tensor(asu_list, dtype=get_int_dtype(), device=device) reduction_indices = torch.full( - (n_asu, n_equiv), -1, dtype=torch.int64, device=device # dtype-ok: reduction index map (-1 sentinel); int64 index tensor required + (n_asu, n_equiv), -1, dtype=get_int_dtype(), device=device ) phase_shifts = torch.zeros((n_asu, n_equiv), dtype=get_float_dtype(), device=device) @@ -446,7 +446,7 @@ def _canonicalize_hkl( empty_hkl = torch.empty((0, 3), dtype=hkl_dtype, device=device) empty_f = torch.empty(0, dtype=get_float_dtype(), device=device) empty_b = torch.empty(0, dtype=torch.bool, device=device) - empty_i = torch.empty(0, dtype=torch.int64, device=device) # dtype-ok: empty index tensor; int64 index dtype required + empty_i = torch.empty(0, dtype=get_int_dtype(), device=device) return empty_hkl, empty_f, empty_b, empty_i # The ASU lookup tables are numpy-backed, so the operations come across to CPU @@ -550,9 +550,9 @@ def _canonicalize_hkl( h_max = int(canonical_hkl.abs().max().item()) + 1 base = 2 * h_max + 1 sort_key = ( - canonical_hkl[:, 0].to(torch.int64) * base * base # dtype-ok: linear HKL hash/key; int64 avoids overflow for indexing - + canonical_hkl[:, 1].to(torch.int64) * base # dtype-ok: linear HKL hash/key; int64 avoids overflow for indexing - + canonical_hkl[:, 2].to(torch.int64) # dtype-ok: linear HKL hash/key; int64 avoids overflow for indexing + canonical_hkl[:, 0].to(torch.int64) * base * base # dtype-ok: composite sort key h*base^2+k*base+l overflows int32 for large Miller indices + + canonical_hkl[:, 1].to(torch.int64) * base # dtype-ok: composite sort key h*base^2+k*base+l overflows int32 for large Miller indices + + canonical_hkl[:, 2].to(torch.int64) # dtype-ok: composite sort key h*base^2+k*base+l overflows int32 for large Miller indices ) sort_indices = torch.argsort(sort_key) diff --git a/torchref/symmetry/symmetry.py b/torchref/symmetry/symmetry.py index 4cb767d8..7bf79e0c 100644 --- a/torchref/symmetry/symmetry.py +++ b/torchref/symmetry/symmetry.py @@ -29,7 +29,7 @@ import torch -from torchref.config import get_float_dtype +from torchref.config import get_float_dtype, get_int_dtype from torchref.utils.device_mixin import DeviceMixin if TYPE_CHECKING: @@ -352,11 +352,11 @@ def expand_reciprocal(self, hkl: torch.Tensor) -> torch.Tensor: Returns ------- torch.Tensor - Shape ``(n_ops, N, 3)``, rounded to ``int64``. Rounding is exact for valid + Shape ``(n_ops, N, 3)``, rounded to the configured int dtype. Rounding is exact for valid operations on integer indices and only mops up float error. """ equivalents = self.reciprocal.apply_rotations(hkl) - return torch.round(equivalents).to(torch.int64) # dtype-ok: rounded Miller equivalents; int64 for exact integer compare/index + return torch.round(equivalents).to(get_int_dtype()) # ========================================================================= # Reflection predicates @@ -381,7 +381,7 @@ def is_centric(self, hkl: torch.Tensor) -> torch.Tensor: with torch.no_grad(): flat = hkl.reshape(-1, 3) equivalents = self.expand_reciprocal(flat) # (n_ops, N, 3) - target = -flat.to(device=equivalents.device, dtype=torch.int64) # dtype-ok: compare target for int64 equivalents; dtype must match + target = -flat.to(device=equivalents.device, dtype=get_int_dtype()) centric = (equivalents == target).all(dim=-1).any(dim=0) return centric.reshape(original_shape).to(hkl.device) @@ -405,7 +405,7 @@ def is_absent(self, hkl: torch.Tensor) -> torch.Tensor: with torch.no_grad(): flat = hkl.reshape(-1, 3) equivalents = self.expand_reciprocal(flat) # (n_ops, N, 3) - target = flat.to(device=equivalents.device, dtype=torch.int64) # dtype-ok: compare target for int64 equivalents; dtype must match + target = flat.to(device=equivalents.device, dtype=get_int_dtype()) maps_to_self = (equivalents == target).all(dim=-1) # (n_ops, N) h_dot_t = torch.matmul( @@ -465,7 +465,7 @@ def epsilon(self, hkl: torch.Tensor, *, friedel: bool = True) -> torch.Tensor: float_dtype = get_float_dtype() with torch.no_grad(): equivalents = self.expand_reciprocal(hkl) # (n_ops, N, 3) - target = hkl.to(device=equivalents.device, dtype=torch.int64) # dtype-ok: compare target for int64 equivalents; dtype must match + target = hkl.to(device=equivalents.device, dtype=get_int_dtype()) fixes = (equivalents == target).all(dim=-1) if friedel: fixes = fixes | (equivalents == -target).all(dim=-1) @@ -494,7 +494,7 @@ def grid_requirements(self) -> dict: # ``Fraction(float)`` would need a tolerance where this is exact. numerators = torch.round( self.translations.detach().cpu().double() * _TRANSLATION_DENOMINATOR - ).to(torch.int64) # dtype-ok: integer translation numerators for exact Fraction recovery + ).to(get_int_dtype()) for op_numerators in numerators.tolist(): for axis, numerator in enumerate(op_numerators): diff --git a/torchref/topology/atom_graph.py b/torchref/topology/atom_graph.py index 05204315..ba50e04d 100644 --- a/torchref/topology/atom_graph.py +++ b/torchref/topology/atom_graph.py @@ -16,6 +16,7 @@ import numpy as np import torch +from torchref.config import get_int_dtype from torchref.topology.edges import EdgeBlock from torchref.utils.device_mixin import DeviceMixin @@ -42,8 +43,8 @@ def _build_csr(bonds: torch.Tensor, n_atoms: int) -> Tuple[torch.Tensor, torch.T device = bonds.device if bonds.numel() == 0: return ( - torch.zeros(n_atoms + 1, dtype=torch.int64, device=device), # dtype-ok: CSR indptr offset array; int64 index required - torch.zeros(0, dtype=torch.int64, device=device), # dtype-ok: empty CSR neighbor index array; int64 index required + torch.zeros(n_atoms + 1, dtype=get_int_dtype(), device=device), + torch.zeros(0, dtype=get_int_dtype(), device=device), ) src = torch.cat([bonds[:, 0], bonds[:, 1]]) @@ -55,9 +56,9 @@ def _build_csr(bonds: torch.Tensor, n_atoms: int) -> Tuple[torch.Tensor, torch.T src, dst = pairs[:, 0], pairs[:, 1] counts = torch.bincount(src, minlength=n_atoms) - indptr = torch.zeros(n_atoms + 1, dtype=torch.int64, device=device) # dtype-ok: CSR indptr offset array; int64 index required + indptr = torch.zeros(n_atoms + 1, dtype=get_int_dtype(), device=device) torch.cumsum(counts, dim=0, out=indptr[1:]) - return indptr, dst.to(torch.int64) # dtype-ok: CSR neighbor (dst) index array; int64 index required + return indptr, dst.to(get_int_dtype()) def _extend_paths( @@ -81,13 +82,13 @@ def _extend_paths( """ device = paths.device if paths.numel() == 0: - return torch.zeros((0, paths.shape[1] + 1), dtype=torch.int64, device=device) # dtype-ok: empty BFS path index array; int64 index required + return torch.zeros((0, paths.shape[1] + 1), dtype=get_int_dtype(), device=device) last, prev = paths[:, -1], paths[:, -2] counts = indptr[last + 1] - indptr[last] total = int(counts.sum()) if total == 0: - return torch.zeros((0, paths.shape[1] + 1), dtype=torch.int64, device=device) # dtype-ok: empty BFS path index array; int64 index required + return torch.zeros((0, paths.shape[1] + 1), dtype=get_int_dtype(), device=device) row = torch.repeat_interleave(torch.arange(len(paths), device=device), counts) # Offset of each slot within its own neighbour list. @@ -204,16 +205,16 @@ def implicit_h_count(self) -> Optional[torch.Tensor]: return None is_h = self.is_hydrogen bonds = self.bonds.indices - present = torch.zeros(self.n_atoms, dtype=torch.int64, device=bonds.device) # dtype-ok: bincount output; int64 + present = torch.zeros(self.n_atoms, dtype=get_int_dtype(), device=bonds.device) if bonds.numel(): heavy_of_h = torch.cat( [bonds[is_h[bonds[:, 1]] & ~is_h[bonds[:, 0]], 0], bonds[is_h[bonds[:, 0]] & ~is_h[bonds[:, 1]], 1]] ) if heavy_of_h.numel(): - present = torch.bincount(heavy_of_h, minlength=self.n_atoms) + present = torch.bincount(heavy_of_h, minlength=self.n_atoms).to(present.dtype) known = self.template_h_count >= 0 - missing = self.template_h_count.to(torch.int64) - present + missing = self.template_h_count - present return torch.where(known, missing.clamp(min=0), torch.zeros_like(missing)) def subset(self, remap: torch.Tensor, residue_remap: torch.Tensor) -> "AtomGraph": diff --git a/torchref/topology/build.py b/torchref/topology/build.py index 87970297..150a80c0 100644 --- a/torchref/topology/build.py +++ b/torchref/topology/build.py @@ -11,6 +11,7 @@ import numpy as np import pandas as pd import torch +from torchref.config import get_int_dtype from torchref.topology.builders import ( InterResidueAngleBuilder, @@ -711,7 +712,7 @@ def _block_with_values( per_origin, arity, edge_type, payload ) block = EdgeBlock( - indices=torch.as_tensor(indices, dtype=torch.int64, device=device), # dtype-ok: atom index tensor for restraint edges; int64 index required + indices=torch.as_tensor(indices, dtype=get_int_dtype(), device=device), origin_bounds=bounds, ) values = { @@ -997,7 +998,7 @@ def build_topology_with_values( np.arange(n_res, dtype=np.int64), nodes["atom_end"] - nodes["atom_start"], ), - dtype=torch.int64, # dtype-ok: atom index tensor; int64 index required + dtype=get_int_dtype(), device=device, ), bonds=bond_block, @@ -1007,7 +1008,7 @@ def build_topology_with_values( planes=plane_blocks, energy_type=energy_type, template_h_count=torch.as_tensor( - template_h_count, dtype=torch.int8, device=device + template_h_count, dtype=get_int_dtype(), device=device ), ) diff --git a/torchref/topology/builders.py b/torchref/topology/builders.py index 14502625..3cb0ab87 100644 --- a/torchref/topology/builders.py +++ b/torchref/topology/builders.py @@ -613,7 +613,7 @@ def build( sigmas = np.where(sigmas == 0, 1e-4, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -720,7 +720,7 @@ def build( sigmas = np.where(sigmas == 0, 1e-4, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -841,7 +841,7 @@ def build( sigmas = np.where(sigmas == 0, 1e-4, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), "periods": torch.tensor(periods, dtype=get_int_dtype(), device=device), @@ -932,7 +932,7 @@ def build( key = f"{n_atoms}_atoms" result[key] = { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -1054,7 +1054,7 @@ def build( sigmas = np.where(sigmas == 0, 1e-4, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), "ideal_volumes": torch.tensor( ideal_volumes, dtype=get_float_dtype(), device=device ), @@ -1309,7 +1309,7 @@ def finalize( sigmas = np.where(sigmas == 0, min_sigma, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -1412,7 +1412,7 @@ def build( sigmas = np.where(sigmas == 0, 1e-4, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -1534,7 +1534,7 @@ def finalize( sigmas = np.where(sigmas == 0, min_sigma, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -1648,7 +1648,7 @@ def build( sigmas = np.where(sigmas == 0, 1e-4, sigmas) return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -1785,7 +1785,7 @@ def finalize_disulfide( periods = periods[sort_order] return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), "references": torch.tensor(references, dtype=get_float_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), "periods": torch.tensor(periods, dtype=get_int_dtype(), device=device), @@ -1952,7 +1952,7 @@ def build( indices = indices[order] periods = periods[order] result["phi"] = { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), "periods": torch.tensor(periods, dtype=get_int_dtype(), device=device), } @@ -1965,7 +1965,7 @@ def build( indices = indices[order] periods = periods[order] result["psi"] = { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), "periods": torch.tensor(periods, dtype=get_int_dtype(), device=device), } @@ -1984,7 +1984,7 @@ def build( periods = periods[order] is_proline = is_proline[order] result["omega"] = { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), "references": torch.tensor( references, dtype=get_float_dtype(), device=device ), @@ -2023,13 +2023,13 @@ def build( stypes = stypes[order] result["ramachandran"] = { "phi_indices": torch.tensor( - phi_idx, dtype=torch.long, device=device # dtype-ok: phi atom-index tensor for dihedral; int64 required + phi_idx, dtype=get_int_dtype(), device=device ), "psi_indices": torch.tensor( - psi_idx, dtype=torch.long, device=device # dtype-ok: psi atom-index tensor for dihedral; int64 required + psi_idx, dtype=get_int_dtype(), device=device ), "surface_type": torch.tensor( - stypes, dtype=torch.long, device=device # dtype-ok: categorical rama surface-type code used as advanced index; int64 + stypes, dtype=get_int_dtype(), device=device ), } @@ -2128,7 +2128,7 @@ def build( key = f"{n_atoms}_atoms" result[key] = { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 + "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } diff --git a/torchref/topology/edges.py b/torchref/topology/edges.py index c70a5f2d..b99288b9 100644 --- a/torchref/topology/edges.py +++ b/torchref/topology/edges.py @@ -16,6 +16,7 @@ import numpy as np import torch +from torchref.config import get_int_dtype from torchref.utils.device_mixin import DeviceMixin @@ -153,7 +154,7 @@ class EdgeBlock(DeviceMixin): def empty(cls, arity: int, device=None) -> "EdgeBlock": """An edge-free block of the given arity.""" return cls( - indices=torch.zeros((0, arity), dtype=torch.int64, device=device), # dtype-ok: empty edge index tensor (0,arity); int64 index required + indices=torch.zeros((0, arity), dtype=get_int_dtype(), device=device), origin_bounds={}, ) @@ -190,7 +191,7 @@ def from_origins( if len(indices) == 0: return cls.empty(arity, device=device) return cls( - indices=torch.as_tensor(indices, dtype=torch.int64, device=device), # dtype-ok: edge atom index tensor; int64 index required + indices=torch.as_tensor(indices, dtype=get_int_dtype(), device=device), origin_bounds=bounds, ) diff --git a/torchref/topology/hydrogens.py b/torchref/topology/hydrogens.py index 5ec85303..0f5e3ddd 100644 --- a/torchref/topology/hydrogens.py +++ b/torchref/topology/hydrogens.py @@ -28,7 +28,7 @@ import numpy as np import torch -from torchref.config import get_float_dtype +from torchref.config import get_float_dtype, get_int_dtype #: Standard heavy-atom valences, one of the two budgets that cap how many hydrogens a #: parent may take. Elements not listed fall back to 4 and are then bounded only by the @@ -1146,17 +1146,17 @@ def to_tensors(self, device=None) -> Dict[str, torch.Tensor]: """Return frame and orientation arrays as tensors, keyed by field name.""" return { "h_row": torch.as_tensor( - self.h_row, dtype=torch.int64, device=device - ), # dtype-ok: row index; int64 required + self.h_row, dtype=get_int_dtype(), device=device + ), "parent_row": torch.as_tensor( - self.parent_row, dtype=torch.int64, device=device - ), # dtype-ok: row index; int64 required + self.parent_row, dtype=get_int_dtype(), device=device + ), "n1_row": torch.as_tensor( - self.n1_row, dtype=torch.int64, device=device - ), # dtype-ok: row index; int64 required + self.n1_row, dtype=get_int_dtype(), device=device + ), "n2_row": torch.as_tensor( - self.n2_row, dtype=torch.int64, device=device - ), # dtype-ok: row index; int64 required + self.n2_row, dtype=get_int_dtype(), device=device + ), "frame_valid": torch.as_tensor( self.frame_valid, dtype=torch.bool, device=device ), diff --git a/torchref/topology/nonbonded.py b/torchref/topology/nonbonded.py index 404e7a6f..236f6f85 100644 --- a/torchref/topology/nonbonded.py +++ b/torchref/topology/nonbonded.py @@ -15,7 +15,7 @@ import numpy as np import torch -from torchref.config import dtypes, get_float_dtype +from torchref.config import dtypes, get_float_dtype, get_int_dtype if TYPE_CHECKING: from torchref.symmetry.cell import Cell @@ -85,8 +85,8 @@ def prefilter_symop_offsets( valid_ops.append(op_idx) valid_offsets.append([dx, dy, dz]) - op_indices = torch.tensor(valid_ops, dtype=torch.long, device=device) # dtype-ok: symmetry-operator index tensor; int64 - cell_offsets = torch.tensor(valid_offsets, dtype=torch.long, device=device) # dtype-ok: integer cell-offset lattice vectors; symmetry-image metadata + op_indices = torch.tensor(valid_ops, dtype=get_int_dtype(), device=device) + cell_offsets = torch.tensor(valid_offsets, dtype=get_int_dtype(), device=device) return op_indices, cell_offsets @@ -146,7 +146,7 @@ def assign_to_grid( gd = grid_dims.to(device=device, dtype=fdtype) cell_ijk = (frac_wrapped * gd[None, None, :]).long() cell_ijk = cell_ijk.clamp( - min=torch.zeros(3, dtype=torch.long, device=device), # dtype-ok: clamp min-bound for long grid-index tensor; matches int64 + min=torch.zeros(3, dtype=get_int_dtype(), device=device), max=(grid_dims - 1).to(device), ) @@ -188,14 +188,14 @@ def build_cell_list( unique_cells, counts = torch.unique_consecutive( sorted_cells, return_counts=True ) - starts = torch.zeros(len(unique_cells) + 1, dtype=torch.long, device=device) # dtype-ok: CSR boundary/offset array; int64 required + starts = torch.zeros(len(unique_cells) + 1, dtype=get_int_dtype(), device=device) starts[1:] = counts.cumsum(0) cell_lookup = torch.full( - (n_grid_total,), -1, dtype=torch.long, device=device # dtype-ok: grid-cell to index lookup table; used for indexing, int64 + (n_grid_total,), -1, dtype=get_int_dtype(), device=device ) cell_lookup[unique_cells] = torch.arange( - len(unique_cells), dtype=torch.long, device=device # dtype-ok: index values written into lookup table; int64 + len(unique_cells), dtype=get_int_dtype(), device=device ) return sort_order, unique_cells, starts, cell_lookup @@ -233,7 +233,7 @@ def _get_canonical_offsets_14(device: torch.device) -> torch.Tensor: offsets.append([dx, dy, dz]) assert len(offsets) == 14, f"expected 14 canonical offsets, got {len(offsets)}" _NEIGHBOR_OFFSETS_14 = torch.tensor( - offsets, dtype=torch.long, device=device # dtype-ok: grid neighbor-cell offset deltas used to compute index; int64 + offsets, dtype=get_int_dtype(), device=device ) return _NEIGHBOR_OFFSETS_14 @@ -493,7 +493,7 @@ def find_pairs_periodic_grid_v2( all_pair_combo_j.append(cj) if not all_pair_atom_i: - empty = torch.tensor([], dtype=torch.long, device=device) # dtype-ok: empty atom-pair index placeholder; int64 required + empty = torch.tensor([], dtype=get_int_dtype(), device=device) return empty, empty, empty return ( @@ -517,11 +517,11 @@ def exclusion_set_to_hash( Hash: min(i,j) * max_idx + max(i,j), sorted for searchsorted. """ if not exclusion_set: - return torch.tensor([], dtype=torch.long, device=device) # dtype-ok: empty exclusion-hash placeholder; int64 + return torch.tensor([], dtype=torch.long, device=device) # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 arr = np.array(list(exclusion_set), dtype=np.int64) hashes = arr[:, 0] * max_idx + arr[:, 1] # already (min, max) hashes.sort() - return torch.tensor(hashes, dtype=torch.long, device=device) # dtype-ok: packed pair-hash key for searchsorted; int64 avoids overflow + return torch.tensor(hashes, dtype=torch.long, device=device) # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 def filter_pairs( @@ -629,11 +629,11 @@ def build_vdw_restraints_gpu( sg = SG(sg) empty_result = { - "indices": torch.zeros(0, 2, dtype=torch.long, device=device), # dtype-ok: atom-pair index tensor; torch indexing requires int64 + "indices": torch.zeros(0, 2, dtype=get_int_dtype(), device=device), "min_distances": torch.zeros(0, dtype=get_float_dtype(), device=device), "sigmas": torch.zeros(0, dtype=get_float_dtype(), device=device), - "symop_indices": torch.zeros(0, dtype=torch.long, device=device), # dtype-ok: symmetry-operator index tensor; int64 - "cell_offsets": torch.zeros(0, 3, dtype=torch.long, device=device), # dtype-ok: integer cell-offset lattice vectors; symmetry-image metadata + "symop_indices": torch.zeros(0, dtype=get_int_dtype(), device=device), + "cell_offsets": torch.zeros(0, 3, dtype=get_int_dtype(), device=device), } # Step 1: prefilter symop combos @@ -655,10 +655,10 @@ def build_vdw_restraints_gpu( if len(identity_indices) == 0: # Identity not in valid combos — should not happen, but add it op_indices = torch.cat([ - torch.zeros(1, dtype=torch.long, device=device), op_indices # dtype-ok: identity prepended to symop-index tensor; int64 + torch.zeros(1, dtype=get_int_dtype(), device=device), op_indices ]) cell_offsets_valid = torch.cat([ - torch.zeros(1, 3, dtype=torch.long, device=device), cell_offsets_valid # dtype-ok: identity prepended to cell-offset tensor; int64 + torch.zeros(1, 3, dtype=get_int_dtype(), device=device), cell_offsets_valid ]) identity_combo = 0 M = len(op_indices) @@ -774,7 +774,7 @@ def build_vdw_restraints_gpu( "valid_op_indices": op_indices, "valid_cell_offsets": cell_offsets_valid, "grid_dims": grid_dims, - "identity_combo": torch.tensor(identity_combo, dtype=torch.long, device=device), # dtype-ok: combo index scalar into symop/offset arrays; int64 + "identity_combo": torch.tensor(identity_combo, dtype=get_int_dtype(), device=device), } if verbose > 0: @@ -856,7 +856,7 @@ def find_h_vdw_pairs_gpu( xyz_all = torch.cat([xyz_heavy, xyz_h], dim=0) # (N_all, 3) n_all = xyz_all.shape[0] - empty = torch.tensor([], dtype=torch.long, device=device) # dtype-ok: empty index placeholder tensor; int64 required + empty = torch.tensor([], dtype=get_int_dtype(), device=device) if n_all == 0: return empty, empty, empty diff --git a/torchref/topology/residue_graph.py b/torchref/topology/residue_graph.py index 0a68a21c..5231f02d 100644 --- a/torchref/topology/residue_graph.py +++ b/torchref/topology/residue_graph.py @@ -14,6 +14,7 @@ import numpy as np import torch +from torchref.config import get_int_dtype #: SG-SG separation below which two cysteines are taken to be disulfide-bonded. DISULFIDE_MAX_DISTANCE = 2.5 @@ -272,7 +273,7 @@ def find_disulfide_links( rows = list(sg_rows) if len(rows) < 2: return [] - idx = torch.as_tensor(rows, dtype=torch.int64, device=xyz.device) # dtype-ok: residue-atom index tensor; int64 index required + idx = torch.as_tensor(rows, dtype=get_int_dtype(), device=xyz.device) dist = torch.cdist(xyz[idx], xyz[idx]) close = (dist > DISULFIDE_MIN_DISTANCE) & (dist < DISULFIDE_MAX_DISTANCE) diff --git a/torchref/topology/restraint_sets.py b/torchref/topology/restraint_sets.py index 45c079ae..a3ae4320 100644 --- a/torchref/topology/restraint_sets.py +++ b/torchref/topology/restraint_sets.py @@ -17,7 +17,7 @@ import numpy as np import torch -from torchref.config import get_float_dtype +from torchref.config import get_float_dtype, get_int_dtype #: Origins making up each edge type's ``all`` group -- what the geometry targets read. #: ``None`` means every origin present. ``phi`` and ``psi`` are conformationally free @@ -41,7 +41,7 @@ def to_tensor(values, prop: str, device=None) -> torch.Tensor: if isinstance(values, torch.Tensor): return values.to(device=device) if device is not None else values if prop in _INTEGER_PROPERTIES: - dtype = torch.int64 # dtype-ok: dtype var for index tensors; int64 index required + dtype = get_int_dtype() elif prop in _BOOL_PROPERTIES: dtype = torch.bool else: diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index 0c49868f..f591009d 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -31,7 +31,7 @@ read_cif, read_link_definitions, ) -from torchref.config import get_float_dtype +from torchref.config import get_float_dtype, get_int_dtype from torchref.utils.debug_utils import DebugMixin from torchref.utils.device_mixin import DeviceMixin @@ -408,7 +408,7 @@ def _find_nearby_pairs_spatial_hash(self, xyz, cutoff=6.0): n_atoms = xyz.shape[0] if n_atoms == 0: - return torch.tensor([], dtype=torch.long, device=device).reshape(0, 2) # dtype-ok: empty atom-pair index tensor; int64 index required + return torch.tensor([], dtype=get_int_dtype(), device=device).reshape(0, 2) # Work on CPU to avoid per-iteration GPU kernel launch overhead coords = xyz.detach().cpu() @@ -433,12 +433,12 @@ def _find_nearby_pairs_spatial_hash(self, xyz, cutoff=6.0): sorted_flat, return_counts=True ) n_unique = len(unique_cells) - starts = torch.zeros(n_unique + 1, dtype=torch.long) # dtype-ok: grid-cell CSR start offsets; int64 index required + starts = torch.zeros(n_unique + 1, dtype=get_int_dtype()) starts[1:] = counts.cumsum(0) # Lookup: flat_cell -> index in unique_cells (-1 if empty) n_grid = gx * gyz - cell_lookup = torch.full((n_grid,), -1, dtype=torch.long) # dtype-ok: cell lookup table (-1 sentinel); int64 index required + cell_lookup = torch.full((n_grid,), -1, dtype=get_int_dtype()) cell_lookup[unique_cells] = torch.arange(n_unique) # 14 unique neighbour offsets: self (0,0,0) + 13 forward neighbours. @@ -524,9 +524,9 @@ def _find_nearby_pairs_spatial_hash(self, xyz, cutoff=6.0): if pair_chunks: all_pairs = np.concatenate(pair_chunks, axis=0) - return torch.from_numpy(all_pairs).to(dtype=torch.long, device=device) # dtype-ok: atom-pair index array from numpy; int64 index required + return torch.from_numpy(all_pairs).to(dtype=get_int_dtype(), device=device) else: - return torch.tensor([], dtype=torch.long, device=device).reshape(0, 2) # dtype-ok: empty atom-pair index tensor; int64 index required + return torch.tensor([], dtype=get_int_dtype(), device=device).reshape(0, 2) def _expand_with_symmetry_mates(self, xyz, cutoff): """Append symmetry-mate positions to ASU ``xyz`` for neighbour search. @@ -659,7 +659,7 @@ def _build_h_exclusion_hash(self, h_topo, device): ``torch.searchsorted`` lookup. """ if h_topo is None or h_topo.n_hydrogens == 0: - return torch.tensor([], dtype=torch.long, device=device) # dtype-ok: empty index tensor; int64 index required + return torch.tensor([], dtype=torch.long, device=device) # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 n_heavy = len(self.pdb) n_h = h_topo.n_hydrogens @@ -684,13 +684,13 @@ def _build_h_exclusion_hash(self, h_topo, device): exclusions.add((min(h_combined, nb), max(h_combined, nb))) if not exclusions: - return torch.tensor([], dtype=torch.long, device=device) # dtype-ok: empty index tensor; int64 index required + return torch.tensor([], dtype=torch.long, device=device) # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 arr = np.array(list(exclusions), dtype=np.int64) max_idx = max(n_heavy + n_h, int(arr.max()) + 1) hashes = arr[:, 0] * max_idx + arr[:, 1] hashes.sort() - return torch.tensor(hashes, dtype=torch.long, device=device) # dtype-ok: grid-cell hash values used as keys/index; int64 required + return torch.tensor(hashes, dtype=torch.long, device=device) # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 def _build_vdw_restraints( self, cutoff=6.0, sigma=0.2, inter_residue_only=True, use_spatial_hash=True @@ -893,17 +893,17 @@ def _build_vdw_restraints_legacy( if dist_sq < cutoff_sq: pairs_list.append([i, j]) nearby_pairs = ( - torch.tensor(pairs_list, dtype=torch.long, device=device) # dtype-ok: atom-pair index tensor; int64 index required + torch.tensor(pairs_list, dtype=get_int_dtype(), device=device) if pairs_list - else torch.tensor([], dtype=torch.long, device=device).reshape(0, 2) # dtype-ok: empty atom-pair index tensor; int64 index required + else torch.tensor([], dtype=get_int_dtype(), device=device).reshape(0, 2) ) empty_result = { - "indices": torch.tensor([], dtype=torch.long, device=device).reshape(0, 2), # dtype-ok: empty atom-pair index tensor; int64 index required + "indices": torch.tensor([], dtype=get_int_dtype(), device=device).reshape(0, 2), "min_distances": torch.tensor([], dtype=get_float_dtype(), device=device), "sigmas": torch.tensor([], dtype=get_float_dtype(), device=device), - "symop_indices": torch.tensor([], dtype=torch.long, device=device), # dtype-ok: empty symop index tensor; int64 index required - "cell_offsets": torch.tensor([], dtype=torch.long, device=device).reshape(0, 3), # dtype-ok: empty cell-offset index tensor; int64 index required + "symop_indices": torch.tensor([], dtype=get_int_dtype(), device=device), + "cell_offsets": torch.tensor([], dtype=get_int_dtype(), device=device).reshape(0, 3), } if len(nearby_pairs) == 0: @@ -1035,7 +1035,7 @@ def _build_vdw_restraints_legacy( # Store results final_pairs = np.stack([final_i1, final_i2], axis=1) self._vdw = { - "indices": torch.tensor(final_pairs, dtype=torch.long, device=device), # dtype-ok: final atom-pair index tensor; int64 index required + "indices": torch.tensor(final_pairs, dtype=get_int_dtype(), device=device), "min_distances": torch.tensor( min_distances, dtype=get_float_dtype(), device=device ), @@ -1043,10 +1043,10 @@ def _build_vdw_restraints_legacy( (len(final_pairs),), sigma, dtype=get_float_dtype(), device=device ), "symop_indices": torch.tensor( - final_symop, dtype=torch.long, device=device # dtype-ok: symop index tensor; int64 index required + final_symop, dtype=get_int_dtype(), device=device ), "cell_offsets": torch.tensor( - final_offsets, dtype=torch.long, device=device # dtype-ok: cell-offset index tensor; int64 index required + final_offsets, dtype=get_int_dtype(), device=device ), } diff --git a/torchref/topology/riding.py b/torchref/topology/riding.py index f36a681c..894476e3 100644 --- a/torchref/topology/riding.py +++ b/torchref/topology/riding.py @@ -22,7 +22,7 @@ import numpy as np import torch -from torchref.config import dtypes, normalize_device +from torchref.config import dtypes, get_int_dtype, normalize_device from torchref.utils.device_resolution import resolve_device from torchref.utils.device_mixin import DeviceMixin @@ -431,17 +431,17 @@ def build_hydrogen_topology( fdtype = dtypes.float if n_h_total == 0: - topo.h_parent_idx = torch.zeros(0, dtype=torch.long, device=device) # dtype-ok: parent atom-index tensor (empty); int64 required + topo.h_parent_idx = torch.zeros(0, dtype=get_int_dtype(), device=device) topo.h_bond_length = torch.zeros(0, dtype=fdtype, device=device) topo.h_vdw_radius = torch.zeros(0, dtype=fdtype, device=device) - topo.h_placement_type = torch.zeros(0, dtype=torch.long, device=device) # dtype-ok: categorical H placement-type code (empty) - topo.h_slot_in_parent = torch.zeros(0, dtype=torch.long, device=device) # dtype-ok: slot index into parent (empty); int64 + topo.h_placement_type = torch.zeros(0, dtype=get_int_dtype(), device=device) + topo.h_slot_in_parent = torch.zeros(0, dtype=get_int_dtype(), device=device) topo.parent_neighbor_idx = torch.zeros( - 0, MAX_HEAVY_NB, dtype=torch.long, device=device # dtype-ok: parent neighbor atom-index tensor (empty); int64 required + 0, MAX_HEAVY_NB, dtype=get_int_dtype(), device=device ) - topo.parent_neighbor_count = torch.zeros(0, dtype=torch.long, device=device) # dtype-ok: per-parent neighbor count (empty); structural int - topo.h_chainid_enc = torch.zeros(0, dtype=torch.long, device=device) # dtype-ok: categorical chain-id encoding (empty) - topo.h_resseq = torch.zeros(0, dtype=torch.long, device=device) # dtype-ok: residue sequence id (empty); categorical + topo.parent_neighbor_count = torch.zeros(0, dtype=get_int_dtype(), device=device) + topo.h_chainid_enc = torch.zeros(0, dtype=get_int_dtype(), device=device) + topo.h_resseq = torch.zeros(0, dtype=get_int_dtype(), device=device) return topo # Sort all topology arrays by placement type for contiguous slicing @@ -466,22 +466,22 @@ def build_hydrogen_topology( idxs = np.where(mask)[0] type_bounds[t] = (int(idxs[0]), int(idxs[-1]) + 1) - topo.h_parent_idx = torch.tensor(acc_parent_idx, dtype=torch.long, device=device) # dtype-ok: parent atom-index tensor; torch indexing requires int64 + topo.h_parent_idx = torch.tensor(acc_parent_idx, dtype=get_int_dtype(), device=device) topo.h_bond_length = torch.tensor(acc_bond_length, dtype=fdtype, device=device) topo.h_vdw_radius = torch.full((n_h_total,), 1.20, dtype=fdtype, device=device) topo.h_placement_type = torch.tensor( - acc_placement_type, dtype=torch.long, device=device # dtype-ok: categorical H placement-type code; used for sort/slice + acc_placement_type, dtype=get_int_dtype(), device=device ) - topo.h_slot_in_parent = torch.tensor(acc_slot, dtype=torch.long, device=device) # dtype-ok: slot index into parent neighbor slots; int64 + topo.h_slot_in_parent = torch.tensor(acc_slot, dtype=get_int_dtype(), device=device) topo.parent_neighbor_idx = torch.tensor( - np.stack(acc_nb_idx), dtype=torch.long, device=device # dtype-ok: parent neighbor atom-index tensor; int64 required + np.stack(acc_nb_idx), dtype=get_int_dtype(), device=device ) topo.parent_neighbor_count = torch.tensor( - acc_nb_count, dtype=torch.long, device=device # dtype-ok: per-parent neighbor count; structural int metadata + acc_nb_count, dtype=get_int_dtype(), device=device ) topo.type_bounds = type_bounds # dict: type_code -> (start, end) - topo.h_chainid_enc = torch.tensor(acc_chainid_enc, dtype=torch.long, device=device) # dtype-ok: categorical chain-id encoding - topo.h_resseq = torch.tensor(acc_resseq, dtype=torch.long, device=device) # dtype-ok: residue sequence id; categorical + topo.h_chainid_enc = torch.tensor(acc_chainid_enc, dtype=get_int_dtype(), device=device) + topo.h_resseq = torch.tensor(acc_resseq, dtype=get_int_dtype(), device=device) if verbose > 0: print(f" Hydrogen topology: {n_h_total} riding H atoms") @@ -736,8 +736,8 @@ def build_h_candidate_pairs( if n_h == 0: for name in ("cand_idx_i", "cand_idx_j", "cand_symop_idx"): - setattr(h_topo, name, torch.zeros(0, dtype=torch.long, device=device)) # dtype-ok: candidate atom/symop index tensors (empty); int64 required - h_topo.cand_cell_offset = torch.zeros(0, 3, dtype=torch.long, device=device) # dtype-ok: integer cell-offset lattice vectors (empty); symmetry metadata + setattr(h_topo, name, torch.zeros(0, dtype=get_int_dtype(), device=device)) + h_topo.cand_cell_offset = torch.zeros(0, 3, dtype=get_int_dtype(), device=device) h_topo.cand_min_dist = torch.zeros(0, dtype=dtypes.float, device=device) return @@ -836,15 +836,15 @@ def _same_res(chain_a, resseq_a, chain_b, resseq_b): if not acc_idx_i: for name in ("cand_idx_i", "cand_idx_j", "cand_symop_idx"): - setattr(h_topo, name, torch.zeros(0, dtype=torch.long, device=device)) # dtype-ok: candidate atom/symop index tensors (empty); int64 required - h_topo.cand_cell_offset = torch.zeros(0, 3, dtype=torch.long, device=device) # dtype-ok: integer cell-offset lattice vectors (empty); symmetry metadata + setattr(h_topo, name, torch.zeros(0, dtype=get_int_dtype(), device=device)) + h_topo.cand_cell_offset = torch.zeros(0, 3, dtype=get_int_dtype(), device=device) h_topo.cand_min_dist = torch.zeros(0, dtype=dtypes.float, device=device) return - cand_i = torch.tensor(acc_idx_i, dtype=torch.long, device=device) # dtype-ok: combined atom-index tensor; torch indexing requires int64 - cand_j = torch.tensor(acc_idx_j, dtype=torch.long, device=device) # dtype-ok: combined atom-index tensor; torch indexing requires int64 - cand_sym = torch.tensor(acc_symop, dtype=torch.long, device=device) # dtype-ok: symmetry-operator index; int64 - cand_off = torch.tensor(np.stack(acc_offset), dtype=torch.long, device=device) # dtype-ok: integer cell-offset lattice vectors; symmetry-image metadata + cand_i = torch.tensor(acc_idx_i, dtype=get_int_dtype(), device=device) + cand_j = torch.tensor(acc_idx_j, dtype=get_int_dtype(), device=device) + cand_sym = torch.tensor(acc_symop, dtype=get_int_dtype(), device=device) + cand_off = torch.tensor(np.stack(acc_offset), dtype=get_int_dtype(), device=device) # Apply 1-2 / 1-3 exclusions for intra-ASU candidates if h_excl_hash is not None and len(h_excl_hash) > 0: @@ -853,7 +853,7 @@ def _same_res(chain_a, resseq_a, chain_b, resseq_b): max_idx = n_heavy + n_h norm_i = torch.minimum(cand_i, cand_j) norm_j = torch.maximum(cand_i, cand_j) - pair_hash = norm_i * max_idx + norm_j + pair_hash = norm_i.to(torch.int64) * max_idx + norm_j.to(torch.int64) # dtype-ok: packed pair key overflows int32; searchsorted needs int64 like the table ins = torch.searchsorted(h_excl_hash, pair_hash).clamp( max=len(h_excl_hash) - 1 ) diff --git a/torchref/topology/topology.py b/torchref/topology/topology.py index 51419991..516e4731 100644 --- a/torchref/topology/topology.py +++ b/torchref/topology/topology.py @@ -13,6 +13,7 @@ import numpy as np import torch +from torchref.config import get_int_dtype from torchref.topology.atom_graph import AtomGraph from torchref.topology.residue_graph import ResidueGraph @@ -89,7 +90,7 @@ def subset(self, keep) -> "Topology": mask = torch.as_tensor(keep) if mask.dtype != torch.bool: selected = torch.zeros(self.n_atoms, dtype=torch.bool) - selected[mask.to(torch.int64)] = True # dtype-ok: boolean-mask->index cast for scatter select; int64 index required + selected[mask.to(get_int_dtype())] = True mask = selected mask = mask.to(device=self.atoms.residue_of.device) @@ -97,8 +98,8 @@ def subset(self, keep) -> "Topology": raise ValueError("subset would keep no atoms") n_kept = int(mask.sum()) - remap = torch.full((self.n_atoms,), -1, dtype=torch.int64, device=mask.device) # dtype-ok: atom remap index array (-1 sentinel); int64 index required - remap[mask] = torch.arange(n_kept, dtype=torch.int64, device=mask.device) # dtype-ok: arange remap indices; int64 index required + remap = torch.full((self.n_atoms,), -1, dtype=get_int_dtype(), device=mask.device) + remap[mask] = torch.arange(n_kept, dtype=get_int_dtype(), device=mask.device) # A residue survives if any of its atoms does. Counting per residue also # gives the new atom ranges, contiguous because the atom order is unchanged. @@ -112,10 +113,10 @@ def subset(self, keep) -> "Topology": atom_start = atom_end - counts residue_remap = torch.full( - (self.n_residues,), -1, dtype=torch.int64, device=mask.device # dtype-ok: residue remap index array (-1 sentinel); int64 index required + (self.n_residues,), -1, dtype=get_int_dtype(), device=mask.device ) residue_remap[torch.as_tensor(residue_keep, device=mask.device)] = torch.arange( - int(residue_keep.sum()), dtype=torch.int64, device=mask.device # dtype-ok: arange residue remap indices; int64 index required + int(residue_keep.sum()), dtype=get_int_dtype(), device=mask.device ) return Topology( From 2422bbab266e60213b21527d5f4986d8db1d7ae4 Mon Sep 17 00:00:00 2001 From: Claude Date: Fri, 18 Sep 2026 11:42:33 +0000 Subject: [PATCH 174/250] Write configured-int lookup tables from sources of the same dtype ``dest[index] = source`` checks that the two dtypes match, unlike slice assignment, so an int64 ``arange`` or ``unique`` inverse can no longer be written into a table that now carries the configured int dtype. The Friedel-mate maps, the grid-cell lookup, the spherical-harmonic row map and the atom remap used by hydrogen generation build their sources in the table's dtype instead. The Legendre shell kernel TORCH_CHECKs int64 shell labels and offsets, so those two tensors keep int64 and their markers say why. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- torchref/experimental/alignment/frf/data_mr.py | 2 +- .../experimental/alignment/frf/kernels/cpu/legendre_shell.py | 5 ++--- torchref/experimental/alignment/sh.py | 2 +- torchref/io/datasets/reflection_data.py | 2 +- torchref/model/model.py | 4 ++-- torchref/topology/restraints.py | 2 +- 6 files changed, 8 insertions(+), 9 deletions(-) diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index 986f13bd..bec9f93a 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -380,7 +380,7 @@ def _group_mean(values, index, n_groups): # index to meet device values. inv_s = inv_s.to(device) n_shells = int(uniq_ks.shape[0]) - shell_of_cluster = torch.zeros(n_clusters, dtype=get_int_dtype(), device=device) + shell_of_cluster = torch.zeros(n_clusters, dtype=torch.int64, device=device) # dtype-ok: the legendre_shell kernel TORCH_CHECKs int64 shell labels shell_of_cluster[inverse] = inv_s shell_smag = _group_mean(s_mag_all.to(comp_real), inv_s, n_shells) diff --git a/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py b/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py index 6a59efd6..32f0c268 100644 --- a/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py +++ b/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py @@ -32,7 +32,6 @@ from typing import Optional, Tuple import torch -from torchref.config import get_int_dtype from torchref.base.electron_density.kernels.cpu._cpp_build import build_extension @@ -242,12 +241,12 @@ def clear_cache() -> None: def shell_offsets(shell: torch.Tensor, n_shells: int) -> torch.Tensor: """Start index of each shell in a shell-sorted cluster array, plus the end. - ``(n_shells + 1,)`` integer offsets. The kernel needs the ranges rather than the + ``(n_shells + 1,)`` int64. The kernel needs the ranges rather than the per-cluster labels so that a thread can own a set of shells outright and write their accumulator rows without atomics. """ counts = torch.bincount(shell, minlength=n_shells) - offsets = torch.zeros(n_shells + 1, dtype=get_int_dtype(), device=shell.device) + offsets = torch.zeros(n_shells + 1, dtype=torch.int64, device=shell.device) # dtype-ok: the kernel TORCH_CHECKs int64 offsets torch.cumsum(counts, dim=0, out=offsets[1:]) return offsets diff --git a/torchref/experimental/alignment/sh.py b/torchref/experimental/alignment/sh.py index f343c886..87a779f9 100644 --- a/torchref/experimental/alignment/sh.py +++ b/torchref/experimental/alignment/sh.py @@ -138,7 +138,7 @@ def _bar_legendre_recurrence( rows = keep_l.to(device=device, dtype=get_int_dtype()) # l -> its position in the output, or -1 when it is not kept. where = torch.full((L,), -1, dtype=get_int_dtype(), device=device) - where[rows] = torch.arange(rows.numel(), device=device) + where[rows] = torch.arange(rows.numel(), device=device, dtype=where.dtype) where_list = where.tolist() out = torch.zeros((*batch_shape, rows.numel(), L), dtype=dtype, device=device) diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index 1b23355e..1b9d117f 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -2507,7 +2507,7 @@ def _build_anomalous_dataframe( uniq = hkl[self._group_representative_rows(inverse, M)] # The (+) member is the unconjugated row, (-) is the Friedel-flagged row. - arange = torch.arange(N) + arange = torch.arange(N, dtype=get_int_dtype()) plus_idx = torch.full((M,), -1, dtype=get_int_dtype()) minus_idx = torch.full((M,), -1, dtype=get_int_dtype()) # A Bijvoet mate only counts as present if it is a real, positive diff --git a/torchref/model/model.py b/torchref/model/model.py index bf547fb4..c835a495 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -3004,8 +3004,8 @@ def _complete_riding_waters(self, frames): source = torch.empty(len(augmented), dtype=get_int_dtype(), device=self.device) old_index = torch.as_tensor(old_rows, device=self.device) new_index = torch.as_tensor(new_rows, device=self.device) - source[old_index] = torch.arange(len(self.pdb), device=self.device) - source[new_index] = torch.as_tensor(plan.parent, device=self.device) + source[old_index] = torch.arange(len(self.pdb), device=self.device, dtype=source.dtype) + source[new_index] = torch.as_tensor(plan.parent, device=self.device, dtype=source.dtype) xyz = ( self.xyz.to_mixed_tensor() if hasattr(self.xyz, "to_mixed_tensor") diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index f591009d..5688b11f 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -439,7 +439,7 @@ def _find_nearby_pairs_spatial_hash(self, xyz, cutoff=6.0): # Lookup: flat_cell -> index in unique_cells (-1 if empty) n_grid = gx * gyz cell_lookup = torch.full((n_grid,), -1, dtype=get_int_dtype()) - cell_lookup[unique_cells] = torch.arange(n_unique) + cell_lookup[unique_cells] = torch.arange(n_unique, dtype=cell_lookup.dtype) # 14 unique neighbour offsets: self (0,0,0) + 13 forward neighbours. # "Forward" = first non-zero component is positive, avoiding double counting. From 560d3f95845981c8a5f30a225998601380621e9b Mon Sep 17 00:00:00 2001 From: Claude Date: Fri, 18 Sep 2026 11:43:34 +0000 Subject: [PATCH 175/250] Name compiled-kernel dtype contracts in the int dtype policy Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- AGENTS.md | 3 ++- docs/changelog.rst | 2 +- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index 7bb4902b..4c410f92 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -55,7 +55,8 @@ Practically: `index_select` and `index_add_` accept it. A literal int dtype survives only where torch or the arithmetic forces it, with a `# dtype-ok:` marker naming the constraint: `scatter`/`gather` indices (int64 on torch < 2.8), `index_copy_`/`index_fill_`/`one_hot` (int64 always), packed - keys such as `i * n + j` that overflow int32, and external-library contracts (TorchMD-Net). + keys such as `i * n + j` that overflow int32, compiled kernels that `TORCH_CHECK` a dtype + (the Legendre shell kernel), and external-library contracts (TorchMD-Net). - `torch.float64` *is* a supported configuration (`TORCHREF_DTYPE_FLOAT=float64`) used as an eager numerical reference and in gradient checks. Code must **work** in float64, must not **require** it, and must not silently downcast (see `tests/integration/test_dtype_config_float64.py`). diff --git a/docs/changelog.rst b/docs/changelog.rst index 58c0b92d..ccf157d1 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,7 +4,7 @@ Changelog Unreleased ---------- -- Integer and index tensors now take the configured int dtype (``get_int_dtype()``, ``TORCHREF_DTYPE_INT``, int32 by default) throughout the package. A hardcoded ``int64`` remains only where a torch op (``scatter``/``gather`` on torch < 2.8, ``index_copy_``), an int32 overflow, or an external library requires it, and each such site says which. +- Integer and index tensors now take the configured int dtype (``get_int_dtype()``, ``TORCHREF_DTYPE_INT``, int32 by default) throughout the package. A hardcoded ``int64`` remains only where a torch op (``scatter``/``gather`` on torch < 2.8, ``index_copy_``), a compiled kernel, an int32 overflow, or an external library requires it, and each such site says which. - Align difference-refinement tests with fixture ownership, integration placement and configured dtype/device conventions; share fresh collection setup and exercise noise statistics on deposited amplitudes. - Read scaled observations directly in subset and collection accessors, avoiding unused sigma/amplitude corrections and removing redundant internal forwarding helpers. - Remove one-off diagnostic scripts and consolidate difference-refinement regression tests while retaining numerical and output-format coverage. From b78a89ddbbeb25599c44e3cef5a1e700eeb59da3 Mon Sep 17 00:00:00 2001 From: Claude Date: Fri, 18 Sep 2026 11:50:46 +0000 Subject: [PATCH 176/250] Reindex anomalous HKL in the configured int dtype ```_reindex_per_reflection``` clones the new HKL table, which carries the configured int dtype, and writes the stored anomalous HKL rows into it. Under ``TORCHREF_DTYPE_INT=int64`` those rows were int32, and indexed assignment rejects the mismatch; this predates the dtype conversion and showed up in the int64 run of the reflection-data tests. Cast the rows to the table's dtype. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- torchref/io/datasets/reflection_data.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index 1b9d117f..729ae345 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -448,7 +448,7 @@ def _reindex_per_reflection( # Present rows keep their signed (anomalous) index; missing rows # fall back to the canonical reference HKL (never a 0,0,0 row). out = new_hkl.clone() - out[present] = val[src_idx] + out[present] = val[src_idx].to(out.dtype) else: fill = self._REINDEX_FILL.get(name, 0) out = torch.full( From e26303b0ba88027ea7ac6a653d5c545fefd188c0 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sat, 19 Sep 2026 12:40:08 +0200 Subject: [PATCH 177/250] Add sigma_D difference weights, difference_sd target and maps in e/A^3 Expected difference-power estimator sigma_D (model_error_estimation/sigma_d.py): per-shell mean(dF_obs^2) - mean(sigma^2) with a fitted F_dark^gamma amplitude law, signed shrinkage toward a decaying curve in d*^2 so a null dataset yields no power, Wiener weight S/(S + sigma^2) and alpha/beta_model with a difference model. Shared shell helpers in _shells.py; estimate_beta unchanged. Registered DED weight schemes (maps/ded_weights.py: none, inverse_variance, sigma_d; default inverse_variance) selected with --ded-weight on difference-map, difference-refine and validate-ded. The difference MTZ writes DF/SIGDF/PHDELWT with W_IVW, W_SD (mean one) and KSCALE; DELFWT is no longer written. mtz2map gains -cw/--column-weight, -ck/--column-scale and --units {sigma,electrons,raw}; ScalerBase.multiplicative_scale() provides the absolute scale. validate-ded reports unweighted, inverse-variance and sigma_D correlations side by side. New collection target difference_sd (CollectionDifferenceSigmaDTarget). Bayes extrapolation shrinks per resolution shell through sigma_D. Co-Authored-By: Claude Fable 5.1 --- docs/changelog.rst | 9 + docs/user_guide/cli.rst | 47 +- docs/user_guide/targets.rst | 10 + tests/helpers/device_cases.py | 1 + tests/integration/test_cli_ded_weights.py | 299 +++++++ tests/integration/test_cli_two_moment_mtz.py | 43 +- tests/unit/maps/test_ded_weights.py | 128 +++ tests/unit/maps/test_map_units.py | 51 ++ .../test_collection_sigma_d_target.py | 126 +++ .../refinement/test_collection_taxonomy.py | 4 +- tests/unit/refinement/test_shells.py | 127 +++ tests/unit/refinement/test_sigma_d.py | 358 ++++++++ .../unit/scaling/test_multiplicative_scale.py | 42 + torchref/cli/_common.py | 37 + torchref/cli/collection_difference_refine.py | 376 +++++++-- torchref/cli/difference_map.py | 29 +- torchref/cli/mtz2map.py | 93 +- torchref/cli/validate_ded.py | 144 +++- torchref/maps/ded_weights.py | 253 ++++++ torchref/maps/difference_map.py | 27 +- torchref/maps/map.py | 18 + .../model_error_estimation/_shells.py | 187 ++++ .../model_error_estimation/sigma_a.py | 60 +- .../model_error_estimation/sigma_d.py | 795 ++++++++++++++++++ torchref/refinement/targets/__init__.py | 2 + .../refinement/targets/collection/__init__.py | 4 + .../refinement/targets/collection/_specs.py | 13 +- .../refinement/targets/collection/_util.py | 13 + .../refinement/targets/collection/base.py | 19 + .../refinement/targets/collection/xray.py | 169 +++- torchref/scaling/scaler_base.py | 28 + 31 files changed, 3311 insertions(+), 201 deletions(-) create mode 100644 tests/integration/test_cli_ded_weights.py create mode 100644 tests/unit/maps/test_ded_weights.py create mode 100644 tests/unit/maps/test_map_units.py create mode 100644 tests/unit/refinement/test_collection_sigma_d_target.py create mode 100644 tests/unit/refinement/test_shells.py create mode 100644 tests/unit/refinement/test_sigma_d.py create mode 100644 tests/unit/scaling/test_multiplicative_scale.py create mode 100644 torchref/maps/ded_weights.py create mode 100644 torchref/refinement/model_error_estimation/_shells.py create mode 100644 torchref/refinement/model_error_estimation/sigma_d.py diff --git a/docs/changelog.rst b/docs/changelog.rst index a1cae177..48350ebb 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,15 @@ Changelog Unreleased ---------- +- ``torchref.difference-map``, ``torchref.difference-refine`` and ``torchref.validate-ded`` gain ``--ded-weight {sigma_d,inverse_variance,none}`` and ``--sigma-d-gamma``. The difference MTZ now carries the unweighted ``DF``/``SIGDF`` on ``PHDELWT`` with one mean-one weight column per scheme, ``W_SD`` and ``W_IVW`` (MTZ type W), and the observed-to-model scale ``KSCALE``; ``DELFWT`` is no longer written, build the map with ``torchref.mtz2map -csf DF -cw W_IVW -cphi PHDELWT``. Registered in ``torchref.maps.ded_weights`` +- Added the ``sigma_D`` estimator (``torchref.refinement.model_error_estimation.sigma_d``): the expected true difference power per resolution shell, ``mean(dF_obs^2) - mean(sigma^2)`` with a fitted ``F_dark^gamma`` amplitude law and DerSimonian-Laird shrinkage of the signed shell power toward a decaying exponential in ``d*^2`` (fitted on all shells, so a dataset without a difference yields no power instead of the positive half of its noise), giving the Wiener weight ``S/(S + sigma^2)`` and, with a difference model, ``alpha``/``beta_model``. Inverse-variance weights suppress the strong reflections whose difference power is 10-70x that of weak ones; on independent half-datasets a Wiener weight with the true power raised map agreement 1.2-1.8x in effective patterns. The single-dataset estimate inherits the calibration of the reported sigmas, and on the campaign TD1 data (sigmas ~1.5x too large at high resolution) it emptied 60-90 % of the shells, so inverse variance stays the default; ``sigma_d`` reports its clamped-shell count and falls back to inverse variance with a warning when every shell is empty +- Added the ``difference_sd`` collection target (``CollectionDifferenceSigmaDTarget``): the difference Gaussian centred on ``alpha * dF_calc`` with variance ``beta_model + sigma_diff^2`` from ``sigma_D`` fitted on the free set. Selected with ``torchref.difference-refine --difference-target difference_sd``; ``difference`` stays the default +- ``torchref.mtz2map`` gains ``--column-weight``/``-cw`` (multiply the amplitudes by a weight column before the FFT) and ``--units {sigma,electrons,raw}``: ``electrons`` writes e/A^3 as ``(1/V) sum_h F(h) exp(-2 pi i h.x)`` with the amplitudes divided by the ``--column-scale``/``-ck`` factor (``KSCALE`` by default). ``-n``/``--normalize`` is a deprecated alias. ``Map`` and ``DifferenceMap`` take ``units`` too, and ``DifferenceMap`` an optional per-reflection ``scale`` +- ``ScalerBase.multiplicative_scale()`` returns the per-reflection ``K_overall * b_overall * anisotropy`` factor, every multiplicative component of ``forward`` and none of the additive solvent term +- ``torchref.validate-ded`` reports the real- and reciprocal-space correlations for unweighted, inverse-variance and ``sigma_D`` weights side by side (``by_weight`` in the JSON, a table in the summary) and records when ``sigma_D`` fell back to inverse variance +- The extrapolated-map Bayes shrinkage estimates its signal variance per resolution shell through ``sigma_D`` instead of one global ``tau^2``; ``tau_sq`` in the summary is now the count-weighted shell mean +- Inverse-variance difference weights floor ``sigma_diff`` at a tenth of its median, so a zero sigma yields a large finite weight rather than an infinite one +- Shell construction, segment sums, interpolation and the shrinkage line fit shared by ``sigma_A`` and ``sigma_D`` live in ``torchref.refinement.model_error_estimation._shells``; ``estimate_beta`` is unchanged - Align difference-refinement tests with fixture ownership, integration placement and configured dtype/device conventions; share fresh collection setup and exercise noise statistics on deposited amplitudes. - Read scaled observations directly in subset and collection accessors, avoiding unused sigma/amplitude corrections and removing redundant internal forwarding helpers. - Remove one-off diagnostic scripts and consolidate difference-refinement regression tests while retaining numerical and output-format coverage. diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index 11784112..171f6e90 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -91,7 +91,9 @@ restraints. ``-dsf``/``--dark-structure-factor``, ``-lsf``/``--light-structure-factor``, ``--fraction`` (light-state population fraction, singular), ``--weight-schedule`` annealing schedule (default ``5,3,2``), -``-n``/``--n-cycles`` macro-cycles. +``-n``/``--n-cycles`` macro-cycles, ``--difference-target {difference,difference_sd}`` +(the difference row the schedule drives; default ``difference``), ``--ded-weight`` and +``--sigma-d-gamma`` for the difference MTZ (see ``torchref.difference-map``). :API: :mod:`torchref.cli.collection_difference_refine` @@ -106,10 +108,17 @@ columns, expands to P1, and computes a real-space map via FFT. .. code-block:: bash - torchref.mtz2map -f refined.mtz -F 2FOFCWT -P PH2FOFCWT -o map.ccp4 + torchref.mtz2map -sf refined.mtz -csf 2FOFCWT -cphi PH2FOFCWT -o map.ccp4 + torchref.mtz2map -sf diff.mtz -csf DF -cw W_IVW -cphi PHDELWT -o diff.ccp4 + torchref.mtz2map -sf diff.mtz -csf DF -cw W_SD -cphi PHDELWT --units electrons -o diff_e.ccp4 -**Key options:** ``--high-res``, ``--low-res`` resolution limits, -``--gridsize`` override, ``-n`` normalize to sigma units. +**Key options:** ``--dmin``/``--dmax`` resolution limits, ``--gridsize`` override, +``-cw``/``--column-weight`` multiplies the amplitudes by a weight column before the +FFT, ``--units {sigma,electrons,raw}`` (``sigma``, the default, gives zero mean and +unit standard deviation; ``electrons`` gives e/A^3 as +``(1/V) sum_h F(h) exp(-2 pi i h.x)`` with the amplitudes divided by the +``-ck``/``--column-scale`` factor, ``KSCALE`` by default). ``-n`` is the deprecated +alias of ``--units sigma``/``raw``. :API: :mod:`torchref.cli.mtz2map` @@ -126,7 +135,11 @@ Computes real-space correlations and resolution-binned reciprocal-space CC. -dm dark.pdb -lm light.pdb **Key options:** ``--fraction``, ``--selection`` (Phenix-style atom -selection), ``--mask-radius``, ``--n-bins``. +selection), ``--mask-radius``, ``--n-bins``, ``--ded-weight`` (the headline weight +scheme; every scheme is also reported side by side, real-space in each mask and +reciprocal-space overall, as the ``by_weight`` block of the JSON and a table in the +summary, and a ``sigma_d`` fallback to inverse variance is recorded under +``weights``). :API: :mod:`torchref.cli.validate_ded` @@ -137,10 +150,14 @@ Compute difference and extrapolated map coefficients without refinement. Uses the same pipeline as ``torchref.difference-refine`` but the input models are kept as-is. -The default output is the weighted difference map ``DELFWT``/``PHDELWT`` -- -the inverse-variance-weighted amplitude difference on the **dark** model's -phases, which is the construction ``torchref.validate-ded`` correlates -against. It needs no light-state model, so ``-lm`` is optional: +The default output is the difference map: the amplitude difference ``DF``/``SIGDF`` +on the **dark** model's phases ``PHDELWT``, with one mean-one weight column per +registered scheme beside it -- ``W_IVW``, the inverse variance ``1/sigma^2`` (the +default), and ``W_SD``, the sigma_D Wiener weight ``S/(S + sigma^2)`` built from the +expected difference power -- and ``KSCALE``, the scaler's factor from model to observed +scale. This is the construction ``torchref.validate-ded`` correlates against. Build the +map with ``torchref.mtz2map -csf DF -cw W_IVW -cphi PHDELWT``, adding +``--units electrons`` for e/A^3. It needs no light-state model, so ``-lm`` is optional: .. code-block:: bash @@ -158,9 +175,15 @@ phase and the extrapolated map ``FWT``/``PHWT``: -dsf dark.mtz -lsf light.mtz \ --fraction 0.37 -o results.mtz -**Key options:** ``--all-columns`` writes every alternative map coefficient -and diagnostic -- the model-phased difference, the two other extrapolations -and the intensity block -- at the cost of two further scale fits. +**Key options:** ``--ded-weight {inverse_variance,sigma_d,none}`` selects the +scheme the model-phased and two-moment difference columns carry (default +``inverse_variance``; ``sigma_d`` needs calibrated sigmas, reports how many shells +it found without difference power, and falls back to inverse variance with a warning +when that is every shell); ``--sigma-d-gamma`` fixes the +dark-amplitude exponent of the sigma_D power law instead of fitting it; +``--all-columns`` writes every alternative map coefficient and diagnostic -- the +model-phased difference, the two other extrapolations and the intensity block -- at +the cost of two further scale fits. :API: :mod:`torchref.cli.difference_map` diff --git a/docs/user_guide/targets.rst b/docs/user_guide/targets.rst index d4874be7..44dbe587 100644 --- a/docs/user_guide/targets.rst +++ b/docs/user_guide/targets.rst @@ -80,6 +80,16 @@ batched over ``(n_datasets, n_hkl)`` on the collection's common HKL grid. - ``difference_i`` — the same on **intensities**. The entire class is one ``observable`` declaration: the difference-from-mean algebra does not care what the observable is. +- ``difference_sd`` — the ``difference`` Gaussian centred on + :math:`\alpha\,\Delta F_{calc}` with variance + :math:`\beta_{model} + \sigma_{\Delta}^2`, where :math:`\alpha` and the unexplained + difference power :math:`\beta_{model}` come from a per-shell moment fit of the + observed differences on the free set (``sigma_D``, + :mod:`torchref.refinement.model_error_estimation.sigma_d`); the expected power + carries an :math:`F_{dark}^{\gamma}` dependence with one fitted :math:`\gamma`. A + poor light model inflates the variance where it fails instead of pulling the + coordinates toward noise. Select it with + ``torchref.difference-refine --difference-target difference_sd``. - ``two_moment`` — merged **intensities** as :math:`|F(\bar\alpha)|^2 + \sigma_\alpha^2 |\Delta F|^2`, accounting for crystal-to-crystal spread in activation. diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index a2675973..d2555971 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -509,6 +509,7 @@ class TargetDeviceCase: # Device-bearing classes deliberately not in CASES. Every entry needs a reason; # "hard to build" is a reason, "didn't get to it" is not. UNCOVERED: Dict[str, str] = { + "CollectionDifferenceSigmaDTarget": "needs a dataset collection", # --- abstract / mixin bases: never instantiated directly ----------------- "Target": "abstract base; covered through its concrete subclasses", "ModelTarget": "abstract base; needs a loaded model", diff --git a/tests/integration/test_cli_ded_weights.py b/tests/integration/test_cli_ded_weights.py new file mode 100644 index 00000000..8679564b --- /dev/null +++ b/tests/integration/test_cli_ded_weights.py @@ -0,0 +1,299 @@ +"""The registered difference weights through the CLIs. + +Pinned: ``torchref.difference-map`` writes ``DF`` with one mean-one weight column per +scheme and ``KSCALE``; ``torchref.mtz2map`` builds the weighted map from those columns +and the electrons map is the volume-normalised synthesis on the absolute scale; +``torchref.validate-ded`` reports every scheme side by side and records a fallback; +``torchref.difference-refine`` runs with the ``difference_sd`` row and reports the fit. +""" + +import json +import os +import subprocess +import sys + +import numpy as np +import pytest + +pytestmark = [pytest.mark.integration, pytest.mark.slow] + +DIFF_COLUMNS = { + "Fo_dark": "SFAmplitude", + "SIGFo_dark": "Stddev", + "Fo_light": "SFAmplitude", + "SIGFo_light": "Stddev", + "DF": "SFAmplitude", + "SIGDF": "Stddev", + "PHDELWT": "Phase", + "W_IVW": "Weight", + "W_SD": "Weight", + "KSCALE": "MTZReal", + "Fc_dark": "SFAmplitude", + "FreeR_flag_dark": "MTZInt", + "FreeR_flag_light": "MTZInt", +} + + +@pytest.fixture(scope="module") +def pair(mtz_dir, pdb_dir, tmp_path_factory): + """A dark/light pair from 1DAW with a perturbed light state. + + The light amplitudes carry an added difference proportional to ``F`` with a + resolution-dependent power, so the sigma_D fit has signal to find; the dark set + keeps the deposited values. The light model is the dark one shifted by 0.2 A. + """ + import torch + + from torchref import ReflectionData + from torchref.config import get_int_dtype + + mtz = mtz_dir / "1DAW.mtz" + pdb = pdb_dir / "1DAW.pdb" + assert mtz.is_file() and pdb.is_file() + out = tmp_path_factory.mktemp("ded_weights_cli") + + data = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz)) + n = len(data) + idx = torch.arange(n, dtype=get_int_dtype()) + dark = data.__select__(idx < int(n * 0.97)) + dark.write_mtz(str(out / "dark.mtz")) + + sel = data.__select__(idx >= int(n * 0.03)) + g = torch.Generator().manual_seed(11) + f = sel.F + dss = 1.0 / sel.resolution**2 + change = 0.08 * f * torch.exp(-2.0 * dss) * torch.randn(len(sel), generator=g) + light = ReflectionData.from_tensors( + hkl=sel.hkl, + F=(f + change).clamp(min=0.0), + F_sigma=sel.F_sigma, + cell=sel.cell, + spacegroup=sel.spacegroup, + rfree_flags=sel.rfree_flags, + device="cpu", + verbose=0, + ) + light.write_mtz(str(out / "light.mtz")) + + lines = [] + for line in pdb.read_text().splitlines(): + if line.startswith(("ATOM", "HETATM")): + x = float(line[30:38]) + 0.2 + line = line[:30] + f"{x:8.3f}" + line[38:] + lines.append(line) + (out / "light.pdb").write_text("\n".join(lines) + "\n") + return {"dir": out, "pdb": pdb, "light_pdb": out / "light.pdb"} + + +def _run(project_root, module, *argv, timeout=1800): + script = project_root / "torchref" / "cli" / module + env = dict(os.environ) + env["PYTHONPATH"] = str(project_root) + os.pathsep + env.get("PYTHONPATH", "") + proc = subprocess.run( + [sys.executable, str(script), *map(str, argv)], + capture_output=True, + text=True, + timeout=timeout, + env=env, + ) + assert proc.returncode == 0, ( + f"{module} failed ({proc.returncode})\nstdout tail:\n{proc.stdout[-2000:]}" + f"\nstderr tail:\n{proc.stderr[-3000:]}" + ) + return proc + + +@pytest.fixture(scope="module") +def diff_mtz(project_root, pair): + out = pair["dir"] / "diff.mtz" + _run( + project_root, + "difference_map.py", + "-dm", + pair["pdb"], + "-dsf", + pair["dir"] / "dark.mtz", + "-lsf", + pair["dir"] / "light.mtz", + "--dmin", + "2.2", + "--device", + "cpu", + "--ded-weight", + "sigma_d", + "-v", + "1", + "-o", + out, + ) + return out + + +def _read(path): + import reciprocalspaceship as rs + + return rs.read_mtz(str(path)) + + +def test_difference_map_writes_df_weights_and_scale(diff_mtz): + df = _read(diff_mtz) + assert {c: str(df.dtypes[c]) for c in df.columns} == DIFF_COLUMNS + for col in ("W_IVW", "W_SD"): + w = df[col].to_numpy().astype(float) + assert np.isfinite(w).all() and (w >= 0).all() + assert abs(w.mean() - 1.0) < 1e-4 + assert (df["KSCALE"].to_numpy().astype(float) > 0).all() + # The sigma_D weights favour the strong reflections, inverse variance does not. + f = df["Fo_dark"].to_numpy().astype(float) + w_sd = df["W_SD"].to_numpy().astype(float) + strong = f > np.median(f) + assert w_sd[strong].mean() > w_sd[~strong].mean() + + +def test_mtz2map_builds_the_weighted_and_electron_maps(project_root, pair, diff_mtz): + import gemmi + + out = pair["dir"] + _run( + project_root, + "mtz2map.py", + "-sf", + diff_mtz, + "-csf", + "DF", + "-cw", + "W_SD", + "-cphi", + "PHDELWT", + "--device", + "cpu", + "-o", + out / "sd_sigma.ccp4", + ) + _run( + project_root, + "mtz2map.py", + "-sf", + diff_mtz, + "-csf", + "DF", + "-cw", + "W_SD", + "-cphi", + "PHDELWT", + "--units", + "raw", + "--device", + "cpu", + "-o", + out / "sd_raw.ccp4", + ) + _run( + project_root, + "mtz2map.py", + "-sf", + diff_mtz, + "-csf", + "DF", + "-cw", + "W_SD", + "-cphi", + "PHDELWT", + "--units", + "electrons", + "--device", + "cpu", + "-o", + out / "sd_e.ccp4", + ) + sigma = np.array(gemmi.read_ccp4_map(str(out / "sd_sigma.ccp4")).grid, copy=False) + raw = np.array(gemmi.read_ccp4_map(str(out / "sd_raw.ccp4")).grid, copy=False) + electrons = np.array(gemmi.read_ccp4_map(str(out / "sd_e.ccp4")).grid, copy=False) + assert abs(sigma.std() - 1.0) < 1e-3 and abs(sigma.mean()) < 1e-3 + # Same map up to normalisation: the correlation is one. + assert np.corrcoef(sigma.ravel(), raw.ravel())[0, 1] > 0.9999 + # Dividing by the per-reflection KSCALE reshapes the map slightly, so electrons is + # highly but not perfectly correlated with the sigma map. + assert 0.9 < np.corrcoef(sigma.ravel(), electrons.ravel())[0, 1] < 0.9999 + ratio = electrons.std() / raw.std() + assert np.isfinite(ratio) and ratio > 0 + + +def test_validate_ded_reports_every_scheme(project_root, pair): + out = pair["dir"] / "val" + proc = _run( + project_root, + "validate_ded.py", + "-dsf", + pair["dir"] / "dark.mtz", + "-lsf", + pair["dir"] / "light.mtz", + "-dm", + pair["pdb"], + "-lm", + pair["light_pdb"], + "--fraction", + "0.3", + "--dmin", + "2.2", + "--device", + "cpu", + "--ded-weight", + "sigma_d", + "-v", + "1", + "-o", + out, + ) + results = json.loads((out / "validate_ded_results.json").read_text()) + assert results["weights"]["requested"] == "sigma_d" + assert results["weights"]["applied"] in ("sigma_d", "inverse_variance") + assert set(results["by_weight"]) == {"none", "inverse_variance", "sigma_d"} + for entry in results["by_weight"].values(): + assert np.isfinite(entry["reciprocal_cc_overall"]) + assert "full_cell" in entry["realspace_correlation"] + headline = results["by_weight"][results["weights"]["applied"]] + assert results["reciprocal_cc_overall"] == pytest.approx( + headline["reciprocal_cc_overall"], abs=1e-3 + ) + assert "weights " in proc.stdout and "sigma_d" in proc.stdout + + +def test_difference_refine_runs_the_sigma_d_row(project_root, pair): + out = pair["dir"] / "refine" + _run( + project_root, + "collection_difference_refine.py", + "-dm", + pair["pdb"], + "-lm", + pair["light_pdb"], + "-dsf", + pair["dir"] / "dark.mtz", + "-lsf", + pair["dir"] / "light.mtz", + "--fraction", + "0.25", + "--difference-target", + "difference_sd", + "--n-cycles", + "1", + "--n-steps", + "1", + "--max-iter", + "3", + "--dmin", + "2.2", + "--device", + "cpu", + "--verbose", + "0", + "-o", + out, + ) + summaries = list(out.glob("*_summary.json")) + assert len(summaries) == 1 + results = json.loads(summaries[0].read_text())["results"] + assert results["ded_weights"]["scheme"] == "inverse_variance" + assert results["ded_weights"]["applied"] == "inverse_variance" + assert "gamma" in results["ded_weights"]["sigma_d"] diff --git a/tests/integration/test_cli_two_moment_mtz.py b/tests/integration/test_cli_two_moment_mtz.py index 9f42aee7..95a46632 100644 --- a/tests/integration/test_cli_two_moment_mtz.py +++ b/tests/integration/test_cli_two_moment_mtz.py @@ -19,9 +19,12 @@ "SIGFo_light": "Stddev", "DF": "SFAmplitude", "SIGDF": "Stddev", - # The difference map. CCP4/Coot open these by name. - "DELFWT": "SFAmplitude", + # The difference map: DF on the dark phases, one weight column per registered + # scheme, and the observed-to-model scale. "PHDELWT": "Phase", + "W_IVW": "Weight", + "W_SD": "Weight", + "KSCALE": "MTZReal", "Fc_dark": "SFAmplitude", # The mixed model, and the extrapolated map to refine against. "FC": "SFAmplitude", @@ -216,14 +219,14 @@ def test_column_layout_and_types(request, fixture, extra): class TestDefaultLayout: - def test_the_difference_map_is_the_weighted_difference_on_dark_phases( + def test_the_difference_columns_carry_df_and_the_registered_weights( self, baseline_mtz ): - """``DELFWT`` must be ``(Fo_light - Fo_dark) * w`` with ``w`` the mean- - normalised inverse variance -- the construction ``torchref.validate-ded`` - correlates against. If these two ever diverge, the map in the file stops being - the map the validation reports on, which is how the output drifted from the - science before.""" + """``DF`` must be ``Fo_light - Fo_dark``, ``W_IVW`` the mean-normalised inverse + variance and ``W_SD`` a mean-one weight -- the constructions + ``torchref.validate-ded`` correlates against. If these ever diverge, the map + built from the file stops being the map the validation reports on, which is how + the output drifted from the science before.""" import numpy as np df = _read(baseline_mtz[0]) @@ -234,13 +237,17 @@ def test_the_difference_map_is_the_weighted_difference_on_dark_phases( df["SIGFo_dark"].to_numpy().astype(float) ** 2 + df["SIGFo_light"].to_numpy().astype(float) ** 2 ) - w = 1 / sig**2 + w = 1 / np.maximum(sig, 0.1 * np.median(sig)) ** 2 w = w / w.mean() - expected = dfo * w - got = df["DELFWT"].to_numpy().astype(float) - scale = max(float(np.abs(expected).max()), 1e-30) - assert np.abs(got - expected).max() / scale < 1e-5 + got_df = df["DF"].to_numpy().astype(float) + scale = max(float(np.abs(dfo).max()), 1e-30) + assert np.abs(got_df - dfo).max() / scale < 1e-5 + got_w = df["W_IVW"].to_numpy().astype(float) + assert np.abs(got_w - w).max() / max(float(np.abs(w).max()), 1e-30) < 1e-4 + w_sd = df["W_SD"].to_numpy().astype(float) + assert np.isfinite(w_sd).all() and abs(w_sd.mean() - 1.0) < 1e-4 + assert (df["KSCALE"].to_numpy().astype(float) > 0).all() # And the phase is the dark model's, not the mixed model's. assert not np.allclose( @@ -255,16 +262,12 @@ def test_the_corrected_difference_map_pairs_with_the_same_phases( self, two_moment_mtz ): """``DELFWT_corr`` is the corrected difference on the *same* dark phases, so it - is opened against ``PHDELWT`` and must be built the same way as ``DELFWT``.""" + is opened against ``PHDELWT`` and carries the selected weight scheme, the + inverse-variance weights ``W_IVW`` by default.""" import numpy as np df = _read(two_moment_mtz[0]) - sig = np.sqrt( - df["SIGFo_dark"].to_numpy().astype(float) ** 2 - + df["SIGFo_light"].to_numpy().astype(float) ** 2 - ) - w = 1 / sig**2 - w = w / w.mean() + w = df["W_IVW"].to_numpy().astype(float) expected = df["DF_corr"].to_numpy().astype(float) * w got = df["DELFWT_corr"].to_numpy().astype(float) diff --git a/tests/unit/maps/test_ded_weights.py b/tests/unit/maps/test_ded_weights.py new file mode 100644 index 00000000..60abe281 --- /dev/null +++ b/tests/unit/maps/test_ded_weights.py @@ -0,0 +1,128 @@ +"""The registered difference-coefficient weight schemes. + +Pinned: the three schemes exist with their MTZ column names; ``none`` is flat; +``inverse_variance`` has mean one and floors a zero sigma; ``sigma_d`` gives strong +reflections more weight than weak ones within a shell where inverse variance cannot; +and an all-noise input falls back to inverse variance with a warning that names why. +""" + +import pytest +import torch + +from tests.unit.refinement.test_sigma_d import synth_diff +from torchref.maps.ded_weights import ( + DEFAULT_SCHEME, + SCHEMES, + WEIGHT_COLUMNS, + DedWeightFallbackWarning, + all_ded_weights, + compute_ded_weights, + normalise_mean_one, +) +from torchref.symmetry import SpaceGroup + + +def _inputs(n=20000, sig_frac=1.0, device="cpu"): + d = synth_diff(n=n, sig_frac=sig_frac, device=device) + g = torch.Generator().manual_seed(5) + hkl = torch.randint(-20, 21, (n, 3), generator=g).to(device) + cell = torch.tensor([40.0, 50.0, 60.0, 90.0, 90.0, 90.0], device=device) + return d, hkl, cell, SpaceGroup("P 1", device=device) + + +@pytest.mark.unit +def test_registry_is_consistent(): + assert DEFAULT_SCHEME in SCHEMES + assert set(WEIGHT_COLUMNS) == set(SCHEMES) - {"none"} + with pytest.raises(ValueError): + compute_ded_weights( + "bogus", + delta_obs=torch.zeros(3), + sigma_diff=torch.ones(3), + hkl=torch.zeros(3, 3), + cell=torch.ones(6), + spacegroup=None, + ) + + +@pytest.mark.unit +def test_normalise_mean_one_handles_nonfinite_and_zero(): + # Non-finite entries drop to zero and count in the mean, so the column mean is one + # however many reflections carry weight. + w = normalise_mean_one(torch.tensor([1.0, 3.0, float("nan"), float("inf")])) + assert torch.allclose(w, torch.tensor([1.0, 3.0, 0.0, 0.0])) + assert w.mean() == pytest.approx(1.0) + half = normalise_mean_one(torch.tensor([0.0, 0.0, 2.0, 6.0])) + assert torch.allclose(half, torch.tensor([0.0, 0.0, 1.0, 3.0])) + z = normalise_mean_one(torch.zeros(4)) + assert torch.equal(z, torch.zeros(4)) + + +@pytest.mark.unit +def test_none_and_inverse_variance(any_device): + d, hkl, cell, sg = _inputs(n=2000, device=any_device) + sig = d["sigma_diff"].clone() + sig[0] = 0.0 + kw = { + "delta_obs": d["delta_obs"], + "sigma_diff": sig, + "hkl": hkl, + "cell": cell, + "spacegroup": sg, + } + flat = compute_ded_weights("none", **kw) + assert torch.equal(flat.weights, torch.ones_like(sig)) + ivw = compute_ded_weights("inverse_variance", **kw) + assert ivw.applied == "inverse_variance" + assert abs(float(ivw.weights.mean()) - 1.0) < 1e-5 + # The zero sigma is floored, so it carries the largest finite weight. + assert torch.isfinite(ivw.weights).all() + assert ( + float(ivw.weights[0]) == float(ivw.weights.max()) > float(ivw.weights[1:].max()) + ) + assert ivw.weights.device == d["delta_obs"].device + + +@pytest.mark.unit +def test_sigma_d_favours_strong_reflections_where_inverse_variance_cannot(any_device): + d, hkl, cell, sg = _inputs(device=any_device) + kw = { + "delta_obs": d["delta_obs"], + "sigma_diff": d["sigma_diff"], + "hkl": hkl, + "cell": cell, + "spacegroup": sg, + "f_dark": d["f_dark"], + } + every = all_ded_weights(**kw) + assert set(every) == set(SCHEMES) + sd = every["sigma_d"] + assert sd.applied == "sigma_d" and abs(float(sd.weights.mean()) - 1.0) < 1e-4 + assert 0.8 < sd.diagnostics["gamma"] < 1.2 + assert sd.diagnostics["n_shell"] > 10 and "shells" in sd.diagnostics + # Within the highest-resolution tenth, the strongest reflections carry more weight. + order = torch.argsort(d["d_star_sq"])[-2000:] + f, w = d["f_dark"][order], sd.weights[order] + strong, weak = f > f.median(), f <= f.median() + assert float(w[strong].mean()) > 1.5 * float(w[weak].mean()) + ivw = every["inverse_variance"].weights[order] + assert abs(float(ivw[strong].mean()) - float(ivw[weak].mean())) < 1e-4 + + +@pytest.mark.unit +def test_all_noise_falls_back_to_inverse_variance_with_a_warning(): + d, hkl, cell, sg = _inputs(n=5000, sig_frac=50.0) + kw = { + "delta_obs": d["delta_obs"], + "sigma_diff": d["sigma_diff"] * 1.2, + "hkl": hkl, + "cell": cell, + "spacegroup": sg, + "f_dark": d["f_dark"], + } + with pytest.warns(DedWeightFallbackWarning, match="inverse-variance"): + sd = compute_ded_weights("sigma_d", **kw) + assert sd.scheme == "sigma_d" and sd.applied == "inverse_variance" + assert "fallback_reason" in sd.diagnostics + ivw = compute_ded_weights("inverse_variance", **kw) + assert torch.allclose(sd.weights, ivw.weights) diff --git a/tests/unit/maps/test_map_units.py b/tests/unit/maps/test_map_units.py new file mode 100644 index 00000000..414e77c8 --- /dev/null +++ b/tests/unit/maps/test_map_units.py @@ -0,0 +1,51 @@ +"""Map units: the electrons-per-cubic-Angstrom synthesis. + +Pinned: ``units="electrons"`` is the ``1/N``-normalised FFT rescaled by ``N / V``, i.e. +``(1/V) sum_h F(h) exp(-2 pi i h.x)``; the default is unchanged; an unknown unit is +rejected; a ``DifferenceMap`` accepts a per-reflection scale and the same units. +""" + +import pytest +import torch + +from torchref.io import ReflectionData +from torchref.maps import DifferenceMap, Map +from torchref.model.model_ft import ModelFT + + +@pytest.fixture(scope="module") +def model_ft_and_data(sample_structure_pair): + model = ModelFT() + model.load_cif(str(sample_structure_pair["model"])) + data = ReflectionData() + data.load_mtz(str(sample_structure_pair["reflections"])) + return model, data + + +@pytest.mark.unit +def test_electrons_is_the_volume_normalised_synthesis(model_ft_and_data): + model, data = model_ft_and_data + normalized = Map(data, model, map_type="Fcalc").calculate() + electrons = Map(data, model, map_type="Fcalc", units="electrons").calculate() + volume = data.cell.volume.to(normalized.dtype) + assert torch.allclose( + electrons, normalized * (normalized.numel() / volume), rtol=1e-5, atol=1e-6 + ) + + +@pytest.mark.unit +def test_unknown_units_are_rejected(model_ft_and_data): + model, data = model_ft_and_data + with pytest.raises(ValueError, match="units must be one of"): + Map(data, model, units="e/A3") + + +@pytest.mark.unit +def test_difference_map_scale_and_units(model_ft_and_data): + model, data = model_ft_and_data + plain = DifferenceMap(data, data, model).calculate() + scale = torch.full((len(data),), 2.0, dtype=plain.dtype, device=plain.device) + scaled = DifferenceMap(data, data, model, scale=scale, units="electrons") + out = scaled.calculate() + assert out.shape == plain.shape and torch.isfinite(out).all() + assert scaled.units == "electrons" and scaled.scale is scale diff --git a/tests/unit/refinement/test_collection_sigma_d_target.py b/tests/unit/refinement/test_collection_sigma_d_target.py new file mode 100644 index 00000000..dddb57d2 --- /dev/null +++ b/tests/unit/refinement/test_collection_sigma_d_target.py @@ -0,0 +1,126 @@ +"""The ``difference_sd`` collection row on a real dark/light pair. + +Pinned on 1DAW with a 0.2 A shifted light model: the loss is finite, gradients reach +the light model through ``dF_calc`` only, the sigma_D estimate is owned by the target, +fitted on free reflections of the timepoint row, cached across forwards and cleared +by ``maintenance()``, and the fit summary reaches ``stats()``. +""" + +import pytest +import torch + +from torchref.refinement.model_error_estimation import sigma_d as sigma_d_module +from torchref.refinement.model_error_estimation.sigma_d import SigmaDEstimator +from torchref.refinement.targets.collection import ( + CollectionDifferenceSigmaDTarget, + CollectionSigmaDLossInputs, +) + +pytestmark = pytest.mark.integration + + +@pytest.fixture +def target(loaded_reflection_data, sample_structure_pair): + """A dark/light collection whose light amplitudes carry a resolution-dependent + difference proportional to ``F``, so the sigma_D coupling is not zero, and a light + model shifted by 0.2 A.""" + from torchref import ReflectionData + from torchref.cli._common import load_model + from torchref.io import DatasetCollection + from torchref.model import ModelCollection + from torchref.scaling import CollectionScaler + + data = loaded_reflection_data + g = torch.Generator().manual_seed(11) + f = data.F + dss = 1.0 / data.resolution**2 + change = ( + 0.08 * f * torch.exp(-2.0 * dss) * torch.randn(len(data), generator=g).to(f) + ) + light = ReflectionData.from_tensors( + hkl=data.hkl, + F=(f + change).clamp(min=0.0), + F_sigma=data.F_sigma, + cell=data.cell, + spacegroup=data.spacegroup, + rfree_flags=data.rfree_flags, + device=str(data.device), + verbose=0, + ) + models = [ + load_model( + str(sample_structure_pair["model"]), + max_res=2.05, + device=data.device, + verbose=0, + ) + for _ in range(2) + ] + with torch.no_grad(): + models[1].xyz.refinable_params += 0.2 + dc = DatasetCollection(device=data.device, verbose=0) + dc.add_dataset("dark", data, set_as_reference=True).add_dataset("light", light) + mc = ModelCollection(models, dark_key="dark", verbose=0) + mc.add_dark().add_timepoint("light", [0.78, 0.22]) + scaler = CollectionScaler(dc, mc, verbose=0).initialize() + return dc, mc, CollectionDifferenceSigmaDTarget(dc, mc, scaler=scaler) + + +def test_forward_is_finite_and_owns_its_estimator(target): + _dc, _mc, t = target + assert isinstance(t._sigma_d, SigmaDEstimator) + loss = t.forward() + assert torch.isfinite(loss) + ctx = t._loss_inputs() + assert isinstance(ctx, CollectionSigmaDLossInputs) + assert ctx.alpha.shape == ctx.beta_model.shape == (ctx.obs.shape[1],) + assert not ctx.alpha.requires_grad and not ctx.beta_model.requires_grad + assert (ctx.beta_model >= 0).all() + shells = t._sigma_d.shells + assert shells.has_model and not shells.all_zero and not shells.degenerate + + +def test_gradient_reaches_the_light_model(target): + _dc, mc, t = target + light = mc.base_models[1] + light.zero_grad(set_to_none=True) + t.forward().backward() + grads = [p.grad for p in light.parameters() if p.grad is not None] + assert grads and any(torch.isfinite(g).all() and g.abs().sum() > 0 for g in grads) + + +def test_estimate_is_cached_until_maintenance(target): + _dc, _mc, t = target + t.forward() + assert t._sigma_d._cache is not None + first = t._sigma_d._cache + t.forward() + assert t._sigma_d._cache is first + t.maintenance() + assert t._sigma_d._cache is None + + +def test_fit_uses_free_reflections_of_the_timepoint_row(target, monkeypatch): + dc, _mc, t = target + seen = {} + real = sigma_d_module.estimate_sigma_d + + def spy(delta_obs, sigma_diff, epsilon, d_star_sq, f_dark, fit_mask, **kw): + seen["fit_mask"] = fit_mask.clone() + seen["n"] = delta_obs.numel() + return real(delta_obs, sigma_diff, epsilon, d_star_sq, f_dark, fit_mask, **kw) + + monkeypatch.setattr(sigma_d_module, "estimate_sigma_d", spy) + t.forward() + n_hkl = dc.hkl.shape[0] + assert seen["n"] == n_hkl + free = dc["light"].free.mask.to(seen["fit_mask"].device) + assert bool((seen["fit_mask"] & ~free).sum() == 0) + assert int(seen["fit_mask"].sum()) > 0 + + +def test_stats_carry_the_fit_summary(target): + _dc, _mc, t = target + t.forward() + stats = t.stats() + assert "sigma_d_gamma" in stats and "sigma_d_tau" in stats diff --git a/tests/unit/refinement/test_collection_taxonomy.py b/tests/unit/refinement/test_collection_taxonomy.py index 9465712e..c0956b82 100644 --- a/tests/unit/refinement/test_collection_taxonomy.py +++ b/tests/unit/refinement/test_collection_taxonomy.py @@ -14,6 +14,7 @@ def test_each_row_has_its_own_class(): ) from torchref.refinement.targets.collection.xray import ( CollectionDifferenceIntensityTarget, + CollectionDifferenceSigmaDTarget, CollectionDifferenceTarget, CollectionMLTarget, ) @@ -21,6 +22,7 @@ def test_each_row_has_its_own_class(): expected = { "difference": CollectionDifferenceTarget, "difference_i": CollectionDifferenceIntensityTarget, + "difference_sd": CollectionDifferenceSigmaDTarget, "two_moment": CollectionTwoMomentIntensityTarget, "ml": CollectionMLTarget, } @@ -56,7 +58,7 @@ def test_the_observable_is_declared_not_passed(): ) assert set(by_obs["intensity"]) == {"difference_i", "two_moment"} - assert set(by_obs["amplitude"]) == {"difference", "ml"} + assert set(by_obs["amplitude"]) == {"difference", "difference_sd", "ml"} with pytest.raises(ValueError, match="observable"): CollectionXrayTargetSpec( diff --git a/tests/unit/refinement/test_shells.py b/tests/unit/refinement/test_shells.py new file mode 100644 index 00000000..c29a83d5 --- /dev/null +++ b/tests/unit/refinement/test_shells.py @@ -0,0 +1,127 @@ +"""Shell helpers shared by sigma_A and sigma_D. + +Pinned: the helpers ``sigma_a`` re-imports are the same objects ``_shells`` defines, so a +fit through either module reduces identically; ``equal_count_shells`` reproduces the +shell construction written out in ``estimate_beta``; the line shrinkage takes the line +outright when the scatter is below the noise, passes through with too few shells, +honours a slope clamp, and treats a shell without a finite variance as undetermined. +""" + +import math + +import pytest +import torch + +from torchref.refinement.model_error_estimation import _shells, sigma_a +from torchref.refinement.model_error_estimation._shells import ( + dl_shrink_to_line, + equal_count_shells, + interp_in_dss, + segsum, +) + + +@pytest.mark.unit +def test_sigma_a_uses_the_shared_helpers(): + assert sigma_a._segsum is _shells.segsum + assert sigma_a._interp_in_dss is _shells.interp_in_dss + assert sigma_a._segment_layout is _shells.segment_layout + + +@pytest.mark.unit +@pytest.mark.parametrize("n,per_bin", [(2000, 140), (37, 140), (9, 140), (5000, 500)]) +def test_equal_count_shells_matches_estimate_beta_construction(n, per_bin, any_device): + """The construction ``estimate_beta`` writes out inline, reproduced independently.""" + g = torch.Generator().manual_seed(3) + dss = (torch.rand(n, generator=g) * 0.3 + 0.02).to(any_device) + min_bins, min_per_bin = 5, 40 + order, seg, seg_lengths, n_bins = equal_count_shells( + dss, per_bin=per_bin, min_bins=min_bins, min_per_bin=min_per_bin + ) + # Oracle: the four lines of estimate_beta. + ref_order = torch.argsort(dss, stable=True) + n_by_count = max(1, n // per_bin) + n_cap = max(1, n // min_per_bin) + ref_bins = max(n_by_count, min(min_bins, n_cap)) + ref_seg = (torch.arange(n, device=dss.device) * ref_bins) // n + assert n_bins == ref_bins + assert torch.equal(order, ref_order) + assert torch.equal(seg, ref_seg) + assert torch.equal(seg_lengths, torch.bincount(ref_seg, minlength=ref_bins)) + assert int(seg_lengths.sum()) == n + assert int(seg_lengths.max() - seg_lengths.min()) <= 1 + + +@pytest.mark.unit +def test_segsum_and_interp_round_trip(any_device): + lengths = torch.tensor([3, 2, 4], device=any_device) + x = torch.arange(9, dtype=torch.float32, device=any_device) + assert torch.equal( + segsum(x, lengths), torch.tensor([3.0, 7.0, 26.0], device=any_device) + ) + bin_dss = torch.tensor([0.1, 0.2, 0.3], device=any_device) + vals = torch.tensor([1.0, 3.0, 5.0], device=any_device) + grid = torch.tensor([0.0, 0.15, 0.25, 0.5], device=any_device) + out = interp_in_dss(grid, bin_dss, vals) + assert torch.allclose(out, torch.tensor([1.0, 2.0, 4.0, 5.0], device=any_device)) + + +@pytest.mark.unit +def test_shrink_passes_through_with_fewer_than_four_shells(): + y = torch.tensor([1.0, 2.0, 3.0]) + var = torch.ones(3) + x = torch.tensor([0.1, 0.2, 0.3]) + out, w, tau_sq, a, b = dl_shrink_to_line(y, var, x) + assert torch.equal(out, y) + assert torch.equal(w, torch.zeros(3)) + assert float(tau_sq) == 0.0 and math.isnan(a) and math.isnan(b) + + +@pytest.mark.unit +def test_shrink_takes_the_line_when_scatter_is_below_noise(): + x = torch.linspace(0.05, 0.35, 8, dtype=torch.float64) + line = 2.0 - 3.0 * x + g = torch.Generator().manual_seed(1) + var = torch.full((8,), 0.04, dtype=torch.float64) + y = line + 0.01 * torch.randn(8, generator=g, dtype=torch.float64) + out, w, tau_sq, a, b = dl_shrink_to_line(y, var, x) + assert float(tau_sq) == 0.0 + assert torch.allclose(w, torch.ones(8, dtype=torch.float64)) + assert torch.allclose(out, a + b * x) + assert abs(a - 2.0) < 0.05 and abs(b + 3.0) < 0.3 + + +@pytest.mark.unit +def test_shrink_keeps_a_real_departure(): + x = torch.linspace(0.05, 0.35, 10, dtype=torch.float64) + y = 1.0 - 2.0 * x + y[4] += 3.0 # one shell far off the line, far beyond its own variance + var = torch.full((10,), 1e-4, dtype=torch.float64) + out, w, tau_sq, _, _ = dl_shrink_to_line(y, var, x) + assert float(tau_sq) > 0.0 + assert float(w[4]) < 0.01 + assert abs(float(out[4] - y[4])) < 0.05 + + +@pytest.mark.unit +def test_shrink_slope_clamp_is_honoured(): + x = torch.linspace(0.05, 0.35, 8, dtype=torch.float64) + y = 1.0 + 4.0 * x + var = torch.full((8,), 0.01, dtype=torch.float64) + _, _, _, _, b = dl_shrink_to_line(y, var, x, slope_max=0.0) + assert b == 0.0 + _, _, _, _, b2 = dl_shrink_to_line(-y, var, x, slope_min=0.0) + assert b2 == 0.0 + + +@pytest.mark.unit +def test_shrink_replaces_undetermined_shells_by_the_line(): + x = torch.linspace(0.05, 0.35, 8, dtype=torch.float64) + y = 2.0 - 3.0 * x + var = torch.full((8,), 0.01, dtype=torch.float64) + y[2] = float("nan") + var[5] = float("inf") + out, w, _, a, b = dl_shrink_to_line(y, var, x) + assert float(w[2]) == 1.0 and float(w[5]) == 1.0 + assert torch.isfinite(out).all() + assert torch.allclose(out[[2, 5]], a + b * x[[2, 5]]) diff --git a/tests/unit/refinement/test_sigma_d.py b/tests/unit/refinement/test_sigma_d.py new file mode 100644 index 00000000..39ec8f72 --- /dev/null +++ b/tests/unit/refinement/test_sigma_d.py @@ -0,0 +1,358 @@ +"""Properties of the sigma_D difference-power estimator. + +Pinned on seeded synthetic differences with a KNOWN power law: the per-shell power is +recovered from a single dataset through ``mean(dF**2) - mean(sigma**2)``, the dark- +amplitude exponent is recovered and can be fixed, the moment identity with a difference +model holds exactly, clamps are counted and an all-noise input is flagged rather than +weighted, degenerate inputs stay finite, per-reflection weights lie in ``[0, 1)`` with the +shell mean of the power preserved, the fit is deterministic and device-independent, and +the cached estimator resets on demand. +""" + +import pytest +import torch + +from torchref.refinement.model_error_estimation._shells import interp_in_dss +from torchref.refinement.model_error_estimation.sigma_d import ( + GAMMA_DEFAULT, + SigmaDConfig, + SigmaDEstimator, + estimate_sigma_d, + sigma_d_per_reflection, +) + +#: Tolerance on the recovered shell power relative to the truth. A shell of 140 +#: reflections estimates ``B`` with a relative sd of ``sqrt(2/140) = 12%``; the line +#: shrinkage pools shells, and the decile means below average ~14 shells, so 10% is +#: ~3 sd of what remains. +POWER_RTOL = 0.10 +#: Tolerance on the fitted exponent. Its standard error at 30 000 reflections is ~0.04; +#: 0.15 is well above that and well below the difference between the pure shell model +#: (0) and the default (1). +GAMMA_ATOL = 0.15 + + +def synth_diff( + n=30000, + gamma=1.0, + sig_frac=1.0, + seed=7, + dtype=torch.float32, + device="cpu", + with_model=False, + alpha_true=0.8, +): + """Signed differences with power ``Sigma_N(d*^2) * (F / )**gamma``. + + ``F`` is Wilson-like (the modulus of a complex normal) so the amplitude classes are + populated realistically; ``Sigma_N`` falls with resolution. The measurement sigma is + ``sig_frac`` times the rms true difference, constant across reflections so the + inverse-variance and sigma_D weights differ only through ``S``. + """ + g = torch.Generator().manual_seed(seed) + dss = torch.linspace(0.02, 0.35, n, dtype=torch.float64) + f = ( + torch.randn(n, generator=g, dtype=torch.float64) ** 2 + + torch.randn(n, generator=g, dtype=torch.float64) ** 2 + ).sqrt() * 10.0 + sigma_n = 0.05 * torch.exp(-3.0 * dss) + s_true = sigma_n * (f / f.mean()) ** gamma + d_true = torch.randn(n, generator=g, dtype=torch.float64) * s_true.sqrt() + sig = torch.full( + (n,), float(sig_frac) * float(s_true.mean().sqrt()), dtype=torch.float64 + ) + d_obs = d_true + torch.randn(n, generator=g, dtype=torch.float64) * sig + out = { + "delta_obs": d_obs, + "sigma_diff": sig, + "d_star_sq": dss, + "f_dark": f, + "s_true": s_true, + "fit_mask": torch.ones(n, dtype=torch.bool), + } + if with_model: + beta_true = 0.3 * s_true + out["delta_calc"] = ( + d_true + torch.randn(n, generator=g, dtype=torch.float64) * beta_true.sqrt() + ) / alpha_true + return { + k: ( + v.to(device=device, dtype=dtype) + if v.dtype.is_floating_point + else v.to(device) + ) + for k, v in out.items() + } + + +def _decile_means(values, dss, n_dec=10): + order = torch.argsort(dss) + chunks = torch.chunk(values[order], n_dec) + return torch.stack([c.mean() for c in chunks]) + + +@pytest.mark.unit +def test_recovers_shell_power_from_one_dataset(any_device): + d = synth_diff(device=any_device) + sh = estimate_sigma_d( + d["delta_obs"], + d["sigma_diff"], + None, + d["d_star_sq"], + d["f_dark"], + d["fit_mask"], + ) + est = sigma_d_per_reflection(sh, d["d_star_sq"], None, d["f_dark"], d["sigma_diff"]) + assert not sh.degenerate and not sh.all_zero + assert (sh.Sigma_N > 0).all() + got = _decile_means(est.S, d["d_star_sq"]) + want = _decile_means(d["s_true"], d["d_star_sq"]) + assert torch.allclose(got, want, rtol=POWER_RTOL) + + +@pytest.mark.unit +@pytest.mark.parametrize("gamma", [1.0, 0.5]) +def test_recovers_the_amplitude_exponent(gamma): + d = synth_diff(gamma=gamma, dtype=torch.float64) + sh = estimate_sigma_d( + d["delta_obs"], + d["sigma_diff"], + None, + d["d_star_sq"], + d["f_dark"], + d["fit_mask"], + ) + assert sh.gamma_fitted and sh.diagnostics["gamma_reason"] == "fitted" + assert abs(sh.gamma - gamma) < GAMMA_ATOL + assert sh.diagnostics["gamma_se"] < GAMMA_ATOL + + +@pytest.mark.unit +def test_fixed_exponent_is_honoured(): + d = synth_diff(n=5000) + sh = estimate_sigma_d( + d["delta_obs"], + d["sigma_diff"], + None, + d["d_star_sq"], + d["f_dark"], + d["fit_mask"], + gamma=0.7, + ) + assert sh.gamma == 0.7 and not sh.gamma_fitted + assert sh.diagnostics["gamma_reason"] == "fixed" + with pytest.raises(ValueError): + SigmaDConfig(gamma=3.0) + + +@pytest.mark.unit +def test_without_dark_amplitude_the_power_is_flat_within_a_shell(): + d = synth_diff(n=5000) + sh = estimate_sigma_d( + d["delta_obs"], d["sigma_diff"], None, d["d_star_sq"], None, d["fit_mask"] + ) + assert sh.gamma == GAMMA_DEFAULT and not sh.gamma_fitted + assert sh.diagnostics["gamma_reason"] == "no_f_dark" + est = sigma_d_per_reflection(sh, d["d_star_sq"], None, None, d["sigma_diff"]) + # Reflections at the same resolution share the power regardless of amplitude. + order = torch.argsort(d["d_star_sq"]) + close = est.S[order][:200] + assert float(close.max() / close.min()) < 1.05 + + +@pytest.mark.unit +@pytest.mark.parametrize("dtype,rtol", [(torch.float32, 1e-4), (torch.float64, 1e-10)]) +def test_moment_identity_with_a_difference_model(dtype, rtol): + d = synth_diff(dtype=dtype, with_model=True) + sh = estimate_sigma_d( + d["delta_obs"], + d["sigma_diff"], + None, + d["d_star_sq"], + d["f_dark"], + d["fit_mask"], + delta_calc=d["delta_calc"], + shrink=False, + ) + assert sh.has_model + assert sh.diagnostics["n_s2_clamped"] == 0 + # Sampling noise can push alpha**2 Sigma_P above Sigma_N in a few shells; the clamp + # there is counted, and the identity is exact everywhere it did not fire. + unclamped = sh.alpha**2 * sh.Sigma_P <= sh.Sigma_N + assert ( + int(unclamped.sum()) + == sh.diagnostics["n_shell"] - sh.diagnostics["n_beta_clamped"] + ) + assert float(unclamped.float().mean()) > 0.8 + lhs = sh.alpha**2 * sh.Sigma_P + sh.beta_model + sh.S2 + assert torch.allclose(lhs[unclamped], sh.B[unclamped], rtol=rtol) + # alpha is the Gaussian coupling S / (S + beta_true) / alpha_true-scaled slope; it must + # be positive and below one for this generator. + assert (sh.alpha > 0).all() and (sh.alpha < 1).all() + + +@pytest.mark.unit +def test_clamps_are_counted_and_all_noise_is_flagged(): + d = synth_diff(n=5000, sig_frac=5.0) + sh = estimate_sigma_d( + d["delta_obs"], + d["sigma_diff"], + None, + d["d_star_sq"], + d["f_dark"], + d["fit_mask"], + ) + assert sh.diagnostics["n_s2_clamped"] > 0 + noise = synth_diff(n=5000, sig_frac=50.0) + sh2 = estimate_sigma_d( + noise["delta_obs"], + noise["sigma_diff"] * 1.2, + None, + noise["d_star_sq"], + noise["f_dark"], + noise["fit_mask"], + shrink=False, + ) + assert sh2.all_zero + est = sigma_d_per_reflection( + sh2, noise["d_star_sq"], None, noise["f_dark"], noise["sigma_diff"] + ) + assert torch.equal(est.w, torch.zeros_like(est.w)) + + +@pytest.mark.unit +def test_pure_noise_with_calibrated_sigma_gets_no_power(): + """Nothing tells the estimator whether a difference exists: on pure noise with + calibrated sigmas the shrinkage must not manufacture power from the positive half + of the noise in ``B - S2``, and the weights collapse onto inverse variance.""" + g = torch.Generator().manual_seed(3) + n = 30000 + dss = torch.linspace(0.02, 0.35, n, dtype=torch.float64) + f = torch.rand(n, generator=g, dtype=torch.float64) * 20.0 + 1.0 + sig = 0.2 + 0.8 * dss + d_obs = torch.randn(n, generator=g, dtype=torch.float64) * sig + mask = torch.ones(n, dtype=torch.bool) + sh = estimate_sigma_d(d_obs, sig, None, dss, f, mask) + assert not sh.degenerate + # Shell power is below a few per cent of the noise power in every shell. + assert (sh.Sigma_N <= 0.05 * sh.S2).all() + est = sigma_d_per_reflection(sh, dss, None, f, sig) + ivw = 1.0 / sig**2 + ivw = ivw / ivw.mean() + w_sd = est.w / est.w.mean().clamp(min=1e-30) + if not sh.all_zero: + assert torch.corrcoef(torch.stack([w_sd, ivw]))[0, 1] > 0.97 + + +@pytest.mark.unit +def test_degenerate_input_stays_finite(): + d = synth_diff(n=100) + mask = torch.zeros(100, dtype=torch.bool) + mask[0] = True + sh = estimate_sigma_d( + d["delta_obs"], d["sigma_diff"], None, d["d_star_sq"], d["f_dark"], mask + ) + assert sh.degenerate and not sh.all_zero + est = sigma_d_per_reflection(sh, d["d_star_sq"], None, d["f_dark"], d["sigma_diff"]) + assert torch.isfinite(est.S).all() and torch.isfinite(est.w).all() + assert est.S.shape == (100,) + + +@pytest.mark.unit +def test_per_reflection_weights_and_shell_mean(any_device): + d = synth_diff(device=any_device) + eps = torch.where( + torch.arange(d["delta_obs"].numel(), device=any_device) % 7 == 0, 2.0, 1.0 + ).to(d["delta_obs"].dtype) + sh = estimate_sigma_d( + d["delta_obs"], d["sigma_diff"], eps, d["d_star_sq"], d["f_dark"], d["fit_mask"] + ) + est = sigma_d_per_reflection(sh, d["d_star_sq"], eps, d["f_dark"], d["sigma_diff"]) + assert (est.w >= 0).all() and (est.w < 1).all() + assert est.S.device == d["delta_obs"].device + # The multiplier has shell mean one, so S / epsilon averages to Sigma_N over a shell. + counts = sh.counts.to(torch.long) # dtype-ok: split sizes; PyTorch requires int64 + order = torch.argsort(d["d_star_sq"]) + per_shell = torch.stack( + [c.mean() for c in torch.split((est.S / eps)[order], counts.tolist())] + ) + assert torch.allclose(per_shell, sh.Sigma_N, rtol=0.15) + # A missing dark amplitude means a multiplier of one. + f_missing = d["f_dark"].clone() + f_missing[:50] = float("nan") + est2 = sigma_d_per_reflection(sh, d["d_star_sq"], eps, f_missing, d["sigma_diff"]) + log_sn = interp_in_dss(d["d_star_sq"][:50], sh.bin_dss, torch.log(sh.Sigma_N)) + assert torch.allclose(est2.S[:50], eps[:50] * torch.exp(log_sn), rtol=1e-4) + + +@pytest.mark.unit +def test_deterministic_and_device_independent(any_device): + d_cpu = synth_diff() + a = estimate_sigma_d( + d_cpu["delta_obs"], + d_cpu["sigma_diff"], + None, + d_cpu["d_star_sq"], + d_cpu["f_dark"], + d_cpu["fit_mask"], + ) + b = estimate_sigma_d( + d_cpu["delta_obs"], + d_cpu["sigma_diff"], + None, + d_cpu["d_star_sq"], + d_cpu["f_dark"], + d_cpu["fit_mask"], + ) + assert torch.equal(a.Sigma_N, b.Sigma_N) and a.gamma == b.gamma + d_dev = synth_diff(device=any_device) + c = estimate_sigma_d( + d_dev["delta_obs"], + d_dev["sigma_diff"], + None, + d_dev["d_star_sq"], + d_dev["f_dark"], + d_dev["fit_mask"], + ) + assert torch.allclose(c.Sigma_N.cpu(), a.Sigma_N, rtol=1e-4) + assert abs(c.gamma - a.gamma) < 1e-3 + + +@pytest.mark.unit +def test_estimator_caches_until_reset_and_remaps(): + d = synth_diff(n=5000) + est = SigmaDEstimator(SigmaDConfig(gamma=1.0)) + first = est.get( + d["delta_obs"], + d["sigma_diff"], + None, + d["d_star_sq"], + d["f_dark"], + d["fit_mask"], + ) + assert ( + est.get( + d["delta_obs"], + d["sigma_diff"], + None, + d["d_star_sq"], + d["f_dark"], + d["fit_mask"], + ) + is first + ) + est.reset() + assert est._cache is None + target = d["d_star_sq"][:1000] + remapped = est.get( + d["delta_obs"], + d["sigma_diff"], + None, + d["d_star_sq"], + d["f_dark"], + d["fit_mask"], + target_dss=target, + out_f_dark=d["f_dark"][:1000], + out_sigma_diff=d["sigma_diff"][:1000], + ) + assert remapped.S.shape == (1000,) and est.shells is not None diff --git a/tests/unit/scaling/test_multiplicative_scale.py b/tests/unit/scaling/test_multiplicative_scale.py new file mode 100644 index 00000000..479e8ed4 --- /dev/null +++ b/tests/unit/scaling/test_multiplicative_scale.py @@ -0,0 +1,42 @@ +"""The scaler's observed-to-model factor without the bulk-solvent term. + +Pinned: an uninitialised scaler reports ones; after initialisation the factor +reproduces ``forward`` exactly once the additive solvent term is removed, so dividing +observed amplitudes by it returns them to the model's absolute scale. +""" + +import pytest +import torch + +from torchref.io import ReflectionData +from torchref.model.model_ft import ModelFT +from torchref.scaling.scaler import Scaler + + +@pytest.fixture +def scaler(sample_structure_pair): + model = ModelFT() + model.load_cif(str(sample_structure_pair["model"])) + data = ReflectionData(verbose=0) + data.load_mtz(str(sample_structure_pair["reflections"])) + return Scaler(model=model, data=data, nbins=10, verbose=0) + + +@pytest.mark.unit +def test_uninitialised_scaler_reports_ones(scaler): + factor = scaler.multiplicative_scale() + assert factor.shape == (int(scaler.bins.numel()),) + assert torch.equal(factor, torch.ones_like(factor)) + assert factor.device == scaler.device + + +@pytest.mark.integration +def test_factor_reproduces_forward_without_solvent(scaler): + scaler.initialize() + fcalc = scaler.compute_fcalc() + factor = scaler.multiplicative_scale() + assert (factor > 0).all() and torch.isfinite(factor).all() + assert not factor.requires_grad + with torch.no_grad(): + scaled = scaler(fcalc, f_sol_override=torch.zeros_like(fcalc)) + assert torch.allclose(scaled, factor.to(scaled.dtype) * fcalc, rtol=1e-5, atol=1e-6) diff --git a/torchref/cli/_common.py b/torchref/cli/_common.py index 8efb3201..135e62a8 100644 --- a/torchref/cli/_common.py +++ b/torchref/cli/_common.py @@ -389,6 +389,43 @@ def add_all_columns_arg(parser: argparse.ArgumentParser) -> None: ) +def add_ded_weight_args(parser: argparse.ArgumentParser) -> None: + """Add ``--ded-weight`` and ``--sigma-d-gamma`` for the difference-map writers. + + Every registered scheme's weight is written to the difference MTZ regardless; the + choice here decides which one the headline products (validate-ded correlations, + model-phased difference columns) carry. + """ + from torchref.maps.ded_weights import DEFAULT_SCHEME, SCHEMES + + parser.add_argument( + "--ded-weight", + choices=list(SCHEMES), + default=DEFAULT_SCHEME, + help="Per-reflection weight for difference coefficients: 'inverse_variance' " + "is 1/sigma^2, 'sigma_d' is the Wiener weight S/(S+sigma^2) from the " + "expected difference power (needs calibrated sigmas; check the reported " + f"clamped-shell count), 'none' is flat (default: {DEFAULT_SCHEME}). All " + "weights are written as columns.", + ) + parser.add_argument( + "--sigma-d-gamma", + type=float, + default=None, + metavar="GAMMA", + help="Fix the dark-amplitude exponent of the sigma_d power law in [0, 2] " + "instead of fitting it (default: fitted).", + ) + + +def sigma_d_config_from_args(args: argparse.Namespace): + """The :class:`~torchref.refinement.model_error_estimation.sigma_d.SigmaDConfig` + selected by ``--sigma-d-gamma``.""" + from torchref.refinement.model_error_estimation.sigma_d import SigmaDConfig + + return SigmaDConfig(gamma=getattr(args, "sigma_d_gamma", None)) + + def add_output_format_args(parser: argparse.ArgumentParser) -> None: """Add ``--output-format`` argument for coordinate file format.""" parser.add_argument( diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index c7242750..d35caf15 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -32,6 +32,7 @@ from torchref.cli._common import ( add_all_columns_arg, + add_ded_weight_args, add_dmin_arg, add_dual_model_args, add_general_args, @@ -46,9 +47,16 @@ parse_device_str, parse_weights, register_timing, + sigma_d_config_from_args, validate_cif_files, validate_files, ) +from torchref.maps.ded_weights import ( + DEFAULT_SCHEME, + WEIGHT_COLUMNS, + all_ded_weights, + reflection_geometry, +) from torchref.utils.serialization import convert_to_serializable configure_unbuffered_output() @@ -59,6 +67,8 @@ DEFAULT_TARGET_WEIGHTS = { "xray/difference": 1.0, + # Selected by --difference-target; the schedule drives whichever row is chosen. + "xray/difference_sd": 0.0, # The absolute channel. Zero by default: the difference refinement fixes the # dark model, so the overall level is already anchored and this term only adds # the systematic errors the difference cancels. @@ -219,9 +229,17 @@ def compute_rfactors(model, data, scaler): return rfactor_work_free(data, torch.abs(fcalc_scaled)) -def setup_loss_state(dataset_collection, model_collection, scaler, - target_weights, device, similarity_alpha=2.0, - two_moment=False): +def setup_loss_state( + dataset_collection, + model_collection, + scaler, + target_weights, + device, + similarity_alpha=2.0, + two_moment=False, + difference_target="difference", + sigma_d_config=None, +): """Build LossState with collection-aware targets. Geometry and ADP restraints are applied only to the light base model @@ -233,10 +251,17 @@ def setup_loss_state(dataset_collection, model_collection, scaler, Also register the two-moment intensity target, which fits merged intensities under ``|F(alpha)|^2 + sigma_alpha^2 |dF|^2``. Requires I/SIGI on every dataset. Default False. + difference_target : {"difference", "difference_sd"}, optional + Which difference row the weight schedule drives. Both are registered, as + ``xray/difference`` and ``xray/difference_sd``; the other keeps the weight in + ``target_weights`` (zero by default). + sigma_d_config : SigmaDConfig, optional + Exponent and shrinkage settings of the ``difference_sd`` row's estimator. """ from torchref.refinement import LossState from torchref.refinement.targets import TotalADPTarget, TotalGeometryTarget from torchref.refinement.targets.collection import ( + CollectionDifferenceSigmaDTarget, CollectionDifferenceTarget, CollectionMLTarget, ) @@ -250,6 +275,15 @@ def setup_loss_state(dataset_collection, model_collection, scaler, diff_target = CollectionDifferenceTarget( dataset_collection, model_collection, scaler=scaler, ) + diff_sd_target = CollectionDifferenceSigmaDTarget( + dataset_collection, + model_collection, + scaler=scaler, + sigma_d_config=sigma_d_config, + ) + selected_diff = {"difference": diff_target, "difference_sd": diff_sd_target}[ + difference_target + ] ml_target = CollectionMLTarget( dataset_collection, model_collection, scaler=scaler, ) @@ -261,6 +295,7 @@ def setup_loss_state(dataset_collection, model_collection, scaler, ) state.register_target("xray/difference", diff_target) + state.register_target("xray/difference_sd", diff_sd_target) state.register_target("xray/ml", ml_target) state.register_target("geometry", geom_target) state.register_target("adp", adp_target) @@ -275,7 +310,7 @@ def setup_loss_state(dataset_collection, model_collection, scaler, # Match intensity and amplitude gradient norms to balance the targets # against the geometry restraints despite their different units. two_moment_target.calibrate_base_weight( - diff_target, list(model_light.parameters()) + selected_diff, list(model_light.parameters()) ) state.register_target("xray/two_moment", two_moment_target) @@ -285,8 +320,16 @@ def setup_loss_state(dataset_collection, model_collection, scaler, def compute_bayes_extrapolated_amplitudes( - Fobs_dark, Fobs_light, sig_ext, phi_dark, phi_mixed, f, - *, tau_sq_floor=1e-4, + Fobs_dark, + Fobs_light, + sig_ext, + phi_dark, + phi_mixed, + f, + *, + tau_sq_floor=1e-4, + epsilon=None, + d_star_sq=None, ): """Empirical Bayes shrinkage estimator for extrapolated SF amplitudes. @@ -295,10 +338,16 @@ def compute_bayes_extrapolated_amplitudes( Fo_dark, regularising noisy high-resolution and weakly-measured reflections:: F_ext = |F_dark*e^(iφ_d) + ΔF/f| (phase-aware amplitude) - τ² = max(<(F_ext - Fo_dark)²> - <σ_ext²>, floor) - w(h) = τ² / (τ² + σ_ext²(h)) + S(h) = expected power of (F_ext - Fo_dark), per resolution shell + w(h) = S(h) / (S(h) + σ_ext²(h)) F_extb = w(h)·F_ext + (1-w(h))·Fo_dark (amplitude shrinkage) + With ``d_star_sq`` the signal power comes per resolution shell from + :func:`~torchref.refinement.model_error_estimation.sigma_d.estimate_sigma_d` + (``<(F_ext - Fo_dark)²> - <σ_ext²>`` per shell, shrunk toward a smooth curve); + without it the single global ``τ² = max(<(F_ext - Fo_dark)²> - <σ_ext²>, floor)`` + is used, which is the one-shell special case. + Parameters ---------- Fobs_dark, Fobs_light : Tensor (N,) @@ -313,14 +362,17 @@ def compute_bayes_extrapolated_amplitudes( f : float or Tensor Excited-state population fraction. tau_sq_floor : float - Floor on the estimated signal variance τ². + Floor on the estimated signal variance. + epsilon, d_star_sq : Tensor (N,), optional + Reflection multiplicity and ``1/d**2`` in A^-2. Given ``d_star_sq`` the signal + power is estimated per resolution shell. Returns ------- tuple ``(F_ext_bayes, var_ext_bayes, w_shrinkage, tau_sq)`` -- the **shrunk** extrapolated amplitude, its posterior variance and the shrinkage weight per - reflection, and the global τ² as a float. + reflection, and the count-weighted mean signal variance as a float. """ F_dark_phased = Fobs_dark * torch.exp(1j * phi_dark) F_light_phased = Fobs_light * torch.exp(1j * phi_mixed) @@ -332,15 +384,34 @@ def compute_bayes_extrapolated_amplitudes( F_ext_complex = F_dark_phased + delta_F / f F_ext = torch.abs(F_ext_complex) - # Estimate signal variance τ² - residuals_sq = (F_ext - Fobs_dark) ** 2 - tau_sq = max((residuals_sq.mean() - sig_sq_ext.mean()).item(), tau_sq_floor) + residual = F_ext - Fobs_dark + if d_star_sq is None: + tau_sq = max( + (residual.square().mean() - sig_sq_ext.mean()).item(), tau_sq_floor + ) + S = torch.full_like(F_ext, tau_sq) + else: + from torchref.refinement.model_error_estimation.sigma_d import ( + estimate_sigma_d, + sigma_d_per_reflection, + ) + + fit = torch.isfinite(residual) & torch.isfinite(sig_ext) + shells = estimate_sigma_d( + residual, sig_ext, epsilon, d_star_sq, None, fit, gamma=0.0 + ) + est = sigma_d_per_reflection(shells, d_star_sq, epsilon, None, sig_ext) + S = est.S.clamp(min=tau_sq_floor) + weight = shells.counts.clamp(min=1.0) + tau_sq = max( + float((shells.Sigma_N * shells.counts).sum() / weight.sum()), tau_sq_floor + ) # Per-reflection shrinkage weight (in [0, 1]) - w = tau_sq / (tau_sq + sig_sq_ext) + w = S / (S + sig_sq_ext) # Posterior variance - var_ext_bayes = (tau_sq * sig_sq_ext) / (tau_sq + sig_sq_ext) + var_ext_bayes = (S * sig_sq_ext) / (S + sig_sq_ext) # Shrink the amplitude toward Fo_dark -- scalar, so no phase interference. F_ext_bayes = w * F_ext + (1 - w) * Fobs_dark @@ -485,16 +556,34 @@ def _np(t): return columns, types -def _difference_columns(data_dark, data_light, mask, hkl_np, *, Fobs_dark, sig_dark, - Fobs_light, sig_light, Fcalc_dark, phases_dark, diff_Fobs, - sig_diff, weights): - """The weighted difference map, and the observations behind it. - - ``DELFWT``/``PHDELWT`` is the inverse-variance-weighted amplitude difference carried - on the **dark** model's phases -- the isomorphous difference Fourier, and the same - construction ``torchref.validate-ded`` correlates against, so the map in this file - and the map the validation reports are one object. CCP4 and Coot recognise the names - and open it as a difference map without being told which columns to use. +def _difference_columns( + data_dark, + data_light, + mask, + hkl_np, + *, + Fobs_dark, + sig_dark, + Fobs_light, + sig_light, + Fcalc_dark, + phases_dark, + diff_Fobs, + sig_diff, + weight_columns, + kscale, +): + """The difference map's amplitudes, phases and weights. + + ``DF``/``SIGDF`` is the signed amplitude difference ``|Fo_light| - |Fo_dark|`` with + its propagated uncertainty, ``PHDELWT`` the **dark** model's phase it is carried on: + the isomorphous difference Fourier, and the construction ``torchref.validate-ded`` + correlates against. One weight column per registered scheme (``W_SD``, ``W_IVW``; + MTZ type ``W``, mean one) sits beside it, so any weighting is ``DF`` times a column + and reproducible from the file: ``torchref.mtz2map -csf DF -cw W_SD -cphi PHDELWT``. + ``KSCALE`` (type ``R``) is the scaler's multiplicative factor from model to observed + scale, so ``DF / KSCALE`` is in electrons and ``mtz2map --units electrons`` gives + e/A^3. This layer needs no light-state model: the amplitude is ``|Fo_light| - |Fo_dark|`` and the phase comes from the dark model. Keeping the light state's model out is the @@ -511,11 +600,18 @@ def _flags(data): return data.rfree_flags[mask].cpu().numpy().astype(int) columns = { - "H": hkl_np[:, 0], "K": hkl_np[:, 1], "L": hkl_np[:, 2], - "Fo_dark": Fobs_dark, "SIGFo_dark": sig_dark, - "Fo_light": Fobs_light, "SIGFo_light": sig_light, - "DF": diff_Fobs, "SIGDF": sig_diff, - "DELFWT": diff_Fobs * weights, "PHDELWT": phases_dark, + "H": hkl_np[:, 0], + "K": hkl_np[:, 1], + "L": hkl_np[:, 2], + "Fo_dark": Fobs_dark, + "SIGFo_dark": sig_dark, + "Fo_light": Fobs_light, + "SIGFo_light": sig_light, + "DF": diff_Fobs, + "SIGDF": sig_diff, + "PHDELWT": phases_dark, + **weight_columns, + "KSCALE": kscale, "Fc_dark": Fcalc_dark, # 1 = work, 0 = free. Both are kept: the two datasets can disagree, and # picking one would silently report an R-free against the wrong test set. @@ -523,13 +619,21 @@ def _flags(data): "FreeR_flag_light": _flags(data_light), } types = { - "H": "H", "K": "H", "L": "H", - "Fo_dark": "F", "SIGFo_dark": "Q", - "Fo_light": "F", "SIGFo_light": "Q", - "DF": "F", "SIGDF": "Q", - "DELFWT": "F", "PHDELWT": "P", + "H": "H", + "K": "H", + "L": "H", + "Fo_dark": "F", + "SIGFo_dark": "Q", + "Fo_light": "F", + "SIGFo_light": "Q", + "DF": "F", + "SIGDF": "Q", + "PHDELWT": "P", + **{name: "W" for name in weight_columns}, + "KSCALE": "R", "Fc_dark": "F", - "FreeR_flag_dark": "I", "FreeR_flag_light": "I", + "FreeR_flag_dark": "I", + "FreeR_flag_light": "I", } return columns, types @@ -604,9 +708,22 @@ def _phasing_columns(mc, scaler, hkl_all, mask, *, fcalc_dark, Fobs_dark_vals, return columns, types, ctx -def _extrapolation_columns(mc, dc, hkl, *, Fobs_dark_vals, Fobs_light_vals, - sig_dark_vals, sig_light_vals, phi_dark, ctx, - rfree_flags_masked, all_columns=False, verbose=1): +def _extrapolation_columns( + mc, + dc, + hkl, + *, + Fobs_dark_vals, + Fobs_light_vals, + sig_dark_vals, + sig_light_vals, + phi_dark, + ctx, + rfree_flags_masked, + all_columns=False, + verbose=1, + geometry=None, +): """Extrapolated light-state amplitudes and the map to refine against. Three constructions of the same quantity, all needing the light model: @@ -654,10 +771,17 @@ def _fit(amp, sig): sig_light_vals**2 + w_dark**2 * sig_dark_vals**2 ) / w_light + eps, dss = geometry if geometry is not None else (None, None) F_ext_bayes_amp, var_ext_bayes, w_shrinkage, tau_sq = ( compute_bayes_extrapolated_amplitudes( - Fobs_dark_vals, Fobs_light_vals, sig_light_extra, - phi_dark, ctx["phi_mixed"], w_light, + Fobs_dark_vals, + Fobs_light_vals, + sig_light_extra, + phi_dark, + ctx["phi_mixed"], + w_light, + epsilon=eps, + d_star_sq=dss, ) ) sig_ext_bayes = torch.sqrt(var_ext_bayes) @@ -724,13 +848,26 @@ def _np(t): return columns, types, diagnostics -def write_results_mtz(dc, dark_model, scaler, filename, *, mc=None, - all_columns=False, verbose=1): +def write_results_mtz( + dc, + dark_model, + scaler, + filename, + *, + mc=None, + all_columns=False, + verbose=1, + ded_weight=DEFAULT_SCHEME, + sigma_d_config=None, +): """Write the difference map, and map coefficients when a light model is given. - The default output is the **weighted difference map**: ``DELFWT``/``PHDELWT``, the - inverse-variance-weighted amplitude difference on the dark model's phases. That needs - no light-state model, which is why ``mc`` is optional -- with a dark model alone this + The default output is the **difference map**: ``DF``/``SIGDF`` on the dark model's + phases ``PHDELWT``, with one mean-one weight column per registered scheme + (``W_SD``, ``W_IVW``) and the observed-to-model scale ``KSCALE``; see + :func:`_difference_columns`. ``ded_weight`` selects the scheme the model-phased + difference columns and the two-moment columns are weighted with. That needs no + light-state model, which is why ``mc`` is optional -- with a dark model alone this writes a difference map and nothing else, and no scale fit is run beyond the one that produced ``scaler``. @@ -752,6 +889,11 @@ def write_results_mtz(dc, dark_model, scaler, filename, *, mc=None, The dark+light collection. Absent means difference map only. filename : str Output MTZ path. + ded_weight : str, optional + Weight scheme for the model-phased and two-moment difference columns; one of + :data:`torchref.maps.ded_weights.SCHEMES`. + sigma_d_config : SigmaDConfig, optional + Exponent and shrinkage settings of the ``sigma_d`` scheme. Returns ------- @@ -801,20 +943,79 @@ def write_results_mtz(dc, dark_model, scaler, filename, *, mc=None, Fcalc_dark = torch.abs(fcalc_dark).detach().cpu().numpy() phases_dark = phi_dark.detach().rad2deg().cpu().numpy() - diff_Fobs = Fobs_light - Fobs_dark - sig_diff = (sig_dark**2 + sig_light**2) ** 0.5 - weights = 1 / sig_diff**2 - weights = weights / weights.mean() + diff_t = Fobs_light_vals - Fobs_dark_vals + sig_diff_t = torch.sqrt(sig_dark_vals**2 + sig_light_vals**2) + all_w = all_ded_weights( + delta_obs=diff_t, + sigma_diff=sig_diff_t, + hkl=hkl, + cell=data_dark.cell, + spacegroup=data_dark.spacegroup, + f_dark=Fobs_dark_vals, + sigma_d_config=sigma_d_config, + ) + selected = all_w[ded_weight] + weights = selected.weights.detach().cpu().numpy() + diff_Fobs = diff_t.detach().cpu().numpy() + sig_diff = sig_diff_t.detach().cpu().numpy() + weight_columns = { + WEIGHT_COLUMNS[name]: all_w[name].weights.detach().cpu().numpy() + for name in WEIGHT_COLUMNS + } + kscale = scaler.multiplicative_scale()[mask].detach().cpu().numpy() + geometry = reflection_geometry( + hkl, data_dark.cell, data_dark.spacegroup, diff_t.device, diff_t.dtype + ) + sd_diag = { + k: v + for k, v in all_w["sigma_d"].diagnostics.items() + if k != "weight_sigma_d_raw" + } + diagnostics = { + "ded_weights": { + "scheme": ded_weight, + "applied": selected.applied, + "sigma_d": sd_diag, + } + } + if verbose > 0: + print(f" Difference weights: {ded_weight} (applied: {selected.applied})") + print( + f" sigma_D: gamma = {sd_diag['gamma']:.3f} ({sd_diag['gamma_reason']}), " + f"tau = {sd_diag['tau']:.3f}, shells = {sd_diag['n_shell']}, " + f"shells without difference power = {sd_diag['n_s2_clamped']}" + ) + if "fallback_reason" in sd_diag: + print(f" sigma_D fallback: {sd_diag['fallback_reason']}") + if verbose > 1 and not sd_diag["degenerate"]: + table = sd_diag["shells"] + print(" sigma_D shells: d(A) n B S2 Sigma_N") + for dss, n, b, s2, sn in zip( + table["d_star_sq"], + table["counts"], + table["B"], + table["S2"], + table["Sigma_N"], + ): + print(f" {dss ** -0.5:6.2f} {int(n):5d} {b:9.4f} {s2:9.4f} {sn:9.4f}") columns, types = _difference_columns( - data_dark, data_light, mask, hkl_np, - Fobs_dark=Fobs_dark, sig_dark=sig_dark, - Fobs_light=Fobs_light, sig_light=sig_light, - Fcalc_dark=Fcalc_dark, phases_dark=phases_dark, - diff_Fobs=diff_Fobs, sig_diff=sig_diff, weights=weights, + data_dark, + data_light, + mask, + hkl_np, + Fobs_dark=Fobs_dark, + sig_dark=sig_dark, + Fobs_light=Fobs_light, + sig_light=sig_light, + Fcalc_dark=Fcalc_dark, + phases_dark=phases_dark, + diff_Fobs=diff_Fobs, + sig_diff=sig_diff, + weight_columns=weight_columns, + kscale=kscale, ) - diagnostics = {} if mc is not None: phase_cols, phase_types, ctx = _phasing_columns( mc, scaler, hkl_all, mask, @@ -825,13 +1026,22 @@ def write_results_mtz(dc, dark_model, scaler, filename, *, mc=None, columns.update(phase_cols) types.update(phase_types) - ext_cols, ext_types, diagnostics = _extrapolation_columns( - mc, dc, hkl, - Fobs_dark_vals=Fobs_dark_vals, Fobs_light_vals=Fobs_light_vals, - sig_dark_vals=sig_dark_vals, sig_light_vals=sig_light_vals, - phi_dark=phi_dark, ctx=ctx, rfree_flags_masked=rfree_flags_masked, - all_columns=all_columns, verbose=verbose, + ext_cols, ext_types, ext_diagnostics = _extrapolation_columns( + mc, + dc, + hkl, + Fobs_dark_vals=Fobs_dark_vals, + Fobs_light_vals=Fobs_light_vals, + sig_dark_vals=sig_dark_vals, + sig_light_vals=sig_light_vals, + phi_dark=phi_dark, + ctx=ctx, + rfree_flags_masked=rfree_flags_masked, + all_columns=all_columns, + verbose=verbose, + geometry=geometry, ) + diagnostics.update(ext_diagnostics) columns.update(ext_cols) types.update(ext_types) @@ -931,8 +1141,18 @@ def main(): add_output_format_args(output) add_metadata_args(output) add_all_columns_arg(output) + add_ded_weight_args(output) refine = parser.add_argument_group("Refinement") + refine.add_argument( + "--difference-target", + choices=("difference", "difference_sd"), + default="difference", + help="Difference row the weight schedule drives: 'difference' is the Gaussian " + "under the measurement variance, 'difference_sd' centres on " + "alpha*dF_calc with the sigma_D unexplained power added to the variance " + "(default: difference).", + ) refine.add_argument( "--weight-schedule", type=str, default="5,3,2", help="Comma-separated difference-target weights applied in " @@ -1037,7 +1257,10 @@ def main(): # --- Parse and merge target weights --- target_weights = dict(DEFAULT_TARGET_WEIGHTS) - target_weights["xray/difference"] = weight_schedule[0] + difference_key = f"xray/{args.difference_target}" + target_weights["xray/difference"] = 0.0 + target_weights["xray/difference_sd"] = 0.0 + target_weights[difference_key] = weight_schedule[0] target_weights["similarity"] = args.similarity_weight target_weights, err = parse_weights(args.weights, defaults=target_weights) if err: @@ -1175,9 +1398,17 @@ def main(): # per-reflection weighting, which needs no intensity data. mc.set_lambda_twin(args.lambda_twin, refinable=args.refine_lambda_twin) - state = setup_loss_state(dc, mc, scaler, target_weights, device, - similarity_alpha=args.similarity_alpha, - two_moment=args.two_moment) + state = setup_loss_state( + dc, + mc, + scaler, + target_weights, + device, + similarity_alpha=args.similarity_alpha, + two_moment=args.two_moment, + difference_target=args.difference_target, + sigma_d_config=sigma_d_config_from_args(args), + ) if args.verbose > 0: print("Initial loss breakdown:") @@ -1217,7 +1448,7 @@ def main(): ) sys.stdout.flush() - state.set_weights({"xray/difference": t_weight}) + state.set_weights({difference_key: t_weight}) optimize_lbfgs( state, params, max_iter=args.max_iter, @@ -1471,8 +1702,15 @@ def _mtz_to_cif(mtz_path, cif_path): print(f" Light SF written to {light_sf_mtz}, {light_sf_cif}") map_diagnostics = write_results_mtz( - dc, mc.dark_model, scaler, diff_mtz_out, - mc=mc, all_columns=args.all_columns, verbose=args.verbose, + dc, + mc.dark_model, + scaler, + diff_mtz_out, + mc=mc, + all_columns=args.all_columns, + verbose=args.verbose, + ded_weight=args.ded_weight, + sigma_d_config=sigma_d_config_from_args(args), ) # --- JSON summary --- diff --git a/torchref/cli/difference_map.py b/torchref/cli/difference_map.py index 8253ee5c..54f4bc6c 100644 --- a/torchref/cli/difference_map.py +++ b/torchref/cli/difference_map.py @@ -5,12 +5,15 @@ Uses the ``torchref.difference-refine`` pipeline but performs **no refinement**: the input models are used as-is. -The default output is the weighted difference map -- the inverse-variance-weighted -amplitude difference ``|Fo_light| - |Fo_dark|`` carried on the **dark** model's phases, -written as ``DELFWT``/``PHDELWT``. That needs no light-state model, so ``-lm`` is -optional. It is also deliberately not a *phased* difference map: putting the light -state's model phases into the observed amplitude biases the map toward the very model -the experiment is testing. +The default output is the difference map: the amplitude difference +``|Fo_light| - |Fo_dark|`` as ``DF``/``SIGDF`` carried on the **dark** model's phases +``PHDELWT``, with the per-reflection weights of every registered scheme beside it as +``W_IVW`` (inverse variance, the default) and ``W_SD`` (sigma_D Wiener weight), and the +observed-to-model scale ``KSCALE``. Build the map with +``torchref.mtz2map -csf DF -cw W_IVW -cphi PHDELWT`` (``--units electrons`` for e/A^3). +That needs no light-state model, so ``-lm`` is optional. It is also deliberately not a +*phased* difference map: putting the light state's model phases into the observed +amplitude biases the map toward the very model the experiment is testing. Given ``-lm``, the light state's amplitude and phase and the extrapolated map follow. ``--all-columns`` adds the alternative constructions of both. @@ -36,6 +39,7 @@ from torchref.cli._common import ( add_all_columns_arg, + add_ded_weight_args, add_dual_model_args, add_dmin_arg, add_general_args, @@ -44,6 +48,7 @@ configure_unbuffered_output, register_timing, parse_device_str, + sigma_d_config_from_args, validate_cif_files, validate_files, ) @@ -81,6 +86,7 @@ def main(): output = parser.add_argument_group("Output") add_output_arg(output, help="Output MTZ file path (e.g. results.mtz)") add_all_columns_arg(output) + add_ded_weight_args(output) res = parser.add_argument_group("Resolution") add_dmin_arg(res) @@ -215,8 +221,15 @@ def main(): with torch.no_grad(): write_results_mtz( - dc, dark_model, scaler, str(out_path), - mc=mc, all_columns=args.all_columns, verbose=args.verbose, + dc, + dark_model, + scaler, + str(out_path), + mc=mc, + all_columns=args.all_columns, + verbose=args.verbose, + ded_weight=args.ded_weight, + sigma_d_config=sigma_d_config_from_args(args), ) if args.verbose > 0: diff --git a/torchref/cli/mtz2map.py b/torchref/cli/mtz2map.py index d02d75ac..e98fe12b 100644 --- a/torchref/cli/mtz2map.py +++ b/torchref/cli/mtz2map.py @@ -59,6 +59,25 @@ def main(): metavar="COL", help="Column name for phases in degrees (e.g. PHWT, PHDELWT, PH2FOFCWT).", ) + inp.add_argument( + "-cw", + "--column-weight", + default=None, + type=str, + metavar="COL", + help="Weight column multiplied into the amplitudes before the FFT " + "(e.g. W_SD, W_IVW from torchref.difference-map). Default: none.", + ) + inp.add_argument( + "-ck", + "--column-scale", + default=None, + type=str, + metavar="COL", + help="Per-reflection observed-to-model scale factor the amplitudes are divided " + "by for --units electrons (e.g. KSCALE from torchref.difference-map). " + "Default: KSCALE when the file has it.", + ) output = parser.add_argument_group("Output") output.add_argument( @@ -74,15 +93,24 @@ def main(): metavar=("NX", "NY", "NZ"), help="Override grid dimensions. Default: auto from cell and resolution.", ) + mapopts.add_argument( + "--units", + type=str, + choices=["sigma", "electrons", "raw"], + default=None, + help="Map units. 'sigma': zero mean and unit standard deviation (default). " + "'electrons': electrons per cubic Angstrom, sum_h F(h) exp(-2 pi i h.x) / V " + "with F divided by the --column-scale factor. 'raw': the plain FFT with the " + "1/N normalisation, no rescaling.", + ) mapopts.add_argument( "-n", "--normalize", type=str, - choices=['True', 'False'], - default='True', - help="Normalize amplitudes to unit variance. Accepts only the literal " - "strings 'True' or 'False' (case-sensitive); pass '-n False' to disable. " - "Default: True.", + choices=["True", "False"], + default=None, + help="Deprecated alias: '-n True' is '--units sigma', '-n False' is " + "'--units raw'.", ) res = parser.add_argument_group("Resolution") @@ -106,7 +134,24 @@ def main(): mtz = rs.read_mtz(args.structure_factor) available = list(mtz.columns) - normalize = args.normalize == 'True' + if args.units is not None and args.normalize is not None: + print("Error: --units and --normalize cannot both be given", file=sys.stderr) + sys.exit(1) + if args.units is not None: + units = args.units + elif args.normalize is not None: + units = "sigma" if args.normalize == "True" else "raw" + else: + units = "sigma" + scale_column = args.column_scale + if scale_column is None and units == "electrons" and "KSCALE" in available: + scale_column = "KSCALE" + if units == "electrons" and scale_column is None: + print( + "Error: --units electrons needs --column-scale (no KSCALE column found).", + file=sys.stderr, + ) + sys.exit(1) if args.column_structure_factor not in available: print( @@ -122,6 +167,14 @@ def main(): file=sys.stderr, ) sys.exit(1) + for label, col in (("weight", args.column_weight), ("scale", scale_column)): + if col is not None and col not in available: + print( + f"Error: {label} column '{col}' not found.\n" + f"Available columns: {available}", + file=sys.stderr, + ) + sys.exit(1) # Extract cell and spacegroup cell = np.array( @@ -141,9 +194,23 @@ def main(): hkl = df[["H", "K", "L"]].to_numpy().astype(np.int32) amplitudes = df[args.column_structure_factor].to_numpy().astype(np.float32) phases_deg = df[args.column_phase].to_numpy().astype(np.float32) + valid = np.isfinite(amplitudes) & np.isfinite(phases_deg) + if args.column_weight is not None: + weights = df[args.column_weight].to_numpy().astype(np.float32) + valid &= np.isfinite(weights) + amplitudes = amplitudes * weights + if args.verbose >= 1: + print(f" Weights: {args.column_weight} (mean {np.nanmean(weights):.3f})") + if units == "electrons": + kscale = df[scale_column].to_numpy().astype(np.float32) + valid &= np.isfinite(kscale) & (kscale > 0) + # Observed amplitudes carry the scaler's overall scale, B and anisotropy; + # dividing by that factor returns them to electrons. + amplitudes = amplitudes / np.where(valid, kscale, 1.0) + if args.verbose >= 1: + print(f" Absolute scale: dividing by {scale_column}") # Drop NaN reflections - valid = np.isfinite(amplitudes) & np.isfinite(phases_deg) if not valid.all(): n_drop = (~valid).sum() if args.verbose >= 1: @@ -215,12 +282,17 @@ def main(): grid = place_on_grid(hkl_p1, coefficients, gridsize, enforce_hermitian=True) - # FFT to real space: rho(r) = sum_h F(h) * exp(-2*pi*i * h.r) + # FFT to real space with the 1/N normalisation: rho_raw(r) = (1/N) sum_h F(h) exp(-2 pi i h.r) real_map = torch.fft.fftn(grid, dim=(0, 1, 2), norm="forward").real - if normalize: + if units == "sigma": real_map = (real_map - real_map.mean()) / real_map.std() - + elif units == "electrons": + # rho(r) = (1/V) sum_h F(h) exp(-2 pi i h.r): undo the 1/N and divide by the + # cell volume, so the map is in electrons per cubic Angstrom. + volume = Cell(cell, device=device).volume.to(real_map.dtype) + real_map = real_map * (real_map.numel() / volume) + # --- Write output --- from torchref.io.cif import write_map @@ -229,6 +301,7 @@ def main(): if args.verbose >= 1: print(f" Written: {args.output}") sigma = float(real_map.std()) + print(f" Units: {units}") print(f" Map sigma: {sigma:.4f}") diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index 8c94a52d..18b6d9a7 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -24,12 +24,14 @@ import argparse import json import sys +import warnings from pathlib import Path import numpy as np import torch from torchref.cli._common import ( + add_ded_weight_args, add_dmin_arg, add_dual_model_args, add_general_args, @@ -40,9 +42,15 @@ load_reflection_data, parse_device_str, register_timing, + sigma_d_config_from_args, validate_cif_files, validate_files, ) +from torchref.maps.ded_weights import ( + DEFAULT_SCHEME, + DedWeightFallbackWarning, + all_ded_weights, +) from torchref.utils.serialization import convert_to_serializable configure_unbuffered_output() @@ -196,6 +204,8 @@ def setup_ded_context( col_light=None, n_bins=20, verbose=0, + ded_weight=DEFAULT_SCHEME, + sigma_d_config=None, ): """Load reflection data and prepare shared state for DED validation. @@ -280,11 +290,20 @@ def setup_ded_context( else: free_mask = work_mask = None - # Weighted difference Fo + # Difference Fo and the registered weights; the selected scheme is the headline. dfo = F_light - F_dark - sig_diff = (sig_dark**2 + sig_light**2) ** 0.5 - weights = 1 / sig_diff**2 - weights = weights / weights.mean() + sig_diff = torch.sqrt(sig_dark**2 + sig_light**2) + all_w = all_ded_weights( + delta_obs=dfo, + sigma_diff=sig_diff, + hkl=hkl, + cell=data_dark.cell, + spacegroup=data_dark.spacegroup, + f_dark=F_dark, + sigma_d_config=sigma_d_config, + ) + selected = all_w[ded_weight] + weights = selected.weights w_dfo = dfo * weights # Cell, spacegroup, d-spacings @@ -309,6 +328,9 @@ def setup_ded_context( ) w_dfo_p1 = w_dfo[orig_idx] weights_p1 = weights[orig_idx] + weights_by_scheme = { + name: (w.weights, w.weights[orig_idx]) for name, w in all_w.items() + } if verbose >= 1: print(f"Matched reflections: {len(hkl)}") @@ -326,6 +348,16 @@ def setup_ded_context( "refl_mask": refl_mask, "w_dfo": w_dfo, "weights": weights, + "dfo": dfo, + "dfo_p1": dfo[orig_idx], + "weights_by_scheme": weights_by_scheme, + "ded_weight": ded_weight, + "ded_weight_applied": selected.applied, + "ded_weight_diagnostics": { + k: v + for k, v in all_w["sigma_d"].diagnostics.items() + if k not in ("weight_sigma_d_raw", "shells") + }, "d_spacing": d_spacing, "cell_t": cell_t, "cell_np": cell_np, @@ -526,10 +558,48 @@ def compute_ded_maps( if cc_work is not None: print(f" Work CC = {cc_work:.4f}, Free CC = {cc_free:.4f}") + # Every registered scheme on the same coefficients, for side-by-side reporting. + by_weight = {} + for name, (w_asu, w_p1) in ctx.get("weights_by_scheme", {}).items(): + with torch.no_grad(): + m_o = compute_map_from_coefficients( + ctx["dfo_p1"] * w_p1, phi_dark_p1, ctx["hkl_p1"], ctx["gridsize"] + ) + m_c = compute_map_from_coefficients( + delta_fcalc * w_p1, phi_dark_p1, ctx["hkl_p1"], ctx["gridsize"] + ) + entry = { + "realspace_correlation": { + mname: round(float(compute_correlation(m_o, m_c, mm)), 4) + for mname, mm in mask_dict.items() + } + } + wo, wc = ctx["dfo"] * w_asu, delta_fcalc_asu * w_asu + entry["reciprocal_cc_overall"] = round( + torch.corrcoef(torch.stack([wo, wc]))[0, 1].item(), 4 + ) + if free_mask is not None and free_mask.sum() > 10: + entry["reciprocal_cc_work"] = round( + torch.corrcoef(torch.stack([wo[work_mask], wc[work_mask]]))[ + 0, 1 + ].item(), + 4, + ) + entry["reciprocal_cc_free"] = round( + torch.corrcoef(torch.stack([wo[free_mask], wc[free_mask]]))[ + 0, 1 + ].item(), + 4, + ) + else: + entry["reciprocal_cc_work"] = entry["reciprocal_cc_free"] = None + by_weight[name] = entry + return { "map_dfo": map_dfo, "map_dfc": map_dfc, "mask_dict": mask_dict, + "by_weight": by_weight, "realspace_correlation": rs_corr, "resolution_bins": bin_results, "reciprocal_cc_overall": round(cc_overall, 4), @@ -567,16 +637,27 @@ def run_validation(args): print(f" Light SF: {args.light_structure_factor}") col_dark, col_light = build_dual_column_names(args) - ctx = setup_ded_context( - args.dark_structure_factor, - args.light_structure_factor, - dmin=args.dmin, - device=device, - col_dark=col_dark, - col_light=col_light, - n_bins=args.n_bins, - verbose=args.verbose, - ) + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always", DedWeightFallbackWarning) + ctx = setup_ded_context( + args.dark_structure_factor, + args.light_structure_factor, + dmin=args.dmin, + device=device, + col_dark=col_dark, + col_light=col_light, + n_bins=args.n_bins, + verbose=args.verbose, + ded_weight=args.ded_weight, + sigma_d_config=sigma_d_config_from_args(args), + ) + fallback_messages = [ + str(w.message) + for w in caught + if issubclass(w.category, DedWeightFallbackWarning) + ] + for message in fallback_messages: + print(f"WARNING: {message}") d_min = ctx["d_min"] # Load models @@ -641,6 +722,23 @@ def run_validation(args): "mask_radius": args.mask_radius, "dmin": d_min, }, + "weights": { + "requested": ctx["ded_weight"], + "applied": ctx["ded_weight_applied"], + **{ + k: ctx["ded_weight_diagnostics"].get(k) + for k in ( + "gamma", + "gamma_fitted", + "gamma_reason", + "tau", + "n_shell", + "n_s2_clamped", + "fallback_reason", + ) + }, + }, + "by_weight": result["by_weight"], "realspace_correlation": result["realspace_correlation"], "reciprocal_cc_overall": result["reciprocal_cc_overall"], "reciprocal_cc_work": result["reciprocal_cc_work"], @@ -689,10 +787,25 @@ def run_validation(args): # Summary if args.verbose >= 1: print(f"\n{'=' * 70}") - print("Summary:") + print( + f"Summary (headline weights: {ctx['ded_weight']}, " + f"applied: {ctx['ded_weight_applied']}):" + ) for name, corr in result["realspace_correlation"].items(): print(f" {name}: CC = {corr['cc']:.4f}") print(f" Reciprocal-space CC (overall): {result['reciprocal_cc_overall']}") + masks = list(result["realspace_correlation"]) + header = " {:<17s}".format("weights") + "".join( + f"{m[:12]:>13s}" for m in masks + ) + print(header + f"{'recip. CC':>13s}") + for name, entry in result["by_weight"].items(): + row = f" {name:<17s}" + "".join( + f"{entry['realspace_correlation'][m]:13.4f}" for m in masks + ) + print(row + f"{entry['reciprocal_cc_overall']:13.4f}") + for message in fallback_messages: + print(f" WARNING: {message}") print(f"{'=' * 70}") return 0 @@ -784,6 +897,7 @@ def main(): action="store_true", help="Write CCP4 map files for WDFo and WDFcalc", ) + add_ded_weight_args(analysis) res = parser.add_argument_group("Resolution") add_dmin_arg(res) diff --git a/torchref/maps/ded_weights.py b/torchref/maps/ded_weights.py new file mode 100644 index 00000000..e09a93b8 --- /dev/null +++ b/torchref/maps/ded_weights.py @@ -0,0 +1,253 @@ +"""Registered per-reflection weights for light-minus-dark difference coefficients. + +A weight scheme turns the observed differences and their uncertainties into one weight +per reflection, normalised to mean one so that maps built from different schemes sit on +comparable scales. Three schemes are registered: + +``none`` + Every reflection weighted equally. +``inverse_variance`` + ``1 / sigma_diff**2``. Weights by precision alone; the right rule for averaging + estimates of one quantity, and the default for difference maps. +``sigma_d`` + The Wiener weight ``S / (S + sigma_diff**2)`` with ``S`` the expected true difference + power from :mod:`torchref.refinement.model_error_estimation.sigma_d`. Weights by the + signal fraction of each coefficient, so strong reflections whose expected difference + is large keep their weight. ``S`` is ``mean(dF**2) - mean(sigma**2)`` per shell, so + it inherits any miscalibration of ``sigma_diff``: where the reported sigmas are too + large the estimate finds no power and the weight vanishes, which turns the scheme + into a resolution cut. The count of such shells is reported as ``n_s2_clamped``; + a large fraction means the sigmas, not the data, are deciding the map. + +Plain tensors in and out. The weights live on the device of ``delta_obs``. The sigma_D +estimator is imported inside the scheme that needs it so that :mod:`torchref.maps` does +not import :mod:`torchref.refinement` at module load. +""" + +import warnings +from dataclasses import dataclass, field +from typing import TYPE_CHECKING + +import torch + +from torchref.base.reciprocal.basis import get_scattering_vectors +from torchref.base.targets.xray_likelihoods import SIGMA_FLOOR_ABS, SIGMA_FLOOR_FRAC + +if TYPE_CHECKING: + from torchref.refinement.model_error_estimation.sigma_d import SigmaDConfig + +#: The selectable schemes, in the order they are reported. +SCHEMES = ("none", "inverse_variance", "sigma_d") +#: Scheme applied when none is named. Inverse variance, because ``sigma_d`` depends on +#: calibrated sigmas: on the 15 Sep campaign TorchSX's TD1 sigmas were ~1.5x too large at +#: high resolution, ``sigma_d`` zeroed 60-90 % of the shells there and the map agreement +#: rose in the bulk solvent as much as in the region of interest. +DEFAULT_SCHEME = "inverse_variance" +#: MTZ column carrying each scheme's weight (type ``W``); ``none`` writes no column. +WEIGHT_COLUMNS = {"inverse_variance": "W_IVW", "sigma_d": "W_SD"} + + +class DedWeightFallbackWarning(UserWarning): + """A requested weight scheme could not be evaluated and another was applied.""" + + +@dataclass(frozen=True) +class DedWeights: + """One scheme's weights and how they came about. + + Attributes + ---------- + scheme + The scheme requested. + applied + The scheme whose weights are in ``weights``; differs from ``scheme`` only after a + fallback, which ``diagnostics["fallback_reason"]`` then names. + weights + Per-reflection weights, shape ``(N,)``, mean one over finite positive entries. + diagnostics + Scheme-specific record: for ``sigma_d`` the fitted exponent, shrinkage sd, + clamp counters and the per-shell table. + """ + + scheme: str + applied: str + weights: torch.Tensor + diagnostics: dict = field(default_factory=dict) + + +def normalise_mean_one(w: torch.Tensor) -> torch.Tensor: + """Divide by the mean over all entries so the column averages one; unchanged when + that mean is not positive. + + Non-finite entries become zero, so a coefficient without a usable uncertainty drops + out of the map rather than poisoning it. Zero weights (shells without difference + power) stay zero and count in the mean, so the column mean is one whatever fraction + of the reflections carries weight. + """ + w = torch.where(torch.isfinite(w), w, torch.zeros_like(w)) + if w.numel() == 0 or not bool((w > 0).any()): + return w + return w / w.mean() + + +def reflection_geometry(hkl, cell, spacegroup, device, dtype): + """``(epsilon, d_star_sq)`` for ``hkl``: the reflection multiplicity from + ``spacegroup`` (ones when ``None``) and ``1/d**2`` in A^-2 from ``cell``, both on + ``device`` in ``dtype``.""" + from torchref.refinement.model_error_estimation.sigma_a import epsilon_from_hkl + + hkl_t = torch.as_tensor(hkl, device=device) + cell_t = cell if torch.is_tensor(cell) else cell.data + cell_t = torch.as_tensor(cell_t, device=device, dtype=dtype) + s = get_scattering_vectors(hkl_t, cell_t) + dss = (s * s).sum(dim=1).to(dtype) + eps = epsilon_from_hkl(hkl_t, spacegroup).to(device=device, dtype=dtype) + return eps, dss + + +def _inverse_variance(sigma_diff: torch.Tensor) -> torch.Tensor: + """``1 / sigma**2`` with sigma floored at a tenth of its median, so a reported zero + uncertainty gives a large finite weight rather than an infinite one.""" + finite = torch.isfinite(sigma_diff) & (sigma_diff >= 0) + positive = finite & (sigma_diff > 0) + if not bool(positive.any()): + return torch.zeros_like(sigma_diff) + floor = (sigma_diff[positive].median() * SIGMA_FLOOR_FRAC).clamp_min( + SIGMA_FLOOR_ABS + ) + sig = torch.where(finite, sigma_diff, torch.full_like(sigma_diff, float("inf"))) + return 1.0 / sig.clamp(min=floor) ** 2 + + +def compute_ded_weights( + scheme: str, + *, + delta_obs: torch.Tensor, + sigma_diff: torch.Tensor, + hkl: torch.Tensor, + cell, + spacegroup, + f_dark: torch.Tensor | None = None, + fit_mask: torch.Tensor | None = None, + sigma_d_config: "SigmaDConfig | None" = None, +) -> DedWeights: + """Per-reflection weights for one scheme. + + Parameters + ---------- + scheme : str + One of :data:`SCHEMES`. + delta_obs, sigma_diff : torch.Tensor + Signed observed differences and their propagated uncertainty, shape ``(N,)``, on + one common amplitude scale. + hkl : torch.Tensor + Miller indices, shape ``(N, 3)``. + cell : Cell or torch.Tensor + Unit cell, as a :class:`~torchref.symmetry.Cell` or its six parameters in A and + degrees. + spacegroup : SpaceGroup or None + For the reflection multiplicity; ``None`` means ones. + f_dark : torch.Tensor, optional + Dark amplitudes, shape ``(N,)``, for the sigma_D amplitude power law. + fit_mask : torch.Tensor, optional + Reflections entering the sigma_D fit; default every finite one. + sigma_d_config : SigmaDConfig, optional + Exponent and shrinkage settings for ``sigma_d``. + + Returns + ------- + DedWeights + Mean-one weights on ``delta_obs.device``. When ``sigma_d`` finds no difference + power in any shell, the inverse-variance weights are returned with + ``applied="inverse_variance"`` and a :class:`DedWeightFallbackWarning`. + """ + if scheme not in SCHEMES: + raise ValueError(f"Unknown DED weight scheme {scheme!r}; choose from {SCHEMES}") + sigma_diff = sigma_diff.reshape(-1).to(delta_obs.device, delta_obs.dtype) + if scheme == "none": + return DedWeights(scheme, scheme, torch.ones_like(sigma_diff)) + if scheme == "inverse_variance": + return DedWeights( + scheme, scheme, normalise_mean_one(_inverse_variance(sigma_diff)) + ) + + from torchref.refinement.model_error_estimation.sigma_d import ( + SigmaDConfig, + estimate_sigma_d, + sigma_d_per_reflection, + ) + + config = sigma_d_config if sigma_d_config is not None else SigmaDConfig() + d = delta_obs.reshape(-1) + eps, dss = reflection_geometry(hkl, cell, spacegroup, d.device, d.dtype) + f = f_dark.reshape(-1).to(d.device, d.dtype) if f_dark is not None else None + mask = ( + fit_mask.reshape(-1).to(d.device, torch.bool) + if fit_mask is not None + else torch.isfinite(d) & torch.isfinite(sigma_diff) + ) + shells = estimate_sigma_d( + d, sigma_diff, eps, dss, f, mask, gamma=config.gamma, shrink=config.shrink + ) + est = sigma_d_per_reflection(shells, dss, eps, f, sigma_diff) + diagnostics = { + "gamma": shells.gamma, + "gamma_fitted": shells.gamma_fitted, + "tau": shells.tau, + "curve_a": shells.curve_a, + "curve_b": shells.curve_b, + "degenerate": shells.degenerate, + "all_zero": shells.all_zero, + **shells.diagnostics, + "shells": { + "d_star_sq": shells.bin_dss.detach().cpu().tolist(), + "counts": shells.counts.detach().cpu().tolist(), + "B": shells.B.detach().cpu().tolist(), + "S2": shells.S2.detach().cpu().tolist(), + "Sigma_N_raw": shells.Sigma_N_raw.detach().cpu().tolist(), + "Sigma_N": shells.Sigma_N.detach().cpu().tolist(), + }, + "weight_sigma_d_raw": est.w.detach(), + } + if shells.all_zero or shells.degenerate: + reason = ( + "no difference power above the measurement variance in any shell " + f"(n_s2_clamped={shells.diagnostics['n_s2_clamped']})" + if shells.all_zero + else "fewer than two usable reflections" + ) + warnings.warn( + f"sigma_d weights: {reason}; applying inverse-variance weights instead", + DedWeightFallbackWarning, + stacklevel=2, + ) + diagnostics["fallback_reason"] = reason + return DedWeights( + scheme, + "inverse_variance", + normalise_mean_one(_inverse_variance(sigma_diff)), + diagnostics, + ) + return DedWeights(scheme, scheme, normalise_mean_one(est.w), diagnostics) + + +def all_ded_weights(**kwargs) -> dict[str, DedWeights]: + """Every registered scheme on the same inputs, keyed by scheme name. + + Takes the keyword arguments of :func:`compute_ded_weights` except ``scheme``. Used for + side-by-side reporting and for writing every weight column at once. + """ + return {scheme: compute_ded_weights(scheme, **kwargs) for scheme in SCHEMES} + + +__all__ = [ + "DEFAULT_SCHEME", + "SCHEMES", + "WEIGHT_COLUMNS", + "DedWeightFallbackWarning", + "DedWeights", + "all_ded_weights", + "compute_ded_weights", + "normalise_mean_one", + "reflection_geometry", +] diff --git a/torchref/maps/difference_map.py b/torchref/maps/difference_map.py index 69bcbeec..4dc010ea 100644 --- a/torchref/maps/difference_map.py +++ b/torchref/maps/difference_map.py @@ -65,8 +65,16 @@ class DifferenceMap(Map): (see :mod:`torchref.maps.map`). """ - def __init__(self, data, data_reference, model, gridsize=None, - device: Optional[torch.device] = None): + def __init__( + self, + data, + data_reference, + model, + gridsize=None, + device: Optional[torch.device] = None, + units: str = "normalized", + scale: Optional[torch.Tensor] = None, + ): # Pin all three inputs onto one device before constructing the # DatasetCollection / super().__init__ — both consume tensors # from data.hkl / model and would otherwise inherit whichever @@ -92,7 +100,11 @@ def __init__(self, data, data_reference, model, gridsize=None, gridsize=gridsize, map_type="Fcalc", # placeholder, calculate() is overridden device=resolved, + units=units, ) + # Per-reflection observed-to-model scale over the reference dataset's full + # reflection list; dividing by it puts the differences in electrons. + self.scale = scale def calculate(self) -> torch.Tensor: """Compute the isomorphous difference map. @@ -122,9 +134,10 @@ def calculate(self) -> torch.Tensor: ) # Map scaled amplitudes to P1 (amplitudes are invariant under symmetry) - fobs_ref_p1 = fobs_ref[orig_idx] - fobs_pert_p1 = fobs_pert[orig_idx] - delta_f_p1 = fobs_pert_p1 - fobs_ref_p1 + delta_f = fobs_pert - fobs_ref + if self.scale is not None: + delta_f = delta_f / self.scale.to(delta_f)[mask_combined] + delta_f_p1 = delta_f[orig_idx] # Compute Fcalc for P1 hkl (for phases) fcalc_p1 = self.model.get_structure_factor(hkl_p1) @@ -143,6 +156,8 @@ def calculate(self) -> torch.Tensor: grid = place_on_grid( hkl_p1, coefficients_p1, gridsize, enforce_hermitian=True ) - self._map = torch.fft.fftn(grid, dim=(0, 1, 2), norm="forward").real + self._map = self._to_units( + torch.fft.fftn(grid, dim=(0, 1, 2), norm="forward").real + ) return self._map diff --git a/torchref/maps/map.py b/torchref/maps/map.py index 85c75b06..204df60c 100644 --- a/torchref/maps/map.py +++ b/torchref/maps/map.py @@ -44,6 +44,11 @@ class Map(DeviceMixin): Default is ``"2Fo-Fc"``. Note ``"2Fo-Fc"`` is a *plain* 2Fo-Fc map (no figure-of-merit ``m`` and no sigma-A coefficient ``D``; i.e. ``m=1``, ``D=1``), not a likelihood-weighted 2mFo-DFc map. + units : str, optional + ``"normalized"`` (default) keeps the FFT's ``1/N`` normalisation; + ``"electrons"`` gives ``(1/V) sum_h F(h) exp(-2 pi i h.x)``, electrons per + cubic Angstrom, which is meaningful only when the coefficients are on the + absolute scale. Attributes ---------- @@ -72,6 +77,7 @@ class Map(DeviceMixin): """ VALID_MAP_TYPES = ("2Fo-Fc", "Fcalc") + VALID_UNITS = ("normalized", "electrons") def __init__( self, @@ -80,11 +86,15 @@ def __init__( gridsize: Optional[Tuple[int, int, int]] = None, map_type: str = "2Fo-Fc", device: Optional[torch.device] = None, + units: str = "normalized", ): if map_type not in self.VALID_MAP_TYPES: raise ValueError( f"map_type must be one of {self.VALID_MAP_TYPES}, got '{map_type}'" ) + if units not in self.VALID_UNITS: + raise ValueError(f"units must be one of {self.VALID_UNITS}, got '{units}'") + self.units = units self.device = resolve_device(data, model, device=device) self.data = data self.model = model @@ -166,9 +176,17 @@ def calculate(self) -> torch.Tensor: # FFT to real space: ρ(r) = (1/N) * sum_h F(h) * exp(-2πi h·r) # (norm="forward" applies the 1/N normalization, N = grid points) self._map = torch.fft.fftn(grid, dim=(0, 1, 2), norm="forward").real + self._map = self._to_units(self._map) return self._map + def _to_units(self, real_map: torch.Tensor) -> torch.Tensor: + """Rescale a ``1/N``-normalised FFT map to the configured units.""" + if self.units == "electrons": + volume = self.data.cell.volume.to(real_map.dtype) + return real_map * (real_map.numel() / volume) + return real_map + def write(self, filepath: str) -> int: """Write the map to a CCP4 file. diff --git a/torchref/refinement/model_error_estimation/_shells.py b/torchref/refinement/model_error_estimation/_shells.py new file mode 100644 index 00000000..3d6e9a7d --- /dev/null +++ b/torchref/refinement/model_error_estimation/_shells.py @@ -0,0 +1,187 @@ +"""Resolution-shell machinery shared by the model-error estimators. + +Equal-count shells over ``d*^2``, atomic-free segment sums, linear interpolation of +per-shell values back to reflections, and DerSimonian-Laird shrinkage of noisy per-shell +estimates toward a weighted straight line. :mod:`.sigma_a` and :mod:`.sigma_d` both +build on these; ``estimate_beta`` keeps its own module-level aliases so that its body +resolves the same globals it always did. + +Plain tensors in and out. Every result lives on the device of its inputs, and float +work happens in the dtype of the inputs, so callers control both by what they pass. +""" + +from functools import lru_cache + +import torch + + +@lru_cache(maxsize=8) +def segment_layout(lengths: tuple[int, ...], device_str: str): + """``(index, mask)`` placing contiguous segments on a padded ``(n_seg, max_len)`` grid. + + Cached: the sigma_A solve reduces ``n_grid * n_stages`` times over one layout. + ``lengths`` is a tuple so it can be a cache key. + """ + device = torch.device(device_str) + # dtype-ok: segment lengths for cumsum offsets/gather index; PyTorch requires int64 + L = torch.tensor(lengths, dtype=torch.long, device=device) + total = int(L.sum()) + max_len = int(L.max()) if L.numel() else 0 + # dtype-ok: zero offset concatenated into gather index; PyTorch requires int64 + zero = torch.zeros(1, dtype=torch.long, device=device) + starts = torch.cat([zero, L.cumsum(0)[:-1]]) + ar = torch.arange(max_len, device=device).reshape(1, max_len) + # Clamp keeps the gather in bounds for the padding slots; `mask` zeroes them anyway. + index = (starts.reshape(-1, 1) + ar).clamp(max=max(total - 1, 0)) + mask = ar < L.reshape(-1, 1) + return index, mask + + +def segsum(x: torch.Tensor, lengths: torch.Tensor) -> torch.Tensor: + """Sum ``x`` over contiguous segments, reducing along a padded trailing axis. + + Replaces ``torch.segment_reduce``, which is unimplemented on MPS. Keeps the properties + that op was chosen for: atomic-free, one fixed reduction order per segment, so the + result is bit-stable run to run and does not depend on ``scatter_add``'s CUDA atomicAdd + accumulation order (see ``tests/unit/refinement/test_estimate_beta_determinism.py``). + + Deliberately NOT ``cumsum[end] - cumsum[start]``, the usual contiguous-segment trick: + that recovers each shell sum by subtracting two running totals of the whole array, + reintroducing the large-minus-large these estimators are written to avoid. + + ``x`` reduces over its last axis, so a leading batch dimension is handled in one call. + Segments differ in length by at most one element, so the padding overhead is at most + ``n_seg`` slots. + """ + index, mask = segment_layout(tuple(int(v) for v in lengths), str(x.device)) + return (x[..., index] * mask.to(x.dtype)).sum(dim=-1) + + +def interp_in_dss( + dss_all: torch.Tensor, bin_dss: torch.Tensor, vals: torch.Tensor +) -> torch.Tensor: + """Linear interpolation of per-bin ``vals`` (at ``bin_dss``) to all reflections by + their ``d_star_sq``; clamp-to-edge outside the range.""" + n_bins = bin_dss.numel() + if n_bins == 1: + return torch.full_like(dss_all, float(vals[0])) + idx = torch.searchsorted(bin_dss, dss_all).clamp(1, n_bins - 1) + x0 = bin_dss[idx - 1] + x1 = bin_dss[idx] + wlin = ((dss_all - x0) / (x1 - x0).clamp(min=1e-30)).clamp(0.0, 1.0) + return (1 - wlin) * vals[idx - 1] + wlin * vals[idx] + + +def equal_count_shells( + dss: torch.Tensor, *, per_bin: int, min_bins: int, min_per_bin: int +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int]: + """Equal-count resolution shells over ``dss``, sorted ascending. + + Parameters + ---------- + dss : torch.Tensor + ``d*^2`` of the reflections entering the fit, shape ``(n,)``, in A^-2. + per_bin : int + Target reflections per shell. + min_bins, min_per_bin : int + Floor on the shell count for sparse sets: at least ``min_bins`` shells as long + as each still holds ``min_per_bin`` reflections. + + Returns + ------- + tuple + ``(order, seg, seg_lengths, n_bins)``. ``order`` sorts ``dss`` ascending (a stable + sort, so tied values bin identically on every backend); ``seg`` is the shell index + of each sorted reflection, a non-decreasing ramp; ``seg_lengths`` the count per + shell. + """ + n = int(dss.numel()) + order = torch.argsort(dss, stable=True) + n_by_count = max(1, n // per_bin) + n_cap = max(1, n // min_per_bin) + n_bins = max(n_by_count, min(min_bins, n_cap)) + seg = ( + torch.arange(n, device=dss.device) * n_bins + ) // n # dtype-ok: bincount input; PyTorch requires int64 + seg_lengths = torch.bincount(seg, minlength=n_bins) + return order, seg, seg_lengths, n_bins + + +def dl_shrink_to_line( + y: torch.Tensor, + var: torch.Tensor, + x: torch.Tensor, + *, + slope_min: float | None = None, + slope_max: float | None = None, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, float, float]: + """Shrink noisy per-shell values toward a weighted straight line in ``x``. + + DerSimonian-Laird shrinkage toward a two-parameter line fitted across all shells, + which the two parameters determine far better than any one shell is determined:: + + fit y = a + b*x, weights 1/var + tau^2 = DL between-shell variance about the line, weights 1/var + w_i = var_i / (var_i + tau^2) + y_i <- (1 - w_i)*y_i + w_i*line_i + + ``tau^2`` is the size of the dataset-specific residual the line does not capture, so + ``w_i -> 0`` where that residual is real and large and ``w_i -> 1`` where the shell is + badly determined. One shot, no iteration: the target is a fixed line. Weights are + ``1/var``, never counts, because count weighting lets high-``var`` shells dominate + ``Q`` and veto shrinkage entirely. + + Parameters + ---------- + y, var, x : torch.Tensor + Per-shell value, its sampling variance and the abscissa (``d*^2``), shape + ``(k,)``. A shell with non-finite ``y`` or ``var`` (or ``var <= 0``) takes no part + in the fit and is replaced by the line outright (``w = 1``). + slope_min, slope_max : float, optional + Clamp on the fitted slope ``b``; ``a`` is refitted after the clamp so the line + still passes through the weighted centroid. + + Returns + ------- + tuple + ``(y_shrunk, w, tau_sq, a, b)``. With fewer than four usable shells, or when the + slope is unidentifiable (all shells at one ``x``), the input is returned + unchanged with ``w = 0``, ``tau_sq = 0`` and NaN line coefficients. + """ + nan = float("nan") + usable = torch.isfinite(y) & torch.isfinite(var) & (var > 0) + k = int(usable.sum()) + # Two fitted parameters need at least two residual degrees of freedom. + if k < 4: + return y, torch.zeros_like(y), y.new_zeros(()), nan, nan + + wt = torch.where(usable, 1.0 / var.clamp(min=1e-30), torch.zeros_like(var)) + yz = torch.where(usable, y, torch.zeros_like(y)) + S = wt.sum() + Sx = (wt * x).sum() + Sxx = (wt * x * x).sum() + Sy = (wt * yz).sum() + Sxy = (wt * x * yz).sum() + det = S * Sxx - Sx * Sx + # Relative, not absolute: `det` is a difference of two ~`S**2 * x**2` terms, so on a + # degenerate input it lands at the cancellation floor, not near zero. + if float(det.abs()) <= 1e-12 * float((S * Sxx).abs()): + return y, torch.zeros_like(y), y.new_zeros(()), nan, nan + b = (S * Sxy - Sx * Sy) / det + if slope_min is not None: + b = b.clamp(min=slope_min) + if slope_max is not None: + b = b.clamp(max=slope_max) + a = (wt * (yz - b * x)).sum() / S.clamp(min=1e-30) + line = a + b * x + + resid = torch.where(usable, yz - line, torch.zeros_like(y)) + Q = (wt * resid * resid).sum() + dof = float(k - 2) # two parameters were fitted + c = (S - (wt * wt).sum() / S.clamp(min=1e-30)).clamp(min=1e-30) + # Q < k-2 means the scatter about the line is SMALLER than the noise alone predicts, + # i.e. no evidence of structure the line is missing -> tau^2 = 0 -> take the line. + tau_sq = ((Q - dof) / c).clamp(min=0.0) + w = torch.where(usable, var / (var + tau_sq).clamp(min=1e-30), torch.ones_like(var)) + out = (1.0 - w) * torch.where(usable, y, line) + w * line + return out, w, tau_sq, float(a), float(b) diff --git a/torchref/refinement/model_error_estimation/sigma_a.py b/torchref/refinement/model_error_estimation/sigma_a.py index a24f9285..dab6d53a 100644 --- a/torchref/refinement/model_error_estimation/sigma_a.py +++ b/torchref/refinement/model_error_estimation/sigma_a.py @@ -17,13 +17,15 @@ import math from dataclasses import dataclass -from functools import lru_cache -from typing import Optional, Tuple +from typing import Optional import torch from torchref.config import get_float_dtype +from ._shells import interp_in_dss as _interp_in_dss +from ._shells import segment_layout as _segment_layout # noqa: F401 +from ._shells import segsum as _segsum def epsilon_from_hkl(hkl: torch.Tensor, spacegroup) -> torch.Tensor: """Per-reflection epsilon, tolerating a missing space group. @@ -255,47 +257,6 @@ def _rice_nll_reduced( return torch.where(centric, cen, acen) -@lru_cache(maxsize=8) -def _segment_layout(lengths: Tuple[int, ...], device_str: str): - """``(index, mask)`` placing contiguous segments on a padded ``(n_seg, max_len)`` grid. - - Cached: ``_solve_sigma_a`` reduces ``n_grid * n_stages`` times over one layout. - ``lengths`` is a tuple so it can be a cache key. - """ - device = torch.device(device_str) - L = torch.tensor(lengths, dtype=torch.long, device=device) # dtype-ok: segment lengths for cumsum offsets/gather index; PyTorch requires int64 - total = int(L.sum()) - max_len = int(L.max()) if L.numel() else 0 - zero = torch.zeros(1, dtype=torch.long, device=device) # dtype-ok: zero offset concatenated into gather index; PyTorch requires int64 - starts = torch.cat([zero, L.cumsum(0)[:-1]]) - ar = torch.arange(max_len, device=device).reshape(1, max_len) - # Clamp keeps the gather in bounds for the padding slots; `mask` zeroes them anyway. - index = (starts.reshape(-1, 1) + ar).clamp(max=max(total - 1, 0)) - mask = ar < L.reshape(-1, 1) - return index, mask - - -def _segsum(x: torch.Tensor, lengths: torch.Tensor) -> torch.Tensor: - """Sum ``x`` over contiguous segments, reducing along a padded trailing axis. - - Replaces ``torch.segment_reduce``, which is unimplemented on MPS. Keeps the properties - that op was chosen for: atomic-free, one fixed reduction order per segment, so the - result is bit-stable run to run and does not depend on ``scatter_add``'s CUDA atomicAdd - accumulation order (the original GPU non-determinism bug -- see - ``tests/unit/refinement/test_estimate_beta_determinism.py``). - - Deliberately NOT ``cumsum[end] - cumsum[start]``, the usual contiguous-segment trick: - that recovers each shell sum by subtracting two running totals of the whole array, - reintroducing the large-minus-large this module is written to avoid. - - ``x`` reduces over its last axis, so a leading batch dimension (the grid candidates) - is handled in one call. Segments here differ in length by at most one element, so the - padding overhead is at most ``n_seg`` slots. - """ - index, mask = _segment_layout(tuple(int(v) for v in lengths), str(x.device)) - return (x[..., index] * mask.to(x.dtype)).sum(dim=-1) - - def _grid_ladder(n: int, ratio: float, device, dtype) -> torch.Tensor: """``ratio ** (k - (n-1)/2)`` for ``k`` in ``[0, n)``, built from Python floats. @@ -735,19 +696,6 @@ def segsum(x): ) -def _interp_in_dss(dss_all, bin_dss, vals): - """Linear interpolation of per-bin ``vals`` (at ``bin_dss``) to all - reflections by their ``d_star_sq``; clamp-to-edge outside the range.""" - n_bins = bin_dss.numel() - if n_bins == 1: - return torch.full_like(dss_all, float(vals[0])) - idx = torch.searchsorted(bin_dss, dss_all).clamp(1, n_bins - 1) - x0 = bin_dss[idx - 1] - x1 = bin_dss[idx] - wlin = ((dss_all - x0) / (x1 - x0).clamp(min=1e-30)).clamp(0.0, 1.0) - return (1 - wlin) * vals[idx - 1] + wlin * vals[idx] - - # ===================================================================== # Stateful estimator (owned by the consuming target, not the scaler) # ===================================================================== diff --git a/torchref/refinement/model_error_estimation/sigma_d.py b/torchref/refinement/model_error_estimation/sigma_d.py new file mode 100644 index 00000000..8ead3617 --- /dev/null +++ b/torchref/refinement/model_error_estimation/sigma_d.py @@ -0,0 +1,795 @@ +"""Difference-driven error estimation: the expected difference power ``sigma_D``. + +A light-minus-dark difference coefficient ``dF_obs = dF_true + noise`` carries a true +signal whose power ``S = E[dF_true**2]`` varies with resolution and with the dark +amplitude, and a measurement noise ``sigma_diff**2`` the merge reports. The best linear +estimate of ``dF_true`` from ``dF_obs`` is ``w * dF_obs`` with the Wiener weight +``w = S / (S + sigma_diff**2)``, so a difference map needs ``S`` per reflection. Inverse +variance alone weights by precision and treats every reflection as carrying the same +expected difference, which suppresses the strong reflections whose difference power is +ten to seventy times that of weak ones. + +``S`` needs no half datasets. Per resolution shell the second moment of the observed +differences is ``B = S + S2`` with ``S2`` the mean measurement variance, so +``Sigma_N = B - S2`` is the expected true difference power, the same identity +:mod:`.sigma_a` uses for amplitudes. Within a shell the power follows the dark amplitude +as ``(F_dark / )**gamma`` with one fitted exponent, carried by the per-reflection +multiplier ``epsilon`` exactly as the reflection multiplicity is. With a difference model +``dF_calc`` the shell moments also give the Gaussian coupling ``alpha = / +`` and the unexplained power ``beta_model = Sigma_N - alpha**2 Sigma_P``, the +extra variance a difference likelihood adds to ``sigma_diff**2``. + +Differences are signed and small, so the statistics are Gaussian throughout: there is no +Rice branch and no centric distinction. Plain tensors in and out, no ``ReflectionData`` +or ``Scaler`` coupling, so :mod:`torchref.maps` and :mod:`torchref.cli` can import this +module without closing an import cycle. Every result lives on the device of its inputs. +""" + +import math +from dataclasses import dataclass + +import torch + +from torchref.config import get_float_dtype + +from ._shells import equal_count_shells, interp_in_dss, segsum +from .sigma_a import SHRINK_ENABLED + +# --- sigma_D estimator constants ------------------------------------------------- +#: Shell construction, matching the ``estimate_beta`` defaults so a sigma_A and a sigma_D +#: fit on the same reflections use the same shells. +PER_BIN = 140 +MIN_BINS = 5 +MIN_PER_BIN = 40 +#: Exponent of the dark-amplitude power law when it cannot be fitted. The fitted value on +#: the small-molecule, TD1 and bacteriorhodopsin A/B campaigns was 0.8-1.2, so 1.0 (power +#: proportional to the amplitude) is the informed default; 0.0 would be the pure shell model. +GAMMA_DEFAULT = 1.0 +#: Bounds on the fitted exponent. Outside [0, 2] the class means are dominated by one +#: amplitude decile and the regression is on noise; 2 is proportionality to the intensity. +GAMMA_BOUNDS = (0.0, 2.0) +#: Dark-amplitude quantile classes per shell for the exponent regression. Four keeps at +#: least 35 reflections per class at the default shell size; more classes did not move +#: the fitted exponent on the campaign data. +N_F_CLASSES = 4 +#: Minimum number of usable (shell, class) cells before the exponent is fitted at all. +GAMMA_MIN_CLASSES = 8 +#: Minimum reflections in a class for its moment to enter the exponent regression. +MIN_PER_CLASS = 5 +#: Largest standard error at which a fitted exponent is used. Above it the class +#: moments are noise (a null dataset gives se ~ 1 or more), and the default is safer than +#: a random exponent that would redistribute weight between amplitude classes. +GAMMA_SE_MAX = 0.5 +#: Gauss-Newton iterations for the decaying-curve fit; the problem is two-parameter and +#: well conditioned, so this is far more than it needs. +CURVE_ITERS = 60 +#: Floor on ``F_dark`` relative to its shell mean before the power law is evaluated, so a +#: zero or near-zero dark amplitude cannot delete the expected difference power. +F_FLOOR_FRAC = 0.05 +#: Positive floor used where a logarithm of a clamped-to-zero power is needed. +_TINY = 1e-30 + + +@dataclass(frozen=True) +class SigmaDConfig: + """The estimator's knobs, as one value. + + ``gamma=None`` fits the dark-amplitude exponent; a float fixes it. ``shrink=None`` + means the module default shared with sigma_A, normalised here so consumers never + handle ``None``. Frozen, so two consumers sharing a config cannot drift apart. + """ + + gamma: float | None = None + shrink: bool | None = None + + def __post_init__(self): + if self.gamma is not None: + g = float(self.gamma) + lo, hi = GAMMA_BOUNDS + if not (lo <= g <= hi): + raise ValueError(f"gamma must lie in {GAMMA_BOUNDS}, got {g}") + object.__setattr__(self, "gamma", g) + object.__setattr__( + self, "shrink", bool(SHRINK_ENABLED if self.shrink is None else self.shrink) + ) + + +@dataclass(frozen=True) +class SigmaDShells: + """Per-shell output of :func:`estimate_sigma_d`. + + All power quantities are ``epsilon``-reduced and in ``F**2`` units of the input. + + Attributes + ---------- + B, S2 + Raw second moment of the observed differences and mean measurement variance. + Sigma_N_raw, Sigma_N + Expected true difference power ``(B - S2)`` clamped at zero, before and after the + shrinkage of the signed value toward a decaying curve ``exp(a + b d*^2)`` fitted + to every shell. ``Sigma_N`` is what the weights use. + Sigma_P, C, alpha, beta_model + Model power, cross moment, Gaussian coupling ``C / Sigma_P`` and unexplained power + ``(Sigma_N - alpha**2 Sigma_P)`` clamped at zero. Without a model ``Sigma_P`` and + ``C`` are zero, ``alpha`` one and ``beta_model == Sigma_N``. + counts, bin_dss + Reflections per shell and its mean ``d*^2`` in A^-2, the interpolation abscissa. + bin_log_fbar, bin_log_z + Log of the shell-mean dark amplitude and of the shell mean of + ``(F / Fbar)**gamma``, so the per-reflection multiplier + ``(F / Fbar)**gamma / Z`` has shell mean one and ``Sigma_N`` stays the shell mean + of the per-reflection power. Zero when no dark amplitude was supplied. + shrink_w, tau, curve_a, curve_b + Shrinkage weight per shell, the between-shell sd about the fitted curve and its + coefficients ``exp(a + b d*^2)`` (NaN when no curve was fitted; the curve is then + zero everywhere). + gamma, gamma_fitted + The exponent used and whether it was fitted rather than fixed or defaulted. + has_model, degenerate, all_zero + Whether ``dF_calc`` was supplied, whether fewer than two usable reflections + existed, and whether every shell's ``Sigma_N`` is zero (weights would vanish). + diagnostics + Counters: ``n_dropped, n_fit, n_shell, n_s2_clamped, n_beta_clamped, + n_f_floored, n_class_dropped, n_class_used, gamma_se, gamma_at_bound, + gamma_reason``. + """ + + B: torch.Tensor + S2: torch.Tensor + Sigma_N_raw: torch.Tensor + Sigma_N: torch.Tensor + Sigma_P: torch.Tensor + C: torch.Tensor + alpha: torch.Tensor + beta_model: torch.Tensor + counts: torch.Tensor + bin_dss: torch.Tensor + bin_log_fbar: torch.Tensor + bin_log_z: torch.Tensor + shrink_w: torch.Tensor + tau: float + curve_a: float + curve_b: float + gamma: float + gamma_fitted: bool + has_model: bool + degenerate: bool + all_zero: bool + diagnostics: dict + + +@dataclass(frozen=True) +class SigmaDEstimate: + """Everything a consumer needs from one estimate, per reflection and detached. + + Attributes + ---------- + S + Expected true difference power ``epsilon * Sigma_N(d*^2) * g(F_dark)``. + sigma_sq + The measurement variance the weight was formed with (``sigma_diff**2``). + w + Wiener weight ``S / (S + sigma_sq)`` in ``[0, 1)``, not normalised. + alpha, beta_model + Coupling and unexplained power, interpolated per shell; ``beta_model`` carries + the same ``epsilon * g`` multiplier as ``S``. + epsilon + The multiplicity actually applied. + shells + The :class:`SigmaDShells` this was interpolated from. + """ + + S: torch.Tensor + sigma_sq: torch.Tensor + w: torch.Tensor + alpha: torch.Tensor + beta_model: torch.Tensor + epsilon: torch.Tensor + shells: SigmaDShells + + +def _working_dtype(t: torch.Tensor) -> torch.dtype: + dtype = torch.promote_types(get_float_dtype(), t.dtype) + # dtype-ok: MPS capability guard, not an allocation + if dtype == torch.float64 and t.device.type == "mps": + raise RuntimeError( + "MPS has no float64; set the defaults float dtype to float32 or use CPU" + ) + return dtype + + +def _degenerate( + delta_obs: torch.Tensor, gamma: float, has_model: bool, diagnostics: dict, out_dtype +) -> SigmaDShells: + """One conservative shell: the mean squared difference as the power, alpha one.""" + ok = torch.isfinite(delta_obs) + b = (delta_obs[ok] ** 2).mean() if bool(ok.any()) else delta_obs.new_ones(()) + one = torch.ones(1, device=delta_obs.device, dtype=out_dtype) + zero = torch.zeros(1, device=delta_obs.device, dtype=out_dtype) + b1 = (one * b).to(out_dtype) + return SigmaDShells( + B=b1, + S2=zero, + Sigma_N_raw=b1, + Sigma_N=b1, + Sigma_P=zero, + C=zero, + alpha=one, + beta_model=b1, + counts=zero, + bin_dss=zero, + bin_log_fbar=zero, + bin_log_z=zero, + shrink_w=zero, + tau=0.0, + curve_a=float("nan"), + curve_b=float("nan"), + gamma=gamma, + gamma_fitted=False, + has_model=has_model, + degenerate=True, + all_zero=False, + diagnostics=diagnostics, + ) + + +def _fit_gamma( + d2e: torch.Tensor, + s2e: torch.Tensor, + log_ratio: torch.Tensor, + seg: torch.Tensor, + n_bins: int, +) -> tuple[float, float, bool, int, int, str]: + """Fit the dark-amplitude exponent from within-shell amplitude classes. + + Each shell is split into ``N_F_CLASSES`` quantile classes of the dark amplitude. A + class contributes ``log( - )`` against its mean log amplitude + ratio when that difference power is positive. One slope is fitted across all shells + with the shell means removed (fixed effects), weighted by ``n_c / 2``: the log of a + mean of ``n`` squared Gaussians has variance ``2 / n``. + + Returns ``(gamma, gamma_se, at_bound, n_used, n_dropped, reason)``; ``reason`` is + ``"fitted"`` or names why the default was taken. + """ + xs, ys, ws, shell_id = [], [], [], [] + n_dropped = 0 + for k in range(n_bins): + in_shell = torch.nonzero(seg == k, as_tuple=True)[0] + n_k = int(in_shell.numel()) + if n_k < N_F_CLASSES * MIN_PER_CLASS: + n_dropped += N_F_CLASSES + continue + order = torch.argsort(log_ratio[in_shell], stable=True) + idx = in_shell[order] + cls = ( + torch.arange(n_k, device=seg.device) * N_F_CLASSES + ) // n_k # dtype-ok: bincount input; PyTorch requires int64 + lengths = torch.bincount(cls, minlength=N_F_CLASSES).to(d2e.dtype) + m = (segsum(d2e[idx], lengths) - segsum(s2e[idx], lengths)) / lengths + xc = segsum(log_ratio[idx], lengths) / lengths + keep = (m > 0) & (lengths >= MIN_PER_CLASS) + n_dropped += int((~keep).sum()) + if int(keep.sum()) < 2: + continue + xs.append(xc[keep]) + ys.append(torch.log(m[keep])) + ws.append(lengths[keep] / 2.0) + shell_id.append(torch.full_like(xc[keep], float(k))) + if not xs: + return GAMMA_DEFAULT, float("nan"), False, 0, n_dropped, "too_few_classes" + x = torch.cat(xs) + y = torch.cat(ys) + w = torch.cat(ws) + sid = torch.cat(shell_id) + n_used = int(x.numel()) + if n_used < GAMMA_MIN_CLASSES: + return GAMMA_DEFAULT, float("nan"), False, n_used, n_dropped, "too_few_classes" + # Remove each shell's weighted mean from x and y: the slope is then estimated from + # within-shell contrasts only, so shell-to-shell differences in power cannot leak in. + xc = x.clone() + yc = y.clone() + n_shells_used = 0 + for k in torch.unique(sid): + s = sid == k + n_shells_used += 1 + wk = w[s] + xc[s] = x[s] - (wk * x[s]).sum() / wk.sum() + yc[s] = y[s] - (wk * y[s]).sum() / wk.sum() + sxx = (w * xc * xc).sum() + if float(sxx) <= 0.0: + return ( + GAMMA_DEFAULT, + float("nan"), + False, + n_used, + n_dropped, + "no_amplitude_spread", + ) + gamma = float((w * xc * yc).sum() / sxx) + dof = n_used - n_shells_used - 1 + if dof > 0: + resid = yc - gamma * xc + s2 = float((w * resid * resid).sum() / dof) + gamma_se = math.sqrt(max(s2, 0.0) / float(sxx)) + else: + gamma_se = float("nan") + if not math.isfinite(gamma_se) or gamma_se > GAMMA_SE_MAX: + return GAMMA_DEFAULT, gamma_se, False, n_used, n_dropped, "too_uncertain" + lo, hi = GAMMA_BOUNDS + clamped = min(max(gamma, lo), hi) + return clamped, gamma_se, clamped != gamma, n_used, n_dropped, "fitted" + + +def _fit_decay(y: torch.Tensor, var: torch.Tensor, x: torch.Tensor): + """Weighted fit of ``exp(a + b x)``, ``b <= 0``, to signed per-shell power. + + Works on the signed ``B - S2`` of every shell, so shells whose power is zero or + negative by sampling noise pull the curve down instead of being ignored: on a null + dataset the curve goes to zero rather than to the winner's curse of the positive + shells. Gauss-Newton with step halving on the weighted least squares; the fit is + two-parameter and well conditioned. + + Returns ``(curve, a, b)``; ``curve`` is zeros with NaN coefficients when fewer than + four shells are usable or the weighted mean power is not positive. + """ + nan = float("nan") + usable = torch.isfinite(y) & torch.isfinite(var) & (var > 0) + if int(usable.sum()) < 4: + return torch.zeros_like(y), nan, nan + w = torch.where(usable, 1.0 / var.clamp(min=_TINY), torch.zeros_like(var)) + yz = torch.where(usable, y, torch.zeros_like(y)) + mean = float((w * yz).sum() / w.sum()) + if mean <= 0.0: + return torch.zeros_like(y), nan, nan + a = torch.tensor(math.log(mean), dtype=y.dtype, device=y.device) + b = torch.zeros((), dtype=y.dtype, device=y.device) + + def loss(a_, b_): + r = yz - torch.exp(a_ + b_ * x) + return float((w * r * r).sum()) + + current = loss(a, b) + for _ in range(CURVE_ITERS): + f = torch.exp(a + b * x) + r = yz - f + # Jacobian of f with respect to (a, b): f and f*x. + j_a, j_b = f, f * x + g = torch.stack([(w * r * j_a).sum(), (w * r * j_b).sum()]) + h = torch.stack( + [ + torch.stack([(w * j_a * j_a).sum(), (w * j_a * j_b).sum()]), + torch.stack([(w * j_a * j_b).sum(), (w * j_b * j_b).sum()]), + ] + ) + h = ( + h + + 1e-12 * torch.eye(2, dtype=h.dtype, device=h.device) * h.diagonal().max() + ) + step = torch.linalg.solve(h, g) + scale = 1.0 + improved = False + for _ in range(12): + a_new = a + scale * step[0] + b_new = (b + scale * step[1]).clamp(max=0.0) + new = loss(a_new, b_new) + if new < current: + a, b, current, improved = a_new, b_new, new, True + break + scale *= 0.5 + if not improved or float(step.abs().max()) < 1e-9: + break + return torch.exp(a + b * x), float(a), float(b) + + +def estimate_sigma_d( + delta_obs: torch.Tensor, + sigma_diff: torch.Tensor, + epsilon: torch.Tensor | None, + d_star_sq: torch.Tensor, + f_dark: torch.Tensor | None, + fit_mask: torch.Tensor, + *, + delta_calc: torch.Tensor | None = None, + gamma: float | None = None, + shrink: bool | None = None, + per_bin: int = PER_BIN, + min_bins: int = MIN_BINS, + min_per_bin: int = MIN_PER_BIN, +) -> SigmaDShells: + """Per-shell expected difference power, with the dark-amplitude exponent and, + given a difference model, its coupling and unexplained power. + + Runs under ``torch.no_grad()``. The working dtype is the wider of the configured + float dtype and ``delta_obs.dtype``; results are cast back to ``delta_obs.dtype``. + + Parameters + ---------- + delta_obs : torch.Tensor + Signed observed differences ``F_light - F_dark``, shape ``(N,)``, on one common + amplitude scale. + sigma_diff : torch.Tensor + Propagated uncertainty of ``delta_obs``, shape ``(N,)``, same units. + epsilon : torch.Tensor or None + Reflection multiplicity, shape ``(N,)``; ``None`` means ones. + d_star_sq : torch.Tensor + ``1/d**2`` per reflection, shape ``(N,)``, in A^-2. + f_dark : torch.Tensor or None + Dark amplitude for the power law, shape ``(N,)``. ``None`` disables the amplitude + dependence (``gamma`` reported as the default with reason ``"no_f_dark"``). + fit_mask : torch.Tensor + Boolean ``(N,)``: which reflections enter the fit. + delta_calc : torch.Tensor, optional + Model differences ``|F_calc_light| - |F_calc_dark|``, shape ``(N,)``, on the + observed scale. Enables ``alpha`` and ``beta_model``. + gamma : float, optional + Fix the exponent instead of fitting it. + shrink : bool, optional + Shrink the signed shell power toward a decaying curve in ``d*^2``; default the + module setting. + per_bin, min_bins, min_per_bin : int, optional + Shell construction, see :func:`~._shells.equal_count_shells`. + + Returns + ------- + SigmaDShells + One frozen record of per-shell quantities plus counters. + """ + device = delta_obs.device + out_dtype = delta_obs.dtype + dtype = _working_dtype(delta_obs) + shrink = bool(SHRINK_ENABLED if shrink is None else shrink) + if gamma is not None: + lo, hi = GAMMA_BOUNDS + if not (lo <= float(gamma) <= hi): + raise ValueError(f"gamma must lie in {GAMMA_BOUNDS}, got {gamma}") + + with torch.no_grad(): + d_all = delta_obs.reshape(-1).to(dtype) + s_all = sigma_diff.reshape(-1).to(dtype) + x_all = d_star_sq.reshape(-1).to(dtype) + e_all = ( + epsilon.reshape(-1).to(dtype) + if epsilon is not None + else torch.ones_like(d_all) + ) + has_f = f_dark is not None + f_all = f_dark.reshape(-1).to(dtype) if has_f else None + has_model = delta_calc is not None + c_all = delta_calc.reshape(-1).to(dtype) if has_model else None + + finite = ( + torch.isfinite(d_all) + & torch.isfinite(s_all) + & torch.isfinite(x_all) + & torch.isfinite(e_all) + & (s_all >= 0.0) + & (e_all > 0.0) + ) + if has_f: + finite &= torch.isfinite(f_all) + if has_model: + finite &= torch.isfinite(c_all) + fit = fit_mask.reshape(-1).to(torch.bool) + usable = fit & finite + n_dropped = int((fit & ~finite).sum()) + idx = torch.nonzero(usable, as_tuple=True)[0] + n_fit = int(idx.numel()) + + diagnostics = { + "n_dropped": n_dropped, + "n_fit": n_fit, + "n_shell": 0, + "n_s2_clamped": 0, + "n_beta_clamped": 0, + "n_f_floored": 0, + "n_class_dropped": 0, + "n_class_used": 0, + "gamma_se": float("nan"), + "gamma_at_bound": False, + "gamma_reason": "degenerate", + } + if n_fit < 2: + g = float(gamma) if gamma is not None else GAMMA_DEFAULT + return _degenerate(d_all, g, has_model, diagnostics, out_dtype) + + order, seg, seg_lengths, n_bins = equal_count_shells( + x_all[idx], per_bin=per_bin, min_bins=min_bins, min_per_bin=min_per_bin + ) + sel = idx[order] + d, s, e, x = d_all[sel], s_all[sel], e_all[sel], x_all[sel] + counts = seg_lengths.to(dtype) + d2e = d * d / e + s2e = s * s / e + + B = segsum(d2e, seg_lengths) / counts + S2 = segsum(s2e, seg_lengths) / counts + Sigma_N_raw = (B - S2).clamp(min=0.0) + n_s2_clamped = int((S2 >= B).sum()) + bin_dss = segsum(x, seg_lengths) / counts + + # --- dark-amplitude power law ------------------------------------------- + if has_f: + f = f_all[sel] + fbar = (segsum(f, seg_lengths) / counts).clamp(min=_TINY) + fbar_h = fbar[seg] + floor = F_FLOOR_FRAC * fbar_h + n_f_floored = int((f < floor).sum()) + f_fl = torch.maximum(f, floor) + log_ratio = torch.log(f_fl) - torch.log(fbar_h) + bin_log_fbar = torch.log(fbar) + else: + n_f_floored = 0 + log_ratio = torch.zeros_like(d) + bin_log_fbar = torch.zeros_like(bin_dss) + + if gamma is not None: + g_used, g_se, at_bound, n_used, n_cls_dropped, reason = ( + float(gamma), + float("nan"), + False, + 0, + 0, + "fixed", + ) + fitted = False + elif not has_f: + g_used, g_se, at_bound, n_used, n_cls_dropped, reason = ( + GAMMA_DEFAULT, + float("nan"), + False, + 0, + 0, + "no_f_dark", + ) + fitted = False + else: + g_used, g_se, at_bound, n_used, n_cls_dropped, reason = _fit_gamma( + d2e, s2e, log_ratio, seg, n_bins + ) + fitted = reason == "fitted" + + if has_f: + g_raw = torch.exp(g_used * log_ratio) + Z = (segsum(g_raw, seg_lengths) / counts).clamp(min=_TINY) + bin_log_z = torch.log(Z) + else: + bin_log_z = torch.zeros_like(bin_dss) + + # --- difference model ------------------------------------------------------ + if has_model: + c = c_all[sel] + Sigma_P = segsum(c * c / e, seg_lengths) / counts + C = segsum(d * c / e, seg_lengths) / counts + alpha = C / Sigma_P.clamp(min=_TINY) + else: + Sigma_P = torch.zeros_like(B) + C = torch.zeros_like(B) + alpha = torch.ones_like(B) + + # --- stability shrinkage of the signed power toward a decaying curve --------- + # The signed B - S2 keeps every shell as evidence: a shell below zero by noise + # says the power there is small, and it must count. var(B) is 2 B**2 / n for + # Gaussian differences, so that is the sampling variance of each shell's value. + signed = B - S2 + var_s = 2.0 * B * B / counts + if shrink: + curve, curve_a, curve_b = _fit_decay(signed, var_s, bin_dss) + resid = signed - curve + prec = 1.0 / var_s.clamp(min=_TINY) + Q = (prec * resid * resid).sum() + dof = float(max(int(signed.numel()) - 2, 1)) + c = (prec.sum() - (prec * prec).sum() / prec.sum().clamp(min=_TINY)).clamp( + min=_TINY + ) + # Q < dof means the shells scatter no more than their noise: take the curve. + tau_sq = ((Q - dof) / c).clamp(min=0.0) + shrink_w = var_s / (var_s + tau_sq).clamp(min=_TINY) + Sigma_N = ((1.0 - shrink_w) * signed + shrink_w * curve).clamp(min=0.0) + else: + Sigma_N = Sigma_N_raw + shrink_w, tau_sq = torch.zeros_like(B), B.new_zeros(()) + curve_a = curve_b = float("nan") + Sigma_N = torch.where(torch.isfinite(Sigma_N), Sigma_N, torch.zeros_like(B)) + + beta_model = (Sigma_N - alpha * alpha * Sigma_P).clamp(min=0.0) + n_beta_clamped = int(((Sigma_N - alpha * alpha * Sigma_P) < 0.0).sum()) + all_zero = bool((Sigma_N <= 0.0).all()) + + diagnostics.update( + n_shell=int(n_bins), + n_s2_clamped=n_s2_clamped, + n_beta_clamped=n_beta_clamped, + n_f_floored=n_f_floored, + n_class_dropped=n_cls_dropped, + n_class_used=n_used, + gamma_se=g_se, + gamma_at_bound=at_bound, + gamma_reason=reason, + ) + + to = lambda t: t.to(out_dtype) + return SigmaDShells( + B=to(B), + S2=to(S2), + Sigma_N_raw=to(Sigma_N_raw), + Sigma_N=to(Sigma_N), + Sigma_P=to(Sigma_P), + C=to(C), + alpha=to(alpha), + beta_model=to(beta_model), + counts=to(counts), + bin_dss=to(bin_dss), + bin_log_fbar=to(bin_log_fbar), + bin_log_z=to(bin_log_z), + shrink_w=to(shrink_w), + tau=float(tau_sq.clamp(min=0.0).sqrt()), + curve_a=curve_a, + curve_b=curve_b, + gamma=float(g_used), + gamma_fitted=fitted, + has_model=has_model, + degenerate=False, + all_zero=all_zero, + diagnostics=diagnostics, + ) + + +def sigma_d_per_reflection( + shells: SigmaDShells, + d_star_sq: torch.Tensor, + epsilon: torch.Tensor | None, + f_dark: torch.Tensor | None, + sigma_diff: torch.Tensor, +) -> SigmaDEstimate: + """Interpolate a shell estimate onto reflections and form the Wiener weights. + + Parameters + ---------- + shells : SigmaDShells + The shell estimate. + d_star_sq : torch.Tensor + ``1/d**2`` of the output reflections, shape ``(M,)``, in A^-2. + epsilon : torch.Tensor or None + Multiplicity of the output reflections, shape ``(M,)``; ``None`` means ones. + f_dark : torch.Tensor or None + Dark amplitude of the output reflections for the power law; reflections with a + missing or non-finite value get a multiplier of one. + sigma_diff : torch.Tensor + Propagated uncertainty of the output differences, shape ``(M,)``. Non-finite + entries give a weight of zero. + + Returns + ------- + SigmaDEstimate + Per-reflection, detached fields all of length ``M``. + """ + with torch.no_grad(): + dtype = shells.Sigma_N.dtype + grid = d_star_sq.reshape(-1).to(dtype) + eps = ( + epsilon.reshape(-1).to(dtype) + if epsilon is not None + else torch.ones_like(grid) + ) + sig = sigma_diff.reshape(-1).to(dtype) + if shells.degenerate or shells.bin_dss.numel() == 0: + sigma_n = torch.full_like(grid, float(shells.Sigma_N[0])) + alpha = torch.full_like(grid, float(shells.alpha[0])) + beta_model = torch.full_like(grid, float(shells.beta_model[0])) + g = torch.ones_like(grid) + else: + log_sn = interp_in_dss( + grid, shells.bin_dss, torch.log(shells.Sigma_N.clamp(min=_TINY)) + ) + sigma_n = torch.exp(log_sn) + sigma_n = torch.where( + sigma_n > 10.0 * _TINY, sigma_n, torch.zeros_like(sigma_n) + ) + alpha = interp_in_dss(grid, shells.bin_dss, shells.alpha) + log_bm = interp_in_dss( + grid, shells.bin_dss, torch.log(shells.beta_model.clamp(min=_TINY)) + ) + beta_model = torch.exp(log_bm) + beta_model = torch.where( + beta_model > 10.0 * _TINY, beta_model, torch.zeros_like(beta_model) + ) + if f_dark is not None and shells.gamma != 0.0: + f = f_dark.reshape(-1).to(dtype) + log_fbar = interp_in_dss(grid, shells.bin_dss, shells.bin_log_fbar) + log_z = interp_in_dss(grid, shells.bin_dss, shells.bin_log_z) + fbar = torch.exp(log_fbar) + f_fl = torch.maximum(f, F_FLOOR_FRAC * fbar) + g = torch.exp(shells.gamma * (torch.log(f_fl) - log_fbar) - log_z) + g = torch.where(torch.isfinite(f) & (fbar > 0), g, torch.ones_like(g)) + else: + g = torch.ones_like(grid) + S = eps * sigma_n * g + sigma_sq = sig * sig + w = torch.where( + torch.isfinite(sigma_sq), + S / (S + sigma_sq).clamp(min=_TINY), + torch.zeros_like(S), + ) + return SigmaDEstimate( + S=S.detach(), + sigma_sq=sigma_sq.detach(), + w=w.detach(), + alpha=alpha.detach(), + beta_model=(eps * beta_model * g).detach(), + epsilon=eps.detach(), + shells=shells, + ) + + +class SigmaDEstimator: + """Lazy, cached difference-power estimate. + + Thin stateful wrapper around :func:`estimate_sigma_d` and + :func:`sigma_d_per_reflection`: caches the detached estimate and re-estimates only + after :meth:`reset`. **The owning target must call :meth:`reset` from its + ``maintenance()`` hook**, otherwise the estimate is frozen for the whole run. Holds + no tensors of its own beyond the cache, so it has no device to move. + + Parameters + ---------- + config : SigmaDConfig, optional + Exponent and shrinkage settings; the module defaults when omitted. + """ + + def __init__(self, config: SigmaDConfig | None = None): + self.config = config if config is not None else SigmaDConfig() + self._cache: SigmaDEstimate | None = None + self._shells: SigmaDShells | None = None + + def reset(self) -> None: + """Invalidate the cache so the next :meth:`get` re-estimates.""" + self._cache = None + + @property + def shells(self) -> SigmaDShells | None: + """Last shell estimate, for diagnostics; ``None`` until the first call.""" + return self._shells + + def get( + self, + delta_obs: torch.Tensor, + sigma_diff: torch.Tensor, + epsilon: torch.Tensor | None, + d_star_sq: torch.Tensor, + f_dark: torch.Tensor | None, + fit_mask: torch.Tensor, + *, + delta_calc: torch.Tensor | None = None, + target_dss: torch.Tensor | None = None, + out_epsilon: torch.Tensor | None = None, + out_f_dark: torch.Tensor | None = None, + out_sigma_diff: torch.Tensor | None = None, + ) -> SigmaDEstimate: + """Return the cached-or-recomputed :class:`SigmaDEstimate`. + + The fit inputs may be a pooled, flattened set (several datasets end to end); the + ``target_*`` / ``out_*`` arguments map the result onto another reflection list, + defaulting to the fit inputs themselves. + """ + if self._cache is not None: + return self._cache + shells = estimate_sigma_d( + delta_obs, + sigma_diff, + epsilon, + d_star_sq, + f_dark, + fit_mask, + delta_calc=delta_calc, + gamma=self.config.gamma, + shrink=self.config.shrink, + ) + self._shells = shells + self._cache = sigma_d_per_reflection( + shells, + d_star_sq if target_dss is None else target_dss, + epsilon if out_epsilon is None else out_epsilon, + f_dark if out_f_dark is None else out_f_dark, + sigma_diff if out_sigma_diff is None else out_sigma_diff, + ) + return self._cache diff --git a/torchref/refinement/targets/__init__.py b/torchref/refinement/targets/__init__.py index 162661b2..4bdaba50 100644 --- a/torchref/refinement/targets/__init__.py +++ b/torchref/refinement/targets/__init__.py @@ -22,6 +22,7 @@ from .collection import ( COLLECTION_XRAY_TARGETS, CollectionDifferenceIntensityTarget, + CollectionDifferenceSigmaDTarget, CollectionDifferenceTarget, CollectionMLTarget, CollectionTwoMomentIntensityTarget, @@ -93,6 +94,7 @@ "CollectionTwoMomentIntensityTarget", "CollectionMLTarget", "CollectionDifferenceIntensityTarget", + "CollectionDifferenceSigmaDTarget", "COLLECTION_XRAY_TARGETS", "MultiModelGeometryTarget", "MultiModelADPTarget", diff --git a/torchref/refinement/targets/collection/__init__.py b/torchref/refinement/targets/collection/__init__.py index 92e25460..adea453f 100644 --- a/torchref/refinement/targets/collection/__init__.py +++ b/torchref/refinement/targets/collection/__init__.py @@ -11,12 +11,14 @@ from .base import ( CollectionLossInputs, CollectionSigmaALossInputs, + CollectionSigmaDLossInputs, CollectionXrayTarget, ) from .intensity import CollectionTwoMomentIntensityTarget from .multimodel import MultiModelADPTarget, MultiModelGeometryTarget from .xray import ( CollectionDifferenceIntensityTarget, + CollectionDifferenceSigmaDTarget, CollectionDifferenceTarget, CollectionMLTarget, ) @@ -33,9 +35,11 @@ "CollectionXrayTarget", "CollectionLossInputs", "CollectionSigmaALossInputs", + "CollectionSigmaDLossInputs", "CollectionTwoMomentIntensityTarget", "CollectionDifferenceTarget", "CollectionDifferenceIntensityTarget", + "CollectionDifferenceSigmaDTarget", "CollectionMLTarget", "MultiModelGeometryTarget", "MultiModelADPTarget", diff --git a/torchref/refinement/targets/collection/_specs.py b/torchref/refinement/targets/collection/_specs.py index 7d589cd3..2d1b9436 100644 --- a/torchref/refinement/targets/collection/_specs.py +++ b/torchref/refinement/targets/collection/_specs.py @@ -4,7 +4,7 @@ invariants checked the same way at import: unique names, and **one class per row**, so dispatch is ``spec.target_cls(**kwargs)`` with nothing to branch on. -Four rows over two axes -- what the loss compares (a difference from the collection mean, +Five rows over two axes -- what the loss compares (a difference from the collection mean, or each dataset absolutely) and in which observable: ==================== =========== ============================================== @@ -12,6 +12,8 @@ ==================== =========== ============================================== ``difference`` amplitude ``F_i - F_mean`` against the model's own spread ``difference_i`` intensity the same, in intensities +``difference_sd`` amplitude ``F_i - F_mean`` against ``alpha dF_calc``, variance + ``beta_model + sigma^2`` from a sigma_D fit ``two_moment`` intensity ``|F(alpha)|^2 + sigma_alpha^2 |dF|^2`` ``ml`` amplitude each dataset absolutely, at a shared Luzzati beta ==================== =========== ============================================== @@ -27,11 +29,11 @@ from .intensity import CollectionTwoMomentIntensityTarget from .xray import ( CollectionDifferenceIntensityTarget, + CollectionDifferenceSigmaDTarget, CollectionDifferenceTarget, CollectionMLTarget, ) - @dataclass(frozen=True) class CollectionXrayTargetSpec: """One selectable collection x-ray target: a name, and the class implementing it. @@ -135,6 +137,13 @@ def by_name(self, name: str) -> CollectionXrayTargetSpec: doc="As 'difference' but on intensities, skipping the French-Wilson " "conversion that reshapes the weak tail.", ), + CollectionXrayTargetSpec( + name="difference_sd", + target_cls=CollectionDifferenceSigmaDTarget, + doc="As 'difference', centred on alpha * dF_calc with the unexplained " + "difference power beta_model (sigma_D, fitted on the free set) added to " + "the measurement variance.", + ), CollectionXrayTargetSpec( name="two_moment", target_cls=CollectionTwoMomentIntensityTarget, diff --git a/torchref/refinement/targets/collection/_util.py b/torchref/refinement/targets/collection/_util.py index 7611ccf1..1740335e 100644 --- a/torchref/refinement/targets/collection/_util.py +++ b/torchref/refinement/targets/collection/_util.py @@ -12,3 +12,16 @@ def _scale_fcalc(scaler, fcalc, model): if hasattr(scaler, "forward_mixed") and hasattr(model, "fractions"): return scaler.forward_mixed(fcalc, model.fractions) return scaler(fcalc) + + +def common_geom(data): + """``(epsilon, d_star_sq)`` on a dataset's HKL: multiplicity and ``1/d**2`` in A^-2.""" + import torch + + from torchref.base.reciprocal import get_scattering_vectors + from torchref.refinement.model_error_estimation.sigma_a import epsilon_from_hkl + + eps = epsilon_from_hkl(data.hkl, getattr(data, "spacegroup", None)) + s = get_scattering_vectors(data.hkl, data.cell) + dss = (torch.norm(s, dim=1) ** 2).to(eps.dtype) + return eps, dss diff --git a/torchref/refinement/targets/collection/base.py b/torchref/refinement/targets/collection/base.py index 23d3ce21..a274ca68 100644 --- a/torchref/refinement/targets/collection/base.py +++ b/torchref/refinement/targets/collection/base.py @@ -88,6 +88,25 @@ class CollectionSigmaALossInputs(NamedTuple): epsilon: torch.Tensor = None +class CollectionSigmaDLossInputs(NamedTuple): + """:class:`CollectionLossInputs` plus the sigma_D difference-error estimate. + + ``alpha`` and ``beta_model`` live on the **common HKL**, shape ``(n_hkl,)``, and + broadcast over the dataset axis: the coupling of the model difference to the true + one and the difference power the model leaves unexplained, fitted once on the pooled + free reflections of the timepoint rows. Detached, so gradients reach the models only + through ``model``. + """ + + obs: torch.Tensor + model: torch.Tensor + sigma: torch.Tensor + mask: torch.Tensor + keys: List[str] + alpha: torch.Tensor = None + beta_model: torch.Tensor = None + + class CollectionXrayTarget(Target): """Base class for multi-dataset X-ray targets. diff --git a/torchref/refinement/targets/collection/xray.py b/torchref/refinement/targets/collection/xray.py index 5c399224..2b4587ec 100644 --- a/torchref/refinement/targets/collection/xray.py +++ b/torchref/refinement/targets/collection/xray.py @@ -7,6 +7,10 @@ Mean-based differences on amplitudes; the primary optimization driver. :class:`CollectionDifferenceIntensityTarget` The same on intensities -- the whole class is one ``observable`` declaration. +:class:`CollectionDifferenceSigmaDTarget` + The amplitude difference centred on ``alpha * dF_calc`` with the unexplained + difference power ``beta_model`` added to the measurement variance, both from a + sigma_D fit on the free set. :class:`CollectionMLTarget` Read MLF per dataset at one shared Luzzati ``beta``, pooled over every dataset's free reflections and owned by the target rather than the scaler. The absolute channel. @@ -28,10 +32,22 @@ gaussian_per_refl, rice_per_refl, ) -from torchref.refinement.model_error_estimation.sigma_a import SigmaAEstimator, epsilon_from_hkl +from torchref.refinement.model_error_estimation.sigma_a import ( + SigmaAEstimator, + epsilon_from_hkl, +) +from torchref.refinement.model_error_estimation.sigma_d import ( + SigmaDConfig, + SigmaDEstimator, +) from torchref.utils.stats import VERBOSITY_STANDARD, StatEntry, stat -from .base import CollectionSigmaALossInputs, CollectionXrayTarget +from ._util import common_geom +from .base import ( + CollectionSigmaALossInputs, + CollectionSigmaDLossInputs, + CollectionXrayTarget, +) if TYPE_CHECKING: from torchref.io.datasets.collection import DatasetCollection @@ -174,6 +190,155 @@ class CollectionDifferenceIntensityTarget(CollectionDifferenceTarget): observable: str = "intensity" +# ========================================================================= +# CollectionDifferenceSigmaDTarget +# ========================================================================= + + +class CollectionDifferenceSigmaDTarget(CollectionDifferenceTarget): + """The difference-from-mean Gaussian with a sigma_D error model. + + The parent compares ``dF_obs`` with ``dF_calc`` under the measurement variance + alone. Here the likelihood is centred on ``alpha * dF_calc`` and its variance is + ``beta_model + sigma_diff**2``: ``alpha`` is the Gaussian coupling of the model + difference to the true one and ``beta_model`` the difference power the model does + not explain, both per resolution shell from + :class:`~torchref.refinement.model_error_estimation.sigma_d.SigmaDEstimator` fitted + on the pooled **free** reflections of the timepoint rows, with the dark-amplitude + power law carried per reflection. A poor light model therefore inflates the + variance where it fails instead of pulling the coordinates toward noise. + + At ``N = 2`` the timepoint row's difference from the mean is half the dark + subtraction; ``S`` and ``sigma_diff**2`` scale together, so the estimate is + invariant to that factor. The estimate is cached until :meth:`maintenance`, which + ``LossState`` calls after each optimizer-step block. + + Parameters + ---------- + sigma_d_config : SigmaDConfig, optional + Exponent and shrinkage settings; the module defaults when omitted. + """ + + name: str = "difference_sigma_d_xray" + + def __init__( + self, + dataset_collection: "DatasetCollection", + model_collection: "ModelCollection", + scaler: "ScalerBase" = None, + normalize: bool = True, + use_work_set: bool = True, + use_set: str = None, + verbose: int = 0, + sigma_d_config: SigmaDConfig = None, + ): + super().__init__( + dataset_collection, + model_collection, + scaler=scaler, + normalize=normalize, + use_work_set=use_work_set, + use_set=use_set, + verbose=verbose, + ) + # Constructed once; the cache lives until maintenance() resets it. + self._sigma_d = SigmaDEstimator(sigma_d_config) + self._eps_common: torch.Tensor = None + self._dss_common: torch.Tensor = None + self._geom_key: int = None + + def _common_geom(self): + """``(epsilon, d_star_sq)`` on the common HKL, cached per dark dataset.""" + data = self._dataset_collection[self._model_collection.dark_key] + key = id(data) + if self._eps_common is None or self._geom_key != key: + self._eps_common, self._dss_common = common_geom(data) + self._geom_key = key + return self._eps_common, self._dss_common + + @staticmethod + def _difference_terms(ctx): + """Difference-from-mean observations, model and propagated sigma, ``(N, n_hkl)``.""" + N = len(ctx.keys) + delta_obs = ctx.obs - ctx.obs.mean(dim=0) + delta_calc = ctx.model - ctx.model.mean(dim=0) + sum_sigma_sq = (ctx.sigma**2).sum(dim=0) + sigma_diff_sq = ctx.sigma**2 * (1 - 2.0 / N) + sum_sigma_sq / (N**2) + return delta_obs, delta_calc, torch.sqrt(sigma_diff_sq.clamp(min=1e-12)) + + def _loss_inputs(self, recalc: bool = False): + """The parent's stack plus ``alpha`` and ``beta_model`` on the common HKL. + + The estimator sees the timepoint rows only (the dark row is the reference the + differences are taken against), their free reflections, and a detached model + difference, so gradients reach the models only through ``ctx.model``. + """ + ctx = super()._loss_inputs(recalc=recalc) + delta_obs, delta_calc, sigma_diff = self._difference_terms(ctx) + delta_calc = delta_calc.detach() + dark = ctx.keys.index(self._model_collection.dark_key) + rows = [i for i in range(len(ctx.keys)) if i != dark] or [dark] + dc = self._dataset_collection + eps, dss = self._common_geom() + dtype = ctx.obs.dtype + eps, dss = eps.to(dtype), dss.to(dtype) + f_dark = ctx.obs[dark] + # The free set, independent of this target's own subset; the estimator drops + # non-finite observations itself. + fit_mask = torch.cat( + [dc[ctx.keys[i]].free.mask.to(ctx.mask.device) for i in rows] + ) + n_rows = len(rows) + est = self._sigma_d.get( + torch.cat([delta_obs[i] for i in rows]), + torch.cat([sigma_diff[i] for i in rows]), + eps.repeat(n_rows), + dss.repeat(n_rows), + f_dark.repeat(n_rows), + fit_mask, + delta_calc=torch.cat([delta_calc[i] for i in rows]), + target_dss=dss, + out_epsilon=eps, + out_f_dark=f_dark, + out_sigma_diff=sigma_diff[rows[0]], + ) + return CollectionSigmaDLossInputs( + *ctx, alpha=est.alpha.to(dtype), beta_model=est.beta_model.to(dtype) + ) + + def _per_refl(self, ctx) -> torch.Tensor: + """Gaussian NLL of the difference about ``alpha * dF_calc`` with variance + ``beta_model + sigma_diff**2``, per reflection and unreduced.""" + mask_all = ctx.mask[0] + delta_obs, delta_calc, sigma_diff = self._difference_terms(ctx) + delta_obs = torch.where(mask_all, delta_obs, torch.zeros_like(delta_obs)) + delta_calc = torch.where(mask_all, delta_calc, torch.zeros_like(delta_calc)) + sigma_diff = torch.where(mask_all, sigma_diff, torch.ones_like(sigma_diff)) + sigma_safe = sigma_diff.clamp(min=self._sigma_floor(sigma_diff, ctx.mask)) + var = ctx.beta_model.unsqueeze(0) + sigma_safe**2 + mean = ctx.alpha.unsqueeze(0) * delta_calc + nll = gaussian_per_refl(delta_obs, mean, var, var_floor=0.0) + # A single NaN would poison the whole gradient; 1e6 lets the step be rejected. + return torch.where(torch.isfinite(nll), nll, torch.full_like(nll, 1e6)) + + def maintenance(self) -> None: + """Invalidate the sigma_D estimate so it is refitted from the updated models on + the next forward (``LossState`` calls this after each optimizer-step block).""" + self._sigma_d.reset() + + def stats(self) -> Dict[str, StatEntry]: + """Base collection X-ray stats plus the sigma_D fit summary.""" + out = super().stats() + sh = self._sigma_d.shells + if sh is not None: + out["sigma_d_gamma"] = stat(float(sh.gamma), VERBOSITY_STANDARD) + out["sigma_d_tau"] = stat(float(sh.tau), VERBOSITY_STANDARD) + out["sigma_d_shells_without_power"] = stat( + float(sh.diagnostics["n_s2_clamped"]), VERBOSITY_STANDARD + ) + return out + + # ========================================================================= # CollectionMLTarget # ========================================================================= diff --git a/torchref/scaling/scaler_base.py b/torchref/scaling/scaler_base.py index 689a49f5..b8350e48 100644 --- a/torchref/scaling/scaler_base.py +++ b/torchref/scaling/scaler_base.py @@ -346,6 +346,34 @@ def get_scale(self) -> float: return torch.exp(self.iso_log_scale().mean()).item() return 1.0 + def multiplicative_scale(self) -> torch.Tensor: + """Per-reflection factor taking model amplitudes to the observed scale. + + ``K_overall * b_overall * anisotropy``: every multiplicative component + :meth:`forward` applies and none of the additive bulk-solvent term, so dividing + observed amplitudes by it returns them to the model's absolute scale, electrons. + Components not yet set up contribute ones. + + Returns + ------- + torch.Tensor + Shape ``(N,)`` over the scaler's full reflection list, detached, on + ``self.device`` in the scale parameters' dtype. + """ + c_iso = getattr(self, "c_iso", None) + dtype = c_iso.dtype if c_iso is not None else get_float_dtype() + factor = torch.ones(int(self.bins.numel()), device=self.device, dtype=dtype) + with torch.no_grad(): + if hasattr(self, "U"): + factor = factor * self.anisotropy_correction().to(factor) + if c_iso is not None: + factor = factor * torch.exp(self.iso_log_scale(self._iso_design)).to( + factor + ) + if getattr(self, "bin_wise_bfactor", None) is not None: + factor = factor * self.bin_wise_bfactor_correction().to(factor) + return factor.detach() + def setup_bin_wise_bfactor(self): """Initialize bin-wise B-factor correction parameters.""" self.bin_wise_bfactor = nn.Parameter( From b4d6f91812aad9791f16728d33413f425cf62f9b Mon Sep 17 00:00:00 2001 From: Claude Date: Sat, 19 Sep 2026 13:53:01 +0000 Subject: [PATCH 178/250] Wrap the converted dtype lines and settle their imports Lines that grew past 88 columns when their literal became get_int_dtype() are wrapped, with each remaining dtype-ok marker on the line above its literal so the dtype guard still sees it. The configured-int imports sit in their first-party groups, the canonical HKL sort key casts once before slicing, and two docstring lines are rewrapped. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- tests/unit/model/test_disorder_field.py | 1 - .../kernels/cpu/jit_reference.py | 5 +- .../kernels/cpu/variable_radius.py | 15 ++-- .../base/electron_density/map_building.py | 3 +- torchref/base/metrics/binwise_scale.py | 3 +- torchref/base/reciprocal/grid_operations.py | 1 + torchref/base/reciprocal/symmetry.py | 3 +- torchref/cli/mtz2map.py | 2 +- torchref/cli/validate_ded.py | 2 +- .../experimental/alignment/frf/data_mr.py | 10 ++- .../experimental/alignment/frf/dense_calc.py | 6 +- .../frf/kernels/cpu/legendre_shell.py | 3 +- .../experimental/alignment/frf/peak_finder.py | 1 + .../alignment/frf/sitelist_ang.py | 7 +- torchref/experimental/alignment/sh.py | 1 + .../experimental/alignment/translation.py | 7 +- .../ensemble/ensemble_amber_kl.py | 2 +- .../ensemble/quasi_crystal_amber.py | 3 +- .../experimental/ensemble/wilson_prior.py | 3 +- .../experimental/targets/forcefield_target.py | 3 +- .../targets/sampled_ml_phase_target.py | 10 ++- torchref/io/datasets/reflection_data.py | 24 ++++-- torchref/model/model.py | 13 ++-- torchref/model/parameter_wrappers.py | 10 ++- torchref/model/riding_xyz.py | 22 ++---- torchref/refinement/optimizers/curvature.py | 2 +- torchref/refinement/targets/adp/rigid_bond.py | 2 +- torchref/refinement/targets/adp/similarity.py | 7 +- torchref/refinement/targets/difference.py | 2 +- .../refinement/targets/geometry/chiral.py | 8 +- .../refinement/targets/geometry/non_bonded.py | 14 ++-- torchref/refinement/targets/similarity.py | 2 +- torchref/scaling/collection_scaler.py | 3 +- torchref/scaling/scaler_base.py | 7 +- torchref/scaling/wilson.py | 1 - torchref/symmetry/map_symmetry.py | 2 +- torchref/symmetry/reciprocal_symmetry.py | 9 +-- torchref/symmetry/symmetry.py | 4 +- torchref/topology/atom_graph.py | 14 +++- torchref/topology/build.py | 2 +- torchref/topology/builders.py | 78 +++++++++++-------- torchref/topology/edges.py | 2 +- torchref/topology/hydrogens.py | 5 +- torchref/topology/nonbonded.py | 29 ++++--- torchref/topology/residue_graph.py | 1 + torchref/topology/restraints.py | 46 ++++------- torchref/topology/riding.py | 23 ++++-- torchref/topology/topology.py | 6 +- 48 files changed, 247 insertions(+), 182 deletions(-) diff --git a/tests/unit/model/test_disorder_field.py b/tests/unit/model/test_disorder_field.py index 12fc6546..506447d4 100644 --- a/tests/unit/model/test_disorder_field.py +++ b/tests/unit/model/test_disorder_field.py @@ -15,7 +15,6 @@ import torch from torchref.config import get_int_dtype - from torchref.model.disorder_field import ( DisorderFieldTensor, build_neighbor_list, diff --git a/torchref/base/electron_density/kernels/cpu/jit_reference.py b/torchref/base/electron_density/kernels/cpu/jit_reference.py index 07b6ba0e..e81e02c2 100644 --- a/torchref/base/electron_density/kernels/cpu/jit_reference.py +++ b/torchref/base/electron_density/kernels/cpu/jit_reference.py @@ -135,9 +135,8 @@ def forward( # Scatter add to density map ny: int = density_map.shape[1] nz: int = density_map.shape[2] - strides = torch.tensor( - [ny * nz, nz, 1], device=voxel_indices.device, dtype=torch.long # dtype-ok: int64 strides make the flat voxel index int64; scatter_add_ requires int64 on torch < 2.8 - ) + # dtype-ok: int64 strides make the flat voxel index int64; scatter_add_ requires int64 on torch < 2.8 + strides = voxel_indices.new_tensor([ny * nz, nz, 1], dtype=torch.long) index_flat = torch.sum(voxel_indices.to(torch.long) * strides, dim=-1).view(-1) # dtype-ok: voxel indices flattened for scatter; indexing requires long density_map.view(-1).scatter_add_(0, index_flat, density.reshape(-1)) diff --git a/torchref/base/electron_density/kernels/cpu/variable_radius.py b/torchref/base/electron_density/kernels/cpu/variable_radius.py index aa3f4a8e..25ae2b9e 100644 --- a/torchref/base/electron_density/kernels/cpu/variable_radius.py +++ b/torchref/base/electron_density/kernels/cpu/variable_radius.py @@ -28,9 +28,9 @@ import math import torch -from torchref.config import get_int_dtype from torchref.base.electron_density.radius_policy import _u6_to_u3 +from torchref.config import get_int_dtype _PI = math.pi _PI_SQ = _PI * _PI @@ -52,8 +52,11 @@ def _bucket_by_radius(radius: torch.Tensor, center_1d: torch.Tensor): order_parts.append(idx) spans.append((float(r), cursor, cursor + idx.numel())) cursor += idx.numel() - order = (torch.cat(order_parts) if order_parts - else torch.zeros(0, dtype=get_int_dtype(), device=radius.device)) + order = ( + torch.cat(order_parts) + if order_parts + else torch.zeros(0, dtype=get_int_dtype(), device=radius.device) + ) return order, spans @@ -112,7 +115,8 @@ def add_isotropic_plain_var(density_map, xyz, adp, occ, A, B, device, dtype = xyz.device, density_map.dtype nx, ny, nz = (int(s) for s in density_map.shape) grid_dims = (nx, ny, nz) - strides = torch.tensor([ny * nz, nz, 1], device=device, dtype=torch.long) # dtype-ok: int64 strides make the flat voxel index int64; scatter_add requires int64 on torch < 2.8 + # dtype-ok: int64 strides make the flat voxel index int64; scatter_add requires int64 on torch < 2.8 + strides = torch.tensor([ny * nz, nz, 1], device=device, dtype=torch.long) grid_shape = torch.tensor(grid_dims, device=device, dtype=get_int_dtype()) order, spans, center_idx, w0 = _canonical_setup( @@ -152,7 +156,8 @@ def add_anisotropic_plain_var(density_map, xyz, u, occ, A, B, device, dtype = xyz.device, density_map.dtype nx, ny, nz = (int(s) for s in density_map.shape) grid_dims = (nx, ny, nz) - strides = torch.tensor([ny * nz, nz, 1], device=device, dtype=torch.long) # dtype-ok: int64 strides make the flat voxel index int64; scatter_add requires int64 on torch < 2.8 + # dtype-ok: int64 strides make the flat voxel index int64; scatter_add requires int64 on torch < 2.8 + strides = torch.tensor([ny * nz, nz, 1], device=device, dtype=torch.long) grid_shape = torch.tensor(grid_dims, device=device, dtype=get_int_dtype()) order, spans, center_idx, w0 = _canonical_setup( diff --git a/torchref/base/electron_density/map_building.py b/torchref/base/electron_density/map_building.py index 8840cebc..4fe71c1c 100644 --- a/torchref/base/electron_density/map_building.py +++ b/torchref/base/electron_density/map_building.py @@ -28,7 +28,8 @@ def scatter_add_nd(source, index, map): """Vectorized n-dimensional scatter-add: ``source`` ``(N,)`` into ``map`` ``(d1..dn)`` at ``index`` ``(N, ndim)``, returning the modified map. """ - map_shape = torch.tensor(map.shape, device=index.device, dtype=torch.int64) # dtype-ok: int64 shape/strides make the flat index int64; scatter_add_ requires int64 on torch < 2.8 + # dtype-ok: int64 shape/strides make the flat index int64; scatter_add_ requires int64 on torch < 2.8 + map_shape = torch.tensor(map.shape, device=index.device, dtype=torch.int64) # Convert n-dimensional indices to flat indices # For shape (d1, d2, d3, ..., dn), flat_index = i0 * (d1*d2*...*dn) + i1 * (d2*d3*...*dn) + ... + in diff --git a/torchref/base/metrics/binwise_scale.py b/torchref/base/metrics/binwise_scale.py index 8e0af305..0bfffc1f 100644 --- a/torchref/base/metrics/binwise_scale.py +++ b/torchref/base/metrics/binwise_scale.py @@ -58,7 +58,8 @@ def binwise_scale( Fo = Fo.reshape(-1) device, dtype = Fc.device, Fc.dtype - bins = bins.reshape(-1).to(device=device, dtype=torch.int64) # dtype-ok: scatter_add index; int64 required on torch < 2.8 + # dtype-ok: scatter_add index; int64 required on torch < 2.8 + bins = bins.reshape(-1).to(device=device, dtype=torch.int64) if nbins is None: nbins = int(bins.max().item()) + 1 if bins.numel() else 0 diff --git a/torchref/base/reciprocal/grid_operations.py b/torchref/base/reciprocal/grid_operations.py index 34555706..ebfa11a7 100644 --- a/torchref/base/reciprocal/grid_operations.py +++ b/torchref/base/reciprocal/grid_operations.py @@ -7,6 +7,7 @@ import math import torch + from torchref.config import get_int_dtype diff --git a/torchref/base/reciprocal/symmetry.py b/torchref/base/reciprocal/symmetry.py index f10c879e..2758474a 100644 --- a/torchref/base/reciprocal/symmetry.py +++ b/torchref/base/reciprocal/symmetry.py @@ -45,7 +45,8 @@ def _equiv_hkls_to_flat_indices( Returns ------- torch.Tensor - Flat indices, shape ``(n_ops * N,)``, in the configured int dtype, wrapped modulo the grid. + Flat indices, shape ``(n_ops * N,)``, in the configured int dtype, wrapped + modulo the grid. """ all_hkl = equiv_hkls.reshape(-1, 3) hi = torch.remainder(all_hkl[:, 0], Nx) diff --git a/torchref/cli/mtz2map.py b/torchref/cli/mtz2map.py index a3d2ba17..621a8336 100644 --- a/torchref/cli/mtz2map.py +++ b/torchref/cli/mtz2map.py @@ -18,13 +18,13 @@ import numpy as np import torch -from torchref.config import get_float_dtype, get_int_dtype from torchref.cli._common import ( add_general_args, add_resolution_args, register_timing, parse_device_str, ) +from torchref.config import get_float_dtype, get_int_dtype def main(): diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index 989f5e48..ea890220 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -29,7 +29,6 @@ import numpy as np import torch -from torchref.config import get_int_dtype from torchref.cli._common import ( add_ded_weight_args, @@ -47,6 +46,7 @@ validate_cif_files, validate_files, ) +from torchref.config import get_int_dtype from torchref.maps.ded_weights import ( DEFAULT_SCHEME, DedWeightFallbackWarning, diff --git a/torchref/experimental/alignment/frf/data_mr.py b/torchref/experimental/alignment/frf/data_mr.py index bec9f93a..9d5d506d 100644 --- a/torchref/experimental/alignment/frf/data_mr.py +++ b/torchref/experimental/alignment/frf/data_mr.py @@ -18,6 +18,7 @@ import time import torch + from torchref.config import get_int_dtype _PROFILE = bool(os.environ.get("FRF_PROFILE")) @@ -348,8 +349,10 @@ def _tick(t0): s_key = s_vectors.detach().cpu().to(torch.float64) # dtype-ok: exact clustering key on the host; the device never sees it s_mag_key = s_key.norm(dim=-1).clamp(min=1e-30) cos_key = (s_key[..., 2] / s_mag_key).clamp(min=-1.0, max=1.0) - k_s = (s_mag_key * _GROUP_SCALE_S).round().to(torch.int64) # dtype-ok: clustering key k_s*(2e7+1)+k_c overflows int32 - k_c = (cos_key * _GROUP_SCALE_COS).round().to(torch.int64) + _GROUP_SCALE_COS # dtype-ok: clustering key k_s*(2e7+1)+k_c overflows int32 + # dtype-ok: clustering key k_s*(2e7+1)+k_c overflows int32 + k_s = (s_mag_key * _GROUP_SCALE_S).round().to(torch.int64) + # dtype-ok: clustering key k_s*(2e7+1)+k_c overflows int32 + k_c = (cos_key * _GROUP_SCALE_COS).round().to(torch.int64) + _GROUP_SCALE_COS key = (k_s * (2 * _GROUP_SCALE_COS + 1) + k_c).to(s_vectors.device) uniq_key, inverse = torch.unique(key, return_inverse=True) n_clusters = int(uniq_key.shape[0]) @@ -380,7 +383,8 @@ def _group_mean(values, index, n_groups): # index to meet device values. inv_s = inv_s.to(device) n_shells = int(uniq_ks.shape[0]) - shell_of_cluster = torch.zeros(n_clusters, dtype=torch.int64, device=device) # dtype-ok: the legendre_shell kernel TORCH_CHECKs int64 shell labels + # dtype-ok: the legendre_shell kernel TORCH_CHECKs int64 shell labels + shell_of_cluster = torch.zeros(n_clusters, dtype=torch.int64, device=device) shell_of_cluster[inverse] = inv_s shell_smag = _group_mean(s_mag_all.to(comp_real), inv_s, n_shells) diff --git a/torchref/experimental/alignment/frf/dense_calc.py b/torchref/experimental/alignment/frf/dense_calc.py index c7be30d0..2c11c01d 100644 --- a/torchref/experimental/alignment/frf/dense_calc.py +++ b/torchref/experimental/alignment/frf/dense_calc.py @@ -80,9 +80,9 @@ def dense_calc_via_box( nmax = int(math.ceil(a / d_min)) idx = torch.arange(-nmax, nmax + 1, device=dev) H, K, Lg = torch.meshgrid(idx, idx, idx, indexing="ij") - hkl = torch.stack( - [H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], dim=-1 - ).to(get_int_dtype()) + hkl = torch.stack([H.reshape(-1), K.reshape(-1), Lg.reshape(-1)], dim=-1).to( + get_int_dtype() + ) # Cubic box: |s| = |hkl| / a. real = get_float_dtype() smag = hkl.to(real).norm(dim=-1) / a diff --git a/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py b/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py index 32f0c268..07f54781 100644 --- a/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py +++ b/torchref/experimental/alignment/frf/kernels/cpu/legendre_shell.py @@ -246,7 +246,8 @@ def shell_offsets(shell: torch.Tensor, n_shells: int) -> torch.Tensor: write their accumulator rows without atomics. """ counts = torch.bincount(shell, minlength=n_shells) - offsets = torch.zeros(n_shells + 1, dtype=torch.int64, device=shell.device) # dtype-ok: the kernel TORCH_CHECKs int64 offsets + # dtype-ok: the kernel TORCH_CHECKs int64 offsets + offsets = torch.zeros(n_shells + 1, dtype=torch.int64, device=shell.device) torch.cumsum(counts, dim=0, out=offsets[1:]) return offsets diff --git a/torchref/experimental/alignment/frf/peak_finder.py b/torchref/experimental/alignment/frf/peak_finder.py index de15f569..ddb8e567 100644 --- a/torchref/experimental/alignment/frf/peak_finder.py +++ b/torchref/experimental/alignment/frf/peak_finder.py @@ -28,6 +28,7 @@ from typing import List, Optional import torch + from torchref.config import get_int_dtype from ....base.alignment.rotation import rotation_matrix_euler_zyz diff --git a/torchref/experimental/alignment/frf/sitelist_ang.py b/torchref/experimental/alignment/frf/sitelist_ang.py index a7a562f2..be9bcb76 100644 --- a/torchref/experimental/alignment/frf/sitelist_ang.py +++ b/torchref/experimental/alignment/frf/sitelist_ang.py @@ -39,6 +39,7 @@ from typing import List, Tuple import torch + from torchref.config import get_int_dtype from ....config import canonical_device @@ -211,8 +212,10 @@ def build_adaptive_sample_list( # original dict scan, but no host sync / Python loop. # Hash the two rounded fracs (each in [0, 1e6]) into one int64 so we # can use the fast 1-D unique instead of a 2-D row lexsort. - a_round = (alpha_frac * 1_000_000).round().to(torch.int64) # dtype-ok: a_round*1_000_001+g_round overflows int32 - g_round = (gamma_frac * 1_000_000).round().to(torch.int64) # dtype-ok: a_round*1_000_001+g_round overflows int32 + # dtype-ok: a_round*1_000_001+g_round overflows int32 + a_round = (alpha_frac * 1_000_000).round().to(torch.int64) + # dtype-ok: a_round*1_000_001+g_round overflows int32 + g_round = (gamma_frac * 1_000_000).round().to(torch.int64) key_hash = a_round * 1_000_001 + g_round _, uniq_idx = torch.unique(key_hash, return_inverse=True) n = uniq_idx.shape[0] diff --git a/torchref/experimental/alignment/sh.py b/torchref/experimental/alignment/sh.py index 87a779f9..6b94c1f5 100644 --- a/torchref/experimental/alignment/sh.py +++ b/torchref/experimental/alignment/sh.py @@ -28,6 +28,7 @@ from typing import Optional, Tuple import torch + from torchref.config import get_int_dtype from ...config import get_float_dtype diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index 7a3ec1f1..0e5bd886 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -38,7 +38,12 @@ import torch from torchref.base.targets.xray_likelihoods import rice_per_refl -from torchref.config import get_complex_dtype, get_default_device, get_float_dtype, get_int_dtype +from torchref.config import ( + get_complex_dtype, + get_default_device, + get_float_dtype, + get_int_dtype, +) from torchref.scaling import WilsonNormaliser from torchref.scaling.weighting import (inverse_variance_weight, normalise_weight, snr_from_amplitude) diff --git a/torchref/experimental/ensemble/ensemble_amber_kl.py b/torchref/experimental/ensemble/ensemble_amber_kl.py index effb10b4..1994cf81 100644 --- a/torchref/experimental/ensemble/ensemble_amber_kl.py +++ b/torchref/experimental/ensemble/ensemble_amber_kl.py @@ -52,8 +52,8 @@ import numpy as np import torch -from torchref.config import get_int_dtype +from torchref.config import get_int_dtype from torchref.experimental.targets.amber_target import AMBER14_STANDARD, AmberTarget if TYPE_CHECKING: diff --git a/torchref/experimental/ensemble/quasi_crystal_amber.py b/torchref/experimental/ensemble/quasi_crystal_amber.py index bb712ece..bf90cd9e 100644 --- a/torchref/experimental/ensemble/quasi_crystal_amber.py +++ b/torchref/experimental/ensemble/quasi_crystal_amber.py @@ -59,8 +59,8 @@ import numpy as np import torch -from torchref.config import get_int_dtype +from torchref.config import get_int_dtype from torchref.experimental.targets.amber_target import ( AmberTarget, _OpenMMAMBERFunction, @@ -611,7 +611,6 @@ def _compose_full_omm_xyz(self, supercell_xyz_nm: torch.Tensor) -> torch.Tensor: ) return supercell_xyz_nm.index_select(1, self._omm_to_model).reshape(-1, 3) - # ------------------------------------------------------------------ # Forward # ------------------------------------------------------------------ diff --git a/torchref/experimental/ensemble/wilson_prior.py b/torchref/experimental/ensemble/wilson_prior.py index 3d9704bf..8fce2428 100644 --- a/torchref/experimental/ensemble/wilson_prior.py +++ b/torchref/experimental/ensemble/wilson_prior.py @@ -172,7 +172,8 @@ def _build_bin_assignment(self) -> None: order = torch.argsort(res) n = res.numel() nbins = min(self.nbins, max(1, n // 50)) - bin_assign = torch.empty(n, dtype=torch.long, device=res.device) # dtype-ok: scatter_add index; int64 required on torch < 2.8 + # dtype-ok: scatter_add index; int64 required on torch < 2.8 + bin_assign = torch.empty(n, dtype=torch.long, device=res.device) edges = torch.linspace(0, n, nbins + 1, device=res.device).round().long() for b in range(nbins): start = int(edges[b].item()) diff --git a/torchref/experimental/targets/forcefield_target.py b/torchref/experimental/targets/forcefield_target.py index 89beb6e6..62290cc8 100644 --- a/torchref/experimental/targets/forcefield_target.py +++ b/torchref/experimental/targets/forcefield_target.py @@ -178,7 +178,8 @@ def forward(self) -> torch.Tensor: Z = self.model.Z # Shape: (n_atoms,) # Ensure Z is long tensor - if Z.dtype != torch.long: # dtype-ok: TorchMD-Net expects a LongTensor Z; external library contract + # dtype-ok: TorchMD-Net expects a LongTensor Z; external library contract + if Z.dtype != torch.long: Z = Z.long() # Create batch tensor (single structure = all zeros) diff --git a/torchref/experimental/targets/sampled_ml_phase_target.py b/torchref/experimental/targets/sampled_ml_phase_target.py index d3510cc3..7d4fa1cd 100644 --- a/torchref/experimental/targets/sampled_ml_phase_target.py +++ b/torchref/experimental/targets/sampled_ml_phase_target.py @@ -16,9 +16,9 @@ import numpy as np import torch -from torchref.config import get_int_dtype from typing import TYPE_CHECKING, Dict, Tuple +from torchref.config import get_int_dtype from torchref.refinement.targets.base import Target from torchref.refinement.targets.xray import XrayTarget from torchref.utils.stats import ( @@ -130,7 +130,9 @@ def __init__( self.name = "xray_sampled_ml_work" if use_work_set else "xray_sampled_ml_test" # Register tunable parameters as buffers for state_dict access - self.register_buffer("_n_samples", torch.tensor(n_samples, dtype=get_int_dtype())) + self.register_buffer( + "_n_samples", torch.tensor(n_samples, dtype=get_int_dtype()) + ) self.register_buffer("_sigma_model_log", torch.tensor(sigma_model_log)) self.register_buffer("_use_analytical", torch.tensor(use_analytical)) self.register_buffer("_use_antithetic", torch.tensor(use_antithetic)) @@ -546,7 +548,9 @@ def __init__( self.add_module("_scaler_dark", scaler_dark) # Tunable parameters as buffers - self.register_buffer("_n_samples", torch.tensor(n_samples, dtype=get_int_dtype())) + self.register_buffer( + "_n_samples", torch.tensor(n_samples, dtype=get_int_dtype()) + ) self.register_buffer("_sigma_model_log", torch.tensor(sigma_model_log)) self.use_work_set = use_work_set diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index 729ae345..9893923b 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -1357,13 +1357,15 @@ def mean_res_per_bin(self) -> torch.Tensor: mean_resolutions = torch.scatter_add( mean_resolutions, 0, - self.bin_indices[mask].to(torch.int64), # dtype-ok: scatter_add index; int64 required on torch < 2.8 + # dtype-ok: scatter_add index; int64 required on torch < 2.8 + self.bin_indices[mask].to(torch.int64), self.resolution[mask], ) count_per_bin = torch.scatter_add( count_per_bin, 0, - self.bin_indices[mask].to(torch.int64), # dtype-ok: scatter_add index; int64 required on torch < 2.8 + # dtype-ok: scatter_add index; int64 required on torch < 2.8 + self.bin_indices[mask].to(torch.int64), torch.ones_like(self.resolution[mask], dtype=dtypes.int), ) mean_resolutions = mean_resolutions / count_per_bin.clamp(min=1).float() @@ -1392,12 +1394,17 @@ def mean_F_per_bin(self) -> torch.Tensor: count_per_bin = torch.zeros(self._n_bins, dtype=dtypes.int, device=self.device) mask = self.masks() mean_F = torch.scatter_add( - mean_F, 0, self.bin_indices[mask].to(torch.int64), self.F[mask] # dtype-ok: scatter_add index; int64 required on torch < 2.8 + mean_F, + 0, + # dtype-ok: scatter_add index; int64 required on torch < 2.8 + self.bin_indices[mask].to(torch.int64), + self.F[mask], ) count_per_bin = torch.scatter_add( count_per_bin, 0, - self.bin_indices[mask].to(torch.int64), # dtype-ok: scatter_add index; int64 required on torch < 2.8 + # dtype-ok: scatter_add index; int64 required on torch < 2.8 + self.bin_indices[mask].to(torch.int64), torch.ones_like(self.F[mask], dtype=dtypes.int), ) mean_F = mean_F / count_per_bin.clamp(min=1).float() @@ -1426,12 +1433,17 @@ def mean_sigma_per_bin(self) -> Optional[torch.Tensor]: count_per_bin = torch.zeros(self._n_bins, dtype=dtypes.int, device=self.device) mask = self.masks() mean_sigma = torch.scatter_add( - mean_sigma, 0, self.bin_indices[mask].to(torch.int64), self.F_sigma[mask] # dtype-ok: scatter_add index; int64 required on torch < 2.8 + mean_sigma, + 0, + # dtype-ok: scatter_add index; int64 required on torch < 2.8 + self.bin_indices[mask].to(torch.int64), + self.F_sigma[mask], ) count_per_bin = torch.scatter_add( count_per_bin, 0, - self.bin_indices[mask].to(torch.int64), # dtype-ok: scatter_add index; int64 required on torch < 2.8 + # dtype-ok: scatter_add index; int64 required on torch < 2.8 + self.bin_indices[mask].to(torch.int64), torch.ones_like(self.F_sigma[mask], dtype=dtypes.int), ) mean_sigma = mean_sigma / count_per_bin.clamp(min=1).float() diff --git a/torchref/model/model.py b/torchref/model/model.py index c835a495..879eff43 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -22,10 +22,10 @@ from torchref.base import math_torch from torchref.config import ( - get_int_dtype, canonical_device, get_default_device, get_float_dtype, + get_int_dtype, normalize_device, ) from torchref.io import cif, pdb @@ -335,7 +335,6 @@ def _iso_covers_all(self) -> bool: def _aniso_is_empty(self) -> bool: return self._sf_partition()[3] - # ========================================================================= # Cell, SpaceGroup, and Symmetry properties # ========================================================================= @@ -2071,7 +2070,6 @@ def shake_adp(self, stddev: float): new_adp, refinable_mask=self.adp.refinable_mask, name="adp" ) - def _new_model_from_df(self, df, *, strip_H=None, add_hydrogens=False): """Build a fresh model of the same class from a DataFrame. @@ -2222,7 +2220,6 @@ def hydrogenate(self, verbose: int = 0, optimize: bool = True) -> "Model": augmented = augment_atom_table(self.pdb, plan, restraints.topology) return self._new_model_from_df(augmented, strip_H=False) - def state_dict(self, destination=None, prefix="", keep_vars=False): """ Return a dictionary containing the complete state of the Model. @@ -3004,8 +3001,12 @@ def _complete_riding_waters(self, frames): source = torch.empty(len(augmented), dtype=get_int_dtype(), device=self.device) old_index = torch.as_tensor(old_rows, device=self.device) new_index = torch.as_tensor(new_rows, device=self.device) - source[old_index] = torch.arange(len(self.pdb), device=self.device, dtype=source.dtype) - source[new_index] = torch.as_tensor(plan.parent, device=self.device, dtype=source.dtype) + source[old_index] = torch.arange( + len(self.pdb), device=self.device, dtype=source.dtype + ) + source[new_index] = torch.as_tensor( + plan.parent, device=self.device, dtype=source.dtype + ) xyz = ( self.xyz.to_mixed_tensor() if hasattr(self.xyz, "to_mixed_tensor") diff --git a/torchref/model/parameter_wrappers.py b/torchref/model/parameter_wrappers.py index b009415d..288919c1 100644 --- a/torchref/model/parameter_wrappers.py +++ b/torchref/model/parameter_wrappers.py @@ -1474,11 +1474,13 @@ def _setup_sharing_groups_and_expansion( # Use sharing_groups directly as the expansion mask if sharing_groups is None: # No sharing - each atom maps to its own index - expansion_mask = torch.arange(n_atoms, dtype=torch.long, device=device) # dtype-ok: expansion_mask is a scatter_add_ index; int64 required on torch < 2.8 + # dtype-ok: expansion_mask is a scatter_add_ index; int64 required on torch < 2.8 + expansion_mask = torch.arange(n_atoms, dtype=torch.long, device=device) self._collapsed_shape = n_atoms else: # Use the provided index tensor - expansion_mask = sharing_groups.to(device=device, dtype=torch.long) # dtype-ok: expansion_mask is a scatter_add_ index; int64 required on torch < 2.8 + # dtype-ok: expansion_mask is a scatter_add_ index; int64 required on torch < 2.8 + expansion_mask = sharing_groups.to(device=device, dtype=torch.long) self._collapsed_shape = expansion_mask.max().item() + 1 self.register_buffer("expansion_mask", expansion_mask) @@ -1539,7 +1541,9 @@ def _setup_sharing_groups_and_expansion( # Create count buffer for vectorized collapse operations # counts[i] = number of atoms that map to collapsed index i - counts = torch.zeros(self._collapsed_shape, dtype=expansion_mask.dtype, device=device) + counts = torch.zeros( + self._collapsed_shape, dtype=expansion_mask.dtype, device=device + ) counts.scatter_add_(0, expansion_mask, torch.ones_like(expansion_mask)) self.register_buffer("collapse_counts", counts) diff --git a/torchref/model/riding_xyz.py b/torchref/model/riding_xyz.py index 7301825b..d2133c61 100644 --- a/torchref/model/riding_xyz.py +++ b/torchref/model/riding_xyz.py @@ -23,7 +23,6 @@ import numpy as np import torch -from torchref.config import get_int_dtype import torch.nn as nn from torchref.base.coordinates.local_frame import ( frame_is_degenerate, @@ -31,6 +30,7 @@ place_local_frame, rotate_vectors, ) +from torchref.config import get_int_dtype from torchref.model.parameter_wrappers import MixedTensor from torchref.topology.hydrogens import HydrogenFrames @@ -66,9 +66,7 @@ def _register_rows(self, n_full: int, frames: HydrogenFrames, device) -> None: if len(rows) and is_riding[rows[rows >= 0]].any(): raise ValueError(f"{name} must reference stored rows, not riding ones") - long = dict( - dtype=get_int_dtype(), device=device - ) + long = dict(dtype=get_int_dtype(), device=device) self.register_buffer("base_row", torch.as_tensor(base, **long)) self.register_buffer("h_row", torch.as_tensor(h, **long)) self.register_buffer( @@ -106,12 +104,8 @@ def _rebuild_row_cache(self) -> None: self._n1_bidx = full_to_base[self.n1_row.clamp(min=0)].clamp(min=0) self._n2_bidx = full_to_base[self.n2_row.clamp(min=0)].clamp(min=0) # ``cat([base, derived])[gather]`` lays the full table out in one gather. - order = torch.empty( - n_full, dtype=get_int_dtype(), device=device - ) - order[base] = torch.arange( - base.numel(), dtype=get_int_dtype(), device=device - ) + order = torch.empty(n_full, dtype=get_int_dtype(), device=device) + order[base] = torch.arange(base.numel(), dtype=get_int_dtype(), device=device) order[self.h_row] = base.numel() + torch.arange( self.h_row.numel(), dtype=get_int_dtype(), @@ -226,9 +220,7 @@ def __init__( for buffer in ("base_row", "h_row", "parent_row", "n1_row", "n2_row"): self.register_buffer( buffer, - torch.zeros( - 0, dtype=get_int_dtype(), device=self.device - ), + torch.zeros(0, dtype=get_int_dtype(), device=self.device), ) self.register_buffer( "frame_valid", torch.zeros(0, dtype=torch.bool, device=self.device) @@ -382,7 +374,9 @@ def _orientation_selection(self, full_mask): rows = getattr(self, "_" + kind + "_h") groups = getattr(self, "_" + kind + "_inverse") selected = full_mask[parents].to(get_int_dtype()) - selected.index_add_(0, groups, full_mask[self.h_row[rows]].to(get_int_dtype())) + selected.index_add_( + 0, groups, full_mask[self.h_row[rows]].to(get_int_dtype()) + ) selections.append(selected > 0) return selections diff --git a/torchref/refinement/optimizers/curvature.py b/torchref/refinement/optimizers/curvature.py index 6a485cc7..84d28fed 100644 --- a/torchref/refinement/optimizers/curvature.py +++ b/torchref/refinement/optimizers/curvature.py @@ -22,8 +22,8 @@ from typing import Callable, Optional, Sequence import torch -from torchref.config import get_int_dtype +from torchref.config import get_int_dtype from torchref.utils import use_portable diff --git a/torchref/refinement/targets/adp/rigid_bond.py b/torchref/refinement/targets/adp/rigid_bond.py index 2e1d3068..1a58df09 100644 --- a/torchref/refinement/targets/adp/rigid_bond.py +++ b/torchref/refinement/targets/adp/rigid_bond.py @@ -2,10 +2,10 @@ import numpy as np import torch -from torchref.config import get_int_dtype from typing import TYPE_CHECKING, Dict from torchref.base.targets.adp import adp_rigid_bond_aniso_math +from torchref.config import get_int_dtype from torchref.utils.stats import ( VERBOSITY_DEBUG, VERBOSITY_DETAILED, diff --git a/torchref/refinement/targets/adp/similarity.py b/torchref/refinement/targets/adp/similarity.py index dde02d0f..bea1fac8 100644 --- a/torchref/refinement/targets/adp/similarity.py +++ b/torchref/refinement/targets/adp/similarity.py @@ -1,9 +1,9 @@ import numpy as np import torch -from torchref.config import get_int_dtype from typing import TYPE_CHECKING, Dict from torchref.base.targets.adp import adp_simu_math, adp_simu_aniso_math +from torchref.config import get_int_dtype from torchref.utils.stats import ( VERBOSITY_DEBUG, VERBOSITY_DETAILED, @@ -95,8 +95,9 @@ def _get_pair_indices(self) -> torch.Tensor: if chunks: cached = torch.cat(chunks, dim=0).contiguous() else: - cached = torch.empty(0, 2, dtype=get_int_dtype(), - device=self.model.xyz().device) + cached = torch.empty( + 0, 2, dtype=get_int_dtype(), device=self.model.xyz().device + ) self._simu_pair_indices_cache = cached return cached diff --git a/torchref/refinement/targets/difference.py b/torchref/refinement/targets/difference.py index 19a7b971..21ef3108 100644 --- a/torchref/refinement/targets/difference.py +++ b/torchref/refinement/targets/difference.py @@ -9,11 +9,11 @@ """ import torch -from torchref.config import get_int_dtype from torch import nn from typing import TYPE_CHECKING, Dict, Literal, Optional, Tuple from .base import Target +from torchref.config import get_int_dtype from torchref.utils.stats import ( VERBOSITY_DEBUG, VERBOSITY_DETAILED, diff --git a/torchref/refinement/targets/geometry/chiral.py b/torchref/refinement/targets/geometry/chiral.py index 3dcd007c..d6e69254 100644 --- a/torchref/refinement/targets/geometry/chiral.py +++ b/torchref/refinement/targets/geometry/chiral.py @@ -1,9 +1,9 @@ import numpy as np import torch -from torchref.config import get_int_dtype from typing import TYPE_CHECKING, Dict from torchref.base.targets.chiral import chiral_math +from torchref.config import get_int_dtype from torchref.utils.stats import ( VERBOSITY_DEBUG, VERBOSITY_DETAILED, @@ -83,9 +83,9 @@ def get_violations(self, threshold: float = 0.5) -> Dict[str, torch.Tensor]: if "chiral" not in self.restraints.restraints: return { - "indices": torch.tensor([], dtype=get_int_dtype(), device=device).reshape( - 0, 4 - ), + "indices": torch.tensor( + [], dtype=get_int_dtype(), device=device + ).reshape(0, 4), "volumes": torch.tensor([], device=device), "ideal_volumes": torch.tensor([], device=device), "deviations": torch.tensor([], device=device), diff --git a/torchref/refinement/targets/geometry/non_bonded.py b/torchref/refinement/targets/geometry/non_bonded.py index 10b27e98..bc7a3f0a 100644 --- a/torchref/refinement/targets/geometry/non_bonded.py +++ b/torchref/refinement/targets/geometry/non_bonded.py @@ -7,9 +7,9 @@ import numpy as np import torch -from torchref.config import get_int_dtype from typing import TYPE_CHECKING, Dict, Tuple +from torchref.config import get_int_dtype from torchref.utils.stats import ( VERBOSITY_DEBUG, VERBOSITY_DETAILED, @@ -347,9 +347,9 @@ def get_violations(self, threshold: float = 0.0) -> Dict[str, torch.Tensor]: if "vdw" not in self.restraints.restraints: return { - "indices": torch.tensor([], dtype=get_int_dtype(), device=device).reshape( - 0, 2 - ), + "indices": torch.tensor( + [], dtype=get_int_dtype(), device=device + ).reshape(0, 2), "violations": torch.tensor([], device=device), "distances": torch.tensor([], device=device), "min_distances": torch.tensor([], device=device), @@ -360,9 +360,9 @@ def get_violations(self, threshold: float = 0.0) -> Dict[str, torch.Tensor]: if indices is None or len(indices) == 0: return { - "indices": torch.tensor([], dtype=get_int_dtype(), device=device).reshape( - 0, 2 - ), + "indices": torch.tensor( + [], dtype=get_int_dtype(), device=device + ).reshape(0, 2), "violations": torch.tensor([], device=device), "distances": torch.tensor([], device=device), "min_distances": torch.tensor([], device=device), diff --git a/torchref/refinement/targets/similarity.py b/torchref/refinement/targets/similarity.py index ab190c86..700883c3 100644 --- a/torchref/refinement/targets/similarity.py +++ b/torchref/refinement/targets/similarity.py @@ -6,10 +6,10 @@ """ import torch -from torchref.config import get_int_dtype from typing import TYPE_CHECKING, Dict from .base import Target +from torchref.config import get_int_dtype from torchref.utils.stats import ( VERBOSITY_DEBUG, VERBOSITY_DETAILED, diff --git a/torchref/scaling/collection_scaler.py b/torchref/scaling/collection_scaler.py index 89759f5a..a1d99f7c 100644 --- a/torchref/scaling/collection_scaler.py +++ b/torchref/scaling/collection_scaler.py @@ -192,7 +192,8 @@ def _calc_initial_scale_joint(self): pos_mask = torch.ones_like(fobs, dtype=torch.bool) mask = (work_mask & pos_mask).to(torch.bool) - bins = self.bins[mask].to(torch.int64) # dtype-ok: scatter_add index; int64 required on torch < 2.8 + # dtype-ok: scatter_add index; int64 required on torch < 2.8 + bins = self.bins[mask].to(torch.int64) log_ratios = ( torch.log(fobs_clamped[mask]) - torch.log(fcalc_amp[mask]) ).to(self.device) diff --git a/torchref/scaling/scaler_base.py b/torchref/scaling/scaler_base.py index 39a57cac..786e67ed 100644 --- a/torchref/scaling/scaler_base.py +++ b/torchref/scaling/scaler_base.py @@ -264,7 +264,9 @@ def calc_initial_scale(self, fcalc: torch.Tensor): initial_log_scale.detach().cpu().numpy(), ) with torch.no_grad(): - target = initial_log_scale.detach().to(self.device)[self.bins.to(get_int_dtype())] + target = initial_log_scale.detach().to(self.device)[ + self.bins.to(get_int_dtype()) + ] design = self._iso_design.to(target.dtype) coeff = torch.linalg.lstsq(design, target.unsqueeze(1)).solution.squeeze(1) self.c_iso = nn.Parameter(coeff.detach().to(self.device)) @@ -422,7 +424,8 @@ def get_binwise_mean_intensity(self, fcalc: torch.Tensor): mean_calc_intensity = torch.zeros(self.nbins, device=self.device, dtype=fobs.dtype) counts = torch.zeros(self.nbins, device=self.device, dtype=fobs.dtype) counts_vals = torch.ones_like(F_calc, device=self.device, dtype=fobs.dtype) - bins_sel = self.bins.to(torch.int64)[sel] # dtype-ok: scatter_add index; int64 required on torch < 2.8 + # dtype-ok: scatter_add index; int64 required on torch < 2.8 + bins_sel = self.bins.to(torch.int64)[sel] mean_obs_intensity = torch.scatter_add( mean_obs_intensity, 0, bins_sel, intensities[sel] ) diff --git a/torchref/scaling/wilson.py b/torchref/scaling/wilson.py index 4f5e6116..e476201b 100644 --- a/torchref/scaling/wilson.py +++ b/torchref/scaling/wilson.py @@ -267,7 +267,6 @@ def _solve_intercept( out[0] = out[0] + torch.log(ratio) return out - def _irls( self, X: torch.Tensor, diff --git a/torchref/symmetry/map_symmetry.py b/torchref/symmetry/map_symmetry.py index 43d26523..b58b65b2 100644 --- a/torchref/symmetry/map_symmetry.py +++ b/torchref/symmetry/map_symmetry.py @@ -20,8 +20,8 @@ from __future__ import annotations import torch -from torchref.config import get_int_dtype +from torchref.config import get_int_dtype from torchref.utils.device_mixin import DeviceMixin diff --git a/torchref/symmetry/reciprocal_symmetry.py b/torchref/symmetry/reciprocal_symmetry.py index d5b05c25..1954fae4 100644 --- a/torchref/symmetry/reciprocal_symmetry.py +++ b/torchref/symmetry/reciprocal_symmetry.py @@ -27,7 +27,6 @@ from torchref.config import get_float_dtype, get_int_dtype - def _expand_hkl( sym, hkl: torch.Tensor, @@ -549,11 +548,9 @@ def _canonicalize_hkl( # Lexicographic sort by (h, k, l) via composite key h_max = int(canonical_hkl.abs().max().item()) + 1 base = 2 * h_max + 1 - sort_key = ( - canonical_hkl[:, 0].to(torch.int64) * base * base # dtype-ok: composite sort key h*base^2+k*base+l overflows int32 for large Miller indices - + canonical_hkl[:, 1].to(torch.int64) * base # dtype-ok: composite sort key h*base^2+k*base+l overflows int32 for large Miller indices - + canonical_hkl[:, 2].to(torch.int64) # dtype-ok: composite sort key h*base^2+k*base+l overflows int32 for large Miller indices - ) + # dtype-ok: composite sort key h*base^2+k*base+l overflows int32 for large Miller indices + hkl64 = canonical_hkl.to(torch.int64) + sort_key = hkl64[:, 0] * base * base + hkl64[:, 1] * base + hkl64[:, 2] sort_indices = torch.argsort(sort_key) return ( diff --git a/torchref/symmetry/symmetry.py b/torchref/symmetry/symmetry.py index 7bf79e0c..428366ce 100644 --- a/torchref/symmetry/symmetry.py +++ b/torchref/symmetry/symmetry.py @@ -352,8 +352,8 @@ def expand_reciprocal(self, hkl: torch.Tensor) -> torch.Tensor: Returns ------- torch.Tensor - Shape ``(n_ops, N, 3)``, rounded to the configured int dtype. Rounding is exact for valid - operations on integer indices and only mops up float error. + Shape ``(n_ops, N, 3)``, rounded to the configured int dtype. Rounding is + exact for valid operations on integer indices and only mops up float error. """ equivalents = self.reciprocal.apply_rotations(hkl) return torch.round(equivalents).to(get_int_dtype()) diff --git a/torchref/topology/atom_graph.py b/torchref/topology/atom_graph.py index ba50e04d..c0121875 100644 --- a/torchref/topology/atom_graph.py +++ b/torchref/topology/atom_graph.py @@ -16,8 +16,8 @@ import numpy as np import torch -from torchref.config import get_int_dtype +from torchref.config import get_int_dtype from torchref.topology.edges import EdgeBlock from torchref.utils.device_mixin import DeviceMixin @@ -82,13 +82,17 @@ def _extend_paths( """ device = paths.device if paths.numel() == 0: - return torch.zeros((0, paths.shape[1] + 1), dtype=get_int_dtype(), device=device) + return torch.zeros( + (0, paths.shape[1] + 1), dtype=get_int_dtype(), device=device + ) last, prev = paths[:, -1], paths[:, -2] counts = indptr[last + 1] - indptr[last] total = int(counts.sum()) if total == 0: - return torch.zeros((0, paths.shape[1] + 1), dtype=get_int_dtype(), device=device) + return torch.zeros( + (0, paths.shape[1] + 1), dtype=get_int_dtype(), device=device + ) row = torch.repeat_interleave(torch.arange(len(paths), device=device), counts) # Offset of each slot within its own neighbour list. @@ -212,7 +216,9 @@ def implicit_h_count(self) -> Optional[torch.Tensor]: bonds[is_h[bonds[:, 0]] & ~is_h[bonds[:, 1]], 1]] ) if heavy_of_h.numel(): - present = torch.bincount(heavy_of_h, minlength=self.n_atoms).to(present.dtype) + present = torch.bincount(heavy_of_h, minlength=self.n_atoms).to( + present.dtype + ) known = self.template_h_count >= 0 missing = self.template_h_count - present return torch.where(known, missing.clamp(min=0), torch.zeros_like(missing)) diff --git a/torchref/topology/build.py b/torchref/topology/build.py index 150a80c0..901ff922 100644 --- a/torchref/topology/build.py +++ b/torchref/topology/build.py @@ -11,8 +11,8 @@ import numpy as np import pandas as pd import torch -from torchref.config import get_int_dtype +from torchref.config import get_int_dtype from torchref.topology.builders import ( InterResidueAngleBuilder, InterResidueBondBuilder, diff --git a/torchref/topology/builders.py b/torchref/topology/builders.py index 3cb0ab87..f554ba8f 100644 --- a/torchref/topology/builders.py +++ b/torchref/topology/builders.py @@ -614,7 +614,9 @@ def build( return { "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), - "references": torch.tensor(references, dtype=get_float_dtype(), device=device), + "references": torch.tensor( + references, dtype=get_float_dtype(), device=device + ), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -721,7 +723,9 @@ def build( return { "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), - "references": torch.tensor(references, dtype=get_float_dtype(), device=device), + "references": torch.tensor( + references, dtype=get_float_dtype(), device=device + ), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -842,7 +846,9 @@ def build( return { "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), - "references": torch.tensor(references, dtype=get_float_dtype(), device=device), + "references": torch.tensor( + references, dtype=get_float_dtype(), device=device + ), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), "periods": torch.tensor(periods, dtype=get_int_dtype(), device=device), } @@ -1310,7 +1316,9 @@ def finalize( return { "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), - "references": torch.tensor(references, dtype=get_float_dtype(), device=device), + "references": torch.tensor( + references, dtype=get_float_dtype(), device=device + ), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -1413,7 +1421,9 @@ def build( return { "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), - "references": torch.tensor(references, dtype=get_float_dtype(), device=device), + "references": torch.tensor( + references, dtype=get_float_dtype(), device=device + ), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -1535,7 +1545,9 @@ def finalize( return { "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), - "references": torch.tensor(references, dtype=get_float_dtype(), device=device), + "references": torch.tensor( + references, dtype=get_float_dtype(), device=device + ), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -1649,7 +1661,9 @@ def build( return { "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), - "references": torch.tensor(references, dtype=get_float_dtype(), device=device), + "references": torch.tensor( + references, dtype=get_float_dtype(), device=device + ), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), } @@ -1786,7 +1800,9 @@ def finalize_disulfide( return { "indices": torch.tensor(indices, dtype=get_int_dtype(), device=device), - "references": torch.tensor(references, dtype=get_float_dtype(), device=device), + "references": torch.tensor( + references, dtype=get_float_dtype(), device=device + ), "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), "periods": torch.tensor(periods, dtype=get_int_dtype(), device=device), } @@ -2077,34 +2093,34 @@ def build( planes_by_size: Dict[int, List[Tuple[np.ndarray, np.ndarray]]] = {} for res_i_idx, res_next_idx in pairs: - for map_i in conf_maps[res_i_idx]: - for map_next in conf_maps[res_next_idx]: + for map_i in conf_maps[res_i_idx]: + for map_next in conf_maps[res_next_idx]: - for plane_data in link_data.planes: - comp_ids = plane_data["comp_ids"] - atom_names = plane_data["atoms"] - sigmas = plane_data["sigmas"] + for plane_data in link_data.planes: + comp_ids = plane_data["comp_ids"] + atom_names = plane_data["atoms"] + sigmas = plane_data["sigmas"] - plane_indices = [] - plane_sigmas = [] - all_found = True + plane_indices = [] + plane_sigmas = [] + all_found = True - for i, (comp_id, atom_name, sigma) in enumerate( + for i, (comp_id, atom_name, sigma) in enumerate( zip(comp_ids, atom_names, sigmas) ): - atom_map = map_i if comp_id == "1" else map_next - if atom_name in atom_map: - plane_indices.append(atom_map[atom_name]) - plane_sigmas.append(sigma) - else: - all_found = False - break - - if all_found and len(plane_indices) >= 3: - n_atoms = len(plane_indices) - if n_atoms not in planes_by_size: - planes_by_size[n_atoms] = [] - planes_by_size[n_atoms].append( + atom_map = map_i if comp_id == "1" else map_next + if atom_name in atom_map: + plane_indices.append(atom_map[atom_name]) + plane_sigmas.append(sigma) + else: + all_found = False + break + + if all_found and len(plane_indices) >= 3: + n_atoms = len(plane_indices) + if n_atoms not in planes_by_size: + planes_by_size[n_atoms] = [] + planes_by_size[n_atoms].append( ( np.array(plane_indices, dtype=np.int64), np.array(plane_sigmas, dtype=np.float64), diff --git a/torchref/topology/edges.py b/torchref/topology/edges.py index b99288b9..50b3877b 100644 --- a/torchref/topology/edges.py +++ b/torchref/topology/edges.py @@ -16,8 +16,8 @@ import numpy as np import torch -from torchref.config import get_int_dtype +from torchref.config import get_int_dtype from torchref.utils.device_mixin import DeviceMixin #: Origin order per edge type. Fixes the block layout so a rebuild on the same diff --git a/torchref/topology/hydrogens.py b/torchref/topology/hydrogens.py index 0f5e3ddd..b4b41ed3 100644 --- a/torchref/topology/hydrogens.py +++ b/torchref/topology/hydrogens.py @@ -28,6 +28,7 @@ import numpy as np import torch + from torchref.config import get_float_dtype, get_int_dtype #: Standard heavy-atom valences, one of the two budgets that cap how many hydrogens a @@ -1145,9 +1146,7 @@ def sorted_by_row(self) -> "HydrogenFrames": def to_tensors(self, device=None) -> Dict[str, torch.Tensor]: """Return frame and orientation arrays as tensors, keyed by field name.""" return { - "h_row": torch.as_tensor( - self.h_row, dtype=get_int_dtype(), device=device - ), + "h_row": torch.as_tensor(self.h_row, dtype=get_int_dtype(), device=device), "parent_row": torch.as_tensor( self.parent_row, dtype=get_int_dtype(), device=device ), diff --git a/torchref/topology/nonbonded.py b/torchref/topology/nonbonded.py index 236f6f85..8128f60e 100644 --- a/torchref/topology/nonbonded.py +++ b/torchref/topology/nonbonded.py @@ -191,9 +191,7 @@ def build_cell_list( starts = torch.zeros(len(unique_cells) + 1, dtype=get_int_dtype(), device=device) starts[1:] = counts.cumsum(0) - cell_lookup = torch.full( - (n_grid_total,), -1, dtype=get_int_dtype(), device=device - ) + cell_lookup = torch.full((n_grid_total,), -1, dtype=get_int_dtype(), device=device) cell_lookup[unique_cells] = torch.arange( len(unique_cells), dtype=get_int_dtype(), device=device ) @@ -517,11 +515,13 @@ def exclusion_set_to_hash( Hash: min(i,j) * max_idx + max(i,j), sorted for searchsorted. """ if not exclusion_set: - return torch.tensor([], dtype=torch.long, device=device) # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 + # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 + return torch.tensor([], dtype=torch.long, device=device) arr = np.array(list(exclusion_set), dtype=np.int64) hashes = arr[:, 0] * max_idx + arr[:, 1] # already (min, max) hashes.sort() - return torch.tensor(hashes, dtype=torch.long, device=device) # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 + # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 + return torch.tensor(hashes, dtype=torch.long, device=device) def filter_pairs( @@ -654,12 +654,15 @@ def build_vdw_restraints_gpu( identity_indices = is_identity.nonzero(as_tuple=True)[0] if len(identity_indices) == 0: # Identity not in valid combos — should not happen, but add it - op_indices = torch.cat([ - torch.zeros(1, dtype=get_int_dtype(), device=device), op_indices - ]) - cell_offsets_valid = torch.cat([ - torch.zeros(1, 3, dtype=get_int_dtype(), device=device), cell_offsets_valid - ]) + op_indices = torch.cat( + [torch.zeros(1, dtype=get_int_dtype(), device=device), op_indices] + ) + cell_offsets_valid = torch.cat( + [ + torch.zeros(1, 3, dtype=get_int_dtype(), device=device), + cell_offsets_valid, + ] + ) identity_combo = 0 M = len(op_indices) else: @@ -774,7 +777,9 @@ def build_vdw_restraints_gpu( "valid_op_indices": op_indices, "valid_cell_offsets": cell_offsets_valid, "grid_dims": grid_dims, - "identity_combo": torch.tensor(identity_combo, dtype=get_int_dtype(), device=device), + "identity_combo": torch.tensor( + identity_combo, dtype=get_int_dtype(), device=device + ), } if verbose > 0: diff --git a/torchref/topology/residue_graph.py b/torchref/topology/residue_graph.py index 5231f02d..0cc77c10 100644 --- a/torchref/topology/residue_graph.py +++ b/torchref/topology/residue_graph.py @@ -14,6 +14,7 @@ import numpy as np import torch + from torchref.config import get_int_dtype #: SG-SG separation below which two cysteines are taken to be disulfide-bonded. diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index 5688b11f..f1b65c72 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -26,17 +26,16 @@ import torch from torch.nn import Module +from torchref.config import get_float_dtype, get_int_dtype from torchref.topology.monomer.cif import ( find_cif_file_in_library, read_cif, read_link_definitions, ) -from torchref.config import get_float_dtype, get_int_dtype from torchref.utils.debug_utils import DebugMixin from torchref.utils.device_mixin import DeviceMixin - class Restraints(DeviceMixin, DebugMixin, Module): """ Restraints handler for crystallographic model refinement. @@ -223,13 +222,6 @@ def get_vdw_radii(self) -> torch.Tensor: # Restraint storage # ========================================================================= - - - - - - - @property def restraints(self) -> dict: """Restraint groups as ``[edge type][origin][property]``. @@ -339,7 +331,6 @@ def _load_cif_dictionaries(self, cif_path): f"and will have no restraints applied: {self.missing_residues}" ) - def _load_rama_surfaces(self, device: torch.device): """Load pre-computed Ramachandran NLL surfaces as a buffer.""" from torchref.topology.ramachandran import load_nll_surfaces @@ -391,12 +382,6 @@ def build_restraints(self): self.debug_on_error(e, context="Restraints.build_restraints") raise - - - - - - def _find_nearby_pairs_spatial_hash(self, xyz, cutoff=6.0): """Atom pairs within ``cutoff`` of each other, as (M, 2) rows with i < j. @@ -651,7 +636,6 @@ def h_topo(self): """Access riding hydrogen topology (None if not built).""" return getattr(self, "_h_topo", None) - def _build_h_exclusion_hash(self, h_topo, device): """Sorted 1-D hash tensor of H-specific 1-2 and 1-3 exclusions. @@ -659,7 +643,8 @@ def _build_h_exclusion_hash(self, h_topo, device): ``torch.searchsorted`` lookup. """ if h_topo is None or h_topo.n_hydrogens == 0: - return torch.tensor([], dtype=torch.long, device=device) # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 + # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 + return torch.tensor([], dtype=torch.long, device=device) n_heavy = len(self.pdb) n_h = h_topo.n_hydrogens @@ -684,13 +669,15 @@ def _build_h_exclusion_hash(self, h_topo, device): exclusions.add((min(h_combined, nb), max(h_combined, nb))) if not exclusions: - return torch.tensor([], dtype=torch.long, device=device) # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 + # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 + return torch.tensor([], dtype=torch.long, device=device) arr = np.array(list(exclusions), dtype=np.int64) max_idx = max(n_heavy + n_h, int(arr.max()) + 1) hashes = arr[:, 0] * max_idx + arr[:, 1] hashes.sort() - return torch.tensor(hashes, dtype=torch.long, device=device) # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 + # dtype-ok: packed pair key min*max_idx+max overflows int32 beyond ~46k atoms; searchsorted needs both sides int64 + return torch.tensor(hashes, dtype=torch.long, device=device) def _build_vdw_restraints( self, cutoff=6.0, sigma=0.2, inter_residue_only=True, use_spatial_hash=True @@ -895,15 +882,21 @@ def _build_vdw_restraints_legacy( nearby_pairs = ( torch.tensor(pairs_list, dtype=get_int_dtype(), device=device) if pairs_list - else torch.tensor([], dtype=get_int_dtype(), device=device).reshape(0, 2) + else torch.tensor([], dtype=get_int_dtype(), device=device).reshape( + 0, 2 + ) ) empty_result = { - "indices": torch.tensor([], dtype=get_int_dtype(), device=device).reshape(0, 2), + "indices": torch.tensor([], dtype=get_int_dtype(), device=device).reshape( + 0, 2 + ), "min_distances": torch.tensor([], dtype=get_float_dtype(), device=device), "sigmas": torch.tensor([], dtype=get_float_dtype(), device=device), "symop_indices": torch.tensor([], dtype=get_int_dtype(), device=device), - "cell_offsets": torch.tensor([], dtype=get_int_dtype(), device=device).reshape(0, 3), + "cell_offsets": torch.tensor( + [], dtype=get_int_dtype(), device=device + ).reshape(0, 3), } if len(nearby_pairs) == 0: @@ -1162,8 +1155,6 @@ def get_count(rtype, origin): f"torsions={n_torsions}, peptide_bonds={n_bonds_peptide})" ) - - def bond_lengths(self, idx, xyz: torch.Tensor = None): """ Compute current bond lengths from atomic coordinates. @@ -1505,7 +1496,6 @@ def _wrap_torsion_periodicity(self, diff_rad, periods): # All periods are 0 or 1, simple wrapping return torch.remainder(diff_rad + torch.pi, 2.0 * torch.pi) - torch.pi - def torsion_deviations_with_sigmas(self, xyz: torch.Tensor = None): """ Compute torsion deviations (wrapped for periodicity) and sigmas. @@ -1538,9 +1528,6 @@ def torsion_deviations_with_sigmas(self, xyz: torch.Tensor = None): return deviations_rad, sigmas_deg - - - def adp_b_differences(self, adp: torch.Tensor = None): """ Compute B-factor differences between bonded atoms. @@ -1571,4 +1558,3 @@ def adp_b_differences(self, adp: torch.Tensor = None): if diffs_list: return torch.cat(diffs_list, dim=0) return torch.tensor([], device=b_factors.device) - diff --git a/torchref/topology/riding.py b/torchref/topology/riding.py index 894476e3..29e1872e 100644 --- a/torchref/topology/riding.py +++ b/torchref/topology/riding.py @@ -439,7 +439,9 @@ def build_hydrogen_topology( topo.parent_neighbor_idx = torch.zeros( 0, MAX_HEAVY_NB, dtype=get_int_dtype(), device=device ) - topo.parent_neighbor_count = torch.zeros(0, dtype=get_int_dtype(), device=device) + topo.parent_neighbor_count = torch.zeros( + 0, dtype=get_int_dtype(), device=device + ) topo.h_chainid_enc = torch.zeros(0, dtype=get_int_dtype(), device=device) topo.h_resseq = torch.zeros(0, dtype=get_int_dtype(), device=device) return topo @@ -466,7 +468,9 @@ def build_hydrogen_topology( idxs = np.where(mask)[0] type_bounds[t] = (int(idxs[0]), int(idxs[-1]) + 1) - topo.h_parent_idx = torch.tensor(acc_parent_idx, dtype=get_int_dtype(), device=device) + topo.h_parent_idx = torch.tensor( + acc_parent_idx, dtype=get_int_dtype(), device=device + ) topo.h_bond_length = torch.tensor(acc_bond_length, dtype=fdtype, device=device) topo.h_vdw_radius = torch.full((n_h_total,), 1.20, dtype=fdtype, device=device) topo.h_placement_type = torch.tensor( @@ -480,7 +484,9 @@ def build_hydrogen_topology( acc_nb_count, dtype=get_int_dtype(), device=device ) topo.type_bounds = type_bounds # dict: type_code -> (start, end) - topo.h_chainid_enc = torch.tensor(acc_chainid_enc, dtype=get_int_dtype(), device=device) + topo.h_chainid_enc = torch.tensor( + acc_chainid_enc, dtype=get_int_dtype(), device=device + ) topo.h_resseq = torch.tensor(acc_resseq, dtype=get_int_dtype(), device=device) if verbose > 0: @@ -737,7 +743,9 @@ def build_h_candidate_pairs( if n_h == 0: for name in ("cand_idx_i", "cand_idx_j", "cand_symop_idx"): setattr(h_topo, name, torch.zeros(0, dtype=get_int_dtype(), device=device)) - h_topo.cand_cell_offset = torch.zeros(0, 3, dtype=get_int_dtype(), device=device) + h_topo.cand_cell_offset = torch.zeros( + 0, 3, dtype=get_int_dtype(), device=device + ) h_topo.cand_min_dist = torch.zeros(0, dtype=dtypes.float, device=device) return @@ -837,7 +845,9 @@ def _same_res(chain_a, resseq_a, chain_b, resseq_b): if not acc_idx_i: for name in ("cand_idx_i", "cand_idx_j", "cand_symop_idx"): setattr(h_topo, name, torch.zeros(0, dtype=get_int_dtype(), device=device)) - h_topo.cand_cell_offset = torch.zeros(0, 3, dtype=get_int_dtype(), device=device) + h_topo.cand_cell_offset = torch.zeros( + 0, 3, dtype=get_int_dtype(), device=device + ) h_topo.cand_min_dist = torch.zeros(0, dtype=dtypes.float, device=device) return @@ -853,7 +863,8 @@ def _same_res(chain_a, resseq_a, chain_b, resseq_b): max_idx = n_heavy + n_h norm_i = torch.minimum(cand_i, cand_j) norm_j = torch.maximum(cand_i, cand_j) - pair_hash = norm_i.to(torch.int64) * max_idx + norm_j.to(torch.int64) # dtype-ok: packed pair key overflows int32; searchsorted needs int64 like the table + # dtype-ok: packed pair key overflows int32; searchsorted needs int64 like the table + pair_hash = norm_i.to(torch.int64) * max_idx + norm_j.to(torch.int64) ins = torch.searchsorted(h_excl_hash, pair_hash).clamp( max=len(h_excl_hash) - 1 ) diff --git a/torchref/topology/topology.py b/torchref/topology/topology.py index 516e4731..ab777ac2 100644 --- a/torchref/topology/topology.py +++ b/torchref/topology/topology.py @@ -13,8 +13,8 @@ import numpy as np import torch -from torchref.config import get_int_dtype +from torchref.config import get_int_dtype from torchref.topology.atom_graph import AtomGraph from torchref.topology.residue_graph import ResidueGraph from torchref.utils.device_mixin import DeviceMixin @@ -98,7 +98,9 @@ def subset(self, keep) -> "Topology": raise ValueError("subset would keep no atoms") n_kept = int(mask.sum()) - remap = torch.full((self.n_atoms,), -1, dtype=get_int_dtype(), device=mask.device) + remap = torch.full( + (self.n_atoms,), -1, dtype=get_int_dtype(), device=mask.device + ) remap[mask] = torch.arange(n_kept, dtype=get_int_dtype(), device=mask.device) # A residue survives if any of its atoms does. Counting per residue also From 31fed868928932f7daff6366a6a0bacd414bc664 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 23 Sep 2026 08:04:01 +0000 Subject: [PATCH 179/250] Build the CPU density kernel's strides with torch.tensor The CPU density kernel is compiled by TorchScript when torchref is imported, and TorchScript has no Tensor.new_tensor, so the import failed on any machine without a cached compiled kernel. The strides go back to torch.tensor with an explicit device and int64 dtype; the compiled kernel is bit-identical to dev's. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- .../base/electron_density/kernels/cpu/jit_reference.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/torchref/base/electron_density/kernels/cpu/jit_reference.py b/torchref/base/electron_density/kernels/cpu/jit_reference.py index e81e02c2..be65d44b 100644 --- a/torchref/base/electron_density/kernels/cpu/jit_reference.py +++ b/torchref/base/electron_density/kernels/cpu/jit_reference.py @@ -135,8 +135,13 @@ def forward( # Scatter add to density map ny: int = density_map.shape[1] nz: int = density_map.shape[2] - # dtype-ok: int64 strides make the flat voxel index int64; scatter_add_ requires int64 on torch < 2.8 - strides = voxel_indices.new_tensor([ny * nz, nz, 1], dtype=torch.long) + # Compiled by TorchScript: no Tensor.new_tensor, no config dtype getters. + strides = torch.tensor( + [ny * nz, nz, 1], + device=voxel_indices.device, + # dtype-ok: int64 strides make the flat voxel index int64; scatter_add_ requires int64 on torch < 2.8 + dtype=torch.long, + ) index_flat = torch.sum(voxel_indices.to(torch.long) * strides, dim=-1).view(-1) # dtype-ok: voxel indices flattened for scatter; indexing requires long density_map.view(-1).scatter_add_(0, index_flat, density.reshape(-1)) From 8ea0e95ec48cba0c4a0e2064d6fd11fdfaeb3c96 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 23 Sep 2026 08:15:28 +0000 Subject: [PATCH 180/250] Keep flat grid indices in int64 h*Ny*Nz + k*Nz + l is a packed key: formed in the configured int32 it wraps once a grid holds more than 2**31 voxels, so place_on_grid and the reciprocal symmetry extractor would address aliased voxels and the translation function would merge bins. Those sites and the CPU variable-radius sort key build the index in int64 again, the packed-key exception to the int dtype policy, and a test pins the extractor's indices on a 2048**3 grid. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- tests/unit/base/test_flat_grid_index.py | 23 +++++++++++++++++++ .../kernels/cpu/variable_radius.py | 5 ++-- torchref/base/reciprocal/grid_operations.py | 13 ++++++----- torchref/base/reciprocal/symmetry.py | 10 ++++---- .../experimental/alignment/translation.py | 3 ++- 5 files changed, 40 insertions(+), 14 deletions(-) create mode 100644 tests/unit/base/test_flat_grid_index.py diff --git a/tests/unit/base/test_flat_grid_index.py b/tests/unit/base/test_flat_grid_index.py new file mode 100644 index 00000000..0e8255e6 --- /dev/null +++ b/tests/unit/base/test_flat_grid_index.py @@ -0,0 +1,23 @@ +"""Flat reciprocal-grid indices stay exact on grids larger than int32 can address. + +``h*Ny*Nz + k*Nz + l`` is a packed key: formed in the default int32 it wraps past +2**31 voxels and addresses aliased voxels, so it is built in int64 whatever dtype the +Miller indices arrive in. +""" + +import pytest +import torch + +from torchref.base.reciprocal.symmetry import _equiv_hkls_to_flat_indices +from torchref.config import get_int_dtype + +pytestmark = pytest.mark.unit + + +def test_flat_indices_are_exact_past_int32(): + n = 2048 # 2048**3 voxels; the helper never allocates the grid + hkl = torch.tensor([[[1000, -3, 7], [-1, 0, 2047]]], dtype=get_int_dtype()) + flat = _equiv_hkls_to_flat_indices(hkl, n, n, n) + expected = [(h % n) * n * n + (k % n) * n + (l % n) for h, k, l in hkl[0].tolist()] + assert flat.dtype == torch.int64 + assert flat.tolist() == expected diff --git a/torchref/base/electron_density/kernels/cpu/variable_radius.py b/torchref/base/electron_density/kernels/cpu/variable_radius.py index 25ae2b9e..3feab135 100644 --- a/torchref/base/electron_density/kernels/cpu/variable_radius.py +++ b/torchref/base/electron_density/kernels/cpu/variable_radius.py @@ -96,8 +96,9 @@ def _canonical_setup(xyz, inv_frac, frac, grid_dims, radius_per_atom, dtype): # w0: atom position relative to its anchor node, in Cartesian. This is what # centres the sphere on the atom rather than on the node. w0 = (xyz_frac - center_idx.to(dtype) / grid_f) @ frac.T - center_1d = ((center_idx[:, 0] % nx) * (ny * nz) - + (center_idx[:, 1] % ny) * nz + (center_idx[:, 2] % nz)) + # dtype-ok: the flat voxel index overflows int32 above 2**31 voxels + c = center_idx.to(torch.int64) + center_1d = (c[:, 0] % nx) * (ny * nz) + (c[:, 1] % ny) * nz + (c[:, 2] % nz) order, spans = _bucket_by_radius(radius_per_atom, center_1d) return order, spans, center_idx[order], w0[order] diff --git a/torchref/base/reciprocal/grid_operations.py b/torchref/base/reciprocal/grid_operations.py index ebfa11a7..2aabbd1c 100644 --- a/torchref/base/reciprocal/grid_operations.py +++ b/torchref/base/reciprocal/grid_operations.py @@ -48,15 +48,16 @@ def place_on_grid( device = structure_factor.device dtype = structure_factor.dtype Nx, Ny, Nz = [int(x) for x in grid_size] - hkls = hkls.to(device=device) - h = hkls[:, 0].to(get_int_dtype()) - k = hkls[:, 1].to(get_int_dtype()) - l = hkls[:, 2].to(get_int_dtype()) + # dtype-ok: the flat index h*Ny*Nz + k*Nz + l overflows int32 above 2**31 voxels + hkls = hkls.to(device=device, dtype=torch.int64) + h = hkls[:, 0] + k = hkls[:, 1] + l = hkls[:, 2] hi = torch.remainder(h, Nx) ki = torch.remainder(k, Ny) li = torch.remainder(l, Nz) - lin = (hi * (Ny * Nz) + ki * Nz + li).to(get_int_dtype()) # (N,) + lin = hi * (Ny * Nz) + ki * Nz + li # (N,) grid = torch.zeros((B, Nx * Ny * Nz), dtype=dtype, device=device) grid = grid.index_add(1, lin, structure_factor) # (B, Nx*Ny*Nz) @@ -64,7 +65,7 @@ def place_on_grid( hi_sym = torch.remainder(-h, Nx) ki_sym = torch.remainder(-k, Ny) li_sym = torch.remainder(-l, Nz) - lin_sym = (hi_sym * (Ny * Nz) + ki_sym * Nz + li_sym).to(get_int_dtype()) + lin_sym = hi_sym * (Ny * Nz) + ki_sym * Nz + li_sym vals_conj = torch.conj(structure_factor) grid = grid.index_add(1, lin_sym, vals_conj) diff --git a/torchref/base/reciprocal/symmetry.py b/torchref/base/reciprocal/symmetry.py index 2758474a..04c6ecdd 100644 --- a/torchref/base/reciprocal/symmetry.py +++ b/torchref/base/reciprocal/symmetry.py @@ -22,7 +22,7 @@ class through import torch -from torchref.config import canonical_device, get_int_dtype +from torchref.config import canonical_device from torchref.utils.autograd_ops import gather_with_index_add from torchref.utils.device_mixin import DeviceMixin @@ -45,14 +45,14 @@ def _equiv_hkls_to_flat_indices( Returns ------- torch.Tensor - Flat indices, shape ``(n_ops * N,)``, in the configured int dtype, wrapped - modulo the grid. + Flat indices, shape ``(n_ops * N,)``, dtype ``int64``, wrapped modulo the grid. """ - all_hkl = equiv_hkls.reshape(-1, 3) + # dtype-ok: the flat index h*Ny*Nz + k*Nz + l overflows int32 above 2**31 voxels + all_hkl = equiv_hkls.reshape(-1, 3).to(torch.int64) hi = torch.remainder(all_hkl[:, 0], Nx) ki = torch.remainder(all_hkl[:, 1], Ny) li = torch.remainder(all_hkl[:, 2], Nz) - return (hi * (Ny * Nz) + ki * Nz + li).to(get_int_dtype()) + return hi * (Ny * Nz) + ki * Nz + li class ReciprocalSymmetryExtractor(DeviceMixin): diff --git a/torchref/experimental/alignment/translation.py b/torchref/experimental/alignment/translation.py index 0e5bd886..0fdb41b5 100644 --- a/torchref/experimental/alignment/translation.py +++ b/torchref/experimental/alignment/translation.py @@ -400,7 +400,8 @@ def fast_translation_function( G = cand.G.to(device=device, dtype=cplx) S, N = G.shape coeff = obs.coeff.to(device=device, dtype=cplx) - h_R_int = cand.h_R.round().to(get_int_dtype()) + # dtype-ok: the flat translation-grid index overflows int32 above 2**31 grid points + h_R_int = cand.h_R.round().to(torch.int64) # The pair (j, i) is the conjugate of (i, j) at -dh, so the map is twice # the real part of the upper triangle's transform plus the diagonal, which From 76564cbccdf8ba3b87ad72ec551fe85d3ca154ca Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Wed, 23 Sep 2026 14:21:49 +0200 Subject: [PATCH 181/250] Annotate difference MTZ columns with named datasets and history Group the difference MTZ columns into observed, difference, light_model, extrapolated_light and two_moment datasets, with one history line each, so FWT/PHWT reads as /torchref/extrapolated_light/FWT. Labels are unchanged, so Coot still auto-opens the extrapolated map. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 1 + docs/user_guide/cli.rst | 7 ++ tests/integration/test_cli_two_moment_mtz.py | 23 ++++++ torchref/cli/collection_difference_refine.py | 80 ++++++++++++++++++++ 4 files changed, 111 insertions(+) diff --git a/docs/changelog.rst b/docs/changelog.rst index 48350ebb..4c09517e 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- The difference MTZ groups its columns into named datasets -- ``observed``, ``difference``, ``light_model``, ``extrapolated_light``, ``two_moment`` -- with one history line describing each, so ``FWT``/``PHWT`` reads as ``/torchref/extrapolated_light/FWT`` (the extrapolated light-state map ``2*FEXT - Fc``). Labels are unchanged and Coot still auto-opens it - ``torchref.difference-map``, ``torchref.difference-refine`` and ``torchref.validate-ded`` gain ``--ded-weight {sigma_d,inverse_variance,none}`` and ``--sigma-d-gamma``. The difference MTZ now carries the unweighted ``DF``/``SIGDF`` on ``PHDELWT`` with one mean-one weight column per scheme, ``W_SD`` and ``W_IVW`` (MTZ type W), and the observed-to-model scale ``KSCALE``; ``DELFWT`` is no longer written, build the map with ``torchref.mtz2map -csf DF -cw W_IVW -cphi PHDELWT``. Registered in ``torchref.maps.ded_weights`` - Added the ``sigma_D`` estimator (``torchref.refinement.model_error_estimation.sigma_d``): the expected true difference power per resolution shell, ``mean(dF_obs^2) - mean(sigma^2)`` with a fitted ``F_dark^gamma`` amplitude law and DerSimonian-Laird shrinkage of the signed shell power toward a decaying exponential in ``d*^2`` (fitted on all shells, so a dataset without a difference yields no power instead of the positive half of its noise), giving the Wiener weight ``S/(S + sigma^2)`` and, with a difference model, ``alpha``/``beta_model``. Inverse-variance weights suppress the strong reflections whose difference power is 10-70x that of weak ones; on independent half-datasets a Wiener weight with the true power raised map agreement 1.2-1.8x in effective patterns. The single-dataset estimate inherits the calibration of the reported sigmas, and on the campaign TD1 data (sigmas ~1.5x too large at high resolution) it emptied 60-90 % of the shells, so inverse variance stays the default; ``sigma_d`` reports its clamped-shell count and falls back to inverse variance with a warning when every shell is empty - Added the ``difference_sd`` collection target (``CollectionDifferenceSigmaDTarget``): the difference Gaussian centred on ``alpha * dF_calc`` with variance ``beta_model + sigma_diff^2`` from ``sigma_D`` fitted on the free set. Selected with ``torchref.difference-refine --difference-target difference_sd``; ``difference`` stays the default diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index 171f6e90..eadaa72d 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -175,6 +175,13 @@ phase and the extrapolated map ``FWT``/``PHWT``: -dsf dark.mtz -lsf light.mtz \ --fraction 0.37 -o results.mtz +``FWT``/``PHWT`` keep the standard labels so Coot auto-opens the map, but here they +are the extrapolated light-state map ``2*FEXT - Fc``, not a ``2mFo-DFc``. The file +records this: the columns sit in named MTZ datasets -- ``observed``, ``difference``, +``light_model``, ``extrapolated_light`` and, when written, ``two_moment`` -- so Coot's +column chooser shows ``/torchref/extrapolated_light/FWT``, and ``gemmi mtz`` prints a +history line per dataset. ``torchref.difference-refine`` writes the same file. + **Key options:** ``--ded-weight {inverse_variance,sigma_d,none}`` selects the scheme the model-phased and two-moment difference columns carry (default ``inverse_variance``; ``sigma_d`` needs calibrated sigmas, reports how many shells diff --git a/tests/integration/test_cli_two_moment_mtz.py b/tests/integration/test_cli_two_moment_mtz.py index 95a46632..01270bae 100644 --- a/tests/integration/test_cli_two_moment_mtz.py +++ b/tests/integration/test_cli_two_moment_mtz.py @@ -217,6 +217,29 @@ def test_column_layout_and_types(request, fixture, extra): assert all(hasattr(dtype, "mtztype") for dtype in frame.dtypes) +def test_map_coefficients_are_grouped_into_described_datasets(two_moment_all_mtz): + """``FWT``/``PHWT`` keep the label Coot auto-opens but sit in the + ``extrapolated_light`` dataset, and every dataset has a history line saying what it + holds -- the file, not the label, says which map a standard name is.""" + import gemmi + + mtz = gemmi.read_mtz_file(str(two_moment_all_mtz[0])) + names = {ds.id: ds.dataset_name for ds in mtz.datasets} + where = {col.label: names[col.dataset_id] for col in mtz.columns} + + assert where["FWT"] == where["PHWT"] == where["FEXT"] == "extrapolated_light" + assert where["FEXT_PHASED"] == where["FEXT_SCALAR"] == "extrapolated_light" + assert where["DF"] == where["PHDELWT"] == where["W_SD"] == "difference" + assert where["FC"] == where["PHIC"] == "light_model" + assert where["DF_corr"] == "two_moment" + assert where["H"] == where["Fo_dark"] == where["FreeR_flag_dark"] == "observed" + assert all(ds.crystal_name == "torchref" for ds in mtz.datasets) + + assert all(len(line) <= 80 for line in mtz.history) + for name in set(names.values()): + assert any(line.startswith(f"{name}:") for line in mtz.history), name + + class TestDefaultLayout: def test_the_difference_columns_carry_df_and_the_registered_weights( diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index d35caf15..a6b36909 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -848,6 +848,73 @@ def _np(t): return columns, types, diagnostics +_DIFFERENCE_DATASET_COLUMNS = ( + "DF", "SIGDF", "PHDELWT", "KSCALE", *WEIGHT_COLUMNS.values() +) + +# One history line per MTZ dataset, in the order they are written. MTZ history lines +# are at most 80 characters. +_MTZ_DATASET_HISTORY = { + "observed": "observed: Fo_dark, Fo_light and flags on the shared scale; Fc_dark", + "difference": "difference: DF/SIGDF on dark phases PHDELWT; weights W_SD, W_IVW", + "light_model": ( + "light_model: FC/PHIC, amplitude and phase of the mixed dark+light model" + ), + "extrapolated_light": ( + "extrapolated_light: FWT/PHWT = 2*FEXT - Fc, the extrapolated light map" + ), + "two_moment": "two_moment: difference columns with the two-moment correction", +} + + +def _annotate_mtz(filename, datasets): + """Group the columns of a written MTZ into named datasets and describe them. + + The column labels are left alone -- Coot's auto-open looks for ``FWT``/``PHWT`` by + label and ignores the dataset -- so the grouping is what tells a reader which map a + standard label belongs to: the column chooser shows ``/torchref//FWT``, + and ``gemmi mtz`` or ``mtzdump`` print the history lines. + + Parameters + ---------- + filename : str + MTZ written by reciprocalspaceship, with all columns in one dataset. Rewritten + in place. + datasets : dict[str, str] + Column label to dataset name. Unlisted columns, Miller indices included, stay + in ``observed``. + """ + import gemmi + + mtz = gemmi.read_mtz_file(filename) + base = mtz.datasets[0] + base.project_name = "torchref" + base.crystal_name = "torchref" + base.dataset_name = "observed" + + for name in dict.fromkeys(datasets.values()): + mtz.add_dataset(name) + # Filled in afterwards: ``add_dataset`` returns a reference into a vector that the + # next call may reallocate, so writes through it can be lost. + base = mtz.datasets[0] + ids = {} + for ds in mtz.datasets: + ds.project_name = base.project_name + ds.crystal_name = base.crystal_name + ds.cell = base.cell + ds.wavelength = base.wavelength + ids[ds.dataset_name] = ds.id + for col in mtz.columns: + if col.label in datasets: + col.dataset_id = ids[datasets[col.label]] + + mtz.history = [ + "torchref difference-refine map coefficients, grouped by dataset:", + *(_MTZ_DATASET_HISTORY[name] for name in ids), + ] + mtz.write_to_file(filename) + + def write_results_mtz( dc, dark_model, @@ -876,6 +943,12 @@ def write_results_mtz( alternatives within each layer -- see :func:`_phasing_columns` and :func:`_extrapolation_columns` for what each contains and why it is gated. + Column labels are the standard CCP4 ones, so Coot auto-opens ``FWT``/``PHWT`` -- + here the *extrapolated light-state* map, not a ``2mFo-DFc``. What each label means + is recorded in the file: the columns are grouped into MTZ datasets (``observed``, + ``difference``, ``light_model``, ``extrapolated_light``, ``two_moment``) with one + history line describing each. + Parameters ---------- dc : DatasetCollection @@ -1015,6 +1088,7 @@ def write_results_mtz( weight_columns=weight_columns, kscale=kscale, ) + datasets = {name: "difference" for name in _DIFFERENCE_DATASET_COLUMNS} if mc is not None: phase_cols, phase_types, ctx = _phasing_columns( @@ -1024,6 +1098,7 @@ def write_results_mtz( Fcalc_dark=Fcalc_dark, weights=weights, all_columns=all_columns, ) columns.update(phase_cols) + datasets.update(dict.fromkeys(phase_cols, "light_model")) types.update(phase_types) ext_cols, ext_types, ext_diagnostics = _extrapolation_columns( @@ -1043,6 +1118,7 @@ def write_results_mtz( ) diagnostics.update(ext_diagnostics) columns.update(ext_cols) + datasets.update(dict.fromkeys(ext_cols, "extrapolated_light")) types.update(ext_types) tm_cols, tm_types = _two_moment_columns( @@ -1054,6 +1130,7 @@ def write_results_mtz( all_columns=all_columns, ) columns.update(tm_cols) + datasets.update(dict.fromkeys(tm_cols, "two_moment")) types.update(tm_types) df = rs.DataSet( @@ -1071,9 +1148,12 @@ def write_results_mtz( df = df.infer_mtz_dtypes() df.set_index(["H", "K", "L"], inplace=True) df.write_mtz(filename) + _annotate_mtz(filename, datasets) if verbose > 0: print(f" Results MTZ written to {filename} ({len(columns)} columns)") + for name in dict.fromkeys(["observed", *datasets.values()]): + print(f" {_MTZ_DATASET_HISTORY[name]}") if mc is not None: fractions = mc["light"].fractions.detach() print(f" w_dark={fractions[0].item():.3f}, " From 7afb608c1f1313efe129fd5f94b8bb1015f29c0e Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 23 Sep 2026 15:10:29 +0200 Subject: [PATCH 182/250] Speed up canonicalize_hkl and add sort=False Map reflections onto the CCP4 reciprocal ASU with threaded torch operations instead of single-threaded numpy, reuse the rotated indices for the Friedel mate, and skip the phase-shift arithmetic in groups without translations. sort=False keeps the input row order and returns None for sort_indices, for callers that only need the per-row mapping. Default outputs are unchanged: canonical indices, phase shifts, Friedel flags and sort order are identical to dev on 2M random reflections in 23 space groups. New tests pin the mapping against gemmi's ReciprocalAsu.to_asu in one group per ASU condition and sort=False against the sorted output. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 1 + tests/unit/symmetry/test_canonicalize_hkl.py | 36 +++- torchref/symmetry/reciprocal_symmetry.py | 179 ++++++++++--------- torchref/symmetry/spacegroup.py | 21 ++- 4 files changed, 148 insertions(+), 89 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 48350ebb..a5cce2f4 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- ``SpaceGroup.canonicalize_hkl`` gains ``sort=False``, which keeps the input row order and returns ``None`` for ``sort_indices``; the ASU mapping runs as threaded torch operations and reuses the rotated indices for the Friedel mate, so large reflection lists map faster with unchanged outputs - ``torchref.difference-map``, ``torchref.difference-refine`` and ``torchref.validate-ded`` gain ``--ded-weight {sigma_d,inverse_variance,none}`` and ``--sigma-d-gamma``. The difference MTZ now carries the unweighted ``DF``/``SIGDF`` on ``PHDELWT`` with one mean-one weight column per scheme, ``W_SD`` and ``W_IVW`` (MTZ type W), and the observed-to-model scale ``KSCALE``; ``DELFWT`` is no longer written, build the map with ``torchref.mtz2map -csf DF -cw W_IVW -cphi PHDELWT``. Registered in ``torchref.maps.ded_weights`` - Added the ``sigma_D`` estimator (``torchref.refinement.model_error_estimation.sigma_d``): the expected true difference power per resolution shell, ``mean(dF_obs^2) - mean(sigma^2)`` with a fitted ``F_dark^gamma`` amplitude law and DerSimonian-Laird shrinkage of the signed shell power toward a decaying exponential in ``d*^2`` (fitted on all shells, so a dataset without a difference yields no power instead of the positive half of its noise), giving the Wiener weight ``S/(S + sigma^2)`` and, with a difference model, ``alpha``/``beta_model``. Inverse-variance weights suppress the strong reflections whose difference power is 10-70x that of weak ones; on independent half-datasets a Wiener weight with the true power raised map agreement 1.2-1.8x in effective patterns. The single-dataset estimate inherits the calibration of the reported sigmas, and on the campaign TD1 data (sigmas ~1.5x too large at high resolution) it emptied 60-90 % of the shells, so inverse variance stays the default; ``sigma_d`` reports its clamped-shell count and falls back to inverse variance with a warning when every shell is empty - Added the ``difference_sd`` collection target (``CollectionDifferenceSigmaDTarget``): the difference Gaussian centred on ``alpha * dF_calc`` with variance ``beta_model + sigma_diff^2`` from ``sigma_D`` fitted on the free set. Selected with ``torchref.difference-refine --difference-target difference_sd``; ``difference`` stays the default diff --git a/tests/unit/symmetry/test_canonicalize_hkl.py b/tests/unit/symmetry/test_canonicalize_hkl.py index 1764baf6..95b6ff0f 100644 --- a/tests/unit/symmetry/test_canonicalize_hkl.py +++ b/tests/unit/symmetry/test_canonicalize_hkl.py @@ -10,9 +10,9 @@ # The HKL verbs live on the space group now. These adapters keep the assertions # below -- which pin the phase-sign contract -- expressed in terms of the space # group specifications the cases are parametrised over. -def canonicalize_hkl(hkl, sg, include_friedel=True, device=None): +def canonicalize_hkl(hkl, sg, include_friedel=True, device=None, sort=True): return SpaceGroup(sg).canonicalize_hkl( - hkl, include_friedel=include_friedel, device=device + hkl, include_friedel=include_friedel, device=device, sort=sort ) @@ -259,6 +259,38 @@ def test_unmappable_reflection_raises(self): with pytest.raises(ValueError, match="could not map"): canonicalize_hkl(hkl, "P1", include_friedel=False) + # One group per CCP4 reciprocal-ASU condition, plus centred settings. + ASU_GROUPS = ["P1", "P21", "C2", "P212121", "I222", "P4", "P41212", "I41/a", + "P3", "P3121", "P3112", "R3", "P6", "P63", "P6122", "P23", "I23", + "P432", "Fm-3m"] + + @pytest.mark.parametrize("sg", ASU_GROUPS) + def test_matches_gemmi_asu(self, sg): + """Canonical indices agree with gemmi's own ASU mapping, row by row.""" + import gemmi + + g = torch.Generator().manual_seed(0) + hkl = torch.randint(-6, 7, (300, 3), generator=g, dtype=torch.int32) + can, _, _, _ = canonicalize_hkl(hkl, sg, sort=False) + group = gemmi.SpaceGroup(SpaceGroup(sg)._gemmi.xhm()) + asu, ops = gemmi.ReciprocalAsu(group), group.operations() + want = torch.tensor( + [asu.to_asu(row, ops)[0] for row in hkl.tolist()], dtype=torch.int32 + ) + assert torch.equal(can, want) + + @pytest.mark.parametrize("sg", ["P1", "P21", "C2", "P43212", "P63", "R3", "I23"]) + def test_unsorted_is_sorted_output_in_input_order(self, sg): + """``sort=False`` returns the sorted outputs un-permuted, and no permutation.""" + g = torch.Generator().manual_seed(1) + hkl = torch.randint(-8, 9, (2000, 3), generator=g, dtype=torch.int32) + can_s, ps_s, ff_s, si = canonicalize_hkl(hkl, sg) + can_u, ps_u, ff_u, none = canonicalize_hkl(hkl, sg, sort=False) + assert none is None + assert torch.equal(can_u[si], can_s) + assert torch.equal(ff_u[si], ff_s) + assert torch.equal(ps_u[si], ps_s) + def test_empty_input(self): """Empty input should return empty tensors without error.""" hkl = torch.empty((0, 3), dtype=torch.int32) diff --git a/torchref/symmetry/reciprocal_symmetry.py b/torchref/symmetry/reciprocal_symmetry.py index f595ee70..8e23a57f 100644 --- a/torchref/symmetry/reciprocal_symmetry.py +++ b/torchref/symmetry/reciprocal_symmetry.py @@ -19,6 +19,7 @@ ``tests/unit/symmetry/test_phase_convention.py``. """ +import math from typing import Optional, Tuple import numpy as np @@ -399,7 +400,8 @@ def _canonicalize_hkl( hkl: torch.Tensor, include_friedel: bool = True, device: Optional[torch.device] = None, -) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + sort: bool = True, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: """Map Miller indices to canonical CCP4 ASU representatives. Selects one representative per reflection under the standard CCP4 asymmetric @@ -418,17 +420,20 @@ def _canonicalize_hkl( Whether Friedel mates are considered equivalent. device : torch.device, optional Computation device. If None, uses hkl's device. + sort : bool, default True + Return the rows sorted lexicographically by canonical (h, k, l). With + ``False`` the rows stay in input order and no permutation is formed. Returns ------- canonical_hkl : torch.Tensor, shape (N, 3), dtype int32 - Remapped indices, sorted lexicographically by (h, k, l). + Remapped indices, sorted lexicographically by (h, k, l) when ``sort``. phase_shifts : torch.Tensor, shape (N,), dtype float32 - Additive phase correction in radians. + Additive phase correction in radians, in the same row order. friedel_flags : torch.Tensor, shape (N,), dtype bool - True where Friedel conjugation was applied. - sort_indices : torch.Tensor, shape (N,), dtype int64 - Permutation from original to sorted order. + True where Friedel conjugation was applied, in the same row order. + sort_indices : torch.Tensor or None, shape (N,), dtype int64 + Permutation from original to sorted order; ``None`` when ``sort=False``. Notes ----- @@ -446,82 +451,89 @@ def _canonicalize_hkl( empty_hkl = torch.empty((0, 3), dtype=hkl_dtype, device=device) empty_f = torch.empty(0, dtype=get_float_dtype(), device=device) empty_b = torch.empty(0, dtype=torch.bool, device=device) - empty_i = torch.empty(0, dtype=torch.int64, device=device) # dtype-ok: empty index tensor; int64 index dtype required + empty_i = torch.empty(0, dtype=torch.int64, device=device) if sort else None # dtype-ok: empty index tensor; int64 index dtype required return empty_hkl, empty_f, empty_b, empty_i - # The ASU lookup tables are numpy-backed, so the operations come across to CPU - # regardless of where ``sym`` lives; only the returned tensors honour ``device``. + # The mapping runs on CPU whatever device ``sym`` or ``hkl`` live on (gemmi's + # scalar ASU test is the fallback); only the returned tensors honour ``device``. + # Torch rather than numpy for the per-row arithmetic: the work is a handful of + # elementwise passes over every reflection, which torch spreads over threads. asu = gemmi.ReciprocalAsu(sym._gemmi) condition_key = asu.condition_str() - recip_mats = sym.reciprocal.matrices.detach().cpu().numpy() # (n_ops, 3, 3) - translations_np = sym.translations.detach().cpu().numpy() # (n_ops, 3) - n_ops = len(recip_mats) - - hkl_np = hkl.cpu().numpy().astype(np.int32) # (N, 3) # Reciprocal-space rotation matrices are always integer-valued (0, ±1). - recip_mats_i = np.round(recip_mats).astype(np.int32) - - # One op (+ its Friedel mate) at a time, so high-symmetry groups exit early: - # most reflections are resolved by the first few operators. - canonical_np = np.empty_like(hkl_np) - op_idx = np.empty(n_refl, dtype=np.int32) - friedel_np = np.zeros(n_refl, dtype=bool) - remaining = np.ones(n_refl, dtype=bool) + recip_ops = torch.round(sym.reciprocal.matrices.detach().cpu()).to(torch.int32) + translations = sym.translations.detach().cpu() # (n_ops, 3) + n_ops = len(recip_ops) + hkl_cpu = hkl.detach().to(device="cpu", dtype=torch.int32) # (N, 3) - for i_op in range(n_ops): - if not remaining.any(): - break - idx = np.where(remaining)[0] - hkl_sub = hkl_np[idx] # (M, 3) - R = recip_mats_i[i_op] # (3, 3) - equiv_sub = hkl_sub @ R.T # (M, 3), int32 matmul — no rounding needed - - # Check non-Friedel - h, k, l = equiv_sub[:, 0], equiv_sub[:, 1], equiv_sub[:, 2] + def in_asu(h, k, l): try: - in_asu_pos = _asu_condition_vectorized(h, k, l, condition_key) + return _asu_condition_vectorized(h, k, l, condition_key) except ValueError: - in_asu_pos = np.array( - [asu.is_in(row.tolist()) for row in equiv_sub], dtype=bool + return torch.tensor( + [ + asu.is_in([a, b, c]) + for a, b, c in zip(h.tolist(), k.tolist(), l.tolist()) + ], + dtype=torch.bool, ) - hit_pos = np.where(in_asu_pos)[0] - if len(hit_pos) > 0: - global_idx = idx[hit_pos] - canonical_np[global_idx] = equiv_sub[hit_pos] - op_idx[global_idx] = i_op - remaining[global_idx] = False - - # Check Friedel mate - if include_friedel and remaining.any(): - # Recompute idx for remaining after non-Friedel hits - idx_f = np.where(remaining)[0] - hkl_sub_f = hkl_np[idx_f] - equiv_neg = -(hkl_sub_f @ R.T) - - h_n, k_n, l_n = equiv_neg[:, 0], equiv_neg[:, 1], equiv_neg[:, 2] - try: - in_asu_neg = _asu_condition_vectorized(h_n, k_n, l_n, condition_key) - except ValueError: - in_asu_neg = np.array( - [asu.is_in(row.tolist()) for row in equiv_neg], dtype=bool - ) + def rotate(row, h, k, l): + """``row . (h, k, l)`` for one row of an integer (0, ±1) rotation.""" + out = None + for coef, col in zip(row, (h, k, l)): + if coef == 0: + continue + term = col if coef == 1 else -col if coef == -1 else col * coef + out = term if out is None else out + term + return torch.zeros_like(h) if out is None else out - hit_neg = np.where(in_asu_neg)[0] - if len(hit_neg) > 0: - global_idx_f = idx_f[hit_neg] - canonical_np[global_idx_f] = equiv_neg[hit_neg] - op_idx[global_idx_f] = i_op - friedel_np[global_idx_f] = True - remaining[global_idx_f] = False + # One op (+ its Friedel mate) at a time, so high-symmetry groups exit early: + # most reflections are resolved by the first few operators. ``todo`` holds the + # still-unmapped rows in increasing order (``None`` while that is all of them). + canonical = torch.empty_like(hkl_cpu) + op_idx = torch.empty(n_refl, dtype=torch.int16) + friedel = torch.zeros(n_refl, dtype=torch.bool) + todo = None + + for i_op in range(n_ops): + if todo is not None and todo.numel() == 0: + break + R = recip_ops[i_op].tolist() + sub = hkl_cpu if todo is None else hkl_cpu.index_select(0, todo) + h, k, l = sub.unbind(1) + eh, ek, el = (rotate(R[i], h, k, l) for i in range(3)) + + hit = in_asu(eh, ek, el) + miss = ~hit + rows = hit.nonzero().squeeze(1) + left = miss.nonzero().squeeze(1) + if todo is not None: + rows, left = todo[rows], todo[left] + if rows.numel(): + canonical.index_copy_(0, rows, torch.stack((eh[hit], ek[hit], el[hit]), 1)) + op_idx.index_fill_(0, rows, i_op) + todo = left + + # The Friedel mate of R h is -(R h): reuse the rotated indices. + if include_friedel and todo.numel(): + nh, nk, nl = -eh[miss], -ek[miss], -el[miss] + hit_f = in_asu(nh, nk, nl) + rows = todo[hit_f] + if rows.numel(): + canonical.index_copy_( + 0, rows, torch.stack((nh[hit_f], nk[hit_f], nl[hit_f]), 1) + ) + op_idx.index_fill_(0, rows, i_op) + friedel.index_fill_(0, rows, True) + todo = todo[~hit_f] - # ``canonical_np``/``op_idx`` are uninitialized ``np.empty`` buffers, so an + # ``canonical``/``op_idx`` are uninitialized ``torch.empty`` buffers, so an # unmapped row would propagate garbage indices and phases. Fail loudly instead. - if remaining.any(): - n_unmapped = int(remaining.sum()) - example = hkl_np[np.where(remaining)[0][0]].tolist() + if todo is not None and todo.numel(): + example = hkl_cpu[todo[0]].tolist() raise ValueError( - f"canonicalize_hkl could not map {n_unmapped} reflection(s) to the " + f"canonicalize_hkl could not map {todo.numel()} reflection(s) to the " f"reciprocal ASU of space group {sym} " f"(include_friedel={include_friedel}); e.g. hkl={example}. With " f"include_friedel=False the Friedel half of reciprocal space has no " @@ -532,19 +544,24 @@ def _canonicalize_hkl( # already negated phi for those rows: -2π h·t normally, +2π h·t for Friedel. # A single uniform sign is wrong for one half and invisible in P21/P212121/C2, # where every shift is 0 or π. tests/unit/symmetry/test_phase_convention.py. - t_selected = translations_np[op_idx] # (N, 3) - friedel_sign = np.where(friedel_np, 1.0, -1.0).astype(np.float32) - phase_shifts_np = ( - friedel_sign - * 2.0 - * np.pi - * np.sum(hkl_np.astype(np.float32) * t_selected, axis=1) - ).astype(np.float32) - - # --- Convert to tensors and sort --- - canonical_hkl = torch.tensor(canonical_np, dtype=hkl_dtype, device=device) - phase_shifts = torch.tensor(phase_shifts_np, dtype=get_float_dtype(), device=device) - friedel_flags = torch.tensor(friedel_np, dtype=torch.bool, device=device) + # h·t is summed left to right so the value does not depend on a backend's + # reduction order; the shift is rounded to float32 like the rest of the output. + if bool(translations.any()): + t_sel = translations.index_select(0, op_idx.long()) + hf = hkl_cpu.to(torch.float32) + h_dot_t = ( + hf[:, 0] * t_sel[:, 0] + hf[:, 1] * t_sel[:, 1] + hf[:, 2] * t_sel[:, 2] + ) + friedel_sign = torch.where(friedel, 1.0, -1.0).to(torch.float32) + phase = (friedel_sign * 2.0 * math.pi * h_dot_t).to(torch.float32) + else: + phase = torch.zeros(n_refl, dtype=torch.float32) + + canonical_hkl = canonical.to(dtype=hkl_dtype, device=device) + phase_shifts = phase.to(dtype=get_float_dtype(), device=device) + friedel_flags = friedel.to(device=device) + if not sort: + return canonical_hkl, phase_shifts, friedel_flags, None # Lexicographic sort by (h, k, l) via composite key h_max = int(canonical_hkl.abs().max().item()) + 1 diff --git a/torchref/symmetry/spacegroup.py b/torchref/symmetry/spacegroup.py index e7c7d5a4..c7a9354e 100644 --- a/torchref/symmetry/spacegroup.py +++ b/torchref/symmetry/spacegroup.py @@ -422,6 +422,8 @@ def canonicalize_hkl( hkl: torch.Tensor, include_friedel: bool = True, device: Optional[torch.device] = None, + *, + sort: bool = True, ): """Map Miller indices onto their canonical CCP4 ASU representatives. @@ -436,17 +438,24 @@ def canonicalize_hkl( device : torch.device, optional Output device. Defaults to ``hkl``'s. The lookup itself runs on CPU whatever device this group is on, because the ASU tables are numpy-backed. + sort : bool, default True + Sort the rows lexicographically by canonical ``(h, k, l)``. ``False`` + keeps the input row order and skips the sort, which callers that only + need the per-row mapping should prefer on large inputs. Returns ------- canonical_hkl : torch.Tensor - Remapped indices sorted lexicographically, shape ``(N, 3)``. + Remapped indices, shape ``(N, 3)``; sorted lexicographically when + ``sort``. phase_shifts : torch.Tensor - Additive phase correction in radians, shape ``(N,)``. + Additive phase correction in radians, shape ``(N,)``, same row order. friedel_flags : torch.Tensor - Boolean, shape ``(N,)``, True where Friedel conjugation was applied. - sort_indices : torch.Tensor - Permutation from original to sorted order, shape ``(N,)``. + Boolean, shape ``(N,)``, True where Friedel conjugation was applied, + same row order. + sort_indices : torch.Tensor or None + Permutation from original to sorted order, shape ``(N,)``; ``None`` + when ``sort=False``. Notes ----- @@ -456,7 +465,7 @@ def canonicalize_hkl( from torchref.symmetry.reciprocal_symmetry import _canonicalize_hkl return _canonicalize_hkl( - self, hkl, include_friedel=include_friedel, device=device + self, hkl, include_friedel=include_friedel, device=device, sort=sort ) # ========================================================================= From 32911445fb348208c246d7ff0f82249da98339ac Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sun, 27 Sep 2026 11:05:36 +0200 Subject: [PATCH 183/250] Drop numba: the restraint matchers are faster as plain Python The four match_*_numba functions in topology/builders_numba.py were the only numba code in TorchRef. Each call handles a single residue, so the JIT's dispatch overhead outweighed the loop it compiled. Building restraints for 2DQ6 (6,933 atoms) spent 0.62 s in the matchers with a warm numba cache and 0.39 s as plain Python; a cold cache added 13.5 s of compilation. - builders_numba.py -> matchers.py and match_*_numba -> match_*, with the function bodies unchanged. - numba removed from the dependencies, the tox environments, the CI workflows and Dependabot, and from the requirement lists in the README and docs. Checked on 11 AlphaFold-start structures, each with and without --add-hydrogens: every topology and restraint-value tensor is bitwise identical to what the numba version builds. 10-cycle refinements (hydrogens off and riding) agree with the old code within its own run-to-run spread. The CPU suite fails the same 7 tests as before this change: the five af_trajectory references, a strict xfail that now passes, and the dtype conformance scan. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01KwtigurBqaYuzV6426n5XG --- .github/dependabot.yml | 2 - .github/workflows/accelerator.yml | 12 ++--- .github/workflows/ci.yml | 4 -- .github/workflows/compatibility.yml | 8 +-- .github/workflows/dev-pr.yml | 1 - README.md | 2 +- docs/installation.rst | 2 +- docs/user_guide/testing.rst | 4 +- pyproject.toml | 4 +- .../unit/io/test_multicomponent_restraints.py | 2 +- torchref/topology/build.py | 22 ++++---- torchref/topology/builders.py | 25 +++++---- .../{builders_numba.py => matchers.py} | 54 ++++--------------- tox.ini | 41 -------------- 14 files changed, 48 insertions(+), 135 deletions(-) rename torchref/topology/{builders_numba.py => matchers.py} (85%) diff --git a/.github/dependabot.yml b/.github/dependabot.yml index a24e37b4..672ccd02 100644 --- a/.github/dependabot.yml +++ b/.github/dependabot.yml @@ -29,8 +29,6 @@ updates: update-types: ["version-update:semver-major"] - dependency-name: "torch" update-types: ["version-update:semver-major"] - - dependency-name: "numba" - update-types: ["version-update:semver-major"] # GitHub Actions - package-ecosystem: "github-actions" diff --git a/.github/workflows/accelerator.yml b/.github/workflows/accelerator.yml index 96d16fef..e8eb3e4c 100644 --- a/.github/workflows/accelerator.yml +++ b/.github/workflows/accelerator.yml @@ -20,13 +20,11 @@ on: # Lets publish.yml gate a PyPI release on this passing. workflow_call: -env: - NUMBA_CACHE_DIR: /tmp/numba_cache - # PYTORCH_ENABLE_MPS_FALLBACK is left alone on purpose. torchref/__init__.py - # setdefault()s it to 1, so CI runs with the same fallback behaviour users - # get. Setting it to 0 here would make an op with no Metal kernel a CI - # failure while it stays a silent CPU round-trip in production -- a stricter - # policy than the package's, and not one to introduce from a workflow file. +# PYTORCH_ENABLE_MPS_FALLBACK is left alone on purpose. torchref/__init__.py +# setdefault()s it to 1, so CI runs with the same fallback behaviour users +# get. Setting it to 0 here would make an op with no Metal kernel a CI +# failure while it stays a silent CPU round-trip in production -- a stricter +# policy than the package's, and not one to introduce from a workflow file. jobs: mps: diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 094dd811..ddd37f6d 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -6,9 +6,6 @@ on: pull_request: branches: [main, master, develop] -env: - NUMBA_CACHE_DIR: /tmp/numba_cache - jobs: test: name: Test - Python ${{ matrix.python-version }} @@ -45,7 +42,6 @@ jobs: python --version python -c "import numpy; print(f'numpy: {numpy.__version__}')" python -c "import torch; print(f'torch: {torch.__version__}')" - python -c "import numba; print(f'numba: {numba.__version__}')" - name: Run tests run: | diff --git a/.github/workflows/compatibility.yml b/.github/workflows/compatibility.yml index 7c46cccc..480b5ffe 100644 --- a/.github/workflows/compatibility.yml +++ b/.github/workflows/compatibility.yml @@ -27,10 +27,6 @@ on: - 'tox.ini' - '.github/workflows/compatibility.yml' -env: - # Avoid numba caching issues in CI - NUMBA_CACHE_DIR: /tmp/numba_cache - jobs: # ========================================================================== # Test with bounded versions (as specified in pyproject.toml) @@ -63,7 +59,6 @@ jobs: python -c "import numpy; print(f'numpy: {numpy.__version__}')" python -c "import pandas; print(f'pandas: {pandas.__version__}')" python -c "import torch; print(f'torch: {torch.__version__}')" - python -c "import numba; print(f'numba: {numba.__version__}')" python -c "import scipy; print(f'scipy: {scipy.__version__}')" python -c "import gemmi; print(f'gemmi: {gemmi.__version__}')" python -c "import reciprocalspaceship; print(f'reciprocalspaceship: {reciprocalspaceship.__version__}')" @@ -97,7 +92,7 @@ jobs: python -m pip install --upgrade pip pip install pytest pytest-cov # Install dependencies without upper bounds - pip install numpy pandas torch numba scipy matplotlib pyarrow tqdm + pip install numpy pandas torch scipy matplotlib pyarrow tqdm pip install gemmi reciprocalspaceship # Install package in editable mode without deps (already installed) pip install -e . --no-deps @@ -108,7 +103,6 @@ jobs: python -c "import numpy; print(f'numpy: {numpy.__version__}')" python -c "import pandas; print(f'pandas: {pandas.__version__}')" python -c "import torch; print(f'torch: {torch.__version__}')" - python -c "import numba; print(f'numba: {numba.__version__}')" python -c "import scipy; print(f'scipy: {scipy.__version__}')" python -c "import gemmi; print(f'gemmi: {gemmi.__version__}')" python -c "import reciprocalspaceship; print(f'reciprocalspaceship: {reciprocalspaceship.__version__}')" diff --git a/.github/workflows/dev-pr.yml b/.github/workflows/dev-pr.yml index 33131672..ea96f29e 100644 --- a/.github/workflows/dev-pr.yml +++ b/.github/workflows/dev-pr.yml @@ -18,7 +18,6 @@ jobs: timeout-minutes: 90 env: TORCHREF_DEVICE: cpu - NUMBA_CACHE_DIR: /tmp/numba_cache steps: - uses: actions/checkout@v7 diff --git a/README.md b/README.md index def677b5..f77689af 100644 --- a/README.md +++ b/README.md @@ -73,7 +73,7 @@ the checkout are fetched on demand, so add paths later with `git sparse-checkout ### Dependencies -Python ≥ 3.10, PyTorch ≥ 2.4, NumPy ≥ 2.0, Pandas ≥ 2.0, SciPy ≥ 1.10, Gemmi ≥ 0.5, reciprocalspaceship ≥ 0.9.18, Numba ≥ 0.59, Matplotlib ≥ 3.7. `pyproject.toml` carries the authoritative pinned ranges; upper bounds are set one minor version above the tested maximum, so a newer dependency will refuse to install rather than fail at runtime. +Python ≥ 3.10, PyTorch ≥ 2.4, NumPy ≥ 2.0, Pandas ≥ 2.0, SciPy ≥ 1.10, Gemmi ≥ 0.5, reciprocalspaceship ≥ 0.9.18, Matplotlib ≥ 3.7. `pyproject.toml` carries the authoritative pinned ranges; upper bounds are set one minor version above the tested maximum, so a newer dependency will refuse to install rather than fail at runtime. ### Testing diff --git a/docs/installation.rst b/docs/installation.rst index c10adc4f..1004f734 100644 --- a/docs/installation.rst +++ b/docs/installation.rst @@ -5,7 +5,7 @@ Requirements ------------ Python ≥ 3.10, PyTorch ≥ 2.4, NumPy ≥ 2.0, Pandas ≥ 2.0, SciPy ≥ 1.10, -Gemmi ≥ 0.5, reciprocalspaceship ≥ 0.9.18, Numba ≥ 0.59, Matplotlib ≥ 3.7. +Gemmi ≥ 0.5, reciprocalspaceship ≥ 0.9.18, Matplotlib ≥ 3.7. ``pyproject.toml`` carries the authoritative pinned ranges. Upper bounds are set one minor version above the tested maximum, so an untested dependency version diff --git a/docs/user_guide/testing.rst b/docs/user_guide/testing.rst index 66e60273..cca19834 100644 --- a/docs/user_guide/testing.rst +++ b/docs/user_guide/testing.rst @@ -177,8 +177,8 @@ CI -- ``tox.ini`` defines the environments: ``py310``–``py313`` against current -dependencies, plus boundary environments pinning NumPy, Numba, PyTorch, Pandas -and Gemmi versions and a ``lowerbounds`` pair at the declared minimums. Refer to +dependencies, plus boundary environments pinning NumPy, PyTorch, Pandas and +Gemmi versions and a ``lowerbounds`` pair at the declared minimums. Refer to the file itself for the authoritative list. .. code-block:: bash diff --git a/pyproject.toml b/pyproject.toml index 4104a2e4..e8f539ea 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -28,12 +28,14 @@ authors = [ # Both runs give an identical 1625 passed / 96 skipped, so the bumps are behaviour-neutral # on CPU. The CUDA/Triton and MPS kernels were NOT exercised -- the validation host had no # accelerator, so those paths are covered only by the GPU CI runners. +# +# 2026-09-27: numba dropped. It only ran the per-residue restraint matchers +# (torchref/topology/matchers.py), which are faster as plain Python. dependencies = [ "numpy>=2.0.0,<2.4.0", "pandas>=2.0.0,<2.4.0", "torch>=2.4.0,<2.14.0", "tqdm>=4.61.0,<4.69.0", - "numba>=0.59.0,<0.67.0", "gemmi>=0.5.0,<0.8.0", "scipy>=1.10.0,<1.18.0", "matplotlib>=3.7.0,<3.11.0", diff --git a/tests/unit/io/test_multicomponent_restraints.py b/tests/unit/io/test_multicomponent_restraints.py index 65d344ab..b09bc069 100644 --- a/tests/unit/io/test_multicomponent_restraints.py +++ b/tests/unit/io/test_multicomponent_restraints.py @@ -243,7 +243,7 @@ def test_short_spellings_are_not_dropped(self): ) signs = PreprocessedCIF({})._preprocess_chirals(chirals)["volume_sign"] - # NaN here is not a rounding detail: builders_numba skips those rows, so + # NaN here is not a rounding detail: match_chirals skips those rows, so # an unrecognised spelling deletes the restraint outright. assert not np.isnan(signs).any() assert signs.tolist() == [1.0, 1.0, -1.0, 0.0] diff --git a/torchref/topology/build.py b/torchref/topology/build.py index 87970297..ff9452a6 100644 --- a/torchref/topology/build.py +++ b/torchref/topology/build.py @@ -1,7 +1,7 @@ """Assemble a :class:`~torchref.topology.topology.Topology` from an atom table. -Intra-residue edges are matched here, template by template, through the Numba matchers -in :mod:`torchref.topology.builders_numba`. Inter-residue edges come from the +Intra-residue edges are matched here, template by template, through the matchers +in :mod:`torchref.topology.matchers`. Inter-residue edges come from the ``InterResidue*Builder`` classes, which already encode the link geometry and are reused rather than reimplemented. """ @@ -19,11 +19,11 @@ InterResidueTorsionBuilder, PreprocessedCIF, ) -from torchref.topology.builders_numba import ( - match_angles_numba, - match_bonds_numba, - match_chirals_numba, - match_torsions_numba, +from torchref.topology.matchers import ( + match_angles, + match_bonds, + match_chirals, + match_torsions, ) from torchref.topology.atom_graph import AtomGraph from torchref.topology.edges import EdgeBlock, assemble_origins @@ -205,7 +205,7 @@ def _match_intra( for names, indices in _conformers(cols, start, end): if key in pp_cif.bonds: b = pp_cif.bonds[key] - n = match_bonds_numba( + n = match_bonds( names, indices, b["atom1"], @@ -225,7 +225,7 @@ def _match_intra( val["bonds"]["sigmas"].append(work["f2"][:n].copy()) if key in pp_cif.angles: a = pp_cif.angles[key] - n = match_angles_numba( + n = match_angles( names, indices, a["atom1"], @@ -253,7 +253,7 @@ def _match_intra( val["angles"]["sigmas"].append(work["f2"][:n].copy()) if key in pp_cif.torsions: t = pp_cif.torsions[key] - n = match_torsions_numba( + n = match_torsions( names, indices, t["atom1"], @@ -287,7 +287,7 @@ def _match_intra( val["torsions"]["periods"].append(work["per"][:n].copy()) if key in pp_cif.chirals: c = pp_cif.chirals[key] - n = match_chirals_numba( + n = match_chirals( names, indices, c["center"], diff --git a/torchref/topology/builders.py b/torchref/topology/builders.py index 14502625..2cbab1d1 100644 --- a/torchref/topology/builders.py +++ b/torchref/topology/builders.py @@ -21,12 +21,12 @@ from torchref.config import get_float_dtype, get_int_dtype -# Import the Numba-accelerated matching functions -from torchref.topology.builders_numba import ( - match_angles_numba, - match_bonds_numba, - match_chirals_numba, - match_torsions_numba, +# Intra-residue restraint matchers +from torchref.topology.matchers import ( + match_angles, + match_bonds, + match_chirals, + match_torsions, ) @@ -349,7 +349,7 @@ def _preprocess_chirals(self, chirals_df: pd.DataFrame) -> Dict[str, np.ndarray] """Convert chirals DataFrame to NumPy arrays.""" # Convert volume_sign strings to floats. The CCP4 library writes both the # full and the truncated spelling ("positiv", "negativ"); an unrecognised - # sign becomes NaN and the restraint is then dropped in builders_numba, + # sign becomes NaN and the restraint is then dropped by match_chirals, # so the short forms have to be matched here or those chirals vanish. volume_signs = [] for sign in chirals_df["volume_sign"].values: @@ -572,8 +572,7 @@ def build( # Iterate over altloc conformations (yields once if no altlocs) for atom_names, atom_indices, _ in pp_pdb.get_altloc_conformations(res_idx): - # Use Numba-accelerated matching - count = match_bonds_numba( + count = match_bonds( atom_names, atom_indices, cif_bonds["atom1"], @@ -676,7 +675,7 @@ def build( # Iterate over altloc conformations (yields once if no altlocs) for atom_names, atom_indices, _ in pp_pdb.get_altloc_conformations(res_idx): - count = match_angles_numba( + count = match_angles( atom_names, atom_indices, cif_angles["atom1"], @@ -789,7 +788,7 @@ def build( # Iterate over altloc conformations (yields once if no altlocs) for atom_names, atom_indices, _ in pp_pdb.get_altloc_conformations(res_idx): - count = match_torsions_numba( + count = match_torsions( atom_names, atom_indices, cif_torsions["atom1"], @@ -999,7 +998,7 @@ def build( # Iterate over altloc conformations (yields once if no altlocs) for atom_names, atom_indices, _ in pp_pdb.get_altloc_conformations(res_idx): - count = match_chirals_numba( + count = match_chirals( atom_names, atom_indices, cif_chirals["center"], @@ -1034,7 +1033,7 @@ def build( # restrains |volume| toward 2.5 (not toward a target of 0). # Note: chirals with an unknown sign were mapped to NaN by # PreprocessedCIF._preprocess_chirals and dropped upstream - # by match_chirals_numba, so they never reach here. + # by match_chirals, so they never reach here. all_ideal_volumes.append(work_signs[:count].copy() * 2.5) all_sigmas.append(work_sigmas[:count].copy()) diff --git a/torchref/topology/builders_numba.py b/torchref/topology/matchers.py similarity index 85% rename from torchref/topology/builders_numba.py rename to torchref/topology/matchers.py index bef7ebbc..aad72972 100644 --- a/torchref/topology/builders_numba.py +++ b/torchref/topology/matchers.py @@ -1,47 +1,18 @@ -"""Numba-accelerated CIF-to-restraint matchers, free of Pandas in the hot loop. +"""CIF-to-restraint matchers, free of Pandas in the hot loop. -Every ``match_*_numba`` shares one calling convention: the caller pre-allocates -the ``out_*`` arrays, the function fills entries ``[0:count]`` in place and -returns ``count``. Anything past ``count`` is stale. +Every ``match_*`` shares one calling convention: the caller pre-allocates the +``out_*`` arrays, the function fills entries ``[0:count]`` in place and returns +``count``. Anything past ``count`` is stale. -Numba is optional -- without it ``njit`` degrades to a no-op decorator and -``prange`` to ``range``, so the same code runs orders of magnitude slower rather -than failing. +These were Numba kernels. Each call covers one residue, so Numba's dispatch +overhead outweighed the loop it compiled -- plain Python is faster, and it has no +cold-cache compile (~13 s on first use per environment). """ -from typing import Any, Dict, Iterator, List, Optional, Tuple - import numpy as np -import pandas as pd -import torch - -try: - import numba - from numba import njit, prange - - HAS_NUMBA = True -except ImportError: - HAS_NUMBA = False - - # Fallback decorator that does nothing - def njit(*args, **kwargs): - def decorator(func): - return func - - if len(args) == 1 and callable(args[0]): - return args[0] - return decorator - - prange = range - - -# ============================================================================= -# Numba-accelerated helper functions -# ============================================================================= -@njit(cache=True) -def match_bonds_numba( +def match_bonds( residue_atom_names: np.ndarray, # atom names for this residue residue_atom_indices: np.ndarray, # global atom indices bond_atom1: np.ndarray, # CIF bond atom1 names @@ -100,8 +71,7 @@ def match_bonds_numba( return count -@njit(cache=True) -def match_angles_numba( +def match_angles( residue_atom_names: np.ndarray, residue_atom_indices: np.ndarray, angle_atom1: np.ndarray, @@ -158,8 +128,7 @@ def match_angles_numba( return count -@njit(cache=True) -def match_torsions_numba( +def match_torsions( residue_atom_names: np.ndarray, residue_atom_indices: np.ndarray, torsion_atom1: np.ndarray, @@ -228,8 +197,7 @@ def match_torsions_numba( return count -@njit(cache=True) -def match_chirals_numba( +def match_chirals( residue_atom_names: np.ndarray, residue_atom_indices: np.ndarray, chiral_center: np.ndarray, diff --git a/tox.ini b/tox.ini index 883126fa..4efe36d4 100644 --- a/tox.ini +++ b/tox.ini @@ -8,9 +8,6 @@ envlist = # NumPy 2.x version boundaries py311-numpy2x py312-numpy2x - # Numba version tests - py311-numba061 - py312-numba063 # Minimum viable modern stack py310-minimum py311-minimum @@ -47,7 +44,6 @@ commands = python -c "import numpy; print(f'numpy: {numpy.__version__}')" python -c "import pandas; print(f'pandas: {pandas.__version__}')" python -c "import torch; print(f'torch: {torch.__version__}')" - python -c "import numba; print(f'numba: {numba.__version__}')" python -c "import scipy; print(f'scipy: {scipy.__version__}')" python -c "import torchref; print(f'torchref: {getattr(torchref, \"__version__\", \"unknown\")}')" pytest tests/ -v --tb=short -m "not gpu and not slow" {posargs} @@ -63,7 +59,6 @@ deps = numpy pandas torch - numba scipy [testenv:py310-minimum] @@ -74,7 +69,6 @@ deps = numpy==2.0.0 pandas==2.0.0 torch==2.4.0 - numba==0.61.0 scipy==1.13.0 [testenv:py310-lowerbounds] @@ -87,7 +81,6 @@ deps = pandas==2.0.0 torch==2.4.0 tqdm==4.61.0 - numba==0.59.0 gemmi==0.5.0 scipy==1.10.0 matplotlib==3.7.0 @@ -105,7 +98,6 @@ deps = numpy pandas torch - numba scipy [testenv:py311-numpy2x] @@ -116,20 +108,8 @@ deps = numpy>=2.0.0 pandas>=2.2.0 torch>=2.4.0 - numba>=0.61.0 scipy>=1.13.0 -[testenv:py311-numba061] -description = Python 3.11 with numba 0.61 (first NumPy 2.x support) -basepython = python3.11 -deps = - {[testenv]deps} - numpy>=2.0.0,<2.1.0 - pandas>=2.1.0 - torch>=2.4.0,<2.5.0 - numba>=0.61.0,<0.62.0 - scipy>=1.13.0,<1.14.0 - [testenv:py311-minimum] description = Python 3.11 with minimum viable versions (numpy 2.0 + torch 2.4) basepython = python3.11 @@ -138,7 +118,6 @@ deps = numpy==2.0.0 pandas==2.0.0 torch==2.4.0 - numba==0.61.0 scipy==1.13.0 # ============================================================================= @@ -152,7 +131,6 @@ deps = numpy pandas torch - numba scipy [testenv:py312-numpy2x] @@ -163,20 +141,8 @@ deps = numpy>=2.0.0 pandas>=2.2.0 torch>=2.4.0 - numba>=0.61.0 scipy>=1.13.0 -[testenv:py312-numba063] -description = Python 3.12 with numba 0.63 (latest) -basepython = python3.12 -deps = - {[testenv]deps} - numpy>=2.0.0 - pandas>=2.2.0 - torch>=2.5.0 - numba>=0.63.0 - scipy>=1.14.0 - # ============================================================================= # Python 3.13 - Latest Python # ============================================================================= @@ -188,7 +154,6 @@ deps = numpy pandas torch - numba scipy # ============================================================================= @@ -202,7 +167,6 @@ deps = numpy>=2.0.0,<2.1.0 pandas>=2.1.0,<2.3.0 torch>=2.4.0,<2.5.0 - numba>=0.61.0,<0.62.0 scipy>=1.13.0,<1.14.0 [testenv:py311-torch26] @@ -213,7 +177,6 @@ deps = numpy>=2.1.0,<2.2.0 pandas>=2.2.0,<2.3.0 torch>=2.6.0,<2.7.0 - numba>=0.61.0,<0.62.0 scipy>=1.14.0,<1.15.0 # ============================================================================= @@ -227,7 +190,6 @@ deps = numpy>=2.0.0,<2.1.0 pandas>=2.0.0,<2.1.0 torch>=2.4.0,<2.5.0 - numba>=0.61.0,<0.62.0 scipy>=1.13.0,<1.14.0 [testenv:py311-pandas22] @@ -238,7 +200,6 @@ deps = numpy>=2.0.0,<2.2.0 pandas>=2.2.0,<2.3.0 torch>=2.4.0,<2.6.0 - numba>=0.61.0,<0.62.0 scipy>=1.13.0,<1.15.0 # ============================================================================= @@ -258,7 +219,6 @@ deps = numpy>=2.0.0,<2.1.0 pandas>=2.0.0,<2.2.0 torch>=2.4.0,<2.5.0 - numba>=0.61.0,<0.62.0 scipy>=1.13.0,<1.14.0 # ============================================================================= @@ -274,7 +234,6 @@ deps = pandas==2.0.0 torch==2.4.0 tqdm==4.61.0 - numba==0.59.0 gemmi==0.5.0 scipy==1.10.0 matplotlib==3.7.0 From e70e9d049e07e05f4e3b854992b3d54915956ed3 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sun, 27 Sep 2026 17:50:48 +0200 Subject: [PATCH 184/250] Cache the parsed hydrogen mask on AtomGraph is_hydrogen re-parsed the whole element column (np.char.strip/upper) on every access, and hydrogen placement and the riding frames read it once per atom, so setting up hydrogens was O(N^2). The mask is now parsed once and cached against the element array it came from; each call still returns a fresh tensor. Switching 4BX9 (20,001 atoms with H) to riding hydrogens went from 132 s to 1.5 s, and 5BOV from 108 s to 1.3 s. Placed hydrogens, riding frames and restraints are identical on 11 AlphaFold-start structures. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01KwtigurBqaYuzV6426n5XG --- torchref/topology/atom_graph.py | 17 ++++++++++++++--- 1 file changed, 14 insertions(+), 3 deletions(-) diff --git a/torchref/topology/atom_graph.py b/torchref/topology/atom_graph.py index 05204315..dcfce9b6 100644 --- a/torchref/topology/atom_graph.py +++ b/torchref/topology/atom_graph.py @@ -152,6 +152,8 @@ class AtomGraph(DeviceMixin): _adj_indptr: Optional[torch.Tensor] = field(default=None, repr=False) _adj_indices: Optional[torch.Tensor] = field(default=None, repr=False) + # (element array it was parsed from, hydrogen flags); see is_hydrogen. + _is_h_cache: Optional[Tuple[np.ndarray, np.ndarray]] = field(default=None, repr=False) def __post_init__(self) -> None: if self._adj_indptr is None: @@ -169,9 +171,18 @@ def n_atoms(self) -> int: @property def is_hydrogen(self) -> torch.Tensor: - """Boolean mask of hydrogen atoms, shape ``(N,)``.""" - flags = np.char.upper(np.char.strip(self.element.astype(str))) == "H" - return torch.as_tensor(flags, device=self.bonds.indices.device) + """Boolean mask of hydrogen atoms, shape ``(N,)``. + + The element strings are parsed once and cached against the ``element`` array + they came from, so replacing that array invalidates the cache. Hydrogen + placement and the riding frames read this once per atom, and re-parsing every + time made them O(N^2). Each call returns a fresh tensor, so callers may modify it. + """ + cache = self._is_h_cache + if cache is None or cache[0] is not self.element: + flags = np.char.upper(np.char.strip(self.element.astype(str))) == "H" + cache = self._is_h_cache = (self.element, flags) + return torch.tensor(cache[1], device=self.bonds.indices.device) def copy(self) -> "AtomGraph": """An independent copy sharing no storage with this one.""" From 3e9758e5a4671ce5167575644e7574ac00ee411c Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sun, 27 Sep 2026 17:50:49 +0200 Subject: [PATCH 185/250] Build the ADP-locality neighbour list with a k-d tree ADPLocalityTarget.stats() rebuilds the k-NN list on every metrics collection (about 100 times in a 10-cycle run), which is also what keeps forward()'s list current. The build was a Python loop over a cell list with one argpartition per atom, and it cost 34 s of a 164 s refinement of 6G9X before hydrogens doubled it. It is now a scipy cKDTree query, with distances recomputed from the coordinates as before. On 11 structures, with and without hydrogens, every atom keeps the same neighbours at the same distances; order differs only among equal distances, and the loss agrees to 1e-7. A build is 4-8x faster (4BX9 with H: 1.94 s -> 0.25 s). Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01KwtigurBqaYuzV6426n5XG --- torchref/refinement/targets/adp/locality.py | 135 +++++--------------- 1 file changed, 29 insertions(+), 106 deletions(-) diff --git a/torchref/refinement/targets/adp/locality.py b/torchref/refinement/targets/adp/locality.py index fbfd4ed0..4c0b3d68 100644 --- a/torchref/refinement/targets/adp/locality.py +++ b/torchref/refinement/targets/adp/locality.py @@ -2,6 +2,7 @@ import numpy as np import torch +from scipy.spatial import cKDTree from typing import TYPE_CHECKING, Dict from torchref.base.targets.adp import adp_locality_aniso_math @@ -127,12 +128,20 @@ def sigma_aniso(self, value: float): self._sigma_aniso.fill_(value) # ------------------------------------------------------------------ - # Spatial-hash k-NN (O(N) memory) + # k-NN list # ------------------------------------------------------------------ def _build_neighbor_list(self) -> None: - """Build the k-NN list via a spatial cell-list: O(N·k) for the output plus - O(N) bookkeeping, instead of the O(N²) of a full distance matrix. + """Build each atom's list of its ``k`` nearest other atoms, nearest first. + + Uses a k-d tree. ``stats()`` rebuilds the list on every call, which is also what + keeps ``forward()``'s list current as atoms move, so this runs every time + metrics are collected and has to be cheap. The per-atom Python loop over a cell + list that it replaces was one of the largest costs of a refinement, and more so + with hydrogens, which double the atom count. + + Distances are recomputed from the coordinates in their own dtype, the way the + loss sees them, rather than taken from the tree. """ xyz = self.model.xyz() device = xyz.device @@ -141,114 +150,28 @@ def _build_neighbor_list(self) -> None: coords = xyz.detach().cpu().numpy() - # Must cover the kth-neighbour distance or neighbours are silently missed; - # for proteins k=50 sits within ~8-10 Å, so 12 Å has margin. - cell_size = 12.0 - - xyz_min = coords.min(axis=0) - cell_idx = ((coords - xyz_min) / cell_size).astype(np.int64) - - grid_dims = cell_idx.max(axis=0) + 1 - gx, gy, gz = int(grid_dims[0]), int(grid_dims[1]), int(grid_dims[2]) - gyz = gy * gz - - flat = cell_idx[:, 0] * gyz + cell_idx[:, 1] * gz + cell_idx[:, 2] - - order = np.argsort(flat) - sorted_flat = flat[order] - - unique_cells, first_idx, counts = np.unique( - sorted_flat, return_index=True, return_counts=True - ) - n_unique = len(unique_cells) - - # start[i] .. start[i+1] are the atoms in unique cell i - starts = np.empty(n_unique + 1, dtype=np.int64) - starts[0] = 0 - starts[1:] = np.cumsum(counts) - - # flat_cell -> unique index (-1 = empty) - n_grid = gx * gyz - cell_lookup = np.full(n_grid, -1, dtype=np.int64) - cell_lookup[unique_cells] = np.arange(n_unique, dtype=np.int64) - - # 27 neighbor offsets (self + all adjacent cells) - offsets = [] - for dx in range(-1, 2): - for dy in range(-1, 2): - for dz in range(-1, 2): - offsets.append((dx, dy, dz, dx * gyz + dy * gz + dz)) - - # For each atom, collect candidate neighbors and keep top-k - all_neighbor_idx = np.zeros((n_atoms, k), dtype=np.int64) - all_neighbor_dist = np.full((n_atoms, k), np.inf, dtype=np.float32) - - # atom_cell[i] = unique-cell index for atom i - atom_cell = np.empty(n_atoms, dtype=np.int64) - atom_cell[order] = np.repeat(np.arange(n_unique), counts) - - for ci in range(n_unique): - cell_flat = int(unique_cells[ci]) - sa, ea = int(starts[ci]), int(starts[ci + 1]) - atoms_a = order[sa:ea] - xyz_a = coords[atoms_a] - - cx = cell_flat // gyz - cy = (cell_flat % gyz) // gz - cz = cell_flat % gz - - # Collect all candidate neighbor atoms from adjacent cells - cand_atoms_list = [] - cand_xyz_list = [] - for dx, dy, dz, _ in offsets: - ncx, ncy, ncz = cx + dx, cy + dy, cz + dz - if ncx < 0 or ncx >= gx or ncy < 0 or ncy >= gy or ncz < 0 or ncz >= gz: - continue - nb_flat = ncx * gyz + ncy * gz + ncz - nb_ci = int(cell_lookup[nb_flat]) - if nb_ci < 0: - continue - sb, eb = int(starts[nb_ci]), int(starts[nb_ci + 1]) - cand_atoms_list.append(order[sb:eb]) - cand_xyz_list.append(coords[order[sb:eb]]) - - if not cand_atoms_list: - continue - - cand_atoms = np.concatenate(cand_atoms_list) - cand_xyz = np.concatenate(cand_xyz_list, axis=0) - - # Distances from each atom in this cell to all candidates - # shape: (len(atoms_a), len(cand_atoms)) - diff = xyz_a[:, None, :] - cand_xyz[None, :, :] - dist = np.sqrt((diff * diff).sum(axis=-1)) - - for li, ai in enumerate(atoms_a): - d = dist[li] - # Mask self - self_mask = cand_atoms == ai - d[self_mask] = np.inf - - if len(d) <= k: - top_k_idx = np.argsort(d)[:k] - else: - top_k_idx = np.argpartition(d, k)[:k] - # Sort the top-k for deterministic order - sub_order = np.argsort(d[top_k_idx]) - top_k_idx = top_k_idx[sub_order] - - n_valid = min(k, len(top_k_idx)) - all_neighbor_idx[ai, :n_valid] = cand_atoms[top_k_idx[:n_valid]] - all_neighbor_dist[ai, :n_valid] = d[top_k_idx[:n_valid]] + if k <= 0: + all_neighbor_idx = np.zeros((n_atoms, 0), dtype=np.int64) + all_neighbor_dist = np.zeros((n_atoms, 0), dtype=np.float32) + else: + # k + 1, because each atom is its own nearest point. + _, idx = cKDTree(coords).query(coords, k=k + 1) + # Drop the atom itself by index, not by column: a coincident atom can sort + # ahead of it. Where it is absent (more than k coincident atoms) the + # farthest candidate goes instead. + keep = idx != np.arange(n_atoms)[:, None] + keep[keep.all(axis=1), -1] = False + all_neighbor_idx = idx[keep].reshape(n_atoms, k).astype(np.int64) + diff = coords[:, None, :] - coords[all_neighbor_idx] + all_neighbor_dist = np.sqrt((diff * diff).sum(axis=-1)).astype(np.float32) self._neighbor_indices = torch.from_numpy(all_neighbor_idx).to(device) self._neighbor_distances = torch.from_numpy(all_neighbor_dist).to(device) - if self.verbose > 1: - mean_dist = float(all_neighbor_dist[all_neighbor_dist < np.inf].mean()) + if self.verbose > 1 and all_neighbor_dist.size: print( - f" Built K-NN list (spatial hash): k={k}, " - f"mean dist={mean_dist:.2f}A" + f" Built K-NN list (k-d tree): k={k}, " + f"mean dist={float(all_neighbor_dist.mean()):.2f}A" ) def forward(self, recompute_neighbors: bool = False) -> torch.Tensor: From b7af8df54f7607665184041dec27220bcd1230e1 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sun, 27 Sep 2026 17:50:49 +0200 Subject: [PATCH 186/250] Search VDW pairs with a k-d tree on CPU find_pairs_periodic_grid_v2 pads every grid cell to the fullest one and takes dense cdist tiles between neighbouring cells, so its cost goes with the square of the peak cell occupancy, which hydrogens roughly double. On CPU the pair search is now find_pairs_kdtree, a cKDTree query of the ASU atoms against all symmetry images, returning the same (i, j, combo_j) triples; the grid stays the accelerator path. A search on 5BOV with hydrogens went from 44 s to 0.6 s. The grid cells are cell_length / cutoff wide along each axis, which in an oblique cell leaves them narrower than the cutoff perpendicular to a face, and pairs between that width and the cutoff were missed. The k-d tree finds them: up to 0.23% more pairs on 5 of 11 test structures, all between 5.2 and 6.0 A and so in the drift margin, not in contact. Otherwise the pair sets agree except for rounding at the cutoff. test_vdw_pair_search checks the k-d tree against brute force and the grid against the k-d tree, which keeps the grid covered on CPU-only CI. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01KwtigurBqaYuzV6426n5XG --- tests/unit/topology/test_vdw_pair_search.py | 128 ++++++++++++++++++++ torchref/topology/nonbonded.py | 115 ++++++++++++++---- 2 files changed, 219 insertions(+), 24 deletions(-) create mode 100644 tests/unit/topology/test_vdw_pair_search.py diff --git a/tests/unit/topology/test_vdw_pair_search.py b/tests/unit/topology/test_vdw_pair_search.py new file mode 100644 index 00000000..e0601eff --- /dev/null +++ b/tests/unit/topology/test_vdw_pair_search.py @@ -0,0 +1,128 @@ +"""The two VDW pair searches: the periodic grid (accelerators) and the k-d tree (CPU). + +Both take the same symmetry-image table and return ``(i, j, combo_j)`` under one +convention, so the rest of ``build_vdw_restraints_gpu`` cannot tell them apart. The k-d +tree is checked against brute force; the grid is checked against the k-d tree, which +also keeps it exercised on a CPU-only CI now that CPU builds no longer call it. + +The grid's cells are ``cell_length / cutoff`` along each axis, which in an oblique cell +leaves them narrower than the cutoff perpendicular to a face. Pairs between that width +and the cutoff can then sit two cells apart and are missed. Both structures here are +oblique (1BYW hexagonal, 3E98 monoclinic), so that shortfall is asserted rather than +assumed away. +""" + +import numpy as np +import pytest +import torch + +from torchref.config import dtypes +from torchref.model.model import Model +from torchref.topology import nonbonded as nb + +CUTOFF = 6.0 + + +def _image_table(path): + """Steps 1-2 of ``build_vdw_restraints_gpu`` for one model.""" + model = Model(verbose=0) + model.load_pdb(str(path)) + cell, sg = model.ctx.cell, model.ctx.spacegroup + xyz_frac = cell.cartesian_to_fractional(model.xyz().detach().to(dtypes.float)) + op_indices, offsets = nb.prefilter_symop_offsets(cell, sg, xyz_frac, CUTOFF) + identity = ((op_indices == 0) & (offsets == 0).all(dim=1)).nonzero()[0].item() + lengths = torch.stack([cell.a, cell.b, cell.c]).to(dtypes.float) + grid_dims = torch.clamp((lengths / CUTOFF).long(), min=1) + flat_cell, atom_idx, combo_idx, cart_pos = nb.assign_to_grid( + xyz_frac, cell, sg, op_indices, offsets, grid_dims + ) + # Perpendicular width of one grid cell along each axis: lattice-plane spacing + # V / |face| over the number of cells. + basis = cell.fractional_to_cartesian(torch.eye(3, dtype=dtypes.float)) + volume = torch.linalg.det(basis).abs() + faces = [(basis[1], basis[2]), (basis[2], basis[0]), (basis[0], basis[1])] + widths = [ + (volume / torch.linalg.cross(u, v).norm() / grid_dims[k]).item() + for k, (u, v) in enumerate(faces) + ] + return dict( + n_atoms=xyz_frac.shape[0], n_combos=len(op_indices), identity=identity, + grid_dims=grid_dims, flat_cell=flat_cell, atom_idx=atom_idx, + combo_idx=combo_idx, cart_pos=cart_pos, min_width=min(widths), + ) + + +def _keys(i, j, c, t): + """``(i, j, combo_j)`` as one integer each, for set comparison.""" + return ((i * t["n_atoms"] + j) * t["n_combos"] + c).cpu().numpy() + + +def _distance(keys, t): + n, m = t["n_atoms"], t["n_combos"] + pos = t["cart_pos"].reshape(n, m, 3).double() + keys = torch.as_tensor(keys) + i, j, c = keys // (n * m), (keys // m) % n, keys % m + return (pos[i, t["identity"]] - pos[j, c]).norm(dim=1).numpy() + + +@pytest.fixture(scope="module", params=["1BYW_af.pdb", "3E98.pdb"]) +def table(request, pdb_dir): + return _image_table(pdb_dir / request.param) + + +def test_kdtree_matches_brute_force(table): + t = table + n, m = t["n_atoms"], t["n_combos"] + # The image table is atom-major; the distance lookup relies on it. + assert torch.equal(t["atom_idx"], torch.arange(n).repeat_interleave(m)) + assert torch.equal(t["combo_idx"], torch.arange(m).repeat(n)) + + pos = t["cart_pos"].reshape(n, m, 3) + asu = pos[:, t["identity"]] + images = pos.reshape(n * m, 3) + expected = [] + for start in range(0, n, 256): + d = torch.cdist(asu[start:start + 256].double(), images.double(), compute_mode="donot_use_mm_for_euclid_dist") + i, e = (d < CUTOFF).nonzero(as_tuple=True) + i = i + start + j, c = e // m, e % m + keep = (c != t["identity"]) | (i < j) + expected.append(_keys(i[keep], j[keep], c[keep], t)) + expected = np.concatenate(expected) + + got = _keys(*nb.find_pairs_kdtree( + t["cart_pos"], t["atom_idx"], t["combo_idx"], CUTOFF, t["identity"] + ), t) + assert np.array_equal(np.sort(expected), np.sort(got)) + + +def test_kdtree_output_convention(table): + t = table + i, j, c = nb.find_pairs_kdtree( + t["cart_pos"], t["atom_idx"], t["combo_idx"], CUTOFF, t["identity"] + ) + keys = _keys(i, j, c, t) + assert np.all(np.diff(keys) > 0), "sorted by (i, j, combo_j), no duplicates" + intra = c == t["identity"] + assert bool((i[intra] < j[intra]).all()), "intra-ASU pairs once, i < j, no self" + assert i.dtype == j.dtype == c.dtype == torch.int64 + + +def test_grid_is_the_kdtree_minus_its_cell_width_shortfall(table): + t = table + order, cells, starts, lookup = nb.build_cell_list(t["flat_cell"], int(t["grid_dims"].prod())) + grid = np.unique(_keys(*nb.find_pairs_periodic_grid_v2( + t["cart_pos"][order], t["atom_idx"][order], t["combo_idx"][order], + cells, starts, lookup, t["grid_dims"], CUTOFF, t["identity"], + ), t)) + tree = _keys(*nb.find_pairs_kdtree( + t["cart_pos"], t["atom_idx"], t["combo_idx"], CUTOFF, t["identity"] + ), t) + + # Allow for the grid's matmul cdist rounding right at the cutoff. + only_grid = np.setdiff1d(grid, tree) + assert np.all(np.abs(_distance(only_grid, t) - CUTOFF) < 1e-4) + only_tree = np.setdiff1d(tree, grid) + assert t["min_width"] < CUTOFF, "both test cells are oblique enough to show it" + assert np.all(_distance(only_tree, t) > t["min_width"] - 1e-4) + assert len(only_tree) < 0.01 * len(tree) diff --git a/torchref/topology/nonbonded.py b/torchref/topology/nonbonded.py index 404e7a6f..9f1b646a 100644 --- a/torchref/topology/nonbonded.py +++ b/torchref/topology/nonbonded.py @@ -4,7 +4,8 @@ Works in fractional space with periodic boundary conditions. Avoids explicit symmetry expansion by assigning (atom, symop+offset) entries to grid cells and using padded batched ``torch.cdist`` -for distance computation. +for distance computation. On CPU the pair search itself is a +k-d tree instead (:func:`find_pairs_kdtree`), with the same output. All operations run under ``torch.no_grad()`` on whatever device the input coordinates live on (CPU or GPU). @@ -503,6 +504,63 @@ def find_pairs_periodic_grid_v2( ) +def find_pairs_kdtree( + cart_pos: torch.Tensor, + atom_idx: torch.Tensor, + combo_idx: torch.Tensor, + cutoff: float, + identity_combo: int, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """CPU counterpart of :func:`find_pairs_periodic_grid_v2`, on a k-d tree. + + Returns the same pairs under the same convention: every ASU atom ``i`` and image + ``(j, combo_j)`` closer than ``cutoff``, intra-ASU pairs once with ``i < j``, no + self-pairs. Sorted by ``(i, j, combo_j)``. + + The grid search pads every cell to the fullest one and takes dense distance tiles + between neighbouring cells, so its cost goes with the square of the peak cell + occupancy -- about 4x for the same volume once hydrogens are added. Here the cost + goes with the pairs actually within the cutoff. A tree also has no cell width to + fall short of the cutoff, which a fractional grid cell can do in an oblique cell. + + Parameters + ---------- + cart_pos : (E, 3) float + Cartesian image positions from :func:`assign_to_grid`, unsorted. + atom_idx, combo_idx : (E,) long + ASU atom index and (symop, offset) combo index per entry. + cutoff : float + Cartesian distance cutoff in Angstrom. + identity_combo : int + Combo index corresponding to the identity (op=0, offset=0). + """ + from scipy.spatial import cKDTree + + device = cart_pos.device + pos = cart_pos.detach().cpu().numpy() + atoms = atom_idx.cpu().numpy() + combos = combo_idx.cpu().numpy() + + asu = np.nonzero(combos == identity_combo)[0] + found = cKDTree(pos[asu]).sparse_distance_matrix( + cKDTree(pos), cutoff, output_type="ndarray" + ) + # sparse_distance_matrix keeps d <= cutoff; the grid search keeps d < cutoff. + found = found[found["v"] < cutoff] + ai = atoms[asu[found["i"]]] + aj = atoms[found["j"]] + cj = combos[found["j"]] + # Intra-ASU contacts come back from both ends; keep i < j, which also drops self. + keep = (cj != identity_combo) | (ai < aj) + ai, aj, cj = ai[keep], aj[keep], cj[keep] + + order = np.lexsort((cj, aj, ai)) + return tuple( + torch.from_numpy(np.ascontiguousarray(a[order], dtype=np.int64)).to(device) + for a in (ai, aj, cj) + ) + + # ------------------------------------------------------------------ # # Step 5 – filtering # ------------------------------------------------------------------ # @@ -678,31 +736,40 @@ def build_vdw_restraints_gpu( xyz_frac, cell, sg, op_indices, cell_offsets_valid, grid_dims ) - n_grid_total = grid_dims[0].item() * grid_dims[1].item() * grid_dims[2].item() + if device.type == "cpu": + # Steps 3-4 on CPU: a k-d tree, whose cost follows the pairs found rather + # than the padded cell tiles of the grid search below. + if verbose > 0: + print(f" Pair search: k-d tree over {cart_pos.shape[0]} images") + pair_atom_i, pair_atom_j, pair_combo_j = find_pairs_kdtree( + cart_pos, atom_idx, combo_idx, cutoff, identity_combo + ) + else: + n_grid_total = grid_dims[0].item() * grid_dims[1].item() * grid_dims[2].item() - # Step 3: sort into cell list - sort_order, unique_cells, starts, cell_lookup = build_cell_list( - flat_cell, n_grid_total - ) - cart_sorted = cart_pos[sort_order] - atom_idx_sorted = atom_idx[sort_order] - combo_idx_sorted = combo_idx[sort_order] + # Step 3: sort into cell list + sort_order, unique_cells, starts, cell_lookup = build_cell_list( + flat_cell, n_grid_total + ) + cart_sorted = cart_pos[sort_order] + atom_idx_sorted = atom_idx[sort_order] + combo_idx_sorted = combo_idx[sort_order] - if verbose > 0: - n_occupied = len(unique_cells) - counts = starts[1:] - starts[:-1] - print(f" Grid: {grid_dims.tolist()}, " - f"{n_occupied}/{n_grid_total} cells occupied, " - f"max {counts.max().item()} entries/cell") - - # Step 4: find pairs via periodic grid + batched cdist. Nearly dedup-free by - # construction, but the hash dedup below still catches intra-cell - # swap-canonicalisation collisions. - pair_atom_i, pair_atom_j, pair_combo_j = find_pairs_periodic_grid_v2( - cart_sorted, atom_idx_sorted, combo_idx_sorted, - unique_cells, starts, cell_lookup, grid_dims, - cutoff, identity_combo, - ) + if verbose > 0: + n_occupied = len(unique_cells) + counts = starts[1:] - starts[:-1] + print(f" Grid: {grid_dims.tolist()}, " + f"{n_occupied}/{n_grid_total} cells occupied, " + f"max {counts.max().item()} entries/cell") + + # Step 4: find pairs via periodic grid + batched cdist. Nearly dedup-free by + # construction, but the hash dedup below still catches intra-cell + # swap-canonicalisation collisions. + pair_atom_i, pair_atom_j, pair_combo_j = find_pairs_periodic_grid_v2( + cart_sorted, atom_idx_sorted, combo_idx_sorted, + unique_cells, starts, cell_lookup, grid_dims, + cutoff, identity_combo, + ) if len(pair_atom_i) == 0: if verbose > 0: From e8ea8d62cebd29f19ce8430defcbcae3bdfae198 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sun, 27 Sep 2026 17:50:49 +0200 Subject: [PATCH 187/250] Skip the pair search in the restraint build that plans hydrogens Adding hydrogens builds restraints over the heavy-atom table only to read its topology, then discards them and builds again over the augmented table. That first build ran the full non-bonded setup: the pair search and the riding-hydrogen stand-in for it. Restraints takes nonbonded=False to skip it, and Model._add_missing_hydrogens uses it, silently, since the build that follows reports. With the k-d tree pair search in place this cuts the discarded build from 3.9 s to 1.7 s for 5BOV and 4BX9 (1.7 s to 1.0 s for 3VRJ). Placed hydrogens, riding frames and restraints are unchanged on 11 structures. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01KwtigurBqaYuzV6426n5XG --- torchref/model/model.py | 29 ++++++++++++++++++++--------- torchref/topology/restraints.py | 16 ++++++++++++---- 2 files changed, 32 insertions(+), 13 deletions(-) diff --git a/torchref/model/model.py b/torchref/model/model.py index 56d9a8db..b66b8d17 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -572,12 +572,20 @@ def _build_restraints(self): "Load data first with load_pdb() or load_cif()." ) - from torchref.topology.restraints import Restraints - if self.ctx.verbose > 0: print("Building restraints...") - self._restraints = Restraints( + self._restraints = self._new_restraints() + + return self._restraints + + def _new_restraints(self, nonbonded: bool = True, verbose: Optional[int] = None): + """An uncached ``Restraints`` over this model's DataFrame, wired to the live + ``xyz`` / ``adp`` / ``vdw_radii`` callables; see :meth:`_build_restraints`. + """ + from torchref.topology.restraints import Restraints + + return Restraints( pdb=self.pdb, cif_path=self.ctx.cif_path, xyz_fn=self.xyz, @@ -586,11 +594,10 @@ def _build_restraints(self): cell=self.ctx.cell, spacegroup=self.ctx.spacegroup, links=self.ctx.links, - verbose=self.ctx.verbose, + verbose=self.ctx.verbose if verbose is None else verbose, + nonbonded=nonbonded, ) - return self._restraints - @property def restraints(self): """Bond/angle/torsion/... restraints, built on first access from the @@ -790,8 +797,10 @@ def _add_missing_hydrogens(self) -> None: fixed point. Costs a restraint build that is then discarded, because the plan needs the - topology and the topology is built over the atoms as loaded. Loading invokes - this only when ``add_hydrogens=True`` is requested. + topology and the topology is built over the atoms as loaded. It skips the + non-bonded pair search, which the plan does not use and which was the largest + part of that build, and it is silent: the build over the augmented table reports. + Loading invokes this only when ``add_hydrogens=True`` is requested. """ from torchref.topology.hydrogens import ( augment_atom_table, @@ -799,7 +808,9 @@ def _add_missing_hydrogens(self) -> None: plan_hydrogens, ) - restraints = self.restraints + restraints = self._restraints + if restraints is None: + restraints = self._new_restraints(nonbonded=False, verbose=0) xyz = self.xyz().detach() plan = plan_hydrogens( restraints.topology, restraints.cif_dict, xyz, verbose=self.ctx.verbose diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index 0c49868f..a79009dc 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -68,6 +68,10 @@ class Restraints(DeviceMixin, DebugMixin, Module): between the two named atoms. verbose : int, default 1 Verbosity level (0=silent, 1=normal, 2=detailed). + nonbonded : bool, default True + Build the non-bonded pair list. ``False`` gives connectivity and ideal values + only, for a caller that needs the topology and nothing else, such as hydrogen + planning; the pair search is the largest part of a build. Attributes ---------- @@ -98,12 +102,14 @@ def __init__( spacegroup=None, links: pd.DataFrame = None, verbose: int = 1, + nonbonded: bool = True, ): """Initialize the Restraints handler.""" super().__init__() self.cif_path = cif_path self.verbose = verbose self.links = links + self._nonbonded = bool(nonbonded) # Store callable functions for coordinate/ADP access self._xyz_fn = xyz_fn @@ -350,7 +356,8 @@ def _load_rama_surfaces(self, device: torch.device): def build_restraints(self): """Build the topology, the values over it, and the non-bonded pair list. - Builds on CPU and moves the result to the ``xyz()`` device at the end. + The pair list is skipped when constructed with ``nonbonded=False``. Builds on + CPU and moves the result to the ``xyz()`` device at the end. """ try: target_device = self.xyz().device @@ -380,9 +387,10 @@ def build_restraints(self): # cutoff sits ~1 Å beyond the largest heavy-atom VDW sum (~3.6 Å) plus # expected drift, so a displacement-triggered rebuild stays inside the # margin and cannot miss a newly-formed contact. - self._build_vdw_restraints( - cutoff=6.0, sigma=0.05, inter_residue_only=False, use_spatial_hash=True - ) + if self._nonbonded: + self._build_vdw_restraints( + cutoff=6.0, sigma=0.05, inter_residue_only=False, use_spatial_hash=True + ) if target_device.type != "cpu": self.to(target_device) From dac79136de83505964ff784249c0ce9584a9234e Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sun, 27 Sep 2026 20:28:04 +0200 Subject: [PATCH 188/250] Pin the VDW pair-search tests to CPU The k-d tree is the CPU search, but the fixture loaded its model on the default device, so on the MPS runner the cell sat on the GPU next to CPU tensors and all six tests errored at setup. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01KwtigurBqaYuzV6426n5XG --- tests/unit/topology/test_vdw_pair_search.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/tests/unit/topology/test_vdw_pair_search.py b/tests/unit/topology/test_vdw_pair_search.py index e0601eff..35370679 100644 --- a/tests/unit/topology/test_vdw_pair_search.py +++ b/tests/unit/topology/test_vdw_pair_search.py @@ -24,8 +24,12 @@ def _image_table(path): - """Steps 1-2 of ``build_vdw_restraints_gpu`` for one model.""" - model = Model(verbose=0) + """Steps 1-2 of ``build_vdw_restraints_gpu`` for one model, on CPU. + + Pinned to CPU, not ``TORCHREF_DEVICE``: the k-d tree is the CPU search, and the + accelerator runners would otherwise put the cell on the GPU next to CPU tensors. + """ + model = Model(verbose=0, device=torch.device("cpu")) model.load_pdb(str(path)) cell, sg = model.ctx.cell, model.ctx.spacegroup xyz_frac = cell.cartesian_to_fractional(model.xyz().detach().to(dtypes.float)) From 7ddf1d8d431afe0ecb13f8c4244df757ef57c349 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sun, 27 Sep 2026 20:28:05 +0200 Subject: [PATCH 189/250] Read a single _struct_conn entry written as key-value pairs A file with one connection writes _struct_conn without loop_, which the CIF parser keeps as {attribute: value} rather than a DataFrame of full tag names, and get_link_records crashed on it with AttributeError. 3GR5.cif is such a file, so load_cif failed on it. test_load_different_cif_files loaded the first three files of an unsorted glob, so whether it reached 3GR5 depended on the filesystem: it failed on GitHub runners and passed on merlin7. It now loads every bundled CIF in sorted order, and test_single_connection_written_as_key_value_pairs covers the key-value form directly, including one surviving covalent link. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01KwtigurBqaYuzV6426n5XG --- tests/integration/test_model_operations.py | 6 ++++- tests/unit/io/test_struct_conn_links.py | 26 ++++++++++++++++++++++ torchref/io/cif_readers.py | 5 +++++ 3 files changed, 36 insertions(+), 1 deletion(-) diff --git a/tests/integration/test_model_operations.py b/tests/integration/test_model_operations.py index 2db27886..08b7a736 100644 --- a/tests/integration/test_model_operations.py +++ b/tests/integration/test_model_operations.py @@ -309,7 +309,11 @@ def test_load_different_cif_files(self, cif_dir): """Test loading different CIF files.""" from torchref.model.model import Model - cif_files = list(cif_dir.glob("*.cif"))[:3] # Load first 3 + # All of them, in a fixed order: "the first three" of an unsorted glob depended on + # the filesystem, and on some runners never reached a file (3GR5) whose single + # _struct_conn entry is written as key-value pairs rather than a loop. + cif_files = sorted(cif_dir.glob("*.cif")) + assert cif_files for cif_file in cif_files: model = Model() diff --git a/tests/unit/io/test_struct_conn_links.py b/tests/unit/io/test_struct_conn_links.py index 2e9bc234..fc14eec3 100644 --- a/tests/unit/io/test_struct_conn_links.py +++ b/tests/unit/io/test_struct_conn_links.py @@ -72,3 +72,29 @@ def test_file_without_struct_conn_gives_empty_table(tmp_path): links = ModelCIFReader(str(path)).links assert len(links) == 0 assert list(links.columns) == list(LINK_COLUMNS) + + +@pytest.mark.unit +def test_single_connection_written_as_key_value_pairs(cif_dir, tmp_path): + """One ``_struct_conn`` entry is written without ``loop_`` and parses as a dict.""" + source = cif_dir / "3GR5.cif" + # 3GR5's only connection is a disulfide, which is left to distance detection. + links = ModelCIFReader(str(source)).links + assert len(links) == 0 + assert list(links.columns) == list(LINK_COLUMNS) + + # The same single entry as a covalent connection comes through as one link. + text = source.read_text() + covalent = text.replace( + "_struct_conn.conn_type_id disulf ", + "_struct_conn.conn_type_id covale ", + ) + assert covalent != text + path = tmp_path / "3GR5_covale.cif" + path.write_text(covalent) + links = ModelCIFReader(str(path)).links + assert len(links) == 1 + row = links.iloc[0] + assert (row["name1"], row["resname1"], row["resseq1"]) == ("SG", "CYS", 136) + assert (row["name2"], row["resname2"], row["resseq2"]) == ("SG", "CYS", 155) + assert abs(row["length"] - 2.062) < 1e-6 diff --git a/torchref/io/cif_readers.py b/torchref/io/cif_readers.py index 221bb955..5f5f4b04 100644 --- a/torchref/io/cif_readers.py +++ b/torchref/io/cif_readers.py @@ -1355,6 +1355,11 @@ def get_link_records(self) -> pd.DataFrame: conn = self.cif.data.get("struct_conn") if conn is None or len(conn) == 0: return empty + if isinstance(conn, dict): + # A file with a single connection writes it as key-value pairs rather than + # a loop, which the parser keeps as {attribute: value}; loops come back as a + # DataFrame of full tag names. + conn = pd.DataFrame([{f"_struct_conn.{k}": v for k, v in conn.items()}]) def column(names, default=""): for name in names: From ee1fa533e83e0ae41ff11b7b9f12386f781b8c18 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sun, 27 Sep 2026 20:28:05 +0200 Subject: [PATCH 190/250] Replace the stale strict xfail on copy() with a real test test_copy_after_refinement_setup_is_broken_for_every_representation was a strict xfail for a copy() failure after Refinement setup. copy() works now, so the XPASS failed the suite; it is now an ordinary test that the copy matches the original. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01KwtigurBqaYuzV6426n5XG --- .../integration/test_adp_field_refinement.py | 22 +++++++++---------- 1 file changed, 10 insertions(+), 12 deletions(-) diff --git a/tests/integration/test_adp_field_refinement.py b/tests/integration/test_adp_field_refinement.py index 4dca8c03..6969bd37 100644 --- a/tests/integration/test_adp_field_refinement.py +++ b/tests/integration/test_adp_field_refinement.py @@ -332,18 +332,16 @@ def test_model_copy_round_trip_on_a_bare_model(): @pytest.mark.integration -@pytest.mark.xfail( - reason="PRE-EXISTING and representation-independent: once a Model has been through " - "Refinement setup, a cache somewhere holds a graph-attached tensor and deepcopy " - "refuses it. Measured identically for adp_mode isotropic, anisotropic and " - "field_aniso, and a bare model copies fine, so the node field is not the cause -- " - "it means no refinement of any kind can currently be checkpointed by copy().", - raises=RuntimeError, - strict=True, -) -@pytest.mark.integration -def test_copy_after_refinement_setup_is_broken_for_every_representation(field_refinement): - field_refinement.model.copy() +def test_copy_after_refinement_setup(field_refinement): + """A model that has been through Refinement setup can still be checkpointed by copy(). + + This used to raise for every ADP representation: a cache held a graph-attached + tensor and deepcopy refused it. + """ + model = field_refinement.model + clone = model.copy() + assert torch.allclose(clone.xyz().detach(), model.xyz().detach()) + assert torch.allclose(clone.adp_u6().detach(), model.adp_u6().detach()) # ---------------------------------------------------------------------------------- From 0d6a837c7a23395f10937b54dd21db823fbc0fe6 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sun, 27 Sep 2026 20:28:05 +0200 Subject: [PATCH 191/250] Justify or drop the unmarked integer dtypes in the hydrogen code test_no_unjustified_hardcoded_dtype flagged 18 literals. Most already had a dtype-ok reason, but on the closing parenthesis below the flagged line, where the scanner does not look; those markers now sit on a comment line directly above. The riding-selection counts use the configured int dtype (int32 by default, as before), implicit_h_count casts to the dtype of the bincount it subtracts, and the two remaining literals (a row-index map and AtomGraph's documented int8 template_h_count) carry reasons. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01KwtigurBqaYuzV6426n5XG --- .../ensemble/quasi_crystal_amber.py | 3 +- torchref/model/model.py | 1 + torchref/model/riding_xyz.py | 30 ++++++++++++------- torchref/topology/atom_graph.py | 2 +- torchref/topology/build.py | 1 + torchref/topology/hydrogens.py | 12 +++++--- 6 files changed, 33 insertions(+), 16 deletions(-) diff --git a/torchref/experimental/ensemble/quasi_crystal_amber.py b/torchref/experimental/ensemble/quasi_crystal_amber.py index ee9986ca..150dda96 100644 --- a/torchref/experimental/ensemble/quasi_crystal_amber.py +++ b/torchref/experimental/ensemble/quasi_crystal_amber.py @@ -575,8 +575,9 @@ def _ensure_torch_buffers(self, device: torch.device, dtype: torch.dtype) -> Non # that survived the special-position filter — used in forward to # subset ``xyz_per_member`` before applying the layout transform. self._keep_atom_idx_torch = torch.from_numpy(self._keep_atom_idx_np).to( + # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 device=device, dtype=torch.long - ) # dtype-ok: atom/copy index tensor for indexing; PyTorch requires int64 + ) self._omm_to_model = self._omm_to_model.to(device) self._buffers_device = device diff --git a/torchref/model/model.py b/torchref/model/model.py index b66b8d17..8380b2f6 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -3011,6 +3011,7 @@ def _complete_riding_waters(self, frames): self.pdb, plan, restraints.topology ) frames = generated.remap(old_rows).fill_planned_rows(new_rows) + # dtype-ok: row index map into the augmented table; torch indexing needs int64 source = torch.empty(len(augmented), dtype=torch.long, device=self.device) old_index = torch.as_tensor(old_rows, device=self.device) new_index = torch.as_tensor(new_rows, device=self.device) diff --git a/torchref/model/riding_xyz.py b/torchref/model/riding_xyz.py index 12771a67..9d02c619 100644 --- a/torchref/model/riding_xyz.py +++ b/torchref/model/riding_xyz.py @@ -30,6 +30,7 @@ place_local_frame, rotate_vectors, ) +from torchref.config import get_int_dtype from torchref.model.parameter_wrappers import MixedTensor from torchref.topology.hydrogens import HydrogenFrames @@ -66,8 +67,9 @@ def _register_rows(self, n_full: int, frames: HydrogenFrames, device) -> None: raise ValueError(f"{name} must reference stored rows, not riding ones") long = dict( + # dtype-ok: row index buffers; int64 index required dtype=torch.int64, device=device - ) # dtype-ok: row index buffers; int64 index required + ) self.register_buffer("base_row", torch.as_tensor(base, **long)) self.register_buffer("h_row", torch.as_tensor(h, **long)) self.register_buffer( @@ -96,25 +98,30 @@ def _rebuild_row_cache(self) -> None: n_full = int(base.numel() + self.h_row.numel()) self._n_full = n_full full_to_base = torch.full( + # dtype-ok: index map; int64 (max(n_full, 1),), -1, dtype=torch.int64, device=device - ) # dtype-ok: index map; int64 + ) full_to_base[base] = torch.arange( + # dtype-ok: index map; int64 base.numel(), dtype=torch.int64, device=device - ) # dtype-ok: index map; int64 + ) self._parent_bidx = full_to_base[self.parent_row.clamp(min=0)].clamp(min=0) self._n1_bidx = full_to_base[self.n1_row.clamp(min=0)].clamp(min=0) self._n2_bidx = full_to_base[self.n2_row.clamp(min=0)].clamp(min=0) # ``cat([base, derived])[gather]`` lays the full table out in one gather. order = torch.empty( + # dtype-ok: gather index; int64 n_full, dtype=torch.int64, device=device - ) # dtype-ok: gather index; int64 + ) order[base] = torch.arange( + # dtype-ok: gather index; int64 base.numel(), dtype=torch.int64, device=device - ) # dtype-ok: gather index; int64 + ) order[self.h_row] = base.numel() + torch.arange( self.h_row.numel(), + # dtype-ok: gather index; int64 dtype=torch.int64, - device=device, # dtype-ok: gather index; int64 + device=device, ) self._gather_order = order @@ -226,8 +233,9 @@ def __init__( self.register_buffer( buffer, torch.zeros( + # dtype-ok: empty row-index buffer; int64 0, dtype=torch.int64, device=self.device - ), # dtype-ok: empty row-index buffer; int64 + ), ) self.register_buffer( "frame_valid", torch.zeros(0, dtype=torch.bool, device=self.device) @@ -261,8 +269,9 @@ def __init__( is_riding = np.zeros(n_full, dtype=bool) is_riding[np.asarray(frames.h_row, dtype=np.int64)] = True base_rows = torch.as_tensor( + # dtype-ok: row index; int64 np.nonzero(~is_riding)[0], dtype=torch.int64, device=device - ) # dtype-ok: row index; int64 + ) if refinable_mask is None: base_mask = None @@ -380,8 +389,9 @@ def _orientation_selection(self, full_mask): parents = getattr(self, "_" + kind + "_parents") rows = getattr(self, "_" + kind + "_h") groups = getattr(self, "_" + kind + "_inverse") - selected = full_mask[parents].to(torch.int32) - selected.index_add_(0, groups, full_mask[self.h_row[rows]].to(torch.int32)) + int_dtype = get_int_dtype() + selected = full_mask[parents].to(int_dtype) + selected.index_add_(0, groups, full_mask[self.h_row[rows]].to(int_dtype)) selections.append(selected > 0) return selections diff --git a/torchref/topology/atom_graph.py b/torchref/topology/atom_graph.py index dcfce9b6..a38eb5be 100644 --- a/torchref/topology/atom_graph.py +++ b/torchref/topology/atom_graph.py @@ -224,7 +224,7 @@ def implicit_h_count(self) -> Optional[torch.Tensor]: if heavy_of_h.numel(): present = torch.bincount(heavy_of_h, minlength=self.n_atoms) known = self.template_h_count >= 0 - missing = self.template_h_count.to(torch.int64) - present + missing = self.template_h_count.to(present.dtype) - present return torch.where(known, missing.clamp(min=0), torch.zeros_like(missing)) def subset(self, remap: torch.Tensor, residue_remap: torch.Tensor) -> "AtomGraph": diff --git a/torchref/topology/build.py b/torchref/topology/build.py index ff9452a6..9c1f6519 100644 --- a/torchref/topology/build.py +++ b/torchref/topology/build.py @@ -1007,6 +1007,7 @@ def build_topology_with_values( planes=plane_blocks, energy_type=energy_type, template_h_count=torch.as_tensor( + # dtype-ok: small per-atom count; int8 is AtomGraph's documented storage template_h_count, dtype=torch.int8, device=device ), ) diff --git a/torchref/topology/hydrogens.py b/torchref/topology/hydrogens.py index 5ec85303..4325836e 100644 --- a/torchref/topology/hydrogens.py +++ b/torchref/topology/hydrogens.py @@ -1146,17 +1146,21 @@ def to_tensors(self, device=None) -> Dict[str, torch.Tensor]: """Return frame and orientation arrays as tensors, keyed by field name.""" return { "h_row": torch.as_tensor( + # dtype-ok: row index; int64 required self.h_row, dtype=torch.int64, device=device - ), # dtype-ok: row index; int64 required + ), "parent_row": torch.as_tensor( + # dtype-ok: row index; int64 required self.parent_row, dtype=torch.int64, device=device - ), # dtype-ok: row index; int64 required + ), "n1_row": torch.as_tensor( + # dtype-ok: row index; int64 required self.n1_row, dtype=torch.int64, device=device - ), # dtype-ok: row index; int64 required + ), "n2_row": torch.as_tensor( + # dtype-ok: row index; int64 required self.n2_row, dtype=torch.int64, device=device - ), # dtype-ok: row index; int64 required + ), "frame_valid": torch.as_tensor( self.frame_valid, dtype=torch.bool, device=device ), From a39165de8c4bc8fbe71cb34f6dbb8822b7333b93 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Sun, 27 Sep 2026 20:28:14 +0200 Subject: [PATCH 192/250] Regenerate the AF trajectory reference on structures that refine The reference was written when hydrogens were generated by default, and three of its five structures ended with R-work above R-free. All five (1VER, 6VHI, 1BYW, 6JZA, 6SXW) have free sets of 40-80 reflections, too few to tell R-free from noise: in the joint two-cycle refinement R-free rose mid-trajectory for 6VHI, 6JZA and 6SXW, and 1BYW ended below R-work. The set is now 1AK5, 1DAW, 3GR5 and 1VER (cubic, centred monoclinic, hexagonal, centred tetragonal), the first three with free sets of 560-2063 reflections and AlphaFold starts added from the Figure 2 benchmark. Each descends in R-work and R-free and ends with R-work below R-free, which test_af_starts_refine now asserts for the reference and test_af_trajectory_matches_reference for every observed run; regeneration refuses a set that does not. The R-work tolerance floor goes from 0.002 to 0.005 and is applied at test time: five runs measured 1VER's spread as 0.00045, yet a sixth run on the same machine deviated 0.0021, and GitHub runners sat up to 0.0016 from a reference written on merlin7. 0.005 stays under a quarter of every structure's descent. Three further runs on another node all pass, with deviations of at most 0.0004. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01KwtigurBqaYuzV6426n5XG --- tests/files/pdb/1AK5_af.pdb | 2519 +++++++++++++++ tests/files/pdb/1DAW_af.pdb | 2736 +++++++++++++++++ tests/files/pdb/3GR5_af.pdb | 1172 +++++++ tests/functional/af_trajectory_reference.json | 110 +- tests/functional/test_af_trajectory.py | 97 +- 5 files changed, 6539 insertions(+), 95 deletions(-) create mode 100644 tests/files/pdb/1AK5_af.pdb create mode 100644 tests/files/pdb/1DAW_af.pdb create mode 100644 tests/files/pdb/3GR5_af.pdb diff --git a/tests/files/pdb/1AK5_af.pdb b/tests/files/pdb/1AK5_af.pdb new file mode 100644 index 00000000..0a7ffbe3 --- /dev/null +++ b/tests/files/pdb/1AK5_af.pdb @@ -0,0 +1,2519 @@ +REMARK TITLE 1AK5 AlphaFold-start MR +REMARK Log-Likelihood Gain: 4582.896 +REMARK RFZ=2.5 TFZ=10.3 PAK=0 LLG=123 TFZ==12.5 LLG=4583 TFZ==52.5 PAK=0 LLG=4583 TFZ==52.5 +REMARK ENSEMBLE e_P50097 EULER 67.44 29.80 95.39 FRAC 0.191 0.085 0.383 +CRYST1 157.250 157.250 157.250 90.00 90.00 90.00 P 4 3 2 24 +SCALE1 0.006359 -0.000000 -0.000000 0.00000 +SCALE2 0.000000 0.006359 -0.000000 0.00000 +SCALE3 0.000000 0.000000 0.006359 0.00000 +ATOM 1 N ALA A 2 2.732 22.199 42.564 1.00 39.24 N +ATOM 2 CA ALA A 2 3.908 21.434 42.983 1.00 39.24 C +ATOM 3 C ALA A 2 4.028 20.119 42.188 1.00 39.24 C +ATOM 4 O ALA A 2 3.018 19.578 41.733 1.00 39.24 O +ATOM 5 CB ALA A 2 3.815 21.182 44.493 1.00 39.24 C +ATOM 6 N LYS A 3 5.252 19.604 42.028 1.00 36.24 N +ATOM 7 CA LYS A 3 5.494 18.245 41.524 1.00 36.24 C +ATOM 8 C LYS A 3 5.214 17.258 42.657 1.00 36.24 C +ATOM 9 O LYS A 3 5.785 17.392 43.735 1.00 36.24 O +ATOM 10 CB LYS A 3 6.941 18.138 41.008 1.00 36.24 C +ATOM 11 CG LYS A 3 7.293 16.767 40.397 1.00 36.24 C +ATOM 12 CD LYS A 3 8.784 16.738 40.016 1.00 36.24 C +ATOM 13 CE LYS A 3 9.223 15.398 39.408 1.00 36.24 C +ATOM 14 NZ LYS A 3 10.708 15.332 39.291 1.00 36.24 N +ATOM 15 N TYR A 4 4.355 16.274 42.408 1.00 35.95 N +ATOM 16 CA TYR A 4 4.053 15.209 43.364 1.00 35.95 C +ATOM 17 C TYR A 4 4.838 13.949 43.003 1.00 35.95 C +ATOM 18 O TYR A 4 4.823 13.512 41.852 1.00 35.95 O +ATOM 19 CB TYR A 4 2.542 14.942 43.416 1.00 35.95 C +ATOM 20 CG TYR A 4 1.745 16.062 44.061 1.00 35.95 C +ATOM 21 CD1 TYR A 4 1.358 15.967 45.412 1.00 35.95 C +ATOM 22 CD2 TYR A 4 1.400 17.207 43.316 1.00 35.95 C +ATOM 23 CE1 TYR A 4 0.633 17.011 46.020 1.00 35.95 C +ATOM 24 CE2 TYR A 4 0.682 18.256 43.918 1.00 35.95 C +ATOM 25 CZ TYR A 4 0.297 18.161 45.271 1.00 35.95 C +ATOM 26 OH TYR A 4 -0.393 19.179 45.847 1.00 35.95 O +ATOM 27 N TYR A 5 5.496 13.357 43.996 1.00 36.30 N +ATOM 28 CA TYR A 5 5.987 11.984 43.921 1.00 36.30 C +ATOM 29 C TYR A 5 4.888 11.083 44.495 1.00 36.30 C +ATOM 30 O TYR A 5 4.572 11.184 45.678 1.00 36.30 O +ATOM 31 CB TYR A 5 7.321 11.865 44.674 1.00 36.30 C +ATOM 32 CG TYR A 5 8.449 12.674 44.053 1.00 36.30 C +ATOM 33 CD1 TYR A 5 9.155 12.156 42.949 1.00 36.30 C +ATOM 34 CD2 TYR A 5 8.790 13.944 44.564 1.00 36.30 C +ATOM 35 CE1 TYR A 5 10.182 12.911 42.348 1.00 36.30 C +ATOM 36 CE2 TYR A 5 9.828 14.695 43.975 1.00 36.30 C +ATOM 37 CZ TYR A 5 10.519 14.182 42.856 1.00 36.30 C +ATOM 38 OH TYR A 5 11.483 14.916 42.233 1.00 36.30 O +ATOM 39 N ASN A 6 4.254 10.271 43.643 1.00 34.63 N +ATOM 40 CA ASN A 6 3.090 9.465 44.038 1.00 34.63 C +ATOM 41 C ASN A 6 3.470 8.196 44.816 1.00 34.63 C +ATOM 42 O ASN A 6 2.635 7.649 45.533 1.00 34.63 O +ATOM 43 CB ASN A 6 2.266 9.114 42.788 1.00 34.63 C +ATOM 44 CG ASN A 6 1.597 10.320 42.154 1.00 34.63 C +ATOM 45 OD1 ASN A 6 1.139 11.242 42.806 1.00 34.63 O +ATOM 46 ND2 ASN A 6 1.515 10.357 40.845 1.00 34.63 N +ATOM 47 N GLU A 7 4.710 7.727 44.682 1.00 34.38 N +ATOM 48 CA GLU A 7 5.233 6.625 45.486 1.00 34.38 C +ATOM 49 C GLU A 7 5.763 7.161 46.825 1.00 34.38 C +ATOM 50 O GLU A 7 6.532 8.131 46.830 1.00 34.38 O +ATOM 51 CB GLU A 7 6.327 5.863 44.728 1.00 34.38 C +ATOM 52 CG GLU A 7 5.745 5.125 43.512 1.00 34.38 C +ATOM 53 CD GLU A 7 6.767 4.236 42.790 1.00 34.38 C +ATOM 54 OE1 GLU A 7 6.321 3.492 41.888 1.00 34.38 O +ATOM 55 OE2 GLU A 7 7.974 4.330 43.106 1.00 34.38 O +ATOM 56 N PRO A 8 5.378 6.560 47.968 1.00 33.76 N +ATOM 57 CA PRO A 8 5.956 6.904 49.260 1.00 33.76 C +ATOM 58 C PRO A 8 7.479 6.757 49.250 1.00 33.76 C +ATOM 59 O PRO A 8 8.030 5.830 48.662 1.00 33.76 O +ATOM 60 CB PRO A 8 5.309 5.955 50.277 1.00 33.76 C +ATOM 61 CG PRO A 8 3.992 5.565 49.612 1.00 33.76 C +ATOM 62 CD PRO A 8 4.352 5.541 48.130 1.00 33.76 C +ATOM 63 N CYS A 9 8.185 7.651 49.940 1.00 32.53 N +ATOM 64 CA CYS A 9 9.605 7.434 50.196 1.00 32.53 C +ATOM 65 C CYS A 9 9.805 6.397 51.312 1.00 32.53 C +ATOM 66 O CYS A 9 9.071 6.402 52.299 1.00 32.53 O +ATOM 67 CB CYS A 9 10.294 8.768 50.504 1.00 32.53 C +ATOM 68 SG CYS A 9 9.633 9.487 52.039 1.00 32.53 S +ATOM 69 N HIS A 10 10.863 5.599 51.204 1.00 31.71 N +ATOM 70 CA HIS A 10 11.188 4.488 52.091 1.00 31.71 C +ATOM 71 C HIS A 10 12.486 4.716 52.879 1.00 31.71 C +ATOM 72 O HIS A 10 13.374 5.497 52.510 1.00 31.71 O +ATOM 73 CB HIS A 10 11.277 3.203 51.260 1.00 31.71 C +ATOM 74 CG HIS A 10 9.984 2.835 50.584 1.00 31.71 C +ATOM 75 ND1 HIS A 10 8.853 2.377 51.217 1.00 31.71 N +ATOM 76 CD2 HIS A 10 9.714 2.869 49.242 1.00 31.71 C +ATOM 77 CE1 HIS A 10 7.928 2.133 50.277 1.00 31.71 C +ATOM 78 NE2 HIS A 10 8.410 2.405 49.054 1.00 31.71 N +ATOM 79 N THR A 11 12.608 4.000 53.988 1.00 31.71 N +ATOM 80 CA THR A 11 13.760 3.951 54.893 1.00 31.71 C +ATOM 81 C THR A 11 14.450 2.585 54.832 1.00 31.71 C +ATOM 82 O THR A 11 13.906 1.634 54.284 1.00 31.71 O +ATOM 83 CB THR A 11 13.324 4.252 56.333 1.00 31.71 C +ATOM 84 OG1 THR A 11 12.502 3.230 56.836 1.00 31.71 O +ATOM 85 CG2 THR A 11 12.586 5.584 56.473 1.00 31.71 C +ATOM 86 N PHE A 12 15.652 2.454 55.407 1.00 31.57 N +ATOM 87 CA PHE A 12 16.382 1.176 55.385 1.00 31.57 C +ATOM 88 C PHE A 12 15.659 0.036 56.124 1.00 31.57 C +ATOM 89 O PHE A 12 15.826 -1.113 55.734 1.00 31.57 O +ATOM 90 CB PHE A 12 17.790 1.366 55.965 1.00 31.57 C +ATOM 91 CG PHE A 12 18.650 2.391 55.249 1.00 31.57 C +ATOM 92 CD1 PHE A 12 18.953 2.237 53.883 1.00 31.57 C +ATOM 93 CD2 PHE A 12 19.171 3.489 55.956 1.00 31.57 C +ATOM 94 CE1 PHE A 12 19.756 3.186 53.228 1.00 31.57 C +ATOM 95 CE2 PHE A 12 19.990 4.428 55.306 1.00 31.57 C +ATOM 96 CZ PHE A 12 20.275 4.280 53.938 1.00 31.57 C +ATOM 97 N ASN A 13 14.830 0.345 57.132 1.00 33.41 N +ATOM 98 CA ASN A 13 14.043 -0.651 57.878 1.00 33.41 C +ATOM 99 C ASN A 13 13.064 -1.433 56.986 1.00 33.41 C +ATOM 100 O ASN A 13 12.619 -2.513 57.360 1.00 33.41 O +ATOM 101 CB ASN A 13 13.211 0.072 58.956 1.00 33.41 C +ATOM 102 CG ASN A 13 13.951 0.431 60.228 1.00 33.41 C +ATOM 103 OD1 ASN A 13 14.867 -0.217 60.691 1.00 33.41 O +ATOM 104 ND2 ASN A 13 13.535 1.481 60.895 1.00 33.41 N +ATOM 105 N GLU A 14 12.685 -0.874 55.838 1.00 32.50 N +ATOM 106 CA GLU A 14 11.671 -1.443 54.948 1.00 32.50 C +ATOM 107 C GLU A 14 12.274 -2.369 53.887 1.00 32.50 C +ATOM 108 O GLU A 14 11.549 -2.854 53.022 1.00 32.50 O +ATOM 109 CB GLU A 14 10.866 -0.302 54.310 1.00 32.50 C +ATOM 110 CG GLU A 14 10.074 0.511 55.343 1.00 32.50 C +ATOM 111 CD GLU A 14 9.629 1.838 54.730 1.00 32.50 C +ATOM 112 OE1 GLU A 14 8.640 1.848 53.965 1.00 32.50 O +ATOM 113 OE2 GLU A 14 10.335 2.845 54.994 1.00 32.50 O +ATOM 114 N TYR A 15 13.583 -2.637 53.938 1.00 31.64 N +ATOM 115 CA TYR A 15 14.277 -3.446 52.939 1.00 31.64 C +ATOM 116 C TYR A 15 15.049 -4.617 53.545 1.00 31.64 C +ATOM 117 O TYR A 15 15.629 -4.517 54.624 1.00 31.64 O +ATOM 118 CB TYR A 15 15.211 -2.564 52.106 1.00 31.64 C +ATOM 119 CG TYR A 15 14.498 -1.569 51.215 1.00 31.64 C +ATOM 120 CD1 TYR A 15 13.952 -1.978 49.985 1.00 31.64 C +ATOM 121 CD2 TYR A 15 14.379 -0.228 51.618 1.00 31.64 C +ATOM 122 CE1 TYR A 15 13.310 -1.050 49.142 1.00 31.64 C +ATOM 123 CE2 TYR A 15 13.748 0.705 50.779 1.00 31.64 C +ATOM 124 CZ TYR A 15 13.209 0.300 49.540 1.00 31.64 C +ATOM 125 OH TYR A 15 12.599 1.211 48.740 1.00 31.64 O +ATOM 126 N LEU A 16 15.141 -5.706 52.779 1.00 31.68 N +ATOM 127 CA LEU A 16 16.048 -6.826 53.034 1.00 31.68 C +ATOM 128 C LEU A 16 16.911 -7.126 51.804 1.00 31.68 C +ATOM 129 O LEU A 16 16.492 -6.922 50.662 1.00 31.68 O +ATOM 130 CB LEU A 16 15.261 -8.078 53.461 1.00 31.68 C +ATOM 131 CG LEU A 16 14.559 -7.989 54.829 1.00 31.68 C +ATOM 132 CD1 LEU A 16 13.839 -9.314 55.100 1.00 31.68 C +ATOM 133 CD2 LEU A 16 15.532 -7.756 55.986 1.00 31.68 C +ATOM 134 N LEU A 17 18.107 -7.663 52.048 1.00 31.30 N +ATOM 135 CA LEU A 17 18.954 -8.261 51.015 1.00 31.30 C +ATOM 136 C LEU A 17 18.542 -9.719 50.796 1.00 31.30 C +ATOM 137 O LEU A 17 18.507 -10.508 51.738 1.00 31.30 O +ATOM 138 CB LEU A 17 20.439 -8.161 51.424 1.00 31.30 C +ATOM 139 CG LEU A 17 21.012 -6.744 51.262 1.00 31.30 C +ATOM 140 CD1 LEU A 17 22.287 -6.566 52.085 1.00 31.30 C +ATOM 141 CD2 LEU A 17 21.358 -6.446 49.804 1.00 31.30 C +ATOM 142 N ILE A 18 18.278 -10.090 49.546 1.00 31.20 N +ATOM 143 CA ILE A 18 18.064 -11.477 49.135 1.00 31.20 C +ATOM 144 C ILE A 18 19.440 -12.110 48.872 1.00 31.20 C +ATOM 145 O ILE A 18 20.173 -11.598 48.019 1.00 31.20 O +ATOM 146 CB ILE A 18 17.155 -11.570 47.889 1.00 31.20 C +ATOM 147 CG1 ILE A 18 15.790 -10.883 48.124 1.00 31.20 C +ATOM 148 CG2 ILE A 18 16.967 -13.046 47.481 1.00 31.20 C +ATOM 149 CD1 ILE A 18 14.905 -10.839 46.872 1.00 31.20 C +ATOM 150 N PRO A 19 19.783 -13.242 49.519 1.00 31.37 N +ATOM 151 CA PRO A 19 21.044 -13.937 49.279 1.00 31.37 C +ATOM 152 C PRO A 19 21.322 -14.220 47.792 1.00 31.37 C +ATOM 153 O PRO A 19 20.425 -14.522 46.991 1.00 31.37 O +ATOM 154 CB PRO A 19 20.982 -15.216 50.117 1.00 31.37 C +ATOM 155 CG PRO A 19 20.076 -14.816 51.280 1.00 31.37 C +ATOM 156 CD PRO A 19 19.066 -13.881 50.617 1.00 31.37 C +ATOM 157 N GLY A 20 22.594 -14.098 47.424 1.00 31.97 N +ATOM 158 CA GLY A 20 23.142 -14.425 46.111 1.00 31.97 C +ATOM 159 C GLY A 20 23.911 -15.747 46.140 1.00 31.97 C +ATOM 160 O GLY A 20 23.982 -16.421 47.163 1.00 31.97 O +ATOM 161 N LEU A 21 24.506 -16.120 45.008 1.00 32.00 N +ATOM 162 CA LEU A 21 25.427 -17.253 44.960 1.00 32.00 C +ATOM 163 C LEU A 21 26.746 -16.868 45.650 1.00 32.00 C +ATOM 164 O LEU A 21 27.481 -16.024 45.136 1.00 32.00 O +ATOM 165 CB LEU A 21 25.641 -17.663 43.491 1.00 32.00 C +ATOM 166 CG LEU A 21 26.617 -18.842 43.303 1.00 32.00 C +ATOM 167 CD1 LEU A 21 26.047 -20.150 43.853 1.00 32.00 C +ATOM 168 CD2 LEU A 21 26.911 -19.027 41.814 1.00 32.00 C +ATOM 169 N SER A 22 27.061 -17.501 46.779 1.00 32.58 N +ATOM 170 CA SER A 22 28.381 -17.432 47.408 1.00 32.58 C +ATOM 171 C SER A 22 29.310 -18.490 46.803 1.00 32.58 C +ATOM 172 O SER A 22 28.935 -19.644 46.609 1.00 32.58 O +ATOM 173 CB SER A 22 28.267 -17.560 48.931 1.00 32.58 C +ATOM 174 OG SER A 22 27.572 -18.735 49.297 1.00 32.58 O +ATOM 175 N THR A 23 30.529 -18.086 46.456 1.00 32.11 N +ATOM 176 CA THR A 23 31.574 -19.002 45.966 1.00 32.11 C +ATOM 177 C THR A 23 32.440 -19.484 47.131 1.00 32.11 C +ATOM 178 O THR A 23 32.324 -18.970 48.241 1.00 32.11 O +ATOM 179 CB THR A 23 32.441 -18.348 44.879 1.00 32.11 C +ATOM 180 OG1 THR A 23 33.194 -17.294 45.423 1.00 32.11 O +ATOM 181 CG2 THR A 23 31.623 -17.801 43.708 1.00 32.11 C +ATOM 182 N VAL A 24 33.351 -20.431 46.886 1.00 31.54 N +ATOM 183 CA VAL A 24 34.336 -20.864 47.898 1.00 31.54 C +ATOM 184 C VAL A 24 35.224 -19.718 48.401 1.00 31.54 C +ATOM 185 O VAL A 24 35.702 -19.773 49.530 1.00 31.54 O +ATOM 186 CB VAL A 24 35.214 -22.016 47.376 1.00 31.54 C +ATOM 187 CG1 VAL A 24 34.364 -23.267 47.113 1.00 31.54 C +ATOM 188 CG2 VAL A 24 35.976 -21.662 46.090 1.00 31.54 C +ATOM 189 N ASP A 25 35.381 -18.652 47.609 1.00 31.93 N +ATOM 190 CA ASP A 25 36.158 -17.467 47.981 1.00 31.93 C +ATOM 191 C ASP A 25 35.385 -16.496 48.882 1.00 31.93 C +ATOM 192 O ASP A 25 35.996 -15.583 49.443 1.00 31.93 O +ATOM 193 CB ASP A 25 36.624 -16.730 46.717 1.00 31.93 C +ATOM 194 CG ASP A 25 37.537 -17.578 45.833 1.00 31.93 C +ATOM 195 OD1 ASP A 25 38.394 -18.292 46.398 1.00 31.93 O +ATOM 196 OD2 ASP A 25 37.361 -17.498 44.597 1.00 31.93 O +ATOM 197 N CYS A 26 34.063 -16.670 49.038 1.00 31.37 N +ATOM 198 CA CYS A 26 33.183 -15.846 49.877 1.00 31.37 C +ATOM 199 C CYS A 26 33.400 -16.110 51.375 1.00 31.37 C +ATOM 200 O CYS A 26 32.491 -16.498 52.108 1.00 31.37 O +ATOM 201 CB CYS A 26 31.720 -15.991 49.434 1.00 31.37 C +ATOM 202 SG CYS A 26 31.473 -15.238 47.796 1.00 31.37 S +ATOM 203 N ILE A 27 34.620 -15.845 51.830 1.00 31.61 N +ATOM 204 CA ILE A 27 35.078 -15.926 53.212 1.00 31.61 C +ATOM 205 C ILE A 27 35.259 -14.486 53.718 1.00 31.61 C +ATOM 206 O ILE A 27 35.887 -13.688 53.022 1.00 31.61 O +ATOM 207 CB ILE A 27 36.391 -16.736 53.281 1.00 31.61 C +ATOM 208 CG1 ILE A 27 36.213 -18.159 52.695 1.00 31.61 C +ATOM 209 CG2 ILE A 27 36.879 -16.834 54.739 1.00 31.61 C +ATOM 210 CD1 ILE A 27 37.540 -18.877 52.424 1.00 31.61 C +ATOM 211 N PRO A 28 34.764 -14.116 54.918 1.00 32.00 N +ATOM 212 CA PRO A 28 34.841 -12.738 55.414 1.00 32.00 C +ATOM 213 C PRO A 28 36.246 -12.111 55.404 1.00 32.00 C +ATOM 214 O PRO A 28 36.372 -10.903 55.214 1.00 32.00 O +ATOM 215 CB PRO A 28 34.272 -12.801 56.834 1.00 32.00 C +ATOM 216 CG PRO A 28 33.249 -13.930 56.744 1.00 32.00 C +ATOM 217 CD PRO A 28 33.920 -14.921 55.795 1.00 32.00 C +ATOM 218 N SER A 29 37.310 -12.905 55.563 1.00 32.82 N +ATOM 219 CA SER A 29 38.704 -12.438 55.475 1.00 32.82 C +ATOM 220 C SER A 29 39.113 -11.957 54.078 1.00 32.82 C +ATOM 221 O SER A 29 40.023 -11.143 53.969 1.00 32.82 O +ATOM 222 CB SER A 29 39.654 -13.561 55.899 1.00 32.82 C +ATOM 223 OG SER A 29 39.446 -14.713 55.101 1.00 32.82 O +ATOM 224 N ASN A 30 38.440 -12.434 53.026 1.00 31.57 N +ATOM 225 CA ASN A 30 38.699 -12.059 51.635 1.00 31.57 C +ATOM 226 C ASN A 30 37.927 -10.797 51.213 1.00 31.57 C +ATOM 227 O ASN A 30 38.130 -10.292 50.111 1.00 31.57 O +ATOM 228 CB ASN A 30 38.347 -13.247 50.719 1.00 31.57 C +ATOM 229 CG ASN A 30 39.212 -14.477 50.924 1.00 31.57 C +ATOM 230 OD1 ASN A 30 40.196 -14.483 51.649 1.00 31.57 O +ATOM 231 ND2 ASN A 30 38.849 -15.565 50.290 1.00 31.57 N +ATOM 232 N VAL A 31 37.026 -10.290 52.063 1.00 31.14 N +ATOM 233 CA VAL A 31 36.220 -9.106 51.755 1.00 31.14 C +ATOM 234 C VAL A 31 37.050 -7.836 51.938 1.00 31.14 C +ATOM 235 O VAL A 31 37.509 -7.527 53.041 1.00 31.14 O +ATOM 236 CB VAL A 31 34.935 -9.052 52.598 1.00 31.14 C +ATOM 237 CG1 VAL A 31 34.119 -7.792 52.281 1.00 31.14 C +ATOM 238 CG2 VAL A 31 34.037 -10.265 52.334 1.00 31.14 C +ATOM 239 N ASN A 32 37.174 -7.056 50.866 1.00 31.33 N +ATOM 240 CA ASN A 32 37.784 -5.734 50.873 1.00 31.33 C +ATOM 241 C ASN A 32 36.705 -4.645 50.978 1.00 31.33 C +ATOM 242 O ASN A 32 35.862 -4.487 50.095 1.00 31.33 O +ATOM 243 CB ASN A 32 38.671 -5.584 49.629 1.00 31.33 C +ATOM 244 CG ASN A 32 39.462 -4.287 49.619 1.00 31.33 C +ATOM 245 OD1 ASN A 32 39.299 -3.399 50.442 1.00 31.33 O +ATOM 246 ND2 ASN A 32 40.371 -4.135 48.688 1.00 31.33 N +ATOM 247 N LEU A 33 36.769 -3.877 52.065 1.00 31.07 N +ATOM 248 CA LEU A 33 35.838 -2.791 52.380 1.00 31.07 C +ATOM 249 C LEU A 33 36.354 -1.400 51.976 1.00 31.07 C +ATOM 250 O LEU A 33 35.713 -0.399 52.297 1.00 31.07 O +ATOM 251 CB LEU A 33 35.519 -2.861 53.880 1.00 31.07 C +ATOM 252 CG LEU A 33 34.766 -4.113 54.350 1.00 31.07 C +ATOM 253 CD1 LEU A 33 34.471 -3.950 55.838 1.00 31.07 C +ATOM 254 CD2 LEU A 33 33.426 -4.296 53.641 1.00 31.07 C +ATOM 255 N SER A 34 37.502 -1.321 51.298 1.00 31.26 N +ATOM 256 CA SER A 34 38.084 -0.043 50.891 1.00 31.26 C +ATOM 257 C SER A 34 37.167 0.691 49.914 1.00 31.26 C +ATOM 258 O SER A 34 36.575 0.072 49.027 1.00 31.26 O +ATOM 259 CB SER A 34 39.470 -0.199 50.258 1.00 31.26 C +ATOM 260 OG SER A 34 40.367 -0.896 51.106 1.00 31.26 O +ATOM 261 N THR A 35 37.067 2.011 50.063 1.00 31.07 N +ATOM 262 CA THR A 35 36.122 2.843 49.307 1.00 31.07 C +ATOM 263 C THR A 35 36.681 4.249 49.053 1.00 31.07 C +ATOM 264 O THR A 35 37.424 4.762 49.902 1.00 31.07 O +ATOM 265 CB THR A 35 34.779 2.898 50.052 1.00 31.07 C +ATOM 266 OG1 THR A 35 33.818 3.541 49.274 1.00 31.07 O +ATOM 267 CG2 THR A 35 34.831 3.611 51.400 1.00 31.07 C +ATOM 268 N PRO A 36 36.377 4.888 47.908 1.00 30.82 N +ATOM 269 CA PRO A 36 36.833 6.245 47.627 1.00 30.82 C +ATOM 270 C PRO A 36 36.147 7.272 48.538 1.00 30.82 C +ATOM 271 O PRO A 36 34.932 7.254 48.718 1.00 30.82 O +ATOM 272 CB PRO A 36 36.507 6.473 46.149 1.00 30.82 C +ATOM 273 CG PRO A 36 35.272 5.602 45.925 1.00 30.82 C +ATOM 274 CD PRO A 36 35.592 4.384 46.782 1.00 30.82 C +ATOM 275 N LEU A 37 36.928 8.210 49.082 1.00 30.76 N +ATOM 276 CA LEU A 37 36.406 9.353 49.838 1.00 30.76 C +ATOM 277 C LEU A 37 36.179 10.585 48.957 1.00 30.76 C +ATOM 278 O LEU A 37 35.319 11.398 49.269 1.00 30.76 O +ATOM 279 CB LEU A 37 37.380 9.677 50.983 1.00 30.76 C +ATOM 280 CG LEU A 37 36.929 10.817 51.920 1.00 30.76 C +ATOM 281 CD1 LEU A 37 35.600 10.521 52.615 1.00 30.76 C +ATOM 282 CD2 LEU A 37 37.999 11.037 52.987 1.00 30.76 C +ATOM 283 N VAL A 38 36.965 10.744 47.888 1.00 30.64 N +ATOM 284 CA VAL A 38 36.961 11.932 47.022 1.00 30.64 C +ATOM 285 C VAL A 38 36.831 11.546 45.550 1.00 30.64 C +ATOM 286 O VAL A 38 37.079 10.402 45.163 1.00 30.64 O +ATOM 287 CB VAL A 38 38.195 12.826 47.261 1.00 30.64 C +ATOM 288 CG1 VAL A 38 38.209 13.365 48.696 1.00 30.64 C +ATOM 289 CG2 VAL A 38 39.530 12.123 46.989 1.00 30.64 C +ATOM 290 N LYS A 39 36.417 12.506 44.723 1.00 30.76 N +ATOM 291 CA LYS A 39 36.183 12.304 43.295 1.00 30.76 C +ATOM 292 C LYS A 39 37.448 11.871 42.560 1.00 30.76 C +ATOM 293 O LYS A 39 38.559 12.257 42.921 1.00 30.76 O +ATOM 294 CB LYS A 39 35.522 13.545 42.665 1.00 30.76 C +ATOM 295 CG LYS A 39 36.513 14.708 42.492 1.00 30.76 C +ATOM 296 CD LYS A 39 35.852 15.993 41.985 1.00 30.76 C +ATOM 297 CE LYS A 39 36.948 17.052 41.812 1.00 30.76 C +ATOM 298 NZ LYS A 39 36.414 18.431 41.786 1.00 30.76 N +ATOM 299 N PHE A 40 37.271 11.116 41.483 1.00 30.76 N +ATOM 300 CA PHE A 40 38.362 10.687 40.608 1.00 30.76 C +ATOM 301 C PHE A 40 37.895 10.525 39.163 1.00 30.76 C +ATOM 302 O PHE A 40 36.705 10.369 38.888 1.00 30.76 O +ATOM 303 CB PHE A 40 38.991 9.396 41.145 1.00 30.76 C +ATOM 304 CG PHE A 40 38.061 8.201 41.214 1.00 30.76 C +ATOM 305 CD1 PHE A 40 37.207 8.026 42.319 1.00 30.76 C +ATOM 306 CD2 PHE A 40 38.051 7.256 40.175 1.00 30.76 C +ATOM 307 CE1 PHE A 40 36.344 6.919 42.370 1.00 30.76 C +ATOM 308 CE2 PHE A 40 37.179 6.156 40.218 1.00 30.76 C +ATOM 309 CZ PHE A 40 36.320 5.989 41.317 1.00 30.76 C +ATOM 310 N GLN A 41 38.837 10.570 38.224 1.00 31.10 N +ATOM 311 CA GLN A 41 38.543 10.346 36.810 1.00 31.10 C +ATOM 312 C GLN A 41 38.473 8.851 36.491 1.00 31.10 C +ATOM 313 O GLN A 41 39.174 8.032 37.083 1.00 31.10 O +ATOM 314 CB GLN A 41 39.589 11.029 35.924 1.00 31.10 C +ATOM 315 CG GLN A 41 39.601 12.555 36.089 1.00 31.10 C +ATOM 316 CD GLN A 41 40.601 13.237 35.159 1.00 31.10 C +ATOM 317 OE1 GLN A 41 41.217 12.642 34.290 1.00 31.10 O +ATOM 318 NE2 GLN A 41 40.813 14.526 35.311 1.00 31.10 N +ATOM 319 N LYS A 42 37.665 8.485 35.499 1.00 31.89 N +ATOM 320 CA LYS A 42 37.533 7.121 34.989 1.00 31.89 C +ATOM 321 C LYS A 42 38.910 6.548 34.643 1.00 31.89 C +ATOM 322 O LYS A 42 39.676 7.149 33.897 1.00 31.89 O +ATOM 323 CB LYS A 42 36.596 7.151 33.775 1.00 31.89 C +ATOM 324 CG LYS A 42 36.249 5.752 33.252 1.00 31.89 C +ATOM 325 CD LYS A 42 35.139 5.868 32.200 1.00 31.89 C +ATOM 326 CE LYS A 42 34.721 4.494 31.669 1.00 31.89 C +ATOM 327 NZ LYS A 42 33.551 4.619 30.763 1.00 31.89 N +ATOM 328 N GLY A 43 39.211 5.369 35.186 1.00 32.30 N +ATOM 329 CA GLY A 43 40.508 4.703 35.018 1.00 32.30 C +ATOM 330 C GLY A 43 41.593 5.144 36.007 1.00 32.30 C +ATOM 331 O GLY A 43 42.663 4.543 36.020 1.00 32.30 O +ATOM 332 N GLN A 44 41.326 6.138 36.857 1.00 32.11 N +ATOM 333 CA GLN A 44 42.212 6.549 37.948 1.00 32.11 C +ATOM 334 C GLN A 44 41.704 6.028 39.299 1.00 32.11 C +ATOM 335 O GLN A 44 40.598 5.500 39.409 1.00 32.11 O +ATOM 336 CB GLN A 44 42.393 8.075 37.947 1.00 32.11 C +ATOM 337 CG GLN A 44 42.945 8.594 36.608 1.00 32.11 C +ATOM 338 CD GLN A 44 43.230 10.091 36.621 1.00 32.11 C +ATOM 339 OE1 GLN A 44 43.161 10.769 37.632 1.00 32.11 O +ATOM 340 NE2 GLN A 44 43.531 10.675 35.483 1.00 32.11 N +ATOM 341 N GLN A 45 42.537 6.151 40.330 1.00 33.95 N +ATOM 342 CA GLN A 45 42.169 5.864 41.715 1.00 33.95 C +ATOM 343 C GLN A 45 41.848 7.171 42.441 1.00 33.95 C +ATOM 344 O GLN A 45 42.450 8.203 42.147 1.00 33.95 O +ATOM 345 CB GLN A 45 43.312 5.113 42.416 1.00 33.95 C +ATOM 346 CG GLN A 45 43.545 3.705 41.844 1.00 33.95 C +ATOM 347 CD GLN A 45 42.367 2.770 42.095 1.00 33.95 C +ATOM 348 OE1 GLN A 45 41.851 2.656 43.193 1.00 33.95 O +ATOM 349 NE2 GLN A 45 41.891 2.062 41.096 1.00 33.95 N +ATOM 350 N SER A 46 40.941 7.114 43.418 1.00 31.54 N +ATOM 351 CA SER A 46 40.740 8.227 44.349 1.00 31.54 C +ATOM 352 C SER A 46 42.039 8.558 45.079 1.00 31.54 C +ATOM 353 O SER A 46 42.755 7.656 45.518 1.00 31.54 O +ATOM 354 CB SER A 46 39.624 7.900 45.341 1.00 31.54 C +ATOM 355 OG SER A 46 39.488 8.950 46.271 1.00 31.54 O +ATOM 356 N GLU A 47 42.333 9.848 45.237 1.00 32.34 N +ATOM 357 CA GLU A 47 43.499 10.319 45.993 1.00 32.34 C +ATOM 358 C GLU A 47 43.402 9.925 47.473 1.00 32.34 C +ATOM 359 O GLU A 47 44.405 9.603 48.105 1.00 32.34 O +ATOM 360 CB GLU A 47 43.597 11.841 45.822 1.00 32.34 C +ATOM 361 CG GLU A 47 44.896 12.433 46.389 1.00 32.34 C +ATOM 362 CD GLU A 47 44.986 13.950 46.164 1.00 32.34 C +ATOM 363 OE1 GLU A 47 45.661 14.627 46.963 1.00 32.34 O +ATOM 364 OE2 GLU A 47 44.348 14.458 45.210 1.00 32.34 O +ATOM 365 N ILE A 48 42.176 9.875 48.005 1.00 31.68 N +ATOM 366 CA ILE A 48 41.896 9.435 49.370 1.00 31.68 C +ATOM 367 C ILE A 48 40.934 8.248 49.321 1.00 31.68 C +ATOM 368 O ILE A 48 39.812 8.359 48.823 1.00 31.68 O +ATOM 369 CB ILE A 48 41.375 10.596 50.250 1.00 31.68 C +ATOM 370 CG1 ILE A 48 42.319 11.824 50.192 1.00 31.68 C +ATOM 371 CG2 ILE A 48 41.233 10.071 51.691 1.00 31.68 C +ATOM 372 CD1 ILE A 48 41.830 13.042 50.986 1.00 31.68 C +ATOM 373 N ASN A 49 41.366 7.105 49.850 1.00 31.64 N +ATOM 374 CA ASN A 49 40.529 5.916 50.014 1.00 31.64 C +ATOM 375 C ASN A 49 40.466 5.545 51.489 1.00 31.64 C +ATOM 376 O ASN A 49 41.509 5.388 52.125 1.00 31.64 O +ATOM 377 CB ASN A 49 41.069 4.736 49.193 1.00 31.64 C +ATOM 378 CG ASN A 49 40.927 4.952 47.705 1.00 31.64 C +ATOM 379 OD1 ASN A 49 39.895 4.690 47.115 1.00 31.64 O +ATOM 380 ND2 ASN A 49 41.953 5.457 47.069 1.00 31.64 N +ATOM 381 N LEU A 50 39.257 5.351 52.005 1.00 31.01 N +ATOM 382 CA LEU A 50 39.053 4.760 53.322 1.00 31.01 C +ATOM 383 C LEU A 50 39.351 3.260 53.255 1.00 31.01 C +ATOM 384 O LEU A 50 39.147 2.634 52.213 1.00 31.01 O +ATOM 385 CB LEU A 50 37.619 5.007 53.810 1.00 31.01 C +ATOM 386 CG LEU A 50 37.116 6.458 53.754 1.00 31.01 C +ATOM 387 CD1 LEU A 50 35.665 6.515 54.233 1.00 31.01 C +ATOM 388 CD2 LEU A 50 37.961 7.372 54.634 1.00 31.01 C +ATOM 389 N LYS A 51 39.806 2.668 54.363 1.00 31.57 N +ATOM 390 CA LYS A 51 39.984 1.206 54.473 1.00 31.57 C +ATOM 391 C LYS A 51 38.705 0.504 54.916 1.00 31.57 C +ATOM 392 O LYS A 51 38.504 -0.668 54.603 1.00 31.57 O +ATOM 393 CB LYS A 51 41.137 0.884 55.434 1.00 31.57 C +ATOM 394 CG LYS A 51 42.487 1.484 55.024 1.00 31.57 C +ATOM 395 CD LYS A 51 42.930 0.969 53.656 1.00 31.57 C +ATOM 396 CE LYS A 51 44.340 1.474 53.373 1.00 31.57 C +ATOM 397 NZ LYS A 51 44.811 0.951 52.072 1.00 31.57 N +ATOM 398 N ILE A 52 37.837 1.235 55.611 1.00 31.01 N +ATOM 399 CA ILE A 52 36.479 0.825 55.967 1.00 31.01 C +ATOM 400 C ILE A 52 35.508 1.957 55.603 1.00 31.01 C +ATOM 401 O ILE A 52 35.880 3.126 55.691 1.00 31.01 O +ATOM 402 CB ILE A 52 36.386 0.423 57.459 1.00 31.01 C +ATOM 403 CG1 ILE A 52 36.737 1.585 58.410 1.00 31.01 C +ATOM 404 CG2 ILE A 52 37.311 -0.782 57.719 1.00 31.01 C +ATOM 405 CD1 ILE A 52 36.386 1.349 59.881 1.00 31.01 C +ATOM 406 N PRO A 53 34.250 1.671 55.238 1.00 30.79 N +ATOM 407 CA PRO A 53 33.323 2.677 54.726 1.00 30.79 C +ATOM 408 C PRO A 53 32.648 3.487 55.843 1.00 30.79 C +ATOM 409 O PRO A 53 31.504 3.904 55.690 1.00 30.79 O +ATOM 410 CB PRO A 53 32.341 1.873 53.868 1.00 30.79 C +ATOM 411 CG PRO A 53 32.230 0.561 54.638 1.00 30.79 C +ATOM 412 CD PRO A 53 33.674 0.347 55.077 1.00 30.79 C +ATOM 413 N LEU A 54 33.312 3.673 56.989 1.00 30.73 N +ATOM 414 CA LEU A 54 32.728 4.299 58.174 1.00 30.73 C +ATOM 415 C LEU A 54 33.322 5.683 58.427 1.00 30.73 C +ATOM 416 O LEU A 54 34.541 5.849 58.505 1.00 30.73 O +ATOM 417 CB LEU A 54 32.891 3.408 59.416 1.00 30.73 C +ATOM 418 CG LEU A 54 32.344 1.975 59.315 1.00 30.73 C +ATOM 419 CD1 LEU A 54 32.382 1.328 60.701 1.00 30.73 C +ATOM 420 CD2 LEU A 54 30.898 1.922 58.826 1.00 30.73 C +ATOM 421 N VAL A 55 32.435 6.655 58.619 1.00 30.64 N +ATOM 422 CA VAL A 55 32.772 8.050 58.903 1.00 30.64 C +ATOM 423 C VAL A 55 31.992 8.537 60.121 1.00 30.64 C +ATOM 424 O VAL A 55 30.790 8.282 60.229 1.00 30.64 O +ATOM 425 CB VAL A 55 32.502 8.920 57.661 1.00 30.64 C +ATOM 426 CG1 VAL A 55 32.830 10.390 57.903 1.00 30.64 C +ATOM 427 CG2 VAL A 55 33.329 8.428 56.466 1.00 30.64 C +ATOM 428 N SER A 56 32.635 9.244 61.051 1.00 30.76 N +ATOM 429 CA SER A 56 31.920 9.836 62.190 1.00 30.76 C +ATOM 430 C SER A 56 31.262 11.169 61.817 1.00 30.76 C +ATOM 431 O SER A 56 31.804 11.962 61.045 1.00 30.76 O +ATOM 432 CB SER A 56 32.809 9.940 63.428 1.00 30.76 C +ATOM 433 OG SER A 56 33.681 11.035 63.316 1.00 30.76 O +ATOM 434 N ALA A 57 30.069 11.422 62.361 1.00 31.20 N +ATOM 435 CA ALA A 57 29.293 12.616 62.032 1.00 31.20 C +ATOM 436 C ALA A 57 29.934 13.924 62.522 1.00 31.20 C +ATOM 437 O ALA A 57 30.570 13.961 63.575 1.00 31.20 O +ATOM 438 CB ALA A 57 27.865 12.456 62.556 1.00 31.20 C +ATOM 439 N ILE A 58 29.697 15.014 61.789 1.00 31.26 N +ATOM 440 CA ILE A 58 30.167 16.377 62.092 1.00 31.26 C +ATOM 441 C ILE A 58 29.369 16.940 63.277 1.00 31.26 C +ATOM 442 O ILE A 58 28.456 17.746 63.110 1.00 31.26 O +ATOM 443 CB ILE A 58 30.051 17.285 60.840 1.00 31.26 C +ATOM 444 CG1 ILE A 58 30.719 16.636 59.613 1.00 31.26 C +ATOM 445 CG2 ILE A 58 30.702 18.654 61.114 1.00 31.26 C +ATOM 446 CD1 ILE A 58 30.524 17.374 58.292 1.00 31.26 C +ATOM 447 N MET A 59 29.639 16.439 64.482 1.00 31.61 N +ATOM 448 CA MET A 59 28.840 16.742 65.667 1.00 31.61 C +ATOM 449 C MET A 59 29.690 16.842 66.930 1.00 31.61 C +ATOM 450 O MET A 59 30.566 16.001 67.159 1.00 31.61 O +ATOM 451 CB MET A 59 27.766 15.670 65.880 1.00 31.61 C +ATOM 452 CG MET A 59 26.773 15.516 64.721 1.00 31.61 C +ATOM 453 SD MET A 59 25.550 14.196 64.943 1.00 31.61 S +ATOM 454 CE MET A 59 24.776 14.763 66.474 1.00 31.61 C +ATOM 455 N GLN A 60 29.360 17.804 67.796 1.00 32.22 N +ATOM 456 CA GLN A 60 30.025 18.006 69.094 1.00 32.22 C +ATOM 457 C GLN A 60 29.986 16.761 69.980 1.00 32.22 C +ATOM 458 O GLN A 60 30.945 16.433 70.674 1.00 32.22 O +ATOM 459 CB GLN A 60 29.320 19.122 69.867 1.00 32.22 C +ATOM 460 CG GLN A 60 29.478 20.490 69.220 1.00 32.22 C +ATOM 461 CD GLN A 60 28.721 21.579 69.957 1.00 32.22 C +ATOM 462 OE1 GLN A 60 28.228 21.424 71.065 1.00 32.22 O +ATOM 463 NE2 GLN A 60 28.596 22.718 69.340 1.00 32.22 N +ATOM 464 N SER A 61 28.862 16.050 69.943 1.00 32.00 N +ATOM 465 CA SER A 61 28.617 14.860 70.754 1.00 32.00 C +ATOM 466 C SER A 61 29.278 13.589 70.205 1.00 32.00 C +ATOM 467 O SER A 61 29.203 12.535 70.838 1.00 32.00 O +ATOM 468 CB SER A 61 27.106 14.687 70.929 1.00 32.00 C +ATOM 469 OG SER A 61 26.462 14.625 69.665 1.00 32.00 O +ATOM 470 N VAL A 62 29.935 13.666 69.041 1.00 31.07 N +ATOM 471 CA VAL A 62 30.459 12.499 68.315 1.00 31.07 C +ATOM 472 C VAL A 62 31.953 12.622 68.039 1.00 31.07 C +ATOM 473 O VAL A 62 32.707 11.720 68.397 1.00 31.07 O +ATOM 474 CB VAL A 62 29.685 12.278 66.998 1.00 31.07 C +ATOM 475 CG1 VAL A 62 30.152 11.002 66.292 1.00 31.07 C +ATOM 476 CG2 VAL A 62 28.177 12.135 67.233 1.00 31.07 C +ATOM 477 N SER A 63 32.390 13.723 67.425 1.00 31.07 N +ATOM 478 CA SER A 63 33.685 13.792 66.734 1.00 31.07 C +ATOM 479 C SER A 63 34.642 14.788 67.379 1.00 31.07 C +ATOM 480 O SER A 63 34.753 15.932 66.944 1.00 31.07 O +ATOM 481 CB SER A 63 33.469 14.115 65.256 1.00 31.07 C +ATOM 482 OG SER A 63 32.720 13.070 64.677 1.00 31.07 O +ATOM 483 N GLY A 64 35.353 14.331 68.410 1.00 32.30 N +ATOM 484 CA GLY A 64 36.542 15.003 68.949 1.00 32.30 C +ATOM 485 C GLY A 64 37.837 14.264 68.596 1.00 32.30 C +ATOM 486 O GLY A 64 37.818 13.259 67.884 1.00 32.30 O +ATOM 487 N GLU A 65 38.959 14.719 69.151 1.00 32.50 N +ATOM 488 CA GLU A 65 40.308 14.185 68.888 1.00 32.50 C +ATOM 489 C GLU A 65 40.408 12.669 69.117 1.00 32.50 C +ATOM 490 O GLU A 65 40.882 11.933 68.256 1.00 32.50 O +ATOM 491 CB GLU A 65 41.302 14.874 69.829 1.00 32.50 C +ATOM 492 CG GLU A 65 41.393 16.394 69.617 1.00 32.50 C +ATOM 493 CD GLU A 65 42.115 17.108 70.768 1.00 32.50 C +ATOM 494 OE1 GLU A 65 42.276 18.338 70.644 1.00 32.50 O +ATOM 495 OE2 GLU A 65 42.415 16.439 71.786 1.00 32.50 O +ATOM 496 N LYS A 66 39.896 12.177 70.257 1.00 32.34 N +ATOM 497 CA LYS A 66 39.929 10.746 70.609 1.00 32.34 C +ATOM 498 C LYS A 66 39.186 9.883 69.589 1.00 32.34 C +ATOM 499 O LYS A 66 39.720 8.862 69.163 1.00 32.34 O +ATOM 500 CB LYS A 66 39.347 10.518 72.010 1.00 32.34 C +ATOM 501 CG LYS A 66 40.237 11.097 73.117 1.00 32.34 C +ATOM 502 CD LYS A 66 39.632 10.791 74.492 1.00 32.34 C +ATOM 503 CE LYS A 66 40.523 11.373 75.592 1.00 32.34 C +ATOM 504 NZ LYS A 66 39.953 11.112 76.936 1.00 32.34 N +ATOM 505 N MET A 67 37.989 10.309 69.180 1.00 30.97 N +ATOM 506 CA MET A 67 37.201 9.647 68.136 1.00 30.97 C +ATOM 507 C MET A 67 37.952 9.649 66.803 1.00 30.97 C +ATOM 508 O MET A 67 37.998 8.629 66.128 1.00 30.97 O +ATOM 509 CB MET A 67 35.841 10.355 67.993 1.00 30.97 C +ATOM 510 CG MET A 67 34.950 9.765 66.890 1.00 30.97 C +ATOM 511 SD MET A 67 34.447 8.041 67.104 1.00 30.97 S +ATOM 512 CE MET A 67 33.494 8.142 68.639 1.00 30.97 C +ATOM 513 N ALA A 68 38.560 10.775 66.425 1.00 31.07 N +ATOM 514 CA ALA A 68 39.331 10.890 65.191 1.00 31.07 C +ATOM 515 C ALA A 68 40.539 9.956 65.148 1.00 31.07 C +ATOM 516 O ALA A 68 40.687 9.201 64.191 1.00 31.07 O +ATOM 517 CB ALA A 68 39.685 12.360 64.982 1.00 31.07 C +ATOM 518 N ILE A 69 41.336 9.925 66.212 1.00 33.45 N +ATOM 519 CA ILE A 69 42.497 9.037 66.315 1.00 33.45 C +ATOM 520 C ILE A 69 42.058 7.569 66.310 1.00 33.45 C +ATOM 521 O ILE A 69 42.645 6.748 65.608 1.00 33.45 O +ATOM 522 CB ILE A 69 43.293 9.385 67.589 1.00 33.45 C +ATOM 523 CG1 ILE A 69 43.875 10.813 67.501 1.00 33.45 C +ATOM 524 CG2 ILE A 69 44.429 8.375 67.839 1.00 33.45 C +ATOM 525 CD1 ILE A 69 44.125 11.408 68.889 1.00 33.45 C +ATOM 526 N ALA A 70 41.029 7.219 67.088 1.00 32.94 N +ATOM 527 CA ALA A 70 40.563 5.840 67.187 1.00 32.94 C +ATOM 528 C ALA A 70 39.950 5.336 65.876 1.00 32.94 C +ATOM 529 O ALA A 70 40.269 4.233 65.447 1.00 32.94 O +ATOM 530 CB ALA A 70 39.583 5.736 68.353 1.00 32.94 C +ATOM 531 N LEU A 71 39.123 6.142 65.207 1.00 31.26 N +ATOM 532 CA LEU A 71 38.517 5.739 63.942 1.00 31.26 C +ATOM 533 C LEU A 71 39.546 5.668 62.805 1.00 31.26 C +ATOM 534 O LEU A 71 39.488 4.735 62.005 1.00 31.26 O +ATOM 535 CB LEU A 71 37.343 6.678 63.622 1.00 31.26 C +ATOM 536 CG LEU A 71 36.533 6.262 62.381 1.00 31.26 C +ATOM 537 CD1 LEU A 71 36.014 4.821 62.449 1.00 31.26 C +ATOM 538 CD2 LEU A 71 35.318 7.169 62.213 1.00 31.26 C +ATOM 539 N ALA A 72 40.522 6.583 62.768 1.00 32.07 N +ATOM 540 CA ALA A 72 41.620 6.526 61.802 1.00 32.07 C +ATOM 541 C ALA A 72 42.480 5.265 61.973 1.00 32.07 C +ATOM 542 O ALA A 72 42.898 4.675 60.978 1.00 32.07 O +ATOM 543 CB ALA A 72 42.462 7.797 61.915 1.00 32.07 C +ATOM 544 N ARG A 73 42.683 4.787 63.209 1.00 34.43 N +ATOM 545 CA ARG A 73 43.365 3.506 63.472 1.00 34.43 C +ATOM 546 C ARG A 73 42.627 2.302 62.892 1.00 34.43 C +ATOM 547 O ARG A 73 43.263 1.393 62.369 1.00 34.43 O +ATOM 548 CB ARG A 73 43.560 3.297 64.976 1.00 34.43 C +ATOM 549 CG ARG A 73 44.707 4.147 65.515 1.00 34.43 C +ATOM 550 CD ARG A 73 44.715 4.057 67.039 1.00 34.43 C +ATOM 551 NE ARG A 73 45.802 4.867 67.609 1.00 34.43 N +ATOM 552 CZ ARG A 73 45.967 5.158 68.885 1.00 34.43 C +ATOM 553 NH1 ARG A 73 45.165 4.685 69.799 1.00 34.43 N +ATOM 554 NH2 ARG A 73 46.938 5.934 69.276 1.00 34.43 N +ATOM 555 N GLU A 74 41.298 2.313 62.950 1.00 32.98 N +ATOM 556 CA GLU A 74 40.461 1.249 62.373 1.00 32.98 C +ATOM 557 C GLU A 74 40.238 1.415 60.857 1.00 32.98 C +ATOM 558 O GLU A 74 39.699 0.521 60.202 1.00 32.98 O +ATOM 559 CB GLU A 74 39.129 1.154 63.137 1.00 32.98 C +ATOM 560 CG GLU A 74 39.276 0.897 64.650 1.00 32.98 C +ATOM 561 CD GLU A 74 40.164 -0.306 65.014 1.00 32.98 C +ATOM 562 OE1 GLU A 74 40.897 -0.220 66.031 1.00 32.98 O +ATOM 563 OE2 GLU A 74 40.093 -1.328 64.298 1.00 32.98 O +ATOM 564 N GLY A 75 40.679 2.540 60.283 1.00 32.27 N +ATOM 565 CA GLY A 75 40.678 2.792 58.844 1.00 32.27 C +ATOM 566 C GLY A 75 39.523 3.625 58.301 1.00 32.27 C +ATOM 567 O GLY A 75 39.365 3.732 57.079 1.00 32.27 O +ATOM 568 N GLY A 76 38.699 4.168 59.198 1.00 31.07 N +ATOM 569 CA GLY A 76 37.687 5.167 58.878 1.00 31.07 C +ATOM 570 C GLY A 76 38.252 6.580 59.005 1.00 31.07 C +ATOM 571 O GLY A 76 39.448 6.769 59.218 1.00 31.07 O +ATOM 572 N ILE A 77 37.389 7.589 58.908 1.00 30.79 N +ATOM 573 CA ILE A 77 37.780 8.988 59.119 1.00 30.79 C +ATOM 574 C ILE A 77 36.734 9.719 59.954 1.00 30.79 C +ATOM 575 O ILE A 77 35.534 9.484 59.815 1.00 30.79 O +ATOM 576 CB ILE A 77 38.085 9.687 57.773 1.00 30.79 C +ATOM 577 CG1 ILE A 77 38.822 11.024 57.992 1.00 30.79 C +ATOM 578 CG2 ILE A 77 36.822 9.868 56.917 1.00 30.79 C +ATOM 579 CD1 ILE A 77 39.329 11.659 56.692 1.00 30.79 C +ATOM 580 N SER A 78 37.181 10.626 60.813 1.00 30.67 N +ATOM 581 CA SER A 78 36.290 11.488 61.589 1.00 30.67 C +ATOM 582 C SER A 78 36.332 12.911 61.068 1.00 30.67 C +ATOM 583 O SER A 78 37.402 13.398 60.708 1.00 30.67 O +ATOM 584 CB SER A 78 36.668 11.468 63.065 1.00 30.67 C +ATOM 585 OG SER A 78 36.542 10.148 63.562 1.00 30.67 O +ATOM 586 N PHE A 79 35.186 13.594 61.076 1.00 30.70 N +ATOM 587 CA PHE A 79 35.112 15.017 60.755 1.00 30.70 C +ATOM 588 C PHE A 79 34.907 15.856 62.019 1.00 30.70 C +ATOM 589 O PHE A 79 33.859 15.776 62.661 1.00 30.70 O +ATOM 590 CB PHE A 79 34.024 15.292 59.720 1.00 30.70 C +ATOM 591 CG PHE A 79 34.340 14.865 58.301 1.00 30.70 C +ATOM 592 CD1 PHE A 79 34.726 15.824 57.344 1.00 30.70 C +ATOM 593 CD2 PHE A 79 34.238 13.514 57.929 1.00 30.70 C +ATOM 594 CE1 PHE A 79 35.027 15.429 56.030 1.00 30.70 C +ATOM 595 CE2 PHE A 79 34.537 13.120 56.612 1.00 30.70 C +ATOM 596 CZ PHE A 79 34.929 14.077 55.663 1.00 30.70 C +ATOM 597 N ILE A 80 35.893 16.686 62.363 1.00 30.79 N +ATOM 598 CA ILE A 80 35.821 17.588 63.520 1.00 30.79 C +ATOM 599 C ILE A 80 34.742 18.649 63.286 1.00 30.79 C +ATOM 600 O ILE A 80 34.729 19.318 62.250 1.00 30.79 O +ATOM 601 CB ILE A 80 37.195 18.235 63.811 1.00 30.79 C +ATOM 602 CG1 ILE A 80 38.329 17.200 63.998 1.00 30.79 C +ATOM 603 CG2 ILE A 80 37.126 19.150 65.045 1.00 30.79 C +ATOM 604 CD1 ILE A 80 38.148 16.197 65.150 1.00 30.79 C +ATOM 605 N PHE A 81 33.840 18.820 64.252 1.00 31.54 N +ATOM 606 CA PHE A 81 32.700 19.736 64.158 1.00 31.54 C +ATOM 607 C PHE A 81 33.114 21.216 64.092 1.00 31.54 C +ATOM 608 O PHE A 81 34.103 21.627 64.699 1.00 31.54 O +ATOM 609 CB PHE A 81 31.736 19.474 65.323 1.00 31.54 C +ATOM 610 CG PHE A 81 32.344 19.751 66.682 1.00 31.54 C +ATOM 611 CD1 PHE A 81 33.047 18.738 67.357 1.00 31.54 C +ATOM 612 CD2 PHE A 81 32.250 21.031 67.257 1.00 31.54 C +ATOM 613 CE1 PHE A 81 33.662 19.001 68.592 1.00 31.54 C +ATOM 614 CE2 PHE A 81 32.856 21.293 68.497 1.00 31.54 C +ATOM 615 CZ PHE A 81 33.563 20.280 69.166 1.00 31.54 C +ATOM 616 N GLY A 82 32.326 22.022 63.374 1.00 32.90 N +ATOM 617 CA GLY A 82 32.591 23.452 63.160 1.00 32.90 C +ATOM 618 C GLY A 82 31.925 24.403 64.163 1.00 32.90 C +ATOM 619 O GLY A 82 32.245 25.586 64.169 1.00 32.90 O +ATOM 620 N SER A 83 31.021 23.916 65.021 1.00 34.43 N +ATOM 621 CA SER A 83 30.282 24.698 66.029 1.00 34.43 C +ATOM 622 C SER A 83 31.133 25.051 67.262 1.00 34.43 C +ATOM 623 O SER A 83 30.784 24.774 68.405 1.00 34.43 O +ATOM 624 CB SER A 83 28.961 23.992 66.354 1.00 34.43 C +ATOM 625 OG SER A 83 29.144 22.598 66.554 1.00 34.43 O +ATOM 626 N GLN A 84 32.289 25.653 66.995 1.00 33.19 N +ATOM 627 CA GLN A 84 33.307 26.137 67.926 1.00 33.19 C +ATOM 628 C GLN A 84 34.126 27.246 67.240 1.00 33.19 C +ATOM 629 O GLN A 84 33.947 27.519 66.044 1.00 33.19 O +ATOM 630 CB GLN A 84 34.204 24.971 68.388 1.00 33.19 C +ATOM 631 CG GLN A 84 34.850 24.179 67.235 1.00 33.19 C +ATOM 632 CD GLN A 84 35.725 23.031 67.724 1.00 33.19 C +ATOM 633 OE1 GLN A 84 36.441 23.133 68.707 1.00 33.19 O +ATOM 634 NE2 GLN A 84 35.712 21.895 67.067 1.00 33.19 N +ATOM 635 N SER A 85 35.036 27.891 67.975 1.00 32.11 N +ATOM 636 CA SER A 85 35.941 28.876 67.372 1.00 32.11 C +ATOM 637 C SER A 85 36.845 28.225 66.315 1.00 32.11 C +ATOM 638 O SER A 85 37.062 27.008 66.327 1.00 32.11 O +ATOM 639 CB SER A 85 36.734 29.646 68.437 1.00 32.11 C +ATOM 640 OG SER A 85 37.991 29.066 68.709 1.00 32.11 O +ATOM 641 N ILE A 86 37.350 29.030 65.378 1.00 31.17 N +ATOM 642 CA ILE A 86 38.242 28.554 64.310 1.00 31.17 C +ATOM 643 C ILE A 86 39.519 27.966 64.920 1.00 31.17 C +ATOM 644 O ILE A 86 39.951 26.886 64.529 1.00 31.17 O +ATOM 645 CB ILE A 86 38.540 29.711 63.325 1.00 31.17 C +ATOM 646 CG1 ILE A 86 37.243 30.106 62.577 1.00 31.17 C +ATOM 647 CG2 ILE A 86 39.642 29.304 62.330 1.00 31.17 C +ATOM 648 CD1 ILE A 86 37.362 31.378 61.731 1.00 31.17 C +ATOM 649 N GLU A 87 40.064 28.631 65.936 1.00 31.01 N +ATOM 650 CA GLU A 87 41.287 28.249 66.642 1.00 31.01 C +ATOM 651 C GLU A 87 41.108 26.916 67.373 1.00 31.01 C +ATOM 652 O GLU A 87 41.958 26.040 67.272 1.00 31.01 O +ATOM 653 CB GLU A 87 41.702 29.328 67.665 1.00 31.01 C +ATOM 654 CG GLU A 87 41.825 30.762 67.119 1.00 31.01 C +ATOM 655 CD GLU A 87 40.471 31.418 66.787 1.00 31.01 C +ATOM 656 OE1 GLU A 87 40.458 32.335 65.943 1.00 31.01 O +ATOM 657 OE2 GLU A 87 39.427 30.945 67.315 1.00 31.01 O +ATOM 658 N SER A 88 39.975 26.733 68.062 1.00 31.23 N +ATOM 659 CA SER A 88 39.659 25.483 68.762 1.00 31.23 C +ATOM 660 C SER A 88 39.527 24.310 67.787 1.00 31.23 C +ATOM 661 O SER A 88 40.073 23.234 68.026 1.00 31.23 O +ATOM 662 CB SER A 88 38.371 25.674 69.567 1.00 31.23 C +ATOM 663 OG SER A 88 38.047 24.526 70.325 1.00 31.23 O +ATOM 664 N GLN A 89 38.857 24.516 66.646 1.00 30.88 N +ATOM 665 CA GLN A 89 38.731 23.472 65.627 1.00 30.88 C +ATOM 666 C GLN A 89 40.085 23.122 64.996 1.00 30.88 C +ATOM 667 O GLN A 89 40.404 21.943 64.838 1.00 30.88 O +ATOM 668 CB GLN A 89 37.744 23.921 64.543 1.00 30.88 C +ATOM 669 CG GLN A 89 37.311 22.726 63.683 1.00 30.88 C +ATOM 670 CD GLN A 89 36.410 23.118 62.524 1.00 30.88 C +ATOM 671 OE1 GLN A 89 36.273 24.274 62.150 1.00 30.88 O +ATOM 672 NE2 GLN A 89 35.750 22.166 61.910 1.00 30.88 N +ATOM 673 N ALA A 90 40.882 24.137 64.652 1.00 31.01 N +ATOM 674 CA ALA A 90 42.221 23.963 64.103 1.00 31.01 C +ATOM 675 C ALA A 90 43.143 23.237 65.095 1.00 31.01 C +ATOM 676 O ALA A 90 43.858 22.321 64.696 1.00 31.01 O +ATOM 677 CB ALA A 90 42.762 25.344 63.712 1.00 31.01 C +ATOM 678 N ALA A 91 43.062 23.558 66.390 1.00 31.61 N +ATOM 679 CA ALA A 91 43.802 22.859 67.438 1.00 31.61 C +ATOM 680 C ALA A 91 43.433 21.370 67.512 1.00 31.61 C +ATOM 681 O ALA A 91 44.332 20.536 67.571 1.00 31.61 O +ATOM 682 CB ALA A 91 43.563 23.568 68.775 1.00 31.61 C +ATOM 683 N MET A 92 42.144 21.016 67.419 1.00 31.20 N +ATOM 684 CA MET A 92 41.715 19.610 67.385 1.00 31.20 C +ATOM 685 C MET A 92 42.229 18.870 66.145 1.00 31.20 C +ATOM 686 O MET A 92 42.684 17.732 66.246 1.00 31.20 O +ATOM 687 CB MET A 92 40.186 19.504 67.414 1.00 31.20 C +ATOM 688 CG MET A 92 39.570 19.898 68.754 1.00 31.20 C +ATOM 689 SD MET A 92 37.769 19.669 68.792 1.00 31.20 S +ATOM 690 CE MET A 92 37.453 20.394 70.418 1.00 31.20 C +ATOM 691 N VAL A 93 42.176 19.502 64.965 1.00 31.07 N +ATOM 692 CA VAL A 93 42.748 18.931 63.732 1.00 31.07 C +ATOM 693 C VAL A 93 44.245 18.697 63.908 1.00 31.07 C +ATOM 694 O VAL A 93 44.725 17.585 63.687 1.00 31.07 O +ATOM 695 CB VAL A 93 42.468 19.844 62.523 1.00 31.07 C +ATOM 696 CG1 VAL A 93 43.279 19.475 61.275 1.00 31.07 C +ATOM 697 CG2 VAL A 93 40.983 19.777 62.144 1.00 31.07 C +ATOM 698 N HIS A 94 44.967 19.714 64.373 1.00 34.19 N +ATOM 699 CA HIS A 94 46.402 19.644 64.612 1.00 34.19 C +ATOM 700 C HIS A 94 46.756 18.549 65.629 1.00 34.19 C +ATOM 701 O HIS A 94 47.666 17.756 65.392 1.00 34.19 O +ATOM 702 CB HIS A 94 46.881 21.022 65.079 1.00 34.19 C +ATOM 703 CG HIS A 94 48.372 21.072 65.237 1.00 34.19 C +ATOM 704 ND1 HIS A 94 49.287 21.046 64.213 1.00 34.19 N +ATOM 705 CD2 HIS A 94 49.072 21.088 66.411 1.00 34.19 C +ATOM 706 CE1 HIS A 94 50.514 21.037 64.758 1.00 34.19 C +ATOM 707 NE2 HIS A 94 50.435 21.056 66.100 1.00 34.19 N +ATOM 708 N ALA A 95 45.986 18.433 66.713 1.00 35.19 N +ATOM 709 CA ALA A 95 46.152 17.386 67.711 1.00 35.19 C +ATOM 710 C ALA A 95 46.036 15.989 67.090 1.00 35.19 C +ATOM 711 O ALA A 95 46.867 15.138 67.386 1.00 35.19 O +ATOM 712 CB ALA A 95 45.128 17.589 68.832 1.00 35.19 C +ATOM 713 N VAL A 96 45.082 15.753 66.180 1.00 33.63 N +ATOM 714 CA VAL A 96 44.968 14.473 65.458 1.00 33.63 C +ATOM 715 C VAL A 96 46.153 14.258 64.515 1.00 33.63 C +ATOM 716 O VAL A 96 46.734 13.171 64.510 1.00 33.63 O +ATOM 717 CB VAL A 96 43.639 14.370 64.687 1.00 33.63 C +ATOM 718 CG1 VAL A 96 43.537 13.037 63.928 1.00 33.63 C +ATOM 719 CG2 VAL A 96 42.444 14.438 65.644 1.00 33.63 C +ATOM 720 N LYS A 97 46.556 15.286 63.755 1.00 35.78 N +ATOM 721 CA LYS A 97 47.700 15.209 62.827 1.00 35.78 C +ATOM 722 C LYS A 97 49.039 15.016 63.537 1.00 35.78 C +ATOM 723 O LYS A 97 49.961 14.475 62.929 1.00 35.78 O +ATOM 724 CB LYS A 97 47.753 16.458 61.928 1.00 35.78 C +ATOM 725 CG LYS A 97 46.567 16.618 60.963 1.00 35.78 C +ATOM 726 CD LYS A 97 46.300 15.431 60.026 1.00 35.78 C +ATOM 727 CE LYS A 97 47.442 15.118 59.057 1.00 35.78 C +ATOM 728 NZ LYS A 97 47.036 14.062 58.096 1.00 35.78 N +ATOM 729 N ASN A 98 49.147 15.383 64.813 1.00 43.38 N +ATOM 730 CA ASN A 98 50.334 15.110 65.616 1.00 43.38 C +ATOM 731 C ASN A 98 50.529 13.620 65.899 1.00 43.38 C +ATOM 732 O ASN A 98 51.649 13.219 66.202 1.00 43.38 O +ATOM 733 CB ASN A 98 50.300 15.954 66.904 1.00 43.38 C +ATOM 734 CG ASN A 98 50.670 17.402 66.637 1.00 43.38 C +ATOM 735 OD1 ASN A 98 51.298 17.732 65.644 1.00 43.38 O +ATOM 736 ND2 ASN A 98 50.349 18.305 67.530 1.00 43.38 N +ATOM 737 N PHE A 99 49.508 12.774 65.749 1.00 50.20 N +ATOM 738 CA PHE A 99 49.675 11.330 65.874 1.00 50.20 C +ATOM 739 C PHE A 99 50.158 10.717 64.559 1.00 50.20 C +ATOM 740 O PHE A 99 49.513 10.828 63.514 1.00 50.20 O +ATOM 741 CB PHE A 99 48.382 10.675 66.345 1.00 50.20 C +ATOM 742 CG PHE A 99 48.009 10.990 67.773 1.00 50.20 C +ATOM 743 CD1 PHE A 99 48.371 10.128 68.824 1.00 50.20 C +ATOM 744 CD2 PHE A 99 47.329 12.181 68.049 1.00 50.20 C +ATOM 745 CE1 PHE A 99 48.058 10.475 70.151 1.00 50.20 C +ATOM 746 CE2 PHE A 99 47.057 12.553 69.375 1.00 50.20 C +ATOM 747 CZ PHE A 99 47.411 11.693 70.427 1.00 50.20 C +ATOM 748 N LYS A 100 51.281 9.999 64.640 1.00 58.98 N +ATOM 749 CA LYS A 100 51.906 9.287 63.525 1.00 58.98 C +ATOM 750 C LYS A 100 51.861 7.782 63.799 1.00 58.98 C +ATOM 751 O LYS A 100 52.207 7.337 64.887 1.00 58.98 O +ATOM 752 CB LYS A 100 53.302 9.900 63.311 1.00 58.98 C +ATOM 753 CG LYS A 100 54.121 9.356 62.131 1.00 58.98 C +ATOM 754 CD LYS A 100 55.495 10.051 62.132 1.00 58.98 C +ATOM 755 CE LYS A 100 56.436 9.612 61.003 1.00 58.98 C +ATOM 756 NZ LYS A 100 57.788 10.219 61.184 1.00 58.98 N +ATOM 757 N ALA A 101 51.408 6.984 62.840 1.00 67.64 N +ATOM 758 CA ALA A 101 51.329 5.531 62.961 1.00 67.64 C +ATOM 759 C ALA A 101 52.676 4.930 63.384 1.00 67.64 C +ATOM 760 O ALA A 101 53.719 5.246 62.808 1.00 67.64 O +ATOM 761 CB ALA A 101 50.860 4.938 61.637 1.00 67.64 C +ATOM 762 N ASN A 223 44.919 -1.618 58.637 1.00 46.97 N +ATOM 763 CA ASN A 223 43.846 -0.733 58.200 1.00 46.97 C +ATOM 764 C ASN A 223 44.054 0.739 58.569 1.00 46.97 C +ATOM 765 O ASN A 223 43.185 1.531 58.234 1.00 46.97 O +ATOM 766 CB ASN A 223 42.502 -1.268 58.720 1.00 46.97 C +ATOM 767 CG ASN A 223 42.149 -2.636 58.169 1.00 46.97 C +ATOM 768 OD1 ASN A 223 42.575 -3.062 57.107 1.00 46.97 O +ATOM 769 ND2 ASN A 223 41.347 -3.382 58.890 1.00 46.97 N +ATOM 770 N GLU A 224 45.157 1.133 59.216 1.00 40.15 N +ATOM 771 CA GLU A 224 45.330 2.534 59.629 1.00 40.15 C +ATOM 772 C GLU A 224 45.236 3.491 58.434 1.00 40.15 C +ATOM 773 O GLU A 224 45.922 3.337 57.414 1.00 40.15 O +ATOM 774 CB GLU A 224 46.657 2.769 60.366 1.00 40.15 C +ATOM 775 CG GLU A 224 46.737 2.002 61.689 1.00 40.15 C +ATOM 776 CD GLU A 224 48.030 2.276 62.472 1.00 40.15 C +ATOM 777 OE1 GLU A 224 47.962 2.324 63.721 1.00 40.15 O +ATOM 778 OE2 GLU A 224 49.092 2.440 61.829 1.00 40.15 O +ATOM 779 N LEU A 225 44.386 4.507 58.571 1.00 35.85 N +ATOM 780 CA LEU A 225 44.198 5.541 57.569 1.00 35.85 C +ATOM 781 C LEU A 225 45.186 6.681 57.830 1.00 35.85 C +ATOM 782 O LEU A 225 44.961 7.536 58.689 1.00 35.85 O +ATOM 783 CB LEU A 225 42.727 5.982 57.563 1.00 35.85 C +ATOM 784 CG LEU A 225 42.408 6.854 56.346 1.00 35.85 C +ATOM 785 CD1 LEU A 225 42.426 6.058 55.041 1.00 35.85 C +ATOM 786 CD2 LEU A 225 41.033 7.495 56.490 1.00 35.85 C +ATOM 787 N VAL A 226 46.295 6.669 57.092 1.00 36.65 N +ATOM 788 CA VAL A 226 47.411 7.610 57.266 1.00 36.65 C +ATOM 789 C VAL A 226 47.872 8.253 55.963 1.00 36.65 C +ATOM 790 O VAL A 226 47.736 7.670 54.887 1.00 36.65 O +ATOM 791 CB VAL A 226 48.610 6.956 57.964 1.00 36.65 C +ATOM 792 CG1 VAL A 226 48.260 6.506 59.378 1.00 36.65 C +ATOM 793 CG2 VAL A 226 49.196 5.763 57.194 1.00 36.65 C +ATOM 794 N ASP A 227 48.442 9.451 56.075 1.00 38.60 N +ATOM 795 CA ASP A 227 49.003 10.201 54.955 1.00 38.60 C +ATOM 796 C ASP A 227 50.414 9.700 54.584 1.00 38.60 C +ATOM 797 O ASP A 227 50.951 8.743 55.158 1.00 38.60 O +ATOM 798 CB ASP A 227 48.917 11.719 55.236 1.00 38.60 C +ATOM 799 CG ASP A 227 49.827 12.254 56.349 1.00 38.60 C +ATOM 800 OD1 ASP A 227 50.807 11.573 56.729 1.00 38.60 O +ATOM 801 OD2 ASP A 227 49.581 13.394 56.813 1.00 38.60 O +ATOM 802 N SER A 228 51.047 10.351 53.606 1.00 51.38 N +ATOM 803 CA SER A 228 52.405 10.013 53.155 1.00 51.38 C +ATOM 804 C SER A 228 53.476 10.187 54.242 1.00 51.38 C +ATOM 805 O SER A 228 54.526 9.544 54.175 1.00 51.38 O +ATOM 806 CB SER A 228 52.762 10.863 51.934 1.00 51.38 C +ATOM 807 OG SER A 228 52.713 12.236 52.266 1.00 51.38 O +ATOM 808 N GLN A 229 53.204 10.997 55.269 1.00 62.96 N +ATOM 809 CA GLN A 229 54.056 11.203 56.442 1.00 62.96 C +ATOM 810 C GLN A 229 53.696 10.261 57.602 1.00 62.96 C +ATOM 811 O GLN A 229 54.280 10.367 58.681 1.00 62.96 O +ATOM 812 CB GLN A 229 54.004 12.675 56.877 1.00 62.96 C +ATOM 813 CG GLN A 229 54.552 13.624 55.802 1.00 62.96 C +ATOM 814 CD GLN A 229 54.576 15.070 56.287 1.00 62.96 C +ATOM 815 OE1 GLN A 229 53.702 15.534 56.995 1.00 62.96 O +ATOM 816 NE2 GLN A 229 55.578 15.844 55.932 1.00 62.96 N +ATOM 817 N LYS A 230 52.784 9.304 57.376 1.00 48.97 N +ATOM 818 CA LYS A 230 52.245 8.367 58.368 1.00 48.97 C +ATOM 819 C LYS A 230 51.412 9.028 59.467 1.00 48.97 C +ATOM 820 O LYS A 230 51.174 8.384 60.482 1.00 48.97 O +ATOM 821 CB LYS A 230 53.343 7.423 58.917 1.00 48.97 C +ATOM 822 CG LYS A 230 53.999 6.524 57.861 1.00 48.97 C +ATOM 823 CD LYS A 230 52.972 5.540 57.290 1.00 48.97 C +ATOM 824 CE LYS A 230 53.620 4.493 56.389 1.00 48.97 C +ATOM 825 NZ LYS A 230 52.587 3.541 55.911 1.00 48.97 N +ATOM 826 N ARG A 231 50.940 10.263 59.294 1.00 39.69 N +ATOM 827 CA ARG A 231 50.020 10.907 60.247 1.00 39.69 C +ATOM 828 C ARG A 231 48.595 10.459 59.966 1.00 39.69 C +ATOM 829 O ARG A 231 48.257 10.249 58.802 1.00 39.69 O +ATOM 830 CB ARG A 231 50.122 12.430 60.169 1.00 39.69 C +ATOM 831 CG ARG A 231 51.560 12.952 60.342 1.00 39.69 C +ATOM 832 CD ARG A 231 51.649 14.456 60.074 1.00 39.69 C +ATOM 833 NE ARG A 231 51.165 14.782 58.724 1.00 39.69 N +ATOM 834 CZ ARG A 231 51.011 15.977 58.198 1.00 39.69 C +ATOM 835 NH1 ARG A 231 51.351 17.069 58.826 1.00 39.69 N +ATOM 836 NH2 ARG A 231 50.450 16.060 57.029 1.00 39.69 N +ATOM 837 N TYR A 232 47.759 10.324 60.991 1.00 34.19 N +ATOM 838 CA TYR A 232 46.359 9.935 60.775 1.00 34.19 C +ATOM 839 C TYR A 232 45.619 10.944 59.895 1.00 34.19 C +ATOM 840 O TYR A 232 45.892 12.146 59.949 1.00 34.19 O +ATOM 841 CB TYR A 232 45.629 9.723 62.104 1.00 34.19 C +ATOM 842 CG TYR A 232 46.121 8.523 62.895 1.00 34.19 C +ATOM 843 CD1 TYR A 232 46.191 7.239 62.313 1.00 34.19 C +ATOM 844 CD2 TYR A 232 46.505 8.701 64.231 1.00 34.19 C +ATOM 845 CE1 TYR A 232 46.714 6.158 63.050 1.00 34.19 C +ATOM 846 CE2 TYR A 232 47.052 7.635 64.964 1.00 34.19 C +ATOM 847 CZ TYR A 232 47.175 6.368 64.366 1.00 34.19 C +ATOM 848 OH TYR A 232 47.747 5.365 65.079 1.00 34.19 O +ATOM 849 N LEU A 233 44.697 10.449 59.068 1.00 31.61 N +ATOM 850 CA LEU A 233 43.801 11.309 58.301 1.00 31.61 C +ATOM 851 C LEU A 233 42.687 11.853 59.195 1.00 31.61 C +ATOM 852 O LEU A 233 42.140 11.136 60.034 1.00 31.61 O +ATOM 853 CB LEU A 233 43.211 10.587 57.078 1.00 31.61 C +ATOM 854 CG LEU A 233 44.052 10.746 55.802 1.00 31.61 C +ATOM 855 CD1 LEU A 233 45.395 10.058 55.918 1.00 31.61 C +ATOM 856 CD2 LEU A 233 43.339 10.167 54.581 1.00 31.61 C +ATOM 857 N VAL A 234 42.310 13.107 58.969 1.00 30.85 N +ATOM 858 CA VAL A 234 41.201 13.759 59.666 1.00 30.85 C +ATOM 859 C VAL A 234 40.454 14.702 58.732 1.00 30.85 C +ATOM 860 O VAL A 234 41.050 15.427 57.934 1.00 30.85 O +ATOM 861 CB VAL A 234 41.706 14.459 60.940 1.00 30.85 C +ATOM 862 CG1 VAL A 234 42.594 15.672 60.657 1.00 30.85 C +ATOM 863 CG2 VAL A 234 40.559 14.894 61.857 1.00 30.85 C +ATOM 864 N GLY A 235 39.129 14.683 58.832 1.00 30.67 N +ATOM 865 CA GLY A 235 38.265 15.637 58.159 1.00 30.67 C +ATOM 866 C GLY A 235 37.865 16.796 59.074 1.00 30.67 C +ATOM 867 O GLY A 235 37.930 16.694 60.300 1.00 30.67 O +ATOM 868 N ALA A 236 37.367 17.886 58.498 1.00 30.64 N +ATOM 869 CA ALA A 236 36.788 18.988 59.264 1.00 30.64 C +ATOM 870 C ALA A 236 35.496 19.516 58.627 1.00 30.64 C +ATOM 871 O ALA A 236 35.408 19.716 57.416 1.00 30.64 O +ATOM 872 CB ALA A 236 37.846 20.077 59.449 1.00 30.64 C +ATOM 873 N GLY A 237 34.476 19.740 59.454 1.00 30.97 N +ATOM 874 CA GLY A 237 33.238 20.391 59.041 1.00 30.97 C +ATOM 875 C GLY A 237 33.407 21.904 58.933 1.00 30.97 C +ATOM 876 O GLY A 237 33.947 22.532 59.842 1.00 30.97 O +ATOM 877 N ILE A 238 32.917 22.503 57.856 1.00 31.10 N +ATOM 878 CA ILE A 238 32.906 23.956 57.645 1.00 31.10 C +ATOM 879 C ILE A 238 31.474 24.440 57.399 1.00 31.10 C +ATOM 880 O ILE A 238 30.609 23.663 57.003 1.00 31.10 O +ATOM 881 CB ILE A 238 33.886 24.376 56.523 1.00 31.10 C +ATOM 882 CG1 ILE A 238 33.482 23.814 55.142 1.00 31.10 C +ATOM 883 CG2 ILE A 238 35.322 23.970 56.905 1.00 31.10 C +ATOM 884 CD1 ILE A 238 34.346 24.331 53.985 1.00 31.10 C +ATOM 885 N ASN A 239 31.217 25.726 57.623 1.00 32.98 N +ATOM 886 CA ASN A 239 29.947 26.372 57.285 1.00 32.98 C +ATOM 887 C ASN A 239 30.117 27.285 56.054 1.00 32.98 C +ATOM 888 O ASN A 239 31.213 27.433 55.522 1.00 32.98 O +ATOM 889 CB ASN A 239 29.412 27.104 58.532 1.00 32.98 C +ATOM 890 CG ASN A 239 30.330 28.235 58.937 1.00 32.98 C +ATOM 891 OD1 ASN A 239 30.412 29.243 58.255 1.00 32.98 O +ATOM 892 ND2 ASN A 239 31.083 28.077 59.990 1.00 32.98 N +ATOM 893 N THR A 240 29.035 27.917 55.603 1.00 32.78 N +ATOM 894 CA THR A 240 29.010 28.807 54.425 1.00 32.78 C +ATOM 895 C THR A 240 29.076 30.300 54.786 1.00 32.78 C +ATOM 896 O THR A 240 28.637 31.159 54.015 1.00 32.78 O +ATOM 897 CB THR A 240 27.768 28.523 53.579 1.00 32.78 C +ATOM 898 OG1 THR A 240 26.616 28.702 54.376 1.00 32.78 O +ATOM 899 CG2 THR A 240 27.711 27.097 53.039 1.00 32.78 C +ATOM 900 N ARG A 241 29.576 30.635 55.982 1.00 34.43 N +ATOM 901 CA ARG A 241 29.621 32.000 56.532 1.00 34.43 C +ATOM 902 C ARG A 241 31.058 32.459 56.774 1.00 34.43 C +ATOM 903 O ARG A 241 31.514 33.360 56.080 1.00 34.43 O +ATOM 904 CB ARG A 241 28.741 32.099 57.799 1.00 34.43 C +ATOM 905 CG ARG A 241 27.274 31.667 57.602 1.00 34.43 C +ATOM 906 CD ARG A 241 26.510 32.626 56.682 1.00 34.43 C +ATOM 907 NE ARG A 241 25.157 32.131 56.370 1.00 34.43 N +ATOM 908 CZ ARG A 241 24.738 31.727 55.181 1.00 34.43 C +ATOM 909 NH1 ARG A 241 25.529 31.647 54.149 1.00 34.43 N +ATOM 910 NH2 ARG A 241 23.503 31.394 54.970 1.00 34.43 N +ATOM 911 N ASP A 242 31.774 31.806 57.688 1.00 32.69 N +ATOM 912 CA ASP A 242 33.154 32.149 58.089 1.00 32.69 C +ATOM 913 C ASP A 242 34.233 31.339 57.341 1.00 32.69 C +ATOM 914 O ASP A 242 35.389 31.256 57.758 1.00 32.69 O +ATOM 915 CB ASP A 242 33.295 32.061 59.620 1.00 32.69 C +ATOM 916 CG ASP A 242 33.189 30.643 60.192 1.00 32.69 C +ATOM 917 OD1 ASP A 242 33.274 29.639 59.456 1.00 32.69 O +ATOM 918 OD2 ASP A 242 32.962 30.501 61.412 1.00 32.69 O +ATOM 919 N PHE A 243 33.868 30.709 56.219 1.00 31.75 N +ATOM 920 CA PHE A 243 34.757 29.804 55.484 1.00 31.75 C +ATOM 921 C PHE A 243 36.032 30.478 54.964 1.00 31.75 C +ATOM 922 O PHE A 243 37.031 29.792 54.754 1.00 31.75 O +ATOM 923 CB PHE A 243 34.001 29.158 54.318 1.00 31.75 C +ATOM 924 CG PHE A 243 33.555 30.116 53.231 1.00 31.75 C +ATOM 925 CD1 PHE A 243 32.287 30.718 53.298 1.00 31.75 C +ATOM 926 CD2 PHE A 243 34.397 30.391 52.139 1.00 31.75 C +ATOM 927 CE1 PHE A 243 31.843 31.567 52.270 1.00 31.75 C +ATOM 928 CE2 PHE A 243 33.961 31.252 51.119 1.00 31.75 C +ATOM 929 CZ PHE A 243 32.681 31.823 51.172 1.00 31.75 C +ATOM 930 N ARG A 244 36.026 31.802 54.754 1.00 31.07 N +ATOM 931 CA ARG A 244 37.194 32.531 54.236 1.00 31.07 C +ATOM 932 C ARG A 244 38.344 32.554 55.237 1.00 31.07 C +ATOM 933 O ARG A 244 39.499 32.519 54.816 1.00 31.07 O +ATOM 934 CB ARG A 244 36.805 33.958 53.834 1.00 31.07 C +ATOM 935 CG ARG A 244 35.902 33.961 52.599 1.00 31.07 C +ATOM 936 CD ARG A 244 35.651 35.397 52.131 1.00 31.07 C +ATOM 937 NE ARG A 244 34.880 35.424 50.873 1.00 31.07 N +ATOM 938 CZ ARG A 244 33.580 35.231 50.748 1.00 31.07 C +ATOM 939 NH1 ARG A 244 32.813 35.006 51.779 1.00 31.07 N +ATOM 940 NH2 ARG A 244 33.025 35.259 49.567 1.00 31.07 N +ATOM 941 N GLU A 245 38.023 32.568 56.524 1.00 30.91 N +ATOM 942 CA GLU A 245 38.947 32.527 57.654 1.00 30.91 C +ATOM 943 C GLU A 245 39.157 31.091 58.152 1.00 30.91 C +ATOM 944 O GLU A 245 40.286 30.681 58.418 1.00 30.91 O +ATOM 945 CB GLU A 245 38.396 33.406 58.793 1.00 30.91 C +ATOM 946 CG GLU A 245 38.201 34.893 58.436 1.00 30.91 C +ATOM 947 CD GLU A 245 36.917 35.236 57.648 1.00 30.91 C +ATOM 948 OE1 GLU A 245 36.769 36.428 57.299 1.00 30.91 O +ATOM 949 OE2 GLU A 245 36.106 34.330 57.335 1.00 30.91 O +ATOM 950 N ARG A 246 38.083 30.292 58.213 1.00 30.85 N +ATOM 951 CA ARG A 246 38.123 28.924 58.742 1.00 30.85 C +ATOM 952 C ARG A 246 38.889 27.958 57.843 1.00 30.85 C +ATOM 953 O ARG A 246 39.692 27.176 58.341 1.00 30.85 O +ATOM 954 CB ARG A 246 36.685 28.456 59.015 1.00 30.85 C +ATOM 955 CG ARG A 246 36.634 27.095 59.727 1.00 30.85 C +ATOM 956 CD ARG A 246 35.191 26.666 60.019 1.00 30.85 C +ATOM 957 NE ARG A 246 34.521 27.586 60.953 1.00 30.85 N +ATOM 958 CZ ARG A 246 34.541 27.557 62.275 1.00 30.85 C +ATOM 959 NH1 ARG A 246 35.189 26.669 62.969 1.00 30.85 N +ATOM 960 NH2 ARG A 246 33.890 28.442 62.962 1.00 30.85 N +ATOM 961 N VAL A 247 38.671 27.998 56.525 1.00 30.67 N +ATOM 962 CA VAL A 247 39.330 27.061 55.597 1.00 30.67 C +ATOM 963 C VAL A 247 40.860 27.191 55.642 1.00 30.67 C +ATOM 964 O VAL A 247 41.503 26.158 55.806 1.00 30.67 O +ATOM 965 CB VAL A 247 38.777 27.176 54.162 1.00 30.67 C +ATOM 966 CG1 VAL A 247 39.612 26.388 53.144 1.00 30.67 C +ATOM 967 CG2 VAL A 247 37.328 26.685 54.091 1.00 30.67 C +ATOM 968 N PRO A 248 41.472 28.393 55.541 1.00 30.70 N +ATOM 969 CA PRO A 248 42.922 28.534 55.678 1.00 30.70 C +ATOM 970 C PRO A 248 43.475 27.945 56.978 1.00 30.70 C +ATOM 971 O PRO A 248 44.424 27.171 56.917 1.00 30.70 O +ATOM 972 CB PRO A 248 43.219 30.032 55.586 1.00 30.70 C +ATOM 973 CG PRO A 248 42.040 30.583 54.798 1.00 30.70 C +ATOM 974 CD PRO A 248 40.880 29.684 55.211 1.00 30.70 C +ATOM 975 N ALA A 249 42.847 28.240 58.122 1.00 30.73 N +ATOM 976 CA ALA A 249 43.299 27.744 59.420 1.00 30.73 C +ATOM 977 C ALA A 249 43.252 26.207 59.511 1.00 30.73 C +ATOM 978 O ALA A 249 44.159 25.582 60.055 1.00 30.73 O +ATOM 979 CB ALA A 249 42.421 28.388 60.496 1.00 30.73 C +ATOM 980 N LEU A 250 42.217 25.574 58.947 1.00 30.67 N +ATOM 981 CA LEU A 250 42.099 24.111 58.930 1.00 30.67 C +ATOM 982 C LEU A 250 43.095 23.448 57.972 1.00 30.67 C +ATOM 983 O LEU A 250 43.620 22.379 58.282 1.00 30.67 O +ATOM 984 CB LEU A 250 40.664 23.716 58.558 1.00 30.67 C +ATOM 985 CG LEU A 250 39.601 24.128 59.588 1.00 30.67 C +ATOM 986 CD1 LEU A 250 38.223 23.815 59.004 1.00 30.67 C +ATOM 987 CD2 LEU A 250 39.746 23.404 60.923 1.00 30.67 C +ATOM 988 N VAL A 251 43.365 24.073 56.823 1.00 31.10 N +ATOM 989 CA VAL A 251 44.387 23.603 55.875 1.00 31.10 C +ATOM 990 C VAL A 251 45.777 23.690 56.508 1.00 31.10 C +ATOM 991 O VAL A 251 46.540 22.733 56.419 1.00 31.10 O +ATOM 992 CB VAL A 251 44.313 24.396 54.554 1.00 31.10 C +ATOM 993 CG1 VAL A 251 45.484 24.097 53.607 1.00 31.10 C +ATOM 994 CG2 VAL A 251 43.026 24.053 53.790 1.00 31.10 C +ATOM 995 N GLU A 252 46.085 24.790 57.198 1.00 31.86 N +ATOM 996 CA GLU A 252 47.337 24.968 57.945 1.00 31.86 C +ATOM 997 C GLU A 252 47.480 23.951 59.086 1.00 31.86 C +ATOM 998 O GLU A 252 48.539 23.347 59.249 1.00 31.86 O +ATOM 999 CB GLU A 252 47.387 26.410 58.464 1.00 31.86 C +ATOM 1000 CG GLU A 252 48.696 26.734 59.198 1.00 31.86 C +ATOM 1001 CD GLU A 252 48.792 28.205 59.630 1.00 31.86 C +ATOM 1002 OE1 GLU A 252 49.880 28.580 60.122 1.00 31.86 O +ATOM 1003 OE2 GLU A 252 47.804 28.954 59.444 1.00 31.86 O +ATOM 1004 N ALA A 253 46.394 23.677 59.817 1.00 32.50 N +ATOM 1005 CA ALA A 253 46.363 22.637 60.845 1.00 32.50 C +ATOM 1006 C ALA A 253 46.556 21.212 60.287 1.00 32.50 C +ATOM 1007 O ALA A 253 46.896 20.297 61.041 1.00 32.50 O +ATOM 1008 CB ALA A 253 45.039 22.755 61.603 1.00 32.50 C +ATOM 1009 N GLY A 254 46.363 21.022 58.977 1.00 32.53 N +ATOM 1010 CA GLY A 254 46.608 19.770 58.267 1.00 32.53 C +ATOM 1011 C GLY A 254 45.370 18.907 58.023 1.00 32.53 C +ATOM 1012 O GLY A 254 45.527 17.697 57.878 1.00 32.53 O +ATOM 1013 N ALA A 255 44.163 19.485 57.982 1.00 31.01 N +ATOM 1014 CA ALA A 255 42.949 18.752 57.611 1.00 31.01 C +ATOM 1015 C ALA A 255 43.077 18.154 56.199 1.00 31.01 C +ATOM 1016 O ALA A 255 43.384 18.865 55.242 1.00 31.01 O +ATOM 1017 CB ALA A 255 41.728 19.681 57.702 1.00 31.01 C +ATOM 1018 N ASP A 256 42.802 16.855 56.057 1.00 30.88 N +ATOM 1019 CA ASP A 256 42.976 16.128 54.792 1.00 30.88 C +ATOM 1020 C ASP A 256 41.762 16.287 53.859 1.00 30.88 C +ATOM 1021 O ASP A 256 41.874 16.186 52.637 1.00 30.88 O +ATOM 1022 CB ASP A 256 43.247 14.647 55.097 1.00 30.88 C +ATOM 1023 CG ASP A 256 44.443 14.439 56.038 1.00 30.88 C +ATOM 1024 OD1 ASP A 256 45.597 14.300 55.578 1.00 30.88 O +ATOM 1025 OD2 ASP A 256 44.228 14.337 57.265 1.00 30.88 O +ATOM 1026 N VAL A 257 40.583 16.545 54.432 1.00 30.67 N +ATOM 1027 CA VAL A 257 39.336 16.769 53.692 1.00 30.67 C +ATOM 1028 C VAL A 257 38.377 17.660 54.482 1.00 30.67 C +ATOM 1029 O VAL A 257 38.235 17.549 55.698 1.00 30.67 O +ATOM 1030 CB VAL A 257 38.694 15.422 53.300 1.00 30.67 C +ATOM 1031 CG1 VAL A 257 38.341 14.558 54.515 1.00 30.67 C +ATOM 1032 CG2 VAL A 257 37.440 15.605 52.440 1.00 30.67 C +ATOM 1033 N LEU A 258 37.678 18.549 53.789 1.00 30.58 N +ATOM 1034 CA LEU A 258 36.649 19.411 54.359 1.00 30.58 C +ATOM 1035 C LEU A 258 35.256 18.903 53.983 1.00 30.58 C +ATOM 1036 O LEU A 258 35.080 18.224 52.973 1.00 30.58 O +ATOM 1037 CB LEU A 258 36.892 20.864 53.918 1.00 30.58 C +ATOM 1038 CG LEU A 258 38.305 21.389 54.245 1.00 30.58 C +ATOM 1039 CD1 LEU A 258 38.428 22.831 53.770 1.00 30.58 C +ATOM 1040 CD2 LEU A 258 38.619 21.360 55.743 1.00 30.58 C +ATOM 1041 N CYS A 259 34.246 19.247 54.773 1.00 30.67 N +ATOM 1042 CA CYS A 259 32.849 18.996 54.430 1.00 30.67 C +ATOM 1043 C CYS A 259 32.004 20.214 54.793 1.00 30.67 C +ATOM 1044 O CYS A 259 31.982 20.619 55.954 1.00 30.67 O +ATOM 1045 CB CYS A 259 32.378 17.736 55.157 1.00 30.67 C +ATOM 1046 SG CYS A 259 30.700 17.278 54.622 1.00 30.67 S +ATOM 1047 N ILE A 260 31.316 20.798 53.813 1.00 30.88 N +ATOM 1048 CA ILE A 260 30.350 21.867 54.066 1.00 30.88 C +ATOM 1049 C ILE A 260 29.131 21.234 54.741 1.00 30.88 C +ATOM 1050 O ILE A 260 28.387 20.486 54.110 1.00 30.88 O +ATOM 1051 CB ILE A 260 29.983 22.629 52.773 1.00 30.88 C +ATOM 1052 CG1 ILE A 260 31.238 23.205 52.078 1.00 30.88 C +ATOM 1053 CG2 ILE A 260 29.017 23.776 53.125 1.00 30.88 C +ATOM 1054 CD1 ILE A 260 30.979 23.700 50.653 1.00 30.88 C +ATOM 1055 N ASP A 261 28.951 21.509 56.031 1.00 33.36 N +ATOM 1056 CA ASP A 261 27.873 20.940 56.834 1.00 33.36 C +ATOM 1057 C ASP A 261 26.703 21.924 56.916 1.00 33.36 C +ATOM 1058 O ASP A 261 26.750 22.930 57.629 1.00 33.36 O +ATOM 1059 CB ASP A 261 28.398 20.514 58.213 1.00 33.36 C +ATOM 1060 CG ASP A 261 27.383 19.658 58.987 1.00 33.36 C +ATOM 1061 OD1 ASP A 261 26.771 18.735 58.386 1.00 33.36 O +ATOM 1062 OD2 ASP A 261 27.215 19.904 60.204 1.00 33.36 O +ATOM 1063 N SER A 262 25.642 21.633 56.166 1.00 35.68 N +ATOM 1064 CA SER A 262 24.396 22.399 56.166 1.00 35.68 C +ATOM 1065 C SER A 262 23.188 21.480 56.338 1.00 35.68 C +ATOM 1066 O SER A 262 23.232 20.299 55.989 1.00 35.68 O +ATOM 1067 CB SER A 262 24.275 23.216 54.887 1.00 35.68 C +ATOM 1068 OG SER A 262 23.193 24.115 55.001 1.00 35.68 O +ATOM 1069 N SER A 263 22.098 22.012 56.897 1.00 39.47 N +ATOM 1070 CA SER A 263 20.812 21.304 56.934 1.00 39.47 C +ATOM 1071 C SER A 263 20.167 21.169 55.552 1.00 39.47 C +ATOM 1072 O SER A 263 19.459 20.189 55.328 1.00 39.47 O +ATOM 1073 CB SER A 263 19.859 22.000 57.902 1.00 39.47 C +ATOM 1074 OG SER A 263 19.674 23.358 57.556 1.00 39.47 O +ATOM 1075 N ASP A 264 20.455 22.100 54.640 1.00 36.18 N +ATOM 1076 CA ASP A 264 20.055 22.062 53.235 1.00 36.18 C +ATOM 1077 C ASP A 264 21.214 22.540 52.351 1.00 36.18 C +ATOM 1078 O ASP A 264 21.586 23.718 52.313 1.00 36.18 O +ATOM 1079 CB ASP A 264 18.796 22.908 53.016 1.00 36.18 C +ATOM 1080 CG ASP A 264 18.315 22.878 51.560 1.00 36.18 C +ATOM 1081 OD1 ASP A 264 18.906 22.123 50.747 1.00 36.18 O +ATOM 1082 OD2 ASP A 264 17.336 23.607 51.296 1.00 36.18 O +ATOM 1083 N GLY A 265 21.827 21.582 51.661 1.00 34.14 N +ATOM 1084 CA GLY A 265 22.966 21.835 50.801 1.00 34.14 C +ATOM 1085 C GLY A 265 22.626 22.246 49.387 1.00 34.14 C +ATOM 1086 O GLY A 265 23.527 22.700 48.683 1.00 34.14 O +ATOM 1087 N PHE A 266 21.360 22.171 48.978 1.00 32.38 N +ATOM 1088 CA PHE A 266 20.950 22.617 47.656 1.00 32.38 C +ATOM 1089 C PHE A 266 20.792 24.147 47.632 1.00 32.38 C +ATOM 1090 O PHE A 266 19.709 24.696 47.457 1.00 32.38 O +ATOM 1091 CB PHE A 266 19.709 21.842 47.200 1.00 32.38 C +ATOM 1092 CG PHE A 266 19.366 21.952 45.723 1.00 32.38 C +ATOM 1093 CD1 PHE A 266 18.204 21.311 45.258 1.00 32.38 C +ATOM 1094 CD2 PHE A 266 20.181 22.652 44.801 1.00 32.38 C +ATOM 1095 CE1 PHE A 266 17.855 21.365 43.899 1.00 32.38 C +ATOM 1096 CE2 PHE A 266 19.821 22.716 43.447 1.00 32.38 C +ATOM 1097 CZ PHE A 266 18.663 22.067 42.993 1.00 32.38 C +ATOM 1098 N SER A 267 21.903 24.860 47.848 1.00 32.90 N +ATOM 1099 CA SER A 267 21.917 26.315 47.995 1.00 32.90 C +ATOM 1100 C SER A 267 23.113 26.982 47.321 1.00 32.90 C +ATOM 1101 O SER A 267 24.232 26.458 47.286 1.00 32.90 O +ATOM 1102 CB SER A 267 21.840 26.705 49.478 1.00 32.90 C +ATOM 1103 OG SER A 267 22.995 26.327 50.206 1.00 32.90 O +ATOM 1104 N GLU A 268 22.896 28.214 46.852 1.00 31.89 N +ATOM 1105 CA GLU A 268 23.951 29.053 46.267 1.00 31.89 C +ATOM 1106 C GLU A 268 25.100 29.313 47.244 1.00 31.89 C +ATOM 1107 O GLU A 268 26.250 29.440 46.831 1.00 31.89 O +ATOM 1108 CB GLU A 268 23.372 30.394 45.802 1.00 31.89 C +ATOM 1109 CG GLU A 268 22.342 30.236 44.677 1.00 31.89 C +ATOM 1110 CD GLU A 268 21.982 31.592 44.061 1.00 31.89 C +ATOM 1111 OE1 GLU A 268 22.039 31.678 42.813 1.00 31.89 O +ATOM 1112 OE2 GLU A 268 21.691 32.523 44.843 1.00 31.89 O +ATOM 1113 N TRP A 269 24.823 29.316 48.549 1.00 31.97 N +ATOM 1114 CA TRP A 269 25.841 29.474 49.583 1.00 31.97 C +ATOM 1115 C TRP A 269 26.925 28.391 49.520 1.00 31.97 C +ATOM 1116 O TRP A 269 28.107 28.703 49.700 1.00 31.97 O +ATOM 1117 CB TRP A 269 25.152 29.478 50.950 1.00 31.97 C +ATOM 1118 CG TRP A 269 24.230 30.630 51.196 1.00 31.97 C +ATOM 1119 CD1 TRP A 269 22.903 30.549 51.444 1.00 31.97 C +ATOM 1120 CD2 TRP A 269 24.557 32.054 51.214 1.00 31.97 C +ATOM 1121 NE1 TRP A 269 22.388 31.820 51.626 1.00 31.97 N +ATOM 1122 CE2 TRP A 269 23.365 32.784 51.497 1.00 31.97 C +ATOM 1123 CE3 TRP A 269 25.751 32.794 51.061 1.00 31.97 C +ATOM 1124 CZ2 TRP A 269 23.351 34.180 51.607 1.00 31.97 C +ATOM 1125 CZ3 TRP A 269 25.747 34.198 51.171 1.00 31.97 C +ATOM 1126 CH2 TRP A 269 24.552 34.890 51.439 1.00 31.97 C +ATOM 1127 N GLN A 270 26.564 27.136 49.224 1.00 31.20 N +ATOM 1128 CA GLN A 270 27.559 26.074 49.041 1.00 31.20 C +ATOM 1129 C GLN A 270 28.340 26.267 47.751 1.00 31.20 C +ATOM 1130 O GLN A 270 29.563 26.172 47.773 1.00 31.20 O +ATOM 1131 CB GLN A 270 26.931 24.678 49.051 1.00 31.20 C +ATOM 1132 CG GLN A 270 26.237 24.413 50.386 1.00 31.20 C +ATOM 1133 CD GLN A 270 26.355 22.992 50.905 1.00 31.20 C +ATOM 1134 OE1 GLN A 270 27.102 22.134 50.467 1.00 31.20 O +ATOM 1135 NE2 GLN A 270 25.667 22.744 51.980 1.00 31.20 N +ATOM 1136 N LYS A 271 27.666 26.623 46.653 1.00 31.14 N +ATOM 1137 CA LYS A 271 28.328 26.912 45.375 1.00 31.14 C +ATOM 1138 C LYS A 271 29.355 28.042 45.505 1.00 31.14 C +ATOM 1139 O LYS A 271 30.472 27.900 45.016 1.00 31.14 O +ATOM 1140 CB LYS A 271 27.261 27.208 44.315 1.00 31.14 C +ATOM 1141 CG LYS A 271 27.883 27.440 42.933 1.00 31.14 C +ATOM 1142 CD LYS A 271 26.789 27.612 41.878 1.00 31.14 C +ATOM 1143 CE LYS A 271 27.424 27.855 40.507 1.00 31.14 C +ATOM 1144 NZ LYS A 271 26.387 27.953 39.451 1.00 31.14 N +ATOM 1145 N ILE A 272 29.017 29.122 46.211 1.00 31.04 N +ATOM 1146 CA ILE A 272 29.936 30.234 46.514 1.00 31.04 C +ATOM 1147 C ILE A 272 31.136 29.739 47.331 1.00 31.04 C +ATOM 1148 O ILE A 272 32.277 30.083 47.027 1.00 31.04 O +ATOM 1149 CB ILE A 272 29.178 31.364 47.253 1.00 31.04 C +ATOM 1150 CG1 ILE A 272 28.165 32.050 46.308 1.00 31.04 C +ATOM 1151 CG2 ILE A 272 30.146 32.419 47.828 1.00 31.04 C +ATOM 1152 CD1 ILE A 272 27.112 32.891 47.044 1.00 31.04 C +ATOM 1153 N THR A 273 30.887 28.918 48.356 1.00 30.88 N +ATOM 1154 CA THR A 273 31.939 28.378 49.231 1.00 30.88 C +ATOM 1155 C THR A 273 32.900 27.474 48.453 1.00 30.88 C +ATOM 1156 O THR A 273 34.112 27.672 48.519 1.00 30.88 O +ATOM 1157 CB THR A 273 31.322 27.620 50.417 1.00 30.88 C +ATOM 1158 OG1 THR A 273 30.447 28.465 51.137 1.00 30.88 O +ATOM 1159 CG2 THR A 273 32.375 27.107 51.397 1.00 30.88 C +ATOM 1160 N ILE A 274 32.377 26.525 47.665 1.00 30.70 N +ATOM 1161 CA ILE A 274 33.186 25.627 46.828 1.00 30.70 C +ATOM 1162 C ILE A 274 33.966 26.436 45.788 1.00 30.70 C +ATOM 1163 O ILE A 274 35.172 26.244 45.655 1.00 30.70 O +ATOM 1164 CB ILE A 274 32.312 24.547 46.144 1.00 30.70 C +ATOM 1165 CG1 ILE A 274 31.589 23.650 47.171 1.00 30.70 C +ATOM 1166 CG2 ILE A 274 33.190 23.655 45.245 1.00 30.70 C +ATOM 1167 CD1 ILE A 274 30.428 22.848 46.566 1.00 30.70 C +ATOM 1168 N GLY A 275 33.310 27.378 45.102 1.00 30.76 N +ATOM 1169 CA GLY A 275 33.945 28.238 44.102 1.00 30.76 C +ATOM 1170 C GLY A 275 35.142 29.005 44.664 1.00 30.76 C +ATOM 1171 O GLY A 275 36.214 28.982 44.069 1.00 30.76 O +ATOM 1172 N TRP A 276 35.003 29.591 45.858 1.00 30.64 N +ATOM 1173 CA TRP A 276 36.107 30.279 46.530 1.00 30.64 C +ATOM 1174 C TRP A 276 37.260 29.333 46.907 1.00 30.64 C +ATOM 1175 O TRP A 276 38.428 29.688 46.744 1.00 30.64 O +ATOM 1176 CB TRP A 276 35.559 30.983 47.770 1.00 30.64 C +ATOM 1177 CG TRP A 276 36.581 31.759 48.544 1.00 30.64 C +ATOM 1178 CD1 TRP A 276 36.938 33.040 48.305 1.00 30.64 C +ATOM 1179 CD2 TRP A 276 37.460 31.293 49.611 1.00 30.64 C +ATOM 1180 NE1 TRP A 276 37.934 33.418 49.185 1.00 30.64 N +ATOM 1181 CE2 TRP A 276 38.295 32.377 50.016 1.00 30.64 C +ATOM 1182 CE3 TRP A 276 37.650 30.059 50.265 1.00 30.64 C +ATOM 1183 CZ2 TRP A 276 39.229 32.258 51.055 1.00 30.64 C +ATOM 1184 CZ3 TRP A 276 38.605 29.917 51.286 1.00 30.64 C +ATOM 1185 CH2 TRP A 276 39.376 31.017 51.699 1.00 30.64 C +ATOM 1186 N ILE A 277 36.960 28.114 47.382 1.00 30.67 N +ATOM 1187 CA ILE A 277 37.992 27.104 47.686 1.00 30.67 C +ATOM 1188 C ILE A 277 38.762 26.734 46.412 1.00 30.67 C +ATOM 1189 O ILE A 277 39.991 26.674 46.439 1.00 30.67 O +ATOM 1190 CB ILE A 277 37.377 25.855 48.369 1.00 30.67 C +ATOM 1191 CG1 ILE A 277 36.869 26.223 49.780 1.00 30.67 C +ATOM 1192 CG2 ILE A 277 38.394 24.697 48.466 1.00 30.67 C +ATOM 1193 CD1 ILE A 277 36.027 25.134 50.456 1.00 30.67 C +ATOM 1194 N ARG A 278 38.055 26.526 45.293 1.00 30.85 N +ATOM 1195 CA ARG A 278 38.662 26.205 43.992 1.00 30.85 C +ATOM 1196 C ARG A 278 39.501 27.358 43.452 1.00 30.85 C +ATOM 1197 O ARG A 278 40.608 27.112 42.988 1.00 30.85 O +ATOM 1198 CB ARG A 278 37.576 25.799 42.977 1.00 30.85 C +ATOM 1199 CG ARG A 278 36.841 24.501 43.339 1.00 30.85 C +ATOM 1200 CD ARG A 278 37.745 23.268 43.266 1.00 30.85 C +ATOM 1201 NE ARG A 278 37.063 22.106 43.855 1.00 30.85 N +ATOM 1202 CZ ARG A 278 37.497 21.356 44.846 1.00 30.85 C +ATOM 1203 NH1 ARG A 278 38.633 21.554 45.454 1.00 30.85 N +ATOM 1204 NH2 ARG A 278 36.744 20.381 45.233 1.00 30.85 N +ATOM 1205 N GLU A 279 39.030 28.595 43.566 1.00 30.82 N +ATOM 1206 CA GLU A 279 39.786 29.789 43.167 1.00 30.82 C +ATOM 1207 C GLU A 279 41.099 29.919 43.954 1.00 30.82 C +ATOM 1208 O GLU A 279 42.152 30.175 43.373 1.00 30.82 O +ATOM 1209 CB GLU A 279 38.891 31.025 43.366 1.00 30.82 C +ATOM 1210 CG GLU A 279 39.532 32.325 42.857 1.00 30.82 C +ATOM 1211 CD GLU A 279 38.676 33.575 43.130 1.00 30.82 C +ATOM 1212 OE1 GLU A 279 39.168 34.680 42.816 1.00 30.82 O +ATOM 1213 OE2 GLU A 279 37.557 33.438 43.683 1.00 30.82 O +ATOM 1214 N LYS A 280 41.061 29.692 45.275 1.00 30.94 N +ATOM 1215 CA LYS A 280 42.222 29.886 46.154 1.00 30.94 C +ATOM 1216 C LYS A 280 43.212 28.717 46.154 1.00 30.94 C +ATOM 1217 O LYS A 280 44.415 28.943 46.273 1.00 30.94 O +ATOM 1218 CB LYS A 280 41.708 30.224 47.561 1.00 30.94 C +ATOM 1219 CG LYS A 280 42.847 30.647 48.501 1.00 30.94 C +ATOM 1220 CD LYS A 280 42.288 31.258 49.788 1.00 30.94 C +ATOM 1221 CE LYS A 280 43.434 31.676 50.715 1.00 30.94 C +ATOM 1222 NZ LYS A 280 42.946 32.443 51.888 1.00 30.94 N +ATOM 1223 N TYR A 281 42.731 27.477 46.055 1.00 30.97 N +ATOM 1224 CA TYR A 281 43.547 26.271 46.262 1.00 30.97 C +ATOM 1225 C TYR A 281 43.541 25.286 45.082 1.00 30.97 C +ATOM 1226 O TYR A 281 44.269 24.287 45.119 1.00 30.97 O +ATOM 1227 CB TYR A 281 43.090 25.567 47.548 1.00 30.97 C +ATOM 1228 CG TYR A 281 43.188 26.407 48.808 1.00 30.97 C +ATOM 1229 CD1 TYR A 281 44.444 26.651 49.398 1.00 30.97 C +ATOM 1230 CD2 TYR A 281 42.022 26.926 49.402 1.00 30.97 C +ATOM 1231 CE1 TYR A 281 44.536 27.416 50.576 1.00 30.97 C +ATOM 1232 CE2 TYR A 281 42.105 27.673 50.592 1.00 30.97 C +ATOM 1233 CZ TYR A 281 43.366 27.918 51.177 1.00 30.97 C +ATOM 1234 OH TYR A 281 43.447 28.647 52.318 1.00 30.97 O +ATOM 1235 N GLY A 282 42.737 25.528 44.043 1.00 31.47 N +ATOM 1236 CA GLY A 282 42.484 24.551 42.984 1.00 31.47 C +ATOM 1237 C GLY A 282 41.898 23.251 43.541 1.00 31.47 C +ATOM 1238 O GLY A 282 41.074 23.252 44.460 1.00 31.47 O +ATOM 1239 N ASP A 283 42.371 22.122 43.017 1.00 32.41 N +ATOM 1240 CA ASP A 283 41.966 20.784 43.470 1.00 32.41 C +ATOM 1241 C ASP A 283 42.821 20.237 44.629 1.00 32.41 C +ATOM 1242 O ASP A 283 42.640 19.088 45.032 1.00 32.41 O +ATOM 1243 CB ASP A 283 41.945 19.820 42.274 1.00 32.41 C +ATOM 1244 CG ASP A 283 40.815 20.116 41.283 1.00 32.41 C +ATOM 1245 OD1 ASP A 283 39.695 20.459 41.734 1.00 32.41 O +ATOM 1246 OD2 ASP A 283 41.075 19.952 40.074 1.00 32.41 O +ATOM 1247 N LYS A 284 43.757 21.027 45.185 1.00 31.82 N +ATOM 1248 CA LYS A 284 44.615 20.574 46.298 1.00 31.82 C +ATOM 1249 C LYS A 284 43.828 20.371 47.592 1.00 31.82 C +ATOM 1250 O LYS A 284 44.085 19.420 48.317 1.00 31.82 O +ATOM 1251 CB LYS A 284 45.762 21.561 46.554 1.00 31.82 C +ATOM 1252 CG LYS A 284 46.775 21.626 45.404 1.00 31.82 C +ATOM 1253 CD LYS A 284 47.905 22.596 45.768 1.00 31.82 C +ATOM 1254 CE LYS A 284 48.950 22.644 44.651 1.00 31.82 C +ATOM 1255 NZ LYS A 284 50.033 23.606 44.974 1.00 31.82 N +ATOM 1256 N VAL A 285 42.874 21.259 47.880 1.00 31.07 N +ATOM 1257 CA VAL A 285 41.999 21.145 49.055 1.00 31.07 C +ATOM 1258 C VAL A 285 40.766 20.341 48.673 1.00 31.07 C +ATOM 1259 O VAL A 285 40.015 20.736 47.777 1.00 31.07 O +ATOM 1260 CB VAL A 285 41.634 22.525 49.630 1.00 31.07 C +ATOM 1261 CG1 VAL A 285 40.580 22.440 50.743 1.00 31.07 C +ATOM 1262 CG2 VAL A 285 42.886 23.180 50.227 1.00 31.07 C +ATOM 1263 N LYS A 286 40.552 19.220 49.364 1.00 30.67 N +ATOM 1264 CA LYS A 286 39.384 18.368 49.143 1.00 30.67 C +ATOM 1265 C LYS A 286 38.196 18.875 49.950 1.00 30.67 C +ATOM 1266 O LYS A 286 38.339 19.145 51.138 1.00 30.67 O +ATOM 1267 CB LYS A 286 39.703 16.896 49.433 1.00 30.67 C +ATOM 1268 CG LYS A 286 40.954 16.371 48.710 1.00 30.67 C +ATOM 1269 CD LYS A 286 40.949 16.609 47.192 1.00 30.67 C +ATOM 1270 CE LYS A 286 42.249 16.032 46.639 1.00 30.67 C +ATOM 1271 NZ LYS A 286 42.432 16.319 45.204 1.00 30.67 N +ATOM 1272 N VAL A 287 37.031 19.004 49.326 1.00 30.64 N +ATOM 1273 CA VAL A 287 35.808 19.510 49.958 1.00 30.64 C +ATOM 1274 C VAL A 287 34.578 18.739 49.488 1.00 30.64 C +ATOM 1275 O VAL A 287 34.222 18.759 48.313 1.00 30.64 O +ATOM 1276 CB VAL A 287 35.668 21.037 49.776 1.00 30.64 C +ATOM 1277 CG1 VAL A 287 35.649 21.524 48.320 1.00 30.64 C +ATOM 1278 CG2 VAL A 287 34.425 21.564 50.501 1.00 30.64 C +ATOM 1279 N GLY A 288 33.903 18.069 50.417 1.00 30.67 N +ATOM 1280 CA GLY A 288 32.564 17.531 50.205 1.00 30.67 C +ATOM 1281 C GLY A 288 31.465 18.537 50.528 1.00 30.67 C +ATOM 1282 O GLY A 288 31.691 19.522 51.235 1.00 30.67 O +ATOM 1283 N ALA A 289 30.264 18.263 50.033 1.00 30.67 N +ATOM 1284 CA ALA A 289 29.116 19.159 50.131 1.00 30.67 C +ATOM 1285 C ALA A 289 27.820 18.401 50.455 1.00 30.67 C +ATOM 1286 O ALA A 289 27.707 17.205 50.186 1.00 30.67 O +ATOM 1287 CB ALA A 289 29.032 19.922 48.807 1.00 30.67 C +ATOM 1288 N GLY A 290 26.839 19.091 51.036 1.00 31.17 N +ATOM 1289 CA GLY A 290 25.548 18.518 51.424 1.00 31.17 C +ATOM 1290 C GLY A 290 24.858 19.280 52.568 1.00 31.17 C +ATOM 1291 O GLY A 290 25.365 20.283 53.051 1.00 31.17 O +ATOM 1292 N ASN A 291 23.692 18.863 53.049 1.00 31.33 N +ATOM 1293 CA ASN A 291 23.011 17.620 52.692 1.00 31.33 C +ATOM 1294 C ASN A 291 21.997 17.815 51.566 1.00 31.33 C +ATOM 1295 O ASN A 291 21.326 18.837 51.514 1.00 31.33 O +ATOM 1296 CB ASN A 291 22.367 16.994 53.938 1.00 31.33 C +ATOM 1297 CG ASN A 291 23.398 16.471 54.920 1.00 31.33 C +ATOM 1298 OD1 ASN A 291 24.531 16.905 54.936 1.00 31.33 O +ATOM 1299 ND2 ASN A 291 23.067 15.516 55.758 1.00 31.33 N +ATOM 1300 N ILE A 292 21.858 16.818 50.695 1.00 31.01 N +ATOM 1301 CA ILE A 292 20.831 16.779 49.639 1.00 31.01 C +ATOM 1302 C ILE A 292 20.015 15.482 49.730 1.00 31.01 C +ATOM 1303 O ILE A 292 20.363 14.585 50.503 1.00 31.01 O +ATOM 1304 CB ILE A 292 21.456 17.050 48.247 1.00 31.01 C +ATOM 1305 CG1 ILE A 292 22.579 16.084 47.813 1.00 31.01 C +ATOM 1306 CG2 ILE A 292 22.011 18.486 48.209 1.00 31.01 C +ATOM 1307 CD1 ILE A 292 22.080 14.694 47.425 1.00 31.01 C +ATOM 1308 N VAL A 293 18.901 15.378 48.996 1.00 31.33 N +ATOM 1309 CA VAL A 293 18.007 14.196 49.036 1.00 31.33 C +ATOM 1310 C VAL A 293 17.478 13.736 47.672 1.00 31.33 C +ATOM 1311 O VAL A 293 16.644 12.826 47.619 1.00 31.33 O +ATOM 1312 CB VAL A 293 16.827 14.397 50.011 1.00 31.33 C +ATOM 1313 CG1 VAL A 293 17.284 14.425 51.470 1.00 31.33 C +ATOM 1314 CG2 VAL A 293 16.005 15.655 49.709 1.00 31.33 C +ATOM 1315 N ASP A 294 17.924 14.347 46.578 1.00 31.40 N +ATOM 1316 CA ASP A 294 17.494 14.028 45.216 1.00 31.40 C +ATOM 1317 C ASP A 294 18.600 14.275 44.177 1.00 31.40 C +ATOM 1318 O ASP A 294 19.666 14.826 44.481 1.00 31.40 O +ATOM 1319 CB ASP A 294 16.208 14.805 44.880 1.00 31.40 C +ATOM 1320 CG ASP A 294 16.334 16.334 44.883 1.00 31.40 C +ATOM 1321 OD1 ASP A 294 17.470 16.856 44.793 1.00 31.40 O +ATOM 1322 OD2 ASP A 294 15.260 16.961 44.990 1.00 31.40 O +ATOM 1323 N GLY A 295 18.350 13.829 42.944 1.00 31.01 N +ATOM 1324 CA GLY A 295 19.292 13.972 41.835 1.00 31.01 C +ATOM 1325 C GLY A 295 19.615 15.429 41.473 1.00 31.01 C +ATOM 1326 O GLY A 295 20.761 15.737 41.153 1.00 31.01 O +ATOM 1327 N GLU A 296 18.658 16.356 41.575 1.00 31.17 N +ATOM 1328 CA GLU A 296 18.898 17.774 41.258 1.00 31.17 C +ATOM 1329 C GLU A 296 19.920 18.399 42.214 1.00 31.17 C +ATOM 1330 O GLU A 296 20.886 19.028 41.767 1.00 31.17 O +ATOM 1331 CB GLU A 296 17.584 18.575 41.281 1.00 31.17 C +ATOM 1332 CG GLU A 296 16.643 18.197 40.122 1.00 31.17 C +ATOM 1333 CD GLU A 296 15.334 19.013 40.090 1.00 31.17 C +ATOM 1334 OE1 GLU A 296 14.371 18.532 39.438 1.00 31.17 O +ATOM 1335 OE2 GLU A 296 15.289 20.115 40.681 1.00 31.17 O +ATOM 1336 N GLY A 297 19.775 18.152 43.520 1.00 30.88 N +ATOM 1337 CA GLY A 297 20.746 18.594 44.517 1.00 30.88 C +ATOM 1338 C GLY A 297 22.127 17.961 44.318 1.00 30.88 C +ATOM 1339 O GLY A 297 23.143 18.647 44.470 1.00 30.88 O +ATOM 1340 N PHE A 298 22.186 16.681 43.925 1.00 30.61 N +ATOM 1341 CA PHE A 298 23.449 16.037 43.552 1.00 30.61 C +ATOM 1342 C PHE A 298 24.108 16.754 42.373 1.00 30.61 C +ATOM 1343 O PHE A 298 25.262 17.171 42.485 1.00 30.61 O +ATOM 1344 CB PHE A 298 23.244 14.544 43.250 1.00 30.61 C +ATOM 1345 CG PHE A 298 24.470 13.872 42.639 1.00 30.61 C +ATOM 1346 CD1 PHE A 298 24.686 13.949 41.251 1.00 30.61 C +ATOM 1347 CD2 PHE A 298 25.436 13.247 43.451 1.00 30.61 C +ATOM 1348 CE1 PHE A 298 25.887 13.489 40.690 1.00 30.61 C +ATOM 1349 CE2 PHE A 298 26.627 12.753 42.883 1.00 30.61 C +ATOM 1350 CZ PHE A 298 26.868 12.911 41.509 1.00 30.61 C +ATOM 1351 N ARG A 299 23.380 16.932 41.261 1.00 30.70 N +ATOM 1352 CA ARG A 299 23.917 17.532 40.034 1.00 30.70 C +ATOM 1353 C ARG A 299 24.439 18.939 40.302 1.00 30.70 C +ATOM 1354 O ARG A 299 25.553 19.263 39.908 1.00 30.70 O +ATOM 1355 CB ARG A 299 22.834 17.520 38.940 1.00 30.70 C +ATOM 1356 CG ARG A 299 23.307 18.099 37.595 1.00 30.70 C +ATOM 1357 CD ARG A 299 24.490 17.345 36.968 1.00 30.70 C +ATOM 1358 NE ARG A 299 24.130 15.954 36.610 1.00 30.70 N +ATOM 1359 CZ ARG A 299 24.683 15.208 35.675 1.00 30.70 C +ATOM 1360 NH1 ARG A 299 25.735 15.555 35.002 1.00 30.70 N +ATOM 1361 NH2 ARG A 299 24.162 14.075 35.343 1.00 30.70 N +ATOM 1362 N TYR A 300 23.683 19.734 41.055 1.00 30.64 N +ATOM 1363 CA TYR A 300 24.067 21.095 41.411 1.00 30.64 C +ATOM 1364 C TYR A 300 25.405 21.168 42.166 1.00 30.64 C +ATOM 1365 O TYR A 300 26.269 21.981 41.829 1.00 30.64 O +ATOM 1366 CB TYR A 300 22.938 21.701 42.248 1.00 30.64 C +ATOM 1367 CG TYR A 300 23.170 23.157 42.585 1.00 30.64 C +ATOM 1368 CD1 TYR A 300 23.800 23.517 43.791 1.00 30.64 C +ATOM 1369 CD2 TYR A 300 22.751 24.149 41.680 1.00 30.64 C +ATOM 1370 CE1 TYR A 300 24.023 24.875 44.088 1.00 30.64 C +ATOM 1371 CE2 TYR A 300 22.941 25.508 41.987 1.00 30.64 C +ATOM 1372 CZ TYR A 300 23.585 25.869 43.188 1.00 30.64 C +ATOM 1373 OH TYR A 300 23.785 27.180 43.473 1.00 30.64 O +ATOM 1374 N LEU A 301 25.597 20.320 43.182 1.00 30.67 N +ATOM 1375 CA LEU A 301 26.828 20.301 43.980 1.00 30.67 C +ATOM 1376 C LEU A 301 28.003 19.644 43.243 1.00 30.67 C +ATOM 1377 O LEU A 301 29.150 20.066 43.419 1.00 30.67 O +ATOM 1378 CB LEU A 301 26.551 19.600 45.321 1.00 30.67 C +ATOM 1379 CG LEU A 301 25.614 20.381 46.261 1.00 30.67 C +ATOM 1380 CD1 LEU A 301 25.401 19.572 47.542 1.00 30.67 C +ATOM 1381 CD2 LEU A 301 26.195 21.751 46.629 1.00 30.67 C +ATOM 1382 N ALA A 302 27.724 18.653 42.395 1.00 30.67 N +ATOM 1383 CA ALA A 302 28.705 18.039 41.510 1.00 30.67 C +ATOM 1384 C ALA A 302 29.274 19.071 40.524 1.00 30.67 C +ATOM 1385 O ALA A 302 30.490 19.249 40.453 1.00 30.67 O +ATOM 1386 CB ALA A 302 28.036 16.866 40.784 1.00 30.67 C +ATOM 1387 N ASP A 303 28.401 19.822 39.846 1.00 30.82 N +ATOM 1388 CA ASP A 303 28.779 20.876 38.898 1.00 30.82 C +ATOM 1389 C ASP A 303 29.481 22.057 39.586 1.00 30.82 C +ATOM 1390 O ASP A 303 30.341 22.709 38.993 1.00 30.82 O +ATOM 1391 CB ASP A 303 27.530 21.376 38.152 1.00 30.82 C +ATOM 1392 CG ASP A 303 26.908 20.351 37.193 1.00 30.82 C +ATOM 1393 OD1 ASP A 303 27.594 19.370 36.826 1.00 30.82 O +ATOM 1394 OD2 ASP A 303 25.746 20.581 36.790 1.00 30.82 O +ATOM 1395 N ALA A 304 29.168 22.324 40.859 1.00 30.88 N +ATOM 1396 CA ALA A 304 29.898 23.300 41.668 1.00 30.88 C +ATOM 1397 C ALA A 304 31.340 22.860 41.989 1.00 30.88 C +ATOM 1398 O ALA A 304 32.168 23.703 42.332 1.00 30.88 O +ATOM 1399 CB ALA A 304 29.099 23.587 42.945 1.00 30.88 C +ATOM 1400 N GLY A 305 31.654 21.564 41.866 1.00 30.97 N +ATOM 1401 CA GLY A 305 33.001 21.015 42.016 1.00 30.97 C +ATOM 1402 C GLY A 305 33.268 20.283 43.333 1.00 30.97 C +ATOM 1403 O GLY A 305 34.443 20.143 43.691 1.00 30.97 O +ATOM 1404 N ALA A 306 32.232 19.824 44.047 1.00 30.67 N +ATOM 1405 CA ALA A 306 32.379 19.018 45.264 1.00 30.67 C +ATOM 1406 C ALA A 306 33.152 17.707 45.007 1.00 30.67 C +ATOM 1407 O ALA A 306 33.004 17.076 43.963 1.00 30.67 O +ATOM 1408 CB ALA A 306 30.988 18.735 45.847 1.00 30.67 C +ATOM 1409 N ASP A 307 33.979 17.287 45.966 1.00 30.58 N +ATOM 1410 CA ASP A 307 34.764 16.047 45.896 1.00 30.58 C +ATOM 1411 C ASP A 307 34.004 14.822 46.407 1.00 30.58 C +ATOM 1412 O ASP A 307 34.337 13.706 46.031 1.00 30.58 O +ATOM 1413 CB ASP A 307 36.085 16.211 46.653 1.00 30.58 C +ATOM 1414 CG ASP A 307 36.979 17.260 46.006 1.00 30.58 C +ATOM 1415 OD1 ASP A 307 37.288 17.149 44.803 1.00 30.58 O +ATOM 1416 OD2 ASP A 307 37.337 18.227 46.707 1.00 30.58 O +ATOM 1417 N PHE A 308 32.974 15.009 47.227 1.00 30.55 N +ATOM 1418 CA PHE A 308 31.982 13.988 47.566 1.00 30.55 C +ATOM 1419 C PHE A 308 30.667 14.665 47.950 1.00 30.55 C +ATOM 1420 O PHE A 308 30.651 15.853 48.289 1.00 30.55 O +ATOM 1421 CB PHE A 308 32.482 13.083 48.698 1.00 30.55 C +ATOM 1422 CG PHE A 308 32.547 13.725 50.072 1.00 30.55 C +ATOM 1423 CD1 PHE A 308 33.755 14.268 50.545 1.00 30.55 C +ATOM 1424 CD2 PHE A 308 31.402 13.750 50.892 1.00 30.55 C +ATOM 1425 CE1 PHE A 308 33.816 14.821 51.836 1.00 30.55 C +ATOM 1426 CE2 PHE A 308 31.463 14.307 52.180 1.00 30.55 C +ATOM 1427 CZ PHE A 308 32.675 14.832 52.655 1.00 30.55 C +ATOM 1428 N ILE A 309 29.563 13.917 47.925 1.00 30.58 N +ATOM 1429 CA ILE A 309 28.233 14.459 48.229 1.00 30.58 C +ATOM 1430 C ILE A 309 27.591 13.711 49.397 1.00 30.58 C +ATOM 1431 O ILE A 309 27.524 12.483 49.414 1.00 30.58 O +ATOM 1432 CB ILE A 309 27.374 14.524 46.948 1.00 30.58 C +ATOM 1433 CG1 ILE A 309 27.962 15.623 46.028 1.00 30.58 C +ATOM 1434 CG2 ILE A 309 25.903 14.842 47.280 1.00 30.58 C +ATOM 1435 CD1 ILE A 309 27.433 15.630 44.597 1.00 30.58 C +ATOM 1436 N LYS A 310 27.118 14.464 50.393 1.00 30.61 N +ATOM 1437 CA LYS A 310 26.457 13.964 51.602 1.00 30.61 C +ATOM 1438 C LYS A 310 24.938 13.960 51.424 1.00 30.61 C +ATOM 1439 O LYS A 310 24.331 14.996 51.155 1.00 30.61 O +ATOM 1440 CB LYS A 310 26.958 14.774 52.814 1.00 30.61 C +ATOM 1441 CG LYS A 310 26.528 14.172 54.159 1.00 30.61 C +ATOM 1442 CD LYS A 310 27.412 14.638 55.331 1.00 30.61 C +ATOM 1443 CE LYS A 310 27.337 16.117 55.753 1.00 30.61 C +ATOM 1444 NZ LYS A 310 26.320 16.361 56.806 1.00 30.61 N +ATOM 1445 N ILE A 311 24.324 12.790 51.577 1.00 30.76 N +ATOM 1446 CA ILE A 311 22.894 12.547 51.366 1.00 30.76 C +ATOM 1447 C ILE A 311 22.198 12.430 52.718 1.00 30.76 C +ATOM 1448 O ILE A 311 22.560 11.572 53.523 1.00 30.76 O +ATOM 1449 CB ILE A 311 22.653 11.264 50.537 1.00 30.76 C +ATOM 1450 CG1 ILE A 311 23.442 11.277 49.209 1.00 30.76 C +ATOM 1451 CG2 ILE A 311 21.141 11.088 50.271 1.00 30.76 C +ATOM 1452 CD1 ILE A 311 23.436 9.912 48.514 1.00 30.76 C +ATOM 1453 N GLY A 312 21.145 13.218 52.930 1.00 31.37 N +ATOM 1454 CA GLY A 312 20.190 12.979 54.010 1.00 31.37 C +ATOM 1455 C GLY A 312 19.664 14.238 54.684 1.00 31.37 C +ATOM 1456 O GLY A 312 20.407 14.970 55.327 1.00 31.37 O +ATOM 1457 N ILE A 313 18.352 14.457 54.622 1.00 32.41 N +ATOM 1458 CA ILE A 313 17.670 15.534 55.354 1.00 32.41 C +ATOM 1459 C ILE A 313 16.587 14.909 56.239 1.00 32.41 C +ATOM 1460 O ILE A 313 15.704 14.169 55.781 1.00 32.41 O +ATOM 1461 CB ILE A 313 17.158 16.637 54.397 1.00 32.41 C +ATOM 1462 CG1 ILE A 313 18.345 17.298 53.651 1.00 32.41 C +ATOM 1463 CG2 ILE A 313 16.353 17.699 55.169 1.00 32.41 C +ATOM 1464 CD1 ILE A 313 17.923 18.305 52.574 1.00 32.41 C +ATOM 1465 N GLY A 314 16.689 15.157 57.547 1.00 35.78 N +ATOM 1466 CA GLY A 314 15.727 14.687 58.549 1.00 35.78 C +ATOM 1467 C GLY A 314 15.754 13.189 58.873 1.00 35.78 C +ATOM 1468 O GLY A 314 14.844 12.714 59.536 1.00 35.78 O +ATOM 1469 N ARG A 329 13.300 13.499 53.023 1.00 32.50 N +ATOM 1470 CA ARG A 329 12.933 12.303 52.247 1.00 32.50 C +ATOM 1471 C ARG A 329 13.444 11.031 52.909 1.00 32.50 C +ATOM 1472 O ARG A 329 14.525 11.049 53.492 1.00 32.50 O +ATOM 1473 CB ARG A 329 13.540 12.446 50.847 1.00 32.50 C +ATOM 1474 CG ARG A 329 12.927 11.489 49.818 1.00 32.50 C +ATOM 1475 CD ARG A 329 13.588 11.783 48.475 1.00 32.50 C +ATOM 1476 NE ARG A 329 12.996 11.030 47.355 1.00 32.50 N +ATOM 1477 CZ ARG A 329 13.486 11.084 46.128 1.00 32.50 C +ATOM 1478 NH1 ARG A 329 14.535 11.797 45.862 1.00 32.50 N +ATOM 1479 NH2 ARG A 329 12.953 10.432 45.136 1.00 32.50 N +ATOM 1480 N GLY A 330 12.698 9.930 52.830 1.00 31.57 N +ATOM 1481 CA GLY A 330 13.143 8.612 53.299 1.00 31.57 C +ATOM 1482 C GLY A 330 14.568 8.298 52.828 1.00 31.57 C +ATOM 1483 O GLY A 330 14.876 8.465 51.649 1.00 31.57 O +ATOM 1484 N GLN A 331 15.450 7.920 53.761 1.00 31.04 N +ATOM 1485 CA GLN A 331 16.895 7.881 53.500 1.00 31.04 C +ATOM 1486 C GLN A 331 17.269 6.854 52.425 1.00 31.04 C +ATOM 1487 O GLN A 331 18.136 7.137 51.605 1.00 31.04 O +ATOM 1488 CB GLN A 331 17.650 7.601 54.812 1.00 31.04 C +ATOM 1489 CG GLN A 331 19.179 7.716 54.670 1.00 31.04 C +ATOM 1490 CD GLN A 331 19.691 9.129 54.412 1.00 31.04 C +ATOM 1491 OE1 GLN A 331 18.966 10.109 54.523 1.00 31.04 O +ATOM 1492 NE2 GLN A 331 20.958 9.260 54.095 1.00 31.04 N +ATOM 1493 N ALA A 332 16.598 5.698 52.389 1.00 30.97 N +ATOM 1494 CA ALA A 332 16.849 4.688 51.363 1.00 30.97 C +ATOM 1495 C ALA A 332 16.482 5.229 49.975 1.00 30.97 C +ATOM 1496 O ALA A 332 17.301 5.190 49.063 1.00 30.97 O +ATOM 1497 CB ALA A 332 16.077 3.412 51.713 1.00 30.97 C +ATOM 1498 N THR A 333 15.296 5.834 49.834 1.00 31.10 N +ATOM 1499 CA THR A 333 14.884 6.472 48.574 1.00 31.10 C +ATOM 1500 C THR A 333 15.812 7.615 48.167 1.00 31.10 C +ATOM 1501 O THR A 333 16.130 7.734 46.988 1.00 31.10 O +ATOM 1502 CB THR A 333 13.450 7.003 48.676 1.00 31.10 C +ATOM 1503 OG1 THR A 333 12.588 5.953 49.019 1.00 31.10 O +ATOM 1504 CG2 THR A 333 12.902 7.566 47.369 1.00 31.10 C +ATOM 1505 N ALA A 334 16.266 8.439 49.115 1.00 30.91 N +ATOM 1506 CA ALA A 334 17.212 9.518 48.833 1.00 30.91 C +ATOM 1507 C ALA A 334 18.548 8.980 48.295 1.00 30.91 C +ATOM 1508 O ALA A 334 19.060 9.503 47.310 1.00 30.91 O +ATOM 1509 CB ALA A 334 17.408 10.349 50.107 1.00 30.91 C +ATOM 1510 N VAL A 335 19.089 7.914 48.901 1.00 30.70 N +ATOM 1511 CA VAL A 335 20.314 7.257 48.418 1.00 30.70 C +ATOM 1512 C VAL A 335 20.116 6.699 47.011 1.00 30.70 C +ATOM 1513 O VAL A 335 20.927 6.993 46.140 1.00 30.70 O +ATOM 1514 CB VAL A 335 20.807 6.177 49.403 1.00 30.70 C +ATOM 1515 CG1 VAL A 335 21.955 5.327 48.837 1.00 30.70 C +ATOM 1516 CG2 VAL A 335 21.334 6.833 50.689 1.00 30.70 C +ATOM 1517 N ILE A 336 19.031 5.961 46.768 1.00 31.10 N +ATOM 1518 CA ILE A 336 18.738 5.351 45.459 1.00 31.10 C +ATOM 1519 C ILE A 336 18.660 6.414 44.352 1.00 31.10 C +ATOM 1520 O ILE A 336 19.286 6.256 43.306 1.00 31.10 O +ATOM 1521 CB ILE A 336 17.438 4.516 45.560 1.00 31.10 C +ATOM 1522 CG1 ILE A 336 17.658 3.277 46.463 1.00 31.10 C +ATOM 1523 CG2 ILE A 336 16.930 4.054 44.179 1.00 31.10 C +ATOM 1524 CD1 ILE A 336 16.354 2.627 46.946 1.00 31.10 C +ATOM 1525 N ASP A 337 17.925 7.505 44.585 1.00 31.07 N +ATOM 1526 CA ASP A 337 17.727 8.564 43.586 1.00 31.07 C +ATOM 1527 C ASP A 337 19.031 9.319 43.277 1.00 31.07 C +ATOM 1528 O ASP A 337 19.407 9.510 42.121 1.00 31.07 O +ATOM 1529 CB ASP A 337 16.636 9.514 44.100 1.00 31.07 C +ATOM 1530 CG ASP A 337 16.141 10.503 43.042 1.00 31.07 C +ATOM 1531 OD1 ASP A 337 16.138 10.164 41.848 1.00 31.07 O +ATOM 1532 OD2 ASP A 337 15.633 11.570 43.457 1.00 31.07 O +ATOM 1533 N VAL A 338 19.786 9.681 44.316 1.00 30.61 N +ATOM 1534 CA VAL A 338 21.072 10.369 44.155 1.00 30.61 C +ATOM 1535 C VAL A 338 22.110 9.479 43.470 1.00 30.61 C +ATOM 1536 O VAL A 338 22.858 9.956 42.620 1.00 30.61 O +ATOM 1537 CB VAL A 338 21.591 10.838 45.519 1.00 30.61 C +ATOM 1538 CG1 VAL A 338 23.033 11.347 45.423 1.00 30.61 C +ATOM 1539 CG2 VAL A 338 20.708 11.958 46.075 1.00 30.61 C +ATOM 1540 N VAL A 339 22.166 8.188 43.809 1.00 30.61 N +ATOM 1541 CA VAL A 339 23.084 7.231 43.173 1.00 30.61 C +ATOM 1542 C VAL A 339 22.773 7.082 41.685 1.00 30.61 C +ATOM 1543 O VAL A 339 23.704 7.016 40.882 1.00 30.61 O +ATOM 1544 CB VAL A 339 23.024 5.868 43.889 1.00 30.61 C +ATOM 1545 CG1 VAL A 339 23.670 4.731 43.083 1.00 30.61 C +ATOM 1546 CG2 VAL A 339 23.753 5.956 45.237 1.00 30.61 C +ATOM 1547 N ALA A 340 21.493 7.057 41.302 1.00 30.67 N +ATOM 1548 CA ALA A 340 21.100 7.002 39.897 1.00 30.67 C +ATOM 1549 C ALA A 340 21.633 8.221 39.122 1.00 30.67 C +ATOM 1550 O ALA A 340 22.267 8.055 38.076 1.00 30.67 O +ATOM 1551 CB ALA A 340 19.575 6.875 39.810 1.00 30.67 C +ATOM 1552 N GLU A 341 21.472 9.428 39.670 1.00 30.58 N +ATOM 1553 CA GLU A 341 22.010 10.643 39.049 1.00 30.58 C +ATOM 1554 C GLU A 341 23.550 10.679 39.062 1.00 30.58 C +ATOM 1555 O GLU A 341 24.160 11.087 38.072 1.00 30.58 O +ATOM 1556 CB GLU A 341 21.403 11.879 39.722 1.00 30.58 C +ATOM 1557 CG GLU A 341 21.799 13.190 39.025 1.00 30.58 C +ATOM 1558 CD GLU A 341 21.340 13.339 37.566 1.00 30.58 C +ATOM 1559 OE1 GLU A 341 21.935 14.192 36.867 1.00 30.58 O +ATOM 1560 OE2 GLU A 341 20.435 12.615 37.097 1.00 30.58 O +ATOM 1561 N ARG A 342 24.203 10.179 40.122 1.00 30.64 N +ATOM 1562 CA ARG A 342 25.668 10.032 40.179 1.00 30.64 C +ATOM 1563 C ARG A 342 26.193 9.138 39.071 1.00 30.64 C +ATOM 1564 O ARG A 342 27.193 9.471 38.438 1.00 30.64 O +ATOM 1565 CB ARG A 342 26.095 9.517 41.564 1.00 30.64 C +ATOM 1566 CG ARG A 342 27.621 9.425 41.765 1.00 30.64 C +ATOM 1567 CD ARG A 342 28.336 8.181 41.198 1.00 30.64 C +ATOM 1568 NE ARG A 342 27.826 6.920 41.761 1.00 30.64 N +ATOM 1569 CZ ARG A 342 28.189 6.390 42.915 1.00 30.64 C +ATOM 1570 NH1 ARG A 342 29.048 6.989 43.690 1.00 30.64 N +ATOM 1571 NH2 ARG A 342 27.670 5.268 43.317 1.00 30.64 N +ATOM 1572 N ASN A 343 25.547 7.999 38.843 1.00 30.76 N +ATOM 1573 CA ASN A 343 25.969 7.055 37.812 1.00 30.76 C +ATOM 1574 C ASN A 343 25.846 7.685 36.418 1.00 30.76 C +ATOM 1575 O ASN A 343 26.796 7.621 35.642 1.00 30.76 O +ATOM 1576 CB ASN A 343 25.154 5.762 37.950 1.00 30.76 C +ATOM 1577 CG ASN A 343 25.455 4.980 39.221 1.00 30.76 C +ATOM 1578 OD1 ASN A 343 26.404 5.221 39.959 1.00 30.76 O +ATOM 1579 ND2 ASN A 343 24.651 3.984 39.507 1.00 30.76 N +ATOM 1580 N LYS A 344 24.743 8.394 36.150 1.00 30.61 N +ATOM 1581 CA LYS A 344 24.563 9.177 34.920 1.00 30.61 C +ATOM 1582 C LYS A 344 25.637 10.262 34.765 1.00 30.61 C +ATOM 1583 O LYS A 344 26.245 10.374 33.705 1.00 30.61 O +ATOM 1584 CB LYS A 344 23.146 9.756 34.937 1.00 30.61 C +ATOM 1585 CG LYS A 344 22.852 10.599 33.693 1.00 30.61 C +ATOM 1586 CD LYS A 344 21.441 11.171 33.801 1.00 30.61 C +ATOM 1587 CE LYS A 344 21.248 12.231 32.725 1.00 30.61 C +ATOM 1588 NZ LYS A 344 20.023 13.009 32.989 1.00 30.61 N +ATOM 1589 N TYR A 345 25.933 11.015 35.826 1.00 30.58 N +ATOM 1590 CA TYR A 345 27.015 12.009 35.835 1.00 30.58 C +ATOM 1591 C TYR A 345 28.368 11.380 35.472 1.00 30.58 C +ATOM 1592 O TYR A 345 29.123 11.942 34.677 1.00 30.58 O +ATOM 1593 CB TYR A 345 27.076 12.680 37.217 1.00 30.58 C +ATOM 1594 CG TYR A 345 28.127 13.768 37.350 1.00 30.58 C +ATOM 1595 CD1 TYR A 345 29.482 13.448 37.556 1.00 30.58 C +ATOM 1596 CD2 TYR A 345 27.747 15.119 37.290 1.00 30.58 C +ATOM 1597 CE1 TYR A 345 30.444 14.475 37.619 1.00 30.58 C +ATOM 1598 CE2 TYR A 345 28.708 16.142 37.288 1.00 30.58 C +ATOM 1599 CZ TYR A 345 30.068 15.820 37.442 1.00 30.58 C +ATOM 1600 OH TYR A 345 30.995 16.812 37.473 1.00 30.58 O +ATOM 1601 N PHE A 346 28.676 10.202 36.022 1.00 30.70 N +ATOM 1602 CA PHE A 346 29.916 9.487 35.722 1.00 30.70 C +ATOM 1603 C PHE A 346 29.985 9.008 34.266 1.00 30.70 C +ATOM 1604 O PHE A 346 31.049 9.081 33.651 1.00 30.70 O +ATOM 1605 CB PHE A 346 30.074 8.306 36.691 1.00 30.70 C +ATOM 1606 CG PHE A 346 31.315 7.462 36.453 1.00 30.70 C +ATOM 1607 CD1 PHE A 346 31.225 6.064 36.327 1.00 30.70 C +ATOM 1608 CD2 PHE A 346 32.574 8.077 36.364 1.00 30.70 C +ATOM 1609 CE1 PHE A 346 32.392 5.295 36.157 1.00 30.70 C +ATOM 1610 CE2 PHE A 346 33.743 7.317 36.215 1.00 30.70 C +ATOM 1611 CZ PHE A 346 33.653 5.919 36.121 1.00 30.70 C +ATOM 1612 N GLU A 347 28.868 8.549 33.699 1.00 30.79 N +ATOM 1613 CA GLU A 347 28.773 8.173 32.283 1.00 30.79 C +ATOM 1614 C GLU A 347 28.995 9.373 31.351 1.00 30.79 C +ATOM 1615 O GLU A 347 29.722 9.252 30.367 1.00 30.79 O +ATOM 1616 CB GLU A 347 27.410 7.522 32.004 1.00 30.79 C +ATOM 1617 CG GLU A 347 27.307 6.111 32.606 1.00 30.79 C +ATOM 1618 CD GLU A 347 25.901 5.498 32.494 1.00 30.79 C +ATOM 1619 OE1 GLU A 347 25.735 4.381 33.038 1.00 30.79 O +ATOM 1620 OE2 GLU A 347 25.019 6.107 31.847 1.00 30.79 O +ATOM 1621 N GLU A 348 28.420 10.532 31.684 1.00 30.79 N +ATOM 1622 CA GLU A 348 28.502 11.766 30.891 1.00 30.79 C +ATOM 1623 C GLU A 348 29.880 12.441 30.958 1.00 30.79 C +ATOM 1624 O GLU A 348 30.381 12.938 29.951 1.00 30.79 O +ATOM 1625 CB GLU A 348 27.444 12.755 31.406 1.00 30.79 C +ATOM 1626 CG GLU A 348 26.004 12.340 31.061 1.00 30.79 C +ATOM 1627 CD GLU A 348 24.941 13.155 31.818 1.00 30.79 C +ATOM 1628 OE1 GLU A 348 23.748 13.029 31.460 1.00 30.79 O +ATOM 1629 OE2 GLU A 348 25.275 13.884 32.787 1.00 30.79 O +ATOM 1630 N THR A 349 30.489 12.491 32.146 1.00 30.82 N +ATOM 1631 CA THR A 349 31.688 13.312 32.408 1.00 30.82 C +ATOM 1632 C THR A 349 32.970 12.499 32.555 1.00 30.82 C +ATOM 1633 O THR A 349 34.068 13.046 32.470 1.00 30.82 O +ATOM 1634 CB THR A 349 31.509 14.162 33.674 1.00 30.82 C +ATOM 1635 OG1 THR A 349 31.460 13.321 34.803 1.00 30.82 O +ATOM 1636 CG2 THR A 349 30.247 15.026 33.651 1.00 30.82 C +ATOM 1637 N GLY A 350 32.857 11.196 32.825 1.00 30.76 N +ATOM 1638 CA GLY A 350 33.983 10.361 33.227 1.00 30.76 C +ATOM 1639 C GLY A 350 34.522 10.670 34.629 1.00 30.76 C +ATOM 1640 O GLY A 350 35.579 10.152 34.975 1.00 30.76 O +ATOM 1641 N ILE A 351 33.843 11.483 35.446 1.00 30.70 N +ATOM 1642 CA ILE A 351 34.255 11.810 36.818 1.00 30.70 C +ATOM 1643 C ILE A 351 33.331 11.093 37.804 1.00 30.70 C +ATOM 1644 O ILE A 351 32.122 11.308 37.814 1.00 30.70 O +ATOM 1645 CB ILE A 351 34.264 13.337 37.038 1.00 30.70 C +ATOM 1646 CG1 ILE A 351 35.242 14.036 36.065 1.00 30.70 C +ATOM 1647 CG2 ILE A 351 34.615 13.669 38.504 1.00 30.70 C +ATOM 1648 CD1 ILE A 351 35.120 15.565 36.073 1.00 30.70 C +ATOM 1649 N TYR A 352 33.893 10.212 38.627 1.00 30.64 N +ATOM 1650 CA TYR A 352 33.150 9.519 39.672 1.00 30.64 C +ATOM 1651 C TYR A 352 33.182 10.361 40.942 1.00 30.64 C +ATOM 1652 O TYR A 352 34.263 10.646 41.455 1.00 30.64 O +ATOM 1653 CB TYR A 352 33.754 8.134 39.914 1.00 30.64 C +ATOM 1654 CG TYR A 352 32.925 7.263 40.837 1.00 30.64 C +ATOM 1655 CD1 TYR A 352 33.078 7.354 42.235 1.00 30.64 C +ATOM 1656 CD2 TYR A 352 32.007 6.346 40.291 1.00 30.64 C +ATOM 1657 CE1 TYR A 352 32.370 6.480 43.084 1.00 30.64 C +ATOM 1658 CE2 TYR A 352 31.289 5.480 41.138 1.00 30.64 C +ATOM 1659 CZ TYR A 352 31.494 5.521 42.533 1.00 30.64 C +ATOM 1660 OH TYR A 352 30.847 4.630 43.330 1.00 30.64 O +ATOM 1661 N ILE A 353 32.014 10.742 41.458 1.00 30.58 N +ATOM 1662 CA ILE A 353 31.879 11.459 42.731 1.00 30.58 C +ATOM 1663 C ILE A 353 31.320 10.480 43.774 1.00 30.58 C +ATOM 1664 O ILE A 353 30.200 9.987 43.594 1.00 30.58 O +ATOM 1665 CB ILE A 353 31.016 12.731 42.572 1.00 30.58 C +ATOM 1666 CG1 ILE A 353 31.658 13.671 41.523 1.00 30.58 C +ATOM 1667 CG2 ILE A 353 30.870 13.445 43.930 1.00 30.58 C +ATOM 1668 CD1 ILE A 353 30.860 14.947 41.258 1.00 30.58 C +ATOM 1669 N PRO A 354 32.072 10.167 44.846 1.00 30.55 N +ATOM 1670 CA PRO A 354 31.573 9.338 45.932 1.00 30.55 C +ATOM 1671 C PRO A 354 30.384 9.973 46.651 1.00 30.55 C +ATOM 1672 O PRO A 354 30.327 11.196 46.832 1.00 30.55 O +ATOM 1673 CB PRO A 354 32.753 9.119 46.880 1.00 30.55 C +ATOM 1674 CG PRO A 354 33.955 9.269 45.961 1.00 30.55 C +ATOM 1675 CD PRO A 354 33.503 10.373 45.012 1.00 30.55 C +ATOM 1676 N VAL A 355 29.449 9.140 47.104 1.00 30.58 N +ATOM 1677 CA VAL A 355 28.276 9.586 47.860 1.00 30.58 C +ATOM 1678 C VAL A 355 28.221 8.960 49.250 1.00 30.58 C +ATOM 1679 O VAL A 355 28.465 7.768 49.441 1.00 30.58 O +ATOM 1680 CB VAL A 355 26.963 9.418 47.080 1.00 30.58 C +ATOM 1681 CG1 VAL A 355 27.010 10.144 45.733 1.00 30.58 C +ATOM 1682 CG2 VAL A 355 26.557 7.964 46.849 1.00 30.58 C +ATOM 1683 N CYS A 356 27.885 9.792 50.233 1.00 30.61 N +ATOM 1684 CA CYS A 356 27.804 9.434 51.639 1.00 30.61 C +ATOM 1685 C CYS A 356 26.349 9.360 52.099 1.00 30.61 C +ATOM 1686 O CYS A 356 25.632 10.352 51.985 1.00 30.61 O +ATOM 1687 CB CYS A 356 28.561 10.488 52.446 1.00 30.61 C +ATOM 1688 SG CYS A 356 28.567 10.016 54.195 1.00 30.61 S +ATOM 1689 N SER A 357 25.920 8.241 52.689 1.00 30.70 N +ATOM 1690 CA SER A 357 24.637 8.209 53.407 1.00 30.70 C +ATOM 1691 C SER A 357 24.820 8.713 54.841 1.00 30.70 C +ATOM 1692 O SER A 357 25.463 8.044 55.651 1.00 30.70 O +ATOM 1693 CB SER A 357 24.018 6.811 53.381 1.00 30.70 C +ATOM 1694 OG SER A 357 22.786 6.821 54.086 1.00 30.70 O +ATOM 1695 N ASP A 358 24.220 9.860 55.165 1.00 31.04 N +ATOM 1696 CA ASP A 358 24.245 10.503 56.485 1.00 31.04 C +ATOM 1697 C ASP A 358 22.859 10.442 57.151 1.00 31.04 C +ATOM 1698 O ASP A 358 21.897 11.080 56.719 1.00 31.04 O +ATOM 1699 CB ASP A 358 24.765 11.944 56.326 1.00 31.04 C +ATOM 1700 CG ASP A 358 24.822 12.759 57.626 1.00 31.04 C +ATOM 1701 OD1 ASP A 358 24.816 12.147 58.719 1.00 31.04 O +ATOM 1702 OD2 ASP A 358 24.894 14.015 57.516 1.00 31.04 O +ATOM 1703 N GLY A 359 22.752 9.652 58.223 1.00 33.95 N +ATOM 1704 CA GLY A 359 21.502 9.429 58.959 1.00 33.95 C +ATOM 1705 C GLY A 359 20.728 8.169 58.539 1.00 33.95 C +ATOM 1706 O GLY A 359 21.114 7.434 57.637 1.00 33.95 O +ATOM 1707 N GLY A 360 19.647 7.855 59.261 1.00 34.67 N +ATOM 1708 CA GLY A 360 18.765 6.701 59.011 1.00 34.67 C +ATOM 1709 C GLY A 360 19.334 5.315 59.362 1.00 34.67 C +ATOM 1710 O GLY A 360 18.606 4.322 59.355 1.00 34.67 O +ATOM 1711 N ILE A 361 20.622 5.216 59.691 1.00 32.62 N +ATOM 1712 CA ILE A 361 21.304 3.953 60.005 1.00 32.62 C +ATOM 1713 C ILE A 361 21.196 3.680 61.510 1.00 32.62 C +ATOM 1714 O ILE A 361 21.797 4.374 62.333 1.00 32.62 O +ATOM 1715 CB ILE A 361 22.750 3.993 59.468 1.00 32.62 C +ATOM 1716 CG1 ILE A 361 22.715 4.084 57.923 1.00 32.62 C +ATOM 1717 CG2 ILE A 361 23.560 2.760 59.917 1.00 32.62 C +ATOM 1718 CD1 ILE A 361 24.050 4.470 57.288 1.00 32.62 C +ATOM 1719 N VAL A 362 20.394 2.674 61.869 1.00 34.48 N +ATOM 1720 CA VAL A 362 20.066 2.318 63.261 1.00 34.48 C +ATOM 1721 C VAL A 362 20.716 0.996 63.656 1.00 34.48 C +ATOM 1722 O VAL A 362 21.278 0.899 64.750 1.00 34.48 O +ATOM 1723 CB VAL A 362 18.537 2.257 63.468 1.00 34.48 C +ATOM 1724 CG1 VAL A 362 18.171 1.906 64.918 1.00 34.48 C +ATOM 1725 CG2 VAL A 362 17.881 3.603 63.132 1.00 34.48 C +ATOM 1726 N TYR A 363 20.660 0.014 62.757 1.00 34.19 N +ATOM 1727 CA TYR A 363 21.188 -1.336 62.935 1.00 34.19 C +ATOM 1728 C TYR A 363 22.369 -1.592 61.995 1.00 34.19 C +ATOM 1729 O TYR A 363 22.455 -0.988 60.927 1.00 34.19 O +ATOM 1730 CB TYR A 363 20.063 -2.341 62.665 1.00 34.19 C +ATOM 1731 CG TYR A 363 18.831 -2.150 63.529 1.00 34.19 C +ATOM 1732 CD1 TYR A 363 18.877 -2.482 64.896 1.00 34.19 C +ATOM 1733 CD2 TYR A 363 17.638 -1.651 62.970 1.00 34.19 C +ATOM 1734 CE1 TYR A 363 17.737 -2.310 65.702 1.00 34.19 C +ATOM 1735 CE2 TYR A 363 16.505 -1.450 63.779 1.00 34.19 C +ATOM 1736 CZ TYR A 363 16.553 -1.777 65.149 1.00 34.19 C +ATOM 1737 OH TYR A 363 15.465 -1.580 65.937 1.00 34.19 O +ATOM 1738 N ASP A 364 23.246 -2.536 62.346 1.00 33.06 N +ATOM 1739 CA ASP A 364 24.420 -2.872 61.522 1.00 33.06 C +ATOM 1740 C ASP A 364 24.033 -3.297 60.096 1.00 33.06 C +ATOM 1741 O ASP A 364 24.720 -2.953 59.141 1.00 33.06 O +ATOM 1742 CB ASP A 364 25.228 -3.998 62.181 1.00 33.06 C +ATOM 1743 CG ASP A 364 25.867 -3.604 63.511 1.00 33.06 C +ATOM 1744 OD1 ASP A 364 26.278 -2.436 63.672 1.00 33.06 O +ATOM 1745 OD2 ASP A 364 25.952 -4.496 64.384 1.00 33.06 O +ATOM 1746 N TYR A 365 22.897 -3.978 59.911 1.00 32.90 N +ATOM 1747 CA TYR A 365 22.452 -4.381 58.574 1.00 32.90 C +ATOM 1748 C TYR A 365 22.016 -3.190 57.696 1.00 32.90 C +ATOM 1749 O TYR A 365 22.091 -3.288 56.472 1.00 32.90 O +ATOM 1750 CB TYR A 365 21.369 -5.463 58.676 1.00 32.90 C +ATOM 1751 CG TYR A 365 20.000 -4.959 59.090 1.00 32.90 C +ATOM 1752 CD1 TYR A 365 19.610 -4.965 60.444 1.00 32.90 C +ATOM 1753 CD2 TYR A 365 19.102 -4.511 58.104 1.00 32.90 C +ATOM 1754 CE1 TYR A 365 18.334 -4.491 60.812 1.00 32.90 C +ATOM 1755 CE2 TYR A 365 17.823 -4.050 58.467 1.00 32.90 C +ATOM 1756 CZ TYR A 365 17.442 -4.019 59.822 1.00 32.90 C +ATOM 1757 OH TYR A 365 16.218 -3.545 60.161 1.00 32.90 O +ATOM 1758 N HIS A 366 21.647 -2.038 58.284 1.00 31.17 N +ATOM 1759 CA HIS A 366 21.418 -0.805 57.515 1.00 31.17 C +ATOM 1760 C HIS A 366 22.708 -0.309 56.865 1.00 31.17 C +ATOM 1761 O HIS A 366 22.645 0.297 55.803 1.00 31.17 O +ATOM 1762 CB HIS A 366 20.875 0.345 58.377 1.00 31.17 C +ATOM 1763 CG HIS A 366 19.505 0.174 58.950 1.00 31.17 C +ATOM 1764 ND1 HIS A 366 18.780 1.145 59.610 1.00 31.17 N +ATOM 1765 CD2 HIS A 366 18.731 -0.944 58.871 1.00 31.17 C +ATOM 1766 CE1 HIS A 366 17.609 0.596 59.957 1.00 31.17 C +ATOM 1767 NE2 HIS A 366 17.543 -0.665 59.523 1.00 31.17 N +ATOM 1768 N MET A 367 23.871 -0.573 57.480 1.00 30.76 N +ATOM 1769 CA MET A 367 25.162 -0.243 56.879 1.00 30.76 C +ATOM 1770 C MET A 367 25.313 -0.986 55.551 1.00 30.76 C +ATOM 1771 O MET A 367 25.560 -0.364 54.523 1.00 30.76 O +ATOM 1772 CB MET A 367 26.331 -0.594 57.816 1.00 30.76 C +ATOM 1773 CG MET A 367 26.207 0.020 59.215 1.00 30.76 C +ATOM 1774 SD MET A 367 27.666 -0.269 60.246 1.00 30.76 S +ATOM 1775 CE MET A 367 27.120 0.504 61.789 1.00 30.76 C +ATOM 1776 N THR A 368 25.072 -2.301 55.548 1.00 30.85 N +ATOM 1777 CA THR A 368 25.141 -3.106 54.322 1.00 30.85 C +ATOM 1778 C THR A 368 24.082 -2.693 53.302 1.00 30.85 C +ATOM 1779 O THR A 368 24.395 -2.627 52.118 1.00 30.85 O +ATOM 1780 CB THR A 368 24.998 -4.607 54.604 1.00 30.85 C +ATOM 1781 OG1 THR A 368 25.758 -4.990 55.727 1.00 30.85 O +ATOM 1782 CG2 THR A 368 25.520 -5.422 53.421 1.00 30.85 C +ATOM 1783 N LEU A 369 22.852 -2.380 53.734 1.00 30.70 N +ATOM 1784 CA LEU A 369 21.798 -1.893 52.837 1.00 30.70 C +ATOM 1785 C LEU A 369 22.181 -0.572 52.165 1.00 30.70 C +ATOM 1786 O LEU A 369 22.074 -0.470 50.949 1.00 30.70 O +ATOM 1787 CB LEU A 369 20.472 -1.726 53.600 1.00 30.70 C +ATOM 1788 CG LEU A 369 19.742 -3.044 53.900 1.00 30.70 C +ATOM 1789 CD1 LEU A 369 18.621 -2.779 54.899 1.00 30.70 C +ATOM 1790 CD2 LEU A 369 19.110 -3.637 52.638 1.00 30.70 C +ATOM 1791 N ALA A 370 22.672 0.412 52.922 1.00 30.67 N +ATOM 1792 CA ALA A 370 23.096 1.696 52.368 1.00 30.67 C +ATOM 1793 C ALA A 370 24.219 1.524 51.330 1.00 30.67 C +ATOM 1794 O ALA A 370 24.150 2.107 50.248 1.00 30.67 O +ATOM 1795 CB ALA A 370 23.526 2.605 53.525 1.00 30.67 C +ATOM 1796 N LEU A 371 25.214 0.677 51.629 1.00 30.58 N +ATOM 1797 CA LEU A 371 26.306 0.358 50.702 1.00 30.58 C +ATOM 1798 C LEU A 371 25.803 -0.377 49.453 1.00 30.58 C +ATOM 1799 O LEU A 371 26.199 -0.040 48.342 1.00 30.58 O +ATOM 1800 CB LEU A 371 27.374 -0.483 51.423 1.00 30.58 C +ATOM 1801 CG LEU A 371 28.096 0.241 52.572 1.00 30.58 C +ATOM 1802 CD1 LEU A 371 28.978 -0.757 53.327 1.00 30.58 C +ATOM 1803 CD2 LEU A 371 28.959 1.403 52.094 1.00 30.58 C +ATOM 1804 N ALA A 372 24.904 -1.352 49.619 1.00 30.64 N +ATOM 1805 CA ALA A 372 24.296 -2.089 48.514 1.00 30.64 C +ATOM 1806 C ALA A 372 23.427 -1.196 47.615 1.00 30.64 C +ATOM 1807 O ALA A 372 23.397 -1.400 46.408 1.00 30.64 O +ATOM 1808 CB ALA A 372 23.481 -3.249 49.095 1.00 30.64 C +ATOM 1809 N MET A 373 22.770 -0.181 48.186 1.00 30.70 N +ATOM 1810 CA MET A 373 21.998 0.831 47.453 1.00 30.70 C +ATOM 1811 C MET A 373 22.876 1.857 46.717 1.00 30.70 C +ATOM 1812 O MET A 373 22.340 2.713 46.019 1.00 30.70 O +ATOM 1813 CB MET A 373 21.022 1.533 48.411 1.00 30.70 C +ATOM 1814 CG MET A 373 19.882 0.611 48.850 1.00 30.70 C +ATOM 1815 SD MET A 373 18.822 1.334 50.131 1.00 30.70 S +ATOM 1816 CE MET A 373 17.649 -0.029 50.351 1.00 30.70 C +ATOM 1817 N GLY A 374 24.206 1.773 46.843 1.00 30.73 N +ATOM 1818 CA GLY A 374 25.150 2.563 46.051 1.00 30.73 C +ATOM 1819 C GLY A 374 25.915 3.642 46.810 1.00 30.73 C +ATOM 1820 O GLY A 374 26.744 4.312 46.192 1.00 30.73 O +ATOM 1821 N ALA A 375 25.684 3.801 48.119 1.00 30.61 N +ATOM 1822 CA ALA A 375 26.552 4.643 48.934 1.00 30.61 C +ATOM 1823 C ALA A 375 27.982 4.084 48.932 1.00 30.61 C +ATOM 1824 O ALA A 375 28.188 2.874 49.040 1.00 30.61 O +ATOM 1825 CB ALA A 375 25.986 4.805 50.351 1.00 30.61 C +ATOM 1826 N ASP A 376 28.969 4.967 48.819 1.00 30.58 N +ATOM 1827 CA ASP A 376 30.384 4.595 48.877 1.00 30.58 C +ATOM 1828 C ASP A 376 30.836 4.452 50.335 1.00 30.58 C +ATOM 1829 O ASP A 376 31.575 3.531 50.682 1.00 30.58 O +ATOM 1830 CB ASP A 376 31.208 5.646 48.122 1.00 30.58 C +ATOM 1831 CG ASP A 376 30.860 5.661 46.634 1.00 30.58 C +ATOM 1832 OD1 ASP A 376 31.406 4.838 45.871 1.00 30.58 O +ATOM 1833 OD2 ASP A 376 30.035 6.507 46.221 1.00 30.58 O +ATOM 1834 N PHE A 377 30.348 5.334 51.207 1.00 30.58 N +ATOM 1835 CA PHE A 377 30.593 5.303 52.646 1.00 30.58 C +ATOM 1836 C PHE A 377 29.393 5.847 53.428 1.00 30.58 C +ATOM 1837 O PHE A 377 28.457 6.420 52.866 1.00 30.58 O +ATOM 1838 CB PHE A 377 31.896 6.050 52.968 1.00 30.58 C +ATOM 1839 CG PHE A 377 31.930 7.505 52.549 1.00 30.58 C +ATOM 1840 CD1 PHE A 377 32.350 7.856 51.251 1.00 30.58 C +ATOM 1841 CD2 PHE A 377 31.593 8.512 53.471 1.00 30.58 C +ATOM 1842 CE1 PHE A 377 32.420 9.207 50.875 1.00 30.58 C +ATOM 1843 CE2 PHE A 377 31.710 9.865 53.107 1.00 30.58 C +ATOM 1844 CZ PHE A 377 32.108 10.211 51.805 1.00 30.58 C +ATOM 1845 N ILE A 378 29.400 5.657 54.745 1.00 30.61 N +ATOM 1846 CA ILE A 378 28.291 6.035 55.624 1.00 30.61 C +ATOM 1847 C ILE A 378 28.774 6.906 56.780 1.00 30.61 C +ATOM 1848 O ILE A 378 29.781 6.608 57.425 1.00 30.61 O +ATOM 1849 CB ILE A 378 27.503 4.794 56.096 1.00 30.61 C +ATOM 1850 CG1 ILE A 378 28.371 3.807 56.908 1.00 30.61 C +ATOM 1851 CG2 ILE A 378 26.860 4.097 54.881 1.00 30.61 C +ATOM 1852 CD1 ILE A 378 27.645 2.537 57.341 1.00 30.61 C +ATOM 1853 N MET A 379 28.036 7.985 57.042 1.00 30.70 N +ATOM 1854 CA MET A 379 28.287 8.892 58.155 1.00 30.70 C +ATOM 1855 C MET A 379 27.343 8.565 59.309 1.00 30.70 C +ATOM 1856 O MET A 379 26.119 8.549 59.153 1.00 30.70 O +ATOM 1857 CB MET A 379 28.206 10.350 57.697 1.00 30.70 C +ATOM 1858 CG MET A 379 28.534 11.321 58.834 1.00 30.70 C +ATOM 1859 SD MET A 379 28.650 13.069 58.353 1.00 30.70 S +ATOM 1860 CE MET A 379 30.191 13.066 57.397 1.00 30.70 C +ATOM 1861 N LEU A 380 27.912 8.277 60.479 1.00 31.01 N +ATOM 1862 CA LEU A 380 27.166 7.787 61.632 1.00 31.01 C +ATOM 1863 C LEU A 380 27.307 8.743 62.824 1.00 31.01 C +ATOM 1864 O LEU A 380 28.406 9.052 63.272 1.00 31.01 O +ATOM 1865 CB LEU A 380 27.604 6.350 61.967 1.00 31.01 C +ATOM 1866 CG LEU A 380 27.615 5.337 60.804 1.00 31.01 C +ATOM 1867 CD1 LEU A 380 28.190 4.003 61.282 1.00 31.01 C +ATOM 1868 CD2 LEU A 380 26.203 5.085 60.281 1.00 31.01 C +ATOM 1869 N GLY A 381 26.181 9.209 63.366 1.00 31.93 N +ATOM 1870 CA GLY A 381 26.154 10.006 64.600 1.00 31.93 C +ATOM 1871 C GLY A 381 25.950 9.121 65.825 1.00 31.93 C +ATOM 1872 O GLY A 381 26.895 8.785 66.540 1.00 31.93 O +ATOM 1873 N ARG A 382 24.703 8.663 66.012 1.00 32.04 N +ATOM 1874 CA ARG A 382 24.269 7.825 67.145 1.00 32.04 C +ATOM 1875 C ARG A 382 25.179 6.625 67.385 1.00 32.04 C +ATOM 1876 O ARG A 382 25.438 6.308 68.534 1.00 32.04 O +ATOM 1877 CB ARG A 382 22.813 7.370 66.921 1.00 32.04 C +ATOM 1878 CG ARG A 382 22.302 6.420 68.021 1.00 32.04 C +ATOM 1879 CD ARG A 382 20.842 6.004 67.807 1.00 32.04 C +ATOM 1880 NE ARG A 382 20.438 4.940 68.753 1.00 32.04 N +ATOM 1881 CZ ARG A 382 20.565 3.629 68.593 1.00 32.04 C +ATOM 1882 NH1 ARG A 382 21.098 3.105 67.519 1.00 32.04 N +ATOM 1883 NH2 ARG A 382 20.150 2.810 69.520 1.00 32.04 N +ATOM 1884 N TYR A 383 25.639 5.948 66.333 1.00 31.33 N +ATOM 1885 CA TYR A 383 26.493 4.765 66.467 1.00 31.33 C +ATOM 1886 C TYR A 383 27.778 5.081 67.248 1.00 31.33 C +ATOM 1887 O TYR A 383 28.019 4.447 68.272 1.00 31.33 O +ATOM 1888 CB TYR A 383 26.808 4.207 65.078 1.00 31.33 C +ATOM 1889 CG TYR A 383 27.687 2.975 65.080 1.00 31.33 C +ATOM 1890 CD1 TYR A 383 29.073 3.116 64.882 1.00 31.33 C +ATOM 1891 CD2 TYR A 383 27.124 1.693 65.248 1.00 31.33 C +ATOM 1892 CE1 TYR A 383 29.897 1.980 64.862 1.00 31.33 C +ATOM 1893 CE2 TYR A 383 27.954 0.552 65.246 1.00 31.33 C +ATOM 1894 CZ TYR A 383 29.346 0.699 65.055 1.00 31.33 C +ATOM 1895 OH TYR A 383 30.184 -0.363 65.072 1.00 31.33 O +ATOM 1896 N PHE A 384 28.522 6.103 66.808 1.00 31.01 N +ATOM 1897 CA PHE A 384 29.793 6.524 67.402 1.00 31.01 C +ATOM 1898 C PHE A 384 29.620 7.294 68.716 1.00 31.01 C +ATOM 1899 O PHE A 384 30.468 7.180 69.597 1.00 31.01 O +ATOM 1900 CB PHE A 384 30.567 7.360 66.371 1.00 31.01 C +ATOM 1901 CG PHE A 384 31.161 6.563 65.225 1.00 31.01 C +ATOM 1902 CD1 PHE A 384 32.104 5.554 65.475 1.00 31.01 C +ATOM 1903 CD2 PHE A 384 30.803 6.840 63.899 1.00 31.01 C +ATOM 1904 CE1 PHE A 384 32.641 4.803 64.416 1.00 31.01 C +ATOM 1905 CE2 PHE A 384 31.342 6.105 62.832 1.00 31.01 C +ATOM 1906 CZ PHE A 384 32.261 5.079 63.095 1.00 31.01 C +ATOM 1907 N ALA A 385 28.505 8.011 68.909 1.00 31.37 N +ATOM 1908 CA ALA A 385 28.201 8.706 70.166 1.00 31.37 C +ATOM 1909 C ALA A 385 28.215 7.774 71.396 1.00 31.37 C +ATOM 1910 O ALA A 385 28.502 8.223 72.503 1.00 31.37 O +ATOM 1911 CB ALA A 385 26.826 9.366 70.040 1.00 31.37 C +ATOM 1912 N ARG A 386 27.938 6.478 71.201 1.00 31.54 N +ATOM 1913 CA ARG A 386 27.888 5.454 72.257 1.00 31.54 C +ATOM 1914 C ARG A 386 29.256 5.013 72.777 1.00 31.54 C +ATOM 1915 O ARG A 386 29.307 4.350 73.813 1.00 31.54 O +ATOM 1916 CB ARG A 386 27.151 4.220 71.725 1.00 31.54 C +ATOM 1917 CG ARG A 386 25.688 4.520 71.381 1.00 31.54 C +ATOM 1918 CD ARG A 386 24.941 3.272 70.909 1.00 31.54 C +ATOM 1919 NE ARG A 386 25.551 2.695 69.695 1.00 31.54 N +ATOM 1920 CZ ARG A 386 25.247 1.530 69.153 1.00 31.54 C +ATOM 1921 NH1 ARG A 386 24.277 0.785 69.614 1.00 31.54 N +ATOM 1922 NH2 ARG A 386 25.928 1.083 68.137 1.00 31.54 N +ATOM 1923 N PHE A 387 30.340 5.314 72.066 1.00 31.33 N +ATOM 1924 CA PHE A 387 31.657 4.760 72.376 1.00 31.33 C +ATOM 1925 C PHE A 387 32.427 5.585 73.398 1.00 31.33 C +ATOM 1926 O PHE A 387 32.156 6.770 73.594 1.00 31.33 O +ATOM 1927 CB PHE A 387 32.457 4.544 71.088 1.00 31.33 C +ATOM 1928 CG PHE A 387 31.801 3.642 70.060 1.00 31.33 C +ATOM 1929 CD1 PHE A 387 30.869 2.652 70.435 1.00 31.33 C +ATOM 1930 CD2 PHE A 387 32.170 3.766 68.711 1.00 31.33 C +ATOM 1931 CE1 PHE A 387 30.275 1.833 69.466 1.00 31.33 C +ATOM 1932 CE2 PHE A 387 31.607 2.918 67.744 1.00 31.33 C +ATOM 1933 CZ PHE A 387 30.652 1.963 68.123 1.00 31.33 C +ATOM 1934 N GLU A 388 33.406 4.959 74.048 1.00 31.75 N +ATOM 1935 CA GLU A 388 34.289 5.598 75.028 1.00 31.75 C +ATOM 1936 C GLU A 388 34.967 6.846 74.455 1.00 31.75 C +ATOM 1937 O GLU A 388 35.065 7.865 75.134 1.00 31.75 O +ATOM 1938 CB GLU A 388 35.342 4.573 75.492 1.00 31.75 C +ATOM 1939 CG GLU A 388 36.234 5.081 76.643 1.00 31.75 C +ATOM 1940 CD GLU A 388 35.456 5.375 77.936 1.00 31.75 C +ATOM 1941 OE1 GLU A 388 35.982 6.033 78.865 1.00 31.75 O +ATOM 1942 OE2 GLU A 388 34.307 4.910 78.090 1.00 31.75 O +ATOM 1943 N GLU A 389 35.370 6.785 73.187 1.00 31.20 N +ATOM 1944 CA GLU A 389 36.126 7.816 72.484 1.00 31.20 C +ATOM 1945 C GLU A 389 35.283 9.025 72.056 1.00 31.20 C +ATOM 1946 O GLU A 389 35.855 10.051 71.677 1.00 31.20 O +ATOM 1947 CB GLU A 389 36.833 7.188 71.270 1.00 31.20 C +ATOM 1948 CG GLU A 389 37.839 6.083 71.652 1.00 31.20 C +ATOM 1949 CD GLU A 389 37.253 4.667 71.823 1.00 31.20 C +ATOM 1950 OE1 GLU A 389 38.046 3.714 71.977 1.00 31.20 O +ATOM 1951 OE2 GLU A 389 36.011 4.499 71.828 1.00 31.20 O +ATOM 1952 N SER A 390 33.946 8.944 72.123 1.00 31.37 N +ATOM 1953 CA SER A 390 33.108 10.116 71.858 1.00 31.37 C +ATOM 1954 C SER A 390 33.258 11.158 72.986 1.00 31.37 C +ATOM 1955 O SER A 390 33.468 10.784 74.145 1.00 31.37 O +ATOM 1956 CB SER A 390 31.644 9.754 71.582 1.00 31.37 C +ATOM 1957 OG SER A 390 30.957 9.408 72.762 1.00 31.37 O +ATOM 1958 N PRO A 391 33.143 12.469 72.702 1.00 32.78 N +ATOM 1959 CA PRO A 391 33.513 13.525 73.655 1.00 32.78 C +ATOM 1960 C PRO A 391 32.655 13.605 74.922 1.00 32.78 C +ATOM 1961 O PRO A 391 33.035 14.268 75.888 1.00 32.78 O +ATOM 1962 CB PRO A 391 33.396 14.835 72.866 1.00 32.78 C +ATOM 1963 CG PRO A 391 33.611 14.400 71.423 1.00 32.78 C +ATOM 1964 CD PRO A 391 32.922 13.047 71.387 1.00 32.78 C +ATOM 1965 N THR A 392 31.472 12.991 74.917 1.00 33.06 N +ATOM 1966 CA THR A 392 30.457 13.228 75.946 1.00 33.06 C +ATOM 1967 C THR A 392 30.684 12.404 77.207 1.00 33.06 C +ATOM 1968 O THR A 392 31.391 11.391 77.221 1.00 33.06 O +ATOM 1969 CB THR A 392 29.038 13.059 75.395 1.00 33.06 C +ATOM 1970 OG1 THR A 392 28.785 11.745 74.980 1.00 33.06 O +ATOM 1971 CG2 THR A 392 28.834 13.961 74.183 1.00 33.06 C +ATOM 1972 N ARG A 393 30.090 12.861 78.312 1.00 34.88 N +ATOM 1973 CA ARG A 393 30.225 12.203 79.612 1.00 34.88 C +ATOM 1974 C ARG A 393 29.393 10.926 79.656 1.00 34.88 C +ATOM 1975 O ARG A 393 28.302 10.861 79.093 1.00 34.88 O +ATOM 1976 CB ARG A 393 29.828 13.166 80.739 1.00 34.88 C +ATOM 1977 CG ARG A 393 30.793 14.357 80.823 1.00 34.88 C +ATOM 1978 CD ARG A 393 30.361 15.312 81.935 1.00 34.88 C +ATOM 1979 NE ARG A 393 31.281 16.459 82.038 1.00 34.88 N +ATOM 1980 CZ ARG A 393 31.190 17.452 82.905 1.00 34.88 C +ATOM 1981 NH1 ARG A 393 30.234 17.504 83.791 1.00 34.88 N +ATOM 1982 NH2 ARG A 393 32.066 18.417 82.898 1.00 34.88 N +ATOM 1983 N LYS A 394 29.905 9.935 80.387 1.00 33.24 N +ATOM 1984 CA LYS A 394 29.114 8.790 80.843 1.00 33.24 C +ATOM 1985 C LYS A 394 28.243 9.237 82.014 1.00 33.24 C +ATOM 1986 O LYS A 394 28.752 9.842 82.956 1.00 33.24 O +ATOM 1987 CB LYS A 394 30.025 7.611 81.220 1.00 33.24 C +ATOM 1988 CG LYS A 394 30.635 6.960 79.969 1.00 33.24 C +ATOM 1989 CD LYS A 394 31.615 5.822 80.286 1.00 33.24 C +ATOM 1990 CE LYS A 394 32.973 6.352 80.763 1.00 33.24 C +ATOM 1991 NZ LYS A 394 34.014 5.303 80.694 1.00 33.24 N +ATOM 1992 N VAL A 395 26.954 8.939 81.950 1.00 33.11 N +ATOM 1993 CA VAL A 395 25.967 9.240 82.992 1.00 33.11 C +ATOM 1994 C VAL A 395 25.162 7.988 83.300 1.00 33.11 C +ATOM 1995 O VAL A 395 24.866 7.204 82.403 1.00 33.11 O +ATOM 1996 CB VAL A 395 25.047 10.417 82.606 1.00 33.11 C +ATOM 1997 CG1 VAL A 395 25.848 11.712 82.431 1.00 33.11 C +ATOM 1998 CG2 VAL A 395 24.241 10.173 81.325 1.00 33.11 C +ATOM 1999 N THR A 396 24.800 7.784 84.562 1.00 32.82 N +ATOM 2000 CA THR A 396 23.946 6.658 84.954 1.00 32.82 C +ATOM 2001 C THR A 396 22.504 7.130 85.041 1.00 32.82 C +ATOM 2002 O THR A 396 22.194 8.014 85.837 1.00 32.82 O +ATOM 2003 CB THR A 396 24.410 6.026 86.270 1.00 32.82 C +ATOM 2004 OG1 THR A 396 25.751 5.611 86.141 1.00 32.82 O +ATOM 2005 CG2 THR A 396 23.595 4.782 86.626 1.00 32.82 C +ATOM 2006 N ILE A 397 21.618 6.534 84.243 1.00 36.07 N +ATOM 2007 CA ILE A 397 20.182 6.833 84.223 1.00 36.07 C +ATOM 2008 C ILE A 397 19.437 5.525 84.438 1.00 36.07 C +ATOM 2009 O ILE A 397 19.608 4.581 83.670 1.00 36.07 O +ATOM 2010 CB ILE A 397 19.765 7.517 82.902 1.00 36.07 C +ATOM 2011 CG1 ILE A 397 20.564 8.823 82.716 1.00 36.07 C +ATOM 2012 CG2 ILE A 397 18.246 7.790 82.897 1.00 36.07 C +ATOM 2013 CD1 ILE A 397 20.273 9.546 81.404 1.00 36.07 C +ATOM 2014 N ASN A 398 18.624 5.454 85.494 1.00 36.02 N +ATOM 2015 CA ASN A 398 17.832 4.266 85.836 1.00 36.02 C +ATOM 2016 C ASN A 398 18.667 2.968 85.865 1.00 36.02 C +ATOM 2017 O ASN A 398 18.246 1.930 85.365 1.00 36.02 O +ATOM 2018 CB ASN A 398 16.598 4.205 84.916 1.00 36.02 C +ATOM 2019 CG ASN A 398 15.725 5.443 85.023 1.00 36.02 C +ATOM 2020 OD1 ASN A 398 15.736 6.164 86.005 1.00 36.02 O +ATOM 2021 ND2 ASN A 398 14.951 5.739 84.007 1.00 36.02 N +ATOM 2022 N GLY A 399 19.889 3.050 86.404 1.00 36.13 N +ATOM 2023 CA GLY A 399 20.826 1.923 86.498 1.00 36.13 C +ATOM 2024 C GLY A 399 21.584 1.582 85.210 1.00 36.13 C +ATOM 2025 O GLY A 399 22.533 0.806 85.266 1.00 36.13 O +ATOM 2026 N SER A 400 21.231 2.182 84.070 1.00 35.78 N +ATOM 2027 CA SER A 400 21.938 1.997 82.799 1.00 35.78 C +ATOM 2028 C SER A 400 22.997 3.077 82.600 1.00 35.78 C +ATOM 2029 O SER A 400 22.760 4.253 82.881 1.00 35.78 O +ATOM 2030 CB SER A 400 20.951 2.008 81.632 1.00 35.78 C +ATOM 2031 OG SER A 400 20.041 0.940 81.777 1.00 35.78 O +ATOM 2032 N VAL A 401 24.169 2.690 82.091 1.00 32.90 N +ATOM 2033 CA VAL A 401 25.214 3.647 81.706 1.00 32.90 C +ATOM 2034 C VAL A 401 24.901 4.180 80.311 1.00 32.90 C +ATOM 2035 O VAL A 401 24.823 3.431 79.338 1.00 32.90 O +ATOM 2036 CB VAL A 401 26.629 3.044 81.778 1.00 32.90 C +ATOM 2037 CG1 VAL A 401 27.695 4.115 81.504 1.00 32.90 C +ATOM 2038 CG2 VAL A 401 26.914 2.459 83.168 1.00 32.90 C +ATOM 2039 N MET A 402 24.740 5.492 80.219 1.00 32.38 N +ATOM 2040 CA MET A 402 24.385 6.225 79.012 1.00 32.38 C +ATOM 2041 C MET A 402 25.470 7.259 78.692 1.00 32.38 C +ATOM 2042 O MET A 402 26.295 7.602 79.540 1.00 32.38 O +ATOM 2043 CB MET A 402 23.021 6.910 79.207 1.00 32.38 C +ATOM 2044 CG MET A 402 21.888 5.980 79.667 1.00 32.38 C +ATOM 2045 SD MET A 402 21.391 4.691 78.496 1.00 32.38 S +ATOM 2046 CE MET A 402 20.432 5.681 77.320 1.00 32.38 C +ATOM 2047 N LYS A 403 25.452 7.795 77.474 1.00 31.93 N +ATOM 2048 CA LYS A 403 26.211 8.976 77.059 1.00 31.93 C +ATOM 2049 C LYS A 403 25.256 10.044 76.545 1.00 31.93 C +ATOM 2050 O LYS A 403 24.234 9.727 75.939 1.00 31.93 O +ATOM 2051 CB LYS A 403 27.277 8.608 76.013 1.00 31.93 C +ATOM 2052 CG LYS A 403 28.362 7.689 76.599 1.00 31.93 C +ATOM 2053 CD LYS A 403 29.660 7.640 75.775 1.00 31.93 C +ATOM 2054 CE LYS A 403 30.397 8.984 75.813 1.00 31.93 C +ATOM 2055 NZ LYS A 403 31.781 8.887 75.296 1.00 31.93 N +ATOM 2056 N GLU A 404 25.583 11.304 76.800 1.00 32.41 N +ATOM 2057 CA GLU A 404 24.832 12.435 76.244 1.00 32.41 C +ATOM 2058 C GLU A 404 25.003 12.488 74.719 1.00 32.41 C +ATOM 2059 O GLU A 404 26.088 12.205 74.202 1.00 32.41 O +ATOM 2060 CB GLU A 404 25.294 13.759 76.861 1.00 32.41 C +ATOM 2061 CG GLU A 404 25.219 13.804 78.391 1.00 32.41 C +ATOM 2062 CD GLU A 404 25.735 15.159 78.873 1.00 32.41 C +ATOM 2063 OE1 GLU A 404 24.910 16.028 79.224 1.00 32.41 O +ATOM 2064 OE2 GLU A 404 26.958 15.408 78.759 1.00 32.41 O +ATOM 2065 N TYR A 405 23.942 12.852 74.002 1.00 32.94 N +ATOM 2066 CA TYR A 405 23.920 12.927 72.545 1.00 32.94 C +ATOM 2067 C TYR A 405 22.850 13.912 72.069 1.00 32.94 C +ATOM 2068 O TYR A 405 21.657 13.677 72.237 1.00 32.94 O +ATOM 2069 CB TYR A 405 23.653 11.527 71.982 1.00 32.94 C +ATOM 2070 CG TYR A 405 23.538 11.466 70.471 1.00 32.94 C +ATOM 2071 CD1 TYR A 405 22.362 10.981 69.872 1.00 32.94 C +ATOM 2072 CD2 TYR A 405 24.610 11.888 69.663 1.00 32.94 C +ATOM 2073 CE1 TYR A 405 22.263 10.912 68.472 1.00 32.94 C +ATOM 2074 CE2 TYR A 405 24.507 11.836 68.260 1.00 32.94 C +ATOM 2075 CZ TYR A 405 23.326 11.359 67.661 1.00 32.94 C +ATOM 2076 OH TYR A 405 23.207 11.309 66.308 1.00 32.94 O +ATOM 2077 N TRP A 406 23.281 14.998 71.431 1.00 32.98 N +ATOM 2078 CA TRP A 406 22.402 16.050 70.920 1.00 32.98 C +ATOM 2079 C TRP A 406 22.677 16.355 69.455 1.00 32.98 C +ATOM 2080 O TRP A 406 23.780 16.115 68.953 1.00 32.98 O +ATOM 2081 CB TRP A 406 22.567 17.320 71.758 1.00 32.98 C +ATOM 2082 CG TRP A 406 23.932 17.932 71.714 1.00 32.98 C +ATOM 2083 CD1 TRP A 406 24.348 18.890 70.859 1.00 32.98 C +ATOM 2084 CD2 TRP A 406 25.082 17.616 72.548 1.00 32.98 C +ATOM 2085 NE1 TRP A 406 25.668 19.210 71.116 1.00 32.98 N +ATOM 2086 CE2 TRP A 406 26.168 18.451 72.151 1.00 32.98 C +ATOM 2087 CE3 TRP A 406 25.308 16.714 73.609 1.00 32.98 C +ATOM 2088 CZ2 TRP A 406 27.413 18.405 72.789 1.00 32.98 C +ATOM 2089 CZ3 TRP A 406 26.565 16.644 74.239 1.00 32.98 C +ATOM 2090 CH2 TRP A 406 27.616 17.484 73.830 1.00 32.98 C +ATOM 2091 N GLY A 407 21.670 16.912 68.784 1.00 36.42 N +ATOM 2092 CA GLY A 407 21.782 17.395 67.412 1.00 36.42 C +ATOM 2093 C GLY A 407 22.378 18.774 67.279 1.00 36.42 C +ATOM 2094 O GLY A 407 22.160 19.637 68.123 1.00 36.42 O +ATOM 2095 N GLU A 408 23.043 19.011 66.153 1.00 35.78 N +ATOM 2096 CA GLU A 408 23.558 20.338 65.802 1.00 35.78 C +ATOM 2097 C GLU A 408 22.434 21.377 65.637 1.00 35.78 C +ATOM 2098 O GLU A 408 22.671 22.569 65.791 1.00 35.78 O +ATOM 2099 CB GLU A 408 24.427 20.239 64.543 1.00 35.78 C +ATOM 2100 CG GLU A 408 25.714 19.419 64.759 1.00 35.78 C +ATOM 2101 CD GLU A 408 26.653 20.008 65.824 1.00 35.78 C +ATOM 2102 OE1 GLU A 408 27.061 19.259 66.744 1.00 35.78 O +ATOM 2103 OE2 GLU A 408 26.969 21.219 65.723 1.00 35.78 O +ATOM 2104 N GLY A 409 21.196 20.930 65.394 1.00 43.29 N +ATOM 2105 CA GLY A 409 19.994 21.772 65.353 1.00 43.29 C +ATOM 2106 C GLY A 409 19.320 22.006 66.698 1.00 43.29 C +ATOM 2107 O GLY A 409 18.259 22.625 66.723 1.00 43.29 O +ATOM 2108 N SER A 410 19.888 21.505 67.795 1.00 41.18 N +ATOM 2109 CA SER A 410 19.402 21.800 69.144 1.00 41.18 C +ATOM 2110 C SER A 410 20.002 23.109 69.665 1.00 41.18 C +ATOM 2111 O SER A 410 21.107 23.510 69.280 1.00 41.18 O +ATOM 2112 CB SER A 410 19.695 20.644 70.107 1.00 41.18 C +ATOM 2113 OG SER A 410 21.086 20.546 70.360 1.00 41.18 O +ATOM 2114 N SER A 411 19.318 23.759 70.606 1.00 43.38 N +ATOM 2115 CA SER A 411 19.841 24.933 71.316 1.00 43.38 C +ATOM 2116 C SER A 411 21.189 24.662 71.991 1.00 43.38 C +ATOM 2117 O SER A 411 22.059 25.538 71.987 1.00 43.38 O +ATOM 2118 CB SER A 411 18.836 25.402 72.372 1.00 43.38 C +ATOM 2119 OG SER A 411 18.456 24.327 73.204 1.00 43.38 O +ATOM 2120 N ARG A 412 21.402 23.439 72.500 1.00 40.46 N +ATOM 2121 CA ARG A 412 22.650 23.009 73.150 1.00 40.46 C +ATOM 2122 C ARG A 412 23.864 23.130 72.233 1.00 40.46 C +ATOM 2123 O ARG A 412 24.903 23.617 72.682 1.00 40.46 O +ATOM 2124 CB ARG A 412 22.477 21.573 73.679 1.00 40.46 C +ATOM 2125 CG ARG A 412 23.750 21.030 74.349 1.00 40.46 C +ATOM 2126 CD ARG A 412 23.442 19.724 75.083 1.00 40.46 C +ATOM 2127 NE ARG A 412 24.637 19.141 75.713 1.00 40.46 N +ATOM 2128 CZ ARG A 412 24.621 18.207 76.650 1.00 40.46 C +ATOM 2129 NH1 ARG A 412 23.545 17.606 77.060 1.00 40.46 N +ATOM 2130 NH2 ARG A 412 25.727 17.852 77.234 1.00 40.46 N +ATOM 2131 N GLY A 432 16.747 19.933 66.018 1.00 40.78 N +ATOM 2132 CA GLY A 432 17.261 18.743 66.674 1.00 40.78 C +ATOM 2133 C GLY A 432 16.982 18.748 68.170 1.00 40.78 C +ATOM 2134 O GLY A 432 16.673 19.783 68.751 1.00 40.78 O +ATOM 2135 N VAL A 433 17.126 17.578 68.781 1.00 37.27 N +ATOM 2136 CA VAL A 433 16.869 17.353 70.212 1.00 37.27 C +ATOM 2137 C VAL A 433 18.159 17.067 70.986 1.00 37.27 C +ATOM 2138 O VAL A 433 19.182 16.700 70.398 1.00 37.27 O +ATOM 2139 CB VAL A 433 15.822 16.240 70.414 1.00 37.27 C +ATOM 2140 CG1 VAL A 433 14.462 16.655 69.843 1.00 37.27 C +ATOM 2141 CG2 VAL A 433 16.233 14.912 69.768 1.00 37.27 C +ATOM 2142 N ASP A 434 18.098 17.231 72.309 1.00 35.46 N +ATOM 2143 CA ASP A 434 19.097 16.752 73.273 1.00 35.46 C +ATOM 2144 C ASP A 434 18.597 15.447 73.913 1.00 35.46 C +ATOM 2145 O ASP A 434 17.411 15.305 74.214 1.00 35.46 O +ATOM 2146 CB ASP A 434 19.396 17.844 74.317 1.00 35.46 C +ATOM 2147 CG ASP A 434 20.614 17.530 75.200 1.00 35.46 C +ATOM 2148 OD1 ASP A 434 21.285 16.499 74.981 1.00 35.46 O +ATOM 2149 OD2 ASP A 434 20.963 18.365 76.063 1.00 35.46 O +ATOM 2150 N SER A 435 19.455 14.437 74.024 1.00 34.63 N +ATOM 2151 CA SER A 435 19.055 13.061 74.332 1.00 34.63 C +ATOM 2152 C SER A 435 20.201 12.240 74.932 1.00 34.63 C +ATOM 2153 O SER A 435 21.333 12.700 75.083 1.00 34.63 O +ATOM 2154 CB SER A 435 18.536 12.394 73.047 1.00 34.63 C +ATOM 2155 OG SER A 435 17.294 12.955 72.674 1.00 34.63 O +ATOM 2156 N TYR A 436 19.902 10.985 75.266 1.00 33.86 N +ATOM 2157 CA TYR A 436 20.879 10.004 75.731 1.00 33.86 C +ATOM 2158 C TYR A 436 20.959 8.815 74.776 1.00 33.86 C +ATOM 2159 O TYR A 436 19.965 8.414 74.170 1.00 33.86 O +ATOM 2160 CB TYR A 436 20.530 9.547 77.150 1.00 33.86 C +ATOM 2161 CG TYR A 436 20.502 10.668 78.166 1.00 33.86 C +ATOM 2162 CD1 TYR A 436 21.710 11.220 78.635 1.00 33.86 C +ATOM 2163 CD2 TYR A 436 19.269 11.175 78.620 1.00 33.86 C +ATOM 2164 CE1 TYR A 436 21.681 12.272 79.571 1.00 33.86 C +ATOM 2165 CE2 TYR A 436 19.239 12.224 79.558 1.00 33.86 C +ATOM 2166 CZ TYR A 436 20.448 12.768 80.041 1.00 33.86 C +ATOM 2167 OH TYR A 436 20.434 13.756 80.972 1.00 33.86 O +ATOM 2168 N VAL A 437 22.138 8.206 74.686 1.00 33.19 N +ATOM 2169 CA VAL A 437 22.367 6.935 73.987 1.00 33.19 C +ATOM 2170 C VAL A 437 23.031 5.927 74.931 1.00 33.19 C +ATOM 2171 O VAL A 437 23.819 6.342 75.777 1.00 33.19 O +ATOM 2172 CB VAL A 437 23.194 7.124 72.702 1.00 33.19 C +ATOM 2173 CG1 VAL A 437 22.429 7.965 71.676 1.00 33.19 C +ATOM 2174 CG2 VAL A 437 24.572 7.754 72.941 1.00 33.19 C +ATOM 2175 N PRO A 438 22.751 4.617 74.824 1.00 33.06 N +ATOM 2176 CA PRO A 438 23.412 3.615 75.665 1.00 33.06 C +ATOM 2177 C PRO A 438 24.925 3.601 75.449 1.00 33.06 C +ATOM 2178 O PRO A 438 25.378 3.636 74.306 1.00 33.06 O +ATOM 2179 CB PRO A 438 22.775 2.270 75.298 1.00 33.06 C +ATOM 2180 CG PRO A 438 21.419 2.665 74.716 1.00 33.06 C +ATOM 2181 CD PRO A 438 21.711 3.992 74.022 1.00 33.06 C +ATOM 2182 N TYR A 439 25.714 3.528 76.521 1.00 31.97 N +ATOM 2183 CA TYR A 439 27.159 3.336 76.404 1.00 31.97 C +ATOM 2184 C TYR A 439 27.463 1.930 75.869 1.00 31.97 C +ATOM 2185 O TYR A 439 26.951 0.945 76.393 1.00 31.97 O +ATOM 2186 CB TYR A 439 27.827 3.569 77.760 1.00 31.97 C +ATOM 2187 CG TYR A 439 29.293 3.186 77.801 1.00 31.97 C +ATOM 2188 CD1 TYR A 439 29.705 2.086 78.576 1.00 31.97 C +ATOM 2189 CD2 TYR A 439 30.237 3.895 77.035 1.00 31.97 C +ATOM 2190 CE1 TYR A 439 31.062 1.717 78.620 1.00 31.97 C +ATOM 2191 CE2 TYR A 439 31.596 3.529 77.077 1.00 31.97 C +ATOM 2192 CZ TYR A 439 32.015 2.450 77.882 1.00 31.97 C +ATOM 2193 OH TYR A 439 33.321 2.082 77.924 1.00 31.97 O +ATOM 2194 N ALA A 440 28.310 1.837 74.841 1.00 32.82 N +ATOM 2195 CA ALA A 440 28.582 0.585 74.131 1.00 32.82 C +ATOM 2196 C ALA A 440 30.026 0.071 74.282 1.00 32.82 C +ATOM 2197 O ALA A 440 30.377 -0.937 73.676 1.00 32.82 O +ATOM 2198 CB ALA A 440 28.154 0.753 72.668 1.00 32.82 C +ATOM 2199 N GLY A 441 30.880 0.720 75.080 1.00 32.27 N +ATOM 2200 CA GLY A 441 32.280 0.304 75.220 1.00 32.27 C +ATOM 2201 C GLY A 441 33.228 1.031 74.266 1.00 32.27 C +ATOM 2202 O GLY A 441 32.965 2.157 73.847 1.00 32.27 O +ATOM 2203 N LYS A 442 34.358 0.395 73.945 1.00 32.04 N +ATOM 2204 CA LYS A 442 35.394 0.959 73.066 1.00 32.04 C +ATOM 2205 C LYS A 442 34.977 0.894 71.598 1.00 32.04 C +ATOM 2206 O LYS A 442 34.326 -0.065 71.175 1.00 32.04 O +ATOM 2207 CB LYS A 442 36.720 0.216 73.272 1.00 32.04 C +ATOM 2208 CG LYS A 442 37.306 0.460 74.668 1.00 32.04 C +ATOM 2209 CD LYS A 442 38.619 -0.310 74.827 1.00 32.04 C +ATOM 2210 CE LYS A 442 39.227 -0.013 76.200 1.00 32.04 C +ATOM 2211 NZ LYS A 442 40.482 -0.775 76.400 1.00 32.04 N +ATOM 2212 N LEU A 443 35.411 1.882 70.814 1.00 31.93 N +ATOM 2213 CA LEU A 443 35.164 1.950 69.374 1.00 31.93 C +ATOM 2214 C LEU A 443 35.657 0.699 68.648 1.00 31.93 C +ATOM 2215 O LEU A 443 34.901 0.126 67.870 1.00 31.93 O +ATOM 2216 CB LEU A 443 35.815 3.234 68.831 1.00 31.93 C +ATOM 2217 CG LEU A 443 35.684 3.434 67.307 1.00 31.93 C +ATOM 2218 CD1 LEU A 443 35.702 4.924 66.994 1.00 31.93 C +ATOM 2219 CD2 LEU A 443 36.827 2.808 66.503 1.00 31.93 C +ATOM 2220 N LYS A 444 36.892 0.269 68.931 1.00 33.32 N +ATOM 2221 CA LYS A 444 37.541 -0.861 68.251 1.00 33.32 C +ATOM 2222 C LYS A 444 36.667 -2.118 68.243 1.00 33.32 C +ATOM 2223 O LYS A 444 36.355 -2.639 67.178 1.00 33.32 O +ATOM 2224 CB LYS A 444 38.911 -1.114 68.895 1.00 33.32 C +ATOM 2225 CG LYS A 444 39.564 -2.351 68.271 1.00 33.32 C +ATOM 2226 CD LYS A 444 40.981 -2.570 68.783 1.00 33.32 C +ATOM 2227 CE LYS A 444 41.458 -3.833 68.076 1.00 33.32 C +ATOM 2228 NZ LYS A 444 42.816 -4.221 68.498 1.00 33.32 N +ATOM 2229 N ASP A 445 36.230 -2.556 69.422 1.00 33.63 N +ATOM 2230 CA ASP A 445 35.494 -3.814 69.589 1.00 33.63 C +ATOM 2231 C ASP A 445 34.148 -3.780 68.842 1.00 33.63 C +ATOM 2232 O ASP A 445 33.734 -4.751 68.206 1.00 33.63 O +ATOM 2233 CB ASP A 445 35.263 -4.076 71.093 1.00 33.63 C +ATOM 2234 CG ASP A 445 36.535 -4.023 71.958 1.00 33.63 C +ATOM 2235 OD1 ASP A 445 37.626 -4.384 71.463 1.00 33.63 O +ATOM 2236 OD2 ASP A 445 36.435 -3.541 73.112 1.00 33.63 O +ATOM 2237 N ASN A 446 33.470 -2.630 68.877 1.00 32.27 N +ATOM 2238 CA ASN A 446 32.177 -2.450 68.224 1.00 32.27 C +ATOM 2239 C ASN A 446 32.303 -2.328 66.703 1.00 32.27 C +ATOM 2240 O ASN A 446 31.519 -2.938 65.977 1.00 32.27 O +ATOM 2241 CB ASN A 446 31.489 -1.212 68.803 1.00 32.27 C +ATOM 2242 CG ASN A 446 30.913 -1.457 70.179 1.00 32.27 C +ATOM 2243 OD1 ASN A 446 29.759 -1.825 70.319 1.00 32.27 O +ATOM 2244 ND2 ASN A 446 31.670 -1.219 71.220 1.00 32.27 N +ATOM 2245 N VAL A 447 33.286 -1.564 66.220 1.00 32.00 N +ATOM 2246 CA VAL A 447 33.562 -1.414 64.788 1.00 32.00 C +ATOM 2247 C VAL A 447 33.953 -2.763 64.190 1.00 32.00 C +ATOM 2248 O VAL A 447 33.393 -3.146 63.166 1.00 32.00 O +ATOM 2249 CB VAL A 447 34.630 -0.329 64.544 1.00 32.00 C +ATOM 2250 CG1 VAL A 447 35.186 -0.344 63.116 1.00 32.00 C +ATOM 2251 CG2 VAL A 447 34.017 1.059 64.777 1.00 32.00 C +ATOM 2252 N GLU A 448 34.822 -3.531 64.851 1.00 33.02 N +ATOM 2253 CA GLU A 448 35.183 -4.885 64.419 1.00 33.02 C +ATOM 2254 C GLU A 448 33.945 -5.791 64.304 1.00 33.02 C +ATOM 2255 O GLU A 448 33.728 -6.421 63.264 1.00 33.02 O +ATOM 2256 CB GLU A 448 36.228 -5.461 65.389 1.00 33.02 C +ATOM 2257 CG GLU A 448 36.711 -6.847 64.935 1.00 33.02 C +ATOM 2258 CD GLU A 448 37.832 -7.445 65.801 1.00 33.02 C +ATOM 2259 OE1 GLU A 448 38.174 -8.615 65.511 1.00 33.02 O +ATOM 2260 OE2 GLU A 448 38.351 -6.760 66.712 1.00 33.02 O +ATOM 2261 N ALA A 449 33.082 -5.804 65.324 1.00 31.89 N +ATOM 2262 CA ALA A 449 31.854 -6.592 65.313 1.00 31.89 C +ATOM 2263 C ALA A 449 30.897 -6.183 64.176 1.00 31.89 C +ATOM 2264 O ALA A 449 30.389 -7.050 63.455 1.00 31.89 O +ATOM 2265 CB ALA A 449 31.185 -6.475 66.687 1.00 31.89 C +ATOM 2266 N SER A 450 30.664 -4.882 63.975 1.00 31.30 N +ATOM 2267 CA SER A 450 29.804 -4.377 62.897 1.00 31.30 C +ATOM 2268 C SER A 450 30.381 -4.700 61.515 1.00 31.30 C +ATOM 2269 O SER A 450 29.656 -5.187 60.648 1.00 31.30 O +ATOM 2270 CB SER A 450 29.591 -2.865 63.033 1.00 31.30 C +ATOM 2271 OG SER A 450 28.852 -2.568 64.205 1.00 31.30 O +ATOM 2272 N LEU A 451 31.687 -4.514 61.305 1.00 31.33 N +ATOM 2273 CA LEU A 451 32.336 -4.830 60.030 1.00 31.33 C +ATOM 2274 C LEU A 451 32.328 -6.333 59.740 1.00 31.33 C +ATOM 2275 O LEU A 451 32.137 -6.720 58.590 1.00 31.33 O +ATOM 2276 CB LEU A 451 33.771 -4.291 60.019 1.00 31.33 C +ATOM 2277 CG LEU A 451 33.884 -2.756 60.060 1.00 31.33 C +ATOM 2278 CD1 LEU A 451 35.356 -2.412 60.250 1.00 31.33 C +ATOM 2279 CD2 LEU A 451 33.395 -2.068 58.788 1.00 31.33 C +ATOM 2280 N ASN A 452 32.463 -7.195 60.751 1.00 31.33 N +ATOM 2281 CA ASN A 452 32.322 -8.642 60.576 1.00 31.33 C +ATOM 2282 C ASN A 452 30.909 -9.031 60.113 1.00 31.33 C +ATOM 2283 O ASN A 452 30.775 -9.895 59.241 1.00 31.33 O +ATOM 2284 CB ASN A 452 32.728 -9.355 61.878 1.00 31.33 C +ATOM 2285 CG ASN A 452 34.236 -9.397 62.076 1.00 31.33 C +ATOM 2286 OD1 ASN A 452 35.012 -9.319 61.132 1.00 31.33 O +ATOM 2287 ND2 ASN A 452 34.689 -9.568 63.296 1.00 31.33 N +ATOM 2288 N LYS A 453 29.861 -8.360 60.614 1.00 30.91 N +ATOM 2289 CA LYS A 453 28.490 -8.531 60.101 1.00 30.91 C +ATOM 2290 C LYS A 453 28.382 -8.073 58.646 1.00 30.91 C +ATOM 2291 O LYS A 453 27.904 -8.847 57.823 1.00 30.91 O +ATOM 2292 CB LYS A 453 27.468 -7.795 60.978 1.00 30.91 C +ATOM 2293 CG LYS A 453 27.350 -8.392 62.389 1.00 30.91 C +ATOM 2294 CD LYS A 453 26.463 -7.483 63.243 1.00 30.91 C +ATOM 2295 CE LYS A 453 26.479 -7.896 64.716 1.00 30.91 C +ATOM 2296 NZ LYS A 453 25.719 -6.920 65.534 1.00 30.91 N +ATOM 2297 N VAL A 454 28.893 -6.883 58.308 1.00 30.82 N +ATOM 2298 CA VAL A 454 28.898 -6.362 56.924 1.00 30.82 C +ATOM 2299 C VAL A 454 29.598 -7.336 55.972 1.00 30.82 C +ATOM 2300 O VAL A 454 29.008 -7.735 54.969 1.00 30.82 O +ATOM 2301 CB VAL A 454 29.540 -4.959 56.844 1.00 30.82 C +ATOM 2302 CG1 VAL A 454 29.692 -4.462 55.398 1.00 30.82 C +ATOM 2303 CG2 VAL A 454 28.695 -3.918 57.590 1.00 30.82 C +ATOM 2304 N LYS A 455 30.813 -7.789 56.310 1.00 30.82 N +ATOM 2305 CA LYS A 455 31.584 -8.755 55.509 1.00 30.82 C +ATOM 2306 C LYS A 455 30.828 -10.068 55.321 1.00 30.82 C +ATOM 2307 O LYS A 455 30.747 -10.575 54.207 1.00 30.82 O +ATOM 2308 CB LYS A 455 32.940 -9.032 56.175 1.00 30.82 C +ATOM 2309 CG LYS A 455 33.894 -7.831 56.126 1.00 30.82 C +ATOM 2310 CD LYS A 455 35.132 -8.113 56.981 1.00 30.82 C +ATOM 2311 CE LYS A 455 36.164 -6.994 56.839 1.00 30.82 C +ATOM 2312 NZ LYS A 455 37.392 -7.319 57.604 1.00 30.82 N +ATOM 2313 N SER A 456 30.224 -10.590 56.387 1.00 31.01 N +ATOM 2314 CA SER A 456 29.429 -11.821 56.322 1.00 31.01 C +ATOM 2315 C SER A 456 28.205 -11.659 55.415 1.00 31.01 C +ATOM 2316 O SER A 456 27.933 -12.519 54.579 1.00 31.01 O +ATOM 2317 CB SER A 456 28.995 -12.255 57.724 1.00 31.01 C +ATOM 2318 OG SER A 456 30.127 -12.443 58.550 1.00 31.01 O +ATOM 2319 N THR A 457 27.490 -10.534 55.507 1.00 30.97 N +ATOM 2320 CA THR A 457 26.357 -10.242 54.618 1.00 30.97 C +ATOM 2321 C THR A 457 26.801 -10.060 53.163 1.00 30.97 C +ATOM 2322 O THR A 457 26.120 -10.550 52.264 1.00 30.97 O +ATOM 2323 CB THR A 457 25.585 -9.001 55.088 1.00 30.97 C +ATOM 2324 OG1 THR A 457 25.216 -9.105 56.443 1.00 30.97 O +ATOM 2325 CG2 THR A 457 24.274 -8.813 54.325 1.00 30.97 C +ATOM 2326 N MET A 458 27.950 -9.424 52.910 1.00 30.70 N +ATOM 2327 CA MET A 458 28.526 -9.313 51.563 1.00 30.70 C +ATOM 2328 C MET A 458 28.839 -10.687 50.964 1.00 30.70 C +ATOM 2329 O MET A 458 28.446 -10.950 49.825 1.00 30.70 O +ATOM 2330 CB MET A 458 29.794 -8.450 51.586 1.00 30.70 C +ATOM 2331 CG MET A 458 29.475 -6.967 51.762 1.00 30.70 C +ATOM 2332 SD MET A 458 30.933 -5.902 51.906 1.00 30.70 S +ATOM 2333 CE MET A 458 31.639 -6.045 50.242 1.00 30.70 C +ATOM 2334 N CYS A 459 29.437 -11.596 51.742 1.00 30.76 N +ATOM 2335 CA CYS A 459 29.652 -12.983 51.324 1.00 30.76 C +ATOM 2336 C CYS A 459 28.335 -13.686 50.965 1.00 30.76 C +ATOM 2337 O CYS A 459 28.260 -14.320 49.915 1.00 30.76 O +ATOM 2338 CB CYS A 459 30.379 -13.753 52.431 1.00 30.76 C +ATOM 2339 SG CYS A 459 32.100 -13.205 52.543 1.00 30.76 S +ATOM 2340 N ASN A 460 27.280 -13.510 51.769 1.00 31.07 N +ATOM 2341 CA ASN A 460 25.947 -14.057 51.477 1.00 31.07 C +ATOM 2342 C ASN A 460 25.320 -13.463 50.202 1.00 31.07 C +ATOM 2343 O ASN A 460 24.473 -14.098 49.579 1.00 31.07 O +ATOM 2344 CB ASN A 460 25.017 -13.821 52.681 1.00 31.07 C +ATOM 2345 CG ASN A 460 25.398 -14.617 53.916 1.00 31.07 C +ATOM 2346 OD1 ASN A 460 26.177 -15.551 53.886 1.00 31.07 O +ATOM 2347 ND2 ASN A 460 24.825 -14.290 55.050 1.00 31.07 N +ATOM 2348 N CYS A 461 25.731 -12.262 49.790 1.00 31.01 N +ATOM 2349 CA CYS A 461 25.342 -11.649 48.517 1.00 31.01 C +ATOM 2350 C CYS A 461 26.250 -12.073 47.342 1.00 31.01 C +ATOM 2351 O CYS A 461 26.013 -11.659 46.209 1.00 31.01 O +ATOM 2352 CB CYS A 461 25.311 -10.121 48.674 1.00 31.01 C +ATOM 2353 SG CYS A 461 24.100 -9.630 49.935 1.00 31.01 S +ATOM 2354 N GLY A 462 27.295 -12.872 47.586 1.00 30.94 N +ATOM 2355 CA GLY A 462 28.285 -13.252 46.575 1.00 30.94 C +ATOM 2356 C GLY A 462 29.272 -12.134 46.215 1.00 30.94 C +ATOM 2357 O GLY A 462 29.803 -12.113 45.099 1.00 30.94 O +ATOM 2358 N ALA A 463 29.495 -11.177 47.119 1.00 30.82 N +ATOM 2359 CA ALA A 463 30.334 -10.002 46.904 1.00 30.82 C +ATOM 2360 C ALA A 463 31.551 -9.977 47.842 1.00 30.82 C +ATOM 2361 O ALA A 463 31.431 -10.210 49.040 1.00 30.82 O +ATOM 2362 CB ALA A 463 29.468 -8.751 47.072 1.00 30.82 C +ATOM 2363 N LEU A 464 32.720 -9.644 47.288 1.00 30.79 N +ATOM 2364 CA LEU A 464 33.987 -9.483 48.013 1.00 30.79 C +ATOM 2365 C LEU A 464 34.481 -8.030 48.043 1.00 30.79 C +ATOM 2366 O LEU A 464 35.442 -7.726 48.742 1.00 30.79 O +ATOM 2367 CB LEU A 464 35.049 -10.402 47.385 1.00 30.79 C +ATOM 2368 CG LEU A 464 34.761 -11.905 47.528 1.00 30.79 C +ATOM 2369 CD1 LEU A 464 35.865 -12.704 46.841 1.00 30.79 C +ATOM 2370 CD2 LEU A 464 34.699 -12.312 48.999 1.00 30.79 C +ATOM 2371 N THR A 465 33.838 -7.125 47.307 1.00 30.73 N +ATOM 2372 CA THR A 465 34.125 -5.682 47.317 1.00 30.73 C +ATOM 2373 C THR A 465 32.823 -4.884 47.321 1.00 30.73 C +ATOM 2374 O THR A 465 31.780 -5.406 46.919 1.00 30.73 O +ATOM 2375 CB THR A 465 35.007 -5.253 46.130 1.00 30.73 C +ATOM 2376 OG1 THR A 465 34.302 -5.319 44.911 1.00 30.73 O +ATOM 2377 CG2 THR A 465 36.267 -6.100 45.949 1.00 30.73 C +ATOM 2378 N ILE A 466 32.856 -3.620 47.753 1.00 30.73 N +ATOM 2379 CA ILE A 466 31.668 -2.745 47.722 1.00 30.73 C +ATOM 2380 C ILE A 466 31.110 -2.582 46.291 1.00 30.73 C +ATOM 2381 O ILE A 466 29.909 -2.789 46.124 1.00 30.73 O +ATOM 2382 CB ILE A 466 31.940 -1.399 48.440 1.00 30.73 C +ATOM 2383 CG1 ILE A 466 32.123 -1.650 49.956 1.00 30.73 C +ATOM 2384 CG2 ILE A 466 30.796 -0.397 48.195 1.00 30.73 C +ATOM 2385 CD1 ILE A 466 32.700 -0.446 50.709 1.00 30.73 C +ATOM 2386 N PRO A 467 31.920 -2.354 45.234 1.00 31.01 N +ATOM 2387 CA PRO A 467 31.408 -2.347 43.860 1.00 31.01 C +ATOM 2388 C PRO A 467 30.732 -3.664 43.450 1.00 31.01 C +ATOM 2389 O PRO A 467 29.695 -3.661 42.785 1.00 31.01 O +ATOM 2390 CB PRO A 467 32.621 -2.050 42.974 1.00 31.01 C +ATOM 2391 CG PRO A 467 33.525 -1.228 43.890 1.00 31.01 C +ATOM 2392 CD PRO A 467 33.296 -1.871 45.255 1.00 31.01 C +ATOM 2393 N GLN A 468 31.271 -4.813 43.883 1.00 30.85 N +ATOM 2394 CA GLN A 468 30.602 -6.093 43.651 1.00 30.85 C +ATOM 2395 C GLN A 468 29.264 -6.166 44.389 1.00 30.85 C +ATOM 2396 O GLN A 468 28.298 -6.633 43.795 1.00 30.85 O +ATOM 2397 CB GLN A 468 31.484 -7.275 44.059 1.00 30.85 C +ATOM 2398 CG GLN A 468 32.624 -7.538 43.067 1.00 30.85 C +ATOM 2399 CD GLN A 468 33.492 -8.708 43.515 1.00 30.85 C +ATOM 2400 OE1 GLN A 468 33.261 -9.340 44.534 1.00 30.85 O +ATOM 2401 NE2 GLN A 468 34.492 -9.085 42.753 1.00 30.85 N +ATOM 2402 N LEU A 469 29.183 -5.685 45.635 1.00 30.67 N +ATOM 2403 CA LEU A 469 27.932 -5.635 46.394 1.00 30.67 C +ATOM 2404 C LEU A 469 26.889 -4.786 45.659 1.00 30.67 C +ATOM 2405 O LEU A 469 25.792 -5.273 45.427 1.00 30.67 O +ATOM 2406 CB LEU A 469 28.202 -5.117 47.820 1.00 30.67 C +ATOM 2407 CG LEU A 469 26.939 -4.990 48.693 1.00 30.67 C +ATOM 2408 CD1 LEU A 469 26.350 -6.358 49.047 1.00 30.67 C +ATOM 2409 CD2 LEU A 469 27.268 -4.222 49.973 1.00 30.67 C +ATOM 2410 N GLN A 470 27.252 -3.579 45.225 1.00 30.88 N +ATOM 2411 CA GLN A 470 26.369 -2.674 44.480 1.00 30.88 C +ATOM 2412 C GLN A 470 25.854 -3.298 43.173 1.00 30.88 C +ATOM 2413 O GLN A 470 24.701 -3.103 42.805 1.00 30.88 O +ATOM 2414 CB GLN A 470 27.145 -1.385 44.177 1.00 30.88 C +ATOM 2415 CG GLN A 470 27.440 -0.565 45.445 1.00 30.88 C +ATOM 2416 CD GLN A 470 28.403 0.595 45.200 1.00 30.88 C +ATOM 2417 OE1 GLN A 470 29.027 0.712 44.157 1.00 30.88 O +ATOM 2418 NE2 GLN A 470 28.582 1.481 46.155 1.00 30.88 N +ATOM 2419 N SER A 471 26.686 -4.089 42.482 1.00 31.30 N +ATOM 2420 CA SER A 471 26.292 -4.757 41.230 1.00 31.30 C +ATOM 2421 C SER A 471 25.491 -6.056 41.414 1.00 31.30 C +ATOM 2422 O SER A 471 24.698 -6.410 40.544 1.00 31.30 O +ATOM 2423 CB SER A 471 27.532 -5.024 40.372 1.00 31.30 C +ATOM 2424 OG SER A 471 28.360 -6.016 40.957 1.00 31.30 O +ATOM 2425 N LYS A 472 25.719 -6.802 42.506 1.00 31.23 N +ATOM 2426 CA LYS A 472 25.167 -8.156 42.721 1.00 31.23 C +ATOM 2427 C LYS A 472 24.010 -8.204 43.713 1.00 31.23 C +ATOM 2428 O LYS A 472 23.274 -9.191 43.728 1.00 31.23 O +ATOM 2429 CB LYS A 472 26.265 -9.109 43.213 1.00 31.23 C +ATOM 2430 CG LYS A 472 27.357 -9.394 42.176 1.00 31.23 C +ATOM 2431 CD LYS A 472 28.373 -10.361 42.792 1.00 31.23 C +ATOM 2432 CE LYS A 472 29.485 -10.722 41.807 1.00 31.23 C +ATOM 2433 NZ LYS A 472 30.387 -11.734 42.411 1.00 31.23 N +ATOM 2434 N ALA A 473 23.887 -7.205 44.584 1.00 31.07 N +ATOM 2435 CA ALA A 473 22.863 -7.185 45.612 1.00 31.07 C +ATOM 2436 C ALA A 473 21.467 -7.228 44.983 1.00 31.07 C +ATOM 2437 O ALA A 473 21.137 -6.472 44.074 1.00 31.07 O +ATOM 2438 CB ALA A 473 23.028 -5.953 46.503 1.00 31.07 C +ATOM 2439 N LYS A 474 20.627 -8.107 45.520 1.00 30.91 N +ATOM 2440 CA LYS A 474 19.194 -8.142 45.243 1.00 30.91 C +ATOM 2441 C LYS A 474 18.501 -7.615 46.487 1.00 30.91 C +ATOM 2442 O LYS A 474 18.694 -8.178 47.562 1.00 30.91 O +ATOM 2443 CB LYS A 474 18.780 -9.580 44.930 1.00 30.91 C +ATOM 2444 CG LYS A 474 19.392 -10.141 43.634 1.00 30.91 C +ATOM 2445 CD LYS A 474 19.028 -11.622 43.443 1.00 30.91 C +ATOM 2446 CE LYS A 474 19.658 -12.482 44.551 1.00 30.91 C +ATOM 2447 NZ LYS A 474 19.237 -13.901 44.500 1.00 30.91 N +ATOM 2448 N ILE A 475 17.729 -6.546 46.362 1.00 31.17 N +ATOM 2449 CA ILE A 475 17.046 -5.897 47.487 1.00 31.17 C +ATOM 2450 C ILE A 475 15.537 -6.018 47.274 1.00 31.17 C +ATOM 2451 O ILE A 475 15.060 -5.824 46.159 1.00 31.17 O +ATOM 2452 CB ILE A 475 17.529 -4.436 47.647 1.00 31.17 C +ATOM 2453 CG1 ILE A 475 19.039 -4.389 47.978 1.00 31.17 C +ATOM 2454 CG2 ILE A 475 16.732 -3.730 48.753 1.00 31.17 C +ATOM 2455 CD1 ILE A 475 19.644 -2.980 47.987 1.00 31.17 C +ATOM 2456 N THR A 476 14.789 -6.332 48.331 1.00 31.54 N +ATOM 2457 CA THR A 476 13.320 -6.361 48.307 1.00 31.54 C +ATOM 2458 C THR A 476 12.748 -5.424 49.353 1.00 31.54 C +ATOM 2459 O THR A 476 13.263 -5.362 50.468 1.00 31.54 O +ATOM 2460 CB THR A 476 12.767 -7.783 48.478 1.00 31.54 C +ATOM 2461 OG1 THR A 476 11.376 -7.794 48.254 1.00 31.54 O +ATOM 2462 CG2 THR A 476 13.041 -8.437 49.832 1.00 31.54 C +ATOM 2463 N LEU A 477 11.670 -4.725 48.996 1.00 31.93 N +ATOM 2464 CA LEU A 477 10.803 -4.065 49.966 1.00 31.93 C +ATOM 2465 C LEU A 477 10.065 -5.150 50.763 1.00 31.93 C +ATOM 2466 O LEU A 477 9.685 -6.180 50.195 1.00 31.93 O +ATOM 2467 CB LEU A 477 9.832 -3.126 49.223 1.00 31.93 C +ATOM 2468 CG LEU A 477 9.029 -2.185 50.144 1.00 31.93 C +ATOM 2469 CD1 LEU A 477 9.888 -1.040 50.678 1.00 31.93 C +ATOM 2470 CD2 LEU A 477 7.871 -1.559 49.367 1.00 31.93 C +ATOM 2471 N VAL A 478 9.877 -4.937 52.060 1.00 33.19 N +ATOM 2472 CA VAL A 478 9.201 -5.887 52.944 1.00 33.19 C +ATOM 2473 C VAL A 478 7.958 -5.283 53.584 1.00 33.19 C +ATOM 2474 O VAL A 478 7.859 -4.078 53.796 1.00 33.19 O +ATOM 2475 CB VAL A 478 10.156 -6.481 53.989 1.00 33.19 C +ATOM 2476 CG1 VAL A 478 11.255 -7.298 53.312 1.00 33.19 C +ATOM 2477 CG2 VAL A 478 10.791 -5.473 54.950 1.00 33.19 C +ATOM 2478 N SER A 479 6.984 -6.140 53.889 1.00 33.24 N +ATOM 2479 CA SER A 479 5.779 -5.724 54.608 1.00 33.24 C +ATOM 2480 C SER A 479 6.091 -5.415 56.076 1.00 33.24 C +ATOM 2481 O SER A 479 7.063 -5.927 56.633 1.00 33.24 O +ATOM 2482 CB SER A 479 4.694 -6.799 54.489 1.00 33.24 C +ATOM 2483 OG SER A 479 4.966 -7.906 55.329 1.00 33.24 O +ATOM 2484 N SER A 480 5.223 -4.659 56.751 1.00 35.78 N +ATOM 2485 CA SER A 480 5.335 -4.440 58.201 1.00 35.78 C +ATOM 2486 C SER A 480 5.323 -5.747 59.005 1.00 35.78 C +ATOM 2487 O SER A 480 5.983 -5.836 60.035 1.00 35.78 O +ATOM 2488 CB SER A 480 4.204 -3.529 58.682 1.00 35.78 C +ATOM 2489 OG SER A 480 2.945 -4.052 58.290 1.00 35.78 O +ATOM 2490 N VAL A 481 4.632 -6.785 58.516 1.00 35.35 N +ATOM 2491 CA VAL A 481 4.609 -8.120 59.137 1.00 35.35 C +ATOM 2492 C VAL A 481 5.964 -8.818 58.992 1.00 35.35 C +ATOM 2493 O VAL A 481 6.469 -9.383 59.956 1.00 35.35 O +ATOM 2494 CB VAL A 481 3.480 -8.981 58.537 1.00 35.35 C +ATOM 2495 CG1 VAL A 481 3.399 -10.358 59.206 1.00 35.35 C +ATOM 2496 CG2 VAL A 481 2.114 -8.296 58.703 1.00 35.35 C +ATOM 2497 N SER A 482 6.595 -8.713 57.822 1.00 35.62 N +ATOM 2498 CA SER A 482 7.939 -9.251 57.576 1.00 35.62 C +ATOM 2499 C SER A 482 9.011 -8.569 58.440 1.00 35.62 C +ATOM 2500 O SER A 482 9.971 -9.213 58.849 1.00 35.62 O +ATOM 2501 CB SER A 482 8.299 -9.053 56.107 1.00 35.62 C +ATOM 2502 OG SER A 482 7.337 -9.629 55.236 1.00 35.62 O +ATOM 2503 N ILE A 483 8.844 -7.278 58.765 1.00 40.61 N +ATOM 2504 CA ILE A 483 9.733 -6.570 59.705 1.00 40.61 C +ATOM 2505 C ILE A 483 9.637 -7.187 61.107 1.00 40.61 C +ATOM 2506 O ILE A 483 10.661 -7.395 61.756 1.00 40.61 O +ATOM 2507 CB ILE A 483 9.420 -5.053 59.735 1.00 40.61 C +ATOM 2508 CG1 ILE A 483 9.732 -4.415 58.363 1.00 40.61 C +ATOM 2509 CG2 ILE A 483 10.210 -4.333 60.848 1.00 40.61 C +ATOM 2510 CD1 ILE A 483 9.302 -2.948 58.231 1.00 40.61 C +END diff --git a/tests/files/pdb/1DAW_af.pdb b/tests/files/pdb/1DAW_af.pdb new file mode 100644 index 00000000..75e969cf --- /dev/null +++ b/tests/files/pdb/1DAW_af.pdb @@ -0,0 +1,2736 @@ +REMARK TITLE 1DAW AlphaFold-start MR +REMARK Log-Likelihood Gain: 2048.362 +REMARK RFZ=9.3 TFZ=7.3 PAK=0 LLG=153 TFZ==9.3 LLG=2048 TFZ==38.0 PAK=0 LLG=2048 TFZ==38.0 +REMARK ENSEMBLE e_P28523 EULER 18.15 74.16 270.75 FRAC 0.133 -0.001 0.274 +CRYST1 143.110 59.200 45.810 90.00 103.56 90.00 C 1 2 1 4 +SCALE1 0.006988 -0.000000 0.001685 0.00000 +SCALE2 0.000000 0.016892 -0.000000 0.00000 +SCALE3 0.000000 0.000000 0.022455 0.00000 +ATOM 1 N SER A 2 18.794 -9.930 -8.756 1.00 40.57 N +ATOM 2 CA SER A 2 17.744 -9.197 -8.072 1.00 40.57 C +ATOM 3 C SER A 2 17.794 -7.707 -8.440 1.00 40.57 C +ATOM 4 O SER A 2 18.868 -7.134 -8.627 1.00 40.57 O +ATOM 5 CB SER A 2 17.951 -9.412 -6.574 1.00 40.57 C +ATOM 6 OG SER A 2 17.040 -8.621 -5.852 1.00 40.57 O +ATOM 7 N LYS A 3 16.627 -7.065 -8.543 1.00 36.78 N +ATOM 8 CA LYS A 3 16.478 -5.609 -8.695 1.00 36.78 C +ATOM 9 C LYS A 3 15.380 -5.145 -7.746 1.00 36.78 C +ATOM 10 O LYS A 3 14.373 -5.830 -7.618 1.00 36.78 O +ATOM 11 CB LYS A 3 16.192 -5.236 -10.167 1.00 36.78 C +ATOM 12 CG LYS A 3 16.094 -3.711 -10.382 1.00 36.78 C +ATOM 13 CD LYS A 3 15.981 -3.282 -11.860 1.00 36.78 C +ATOM 14 CE LYS A 3 15.793 -1.754 -11.916 1.00 36.78 C +ATOM 15 NZ LYS A 3 15.786 -1.149 -13.272 1.00 36.78 N +ATOM 16 N ALA A 4 15.551 -3.991 -7.109 1.00 35.62 N +ATOM 17 CA ALA A 4 14.529 -3.430 -6.234 1.00 35.62 C +ATOM 18 C ALA A 4 13.197 -3.250 -6.980 1.00 35.62 C +ATOM 19 O ALA A 4 13.183 -2.780 -8.119 1.00 35.62 O +ATOM 20 CB ALA A 4 15.032 -2.089 -5.699 1.00 35.62 C +ATOM 21 N ARG A 5 12.067 -3.570 -6.336 1.00 35.59 N +ATOM 22 CA ARG A 5 10.728 -3.360 -6.928 1.00 35.59 C +ATOM 23 C ARG A 5 10.376 -1.884 -7.105 1.00 35.59 C +ATOM 24 O ARG A 5 9.622 -1.534 -8.006 1.00 35.59 O +ATOM 25 CB ARG A 5 9.637 -4.028 -6.082 1.00 35.59 C +ATOM 26 CG ARG A 5 9.797 -5.549 -6.015 1.00 35.59 C +ATOM 27 CD ARG A 5 8.540 -6.226 -5.467 1.00 35.59 C +ATOM 28 NE ARG A 5 8.303 -5.913 -4.044 1.00 35.59 N +ATOM 29 CZ ARG A 5 7.384 -5.112 -3.537 1.00 35.59 C +ATOM 30 NH1 ARG A 5 6.576 -4.400 -4.265 1.00 35.59 N +ATOM 31 NH2 ARG A 5 7.218 -5.005 -2.264 1.00 35.59 N +ATOM 32 N VAL A 6 10.909 -1.031 -6.234 1.00 35.56 N +ATOM 33 CA VAL A 6 10.679 0.419 -6.209 1.00 35.56 C +ATOM 34 C VAL A 6 12.007 1.158 -6.137 1.00 35.56 C +ATOM 35 O VAL A 6 13.003 0.599 -5.691 1.00 35.56 O +ATOM 36 CB VAL A 6 9.760 0.844 -5.049 1.00 35.56 C +ATOM 37 CG1 VAL A 6 8.352 0.271 -5.252 1.00 35.56 C +ATOM 38 CG2 VAL A 6 10.304 0.437 -3.674 1.00 35.56 C +ATOM 39 N TYR A 7 12.019 2.409 -6.603 1.00 35.56 N +ATOM 40 CA TYR A 7 13.173 3.322 -6.547 1.00 35.56 C +ATOM 41 C TYR A 7 14.469 2.809 -7.199 1.00 35.56 C +ATOM 42 O TYR A 7 15.536 3.391 -7.021 1.00 35.56 O +ATOM 43 CB TYR A 7 13.377 3.807 -5.103 1.00 35.56 C +ATOM 44 CG TYR A 7 12.110 4.336 -4.462 1.00 35.56 C +ATOM 45 CD1 TYR A 7 11.378 5.369 -5.085 1.00 35.56 C +ATOM 46 CD2 TYR A 7 11.655 3.784 -3.248 1.00 35.56 C +ATOM 47 CE1 TYR A 7 10.186 5.839 -4.502 1.00 35.56 C +ATOM 48 CE2 TYR A 7 10.465 4.254 -2.664 1.00 35.56 C +ATOM 49 CZ TYR A 7 9.727 5.275 -3.292 1.00 35.56 C +ATOM 50 OH TYR A 7 8.573 5.706 -2.727 1.00 35.56 O +ATOM 51 N ALA A 8 14.378 1.762 -8.019 1.00 35.74 N +ATOM 52 CA ALA A 8 15.537 0.997 -8.452 1.00 35.74 C +ATOM 53 C ALA A 8 16.554 1.793 -9.281 1.00 35.74 C +ATOM 54 O ALA A 8 17.750 1.535 -9.203 1.00 35.74 O +ATOM 55 CB ALA A 8 15.012 -0.191 -9.245 1.00 35.74 C +ATOM 56 N ASP A 9 16.081 2.768 -10.057 1.00 35.74 N +ATOM 57 CA ASP A 9 16.918 3.563 -10.954 1.00 35.74 C +ATOM 58 C ASP A 9 17.134 4.999 -10.437 1.00 35.74 C +ATOM 59 O ASP A 9 17.725 5.815 -11.137 1.00 35.74 O +ATOM 60 CB ASP A 9 16.354 3.481 -12.384 1.00 35.74 C +ATOM 61 CG ASP A 9 16.258 2.027 -12.883 1.00 35.74 C +ATOM 62 OD1 ASP A 9 17.248 1.268 -12.811 1.00 35.74 O +ATOM 63 OD2 ASP A 9 15.165 1.576 -13.288 1.00 35.74 O +ATOM 64 N VAL A 10 16.719 5.328 -9.201 1.00 35.83 N +ATOM 65 CA VAL A 10 16.811 6.704 -8.663 1.00 35.83 C +ATOM 66 C VAL A 10 18.237 7.244 -8.749 1.00 35.83 C +ATOM 67 O VAL A 10 18.447 8.327 -9.285 1.00 35.83 O +ATOM 68 CB VAL A 10 16.280 6.794 -7.216 1.00 35.83 C +ATOM 69 CG1 VAL A 10 16.596 8.141 -6.546 1.00 35.83 C +ATOM 70 CG2 VAL A 10 14.756 6.632 -7.207 1.00 35.83 C +ATOM 71 N ASN A 11 19.238 6.483 -8.298 1.00 35.59 N +ATOM 72 CA ASN A 11 20.632 6.931 -8.373 1.00 35.59 C +ATOM 73 C ASN A 11 21.213 6.881 -9.793 1.00 35.59 C +ATOM 74 O ASN A 11 22.167 7.599 -10.068 1.00 35.59 O +ATOM 75 CB ASN A 11 21.491 6.118 -7.397 1.00 35.59 C +ATOM 76 CG ASN A 11 21.170 6.419 -5.949 1.00 35.59 C +ATOM 77 OD1 ASN A 11 20.891 7.537 -5.559 1.00 35.59 O +ATOM 78 ND2 ASN A 11 21.201 5.420 -5.110 1.00 35.59 N +ATOM 79 N VAL A 12 20.648 6.074 -10.697 1.00 35.96 N +ATOM 80 CA VAL A 12 21.040 6.061 -12.121 1.00 35.96 C +ATOM 81 C VAL A 12 20.605 7.361 -12.800 1.00 35.96 C +ATOM 82 O VAL A 12 21.331 7.901 -13.626 1.00 35.96 O +ATOM 83 CB VAL A 12 20.440 4.843 -12.860 1.00 35.96 C +ATOM 84 CG1 VAL A 12 20.845 4.800 -14.338 1.00 35.96 C +ATOM 85 CG2 VAL A 12 20.884 3.522 -12.215 1.00 35.96 C +ATOM 86 N LEU A 13 19.435 7.878 -12.420 1.00 35.86 N +ATOM 87 CA LEU A 13 18.852 9.104 -12.969 1.00 35.86 C +ATOM 88 C LEU A 13 19.393 10.386 -12.311 1.00 35.86 C +ATOM 89 O LEU A 13 19.206 11.478 -12.846 1.00 35.86 O +ATOM 90 CB LEU A 13 17.323 9.011 -12.822 1.00 35.86 C +ATOM 91 CG LEU A 13 16.665 7.845 -13.588 1.00 35.86 C +ATOM 92 CD1 LEU A 13 15.183 7.774 -13.222 1.00 35.86 C +ATOM 93 CD2 LEU A 13 16.793 7.998 -15.104 1.00 35.86 C +ATOM 94 N ARG A 14 20.040 10.285 -11.143 1.00 35.90 N +ATOM 95 CA ARG A 14 20.640 11.431 -10.442 1.00 35.90 C +ATOM 96 C ARG A 14 22.037 11.756 -10.991 1.00 35.90 C +ATOM 97 O ARG A 14 22.750 10.847 -11.418 1.00 35.90 O +ATOM 98 CB ARG A 14 20.695 11.163 -8.929 1.00 35.90 C +ATOM 99 CG ARG A 14 19.329 11.226 -8.226 1.00 35.90 C +ATOM 100 CD ARG A 14 18.805 12.652 -8.039 1.00 35.90 C +ATOM 101 NE ARG A 14 17.639 12.658 -7.133 1.00 35.90 N +ATOM 102 CZ ARG A 14 16.941 13.709 -6.743 1.00 35.90 C +ATOM 103 NH1 ARG A 14 17.147 14.896 -7.245 1.00 35.90 N +ATOM 104 NH2 ARG A 14 16.029 13.590 -5.819 1.00 35.90 N +ATOM 105 N PRO A 15 22.475 13.030 -10.908 1.00 35.93 N +ATOM 106 CA PRO A 15 23.855 13.405 -11.205 1.00 35.93 C +ATOM 107 C PRO A 15 24.846 12.573 -10.387 1.00 35.93 C +ATOM 108 O PRO A 15 24.561 12.207 -9.243 1.00 35.93 O +ATOM 109 CB PRO A 15 23.975 14.894 -10.867 1.00 35.93 C +ATOM 110 CG PRO A 15 22.538 15.399 -10.948 1.00 35.93 C +ATOM 111 CD PRO A 15 21.721 14.199 -10.482 1.00 35.93 C +ATOM 112 N LYS A 16 26.023 12.295 -10.950 1.00 36.29 N +ATOM 113 CA LYS A 16 27.037 11.443 -10.314 1.00 36.29 C +ATOM 114 C LYS A 16 27.438 11.964 -8.930 1.00 36.29 C +ATOM 115 O LYS A 16 27.617 11.182 -8.004 1.00 36.29 O +ATOM 116 CB LYS A 16 28.231 11.334 -11.270 1.00 36.29 C +ATOM 117 CG LYS A 16 29.312 10.396 -10.725 1.00 36.29 C +ATOM 118 CD LYS A 16 30.419 10.175 -11.756 1.00 36.29 C +ATOM 119 CE LYS A 16 31.491 9.304 -11.103 1.00 36.29 C +ATOM 120 NZ LYS A 16 32.689 9.167 -11.958 1.00 36.29 N +ATOM 121 N GLU A 17 27.480 13.281 -8.755 1.00 36.33 N +ATOM 122 CA GLU A 17 27.826 13.977 -7.511 1.00 36.33 C +ATOM 123 C GLU A 17 26.830 13.706 -6.372 1.00 36.33 C +ATOM 124 O GLU A 17 27.128 13.982 -5.206 1.00 36.33 O +ATOM 125 CB GLU A 17 27.872 15.502 -7.741 1.00 36.33 C +ATOM 126 CG GLU A 17 28.785 15.993 -8.881 1.00 36.33 C +ATOM 127 CD GLU A 17 28.260 15.673 -10.292 1.00 36.33 C +ATOM 128 OE1 GLU A 17 29.100 15.549 -11.206 1.00 36.33 O +ATOM 129 OE2 GLU A 17 27.037 15.437 -10.435 1.00 36.33 O +ATOM 130 N TYR A 18 25.636 13.193 -6.680 1.00 35.65 N +ATOM 131 CA TYR A 18 24.636 12.824 -5.683 1.00 35.65 C +ATOM 132 C TYR A 18 25.061 11.583 -4.886 1.00 35.65 C +ATOM 133 O TYR A 18 24.962 11.590 -3.658 1.00 35.65 O +ATOM 134 CB TYR A 18 23.288 12.600 -6.373 1.00 35.65 C +ATOM 135 CG TYR A 18 22.175 12.258 -5.409 1.00 35.65 C +ATOM 136 CD1 TYR A 18 21.844 10.908 -5.180 1.00 35.65 C +ATOM 137 CD2 TYR A 18 21.502 13.281 -4.712 1.00 35.65 C +ATOM 138 CE1 TYR A 18 20.845 10.578 -4.250 1.00 35.65 C +ATOM 139 CE2 TYR A 18 20.499 12.950 -3.780 1.00 35.65 C +ATOM 140 CZ TYR A 18 20.177 11.597 -3.549 1.00 35.65 C +ATOM 141 OH TYR A 18 19.236 11.269 -2.636 1.00 35.65 O +ATOM 142 N TRP A 19 25.562 10.547 -5.562 1.00 35.83 N +ATOM 143 CA TRP A 19 25.874 9.242 -4.963 1.00 35.83 C +ATOM 144 C TRP A 19 27.373 8.915 -4.935 1.00 35.83 C +ATOM 145 O TRP A 19 27.785 8.052 -4.157 1.00 35.83 O +ATOM 146 CB TRP A 19 25.075 8.159 -5.696 1.00 35.83 C +ATOM 147 CG TRP A 19 25.267 8.133 -7.181 1.00 35.83 C +ATOM 148 CD1 TRP A 19 24.406 8.635 -8.093 1.00 35.83 C +ATOM 149 CD2 TRP A 19 26.390 7.597 -7.948 1.00 35.83 C +ATOM 150 NE1 TRP A 19 24.900 8.433 -9.364 1.00 35.83 N +ATOM 151 CE2 TRP A 19 26.124 7.801 -9.336 1.00 35.83 C +ATOM 152 CE3 TRP A 19 27.592 6.934 -7.615 1.00 35.83 C +ATOM 153 CZ2 TRP A 19 27.003 7.370 -10.338 1.00 35.83 C +ATOM 154 CZ3 TRP A 19 28.494 6.517 -8.612 1.00 35.83 C +ATOM 155 CH2 TRP A 19 28.198 6.728 -9.971 1.00 35.83 C +ATOM 156 N ASP A 20 28.212 9.597 -5.723 1.00 35.96 N +ATOM 157 CA ASP A 20 29.672 9.437 -5.705 1.00 35.96 C +ATOM 158 C ASP A 20 30.288 10.156 -4.496 1.00 35.96 C +ATOM 159 O ASP A 20 30.921 11.211 -4.577 1.00 35.96 O +ATOM 160 CB ASP A 20 30.312 9.858 -7.036 1.00 35.96 C +ATOM 161 CG ASP A 20 31.810 9.517 -7.097 1.00 35.96 C +ATOM 162 OD1 ASP A 20 32.330 8.896 -6.136 1.00 35.96 O +ATOM 163 OD2 ASP A 20 32.431 9.845 -8.135 1.00 35.96 O +ATOM 164 N TYR A 21 30.088 9.565 -3.322 1.00 35.90 N +ATOM 165 CA TYR A 21 30.616 10.082 -2.065 1.00 35.90 C +ATOM 166 C TYR A 21 32.148 10.040 -1.985 1.00 35.90 C +ATOM 167 O TYR A 21 32.722 10.673 -1.099 1.00 35.90 O +ATOM 168 CB TYR A 21 29.990 9.308 -0.904 1.00 35.90 C +ATOM 169 CG TYR A 21 30.177 7.804 -0.976 1.00 35.90 C +ATOM 170 CD1 TYR A 21 29.130 7.003 -1.461 1.00 35.90 C +ATOM 171 CD2 TYR A 21 31.380 7.202 -0.564 1.00 35.90 C +ATOM 172 CE1 TYR A 21 29.259 5.606 -1.510 1.00 35.90 C +ATOM 173 CE2 TYR A 21 31.510 5.800 -0.612 1.00 35.90 C +ATOM 174 CZ TYR A 21 30.447 4.997 -1.072 1.00 35.90 C +ATOM 175 OH TYR A 21 30.521 3.640 -1.012 1.00 35.90 O +ATOM 176 N GLU A 22 32.851 9.319 -2.862 1.00 36.57 N +ATOM 177 CA GLU A 22 34.320 9.317 -2.853 1.00 36.57 C +ATOM 178 C GLU A 22 34.879 10.661 -3.331 1.00 36.57 C +ATOM 179 O GLU A 22 35.879 11.141 -2.776 1.00 36.57 O +ATOM 180 CB GLU A 22 34.883 8.137 -3.659 1.00 36.57 C +ATOM 181 CG GLU A 22 34.436 6.811 -3.033 1.00 36.57 C +ATOM 182 CD GLU A 22 35.255 5.605 -3.500 1.00 36.57 C +ATOM 183 OE1 GLU A 22 34.636 4.586 -3.886 1.00 36.57 O +ATOM 184 OE2 GLU A 22 36.469 5.562 -3.211 1.00 36.57 O +ATOM 185 N ARG A 75 30.955 8.670 2.933 1.00 35.74 N +ATOM 186 CA ARG A 75 29.664 8.016 2.671 1.00 35.74 C +ATOM 187 C ARG A 75 28.599 8.426 3.684 1.00 35.74 C +ATOM 188 O ARG A 75 27.521 8.834 3.275 1.00 35.74 O +ATOM 189 CB ARG A 75 29.837 6.494 2.634 1.00 35.74 C +ATOM 190 CG ARG A 75 28.545 5.805 2.168 1.00 35.74 C +ATOM 191 CD ARG A 75 28.705 4.289 2.130 1.00 35.74 C +ATOM 192 NE ARG A 75 28.872 3.741 3.487 1.00 35.74 N +ATOM 193 CZ ARG A 75 29.294 2.533 3.784 1.00 35.74 C +ATOM 194 NH1 ARG A 75 29.672 1.682 2.872 1.00 35.74 N +ATOM 195 NH2 ARG A 75 29.323 2.141 5.029 1.00 35.74 N +ATOM 196 N GLU A 76 28.903 8.329 4.979 1.00 35.86 N +ATOM 197 CA GLU A 76 27.963 8.702 6.044 1.00 35.86 C +ATOM 198 C GLU A 76 27.579 10.185 5.953 1.00 35.86 C +ATOM 199 O GLU A 76 26.395 10.505 5.964 1.00 35.86 O +ATOM 200 CB GLU A 76 28.575 8.336 7.406 1.00 35.86 C +ATOM 201 CG GLU A 76 27.625 8.618 8.583 1.00 35.86 C +ATOM 202 CD GLU A 76 28.102 8.022 9.919 1.00 35.86 C +ATOM 203 OE1 GLU A 76 27.354 8.144 10.917 1.00 35.86 O +ATOM 204 OE2 GLU A 76 29.175 7.375 9.978 1.00 35.86 O +ATOM 205 N ILE A 77 28.561 11.073 5.755 1.00 35.77 N +ATOM 206 CA ILE A 77 28.318 12.508 5.559 1.00 35.77 C +ATOM 207 C ILE A 77 27.425 12.753 4.341 1.00 35.77 C +ATOM 208 O ILE A 77 26.424 13.453 4.455 1.00 35.77 O +ATOM 209 CB ILE A 77 29.652 13.275 5.425 1.00 35.77 C +ATOM 210 CG1 ILE A 77 30.392 13.290 6.780 1.00 35.77 C +ATOM 211 CG2 ILE A 77 29.401 14.714 4.938 1.00 35.77 C +ATOM 212 CD1 ILE A 77 31.843 13.773 6.673 1.00 35.77 C +ATOM 213 N LYS A 78 27.766 12.180 3.178 1.00 35.65 N +ATOM 214 CA LYS A 78 27.020 12.423 1.938 1.00 35.65 C +ATOM 215 C LYS A 78 25.569 11.965 2.055 1.00 35.65 C +ATOM 216 O LYS A 78 24.667 12.667 1.615 1.00 35.65 O +ATOM 217 CB LYS A 78 27.714 11.724 0.758 1.00 35.65 C +ATOM 218 CG LYS A 78 27.079 12.061 -0.601 1.00 35.65 C +ATOM 219 CD LYS A 78 27.147 13.556 -0.941 1.00 35.65 C +ATOM 220 CE LYS A 78 26.456 13.771 -2.282 1.00 35.65 C +ATOM 221 NZ LYS A 78 26.463 15.189 -2.698 1.00 35.65 N +ATOM 222 N ILE A 79 25.351 10.803 2.669 1.00 35.44 N +ATOM 223 CA ILE A 79 24.013 10.272 2.938 1.00 35.44 C +ATOM 224 C ILE A 79 23.237 11.215 3.860 1.00 35.44 C +ATOM 225 O ILE A 79 22.123 11.597 3.522 1.00 35.44 O +ATOM 226 CB ILE A 79 24.130 8.836 3.492 1.00 35.44 C +ATOM 227 CG1 ILE A 79 24.463 7.897 2.311 1.00 35.44 C +ATOM 228 CG2 ILE A 79 22.863 8.395 4.241 1.00 35.44 C +ATOM 229 CD1 ILE A 79 24.749 6.445 2.703 1.00 35.44 C +ATOM 230 N LEU A 80 23.827 11.646 4.977 1.00 35.56 N +ATOM 231 CA LEU A 80 23.176 12.578 5.903 1.00 35.56 C +ATOM 232 C LEU A 80 22.853 13.928 5.249 1.00 35.56 C +ATOM 233 O LEU A 80 21.775 14.467 5.477 1.00 35.56 O +ATOM 234 CB LEU A 80 24.087 12.781 7.119 1.00 35.56 C +ATOM 235 CG LEU A 80 24.086 11.598 8.099 1.00 35.56 C +ATOM 236 CD1 LEU A 80 25.269 11.738 9.055 1.00 35.56 C +ATOM 237 CD2 LEU A 80 22.791 11.557 8.913 1.00 35.56 C +ATOM 238 N GLN A 81 23.745 14.450 4.403 1.00 35.90 N +ATOM 239 CA GLN A 81 23.496 15.667 3.625 1.00 35.90 C +ATOM 240 C GLN A 81 22.323 15.487 2.655 1.00 35.90 C +ATOM 241 O GLN A 81 21.440 16.339 2.608 1.00 35.90 O +ATOM 242 CB GLN A 81 24.762 16.063 2.850 1.00 35.90 C +ATOM 243 CG GLN A 81 25.840 16.677 3.756 1.00 35.90 C +ATOM 244 CD GLN A 81 27.123 17.012 2.998 1.00 35.90 C +ATOM 245 OE1 GLN A 81 27.468 16.437 1.973 1.00 35.90 O +ATOM 246 NE2 GLN A 81 27.895 17.960 3.485 1.00 35.90 N +ATOM 247 N ASN A 82 22.279 14.372 1.920 1.00 35.62 N +ATOM 248 CA ASN A 82 21.203 14.091 0.967 1.00 35.62 C +ATOM 249 C ASN A 82 19.850 13.838 1.657 1.00 35.62 C +ATOM 250 O ASN A 82 18.806 14.116 1.075 1.00 35.62 O +ATOM 251 CB ASN A 82 21.587 12.874 0.111 1.00 35.62 C +ATOM 252 CG ASN A 82 22.718 13.086 -0.885 1.00 35.62 C +ATOM 253 OD1 ASN A 82 23.358 14.122 -1.013 1.00 35.62 O +ATOM 254 ND2 ASN A 82 22.983 12.053 -1.648 1.00 35.62 N +ATOM 255 N LEU A 83 19.855 13.304 2.882 1.00 35.50 N +ATOM 256 CA LEU A 83 18.639 13.018 3.647 1.00 35.50 C +ATOM 257 C LEU A 83 18.195 14.169 4.560 1.00 35.50 C +ATOM 258 O LEU A 83 17.136 14.072 5.185 1.00 35.50 O +ATOM 259 CB LEU A 83 18.831 11.727 4.456 1.00 35.50 C +ATOM 260 CG LEU A 83 19.055 10.458 3.625 1.00 35.50 C +ATOM 261 CD1 LEU A 83 19.269 9.291 4.578 1.00 35.50 C +ATOM 262 CD2 LEU A 83 17.862 10.128 2.733 1.00 35.50 C +ATOM 263 N CYS A 84 18.978 15.244 4.654 1.00 36.29 N +ATOM 264 CA CYS A 84 18.672 16.383 5.510 1.00 36.29 C +ATOM 265 C CYS A 84 17.295 16.980 5.168 1.00 36.29 C +ATOM 266 O CYS A 84 16.952 17.149 3.999 1.00 36.29 O +ATOM 267 CB CYS A 84 19.800 17.412 5.386 1.00 36.29 C +ATOM 268 SG CYS A 84 19.586 18.706 6.644 1.00 36.29 S +ATOM 269 N GLY A 85 16.492 17.275 6.196 1.00 36.50 N +ATOM 270 CA GLY A 85 15.111 17.758 6.046 1.00 36.50 C +ATOM 271 C GLY A 85 14.064 16.660 5.822 1.00 36.50 C +ATOM 272 O GLY A 85 12.873 16.962 5.802 1.00 36.50 O +ATOM 273 N GLY A 86 14.479 15.397 5.697 1.00 35.62 N +ATOM 274 CA GLY A 86 13.575 14.258 5.580 1.00 35.62 C +ATOM 275 C GLY A 86 12.743 13.984 6.838 1.00 35.62 C +ATOM 276 O GLY A 86 13.222 14.214 7.954 1.00 35.62 O +ATOM 277 N PRO A 87 11.521 13.436 6.697 1.00 35.65 N +ATOM 278 CA PRO A 87 10.686 13.067 7.831 1.00 35.65 C +ATOM 279 C PRO A 87 11.392 12.038 8.714 1.00 35.65 C +ATOM 280 O PRO A 87 11.831 10.988 8.247 1.00 35.65 O +ATOM 281 CB PRO A 87 9.391 12.507 7.234 1.00 35.65 C +ATOM 282 CG PRO A 87 9.836 11.973 5.873 1.00 35.65 C +ATOM 283 CD PRO A 87 10.901 12.986 5.460 1.00 35.65 C +ATOM 284 N ASN A 88 11.472 12.341 10.009 1.00 35.50 N +ATOM 285 CA ASN A 88 12.018 11.457 11.039 1.00 35.50 C +ATOM 286 C ASN A 88 13.468 10.994 10.796 1.00 35.50 C +ATOM 287 O ASN A 88 13.916 10.028 11.410 1.00 35.50 O +ATOM 288 CB ASN A 88 11.013 10.320 11.300 1.00 35.50 C +ATOM 289 CG ASN A 88 9.660 10.844 11.732 1.00 35.50 C +ATOM 290 OD1 ASN A 88 9.554 11.880 12.375 1.00 35.50 O +ATOM 291 ND2 ASN A 88 8.605 10.159 11.373 1.00 35.50 N +ATOM 292 N ILE A 89 14.232 11.703 9.962 1.00 35.47 N +ATOM 293 CA ILE A 89 15.681 11.525 9.824 1.00 35.47 C +ATOM 294 C ILE A 89 16.374 12.451 10.820 1.00 35.47 C +ATOM 295 O ILE A 89 16.003 13.618 10.934 1.00 35.47 O +ATOM 296 CB ILE A 89 16.146 11.826 8.383 1.00 35.47 C +ATOM 297 CG1 ILE A 89 15.336 11.072 7.304 1.00 35.47 C +ATOM 298 CG2 ILE A 89 17.650 11.533 8.242 1.00 35.47 C +ATOM 299 CD1 ILE A 89 15.432 9.545 7.341 1.00 35.47 C +ATOM 300 N VAL A 90 17.379 11.950 11.541 1.00 35.65 N +ATOM 301 CA VAL A 90 18.182 12.805 12.425 1.00 35.65 C +ATOM 302 C VAL A 90 18.873 13.922 11.646 1.00 35.65 C +ATOM 303 O VAL A 90 19.531 13.685 10.632 1.00 35.65 O +ATOM 304 CB VAL A 90 19.173 11.978 13.248 1.00 35.65 C +ATOM 305 CG1 VAL A 90 20.307 11.330 12.438 1.00 35.65 C +ATOM 306 CG2 VAL A 90 19.780 12.804 14.388 1.00 35.65 C +ATOM 307 N LYS A 91 18.746 15.152 12.137 1.00 35.77 N +ATOM 308 CA LYS A 91 19.389 16.317 11.536 1.00 35.77 C +ATOM 309 C LYS A 91 20.886 16.330 11.852 1.00 35.77 C +ATOM 310 O LYS A 91 21.275 16.390 13.018 1.00 35.77 O +ATOM 311 CB LYS A 91 18.661 17.576 12.019 1.00 35.77 C +ATOM 312 CG LYS A 91 19.076 18.821 11.228 1.00 35.77 C +ATOM 313 CD LYS A 91 18.316 20.049 11.746 1.00 35.77 C +ATOM 314 CE LYS A 91 18.777 21.305 11.000 1.00 35.77 C +ATOM 315 NZ LYS A 91 18.194 22.535 11.592 1.00 35.77 N +ATOM 316 N LEU A 92 21.715 16.311 10.809 1.00 35.68 N +ATOM 317 CA LEU A 92 23.134 16.666 10.885 1.00 35.68 C +ATOM 318 C LEU A 92 23.233 18.198 10.933 1.00 35.68 C +ATOM 319 O LEU A 92 22.816 18.869 9.991 1.00 35.68 O +ATOM 320 CB LEU A 92 23.864 16.058 9.671 1.00 35.68 C +ATOM 321 CG LEU A 92 25.376 16.352 9.612 1.00 35.68 C +ATOM 322 CD1 LEU A 92 26.157 15.614 10.699 1.00 35.68 C +ATOM 323 CD2 LEU A 92 25.937 15.891 8.265 1.00 35.68 C +ATOM 324 N LEU A 93 23.726 18.735 12.046 1.00 36.09 N +ATOM 325 CA LEU A 93 23.887 20.170 12.276 1.00 36.09 C +ATOM 326 C LEU A 93 25.213 20.681 11.716 1.00 36.09 C +ATOM 327 O LEU A 93 25.237 21.740 11.101 1.00 36.09 O +ATOM 328 CB LEU A 93 23.802 20.472 13.783 1.00 36.09 C +ATOM 329 CG LEU A 93 22.477 20.091 14.463 1.00 36.09 C +ATOM 330 CD1 LEU A 93 22.551 20.447 15.946 1.00 36.09 C +ATOM 331 CD2 LEU A 93 21.274 20.816 13.854 1.00 36.09 C +ATOM 332 N ASN A 113 15.623 10.980 24.168 1.00 37.04 N +ATOM 333 CA ASN A 113 15.307 9.573 23.966 1.00 37.04 C +ATOM 334 C ASN A 113 14.983 8.836 25.275 1.00 37.04 C +ATOM 335 O ASN A 113 15.781 8.845 26.213 1.00 37.04 O +ATOM 336 CB ASN A 113 16.473 8.908 23.225 1.00 37.04 C +ATOM 337 CG ASN A 113 16.191 7.434 23.028 1.00 37.04 C +ATOM 338 OD1 ASN A 113 15.313 7.054 22.278 1.00 37.04 O +ATOM 339 ND2 ASN A 113 16.859 6.560 23.741 1.00 37.04 N +ATOM 340 N THR A 114 13.891 8.070 25.277 1.00 37.75 N +ATOM 341 CA THR A 114 13.651 7.024 26.284 1.00 37.75 C +ATOM 342 C THR A 114 14.257 5.699 25.811 1.00 37.75 C +ATOM 343 O THR A 114 14.006 5.272 24.685 1.00 37.75 O +ATOM 344 CB THR A 114 12.153 6.864 26.564 1.00 37.75 C +ATOM 345 OG1 THR A 114 11.600 8.119 26.883 1.00 37.75 O +ATOM 346 CG2 THR A 114 11.876 5.952 27.761 1.00 37.75 C +ATOM 347 N ASP A 115 15.059 5.026 26.645 1.00 37.71 N +ATOM 348 CA ASP A 115 15.652 3.729 26.281 1.00 37.71 C +ATOM 349 C ASP A 115 14.563 2.707 25.910 1.00 37.71 C +ATOM 350 O ASP A 115 13.588 2.522 26.642 1.00 37.71 O +ATOM 351 CB ASP A 115 16.551 3.188 27.403 1.00 37.71 C +ATOM 352 CG ASP A 115 17.265 1.903 26.961 1.00 37.71 C +ATOM 353 OD1 ASP A 115 16.649 0.820 27.063 1.00 37.71 O +ATOM 354 OD2 ASP A 115 18.402 2.002 26.455 1.00 37.71 O +ATOM 355 N PHE A 116 14.729 2.025 24.775 1.00 37.27 N +ATOM 356 CA PHE A 116 13.706 1.126 24.239 1.00 37.27 C +ATOM 357 C PHE A 116 13.384 -0.052 25.171 1.00 37.27 C +ATOM 358 O PHE A 116 12.264 -0.555 25.132 1.00 37.27 O +ATOM 359 CB PHE A 116 14.125 0.618 22.853 1.00 37.27 C +ATOM 360 CG PHE A 116 15.173 -0.474 22.894 1.00 37.27 C +ATOM 361 CD1 PHE A 116 16.531 -0.139 23.017 1.00 37.27 C +ATOM 362 CD2 PHE A 116 14.783 -1.827 22.883 1.00 37.27 C +ATOM 363 CE1 PHE A 116 17.494 -1.150 23.148 1.00 37.27 C +ATOM 364 CE2 PHE A 116 15.753 -2.839 22.984 1.00 37.27 C +ATOM 365 CZ PHE A 116 17.108 -2.502 23.120 1.00 37.27 C +ATOM 366 N LYS A 117 14.317 -0.487 26.034 1.00 37.23 N +ATOM 367 CA LYS A 117 14.057 -1.559 27.010 1.00 37.23 C +ATOM 368 C LYS A 117 13.106 -1.111 28.113 1.00 37.23 C +ATOM 369 O LYS A 117 12.453 -1.956 28.716 1.00 37.23 O +ATOM 370 CB LYS A 117 15.354 -2.038 27.665 1.00 37.23 C +ATOM 371 CG LYS A 117 16.368 -2.598 26.667 1.00 37.23 C +ATOM 372 CD LYS A 117 17.615 -3.027 27.437 1.00 37.23 C +ATOM 373 CE LYS A 117 18.634 -3.629 26.479 1.00 37.23 C +ATOM 374 NZ LYS A 117 19.876 -3.963 27.212 1.00 37.23 N +ATOM 375 N VAL A 118 13.047 0.193 28.377 1.00 36.43 N +ATOM 376 CA VAL A 118 12.092 0.806 29.305 1.00 36.43 C +ATOM 377 C VAL A 118 10.798 1.146 28.572 1.00 36.43 C +ATOM 378 O VAL A 118 9.722 0.854 29.081 1.00 36.43 O +ATOM 379 CB VAL A 118 12.698 2.054 29.976 1.00 36.43 C +ATOM 380 CG1 VAL A 118 11.728 2.677 30.987 1.00 36.43 C +ATOM 381 CG2 VAL A 118 13.996 1.707 30.720 1.00 36.43 C +ATOM 382 N LEU A 119 10.896 1.711 27.366 1.00 36.22 N +ATOM 383 CA LEU A 119 9.743 2.194 26.608 1.00 36.22 C +ATOM 384 C LEU A 119 8.887 1.066 26.014 1.00 36.22 C +ATOM 385 O LEU A 119 7.672 1.086 26.153 1.00 36.22 O +ATOM 386 CB LEU A 119 10.249 3.159 25.521 1.00 36.22 C +ATOM 387 CG LEU A 119 9.132 3.783 24.667 1.00 36.22 C +ATOM 388 CD1 LEU A 119 8.177 4.644 25.491 1.00 36.22 C +ATOM 389 CD2 LEU A 119 9.745 4.652 23.571 1.00 36.22 C +ATOM 390 N TYR A 120 9.472 0.077 25.335 1.00 36.39 N +ATOM 391 CA TYR A 120 8.685 -0.917 24.584 1.00 36.39 C +ATOM 392 C TYR A 120 7.710 -1.725 25.457 1.00 36.39 C +ATOM 393 O TYR A 120 6.583 -1.944 25.011 1.00 36.39 O +ATOM 394 CB TYR A 120 9.588 -1.844 23.758 1.00 36.39 C +ATOM 395 CG TYR A 120 10.259 -1.242 22.534 1.00 36.39 C +ATOM 396 CD1 TYR A 120 10.203 0.140 22.239 1.00 36.39 C +ATOM 397 CD2 TYR A 120 10.942 -2.110 21.660 1.00 36.39 C +ATOM 398 CE1 TYR A 120 10.844 0.650 21.098 1.00 36.39 C +ATOM 399 CE2 TYR A 120 11.574 -1.605 20.511 1.00 36.39 C +ATOM 400 CZ TYR A 120 11.540 -0.221 20.236 1.00 36.39 C +ATOM 401 OH TYR A 120 12.149 0.273 19.126 1.00 36.39 O +ATOM 402 N PRO A 121 8.057 -2.124 26.698 1.00 36.93 N +ATOM 403 CA PRO A 121 7.101 -2.773 27.591 1.00 36.93 C +ATOM 404 C PRO A 121 5.876 -1.914 27.948 1.00 36.93 C +ATOM 405 O PRO A 121 4.823 -2.485 28.246 1.00 36.93 O +ATOM 406 CB PRO A 121 7.900 -3.135 28.847 1.00 36.93 C +ATOM 407 CG PRO A 121 9.330 -3.277 28.332 1.00 36.93 C +ATOM 408 CD PRO A 121 9.395 -2.180 27.277 1.00 36.93 C +ATOM 409 N THR A 122 5.991 -0.579 27.917 1.00 36.46 N +ATOM 410 CA THR A 122 4.912 0.351 28.299 1.00 36.46 C +ATOM 411 C THR A 122 4.020 0.769 27.134 1.00 36.46 C +ATOM 412 O THR A 122 2.952 1.328 27.374 1.00 36.46 O +ATOM 413 CB THR A 122 5.440 1.616 29.004 1.00 36.46 C +ATOM 414 OG1 THR A 122 6.119 2.480 28.127 1.00 36.46 O +ATOM 415 CG2 THR A 122 6.396 1.299 30.152 1.00 36.46 C +ATOM 416 N LEU A 123 4.424 0.500 25.887 1.00 35.86 N +ATOM 417 CA LEU A 123 3.667 0.914 24.708 1.00 35.86 C +ATOM 418 C LEU A 123 2.305 0.213 24.624 1.00 35.86 C +ATOM 419 O LEU A 123 2.182 -1.001 24.811 1.00 35.86 O +ATOM 420 CB LEU A 123 4.473 0.675 23.420 1.00 35.86 C +ATOM 421 CG LEU A 123 5.776 1.483 23.289 1.00 35.86 C +ATOM 422 CD1 LEU A 123 6.374 1.270 21.898 1.00 35.86 C +ATOM 423 CD2 LEU A 123 5.592 2.982 23.525 1.00 35.86 C +ATOM 424 N THR A 124 1.284 0.995 24.278 1.00 35.77 N +ATOM 425 CA THR A 124 -0.032 0.481 23.882 1.00 35.77 C +ATOM 426 C THR A 124 -0.006 -0.030 22.435 1.00 35.77 C +ATOM 427 O THR A 124 0.914 0.287 21.677 1.00 35.77 O +ATOM 428 CB THR A 124 -1.109 1.564 24.053 1.00 35.77 C +ATOM 429 OG1 THR A 124 -0.922 2.614 23.133 1.00 35.77 O +ATOM 430 CG2 THR A 124 -1.146 2.159 25.461 1.00 35.77 C +ATOM 431 N ASP A 125 -1.035 -0.770 21.999 1.00 35.65 N +ATOM 432 CA ASP A 125 -1.176 -1.155 20.580 1.00 35.65 C +ATOM 433 C ASP A 125 -1.157 0.080 19.661 1.00 35.65 C +ATOM 434 O ASP A 125 -0.461 0.102 18.646 1.00 35.65 O +ATOM 435 CB ASP A 125 -2.470 -1.964 20.376 1.00 35.65 C +ATOM 436 CG ASP A 125 -2.737 -2.312 18.898 1.00 35.65 C +ATOM 437 OD1 ASP A 125 -1.817 -2.737 18.166 1.00 35.65 O +ATOM 438 OD2 ASP A 125 -3.877 -2.102 18.421 1.00 35.65 O +ATOM 439 N TYR A 126 -1.845 1.159 20.055 1.00 35.65 N +ATOM 440 CA TYR A 126 -1.840 2.399 19.282 1.00 35.65 C +ATOM 441 C TYR A 126 -0.447 3.040 19.210 1.00 35.65 C +ATOM 442 O TYR A 126 -0.060 3.544 18.155 1.00 35.65 O +ATOM 443 CB TYR A 126 -2.867 3.385 19.848 1.00 35.65 C +ATOM 444 CG TYR A 126 -2.976 4.638 19.001 1.00 35.65 C +ATOM 445 CD1 TYR A 126 -2.218 5.781 19.329 1.00 35.65 C +ATOM 446 CD2 TYR A 126 -3.783 4.634 17.845 1.00 35.65 C +ATOM 447 CE1 TYR A 126 -2.261 6.918 18.500 1.00 35.65 C +ATOM 448 CE2 TYR A 126 -3.831 5.771 17.016 1.00 35.65 C +ATOM 449 CZ TYR A 126 -3.067 6.910 17.341 1.00 35.65 C +ATOM 450 OH TYR A 126 -3.103 7.999 16.529 1.00 35.65 O +ATOM 451 N ASP A 127 0.336 3.003 20.291 1.00 35.65 N +ATOM 452 CA ASP A 127 1.700 3.532 20.260 1.00 35.65 C +ATOM 453 C ASP A 127 2.617 2.724 19.347 1.00 35.65 C +ATOM 454 O ASP A 127 3.393 3.321 18.600 1.00 35.65 O +ATOM 455 CB ASP A 127 2.316 3.596 21.653 1.00 35.65 C +ATOM 456 CG ASP A 127 1.583 4.570 22.556 1.00 35.65 C +ATOM 457 OD1 ASP A 127 1.397 5.731 22.109 1.00 35.65 O +ATOM 458 OD2 ASP A 127 1.205 4.131 23.663 1.00 35.65 O +ATOM 459 N ILE A 128 2.505 1.391 19.348 1.00 35.53 N +ATOM 460 CA ILE A 128 3.262 0.537 18.423 1.00 35.53 C +ATOM 461 C ILE A 128 2.889 0.882 16.978 1.00 35.53 C +ATOM 462 O ILE A 128 3.780 1.153 16.170 1.00 35.53 O +ATOM 463 CB ILE A 128 3.048 -0.962 18.735 1.00 35.53 C +ATOM 464 CG1 ILE A 128 3.608 -1.307 20.133 1.00 35.53 C +ATOM 465 CG2 ILE A 128 3.733 -1.834 17.666 1.00 35.53 C +ATOM 466 CD1 ILE A 128 3.256 -2.723 20.603 1.00 35.53 C +ATOM 467 N ARG A 129 1.588 0.956 16.653 1.00 35.50 N +ATOM 468 CA ARG A 129 1.115 1.385 15.323 1.00 35.50 C +ATOM 469 C ARG A 129 1.712 2.734 14.936 1.00 35.50 C +ATOM 470 O ARG A 129 2.239 2.873 13.836 1.00 35.50 O +ATOM 471 CB ARG A 129 -0.420 1.484 15.290 1.00 35.50 C +ATOM 472 CG ARG A 129 -1.113 0.124 15.359 1.00 35.50 C +ATOM 473 CD ARG A 129 -2.629 0.267 15.494 1.00 35.50 C +ATOM 474 NE ARG A 129 -3.244 -1.010 15.882 1.00 35.50 N +ATOM 475 CZ ARG A 129 -4.016 -1.823 15.195 1.00 35.50 C +ATOM 476 NH1 ARG A 129 -4.303 -1.659 13.937 1.00 35.50 N +ATOM 477 NH2 ARG A 129 -4.520 -2.847 15.812 1.00 35.50 N +ATOM 478 N TYR A 130 1.647 3.712 15.837 1.00 35.56 N +ATOM 479 CA TYR A 130 2.142 5.063 15.601 1.00 35.56 C +ATOM 480 C TYR A 130 3.651 5.085 15.340 1.00 35.56 C +ATOM 481 O TYR A 130 4.080 5.602 14.315 1.00 35.56 O +ATOM 482 CB TYR A 130 1.761 5.951 16.791 1.00 35.56 C +ATOM 483 CG TYR A 130 2.272 7.372 16.677 1.00 35.56 C +ATOM 484 CD1 TYR A 130 3.496 7.732 17.274 1.00 35.56 C +ATOM 485 CD2 TYR A 130 1.532 8.327 15.959 1.00 35.56 C +ATOM 486 CE1 TYR A 130 3.983 9.047 17.147 1.00 35.56 C +ATOM 487 CE2 TYR A 130 2.014 9.644 15.835 1.00 35.56 C +ATOM 488 CZ TYR A 130 3.242 10.007 16.426 1.00 35.56 C +ATOM 489 OH TYR A 130 3.666 11.296 16.349 1.00 35.56 O +ATOM 490 N TYR A 131 4.469 4.497 16.215 1.00 35.47 N +ATOM 491 CA TYR A 131 5.925 4.561 16.063 1.00 35.47 C +ATOM 492 C TYR A 131 6.449 3.732 14.896 1.00 35.47 C +ATOM 493 O TYR A 131 7.399 4.154 14.239 1.00 35.47 O +ATOM 494 CB TYR A 131 6.621 4.166 17.367 1.00 35.47 C +ATOM 495 CG TYR A 131 6.487 5.216 18.447 1.00 35.47 C +ATOM 496 CD1 TYR A 131 6.941 6.527 18.199 1.00 35.47 C +ATOM 497 CD2 TYR A 131 5.904 4.890 19.686 1.00 35.47 C +ATOM 498 CE1 TYR A 131 6.789 7.521 19.176 1.00 35.47 C +ATOM 499 CE2 TYR A 131 5.744 5.884 20.666 1.00 35.47 C +ATOM 500 CZ TYR A 131 6.183 7.199 20.403 1.00 35.47 C +ATOM 501 OH TYR A 131 6.026 8.167 21.329 1.00 35.47 O +ATOM 502 N ILE A 132 5.825 2.594 14.584 1.00 35.44 N +ATOM 503 CA ILE A 132 6.165 1.854 13.367 1.00 35.44 C +ATOM 504 C ILE A 132 5.794 2.670 12.126 1.00 35.44 C +ATOM 505 O ILE A 132 6.583 2.719 11.186 1.00 35.44 O +ATOM 506 CB ILE A 132 5.515 0.457 13.372 1.00 35.44 C +ATOM 507 CG1 ILE A 132 6.069 -0.442 14.501 1.00 35.44 C +ATOM 508 CG2 ILE A 132 5.707 -0.240 12.015 1.00 35.44 C +ATOM 509 CD1 ILE A 132 7.548 -0.845 14.380 1.00 35.44 C +ATOM 510 N TYR A 133 4.652 3.361 12.122 1.00 35.44 N +ATOM 511 CA TYR A 133 4.272 4.248 11.022 1.00 35.44 C +ATOM 512 C TYR A 133 5.251 5.422 10.852 1.00 35.44 C +ATOM 513 O TYR A 133 5.688 5.708 9.738 1.00 35.44 O +ATOM 514 CB TYR A 133 2.840 4.733 11.250 1.00 35.44 C +ATOM 515 CG TYR A 133 2.251 5.452 10.063 1.00 35.44 C +ATOM 516 CD1 TYR A 133 2.187 6.857 10.054 1.00 35.44 C +ATOM 517 CD2 TYR A 133 1.760 4.712 8.969 1.00 35.44 C +ATOM 518 CE1 TYR A 133 1.611 7.517 8.958 1.00 35.44 C +ATOM 519 CE2 TYR A 133 1.198 5.382 7.866 1.00 35.44 C +ATOM 520 CZ TYR A 133 1.127 6.788 7.858 1.00 35.44 C +ATOM 521 OH TYR A 133 0.590 7.449 6.801 1.00 35.44 O +ATOM 522 N GLU A 134 5.679 6.050 11.950 1.00 35.47 N +ATOM 523 CA GLU A 134 6.691 7.113 11.923 1.00 35.47 C +ATOM 524 C GLU A 134 8.056 6.618 11.422 1.00 35.47 C +ATOM 525 O GLU A 134 8.739 7.332 10.683 1.00 35.47 O +ATOM 526 CB GLU A 134 6.849 7.734 13.323 1.00 35.47 C +ATOM 527 CG GLU A 134 5.673 8.614 13.768 1.00 35.47 C +ATOM 528 CD GLU A 134 5.392 9.711 12.740 1.00 35.47 C +ATOM 529 OE1 GLU A 134 4.442 9.550 11.949 1.00 35.47 O +ATOM 530 OE2 GLU A 134 6.200 10.658 12.632 1.00 35.47 O +ATOM 531 N LEU A 135 8.451 5.391 11.781 1.00 35.41 N +ATOM 532 CA LEU A 135 9.660 4.758 11.256 1.00 35.41 C +ATOM 533 C LEU A 135 9.528 4.441 9.761 1.00 35.41 C +ATOM 534 O LEU A 135 10.476 4.655 9.008 1.00 35.41 O +ATOM 535 CB LEU A 135 9.971 3.501 12.087 1.00 35.41 C +ATOM 536 CG LEU A 135 11.233 2.739 11.637 1.00 35.41 C +ATOM 537 CD1 LEU A 135 12.497 3.598 11.701 1.00 35.41 C +ATOM 538 CD2 LEU A 135 11.432 1.520 12.541 1.00 35.41 C +ATOM 539 N LEU A 136 8.356 3.985 9.310 1.00 35.38 N +ATOM 540 CA LEU A 136 8.095 3.736 7.893 1.00 35.38 C +ATOM 541 C LEU A 136 8.225 5.009 7.051 1.00 35.38 C +ATOM 542 O LEU A 136 8.784 4.928 5.965 1.00 35.38 O +ATOM 543 CB LEU A 136 6.706 3.115 7.702 1.00 35.38 C +ATOM 544 CG LEU A 136 6.590 1.636 8.091 1.00 35.38 C +ATOM 545 CD1 LEU A 136 5.119 1.224 8.139 1.00 35.38 C +ATOM 546 CD2 LEU A 136 7.295 0.749 7.068 1.00 35.38 C +ATOM 547 N LYS A 137 7.816 6.181 7.558 1.00 35.41 N +ATOM 548 CA LYS A 137 8.055 7.468 6.872 1.00 35.41 C +ATOM 549 C LYS A 137 9.541 7.737 6.639 1.00 35.41 C +ATOM 550 O LYS A 137 9.918 8.176 5.556 1.00 35.41 O +ATOM 551 CB LYS A 137 7.477 8.635 7.679 1.00 35.41 C +ATOM 552 CG LYS A 137 5.951 8.652 7.705 1.00 35.41 C +ATOM 553 CD LYS A 137 5.484 9.785 8.617 1.00 35.41 C +ATOM 554 CE LYS A 137 3.965 9.787 8.683 1.00 35.41 C +ATOM 555 NZ LYS A 137 3.505 10.721 9.722 1.00 35.41 N +ATOM 556 N ALA A 138 10.385 7.467 7.639 1.00 35.38 N +ATOM 557 CA ALA A 138 11.833 7.631 7.510 1.00 35.38 C +ATOM 558 C ALA A 138 12.421 6.663 6.472 1.00 35.38 C +ATOM 559 O ALA A 138 13.267 7.060 5.669 1.00 35.38 O +ATOM 560 CB ALA A 138 12.496 7.425 8.879 1.00 35.38 C +ATOM 561 N LEU A 139 11.965 5.405 6.484 1.00 35.38 N +ATOM 562 CA LEU A 139 12.423 4.381 5.544 1.00 35.38 C +ATOM 563 C LEU A 139 11.965 4.669 4.116 1.00 35.38 C +ATOM 564 O LEU A 139 12.793 4.666 3.215 1.00 35.38 O +ATOM 565 CB LEU A 139 11.973 2.983 6.002 1.00 35.38 C +ATOM 566 CG LEU A 139 12.612 2.506 7.317 1.00 35.38 C +ATOM 567 CD1 LEU A 139 12.095 1.108 7.650 1.00 35.38 C +ATOM 568 CD2 LEU A 139 14.140 2.453 7.259 1.00 35.38 C +ATOM 569 N ASP A 140 10.692 4.999 3.906 1.00 35.38 N +ATOM 570 CA ASP A 140 10.180 5.354 2.581 1.00 35.38 C +ATOM 571 C ASP A 140 10.917 6.567 2.012 1.00 35.38 C +ATOM 572 O ASP A 140 11.362 6.559 0.863 1.00 35.38 O +ATOM 573 CB ASP A 140 8.675 5.636 2.655 1.00 35.38 C +ATOM 574 CG ASP A 140 8.091 5.720 1.246 1.00 35.38 C +ATOM 575 OD1 ASP A 140 8.214 4.695 0.535 1.00 35.38 O +ATOM 576 OD2 ASP A 140 7.524 6.784 0.912 1.00 35.38 O +ATOM 577 N TYR A 141 11.160 7.578 2.854 1.00 35.38 N +ATOM 578 CA TYR A 141 11.934 8.737 2.446 1.00 35.38 C +ATOM 579 C TYR A 141 13.352 8.351 2.018 1.00 35.38 C +ATOM 580 O TYR A 141 13.751 8.685 0.902 1.00 35.38 O +ATOM 581 CB TYR A 141 11.964 9.788 3.553 1.00 35.38 C +ATOM 582 CG TYR A 141 12.732 11.019 3.130 1.00 35.38 C +ATOM 583 CD1 TYR A 141 14.029 11.251 3.623 1.00 35.38 C +ATOM 584 CD2 TYR A 141 12.148 11.924 2.224 1.00 35.38 C +ATOM 585 CE1 TYR A 141 14.740 12.397 3.215 1.00 35.38 C +ATOM 586 CE2 TYR A 141 12.847 13.081 1.835 1.00 35.38 C +ATOM 587 CZ TYR A 141 14.141 13.327 2.339 1.00 35.38 C +ATOM 588 OH TYR A 141 14.783 14.471 1.986 1.00 35.38 O +ATOM 589 N CYS A 142 14.117 7.627 2.843 1.00 35.38 N +ATOM 590 CA CYS A 142 15.494 7.293 2.476 1.00 35.38 C +ATOM 591 C CYS A 142 15.572 6.329 1.282 1.00 35.38 C +ATOM 592 O CYS A 142 16.426 6.519 0.410 1.00 35.38 O +ATOM 593 CB CYS A 142 16.288 6.835 3.706 1.00 35.38 C +ATOM 594 SG CYS A 142 15.844 5.157 4.237 1.00 35.38 S +ATOM 595 N HIS A 143 14.634 5.383 1.166 1.00 35.38 N +ATOM 596 CA HIS A 143 14.504 4.493 0.010 1.00 35.38 C +ATOM 597 C HIS A 143 14.210 5.292 -1.264 1.00 35.38 C +ATOM 598 O HIS A 143 14.892 5.096 -2.272 1.00 35.38 O +ATOM 599 CB HIS A 143 13.417 3.434 0.277 1.00 35.38 C +ATOM 600 CG HIS A 143 13.697 2.456 1.403 1.00 35.38 C +ATOM 601 ND1 HIS A 143 12.816 1.457 1.832 1.00 35.38 N +ATOM 602 CD2 HIS A 143 14.845 2.357 2.137 1.00 35.38 C +ATOM 603 CE1 HIS A 143 13.452 0.800 2.815 1.00 35.38 C +ATOM 604 NE2 HIS A 143 14.680 1.310 3.010 1.00 35.38 N +ATOM 605 N SER A 144 13.297 6.270 -1.202 1.00 35.44 N +ATOM 606 CA SER A 144 12.986 7.177 -2.320 1.00 35.44 C +ATOM 607 C SER A 144 14.157 8.085 -2.714 1.00 35.44 C +ATOM 608 O SER A 144 14.264 8.499 -3.867 1.00 35.44 O +ATOM 609 CB SER A 144 11.742 8.024 -2.016 1.00 35.44 C +ATOM 610 OG SER A 144 12.060 9.135 -1.194 1.00 35.44 O +ATOM 611 N GLN A 145 15.076 8.348 -1.779 1.00 35.44 N +ATOM 612 CA GLN A 145 16.344 9.040 -2.018 1.00 35.44 C +ATOM 613 C GLN A 145 17.470 8.102 -2.487 1.00 35.44 C +ATOM 614 O GLN A 145 18.623 8.529 -2.595 1.00 35.44 O +ATOM 615 CB GLN A 145 16.744 9.849 -0.767 1.00 35.44 C +ATOM 616 CG GLN A 145 15.815 11.023 -0.414 1.00 35.44 C +ATOM 617 CD GLN A 145 15.287 11.791 -1.621 1.00 35.44 C +ATOM 618 OE1 GLN A 145 16.012 12.401 -2.398 1.00 35.44 O +ATOM 619 NE2 GLN A 145 13.997 11.736 -1.868 1.00 35.44 N +ATOM 620 N GLY A 146 17.159 6.838 -2.789 1.00 35.47 N +ATOM 621 CA GLY A 146 18.114 5.865 -3.303 1.00 35.47 C +ATOM 622 C GLY A 146 19.057 5.301 -2.239 1.00 35.47 C +ATOM 623 O GLY A 146 20.144 4.842 -2.589 1.00 35.47 O +ATOM 624 N ILE A 147 18.696 5.344 -0.957 1.00 35.38 N +ATOM 625 CA ILE A 147 19.568 4.972 0.164 1.00 35.38 C +ATOM 626 C ILE A 147 18.930 3.842 0.977 1.00 35.38 C +ATOM 627 O ILE A 147 17.765 3.913 1.342 1.00 35.38 O +ATOM 628 CB ILE A 147 19.880 6.225 1.016 1.00 35.38 C +ATOM 629 CG1 ILE A 147 20.741 7.229 0.208 1.00 35.38 C +ATOM 630 CG2 ILE A 147 20.614 5.842 2.316 1.00 35.38 C +ATOM 631 CD1 ILE A 147 20.634 8.675 0.704 1.00 35.38 C +ATOM 632 N MET A 148 19.713 2.811 1.295 1.00 35.41 N +ATOM 633 CA MET A 148 19.353 1.760 2.257 1.00 35.41 C +ATOM 634 C MET A 148 20.010 2.064 3.603 1.00 35.41 C +ATOM 635 O MET A 148 21.200 2.396 3.627 1.00 35.41 O +ATOM 636 CB MET A 148 19.850 0.396 1.764 1.00 35.41 C +ATOM 637 CG MET A 148 19.358 0.016 0.369 1.00 35.41 C +ATOM 638 SD MET A 148 20.074 -1.538 -0.217 1.00 35.41 S +ATOM 639 CE MET A 148 21.718 -1.004 -0.742 1.00 35.41 C +ATOM 640 N HIS A 149 19.300 1.897 4.719 1.00 35.41 N +ATOM 641 CA HIS A 149 19.846 2.158 6.055 1.00 35.41 C +ATOM 642 C HIS A 149 20.807 1.044 6.515 1.00 35.41 C +ATOM 643 O HIS A 149 21.914 1.312 6.984 1.00 35.41 O +ATOM 644 CB HIS A 149 18.690 2.368 7.040 1.00 35.41 C +ATOM 645 CG HIS A 149 19.177 2.807 8.395 1.00 35.41 C +ATOM 646 ND1 HIS A 149 19.522 1.999 9.456 1.00 35.41 N +ATOM 647 CD2 HIS A 149 19.429 4.094 8.775 1.00 35.41 C +ATOM 648 CE1 HIS A 149 20.009 2.785 10.433 1.00 35.41 C +ATOM 649 NE2 HIS A 149 19.960 4.067 10.059 1.00 35.41 N +ATOM 650 N ARG A 150 20.421 -0.226 6.325 1.00 35.47 N +ATOM 651 CA ARG A 150 21.200 -1.455 6.597 1.00 35.47 C +ATOM 652 C ARG A 150 21.609 -1.707 8.053 1.00 35.47 C +ATOM 653 O ARG A 150 22.470 -2.552 8.312 1.00 35.47 O +ATOM 654 CB ARG A 150 22.411 -1.570 5.657 1.00 35.47 C +ATOM 655 CG ARG A 150 22.093 -1.327 4.176 1.00 35.47 C +ATOM 656 CD ARG A 150 23.246 -1.813 3.296 1.00 35.47 C +ATOM 657 NE ARG A 150 24.503 -1.111 3.609 1.00 35.47 N +ATOM 658 CZ ARG A 150 25.672 -1.315 3.046 1.00 35.47 C +ATOM 659 NH1 ARG A 150 25.884 -2.296 2.218 1.00 35.47 N +ATOM 660 NH2 ARG A 150 26.659 -0.506 3.308 1.00 35.47 N +ATOM 661 N ASP A 151 21.026 -0.989 9.006 1.00 35.71 N +ATOM 662 CA ASP A 151 21.160 -1.255 10.452 1.00 35.71 C +ATOM 663 C ASP A 151 19.938 -0.757 11.242 1.00 35.71 C +ATOM 664 O ASP A 151 20.069 -0.132 12.294 1.00 35.71 O +ATOM 665 CB ASP A 151 22.489 -0.688 10.987 1.00 35.71 C +ATOM 666 CG ASP A 151 22.956 -1.306 12.321 1.00 35.71 C +ATOM 667 OD1 ASP A 151 22.691 -2.498 12.630 1.00 35.71 O +ATOM 668 OD2 ASP A 151 23.688 -0.628 13.076 1.00 35.71 O +ATOM 669 N VAL A 152 18.733 -0.967 10.700 1.00 35.44 N +ATOM 670 CA VAL A 152 17.479 -0.663 11.408 1.00 35.44 C +ATOM 671 C VAL A 152 17.340 -1.613 12.601 1.00 35.44 C +ATOM 672 O VAL A 152 17.400 -2.830 12.445 1.00 35.44 O +ATOM 673 CB VAL A 152 16.262 -0.771 10.471 1.00 35.44 C +ATOM 674 CG1 VAL A 152 14.944 -0.475 11.197 1.00 35.44 C +ATOM 675 CG2 VAL A 152 16.389 0.190 9.289 1.00 35.44 C +ATOM 676 N LYS A 153 17.223 -1.050 13.804 1.00 35.86 N +ATOM 677 CA LYS A 153 17.106 -1.771 15.082 1.00 35.86 C +ATOM 678 C LYS A 153 16.652 -0.801 16.181 1.00 35.86 C +ATOM 679 O LYS A 153 16.875 0.399 16.020 1.00 35.86 O +ATOM 680 CB LYS A 153 18.455 -2.419 15.435 1.00 35.86 C +ATOM 681 CG LYS A 153 19.543 -1.391 15.771 1.00 35.86 C +ATOM 682 CD LYS A 153 20.920 -2.042 15.735 1.00 35.86 C +ATOM 683 CE LYS A 153 21.967 -1.020 16.178 1.00 35.86 C +ATOM 684 NZ LYS A 153 23.301 -1.353 15.650 1.00 35.86 N +ATOM 685 N PRO A 154 16.155 -1.274 17.339 1.00 35.90 N +ATOM 686 CA PRO A 154 15.642 -0.403 18.400 1.00 35.90 C +ATOM 687 C PRO A 154 16.648 0.648 18.885 1.00 35.90 C +ATOM 688 O PRO A 154 16.297 1.802 19.073 1.00 35.90 O +ATOM 689 CB PRO A 154 15.236 -1.350 19.531 1.00 35.90 C +ATOM 690 CG PRO A 154 14.876 -2.638 18.798 1.00 35.90 C +ATOM 691 CD PRO A 154 15.914 -2.666 17.682 1.00 35.90 C +ATOM 692 N HIS A 155 17.932 0.283 18.993 1.00 36.75 N +ATOM 693 CA HIS A 155 18.996 1.209 19.407 1.00 36.75 C +ATOM 694 C HIS A 155 19.224 2.380 18.438 1.00 36.75 C +ATOM 695 O HIS A 155 19.819 3.373 18.840 1.00 36.75 O +ATOM 696 CB HIS A 155 20.322 0.450 19.547 1.00 36.75 C +ATOM 697 CG HIS A 155 20.407 -0.489 20.720 1.00 36.75 C +ATOM 698 ND1 HIS A 155 19.885 -1.761 20.792 1.00 36.75 N +ATOM 699 CD2 HIS A 155 21.110 -0.275 21.876 1.00 36.75 C +ATOM 700 CE1 HIS A 155 20.285 -2.305 21.952 1.00 36.75 C +ATOM 701 NE2 HIS A 155 21.042 -1.444 22.645 1.00 36.75 N +ATOM 702 N ASN A 156 18.802 2.254 17.176 1.00 36.03 N +ATOM 703 CA ASN A 156 18.959 3.294 16.158 1.00 36.03 C +ATOM 704 C ASN A 156 17.653 4.073 15.918 1.00 36.03 C +ATOM 705 O ASN A 156 17.585 4.886 15.000 1.00 36.03 O +ATOM 706 CB ASN A 156 19.549 2.677 14.875 1.00 36.03 C +ATOM 707 CG ASN A 156 21.010 2.272 14.998 1.00 36.03 C +ATOM 708 OD1 ASN A 156 21.697 2.506 15.979 1.00 36.03 O +ATOM 709 ND2 ASN A 156 21.529 1.561 14.025 1.00 36.03 N +ATOM 710 N VAL A 157 16.620 3.845 16.733 1.00 35.62 N +ATOM 711 CA VAL A 157 15.336 4.546 16.660 1.00 35.62 C +ATOM 712 C VAL A 157 15.154 5.331 17.951 1.00 35.62 C +ATOM 713 O VAL A 157 14.763 4.776 18.975 1.00 35.62 O +ATOM 714 CB VAL A 157 14.190 3.556 16.393 1.00 35.62 C +ATOM 715 CG1 VAL A 157 12.847 4.281 16.298 1.00 35.62 C +ATOM 716 CG2 VAL A 157 14.410 2.838 15.054 1.00 35.62 C +ATOM 717 N MET A 158 15.478 6.621 17.903 1.00 35.83 N +ATOM 718 CA MET A 158 15.330 7.515 19.047 1.00 35.83 C +ATOM 719 C MET A 158 13.873 7.958 19.165 1.00 35.83 C +ATOM 720 O MET A 158 13.308 8.449 18.185 1.00 35.83 O +ATOM 721 CB MET A 158 16.253 8.733 18.922 1.00 35.83 C +ATOM 722 CG MET A 158 17.746 8.395 18.887 1.00 35.83 C +ATOM 723 SD MET A 158 18.433 7.706 20.417 1.00 35.83 S +ATOM 724 CE MET A 158 18.576 5.962 19.963 1.00 35.83 C +ATOM 725 N ILE A 159 13.276 7.801 20.347 1.00 35.59 N +ATOM 726 CA ILE A 159 11.882 8.167 20.610 1.00 35.59 C +ATOM 727 C ILE A 159 11.809 9.112 21.805 1.00 35.59 C +ATOM 728 O ILE A 159 12.147 8.744 22.931 1.00 35.59 O +ATOM 729 CB ILE A 159 10.980 6.925 20.797 1.00 35.59 C +ATOM 730 CG1 ILE A 159 10.961 6.074 19.505 1.00 35.59 C +ATOM 731 CG2 ILE A 159 9.555 7.377 21.173 1.00 35.59 C +ATOM 732 CD1 ILE A 159 10.218 4.737 19.625 1.00 35.59 C +ATOM 733 N ASP A 160 11.291 10.307 21.541 1.00 36.03 N +ATOM 734 CA ASP A 160 10.755 11.201 22.559 1.00 36.03 C +ATOM 735 C ASP A 160 9.267 10.869 22.719 1.00 36.03 C +ATOM 736 O ASP A 160 8.435 11.192 21.859 1.00 36.03 O +ATOM 737 CB ASP A 160 11.016 12.651 22.144 1.00 36.03 C +ATOM 738 CG ASP A 160 10.580 13.666 23.200 1.00 36.03 C +ATOM 739 OD1 ASP A 160 9.627 13.371 23.956 1.00 36.03 O +ATOM 740 OD2 ASP A 160 11.168 14.769 23.176 1.00 36.03 O +ATOM 741 N HIS A 161 8.948 10.117 23.775 1.00 36.19 N +ATOM 742 CA HIS A 161 7.594 9.611 23.981 1.00 36.19 C +ATOM 743 C HIS A 161 6.624 10.711 24.419 1.00 36.19 C +ATOM 744 O HIS A 161 5.474 10.711 23.984 1.00 36.19 O +ATOM 745 CB HIS A 161 7.604 8.416 24.941 1.00 36.19 C +ATOM 746 CG HIS A 161 6.274 7.701 25.010 1.00 36.19 C +ATOM 747 ND1 HIS A 161 5.536 7.241 23.942 1.00 36.19 N +ATOM 748 CD2 HIS A 161 5.571 7.376 26.138 1.00 36.19 C +ATOM 749 CE1 HIS A 161 4.419 6.661 24.408 1.00 36.19 C +ATOM 750 NE2 HIS A 161 4.403 6.712 25.745 1.00 36.19 N +ATOM 751 N GLU A 162 7.112 11.681 25.193 1.00 36.86 N +ATOM 752 CA GLU A 162 6.336 12.828 25.666 1.00 36.86 C +ATOM 753 C GLU A 162 5.876 13.698 24.492 1.00 36.86 C +ATOM 754 O GLU A 162 4.690 14.001 24.359 1.00 36.86 O +ATOM 755 CB GLU A 162 7.211 13.615 26.653 1.00 36.86 C +ATOM 756 CG GLU A 162 6.468 14.791 27.299 1.00 36.86 C +ATOM 757 CD GLU A 162 7.329 15.567 28.311 1.00 36.86 C +ATOM 758 OE1 GLU A 162 6.782 16.535 28.885 1.00 36.86 O +ATOM 759 OE2 GLU A 162 8.511 15.201 28.511 1.00 36.86 O +ATOM 760 N LEU A 163 6.790 14.015 23.571 1.00 36.57 N +ATOM 761 CA LEU A 163 6.483 14.821 22.386 1.00 36.57 C +ATOM 762 C LEU A 163 5.968 13.998 21.196 1.00 36.57 C +ATOM 763 O LEU A 163 5.681 14.562 20.139 1.00 36.57 O +ATOM 764 CB LEU A 163 7.714 15.659 22.003 1.00 36.57 C +ATOM 765 CG LEU A 163 8.226 16.601 23.110 1.00 36.57 C +ATOM 766 CD1 LEU A 163 9.343 17.473 22.530 1.00 36.57 C +ATOM 767 CD2 LEU A 163 7.134 17.543 23.627 1.00 36.57 C +ATOM 768 N ARG A 164 5.880 12.669 21.334 1.00 36.67 N +ATOM 769 CA ARG A 164 5.563 11.718 20.256 1.00 36.67 C +ATOM 770 C ARG A 164 6.425 11.899 18.997 1.00 36.67 C +ATOM 771 O ARG A 164 5.921 11.774 17.877 1.00 36.67 O +ATOM 772 CB ARG A 164 4.061 11.724 19.938 1.00 36.67 C +ATOM 773 CG ARG A 164 3.180 11.381 21.142 1.00 36.67 C +ATOM 774 CD ARG A 164 1.748 11.100 20.666 1.00 36.67 C +ATOM 775 NE ARG A 164 1.655 9.832 19.904 1.00 36.67 N +ATOM 776 CZ ARG A 164 1.525 8.629 20.442 1.00 36.67 C +ATOM 777 NH1 ARG A 164 1.474 8.427 21.722 1.00 36.67 N +ATOM 778 NH2 ARG A 164 1.440 7.556 19.720 1.00 36.67 N +ATOM 779 N LYS A 165 7.723 12.163 19.166 1.00 36.33 N +ATOM 780 CA LYS A 165 8.682 12.333 18.060 1.00 36.33 C +ATOM 781 C LYS A 165 9.591 11.118 17.922 1.00 36.33 C +ATOM 782 O LYS A 165 10.018 10.537 18.916 1.00 36.33 O +ATOM 783 CB LYS A 165 9.505 13.616 18.237 1.00 36.33 C +ATOM 784 CG LYS A 165 8.655 14.883 18.074 1.00 36.33 C +ATOM 785 CD LYS A 165 9.511 16.133 18.302 1.00 36.33 C +ATOM 786 CE LYS A 165 8.652 17.393 18.152 1.00 36.33 C +ATOM 787 NZ LYS A 165 9.424 18.612 18.501 1.00 36.33 N +ATOM 788 N LEU A 166 9.916 10.774 16.678 1.00 35.56 N +ATOM 789 CA LEU A 166 10.840 9.698 16.333 1.00 35.56 C +ATOM 790 C LEU A 166 11.958 10.228 15.426 1.00 35.56 C +ATOM 791 O LEU A 166 11.728 11.083 14.563 1.00 35.56 O +ATOM 792 CB LEU A 166 10.026 8.535 15.736 1.00 35.56 C +ATOM 793 CG LEU A 166 10.830 7.258 15.409 1.00 35.56 C +ATOM 794 CD1 LEU A 166 9.915 6.038 15.533 1.00 35.56 C +ATOM 795 CD2 LEU A 166 11.378 7.263 13.977 1.00 35.56 C +ATOM 796 N ARG A 167 13.178 9.720 15.620 1.00 35.50 N +ATOM 797 CA ARG A 167 14.317 9.929 14.718 1.00 35.50 C +ATOM 798 C ARG A 167 15.049 8.620 14.446 1.00 35.50 C +ATOM 799 O ARG A 167 15.460 7.933 15.378 1.00 35.50 O +ATOM 800 CB ARG A 167 15.289 10.977 15.290 1.00 35.50 C +ATOM 801 CG ARG A 167 14.726 12.402 15.352 1.00 35.50 C +ATOM 802 CD ARG A 167 14.459 12.963 13.956 1.00 35.50 C +ATOM 803 NE ARG A 167 14.056 14.366 14.039 1.00 35.50 N +ATOM 804 CZ ARG A 167 12.823 14.826 14.048 1.00 35.50 C +ATOM 805 NH1 ARG A 167 11.774 14.037 14.013 1.00 35.50 N +ATOM 806 NH2 ARG A 167 12.647 16.114 14.099 1.00 35.50 N +ATOM 807 N LEU A 168 15.258 8.307 13.172 1.00 35.44 N +ATOM 808 CA LEU A 168 16.174 7.263 12.724 1.00 35.44 C +ATOM 809 C LEU A 168 17.606 7.821 12.692 1.00 35.44 C +ATOM 810 O LEU A 168 17.882 8.830 12.031 1.00 35.44 O +ATOM 811 CB LEU A 168 15.699 6.727 11.361 1.00 35.44 C +ATOM 812 CG LEU A 168 16.541 5.567 10.792 1.00 35.44 C +ATOM 813 CD1 LEU A 168 16.572 4.345 11.711 1.00 35.44 C +ATOM 814 CD2 LEU A 168 15.944 5.120 9.457 1.00 35.44 C +ATOM 815 N ILE A 169 18.506 7.164 13.425 1.00 35.80 N +ATOM 816 CA ILE A 169 19.901 7.576 13.618 1.00 35.80 C +ATOM 817 C ILE A 169 20.893 6.531 13.077 1.00 35.80 C +ATOM 818 O ILE A 169 20.519 5.428 12.709 1.00 35.80 O +ATOM 819 CB ILE A 169 20.190 7.899 15.108 1.00 35.80 C +ATOM 820 CG1 ILE A 169 20.168 6.656 16.018 1.00 35.80 C +ATOM 821 CG2 ILE A 169 19.252 8.972 15.679 1.00 35.80 C +ATOM 822 CD1 ILE A 169 20.993 6.838 17.295 1.00 35.80 C +ATOM 823 N ASP A 170 22.186 6.868 13.134 1.00 36.67 N +ATOM 824 CA ASP A 170 23.329 5.999 12.798 1.00 36.67 C +ATOM 825 C ASP A 170 23.363 5.475 11.351 1.00 36.67 C +ATOM 826 O ASP A 170 23.361 4.278 11.065 1.00 36.67 O +ATOM 827 CB ASP A 170 23.568 4.902 13.857 1.00 36.67 C +ATOM 828 CG ASP A 170 24.993 4.312 13.787 1.00 36.67 C +ATOM 829 OD1 ASP A 170 25.861 4.890 13.078 1.00 36.67 O +ATOM 830 OD2 ASP A 170 25.294 3.354 14.536 1.00 36.67 O +ATOM 831 N TRP A 171 23.549 6.412 10.424 1.00 35.86 N +ATOM 832 CA TRP A 171 23.681 6.166 8.985 1.00 35.86 C +ATOM 833 C TRP A 171 25.069 5.639 8.564 1.00 35.86 C +ATOM 834 O TRP A 171 25.387 5.565 7.377 1.00 35.86 O +ATOM 835 CB TRP A 171 23.268 7.449 8.251 1.00 35.86 C +ATOM 836 CG TRP A 171 21.838 7.829 8.494 1.00 35.86 C +ATOM 837 CD1 TRP A 171 21.358 8.506 9.567 1.00 35.86 C +ATOM 838 CD2 TRP A 171 20.667 7.450 7.711 1.00 35.86 C +ATOM 839 NE1 TRP A 171 19.982 8.560 9.514 1.00 35.86 N +ATOM 840 CE2 TRP A 171 19.507 7.909 8.400 1.00 35.86 C +ATOM 841 CE3 TRP A 171 20.461 6.725 6.515 1.00 35.86 C +ATOM 842 CZ2 TRP A 171 18.215 7.636 7.949 1.00 35.86 C +ATOM 843 CZ3 TRP A 171 19.160 6.457 6.041 1.00 35.86 C +ATOM 844 CH2 TRP A 171 18.039 6.912 6.760 1.00 35.86 C +ATOM 845 N GLY A 172 25.925 5.225 9.507 1.00 36.60 N +ATOM 846 CA GLY A 172 27.301 4.793 9.223 1.00 36.60 C +ATOM 847 C GLY A 172 27.413 3.477 8.437 1.00 36.60 C +ATOM 848 O GLY A 172 28.453 3.184 7.824 1.00 36.60 O +ATOM 849 N LEU A 173 26.354 2.660 8.435 1.00 36.12 N +ATOM 850 CA LEU A 173 26.241 1.453 7.608 1.00 36.12 C +ATOM 851 C LEU A 173 25.395 1.649 6.347 1.00 36.12 C +ATOM 852 O LEU A 173 25.437 0.770 5.483 1.00 36.12 O +ATOM 853 CB LEU A 173 25.753 0.262 8.455 1.00 36.12 C +ATOM 854 CG LEU A 173 26.745 -0.206 9.539 1.00 36.12 C +ATOM 855 CD1 LEU A 173 26.235 -1.473 10.222 1.00 36.12 C +ATOM 856 CD2 LEU A 173 28.138 -0.536 8.976 1.00 36.12 C +ATOM 857 N ALA A 174 24.737 2.796 6.187 1.00 35.56 N +ATOM 858 CA ALA A 174 23.901 3.083 5.033 1.00 35.56 C +ATOM 859 C ALA A 174 24.697 3.090 3.716 1.00 35.56 C +ATOM 860 O ALA A 174 25.932 3.191 3.707 1.00 35.56 O +ATOM 861 CB ALA A 174 23.156 4.396 5.288 1.00 35.56 C +ATOM 862 N GLU A 175 24.001 2.917 2.593 1.00 35.50 N +ATOM 863 CA GLU A 175 24.602 2.868 1.256 1.00 35.50 C +ATOM 864 C GLU A 175 23.622 3.284 0.165 1.00 35.50 C +ATOM 865 O GLU A 175 22.419 3.060 0.285 1.00 35.50 O +ATOM 866 CB GLU A 175 25.095 1.441 0.992 1.00 35.50 C +ATOM 867 CG GLU A 175 25.989 1.229 -0.236 1.00 35.50 C +ATOM 868 CD GLU A 175 27.335 1.945 -0.120 1.00 35.50 C +ATOM 869 OE1 GLU A 175 27.550 2.965 -0.806 1.00 35.50 O +ATOM 870 OE2 GLU A 175 28.202 1.451 0.644 1.00 35.50 O +ATOM 871 N PHE A 176 24.163 3.825 -0.926 1.00 35.47 N +ATOM 872 CA PHE A 176 23.409 4.100 -2.143 1.00 35.47 C +ATOM 873 C PHE A 176 23.086 2.806 -2.903 1.00 35.47 C +ATOM 874 O PHE A 176 23.973 1.992 -3.211 1.00 35.47 O +ATOM 875 CB PHE A 176 24.186 5.093 -3.013 1.00 35.47 C +ATOM 876 CG PHE A 176 24.308 6.465 -2.382 1.00 35.47 C +ATOM 877 CD1 PHE A 176 23.269 7.399 -2.534 1.00 35.47 C +ATOM 878 CD2 PHE A 176 25.445 6.801 -1.622 1.00 35.47 C +ATOM 879 CE1 PHE A 176 23.366 8.662 -1.932 1.00 35.47 C +ATOM 880 CE2 PHE A 176 25.543 8.070 -1.020 1.00 35.47 C +ATOM 881 CZ PHE A 176 24.499 8.999 -1.172 1.00 35.47 C +ATOM 882 N TYR A 177 21.808 2.623 -3.225 1.00 35.44 N +ATOM 883 CA TYR A 177 21.337 1.524 -4.057 1.00 35.44 C +ATOM 884 C TYR A 177 21.656 1.762 -5.542 1.00 35.44 C +ATOM 885 O TYR A 177 21.425 2.849 -6.066 1.00 35.44 O +ATOM 886 CB TYR A 177 19.851 1.272 -3.849 1.00 35.44 C +ATOM 887 CG TYR A 177 19.361 0.122 -4.709 1.00 35.44 C +ATOM 888 CD1 TYR A 177 18.686 0.376 -5.920 1.00 35.44 C +ATOM 889 CD2 TYR A 177 19.693 -1.201 -4.357 1.00 35.44 C +ATOM 890 CE1 TYR A 177 18.333 -0.692 -6.765 1.00 35.44 C +ATOM 891 CE2 TYR A 177 19.347 -2.266 -5.204 1.00 35.44 C +ATOM 892 CZ TYR A 177 18.651 -2.012 -6.403 1.00 35.44 C +ATOM 893 OH TYR A 177 18.319 -3.047 -7.216 1.00 35.44 O +ATOM 894 N HIS A 178 22.162 0.747 -6.231 1.00 35.65 N +ATOM 895 CA HIS A 178 22.364 0.713 -7.671 1.00 35.65 C +ATOM 896 C HIS A 178 21.964 -0.683 -8.166 1.00 35.65 C +ATOM 897 O HIS A 178 22.380 -1.678 -7.563 1.00 35.65 O +ATOM 898 CB HIS A 178 23.824 1.009 -8.032 1.00 35.65 C +ATOM 899 CG HIS A 178 24.261 2.418 -7.727 1.00 35.65 C +ATOM 900 ND1 HIS A 178 24.101 3.520 -8.536 1.00 35.65 N +ATOM 901 CD2 HIS A 178 24.894 2.842 -6.592 1.00 35.65 C +ATOM 902 CE1 HIS A 178 24.647 4.576 -7.910 1.00 35.65 C +ATOM 903 NE2 HIS A 178 25.152 4.208 -6.718 1.00 35.65 N +ATOM 904 N PRO A 179 21.188 -0.785 -9.252 1.00 35.96 N +ATOM 905 CA PRO A 179 20.735 -2.072 -9.758 1.00 35.96 C +ATOM 906 C PRO A 179 21.928 -2.911 -10.232 1.00 35.96 C +ATOM 907 O PRO A 179 22.830 -2.400 -10.895 1.00 35.96 O +ATOM 908 CB PRO A 179 19.742 -1.733 -10.874 1.00 35.96 C +ATOM 909 CG PRO A 179 20.217 -0.372 -11.380 1.00 35.96 C +ATOM 910 CD PRO A 179 20.771 0.299 -10.128 1.00 35.96 C +ATOM 911 N GLY A 180 21.942 -4.198 -9.873 1.00 36.71 N +ATOM 912 CA GLY A 180 23.017 -5.137 -10.217 1.00 36.71 C +ATOM 913 C GLY A 180 24.321 -4.961 -9.429 1.00 36.71 C +ATOM 914 O GLY A 180 25.278 -5.687 -9.680 1.00 36.71 O +ATOM 915 N LYS A 181 24.389 -4.015 -8.482 1.00 36.15 N +ATOM 916 CA LYS A 181 25.553 -3.862 -7.606 1.00 36.15 C +ATOM 917 C LYS A 181 25.530 -4.917 -6.504 1.00 36.15 C +ATOM 918 O LYS A 181 24.534 -5.067 -5.805 1.00 36.15 O +ATOM 919 CB LYS A 181 25.590 -2.435 -7.049 1.00 36.15 C +ATOM 920 CG LYS A 181 26.778 -2.179 -6.111 1.00 36.15 C +ATOM 921 CD LYS A 181 26.859 -0.683 -5.795 1.00 36.15 C +ATOM 922 CE LYS A 181 27.891 -0.393 -4.708 1.00 36.15 C +ATOM 923 NZ LYS A 181 27.871 1.052 -4.356 1.00 36.15 N +ATOM 924 N GLU A 182 26.666 -5.570 -6.300 1.00 35.83 N +ATOM 925 CA GLU A 182 26.890 -6.444 -5.151 1.00 35.83 C +ATOM 926 C GLU A 182 27.315 -5.627 -3.923 1.00 35.83 C +ATOM 927 O GLU A 182 28.161 -4.724 -3.979 1.00 35.83 O +ATOM 928 CB GLU A 182 27.929 -7.515 -5.494 1.00 35.83 C +ATOM 929 CG GLU A 182 27.392 -8.509 -6.538 1.00 35.83 C +ATOM 930 CD GLU A 182 28.435 -9.551 -6.967 1.00 35.83 C +ATOM 931 OE1 GLU A 182 28.014 -10.546 -7.597 1.00 35.83 O +ATOM 932 OE2 GLU A 182 29.643 -9.312 -6.729 1.00 35.83 O +ATOM 933 N TYR A 183 26.716 -5.949 -2.783 1.00 35.74 N +ATOM 934 CA TYR A 183 26.925 -5.286 -1.509 1.00 35.74 C +ATOM 935 C TYR A 183 27.505 -6.241 -0.474 1.00 35.74 C +ATOM 936 O TYR A 183 27.176 -7.418 -0.393 1.00 35.74 O +ATOM 937 CB TYR A 183 25.605 -4.708 -1.002 1.00 35.74 C +ATOM 938 CG TYR A 183 25.021 -3.630 -1.873 1.00 35.74 C +ATOM 939 CD1 TYR A 183 25.513 -2.316 -1.776 1.00 35.74 C +ATOM 940 CD2 TYR A 183 24.000 -3.948 -2.786 1.00 35.74 C +ATOM 941 CE1 TYR A 183 24.990 -1.316 -2.607 1.00 35.74 C +ATOM 942 CE2 TYR A 183 23.491 -2.956 -3.639 1.00 35.74 C +ATOM 943 CZ TYR A 183 23.998 -1.650 -3.546 1.00 35.74 C +ATOM 944 OH TYR A 183 23.571 -0.676 -4.365 1.00 35.74 O +ATOM 945 N ASN A 184 28.324 -5.689 0.421 1.00 36.06 N +ATOM 946 CA ASN A 184 28.844 -6.443 1.555 1.00 36.06 C +ATOM 947 C ASN A 184 27.700 -6.900 2.482 1.00 36.06 C +ATOM 948 O ASN A 184 26.942 -6.064 2.977 1.00 36.06 O +ATOM 949 CB ASN A 184 29.871 -5.552 2.273 1.00 36.06 C +ATOM 950 CG ASN A 184 30.673 -6.297 3.320 1.00 36.06 C +ATOM 951 OD1 ASN A 184 30.207 -7.158 4.039 1.00 36.06 O +ATOM 952 ND2 ASN A 184 31.938 -5.986 3.460 1.00 36.06 N +ATOM 953 N VAL A 185 27.623 -8.201 2.777 1.00 35.96 N +ATOM 954 CA VAL A 185 26.596 -8.793 3.656 1.00 35.96 C +ATOM 955 C VAL A 185 26.894 -8.663 5.158 1.00 35.96 C +ATOM 956 O VAL A 185 26.043 -8.888 6.024 1.00 35.96 O +ATOM 957 CB VAL A 185 26.343 -10.264 3.287 1.00 35.96 C +ATOM 958 CG1 VAL A 185 25.781 -10.354 1.866 1.00 35.96 C +ATOM 959 CG2 VAL A 185 27.612 -11.111 3.439 1.00 35.96 C +ATOM 960 N ARG A 186 28.107 -8.232 5.525 1.00 36.29 N +ATOM 961 CA ARG A 186 28.547 -8.004 6.915 1.00 36.29 C +ATOM 962 C ARG A 186 28.056 -6.654 7.461 1.00 36.29 C +ATOM 963 O ARG A 186 28.812 -5.910 8.086 1.00 36.29 O +ATOM 964 CB ARG A 186 30.068 -8.204 7.060 1.00 36.29 C +ATOM 965 CG ARG A 186 30.560 -9.564 6.539 1.00 36.29 C +ATOM 966 CD ARG A 186 32.066 -9.711 6.782 1.00 36.29 C +ATOM 967 NE ARG A 186 32.586 -10.966 6.207 1.00 36.29 N +ATOM 968 CZ ARG A 186 33.826 -11.422 6.265 1.00 36.29 C +ATOM 969 NH1 ARG A 186 34.775 -10.779 6.893 1.00 36.29 N +ATOM 970 NH2 ARG A 186 34.140 -12.544 5.685 1.00 36.29 N +ATOM 971 N VAL A 187 26.780 -6.359 7.231 1.00 36.26 N +ATOM 972 CA VAL A 187 26.020 -5.199 7.732 1.00 36.26 C +ATOM 973 C VAL A 187 24.902 -5.656 8.669 1.00 36.26 C +ATOM 974 O VAL A 187 24.706 -6.858 8.823 1.00 36.26 O +ATOM 975 CB VAL A 187 25.464 -4.360 6.566 1.00 36.26 C +ATOM 976 CG1 VAL A 187 26.615 -3.796 5.727 1.00 36.26 C +ATOM 977 CG2 VAL A 187 24.508 -5.155 5.669 1.00 36.26 C +ATOM 978 N ALA A 188 24.186 -4.726 9.300 1.00 36.64 N +ATOM 979 CA ALA A 188 23.172 -4.980 10.321 1.00 36.64 C +ATOM 980 C ALA A 188 23.665 -5.715 11.586 1.00 36.64 C +ATOM 981 O ALA A 188 24.657 -6.456 11.614 1.00 36.64 O +ATOM 982 CB ALA A 188 21.935 -5.639 9.686 1.00 36.64 C +ATOM 983 N SER A 189 22.927 -5.512 12.671 1.00 36.64 N +ATOM 984 CA SER A 189 23.095 -6.246 13.926 1.00 36.64 C +ATOM 985 C SER A 189 22.443 -7.631 13.822 1.00 36.64 C +ATOM 986 O SER A 189 21.391 -7.761 13.205 1.00 36.64 O +ATOM 987 CB SER A 189 22.519 -5.409 15.066 1.00 36.64 C +ATOM 988 OG SER A 189 23.186 -4.157 15.076 1.00 36.64 O +ATOM 989 N ARG A 190 23.054 -8.669 14.420 1.00 35.71 N +ATOM 990 CA ARG A 190 22.719 -10.095 14.187 1.00 35.71 C +ATOM 991 C ARG A 190 21.217 -10.393 14.155 1.00 35.71 C +ATOM 992 O ARG A 190 20.744 -10.979 13.191 1.00 35.71 O +ATOM 993 CB ARG A 190 23.405 -10.971 15.247 1.00 35.71 C +ATOM 994 CG ARG A 190 23.226 -12.479 14.974 1.00 35.71 C +ATOM 995 CD ARG A 190 23.687 -13.306 16.172 1.00 35.71 C +ATOM 996 NE ARG A 190 25.119 -13.101 16.468 1.00 35.71 N +ATOM 997 CZ ARG A 190 25.654 -13.064 17.675 1.00 35.71 C +ATOM 998 NH1 ARG A 190 24.927 -13.161 18.752 1.00 35.71 N +ATOM 999 NH2 ARG A 190 26.945 -12.945 17.818 1.00 35.71 N +ATOM 1000 N TYR A 191 20.483 -9.951 15.173 1.00 35.50 N +ATOM 1001 CA TYR A 191 19.061 -10.265 15.343 1.00 35.50 C +ATOM 1002 C TYR A 191 18.134 -9.628 14.294 1.00 35.50 C +ATOM 1003 O TYR A 191 16.985 -10.039 14.153 1.00 35.50 O +ATOM 1004 CB TYR A 191 18.644 -9.859 16.759 1.00 35.50 C +ATOM 1005 CG TYR A 191 19.550 -10.382 17.865 1.00 35.50 C +ATOM 1006 CD1 TYR A 191 19.971 -11.727 17.858 1.00 35.50 C +ATOM 1007 CD2 TYR A 191 19.958 -9.530 18.911 1.00 35.50 C +ATOM 1008 CE1 TYR A 191 20.791 -12.217 18.886 1.00 35.50 C +ATOM 1009 CE2 TYR A 191 20.769 -10.026 19.954 1.00 35.50 C +ATOM 1010 CZ TYR A 191 21.189 -11.374 19.938 1.00 35.50 C +ATOM 1011 OH TYR A 191 21.981 -11.867 20.925 1.00 35.50 O +ATOM 1012 N PHE A 192 18.648 -8.652 13.546 1.00 35.47 N +ATOM 1013 CA PHE A 192 17.936 -7.894 12.519 1.00 35.47 C +ATOM 1014 C PHE A 192 18.461 -8.203 11.109 1.00 35.47 C +ATOM 1015 O PHE A 192 17.994 -7.620 10.140 1.00 35.47 O +ATOM 1016 CB PHE A 192 18.021 -6.396 12.863 1.00 35.47 C +ATOM 1017 CG PHE A 192 17.535 -6.078 14.266 1.00 35.47 C +ATOM 1018 CD1 PHE A 192 16.161 -5.913 14.513 1.00 35.47 C +ATOM 1019 CD2 PHE A 192 18.445 -6.031 15.340 1.00 35.47 C +ATOM 1020 CE1 PHE A 192 15.696 -5.750 15.829 1.00 35.47 C +ATOM 1021 CE2 PHE A 192 17.983 -5.844 16.656 1.00 35.47 C +ATOM 1022 CZ PHE A 192 16.603 -5.722 16.901 1.00 35.47 C +ATOM 1023 N LYS A 193 19.427 -9.124 10.963 1.00 35.44 N +ATOM 1024 CA LYS A 193 19.922 -9.541 9.647 1.00 35.44 C +ATOM 1025 C LYS A 193 18.816 -10.242 8.858 1.00 35.44 C +ATOM 1026 O LYS A 193 18.253 -11.223 9.346 1.00 35.44 O +ATOM 1027 CB LYS A 193 21.129 -10.476 9.788 1.00 35.44 C +ATOM 1028 CG LYS A 193 22.395 -9.730 10.217 1.00 35.44 C +ATOM 1029 CD LYS A 193 23.551 -10.728 10.355 1.00 35.44 C +ATOM 1030 CE LYS A 193 24.884 -10.013 10.562 1.00 35.44 C +ATOM 1031 NZ LYS A 193 25.308 -9.369 9.299 1.00 35.44 N +ATOM 1032 N GLY A 194 18.562 -9.750 7.646 1.00 35.53 N +ATOM 1033 CA GLY A 194 17.739 -10.430 6.647 1.00 35.53 C +ATOM 1034 C GLY A 194 18.355 -11.760 6.206 1.00 35.53 C +ATOM 1035 O GLY A 194 19.587 -11.882 6.239 1.00 35.53 O +ATOM 1036 N PRO A 195 17.541 -12.750 5.798 1.00 35.53 N +ATOM 1037 CA PRO A 195 18.009 -14.018 5.250 1.00 35.53 C +ATOM 1038 C PRO A 195 19.067 -13.848 4.162 1.00 35.53 C +ATOM 1039 O PRO A 195 20.063 -14.559 4.203 1.00 35.53 O +ATOM 1040 CB PRO A 195 16.759 -14.719 4.710 1.00 35.53 C +ATOM 1041 CG PRO A 195 15.667 -14.199 5.641 1.00 35.53 C +ATOM 1042 CD PRO A 195 16.092 -12.758 5.886 1.00 35.53 C +ATOM 1043 N GLU A 196 18.926 -12.859 3.276 1.00 35.56 N +ATOM 1044 CA GLU A 196 19.898 -12.513 2.233 1.00 35.56 C +ATOM 1045 C GLU A 196 21.317 -12.303 2.781 1.00 35.56 C +ATOM 1046 O GLU A 196 22.282 -12.758 2.179 1.00 35.56 O +ATOM 1047 CB GLU A 196 19.416 -11.287 1.425 1.00 35.56 C +ATOM 1048 CG GLU A 196 19.232 -9.958 2.189 1.00 35.56 C +ATOM 1049 CD GLU A 196 17.835 -9.759 2.789 1.00 35.56 C +ATOM 1050 OE1 GLU A 196 17.313 -8.628 2.757 1.00 35.56 O +ATOM 1051 OE2 GLU A 196 17.257 -10.708 3.367 1.00 35.56 O +ATOM 1052 N LEU A 197 21.465 -11.715 3.974 1.00 35.62 N +ATOM 1053 CA LEU A 197 22.777 -11.514 4.592 1.00 35.62 C +ATOM 1054 C LEU A 197 23.343 -12.793 5.214 1.00 35.62 C +ATOM 1055 O LEU A 197 24.553 -12.895 5.417 1.00 35.62 O +ATOM 1056 CB LEU A 197 22.693 -10.444 5.693 1.00 35.62 C +ATOM 1057 CG LEU A 197 22.091 -9.090 5.294 1.00 35.62 C +ATOM 1058 CD1 LEU A 197 22.263 -8.140 6.483 1.00 35.62 C +ATOM 1059 CD2 LEU A 197 22.776 -8.455 4.090 1.00 35.62 C +ATOM 1060 N LEU A 198 22.465 -13.714 5.610 1.00 35.74 N +ATOM 1061 CA LEU A 198 22.811 -14.963 6.288 1.00 35.74 C +ATOM 1062 C LEU A 198 23.135 -16.084 5.295 1.00 35.74 C +ATOM 1063 O LEU A 198 23.833 -17.017 5.670 1.00 35.74 O +ATOM 1064 CB LEU A 198 21.658 -15.377 7.219 1.00 35.74 C +ATOM 1065 CG LEU A 198 21.288 -14.340 8.294 1.00 35.74 C +ATOM 1066 CD1 LEU A 198 20.074 -14.822 9.083 1.00 35.74 C +ATOM 1067 CD2 LEU A 198 22.433 -14.090 9.281 1.00 35.74 C +ATOM 1068 N VAL A 199 22.656 -15.963 4.054 1.00 36.09 N +ATOM 1069 CA VAL A 199 22.926 -16.894 2.946 1.00 36.09 C +ATOM 1070 C VAL A 199 23.871 -16.307 1.884 1.00 36.09 C +ATOM 1071 O VAL A 199 23.989 -16.855 0.799 1.00 36.09 O +ATOM 1072 CB VAL A 199 21.619 -17.444 2.335 1.00 36.09 C +ATOM 1073 CG1 VAL A 199 20.673 -17.990 3.417 1.00 36.09 C +ATOM 1074 CG2 VAL A 199 20.847 -16.420 1.492 1.00 36.09 C +ATOM 1075 N ASP A 200 24.520 -15.179 2.196 1.00 36.53 N +ATOM 1076 CA ASP A 200 25.484 -14.465 1.338 1.00 36.53 C +ATOM 1077 C ASP A 200 24.952 -14.043 -0.052 1.00 36.53 C +ATOM 1078 O ASP A 200 25.690 -14.010 -1.035 1.00 36.53 O +ATOM 1079 CB ASP A 200 26.839 -15.205 1.304 1.00 36.53 C +ATOM 1080 CG ASP A 200 28.047 -14.320 0.928 1.00 36.53 C +ATOM 1081 OD1 ASP A 200 27.972 -13.067 1.036 1.00 36.53 O +ATOM 1082 OD2 ASP A 200 29.133 -14.880 0.669 1.00 36.53 O +ATOM 1083 N LEU A 201 23.672 -13.664 -0.148 1.00 36.39 N +ATOM 1084 CA LEU A 201 23.129 -13.031 -1.354 1.00 36.39 C +ATOM 1085 C LEU A 201 23.502 -11.543 -1.366 1.00 36.39 C +ATOM 1086 O LEU A 201 23.037 -10.777 -0.521 1.00 36.39 O +ATOM 1087 CB LEU A 201 21.608 -13.243 -1.437 1.00 36.39 C +ATOM 1088 CG LEU A 201 21.003 -12.697 -2.747 1.00 36.39 C +ATOM 1089 CD1 LEU A 201 21.482 -13.487 -3.968 1.00 36.39 C +ATOM 1090 CD2 LEU A 201 19.479 -12.752 -2.708 1.00 36.39 C +ATOM 1091 N GLN A 202 24.340 -11.126 -2.317 1.00 35.96 N +ATOM 1092 CA GLN A 202 24.979 -9.802 -2.296 1.00 35.96 C +ATOM 1093 C GLN A 202 24.236 -8.724 -3.096 1.00 35.96 C +ATOM 1094 O GLN A 202 24.428 -7.541 -2.824 1.00 35.96 O +ATOM 1095 CB GLN A 202 26.442 -9.925 -2.751 1.00 35.96 C +ATOM 1096 CG GLN A 202 27.210 -10.910 -1.856 1.00 35.96 C +ATOM 1097 CD GLN A 202 28.718 -10.897 -2.039 1.00 35.96 C +ATOM 1098 OE1 GLN A 202 29.280 -10.334 -2.961 1.00 35.96 O +ATOM 1099 NE2 GLN A 202 29.442 -11.503 -1.124 1.00 35.96 N +ATOM 1100 N ASP A 203 23.367 -9.083 -4.039 1.00 36.57 N +ATOM 1101 CA ASP A 203 22.589 -8.149 -4.870 1.00 36.57 C +ATOM 1102 C ASP A 203 21.251 -7.728 -4.216 1.00 36.57 C +ATOM 1103 O ASP A 203 20.225 -7.552 -4.881 1.00 36.57 O +ATOM 1104 CB ASP A 203 22.473 -8.688 -6.314 1.00 36.57 C +ATOM 1105 CG ASP A 203 21.672 -9.990 -6.478 1.00 36.57 C +ATOM 1106 OD1 ASP A 203 21.407 -10.663 -5.461 1.00 36.57 O +ATOM 1107 OD2 ASP A 203 21.300 -10.316 -7.638 1.00 36.57 O +ATOM 1108 N TYR A 204 21.265 -7.554 -2.887 1.00 35.71 N +ATOM 1109 CA TYR A 204 20.113 -7.128 -2.086 1.00 35.71 C +ATOM 1110 C TYR A 204 19.785 -5.631 -2.234 1.00 35.71 C +ATOM 1111 O TYR A 204 20.564 -4.839 -2.770 1.00 35.71 O +ATOM 1112 CB TYR A 204 20.314 -7.534 -0.616 1.00 35.71 C +ATOM 1113 CG TYR A 204 21.444 -6.834 0.123 1.00 35.71 C +ATOM 1114 CD1 TYR A 204 22.676 -7.488 0.305 1.00 35.71 C +ATOM 1115 CD2 TYR A 204 21.264 -5.540 0.648 1.00 35.71 C +ATOM 1116 CE1 TYR A 204 23.724 -6.870 1.012 1.00 35.71 C +ATOM 1117 CE2 TYR A 204 22.316 -4.902 1.331 1.00 35.71 C +ATOM 1118 CZ TYR A 204 23.545 -5.567 1.523 1.00 35.71 C +ATOM 1119 OH TYR A 204 24.559 -4.946 2.193 1.00 35.71 O +ATOM 1120 N ASP A 205 18.612 -5.235 -1.735 1.00 35.44 N +ATOM 1121 CA ASP A 205 18.066 -3.883 -1.879 1.00 35.44 C +ATOM 1122 C ASP A 205 17.388 -3.356 -0.592 1.00 35.44 C +ATOM 1123 O ASP A 205 17.596 -3.875 0.508 1.00 35.44 O +ATOM 1124 CB ASP A 205 17.148 -3.859 -3.114 1.00 35.44 C +ATOM 1125 CG ASP A 205 15.913 -4.749 -3.000 1.00 35.44 C +ATOM 1126 OD1 ASP A 205 15.152 -4.567 -2.022 1.00 35.44 O +ATOM 1127 OD2 ASP A 205 15.702 -5.577 -3.911 1.00 35.44 O +ATOM 1128 N TYR A 206 16.594 -2.288 -0.727 1.00 35.44 N +ATOM 1129 CA TYR A 206 15.849 -1.606 0.340 1.00 35.44 C +ATOM 1130 C TYR A 206 15.012 -2.536 1.231 1.00 35.44 C +ATOM 1131 O TYR A 206 14.817 -2.263 2.417 1.00 35.44 O +ATOM 1132 CB TYR A 206 14.899 -0.601 -0.321 1.00 35.44 C +ATOM 1133 CG TYR A 206 15.491 0.249 -1.424 1.00 35.44 C +ATOM 1134 CD1 TYR A 206 16.389 1.291 -1.123 1.00 35.44 C +ATOM 1135 CD2 TYR A 206 15.114 0.004 -2.756 1.00 35.44 C +ATOM 1136 CE1 TYR A 206 16.897 2.097 -2.158 1.00 35.44 C +ATOM 1137 CE2 TYR A 206 15.631 0.797 -3.792 1.00 35.44 C +ATOM 1138 CZ TYR A 206 16.503 1.856 -3.493 1.00 35.44 C +ATOM 1139 OH TYR A 206 16.928 2.659 -4.497 1.00 35.44 O +ATOM 1140 N SER A 207 14.540 -3.655 0.682 1.00 35.41 N +ATOM 1141 CA SER A 207 13.761 -4.679 1.389 1.00 35.41 C +ATOM 1142 C SER A 207 14.496 -5.327 2.577 1.00 35.41 C +ATOM 1143 O SER A 207 13.848 -5.925 3.443 1.00 35.41 O +ATOM 1144 CB SER A 207 13.341 -5.755 0.390 1.00 35.41 C +ATOM 1145 OG SER A 207 14.493 -6.274 -0.228 1.00 35.41 O +ATOM 1146 N LEU A 208 15.821 -5.164 2.688 1.00 35.41 N +ATOM 1147 CA LEU A 208 16.593 -5.531 3.881 1.00 35.41 C +ATOM 1148 C LEU A 208 16.148 -4.742 5.125 1.00 35.41 C +ATOM 1149 O LEU A 208 16.043 -5.290 6.229 1.00 35.41 O +ATOM 1150 CB LEU A 208 18.083 -5.284 3.585 1.00 35.41 C +ATOM 1151 CG LEU A 208 19.006 -5.496 4.800 1.00 35.41 C +ATOM 1152 CD1 LEU A 208 18.980 -6.938 5.311 1.00 35.41 C +ATOM 1153 CD2 LEU A 208 20.440 -5.114 4.433 1.00 35.41 C +ATOM 1154 N ASP A 209 15.866 -3.448 4.958 1.00 35.38 N +ATOM 1155 CA ASP A 209 15.394 -2.599 6.054 1.00 35.38 C +ATOM 1156 C ASP A 209 13.985 -3.021 6.497 1.00 35.38 C +ATOM 1157 O ASP A 209 13.669 -2.983 7.685 1.00 35.38 O +ATOM 1158 CB ASP A 209 15.408 -1.124 5.632 1.00 35.38 C +ATOM 1159 CG ASP A 209 16.812 -0.557 5.388 1.00 35.38 C +ATOM 1160 OD1 ASP A 209 17.777 -0.982 6.066 1.00 35.38 O +ATOM 1161 OD2 ASP A 209 16.936 0.368 4.553 1.00 35.38 O +ATOM 1162 N MET A 210 13.168 -3.511 5.560 1.00 35.38 N +ATOM 1163 CA MET A 210 11.805 -3.985 5.826 1.00 35.38 C +ATOM 1164 C MET A 210 11.789 -5.286 6.633 1.00 35.38 C +ATOM 1165 O MET A 210 10.970 -5.435 7.539 1.00 35.38 O +ATOM 1166 CB MET A 210 11.031 -4.138 4.508 1.00 35.38 C +ATOM 1167 CG MET A 210 10.983 -2.822 3.715 1.00 35.38 C +ATOM 1168 SD MET A 210 10.326 -1.385 4.613 1.00 35.38 S +ATOM 1169 CE MET A 210 8.668 -1.991 5.023 1.00 35.38 C +ATOM 1170 N TRP A 211 12.739 -6.196 6.386 1.00 35.38 N +ATOM 1171 CA TRP A 211 12.954 -7.353 7.264 1.00 35.38 C +ATOM 1172 C TRP A 211 13.336 -6.916 8.680 1.00 35.38 C +ATOM 1173 O TRP A 211 12.764 -7.384 9.665 1.00 35.38 O +ATOM 1174 CB TRP A 211 14.060 -8.245 6.702 1.00 35.38 C +ATOM 1175 CG TRP A 211 14.426 -9.363 7.619 1.00 35.38 C +ATOM 1176 CD1 TRP A 211 15.326 -9.304 8.627 1.00 35.38 C +ATOM 1177 CD2 TRP A 211 13.854 -10.700 7.666 1.00 35.38 C +ATOM 1178 NE1 TRP A 211 15.375 -10.523 9.271 1.00 35.38 N +ATOM 1179 CE2 TRP A 211 14.452 -11.401 8.753 1.00 35.38 C +ATOM 1180 CE3 TRP A 211 12.897 -11.391 6.892 1.00 35.38 C +ATOM 1181 CZ2 TRP A 211 14.106 -12.717 9.071 1.00 35.38 C +ATOM 1182 CZ3 TRP A 211 12.558 -12.722 7.190 1.00 35.38 C +ATOM 1183 CH2 TRP A 211 13.155 -13.379 8.278 1.00 35.38 C +ATOM 1184 N SER A 212 14.292 -5.990 8.778 1.00 35.38 N +ATOM 1185 CA SER A 212 14.780 -5.472 10.056 1.00 35.38 C +ATOM 1186 C SER A 212 13.646 -4.826 10.869 1.00 35.38 C +ATOM 1187 O SER A 212 13.515 -5.086 12.068 1.00 35.38 O +ATOM 1188 CB SER A 212 15.889 -4.452 9.803 1.00 35.38 C +ATOM 1189 OG SER A 212 16.959 -4.985 9.041 1.00 35.38 O +ATOM 1190 N LEU A 213 12.774 -4.058 10.204 1.00 35.38 N +ATOM 1191 CA LEU A 213 11.548 -3.504 10.781 1.00 35.38 C +ATOM 1192 C LEU A 213 10.573 -4.606 11.217 1.00 35.38 C +ATOM 1193 O LEU A 213 10.023 -4.519 12.312 1.00 35.38 O +ATOM 1194 CB LEU A 213 10.923 -2.529 9.767 1.00 35.38 C +ATOM 1195 CG LEU A 213 9.646 -1.821 10.275 1.00 35.38 C +ATOM 1196 CD1 LEU A 213 9.548 -0.431 9.659 1.00 35.38 C +ATOM 1197 CD2 LEU A 213 8.362 -2.564 9.898 1.00 35.38 C +ATOM 1198 N GLY A 214 10.401 -5.666 10.424 1.00 35.41 N +ATOM 1199 CA GLY A 214 9.602 -6.835 10.800 1.00 35.41 C +ATOM 1200 C GLY A 214 10.099 -7.513 12.081 1.00 35.41 C +ATOM 1201 O GLY A 214 9.297 -7.840 12.953 1.00 35.41 O +ATOM 1202 N CYS A 215 11.418 -7.661 12.251 1.00 35.41 N +ATOM 1203 CA CYS A 215 12.011 -8.182 13.486 1.00 35.41 C +ATOM 1204 C CYS A 215 11.719 -7.283 14.696 1.00 35.41 C +ATOM 1205 O CYS A 215 11.385 -7.800 15.763 1.00 35.41 O +ATOM 1206 CB CYS A 215 13.529 -8.334 13.318 1.00 35.41 C +ATOM 1207 SG CYS A 215 13.926 -9.732 12.242 1.00 35.41 S +ATOM 1208 N MET A 216 11.828 -5.958 14.535 1.00 35.56 N +ATOM 1209 CA MET A 216 11.471 -4.994 15.585 1.00 35.56 C +ATOM 1210 C MET A 216 9.989 -5.087 15.946 1.00 35.56 C +ATOM 1211 O MET A 216 9.642 -5.191 17.119 1.00 35.56 O +ATOM 1212 CB MET A 216 11.769 -3.556 15.135 1.00 35.56 C +ATOM 1213 CG MET A 216 13.264 -3.258 15.094 1.00 35.56 C +ATOM 1214 SD MET A 216 13.691 -1.555 14.633 1.00 35.56 S +ATOM 1215 CE MET A 216 12.816 -0.622 15.919 1.00 35.56 C +ATOM 1216 N PHE A 217 9.119 -5.094 14.937 1.00 35.47 N +ATOM 1217 CA PHE A 217 7.676 -5.154 15.125 1.00 35.47 C +ATOM 1218 C PHE A 217 7.251 -6.449 15.831 1.00 35.47 C +ATOM 1219 O PHE A 217 6.504 -6.392 16.805 1.00 35.47 O +ATOM 1220 CB PHE A 217 6.995 -4.985 13.762 1.00 35.47 C +ATOM 1221 CG PHE A 217 5.485 -5.016 13.842 1.00 35.47 C +ATOM 1222 CD1 PHE A 217 4.753 -5.956 13.097 1.00 35.47 C +ATOM 1223 CD2 PHE A 217 4.809 -4.121 14.691 1.00 35.47 C +ATOM 1224 CE1 PHE A 217 3.356 -6.017 13.221 1.00 35.47 C +ATOM 1225 CE2 PHE A 217 3.409 -4.162 14.786 1.00 35.47 C +ATOM 1226 CZ PHE A 217 2.682 -5.114 14.056 1.00 35.47 C +ATOM 1227 N ALA A 218 7.801 -7.602 15.427 1.00 35.47 N +ATOM 1228 CA ALA A 218 7.595 -8.871 16.123 1.00 35.47 C +ATOM 1229 C ALA A 218 8.065 -8.817 17.586 1.00 35.47 C +ATOM 1230 O ALA A 218 7.377 -9.327 18.471 1.00 35.47 O +ATOM 1231 CB ALA A 218 8.343 -9.980 15.375 1.00 35.47 C +ATOM 1232 N GLY A 219 9.220 -8.193 17.846 1.00 35.74 N +ATOM 1233 CA GLY A 219 9.749 -8.016 19.197 1.00 35.74 C +ATOM 1234 C GLY A 219 8.812 -7.210 20.097 1.00 35.74 C +ATOM 1235 O GLY A 219 8.564 -7.614 21.236 1.00 35.74 O +ATOM 1236 N MET A 220 8.218 -6.143 19.551 1.00 35.71 N +ATOM 1237 CA MET A 220 7.228 -5.305 20.233 1.00 35.71 C +ATOM 1238 C MET A 220 5.921 -6.062 20.511 1.00 35.71 C +ATOM 1239 O MET A 220 5.545 -6.209 21.674 1.00 35.71 O +ATOM 1240 CB MET A 220 6.936 -4.042 19.412 1.00 35.71 C +ATOM 1241 CG MET A 220 8.115 -3.075 19.280 1.00 35.71 C +ATOM 1242 SD MET A 220 7.740 -1.771 18.079 1.00 35.71 S +ATOM 1243 CE MET A 220 9.267 -0.818 18.123 1.00 35.71 C +ATOM 1244 N ILE A 221 5.237 -6.576 19.479 1.00 35.62 N +ATOM 1245 CA ILE A 221 3.880 -7.136 19.647 1.00 35.62 C +ATOM 1246 C ILE A 221 3.868 -8.452 20.419 1.00 35.62 C +ATOM 1247 O ILE A 221 2.898 -8.739 21.108 1.00 35.62 O +ATOM 1248 CB ILE A 221 3.114 -7.291 18.316 1.00 35.62 C +ATOM 1249 CG1 ILE A 221 3.694 -8.404 17.410 1.00 35.62 C +ATOM 1250 CG2 ILE A 221 3.012 -5.932 17.608 1.00 35.62 C +ATOM 1251 CD1 ILE A 221 2.879 -8.653 16.137 1.00 35.62 C +ATOM 1252 N PHE A 222 4.942 -9.248 20.344 1.00 35.71 N +ATOM 1253 CA PHE A 222 5.046 -10.491 21.109 1.00 35.71 C +ATOM 1254 C PHE A 222 5.766 -10.320 22.448 1.00 35.71 C +ATOM 1255 O PHE A 222 5.904 -11.305 23.177 1.00 35.71 O +ATOM 1256 CB PHE A 222 5.660 -11.608 20.257 1.00 35.71 C +ATOM 1257 CG PHE A 222 4.842 -11.972 19.035 1.00 35.71 C +ATOM 1258 CD1 PHE A 222 3.492 -12.350 19.176 1.00 35.71 C +ATOM 1259 CD2 PHE A 222 5.426 -11.941 17.757 1.00 35.71 C +ATOM 1260 CE1 PHE A 222 2.729 -12.682 18.043 1.00 35.71 C +ATOM 1261 CE2 PHE A 222 4.673 -12.305 16.629 1.00 35.71 C +ATOM 1262 CZ PHE A 222 3.322 -12.659 16.770 1.00 35.71 C +ATOM 1263 N ARG A 223 6.223 -9.104 22.786 1.00 36.43 N +ATOM 1264 CA ARG A 223 7.042 -8.811 23.977 1.00 36.43 C +ATOM 1265 C ARG A 223 8.234 -9.763 24.115 1.00 36.43 C +ATOM 1266 O ARG A 223 8.523 -10.296 25.186 1.00 36.43 O +ATOM 1267 CB ARG A 223 6.163 -8.748 25.234 1.00 36.43 C +ATOM 1268 CG ARG A 223 5.209 -7.549 25.198 1.00 36.43 C +ATOM 1269 CD ARG A 223 4.443 -7.473 26.518 1.00 36.43 C +ATOM 1270 NE ARG A 223 3.814 -6.154 26.706 1.00 36.43 N +ATOM 1271 CZ ARG A 223 3.135 -5.785 27.777 1.00 36.43 C +ATOM 1272 NH1 ARG A 223 2.772 -6.638 28.690 1.00 36.43 N +ATOM 1273 NH2 ARG A 223 2.799 -4.540 27.954 1.00 36.43 N +ATOM 1274 N LYS A 224 8.927 -9.992 23.002 1.00 36.53 N +ATOM 1275 CA LYS A 224 10.054 -10.923 22.903 1.00 36.53 C +ATOM 1276 C LYS A 224 11.207 -10.256 22.169 1.00 36.53 C +ATOM 1277 O LYS A 224 11.312 -10.359 20.952 1.00 36.53 O +ATOM 1278 CB LYS A 224 9.573 -12.213 22.227 1.00 36.53 C +ATOM 1279 CG LYS A 224 10.678 -13.278 22.249 1.00 36.53 C +ATOM 1280 CD LYS A 224 10.299 -14.447 21.348 1.00 36.53 C +ATOM 1281 CE LYS A 224 11.443 -15.457 21.339 1.00 36.53 C +ATOM 1282 NZ LYS A 224 11.418 -16.239 20.087 1.00 36.53 N +ATOM 1283 N GLU A 225 12.078 -9.594 22.922 1.00 38.61 N +ATOM 1284 CA GLU A 225 13.162 -8.777 22.377 1.00 38.61 C +ATOM 1285 C GLU A 225 14.540 -9.414 22.656 1.00 38.61 C +ATOM 1286 O GLU A 225 14.934 -9.532 23.820 1.00 38.61 O +ATOM 1287 CB GLU A 225 13.052 -7.362 22.964 1.00 38.61 C +ATOM 1288 CG GLU A 225 14.019 -6.347 22.336 1.00 38.61 C +ATOM 1289 CD GLU A 225 13.785 -6.098 20.839 1.00 38.61 C +ATOM 1290 OE1 GLU A 225 14.798 -5.877 20.134 1.00 38.61 O +ATOM 1291 OE2 GLU A 225 12.621 -6.179 20.399 1.00 38.61 O +ATOM 1292 N PRO A 226 15.307 -9.807 21.623 1.00 35.93 N +ATOM 1293 CA PRO A 226 14.943 -9.852 20.204 1.00 35.93 C +ATOM 1294 C PRO A 226 14.056 -11.065 19.861 1.00 35.93 C +ATOM 1295 O PRO A 226 14.124 -12.116 20.506 1.00 35.93 O +ATOM 1296 CB PRO A 226 16.278 -9.952 19.471 1.00 35.93 C +ATOM 1297 CG PRO A 226 17.076 -10.827 20.425 1.00 35.93 C +ATOM 1298 CD PRO A 226 16.653 -10.317 21.802 1.00 35.93 C +ATOM 1299 N PHE A 227 13.285 -10.964 18.773 1.00 35.59 N +ATOM 1300 CA PHE A 227 12.362 -12.035 18.380 1.00 35.59 C +ATOM 1301 C PHE A 227 13.093 -13.306 17.903 1.00 35.59 C +ATOM 1302 O PHE A 227 12.779 -14.424 18.342 1.00 35.59 O +ATOM 1303 CB PHE A 227 11.388 -11.505 17.322 1.00 35.59 C +ATOM 1304 CG PHE A 227 10.312 -12.513 16.981 1.00 35.59 C +ATOM 1305 CD1 PHE A 227 10.388 -13.264 15.796 1.00 35.59 C +ATOM 1306 CD2 PHE A 227 9.247 -12.726 17.876 1.00 35.59 C +ATOM 1307 CE1 PHE A 227 9.408 -14.227 15.505 1.00 35.59 C +ATOM 1308 CE2 PHE A 227 8.273 -13.700 17.592 1.00 35.59 C +ATOM 1309 CZ PHE A 227 8.349 -14.446 16.405 1.00 35.59 C +ATOM 1310 N PHE A 228 14.107 -13.129 17.049 1.00 35.50 N +ATOM 1311 CA PHE A 228 15.013 -14.181 16.584 1.00 35.50 C +ATOM 1312 C PHE A 228 16.362 -14.077 17.304 1.00 35.50 C +ATOM 1313 O PHE A 228 17.214 -13.275 16.926 1.00 35.50 O +ATOM 1314 CB PHE A 228 15.201 -14.082 15.064 1.00 35.50 C +ATOM 1315 CG PHE A 228 13.951 -14.250 14.236 1.00 35.50 C +ATOM 1316 CD1 PHE A 228 13.290 -15.489 14.195 1.00 35.50 C +ATOM 1317 CD2 PHE A 228 13.478 -13.177 13.459 1.00 35.50 C +ATOM 1318 CE1 PHE A 228 12.158 -15.645 13.381 1.00 35.50 C +ATOM 1319 CE2 PHE A 228 12.349 -13.339 12.640 1.00 35.50 C +ATOM 1320 CZ PHE A 228 11.685 -14.577 12.601 1.00 35.50 C +ATOM 1321 N TYR A 229 16.556 -14.895 18.341 1.00 35.65 N +ATOM 1322 CA TYR A 229 17.753 -14.884 19.187 1.00 35.65 C +ATOM 1323 C TYR A 229 18.766 -15.960 18.761 1.00 35.65 C +ATOM 1324 O TYR A 229 18.739 -17.071 19.290 1.00 35.65 O +ATOM 1325 CB TYR A 229 17.322 -15.030 20.656 1.00 35.65 C +ATOM 1326 CG TYR A 229 18.440 -14.761 21.647 1.00 35.65 C +ATOM 1327 CD1 TYR A 229 19.271 -15.793 22.124 1.00 35.65 C +ATOM 1328 CD2 TYR A 229 18.651 -13.451 22.100 1.00 35.65 C +ATOM 1329 CE1 TYR A 229 20.306 -15.499 23.035 1.00 35.65 C +ATOM 1330 CE2 TYR A 229 19.687 -13.142 22.992 1.00 35.65 C +ATOM 1331 CZ TYR A 229 20.521 -14.169 23.461 1.00 35.65 C +ATOM 1332 OH TYR A 229 21.516 -13.867 24.334 1.00 35.65 O +ATOM 1333 N GLY A 230 19.620 -15.661 17.778 1.00 35.71 N +ATOM 1334 CA GLY A 230 20.718 -16.533 17.341 1.00 35.71 C +ATOM 1335 C GLY A 230 22.033 -16.291 18.094 1.00 35.71 C +ATOM 1336 O GLY A 230 22.436 -15.145 18.324 1.00 35.71 O +ATOM 1337 N HIS A 231 22.745 -17.358 18.456 1.00 36.26 N +ATOM 1338 CA HIS A 231 24.043 -17.255 19.144 1.00 36.26 C +ATOM 1339 C HIS A 231 25.166 -16.729 18.234 1.00 36.26 C +ATOM 1340 O HIS A 231 26.035 -15.974 18.678 1.00 36.26 O +ATOM 1341 CB HIS A 231 24.396 -18.614 19.758 1.00 36.26 C +ATOM 1342 CG HIS A 231 23.401 -19.038 20.810 1.00 36.26 C +ATOM 1343 ND1 HIS A 231 23.144 -18.385 21.994 1.00 36.26 N +ATOM 1344 CD2 HIS A 231 22.550 -20.108 20.748 1.00 36.26 C +ATOM 1345 CE1 HIS A 231 22.166 -19.049 22.631 1.00 36.26 C +ATOM 1346 NE2 HIS A 231 21.758 -20.100 21.904 1.00 36.26 N +ATOM 1347 N ASP A 232 25.089 -17.038 16.943 1.00 36.15 N +ATOM 1348 CA ASP A 232 25.946 -16.550 15.864 1.00 36.15 C +ATOM 1349 C ASP A 232 25.114 -16.370 14.575 1.00 36.15 C +ATOM 1350 O ASP A 232 23.884 -16.415 14.621 1.00 36.15 O +ATOM 1351 CB ASP A 232 27.144 -17.498 15.696 1.00 36.15 C +ATOM 1352 CG ASP A 232 26.720 -18.937 15.412 1.00 36.15 C +ATOM 1353 OD1 ASP A 232 25.842 -19.100 14.535 1.00 36.15 O +ATOM 1354 OD2 ASP A 232 27.267 -19.836 16.073 1.00 36.15 O +ATOM 1355 N ASN A 233 25.750 -16.088 13.432 1.00 36.26 N +ATOM 1356 CA ASN A 233 25.015 -15.885 12.177 1.00 36.26 C +ATOM 1357 C ASN A 233 24.406 -17.186 11.624 1.00 36.26 C +ATOM 1358 O ASN A 233 23.345 -17.121 11.008 1.00 36.26 O +ATOM 1359 CB ASN A 233 25.929 -15.239 11.125 1.00 36.26 C +ATOM 1360 CG ASN A 233 26.322 -13.796 11.391 1.00 36.26 C +ATOM 1361 OD1 ASN A 233 25.820 -13.079 12.250 1.00 36.26 O +ATOM 1362 ND2 ASN A 233 27.275 -13.309 10.633 1.00 36.26 N +ATOM 1363 N HIS A 234 25.027 -18.345 11.854 1.00 36.26 N +ATOM 1364 CA HIS A 234 24.488 -19.628 11.401 1.00 36.26 C +ATOM 1365 C HIS A 234 23.243 -19.991 12.217 1.00 36.26 C +ATOM 1366 O HIS A 234 22.171 -20.208 11.654 1.00 36.26 O +ATOM 1367 CB HIS A 234 25.575 -20.712 11.492 1.00 36.26 C +ATOM 1368 CG HIS A 234 26.767 -20.493 10.590 1.00 36.26 C +ATOM 1369 ND1 HIS A 234 28.045 -20.943 10.832 1.00 36.26 N +ATOM 1370 CD2 HIS A 234 26.791 -19.869 9.368 1.00 36.26 C +ATOM 1371 CE1 HIS A 234 28.815 -20.602 9.786 1.00 36.26 C +ATOM 1372 NE2 HIS A 234 28.107 -19.893 8.899 1.00 36.26 N +ATOM 1373 N ASP A 235 23.323 -19.927 13.550 1.00 35.99 N +ATOM 1374 CA ASP A 235 22.166 -20.157 14.424 1.00 35.99 C +ATOM 1375 C ASP A 235 21.076 -19.085 14.237 1.00 35.99 C +ATOM 1376 O ASP A 235 19.895 -19.376 14.418 1.00 35.99 O +ATOM 1377 CB ASP A 235 22.624 -20.274 15.885 1.00 35.99 C +ATOM 1378 CG ASP A 235 21.476 -20.565 16.866 1.00 35.99 C +ATOM 1379 OD1 ASP A 235 20.718 -21.555 16.727 1.00 35.99 O +ATOM 1380 OD2 ASP A 235 21.319 -19.784 17.833 1.00 35.99 O +ATOM 1381 N GLN A 236 21.419 -17.864 13.805 1.00 35.53 N +ATOM 1382 CA GLN A 236 20.416 -16.857 13.443 1.00 35.53 C +ATOM 1383 C GLN A 236 19.497 -17.349 12.313 1.00 35.53 C +ATOM 1384 O GLN A 236 18.276 -17.219 12.437 1.00 35.53 O +ATOM 1385 CB GLN A 236 21.092 -15.527 13.066 1.00 35.53 C +ATOM 1386 CG GLN A 236 20.083 -14.387 12.861 1.00 35.53 C +ATOM 1387 CD GLN A 236 19.353 -13.993 14.140 1.00 35.53 C +ATOM 1388 OE1 GLN A 236 19.841 -14.127 15.254 1.00 35.53 O +ATOM 1389 NE2 GLN A 236 18.161 -13.460 14.042 1.00 35.53 N +ATOM 1390 N LEU A 237 20.050 -17.957 11.255 1.00 35.62 N +ATOM 1391 CA LEU A 237 19.249 -18.537 10.169 1.00 35.62 C +ATOM 1392 C LEU A 237 18.388 -19.698 10.680 1.00 35.62 C +ATOM 1393 O LEU A 237 17.213 -19.798 10.331 1.00 35.62 O +ATOM 1394 CB LEU A 237 20.174 -18.976 9.019 1.00 35.62 C +ATOM 1395 CG LEU A 237 19.427 -19.498 7.773 1.00 35.62 C +ATOM 1396 CD1 LEU A 237 18.529 -18.431 7.140 1.00 35.62 C +ATOM 1397 CD2 LEU A 237 20.435 -19.958 6.722 1.00 35.62 C +ATOM 1398 N VAL A 238 18.925 -20.518 11.588 1.00 35.71 N +ATOM 1399 CA VAL A 238 18.185 -21.618 12.232 1.00 35.71 C +ATOM 1400 C VAL A 238 16.984 -21.094 13.023 1.00 35.71 C +ATOM 1401 O VAL A 238 15.895 -21.666 12.944 1.00 35.71 O +ATOM 1402 CB VAL A 238 19.097 -22.418 13.177 1.00 35.71 C +ATOM 1403 CG1 VAL A 238 18.365 -23.590 13.848 1.00 35.71 C +ATOM 1404 CG2 VAL A 238 20.322 -22.978 12.462 1.00 35.71 C +ATOM 1405 N LYS A 239 17.132 -20.000 13.788 1.00 35.56 N +ATOM 1406 CA LYS A 239 15.998 -19.409 14.523 1.00 35.56 C +ATOM 1407 C LYS A 239 14.906 -18.921 13.580 1.00 35.56 C +ATOM 1408 O LYS A 239 13.730 -19.085 13.900 1.00 35.56 O +ATOM 1409 CB LYS A 239 16.415 -18.233 15.417 1.00 35.56 C +ATOM 1410 CG LYS A 239 17.415 -18.526 16.535 1.00 35.56 C +ATOM 1411 CD LYS A 239 17.233 -19.840 17.309 1.00 35.56 C +ATOM 1412 CE LYS A 239 18.281 -19.812 18.424 1.00 35.56 C +ATOM 1413 NZ LYS A 239 18.880 -21.121 18.738 1.00 35.56 N +ATOM 1414 N ILE A 240 15.286 -18.352 12.437 1.00 35.47 N +ATOM 1415 CA ILE A 240 14.344 -17.922 11.401 1.00 35.47 C +ATOM 1416 C ILE A 240 13.628 -19.144 10.810 1.00 35.47 C +ATOM 1417 O ILE A 240 12.396 -19.190 10.818 1.00 35.47 O +ATOM 1418 CB ILE A 240 15.071 -17.066 10.343 1.00 35.47 C +ATOM 1419 CG1 ILE A 240 15.584 -15.749 10.973 1.00 35.47 C +ATOM 1420 CG2 ILE A 240 14.107 -16.761 9.188 1.00 35.47 C +ATOM 1421 CD1 ILE A 240 16.631 -15.033 10.109 1.00 35.47 C +ATOM 1422 N ALA A 241 14.379 -20.177 10.415 1.00 35.59 N +ATOM 1423 CA ALA A 241 13.855 -21.430 9.869 1.00 35.59 C +ATOM 1424 C ALA A 241 12.878 -22.148 10.801 1.00 35.59 C +ATOM 1425 O ALA A 241 11.850 -22.661 10.357 1.00 35.59 O +ATOM 1426 CB ALA A 241 15.034 -22.317 9.465 1.00 35.59 C +ATOM 1427 N LYS A 242 13.123 -22.115 12.113 1.00 35.68 N +ATOM 1428 CA LYS A 242 12.214 -22.680 13.122 1.00 35.68 C +ATOM 1429 C LYS A 242 10.842 -22.004 13.182 1.00 35.68 C +ATOM 1430 O LYS A 242 9.904 -22.619 13.695 1.00 35.68 O +ATOM 1431 CB LYS A 242 12.903 -22.644 14.495 1.00 35.68 C +ATOM 1432 CG LYS A 242 13.884 -23.813 14.592 1.00 35.68 C +ATOM 1433 CD LYS A 242 14.707 -23.822 15.882 1.00 35.68 C +ATOM 1434 CE LYS A 242 15.554 -25.100 15.832 1.00 35.68 C +ATOM 1435 NZ LYS A 242 16.550 -25.193 16.926 1.00 35.68 N +ATOM 1436 N VAL A 243 10.722 -20.768 12.693 1.00 35.53 N +ATOM 1437 CA VAL A 243 9.466 -20.003 12.663 1.00 35.53 C +ATOM 1438 C VAL A 243 8.860 -19.986 11.267 1.00 35.53 C +ATOM 1439 O VAL A 243 7.716 -20.400 11.105 1.00 35.53 O +ATOM 1440 CB VAL A 243 9.669 -18.569 13.181 1.00 35.53 C +ATOM 1441 CG1 VAL A 243 8.357 -17.777 13.167 1.00 35.53 C +ATOM 1442 CG2 VAL A 243 10.229 -18.583 14.613 1.00 35.53 C +ATOM 1443 N LEU A 244 9.612 -19.548 10.259 1.00 35.50 N +ATOM 1444 CA LEU A 244 9.109 -19.399 8.891 1.00 35.50 C +ATOM 1445 C LEU A 244 9.058 -20.726 8.118 1.00 35.50 C +ATOM 1446 O LEU A 244 8.416 -20.806 7.073 1.00 35.50 O +ATOM 1447 CB LEU A 244 9.943 -18.340 8.154 1.00 35.50 C +ATOM 1448 CG LEU A 244 9.918 -16.923 8.749 1.00 35.50 C +ATOM 1449 CD1 LEU A 244 10.640 -15.980 7.791 1.00 35.50 C +ATOM 1450 CD2 LEU A 244 8.501 -16.390 8.966 1.00 35.50 C +ATOM 1451 N GLY A 245 9.667 -21.782 8.659 1.00 35.62 N +ATOM 1452 CA GLY A 245 9.749 -23.096 8.036 1.00 35.62 C +ATOM 1453 C GLY A 245 10.874 -23.190 7.008 1.00 35.62 C +ATOM 1454 O GLY A 245 11.383 -22.182 6.518 1.00 35.62 O +ATOM 1455 N THR A 246 11.268 -24.418 6.691 1.00 35.71 N +ATOM 1456 CA THR A 246 12.372 -24.710 5.762 1.00 35.71 C +ATOM 1457 C THR A 246 11.908 -24.810 4.316 1.00 35.71 C +ATOM 1458 O THR A 246 12.680 -24.516 3.416 1.00 35.71 O +ATOM 1459 CB THR A 246 13.119 -25.980 6.189 1.00 35.71 C +ATOM 1460 OG1 THR A 246 12.210 -27.042 6.430 1.00 35.71 O +ATOM 1461 CG2 THR A 246 13.876 -25.684 7.484 1.00 35.71 C +ATOM 1462 N ASP A 247 10.637 -25.134 4.066 1.00 35.80 N +ATOM 1463 CA ASP A 247 10.115 -25.242 2.695 1.00 35.80 C +ATOM 1464 C ASP A 247 10.187 -23.889 1.964 1.00 35.80 C +ATOM 1465 O ASP A 247 10.676 -23.812 0.839 1.00 35.80 O +ATOM 1466 CB ASP A 247 8.677 -25.787 2.715 1.00 35.80 C +ATOM 1467 CG ASP A 247 8.538 -27.179 3.355 1.00 35.80 C +ATOM 1468 OD1 ASP A 247 9.550 -27.912 3.460 1.00 35.80 O +ATOM 1469 OD2 ASP A 247 7.419 -27.483 3.821 1.00 35.80 O +ATOM 1470 N GLY A 248 9.790 -22.799 2.635 1.00 35.93 N +ATOM 1471 CA GLY A 248 9.911 -21.442 2.090 1.00 35.93 C +ATOM 1472 C GLY A 248 11.364 -20.987 1.909 1.00 35.93 C +ATOM 1473 O GLY A 248 11.665 -20.293 0.940 1.00 35.93 O +ATOM 1474 N LEU A 249 12.270 -21.416 2.797 1.00 35.68 N +ATOM 1475 CA LEU A 249 13.706 -21.164 2.652 1.00 35.68 C +ATOM 1476 C LEU A 249 14.247 -21.865 1.401 1.00 35.68 C +ATOM 1477 O LEU A 249 14.922 -21.234 0.599 1.00 35.68 O +ATOM 1478 CB LEU A 249 14.443 -21.611 3.930 1.00 35.68 C +ATOM 1479 CG LEU A 249 15.975 -21.466 3.883 1.00 35.68 C +ATOM 1480 CD1 LEU A 249 16.433 -20.026 3.649 1.00 35.68 C +ATOM 1481 CD2 LEU A 249 16.558 -21.941 5.213 1.00 35.68 C +ATOM 1482 N ASN A 250 13.892 -23.132 1.184 1.00 35.90 N +ATOM 1483 CA ASN A 250 14.342 -23.907 0.028 1.00 35.90 C +ATOM 1484 C ASN A 250 13.831 -23.325 -1.296 1.00 35.90 C +ATOM 1485 O ASN A 250 14.582 -23.253 -2.266 1.00 35.90 O +ATOM 1486 CB ASN A 250 13.894 -25.367 0.207 1.00 35.90 C +ATOM 1487 CG ASN A 250 14.605 -26.067 1.352 1.00 35.90 C +ATOM 1488 OD1 ASN A 250 15.618 -25.632 1.863 1.00 35.90 O +ATOM 1489 ND2 ASN A 250 14.101 -27.194 1.797 1.00 35.90 N +ATOM 1490 N VAL A 251 12.575 -22.860 -1.345 1.00 35.80 N +ATOM 1491 CA VAL A 251 12.037 -22.154 -2.524 1.00 35.80 C +ATOM 1492 C VAL A 251 12.848 -20.891 -2.821 1.00 35.80 C +ATOM 1493 O VAL A 251 13.218 -20.661 -3.972 1.00 35.80 O +ATOM 1494 CB VAL A 251 10.545 -21.817 -2.337 1.00 35.80 C +ATOM 1495 CG1 VAL A 251 9.999 -20.901 -3.442 1.00 35.80 C +ATOM 1496 CG2 VAL A 251 9.702 -23.099 -2.358 1.00 35.80 C +ATOM 1497 N TYR A 252 13.165 -20.103 -1.791 1.00 35.71 N +ATOM 1498 CA TYR A 252 13.983 -18.900 -1.918 1.00 35.71 C +ATOM 1499 C TYR A 252 15.404 -19.213 -2.420 1.00 35.71 C +ATOM 1500 O TYR A 252 15.838 -18.641 -3.419 1.00 35.71 O +ATOM 1501 CB TYR A 252 13.981 -18.169 -0.569 1.00 35.71 C +ATOM 1502 CG TYR A 252 15.009 -17.067 -0.460 1.00 35.71 C +ATOM 1503 CD1 TYR A 252 16.110 -17.212 0.405 1.00 35.71 C +ATOM 1504 CD2 TYR A 252 14.876 -15.911 -1.246 1.00 35.71 C +ATOM 1505 CE1 TYR A 252 17.072 -16.191 0.495 1.00 35.71 C +ATOM 1506 CE2 TYR A 252 15.840 -14.889 -1.164 1.00 35.71 C +ATOM 1507 CZ TYR A 252 16.941 -15.030 -0.296 1.00 35.71 C +ATOM 1508 OH TYR A 252 17.885 -14.063 -0.227 1.00 35.71 O +ATOM 1509 N LEU A 253 16.099 -20.169 -1.795 1.00 35.96 N +ATOM 1510 CA LEU A 253 17.446 -20.590 -2.197 1.00 35.96 C +ATOM 1511 C LEU A 253 17.478 -21.061 -3.658 1.00 35.96 C +ATOM 1512 O LEU A 253 18.308 -20.597 -4.439 1.00 35.96 O +ATOM 1513 CB LEU A 253 17.925 -21.712 -1.262 1.00 35.96 C +ATOM 1514 CG LEU A 253 18.178 -21.279 0.193 1.00 35.96 C +ATOM 1515 CD1 LEU A 253 18.435 -22.521 1.040 1.00 35.96 C +ATOM 1516 CD2 LEU A 253 19.359 -20.324 0.336 1.00 35.96 C +ATOM 1517 N ASN A 254 16.513 -21.894 -4.061 1.00 35.96 N +ATOM 1518 CA ASN A 254 16.396 -22.378 -5.437 1.00 35.96 C +ATOM 1519 C ASN A 254 16.129 -21.244 -6.437 1.00 35.96 C +ATOM 1520 O ASN A 254 16.749 -21.201 -7.502 1.00 35.96 O +ATOM 1521 CB ASN A 254 15.277 -23.431 -5.501 1.00 35.96 C +ATOM 1522 CG ASN A 254 15.675 -24.758 -4.880 1.00 35.96 C +ATOM 1523 OD1 ASN A 254 16.826 -25.143 -4.845 1.00 35.96 O +ATOM 1524 ND2 ASN A 254 14.729 -25.538 -4.413 1.00 35.96 N +ATOM 1525 N LYS A 255 15.235 -20.300 -6.102 1.00 35.86 N +ATOM 1526 CA LYS A 255 14.898 -19.155 -6.966 1.00 35.86 C +ATOM 1527 C LYS A 255 16.126 -18.299 -7.284 1.00 35.86 C +ATOM 1528 O LYS A 255 16.290 -17.874 -8.428 1.00 35.86 O +ATOM 1529 CB LYS A 255 13.792 -18.325 -6.294 1.00 35.86 C +ATOM 1530 CG LYS A 255 13.388 -17.095 -7.124 1.00 35.86 C +ATOM 1531 CD LYS A 255 12.270 -16.320 -6.423 1.00 35.86 C +ATOM 1532 CE LYS A 255 11.952 -15.023 -7.175 1.00 35.86 C +ATOM 1533 NZ LYS A 255 10.937 -14.232 -6.443 1.00 35.86 N +ATOM 1534 N TYR A 256 16.984 -18.066 -6.291 1.00 36.09 N +ATOM 1535 CA TYR A 256 18.186 -17.239 -6.430 1.00 36.09 C +ATOM 1536 C TYR A 256 19.467 -18.039 -6.703 1.00 36.09 C +ATOM 1537 O TYR A 256 20.523 -17.431 -6.851 1.00 36.09 O +ATOM 1538 CB TYR A 256 18.296 -16.289 -5.231 1.00 36.09 C +ATOM 1539 CG TYR A 256 17.199 -15.244 -5.233 1.00 36.09 C +ATOM 1540 CD1 TYR A 256 17.322 -14.101 -6.046 1.00 36.09 C +ATOM 1541 CD2 TYR A 256 16.034 -15.442 -4.472 1.00 36.09 C +ATOM 1542 CE1 TYR A 256 16.269 -13.169 -6.116 1.00 36.09 C +ATOM 1543 CE2 TYR A 256 14.972 -14.525 -4.549 1.00 36.09 C +ATOM 1544 CZ TYR A 256 15.088 -13.396 -5.376 1.00 36.09 C +ATOM 1545 OH TYR A 256 14.052 -12.535 -5.485 1.00 36.09 O +ATOM 1546 N ARG A 257 19.371 -19.369 -6.853 1.00 36.39 N +ATOM 1547 CA ARG A 257 20.500 -20.285 -7.104 1.00 36.39 C +ATOM 1548 C ARG A 257 21.597 -20.168 -6.039 1.00 36.39 C +ATOM 1549 O ARG A 257 22.775 -20.057 -6.363 1.00 36.39 O +ATOM 1550 CB ARG A 257 21.054 -20.108 -8.527 1.00 36.39 C +ATOM 1551 CG ARG A 257 19.987 -20.215 -9.620 1.00 36.39 C +ATOM 1552 CD ARG A 257 20.664 -20.008 -10.973 1.00 36.39 C +ATOM 1553 NE ARG A 257 19.677 -20.003 -12.065 1.00 36.39 N +ATOM 1554 CZ ARG A 257 19.950 -19.871 -13.350 1.00 36.39 C +ATOM 1555 NH1 ARG A 257 21.171 -19.713 -13.778 1.00 36.39 N +ATOM 1556 NH2 ARG A 257 18.990 -19.896 -14.232 1.00 36.39 N +ATOM 1557 N ILE A 258 21.181 -20.158 -4.777 1.00 36.78 N +ATOM 1558 CA ILE A 258 22.066 -20.067 -3.615 1.00 36.78 C +ATOM 1559 C ILE A 258 22.243 -21.469 -3.037 1.00 36.78 C +ATOM 1560 O ILE A 258 21.263 -22.106 -2.651 1.00 36.78 O +ATOM 1561 CB ILE A 258 21.490 -19.098 -2.563 1.00 36.78 C +ATOM 1562 CG1 ILE A 258 21.207 -17.696 -3.150 1.00 36.78 C +ATOM 1563 CG2 ILE A 258 22.448 -18.992 -1.361 1.00 36.78 C +ATOM 1564 CD1 ILE A 258 20.347 -16.840 -2.216 1.00 36.78 C +ATOM 1565 N GLU A 259 23.487 -21.926 -2.942 1.00 37.62 N +ATOM 1566 CA GLU A 259 23.846 -23.160 -2.242 1.00 37.62 C +ATOM 1567 C GLU A 259 24.312 -22.813 -0.825 1.00 37.62 C +ATOM 1568 O GLU A 259 25.213 -21.992 -0.641 1.00 37.62 O +ATOM 1569 CB GLU A 259 24.914 -23.939 -3.025 1.00 37.62 C +ATOM 1570 CG GLU A 259 24.354 -24.489 -4.349 1.00 37.62 C +ATOM 1571 CD GLU A 259 25.369 -25.301 -5.173 1.00 37.62 C +ATOM 1572 OE1 GLU A 259 24.969 -25.761 -6.268 1.00 37.62 O +ATOM 1573 OE2 GLU A 259 26.535 -25.446 -4.737 1.00 37.62 O +ATOM 1574 N LEU A 260 23.681 -23.413 0.187 1.00 38.34 N +ATOM 1575 CA LEU A 260 24.167 -23.310 1.560 1.00 38.34 C +ATOM 1576 C LEU A 260 25.397 -24.195 1.732 1.00 38.34 C +ATOM 1577 O LEU A 260 25.502 -25.255 1.115 1.00 38.34 O +ATOM 1578 CB LEU A 260 23.078 -23.700 2.572 1.00 38.34 C +ATOM 1579 CG LEU A 260 21.855 -22.773 2.604 1.00 38.34 C +ATOM 1580 CD1 LEU A 260 20.885 -23.268 3.680 1.00 38.34 C +ATOM 1581 CD2 LEU A 260 22.205 -21.319 2.924 1.00 38.34 C +ATOM 1582 N ASP A 261 26.309 -23.799 2.620 1.00 40.79 N +ATOM 1583 CA ASP A 261 27.359 -24.723 3.015 1.00 40.79 C +ATOM 1584 C ASP A 261 26.741 -25.947 3.731 1.00 40.79 C +ATOM 1585 O ASP A 261 25.752 -25.805 4.465 1.00 40.79 O +ATOM 1586 CB ASP A 261 28.444 -24.029 3.843 1.00 40.79 C +ATOM 1587 CG ASP A 261 28.052 -23.899 5.308 1.00 40.79 C +ATOM 1588 OD1 ASP A 261 28.296 -24.878 6.051 1.00 40.79 O +ATOM 1589 OD2 ASP A 261 27.510 -22.837 5.670 1.00 40.79 O +ATOM 1590 N PRO A 262 27.315 -27.153 3.566 1.00 38.84 N +ATOM 1591 CA PRO A 262 26.729 -28.372 4.123 1.00 38.84 C +ATOM 1592 C PRO A 262 26.570 -28.366 5.651 1.00 38.84 C +ATOM 1593 O PRO A 262 25.690 -29.047 6.182 1.00 38.84 O +ATOM 1594 CB PRO A 262 27.669 -29.498 3.683 1.00 38.84 C +ATOM 1595 CG PRO A 262 28.276 -28.972 2.384 1.00 38.84 C +ATOM 1596 CD PRO A 262 28.401 -27.477 2.649 1.00 38.84 C +ATOM 1597 N GLN A 263 27.412 -27.626 6.385 1.00 38.75 N +ATOM 1598 CA GLN A 263 27.311 -27.555 7.846 1.00 38.75 C +ATOM 1599 C GLN A 263 26.111 -26.702 8.259 1.00 38.75 C +ATOM 1600 O GLN A 263 25.366 -27.089 9.161 1.00 38.75 O +ATOM 1601 CB GLN A 263 28.597 -27.002 8.486 1.00 38.75 C +ATOM 1602 CG GLN A 263 29.858 -27.796 8.119 1.00 38.75 C +ATOM 1603 CD GLN A 263 31.127 -27.248 8.770 1.00 38.75 C +ATOM 1604 OE1 GLN A 263 31.166 -26.241 9.458 1.00 38.75 O +ATOM 1605 NE2 GLN A 263 32.245 -27.918 8.592 1.00 38.75 N +ATOM 1606 N LEU A 264 25.893 -25.570 7.585 1.00 39.27 N +ATOM 1607 CA LEU A 264 24.724 -24.725 7.804 1.00 39.27 C +ATOM 1608 C LEU A 264 23.431 -25.418 7.379 1.00 39.27 C +ATOM 1609 O LEU A 264 22.449 -25.347 8.117 1.00 39.27 O +ATOM 1610 CB LEU A 264 24.927 -23.401 7.061 1.00 39.27 C +ATOM 1611 CG LEU A 264 23.785 -22.382 7.183 1.00 39.27 C +ATOM 1612 CD1 LEU A 264 23.504 -22.018 8.643 1.00 39.27 C +ATOM 1613 CD2 LEU A 264 24.185 -21.104 6.447 1.00 39.27 C +ATOM 1614 N GLU A 265 23.423 -26.127 6.249 1.00 38.13 N +ATOM 1615 CA GLU A 265 22.256 -26.900 5.811 1.00 38.13 C +ATOM 1616 C GLU A 265 21.857 -27.950 6.862 1.00 38.13 C +ATOM 1617 O GLU A 265 20.698 -28.010 7.290 1.00 38.13 O +ATOM 1618 CB GLU A 265 22.548 -27.546 4.450 1.00 38.13 C +ATOM 1619 CG GLU A 265 21.282 -28.201 3.877 1.00 38.13 C +ATOM 1620 CD GLU A 265 21.480 -28.833 2.493 1.00 38.13 C +ATOM 1621 OE1 GLU A 265 20.438 -29.130 1.863 1.00 38.13 O +ATOM 1622 OE2 GLU A 265 22.644 -29.067 2.103 1.00 38.13 O +ATOM 1623 N ALA A 266 22.833 -28.714 7.368 1.00 37.39 N +ATOM 1624 CA ALA A 266 22.605 -29.686 8.434 1.00 37.39 C +ATOM 1625 C ALA A 266 22.091 -29.030 9.728 1.00 37.39 C +ATOM 1626 O ALA A 266 21.218 -29.583 10.401 1.00 37.39 O +ATOM 1627 CB ALA A 266 23.909 -30.454 8.676 1.00 37.39 C +ATOM 1628 N LEU A 267 22.602 -27.843 10.073 1.00 37.30 N +ATOM 1629 CA LEU A 267 22.187 -27.099 11.263 1.00 37.30 C +ATOM 1630 C LEU A 267 20.751 -26.558 11.148 1.00 37.30 C +ATOM 1631 O LEU A 267 20.009 -26.557 12.136 1.00 37.30 O +ATOM 1632 CB LEU A 267 23.199 -25.964 11.496 1.00 37.30 C +ATOM 1633 CG LEU A 267 23.036 -25.273 12.858 1.00 37.30 C +ATOM 1634 CD1 LEU A 267 23.532 -26.136 14.016 1.00 37.30 C +ATOM 1635 CD2 LEU A 267 23.795 -23.944 12.871 1.00 37.30 C +ATOM 1636 N VAL A 268 20.356 -26.092 9.960 1.00 36.71 N +ATOM 1637 CA VAL A 268 19.001 -25.596 9.668 1.00 36.71 C +ATOM 1638 C VAL A 268 17.970 -26.720 9.791 1.00 36.71 C +ATOM 1639 O VAL A 268 16.895 -26.499 10.356 1.00 36.71 O +ATOM 1640 CB VAL A 268 18.969 -24.936 8.274 1.00 36.71 C +ATOM 1641 CG1 VAL A 268 17.547 -24.680 7.767 1.00 36.71 C +ATOM 1642 CG2 VAL A 268 19.677 -23.572 8.313 1.00 36.71 C +ATOM 1643 N GLY A 269 18.302 -27.923 9.316 1.00 37.62 N +ATOM 1644 CA GLY A 269 17.444 -29.101 9.423 1.00 37.62 C +ATOM 1645 C GLY A 269 16.104 -28.938 8.697 1.00 37.62 C +ATOM 1646 O GLY A 269 16.014 -28.279 7.666 1.00 37.62 O +ATOM 1647 N ARG A 270 15.036 -29.554 9.229 1.00 36.36 N +ATOM 1648 CA ARG A 270 13.677 -29.480 8.662 1.00 36.36 C +ATOM 1649 C ARG A 270 12.683 -28.939 9.680 1.00 36.36 C +ATOM 1650 O ARG A 270 12.519 -29.495 10.768 1.00 36.36 O +ATOM 1651 CB ARG A 270 13.265 -30.848 8.094 1.00 36.36 C +ATOM 1652 CG ARG A 270 11.990 -30.748 7.242 1.00 36.36 C +ATOM 1653 CD ARG A 270 11.666 -32.096 6.593 1.00 36.36 C +ATOM 1654 NE ARG A 270 10.491 -31.998 5.706 1.00 36.36 N +ATOM 1655 CZ ARG A 270 9.966 -32.977 4.991 1.00 36.36 C +ATOM 1656 NH1 ARG A 270 10.449 -34.190 5.016 1.00 36.36 N +ATOM 1657 NH2 ARG A 270 8.939 -32.743 4.223 1.00 36.36 N +ATOM 1658 N HIS A 271 11.959 -27.886 9.310 1.00 35.96 N +ATOM 1659 CA HIS A 271 10.982 -27.237 10.181 1.00 35.96 C +ATOM 1660 C HIS A 271 9.712 -26.851 9.422 1.00 35.96 C +ATOM 1661 O HIS A 271 9.766 -26.269 8.342 1.00 35.96 O +ATOM 1662 CB HIS A 271 11.624 -26.015 10.858 1.00 35.96 C +ATOM 1663 CG HIS A 271 12.756 -26.358 11.798 1.00 35.96 C +ATOM 1664 ND1 HIS A 271 12.690 -27.269 12.829 1.00 35.96 N +ATOM 1665 CD2 HIS A 271 14.039 -25.875 11.771 1.00 35.96 C +ATOM 1666 CE1 HIS A 271 13.898 -27.326 13.409 1.00 35.96 C +ATOM 1667 NE2 HIS A 271 14.735 -26.458 12.839 1.00 35.96 N +ATOM 1668 N SER A 272 8.549 -27.120 10.022 1.00 35.93 N +ATOM 1669 CA SER A 272 7.273 -26.589 9.538 1.00 35.93 C +ATOM 1670 C SER A 272 7.091 -25.130 9.950 1.00 35.93 C +ATOM 1671 O SER A 272 7.392 -24.776 11.098 1.00 35.93 O +ATOM 1672 CB SER A 272 6.094 -27.440 10.026 1.00 35.93 C +ATOM 1673 OG SER A 272 6.016 -27.551 11.446 1.00 35.93 O +ATOM 1674 N ARG A 273 6.547 -24.309 9.039 1.00 35.62 N +ATOM 1675 CA ARG A 273 6.151 -22.922 9.324 1.00 35.62 C +ATOM 1676 C ARG A 273 5.188 -22.891 10.508 1.00 35.62 C +ATOM 1677 O ARG A 273 4.247 -23.682 10.578 1.00 35.62 O +ATOM 1678 CB ARG A 273 5.535 -22.276 8.070 1.00 35.62 C +ATOM 1679 CG ARG A 273 5.201 -20.782 8.253 1.00 35.62 C +ATOM 1680 CD ARG A 273 4.767 -20.187 6.908 1.00 35.62 C +ATOM 1681 NE ARG A 273 4.430 -18.751 6.988 1.00 35.62 N +ATOM 1682 CZ ARG A 273 3.886 -18.039 6.013 1.00 35.62 C +ATOM 1683 NH1 ARG A 273 3.498 -18.572 4.890 1.00 35.62 N +ATOM 1684 NH2 ARG A 273 3.719 -16.765 6.113 1.00 35.62 N +ATOM 1685 N LYS A 274 5.442 -22.003 11.465 1.00 35.71 N +ATOM 1686 CA LYS A 274 4.616 -21.840 12.663 1.00 35.71 C +ATOM 1687 C LYS A 274 3.576 -20.748 12.426 1.00 35.71 C +ATOM 1688 O LYS A 274 3.931 -19.699 11.895 1.00 35.71 O +ATOM 1689 CB LYS A 274 5.486 -21.537 13.896 1.00 35.71 C +ATOM 1690 CG LYS A 274 6.586 -22.579 14.159 1.00 35.71 C +ATOM 1691 CD LYS A 274 6.058 -24.018 14.250 1.00 35.71 C +ATOM 1692 CE LYS A 274 7.226 -24.987 14.421 1.00 35.71 C +ATOM 1693 NZ LYS A 274 6.836 -26.356 14.004 1.00 35.71 N +ATOM 1694 N PRO A 275 2.315 -20.955 12.841 1.00 35.86 N +ATOM 1695 CA PRO A 275 1.345 -19.871 12.856 1.00 35.86 C +ATOM 1696 C PRO A 275 1.781 -18.814 13.873 1.00 35.86 C +ATOM 1697 O PRO A 275 2.203 -19.157 14.980 1.00 35.86 O +ATOM 1698 CB PRO A 275 0.013 -20.526 13.232 1.00 35.86 C +ATOM 1699 CG PRO A 275 0.429 -21.730 14.077 1.00 35.86 C +ATOM 1700 CD PRO A 275 1.756 -22.156 13.452 1.00 35.86 C +ATOM 1701 N TRP A 276 1.645 -17.533 13.528 1.00 35.68 N +ATOM 1702 CA TRP A 276 2.038 -16.428 14.407 1.00 35.68 C +ATOM 1703 C TRP A 276 1.307 -16.429 15.754 1.00 35.68 C +ATOM 1704 O TRP A 276 1.896 -16.084 16.774 1.00 35.68 O +ATOM 1705 CB TRP A 276 1.795 -15.107 13.689 1.00 35.68 C +ATOM 1706 CG TRP A 276 2.646 -14.885 12.485 1.00 35.68 C +ATOM 1707 CD1 TRP A 276 2.223 -14.810 11.203 1.00 35.68 C +ATOM 1708 CD2 TRP A 276 4.093 -14.701 12.438 1.00 35.68 C +ATOM 1709 NE1 TRP A 276 3.303 -14.605 10.375 1.00 35.68 N +ATOM 1710 CE2 TRP A 276 4.482 -14.534 11.077 1.00 35.68 C +ATOM 1711 CE3 TRP A 276 5.114 -14.660 13.413 1.00 35.68 C +ATOM 1712 CZ2 TRP A 276 5.818 -14.361 10.692 1.00 35.68 C +ATOM 1713 CZ3 TRP A 276 6.459 -14.485 13.038 1.00 35.68 C +ATOM 1714 CH2 TRP A 276 6.812 -14.352 11.682 1.00 35.68 C +ATOM 1715 N LEU A 277 0.063 -16.922 15.782 1.00 36.82 N +ATOM 1716 CA LEU A 277 -0.728 -17.086 17.008 1.00 36.82 C +ATOM 1717 C LEU A 277 -0.042 -17.971 18.063 1.00 36.82 C +ATOM 1718 O LEU A 277 -0.368 -17.871 19.239 1.00 36.82 O +ATOM 1719 CB LEU A 277 -2.110 -17.665 16.648 1.00 36.82 C +ATOM 1720 CG LEU A 277 -2.992 -16.755 15.773 1.00 36.82 C +ATOM 1721 CD1 LEU A 277 -4.285 -17.490 15.421 1.00 36.82 C +ATOM 1722 CD2 LEU A 277 -3.350 -15.444 16.474 1.00 36.82 C +ATOM 1723 N LYS A 278 0.950 -18.791 17.685 1.00 36.15 N +ATOM 1724 CA LYS A 278 1.761 -19.567 18.635 1.00 36.15 C +ATOM 1725 C LYS A 278 2.577 -18.687 19.595 1.00 36.15 C +ATOM 1726 O LYS A 278 2.960 -19.158 20.661 1.00 36.15 O +ATOM 1727 CB LYS A 278 2.668 -20.524 17.843 1.00 36.15 C +ATOM 1728 CG LYS A 278 3.482 -21.419 18.784 1.00 36.15 C +ATOM 1729 CD LYS A 278 4.236 -22.534 18.065 1.00 36.15 C +ATOM 1730 CE LYS A 278 4.970 -23.306 19.165 1.00 36.15 C +ATOM 1731 NZ LYS A 278 5.619 -24.536 18.661 1.00 36.15 N +ATOM 1732 N PHE A 279 2.888 -17.452 19.210 1.00 35.83 N +ATOM 1733 CA PHE A 279 3.663 -16.514 20.029 1.00 35.83 C +ATOM 1734 C PHE A 279 2.773 -15.606 20.894 1.00 35.83 C +ATOM 1735 O PHE A 279 3.295 -14.799 21.662 1.00 35.83 O +ATOM 1736 CB PHE A 279 4.613 -15.718 19.126 1.00 35.83 C +ATOM 1737 CG PHE A 279 5.503 -16.586 18.258 1.00 35.83 C +ATOM 1738 CD1 PHE A 279 6.600 -17.266 18.818 1.00 35.83 C +ATOM 1739 CD2 PHE A 279 5.206 -16.744 16.893 1.00 35.83 C +ATOM 1740 CE1 PHE A 279 7.387 -18.114 18.014 1.00 35.83 C +ATOM 1741 CE2 PHE A 279 5.980 -17.597 16.091 1.00 35.83 C +ATOM 1742 CZ PHE A 279 7.068 -18.286 16.655 1.00 35.83 C +ATOM 1743 N MET A 280 1.448 -15.748 20.788 1.00 36.12 N +ATOM 1744 CA MET A 280 0.487 -15.031 21.618 1.00 36.12 C +ATOM 1745 C MET A 280 0.366 -15.686 23.001 1.00 36.12 C +ATOM 1746 O MET A 280 0.330 -16.912 23.117 1.00 36.12 O +ATOM 1747 CB MET A 280 -0.857 -14.968 20.887 1.00 36.12 C +ATOM 1748 CG MET A 280 -1.842 -14.001 21.550 1.00 36.12 C +ATOM 1749 SD MET A 280 -3.379 -13.783 20.613 1.00 36.12 S +ATOM 1750 CE MET A 280 -4.093 -15.446 20.737 1.00 36.12 C +ATOM 1751 N ASN A 281 0.307 -14.873 24.049 1.00 36.26 N +ATOM 1752 CA ASN A 281 0.113 -15.285 25.437 1.00 36.26 C +ATOM 1753 C ASN A 281 -0.682 -14.213 26.208 1.00 36.26 C +ATOM 1754 O ASN A 281 -1.041 -13.173 25.655 1.00 36.26 O +ATOM 1755 CB ASN A 281 1.491 -15.606 26.056 1.00 36.26 C +ATOM 1756 CG ASN A 281 2.430 -14.412 26.131 1.00 36.26 C +ATOM 1757 OD1 ASN A 281 2.036 -13.288 26.372 1.00 36.26 O +ATOM 1758 ND2 ASN A 281 3.713 -14.617 25.956 1.00 36.26 N +ATOM 1759 N ALA A 282 -0.958 -14.461 27.491 1.00 36.22 N +ATOM 1760 CA ALA A 282 -1.737 -13.545 28.327 1.00 36.22 C +ATOM 1761 C ALA A 282 -1.109 -12.141 28.450 1.00 36.22 C +ATOM 1762 O ALA A 282 -1.841 -11.163 28.583 1.00 36.22 O +ATOM 1763 CB ALA A 282 -1.911 -14.193 29.706 1.00 36.22 C +ATOM 1764 N ASP A 283 0.220 -12.035 28.345 1.00 37.08 N +ATOM 1765 CA ASP A 283 0.962 -10.784 28.531 1.00 37.08 C +ATOM 1766 C ASP A 283 1.025 -9.917 27.271 1.00 37.08 C +ATOM 1767 O ASP A 283 1.440 -8.759 27.365 1.00 37.08 O +ATOM 1768 CB ASP A 283 2.405 -11.078 28.974 1.00 37.08 C +ATOM 1769 CG ASP A 283 2.508 -11.859 30.281 1.00 37.08 C +ATOM 1770 OD1 ASP A 283 1.645 -11.648 31.159 1.00 37.08 O +ATOM 1771 OD2 ASP A 283 3.469 -12.652 30.384 1.00 37.08 O +ATOM 1772 N ASN A 284 0.699 -10.463 26.096 1.00 35.96 N +ATOM 1773 CA ASN A 284 0.849 -9.772 24.814 1.00 35.96 C +ATOM 1774 C ASN A 284 -0.402 -9.805 23.919 1.00 35.96 C +ATOM 1775 O ASN A 284 -0.446 -9.081 22.928 1.00 35.96 O +ATOM 1776 CB ASN A 284 2.128 -10.289 24.125 1.00 35.96 C +ATOM 1777 CG ASN A 284 2.032 -11.689 23.538 1.00 35.96 C +ATOM 1778 OD1 ASN A 284 0.972 -12.237 23.291 1.00 35.96 O +ATOM 1779 ND2 ASN A 284 3.150 -12.326 23.266 1.00 35.96 N +ATOM 1780 N GLN A 285 -1.437 -10.579 24.267 1.00 35.99 N +ATOM 1781 CA GLN A 285 -2.658 -10.716 23.459 1.00 35.99 C +ATOM 1782 C GLN A 285 -3.348 -9.380 23.140 1.00 35.99 C +ATOM 1783 O GLN A 285 -3.956 -9.245 22.085 1.00 35.99 O +ATOM 1784 CB GLN A 285 -3.641 -11.680 24.148 1.00 35.99 C +ATOM 1785 CG GLN A 285 -4.110 -11.194 25.533 1.00 35.99 C +ATOM 1786 CD GLN A 285 -5.091 -12.142 26.214 1.00 35.99 C +ATOM 1787 OE1 GLN A 285 -5.478 -13.185 25.714 1.00 35.99 O +ATOM 1788 NE2 GLN A 285 -5.531 -11.810 27.407 1.00 35.99 N +ATOM 1789 N HIS A 286 -3.224 -8.368 24.005 1.00 35.96 N +ATOM 1790 CA HIS A 286 -3.790 -7.031 23.782 1.00 35.96 C +ATOM 1791 C HIS A 286 -3.043 -6.207 22.720 1.00 35.96 C +ATOM 1792 O HIS A 286 -3.537 -5.158 22.316 1.00 35.96 O +ATOM 1793 CB HIS A 286 -3.829 -6.282 25.119 1.00 35.96 C +ATOM 1794 CG HIS A 286 -2.458 -6.011 25.682 1.00 35.96 C +ATOM 1795 ND1 HIS A 286 -1.622 -6.971 26.253 1.00 35.96 N +ATOM 1796 CD2 HIS A 286 -1.816 -4.809 25.681 1.00 35.96 C +ATOM 1797 CE1 HIS A 286 -0.505 -6.318 26.600 1.00 35.96 C +ATOM 1798 NE2 HIS A 286 -0.589 -5.019 26.268 1.00 35.96 N +ATOM 1799 N LEU A 287 -1.866 -6.659 22.278 1.00 35.74 N +ATOM 1800 CA LEU A 287 -1.060 -6.051 21.211 1.00 35.74 C +ATOM 1801 C LEU A 287 -1.189 -6.805 19.878 1.00 35.74 C +ATOM 1802 O LEU A 287 -0.661 -6.364 18.858 1.00 35.74 O +ATOM 1803 CB LEU A 287 0.417 -6.005 21.651 1.00 35.74 C +ATOM 1804 CG LEU A 287 0.688 -5.327 23.004 1.00 35.74 C +ATOM 1805 CD1 LEU A 287 2.175 -5.419 23.345 1.00 35.74 C +ATOM 1806 CD2 LEU A 287 0.260 -3.857 23.016 1.00 35.74 C +ATOM 1807 N VAL A 288 -1.862 -7.960 19.872 1.00 35.86 N +ATOM 1808 CA VAL A 288 -1.957 -8.849 18.712 1.00 35.86 C +ATOM 1809 C VAL A 288 -3.370 -8.786 18.134 1.00 35.86 C +ATOM 1810 O VAL A 288 -4.335 -9.228 18.747 1.00 35.86 O +ATOM 1811 CB VAL A 288 -1.526 -10.283 19.078 1.00 35.86 C +ATOM 1812 CG1 VAL A 288 -1.621 -11.205 17.856 1.00 35.86 C +ATOM 1813 CG2 VAL A 288 -0.066 -10.329 19.559 1.00 35.86 C +ATOM 1814 N SER A 289 -3.485 -8.266 16.913 1.00 35.96 N +ATOM 1815 CA SER A 289 -4.725 -8.225 16.121 1.00 35.96 C +ATOM 1816 C SER A 289 -4.547 -8.968 14.791 1.00 35.96 C +ATOM 1817 O SER A 289 -3.404 -9.165 14.360 1.00 35.96 O +ATOM 1818 CB SER A 289 -5.166 -6.773 15.886 1.00 35.96 C +ATOM 1819 OG SER A 289 -4.124 -6.000 15.312 1.00 35.96 O +ATOM 1820 N PRO A 290 -5.633 -9.373 14.102 1.00 35.68 N +ATOM 1821 CA PRO A 290 -5.539 -9.958 12.763 1.00 35.68 C +ATOM 1822 C PRO A 290 -4.726 -9.088 11.794 1.00 35.68 C +ATOM 1823 O PRO A 290 -3.882 -9.606 11.065 1.00 35.68 O +ATOM 1824 CB PRO A 290 -6.986 -10.128 12.292 1.00 35.68 C +ATOM 1825 CG PRO A 290 -7.764 -10.284 13.598 1.00 35.68 C +ATOM 1826 CD PRO A 290 -7.024 -9.336 14.538 1.00 35.68 C +ATOM 1827 N GLU A 291 -4.899 -7.765 11.847 1.00 35.53 N +ATOM 1828 CA GLU A 291 -4.159 -6.813 11.018 1.00 35.53 C +ATOM 1829 C GLU A 291 -2.672 -6.754 11.386 1.00 35.53 C +ATOM 1830 O GLU A 291 -1.841 -6.594 10.492 1.00 35.53 O +ATOM 1831 CB GLU A 291 -4.746 -5.394 11.139 1.00 35.53 C +ATOM 1832 CG GLU A 291 -6.212 -5.239 10.703 1.00 35.53 C +ATOM 1833 CD GLU A 291 -7.239 -5.797 11.704 1.00 35.53 C +ATOM 1834 OE1 GLU A 291 -8.406 -5.950 11.296 1.00 35.53 O +ATOM 1835 OE2 GLU A 291 -6.851 -6.075 12.865 1.00 35.53 O +ATOM 1836 N ALA A 292 -2.323 -6.888 12.673 1.00 35.53 N +ATOM 1837 CA ALA A 292 -0.929 -6.951 13.118 1.00 35.53 C +ATOM 1838 C ALA A 292 -0.240 -8.210 12.580 1.00 35.53 C +ATOM 1839 O ALA A 292 0.879 -8.143 12.074 1.00 35.53 O +ATOM 1840 CB ALA A 292 -0.865 -6.936 14.652 1.00 35.53 C +ATOM 1841 N ILE A 293 -0.921 -9.357 12.659 1.00 35.50 N +ATOM 1842 CA ILE A 293 -0.395 -10.631 12.168 1.00 35.50 C +ATOM 1843 C ILE A 293 -0.218 -10.610 10.650 1.00 35.50 C +ATOM 1844 O ILE A 293 0.835 -11.004 10.157 1.00 35.50 O +ATOM 1845 CB ILE A 293 -1.299 -11.794 12.634 1.00 35.50 C +ATOM 1846 CG1 ILE A 293 -1.278 -11.978 14.170 1.00 35.50 C +ATOM 1847 CG2 ILE A 293 -0.895 -13.112 11.950 1.00 35.50 C +ATOM 1848 CD1 ILE A 293 0.120 -12.011 14.801 1.00 35.50 C +ATOM 1849 N ASP A 294 -1.215 -10.123 9.916 1.00 35.44 N +ATOM 1850 CA ASP A 294 -1.156 -9.995 8.459 1.00 35.44 C +ATOM 1851 C ASP A 294 -0.049 -9.021 8.015 1.00 35.44 C +ATOM 1852 O ASP A 294 0.724 -9.325 7.103 1.00 35.44 O +ATOM 1853 CB ASP A 294 -2.549 -9.565 7.990 1.00 35.44 C +ATOM 1854 CG ASP A 294 -2.636 -9.351 6.484 1.00 35.44 C +ATOM 1855 OD1 ASP A 294 -2.234 -10.222 5.684 1.00 35.44 O +ATOM 1856 OD2 ASP A 294 -3.097 -8.271 6.071 1.00 35.44 O +ATOM 1857 N PHE A 295 0.108 -7.895 8.720 1.00 35.41 N +ATOM 1858 CA PHE A 295 1.200 -6.953 8.473 1.00 35.41 C +ATOM 1859 C PHE A 295 2.574 -7.598 8.695 1.00 35.41 C +ATOM 1860 O PHE A 295 3.444 -7.520 7.824 1.00 35.41 O +ATOM 1861 CB PHE A 295 1.043 -5.723 9.378 1.00 35.41 C +ATOM 1862 CG PHE A 295 2.126 -4.685 9.180 1.00 35.41 C +ATOM 1863 CD1 PHE A 295 2.883 -4.243 10.276 1.00 35.41 C +ATOM 1864 CD2 PHE A 295 2.419 -4.203 7.893 1.00 35.41 C +ATOM 1865 CE1 PHE A 295 3.952 -3.353 10.089 1.00 35.41 C +ATOM 1866 CE2 PHE A 295 3.480 -3.308 7.706 1.00 35.41 C +ATOM 1867 CZ PHE A 295 4.244 -2.871 8.803 1.00 35.41 C +ATOM 1868 N LEU A 296 2.762 -8.273 9.836 1.00 35.41 N +ATOM 1869 CA LEU A 296 4.016 -8.947 10.165 1.00 35.41 C +ATOM 1870 C LEU A 296 4.360 -10.034 9.143 1.00 35.41 C +ATOM 1871 O LEU A 296 5.514 -10.145 8.720 1.00 35.41 O +ATOM 1872 CB LEU A 296 3.915 -9.551 11.579 1.00 35.41 C +ATOM 1873 CG LEU A 296 5.200 -10.279 12.011 1.00 35.41 C +ATOM 1874 CD1 LEU A 296 6.386 -9.317 12.107 1.00 35.41 C +ATOM 1875 CD2 LEU A 296 5.007 -10.930 13.374 1.00 35.41 C +ATOM 1876 N ASP A 297 3.363 -10.817 8.725 1.00 35.44 N +ATOM 1877 CA ASP A 297 3.545 -11.903 7.767 1.00 35.44 C +ATOM 1878 C ASP A 297 4.074 -11.411 6.421 1.00 35.44 C +ATOM 1879 O ASP A 297 4.891 -12.080 5.790 1.00 35.44 O +ATOM 1880 CB ASP A 297 2.218 -12.633 7.533 1.00 35.44 C +ATOM 1881 CG ASP A 297 2.513 -13.977 6.884 1.00 35.44 C +ATOM 1882 OD1 ASP A 297 2.955 -14.854 7.662 1.00 35.44 O +ATOM 1883 OD2 ASP A 297 2.398 -14.146 5.641 1.00 35.44 O +ATOM 1884 N LYS A 298 3.628 -10.230 5.989 1.00 35.41 N +ATOM 1885 CA LYS A 298 4.044 -9.599 4.733 1.00 35.41 C +ATOM 1886 C LYS A 298 5.409 -8.898 4.819 1.00 35.41 C +ATOM 1887 O LYS A 298 5.960 -8.530 3.786 1.00 35.41 O +ATOM 1888 CB LYS A 298 2.943 -8.632 4.296 1.00 35.41 C +ATOM 1889 CG LYS A 298 1.633 -9.305 3.873 1.00 35.41 C +ATOM 1890 CD LYS A 298 0.539 -8.237 3.773 1.00 35.41 C +ATOM 1891 CE LYS A 298 -0.729 -8.818 3.150 1.00 35.41 C +ATOM 1892 NZ LYS A 298 -1.929 -8.169 3.714 1.00 35.41 N +ATOM 1893 N LEU A 299 5.976 -8.726 6.017 1.00 35.41 N +ATOM 1894 CA LEU A 299 7.342 -8.222 6.217 1.00 35.41 C +ATOM 1895 C LEU A 299 8.365 -9.354 6.351 1.00 35.41 C +ATOM 1896 O LEU A 299 9.423 -9.318 5.724 1.00 35.41 O +ATOM 1897 CB LEU A 299 7.391 -7.339 7.473 1.00 35.41 C +ATOM 1898 CG LEU A 299 6.630 -6.009 7.379 1.00 35.41 C +ATOM 1899 CD1 LEU A 299 6.706 -5.356 8.755 1.00 35.41 C +ATOM 1900 CD2 LEU A 299 7.253 -5.058 6.357 1.00 35.41 C +ATOM 1901 N LEU A 300 8.064 -10.369 7.168 1.00 35.41 N +ATOM 1902 CA LEU A 300 8.979 -11.475 7.468 1.00 35.41 C +ATOM 1903 C LEU A 300 8.879 -12.597 6.425 1.00 35.41 C +ATOM 1904 O LEU A 300 8.481 -13.723 6.721 1.00 35.41 O +ATOM 1905 CB LEU A 300 8.785 -11.955 8.919 1.00 35.41 C +ATOM 1906 CG LEU A 300 9.202 -10.940 10.001 1.00 35.41 C +ATOM 1907 CD1 LEU A 300 9.079 -11.606 11.373 1.00 35.41 C +ATOM 1908 CD2 LEU A 300 10.653 -10.474 9.848 1.00 35.41 C +ATOM 1909 N ARG A 301 9.281 -12.277 5.191 1.00 35.41 N +ATOM 1910 CA ARG A 301 9.434 -13.224 4.074 1.00 35.41 C +ATOM 1911 C ARG A 301 10.906 -13.500 3.791 1.00 35.41 C +ATOM 1912 O ARG A 301 11.724 -12.584 3.857 1.00 35.41 O +ATOM 1913 CB ARG A 301 8.734 -12.689 2.809 1.00 35.41 C +ATOM 1914 CG ARG A 301 7.228 -12.438 2.977 1.00 35.41 C +ATOM 1915 CD ARG A 301 6.492 -13.751 3.267 1.00 35.41 C +ATOM 1916 NE ARG A 301 5.053 -13.535 3.430 1.00 35.41 N +ATOM 1917 CZ ARG A 301 4.134 -13.536 2.491 1.00 35.41 C +ATOM 1918 NH1 ARG A 301 4.411 -13.723 1.223 1.00 35.41 N +ATOM 1919 NH2 ARG A 301 2.893 -13.370 2.859 1.00 35.41 N +ATOM 1920 N TYR A 302 11.235 -14.750 3.449 1.00 35.44 N +ATOM 1921 CA TYR A 302 12.582 -15.095 2.983 1.00 35.44 C +ATOM 1922 C TYR A 302 12.938 -14.301 1.739 1.00 35.44 C +ATOM 1923 O TYR A 302 13.914 -13.555 1.750 1.00 35.44 O +ATOM 1924 CB TYR A 302 12.710 -16.587 2.676 1.00 35.44 C +ATOM 1925 CG TYR A 302 12.722 -17.444 3.909 1.00 35.44 C +ATOM 1926 CD1 TYR A 302 13.830 -17.388 4.772 1.00 35.44 C +ATOM 1927 CD2 TYR A 302 11.649 -18.312 4.166 1.00 35.44 C +ATOM 1928 CE1 TYR A 302 13.877 -18.218 5.902 1.00 35.44 C +ATOM 1929 CE2 TYR A 302 11.698 -19.159 5.282 1.00 35.44 C +ATOM 1930 CZ TYR A 302 12.817 -19.112 6.136 1.00 35.44 C +ATOM 1931 OH TYR A 302 12.900 -19.995 7.147 1.00 35.44 O +ATOM 1932 N ASP A 303 12.087 -14.421 0.724 1.00 35.50 N +ATOM 1933 CA ASP A 303 12.221 -13.694 -0.520 1.00 35.50 C +ATOM 1934 C ASP A 303 12.040 -12.198 -0.279 1.00 35.50 C +ATOM 1935 O ASP A 303 10.955 -11.711 0.049 1.00 35.50 O +ATOM 1936 CB ASP A 303 11.225 -14.255 -1.532 1.00 35.50 C +ATOM 1937 CG ASP A 303 11.536 -13.771 -2.941 1.00 35.50 C +ATOM 1938 OD1 ASP A 303 12.058 -12.650 -3.106 1.00 35.50 O +ATOM 1939 OD2 ASP A 303 11.293 -14.550 -3.883 1.00 35.50 O +ATOM 1940 N HIS A 304 13.139 -11.469 -0.425 1.00 35.50 N +ATOM 1941 CA HIS A 304 13.207 -10.031 -0.236 1.00 35.50 C +ATOM 1942 C HIS A 304 12.340 -9.266 -1.240 1.00 35.50 C +ATOM 1943 O HIS A 304 11.836 -8.195 -0.909 1.00 35.50 O +ATOM 1944 CB HIS A 304 14.691 -9.651 -0.268 1.00 35.50 C +ATOM 1945 CG HIS A 304 15.373 -10.012 -1.555 1.00 35.50 C +ATOM 1946 ND1 HIS A 304 15.813 -11.297 -1.862 1.00 35.50 N +ATOM 1947 CD2 HIS A 304 15.663 -9.175 -2.590 1.00 35.50 C +ATOM 1948 CE1 HIS A 304 16.385 -11.212 -3.064 1.00 35.50 C +ATOM 1949 NE2 HIS A 304 16.307 -9.955 -3.523 1.00 35.50 N +ATOM 1950 N GLN A 305 12.046 -9.850 -2.406 1.00 35.65 N +ATOM 1951 CA GLN A 305 11.128 -9.270 -3.388 1.00 35.65 C +ATOM 1952 C GLN A 305 9.656 -9.392 -2.976 1.00 35.65 C +ATOM 1953 O GLN A 305 8.832 -8.620 -3.457 1.00 35.65 O +ATOM 1954 CB GLN A 305 11.336 -9.933 -4.758 1.00 35.65 C +ATOM 1955 CG GLN A 305 12.746 -9.710 -5.322 1.00 35.65 C +ATOM 1956 CD GLN A 305 12.987 -8.270 -5.735 1.00 35.65 C +ATOM 1957 OE1 GLN A 305 12.111 -7.601 -6.257 1.00 35.65 O +ATOM 1958 NE2 GLN A 305 14.179 -7.758 -5.566 1.00 35.65 N +ATOM 1959 N GLU A 306 9.296 -10.319 -2.088 1.00 35.62 N +ATOM 1960 CA GLU A 306 7.909 -10.483 -1.622 1.00 35.62 C +ATOM 1961 C GLU A 306 7.553 -9.594 -0.427 1.00 35.62 C +ATOM 1962 O GLU A 306 6.377 -9.468 -0.078 1.00 35.62 O +ATOM 1963 CB GLU A 306 7.644 -11.937 -1.230 1.00 35.62 C +ATOM 1964 CG GLU A 306 7.670 -12.907 -2.414 1.00 35.62 C +ATOM 1965 CD GLU A 306 7.238 -14.322 -1.994 1.00 35.62 C +ATOM 1966 OE1 GLU A 306 7.359 -15.223 -2.853 1.00 35.62 O +ATOM 1967 OE2 GLU A 306 6.749 -14.496 -0.842 1.00 35.62 O +ATOM 1968 N ARG A 307 8.552 -8.999 0.230 1.00 35.44 N +ATOM 1969 CA ARG A 307 8.336 -8.132 1.396 1.00 35.44 C +ATOM 1970 C ARG A 307 7.666 -6.841 0.959 1.00 35.44 C +ATOM 1971 O ARG A 307 8.035 -6.309 -0.086 1.00 35.44 O +ATOM 1972 CB ARG A 307 9.660 -7.814 2.099 1.00 35.44 C +ATOM 1973 CG ARG A 307 10.404 -9.082 2.524 1.00 35.44 C +ATOM 1974 CD ARG A 307 11.721 -8.725 3.211 1.00 35.44 C +ATOM 1975 NE ARG A 307 12.595 -9.907 3.308 1.00 35.44 N +ATOM 1976 CZ ARG A 307 13.910 -9.910 3.242 1.00 35.44 C +ATOM 1977 NH1 ARG A 307 14.612 -8.813 3.234 1.00 35.44 N +ATOM 1978 NH2 ARG A 307 14.550 -11.033 3.160 1.00 35.44 N +ATOM 1979 N LEU A 308 6.754 -6.303 1.768 1.00 35.41 N +ATOM 1980 CA LEU A 308 6.185 -4.967 1.534 1.00 35.41 C +ATOM 1981 C LEU A 308 7.295 -3.922 1.378 1.00 35.41 C +ATOM 1982 O LEU A 308 8.279 -3.944 2.118 1.00 35.41 O +ATOM 1983 CB LEU A 308 5.266 -4.537 2.693 1.00 35.41 C +ATOM 1984 CG LEU A 308 4.025 -5.415 2.864 1.00 35.41 C +ATOM 1985 CD1 LEU A 308 3.335 -5.114 4.197 1.00 35.41 C +ATOM 1986 CD2 LEU A 308 2.996 -5.223 1.750 1.00 35.41 C +ATOM 1987 N THR A 309 7.110 -2.974 0.463 1.00 35.38 N +ATOM 1988 CA THR A 309 7.872 -1.718 0.482 1.00 35.38 C +ATOM 1989 C THR A 309 7.441 -0.858 1.671 1.00 35.38 C +ATOM 1990 O THR A 309 6.394 -1.103 2.274 1.00 35.38 O +ATOM 1991 CB THR A 309 7.720 -0.916 -0.818 1.00 35.38 C +ATOM 1992 OG1 THR A 309 6.420 -0.398 -0.952 1.00 35.38 O +ATOM 1993 CG2 THR A 309 8.001 -1.749 -2.061 1.00 35.38 C +ATOM 1994 N ALA A 310 8.220 0.170 2.013 1.00 35.41 N +ATOM 1995 CA ALA A 310 7.841 1.099 3.075 1.00 35.41 C +ATOM 1996 C ALA A 310 6.492 1.791 2.779 1.00 35.41 C +ATOM 1997 O ALA A 310 5.614 1.803 3.642 1.00 35.41 O +ATOM 1998 CB ALA A 310 8.988 2.085 3.284 1.00 35.41 C +ATOM 1999 N LEU A 311 6.270 2.250 1.542 1.00 35.50 N +ATOM 2000 CA LEU A 311 4.987 2.814 1.113 1.00 35.50 C +ATOM 2001 C LEU A 311 3.827 1.803 1.170 1.00 35.50 C +ATOM 2002 O LEU A 311 2.757 2.113 1.699 1.00 35.50 O +ATOM 2003 CB LEU A 311 5.158 3.385 -0.306 1.00 35.50 C +ATOM 2004 CG LEU A 311 3.939 4.163 -0.827 1.00 35.50 C +ATOM 2005 CD1 LEU A 311 3.605 5.372 0.050 1.00 35.50 C +ATOM 2006 CD2 LEU A 311 4.224 4.659 -2.245 1.00 35.50 C +ATOM 2007 N GLU A 312 4.021 0.575 0.676 1.00 35.41 N +ATOM 2008 CA GLU A 312 3.002 -0.485 0.770 1.00 35.41 C +ATOM 2009 C GLU A 312 2.670 -0.790 2.239 1.00 35.41 C +ATOM 2010 O GLU A 312 1.500 -0.878 2.613 1.00 35.41 O +ATOM 2011 CB GLU A 312 3.493 -1.765 0.080 1.00 35.41 C +ATOM 2012 CG GLU A 312 3.518 -1.691 -1.457 1.00 35.41 C +ATOM 2013 CD GLU A 312 4.351 -2.823 -2.092 1.00 35.41 C +ATOM 2014 OE1 GLU A 312 4.429 -2.902 -3.336 1.00 35.41 O +ATOM 2015 OE2 GLU A 312 5.025 -3.590 -1.369 1.00 35.41 O +ATOM 2016 N ALA A 313 3.684 -0.858 3.102 1.00 35.41 N +ATOM 2017 CA ALA A 313 3.515 -1.022 4.538 1.00 35.41 C +ATOM 2018 C ALA A 313 2.704 0.126 5.159 1.00 35.41 C +ATOM 2019 O ALA A 313 1.754 -0.142 5.891 1.00 35.41 O +ATOM 2020 CB ALA A 313 4.902 -1.180 5.169 1.00 35.41 C +ATOM 2021 N MET A 314 2.984 1.387 4.818 1.00 35.47 N +ATOM 2022 CA MET A 314 2.224 2.544 5.317 1.00 35.47 C +ATOM 2023 C MET A 314 0.747 2.524 4.911 1.00 35.47 C +ATOM 2024 O MET A 314 -0.099 3.057 5.632 1.00 35.47 O +ATOM 2025 CB MET A 314 2.853 3.848 4.814 1.00 35.47 C +ATOM 2026 CG MET A 314 4.157 4.155 5.542 1.00 35.47 C +ATOM 2027 SD MET A 314 4.995 5.664 4.996 1.00 35.47 S +ATOM 2028 CE MET A 314 3.804 6.914 5.541 1.00 35.47 C +ATOM 2029 N THR A 315 0.406 1.898 3.783 1.00 35.59 N +ATOM 2030 CA THR A 315 -0.988 1.766 3.318 1.00 35.59 C +ATOM 2031 C THR A 315 -1.740 0.585 3.941 1.00 35.59 C +ATOM 2032 O THR A 315 -2.968 0.526 3.851 1.00 35.59 O +ATOM 2033 CB THR A 315 -1.090 1.743 1.786 1.00 35.59 C +ATOM 2034 OG1 THR A 315 -0.395 0.664 1.224 1.00 35.59 O +ATOM 2035 CG2 THR A 315 -0.575 3.035 1.154 1.00 35.59 C +ATOM 2036 N HIS A 316 -1.051 -0.301 4.670 1.00 35.41 N +ATOM 2037 CA HIS A 316 -1.641 -1.485 5.297 1.00 35.41 C +ATOM 2038 C HIS A 316 -2.808 -1.141 6.258 1.00 35.41 C +ATOM 2039 O HIS A 316 -2.741 -0.118 6.956 1.00 35.41 O +ATOM 2040 CB HIS A 316 -0.543 -2.267 6.037 1.00 35.41 C +ATOM 2041 CG HIS A 316 -0.944 -3.686 6.317 1.00 35.41 C +ATOM 2042 ND1 HIS A 316 -1.509 -4.195 7.467 1.00 35.41 N +ATOM 2043 CD2 HIS A 316 -0.858 -4.706 5.416 1.00 35.41 C +ATOM 2044 CE1 HIS A 316 -1.788 -5.493 7.242 1.00 35.41 C +ATOM 2045 NE2 HIS A 316 -1.394 -5.843 6.007 1.00 35.41 N +ATOM 2046 N PRO A 317 -3.863 -1.982 6.363 1.00 35.50 N +ATOM 2047 CA PRO A 317 -4.978 -1.773 7.299 1.00 35.50 C +ATOM 2048 C PRO A 317 -4.560 -1.580 8.763 1.00 35.50 C +ATOM 2049 O PRO A 317 -5.189 -0.807 9.481 1.00 35.50 O +ATOM 2050 CB PRO A 317 -5.872 -3.006 7.145 1.00 35.50 C +ATOM 2051 CG PRO A 317 -5.659 -3.400 5.687 1.00 35.50 C +ATOM 2052 CD PRO A 317 -4.177 -3.098 5.473 1.00 35.50 C +ATOM 2053 N TYR A 318 -3.446 -2.189 9.187 1.00 35.47 N +ATOM 2054 CA TYR A 318 -2.897 -2.021 10.542 1.00 35.47 C +ATOM 2055 C TYR A 318 -2.682 -0.542 10.931 1.00 35.47 C +ATOM 2056 O TYR A 318 -2.874 -0.178 12.091 1.00 35.47 O +ATOM 2057 CB TYR A 318 -1.593 -2.831 10.678 1.00 35.47 C +ATOM 2058 CG TYR A 318 -1.028 -2.874 12.090 1.00 35.47 C +ATOM 2059 CD1 TYR A 318 0.225 -2.298 12.379 1.00 35.47 C +ATOM 2060 CD2 TYR A 318 -1.764 -3.484 13.125 1.00 35.47 C +ATOM 2061 CE1 TYR A 318 0.727 -2.318 13.697 1.00 35.47 C +ATOM 2062 CE2 TYR A 318 -1.283 -3.466 14.449 1.00 35.47 C +ATOM 2063 CZ TYR A 318 -0.035 -2.886 14.743 1.00 35.47 C +ATOM 2064 OH TYR A 318 0.407 -2.823 16.029 1.00 35.47 O +ATOM 2065 N PHE A 319 -2.389 0.341 9.970 1.00 35.47 N +ATOM 2066 CA PHE A 319 -2.171 1.773 10.223 1.00 35.47 C +ATOM 2067 C PHE A 319 -3.388 2.664 9.943 1.00 35.47 C +ATOM 2068 O PHE A 319 -3.270 3.886 10.005 1.00 35.47 O +ATOM 2069 CB PHE A 319 -0.932 2.256 9.464 1.00 35.47 C +ATOM 2070 CG PHE A 319 0.296 1.433 9.752 1.00 35.47 C +ATOM 2071 CD1 PHE A 319 0.920 1.518 11.008 1.00 35.47 C +ATOM 2072 CD2 PHE A 319 0.795 0.562 8.772 1.00 35.47 C +ATOM 2073 CE1 PHE A 319 2.052 0.737 11.278 1.00 35.47 C +ATOM 2074 CE2 PHE A 319 1.922 -0.225 9.047 1.00 35.47 C +ATOM 2075 CZ PHE A 319 2.548 -0.138 10.298 1.00 35.47 C +ATOM 2076 N GLN A 320 -4.566 2.100 9.653 1.00 35.77 N +ATOM 2077 CA GLN A 320 -5.760 2.892 9.327 1.00 35.77 C +ATOM 2078 C GLN A 320 -6.103 3.919 10.418 1.00 35.77 C +ATOM 2079 O GLN A 320 -6.396 5.068 10.096 1.00 35.77 O +ATOM 2080 CB GLN A 320 -6.937 1.943 9.065 1.00 35.77 C +ATOM 2081 CG GLN A 320 -8.225 2.691 8.673 1.00 35.77 C +ATOM 2082 CD GLN A 320 -9.381 1.747 8.355 1.00 35.77 C +ATOM 2083 OE1 GLN A 320 -9.255 0.539 8.340 1.00 35.77 O +ATOM 2084 NE2 GLN A 320 -10.555 2.264 8.071 1.00 35.77 N +ATOM 2085 N GLN A 321 -6.015 3.533 11.694 1.00 36.26 N +ATOM 2086 CA GLN A 321 -6.288 4.430 12.824 1.00 36.26 C +ATOM 2087 C GLN A 321 -5.279 5.583 12.909 1.00 36.26 C +ATOM 2088 O GLN A 321 -5.671 6.728 13.120 1.00 36.26 O +ATOM 2089 CB GLN A 321 -6.276 3.637 14.136 1.00 36.26 C +ATOM 2090 CG GLN A 321 -7.402 2.594 14.232 1.00 36.26 C +ATOM 2091 CD GLN A 321 -7.352 1.821 15.549 1.00 36.26 C +ATOM 2092 OE1 GLN A 321 -6.565 2.099 16.436 1.00 36.26 O +ATOM 2093 NE2 GLN A 321 -8.177 0.812 15.722 1.00 36.26 N +ATOM 2094 N VAL A 322 -3.990 5.300 12.683 1.00 35.90 N +ATOM 2095 CA VAL A 322 -2.926 6.317 12.690 1.00 35.90 C +ATOM 2096 C VAL A 322 -3.143 7.316 11.554 1.00 35.90 C +ATOM 2097 O VAL A 322 -3.162 8.521 11.791 1.00 35.90 O +ATOM 2098 CB VAL A 322 -1.530 5.669 12.596 1.00 35.90 C +ATOM 2099 CG1 VAL A 322 -0.418 6.721 12.623 1.00 35.90 C +ATOM 2100 CG2 VAL A 322 -1.291 4.701 13.761 1.00 35.90 C +ATOM 2101 N ARG A 323 -3.415 6.828 10.334 1.00 36.06 N +ATOM 2102 CA ARG A 323 -3.719 7.689 9.176 1.00 36.06 C +ATOM 2103 C ARG A 323 -4.964 8.549 9.401 1.00 36.06 C +ATOM 2104 O ARG A 323 -4.981 9.718 9.027 1.00 36.06 O +ATOM 2105 CB ARG A 323 -3.891 6.841 7.906 1.00 36.06 C +ATOM 2106 CG ARG A 323 -2.588 6.152 7.472 1.00 36.06 C +ATOM 2107 CD ARG A 323 -2.674 5.575 6.051 1.00 36.06 C +ATOM 2108 NE ARG A 323 -3.789 4.615 5.882 1.00 36.06 N +ATOM 2109 CZ ARG A 323 -3.748 3.306 6.069 1.00 36.06 C +ATOM 2110 NH1 ARG A 323 -2.680 2.676 6.442 1.00 36.06 N +ATOM 2111 NH2 ARG A 323 -4.794 2.562 5.857 1.00 36.06 N +ATOM 2112 N ALA A 324 -6.009 7.991 10.014 1.00 36.46 N +ATOM 2113 CA ALA A 324 -7.222 8.739 10.340 1.00 36.46 C +ATOM 2114 C ALA A 324 -6.948 9.865 11.355 1.00 36.46 C +ATOM 2115 O ALA A 324 -7.437 10.985 11.178 1.00 36.46 O +ATOM 2116 CB ALA A 324 -8.286 7.757 10.845 1.00 36.46 C +ATOM 2117 N ALA A 325 -6.131 9.594 12.377 1.00 37.47 N +ATOM 2118 CA ALA A 325 -5.729 10.589 13.364 1.00 37.47 C +ATOM 2119 C ALA A 325 -4.879 11.716 12.747 1.00 37.47 C +ATOM 2120 O ALA A 325 -5.115 12.886 13.046 1.00 37.47 O +ATOM 2121 CB ALA A 325 -4.993 9.875 14.503 1.00 37.47 C +ATOM 2122 N GLU A 326 -3.944 11.396 11.847 1.00 40.46 N +ATOM 2123 CA GLU A 326 -3.152 12.404 11.123 1.00 40.46 C +ATOM 2124 C GLU A 326 -4.012 13.292 10.230 1.00 40.46 C +ATOM 2125 O GLU A 326 -3.943 14.514 10.331 1.00 40.46 O +ATOM 2126 CB GLU A 326 -2.077 11.736 10.263 1.00 40.46 C +ATOM 2127 CG GLU A 326 -0.951 11.214 11.146 1.00 40.46 C +ATOM 2128 CD GLU A 326 0.229 10.687 10.338 1.00 40.46 C +ATOM 2129 OE1 GLU A 326 1.256 10.433 10.988 1.00 40.46 O +ATOM 2130 OE2 GLU A 326 0.179 10.557 9.093 1.00 40.46 O +ATOM 2131 N ASN A 327 -4.885 12.695 9.415 1.00 42.61 N +ATOM 2132 CA ASN A 327 -5.769 13.446 8.523 1.00 42.61 C +ATOM 2133 C ASN A 327 -6.712 14.391 9.279 1.00 42.61 C +ATOM 2134 O ASN A 327 -7.091 15.434 8.750 1.00 42.61 O +ATOM 2135 CB ASN A 327 -6.581 12.447 7.687 1.00 42.61 C +ATOM 2136 CG ASN A 327 -5.761 11.777 6.603 1.00 42.61 C +ATOM 2137 OD1 ASN A 327 -4.758 12.275 6.129 1.00 42.61 O +ATOM 2138 ND2 ASN A 327 -6.200 10.642 6.117 1.00 42.61 N +ATOM 2139 N SER A 328 -7.093 14.030 10.506 1.00 43.08 N +ATOM 2140 CA SER A 328 -7.922 14.880 11.362 1.00 43.08 C +ATOM 2141 C SER A 328 -7.139 16.081 11.902 1.00 43.08 C +ATOM 2142 O SER A 328 -7.699 17.167 11.995 1.00 43.08 O +ATOM 2143 CB SER A 328 -8.509 14.062 12.514 1.00 43.08 C +ATOM 2144 OG SER A 328 -9.235 12.952 12.014 1.00 43.08 O +ATOM 2145 N ALA B 23 34.168 11.310 -4.257 1.00 36.93 N +ATOM 2146 CA ALA B 23 34.467 12.641 -4.771 1.00 36.93 C +ATOM 2147 C ALA B 23 34.080 13.788 -3.814 1.00 36.93 C +ATOM 2148 O ALA B 23 34.485 14.924 -4.040 1.00 36.93 O +ATOM 2149 CB ALA B 23 33.779 12.778 -6.133 1.00 36.93 C +ATOM 2150 N LEU B 24 33.352 13.520 -2.718 1.00 36.60 N +ATOM 2151 CA LEU B 24 32.918 14.563 -1.778 1.00 36.60 C +ATOM 2152 C LEU B 24 34.114 15.316 -1.164 1.00 36.60 C +ATOM 2153 O LEU B 24 34.985 14.709 -0.527 1.00 36.60 O +ATOM 2154 CB LEU B 24 32.022 13.945 -0.683 1.00 36.60 C +ATOM 2155 CG LEU B 24 31.565 14.941 0.404 1.00 36.60 C +ATOM 2156 CD1 LEU B 24 30.627 16.013 -0.147 1.00 36.60 C +ATOM 2157 CD2 LEU B 24 30.851 14.211 1.536 1.00 36.60 C +ATOM 2158 N THR B 25 34.100 16.644 -1.279 1.00 36.82 N +ATOM 2159 CA THR B 25 34.960 17.554 -0.511 1.00 36.82 C +ATOM 2160 C THR B 25 34.204 18.007 0.733 1.00 36.82 C +ATOM 2161 O THR B 25 33.169 18.657 0.633 1.00 36.82 O +ATOM 2162 CB THR B 25 35.405 18.753 -1.360 1.00 36.82 C +ATOM 2163 OG1 THR B 25 36.138 18.262 -2.458 1.00 36.82 O +ATOM 2164 CG2 THR B 25 36.339 19.698 -0.604 1.00 36.82 C +ATOM 2165 N VAL B 26 34.707 17.630 1.909 1.00 37.79 N +ATOM 2166 CA VAL B 26 34.093 17.975 3.199 1.00 37.79 C +ATOM 2167 C VAL B 26 34.607 19.341 3.644 1.00 37.79 C +ATOM 2168 O VAL B 26 35.820 19.524 3.763 1.00 37.79 O +ATOM 2169 CB VAL B 26 34.378 16.899 4.265 1.00 37.79 C +ATOM 2170 CG1 VAL B 26 33.704 17.223 5.599 1.00 37.79 C +ATOM 2171 CG2 VAL B 26 33.884 15.514 3.813 1.00 37.79 C +ATOM 2172 N GLN B 27 33.686 20.269 3.894 1.00 37.87 N +ATOM 2173 CA GLN B 27 33.957 21.500 4.631 1.00 37.87 C +ATOM 2174 C GLN B 27 33.898 21.170 6.121 1.00 37.87 C +ATOM 2175 O GLN B 27 32.899 20.623 6.580 1.00 37.87 O +ATOM 2176 CB GLN B 27 32.931 22.576 4.250 1.00 37.87 C +ATOM 2177 CG GLN B 27 33.077 23.018 2.784 1.00 37.87 C +ATOM 2178 CD GLN B 27 32.074 24.098 2.383 1.00 37.87 C +ATOM 2179 OE1 GLN B 27 31.163 24.458 3.101 1.00 37.87 O +ATOM 2180 NE2 GLN B 27 32.183 24.649 1.194 1.00 37.87 N +ATOM 2181 N TRP B 28 34.995 21.415 6.833 1.00 37.30 N +ATOM 2182 CA TRP B 28 35.108 21.121 8.260 1.00 37.30 C +ATOM 2183 C TRP B 28 34.841 22.396 9.055 1.00 37.30 C +ATOM 2184 O TRP B 28 35.487 23.406 8.779 1.00 37.30 O +ATOM 2185 CB TRP B 28 36.500 20.549 8.573 1.00 37.30 C +ATOM 2186 CG TRP B 28 36.857 19.290 7.840 1.00 37.30 C +ATOM 2187 CD1 TRP B 28 37.592 19.216 6.709 1.00 37.30 C +ATOM 2188 CD2 TRP B 28 36.422 17.929 8.119 1.00 37.30 C +ATOM 2189 NE1 TRP B 28 37.653 17.907 6.274 1.00 37.30 N +ATOM 2190 CE2 TRP B 28 36.963 17.064 7.121 1.00 37.30 C +ATOM 2191 CE3 TRP B 28 35.593 17.344 9.094 1.00 37.30 C +ATOM 2192 CZ2 TRP B 28 36.710 15.682 7.106 1.00 37.30 C +ATOM 2193 CZ3 TRP B 28 35.332 15.964 9.094 1.00 37.30 C +ATOM 2194 CH2 TRP B 28 35.886 15.131 8.106 1.00 37.30 C +ATOM 2195 N GLY B 29 33.920 22.331 10.012 1.00 37.47 N +ATOM 2196 CA GLY B 29 33.741 23.345 11.049 1.00 37.47 C +ATOM 2197 C GLY B 29 34.752 23.199 12.189 1.00 37.47 C +ATOM 2198 O GLY B 29 35.539 22.240 12.233 1.00 37.47 O +ATOM 2199 N GLU B 30 34.703 24.145 13.124 1.00 39.47 N +ATOM 2200 CA GLU B 30 35.582 24.196 14.291 1.00 39.47 C +ATOM 2201 C GLU B 30 35.010 23.352 15.433 1.00 39.47 C +ATOM 2202 O GLU B 30 33.866 23.513 15.849 1.00 39.47 O +ATOM 2203 CB GLU B 30 35.813 25.655 14.724 1.00 39.47 C +ATOM 2204 CG GLU B 30 36.558 26.498 13.670 1.00 39.47 C +ATOM 2205 CD GLU B 30 37.959 25.963 13.320 1.00 39.47 C +ATOM 2206 OE1 GLU B 30 38.478 26.299 12.232 1.00 39.47 O +ATOM 2207 OE2 GLU B 30 38.539 25.187 14.117 1.00 39.47 O +ATOM 2208 N GLN B 31 35.803 22.411 15.955 1.00 38.52 N +ATOM 2209 CA GLN B 31 35.336 21.555 17.051 1.00 38.52 C +ATOM 2210 C GLN B 31 35.195 22.327 18.371 1.00 38.52 C +ATOM 2211 O GLN B 31 34.383 21.939 19.207 1.00 38.52 O +ATOM 2212 CB GLN B 31 36.266 20.341 17.188 1.00 38.52 C +ATOM 2213 CG GLN B 31 35.674 19.289 18.137 1.00 38.52 C +ATOM 2214 CD GLN B 31 36.570 18.072 18.272 1.00 38.52 C +ATOM 2215 OE1 GLN B 31 36.665 17.228 17.386 1.00 38.52 O +ATOM 2216 NE2 GLN B 31 37.248 17.927 19.388 1.00 38.52 N +ATOM 2217 N ASP B 32 35.948 23.417 18.535 1.00 38.04 N +ATOM 2218 CA ASP B 32 35.962 24.241 19.749 1.00 38.04 C +ATOM 2219 C ASP B 32 34.642 25.006 19.967 1.00 38.04 C +ATOM 2220 O ASP B 32 34.375 25.464 21.075 1.00 38.04 O +ATOM 2221 CB ASP B 32 37.176 25.189 19.706 1.00 38.04 C +ATOM 2222 CG ASP B 32 38.537 24.471 19.739 1.00 38.04 C +ATOM 2223 OD1 ASP B 32 38.582 23.258 20.063 1.00 38.04 O +ATOM 2224 OD2 ASP B 32 39.551 25.132 19.421 1.00 38.04 O +ATOM 2225 N ASP B 33 33.769 25.055 18.952 1.00 36.86 N +ATOM 2226 CA ASP B 33 32.392 25.552 19.070 1.00 36.86 C +ATOM 2227 C ASP B 33 31.487 24.628 19.915 1.00 36.86 C +ATOM 2228 O ASP B 33 30.337 24.972 20.204 1.00 36.86 O +ATOM 2229 CB ASP B 33 31.788 25.723 17.663 1.00 36.86 C +ATOM 2230 CG ASP B 33 32.363 26.893 16.857 1.00 36.86 C +ATOM 2231 OD1 ASP B 33 32.827 27.872 17.480 1.00 36.86 O +ATOM 2232 OD2 ASP B 33 32.208 26.872 15.616 1.00 36.86 O +ATOM 2233 N TYR B 34 31.970 23.443 20.306 1.00 35.93 N +ATOM 2234 CA TYR B 34 31.175 22.420 20.980 1.00 35.93 C +ATOM 2235 C TYR B 34 31.846 21.911 22.257 1.00 35.93 C +ATOM 2236 O TYR B 34 32.911 21.293 22.227 1.00 35.93 O +ATOM 2237 CB TYR B 34 30.882 21.267 20.013 1.00 35.93 C +ATOM 2238 CG TYR B 34 30.166 21.697 18.748 1.00 35.93 C +ATOM 2239 CD1 TYR B 34 28.763 21.817 18.740 1.00 35.93 C +ATOM 2240 CD2 TYR B 34 30.910 22.052 17.606 1.00 35.93 C +ATOM 2241 CE1 TYR B 34 28.106 22.268 17.579 1.00 35.93 C +ATOM 2242 CE2 TYR B 34 30.252 22.516 16.453 1.00 35.93 C +ATOM 2243 CZ TYR B 34 28.849 22.604 16.427 1.00 35.93 C +ATOM 2244 OH TYR B 34 28.213 23.008 15.299 1.00 35.93 O +ATOM 2245 N GLU B 35 31.160 22.072 23.386 1.00 36.33 N +ATOM 2246 CA GLU B 35 31.604 21.554 24.679 1.00 36.33 C +ATOM 2247 C GLU B 35 30.961 20.191 24.976 1.00 36.33 C +ATOM 2248 O GLU B 35 29.744 20.014 24.872 1.00 36.33 O +ATOM 2249 CB GLU B 35 31.330 22.584 25.781 1.00 36.33 C +ATOM 2250 CG GLU B 35 32.011 22.179 27.098 1.00 36.33 C +ATOM 2251 CD GLU B 35 31.755 23.162 28.251 1.00 36.33 C +ATOM 2252 OE1 GLU B 35 32.009 22.751 29.407 1.00 36.33 O +ATOM 2253 OE2 GLU B 35 31.299 24.296 27.988 1.00 36.33 O +ATOM 2254 N VAL B 36 31.775 19.205 25.367 1.00 36.26 N +ATOM 2255 CA VAL B 36 31.300 17.877 25.782 1.00 36.26 C +ATOM 2256 C VAL B 36 30.783 17.936 27.218 1.00 36.26 C +ATOM 2257 O VAL B 36 31.545 18.186 28.143 1.00 36.26 O +ATOM 2258 CB VAL B 36 32.412 16.815 25.651 1.00 36.26 C +ATOM 2259 CG1 VAL B 36 32.008 15.468 26.272 1.00 36.26 C +ATOM 2260 CG2 VAL B 36 32.755 16.557 24.178 1.00 36.26 C +ATOM 2261 N VAL B 37 29.512 17.587 27.419 1.00 35.99 N +ATOM 2262 CA VAL B 37 28.890 17.528 28.751 1.00 35.99 C +ATOM 2263 C VAL B 37 29.049 16.140 29.365 1.00 35.99 C +ATOM 2264 O VAL B 37 29.523 15.984 30.487 1.00 35.99 O +ATOM 2265 CB VAL B 37 27.406 17.934 28.671 1.00 35.99 C +ATOM 2266 CG1 VAL B 37 26.682 17.781 30.016 1.00 35.99 C +ATOM 2267 CG2 VAL B 37 27.260 19.390 28.221 1.00 35.99 C +ATOM 2268 N ARG B 38 28.629 15.094 28.643 1.00 36.12 N +ATOM 2269 CA ARG B 38 28.696 13.712 29.143 1.00 36.12 C +ATOM 2270 C ARG B 38 28.710 12.693 28.019 1.00 36.12 C +ATOM 2271 O ARG B 38 28.138 12.895 26.952 1.00 36.12 O +ATOM 2272 CB ARG B 38 27.542 13.428 30.128 1.00 36.12 C +ATOM 2273 CG ARG B 38 26.160 13.515 29.472 1.00 36.12 C +ATOM 2274 CD ARG B 38 25.040 13.206 30.462 1.00 36.12 C +ATOM 2275 NE ARG B 38 23.748 13.246 29.763 1.00 36.12 N +ATOM 2276 CZ ARG B 38 22.955 12.253 29.422 1.00 36.12 C +ATOM 2277 NH1 ARG B 38 23.200 11.010 29.727 1.00 36.12 N +ATOM 2278 NH2 ARG B 38 21.892 12.523 28.730 1.00 36.12 N +ATOM 2279 N LYS B 39 29.297 11.530 28.290 1.00 36.36 N +ATOM 2280 CA LYS B 39 29.219 10.377 27.388 1.00 36.36 C +ATOM 2281 C LYS B 39 27.819 9.764 27.439 1.00 36.36 C +ATOM 2282 O LYS B 39 27.318 9.477 28.523 1.00 36.36 O +ATOM 2283 CB LYS B 39 30.322 9.382 27.760 1.00 36.36 C +ATOM 2284 CG LYS B 39 30.535 8.325 26.668 1.00 36.36 C +ATOM 2285 CD LYS B 39 31.815 7.535 26.966 1.00 36.36 C +ATOM 2286 CE LYS B 39 32.200 6.652 25.776 1.00 36.36 C +ATOM 2287 NZ LYS B 39 33.639 6.286 25.837 1.00 36.36 N +ATOM 2288 N VAL B 40 27.214 9.529 26.276 1.00 37.75 N +ATOM 2289 CA VAL B 40 25.878 8.908 26.163 1.00 37.75 C +ATOM 2290 C VAL B 40 25.913 7.539 25.498 1.00 37.75 C +ATOM 2291 O VAL B 40 25.022 6.727 25.717 1.00 37.75 O +ATOM 2292 CB VAL B 40 24.863 9.826 25.455 1.00 37.75 C +ATOM 2293 CG1 VAL B 40 24.679 11.130 26.233 1.00 37.75 C +ATOM 2294 CG2 VAL B 40 25.246 10.162 24.012 1.00 37.75 C +ATOM 2295 N GLY B 41 26.951 7.235 24.716 1.00 40.96 N +ATOM 2296 CA GLY B 41 27.029 5.953 24.030 1.00 40.96 C +ATOM 2297 C GLY B 41 28.410 5.615 23.490 1.00 40.96 C +ATOM 2298 O GLY B 41 29.298 6.459 23.348 1.00 40.96 O +ATOM 2299 N ARG B 42 28.597 4.335 23.167 1.00 40.18 N +ATOM 2300 CA ARG B 42 29.803 3.830 22.508 1.00 40.18 C +ATOM 2301 C ARG B 42 29.414 2.823 21.438 1.00 40.18 C +ATOM 2302 O ARG B 42 28.839 1.783 21.740 1.00 40.18 O +ATOM 2303 CB ARG B 42 30.745 3.225 23.558 1.00 40.18 C +ATOM 2304 CG ARG B 42 32.130 2.922 22.974 1.00 40.18 C +ATOM 2305 CD ARG B 42 33.014 2.278 24.047 1.00 40.18 C +ATOM 2306 NE ARG B 42 34.438 2.272 23.661 1.00 40.18 N +ATOM 2307 CZ ARG B 42 35.340 1.386 24.061 1.00 40.18 C +ATOM 2308 NH1 ARG B 42 35.063 0.446 24.920 1.00 40.18 N +ATOM 2309 NH2 ARG B 42 36.563 1.439 23.619 1.00 40.18 N +ATOM 2310 N GLY B 43 29.774 3.121 20.197 1.00 47.54 N +ATOM 2311 CA GLY B 43 29.597 2.228 19.060 1.00 47.54 C +ATOM 2312 C GLY B 43 30.897 1.530 18.663 1.00 47.54 C +ATOM 2313 O GLY B 43 31.998 1.834 19.143 1.00 47.54 O +ATOM 2314 N LYS B 44 30.790 0.611 17.700 1.00 46.58 N +ATOM 2315 CA LYS B 44 31.960 -0.027 17.075 1.00 46.58 C +ATOM 2316 C LYS B 44 32.865 1.002 16.387 1.00 46.58 C +ATOM 2317 O LYS B 44 34.083 0.887 16.472 1.00 46.58 O +ATOM 2318 CB LYS B 44 31.476 -1.109 16.099 1.00 46.58 C +ATOM 2319 CG LYS B 44 32.643 -1.888 15.471 1.00 46.58 C +ATOM 2320 CD LYS B 44 32.114 -3.067 14.647 1.00 46.58 C +ATOM 2321 CE LYS B 44 33.274 -3.825 13.991 1.00 46.58 C +ATOM 2322 NZ LYS B 44 32.773 -4.932 13.139 1.00 46.58 N +ATOM 2323 N TYR B 45 32.268 2.009 15.752 1.00 40.67 N +ATOM 2324 CA TYR B 45 32.966 2.990 14.914 1.00 40.67 C +ATOM 2325 C TYR B 45 32.959 4.415 15.476 1.00 40.67 C +ATOM 2326 O TYR B 45 33.548 5.299 14.866 1.00 40.67 O +ATOM 2327 CB TYR B 45 32.365 2.951 13.503 1.00 40.67 C +ATOM 2328 CG TYR B 45 32.384 1.571 12.871 1.00 40.67 C +ATOM 2329 CD1 TYR B 45 33.612 0.996 12.496 1.00 40.67 C +ATOM 2330 CD2 TYR B 45 31.187 0.849 12.697 1.00 40.67 C +ATOM 2331 CE1 TYR B 45 33.651 -0.291 11.930 1.00 40.67 C +ATOM 2332 CE2 TYR B 45 31.220 -0.449 12.150 1.00 40.67 C +ATOM 2333 CZ TYR B 45 32.454 -1.019 11.764 1.00 40.67 C +ATOM 2334 OH TYR B 45 32.498 -2.259 11.211 1.00 40.67 O +ATOM 2335 N SER B 46 32.323 4.647 16.624 1.00 37.71 N +ATOM 2336 CA SER B 46 32.205 5.984 17.204 1.00 37.71 C +ATOM 2337 C SER B 46 32.113 5.959 18.724 1.00 37.71 C +ATOM 2338 O SER B 46 31.836 4.923 19.340 1.00 37.71 O +ATOM 2339 CB SER B 46 30.984 6.711 16.628 1.00 37.71 C +ATOM 2340 OG SER B 46 29.797 5.988 16.911 1.00 37.71 O +ATOM 2341 N GLU B 47 32.336 7.118 19.322 1.00 36.53 N +ATOM 2342 CA GLU B 47 31.928 7.445 20.684 1.00 36.53 C +ATOM 2343 C GLU B 47 30.949 8.617 20.618 1.00 36.53 C +ATOM 2344 O GLU B 47 31.087 9.481 19.757 1.00 36.53 O +ATOM 2345 CB GLU B 47 33.151 7.721 21.567 1.00 36.53 C +ATOM 2346 CG GLU B 47 34.120 6.529 21.542 1.00 36.53 C +ATOM 2347 CD GLU B 47 34.946 6.429 22.821 1.00 36.53 C +ATOM 2348 OE1 GLU B 47 34.801 5.378 23.508 1.00 36.53 O +ATOM 2349 OE2 GLU B 47 35.666 7.378 23.156 1.00 36.53 O +ATOM 2350 N VAL B 48 29.910 8.590 21.451 1.00 36.06 N +ATOM 2351 CA VAL B 48 28.791 9.533 21.357 1.00 36.06 C +ATOM 2352 C VAL B 48 28.635 10.255 22.684 1.00 36.06 C +ATOM 2353 O VAL B 48 28.599 9.620 23.745 1.00 36.06 O +ATOM 2354 CB VAL B 48 27.477 8.845 20.938 1.00 36.06 C +ATOM 2355 CG1 VAL B 48 26.410 9.886 20.590 1.00 36.06 C +ATOM 2356 CG2 VAL B 48 27.664 7.960 19.696 1.00 36.06 C +ATOM 2357 N PHE B 49 28.526 11.574 22.606 1.00 35.80 N +ATOM 2358 CA PHE B 49 28.440 12.471 23.748 1.00 35.80 C +ATOM 2359 C PHE B 49 27.249 13.415 23.590 1.00 35.80 C +ATOM 2360 O PHE B 49 26.901 13.794 22.475 1.00 35.80 O +ATOM 2361 CB PHE B 49 29.761 13.238 23.895 1.00 35.80 C +ATOM 2362 CG PHE B 49 30.991 12.355 24.024 1.00 35.80 C +ATOM 2363 CD1 PHE B 49 31.526 12.059 25.290 1.00 35.80 C +ATOM 2364 CD2 PHE B 49 31.608 11.832 22.871 1.00 35.80 C +ATOM 2365 CE1 PHE B 49 32.662 11.235 25.406 1.00 35.80 C +ATOM 2366 CE2 PHE B 49 32.733 11.001 22.988 1.00 35.80 C +ATOM 2367 CZ PHE B 49 33.265 10.701 24.254 1.00 35.80 C +ATOM 2368 N GLU B 50 26.629 13.787 24.706 1.00 35.71 N +ATOM 2369 CA GLU B 50 25.817 15.001 24.783 1.00 35.71 C +ATOM 2370 C GLU B 50 26.773 16.181 24.940 1.00 35.71 C +ATOM 2371 O GLU B 50 27.740 16.101 25.707 1.00 35.71 O +ATOM 2372 CB GLU B 50 24.813 14.905 25.947 1.00 35.71 C +ATOM 2373 CG GLU B 50 23.934 16.163 26.144 1.00 35.71 C +ATOM 2374 CD GLU B 50 22.887 16.025 27.274 1.00 35.71 C +ATOM 2375 OE1 GLU B 50 22.189 17.012 27.618 1.00 35.71 O +ATOM 2376 OE2 GLU B 50 22.730 14.904 27.816 1.00 35.71 O +ATOM 2377 N GLY B 51 26.500 17.253 24.208 1.00 35.74 N +ATOM 2378 CA GLY B 51 27.261 18.487 24.266 1.00 35.74 C +ATOM 2379 C GLY B 51 26.378 19.716 24.115 1.00 35.74 C +ATOM 2380 O GLY B 51 25.156 19.611 23.968 1.00 35.74 O +ATOM 2381 N ILE B 52 27.011 20.881 24.146 1.00 35.68 N +ATOM 2382 CA ILE B 52 26.380 22.187 23.955 1.00 35.68 C +ATOM 2383 C ILE B 52 27.173 22.934 22.885 1.00 35.68 C +ATOM 2384 O ILE B 52 28.402 22.912 22.908 1.00 35.68 O +ATOM 2385 CB ILE B 52 26.311 22.958 25.294 1.00 35.68 C +ATOM 2386 CG1 ILE B 52 25.418 22.199 26.307 1.00 35.68 C +ATOM 2387 CG2 ILE B 52 25.767 24.379 25.076 1.00 35.68 C +ATOM 2388 CD1 ILE B 52 25.452 22.768 27.730 1.00 35.68 C +ATOM 2389 N ASN B 53 26.480 23.573 21.945 1.00 35.77 N +ATOM 2390 CA ASN B 53 27.104 24.552 21.061 1.00 35.77 C +ATOM 2391 C ASN B 53 27.294 25.860 21.846 1.00 35.77 C +ATOM 2392 O ASN B 53 26.316 26.474 22.273 1.00 35.77 O +ATOM 2393 CB ASN B 53 26.240 24.689 19.800 1.00 35.77 C +ATOM 2394 CG ASN B 53 26.765 25.728 18.817 1.00 35.77 C +ATOM 2395 OD1 ASN B 53 27.338 26.746 19.162 1.00 35.77 O +ATOM 2396 ND2 ASN B 53 26.594 25.508 17.535 1.00 35.77 N +ATOM 2397 N VAL B 54 28.541 26.280 22.059 1.00 36.33 N +ATOM 2398 CA VAL B 54 28.876 27.419 22.935 1.00 36.33 C +ATOM 2399 C VAL B 54 28.456 28.768 22.349 1.00 36.33 C +ATOM 2400 O VAL B 54 28.293 29.738 23.085 1.00 36.33 O +ATOM 2401 CB VAL B 54 30.373 27.437 23.301 1.00 36.33 C +ATOM 2402 CG1 VAL B 54 30.799 26.121 23.967 1.00 36.33 C +ATOM 2403 CG2 VAL B 54 31.277 27.725 22.095 1.00 36.33 C +ATOM 2404 N ASN B 55 28.228 28.836 21.034 1.00 36.60 N +ATOM 2405 CA ASN B 55 27.842 30.066 20.350 1.00 36.60 C +ATOM 2406 C ASN B 55 26.356 30.406 20.538 1.00 36.60 C +ATOM 2407 O ASN B 55 25.979 31.574 20.452 1.00 36.60 O +ATOM 2408 CB ASN B 55 28.214 29.942 18.861 1.00 36.60 C +ATOM 2409 CG ASN B 55 29.713 29.811 18.658 1.00 36.60 C +ATOM 2410 OD1 ASN B 55 30.481 30.596 19.187 1.00 36.60 O +ATOM 2411 ND2 ASN B 55 30.163 28.850 17.888 1.00 36.60 N +ATOM 2412 N ASN B 56 25.499 29.406 20.781 1.00 36.43 N +ATOM 2413 CA ASN B 56 24.050 29.600 20.910 1.00 36.43 C +ATOM 2414 C ASN B 56 23.405 28.890 22.118 1.00 36.43 C +ATOM 2415 O ASN B 56 22.200 29.022 22.323 1.00 36.43 O +ATOM 2416 CB ASN B 56 23.380 29.240 19.568 1.00 36.43 C +ATOM 2417 CG ASN B 56 23.391 27.758 19.241 1.00 36.43 C +ATOM 2418 OD1 ASN B 56 23.857 26.926 19.998 1.00 36.43 O +ATOM 2419 ND2 ASN B 56 22.821 27.381 18.121 1.00 36.43 N +ATOM 2420 N ASN B 57 24.187 28.173 22.932 1.00 36.64 N +ATOM 2421 CA ASN B 57 23.748 27.366 24.074 1.00 36.64 C +ATOM 2422 C ASN B 57 22.743 26.245 23.739 1.00 36.64 C +ATOM 2423 O ASN B 57 22.067 25.727 24.633 1.00 36.64 O +ATOM 2424 CB ASN B 57 23.302 28.277 25.235 1.00 36.64 C +ATOM 2425 CG ASN B 57 24.454 29.042 25.847 1.00 36.64 C +ATOM 2426 OD1 ASN B 57 25.515 28.510 26.113 1.00 36.64 O +ATOM 2427 ND2 ASN B 57 24.277 30.308 26.139 1.00 36.64 N +ATOM 2428 N GLU B 58 22.636 25.826 22.476 1.00 35.96 N +ATOM 2429 CA GLU B 58 21.772 24.708 22.099 1.00 35.96 C +ATOM 2430 C GLU B 58 22.424 23.358 22.420 1.00 35.96 C +ATOM 2431 O GLU B 58 23.607 23.121 22.158 1.00 35.96 O +ATOM 2432 CB GLU B 58 21.354 24.770 20.624 1.00 35.96 C +ATOM 2433 CG GLU B 58 20.395 25.936 20.332 1.00 35.96 C +ATOM 2434 CD GLU B 58 19.871 25.921 18.886 1.00 35.96 C +ATOM 2435 OE1 GLU B 58 18.802 26.526 18.647 1.00 35.96 O +ATOM 2436 OE2 GLU B 58 20.531 25.309 18.011 1.00 35.96 O +ATOM 2437 N LYS B 59 21.624 22.428 22.955 1.00 35.80 N +ATOM 2438 CA LYS B 59 22.053 21.042 23.165 1.00 35.80 C +ATOM 2439 C LYS B 59 22.253 20.325 21.833 1.00 35.80 C +ATOM 2440 O LYS B 59 21.415 20.411 20.936 1.00 35.80 O +ATOM 2441 CB LYS B 59 21.042 20.262 24.008 1.00 35.80 C +ATOM 2442 CG LYS B 59 21.005 20.724 25.467 1.00 35.80 C +ATOM 2443 CD LYS B 59 20.162 19.728 26.265 1.00 35.80 C +ATOM 2444 CE LYS B 59 20.258 20.016 27.763 1.00 35.80 C +ATOM 2445 NZ LYS B 59 20.148 18.748 28.515 1.00 35.80 N +ATOM 2446 N CYS B 60 23.300 19.514 21.755 1.00 35.62 N +ATOM 2447 CA CYS B 60 23.604 18.694 20.588 1.00 35.62 C +ATOM 2448 C CYS B 60 24.119 17.302 20.976 1.00 35.62 C +ATOM 2449 O CYS B 60 24.403 17.005 22.141 1.00 35.62 O +ATOM 2450 CB CYS B 60 24.592 19.456 19.689 1.00 35.62 C +ATOM 2451 SG CYS B 60 26.201 19.644 20.507 1.00 35.62 S +ATOM 2452 N ILE B 61 24.236 16.434 19.972 1.00 35.65 N +ATOM 2453 CA ILE B 61 24.909 15.142 20.087 1.00 35.65 C +ATOM 2454 C ILE B 61 26.190 15.166 19.259 1.00 35.65 C +ATOM 2455 O ILE B 61 26.155 15.382 18.050 1.00 35.65 O +ATOM 2456 CB ILE B 61 23.962 13.994 19.685 1.00 35.65 C +ATOM 2457 CG1 ILE B 61 22.739 13.883 20.619 1.00 35.65 C +ATOM 2458 CG2 ILE B 61 24.710 12.654 19.612 1.00 35.65 C +ATOM 2459 CD1 ILE B 61 23.040 13.578 22.093 1.00 35.65 C +ATOM 2460 N ILE B 62 27.315 14.879 19.907 1.00 35.74 N +ATOM 2461 CA ILE B 62 28.646 14.842 19.302 1.00 35.74 C +ATOM 2462 C ILE B 62 29.024 13.377 19.076 1.00 35.74 C +ATOM 2463 O ILE B 62 29.297 12.630 20.021 1.00 35.74 O +ATOM 2464 CB ILE B 62 29.676 15.585 20.184 1.00 35.74 C +ATOM 2465 CG1 ILE B 62 29.205 17.014 20.553 1.00 35.74 C +ATOM 2466 CG2 ILE B 62 31.040 15.614 19.465 1.00 35.74 C +ATOM 2467 CD1 ILE B 62 30.078 17.711 21.602 1.00 35.74 C +ATOM 2468 N LYS B 63 29.043 12.939 17.814 1.00 35.93 N +ATOM 2469 CA LYS B 63 29.515 11.607 17.413 1.00 35.93 C +ATOM 2470 C LYS B 63 30.961 11.706 16.945 1.00 35.93 C +ATOM 2471 O LYS B 63 31.223 12.001 15.781 1.00 35.93 O +ATOM 2472 CB LYS B 63 28.584 11.039 16.335 1.00 35.93 C +ATOM 2473 CG LYS B 63 28.926 9.593 15.931 1.00 35.93 C +ATOM 2474 CD LYS B 63 28.014 9.112 14.792 1.00 35.93 C +ATOM 2475 CE LYS B 63 28.211 7.623 14.465 1.00 35.93 C +ATOM 2476 NZ LYS B 63 27.277 7.187 13.400 1.00 35.93 N +ATOM 2477 N ILE B 64 31.893 11.394 17.837 1.00 36.36 N +ATOM 2478 CA ILE B 64 33.321 11.320 17.519 1.00 36.36 C +ATOM 2479 C ILE B 64 33.580 10.023 16.751 1.00 36.36 C +ATOM 2480 O ILE B 64 33.304 8.924 17.249 1.00 36.36 O +ATOM 2481 CB ILE B 64 34.189 11.426 18.791 1.00 36.36 C +ATOM 2482 CG1 ILE B 64 33.797 12.674 19.617 1.00 36.36 C +ATOM 2483 CG2 ILE B 64 35.675 11.453 18.387 1.00 36.36 C +ATOM 2484 CD1 ILE B 64 34.700 12.962 20.823 1.00 36.36 C +ATOM 2485 N LEU B 65 34.095 10.126 15.527 1.00 37.16 N +ATOM 2486 CA LEU B 65 34.362 8.967 14.679 1.00 37.16 C +ATOM 2487 C LEU B 65 35.731 8.370 15.022 1.00 37.16 C +ATOM 2488 O LEU B 65 36.764 9.031 14.933 1.00 37.16 O +ATOM 2489 CB LEU B 65 34.254 9.357 13.197 1.00 37.16 C +ATOM 2490 CG LEU B 65 32.901 9.965 12.774 1.00 37.16 C +ATOM 2491 CD1 LEU B 65 32.962 10.329 11.292 1.00 37.16 C +ATOM 2492 CD2 LEU B 65 31.738 8.989 12.977 1.00 37.16 C +ATOM 2493 N LYS B 66 35.754 7.080 15.374 1.00 38.75 N +ATOM 2494 CA LYS B 66 37.003 6.332 15.583 1.00 38.75 C +ATOM 2495 C LYS B 66 37.777 6.221 14.263 1.00 38.75 C +ATOM 2496 O LYS B 66 37.165 6.334 13.199 1.00 38.75 O +ATOM 2497 CB LYS B 66 36.711 4.939 16.162 1.00 38.75 C +ATOM 2498 CG LYS B 66 36.115 5.022 17.571 1.00 38.75 C +ATOM 2499 CD LYS B 66 35.845 3.620 18.121 1.00 38.75 C +ATOM 2500 CE LYS B 66 35.161 3.742 19.483 1.00 38.75 C +ATOM 2501 NZ LYS B 66 34.646 2.435 19.949 1.00 38.75 N +ATOM 2502 N PRO B 67 39.088 5.910 14.294 1.00 39.47 N +ATOM 2503 CA PRO B 67 39.868 5.688 13.080 1.00 39.47 C +ATOM 2504 C PRO B 67 39.207 4.665 12.141 1.00 39.47 C +ATOM 2505 O PRO B 67 39.157 3.465 12.414 1.00 39.47 O +ATOM 2506 CB PRO B 67 41.257 5.253 13.563 1.00 39.47 C +ATOM 2507 CG PRO B 67 41.376 5.932 14.926 1.00 39.47 C +ATOM 2508 CD PRO B 67 39.951 5.854 15.468 1.00 39.47 C +ATOM 2509 N VAL B 68 38.674 5.153 11.022 1.00 40.96 N +ATOM 2510 CA VAL B 68 38.045 4.360 9.959 1.00 40.96 C +ATOM 2511 C VAL B 68 38.474 4.894 8.593 1.00 40.96 C +ATOM 2512 O VAL B 68 38.956 6.020 8.465 1.00 40.96 O +ATOM 2513 CB VAL B 68 36.502 4.315 10.075 1.00 40.96 C +ATOM 2514 CG1 VAL B 68 36.021 3.576 11.326 1.00 40.96 C +ATOM 2515 CG2 VAL B 68 35.838 5.691 10.012 1.00 40.96 C +ATOM 2516 N LYS B 69 38.299 4.091 7.536 1.00 39.32 N +ATOM 2517 CA LYS B 69 38.620 4.510 6.163 1.00 39.32 C +ATOM 2518 C LYS B 69 37.864 5.798 5.812 1.00 39.32 C +ATOM 2519 O LYS B 69 36.641 5.831 5.944 1.00 39.32 O +ATOM 2520 CB LYS B 69 38.277 3.397 5.151 1.00 39.32 C +ATOM 2521 CG LYS B 69 39.141 2.134 5.312 1.00 39.32 C +ATOM 2522 CD LYS B 69 38.866 1.111 4.194 1.00 39.32 C +ATOM 2523 CE LYS B 69 39.778 -0.118 4.346 1.00 39.32 C +ATOM 2524 NZ LYS B 69 39.636 -1.084 3.222 1.00 39.32 N +ATOM 2525 N LYS B 70 38.563 6.803 5.261 1.00 38.93 N +ATOM 2526 CA LYS B 70 37.972 8.085 4.813 1.00 38.93 C +ATOM 2527 C LYS B 70 36.733 7.892 3.931 1.00 38.93 C +ATOM 2528 O LYS B 70 35.750 8.600 4.109 1.00 38.93 O +ATOM 2529 CB LYS B 70 39.012 8.929 4.049 1.00 38.93 C +ATOM 2530 CG LYS B 70 40.131 9.495 4.940 1.00 38.93 C +ATOM 2531 CD LYS B 70 41.085 10.376 4.114 1.00 38.93 C +ATOM 2532 CE LYS B 70 42.202 10.970 4.985 1.00 38.93 C +ATOM 2533 NZ LYS B 70 43.173 11.765 4.181 1.00 38.93 N +ATOM 2534 N LYS B 71 36.748 6.880 3.053 1.00 37.19 N +ATOM 2535 CA LYS B 71 35.604 6.479 2.216 1.00 37.19 C +ATOM 2536 C LYS B 71 34.311 6.272 3.022 1.00 37.19 C +ATOM 2537 O LYS B 71 33.256 6.724 2.602 1.00 37.19 O +ATOM 2538 CB LYS B 71 35.993 5.218 1.416 1.00 37.19 C +ATOM 2539 CG LYS B 71 34.849 4.740 0.511 1.00 37.19 C +ATOM 2540 CD LYS B 71 35.236 3.606 -0.449 1.00 37.19 C +ATOM 2541 CE LYS B 71 33.998 3.145 -1.237 1.00 37.19 C +ATOM 2542 NZ LYS B 71 34.360 2.399 -2.464 1.00 37.19 N +ATOM 2543 N LYS B 72 34.389 5.629 4.194 1.00 37.19 N +ATOM 2544 CA LYS B 72 33.223 5.372 5.055 1.00 37.19 C +ATOM 2545 C LYS B 72 32.685 6.661 5.686 1.00 37.19 C +ATOM 2546 O LYS B 72 31.475 6.834 5.742 1.00 37.19 O +ATOM 2547 CB LYS B 72 33.580 4.315 6.114 1.00 37.19 C +ATOM 2548 CG LYS B 72 32.332 3.868 6.886 1.00 37.19 C +ATOM 2549 CD LYS B 72 32.616 2.718 7.858 1.00 37.19 C +ATOM 2550 CE LYS B 72 31.290 2.380 8.547 1.00 37.19 C +ATOM 2551 NZ LYS B 72 31.390 1.244 9.490 1.00 37.19 N +ATOM 2552 N ILE B 73 33.583 7.556 6.100 1.00 37.08 N +ATOM 2553 CA ILE B 73 33.236 8.869 6.666 1.00 37.08 C +ATOM 2554 C ILE B 73 32.522 9.717 5.614 1.00 37.08 C +ATOM 2555 O ILE B 73 31.425 10.205 5.861 1.00 37.08 O +ATOM 2556 CB ILE B 73 34.495 9.592 7.201 1.00 37.08 C +ATOM 2557 CG1 ILE B 73 35.234 8.686 8.208 1.00 37.08 C +ATOM 2558 CG2 ILE B 73 34.125 10.955 7.806 1.00 37.08 C +ATOM 2559 CD1 ILE B 73 36.489 9.302 8.829 1.00 37.08 C +ATOM 2560 N LYS B 74 33.103 9.828 4.410 1.00 36.03 N +ATOM 2561 CA LYS B 74 32.483 10.569 3.306 1.00 36.03 C +ATOM 2562 C LYS B 74 31.115 10.000 2.925 1.00 36.03 C +ATOM 2563 O LYS B 74 30.205 10.774 2.659 1.00 36.03 O +ATOM 2564 CB LYS B 74 33.372 10.588 2.062 1.00 36.03 C +ATOM 2565 CG LYS B 74 34.698 11.354 2.190 1.00 36.03 C +ATOM 2566 CD LYS B 74 35.293 11.498 0.781 1.00 36.03 C +ATOM 2567 CE LYS B 74 36.618 12.261 0.725 1.00 36.03 C +ATOM 2568 NZ LYS B 74 36.879 12.693 -0.671 1.00 36.03 N +ATOM 2569 N ASP B 94 26.289 19.920 11.919 1.00 35.96 N +ATOM 2570 CA ASP B 94 27.644 20.333 11.558 1.00 35.96 C +ATOM 2571 C ASP B 94 28.586 19.124 11.413 1.00 35.96 C +ATOM 2572 O ASP B 94 28.298 18.010 11.871 1.00 35.96 O +ATOM 2573 CB ASP B 94 28.155 21.320 12.621 1.00 35.96 C +ATOM 2574 CG ASP B 94 29.331 22.189 12.174 1.00 35.96 C +ATOM 2575 OD1 ASP B 94 29.777 22.043 11.012 1.00 35.96 O +ATOM 2576 OD2 ASP B 94 29.786 22.964 13.038 1.00 35.96 O +ATOM 2577 N ILE B 95 29.721 19.341 10.755 1.00 35.96 N +ATOM 2578 CA ILE B 95 30.783 18.366 10.527 1.00 35.96 C +ATOM 2579 C ILE B 95 32.093 19.017 10.953 1.00 35.96 C +ATOM 2580 O ILE B 95 32.688 19.789 10.205 1.00 35.96 O +ATOM 2581 CB ILE B 95 30.839 17.950 9.043 1.00 35.96 C +ATOM 2582 CG1 ILE B 95 29.499 17.412 8.501 1.00 35.96 C +ATOM 2583 CG2 ILE B 95 31.935 16.883 8.862 1.00 35.96 C +ATOM 2584 CD1 ILE B 95 29.430 17.522 6.975 1.00 35.96 C +ATOM 2585 N VAL B 96 32.582 18.664 12.133 1.00 36.09 N +ATOM 2586 CA VAL B 96 33.757 19.311 12.723 1.00 36.09 C +ATOM 2587 C VAL B 96 34.961 18.389 12.745 1.00 36.09 C +ATOM 2588 O VAL B 96 34.859 17.165 12.577 1.00 36.09 O +ATOM 2589 CB VAL B 96 33.458 19.897 14.109 1.00 36.09 C +ATOM 2590 CG1 VAL B 96 32.289 20.871 14.055 1.00 36.09 C +ATOM 2591 CG2 VAL B 96 33.165 18.822 15.156 1.00 36.09 C +ATOM 2592 N ARG B 97 36.137 18.976 12.944 1.00 37.62 N +ATOM 2593 CA ARG B 97 37.380 18.223 13.070 1.00 37.62 C +ATOM 2594 C ARG B 97 38.289 18.862 14.100 1.00 37.62 C +ATOM 2595 O ARG B 97 38.619 20.032 13.972 1.00 37.62 O +ATOM 2596 CB ARG B 97 38.039 18.169 11.692 1.00 37.62 C +ATOM 2597 CG ARG B 97 39.144 17.121 11.618 1.00 37.62 C +ATOM 2598 CD ARG B 97 39.669 17.089 10.184 1.00 37.62 C +ATOM 2599 NE ARG B 97 40.589 15.962 9.984 1.00 37.62 N +ATOM 2600 CZ ARG B 97 41.357 15.761 8.937 1.00 37.62 C +ATOM 2601 NH1 ARG B 97 41.368 16.594 7.935 1.00 37.62 N +ATOM 2602 NH2 ARG B 97 42.134 14.716 8.895 1.00 37.62 N +ATOM 2603 N ASP B 98 38.772 18.063 15.044 1.00 39.32 N +ATOM 2604 CA ASP B 98 39.797 18.523 15.976 1.00 39.32 C +ATOM 2605 C ASP B 98 41.063 18.956 15.216 1.00 39.32 C +ATOM 2606 O ASP B 98 41.601 18.204 14.390 1.00 39.32 O +ATOM 2607 CB ASP B 98 40.115 17.418 16.985 1.00 39.32 C +ATOM 2608 CG ASP B 98 41.265 17.828 17.899 1.00 39.32 C +ATOM 2609 OD1 ASP B 98 42.174 16.988 18.058 1.00 39.32 O +ATOM 2610 OD2 ASP B 98 41.331 18.987 18.326 1.00 39.32 O +ATOM 2611 N GLN B 99 41.559 20.162 15.491 1.00 42.61 N +ATOM 2612 CA GLN B 99 42.705 20.715 14.777 1.00 42.61 C +ATOM 2613 C GLN B 99 44.002 19.958 15.086 1.00 42.61 C +ATOM 2614 O GLN B 99 44.845 19.836 14.192 1.00 42.61 O +ATOM 2615 CB GLN B 99 42.883 22.203 15.087 1.00 42.61 C +ATOM 2616 CG GLN B 99 41.755 23.100 14.540 1.00 42.61 C +ATOM 2617 CD GLN B 99 42.141 24.581 14.566 1.00 42.61 C +ATOM 2618 OE1 GLN B 99 43.278 24.942 14.833 1.00 42.61 O +ATOM 2619 NE2 GLN B 99 41.256 25.475 14.210 1.00 42.61 N +ATOM 2620 N HIS B 100 44.152 19.398 16.290 1.00 40.91 N +ATOM 2621 CA HIS B 100 45.373 18.705 16.706 1.00 40.91 C +ATOM 2622 C HIS B 100 45.394 17.240 16.252 1.00 40.91 C +ATOM 2623 O HIS B 100 46.227 16.845 15.435 1.00 40.91 O +ATOM 2624 CB HIS B 100 45.547 18.842 18.224 1.00 40.91 C +ATOM 2625 CG HIS B 100 45.954 20.232 18.634 1.00 40.91 C +ATOM 2626 ND1 HIS B 100 45.117 21.256 19.020 1.00 40.91 N +ATOM 2627 CD2 HIS B 100 47.233 20.718 18.676 1.00 40.91 C +ATOM 2628 CE1 HIS B 100 45.877 22.333 19.285 1.00 40.91 C +ATOM 2629 NE2 HIS B 100 47.177 22.050 19.088 1.00 40.91 N +ATOM 2630 N SER B 101 44.453 16.426 16.730 1.00 39.67 N +ATOM 2631 CA SER B 101 44.360 14.993 16.421 1.00 39.67 C +ATOM 2632 C SER B 101 43.822 14.705 15.023 1.00 39.67 C +ATOM 2633 O SER B 101 43.928 13.576 14.540 1.00 39.67 O +ATOM 2634 CB SER B 101 43.490 14.262 17.449 1.00 39.67 C +ATOM 2635 OG SER B 101 42.123 14.584 17.273 1.00 39.67 O +ATOM 2636 N LYS B 102 43.224 15.706 14.359 1.00 40.03 N +ATOM 2637 CA LYS B 102 42.567 15.558 13.056 1.00 40.03 C +ATOM 2638 C LYS B 102 41.410 14.549 13.073 1.00 40.03 C +ATOM 2639 O LYS B 102 41.035 14.049 12.002 1.00 40.03 O +ATOM 2640 CB LYS B 102 43.629 15.290 11.973 1.00 40.03 C +ATOM 2641 CG LYS B 102 44.749 16.334 11.843 1.00 40.03 C +ATOM 2642 CD LYS B 102 44.214 17.756 11.612 1.00 40.03 C +ATOM 2643 CE LYS B 102 45.372 18.723 11.336 1.00 40.03 C +ATOM 2644 NZ LYS B 102 44.952 20.140 11.490 1.00 40.03 N +ATOM 2645 N THR B 103 40.841 14.275 14.248 1.00 38.21 N +ATOM 2646 CA THR B 103 39.703 13.369 14.441 1.00 38.21 C +ATOM 2647 C THR B 103 38.406 14.050 13.996 1.00 38.21 C +ATOM 2648 O THR B 103 38.123 15.153 14.456 1.00 38.21 O +ATOM 2649 CB THR B 103 39.578 12.915 15.902 1.00 38.21 C +ATOM 2650 OG1 THR B 103 40.799 12.335 16.297 1.00 38.21 O +ATOM 2651 CG2 THR B 103 38.506 11.839 16.081 1.00 38.21 C +ATOM 2652 N PRO B 104 37.628 13.441 13.085 1.00 36.60 N +ATOM 2653 CA PRO B 104 36.365 14.006 12.636 1.00 36.60 C +ATOM 2654 C PRO B 104 35.216 13.662 13.589 1.00 36.60 C +ATOM 2655 O PRO B 104 35.111 12.527 14.068 1.00 36.60 O +ATOM 2656 CB PRO B 104 36.158 13.430 11.239 1.00 36.60 C +ATOM 2657 CG PRO B 104 36.858 12.074 11.287 1.00 36.60 C +ATOM 2658 CD PRO B 104 37.946 12.231 12.345 1.00 36.60 C +ATOM 2659 N SER B 105 34.313 14.618 13.778 1.00 36.06 N +ATOM 2660 CA SER B 105 33.100 14.467 14.578 1.00 36.06 C +ATOM 2661 C SER B 105 31.884 14.957 13.797 1.00 36.06 C +ATOM 2662 O SER B 105 31.958 15.921 13.038 1.00 36.06 O +ATOM 2663 CB SER B 105 33.231 15.202 15.917 1.00 36.06 C +ATOM 2664 OG SER B 105 34.361 14.731 16.632 1.00 36.06 O +ATOM 2665 N LEU B 106 30.758 14.270 13.963 1.00 35.80 N +ATOM 2666 CA LEU B 106 29.472 14.670 13.392 1.00 35.80 C +ATOM 2667 C LEU B 106 28.594 15.228 14.508 1.00 35.80 C +ATOM 2668 O LEU B 106 28.442 14.573 15.543 1.00 35.80 O +ATOM 2669 CB LEU B 106 28.805 13.475 12.692 1.00 35.80 C +ATOM 2670 CG LEU B 106 29.633 12.798 11.589 1.00 35.80 C +ATOM 2671 CD1 LEU B 106 28.839 11.640 10.980 1.00 35.80 C +ATOM 2672 CD2 LEU B 106 30.019 13.756 10.466 1.00 35.80 C +ATOM 2673 N ILE B 107 28.031 16.413 14.290 1.00 35.65 N +ATOM 2674 CA ILE B 107 27.170 17.094 15.254 1.00 35.65 C +ATOM 2675 C ILE B 107 25.719 16.908 14.827 1.00 35.65 C +ATOM 2676 O ILE B 107 25.355 17.235 13.701 1.00 35.65 O +ATOM 2677 CB ILE B 107 27.531 18.587 15.363 1.00 35.65 C +ATOM 2678 CG1 ILE B 107 29.046 18.847 15.545 1.00 35.65 C +ATOM 2679 CG2 ILE B 107 26.712 19.232 16.492 1.00 35.65 C +ATOM 2680 CD1 ILE B 107 29.685 18.176 16.766 1.00 35.65 C +ATOM 2681 N PHE B 108 24.882 16.393 15.718 1.00 35.62 N +ATOM 2682 CA PHE B 108 23.472 16.124 15.450 1.00 35.62 C +ATOM 2683 C PHE B 108 22.555 16.906 16.384 1.00 35.62 C +ATOM 2684 O PHE B 108 22.962 17.311 17.477 1.00 35.62 O +ATOM 2685 CB PHE B 108 23.195 14.622 15.563 1.00 35.62 C +ATOM 2686 CG PHE B 108 23.971 13.762 14.591 1.00 35.62 C +ATOM 2687 CD1 PHE B 108 23.477 13.556 13.292 1.00 35.62 C +ATOM 2688 CD2 PHE B 108 25.177 13.159 14.986 1.00 35.62 C +ATOM 2689 CE1 PHE B 108 24.177 12.737 12.393 1.00 35.62 C +ATOM 2690 CE2 PHE B 108 25.869 12.330 14.087 1.00 35.62 C +ATOM 2691 CZ PHE B 108 25.377 12.122 12.788 1.00 35.62 C +ATOM 2692 N GLU B 109 21.294 17.060 15.969 1.00 36.03 N +ATOM 2693 CA GLU B 109 20.234 17.522 16.866 1.00 36.03 C +ATOM 2694 C GLU B 109 20.156 16.639 18.120 1.00 36.03 C +ATOM 2695 O GLU B 109 20.328 15.417 18.064 1.00 36.03 O +ATOM 2696 CB GLU B 109 18.870 17.628 16.147 1.00 36.03 C +ATOM 2697 CG GLU B 109 18.269 16.279 15.706 1.00 36.03 C +ATOM 2698 CD GLU B 109 16.918 16.380 14.970 1.00 36.03 C +ATOM 2699 OE1 GLU B 109 16.690 15.546 14.062 1.00 36.03 O +ATOM 2700 OE2 GLU B 109 16.060 17.230 15.293 1.00 36.03 O +ATOM 2701 N TYR B 110 19.888 17.263 19.265 1.00 35.90 N +ATOM 2702 CA TYR B 110 19.574 16.530 20.480 1.00 35.90 C +ATOM 2703 C TYR B 110 18.156 15.953 20.406 1.00 35.90 C +ATOM 2704 O TYR B 110 17.203 16.657 20.076 1.00 35.90 O +ATOM 2705 CB TYR B 110 19.779 17.435 21.696 1.00 35.90 C +ATOM 2706 CG TYR B 110 19.440 16.752 23.003 1.00 35.90 C +ATOM 2707 CD1 TYR B 110 18.143 16.873 23.538 1.00 35.90 C +ATOM 2708 CD2 TYR B 110 20.411 15.977 23.664 1.00 35.90 C +ATOM 2709 CE1 TYR B 110 17.817 16.223 24.741 1.00 35.90 C +ATOM 2710 CE2 TYR B 110 20.087 15.325 24.869 1.00 35.90 C +ATOM 2711 CZ TYR B 110 18.788 15.450 25.407 1.00 35.90 C +ATOM 2712 OH TYR B 110 18.478 14.854 26.586 1.00 35.90 O +ATOM 2713 N VAL B 111 18.011 14.673 20.758 1.00 36.82 N +ATOM 2714 CA VAL B 111 16.717 13.986 20.865 1.00 36.82 C +ATOM 2715 C VAL B 111 16.574 13.454 22.286 1.00 36.82 C +ATOM 2716 O VAL B 111 17.429 12.694 22.746 1.00 36.82 O +ATOM 2717 CB VAL B 111 16.581 12.853 19.826 1.00 36.82 C +ATOM 2718 CG1 VAL B 111 15.173 12.239 19.869 1.00 36.82 C +ATOM 2719 CG2 VAL B 111 16.833 13.342 18.393 1.00 36.82 C +ATOM 2720 N ASN B 112 15.494 13.833 22.976 1.00 37.39 N +ATOM 2721 CA ASN B 112 15.198 13.377 24.335 1.00 37.39 C +ATOM 2722 C ASN B 112 14.696 11.922 24.325 1.00 37.39 C +ATOM 2723 O ASN B 112 13.500 11.646 24.375 1.00 37.39 O +ATOM 2724 CB ASN B 112 14.201 14.346 24.989 1.00 37.39 C +ATOM 2725 CG ASN B 112 13.964 14.001 26.447 1.00 37.39 C +ATOM 2726 OD1 ASN B 112 14.843 13.503 27.136 1.00 37.39 O +ATOM 2727 ND2 ASN B 112 12.788 14.278 26.956 1.00 37.39 N +END diff --git a/tests/files/pdb/3GR5_af.pdb b/tests/files/pdb/3GR5_af.pdb new file mode 100644 index 00000000..a857b161 --- /dev/null +++ b/tests/files/pdb/3GR5_af.pdb @@ -0,0 +1,1172 @@ +REMARK TITLE 3GR5 AlphaFold-start MR +REMARK Log-Likelihood Gain: 961.039 +REMARK RFZ=4.3 TFZ=6.6 PAK=3 LLG=91 TFZ==11.1 LLG=961 TFZ==34.3 PAK=3 LLG=961 TFZ==34.3 +REMARK ENSEMBLE e_O52135 EULER 235.43 150.60 357.58 FRAC -0.344 -0.741 -0.238 +CRYST1 90.670 90.670 133.440 90.00 90.00 120.00 P 65 2 2 12 +SCALE1 0.011029 0.006368 -0.000000 0.00000 +SCALE2 0.000000 0.012735 -0.000000 0.00000 +SCALE3 0.000000 0.000000 0.007494 0.00000 +ATOM 1 N LEU A 24 35.139 -15.275 -5.797 1.00 67.95 N +ATOM 2 CA LEU A 24 36.527 -14.831 -5.774 1.00 67.95 C +ATOM 3 C LEU A 24 36.717 -13.478 -5.054 1.00 67.95 C +ATOM 4 O LEU A 24 37.649 -13.330 -4.269 1.00 67.95 O +ATOM 5 CB LEU A 24 37.011 -14.790 -7.236 1.00 67.95 C +ATOM 6 CG LEU A 24 38.502 -14.459 -7.376 1.00 67.95 C +ATOM 7 CD1 LEU A 24 39.421 -15.519 -6.767 1.00 67.95 C +ATOM 8 CD2 LEU A 24 38.882 -14.281 -8.841 1.00 67.95 C +ATOM 9 N GLU A 25 35.829 -12.502 -5.261 1.00 61.74 N +ATOM 10 CA GLU A 25 35.888 -11.184 -4.603 1.00 61.74 C +ATOM 11 C GLU A 25 35.680 -11.280 -3.084 1.00 61.74 C +ATOM 12 O GLU A 25 36.308 -10.544 -2.323 1.00 61.74 O +ATOM 13 CB GLU A 25 34.849 -10.257 -5.251 1.00 61.74 C +ATOM 14 CG GLU A 25 34.898 -8.808 -4.737 1.00 61.74 C +ATOM 15 CD GLU A 25 33.967 -7.866 -5.519 1.00 61.74 C +ATOM 16 OE1 GLU A 25 34.078 -6.636 -5.291 1.00 61.74 O +ATOM 17 OE2 GLU A 25 33.143 -8.379 -6.314 1.00 61.74 O +ATOM 18 N LYS A 26 34.859 -12.233 -2.624 1.00 64.09 N +ATOM 19 CA LYS A 26 34.713 -12.518 -1.188 1.00 64.09 C +ATOM 20 C LYS A 26 35.998 -13.083 -0.576 1.00 64.09 C +ATOM 21 O LYS A 26 36.311 -12.744 0.562 1.00 64.09 O +ATOM 22 CB LYS A 26 33.543 -13.477 -0.941 1.00 64.09 C +ATOM 23 CG LYS A 26 32.179 -12.831 -1.222 1.00 64.09 C +ATOM 24 CD LYS A 26 31.063 -13.835 -0.912 1.00 64.09 C +ATOM 25 CE LYS A 26 29.693 -13.253 -1.267 1.00 64.09 C +ATOM 26 NZ LYS A 26 28.620 -14.261 -1.078 1.00 64.09 N +ATOM 27 N ARG A 27 36.738 -13.922 -1.312 1.00 51.43 N +ATOM 28 CA ARG A 27 37.956 -14.585 -0.815 1.00 51.43 C +ATOM 29 C ARG A 27 39.210 -13.713 -0.912 1.00 51.43 C +ATOM 30 O ARG A 27 40.042 -13.768 -0.018 1.00 51.43 O +ATOM 31 CB ARG A 27 38.124 -15.934 -1.527 1.00 51.43 C +ATOM 32 CG ARG A 27 39.120 -16.847 -0.790 1.00 51.43 C +ATOM 33 CD ARG A 27 39.106 -18.225 -1.445 1.00 51.43 C +ATOM 34 NE ARG A 27 40.098 -19.161 -0.878 1.00 51.43 N +ATOM 35 CZ ARG A 27 40.135 -20.473 -1.062 1.00 51.43 C +ATOM 36 NH1 ARG A 27 39.175 -21.137 -1.626 1.00 51.43 N +ATOM 37 NH2 ARG A 27 41.168 -21.172 -0.694 1.00 51.43 N +ATOM 38 N LEU A 28 39.321 -12.872 -1.942 1.00 46.26 N +ATOM 39 CA LEU A 28 40.435 -11.926 -2.117 1.00 46.26 C +ATOM 40 C LEU A 28 40.441 -10.785 -1.076 1.00 46.26 C +ATOM 41 O LEU A 28 41.436 -10.079 -0.934 1.00 46.26 O +ATOM 42 CB LEU A 28 40.387 -11.369 -3.551 1.00 46.26 C +ATOM 43 CG LEU A 28 40.864 -12.345 -4.644 1.00 46.26 C +ATOM 44 CD1 LEU A 28 40.471 -11.803 -6.018 1.00 46.26 C +ATOM 45 CD2 LEU A 28 42.385 -12.509 -4.640 1.00 46.26 C +ATOM 46 N GLY A 29 39.356 -10.600 -0.317 1.00 61.74 N +ATOM 47 CA GLY A 29 39.280 -9.592 0.739 1.00 61.74 C +ATOM 48 C GLY A 29 39.378 -8.149 0.223 1.00 61.74 C +ATOM 49 O GLY A 29 39.233 -7.864 -0.966 1.00 61.74 O +ATOM 50 N LYS A 30 39.584 -7.197 1.144 1.00 67.68 N +ATOM 51 CA LYS A 30 39.664 -5.754 0.831 1.00 67.68 C +ATOM 52 C LYS A 30 41.091 -5.204 0.789 1.00 67.68 C +ATOM 53 O LYS A 30 41.256 -4.025 0.491 1.00 67.68 O +ATOM 54 CB LYS A 30 38.798 -4.944 1.808 1.00 67.68 C +ATOM 55 CG LYS A 30 37.297 -5.186 1.607 1.00 67.68 C +ATOM 56 CD LYS A 30 36.493 -4.296 2.563 1.00 67.68 C +ATOM 57 CE LYS A 30 34.991 -4.547 2.394 1.00 67.68 C +ATOM 58 NZ LYS A 30 34.198 -3.757 3.369 1.00 67.68 N +ATOM 59 N ASN A 31 42.090 -6.027 1.103 1.00 38.63 N +ATOM 60 CA ASN A 31 43.474 -5.579 1.208 1.00 38.63 C +ATOM 61 C ASN A 31 44.012 -5.231 -0.180 1.00 38.63 C +ATOM 62 O ASN A 31 43.859 -6.012 -1.121 1.00 38.63 O +ATOM 63 CB ASN A 31 44.322 -6.645 1.921 1.00 38.63 C +ATOM 64 CG ASN A 31 43.912 -6.874 3.365 1.00 38.63 C +ATOM 65 OD1 ASN A 31 43.092 -6.174 3.939 1.00 38.63 O +ATOM 66 ND2 ASN A 31 44.450 -7.890 3.994 1.00 38.63 N +ATOM 67 N GLU A 32 44.609 -4.053 -0.302 1.00 28.08 N +ATOM 68 CA GLU A 32 45.139 -3.536 -1.561 1.00 28.08 C +ATOM 69 C GLU A 32 46.424 -4.272 -1.956 1.00 28.08 C +ATOM 70 O GLU A 32 47.208 -4.684 -1.101 1.00 28.08 O +ATOM 71 CB GLU A 32 45.369 -2.021 -1.445 1.00 28.08 C +ATOM 72 CG GLU A 32 44.055 -1.254 -1.211 1.00 28.08 C +ATOM 73 CD GLU A 32 44.231 0.258 -0.984 1.00 28.08 C +ATOM 74 OE1 GLU A 32 43.194 0.955 -1.073 1.00 28.08 O +ATOM 75 OE2 GLU A 32 45.363 0.706 -0.681 1.00 28.08 O +ATOM 76 N TYR A 33 46.635 -4.429 -3.263 1.00 26.41 N +ATOM 77 CA TYR A 33 47.849 -5.017 -3.820 1.00 26.41 C +ATOM 78 C TYR A 33 48.711 -3.929 -4.457 1.00 26.41 C +ATOM 79 O TYR A 33 48.198 -3.064 -5.178 1.00 26.41 O +ATOM 80 CB TYR A 33 47.480 -6.103 -4.834 1.00 26.41 C +ATOM 81 CG TYR A 33 48.667 -6.855 -5.394 1.00 26.41 C +ATOM 82 CD1 TYR A 33 49.124 -6.590 -6.699 1.00 26.41 C +ATOM 83 CD2 TYR A 33 49.309 -7.830 -4.606 1.00 26.41 C +ATOM 84 CE1 TYR A 33 50.198 -7.329 -7.235 1.00 26.41 C +ATOM 85 CE2 TYR A 33 50.395 -8.554 -5.130 1.00 26.41 C +ATOM 86 CZ TYR A 33 50.835 -8.314 -6.448 1.00 26.41 C +ATOM 87 OH TYR A 33 51.855 -9.051 -6.952 1.00 26.41 O +ATOM 88 N PHE A 34 50.021 -3.994 -4.223 1.00 25.86 N +ATOM 89 CA PHE A 34 50.997 -3.137 -4.885 1.00 25.86 C +ATOM 90 C PHE A 34 52.183 -3.959 -5.390 1.00 25.86 C +ATOM 91 O PHE A 34 52.607 -4.919 -4.753 1.00 25.86 O +ATOM 92 CB PHE A 34 51.430 -1.979 -3.970 1.00 25.86 C +ATOM 93 CG PHE A 34 52.321 -2.382 -2.807 1.00 25.86 C +ATOM 94 CD1 PHE A 34 51.754 -2.794 -1.587 1.00 25.86 C +ATOM 95 CD2 PHE A 34 53.723 -2.361 -2.953 1.00 25.86 C +ATOM 96 CE1 PHE A 34 52.581 -3.193 -0.523 1.00 25.86 C +ATOM 97 CE2 PHE A 34 54.550 -2.763 -1.889 1.00 25.86 C +ATOM 98 CZ PHE A 34 53.979 -3.181 -0.675 1.00 25.86 C +ATOM 99 N ILE A 35 52.723 -3.575 -6.545 1.00 25.34 N +ATOM 100 CA ILE A 35 53.951 -4.157 -7.094 1.00 25.34 C +ATOM 101 C ILE A 35 54.679 -3.118 -7.945 1.00 25.34 C +ATOM 102 O ILE A 35 54.049 -2.367 -8.690 1.00 25.34 O +ATOM 103 CB ILE A 35 53.675 -5.477 -7.861 1.00 25.34 C +ATOM 104 CG1 ILE A 35 55.005 -6.138 -8.284 1.00 25.34 C +ATOM 105 CG2 ILE A 35 52.732 -5.288 -9.063 1.00 25.34 C +ATOM 106 CD1 ILE A 35 54.850 -7.566 -8.814 1.00 25.34 C +ATOM 107 N ILE A 36 56.007 -3.089 -7.842 1.00 24.41 N +ATOM 108 CA ILE A 36 56.893 -2.239 -8.643 1.00 24.41 C +ATOM 109 C ILE A 36 57.961 -3.146 -9.245 1.00 24.41 C +ATOM 110 O ILE A 36 58.743 -3.746 -8.510 1.00 24.41 O +ATOM 111 CB ILE A 36 57.536 -1.125 -7.783 1.00 24.41 C +ATOM 112 CG1 ILE A 36 56.472 -0.249 -7.089 1.00 24.41 C +ATOM 113 CG2 ILE A 36 58.441 -0.247 -8.672 1.00 24.41 C +ATOM 114 CD1 ILE A 36 57.023 0.715 -6.033 1.00 24.41 C +ATOM 115 N THR A 37 58.012 -3.254 -10.573 1.00 23.07 N +ATOM 116 CA THR A 37 59.039 -4.059 -11.248 1.00 23.07 C +ATOM 117 C THR A 37 59.467 -3.460 -12.584 1.00 23.07 C +ATOM 118 O THR A 37 58.686 -2.817 -13.285 1.00 23.07 O +ATOM 119 CB THR A 37 58.611 -5.526 -11.395 1.00 23.07 C +ATOM 120 OG1 THR A 37 59.684 -6.219 -11.990 1.00 23.07 O +ATOM 121 CG2 THR A 37 57.364 -5.726 -12.253 1.00 23.07 C +ATOM 122 N LYS A 38 60.743 -3.666 -12.938 1.00 28.98 N +ATOM 123 CA LYS A 38 61.327 -3.235 -14.220 1.00 28.98 C +ATOM 124 C LYS A 38 61.092 -4.250 -15.339 1.00 28.98 C +ATOM 125 O LYS A 38 60.912 -3.838 -16.478 1.00 28.98 O +ATOM 126 CB LYS A 38 62.832 -2.972 -14.066 1.00 28.98 C +ATOM 127 CG LYS A 38 63.141 -1.742 -13.201 1.00 28.98 C +ATOM 128 CD LYS A 38 64.654 -1.487 -13.172 1.00 28.98 C +ATOM 129 CE LYS A 38 64.970 -0.240 -12.340 1.00 28.98 C +ATOM 130 NZ LYS A 38 66.431 0.026 -12.294 1.00 28.98 N +ATOM 131 N SER A 39 61.124 -5.542 -15.028 1.00 23.27 N +ATOM 132 CA SER A 39 60.750 -6.635 -15.926 1.00 23.27 C +ATOM 133 C SER A 39 60.594 -7.897 -15.087 1.00 23.27 C +ATOM 134 O SER A 39 61.539 -8.333 -14.426 1.00 23.27 O +ATOM 135 CB SER A 39 61.805 -6.887 -17.010 1.00 23.27 C +ATOM 136 OG SER A 39 61.265 -7.833 -17.906 1.00 23.27 O +ATOM 137 N SER A 40 59.393 -8.460 -15.045 1.00 21.49 N +ATOM 138 CA SER A 40 59.146 -9.748 -14.399 1.00 21.49 C +ATOM 139 C SER A 40 58.102 -10.529 -15.181 1.00 21.49 C +ATOM 140 O SER A 40 57.132 -9.920 -15.644 1.00 21.49 O +ATOM 141 CB SER A 40 58.687 -9.566 -12.950 1.00 21.49 C +ATOM 142 OG SER A 40 59.706 -8.920 -12.214 1.00 21.49 O +ATOM 143 N PRO A 41 58.251 -11.860 -15.298 1.00 19.55 N +ATOM 144 CA PRO A 41 57.278 -12.669 -16.008 1.00 19.55 C +ATOM 145 C PRO A 41 55.926 -12.559 -15.310 1.00 19.55 C +ATOM 146 O PRO A 41 55.821 -12.722 -14.091 1.00 19.55 O +ATOM 147 CB PRO A 41 57.844 -14.092 -16.029 1.00 19.55 C +ATOM 148 CG PRO A 41 58.793 -14.128 -14.832 1.00 19.55 C +ATOM 149 CD PRO A 41 59.308 -12.693 -14.742 1.00 19.55 C +ATOM 150 N VAL A 42 54.869 -12.324 -16.087 1.00 20.70 N +ATOM 151 CA VAL A 42 53.490 -12.184 -15.595 1.00 20.70 C +ATOM 152 C VAL A 42 53.082 -13.405 -14.768 1.00 20.70 C +ATOM 153 O VAL A 42 52.393 -13.263 -13.763 1.00 20.70 O +ATOM 154 CB VAL A 42 52.531 -11.954 -16.782 1.00 20.70 C +ATOM 155 CG1 VAL A 42 51.051 -12.025 -16.382 1.00 20.70 C +ATOM 156 CG2 VAL A 42 52.756 -10.565 -17.387 1.00 20.70 C +ATOM 157 N ARG A 43 53.580 -14.598 -15.120 1.00 19.87 N +ATOM 158 CA ARG A 43 53.400 -15.826 -14.331 1.00 19.87 C +ATOM 159 C ARG A 43 53.902 -15.685 -12.888 1.00 19.87 C +ATOM 160 O ARG A 43 53.205 -16.110 -11.976 1.00 19.87 O +ATOM 161 CB ARG A 43 54.081 -17.000 -15.057 1.00 19.87 C +ATOM 162 CG ARG A 43 53.892 -18.338 -14.318 1.00 19.87 C +ATOM 163 CD ARG A 43 54.574 -19.515 -15.029 1.00 19.87 C +ATOM 164 NE ARG A 43 53.909 -19.856 -16.303 1.00 19.87 N +ATOM 165 CZ ARG A 43 54.457 -19.943 -17.504 1.00 19.87 C +ATOM 166 NH1 ARG A 43 55.732 -19.833 -17.732 1.00 19.87 N +ATOM 167 NH2 ARG A 43 53.712 -20.136 -18.544 1.00 19.87 N +ATOM 168 N ALA A 44 55.081 -15.098 -12.675 1.00 21.22 N +ATOM 169 CA ALA A 44 55.617 -14.891 -11.329 1.00 21.22 C +ATOM 170 C ALA A 44 54.773 -13.875 -10.554 1.00 21.22 C +ATOM 171 O ALA A 44 54.387 -14.146 -9.425 1.00 21.22 O +ATOM 172 CB ALA A 44 57.086 -14.464 -11.409 1.00 21.22 C +ATOM 173 N ILE A 45 54.376 -12.774 -11.196 1.00 21.44 N +ATOM 174 CA ILE A 45 53.526 -11.749 -10.571 1.00 21.44 C +ATOM 175 C ILE A 45 52.174 -12.322 -10.148 1.00 21.44 C +ATOM 176 O ILE A 45 51.676 -12.003 -9.074 1.00 21.44 O +ATOM 177 CB ILE A 45 53.327 -10.561 -11.533 1.00 21.44 C +ATOM 178 CG1 ILE A 45 54.700 -9.931 -11.844 1.00 21.44 C +ATOM 179 CG2 ILE A 45 52.333 -9.529 -10.956 1.00 21.44 C +ATOM 180 CD1 ILE A 45 54.600 -8.799 -12.854 1.00 21.44 C +ATOM 181 N LEU A 46 51.564 -13.157 -10.988 1.00 20.95 N +ATOM 182 CA LEU A 46 50.295 -13.812 -10.677 1.00 20.95 C +ATOM 183 C LEU A 46 50.443 -14.861 -9.568 1.00 20.95 C +ATOM 184 O LEU A 46 49.535 -14.992 -8.750 1.00 20.95 O +ATOM 185 CB LEU A 46 49.740 -14.444 -11.957 1.00 20.95 C +ATOM 186 CG LEU A 46 49.221 -13.435 -12.995 1.00 20.95 C +ATOM 187 CD1 LEU A 46 48.908 -14.200 -14.280 1.00 20.95 C +ATOM 188 CD2 LEU A 46 47.969 -12.695 -12.519 1.00 20.95 C +ATOM 189 N ASN A 47 51.576 -15.568 -9.512 1.00 20.85 N +ATOM 190 CA ASN A 47 51.900 -16.467 -8.403 1.00 20.85 C +ATOM 191 C ASN A 47 52.059 -15.700 -7.087 1.00 20.85 C +ATOM 192 O ASN A 47 51.444 -16.081 -6.095 1.00 20.85 O +ATOM 193 CB ASN A 47 53.171 -17.268 -8.723 1.00 20.85 C +ATOM 194 CG ASN A 47 52.944 -18.437 -9.659 1.00 20.85 C +ATOM 195 OD1 ASN A 47 51.871 -18.998 -9.783 1.00 20.85 O +ATOM 196 ND2 ASN A 47 53.979 -18.905 -10.313 1.00 20.85 N +ATOM 197 N ASP A 48 52.817 -14.604 -7.085 1.00 22.11 N +ATOM 198 CA ASP A 48 53.033 -13.766 -5.901 1.00 22.11 C +ATOM 199 C ASP A 48 51.726 -13.111 -5.442 1.00 22.11 C +ATOM 200 O ASP A 48 51.405 -13.102 -4.251 1.00 22.11 O +ATOM 201 CB ASP A 48 54.082 -12.686 -6.212 1.00 22.11 C +ATOM 202 CG ASP A 48 55.493 -13.232 -6.457 1.00 22.11 C +ATOM 203 OD1 ASP A 48 55.770 -14.384 -6.050 1.00 22.11 O +ATOM 204 OD2 ASP A 48 56.296 -12.473 -7.047 1.00 22.11 O +ATOM 205 N PHE A 49 50.920 -12.632 -6.394 1.00 21.77 N +ATOM 206 CA PHE A 49 49.577 -12.125 -6.139 1.00 21.77 C +ATOM 207 C PHE A 49 48.714 -13.188 -5.453 1.00 21.77 C +ATOM 208 O PHE A 49 48.134 -12.928 -4.402 1.00 21.77 O +ATOM 209 CB PHE A 49 48.947 -11.665 -7.461 1.00 21.77 C +ATOM 210 CG PHE A 49 47.507 -11.213 -7.332 1.00 21.77 C +ATOM 211 CD1 PHE A 49 46.455 -12.125 -7.537 1.00 21.77 C +ATOM 212 CD2 PHE A 49 47.222 -9.882 -6.988 1.00 21.77 C +ATOM 213 CE1 PHE A 49 45.122 -11.704 -7.395 1.00 21.77 C +ATOM 214 CE2 PHE A 49 45.890 -9.454 -6.860 1.00 21.77 C +ATOM 215 CZ PHE A 49 44.841 -10.368 -7.064 1.00 21.77 C +ATOM 216 N ALA A 50 48.644 -14.398 -6.009 1.00 22.23 N +ATOM 217 CA ALA A 50 47.825 -15.467 -5.452 1.00 22.23 C +ATOM 218 C ALA A 50 48.332 -15.936 -4.073 1.00 22.23 C +ATOM 219 O ALA A 50 47.528 -16.142 -3.160 1.00 22.23 O +ATOM 220 CB ALA A 50 47.778 -16.593 -6.481 1.00 22.23 C +ATOM 221 N ALA A 51 49.655 -16.008 -3.891 1.00 23.86 N +ATOM 222 CA ALA A 51 50.296 -16.328 -2.618 1.00 23.86 C +ATOM 223 C ALA A 51 49.968 -15.295 -1.527 1.00 23.86 C +ATOM 224 O ALA A 51 49.620 -15.683 -0.412 1.00 23.86 O +ATOM 225 CB ALA A 51 51.809 -16.444 -2.848 1.00 23.86 C +ATOM 226 N ASN A 52 49.981 -13.997 -1.854 1.00 22.77 N +ATOM 227 CA ASN A 52 49.625 -12.918 -0.924 1.00 22.77 C +ATOM 228 C ASN A 52 48.194 -13.056 -0.370 1.00 22.77 C +ATOM 229 O ASN A 52 47.934 -12.733 0.786 1.00 22.77 O +ATOM 230 CB ASN A 52 49.811 -11.582 -1.657 1.00 22.77 C +ATOM 231 CG ASN A 52 49.577 -10.390 -0.748 1.00 22.77 C +ATOM 232 OD1 ASN A 52 50.130 -10.274 0.330 1.00 22.77 O +ATOM 233 ND2 ASN A 52 48.750 -9.457 -1.158 1.00 22.77 N +ATOM 234 N TYR A 53 47.266 -13.582 -1.176 1.00 23.66 N +ATOM 235 CA TYR A 53 45.878 -13.827 -0.770 1.00 23.66 C +ATOM 236 C TYR A 53 45.602 -15.277 -0.344 1.00 23.66 C +ATOM 237 O TYR A 53 44.448 -15.627 -0.099 1.00 23.66 O +ATOM 238 CB TYR A 53 44.927 -13.324 -1.863 1.00 23.66 C +ATOM 239 CG TYR A 53 45.015 -11.824 -2.074 1.00 23.66 C +ATOM 240 CD1 TYR A 53 44.445 -10.946 -1.133 1.00 23.66 C +ATOM 241 CD2 TYR A 53 45.679 -11.302 -3.197 1.00 23.66 C +ATOM 242 CE1 TYR A 53 44.541 -9.552 -1.314 1.00 23.66 C +ATOM 243 CE2 TYR A 53 45.823 -9.912 -3.357 1.00 23.66 C +ATOM 244 CZ TYR A 53 45.242 -9.035 -2.422 1.00 23.66 C +ATOM 245 OH TYR A 53 45.319 -7.696 -2.623 1.00 23.66 O +ATOM 246 N SER A 54 46.633 -16.124 -0.221 1.00 24.26 N +ATOM 247 CA SER A 54 46.502 -17.543 0.151 1.00 24.26 C +ATOM 248 C SER A 54 45.539 -18.331 -0.757 1.00 24.26 C +ATOM 249 O SER A 54 44.770 -19.181 -0.298 1.00 24.26 O +ATOM 250 CB SER A 54 46.142 -17.683 1.635 1.00 24.26 C +ATOM 251 OG SER A 54 47.111 -17.039 2.436 1.00 24.26 O +ATOM 252 N ILE A 55 45.565 -18.039 -2.059 1.00 22.83 N +ATOM 253 CA ILE A 55 44.779 -18.742 -3.076 1.00 22.83 C +ATOM 254 C ILE A 55 45.740 -19.585 -3.921 1.00 22.83 C +ATOM 255 O ILE A 55 46.652 -19.030 -4.530 1.00 22.83 O +ATOM 256 CB ILE A 55 43.937 -17.763 -3.927 1.00 22.83 C +ATOM 257 CG1 ILE A 55 43.020 -16.911 -3.018 1.00 22.83 C +ATOM 258 CG2 ILE A 55 43.093 -18.548 -4.951 1.00 22.83 C +ATOM 259 CD1 ILE A 55 42.142 -15.899 -3.757 1.00 22.83 C +ATOM 260 N PRO A 56 45.567 -20.914 -3.990 1.00 21.28 N +ATOM 261 CA PRO A 56 46.280 -21.734 -4.961 1.00 21.28 C +ATOM 262 C PRO A 56 45.979 -21.259 -6.386 1.00 21.28 C +ATOM 263 O PRO A 56 44.821 -21.016 -6.726 1.00 21.28 O +ATOM 264 CB PRO A 56 45.784 -23.164 -4.732 1.00 21.28 C +ATOM 265 CG PRO A 56 45.252 -23.154 -3.300 1.00 21.28 C +ATOM 266 CD PRO A 56 44.721 -21.733 -3.140 1.00 21.28 C +ATOM 267 N VAL A 57 47.000 -21.142 -7.235 1.00 20.21 N +ATOM 268 CA VAL A 57 46.837 -20.691 -8.622 1.00 20.21 C +ATOM 269 C VAL A 57 47.405 -21.715 -9.595 1.00 20.21 C +ATOM 270 O VAL A 57 48.477 -22.279 -9.383 1.00 20.21 O +ATOM 271 CB VAL A 57 47.400 -19.270 -8.820 1.00 20.21 C +ATOM 272 CG1 VAL A 57 48.922 -19.204 -8.691 1.00 20.21 C +ATOM 273 CG2 VAL A 57 46.997 -18.622 -10.150 1.00 20.21 C +ATOM 274 N PHE A 58 46.673 -21.947 -10.678 1.00 19.37 N +ATOM 275 CA PHE A 58 47.143 -22.660 -11.851 1.00 19.37 C +ATOM 276 C PHE A 58 47.280 -21.662 -12.999 1.00 19.37 C +ATOM 277 O PHE A 58 46.319 -20.985 -13.360 1.00 19.37 O +ATOM 278 CB PHE A 58 46.178 -23.795 -12.194 1.00 19.37 C +ATOM 279 CG PHE A 58 46.565 -24.513 -13.472 1.00 19.37 C +ATOM 280 CD1 PHE A 58 46.086 -24.053 -14.714 1.00 19.37 C +ATOM 281 CD2 PHE A 58 47.460 -25.596 -13.428 1.00 19.37 C +ATOM 282 CE1 PHE A 58 46.491 -24.680 -15.903 1.00 19.37 C +ATOM 283 CE2 PHE A 58 47.862 -26.227 -14.619 1.00 19.37 C +ATOM 284 CZ PHE A 58 47.374 -25.772 -15.857 1.00 19.37 C +ATOM 285 N ILE A 59 48.467 -21.578 -13.595 1.00 19.63 N +ATOM 286 CA ILE A 59 48.748 -20.658 -14.700 1.00 19.63 C +ATOM 287 C ILE A 59 49.168 -21.480 -15.911 1.00 19.63 C +ATOM 288 O ILE A 59 50.124 -22.253 -15.830 1.00 19.63 O +ATOM 289 CB ILE A 59 49.817 -19.613 -14.306 1.00 19.63 C +ATOM 290 CG1 ILE A 59 49.421 -18.871 -13.009 1.00 19.63 C +ATOM 291 CG2 ILE A 59 50.012 -18.620 -15.472 1.00 19.63 C +ATOM 292 CD1 ILE A 59 50.507 -17.942 -12.469 1.00 19.63 C +ATOM 293 N SER A 60 48.487 -21.287 -17.041 1.00 19.23 N +ATOM 294 CA SER A 60 48.823 -21.979 -18.287 1.00 19.23 C +ATOM 295 C SER A 60 50.280 -21.734 -18.708 1.00 19.23 C +ATOM 296 O SER A 60 50.834 -20.634 -18.566 1.00 19.23 O +ATOM 297 CB SER A 60 47.859 -21.592 -19.415 1.00 19.23 C +ATOM 298 OG SER A 60 48.259 -22.267 -20.592 1.00 19.23 O +ATOM 299 N SER A 61 50.901 -22.756 -19.302 1.00 22.17 N +ATOM 300 CA SER A 61 52.217 -22.657 -19.946 1.00 22.17 C +ATOM 301 C SER A 61 52.253 -21.614 -21.071 1.00 22.17 C +ATOM 302 O SER A 61 53.326 -21.124 -21.411 1.00 22.17 O +ATOM 303 CB SER A 61 52.645 -24.019 -20.491 1.00 22.17 C +ATOM 304 OG SER A 61 51.627 -24.583 -21.295 1.00 22.17 O +ATOM 305 N SER A 62 51.096 -21.235 -21.619 1.00 22.23 N +ATOM 306 CA SER A 62 50.964 -20.224 -22.677 1.00 22.23 C +ATOM 307 C SER A 62 51.096 -18.777 -22.180 1.00 22.23 C +ATOM 308 O SER A 62 51.204 -17.854 -22.988 1.00 22.23 O +ATOM 309 CB SER A 62 49.617 -20.405 -23.377 1.00 22.23 C +ATOM 310 OG SER A 62 49.491 -21.737 -23.833 1.00 22.23 O +ATOM 311 N VAL A 63 51.091 -18.546 -20.861 1.00 21.44 N +ATOM 312 CA VAL A 63 51.296 -17.218 -20.256 1.00 21.44 C +ATOM 313 C VAL A 63 52.794 -16.959 -20.074 1.00 21.44 C +ATOM 314 O VAL A 63 53.360 -17.284 -19.032 1.00 21.44 O +ATOM 315 CB VAL A 63 50.535 -17.078 -18.921 1.00 21.44 C +ATOM 316 CG1 VAL A 63 50.655 -15.649 -18.369 1.00 21.44 C +ATOM 317 CG2 VAL A 63 49.040 -17.391 -19.074 1.00 21.44 C +ATOM 318 N ASN A 64 53.439 -16.390 -21.092 1.00 22.29 N +ATOM 319 CA ASN A 64 54.888 -16.138 -21.112 1.00 22.29 C +ATOM 320 C ASN A 64 55.244 -14.653 -21.308 1.00 22.29 C +ATOM 321 O ASN A 64 56.346 -14.341 -21.743 1.00 22.29 O +ATOM 322 CB ASN A 64 55.550 -17.055 -22.159 1.00 22.29 C +ATOM 323 CG ASN A 64 55.361 -18.532 -21.865 1.00 22.29 C +ATOM 324 OD1 ASN A 64 55.343 -18.992 -20.727 1.00 22.29 O +ATOM 325 ND2 ASN A 64 55.196 -19.338 -22.885 1.00 22.29 N +ATOM 326 N ASP A 65 54.308 -13.743 -21.036 1.00 21.33 N +ATOM 327 CA ASP A 65 54.536 -12.303 -21.168 1.00 21.33 C +ATOM 328 C ASP A 65 55.324 -11.740 -19.980 1.00 21.33 C +ATOM 329 O ASP A 65 55.243 -12.266 -18.866 1.00 21.33 O +ATOM 330 CB ASP A 65 53.210 -11.543 -21.296 1.00 21.33 C +ATOM 331 CG ASP A 65 52.169 -12.283 -22.121 1.00 21.33 C +ATOM 332 OD1 ASP A 65 51.992 -12.027 -23.330 1.00 21.33 O +ATOM 333 OD2 ASP A 65 51.450 -13.119 -21.531 1.00 21.33 O +ATOM 334 N ASP A 66 55.998 -10.612 -20.202 1.00 20.45 N +ATOM 335 CA ASP A 66 56.677 -9.840 -19.162 1.00 20.45 C +ATOM 336 C ASP A 66 55.900 -8.560 -18.823 1.00 20.45 C +ATOM 337 O ASP A 66 55.307 -7.899 -19.684 1.00 20.45 O +ATOM 338 CB ASP A 66 58.120 -9.522 -19.577 1.00 20.45 C +ATOM 339 CG ASP A 66 59.034 -10.753 -19.610 1.00 20.45 C +ATOM 340 OD1 ASP A 66 58.867 -11.631 -18.733 1.00 20.45 O +ATOM 341 OD2 ASP A 66 59.936 -10.767 -20.476 1.00 20.45 O +ATOM 342 N PHE A 67 55.921 -8.189 -17.545 1.00 21.60 N +ATOM 343 CA PHE A 67 55.363 -6.940 -17.041 1.00 21.60 C +ATOM 344 C PHE A 67 56.457 -6.005 -16.552 1.00 21.60 C +ATOM 345 O PHE A 67 57.347 -6.382 -15.787 1.00 21.60 O +ATOM 346 CB PHE A 67 54.303 -7.224 -15.977 1.00 21.60 C +ATOM 347 CG PHE A 67 53.693 -6.003 -15.305 1.00 21.60 C +ATOM 348 CD1 PHE A 67 54.199 -5.540 -14.079 1.00 21.60 C +ATOM 349 CD2 PHE A 67 52.621 -5.317 -15.903 1.00 21.60 C +ATOM 350 CE1 PHE A 67 53.639 -4.409 -13.456 1.00 21.60 C +ATOM 351 CE2 PHE A 67 52.056 -4.188 -15.288 1.00 21.60 C +ATOM 352 CZ PHE A 67 52.567 -3.732 -14.063 1.00 21.60 C +ATOM 353 N SER A 68 56.347 -4.759 -16.997 1.00 24.41 N +ATOM 354 CA SER A 68 57.180 -3.638 -16.593 1.00 24.41 C +ATOM 355 C SER A 68 56.262 -2.489 -16.211 1.00 24.41 C +ATOM 356 O SER A 68 55.400 -2.090 -16.997 1.00 24.41 O +ATOM 357 CB SER A 68 58.089 -3.226 -17.748 1.00 24.41 C +ATOM 358 OG SER A 68 59.003 -2.247 -17.304 1.00 24.41 O +ATOM 359 N GLY A 69 56.415 -1.980 -14.993 1.00 26.09 N +ATOM 360 CA GLY A 69 55.584 -0.903 -14.485 1.00 26.09 C +ATOM 361 C GLY A 69 55.259 -1.044 -13.007 1.00 26.09 C +ATOM 362 O GLY A 69 55.863 -1.818 -12.262 1.00 26.09 O +ATOM 363 N GLU A 70 54.283 -0.250 -12.587 1.00 23.72 N +ATOM 364 CA GLU A 70 53.881 -0.129 -11.197 1.00 23.72 C +ATOM 365 C GLU A 70 52.356 -0.235 -11.063 1.00 23.72 C +ATOM 366 O GLU A 70 51.594 0.352 -11.836 1.00 23.72 O +ATOM 367 CB GLU A 70 54.461 1.177 -10.642 1.00 23.72 C +ATOM 368 CG GLU A 70 53.949 1.485 -9.238 1.00 23.72 C +ATOM 369 CD GLU A 70 54.691 2.656 -8.588 1.00 23.72 C +ATOM 370 OE1 GLU A 70 54.731 2.693 -7.342 1.00 23.72 O +ATOM 371 OE2 GLU A 70 55.051 3.628 -9.298 1.00 23.72 O +ATOM 372 N ILE A 71 51.916 -0.987 -10.056 1.00 26.33 N +ATOM 373 CA ILE A 71 50.526 -1.088 -9.618 1.00 26.33 C +ATOM 374 C ILE A 71 50.481 -0.500 -8.209 1.00 26.33 C +ATOM 375 O ILE A 71 51.081 -1.064 -7.297 1.00 26.33 O +ATOM 376 CB ILE A 71 50.038 -2.555 -9.667 1.00 26.33 C +ATOM 377 CG1 ILE A 71 50.263 -3.162 -11.075 1.00 26.33 C +ATOM 378 CG2 ILE A 71 48.566 -2.651 -9.230 1.00 26.33 C +ATOM 379 CD1 ILE A 71 49.743 -4.590 -11.276 1.00 26.33 C +ATOM 380 N LYS A 72 49.809 0.646 -8.035 1.00 30.72 N +ATOM 381 CA LYS A 72 49.692 1.333 -6.738 1.00 30.72 C +ATOM 382 C LYS A 72 48.325 1.077 -6.127 1.00 30.72 C +ATOM 383 O LYS A 72 47.341 1.595 -6.647 1.00 30.72 O +ATOM 384 CB LYS A 72 49.899 2.854 -6.865 1.00 30.72 C +ATOM 385 CG LYS A 72 51.304 3.231 -7.323 1.00 30.72 C +ATOM 386 CD LYS A 72 51.548 4.749 -7.280 1.00 30.72 C +ATOM 387 CE LYS A 72 52.992 4.916 -7.728 1.00 30.72 C +ATOM 388 NZ LYS A 72 53.499 6.286 -7.899 1.00 30.72 N +ATOM 389 N ASN A 73 48.293 0.363 -5.004 1.00 40.06 N +ATOM 390 CA ASN A 73 47.154 0.316 -4.085 1.00 40.06 C +ATOM 391 C ASN A 73 45.800 0.061 -4.773 1.00 40.06 C +ATOM 392 O ASN A 73 44.807 0.750 -4.548 1.00 40.06 O +ATOM 393 CB ASN A 73 47.190 1.570 -3.201 1.00 40.06 C +ATOM 394 CG ASN A 73 48.253 1.476 -2.126 1.00 40.06 C +ATOM 395 OD1 ASN A 73 49.366 1.023 -2.344 1.00 40.06 O +ATOM 396 ND2 ASN A 73 47.936 1.887 -0.929 1.00 40.06 N +ATOM 397 N GLU A 74 45.762 -0.931 -5.662 1.00 30.32 N +ATOM 398 CA GLU A 74 44.533 -1.306 -6.354 1.00 30.32 C +ATOM 399 C GLU A 74 43.854 -2.481 -5.637 1.00 30.32 C +ATOM 400 O GLU A 74 44.498 -3.342 -5.032 1.00 30.32 O +ATOM 401 CB GLU A 74 44.796 -1.581 -7.846 1.00 30.32 C +ATOM 402 CG GLU A 74 45.072 -0.309 -8.679 1.00 30.32 C +ATOM 403 CD GLU A 74 45.176 -0.573 -10.198 1.00 30.32 C +ATOM 404 OE1 GLU A 74 45.866 0.189 -10.925 1.00 30.32 O +ATOM 405 OE2 GLU A 74 44.535 -1.526 -10.689 1.00 30.32 O +ATOM 406 N LYS A 75 42.518 -2.540 -5.721 1.00 25.25 N +ATOM 407 CA LYS A 75 41.761 -3.703 -5.237 1.00 25.25 C +ATOM 408 C LYS A 75 42.166 -4.955 -6.026 1.00 25.25 C +ATOM 409 O LYS A 75 42.300 -4.861 -7.247 1.00 25.25 O +ATOM 410 CB LYS A 75 40.248 -3.481 -5.350 1.00 25.25 C +ATOM 411 CG LYS A 75 39.736 -2.436 -4.352 1.00 25.25 C +ATOM 412 CD LYS A 75 38.205 -2.387 -4.402 1.00 25.25 C +ATOM 413 CE LYS A 75 37.672 -1.380 -3.381 1.00 25.25 C +ATOM 414 NZ LYS A 75 36.188 -1.390 -3.355 1.00 25.25 N +ATOM 415 N PRO A 76 42.241 -6.136 -5.393 1.00 27.13 N +ATOM 416 CA PRO A 76 42.759 -7.350 -6.026 1.00 27.13 C +ATOM 417 C PRO A 76 42.005 -7.751 -7.299 1.00 27.13 C +ATOM 418 O PRO A 76 42.621 -8.092 -8.304 1.00 27.13 O +ATOM 419 CB PRO A 76 42.675 -8.428 -4.943 1.00 27.13 C +ATOM 420 CG PRO A 76 41.700 -7.884 -3.900 1.00 27.13 C +ATOM 421 CD PRO A 76 41.908 -6.382 -3.999 1.00 27.13 C +ATOM 422 N VAL A 77 40.675 -7.623 -7.307 1.00 27.48 N +ATOM 423 CA VAL A 77 39.862 -7.892 -8.507 1.00 27.48 C +ATOM 424 C VAL A 77 40.194 -6.920 -9.642 1.00 27.48 C +ATOM 425 O VAL A 77 40.325 -7.343 -10.785 1.00 27.48 O +ATOM 426 CB VAL A 77 38.358 -7.846 -8.175 1.00 27.48 C +ATOM 427 CG1 VAL A 77 37.485 -8.126 -9.404 1.00 27.48 C +ATOM 428 CG2 VAL A 77 38.007 -8.887 -7.104 1.00 27.48 C +ATOM 429 N LYS A 78 40.414 -5.632 -9.339 1.00 26.18 N +ATOM 430 CA LYS A 78 40.792 -4.630 -10.349 1.00 26.18 C +ATOM 431 C LYS A 78 42.177 -4.899 -10.927 1.00 26.18 C +ATOM 432 O LYS A 78 42.362 -4.740 -12.130 1.00 26.18 O +ATOM 433 CB LYS A 78 40.758 -3.211 -9.778 1.00 26.18 C +ATOM 434 CG LYS A 78 39.335 -2.715 -9.510 1.00 26.18 C +ATOM 435 CD LYS A 78 39.398 -1.233 -9.131 1.00 26.18 C +ATOM 436 CE LYS A 78 37.992 -0.654 -8.981 1.00 26.18 C +ATOM 437 NZ LYS A 78 38.056 0.813 -8.768 1.00 26.18 N +ATOM 438 N VAL A 79 43.121 -5.346 -10.099 1.00 24.68 N +ATOM 439 CA VAL A 79 44.457 -5.763 -10.551 1.00 24.68 C +ATOM 440 C VAL A 79 44.342 -6.928 -11.526 1.00 24.68 C +ATOM 441 O VAL A 79 44.900 -6.877 -12.621 1.00 24.68 O +ATOM 442 CB VAL A 79 45.353 -6.151 -9.360 1.00 24.68 C +ATOM 443 CG1 VAL A 79 46.695 -6.738 -9.808 1.00 24.68 C +ATOM 444 CG2 VAL A 79 45.626 -4.910 -8.513 1.00 24.68 C +ATOM 445 N LEU A 80 43.567 -7.947 -11.158 1.00 25.56 N +ATOM 446 CA LEU A 80 43.337 -9.119 -11.992 1.00 25.56 C +ATOM 447 C LEU A 80 42.633 -8.742 -13.309 1.00 25.56 C +ATOM 448 O LEU A 80 43.050 -9.187 -14.374 1.00 25.56 O +ATOM 449 CB LEU A 80 42.530 -10.121 -11.158 1.00 25.56 C +ATOM 450 CG LEU A 80 42.611 -11.554 -11.709 1.00 25.56 C +ATOM 451 CD1 LEU A 80 43.626 -12.397 -10.936 1.00 25.56 C +ATOM 452 CD2 LEU A 80 41.260 -12.223 -11.541 1.00 25.56 C +ATOM 453 N GLU A 81 41.628 -7.861 -13.271 1.00 26.41 N +ATOM 454 CA GLU A 81 40.958 -7.321 -14.460 1.00 26.41 C +ATOM 455 C GLU A 81 41.923 -6.538 -15.360 1.00 26.41 C +ATOM 456 O GLU A 81 41.959 -6.750 -16.574 1.00 26.41 O +ATOM 457 CB GLU A 81 39.801 -6.398 -14.056 1.00 26.41 C +ATOM 458 CG GLU A 81 38.550 -7.170 -13.617 1.00 26.41 C +ATOM 459 CD GLU A 81 37.400 -6.243 -13.193 1.00 26.41 C +ATOM 460 OE1 GLU A 81 36.318 -6.789 -12.888 1.00 26.41 O +ATOM 461 OE2 GLU A 81 37.585 -5.000 -13.189 1.00 26.41 O +ATOM 462 N LYS A 82 42.751 -5.668 -14.778 1.00 24.33 N +ATOM 463 CA LYS A 82 43.742 -4.860 -15.496 1.00 24.33 C +ATOM 464 C LYS A 82 44.780 -5.742 -16.183 1.00 24.33 C +ATOM 465 O LYS A 82 45.014 -5.576 -17.379 1.00 24.33 O +ATOM 466 CB LYS A 82 44.372 -3.883 -14.495 1.00 24.33 C +ATOM 467 CG LYS A 82 45.287 -2.828 -15.133 1.00 24.33 C +ATOM 468 CD LYS A 82 45.831 -1.938 -14.008 1.00 24.33 C +ATOM 469 CE LYS A 82 46.623 -0.730 -14.507 1.00 24.33 C +ATOM 470 NZ LYS A 82 47.133 0.033 -13.338 1.00 24.33 N +ATOM 471 N LEU A 83 45.338 -6.719 -15.466 1.00 23.66 N +ATOM 472 CA LEU A 83 46.268 -7.700 -16.028 1.00 23.66 C +ATOM 473 C LEU A 83 45.590 -8.557 -17.101 1.00 23.66 C +ATOM 474 O LEU A 83 46.168 -8.775 -18.165 1.00 23.66 O +ATOM 475 CB LEU A 83 46.843 -8.582 -14.905 1.00 23.66 C +ATOM 476 CG LEU A 83 47.822 -7.863 -13.958 1.00 23.66 C +ATOM 477 CD1 LEU A 83 48.240 -8.818 -12.839 1.00 23.66 C +ATOM 478 CD2 LEU A 83 49.087 -7.390 -14.681 1.00 23.66 C +ATOM 479 N SER A 84 44.342 -8.973 -16.876 1.00 22.05 N +ATOM 480 CA SER A 84 43.591 -9.761 -17.852 1.00 22.05 C +ATOM 481 C SER A 84 43.346 -9.006 -19.157 1.00 22.05 C +ATOM 482 O SER A 84 43.502 -9.571 -20.239 1.00 22.05 O +ATOM 483 CB SER A 84 42.289 -10.285 -17.237 1.00 22.05 C +ATOM 484 OG SER A 84 41.301 -9.285 -17.118 1.00 22.05 O +ATOM 485 N LYS A 85 43.058 -7.702 -19.075 1.00 23.01 N +ATOM 486 CA LYS A 85 42.844 -6.846 -20.241 1.00 23.01 C +ATOM 487 C LYS A 85 44.139 -6.550 -20.997 1.00 23.01 C +ATOM 488 O LYS A 85 44.117 -6.548 -22.222 1.00 23.01 O +ATOM 489 CB LYS A 85 42.135 -5.566 -19.781 1.00 23.01 C +ATOM 490 CG LYS A 85 41.637 -4.737 -20.971 1.00 23.01 C +ATOM 491 CD LYS A 85 40.847 -3.522 -20.478 1.00 23.01 C +ATOM 492 CE LYS A 85 40.333 -2.726 -21.680 1.00 23.01 C +ATOM 493 NZ LYS A 85 39.567 -1.532 -21.250 1.00 23.01 N +ATOM 494 N LEU A 86 45.240 -6.303 -20.284 1.00 22.05 N +ATOM 495 CA LEU A 86 46.534 -5.969 -20.890 1.00 22.05 C +ATOM 496 C LEU A 86 47.188 -7.171 -21.578 1.00 22.05 C +ATOM 497 O LEU A 86 47.725 -7.025 -22.669 1.00 22.05 O +ATOM 498 CB LEU A 86 47.473 -5.392 -19.814 1.00 22.05 C +ATOM 499 CG LEU A 86 47.106 -3.972 -19.345 1.00 22.05 C +ATOM 500 CD1 LEU A 86 47.963 -3.596 -18.134 1.00 22.05 C +ATOM 501 CD2 LEU A 86 47.329 -2.919 -20.433 1.00 22.05 C +ATOM 502 N TYR A 87 47.111 -8.353 -20.964 1.00 20.65 N +ATOM 503 CA TYR A 87 47.806 -9.556 -21.437 1.00 20.65 C +ATOM 504 C TYR A 87 46.890 -10.590 -22.093 1.00 20.65 C +ATOM 505 O TYR A 87 47.325 -11.710 -22.368 1.00 20.65 O +ATOM 506 CB TYR A 87 48.627 -10.143 -20.285 1.00 20.65 C +ATOM 507 CG TYR A 87 49.751 -9.225 -19.885 1.00 20.65 C +ATOM 508 CD1 TYR A 87 50.866 -9.111 -20.734 1.00 20.65 C +ATOM 509 CD2 TYR A 87 49.664 -8.457 -18.708 1.00 20.65 C +ATOM 510 CE1 TYR A 87 51.915 -8.243 -20.400 1.00 20.65 C +ATOM 511 CE2 TYR A 87 50.702 -7.565 -18.385 1.00 20.65 C +ATOM 512 CZ TYR A 87 51.827 -7.465 -19.234 1.00 20.65 C +ATOM 513 OH TYR A 87 52.843 -6.624 -18.955 1.00 20.65 O +ATOM 514 N HIS A 88 45.632 -10.224 -22.359 1.00 20.75 N +ATOM 515 CA HIS A 88 44.611 -11.118 -22.910 1.00 20.75 C +ATOM 516 C HIS A 88 44.504 -12.419 -22.106 1.00 20.75 C +ATOM 517 O HIS A 88 44.617 -13.518 -22.653 1.00 20.75 O +ATOM 518 CB HIS A 88 44.856 -11.353 -24.408 1.00 20.75 C +ATOM 519 CG HIS A 88 44.979 -10.084 -25.204 1.00 20.75 C +ATOM 520 ND1 HIS A 88 43.969 -9.186 -25.456 1.00 20.75 N +ATOM 521 CD2 HIS A 88 46.112 -9.599 -25.802 1.00 20.75 C +ATOM 522 CE1 HIS A 88 44.479 -8.187 -26.195 1.00 20.75 C +ATOM 523 NE2 HIS A 88 45.777 -8.405 -26.442 1.00 20.75 N +ATOM 524 N LEU A 89 44.339 -12.284 -20.791 1.00 21.06 N +ATOM 525 CA LEU A 89 44.126 -13.413 -19.893 1.00 21.06 C +ATOM 526 C LEU A 89 42.636 -13.576 -19.617 1.00 21.06 C +ATOM 527 O LEU A 89 41.872 -12.614 -19.599 1.00 21.06 O +ATOM 528 CB LEU A 89 44.948 -13.305 -18.594 1.00 21.06 C +ATOM 529 CG LEU A 89 46.422 -12.910 -18.788 1.00 21.06 C +ATOM 530 CD1 LEU A 89 47.067 -12.615 -17.438 1.00 21.06 C +ATOM 531 CD2 LEU A 89 47.204 -14.028 -19.475 1.00 21.06 C +ATOM 532 N THR A 90 42.229 -14.808 -19.365 1.00 21.65 N +ATOM 533 CA THR A 90 40.911 -15.153 -18.846 1.00 21.65 C +ATOM 534 C THR A 90 41.122 -16.028 -17.624 1.00 21.65 C +ATOM 535 O THR A 90 41.902 -16.978 -17.659 1.00 21.65 O +ATOM 536 CB THR A 90 40.070 -15.860 -19.910 1.00 21.65 C +ATOM 537 OG1 THR A 90 39.886 -14.987 -21.003 1.00 21.65 O +ATOM 538 CG2 THR A 90 38.677 -16.234 -19.401 1.00 21.65 C +ATOM 539 N TRP A 91 40.443 -15.694 -16.536 1.00 24.54 N +ATOM 540 CA TRP A 91 40.562 -16.410 -15.277 1.00 24.54 C +ATOM 541 C TRP A 91 39.255 -17.128 -14.946 1.00 24.54 C +ATOM 542 O TRP A 91 38.170 -16.655 -15.278 1.00 24.54 O +ATOM 543 CB TRP A 91 41.015 -15.443 -14.184 1.00 24.54 C +ATOM 544 CG TRP A 91 40.124 -14.253 -14.019 1.00 24.54 C +ATOM 545 CD1 TRP A 91 40.280 -13.055 -14.629 1.00 24.54 C +ATOM 546 CD2 TRP A 91 38.925 -14.135 -13.196 1.00 24.54 C +ATOM 547 NE1 TRP A 91 39.257 -12.207 -14.251 1.00 24.54 N +ATOM 548 CE2 TRP A 91 38.410 -12.812 -13.344 1.00 24.54 C +ATOM 549 CE3 TRP A 91 38.230 -15.007 -12.329 1.00 24.54 C +ATOM 550 CZ2 TRP A 91 37.284 -12.364 -12.639 1.00 24.54 C +ATOM 551 CZ3 TRP A 91 37.082 -14.575 -11.636 1.00 24.54 C +ATOM 552 CH2 TRP A 91 36.618 -13.253 -11.777 1.00 24.54 C +ATOM 553 N TYR A 92 39.364 -18.279 -14.299 1.00 24.33 N +ATOM 554 CA TYR A 92 38.255 -19.087 -13.812 1.00 24.33 C +ATOM 555 C TYR A 92 38.564 -19.514 -12.385 1.00 24.33 C +ATOM 556 O TYR A 92 39.698 -19.872 -12.080 1.00 24.33 O +ATOM 557 CB TYR A 92 38.060 -20.297 -14.728 1.00 24.33 C +ATOM 558 CG TYR A 92 37.033 -21.299 -14.236 1.00 24.33 C +ATOM 559 CD1 TYR A 92 37.446 -22.559 -13.756 1.00 24.33 C +ATOM 560 CD2 TYR A 92 35.666 -20.959 -14.226 1.00 24.33 C +ATOM 561 CE1 TYR A 92 36.494 -23.479 -13.277 1.00 24.33 C +ATOM 562 CE2 TYR A 92 34.713 -21.879 -13.753 1.00 24.33 C +ATOM 563 CZ TYR A 92 35.125 -23.140 -13.272 1.00 24.33 C +ATOM 564 OH TYR A 92 34.194 -24.012 -12.808 1.00 24.33 O +ATOM 565 N TYR A 93 37.575 -19.441 -11.508 1.00 26.97 N +ATOM 566 CA TYR A 93 37.742 -19.772 -10.103 1.00 26.97 C +ATOM 567 C TYR A 93 36.696 -20.807 -9.706 1.00 26.97 C +ATOM 568 O TYR A 93 35.502 -20.534 -9.817 1.00 26.97 O +ATOM 569 CB TYR A 93 37.652 -18.482 -9.283 1.00 26.97 C +ATOM 570 CG TYR A 93 37.862 -18.724 -7.811 1.00 26.97 C +ATOM 571 CD1 TYR A 93 36.831 -18.457 -6.890 1.00 26.97 C +ATOM 572 CD2 TYR A 93 39.088 -19.262 -7.380 1.00 26.97 C +ATOM 573 CE1 TYR A 93 37.034 -18.730 -5.529 1.00 26.97 C +ATOM 574 CE2 TYR A 93 39.279 -19.583 -6.031 1.00 26.97 C +ATOM 575 CZ TYR A 93 38.242 -19.323 -5.119 1.00 26.97 C +ATOM 576 OH TYR A 93 38.386 -19.723 -3.849 1.00 26.97 O +ATOM 577 N ASP A 94 37.149 -21.976 -9.264 1.00 38.20 N +ATOM 578 CA ASP A 94 36.318 -23.141 -8.921 1.00 38.20 C +ATOM 579 C ASP A 94 36.052 -23.262 -7.408 1.00 38.20 C +ATOM 580 O ASP A 94 35.791 -24.345 -6.901 1.00 38.20 O +ATOM 581 CB ASP A 94 36.972 -24.413 -9.491 1.00 38.20 C +ATOM 582 CG ASP A 94 38.219 -24.849 -8.710 1.00 38.20 C +ATOM 583 OD1 ASP A 94 38.760 -24.022 -7.939 1.00 38.20 O +ATOM 584 OD2 ASP A 94 38.674 -25.990 -8.921 1.00 38.20 O +ATOM 585 N GLU A 95 36.154 -22.146 -6.677 1.00 34.70 N +ATOM 586 CA GLU A 95 36.099 -22.071 -5.208 1.00 34.70 C +ATOM 587 C GLU A 95 37.325 -22.615 -4.465 1.00 34.70 C +ATOM 588 O GLU A 95 37.423 -22.423 -3.250 1.00 34.70 O +ATOM 589 CB GLU A 95 34.796 -22.634 -4.622 1.00 34.70 C +ATOM 590 CG GLU A 95 33.547 -22.101 -5.334 1.00 34.70 C +ATOM 591 CD GLU A 95 32.254 -22.432 -4.582 1.00 34.70 C +ATOM 592 OE1 GLU A 95 31.276 -21.689 -4.838 1.00 34.70 O +ATOM 593 OE2 GLU A 95 32.263 -23.334 -3.716 1.00 34.70 O +ATOM 594 N ASN A 96 38.317 -23.189 -5.152 1.00 29.54 N +ATOM 595 CA ASN A 96 39.544 -23.669 -4.521 1.00 29.54 C +ATOM 596 C ASN A 96 40.812 -23.114 -5.181 1.00 29.54 C +ATOM 597 O ASN A 96 41.635 -22.506 -4.499 1.00 29.54 O +ATOM 598 CB ASN A 96 39.510 -25.202 -4.498 1.00 29.54 C +ATOM 599 CG ASN A 96 40.587 -25.782 -3.603 1.00 29.54 C +ATOM 600 OD1 ASN A 96 41.028 -25.184 -2.628 1.00 29.54 O +ATOM 601 ND2 ASN A 96 41.033 -26.980 -3.893 1.00 29.54 N +ATOM 602 N ILE A 97 40.935 -23.270 -6.496 1.00 23.86 N +ATOM 603 CA ILE A 97 42.087 -22.904 -7.311 1.00 23.86 C +ATOM 604 C ILE A 97 41.684 -21.796 -8.291 1.00 23.86 C +ATOM 605 O ILE A 97 40.638 -21.825 -8.940 1.00 23.86 O +ATOM 606 CB ILE A 97 42.663 -24.149 -8.031 1.00 23.86 C +ATOM 607 CG1 ILE A 97 43.004 -25.273 -7.017 1.00 23.86 C +ATOM 608 CG2 ILE A 97 43.896 -23.770 -8.878 1.00 23.86 C +ATOM 609 CD1 ILE A 97 43.550 -26.565 -7.633 1.00 23.86 C +ATOM 610 N LEU A 98 42.537 -20.783 -8.410 1.00 22.52 N +ATOM 611 CA LEU A 98 42.420 -19.765 -9.444 1.00 22.52 C +ATOM 612 C LEU A 98 43.136 -20.242 -10.705 1.00 22.52 C +ATOM 613 O LEU A 98 44.358 -20.330 -10.741 1.00 22.52 O +ATOM 614 CB LEU A 98 42.967 -18.440 -8.897 1.00 22.52 C +ATOM 615 CG LEU A 98 42.886 -17.278 -9.901 1.00 22.52 C +ATOM 616 CD1 LEU A 98 41.441 -16.928 -10.271 1.00 22.52 C +ATOM 617 CD2 LEU A 98 43.539 -16.045 -9.281 1.00 22.52 C +ATOM 618 N TYR A 99 42.380 -20.522 -11.755 1.00 20.95 N +ATOM 619 CA TYR A 99 42.920 -20.899 -13.049 1.00 20.95 C +ATOM 620 C TYR A 99 43.072 -19.673 -13.939 1.00 20.95 C +ATOM 621 O TYR A 99 42.123 -18.909 -14.109 1.00 20.95 O +ATOM 622 CB TYR A 99 42.022 -21.930 -13.718 1.00 20.95 C +ATOM 623 CG TYR A 99 41.919 -23.240 -12.978 1.00 20.95 C +ATOM 624 CD1 TYR A 99 42.757 -24.311 -13.342 1.00 20.95 C +ATOM 625 CD2 TYR A 99 41.002 -23.388 -11.921 1.00 20.95 C +ATOM 626 CE1 TYR A 99 42.709 -25.522 -12.630 1.00 20.95 C +ATOM 627 CE2 TYR A 99 40.946 -24.601 -11.217 1.00 20.95 C +ATOM 628 CZ TYR A 99 41.800 -25.668 -11.561 1.00 20.95 C +ATOM 629 OH TYR A 99 41.720 -26.846 -10.897 1.00 20.95 O +ATOM 630 N ILE A 100 44.244 -19.498 -14.542 1.00 20.30 N +ATOM 631 CA ILE A 100 44.545 -18.384 -15.440 1.00 20.30 C +ATOM 632 C ILE A 100 45.052 -18.933 -16.769 1.00 20.30 C +ATOM 633 O ILE A 100 46.084 -19.605 -16.842 1.00 20.30 O +ATOM 634 CB ILE A 100 45.527 -17.382 -14.797 1.00 20.30 C +ATOM 635 CG1 ILE A 100 44.998 -16.915 -13.418 1.00 20.30 C +ATOM 636 CG2 ILE A 100 45.719 -16.195 -15.763 1.00 20.30 C +ATOM 637 CD1 ILE A 100 45.908 -15.936 -12.677 1.00 20.30 C +ATOM 638 N TYR A 101 44.322 -18.598 -17.825 1.00 19.73 N +ATOM 639 CA TYR A 101 44.580 -19.010 -19.199 1.00 19.73 C +ATOM 640 C TYR A 101 44.716 -17.789 -20.103 1.00 19.73 C +ATOM 641 O TYR A 101 44.293 -16.683 -19.753 1.00 19.73 O +ATOM 642 CB TYR A 101 43.436 -19.916 -19.665 1.00 19.73 C +ATOM 643 CG TYR A 101 43.381 -21.241 -18.934 1.00 19.73 C +ATOM 644 CD1 TYR A 101 44.196 -22.310 -19.353 1.00 19.73 C +ATOM 645 CD2 TYR A 101 42.505 -21.413 -17.847 1.00 19.73 C +ATOM 646 CE1 TYR A 101 44.168 -23.540 -18.672 1.00 19.73 C +ATOM 647 CE2 TYR A 101 42.416 -22.668 -17.219 1.00 19.73 C +ATOM 648 CZ TYR A 101 43.278 -23.714 -17.594 1.00 19.73 C +ATOM 649 OH TYR A 101 43.217 -24.888 -16.922 1.00 19.73 O +ATOM 650 N LYS A 102 45.268 -17.984 -21.298 1.00 19.96 N +ATOM 651 CA LYS A 102 45.185 -16.979 -22.362 1.00 19.96 C +ATOM 652 C LYS A 102 43.802 -17.023 -23.013 1.00 19.96 C +ATOM 653 O LYS A 102 43.180 -18.076 -23.130 1.00 19.96 O +ATOM 654 CB LYS A 102 46.299 -17.225 -23.383 1.00 19.96 C +ATOM 655 CG LYS A 102 47.685 -16.747 -22.935 1.00 19.96 C +ATOM 656 CD LYS A 102 47.847 -15.245 -23.200 1.00 19.96 C +ATOM 657 CE LYS A 102 49.246 -14.768 -22.823 1.00 19.96 C +ATOM 658 NZ LYS A 102 49.455 -13.365 -23.243 1.00 19.96 N +ATOM 659 N THR A 103 43.319 -15.887 -23.508 1.00 22.95 N +ATOM 660 CA THR A 103 42.008 -15.799 -24.178 1.00 22.95 C +ATOM 661 C THR A 103 41.932 -16.655 -25.452 1.00 22.95 C +ATOM 662 O THR A 103 40.843 -17.071 -25.841 1.00 22.95 O +ATOM 663 CB THR A 103 41.673 -14.328 -24.478 1.00 22.95 C +ATOM 664 OG1 THR A 103 41.697 -13.594 -23.274 1.00 22.95 O +ATOM 665 CG2 THR A 103 40.278 -14.102 -25.058 1.00 22.95 C +ATOM 666 N ASN A 104 43.064 -16.969 -26.093 1.00 21.72 N +ATOM 667 CA ASN A 104 43.124 -17.867 -27.255 1.00 21.72 C +ATOM 668 C ASN A 104 42.994 -19.362 -26.897 1.00 21.72 C +ATOM 669 O ASN A 104 42.805 -20.174 -27.796 1.00 21.72 O +ATOM 670 CB ASN A 104 44.403 -17.586 -28.070 1.00 21.72 C +ATOM 671 CG ASN A 104 45.687 -17.964 -27.349 1.00 21.72 C +ATOM 672 OD1 ASN A 104 45.710 -18.228 -26.165 1.00 21.72 O +ATOM 673 ND2 ASN A 104 46.812 -17.947 -28.018 1.00 21.72 N +ATOM 674 N GLU A 105 43.060 -19.729 -25.614 1.00 20.95 N +ATOM 675 CA GLU A 105 42.858 -21.103 -25.130 1.00 20.95 C +ATOM 676 C GLU A 105 41.382 -21.406 -24.830 1.00 20.95 C +ATOM 677 O GLU A 105 41.035 -22.525 -24.453 1.00 20.95 O +ATOM 678 CB GLU A 105 43.725 -21.360 -23.887 1.00 20.95 C +ATOM 679 CG GLU A 105 45.222 -21.217 -24.182 1.00 20.95 C +ATOM 680 CD GLU A 105 46.078 -21.408 -22.929 1.00 20.95 C +ATOM 681 OE1 GLU A 105 47.096 -22.126 -22.997 1.00 20.95 O +ATOM 682 OE2 GLU A 105 45.797 -20.788 -21.881 1.00 20.95 O +ATOM 683 N ILE A 106 40.495 -20.417 -25.000 1.00 21.94 N +ATOM 684 CA ILE A 106 39.050 -20.618 -24.891 1.00 21.94 C +ATOM 685 C ILE A 106 38.598 -21.511 -26.041 1.00 21.94 C +ATOM 686 O ILE A 106 38.690 -21.142 -27.213 1.00 21.94 O +ATOM 687 CB ILE A 106 38.278 -19.284 -24.878 1.00 21.94 C +ATOM 688 CG1 ILE A 106 38.678 -18.459 -23.635 1.00 21.94 C +ATOM 689 CG2 ILE A 106 36.751 -19.534 -24.898 1.00 21.94 C +ATOM 690 CD1 ILE A 106 38.060 -17.057 -23.604 1.00 21.94 C +ATOM 691 N SER A 107 38.029 -22.655 -25.689 1.00 21.22 N +ATOM 692 CA SER A 107 37.500 -23.619 -26.644 1.00 21.22 C +ATOM 693 C SER A 107 36.001 -23.828 -26.438 1.00 21.22 C +ATOM 694 O SER A 107 35.355 -23.192 -25.595 1.00 21.22 O +ATOM 695 CB SER A 107 38.333 -24.904 -26.594 1.00 21.22 C +ATOM 696 OG SER A 107 38.238 -25.525 -25.335 1.00 21.22 O +ATOM 697 N ARG A 108 35.419 -24.660 -27.300 1.00 20.80 N +ATOM 698 CA ARG A 108 34.013 -25.053 -27.251 1.00 20.80 C +ATOM 699 C ARG A 108 33.936 -26.568 -27.265 1.00 20.80 C +ATOM 700 O ARG A 108 34.682 -27.204 -28.004 1.00 20.80 O +ATOM 701 CB ARG A 108 33.229 -24.468 -28.429 1.00 20.80 C +ATOM 702 CG ARG A 108 33.212 -22.935 -28.441 1.00 20.80 C +ATOM 703 CD ARG A 108 32.514 -22.492 -29.722 1.00 20.80 C +ATOM 704 NE ARG A 108 32.528 -21.028 -29.896 1.00 20.80 N +ATOM 705 CZ ARG A 108 31.564 -20.342 -30.482 1.00 20.80 C +ATOM 706 NH1 ARG A 108 30.394 -20.859 -30.731 1.00 20.80 N +ATOM 707 NH2 ARG A 108 31.772 -19.102 -30.837 1.00 20.80 N +ATOM 708 N SER A 109 33.020 -27.116 -26.485 1.00 20.39 N +ATOM 709 CA SER A 109 32.708 -28.540 -26.462 1.00 20.39 C +ATOM 710 C SER A 109 31.200 -28.735 -26.513 1.00 20.39 C +ATOM 711 O SER A 109 30.453 -27.962 -25.913 1.00 20.39 O +ATOM 712 CB SER A 109 33.292 -29.190 -25.210 1.00 20.39 C +ATOM 713 OG SER A 109 33.111 -30.585 -25.297 1.00 20.39 O +ATOM 714 N ILE A 110 30.761 -29.760 -27.238 1.00 20.21 N +ATOM 715 CA ILE A 110 29.362 -30.178 -27.280 1.00 20.21 C +ATOM 716 C ILE A 110 29.253 -31.464 -26.470 1.00 20.21 C +ATOM 717 O ILE A 110 29.992 -32.414 -26.720 1.00 20.21 O +ATOM 718 CB ILE A 110 28.847 -30.340 -28.727 1.00 20.21 C +ATOM 719 CG1 ILE A 110 28.946 -28.995 -29.485 1.00 20.21 C +ATOM 720 CG2 ILE A 110 27.392 -30.849 -28.712 1.00 20.21 C +ATOM 721 CD1 ILE A 110 28.565 -29.082 -30.969 1.00 20.21 C +ATOM 722 N ILE A 111 28.329 -31.487 -25.514 1.00 21.55 N +ATOM 723 CA ILE A 111 28.070 -32.648 -24.660 1.00 21.55 C +ATOM 724 C ILE A 111 26.642 -33.118 -24.925 1.00 21.55 C +ATOM 725 O ILE A 111 25.701 -32.326 -24.847 1.00 21.55 O +ATOM 726 CB ILE A 111 28.316 -32.306 -23.176 1.00 21.55 C +ATOM 727 CG1 ILE A 111 29.759 -31.799 -22.930 1.00 21.55 C +ATOM 728 CG2 ILE A 111 28.024 -33.530 -22.287 1.00 21.55 C +ATOM 729 CD1 ILE A 111 29.910 -31.092 -21.581 1.00 21.55 C +ATOM 730 N THR A 112 26.491 -34.399 -25.250 1.00 20.95 N +ATOM 731 CA THR A 112 25.222 -35.045 -25.623 1.00 20.95 C +ATOM 732 C THR A 112 24.907 -36.179 -24.638 1.00 20.95 C +ATOM 733 O THR A 112 25.314 -37.317 -24.893 1.00 20.95 O +ATOM 734 CB THR A 112 25.305 -35.600 -27.061 1.00 20.95 C +ATOM 735 OG1 THR A 112 26.454 -36.401 -27.207 1.00 20.95 O +ATOM 736 CG2 THR A 112 25.379 -34.531 -28.153 1.00 20.95 C +ATOM 737 N PRO A 113 24.248 -35.905 -23.495 1.00 23.79 N +ATOM 738 CA PRO A 113 23.811 -36.963 -22.587 1.00 23.79 C +ATOM 739 C PRO A 113 22.767 -37.869 -23.252 1.00 23.79 C +ATOM 740 O PRO A 113 21.949 -37.416 -24.061 1.00 23.79 O +ATOM 741 CB PRO A 113 23.257 -36.254 -21.351 1.00 23.79 C +ATOM 742 CG PRO A 113 22.802 -34.898 -21.892 1.00 23.79 C +ATOM 743 CD PRO A 113 23.802 -34.605 -23.010 1.00 23.79 C +ATOM 744 N THR A 114 22.801 -39.157 -22.908 1.00 25.41 N +ATOM 745 CA THR A 114 21.937 -40.178 -23.525 1.00 25.41 C +ATOM 746 C THR A 114 20.724 -40.543 -22.674 1.00 25.41 C +ATOM 747 O THR A 114 19.667 -40.831 -23.229 1.00 25.41 O +ATOM 748 CB THR A 114 22.730 -41.440 -23.893 1.00 25.41 C +ATOM 749 OG1 THR A 114 23.502 -41.961 -22.831 1.00 25.41 O +ATOM 750 CG2 THR A 114 23.695 -41.190 -25.050 1.00 25.41 C +ATOM 751 N TYR A 115 20.853 -40.499 -21.347 1.00 31.03 N +ATOM 752 CA TYR A 115 19.810 -40.877 -20.387 1.00 31.03 C +ATOM 753 C TYR A 115 19.348 -39.703 -19.517 1.00 31.03 C +ATOM 754 O TYR A 115 18.255 -39.751 -18.956 1.00 31.03 O +ATOM 755 CB TYR A 115 20.334 -42.013 -19.495 1.00 31.03 C +ATOM 756 CG TYR A 115 20.750 -43.260 -20.250 1.00 31.03 C +ATOM 757 CD1 TYR A 115 19.799 -44.239 -20.595 1.00 31.03 C +ATOM 758 CD2 TYR A 115 22.091 -43.418 -20.641 1.00 31.03 C +ATOM 759 CE1 TYR A 115 20.192 -45.371 -21.335 1.00 31.03 C +ATOM 760 CE2 TYR A 115 22.480 -44.525 -21.415 1.00 31.03 C +ATOM 761 CZ TYR A 115 21.528 -45.503 -21.763 1.00 31.03 C +ATOM 762 OH TYR A 115 21.916 -46.586 -22.486 1.00 31.03 O +ATOM 763 N LEU A 116 20.173 -38.662 -19.377 1.00 26.18 N +ATOM 764 CA LEU A 116 19.865 -37.476 -18.586 1.00 26.18 C +ATOM 765 C LEU A 116 19.193 -36.378 -19.426 1.00 26.18 C +ATOM 766 O LEU A 116 19.653 -36.016 -20.511 1.00 26.18 O +ATOM 767 CB LEU A 116 21.159 -36.999 -17.908 1.00 26.18 C +ATOM 768 CG LEU A 116 20.966 -35.882 -16.871 1.00 26.18 C +ATOM 769 CD1 LEU A 116 20.249 -36.372 -15.609 1.00 26.18 C +ATOM 770 CD2 LEU A 116 22.338 -35.357 -16.463 1.00 26.18 C +ATOM 771 N ASP A 117 18.125 -35.808 -18.872 1.00 24.33 N +ATOM 772 CA ASP A 117 17.437 -34.640 -19.421 1.00 24.33 C +ATOM 773 C ASP A 117 18.255 -33.345 -19.230 1.00 24.33 C +ATOM 774 O ASP A 117 18.846 -33.110 -18.170 1.00 24.33 O +ATOM 775 CB ASP A 117 16.052 -34.554 -18.772 1.00 24.33 C +ATOM 776 CG ASP A 117 15.368 -33.247 -19.137 1.00 24.33 C +ATOM 777 OD1 ASP A 117 15.246 -32.984 -20.351 1.00 24.33 O +ATOM 778 OD2 ASP A 117 15.115 -32.476 -18.182 1.00 24.33 O +ATOM 779 N ILE A 118 18.280 -32.478 -20.248 1.00 24.06 N +ATOM 780 CA ILE A 118 19.091 -31.247 -20.233 1.00 24.06 C +ATOM 781 C ILE A 118 18.590 -30.240 -19.225 1.00 24.06 C +ATOM 782 O ILE A 118 19.410 -29.577 -18.595 1.00 24.06 O +ATOM 783 CB ILE A 118 19.096 -30.549 -21.598 1.00 24.06 C +ATOM 784 CG1 ILE A 118 19.915 -31.395 -22.556 1.00 24.06 C +ATOM 785 CG2 ILE A 118 19.703 -29.123 -21.583 1.00 24.06 C +ATOM 786 CD1 ILE A 118 19.561 -30.970 -23.964 1.00 24.06 C +ATOM 787 N ASP A 119 17.277 -30.071 -19.087 1.00 24.41 N +ATOM 788 CA ASP A 119 16.754 -29.067 -18.163 1.00 24.41 C +ATOM 789 C ASP A 119 17.125 -29.453 -16.720 1.00 24.41 C +ATOM 790 O ASP A 119 17.521 -28.603 -15.918 1.00 24.41 O +ATOM 791 CB ASP A 119 15.245 -28.879 -18.382 1.00 24.41 C +ATOM 792 CG ASP A 119 14.894 -28.176 -19.710 1.00 24.41 C +ATOM 793 OD1 ASP A 119 15.590 -27.198 -20.105 1.00 24.41 O +ATOM 794 OD2 ASP A 119 13.861 -28.530 -20.313 1.00 24.41 O +ATOM 795 N SER A 120 17.136 -30.758 -16.432 1.00 25.25 N +ATOM 796 CA SER A 120 17.678 -31.313 -15.190 1.00 25.25 C +ATOM 797 C SER A 120 19.189 -31.073 -15.048 1.00 25.25 C +ATOM 798 O SER A 120 19.634 -30.594 -14.005 1.00 25.25 O +ATOM 799 CB SER A 120 17.357 -32.807 -15.103 1.00 25.25 C +ATOM 800 OG SER A 120 15.956 -33.013 -15.174 1.00 25.25 O +ATOM 801 N LEU A 121 19.987 -31.346 -16.087 1.00 24.33 N +ATOM 802 CA LEU A 121 21.438 -31.108 -16.083 1.00 24.33 C +ATOM 803 C LEU A 121 21.784 -29.627 -15.869 1.00 24.33 C +ATOM 804 O LEU A 121 22.648 -29.307 -15.057 1.00 24.33 O +ATOM 805 CB LEU A 121 22.036 -31.634 -17.401 1.00 24.33 C +ATOM 806 CG LEU A 121 23.558 -31.430 -17.552 1.00 24.33 C +ATOM 807 CD1 LEU A 121 24.378 -32.165 -16.489 1.00 24.33 C +ATOM 808 CD2 LEU A 121 24.002 -31.920 -18.929 1.00 24.33 C +ATOM 809 N LEU A 122 21.101 -28.711 -16.560 1.00 24.76 N +ATOM 810 CA LEU A 122 21.306 -27.270 -16.416 1.00 24.76 C +ATOM 811 C LEU A 122 20.980 -26.786 -15.008 1.00 24.76 C +ATOM 812 O LEU A 122 21.688 -25.918 -14.507 1.00 24.76 O +ATOM 813 CB LEU A 122 20.438 -26.505 -17.424 1.00 24.76 C +ATOM 814 CG LEU A 122 20.946 -26.552 -18.869 1.00 24.76 C +ATOM 815 CD1 LEU A 122 19.947 -25.801 -19.746 1.00 24.76 C +ATOM 816 CD2 LEU A 122 22.314 -25.881 -19.047 1.00 24.76 C +ATOM 817 N LYS A 123 19.959 -27.363 -14.368 1.00 25.71 N +ATOM 818 CA LYS A 123 19.634 -27.083 -12.969 1.00 25.71 C +ATOM 819 C LYS A 123 20.761 -27.529 -12.029 1.00 25.71 C +ATOM 820 O LYS A 123 21.200 -26.745 -11.194 1.00 25.71 O +ATOM 821 CB LYS A 123 18.290 -27.739 -12.643 1.00 25.71 C +ATOM 822 CG LYS A 123 17.812 -27.355 -11.242 1.00 25.71 C +ATOM 823 CD LYS A 123 16.473 -28.030 -10.954 1.00 25.71 C +ATOM 824 CE LYS A 123 16.055 -27.660 -9.533 1.00 25.71 C +ATOM 825 NZ LYS A 123 14.786 -28.328 -9.165 1.00 25.71 N +ATOM 826 N TYR A 124 21.292 -28.743 -12.205 1.00 27.73 N +ATOM 827 CA TYR A 124 22.450 -29.199 -11.423 1.00 27.73 C +ATOM 828 C TYR A 124 23.685 -28.319 -11.650 1.00 27.73 C +ATOM 829 O TYR A 124 24.397 -27.999 -10.699 1.00 27.73 O +ATOM 830 CB TYR A 124 22.784 -30.664 -11.745 1.00 27.73 C +ATOM 831 CG TYR A 124 21.805 -31.667 -11.169 1.00 27.73 C +ATOM 832 CD1 TYR A 124 21.703 -31.820 -9.773 1.00 27.73 C +ATOM 833 CD2 TYR A 124 21.018 -32.462 -12.022 1.00 27.73 C +ATOM 834 CE1 TYR A 124 20.795 -32.750 -9.230 1.00 27.73 C +ATOM 835 CE2 TYR A 124 20.098 -33.382 -11.486 1.00 27.73 C +ATOM 836 CZ TYR A 124 19.987 -33.526 -10.086 1.00 27.73 C +ATOM 837 OH TYR A 124 19.101 -34.417 -9.567 1.00 27.73 O +ATOM 838 N LEU A 125 23.931 -27.888 -12.891 1.00 28.90 N +ATOM 839 CA LEU A 125 25.051 -27.006 -13.224 1.00 28.90 C +ATOM 840 C LEU A 125 24.869 -25.587 -12.671 1.00 28.90 C +ATOM 841 O LEU A 125 25.848 -24.992 -12.229 1.00 28.90 O +ATOM 842 CB LEU A 125 25.251 -26.977 -14.750 1.00 28.90 C +ATOM 843 CG LEU A 125 25.799 -28.284 -15.354 1.00 28.90 C +ATOM 844 CD1 LEU A 125 25.806 -28.171 -16.881 1.00 28.90 C +ATOM 845 CD2 LEU A 125 27.227 -28.578 -14.896 1.00 28.90 C +ATOM 846 N SER A 126 23.648 -25.042 -12.659 1.00 34.95 N +ATOM 847 CA SER A 126 23.382 -23.720 -12.080 1.00 34.95 C +ATOM 848 C SER A 126 23.524 -23.696 -10.563 1.00 34.95 C +ATOM 849 O SER A 126 23.908 -22.667 -10.014 1.00 34.95 O +ATOM 850 CB SER A 126 21.996 -23.198 -12.475 1.00 34.95 C +ATOM 851 OG SER A 126 20.917 -23.988 -12.010 1.00 34.95 O +ATOM 852 N ASP A 127 23.230 -24.816 -9.900 1.00 51.20 N +ATOM 853 CA ASP A 127 23.321 -24.935 -8.442 1.00 51.20 C +ATOM 854 C ASP A 127 24.768 -25.163 -7.966 1.00 51.20 C +ATOM 855 O ASP A 127 25.108 -24.798 -6.843 1.00 51.20 O +ATOM 856 CB ASP A 127 22.383 -26.064 -7.965 1.00 51.20 C +ATOM 857 CG ASP A 127 20.880 -25.755 -8.107 1.00 51.20 C +ATOM 858 OD1 ASP A 127 20.522 -24.588 -8.392 1.00 51.20 O +ATOM 859 OD2 ASP A 127 20.069 -26.692 -7.902 1.00 51.20 O +ATOM 860 N SER A 135 34.244 -16.522 -18.615 1.00 41.26 N +ATOM 861 CA SER A 135 35.237 -17.604 -18.500 1.00 41.26 C +ATOM 862 C SER A 135 34.643 -18.996 -18.698 1.00 41.26 C +ATOM 863 O SER A 135 35.353 -19.910 -19.119 1.00 41.26 O +ATOM 864 CB SER A 135 35.923 -17.512 -17.139 1.00 41.26 C +ATOM 865 OG SER A 135 34.972 -17.687 -16.106 1.00 41.26 O +ATOM 866 N CYS A 136 33.350 -19.160 -18.412 1.00 26.65 N +ATOM 867 CA CYS A 136 32.617 -20.380 -18.701 1.00 26.65 C +ATOM 868 C CYS A 136 31.141 -20.071 -18.951 1.00 26.65 C +ATOM 869 O CYS A 136 30.463 -19.487 -18.112 1.00 26.65 O +ATOM 870 CB CYS A 136 32.806 -21.372 -17.556 1.00 26.65 C +ATOM 871 SG CYS A 136 32.412 -23.043 -18.089 1.00 26.65 S +ATOM 872 N ASN A 137 30.645 -20.433 -20.127 1.00 24.82 N +ATOM 873 CA ASN A 137 29.281 -20.197 -20.560 1.00 24.82 C +ATOM 874 C ASN A 137 28.690 -21.507 -21.085 1.00 24.82 C +ATOM 875 O ASN A 137 29.307 -22.172 -21.915 1.00 24.82 O +ATOM 876 CB ASN A 137 29.295 -19.092 -21.627 1.00 24.82 C +ATOM 877 CG ASN A 137 27.894 -18.629 -21.964 1.00 24.82 C +ATOM 878 OD1 ASN A 137 27.173 -19.224 -22.740 1.00 24.82 O +ATOM 879 ND2 ASN A 137 27.454 -17.536 -21.388 1.00 24.82 N +ATOM 880 N VAL A 138 27.498 -21.856 -20.606 1.00 23.93 N +ATOM 881 CA VAL A 138 26.767 -23.053 -21.029 1.00 23.93 C +ATOM 882 C VAL A 138 25.485 -22.613 -21.719 1.00 23.93 C +ATOM 883 O VAL A 138 24.740 -21.784 -21.191 1.00 23.93 O +ATOM 884 CB VAL A 138 26.475 -23.996 -19.844 1.00 23.93 C +ATOM 885 CG1 VAL A 138 25.831 -25.301 -20.330 1.00 23.93 C +ATOM 886 CG2 VAL A 138 27.760 -24.355 -19.087 1.00 23.93 C +ATOM 887 N ARG A 139 25.220 -23.157 -22.909 1.00 25.25 N +ATOM 888 CA ARG A 139 23.999 -22.893 -23.677 1.00 25.25 C +ATOM 889 C ARG A 139 23.367 -24.203 -24.109 1.00 25.25 C +ATOM 890 O ARG A 139 24.058 -25.092 -24.592 1.00 25.25 O +ATOM 891 CB ARG A 139 24.298 -22.014 -24.902 1.00 25.25 C +ATOM 892 CG ARG A 139 24.915 -20.666 -24.506 1.00 25.25 C +ATOM 893 CD ARG A 139 25.134 -19.745 -25.707 1.00 25.25 C +ATOM 894 NE ARG A 139 23.864 -19.230 -26.259 1.00 25.25 N +ATOM 895 CZ ARG A 139 23.744 -18.388 -27.270 1.00 25.25 C +ATOM 896 NH1 ARG A 139 24.788 -17.930 -27.900 1.00 25.25 N +ATOM 897 NH2 ARG A 139 22.567 -17.996 -27.675 1.00 25.25 N +ATOM 898 N LYS A 140 22.044 -24.306 -23.988 1.00 23.01 N +ATOM 899 CA LYS A 140 21.310 -25.436 -24.565 1.00 23.01 C +ATOM 900 C LYS A 140 21.195 -25.305 -26.075 1.00 23.01 C +ATOM 901 O LYS A 140 20.948 -24.209 -26.584 1.00 23.01 O +ATOM 902 CB LYS A 140 19.946 -25.643 -23.905 1.00 23.01 C +ATOM 903 CG LYS A 140 18.906 -24.519 -24.074 1.00 23.01 C +ATOM 904 CD LYS A 140 17.658 -24.954 -23.293 1.00 23.01 C +ATOM 905 CE LYS A 140 16.448 -24.023 -23.391 1.00 23.01 C +ATOM 906 NZ LYS A 140 15.328 -24.618 -22.605 1.00 23.01 N +ATOM 907 N ILE A 141 21.311 -26.432 -26.764 1.00 25.71 N +ATOM 908 CA ILE A 141 20.990 -26.554 -28.180 1.00 25.71 C +ATOM 909 C ILE A 141 19.590 -27.160 -28.244 1.00 25.71 C +ATOM 910 O ILE A 141 19.405 -28.314 -27.911 1.00 25.71 O +ATOM 911 CB ILE A 141 22.064 -27.414 -28.881 1.00 25.71 C +ATOM 912 CG1 ILE A 141 23.450 -26.723 -28.820 1.00 25.71 C +ATOM 913 CG2 ILE A 141 21.649 -27.664 -30.343 1.00 25.71 C +ATOM 914 CD1 ILE A 141 24.623 -27.639 -29.196 1.00 25.71 C +ATOM 915 N THR A 142 18.572 -26.393 -28.624 1.00 50.20 N +ATOM 916 CA THR A 142 17.170 -26.854 -28.536 1.00 50.20 C +ATOM 917 C THR A 142 16.800 -27.940 -29.545 1.00 50.20 C +ATOM 918 O THR A 142 15.763 -28.573 -29.399 1.00 50.20 O +ATOM 919 CB THR A 142 16.204 -25.676 -28.714 1.00 50.20 C +ATOM 920 OG1 THR A 142 16.541 -24.941 -29.872 1.00 50.20 O +ATOM 921 CG2 THR A 142 16.278 -24.712 -27.529 1.00 50.20 C +ATOM 922 N THR A 143 17.606 -28.134 -30.589 1.00 38.34 N +ATOM 923 CA THR A 143 17.338 -29.120 -31.646 1.00 38.34 C +ATOM 924 C THR A 143 17.810 -30.526 -31.297 1.00 38.34 C +ATOM 925 O THR A 143 17.345 -31.483 -31.905 1.00 38.34 O +ATOM 926 CB THR A 143 18.003 -28.695 -32.961 1.00 38.34 C +ATOM 927 OG1 THR A 143 19.374 -28.432 -32.759 1.00 38.34 O +ATOM 928 CG2 THR A 143 17.379 -27.419 -33.524 1.00 38.34 C +ATOM 929 N PHE A 144 18.738 -30.665 -30.352 1.00 45.38 N +ATOM 930 CA PHE A 144 19.329 -31.942 -29.968 1.00 45.38 C +ATOM 931 C PHE A 144 19.332 -32.051 -28.450 1.00 45.38 C +ATOM 932 O PHE A 144 19.362 -31.032 -27.771 1.00 45.38 O +ATOM 933 CB PHE A 144 20.749 -32.046 -30.546 1.00 45.38 C +ATOM 934 CG PHE A 144 20.795 -31.976 -32.061 1.00 45.38 C +ATOM 935 CD1 PHE A 144 20.216 -33.005 -32.827 1.00 45.38 C +ATOM 936 CD2 PHE A 144 21.373 -30.866 -32.710 1.00 45.38 C +ATOM 937 CE1 PHE A 144 20.203 -32.921 -34.230 1.00 45.38 C +ATOM 938 CE2 PHE A 144 21.354 -30.780 -34.114 1.00 45.38 C +ATOM 939 CZ PHE A 144 20.767 -31.807 -34.873 1.00 45.38 C +ATOM 940 N ASN A 145 19.345 -33.269 -27.901 1.00 26.97 N +ATOM 941 CA ASN A 145 19.579 -33.433 -26.470 1.00 26.97 C +ATOM 942 C ASN A 145 21.071 -33.140 -26.150 1.00 26.97 C +ATOM 943 O ASN A 145 21.824 -34.039 -25.789 1.00 26.97 O +ATOM 944 CB ASN A 145 19.068 -34.791 -25.946 1.00 26.97 C +ATOM 945 CG ASN A 145 18.910 -34.782 -24.426 1.00 26.97 C +ATOM 946 OD1 ASN A 145 18.299 -33.887 -23.869 1.00 26.97 O +ATOM 947 ND2 ASN A 145 19.436 -35.751 -23.715 1.00 26.97 N +ATOM 948 N SER A 146 21.524 -31.890 -26.344 1.00 21.49 N +ATOM 949 CA SER A 146 22.896 -31.450 -26.113 1.00 21.49 C +ATOM 950 C SER A 146 23.053 -30.035 -25.542 1.00 21.49 C +ATOM 951 O SER A 146 22.207 -29.145 -25.686 1.00 21.49 O +ATOM 952 CB SER A 146 23.689 -31.597 -27.411 1.00 21.49 C +ATOM 953 OG SER A 146 23.281 -30.692 -28.416 1.00 21.49 O +ATOM 954 N ILE A 147 24.205 -29.823 -24.902 1.00 21.33 N +ATOM 955 CA ILE A 147 24.667 -28.522 -24.414 1.00 21.33 C +ATOM 956 C ILE A 147 25.963 -28.116 -25.124 1.00 21.33 C +ATOM 957 O ILE A 147 26.845 -28.943 -25.358 1.00 21.33 O +ATOM 958 CB ILE A 147 24.791 -28.499 -22.871 1.00 21.33 C +ATOM 959 CG1 ILE A 147 25.811 -29.529 -22.333 1.00 21.33 C +ATOM 960 CG2 ILE A 147 23.404 -28.705 -22.235 1.00 21.33 C +ATOM 961 CD1 ILE A 147 26.126 -29.390 -20.838 1.00 21.33 C +ATOM 962 N GLU A 148 26.088 -26.830 -25.453 1.00 21.16 N +ATOM 963 CA GLU A 148 27.340 -26.197 -25.874 1.00 21.16 C +ATOM 964 C GLU A 148 27.990 -25.553 -24.647 1.00 21.16 C +ATOM 965 O GLU A 148 27.421 -24.656 -24.017 1.00 21.16 O +ATOM 966 CB GLU A 148 27.096 -25.154 -26.986 1.00 21.16 C +ATOM 967 CG GLU A 148 28.406 -24.581 -27.579 1.00 21.16 C +ATOM 968 CD GLU A 148 28.214 -23.409 -28.576 1.00 21.16 C +ATOM 969 OE1 GLU A 148 29.208 -23.012 -29.244 1.00 21.16 O +ATOM 970 OE2 GLU A 148 27.114 -22.815 -28.624 1.00 21.16 O +ATOM 971 N VAL A 149 29.201 -25.994 -24.324 1.00 21.11 N +ATOM 972 CA VAL A 149 30.039 -25.408 -23.281 1.00 21.11 C +ATOM 973 C VAL A 149 31.148 -24.613 -23.952 1.00 21.11 C +ATOM 974 O VAL A 149 31.944 -25.161 -24.711 1.00 21.11 O +ATOM 975 CB VAL A 149 30.621 -26.485 -22.352 1.00 21.11 C +ATOM 976 CG1 VAL A 149 31.391 -25.833 -21.196 1.00 21.11 C +ATOM 977 CG2 VAL A 149 29.524 -27.369 -21.751 1.00 21.11 C +ATOM 978 N ARG A 150 31.223 -23.316 -23.658 1.00 21.60 N +ATOM 979 CA ARG A 150 32.292 -22.423 -24.112 1.00 21.60 C +ATOM 980 C ARG A 150 33.040 -21.881 -22.907 1.00 21.60 C +ATOM 981 O ARG A 150 32.453 -21.163 -22.101 1.00 21.60 O +ATOM 982 CB ARG A 150 31.683 -21.310 -24.973 1.00 21.60 C +ATOM 983 CG ARG A 150 32.768 -20.342 -25.462 1.00 21.60 C +ATOM 984 CD ARG A 150 32.177 -19.279 -26.387 1.00 21.60 C +ATOM 985 NE ARG A 150 33.235 -18.357 -26.851 1.00 21.60 N +ATOM 986 CZ ARG A 150 33.455 -17.116 -26.448 1.00 21.60 C +ATOM 987 NH1 ARG A 150 32.707 -16.521 -25.562 1.00 21.60 N +ATOM 988 NH2 ARG A 150 34.465 -16.447 -26.931 1.00 21.60 N +ATOM 989 N GLY A 151 34.335 -22.145 -22.802 1.00 21.44 N +ATOM 990 CA GLY A 151 35.106 -21.670 -21.660 1.00 21.44 C +ATOM 991 C GLY A 151 36.568 -22.073 -21.683 1.00 21.44 C +ATOM 992 O GLY A 151 37.084 -22.557 -22.691 1.00 21.44 O +ATOM 993 N VAL A 152 37.230 -21.845 -20.552 1.00 20.54 N +ATOM 994 CA VAL A 152 38.600 -22.316 -20.324 1.00 20.54 C +ATOM 995 C VAL A 152 38.653 -23.852 -20.247 1.00 20.54 C +ATOM 996 O VAL A 152 37.656 -24.468 -19.850 1.00 20.54 O +ATOM 997 CB VAL A 152 39.222 -21.680 -19.071 1.00 20.54 C +ATOM 998 CG1 VAL A 152 39.304 -20.153 -19.202 1.00 20.54 C +ATOM 999 CG2 VAL A 152 38.485 -22.035 -17.776 1.00 20.54 C +ATOM 1000 N PRO A 153 39.799 -24.485 -20.572 1.00 20.65 N +ATOM 1001 CA PRO A 153 39.922 -25.944 -20.627 1.00 20.65 C +ATOM 1002 C PRO A 153 39.458 -26.677 -19.361 1.00 20.65 C +ATOM 1003 O PRO A 153 38.756 -27.681 -19.470 1.00 20.65 O +ATOM 1004 CB PRO A 153 41.402 -26.210 -20.918 1.00 20.65 C +ATOM 1005 CG PRO A 153 41.817 -24.989 -21.734 1.00 20.65 C +ATOM 1006 CD PRO A 153 41.020 -23.864 -21.084 1.00 20.65 C +ATOM 1007 N GLU A 154 39.777 -26.160 -18.170 1.00 21.55 N +ATOM 1008 CA GLU A 154 39.370 -26.809 -16.917 1.00 21.55 C +ATOM 1009 C GLU A 154 37.850 -26.797 -16.714 1.00 21.55 C +ATOM 1010 O GLU A 154 37.280 -27.796 -16.282 1.00 21.55 O +ATOM 1011 CB GLU A 154 40.077 -26.153 -15.720 1.00 21.55 C +ATOM 1012 CG GLU A 154 39.990 -27.032 -14.460 1.00 21.55 C +ATOM 1013 CD GLU A 154 40.725 -28.385 -14.597 1.00 21.55 C +ATOM 1014 OE1 GLU A 154 40.394 -29.338 -13.858 1.00 21.55 O +ATOM 1015 OE2 GLU A 154 41.602 -28.522 -15.484 1.00 21.55 O +ATOM 1016 N CYS A 155 37.165 -25.713 -17.101 1.00 22.35 N +ATOM 1017 CA CYS A 155 35.709 -25.657 -16.982 1.00 22.35 C +ATOM 1018 C CYS A 155 35.035 -26.680 -17.902 1.00 22.35 C +ATOM 1019 O CYS A 155 34.116 -27.386 -17.490 1.00 22.35 O +ATOM 1020 CB CYS A 155 35.166 -24.257 -17.285 1.00 22.35 C +ATOM 1021 SG CYS A 155 33.430 -24.173 -16.760 1.00 22.35 S +ATOM 1022 N ILE A 156 35.515 -26.794 -19.146 1.00 21.49 N +ATOM 1023 CA ILE A 156 35.000 -27.785 -20.097 1.00 21.49 C +ATOM 1024 C ILE A 156 35.222 -29.193 -19.553 1.00 21.49 C +ATOM 1025 O ILE A 156 34.290 -29.990 -19.534 1.00 21.49 O +ATOM 1026 CB ILE A 156 35.647 -27.587 -21.483 1.00 21.49 C +ATOM 1027 CG1 ILE A 156 35.112 -26.276 -22.097 1.00 21.49 C +ATOM 1028 CG2 ILE A 156 35.350 -28.780 -22.412 1.00 21.49 C +ATOM 1029 CD1 ILE A 156 35.885 -25.815 -23.330 1.00 21.49 C +ATOM 1030 N LYS A 157 36.429 -29.489 -19.061 1.00 21.01 N +ATOM 1031 CA LYS A 157 36.769 -30.798 -18.500 1.00 21.01 C +ATOM 1032 C LYS A 157 35.868 -31.171 -17.321 1.00 21.01 C +ATOM 1033 O LYS A 157 35.382 -32.298 -17.276 1.00 21.01 O +ATOM 1034 CB LYS A 157 38.242 -30.762 -18.096 1.00 21.01 C +ATOM 1035 CG LYS A 157 38.752 -32.118 -17.596 1.00 21.01 C +ATOM 1036 CD LYS A 157 40.189 -31.928 -17.121 1.00 21.01 C +ATOM 1037 CE LYS A 157 40.739 -33.206 -16.496 1.00 21.01 C +ATOM 1038 NZ LYS A 157 41.989 -32.879 -15.774 1.00 21.01 N +ATOM 1039 N TYR A 158 35.633 -30.233 -16.405 1.00 23.40 N +ATOM 1040 CA TYR A 158 34.760 -30.433 -15.252 1.00 23.40 C +ATOM 1041 C TYR A 158 33.308 -30.694 -15.667 1.00 23.40 C +ATOM 1042 O TYR A 158 32.700 -31.664 -15.220 1.00 23.40 O +ATOM 1043 CB TYR A 158 34.853 -29.210 -14.333 1.00 23.40 C +ATOM 1044 CG TYR A 158 33.962 -29.328 -13.114 1.00 23.40 C +ATOM 1045 CD1 TYR A 158 32.713 -28.678 -13.078 1.00 23.40 C +ATOM 1046 CD2 TYR A 158 34.371 -30.129 -12.030 1.00 23.40 C +ATOM 1047 CE1 TYR A 158 31.877 -28.832 -11.954 1.00 23.40 C +ATOM 1048 CE2 TYR A 158 33.538 -30.278 -10.906 1.00 23.40 C +ATOM 1049 CZ TYR A 158 32.288 -29.626 -10.869 1.00 23.40 C +ATOM 1050 OH TYR A 158 31.471 -29.755 -9.791 1.00 23.40 O +ATOM 1051 N ILE A 159 32.752 -29.870 -16.561 1.00 22.41 N +ATOM 1052 CA ILE A 159 31.366 -30.042 -17.011 1.00 22.41 C +ATOM 1053 C ILE A 159 31.210 -31.360 -17.777 1.00 22.41 C +ATOM 1054 O ILE A 159 30.220 -32.057 -17.566 1.00 22.41 O +ATOM 1055 CB ILE A 159 30.881 -28.822 -17.825 1.00 22.41 C +ATOM 1056 CG1 ILE A 159 30.831 -27.575 -16.914 1.00 22.41 C +ATOM 1057 CG2 ILE A 159 29.472 -29.091 -18.397 1.00 22.41 C +ATOM 1058 CD1 ILE A 159 30.639 -26.255 -17.665 1.00 22.41 C +ATOM 1059 N THR A 160 32.178 -31.741 -18.614 1.00 20.65 N +ATOM 1060 CA THR A 160 32.163 -33.030 -19.318 1.00 20.65 C +ATOM 1061 C THR A 160 32.170 -34.202 -18.339 1.00 20.65 C +ATOM 1062 O THR A 160 31.289 -35.055 -18.417 1.00 20.65 O +ATOM 1063 CB THR A 160 33.337 -33.147 -20.301 1.00 20.65 C +ATOM 1064 OG1 THR A 160 33.288 -32.123 -21.269 1.00 20.65 O +ATOM 1065 CG2 THR A 160 33.312 -34.448 -21.098 1.00 20.65 C +ATOM 1066 N SER A 161 33.098 -34.237 -17.378 1.00 21.22 N +ATOM 1067 CA SER A 161 33.203 -35.358 -16.434 1.00 21.22 C +ATOM 1068 C SER A 161 31.994 -35.467 -15.499 1.00 21.22 C +ATOM 1069 O SER A 161 31.530 -36.573 -15.207 1.00 21.22 O +ATOM 1070 CB SER A 161 34.494 -35.245 -15.619 1.00 21.22 C +ATOM 1071 OG SER A 161 34.469 -34.110 -14.776 1.00 21.22 O +ATOM 1072 N LEU A 162 31.446 -34.329 -15.062 1.00 22.29 N +ATOM 1073 CA LEU A 162 30.228 -34.282 -14.261 1.00 22.29 C +ATOM 1074 C LEU A 162 29.021 -34.766 -15.069 1.00 22.29 C +ATOM 1075 O LEU A 162 28.268 -35.609 -14.586 1.00 22.29 O +ATOM 1076 CB LEU A 162 30.027 -32.851 -13.735 1.00 22.29 C +ATOM 1077 CG LEU A 162 28.769 -32.675 -12.863 1.00 22.29 C +ATOM 1078 CD1 LEU A 162 28.799 -33.550 -11.607 1.00 22.29 C +ATOM 1079 CD2 LEU A 162 28.655 -31.216 -12.430 1.00 22.29 C +ATOM 1080 N SER A 163 28.869 -34.288 -16.307 1.00 21.33 N +ATOM 1081 CA SER A 163 27.781 -34.709 -17.197 1.00 21.33 C +ATOM 1082 C SER A 163 27.831 -36.214 -17.459 1.00 21.33 C +ATOM 1083 O SER A 163 26.811 -36.882 -17.338 1.00 21.33 O +ATOM 1084 CB SER A 163 27.826 -33.968 -18.535 1.00 21.33 C +ATOM 1085 OG SER A 163 27.743 -32.569 -18.351 1.00 21.33 O +ATOM 1086 N GLU A 164 29.014 -36.772 -17.735 1.00 20.39 N +ATOM 1087 CA GLU A 164 29.199 -38.219 -17.909 1.00 20.39 C +ATOM 1088 C GLU A 164 28.852 -39.015 -16.644 1.00 20.39 C +ATOM 1089 O GLU A 164 28.279 -40.102 -16.726 1.00 20.39 O +ATOM 1090 CB GLU A 164 30.662 -38.518 -18.261 1.00 20.39 C +ATOM 1091 CG GLU A 164 31.043 -38.165 -19.704 1.00 20.39 C +ATOM 1092 CD GLU A 164 32.531 -38.439 -19.990 1.00 20.39 C +ATOM 1093 OE1 GLU A 164 33.007 -37.985 -21.052 1.00 20.39 O +ATOM 1094 OE2 GLU A 164 33.184 -39.128 -19.165 1.00 20.39 O +ATOM 1095 N SER A 165 29.200 -38.495 -15.465 1.00 21.65 N +ATOM 1096 CA SER A 165 28.922 -39.164 -14.189 1.00 21.65 C +ATOM 1097 C SER A 165 27.426 -39.164 -13.866 1.00 21.65 C +ATOM 1098 O SER A 165 26.879 -40.199 -13.487 1.00 21.65 O +ATOM 1099 CB SER A 165 29.714 -38.508 -13.059 1.00 21.65 C +ATOM 1100 OG SER A 165 31.101 -38.583 -13.337 1.00 21.65 O +ATOM 1101 N LEU A 166 26.751 -38.031 -14.077 1.00 21.44 N +ATOM 1102 CA LEU A 166 25.306 -37.908 -13.891 1.00 21.44 C +ATOM 1103 C LEU A 166 24.525 -38.731 -14.926 1.00 21.44 C +ATOM 1104 O LEU A 166 23.520 -39.345 -14.572 1.00 21.44 O +ATOM 1105 CB LEU A 166 24.904 -36.424 -13.960 1.00 21.44 C +ATOM 1106 CG LEU A 166 25.416 -35.545 -12.804 1.00 21.44 C +ATOM 1107 CD1 LEU A 166 25.041 -34.087 -13.082 1.00 21.44 C +ATOM 1108 CD2 LEU A 166 24.811 -35.944 -11.458 1.00 21.44 C +ATOM 1109 N ASP A 167 24.991 -38.803 -16.178 1.00 21.72 N +ATOM 1110 CA ASP A 167 24.375 -39.640 -17.216 1.00 21.72 C +ATOM 1111 C ASP A 167 24.494 -41.134 -16.872 1.00 21.72 C +ATOM 1112 O ASP A 167 23.527 -41.884 -17.013 1.00 21.72 O +ATOM 1113 CB ASP A 167 24.987 -39.320 -18.594 1.00 21.72 C +ATOM 1114 CG ASP A 167 24.105 -39.779 -19.766 1.00 21.72 C +ATOM 1115 OD1 ASP A 167 22.901 -39.451 -19.758 1.00 21.72 O +ATOM 1116 OD2 ASP A 167 24.604 -40.379 -20.746 1.00 21.72 O +ATOM 1117 N LYS A 168 25.639 -41.570 -16.321 1.00 23.27 N +ATOM 1118 CA LYS A 168 25.809 -42.937 -15.794 1.00 23.27 C +ATOM 1119 C LYS A 168 24.868 -43.230 -14.625 1.00 23.27 C +ATOM 1120 O LYS A 168 24.285 -44.310 -14.572 1.00 23.27 O +ATOM 1121 CB LYS A 168 27.264 -43.170 -15.362 1.00 23.27 C +ATOM 1122 CG LYS A 168 28.195 -43.386 -16.561 1.00 23.27 C +ATOM 1123 CD LYS A 168 29.655 -43.417 -16.095 1.00 23.27 C +ATOM 1124 CE LYS A 168 30.575 -43.499 -17.315 1.00 23.27 C +ATOM 1125 NZ LYS A 168 31.996 -43.302 -16.944 1.00 23.27 N +ATOM 1126 N GLU A 169 24.685 -42.288 -13.702 1.00 24.82 N +ATOM 1127 CA GLU A 169 23.738 -42.458 -12.594 1.00 24.82 C +ATOM 1128 C GLU A 169 22.288 -42.539 -13.099 1.00 24.82 C +ATOM 1129 O GLU A 169 21.523 -43.406 -12.665 1.00 24.82 O +ATOM 1130 CB GLU A 169 23.900 -41.316 -11.577 1.00 24.82 C +ATOM 1131 CG GLU A 169 23.138 -41.634 -10.279 1.00 24.82 C +ATOM 1132 CD GLU A 169 23.071 -40.456 -9.300 1.00 24.82 C +ATOM 1133 OE1 GLU A 169 22.060 -40.414 -8.551 1.00 24.82 O +ATOM 1134 OE2 GLU A 169 23.969 -39.593 -9.327 1.00 24.82 O +ATOM 1135 N ALA A 170 21.916 -41.680 -14.054 1.00 27.81 N +ATOM 1136 CA ALA A 170 20.612 -41.716 -14.712 1.00 27.81 C +ATOM 1137 C ALA A 170 20.384 -43.056 -15.427 1.00 27.81 C +ATOM 1138 O ALA A 170 19.322 -43.662 -15.277 1.00 27.81 O +ATOM 1139 CB ALA A 170 20.510 -40.527 -15.676 1.00 27.81 C +ATOM 1140 N GLN A 171 21.406 -43.580 -16.110 1.00 30.22 N +ATOM 1141 CA GLN A 171 21.373 -44.904 -16.725 1.00 30.22 C +ATOM 1142 C GLN A 171 21.136 -46.013 -15.688 1.00 30.22 C +ATOM 1143 O GLN A 171 20.316 -46.905 -15.914 1.00 30.22 O +ATOM 1144 CB GLN A 171 22.687 -45.145 -17.481 1.00 30.22 C +ATOM 1145 CG GLN A 171 22.613 -46.433 -18.309 1.00 30.22 C +ATOM 1146 CD GLN A 171 23.927 -46.786 -18.992 1.00 30.22 C +ATOM 1147 OE1 GLN A 171 25.012 -46.695 -18.448 1.00 30.22 O +ATOM 1148 NE2 GLN A 171 23.881 -47.350 -20.178 1.00 30.22 N +ATOM 1149 N SER A 172 21.829 -45.980 -14.547 1.00 37.66 N +ATOM 1150 CA SER A 172 21.633 -46.958 -13.468 1.00 37.66 C +ATOM 1151 C SER A 172 20.227 -46.890 -12.871 1.00 37.66 C +ATOM 1152 O SER A 172 19.623 -47.932 -12.619 1.00 37.66 O +ATOM 1153 CB SER A 172 22.672 -46.747 -12.366 1.00 37.66 C +ATOM 1154 OG SER A 172 23.948 -47.143 -12.827 1.00 37.66 O +ATOM 1155 N LYS A 173 19.670 -45.686 -12.692 1.00 42.85 N +ATOM 1156 CA LYS A 173 18.280 -45.510 -12.239 1.00 42.85 C +ATOM 1157 C LYS A 173 17.283 -46.058 -13.258 1.00 42.85 C +ATOM 1158 O LYS A 173 16.374 -46.782 -12.868 1.00 42.85 O +ATOM 1159 CB LYS A 173 17.997 -44.031 -11.930 1.00 42.85 C +ATOM 1160 CG LYS A 173 18.656 -43.606 -10.612 1.00 42.85 C +ATOM 1161 CD LYS A 173 18.432 -42.117 -10.315 1.00 42.85 C +ATOM 1162 CE LYS A 173 19.219 -41.765 -9.047 1.00 42.85 C +ATOM 1163 NZ LYS A 173 19.398 -40.308 -8.853 1.00 42.85 N +END diff --git a/tests/functional/af_trajectory_reference.json b/tests/functional/af_trajectory_reference.json index 1cc06d09..6b380f4d 100644 --- a/tests/functional/af_trajectory_reference.json +++ b/tests/functional/af_trajectory_reference.json @@ -2,115 +2,93 @@ "cycles": 2, "spread_runs": 5, "structures": { - "1VER": { + "1AK5": { "trajectory": [ [ - 0.378484, - 0.332842 + 0.368808, + 0.375264 ], [ - 0.336894, - 0.317116 + 0.270958, + 0.284292 ], [ - 0.330538, - 0.308256 + 0.268378, + 0.281478 ], [ - 0.312144, - 0.305218 + 0.228356, + 0.252214 ] ], - "spread": 7.1e-05, - "tolerance": 0.002 + "spread": 0.000656, + "tolerance": 0.005 }, - "6VHI": { + "1DAW": { "trajectory": [ [ - 0.332963, - 0.344154 + 0.386278, + 0.378892 ], [ - 0.287595, - 0.332427 + 0.314874, + 0.354788 ], [ - 0.282997, - 0.309342 + 0.314025, + 0.3516 ], [ - 0.269246, - 0.324271 + 0.287673, + 0.342073 ] ], - "spread": 0.000567, - "tolerance": 0.002 + "spread": 0.000173, + "tolerance": 0.005 }, - "1BYW": { + "3GR5": { "trajectory": [ [ - 0.396823, - 0.352145 + 0.468181, + 0.457217 ], [ - 0.356685, - 0.322779 + 0.400075, + 0.415961 ], [ - 0.359504, - 0.321488 + 0.398326, + 0.414747 ], [ - 0.336396, - 0.309443 + 0.344741, + 0.37339 ] ], - "spread": 0.000107, - "tolerance": 0.002 + "spread": 0.000958, + "tolerance": 0.005 }, - "6JZA": { - "trajectory": [ - [ - 0.416852, - 0.378224 - ], - [ - 0.36559, - 0.35986 - ], - [ - 0.366328, - 0.363104 - ], - [ - 0.329881, - 0.325672 - ] - ], - "spread": 0.004474, - "tolerance": 0.013422 - }, - "6SXW": { + "1VER": { "trajectory": [ [ - 0.430794, - 0.4249 + 0.372527, + 0.33176 ], [ - 0.383688, - 0.430362 + 0.318186, + 0.328353 ], [ - 0.387311, - 0.431517 + 0.304089, + 0.314967 ], [ - 0.340972, - 0.407762 + 0.297222, + 0.314039 ] ], - "spread": 9.8e-05, - "tolerance": 0.002 + "spread": 0.00045, + "tolerance": 0.005 } } } diff --git a/tests/functional/test_af_trajectory.py b/tests/functional/test_af_trajectory.py index b242acb6..de5eb5bd 100644 --- a/tests/functional/test_af_trajectory.py +++ b/tests/functional/test_af_trajectory.py @@ -9,16 +9,17 @@ means nothing on its own. Each structure therefore carries its **own** tolerance, measured over :data:`SPREAD_RUNS` independent runs when the reference was written, and committed alongside it. Sizing the tolerance from a couple of runs at test time does -not work: 6JZA is bimodal at this cycle count -- its trajectories land in one of two -basins about 0.0074 apart -- and two runs that happen to pick the same basin report a -spread 140 times too small. - -R-work and R-free are held to different bounds. R-work is reproducible to a few parts in -ten thousand, so it keeps the measured tolerance and is what catches a change that -actually moves the refinement. R-free is computed on the small free set and has a rare -second basin of its own, a few thousandths wide, that a handful of runs will usually -miss; it therefore carries :data:`RFREE_TOLERANCE_FLOOR`, wide enough to sit outside -that basin and still far inside the descent the trajectory shows. +not work: 6JZA, in an earlier version of this set, was bimodal at this cycle count -- +its trajectories landed in one of two basins about 0.0074 apart -- and two runs that +happened to pick the same basin reported a spread 140 times too small. + +R-work and R-free are held to different bounds. R-work is reproducible to a few +thousandths, run to run and across machines, so it has the tighter bound +(:data:`TOLERANCE_FLOOR` unless its measured spread is wider) and is what catches a +change that actually moves the refinement. R-free is computed on the free set and has a +rare second basin of its own, a few thousandths wide, that a handful of runs will +usually miss; it therefore carries :data:`RFREE_TOLERANCE_FLOOR`, wide enough to sit +outside that basin and still far inside the descent the trajectory shows. Regenerate deliberately, after a change meant to move these numbers:: @@ -32,12 +33,20 @@ import pytest import torch -#: AlphaFold-start structures, chosen so all five actually descend under refinement and -#: the space groups span centred tetragonal, centred monoclinic, hexagonal and trigonal. -CODES = ["1VER", "6VHI", "1BYW", "6JZA", "6SXW"] +#: AlphaFold-start structures that actually refine in two cycles: R-work and R-free both +#: fall and the run ends with R-work below R-free (see ``test_af_starts_refine``). The +#: space groups span cubic, centred monoclinic, hexagonal and centred tetragonal. The +#: earlier set (6VHI, 1BYW, 6JZA, 6SXW) had free sets of 40-80 reflections, too few to +#: tell R-free from noise: R-free rose mid-trajectory or ended below R-work. +CODES = ["1AK5", "1DAW", "3GR5", "1VER"] + +#: What "refines" means for a reference trajectory: how far R-work and R-free must fall +#: from first to last stage. +MIN_RWORK_DESCENT = 0.02 +MIN_RFREE_DESCENT = 0.01 #: Macro-cycles per trajectory. Two is enough for the descent to be visible while -#: keeping the whole test at a few seconds per structure. +#: keeping the whole test at well under a minute per structure. CYCLES = 2 #: Independent runs used to size each structure's tolerance at regeneration time. Enough @@ -48,8 +57,12 @@ SPREAD_MULTIPLE = 3.0 #: Tolerance floor for R-work, so a structure whose runs agree very closely is not held -#: to an unreasonably tight bound. -TOLERANCE_FLOOR = 0.002 +#: to an unreasonably tight bound. Five runs underestimate the spread: 1VER measured +#: 0.00045 yet deviated 0.0021 from its mean on the machine that wrote the reference, +#: and GitHub runners sat up to 0.0016 from a reference written elsewhere. 0.005 covers +#: both together and stays under a quarter of every structure's descent. Applied at test +#: time as well, so raising it does not need a regeneration. +TOLERANCE_FLOOR = 0.005 #: Tolerance floor for R-free, which moves in discrete basins rather than jitter. Set #: above the widest basin separation seen on this set and kept well inside every @@ -100,6 +113,24 @@ def _deviations(a, b): ) +def _refinement_problems(series): + """Why a trajectory does not count as refining, or an empty list if it does.""" + (work0, free0), (work1, free1) = series[0], series[-1] + problems = [] + if work1 > work0 - MIN_RWORK_DESCENT: + problems.append(f"R-work only moves from {work0:.4f} to {work1:.4f}") + if free1 > free0 - MIN_RFREE_DESCENT: + problems.append(f"R-free only moves from {free0:.4f} to {free1:.4f}") + if work1 >= free1: + problems.append(f"it ends with R-work {work1:.4f} not below R-free {free1:.4f}") + return problems + + +def _rwork_tolerance(entry): + """The R-work bound: this structure's measured tolerance, floored.""" + return max(float(entry["tolerance"]), TOLERANCE_FLOOR) + + def _rfree_tolerance(entry): """The R-free bound: this structure's measured tolerance, floored.""" return max(float(entry["tolerance"]), RFREE_TOLERANCE_FLOOR) @@ -131,13 +162,15 @@ def test_af_trajectory_matches_reference(code, reference, test_files_dir): entry = reference["structures"][code] expected = [tuple(point) for point in entry["trajectory"]] - rwork_tolerance = float(entry["tolerance"]) + rwork_tolerance = _rwork_tolerance(entry) rfree_tolerance = _rfree_tolerance(entry) observed = trajectory(pdb_path, mtz_path) assert len(observed) == len( expected ), f"{code}: trajectory has {len(observed)} stages, reference has {len(expected)}" + problems = _refinement_problems(observed) + assert not problems, f"{code} does not refine: {'; '.join(problems)}\n{observed}" dev_work, dev_free = _deviations(observed, expected) for label, deviation, tolerance in ( @@ -152,19 +185,17 @@ def test_af_trajectory_matches_reference(code, reference, test_files_dir): @pytest.mark.integration -def test_af_starts_actually_descend(reference): - """Each reference trajectory improves R-work, so a regression has signal to lose. +def test_af_starts_refine(reference): + """Each reference trajectory refines, so a regression has signal to lose. - A structure that barely moves under refinement cannot show a trajectory regression, - so the set is only useful while every member descends. + R-work and R-free both have to fall, and the run has to end with R-work below + R-free. A structure that barely moves cannot show a trajectory regression, and one + whose R-free rises or ends below R-work is tracking noise, not refinement. """ + assert set(reference["structures"]) == set(CODES) for code in CODES: - series = reference["structures"][code]["trajectory"] - first, last = series[0][0], series[-1][0] - assert last < first - 0.02, ( - f"{code} only moves R-work from {first:.4f} to {last:.4f}; it is too flat " - f"to serve as a trajectory probe" - ) + problems = _refinement_problems(reference["structures"][code]["trajectory"]) + assert not problems, f"{code} is not a refinement probe: {'; '.join(problems)}" @pytest.mark.integration @@ -181,7 +212,7 @@ def test_tolerances_are_tight_enough_to_detect_something(reference): descent = series[0][0] - series[-1][0] # Checks the widest bound actually applied, not the stored one: flooring R-free # would otherwise widen the real bound without this guard seeing it. - widest = max(float(entry["tolerance"]), _rfree_tolerance(entry)) + widest = max(_rwork_tolerance(entry), _rfree_tolerance(entry)) assert widest < descent / 4.0, ( f"{code}: tolerance {widest:.4f} is not small against its " f"own R-work descent of {descent:.4f}" @@ -203,13 +234,21 @@ def _write_reference(): ] spread = max(_max_deviation(run, mean) for run in runs) tolerance = max(TOLERANCE_FLOOR, SPREAD_MULTIPLE * spread) + problems = _refinement_problems(mean) + if problems: + raise SystemExit(f"{code} does not refine: {'; '.join(problems)}") structures[code] = { "trajectory": [[round(w, 6), round(f, 6)] for w, f in mean], "spread": round(spread, 6), "tolerance": round(tolerance, 6), } - print(f"{code}: spread={spread:.6f} tolerance={tolerance:.6f}", flush=True) + print( + f"{code}: R-work/R-free {mean[0][0]:.4f}/{mean[0][1]:.4f} -> " + f"{mean[-1][0]:.4f}/{mean[-1][1]:.4f}, spread={spread:.6f} " + f"tolerance={tolerance:.6f}", + flush=True, + ) REFERENCE.write_text( json.dumps( From afe81380e61f07ca4f47f966e7d67d022f9e75a2 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 14:02:12 +0200 Subject: [PATCH 193/250] Reuse cached components in the two-moment lambda comparisons Two tests compared intensity models at different lambda_twin values at 1e-6, but recomputed the component structure factors for each. On MPS a second structure-factor pass is not bit-reproducible -- the Metal density kernel accumulates with atomics, so two identical passes agree on only ~20% of values, differing by up to 5e-4 relative -- and the noise swamped the comparison. On CPU the passes are bit-identical, which hid it. lambda enters only downstream of the components, so the second call now reads the cache (recalc=False), as test_lambda_zero_reduces_to_the_squared_mean already does. The comparisons keep their tolerance and no longer depend on the kernel being deterministic. Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/integration/test_two_moment_intensity.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/tests/integration/test_two_moment_intensity.py b/tests/integration/test_two_moment_intensity.py index 9a4c2b1e..16e6ce42 100644 --- a/tests/integration/test_two_moment_intensity.py +++ b/tests/integration/test_two_moment_intensity.py @@ -64,8 +64,11 @@ def test_a_nonzero_lambda_changes_the_prediction(self, difference_collection): mc.set_lambda_twin(0.0) coherent = _target(dc, mc, scaler).intensity_model(recalc=True) + # Reuse the cached components: lambda enters only downstream of them, and + # a second structure-factor pass is not bit-reproducible on accelerators + # (atomic splatting), which would swamp the 1e-6 comparison below. mc.set_lambda_twin(0.5) - dispersed = _target(dc, mc, scaler).intensity_model(recalc=True) + dispersed = _target(dc, mc, scaler).intensity_model(recalc=False) assert not torch.allclose(coherent, dispersed) # Strictly positive: |dF|^2 has no sign. @@ -105,8 +108,9 @@ def test_the_reference_row_carries_no_variance(self, difference_collection): mc.set_lambda_twin(0.0) coherent = _target(dc, mc, scaler).intensity_model(recalc=True)[0] + # Same components for both: see test_a_nonzero_lambda_changes_the_prediction. mc.set_lambda_twin(0.9) - dispersed = _target(dc, mc, scaler).intensity_model(recalc=True)[0] + dispersed = _target(dc, mc, scaler).intensity_model(recalc=False)[0] assert torch.allclose(coherent, dispersed, rtol=1e-6) From bd19b44b703bbf4b863d6097335055d1e4d54cdb Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 14:02:12 +0200 Subject: [PATCH 194/250] Record the starting model by file name, not local path The refinement header wrote the input model's full path into REMARK 3 STARTING MODEL and _refine.pdbx_starting_model. A local path means nothing to a reader of the output and exposes the refiner's filesystem, and in the PDB header a long one ran past column 80 unwrapped. The file name alone is now recorded when metadata is collected and again when either format is rendered, matching _pdbx_initial_refinement_model, which already used the basename. A name too long for the line wraps onto continuation lines aligned under the value. The header tests located 3GR5.pdb by a cwd-relative path, so they failed when pytest ran from inside tests/, and the short relative path is what kept the overflow hidden. They now resolve it from __file__, and a new test checks that no part of a local path reaches either output. Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/io/test_refinement_header.py | 18 +++++++++++++-- torchref/io/metadata.py | 29 ++++++++++++++++++++----- 2 files changed, 40 insertions(+), 7 deletions(-) diff --git a/tests/unit/io/test_refinement_header.py b/tests/unit/io/test_refinement_header.py index 9da010e0..4a6fb5fc 100644 --- a/tests/unit/io/test_refinement_header.py +++ b/tests/unit/io/test_refinement_header.py @@ -15,6 +15,8 @@ previous program's output. """ +from pathlib import Path + import pandas as pd import pytest @@ -23,7 +25,7 @@ # 3GR5 was refined with REFMAC 5.1.24 and carries a full deposition header: # 420 lines including REMARK 2/3/500, JRNL, AUTHOR, SEQRES, SSBOND and SITE. -INPUT_PDB = "tests/files/pdb/3GR5.pdb" +INPUT_PDB = str(Path(__file__).resolve().parents[2] / "files" / "pdb" / "3GR5.pdb") #: PDB record order, abridged to the records this writer can emit. The format #: mandates this sequence; TITLE used to be written *after* REMARK 900. @@ -397,7 +399,19 @@ def test_starting_model_is_recorded_as_an_accession(tmp_path): initial = cats["_pdbx_initial_refinement_model"] assert initial["_pdbx_initial_refinement_model.accession_code"] == "3GR5" assert initial["_pdbx_initial_refinement_model.type"] == "experimental model" - assert cats["_refine"]["_refine.pdbx_starting_model"] == INPUT_PDB + assert cats["_refine"]["_refine.pdbx_starting_model"] == "3GR5.pdb" + + +@pytest.mark.unit +def test_starting_model_is_named_without_its_local_path(): + """The input's directory is the refiner's filesystem, not provenance.""" + meta = _refined_metadata() + meta.starting_model = "/home/someone/projects/secret_project/run_07/model.pdb" + cats = meta.render_cif_categories() + header = meta.render_pdb_header() + assert "/" not in cats["_refine"]["_refine.pdbx_starting_model"] + assert "secret_project" not in header + assert "REMARK 3 STARTING MODEL: model.pdb" in header.splitlines() @pytest.mark.unit diff --git a/torchref/io/metadata.py b/torchref/io/metadata.py index 4a1f378c..e9b0dc3d 100644 --- a/torchref/io/metadata.py +++ b/torchref/io/metadata.py @@ -12,6 +12,7 @@ from __future__ import annotations import json +import os from dataclasses import asdict, dataclass, field, fields from datetime import date from typing import Any, Dict, List, Optional @@ -83,7 +84,9 @@ class RefinementMetadata: title, authors Structure title and author names. starting_model : str, optional - Input model this refinement started from (path or PDB ID). + Input model this refinement started from (file name or PDB ID). Only + the file name is ever written: a local path means nothing to a reader + of the output and exposes the refiner's filesystem. rfree_selection : str, optional Where the free-set flags came from. Filled from ``ReflectionData.rfree_source``, so the values are that field's: @@ -310,7 +313,7 @@ def from_refinement(cls, refinement) -> RefinementMetadata: try: input_file = refinement.model.ctx.input_file if input_file: - meta.starting_model = str(input_file) + meta.starting_model = os.path.basename(str(input_file)) except Exception: pass @@ -636,7 +639,7 @@ def _render_remark3(self) -> List[str]: lines.append("REMARK 3") if self.starting_model: - lines.append(f"REMARK 3 STARTING MODEL: {self.starting_model}") + _wrap_starting_model(lines, os.path.basename(self.starting_model)) lines.append("REMARK 3") # The only free text in the block, and the caller wrote all of it. @@ -743,7 +746,7 @@ def render_cif_categories(self) -> Dict[str, Dict[str, str]]: if self.rfree_selection: ref["_refine.pdbx_R_Free_selection_details"] = self.rfree_selection if self.starting_model: - ref["_refine.pdbx_starting_model"] = self.starting_model + ref["_refine.pdbx_starting_model"] = os.path.basename(self.starting_model) if self.refinement_method: ref["_refine.pdbx_method_to_determine_struct"] = self.refinement_method if self.output_remarks: @@ -803,7 +806,6 @@ def _initial_model_category(starting_model: str) -> Dict[str, str]: alphanumerics, e.g. ``3GR5.pdb``) is reported as an accession code; anything else is named in ``details`` and left unaccessioned rather than guessed at. """ - import os import re basename = os.path.basename(starting_model) @@ -863,6 +865,23 @@ def _ident(lines: List[str], label: str, value: str) -> None: lines.append(prefix + current) +def _wrap_starting_model(lines: List[str], name: str) -> None: + """Append ``REMARK 3 STARTING MODEL: name``, wrapped if long. + + A file name has no spaces for a word wrap to break at, so a long one would + overrun column 80. Overflow continues on lines indented to the value column, + cut mid-string; nothing of the name is dropped. + """ + head = "REMARK 3 STARTING MODEL: " + cont = "REMARK 3" + " " * (len(head) - len("REMARK 3")) + width = 80 - len(head) + prefix, rest = head, name + while len(rest) > width: + lines.append(prefix + rest[:width]) + prefix, rest = cont, rest[width:] + lines.append(prefix + rest) + + def _wrap_remark3_text(lines: List[str], text: str) -> None: """Append free text as continuation-free ``REMARK 3`` lines. From 08cc5175964f2725c33086499998c6da3e770aeb Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 17:36:54 +0200 Subject: [PATCH 195/250] Restore cif_path with a ModelFT and drop the stray restraints in the refinement restore ModelFT.create_from_state_dict left cif_path in the state dict, where the non-strict load discarded it. Refinement.create_from_state_dict built Restraints(model), passing the model where the atom table belongs. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 1 + tests/unit/model/test_hydrogen_default.py | 9 ++++++--- torchref/model/model_ft.py | 2 ++ torchref/refinement/base_refinement.py | 14 +++----------- 4 files changed, 12 insertions(+), 14 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 4c09517e..8ce89eec 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- ``ModelFT.create_from_state_dict`` restores the restraint dictionary path (``cif_path``) as ``Model`` does, and ``Refinement.create_from_state_dict`` no longer builds a stray ``Restraints`` from the model; the restored model builds its own on first access. - The difference MTZ groups its columns into named datasets -- ``observed``, ``difference``, ``light_model``, ``extrapolated_light``, ``two_moment`` -- with one history line describing each, so ``FWT``/``PHWT`` reads as ``/torchref/extrapolated_light/FWT`` (the extrapolated light-state map ``2*FEXT - Fc``). Labels are unchanged and Coot still auto-opens it - ``torchref.difference-map``, ``torchref.difference-refine`` and ``torchref.validate-ded`` gain ``--ded-weight {sigma_d,inverse_variance,none}`` and ``--sigma-d-gamma``. The difference MTZ now carries the unweighted ``DF``/``SIGDF`` on ``PHDELWT`` with one mean-one weight column per scheme, ``W_SD`` and ``W_IVW`` (MTZ type W), and the observed-to-model scale ``KSCALE``; ``DELFWT`` is no longer written, build the map with ``torchref.mtz2map -csf DF -cw W_IVW -cphi PHDELWT``. Registered in ``torchref.maps.ded_weights`` - Added the ``sigma_D`` estimator (``torchref.refinement.model_error_estimation.sigma_d``): the expected true difference power per resolution shell, ``mean(dF_obs^2) - mean(sigma^2)`` with a fitted ``F_dark^gamma`` amplitude law and DerSimonian-Laird shrinkage of the signed shell power toward a decaying exponential in ``d*^2`` (fitted on all shells, so a dataset without a difference yields no power instead of the positive half of its noise), giving the Wiener weight ``S/(S + sigma^2)`` and, with a difference model, ``alpha``/``beta_model``. Inverse-variance weights suppress the strong reflections whose difference power is 10-70x that of weak ones; on independent half-datasets a Wiener weight with the true power raised map agreement 1.2-1.8x in effective patterns. The single-dataset estimate inherits the calibration of the reported sigmas, and on the campaign TD1 data (sigmas ~1.5x too large at high resolution) it emptied 60-90 % of the shells, so inverse variance stays the default; ``sigma_d`` reports its clamped-shell count and falls back to inverse variance with a warning when every shell is empty diff --git a/tests/unit/model/test_hydrogen_default.py b/tests/unit/model/test_hydrogen_default.py index 5203331c..de2ad847 100644 --- a/tests/unit/model/test_hydrogen_default.py +++ b/tests/unit/model/test_hydrogen_default.py @@ -258,10 +258,13 @@ def test_derived_models_keep_the_restraint_cif(pdb_dir, renamed_glu_cif): @pytest.mark.unit -def test_state_dict_round_trips_the_restraint_cif(pdb_dir, renamed_glu_cif): - model = Model(verbose=0, cif_path=renamed_glu_cif) +@pytest.mark.parametrize("model_class", [Model, ModelFT]) +def test_state_dict_round_trips_the_restraint_cif( + pdb_dir, renamed_glu_cif, model_class +): + model = model_class(verbose=0, cif_path=renamed_glu_cif) model.load_pdb(str(pdb_dir / "1DAW.pdb")) - restored = Model.create_from_state_dict(model.state_dict(), verbose=0) + restored = model_class.create_from_state_dict(model.state_dict(), verbose=0) assert restored.ctx.cif_path == renamed_glu_cif diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index 0c3299a1..0e83cfd8 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -920,6 +920,7 @@ def create_from_state_dict( saved_dtype = state_dict.pop("dtype_float", dtype_float) state_dict.pop("device", None) # Remove but don't use (use provided device) strip_H = state_dict.pop("strip_H", True) + cif_path = state_dict.pop("cif_path", None) altloc_pairs = state_dict.pop("altloc_pairs", []) hydrogens_in_xray = state_dict.pop("hydrogens_in_xray", True) hydrogen_mode = state_dict.pop("hydrogen_mode", None) @@ -938,6 +939,7 @@ def create_from_state_dict( verbose=verbose, device=device, strip_H=strip_H, + cif_path=cif_path, max_res=max_res, gridsize=explicit_gridsize, wavelength=wavelength, diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index 96ee5799..2547d939 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -1218,9 +1218,8 @@ def create_from_state_dict( The recommended restore path: it rebuilds reflection data, model and scaler through their own factories before calling ``load_state_dict``, which - :meth:`load_state` cannot do. Restraints are normally lazy via - ``model.restraints``; the standalone handling here is a legacy state-dict - path and does not make them a first-class persisted submodule. + :meth:`load_state` cannot do. Restraints are not persisted; the + restored model rebuilds them on first access to ``model.restraints``. Parameters ---------- @@ -1253,13 +1252,12 @@ def extract_submodule_state(state_dict: dict, prefix: str) -> dict: model_state = extract_submodule_state(state_dict, "model") reflection_data_state = extract_submodule_state(state_dict, "reflection_data") scaler_state = extract_submodule_state(state_dict, "scaler") - restraints_state = extract_submodule_state(state_dict, "restraints") weighter_state = extract_submodule_state(state_dict, "weighter") if verbose > 0: print( f"Extracted state dict sizes: model={len(model_state)}, data={len(reflection_data_state)}, " - f"scaler={len(scaler_state)}, restraints={len(restraints_state)}" + f"scaler={len(scaler_state)}" ) # Create submodules using their factory methods @@ -1276,11 +1274,6 @@ def extract_submodule_state(state_dict: dict, prefix: str) -> dict: # Create Scaler with model and data (required for proper setup) scaler = Scaler(model, reflection_data, verbose=verbose, device=device) - # Create Restraints with model (required for proper setup) - from torchref.topology.restraints import Restraints - - restraints = Restraints(model, verbose=verbose) - # Create empty instance instance = cls.__new__(cls) nnModule.__init__(instance) @@ -1301,7 +1294,6 @@ def extract_submodule_state(state_dict: dict, prefix: str) -> dict: instance.reflection_data = reflection_data instance.model = model instance.scaler = scaler - instance.restraints = restraints instance.weighter = None # Now load the state dict - PyTorch's default will fill in values From 67e774a818675cc8c0ad8ad67e0c6ac2e29c26ad Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 18:10:11 +0200 Subject: [PATCH 196/250] Keep restraints on the model context and pass coordinates in Restraints no longer borrow xyz/adp/vdw-radius accessors from the model. The constructor takes the coordinates to build over, every evaluation takes the coordinates or B-factors it scores, and the pair-list rebuild takes the current coordinates from the non-bonded target. They live on ModelContext (ctx.restraints, ctx.build_restraints, ctx.set_cif_path), so Model.copy and device moves carry them with the context and _repoint_coordinate_accessors no longer has to patch them. The no-symmetry pair search is replaced by the periodic search in an isolated P1 box, which removes the older spatial-hash builder and its symmetry-expansion helper. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 1 + docs/user_guide/restraints.rst | 24 +- tests/fixtures/objects.py | 4 +- .../functional/test_restraints_functional.py | 36 +- tests/functional/test_targets_functional.py | 10 +- tests/integration/test_refinement_pipeline.py | 3 +- tests/unit/model/test_hydrogen_mode.py | 12 +- .../model/test_riding_water_completion.py | 2 +- tests/unit/topology/test_equivalence.py | 4 +- tests/unit/topology/test_hydrogens.py | 6 +- tests/unit/topology/test_insertion_codes.py | 2 +- tests/unit/topology/test_links.py | 2 +- tests/unit/topology/test_storage.py | 29 +- tests/unit/topology/test_subset.py | 2 +- torchref/cli/collection_difference_refine.py | 6 +- torchref/experimental/kinetic/refinement.py | 3 +- torchref/io/metadata.py | 6 +- torchref/model/context.py | 72 +- torchref/model/model.py | 174 +--- torchref/refinement/base_refinement.py | 5 +- torchref/refinement/targets/adp/similarity.py | 2 +- .../refinement/targets/geometry/angles.py | 2 +- torchref/refinement/targets/geometry/bonds.py | 2 +- .../refinement/targets/geometry/non_bonded.py | 2 +- .../refinement/targets/geometry/torsions.py | 10 +- torchref/topology/nonbonded.py | 56 +- torchref/topology/restraints.py | 774 +++--------------- 27 files changed, 336 insertions(+), 915 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 8ce89eec..504adc3c 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- ``Restraints`` no longer borrow accessors from the model and live on its context (``model.ctx.restraints``; ``model.restraints`` still builds them on first access). The constructor takes the coordinates to build over (``xyz=``) instead of ``xyz_fn``/``adp_fn``/``vdw_radii_fn``, and every evaluation takes the coordinates or B-factors it scores, e.g. ``bond_deviations(model.xyz())``. ``Model.set_restraints_cif`` is replaced by ``model.ctx.set_cif_path``, and ``Model.bond_deviations``/``angle_deviations``/``torsion_deviations_with_sigmas`` are removed. Without a cell and space group the pair list is searched in an isolated P1 box - ``ModelFT.create_from_state_dict`` restores the restraint dictionary path (``cif_path``) as ``Model`` does, and ``Refinement.create_from_state_dict`` no longer builds a stray ``Restraints`` from the model; the restored model builds its own on first access. - The difference MTZ groups its columns into named datasets -- ``observed``, ``difference``, ``light_model``, ``extrapolated_light``, ``two_moment`` -- with one history line describing each, so ``FWT``/``PHWT`` reads as ``/torchref/extrapolated_light/FWT`` (the extrapolated light-state map ``2*FEXT - Fc``). Labels are unchanged and Coot still auto-opens it - ``torchref.difference-map``, ``torchref.difference-refine`` and ``torchref.validate-ded`` gain ``--ded-weight {sigma_d,inverse_variance,none}`` and ``--sigma-d-gamma``. The difference MTZ now carries the unweighted ``DF``/``SIGDF`` on ``PHDELWT`` with one mean-one weight column per scheme, ``W_SD`` and ``W_IVW`` (MTZ type W), and the observed-to-model scale ``KSCALE``; ``DELFWT`` is no longer written, build the map with ``torchref.mtz2map -csf DF -cw W_IVW -cphi PHDELWT``. Registered in ``torchref.maps.ded_weights`` diff --git a/docs/user_guide/restraints.rst b/docs/user_guide/restraints.rst index 96e149b0..de5128e8 100644 --- a/docs/user_guide/restraints.rst +++ b/docs/user_guide/restraints.rst @@ -2,26 +2,28 @@ Geometry Restraints =================== Geometry restraints keep the model chemically reasonable during refinement. -:class:`~torchref.restraints.Restraints` (the exported alias of -``RestraintsNew``) builds and holds bond, angle, torsion, planarity, chirality, -and non-bonded (VDW) restraints. +:class:`~torchref.topology.Restraints` builds and holds bond, angle, torsion, +planarity, chirality, and non-bonded (VDW) restraints. Restraint Setup --------------- -You do not normally construct ``Restraints`` yourself. ``model.restraints`` is a -lazy property that builds them from the monomer library — fetched per monomer on -demand — on first access. Point it at extra CIF definitions *before* that first -access: +You do not normally construct ``Restraints`` yourself. They live on the model's +context, ``model.ctx.restraints``, and ``model.restraints`` builds them from the +monomer library -- fetched per monomer on demand -- on first access. Point it at +extra CIF definitions *before* that first access: .. code-block:: python - model.set_restraints_cif("ligand.cif") # or a list of paths; chainable + model.ctx.set_cif_path("ligand.cif") # or a list of paths restraints = model.restraints # built here, on first access + deviations, sigmas = restraints.bond_deviations(model.xyz()) -``Restraints.__init__`` takes a PDB DataFrame plus accessor callables -(``pdb, cif_path, xyz_fn, adp_fn, vdw_radii_fn, cell, spacegroup, links, -verbose``), not a model — that is what the lazy property assembles for you. +Restraints hold no reference to the model: every evaluation takes the +coordinates (or B-factors) it scores, and the non-bonded pair list is rebuilt +from the coordinates the non-bonded target passes in. ``Restraints.__init__`` +takes an atom table and the coordinates to build over (``pdb, cif_path, xyz, +cell, spacegroup, links, verbose, nonbonded``). Residues for which no restraints could be built are frozen in ``xyz`` rather than refined unrestrained, so a missing ligand definition shows up as an diff --git a/tests/fixtures/objects.py b/tests/fixtures/objects.py index 3eb7628e..9bdedc87 100644 --- a/tests/fixtures/objects.py +++ b/tests/fixtures/objects.py @@ -114,11 +114,9 @@ def model_with_restraints(loaded_model: Model) -> dict[str, Any]: restraints = Restraints( pdb=loaded_model.pdb, - xyz_fn=loaded_model.xyz, - vdw_radii_fn=loaded_model.get_vdw_radii, + xyz=loaded_model.xyz(), verbose=0, ) - restraints.build_restraints() return {"model": loaded_model, "restraints": restraints} diff --git a/tests/functional/test_restraints_functional.py b/tests/functional/test_restraints_functional.py index 841022cc..146e0fcc 100644 --- a/tests/functional/test_restraints_functional.py +++ b/tests/functional/test_restraints_functional.py @@ -21,9 +21,8 @@ def test_build_restraints_from_cif(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 + pdb=model.pdb, xyz=model.xyz(), verbose=0 ) - restraints.build_restraints() # Should have built some restraints assert restraints.restraints is not None @@ -39,9 +38,8 @@ def test_bond_restraints_built(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 + pdb=model.pdb, xyz=model.xyz(), verbose=0 ) - restraints.build_restraints() # Check bond restraints exist assert "bond" in restraints.restraints @@ -77,9 +75,8 @@ def test_angle_restraints_built(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 + pdb=model.pdb, xyz=model.xyz(), verbose=0 ) - restraints.build_restraints() # Check angle restraints exist assert "angle" in restraints.restraints @@ -108,9 +105,8 @@ def test_torsion_restraints_built(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 + pdb=model.pdb, xyz=model.xyz(), verbose=0 ) - restraints.build_restraints() # Check torsion restraints exist assert "torsion" in restraints.restraints @@ -137,9 +133,8 @@ def test_plane_restraints_built(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 + pdb=model.pdb, xyz=model.xyz(), verbose=0 ) - restraints.build_restraints() # Check plane restraints exist assert "plane" in restraints.restraints @@ -169,13 +164,12 @@ def test_bond_deviations(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 + pdb=model.pdb, xyz=model.xyz(), verbose=0 ) - restraints.build_restraints() # Compute bond deviations if hasattr(restraints, "bond_deviations"): - deviations, sigmas = restraints.bond_deviations() + deviations, sigmas = restraints.bond_deviations(model.xyz()) assert torch.all(torch.isfinite(deviations)) assert torch.all(sigmas > 0) @@ -193,13 +187,12 @@ def test_angle_deviations(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 + pdb=model.pdb, xyz=model.xyz(), verbose=0 ) - restraints.build_restraints() # Compute angle deviations if hasattr(restraints, "angle_deviations"): - deviations, sigmas = restraints.angle_deviations() + deviations, sigmas = restraints.angle_deviations(model.xyz()) assert torch.all(torch.isfinite(deviations)) assert torch.all(sigmas > 0) @@ -217,11 +210,9 @@ def test_restraints_multiple_cif_files(self, compatibility_model): model = compatibility_model restraints = Restraints( pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, + xyz=model.xyz(), verbose=0, ) - restraints.build_restraints() assert "bond" in restraints.restraints assert "angle" in restraints.restraints @@ -239,9 +230,8 @@ def test_restraints_device_movement(self, sample_cif_file, cpu_device): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 + pdb=model.pdb, xyz=model.xyz(), verbose=0 ) - restraints.build_restraints() # Check that tensors are on the correct device if "bond" in restraints.restraints and "intra" in restraints.restraints["bond"]: @@ -262,7 +252,7 @@ def test_cif_dict_loaded(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 + pdb=model.pdb, xyz=model.xyz(), verbose=0 ) # CIF dict should be populated with residue restraints @@ -288,7 +278,7 @@ def test_unique_residues_detected(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 + pdb=model.pdb, xyz=model.xyz(), verbose=0 ) # Should have detected unique residues diff --git a/tests/functional/test_targets_functional.py b/tests/functional/test_targets_functional.py index c3f2f70d..b3c1d00d 100644 --- a/tests/functional/test_targets_functional.py +++ b/tests/functional/test_targets_functional.py @@ -154,7 +154,7 @@ def test_bond_target_with_real_structure(self, sample_cif_file, external_monomer model.load_cif(str(sample_cif_file)) # Use new model-based restraints API - model.set_restraints_cif(str(external_monomer_library)) + model.ctx.set_cif_path(str(external_monomer_library)) restraints = model.restraints # Calculate bond deviations manually @@ -191,7 +191,7 @@ def test_angle_target_with_real_structure(self, sample_cif_file, external_monome model.load_cif(str(sample_cif_file)) # Use new model-based restraints API - model.set_restraints_cif(str(external_monomer_library)) + model.ctx.set_cif_path(str(external_monomer_library)) restraints = model.restraints # Calculate angle deviations @@ -613,7 +613,7 @@ def test_bond_deviation_calculation(self, sample_cif_file, external_monomer_libr model.load_cif(str(sample_cif_file)) # Use new model-based restraints API - model.set_restraints_cif(str(external_monomer_library)) + model.ctx.set_cif_path(str(external_monomer_library)) restraints = model.restraints if 'bond' in restraints.restraints and 'intra' in restraints.restraints['bond']: @@ -644,7 +644,7 @@ def test_angle_deviation_calculation(self, sample_cif_file, external_monomer_lib model.load_cif(str(sample_cif_file)) # Use new model-based restraints API - model.set_restraints_cif(str(external_monomer_library)) + model.ctx.set_cif_path(str(external_monomer_library)) restraints = model.restraints if 'angle' in restraints.restraints and 'intra' in restraints.restraints['angle']: @@ -695,7 +695,7 @@ def test_xray_plus_geometry_loss(self, sample_structure_pair, external_monomer_l data.load_mtz(str(sample_structure_pair["reflections"])) # Use new model-based restraints API - model.set_restraints_cif(str(external_monomer_library)) + model.ctx.set_cif_path(str(external_monomer_library)) restraints = model.restraints # X-ray loss diff --git a/tests/integration/test_refinement_pipeline.py b/tests/integration/test_refinement_pipeline.py index 94cb9258..bf193d54 100644 --- a/tests/integration/test_refinement_pipeline.py +++ b/tests/integration/test_refinement_pipeline.py @@ -50,8 +50,7 @@ def test_restraints_from_model(self, sample_cif_file): model.load_cif(str(sample_cif_file)) # Build restraints - restraints = Restraints(pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii) - restraints.build_restraints() + restraints = Restraints(pdb=model.pdb, xyz=model.xyz()) # Should have some restraints assert restraints.restraints is not None diff --git a/tests/unit/model/test_hydrogen_mode.py b/tests/unit/model/test_hydrogen_mode.py index c7676a5a..52710726 100644 --- a/tests/unit/model/test_hydrogen_mode.py +++ b/tests/unit/model/test_hydrogen_mode.py @@ -49,14 +49,16 @@ def test_switch_back_to_free_restores_per_atom_wrapper(free_model): @pytest.mark.unit -def test_restraints_read_the_installed_wrapper(free_model): +def test_restraints_survive_the_switch(free_model): + """The atom table is unchanged, so the restraints are too; they score whatever + coordinates they are handed, including the riding wrapper's.""" model = free_model restraints = model.restraints model.set_hydrogen_mode("riding") - assert restraints._xyz_fn is model.xyz - with torch.no_grad(): - model.xyz.refinable_params.add_(0.1) - assert torch.equal(restraints.xyz(), model.xyz()) + assert model.restraints is restraints + deviations, _ = restraints.bond_deviations(model.xyz()) + deviations.sum().backward() + assert model.xyz.refinable_params.grad is not None @pytest.mark.unit diff --git a/tests/unit/model/test_riding_water_completion.py b/tests/unit/model/test_riding_water_completion.py index 418b75af..7b89d7c5 100644 --- a/tests/unit/model/test_riding_water_completion.py +++ b/tests/unit/model/test_riding_water_completion.py @@ -52,7 +52,7 @@ def test_riding_completes_only_waters_and_preserves_live_atoms( assert not model.occupancy.get_refinable_atoms().any() assert model.ctx.links is links assert model.xyz.rotations.shape[0] == int(is_h.sum()) // 2 - assert model.restraints.xyz().shape == model.xyz.shape + assert model.restraints.topology.n_atoms == model.xyz.shape[0] if isinstance(model, ModelFT): assert torch.isfinite(model(hkl)).all() restored = model_class.create_from_state_dict(model.state_dict(), device="cpu") diff --git a/tests/unit/topology/test_equivalence.py b/tests/unit/topology/test_equivalence.py index 0e276804..0f67843c 100644 --- a/tests/unit/topology/test_equivalence.py +++ b/tests/unit/topology/test_equivalence.py @@ -68,7 +68,7 @@ def _build(code): pytest.skip(f"{code}.pdb not bundled") model = Model(verbose=0) model.load_pdb(str(path)) - model.set_restraints_cif(None) + model.ctx.set_cif_path(None) restraints = model.restraints topology = build_topology( model.pdb, @@ -200,7 +200,7 @@ def test_layout_is_reproducible(built, pdb_dir): model = Model(verbose=0) model.load_pdb(str(pdb_dir / "7L84.pdb")) - model.set_restraints_cif(None) + model.ctx.set_cif_path(None) restraints = model.restraints topology_b = build_topology( model.pdb, diff --git a/tests/unit/topology/test_hydrogens.py b/tests/unit/topology/test_hydrogens.py index bcd2620d..59bbe307 100644 --- a/tests/unit/topology/test_hydrogens.py +++ b/tests/unit/topology/test_hydrogens.py @@ -34,7 +34,7 @@ def _build(code): # has to arrive without the hydrogens the loader would otherwise add. model = Model(verbose=0, add_hydrogens=False, strip_H=True) model.load_pdb(str(pdb_dir / f"{code}.pdb")) - model.set_restraints_cif(None) + model.ctx.set_cif_path(None) restraints = model.restraints plan = plan_hydrogens( restraints.topology, restraints.cif_dict, model.xyz().detach() @@ -276,7 +276,7 @@ def test_hydrogenate_returns_a_consistent_model(pdb_dir): """The end-to-end path yields a model whose tensors, table and restraints agree.""" model = Model(verbose=0, add_hydrogens=False, strip_H=True) model.load_pdb(str(pdb_dir / "7L84.pdb")) - model.set_restraints_cif(None) + model.ctx.set_cif_path(None) n_heavy = len(model.pdb) hydrogenated = model.hydrogenate(verbose=0) @@ -387,7 +387,7 @@ def test_acetyl_cap_carbon_gets_no_hydrogen(tmp_path): path.write_text(_ACE_MET) model = Model(verbose=0, add_hydrogens=False, strip_H=True) model.load_pdb(str(path)) - model.set_restraints_cif(None) + model.ctx.set_cif_path(None) restraints = model.restraints plan = plan_hydrogens( restraints.topology, restraints.cif_dict, model.xyz().detach() diff --git a/tests/unit/topology/test_insertion_codes.py b/tests/unit/topology/test_insertion_codes.py index 0e5aec1b..6dafac8f 100644 --- a/tests/unit/topology/test_insertion_codes.py +++ b/tests/unit/topology/test_insertion_codes.py @@ -62,7 +62,7 @@ def inserted(pdb_dir, tmp_path_factory): model = Model(verbose=0, strip_H=True, add_hydrogens=False) model.load_pdb(str(path)) - model.set_restraints_cif(None) + model.ctx.set_cif_path(None) restraints = model.restraints topology = build_topology( diff --git a/tests/unit/topology/test_links.py b/tests/unit/topology/test_links.py index 0a9c2aec..f1ee1da6 100644 --- a/tests/unit/topology/test_links.py +++ b/tests/unit/topology/test_links.py @@ -52,7 +52,7 @@ def _row(model, chain, resseq, name, altloc=""): def _load(path): model = Model(verbose=0, add_hydrogens=False, strip_H=True) model.load_pdb(str(path)) if str(path).endswith(".pdb") else model.load_cif(str(path)) - model.set_restraints_cif(None) + model.ctx.set_cif_path(None) return model diff --git a/tests/unit/topology/test_storage.py b/tests/unit/topology/test_storage.py index d3423297..56f89745 100644 --- a/tests/unit/topology/test_storage.py +++ b/tests/unit/topology/test_storage.py @@ -10,6 +10,7 @@ import pytest import torch +from torchref.config import get_float_dtype from torchref.model.model import Model from torchref.utils.caching import ParameterFingerprint @@ -40,7 +41,7 @@ def restraints(pdb_dir): """Restraints for a structure with altlocs, disulfides and peptide links.""" model = Model(verbose=0) model.load_pdb(str(pdb_dir / "7L84.pdb")) - model.set_restraints_cif(None) + model.ctx.set_cif_path(None) return model.restraints @@ -171,7 +172,14 @@ def test_blocks_are_untouched_by_a_refinement_step(restraints): blocks = [restraints.topology.edge_block(t).indices for t in KEYED_TYPES] fingerprint = ParameterFingerprint(blocks) - loss = restraints.nll_bonds().sum() + restraints.nll_angles().sum() + block = restraints.topology.atoms.bonds.indices + xyz = torch.tensor( + restraints.pdb[["x", "y", "z"]].values, + dtype=get_float_dtype(), + device=block.device, + requires_grad=True, + ) + loss = restraints.nll_bonds(xyz).sum() + restraints.nll_angles(xyz).sum() loss.backward() assert fingerprint.matches( @@ -185,13 +193,6 @@ def test_rebuilding_entries_reslices_onto_the_current_blocks(restraints): This is the operation ``_apply`` and ``copy`` both rely on, and the one that has to stay cheap: it re-slices rather than recomputing anything. - - ``Restraints.copy`` is not exercised here because it cannot run at all -- it is - ``deepcopy``, which walks the *borrowed* ``_xyz_fn`` wrapper, whose cache holds a - graph-attached tensor once ``xyz()`` has been evaluated. Verified to fail - identically at the commit before this change, so it is pre-existing rather than a - regression, and it is reached only through ``Model.copy`` on a model whose lazy - restraints have already been built. """ block = restraints.topology.atoms.bonds.indices before = restraints.restraints["bond"]["all"]["indices"].clone() @@ -205,3 +206,13 @@ def test_rebuilding_entries_reslices_onto_the_current_blocks(restraints): entry = restraints.restraints["bond"][origin]["indices"] assert entry.shape[0] == bounds[1] - bounds[0] assert restraints.restraints["vdw"].get("indices") is not None + + +@pytest.mark.unit +def test_copy_aliases_its_own_blocks(restraints): + """A copy re-slices its entries onto its own blocks, not the original's.""" + duplicate = restraints.copy() + block = duplicate.topology.atoms.bonds.indices + entry = duplicate.restraints["bond"]["all"]["indices"] + assert entry.data_ptr() == block.data_ptr() + assert block.data_ptr() != restraints.topology.atoms.bonds.indices.data_ptr() diff --git a/tests/unit/topology/test_subset.py b/tests/unit/topology/test_subset.py index 60b36126..001df453 100644 --- a/tests/unit/topology/test_subset.py +++ b/tests/unit/topology/test_subset.py @@ -22,7 +22,7 @@ def topology(pdb_dir): """A topology with altlocs, disulfides, peptide links and hydrogens.""" model = Model(verbose=0, add_hydrogens=False, strip_H=True) model.load_pdb(str(pdb_dir / "7L84.pdb")) - model.set_restraints_cif(None) + model.ctx.set_cif_path(None) return model.restraints.topology diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index a6b36909..dade6c23 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -1639,14 +1639,14 @@ def _build_metadata(model, data, r_work, r_free): meta.n_atoms_solvent = int((pdb["ATOM"] == "HETATM").sum()) # Geometry deviations - if model.ctx.initialized and model._restraints is not None: + if model.ctx.initialized and model.ctx.restraints is not None: restraints = model.restraints with torch.no_grad(): if hasattr(restraints, "bond_deviations"): - bond_devs, _ = restraints.bond_deviations() + bond_devs, _ = restraints.bond_deviations(model.xyz()) meta.rmsd_bond_lengths = float(torch.sqrt((bond_devs**2).mean())) if hasattr(restraints, "angle_deviations"): - angle_devs, _ = restraints.angle_deviations() + angle_devs, _ = restraints.angle_deviations(model.xyz()) meta.rmsd_bond_angles = float(torch.sqrt((angle_devs**2).mean())) # Solvent model from CollectionScaler diff --git a/torchref/experimental/kinetic/refinement.py b/torchref/experimental/kinetic/refinement.py index 2aebd182..286bba05 100644 --- a/torchref/experimental/kinetic/refinement.py +++ b/torchref/experimental/kinetic/refinement.py @@ -154,8 +154,7 @@ def setup( # ---- CIF restraints on base models ---- if cif_paths: for model in mc.base_models: - if hasattr(model, "set_restraints_cif"): - model.set_restraints_cif(cif_paths) + model.ctx.set_cif_path(cif_paths) # ---- Scalers ---- self._setup_scalers() diff --git a/torchref/io/metadata.py b/torchref/io/metadata.py index e9b0dc3d..e7ba8760 100644 --- a/torchref/io/metadata.py +++ b/torchref/io/metadata.py @@ -263,17 +263,17 @@ def from_refinement(cls, refinement) -> RefinementMetadata: # --- Geometry deviations (silently skip if no restraints) --- try: model = refinement.model - if model.ctx.initialized and model._restraints is not None: + if model.ctx.initialized and model.ctx.restraints is not None: restraints = model.restraints if hasattr(restraints, "bond_deviations"): with torch.no_grad(): - bond_devs, _ = restraints.bond_deviations() + bond_devs, _ = restraints.bond_deviations(model.xyz()) meta.rmsd_bond_lengths = float( torch.sqrt((bond_devs**2).mean()) ) if hasattr(restraints, "angle_deviations"): with torch.no_grad(): - angle_devs, _ = restraints.angle_deviations() + angle_devs, _ = restraints.angle_deviations(model.xyz()) meta.rmsd_bond_angles = float( torch.sqrt((angle_devs**2).mean()) ) diff --git a/torchref/model/context.py b/torchref/model/context.py index 1b9fc66e..5b0da878 100644 --- a/torchref/model/context.py +++ b/torchref/model/context.py @@ -3,7 +3,8 @@ :class:`ModelContext` holds what a model *is loaded from* and *sits in* -- the unit cell, the space group, the atom table, the link records and the provenance -- as opposed to what is being refined, which stays on the model as parameter wrappers and -per-atom buffers. +per-atom buffers. The geometry restraints belong here too: they are fixed by the atom +set and the dictionaries, and are evaluated against coordinates the caller passes in. Splitting it out means the crystallographic context can be passed to code that needs only that (structure-factor engines, scalers, most targets) without handing over the @@ -21,8 +22,10 @@ if TYPE_CHECKING: import pandas + import torch from torchref.symmetry import Cell, SpaceGroup + from torchref.topology.restraints import Restraints @dataclass(eq=False, repr=False) @@ -63,6 +66,10 @@ class ModelContext(DeviceMixin): atoms) or ``"none"`` (the table holds no hydrogens). initialized : bool, default False Whether a structure has been loaded. ``if model:`` tests this. + restraints : Restraints or None + Geometry restraints over ``pdb``, or None until :meth:`build_restraints` runs. + Reset to None whenever the atom table or ``cif_path`` changes; read them + through ``Model.restraints``, which builds on first access. Notes ----- @@ -88,12 +95,60 @@ class ModelContext(DeviceMixin): add_hydrogens: bool = False hydrogen_mode: str = "free" initialized: bool = False + restraints: Optional["Restraints"] = None + + def set_cif_path(self, cif_path) -> None: + """Replace the restraint dictionary path and drop restraints built over the old one. + + Parameters + ---------- + cif_path : str or list of str or None + Restraint dictionary file(s). + """ + self.cif_path = cif_path + self.restraints = None + + def build_restraints( + self, xyz: "torch.Tensor", *, nonbonded: bool = True, verbose=None + ) -> "Restraints": + """Build restraints over the atom table and store them on :attr:`restraints`. + + Parameters + ---------- + xyz : torch.Tensor + Current Cartesian coordinates in Å, shape ``(n_atoms, 3)``; the atom table's + own columns are stale during refinement. The restraints land on its device. + nonbonded : bool, default True + Build the non-bonded pair list. False is for a throwaway build that needs + only the topology, and is then **not** stored. + verbose : int, optional + Defaults to :attr:`verbose`. + + Returns + ------- + Restraints + """ + from torchref.topology.restraints import Restraints + + restraints = Restraints( + pdb=self.pdb, + cif_path=self.cif_path, + xyz=xyz.detach(), + cell=self.cell, + spacegroup=self.spacegroup, + links=self.links, + verbose=self.verbose if verbose is None else verbose, + nonbonded=nonbonded, + ) + if nonbonded: + self.restraints = restraints + return restraints def copy(self) -> "ModelContext": """An independent copy. - The atom table is deep-copied and the cell and space group are cloned, so - nothing is shared with the original. Cloning the space group matters now that + The atom table is deep-copied, the cell and space group are cloned and built + restraints are copied, so nothing is shared with the original. Cloning the space group matters now that it is a mutable dataclass: sharing the reference would let an edit through one model's context reach every model that was copied from it. @@ -102,7 +157,7 @@ def copy(self) -> "ModelContext": ModelContext New context sharing no mutable state with this one. """ - return ModelContext( + duplicate = ModelContext( cell=self.cell.clone() if self.cell is not None else None, spacegroup=( self.spacegroup.copy() if self.spacegroup is not None else None @@ -121,6 +176,15 @@ def copy(self) -> "ModelContext": hydrogen_mode=self.hydrogen_mode, initialized=self.initialized, ) + if self.restraints is not None: + restraints = self.restraints.copy() + # Point at the copied table and crystal rather than the deep-copied + # duplicates, so the new context is the single owner of both. + restraints.pdb = duplicate.pdb + restraints._cell = duplicate.cell + restraints._spacegroup = duplicate.spacegroup + duplicate.restraints = restraints + return duplicate @property def crystal_key(self): diff --git a/torchref/model/model.py b/torchref/model/model.py index 8380b2f6..500fcd35 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -161,9 +161,9 @@ def __init__( Generate missing hydrogens on load when True. Default False; ignored when ``strip_H`` is set. cif_path : str or list of str, optional - Restraint dictionary file(s); see the class docstring. :meth:`set_restraints_cif` - can still change it after loading, but generation on load only sees the value - given here. + Restraint dictionary file(s); see the class docstring. + :meth:`ModelContext.set_cif_path` can still change it after loading, but + generation on load only sees the value given here. hydrogens_in_xray : bool, optional Whether hydrogens contribute to the structure factors. Default True. They stay in the restraints either way; see :attr:`hydrogens_in_xray`. @@ -199,9 +199,6 @@ def __init__( # Scattering factor parametrization (built lazily on first access) self._parametrization = None - # Restraints (built lazily on first access) - self._restraints = None - def __bool__(self): """Return the initialization status when used in boolean context. @@ -540,113 +537,24 @@ def get_scattering_params_aniso(self): # Restraints (Geometry Restraints) # ========================================================================= - def set_restraints_cif(self, cif_path): - """ - Set CIF path for lazy restraint building. - - Parameters - ---------- - cif_path : str or list of str - Path(s) to CIF restraints dictionary file(s). - - Returns - ------- - Model - Self, for method chaining. - """ - self.ctx.cif_path = cif_path - # Reset restraints so they will be rebuilt on next access - self._restraints = None - return self - - def _build_restraints(self): - """Build and cache ``Restraints`` over this model's DataFrame, wiring in - the live ``xyz`` / ``adp`` / ``vdw_radii`` callables. - """ - if self._restraints is not None: - return self._restraints - - if not self.ctx.initialized: - raise RuntimeError( - "Cannot build restraints: model not initialized. " - "Load data first with load_pdb() or load_cif()." - ) - - if self.ctx.verbose > 0: - print("Building restraints...") - - self._restraints = self._new_restraints() - - return self._restraints - - def _new_restraints(self, nonbonded: bool = True, verbose: Optional[int] = None): - """An uncached ``Restraints`` over this model's DataFrame, wired to the live - ``xyz`` / ``adp`` / ``vdw_radii`` callables; see :meth:`_build_restraints`. - """ - from torchref.topology.restraints import Restraints - - return Restraints( - pdb=self.pdb, - cif_path=self.ctx.cif_path, - xyz_fn=self.xyz, - adp_fn=self.adp, - vdw_radii_fn=self.get_vdw_radii, - cell=self.ctx.cell, - spacegroup=self.ctx.spacegroup, - links=self.ctx.links, - verbose=self.ctx.verbose if verbose is None else verbose, - nonbonded=nonbonded, - ) - @property def restraints(self): - """Bond/angle/torsion/... restraints, built on first access from the - DataFrame and the CIF path given to :meth:`set_restraints_cif`. - """ - return self._build_restraints() + """Geometry restraints over the atom table, on :attr:`ctx`. - # ========================================================================= - # Restraint Evaluation Wrappers - # ========================================================================= - - def bond_deviations(self): - """ - Compute bond length deviations using current xyz coordinates. - - Returns - ------- - deviations : torch.Tensor - Calculated minus expected bond lengths in Angstroms. - sigmas : torch.Tensor - Standard deviations from CIF library in Angstroms. - """ - return self.restraints.bond_deviations(self.xyz()) - - def angle_deviations(self): - """ - Compute angle deviations using current xyz coordinates. - - Returns - ------- - deviations : torch.Tensor - Calculated minus expected angles in radians. - sigmas : torch.Tensor - Standard deviations in radians. - """ - return self.restraints.angle_deviations(self.xyz()) - - def torsion_deviations_with_sigmas(self): + Built on first access over the current coordinates and cached on the context + until the atom table or ``ctx.cif_path`` changes. Evaluations take the + coordinates as an argument, e.g. ``model.restraints.bond_deviations(model.xyz())``. """ - Compute torsion deviations (wrapped for periodicity) and sigmas. - - Returns - ------- - deviations_rad : torch.Tensor - Wrapped deviations in radians. - sigmas_deg : torch.Tensor - Standard deviations in degrees (for von Mises NLL). - """ - return self.restraints.torsion_deviations_with_sigmas(self.xyz()) + if self.ctx.restraints is None: + if not self.ctx.initialized: + raise RuntimeError( + "Cannot build restraints: model not initialized. " + "Load data first with load_pdb() or load_cif()." + ) + if self.ctx.verbose > 0: + print("Building restraints...") + self.ctx.build_restraints(self.xyz()) + return self.ctx.restraints #: Per-atom buffers built lazily on first use and cached. Each is sized to the atom #: table, so all of them go stale the moment the atom set changes. @@ -672,6 +580,7 @@ def _invalidate_atom_derived_caches(self) -> None: if hasattr(self, name): delattr(self, name) self._parametrization = None + self.ctx.restraints = None def load(self, reader, add_hydrogens: bool = None): """ @@ -808,10 +717,10 @@ def _add_missing_hydrogens(self) -> None: plan_hydrogens, ) - restraints = self._restraints - if restraints is None: - restraints = self._new_restraints(nonbonded=False, verbose=0) xyz = self.xyz().detach() + restraints = self.ctx.restraints + if restraints is None: + restraints = self.ctx.build_restraints(xyz, nonbonded=False, verbose=0) plan = plan_hydrogens( restraints.topology, restraints.cif_dict, xyz, verbose=self.ctx.verbose ) @@ -829,8 +738,6 @@ def _add_missing_hydrogens(self) -> None: if self.ctx.verbose > 0: print(f"Generated {plan.n_hydrogens} hydrogens") - # The topology and every per-atom tensor are sized for the old atom set. - self._restraints = None cell, spacegroup = self.cell, self.spacegroup links = self.ctx.links @@ -1103,35 +1010,11 @@ def get_vdw_radii(self): torch.Tensor Van der Waals radii for each atom with shape (n_atoms,). """ - import os - - import pandas as pd - - from torchref import PATH_TORCHREF_DATA + from torchref.topology.nonbonded import vdw_radii_for_elements if hasattr(self, "vdw_radii"): return self.vdw_radii - elements = self.pdb.loc[:, "element"] - path = os.path.join( - PATH_TORCHREF_DATA, - "atomic_vdw_radii.csv", - ) - vdw_df = pd.read_csv(path, comment="#") - vdw_df["element"] = vdw_df["element"].str.strip().str.capitalize() - elements = elements.str.strip().str.capitalize() - elements_not_in = elements[~elements.isin(vdw_df["element"])] - if len(elements_not_in) > 0: - # Add missing elements with default vdW radius 1.9 Å - missing = sorted(set(e.strip().capitalize() for e in elements_not_in)) - if missing: - add_df = pd.DataFrame( - {"element": missing, "vdW_Radius_Angstrom": [1.9] * len(missing)} - ) - vdw_df = pd.concat([vdw_df, add_df], ignore_index=True) - - vdw_radii = ( - vdw_df.set_index("element").loc[elements]["vdW_Radius_Angstrom"].values - ) + vdw_radii = vdw_radii_for_elements(self.pdb["element"]) self.register_buffer( "vdw_radii", torch.tensor(vdw_radii, dtype=self.dtype_float, device=self.device), @@ -2952,15 +2835,9 @@ def hydrogen_frames(self): def _repoint_coordinate_accessors(self) -> None: """Make every borrowed coordinate accessor read the current ``xyz`` wrapper. - The restraints keep ``xyz_fn`` for pair-list maintenance and the ADP node - field borrows the coordinates through ``set_xyz_fn``; after the wrapper slot - is replaced both would otherwise keep reading a dead module. + The ADP node field borrows the coordinates through ``set_xyz_fn``; after the + wrapper slot is replaced it would otherwise keep reading a dead module. """ - restraints = self._restraints - if restraints is not None: - restraints._xyz_fn = self.xyz - restraints._adp_fn = self.adp - restraints._vdw_radii_fn = self.get_vdw_radii for module in self._modules.values(): if module is not None and hasattr(module, "set_xyz_fn"): module.set_xyz_fn(self.xyz) @@ -3049,7 +2926,6 @@ def reader(): reader.links = links strip_h = self.ctx.strip_H self.ctx.strip_H = False - self._restraints = None try: self.load(reader, add_hydrogens=False) finally: diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index 2547d939..27d85c6b 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -363,7 +363,7 @@ def __init__( ) self.setup_scaler() # The CIF path went in at construction; build the restraints over it now. - self.model._build_restraints() + self.model.restraints self._freeze_unrestrained_residues() # Initialize target functions (instantiated once, evaluated each iteration) @@ -387,7 +387,8 @@ def _freeze_unrestrained_residues(self): model = self.model pdb = getattr(model, "pdb", None) - acc = getattr(getattr(model, "_restraints", None), "restraints", None) + restraints = getattr(getattr(model, "ctx", None), "restraints", None) + acc = None if restraints is None else restraints.restraints if pdb is None or acc is None: return n = len(pdb) diff --git a/torchref/refinement/targets/adp/similarity.py b/torchref/refinement/targets/adp/similarity.py index 191d2379..a46db18a 100644 --- a/torchref/refinement/targets/adp/similarity.py +++ b/torchref/refinement/targets/adp/similarity.py @@ -116,7 +116,7 @@ def forward(self) -> torch.Tensor: def stats(self) -> Dict[str, any]: """Get SIMU restraint statistics.""" - b_diffs = self.restraints.adp_b_differences() + b_diffs = self.restraints.adp_b_differences(self.model.adp()) if len(b_diffs) == 0: return {} diff --git a/torchref/refinement/targets/geometry/angles.py b/torchref/refinement/targets/geometry/angles.py index 11281f2c..7b6b6e05 100644 --- a/torchref/refinement/targets/geometry/angles.py +++ b/torchref/refinement/targets/geometry/angles.py @@ -50,7 +50,7 @@ def forward(self) -> torch.Tensor: def stats(self) -> Dict[str, StatEntry]: """Get angle restraint statistics.""" - deviations_rad, sigmas_rad = self.restraints.angle_deviations() + deviations_rad, sigmas_rad = self.restraints.angle_deviations(self.model.xyz()) if len(deviations_rad) == 0: return {} diff --git a/torchref/refinement/targets/geometry/bonds.py b/torchref/refinement/targets/geometry/bonds.py index efe79b28..2b18e301 100644 --- a/torchref/refinement/targets/geometry/bonds.py +++ b/torchref/refinement/targets/geometry/bonds.py @@ -46,7 +46,7 @@ def forward(self) -> torch.Tensor: def stats(self) -> Dict[str, StatEntry]: """Get bond restraint statistics.""" - deviations, sigmas = self.restraints.bond_deviations() + deviations, sigmas = self.restraints.bond_deviations(self.model.xyz()) if len(deviations) == 0: return {} diff --git a/torchref/refinement/targets/geometry/non_bonded.py b/torchref/refinement/targets/geometry/non_bonded.py index 637be377..e61b001d 100644 --- a/torchref/refinement/targets/geometry/non_bonded.py +++ b/torchref/refinement/targets/geometry/non_bonded.py @@ -210,7 +210,7 @@ def maintenance(self) -> None: f" VDW rebuild: max drift {max_disp:.2f} Å > " f"threshold {thresh:.2f} Å" ) - r.rebuild_vdw_restraints() + r.rebuild_vdw_restraints(self._model.xyz().detach()) def _compute_positions( self, xyz: torch.Tensor diff --git a/torchref/refinement/targets/geometry/torsions.py b/torchref/refinement/targets/geometry/torsions.py index 8ee70047..b0c121cc 100644 --- a/torchref/refinement/targets/geometry/torsions.py +++ b/torchref/refinement/targets/geometry/torsions.py @@ -139,7 +139,9 @@ def forward(self) -> torch.Tensor: tdata["sigmas"], tdata["periods"], ) else: - deviations_rad, sigmas_deg = self.restraints.torsion_deviations_with_sigmas() + deviations_rad, sigmas_deg = self.restraints.torsion_deviations_with_sigmas( + xyz + ) if len(deviations_rad) > 0: total = total + _von_mises_nll(deviations_rad, sigmas_deg).sum() @@ -163,7 +165,9 @@ def stats(self) -> Dict[str, StatEntry]: result = {} # --- Intra-residue + disulfide stats --- - deviations_rad, sigmas_deg = self.restraints.torsion_deviations_with_sigmas() + deviations_rad, sigmas_deg = self.restraints.torsion_deviations_with_sigmas( + self.model.xyz() + ) if len(deviations_rad) > 0: deviations_deg = deviations_rad * (180.0 / np.pi) sigmas_rad = sigmas_deg * (np.pi / 180.0) @@ -184,7 +188,7 @@ def stats(self) -> Dict[str, StatEntry]: with torch.no_grad(): indices = omega_data["indices"] is_proline = omega_data["is_proline"] - omega_deg = self.restraints.torsions(indices) + omega_deg = self.restraints.torsions(indices, self.model.xyz()) is_cis = torch.abs(omega_deg) < 90.0 n_cis = int(is_cis.sum().item()) diff --git a/torchref/topology/nonbonded.py b/torchref/topology/nonbonded.py index 9f1b646a..25a658e7 100644 --- a/torchref/topology/nonbonded.py +++ b/torchref/topology/nonbonded.py @@ -8,7 +8,8 @@ k-d tree instead (:func:`find_pairs_kdtree`), with the same output. All operations run under ``torch.no_grad()`` on whatever device -the input coordinates live on (CPU or GPU). +the input coordinates live on (CPU or GPU). :func:`vdw_radii_for_elements` +gives the per-atom radii the contact distances are summed from. """ from typing import TYPE_CHECKING, Dict, List, Optional, Set, Tuple @@ -19,9 +20,49 @@ from torchref.config import dtypes, get_float_dtype if TYPE_CHECKING: + import pandas + from torchref.symmetry.cell import Cell from torchref.symmetry.spacegroup import SpaceGroup +#: Radius in Å for an element missing from ``atomic_vdw_radii.csv``. +_DEFAULT_VDW_RADIUS = 1.9 + + +def vdw_radii_for_elements(elements: "pandas.Series") -> np.ndarray: + """Van der Waals radius of each atom, looked up by element. + + Parameters + ---------- + elements : pandas.Series + Element symbols, one per atom; case and surrounding whitespace are ignored. + + Returns + ------- + numpy.ndarray + Radii in Å, shape ``(n_atoms,)``, float64. Elements the table does not list + get 1.9 Å. + """ + import os + + import pandas as pd + + from torchref import PATH_TORCHREF_DATA + + table = pd.read_csv( + os.path.join(PATH_TORCHREF_DATA, "atomic_vdw_radii.csv"), comment="#" + ) + radius = dict( + zip( + table["element"].str.strip().str.capitalize(), + table["vdW_Radius_Angstrom"], + ) + ) + symbols = elements.astype(str).str.strip().str.capitalize() + return np.array( + [radius.get(e, _DEFAULT_VDW_RADIUS) for e in symbols], dtype=np.float64 + ) + # ------------------------------------------------------------------ # # Step 1 – centroid pre-filter @@ -644,8 +685,8 @@ def filter_pairs( @torch.no_grad() def build_vdw_restraints_gpu( - xyz_fn, - vdw_radii_fn, + xyz: torch.Tensor, + vdw_radii: torch.Tensor, cell: "Cell", sg: "SpaceGroup", pdb, @@ -659,8 +700,10 @@ def build_vdw_restraints_gpu( Parameters ---------- - xyz_fn : callable returns (N, 3) Cartesian coordinates - vdw_radii_fn : callable returns (N,) VDW radii + xyz : torch.Tensor + ``(N, 3)`` Cartesian ASU coordinates in Å. + vdw_radii : torch.Tensor + ``(N,)`` van der Waals radii in Å. cell : Cell sg : SpaceGroup pdb : DataFrame @@ -678,7 +721,6 @@ def build_vdw_restraints_gpu( """ from torchref.symmetry.spacegroup import SpaceGroup as SG - xyz = xyz_fn() device = xyz.device fdtype = dtypes.float n_asu = xyz.shape[0] @@ -822,8 +864,6 @@ def build_vdw_restraints_gpu( symop_indices = op_indices[pair_combo_j] pair_cell_offsets = cell_offsets_valid[pair_combo_j] - # VDW radii - vdw_radii = vdw_radii_fn() min_distances = vdw_radii[pair_atom_i] + vdw_radii[pair_atom_j] # Build output diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index a79009dc..b0dd4590 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -14,13 +14,12 @@ it is held apart from the rest; * the Ramachandran map, a residue-level product of the same build. -Deliberately decoupled from :class:`~torchref.model.Model`: it takes an atom table plus -callables for coordinates, ADPs and van der Waals radii, so it can be built and tested -without one. +Deliberately decoupled from :class:`~torchref.model.Model`: it takes an atom table and +holds no reference back to whatever owns the coordinates. Every evaluation takes the +coordinates (or ADPs) it scores as an argument, and the pair list is rebuilt from the +coordinates it is handed, so the same object serves any model that shares the atom set. """ -from typing import Callable - import numpy as np import pandas as pd import torch @@ -41,26 +40,20 @@ class Restraints(DeviceMixin, DebugMixin, Module): """ Restraints handler for crystallographic model refinement. - Builds restraint tensors via the builder classes in ``builders_fast``. - Decoupled from Model: takes a pdb DataFrame plus callables for coordinates, - ADPs and VDW radii. - Parameters ---------- pdb : pd.DataFrame, optional DataFrame containing atomic structure data. If None, creates empty shell. cif_path : str or list of str, optional Path to the CIF restraints dictionary file(s). - xyz_fn : callable, optional - Returns current xyz coordinates. Required to build/evaluate when ``pdb`` - is provided. - adp_fn : callable, optional - Returns current ADP values. Required for ADP-based restraints. - vdw_radii_fn : callable, optional - Returns VDW radii. Required for VDW restraints. + xyz : torch.Tensor, optional + Cartesian coordinates in Å, shape ``(n_atoms, 3)``, that the topology and the + first pair list are built over. Defaults to the ``x``/``y``/``z`` columns of + ``pdb``. The build lands on this tensor's device. Not retained. cell : Cell, optional Crystallographic unit cell. Together with ``spacegroup``, enables - symmetry-aware VDW restraints (contacts with symmetry mates). + symmetry-aware VDW restraints (contacts with symmetry mates). Without both, + the pair list is searched in an isolated P1 box. spacegroup : SpaceGroup or str, optional Space group. Together with ``cell``, enables symmetry-aware VDW restraints. links : pd.DataFrame, optional @@ -95,9 +88,7 @@ def __init__( self, pdb: pd.DataFrame = None, cif_path=None, - xyz_fn: Callable[[], torch.Tensor] = None, - adp_fn: Callable[[], torch.Tensor] = None, - vdw_radii_fn: Callable[[], torch.Tensor] = None, + xyz: torch.Tensor = None, cell=None, spacegroup=None, links: pd.DataFrame = None, @@ -111,11 +102,6 @@ def __init__( self.links = links self._nonbonded = bool(nonbonded) - # Store callable functions for coordinate/ADP access - self._xyz_fn = xyz_fn - self._adp_fn = adp_fn - self._vdw_radii_fn = vdw_radii_fn - # Store crystallographic info for symmetry VDW restraints self._cell = cell self._spacegroup = spacegroup @@ -138,7 +124,14 @@ def __init__( return # Full initialization with pdb + from torchref.topology.nonbonded import vdw_radii_for_elements + self.pdb = pdb + if xyz is None: + xyz = torch.tensor(pdb[["x", "y", "z"]].values, dtype=get_float_dtype()) + self._vdw_radii = torch.tensor( + vdw_radii_for_elements(pdb["element"]), dtype=get_float_dtype() + ) self.unique_residues = pdb.resname.unique() self.unique_residues = [ residue @@ -156,86 +149,10 @@ def __init__( if verbose > 1: print(f"Loaded {len(self.link_dict)} link types") - # Build restraints using the new builder pattern - self.build_restraints() + self.build_restraints(xyz) if self.verbose > 0: self.summary() - def xyz(self, xyz: torch.Tensor = None) -> torch.Tensor: - """ - Get current xyz coordinates. - - Parameters - ---------- - xyz : torch.Tensor, optional - If provided, returns this tensor directly. - Otherwise calls the stored xyz_fn callable. - - Returns - ------- - torch.Tensor - Current xyz coordinates of shape (n_atoms, 3). - """ - if xyz is not None: - return xyz - if self._xyz_fn is None: - raise RuntimeError( - "No xyz callable provided. Initialize with xyz_fn or pass xyz argument." - ) - return self._xyz_fn() - - def adp(self, adp: torch.Tensor = None) -> torch.Tensor: - """ - Get current ADP values. - - Parameters - ---------- - adp : torch.Tensor, optional - If provided, returns this tensor directly. - Otherwise calls the stored adp_fn callable. - - Returns - ------- - torch.Tensor - Current ADP values. Shape ``(n_atoms,)`` for isotropic B-factors, - or ``(n_atoms, 6)`` for anisotropic ADPs (the six unique - components of the U tensor per atom), depending on what the stored - ``adp_fn`` (or the ``adp`` argument) supplies. - """ - if adp is not None: - return adp - if self._adp_fn is None: - raise RuntimeError( - "No adp callable provided. Initialize with adp_fn or pass adp argument." - ) - return self._adp_fn() - - def get_vdw_radii(self) -> torch.Tensor: - """ - Get VDW radii for all atoms. - - Returns - ------- - torch.Tensor - VDW radii of shape (n_atoms,). - """ - if self._vdw_radii_fn is None: - raise RuntimeError( - "No vdw_radii callable provided. Initialize with vdw_radii_fn." - ) - return self._vdw_radii_fn() - - # ========================================================================= - # Restraint storage - # ========================================================================= - - - - - - - - @property def restraints(self) -> dict: """Restraint groups as ``[edge type][origin][property]``. @@ -353,14 +270,18 @@ def _load_rama_surfaces(self, device: torch.device): surfaces = load_nll_surfaces(device) self.register_buffer("_rama_surfaces", surfaces) - def build_restraints(self): + def build_restraints(self, xyz: torch.Tensor): """Build the topology, the values over it, and the non-bonded pair list. - The pair list is skipped when constructed with ``nonbonded=False``. Builds on - CPU and moves the result to the ``xyz()`` device at the end. + Parameters + ---------- + xyz : torch.Tensor + Cartesian coordinates in Å, shape ``(n_atoms, 3)``. The pair list is + skipped when constructed with ``nonbonded=False``. Builds on CPU and moves + the result to ``xyz``'s device at the end. """ try: - target_device = self.xyz().device + target_device = xyz.device device = torch.device("cpu") from torchref.topology import build_topology_with_values @@ -371,7 +292,7 @@ def build_restraints(self): link_dict=self.link_dict, link_list=self.link_list, links=self.links, - xyz=self.xyz().detach().to(device), + xyz=xyz.detach().to(device), device=device, verbose=self.verbose, ) @@ -389,7 +310,7 @@ def build_restraints(self): # margin and cannot miss a newly-formed contact. if self._nonbonded: self._build_vdw_restraints( - cutoff=6.0, sigma=0.05, inter_residue_only=False, use_spatial_hash=True + xyz, cutoff=6.0, sigma=0.05, inter_residue_only=False ) if target_device.type != "cpu": @@ -399,261 +320,6 @@ def build_restraints(self): self.debug_on_error(e, context="Restraints.build_restraints") raise - - - - - - - def _find_nearby_pairs_spatial_hash(self, xyz, cutoff=6.0): - """Atom pairs within ``cutoff`` of each other, as (M, 2) rows with i < j. - - Cell-list search over cubic cells of side ``cutoff``, checking only the 14 - unique offsets (self + 13 forward neighbours): O(N) memory instead of the - O(N^2) distance matrix. Runs on CPU regardless of ``xyz``'s device. - """ - device = xyz.device - n_atoms = xyz.shape[0] - - if n_atoms == 0: - return torch.tensor([], dtype=torch.long, device=device).reshape(0, 2) # dtype-ok: empty atom-pair index tensor; int64 index required - - # Work on CPU to avoid per-iteration GPU kernel launch overhead - coords = xyz.detach().cpu() - cell_size = cutoff - - # Assign each atom to a cubic cell - xyz_min = coords.min(dim=0).values - cell_idx = ((coords - xyz_min) / cell_size).long() # (N, 3) - - grid_dims = cell_idx.max(dim=0).values + 1 - gx, gy, gz = grid_dims[0].item(), grid_dims[1].item(), grid_dims[2].item() - gyz = gy * gz - - # Flat cell index per atom - flat = cell_idx[:, 0] * gyz + cell_idx[:, 1] * gz + cell_idx[:, 2] - - # Sort atoms by cell so each cell's atoms are contiguous - order = flat.argsort() - sorted_flat = flat[order] - - unique_cells, counts = torch.unique_consecutive( - sorted_flat, return_counts=True - ) - n_unique = len(unique_cells) - starts = torch.zeros(n_unique + 1, dtype=torch.long) # dtype-ok: grid-cell CSR start offsets; int64 index required - starts[1:] = counts.cumsum(0) - - # Lookup: flat_cell -> index in unique_cells (-1 if empty) - n_grid = gx * gyz - cell_lookup = torch.full((n_grid,), -1, dtype=torch.long) # dtype-ok: cell lookup table (-1 sentinel); int64 index required - cell_lookup[unique_cells] = torch.arange(n_unique) - - # 14 unique neighbour offsets: self (0,0,0) + 13 forward neighbours. - # "Forward" = first non-zero component is positive, avoiding double counting. - offsets_list = [] - for dx in range(-1, 2): - for dy in range(-1, 2): - for dz in range(-1, 2): - if ( - dx > 0 - or (dx == 0 and dy > 0) - or (dx == 0 and dy == 0 and dz >= 0) - ): - offsets_list.append( - (dx, dy, dz, dx * gyz + dy * gz + dz) - ) - - cutoff_sq = cutoff * cutoff - pair_chunks = [] - - # Move to numpy for tight loop (faster item access than torch on CPU) - unique_np = unique_cells.numpy() - starts_np = starts.numpy() - order_np = order.numpy() - coords_np = coords.numpy() - - for ci in range(n_unique): - cell_flat = int(unique_np[ci]) - sa, ea = int(starts_np[ci]), int(starts_np[ci + 1]) - atoms_a = order_np[sa:ea] - xyz_a = coords_np[atoms_a] # (na, 3) - - cx = cell_flat // gyz - cy = (cell_flat % gyz) // gz - cz = cell_flat % gz - - for dx, dy, dz, off_flat in offsets_list: - ncx, ncy, ncz = cx + dx, cy + dy, cz + dz - if ( - ncx < 0 or ncx >= gx - or ncy < 0 or ncy >= gy - or ncz < 0 or ncz >= gz - ): - continue - - nb_flat = ncx * gyz + ncy * gz + ncz - nb_ci = int(cell_lookup[nb_flat]) - if nb_ci < 0: - continue - - sb, eb = int(starts_np[nb_ci]), int(starts_np[nb_ci + 1]) - atoms_b = order_np[sb:eb] - xyz_b = coords_np[atoms_b] # (nb, 3) - - # Vectorised distance² via broadcasting: (na, nb, 3) - diff = xyz_a[:, None, :] - xyz_b[None, :, :] - dist_sq = (diff * diff).sum(axis=-1) # (na, nb) - - if off_flat == 0: - # Self-cell: upper triangle only - na = len(atoms_a) - if na < 2: - continue - ii, jj = np.triu_indices(na, k=1) - mask = dist_sq[ii, jj] < cutoff_sq - if mask.any(): - ai = atoms_a[ii[mask]] - aj = atoms_a[jj[mask]] - pairs = np.stack( - [np.minimum(ai, aj), np.maximum(ai, aj)], axis=1 - ) - pair_chunks.append(pairs) - else: - # Inter-cell: all pairs - ii, jj = np.where(dist_sq < cutoff_sq) - if len(ii) > 0: - ai = atoms_a[ii] - bj = atoms_b[jj] - pairs = np.stack( - [np.minimum(ai, bj), np.maximum(ai, bj)], axis=1 - ) - pair_chunks.append(pairs) - - if pair_chunks: - all_pairs = np.concatenate(pair_chunks, axis=0) - return torch.from_numpy(all_pairs).to(dtype=torch.long, device=device) # dtype-ok: atom-pair index array from numpy; int64 index required - else: - return torch.tensor([], dtype=torch.long, device=device).reshape(0, 2) # dtype-ok: empty atom-pair index tensor; int64 index required - - def _expand_with_symmetry_mates(self, xyz, cutoff): - """Append symmetry-mate positions to ASU ``xyz`` for neighbour search. - - Centroid pre-filtering skips mates that cannot reach within ``cutoff``. - Returns ``(combined_xyz, provenance)``; the ASU occupies rows ``[:N]`` and - ``provenance`` maps every row back with numpy ``asu_source_indices``, - ``symop_indices`` and ``(N, 3)`` ``cell_offsets``. - """ - from torchref.config import dtypes - from torchref.symmetry import SpaceGroup - - cell = self._cell - sg = self._spacegroup - if not isinstance(sg, SpaceGroup): - sg = SpaceGroup(sg) - - n_asu = xyz.shape[0] - device = xyz.device - fdtype = dtypes.float - - # Work on the model's device throughout - xyz_det = xyz.detach().to(fdtype) - xyz_frac = cell.cartesian_to_fractional(xyz_det) - - # Compute centroid and molecule radius for pre-filtering - centroid_frac = xyz_frac.mean(dim=0) - centroid_cart = xyz_det.mean(dim=0) - molecule_radius = (xyz_det - centroid_cart).norm(dim=1).max().item() - threshold = 2 * molecule_radius + cutoff - - B = cell.fractional_matrix.to(device=device, dtype=fdtype) - I_mat = torch.eye(3, dtype=fdtype, device=device) - - # Phase 1: centroid pre-filter to find which (symop, offset) combos - # can produce contacts. This is a small loop over scalar ops. - n_ops = sg.n_ops - matrices = sg.matrices.to(device=device, dtype=fdtype) - translations = sg.translations.to(device=device, dtype=fdtype) - - valid_ops = [] # list of (op_idx, dx, dy, dz) - for op_idx in range(n_ops): - R = matrices[op_idx] - t = translations[op_idx] - for dx in range(-1, 2): - for dy in range(-1, 2): - for dz in range(-1, 2): - if op_idx == 0 and dx == 0 and dy == 0 and dz == 0: - continue - offset = torch.tensor([dx, dy, dz], dtype=fdtype, - device=device) - d_frac = (R - I_mat) @ centroid_frac + t + offset - d_cart = B @ d_frac - if d_cart.norm().item() <= threshold: - valid_ops.append((op_idx, dx, dy, dz)) - - if not valid_ops: - provenance = { - "asu_source_indices": np.arange(n_asu, dtype=np.int64), - "symop_indices": np.zeros(n_asu, dtype=np.int64), - "cell_offsets": np.zeros((n_asu, 3), dtype=np.int64), - } - if self.verbose > 0: - print(" Symmetry expansion: 0 mate(s) within range " - f"({n_asu} total atoms for neighbor search)") - return xyz_det, provenance - - # Phase 2: batch-generate all mate coordinates in one go - n_valid = len(valid_ops) - op_indices = [v[0] for v in valid_ops] - cell_offs = torch.tensor( - [[v[1], v[2], v[3]] for v in valid_ops], dtype=fdtype, - device=device, - ) # (n_valid, 3) - - # Gather rotation matrices and translations for valid ops - R_batch = matrices[op_indices] # (n_valid, 3, 3) - t_batch = translations[op_indices] # (n_valid, 3) - - # Batched transform: for each valid op, compute R @ xyz_frac.T + t + offset - # xyz_frac: (N, 3), R_batch: (n_valid, 3, 3) - # -> (n_valid, 3, N) via batched matmul, then transpose to (n_valid, N, 3) - xyz_frac_T = xyz_frac.T.unsqueeze(0).expand(n_valid, -1, -1) # (n_valid, 3, N) - mate_frac_all = torch.bmm(R_batch, xyz_frac_T).permute(0, 2, 1) # (n_valid, N, 3) - mate_frac_all = mate_frac_all + t_batch.unsqueeze(1) + cell_offs.unsqueeze(1) - - # Convert all to Cartesian: (n_valid * N, 3) - mate_frac_flat = mate_frac_all.reshape(-1, 3) - mate_cart_flat = cell.fractional_to_cartesian(mate_frac_flat) - - # Build combined coordinate array: ASU + all mates - combined_xyz = torch.cat( - [xyz_det, mate_cart_flat], dim=0 - ) - - # Build provenance arrays - asu_source = np.arange(n_asu, dtype=np.int64) - # ASU block - all_asu_sources = [asu_source] - all_symops = [np.zeros(n_asu, dtype=np.int64)] - all_offsets = [np.zeros((n_asu, 3), dtype=np.int64)] - # Mate blocks (each has n_asu atoms) - for op_idx, dx, dy, dz in valid_ops: - all_asu_sources.append(asu_source) - all_symops.append(np.full(n_asu, op_idx, dtype=np.int64)) - all_offsets.append(np.tile([dx, dy, dz], (n_asu, 1)).astype(np.int64)) - - provenance = { - "asu_source_indices": np.concatenate(all_asu_sources), - "symop_indices": np.concatenate(all_symops), - "cell_offsets": np.concatenate(all_offsets), - } - - if self.verbose > 0: - print(f" Symmetry expansion: {n_valid} mate(s) within range " - f"({combined_xyz.shape[0]} total atoms for neighbor search)") - - return combined_xyz, provenance - @property def h_topo(self): """Access riding hydrogen topology (None if not built).""" @@ -701,17 +367,19 @@ def _build_h_exclusion_hash(self, h_topo, device): return torch.tensor(hashes, dtype=torch.long, device=device) # dtype-ok: grid-cell hash values used as keys/index; int64 required def _build_vdw_restraints( - self, cutoff=6.0, sigma=0.2, inter_residue_only=True, use_spatial_hash=True + self, xyz, cutoff=6.0, sigma=0.2, inter_residue_only=True ): """Build van der Waals (non-bonded contact) restraints. - With cell and spacegroup present, includes contacts to symmetry mates via - the GPU-native periodic grid search; otherwise falls back to - :meth:`_build_vdw_restraints_legacy` (whose own default cutoff is 5.0). + With cell and spacegroup present, includes contacts to symmetry mates. + Without them the same periodic search runs in an isolated P1 box, one + cutoff wider than the model on every side, so no image comes within range. Also builds the riding-hydrogen topology for H-VDW evaluation. Parameters ---------- + xyz : torch.Tensor + Cartesian coordinates in Å, shape ``(n_atoms, 3)``. cutoff : float, default 6.0 Contact-search cutoff in Angstroms. Keep it ~1 Å beyond the largest heavy-atom VDW sum so the rebuild threshold has margin. @@ -719,8 +387,6 @@ def _build_vdw_restraints( Restraint sigma in Angstroms (the production caller passes 0.05). inter_residue_only : bool, default True If True, only build contacts between atoms in different residues. - use_spatial_hash : bool, default True - Use the spatial-hash neighbour search in the legacy (no-symmetry) path. Notes ----- @@ -732,64 +398,49 @@ def _build_vdw_restraints( cutoff=cutoff, sigma=sigma, inter_residue_only=inter_residue_only, - use_spatial_hash=use_spatial_hash, ) if self.verbose > 0: print("\nBuilding VDW (non-bonded) restraints...") - has_symmetry = ( - self._cell is not None - and self._spacegroup is not None - ) - # The build (neighbour search, H topology, exclusion hashing) runs on CPU; # everything it registers is migrated to target_device at the end, which # the maintenance-triggered rebuild path depends on. cpu = torch.device("cpu") - target_device = self.xyz().device if self._xyz_fn is not None else cpu - - def xyz_cpu(): - return self.xyz().detach().to(cpu) - - def vdw_radii_cpu(): - return self.get_vdw_radii().detach().to(cpu) - - # Construct fresh CPU copies — Cell/SpaceGroup ``.to()`` mutates - # in place, which would silently relocate the model's own Cell/SG. - if self._cell is not None: - from torchref.symmetry.cell import Cell - cell_cpu = Cell(self._cell._data.detach(), device=cpu, - dtype=self._cell.dtype) - else: - cell_cpu = None - if self._spacegroup is not None: - sg_cpu = self._spacegroup.copy().to(cpu) - else: - sg_cpu = None + target_device = xyz.device + xyz_cpu = xyz.detach().to(cpu) + radii_cpu = self._vdw_radii.to(cpu) - if has_symmetry: - from torchref.topology.nonbonded import build_vdw_restraints_gpu + from torchref.symmetry import SpaceGroup + from torchref.symmetry.cell import Cell - exclusions = self.topology.atoms.exclusions_from_restraint_edges() - self._vdw = build_vdw_restraints_gpu( - xyz_fn=xyz_cpu, - vdw_radii_fn=vdw_radii_cpu, - cell=cell_cpu, - sg=sg_cpu, - pdb=self.pdb, - exclusion_set=exclusions, - cutoff=cutoff, - sigma=sigma, - inter_residue_only=inter_residue_only, - verbose=self.verbose, + # Fresh CPU copies: Cell/SpaceGroup ``.to()`` mutates in place, which would + # silently relocate the model's own Cell/SG. + if self._cell is not None and self._spacegroup is not None: + cell_cpu = Cell( + self._cell._data.detach(), device=cpu, dtype=self._cell.dtype ) + sg_cpu = self._spacegroup.copy().to(cpu) else: - self._build_vdw_restraints_legacy( - cutoff=cutoff, sigma=sigma, - inter_residue_only=inter_residue_only, - use_spatial_hash=use_spatial_hash, - ) + extent = float((xyz_cpu.max(dim=0).values - xyz_cpu.min(dim=0).values).max()) + side = extent + 2.0 * cutoff + cell_cpu = Cell([side, side, side, 90.0, 90.0, 90.0], device=cpu) + sg_cpu = SpaceGroup("P 1", device=cpu) + + from torchref.topology.nonbonded import build_vdw_restraints_gpu + + self._vdw = build_vdw_restraints_gpu( + xyz=xyz_cpu, + vdw_radii=radii_cpu, + cell=cell_cpu, + sg=sg_cpu, + pdb=self.pdb, + exclusion_set=self.topology.atoms.exclusions_from_restraint_edges(), + cutoff=cutoff, + sigma=sigma, + inter_residue_only=inter_residue_only, + verbose=self.verbose, + ) # Publish the new pair list before anything reads it back below. Unlike the # geometry edges it is not derived from the topology, so it is held separately @@ -833,7 +484,7 @@ def vdw_radii_cpu(): ) # Fill in VDW min distances using combined radii array if self._h_topo.has_candidates: - heavy_radii = vdw_radii_cpu() # (N_heavy,) + heavy_radii = radii_cpu # (N_heavy,) h_radii = self._h_topo.h_vdw_radius # (N_h,) on CPU all_radii = torch.cat([heavy_radii, h_radii]) self._h_topo.cand_min_dist = ( @@ -843,229 +494,31 @@ def vdw_radii_cpu(): # Snapshot at build time so maintenance() can diff current positions # against it; kept on the model device so the compare is one op. - if self._xyz_fn is not None: - self._last_vdw_build_xyz = self.xyz().detach().clone() + self._last_vdw_build_xyz = xyz.detach().clone() # Move the CPU-built pair list, h_topo and h_excl_hash to the model device. # The rebuild path has no surrounding migration, so this cannot be dropped. if target_device.type != "cpu": self.to(target_device) - def rebuild_vdw_restraints(self) -> None: - """Refresh the VDW pair list with the kwargs the initial build was given. + def rebuild_vdw_restraints(self, xyz: torch.Tensor) -> None: + """Refresh the VDW pair list over ``xyz`` with the initial build's kwargs. Called by :meth:`NonBondedTarget.maintenance` once max atomic displacement since the last build exceeds its threshold. Raises ``RuntimeError`` if no initial build has run. + + Parameters + ---------- + xyz : torch.Tensor + Current Cartesian coordinates in Å, shape ``(n_atoms, 3)``. """ if not hasattr(self, "_vdw_build_kwargs"): raise RuntimeError( "rebuild_vdw_restraints called before initial build " "— _vdw_build_kwargs is missing" ) - self._build_vdw_restraints(**self._vdw_build_kwargs) - - def _build_vdw_restraints_legacy( - self, cutoff=5.0, sigma=0.2, inter_residue_only=True, use_spatial_hash=True - ): - """Legacy VDW restraint builder (no symmetry or CPU fallback).""" - - exclusions = self.topology.atoms.exclusions_from_restraint_edges() - vdw_radii = self.get_vdw_radii() - xyz = self.xyz() - device = xyz.device - pdb = self.pdb - n_asu = xyz.shape[0] - - # Expand with symmetry mates if crystallographic info is available - has_symmetry = ( - self._cell is not None - and self._spacegroup is not None - ) - if has_symmetry: - combined_xyz, provenance = self._expand_with_symmetry_mates(xyz, cutoff) - else: - combined_xyz = xyz - provenance = None - - # Find nearby pairs in the (potentially expanded) coordinate set - if use_spatial_hash: - nearby_pairs = self._find_nearby_pairs_spatial_hash(combined_xyz, cutoff) - else: - n_total = combined_xyz.shape[0] - pairs_list = [] - cutoff_sq = cutoff**2 - for i in range(n_total): - for j in range(i + 1, n_total): - dist_sq = ((combined_xyz[i] - combined_xyz[j]) ** 2).sum() - if dist_sq < cutoff_sq: - pairs_list.append([i, j]) - nearby_pairs = ( - torch.tensor(pairs_list, dtype=torch.long, device=device) # dtype-ok: atom-pair index tensor; int64 index required - if pairs_list - else torch.tensor([], dtype=torch.long, device=device).reshape(0, 2) # dtype-ok: empty atom-pair index tensor; int64 index required - ) - - empty_result = { - "indices": torch.tensor([], dtype=torch.long, device=device).reshape(0, 2), # dtype-ok: empty atom-pair index tensor; int64 index required - "min_distances": torch.tensor([], dtype=get_float_dtype(), device=device), - "sigmas": torch.tensor([], dtype=get_float_dtype(), device=device), - "symop_indices": torch.tensor([], dtype=torch.long, device=device), # dtype-ok: empty symop index tensor; int64 index required - "cell_offsets": torch.tensor([], dtype=torch.long, device=device).reshape(0, 3), # dtype-ok: empty cell-offset index tensor; int64 index required - } - - if len(nearby_pairs) == 0: - self._vdw = empty_result - return - - pairs_np = nearby_pairs.cpu().numpy() - - # Map indices through provenance to get ASU source atoms and symop info - if provenance is not None: - prov_asu = provenance["asu_source_indices"] - prov_sym = provenance["symop_indices"] - prov_off = provenance["cell_offsets"] - - # Get provenance for each atom in each pair - idx0 = pairs_np[:, 0] - idx1 = pairs_np[:, 1] - - asu_src_0 = prov_asu[idx0] - asu_src_1 = prov_asu[idx1] - sym_0 = prov_sym[idx0] - sym_1 = prov_sym[idx1] - off_0 = prov_off[idx0] - off_1 = prov_off[idx1] - - is_asu_0 = (sym_0 == 0) & (off_0 == 0).all(axis=1) - is_asu_1 = (sym_1 == 0) & (off_1 == 0).all(axis=1) - - # Keep only pairs where at least one atom is from the ASU - has_asu = is_asu_0 | is_asu_1 - pairs_np = pairs_np[has_asu] - asu_src_0 = asu_src_0[has_asu] - asu_src_1 = asu_src_1[has_asu] - sym_0 = sym_0[has_asu] - sym_1 = sym_1[has_asu] - off_0 = off_0[has_asu] - off_1 = off_1[has_asu] - is_asu_0 = is_asu_0[has_asu] - is_asu_1 = is_asu_1[has_asu] - - # Normalize: put the ASU atom in position 0, mate in position 1 - # For intra-ASU pairs (both ASU), keep as-is (both are ASU anyway) - # For symmetry pairs: swap so ASU is first - swap = ~is_asu_0 & is_asu_1 - if swap.any(): - asu_src_0[swap], asu_src_1[swap] = asu_src_1[swap].copy(), asu_src_0[swap].copy() - sym_0[swap], sym_1[swap] = sym_1[swap].copy(), sym_0[swap].copy() - off_0[swap], off_1[swap] = off_1[swap].copy(), off_0[swap].copy() - is_asu_0[swap] = True - is_asu_1[swap] = False - - # Final indices: ASU atom indices for both atoms in each pair - final_i1 = asu_src_0 - final_i2 = asu_src_1 - # Symmetry info comes from the mate atom (position 1) - final_symop = sym_1 - final_offsets = off_1 - - is_both_asu = is_asu_0 & is_asu_1 - else: - # No symmetry: all pairs are intra-ASU - final_i1 = pairs_np[:, 0] - final_i2 = pairs_np[:, 1] - final_symop = np.zeros(len(pairs_np), dtype=np.int64) - final_offsets = np.zeros((len(pairs_np), 3), dtype=np.int64) - is_both_asu = np.ones(len(pairs_np), dtype=bool) - - # --- Filtering --- - # Bonded exclusions, same-residue, and altloc filters apply only to - # intra-ASU pairs. Symmetry pairs cannot be bonded. - - # Start with all pairs kept - keep_mask = np.ones(len(final_i1), dtype=bool) - - # Exclusion mask (bonded 1-2, 1-3, 1-4) -- intra-ASU only - if exclusions and is_both_asu.any(): - exclusion_arr = np.array(list(exclusions), dtype=np.int64) - max_idx = max( - pdb["index"].max() + 1, - final_i1[is_both_asu].max() + 1, - final_i2[is_both_asu].max() + 1, - ) - # Normalize pair order for comparison - norm_i1 = np.minimum(final_i1, final_i2) - norm_i2 = np.maximum(final_i1, final_i2) - pair_hash = norm_i1 * max_idx + norm_i2 - excl_hash = exclusion_arr[:, 0] * max_idx + exclusion_arr[:, 1] - is_excluded = np.isin(pair_hash, excl_hash) - # Only apply to intra-ASU pairs - keep_mask &= ~(is_excluded & is_both_asu) - - # Inter-residue mask -- intra-ASU only - if inter_residue_only: - chainid_array = pdb["chainid"].values - resseq_array = pdb["resseq"].values - same_residue = ( - (chainid_array[final_i1] == chainid_array[final_i2]) - & (resseq_array[final_i1] == resseq_array[final_i2]) - ) - keep_mask &= ~(same_residue & is_both_asu) - - # Altloc compatibility -- intra-ASU only - if "altloc" in pdb.columns: - altloc_array = pdb["altloc"].values.astype(str) - altloc_array = np.where( - np.isin(altloc_array, ["", " "]), " ", altloc_array - ) - altloc_i = altloc_array[final_i1] - altloc_j = altloc_array[final_i2] - incompatible_altloc = ( - (altloc_i != " ") & (altloc_j != " ") & (altloc_i != altloc_j) - ) - keep_mask &= ~(incompatible_altloc & is_both_asu) - - # Apply filter - final_i1 = final_i1[keep_mask] - final_i2 = final_i2[keep_mask] - final_symop = final_symop[keep_mask] - final_offsets = final_offsets[keep_mask] - - if len(final_i1) == 0: - self._vdw = empty_result - return - - # Compute min distances using VDW radii of ASU source atoms. - vdw_np = vdw_radii.cpu().numpy() - min_distances = vdw_np[final_i1] + vdw_np[final_i2] - - # Store results - final_pairs = np.stack([final_i1, final_i2], axis=1) - self._vdw = { - "indices": torch.tensor(final_pairs, dtype=torch.long, device=device), # dtype-ok: final atom-pair index tensor; int64 index required - "min_distances": torch.tensor( - min_distances, dtype=get_float_dtype(), device=device - ), - "sigmas": torch.full( - (len(final_pairs),), sigma, dtype=get_float_dtype(), device=device - ), - "symop_indices": torch.tensor( - final_symop, dtype=torch.long, device=device # dtype-ok: symop index tensor; int64 index required - ), - "cell_offsets": torch.tensor( - final_offsets, dtype=torch.long, device=device # dtype-ok: cell-offset index tensor; int64 index required - ), - } - - if self.verbose > 0: - scope = "inter-residue" if inter_residue_only else "all" - msg = f" Built {len(final_pairs)} VDW restraints ({scope} contacts)" - if has_symmetry: - is_sym_pair = (final_symop != 0) | (final_offsets != 0).any(axis=1) - n_sym_count = int(is_sym_pair.sum()) - msg += f", {n_sym_count} symmetry contacts" - print(msg) + self._build_vdw_restraints(xyz, **self._vdw_build_kwargs) # Device movement goes through DeviceMixin: the topology and the value tensors are # walked and moved, and _apply re-slices the derived entry views afterwards. @@ -1172,7 +625,7 @@ def get_count(rtype, origin): - def bond_lengths(self, idx, xyz: torch.Tensor = None): + def bond_lengths(self, idx, xyz: torch.Tensor): """ Compute current bond lengths from atomic coordinates. @@ -1180,16 +633,14 @@ def bond_lengths(self, idx, xyz: torch.Tensor = None): ---------- idx : torch.Tensor Bond indices tensor of shape (N, 2). - xyz : torch.Tensor, optional - Coordinates tensor of shape (n_atoms, 3). - If None, uses the stored xyz_fn callable. + xyz : torch.Tensor + Cartesian coordinates in Å, shape (n_atoms, 3). Returns ------- torch.Tensor Tensor of bond lengths of shape (N,). """ - xyz = self.xyz(xyz) if idx is None: return torch.tensor([], device=xyz.device) pos1 = xyz[idx[:, 0], :] @@ -1211,32 +662,18 @@ def copy(self): """ import copy - # The coordinate and ADP accessors are borrowed from the model, not owned: - # duplicating them would hand the copy a third, orphaned parameter set (and - # deep-copying a wrapper with a cached forward fails on its graph tensor). - # They are carried across by reference; the owning model re-points them. - borrowed = ("_xyz_fn", "_adp_fn", "_vdw_radii_fn") - saved = {name: getattr(self, name, None) for name in borrowed} - for name in borrowed: - setattr(self, name, None) - try: - duplicate = copy.deepcopy(self) - finally: - for name, value in saved.items(): - setattr(self, name, value) - for name, value in saved.items(): - setattr(duplicate, name, value) + duplicate = copy.deepcopy(self) duplicate._rebuild_entries() return duplicate - def bond_deviations(self, xyz: torch.Tensor = None): + def bond_deviations(self, xyz: torch.Tensor): """ Compute bond length deviations and sigmas. Parameters ---------- - xyz : torch.Tensor, optional - Coordinates tensor. If None, uses the stored xyz_fn callable. + xyz : torch.Tensor + Cartesian coordinates in Å, shape (n_atoms, 3). Returns ------- @@ -1258,7 +695,7 @@ def bond_deviations(self, xyz: torch.Tensor = None): return deviations, sigmas - def nll_bonds(self, xyz: torch.Tensor = None): + def nll_bonds(self, xyz: torch.Tensor): """ Compute negative log-likelihood for bond length restraints. @@ -1269,8 +706,8 @@ def nll_bonds(self, xyz: torch.Tensor = None): Parameters ---------- - xyz : torch.Tensor, optional - Coordinates tensor. If None, uses the stored xyz_fn callable. + xyz : torch.Tensor + Cartesian coordinates in Å, shape (n_atoms, 3). Returns ------- @@ -1282,7 +719,7 @@ def nll_bonds(self, xyz: torch.Tensor = None): deviations, sigmas = self.bond_deviations(xyz) return gaussian_nll(deviations, sigmas) - def angles(self, idx, xyz: torch.Tensor = None): + def angles(self, idx, xyz: torch.Tensor): """ Compute current angle values for all angle restraints. @@ -1290,15 +727,14 @@ def angles(self, idx, xyz: torch.Tensor = None): ---------- idx : torch.Tensor Angle indices tensor of shape (N, 3). - xyz : torch.Tensor, optional - Coordinates tensor. If None, uses the stored xyz_fn callable. + xyz : torch.Tensor + Cartesian coordinates in Å, shape (n_atoms, 3). Returns ------- torch.Tensor Tensor of shape (n_angles,) with current angle values in degrees. """ - xyz = self.xyz(xyz) pos1 = xyz[idx[:, 0], :] pos2 = xyz[idx[:, 1], :] pos3 = xyz[idx[:, 2], :] @@ -1322,14 +758,14 @@ def angles(self, idx, xyz: torch.Tensor = None): return angles_deg - def angle_deviations(self, xyz: torch.Tensor = None): + def angle_deviations(self, xyz: torch.Tensor): """ Compute angle deviations and sigmas. Parameters ---------- - xyz : torch.Tensor, optional - Coordinates tensor. If None, uses the stored xyz_fn callable. + xyz : torch.Tensor + Cartesian coordinates in Å, shape (n_atoms, 3). Returns ------- @@ -1354,7 +790,7 @@ def angle_deviations(self, xyz: torch.Tensor = None): return deviations, sigmas_rad - def nll_angles(self, xyz: torch.Tensor = None): + def nll_angles(self, xyz: torch.Tensor): """ Compute negative log-likelihood for angle restraints. @@ -1365,8 +801,8 @@ def nll_angles(self, xyz: torch.Tensor = None): Parameters ---------- - xyz : torch.Tensor, optional - Coordinates tensor. If None, uses the stored xyz_fn callable. + xyz : torch.Tensor + Cartesian coordinates in Å, shape (n_atoms, 3). Returns ------- @@ -1394,7 +830,7 @@ def cat_dict(self): if self.topology is not None and "all" not in self._entries.get("bond", {}): self._rebuild_entries() - def torsions(self, idx, xyz: torch.Tensor = None): + def torsions(self, idx, xyz: torch.Tensor): """ Compute current torsion angle values for all torsion restraints. @@ -1402,16 +838,14 @@ def torsions(self, idx, xyz: torch.Tensor = None): ---------- idx : torch.Tensor Torsion indices tensor of shape (N, 4). - xyz : torch.Tensor, optional - Coordinates tensor. If None, uses the stored xyz_fn callable. + xyz : torch.Tensor + Cartesian coordinates in Å, shape (n_atoms, 3). Returns ------- torch.Tensor Tensor of shape (n_torsions,) with current torsion values in degrees. """ - xyz = self.xyz(xyz) - pos1 = xyz[idx[:, 0], :] pos2 = xyz[idx[:, 1], :] pos3 = xyz[idx[:, 2], :] @@ -1514,14 +948,14 @@ def _wrap_torsion_periodicity(self, diff_rad, periods): return torch.remainder(diff_rad + torch.pi, 2.0 * torch.pi) - torch.pi - def torsion_deviations_with_sigmas(self, xyz: torch.Tensor = None): + def torsion_deviations_with_sigmas(self, xyz: torch.Tensor): """ Compute torsion deviations (wrapped for periodicity) and sigmas. Parameters ---------- - xyz : torch.Tensor, optional - Coordinates tensor. If None, uses the stored xyz_fn callable. + xyz : torch.Tensor + Cartesian coordinates in Å, shape (n_atoms, 3). Returns ------- @@ -1549,21 +983,21 @@ def torsion_deviations_with_sigmas(self, xyz: torch.Tensor = None): - def adp_b_differences(self, adp: torch.Tensor = None): + def adp_b_differences(self, adp: torch.Tensor): """ Compute B-factor differences between bonded atoms. Parameters ---------- - adp : torch.Tensor, optional - ADP values. If None, uses the stored adp_fn callable. + adp : torch.Tensor + Isotropic B-factors in Ų, shape (n_atoms,). Returns ------- torch.Tensor Tensor of B-factor differences (B_i - B_j) for all bonds. """ - b_factors = self.adp(adp) + b_factors = adp diffs_list = [] if "bond" in self.restraints: From a0aeec6f9bccc68352c49a2ade61229f344b6ed7 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 21:51:12 +0200 Subject: [PATCH 197/250] Settle hydrogens on the context and build every model through one install path The hydrogen policy is two settings on ModelContext, applied once when the atom table is settled: hydrogens (keep / add / strip) and hydrogen_mode (atoms / riding); strip with riding raises. They replace strip_H, add_hydrogens and the "free" mode on Model, ModelFT, Refinement, EnsembleModel, load_model and the refine CLI (--hydrogens / --hydrogen-mode). Water hydrogens are generated at load, so set_hydrogen_mode only swaps the coordinate wrapper. ModelContext.from_atoms is the one place an atom table is settled and Model._install_parameters the one place wrappers are built. load, select, copy, create_from_state_dict and the strip/hydrogenate helpers all go through them, so ModelFT keeps only its own settings and the duplicated copy/restore paths, the re-entrant loads and the water completion mask surgery are gone. Atom-table queries (occupancy groups, altlocs, chain sequences) moved to the context. Co-Authored-By: Claude Opus 5.5 (1M context) --- AGENTS.md | 4 +- docs/changelog.rst | 4 + docs/user_guide/cli.rst | 8 +- tests/benchmarks/compare_amber_classic.py | 7 +- tests/integration/test_cli_hydrogens.py | 48 +- tests/integration/test_io_cif.py | 2 +- tests/integration/test_model_operations.py | 2 +- tests/unit/io/test_hkl_convention.py | 2 +- tests/unit/model/test_hydrogen_default.py | 50 +- tests/unit/model/test_hydrogen_mode.py | 63 +- tests/unit/model/test_hydrogens_in_xray.py | 14 +- tests/unit/model/test_model.py | 10 +- tests/unit/model/test_riding_orientations.py | 2 +- .../model/test_riding_water_completion.py | 215 ++- tests/unit/model/test_riding_xyz.py | 2 +- tests/unit/model/test_sf_grid_key.py | 2 +- tests/unit/monomer/test_link_modifications.py | 2 +- tests/unit/refinement/test_amber_target.py | 16 +- .../unit/refinement/test_forcefield_target.py | 6 +- tests/unit/topology/test_energy_types.py | 4 +- tests/unit/topology/test_hydrogen_frames.py | 4 +- tests/unit/topology/test_hydrogens.py | 10 +- tests/unit/topology/test_insertion_codes.py | 2 +- tests/unit/topology/test_links.py | 2 +- tests/unit/topology/test_subset.py | 2 +- torchref/cli/_common.py | 12 +- torchref/cli/refine.py | 22 +- .../experimental/ensemble/ensemble_model.py | 50 +- .../experimental/targets/forcefield_target.py | 6 +- torchref/io/ihm.py | 2 +- torchref/model/context.py | 611 +++++++- torchref/model/model.py | 1287 ++++------------- torchref/model/model_ft.py | 354 +---- torchref/refinement/base_refinement.py | 27 +- 34 files changed, 1242 insertions(+), 1612 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index d02be91b..318346ae 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -186,9 +186,9 @@ Black, 88 columns, `isort` with the black profile. Ruff lint with |---|---| | `base/` | Low-level math and crystallography. `coordinates/` (Cartesian↔fractional), `reciprocal/` (basis, HKL, d-spacing, interpolation, symmetry), `direct_summation/` (F_calc by summation; eager + Triton), `electron_density/` (real-space splatting with CPU/CUDA/MPS kernels, solvent mask, radius policy), `fourier/` (FFT and grids), `scattering/` (form-factor and anomalous tables), `metrics/` (R-factors, binwise scale, loss), `targets/` (the *kernels* behind refinement targets, eager + `triton/`), `french_wilson.py`, `math_torch.py`, `alignment/` | | `io/` | `ReflectionData`, `DatasetCollection`, `FcalcDataset`; MTZ / PDB / CIF / IHM readers and writers; `read_mtz` / `read_pdb` / `read_cif` | -| `model/` | `Model` (refinable atomic parameters), `ModelContext` (the cell, space group, atom table, links and provenance a model is loaded with — `model.cell` / `.spacegroup` / `.pdb` forward to it, the rest is `model.ctx.*`), `ModelFT` (adds F_calc via `SfFFT` or `SfDS`), `MixedModel`, `ModelCollection`, and the parametrizations in `parameter_wrappers.py` / `rigid_xyz.py` that decide what is refinable | +| `model/` | `Model` (refinable atomic parameters), `ModelContext` (the cell, space group, atom table, links, provenance, hydrogen policy and geometry restraints a model is loaded with — `model.cell` / `.spacegroup` / `.pdb` / `.restraints` forward to it, the rest is `model.ctx.*`; `ModelContext.from_atoms` is the one place an atom table is settled and `Model._install_parameters` the one place wrappers are built), `ModelFT` (adds F_calc via `SfFFT` or `SfDS`), `MixedModel`, `ModelCollection`, and the parametrizations in `parameter_wrappers.py` / `rigid_xyz.py` that decide what is refinable | | `refinement/` | Drivers (`Refinement`, `LBFGSRefinement`, `RigidBodyRefinementStep`), `targets/` (`xray/`, `geometry/`, `adp/`, `collection/`, `combined.py`), `weighting/`, `optimizers/` (annealing, Langevin, preconditioned/seeded L-BFGS), `model_error_estimation/` (σ_A, σ_M), `loss_state.py`, `logger.py` | -| `restraints/` | Bonds, angles, torsions, planes, chirals, VDW. Built from the CCP4 Monomer Library, resolved lazily via `get_library_manager()` — importing this package must not trigger a library download | +| `topology/` | The connectivity graph (`Topology`, `AtomGraph`, `ResidueGraph`) and the restraint layer over it (`Restraints`: bonds, angles, torsions, planes, chirals, VDW pair list), hydrogen generation (`hydrogens.py`) and riding frames. Built from the CCP4 Monomer Library, resolved lazily via `monomer.library.get_library_manager()` — importing this package must not trigger a library download. `Restraints` holds no reference to a model: evaluations take the coordinates they score | | `scaling/` | `ScalerBase` (model-independent), `Scaler`, `CollectionScaler`, `SolventModel` (k_sol, B_sol) | | `symmetry/` | `Symmetry` (operations plus everything derived from them), `SpaceGroup` (adds the crystallographic identity and the CCP4 ASU verbs), `Cell`. All dataclasses over `DeviceMixin`, not `nn.Module` — they hold no refinable parameters. Map and reciprocal-grid operators are private, reached through `Symmetry` | | `maps/` | `Map` (2Fo−Fc, Fcalc), `DifferenceMap` | diff --git a/docs/changelog.rst b/docs/changelog.rst index 504adc3c..f7a6f115 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,10 @@ Changelog Unreleased ---------- +- Hydrogens are one policy on the model context, set at construction: ``hydrogens`` (``keep`` / ``add`` / ``strip``) settles the atom table when it loads and ``hydrogen_mode`` (``atoms`` / ``riding``) how hydrogen rows are parametrised; ``riding`` with ``strip`` raises. It replaces ``strip_H``, ``add_hydrogens`` and ``hydrogen_mode="free"`` on ``Model``, ``ModelFT``, ``Refinement``, ``EnsembleModel`` and ``cli._common.load_model`` (checkpoints written with them still restore), and ``torchref.refine`` takes ``--hydrogens`` / ``--hydrogen-mode`` instead of ``--add-hydrogens``. The deprecated ``exclude_H_from_sf`` alias is removed +- Missing water hydrogens are generated at load with ``hydrogens="add"``, never by ``set_hydrogen_mode("riding")``, which now only swaps the coordinate wrapper and leaves the atom table alone; ``hydrogenate()`` no longer takes ``optimize`` (free torsions are always scanned) +- ``Model.select`` keeps the anisotropic ``u`` in its Cholesky parametrization, and ``select``/``copy``/``create_from_state_dict`` and the strip/hydrogenate helpers are shared by ``Model`` and ``ModelFT``, so ``copy`` of a ``ModelFT`` carries ``apply_bijvoet`` +- The per-atom-table queries moved to the context: ``model.ctx.chain_sequences``, ``model.ctx.chain_residues`` (was ``Model.get_chain_residues``), ``model.ctx.occupancy_groups`` and ``model.ctx.register_altlocs`` - ``Restraints`` no longer borrow accessors from the model and live on its context (``model.ctx.restraints``; ``model.restraints`` still builds them on first access). The constructor takes the coordinates to build over (``xyz=``) instead of ``xyz_fn``/``adp_fn``/``vdw_radii_fn``, and every evaluation takes the coordinates or B-factors it scores, e.g. ``bond_deviations(model.xyz())``. ``Model.set_restraints_cif`` is replaced by ``model.ctx.set_cif_path``, and ``Model.bond_deviations``/``angle_deviations``/``torsion_deviations_with_sigmas`` are removed. Without a cell and space group the pair list is searched in an isolated P1 box - ``ModelFT.create_from_state_dict`` restores the restraint dictionary path (``cif_path``) as ``Model`` does, and ``Refinement.create_from_state_dict`` no longer builds a stray ``Restraints`` from the model; the restored model builds its own on first access. - The difference MTZ groups its columns into named datasets -- ``observed``, ``difference``, ``light_model``, ``extrapolated_light``, ``two_moment`` -- with one history line describing each, so ``FWT``/``PHWT`` reads as ``/torchref/extrapolated_light/FWT`` (the extrapolated light-state map ``2*FEXT - Fc``). Labels are unchanged and Coot still auto-opens it diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index eadaa72d..be8b99a5 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -29,8 +29,12 @@ and a ``refinement_history.json`` log. **Key options:** * ``-n`` / ``--n-cycles`` number of macro cycles (default 5) -* ``--add-hydrogens`` generate missing hydrogens on model loading (default off). - Hydrogens already present in the input are retained with or without this flag +* ``--hydrogens {keep,add,strip}`` what loading does with the model's hydrogens: + keep the ones the file has (default), also generate the missing ones (waters + included), or strip them all +* ``--hydrogen-mode {atoms,riding}`` refine hydrogens as ordinary atoms (default) or + let them ride on their parent heavy atoms; ``riding`` with ``--hydrogens strip`` is + an error * ``--hydrogens-in-xray`` / ``--no-hydrogens-in-xray`` include hydrogen atoms in the structure-factor calculation (default on). Off keeps them in the restraints only; the bulk-solvent mask is built from heavy atoms in either case diff --git a/tests/benchmarks/compare_amber_classic.py b/tests/benchmarks/compare_amber_classic.py index fd678d57..1231135d 100644 --- a/tests/benchmarks/compare_amber_classic.py +++ b/tests/benchmarks/compare_amber_classic.py @@ -32,7 +32,7 @@ def _prepare(code: str, output: Path, seed: int) -> dict: source = FILES / "pdb" / f"{code}_af.pdb" original = ( - Model(device="cpu", verbose=0, strip_H=True, add_hydrogens=False) + Model(device="cpu", verbose=0, hydrogens="strip") .load_pdb(str(source)) .strip_altlocs() ) @@ -44,7 +44,7 @@ def _prepare(code: str, output: Path, seed: int) -> dict: fixed_path = output / f"{code}_heavy_completed.pdb" with fixed_path.open("w") as handle: app.PDBFile.writeFile(fixer.topology, fixer.positions, handle, keepIds=True) - fixed = Model(device="cpu", verbose=0, add_hydrogens=False).load_pdb( + fixed = Model(device="cpu", verbose=0).load_pdb( str(fixed_path) ) original_rows = { @@ -70,7 +70,7 @@ def _prepare(code: str, output: Path, seed: int) -> dict: frame.at[i, "tempfactor"] = float(same_residue.tempfactor.mean()) frame.at[i, "occupancy"] = 1.0 assert found == set(original_rows), "Heavy-atom preparation dropped original atoms" - model = original._new_model_from_df(frame, strip_H=False) + model = original._derive(frame, hydrogens="keep") torch.manual_seed(seed) model = model.hydrogenate() model.set_hydrogen_mode("riding") @@ -230,7 +230,6 @@ def _run( data_file=str(FILES / "mtz" / f"{code}.mtz"), device=torch.device("cpu"), verbose=0, - add_hydrogens=False, hydrogens_in_xray=True, target_mode="ml", ) diff --git a/tests/integration/test_cli_hydrogens.py b/tests/integration/test_cli_hydrogens.py index f29b2235..13ccfc37 100644 --- a/tests/integration/test_cli_hydrogens.py +++ b/tests/integration/test_cli_hydrogens.py @@ -1,4 +1,4 @@ -"""Exercise hydrogen opt-in from CLI parsing through deposited-model loading.""" +"""Exercise the hydrogen flags from CLI parsing through deposited-model loading.""" import sys from pathlib import Path @@ -23,7 +23,7 @@ def test_cli_hydrogen_generation_is_opt_in( add_hydrogens: bool, model_format: str, ) -> None: - """The CLI flag generates missing hydrogens for PDB and mmCIF inputs.""" + """``--hydrogens add`` generates missing hydrogens for PDB and mmCIF inputs.""" from torchref.cli import refine loaded = [] @@ -45,14 +45,14 @@ def stop_after_loading(refinement: Refinement) -> None: "0", ] if add_hydrogens: - argv.append("--add-hydrogens") + argv += ["--hydrogens", "add"] monkeypatch.setattr(sys, "argv", argv) with pytest.raises(_ModelLoaded): refine.main() (model,) = loaded - assert model.ctx.add_hydrogens is add_hydrogens + assert model.ctx.hydrogens == ("add" if add_hydrogens else "keep") assert len(model.pdb) > 0 n_hydrogens = int(model.pdb["element"].str.strip().eq("H").sum()) assert (n_hydrogens > 0) is add_hydrogens @@ -83,7 +83,8 @@ def stop_after_loading(refinement: Refinement) -> None: str(tmp_path / "refined"), "-v", "0", - "--add-hydrogens", + "--hydrogens", + "add", "--cif", cif, ] @@ -100,8 +101,35 @@ def stop_after_loading(refinement: Refinement) -> None: @pytest.mark.unit -@pytest.mark.parametrize("add_hydrogens", [False, True]) -def test_empty_refinement_preserves_hydrogen_setting(add_hydrogens: bool) -> None: - """An empty refinement shell forwards the setting to its model too.""" - refinement = LBFGSRefinement(verbose=0, add_hydrogens=add_hydrogens) - assert refinement.model.ctx.add_hydrogens is add_hydrogens +@pytest.mark.parametrize("hydrogens", ["keep", "add", "strip"]) +def test_empty_refinement_preserves_hydrogen_setting(hydrogens: str) -> None: + """An empty refinement shell forwards the policy to its model too.""" + refinement = LBFGSRefinement(verbose=0, hydrogens=hydrogens, hydrogen_mode="atoms") + assert refinement.model.ctx.hydrogens == hydrogens + + +@pytest.mark.integration +def test_cli_refuses_riding_on_stripped_hydrogens( + test_files_dir: Path, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """``--hydrogens strip --hydrogen-mode riding`` fails before anything loads.""" + from torchref.cli import refine + + argv = [ + "torchref.refine", + "-m", + str(test_files_dir / "pdb" / "1DAW.pdb"), + "-sf", + str(test_files_dir / "mtz" / "1DAW.mtz"), + "-o", + str(tmp_path / "refined"), + "-v", + "0", + "--hydrogens", + "strip", + "--hydrogen-mode", + "riding", + ] + monkeypatch.setattr(sys, "argv", argv) + with pytest.raises(ValueError, match="Nothing is left to ride"): + refine.main() diff --git a/tests/integration/test_io_cif.py b/tests/integration/test_io_cif.py index 2bd47fa1..d91df0bc 100644 --- a/tests/integration/test_io_cif.py +++ b/tests/integration/test_io_cif.py @@ -63,7 +63,7 @@ def test_save_and_reload_cif(self, sample_cif_file, tmp_path): # legitimately differ, because ``write_pdb`` does not emit LINK records -- so a # metal-coordinated nitrogen comes back with a free valence and takes a hydrogen # it did not have before. - model2 = Model(add_hydrogens=False) + model2 = Model() model2.load_pdb(str(output_path)) n_atoms2 = model2.xyz().shape[0] diff --git a/tests/integration/test_model_operations.py b/tests/integration/test_model_operations.py index 08b7a736..55980f75 100644 --- a/tests/integration/test_model_operations.py +++ b/tests/integration/test_model_operations.py @@ -294,7 +294,7 @@ def test_model_roundtrip_pdb(self, sample_cif_file, tmp_path): # legitimately differ, because ``write_pdb`` does not emit LINK records -- so a # metal-coordinated nitrogen comes back with a free valence and takes a hydrogen # it did not have before. - model2 = Model(add_hydrogens=False) + model2 = Model() model2.load_pdb(str(output_path)) n_atoms2 = model2.xyz().shape[0] diff --git a/tests/unit/io/test_hkl_convention.py b/tests/unit/io/test_hkl_convention.py index 44636ac4..c04db5ff 100644 --- a/tests/unit/io/test_hkl_convention.py +++ b/tests/unit/io/test_hkl_convention.py @@ -60,7 +60,7 @@ def _model(pdb_dir, data): # strip_H: what is under test is the phase convention, and the absolute check # compares against a gemmi calculation that calls ``remove_hydrogens``. Letting # torchref generate hydrogens would have it computing a different structure. - m = ModelFT(verbose=0, max_res=2.0, strip_H=True) + m = ModelFT(verbose=0, max_res=2.0, hydrogens="strip") m.load_pdb(str(pdb_dir / f"{CODE}.pdb")) m.cell, m.spacegroup = data.cell, data.spacegroup return m diff --git a/tests/unit/model/test_hydrogen_default.py b/tests/unit/model/test_hydrogen_default.py index de2ad847..8ef331bb 100644 --- a/tests/unit/model/test_hydrogen_default.py +++ b/tests/unit/model/test_hydrogen_default.py @@ -10,6 +10,7 @@ import numpy as np import pytest +from torchref.model.context import ModelContext from torchref.model.model import Model from torchref.model.model_ft import ModelFT @@ -39,10 +40,10 @@ def test_default_preserves_deposited_atoms( path = pdb_dir / filename deposited, _, _ = PDBReader(verbose=0).read(str(path))() - def unexpected_generation(self: Model) -> None: + def unexpected_generation(self: ModelContext, dtype) -> None: pytest.fail("Default loading must not generate hydrogens") - monkeypatch.setattr(Model, "_add_missing_hydrogens", unexpected_generation) + monkeypatch.setattr(ModelContext, "_add_missing_hydrogens", unexpected_generation) model = model_class(verbose=0).load_pdb(str(path)) np.testing.assert_array_equal(_elements(model), deposited["element"].str.strip()) @@ -54,10 +55,10 @@ def test_default_cif_load_does_not_generate_hydrogens( ) -> None: """mmCIF loading also leaves missing hydrogens absent by default.""" - def unexpected_generation(self: Model) -> None: + def unexpected_generation(self: ModelContext, dtype) -> None: pytest.fail("Default mmCIF loading must not generate hydrogens") - monkeypatch.setattr(Model, "_add_missing_hydrogens", unexpected_generation) + monkeypatch.setattr(ModelContext, "_add_missing_hydrogens", unexpected_generation) model = Model(verbose=0).load_cif(str(cif_dir / "1DAW.cif")) total, n_h = _counts(model) assert total > 0 @@ -66,16 +67,15 @@ def unexpected_generation(self: Model) -> None: @pytest.mark.unit def test_context_defaults_to_no_hydrogen_generation() -> None: - """A standalone model context leaves hydrogen generation disabled.""" - from torchref.model.context import ModelContext - - assert ModelContext().add_hydrogens is False + """A standalone model context keeps the file's hydrogens as atoms.""" + assert ModelContext().hydrogens == "keep" + assert ModelContext().hydrogen_mode == "atoms" @pytest.mark.unit def test_a_file_without_hydrogens_gets_them(pdb_dir): """1DAW ships none, so every hydrogen here is generated.""" - model = Model(verbose=0, add_hydrogens=True) + model = Model(verbose=0, hydrogens="add") model.load_pdb(str(pdb_dir / "1DAW.pdb")) total, n_h = _counts(model) @@ -94,11 +94,11 @@ def test_a_partially_hydrogenated_file_is_topped_up(pdb_dir): names and the model lacks -- so a file that already has some still gets the rest. A does-the-table-contain-any test would have left this structure as deposited. """ - kept = Model(verbose=0, add_hydrogens=False) + kept = Model(verbose=0) kept.load_pdb(str(pdb_dir / "1AK5_with_H.pdb")) _, n_kept = _counts(kept) - topped = Model(verbose=0, add_hydrogens=True) + topped = Model(verbose=0, hydrogens="add") topped.load_pdb(str(pdb_dir / "1AK5_with_H.pdb")) _, n_topped = _counts(topped) @@ -109,24 +109,24 @@ def test_a_partially_hydrogenated_file_is_topped_up(pdb_dir): @pytest.mark.unit -def test_strip_H_still_removes_everything(pdb_dir): +def test_strip_removes_everything(pdb_dir): """The opt-out is unaffected: no hydrogen survives, generated or deposited.""" for name in ("1DAW.pdb", "7L84.pdb"): - model = Model(verbose=0, strip_H=True, add_hydrogens=True) + model = Model(verbose=0, hydrogens="strip") model.load_pdb(str(pdb_dir / name)) _, n_h = _counts(model) - assert n_h == 0, f"{name} kept {n_h} hydrogens under strip_H" + assert n_h == 0, f"{name} kept {n_h} hydrogens under hydrogens='strip'" @pytest.mark.unit -def test_add_hydrogens_false_keeps_the_file_as_it_is(pdb_dir): +def test_keep_keeps_the_file_as_it_is(pdb_dir): """Generation off, stripping off: exactly what the reader produced.""" - model = Model(verbose=0, add_hydrogens=False) + model = Model(verbose=0) model.load_pdb(str(pdb_dir / "7L84.pdb")) total, n_h = _counts(model) assert n_h > 0, "7L84 ships hydrogens, so they should have been kept" - generated = Model(verbose=0, add_hydrogens=True) + generated = Model(verbose=0, hydrogens="add") generated.load_pdb(str(pdb_dir / "7L84.pdb")) assert _counts(generated)[0] >= total @@ -140,7 +140,7 @@ def test_per_atom_buffers_are_rebuilt_for_the_new_atom_set(pdb_dir): place left the van der Waals radii at the heavy-atom count while the pair list indexed the full set, and the non-bonded build raised ``IndexError``. """ - model = Model(verbose=0, add_hydrogens=True) + model = Model(verbose=0, hydrogens="add") model.load_pdb(str(pdb_dir / "1DAW.pdb")) n_atoms = len(model.pdb) @@ -157,7 +157,7 @@ def test_per_atom_buffers_are_rebuilt_for_the_new_atom_set(pdb_dir): @pytest.mark.unit def test_restraints_build_over_the_hydrogenated_model(pdb_dir): """Restraints cover the hydrogens, and each carries exactly one bond.""" - model = Model(verbose=0, add_hydrogens=True) + model = Model(verbose=0, hydrogens="add") model.load_pdb(str(pdb_dir / "1DAW.pdb")) restraints = model.restraints @@ -184,14 +184,14 @@ def test_riding_hydrogens_are_not_placed_when_real_ones_exist(pdb_dir): because the riding builder counts bonded neighbours by distance while the generator reads them off the bond graph. """ - model = Model(verbose=0, add_hydrogens=True) + model = Model(verbose=0, hydrogens="add") model.load_pdb(str(pdb_dir / "1DAW.pdb")) restraints = model.restraints assert restraints.h_topo is not None assert restraints.h_topo.n_hydrogens == 0 - stripped = Model(verbose=0, strip_H=True) + stripped = Model(verbose=0, hydrogens="strip") stripped.load_pdb(str(pdb_dir / "1DAW.pdb")) assert ( stripped.restraints.h_topo.n_hydrogens > 0 @@ -225,7 +225,7 @@ def test_generation_reads_the_cif_given_at_construction(pdb_dir, renamed_glu_cif hydrogens are the ones the restraints know: a hydrogen generated from one dictionary and restrained by another has no bond edge at all. """ - model = Model(verbose=0, add_hydrogens=True, cif_path=renamed_glu_cif) + model = Model(verbose=0, hydrogens="add", cif_path=renamed_glu_cif) model.load_pdb(str(pdb_dir / "1DAW.pdb")) assert model.ctx.cif_path == renamed_glu_cif names = _glu_hydrogen_names(model) @@ -241,7 +241,7 @@ def test_generation_reads_the_cif_given_at_construction(pdb_dir, renamed_glu_cif @pytest.mark.unit def test_derived_models_keep_the_restraint_cif(pdb_dir, renamed_glu_cif): """hydrogenate, strip_hydrogens and select all carry the dictionary along.""" - model = Model(verbose=0, add_hydrogens=False, cif_path=renamed_glu_cif) + model = Model(verbose=0, cif_path=renamed_glu_cif) model.load_pdb(str(pdb_dir / "1DAW.pdb")) hydrogenated = model.hydrogenate() @@ -274,7 +274,7 @@ def test_load_model_registers_the_cif_before_loading(pdb_dir, renamed_glu_cif): from torchref.cli._common import load_model model = load_model( - str(pdb_dir / "1DAW.pdb"), verbose=0, cif=renamed_glu_cif, add_hydrogens=True + str(pdb_dir / "1DAW.pdb"), verbose=0, cif=renamed_glu_cif, hydrogens="add" ) assert model.ctx.cif_path == renamed_glu_cif assert RENAMED_GLU_H <= _glu_hydrogen_names(model) @@ -292,7 +292,7 @@ def test_generation_reads_every_compound_of_a_multi_block_cif(pdb_dir, test_file blocks = [l for l in cif.read_text().splitlines() if l.startswith("data_comp_")] assert len(blocks) == 3, blocks # comp_list + GLU + ASP: the fixture is really multi-block - model = Model(verbose=0, add_hydrogens=True, cif_path=str(cif)) + model = Model(verbose=0, hydrogens="add", cif_path=str(cif)) model.load_pdb(str(pdb_dir / "1DAW.pdb")) pdb = model.pdb is_h = pdb["element"].astype(str).str.strip() == "H" diff --git a/tests/unit/model/test_hydrogen_mode.py b/tests/unit/model/test_hydrogen_mode.py index 52710726..bfeea83f 100644 --- a/tests/unit/model/test_hydrogen_mode.py +++ b/tests/unit/model/test_hydrogen_mode.py @@ -1,9 +1,11 @@ -"""Switching a loaded model between riding and free hydrogens. - -The switch replaces the coordinate wrapper and nothing else: coordinates are -unchanged, the heavy-atom refinable set carries over, hydrogens leave or rejoin the -refinable set, the restraints keep reading the live wrapper, and the mode survives -copies, selections and state dicts. +"""The hydrogen policy, and switching a loaded model between riding and atom hydrogens. + +``hydrogens`` (keep / add / strip) settles the atom table at load and +``hydrogen_mode`` (atoms / riding) how its hydrogen rows are parametrised; strip with +riding is refused. The switch replaces the coordinate wrapper and nothing else: +coordinates are unchanged, the heavy-atom refinable set carries over, hydrogens leave or +rejoin the refinable set, the restraints are untouched, and the mode survives copies, +selections and state dicts. """ import pytest @@ -16,7 +18,7 @@ @pytest.fixture def free_model(pdb_dir): - model = Model(verbose=0, add_hydrogens=True) + model = Model(verbose=0, hydrogens="add") model.load_pdb(str(pdb_dir / "1DAW.pdb")) return model @@ -39,11 +41,11 @@ def test_switch_to_riding_keeps_coordinates_and_drops_hydrogen_parameters(free_m @pytest.mark.unit -def test_switch_back_to_free_restores_per_atom_wrapper(free_model): +def test_switch_back_to_atoms_restores_per_atom_wrapper(free_model): model = free_model model.set_hydrogen_mode("riding") - model.set_hydrogen_mode("free") - assert model.hydrogen_mode == "free" + model.set_hydrogen_mode("atoms") + assert model.hydrogen_mode == "atoms" assert not isinstance(model.xyz, RidingXYZTensor) assert model.parameters_of_types(("xyz",))[0].shape[0] == len(model.pdb) @@ -92,7 +94,7 @@ def test_mode_survives_copy_select_and_shake(free_model): @pytest.mark.parametrize("model_class", [Model, ModelFT]) def test_riding_mode_round_trips_through_state_dict(pdb_dir, model_class): kwargs = {"max_res": 3.0} if model_class is ModelFT else {} - model = model_class(verbose=0, add_hydrogens=True, **kwargs) + model = model_class(verbose=0, hydrogens="add", **kwargs) model.load_pdb(str(pdb_dir / "1DAW.pdb")) model.set_hydrogen_mode("riding") with torch.no_grad(): @@ -106,6 +108,39 @@ def test_riding_mode_round_trips_through_state_dict(pdb_dir, model_class): @pytest.mark.unit -def test_none_mode_is_refused(free_model): - with pytest.raises(ValueError): - free_model.set_hydrogen_mode("none") +@pytest.mark.parametrize("mode", ["none", "free"]) +def test_unknown_modes_are_refused(free_model, mode): + with pytest.raises(ValueError, match="hydrogen_mode must be one of"): + free_model.set_hydrogen_mode(mode) + + +@pytest.mark.unit +@pytest.mark.parametrize("hydrogens", ["keep", "add", "strip"]) +@pytest.mark.parametrize("hydrogen_mode", ["atoms", "riding"]) +def test_policy_matrix(pdb_dir, hydrogens, hydrogen_mode): + """Each valid pair loads with the atom set and wrapper it names; strip+riding raises.""" + if hydrogens == "strip" and hydrogen_mode == "riding": + with pytest.raises(ValueError, match="Nothing is left to ride"): + Model(verbose=0, hydrogens=hydrogens, hydrogen_mode=hydrogen_mode) + return + + model = Model(verbose=0, hydrogens=hydrogens, hydrogen_mode=hydrogen_mode) + model.load_pdb(str(pdb_dir / "1AK5_with_H.pdb")) + deposited = Model(verbose=0).load_pdb(str(pdb_dir / "1AK5_with_H.pdb")) + + n_h = _n_h(model) + if hydrogens == "strip": + assert n_h == 0 + elif hydrogens == "keep": + assert n_h == _n_h(deposited) + else: + assert n_h > _n_h(deposited) + assert isinstance(model.xyz, RidingXYZTensor) is (hydrogen_mode == "riding") + assert model.xyz().shape[0] == len(model.pdb) + + +@pytest.mark.unit +def test_riding_is_refused_on_a_stripped_model(pdb_dir): + model = Model(verbose=0, hydrogens="strip").load_pdb(str(pdb_dir / "1DAW.pdb")) + with pytest.raises(ValueError, match="Nothing is left to ride"): + model.set_hydrogen_mode("riding") diff --git a/tests/unit/model/test_hydrogens_in_xray.py b/tests/unit/model/test_hydrogens_in_xray.py index ff709a7d..14d262a6 100644 --- a/tests/unit/model/test_hydrogens_in_xray.py +++ b/tests/unit/model/test_hydrogens_in_xray.py @@ -55,7 +55,7 @@ def test_off_matches_a_stripped_model(pdb_dir): """Excluding hydrogens from Fcalc equals computing Fcalc without them.""" full = ModelFT(verbose=0, max_res=2.5, hydrogens_in_xray=False) full.load_pdb(str(pdb_dir / "7L84.pdb")) - heavy = ModelFT(verbose=0, max_res=2.5, strip_H=True) + heavy = ModelFT(verbose=0, max_res=2.5, hydrogens="strip") heavy.load_pdb(str(pdb_dir / "7L84.pdb")) grid = torch.arange(-3, 4) hkl = torch.cartesian_prod(grid, grid, grid) @@ -89,16 +89,6 @@ def test_setting_round_trips_through_state_dict(pdb_dir, model_class): assert legacy.hydrogens_in_xray is True -@pytest.mark.unit -def test_deprecated_alias_is_inverted_and_warns(): - model = Model(verbose=0) - with pytest.warns(DeprecationWarning): - model.exclude_H_from_sf = True - assert model.hydrogens_in_xray is False - with pytest.warns(DeprecationWarning): - assert model.exclude_H_from_sf is True - - @pytest.mark.unit def test_solvent_mask_ignores_hydrogens(pdb_dir): """The bulk-solvent mask is the same with and without hydrogen rows.""" @@ -106,7 +96,7 @@ def test_solvent_mask_ignores_hydrogens(pdb_dir): full = ModelFT(verbose=0, max_res=2.5) full.load_pdb(str(pdb_dir / "7L84.pdb")) - heavy = ModelFT(verbose=0, max_res=2.5, strip_H=True) + heavy = ModelFT(verbose=0, max_res=2.5, hydrogens="strip") heavy.load_pdb(str(pdb_dir / "7L84.pdb")) assert _n_h(full) > 0 and _n_h(heavy) == 0 mask_full = SolventModel(full, verbose=0).get_solvent_mask() diff --git a/tests/unit/model/test_model.py b/tests/unit/model/test_model.py index da9a28bb..b556d492 100644 --- a/tests/unit/model/test_model.py +++ b/tests/unit/model/test_model.py @@ -53,14 +53,14 @@ def test_model_custom_dtype(self): assert model.dtype_float == torch.float64 @pytest.mark.unit - def test_model_strip_h_default(self): - """Hydrogen stripping and generation are both opt-in.""" + def test_model_hydrogen_default(self): + """Hydrogen stripping and generation are both opt-in; hydrogens are atoms.""" from torchref.model.model import Model model = Model() - assert model.ctx.strip_H is False - assert model.ctx.add_hydrogens is False + assert model.ctx.hydrogens == "keep" + assert model.ctx.hydrogen_mode == "atoms" @pytest.mark.unit def test_model_bool_uninitialized(self): @@ -145,7 +145,7 @@ def test_dropped_rows_leave_a_positional_index(pdb_dir, tmp_path): sg = src.spacegroup model = Model(verbose=0) - model.load(lambda: (df, cell, sg), add_hydrogens=False) + model.load(lambda: (df, cell, sg)) assert len(model.pdb) == n_before - len(victims) idx = model.pdb["index"].to_numpy() diff --git a/tests/unit/model/test_riding_orientations.py b/tests/unit/model/test_riding_orientations.py index 633af89f..8d7b4827 100644 --- a/tests/unit/model/test_riding_orientations.py +++ b/tests/unit/model/test_riding_orientations.py @@ -15,7 +15,7 @@ def oriented_model(pdb_dir): """A deposited protein with explicitly generated protein and water hydrogens.""" with torch.random.fork_rng(): torch.manual_seed(19) - model = Model(verbose=0, device="cpu", add_hydrogens=True) + model = Model(verbose=0, device="cpu", hydrogens="add") model.load_pdb(str(pdb_dir / "1DAW.pdb")) model.set_hydrogen_mode("riding") return model diff --git a/tests/unit/model/test_riding_water_completion.py b/tests/unit/model/test_riding_water_completion.py index 7b89d7c5..0f7df251 100644 --- a/tests/unit/model/test_riding_water_completion.py +++ b/tests/unit/model/test_riding_water_completion.py @@ -1,4 +1,9 @@ -"""Riding-mode water completion respects the model's hydrogen generation setting.""" +"""Water hydrogens are completed at load, and only when generation is asked for. + +``hydrogens="add"`` gives every HOH its two hydrogens when the table is settled, so a +riding model starts with a water rotation per water. Switching modes afterwards never +changes the atom table: with ``hydrogens="keep"`` an oxygen-only water stays oxygen-only. +""" import pytest import torch @@ -6,158 +11,81 @@ from torchref import Model, ModelFT +def _n_waters(model): + pdb = model.pdb + return int((pdb.resname.str.strip().eq("HOH") & pdb.element.str.strip().eq("O")).sum()) + + @pytest.fixture(scope="module") def heavy_model(pdb_dir): - """Deposited 1DAW loaded without hydrogen generation.""" - return Model(device="cpu", verbose=0, strip_H=True, add_hydrogens=False).load_pdb( + """Deposited 1DAW with every hydrogen stripped.""" + return Model(device="cpu", verbose=0, hydrogens="strip").load_pdb( str(pdb_dir / "1DAW.pdb") ) @pytest.mark.unit @pytest.mark.parametrize("model_class", [Model, ModelFT]) -def test_riding_completes_only_waters_and_preserves_live_atoms( - heavy_model, model_class -): - """Water completion preserves current coordinates, ADPs, selections and links.""" - model = model_class(device="cpu", verbose=0, add_hydrogens=False, strip_H=False) - model.load( - lambda: (heavy_model.pdb.copy(), heavy_model.cell.data, heavy_model.spacegroup) - ) - with torch.no_grad(): - model.xyz.refinable_params.add_(0.25) - adp = model.adp().detach().clone() - mask = torch.arange(len(model.pdb)) % 2 == 0 - model.xyz.update_refinable_mask(mask) - model.adp.update_refinable_mask(mask) - model.occupancy.freeze_all() - before = model.xyz().detach().clone() - links = model.ctx.links - hkl = torch.tensor([[1, 0, 0], [0, 1, 0], [1, 1, 1]]) - if isinstance(model, ModelFT): - model(hkl) - model.ctx.add_hydrogens = True - returned = model.set_hydrogen_mode("riding") - assert returned is model - is_h = torch.as_tensor(model.pdb.element.str.strip().eq("H").to_numpy()) - is_water = model.pdb.resname.str.strip().eq("HOH") - assert int(is_h.sum()) == 2 * int( - heavy_model.pdb.resname.str.strip().eq("HOH").sum() - ) - assert is_water[is_h.numpy()].all() - assert torch.allclose(model.xyz()[~is_h], before, atol=1e-5) - assert torch.allclose(model.adp()[~is_h], adp) - assert torch.equal(model.xyz.full_refinable_mask[~is_h], mask) - assert torch.equal(model.adp.refinable_mask[~is_h], mask) - assert not model.occupancy.get_refinable_atoms().any() - assert model.ctx.links is links - assert model.xyz.rotations.shape[0] == int(is_h.sum()) // 2 +def test_add_with_riding_completes_every_water(heavy_model, model_class): + """Each water gets two riding hydrogens and one rotation, and the model round-trips.""" + table = heavy_model.pdb[heavy_model.pdb.resname.str.strip().eq("HOH")].copy() + model = model_class(device="cpu", verbose=0, hydrogens="add", hydrogen_mode="riding") + model.load(lambda: (table, heavy_model.cell.data, heavy_model.spacegroup)) + + is_h = model.pdb.element.str.strip().eq("H").to_numpy() + assert int(is_h.sum()) == 2 * len(table) + assert model.xyz.n_hydrogens == int(is_h.sum()) + assert model.xyz.rotations.shape == (len(table), 3) assert model.restraints.topology.n_atoms == model.xyz.shape[0] + heavy_rows = torch.as_tensor(~is_h) + expected = torch.tensor(table[["x", "y", "z"]].values, dtype=model.xyz().dtype) + assert torch.allclose(model.xyz()[heavy_rows], expected, atol=1e-5) if isinstance(model, ModelFT): + hkl = torch.tensor([[1, 0, 0], [0, 1, 0], [1, 1, 1]]) assert torch.isfinite(model(hkl)).all() + restored = model_class.create_from_state_dict(model.state_dict(), device="cpu") + assert restored.hydrogen_mode == "riding" assert torch.allclose(restored.xyz(), model.xyz(), atol=1e-5) @pytest.mark.unit -def test_partial_water_is_completed_without_moving_existing_hydrogen(heavy_model): - """A deposited oxygen with one supplied hydrogen receives just its missing partner.""" +def test_partial_water_is_completed_without_moving_its_hydrogen(heavy_model): + """A water that arrives with one hydrogen receives just its missing partner.""" water = ( heavy_model.pdb[heavy_model.pdb.resname.str.strip().eq("HOH")].iloc[:1].copy() ) - model = heavy_model._new_model_from_df(water, strip_H=False) - model.ctx.add_hydrogens = True - model.set_hydrogen_mode("riding") - model.update_pdb() - partial = model._new_model_from_df(model.pdb.iloc[:2].copy(), strip_H=False) - before = partial.xyz().detach().clone() - partial.set_hydrogen_mode("riding") - assert len(partial.pdb) == 2 - assert torch.allclose(partial.xyz(), before, atol=1e-5) - partial.ctx.add_hydrogens = True - partial.set_hydrogen_mode("riding") - assert len(partial.pdb) == 3 - assert torch.allclose(partial.xyz()[:2], before, atol=1e-5) - with torch.no_grad(): - partial.xyz.rotations.refinable_params.fill_(0.3) - wrapper = partial.xyz - coords = partial.xyz().detach().clone() - partial.set_hydrogen_mode("riding") - assert partial.xyz is wrapper - assert torch.equal(partial.xyz(), coords) - partial.set_hydrogen_mode("free").set_hydrogen_mode("riding") - assert len(partial.pdb) == 3 - assert torch.allclose(partial.xyz(), coords, atol=1e-5) + complete = heavy_model._derive(water, hydrogens="add") + complete.update_pdb() + partial_table = complete.pdb.iloc[:2].copy() + kept = heavy_model._derive(partial_table, hydrogens="keep") + assert len(kept.pdb) == 2 + before = kept.xyz().detach().clone() -@pytest.mark.unit -def test_supplied_frames_are_remapped_when_waters_are_completed(heavy_model): - """Explicit frames for the original atom table coexist with generated water frames.""" - water = ( - heavy_model.pdb[heavy_model.pdb.resname.str.strip().eq("HOH")].iloc[:2].copy() - ) - model = heavy_model._new_model_from_df(water, strip_H=False) - frames = model.hydrogen_frames() - model.ctx.add_hydrogens = True - model.set_hydrogen_mode("riding", frames=frames) - assert len(model.pdb) == 6 - assert model.xyz.n_hydrogens == 4 - assert model.xyz.rotations.shape == (2, 3) - - -@pytest.mark.unit -def test_water_completion_preserves_adp_field(heavy_model): - """Adding water hydrogens retains the node parametrization and its parameters.""" - model = heavy_model._new_model_from_df(heavy_model.pdb.copy(), strip_H=False) - model.set_adp_mode("field", n_nodes=8, k_neighbors=4) - field = model.adp - values = field().detach().clone() - parameters = field.refinable_params - model.ctx.add_hydrogens = True - model.set_hydrogen_mode("riding") - heavy = torch.as_tensor(~model.pdb.element.str.strip().eq("H").to_numpy()) - assert model.adp is field - assert field.refinable_params is parameters - assert torch.allclose(field()[heavy], values, atol=1e-5) - assert field().shape == (len(model.pdb),) - - -@pytest.mark.integration -@pytest.mark.parametrize("generate", [False, True]) -def test_refinement_targets_follow_completed_atom_table(pdb_dir, mtz_dir, generate): - """Refinement adds water H and refreshes targets only when generation is enabled.""" - from torchref.refinement.base_refinement import Refinement + topped = heavy_model._derive(partial_table, hydrogens="add", hydrogen_mode="riding") + assert len(topped.pdb) == 3 + assert torch.allclose(topped.xyz()[:2], before, atol=1e-5) - refinement = Refinement( - pdb=str(pdb_dir / "1DAW.pdb"), - data_file=str(mtz_dir / "1DAW.mtz"), - device="cpu", - verbose=0, - max_res=3.0, - add_hydrogens=False, - ) - previous = refinement.adp_target - n_atoms = len(refinement.model.pdb) - refinement.model.ctx.add_hydrogens = generate - refinement.set_hydrogen_mode("riding") - assert (refinement.adp_target is not previous) == generate - assert (len(refinement.model.pdb) > n_atoms) == generate - geometry = refinement.geometry_target() - assert torch.isfinite(geometry) - geometry.backward() - if generate: - gradient = refinement.model.xyz.rotations.refinable_params.grad - assert gradient is not None and torch.isfinite(gradient).all() + with torch.no_grad(): + topped.xyz.rotations.refinable_params.fill_(0.3) + wrapper = topped.xyz + coords = topped.xyz().detach().clone() + topped.set_hydrogen_mode("riding") + assert topped.xyz is wrapper + topped.set_hydrogen_mode("atoms").set_hydrogen_mode("riding") + assert len(topped.pdb) == 3 + assert torch.allclose(topped.xyz(), coords, atol=1e-5) @pytest.mark.unit @pytest.mark.parametrize("model_class", [Model, ModelFT]) @pytest.mark.parametrize("explicit_frames", [False, True]) -def test_disabled_generation_keeps_atom_table( +def test_switching_to_riding_never_changes_the_atom_table( heavy_model, model_class, explicit_frames ): - """Riding mode leaves oxygen-only waters untouched when add_hydrogens is False.""" - model = model_class(device="cpu", verbose=0, add_hydrogens=False) + """With ``hydrogens="keep"``, oxygen-only waters stay oxygen-only under riding.""" + model = model_class(device="cpu", verbose=0) model.load( lambda: (heavy_model.pdb.copy(), heavy_model.cell.data, heavy_model.spacegroup) ) @@ -173,12 +101,33 @@ def test_disabled_generation_keeps_atom_table( assert model.xyz.rotations.shape == (0, 3) -@pytest.mark.unit -def test_stripping_prevents_water_completion(heavy_model): - """The stripping preference takes precedence even when generation is enabled.""" - model = heavy_model._new_model_from_df(heavy_model.pdb.copy(), strip_H=True) - model.ctx.add_hydrogens = True - n_atoms = len(model.pdb) - model.set_hydrogen_mode("riding") - assert len(model.pdb) == n_atoms - assert model.xyz.n_hydrogens == 0 +@pytest.mark.integration +@pytest.mark.parametrize("hydrogens", ["keep", "add"]) +def test_refinement_targets_see_the_riding_waters(pdb_dir, mtz_dir, hydrogens): + """Water rotations reach the geometry gradient when the waters were completed.""" + from torchref.refinement.base_refinement import Refinement + + refinement = Refinement( + pdb=str(pdb_dir / "1DAW.pdb"), + data_file=str(mtz_dir / "1DAW.mtz"), + device="cpu", + verbose=0, + max_res=3.0, + hydrogens=hydrogens, + ) + n_atoms = len(refinement.model.pdb) + adp_target = refinement.adp_target + refinement.set_hydrogen_mode("riding") + assert len(refinement.model.pdb) == n_atoms + assert refinement.adp_target is adp_target + + geometry = refinement.geometry_target() + assert torch.isfinite(geometry) + geometry.backward() + rotations = refinement.model.xyz.rotations + assert rotations.shape[0] == ( + _n_waters(refinement.model) if hydrogens == "add" else 0 + ) + if hydrogens == "add": + gradient = rotations.refinable_params.grad + assert gradient is not None and torch.isfinite(gradient).all() diff --git a/tests/unit/model/test_riding_xyz.py b/tests/unit/model/test_riding_xyz.py index ba5b9414..6b4ba9b6 100644 --- a/tests/unit/model/test_riding_xyz.py +++ b/tests/unit/model/test_riding_xyz.py @@ -20,7 +20,7 @@ @pytest.fixture(scope="module") def hydrogenated(pdb_dir): """1DAW with generated hydrogens, its frames, and its full coordinate table.""" - model = Model(verbose=0, add_hydrogens=True) + model = Model(verbose=0, hydrogens="add") model.load_pdb(str(pdb_dir / "1DAW.pdb")) frames = hydrogen_frames(model.restraints.topology) return model, frames, model.xyz().detach() diff --git a/tests/unit/model/test_sf_grid_key.py b/tests/unit/model/test_sf_grid_key.py index c43b71df..97cb3621 100644 --- a/tests/unit/model/test_sf_grid_key.py +++ b/tests/unit/model/test_sf_grid_key.py @@ -49,7 +49,7 @@ def test_one_engine_and_one_spacegroup_per_load(pdb_path, monkeypatch, strip_H): engines = _count_calls(monkeypatch, model_ft_module.SfFFT, "__init__") spacegroups = _count_calls(monkeypatch, SpaceGroup, "__init__") - model = ModelFT(max_res=2.5, verbose=0, device="cpu", strip_H=strip_H) + model = ModelFT(max_res=2.5, verbose=0, device="cpu", hydrogens="strip" if strip_H else "keep") model.load_pdb(pdb_path) assert engines["n"] == 1 assert spacegroups["n"] == 1 diff --git a/tests/unit/monomer/test_link_modifications.py b/tests/unit/monomer/test_link_modifications.py index 63b73854..95ffc459 100644 --- a/tests/unit/monomer/test_link_modifications.py +++ b/tests/unit/monomer/test_link_modifications.py @@ -233,7 +233,7 @@ def _built(pdb_path, strip_H=True): """ from torchref import Model - model = Model(verbose=0, strip_H=strip_H, add_hydrogens=False) + model = Model(verbose=0, hydrogens="strip" if strip_H else "keep") model.load_pdb(str(pdb_path)) return model, model.restraints.restraints diff --git a/tests/unit/refinement/test_amber_target.py b/tests/unit/refinement/test_amber_target.py index 1b88a35a..a05d62bf 100644 --- a/tests/unit/refinement/test_amber_target.py +++ b/tests/unit/refinement/test_amber_target.py @@ -20,7 +20,7 @@ @pytest.fixture(scope="module") def protein(): """A deposited, hydrogenated protein with methyl orientation parameters.""" - model = Model(verbose=0, device="cpu", add_hydrogens=True).load_pdb(TEST_PDB) + model = Model(verbose=0, device="cpu", hydrogens="add").load_pdb(TEST_PDB) model = model.strip_altlocs() model.set_hydrogen_mode("riding") return model @@ -118,7 +118,7 @@ def test_permuted_positions_and_gradients(target, protein): def test_missing_hydrogens_are_not_added(): """Disabled generation leaves a heavy-only model untouched on AMBER rejection.""" model = ( - Model(verbose=0, device="cpu", strip_H=True, add_hydrogens=False) + Model(verbose=0, device="cpu", hydrogens="strip") .load_pdb(TEST_PDB) .strip_altlocs() ) @@ -133,7 +133,7 @@ def test_missing_hydrogens_are_not_added(): @pytest.fixture(scope="module") def water_target(pdb_dir): """Two nearby deposited waters with TorchRef-generated rotating hydrogens.""" - model = Model(verbose=0, device="cpu", add_hydrogens=False).load_pdb( + model = Model(verbose=0, device="cpu").load_pdb( str(pdb_dir / "1DAW.pdb") ) waters = model.pdb[model.pdb.resname.str.strip().eq("HOH")].copy() @@ -141,11 +141,11 @@ def water_target(pdb_dir): distances = np.linalg.norm(coords[:, None] - coords[None, :], axis=-1) np.fill_diagonal(distances, np.inf) i, j = np.unravel_index(np.argmin(distances), distances.shape) - model = model._new_model_from_df(waters.iloc[sorted([i, j])].copy(), strip_H=False) - model.ctx.add_hydrogens = True with torch.random.fork_rng(): torch.manual_seed(42) - model.set_hydrogen_mode("riding") + model = model._derive( + waters.iloc[sorted([i, j])].copy(), hydrogens="add", hydrogen_mode="riding" + ) return AmberTarget(model=model, normalize_by_atoms=False) @@ -306,7 +306,7 @@ def test_bridge_follows_coordinate_device(water_target, any_device): def test_torchref_hydrogenation_prepares_compatible_protein(): """TorchRef can prepare all model hydrogens before AMBER construction.""" model = ( - Model(verbose=0, device="cpu", strip_H=True, add_hydrogens=False) + Model(verbose=0, device="cpu", hydrogens="strip") .load_pdb(TEST_PDB) .strip_altlocs() .hydrogenate() @@ -355,7 +355,7 @@ def test_partial_terminal_hydrogens_preserve_h1_alias(protein): & (pdb.icode == first.icode) ) missing = residue & pdb.name.str.strip().isin(["H2", "H3"]) - partial = protein._new_model_from_df(pdb.loc[~missing].copy(), strip_H=False) + partial = protein._derive(pdb.loc[~missing].copy(), hydrogens="keep") prepared = partial.hydrogenate() first_residue = prepared.pdb[ (prepared.pdb.chainid == first.chainid) & (prepared.pdb.resseq == first.resseq) diff --git a/tests/unit/refinement/test_forcefield_target.py b/tests/unit/refinement/test_forcefield_target.py index df5db224..8d185ad4 100644 --- a/tests/unit/refinement/test_forcefield_target.py +++ b/tests/unit/refinement/test_forcefield_target.py @@ -302,7 +302,7 @@ def test_forward_with_real_pdb(self): from torchref.model import Model from torchref.experimental.targets import ForceFieldTarget # Load model WITH hydrogens - model = Model(strip_H=False) + model = Model() model.load_pdb(str(TEST_PDB_WITH_H)) # Check we have hydrogens @@ -333,7 +333,7 @@ def test_gradient_flow(self): from torchref.model import Model from torchref.experimental.targets import ForceFieldTarget # Load model WITH hydrogens - model = Model(strip_H=False) + model = Model() model.load_pdb(str(TEST_PDB_WITH_H)) # Create target @@ -360,7 +360,7 @@ def test_stats_with_real_model(self): """Test stats() with real model.""" from torchref.model import Model from torchref.experimental.targets import ForceFieldTarget - model = Model(strip_H=False) + model = Model() model.load_pdb(str(TEST_PDB_WITH_H)) target = ForceFieldTarget( diff --git a/tests/unit/topology/test_energy_types.py b/tests/unit/topology/test_energy_types.py index d296a038..68f3c788 100644 --- a/tests/unit/topology/test_energy_types.py +++ b/tests/unit/topology/test_energy_types.py @@ -19,7 +19,7 @@ @pytest.fixture(scope="module") def heavy_1daw(pdb_dir): - model = Model(verbose=0, strip_H=True, add_hydrogens=False) + model = Model(verbose=0, hydrogens="strip") model.load_pdb(str(pdb_dir / "1DAW.pdb")) return model @@ -87,7 +87,7 @@ def test_template_h_count_and_implicit_hydrogens(heavy_1daw): @pytest.mark.unit def test_hydrogenated_model_completes_charged_amines(pdb_dir): """Explicit hydrogen generation fills the polymer's template hydrogen counts.""" - model = Model(verbose=0, add_hydrogens=True) + model = Model(verbose=0, hydrogens="add") model.load_pdb(str(pdb_dir / "1DAW.pdb")) atoms = model.restraints.topology.atoms missing = atoms.implicit_h_count().cpu().numpy() diff --git a/tests/unit/topology/test_hydrogen_frames.py b/tests/unit/topology/test_hydrogen_frames.py index 268f6a52..f3a4e5d9 100644 --- a/tests/unit/topology/test_hydrogen_frames.py +++ b/tests/unit/topology/test_hydrogen_frames.py @@ -28,7 +28,7 @@ @pytest.fixture(scope="module") def heavy_and_plan(pdb_dir): """Heavy-only 1DAW with its hydrogen plan.""" - model = Model(verbose=0, strip_H=True, add_hydrogens=False) + model = Model(verbose=0, hydrogens="strip") model.load_pdb(str(pdb_dir / "1DAW.pdb")) restraints = model.restraints xyz = model.xyz().detach() @@ -44,7 +44,7 @@ def hydrogenated(heavy_and_plan): augmented, old_to_new, plan_to_new = augment_atom_table_with_maps( model.pdb, plan, restraints.topology ) - full = Model(verbose=0, strip_H=False, add_hydrogens=False) + full = Model(verbose=0) cell, spacegroup = model.cell.data.cpu().numpy(), model.spacegroup def reader(): diff --git a/tests/unit/topology/test_hydrogens.py b/tests/unit/topology/test_hydrogens.py index 59bbe307..b0c1e2a5 100644 --- a/tests/unit/topology/test_hydrogens.py +++ b/tests/unit/topology/test_hydrogens.py @@ -32,7 +32,7 @@ def _build(code): if code not in cache: # add_hydrogens=False: these tests exercise generation itself, so the model # has to arrive without the hydrogens the loader would otherwise add. - model = Model(verbose=0, add_hydrogens=False, strip_H=True) + model = Model(verbose=0, hydrogens="strip") model.load_pdb(str(pdb_dir / f"{code}.pdb")) model.ctx.set_cif_path(None) restraints = model.restraints @@ -274,14 +274,14 @@ def test_waters_receive_two_hydrogens_with_initial_orientations(built): @pytest.mark.unit def test_hydrogenate_returns_a_consistent_model(pdb_dir): """The end-to-end path yields a model whose tensors, table and restraints agree.""" - model = Model(verbose=0, add_hydrogens=False, strip_H=True) + model = Model(verbose=0, hydrogens="strip") model.load_pdb(str(pdb_dir / "7L84.pdb")) model.ctx.set_cif_path(None) n_heavy = len(model.pdb) hydrogenated = model.hydrogenate(verbose=0) - assert hydrogenated.ctx.strip_H is False + assert hydrogenated.ctx.hydrogens == "add" assert len(hydrogenated.pdb) > n_heavy assert hydrogenated.xyz().shape[0] == len(hydrogenated.pdb) assert hydrogenated.adp().shape[0] == len(hydrogenated.pdb) @@ -308,7 +308,7 @@ def test_hydrogenate_returns_a_consistent_model(pdb_dir): @pytest.mark.unit def test_strip_H_removes_deposited_hydrogens(pdb_dir): """The opt-out drops the hydrogens the file carries, as it always did.""" - model = Model(verbose=0, strip_H=True) + model = Model(verbose=0, hydrogens="strip") model.load_pdb(str(pdb_dir / "1AK5_with_H.pdb")) elements = model.pdb["element"].astype(str).str.strip().values assert not (elements == "H").any() @@ -385,7 +385,7 @@ def test_acetyl_cap_carbon_gets_no_hydrogen(tmp_path): pytest.skip("ACE not in the monomer library") path = tmp_path / "ace_met.pdb" path.write_text(_ACE_MET) - model = Model(verbose=0, add_hydrogens=False, strip_H=True) + model = Model(verbose=0, hydrogens="strip") model.load_pdb(str(path)) model.ctx.set_cif_path(None) restraints = model.restraints diff --git a/tests/unit/topology/test_insertion_codes.py b/tests/unit/topology/test_insertion_codes.py index 6dafac8f..3e52a8ea 100644 --- a/tests/unit/topology/test_insertion_codes.py +++ b/tests/unit/topology/test_insertion_codes.py @@ -60,7 +60,7 @@ def inserted(pdb_dir, tmp_path_factory): path = tmp_path_factory.mktemp("icode") / f"{BASE}_icode.pdb" expected = _rewrite_with_insertion_codes(pdb_dir / f"{BASE}.pdb", path) - model = Model(verbose=0, strip_H=True, add_hydrogens=False) + model = Model(verbose=0, hydrogens="strip") model.load_pdb(str(path)) model.ctx.set_cif_path(None) restraints = model.restraints diff --git a/tests/unit/topology/test_links.py b/tests/unit/topology/test_links.py index f1ee1da6..33f46297 100644 --- a/tests/unit/topology/test_links.py +++ b/tests/unit/topology/test_links.py @@ -50,7 +50,7 @@ def _row(model, chain, resseq, name, altloc=""): def _load(path): - model = Model(verbose=0, add_hydrogens=False, strip_H=True) + model = Model(verbose=0, hydrogens="strip") model.load_pdb(str(path)) if str(path).endswith(".pdb") else model.load_cif(str(path)) model.ctx.set_cif_path(None) return model diff --git a/tests/unit/topology/test_subset.py b/tests/unit/topology/test_subset.py index 001df453..bd2edae4 100644 --- a/tests/unit/topology/test_subset.py +++ b/tests/unit/topology/test_subset.py @@ -20,7 +20,7 @@ @pytest.fixture(scope="module") def topology(pdb_dir): """A topology with altlocs, disulfides, peptide links and hydrogens.""" - model = Model(verbose=0, add_hydrogens=False, strip_H=True) + model = Model(verbose=0, hydrogens="strip") model.load_pdb(str(pdb_dir / "7L84.pdb")) model.ctx.set_cif_path(None) return model.restraints.topology diff --git a/torchref/cli/_common.py b/torchref/cli/_common.py index 135e62a8..3882a43f 100644 --- a/torchref/cli/_common.py +++ b/torchref/cli/_common.py @@ -711,7 +711,8 @@ def load_model( device: Union[str, "torch.device", None] = None, verbose: int = 0, cif: Optional[Union[str, List[str]]] = None, - add_hydrogens: bool = False, + hydrogens: str = "keep", + hydrogen_mode: str = "atoms", hydrogens_in_xray: bool = True, ) -> "ModelFT": """Load a model from PDB or CIF, auto-detected by file extension. @@ -729,8 +730,10 @@ def load_model( cif : str or list of str, optional CIF restraint file(s), registered on the model before it loads so that hydrogen generation and the restraints read the same dictionary. - add_hydrogens : bool, optional - Generate missing hydrogens on load. Default False. + hydrogens : {"keep", "add", "strip"}, optional + What loading does with the file's hydrogens. Default ``"keep"``. + hydrogen_mode : {"atoms", "riding"}, optional + Hydrogens as refinable atoms or riding on their parents. Default ``"atoms"``. hydrogens_in_xray : bool, optional Whether hydrogens contribute to the structure factors. Default True. @@ -747,7 +750,8 @@ def load_model( device=device, verbose=verbose, cif_path=cif, - add_hydrogens=add_hydrogens, + hydrogens=hydrogens, + hydrogen_mode=hydrogen_mode, hydrogens_in_xray=hydrogens_in_xray, ) suffix = Path(path).suffix.lower() diff --git a/torchref/cli/refine.py b/torchref/cli/refine.py index 7aed867c..d93a3c41 100644 --- a/torchref/cli/refine.py +++ b/torchref/cli/refine.py @@ -115,10 +115,19 @@ def main(): refine_group = parser.add_argument_group("Refinement") add_n_cycles_arg(refine_group) refine_group.add_argument( - "--add-hydrogens", - action="store_true", - help="Generate missing hydrogens when loading the model (default: off). " - "Hydrogens already present in the input are retained either way.", + "--hydrogens", + choices=["keep", "add", "strip"], + default="keep", + help="What to do with the model's hydrogens on load: keep the ones the file " + "has (default), also generate the missing ones, or strip them all.", + ) + refine_group.add_argument( + "--hydrogen-mode", + dest="hydrogen_mode", + choices=["atoms", "riding"], + default="atoms", + help="Refine hydrogens as ordinary atoms (default) or let them ride on their " + "parent heavy atoms. 'riding' cannot be combined with --hydrogens strip.", ) refine_group.add_argument( "--hydrogens-in-xray", @@ -262,7 +271,7 @@ def main(): print(f"Refinement mode: {args.mode}") print(f"X-ray target: {args.xray_mode}") print(f"Refinement cycles: {args.n_cycles}") - print(f"Add hydrogens: {'on' if args.add_hydrogens else 'off'}") + print(f"Hydrogens: {args.hydrogens} ({args.hydrogen_mode})") print(f"Hydrogens in Fcalc: {'on' if args.hydrogens_in_xray else 'off'}") if args.with_rigid_body: print(f"Rigid-body step: on (iterations/cutoff = {args.rigid_body_iter})") @@ -320,7 +329,8 @@ def main(): reflections_per_adp_parameter=args.reflections_per_adp_parameter, aniso_selection=args.anisotropic_selection, wavelength=args.wavelength, - add_hydrogens=args.add_hydrogens, + hydrogens=args.hydrogens, + hydrogen_mode=args.hydrogen_mode, hydrogens_in_xray=args.hydrogens_in_xray, ) diff --git a/torchref/experimental/ensemble/ensemble_model.py b/torchref/experimental/ensemble/ensemble_model.py index de52ce11..fed3bf93 100644 --- a/torchref/experimental/ensemble/ensemble_model.py +++ b/torchref/experimental/ensemble/ensemble_model.py @@ -291,7 +291,7 @@ def build_single_copy_model(ensemble, atom_idx=None, verbose: int = 0): df = ensemble._pdb_single if atom_idx is not None: df = df.iloc[np.asarray(atom_idx)] - chem = Model(verbose=verbose, strip_H=False, device=ensemble.device) + chem = Model(verbose=verbose, hydrogens="keep", device=ensemble.device) chem.load( _SyntheticPDBReader( df.reset_index(drop=True).copy(), @@ -322,8 +322,10 @@ class EnsembleModel(ModelFT): Verbosity. device : torch.device Computation device. - strip_H : bool - Whether to strip hydrogens (inherited). + hydrogens : {"keep", "strip"} + Hydrogen policy on load (inherited). ``"add"`` is refused: the atom set is + the replicated single copy the factories build, and ``_finalize_ensemble`` + reshapes by ``n_atoms_per_member``, which generated hydrogens would break. max_res : float FFT grid target resolution (inherited). @@ -339,18 +341,20 @@ def __init__( dtype_float=None, verbose: int = 1, device=None, - strip_H: bool = True, - # An ensemble's atom set is the replicated single copy its factories build, and - # _finalize_ensemble reshapes by n_atoms_per_member, so generating hydrogens on - # load would invalidate that. Off by default here, unlike on the base class. - add_hydrogens: bool = False, + hydrogens: str = "keep", max_res: float = 1.0, gridsize: Optional[Tuple[int, int, int]] = None, wavelength: float = 1.0, anomalous_threshold: float = 0.5, + apply_bijvoet: bool = False, cif_path=None, hydrogens_in_xray: bool = True, ): + if hydrogens == "add": + raise ValueError( + "EnsembleModel cannot generate hydrogens: its atom set is the " + "replicated single copy; hydrogenate the input first." + ) if dtype_float is None: dtype_float = get_float_dtype() if device is None: @@ -359,12 +363,12 @@ def __init__( dtype_float=dtype_float, verbose=verbose, device=device, - strip_H=strip_H, - add_hydrogens=add_hydrogens, + hydrogens=hydrogens, max_res=max_res, gridsize=gridsize, wavelength=wavelength, anomalous_threshold=anomalous_threshold, + apply_bijvoet=apply_bijvoet, cif_path=cif_path, hydrogens_in_xray=hydrogens_in_xray, ) @@ -397,7 +401,7 @@ def from_single( seed: Optional[int] = None, verbose: int = 1, device=None, - strip_H: bool = True, + hydrogens: str = "strip", max_res: float = 1.0, n_max: Optional[int] = None, **modelft_kwargs, @@ -430,8 +434,8 @@ def from_single( Verbosity. device : torch.device, optional Computation device. - strip_H : bool - Strip hydrogens before replication (default True). + hydrogens : {"strip", "keep"} + Strip hydrogens before replication (default) or keep the file's. max_res : float FFT grid target resolution (Å), forwarded to ``ModelFT``. n_max : int, optional @@ -444,7 +448,7 @@ def from_single( """ reader = pdb_io.PDBReader(verbose=verbose).read(pdb_path) df, cell, spacegroup = reader() - if strip_H: + if hydrogens == "strip": df = df.loc[df["element"].astype(str).str.strip() != "H"].reset_index(drop=True) # Strip alternate conformations: the ensemble IS the disorder model, # so per-residue altlocs would double-count atoms in OpenMM topology, @@ -460,10 +464,9 @@ def from_single( ) model = cls( - verbose=verbose, device=device, strip_H=False, # already stripped - # The replicated table is the atom set; _finalize_ensemble reshapes by - # n_atoms_per_member, so generating hydrogens here would invalidate it. - add_hydrogens=False, + verbose=verbose, + device=device, + hydrogens="keep", # already stripped, if asked max_res=max_res, **modelft_kwargs, ) @@ -485,7 +488,7 @@ def from_multimodel_pdb( seed: Optional[int] = None, verbose: int = 1, device=None, - strip_H: bool = True, + hydrogens: str = "strip", max_res: float = 1.0, n_max: Optional[int] = None, **modelft_kwargs, @@ -504,7 +507,7 @@ def from_multimodel_pdb( ready for bifurcation to reactivate). Default ``n_max = n_members`` (no spare slots; bifurcation can only reuse slots freed by deaths). """ - models = _parse_multi_model_pdb(pdb_path, strip_H=strip_H) + models = _parse_multi_model_pdb(pdb_path, strip_H=hydrogens == "strip") if len(models) == 0: raise ValueError(f"No usable atomic models parsed from {pdb_path}") if n_members is None: @@ -549,10 +552,9 @@ def from_multimodel_pdb( replicated = pd.concat(pieces, ignore_index=True) model = cls( - verbose=verbose, device=device, strip_H=False, - # See from_single: the replicated table is the atom set, and - # _finalize_ensemble reshapes by n_atoms_per_member. - add_hydrogens=False, + verbose=verbose, + device=device, + hydrogens="keep", # already stripped, if asked max_res=max_res, **modelft_kwargs, ) diff --git a/torchref/experimental/targets/forcefield_target.py b/torchref/experimental/targets/forcefield_target.py index 09f0e1de..2e6106d5 100644 --- a/torchref/experimental/targets/forcefield_target.py +++ b/torchref/experimental/targets/forcefield_target.py @@ -44,7 +44,7 @@ class ForceFieldTarget(ModelTarget): ---------- model : Model, optional Reference to the Model object. Should include hydrogens for accurate - energies (load with ``strip_H=False``); a hydrogen-less model is not + energies (load with ``hydrogens="keep"`` or ``"add"``); a hydrogen-less model is not rejected, only flagged via a warning when ``verbose > 0``. model_path : str, optional Path to TorchMD-Net checkpoint file (.ckpt). @@ -63,7 +63,7 @@ class ForceFieldTarget(ModelTarget): >>> from torchref.experimental.targets import ForceFieldTarget >>> >>> # Load model WITH hydrogens - >>> model = Model(strip_H=False) + >>> model = Model(hydrogens="add") >>> model.load_pdb('structure_with_H.pdb') >>> >>> # Create force field target @@ -152,7 +152,7 @@ def _validate_hydrogens(self) -> None: warnings.warn( "Model appears to have no hydrogen atoms. " "TorchMD-Net typically requires all-atom structures. " - "Load with Model(strip_H=False) if hydrogens are needed.", + "Load with Model(hydrogens='add') if hydrogens are needed.", UserWarning ) diff --git a/torchref/io/ihm.py b/torchref/io/ihm.py index 06d4de65..97a8c5f2 100644 --- a/torchref/io/ihm.py +++ b/torchref/io/ihm.py @@ -731,7 +731,7 @@ def write(self, filepath: str) -> None: if mc.n_base_models > 0: model0 = mc.base_models[0] - for chain_id, seq_str in model0.chain_sequences: + for chain_id, seq_str in model0.ctx.chain_sequences: seq = [] for char in seq_str: if char == "?": diff --git a/torchref/model/context.py b/torchref/model/context.py index 5b0da878..2f402335 100644 --- a/torchref/model/context.py +++ b/torchref/model/context.py @@ -1,10 +1,16 @@ """The information half of a :class:`~torchref.model.model.Model`. :class:`ModelContext` holds what a model *is loaded from* and *sits in* -- the unit -cell, the space group, the atom table, the link records and the provenance -- as -opposed to what is being refined, which stays on the model as parameter wrappers and -per-atom buffers. The geometry restraints belong here too: they are fixed by the atom -set and the dictionaries, and are evaluated against coordinates the caller passes in. +cell, the space group, the atom table, the link records, the provenance and the +hydrogen policy -- as opposed to what is being refined, which stays on the model as +parameter wrappers and per-atom buffers. The geometry restraints belong here too: they +are fixed by the atom set and the dictionaries, and are evaluated against coordinates +the caller passes in. + +:meth:`ModelContext.from_atoms` is the one place an atom table is settled: hydrogens +stripped or generated, unusable rows dropped, the crystal built. Every way of making a +model -- loading a file, selecting, stripping, hydrogenating, restoring a state dict -- +produces a context first and only then installs parameter wrappers over it. Splitting it out means the crystallographic context can be passed to code that needs only that (structure-factor engines, scalers, most targets) without handing over the @@ -16,17 +22,118 @@ from __future__ import annotations from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple + +import torch from torchref.utils.device_mixin import DeviceMixin if TYPE_CHECKING: import pandas - import torch from torchref.symmetry import Cell, SpaceGroup from torchref.topology.restraints import Restraints +#: Three-letter residue code to one-letter code, modified residues included. +THREE_TO_ONE = { + "ALA": "A", + "ARG": "R", + "ASN": "N", + "ASP": "D", + "CYS": "C", + "GLN": "Q", + "GLU": "E", + "GLY": "G", + "HIS": "H", + "ILE": "I", + "LEU": "L", + "LYS": "K", + "MET": "M", + "PHE": "F", + "PRO": "P", + "SER": "S", + "THR": "T", + "TRP": "W", + "TYR": "Y", + "VAL": "V", + "SEC": "U", + "PYL": "O", + # Common modified residues + "MSE": "M", + "CSE": "C", + "SEP": "S", + "TPO": "T", + "PTR": "Y", +} + +#: What to do with the hydrogens of the input atom table. +HYDROGEN_SOURCES = ("keep", "add", "strip") + +#: How hydrogen rows are parametrised. +HYDROGEN_MODES = ("atoms", "riding") + +#: Fields a derived context inherits from the one it was derived from. +_SETTINGS = ( + "cif_path", + "verbose", + "hydrogens", + "hydrogen_mode", + "hydrogens_in_xray", + "input_file", +) + + +def check_hydrogen_policy(hydrogens: str, hydrogen_mode: str) -> None: + """Raise ``ValueError`` unless ``(hydrogens, hydrogen_mode)`` is a valid pair. + + Parameters + ---------- + hydrogens : {"keep", "add", "strip"} + hydrogen_mode : {"atoms", "riding"} + """ + if hydrogens not in HYDROGEN_SOURCES: + raise ValueError( + f"hydrogens must be one of {HYDROGEN_SOURCES}, got {hydrogens!r}" + ) + if hydrogen_mode not in HYDROGEN_MODES: + raise ValueError( + f"hydrogen_mode must be one of {HYDROGEN_MODES}, got {hydrogen_mode!r}" + ) + if hydrogens == "strip" and hydrogen_mode == "riding": + raise ValueError( + "hydrogen_mode='riding' with hydrogens='strip': you threw the hydrogens " + "overboard and then asked them to ride. Nothing is left to ride -- use " + "hydrogens='keep' or hydrogens='add'." + ) + + +def own_spacegroup(value, dtype: torch.dtype, device) -> Optional["SpaceGroup"]: + """A space group owned by the caller, on ``device`` and in ``dtype``. + + An incoming :class:`~torchref.symmetry.SpaceGroup` is copied rather than shared, + because ``.to()`` moves in place and would otherwise relocate the caller's object. + + Parameters + ---------- + value : SpaceGroup, gemmi.SpaceGroup, str, int or None + Anything :class:`~torchref.symmetry.SpaceGroup` accepts. + dtype : torch.dtype + device : torch.device + + Returns + ------- + SpaceGroup or None + """ + from torchref.symmetry import SpaceGroup + + if value is None: + return None + if isinstance(value, SpaceGroup): + return value.copy().to(device=device, dtype=dtype) + # SpaceGroup falls back to the global default device otherwise, which would plant + # accelerator-resident matrices on a CPU-pinned model. + return SpaceGroup(value, dtype=dtype, device=device) + @dataclass(eq=False, repr=False) class ModelContext(DeviceMixin): @@ -44,26 +151,27 @@ class ModelContext(DeviceMixin): links : list or None Link records from the reader, used to build inter-residue restraints. altloc_pairs : list - Index groups of alternative conformations, rebuilt by - ``Model.register_alternative_conformations``. + Index groups of alternative conformations, one tuple of index tensors per + residue with more than one conformation; rebuilt by :meth:`register_altlocs`. input_file : str or None Path the structure was loaded from. - cif_path : str or None - Restraint dictionary path, if one was set. + cif_path : str or list of str or None + Restraint dictionary path(s). Change it with :meth:`set_cif_path`, which drops + restraints built over the old dictionaries. verbose : int, default 1 Verbosity level. - strip_H : bool, default True - Whether hydrogens were stripped on load. + hydrogens : {"keep", "add", "strip"}, default "keep" + What :meth:`from_atoms` does with the input's hydrogens: keep what the file + has, additionally generate the ones the monomer templates name and the file + lacks (waters included), or remove them all. + hydrogen_mode : {"atoms", "riding"}, default "atoms" + How hydrogen rows are parametrised: as ordinary refinable atoms, or riding on + their parent heavy atoms (rebuilt from them every forward, not refined). + ``"riding"`` with ``hydrogens="strip"`` raises ``ValueError``: there is + nothing left to ride. hydrogens_in_xray : bool, default True Whether hydrogens enter the structure-factor calculation. Restraints and the non-bonded term see them either way; the bulk-solvent mask never does. - add_hydrogens : bool, default False - Generate hydrogens on load for residues that arrive without them. Ignored when - ``strip_H`` is set, which removes them again. - hydrogen_mode : str, default "free" - How hydrogen rows are parametrised: ``"riding"`` (positions derived from the - parent heavy atoms each forward, not refined), ``"free"`` (ordinary refinable - atoms) or ``"none"`` (the table holds no hydrogens). initialized : bool, default False Whether a structure has been loaded. ``if model:`` tests this. restraints : Restraints or None @@ -88,15 +196,171 @@ class ModelContext(DeviceMixin): links: Optional[List[Any]] = None altloc_pairs: List[Any] = field(default_factory=list) input_file: Optional[str] = None - cif_path: Optional[str] = None + cif_path: Optional[Any] = None verbose: int = 1 - strip_H: bool = True + hydrogens: str = "keep" + hydrogen_mode: str = "atoms" hydrogens_in_xray: bool = True - add_hydrogens: bool = False - hydrogen_mode: str = "free" initialized: bool = False restraints: Optional["Restraints"] = None + def __post_init__(self) -> None: + check_hydrogen_policy(self.hydrogens, self.hydrogen_mode) + + # ------------------------------------------------------------------ + # Building + # ------------------------------------------------------------------ + + def settings(self) -> Dict[str, Any]: + """The policy fields a context derived from this one inherits. + + Returns + ------- + dict + ``cif_path``, ``verbose``, ``hydrogens``, ``hydrogen_mode``, + ``hydrogens_in_xray`` and ``input_file``, ready to pass to + :meth:`from_atoms`. + """ + return {name: getattr(self, name) for name in _SETTINGS} + + @classmethod + def from_atoms( + cls, + pdb: "pandas.DataFrame", + cell, + spacegroup, + *, + dtype: torch.dtype, + device, + links=None, + **settings, + ) -> "ModelContext": + """Settle an atom table and build the context around it. + + In order: strip hydrogens (``hydrogens="strip"``); drop rows without + coordinates, B-factor or occupancy and renumber the ``index`` column; build the + cell and space group; generate missing hydrogens (``hydrogens="add"``); record + the alternative conformations. + + Parameters + ---------- + pdb : pandas.DataFrame + Atom table as read. Not modified. + cell : Cell or array-like + ``[a, b, c, alpha, beta, gamma]`` in Å and degrees, or a Cell to copy. + spacegroup : SpaceGroup, gemmi.SpaceGroup, str or int + dtype : torch.dtype + Float dtype of the cell and space-group tensors. + device : torch.device + Where the cell and space-group tensors live. + links : list, optional + Link records from the reader. + **settings + Any of the fields :meth:`settings` returns. + + Returns + ------- + ModelContext + Initialized, with no restraints built yet unless hydrogen generation + needed them (in which case they were built over the table *before* + generation and discarded). + + Raises + ------ + ValueError + For an invalid hydrogen policy, including ``strip`` with ``riding``. + """ + from torchref.symmetry import Cell + + ctx = cls(links=links, **settings) + if ctx.hydrogens == "strip": + pdb = pdb.loc[pdb["element"].str.strip() != "H"] + # Renumber before deriving ``index``: every consumer uses it to address length-N + # per-atom tensors positionally, so a gapped index from the drop sends them past + # the end (roughly one PDB-REDO entry in six loses rows here). + pdb = pdb.dropna(subset=["x", "y", "z", "tempfactor", "occupancy"]) + ctx.pdb = cls._renumbered(pdb) + ctx.cell = Cell( + cell.data if isinstance(cell, Cell) else cell, dtype=dtype, device=device + ) + ctx.spacegroup = own_spacegroup(spacegroup, dtype, device) + if ctx.hydrogens == "add": + ctx._add_missing_hydrogens(dtype) + ctx.register_altlocs() + ctx.initialized = True + return ctx + + def derive(self, pdb: "pandas.DataFrame", **overrides) -> "ModelContext": + """A new context over ``pdb`` in this one's crystal, with its settings. + + Parameters + ---------- + pdb : pandas.DataFrame + The new atom table; see :meth:`from_atoms`. + **overrides + Settings to change, e.g. ``hydrogens="strip"``. + + Returns + ------- + ModelContext + """ + settings = {**self.settings(), **overrides} + return ModelContext.from_atoms( + pdb, + self.cell, + self.spacegroup, + dtype=self.cell.dtype, + device=self.cell.device, + links=self.links, + **settings, + ) + + @staticmethod + def _renumbered(pdb: "pandas.DataFrame") -> "pandas.DataFrame": + pdb = pdb.reset_index(drop=True) + pdb["index"] = pdb.index.to_numpy(dtype=int) + return pdb + + def _add_missing_hydrogens(self, dtype: torch.dtype) -> None: + """Top up the hydrogens the atom table is missing. + + Per parent, not per file: a structure deposited with some hydrogens gets the + rest, because the plan only ever proposes a hydrogen the template names and the + table does not have (1AK5 arrives with 675 of roughly 2500). + + Costs a restraint build without the pair list, over the table as loaded, + because the plan needs its topology; it is discarded afterwards. + """ + from torchref.topology.hydrogens import ( + augment_atom_table, + optimise_free_torsions, + plan_hydrogens, + ) + + xyz = torch.tensor(self.pdb[["x", "y", "z"]].values, dtype=dtype) + restraints = self.build_restraints(xyz, nonbonded=False, verbose=0) + plan = plan_hydrogens( + restraints.topology, restraints.cif_dict, xyz, verbose=self.verbose + ) + if self.verbose > 0 and restraints.missing_residues: + print( + "No restraint dictionary for " + f"{sorted(restraints.missing_residues)}: not hydrogenated. Pass one " + "with cif_path / --cif." + ) + if plan.n_hydrogens == 0: + return + optimise_free_torsions(plan, restraints.topology, xyz) + self.pdb = self._renumbered( + augment_atom_table(self.pdb, plan, restraints.topology) + ) + if self.verbose > 0: + print(f"Generated {plan.n_hydrogens} hydrogens") + + # ------------------------------------------------------------------ + # Restraints + # ------------------------------------------------------------------ + def set_cif_path(self, cif_path) -> None: """Replace the restraint dictionary path and drop restraints built over the old one. @@ -109,7 +373,7 @@ def set_cif_path(self, cif_path) -> None: self.restraints = None def build_restraints( - self, xyz: "torch.Tensor", *, nonbonded: bool = True, verbose=None + self, xyz: torch.Tensor, *, nonbonded: bool = True, verbose=None ) -> "Restraints": """Build restraints over the atom table and store them on :attr:`restraints`. @@ -144,13 +408,220 @@ def build_restraints( self.restraints = restraints return restraints + # ------------------------------------------------------------------ + # Atom-table queries + # ------------------------------------------------------------------ + + def occupancy_groups(self, initial_occ): + """``(sharing_groups, altloc_groups, refinable_mask)`` for an + :class:`~torchref.model.parameter_wrappers.OccupancyTensor` over this table. + + Altloc conformations share one collapsed index each; other residues share + one only when their occupancies agree to within 0.01, and an occupancy is + refinable only if it differs from 1.0 by more than that same deadband. + """ + n_atoms = len(initial_occ) + altloc_groups = [] + refinable_mask = torch.zeros(n_atoms, dtype=torch.bool) + + sharing_groups_tensor = torch.arange(n_atoms, dtype=torch.long) # dtype-ok: arange atom indices (sharing groups); index requires long + collapsed_idx = 0 + + # First pass: altlocs. ALL atoms of one conformation must share a collapsed + # index whatever their individual occupancies, or the sum-to-1 + # normalization in OccupancyTensor.forward() acts on the wrong group. + pdb_with_altlocs = self.pdb[self.pdb["altloc"] != ""] + altloc_residues = set() + + if len(pdb_with_altlocs) > 0: + grouped_by_residue = pdb_with_altlocs.groupby( + ["resname", "resseq", "chainid"] + ) + + for (resname, resseq, chainid), group in grouped_by_residue: + unique_altlocs = sorted(group["altloc"].unique()) + + if len(unique_altlocs) > 1: + altloc_residues.add((resname, resseq, chainid)) + conformation_atom_lists = [] + + for altloc in unique_altlocs: + altloc_atoms = group[group["altloc"] == altloc] + indices = altloc_atoms["index"].tolist() + + sharing_groups_tensor[indices] = collapsed_idx + + for idx in indices: + if abs(initial_occ[idx].item() - 1.0) > 0.01: + refinable_mask[idx] = True + + conformation_atom_lists.append(indices) + collapsed_idx += 1 + + altloc_groups.append(tuple(conformation_atom_lists)) + + # Second pass: non-altloc residues, sharing by occupancy similarity. + grouped = self.pdb.groupby(["resname", "resseq", "chainid", "altloc"]) + + for (resname, resseq, chainid, altloc), group in grouped: + if (resname, resseq, chainid) in altloc_residues: + continue + + indices = group["index"].tolist() + + if len(indices) == 0: + continue + + residue_occs = initial_occ[indices] + + occ_min = residue_occs.min().item() + occ_max = residue_occs.max().item() + occ_mean = residue_occs.mean().item() + + if (occ_max - occ_min) <= 0.01: + sharing_groups_tensor[indices] = collapsed_idx + collapsed_idx += 1 + + if abs(occ_mean - 1.0) > 0.01: + for idx in indices: + refinable_mask[idx] = True + else: + # Occupancies disagree within the residue: keep atoms independent. + for idx in indices: + if abs(initial_occ[idx].item() - 1.0) > 0.01: + refinable_mask[idx] = True + + # Compact to contiguous indices 0..n_collapsed-1. + unique_indices = torch.unique(sharing_groups_tensor, sorted=True) + index_map = torch.zeros(n_atoms, dtype=torch.long) # dtype-ok: index_map atom-index remap; indexing requires long + for new_idx, old_idx in enumerate(unique_indices): + mask = sharing_groups_tensor == old_idx + sharing_groups_tensor[mask] = new_idx + + n_collapsed = len(unique_indices) + + if self.verbose > 1: + n_groups = n_collapsed + n_independent = n_atoms - n_collapsed + n_refinable = refinable_mask.sum().item() + n_altloc_groups = len(altloc_groups) + + print("\nOccupancy Setup:") + print(f" Total atoms: {n_atoms}") + print(f" Collapsed indices: {n_collapsed}") + print(f" Alternative conformation groups: {n_altloc_groups}") + print(f" Refinable atoms: {n_refinable}") + print(f" Compression ratio: {n_atoms / n_collapsed:.2f}x") + + return sharing_groups_tensor, altloc_groups, refinable_mask + + def register_altlocs(self) -> None: + """ + Rebuild ``self.altloc_pairs`` from the ``altloc`` column. + + One tuple per residue that has multiple conformations, holding one + index tensor per conformation (in sorted altloc order), e.g. + ``[(tensor([100, 101]), tensor([110, 111])), ...]``. Overwrites any + previous content, so call it after the atom numbering changes. + """ + self.altloc_pairs = [] + + pdb_with_altlocs = self.pdb[self.pdb["altloc"] != ""] + + if len(pdb_with_altlocs) == 0: + return + + grouped = pdb_with_altlocs.groupby(["resname", "resseq", "chainid"]) + + for (resname, resseq, chainid), group in grouped: + unique_altlocs = sorted(group["altloc"].unique()) + + # A lone altloc label is not an alternative conformation. + if len(unique_altlocs) > 1: + conformation_tensors = [] + for altloc in unique_altlocs: + altloc_atoms = group[group["altloc"] == altloc] + indices = torch.tensor( + altloc_atoms["index"].tolist(), dtype=torch.long # dtype-ok: altloc atom indices; indexing requires long + ) + conformation_tensors.append(indices) + + self.altloc_pairs.append(tuple(conformation_tensors)) + + @property + def chain_sequences(self) -> List[Tuple[str, str]]: + """Per-chain one-letter sequences, ``[(chain_id, sequence), ...]``. + + HETATM records are excluded, numbering gaps become ``?`` and unrecognized + residues ``X``. + """ + if self.pdb is None: + return [] + + atom_df = self.pdb[self.pdb["ATOM"] == "ATOM"] + result = [] + + for chain in atom_df["chainid"].unique(): + chain_df = atom_df[atom_df["chainid"] == chain] + residues = chain_df.drop_duplicates(subset=["resseq", "icode"]).sort_values( + "resseq" + ) + resseqs = residues["resseq"].values + resnames = residues["resname"].values + + seq_chars = [] + for i, (rseq, rname) in enumerate(zip(resseqs, resnames)): + if i > 0: + gap = int(rseq) - int(resseqs[i - 1]) - 1 + if gap > 0: + seq_chars.extend(["?"] * gap) + code = THREE_TO_ONE.get(str(rname).strip(), "X") + seq_chars.append(code) + + result.append((str(chain), "".join(seq_chars))) + + return result + + @property + def chain_residues(self) -> List[Tuple[str, List[str]]]: + """ + Per-chain residue names as 3-letter codes (for IHM/CIF writing). + + Excludes HETATM records. Unlike :attr:`chain_sequences`, returns + the raw 3-letter codes without gap filling. + + Returns + ------- + list of (str, list of str) + Ordered list of ``(chain_id, [resname, ...])``. + """ + if self.pdb is None: + return [] + + atom_df = self.pdb[self.pdb["ATOM"] == "ATOM"] + result = [] + + for chain in atom_df["chainid"].unique(): + chain_df = atom_df[atom_df["chainid"] == chain] + residues = chain_df.drop_duplicates(subset=["resseq", "icode"]).sort_values( + "resseq" + ) + resnames = [str(r).strip() for r in residues["resname"].values] + result.append((str(chain), resnames)) + + return result + + # ------------------------------------------------------------------ + # Copying and persistence + # ------------------------------------------------------------------ + def copy(self) -> "ModelContext": """An independent copy. The atom table is deep-copied, the cell and space group are cloned and built - restraints are copied, so nothing is shared with the original. Cloning the space group matters now that - it is a mutable dataclass: sharing the reference would let an edit through one - model's context reach every model that was copied from it. + restraints are copied, so nothing is shared with the original. Cloning the + space group matters because it is a mutable dataclass: sharing the reference + would let an edit through one model's context reach every model copied from it. Returns ------- @@ -167,14 +638,8 @@ def copy(self) -> "ModelContext": altloc_pairs=[ tuple(t.clone() for t in group) for group in self.altloc_pairs ], - input_file=self.input_file, - cif_path=self.cif_path, - verbose=self.verbose, - strip_H=self.strip_H, - hydrogens_in_xray=self.hydrogens_in_xray, - add_hydrogens=self.add_hydrogens, - hydrogen_mode=self.hydrogen_mode, initialized=self.initialized, + **self.settings(), ) if self.restraints is not None: restraints = self.restraints.copy() @@ -186,6 +651,81 @@ def copy(self) -> "ModelContext": duplicate.restraints = restraints return duplicate + def state(self) -> Dict[str, Any]: + """What :meth:`from_state` needs, as picklable entries for a model state dict. + + Returns + ------- + dict + The atom table, the cell as a CPU tensor, the space group as its extended + Hermann-Mauguin symbol (``gemmi.SpaceGroup`` is not picklable), the + altloc groups and the settings. Restraints are not saved; they rebuild. + """ + return { + "pdb": self.pdb.copy() if self.pdb is not None else None, + "cell": self.cell.data.cpu() if self.cell is not None else None, + "spacegroup": self.spacegroup.xhm if self.spacegroup else None, + "initialized": self.initialized, + "cif_path": self.cif_path, + "altloc_pairs": self.altloc_pairs, + "hydrogens": self.hydrogens, + "hydrogen_mode": self.hydrogen_mode, + "hydrogens_in_xray": self.hydrogens_in_xray, + } + + @classmethod + def from_state( + cls, state: Dict[str, Any], *, dtype: torch.dtype, device, verbose: int = 1 + ) -> "ModelContext": + """Rebuild a context from the entries :meth:`state` wrote. + + Consumes them: every key read is popped off ``state``, so what remains is for + ``load_state_dict``. The atom table is taken as saved -- the hydrogen policy is + recorded, not re-applied. Checkpoints that predate the policy are mapped: + ``strip_H`` becomes ``hydrogens``, ``"free"`` becomes ``"atoms"``, and a saved + riding wrapper (an ``xyz.h_row`` entry) implies ``"riding"``. + + Parameters + ---------- + state : dict + A model state dict. + dtype : torch.dtype + Float dtype for the cell and space group. + device : torch.device + verbose : int, default 1 + + Returns + ------- + ModelContext + """ + from torchref.symmetry import Cell + + hydrogens = state.pop("hydrogens", None) + strip_h = state.pop("strip_H", True) + state.pop("add_hydrogens", None) + if hydrogens is None: + hydrogens = "strip" if strip_h else "keep" + mode = state.pop("hydrogen_mode", None) + if mode not in HYDROGEN_MODES: + mode = "riding" if state.get("xyz.h_row") is not None else "atoms" + if hydrogens == "strip" and mode == "riding": + hydrogens = "keep" + + cell = state.pop("cell", None) + ctx = cls( + pdb=state.pop("pdb", None), + cell=Cell(cell, dtype=dtype, device=device) if cell is not None else None, + spacegroup=own_spacegroup(state.pop("spacegroup", None), dtype, device), + initialized=state.pop("initialized", False), + cif_path=state.pop("cif_path", None), + altloc_pairs=state.pop("altloc_pairs", []), + hydrogens=hydrogens, + hydrogen_mode=mode, + hydrogens_in_xray=state.pop("hydrogens_in_xray", True), + verbose=verbose, + ) + return ctx + @property def crystal_key(self): """Value identity of the crystal, or None while cell or space group is unset. @@ -205,8 +745,9 @@ def __repr__(self) -> str: sg = None if self.spacegroup is None else self.spacegroup.name return ( f"ModelContext(spacegroup={sg!r}, n_atoms={n_atoms}, " + f"hydrogens={self.hydrogens!r}, hydrogen_mode={self.hydrogen_mode!r}, " f"initialized={self.initialized})" ) -__all__ = ["ModelContext"] +__all__ = ["ModelContext", "HYDROGEN_SOURCES", "HYDROGEN_MODES", "check_hydrogen_policy"] diff --git a/torchref/model/model.py b/torchref/model/model.py index 500fcd35..6d5f082c 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -14,8 +14,6 @@ from typing import Dict, Iterable, List, Optional, Tuple, Union -import warnings - import gemmi import torch import torch.nn as nn @@ -28,7 +26,7 @@ normalize_device, ) from torchref.io import cif, pdb -from torchref.model.context import ModelContext +from torchref.model.context import ModelContext, own_spacegroup from torchref.model.parameter_wrappers import ( CholeskyMixedTensor, MixedTensor, @@ -40,38 +38,6 @@ from torchref.utils.device_mixin import DeviceMovementMixin from torchref.utils.utils import sanitize_pdb_dataframe -# Standard 3-letter to 1-letter amino acid code mapping -_THREE_TO_ONE = { - "ALA": "A", - "ARG": "R", - "ASN": "N", - "ASP": "D", - "CYS": "C", - "GLN": "Q", - "GLU": "E", - "GLY": "G", - "HIS": "H", - "ILE": "I", - "LEU": "L", - "LYS": "K", - "MET": "M", - "PHE": "F", - "PRO": "P", - "SER": "S", - "THR": "T", - "TRP": "W", - "TYR": "Y", - "VAL": "V", - "SEC": "U", - "PYL": "O", - # Common modified residues - "MSE": "M", - "CSE": "C", - "SEP": "S", - "TPO": "T", - "PTR": "Y", -} - class Model(DeviceMovementMixin, DebugMixin, nn.Module): """ @@ -92,16 +58,18 @@ class Model(DeviceMovementMixin, DebugMixin, nn.Module): Verbosity level for logging. Default is 1. device : torch.device, optional Computation device. Defaults to the configured device.current. - strip_H : bool, optional - Whether to strip hydrogen atoms when loading. Default False: hydrogens are kept - where the file has them. - add_hydrogens : bool, optional - Generate missing hydrogens on load when True. Default False; - ignored when ``strip_H`` is set. + hydrogens : {"keep", "add", "strip"}, optional + What loading does with the file's hydrogens: keep them (default), also generate + the missing ones from the monomer templates, or remove them all. + hydrogen_mode : {"atoms", "riding"}, optional + Hydrogens as ordinary refinable atoms (default) or riding on their parents. + ``"riding"`` with ``hydrogens="strip"`` raises ``ValueError``. cif_path : str or list of str, optional Restraint dictionary file(s) for residues the monomer library does not know, or whose library entry should be overridden. Given here rather than after loading so that hydrogen generation on load reads the same dictionary the restraints will. + hydrogens_in_xray : bool, optional + Whether hydrogens contribute to the structure factors. Default True. Attributes ---------- @@ -117,7 +85,7 @@ class Model(DeviceMovementMixin, DebugMixin, nn.Module): ctx : ModelContext The unit cell, space group, atom table, link records, provenance and configuration. The fields not forwarded below are reached through it, e.g. - ``model.ctx.strip_H`` and ``model.ctx.initialized``. + ``model.ctx.hydrogens`` and ``model.ctx.initialized``. pdb : pandas.DataFrame Atom table, forwarded to :attr:`ctx`. Only refreshed from the tensors by :meth:`update_pdb`. @@ -135,38 +103,15 @@ def __init__( dtype_float=None, verbose=1, device=None, - strip_H: bool = False, - add_hydrogens: bool = False, + hydrogens: str = "keep", + hydrogen_mode: str = "atoms", cif_path: Optional[Union[str, List[str]]] = None, hydrogens_in_xray: bool = True, ): - """ - Initialize an empty Model shell. + """Initialize an empty Model shell; see the class docstring for the arguments. - Creates a model shell ready for file loading via load_pdb()/load_cif() - or state restoration via load_state_dict(). - - Parameters - ---------- - dtype_float : torch.dtype, optional - Data type for floating point tensors. Defaults to the configured dtypes.float. - verbose : int, optional - Verbosity level for logging. Default is 1. - device : torch.device, optional - Computation device. Defaults to the configured device.current. - strip_H : bool, optional - Whether to strip hydrogen atoms when loading. Default False: hydrogens are - kept where the file has them. - add_hydrogens : bool, optional - Generate missing hydrogens on load when True. Default False; - ignored when ``strip_H`` is set. - cif_path : str or list of str, optional - Restraint dictionary file(s); see the class docstring. - :meth:`ModelContext.set_cif_path` can still change it after loading, but - generation on load only sees the value given here. - hydrogens_in_xray : bool, optional - Whether hydrogens contribute to the structure factors. Default True. They - stay in the restraints either way; see :attr:`hydrogens_in_xray`. + Load a structure with :meth:`load_pdb` / :meth:`load_cif`, or restore one with + :meth:`create_from_state_dict`. """ super().__init__() # Resolve dtype/device at call time (not import time) so a runtime @@ -180,17 +125,17 @@ def __init__( self.dtype_float = dtype_float self.device = device - # Everything the model is loaded from and sits in, as opposed to what is - # refined. Populated by load() / create_from_state_dict(). + # Settings only until a structure is loaded, which replaces the context with + # one built by ``ModelContext.from_atoms``. self.ctx = ModelContext( verbose=verbose, - strip_H=strip_H, - add_hydrogens=add_hydrogens, + hydrogens=hydrogens, + hydrogen_mode=hydrogen_mode, cif_path=cif_path, hydrogens_in_xray=hydrogens_in_xray, ) - # Submodules (created during load or load_state_dict) + # Parameter wrappers, installed by _install_parameters. self.xyz = None self.adp = None self.u = None @@ -222,29 +167,6 @@ def hydrogens_in_xray(self) -> bool: def hydrogens_in_xray(self, value: bool): self.ctx.hydrogens_in_xray = bool(value) - @property - def exclude_H_from_sf(self) -> bool: - """Inverse of :attr:`hydrogens_in_xray`. - - .. deprecated:: - Use ``hydrogens_in_xray`` instead. - """ - warnings.warn( - "exclude_H_from_sf is deprecated; use hydrogens_in_xray", - DeprecationWarning, - stacklevel=2, - ) - return not self.ctx.hydrogens_in_xray - - @exclude_H_from_sf.setter - def exclude_H_from_sf(self, value: bool): - warnings.warn( - "exclude_H_from_sf is deprecated; use hydrogens_in_xray", - DeprecationWarning, - stacklevel=2, - ) - self.ctx.hydrogens_in_xray = not bool(value) - def _sf_atom_mask(self) -> Optional[torch.Tensor]: """Atoms that enter Fcalc, or None when every atom does. @@ -365,22 +287,10 @@ def spacegroup(self, value): """Set the space group from a SpaceGroup, gemmi object, name or number. The model owns its space group: an incoming ``SpaceGroup`` is copied rather - than shared, because ``.to()`` moves in place and would otherwise relocate - the caller's object. The copy lands on the model's device and float dtype. - """ - if value is None: - self.ctx.spacegroup = None - elif isinstance(value, SpaceGroup): - self.ctx.spacegroup = value.copy().to( - device=self.device, dtype=self.dtype_float - ) - else: - # ``device=self.device``: SpaceGroup falls back to the global - # default otherwise, so setting a spacegroup on a CPU-pinned Model - # would silently plant accelerator-resident matrices on it. - self.ctx.spacegroup = SpaceGroup( - value, dtype=self.dtype_float, device=self.device - ) + than shared (see :func:`~torchref.model.context.own_spacegroup`), and lands on + the model's device and float dtype. + """ + self.ctx.spacegroup = own_spacegroup(value, self.dtype_float, self.device) # ========================================================================= # Crystallographic matrix properties (delegated to Cell) @@ -580,174 +490,142 @@ def _invalidate_atom_derived_caches(self) -> None: if hasattr(self, name): delattr(self, name) self._parametrization = None - self.ctx.restraints = None - def load(self, reader, add_hydrogens: bool = None): + def load(self, reader): """ Populate the model from a reader callable. - The central loader that ``load_pdb`` / ``load_cif`` / - ``_new_model_from_df`` funnel through: it strips hydrogens (when - ``strip_H``), drops rows with NaN coordinates / B-factors / occupancies, - builds the cell and space group, and constructs the four parameter - wrappers (``u`` as a :class:`CholeskyMixedTensor`, as - :meth:`create_from_state_dict` also does, so the parametrization - round-trips). + The central loader that ``load_pdb`` / ``load_cif`` funnel through. The context + is built by :meth:`ModelContext.from_atoms`, which applies the hydrogen policy, + drops rows without coordinates, B-factor or occupancy, and builds the cell and + space group; the parameter wrappers are then installed over it. Parameters ---------- reader : callable Zero-argument callable returning ``(pdb_df, cell, spacegroup)``. An - optional ``.links`` attribute on it is stored on ``self.ctx.links``. - add_hydrogens : bool, optional - Whether to top up missing hydrogens once the model is built. Defaults to the - context's setting, and is forced off for the re-entry that - :meth:`_add_missing_hydrogens` makes, so generation happens once per load. + optional ``.links`` attribute on it is kept as ``ctx.links``. Returns ------- Model Self, for method chaining. - - Notes - ----- - Side effects: sets ``pdb``, ``links``, ``cell``, ``spacegroup``, the - ``aniso_flag`` buffer, the four wrappers, the default masks, the altloc - registration and ``initialized = True``. """ - if add_hydrogens is None: - add_hydrogens = self.ctx.add_hydrogens and not self.ctx.strip_H + pdb, cell, spacegroup = reader() self._invalidate_atom_derived_caches() - self.pdb, cell, spacegroup = reader() - self.ctx.links = getattr(reader, "links", None) - - self.pdb = ( - self.pdb.loc[self.pdb["element"] != "H"].reset_index(drop=True) - if self.ctx.strip_H - else self.pdb + self.ctx = ModelContext.from_atoms( + pdb, + cell, + spacegroup, + dtype=self.dtype_float, + device=self.device, + links=getattr(reader, "links", None), + **self.ctx.settings(), ) - self.pdb.dropna(subset=["x", "y", "z", "tempfactor", "occupancy"], inplace=True) - # Reindex before deriving the ``index`` column: every consumer uses it to - # address length-N per-atom tensors positionally (see - # ``_create_occupancy_groups``), so a gapped index from the drop above sends - # them past the end. Only the strip_H branch reset, so a model losing rows to - # the dropna instead -- an atom with no coordinates or no B -- raised - # IndexError at load. Hit on roughly one PDB-REDO entry in six. - self.pdb.reset_index(drop=True, inplace=True) - self.pdb["index"] = self.pdb.index.to_numpy(dtype=int) + self._install_parameters() + return self + + def _install_parameters(self, state: Optional[dict] = None, xyz=None) -> None: + """Build the parameter wrappers and per-atom buffers over ``ctx.pdb``. - self.cell = Cell(cell, dtype=self.dtype_float, device=self.device) + The only place the wrappers are constructed: load, restore, select and every + derived model come through here. Values are read off the atom table. - # Setter also updates symmetry. - self.spacegroup = spacegroup + Parameters + ---------- + state : dict, optional + A state dict about to be loaded. Its saved refinable masks, riding frames + and node-field ADP layout fix the wrappers' shapes, and the default masks + are **not** applied; ``load_state_dict`` supplies the values afterwards. + xyz : MixedTensor, optional + Coordinate wrapper to install as is, instead of building one from the table. + """ + from torchref.model.riding_xyz import RidingXYZTensor + + pdb, dtype = self.pdb, self.dtype_float + restoring = state is not None + state = {} if state is None else state self.register_buffer( "aniso_flag", torch.tensor( - self.pdb["anisou_flag"].values, dtype=torch.bool, device=self.device - ), - ) - - self.xyz = MixedTensor( - torch.tensor(self.pdb[["x", "y", "z"]].values, dtype=self.dtype_float), - name="xyz", - device=self.device, - ) - self.adp = PositiveMixedTensor( - torch.tensor(self.pdb["tempfactor"].values, dtype=self.dtype_float), - name="adp", - device=self.device, - ) - # Cholesky parametrization keeps U positive-definite by construction - # (U = L Lᵀ), so refinement cannot drive it indefinite and NaN the FFT. - self.u = CholeskyMixedTensor( - torch.tensor( - self.pdb[["u11", "u22", "u33", "u12", "u13", "u23"]].values, - dtype=self.dtype_float, + pdb["anisou_flag"].values, dtype=torch.bool, device=self.device ), - name="aniso_U", - device=self.device, ) + self.xyz = self._build_xyz(state) if xyz is None else xyz + self.adp = self._restore_adp_slot("adp", state, pdb, dtype, self.xyz, self.device) + self.u = self._restore_adp_slot("u", state, pdb, dtype, self.xyz, self.device) # Residue-level sharing plus altloc sum-to-1 groups. - initial_occ = torch.tensor(self.pdb["occupancy"].values, dtype=self.dtype_float) - sharing_groups, altloc_groups, refinable_mask = self._create_occupancy_groups( - self.pdb, initial_occ + initial_occ = torch.tensor(pdb["occupancy"].values, dtype=dtype) + sharing_groups, altloc_groups, refinable_mask = self.ctx.occupancy_groups( + initial_occ ) + saved_occ_mask = state.get("occupancy.refinable_mask") + if saved_occ_mask is not None: + # Saved in group space; expanded back over atoms. + refinable_mask = saved_occ_mask.to(sharing_groups.device)[sharing_groups] self.occupancy = OccupancyTensor( initial_values=initial_occ, sharing_groups=sharing_groups, altloc_groups=altloc_groups, refinable_mask=refinable_mask, - dtype=self.dtype_float, + dtype=dtype, device=self.device, name="occupancy", ) - self.set_default_masks() - self.register_alternative_conformations() - self.ctx.initialized = True - - if add_hydrogens: - self._add_missing_hydrogens() - return self - - def _add_missing_hydrogens(self) -> None: - """Top up the hydrogens the atom table is missing, in place. - - Per parent, not per file: a structure deposited with some hydrogens gets the - rest, because the plan only ever proposes a hydrogen the template names and the - model does not have. 1AK5 arrives with 675 of roughly 2500, and a - does-it-have-any test would have left it there. - - Re-enters :meth:`load` on the augmented atom table, which rebuilds the parameter - wrappers and per-atom buffers at the new size. The re-entry is told not to - consider hydrogens again, so this runs once per load rather than recursing to a - fixed point. - - Costs a restraint build that is then discarded, because the plan needs the - topology and the topology is built over the atoms as loaded. It skips the - non-bonded pair search, which the plan does not use and which was the largest - part of that build, and it is silent: the build over the augmented table reports. - Loading invokes this only when ``add_hydrogens=True`` is requested. - """ - from torchref.topology.hydrogens import ( - augment_atom_table, - optimise_free_torsions, - plan_hydrogens, - ) - - xyz = self.xyz().detach() - restraints = self.ctx.restraints - if restraints is None: - restraints = self.ctx.build_restraints(xyz, nonbonded=False, verbose=0) - plan = plan_hydrogens( - restraints.topology, restraints.cif_dict, xyz, verbose=self.ctx.verbose - ) - if self.ctx.verbose > 0 and restraints.missing_residues: - print( - "No restraint dictionary for " - f"{sorted(restraints.missing_residues)}: not hydrogenated. Pass one " - "with cif_path / --cif." - ) - if plan.n_hydrogens == 0: + if restoring: + # Placeholders: the saved masks arrive with load_state_dict, and applying + # the defaults here would resize the refinable sets it has to match. + for mask_name in ("xyz_mask", "adp_mask", "u_mask", "occupancy_mask"): + self.register_buffer( + mask_name, torch.ones(len(pdb), dtype=torch.bool, device=self.device) + ) + if state.get("vdw_radii") is not None: + self.register_buffer( + "vdw_radii", torch.zeros_like(state["vdw_radii"], device=self.device) + ) return - optimise_free_torsions(plan, restraints.topology, xyz) - augmented = augment_atom_table(self.pdb, plan, restraints.topology) - if self.ctx.verbose > 0: - print(f"Generated {plan.n_hydrogens} hydrogens") + self.set_default_masks() + if self.ctx.hydrogen_mode == "riding" and not isinstance( + self.xyz, RidingXYZTensor + ): + self.xyz = RidingXYZTensor.from_mixed_tensor(self.xyz, self.hydrogen_frames()) + self._repoint_coordinate_accessors() - cell, spacegroup = self.cell, self.spacegroup - links = self.ctx.links + def _build_xyz(self, state: dict): + """The coordinate wrapper over the atom table, riding if ``state`` saved one. - def reader(): - return augmented, cell.data.cpu().numpy(), spacegroup + A saved riding wrapper is recognised by its frame buffers, never by shape: its + storage is ``(n_base, 3)`` and a plain wrapper's ``(n_atoms, 3)``, both 2-D. + """ + values = torch.tensor(self.pdb[["x", "y", "z"]].values, dtype=self.dtype_float) + mask = state.get("xyz.refinable_mask") + if state.get("xyz.h_row") is None: + return MixedTensor(values, refinable_mask=mask, name="xyz", device=self.device) - # Carried explicitly: ``load`` reads links off the reader, so a bare callable - # would drop the LINK records the first read resolved. - reader.links = links - self.load(reader, add_hydrogens=False) + from torchref.model.riding_xyz import RidingXYZTensor + from torchref.topology.hydrogens import HydrogenFrames + + frames = HydrogenFrames.from_tensors( + state["xyz.h_row"], + state["xyz.parent_row"], + state["xyz.n1_row"], + state["xyz.n2_row"], + state["xyz.frame_valid"], + state.get("xyz.torsion_group"), + state.get("xyz.rotation_group"), + ) + return RidingXYZTensor( + values, + frames, + refinable_mask=mask, + mask_in_base_space=True, + name="xyz", + device=self.device, + ) def load_pdb(self, file): """ @@ -790,171 +668,6 @@ def load_cif(self, file): return self.load(cif_reader) - @property - def chain_sequences(self) -> List[Tuple[str, str]]: - """Per-chain one-letter sequences, ``[(chain_id, sequence), ...]``. - - HETATM records are excluded, numbering gaps become ``?`` and unrecognized - residues ``X``. - """ - if self.pdb is None: - return [] - - atom_df = self.pdb[self.pdb["ATOM"] == "ATOM"] - result = [] - - for chain in atom_df["chainid"].unique(): - chain_df = atom_df[atom_df["chainid"] == chain] - residues = chain_df.drop_duplicates(subset=["resseq", "icode"]).sort_values( - "resseq" - ) - resseqs = residues["resseq"].values - resnames = residues["resname"].values - - seq_chars = [] - for i, (rseq, rname) in enumerate(zip(resseqs, resnames)): - if i > 0: - gap = int(rseq) - int(resseqs[i - 1]) - 1 - if gap > 0: - seq_chars.extend(["?"] * gap) - code = _THREE_TO_ONE.get(str(rname).strip(), "X") - seq_chars.append(code) - - result.append((str(chain), "".join(seq_chars))) - - return result - - def get_chain_residues(self) -> List[Tuple[str, List[str]]]: - """ - Per-chain residue names as 3-letter codes (for IHM/CIF writing). - - Excludes HETATM records. Unlike :attr:`chain_sequences`, returns - the raw 3-letter codes without gap filling. - - Returns - ------- - list of (str, list of str) - Ordered list of ``(chain_id, [resname, ...])``. - """ - if self.pdb is None: - return [] - - atom_df = self.pdb[self.pdb["ATOM"] == "ATOM"] - result = [] - - for chain in atom_df["chainid"].unique(): - chain_df = atom_df[atom_df["chainid"] == chain] - residues = chain_df.drop_duplicates(subset=["resseq", "icode"]).sort_values( - "resseq" - ) - resnames = [str(r).strip() for r in residues["resname"].values] - result.append((str(chain), resnames)) - - return result - - def _create_occupancy_groups(self, pdb_df, initial_occ): - """Build ``(sharing_groups, altloc_groups, refinable_mask)`` for - :class:`OccupancyTensor`. - - Altloc conformations share one collapsed index each; other residues share - one only when their occupancies agree to within 0.01, and an occupancy is - refinable only if it differs from 1.0 by more than that same deadband. - """ - n_atoms = len(initial_occ) - altloc_groups = [] - refinable_mask = torch.zeros(n_atoms, dtype=torch.bool) - - sharing_groups_tensor = torch.arange(n_atoms, dtype=torch.long) # dtype-ok: arange atom indices (sharing groups); index requires long - collapsed_idx = 0 - - # First pass: altlocs. ALL atoms of one conformation must share a collapsed - # index whatever their individual occupancies, or the sum-to-1 - # normalization in OccupancyTensor.forward() acts on the wrong group. - pdb_with_altlocs = pdb_df[pdb_df["altloc"] != ""] - altloc_residues = set() - - if len(pdb_with_altlocs) > 0: - grouped_by_residue = pdb_with_altlocs.groupby( - ["resname", "resseq", "chainid"] - ) - - for (resname, resseq, chainid), group in grouped_by_residue: - unique_altlocs = sorted(group["altloc"].unique()) - - if len(unique_altlocs) > 1: - altloc_residues.add((resname, resseq, chainid)) - conformation_atom_lists = [] - - for altloc in unique_altlocs: - altloc_atoms = group[group["altloc"] == altloc] - indices = altloc_atoms["index"].tolist() - - sharing_groups_tensor[indices] = collapsed_idx - - for idx in indices: - if abs(initial_occ[idx].item() - 1.0) > 0.01: - refinable_mask[idx] = True - - conformation_atom_lists.append(indices) - collapsed_idx += 1 - - altloc_groups.append(tuple(conformation_atom_lists)) - - # Second pass: non-altloc residues, sharing by occupancy similarity. - grouped = pdb_df.groupby(["resname", "resseq", "chainid", "altloc"]) - - for (resname, resseq, chainid, altloc), group in grouped: - if (resname, resseq, chainid) in altloc_residues: - continue - - indices = group["index"].tolist() - - if len(indices) == 0: - continue - - residue_occs = initial_occ[indices] - - occ_min = residue_occs.min().item() - occ_max = residue_occs.max().item() - occ_mean = residue_occs.mean().item() - - if (occ_max - occ_min) <= 0.01: - sharing_groups_tensor[indices] = collapsed_idx - collapsed_idx += 1 - - if abs(occ_mean - 1.0) > 0.01: - for idx in indices: - refinable_mask[idx] = True - else: - # Occupancies disagree within the residue: keep atoms independent. - for idx in indices: - if abs(initial_occ[idx].item() - 1.0) > 0.01: - refinable_mask[idx] = True - - # Compact to contiguous indices 0..n_collapsed-1. - unique_indices = torch.unique(sharing_groups_tensor, sorted=True) - index_map = torch.zeros(n_atoms, dtype=torch.long) # dtype-ok: index_map atom-index remap; indexing requires long - for new_idx, old_idx in enumerate(unique_indices): - mask = sharing_groups_tensor == old_idx - sharing_groups_tensor[mask] = new_idx - - n_collapsed = len(unique_indices) - - if self.ctx.verbose > 1: - n_groups = n_collapsed - n_independent = n_atoms - n_collapsed - n_refinable = refinable_mask.sum().item() - n_altloc_groups = len(altloc_groups) - - print("\nOccupancy Setup:") - print(f" Total atoms: {n_atoms}") - print(f" Collapsed indices: {n_collapsed}") - print(f" Alternative conformation groups: {n_altloc_groups}") - print(f" Refinable atoms: {n_refinable}") - print(f" Compression ratio: {n_atoms / n_collapsed:.2f}x") - - return sharing_groups_tensor, altloc_groups, refinable_mask - def update_pdb(self): """ Write the current refinable parameters back into ``self.pdb``. @@ -1041,52 +754,87 @@ def _after_device_apply( def copy(self): """ - Create a deep copy of the Model. + Create a deep copy of the model, of the same class. - Independent in every part: the context is copied via + Independent in every part: the context -- restraints included -- is copied via :meth:`~torchref.model.context.ModelContext.copy`, buffers are cloned and each parameter wrapper is copied through its own ``copy`` so its parametrization - survives. + survives. Subclass settings carry over through :meth:`_subclass_kwargs`. Returns ------- Model - A new, fully independent Model instance with copied data. + A new, fully independent instance with copied data. """ + import copy as copy_module + if not self.ctx.initialized: raise RuntimeError("Cannot copy an uninitialized Model. Load data first.") - model_copy = Model( + duplicate = self._spawn(self.ctx.copy()) + for name, buffer in self._buffers.items(): + if buffer is not None: + duplicate.register_buffer(name, buffer.clone().detach()) + for name, module in self._modules.items(): + # Submodules the constructor already built (ModelFT's engine) derive from + # the context and are not copied. + if module is None or name in duplicate._modules or not hasattr(module, "copy"): + continue + setattr(duplicate, name, module.copy()) + if self._parametrization is not None: + duplicate._parametrization = copy_module.deepcopy(self._parametrization) + + # Anything that borrows the coordinates -- the ADP node field -- carries the + # reference through its own ``copy`` and still points at THIS model's ``xyz``. + duplicate._repoint_coordinate_accessors() + if hasattr(duplicate, "reset_cache"): + duplicate.reset_cache() + + if self.ctx.verbose > 0: + print(f"Copied {type(self).__name__} ({len(duplicate.pdb)} atoms)") + return duplicate + + def _spawn(self, ctx: ModelContext) -> "Model": + """An empty instance of this class around ``ctx``. + + Carries this model's dtype, device and subclass settings + (:meth:`_subclass_kwargs`); the caller installs the parameters. + """ + model = type(self)( dtype_float=self.dtype_float, - verbose=self.ctx.verbose, + verbose=ctx.verbose, device=self.device, - strip_H=self.ctx.strip_H, + **self._subclass_kwargs(), ) + model.ctx = ctx + return model - # One call carries the atom table, cell, space group, altloc groups and - # provenance, each deep-copied or cloned -- see ``ModelContext.copy``. - model_copy.ctx = self.ctx.copy() + def _subclass_kwargs(self) -> dict: + """Constructor arguments a subclass adds, as this instance holds them. - for buffer_name, buffer_value in self._buffers.items(): - if buffer_value is not None: - model_copy.register_buffer(buffer_name, buffer_value.clone()) - - # Parameter wrappers via their own .copy(), which preserves each - # wrapper's parametrization (log-space, Cholesky, collapsed logits). - for module_name, module in self._modules.items(): - if module is not None and hasattr(module, "copy"): - setattr(model_copy, module_name, module.copy()) + :meth:`copy`, :meth:`select` and the derived-model helpers build new instances + through it. Model adds none. + """ + return {} - # Anything that borrows the coordinates -- the ADP node field, the restraints' - # pair-list maintenance -- carries the reference through its own ``copy`` and - # still points at THIS model's ``xyz``. Re-point it, or the two models silently - # share coordinates and the copy is not independent. - model_copy._repoint_coordinate_accessors() + def _derive(self, pdb, **overrides) -> "Model": + """A new, quiet model of this class over ``pdb`` in this crystal. - if self.ctx.verbose > 0: - print(f"✓ Model copied successfully ({len(model_copy.pdb)} atoms)") + Parameters + ---------- + pdb : pandas.DataFrame + Atom table for the new model; settled by + :meth:`~torchref.model.context.ModelContext.derive`. + **overrides + Context settings to change, e.g. ``hydrogens="strip"``. + """ + model = self._spawn(self.ctx.derive(pdb, **{"verbose": 0, **overrides})) + model._install_parameters() + return model - return model_copy + def _kept_hydrogens(self) -> str: + """The policy for a table derived from this one: never generate again.""" + return "strip" if self.ctx.hydrogens == "strip" else "keep" def write_pdb(self, filename, metadata=None): """Write model to PDB file with optional metadata header. @@ -1157,7 +905,7 @@ def set_default_masks(self): B-factors), ``u_mask`` (atoms with no NaN U component), and ``occupancy_mask`` (occupancies below 0.999), then pushes each mask into the corresponding parameter wrapper via ``update_refinable_mask``. - Called from :meth:`load` after the wrappers are constructed. + Called from :meth:`_install_parameters` after the wrappers are constructed. """ self.register_buffer( "xyz_mask", torch.ones(len(self.pdb), dtype=torch.bool, device=self.device) @@ -1896,39 +1644,6 @@ def print_parameters_info(self): ) print("=" * 80) - def register_alternative_conformations(self): - """ - Rebuild ``self.ctx.altloc_pairs`` from the ``altloc`` column. - - One tuple per residue that has multiple conformations, holding one - index tensor per conformation (in sorted altloc order), e.g. - ``[(tensor([100, 101]), tensor([110, 111])), ...]``. Overwrites any - previous content, so call it after the atom numbering changes. - """ - self.ctx.altloc_pairs = [] - - pdb_with_altlocs = self.pdb[self.pdb["altloc"] != ""] - - if len(pdb_with_altlocs) == 0: - return - - grouped = pdb_with_altlocs.groupby(["resname", "resseq", "chainid"]) - - for (resname, resseq, chainid), group in grouped: - unique_altlocs = sorted(group["altloc"].unique()) - - # A lone altloc label is not an alternative conformation. - if len(unique_altlocs) > 1: - conformation_tensors = [] - for altloc in unique_altlocs: - altloc_atoms = group[group["altloc"] == altloc] - indices = torch.tensor( - altloc_atoms["index"].tolist(), dtype=torch.long # dtype-ok: altloc atom indices; indexing requires long - ) - conformation_tensors.append(indices) - - self.ctx.altloc_pairs.append(tuple(conformation_tensors)) - def shake_coords(self, stddev: float): """ Perturb every atom's coordinates with Gaussian noise of width *stddev* (Å). @@ -1965,44 +1680,6 @@ def shake_adp(self, stddev: float): ) - def _new_model_from_df(self, df, *, strip_H=None, add_hydrogens=False): - """Build a fresh model of the same class from a DataFrame. - - ``add_hydrogens`` defaults to False: the caller has - already settled which atoms the table holds, and generating more would fight - that. :meth:`hydrogenate` passes an already-augmented table for the same reason. - """ - import inspect - - sh = self.ctx.strip_H if strip_H is None else strip_H - ctor_kw = dict( - dtype_float=self.dtype_float, - verbose=0, - device=self.device, - strip_H=sh, - add_hydrogens=add_hydrogens, - cif_path=self.ctx.cif_path, - ) - sig = inspect.signature(self.__class__.__init__) - for pname, param in sig.parameters.items(): - if pname in ("self",) or pname in ctor_kw: - continue - if param.kind in (param.VAR_POSITIONAL, param.VAR_KEYWORD): - continue - if pname == "gridsize": - # The constructor argument is the explicit override, not the - # derived grid a ``gridsize`` attribute would return. - if hasattr(self, "explicit_gridsize"): - ctor_kw[pname] = self.explicit_gridsize - continue - if hasattr(self, pname): - ctor_kw[pname] = getattr(self, pname) - - new_model = self.__class__(**ctor_kw) - sg_str = self.spacegroup.xhm if self.spacegroup else "P 1" - new_model.load(lambda: (df, self.pdb.attrs.get("cell"), sg_str)) - return new_model - def strip_altlocs(self) -> "Model": """Return a new model with alternate conformations removed. @@ -2016,7 +1693,7 @@ def strip_altlocs(self) -> "Model": pdb = self.pdb.copy() has_altloc = pdb["altloc"].astype(str).str.strip() != "" if not has_altloc.any(): - return self._new_model_from_df(pdb) + return self._derive(pdb, hydrogens=self._kept_hydrogens()) drop_idx = [] res_cols = ["chainid", "resseq", "icode", "resname"] @@ -2044,14 +1721,13 @@ def strip_altlocs(self) -> "Model": # Preserve DataFrame attrs filtered.attrs = pdb.attrs.copy() - return self._new_model_from_df(filtered) + return self._derive(filtered, hydrogens=self._kept_hydrogens()) def strip_hydrogens(self) -> "Model": """Return a new model with hydrogen atoms removed. - The returned model has consistent DataFrame and tensors (xyz, adp, - occupancy) with H atoms excluded. The original model is not - modified. + Built from the current parameter values with ``hydrogens="strip"`` and + ``hydrogen_mode="atoms"``. The original model is not modified. Returns ------- @@ -2059,71 +1735,37 @@ def strip_hydrogens(self) -> "Model": New model without hydrogen atoms. """ self.update_pdb() - pdb = self.pdb.copy() - h_mask = pdb["element"].str.strip() == "H" - if not h_mask.any(): - return self._new_model_from_df(pdb, strip_H=True) - - filtered = pdb[~h_mask].reset_index(drop=True) - filtered["index"] = range(len(filtered)) - filtered.attrs = pdb.attrs.copy() - return self._new_model_from_df(filtered, strip_H=True) + return self._derive(self.pdb.copy(), hydrogens="strip", hydrogen_mode="atoms") - def hydrogenate(self, verbose: int = 0, optimize: bool = True) -> "Model": - """Return a new model with hydrogens added from the monomer templates. + def hydrogenate(self, verbose: int = 0) -> "Model": + """Return a new model with the missing hydrogens added from the monomer templates. - Hydrogen generation is template instantiation over the topology: each residue's + Built from the current parameter values with ``hydrogens="add"``: each residue's library template is aligned onto the heavy atoms present and its hydrogens read - off, and the bond graph decides how many hydrogens a parent can carry and which - of them have a free torsion. Missing HOH hydrogens use the water dictionary - geometry with a random initial orientation, controlled by ``torch.manual_seed``. - Existing hydrogen coordinates are retained. The original model is not modified. + off, and every free torsion (hydroxyl, thiol, amine, methyl) is scanned for the + least-clashing angle. Missing HOH hydrogens get the dictionary geometry with a + random orientation drawn from ``torch.manual_seed``. Existing hydrogens are + retained. The original model is not modified. Parameters ---------- verbose : int, default 0 Verbosity level. - optimize : bool, default True - Scan each free torsion -- hydroxyl, thiol, amine, methyl -- for the - least-clashing angle. The template's dihedral for those is arbitrary, so this - is on by default; it is a rotation about one bond and costs little. Returns ------- Model - New model with hydrogens, built with ``strip_H=False`` so they survive the - load. """ - from torchref.topology.hydrogens import ( - augment_atom_table, - optimise_free_torsions, - plan_hydrogens, - ) - self.update_pdb() - restraints = self.restraints # builds the topology this reads - xyz = self.xyz().detach() - - plan = plan_hydrogens( - restraints.topology, restraints.cif_dict, xyz, verbose=verbose - ) - if optimize: - optimise_free_torsions(plan, restraints.topology, xyz) - - if verbose > 0: - print(f"Adding {plan.n_hydrogens} hydrogens") - augmented = augment_atom_table(self.pdb, plan, restraints.topology) - return self._new_model_from_df(augmented, strip_H=False) - + return self._derive(self.pdb.copy(), hydrogens="add", verbose=verbose) def state_dict(self, destination=None, prefix="", keep_vars=False): """ Return a dictionary containing the complete state of the Model. - Registered buffers, the four parameter wrappers, the PDB DataFrame and the - metadata (space group as a string, cell as a CPU tensor, dtype, device, - ``strip_H``, altloc pairs). Restore with :meth:`create_from_state_dict`, - which is what knows how to rebuild the wrappers. + Registered buffers, the four parameter wrappers, the context's entries + (:meth:`ModelContext.state`), dtype and device. Restore with + :meth:`create_from_state_dict`, which is what knows how to rebuild the wrappers. Parameters ---------- @@ -2143,18 +1785,10 @@ def state_dict(self, destination=None, prefix="", keep_vars=False): destination=destination, prefix=prefix, keep_vars=keep_vars ) - state[prefix + "pdb"] = self.pdb.copy() if self.pdb is not None else None - state[prefix + "cell"] = self.cell.data.cpu() if self.cell is not None else None - # As a string: gemmi.SpaceGroup is not picklable. - state[prefix + "spacegroup"] = self.spacegroup.xhm if self.spacegroup else None - state[prefix + "initialized"] = self.ctx.initialized + for key, value in self.ctx.state().items(): + state[prefix + key] = value state[prefix + "dtype_float"] = self.dtype_float state[prefix + "device"] = self.device - state[prefix + "strip_H"] = self.ctx.strip_H - state[prefix + "cif_path"] = self.ctx.cif_path - state[prefix + "altloc_pairs"] = self.ctx.altloc_pairs - state[prefix + "hydrogens_in_xray"] = self.ctx.hydrogens_in_xray - state[prefix + "hydrogen_mode"] = self.ctx.hydrogen_mode return state @@ -2197,11 +1831,11 @@ def load_state(self, path: str, strict: bool = True, device=None): print(f"Loaded model state from {path}") @staticmethod - def _restore_adp_slot(prefix, state_dict, pdb, saved_dtype, xyz_wrapper): - """Rebuild the ``adp`` or ``u`` wrapper, as a node field when the state was one. + def _restore_adp_slot(prefix, state_dict, pdb, saved_dtype, xyz_wrapper, device): + """Build the ``adp`` or ``u`` wrapper, as a node field when the state was one. - Built from the PDB for its shapes and masks only; ``load_state_dict`` overwrites - every value afterwards. + With an empty ``state_dict`` this is the per-atom wrapper a fresh load uses. + Otherwise the values are placeholders that ``load_state_dict`` overwrites. A saved :class:`~torchref.model.disorder_field.DisorderFieldTensor` is recognised by its ``neighbor_list``, not by the shape of its storage: the ``u`` slot holds a @@ -2219,8 +1853,9 @@ def _restore_adp_slot(prefix, state_dict, pdb, saved_dtype, xyz_wrapper): saved_dtype : torch.dtype Float dtype the state was saved in. xyz_wrapper : MixedTensor - The already-rebuilt coordinate wrapper; a node field derives its node + The already-built coordinate wrapper; a node field derives its node positions from it. + device : torch.device """ from torchref.model.parameter_wrappers import ( CholeskyMixedTensor, @@ -2244,7 +1879,7 @@ def _restore_adp_slot(prefix, state_dict, pdb, saved_dtype, xyz_wrapper): # model refines it in the same positive-definite-by-construction # parametrization as a freshly-loaded one. wrapper = CholeskyMixedTensor if aniso else PositiveMixedTensor - return wrapper(initial, refinable_mask=mask, name=name) + return wrapper(initial, refinable_mask=mask, name=name, device=device) from torchref.model.disorder_field import ( AnisotropicPayload, @@ -2285,100 +1920,9 @@ def _restore_adp_slot(prefix, state_dict, pdb, saved_dtype, xyz_wrapper): mask_in_node_space=True, name=name, dtype=saved_dtype, - ) - - @classmethod - def _rebuild_wrappers_from_pdb(cls, instance, pdb, state_dict, saved_dtype, device): - """Give ``instance`` parameter wrappers and per-atom buffers of the right shape. - - The half of :meth:`create_from_state_dict` that every subclass needs - identically, so subclasses call this rather than restating it: a per-class copy - drifts, and a restore that rebuilds the wrong wrapper type fails on a shape - mismatch rather than on anything that names the real cause. - - Values are placeholders throughout --- the caller's ``load_state_dict`` is what - puts the saved numbers in. Only shapes, masks and dtypes matter here. - """ - from torchref.model.parameter_wrappers import MixedTensor, OccupancyTensor - - n_atoms = len(pdb) - - xyz_values = torch.tensor(pdb[["x", "y", "z"]].values, dtype=saved_dtype) - if state_dict.get("xyz.h_row") is not None: - # A saved riding wrapper is recognised by its frame buffers, never by - # shape: its storage is (n_base, 3), a plain wrapper's (n_atoms, 3), and - # both are 2-D. The frames restore from the buffers, so no topology is - # needed here. - from torchref.model.riding_xyz import RidingXYZTensor - from torchref.topology.hydrogens import HydrogenFrames - - frames = HydrogenFrames.from_tensors( - state_dict["xyz.h_row"], - state_dict["xyz.parent_row"], - state_dict["xyz.n1_row"], - state_dict["xyz.n2_row"], - state_dict["xyz.frame_valid"], - state_dict.get("xyz.torsion_group"), - state_dict.get("xyz.rotation_group"), - ) - instance.xyz = RidingXYZTensor( - xyz_values, - frames, - refinable_mask=state_dict.get("xyz.refinable_mask"), - mask_in_base_space=True, - name="xyz", - ) - else: - instance.xyz = MixedTensor( - xyz_values, - refinable_mask=state_dict.get("xyz.refinable_mask"), - name="xyz", - ) - instance.adp = cls._restore_adp_slot( - "adp", state_dict, pdb, saved_dtype, instance.xyz - ) - instance.u = cls._restore_adp_slot( - "u", state_dict, pdb, saved_dtype, instance.xyz - ) - - initial_occ = torch.tensor(pdb["occupancy"].values, dtype=saved_dtype) - sharing_groups, altloc_groups, refinable_mask = ( - instance._create_occupancy_groups(pdb, initial_occ) - ) - # A saved mask is in group space; expand it back over atoms. - saved_occ_mask = state_dict.get("occupancy.refinable_mask") - if saved_occ_mask is not None: - if saved_occ_mask.device != sharing_groups.device: - saved_occ_mask = saved_occ_mask.to(sharing_groups.device) - refinable_mask = saved_occ_mask[sharing_groups] - - instance.occupancy = OccupancyTensor( - initial_values=initial_occ, - sharing_groups=sharing_groups, - altloc_groups=altloc_groups, - refinable_mask=refinable_mask, - dtype=saved_dtype, device=device, - name="occupancy", ) - if "aniso_flag" not in instance._buffers or instance.aniso_flag is None: - instance.register_buffer( - "aniso_flag", - torch.tensor(pdb["anisou_flag"].values, dtype=torch.bool), - ) - for mask_name in ("xyz_mask", "adp_mask", "u_mask", "occupancy_mask"): - instance.register_buffer( - mask_name, torch.ones(n_atoms, dtype=torch.bool, device=device) - ) - - # Note: inv_fractional_matrix, fractional_matrix and recB are properties - # delegating to Cell, so they are not registered as buffers. - if state_dict.get("vdw_radii") is not None: - instance.register_buffer( - "vdw_radii", torch.zeros_like(state_dict["vdw_radii"], device=device) - ) - @classmethod def create_from_state_dict( cls, @@ -2388,24 +1932,22 @@ def create_from_state_dict( dtype_float: torch.dtype = None, ) -> "Model": """ - Create a fully initialized Model from a state dictionary. - - This is the recommended way to restore a Model from a saved state. - Creates an instance with properly initialized submodules, then loads the state. + Create a fully initialized model of this class from a state dictionary. Parameters ---------- state_dict : dict - State dictionary from torch.save(model.state_dict(), ...). + State dictionary from ``torch.save(model.state_dict(), ...)``. device : torch.device, optional Move the restored model here once it is built. The restore itself always runs on CPU; ``None`` then moves it to the configured default device (``get_default_device()``), so a round-trip lands beside a same-config - model rather than stranding itself on CPU. Pass a device to override. + model rather than stranding itself on CPU. verbose : int, optional Verbosity level. Default is 1. dtype_float : torch.dtype, optional - Float dtype for tensors. Defaults to the configured dtypes.float. + Float dtype for tensors when the state does not record one. Defaults to + the configured dtypes.float. Returns ------- @@ -2414,83 +1956,60 @@ def create_from_state_dict( Notes ----- - Consumes ``state_dict``: the metadata keys are popped off it. The - anisotropic ``u`` is rebuilt as a :class:`CholeskyMixedTensor`, matching - :meth:`load`, so the positive-definite parametrization round-trips. - """ - # Build on CPU throughout, then move once at the end -- to the caller's device - # if they named one, otherwise to the configured default device, so a restore - # lands beside a same-config model instead of stranding itself on CPU. One - # device for the whole model is the invariant that matters: the wrappers are - # built from the atom table and land on CPU whatever is asked for, so resolving - # an accelerator up front splits the model rather than placing it. + Consumes ``state_dict``: the metadata keys are popped off it. Checkpoints + written before the hydrogen policy existed are mapped onto it; see + :meth:`ModelContext.from_state`. + """ + # Build on CPU throughout, then move once: the wrappers are built from the atom + # table and land on CPU whatever is asked for, so resolving an accelerator up + # front would split the model rather than place it. target_device = ( canonical_device(device) if device is not None else get_default_device() ) - device = torch.device("cpu") + cpu = torch.device("cpu") if dtype_float is None: dtype_float = get_float_dtype() - pdb = state_dict.pop("pdb", None) - cell_tensor = state_dict.pop("cell", None) - spacegroup = state_dict.pop("spacegroup", None) - initialized = state_dict.pop("initialized", False) saved_dtype = state_dict.pop("dtype_float", dtype_float) - state_dict.pop("device", None) # popped so it never reaches load_state_dict - strip_H = state_dict.pop("strip_H", True) - cif_path = state_dict.pop("cif_path", None) - altloc_pairs = state_dict.pop("altloc_pairs", []) - hydrogens_in_xray = state_dict.pop("hydrogens_in_xray", True) - hydrogen_mode = state_dict.pop("hydrogen_mode", None) + state_dict.pop("device", None) instance = cls( dtype_float=saved_dtype, verbose=verbose, - device=device, - strip_H=strip_H, - cif_path=cif_path, - hydrogens_in_xray=hydrogens_in_xray, + device=cpu, + **cls._pop_subclass_state(state_dict), ) - if hydrogen_mode is None: - # Older checkpoints: riding wrappers did not exist, so any hydrogens - # present were free parameters. - hydrogen_mode = "riding" if state_dict.get("xyz.h_row") is not None else "free" - instance.ctx.hydrogen_mode = hydrogen_mode - - instance.pdb = pdb - instance.ctx.initialized = initialized - instance.ctx.altloc_pairs = altloc_pairs - - # Setter also sets symmetry. - instance.spacegroup = spacegroup - - if cell_tensor is not None: - instance.cell = Cell(cell_tensor, dtype=saved_dtype, device=device) - - # The wrappers are built from the PDB purely to get the right shapes and - # masks; load_state_dict below overwrites their values. - if pdb is not None: - cls._rebuild_wrappers_from_pdb(instance, pdb, state_dict, saved_dtype, device) - - # Drop only empty-in-dim-0 tensors (placeholders from an atom-less state); - # scalars and non-tensor entries must survive for load_state_dict. - state_dict = { - k: v - for k, v in state_dict.items() - if not (torch.is_tensor(v) and v.ndim >= 1 and v.shape[0] == 0) - } - instance.load_state_dict(state_dict, strict=False) - - # Always placed: target_device is the caller's device or the configured default, - # never None. Without this the restore used to stay on CPU and split a - # round-trip's restored model from its (default-device) source. + instance.ctx = ModelContext.from_state( + state_dict, dtype=saved_dtype, device=cpu, verbose=verbose + ) + if instance.pdb is not None: + instance._install_parameters(state=state_dict) + instance.load_state_dict(instance._restorable_entries(state_dict), strict=False) instance.to(target_device) + if hasattr(instance, "reset_cache"): + instance.reset_cache() if verbose > 0: n_atoms = len(instance.pdb) if instance.pdb is not None else 0 - print(f"Created Model from state_dict: {n_atoms} atoms") - + print(f"Created {cls.__name__} from state_dict: {n_atoms} atoms") return instance + @classmethod + def _pop_subclass_state(cls, state_dict: dict) -> dict: + """Pop a subclass's own metadata keys and return them as constructor kwargs.""" + return {} + + def _restorable_entries(self, state_dict: dict) -> dict: + """The entries of ``state_dict`` that ``load_state_dict`` should see. + + Drops tensors empty along dim 0 (placeholders from an atom-less state); scalars + and non-tensor entries survive. + """ + return { + k: v + for k, v in state_dict.items() + if not (torch.is_tensor(v) and v.ndim >= 1 and v.shape[0] == 0) + } + def get_selection_mask(self, selection: str) -> torch.Tensor: """ Return a boolean mask for atoms matching a Phenix-style selection. @@ -2535,10 +2054,7 @@ def get_selection_mask(self, selection: str) -> torch.Tensor: def select(self, selection: str) -> "Model": """ - Return a new Model containing only atoms matching the Phenix-style selection. - - An independent model with every per-atom tensor, buffer and metadata field - subsetted, built as ``type(self)`` so subclasses return their own class. + Return a new model of the same class holding only the atoms a selection matches. Parameters ---------- @@ -2548,7 +2064,9 @@ def select(self, selection: str) -> "Model": Returns ------- Model - New instance of the same class holding only the selected atoms. + Built from the current parameter values, with default refinable masks. A + riding wrapper keeps its frames and orientations, and a hydrogen whose + parent is cut becomes an ordinary row. Restraints rebuild on first access. Raises ------ @@ -2556,22 +2074,6 @@ def select(self, selection: str) -> "Model": If the model has not been initialized. ValueError If selection syntax is invalid or no atoms are selected. - - Notes - ----- - The subclass constructor is called with the base kwargs only, so - subclass-specific settings fall back to their defaults (see - :meth:`ModelFT.select`). ``u`` is rebuilt as a plain - :class:`MixedTensor`, *not* a :class:`CholeskyMixedTensor`, so the - selected model loses the positive-definite parametrization of its - anisotropic ADPs. - - Examples - -------- - :: - - chain_a = model.select("chain A") - no_water = model.select("not resname HOH") """ from torchref.utils.utils import parse_phenix_selection @@ -2580,102 +2082,34 @@ def select(self, selection: str) -> "Model": "Cannot select from an uninitialized Model. Load data first." ) - selection_mask = parse_phenix_selection(selection, self.pdb) - - n_selected = selection_mask.sum().item() + mask = parse_phenix_selection(selection, self.pdb) + n_selected = int(mask.sum()) if n_selected == 0: raise ValueError(f"Selection '{selection}' matched no atoms.") - selected_indices = torch.where(selection_mask)[0] - - # type(self), so a subclass returns its own type. - selected_model = type(self)( - dtype_float=self.dtype_float, - verbose=self.ctx.verbose, - device=self.device, - strip_H=self.ctx.strip_H, - cif_path=self.ctx.cif_path, - hydrogens_in_xray=self.ctx.hydrogens_in_xray, - ) - - # ``index`` must be renumbered: the occupancy grouping below reads it. - mask_np = selection_mask.cpu().numpy() - selected_model.pdb = self.pdb.loc[mask_np].copy() - selected_model.pdb = selected_model.pdb.reset_index(drop=True) - selected_model.pdb["index"] = selected_model.pdb.index.to_numpy(dtype=int) - - # The setter rebuilds a SpaceGroup, so the selection gets its own. - selected_model.spacegroup = self.spacegroup - - # The fractional / reciprocal matrices are properties over the Cell, so - # cloning the Cell carries all of them. - if self.cell is not None: - selected_model.cell = self.cell.clone() - - if hasattr(self, "aniso_flag") and self.aniso_flag is not None: - selected_model.register_buffer( - "aniso_flag", self.aniso_flag[selection_mask].clone() - ) - - if hasattr(self.xyz, "select_rows"): - # Riding wrapper: frames are remapped, a hydrogen whose parent is cut - # becomes an ordinary row. - selected_model.xyz = self.xyz.select_rows(selection_mask) - else: - selected_model.xyz = MixedTensor( - self.xyz()[selection_mask].clone().detach(), - refinable_mask=( - self.xyz.refinable_mask[selection_mask] - if self.xyz.refinable_mask is not None - else None - ), - name="xyz", - ) - - selected_model.adp = PositiveMixedTensor( - self.adp()[selection_mask].clone().detach(), - refinable_mask=( - self.adp.refinable_mask[selection_mask] - if self.adp.refinable_mask is not None - else None - ), - name="adp", - ) - - selected_model.u = MixedTensor( - self.u()[selection_mask].clone().detach(), - refinable_mask=( - self.u.refinable_mask[selection_mask] - if self.u.refinable_mask is not None - else None - ), - name="aniso_U", - ) - - # Occupancy sharing/altloc groups must be rebuilt for the new numbering. - initial_occ = self.occupancy()[selection_mask].clone().detach() - sharing_groups, altloc_groups, refinable_mask = ( - selected_model._create_occupancy_groups(selected_model.pdb, initial_occ) - ) - selected_model.occupancy = OccupancyTensor( - initial_values=initial_occ, - sharing_groups=sharing_groups, - altloc_groups=altloc_groups, - refinable_mask=refinable_mask, - dtype=self.dtype_float, - device=self.device, - name="occupancy", - ) - - selected_model.set_default_masks() - selected_model.register_alternative_conformations() - selected_model.ctx.initialized = True - selected_model.ctx.hydrogen_mode = self.ctx.hydrogen_mode + table = self._table_with_current_values().loc[mask.cpu().numpy()] + selected = self._spawn(self.ctx.derive(table, hydrogens=self._kept_hydrogens())) + riding_xyz = self.xyz.select_rows(mask) if hasattr(self.xyz, "select_rows") else None + selected._install_parameters(xyz=riding_xyz) if self.ctx.verbose > 0: print(f"Selected {n_selected}/{len(self.pdb)} atoms with '{selection}'") + return selected + + def _table_with_current_values(self): + """A copy of the atom table carrying the wrappers' current values. - return selected_model + Unlike :meth:`update_pdb` it leaves ``self.pdb`` alone, and ``tempfactor`` is + the isotropic wrapper's value rather than the anisotropic B_eq. + """ + table = self.pdb.copy() + table[["x", "y", "z"]] = self.xyz().detach().cpu().numpy() + table["tempfactor"] = self.adp().detach().cpu().numpy() + table[["u11", "u22", "u33", "u12", "u13", "u23"]] = ( + self.u().detach().cpu().numpy() + ) + table["occupancy"] = self.occupancy().detach().cpu().numpy() + return table def xyz_fractional(self) -> torch.Tensor: """ @@ -2810,7 +2244,7 @@ def get_centroid(self) -> torch.Tensor: @property def hydrogen_mode(self) -> str: - """``"riding"``, ``"free"`` or ``"none"``; see :class:`ModelContext`.""" + """``"atoms"`` or ``"riding"``; see :class:`ModelContext`.""" return self.ctx.hydrogen_mode def hydrogen_frames(self): @@ -2825,9 +2259,6 @@ def hydrogen_frames(self): """ if hasattr(self.xyz, "hydrogen_frames"): return self.xyz.hydrogen_frames() - frames = getattr(self, "_hydrogen_frames", None) - if frames is not None and frames.n_hydrogens >= 0: - return frames from torchref.topology.hydrogens import hydrogen_frames return hydrogen_frames(self.restraints.topology) @@ -2842,177 +2273,61 @@ def _repoint_coordinate_accessors(self) -> None: if module is not None and hasattr(module, "set_xyz_fn"): module.set_xyz_fn(self.xyz) - def _complete_riding_waters(self, frames): - """Complete HOH residues once and remap frames and refinement selections.""" - from dataclasses import fields - - import numpy as np - - from torchref.topology.hydrogens import ( - HydrogenFrames, - augment_atom_table_with_maps, - hydrogen_frames, - plan_hydrogens, - ) - - if not self.ctx.add_hydrogens or self.ctx.strip_H: - return frames - if not self.pdb["resname"].str.strip().eq("HOH").any(): - return frames - restraints = self.restraints - dictionaries = { - key: value for key, value in restraints.cif_dict.items() if key == "HOH" - } - plan = plan_hydrogens(restraints.topology, dictionaries, self.xyz().detach()) - if plan.n_hydrogens == 0: - return frames - - generated = hydrogen_frames(restraints.topology, plan) - if frames is not None: - # Water groups include both existing and planned H atoms; keep custom - # frames for the rest of the table and give the water groups fresh IDs. - water = self.pdb["resname"].str.strip().eq("HOH").to_numpy() - keep = ~water[frames.parent_row] - take = water[generated.parent_row] - arrays = {} - for field in fields(HydrogenFrames): - existing = getattr(frames, field.name)[keep] - added = getattr(generated, field.name)[take].copy() - if field.name in ("torsion_group", "rotation_group"): - added[added >= 0] += int(existing.max(initial=-1)) + 1 - arrays[field.name] = np.concatenate((existing, added)) - generated = HydrogenFrames(**arrays) - - self.update_pdb() - augmented, old_rows, new_rows = augment_atom_table_with_maps( - self.pdb, plan, restraints.topology - ) - frames = generated.remap(old_rows).fill_planned_rows(new_rows) - # dtype-ok: row index map into the augmented table; torch indexing needs int64 - source = torch.empty(len(augmented), dtype=torch.long, device=self.device) - old_index = torch.as_tensor(old_rows, device=self.device) - new_index = torch.as_tensor(new_rows, device=self.device) - source[old_index] = torch.arange(len(self.pdb), device=self.device) - source[new_index] = torch.as_tensor(plan.parent, device=self.device) - xyz = ( - self.xyz.to_mixed_tensor() - if hasattr(self.xyz, "to_mixed_tensor") - else self.xyz - ) - masks = { - "xyz": xyz.refinable_mask[source], - "occupancy": self.occupancy.get_refinable_atoms()[source], - } - from torchref.model.disorder_field import DisorderFieldTensor - - adp_fields = {} - for name in ("adp", "u"): - wrapper = getattr(self, name) - if isinstance(wrapper, DisorderFieldTensor): - adp_fields[name] = wrapper - else: - masks[name] = wrapper.refinable_mask[source] - if "u" in masks: - masks["u"][new_index] = False - gradients = { - name: getattr(self, name).refinable_params.requires_grad for name in masks - } - adp = self.adp().detach() - cell, spacegroup, links = self.cell, self.spacegroup, self.ctx.links - - def reader(): - return augmented, cell.data.cpu().numpy(), spacegroup - - reader.links = links - strip_h = self.ctx.strip_H - self.ctx.strip_H = False - try: - self.load(reader, add_hydrogens=False) - finally: - self.ctx.strip_H = strip_h - if "adp" in masks: - self.adp[old_index] = adp - for name, mask in masks.items(): - wrapper = getattr(self, name) - wrapper.update_refinable_mask(mask) - wrapper.refinable_params.requires_grad_(gradients[name]) - for name, field in adp_fields.items(): - field.anchor_atom = old_index[field.anchor_atom] - field.neighbor_list = field.neighbor_list[source] - field._full_shape = len(augmented) - field.set_xyz_fn(self.xyz) - setattr(self, name, field) - self._hydrogen_frames = frames - return frames - def set_hydrogen_mode(self, mode: str, frames=None) -> "Model": - """Switch the hydrogen parametrisation of the current atom table. + """Switch how the hydrogen rows of the current atom table are parametrised. Parameters ---------- - mode : str + mode : {"atoms", "riding"} ``"riding"``: hydrogen coordinates derive from their parents each forward; - rotatable groups retain shared torsion or orientation parameters. - Missing HOH hydrogens are completed only when ``ctx.add_hydrogens`` - is True and ``ctx.strip_H`` is False. With hydrogen generation disabled, - the atom table is unchanged. - ``"free"``: hydrogens are ordinary refinable atoms again. - ``"none"`` is a different atom table; use - :meth:`strip_hydrogens`. + rotatable groups keep shared torsion or orientation parameters. + ``"atoms"``: hydrogens are ordinary refinable atoms. frames : HydrogenFrames, optional Riding frames for the current table; default :meth:`hydrogen_frames`. - Water frames are completed and row indices remapped if atoms are added. Returns ------- Model Self, for chaining. + Raises + ------ + ValueError + For an unknown mode, or ``"riding"`` on a model loaded with + ``hydrogens="strip"``. + Notes ----- - Replaces the ``xyz`` wrapper, so any optimizer or ``LossState`` built over the - old parameters is stale; :meth:`Refinement.set_hydrogen_mode` does the - engine-side reset. The refinable set carries over row for row (a hydrogen - released to ``"free"`` follows its parent's mask). Existing atom coordinates - are preserved. Completing waters rebuilds the per-atom wrappers and - restraints; new hydrogens inherit their oxygen's refinement selections. - Water initialization follows ``torch.manual_seed`` and never runs in forward. - """ + The atom table never changes here: hydrogens a table lacks are generated only + at load, with ``hydrogens="add"``. Replaces the ``xyz`` wrapper, so any + optimizer or ``LossState`` built over the old parameters is stale; + :meth:`Refinement.set_hydrogen_mode` does the engine-side reset. The refinable + set carries over row for row (a hydrogen released to ``"atoms"`` follows its + parent's mask). + """ + from torchref.model.context import check_hydrogen_policy from torchref.model.riding_xyz import RidingXYZTensor if not self.ctx.initialized: raise RuntimeError("Load a structure before setting the hydrogen mode.") - if mode == "none": - raise ValueError( - "hydrogen_mode 'none' changes the atom table; use strip_hydrogens()" - ) - if mode not in ("riding", "free"): - raise ValueError(f"unknown hydrogen_mode {mode!r}") + check_hydrogen_policy(self.ctx.hydrogens, mode) - if mode == "riding": - frames = self._complete_riding_waters(frames) - if isinstance(self.xyz, RidingXYZTensor) and frames is None: - return self + riding = isinstance(self.xyz, RidingXYZTensor) + if mode == "atoms": + new_xyz = self.xyz.to_mixed_tensor() if riding else self.xyz + elif riding and frames is None: + new_xyz = self.xyz + else: + base = self.xyz.to_mixed_tensor() if riding else self.xyz if frames is None: frames = self.hydrogen_frames() - if isinstance(self.xyz, RidingXYZTensor): - current = self.xyz.to_mixed_tensor() - else: - current = self.xyz - new_xyz = RidingXYZTensor.from_mixed_tensor(current, frames) - else: - if isinstance(self.xyz, RidingXYZTensor): - frames = self.xyz.hydrogen_frames() - new_xyz = self.xyz.to_mixed_tensor() - else: - new_xyz = self.xyz + new_xyz = RidingXYZTensor.from_mixed_tensor(base, frames) if new_xyz is not self.xyz: # Pop first so the new wrapper registers as a fresh submodule. self._modules.pop("xyz") self.xyz = new_xyz self._repoint_coordinate_accessors() - self._hydrogen_frames = frames self.ctx.hydrogen_mode = mode if hasattr(self, "reset_cache"): self.reset_cache() diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index 0e83cfd8..bd520518 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -13,7 +13,7 @@ import torch from torchref.base.fourier import fft, ifft -from torchref.config import canonical_device, dtypes, get_default_device, get_float_dtype +from torchref.config import dtypes from torchref.model.model import Model from torchref.model.sf_fft import SfFFT from torchref.symmetry import SpaceGroup @@ -183,75 +183,6 @@ def _fingerprint_state(self): """ return super()._fingerprint_state() + (self.fft.grid_key,) - def load_pdb(self, filename): - """ - Load a PDB file and initialize the model with FT-specific setup. - - Parameters - ---------- - filename : str - Path to the PDB file. - - Returns - ------- - ModelFT - Self, for method chaining. - """ - super().load_pdb(filename) - return self - - def select(self, selection): - """ - Return a new ModelFT containing only the selected atoms. - - Extends :meth:`Model.select` with the FT-specific setup: rebuilding - the ITC92 parametrization and carrying ``max_res`` and - ``explicit_gridsize`` across, so the selection sizes its grid the same way. - - Parameters - ---------- - selection : array-like or str - Atom selection forwarded to :meth:`Model.select`. - - Returns - ------- - ModelFT - A new model holding the selected atoms. - - Notes - ----- - ``wavelength`` and ``anomalous_threshold`` are **not** propagated: - :meth:`Model.select` passes only the base kwargs, so the returned model - carries the ModelFT defaults for those. - """ - selection = super().select(selection) - selection._build_parametrization() - selection.max_res = self.max_res - selection.explicit_gridsize = self.explicit_gridsize - return selection - - def load_cif(self, filename): - """ - Load a CIF file and initialize the model with FT-specific setup. - - Parameters - ---------- - filename : str - Path to the CIF/mmCIF file. - - Returns - ------- - ModelFT - Self, for method chaining. - """ - super().load_cif(filename) - self._build_parametrization() - return self - - def _build_parametrization(self): - """Build the ITC92 parametrization (delegates to :class:`Model`).""" - return super()._build_parametrization() - # ========================================================================= # Backward-compatible properties for scattering parameters # ========================================================================= @@ -518,12 +449,6 @@ def get_map_statistics(self): } return stats - def update_pdb(self): - """ - Update PDB with current atomic parameters. - """ - return super().update_pdb() - def reset_cache(self): """Reset SF cache, anomalous cache, and all wrapper forward caches.""" self.reset_forward_cache() @@ -746,80 +671,6 @@ def forward(self, hkl, apply_anomalous: bool = True) -> torch.Tensor: return sf - def copy(self, detach: bool = True) -> "ModelFT": - """ - Create a deep copy of the ModelFT. - - Creates a complete independent copy including all Model base class data, - the grid inputs (``max_res``, ``explicit_gridsize``; the grid itself is - re-derived from the copied context), the ITC92 parametrization, and - scalar attributes. - Cache is reset to empty. - - Parameters - ---------- - detach : bool, optional - If True, the copy's parameters will be detached from the - computation graph (default: True). - Returns - ------- - ModelFT - A new, fully independent ModelFT instance with copied data. - """ - if not self.ctx.initialized: - raise RuntimeError("Cannot copy an uninitialized ModelFT. Load data first.") - - model_copy = ModelFT( - dtype_float=self.dtype_float, - verbose=self.ctx.verbose, - device=self.device, - strip_H=self.ctx.strip_H, - max_res=self.max_res, - gridsize=self.explicit_gridsize, - wavelength=self.wavelength, - anomalous_threshold=self.anomalous_threshold, - ) - - # Carries the atom table, cell, space group, altloc groups and provenance. - model_copy.ctx = self.ctx.copy() - - # Own buffers only; the engine's grid buffers are derived, not copied. - for buffer_name, buffer_value in self._buffers.items(): - if buffer_value is not None: - if detach: - model_copy.register_buffer( - buffer_name, buffer_value.clone().detach() - ) - else: - model_copy.register_buffer(buffer_name, buffer_value.clone()) - - # Parameter wrappers via their own .copy(); the engine came from the ctor. - skip_modules = {"_fft"} - for module_name, module in self._modules.items(): - if module_name in skip_modules: - continue - if module is not None and hasattr(module, "copy"): - setattr(model_copy, module_name, module.copy()) - - if hasattr(self, "_parametrization") and self._parametrization is not None: - import copy as copy_module - - model_copy._parametrization = copy_module.deepcopy(self._parametrization) - - # Borrowed coordinate accessors (restraints, ADP node field) still point at - # THIS model's wrappers after their own ``copy``; re-point them. - model_copy._repoint_coordinate_accessors() - - # Don't share cached structure factors with the original. - model_copy.reset_cache() - # The iso/aniso partition is derived state, not a buffer, so it is not - # carried by the buffer loop above; get_iso()/get_aniso() read it. - - if self.ctx.verbose > 0: - print(f"✓ ModelFT copied successfully ({len(model_copy.pdb)} atoms)") - - return model_copy - def state_dict(self, destination=None, prefix="", keep_vars=False): """ Return a dictionary containing the complete state of the ModelFT. @@ -857,163 +708,62 @@ def state_dict(self, destination=None, prefix="", keep_vars=False): # _cache, _anomalous_cache (from the element list). return state - @classmethod - def create_from_state_dict( - cls, - state_dict: dict, - device: torch.device = None, - verbose: int = 1, - dtype_float: torch.dtype = None, - ) -> "ModelFT": - """ - Create a fully initialized ModelFT from a state dictionary. - - This is the recommended way to restore a ModelFT from a saved state. - Creates an instance with properly initialized submodules, then loads the state. - - Parameters - ---------- - state_dict : dict - State dictionary from torch.save(model.state_dict(), ...). - device : torch.device, optional - Move the restored model here once it is built. The restore itself always - runs on CPU; ``None`` then moves it to the configured default device; see - :meth:`Model.create_from_state_dict`. - verbose : int, optional - Verbosity level. Default is 1. - dtype_float : torch.dtype, optional - Float dtype for tensors. Default is dtypes.float. + def _subclass_kwargs(self) -> dict: + """The grid, wavelength and Bijvoet settings a new instance must share.""" + return { + "max_res": self.max_res, + "gridsize": self.explicit_gridsize, + "wavelength": self.wavelength, + "anomalous_threshold": self.anomalous_threshold, + "apply_bijvoet": bool(self.anomalous_bijvoet), + } - Returns - ------- - ModelFT - Fully initialized instance with restored state. + @classmethod + def _pop_subclass_state(cls, state_dict: dict) -> dict: + """Pop the FT settings :meth:`state_dict` wrote, as constructor kwargs. - Notes - ----- - Legacy state_dicts are accepted: the obsolete ``radius_angstrom`` key is - ignored and old-style ``A`` / ``B`` buffers are remapped to ``_A`` / ``_B``. - The anisotropic ``u`` is rebuilt as a :class:`CholeskyMixedTensor`, as in - :meth:`load`, so the positive-definite parametrization round-trips. + The ``radius_angstrom`` key of older checkpoints is dropped unused. """ - # Build on CPU throughout and move once at the end, as Model does; the grid - # setup below otherwise sizes an accelerator allocation before the model is - # placed. The final target is the caller's device, or the configured default - # when they name none, so a restore lands beside a same-config model. - target_device = ( - canonical_device(device) if device is not None else get_default_device() - ) - device = torch.device("cpu") - if dtype_float is None: - dtype_float = get_float_dtype() - - max_res = state_dict.pop("max_res", 1.0) - explicit_gridsize = state_dict.pop("explicit_gridsize", None) - state_dict.pop("radius_angstrom", None) # legacy key, no longer used - wavelength = state_dict.pop("wavelength", 1.0) - anomalous_threshold = state_dict.pop("anomalous_threshold", 0.5) - - pdb = state_dict.pop("pdb", None) - spacegroup_str = state_dict.pop("spacegroup", None) - cell_tensor = state_dict.pop("cell", None) - initialized = state_dict.pop("initialized", False) - saved_dtype = state_dict.pop("dtype_float", dtype_float) - state_dict.pop("device", None) # Remove but don't use (use provided device) - strip_H = state_dict.pop("strip_H", True) - cif_path = state_dict.pop("cif_path", None) - altloc_pairs = state_dict.pop("altloc_pairs", []) - hydrogens_in_xray = state_dict.pop("hydrogens_in_xray", True) - hydrogen_mode = state_dict.pop("hydrogen_mode", None) - - # Checkpoints written while the grid was stored state carry its buffers - # ("_fft." prefixed, or flat in older ones). The size is adopted below only - # when it differs from what the crystal and max_res give. - legacy_gridsize = state_dict.pop("_fft.gridsize", None) - if legacy_gridsize is None: - legacy_gridsize = state_dict.pop("gridsize", None) - state_dict.pop("_fft.voxel_size", None) - state_dict.pop("voxel_size", None) - - instance = cls( - dtype_float=saved_dtype, - verbose=verbose, - device=device, - strip_H=strip_H, - cif_path=cif_path, - max_res=max_res, - gridsize=explicit_gridsize, - wavelength=wavelength, - anomalous_threshold=anomalous_threshold, - hydrogens_in_xray=hydrogens_in_xray, - ) - - instance.pdb = pdb - instance.ctx.initialized = initialized - instance.ctx.altloc_pairs = altloc_pairs - if hydrogen_mode is None: - hydrogen_mode = "riding" if state_dict.get("xyz.h_row") is not None else "free" - instance.ctx.hydrogen_mode = hydrogen_mode - - # The engine reads both off the context; nothing further to build. - instance.spacegroup = spacegroup_str - - from torchref.symmetry import Cell - - if cell_tensor is not None: - instance.cell = Cell(cell_tensor, dtype=saved_dtype, device=device) - - # Wrappers and per-atom buffers: shared with Model so the two restores cannot - # drift apart again. ModelFT adds only its own scattering buffers below. - if pdb is not None: - cls._rebuild_wrappers_from_pdb( - instance, pdb, state_dict, saved_dtype, device - ) - - # Scattering buffers: accept both old-style (A, B) and new (_A, _B). - a_key = "_A" if "_A" in state_dict else "A" if "A" in state_dict else None - b_key = "_B" if "_B" in state_dict else "B" if "B" in state_dict else None + state_dict.pop("radius_angstrom", None) + return { + "max_res": state_dict.pop("max_res", 1.0), + "gridsize": state_dict.pop("explicit_gridsize", None), + "wavelength": state_dict.pop("wavelength", 1.0), + "anomalous_threshold": state_dict.pop("anomalous_threshold", 0.5), + } - if a_key and state_dict[a_key] is not None: - instance.register_buffer( - "_A", torch.zeros_like(state_dict[a_key], device=device) - ) - if b_key and state_dict[b_key] is not None: - instance.register_buffer( - "_B", torch.zeros_like(state_dict[b_key], device=device) + def _restorable_entries(self, state_dict: dict) -> dict: + """Register the scattering buffers and adopt a legacy stored grid size. + + Old checkpoints name the scattering buffers ``A`` / ``B`` rather than + ``_A`` / ``_B``, and those written while the grid was stored state carry its + size (``_fft.gridsize``, or a flat ``gridsize``). That size is adopted only + when it differs from what the crystal and ``max_res`` give. + """ + for old, new in (("A", "_A"), ("B", "_B")): + if old in state_dict and new not in state_dict: + state_dict[new] = state_dict.pop(old) + for name in ("_A", "_B"): + if state_dict.get(name) is not None and self.pdb is not None: + self.register_buffer( + name, torch.zeros_like(state_dict[name], device=self.device) ) + legacy = state_dict.pop("_fft.gridsize", None) + if legacy is None: + legacy = state_dict.pop("gridsize", None) + state_dict.pop("_fft.voxel_size", None) + state_dict.pop("voxel_size", None) if ( - legacy_gridsize is not None - and explicit_gridsize is None - and instance.ctx.crystal_key is not None - and instance.max_res is not None + legacy is not None + and self.explicit_gridsize is None + and self.ctx.crystal_key is not None + and self.max_res is not None ): - if isinstance(legacy_gridsize, torch.Tensor): - legacy_gridsize = legacy_gridsize.tolist() - legacy = tuple(int(x) for x in legacy_gridsize) - if legacy != instance.fft.compute_optimal_gridsize(instance.max_res): - instance.explicit_gridsize = legacy - - # Drop empty placeholders, remapping old-style A/B keys to _A/_B. - filtered_state_dict = {} - for k, v in state_dict.items(): - if not hasattr(v, "shape") or v.numel() > 0: - if k == "A": - filtered_state_dict["_A"] = v - elif k == "B": - filtered_state_dict["_B"] = v - else: - filtered_state_dict[k] = v - - instance.load_state_dict(filtered_state_dict, strict=False) - - # Always placed: target_device is the caller's device or the configured default. - instance.to(target_device) - - instance.reset_cache() - - if verbose > 0: - n_atoms = len(instance.pdb) if instance.pdb is not None else 0 - print(f"Created ModelFT from state_dict: {n_atoms} atoms") - - return instance + if isinstance(legacy, torch.Tensor): + legacy = legacy.tolist() + legacy = tuple(int(x) for x in legacy) + if legacy != self.fft.compute_optimal_gridsize(self.max_res): + self.explicit_gridsize = legacy + return super()._restorable_entries(state_dict) + diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index 27d85c6b..67028116 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -137,7 +137,8 @@ def __init__( shrink: bool = SHRINK_ENABLED, scale_target: str = DEFAULT_SCALE_TARGET, aniso_selection: Optional[str] = None, - add_hydrogens: bool = False, + hydrogens: str = "keep", + hydrogen_mode: str = "atoms", hydrogens_in_xray: bool = True, ): """Initialize Refinement, fully if ``data_file`` and ``pdb`` are given. @@ -207,9 +208,11 @@ def __init__( aniso_selection : str, optional Phenix-style selection of atoms refined anisotropically when ``adp_mode="anisotropic"``. Defaults to all non-water heavy atoms. - add_hydrogens : bool, optional - Generate missing hydrogens when loading the model. Default False. - Hydrogens already present in the input are retained either way. + hydrogens : {"keep", "add", "strip"}, optional + What loading the model does with its hydrogens: keep the file's (default), + also generate the missing ones, or remove them all. + hydrogen_mode : {"atoms", "riding"}, optional + Hydrogens as refinable atoms (default) or riding on their parents. hydrogens_in_xray : bool, optional Whether hydrogens contribute to the structure factors. Default True. They take part in the restraints either way. @@ -283,7 +286,8 @@ def __init__( device=self.device, wavelength=self.wavelength, anomalous_threshold=self.anomalous_threshold, - add_hydrogens=add_hydrogens, + hydrogens=hydrogens, + hydrogen_mode=hydrogen_mode, cif_path=cif, hydrogens_in_xray=hydrogens_in_xray, ) @@ -332,7 +336,8 @@ def __init__( device=self.device, wavelength=self.wavelength, anomalous_threshold=self.anomalous_threshold, - add_hydrogens=add_hydrogens, + hydrogens=hydrogens, + hydrogen_mode=hydrogen_mode, hydrogens_in_xray=hydrogens_in_xray, # Before load, not after: generation on load reads this dictionary. cif_path=cif, @@ -758,8 +763,8 @@ def set_hydrogen_mode(self, mode: str) -> "Refinement": Parameters ---------- - mode : str - ``"riding"`` or ``"free"``; see :meth:`Model.set_hydrogen_mode`. + mode : {"atoms", "riding"} + See :meth:`Model.set_hydrogen_mode`. Returns ------- @@ -772,13 +777,7 @@ def set_hydrogen_mode(self, mode: str) -> "Refinement": ``LossState`` are dropped and rebuilt on the next step. Call between macro cycles, never inside one. """ - n_atoms = len(self.model.pdb) self.model.set_hydrogen_mode(mode) - if ( - len(self.model.pdb) != n_atoms - and getattr(self, "adp_target", None) is not None - ): - self._init_targets() persistent = getattr(self, "_persistent_optimizers", None) if persistent is not None: persistent.clear() From d2c5ca3bb7e16374633e9f9a665566cf320f769b Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 22:06:17 +0200 Subject: [PATCH 198/250] Give the topology an identity layer that needs no dictionaries Topology.from_table builds a node-only topology from an atom table's identity columns (names, elements, altlocs, residues, chains, record types, charges), with connected=False until the restraint build adds edges. It selects atoms (Topology.select, over a recursive-descent parser in utils.selection that parse_phenix_selection now shares), exposes water/polymer masks and cached atomic numbers and vdW radii, and inserts a hydrogen plan row for row as the table-level insertion does. The old parser could resolve a selection with two parenthesised groups to every atom; the shared one does not. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 1 + tests/unit/topology/test_topology_identity.py | 138 ++++++++ torchref/topology/__init__.py | 9 +- torchref/topology/atom_graph.py | 73 ++++- torchref/topology/build.py | 2 +- torchref/topology/nonbonded.py | 8 +- torchref/topology/residue_graph.py | 26 +- torchref/topology/topology.py | 299 +++++++++++++++++- torchref/utils/selection.py | 144 +++++++++ torchref/utils/utils.py | 170 +--------- 10 files changed, 691 insertions(+), 179 deletions(-) create mode 100644 tests/unit/topology/test_topology_identity.py create mode 100644 torchref/utils/selection.py diff --git a/docs/changelog.rst b/docs/changelog.rst index f7a6f115..214c993b 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Atom identity is available without building restraints: ``Topology.from_table`` gives a node-only topology (names, elements, altlocs, residues, chains, record types and charges; ``connected=False``), with ``Topology.select`` for Phenix-style selections, ``is_water`` / ``is_polymer`` masks, cached ``AtomGraph.atomic_number`` / ``vdw_radii`` and ``Topology.with_hydrogens`` for inserting a hydrogen plan. Selections are evaluated by one recursive-descent parser (``torchref.utils.selection``); a selection with more than one parenthesised group, which the old parser could resolve to every atom, now selects what it says - Hydrogens are one policy on the model context, set at construction: ``hydrogens`` (``keep`` / ``add`` / ``strip``) settles the atom table when it loads and ``hydrogen_mode`` (``atoms`` / ``riding``) how hydrogen rows are parametrised; ``riding`` with ``strip`` raises. It replaces ``strip_H``, ``add_hydrogens`` and ``hydrogen_mode="free"`` on ``Model``, ``ModelFT``, ``Refinement``, ``EnsembleModel`` and ``cli._common.load_model`` (checkpoints written with them still restore), and ``torchref.refine`` takes ``--hydrogens`` / ``--hydrogen-mode`` instead of ``--add-hydrogens``. The deprecated ``exclude_H_from_sf`` alias is removed - Missing water hydrogens are generated at load with ``hydrogens="add"``, never by ``set_hydrogen_mode("riding")``, which now only swaps the coordinate wrapper and leaves the atom table alone; ``hydrogenate()`` no longer takes ``optimize`` (free torsions are always scanned) - ``Model.select`` keeps the anisotropic ``u`` in its Cholesky parametrization, and ``select``/``copy``/``create_from_state_dict`` and the strip/hydrogenate helpers are shared by ``Model`` and ``ModelFT``, so ``copy`` of a ``ModelFT`` carries ``apply_bijvoet`` diff --git a/tests/unit/topology/test_topology_identity.py b/tests/unit/topology/test_topology_identity.py new file mode 100644 index 00000000..0ba0ea21 --- /dev/null +++ b/tests/unit/topology/test_topology_identity.py @@ -0,0 +1,138 @@ +"""The identity half of a topology: built from an atom table, no dictionaries needed. + +A node-only topology must describe the same atoms and residues as the connected one the +restraint build produces, select the same atoms a direct reading of the table would, +and insert hydrogens exactly where the table-level insertion does. +""" + +import numpy as np +import pytest +import torch + +from torchref.io.pdb import PDBReader +from torchref.model.model import Model +from torchref.topology.hydrogens import augment_atom_table_with_maps, plan_hydrogens +from torchref.topology.topology import IDENTITY_COLUMNS, Topology + +STRUCTURES = ["1DAW", "7L84", "1AK5_with_H"] + + +def _table(pdb_dir, code): + df, _, _ = PDBReader(verbose=0).read(str(pdb_dir / f"{code}.pdb"))() + return df.dropna(subset=["x", "y", "z", "tempfactor", "occupancy"]).reset_index( + drop=True + ) + + +@pytest.mark.unit +@pytest.mark.parametrize("code", STRUCTURES) +def test_node_topology_matches_the_connected_one(pdb_dir, code): + """Atoms and residues agree with the restraint build's topology, edges aside.""" + model = Model(verbose=0).load_pdb(str(pdb_dir / f"{code}.pdb")) + connected = model.restraints.topology + nodes = Topology.from_table(model.pdb) + + assert connected.connected and not nodes.connected + assert nodes.atoms.bonds.n_edges == 0 + for field in ("name", "element", "altloc"): + np.testing.assert_array_equal( + getattr(nodes.atoms, field), getattr(connected.atoms, field) + ) + assert torch.equal(nodes.atoms.residue_of.cpu(), connected.atoms.residue_of.cpu()) + for field in ("chain", "resseq", "icode", "resname", "atom_start", "atom_end"): + np.testing.assert_array_equal( + getattr(nodes.residues, field), getattr(connected.residues, field) + ) + + +@pytest.mark.unit +@pytest.mark.parametrize("code", STRUCTURES) +def test_columns_round_trip(pdb_dir, code): + nodes = Topology.from_table(_table(pdb_dir, code)) + again = Topology.from_columns(nodes.columns()) + for key, value in nodes.columns().items(): + np.testing.assert_array_equal(again.columns()[key], value) + assert set(nodes.columns()) == set(IDENTITY_COLUMNS) + + +@pytest.mark.unit +@pytest.mark.parametrize("code", STRUCTURES) +def test_selection_matches_the_table(pdb_dir, code): + """Every keyword and operator, checked against the same condition on the table.""" + df = _table(pdb_dir, code) + nodes = Topology.from_table(df) + water = df.resname == "HOH" + blank = df.altloc.astype(str).str.strip() == "" + cases = { + "all": np.ones(len(df), bool), + "chain A": df.chainid == "A", + "resseq 10": df.resseq == 10, + "resseq 10:40": df.resseq.between(10, 40), + "resname hoh": water, + "name ca": df.name == "CA", + "element c": df.element.str.strip().str.capitalize() == "C", + "altloc A": df.altloc == "A", + "not resname HOH": ~water, + "chain A and not resname HOH": (df.chainid == "A") & ~water, + "name CA or name CB": df.name.isin(["CA", "CB"]), + "not (resname HOH or element H)": ~(water | (df.element.str.strip() == "H")), + "(chain A and resseq 1:50) or (resname HOH and not name O)": ( + (df.chainid == "A") & df.resseq.between(1, 50) + ) + | (water & (df.name != "O")), + "NOT resname HOH AND name CA": ~water & (df.name == "CA"), + "resseq 1:20 or resseq 30:40 and name N": df.resseq.between(1, 20) + | (df.resseq.between(30, 40) & (df.name == "N")), + "not not name CA": df.name == "CA", + } + for selection, expected in cases.items(): + got = nodes.select(selection).numpy() + np.testing.assert_array_equal(got, np.asarray(expected), err_msg=selection) + assert not nodes.select("altloc A").numpy()[blank.to_numpy()].any() + + +@pytest.mark.unit +@pytest.mark.parametrize( + "selection", ["", "chain", "resname HOH and", "(name CA", "name CA)", "bogus X"] +) +def test_malformed_selections_raise(pdb_dir, selection): + nodes = Topology.from_table(_table(pdb_dir, "1DAW")) + with pytest.raises(ValueError): + nodes.select(selection) + + +@pytest.mark.unit +def test_water_and_polymer_masks(pdb_dir): + df = _table(pdb_dir, "1DAW") + nodes = Topology.from_table(df) + np.testing.assert_array_equal(nodes.is_water, (df.resname == "HOH").to_numpy()) + np.testing.assert_array_equal(nodes.is_polymer, (df.ATOM == "ATOM").to_numpy()) + np.testing.assert_array_equal( + nodes.atoms.is_hydrogen.cpu().numpy(), (df.element.str.strip() == "H").to_numpy() + ) + + +@pytest.mark.unit +@pytest.mark.parametrize("code", ["1DAW", "7L84"]) +def test_hydrogen_insertion_matches_the_table_insertion(pdb_dir, code): + """Same row maps and the same identity, row for row, as the table-level insertion.""" + model = Model(verbose=0, hydrogens="strip").load_pdb(str(pdb_dir / f"{code}.pdb")) + restraints = model.ctx.build_restraints(model.xyz(), nonbonded=False, verbose=0) + plan = plan_hydrogens(restraints.topology, restraints.cif_dict, model.xyz().detach()) + assert plan.n_hydrogens > 0 + + augmented, old_to_new, plan_to_new = augment_atom_table_with_maps( + model.pdb, plan, restraints.topology + ) + nodes, source, old_to_new_t, plan_to_new_t = Topology.from_table( + model.pdb + ).with_hydrogens(plan) + + np.testing.assert_array_equal(old_to_new_t, old_to_new) + np.testing.assert_array_equal(plan_to_new_t, plan_to_new) + np.testing.assert_array_equal(source[old_to_new], np.arange(len(model.pdb))) + np.testing.assert_array_equal(source[plan_to_new], plan.parent) + expected = Topology.from_table(augmented) + for key, value in expected.columns().items(): + np.testing.assert_array_equal(nodes.columns()[key], value, err_msg=key) + np.testing.assert_array_equal(nodes.residues.atom_start, expected.residues.atom_start) diff --git a/torchref/topology/__init__.py b/torchref/topology/__init__.py index 9fc6c49c..d698fcd9 100644 --- a/torchref/topology/__init__.py +++ b/torchref/topology/__init__.py @@ -1,5 +1,10 @@ """Model topology as a graph: residues over atoms, connectivity over restraints. +The topology is where a model's atom identity lives -- names, elements, altlocs, +residues, chains -- from the moment its atom table is read +(:meth:`Topology.from_table`, :meth:`Topology.select`). Connectivity is added later, +against the monomer dictionaries. + :class:`Topology` holds two levels. :class:`ResidueGraph` is the sequence -- residues as template instances, inter-residue links as edges. :class:`AtomGraph` is the expansion -- atoms as nodes, typed :class:`EdgeBlock` sets over them, and a CSR bond adjacency that @@ -38,10 +43,12 @@ ) from .restraint_sets import assemble_entries, max_period from .templates import resolve_template_keys -from .topology import Topology +from .topology import IDENTITY_COLUMNS, Topology, identity_columns __all__ = [ "Topology", + "identity_columns", + "IDENTITY_COLUMNS", "Restraints", "ResidueGraph", "AtomGraph", diff --git a/torchref/topology/atom_graph.py b/torchref/topology/atom_graph.py index a38eb5be..8d20d896 100644 --- a/torchref/topology/atom_graph.py +++ b/torchref/topology/atom_graph.py @@ -9,6 +9,10 @@ Every indexing structure here is a tensor, so it moves with ``.to(device)`` alongside the edge blocks. Only the per-atom identifiers are NumPy, because they are strings. + +The identity (names, elements, altlocs, record type, charge, residue membership) exists +from the moment an atom table is read; the edge blocks stay empty until the graph is +connected against the monomer dictionaries. """ from dataclasses import dataclass, field @@ -110,11 +114,16 @@ class AtomGraph(DeviceMixin): name, element, altloc : numpy.ndarray Per-atom identifiers, shape ``(N,)``. Strings, so NumPy rather than tensors; residue-level identity is reached through ``residue_of`` rather than duplicated - here. + here. ``altloc`` is ``' '`` for atoms in no alternative conformation. residue_of : torch.Tensor Residue index per atom, shape ``(N,)``, dtype ``int64``. - bonds, angles, torsions, chirals : EdgeBlock - Typed edge blocks. ``bonds`` also backs the adjacency. + is_hetatm : numpy.ndarray, optional + True for HETATM records, shape ``(N,)``. Defaults to all False. + charge : numpy.ndarray, optional + Formal charge per atom, shape ``(N,)``, integer. Defaults to zeros. + bonds, angles, torsions, chirals : EdgeBlock, optional + Typed edge blocks, empty until the graph is connected. ``bonds`` also backs the + adjacency. planes : dict ``{n_atoms_in_plane: EdgeBlock}`` -- planes are ragged, so they are grouped by atom count the way the plane restraints already are. @@ -141,10 +150,12 @@ class AtomGraph(DeviceMixin): element: np.ndarray altloc: np.ndarray residue_of: torch.Tensor - bonds: EdgeBlock - angles: EdgeBlock - torsions: EdgeBlock - chirals: EdgeBlock + is_hetatm: Optional[np.ndarray] = None + charge: Optional[np.ndarray] = None + bonds: Optional[EdgeBlock] = None + angles: Optional[EdgeBlock] = None + torsions: Optional[EdgeBlock] = None + chirals: Optional[EdgeBlock] = None planes: Dict[int, EdgeBlock] = field(default_factory=dict) energy_type: Optional[np.ndarray] = None template_h_count: Optional[torch.Tensor] = None @@ -154,8 +165,19 @@ class AtomGraph(DeviceMixin): _adj_indices: Optional[torch.Tensor] = field(default=None, repr=False) # (element array it was parsed from, hydrogen flags); see is_hydrogen. _is_h_cache: Optional[Tuple[np.ndarray, np.ndarray]] = field(default=None, repr=False) + # (element array, symbols, atomic numbers, van der Waals radii); see _element_table. + _element_cache: Optional[Tuple[np.ndarray, ...]] = field(default=None, repr=False) def __post_init__(self) -> None: + n = len(self.name) + device = self.residue_of.device + if self.is_hetatm is None: + self.is_hetatm = np.zeros(n, dtype=bool) + if self.charge is None: + self.charge = np.zeros(n, dtype=np.int64) + for edge, arity in (("bonds", 2), ("angles", 3), ("torsions", 4), ("chirals", 4)): + if getattr(self, edge) is None: + setattr(self, edge, EdgeBlock.empty(arity, device=device)) if self._adj_indptr is None: self.rebuild_adjacency() @@ -184,6 +206,39 @@ def is_hydrogen(self) -> torch.Tensor: cache = self._is_h_cache = (self.element, flags) return torch.tensor(cache[1], device=self.bonds.indices.device) + def _element_table(self) -> Tuple[np.ndarray, ...]: + """``(symbols, atomic numbers, vdW radii)``, parsed once per ``element`` array.""" + cache = self._element_cache + if cache is None or cache[0] is not self.element: + import gemmi + + from torchref.topology.nonbonded import vdw_radii_for_elements + + symbols = np.char.capitalize(np.char.strip(self.element.astype(str))) + numbers = np.array([gemmi.Element(s).atomic_number for s in symbols]) + cache = self._element_cache = ( + self.element, + symbols, + numbers.astype(np.int64), + vdw_radii_for_elements(symbols), + ) + return cache[1:] + + @property + def symbols(self) -> np.ndarray: + """Element symbols normalised to ``'C'``, ``'Fe'``, ..., shape ``(N,)``.""" + return self._element_table()[0] + + @property + def atomic_number(self) -> np.ndarray: + """Atomic number per atom, shape ``(N,)``, int64; 0 for an unknown element.""" + return self._element_table()[1] + + @property + def vdw_radii(self) -> np.ndarray: + """Van der Waals radius per atom in Å, shape ``(N,)``, float64.""" + return self._element_table()[2] + def copy(self) -> "AtomGraph": """An independent copy sharing no storage with this one.""" return AtomGraph( @@ -191,6 +246,8 @@ def copy(self) -> "AtomGraph": element=self.element.copy(), altloc=self.altloc.copy(), residue_of=self.residue_of.clone(), + is_hetatm=self.is_hetatm.copy(), + charge=self.charge.copy(), bonds=self.bonds.copy(), angles=self.angles.copy(), torsions=self.torsions.copy(), @@ -256,6 +313,8 @@ def subset(self, remap: torch.Tensor, residue_remap: torch.Tensor) -> "AtomGraph element=self.element[keep], altloc=self.altloc[keep], residue_of=residue_remap[self.residue_of[keep_t]], + is_hetatm=self.is_hetatm[keep], + charge=self.charge[keep], bonds=self.bonds.subset(remap), angles=self.angles.subset(remap), torsions=self.torsions.subset(remap), diff --git a/torchref/topology/build.py b/torchref/topology/build.py index 9c1f6519..936d0adb 100644 --- a/torchref/topology/build.py +++ b/torchref/topology/build.py @@ -1020,7 +1020,7 @@ def build_topology_with_values( "chiral": chiral_values.get("intra", {}), "plane": plane_values, } - return Topology(residues=residues, atoms=atoms), values, extras + return Topology(residues=residues, atoms=atoms, connected=True), values, extras __all__ = ["build_topology", "build_topology_with_values"] diff --git a/torchref/topology/nonbonded.py b/torchref/topology/nonbonded.py index 25a658e7..7fe74309 100644 --- a/torchref/topology/nonbonded.py +++ b/torchref/topology/nonbonded.py @@ -20,8 +20,6 @@ from torchref.config import dtypes, get_float_dtype if TYPE_CHECKING: - import pandas - from torchref.symmetry.cell import Cell from torchref.symmetry.spacegroup import SpaceGroup @@ -29,12 +27,12 @@ _DEFAULT_VDW_RADIUS = 1.9 -def vdw_radii_for_elements(elements: "pandas.Series") -> np.ndarray: +def vdw_radii_for_elements(elements) -> np.ndarray: """Van der Waals radius of each atom, looked up by element. Parameters ---------- - elements : pandas.Series + elements : array-like of str Element symbols, one per atom; case and surrounding whitespace are ignored. Returns @@ -58,7 +56,7 @@ def vdw_radii_for_elements(elements: "pandas.Series") -> np.ndarray: table["vdW_Radius_Angstrom"], ) ) - symbols = elements.astype(str).str.strip().str.capitalize() + symbols = np.char.capitalize(np.char.strip(np.asarray(elements).astype(str))) return np.array( [radius.get(e, _DEFAULT_VDW_RADIUS) for e in symbols], dtype=np.float64 ) diff --git a/torchref/topology/residue_graph.py b/torchref/topology/residue_graph.py index 0a68a21c..9e768c34 100644 --- a/torchref/topology/residue_graph.py +++ b/torchref/topology/residue_graph.py @@ -10,11 +10,16 @@ """ from dataclasses import dataclass, field -from typing import Dict, List, Sequence, Tuple +from typing import Dict, List, Optional, Sequence, Tuple import numpy as np import torch +#: Residue names treated as water, whichever naming convention the file follows. +WATER_RESNAMES = frozenset( + {"HOH", "WAT", "DOD", "H2O", "SOL", "TIP", "TIP3", "TIP4"} +) + #: SG-SG separation below which two cysteines are taken to be disulfide-bonded. DISULFIDE_MAX_DISTANCE = 2.5 @@ -60,11 +65,12 @@ class ResidueGraph: ---------- chain, resseq, icode, resname : numpy.ndarray Per-residue identity, shape ``(R,)``. - template_key : numpy.ndarray - Restraint-dictionary key per residue, shape ``(R,)``. Either the residue name - or a link-modified variant such as ``'ALA:DEL-HN1+DEL-OXT'``. atom_start, atom_end : numpy.ndarray Half-open row range of each residue's atoms, shape ``(R,)``. + template_key : numpy.ndarray, optional + Restraint-dictionary key per residue, shape ``(R,)``. Either the residue name + or a link-modified variant such as ``'ALA:DEL-HN1+DEL-OXT'``. Defaults to the + residue name until the graph is connected. link_pairs : numpy.ndarray Residue index pairs, shape ``(L, 2)``. For a peptide link the first entry donates its ``C`` and the second its ``N``. @@ -77,14 +83,23 @@ class ResidueGraph: resseq: np.ndarray icode: np.ndarray resname: np.ndarray - template_key: np.ndarray atom_start: np.ndarray atom_end: np.ndarray + template_key: Optional[np.ndarray] = None link_pairs: np.ndarray = field( default_factory=lambda: np.zeros((0, 2), dtype=np.int64) ) link_kind: np.ndarray = field(default_factory=lambda: np.zeros(0, dtype=" None: + if self.template_key is None: + self.template_key = np.asarray(self.resname, dtype=object).copy() + + @property + def is_water(self) -> np.ndarray: + """True for water residues (:data:`WATER_RESNAMES`), shape ``(R,)``.""" + return np.isin(np.char.strip(self.resname.astype(str)), list(WATER_RESNAMES)) + @property def n_residues(self) -> int: """Number of residue nodes.""" @@ -293,4 +308,5 @@ def find_disulfide_links( "find_disulfide_links", "DISULFIDE_MAX_DISTANCE", "DISULFIDE_MIN_DISTANCE", + "WATER_RESNAMES", ] diff --git a/torchref/topology/topology.py b/torchref/topology/topology.py index 51419991..44b1a2cd 100644 --- a/torchref/topology/topology.py +++ b/torchref/topology/topology.py @@ -1,15 +1,22 @@ """The topology container: a residue graph over an atom graph. -:class:`Topology` is what the model's connectivity lives in. The residue level carries -the sequence and the inter-residue links; the atom level carries the atoms, the typed -edge blocks and the bond adjacency. Per-atom residue identity is reached through -``atoms.residue_of`` rather than duplicated per atom. +:class:`Topology` is where a model's atom identity and connectivity live. The residue +level carries the sequence and the inter-residue links; the atom level carries the +atoms, the typed edge blocks and the bond adjacency. Per-atom residue identity is +reached through ``atoms.residue_of`` rather than duplicated per atom. + +Identity comes first. :meth:`Topology.from_table` is the one place an atom table's +identity columns become arrays; the result is a node-only topology -- names, elements, +altlocs, residues, chains, record types -- with empty edge blocks and +``connected=False``. Connecting it against the monomer dictionaries is the restraint +build's job. Refinable values (coordinates, B-factors, occupancies) never live here; +they belong to the model's parameter wrappers. Mutable by design; prefer :meth:`Topology.copy` over editing in place. """ from dataclasses import dataclass -from typing import Dict, Set, Tuple +from typing import TYPE_CHECKING, Dict, Mapping, Set, Tuple import numpy as np import torch @@ -18,6 +25,65 @@ from torchref.topology.residue_graph import ResidueGraph from torchref.utils.device_mixin import DeviceMixin +if TYPE_CHECKING: + import pandas + +#: Per-atom identity columns, as :meth:`Topology.columns` returns them. +IDENTITY_COLUMNS = ( + "name", + "element", + "altloc", + "chain", + "resseq", + "icode", + "resname", + "is_hetatm", + "charge", +) + + +def identity_columns(pdb: "pandas.DataFrame") -> Dict[str, np.ndarray]: + """The identity half of an atom table, as per-atom arrays. + + Parameters + ---------- + pdb : pandas.DataFrame + Atom table with ``name``, ``chainid``, ``resseq`` and ``resname`` columns; + ``element``, ``altloc``, ``icode``, ``ATOM`` and ``charge`` are optional. + + Returns + ------- + dict + :data:`IDENTITY_COLUMNS`, each shape ``(N,)``. A blank altloc reads as + ``' '``; a missing or non-numeric charge as 0. + """ + import pandas as pd + + n = len(pdb) + + def text(column: str, default: str) -> np.ndarray: + if column in pdb.columns: + return pdb[column].values.astype(str) + return np.full(n, default) + + altloc = text("altloc", "") + charge = ( + pd.to_numeric(pdb["charge"], errors="coerce").fillna(0).to_numpy() + if "charge" in pdb.columns + else np.zeros(n) + ) + return { + "name": text("name", ""), + "element": text("element", ""), + "altloc": np.where(np.char.strip(altloc) == "", " ", altloc), + "chain": text("chainid", ""), + "resseq": pdb["resseq"].values.astype(np.int64), + "icode": text("icode", ""), + "resname": text("resname", ""), + "is_hetatm": text("ATOM", "ATOM") == "HETATM", + "charge": charge.astype(np.int64), + } + @dataclass(eq=False, repr=False) class Topology(DeviceMixin): @@ -29,6 +95,9 @@ class Topology(DeviceMixin): Sequence level -- residues as nodes, links as edges. atoms : AtomGraph Atom level -- atoms as nodes, typed edge blocks, bond adjacency. + connected : bool, default False + Whether the edge blocks and residue links have been built. A node-only topology + (from :meth:`from_table`) answers every identity question but has no edges. Notes ----- @@ -39,6 +108,217 @@ class Topology(DeviceMixin): residues: ResidueGraph atoms: AtomGraph + connected: bool = False + + # ------------------------------------------------------------------ + # Identity + # ------------------------------------------------------------------ + + @classmethod + def from_table(cls, pdb: "pandas.DataFrame", device=None) -> "Topology": + """A node-only topology over an atom table's identity columns. + + Parameters + ---------- + pdb : pandas.DataFrame + Atom table; see :func:`identity_columns`. Row order is atom order. + device : torch.device, optional + Where ``residue_of`` and the (empty) edge blocks live. + + Returns + ------- + Topology + ``connected=False``. + """ + return cls.from_columns(identity_columns(pdb), device=device) + + @classmethod + def from_columns(cls, columns: Mapping[str, np.ndarray], device=None) -> "Topology": + """A node-only topology over per-atom identity arrays. + + Parameters + ---------- + columns : mapping + :data:`IDENTITY_COLUMNS`, each shape ``(N,)``. Residues are the contiguous + runs of ``(chain, resseq, icode)``. + device : torch.device, optional + + Returns + ------- + Topology + ``connected=False``. + """ + from torchref.topology.residue_graph import build_residue_nodes + + nodes = build_residue_nodes( + columns["chain"], columns["resseq"], columns["icode"], columns["resname"] + ) + n_residues = len(nodes["chain"]) + residues = ResidueGraph( + chain=nodes["chain"], + resseq=nodes["resseq"], + icode=nodes["icode"], + resname=nodes["resname"], + atom_start=nodes["atom_start"], + atom_end=nodes["atom_end"], + ) + residue_of = torch.as_tensor( + np.repeat( + np.arange(n_residues, dtype=np.int64), + nodes["atom_end"] - nodes["atom_start"], + ), + dtype=torch.int64, # dtype-ok: residue index per atom; int64 index required + device=device, + ) + atoms = AtomGraph( + name=np.asarray(columns["name"]), + element=np.asarray(columns["element"]), + altloc=np.asarray(columns["altloc"]), + residue_of=residue_of, + is_hetatm=np.asarray(columns["is_hetatm"], dtype=bool), + charge=np.asarray(columns["charge"], dtype=np.int64), + ) + return cls(residues=residues, atoms=atoms) + + def columns(self) -> Dict[str, np.ndarray]: + """Per-atom identity arrays, residue fields broadcast to atoms. + + Returns + ------- + dict + :data:`IDENTITY_COLUMNS`, each shape ``(N,)``, freshly allocated. + """ + of = self.atoms.residue_of.cpu().numpy() + return { + "name": self.atoms.name.copy(), + "element": self.atoms.element.copy(), + "altloc": self.atoms.altloc.copy(), + "chain": self.residues.chain[of], + "resseq": self.residues.resseq[of], + "icode": self.residues.icode[of], + "resname": self.residues.resname[of], + "is_hetatm": self.atoms.is_hetatm.copy(), + "charge": self.atoms.charge.copy(), + } + + def gather(self, rows: np.ndarray) -> "Topology": + """A node-only topology whose atom ``i`` is this one's atom ``rows[i]``. + + Parameters + ---------- + rows : numpy.ndarray + Source atom per new atom, shape ``(N_new,)``. Rows may repeat -- a hydrogen + gathered from its parent -- and the caller overwrites what differs. + + Returns + ------- + Topology + ``connected=False``: edges are not carried, and residues are re-derived from + the gathered order, so the gathered atoms of one residue must stay + contiguous. + """ + rows = np.asarray(rows, dtype=np.int64) + columns = {key: value[rows] for key, value in self.columns().items()} + return Topology.from_columns(columns, device=self.atoms.residue_of.device) + + def select(self, selection: str) -> torch.Tensor: + """Atoms matching a Phenix-style selection. + + Parameters + ---------- + selection : str + Grammar in :mod:`torchref.utils.selection`, e.g. + ``"chain A and not resname HOH"``. + + Returns + ------- + torch.Tensor + Boolean mask, shape ``(N,)``, on the CPU. + """ + from torchref.utils.selection import select_atoms + + return select_atoms(self.columns(), selection) + + @property + def is_water(self) -> np.ndarray: + """True for atoms of water residues, shape ``(N,)``.""" + return self.residues.is_water[self.atoms.residue_of.cpu().numpy()] + + @property + def is_polymer(self) -> np.ndarray: + """True for atoms of polymer residues, shape ``(N,)``. + + A residue is polymer when its first atom is an ATOM record, the same rule the + peptide-link search uses. + """ + first = self.residues.atom_start.astype(np.int64) + per_residue = ~self.atoms.is_hetatm[first] if len(first) else np.zeros(0, bool) + return per_residue[self.atoms.residue_of.cpu().numpy()] + + def with_hydrogens(self, plan) -> Tuple["Topology", np.ndarray, np.ndarray, np.ndarray]: + """This topology's atoms with a hydrogen plan's atoms inserted. + + Each residue's planned hydrogens go immediately after its own atoms, never at + the end: residues are contiguous runs, so appending would split every + hydrogenated residue into two nodes. A hydrogen inherits its parent's identity + and takes the plan's ``name``, ``element`` and ``altloc``. + + Parameters + ---------- + plan : HydrogenPlan + From :func:`torchref.topology.hydrogens.plan_hydrogens` over this topology. + + Returns + ------- + topology : Topology + Node-only, ``connected=False``. + source : numpy.ndarray + Row of this topology each new atom was gathered from, ``(N_new,)`` -- the + parent for a hydrogen. Gather per-atom values with it too. + old_to_new : numpy.ndarray + New row of every existing atom, ``(N,)``. + plan_to_new : numpy.ndarray + New row of every planned hydrogen, ``(plan.n_hydrogens,)``. + """ + n_old = self.n_atoms + if plan.n_hydrogens == 0: + rows = np.arange(n_old, dtype=np.int64) + return self.gather(rows), rows, rows.copy(), np.zeros(0, dtype=np.int64) + + by_residue: Dict[int, list] = {} + for i, residue in enumerate(np.asarray(plan.residue).tolist()): + by_residue.setdefault(int(residue), []).append(i) + + source, old_to_new = [], np.empty(n_old, dtype=np.int64) + plan_to_new = np.empty(plan.n_hydrogens, dtype=np.int64) + offset = 0 + for residue in range(self.n_residues): + start = int(self.residues.atom_start[residue]) + end = int(self.residues.atom_end[residue]) + source.append(np.arange(start, end, dtype=np.int64)) + old_to_new[start:end] = offset + np.arange(end - start) + offset += end - start + members = by_residue.get(residue) + if members: + source.append(np.asarray(plan.parent, dtype=np.int64)[members]) + plan_to_new[members] = offset + np.arange(len(members)) + offset += len(members) + source = np.concatenate(source) + + columns = {key: value[source] for key, value in self.columns().items()} + altloc = np.asarray(plan.altloc).astype(str) + for key, values in ( + ("name", np.asarray(plan.name).astype(str)), + ("element", np.asarray(plan.element).astype(str)), + ("altloc", np.where(np.char.strip(altloc) == "", " ", altloc)), + ): + column = columns[key].astype( + np.result_type(columns[key].dtype, values.dtype) + ) + column[plan_to_new] = values + columns[key] = column + topology = Topology.from_columns(columns, device=self.atoms.residue_of.device) + return topology, source, old_to_new, plan_to_new @property def device(self) -> torch.device: @@ -57,7 +337,11 @@ def n_residues(self) -> int: def copy(self) -> "Topology": """An independent copy sharing no storage with this one.""" - return Topology(residues=self.residues.copy(), atoms=self.atoms.copy()) + return Topology( + residues=self.residues.copy(), + atoms=self.atoms.copy(), + connected=self.connected, + ) def subset(self, keep) -> "Topology": """The topology over a subset of the atoms. @@ -125,6 +409,7 @@ def subset(self, keep) -> "Topology": atom_end.astype(np.int64), ), atoms=self.atoms.subset(remap, residue_remap), + connected=self.connected, ) def neighbors(self, i: int) -> torch.Tensor: @@ -175,4 +460,4 @@ def __repr__(self) -> str: return f"Topology({self.residues!r}, {self.atoms!r})" -__all__ = ["Topology"] +__all__ = ["Topology", "identity_columns", "IDENTITY_COLUMNS"] diff --git a/torchref/utils/selection.py b/torchref/utils/selection.py new file mode 100644 index 00000000..7cb842eb --- /dev/null +++ b/torchref/utils/selection.py @@ -0,0 +1,144 @@ +"""Phenix-style atom selections evaluated over per-atom identity arrays. + +The grammar is the contract. Terms: ``chain ``, ``resseq ``, +``resseq :`` (inclusive), ``resname ``, ``name ``, +``element ``, ``altloc `` and ``all``. They combine with ``not``, ``and`` and +``or`` -- binding in that order, tightest first -- and ``(...)`` groups, e.g. +``"chain A and (name CA or name CB)"``. ``resname``, ``name`` and ``element`` match +case-insensitively; ``chain`` and ``altloc`` do not. + +A selection is evaluated against a mapping of per-atom arrays rather than an atom table, +so the same parser serves :meth:`torchref.topology.Topology.select` and anything else +that can produce the columns. +""" + +import re +from typing import List, Mapping + +import numpy as np +import torch + +#: Columns a selection may read, each an array of shape ``(N,)``. +SELECTION_COLUMNS = ("chain", "resseq", "resname", "name", "element", "altloc") + +_TOKEN = re.compile(r"\(|\)|[^\s()]+") + + +def select_atoms(columns: Mapping[str, np.ndarray], selection: str) -> torch.Tensor: + """Evaluate a Phenix-style selection. + + Parameters + ---------- + columns : mapping of str to numpy.ndarray + Per-atom arrays keyed by :data:`SELECTION_COLUMNS`, all of shape ``(N,)``; + ``resseq`` is integer, the rest strings. + selection : str + Selection string; see the module docstring for the grammar. + + Returns + ------- + torch.Tensor + Boolean mask of shape ``(N,)``, on the CPU. + + Raises + ------ + ValueError + On an empty selection, an unknown keyword, a term without a value, unbalanced + parentheses or a trailing token. + """ + tokens = _TOKEN.findall(selection) + if not tokens: + raise ValueError("Selection string cannot be empty") + parser = _Parser(tokens, columns, selection) + mask = parser.expression() + if parser.pos != len(tokens): + raise ValueError( + f"Invalid selection syntax: unexpected {tokens[parser.pos]!r} in " + f"{selection!r}" + ) + return torch.as_tensor(mask, dtype=torch.bool) + + +class _Parser: + """Recursive descent: ``or`` over ``and`` over ``not`` over terms and groups.""" + + def __init__(self, tokens: List[str], columns: Mapping[str, np.ndarray], text: str): + self.tokens = tokens + self.columns = columns + self.text = text + self.pos = 0 + self.n = len(columns["name"]) + + def _peek(self) -> str: + return self.tokens[self.pos].lower() if self.pos < len(self.tokens) else "" + + def _take(self) -> str: + if self.pos >= len(self.tokens): + raise ValueError(f"Invalid selection syntax: {self.text!r} ends early") + token = self.tokens[self.pos] + self.pos += 1 + return token + + def expression(self) -> np.ndarray: + mask = self._conjunction() + while self._peek() == "or": + self._take() + mask = mask | self._conjunction() + return mask + + def _conjunction(self) -> np.ndarray: + mask = self._unary() + while self._peek() == "and": + self._take() + mask = mask & self._unary() + return mask + + def _unary(self) -> np.ndarray: + token = self._peek() + if token == "not": + self._take() + return ~self._unary() + if token == "(": + self._take() + mask = self.expression() + if self._take() != ")": + raise ValueError(f"Unbalanced parentheses in {self.text!r}") + return mask + if token == ")": + raise ValueError(f"Unbalanced parentheses in {self.text!r}") + return self._term() + + def _term(self) -> np.ndarray: + keyword = self._take().lower() + if keyword == "all": + return np.ones(self.n, dtype=bool) + if keyword in ("and", "or"): + raise ValueError(f"Invalid selection syntax: {self.text!r}") + if self._peek() in ("", "and", "or", ")", "("): + raise ValueError(f"Invalid selection syntax: '{keyword}' has no value") + value = self._take() + cols = self.columns + if keyword == "chain": + return cols["chain"] == value + if keyword == "resseq": + resseq = np.asarray(cols["resseq"]) + if ":" in value: + start, end = (int(v) for v in value.split(":")) + return (resseq >= start) & (resseq <= end) + return resseq == int(value) + if keyword in ("resname", "name"): + return _upper(cols[keyword]) == value.upper() + if keyword == "element": + return np.char.capitalize(np.char.strip(cols["element"].astype(str))) == ( + value.capitalize() + ) + if keyword == "altloc": + return cols["altloc"] == value + raise ValueError(f"Unknown selection keyword: '{keyword}'") + + +def _upper(values: np.ndarray) -> np.ndarray: + return np.char.upper(np.char.strip(np.asarray(values).astype(str))) + + +__all__ = ["select_atoms", "SELECTION_COLUMNS"] diff --git a/torchref/utils/utils.py b/torchref/utils/utils.py index a3f1fc62..f5a1e669 100644 --- a/torchref/utils/utils.py +++ b/torchref/utils/utils.py @@ -425,177 +425,41 @@ def sanitize_pdb_dataframe(pdb: pd.DataFrame, verbose: int = 0) -> pd.DataFrame: return pdb -def _parse_with_parentheses( - selection_string: str, pdb_df: pd.DataFrame -) -> torch.Tensor: - """ - Helper function to handle parentheses in selection strings. - Recursively evaluates innermost parentheses first. - """ - import re - - # Find innermost parentheses - while True: - match = re.search(r"\(([^()]+)\)", selection_string) - if not match: - break - - # Evaluate the innermost parenthesized expression - inner = match.group(1) - inner_mask = _parse_without_parentheses(inner, pdb_df) - - # Replace with a placeholder that we'll substitute back - # Use a unique placeholder that won't appear in normal selection - placeholder = f"__MASK_{id(inner_mask)}__" - selection_string = ( - selection_string[: match.start()] - + placeholder - + selection_string[match.end() :] - ) - - # Store the mask result in a temporary global dict - # (not ideal but works for this recursive evaluation) - if not hasattr(_parse_with_parentheses, "_mask_cache"): - _parse_with_parentheses._mask_cache = {} - _parse_with_parentheses._mask_cache[placeholder] = inner_mask - - # Now parse the expression without parentheses, substituting cached masks - return _parse_without_parentheses(selection_string, pdb_df) - - -def _parse_without_parentheses( - selection_string: str, pdb_df: pd.DataFrame -) -> torch.Tensor: - """ - Parse selection string without parentheses. - Handles logical operators and basic keywords. - """ - import re - - selection_string = selection_string.strip() - - if not selection_string: - raise ValueError("Selection string cannot be empty") - - if selection_string.startswith("__MASK_") and selection_string.endswith("__"): - if hasattr(_parse_with_parentheses, "_mask_cache"): - return _parse_with_parentheses._mask_cache.get( - selection_string, torch.ones(len(pdb_df), dtype=torch.bool) - ) - return torch.ones(len(pdb_df), dtype=torch.bool) - - if selection_string.lower() == "all": - return torch.ones(len(pdb_df), dtype=torch.bool) - - # Priority: not > and > or - - if " or " in selection_string.lower(): - parts = re.split(r"\s+or\s+", selection_string, flags=re.IGNORECASE) - masks = [_parse_without_parentheses(part.strip(), pdb_df) for part in parts] - result = masks[0] - for mask in masks[1:]: - result = result | mask - return result - - if " and " in selection_string.lower(): - parts = re.split(r"\s+and\s+", selection_string, flags=re.IGNORECASE) - masks = [_parse_without_parentheses(part.strip(), pdb_df) for part in parts] - result = masks[0] - for mask in masks[1:]: - result = result & mask - return result - - if selection_string.lower().startswith("not "): - inner_selection = selection_string[4:].strip() - return ~_parse_without_parentheses(inner_selection, pdb_df) - - parts = selection_string.split(None, 1) - if len(parts) < 2: - raise ValueError(f"Invalid selection syntax: '{selection_string}'") - - keyword, value = parts[0].lower(), parts[1] - - mask = torch.zeros(len(pdb_df), dtype=torch.bool) - - if keyword == "chain": - chain_id = value.strip() - selected = pdb_df["chainid"] == chain_id - mask = torch.tensor(selected.values, dtype=torch.bool) - - elif keyword == "resseq": - if ":" in value: - start, end = value.split(":") - start, end = int(start.strip()), int(end.strip()) - selected = (pdb_df["resseq"] >= start) & (pdb_df["resseq"] <= end) - else: - resseq_num = int(value.strip()) - selected = pdb_df["resseq"] == resseq_num - mask = torch.tensor(selected.values, dtype=torch.bool) - - elif keyword == "resname": - resname = value.strip().upper() - selected = pdb_df["resname"].str.upper() == resname - mask = torch.tensor(selected.values, dtype=torch.bool) - - elif keyword == "name": - atom_name = value.strip().upper() - selected = pdb_df["name"].str.upper() == atom_name - mask = torch.tensor(selected.values, dtype=torch.bool) - - elif keyword == "element": - element = value.strip().capitalize() - selected = pdb_df["element"].str.capitalize() == element - mask = torch.tensor(selected.values, dtype=torch.bool) - - elif keyword == "altloc": - altloc = value.strip() - selected = pdb_df["altloc"] == altloc - mask = torch.tensor(selected.values, dtype=torch.bool) - - else: - raise ValueError(f"Unknown selection keyword: '{keyword}'") - - return mask - - def parse_phenix_selection(selection_string: str, pdb_df: pd.DataFrame) -> torch.Tensor: - """ - Parse Phenix-style atom selection syntax and return a boolean mask. + """Evaluate a Phenix-style selection against an atom table. - The grammar is the contract, so it is spelled out. Terms: - ``chain ``, ``resseq ``, ``resseq :`` (inclusive), - ``resname ``, ``name ``, ``element ``, ``altloc ``, ``all``. - Combined with ``not``, ``and``, ``or`` (that precedence) and ``(...)`` for grouping, - e.g. ``"chain A and (name CA or name CB)"``. ``resname``/``name`` match - case-insensitively, ``chain``/``altloc`` do not. + The grammar is documented in :mod:`torchref.utils.selection`; a model's own atoms + are selected with ``model.ctx.topology.select``. Parameters ---------- selection_string : str Phenix-style selection string. pdb_df : pandas.DataFrame - Atomic data with columns 'chainid', 'resseq', 'resname', 'name', 'element', - 'altloc'. + Atom table with ``chainid``, ``resseq``, ``resname``, ``name``, ``element`` and + ``altloc`` columns. Returns ------- torch.Tensor - Boolean tensor of shape (n_atoms,), on the CPU regardless of where ``pdb_df``'s - consumers live. + Boolean tensor of shape (n_atoms,), on the CPU. Raises ------ ValueError On an unknown keyword, an empty selection, or a bare term with no value. """ - # Clear any cached masks from previous calls - if hasattr(_parse_with_parentheses, "_mask_cache"): - _parse_with_parentheses._mask_cache.clear() - - if "(" in selection_string: - return _parse_with_parentheses(selection_string, pdb_df) - else: - return _parse_without_parentheses(selection_string, pdb_df) + from torchref.utils.selection import select_atoms + + columns = { + "chain": pdb_df["chainid"].values.astype(str), + "resseq": pdb_df["resseq"].values, + "resname": pdb_df["resname"].values.astype(str), + "name": pdb_df["name"].values.astype(str), + "element": pdb_df["element"].values.astype(str), + "altloc": pdb_df["altloc"].values.astype(str), + } + return select_atoms(columns, selection_string) def create_selection_mask( From 158913f2b8154752ac7e440314e90a0d7195d759 Mon Sep 17 00:00:00 2001 From: Claude Date: Mon, 28 Sep 2026 20:09:53 +0000 Subject: [PATCH 199/250] Take canonicalize_hkl's working tensors from the configured dtypes dev's threaded canonicalize_hkl hardcodes int32 for the Miller indices and rotation operators, int16 for the per-row operator index, and float32 for the phase-shift arithmetic. None of them is marked, so test_no_unjustified_hardcoded_dtype fails on dev. They now take get_int_dtype() and get_float_dtype(). The operator index goes to index_select directly, which accepts the configured int dtype, so its .long() copy is dropped. Under the default int32/float32 config the outputs are bit-identical to dev's on 229,633 reflections in 36 space groups, sorted and unsorted, and so they are under TORCHREF_DTYPE_INT=int64. Under TORCHREF_DTYPE_FLOAT=float64 the phase arithmetic no longer runs through float32: the shifts land within 1e-13 of the exact multiples of 2*pi/12, where dev's were off by up to 9e-5 rad. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- docs/changelog.rst | 1 + torchref/symmetry/reciprocal_symmetry.py | 18 +++++++++--------- 2 files changed, 10 insertions(+), 9 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 50349467..dcb54aa4 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -7,6 +7,7 @@ Unreleased - Integer and index tensors now take the configured int dtype (``get_int_dtype()``, ``TORCHREF_DTYPE_INT``, int32 by default) throughout the package. A hardcoded ``int64`` remains only where a torch op (``scatter``/``gather`` on torch < 2.8, ``index_copy_``), a compiled kernel, an int32 overflow, or an external library requires it, and each such site says which. - The difference MTZ groups its columns into named datasets -- ``observed``, ``difference``, ``light_model``, ``extrapolated_light``, ``two_moment`` -- with one history line describing each, so ``FWT``/``PHWT`` reads as ``/torchref/extrapolated_light/FWT`` (the extrapolated light-state map ``2*FEXT - Fc``). Labels are unchanged and Coot still auto-opens it - ``SpaceGroup.canonicalize_hkl`` gains ``sort=False``, which keeps the input row order and returns ``None`` for ``sort_indices``; the ASU mapping runs as threaded torch operations and reuses the rotated indices for the Friedel mate, so large reflection lists map faster with unchanged outputs +- ``SpaceGroup.canonicalize_hkl`` computes its phase shifts in the configured float dtype, so under ``TORCHREF_DTYPE_FLOAT=float64`` they are exact to float64 instead of rounded to float32 - ``torchref.difference-map``, ``torchref.difference-refine`` and ``torchref.validate-ded`` gain ``--ded-weight {sigma_d,inverse_variance,none}`` and ``--sigma-d-gamma``. The difference MTZ now carries the unweighted ``DF``/``SIGDF`` on ``PHDELWT`` with one mean-one weight column per scheme, ``W_SD`` and ``W_IVW`` (MTZ type W), and the observed-to-model scale ``KSCALE``; ``DELFWT`` is no longer written, build the map with ``torchref.mtz2map -csf DF -cw W_IVW -cphi PHDELWT``. Registered in ``torchref.maps.ded_weights`` - Added the ``sigma_D`` estimator (``torchref.refinement.model_error_estimation.sigma_d``): the expected true difference power per resolution shell, ``mean(dF_obs^2) - mean(sigma^2)`` with a fitted ``F_dark^gamma`` amplitude law and DerSimonian-Laird shrinkage of the signed shell power toward a decaying exponential in ``d*^2`` (fitted on all shells, so a dataset without a difference yields no power instead of the positive half of its noise), giving the Wiener weight ``S/(S + sigma^2)`` and, with a difference model, ``alpha``/``beta_model``. Inverse-variance weights suppress the strong reflections whose difference power is 10-70x that of weak ones; on independent half-datasets a Wiener weight with the true power raised map agreement 1.2-1.8x in effective patterns. The single-dataset estimate inherits the calibration of the reported sigmas, and on the campaign TD1 data (sigmas ~1.5x too large at high resolution) it emptied 60-90 % of the shells, so inverse variance stays the default; ``sigma_d`` reports its clamped-shell count and falls back to inverse variance with a warning when every shell is empty - Added the ``difference_sd`` collection target (``CollectionDifferenceSigmaDTarget``): the difference Gaussian centred on ``alpha * dF_calc`` with variance ``beta_model + sigma_diff^2`` from ``sigma_D`` fitted on the free set. Selected with ``torchref.difference-refine --difference-target difference_sd``; ``difference`` stays the default diff --git a/torchref/symmetry/reciprocal_symmetry.py b/torchref/symmetry/reciprocal_symmetry.py index 2285186e..e8b0c170 100644 --- a/torchref/symmetry/reciprocal_symmetry.py +++ b/torchref/symmetry/reciprocal_symmetry.py @@ -461,10 +461,10 @@ def _canonicalize_hkl( asu = gemmi.ReciprocalAsu(sym._gemmi) condition_key = asu.condition_str() # Reciprocal-space rotation matrices are always integer-valued (0, ±1). - recip_ops = torch.round(sym.reciprocal.matrices.detach().cpu()).to(torch.int32) + recip_ops = torch.round(sym.reciprocal.matrices.detach().cpu()).to(get_int_dtype()) translations = sym.translations.detach().cpu() # (n_ops, 3) n_ops = len(recip_ops) - hkl_cpu = hkl.detach().to(device="cpu", dtype=torch.int32) # (N, 3) + hkl_cpu = hkl.detach().to(device="cpu", dtype=get_int_dtype()) # (N, 3) def in_asu(h, k, l): try: @@ -492,7 +492,7 @@ def rotate(row, h, k, l): # most reflections are resolved by the first few operators. ``todo`` holds the # still-unmapped rows in increasing order (``None`` while that is all of them). canonical = torch.empty_like(hkl_cpu) - op_idx = torch.empty(n_refl, dtype=torch.int16) + op_idx = torch.empty(n_refl, dtype=get_int_dtype()) friedel = torch.zeros(n_refl, dtype=torch.bool) todo = None @@ -545,17 +545,17 @@ def rotate(row, h, k, l): # A single uniform sign is wrong for one half and invisible in P21/P212121/C2, # where every shift is 0 or π. tests/unit/symmetry/test_phase_convention.py. # h·t is summed left to right so the value does not depend on a backend's - # reduction order; the shift is rounded to float32 like the rest of the output. + # reduction order. if bool(translations.any()): - t_sel = translations.index_select(0, op_idx.long()) - hf = hkl_cpu.to(torch.float32) + t_sel = translations.index_select(0, op_idx) + hf = hkl_cpu.to(get_float_dtype()) h_dot_t = ( hf[:, 0] * t_sel[:, 0] + hf[:, 1] * t_sel[:, 1] + hf[:, 2] * t_sel[:, 2] ) - friedel_sign = torch.where(friedel, 1.0, -1.0).to(torch.float32) - phase = (friedel_sign * 2.0 * math.pi * h_dot_t).to(torch.float32) + friedel_sign = torch.where(friedel, 1.0, -1.0).to(get_float_dtype()) + phase = friedel_sign * 2.0 * math.pi * h_dot_t else: - phase = torch.zeros(n_refl, dtype=torch.float32) + phase = torch.zeros(n_refl, dtype=get_float_dtype()) canonical_hkl = canonical.to(dtype=hkl_dtype, device=device) phase_shifts = phase.to(dtype=get_float_dtype(), device=device) From 4ee15bfef11880b07659a125a77f1ab898a5d8e9 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Tue, 29 Sep 2026 10:00:58 +0200 Subject: [PATCH 200/250] Build restraints from the topology instead of the atom table The builders read a node-only Topology plus coordinates. The peptide-link builders take a prepared PeptideResidues (conformer maps over residue ranges, the residue graph's peptide pairs, coordinates for the omega classification) instead of a DataFrame grouped on (chain, resseq), so an insertion-code step is linked like any other residue step. Disulfide paths and LINK records resolve atoms through residue ranges; the pair filter reads residue membership and altlocs from the topology. On the twelve deposited structures and all bundled mmCIF inputs the edges, links, template keys, pair lists, Ramachandran pairing and proline flags are unchanged. Identity strings are stripped at the boundary, and missing icode/ATOM columns fall back to defaults instead of failing. The dead per-residue builders, PreprocessedPDB and the unused H pair search go. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 1 + tests/fixtures/objects.py | 3 +- .../functional/test_restraints_functional.py | 33 +- tests/integration/test_refinement_pipeline.py | 5 +- tests/unit/monomer/test_restraints.py | 4 +- tests/unit/topology/test_equivalence.py | 10 +- tests/unit/topology/test_insertion_codes.py | 138 +- tests/unit/topology/test_storage.py | 9 +- tests/unit/topology/test_topology_identity.py | 36 + torchref/model/context.py | 9 +- torchref/topology/build.py | 225 ++- torchref/topology/builders.py | 1214 ++--------------- torchref/topology/nonbonded.py | 259 +--- torchref/topology/restraints.py | 91 +- torchref/topology/topology.py | 7 +- 15 files changed, 425 insertions(+), 1619 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 214c993b..6ea02354 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Restraints are built from a topology rather than an atom table: ``Restraints(topology=..., xyz=...)`` and ``build_topology(_with_values)(topology, cif_dict, xyz, ...)`` take a node-only ``Topology`` (``Topology.from_table``), and the peptide-link builders pair residues along the residue graph's links, so residues with insertion codes (100, 100A, 101) are now peptide-linked and restrained like any other; structures without insertion codes get identical edges and pair lists. A table without ``icode`` or ``ATOM`` columns no longer fails, and padded names read like clean ones. The unused per-residue restraint builders, ``build_all_restraints``, ``ResidueIterator``, ``PreprocessedPDB`` and ``find_h_vdw_pairs_gpu`` are removed - Atom identity is available without building restraints: ``Topology.from_table`` gives a node-only topology (names, elements, altlocs, residues, chains, record types and charges; ``connected=False``), with ``Topology.select`` for Phenix-style selections, ``is_water`` / ``is_polymer`` masks, cached ``AtomGraph.atomic_number`` / ``vdw_radii`` and ``Topology.with_hydrogens`` for inserting a hydrogen plan. Selections are evaluated by one recursive-descent parser (``torchref.utils.selection``); a selection with more than one parenthesised group, which the old parser could resolve to every atom, now selects what it says - Hydrogens are one policy on the model context, set at construction: ``hydrogens`` (``keep`` / ``add`` / ``strip``) settles the atom table when it loads and ``hydrogen_mode`` (``atoms`` / ``riding``) how hydrogen rows are parametrised; ``riding`` with ``strip`` raises. It replaces ``strip_H``, ``add_hydrogens`` and ``hydrogen_mode="free"`` on ``Model``, ``ModelFT``, ``Refinement``, ``EnsembleModel`` and ``cli._common.load_model`` (checkpoints written with them still restore), and ``torchref.refine`` takes ``--hydrogens`` / ``--hydrogen-mode`` instead of ``--add-hydrogens``. The deprecated ``exclude_H_from_sf`` alias is removed - Missing water hydrogens are generated at load with ``hydrogens="add"``, never by ``set_hydrogen_mode("riding")``, which now only swaps the coordinate wrapper and leaves the atom table alone; ``hydrogenate()`` no longer takes ``optimize`` (free torsions are always scanned) diff --git a/tests/fixtures/objects.py b/tests/fixtures/objects.py index 9bdedc87..47f67141 100644 --- a/tests/fixtures/objects.py +++ b/tests/fixtures/objects.py @@ -111,9 +111,10 @@ def initialized_scaler(model_and_data: dict[str, Any]) -> Scaler: def model_with_restraints(loaded_model: Model) -> dict[str, Any]: """Build restraints around a fresh model.""" from torchref.topology.restraints import Restraints + from torchref.topology.topology import Topology restraints = Restraints( - pdb=loaded_model.pdb, + topology=Topology.from_table(loaded_model.pdb), xyz=loaded_model.xyz(), verbose=0, ) diff --git a/tests/functional/test_restraints_functional.py b/tests/functional/test_restraints_functional.py index 146e0fcc..44ae43f4 100644 --- a/tests/functional/test_restraints_functional.py +++ b/tests/functional/test_restraints_functional.py @@ -16,12 +16,13 @@ def test_build_restraints_from_cif(self, sample_cif_file): """Test building restraints from a real CIF file.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints + from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz=model.xyz(), verbose=0 + topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 ) # Should have built some restraints @@ -33,12 +34,13 @@ def test_bond_restraints_built(self, sample_cif_file): """Test that bond restraints are built correctly.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints + from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz=model.xyz(), verbose=0 + topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 ) # Check bond restraints exist @@ -70,12 +72,13 @@ def test_angle_restraints_built(self, sample_cif_file): """Test that angle restraints are built correctly.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints + from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz=model.xyz(), verbose=0 + topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 ) # Check angle restraints exist @@ -100,12 +103,13 @@ def test_torsion_restraints_built(self, sample_cif_file): """Test that torsion restraints are built correctly.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints + from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz=model.xyz(), verbose=0 + topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 ) # Check torsion restraints exist @@ -128,12 +132,13 @@ def test_plane_restraints_built(self, sample_cif_file): """Test that plane restraints are built correctly.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints + from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz=model.xyz(), verbose=0 + topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 ) # Check plane restraints exist @@ -159,12 +164,13 @@ def test_bond_deviations(self, sample_cif_file): """Test computing bond length deviations.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints + from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz=model.xyz(), verbose=0 + topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 ) # Compute bond deviations @@ -182,12 +188,13 @@ def test_angle_deviations(self, sample_cif_file): """Test computing angle deviations.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints + from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz=model.xyz(), verbose=0 + topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 ) # Compute angle deviations @@ -206,10 +213,11 @@ class TestRestraintsMultipleStructures: def test_restraints_multiple_cif_files(self, compatibility_model): """Each extended crystal supplies bond and angle restraints.""" from torchref.topology.restraints import Restraints + from torchref.topology.topology import Topology model = compatibility_model restraints = Restraints( - pdb=model.pdb, + topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0, ) @@ -225,12 +233,13 @@ def test_restraints_device_movement(self, sample_cif_file, cpu_device): """Test moving restraints to different devices.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints + from torchref.topology.topology import Topology model = Model(device=cpu_device) model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz=model.xyz(), verbose=0 + topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 ) # Check that tensors are on the correct device @@ -247,12 +256,13 @@ def test_cif_dict_loaded(self, sample_cif_file): """Test that CIF dictionary is loaded correctly.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints + from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz=model.xyz(), verbose=0 + topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 ) # CIF dict should be populated with residue restraints @@ -273,12 +283,13 @@ def test_unique_residues_detected(self, sample_cif_file): """Test that unique residues are detected from model.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints + from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, xyz=model.xyz(), verbose=0 + topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 ) # Should have detected unique residues diff --git a/tests/integration/test_refinement_pipeline.py b/tests/integration/test_refinement_pipeline.py index bf193d54..2947f4a7 100644 --- a/tests/integration/test_refinement_pipeline.py +++ b/tests/integration/test_refinement_pipeline.py @@ -45,12 +45,15 @@ def test_restraints_from_model(self, sample_cif_file): """Test building restraints from a loaded model.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints + from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) # Build restraints - restraints = Restraints(pdb=model.pdb, xyz=model.xyz()) + restraints = Restraints( + topology=Topology.from_table(model.pdb), xyz=model.xyz() + ) # Should have some restraints assert restraints.restraints is not None diff --git a/tests/unit/monomer/test_restraints.py b/tests/unit/monomer/test_restraints.py index 97ea6c1a..2e558dbb 100644 --- a/tests/unit/monomer/test_restraints.py +++ b/tests/unit/monomer/test_restraints.py @@ -20,8 +20,8 @@ def test_restraints_empty_init(self): from torchref.topology.restraints import Restraints restraints = Restraints() - - assert restraints.pdb is None + + assert restraints.topology is None @pytest.mark.unit def test_restraints_is_nn_module(self): diff --git a/tests/unit/topology/test_equivalence.py b/tests/unit/topology/test_equivalence.py index 0f67843c..c832514b 100644 --- a/tests/unit/topology/test_equivalence.py +++ b/tests/unit/topology/test_equivalence.py @@ -7,6 +7,8 @@ """ import pytest + +from torchref.topology.topology import Topology import torch from torchref.model.model import Model @@ -71,12 +73,12 @@ def _build(code): model.ctx.set_cif_path(None) restraints = model.restraints topology = build_topology( - model.pdb, + Topology.from_table(model.pdb), restraints.cif_dict, + model.xyz().detach(), link_dict=getattr(restraints, "link_dict", None), link_list=getattr(restraints, "link_list", None), links=restraints.links, - xyz=model.xyz().detach(), verbose=0, ) cache[code] = (topology, _current_edges(restraints)) @@ -203,12 +205,12 @@ def test_layout_is_reproducible(built, pdb_dir): model.ctx.set_cif_path(None) restraints = model.restraints topology_b = build_topology( - model.pdb, + Topology.from_table(model.pdb), restraints.cif_dict, + model.xyz().detach(), link_dict=getattr(restraints, "link_dict", None), link_list=getattr(restraints, "link_list", None), links=restraints.links, - xyz=model.xyz().detach(), verbose=0, ) diff --git a/tests/unit/topology/test_insertion_codes.py b/tests/unit/topology/test_insertion_codes.py index 3e52a8ea..738aa160 100644 --- a/tests/unit/topology/test_insertion_codes.py +++ b/tests/unit/topology/test_insertion_codes.py @@ -2,10 +2,9 @@ A deposited structure may number two residues 100 and 100A. They are different residues with different chemistry, and the only thing separating them is the insertion code. The -topology keys residues on ``(chain, resseq, icode)`` for that reason; the restraint -builders key on ``(chain, resseq)`` alone, which merges them into one residue whose -atom names then collide, so the name-to-index map keeps the first of each and every -restraint belonging to the later residues is silently lost. +topology keys residues on ``(chain, resseq, icode)`` for that reason, and every builder +-- intra-residue matching and the peptide links alike -- works on those residues, so +each inserted residue gets its own geometry and the chain runs through the insertion. No bundled structure has an insertion code, so the case is synthesised here rather than shipped as another data file: the rewrite is then visible, and it is obvious that @@ -16,6 +15,7 @@ from torchref.model.model import Model from torchref.topology import build_topology +from torchref.topology.topology import Topology #: Base structure: chain A, no altlocs, no insertion codes anywhere. BASE = "3GR5" @@ -66,12 +66,12 @@ def inserted(pdb_dir, tmp_path_factory): restraints = model.restraints topology = build_topology( - model.pdb, + Topology.from_table(model.pdb), restraints.cif_dict, + model.xyz().detach(), link_dict=restraints.link_dict, link_list=restraints.link_list, links=restraints.links, - xyz=model.xyz().detach(), verbose=0, ) return topology, restraints, expected, model @@ -105,49 +105,6 @@ def test_the_graph_keeps_them_apart(inserted): assert len({key[2] for key in found}) == 3, "insertion codes were not distinguished" -@pytest.mark.unit -def test_the_builders_merge_them(inserted): - """The comparison only means something if the old grouping really does merge. - - ``PreprocessedPDB`` groups on ``(chain, resseq)``, so the three residues become one - with three sets of backbone atom names. - """ - from torchref.topology.builders import PreprocessedPDB - - _, _, expected, model = inserted - preprocessed = PreprocessedPDB(model.pdb) - - merged = [ - i - for i in range(preprocessed.n_residues) - if int(preprocessed.residue_resseqs[i]) == expected[0][0] - ] - assert len(merged) == 1, "the builders did not merge the inserted residues" - assert preprocessed.has_duplicate_atoms(merged[0]), ( - "the merged residue should carry duplicate atom names, which is what makes the " - "name-to-index map lose the later residues" - ) - - -def _legacy_intra_bonds(model, restraints): - """Intra-residue bonds as the ``(chain, resseq)``-keyed builder produces them. - - ``restraints.py`` now builds from the topology, so it cannot serve as the - baseline -- it *is* the graph. ``BondRestraintBuilder`` is the original path, still - keying residues on ``(chain, resseq)``, which is the behaviour under test. - """ - import torch - - from torchref.topology.builders import BondRestraintBuilder - - built = BondRestraintBuilder(verbose=0).build( - model.pdb, restraints.cif_dict, torch.device("cpu") - ) - if not built: - return set() - return {tuple(int(v) for v in row) for row in built["indices"].cpu().numpy()} - - def _inserted_residue_indices(topology, expected): wanted = {(seq, code) for seq, code in expected} return [ @@ -164,59 +121,6 @@ def _bonds_within(edges, topology, residue): return {e for e in edges if all(start <= int(a) < end for a in e)} -@pytest.mark.unit -def test_the_legacy_grouping_loses_the_later_residues(inserted): - """The merged residue gets restraints for its first component only. - - This is the defect keying on ``(chain, resseq, icode)`` fixes. The three residues - become one, their backbone atom names collide, the name-to-index map keeps the first - of each, and the second and third end up with no intra-residue geometry at all. - """ - topology, restraints, expected, model = inserted - legacy = _legacy_intra_bonds(model, restraints) - residues = _inserted_residue_indices(topology, expected) - assert len(residues) == 3 - - counts = [len(_bonds_within(legacy, topology, r)) for r in residues] - assert counts[0] > 0, "even the first component lost its bonds; check the fixture" - assert counts[1:] == [0, 0], ( - f"the legacy grouping was expected to lose the second and third residues, " - f"but found {counts} bonds in them" - ) - - -@pytest.mark.unit -def test_the_graph_finds_what_the_legacy_grouping_lost(inserted): - """Every inserted residue gets its own bonds, and the graph is a strict superset. - - Localised, not merely larger: the bonds the graph adds all lie inside the inserted - residues, so this is the insertion-code fix rather than a general difference. - """ - topology, restraints, expected, model = inserted - legacy = _legacy_intra_bonds(model, restraints) - graph = topology.atoms.bonds.tuple_set("intra") - residues = _inserted_residue_indices(topology, expected) - - for residue in residues: - assert _bonds_within( - graph, topology, residue - ), f"residue {topology.residues.key(residue)} has no intra-residue bonds" - - gained = graph - legacy - assert gained, "the graph found nothing the legacy grouping missed" - - inserted_set = set(residues) - stray = [ - edge - for edge in gained - if not ({topology.residue_of_atom(a) for a in edge} & inserted_set) - ] - assert not stray, ( - f"{len(stray)} gained bonds lie outside the inserted residues, so the " - f"difference is not localised to the insertion codes: {stray[:3]}" - ) - - @pytest.mark.unit def test_the_inserted_residues_get_their_own_intra_restraints(inserted): """Each of the three carries bonds of its own, not just the first. @@ -281,3 +185,33 @@ def test_the_inserted_stretch_is_peptide_linked(inserted): f"expected two peptide links inside the three inserted residues, got " f"{len(internal)}" ) + + +def _atom(topology, residue, name): + rows = topology.residues.atom_rows(residue) + return next(r for r in rows if str(topology.atoms.name[r]).strip() == name) + + +@pytest.mark.unit +def test_peptide_edges_run_through_the_insertion(inserted): + """C(i)-N(i+1) bonds, and the phi/psi torsions, exist at every insertion-code step. + + The peptide-link builders pair residues along the residue graph's links, so 23 to + 23A and 23A to 23B are linked like any other step, and the middle residue gets + both its phi and its psi. + """ + topology, _, expected, _ = inserted + residues = _inserted_residue_indices(topology, expected) + peptide = topology.atoms.bonds.tuple_set("peptide") + phi = topology.atoms.torsions.tuple_set("phi") + psi = topology.atoms.torsions.tuple_set("psi") + + for first, second in zip(residues, residues[1:]): + c, n = _atom(topology, first, "C"), _atom(topology, second, "N") + assert (min(c, n), max(c, n)) in {tuple(sorted(e)) for e in peptide}, ( + f"no peptide bond from {topology.residues.key(first)} to " + f"{topology.residues.key(second)}" + ) + middle = residues[1] + assert any(row[1] == _atom(topology, middle, "N") for row in phi) + assert any(row[0] == _atom(topology, middle, "N") for row in psi) diff --git a/tests/unit/topology/test_storage.py b/tests/unit/topology/test_storage.py index 56f89745..34bde497 100644 --- a/tests/unit/topology/test_storage.py +++ b/tests/unit/topology/test_storage.py @@ -10,7 +10,6 @@ import pytest import torch -from torchref.config import get_float_dtype from torchref.model.model import Model from torchref.utils.caching import ParameterFingerprint @@ -172,13 +171,7 @@ def test_blocks_are_untouched_by_a_refinement_step(restraints): blocks = [restraints.topology.edge_block(t).indices for t in KEYED_TYPES] fingerprint = ParameterFingerprint(blocks) - block = restraints.topology.atoms.bonds.indices - xyz = torch.tensor( - restraints.pdb[["x", "y", "z"]].values, - dtype=get_float_dtype(), - device=block.device, - requires_grad=True, - ) + xyz = restraints._last_vdw_build_xyz.clone().requires_grad_(True) loss = restraints.nll_bonds(xyz).sum() + restraints.nll_angles(xyz).sum() loss.backward() diff --git a/tests/unit/topology/test_topology_identity.py b/tests/unit/topology/test_topology_identity.py index 0ba0ea21..d49aca64 100644 --- a/tests/unit/topology/test_topology_identity.py +++ b/tests/unit/topology/test_topology_identity.py @@ -136,3 +136,39 @@ def test_hydrogen_insertion_matches_the_table_insertion(pdb_dir, code): for key, value in expected.columns().items(): np.testing.assert_array_equal(nodes.columns()[key], value, err_msg=key) np.testing.assert_array_equal(nodes.residues.atom_start, expected.residues.atom_start) + + +@pytest.mark.unit +def test_padded_strings_read_like_clean_ones(pdb_dir): + """Whitespace around names, residue names, elements and icodes is not identity.""" + df = _table(pdb_dir, "1DAW") + padded = df.copy() + for column in ("name", "resname", "element", "icode"): + padded[column] = " " + padded[column].astype(str) + " " + clean, noisy = Topology.from_table(df), Topology.from_table(padded) + for key, value in clean.columns().items(): + np.testing.assert_array_equal(noisy.columns()[key], value, err_msg=key) + + +@pytest.mark.unit +@pytest.mark.parametrize("dropped", [["icode"], ["altloc"], ["ATOM"], ["charge"], ["element"]]) +def test_optional_columns_fall_back_to_defaults(pdb_dir, dropped): + df = _table(pdb_dir, "1DAW") + full = Topology.from_table(df) + reduced = Topology.from_table(df.drop(columns=dropped)) + assert reduced.n_atoms == full.n_atoms + defaults = {"icode": "", "altloc": " ", "ATOM": False, "charge": 0, "element": ""} + column = {"ATOM": "is_hetatm"}.get(dropped[0], dropped[0]) + assert (reduced.columns()[column] == defaults[dropped[0]]).all() + if dropped[0] in ("charge", "element"): + np.testing.assert_array_equal(reduced.residues.atom_start, full.residues.atom_start) + + +@pytest.mark.unit +def test_charges_are_coerced(pdb_dir): + df = _table(pdb_dir, "1DAW") + df["charge"] = df["charge"].astype(object) + df.loc[::7, "charge"] = "junk" + df.loc[1, "charge"] = 2 + charge = Topology.from_table(df).atoms.charge + assert charge.dtype == np.int64 and charge[0] == 0 and charge[1] == 2 diff --git a/torchref/model/context.py b/torchref/model/context.py index 2f402335..0d152190 100644 --- a/torchref/model/context.py +++ b/torchref/model/context.py @@ -394,8 +394,10 @@ def build_restraints( """ from torchref.topology.restraints import Restraints + from torchref.topology import Topology + restraints = Restraints( - pdb=self.pdb, + topology=Topology.from_table(self.pdb), cif_path=self.cif_path, xyz=xyz.detach(), cell=self.cell, @@ -643,9 +645,8 @@ def copy(self) -> "ModelContext": ) if self.restraints is not None: restraints = self.restraints.copy() - # Point at the copied table and crystal rather than the deep-copied - # duplicates, so the new context is the single owner of both. - restraints.pdb = duplicate.pdb + # Point at the copied crystal rather than the deep-copied duplicates, so the + # new context is its single owner. restraints._cell = duplicate.cell restraints._spacegroup = duplicate.spacegroup duplicate.restraints = restraints diff --git a/torchref/topology/build.py b/torchref/topology/build.py index 936d0adb..da2d6f29 100644 --- a/torchref/topology/build.py +++ b/torchref/topology/build.py @@ -1,15 +1,16 @@ -"""Assemble a :class:`~torchref.topology.topology.Topology` from an atom table. - -Intra-residue edges are matched here, template by template, through the matchers -in :mod:`torchref.topology.matchers`. Inter-residue edges come from the -``InterResidue*Builder`` classes, which already encode the link geometry and are reused -rather than reimplemented. +"""Connect a node-only :class:`~torchref.topology.topology.Topology` against the dictionaries. + +The input carries identity only (:meth:`Topology.from_table`); this module adds the +edges. Intra-residue edges are matched template by template through the matchers in +:mod:`torchref.topology.matchers`. Inter-residue edges come from the +``InterResidue*Builder`` classes over the residue graph's peptide links, disulfides are +found by SG-SG distance, and ``LINK`` records are resolved by residue identity. Every +edge index is an atom row of the topology. """ from typing import Dict, List, Optional, Sequence, Tuple import numpy as np -import pandas as pd import torch from torchref.topology.builders import ( @@ -17,6 +18,7 @@ InterResidueBondBuilder, InterResiduePlaneBuilder, InterResidueTorsionBuilder, + PeptideResidues, PreprocessedCIF, ) from torchref.topology.matchers import ( @@ -29,7 +31,6 @@ from torchref.topology.edges import EdgeBlock, assemble_origins from torchref.topology.residue_graph import ( ResidueGraph, - build_residue_nodes, find_disulfide_links, find_peptide_links, ) @@ -41,38 +42,15 @@ _WORK = 64 -def _atom_columns(pdb: pd.DataFrame) -> Dict[str, np.ndarray]: - """Per-atom identity arrays, with altlocs normalised so blank reads as ``' '``.""" - altloc = pdb["altloc"].values.astype(str) if "altloc" in pdb.columns else None - if altloc is None: - altloc = np.full(len(pdb), " ", dtype=" Dict[str, np.ndarray]: + """Per-atom identity arrays the matchers read, ``record`` and ``index`` included. + ``index`` is the atom row: edge indices are rows of the topology. + """ + cols = topology.columns() + cols["record"] = np.where(cols.pop("is_hetatm"), "HETATM", "ATOM") + cols["index"] = np.arange(topology.n_atoms, dtype=np.int64) + return cols def _conformers( cols: Dict[str, np.ndarray], start: int, end: int @@ -390,13 +368,15 @@ def _match_intra_planes( def _inter_residue_edges( - pdb: pd.DataFrame, + residues: PeptideResidues, link_dict: Optional[Dict], verbose: int, ) -> Tuple[Dict[str, Dict[str, np.ndarray]], Dict[str, Dict], Dict[str, Dict]]: """Peptide edges, their values, and the Ramachandran pairing, from the builders. - Reuses ``InterResidue*Builder`` rather than reimplementing the link geometry. + Reuses ``InterResidue*Builder`` rather than reimplementing the link geometry. The + pairs are the residue graph's peptide links, so an insertion-code step (100 to + 100A) is linked like any other. Returns ------- @@ -436,7 +416,7 @@ def split(group): return rows, rest bond = InterResidueBondBuilder(verbose=verbose).build( - pdb, trans, cpu, filter_atom_type="ATOM" + residues, trans, cpu ) if bond: indices["bond"]["peptide"], values["bond"]["peptide"] = split(bond) @@ -447,14 +427,14 @@ def split(group): # and excluded from the TRANS pass to avoid two restraints on the same atoms. groups = [ ab.build( - pdb, trans, cpu, filter_atom_type="ATOM", exclude_next_resname="PRO" + residues, trans, cpu, exclude_next_resname="PRO" ), ab.build( - pdb, ptrans, cpu, filter_atom_type="ATOM", next_resname_filter="PRO" + residues, ptrans, cpu, next_resname_filter="PRO" ), ] else: - groups = [ab.build(pdb, trans, cpu, filter_atom_type="ATOM")] + groups = [ab.build(residues, trans, cpu)] parts = [split(g) for g in groups if g] if parts: indices["angle"]["peptide"] = np.concatenate([p[0] for p in parts], axis=0) @@ -464,7 +444,7 @@ def split(group): } tors = InterResidueTorsionBuilder(verbose=verbose).build( - pdb, trans, cpu, filter_atom_type="ATOM" + residues, trans, cpu ) if tors: for origin in ("phi", "psi", "omega"): @@ -476,7 +456,7 @@ def split(group): extras["ramachandran"] = tors["ramachandran"] planes = InterResiduePlaneBuilder(verbose=verbose).build( - pdb, trans, cpu, filter_atom_type="ATOM" + residues, trans, cpu ) if planes: for key, group in planes.items(): @@ -502,8 +482,7 @@ def _origins( def _disulfide_edges( - pdb: pd.DataFrame, - nodes: Dict[str, np.ndarray], + topology: Topology, cols: Dict[str, np.ndarray], residue_of_row: Dict[int, int], pairs: Sequence[Tuple[int, int]], @@ -543,22 +522,15 @@ def _disulfide_edges( torsion_builder = InterResidueTorsionBuilder(verbose=verbose) for row_a, row_b in pairs: - # The edge indices are the atom table's ``index`` column, not its row number. - bond_builder.process_disulfide_bond( - int(cols["index"][row_a]), int(cols["index"][row_b]), length, sigma - ) + bond_builder.process_disulfide_bond(int(row_a), int(row_b), length, sigma) res_a, res_b = residue_of_row[row_a], residue_of_row[row_b] - atoms_a = pdb.iloc[ - int(nodes["atom_start"][res_a]) : int(nodes["atom_end"][res_a]) - ] - atoms_b = pdb.iloc[ - int(nodes["atom_start"][res_b]) : int(nodes["atom_end"][res_b]) - ] if disulf.get("angles") is not None: - angle_builder.process_disulfide_angles(atoms_a, atoms_b, disulf["angles"]) + angle_builder.process_disulfide_angles( + topology, res_a, res_b, disulf["angles"] + ) if disulf.get("torsions") is not None: torsion_builder.process_disulfide_torsions( - atoms_a, atoms_b, disulf["torsions"] + topology, res_a, res_b, disulf["torsions"] ) values: Dict[str, Dict[str, np.ndarray]] = {} @@ -579,7 +551,8 @@ def _disulfide_edges( def _lookup_link_atom( - pdb: pd.DataFrame, + topology: Topology, + residue_by_key: Dict[Tuple[str, int, str], List[int]], chainid: str, resseq: int, icode: str, @@ -587,40 +560,39 @@ def _lookup_link_atom( name: str, altloc: str, ): - """Resolve one ``LINK`` record's atom to a row of the atom table, or None. + """Resolve one ``LINK`` record's atom to a row of the topology, or None. Matches on ``(chainid, resseq, icode, name)`` with ``resname`` as a tie-breaker. Where a residue has alternative conformations the requested altloc wins, then the blank one, then ``'A'``, then whatever is left -- a LINK naming a specific conformer should reach that conformer, but one naming none should still resolve. """ - sel = pdb[ - (pdb["chainid"].astype(str) == str(chainid)) - & (pdb["resseq"].astype(int) == int(resseq)) - & (pdb["icode"].astype(str) == str(icode)) - & (pdb["name"].astype(str).str.strip() == str(name).strip()) + key = (str(chainid), int(resseq), str(icode).strip()) + candidates = residue_by_key.get(key, []) + wanted = str(resname).strip() if resname else "" + if wanted: + tied = [ + r for r in candidates if str(topology.residues.resname[r]).strip() == wanted + ] + candidates = tied or candidates + rows = [ + row + for r in candidates + for row in topology.residues.atom_rows(r) + if str(topology.atoms.name[row]).strip() == str(name).strip() ] - if len(sel) == 0: + if not rows: return None - if resname: - tied = sel[sel["resname"].astype(str).str.strip() == str(resname).strip()] - if len(tied) > 0: - sel = tied - - if altloc: - for candidate in (altloc, ""): - hit = sel[sel["altloc"].astype(str) == candidate] - if len(hit) > 0: - return int(hit.iloc[0]["index"]) - for candidate in ("", "A"): - hit = sel[sel["altloc"].astype(str) == candidate] - if len(hit) > 0: - return int(hit.iloc[0]["index"]) - return int(sel.iloc[0]["index"]) - + altlocs = [str(topology.atoms.altloc[row]) for row in rows] + requested = str(altloc).strip() if altloc else "" + for candidate in ((requested, " ") if requested else ()) + (" ", "A"): + for row, alt in zip(rows, altlocs): + if alt == candidate: + return int(row) + return int(rows[0]) def _link_record_edges( - pdb: pd.DataFrame, + topology: Topology, links, disulfide_bonds: Optional[np.ndarray], verbose: int, @@ -649,12 +621,18 @@ def _link_record_edges( for a, b in disulfide_bonds: existing.add((min(int(a), int(b)), max(int(a), int(b)))) + residue_by_key: Dict[Tuple[str, int, str], List[int]] = {} + for r in range(topology.n_residues): + chain, resseq, icode = topology.residues.key(r) + residue_by_key.setdefault((chain, resseq, icode.strip()), []).append(r) + rows: List[Tuple[int, int]] = [] lengths: List[float] = [] n_unresolved = 0 for _, link in links.iterrows(): idx1 = _lookup_link_atom( - pdb, + topology, + residue_by_key, chainid=link["chainid1"], resseq=int(link["resseq1"]), icode=link["icode1"], @@ -663,7 +641,8 @@ def _link_record_edges( altloc=link["altloc1"], ) idx2 = _lookup_link_atom( - pdb, + topology, + residue_by_key, chainid=link["chainid2"], resseq=int(link["resseq2"]), icode=link["icode2"], @@ -725,12 +704,12 @@ def _block_with_values( def build_topology( - pdb: pd.DataFrame, + topology: Topology, cif_dict: Dict, + xyz: torch.Tensor, link_dict: Optional[Dict] = None, link_list=None, links=None, - xyz: Optional[torch.Tensor] = None, device=None, verbose: int = 0, ) -> Topology: @@ -740,12 +719,12 @@ def build_topology( half on its own, for callers that need the graph and no ideal geometry. """ topology, _, _ = build_topology_with_values( - pdb, + topology, cif_dict, + xyz, link_dict=link_dict, link_list=link_list, links=links, - xyz=xyz, device=device, verbose=verbose, ) @@ -753,24 +732,27 @@ def build_topology( def build_topology_with_values( - pdb: pd.DataFrame, + topology: Topology, cif_dict: Dict, + xyz: torch.Tensor, link_dict: Optional[Dict] = None, link_list=None, links=None, - xyz: Optional[torch.Tensor] = None, device=None, verbose: int = 0, ) -> Tuple[Topology, Dict[str, Dict], Dict[str, Dict]]: - """Build a topology from an atom table and the restraint dictionaries. + """Connect a node-only topology against the restraint dictionaries. Parameters ---------- - pdb : pandas.DataFrame - Atom table, with ``name``, ``element``, ``altloc``, ``chainid``, ``resseq``, - ``icode``, ``resname``, ``ATOM`` and ``index`` columns. + topology : Topology + Identity to connect, e.g. from :meth:`Topology.from_table`. Not modified; its + edges, if any, are ignored. cif_dict : dict Restraint dictionary keyed by residue name. + xyz : torch.Tensor + Cartesian coordinates in Å, shape ``(N, 3)``. Disulfides are detected by SG-SG + distance and proline omega classified cis or trans from them. link_dict : dict, optional Link-type definitions. Without it no inter-residue edges are built. link_list : pandas.DataFrame, optional @@ -778,9 +760,6 @@ def build_topology_with_values( links : pandas.DataFrame, optional Parsed PDB ``LINK`` records. Each record that resolves to two distinct atoms and does not duplicate an auto-detected disulfide contributes one bond edge. - xyz : torch.Tensor, optional - Coordinates, shape ``(N, 3)``. Needed only to detect disulfide links, which are - found by SG-SG distance. device : torch.device, optional Where to place the edge blocks. verbose : int, default 0 @@ -789,7 +768,7 @@ def build_topology_with_values( Returns ------- topology : Topology - The connectivity. + A new, connected topology over the same atoms. values : dict ``{edge_type: {origin: {property: tensor}}}`` for bonds, angles and torsions; ``{'chiral': {property: tensor}}`` and ``{'plane': {size: {property: tensor}}}`` @@ -797,10 +776,12 @@ def build_topology_with_values( extras : dict Products of the same pass that are not edges -- currently ``ramachandran``. """ - cols = _atom_columns(pdb) - nodes = build_residue_nodes( - cols["chain"], cols["resseq"], cols["icode"], cols["resname"] - ) + cols = _atom_columns(topology) + residue_nodes = topology.residues + nodes = { + field: getattr(residue_nodes, field) + for field in ("chain", "resseq", "icode", "resname", "atom_start", "atom_end") + } n_res = len(nodes["chain"]) names_by_residue = [ @@ -848,29 +829,29 @@ def build_topology_with_values( intra_planes, intra_plane_values = _match_intra_planes( match_cols, nodes, template_key, pp_cif ) - inter, inter_values, extras = _inter_residue_edges(pdb, link_dict, verbose) + inter, inter_values, extras = _inter_residue_edges( + PeptideResidues(topology, peptide_pairs, xyz.detach().cpu().numpy()), + link_dict, + verbose, + ) residue_of_row = {} for r in range(n_res): for row in range(int(nodes["atom_start"][r]), int(nodes["atom_end"][r])): residue_of_row[row] = r - disulfide_pairs: List[Tuple[int, int]] = [] - disulfide: Dict[str, np.ndarray] = {} - disulfide_values: Dict[str, Dict[str, np.ndarray]] = {} - if xyz is not None: - sg_rows = [ - row - for row in range(len(cols["name"])) - if cols["name"][row] == "SG" and cols["record"][row] == "ATOM" - ] - disulfide_pairs = find_disulfide_links(sg_rows, residue_of_row, xyz) - disulfide, disulfide_values = _disulfide_edges( - pdb, nodes, cols, residue_of_row, disulfide_pairs, link_dict, verbose - ) + sg_rows = [ + row + for row in range(len(cols["name"])) + if cols["name"][row] == "SG" and cols["record"][row] == "ATOM" + ] + disulfide_pairs = find_disulfide_links(sg_rows, residue_of_row, xyz) + disulfide, disulfide_values = _disulfide_edges( + topology, cols, residue_of_row, disulfide_pairs, link_dict, verbose + ) link_edges, link_atom_pairs, link_values = _link_record_edges( - pdb, links, disulfide.get("bond"), verbose + topology, links, disulfide.get("bond"), verbose ) # LINK edges carry ``index`` values, so lift them through that column. @@ -992,6 +973,8 @@ def build_topology_with_values( name=cols["name"], element=cols["element"], altloc=cols["altloc"], + is_hetatm=topology.atoms.is_hetatm.copy(), + charge=topology.atoms.charge.copy(), residue_of=torch.as_tensor( np.repeat( np.arange(n_res, dtype=np.int64), diff --git a/torchref/topology/builders.py b/torchref/topology/builders.py index 2cbab1d1..a0f457f3 100644 --- a/torchref/topology/builders.py +++ b/torchref/topology/builders.py @@ -1,19 +1,18 @@ -"""Restraint builders that walk the whole structure inside a single ``build()``. +"""Inter-residue restraint builders, and the dictionary preprocessing they share. -One ``build(pdb, cif_dict, device)`` call per restraint type handles every residue -internally -- callers do not loop -- or :func:`build_all_restraints` does all the -intra-residue types at once. Inter-residue links go through the -``InterResidue*Builder`` classes, whose ``build()`` returns directly like the others; -only their disulfide path is stateful -- it accumulates over ``process_disulfide_*`` -calls and emits nothing until ``finalize()`` (``finalize_disulfide()`` on the torsion -builder). +The ``InterResidue*Builder`` classes turn a link definition (``TRANS``, ``PTRANS``, +``disulf``) into edges over a topology. Their ``build()`` reads a +:class:`PeptideResidues` -- the peptide-linked residue pairs and each residue's +conformer maps, prepared once from the topology -- and returns directly; only the +disulfide path is stateful, accumulating over ``process_disulfide_*`` calls until +``finalize()`` (``finalize_disulfide()`` on the torsion builder). Edge indices are atom +rows of the topology. Nothing here is re-exported at the package level; import from ``torchref.topology.builders``. """ -from abc import ABC, abstractmethod -from typing import Any, Dict, Iterator, List, Mapping, Optional, Tuple +from typing import Dict, List, Optional, Tuple import numpy as np import pandas as pd @@ -21,205 +20,80 @@ from torchref.config import get_float_dtype, get_int_dtype -# Intra-residue restraint matchers -from torchref.topology.matchers import ( - match_angles, - match_bonds, - match_chirals, - match_torsions, -) - # ============================================================================= # Pre-processing utilities # ============================================================================= -class PreprocessedPDB: - """ - Pre-processed PDB data as NumPy arrays for fast iteration. - - Converts DataFrame to arrays once, computes residue boundaries, - enabling O(1) access to residue data without DataFrame operations. +def _conformer_maps(topology, residue: int) -> List[Dict[str, int]]: + """Atom-name-to-row maps for one residue, one per alternative conformation. - Supports altloc expansion: residues with alternate conformations can be - expanded into multiple conformations, each with common atoms plus the - specific altloc atoms. The constructor only normalizes altlocs and - computes residue boundaries; the expansion itself is performed on demand - by :meth:`get_altloc_conformations`, not at preprocessing time. + A residue without altlocs gives one map. One with altlocs gives one per altloc, + each holding the residue's blank-altloc atoms plus that altloc's own; one with no + blank atoms gives one per altloc on its own. """ + start = int(topology.residues.atom_start[residue]) + end = int(topology.residues.atom_end[residue]) + names = topology.atoms.name[start:end] + rows = np.arange(start, end, dtype=np.int64) + altlocs = topology.atoms.altloc[start:end] + unique = np.unique(altlocs) + if len(unique) == 1 and unique[0] == " ": + return [dict(zip(names, rows))] + common = altlocs == " " + maps = [] + for alt in unique: + if alt == " ": + continue + chosen = common | (altlocs == alt) + maps.append(dict(zip(names[chosen], rows[chosen]))) + return maps + + +def _atom_row(topology, residue: int, name: str) -> Optional[int]: + """Row of atom ``name`` in ``residue``: the blank altloc, else ``'A'``, else the first.""" + start = int(topology.residues.atom_start[residue]) + end = int(topology.residues.atom_end[residue]) + hits = np.nonzero(topology.atoms.name[start:end] == name)[0] + if len(hits) == 0: + return None + altlocs = topology.atoms.altloc[start:end][hits] + for wanted in (" ", "A"): + chosen = hits[altlocs == wanted] + if len(chosen): + return start + int(chosen[0]) + return start + int(hits[0]) + + +class PeptideResidues: + """What the peptide-link builders read, prepared once from a topology. - def __init__(self, pdb: pd.DataFrame): - """ - Initialize from PDB DataFrame. - - Parameters - ---------- - pdb : pd.DataFrame - PDB DataFrame with standard columns. - """ - self.n_atoms = len(pdb) - - # Core arrays - self.atom_names = pdb["name"].values.astype(str) - self.atom_indices = pdb["index"].values.astype(np.int64) - self.chain_ids = pdb["chainid"].values.astype(str) - self.resseqs = pdb["resseq"].values.astype(np.int64) - self.resnames = pdb["resname"].values.astype(str) - - # Optional columns - normalize altloc (treat '' as ' ') - if "altloc" in pdb.columns: - altlocs = pdb["altloc"].values.astype(str) - # Normalize: treat '' as ' ' (no altloc) - altlocs = np.where(altlocs == "", " ", altlocs) - self.altlocs = altlocs - else: - self.altlocs = np.full(self.n_atoms, " ", dtype=" Tuple[np.ndarray, np.ndarray, str]: - """ - Get atom data for a residue. - - Returns - ------- - atom_names : np.ndarray - Atom names for this residue. - atom_indices : np.ndarray - Global atom indices. - resname : str - Residue name. - """ - start = self.residue_starts[residue_idx] - end = self.residue_ends[residue_idx] - return ( - self.atom_names[start:end], - self.atom_indices[start:end], - self.residue_resnames[residue_idx], - ) - - def residue_keys( - self, mapping: Optional[Mapping[Tuple[str, int], str]] = None - ) -> List[str]: - """Return the restraint-dictionary key of each residue, by residue index. - - Parameters - ---------- - mapping : mapping, optional - ``{(chain_id, resseq): key}`` overriding the residue name for those - residues -- how a linked residue is pointed at a modified copy of its - component (see :mod:`torchref.topology.monomer.modifications`). Residues - absent from it, and every residue when this is None, key on their own - residue name, which is the unmodified behaviour. + Parameters + ---------- + topology : Topology + Supplies atom names, altlocs and residue ranges; edges are not needed. + pairs : sequence of tuple of int + ``(residue donating C, residue donating N)`` pairs, from + :func:`~torchref.topology.residue_graph.find_peptide_links`. + xyz : numpy.ndarray + Cartesian coordinates in Å, shape ``(N, 3)``; the torsion builder classifies + each proline's omega as cis or trans from them. + + Attributes + ---------- + conformer_maps : dict + ``{residue: [ {atom name: row}, ... ]}`` for every residue in a pair. + resnames : numpy.ndarray + Residue name per residue, shape ``(R,)``. + """ - Returns - ------- - list of str - One key per residue, indexed as ``residue_resnames``. - """ - if not mapping: - return list(self.residue_resnames) - return [ - mapping.get( - (str(self.residue_chain_ids[i]), int(self.residue_resseqs[i])), - self.residue_resnames[i], - ) - for i in range(self.n_residues) - ] - - def has_duplicate_atoms(self, residue_idx: int) -> bool: - """Check if residue has duplicate atom names (altlocs).""" - start = self.residue_starts[residue_idx] - end = self.residue_ends[residue_idx] - names = self.atom_names[start:end] - return len(names) != len(set(names)) - - def has_altlocs(self, residue_idx: int) -> bool: - """Check if residue has any alternate conformations.""" - start = self.residue_starts[residue_idx] - end = self.residue_ends[residue_idx] - altlocs = self.altlocs[start:end] - # Has altlocs if any altloc is not ' ' (normalized no-altloc marker) - return np.any(altlocs != " ") - - def get_altloc_conformations( - self, residue_idx: int - ) -> Iterator[Tuple[np.ndarray, np.ndarray, str]]: - """ - Iterate over altloc conformations for a residue. - - For residues without altlocs, yields once with all atoms. - For residues with altlocs, yields once per unique altloc, - each time with common atoms (no altloc) + altloc-specific atoms. - - Yields - ------ - atom_names : np.ndarray - Atom names for this conformation. - atom_indices : np.ndarray - Global atom indices. - resname : str - Residue name. - """ - start = self.residue_starts[residue_idx] - end = self.residue_ends[residue_idx] - - names = self.atom_names[start:end] - indices = self.atom_indices[start:end] - altlocs = self.altlocs[start:end] - resname = self.residue_resnames[residue_idx] - - unique_altlocs = np.unique(altlocs) - - if len(unique_altlocs) == 1 and unique_altlocs[0] == " ": - # No altlocs - yield all atoms once - yield names, indices, resname - elif " " in unique_altlocs: - # Has common atoms (no altloc) and altloc-specific atoms - # Common atoms mask - common_mask = altlocs == " " - common_names = names[common_mask] - common_indices = indices[common_mask] - - # Yield once per specific altloc (A, B, etc.) - for alt in unique_altlocs: - if alt == " ": - continue - alt_mask = altlocs == alt - # Combine common + altloc-specific - combined_names = np.concatenate([common_names, names[alt_mask]]) - combined_indices = np.concatenate([common_indices, indices[alt_mask]]) - yield combined_names, combined_indices, resname - else: - # No common atoms - yield each altloc separately - for alt in unique_altlocs: - alt_mask = altlocs == alt - yield names[alt_mask], indices[alt_mask], resname + def __init__(self, topology, pairs, xyz): + self.pairs = [(int(a), int(b)) for a, b in pairs] + self.resnames = np.char.strip(np.asarray(topology.residues.resname).astype(str)) + self.xyz = np.asarray(xyz, dtype=np.float64) + involved = sorted({r for pair in self.pairs for r in pair}) + self.conformer_maps = {r: _conformer_maps(topology, r) for r in involved} class PreprocessedCIF: @@ -379,743 +253,6 @@ def _preprocess_chirals(self, chirals_df: pd.DataFrame) -> Dict[str, np.ndarray] } -# ============================================================================= -# Residue pairing shared by the inter-residue builders -# ============================================================================= - - -def build_residue_conformation_maps( - pp_pdb: PreprocessedPDB, -) -> List[List[Dict[str, int]]]: - """Atom-name-to-index maps per residue per conformer (outer, then inner list). - - One map for a residue without altlocs; otherwise one per conformer, each - holding the common atoms plus that conformer's own. - """ - all_maps = [] - for res_idx in range(pp_pdb.n_residues): - res_maps = [] - for atom_names, atom_indices, _ in pp_pdb.get_altloc_conformations(res_idx): - res_maps.append(dict(zip(atom_names, atom_indices))) - all_maps.append(res_maps) - return all_maps - - -def find_consecutive_residue_pairs(pp_pdb: PreprocessedPDB) -> List[Tuple[int, int]]: - """Find pairs of residues numbered consecutively within one chain. - - Sequence numbering is the only criterion: no distance check, and insertion - codes are invisible here because :class:`PreprocessedPDB` groups residues on - ``(chain, resseq)`` alone. - """ - pairs = [] - by_chain: Dict[Any, List[Tuple[int, int]]] = {} - for res_idx in range(pp_pdb.n_residues): - chain = pp_pdb.residue_chain_ids[res_idx] - by_chain.setdefault(chain, []).append( - (pp_pdb.residue_resseqs[res_idx], res_idx) - ) - - for residues in by_chain.values(): - residues_sorted = sorted(residues, key=lambda x: x[0]) - for i in range(len(residues_sorted) - 1): - resseq_i, idx_i = residues_sorted[i] - resseq_next, idx_next = residues_sorted[i + 1] - if resseq_next == resseq_i + 1: - pairs.append((idx_i, idx_next)) - return pairs - - -def find_peptide_link_pairs(pp_pdb: PreprocessedPDB) -> List[Tuple[int, int]]: - """Consecutive residue pairs that actually carry a C-N peptide bond. - - Narrows :func:`find_consecutive_residue_pairs` to pairs where the first - residue has a ``C`` and the second an ``N`` -- the condition - :meth:`InterResidueBondBuilder.build` applies implicitly when it looks the two - atoms up. Deciding which residues are peptide-linked has to agree with that - builder exactly, or a residue could be given linked restraint targets without - getting the link itself. - - Parameters - ---------- - pp_pdb : PreprocessedPDB - Preprocessed atoms, already filtered to polymer atoms by the caller. - - Returns - ------- - list of tuple of int - ``(residue index donating C, residue index donating N)`` pairs. - """ - pairs = [] - for res_i, res_next in find_consecutive_residue_pairs(pp_pdb): - names_i, _, _ = pp_pdb.get_residue_data(res_i) - names_next, _, _ = pp_pdb.get_residue_data(res_next) - if "C" in names_i and "N" in names_next: - pairs.append((res_i, res_next)) - return pairs - - -# ============================================================================= -# Builder Base Class -# ============================================================================= - - -class RestraintBuilder(ABC): - """ - Abstract base class for restraint builders. - - All builders share the same API: - builder = SomeRestraintBuilder(verbose=0) - result = builder.build(pdb, cif_dict, device) - """ - - def __init__(self, verbose: int = 0): - """Initialize builder.""" - self.verbose = verbose - - @abstractmethod - def build( - self, - pdb: pd.DataFrame, - cif_dict: Dict, - device: torch.device, - sort_indices: bool = True, - residue_keys: Optional[Mapping[Tuple[str, int], str]] = None, - ) -> Optional[Dict[str, torch.Tensor]]: - """ - Build restraints from PDB and CIF data. - - Parameters - ---------- - pdb : pd.DataFrame - PDB DataFrame with atom data. - cif_dict : dict - CIF dictionary with restraints per residue type. - device : torch.device - Target device for output tensors. - sort_indices : bool, default True - Whether to sort by first atom index for cache efficiency. - residue_keys : mapping, optional - ``{(chain_id, resseq): cif_dict key}`` for residues whose restraints - come from somewhere other than their residue name -- a peptide-linked - residue draws from a modified copy of its component. Residues absent - from it key on their residue name, so None reproduces the plain - per-residue-type lookup. - - Returns - ------- - dict or None - Dictionary with restraint tensors, or None if no restraints found. - """ - pass - - -# ============================================================================= -# Bond Builder -# ============================================================================= - - -class BondRestraintBuilder(RestraintBuilder): - """ - Fast bond restraint builder. - - Usage: - builder = BondRestraintBuilder() - result = builder.build(pdb, cif_dict, device) - # result = {'indices': tensor, 'references': tensor, 'sigmas': tensor} - """ - - def build( - self, - pdb: pd.DataFrame, - cif_dict: Dict, - device: torch.device, - sort_indices: bool = True, - residue_keys: Optional[Mapping[Tuple[str, int], str]] = None, - ) -> Optional[Dict[str, torch.Tensor]]: - """Build all bond restraints.""" - # Pre-process data - pp_pdb = PreprocessedPDB(pdb) - pp_cif = PreprocessedCIF(cif_dict) - keys = pp_pdb.residue_keys(residue_keys) - - # Allocate work arrays - max_per_residue = 50 - work_idx1 = np.zeros(max_per_residue, dtype=np.int64) - work_idx2 = np.zeros(max_per_residue, dtype=np.int64) - work_refs = np.zeros(max_per_residue, dtype=np.float64) - work_sigmas = np.zeros(max_per_residue, dtype=np.float64) - - # Accumulate results - all_indices = [] - all_refs = [] - all_sigmas = [] - - # Process all residues (with altloc expansion) - for res_idx in range(pp_pdb.n_residues): - key = keys[res_idx] - - # Skip if no bond restraints for this residue type - if key not in pp_cif.bonds: - continue - - cif_bonds = pp_cif.bonds[key] - n_cif = len(cif_bonds["atom1"]) - - # Resize work arrays if needed - if n_cif > max_per_residue: - max_per_residue = n_cif * 2 - work_idx1 = np.zeros(max_per_residue, dtype=np.int64) - work_idx2 = np.zeros(max_per_residue, dtype=np.int64) - work_refs = np.zeros(max_per_residue, dtype=np.float64) - work_sigmas = np.zeros(max_per_residue, dtype=np.float64) - - # Iterate over altloc conformations (yields once if no altlocs) - for atom_names, atom_indices, _ in pp_pdb.get_altloc_conformations(res_idx): - count = match_bonds( - atom_names, - atom_indices, - cif_bonds["atom1"], - cif_bonds["atom2"], - cif_bonds["value"], - cif_bonds["sigma"], - work_idx1, - work_idx2, - work_refs, - work_sigmas, - ) - - if count > 0: - all_indices.append( - np.column_stack( - [work_idx1[:count].copy(), work_idx2[:count].copy()] - ) - ) - all_refs.append(work_refs[:count].copy()) - all_sigmas.append(work_sigmas[:count].copy()) - - # Finalize - if not all_indices: - return None - - indices = np.concatenate(all_indices, axis=0) - references = np.concatenate(all_refs) - sigmas = np.concatenate(all_sigmas) - - if sort_indices and len(indices) > 0: - order = np.argsort(indices[:, 0]) - indices = indices[order] - references = references[order] - sigmas = sigmas[order] - - # Replace zero sigmas - sigmas = np.where(sigmas == 0, 1e-4, sigmas) - - return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 - "references": torch.tensor(references, dtype=get_float_dtype(), device=device), - "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), - } - - -# ============================================================================= -# Angle Builder -# ============================================================================= - - -class AngleRestraintBuilder(RestraintBuilder): - """ - Fast angle restraint builder. - - Usage: - builder = AngleRestraintBuilder() - result = builder.build(pdb, cif_dict, device) - """ - - def build( - self, - pdb: pd.DataFrame, - cif_dict: Dict, - device: torch.device, - sort_indices: bool = True, - residue_keys: Optional[Mapping[Tuple[str, int], str]] = None, - ) -> Optional[Dict[str, torch.Tensor]]: - """Build all angle restraints.""" - pp_pdb = PreprocessedPDB(pdb) - pp_cif = PreprocessedCIF(cif_dict) - keys = pp_pdb.residue_keys(residue_keys) - - max_per_residue = 100 - work_idx1 = np.zeros(max_per_residue, dtype=np.int64) - work_idx2 = np.zeros(max_per_residue, dtype=np.int64) - work_idx3 = np.zeros(max_per_residue, dtype=np.int64) - work_refs = np.zeros(max_per_residue, dtype=np.float64) - work_sigmas = np.zeros(max_per_residue, dtype=np.float64) - - all_indices = [] - all_refs = [] - all_sigmas = [] - - for res_idx in range(pp_pdb.n_residues): - key = keys[res_idx] - - if key not in pp_cif.angles: - continue - - cif_angles = pp_cif.angles[key] - n_cif = len(cif_angles["atom1"]) - - if n_cif > max_per_residue: - max_per_residue = n_cif * 2 - work_idx1 = np.zeros(max_per_residue, dtype=np.int64) - work_idx2 = np.zeros(max_per_residue, dtype=np.int64) - work_idx3 = np.zeros(max_per_residue, dtype=np.int64) - work_refs = np.zeros(max_per_residue, dtype=np.float64) - work_sigmas = np.zeros(max_per_residue, dtype=np.float64) - - # Iterate over altloc conformations (yields once if no altlocs) - for atom_names, atom_indices, _ in pp_pdb.get_altloc_conformations(res_idx): - count = match_angles( - atom_names, - atom_indices, - cif_angles["atom1"], - cif_angles["atom2"], - cif_angles["atom3"], - cif_angles["value"], - cif_angles["sigma"], - work_idx1, - work_idx2, - work_idx3, - work_refs, - work_sigmas, - ) - - if count > 0: - all_indices.append( - np.column_stack( - [ - work_idx1[:count].copy(), - work_idx2[:count].copy(), - work_idx3[:count].copy(), - ] - ) - ) - all_refs.append(work_refs[:count].copy()) - all_sigmas.append(work_sigmas[:count].copy()) - - if not all_indices: - return None - - indices = np.concatenate(all_indices, axis=0) - references = np.concatenate(all_refs) - sigmas = np.concatenate(all_sigmas) - - if sort_indices and len(indices) > 0: - order = np.argsort(indices[:, 0]) - indices = indices[order] - references = references[order] - sigmas = sigmas[order] - - sigmas = np.where(sigmas == 0, 1e-4, sigmas) - - return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 - "references": torch.tensor(references, dtype=get_float_dtype(), device=device), - "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), - } - - -# ============================================================================= -# Torsion Builder -# ============================================================================= - - -class TorsionRestraintBuilder(RestraintBuilder): - """ - Fast torsion restraint builder. - - Usage: - builder = TorsionRestraintBuilder() - result = builder.build(pdb, cif_dict, device) - # result includes 'periods' tensor - """ - - def build( - self, - pdb: pd.DataFrame, - cif_dict: Dict, - device: torch.device, - sort_indices: bool = True, - residue_keys: Optional[Mapping[Tuple[str, int], str]] = None, - ) -> Optional[Dict[str, torch.Tensor]]: - """Build all torsion restraints.""" - pp_pdb = PreprocessedPDB(pdb) - pp_cif = PreprocessedCIF(cif_dict) - keys = pp_pdb.residue_keys(residue_keys) - - max_per_residue = 50 - work_idx1 = np.zeros(max_per_residue, dtype=np.int64) - work_idx2 = np.zeros(max_per_residue, dtype=np.int64) - work_idx3 = np.zeros(max_per_residue, dtype=np.int64) - work_idx4 = np.zeros(max_per_residue, dtype=np.int64) - work_refs = np.zeros(max_per_residue, dtype=np.float64) - work_sigmas = np.zeros(max_per_residue, dtype=np.float64) - work_periods = np.zeros(max_per_residue, dtype=np.int64) - - all_indices = [] - all_refs = [] - all_sigmas = [] - all_periods = [] - - for res_idx in range(pp_pdb.n_residues): - key = keys[res_idx] - - if key not in pp_cif.torsions: - continue - - cif_torsions = pp_cif.torsions[key] - n_cif = len(cif_torsions["atom1"]) - - if n_cif > max_per_residue: - max_per_residue = n_cif * 2 - work_idx1 = np.zeros(max_per_residue, dtype=np.int64) - work_idx2 = np.zeros(max_per_residue, dtype=np.int64) - work_idx3 = np.zeros(max_per_residue, dtype=np.int64) - work_idx4 = np.zeros(max_per_residue, dtype=np.int64) - work_refs = np.zeros(max_per_residue, dtype=np.float64) - work_sigmas = np.zeros(max_per_residue, dtype=np.float64) - work_periods = np.zeros(max_per_residue, dtype=np.int64) - - # Iterate over altloc conformations (yields once if no altlocs) - for atom_names, atom_indices, _ in pp_pdb.get_altloc_conformations(res_idx): - count = match_torsions( - atom_names, - atom_indices, - cif_torsions["atom1"], - cif_torsions["atom2"], - cif_torsions["atom3"], - cif_torsions["atom4"], - cif_torsions["value"], - cif_torsions["sigma"], - cif_torsions["period"], - work_idx1, - work_idx2, - work_idx3, - work_idx4, - work_refs, - work_sigmas, - work_periods, - ) - - if count > 0: - all_indices.append( - np.column_stack( - [ - work_idx1[:count].copy(), - work_idx2[:count].copy(), - work_idx3[:count].copy(), - work_idx4[:count].copy(), - ] - ) - ) - all_refs.append(work_refs[:count].copy()) - all_sigmas.append(work_sigmas[:count].copy()) - all_periods.append(work_periods[:count].copy()) - - if not all_indices: - return None - - indices = np.concatenate(all_indices, axis=0) - references = np.concatenate(all_refs) - sigmas = np.concatenate(all_sigmas) - periods = np.concatenate(all_periods) - - if sort_indices and len(indices) > 0: - order = np.argsort(indices[:, 0]) - indices = indices[order] - references = references[order] - sigmas = sigmas[order] - periods = periods[order] - - sigmas = np.where(sigmas == 0, 1e-4, sigmas) - - return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 - "references": torch.tensor(references, dtype=get_float_dtype(), device=device), - "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), - "periods": torch.tensor(periods, dtype=get_int_dtype(), device=device), - } - - -# ============================================================================= -# Plane Builder -# ============================================================================= - - -class PlaneRestraintBuilder(RestraintBuilder): - """ - Fast plane restraint builder. - - Returns planes grouped by atom count (e.g., '4_atoms', '5_atoms'). - - Usage: - builder = PlaneRestraintBuilder() - result = builder.build(pdb, cif_dict, device) - # result = {'4_atoms': {'indices': ..., 'sigmas': ...}, '5_atoms': {...}} - """ - - def build( - self, - pdb: pd.DataFrame, - cif_dict: Dict, - device: torch.device, - sort_indices: bool = True, - residue_keys: Optional[Mapping[Tuple[str, int], str]] = None, - ) -> Optional[Dict[str, Dict[str, torch.Tensor]]]: - """Build all plane restraints, grouped by atom count.""" - pp_pdb = PreprocessedPDB(pdb) - pp_cif = PreprocessedCIF(cif_dict) - keys = pp_pdb.residue_keys(residue_keys) - - # Group planes by size: {n_atoms: [(indices_array, sigmas_array), ...]} - planes_by_size: Dict[int, List[Tuple[np.ndarray, np.ndarray]]] = {} - - for res_idx in range(pp_pdb.n_residues): - key = keys[res_idx] - - if key not in pp_cif.planes: - continue - - # Iterate over altloc conformations (yields once if no altlocs) - for atom_names, atom_indices, _ in pp_pdb.get_altloc_conformations(res_idx): - # Build name to index map - name_to_idx = {name: idx for name, idx in zip(atom_names, atom_indices)} - - for plane_data in pp_cif.planes[key]: - plane_atom_names = plane_data["atoms"] - plane_sigmas = plane_data["sigmas"] - - # Find atoms in this plane - plane_indices = [] - plane_sigma_values = [] - for i, atom_name in enumerate(plane_atom_names): - if atom_name in name_to_idx: - plane_indices.append(name_to_idx[atom_name]) - plane_sigma_values.append(plane_sigmas[i]) - - # Need at least 3 atoms for a plane - if len(plane_indices) >= 3: - n_atoms = len(plane_indices) - indices_array = np.array(plane_indices, dtype=np.int64) - sigmas_array = np.array(plane_sigma_values, dtype=np.float64) - - if n_atoms not in planes_by_size: - planes_by_size[n_atoms] = [] - planes_by_size[n_atoms].append((indices_array, sigmas_array)) - - if not planes_by_size: - return None - - # Finalize each size group - result = {} - for n_atoms, planes_list in planes_by_size.items(): - indices = np.stack([p[0] for p in planes_list], axis=0) - sigmas = np.stack([p[1] for p in planes_list], axis=0) - - if sort_indices and len(indices) > 0: - order = np.argsort(indices[:, 0]) - indices = indices[order] - sigmas = sigmas[order] - - sigmas = np.where(sigmas == 0, 1e-4, sigmas) - - key = f"{n_atoms}_atoms" - result[key] = { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 - "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), - } - - return result - - -# ============================================================================= -# Chiral Builder -# ============================================================================= - - -class ChiralRestraintBuilder(RestraintBuilder): - """ - Fast chiral restraint builder. - - Usage: - builder = ChiralRestraintBuilder() - result = builder.build(pdb, cif_dict, device) - # result includes 'ideal_volumes' tensor - """ - - def build( - self, - pdb: pd.DataFrame, - cif_dict: Dict, - device: torch.device, - sort_indices: bool = True, - residue_keys: Optional[Mapping[Tuple[str, int], str]] = None, - ) -> Optional[Dict[str, torch.Tensor]]: - """Build all chiral restraints.""" - pp_pdb = PreprocessedPDB(pdb) - pp_cif = PreprocessedCIF(cif_dict) - keys = pp_pdb.residue_keys(residue_keys) - - max_per_residue = 20 - work_center = np.zeros(max_per_residue, dtype=np.int64) - work_idx1 = np.zeros(max_per_residue, dtype=np.int64) - work_idx2 = np.zeros(max_per_residue, dtype=np.int64) - work_idx3 = np.zeros(max_per_residue, dtype=np.int64) - work_signs = np.zeros(max_per_residue, dtype=np.float64) - work_sigmas = np.zeros(max_per_residue, dtype=np.float64) - - all_indices = [] - all_ideal_volumes = [] - all_sigmas = [] - - for res_idx in range(pp_pdb.n_residues): - key = keys[res_idx] - - if key not in pp_cif.chirals: - continue - - cif_chirals = pp_cif.chirals[key] - n_cif = len(cif_chirals["center"]) - - if n_cif > max_per_residue: - max_per_residue = n_cif * 2 - work_center = np.zeros(max_per_residue, dtype=np.int64) - work_idx1 = np.zeros(max_per_residue, dtype=np.int64) - work_idx2 = np.zeros(max_per_residue, dtype=np.int64) - work_idx3 = np.zeros(max_per_residue, dtype=np.int64) - work_signs = np.zeros(max_per_residue, dtype=np.float64) - work_sigmas = np.zeros(max_per_residue, dtype=np.float64) - - # Iterate over altloc conformations (yields once if no altlocs) - for atom_names, atom_indices, _ in pp_pdb.get_altloc_conformations(res_idx): - count = match_chirals( - atom_names, - atom_indices, - cif_chirals["center"], - cif_chirals["atom1"], - cif_chirals["atom2"], - cif_chirals["atom3"], - cif_chirals["volume_sign"], - cif_chirals["sigma"], - work_center, - work_idx1, - work_idx2, - work_idx3, - work_signs, - work_sigmas, - ) - - if count > 0: - all_indices.append( - np.column_stack( - [ - work_center[:count].copy(), - work_idx1[:count].copy(), - work_idx2[:count].copy(), - work_idx3[:count].copy(), - ] - ) - ) - # Ideal volume = sign * 2.5 (typical tetrahedral volume). - # For volume_sign 'both'/'either' the sign is 0.0, so this - # stores exactly 0.0. That 0.0 is a sentinel: the chiral - # target treats ideal_volume == 0 as an achiral centre and - # restrains |volume| toward 2.5 (not toward a target of 0). - # Note: chirals with an unknown sign were mapped to NaN by - # PreprocessedCIF._preprocess_chirals and dropped upstream - # by match_chirals, so they never reach here. - all_ideal_volumes.append(work_signs[:count].copy() * 2.5) - all_sigmas.append(work_sigmas[:count].copy()) - - if not all_indices: - return None - - indices = np.concatenate(all_indices, axis=0) - ideal_volumes = np.concatenate(all_ideal_volumes) - sigmas = np.concatenate(all_sigmas) - - if sort_indices and len(indices) > 0: - order = np.argsort(indices[:, 0]) - indices = indices[order] - ideal_volumes = ideal_volumes[order] - sigmas = sigmas[order] - - sigmas = np.where(sigmas == 0, 1e-4, sigmas) - - return { - "indices": torch.tensor(indices, dtype=torch.long, device=device), # dtype-ok: atom-index restraint tensor; torch indexing requires int64 - "ideal_volumes": torch.tensor( - ideal_volumes, dtype=get_float_dtype(), device=device - ), - "sigmas": torch.tensor(sigmas, dtype=get_float_dtype(), device=device), - } - - -# ============================================================================= -# Convenience function to build all restraints at once -# ============================================================================= - - -def build_all_restraints( - pdb: pd.DataFrame, cif_dict: Dict, device: torch.device, verbose: int = 0 -) -> Dict[str, Any]: - """ - Build every intra-residue restraint type at once, on ``device``. - - Returns - ------- - dict - Present keys only, so a type with no matches is absent rather than empty: - ``bond``/``angle`` as ``{indices, references, sigmas}``, ``torsion`` also - with ``periods``, ``chiral`` also with ``ideal_volumes``, and ``plane`` - nested by atom count (``{'4_atoms': {...}, '5_atoms': {...}}``). - """ - result = {} - - bond_result = BondRestraintBuilder(verbose).build(pdb, cif_dict, device) - if bond_result: - result["bond"] = bond_result - if verbose > 0: - print(f"Built {bond_result['indices'].shape[0]} bond restraints") - - angle_result = AngleRestraintBuilder(verbose).build(pdb, cif_dict, device) - if angle_result: - result["angle"] = angle_result - if verbose > 0: - print(f"Built {angle_result['indices'].shape[0]} angle restraints") - - torsion_result = TorsionRestraintBuilder(verbose).build(pdb, cif_dict, device) - if torsion_result: - result["torsion"] = torsion_result - if verbose > 0: - print(f"Built {torsion_result['indices'].shape[0]} torsion restraints") - - plane_result = PlaneRestraintBuilder(verbose).build(pdb, cif_dict, device) - if plane_result: - result["plane"] = plane_result - if verbose > 0: - n_planes = sum(v["indices"].shape[0] for v in plane_result.values()) - print(f"Built {n_planes} plane restraints") - - chiral_result = ChiralRestraintBuilder(verbose).build(pdb, cif_dict, device) - if chiral_result: - result["chiral"] = chiral_result - if verbose > 0: - print(f"Built {chiral_result['indices'].shape[0]} chiral restraints") - - return result - - # ============================================================================= # Fast Inter-Residue Builders # ============================================================================= @@ -1219,7 +356,7 @@ class InterResidueBondBuilder: Usage: builder = InterResidueBondBuilder() - result = builder.build(pdb, link_dict, device) + result = builder.build(residues, link_dict, device) # Or for disulfides (incremental): builder = InterResidueBondBuilder() @@ -1320,10 +457,9 @@ def count(self) -> int: def build( self, - pdb: pd.DataFrame, + residues: "PeptideResidues", link_dict: Dict, device: torch.device, - filter_atom_type: str = "ATOM", sort_indices: bool = True, ) -> Optional[Dict[str, torch.Tensor]]: """ @@ -1331,14 +467,12 @@ def build( Parameters ---------- - pdb : pd.DataFrame - PDB DataFrame. + residues : PeptideResidues + The linked residue pairs and their atoms. link_dict : dict Link dictionary with 'bonds' DataFrame. device : torch.device Target device. - filter_atom_type : str, optional - Filter to only this atom type (e.g., 'ATOM' for protein). sort_indices : bool Whether to sort output by first atom index. @@ -1356,17 +490,9 @@ def build( return None # Pre-process PDB - if filter_atom_type: - pdb = pdb[pdb["ATOM"] == filter_atom_type] - if pdb.empty: + conf_maps, pairs = residues.conformer_maps, residues.pairs + if not pairs: return None - pp_pdb = PreprocessedPDB(pdb) - - # Build per-conformation maps for each residue (altloc-aware) - conf_maps = build_residue_conformation_maps(pp_pdb) - - # Find consecutive residue pairs - pairs = find_consecutive_residue_pairs(pp_pdb) # Accumulate restraints all_indices = [] @@ -1423,11 +549,11 @@ class InterResidueAngleBuilder: Usage: builder = InterResidueAngleBuilder() - result = builder.build(pdb, link_dict, device) + result = builder.build(residues, link_dict, device) # Or for disulfides (incremental): builder = InterResidueAngleBuilder() - builder.process_disulfide_angles(res1_atoms, res2_atoms, link_angles) + builder.process_disulfide_angles(topology, res1, res2, link_angles) result = builder.finalize(device) """ @@ -1447,23 +573,11 @@ def reset(self): self._sigmas.clear() self._count = 0 - @staticmethod - def _get_atom_index(residue: pd.DataFrame, atom_name: str) -> Optional[int]: - """Get atom index from residue, handling alternate conformations.""" - atoms = residue[residue["name"] == atom_name] - if len(atoms) == 0: - return None - if " " in atoms["altloc"].values: - return int(atoms[atoms["altloc"] == " "].iloc[0]["index"]) - elif "A" in atoms["altloc"].values: - return int(atoms[atoms["altloc"] == "A"].iloc[0]["index"]) - else: - return int(atoms.iloc[0]["index"]) - def process_disulfide_angles( self, - res1_atoms: pd.DataFrame, - res2_atoms: pd.DataFrame, + topology, + res1_atoms: int, + res2_atoms: int, link_angles: pd.DataFrame, ) -> int: """ @@ -1471,10 +585,10 @@ def process_disulfide_angles( Parameters ---------- - res1_atoms : pd.DataFrame - First cysteine residue atoms. - res2_atoms : pd.DataFrame - Second cysteine residue atoms. + topology : Topology + Supplies the two residues' atom names and altlocs. + res1_atoms, res2_atoms : int + Residue indices of the two cysteines. link_angles : pd.DataFrame Angle definitions from disulfide link. @@ -1496,9 +610,9 @@ def process_disulfide_angles( res2 = res1_atoms if comp2 == "1" else res2_atoms res3 = res1_atoms if comp3 == "1" else res2_atoms - idx1 = self._get_atom_index(res1, atom1_name) - idx2 = self._get_atom_index(res2, atom2_name) - idx3 = self._get_atom_index(res3, atom3_name) + idx1 = _atom_row(topology, res1, atom1_name) + idx2 = _atom_row(topology, res2, atom2_name) + idx3 = _atom_row(topology, res3, atom3_name) if idx1 is not None and idx2 is not None and idx3 is not None: self._indices.append(np.array([[idx1, idx2, idx3]], dtype=np.int64)) @@ -1545,10 +659,9 @@ def count(self) -> int: def build( self, - pdb: pd.DataFrame, + residues: "PeptideResidues", link_dict: Dict, device: torch.device, - filter_atom_type: str = "ATOM", sort_indices: bool = True, next_resname_filter: Optional[str] = None, exclude_next_resname: Optional[str] = None, @@ -1557,14 +670,12 @@ def build( Parameters ---------- - pdb : pd.DataFrame - Atom DataFrame. + residues : PeptideResidues + The linked residue pairs and their atoms. link_dict : Dict Link definition dictionary containing angle parameters. device : torch.device Target device for tensors. - filter_atom_type : str, optional - Filter to this ATOM type (default "ATOM"). sort_indices : bool, optional Sort output by first atom index (default True). next_resname_filter : str, optional @@ -1582,13 +693,9 @@ def build( if link_data.angles is None: return None - if filter_atom_type: - pdb = pdb[pdb["ATOM"] == filter_atom_type] - if pdb.empty: + conf_maps, pairs = residues.conformer_maps, residues.pairs + if not pairs: return None - pp_pdb = PreprocessedPDB(pdb) - conf_maps = build_residue_conformation_maps(pp_pdb) - pairs = find_consecutive_residue_pairs(pp_pdb) all_indices = [] all_refs = [] @@ -1600,10 +707,10 @@ def build( for res_i_idx, res_next_idx in pairs: # Filter by next residue name if requested if next_resname_filter is not None: - if pp_pdb.residue_resnames[res_next_idx] != next_resname_filter: + if residues.resnames[res_next_idx] != next_resname_filter: continue if exclude_next_resname is not None: - if pp_pdb.residue_resnames[res_next_idx] == exclude_next_resname: + if residues.resnames[res_next_idx] == exclude_next_resname: continue for map_i in conf_maps[res_i_idx]: @@ -1659,13 +766,13 @@ class InterResidueTorsionBuilder: Usage: builder = InterResidueTorsionBuilder() - result = builder.build(pdb, link_dict, device) + result = builder.build(residues, link_dict, device) # result = {'phi': {...}, 'psi': {...}, 'omega': {...}, # 'ramachandran': {...}} # Or for disulfides (incremental): builder = InterResidueTorsionBuilder() - builder.process_disulfide_torsions(res1_atoms, res2_atoms, link_torsions) + builder.process_disulfide_torsions(topology, res1, res2, link_torsions) result = builder.finalize_disulfide(device) """ @@ -1687,23 +794,11 @@ def reset(self): self._disulfide_periods.clear() self._disulfide_count = 0 - @staticmethod - def _get_atom_index(residue: pd.DataFrame, atom_name: str) -> Optional[int]: - """Get atom index from residue, handling alternate conformations.""" - atoms = residue[residue["name"] == atom_name] - if len(atoms) == 0: - return None - if " " in atoms["altloc"].values: - return int(atoms[atoms["altloc"] == " "].iloc[0]["index"]) - elif "A" in atoms["altloc"].values: - return int(atoms[atoms["altloc"] == "A"].iloc[0]["index"]) - else: - return int(atoms.iloc[0]["index"]) - def process_disulfide_torsions( self, - res1_atoms: pd.DataFrame, - res2_atoms: pd.DataFrame, + topology, + res1_atoms: int, + res2_atoms: int, link_torsions: pd.DataFrame, ) -> int: """ @@ -1711,10 +806,10 @@ def process_disulfide_torsions( Parameters ---------- - res1_atoms : pd.DataFrame - First cysteine residue atoms. - res2_atoms : pd.DataFrame - Second cysteine residue atoms. + topology : Topology + Supplies the two residues' atom names and altlocs. + res1_atoms, res2_atoms : int + Residue indices of the two cysteines. link_torsions : pd.DataFrame Torsion definitions from disulfide link. @@ -1739,10 +834,10 @@ def process_disulfide_torsions( res3 = res1_atoms if comp3 == "1" else res2_atoms res4 = res1_atoms if comp4 == "1" else res2_atoms - idx1 = self._get_atom_index(res1, atom1_name) - idx2 = self._get_atom_index(res2, atom2_name) - idx3 = self._get_atom_index(res3, atom3_name) - idx4 = self._get_atom_index(res4, atom4_name) + idx1 = _atom_row(topology, res1, atom1_name) + idx2 = _atom_row(topology, res2, atom2_name) + idx3 = _atom_row(topology, res3, atom3_name) + idx4 = _atom_row(topology, res4, atom4_name) if idx1 is None or idx2 is None or idx3 is None or idx4 is None: continue @@ -1810,10 +905,9 @@ def _torsion_angle_np(coords: np.ndarray, i1, i2, i3, i4) -> float: def build( self, - pdb: pd.DataFrame, + residues: "PeptideResidues", link_dict: Dict, device: torch.device, - filter_atom_type: str = "ATOM", sort_indices: bool = True, ) -> Optional[Dict[str, Dict[str, torch.Tensor]]]: """ @@ -1828,18 +922,11 @@ def build( if link_data.torsions is None: return None - if filter_atom_type: - pdb = pdb[pdb["ATOM"] == filter_atom_type] - if pdb.empty: + conf_maps, pairs = residues.conformer_maps, residues.pairs + if not pairs: return None - pp_pdb = PreprocessedPDB(pdb) - conf_maps = build_residue_conformation_maps(pp_pdb) - pairs = find_consecutive_residue_pairs(pp_pdb) - # Build coordinate array for omega angle computation (cis/trans PRO) - max_idx = int(pdb["index"].max()) + 1 - coords_np = np.zeros((max_idx, 3)) - coords_np[pdb["index"].values] = pdb[["x", "y", "z"]].values + coords_np = residues.xyz # Separate accumulators for phi, psi, omega phi_data = {"indices": [], "periods": []} @@ -1866,8 +953,8 @@ def build( from torchref.topology.ramachandran import classify_residue for res_i_idx, res_next_idx in pairs: - resname_i = pp_pdb.residue_resnames[res_i_idx] - resname_next = pp_pdb.residue_resnames[res_next_idx] + resname_i = residues.resnames[res_i_idx] + resname_next = residues.resnames[res_next_idx] is_proline = resname_next == "PRO" for map_i in conf_maps[res_i_idx]: @@ -2041,7 +1128,7 @@ class InterResiduePlaneBuilder: Usage: builder = InterResiduePlaneBuilder() - result = builder.build(pdb, link_dict, device) + result = builder.build(residues, link_dict, device) """ def __init__(self, verbose: int = 0): @@ -2050,10 +1137,9 @@ def __init__(self, verbose: int = 0): def build( self, - pdb: pd.DataFrame, + residues: "PeptideResidues", link_dict: Dict, device: torch.device, - filter_atom_type: str = "ATOM", sort_indices: bool = True, ) -> Optional[Dict[str, Dict[str, torch.Tensor]]]: """Build all inter-residue plane restraints, grouped by atom count.""" @@ -2064,13 +1150,9 @@ def build( if link_data.planes is None: return None - if filter_atom_type: - pdb = pdb[pdb["ATOM"] == filter_atom_type] - if pdb.empty: + conf_maps, pairs = residues.conformer_maps, residues.pairs + if not pairs: return None - pp_pdb = PreprocessedPDB(pdb) - conf_maps = build_residue_conformation_maps(pp_pdb) - pairs = find_consecutive_residue_pairs(pp_pdb) # Group planes by atom count planes_by_size: Dict[int, List[Tuple[np.ndarray, np.ndarray]]] = {} @@ -2134,57 +1216,3 @@ def build( return result -# ============================================================================= -# Legacy-compatible ResidueIterator (for code that still needs it) -# ============================================================================= - - -class ResidueIterator: - """ - Efficient iterator over residues. - - .. deprecated:: - Retained only for backward compatibility with code that still relies - on residue-by-residue iteration. Prefer :func:`build_all_restraints` - or the individual ``builder.build()`` methods instead. - """ - - def __init__(self, pdb: pd.DataFrame, filter_atom_type: Optional[str] = None): - """Initialize with pre-grouping of residues.""" - if filter_atom_type is not None: - pdb = pdb[pdb["ATOM"] == filter_atom_type] - - self.pdb = pdb - self._grouped = pdb.groupby(["chainid", "resseq"], sort=False) - self.groups = list(self._grouped.groups.keys()) - - def __iter__(self) -> Iterator[Tuple[str, int, pd.DataFrame]]: - """Iterate over residues.""" - for chain_id, resseq in self.groups: - residue = self._grouped.get_group((chain_id, resseq)) - yield chain_id, resseq, residue - - def __len__(self) -> int: - """Return number of residues.""" - return len(self.groups) - - def get_consecutive_pairs(self) -> Iterator[Tuple[pd.DataFrame, pd.DataFrame]]: - """Iterate over consecutive residue pairs within each chain.""" - by_chain = {} - for chain_id, resseq in self.groups: - if chain_id not in by_chain: - by_chain[chain_id] = [] - by_chain[chain_id].append(resseq) - - for chain_id, resseqs in by_chain.items(): - resseqs_sorted = sorted(resseqs) - for i in range(len(resseqs_sorted) - 1): - resseq_i = resseqs_sorted[i] - resseq_next = resseqs_sorted[i + 1] - - if resseq_next != resseq_i + 1: - continue - - residue_i = self._grouped.get_group((chain_id, resseq_i)) - residue_next = self._grouped.get_group((chain_id, resseq_next)) - yield residue_i, residue_next diff --git a/torchref/topology/nonbonded.py b/torchref/topology/nonbonded.py index 7fe74309..b5b442c2 100644 --- a/torchref/topology/nonbonded.py +++ b/torchref/topology/nonbonded.py @@ -628,10 +628,14 @@ def filter_pairs( identity_combo: int, excl_hash: torch.Tensor, max_idx: int, - pdb, + topology, inter_residue_only: bool = True, ) -> torch.Tensor: - """Apply exclusion, residue, and altloc filters. Returns keep mask.""" + """Apply exclusion, residue, and altloc filters. Returns keep mask. + + Residues are the topology's ``(chain, resseq, icode)`` nodes, so atoms of residues + 100 and 100A are in different residues. + """ device = pair_atom_i.device N = len(pair_atom_i) keep = torch.ones(N, dtype=torch.bool, device=device) @@ -650,29 +654,21 @@ def filter_pairs( keep &= ~(is_excluded & is_intra_asu) # Same-residue filter – intra-ASU only + ai_np = pair_atom_i.cpu().numpy() + aj_np = pair_atom_j.cpu().numpy() if inter_residue_only: - chainid = pdb["chainid"].values - resseq = pdb["resseq"].values - ai_np = pair_atom_i.cpu().numpy() - aj_np = pair_atom_j.cpu().numpy() - same_res = ( - (chainid[ai_np] == chainid[aj_np]) - & (resseq[ai_np] == resseq[aj_np]) - ) + residue_of = topology.atoms.residue_of.cpu().numpy() + same_res = residue_of[ai_np] == residue_of[aj_np] same_res_t = torch.tensor(same_res, dtype=torch.bool, device=device) keep &= ~(same_res_t & is_intra_asu) # Altloc compatibility – intra-ASU only - if "altloc" in pdb.columns: - altloc = pdb["altloc"].values.astype(str) - altloc = np.where(np.isin(altloc, ["", " "]), " ", altloc) - ai_np = pair_atom_i.cpu().numpy() - aj_np = pair_atom_j.cpu().numpy() - alt_i = altloc[ai_np] - alt_j = altloc[aj_np] - incompat = (alt_i != " ") & (alt_j != " ") & (alt_i != alt_j) - incompat_t = torch.tensor(incompat, dtype=torch.bool, device=device) - keep &= ~(incompat_t & is_intra_asu) + altloc = topology.atoms.altloc + alt_i = altloc[ai_np] + alt_j = altloc[aj_np] + incompat = (alt_i != " ") & (alt_j != " ") & (alt_i != alt_j) + incompat_t = torch.tensor(incompat, dtype=torch.bool, device=device) + keep &= ~(incompat_t & is_intra_asu) return keep @@ -687,7 +683,7 @@ def build_vdw_restraints_gpu( vdw_radii: torch.Tensor, cell: "Cell", sg: "SpaceGroup", - pdb, + topology, exclusion_set: Set[Tuple[int, int]], cutoff: float = 5.0, sigma: float = 0.2, @@ -704,7 +700,8 @@ def build_vdw_restraints_gpu( ``(N,)`` van der Waals radii in Å. cell : Cell sg : SpaceGroup - pdb : DataFrame + topology : Topology + Residue membership and altlocs for the same-residue and altloc filters. exclusion_set : set of (int, int) bonded exclusion pairs cutoff : float Contact distance cutoff in Angstrom. @@ -824,7 +821,7 @@ def build_vdw_restraints_gpu( keep = filter_pairs( pair_atom_i, pair_atom_j, pair_combo_j, identity_combo, excl_hash, max_idx, - pdb, inter_residue_only, + topology, inter_residue_only, ) pair_atom_i = pair_atom_i[keep] @@ -891,219 +888,3 @@ def build_vdw_restraints_gpu( return result -# ------------------------------------------------------------------ # -# H-involving pair search (forward-time, called every evaluation) -# ------------------------------------------------------------------ # - -@torch.no_grad() -def find_h_vdw_pairs_gpu( - xyz_heavy: torch.Tensor, - xyz_h: torch.Tensor, - cell: "Cell", - sg: "SpaceGroup", - op_indices: torch.Tensor, - cell_offsets_valid: torch.Tensor, - grid_dims: torch.Tensor, - identity_combo: int, - n_heavy: int, - cutoff: float = 3.5, - h_excl_hash: Optional[torch.Tensor] = None, - pdb=None, - h_chainid_enc: Optional[torch.Tensor] = None, - h_resseq: Optional[torch.Tensor] = None, - inter_residue_only: bool = True, -) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Find VDW pairs involving at least one hydrogen atom. - - The heavy-atom search run over a combined (heavy + H) coordinate set, keeping - only pairs with a participant at index >= ``n_heavy``. Called on every forward - evaluation, so it reuses the cached grid and symop combos rather than - re-deriving them. - - Parameters - ---------- - xyz_heavy : (N_heavy, 3) Cartesian ASU heavy-atom positions - xyz_h : (N_h, 3) Cartesian ASU hydrogen positions - cell, sg : Cell, SpaceGroup - op_indices : (M,) cached valid symop indices - cell_offsets_valid : (M, 3) cached valid cell translations - grid_dims : (3,) cached grid dimensions - identity_combo : int - n_heavy : int - Number of heavy atoms (indices 0..n_heavy-1 are heavy) - cutoff : float - Cartesian cutoff (Å), tighter than heavy-atom search - h_excl_hash : (E,) sorted long - Hash tensor for H-specific 1-2/1-3 exclusions - pdb : DataFrame - Heavy-atom pdb for same-residue filtering - h_chainid_enc : (N_h,) long - Chain ID encoding for H atoms - h_resseq : (N_h,) long - Residue sequence number for H atoms - inter_residue_only : bool - - Returns - ------- - pair_atom_i, pair_atom_j, pair_combo_j : each (P,) long - Indices into the combined (heavy + H) array. atom_i is always - from the ASU (identity combo). - """ - from torchref.symmetry.spacegroup import SpaceGroup as SG - - device = xyz_heavy.device - fdtype = dtypes.float - - if not isinstance(sg, SG): - sg = SG(sg) - - # Combine heavy + H into a single coordinate set - xyz_all = torch.cat([xyz_heavy, xyz_h], dim=0) # (N_all, 3) - n_all = xyz_all.shape[0] - - empty = torch.tensor([], dtype=torch.long, device=device) # dtype-ok: empty index placeholder tensor; int64 required - if n_all == 0: - return empty, empty, empty - - # Convert to fractional - xyz_frac = cell.cartesian_to_fractional(xyz_all.detach().to(fdtype)) - - M = op_indices.shape[0] - - # Step 2: assign to grid (reusing cached grid_dims and symop combos) - flat_cell, atom_idx, combo_idx, cart_pos = assign_to_grid( - xyz_frac, cell, sg, op_indices, cell_offsets_valid, grid_dims - ) - - n_grid_total = grid_dims[0].item() * grid_dims[1].item() * grid_dims[2].item() - - # Step 3: sort into cell list - sort_order, unique_cells, starts, cell_lookup = build_cell_list( - flat_cell, n_grid_total - ) - cart_sorted = cart_pos[sort_order] - atom_idx_sorted = atom_idx[sort_order] - combo_idx_sorted = combo_idx[sort_order] - - # Step 4: find pairs (use the canonical-14-offset fast path) - pair_atom_i, pair_atom_j, pair_combo_j = find_pairs_periodic_grid_v2( - cart_sorted, atom_idx_sorted, combo_idx_sorted, - unique_cells, starts, cell_lookup, grid_dims, - cutoff, identity_combo, - ) - - if len(pair_atom_i) == 0: - return empty, empty, empty - - # Filter to keep only pairs involving at least one H - has_h = (pair_atom_i >= n_heavy) | (pair_atom_j >= n_heavy) - pair_atom_i = pair_atom_i[has_h] - pair_atom_j = pair_atom_j[has_h] - pair_combo_j = pair_combo_j[has_h] - - if len(pair_atom_i) == 0: - return empty, empty, empty - - # Step 5: filtering - - is_intra_asu = pair_combo_j == identity_combo - - # Bonded exclusions (1-2 H-parent, 1-3 H-parent_neighbor) — intra-ASU - if h_excl_hash is not None and len(h_excl_hash) > 0 and is_intra_asu.any(): - max_idx = max(n_all, int(pair_atom_i.max().item()) + 1, - int(pair_atom_j.max().item()) + 1) - norm_i = torch.minimum(pair_atom_i, pair_atom_j) - norm_j = torch.maximum(pair_atom_i, pair_atom_j) - pair_hash = norm_i * max_idx + norm_j - ins = torch.searchsorted(h_excl_hash, pair_hash) - ins = ins.clamp(max=len(h_excl_hash) - 1) - is_excluded = h_excl_hash[ins] == pair_hash - keep = ~(is_excluded & is_intra_asu) - pair_atom_i = pair_atom_i[keep] - pair_atom_j = pair_atom_j[keep] - pair_combo_j = pair_combo_j[keep] - is_intra_asu = is_intra_asu[keep] - - if len(pair_atom_i) == 0: - return empty, empty, empty - - # Same-residue filter — intra-ASU only - if inter_residue_only and pdb is not None: - pdb_chainid = pdb["chainid"].values - pdb_resseq = pdb["resseq"].values.astype(np.int64) - - # Build combined chain/resseq arrays (heavy from pdb, H from topology) - if h_chainid_enc is not None and h_resseq is not None: - # For heavy atoms, encode chain IDs consistently - chain_vals = pdb_chainid.astype(str) - unique_chains = np.unique(chain_vals) - chain_to_int = {c: i for i, c in enumerate(unique_chains)} - heavy_chain_enc = np.array([chain_to_int.get(c, -1) for c in chain_vals], - dtype=np.int64) - heavy_resseq = pdb_resseq - - all_chain_enc = np.concatenate([ - heavy_chain_enc, - h_chainid_enc.cpu().numpy(), - ]) - all_resseq = np.concatenate([ - heavy_resseq, - h_resseq.cpu().numpy(), - ]) - else: - all_chain_enc = None - - if all_chain_enc is not None: - ai_np = pair_atom_i.cpu().numpy() - aj_np = pair_atom_j.cpu().numpy() - same_res = ( - (all_chain_enc[ai_np] == all_chain_enc[aj_np]) - & (all_resseq[ai_np] == all_resseq[aj_np]) - ) - same_res_t = torch.tensor(same_res, dtype=torch.bool, device=device) - keep = ~(same_res_t & is_intra_asu) - pair_atom_i = pair_atom_i[keep] - pair_atom_j = pair_atom_j[keep] - pair_combo_j = pair_combo_j[keep] - - if len(pair_atom_i) == 0: - return empty, empty, empty - - # Altloc compatibility — intra-ASU, only relevant for heavy atoms - if pdb is not None and "altloc" in pdb.columns: - is_intra_asu = pair_combo_j == identity_combo - # Only check altloc for pairs where both are heavy atoms - both_heavy = (pair_atom_i < n_heavy) & (pair_atom_j < n_heavy) & is_intra_asu - if both_heavy.any(): - altloc = pdb["altloc"].values.astype(str) - altloc = np.where(np.isin(altloc, ["", " "]), " ", altloc) - ai_np = pair_atom_i[both_heavy].cpu().numpy() - aj_np = pair_atom_j[both_heavy].cpu().numpy() - incompat = (altloc[ai_np] != " ") & (altloc[aj_np] != " ") & (altloc[ai_np] != altloc[aj_np]) - reject = torch.zeros(len(pair_atom_i), dtype=torch.bool, device=device) - reject[both_heavy] = torch.tensor(incompat, dtype=torch.bool, device=device) - keep = ~reject - pair_atom_i = pair_atom_i[keep] - pair_atom_j = pair_atom_j[keep] - pair_combo_j = pair_combo_j[keep] - - if len(pair_atom_i) == 0: - return empty, empty, empty - - # Deduplicate - dedup_hash = pair_atom_i * (n_all * M) + pair_atom_j * M + pair_combo_j - _, inverse, counts = torch.unique( - dedup_hash, return_inverse=True, return_counts=True - ) - # MPS does not support int64 scatter_reduce; use the configured int dtype. - _int_dtype = dtypes.int - inverse_i = inverse.to(_int_dtype) - perm = torch.arange(len(inverse), device=device, dtype=_int_dtype) - first_occ = torch.full( - (counts.shape[0],), len(inverse), device=device, dtype=_int_dtype - ) - first_occ.scatter_reduce_(0, inverse_i, perm, reduce="amin") - first_mask = torch.zeros(len(pair_atom_i), dtype=torch.bool, device=device) - first_mask[first_occ.long()] = True - - return pair_atom_i[first_mask], pair_atom_j[first_mask], pair_combo_j[first_mask] diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index b0dd4590..c252d02f 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -42,14 +42,16 @@ class Restraints(DeviceMixin, DebugMixin, Module): Parameters ---------- - pdb : pd.DataFrame, optional - DataFrame containing atomic structure data. If None, creates empty shell. + topology : Topology, optional + The atoms to restrain -- a node-only topology is enough + (:meth:`~torchref.topology.Topology.from_table`); it is connected here. If + None, creates an empty shell. cif_path : str or list of str, optional Path to the CIF restraints dictionary file(s). xyz : torch.Tensor, optional Cartesian coordinates in Å, shape ``(n_atoms, 3)``, that the topology and the - first pair list are built over. Defaults to the ``x``/``y``/``z`` columns of - ``pdb``. The build lands on this tensor's device. Not retained. + first pair list are built over. Required with ``topology``. The build lands on + this tensor's device. Not retained. cell : Cell, optional Crystallographic unit cell. Together with ``spacegroup``, enables symmetry-aware VDW restraints (contacts with symmetry mates). Without both, @@ -72,7 +74,8 @@ class Restraints(DeviceMixin, DebugMixin, Module): Restraint groups as ``restraints["bond"]["intra"]["indices"]``. A plain nested dict; the per-origin indices are views into ``topology``'s edge blocks. topology : Topology - The connectivity the geometry restraints are defined over. + The connected topology the geometry restraints are defined over; a new object, + the input is not modified. cif_dict : dict Parsed CIF restraints keyed by residue type; ``missing_residues`` lists the types that could not be resolved. @@ -80,13 +83,13 @@ class Restraints(DeviceMixin, DebugMixin, Module): Riding-hydrogen map, built only when the model carries no hydrogens of its own. Empty otherwise; see :mod:`torchref.topology.riding`. link_dict, link_list - Link-type definitions from the monomer library, set only when ``pdb`` + Link-type definitions from the monomer library, set only when ``topology`` was provided. """ def __init__( self, - pdb: pd.DataFrame = None, + topology=None, cif_path=None, xyz: torch.Tensor = None, cell=None, @@ -117,27 +120,18 @@ def __init__( self._torsion_max_period = 1 # Empty initialization - if pdb is None: - self.pdb = None + if topology is None: self.cif_dict = {} self.unique_residues = [] return - - # Full initialization with pdb - from torchref.topology.nonbonded import vdw_radii_for_elements - - self.pdb = pdb if xyz is None: - xyz = torch.tensor(pdb[["x", "y", "z"]].values, dtype=get_float_dtype()) + raise ValueError("Restraints over a topology need the coordinates, xyz=") + + self._nodes = topology self._vdw_radii = torch.tensor( - vdw_radii_for_elements(pdb["element"]), dtype=get_float_dtype() + topology.atoms.vdw_radii, dtype=get_float_dtype() ) - self.unique_residues = pdb.resname.unique() - self.unique_residues = [ - residue - for residue in self.unique_residues - if self.pdb.loc[self.pdb["resname"] == residue, "name"].nunique() > 1 - ] + self.unique_residues = self._multi_atom_resnames(topology) # Parse CIF files self._load_cif_dictionaries(cif_path) @@ -153,6 +147,43 @@ def __init__( if self.verbose > 0: self.summary() + @staticmethod + def _multi_atom_resnames(topology) -> list: + """Residue names, in first-seen order, whose atoms carry more than one name. + + Single-atom residues (ions, lone waters) need no dictionary lookup. + """ + names_by_resname: dict = {} + resnames = np.char.strip(topology.residues.resname.astype(str)) + for r, resname in enumerate(resnames): + rows = topology.residues.atom_rows(r) + names_by_resname.setdefault(str(resname), set()).update( + topology.atoms.name[rows.start : rows.stop].tolist() + ) + return [name for name, atoms in names_by_resname.items() if len(atoms) > 1] + + def _riding_table(self, xyz: torch.Tensor) -> pd.DataFrame: + """The identity-plus-coordinates table :mod:`torchref.topology.riding` reads. + + That module still takes an atom table; it goes when the phantom-hydrogen path + is deleted. + """ + columns = self.topology.columns() + coords = xyz.detach().cpu().numpy() + return pd.DataFrame( + { + "name": columns["name"], + "element": columns["element"], + "resname": columns["resname"], + "chainid": columns["chain"], + "resseq": columns["resseq"], + "icode": columns["icode"], + "x": coords[:, 0], + "y": coords[:, 1], + "z": coords[:, 2], + } + ) + @property def restraints(self) -> dict: """Restraint groups as ``[edge type][origin][property]``. @@ -287,12 +318,12 @@ def build_restraints(self, xyz: torch.Tensor): from torchref.topology import build_topology_with_values self.topology, self._values, extras = build_topology_with_values( - self.pdb, + self._nodes, self.cif_dict, + xyz.detach().to(device), link_dict=self.link_dict, link_list=self.link_list, links=self.links, - xyz=xyz.detach().to(device), device=device, verbose=self.verbose, ) @@ -335,7 +366,7 @@ def _build_h_exclusion_hash(self, h_topo, device): if h_topo is None or h_topo.n_hydrogens == 0: return torch.tensor([], dtype=torch.long, device=device) # dtype-ok: empty index tensor; int64 index required - n_heavy = len(self.pdb) + n_heavy = self.topology.n_atoms n_h = h_topo.n_hydrogens exclusions = set() @@ -434,7 +465,7 @@ def _build_vdw_restraints( vdw_radii=radii_cpu, cell=cell_cpu, sg=sg_cpu, - pdb=self.pdb, + topology=self.topology, exclusion_set=self.topology.atoms.exclusions_from_restraint_edges(), cutoff=cutoff, sigma=sigma, @@ -460,12 +491,12 @@ def _build_vdw_restraints( build_hydrogen_topology, ) - elements = self.pdb["element"].astype(str).str.strip().values - if (elements == "H").any(): + if bool(self.topology.atoms.is_hydrogen.any()): self._h_topo = HydrogenTopology(device=cpu) else: + riding_table = self._riding_table(xyz) self._h_topo = build_hydrogen_topology( - pdb=self.pdb, + pdb=riding_table, device=cpu, verbose=self.verbose, ) @@ -477,7 +508,7 @@ def _build_vdw_restraints( build_h_candidate_pairs( h_topo=self._h_topo, vdw_data=vdw_data, - pdb=self.pdb, + pdb=riding_table, h_excl_hash=self._h_excl_hash, device=cpu, verbose=self.verbose, diff --git a/torchref/topology/topology.py b/torchref/topology/topology.py index 44b1a2cd..235576eb 100644 --- a/torchref/topology/topology.py +++ b/torchref/topology/topology.py @@ -54,8 +54,9 @@ def identity_columns(pdb: "pandas.DataFrame") -> Dict[str, np.ndarray]: Returns ------- dict - :data:`IDENTITY_COLUMNS`, each shape ``(N,)``. A blank altloc reads as - ``' '``; a missing or non-numeric charge as 0. + :data:`IDENTITY_COLUMNS`, each shape ``(N,)``. Strings are stripped, so a + padded ``' ALA'`` reads as ``'ALA'``; a blank altloc reads as ``' '``; a + missing or non-numeric charge as 0. """ import pandas as pd @@ -63,7 +64,7 @@ def identity_columns(pdb: "pandas.DataFrame") -> Dict[str, np.ndarray]: def text(column: str, default: str) -> np.ndarray: if column in pdb.columns: - return pdb[column].values.astype(str) + return np.char.strip(pdb[column].values.astype(str)) return np.full(n, default) altloc = text("altloc", "") From c2fb2d58e8bef3d99d2be497d82bda30d8a90a4d Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Tue, 29 Sep 2026 12:06:24 +0200 Subject: [PATCH 201/250] Keep atom identity on the topology and values on the wrappers A model no longer holds an atom table. ModelContext keeps a node-only topology; ModelContext.from_atoms splits the table read at construction into that topology and AtomValues, which _install_parameters consumes. Model.to_dataframe joins identity and the wrappers' current values for the writers, metadata, IHM and checkpoints, so update_pdb is gone and Model.pdb is a deprecated read-only view. Selections, occupancy and altloc groups, the rigid-body filter, strip/select/hydrogenate and every consumer outside the model classes read the topology. Occupancy and altloc grouping keep pandas' sorted-key order, so checkpoints restore with the same occupancy groups; written PDB and mmCIF files are byte-identical to before. Copying a context copies its LINK records as a table instead of turning them into column names. Co-Authored-By: Claude Opus 5.5 (1M context) --- AGENTS.md | 2 +- docs/changelog.rst | 3 + docs/quickstart.rst | 2 +- docs/user_guide/testing.rst | 2 +- .../integration/test_adp_field_refinement.py | 1 - .../model/test_riding_water_completion.py | 8 +- torchref/cli/collection_difference_refine.py | 20 +- torchref/cli/validate_ded.py | 4 +- .../ensemble/ensemble_amber_kl.py | 2 +- .../experimental/ensemble/ensemble_model.py | 2 +- .../ensemble/ensemble_refinement.py | 2 +- torchref/experimental/targets/amber_target.py | 19 +- torchref/io/ihm.py | 6 +- torchref/io/metadata.py | 16 +- torchref/model/context.py | 517 +++++++++++------- torchref/model/model.py | 494 +++++++++-------- torchref/model/model_ft.py | 4 +- torchref/refinement/base_refinement.py | 26 +- torchref/refinement/targets/similarity.py | 32 +- torchref/scaling/solvent.py | 5 +- 20 files changed, 649 insertions(+), 518 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index 318346ae..9346e41f 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -186,7 +186,7 @@ Black, 88 columns, `isort` with the black profile. Ruff lint with |---|---| | `base/` | Low-level math and crystallography. `coordinates/` (Cartesian↔fractional), `reciprocal/` (basis, HKL, d-spacing, interpolation, symmetry), `direct_summation/` (F_calc by summation; eager + Triton), `electron_density/` (real-space splatting with CPU/CUDA/MPS kernels, solvent mask, radius policy), `fourier/` (FFT and grids), `scattering/` (form-factor and anomalous tables), `metrics/` (R-factors, binwise scale, loss), `targets/` (the *kernels* behind refinement targets, eager + `triton/`), `french_wilson.py`, `math_torch.py`, `alignment/` | | `io/` | `ReflectionData`, `DatasetCollection`, `FcalcDataset`; MTZ / PDB / CIF / IHM readers and writers; `read_mtz` / `read_pdb` / `read_cif` | -| `model/` | `Model` (refinable atomic parameters), `ModelContext` (the cell, space group, atom table, links, provenance, hydrogen policy and geometry restraints a model is loaded with — `model.cell` / `.spacegroup` / `.pdb` / `.restraints` forward to it, the rest is `model.ctx.*`; `ModelContext.from_atoms` is the one place an atom table is settled and `Model._install_parameters` the one place wrappers are built), `ModelFT` (adds F_calc via `SfFFT` or `SfDS`), `MixedModel`, `ModelCollection`, and the parametrizations in `parameter_wrappers.py` / `rigid_xyz.py` that decide what is refinable | +| `model/` | `Model` (refinable atomic parameters), `ModelContext` (the cell, space group, atom identity as a node-only `ctx.topology`, links, provenance, hydrogen policy and geometry restraints a model is loaded with — `model.cell` / `.spacegroup` / `.restraints` forward to it, the rest is `model.ctx.*`). **Identity is read from `ctx.topology`, values only through the wrappers (`model.xyz()`, `.adp()`, `.u()`, `.occupancy()`, `aniso_flag`); a pandas atom table appears only at construction (`ModelContext.from_atoms`, the one place a table is settled) and output (`Model.to_dataframe()`) — never read `model.pdb`, a deprecated view.** `Model._install_parameters` is the one place wrappers are built. `ModelFT` (adds F_calc via `SfFFT` or `SfDS`), `MixedModel`, `ModelCollection`, and the parametrizations in `parameter_wrappers.py` / `rigid_xyz.py` that decide what is refinable | | `refinement/` | Drivers (`Refinement`, `LBFGSRefinement`, `RigidBodyRefinementStep`), `targets/` (`xray/`, `geometry/`, `adp/`, `collection/`, `combined.py`), `weighting/`, `optimizers/` (annealing, Langevin, preconditioned/seeded L-BFGS), `model_error_estimation/` (σ_A, σ_M), `loss_state.py`, `logger.py` | | `topology/` | The connectivity graph (`Topology`, `AtomGraph`, `ResidueGraph`) and the restraint layer over it (`Restraints`: bonds, angles, torsions, planes, chirals, VDW pair list), hydrogen generation (`hydrogens.py`) and riding frames. Built from the CCP4 Monomer Library, resolved lazily via `monomer.library.get_library_manager()` — importing this package must not trigger a library download. `Restraints` holds no reference to a model: evaluations take the coordinates they score | | `scaling/` | `ScalerBase` (model-independent), `Scaler`, `CollectionScaler`, `SolventModel` (k_sol, B_sol) | diff --git a/docs/changelog.rst b/docs/changelog.rst index 6ea02354..c9a7653b 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,9 @@ Changelog Unreleased ---------- +- A model no longer keeps an atom table. Atom identity lives on ``model.ctx.topology`` (a node-only ``Topology``) and every refinable value only on the parameter wrappers; the table is read once at construction (``ModelContext.from_atoms`` splits it into the topology and ``AtomValues``) and written by ``Model.to_dataframe()``, which joins identity and current values afresh on every call. ``Model.update_pdb`` is removed, ``Model.pdb`` is a deprecated read-only view of ``to_dataframe()`` (writing into it changes nothing), ``model.n_atoms`` replaces ``len(model.pdb)``, and checkpoints keep storing the table under ``"pdb"`` so older ones still restore +- ``Model.strip_altlocs`` picks the kept conformer by current occupancy rather than the occupancies the file was loaded with +- Copying a model's context copies its LINK records as a table; it previously turned them into a list of column names, which broke a later restraint rebuild on the copy - Restraints are built from a topology rather than an atom table: ``Restraints(topology=..., xyz=...)`` and ``build_topology(_with_values)(topology, cif_dict, xyz, ...)`` take a node-only ``Topology`` (``Topology.from_table``), and the peptide-link builders pair residues along the residue graph's links, so residues with insertion codes (100, 100A, 101) are now peptide-linked and restrained like any other; structures without insertion codes get identical edges and pair lists. A table without ``icode`` or ``ATOM`` columns no longer fails, and padded names read like clean ones. The unused per-residue restraint builders, ``build_all_restraints``, ``ResidueIterator``, ``PreprocessedPDB`` and ``find_h_vdw_pairs_gpu`` are removed - Atom identity is available without building restraints: ``Topology.from_table`` gives a node-only topology (names, elements, altlocs, residues, chains, record types and charges; ``connected=False``), with ``Topology.select`` for Phenix-style selections, ``is_water`` / ``is_polymer`` masks, cached ``AtomGraph.atomic_number`` / ``vdw_radii`` and ``Topology.with_hydrogens`` for inserting a hydrogen plan. Selections are evaluated by one recursive-descent parser (``torchref.utils.selection``); a selection with more than one parenthesised group, which the old parser could resolve to every atom, now selects what it says - Hydrogens are one policy on the model context, set at construction: ``hydrogens`` (``keep`` / ``add`` / ``strip``) settles the atom table when it loads and ``hydrogen_mode`` (``atoms`` / ``riding``) how hydrogen rows are parametrised; ``riding`` with ``strip`` raises. It replaces ``strip_H``, ``add_hydrogens`` and ``hydrogen_mode="free"`` on ``Model``, ``ModelFT``, ``Refinement``, ``EnsembleModel`` and ``cli._common.load_model`` (checkpoints written with them still restore), and ``torchref.refine`` takes ``--hydrogens`` / ``--hydrogen-mode`` instead of ``--add-hydrogens``. The deprecated ``exclude_H_from_sf`` alias is removed diff --git a/docs/quickstart.rst b/docs/quickstart.rst index dfba917f..4dac6033 100644 --- a/docs/quickstart.rst +++ b/docs/quickstart.rst @@ -68,7 +68,7 @@ TorchRef supports multiple file formats: model = read_pdb(f"{ROOT_TORCHREF}/example_notebooks/1DAW.pdb") - print(f"Number of atoms: {len(model.pdb)}") + print(f"Number of atoms: {model.n_atoms}") .. testoutput:: :options: +ELLIPSIS diff --git a/docs/user_guide/testing.rst b/docs/user_guide/testing.rst index cca19834..3cf7e36f 100644 --- a/docs/user_guide/testing.rst +++ b/docs/user_guide/testing.rst @@ -168,7 +168,7 @@ free of file I/O: model = Model() model.load_cif(str(sample_cif_file)) assert model.initialized - assert len(model.pdb) > 0 + assert model.n_atoms > 0 Cover the edge cases that actually bite here: empty selections, degenerate geometry (which is where gradients go NaN), and non-default dtype/device. diff --git a/tests/integration/test_adp_field_refinement.py b/tests/integration/test_adp_field_refinement.py index 6969bd37..88ace7f1 100644 --- a/tests/integration/test_adp_field_refinement.py +++ b/tests/integration/test_adp_field_refinement.py @@ -389,7 +389,6 @@ def test_cli_field_mode_end_to_end(files, tmp_path): assert ref.model.adp_field.n_nodes == 9 ref.refine_adp() out = tmp_path / "out.pdb" - ref.model.update_pdb() ref.model.write_pdb(str(out)) assert out.exists() and out.stat().st_size > 0 text = out.read_text() diff --git a/tests/unit/model/test_riding_water_completion.py b/tests/unit/model/test_riding_water_completion.py index 0f7df251..22ffa20f 100644 --- a/tests/unit/model/test_riding_water_completion.py +++ b/tests/unit/model/test_riding_water_completion.py @@ -52,12 +52,10 @@ def test_add_with_riding_completes_every_water(heavy_model, model_class): @pytest.mark.unit def test_partial_water_is_completed_without_moving_its_hydrogen(heavy_model): """A water that arrives with one hydrogen receives just its missing partner.""" - water = ( - heavy_model.pdb[heavy_model.pdb.resname.str.strip().eq("HOH")].iloc[:1].copy() - ) + table = heavy_model.to_dataframe() + water = table[table.resname.eq("HOH")].iloc[:1].copy() complete = heavy_model._derive(water, hydrogens="add") - complete.update_pdb() - partial_table = complete.pdb.iloc[:2].copy() + partial_table = complete.to_dataframe().iloc[:2].copy() kept = heavy_model._derive(partial_table, hydrogens="keep") assert len(kept.pdb) == 2 diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index dade6c23..adb01973 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -1626,17 +1626,17 @@ def _build_metadata(model, data, r_work, r_free): meta.n_reflections_all = n_all meta.percent_free = 100.0 * n_test / n_all if n_all > 0 else None - # B-factor statistics - pdb = model.pdb - bvals = pdb["tempfactor"] + # B-factor statistics, from the written B column + bvals = model.to_dataframe()["tempfactor"] meta.b_mean_overall = float(bvals.mean()) meta.b_min = float(bvals.min()) meta.b_max = float(bvals.max()) # Atom counts - meta.n_atoms_total = len(pdb) - meta.n_atoms_protein = int((pdb["ATOM"] == "ATOM").sum()) - meta.n_atoms_solvent = int((pdb["ATOM"] == "HETATM").sum()) + is_hetatm = model.ctx.topology.atoms.is_hetatm + meta.n_atoms_total = len(is_hetatm) + meta.n_atoms_protein = int((~is_hetatm).sum()) + meta.n_atoms_solvent = int(is_hetatm.sum()) # Geometry deviations if model.ctx.initialized and model.ctx.restraints is not None: @@ -1682,8 +1682,8 @@ def _build_metadata(model, data, r_work, r_free): # --- Write merged deposition CIF (if no altlocs) --- merged_cif_out = str(outdir / f"{prefix}_merged.cif") - has_altloc_dark = (model_dark.pdb["altloc"].astype(str).str.strip() != "").any() - has_altloc_light = (model_light.pdb["altloc"].astype(str).str.strip() != "").any() + has_altloc_dark = bool((model_dark.ctx.topology.atoms.altloc != " ").any()) + has_altloc_light = bool((model_light.ctx.topology.atoms.altloc != " ").any()) if not has_altloc_dark and not has_altloc_light: import pandas as pd @@ -1691,11 +1691,11 @@ def _build_metadata(model, data, r_work, r_free): from torchref import __version__ from torchref.io.metadata import RefinementMetadata - dark_df = model_dark.pdb.copy() + dark_df = model_dark.to_dataframe() dark_df["altloc"] = "A" dark_df["occupancy"] = fractions[0] - light_df = model_light.pdb.copy() + light_df = model_light.to_dataframe() light_df["altloc"] = "B" light_df["occupancy"] = fractions[1] diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index 18b6d9a7..77f1db04 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -672,8 +672,8 @@ def run_validation(args): ) if args.verbose >= 1: - print(f" Dark model: {len(model_dark.pdb)} atoms") - print(f" Light model: {len(model_light.pdb)} atoms") + print(f" Dark model: {model_dark.n_atoms} atoms") + print(f" Light model: {model_light.n_atoms} atoms") print(f" Fraction: {args.fraction}") # R-factors (verbose only, before DED computation) diff --git a/torchref/experimental/ensemble/ensemble_amber_kl.py b/torchref/experimental/ensemble/ensemble_amber_kl.py index ebb5cfc3..f2e8e4b3 100644 --- a/torchref/experimental/ensemble/ensemble_amber_kl.py +++ b/torchref/experimental/ensemble/ensemble_amber_kl.py @@ -172,7 +172,7 @@ def _make_chem_model(self, ensemble: "EnsembleModel", verbose: int): def _member_xyz(self, i: int) -> torch.Tensor: """Member ``i`` coordinates ``(n_chem_atoms, 3)``, subset to kept atoms. - The returned ordering matches ``self._chem_model.pdb`` (what the OpenMM + The returned ordering matches ``self._chem_model.to_dataframe()`` (what the OpenMM atom map was built on), so it can be fed straight to :meth:`AmberTarget._energy`. """ diff --git a/torchref/experimental/ensemble/ensemble_model.py b/torchref/experimental/ensemble/ensemble_model.py index fed3bf93..43521193 100644 --- a/torchref/experimental/ensemble/ensemble_model.py +++ b/torchref/experimental/ensemble/ensemble_model.py @@ -283,7 +283,7 @@ def build_single_copy_model(ensemble, atom_idx=None, verbose: int = 0): Returns ------- Model - A single-conformation model exposing ``.pdb`` / ``.update_pdb()`` / + A single-conformation model exposing ``.to_dataframe()`` / ``.ctx.topology`` / ``.xyz()`` / ``.device`` over the selected atoms. """ from torchref.model.model import Model diff --git a/torchref/experimental/ensemble/ensemble_refinement.py b/torchref/experimental/ensemble/ensemble_refinement.py index a36563e7..68f26d74 100644 --- a/torchref/experimental/ensemble/ensemble_refinement.py +++ b/torchref/experimental/ensemble/ensemble_refinement.py @@ -1708,7 +1708,7 @@ def _physical_masses_for_xyz(self) -> Optional[Dict[int, torch.Tensor]]: if flat is None or flat.dim() != 2 or flat.shape[1] != 3: return None n_rows = int(flat.shape[0]) - elements = self.model.pdb["element"].astype(str).str.strip().tolist() + elements = self.model.ctx.topology.atoms.element.tolist() if len(elements) != n_rows: return None import gemmi diff --git a/torchref/experimental/targets/amber_target.py b/torchref/experimental/targets/amber_target.py index ca6383f2..8544f332 100644 --- a/torchref/experimental/targets/amber_target.py +++ b/torchref/experimental/targets/amber_target.py @@ -309,8 +309,7 @@ def _build(self) -> None: """ # Reject models with alternate conformations — OpenMM only handles # a single conformation. Call model.strip_altlocs() first. - altlocs = self._chem_model.pdb["altloc"].astype(str).str.strip() - if (altlocs != "").any(): + if (self._chem_model.ctx.topology.atoms.altloc != " ").any(): raise ValueError( "[AmberTarget] Model contains alternate conformations. " "OpenMM requires a single conformation.\n" @@ -338,7 +337,7 @@ def _build(self) -> None: self._build_context(positions_nm) self._pos_buf = positions_nm.copy() - self._n_model_atoms = len(self._chem_model.pdb) + self._n_model_atoms = self._chem_model.n_atoms if self.verbose >= 1: print( @@ -356,7 +355,7 @@ def _detect_nonstandard_residues(self) -> List[Tuple[str, int]]: Return ``(resname, net_charge)`` for HETATM residues not in :data:`AMBER14_STANDARD`. ATOM records with unknown resnames warn. """ - pdb = self._chem_model.pdb + pdb = self._chem_model.to_dataframe() nonstandard: List[Tuple[str, int]] = [] seen: set = set() @@ -441,7 +440,7 @@ def _run_antechamber_one( Cache is checked first. On a miss, work happens in a temp dir and results are atomically moved to the cache (write-then-rename). """ - pdb = self._chem_model.pdb.copy() + pdb = self._chem_model.to_dataframe() pdb[["x", "y", "z"]] = self._chem_model.xyz().detach().cpu().numpy() res_atoms = pdb[pdb["resname"].astype(str).str.strip() == resname] first = res_atoms.iloc[0] @@ -587,7 +586,7 @@ def _filter_pdb_for_tleap(self): The resulting topology must map back to every model atom, including hydrogens and terminal oxygens, before a context can be constructed. """ - pdb = self._chem_model.pdb.copy() + pdb = self._chem_model.to_dataframe() pdb[["x", "y", "z"]] = self._chem_model.xyz().detach().cpu().numpy() mask = pdb["altloc"].astype(str).str.strip().isin(["", "A"]) @@ -633,7 +632,7 @@ def _build_standard(self, cutoff_A: float, app, unit) -> Tuple: from torchref.io import pdb as pdbio self._tleap_residue_map = None - pdb = self._chem_model.pdb.copy() + pdb = self._chem_model.to_dataframe() xyz = self._chem_model.xyz().detach().cpu().numpy() pdb[["x", "y", "z"]] = xyz pdb["serial"] = np.arange(1, len(pdb) + 1) @@ -718,7 +717,7 @@ def _build_gaff2( ligand_copies = [] ligand_keys = [] - pdb = self._chem_model.pdb + pdb = self._chem_model.to_dataframe() for rn in lig_names: rows = pdb[pdb["resname"].astype(str).str.strip() == rn] for key, _ in rows.groupby(["chainid", "resseq", "icode"], sort=False): @@ -784,7 +783,7 @@ def _is_hydrogen(omm_atom) -> bool: def _build_atom_map(self) -> None: """Require a bijection between model rows and all OpenMM particles.""" - pdb = self._chem_model.pdb + pdb = self._chem_model.to_dataframe() n_model = len(pdb) atoms = list(self._topology.atoms()) mapping = np.full(n_model, -1, dtype=np.int32) @@ -831,7 +830,7 @@ def _map_gaff2_atoms(self, mapping: np.ndarray, atoms: list) -> None: """Match residue instances using heavy anchors, then names and H parents.""" from scipy.spatial import cKDTree - pdb = self._chem_model.pdb + pdb = self._chem_model.to_dataframe() keys = [ tuple(row) for row in pdb[["chainid", "resseq", "icode"]].itertuples( diff --git a/torchref/io/ihm.py b/torchref/io/ihm.py index 97a8c5f2..6f63a49d 100644 --- a/torchref/io/ihm.py +++ b/torchref/io/ihm.py @@ -881,11 +881,7 @@ def _append_atom_site( break model = mc.base_models[i] - # Update PDB DataFrame with current refined coordinates - if hasattr(model, "update_pdb"): - model.update_pdb() - - pdb_df = model.pdb + pdb_df = model.to_dataframe() for _, row in pdb_df.iterrows(): atom_name = str(row.get("name", "CA")) diff --git a/torchref/io/metadata.py b/torchref/io/metadata.py index e7ba8760..962e1503 100644 --- a/torchref/io/metadata.py +++ b/torchref/io/metadata.py @@ -250,10 +250,8 @@ def from_refinement(cls, refinement) -> RefinementMetadata: # --- B-factor statistics from model --- try: - model = refinement.model - model.update_pdb() - pdb = model.pdb - bvals = pdb["tempfactor"] + # The written B column: B_eq for anisotropic atoms, as in the file. + bvals = refinement.model.to_dataframe()["tempfactor"] meta.b_mean_overall = float(bvals.mean()) meta.b_min = float(bvals.min()) meta.b_max = float(bvals.max()) @@ -282,12 +280,10 @@ def from_refinement(cls, refinement) -> RefinementMetadata: # --- Atom counts --- try: - pdb = refinement.model.pdb - meta.n_atoms_total = len(pdb) - protein_mask = pdb["ATOM"] == "ATOM" - meta.n_atoms_protein = int(protein_mask.sum()) - solvent_mask = pdb["ATOM"] == "HETATM" - meta.n_atoms_solvent = int(solvent_mask.sum()) + is_hetatm = refinement.model.ctx.topology.atoms.is_hetatm + meta.n_atoms_total = len(is_hetatm) + meta.n_atoms_protein = int((~is_hetatm).sum()) + meta.n_atoms_solvent = int(is_hetatm.sum()) except Exception: pass diff --git a/torchref/model/context.py b/torchref/model/context.py index 0d152190..05146313 100644 --- a/torchref/model/context.py +++ b/torchref/model/context.py @@ -7,10 +7,14 @@ are fixed by the atom set and the dictionaries, and are evaluated against coordinates the caller passes in. -:meth:`ModelContext.from_atoms` is the one place an atom table is settled: hydrogens -stripped or generated, unusable rows dropped, the crystal built. Every way of making a -model -- loading a file, selecting, stripping, hydrogenating, restoring a state dict -- -produces a context first and only then installs parameter wrappers over it. +Atom identity lives on :attr:`ModelContext.topology`, a node-only +:class:`~torchref.topology.Topology`; refinable values never live here. An atom table +(a pandas DataFrame) is read only at construction: :meth:`ModelContext.from_atoms` +settles it -- unusable rows dropped, hydrogens stripped or generated, the crystal built +-- and splits it into the topology and an :class:`AtomValues` bundle of starting values +that the model's parameter wrappers are built from. Every way of making a model -- +loading a file, selecting, stripping, hydrogenating, restoring a state dict -- produces +a context and values first, and only then installs wrappers over them. Splitting it out means the crystallographic context can be passed to code that needs only that (structure-factor engines, scalers, most targets) without handing over the @@ -24,6 +28,7 @@ from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +import numpy as np import torch from torchref.utils.device_mixin import DeviceMixin @@ -32,6 +37,7 @@ import pandas from torchref.symmetry import Cell, SpaceGroup + from torchref.topology import Topology from torchref.topology.restraints import Restraints #: Three-letter residue code to one-letter code, modified residues included. @@ -107,6 +113,13 @@ def check_hydrogen_policy(hydrogens: str, hydrogen_mode: str) -> None: ) +def _copy_links(links): + """An independent copy of a reader's LINK records (a DataFrame or a sequence).""" + if links is None: + return None + return links.copy() if hasattr(links, "copy") else list(links) + + def own_spacegroup(value, dtype: torch.dtype, device) -> Optional["SpaceGroup"]: """A space group owned by the caller, on ``device`` and in ``dtype``. @@ -135,6 +148,75 @@ def own_spacegroup(value, dtype: torch.dtype, device) -> Optional["SpaceGroup"]: return SpaceGroup(value, dtype=dtype, device=device) +#: Columns of an atom table that are parameter values rather than identity. +_U_COLUMNS = ("u11", "u22", "u33", "u12", "u13", "u23") + + +@dataclass(eq=False) +class AtomValues: + """Starting values for the parameter wrappers, one row per atom. + + Read from an atom table at construction and consumed by + ``Model._install_parameters``; afterwards the wrappers are the only source of these + values. + + Parameters + ---------- + xyz : numpy.ndarray + Cartesian coordinates in Å, shape ``(N, 3)``. + b : numpy.ndarray + Isotropic B-factors in Ų, shape ``(N,)``. + u : numpy.ndarray + Anisotropic U in Ų, shape ``(N, 6)`` as ``u11 u22 u33 u12 u13 u23``; NaN for + isotropic atoms. + occupancy : numpy.ndarray + Occupancies, shape ``(N,)``. + aniso : numpy.ndarray + True for atoms carrying an ANISOU record, shape ``(N,)``. + """ + + xyz: np.ndarray + b: np.ndarray + u: np.ndarray + occupancy: np.ndarray + aniso: np.ndarray + + @classmethod + def from_table(cls, pdb: "pandas.DataFrame") -> "AtomValues": + """The value columns of an atom table; missing ANISOU columns read as NaN.""" + n = len(pdb) + u = np.full((n, 6), np.nan) + for i, column in enumerate(_U_COLUMNS): + if column in pdb.columns: + u[:, i] = pdb[column].to_numpy(dtype=np.float64) + aniso = ( + pdb["anisou_flag"].to_numpy(dtype=bool) + if "anisou_flag" in pdb.columns + else np.zeros(n, dtype=bool) + ) + return cls( + xyz=pdb[["x", "y", "z"]].to_numpy(dtype=np.float64), + b=pdb["tempfactor"].to_numpy(dtype=np.float64), + u=u, + occupancy=pdb["occupancy"].to_numpy(dtype=np.float64), + aniso=aniso, + ) + + def __len__(self) -> int: + return len(self.b) + + def gather(self, rows: np.ndarray) -> "AtomValues": + """Values for the atoms ``rows`` names, in that order; rows may repeat.""" + rows = np.asarray(rows, dtype=np.int64) + return AtomValues( + xyz=self.xyz[rows].copy(), + b=self.b[rows].copy(), + u=self.u[rows].copy(), + occupancy=self.occupancy[rows].copy(), + aniso=self.aniso[rows].copy(), + ) + + @dataclass(eq=False, repr=False) class ModelContext(DeviceMixin): """Crystallographic context, atom bookkeeping and provenance for one model. @@ -145,9 +227,10 @@ class ModelContext(DeviceMixin): Unit cell, or None before a structure is loaded. spacegroup : SpaceGroup or None Space group, or None before a structure is loaded. - pdb : pandas.DataFrame or None - The atom table. Refreshed from the model's tensors only by - ``Model.update_pdb``, so it is stale between refinement steps by design. + topology : Topology or None + Atom identity -- names, elements, altlocs, residues, chains, record types -- + as a node-only topology (no edges). Replaced, never edited, when the atom set + changes. links : list or None Link records from the reader, used to build inter-residue restraints. altloc_pairs : list @@ -155,6 +238,8 @@ class ModelContext(DeviceMixin): residue with more than one conformation; rebuilt by :meth:`register_altlocs`. input_file : str or None Path the structure was loaded from. + z_value : int or None + The CRYST1 Z of the input file, written back unchanged. cif_path : str or list of str or None Restraint dictionary path(s). Change it with :meth:`set_cif_path`, which drops restraints built over the old dictionaries. @@ -175,7 +260,8 @@ class ModelContext(DeviceMixin): initialized : bool, default False Whether a structure has been loaded. ``if model:`` tests this. restraints : Restraints or None - Geometry restraints over ``pdb``, or None until :meth:`build_restraints` runs. + Geometry restraints over ``topology``, or None until :meth:`build_restraints` + runs. Reset to None whenever the atom table or ``cif_path`` changes; read them through ``Model.restraints``, which builds on first access. @@ -192,10 +278,11 @@ class ModelContext(DeviceMixin): cell: Optional["Cell"] = None spacegroup: Optional["SpaceGroup"] = None - pdb: Optional["pandas.DataFrame"] = None + topology: Optional["Topology"] = None links: Optional[List[Any]] = None altloc_pairs: List[Any] = field(default_factory=list) input_file: Optional[str] = None + z_value: Optional[int] = None cif_path: Optional[Any] = None verbose: int = 1 hydrogens: str = "keep" @@ -234,13 +321,13 @@ def from_atoms( device, links=None, **settings, - ) -> "ModelContext": - """Settle an atom table and build the context around it. + ) -> Tuple["ModelContext", AtomValues]: + """Settle an atom table and split it into a context and starting values. - In order: strip hydrogens (``hydrogens="strip"``); drop rows without - coordinates, B-factor or occupancy and renumber the ``index`` column; build the - cell and space group; generate missing hydrogens (``hydrogens="add"``); record - the alternative conformations. + Rows without coordinates, B-factor or occupancy are dropped, the table is split + into identity (:meth:`Topology.from_table`) and :class:`AtomValues`, the cell + and space group are built, and the hydrogen policy is applied (see + :meth:`derive`). This is the only place a model's atoms are read from a table. Parameters ---------- @@ -260,10 +347,10 @@ def from_atoms( Returns ------- - ModelContext - Initialized, with no restraints built yet unless hydrogen generation - needed them (in which case they were built over the table *before* - generation and discarded). + ctx : ModelContext + Initialized, with no restraints built yet. + values : AtomValues + Starting values, row-aligned with ``ctx.topology``. Raises ------ @@ -271,73 +358,83 @@ def from_atoms( For an invalid hydrogen policy, including ``strip`` with ``riding``. """ from torchref.symmetry import Cell + from torchref.topology import Topology - ctx = cls(links=links, **settings) - if ctx.hydrogens == "strip": - pdb = pdb.loc[pdb["element"].str.strip() != "H"] - # Renumber before deriving ``index``: every consumer uses it to address length-N - # per-atom tensors positionally, so a gapped index from the drop sends them past - # the end (roughly one PDB-REDO entry in six loses rows here). + z_value = getattr(pdb, "attrs", {}).get("z") + ctx = cls(links=links, z_value=z_value, **settings) pdb = pdb.dropna(subset=["x", "y", "z", "tempfactor", "occupancy"]) - ctx.pdb = cls._renumbered(pdb) + pdb = pdb.reset_index(drop=True) ctx.cell = Cell( cell.data if isinstance(cell, Cell) else cell, dtype=dtype, device=device ) ctx.spacegroup = own_spacegroup(spacegroup, dtype, device) - if ctx.hydrogens == "add": - ctx._add_missing_hydrogens(dtype) - ctx.register_altlocs() - ctx.initialized = True - return ctx + ctx.topology = Topology.from_table(pdb) + values = ctx._settle(AtomValues.from_table(pdb), dtype) + return ctx, values + + def derive( + self, topology: "Topology", values: AtomValues, **overrides + ) -> Tuple["ModelContext", AtomValues]: + """A new context over ``topology`` in this one's crystal, with its settings. - def derive(self, pdb: "pandas.DataFrame", **overrides) -> "ModelContext": - """A new context over ``pdb`` in this one's crystal, with its settings. + The hydrogen policy is applied to the new atoms: ``"strip"`` removes every + hydrogen, ``"add"`` generates the ones the monomer templates name and the atoms + lack (waters included), ``"keep"`` leaves them as they are. Parameters ---------- - pdb : pandas.DataFrame - The new atom table; see :meth:`from_atoms`. + topology : Topology + Identity of the new atom set, node-only. + values : AtomValues + Its starting values, row-aligned with ``topology``. **overrides Settings to change, e.g. ``hydrogens="strip"``. Returns ------- - ModelContext + ctx : ModelContext + values : AtomValues + Row-aligned with ``ctx.topology``, which differs from ``topology`` when the + policy added or removed atoms. """ - settings = {**self.settings(), **overrides} - return ModelContext.from_atoms( - pdb, - self.cell, - self.spacegroup, - dtype=self.cell.dtype, - device=self.cell.device, - links=self.links, - **settings, + ctx = ModelContext( + cell=self.cell.clone() if self.cell is not None else None, + spacegroup=self.spacegroup.copy() if self.spacegroup is not None else None, + links=_copy_links(self.links), + topology=topology, + z_value=self.z_value, + **{**self.settings(), **overrides}, ) - - @staticmethod - def _renumbered(pdb: "pandas.DataFrame") -> "pandas.DataFrame": - pdb = pdb.reset_index(drop=True) - pdb["index"] = pdb.index.to_numpy(dtype=int) - return pdb - - def _add_missing_hydrogens(self, dtype: torch.dtype) -> None: - """Top up the hydrogens the atom table is missing. + return ctx, ctx._settle(values, self.cell.dtype) + + def _settle(self, values: AtomValues, dtype: torch.dtype) -> AtomValues: + """Apply the hydrogen policy to ``topology`` and ``values``; finish the context.""" + if self.hydrogens == "strip": + keep = ~self.topology.atoms.is_hydrogen.cpu().numpy() + if not keep.all(): + rows = np.nonzero(keep)[0] + self.topology = self.topology.gather(rows) + values = values.gather(rows) + if self.hydrogens == "add": + values = self._add_missing_hydrogens(values, dtype) + self.register_altlocs() + self.initialized = True + return values + + def _add_missing_hydrogens(self, values: AtomValues, dtype: torch.dtype) -> AtomValues: + """Top up the hydrogens the atoms are missing; returns the extended values. Per parent, not per file: a structure deposited with some hydrogens gets the rest, because the plan only ever proposes a hydrogen the template names and the - table does not have (1AK5 arrives with 675 of roughly 2500). + atoms do not have (1AK5 arrives with 675 of roughly 2500). A new hydrogen takes + its parent's occupancy and B-factor, and is isotropic. - Costs a restraint build without the pair list, over the table as loaded, - because the plan needs its topology; it is discarded afterwards. + Costs a restraint build without the pair list, because the plan needs the + connected topology; it is discarded afterwards. """ - from torchref.topology.hydrogens import ( - augment_atom_table, - optimise_free_torsions, - plan_hydrogens, - ) + from torchref.topology.hydrogens import optimise_free_torsions, plan_hydrogens - xyz = torch.tensor(self.pdb[["x", "y", "z"]].values, dtype=dtype) + xyz = torch.tensor(values.xyz, dtype=dtype) restraints = self.build_restraints(xyz, nonbonded=False, verbose=0) plan = plan_hydrogens( restraints.topology, restraints.cif_dict, xyz, verbose=self.verbose @@ -349,17 +446,21 @@ def _add_missing_hydrogens(self, dtype: torch.dtype) -> None: "with cif_path / --cif." ) if plan.n_hydrogens == 0: - return + return values optimise_free_torsions(plan, restraints.topology, xyz) - self.pdb = self._renumbered( - augment_atom_table(self.pdb, plan, restraints.topology) - ) + self.topology, source, _, plan_rows = self.topology.with_hydrogens(plan) + values = values.gather(source) + values.xyz[plan_rows] = np.asarray(plan.position, dtype=np.float64) + values.u[plan_rows] = np.nan + values.aniso[plan_rows] = False if self.verbose > 0: print(f"Generated {plan.n_hydrogens} hydrogens") + return values - # ------------------------------------------------------------------ - # Restraints - # ------------------------------------------------------------------ + @property + def n_atoms(self) -> int: + """Number of atoms; 0 before a structure is loaded.""" + return 0 if self.topology is None else self.topology.n_atoms def set_cif_path(self, cif_path) -> None: """Replace the restraint dictionary path and drop restraints built over the old one. @@ -375,13 +476,13 @@ def set_cif_path(self, cif_path) -> None: def build_restraints( self, xyz: torch.Tensor, *, nonbonded: bool = True, verbose=None ) -> "Restraints": - """Build restraints over the atom table and store them on :attr:`restraints`. + """Build restraints over :attr:`topology` and store them on :attr:`restraints`. Parameters ---------- xyz : torch.Tensor - Current Cartesian coordinates in Å, shape ``(n_atoms, 3)``; the atom table's - own columns are stale during refinement. The restraints land on its device. + Current Cartesian coordinates in Å, shape ``(n_atoms, 3)``, from the + model's ``xyz`` wrapper. The restraints land on its device. nonbonded : bool, default True Build the non-bonded pair list. False is for a throwaway build that needs only the topology, and is then **not** stored. @@ -394,10 +495,8 @@ def build_restraints( """ from torchref.topology.restraints import Restraints - from torchref.topology import Topology - restraints = Restraints( - topology=Topology.from_table(self.pdb), + topology=self.topology, cif_path=self.cif_path, xyz=xyz.detach(), cell=self.cell, @@ -411,12 +510,51 @@ def build_restraints( return restraints # ------------------------------------------------------------------ - # Atom-table queries + # Identity queries # ------------------------------------------------------------------ + def _residue_groups(self, with_altloc: bool) -> Dict[tuple, List[int]]: + """Atom rows grouped by ``(resname, resseq, chain[, altloc])``, keys sorted. + + The key order is the one pandas' sorted ``groupby`` gives. Occupancy groups are + numbered in it, and checkpoints store occupancies in group space, so it must + not change. + """ + columns = self.topology.columns() + altloc = np.where(columns["altloc"] == " ", "", columns["altloc"]) + keys: Dict[tuple, List[int]] = {} + for row in range(self.topology.n_atoms): + key = ( + str(columns["resname"][row]), + int(columns["resseq"][row]), + str(columns["chain"][row]), + ) + if with_altloc: + key = key + (str(altloc[row]),) + keys.setdefault(key, []).append(row) + return {key: keys[key] for key in sorted(keys)} + + def _altloc_residues(self) -> List[Tuple[tuple, List[str], Dict[str, List[int]]]]: + """Residues with more than one altloc: ``(key, sorted altlocs, rows per altloc)``. + + Keys are ``(resname, resseq, chain)``, sorted; blank-altloc atoms are not part + of any conformer. + """ + altloc = self.topology.atoms.altloc + out = [] + for key, rows in self._residue_groups(with_altloc=False).items(): + by_altloc: Dict[str, List[int]] = {} + for row in rows: + if altloc[row] != " ": + by_altloc.setdefault(str(altloc[row]), []).append(row) + if len(by_altloc) > 1: + labels = sorted(by_altloc) + out.append((key, labels, {a: by_altloc[a] for a in labels})) + return out + def occupancy_groups(self, initial_occ): """``(sharing_groups, altloc_groups, refinable_mask)`` for an - :class:`~torchref.model.parameter_wrappers.OccupancyTensor` over this table. + :class:`~torchref.model.parameter_wrappers.OccupancyTensor` over these atoms. Altloc conformations share one collapsed index each; other residues share one only when their occupancies agree to within 0.01, and an occupancy is @@ -432,46 +570,23 @@ def occupancy_groups(self, initial_occ): # First pass: altlocs. ALL atoms of one conformation must share a collapsed # index whatever their individual occupancies, or the sum-to-1 # normalization in OccupancyTensor.forward() acts on the wrong group. - pdb_with_altlocs = self.pdb[self.pdb["altloc"] != ""] altloc_residues = set() - - if len(pdb_with_altlocs) > 0: - grouped_by_residue = pdb_with_altlocs.groupby( - ["resname", "resseq", "chainid"] - ) - - for (resname, resseq, chainid), group in grouped_by_residue: - unique_altlocs = sorted(group["altloc"].unique()) - - if len(unique_altlocs) > 1: - altloc_residues.add((resname, resseq, chainid)) - conformation_atom_lists = [] - - for altloc in unique_altlocs: - altloc_atoms = group[group["altloc"] == altloc] - indices = altloc_atoms["index"].tolist() - - sharing_groups_tensor[indices] = collapsed_idx - - for idx in indices: - if abs(initial_occ[idx].item() - 1.0) > 0.01: - refinable_mask[idx] = True - - conformation_atom_lists.append(indices) - collapsed_idx += 1 - - altloc_groups.append(tuple(conformation_atom_lists)) + for key, labels, rows_by_altloc in self._altloc_residues(): + altloc_residues.add(key) + conformation_atom_lists = [] + for label in labels: + indices = rows_by_altloc[label] + sharing_groups_tensor[indices] = collapsed_idx + for idx in indices: + if abs(initial_occ[idx].item() - 1.0) > 0.01: + refinable_mask[idx] = True + conformation_atom_lists.append(indices) + collapsed_idx += 1 + altloc_groups.append(tuple(conformation_atom_lists)) # Second pass: non-altloc residues, sharing by occupancy similarity. - grouped = self.pdb.groupby(["resname", "resseq", "chainid", "altloc"]) - - for (resname, resseq, chainid, altloc), group in grouped: - if (resname, resseq, chainid) in altloc_residues: - continue - - indices = group["index"].tolist() - - if len(indices) == 0: + for key, indices in self._residue_groups(with_altloc=True).items(): + if key[:3] in altloc_residues: continue residue_occs = initial_occ[indices] @@ -518,37 +633,20 @@ def occupancy_groups(self, initial_occ): return sharing_groups_tensor, altloc_groups, refinable_mask def register_altlocs(self) -> None: - """ - Rebuild ``self.altloc_pairs`` from the ``altloc`` column. + """Rebuild :attr:`altloc_pairs` from the topology's altlocs. - One tuple per residue that has multiple conformations, holding one - index tensor per conformation (in sorted altloc order), e.g. - ``[(tensor([100, 101]), tensor([110, 111])), ...]``. Overwrites any - previous content, so call it after the atom numbering changes. + One tuple per residue that has multiple conformations, holding one index tensor + per conformation (in sorted altloc order), e.g. + ``[(tensor([100, 101]), tensor([110, 111])), ...]``. Overwrites any previous + content, so call it after the atom numbering changes. """ - self.altloc_pairs = [] - - pdb_with_altlocs = self.pdb[self.pdb["altloc"] != ""] - - if len(pdb_with_altlocs) == 0: - return - - grouped = pdb_with_altlocs.groupby(["resname", "resseq", "chainid"]) - - for (resname, resseq, chainid), group in grouped: - unique_altlocs = sorted(group["altloc"].unique()) - - # A lone altloc label is not an alternative conformation. - if len(unique_altlocs) > 1: - conformation_tensors = [] - for altloc in unique_altlocs: - altloc_atoms = group[group["altloc"] == altloc] - indices = torch.tensor( - altloc_atoms["index"].tolist(), dtype=torch.long # dtype-ok: altloc atom indices; indexing requires long - ) - conformation_tensors.append(indices) - - self.altloc_pairs.append(tuple(conformation_tensors)) + self.altloc_pairs = [ + tuple( + torch.tensor(rows_by_altloc[label], dtype=torch.long) # dtype-ok: altloc atom indices; indexing requires long + for label in labels + ) + for _, labels, rows_by_altloc in self._altloc_residues() + ] @property def chain_sequences(self) -> List[Tuple[str, str]]: @@ -557,71 +655,56 @@ def chain_sequences(self) -> List[Tuple[str, str]]: HETATM records are excluded, numbering gaps become ``?`` and unrecognized residues ``X``. """ - if self.pdb is None: - return [] - - atom_df = self.pdb[self.pdb["ATOM"] == "ATOM"] result = [] - - for chain in atom_df["chainid"].unique(): - chain_df = atom_df[atom_df["chainid"] == chain] - residues = chain_df.drop_duplicates(subset=["resseq", "icode"]).sort_values( - "resseq" - ) - resseqs = residues["resseq"].values - resnames = residues["resname"].values - + for chain, residues in self._polymer_residues(): seq_chars = [] - for i, (rseq, rname) in enumerate(zip(resseqs, resnames)): + for i, (resseq, resname) in enumerate(residues): if i > 0: - gap = int(rseq) - int(resseqs[i - 1]) - 1 + gap = resseq - residues[i - 1][0] - 1 if gap > 0: seq_chars.extend(["?"] * gap) - code = THREE_TO_ONE.get(str(rname).strip(), "X") - seq_chars.append(code) - - result.append((str(chain), "".join(seq_chars))) - + seq_chars.append(THREE_TO_ONE.get(resname, "X")) + result.append((chain, "".join(seq_chars))) return result - @property - def chain_residues(self) -> List[Tuple[str, List[str]]]: - """ - Per-chain residue names as 3-letter codes (for IHM/CIF writing). - - Excludes HETATM records. Unlike :attr:`chain_sequences`, returns - the raw 3-letter codes without gap filling. + def _polymer_residues(self) -> List[Tuple[str, List[Tuple[int, str]]]]: + """``(chain, [(resseq, resname), ...])`` over ATOM records, chains in file order. - Returns - ------- - list of (str, list of str) - Ordered list of ``(chain_id, [resname, ...])``. + One entry per ``(resseq, icode)``, sorted by ``resseq`` (stably, so insertion + codes keep their file order). """ - if self.pdb is None: + if self.topology is None: return [] + residues = self.topology.residues + first = residues.atom_start.astype(np.int64) + polymer = ~self.topology.atoms.is_hetatm[first] if len(first) else [] + chains: Dict[str, Dict[tuple, Tuple[int, str]]] = {} + for r in np.nonzero(polymer)[0]: + chain, resseq, icode = residues.key(int(r)) + seen = chains.setdefault(chain, {}) + seen.setdefault((resseq, icode), (resseq, str(residues.resname[r]))) + return [ + (chain, sorted(seen.values(), key=lambda item: item[0])) + for chain, seen in chains.items() + ] - atom_df = self.pdb[self.pdb["ATOM"] == "ATOM"] - result = [] - - for chain in atom_df["chainid"].unique(): - chain_df = atom_df[atom_df["chainid"] == chain] - residues = chain_df.drop_duplicates(subset=["resseq", "icode"]).sort_values( - "resseq" - ) - resnames = [str(r).strip() for r in residues["resname"].values] - result.append((str(chain), resnames)) - - return result + @property + def chain_residues(self) -> List[Tuple[str, List[str]]]: + """Per-chain residue names as 3-letter codes, ``[(chain_id, [resname, ...])]``. - # ------------------------------------------------------------------ - # Copying and persistence - # ------------------------------------------------------------------ + Excludes HETATM records. Unlike :attr:`chain_sequences`, the raw 3-letter codes + without gap filling; used by the IHM and mmCIF writers. + """ + return [ + (chain, [resname for _, resname in residues]) + for chain, residues in self._polymer_residues() + ] def copy(self) -> "ModelContext": """An independent copy. - The atom table is deep-copied, the cell and space group are cloned and built - restraints are copied, so nothing is shared with the original. Cloning the + The topology is copied, the cell and space group are cloned and built restraints + are copied, so nothing is shared with the original. Cloning the space group matters because it is a mutable dataclass: sharing the reference would let an edit through one model's context reach every model copied from it. @@ -635,12 +718,13 @@ def copy(self) -> "ModelContext": spacegroup=( self.spacegroup.copy() if self.spacegroup is not None else None ), - pdb=self.pdb.copy(deep=True) if self.pdb is not None else None, - links=list(self.links) if self.links is not None else None, + topology=self.topology.copy() if self.topology is not None else None, + links=_copy_links(self.links), altloc_pairs=[ tuple(t.clone() for t in group) for group in self.altloc_pairs ], initialized=self.initialized, + z_value=self.z_value, **self.settings(), ) if self.restraints is not None: @@ -653,17 +737,17 @@ def copy(self) -> "ModelContext": return duplicate def state(self) -> Dict[str, Any]: - """What :meth:`from_state` needs, as picklable entries for a model state dict. + """What :meth:`from_state` needs besides the atom table, as picklable entries. Returns ------- dict - The atom table, the cell as a CPU tensor, the space group as its extended - Hermann-Mauguin symbol (``gemmi.SpaceGroup`` is not picklable), the - altloc groups and the settings. Restraints are not saved; they rebuild. + The cell as a CPU tensor, the space group as its extended Hermann-Mauguin + symbol (``gemmi.SpaceGroup`` is not picklable), the altloc groups and the + settings. The atom table itself is written by the model, which alone has + the current values; restraints are not saved, they rebuild. """ return { - "pdb": self.pdb.copy() if self.pdb is not None else None, "cell": self.cell.data.cpu() if self.cell is not None else None, "spacegroup": self.spacegroup.xhm if self.spacegroup else None, "initialized": self.initialized, @@ -677,12 +761,13 @@ def state(self) -> Dict[str, Any]: @classmethod def from_state( cls, state: Dict[str, Any], *, dtype: torch.dtype, device, verbose: int = 1 - ) -> "ModelContext": - """Rebuild a context from the entries :meth:`state` wrote. + ) -> Tuple["ModelContext", Optional[AtomValues]]: + """Rebuild a context, and the saved values, from a model state dict. - Consumes them: every key read is popped off ``state``, so what remains is for - ``load_state_dict``. The atom table is taken as saved -- the hydrogen policy is - recorded, not re-applied. Checkpoints that predate the policy are mapped: + Consumes the entries: every key read is popped off ``state``, so what remains is + for ``load_state_dict``. The saved atom table (``"pdb"``) is split as at + construction, but taken as saved -- the hydrogen policy is recorded, not + re-applied. Checkpoints that predate the policy are mapped: ``strip_H`` becomes ``hydrogens``, ``"free"`` becomes ``"atoms"``, and a saved riding wrapper (an ``xyz.h_row`` entry) implies ``"riding"``. @@ -697,9 +782,12 @@ def from_state( Returns ------- - ModelContext + ctx : ModelContext + values : AtomValues or None + None when the state holds no atoms. """ from torchref.symmetry import Cell + from torchref.topology import Topology hydrogens = state.pop("hydrogens", None) strip_h = state.pop("strip_H", True) @@ -713,8 +801,10 @@ def from_state( hydrogens = "keep" cell = state.pop("cell", None) + table = state.pop("pdb", None) + z_value = None if table is None else getattr(table, "attrs", {}).get("z") ctx = cls( - pdb=state.pop("pdb", None), + topology=None if table is None else Topology.from_table(table), cell=Cell(cell, dtype=dtype, device=device) if cell is not None else None, spacegroup=own_spacegroup(state.pop("spacegroup", None), dtype, device), initialized=state.pop("initialized", False), @@ -724,8 +814,9 @@ def from_state( hydrogen_mode=mode, hydrogens_in_xray=state.pop("hydrogens_in_xray", True), verbose=verbose, + z_value=z_value, ) - return ctx + return ctx, None if table is None else AtomValues.from_table(table) @property def crystal_key(self): @@ -742,7 +833,7 @@ def crystal_key(self): return (self.cell.key, self.spacegroup.key) def __repr__(self) -> str: - n_atoms = 0 if self.pdb is None else len(self.pdb) + n_atoms = self.n_atoms sg = None if self.spacegroup is None else self.spacegroup.name return ( f"ModelContext(spacegroup={sg!r}, n_atoms={n_atoms}, " @@ -751,4 +842,10 @@ def __repr__(self) -> str: ) -__all__ = ["ModelContext", "HYDROGEN_SOURCES", "HYDROGEN_MODES", "check_hydrogen_policy"] +__all__ = [ + "ModelContext", + "AtomValues", + "HYDROGEN_SOURCES", + "HYDROGEN_MODES", + "check_hydrogen_policy", +] diff --git a/torchref/model/model.py b/torchref/model/model.py index 6d5f082c..4e7b804b 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -12,9 +12,11 @@ - f_calc/f_obs: Complex structure factors (lowercase = complex) """ +import warnings from typing import Dict, Iterable, List, Optional, Tuple, Union import gemmi +import numpy as np import torch import torch.nn as nn @@ -26,7 +28,7 @@ normalize_device, ) from torchref.io import cif, pdb -from torchref.model.context import ModelContext, own_spacegroup +from torchref.model.context import AtomValues, ModelContext, own_spacegroup from torchref.model.parameter_wrappers import ( CholeskyMixedTensor, MixedTensor, @@ -83,12 +85,11 @@ class Model(DeviceMovementMixin, DebugMixin, nn.Module): occupancy : OccupancyTensor Atomic occupancies with values in [0, 1]. ctx : ModelContext - The unit cell, space group, atom table, link records, provenance and - configuration. The fields not forwarded below are reached through it, e.g. - ``model.ctx.hydrogens`` and ``model.ctx.initialized``. - pdb : pandas.DataFrame - Atom table, forwarded to :attr:`ctx`. Only refreshed from the tensors by - :meth:`update_pdb`. + The unit cell, space group, atom identity (``ctx.topology``), link records, + provenance and configuration. The fields not forwarded below are reached + through it, e.g. ``model.ctx.hydrogens`` and ``model.ctx.initialized``. + n_atoms : int + Number of atoms. cell : Cell Unit cell, forwarded to :attr:`ctx`. spacegroup : SpaceGroup @@ -174,16 +175,12 @@ def _sf_atom_mask(self) -> Optional[torch.Tensor]: buffer, which is dropped with the other per-atom caches when the atom set changes, so it never outlives the table it was built for. """ - if self.ctx.hydrogens_in_xray or self.pdb is None: + if self.ctx.hydrogens_in_xray or self.ctx.topology is None: return None if getattr(self, "_heavy_atom_mask", None) is None: self.register_buffer( "_heavy_atom_mask", - torch.tensor( - (self.pdb["element"].str.strip().str.upper() != "H").values, - dtype=torch.bool, - device=self.device, - ), + ~self.ctx.topology.atoms.is_hydrogen.to(self.device), ) return self._heavy_atom_mask @@ -210,7 +207,7 @@ def _sf_partition(self): (flag.data_ptr(), flag._version) if flag is not None else None, bool(self.ctx.hydrogens_in_xray), None if heavy is None else (heavy.data_ptr(), heavy._version), - 0 if self.pdb is None else len(self.pdb), + self.n_atoms, ) cached = getattr(self, "_sf_partition_cache", None) if cached is not None and self._sf_partition_fp == fp: @@ -259,13 +256,26 @@ def _aniso_is_empty(self) -> bool: # ========================================================================= @property - def pdb(self) -> Optional["pandas.DataFrame"]: - """Atom table. Only refreshed from the tensors by :meth:`update_pdb`.""" - return self.ctx.pdb + def n_atoms(self) -> int: + """Number of atoms; 0 before a structure is loaded.""" + return self.ctx.n_atoms - @pdb.setter - def pdb(self, value): - self.ctx.pdb = value + @property + def pdb(self) -> Optional["pandas.DataFrame"]: + """The atom table, freshly joined from identity and current values. + + .. deprecated:: + Use :meth:`to_dataframe` for output and ``model.ctx.topology`` for atom + identity. Each access builds a new table, so writing into it changes + nothing. + """ + warnings.warn( + "Model.pdb is deprecated: use model.to_dataframe() for a table and " + "model.ctx.topology for atom identity", + DeprecationWarning, + stacklevel=2, + ) + return self.to_dataframe() if self.ctx.topology is not None else None @property def cell(self) -> Optional[Cell]: @@ -328,7 +338,7 @@ def _build_z_tensor(self) -> torch.Tensor: if hasattr(self, "_Z") and self._Z is not None: return self._Z - if not self.ctx.initialized or self.pdb is None: + if not self.ctx.initialized: raise RuntimeError( "Cannot build Z tensor: model not initialized. " "Load data first with load_pdb() or load_cif()." @@ -339,7 +349,7 @@ def _build_z_tensor(self) -> torch.Tensor: element_to_z = get_element_to_z_mapping() z_values = [ element_to_z.get(elem.strip().capitalize(), 0) - for elem in self.pdb["element"] + for elem in self.ctx.topology.atoms.element ] self.register_buffer( "_Z", torch.tensor(z_values, dtype=torch.int32, device=self.device) # dtype-ok: atomic-number Z categorical codes buffer; fixed int32 lookup keys @@ -358,7 +368,7 @@ def _build_parametrization(self): if self._parametrization is not None: return self._parametrization - if not self.ctx.initialized or self.pdb is None: + if not self.ctx.initialized: raise RuntimeError( "Cannot build parametrization: model not initialized. " "Load data first with load_pdb() or load_cif()." @@ -378,7 +388,7 @@ def _build_parametrization(self): self.register_buffer("_B", B) # Legacy per-element view: one representative row per element. - elements = self.pdb.element.tolist() + elements = self.ctx.topology.atoms.element.tolist() unique_elements = list(set(elements)) self._parametrization = {} @@ -496,9 +506,10 @@ def load(self, reader): Populate the model from a reader callable. The central loader that ``load_pdb`` / ``load_cif`` funnel through. The context - is built by :meth:`ModelContext.from_atoms`, which applies the hydrogen policy, - drops rows without coordinates, B-factor or occupancy, and builds the cell and - space group; the parameter wrappers are then installed over it. + and starting values are built by :meth:`ModelContext.from_atoms` -- which drops + rows without coordinates, B-factor or occupancy, applies the hydrogen policy and + builds the cell and space group -- and the parameter wrappers are installed over + them. The table is not kept. Parameters ---------- @@ -513,7 +524,7 @@ def load(self, reader): """ pdb, cell, spacegroup = reader() self._invalidate_atom_derived_caches() - self.ctx = ModelContext.from_atoms( + self.ctx, values = ModelContext.from_atoms( pdb, cell, spacegroup, @@ -522,42 +533,45 @@ def load(self, reader): links=getattr(reader, "links", None), **self.ctx.settings(), ) - self._install_parameters() + self._install_parameters(values) return self - def _install_parameters(self, state: Optional[dict] = None, xyz=None) -> None: - """Build the parameter wrappers and per-atom buffers over ``ctx.pdb``. + def _install_parameters( + self, values: AtomValues, state: Optional[dict] = None, xyz=None + ) -> None: + """Build the parameter wrappers and per-atom buffers over ``ctx.topology``. The only place the wrappers are constructed: load, restore, select and every - derived model come through here. Values are read off the atom table. + derived model come through here. Parameters ---------- + values : AtomValues + Starting values, row-aligned with ``ctx.topology``. state : dict, optional A state dict about to be loaded. Its saved refinable masks, riding frames and node-field ADP layout fix the wrappers' shapes, and the default masks are **not** applied; ``load_state_dict`` supplies the values afterwards. xyz : MixedTensor, optional - Coordinate wrapper to install as is, instead of building one from the table. + Coordinate wrapper to install as is, instead of building one from + ``values``. """ from torchref.model.riding_xyz import RidingXYZTensor - pdb, dtype = self.pdb, self.dtype_float + dtype = self.dtype_float restoring = state is not None state = {} if state is None else state self.register_buffer( "aniso_flag", - torch.tensor( - pdb["anisou_flag"].values, dtype=torch.bool, device=self.device - ), + torch.as_tensor(values.aniso, dtype=torch.bool, device=self.device), ) - self.xyz = self._build_xyz(state) if xyz is None else xyz - self.adp = self._restore_adp_slot("adp", state, pdb, dtype, self.xyz, self.device) - self.u = self._restore_adp_slot("u", state, pdb, dtype, self.xyz, self.device) + self.xyz = self._build_xyz(values, state) if xyz is None else xyz + self.adp = self._restore_adp_slot("adp", state, values, dtype, self.xyz, self.device) + self.u = self._restore_adp_slot("u", state, values, dtype, self.xyz, self.device) # Residue-level sharing plus altloc sum-to-1 groups. - initial_occ = torch.tensor(pdb["occupancy"].values, dtype=dtype) + initial_occ = torch.tensor(values.occupancy, dtype=dtype) sharing_groups, altloc_groups, refinable_mask = self.ctx.occupancy_groups( initial_occ ) @@ -580,7 +594,8 @@ def _install_parameters(self, state: Optional[dict] = None, xyz=None) -> None: # the defaults here would resize the refinable sets it has to match. for mask_name in ("xyz_mask", "adp_mask", "u_mask", "occupancy_mask"): self.register_buffer( - mask_name, torch.ones(len(pdb), dtype=torch.bool, device=self.device) + mask_name, + torch.ones(self.n_atoms, dtype=torch.bool, device=self.device), ) if state.get("vdw_radii") is not None: self.register_buffer( @@ -595,16 +610,16 @@ def _install_parameters(self, state: Optional[dict] = None, xyz=None) -> None: self.xyz = RidingXYZTensor.from_mixed_tensor(self.xyz, self.hydrogen_frames()) self._repoint_coordinate_accessors() - def _build_xyz(self, state: dict): - """The coordinate wrapper over the atom table, riding if ``state`` saved one. + def _build_xyz(self, values: AtomValues, state: dict): + """The coordinate wrapper over ``values``, riding if ``state`` saved one. A saved riding wrapper is recognised by its frame buffers, never by shape: its storage is ``(n_base, 3)`` and a plain wrapper's ``(n_atoms, 3)``, both 2-D. """ - values = torch.tensor(self.pdb[["x", "y", "z"]].values, dtype=self.dtype_float) + coords = torch.tensor(values.xyz, dtype=self.dtype_float) mask = state.get("xyz.refinable_mask") if state.get("xyz.h_row") is None: - return MixedTensor(values, refinable_mask=mask, name="xyz", device=self.device) + return MixedTensor(coords, refinable_mask=mask, name="xyz", device=self.device) from torchref.model.riding_xyz import RidingXYZTensor from torchref.topology.hydrogens import HydrogenFrames @@ -619,7 +634,7 @@ def _build_xyz(self, state: dict): state.get("xyz.rotation_group"), ) return RidingXYZTensor( - values, + coords, frames, refinable_mask=mask, mask_in_base_space=True, @@ -668,49 +683,92 @@ def load_cif(self, file): return self.load(cif_reader) - def update_pdb(self): - """ - Write the current refinable parameters back into ``self.pdb``. + def to_dataframe(self) -> "pandas.DataFrame": + """The atom table: identity from the topology, values from the wrappers. - Copies the live values of ``xyz`` (x/y/z), ``u`` (u11..u23) and - ``occupancy`` from the parameter wrappers into the corresponding columns of - the ``self.pdb`` DataFrame. Called by every writer and by ``hydrogenate`` - before output. + Built fresh on every call and never kept, so it cannot go stale and editing it + changes nothing. Columns are the readers' (``ATOM``, ``serial``, ``name``, + ``altloc``, ``resname``, ``chainid``, ``resseq``, ``icode``, ``x``/``y``/``z``, + ``occupancy``, ``tempfactor``, ``element``, ``charge``, ``anisou_flag``, + ``u11`` ... ``u23``, ``index``); ``serial`` runs 1..N and a blank altloc is + ``''``. ``attrs`` carries the cell and the space group. - ``tempfactor`` is the equivalent isotropic B whenever any atom is - anisotropic, so the column agrees with the ANISOU records written beside it; - with no anisotropic atoms it is the isotropic wrapper directly. + ``tempfactor`` is the equivalent isotropic B whenever any atom is anisotropic, + so the column agrees with the ANISOU records written beside it: for an + anisotropic atom the PDB convention is B_eq = (8 pi^2 / 3) tr(U), not whatever + the isotropic wrapper still holds, which stops being refined the moment an atom + goes anisotropic. Returns ------- pandas.DataFrame - The updated ``self.pdb`` DataFrame. - - Notes - ----- - This does **not** touch ``anisou_flag``: the iso/aniso classification - of each atom is left unchanged (see ``_apply_adp_partition`` for the - partition logic that owns that flag). """ - self.pdb.loc[:, ["x", "y", "z"]] = self.xyz().cpu().detach().numpy() - self.pdb.loc[:, ["u11", "u22", "u33", "u12", "u13", "u23"]] = ( - self.u().cpu().detach().numpy() - ) - # The B column must agree with the ANISOU records beside it: for an - # anisotropic atom the PDB convention is B_eq = (8 pi^2 / 3) tr(U), not - # whatever the isotropic wrapper still happens to hold. That wrapper stops - # being refined the moment an atom goes anisotropic, so writing it directly - # emits a stale B alongside a live U. + import pandas as pd + + identity = self.ctx.topology.columns() + # Detached, not under no_grad: the wrappers cache their forward, and a + # gradient-free result cached here would be served to the next loss. + xyz = self.xyz().detach().cpu().numpy() + u = self.u().detach().cpu().numpy() if getattr(self, "_aniso_is_empty", True): - self.pdb.loc[:, "tempfactor"] = self.adp().cpu().detach().numpy() + b = self.adp().detach().cpu().numpy() else: from torchref.base.targets.adp import u6_b_eq - self.pdb.loc[:, "tempfactor"] = ( - u6_b_eq(self.adp_u6()).cpu().detach().numpy() - ) - self.pdb.loc[:, "occupancy"] = self.occupancy().cpu().detach().numpy() - return self.pdb + b = u6_b_eq(self.adp_u6()).detach().cpu().numpy() + occupancy = self.occupancy().detach().cpu().numpy() + n = self.n_atoms + table = pd.DataFrame( + { + "ATOM": np.where(identity["is_hetatm"], "HETATM", "ATOM"), + "serial": np.arange(1, n + 1), + "name": identity["name"], + "altloc": np.where(identity["altloc"] == " ", "", identity["altloc"]), + "resname": identity["resname"], + "chainid": identity["chain"], + "resseq": identity["resseq"], + "icode": identity["icode"], + "x": xyz[:, 0], + "y": xyz[:, 1], + "z": xyz[:, 2], + "occupancy": occupancy, + "tempfactor": b, + "element": identity["element"], + "charge": identity["charge"], + "anisou_flag": self.aniso_flag.detach().cpu().numpy(), + **{ + column: u[:, i] + for i, column in enumerate( + ("u11", "u22", "u33", "u12", "u13", "u23") + ) + }, + "index": np.arange(n), + } + ) + if self.cell is not None: + # The shortest decimal that round-trips the stored precision, so a float32 + # cell writes as 143.11 rather than 143.110001. + cell = self.cell.data.detach().cpu().numpy() + table.attrs["cell"] = [float(str(v)) for v in cell] + table.attrs["spacegroup"] = self.spacegroup.hm if self.spacegroup else "P 1" + if self.ctx.z_value is not None: + table.attrs["z"] = self.ctx.z_value + return table + + def _current_values(self) -> AtomValues: + """The wrappers' current values, as starting values for a derived model. + + ``b`` is the isotropic wrapper's value, not the anisotropic B_eq that + :meth:`to_dataframe` writes. + """ + # Detached rather than under no_grad; see to_dataframe. + return AtomValues( + xyz=self.xyz().detach().cpu().numpy().astype(np.float64), + b=self.adp().detach().cpu().numpy().astype(np.float64), + u=self.u().detach().cpu().numpy().astype(np.float64), + occupancy=self.occupancy().detach().cpu().numpy().astype(np.float64), + aniso=self.aniso_flag.detach().cpu().numpy().astype(bool), + ) def get_vdw_radii(self): """ @@ -727,14 +785,11 @@ def get_vdw_radii(self): if hasattr(self, "vdw_radii"): return self.vdw_radii - vdw_radii = vdw_radii_for_elements(self.pdb["element"]) + vdw_radii = vdw_radii_for_elements(self.ctx.topology.atoms.element) self.register_buffer( "vdw_radii", torch.tensor(vdw_radii, dtype=self.dtype_float, device=self.device), ) - assert len(self.vdw_radii) == len( - self.pdb - ), f"vdW radii length mismatch with number of atoms {len(self.vdw_radii)} != {len(self.pdb)}" return self.vdw_radii def _after_device_apply( @@ -791,7 +846,7 @@ def copy(self): duplicate.reset_cache() if self.ctx.verbose > 0: - print(f"Copied {type(self).__name__} ({len(duplicate.pdb)} atoms)") + print(f"Copied {type(self).__name__} ({duplicate.n_atoms} atoms)") return duplicate def _spawn(self, ctx: ModelContext) -> "Model": @@ -818,18 +873,40 @@ def _subclass_kwargs(self) -> dict: return {} def _derive(self, pdb, **overrides) -> "Model": - """A new, quiet model of this class over ``pdb`` in this crystal. + """A new, quiet model of this class built from an atom table in this crystal. + + Construction, so the table is read: see :meth:`ModelContext.from_atoms`. Parameters ---------- pdb : pandas.DataFrame - Atom table for the new model; settled by - :meth:`~torchref.model.context.ModelContext.derive`. + Atom table for the new model. **overrides Context settings to change, e.g. ``hydrogens="strip"``. """ - model = self._spawn(self.ctx.derive(pdb, **{"verbose": 0, **overrides})) - model._install_parameters() + from torchref.topology import Topology + + values = AtomValues.from_table(pdb.reset_index(drop=True)) + return self._derive_from(Topology.from_table(pdb), values, **overrides) + + def _derive_from(self, topology, values: AtomValues, xyz=None, **overrides) -> "Model": + """A new model of this class over ``topology`` and ``values`` in this crystal. + + Parameters + ---------- + topology : Topology + Node-only identity of the new atoms. + values : AtomValues + Their starting values. + xyz : MixedTensor, optional + A coordinate wrapper to install instead of one built from ``values``; only + valid when the hydrogen policy leaves the atom set unchanged. + **overrides + Context settings to change; ``verbose`` defaults to 0. + """ + ctx, values = self.ctx.derive(topology, values, **{"verbose": 0, **overrides}) + model = self._spawn(ctx) + model._install_parameters(values, xyz=xyz) return model def _kept_hydrogens(self) -> str: @@ -846,10 +923,9 @@ def write_pdb(self, filename, metadata=None): metadata : RefinementMetadata, optional Metadata to render as PDB header (REMARK 3, TITLE, etc.). """ - self.update_pdb() - self.pdb = sanitize_pdb_dataframe(self.pdb) - self.pdb.attrs["spacegroup"] = self.spacegroup.hm if self.spacegroup else "P 1" - pdb.write(self.pdb, filename, metadata=metadata) + table = sanitize_pdb_dataframe(self.to_dataframe()) + table.attrs["spacegroup"] = self.spacegroup.hm if self.spacegroup else "P 1" + pdb.write(table, filename, metadata=metadata) def write_cif(self, filename, metadata=None): """Write model to mmCIF file with optional metadata. @@ -861,10 +937,9 @@ def write_cif(self, filename, metadata=None): metadata : RefinementMetadata, optional Metadata to include (refinement statistics, title, etc.). """ - self.update_pdb() - self.pdb = sanitize_pdb_dataframe(self.pdb) - self.pdb.attrs["spacegroup"] = self.spacegroup.hm if self.spacegroup else "P 1" - cif.write_model(self.pdb, filename, metadata=metadata) + table = sanitize_pdb_dataframe(self.to_dataframe()) + table.attrs["spacegroup"] = self.spacegroup.hm if self.spacegroup else "P 1" + cif.write_model(table, filename, metadata=metadata) def get_iso(self): """ @@ -908,7 +983,7 @@ def set_default_masks(self): Called from :meth:`_install_parameters` after the wrappers are constructed. """ self.register_buffer( - "xyz_mask", torch.ones(len(self.pdb), dtype=torch.bool, device=self.device) + "xyz_mask", torch.ones(self.n_atoms, dtype=torch.bool, device=self.device) ) self.xyz.update_refinable_mask(self.xyz_mask) self.register_buffer("adp_mask", ~self.adp().detach().isnan()) @@ -1080,7 +1155,7 @@ def set_adp_mode( which a field evaluates per atom, so the field materialises into a per-atom wrapper on the way out. """ - if not self.ctx.initialized or self.pdb is None: + if not self.ctx.initialized: return if mode == "preserve": # Leave the ADPs exactly as loaded. Constructing a Refinement otherwise @@ -1103,18 +1178,15 @@ def set_adp_mode( # asked for. if aniso_selection is None: target_mask = torch.ones( - len(self.pdb), dtype=torch.bool, device=self.device + self.n_atoms, dtype=torch.bool, device=self.device ) else: - from torchref.utils.utils import create_selection_mask - - target_mask = torch.as_tensor( - create_selection_mask(aniso_selection, self.pdb), - dtype=torch.bool, - ).to(self.device) + target_mask = self.get_selection_mask(aniso_selection).to( + self.device + ) else: target_mask = torch.zeros( - len(self.pdb), dtype=torch.bool, device=self.device + self.n_atoms, dtype=torch.bool, device=self.device ) self._apply_adp_partition(target_mask) self._install_disorder_field( @@ -1128,15 +1200,11 @@ def set_adp_mode( return if mode == "isotropic": aniso_mask = torch.zeros( - len(self.pdb), dtype=torch.bool, device=self.device + self.n_atoms, dtype=torch.bool, device=self.device ) elif mode == "anisotropic": - from torchref.utils.utils import create_selection_mask - sel = aniso_selection or "not resname HOH and not element H" - aniso_mask = torch.as_tensor( - create_selection_mask(sel, self.pdb), dtype=torch.bool - ).to(self.device) + aniso_mask = self.get_selection_mask(sel).to(self.device) else: raise ValueError( f"Unknown ADP mode: {mode!r}. Use 'isotropic', 'anisotropic', " @@ -1244,12 +1312,12 @@ def _install_disorder_field( ) B = target if n_nodes is None: - n_nodes = max(4, int(round(len(self.pdb) / 25.0))) + n_nodes = max(4, int(round(self.n_atoms / 25.0))) # Anchor on density clusters, not single atoms: a node placed exactly on an atom # can isolate that atom by narrowing its kernel, which is per-atom refinement # wearing a node's clothes. - anchor_rows = density_anchor_rows(xyz, min(n_nodes, len(self.pdb))) + anchor_rows = density_anchor_rows(xyz, min(n_nodes, self.n_atoms)) if mode_set is not None: payload = ModeCovariancePayload(mode_set) @@ -1280,7 +1348,7 @@ def _install_disorder_field( if self.ctx.verbose > 0: kind = mode_set if mode_set else ("aniso U" if anisotropic else "iso B") - was = len(self.pdb) * (6 if anisotropic else 1) + was = self.n_atoms * (6 if anisotropic else 1) print( f"ADP field ({kind}): {field.n_nodes} nodes, k={k_neighbors}, " f"{int(field.get_refinable_count())} refinable nodes, " @@ -1294,7 +1362,8 @@ def _apply_adp_partition(self, aniso_mask: torch.Tensor): """Convert ADP storage to match a target anisotropic-atom mask. The body of :meth:`set_adp_mode`: rebuilds both wrappers and refreshes - ``aniso_flag``, the SF index cache, the masks, ``anisou_flag`` and caches. + ``aniso_flag`` (which the writers' ``anisou_flag`` column comes from), the SF + index cache, the masks and caches. """ import math @@ -1337,11 +1406,6 @@ def _apply_adp_partition(self, aniso_mask: torch.Tensor): self.adp.update_refinable_mask(self.adp_mask) self.u.update_refinable_mask(self.u_mask) - # The PDB/mmCIF writers gate ANISOU on this column; update_pdb() does not - # touch it, so keep it in sync with the chosen parametrization. - if self.pdb is not None: - self.pdb["anisou_flag"] = aniso_mask.detach().cpu().numpy() - # Anisotropy change invalidates structure-factor + wrapper forward caches. if hasattr(self, "reset_cache"): self.reset_cache() @@ -1359,7 +1423,7 @@ def update_mask_from_selection( Parameters ---------- selection_string : str - Phenix-style selection string (see parse_phenix_selection docs). + Phenix-style selection; grammar in :mod:`torchref.utils.selection`. target : str Parameter to update: 'xyz', 'adp', 'u', or 'occupancy'. mode : str, optional @@ -1381,8 +1445,6 @@ def update_mask_from_selection( model.update_mask_from_selection("chain A", "xyz", freeze=True) model.apply_mask_to_parameter("xyz") """ - from torchref.utils.utils import create_selection_mask - mask_map = { "xyz": "xyz_mask", "adp": "adp_mask", @@ -1398,12 +1460,15 @@ def update_mask_from_selection( mask_name = mask_map[target] current_mask = getattr(self, mask_name) - selection_mask = create_selection_mask( - selection_string, - self.pdb, - current_mask=current_mask if mode != "set" else None, - mode=mode, - ) + selected = self.get_selection_mask(selection_string).to(current_mask.device) + if mode == "set": + selection_mask = selected + elif mode == "add": + selection_mask = current_mask | selected + elif mode == "remove": + selection_mask = current_mask & ~selected + else: + raise ValueError(f"mode must be 'set', 'add' or 'remove', got {mode!r}") # Masks name the REFINABLE atoms, so freezing clears the selection. if freeze: @@ -1421,7 +1486,7 @@ def update_mask_from_selection( f"Selection '{selection_string}' ({n_selected} atoms) {action} for {target}" ) print( - f" Total refinable atoms for {target}: {n_refinable}/{len(self.pdb)}" + f" Total refinable atoms for {target}: {n_refinable}/{self.n_atoms}" ) def apply_mask_to_parameter(self, target: str): @@ -1683,45 +1748,30 @@ def shake_adp(self, stddev: float): def strip_altlocs(self) -> "Model": """Return a new model with alternate conformations removed. - For each residue that has multiple altlocs, the conformer with - highest average occupancy is kept (ties broken alphabetically). - The ``altloc`` column is cleared to ``""`` in the returned model. - The original model is not modified. - """ - import pandas as pd - - pdb = self.pdb.copy() - has_altloc = pdb["altloc"].astype(str).str.strip() != "" - if not has_altloc.any(): - return self._derive(pdb, hydrogens=self._kept_hydrogens()) - - drop_idx = [] - res_cols = ["chainid", "resseq", "icode", "resname"] - altloc_rows = pdb.loc[has_altloc] - for _, grp in altloc_rows.groupby(res_cols): - altlocs = sorted(grp["altloc"].unique()) - if len(altlocs) <= 1: - continue - # Pick conformer with highest mean occupancy - best, best_occ = altlocs[0], -1.0 - for al in altlocs: - occ = grp.loc[grp["altloc"] == al, "occupancy"].mean() - if occ > best_occ: - best, best_occ = al, occ - # Drop rows belonging to non-best conformers - for al in altlocs: - if al != best: - drop_idx.extend(grp.index[grp["altloc"] == al].tolist()) - - filtered = pdb.drop(index=drop_idx).reset_index(drop=True) - filtered["altloc"] = "" - filtered["serial"] = range(1, len(filtered) + 1) - filtered["index"] = range(len(filtered)) - - # Preserve DataFrame attrs - filtered.attrs = pdb.attrs.copy() - - return self._derive(filtered, hydrogens=self._kept_hydrogens()) + For each residue that has multiple altlocs, the conformer with the highest mean + current occupancy is kept (ties to the first in sorted order), together with the + residue's blank-altloc atoms. The returned model has no altlocs; the original is + not modified. + """ + topology = self.ctx.topology + altloc = topology.atoms.altloc + keep = np.ones(self.n_atoms, dtype=bool) + occupancy = self.occupancy().detach().cpu().numpy() + for _, labels, rows_by_altloc in self.ctx._altloc_residues(): + best = max(labels, key=lambda a: (occupancy[rows_by_altloc[a]].mean(), -labels.index(a))) + for label in labels: + if label != best: + keep[rows_by_altloc[label]] = False + rows = np.nonzero(keep)[0] + columns = {key: value[rows] for key, value in topology.columns().items()} + columns["altloc"] = np.full(len(rows), " ") + from torchref.topology import Topology + + return self._derive_from( + Topology.from_columns(columns), + self._current_values().gather(rows), + hydrogens=self._kept_hydrogens(), + ) def strip_hydrogens(self) -> "Model": """Return a new model with hydrogen atoms removed. @@ -1734,8 +1784,12 @@ def strip_hydrogens(self) -> "Model": Model New model without hydrogen atoms. """ - self.update_pdb() - return self._derive(self.pdb.copy(), hydrogens="strip", hydrogen_mode="atoms") + return self._derive_from( + self.ctx.topology, + self._current_values(), + hydrogens="strip", + hydrogen_mode="atoms", + ) def hydrogenate(self, verbose: int = 0) -> "Model": """Return a new model with the missing hydrogens added from the monomer templates. @@ -1756,15 +1810,17 @@ def hydrogenate(self, verbose: int = 0) -> "Model": ------- Model """ - self.update_pdb() - return self._derive(self.pdb.copy(), hydrogens="add", verbose=verbose) + return self._derive_from( + self.ctx.topology, self._current_values(), hydrogens="add", verbose=verbose + ) def state_dict(self, destination=None, prefix="", keep_vars=False): """ Return a dictionary containing the complete state of the Model. Registered buffers, the four parameter wrappers, the context's entries - (:meth:`ModelContext.state`), dtype and device. Restore with + (:meth:`ModelContext.state`), the atom table (:meth:`to_dataframe`), dtype and + device. Restore with :meth:`create_from_state_dict`, which is what knows how to rebuild the wrappers. Parameters @@ -1787,6 +1843,11 @@ def state_dict(self, destination=None, prefix="", keep_vars=False): for key, value in self.ctx.state().items(): state[prefix + key] = value + # The atom table is the checkpoint format for identity and starting values; + # restoring is construction, so it is split again there. + state[prefix + "pdb"] = ( + self.to_dataframe() if self.ctx.topology is not None else None + ) state[prefix + "dtype_float"] = self.dtype_float state[prefix + "device"] = self.device @@ -1831,7 +1892,7 @@ def load_state(self, path: str, strict: bool = True, device=None): print(f"Loaded model state from {path}") @staticmethod - def _restore_adp_slot(prefix, state_dict, pdb, saved_dtype, xyz_wrapper, device): + def _restore_adp_slot(prefix, state_dict, values, saved_dtype, xyz_wrapper, device): """Build the ``adp`` or ``u`` wrapper, as a node field when the state was one. With an empty ``state_dict`` this is the per-atom wrapper a fresh load uses. @@ -1848,8 +1909,8 @@ def _restore_adp_slot(prefix, state_dict, pdb, saved_dtype, xyz_wrapper, device) Which slot to rebuild. ``"u"`` carries the anisotropic representation. state_dict : dict The state being restored, read but not consumed. - pdb : pandas.DataFrame - Atom table supplying the initial values. + values : AtomValues + Supplies the initial values. saved_dtype : torch.dtype Float dtype the state was saved in. xyz_wrapper : MixedTensor @@ -1866,12 +1927,9 @@ def _restore_adp_slot(prefix, state_dict, pdb, saved_dtype, xyz_wrapper, device) name = "aniso_U" if aniso else "adp" mask = state_dict.get(f"{prefix}.refinable_mask") if aniso: - initial = torch.tensor( - pdb[["u11", "u22", "u33", "u12", "u13", "u23"]].values, - dtype=saved_dtype, - ) + initial = torch.tensor(values.u, dtype=saved_dtype) else: - initial = torch.tensor(pdb["tempfactor"].values, dtype=saved_dtype) + initial = torch.tensor(values.b, dtype=saved_dtype) saved_nl = state_dict.get(f"{prefix}.neighbor_list") if saved_nl is None: @@ -1978,19 +2036,18 @@ def create_from_state_dict( device=cpu, **cls._pop_subclass_state(state_dict), ) - instance.ctx = ModelContext.from_state( + instance.ctx, values = ModelContext.from_state( state_dict, dtype=saved_dtype, device=cpu, verbose=verbose ) - if instance.pdb is not None: - instance._install_parameters(state=state_dict) + if values is not None: + instance._install_parameters(values, state=state_dict) instance.load_state_dict(instance._restorable_entries(state_dict), strict=False) instance.to(target_device) if hasattr(instance, "reset_cache"): instance.reset_cache() if verbose > 0: - n_atoms = len(instance.pdb) if instance.pdb is not None else 0 - print(f"Created {cls.__name__} from state_dict: {n_atoms} atoms") + print(f"Created {cls.__name__} from state_dict: {instance.n_atoms} atoms") return instance @classmethod @@ -2014,8 +2071,8 @@ def get_selection_mask(self, selection: str) -> torch.Tensor: """ Return a boolean mask for atoms matching a Phenix-style selection. - Wraps :func:`~torchref.utils.utils.parse_phenix_selection`; the result can - be handed straight to ``MixedTensor.set()``. + Evaluated on the topology (:meth:`Topology.select`); the result can be handed + straight to ``MixedTensor.set()``. Parameters ---------- @@ -2043,14 +2100,12 @@ def get_selection_mask(self, selection: str) -> torch.Tensor: mask = model.get_selection_mask("chain A and (resname ALA or resname GLY)") model.xyz.set(model.xyz()[mask] + translation, mask) """ - from torchref.utils.utils import parse_phenix_selection - if not self.ctx.initialized: raise RuntimeError( "Cannot get selection mask from an uninitialized Model. Load data first." ) - return parse_phenix_selection(selection, self.pdb) + return self.ctx.topology.select(selection) def select(self, selection: str) -> "Model": """ @@ -2075,42 +2130,30 @@ def select(self, selection: str) -> "Model": ValueError If selection syntax is invalid or no atoms are selected. """ - from torchref.utils.utils import parse_phenix_selection - if not self.ctx.initialized: raise RuntimeError( "Cannot select from an uninitialized Model. Load data first." ) - mask = parse_phenix_selection(selection, self.pdb) + mask = self.get_selection_mask(selection) n_selected = int(mask.sum()) if n_selected == 0: raise ValueError(f"Selection '{selection}' matched no atoms.") - table = self._table_with_current_values().loc[mask.cpu().numpy()] - selected = self._spawn(self.ctx.derive(table, hydrogens=self._kept_hydrogens())) + rows = np.nonzero(mask.cpu().numpy())[0] riding_xyz = self.xyz.select_rows(mask) if hasattr(self.xyz, "select_rows") else None - selected._install_parameters(xyz=riding_xyz) + selected = self._derive_from( + self.ctx.topology.gather(rows), + self._current_values().gather(rows), + xyz=riding_xyz, + hydrogens=self._kept_hydrogens(), + verbose=self.ctx.verbose, + ) if self.ctx.verbose > 0: - print(f"Selected {n_selected}/{len(self.pdb)} atoms with '{selection}'") + print(f"Selected {n_selected}/{self.n_atoms} atoms with '{selection}'") return selected - def _table_with_current_values(self): - """A copy of the atom table carrying the wrappers' current values. - - Unlike :meth:`update_pdb` it leaves ``self.pdb`` alone, and ``tempfactor`` is - the isotropic wrapper's value rather than the anisotropic B_eq. - """ - table = self.pdb.copy() - table[["x", "y", "z"]] = self.xyz().detach().cpu().numpy() - table["tempfactor"] = self.adp().detach().cpu().numpy() - table[["u11", "u22", "u33", "u12", "u13", "u23"]] = ( - self.u().detach().cpu().numpy() - ) - table["occupancy"] = self.occupancy().detach().cpu().numpy() - return table - def xyz_fractional(self) -> torch.Tensor: """ Return atomic coordinates in fractional space. @@ -2340,7 +2383,7 @@ def use_rigid_xyz(self) -> "Model": Swap ``self.xyz`` for a per-chain :class:`RigidXYZTensor`. The only refinable leaves become per-chain Euler angles and translations, - with chains auto-detected from ``self.pdb["chainid"]`` (waters and + with chains auto-detected from the topology's chain ids (waters and single-atom non-polymer residues are held fixed). The original container is stashed for :meth:`restore_xyz_from_rigid`. @@ -2365,12 +2408,13 @@ def use_rigid_xyz(self) -> "Model": with torch.no_grad(): current_xyz = self.xyz().detach().clone() - chain_ids = list(self.pdb["chainid"].values) + topology = self.ctx.topology + residue_of = topology.atoms.residue_of.cpu().numpy() + chain_ids = list(topology.residues.chain[residue_of]) # Phenix-style polymer filter: drop waters and single-atom non-peptide # residues (ions), keep multi-atom HET ligands so they ride along with # their parent chain. - _WATERS = {"HOH", "WAT", "DOD", "H2O"} _STD_POLYMER = { "ALA", "ARG", "ASN", "ASP", "CYS", "GLN", "GLU", "GLY", "HIS", "ILE", "LEU", "LYS", "MET", "PHE", "PRO", "SER", "THR", "TRP", @@ -2378,12 +2422,11 @@ def use_rigid_xyz(self) -> "Model": "A", "C", "G", "T", "U", "I", "DA", "DC", "DG", "DT", "DU", "DI", } - resname = self.pdb["resname"].astype(str).str.strip() - is_water = resname.isin(_WATERS).values - is_std = resname.isin(_STD_POLYMER).values - residue_atom_count = ( - self.pdb.groupby(["chainid", "resseq", "icode"])["serial"].transform("count").values - ) + is_water = topology.is_water + is_std = np.isin(topology.residues.resname[residue_of], list(_STD_POLYMER)) + residue_atom_count = (topology.residues.atom_end - topology.residues.atom_start)[ + residue_of + ] is_single_atom = residue_atom_count == 1 drop = is_water | (is_single_atom & ~is_std) mobile_arr = ~drop @@ -2482,7 +2525,6 @@ def restore_xyz_from_rigid(self, commit: bool = True) -> "Model": xyz_mask = getattr(self, "xyz_mask", None) if xyz_mask is not None and xyz_mask.shape[0] == new_xyz.shape[0]: self.xyz.update_refinable_mask(xyz_mask) - self.pdb.loc[:, ["x", "y", "z"]] = current.cpu().numpy() else: original = getattr(self, "_rigid_original_xyz_container", None) if original is None: diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index bd520518..005d4f03 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -481,7 +481,7 @@ def _get_anomalous_cache( get_significant_elements, ) - element_list = self.pdb["element"].tolist() + element_list = self.ctx.topology.atoms.element.tolist() elements_hash = hash(tuple(element_list)) if ( @@ -744,7 +744,7 @@ def _restorable_entries(self, state_dict: dict) -> dict: if old in state_dict and new not in state_dict: state_dict[new] = state_dict.pop(old) for name in ("_A", "_B"): - if state_dict.get(name) is not None and self.pdb is not None: + if state_dict.get(name) is not None and self.ctx.topology is not None: self.register_buffer( name, torch.zeros_like(state_dict[name], device=self.device) ) diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index 67028116..c9e57776 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -388,15 +388,14 @@ def _freeze_unrestrained_residues(self): are exempt; B-factors and occupancy stay refinable. Must run after restraints are built. """ - import pandas as pd - model = self.model - pdb = getattr(model, "pdb", None) - restraints = getattr(getattr(model, "ctx", None), "restraints", None) + ctx = getattr(model, "ctx", None) + topology = getattr(ctx, "topology", None) + restraints = getattr(ctx, "restraints", None) acc = None if restraints is None else restraints.restraints - if pdb is None or acc is None: + if topology is None or acc is None: return - n = len(pdb) + n = topology.n_atoms # 1. atoms that appear in at least one geometry restraint restrained = set() @@ -423,11 +422,11 @@ def mark(idx): pass # 2. group atoms into residues (positional, aligned with xyz) - resname = pdb["resname"].astype(str).str.strip().tolist() - icode = (pdb["icode"].astype(str).tolist() if "icode" in pdb.columns - else [""] * n) - chainid = pdb["chainid"].astype(str).tolist() - resseq = pdb["resseq"].astype(str).tolist() + columns = topology.columns() + resname = columns["resname"].tolist() + icode = columns["icode"].tolist() + chainid = columns["chain"].tolist() + resseq = [str(r) for r in columns["resseq"].tolist()] res_atoms = {} for i in range(n): res_atoms.setdefault( @@ -435,7 +434,8 @@ def mark(idx): ).append(i) # 3. residues with an unrestrained atom (skip water + single-atom residues) - WATER = {"HOH", "WAT", "DOD", "H2O", "SOL", "TIP", "TIP3", "TIP4"} + from torchref.topology.residue_graph import WATER_RESNAMES as WATER + freeze_idx, frozen_res = [], [] for (c, rs, ic, rn), atoms in res_atoms.items(): if rn in WATER or len(atoms) <= 1: @@ -1312,7 +1312,7 @@ def extract_submodule_state(state_dict: dict, prefix: str) -> dict: print(f"Note: Could not initialize targets: {e}") if verbose > 0: - n_atoms = len(instance.model.pdb) if instance.model.pdb is not None else 0 + n_atoms = instance.model.n_atoms n_refl = ( instance.reflection_data.hkl.shape[0] if instance.reflection_data.hkl is not None diff --git a/torchref/refinement/targets/similarity.py b/torchref/refinement/targets/similarity.py index d01c3ef8..7ff9e391 100644 --- a/torchref/refinement/targets/similarity.py +++ b/torchref/refinement/targets/similarity.py @@ -104,22 +104,26 @@ def _build_atom_map(self): import pandas as pd import warnings - pdb_dark = self._model_dark.pdb.copy() - pdb_light = self._model_light.pdb.copy() - - for df in (pdb_dark, pdb_light): - df["_key"] = ( - df["chainid"].astype(str) - + "_" - + df["resseq"].astype(str) - + "_" - + df["icode"].astype(str).str.strip() - + "_" - + df["name"].astype(str).str.strip() - + "_" - + df["altloc"].astype(str).str.strip() + def identity(model): + columns = model.ctx.topology.columns() + return pd.DataFrame( + { + "_key": [ + f"{c}_{r}_{i.strip()}_{n.strip()}_{a.strip()}" + for c, r, i, n, a in zip( + columns["chain"], + columns["resseq"], + columns["icode"], + columns["name"], + columns["altloc"], + ) + ] + } ) + pdb_dark = identity(self._model_dark) + pdb_light = identity(self._model_light) + pdb_dark["_idx"] = range(len(pdb_dark)) pdb_light["_idx"] = range(len(pdb_light)) diff --git a/torchref/scaling/solvent.py b/torchref/scaling/solvent.py index 87f906ee..28a5e7fa 100644 --- a/torchref/scaling/solvent.py +++ b/torchref/scaling/solvent.py @@ -360,10 +360,7 @@ def get_solvent_mask(self): if self.ignore_hydrogens: # Heavy-atom radii are calibrated for masks built without hydrogens, so # adding hydrogen spheres on top would exclude solvent twice. - heavy = torch.as_tensor( - (self.model.pdb["element"].str.strip().str.upper() != "H").values, - device=xyz.device, - ) + heavy = ~self.model.ctx.topology.atoms.is_hydrogen.to(xyz.device) if not bool(heavy.all()): xyz = xyz[heavy] vdw_radii = vdw_radii[heavy] From e533af80a5ebbf204106b4680c3a951e8221b11e Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Tue, 29 Sep 2026 13:54:45 +0200 Subject: [PATCH 202/250] Keep insertion-coded residues apart when stripping altlocs strip_altlocs compares conformers within one topology residue, (chain, resseq, icode), so residues 100 and 100A no longer compete as conformers and alternates carrying different residue names resolve to one. The hydrogen-policy error leads with the requirement, the construction split is described once in the context module, and added lines are wrapped to 88 columns. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 2 +- tests/integration/test_cli_hydrogens.py | 2 +- tests/unit/model/test_hydrogen_mode.py | 4 +- tests/unit/model/test_strip_altlocs.py | 106 ++++++++++++++++++ .../ensemble/ensemble_amber_kl.py | 4 +- .../experimental/targets/forcefield_target.py | 4 +- torchref/model/context.py | 49 +++----- torchref/model/model.py | 87 ++++++++------ torchref/model/model_ft.py | 1 - torchref/topology/atom_graph.py | 5 +- torchref/topology/build.py | 2 +- torchref/topology/builders.py | 4 +- torchref/topology/nonbonded.py | 2 - torchref/topology/restraints.py | 3 +- torchref/topology/topology.py | 4 +- 15 files changed, 195 insertions(+), 84 deletions(-) create mode 100644 tests/unit/model/test_strip_altlocs.py diff --git a/docs/changelog.rst b/docs/changelog.rst index c9a7653b..8e9365c4 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -5,7 +5,7 @@ Changelog Unreleased ---------- - A model no longer keeps an atom table. Atom identity lives on ``model.ctx.topology`` (a node-only ``Topology``) and every refinable value only on the parameter wrappers; the table is read once at construction (``ModelContext.from_atoms`` splits it into the topology and ``AtomValues``) and written by ``Model.to_dataframe()``, which joins identity and current values afresh on every call. ``Model.update_pdb`` is removed, ``Model.pdb`` is a deprecated read-only view of ``to_dataframe()`` (writing into it changes nothing), ``model.n_atoms`` replaces ``len(model.pdb)``, and checkpoints keep storing the table under ``"pdb"`` so older ones still restore -- ``Model.strip_altlocs`` picks the kept conformer by current occupancy rather than the occupancies the file was loaded with +- ``Model.strip_altlocs`` compares conformers within one residue, ``(chain, resseq, icode)``, so residues 100 and 100A never compete and alternates with different residue names are resolved to one; the kept conformer is chosen by current occupancy rather than the occupancies the file was loaded with - Copying a model's context copies its LINK records as a table; it previously turned them into a list of column names, which broke a later restraint rebuild on the copy - Restraints are built from a topology rather than an atom table: ``Restraints(topology=..., xyz=...)`` and ``build_topology(_with_values)(topology, cif_dict, xyz, ...)`` take a node-only ``Topology`` (``Topology.from_table``), and the peptide-link builders pair residues along the residue graph's links, so residues with insertion codes (100, 100A, 101) are now peptide-linked and restrained like any other; structures without insertion codes get identical edges and pair lists. A table without ``icode`` or ``ATOM`` columns no longer fails, and padded names read like clean ones. The unused per-residue restraint builders, ``build_all_restraints``, ``ResidueIterator``, ``PreprocessedPDB`` and ``find_h_vdw_pairs_gpu`` are removed - Atom identity is available without building restraints: ``Topology.from_table`` gives a node-only topology (names, elements, altlocs, residues, chains, record types and charges; ``connected=False``), with ``Topology.select`` for Phenix-style selections, ``is_water`` / ``is_polymer`` masks, cached ``AtomGraph.atomic_number`` / ``vdw_radii`` and ``Topology.with_hydrogens`` for inserting a hydrogen plan. Selections are evaluated by one recursive-descent parser (``torchref.utils.selection``); a selection with more than one parenthesised group, which the old parser could resolve to every atom, now selects what it says diff --git a/tests/integration/test_cli_hydrogens.py b/tests/integration/test_cli_hydrogens.py index 13ccfc37..bfd3820f 100644 --- a/tests/integration/test_cli_hydrogens.py +++ b/tests/integration/test_cli_hydrogens.py @@ -131,5 +131,5 @@ def test_cli_refuses_riding_on_stripped_hydrogens( "riding", ] monkeypatch.setattr(sys, "argv", argv) - with pytest.raises(ValueError, match="Nothing is left to ride"): + with pytest.raises(ValueError, match="requires hydrogens=.keep. or .add."): refine.main() diff --git a/tests/unit/model/test_hydrogen_mode.py b/tests/unit/model/test_hydrogen_mode.py index bfeea83f..2e5911bf 100644 --- a/tests/unit/model/test_hydrogen_mode.py +++ b/tests/unit/model/test_hydrogen_mode.py @@ -120,7 +120,7 @@ def test_unknown_modes_are_refused(free_model, mode): def test_policy_matrix(pdb_dir, hydrogens, hydrogen_mode): """Each valid pair loads with the atom set and wrapper it names; strip+riding raises.""" if hydrogens == "strip" and hydrogen_mode == "riding": - with pytest.raises(ValueError, match="Nothing is left to ride"): + with pytest.raises(ValueError, match="requires hydrogens=.keep. or .add."): Model(verbose=0, hydrogens=hydrogens, hydrogen_mode=hydrogen_mode) return @@ -142,5 +142,5 @@ def test_policy_matrix(pdb_dir, hydrogens, hydrogen_mode): @pytest.mark.unit def test_riding_is_refused_on_a_stripped_model(pdb_dir): model = Model(verbose=0, hydrogens="strip").load_pdb(str(pdb_dir / "1DAW.pdb")) - with pytest.raises(ValueError, match="Nothing is left to ride"): + with pytest.raises(ValueError, match="requires hydrogens=.keep. or .add."): model.set_hydrogen_mode("riding") diff --git a/tests/unit/model/test_strip_altlocs.py b/tests/unit/model/test_strip_altlocs.py new file mode 100644 index 00000000..2b508e4f --- /dev/null +++ b/tests/unit/model/test_strip_altlocs.py @@ -0,0 +1,106 @@ +"""``strip_altlocs`` compares conformers within one residue, insertion code included. + +Residues 100 and 100A are different residues even when both carry an altloc label, so +neither may be dropped as the other's losing conformer. Within one residue the +conformer with the higher occupancy survives, whatever residue name it carries. +""" + +import numpy as np +import pandas as pd +import pytest + +from torchref.io.pdb import PDBReader +from torchref.model.model import Model + + +@pytest.fixture(scope="module") +def table(pdb_dir): + df, cell, sg = PDBReader(verbose=0).read(str(pdb_dir / "1DAW.pdb"))() + return df, np.asarray(cell), sg + + +def _model(df, cell, sg): + return Model(verbose=0, device="cpu").load(lambda: (df, cell, sg)) + + +@pytest.fixture(scope="module") +def baseline(table): + """Atoms left after stripping 1DAW as deposited; it has altlocs of its own.""" + return _model(*table).strip_altlocs().n_atoms + + +def _residue_rows(df, resseq): + return df.index[(df.chainid == "A") & (df.resseq == resseq)] + + +@pytest.mark.unit +def test_insertion_coded_residues_are_not_conformers_of_each_other(table, baseline): + """100 (altloc A) and 100A (altloc B) are two residues; both survive. + + Same residue name on purpose: only the insertion code tells them apart. + """ + df, cell, sg = table + df = df.copy() + first, second = _residue_rows(df, 10), _residue_rows(df, 11) + df.loc[first, ["altloc", "occupancy"]] = ["A", 0.6] + df.loc[second, ["resseq", "icode", "altloc", "occupancy", "resname"]] = [ + 10, "A", "B", 0.4, df.loc[first[0], "resname"] + ] + + model = _model(df, cell, sg) + stripped = model.strip_altlocs() + + assert stripped.n_atoms == baseline + assert (stripped.ctx.topology.atoms.altloc == " ").all() + residues = stripped.ctx.topology.residues + keys = {residues.key(r) for r in range(residues.n_residues)} + assert {("A", 10, ""), ("A", 10, "A")} <= keys + + +@pytest.mark.unit +def test_the_higher_occupancy_conformer_survives(table, baseline): + """A real two-conformer residue keeps its better conformer and its shared atoms.""" + df, cell, sg = table + rows = _residue_rows(df, 20) + side = rows[~df.loc[rows, "name"].isin(["N", "CA", "C", "O"]).to_numpy()] + minor = df.loc[side].copy() + minor[["altloc", "occupancy"]] = ["B", 0.3] + minor[["x", "y", "z"]] += 0.5 + df = df.copy() + df.loc[side, ["altloc", "occupancy"]] = ["A", 0.7] + df = pd.concat([df.loc[: rows[-1]], minor, df.loc[rows[-1] + 1 :]]).reset_index( + drop=True + ) + + model = _model(df, cell, sg) + stripped = model.strip_altlocs() + + assert model.n_atoms > baseline + assert stripped.n_atoms == baseline + kept = stripped.to_dataframe() + kept = kept[(kept.chainid == "A") & (kept.resseq == 20)] + original = df[(df.chainid == "A") & (df.resseq == 20) & (df.altloc != "B")] + np.testing.assert_allclose( + kept[["x", "y", "z"]].to_numpy(), original[["x", "y", "z"]].to_numpy(), atol=1e-4 + ) + + +@pytest.mark.unit +def test_microheterogeneity_keeps_one_residue(table): + """Alternates with different residue names at one position are alternates too.""" + df, cell, sg = table + rows = _residue_rows(df, 30) + other = df.loc[rows].copy() + other[["altloc", "occupancy", "resname"]] = ["B", 0.35, "XAA"] + df = df.copy() + df.loc[rows, ["altloc", "occupancy"]] = ["A", 0.65] + df = pd.concat([df.loc[: rows[-1]], other, df.loc[rows[-1] + 1 :]]).reset_index( + drop=True + ) + + stripped = _model(df, cell, sg).strip_altlocs() + + kept = stripped.to_dataframe() + kept = kept[(kept.chainid == "A") & (kept.resseq == 30)] + assert len(kept) == len(rows) + assert (kept.resname != "XAA").all() diff --git a/torchref/experimental/ensemble/ensemble_amber_kl.py b/torchref/experimental/ensemble/ensemble_amber_kl.py index f2e8e4b3..d1dc68e2 100644 --- a/torchref/experimental/ensemble/ensemble_amber_kl.py +++ b/torchref/experimental/ensemble/ensemble_amber_kl.py @@ -172,8 +172,8 @@ def _make_chem_model(self, ensemble: "EnsembleModel", verbose: int): def _member_xyz(self, i: int) -> torch.Tensor: """Member ``i`` coordinates ``(n_chem_atoms, 3)``, subset to kept atoms. - The returned ordering matches ``self._chem_model.to_dataframe()`` (what the OpenMM - atom map was built on), so it can be fed straight to + The returned ordering matches the rows of ``self._chem_model.to_dataframe()`` + (what the OpenMM atom map was built on), so it can be fed straight to :meth:`AmberTarget._energy`. """ xyz = self._model.xyz_per_member[i] diff --git a/torchref/experimental/targets/forcefield_target.py b/torchref/experimental/targets/forcefield_target.py index 2e6106d5..05977eb5 100644 --- a/torchref/experimental/targets/forcefield_target.py +++ b/torchref/experimental/targets/forcefield_target.py @@ -44,8 +44,8 @@ class ForceFieldTarget(ModelTarget): ---------- model : Model, optional Reference to the Model object. Should include hydrogens for accurate - energies (load with ``hydrogens="keep"`` or ``"add"``); a hydrogen-less model is not - rejected, only flagged via a warning when ``verbose > 0``. + energies (load with ``hydrogens="keep"`` or ``"add"``); a hydrogen-less model is + not rejected, only flagged via a warning when ``verbose > 0``. model_path : str, optional Path to TorchMD-Net checkpoint file (.ckpt). cutoff : float, optional diff --git a/torchref/model/context.py b/torchref/model/context.py index 05146313..d0f6a617 100644 --- a/torchref/model/context.py +++ b/torchref/model/context.py @@ -1,24 +1,12 @@ """The information half of a :class:`~torchref.model.model.Model`. -:class:`ModelContext` holds what a model *is loaded from* and *sits in* -- the unit -cell, the space group, the atom table, the link records, the provenance and the -hydrogen policy -- as opposed to what is being refined, which stays on the model as -parameter wrappers and per-atom buffers. The geometry restraints belong here too: they -are fixed by the atom set and the dictionaries, and are evaluated against coordinates -the caller passes in. - -Atom identity lives on :attr:`ModelContext.topology`, a node-only -:class:`~torchref.topology.Topology`; refinable values never live here. An atom table -(a pandas DataFrame) is read only at construction: :meth:`ModelContext.from_atoms` -settles it -- unusable rows dropped, hydrogens stripped or generated, the crystal built --- and splits it into the topology and an :class:`AtomValues` bundle of starting values -that the model's parameter wrappers are built from. Every way of making a model -- -loading a file, selecting, stripping, hydrogenating, restoring a state dict -- produces -a context and values first, and only then installs wrappers over them. - -Splitting it out means the crystallographic context can be passed to code that needs -only that (structure-factor engines, scalers, most targets) without handing over the -refinable state, and it keeps the model's own surface to parameters and behaviour. +:class:`ModelContext` holds what a model is loaded from and sits in -- cell, space +group, atom identity (:attr:`ModelContext.topology`, node-only), link records, +provenance, hydrogen policy and the geometry restraints -- as opposed to what is +refined, which lives only in the model's parameter wrappers. An atom table is read once, +by :meth:`ModelContext.from_atoms`, which splits it into the topology and the +:class:`AtomValues` the wrappers are built from; every other way of making a model goes +through :meth:`ModelContext.derive`. Mutable by design; prefer :meth:`ModelContext.copy` over editing in place. """ @@ -107,9 +95,8 @@ def check_hydrogen_policy(hydrogens: str, hydrogen_mode: str) -> None: ) if hydrogens == "strip" and hydrogen_mode == "riding": raise ValueError( - "hydrogen_mode='riding' with hydrogens='strip': you threw the hydrogens " - "overboard and then asked them to ride. Nothing is left to ride -- use " - "hydrogens='keep' or hydrogens='add'." + "hydrogen_mode='riding' requires hydrogens='keep' or 'add': you threw " + "the hydrogens overboard and then asked them to ride." ) @@ -156,10 +143,6 @@ def own_spacegroup(value, dtype: torch.dtype, device) -> Optional["SpaceGroup"]: class AtomValues: """Starting values for the parameter wrappers, one row per atom. - Read from an atom table at construction and consumed by - ``Model._install_parameters``; afterwards the wrappers are the only source of these - values. - Parameters ---------- xyz : numpy.ndarray @@ -327,7 +310,7 @@ def from_atoms( Rows without coordinates, B-factor or occupancy are dropped, the table is split into identity (:meth:`Topology.from_table`) and :class:`AtomValues`, the cell and space group are built, and the hydrogen policy is applied (see - :meth:`derive`). This is the only place a model's atoms are read from a table. + :meth:`derive`). Parameters ---------- @@ -408,7 +391,7 @@ def derive( return ctx, ctx._settle(values, self.cell.dtype) def _settle(self, values: AtomValues, dtype: torch.dtype) -> AtomValues: - """Apply the hydrogen policy to ``topology`` and ``values``; finish the context.""" + """Apply the hydrogen policy to ``topology`` and ``values``; finish up.""" if self.hydrogens == "strip": keep = ~self.topology.atoms.is_hydrogen.cpu().numpy() if not keep.all(): @@ -421,7 +404,9 @@ def _settle(self, values: AtomValues, dtype: torch.dtype) -> AtomValues: self.initialized = True return values - def _add_missing_hydrogens(self, values: AtomValues, dtype: torch.dtype) -> AtomValues: + def _add_missing_hydrogens( + self, values: AtomValues, dtype: torch.dtype + ) -> AtomValues: """Top up the hydrogens the atoms are missing; returns the extended values. Per parent, not per file: a structure deposited with some hydrogens gets the @@ -463,7 +448,7 @@ def n_atoms(self) -> int: return 0 if self.topology is None else self.topology.n_atoms def set_cif_path(self, cif_path) -> None: - """Replace the restraint dictionary path and drop restraints built over the old one. + """Replace the restraint dictionary path and drop restraints built over it. Parameters ---------- @@ -535,7 +520,7 @@ def _residue_groups(self, with_altloc: bool) -> Dict[tuple, List[int]]: return {key: keys[key] for key in sorted(keys)} def _altloc_residues(self) -> List[Tuple[tuple, List[str], Dict[str, List[int]]]]: - """Residues with more than one altloc: ``(key, sorted altlocs, rows per altloc)``. + """Residues with several altlocs: ``(key, sorted altlocs, rows per altloc)``. Keys are ``(resname, resseq, chain)``, sorted; blank-altloc atoms are not part of any conformer. @@ -668,7 +653,7 @@ def chain_sequences(self) -> List[Tuple[str, str]]: return result def _polymer_residues(self) -> List[Tuple[str, List[Tuple[int, str]]]]: - """``(chain, [(resseq, resname), ...])`` over ATOM records, chains in file order. + """``(chain, [(resseq, resname), ...])`` over ATOM records, in file order. One entry per ``(resseq, icode)``, sorted by ``resseq`` (stably, so insertion codes keep their file order). diff --git a/torchref/model/model.py b/torchref/model/model.py index 4e7b804b..89583382 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -463,7 +463,8 @@ def restraints(self): Built on first access over the current coordinates and cached on the context until the atom table or ``ctx.cif_path`` changes. Evaluations take the - coordinates as an argument, e.g. ``model.restraints.bond_deviations(model.xyz())``. + coordinates as an argument, e.g. + ``model.restraints.bond_deviations(model.xyz())``. """ if self.ctx.restraints is None: if not self.ctx.initialized: @@ -503,13 +504,8 @@ def _invalidate_atom_derived_caches(self) -> None: def load(self, reader): """ - Populate the model from a reader callable. - - The central loader that ``load_pdb`` / ``load_cif`` funnel through. The context - and starting values are built by :meth:`ModelContext.from_atoms` -- which drops - rows without coordinates, B-factor or occupancy, applies the hydrogen policy and - builds the cell and space group -- and the parameter wrappers are installed over - them. The table is not kept. + Populate the model from a reader callable, through + :meth:`ModelContext.from_atoms`; ``load_pdb`` / ``load_cif`` come through here. Parameters ---------- @@ -567,8 +563,12 @@ def _install_parameters( torch.as_tensor(values.aniso, dtype=torch.bool, device=self.device), ) self.xyz = self._build_xyz(values, state) if xyz is None else xyz - self.adp = self._restore_adp_slot("adp", state, values, dtype, self.xyz, self.device) - self.u = self._restore_adp_slot("u", state, values, dtype, self.xyz, self.device) + self.adp = self._restore_adp_slot( + "adp", state, values, dtype, self.xyz, self.device + ) + self.u = self._restore_adp_slot( + "u", state, values, dtype, self.xyz, self.device + ) # Residue-level sharing plus altloc sum-to-1 groups. initial_occ = torch.tensor(values.occupancy, dtype=dtype) @@ -599,7 +599,8 @@ def _install_parameters( ) if state.get("vdw_radii") is not None: self.register_buffer( - "vdw_radii", torch.zeros_like(state["vdw_radii"], device=self.device) + "vdw_radii", + torch.zeros_like(state["vdw_radii"], device=self.device), ) return @@ -607,7 +608,9 @@ def _install_parameters( if self.ctx.hydrogen_mode == "riding" and not isinstance( self.xyz, RidingXYZTensor ): - self.xyz = RidingXYZTensor.from_mixed_tensor(self.xyz, self.hydrogen_frames()) + self.xyz = RidingXYZTensor.from_mixed_tensor( + self.xyz, self.hydrogen_frames() + ) self._repoint_coordinate_accessors() def _build_xyz(self, values: AtomValues, state: dict): @@ -619,7 +622,9 @@ def _build_xyz(self, values: AtomValues, state: dict): coords = torch.tensor(values.xyz, dtype=self.dtype_float) mask = state.get("xyz.refinable_mask") if state.get("xyz.h_row") is None: - return MixedTensor(coords, refinable_mask=mask, name="xyz", device=self.device) + return MixedTensor( + coords, refinable_mask=mask, name="xyz", device=self.device + ) from torchref.model.riding_xyz import RidingXYZTensor from torchref.topology.hydrogens import HydrogenFrames @@ -833,7 +838,11 @@ def copy(self): for name, module in self._modules.items(): # Submodules the constructor already built (ModelFT's engine) derive from # the context and are not copied. - if module is None or name in duplicate._modules or not hasattr(module, "copy"): + if ( + module is None + or name in duplicate._modules + or not hasattr(module, "copy") + ): continue setattr(duplicate, name, module.copy()) if self._parametrization is not None: @@ -889,7 +898,9 @@ def _derive(self, pdb, **overrides) -> "Model": values = AtomValues.from_table(pdb.reset_index(drop=True)) return self._derive_from(Topology.from_table(pdb), values, **overrides) - def _derive_from(self, topology, values: AtomValues, xyz=None, **overrides) -> "Model": + def _derive_from( + self, topology, values: AtomValues, xyz=None, **overrides + ) -> "Model": """A new model of this class over ``topology`` and ``values`` in this crystal. Parameters @@ -1744,29 +1755,38 @@ def shake_adp(self, stddev: float): new_adp, refinable_mask=self.adp.refinable_mask, name="adp" ) - def strip_altlocs(self) -> "Model": """Return a new model with alternate conformations removed. - For each residue that has multiple altlocs, the conformer with the highest mean - current occupancy is kept (ties to the first in sorted order), together with the - residue's blank-altloc atoms. The returned model has no altlocs; the original is - not modified. + Conformers are compared within one topology residue, ``(chain, resseq, + icode)``, so residues 100 and 100A never compete, and alternates carrying + different residue names (microheterogeneity) are treated as the alternates they + are. In each residue with more than one altloc the conformer with the highest + mean current occupancy is kept (ties to the first in sorted order), together + with the residue's blank-altloc atoms. The returned model has no altlocs; the + original is not modified. """ + from torchref.topology import Topology + topology = self.ctx.topology altloc = topology.atoms.altloc - keep = np.ones(self.n_atoms, dtype=bool) occupancy = self.occupancy().detach().cpu().numpy() - for _, labels, rows_by_altloc in self.ctx._altloc_residues(): - best = max(labels, key=lambda a: (occupancy[rows_by_altloc[a]].mean(), -labels.index(a))) - for label in labels: - if label != best: - keep[rows_by_altloc[label]] = False + keep = np.ones(self.n_atoms, dtype=bool) + for residue in range(topology.n_residues): + rows = np.arange( + int(topology.residues.atom_start[residue]), + int(topology.residues.atom_end[residue]), + ) + labels = sorted(set(altloc[rows].tolist()) - {" "}) + if len(labels) < 2: + continue + means = [occupancy[rows[altloc[rows] == label]].mean() for label in labels] + best = labels[int(np.argmax(means))] + keep[rows[(altloc[rows] != " ") & (altloc[rows] != best)]] = False + rows = np.nonzero(keep)[0] columns = {key: value[rows] for key, value in topology.columns().items()} columns["altloc"] = np.full(len(rows), " ") - from torchref.topology import Topology - return self._derive_from( Topology.from_columns(columns), self._current_values().gather(rows), @@ -1792,7 +1812,7 @@ def strip_hydrogens(self) -> "Model": ) def hydrogenate(self, verbose: int = 0) -> "Model": - """Return a new model with the missing hydrogens added from the monomer templates. + """Return a new model with missing hydrogens added from the monomer templates. Built from the current parameter values with ``hydrogens="add"``: each residue's library template is aligned onto the heavy atoms present and its hydrogens read @@ -2141,7 +2161,9 @@ def select(self, selection: str) -> "Model": raise ValueError(f"Selection '{selection}' matched no atoms.") rows = np.nonzero(mask.cpu().numpy())[0] - riding_xyz = self.xyz.select_rows(mask) if hasattr(self.xyz, "select_rows") else None + riding_xyz = ( + self.xyz.select_rows(mask) if hasattr(self.xyz, "select_rows") else None + ) selected = self._derive_from( self.ctx.topology.gather(rows), self._current_values().gather(rows), @@ -2424,9 +2446,8 @@ def use_rigid_xyz(self) -> "Model": } is_water = topology.is_water is_std = np.isin(topology.residues.resname[residue_of], list(_STD_POLYMER)) - residue_atom_count = (topology.residues.atom_end - topology.residues.atom_start)[ - residue_of - ] + residue_sizes = topology.residues.atom_end - topology.residues.atom_start + residue_atom_count = residue_sizes[residue_of] is_single_atom = residue_atom_count == 1 drop = is_water | (is_single_atom & ~is_std) mobile_arr = ~drop diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index 005d4f03..a58814eb 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -766,4 +766,3 @@ def _restorable_entries(self, state_dict: dict) -> dict: if legacy != self.fft.compute_optimal_gridsize(self.max_res): self.explicit_gridsize = legacy return super()._restorable_entries(state_dict) - diff --git a/torchref/topology/atom_graph.py b/torchref/topology/atom_graph.py index 8d20d896..0c557eff 100644 --- a/torchref/topology/atom_graph.py +++ b/torchref/topology/atom_graph.py @@ -175,7 +175,8 @@ def __post_init__(self) -> None: self.is_hetatm = np.zeros(n, dtype=bool) if self.charge is None: self.charge = np.zeros(n, dtype=np.int64) - for edge, arity in (("bonds", 2), ("angles", 3), ("torsions", 4), ("chirals", 4)): + arities = (("bonds", 2), ("angles", 3), ("torsions", 4), ("chirals", 4)) + for edge, arity in arities: if getattr(self, edge) is None: setattr(self, edge, EdgeBlock.empty(arity, device=device)) if self._adj_indptr is None: @@ -207,7 +208,7 @@ def is_hydrogen(self) -> torch.Tensor: return torch.tensor(cache[1], device=self.bonds.indices.device) def _element_table(self) -> Tuple[np.ndarray, ...]: - """``(symbols, atomic numbers, vdW radii)``, parsed once per ``element`` array.""" + """``(symbols, atomic numbers, vdW radii)``, parsed once per ``element``.""" cache = self._element_cache if cache is None or cache[0] is not self.element: import gemmi diff --git a/torchref/topology/build.py b/torchref/topology/build.py index da2d6f29..de5f5932 100644 --- a/torchref/topology/build.py +++ b/torchref/topology/build.py @@ -1,4 +1,4 @@ -"""Connect a node-only :class:`~torchref.topology.topology.Topology` against the dictionaries. +"""Connect a node-only :class:`~torchref.topology.Topology` to the dictionaries. The input carries identity only (:meth:`Topology.from_table`); this module adds the edges. Intra-residue edges are matched template by template through the matchers in diff --git a/torchref/topology/builders.py b/torchref/topology/builders.py index a0f457f3..2e2a95c0 100644 --- a/torchref/topology/builders.py +++ b/torchref/topology/builders.py @@ -52,7 +52,7 @@ def _conformer_maps(topology, residue: int) -> List[Dict[str, int]]: def _atom_row(topology, residue: int, name: str) -> Optional[int]: - """Row of atom ``name`` in ``residue``: the blank altloc, else ``'A'``, else the first.""" + """Row of atom ``name`` in ``residue``: blank altloc, else ``'A'``, else first.""" start = int(topology.residues.atom_start[residue]) end = int(topology.residues.atom_end[residue]) hits = np.nonzero(topology.atoms.name[start:end] == name)[0] @@ -1214,5 +1214,3 @@ def build( } return result - - diff --git a/torchref/topology/nonbonded.py b/torchref/topology/nonbonded.py index b5b442c2..45a05f17 100644 --- a/torchref/topology/nonbonded.py +++ b/torchref/topology/nonbonded.py @@ -886,5 +886,3 @@ def build_vdw_restraints_gpu( print(f" Built {len(indices)} VDW restraints, {n_sym} symmetry contacts") return result - - diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index c252d02f..95b10e2b 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -453,7 +453,8 @@ def _build_vdw_restraints( ) sg_cpu = self._spacegroup.copy().to(cpu) else: - extent = float((xyz_cpu.max(dim=0).values - xyz_cpu.min(dim=0).values).max()) + span = xyz_cpu.max(dim=0).values - xyz_cpu.min(dim=0).values + extent = float(span.max()) side = extent + 2.0 * cutoff cell_cpu = Cell([side, side, side, 90.0, 90.0, 90.0], device=cpu) sg_cpu = SpaceGroup("P 1", device=cpu) diff --git a/torchref/topology/topology.py b/torchref/topology/topology.py index 235576eb..5561cdd9 100644 --- a/torchref/topology/topology.py +++ b/torchref/topology/topology.py @@ -256,7 +256,9 @@ def is_polymer(self) -> np.ndarray: per_residue = ~self.atoms.is_hetatm[first] if len(first) else np.zeros(0, bool) return per_residue[self.atoms.residue_of.cpu().numpy()] - def with_hydrogens(self, plan) -> Tuple["Topology", np.ndarray, np.ndarray, np.ndarray]: + def with_hydrogens( + self, plan + ) -> Tuple["Topology", np.ndarray, np.ndarray, np.ndarray]: """This topology's atoms with a hydrogen plan's atoms inserted. Each residue's planned hydrogens go immediately after its own atoms, never at From a8e7a198498822722d67971e066e372d1dc405d1 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Wed, 30 Sep 2026 08:39:50 +0200 Subject: [PATCH 203/250] Add torchref.uniform-rfree for a shared free set across datasets Time-resolved datasets of one crystal form must share a free set, or a reflection free in one dataset is work in another and cross-dataset R-free is biased. The new CLI gives any number of MTZ / SF-mmCIF files one CCP4 FreeR_flag while keeping every input column. An existing free set (CCP4, Phenix or mmCIF convention) is inherited by default and extended at its own fraction; the extension is seeded by a hash of the reference partition, so every extension of one free set is identical. New sets are stratified by resolution shell on the complete ASU and depend only on cell, space group and seed. --check reports whether the inputs agree, including conflicting equivalents within a file; --max-free caps a new set; --scale optionally scales the data with DatasetCollection.scale. Excluded reflections (-1 / x) are kept per file. The MTZ reader now keeps flags as integers so -1 reflections are masked on load instead of becoming work reflections. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 5 + docs/user_guide/cli.rst | 69 ++ pyproject.toml | 1 + tests/integration/test_cli_uniform_rfree.py | 201 ++++++ tests/unit/io/test_uniform_rfree.py | 242 +++++++ torchref/cli/__init__.py | 1 + torchref/cli/uniform_rfree.py | 584 ++++++++++++++++ torchref/io/mtz.py | 3 +- torchref/io/rfree.py | 714 ++++++++++++++++++++ 9 files changed, 1819 insertions(+), 1 deletion(-) create mode 100644 tests/integration/test_cli_uniform_rfree.py create mode 100644 tests/unit/io/test_uniform_rfree.py create mode 100644 torchref/cli/uniform_rfree.py create mode 100644 torchref/io/rfree.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 8ce89eec..951a9786 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,11 @@ Changelog Unreleased ---------- +- Added ``torchref.uniform-rfree``: gives any number of MTZ / SF-mmCIF files of one cell and space group a shared CCP4 ``FreeR_flag``. + - An existing free set (CCP4, Phenix or mmCIF convention) is inherited by default and extended, at its own fraction, to reflections it lacks, including those beyond its resolution. That extension is seeded with a hash of the reference's free/work partition, so it is reproducible whatever the file format or convention. New sets are stratified by resolution shell on the complete ASU and depend only on cell, space group and seed. + - Options: ``--check`` reports whether the inputs' free sets agree; ``--max-free`` caps the size of a new set; ``--scale`` optionally scales the datasets together with ``DatasetCollection.scale``. + - Excluded reflections (``-1`` / ``x``) are preserved per file. All input columns are kept, and output is MTZ and/or mmCIF. Library helpers are in ``torchref.io.rfree``. +- The MTZ reader keeps R-free flags as integers, so ``-1`` (excluded) reflections are masked on load instead of becoming work reflections. - ``ModelFT.create_from_state_dict`` restores the restraint dictionary path (``cif_path``) as ``Model`` does, and ``Refinement.create_from_state_dict`` no longer builds a stray ``Restraints`` from the model; the restored model builds its own on first access. - The difference MTZ groups its columns into named datasets -- ``observed``, ``difference``, ``light_model``, ``extrapolated_light``, ``two_moment`` -- with one history line describing each, so ``FWT``/``PHWT`` reads as ``/torchref/extrapolated_light/FWT`` (the extrapolated light-state map ``2*FEXT - Fc``). Labels are unchanged and Coot still auto-opens it - ``torchref.difference-map``, ``torchref.difference-refine`` and ``torchref.validate-ded`` gain ``--ded-weight {sigma_d,inverse_variance,none}`` and ``--sigma-d-gamma``. The difference MTZ now carries the unweighted ``DF``/``SIGDF`` on ``PHDELWT`` with one mean-one weight column per scheme, ``W_SD`` and ``W_IVW`` (MTZ type W), and the observed-to-model scale ``KSCALE``; ``DELFWT`` is no longer written, build the map with ``torchref.mtz2map -csf DF -cw W_IVW -cphi PHDELWT``. Registered in ``torchref.maps.ded_weights`` diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index eadaa72d..7118cc3d 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -97,6 +97,75 @@ restraints. :API: :mod:`torchref.cli.collection_difference_refine` +Data Utilities +-------------- + +``torchref.uniform-rfree`` +~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Give any number of structure-factor files (MTZ or SF-mmCIF) of one crystal +form a single shared R-free set, for example before refining and +difference-refining the dark and light datasets of a time-resolved experiment. +Run it before ``torchref.refine`` / ``torchref.difference-refine`` so that no +reflection is free in one dataset and work in another. + +- **Existing flags are kept by default.** If any input already has an R-free + column (CCP4 ``0 = free``, Phenix ``1 = free`` or mmCIF ``status``), its + free set is inherited and extended to the reflections it lacks. The source + is the first input with flags, or the file named with ``--reference``. + Without any flags a new set is generated. ``--fresh`` always generates one. + Replacing a set that a model was already refined against makes that model's + R-free meaningless. +- **Mixed resolution cutoffs.** If the reference stops short of the data + resolution, a warning is printed and the higher-resolution shells, plus any + gaps in the reference, are generated at the reference's free fraction, + stratified by shell. By default these shells are seeded with a hash of the + reference's free/work partition. Every extension of the same free set is + therefore identical, whatever the file format, row order, flag convention + or the other inputs. A dataset cut at lower resolution gets exactly the + matching subset of the shared flags. Datasets with fewer than 500 free + reflections are reported. +- **Reproducibility.** New flags are drawn on the complete reciprocal ASU + (Friedel mates and symmetry equivalents share a flag), with exactly the free + fraction in every resolution shell of ``--shell-size`` reflections. Each flag + depends only on cell, space group, fraction and ``--seed``, so a dataset + added later gets the same flags. +- **Excluded reflections** (MTZ flag ``-1``, mmCIF ``status x``) stay excluded + in the file that marked them, as ``-1`` / ``x``. They are not copied to the + other files. +- **Output.** Every input column is kept; existing flag columns are replaced + by a CCP4 ``FreeR_flag`` (``0..N-1``, ``0`` = free). mmCIF output writes + ``_refln.status`` ``f``/``o``/``x``, so only the free/work split is kept. + +.. code-block:: bash + + # do the existing free sets agree? (writes nothing; exit code 2 if not) + torchref.uniform-rfree dark.mtz light_*.mtz --check + + # shared free set (inherited if any input has one), MTZ and mmCIF output + torchref.uniform-rfree dark.mtz light_*.mtz --format mtz cif -o flagged/ + + # new set capped at 2000 free reflections, all light data scaled onto dark + torchref.uniform-rfree dark.mtz light_*.mtz --fresh --max-free 2000 \ + --scale --scale-reference dark -o flagged/ + +With ``--scale`` the datasets are scaled together (overall plus anisotropic, on +work reflections only) with ``DatasetCollection.scale``. The fitted factor +multiplies the observed amplitude columns and its square multiplies the +observed intensity columns; map coefficients such as ``FWT``/``PHWT`` are left +alone. By default everything goes onto the shared consensus scale; +``--scale-reference`` leaves one input unchanged instead. + +**Key options:** ``--check``, ``--reference {auto,FILE}``/``--fresh``, +``--reference-column``, ``--free-fraction`` (default: the reference's, else +0.05), ``--max-free`` (because the cap depends on resolution, the run prints +the ``--free-fraction`` that reproduces it), ``--seed`` (default: 0 for a new +set, the reference hash when extending), ``--shell-size``, +``--format {mtz,cif}``, ``--suffix``, ``--keep-old-flags``, +``--length-tol``/``--angle-tol``/``--force`` for the cell/space-group check. + +:API: :mod:`torchref.cli.uniform_rfree` + Map & Validation Utilities -------------------------- diff --git a/pyproject.toml b/pyproject.toml index e8f539ea..1a6c1a30 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -54,6 +54,7 @@ dependencies = [ "torchref.phased-difference-map" = "torchref.cli.difference_map:main" "torchref.add-metadata" = "torchref.cli.add_metadata:main" "torchref.strip-altlocs" = "torchref.cli.strip_altlocs:main" +"torchref.uniform-rfree" = "torchref.cli.uniform_rfree:main" [project.optional-dependencies] diff --git a/tests/integration/test_cli_uniform_rfree.py b/tests/integration/test_cli_uniform_rfree.py new file mode 100644 index 00000000..9a0f628f --- /dev/null +++ b/tests/integration/test_cli_uniform_rfree.py @@ -0,0 +1,201 @@ +"""End-to-end tests for ``torchref.uniform-rfree``.""" + +import subprocess +import sys + +import gemmi +import numpy as np +import pytest + +rs = pytest.importorskip("reciprocalspaceship") + +from torchref.io import rfree # noqa: E402 + +pytestmark = pytest.mark.integration + + +def _run(*args): + return subprocess.run( + [sys.executable, "-m", "torchref.cli.uniform_rfree", *map(str, args)], + capture_output=True, + text=True, + ) + + +@pytest.fixture +def inputs(mtz_dir, cif_sf_dir, tmp_path): + mtz, cif = mtz_dir / "3GR5.mtz", cif_sf_dir / "3GR5-sf.cif" + if not mtz.exists() or not cif.exists(): + pytest.skip("3GR5 test files not found") + full = rs.read_mtz(str(mtz)) + d = full.compute_dHKL()["dHKL"].to_numpy() + rng = np.random.default_rng(0) + dark = full[(d > 2.3) & (rng.random(len(full)) < 0.9)] + light = full[rng.random(len(full)) < 0.8].copy() + # a scale and anisotropy difference the scaler should remove + L = light.get_hkls()[:, 2].astype(float) + f = 1.7 * np.exp(-2.0 * (L / L.max()) ** 2) + light["FP"] = rs.DataSeries(light.FP.to_numpy() * f, index=light.index, dtype="F") + light["SIGFP"] = rs.DataSeries(light.SIGFP.to_numpy() * f, index=light.index, dtype="Q") + paths = [tmp_path / "dark.mtz", tmp_path / "light.mtz"] + dark.write_mtz(str(paths[0])) + light.write_mtz(str(paths[1])) + return full, paths + [cif] + + +def _free_tables(paths): + tables = [] + for p in paths: + ds = rfree.read_sf_file(str(p)) + keys = rfree.hkl_keys(rfree.asu_hkl(ds)).tolist() + tables.append(dict(zip(keys, (ds["FreeR_flag"].to_numpy() == 0).tolist()))) + return tables + + +def test_uniform_flags_mtz_and_cif(inputs, tmp_path): + _, paths = inputs + out = tmp_path / "out" + res = _run(*paths, "-o", out, "--format", "mtz", "cif") + assert res.returncode == 0, res.stderr + written = sorted(out.iterdir()) + assert len(written) == 6 + tables = _free_tables(written) + common = set.intersection(*(set(t) for t in tables)) + assert common + assert all(len({t[k] for t in tables}) == 1 for k in common) + dark = rs.read_mtz(str(out / "dark_rfree.mtz")) + assert {"FP", "SIGFP", "FreeR_flag"} <= set(dark.columns) + + +def test_scale_onto_reference(inputs, tmp_path): + full, paths = inputs + out = tmp_path / "out" + res = _run(*paths[:2], "-o", out, "--scale", "--scale-reference", "dark", "--device", "cpu") + assert res.returncode == 0, res.stderr + light = rs.read_mtz(str(out / "light_rfree.mtz")) + ratio = light.FP.to_numpy() / full.loc[light.index].FP.to_numpy() + assert abs(np.median(ratio) - 1) < 0.02 + + +def test_torchref_reads_flags(inputs, tmp_path): + from torchref.io.datasets.reflection_data import ReflectionData + + _, paths = inputs + out = tmp_path / "out" + assert _run(paths[0], "-o", out, "--fresh").returncode == 0 + data = ReflectionData(device="cpu", verbose=0).load_mtz(str(out / "dark_rfree.mtz")) + assert not str(data.rfree_source).startswith("Generated") + free_fraction = 1 - data.rfree_flags.float().mean().item() + assert abs(free_fraction - 0.05) < 0.01 + + +def test_mismatched_cell_fails(inputs, tmp_path): + _, paths = inputs + other = rs.read_mtz(str(paths[1])) + other.cell = gemmi.UnitCell(*(np.array(other.cell.parameters) * [1.05, 1, 1, 1, 1, 1])) + bad = tmp_path / "bad.mtz" + other.write_mtz(str(bad)) + res = _run(paths[0], bad, "-o", tmp_path / "out") + assert res.returncode == 1 + assert "cell" in res.stderr + + +def _strip_flags(src, dst): + ds = rs.read_mtz(str(src)) + ds.drop(columns=[c for c in ds.columns if c in rfree.FLAG_COLUMN_NAMES]).write_mtz(str(dst)) + + +def test_check_mode(inputs, tmp_path): + _, paths = inputs + # subsets of one file share its deposited free set + res = _run(*paths[:2], "--check") + assert res.returncode == 0, res.stderr + assert "consistent" in res.stdout + assert not any(tmp_path.glob("*_rfree.*")) + fresh = tmp_path / "fresh" + assert _run(paths[1], "--fresh", "--seed", "5", "-o", fresh).returncode == 0 + res = _run(paths[0], fresh / "light_rfree.mtz", "--check") + assert res.returncode == 2 + assert "DISAGREE" in res.stdout + + +def test_auto_inherits_or_generates(inputs, tmp_path): + full, paths = inputs + out = tmp_path / "out" + bare = tmp_path / "bare.mtz" + _strip_flags(paths[1], bare) + res = _run(bare, paths[0], "-o", out) # flags only in the second input + assert res.returncode == 0, res.stderr + assert "inherited from dark" in res.stdout + orig = _free_tables([paths[0]])[0] # the reference: dark's own free set + new = _free_tables([out / "bare_rfree.mtz"])[0] + common = set(new) & set(orig) + assert common + assert all(new[k] == orig[k] for k in common) + + none = tmp_path / "none" + _strip_flags(paths[0], tmp_path / "dark_bare.mtz") + res = _run(tmp_path / "dark_bare.mtz", bare, "-o", none) + assert res.returncode == 0, res.stderr + assert "generating a new free set" in res.stdout + + +def test_mixed_resolution_and_max_free(inputs, tmp_path): + _, paths = inputs + out = tmp_path / "out" + # dark (reference) stops at 2.3 A, light extends to 2.05 A + res = _run(*paths[:2], "-o", out, "--max-free", "500", "--fresh") + assert res.returncode == 0, res.stderr + assert "--free-fraction" in res.stdout # reproducibility hint for the cap + tables = _free_tables([out / "dark_rfree.mtz", out / "light_rfree.mtz"]) + assert sum(tables[1].values()) <= 500 + res = _run(*paths[:2], "-o", tmp_path / "inh") + assert "reference ends at" in res.stdout + assert "seed from reference free set" in res.stdout + light = rs.read_mtz(str(tmp_path / "inh" / "light_rfree.mtz")) + d = light.compute_dHKL()["dHKL"].to_numpy() + free = light["FreeR_flag"].to_numpy() == 0 + # extension beyond the reference keeps the reference's free fraction + assert abs(free[d < 2.3].mean() - free[d >= 2.3].mean()) < 0.02 + + +def test_excluded_flags_survive(inputs, tmp_path): + from torchref.io.datasets.reflection_data import ReflectionData + + _, paths = inputs + ds = rs.read_mtz(str(paths[0])) + flags = ds["FreeR_flag"].to_numpy().astype(int) + flags[:100] = -1 + ds["FreeR_flag"] = rs.DataSeries(flags, index=ds.index, dtype="I") + src = tmp_path / "excl.mtz" + ds.write_mtz(str(src)) + out = tmp_path / "out" + res = _run(src, paths[1], "-o", out, "--format", "mtz", "cif") + assert res.returncode == 0, res.stderr + back = rs.read_mtz(str(out / "excl_rfree.mtz")) + assert (back["FreeR_flag"].to_numpy()[:100] == -1).all() + # the light file does not inherit dark's exclusions + assert (rs.read_mtz(str(out / "light_rfree.mtz"))["FreeR_flag"].to_numpy() >= 0).all() + for f in ("excl_rfree.mtz", "excl_rfree.cif"): + data = ReflectionData(device="cpu", verbose=0) + data.load_mtz(str(out / f)) if f.endswith("mtz") else data.load_cif(str(out / f)) + assert int((~data.masks["flagged_initial"]).sum()) == 100, f + + +def test_unusable_reference_is_skipped_or_rejected(inputs, tmp_path): + _, paths = inputs + ds = rs.read_mtz(str(paths[0])) + ds["FreeR_flag"] = rs.DataSeries(np.ones(len(ds)), index=ds.index, dtype="I") + allwork = tmp_path / "allwork.mtz" + ds.write_mtz(str(allwork)) + bare = tmp_path / "bare.mtz" + _strip_flags(paths[1], bare) + # auto: skipped with a warning, new set generated + res = _run(allwork, bare, "-o", tmp_path / "auto") + assert res.returncode == 0, res.stderr + assert "not inheriting from 'allwork'" in res.stderr + assert "generating a new free set" in res.stdout + # explicit: clean error, no traceback + res = _run(allwork, bare, "--reference", "allwork", "-o", tmp_path / "explicit") + assert res.returncode == 1 + assert "no reflection as free" in res.stderr and "Traceback" not in res.stderr diff --git a/tests/unit/io/test_uniform_rfree.py b/tests/unit/io/test_uniform_rfree.py new file mode 100644 index 00000000..60142f96 --- /dev/null +++ b/tests/unit/io/test_uniform_rfree.py @@ -0,0 +1,242 @@ +"""Unit tests for shared R-free assignment across datasets (torchref.io.rfree).""" + +import gemmi +import numpy as np +import pytest + +rs = pytest.importorskip("reciprocalspaceship") + +from torchref.io import rfree # noqa: E402 + + +@pytest.fixture(scope="module") +def full(mtz_dir): + path = mtz_dir / "3GR5.mtz" + if not path.exists(): + pytest.skip("3GR5.mtz not found") + return rs.read_mtz(str(path)) + + +def _subset(ds, fraction, seed, dmin=None): + rng = np.random.default_rng(seed) + keep = rng.random(len(ds)) < fraction + if dmin is not None: + keep &= ds.compute_dHKL()["dHKL"].to_numpy() > dmin + return ds[keep].copy() + + +def _table(ds, flags): + return dict(zip(rfree.hkl_keys(rfree.asu_hkl(ds)).tolist(), flags.tolist())) + + +def test_complete_table_exact_per_shell(full): + keys, flags = rfree.complete_flag_table(full.cell, full.spacegroup, 2.0, 20, shell_size=1000) + assert len(flags) % 1000 == 0 + hkl = rfree._unkey(keys) + np.testing.assert_array_equal(rfree.hkl_keys(hkl), keys) + order = np.argsort(1.0 / rs.utils.compute_dHKL(hkl, full.cell) ** 2, kind="stable") + shells = flags[order].reshape(-1, 1000) + # every shell holds exactly 5% free (up to ties in resolution at edges) + assert np.all(np.abs((shells == 0).sum(1) - 50) <= 2) + + +def test_table_is_prefix_stable(full): + """Extending to higher resolution never changes existing flags.""" + k1, f1 = rfree.complete_flag_table(full.cell, full.spacegroup, 2.5, 20, seed=3) + k2, f2 = rfree.complete_flag_table(full.cell, full.spacegroup, 2.0, 20, seed=3) + values, found = rfree._lookup(k2, f2, k1) + assert found.all() + np.testing.assert_array_equal(values, f1) + + +def test_common_reflections_share_flags(full): + sets = { + "dark": _subset(full, 0.9, 1, dmin=2.3), + "light1": _subset(full, 0.8, 2), + "light2": _subset(full, 0.7, 3), + } + flags, info = rfree.uniform_rfree(sets, free_fraction=0.05, seed=0) + assert info["n_flags"] == 20 + tables = [_table(ds, flags[n]) for n, ds in sets.items()] + common = set.intersection(*(set(t) for t in tables)) + assert common + assert all(len({t[k] for t in tables}) == 1 for k in common) + for n, ds in sets.items(): + assert len(flags[n]) == len(ds) + assert abs((flags[n] == 0).mean() - 0.05) < 0.01 + + +def test_symmetry_and_friedel_mates_share_flag(full): + ds = _subset(full, 1.0, 0) + # add a copy of every reflection as its Friedel mate + mate = ds.copy().reset_index() + mate[["H", "K", "L"]] *= -1 + both = rs.concat([ds, mate.set_index(["H", "K", "L"])]) + flags, _ = rfree.uniform_rfree({"x": both}, seed=0) + n = len(ds) + np.testing.assert_array_equal(flags["x"][:n], flags["x"][n:]) + + +def test_deterministic_and_independent_of_coverage(full): + a = _subset(full, 0.6, 5) + b = _subset(full, 0.9, 6) + fa, _ = rfree.uniform_rfree({"a": a}, seed=7) + fb, _ = rfree.uniform_rfree({"b": b}, seed=7, dmin=1.8) + ta, tb = _table(a, fa["a"]), _table(b, fb["b"]) + common = set(ta) & set(tb) + assert all(ta[k] == tb[k] for k in common) + fa2, _ = rfree.uniform_rfree({"a": a}, seed=7) + np.testing.assert_array_equal(fa["a"], fa2["a"]) + + +@pytest.mark.parametrize("convention", ["ccp4", "phenix"]) +def test_reference_free_set_is_inherited(full, convention): + ref = _subset(full, 0.8, 11) + free = np.random.default_rng(3).random(len(ref)) < 0.1 + if convention == "ccp4": + values = np.where(free, 0, np.random.default_rng(4).integers(1, 10, len(ref))) + else: + values = free.astype(int) # Phenix: 1 = free + ref["R-free-flags"] = rs.DataSeries(values, index=ref.index, dtype="I") + + target = _subset(full, 0.9, 12) + flags, info = rfree.uniform_rfree({"t": target}, reference=ref, seed=0) + assert info["reference"]["column"] == "R-free-flags" + assert info["reference"]["convention"].startswith( + "ccp4" if convention == "ccp4" else "binary" + ) + ref_free = _table(ref, free.astype(int)) + t = _table(target, flags["t"]) + for k in set(ref_free) & set(t): + assert (t[k] == 0) == bool(ref_free[k]) + + +def test_check_compatible_flags_mismatch(full): + other = full.copy() + other.cell = gemmi.UnitCell(*(np.array(full.cell.parameters) * [1.05, 1, 1, 1, 1, 1])) + assert rfree.check_compatible({"a": full, "b": full.copy()}) == [] + assert rfree.check_compatible({"a": full, "b": other}) + + +def test_apply_flags_replaces_existing_columns(full): + ds = full.copy() + flags = np.zeros(len(ds), dtype=np.int32) + out = rfree.apply_flags(ds, flags) + assert list(out.columns).count("FreeR_flag") == 1 + assert "FreeR_flag_orig" not in out.columns + kept = rfree.apply_flags(ds, flags, keep_old=True) + assert "FreeR_flag_orig" in kept.columns + assert set(ds.columns) - {"FreeR_flag"} <= set(out.columns) + + +def test_scale_columns_amplitude_and_intensity(): + ds = rs.DataSet( + { + "H": [1, 2], "K": [0, 0], "L": [0, 0], + "FP": [10.0, 20.0], "SIGFP": [1.0, 2.0], + "I": [100.0, 400.0], "SIGI": [10.0, 20.0], + "PHI": [30.0, 40.0], + "FWT": [5.0, 6.0], "PHWT": [0.0, 90.0], + }, + cell=[50, 50, 50, 90, 90, 90], spacegroup=1, + ).set_index(["H", "K", "L"]) + ds = ds.astype({"FP": "F", "SIGFP": "Q", "I": "J", "SIGI": "Q", "PHI": "P", "FWT": "F", "PHWT": "P"}) + out, cols = rfree.scale_columns(ds, np.array([2.0, 0.5])) + assert cols == ["FP", "SIGFP", "I", "SIGI"] + np.testing.assert_allclose(out.FP, [20, 10]) + np.testing.assert_allclose(out.SIGFP, [2, 1]) + np.testing.assert_allclose(out.I, [400, 100]) + np.testing.assert_allclose(out.SIGI, [40, 5]) + np.testing.assert_allclose(out.PHI, [30, 40]) + np.testing.assert_allclose(out.FWT, [5, 6]) # map coefficients untouched + + +def _with_flags(ds, free, phenix=False): + out = ds.copy() + for c in [c for c in out.columns if c in rfree.FLAG_COLUMN_NAMES]: + out = out.drop(columns=c) + values = free.astype(int) if phenix else (~free).astype(int) + out["R-free-flags" if phenix else "FreeR_flag"] = rs.DataSeries( + values, index=out.index, dtype="I" + ) + return out + + +def test_extension_seed_depends_only_on_reference_partition(full): + """Extending one free set gives the same new flags whatever its format.""" + d = full.compute_dHKL()["dHKL"].to_numpy() + low = full[d > 2.6] + free = np.random.default_rng(0).random(len(low)) < 0.05 + ccp4 = _with_flags(low, free) + phenix = _with_flags(low, free, phenix=True).sample(frac=1.0, random_state=1) + target = full.copy() + f1, i1 = rfree.uniform_rfree({"t": target}, reference=ccp4) + f2, i2 = rfree.uniform_rfree({"t": target}, reference=phenix) + assert i1["seed_source"] == "reference free set" + assert i1["seed"] == i2["seed"] + np.testing.assert_array_equal(f1["t"] == 0, f2["t"] == 0) + assert i1["n_generated_beyond_reference"] > 0 + + # a different reference partition changes the extension + other = _with_flags(low, np.roll(free, 1)) + f3, i3 = rfree.uniform_rfree({"t": target}, reference=other) + assert i3["seed"] != i1["seed"] + beyond = d <= 2.6 + assert ((f1["t"] == 0) != (f3["t"] == 0))[beyond].any() + + # an explicit seed still wins + _, i4 = rfree.uniform_rfree({"t": target}, reference=ccp4, seed=5) + assert i4["seed"] == 5 and i4["seed_source"] == "user" + + +def test_extension_is_identical_across_targets(full): + d = full.compute_dHKL()["dHKL"].to_numpy() + low = full[d > 2.6] + ref = _with_flags(low, np.random.default_rng(2).random(len(low)) < 0.05) + a, b = _subset(full, 0.7, 21), _subset(full, 0.7, 22) + fa, _ = rfree.uniform_rfree({"a": a}, reference=ref) + fb, _ = rfree.uniform_rfree({"b": b}, reference=ref) + ta, tb = _table(a, fa["a"] == 0), _table(b, fb["b"] == 0) + common = set(ta) & set(tb) + assert all(ta[k] == tb[k] for k in common) + + +@pytest.fixture +def small(mtz_dir): + """30 rows of deposited 1DAW (the reviewer's probe set).""" + path = mtz_dir / "1DAW.mtz" + if not path.exists(): + pytest.skip("1DAW.mtz not found") + return rs.read_mtz(str(path)).iloc[:30].copy() + + +def test_all_excluded_rows_stay_excluded(small): + small["FreeR_flag"] = rs.DataSeries(np.full(len(small), -1), index=small.index, dtype="I") + flags, info = rfree.uniform_rfree({"x": small}, seed=0) + assert info["n_excluded"]["x"] == len(small) + assert (flags["x"] == -1).all() + + +def test_conflicting_equivalents_are_inconsistent(small): + free_row = small[small["FreeR_flag"] == 0].iloc[:1] + if free_row.empty: + free_row = small.iloc[:1].copy() + small.loc[free_row.index, "FreeR_flag"] = 0 + free_row = small.loc[free_row.index] + dup = free_row.copy() + dup["FreeR_flag"] = rs.DataSeries([1], index=dup.index, dtype="I") # work copy + bad = rs.concat([small, dup]) + report = rfree.compare_free_sets({"good": small, "bad": bad}) + assert report["files"]["bad"]["n_conflicting"] == 1 + assert not report["consistent"] + # a single inconsistent file is not consistent either + assert not rfree.compare_free_sets({"bad": bad})["consistent"] + assert rfree.compare_free_sets({"good": small})["consistent"] + + +def test_reference_without_free_reflections_is_rejected(small): + small["FreeR_flag"] = rs.DataSeries(np.ones(len(small)), index=small.index, dtype="I") + with pytest.raises(ValueError, match="no reflection as free"): + rfree.uniform_rfree({"x": small}, reference=small) + report = rfree.compare_free_sets({"x": small}) + assert report["files"]["x"]["n_free"] == 0 and not report["consistent"] diff --git a/torchref/cli/__init__.py b/torchref/cli/__init__.py index efba6bad..8fba6fc9 100644 --- a/torchref/cli/__init__.py +++ b/torchref/cli/__init__.py @@ -9,5 +9,6 @@ "mtz2map", "refine", "strip_altlocs", + "uniform_rfree", "validate_ded", ] diff --git a/torchref/cli/uniform_rfree.py b/torchref/cli/uniform_rfree.py new file mode 100644 index 00000000..6ba911f5 --- /dev/null +++ b/torchref/cli/uniform_rfree.py @@ -0,0 +1,584 @@ +#!/usr/bin/env python3 -u + +"""Give several structure-factor files one shared R-free set. + +Intended as a pre-processing step for time-resolved (TR-SFX) data: every +dataset of one crystal form (dark, light, time points) gets the same CCP4 +``FreeR_flag`` column before refinement / difference refinement, so no +reflection is free in one dataset and work in another. Optionally, the +datasets are also put on a common scale with the joint dataset scaler +(:meth:`torchref.io.datasets.collection.DatasetCollection.scale`). + +If any input already carries an R-free column, its free set is inherited and +extended to the reflections it lacks (``--reference auto``, the default); +otherwise a new set is generated. ``--check`` only reports whether the +inputs' existing free sets agree. + +Usage:: + + torchref.uniform-rfree dark.mtz light_*.mtz --check + torchref.uniform-rfree dark.mtz light_*.mtz -o flagged/ + torchref.uniform-rfree dark.mtz light.cif --reference deposited-sf.cif \\ + --format mtz cif --scale --scale-reference dark -o flagged/ +""" + +import argparse +import sys +import tempfile +from pathlib import Path + +import numpy as np + +from torchref.cli._common import ( + add_device_arg, + add_outdir_arg, + add_verbose_arg, + configure_unbuffered_output, +) + +configure_unbuffered_output() + + +def _parse_args(argv=None): + parser = argparse.ArgumentParser( + prog="torchref.uniform-rfree", + description=( + "Assign one uniform R-free set (CCP4 FreeR_flag, 0 = free) to any " + "number of MTZ / SF-mmCIF files sharing a cell and space group, " + "optionally scaling them together." + ), + formatter_class=argparse.RawDescriptionHelpFormatter, + epilog=""" +Examples: + # Do the existing free sets agree? (writes nothing; exit code 2 if not) + torchref.uniform-rfree dark.mtz light_*.mtz --check + + # Shared free set: inherited from the first input that has one, else new + torchref.uniform-rfree dark.mtz light_*.mtz -o flagged/ + + # Inherit the free set of a deposited dark structure, extend to new reflections + torchref.uniform-rfree dark.mtz light.mtz --reference 1abc-sf.cif -o flagged/ + + # Ignore all existing flags; new 5% set capped at 2000 free reflections + torchref.uniform-rfree dark.mtz light_*.mtz --fresh --max-free 2000 -o flagged/ + + # Write MTZ and mmCIF, and scale all light datasets onto dark + torchref.uniform-rfree dark.mtz light_*.mtz --format mtz cif \\ + --scale --scale-reference dark -o flagged/ +""", + ) + + inp = parser.add_argument_group("Input") + inp.add_argument( + "files", nargs="+", help="Structure-factor files (.mtz or .cif)" + ) + inp.add_argument( + "--cif-block", default=None, help="Data block to read from CIF inputs" + ) + inp.add_argument( + "--length-tol", + type=float, + default=0.01, + help="Allowed relative cell-length deviation (default: 0.01)", + ) + inp.add_argument( + "--angle-tol", + type=float, + default=0.5, + help="Allowed cell-angle deviation in degrees (default: 0.5)", + ) + inp.add_argument( + "--force", + action="store_true", + help="Proceed despite cell / space-group mismatches", + ) + + flg = parser.add_argument_group("Flags") + flg.add_argument( + "--check", + action="store_true", + help=( + "Only report the inputs' existing free sets and whether they agree; " + "write nothing (exit code 0 if consistent, 2 if not)" + ), + ) + flg.add_argument( + "--free-fraction", + type=float, + default=None, + help=( + "Free-set fraction; FreeR_flag takes round(1/f) values (default: the " + "reference's own fraction when inheriting, else 0.05)" + ), + ) + flg.add_argument( + "--max-free", + type=int, + default=None, + help=( + "Cap the free set at this many reflections of the complete set to the " + "best resolution (e.g. 2000, as in Phenix); lowers the fraction" + ), + ) + flg.add_argument( + "--shell-size", + type=int, + default=1000, + help=( + "Reflections per resolution shell; each shell holds exactly the " + "free fraction (default: 1000)" + ), + ) + flg.add_argument( + "--seed", + type=int, + default=None, + help=( + "Random seed (default: 0 for a new set; when extending a reference, a " + "hash of its free set, so every extension of it is identical)" + ), + ) + flg.add_argument( + "--dmin", + type=float, + default=None, + help=( + "Extend the flag table to at least this resolution (default: best " + "resolution of the inputs). Flags never depend on this value; it " + "only matters for what gets reported." + ), + ) + ref = flg.add_mutually_exclusive_group() + ref.add_argument( + "--reference", + default="auto", + help=( + "File whose existing R-free flags are inherited and extended (may be " + "one of the inputs). 'auto' (default): the first input with an R-free " + "column; a new set is generated if none has one" + ), + ) + ref.add_argument( + "--fresh", + action="store_true", + help="Ignore existing flags and generate a new free set", + ) + flg.add_argument( + "--reference-column", + default=None, + help="R-free column in --reference (default: auto-detect)", + ) + + scl = parser.add_argument_group("Scaling") + scl.add_argument( + "--scale", + action="store_true", + help=( + "Jointly scale the datasets (overall + anisotropic) with the " + "dataset scaler; work reflections only" + ), + ) + scl.add_argument( + "--scale-reference", + default=None, + help=( + "Input (file name or stem) kept unscaled; others are put on its scale. " + "Default: the consensus scale of all inputs" + ), + ) + scl.add_argument("--scale-nsteps", type=int, default=10, help=argparse.SUPPRESS) + scl.add_argument("--scale-max-iter", type=int, default=100, help=argparse.SUPPRESS) + add_device_arg(scl) + + out = parser.add_argument_group("Output") + add_outdir_arg(out, required=False, help="Output directory (required unless --check)") + out.add_argument( + "--format", + nargs="+", + choices=["mtz", "cif"], + default=["mtz"], + help="Output format(s) (default: mtz)", + ) + out.add_argument( + "--suffix", + default="_rfree", + help="Suffix appended to each output file stem (default: _rfree)", + ) + out.add_argument( + "--keep-old-flags", + action="store_true", + help="Keep existing flag columns renamed to _orig instead of dropping them", + ) + add_verbose_arg(out) + return parser.parse_args(argv) + + +def _unique_names(paths): + """Short, unique dataset names from file stems.""" + names = {} + for p in paths: + stem = Path(p).stem + name, i = stem, 1 + while name in names: + i += 1 + name = f"{stem}_{i}" + names[name] = p + return names + + +def _resolve_input(key, names): + """Match a user-given file name / stem against the dataset names.""" + for name, path in names.items(): + if key in (name, path, Path(path).name, str(Path(path).resolve())): + return name + return None + + +MIN_FREE_WARN = 500 + + +def _existing_report(report): + """Print the existing free set of every input and pairwise agreement.""" + print("\nExisting R-free flags:") + for name, f in report["files"].items(): + if f["column"] is None: + problem = f.get("problem", "") + why = "" if "no recognised" in problem else f" ({problem})" + print(f" {name:<24s} none{why} dmin {f['dmin']:5.2f} A") + continue + print( + f" {name:<24s} {f['column']:<12s} {f['convention']:<18s} " + f"dmin {f['dmin']:5.2f} A free {f['n_free']:>7d} " + f"({100 * f['n_free'] / max(f['n'], 1):5.2f} %) excluded {f['n_excluded']}" + ) + if f["n_free"] == 0: + print(" Warning: no free reflections") + if f["n_conflicting"]: + print( + f" Warning: {f['n_conflicting']} reflections have symmetry " + "equivalents with different flags in this file" + ) + for (a, b), (n_common, n_bad) in report["pairs"].items(): + status = "agree" if n_bad == 0 else f"DISAGREE on {n_bad}" + print(f" {a} vs {b}: {n_common} common reflections, {status}") + if report["consistent"]: + print(" -> free sets are consistent") + else: + files = report["files"] + missing = [n for n, f in files.items() if f["column"] is None] + bad = [ + n + for n, f in files.items() + if f["column"] is not None and (f["n_free"] == 0 or f["n_conflicting"]) + ] + if missing: + why = f"no flags in {', '.join(missing)}" + elif bad: + why = f"invalid free set in {', '.join(bad)}" + else: + why = "flags disagree" + print(f" -> free sets are NOT consistent ({why})") + + +def _new_report(datasets, flags, info, ref_label, args): + """Print the assigned free set: summary, warnings, one line per file.""" + pct = f"{100 * info['free_fraction']:.2f} %" + print( + f"\nFreeR_flag: {info['n_flags']} values, {pct} free " + f"({info['fraction_source']}), to {info['dmin']:.2f} A" + ) + if info["fraction_source"].startswith("max_free"): + print(f" reproduce with --free-fraction {info['free_fraction']:.6g}") + if args.verbose > 1: + print(f" seed {info['seed']} ({info['seed_source']})") + if "reference" in info: + r = info["reference"] + print( + f" inherited from {ref_label} ({r['column']}, {r['convention']}, " + f"{r['dmin']:.2f} A): {info['n_inherited']} kept, {info['n_generated']} new" + ) + n_beyond, n_gaps = info["n_generated_beyond_reference"], info["n_gaps_in_reference"] + if n_beyond: + print( + f" Warning: reference ends at {r['dmin']:.2f} A; {n_beyond} reflections " + f"to {info['dmin']:.2f} A newly flagged ({pct} free, " + f"seed from {info['seed_source']})" + ) + if n_gaps: + frac = n_gaps / max(n_gaps + info["n_inherited"], 1) + print( + (" Warning: " if frac > 0.05 else " ") + + f"{n_gaps} reflections ({100 * frac:.1f} %) missing within the " + "reference's range, newly flagged" + ) + if r["n_inconsistent"]: + print( + f" Warning: {r['n_inconsistent']} reference reflections have " + "conflicting equivalents (free wins)" + ) + if info["n_off_asu"]: + print( + f" Warning: {info['n_off_asu']} reflections outside the complete ASU " + "(systematic absences?), flagged by hash" + ) + for name, ds in datasets.items(): + _flag_report(name, flags[name], ds, info, args.verbose) + if len(datasets) > 1 and args.verbose > 1: + common = set.intersection( + *(set(_rfree_keys(ds).tolist()) for ds in datasets.values()) + ) + print(f" {len(common)} unique reflections common to all inputs") + + +def _rfree_keys(ds): + """Unique-ASU keys of every row of ``ds``.""" + from torchref.io import rfree + + return rfree.hkl_keys(rfree.asu_hkl(ds)) + + +def _flag_report(name, flags, ds, info, verbose): + from torchref.io import rfree + + n = len(flags) + n_free = int((flags == 0).sum()) + n_excl = info["n_excluded"].get(name, 0) + dmin = ds.compute_dHKL()["dHKL"].min() + print( + f" {name:<24s} {n:>9d} rows dmin {dmin:5.2f} A free {n_free:>7d} " + f"({100 * n_free / max(n - n_excl, 1):5.2f} %)" + + (f" excluded {n_excl} (kept as -1)" if n_excl else "") + ) + if n_free < MIN_FREE_WARN: + print(f" Warning: only {n_free} free reflections; R-free will be noisy") + if verbose > 1: + dstar2 = 1.0 / ds.compute_dHKL()["dHKL"].to_numpy() ** 2 + bins = rfree.resolution_bins(dstar2, 10) + for b in range(bins.max() + 1): + sel = bins == b + d_lo = 1 / np.sqrt(dstar2[sel].min()) + d_hi = 1 / np.sqrt(dstar2[sel].max()) + frac = (flags[sel] == 0).sum() / max((flags[sel] >= 0).sum(), 1) + print(f" {d_lo:6.2f} - {d_hi:5.2f} A free {100 * frac:5.2f} %") + + +def _scale_datasets(flagged, args, device): + """Jointly scale flagged datasets; returns name -> per-row amplitude factor.""" + import torch + + from torchref.cli._common import load_reflection_data + from torchref.config import get_int_dtype + from torchref.io import rfree + from torchref.io.datasets.collection import DatasetCollection + + collection = DatasetCollection(verbose=args.verbose, device=device) + with tempfile.TemporaryDirectory() as tmp: + for name, ds in flagged.items(): + path = Path(tmp) / f"{name}.mtz" + rfree.write_sf_file(ds, str(path)) + data = load_reflection_data(str(path), device=device, verbose=0) + collection.add_dataset(name, data) + collection.scale(nsteps=args.scale_nsteps, max_iter=args.scale_max_iter) + + scaler = collection.scaler + + def row_factors(key, ds, sg): + # the scaler's anisotropy lives in the canonical-ASU setting + hkl = torch.as_tensor(ds.get_hkls(), dtype=get_int_dtype()) + canonical, _, _, order = sg.canonicalize_hkl(hkl) + with torch.no_grad(): + f_sorted = scaler(key, canonical).cpu().numpy() + f = np.empty(len(hkl)) + f[order.cpu().numpy()] = f_sorted + return f + + factors = {} + for name, ds in flagged.items(): + sg = collection[name].spacegroup + f = row_factors(name, ds, sg) + if args.scale_reference is not None: + f = f / row_factors(args.scale_reference, ds, sg) + factors[name] = f + return factors, collection.scaling_metrics + + +def main(argv=None): + """Entry point for ``torchref.uniform-rfree``; returns the process exit code.""" + args = _parse_args(argv) + from torchref.io import rfree + + names = _unique_names(args.files) + for name, path in names.items(): + if not Path(path).is_file(): + print(f"Error: file not found: {path}", file=sys.stderr) + return 1 + + # ---- read ------------------------------------------------------------- + datasets = {} + for name, path in names.items(): + try: + datasets[name] = rfree.read_sf_file(path, cif_block=args.cif_block) + except Exception as exc: # noqa: BLE001 - report and exit cleanly + print(f"Error: cannot read {path}: {exc}", file=sys.stderr) + return 1 + if args.verbose: + ds = datasets[name] + dmin = ds.compute_dHKL()["dHKL"].min() + print( + f"Read {path}: {len(ds)} rows, {ds.spacegroup.xhm()}, " + f"cell {tuple(round(x, 3) for x in ds.cell.parameters)}, dmin {dmin:.2f} A" + ) + + problems = rfree.check_compatible(datasets, args.length_tol, args.angle_tol) + if problems: + for p in problems: + print(("Warning: " if args.force else "Error: ") + p, file=sys.stderr) + if not args.force: + print("Use --force to proceed anyway.", file=sys.stderr) + return 1 + + # ---- existing flags ------------------------------------------------- + report = rfree.compare_free_sets(datasets) + if args.check or args.verbose > 1 or (args.verbose and not report["consistent"]): + _existing_report(report) + elif args.verbose: + print(f"\nExisting free sets are consistent across all {len(datasets)} inputs.") + if args.check: + return 0 if report["consistent"] else 2 + if args.outdir is None: + print("Error: -o/--outdir is required (or use --check)", file=sys.stderr) + return 1 + + reference, ref_label = None, None + with_flags = [n for n, f in report["files"].items() if f["column"] is not None] + usable = [ + n + for n in with_flags + if report["files"][n]["n_free"] > 0 and report["files"][n]["n_conflicting"] == 0 + ] + if args.fresh: + if with_flags and args.verbose: + print( + f"--fresh: replacing the free set(s) of {', '.join(with_flags)}; " + "R-free of models refined against them becomes biased." + ) + elif args.reference == "auto": + for name in [n for n in with_flags if n not in usable]: + print( + f"Warning: not inheriting from {name!r}: its free set is empty or " + "has conflicting symmetry equivalents", + file=sys.stderr, + ) + if usable: + ref_label = usable[0] + reference = datasets[ref_label] + if not report["consistent"] and len(with_flags) > 1: + print( + f"Warning: existing free sets disagree; inheriting from {ref_label!r} " + "(the first input with flags). Pass --reference to choose another.", + file=sys.stderr, + ) + elif args.verbose: + print("No input carries a usable free set; generating a new free set.") + else: + ref_label = _resolve_input(args.reference, names) + try: + reference = ( + datasets[ref_label] + if ref_label is not None + else rfree.read_sf_file(args.reference, cif_block=args.cif_block) + ) + except Exception as exc: # noqa: BLE001 + print(f"Error: cannot read reference {args.reference}: {exc}", file=sys.stderr) + return 1 + ref_label = ref_label or args.reference + if rfree.flag_column(reference, args.reference_column) is None: + print(f"Error: reference {args.reference} has no R-free column", file=sys.stderr) + return 1 + problems = rfree.check_compatible( + {"inputs": next(iter(datasets.values())), "reference": reference}, + args.length_tol, + args.angle_tol, + ) + if problems and not args.force: + for p in problems: + print("Error: " + p, file=sys.stderr) + return 1 + + # ---- flags ------------------------------------------------------------ + try: + flags, info = rfree.uniform_rfree( + datasets, + free_fraction=args.free_fraction, + shell_size=args.shell_size, + seed=args.seed, + dmin=args.dmin, + reference=reference, + reference_column=args.reference_column, + max_free=args.max_free, + ) + except ValueError as exc: + print(f"Error: {exc}", file=sys.stderr) + return 1 + + if args.verbose: + _new_report(datasets, flags, info, ref_label, args) + + flagged = { + name: rfree.apply_flags(ds, flags[name], keep_old=args.keep_old_flags) + for name, ds in datasets.items() + } + + # ---- optional scaling ------------------------------------------------- + if args.scale: + if len(flagged) < 2: + print("Error: --scale needs at least two inputs", file=sys.stderr) + return 1 + if args.scale_reference is not None: + key = _resolve_input(args.scale_reference, names) + if key is None: + print( + f"Error: --scale-reference {args.scale_reference!r} is not an input", + file=sys.stderr, + ) + return 1 + args.scale_reference = key + from torchref.config import normalize_device + + device = normalize_device(args.device) + if args.verbose: + print(f"\nScaling {len(flagged)} datasets jointly on {device} ...") + try: + factors, metrics = _scale_datasets(flagged, args, device) + except Exception as exc: # noqa: BLE001 + print(f"Error: scaling failed: {exc}", file=sys.stderr) + return 1 + for name in flagged: + flagged[name], cols = rfree.scale_columns(flagged[name], factors[name]) + if args.verbose: + f = factors[name] + print( + f" {name:<24s} factor median {np.median(f):.4f} " + f"[{f.min():.4f}, {f.max():.4f}] columns: {', '.join(cols) or '-'}" + ) + if args.verbose > 1 and metrics: + print(f" metrics: {metrics}") + + # ---- write ------------------------------------------------------------ + outdir = Path(args.outdir) + outdir.mkdir(parents=True, exist_ok=True) + for name, ds in flagged.items(): + for fmt in args.format: + path = outdir / f"{name}{args.suffix}.{fmt}" + try: + rfree.write_sf_file(ds, str(path)) + except Exception as exc: # noqa: BLE001 + print(f"Error: cannot write {path}: {exc}", file=sys.stderr) + return 1 + if args.verbose: + print(f"Wrote {path}") + return 0 + + +if __name__ == "__main__": + sys.exit(main() or 0) diff --git a/torchref/io/mtz.py b/torchref/io/mtz.py index ba11fbc7..a06052ca 100644 --- a/torchref/io/mtz.py +++ b/torchref/io/mtz.py @@ -440,7 +440,8 @@ def _extract_rfree_flags(self) -> None: free_pct = 100.0 * n_free / len(rfree_flags) print(f" After flip: free={n_free} ({free_pct:.1f}%)") - self.data["R-free-flags"] = rfree_flags.astype(bool) + # keep int: -1 (excluded) is masked by ReflectionData.load + self.data["R-free-flags"] = rfree_flags self.data["R-free-source"] = col return diff --git a/torchref/io/rfree.py b/torchref/io/rfree.py new file mode 100644 index 00000000..1d0aac26 --- /dev/null +++ b/torchref/io/rfree.py @@ -0,0 +1,714 @@ +""" +Uniform R-free flag assignment across several structure-factor files. + +Time-resolved experiments refine many datasets of one crystal form (dark, +light, time points) and difference-refine them against each other. Every file +must then share one free set, otherwise reflections that are free in one +dataset are work reflections in another and cross-dataset R-free is biased. + +Flags are assigned per unique reciprocal-ASU index (Friedel mates merged), so +symmetry equivalents, Friedel mates and F(+)/F(-) rows always share a flag. +Output follows the CCP4 ``FreeR_flag`` convention: integers ``0..N-1`` with +``0`` = free, so a different test set ``k`` can still be selected later. + +All functions here operate on :class:`reciprocalspaceship.DataSet` objects so +that every original column of the input files is preserved on output. +""" + +import hashlib +from pathlib import Path +from typing import Dict, List, Optional, Tuple + +import gemmi +import numpy as np +import reciprocalspaceship as rs + +from torchref.io.mtz import MTZReader + +FREE_COLUMN = "FreeR_flag" + +# Existing flag columns that are replaced on output. +FLAG_COLUMN_NAMES = tuple(dict.fromkeys([*MTZReader.RFREE_FLAG_NAMES, FREE_COLUMN])) + +_KEY_OFFSET = 1 << 10 # |h|, |k|, |l| < 1024 + + +# --------------------------------------------------------------------------- +# I/O +# --------------------------------------------------------------------------- + + +def read_sf_file(path: str, cif_block: Optional[str] = None) -> rs.DataSet: + """Read an MTZ or SF-mmCIF file into a DataSet, keeping every column. + + Parameters + ---------- + path : str + ``.mtz`` or ``.cif`` / ``.mmcif`` file. + cif_block : str, optional + Data block name for multi-block CIF files. Defaults to the first + block that carries reflections. + + Returns + ------- + rs.DataSet + Reflections indexed by H, K, L with cell and space group attached. + """ + suffix = Path(path).suffix.lower() + if suffix == ".mtz": + return rs.read_mtz(str(path)) + if suffix in (".cif", ".mmcif", ".ent"): + blocks = gemmi.as_refln_blocks(gemmi.cif.read(str(path))) + if cif_block is not None: + blocks = [b for b in blocks if b.block.name == cif_block] + if not blocks: + raise ValueError(f"No reflection block found in {path}") + mtz = gemmi.CifToMtz().convert_block_to_mtz(blocks[0]) + return rs.io.from_gemmi(mtz) + raise ValueError(f"Unsupported structure-factor format: {path}") + + +def write_sf_file(ds: rs.DataSet, path: str) -> None: + """Write a DataSet as MTZ or SF-mmCIF depending on the extension. + + CIF output uses gemmi's MTZ-to-mmCIF conversion, with ``FreeR_flag == 0`` + written as ``_refln.status 'f'`` and negative (excluded) flags as ``'x'``. + """ + suffix = Path(path).suffix.lower() + if suffix == ".mtz": + ds.write_mtz(str(path)) + elif suffix in (".cif", ".mmcif"): + converter = gemmi.MtzToCif() + converter.free_flag_value = 0 + text = converter.write_cif_to_string(ds.to_gemmi()) + if FREE_COLUMN in ds.columns: + text = _mark_excluded(text, ds) + Path(path).write_text(text) + else: + raise ValueError(f"Unsupported output format: {path}") + + +def _mark_excluded(text: str, ds: rs.DataSet) -> str: + """Set ``_refln.status`` to ``x`` for rows whose ``FreeR_flag`` is negative. + + gemmi writes every non-free flag as ``o``, which would turn excluded + reflections back into work reflections on the next read. + """ + flags = ds[FREE_COLUMN].to_numpy(dtype=float) + excluded = set(map(tuple, ds.get_hkls()[np.nan_to_num(flags, nan=0) < 0].tolist())) + if not excluded: + return text + doc = gemmi.cif.read_string(text) + for block in doc: + table = block.find("_refln.", ["index_h", "index_k", "index_l", "status"]) + for row in table: + if (int(row[0]), int(row[1]), int(row[2])) in excluded: + row[3] = "x" + return doc.as_string() + + +# --------------------------------------------------------------------------- +# Existing flags +# --------------------------------------------------------------------------- + + +def flag_column(ds: rs.DataSet, column: Optional[str] = None) -> Optional[str]: + """Name of the R-free column in ``ds`` (``column`` if given), or None.""" + if column is not None: + return column if column in ds.columns else None + return next((c for c in FLAG_COLUMN_NAMES if c in ds.columns), None) + + +def excluded_rows(ds: rs.DataSet, column: Optional[str] = None) -> np.ndarray: + """Rows marked excluded (negative or missing flag, CIF ``x``), shape (N,). + + All-false when ``ds`` has no R-free column. Independent of whether the + remaining rows form a valid partition, so a column that excludes every + row still excludes every row. + """ + column = flag_column(ds, column) + if column is None: + return np.zeros(len(ds), dtype=bool) + values = ds[column].to_numpy(dtype=float) + return ~np.isfinite(values) | (np.nan_to_num(values, nan=-1) < 0) + + +def read_free_set(ds: rs.DataSet, column: Optional[str] = None) -> dict: + """Interpret an existing R-free column row by row. + + Negative or missing values (MTZ ``-1``, CIF ``x``) are *excluded*. Among the + remaining rows, a column with more than two values is CCP4 ``0..K`` + (0 = free); a binary column takes its majority value as work, which covers + both CCP4 ``0 = free`` and Phenix ``1 = free``. + + Returns + ------- + dict + ``column``, ``convention``, ``raw`` (int, -1 where excluded), + ``free`` and ``excluded`` (bool per row). + """ + column = flag_column(ds, column) + if column is None: + raise ValueError("no recognised R-free column") + values = ds[column].to_numpy(dtype=float) + excluded = excluded_rows(ds, column) + raw = np.where(excluded, -1, np.nan_to_num(values, nan=-1)).astype(np.int64) + uvals, counts = np.unique(raw[~excluded], return_counts=True) + if len(uvals) == 0: + raise ValueError(f"R-free column {column!r} has no valid values") + if len(uvals) > 2: + convention = "ccp4" + free = raw == 0 + else: + work_value = uvals[np.argmax(counts)] + convention = f"binary (work={int(work_value)})" + free = ~excluded & (raw != work_value) + return { + "column": column, + "convention": convention, + "raw": raw, + "free": free, + "excluded": excluded, + } + + +def _group_free(keys: np.ndarray, free: np.ndarray): + """Reduce rows to unique ASU keys; returns (keys, free, conflicting). + + ``conflicting`` marks keys whose rows (symmetry equivalents, Friedel mates + or duplicates) are not all free or all work. + """ + ukeys, inverse = np.unique(keys, return_inverse=True) + n_free = np.bincount(inverse, weights=free, minlength=len(ukeys)) + n_rows = np.bincount(inverse, minlength=len(ukeys)) + return ukeys, n_free > 0, (n_free > 0) & (n_free < n_rows) + + +def compare_free_sets( + datasets: Dict[str, rs.DataSet], column: Optional[str] = None +) -> dict: + """Report whether the existing free sets of several datasets agree. + + Returns + ------- + dict + ``files``: name to ``column``/``convention``/``n``/``n_free``/ + ``n_excluded``/``n_conflicting``/``dmin`` (or ``column: None`` without + a usable free set, with ``problem`` saying why). ``n_conflicting`` + counts unique reflections whose equivalents within that file disagree. + ``pairs``: ``(a, b)`` to ``(n_common, n_disagree)`` over unique + reflections that are present, not excluded and not conflicting in + both. ``consistent`` requires every file to have a free set with free + reflections, no internal conflicts and no pairwise disagreement. + """ + files, tables = {}, {} + for name, ds in datasets.items(): + dmin = float(ds.compute_dHKL()["dHKL"].min()) + try: + fs = read_free_set(ds, column) + except ValueError as exc: + files[name] = {"column": None, "n": len(ds), "dmin": dmin, "problem": str(exc)} + continue + keep = ~fs["excluded"] + ukeys, ufree, conflict = _group_free(hkl_keys(asu_hkl(ds))[keep], fs["free"][keep]) + tables[name] = (ukeys[~conflict], ufree[~conflict]) + files[name] = { + "column": fs["column"], + "convention": fs["convention"], + "n": len(ds), + "n_free": int(fs["free"].sum()), + "n_excluded": int(fs["excluded"].sum()), + "n_conflicting": int(conflict.sum()), + "dmin": dmin, + } + pairs = {} + names = list(tables) + for i, a in enumerate(names): + for b in names[i + 1 :]: + ka, fa = tables[a] + kb, fb = tables[b] + common, ia, ib = np.intersect1d(ka, kb, assume_unique=True, return_indices=True) + pairs[(a, b)] = (len(common), int((fa[ia] != fb[ib]).sum())) + consistent = ( + len(tables) == len(datasets) + and all(files[n]["n_free"] > 0 and files[n]["n_conflicting"] == 0 for n in tables) + and all(d == 0 for _, d in pairs.values()) + ) + return {"files": files, "pairs": pairs, "consistent": consistent} + + +# --------------------------------------------------------------------------- +# Consistency +# --------------------------------------------------------------------------- + + +def check_compatible( + datasets: Dict[str, rs.DataSet], + length_tol: float = 0.01, + angle_tol: float = 0.5, +) -> List[str]: + """Check that all datasets share a space group and (nearly) a cell. + + Parameters + ---------- + datasets : dict + Name to DataSet. + length_tol : float + Allowed relative deviation of a, b, c from the first dataset. + angle_tol : float + Allowed absolute deviation of alpha, beta, gamma in degrees. + + Returns + ------- + list of str + Human-readable problems; empty when all datasets are compatible. + """ + names = list(datasets) + ref_name = names[0] + ref = datasets[ref_name] + ref_cell = np.array(ref.cell.parameters) + problems = [] + for name in names[1:]: + ds = datasets[name] + if ds.spacegroup.xhm() != ref.spacegroup.xhm(): + problems.append( + f"{name}: space group {ds.spacegroup.xhm()!r} != " + f"{ref.spacegroup.xhm()!r} ({ref_name})" + ) + cell = np.array(ds.cell.parameters) + rel = np.abs(cell[:3] - ref_cell[:3]) / ref_cell[:3] + dang = np.abs(cell[3:] - ref_cell[3:]) + if (rel > length_tol).any() or (dang > angle_tol).any(): + problems.append( + f"{name}: cell {tuple(np.round(cell, 3))} differs from " + f"{tuple(np.round(ref_cell, 3))} ({ref_name})" + ) + return problems + + +# --------------------------------------------------------------------------- +# Miller-index keys +# --------------------------------------------------------------------------- + + +def asu_hkl(ds: rs.DataSet) -> np.ndarray: + """Friedel-merged reciprocal-ASU indices for every row, shape (N, 3).""" + hkl = ds.get_hkls() + return rs.utils.hkl_to_asu(hkl, ds.spacegroup)[0].astype(np.int64) + + +def hkl_keys(hkl: np.ndarray) -> np.ndarray: + """Encode integer Miller indices (N, 3) as unique int64 scalars.""" + h = hkl.astype(np.int64) + _KEY_OFFSET + span = 2 * _KEY_OFFSET + return (h[:, 0] * span + h[:, 1]) * span + h[:, 2] + + +def _unkey(keys: np.ndarray) -> np.ndarray: + """Inverse of :func:`hkl_keys`.""" + span = 2 * _KEY_OFFSET + return np.stack([keys // span**2, keys // span % span, keys % span], 1) - _KEY_OFFSET + + +def _lookup(table_keys: np.ndarray, table_values: np.ndarray, keys: np.ndarray): + """Look up ``keys`` in sorted ``table_keys``; returns (values, found).""" + pos = np.searchsorted(table_keys, keys) + pos = np.clip(pos, 0, len(table_keys) - 1) + found = table_keys[pos] == keys + return np.where(found, table_values[pos], -1), found + + +# --------------------------------------------------------------------------- +# Flag generation +# --------------------------------------------------------------------------- + + +def resolution_bins(dstar2: np.ndarray, n_bins: int) -> np.ndarray: + """Equal-count resolution bin index (0..n_bins-1) for each 1/d^2 value.""" + n_bins = max(1, min(n_bins, len(dstar2))) + order = np.argsort(dstar2, kind="stable") + bins = np.empty(len(dstar2), dtype=np.int64) + bins[order] = np.arange(len(dstar2)) * n_bins // max(len(dstar2), 1) + return bins + + +def _hash_flags(hkl: np.ndarray, n_flags: int, seed: int) -> np.ndarray: + """Deterministic pseudo-random flag from the Miller index alone.""" + k = hkl_keys(hkl).astype(np.uint64) ^ np.uint64(seed * 0x9E3779B97F4A7C15 % 2**64) + k ^= k >> np.uint64(33) + k *= np.uint64(0xFF51AFD7ED558CCD) + k ^= k >> np.uint64(33) + return (k % np.uint64(n_flags)).astype(np.int32) + + +def partition_seed(keys: np.ndarray, free: np.ndarray) -> int: + """Seed derived from a free/work partition (SHA-256 of keys and free mask). + + Depends only on which unique reflections are free and which are work, not + on file format, row order or flag convention, so every extension of the + same deposited free set is identical. + """ + order = np.argsort(keys) + digest = hashlib.sha256( + np.ascontiguousarray(keys[order], dtype=" Tuple[np.ndarray, np.ndarray]: + """Stratified CCP4 flags on the complete reciprocal ASU out to ``dmin``. + + The complete ASU is sorted by resolution (ties by index) and cut into + consecutive shells of ``shell_size`` reflections. Each shell is shuffled + with its own seed ``(seed, shell)`` and dealt flag values round-robin, so + every value (in particular the free value 0) holds exactly + ``1 / n_flags`` of every shell. The last shell is always completed with + reflections beyond ``dmin``, so the flag of any reflection depends only on + cell, space group, ``n_flags``, ``shell_size`` and ``seed``: a larger + ``dmin`` (a later, better dataset) never changes existing flags. + + Returns + ------- + keys : np.ndarray + Sorted ASU keys (see :func:`hkl_keys`), shape (M,). + flags : np.ndarray + Flag per key, shape (M,). + """ + shell_size = max(n_flags, shell_size - shell_size % n_flags) + d = dmin + while True: + hkl = rs.utils.generate_reciprocal_asu(cell, spacegroup, d, anomalous=False) + hkl = hkl.astype(np.int64) + dstar2 = np.round(1.0 / rs.utils.compute_dHKL(hkl, cell) ** 2, 10) + n_needed = int((dstar2 <= 1.0 / dmin**2).sum()) + n_total = -(-n_needed // shell_size) * shell_size + if len(hkl) >= n_total: + break + d *= 0.95 + order = np.lexsort((hkl[:, 2], hkl[:, 1], hkl[:, 0], dstar2))[:n_total] + hkl = hkl[order] + flags = np.empty(len(hkl), dtype=np.int32) + deal = np.arange(shell_size) % n_flags + for shell in range(n_total // shell_size): + rng = np.random.default_rng([seed, shell]) + flags[shell * shell_size + rng.permutation(shell_size)] = deal + keys = hkl_keys(hkl) + idx = np.argsort(keys) + return keys[idx], flags[idx] + + +def reference_flags( + ref: rs.DataSet, + n_flags: int, + column: Optional[str] = None, + seed: int = 0, +) -> Tuple[np.ndarray, np.ndarray, dict]: + """Extract a CCP4-style flag per unique ASU reflection from a reference. + + Conventions are read by :func:`read_free_set`; excluded reflections are + skipped. Multi-valued columns (CCP4 ``0..K``) are kept as-is. For binary + columns free becomes ``0`` and work reflections are dealt pseudo-random + values ``1..n_flags-1``. + + Returns + ------- + keys : np.ndarray + Sorted unique ASU keys, shape (M,). + flags : np.ndarray + Flag per key, shape (M,). + info : dict + ``column``, ``convention``, ``n_inconsistent`` (unique reflections + whose symmetry equivalents carried different flags), ``n_values``, + ``free_fraction`` (per unique reflection) and ``dmin``. + """ + fs = read_free_set(ref, column) + column, convention = fs["column"], fs["convention"] + keep = ~fs["excluded"] + raw = fs["raw"][keep] + free_rows = fs["free"][keep] + ref_hkl = asu_hkl(ref)[keep] + keys = hkl_keys(ref_hkl) + ref_dmin = float(rs.utils.compute_dHKL(ref_hkl, ref.cell).min()) + + if not free_rows.any(): + raise ValueError( + f"reference column {column!r} ({convention}) marks no reflection as " + "free; choose another --reference or use --fresh" + ) + ukeys, inverse = np.unique(keys, return_inverse=True) + # a unique reflection is inconsistent if its equivalents disagree + lo = np.full(len(ukeys), np.iinfo(np.int64).max) + hi = np.full(len(ukeys), np.iinfo(np.int64).min) + np.minimum.at(lo, inverse, raw) + np.maximum.at(hi, inverse, raw) + n_inconsistent = int((lo != hi).sum()) + + if convention == "ccp4": + # conflicting equivalents: free (0) wins, else the smallest value + uflags = lo + else: + free_any = np.zeros(len(ukeys), dtype=bool) + np.logical_or.at(free_any, inverse, free_rows) + work = 1 + _hash_flags(_unkey(ukeys), n_flags - 1, seed) + uflags = np.where(free_any, 0, work) + info = { + "column": column, + "convention": convention, + "n_inconsistent": n_inconsistent, + "n_values": int(uflags.max()) + 1, + "free_fraction": float((uflags == 0).mean()), + "dmin": ref_dmin, + } + return ukeys, uflags.astype(np.int32), info + + +def uniform_rfree( + datasets: Dict[str, rs.DataSet], + free_fraction: Optional[float] = None, + shell_size: int = 1000, + seed: Optional[int] = None, + dmin: Optional[float] = None, + reference: Optional[rs.DataSet] = None, + reference_column: Optional[str] = None, + max_free: Optional[int] = None, + keep_excluded: bool = True, +) -> Tuple[Dict[str, np.ndarray], dict]: + """Compute one shared CCP4 ``FreeR_flag`` column for several datasets. + + Flags come from :func:`complete_flag_table` on the reference's (else the + first dataset's) cell, + so they depend only on cell, space group and settings, not on which + reflections happen to be measured. A dataset processed later with the + same settings therefore gets identical flags for every reflection. + + Parameters + ---------- + datasets : dict + Name to DataSet (same cell / space group). + free_fraction : float, optional + Target free fraction; the number of flag values is ``round(1/f)``. + Defaults to the reference's own fraction when inheriting, else 0.05. + shell_size : int + Reflections per stratification shell (see :func:`complete_flag_table`). + seed : int, optional + Random seed. By default ``0`` for a new set and, when extending a + reference, :func:`partition_seed` of the reference's free set, so the + flags assigned to reflections the reference lacks (e.g. its missing + high-resolution shells) are a deterministic function of the reference. + dmin : float, optional + High-resolution limit of the flag table; the best resolution of the + inputs is used if it is finer. + reference : rs.DataSet, optional + Dataset whose existing flags are inherited; reflections it lacks are + newly assigned. + reference_column : str, optional + Flag column in ``reference`` (auto-detected by default). + max_free : int, optional + Cap on the number of free reflections in the complete set to ``dmin`` + (Phenix-style); lowers the fraction for large datasets. The cap depends + on ``dmin``, so reuse the reported fraction to reproduce a flag set. + keep_excluded : bool + Rows a dataset itself marks as excluded (negative flag, CIF ``x``) + stay ``-1`` in that dataset's output. + + Returns + ------- + flags : dict + Name to per-row ``FreeR_flag`` array aligned with each DataSet. + info : dict + Summary: ``n_flags``, ``free_fraction``, ``fraction_source``, ``seed``, + ``seed_source``, ``dmin``, ``n_gaps_in_reference``, + ``n_unique``, ``n_off_asu``, ``n_inherited``, ``n_generated``, + ``n_generated_beyond_reference``, ``n_excluded`` (per file) and the + reference ``info``. + """ + if free_fraction is not None and not 0 < free_fraction < 1: + raise ValueError("free_fraction must be between 0 and 1") + + # the flag table lives on the reference's cell when there is one, so the + # result does not depend on input order + first = reference if reference is not None else next(iter(datasets.values())) + cell, sg = first.cell, first.spacegroup + row_hkl = {name: asu_hkl(ds) for name, ds in datasets.items()} + observed = np.unique(np.concatenate(list(row_hkl.values())), axis=0) + data_dmin = float(rs.utils.compute_dHKL(observed, cell).min()) + dmin = min(dmin, data_dmin) if dmin is not None else data_dmin + + rkeys = rflags = rinfo = None + if reference is not None: + # n_flags and seed only affect the pseudo-random work values here + rkeys, rflags, rinfo = reference_flags( + reference, 20, column=reference_column, seed=0 + ) + if seed is not None: + seed_source = "user" + elif rinfo is not None: + seed, seed_source = partition_seed(rkeys, rflags == 0), "reference free set" + else: + seed, seed_source = 0, "default" + if free_fraction is not None: + n_flags, source = max(2, int(round(1.0 / free_fraction))), "user" + elif rinfo is not None and rinfo["convention"] == "ccp4": + n_flags, source = max(2, rinfo["n_values"]), "reference" + elif rinfo is not None: + n_flags, source = max(2, int(round(1.0 / rinfo["free_fraction"]))), "reference" + else: + n_flags, source = 20, "default" + if max_free is not None: + n_complete = len(rs.utils.generate_reciprocal_asu(cell, sg, dmin, anomalous=False)) + if n_complete / n_flags > max_free: + n_flags, source = int(np.ceil(n_complete / max_free)), f"max_free={max_free}" + if rinfo is not None and rinfo["convention"] != "ccp4": + rkeys, rflags, rinfo = reference_flags( + reference, n_flags, column=reference_column, seed=seed + ) + ukeys, flags = complete_flag_table(cell, sg, dmin, n_flags, seed, shell_size) + + # observed indices outside the complete set (e.g. systematic absences) + okeys = hkl_keys(observed) + _, found = _lookup(ukeys, flags, okeys) + if not found.all(): + ukeys = np.concatenate([ukeys, okeys[~found]]) + flags = np.concatenate([flags, _hash_flags(observed[~found], n_flags, seed)]) + idx = np.argsort(ukeys) + ukeys, flags = ukeys[idx], flags[idx] + info = { + "n_flags": n_flags, + "free_fraction": 1.0 / n_flags, + "fraction_source": source, + "seed": seed, + "seed_source": seed_source, + "dmin": dmin, + "n_unique": len(ukeys), + "n_off_asu": int((~found).sum()), + } + + if rinfo is not None: + inherited, found = _lookup(rkeys, rflags, ukeys) + flags = np.where(found, inherited, flags).astype(np.int32) + # restrict counts to reflections some input actually has + seen = np.isin(ukeys, okeys) + d = rs.utils.compute_dHKL(_unkey(ukeys[seen & ~found]), cell) + n_beyond = int((d < rinfo["dmin"] - 1e-6).sum()) + info.update( + reference=rinfo, + n_inherited=int((found & seen).sum()), + n_generated=int((~found & seen).sum()), + n_generated_beyond_reference=n_beyond, + n_gaps_in_reference=int((~found & seen).sum()) - n_beyond, + ) + else: + info.update( + n_inherited=0, + n_generated=len(okeys), + n_generated_beyond_reference=0, + n_gaps_in_reference=0, + ) + + per_file, n_excluded = {}, {} + for name, hkl in row_hkl.items(): + values, found = _lookup(ukeys, flags, hkl_keys(hkl)) + assert found.all(), "every observed reflection is in the universe" + values = values.astype(np.int32) + if keep_excluded: + excluded = excluded_rows(datasets[name]) + else: + excluded = np.zeros(len(values), dtype=bool) + values[excluded] = -1 + n_excluded[name] = int(excluded.sum()) + per_file[name] = values + info["n_excluded"] = n_excluded + return per_file, info + + +def apply_flags( + ds: rs.DataSet, flags: np.ndarray, keep_old: bool = False +) -> rs.DataSet: + """Return a copy of ``ds`` whose only R-free column is ``FreeR_flag``. + + Existing flag columns are dropped, or renamed ``_orig`` with + ``keep_old``. Row order and all other columns are unchanged. + """ + out = ds.copy() + for col in [c for c in out.columns if c in FLAG_COLUMN_NAMES]: + if keep_old: + out = out.rename(columns={col: f"{col}_orig"}) + else: + out = out.drop(columns=col) + out[FREE_COLUMN] = rs.DataSeries(flags, index=out.index, dtype="I") + return out + + +# --------------------------------------------------------------------------- +# Scale application +# --------------------------------------------------------------------------- + +_AMPLITUDE_TYPES = {"F", "G", "D", "L"} # F, F(+/-), anomalous diff, sigma F(+/-) +_INTENSITY_TYPES = {"J", "K", "M"} # I, I(+/-), sigma I(+/-) + + +_CALC_PREFIXES = ("FC", "FCALC", "FMODEL", "F-MODEL", "FCAL", "FWT", "DELFWT", "2FOFC", "FOFC") + + +def _is_calculated(name: str) -> bool: + """Amplitude column names that hold model or map values, not observations.""" + n = name.upper() + return n.startswith(_CALC_PREFIXES) or n.endswith("WT") + + +def scale_columns(ds: rs.DataSet, factor: np.ndarray) -> Tuple[rs.DataSet, List[str]]: + """Multiply amplitude columns by ``factor`` and intensity columns by its square. + + Generic standard deviations (MTZ type ``Q``) follow the preceding column + or their ``SIG`` partner. Map coefficients and calculated amplitudes + (an amplitude directly followed by a phase column, e.g. ``FWT``/``PHWT``, + or a ``FC``/``FCALC``/``FMODEL``-style name), phases, weights, flags and + other columns are untouched. + + Returns + ------- + rs.DataSet + Scaled copy. + list of str + Names of the columns that were scaled. + """ + out = ds.copy() + columns = list(out.columns) + types = {c: out[c].dtype.mtztype for c in columns} + scaled = [] + prev = None + for i, col in enumerate(columns): + t = types[col] + power = None + nxt = columns[i + 1] if i + 1 < len(columns) else None + if t == "F" and (types.get(nxt) == "P" or _is_calculated(col)): + pass # map coefficient or model amplitude + elif t in _AMPLITUDE_TYPES: + power = 1 + elif t in _INTENSITY_TYPES: + power = 2 + elif t == "Q": + partner = col[3:] if col.upper().startswith("SIG") else None + ref_type = types.get(partner) or (types.get(prev) if prev else None) + if ref_type in _AMPLITUDE_TYPES: + power = 1 + elif ref_type in _INTENSITY_TYPES: + power = 2 + if power is not None: + dtype = out[col].dtype + out[col] = rs.DataSeries( + out[col].to_numpy(dtype=float) * factor**power, + index=out.index, + dtype=dtype, + ) + scaled.append(col) + prev = col + return out, scaled From 81c53f499efdf64b4440d99e2e5b310562997f1f Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 21:38:38 +0200 Subject: [PATCH 204/250] Rename difference MTZ columns to dF/SIGdF and the inverse-variance weight to W_InVa The amplitude differences use a lower-case delta (DF -> dF, SIGDF -> SIGdF, DF_corr -> dF_corr, DDF -> ddF, DFc -> dFc, DFc_phased -> dFc_phased) and the inverse-variance weight column W_IVW is W_InVa. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/user_guide/cli.rst | 12 ++--- tests/integration/test_cli_ded_weights.py | 16 +++--- tests/integration/test_cli_two_moment_mtz.py | 40 +++++++-------- torchref/cli/collection_difference_refine.py | 54 ++++++++++---------- torchref/cli/difference_map.py | 6 +-- torchref/cli/mtz2map.py | 2 +- torchref/cli/validate_ded.py | 12 ++--- torchref/io/datasets/reflection_data.py | 2 +- torchref/io/mtz.py | 2 +- torchref/maps/ded_weights.py | 2 +- torchref/maps/difference_map.py | 4 +- 11 files changed, 76 insertions(+), 76 deletions(-) diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index eadaa72d..caa7d48a 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -109,8 +109,8 @@ columns, expands to P1, and computes a real-space map via FFT. .. code-block:: bash torchref.mtz2map -sf refined.mtz -csf 2FOFCWT -cphi PH2FOFCWT -o map.ccp4 - torchref.mtz2map -sf diff.mtz -csf DF -cw W_IVW -cphi PHDELWT -o diff.ccp4 - torchref.mtz2map -sf diff.mtz -csf DF -cw W_SD -cphi PHDELWT --units electrons -o diff_e.ccp4 + torchref.mtz2map -sf diff.mtz -csf dF -cw W_InVa -cphi PHDELWT -o diff.ccp4 + torchref.mtz2map -sf diff.mtz -csf dF -cw W_SD -cphi PHDELWT --units electrons -o diff_e.ccp4 **Key options:** ``--dmin``/``--dmax`` resolution limits, ``--gridsize`` override, ``-cw``/``--column-weight`` multiplies the amplitudes by a weight column before the @@ -125,7 +125,7 @@ alias of ``--units sigma``/``raw``. ``torchref.validate-ded`` ~~~~~~~~~~~~~~~~~~~~~~~~~ -Validate difference electron density by correlating DFo and DFc maps. +Validate difference electron density by correlating dFo and dFc maps. Computes real-space correlations and resolution-binned reciprocal-space CC. .. code-block:: bash @@ -150,13 +150,13 @@ Compute difference and extrapolated map coefficients without refinement. Uses the same pipeline as ``torchref.difference-refine`` but the input models are kept as-is. -The default output is the difference map: the amplitude difference ``DF``/``SIGDF`` +The default output is the difference map: the amplitude difference ``dF``/``SIGdF`` on the **dark** model's phases ``PHDELWT``, with one mean-one weight column per -registered scheme beside it -- ``W_IVW``, the inverse variance ``1/sigma^2`` (the +registered scheme beside it -- ``W_InVa``, the inverse variance ``1/sigma^2`` (the default), and ``W_SD``, the sigma_D Wiener weight ``S/(S + sigma^2)`` built from the expected difference power -- and ``KSCALE``, the scaler's factor from model to observed scale. This is the construction ``torchref.validate-ded`` correlates against. Build the -map with ``torchref.mtz2map -csf DF -cw W_IVW -cphi PHDELWT``, adding +map with ``torchref.mtz2map -csf dF -cw W_InVa -cphi PHDELWT``, adding ``--units electrons`` for e/A^3. It needs no light-state model, so ``-lm`` is optional: .. code-block:: bash diff --git a/tests/integration/test_cli_ded_weights.py b/tests/integration/test_cli_ded_weights.py index 8679564b..116a3b40 100644 --- a/tests/integration/test_cli_ded_weights.py +++ b/tests/integration/test_cli_ded_weights.py @@ -1,6 +1,6 @@ """The registered difference weights through the CLIs. -Pinned: ``torchref.difference-map`` writes ``DF`` with one mean-one weight column per +Pinned: ``torchref.difference-map`` writes ``dF`` with one mean-one weight column per scheme and ``KSCALE``; ``torchref.mtz2map`` builds the weighted map from those columns and the electrons map is the volume-normalised synthesis on the absolute scale; ``torchref.validate-ded`` reports every scheme side by side and records a fallback; @@ -22,10 +22,10 @@ "SIGFo_dark": "Stddev", "Fo_light": "SFAmplitude", "SIGFo_light": "Stddev", - "DF": "SFAmplitude", - "SIGDF": "Stddev", + "dF": "SFAmplitude", + "SIGdF": "Stddev", "PHDELWT": "Phase", - "W_IVW": "Weight", + "W_InVa": "Weight", "W_SD": "Weight", "KSCALE": "MTZReal", "Fc_dark": "SFAmplitude", @@ -138,7 +138,7 @@ def _read(path): def test_difference_map_writes_df_weights_and_scale(diff_mtz): df = _read(diff_mtz) assert {c: str(df.dtypes[c]) for c in df.columns} == DIFF_COLUMNS - for col in ("W_IVW", "W_SD"): + for col in ("W_InVa", "W_SD"): w = df[col].to_numpy().astype(float) assert np.isfinite(w).all() and (w >= 0).all() assert abs(w.mean() - 1.0) < 1e-4 @@ -160,7 +160,7 @@ def test_mtz2map_builds_the_weighted_and_electron_maps(project_root, pair, diff_ "-sf", diff_mtz, "-csf", - "DF", + "dF", "-cw", "W_SD", "-cphi", @@ -176,7 +176,7 @@ def test_mtz2map_builds_the_weighted_and_electron_maps(project_root, pair, diff_ "-sf", diff_mtz, "-csf", - "DF", + "dF", "-cw", "W_SD", "-cphi", @@ -194,7 +194,7 @@ def test_mtz2map_builds_the_weighted_and_electron_maps(project_root, pair, diff_ "-sf", diff_mtz, "-csf", - "DF", + "dF", "-cw", "W_SD", "-cphi", diff --git a/tests/integration/test_cli_two_moment_mtz.py b/tests/integration/test_cli_two_moment_mtz.py index 01270bae..99ec6703 100644 --- a/tests/integration/test_cli_two_moment_mtz.py +++ b/tests/integration/test_cli_two_moment_mtz.py @@ -17,12 +17,12 @@ "SIGFo_dark": "Stddev", "Fo_light": "SFAmplitude", "SIGFo_light": "Stddev", - "DF": "SFAmplitude", - "SIGDF": "Stddev", - # The difference map: DF on the dark phases, one weight column per registered + "dF": "SFAmplitude", + "SIGdF": "Stddev", + # The difference map: dF on the dark phases, one weight column per registered # scheme, and the observed-to-model scale. "PHDELWT": "Phase", - "W_IVW": "Weight", + "W_InVa": "Weight", "W_SD": "Weight", "KSCALE": "MTZReal", "Fc_dark": "SFAmplitude", @@ -40,9 +40,9 @@ "DELFWT_corr": "SFAmplitude", "Fo_light_corr": "SFAmplitude", "SIGFo_light_corr": "Stddev", - "DF_corr": "SFAmplitude", - "SIGDF_corr": "Stddev", - "DDF": "SFAmplitude", + "dF_corr": "SFAmplitude", + "SIGdF_corr": "Stddev", + "ddF": "SFAmplitude", } # What ``--all-columns`` adds on top, given a light model. @@ -50,8 +50,8 @@ "2mDFop-DFc": "SFAmplitude", "mDFop-DFc": "SFAmplitude", "PHIC_diff": "Phase", - "DFc": "SFAmplitude", - "DFc_phased": "SFAmplitude", + "dFc": "SFAmplitude", + "dFc_phased": "SFAmplitude", "FEXT_PHASED": "SFAmplitude", "SIGFEXT_PHASED": "Stddev", "2FEXT_PHASED-Fc": "SFAmplitude", @@ -229,9 +229,9 @@ def test_map_coefficients_are_grouped_into_described_datasets(two_moment_all_mtz assert where["FWT"] == where["PHWT"] == where["FEXT"] == "extrapolated_light" assert where["FEXT_PHASED"] == where["FEXT_SCALAR"] == "extrapolated_light" - assert where["DF"] == where["PHDELWT"] == where["W_SD"] == "difference" + assert where["dF"] == where["PHDELWT"] == where["W_SD"] == "difference" assert where["FC"] == where["PHIC"] == "light_model" - assert where["DF_corr"] == "two_moment" + assert where["dF_corr"] == "two_moment" assert where["H"] == where["Fo_dark"] == where["FreeR_flag_dark"] == "observed" assert all(ds.crystal_name == "torchref" for ds in mtz.datasets) @@ -245,7 +245,7 @@ class TestDefaultLayout: def test_the_difference_columns_carry_df_and_the_registered_weights( self, baseline_mtz ): - """``DF`` must be ``Fo_light - Fo_dark``, ``W_IVW`` the mean-normalised inverse + """``dF`` must be ``Fo_light - Fo_dark``, ``W_InVa`` the mean-normalised inverse variance and ``W_SD`` a mean-one weight -- the constructions ``torchref.validate-ded`` correlates against. If these ever diverge, the map built from the file stops being the map the validation reports on, which is how @@ -263,10 +263,10 @@ def test_the_difference_columns_carry_df_and_the_registered_weights( w = 1 / np.maximum(sig, 0.1 * np.median(sig)) ** 2 w = w / w.mean() - got_df = df["DF"].to_numpy().astype(float) + got_df = df["dF"].to_numpy().astype(float) scale = max(float(np.abs(dfo).max()), 1e-30) assert np.abs(got_df - dfo).max() / scale < 1e-5 - got_w = df["W_IVW"].to_numpy().astype(float) + got_w = df["W_InVa"].to_numpy().astype(float) assert np.abs(got_w - w).max() / max(float(np.abs(w).max()), 1e-30) < 1e-4 w_sd = df["W_SD"].to_numpy().astype(float) assert np.isfinite(w_sd).all() and abs(w_sd.mean() - 1.0) < 1e-4 @@ -286,13 +286,13 @@ def test_the_corrected_difference_map_pairs_with_the_same_phases( ): """``DELFWT_corr`` is the corrected difference on the *same* dark phases, so it is opened against ``PHDELWT`` and carries the selected weight scheme, the - inverse-variance weights ``W_IVW`` by default.""" + inverse-variance weights ``W_InVa`` by default.""" import numpy as np df = _read(two_moment_mtz[0]) - w = df["W_IVW"].to_numpy().astype(float) + w = df["W_InVa"].to_numpy().astype(float) - expected = df["DF_corr"].to_numpy().astype(float) * w + expected = df["dF_corr"].to_numpy().astype(float) * w got = df["DELFWT_corr"].to_numpy().astype(float) scale = max(float(np.abs(expected).max()), 1e-30) assert np.abs(got - expected).max() / scale < 1e-5 @@ -310,7 +310,7 @@ def test_ivar_alpha_is_sigma_sq_times_the_squared_difference( results = json.loads(summary.read_text())["results"] sigma_sq = results["sigma_alpha_sq"] - dfc = df["DFc_phased"].to_numpy().astype(float) + dfc = df["dFc_phased"].to_numpy().astype(float) ivar = df["IVAR_ALPHA"].to_numpy().astype(float) expected = sigma_sq * dfc**2 @@ -358,12 +358,12 @@ def test_the_weight_is_the_contamination_ratio(self, two_moment_all_mtz): assert (ivar > 0).any() def test_the_correction_moves_the_difference_amplitudes(self, two_moment_mtz): - """DDF is the diagnostic; if it were identically zero the whole column set + """ddF is the diagnostic; if it were identically zero the whole column set would be decorative.""" import numpy as np df = _read(two_moment_mtz[0]) - ddf = df["DDF"].to_numpy().astype(float) + ddf = df["ddF"].to_numpy().astype(float) assert np.count_nonzero(ddf) > 0.5 * len(ddf) # Subtracting a positive contamination lowers the light amplitude on average. assert ddf.mean() < 0.0 diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index a6b36909..1517e552 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -429,7 +429,7 @@ def _two_moment_columns(mc, dc, mask, fcalc_dark_full, fcalc_mixed_full, the model's estimate of it and converting back to an amplitude gives a difference amplitude that is comparable across datasets, which the raw one is not. - ``DDF`` is the diagnostic that matters: smooth and featureless against resolution + ``ddF`` is the diagnostic that matters: smooth and featureless against resolution means the correction is collinear with a scale or overall-B error and should be distrusted; structure in it is the signal. @@ -495,9 +495,9 @@ def _np(t): sig_F_corr = _np(sig_F_corr_full) I_two_moment = I_coherent + variance - DF_corr = F_corr - Fobs_dark - DDF = DF_corr - diff_Fobs - sig_DF_corr = np.sqrt(sig_F_corr**2 + sig_dark**2) + dF_corr = F_corr - Fobs_dark + ddF = dF_corr - diff_Fobs + sig_dF_corr = np.sqrt(sig_F_corr**2 + sig_dark**2) # Use the modulus of the complex vector difference so the corrected # coefficient retains the phase rotation between the dark and light states. @@ -517,20 +517,20 @@ def _np(t): columns = { # The corrected difference map, on the same dark phases as DELFWT. - "DELFWT_corr": DF_corr * weights, + "DELFWT_corr": dF_corr * weights, "Fo_light_corr": F_corr, "SIGFo_light_corr": sig_F_corr, - "DF_corr": DF_corr, - "SIGDF_corr": sig_DF_corr, - "DDF": DDF, + "dF_corr": dF_corr, + "SIGdF_corr": sig_dF_corr, + "ddF": ddF, } types = { "DELFWT_corr": "F", "Fo_light_corr": "F", "SIGFo_light_corr": "Q", - "DF_corr": "F", - "SIGDF_corr": "Q", - "DDF": "F", + "dF_corr": "F", + "SIGdF_corr": "Q", + "ddF": "F", } if all_columns: columns.update({ @@ -575,14 +575,14 @@ def _difference_columns( ): """The difference map's amplitudes, phases and weights. - ``DF``/``SIGDF`` is the signed amplitude difference ``|Fo_light| - |Fo_dark|`` with + ``dF``/``SIGdF`` is the signed amplitude difference ``|Fo_light| - |Fo_dark|`` with its propagated uncertainty, ``PHDELWT`` the **dark** model's phase it is carried on: the isomorphous difference Fourier, and the construction ``torchref.validate-ded`` - correlates against. One weight column per registered scheme (``W_SD``, ``W_IVW``; - MTZ type ``W``, mean one) sits beside it, so any weighting is ``DF`` times a column - and reproducible from the file: ``torchref.mtz2map -csf DF -cw W_SD -cphi PHDELWT``. + correlates against. One weight column per registered scheme (``W_SD``, ``W_InVa``; + MTZ type ``W``, mean one) sits beside it, so any weighting is ``dF`` times a column + and reproducible from the file: ``torchref.mtz2map -csf dF -cw W_SD -cphi PHDELWT``. ``KSCALE`` (type ``R``) is the scaler's multiplicative factor from model to observed - scale, so ``DF / KSCALE`` is in electrons and ``mtz2map --units electrons`` gives + scale, so ``dF / KSCALE`` is in electrons and ``mtz2map --units electrons`` gives e/A^3. This layer needs no light-state model: the amplitude is ``|Fo_light| - |Fo_dark|`` @@ -607,8 +607,8 @@ def _flags(data): "SIGFo_dark": sig_dark, "Fo_light": Fobs_light, "SIGFo_light": sig_light, - "DF": diff_Fobs, - "SIGDF": sig_diff, + "dF": diff_Fobs, + "SIGdF": sig_diff, "PHDELWT": phases_dark, **weight_columns, "KSCALE": kscale, @@ -626,8 +626,8 @@ def _flags(data): "SIGFo_dark": "Q", "Fo_light": "F", "SIGFo_light": "Q", - "DF": "F", - "SIGDF": "Q", + "dF": "F", + "SIGdF": "Q", "PHDELWT": "P", **{name: "W" for name in weight_columns}, "KSCALE": "R", @@ -688,14 +688,14 @@ def _phasing_columns(mc, scaler, hkl_all, mask, *, fcalc_dark, Fobs_dark_vals, "2mDFop-DFc": (2 * Fobs_diff_phased - Fcalc_diff_amp) * weights, "mDFop-DFc": (Fobs_diff_phased - Fcalc_diff_amp) * weights, "PHIC_diff": torch.angle(fcalc_diff).detach().rad2deg().cpu().numpy(), - "DFc": Fcalc_light - Fcalc_dark, + "dFc": Fcalc_light - Fcalc_dark, # This column holds the real modulus of the complex vector difference. - "DFc_phased": Fcalc_diff_amp, + "dFc_phased": Fcalc_diff_amp, } ) types.update({ "2mDFop-DFc": "F", "mDFop-DFc": "F", "PHIC_diff": "P", - "DFc": "F", "DFc_phased": "F", + "dFc": "F", "dFc_phased": "F", }) ctx = { @@ -849,14 +849,14 @@ def _np(t): _DIFFERENCE_DATASET_COLUMNS = ( - "DF", "SIGDF", "PHDELWT", "KSCALE", *WEIGHT_COLUMNS.values() + "dF", "SIGdF", "PHDELWT", "KSCALE", *WEIGHT_COLUMNS.values() ) # One history line per MTZ dataset, in the order they are written. MTZ history lines # are at most 80 characters. _MTZ_DATASET_HISTORY = { "observed": "observed: Fo_dark, Fo_light and flags on the shared scale; Fc_dark", - "difference": "difference: DF/SIGDF on dark phases PHDELWT; weights W_SD, W_IVW", + "difference": "difference: dF/SIGdF on dark phases PHDELWT; weights W_SD, W_InVa", "light_model": ( "light_model: FC/PHIC, amplitude and phase of the mixed dark+light model" ), @@ -929,9 +929,9 @@ def write_results_mtz( ): """Write the difference map, and map coefficients when a light model is given. - The default output is the **difference map**: ``DF``/``SIGDF`` on the dark model's + The default output is the **difference map**: ``dF``/``SIGdF`` on the dark model's phases ``PHDELWT``, with one mean-one weight column per registered scheme - (``W_SD``, ``W_IVW``) and the observed-to-model scale ``KSCALE``; see + (``W_SD``, ``W_InVa``) and the observed-to-model scale ``KSCALE``; see :func:`_difference_columns`. ``ded_weight`` selects the scheme the model-phased difference columns and the two-moment columns are weighted with. That needs no light-state model, which is why ``mc`` is optional -- with a dark model alone this diff --git a/torchref/cli/difference_map.py b/torchref/cli/difference_map.py index 54f4bc6c..0d4ac24e 100644 --- a/torchref/cli/difference_map.py +++ b/torchref/cli/difference_map.py @@ -6,11 +6,11 @@ models are used as-is. The default output is the difference map: the amplitude difference -``|Fo_light| - |Fo_dark|`` as ``DF``/``SIGDF`` carried on the **dark** model's phases +``|Fo_light| - |Fo_dark|`` as ``dF``/``SIGdF`` carried on the **dark** model's phases ``PHDELWT``, with the per-reflection weights of every registered scheme beside it as -``W_IVW`` (inverse variance, the default) and ``W_SD`` (sigma_D Wiener weight), and the +``W_InVa`` (inverse variance, the default) and ``W_SD`` (sigma_D Wiener weight), and the observed-to-model scale ``KSCALE``. Build the map with -``torchref.mtz2map -csf DF -cw W_IVW -cphi PHDELWT`` (``--units electrons`` for e/A^3). +``torchref.mtz2map -csf dF -cw W_InVa -cphi PHDELWT`` (``--units electrons`` for e/A^3). That needs no light-state model, so ``-lm`` is optional. It is also deliberately not a *phased* difference map: putting the light state's model phases into the observed amplitude biases the map toward the very model the experiment is testing. diff --git a/torchref/cli/mtz2map.py b/torchref/cli/mtz2map.py index e98fe12b..45c75e1a 100644 --- a/torchref/cli/mtz2map.py +++ b/torchref/cli/mtz2map.py @@ -66,7 +66,7 @@ def main(): type=str, metavar="COL", help="Weight column multiplied into the amplitudes before the FFT " - "(e.g. W_SD, W_IVW from torchref.difference-map). Default: none.", + "(e.g. W_SD, W_InVa from torchref.difference-map). Default: none.", ) inp.add_argument( "-ck", diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index 18b6d9a7..7614f104 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -1,8 +1,8 @@ #!/usr/bin/env python3 -u -"""Validate difference electron density (DED) by correlating DFo and DFc maps. +"""Validate difference electron density (DED) by correlating dFo and dFc maps. Takes separate dark and light MTZ files, computes weighted difference amplitudes -internally, then compares the weighted DFo and DFcalc maps using dark-state phases. +internally, then compares the weighted dFo and dFcalc maps using dark-state phases. Phenix-style atom selections give regional correlations, e.g. around a ligand site. Examples @@ -209,7 +209,7 @@ def setup_ded_context( ): """Load reflection data and prepare shared state for DED validation. - This sets up the observation side (weighted DFo, P1 expansion, resolution + This sets up the observation side (weighted dFo, P1 expansion, resolution bins, free/work masks) that is independent of any particular model. Parameters @@ -451,12 +451,12 @@ def compute_ded_maps( w_delta_fcalc = delta_fcalc * ctx["weights_p1"] phi_dark_p1 = torch.angle(fcalc_dark_p1) - # ASU-level weighted DFcalc + # ASU-level weighted dFcalc delta_fcalc_asu = fcalc_mixed_asu.abs() - fcalc_dark_asu.abs() w_delta_fcalc_asu = delta_fcalc_asu * ctx["weights"] if verbose >= 1: - print(f" |DFcalc| mean: {delta_fcalc.abs().mean():.3f}") + print(f" |dFcalc| mean: {delta_fcalc.abs().mean():.3f}") print(f" |WDFcalc| mean: {w_delta_fcalc.abs().mean():.3f}") # Compute maps @@ -820,7 +820,7 @@ def main(): """Entry point for ``torchref.validate-ded``; returns the exit code.""" parser = argparse.ArgumentParser( description="Validate difference electron density by correlating " - "weighted DFo and DFcalc maps.", + "weighted dFo and dFcalc maps.", formatter_class=argparse.RawDescriptionHelpFormatter, epilog=""" Examples: diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index 58722edc..0b873191 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -1009,7 +1009,7 @@ def load_mtz( column_names : dict, optional Explicit column name mapping to override automatic detection. Supported keys: ``"F"``, ``"SIGF"``, ``"I"``, ``"SIGI"``. - Example: ``{"F": "DFo", "SIGF": "sig_DFo"}``. + Example: ``{"F": "dFo", "SIGF": "sig_dFo"}``. french_wilson : bool, optional Whether to derive amplitudes from intensities via French-Wilson. Default True. Set False to use existing French-Wilson-corrected diff --git a/torchref/io/mtz.py b/torchref/io/mtz.py index ba11fbc7..13cb1b5e 100644 --- a/torchref/io/mtz.py +++ b/torchref/io/mtz.py @@ -118,7 +118,7 @@ def __init__( column_names : dict, optional Explicit column name mapping to override automatic detection. Supported keys: ``"F"``, ``"SIGF"``, ``"I"``, ``"SIGI"``. - Example: ``{"F": "DFo", "SIGF": "sig_DFo"}``. + Example: ``{"F": "dFo", "SIGF": "sig_dFo"}``. anomalous : bool, optional None (default) stacks ``F(+)/F(-)`` (or ``I(+)/I(-)``) into explicit Friedel pairs when such columns exist; True forces that (warning if diff --git a/torchref/maps/ded_weights.py b/torchref/maps/ded_weights.py index e09a93b8..258831e6 100644 --- a/torchref/maps/ded_weights.py +++ b/torchref/maps/ded_weights.py @@ -44,7 +44,7 @@ #: rose in the bulk solvent as much as in the region of interest. DEFAULT_SCHEME = "inverse_variance" #: MTZ column carrying each scheme's weight (type ``W``); ``none`` writes no column. -WEIGHT_COLUMNS = {"inverse_variance": "W_IVW", "sigma_d": "W_SD"} +WEIGHT_COLUMNS = {"inverse_variance": "W_InVa", "sigma_d": "W_SD"} class DedWeightFallbackWarning(UserWarning): diff --git a/torchref/maps/difference_map.py b/torchref/maps/difference_map.py index 4dc010ea..22f84b4b 100644 --- a/torchref/maps/difference_map.py +++ b/torchref/maps/difference_map.py @@ -1,7 +1,7 @@ """ Isomorphous difference map from two datasets. -Computes a difference Fourier map using DF = F_data - F_reference with +Computes a difference Fourier map using dF = F_data - F_reference with phases from a model, after scaling both datasets to a common reference. """ @@ -23,7 +23,7 @@ class DifferenceMap(Map): Scales both datasets to a common reference using ``DatasetCollection``, then computes difference Fourier coefficients: - ``DF * exp(i * phi_calc)`` where ``DF = F_data - F_reference``. + ``dF * exp(i * phi_calc)`` where ``dF = F_data - F_reference``. Parameters ---------- From 751c2f21d49b5e2c7ba87a89d81e1313e1e1a034 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 22:03:34 +0200 Subject: [PATCH 205/250] Replace the shell sigma_D difference weight with a shell-free q-weight fit_difference_power fits the expected difference power by per-reflection maximum likelihood: log S as a Chebyshev series in sin(theta)/lambda plus gamma log F_dark, with a fitted sigma scale k and centric factor. No resolution shells, no clamped moments. The q scheme (column W_Q, now the default) is the Wiener weight (snr + b) / (snr + 1 + b), whose floor keeps every reflection at b/(1+b) of the full weight, so a noisy resolution range is down-weighted rather than removed. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/user_guide/cli.rst | 11 +- tests/integration/test_cli_ded_weights.py | 35 +- tests/integration/test_cli_two_moment_mtz.py | 14 +- tests/unit/maps/test_ded_weights.py | 58 ++- .../unit/refinement/test_difference_power.py | 92 +++++ torchref/cli/_common.py | 15 +- torchref/cli/collection_difference_refine.py | 48 +-- torchref/cli/difference_map.py | 4 +- torchref/cli/mtz2map.py | 2 +- torchref/cli/validate_ded.py | 12 +- torchref/maps/ded_weights.py | 162 ++++---- .../difference_power.py | 383 ++++++++++++++++++ 12 files changed, 660 insertions(+), 176 deletions(-) create mode 100644 tests/unit/refinement/test_difference_power.py create mode 100644 torchref/refinement/model_error_estimation/difference_power.py diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index caa7d48a..df43aab0 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -109,8 +109,8 @@ columns, expands to P1, and computes a real-space map via FFT. .. code-block:: bash torchref.mtz2map -sf refined.mtz -csf 2FOFCWT -cphi PH2FOFCWT -o map.ccp4 - torchref.mtz2map -sf diff.mtz -csf dF -cw W_InVa -cphi PHDELWT -o diff.ccp4 - torchref.mtz2map -sf diff.mtz -csf dF -cw W_SD -cphi PHDELWT --units electrons -o diff_e.ccp4 + torchref.mtz2map -sf diff.mtz -csf dF -cw W_Q -cphi PHDELWT -o diff.ccp4 + torchref.mtz2map -sf diff.mtz -csf dF -cw W_Q -cphi PHDELWT --units electrons -o diff_e.ccp4 **Key options:** ``--dmin``/``--dmax`` resolution limits, ``--gridsize`` override, ``-cw``/``--column-weight`` multiplies the amplitudes by a weight column before the @@ -152,11 +152,10 @@ models are kept as-is. The default output is the difference map: the amplitude difference ``dF``/``SIGdF`` on the **dark** model's phases ``PHDELWT``, with one mean-one weight column per -registered scheme beside it -- ``W_InVa``, the inverse variance ``1/sigma^2`` (the -default), and ``W_SD``, the sigma_D Wiener weight ``S/(S + sigma^2)`` built from the -expected difference power -- and ``KSCALE``, the scaler's factor from model to observed +registered scheme beside it -- ``W_Q``, the q-weight (the default), and ``W_InVa``, +the inverse variance ``1/sigma^2`` -- and ``KSCALE``, the scaler's factor from model to observed scale. This is the construction ``torchref.validate-ded`` correlates against. Build the -map with ``torchref.mtz2map -csf dF -cw W_InVa -cphi PHDELWT``, adding +map with ``torchref.mtz2map -csf dF -cw W_Q -cphi PHDELWT``, adding ``--units electrons`` for e/A^3. It needs no light-state model, so ``-lm`` is optional: .. code-block:: bash diff --git a/tests/integration/test_cli_ded_weights.py b/tests/integration/test_cli_ded_weights.py index 116a3b40..e33e8b8f 100644 --- a/tests/integration/test_cli_ded_weights.py +++ b/tests/integration/test_cli_ded_weights.py @@ -26,7 +26,7 @@ "SIGdF": "Stddev", "PHDELWT": "Phase", "W_InVa": "Weight", - "W_SD": "Weight", + "W_Q": "Weight", "KSCALE": "MTZReal", "Fc_dark": "SFAmplitude", "FreeR_flag_dark": "MTZInt", @@ -120,7 +120,7 @@ def diff_mtz(project_root, pair): "--device", "cpu", "--ded-weight", - "sigma_d", + "q", "-v", "1", "-o", @@ -138,16 +138,17 @@ def _read(path): def test_difference_map_writes_df_weights_and_scale(diff_mtz): df = _read(diff_mtz) assert {c: str(df.dtypes[c]) for c in df.columns} == DIFF_COLUMNS - for col in ("W_InVa", "W_SD"): + for col in ("W_InVa", "W_Q"): w = df[col].to_numpy().astype(float) assert np.isfinite(w).all() and (w >= 0).all() assert abs(w.mean() - 1.0) < 1e-4 assert (df["KSCALE"].to_numpy().astype(float) > 0).all() - # The sigma_D weights favour the strong reflections, inverse variance does not. + # The q-weights favour the strong reflections and never fall to zero. f = df["Fo_dark"].to_numpy().astype(float) - w_sd = df["W_SD"].to_numpy().astype(float) + w_q = df["W_Q"].to_numpy().astype(float) strong = f > np.median(f) - assert w_sd[strong].mean() > w_sd[~strong].mean() + assert w_q[strong].mean() > w_q[~strong].mean() + assert (w_q > 0).all() def test_mtz2map_builds_the_weighted_and_electron_maps(project_root, pair, diff_mtz): @@ -162,7 +163,7 @@ def test_mtz2map_builds_the_weighted_and_electron_maps(project_root, pair, diff_ "-csf", "dF", "-cw", - "W_SD", + "W_Q", "-cphi", "PHDELWT", "--device", @@ -178,7 +179,7 @@ def test_mtz2map_builds_the_weighted_and_electron_maps(project_root, pair, diff_ "-csf", "dF", "-cw", - "W_SD", + "W_Q", "-cphi", "PHDELWT", "--units", @@ -196,7 +197,7 @@ def test_mtz2map_builds_the_weighted_and_electron_maps(project_root, pair, diff_ "-csf", "dF", "-cw", - "W_SD", + "W_Q", "-cphi", "PHDELWT", "--units", @@ -239,16 +240,16 @@ def test_validate_ded_reports_every_scheme(project_root, pair): "--device", "cpu", "--ded-weight", - "sigma_d", + "q", "-v", "1", "-o", out, ) results = json.loads((out / "validate_ded_results.json").read_text()) - assert results["weights"]["requested"] == "sigma_d" - assert results["weights"]["applied"] in ("sigma_d", "inverse_variance") - assert set(results["by_weight"]) == {"none", "inverse_variance", "sigma_d"} + assert results["weights"]["requested"] == "q" + assert results["weights"]["applied"] in ("q", "inverse_variance") + assert set(results["by_weight"]) == {"none", "inverse_variance", "q"} for entry in results["by_weight"].values(): assert np.isfinite(entry["reciprocal_cc_overall"]) assert "full_cell" in entry["realspace_correlation"] @@ -256,7 +257,7 @@ def test_validate_ded_reports_every_scheme(project_root, pair): assert results["reciprocal_cc_overall"] == pytest.approx( headline["reciprocal_cc_overall"], abs=1e-3 ) - assert "weights " in proc.stdout and "sigma_d" in proc.stdout + assert "headline weights: q" in proc.stdout def test_difference_refine_runs_the_sigma_d_row(project_root, pair): @@ -294,6 +295,6 @@ def test_difference_refine_runs_the_sigma_d_row(project_root, pair): summaries = list(out.glob("*_summary.json")) assert len(summaries) == 1 results = json.loads(summaries[0].read_text())["results"] - assert results["ded_weights"]["scheme"] == "inverse_variance" - assert results["ded_weights"]["applied"] == "inverse_variance" - assert "gamma" in results["ded_weights"]["sigma_d"] + assert results["ded_weights"]["scheme"] == "q" + assert results["ded_weights"]["applied"] == "q" + assert "sigma_scale" in results["ded_weights"]["q"] diff --git a/tests/integration/test_cli_two_moment_mtz.py b/tests/integration/test_cli_two_moment_mtz.py index 99ec6703..cf9b3498 100644 --- a/tests/integration/test_cli_two_moment_mtz.py +++ b/tests/integration/test_cli_two_moment_mtz.py @@ -23,7 +23,7 @@ # scheme, and the observed-to-model scale. "PHDELWT": "Phase", "W_InVa": "Weight", - "W_SD": "Weight", + "W_Q": "Weight", "KSCALE": "MTZReal", "Fc_dark": "SFAmplitude", # The mixed model, and the extrapolated map to refine against. @@ -229,7 +229,7 @@ def test_map_coefficients_are_grouped_into_described_datasets(two_moment_all_mtz assert where["FWT"] == where["PHWT"] == where["FEXT"] == "extrapolated_light" assert where["FEXT_PHASED"] == where["FEXT_SCALAR"] == "extrapolated_light" - assert where["dF"] == where["PHDELWT"] == where["W_SD"] == "difference" + assert where["dF"] == where["PHDELWT"] == where["W_Q"] == "difference" assert where["FC"] == where["PHIC"] == "light_model" assert where["dF_corr"] == "two_moment" assert where["H"] == where["Fo_dark"] == where["FreeR_flag_dark"] == "observed" @@ -246,7 +246,7 @@ def test_the_difference_columns_carry_df_and_the_registered_weights( self, baseline_mtz ): """``dF`` must be ``Fo_light - Fo_dark``, ``W_InVa`` the mean-normalised inverse - variance and ``W_SD`` a mean-one weight -- the constructions + variance and ``W_Q`` a mean-one weight -- the constructions ``torchref.validate-ded`` correlates against. If these ever diverge, the map built from the file stops being the map the validation reports on, which is how the output drifted from the science before.""" @@ -268,8 +268,8 @@ def test_the_difference_columns_carry_df_and_the_registered_weights( assert np.abs(got_df - dfo).max() / scale < 1e-5 got_w = df["W_InVa"].to_numpy().astype(float) assert np.abs(got_w - w).max() / max(float(np.abs(w).max()), 1e-30) < 1e-4 - w_sd = df["W_SD"].to_numpy().astype(float) - assert np.isfinite(w_sd).all() and abs(w_sd.mean() - 1.0) < 1e-4 + w_q = df["W_Q"].to_numpy().astype(float) + assert np.isfinite(w_q).all() and abs(w_q.mean() - 1.0) < 1e-4 assert (df["KSCALE"].to_numpy().astype(float) > 0).all() # And the phase is the dark model's, not the mixed model's. @@ -286,11 +286,11 @@ def test_the_corrected_difference_map_pairs_with_the_same_phases( ): """``DELFWT_corr`` is the corrected difference on the *same* dark phases, so it is opened against ``PHDELWT`` and carries the selected weight scheme, the - inverse-variance weights ``W_InVa`` by default.""" + q-weights ``W_Q`` by default.""" import numpy as np df = _read(two_moment_mtz[0]) - w = df["W_InVa"].to_numpy().astype(float) + w = df["W_Q"].to_numpy().astype(float) expected = df["dF_corr"].to_numpy().astype(float) * w got = df["DELFWT_corr"].to_numpy().astype(float) diff --git a/tests/unit/maps/test_ded_weights.py b/tests/unit/maps/test_ded_weights.py index 60abe281..6bba767a 100644 --- a/tests/unit/maps/test_ded_weights.py +++ b/tests/unit/maps/test_ded_weights.py @@ -1,9 +1,10 @@ """The registered difference-coefficient weight schemes. Pinned: the three schemes exist with their MTZ column names; ``none`` is flat; -``inverse_variance`` has mean one and floors a zero sigma; ``sigma_d`` gives strong -reflections more weight than weak ones within a shell where inverse variance cannot; -and an all-noise input falls back to inverse variance with a warning that names why. +``inverse_variance`` has mean one and floors a zero sigma; ``q`` gives strong +reflections more weight than weak ones where inverse variance cannot, keeps every +reflection at or above its floor when noise dominates, and falls back to inverse +variance with a warning that names why when too few reflections exist to fit. """ import pytest @@ -84,7 +85,7 @@ def test_none_and_inverse_variance(any_device): @pytest.mark.unit -def test_sigma_d_favours_strong_reflections_where_inverse_variance_cannot(any_device): +def test_q_favours_strong_reflections_where_inverse_variance_cannot(any_device): d, hkl, cell, sg = _inputs(device=any_device) kw = { "delta_obs": d["delta_obs"], @@ -96,33 +97,52 @@ def test_sigma_d_favours_strong_reflections_where_inverse_variance_cannot(any_de } every = all_ded_weights(**kw) assert set(every) == set(SCHEMES) - sd = every["sigma_d"] - assert sd.applied == "sigma_d" and abs(float(sd.weights.mean()) - 1.0) < 1e-4 - assert 0.8 < sd.diagnostics["gamma"] < 1.2 - assert sd.diagnostics["n_shell"] > 10 and "shells" in sd.diagnostics - # Within the highest-resolution tenth, the strongest reflections carry more weight. - order = torch.argsort(d["d_star_sq"])[-2000:] - f, w = d["f_dark"][order], sd.weights[order] + q = every["q"] + assert q.applied == "q" and abs(float(q.weights.mean()) - 1.0) < 1e-4 + assert 0.8 < q.diagnostics["gamma"] < 1.2 + assert q.diagnostics["converged"] + # The strong half carries more weight; inverse variance cannot tell the halves apart + # because the sigmas are constant. The floor compresses the weights into at most a + # factor three, so the margin is smaller than an unbounded Wiener weight would give. + f, w = d["f_dark"], q.weights strong, weak = f > f.median(), f <= f.median() - assert float(w[strong].mean()) > 1.5 * float(w[weak].mean()) - ivw = every["inverse_variance"].weights[order] + assert float(w[strong].mean()) > 1.2 * float(w[weak].mean()) + ivw = every["inverse_variance"].weights assert abs(float(ivw[strong].mean()) - float(ivw[weak].mean())) < 1e-4 @pytest.mark.unit -def test_all_noise_falls_back_to_inverse_variance_with_a_warning(): +def test_q_never_removes_a_reflection_when_noise_dominates(): d, hkl, cell, sg = _inputs(n=5000, sig_frac=50.0) + q = compute_ded_weights( + "q", + delta_obs=d["delta_obs"], + sigma_diff=d["sigma_diff"] * 1.2, + hkl=hkl, + cell=cell, + spacegroup=sg, + f_dark=d["f_dark"], + ) + assert q.applied == "q" + floor = q.diagnostics["snr_floor"] / (1.0 + q.diagnostics["snr_floor"]) + assert q.diagnostics["weight_min"] >= floor - 1e-6 + assert bool((q.weights > 0).all()) + + +@pytest.mark.unit +def test_too_few_reflections_fall_back_to_inverse_variance_with_a_warning(): + d, hkl, cell, sg = _inputs(n=5) kw = { "delta_obs": d["delta_obs"], - "sigma_diff": d["sigma_diff"] * 1.2, + "sigma_diff": d["sigma_diff"], "hkl": hkl, "cell": cell, "spacegroup": sg, "f_dark": d["f_dark"], } with pytest.warns(DedWeightFallbackWarning, match="inverse-variance"): - sd = compute_ded_weights("sigma_d", **kw) - assert sd.scheme == "sigma_d" and sd.applied == "inverse_variance" - assert "fallback_reason" in sd.diagnostics + q = compute_ded_weights("q", **kw) + assert q.scheme == "q" and q.applied == "inverse_variance" + assert "fallback_reason" in q.diagnostics ivw = compute_ded_weights("inverse_variance", **kw) - assert torch.allclose(sd.weights, ivw.weights) + assert torch.allclose(q.weights, ivw.weights) diff --git a/tests/unit/refinement/test_difference_power.py b/tests/unit/refinement/test_difference_power.py new file mode 100644 index 00000000..65bc298b --- /dev/null +++ b/tests/unit/refinement/test_difference_power.py @@ -0,0 +1,92 @@ +"""Properties of the shell-free difference-power fit. + +Pinned on seeded synthetic differences with a known power law: the power, the +dark-amplitude exponent and the sigma scale are recovered from one dataset, including +when the reported sigmas are uniformly inflated; a fixed exponent stays fixed; the +bounded Wiener weight never falls below its floor, so no reflection or resolution range +is removed even when the data hold no signal; the fit runs under ``torch.no_grad()`` and +on every available device. +""" + +import pytest +import torch + +from torchref.refinement.model_error_estimation.difference_power import ( + bounded_wiener_weight, + fit_difference_power, +) + +#: Median absolute log error of the recovered power. The fit has seven parameters +#: against 40 000 reflections; 0.15 is several times the scatter observed across seeds. +LOG_POWER_ATOL = 0.15 +#: Tolerance on the exponent and on the relative sigma scale. +GAMMA_ATOL = 0.1 +K_RTOL = 0.05 + + +def synth(n=40000, sigma_inflation=1.0, signal=1.0, seed=0, device="cpu"): + """Differences with power ``4 exp(-8 d*^2) (F / )``; reported sigmas vary + many-fold within a resolution, so the sigma scale is identifiable.""" + g = torch.Generator().manual_seed(seed) + dss = torch.rand(n, generator=g) / 1.6**2 + f = torch.exp(torch.randn(n, generator=g) * 0.6) * 100 * torch.exp(-10 * dss) + s_true = signal * 4.0 * torch.exp(-8.0 * dss) * (f / f.mean()) + sig = 0.5 + 6 * dss / dss.max() * torch.exp(torch.randn(n, generator=g) * 0.4) + d = torch.randn(n, generator=g) * s_true.sqrt() + torch.randn(n, generator=g) * sig + out = dict(delta=d, sigma=sig * sigma_inflation, dss=dss, f=f, s_true=s_true) + return {k: v.to(device) for k, v in out.items()} + + +@pytest.mark.unit +@pytest.mark.parametrize("inflation", [1.0, 1.5]) +def test_recovers_power_exponent_and_sigma_scale(inflation): + s = synth(sigma_inflation=inflation) + fit = fit_difference_power(s["delta"], s["sigma"], s["dss"], f_dark=s["f"]) + assert fit.converged + assert fit.gamma == pytest.approx(1.0, abs=GAMMA_ATOL) + assert fit.sigma_scale == pytest.approx(1.0 / inflation, rel=K_RTOL) + power = fit.signal_power(s["dss"], f_dark=s["f"]) + err = (power / s["s_true"]).log().abs().median() + assert float(err) < LOG_POWER_ATOL + assert torch.isfinite(fit.stderr[: len(fit.coeffs)]).all() + + +@pytest.mark.unit +def test_fixed_gamma_is_kept(): + s = synth() + fit = fit_difference_power(s["delta"], s["sigma"], s["dss"], f_dark=s["f"], gamma=0.0) + assert fit.gamma == 0.0 + fit = fit_difference_power(s["delta"], s["sigma"], s["dss"], f_dark=s["f"], gamma=2.0) + assert fit.gamma == 2.0 + + +@pytest.mark.unit +def test_weight_never_removes_a_reflection_without_signal(): + s = synth(signal=0.0) + fit = fit_difference_power( + s["delta"], s["sigma"], s["dss"], f_dark=s["f"], fit_sigma_scale=False + ) + snr = fit.snr(s["sigma"], d_star_sq=s["dss"], f_dark=s["f"]) + w = bounded_wiener_weight(snr, 0.5) + assert float(w.min()) >= 1.0 / 3.0 - 1e-6 + assert float(w.max()) < 1.0 + assert float(bounded_wiener_weight(torch.zeros(1), 0.0)) == 0.0 + with pytest.raises(ValueError): + bounded_wiener_weight(snr, -0.1) + + +@pytest.mark.unit +def test_runs_on_device(any_device): + s = synth(n=5000, device=any_device) + fit = fit_difference_power(s["delta"], s["sigma"], s["dss"], f_dark=s["f"]) + power = fit.signal_power(s["dss"], f_dark=s["f"]) + assert power.device.type == any_device.type + assert torch.isfinite(power).all() and bool((power > 0).all()) + + +@pytest.mark.unit +def test_fits_under_no_grad(): + s = synth(n=5000) + with torch.no_grad(): + fit = fit_difference_power(s["delta"], s["sigma"], s["dss"], f_dark=s["f"]) + assert fit.converged diff --git a/torchref/cli/_common.py b/torchref/cli/_common.py index 135e62a8..0be78058 100644 --- a/torchref/cli/_common.py +++ b/torchref/cli/_common.py @@ -402,19 +402,20 @@ def add_ded_weight_args(parser: argparse.ArgumentParser) -> None: "--ded-weight", choices=list(SCHEMES), default=DEFAULT_SCHEME, - help="Per-reflection weight for difference coefficients: 'inverse_variance' " - "is 1/sigma^2, 'sigma_d' is the Wiener weight S/(S+sigma^2) from the " - "expected difference power (needs calibrated sigmas; check the reported " - f"clamped-shell count), 'none' is flat (default: {DEFAULT_SCHEME}). All " - "weights are written as columns.", + help="Per-reflection weight for difference coefficients: 'q' is the " + "q-weight, a Wiener weight from a shell-free fit of the expected difference " + "power that down-weights noisy reflections to no less than a third, " + "'inverse_variance' is 1/sigma^2, 'none' is flat (default: " + f"{DEFAULT_SCHEME}). All weights are written as columns.", ) parser.add_argument( "--sigma-d-gamma", type=float, default=None, metavar="GAMMA", - help="Fix the dark-amplitude exponent of the sigma_d power law in [0, 2] " - "instead of fitting it (default: fitted).", + help="Fix the dark-amplitude exponent of the difference power law in [0, 2] " + "instead of fitting it; used by the q-weight and the sigma_D difference " + "target (default: fitted).", ) diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index 1517e552..44824ecc 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -578,9 +578,9 @@ def _difference_columns( ``dF``/``SIGdF`` is the signed amplitude difference ``|Fo_light| - |Fo_dark|`` with its propagated uncertainty, ``PHDELWT`` the **dark** model's phase it is carried on: the isomorphous difference Fourier, and the construction ``torchref.validate-ded`` - correlates against. One weight column per registered scheme (``W_SD``, ``W_InVa``; + correlates against. One weight column per registered scheme (``W_Q``, ``W_InVa``; MTZ type ``W``, mean one) sits beside it, so any weighting is ``dF`` times a column - and reproducible from the file: ``torchref.mtz2map -csf dF -cw W_SD -cphi PHDELWT``. + and reproducible from the file: ``torchref.mtz2map -csf dF -cw W_Q -cphi PHDELWT``. ``KSCALE`` (type ``R``) is the scaler's multiplicative factor from model to observed scale, so ``dF / KSCALE`` is in electrons and ``mtz2map --units electrons`` gives e/A^3. @@ -856,7 +856,7 @@ def _np(t): # are at most 80 characters. _MTZ_DATASET_HISTORY = { "observed": "observed: Fo_dark, Fo_light and flags on the shared scale; Fc_dark", - "difference": "difference: dF/SIGdF on dark phases PHDELWT; weights W_SD, W_InVa", + "difference": "difference: dF/SIGdF on dark phases PHDELWT; weights W_Q, W_InVa", "light_model": ( "light_model: FC/PHIC, amplitude and phase of the mixed dark+light model" ), @@ -931,7 +931,7 @@ def write_results_mtz( The default output is the **difference map**: ``dF``/``SIGdF`` on the dark model's phases ``PHDELWT``, with one mean-one weight column per registered scheme - (``W_SD``, ``W_InVa``) and the observed-to-model scale ``KSCALE``; see + (``W_Q``, ``W_InVa``) and the observed-to-model scale ``KSCALE``; see :func:`_difference_columns`. ``ded_weight`` selects the scheme the model-phased difference columns and the two-moment columns are weighted with. That needs no light-state model, which is why ``mc`` is optional -- with a dark model alone this @@ -966,7 +966,7 @@ def write_results_mtz( Weight scheme for the model-phased and two-moment difference columns; one of :data:`torchref.maps.ded_weights.SCHEMES`. sigma_d_config : SigmaDConfig, optional - Exponent and shrinkage settings of the ``sigma_d`` scheme. + Its ``gamma`` fixes the ``F_dark`` exponent of the ``q`` scheme. Returns ------- @@ -1025,7 +1025,7 @@ def write_results_mtz( cell=data_dark.cell, spacegroup=data_dark.spacegroup, f_dark=Fobs_dark_vals, - sigma_d_config=sigma_d_config, + gamma=sigma_d_config.gamma if sigma_d_config is not None else None, ) selected = all_w[ded_weight] weights = selected.weights.detach().cpu().numpy() @@ -1039,38 +1039,26 @@ def write_results_mtz( geometry = reflection_geometry( hkl, data_dark.cell, data_dark.spacegroup, diff_t.device, diff_t.dtype ) - sd_diag = { - k: v - for k, v in all_w["sigma_d"].diagnostics.items() - if k != "weight_sigma_d_raw" - } + q_diag = all_w["q"].diagnostics diagnostics = { "ded_weights": { "scheme": ded_weight, "applied": selected.applied, - "sigma_d": sd_diag, + "q": q_diag, } } if verbose > 0: print(f" Difference weights: {ded_weight} (applied: {selected.applied})") - print( - f" sigma_D: gamma = {sd_diag['gamma']:.3f} ({sd_diag['gamma_reason']}), " - f"tau = {sd_diag['tau']:.3f}, shells = {sd_diag['n_shell']}, " - f"shells without difference power = {sd_diag['n_s2_clamped']}" - ) - if "fallback_reason" in sd_diag: - print(f" sigma_D fallback: {sd_diag['fallback_reason']}") - if verbose > 1 and not sd_diag["degenerate"]: - table = sd_diag["shells"] - print(" sigma_D shells: d(A) n B S2 Sigma_N") - for dss, n, b, s2, sn in zip( - table["d_star_sq"], - table["counts"], - table["B"], - table["S2"], - table["Sigma_N"], - ): - print(f" {dss ** -0.5:6.2f} {int(n):5d} {b:9.4f} {s2:9.4f} {sn:9.4f}") + if "fallback_reason" in q_diag: + print(f" q-weight fallback: {q_diag['fallback_reason']}") + else: + print( + f" q-weight fit: gamma = {q_diag['gamma']:.3f}, " + f"sigma scale k = {q_diag['sigma_scale']:.3f}, " + f"centric factor = {q_diag['centric_factor']:.3f}, " + f"weights {q_diag['weight_min']:.3f}-{q_diag['weight_max']:.3f} " + f"before normalisation" + ) columns, types = _difference_columns( data_dark, diff --git a/torchref/cli/difference_map.py b/torchref/cli/difference_map.py index 0d4ac24e..e2c2d2a6 100644 --- a/torchref/cli/difference_map.py +++ b/torchref/cli/difference_map.py @@ -8,9 +8,9 @@ The default output is the difference map: the amplitude difference ``|Fo_light| - |Fo_dark|`` as ``dF``/``SIGdF`` carried on the **dark** model's phases ``PHDELWT``, with the per-reflection weights of every registered scheme beside it as -``W_InVa`` (inverse variance, the default) and ``W_SD`` (sigma_D Wiener weight), and the +``W_Q`` (q-weight, the default) and ``W_InVa`` (inverse variance), and the observed-to-model scale ``KSCALE``. Build the map with -``torchref.mtz2map -csf dF -cw W_InVa -cphi PHDELWT`` (``--units electrons`` for e/A^3). +``torchref.mtz2map -csf dF -cw W_Q -cphi PHDELWT`` (``--units electrons`` for e/A^3). That needs no light-state model, so ``-lm`` is optional. It is also deliberately not a *phased* difference map: putting the light state's model phases into the observed amplitude biases the map toward the very model the experiment is testing. diff --git a/torchref/cli/mtz2map.py b/torchref/cli/mtz2map.py index 45c75e1a..3f39bc5e 100644 --- a/torchref/cli/mtz2map.py +++ b/torchref/cli/mtz2map.py @@ -66,7 +66,7 @@ def main(): type=str, metavar="COL", help="Weight column multiplied into the amplitudes before the FFT " - "(e.g. W_SD, W_InVa from torchref.difference-map). Default: none.", + "(e.g. W_Q, W_InVa from torchref.difference-map). Default: none.", ) inp.add_argument( "-ck", diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index 7614f104..d00c78bd 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -300,7 +300,7 @@ def setup_ded_context( cell=data_dark.cell, spacegroup=data_dark.spacegroup, f_dark=F_dark, - sigma_d_config=sigma_d_config, + gamma=sigma_d_config.gamma if sigma_d_config is not None else None, ) selected = all_w[ded_weight] weights = selected.weights @@ -355,8 +355,7 @@ def setup_ded_context( "ded_weight_applied": selected.applied, "ded_weight_diagnostics": { k: v - for k, v in all_w["sigma_d"].diagnostics.items() - if k not in ("weight_sigma_d_raw", "shells") + for k, v in all_w["q"].diagnostics.items() }, "d_spacing": d_spacing, "cell_t": cell_t, @@ -730,10 +729,9 @@ def run_validation(args): for k in ( "gamma", "gamma_fitted", - "gamma_reason", - "tau", - "n_shell", - "n_s2_clamped", + "sigma_scale", + "centric_factor", + "snr_floor", "fallback_reason", ) }, diff --git a/torchref/maps/ded_weights.py b/torchref/maps/ded_weights.py index 258831e6..0c9348fd 100644 --- a/torchref/maps/ded_weights.py +++ b/torchref/maps/ded_weights.py @@ -8,43 +8,41 @@ Every reflection weighted equally. ``inverse_variance`` ``1 / sigma_diff**2``. Weights by precision alone; the right rule for averaging - estimates of one quantity, and the default for difference maps. -``sigma_d`` - The Wiener weight ``S / (S + sigma_diff**2)`` with ``S`` the expected true difference - power from :mod:`torchref.refinement.model_error_estimation.sigma_d`. Weights by the - signal fraction of each coefficient, so strong reflections whose expected difference - is large keep their weight. ``S`` is ``mean(dF**2) - mean(sigma**2)`` per shell, so - it inherits any miscalibration of ``sigma_diff``: where the reported sigmas are too - large the estimate finds no power and the weight vanishes, which turns the scheme - into a resolution cut. The count of such shells is reported as ``n_s2_clamped``; - a large fraction means the sigmas, not the data, are deciding the map. - -Plain tensors in and out. The weights live on the device of ``delta_obs``. The sigma_D -estimator is imported inside the scheme that needs it so that :mod:`torchref.maps` does -not import :mod:`torchref.refinement` at module load. + estimates of one quantity. On French-Wilson amplitudes it *up*-weights the weak + high-resolution reflections, whose posterior sigma the prior keeps small. +``q`` + The q-weight: a Wiener weight ``(snr + b) / (snr + 1 + b)`` with + ``snr = S / (k sigma_diff)**2``, the default. ``S`` and the sigma scale ``k`` come + from :func:`~torchref.refinement.model_error_estimation.difference_power. + fit_difference_power`, a per-reflection maximum-likelihood fit of ``log S`` as a + Chebyshev series in resolution plus ``gamma log F_dark``: no resolution shells. The + floor ``b`` keeps a reflection without signal at ``b / (1 + b)`` of the full weight + (one third at the default), so a noisy resolution range is down-weighted, never + removed. ``k`` absorbs a uniform miscalibration of the sigmas; French-Wilson sigmas + of two datasets overstate the error of their difference, so ``k < 1`` is normal. + +Plain tensors in and out. The weights live on the device of ``delta_obs``. The +difference-power fit is imported inside the scheme that needs it so that +:mod:`torchref.maps` does not import :mod:`torchref.refinement` at module load. """ import warnings from dataclasses import dataclass, field -from typing import TYPE_CHECKING import torch from torchref.base.reciprocal.basis import get_scattering_vectors from torchref.base.targets.xray_likelihoods import SIGMA_FLOOR_ABS, SIGMA_FLOOR_FRAC -if TYPE_CHECKING: - from torchref.refinement.model_error_estimation.sigma_d import SigmaDConfig - #: The selectable schemes, in the order they are reported. -SCHEMES = ("none", "inverse_variance", "sigma_d") -#: Scheme applied when none is named. Inverse variance, because ``sigma_d`` depends on -#: calibrated sigmas: on the 15 Sep campaign TorchSX's TD1 sigmas were ~1.5x too large at -#: high resolution, ``sigma_d`` zeroed 60-90 % of the shells there and the map agreement -#: rose in the bulk solvent as much as in the region of interest. -DEFAULT_SCHEME = "inverse_variance" +SCHEMES = ("none", "inverse_variance", "q") +#: Scheme applied when none is named. On simulated dark/light half datasets of 1DAW +#: (ligand moved at 0.3 and 0.6 occupancy, calibrated and 1.5x-inflated sigmas) ``q`` +#: recovered the most contour-level information about the true difference density; +#: ``inverse_variance`` recovered less than no weighting. +DEFAULT_SCHEME = "q" #: MTZ column carrying each scheme's weight (type ``W``); ``none`` writes no column. -WEIGHT_COLUMNS = {"inverse_variance": "W_InVa", "sigma_d": "W_SD"} +WEIGHT_COLUMNS = {"inverse_variance": "W_InVa", "q": "W_Q"} class DedWeightFallbackWarning(UserWarning): @@ -65,8 +63,8 @@ class DedWeights: weights Per-reflection weights, shape ``(N,)``, mean one over finite positive entries. diagnostics - Scheme-specific record: for ``sigma_d`` the fitted exponent, shrinkage sd, - clamp counters and the per-shell table. + Scheme-specific record: for ``q`` the fitted exponent, sigma scale, centric + factor, Chebyshev coefficients and their standard errors, and the weight range. """ scheme: str @@ -80,9 +78,8 @@ def normalise_mean_one(w: torch.Tensor) -> torch.Tensor: that mean is not positive. Non-finite entries become zero, so a coefficient without a usable uncertainty drops - out of the map rather than poisoning it. Zero weights (shells without difference - power) stay zero and count in the mean, so the column mean is one whatever fraction - of the reflections carries weight. + out of the map rather than poisoning it. Zero weights stay zero and count in the + mean, so the column mean is one whatever fraction of the reflections carries weight. """ w = torch.where(torch.isfinite(w), w, torch.zeros_like(w)) if w.numel() == 0 or not bool((w > 0).any()): @@ -129,7 +126,8 @@ def compute_ded_weights( spacegroup, f_dark: torch.Tensor | None = None, fit_mask: torch.Tensor | None = None, - sigma_d_config: "SigmaDConfig | None" = None, + gamma: float | None = None, + snr_floor: float | None = None, ) -> DedWeights: """Per-reflection weights for one scheme. @@ -146,19 +144,23 @@ def compute_ded_weights( Unit cell, as a :class:`~torchref.symmetry.Cell` or its six parameters in A and degrees. spacegroup : SpaceGroup or None - For the reflection multiplicity; ``None`` means ones. + For the reflection multiplicity and centric flags; ``None`` means P1. f_dark : torch.Tensor, optional - Dark amplitudes, shape ``(N,)``, for the sigma_D amplitude power law. + Dark amplitudes, shape ``(N,)``, for the ``q`` power law in ``F_dark``. fit_mask : torch.Tensor, optional - Reflections entering the sigma_D fit; default every finite one. - sigma_d_config : SigmaDConfig, optional - Exponent and shrinkage settings for ``sigma_d``. + Reflections entering the ``q`` fit; default every finite one. + gamma : float, optional + Fix the ``F_dark`` exponent of the ``q`` fit instead of fitting it. + snr_floor : float, optional + Signal-to-noise floor of the ``q`` weight; default + :data:`~torchref.refinement.model_error_estimation.difference_power. + DEFAULT_SNR_FLOOR`. Returns ------- DedWeights - Mean-one weights on ``delta_obs.device``. When ``sigma_d`` finds no difference - power in any shell, the inverse-variance weights are returned with + Mean-one weights on ``delta_obs.device``. When the ``q`` fit has too few + usable reflections, the inverse-variance weights are returned with ``applied="inverse_variance"`` and a :class:`DedWeightFallbackWarning`. """ if scheme not in SCHEMES: @@ -171,64 +173,64 @@ def compute_ded_weights( scheme, scheme, normalise_mean_one(_inverse_variance(sigma_diff)) ) - from torchref.refinement.model_error_estimation.sigma_d import ( - SigmaDConfig, - estimate_sigma_d, - sigma_d_per_reflection, + from torchref.refinement.model_error_estimation.difference_power import ( + DEFAULT_SNR_FLOOR, + bounded_wiener_weight, + fit_difference_power, ) - config = sigma_d_config if sigma_d_config is not None else SigmaDConfig() + floor = DEFAULT_SNR_FLOOR if snr_floor is None else float(snr_floor) d = delta_obs.reshape(-1) eps, dss = reflection_geometry(hkl, cell, spacegroup, d.device, d.dtype) f = f_dark.reshape(-1).to(d.device, d.dtype) if f_dark is not None else None - mask = ( - fit_mask.reshape(-1).to(d.device, torch.bool) - if fit_mask is not None - else torch.isfinite(d) & torch.isfinite(sigma_diff) - ) - shells = estimate_sigma_d( - d, sigma_diff, eps, dss, f, mask, gamma=config.gamma, shrink=config.shrink + centric = ( + spacegroup.is_centric(torch.as_tensor(hkl, device=d.device)).to(d.device) + if spacegroup is not None + else None ) - est = sigma_d_per_reflection(shells, dss, eps, f, sigma_diff) - diagnostics = { - "gamma": shells.gamma, - "gamma_fitted": shells.gamma_fitted, - "tau": shells.tau, - "curve_a": shells.curve_a, - "curve_b": shells.curve_b, - "degenerate": shells.degenerate, - "all_zero": shells.all_zero, - **shells.diagnostics, - "shells": { - "d_star_sq": shells.bin_dss.detach().cpu().tolist(), - "counts": shells.counts.detach().cpu().tolist(), - "B": shells.B.detach().cpu().tolist(), - "S2": shells.S2.detach().cpu().tolist(), - "Sigma_N_raw": shells.Sigma_N_raw.detach().cpu().tolist(), - "Sigma_N": shells.Sigma_N.detach().cpu().tolist(), - }, - "weight_sigma_d_raw": est.w.detach(), - } - if shells.all_zero or shells.degenerate: - reason = ( - "no difference power above the measurement variance in any shell " - f"(n_s2_clamped={shells.diagnostics['n_s2_clamped']})" - if shells.all_zero - else "fewer than two usable reflections" + try: + fit = fit_difference_power( + d, + sigma_diff, + dss, + epsilon=eps, + f_dark=f, + centric=centric, + fit_mask=fit_mask, + gamma=gamma, ) + except ValueError as err: + reason = str(err) warnings.warn( - f"sigma_d weights: {reason}; applying inverse-variance weights instead", + f"q weights: {reason}; applying inverse-variance weights instead", DedWeightFallbackWarning, stacklevel=2, ) - diagnostics["fallback_reason"] = reason return DedWeights( scheme, "inverse_variance", normalise_mean_one(_inverse_variance(sigma_diff)), - diagnostics, + {"fallback_reason": reason}, ) - return DedWeights(scheme, scheme, normalise_mean_one(est.w), diagnostics) + snr = fit.snr(sigma_diff, d_star_sq=dss, epsilon=eps, f_dark=f, centric=centric) + w = bounded_wiener_weight(snr, floor) + ok = torch.isfinite(w) + diagnostics = { + "gamma": fit.gamma, + "gamma_fitted": gamma is None and f is not None, + "sigma_scale": fit.sigma_scale, + "centric_factor": fit.centric_factor, + "order": len(fit.coeffs) - 1, + "coeffs": fit.coeffs.detach().cpu().tolist(), + "stderr": fit.stderr.detach().cpu().tolist(), + "stol_range": list(fit.stol_range), + "converged": fit.converged, + "n_fit": fit.n_fit, + "snr_floor": floor, + "weight_min": float(w[ok].min()) if bool(ok.any()) else float("nan"), + "weight_max": float(w[ok].max()) if bool(ok.any()) else float("nan"), + } + return DedWeights(scheme, scheme, normalise_mean_one(w), diagnostics) def all_ded_weights(**kwargs) -> dict[str, DedWeights]: diff --git a/torchref/refinement/model_error_estimation/difference_power.py b/torchref/refinement/model_error_estimation/difference_power.py new file mode 100644 index 00000000..2f81b76e --- /dev/null +++ b/torchref/refinement/model_error_estimation/difference_power.py @@ -0,0 +1,383 @@ +"""Expected difference power by maximum likelihood over reflections, without shells. + +A light-minus-dark amplitude difference ``dF_obs = dF_true + noise`` is modelled as + +.. math:: + + dF_h \\sim N(0, V_h), \\qquad V_h = \\epsilon_h c_h S_h + k^2 \\sigma_h^2, + + \\log S_h = \\sum_{j=0}^{n} a_j T_j(x_h) + \\gamma \\log F_{dark,h}, + +with :math:`T_j` the Chebyshev polynomials of :func:`torchref.scaling.basis. +chebyshev_design` in :math:`x_h = \\sin\\theta/\\lambda` over the fitted range (the +abscissa every smooth resolution curve in the package uses), :math:`k` a single scale on +the reported sigmas and :math:`c_h` a free factor on centric reflections. Every +parameter is fitted jointly by Newton's method on the per-reflection negative +log-likelihood, so there are no resolution shells, no clamped moments and no +interpolation: ``S`` is positive and smooth by construction. + +The ``gamma * log F_dark`` term needs no per-resolution normalisation of ``F_dark``: +``log (d*^2)`` is itself a smooth function of resolution, so the polynomial +absorbs it. The sigma scale ``k`` is identifiable because the reported sigmas vary +many-fold between reflections at one resolution while ``S`` does not; it is what keeps +inflated sigmas from reading as an absence of signal. + +:func:`bounded_wiener_weight` turns a fit into a weight that down-weights noisy +reflections but never removes one, the resolution-continuous counterpart of the +q-weight's floor. + +Plain tensors in and out; every result lives on the device of ``delta_obs``. +""" + +import math +from dataclasses import dataclass + +import torch + +from torchref.config import get_float_dtype +from torchref.scaling.basis import chebyshev_design + +#: Chebyshev order of ``log S`` in ``sin(theta)/lambda``. Four covers a Wilson-like +#: fall-off (quadratic in this abscissa) plus low-resolution curvature; the fit +#: carries the standard errors to judge whether a higher order is supported. +DEFAULT_ORDER = 4 +#: Signal-to-noise floor of :func:`bounded_wiener_weight`. One half reproduces the +#: q-weight's noise-only limit (``S`` floored at half the raw difference power gives +#: ``w = 1/3``), so a reflection without signal keeps a third of the full weight. +DEFAULT_SNR_FLOOR = 0.5 +#: Degrees of freedom of the Student-t likelihood when ``robust=True``. +DEFAULT_NU = 4.0 +#: Newton iterations; the problem has at most eight parameters and converges in ~10. +MAX_ITER = 60 +_GAMMA_BOUNDS = (-1.0, 3.0) + + +@dataclass(frozen=True) +class DifferencePowerFit: + """A fitted difference-power model; evaluate it with :meth:`signal_power`. + + Attributes + ---------- + coeffs : torch.Tensor + Chebyshev coefficients of ``log S`` in the standardised amplitude units, shape + ``(order + 1,)``. + gamma : float + Exponent on ``F_dark``; ``0`` when no dark amplitude was used. + sigma_scale : float + The factor ``k`` on the reported sigmas; ``1`` when not fitted. + centric_factor : float + Power of a centric reflection relative to an acentric one at equal resolution; + ``1`` when no centric flags were given. + stol_range : tuple of float + The ``sin(theta)/lambda`` range (A^-1) the polynomial is defined on. Outside it + the curve is held at its endpoint value, because a Chebyshev series diverges + beyond ``[-1, 1]``. + amp_scale, log_f_ref : float + The amplitude unit the fit was done in and the mean ``log F_dark`` it was + centred on. + stderr : torch.Tensor + Standard errors of ``(coeffs, gamma, log k, log centric_factor)`` from the + inverse Hessian; NaN where a parameter was fixed. + nll : float + Mean negative log-likelihood per reflection at the optimum. + converged : bool + Whether the Newton step fell below tolerance. + n_fit : int + Reflections in the fit. + """ + + coeffs: torch.Tensor + gamma: float + sigma_scale: float + centric_factor: float + stol_range: tuple + amp_scale: float + log_f_ref: float + stderr: torch.Tensor + nll: float + converged: bool + n_fit: int + + def signal_power( + self, + d_star_sq: torch.Tensor, + epsilon: torch.Tensor | None = None, + f_dark: torch.Tensor | None = None, + centric: torch.Tensor | None = None, + ) -> torch.Tensor: + """Expected true difference power ``epsilon * c * S`` per reflection. + + Parameters + ---------- + d_star_sq : torch.Tensor + ``1/d**2`` in A^-2, shape ``(N,)``. + epsilon : torch.Tensor, optional + Reflection multiplicity, shape ``(N,)``; ones when omitted. + f_dark : torch.Tensor, optional + Dark amplitudes, shape ``(N,)``, in the units of the fitted differences. + Required when the fit used them (``gamma != 0``). + centric : torch.Tensor, optional + Boolean centric flags, shape ``(N,)``. + + Returns + ------- + torch.Tensor + Power in squared units of the fitted differences, shape ``(N,)``. + """ + dss = d_star_sq.to(self.coeffs) + log_s = _design(dss, len(self.coeffs), self.stol_range) @ self.coeffs + if self.gamma != 0.0: + if f_dark is None: + raise ValueError("this fit uses F_dark; pass f_dark") + log_s = log_s + self.gamma * (_log_amp(f_dark.to(dss), self.amp_scale) + - self.log_f_ref) + if centric is not None and self.centric_factor != 1.0: + log_s = log_s + math.log(self.centric_factor) * centric.to(dss) + power = torch.exp(log_s) * self.amp_scale**2 + if epsilon is not None: + power = power * epsilon.to(dss) + return power + + def snr(self, sigma: torch.Tensor, **kwargs) -> torch.Tensor: + """``signal_power / (k * sigma)**2`` per reflection; keywords as + :meth:`signal_power`.""" + s = self.signal_power(**kwargs) + noise = (self.sigma_scale * sigma.to(s)) ** 2 + return s / noise.clamp(min=torch.finfo(s.dtype).tiny) + + +def _design(dss: torch.Tensor, n_coeff: int, stol_range: tuple) -> torch.Tensor: + """Chebyshev design in ``sin(theta)/lambda = sqrt(d*^2) / 2``.""" + lo, hi = stol_range + return chebyshev_design(dss.clamp(min=0.0).sqrt() / 2.0, n_coeff, lo=lo, hi=hi) + + +def _log_amp(f: torch.Tensor, amp_scale: float) -> torch.Tensor: + # A zero dark amplitude would send log S to -inf and delete the reflection's power; + # a floor at 1 % of the amplitude unit keeps it finite without moving real values. + return torch.log((f / amp_scale).clamp(min=1e-2)) + + +# The Newton steps take the Hessian by autograd, so the fit enables gradients itself: +# its callers (map writers) run under ``torch.no_grad()``. +@torch.enable_grad() +def fit_difference_power( + delta_obs: torch.Tensor, + sigma_diff: torch.Tensor, + d_star_sq: torch.Tensor, + *, + epsilon: torch.Tensor | None = None, + f_dark: torch.Tensor | None = None, + centric: torch.Tensor | None = None, + fit_mask: torch.Tensor | None = None, + order: int = DEFAULT_ORDER, + gamma: float | None = None, + fit_sigma_scale: bool = True, + robust: bool = False, + nu: float = DEFAULT_NU, +) -> DifferencePowerFit: + """Fit the expected difference power by per-reflection maximum likelihood. + + Parameters + ---------- + delta_obs, sigma_diff : torch.Tensor + Signed observed differences and their reported uncertainty, shape ``(N,)``, on + one amplitude scale. + d_star_sq : torch.Tensor + ``1/d**2`` in A^-2, shape ``(N,)``. + epsilon : torch.Tensor, optional + Reflection multiplicity; ones when omitted. + f_dark : torch.Tensor, optional + Dark amplitudes for the ``F_dark**gamma`` term; without them ``gamma`` is 0. + centric : torch.Tensor, optional + Boolean centric flags; given, a centric power factor is fitted. + fit_mask : torch.Tensor, optional + Reflections entering the fit; default every finite one with positive sigma. + order : int + Chebyshev order of ``log S`` in ``sin(theta)/lambda``. + gamma : float, optional + Fix the ``F_dark`` exponent instead of fitting it. + fit_sigma_scale : bool + Fit the scale ``k`` on the reported sigmas. Off, the sigmas are taken as + calibrated. + robust : bool + Use a Student-t likelihood with ``nu`` degrees of freedom, so a large + difference loses influence on the fit smoothly instead of dominating it. + nu : float + Student-t degrees of freedom. + + Returns + ------- + DifferencePowerFit + The fitted model, detached from the inputs. Standard errors are NaN for fixed + parameters. Runs with gradients enabled internally, so it works under + ``torch.no_grad()``. + + Raises + ------ + ValueError + If fewer than ``order + 4`` reflections are usable. + """ + dtype = torch.promote_types(get_float_dtype(), delta_obs.dtype) + dev = delta_obs.device + d = delta_obs.detach().reshape(-1).to(dev, dtype) + sig = sigma_diff.detach().reshape(-1).to(dev, dtype) + dss = d_star_sq.detach().reshape(-1).to(dev, dtype) + ok = torch.isfinite(d) & torch.isfinite(sig) & (sig > 0) & torch.isfinite(dss) + if fit_mask is not None: + ok = ok & fit_mask.reshape(-1).to(dev, torch.bool) + use_f = f_dark is not None and gamma != 0.0 + if use_f: + f = f_dark.detach().reshape(-1).to(dev, dtype) + ok = ok & torch.isfinite(f) + n_fit = int(ok.sum()) + if n_fit < order + 4: + raise ValueError(f"need at least {order + 4} usable reflections, got {n_fit}") + + d, sig, dss = d[ok], sig[ok], dss[ok] + # Work in units of the rms difference so every term of the likelihood is O(1) and + # the Hessian is well conditioned in float32. + amp_scale = float(d.square().mean().sqrt().clamp(min=1e-12)) + d2 = (d / amp_scale).square() + log_sig2 = 2.0 * torch.log(sig / amp_scale) + log_eps = ( + torch.log(epsilon.reshape(-1).to(dev, dtype)[ok]) + if epsilon is not None + else torch.zeros_like(d) + ) + stol = dss.clamp(min=0.0).sqrt() / 2.0 + stol_range = (float(stol.min()), float(stol.max())) + basis = _design(dss, order + 1, stol_range) + if use_f: + log_f = _log_amp(f[ok], amp_scale) + log_f_ref = float(log_f.mean()) + log_f = log_f - log_f_ref + else: + log_f, log_f_ref = torch.zeros_like(d), 0.0 + has_centric = centric is not None and bool(centric.reshape(-1)[ok].any()) + cen = centric.reshape(-1).to(dev)[ok].to(dtype) if has_centric else None + + n_c = order + 1 + fit_gamma = use_f and gamma is None + free = torch.zeros(n_c + 3, dtype=torch.bool, device=dev) + free[:n_c] = True + free[n_c] = fit_gamma + free[n_c + 1] = fit_sigma_scale + free[n_c + 2] = has_centric + + theta = torch.zeros(n_c + 3, dtype=dtype, device=dev) + excess = float((d2.mean() - log_sig2.exp().mean())) + theta[0] = math.log(max(excess, 0.1 * float(d2.mean()))) + theta[n_c] = float(gamma) if (use_f and gamma is not None) else (1.0 if use_f else 0.0) + + def nll(t): + log_s = basis @ t[:n_c] + t[n_c] * log_f + if cen is not None: + log_s = log_s + t[n_c + 2] * cen + # log V = log(eps S + k^2 sigma^2), formed in log space so neither term can + # underflow the sum. + log_v = torch.logaddexp(log_eps + log_s, 2.0 * t[n_c + 1] + log_sig2) + z = d2 * torch.exp(-log_v) + if robust: + per = 0.5 * log_v + 0.5 * (nu + 1.0) * torch.log1p(z / nu) + else: + per = 0.5 * (log_v + z) + return per.mean() + + idx = torch.nonzero(free, as_tuple=True)[0] + current = float(nll(theta)) + converged = False + lam = 1e-3 + for _ in range(MAX_ITER): + t = theta.detach().requires_grad_(True) + g_full = torch.autograd.grad(nll(t), t, create_graph=True)[0] + h_full = torch.stack( + [torch.autograd.grad(g_full[i], t, retain_graph=True)[0] for i in idx] + ) + g = g_full[idx].detach() + h = h_full[:, idx].detach() + eye = torch.eye(len(idx), dtype=dtype, device=dev) + improved = False + for _ in range(20): + # Levenberg damping on the diagonal: the Hessian can be indefinite far from + # the optimum, where log V is a log-sum-exp of two linear forms. + step = torch.linalg.solve(h + lam * eye * h.diagonal().abs().max(), g) + trial = theta.clone() + trial[idx] = trial[idx] - step + trial[n_c] = trial[n_c].clamp(*_GAMMA_BOUNDS) + new = float(nll(trial)) + if math.isfinite(new) and new <= current: + theta, improved = trial, True + lam = max(lam / 10.0, 1e-9) + break + lam *= 10.0 + if not improved: + break + delta = current - new + current = new + if float(step.abs().max()) < 1e-6 or delta < 1e-10: + converged = True + break + + t = theta.detach().requires_grad_(True) + g_full = torch.autograd.grad(nll(t), t, create_graph=True)[0] + h = torch.stack( + [torch.autograd.grad(g_full[i], t, retain_graph=True)[0][idx] for i in idx] + ).detach() + stderr = torch.full_like(theta, float("nan")) + try: + # The objective is a mean, so the per-reflection Hessian is n_fit times it. + cov = torch.linalg.inv(h) / n_fit + stderr[idx] = cov.diagonal().clamp(min=0.0).sqrt() + except RuntimeError: + pass + + theta = theta.detach() + return DifferencePowerFit( + coeffs=theta[:n_c].clone(), + gamma=float(theta[n_c]) if use_f else 0.0, + sigma_scale=float(theta[n_c + 1].exp()), + centric_factor=float(theta[n_c + 2].exp()) if has_centric else 1.0, + stol_range=stol_range, + amp_scale=amp_scale, + log_f_ref=log_f_ref, + stderr=stderr, + nll=current, + converged=converged, + n_fit=n_fit, + ) + + +def bounded_wiener_weight( + snr: torch.Tensor, snr_floor: float = DEFAULT_SNR_FLOOR +) -> torch.Tensor: + """Wiener weight with the signal-to-noise ratio floored smoothly at ``snr_floor``. + + ``w = (snr + snr_floor) / (snr + 1 + snr_floor)``, which lies in + ``[snr_floor / (1 + snr_floor), 1)``: a reflection is down-weighted by its noise + fraction but never removed. ``snr_floor = 0`` is the plain Wiener weight. + + Parameters + ---------- + snr : torch.Tensor + Per-reflection ``S / (k sigma)**2``, e.g. from :meth:`DifferencePowerFit.snr`. + snr_floor : float + Non-negative floor. + + Returns + ------- + torch.Tensor + Weights, same shape as ``snr``, not normalised. + """ + if snr_floor < 0: + raise ValueError("snr_floor must be non-negative") + return (snr + snr_floor) / (snr + 1.0 + snr_floor) + + +__all__ = [ + "DEFAULT_ORDER", + "DEFAULT_SNR_FLOOR", + "DifferencePowerFit", + "bounded_wiener_weight", + "fit_difference_power", +] From d278d4cf027e9e0fc68541e3a70e1e845c8a3921 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Tue, 29 Sep 2026 10:51:24 +0200 Subject: [PATCH 206/250] Fit the difference SNR on intensities and shrink the extrapolation by it difference_snr fits the expected difference power on I_light - I_dark divided by 2 F_dark when the data carry intensities: French-Wilson amplitude sigmas overstate the noise of a difference, increasingly toward high resolution, while intensity sigmas are the measurement noise. The q weight and the extrapolated FEXT now share this SNR; FEXT shrinks by w = snr / (1 + snr), positive wherever the fitted power is, in place of the per-shell power that zeroed whole resolution ranges. SIGFEXT carries the dark measurement sigma. The sigma scale and centric factor of the fit are bounded so data without noise or signal cannot drive them to zero. Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/integration/test_cli_two_moment_mtz.py | 12 +- tests/unit/maps/test_ded_weights.py | 73 ++++++- .../unit/refinement/test_difference_power.py | 22 ++- torchref/cli/_common.py | 29 +++ torchref/cli/collection_difference_refine.py | 162 +++++++--------- torchref/cli/validate_ded.py | 5 + torchref/maps/ded_weights.py | 181 +++++++++++++++--- .../difference_power.py | 21 +- 8 files changed, 371 insertions(+), 134 deletions(-) diff --git a/tests/integration/test_cli_two_moment_mtz.py b/tests/integration/test_cli_two_moment_mtz.py index cf9b3498..446c631d 100644 --- a/tests/integration/test_cli_two_moment_mtz.py +++ b/tests/integration/test_cli_two_moment_mtz.py @@ -376,9 +376,13 @@ def test_summary_reports_the_activation_moments(self, two_moment_mtz): assert results["lambda_twin"] == pytest.approx(LAMBDA_TWIN) def test_summary_reports_the_shrinkage_diagnostics(self, two_moment_mtz): - """``tau_sq`` and mean ``w(h)`` say whether the default extrapolated map is - over-shrunk, so they belong in the summary rather than only in a print.""" + """The shrinkage's source and mean and least ``w(h)`` say whether the default + extrapolated map is over-shrunk, so they belong in the summary rather than only + in a print. This pair carries intensities, so they supply the SNR. Its dark and + light halves are the same deposited data, so there is no difference signal and + the weights may reach zero here; positivity with signal is pinned in the unit + tests.""" _, summary = two_moment_mtz results = json.loads(summary.read_text())["results"] - assert "tau_sq" in results and "w_shrinkage_mean" in results - assert 0.0 < results["w_shrinkage_mean"] <= 1.0 + assert results["shrinkage_source"] == "intensity" + assert 0.0 <= results["w_shrinkage_min"] <= results["w_shrinkage_mean"] < 1.0 diff --git a/tests/unit/maps/test_ded_weights.py b/tests/unit/maps/test_ded_weights.py index 6bba767a..9bcc9ea7 100644 --- a/tests/unit/maps/test_ded_weights.py +++ b/tests/unit/maps/test_ded_weights.py @@ -4,7 +4,9 @@ ``inverse_variance`` has mean one and floors a zero sigma; ``q`` gives strong reflections more weight than weak ones where inverse variance cannot, keeps every reflection at or above its floor when noise dominates, and falls back to inverse -variance with a warning that names why when too few reflections exist to fit. +variance with a warning that names why when too few reflections exist to fit. The +SNR prefers intensity differences and calibrates their sigmas; the extrapolated +shrinkage keeps every reflection and its weight does not depend on the occupancy. """ import pytest @@ -18,6 +20,7 @@ DedWeightFallbackWarning, all_ded_weights, compute_ded_weights, + difference_snr, normalise_mean_one, ) from torchref.symmetry import SpaceGroup @@ -146,3 +149,71 @@ def test_too_few_reflections_fall_back_to_inverse_variance_with_a_warning(): assert "fallback_reason" in q.diagnostics ivw = compute_ded_weights("inverse_variance", **kw) assert torch.allclose(q.weights, ivw.weights) + + +def _intensity_inputs(n=20000, inflation=1.0, seed=3): + """Dark/light intensities with known measurement noise, and crude amplitudes.""" + g = torch.Generator().manual_seed(seed) + f_dark = (torch.randn(n, generator=g) ** 2 + torch.randn(n, generator=g) ** 2).sqrt() + f_dark = 100.0 * f_dark + f_light = f_dark + torch.randn(n, generator=g) * 5.0 + sig_i = 200.0 + 0.05 * f_dark**2 * torch.exp(0.3 * torch.randn(n, generator=g)) + i_dark = f_dark**2 + torch.randn(n, generator=g) * sig_i + i_light = f_light**2 + torch.randn(n, generator=g) * sig_i + f_d_obs, f_l_obs = i_dark.clamp(min=1.0).sqrt(), i_light.clamp(min=1.0).sqrt() + hkl = torch.randint(-20, 21, (n, 3), generator=g) + return dict( + delta_obs=f_l_obs - f_d_obs, + sigma_diff=torch.full((n,), 5.0), + hkl=hkl, + cell=torch.tensor([40.0, 50.0, 60.0, 90.0, 90.0, 90.0]), + spacegroup=SpaceGroup("P 1"), + f_dark=f_d_obs, + delta_intensity=i_light - i_dark, + sigma_delta_intensity=inflation * sig_i * 2**0.5, + ) + + +@pytest.mark.unit +@pytest.mark.parametrize("inflation", [1.0, 1.5]) +def test_snr_prefers_intensities_and_calibrates_their_sigmas(inflation): + est = difference_snr(**_intensity_inputs(inflation=inflation)) + assert est.source == "intensity" + assert est.fit.gamma == 0.0 + assert est.fit.sigma_scale == pytest.approx(1.0 / inflation, rel=0.1) + assert bool((est.snr > 0).all()) and bool(torch.isfinite(est.snr).all()) + kw = _intensity_inputs(inflation=inflation) + q = compute_ded_weights("q", **kw) + assert q.diagnostics["source"] == "intensity" + del kw["delta_intensity"], kw["sigma_delta_intensity"] + assert difference_snr(**kw).source == "amplitude" + + +@pytest.mark.unit +def test_extrapolated_shrinkage_keeps_every_reflection_and_ignores_occupancy(): + from torchref.cli.collection_difference_refine import ( + compute_bayes_extrapolated_amplitudes, + ) + + g = torch.Generator().manual_seed(0) + n = 1000 + f_dark = 50 + 10 * torch.rand(n, generator=g) + f_light = f_dark + torch.randn(n, generator=g) + phi = torch.zeros(n) + snr = torch.logspace(-3, 3, n) + noise = torch.ones(n) + sig_dark = torch.full((n,), 0.5) + out = { + f: compute_bayes_extrapolated_amplitudes( + f_dark, f_light, sig_dark, phi, phi, f, snr=snr, noise=noise + ) + for f in (0.2, 0.5) + } + for f, (f_ext_b, var, w) in out.items(): + f_ext = f_dark + (f_light - f_dark) / f + assert bool((w > 0).all()) and bool((w < 1).all()) + lo, hi = torch.minimum(f_dark, f_ext), torch.maximum(f_dark, f_ext) + assert bool((f_ext_b >= lo - 1e-4).all()) and bool((f_ext_b <= hi + 1e-4).all()) + assert torch.allclose(var, sig_dark**2 + w * (noise / f) ** 2) + assert bool((var > 0).all()) + assert torch.equal(out[0.2][2], out[0.5][2]) diff --git a/tests/unit/refinement/test_difference_power.py b/tests/unit/refinement/test_difference_power.py index 65bc298b..3ae1159d 100644 --- a/tests/unit/refinement/test_difference_power.py +++ b/tests/unit/refinement/test_difference_power.py @@ -4,14 +4,15 @@ dark-amplitude exponent and the sigma scale are recovered from one dataset, including when the reported sigmas are uniformly inflated; a fixed exponent stays fixed; the bounded Wiener weight never falls below its floor, so no reflection or resolution range -is removed even when the data hold no signal; the fit runs under ``torch.no_grad()`` and -on every available device. +is removed even when the data hold no signal; the sigma scale stays within its bounds when the differences hold no noise; the fit runs +under ``torch.no_grad()`` and on every available device. """ import pytest import torch from torchref.refinement.model_error_estimation.difference_power import ( + SIGMA_SCALE_BOUNDS, bounded_wiener_weight, fit_difference_power, ) @@ -90,3 +91,20 @@ def test_fits_under_no_grad(): with torch.no_grad(): fit = fit_difference_power(s["delta"], s["sigma"], s["dss"], f_dark=s["f"]) assert fit.converged + + +@pytest.mark.unit +def test_sigma_scale_stays_bounded_when_the_differences_hold_no_noise(): + # Identical datasets: every difference is zero, so the fit would drive k to zero. + s = synth(n=5000) + zeros = torch.zeros_like(s["delta"]) + 1e-3 * torch.randn( + len(s["delta"]), generator=torch.Generator().manual_seed(1) + ) + fit = fit_difference_power(zeros, s["sigma"], s["dss"], f_dark=s["f"]) + assert fit.sigma_scale == pytest.approx(SIGMA_SCALE_BOUNDS[0], rel=1e-3) + assert fit.sigma_scale_at_bound + snr = fit.snr(s["sigma"], d_star_sq=s["dss"], f_dark=s["f"]) + assert bool(torch.isfinite(snr).all()) + assert not fit_difference_power( + s["delta"], s["sigma"], s["dss"], f_dark=s["f"] + ).sigma_scale_at_bound diff --git a/torchref/cli/_common.py b/torchref/cli/_common.py index 0be78058..8285df4f 100644 --- a/torchref/cli/_common.py +++ b/torchref/cli/_common.py @@ -419,6 +419,35 @@ def add_ded_weight_args(parser: argparse.ArgumentParser) -> None: ) +def intensity_difference(data_dark, data_light, mask=None): + """``(I_light - I_dark, sigma)`` on the shared scale, or ``(None, None)``. + + Parameters + ---------- + data_dark, data_light : ReflectionData + Datasets on one HKL list, already inter-scaled. + mask : torch.Tensor, optional + Boolean selection applied to both. + + Returns + ------- + tuple of torch.Tensor or None + ``(None, None)`` when either dataset carries no intensities, so callers fall + back to the amplitude differences. + """ + try: + I_dark, sig_dark = data_dark.get_corrected_intensities() + I_light, sig_light = data_light.get_corrected_intensities() + except ValueError: + return None, None + if I_dark is None or I_light is None or sig_dark is None or sig_light is None: + return None, None + if mask is not None: + I_dark, sig_dark = I_dark[mask], sig_dark[mask] + I_light, sig_light = I_light[mask], sig_light[mask] + return I_light - I_dark, (sig_dark**2 + sig_light**2).sqrt() + + def sigma_d_config_from_args(args: argparse.Namespace): """The :class:`~torchref.refinement.model_error_estimation.sigma_d.SigmaDConfig` selected by ``--sigma-d-gamma``.""" diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index 44824ecc..4f942757 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -47,6 +47,7 @@ parse_device_str, parse_weights, register_timing, + intensity_difference, sigma_d_config_from_args, validate_cif_files, validate_files, @@ -55,7 +56,7 @@ DEFAULT_SCHEME, WEIGHT_COLUMNS, all_ded_weights, - reflection_geometry, + difference_snr, ) from torchref.utils.serialization import convert_to_serializable @@ -322,101 +323,62 @@ def setup_loss_state( def compute_bayes_extrapolated_amplitudes( Fobs_dark, Fobs_light, - sig_ext, + sig_dark, phi_dark, phi_mixed, f, *, - tau_sq_floor=1e-4, - epsilon=None, - d_star_sq=None, + snr, + noise, ): - """Empirical Bayes shrinkage estimator for extrapolated SF amplitudes. + """Shrink the extrapolated amplitudes toward ``Fo_dark`` by their signal fraction. - Estimates per-reflection shrinkage weights from the propagated variance of the - extrapolation, then shrinks the phase-aware extrapolated amplitude toward - Fo_dark, regularising noisy high-resolution and weakly-measured reflections:: + The posterior mean of the extrapolated deviation under a Gaussian prior of the + fitted difference power:: - F_ext = |F_dark*e^(iφ_d) + ΔF/f| (phase-aware amplitude) - S(h) = expected power of (F_ext - Fo_dark), per resolution shell - w(h) = S(h) / (S(h) + σ_ext²(h)) - F_extb = w(h)·F_ext + (1-w(h))·Fo_dark (amplitude shrinkage) + F_ext = |F_dark e^(i phi_d) + dF / f| (phase-aware amplitude) + r = F_ext - Fo_dark (~ dF / f) + w = snr / (1 + snr) + F_extb = Fo_dark + w r - With ``d_star_sq`` the signal power comes per resolution shell from - :func:`~torchref.refinement.model_error_estimation.sigma_d.estimate_sigma_d` - (``<(F_ext - Fo_dark)²> - <σ_ext²>`` per shell, shrunk toward a smooth curve); - without it the single global ``τ² = max(<(F_ext - Fo_dark)²> - <σ_ext²>, floor)`` - is used, which is the one-shell special case. + ``r`` is the observed difference scaled by ``1/f``, signal and noise alike, so its + signal-to-noise ratio is the difference's and the occupancy does not enter ``w``. + ``w`` is positive wherever the fitted power is, so no reflection is removed; a noisy + one keeps a small share of its deviation. The result is for viewing: it is biased + toward the dark state by construction and is not a refinement target. Parameters ---------- Fobs_dark, Fobs_light : Tensor (N,) Observed amplitudes. - sig_ext : Tensor (N,) - Propagated uncertainty of the extrapolated amplitude. Taken from the caller - rather than rebuilt here: ``F_ext`` is linear in the observations with - ``dF_ext/dF_light = 1/f`` and ``dF_ext/dF_dark = 1 - 1/f = -(1-f)/f``, so the - dark term carries a ``(1-f)**2`` weight. + sig_dark : Tensor (N,) + Sigma of ``Fobs_dark``, which the shrunk amplitude sits on. phi_dark, phi_mixed : Tensor (N,) Calculated phases (radians) for the dark and mixed models. f : float or Tensor Excited-state population fraction. - tau_sq_floor : float - Floor on the estimated signal variance. - epsilon, d_star_sq : Tensor (N,), optional - Reflection multiplicity and ``1/d**2`` in A^-2. Given ``d_star_sq`` the signal - power is estimated per resolution shell. + snr : Tensor (N,) + Per-reflection signal-to-noise ratio of the difference, from + :func:`torchref.maps.ded_weights.difference_snr`. + noise : Tensor (N,) + Calibrated noise of the amplitude difference, from the same call. Returns ------- tuple - ``(F_ext_bayes, var_ext_bayes, w_shrinkage, tau_sq)`` -- the **shrunk** - extrapolated amplitude, its posterior variance and the shrinkage weight per - reflection, and the count-weighted mean signal variance as a float. + ``(F_ext_bayes, var_ext_bayes, w_shrinkage)`` -- the shrunk extrapolated + amplitude, its variance ``sig_dark**2 + w (noise / f)**2`` (the dark + measurement plus the posterior variance of the deviation) and the weight per + reflection. """ F_dark_phased = Fobs_dark * torch.exp(1j * phi_dark) F_light_phased = Fobs_light * torch.exp(1j * phi_mixed) - delta_F = F_light_phased - F_dark_phased - - sig_sq_ext = sig_ext**2 - - # Phase-aware extrapolated amplitude - F_ext_complex = F_dark_phased + delta_F / f - F_ext = torch.abs(F_ext_complex) - - residual = F_ext - Fobs_dark - if d_star_sq is None: - tau_sq = max( - (residual.square().mean() - sig_sq_ext.mean()).item(), tau_sq_floor - ) - S = torch.full_like(F_ext, tau_sq) - else: - from torchref.refinement.model_error_estimation.sigma_d import ( - estimate_sigma_d, - sigma_d_per_reflection, - ) - - fit = torch.isfinite(residual) & torch.isfinite(sig_ext) - shells = estimate_sigma_d( - residual, sig_ext, epsilon, d_star_sq, None, fit, gamma=0.0 - ) - est = sigma_d_per_reflection(shells, d_star_sq, epsilon, None, sig_ext) - S = est.S.clamp(min=tau_sq_floor) - weight = shells.counts.clamp(min=1.0) - tau_sq = max( - float((shells.Sigma_N * shells.counts).sum() / weight.sum()), tau_sq_floor - ) - - # Per-reflection shrinkage weight (in [0, 1]) - w = S / (S + sig_sq_ext) - - # Posterior variance - var_ext_bayes = (S * sig_sq_ext) / (S + sig_sq_ext) - + F_ext = torch.abs(F_dark_phased + (F_light_phased - F_dark_phased) / f) + w = snr / (1.0 + snr) + var_ext_bayes = sig_dark**2 + w * (noise / f) ** 2 # Shrink the amplitude toward Fo_dark -- scalar, so no phase interference. - F_ext_bayes = w * F_ext + (1 - w) * Fobs_dark - - return F_ext_bayes, var_ext_bayes, w, tau_sq + F_ext_bayes = Fobs_dark + w * (F_ext - Fobs_dark) + return F_ext_bayes, var_ext_bayes, w def _two_moment_columns(mc, dc, mask, fcalc_dark_full, fcalc_mixed_full, @@ -720,18 +682,21 @@ def _extrapolation_columns( phi_dark, ctx, rfree_flags_masked, + snr_est, all_columns=False, verbose=1, - geometry=None, ): - """Extrapolated light-state amplitudes and the map to refine against. + """Extrapolated light-state amplitudes and the map to view them in. + + ``snr_est`` is the :class:`~torchref.maps.ded_weights.DifferenceSNR` of the + light-minus-dark differences on the same reflections. Three constructions of the same quantity, all needing the light model: - ``FEXT`` (default, Bayes-shrunk) - The phase-aware amplitude shrunk toward ``Fo_dark`` by a per-reflection weight - ``w(h) = tau^2 / (tau^2 + sigma_ext^2(h))``, which quiets the weak and - high-resolution reflections where the extrapolation is noisiest. + ``FEXT`` (default, shrunk) + The phase-aware amplitude shrunk toward ``Fo_dark`` by the signal fraction of + each reflection's difference (:func:`compute_bayes_extrapolated_amplitudes`), + which quiets the reflections where the extrapolation is noisiest. ``FEXT_PHASED`` (``all_columns``) The unshrunk phase-aware amplitude. ``FEXT_SCALAR`` (``all_columns``) @@ -771,18 +736,15 @@ def _fit(amp, sig): sig_light_vals**2 + w_dark**2 * sig_dark_vals**2 ) / w_light - eps, dss = geometry if geometry is not None else (None, None) - F_ext_bayes_amp, var_ext_bayes, w_shrinkage, tau_sq = ( - compute_bayes_extrapolated_amplitudes( - Fobs_dark_vals, - Fobs_light_vals, - sig_light_extra, - phi_dark, - ctx["phi_mixed"], - w_light, - epsilon=eps, - d_star_sq=dss, - ) + F_ext_bayes_amp, var_ext_bayes, w_shrinkage = compute_bayes_extrapolated_amplitudes( + Fobs_dark_vals, + Fobs_light_vals, + sig_dark_vals, + phi_dark, + ctx["phi_mixed"], + w_light, + snr=snr_est.snr, + noise=snr_est.noise, ) sig_ext_bayes = torch.sqrt(var_ext_bayes) @@ -803,8 +765,9 @@ def _np(t): if verbose > 0: print(" Bayes extrapolation rfactors:", rfactor_work_free(data_bayes, amp_calc_bayes)) - print(f" Bayes: tau^2 = {tau_sq:.4f}, " - f"mean w(h) = {w_shrinkage.mean().item():.3f}") + print(f" Shrinkage: SNR from {snr_est.source} differences, " + f"mean w(h) = {w_shrinkage.mean().item():.3f}, " + f"min w(h) = {w_shrinkage.min().item():.3g}") if all_columns: amp_phased = torch.abs(F_light_extra) @@ -842,8 +805,9 @@ def _np(t): rfactor_work_free(data_scalar, amp_calc_scalar)) diagnostics = { - "tau_sq": float(tau_sq), + "shrinkage_source": snr_est.source, "w_shrinkage_mean": float(w_shrinkage.mean().item()), + "w_shrinkage_min": float(w_shrinkage.min().item()), } return columns, types, diagnostics @@ -1018,15 +982,19 @@ def write_results_mtz( diff_t = Fobs_light_vals - Fobs_dark_vals sig_diff_t = torch.sqrt(sig_dark_vals**2 + sig_light_vals**2) - all_w = all_ded_weights( + delta_I, sig_delta_I = intensity_difference(data_dark, data_light, mask) + snr_inputs = dict( delta_obs=diff_t, sigma_diff=sig_diff_t, hkl=hkl, cell=data_dark.cell, spacegroup=data_dark.spacegroup, f_dark=Fobs_dark_vals, + delta_intensity=delta_I, + sigma_delta_intensity=sig_delta_I, gamma=sigma_d_config.gamma if sigma_d_config is not None else None, ) + all_w = all_ded_weights(**snr_inputs) selected = all_w[ded_weight] weights = selected.weights.detach().cpu().numpy() diff_Fobs = diff_t.detach().cpu().numpy() @@ -1036,9 +1004,6 @@ def write_results_mtz( for name in WEIGHT_COLUMNS } kscale = scaler.multiplicative_scale()[mask].detach().cpu().numpy() - geometry = reflection_geometry( - hkl, data_dark.cell, data_dark.spacegroup, diff_t.device, diff_t.dtype - ) q_diag = all_w["q"].diagnostics diagnostics = { "ded_weights": { @@ -1053,7 +1018,8 @@ def write_results_mtz( print(f" q-weight fallback: {q_diag['fallback_reason']}") else: print( - f" q-weight fit: gamma = {q_diag['gamma']:.3f}, " + f" q-weight fit on {q_diag['source']} differences: " + f"gamma = {q_diag['gamma']:.3f}, " f"sigma scale k = {q_diag['sigma_scale']:.3f}, " f"centric factor = {q_diag['centric_factor']:.3f}, " f"weights {q_diag['weight_min']:.3f}-{q_diag['weight_max']:.3f} " @@ -1100,9 +1066,9 @@ def write_results_mtz( phi_dark=phi_dark, ctx=ctx, rfree_flags_masked=rfree_flags_masked, + snr_est=difference_snr(**snr_inputs), all_columns=all_columns, verbose=verbose, - geometry=geometry, ) diagnostics.update(ext_diagnostics) columns.update(ext_cols) diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index d00c78bd..2cf3c73a 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -38,6 +38,7 @@ add_outdir_arg, build_dual_column_names, configure_unbuffered_output, + intensity_difference, load_model, load_reflection_data, parse_device_str, @@ -293,9 +294,12 @@ def setup_ded_context( # Difference Fo and the registered weights; the selected scheme is the headline. dfo = F_light - F_dark sig_diff = torch.sqrt(sig_dark**2 + sig_light**2) + delta_I, sig_delta_I = intensity_difference(data_dark, data_light, refl_mask) all_w = all_ded_weights( delta_obs=dfo, sigma_diff=sig_diff, + delta_intensity=delta_I, + sigma_delta_intensity=sig_delta_I, hkl=hkl, cell=data_dark.cell, spacegroup=data_dark.spacegroup, @@ -727,6 +731,7 @@ def run_validation(args): **{ k: ctx["ded_weight_diagnostics"].get(k) for k in ( + "source", "gamma", "gamma_fitted", "sigma_scale", diff --git a/torchref/maps/ded_weights.py b/torchref/maps/ded_weights.py index 0c9348fd..60cdb418 100644 --- a/torchref/maps/ded_weights.py +++ b/torchref/maps/ded_weights.py @@ -11,15 +11,16 @@ estimates of one quantity. On French-Wilson amplitudes it *up*-weights the weak high-resolution reflections, whose posterior sigma the prior keeps small. ``q`` - The q-weight: a Wiener weight ``(snr + b) / (snr + 1 + b)`` with - ``snr = S / (k sigma_diff)**2``, the default. ``S`` and the sigma scale ``k`` come - from :func:`~torchref.refinement.model_error_estimation.difference_power. - fit_difference_power`, a per-reflection maximum-likelihood fit of ``log S`` as a - Chebyshev series in resolution plus ``gamma log F_dark``: no resolution shells. The - floor ``b`` keeps a reflection without signal at ``b / (1 + b)`` of the full weight - (one third at the default), so a noisy resolution range is down-weighted, never - removed. ``k`` absorbs a uniform miscalibration of the sigmas; French-Wilson sigmas - of two datasets overstate the error of their difference, so ``k < 1`` is normal. + The q-weight: a Wiener weight ``(snr + b) / (snr + 1 + b)`` on the per-reflection + signal-to-noise ratio of :func:`difference_snr`, the default. The floor ``b`` keeps + a reflection without signal at ``b / (1 + b)`` of the full weight (one third at the + default), so a noisy resolution range is down-weighted, never removed. + +:func:`difference_snr` is the one estimate of how much of each observed difference is +signal; the extrapolated amplitudes of ``torchref.difference-map`` shrink by the same +ratio. It prefers intensities: French-Wilson amplitude sigmas describe the posterior of +one amplitude, not the noise of a difference between two, and overstate it increasingly +toward high resolution, where intensity sigmas are the measurement noise itself. Plain tensors in and out. The weights live on the device of ``delta_obs``. The difference-power fit is imported inside the scheme that needs it so that @@ -102,6 +103,128 @@ def reflection_geometry(hkl, cell, spacegroup, device, dtype): return eps, dss +@dataclass(frozen=True) +class DifferenceSNR: + """Per-reflection signal-to-noise ratio of observed differences. + + Attributes + ---------- + snr : torch.Tensor + ``S / noise**2``, shape ``(N,)``, the same in amplitude and intensity terms. + noise : torch.Tensor + Calibrated noise ``k * sigma`` of the amplitude difference, shape ``(N,)``, in + the amplitude units of ``delta_obs``. + fit : DifferencePowerFit + The fit ``snr`` came from. + source : str + ``"intensity"`` or ``"amplitude"``: which observations were fitted. + """ + + snr: torch.Tensor + noise: torch.Tensor + fit: object + source: str + + +def difference_snr( + *, + delta_obs: torch.Tensor, + sigma_diff: torch.Tensor, + hkl: torch.Tensor, + cell, + spacegroup, + f_dark: torch.Tensor | None = None, + delta_intensity: torch.Tensor | None = None, + sigma_delta_intensity: torch.Tensor | None = None, + fit_mask: torch.Tensor | None = None, + gamma: float | None = None, +) -> DifferenceSNR: + """Fit the expected difference power and return each reflection's SNR. + + Given intensity differences and ``f_dark``, the fit runs on + ``delta_intensity / (2 F_dark)`` with sigma ``sigma_delta_intensity / (2 F_dark)``: + the ratio of an intensity difference to its measurement noise, carried on the + amplitude scale so the likelihood is as well conditioned as an amplitude fit. The + ``F_dark`` exponent then defaults to 0, because ``S`` and ``F_dark`` enter the + divided difference together and the fitted exponent is not identified. Otherwise + the amplitude differences and their sigmas are fitted. + + Parameters + ---------- + delta_obs, sigma_diff : torch.Tensor + Signed amplitude differences and their propagated sigma, shape ``(N,)``. + hkl : torch.Tensor + Miller indices, shape ``(N, 3)``. + cell : Cell or torch.Tensor + Unit cell, or its six parameters in A and degrees. + spacegroup : SpaceGroup or None + For the multiplicity and centric flags; ``None`` means P1. + f_dark : torch.Tensor, optional + Dark amplitudes, shape ``(N,)``; required for the intensity path. + delta_intensity, sigma_delta_intensity : torch.Tensor, optional + Intensity differences ``I_light - I_dark`` and their propagated sigma, shape + ``(N,)``, on the scale whose square root is the amplitude scale of + ``delta_obs``. + fit_mask : torch.Tensor, optional + Reflections entering the fit; default every finite one. + gamma : float, optional + Fix the ``F_dark`` exponent. + + Returns + ------- + DifferenceSNR + + Raises + ------ + ValueError + If too few reflections are usable for the fit. + """ + from torchref.refinement.model_error_estimation.difference_power import ( + fit_difference_power, + ) + + d = delta_obs.reshape(-1) + sig = sigma_diff.reshape(-1).to(d.device, d.dtype) + eps, dss = reflection_geometry(hkl, cell, spacegroup, d.device, d.dtype) + f = f_dark.reshape(-1).to(d.device, d.dtype) if f_dark is not None else None + centric = ( + spacegroup.is_centric(torch.as_tensor(hkl, device=d.device)).to(d.device) + if spacegroup is not None + else None + ) + use_intensity = ( + delta_intensity is not None and sigma_delta_intensity is not None and f is not None + ) + if use_intensity: + # A floor on the divisor: a near-zero dark amplitude would send both terms to + # infinity. Their ratio, the SNR, does not depend on it. + ok_f = torch.isfinite(f) & (f > 0) + f_floor = 0.1 * float(f[ok_f].median()) if bool(ok_f.any()) else 1.0 + two_f = 2.0 * f.clamp(min=f_floor) + values = delta_intensity.reshape(-1).to(d.device, d.dtype) / two_f + sigma = sigma_delta_intensity.reshape(-1).to(d.device, d.dtype) / two_f + g = 0.0 if gamma is None else gamma + else: + values, sigma, g = d, sig, gamma + fit = fit_difference_power( + values, + sigma, + dss, + epsilon=eps, + f_dark=f, + centric=centric, + fit_mask=fit_mask, + gamma=g, + ) + snr = fit.snr(sigma, d_star_sq=dss, epsilon=eps, f_dark=f, centric=centric) + return DifferenceSNR( + snr=snr, + noise=fit.sigma_scale * sigma, + fit=fit, + source="intensity" if use_intensity else "amplitude", + ) + + def _inverse_variance(sigma_diff: torch.Tensor) -> torch.Tensor: """``1 / sigma**2`` with sigma floored at a tenth of its median, so a reported zero uncertainty gives a large finite weight rather than an infinite one.""" @@ -125,6 +248,8 @@ def compute_ded_weights( cell, spacegroup, f_dark: torch.Tensor | None = None, + delta_intensity: torch.Tensor | None = None, + sigma_delta_intensity: torch.Tensor | None = None, fit_mask: torch.Tensor | None = None, gamma: float | None = None, snr_floor: float | None = None, @@ -147,6 +272,9 @@ def compute_ded_weights( For the reflection multiplicity and centric flags; ``None`` means P1. f_dark : torch.Tensor, optional Dark amplitudes, shape ``(N,)``, for the ``q`` power law in ``F_dark``. + delta_intensity, sigma_delta_intensity : torch.Tensor, optional + Intensity differences and their sigma; given with ``f_dark``, the ``q`` SNR is + fitted on them (see :func:`difference_snr`). fit_mask : torch.Tensor, optional Reflections entering the ``q`` fit; default every finite one. gamma : float, optional @@ -176,26 +304,19 @@ def compute_ded_weights( from torchref.refinement.model_error_estimation.difference_power import ( DEFAULT_SNR_FLOOR, bounded_wiener_weight, - fit_difference_power, ) floor = DEFAULT_SNR_FLOOR if snr_floor is None else float(snr_floor) - d = delta_obs.reshape(-1) - eps, dss = reflection_geometry(hkl, cell, spacegroup, d.device, d.dtype) - f = f_dark.reshape(-1).to(d.device, d.dtype) if f_dark is not None else None - centric = ( - spacegroup.is_centric(torch.as_tensor(hkl, device=d.device)).to(d.device) - if spacegroup is not None - else None - ) try: - fit = fit_difference_power( - d, - sigma_diff, - dss, - epsilon=eps, - f_dark=f, - centric=centric, + est = difference_snr( + delta_obs=delta_obs, + sigma_diff=sigma_diff, + hkl=hkl, + cell=cell, + spacegroup=spacegroup, + f_dark=f_dark, + delta_intensity=delta_intensity, + sigma_delta_intensity=sigma_delta_intensity, fit_mask=fit_mask, gamma=gamma, ) @@ -212,13 +333,15 @@ def compute_ded_weights( normalise_mean_one(_inverse_variance(sigma_diff)), {"fallback_reason": reason}, ) - snr = fit.snr(sigma_diff, d_star_sq=dss, epsilon=eps, f_dark=f, centric=centric) - w = bounded_wiener_weight(snr, floor) + fit = est.fit + w = bounded_wiener_weight(est.snr, floor) ok = torch.isfinite(w) diagnostics = { + "source": est.source, "gamma": fit.gamma, - "gamma_fitted": gamma is None and f is not None, + "gamma_fitted": gamma is None and est.source == "amplitude" and f_dark is not None, "sigma_scale": fit.sigma_scale, + "sigma_scale_at_bound": fit.sigma_scale_at_bound, "centric_factor": fit.centric_factor, "order": len(fit.coeffs) - 1, "coeffs": fit.coeffs.detach().cpu().tolist(), @@ -248,8 +371,10 @@ def all_ded_weights(**kwargs) -> dict[str, DedWeights]: "WEIGHT_COLUMNS", "DedWeightFallbackWarning", "DedWeights", + "DifferenceSNR", "all_ded_weights", "compute_ded_weights", + "difference_snr", "normalise_mean_one", "reflection_geometry", ] diff --git a/torchref/refinement/model_error_estimation/difference_power.py b/torchref/refinement/model_error_estimation/difference_power.py index 2f81b76e..a2552c7a 100644 --- a/torchref/refinement/model_error_estimation/difference_power.py +++ b/torchref/refinement/model_error_estimation/difference_power.py @@ -50,6 +50,13 @@ #: Newton iterations; the problem has at most eight parameters and converges in ~10. MAX_ITER = 60 _GAMMA_BOUNDS = (-1.0, 3.0) +#: Bounds on the sigma scale ``k``. No merge misreports its sigmas tenfold; outside +#: these the data hold no noise to calibrate against (identical datasets drive ``k`` +#: to zero), and a zero noise would give an infinite SNR and zero sigmas downstream. +SIGMA_SCALE_BOUNDS = (0.1, 10.0) +_LOG_K_BOUNDS = tuple(math.log(b) for b in SIGMA_SCALE_BOUNDS) +# Bounds on the log centric factor, so a fit without signal cannot underflow it to zero. +_LOG_CENTRIC_BOUNDS = (-7.0, 7.0) @dataclass(frozen=True) @@ -64,7 +71,11 @@ class DifferencePowerFit: gamma : float Exponent on ``F_dark``; ``0`` when no dark amplitude was used. sigma_scale : float - The factor ``k`` on the reported sigmas; ``1`` when not fitted. + The factor ``k`` on the reported sigmas, within :data:`SIGMA_SCALE_BOUNDS`; + ``1`` when not fitted. + sigma_scale_at_bound : bool + Whether ``k`` stopped at a bound: the reported sigmas and the scatter of the + differences disagree beyond any plausible miscalibration. centric_factor : float Power of a centric reflection relative to an acentric one at equal resolution; ``1`` when no centric flags were given. @@ -89,6 +100,7 @@ class DifferencePowerFit: coeffs: torch.Tensor gamma: float sigma_scale: float + sigma_scale_at_bound: bool centric_factor: float stol_range: tuple amp_scale: float @@ -305,6 +317,8 @@ def nll(t): trial = theta.clone() trial[idx] = trial[idx] - step trial[n_c] = trial[n_c].clamp(*_GAMMA_BOUNDS) + trial[n_c + 1] = trial[n_c + 1].clamp(*_LOG_K_BOUNDS) + trial[n_c + 2] = trial[n_c + 2].clamp(*_LOG_CENTRIC_BOUNDS) new = float(nll(trial)) if math.isfinite(new) and new <= current: theta, improved = trial, True @@ -337,6 +351,10 @@ def nll(t): coeffs=theta[:n_c].clone(), gamma=float(theta[n_c]) if use_f else 0.0, sigma_scale=float(theta[n_c + 1].exp()), + sigma_scale_at_bound=bool( + fit_sigma_scale + and min(abs(float(theta[n_c + 1]) - b) for b in _LOG_K_BOUNDS) < 1e-4 + ), centric_factor=float(theta[n_c + 2].exp()) if has_centric else 1.0, stol_range=stol_range, amp_scale=amp_scale, @@ -377,6 +395,7 @@ def bounded_wiener_weight( __all__ = [ "DEFAULT_ORDER", "DEFAULT_SNR_FLOOR", + "SIGMA_SCALE_BOUNDS", "DifferencePowerFit", "bounded_wiener_weight", "fit_difference_power", From f4852aa6f3f2007b335fff66b089dbc297e606b8 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Tue, 29 Sep 2026 11:12:26 +0200 Subject: [PATCH 207/250] Fit the difference target's coupling and model error without shells fit_difference_power takes an optional model difference: the mean becomes alpha * dF_calc with alpha a low-order Chebyshev series in resolution, and S the power the model leaves unexplained. CollectionDifferenceSigmaDTarget takes alpha and beta_model from one such fit on the free reflections, with the reported sigmas taken as calibrated as the likelihood itself uses them, in place of the per-shell cross moments of the sigma_D estimator. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../test_collection_sigma_d_target.py | 50 ++++--- .../unit/refinement/test_difference_power.py | 32 ++++- .../difference_power.py | 63 ++++++++- .../refinement/targets/collection/xray.py | 124 ++++++++++-------- 4 files changed, 185 insertions(+), 84 deletions(-) diff --git a/tests/unit/refinement/test_collection_sigma_d_target.py b/tests/unit/refinement/test_collection_sigma_d_target.py index dddb57d2..f7998848 100644 --- a/tests/unit/refinement/test_collection_sigma_d_target.py +++ b/tests/unit/refinement/test_collection_sigma_d_target.py @@ -1,16 +1,19 @@ """The ``difference_sd`` collection row on a real dark/light pair. Pinned on 1DAW with a 0.2 A shifted light model: the loss is finite, gradients reach -the light model through ``dF_calc`` only, the sigma_D estimate is owned by the target, -fitted on free reflections of the timepoint row, cached across forwards and cleared -by ``maintenance()``, and the fit summary reaches ``stats()``. +the light model through ``dF_calc`` only, the difference-power fit is owned by the +target, fitted on free reflections of the timepoint row with the reported sigmas taken +as calibrated, gives a positive unexplained power everywhere, is cached across forwards +and cleared by ``maintenance()``, and the fit summary reaches ``stats()``. """ import pytest import torch -from torchref.refinement.model_error_estimation import sigma_d as sigma_d_module -from torchref.refinement.model_error_estimation.sigma_d import SigmaDEstimator +from torchref.refinement.model_error_estimation.difference_power import ( + DifferencePowerFit, +) +from torchref.refinement.targets.collection import xray as xray_module from torchref.refinement.targets.collection import ( CollectionDifferenceSigmaDTarget, CollectionSigmaDLossInputs, @@ -66,18 +69,18 @@ def target(loaded_reflection_data, sample_structure_pair): return dc, mc, CollectionDifferenceSigmaDTarget(dc, mc, scaler=scaler) -def test_forward_is_finite_and_owns_its_estimator(target): +def test_forward_is_finite_and_owns_its_fit(target): _dc, _mc, t = target - assert isinstance(t._sigma_d, SigmaDEstimator) loss = t.forward() assert torch.isfinite(loss) + assert isinstance(t._fit, DifferencePowerFit) + assert t._fit.sigma_scale == 1.0 and len(t._fit.alpha_coeffs) > 0 ctx = t._loss_inputs() assert isinstance(ctx, CollectionSigmaDLossInputs) assert ctx.alpha.shape == ctx.beta_model.shape == (ctx.obs.shape[1],) assert not ctx.alpha.requires_grad and not ctx.beta_model.requires_grad - assert (ctx.beta_model >= 0).all() - shells = t._sigma_d.shells - assert shells.has_model and not shells.all_zero and not shells.degenerate + assert torch.isfinite(ctx.alpha).all() and torch.isfinite(ctx.beta_model).all() + assert (ctx.beta_model > 0).all() def test_gradient_reaches_the_light_model(target): @@ -89,28 +92,28 @@ def test_gradient_reaches_the_light_model(target): assert grads and any(torch.isfinite(g).all() and g.abs().sum() > 0 for g in grads) -def test_estimate_is_cached_until_maintenance(target): +def test_fit_is_cached_until_maintenance(target): _dc, _mc, t = target t.forward() - assert t._sigma_d._cache is not None - first = t._sigma_d._cache + first = t._fit + assert first is not None t.forward() - assert t._sigma_d._cache is first + assert t._fit is first t.maintenance() - assert t._sigma_d._cache is None + assert t._fit is None def test_fit_uses_free_reflections_of_the_timepoint_row(target, monkeypatch): dc, _mc, t = target seen = {} - real = sigma_d_module.estimate_sigma_d + real = xray_module.fit_difference_power - def spy(delta_obs, sigma_diff, epsilon, d_star_sq, f_dark, fit_mask, **kw): - seen["fit_mask"] = fit_mask.clone() + def spy(delta_obs, sigma_diff, d_star_sq, **kw): + seen["fit_mask"] = kw["fit_mask"].clone() seen["n"] = delta_obs.numel() - return real(delta_obs, sigma_diff, epsilon, d_star_sq, f_dark, fit_mask, **kw) + return real(delta_obs, sigma_diff, d_star_sq, **kw) - monkeypatch.setattr(sigma_d_module, "estimate_sigma_d", spy) + monkeypatch.setattr(xray_module, "fit_difference_power", spy) t.forward() n_hkl = dc.hkl.shape[0] assert seen["n"] == n_hkl @@ -123,4 +126,9 @@ def test_stats_carry_the_fit_summary(target): _dc, _mc, t = target t.forward() stats = t.stats() - assert "sigma_d_gamma" in stats and "sigma_d_tau" in stats + for key in ( + "difference_gamma", + "difference_alpha_low_res", + "difference_alpha_high_res", + ): + assert key in stats diff --git a/tests/unit/refinement/test_difference_power.py b/tests/unit/refinement/test_difference_power.py index 3ae1159d..39751ec6 100644 --- a/tests/unit/refinement/test_difference_power.py +++ b/tests/unit/refinement/test_difference_power.py @@ -2,10 +2,12 @@ Pinned on seeded synthetic differences with a known power law: the power, the dark-amplitude exponent and the sigma scale are recovered from one dataset, including -when the reported sigmas are uniformly inflated; a fixed exponent stays fixed; the -bounded Wiener weight never falls below its floor, so no reflection or resolution range -is removed even when the data hold no signal; the sigma scale stays within its bounds when the differences hold no noise; the fit runs -under ``torch.no_grad()`` and on every available device. +when the reported sigmas are uniformly inflated; a fixed exponent stays fixed; a model +difference's resolution-dependent coupling and the power it leaves unexplained are +recovered together; the bounded Wiener weight never falls below its floor, so no +reflection or resolution range is removed even when the data hold no signal; the sigma +scale stays within its bounds when the differences hold no noise; the fit runs under +``torch.no_grad()`` and on every available device. """ import pytest @@ -108,3 +110,25 @@ def test_sigma_scale_stays_bounded_when_the_differences_hold_no_noise(): assert not fit_difference_power( s["delta"], s["sigma"], s["dss"], f_dark=s["f"] ).sigma_scale_at_bound + + +@pytest.mark.unit +def test_recovers_the_model_coupling_and_the_unexplained_power(): + g = torch.Generator().manual_seed(4) + s = synth() + n = len(s["delta"]) + stol = s["dss"].sqrt() / 2.0 + alpha_true = 0.8 - 0.3 * stol / stol.max() + # The model explains part of the difference; the rest is the planted s_true. + delta_calc = torch.randn(n, generator=g) * 3.0 + delta = s["delta"] + alpha_true * delta_calc + fit = fit_difference_power( + delta, s["sigma"], s["dss"], f_dark=s["f"], delta_calc=delta_calc + ) + assert fit.converged + alpha = fit.alpha_at(s["dss"]) + assert float((alpha - alpha_true).abs().max()) < 0.05 + beta = fit.signal_power(s["dss"], f_dark=s["f"]) + assert float((beta / s["s_true"]).log().abs().median()) < LOG_POWER_ATOL + no_model = fit_difference_power(s["delta"], s["sigma"], s["dss"], f_dark=s["f"]) + assert torch.equal(no_model.alpha_at(s["dss"]), torch.ones_like(s["dss"])) diff --git a/torchref/refinement/model_error_estimation/difference_power.py b/torchref/refinement/model_error_estimation/difference_power.py index a2552c7a..bf50f46a 100644 --- a/torchref/refinement/model_error_estimation/difference_power.py +++ b/torchref/refinement/model_error_estimation/difference_power.py @@ -22,6 +22,11 @@ many-fold between reflections at one resolution while ``S`` does not; it is what keeps inflated sigmas from reading as an absence of signal. +Given a model difference ``dF_calc``, the mean becomes :math:`\\alpha(x_h) dF_{calc,h}` +with :math:`\\alpha` a low-order Chebyshev series in the same abscissa, and ``S`` is the +power the model does not explain -- the coupling and unexplained power a difference +likelihood needs, fitted in one pass instead of from per-shell cross moments. + :func:`bounded_wiener_weight` turns a fit into a weight that down-weights noisy reflections but never removes one, the resolution-continuous counterpart of the q-weight's floor. @@ -45,6 +50,9 @@ #: q-weight's noise-only limit (``S`` floored at half the raw difference power gives #: ``w = 1/3``), so a reflection without signal keeps a third of the full weight. DEFAULT_SNR_FLOOR = 0.5 +#: Chebyshev order of the model coupling ``alpha``. Quadratic follows the fall of the +#: coupling with resolution, which is smooth and far less structured than ``S``. +DEFAULT_ALPHA_ORDER = 2 #: Degrees of freedom of the Student-t likelihood when ``robust=True``. DEFAULT_NU = 4.0 #: Newton iterations; the problem has at most eight parameters and converges in ~10. @@ -86,9 +94,12 @@ class DifferencePowerFit: amp_scale, log_f_ref : float The amplitude unit the fit was done in and the mean ``log F_dark`` it was centred on. + alpha_coeffs : torch.Tensor + Chebyshev coefficients of the model coupling ``alpha``; empty when no model + difference was fitted. Evaluate with :meth:`alpha_at`. stderr : torch.Tensor - Standard errors of ``(coeffs, gamma, log k, log centric_factor)`` from the - inverse Hessian; NaN where a parameter was fixed. + Standard errors of ``(coeffs, gamma, log k, log centric_factor, alpha_coeffs)`` + from the inverse Hessian; NaN where a parameter was fixed. nll : float Mean negative log-likelihood per reflection at the optimum. converged : bool @@ -103,6 +114,7 @@ class DifferencePowerFit: sigma_scale_at_bound: bool centric_factor: float stol_range: tuple + alpha_coeffs: torch.Tensor amp_scale: float log_f_ref: float stderr: torch.Tensor @@ -150,6 +162,14 @@ def signal_power( power = power * epsilon.to(dss) return power + def alpha_at(self, d_star_sq: torch.Tensor) -> torch.Tensor: + """Model coupling ``alpha`` at each ``1/d**2`` (A^-2); ones without a model.""" + dss = d_star_sq.to(self.coeffs) + if len(self.alpha_coeffs) == 0: + return torch.ones_like(dss) + n = len(self.alpha_coeffs) + return _design(dss, n, self.stol_range) @ self.alpha_coeffs + def snr(self, sigma: torch.Tensor, **kwargs) -> torch.Tensor: """``signal_power / (k * sigma)**2`` per reflection; keywords as :meth:`signal_power`.""" @@ -182,7 +202,9 @@ def fit_difference_power( f_dark: torch.Tensor | None = None, centric: torch.Tensor | None = None, fit_mask: torch.Tensor | None = None, + delta_calc: torch.Tensor | None = None, order: int = DEFAULT_ORDER, + alpha_order: int = DEFAULT_ALPHA_ORDER, gamma: float | None = None, fit_sigma_scale: bool = True, robust: bool = False, @@ -205,8 +227,13 @@ def fit_difference_power( Boolean centric flags; given, a centric power factor is fitted. fit_mask : torch.Tensor, optional Reflections entering the fit; default every finite one with positive sigma. + delta_calc : torch.Tensor, optional + Model difference, shape ``(N,)``, on the scale of ``delta_obs``. Given, the + mean is ``alpha * delta_calc`` and ``S`` is the unexplained power. order : int Chebyshev order of ``log S`` in ``sin(theta)/lambda``. + alpha_order : int + Chebyshev order of ``alpha``; used only with ``delta_calc``. gamma : float, optional Fix the ``F_dark`` exponent instead of fitting it. fit_sigma_scale : bool @@ -242,6 +269,9 @@ def fit_difference_power( if use_f: f = f_dark.detach().reshape(-1).to(dev, dtype) ok = ok & torch.isfinite(f) + if delta_calc is not None: + c_all = delta_calc.detach().reshape(-1).to(dev, dtype) + ok = ok & torch.isfinite(c_all) n_fit = int(ok.sum()) if n_fit < order + 4: raise ValueError(f"need at least {order + 4} usable reflections, got {n_fit}") @@ -250,7 +280,8 @@ def fit_difference_power( # Work in units of the rms difference so every term of the likelihood is O(1) and # the Hessian is well conditioned in float32. amp_scale = float(d.square().mean().sqrt().clamp(min=1e-12)) - d2 = (d / amp_scale).square() + d_std = d / amp_scale + d2 = d_std.square() log_sig2 = 2.0 * torch.log(sig / amp_scale) log_eps = ( torch.log(epsilon.reshape(-1).to(dev, dtype)[ok]) @@ -270,14 +301,26 @@ def fit_difference_power( cen = centric.reshape(-1).to(dev)[ok].to(dtype) if has_centric else None n_c = order + 1 + n_a = alpha_order + 1 if delta_calc is not None else 0 + i_a = n_c + 3 + if n_a: + c_std = c_all[ok] / amp_scale + basis_a = _design(dss, n_a, stol_range) fit_gamma = use_f and gamma is None - free = torch.zeros(n_c + 3, dtype=torch.bool, device=dev) + free = torch.zeros(i_a + n_a, dtype=torch.bool, device=dev) free[:n_c] = True free[n_c] = fit_gamma free[n_c + 1] = fit_sigma_scale free[n_c + 2] = has_centric - - theta = torch.zeros(n_c + 3, dtype=dtype, device=dev) + free[i_a:] = True + + theta = torch.zeros(i_a + n_a, dtype=dtype, device=dev) + if n_a: + # Start from the global least-squares coupling, so the power starts from the + # residual rather than from the whole difference. + cc = float(c_std.square().mean()) + theta[i_a] = float((d_std * c_std).mean()) / cc if cc > 0 else 1.0 + d2 = (d_std - theta[i_a] * c_std).square() excess = float((d2.mean() - log_sig2.exp().mean())) theta[0] = math.log(max(excess, 0.1 * float(d2.mean()))) theta[n_c] = float(gamma) if (use_f and gamma is not None) else (1.0 if use_f else 0.0) @@ -289,7 +332,11 @@ def nll(t): # log V = log(eps S + k^2 sigma^2), formed in log space so neither term can # underflow the sum. log_v = torch.logaddexp(log_eps + log_s, 2.0 * t[n_c + 1] + log_sig2) - z = d2 * torch.exp(-log_v) + if n_a: + resid2 = (d_std - (basis_a @ t[i_a:]) * c_std).square() + else: + resid2 = d2 + z = resid2 * torch.exp(-log_v) if robust: per = 0.5 * log_v + 0.5 * (nu + 1.0) * torch.log1p(z / nu) else: @@ -357,6 +404,7 @@ def nll(t): ), centric_factor=float(theta[n_c + 2].exp()) if has_centric else 1.0, stol_range=stol_range, + alpha_coeffs=theta[i_a:].clone(), amp_scale=amp_scale, log_f_ref=log_f_ref, stderr=stderr, @@ -393,6 +441,7 @@ def bounded_wiener_weight( __all__ = [ + "DEFAULT_ALPHA_ORDER", "DEFAULT_ORDER", "DEFAULT_SNR_FLOOR", "SIGMA_SCALE_BOUNDS", diff --git a/torchref/refinement/targets/collection/xray.py b/torchref/refinement/targets/collection/xray.py index 2b4587ec..0e6a40dd 100644 --- a/torchref/refinement/targets/collection/xray.py +++ b/torchref/refinement/targets/collection/xray.py @@ -36,10 +36,11 @@ SigmaAEstimator, epsilon_from_hkl, ) -from torchref.refinement.model_error_estimation.sigma_d import ( - SigmaDConfig, - SigmaDEstimator, +from torchref.refinement.model_error_estimation.difference_power import ( + DifferencePowerFit, + fit_difference_power, ) +from torchref.refinement.model_error_estimation.sigma_d import SigmaDConfig from torchref.utils.stats import VERBOSITY_STANDARD, StatEntry, stat from ._util import common_geom @@ -196,27 +197,30 @@ class CollectionDifferenceIntensityTarget(CollectionDifferenceTarget): class CollectionDifferenceSigmaDTarget(CollectionDifferenceTarget): - """The difference-from-mean Gaussian with a sigma_D error model. + """The difference-from-mean Gaussian with a fitted model-error term. The parent compares ``dF_obs`` with ``dF_calc`` under the measurement variance alone. Here the likelihood is centred on ``alpha * dF_calc`` and its variance is ``beta_model + sigma_diff**2``: ``alpha`` is the Gaussian coupling of the model difference to the true one and ``beta_model`` the difference power the model does - not explain, both per resolution shell from - :class:`~torchref.refinement.model_error_estimation.sigma_d.SigmaDEstimator` fitted - on the pooled **free** reflections of the timepoint rows, with the dark-amplitude - power law carried per reflection. A poor light model therefore inflates the - variance where it fails instead of pulling the coordinates toward noise. + not explain. Both come from one + :func:`~torchref.refinement.model_error_estimation.difference_power. + fit_difference_power` on the pooled **free** reflections of the timepoint rows -- + ``alpha`` a smooth function of resolution, ``beta_model`` a smooth function of + resolution times the dark-amplitude power law, no resolution shells. The reported + sigmas are taken as calibrated (``k = 1``), as the likelihood itself uses them. A + poor light model therefore inflates the variance where it fails instead of pulling + the coordinates toward noise. At ``N = 2`` the timepoint row's difference from the mean is half the dark - subtraction; ``S`` and ``sigma_diff**2`` scale together, so the estimate is - invariant to that factor. The estimate is cached until :meth:`maintenance`, which - ``LossState`` calls after each optimizer-step block. + subtraction; ``beta_model`` and ``sigma_diff**2`` scale together and ``alpha`` is a + ratio, so the fit is invariant to that factor. The fit is cached until + :meth:`maintenance`, which ``LossState`` calls after each optimizer-step block. Parameters ---------- sigma_d_config : SigmaDConfig, optional - Exponent and shrinkage settings; the module defaults when omitted. + Its ``gamma`` fixes the dark-amplitude exponent; fitted when omitted. """ name: str = "difference_sigma_d_xray" @@ -241,20 +245,25 @@ def __init__( use_set=use_set, verbose=verbose, ) - # Constructed once; the cache lives until maintenance() resets it. - self._sigma_d = SigmaDEstimator(sigma_d_config) + self._config = sigma_d_config if sigma_d_config is not None else SigmaDConfig() + # Fitted on first use; maintenance() clears it. + self._fit: DifferencePowerFit = None self._eps_common: torch.Tensor = None self._dss_common: torch.Tensor = None + self._centric_common: torch.Tensor = None self._geom_key: int = None def _common_geom(self): - """``(epsilon, d_star_sq)`` on the common HKL, cached per dark dataset.""" + """``(epsilon, d_star_sq, centric)`` on the common HKL, cached per dark + dataset.""" data = self._dataset_collection[self._model_collection.dark_key] key = id(data) if self._eps_common is None or self._geom_key != key: self._eps_common, self._dss_common = common_geom(data) + sg = getattr(data, "spacegroup", None) + self._centric_common = sg.is_centric(data.hkl) if sg is not None else None self._geom_key = key - return self._eps_common, self._dss_common + return self._eps_common, self._dss_common, self._centric_common @staticmethod def _difference_terms(ctx): @@ -269,41 +278,51 @@ def _difference_terms(ctx): def _loss_inputs(self, recalc: bool = False): """The parent's stack plus ``alpha`` and ``beta_model`` on the common HKL. - The estimator sees the timepoint rows only (the dark row is the reference the + The fit sees the timepoint rows only (the dark row is the reference the differences are taken against), their free reflections, and a detached model difference, so gradients reach the models only through ``ctx.model``. """ ctx = super()._loss_inputs(recalc=recalc) - delta_obs, delta_calc, sigma_diff = self._difference_terms(ctx) - delta_calc = delta_calc.detach() dark = ctx.keys.index(self._model_collection.dark_key) - rows = [i for i in range(len(ctx.keys)) if i != dark] or [dark] - dc = self._dataset_collection - eps, dss = self._common_geom() dtype = ctx.obs.dtype + eps, dss, centric = self._common_geom() eps, dss = eps.to(dtype), dss.to(dtype) f_dark = ctx.obs[dark] - # The free set, independent of this target's own subset; the estimator drops - # non-finite observations itself. - fit_mask = torch.cat( - [dc[ctx.keys[i]].free.mask.to(ctx.mask.device) for i in rows] + if self._fit is None: + delta_obs, delta_calc, sigma_diff = self._difference_terms(ctx) + rows = [i for i in range(len(ctx.keys)) if i != dark] or [dark] + dc = self._dataset_collection + n_rows = len(rows) + # The free set, independent of this target's own subset; the fit drops + # non-finite observations itself. + fit_mask = torch.cat( + [dc[ctx.keys[i]].free.mask.to(ctx.mask.device) for i in rows] + ) + self._fit = fit_difference_power( + torch.cat([delta_obs[i] for i in rows]), + torch.cat([sigma_diff[i] for i in rows]), + dss.repeat(n_rows), + epsilon=eps.repeat(n_rows), + f_dark=f_dark.repeat(n_rows), + centric=centric.repeat(n_rows) if centric is not None else None, + fit_mask=fit_mask, + delta_calc=torch.cat([delta_calc[i].detach() for i in rows]), + gamma=self._config.gamma, + fit_sigma_scale=False, + ) + # A reflection missing from the dark row has no amplitude for the power law; + # evaluate it at the median instead of letting a NaN reach the variance, where + # the masked-out branch of the loss would still turn it into a NaN gradient. + finite = torch.isfinite(f_dark) + f_eval = torch.where( + finite, f_dark, f_dark[finite].median() if bool(finite.any()) else 1.0 ) - n_rows = len(rows) - est = self._sigma_d.get( - torch.cat([delta_obs[i] for i in rows]), - torch.cat([sigma_diff[i] for i in rows]), - eps.repeat(n_rows), - dss.repeat(n_rows), - f_dark.repeat(n_rows), - fit_mask, - delta_calc=torch.cat([delta_calc[i] for i in rows]), - target_dss=dss, - out_epsilon=eps, - out_f_dark=f_dark, - out_sigma_diff=sigma_diff[rows[0]], + alpha = self._fit.alpha_at(dss) + beta = self._fit.signal_power( + dss, epsilon=eps, f_dark=f_eval, centric=centric ) return CollectionSigmaDLossInputs( - *ctx, alpha=est.alpha.to(dtype), beta_model=est.beta_model.to(dtype) + *ctx, alpha=alpha.to(dtype).detach(), beta_model=beta.to(dtype).detach() ) def _per_refl(self, ctx) -> torch.Tensor: @@ -322,20 +341,21 @@ def _per_refl(self, ctx) -> torch.Tensor: return torch.where(torch.isfinite(nll), nll, torch.full_like(nll, 1e6)) def maintenance(self) -> None: - """Invalidate the sigma_D estimate so it is refitted from the updated models on - the next forward (``LossState`` calls this after each optimizer-step block).""" - self._sigma_d.reset() + """Clear the fit so it is redone from the updated models on the next forward + (``LossState`` calls this after each optimizer-step block).""" + self._fit = None def stats(self) -> Dict[str, StatEntry]: - """Base collection X-ray stats plus the sigma_D fit summary.""" + """Base collection X-ray stats plus the difference-power fit summary.""" out = super().stats() - sh = self._sigma_d.shells - if sh is not None: - out["sigma_d_gamma"] = stat(float(sh.gamma), VERBOSITY_STANDARD) - out["sigma_d_tau"] = stat(float(sh.tau), VERBOSITY_STANDARD) - out["sigma_d_shells_without_power"] = stat( - float(sh.diagnostics["n_s2_clamped"]), VERBOSITY_STANDARD - ) + fit = self._fit + if fit is not None: + lo, hi = fit.stol_range + ends = (2.0 * torch.tensor([lo, hi], dtype=fit.coeffs.dtype)) ** 2 + alpha = fit.alpha_at(ends.to(fit.coeffs.device)) + out["difference_gamma"] = stat(float(fit.gamma), VERBOSITY_STANDARD) + out["difference_alpha_low_res"] = stat(float(alpha[0]), VERBOSITY_STANDARD) + out["difference_alpha_high_res"] = stat(float(alpha[1]), VERBOSITY_STANDARD) return out From 1f9308fd8ef3607f0669ee0cced3943f102f753f Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Tue, 29 Sep 2026 11:49:17 +0200 Subject: [PATCH 208/250] Remove the shell sigma_D estimator; add DifferencePowerEstimator Nothing uses the per-shell sigma_D estimator any more, so the module and its tests go. DifferencePowerConfig (the dark-amplitude exponent) replaces SigmaDConfig, and --sigma-d-gamma is --difference-gamma. The new DifferencePowerEstimator caches one fit_difference_power result until reset, the counterpart of SigmaAEstimator; the difference_sd target owns one and resets it from maintenance(). Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/user_guide/cli.rst | 16 +- docs/user_guide/targets.rst | 9 +- tests/integration/test_cli_ded_weights.py | 2 +- tests/unit/maps/test_ded_weights.py | 20 +- .../test_collection_sigma_d_target.py | 23 +- .../unit/refinement/test_difference_power.py | 21 +- tests/unit/refinement/test_shells.py | 2 +- tests/unit/refinement/test_sigma_d.py | 358 -------- torchref/cli/_common.py | 22 +- torchref/cli/collection_difference_refine.py | 22 +- torchref/cli/difference_map.py | 4 +- torchref/cli/validate_ded.py | 8 +- .../model_error_estimation/__init__.py | 9 +- .../model_error_estimation/_shells.py | 6 +- .../difference_power.py | 71 ++ .../model_error_estimation/sigma_d.py | 795 ------------------ .../refinement/targets/collection/_specs.py | 4 +- .../refinement/targets/collection/base.py | 2 +- .../refinement/targets/collection/xray.py | 36 +- 19 files changed, 197 insertions(+), 1233 deletions(-) delete mode 100644 tests/unit/refinement/test_sigma_d.py delete mode 100644 torchref/refinement/model_error_estimation/sigma_d.py diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index df43aab0..786bfe4f 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -93,7 +93,7 @@ restraints. ``--weight-schedule`` annealing schedule (default ``5,3,2``), ``-n``/``--n-cycles`` macro-cycles, ``--difference-target {difference,difference_sd}`` (the difference row the schedule drives; default ``difference``), ``--ded-weight`` and -``--sigma-d-gamma`` for the difference MTZ (see ``torchref.difference-map``). +``--difference-gamma`` for the difference MTZ (see ``torchref.difference-map``). :API: :mod:`torchref.cli.collection_difference_refine` @@ -138,8 +138,7 @@ Computes real-space correlations and resolution-binned reciprocal-space CC. selection), ``--mask-radius``, ``--n-bins``, ``--ded-weight`` (the headline weight scheme; every scheme is also reported side by side, real-space in each mask and reciprocal-space overall, as the ``by_weight`` block of the JSON and a table in the -summary, and a ``sigma_d`` fallback to inverse variance is recorded under -``weights``). +summary, and a ``q`` fallback to inverse variance is recorded under ``weights``). :API: :mod:`torchref.cli.validate_ded` @@ -181,12 +180,11 @@ records this: the columns sit in named MTZ datasets -- ``observed``, ``differenc column chooser shows ``/torchref/extrapolated_light/FWT``, and ``gemmi mtz`` prints a history line per dataset. ``torchref.difference-refine`` writes the same file. -**Key options:** ``--ded-weight {inverse_variance,sigma_d,none}`` selects the -scheme the model-phased and two-moment difference columns carry (default -``inverse_variance``; ``sigma_d`` needs calibrated sigmas, reports how many shells -it found without difference power, and falls back to inverse variance with a warning -when that is every shell); ``--sigma-d-gamma`` fixes the -dark-amplitude exponent of the sigma_D power law instead of fitting it; +**Key options:** ``--ded-weight {q,inverse_variance,none}`` selects the scheme the +model-phased and two-moment difference columns carry (default ``q``; its fit uses the +intensity differences when the data carry ``I``/``SIGI``, and falls back to inverse +variance with a warning when too few reflections exist to fit); ``--difference-gamma`` +fixes the dark-amplitude exponent of the difference power law instead of fitting it; ``--all-columns`` writes every alternative map coefficient and diagnostic -- the model-phased difference, the two other extrapolations and the intensity block -- at the cost of two further scale fits. diff --git a/docs/user_guide/targets.rst b/docs/user_guide/targets.rst index 44dbe587..8a5626c0 100644 --- a/docs/user_guide/targets.rst +++ b/docs/user_guide/targets.rst @@ -83,10 +83,11 @@ batched over ``(n_datasets, n_hkl)`` on the collection's common HKL grid. - ``difference_sd`` — the ``difference`` Gaussian centred on :math:`\alpha\,\Delta F_{calc}` with variance :math:`\beta_{model} + \sigma_{\Delta}^2`, where :math:`\alpha` and the unexplained - difference power :math:`\beta_{model}` come from a per-shell moment fit of the - observed differences on the free set (``sigma_D``, - :mod:`torchref.refinement.model_error_estimation.sigma_d`); the expected power - carries an :math:`F_{dark}^{\gamma}` dependence with one fitted :math:`\gamma`. A + difference power :math:`\beta_{model}` come from one maximum-likelihood fit of the + observed differences on the free set, with no resolution shells + (:func:`~torchref.refinement.model_error_estimation.difference_power.fit_difference_power`): + :math:`\alpha` is a smooth function of resolution and the unexplained power carries an + :math:`F_{dark}^{\gamma}` dependence with one fitted :math:`\gamma`. A poor light model inflates the variance where it fails instead of pulling the coordinates toward noise. Select it with ``torchref.difference-refine --difference-target difference_sd``. diff --git a/tests/integration/test_cli_ded_weights.py b/tests/integration/test_cli_ded_weights.py index e33e8b8f..d222a357 100644 --- a/tests/integration/test_cli_ded_weights.py +++ b/tests/integration/test_cli_ded_weights.py @@ -39,7 +39,7 @@ def pair(mtz_dir, pdb_dir, tmp_path_factory): """A dark/light pair from 1DAW with a perturbed light state. The light amplitudes carry an added difference proportional to ``F`` with a - resolution-dependent power, so the sigma_D fit has signal to find; the dark set + resolution-dependent power, so the difference-power fit has signal to find; the dark set keeps the deposited values. The light model is the dark one shifted by 0.2 A. """ import torch diff --git a/tests/unit/maps/test_ded_weights.py b/tests/unit/maps/test_ded_weights.py index 9bcc9ea7..802c0a03 100644 --- a/tests/unit/maps/test_ded_weights.py +++ b/tests/unit/maps/test_ded_weights.py @@ -12,7 +12,6 @@ import pytest import torch -from tests.unit.refinement.test_sigma_d import synth_diff from torchref.maps.ded_weights import ( DEFAULT_SCHEME, SCHEMES, @@ -26,6 +25,25 @@ from torchref.symmetry import SpaceGroup +def synth_diff(n=30000, gamma=1.0, sig_frac=1.0, seed=7, device="cpu"): + """Signed differences with power ``0.05 exp(-3 d*^2) (F / )**gamma``. + + ``F`` is Wilson-like (the modulus of a complex normal). The measurement sigma is + ``sig_frac`` times the rms true difference, constant across reflections so the + inverse-variance and q weights differ only through ``S``. + """ + g = torch.Generator().manual_seed(seed) + dss = torch.linspace(0.02, 0.35, n) + f = (torch.randn(n, generator=g) ** 2 + torch.randn(n, generator=g) ** 2).sqrt() + f = 10.0 * f + s_true = 0.05 * torch.exp(-3.0 * dss) * (f / f.mean()) ** gamma + d_true = torch.randn(n, generator=g) * s_true.sqrt() + sig = torch.full((n,), float(sig_frac) * float(s_true.mean().sqrt())) + d_obs = d_true + torch.randn(n, generator=g) * sig + out = {"delta_obs": d_obs, "sigma_diff": sig, "d_star_sq": dss, "f_dark": f} + return {k: v.to(device) for k, v in out.items()} + + def _inputs(n=20000, sig_frac=1.0, device="cpu"): d = synth_diff(n=n, sig_frac=sig_frac, device=device) g = torch.Generator().manual_seed(5) diff --git a/tests/unit/refinement/test_collection_sigma_d_target.py b/tests/unit/refinement/test_collection_sigma_d_target.py index f7998848..3165024d 100644 --- a/tests/unit/refinement/test_collection_sigma_d_target.py +++ b/tests/unit/refinement/test_collection_sigma_d_target.py @@ -10,10 +10,13 @@ import pytest import torch +from torchref.refinement.model_error_estimation import ( + difference_power as difference_power_module, +) from torchref.refinement.model_error_estimation.difference_power import ( + DifferencePowerEstimator, DifferencePowerFit, ) -from torchref.refinement.targets.collection import xray as xray_module from torchref.refinement.targets.collection import ( CollectionDifferenceSigmaDTarget, CollectionSigmaDLossInputs, @@ -25,7 +28,7 @@ @pytest.fixture def target(loaded_reflection_data, sample_structure_pair): """A dark/light collection whose light amplitudes carry a resolution-dependent - difference proportional to ``F``, so the sigma_D coupling is not zero, and a light + difference proportional to ``F``, so the fitted coupling is not zero, and a light model shifted by 0.2 A.""" from torchref import ReflectionData from torchref.cli._common import load_model @@ -73,8 +76,10 @@ def test_forward_is_finite_and_owns_its_fit(target): _dc, _mc, t = target loss = t.forward() assert torch.isfinite(loss) - assert isinstance(t._fit, DifferencePowerFit) - assert t._fit.sigma_scale == 1.0 and len(t._fit.alpha_coeffs) > 0 + assert isinstance(t._estimator, DifferencePowerEstimator) + fit = t._estimator.fit + assert isinstance(fit, DifferencePowerFit) + assert fit.sigma_scale == 1.0 and len(fit.alpha_coeffs) > 0 ctx = t._loss_inputs() assert isinstance(ctx, CollectionSigmaDLossInputs) assert ctx.alpha.shape == ctx.beta_model.shape == (ctx.obs.shape[1],) @@ -95,25 +100,25 @@ def test_gradient_reaches_the_light_model(target): def test_fit_is_cached_until_maintenance(target): _dc, _mc, t = target t.forward() - first = t._fit + first = t._estimator.fit assert first is not None t.forward() - assert t._fit is first + assert t._estimator.fit is first t.maintenance() - assert t._fit is None + assert t._estimator.fit is None def test_fit_uses_free_reflections_of_the_timepoint_row(target, monkeypatch): dc, _mc, t = target seen = {} - real = xray_module.fit_difference_power + real = difference_power_module.fit_difference_power def spy(delta_obs, sigma_diff, d_star_sq, **kw): seen["fit_mask"] = kw["fit_mask"].clone() seen["n"] = delta_obs.numel() return real(delta_obs, sigma_diff, d_star_sq, **kw) - monkeypatch.setattr(xray_module, "fit_difference_power", spy) + monkeypatch.setattr(difference_power_module, "fit_difference_power", spy) t.forward() n_hkl = dc.hkl.shape[0] assert seen["n"] == n_hkl diff --git a/tests/unit/refinement/test_difference_power.py b/tests/unit/refinement/test_difference_power.py index 39751ec6..0b364fb8 100644 --- a/tests/unit/refinement/test_difference_power.py +++ b/tests/unit/refinement/test_difference_power.py @@ -6,7 +6,8 @@ difference's resolution-dependent coupling and the power it leaves unexplained are recovered together; the bounded Wiener weight never falls below its floor, so no reflection or resolution range is removed even when the data hold no signal; the sigma -scale stays within its bounds when the differences hold no noise; the fit runs under +scale stays within its bounds when the differences hold no noise; the estimator caches +one fit until reset and applies its configured exponent; the fit runs under ``torch.no_grad()`` and on every available device. """ @@ -15,6 +16,8 @@ from torchref.refinement.model_error_estimation.difference_power import ( SIGMA_SCALE_BOUNDS, + DifferencePowerConfig, + DifferencePowerEstimator, bounded_wiener_weight, fit_difference_power, ) @@ -132,3 +135,19 @@ def test_recovers_the_model_coupling_and_the_unexplained_power(): assert float((beta / s["s_true"]).log().abs().median()) < LOG_POWER_ATOL no_model = fit_difference_power(s["delta"], s["sigma"], s["dss"], f_dark=s["f"]) assert torch.equal(no_model.alpha_at(s["dss"]), torch.ones_like(s["dss"])) + + +@pytest.mark.unit +def test_estimator_caches_until_reset_and_applies_its_config(): + s = synth(n=5000) + est = DifferencePowerEstimator(DifferencePowerConfig(gamma=0.0)) + assert est.fit is None + first = est.get(s["delta"], s["sigma"], s["dss"], f_dark=s["f"]) + assert first.gamma == 0.0 and est.fit is first + # Cached: different arguments are ignored until reset. + assert est.get(2 * s["delta"], s["sigma"], s["dss"], f_dark=s["f"]) is first + est.reset() + assert est.fit is None + assert est.get(s["delta"], s["sigma"], s["dss"], f_dark=s["f"]) is not first + with pytest.raises(ValueError): + DifferencePowerConfig(gamma=10.0) diff --git a/tests/unit/refinement/test_shells.py b/tests/unit/refinement/test_shells.py index c29a83d5..0df13f9d 100644 --- a/tests/unit/refinement/test_shells.py +++ b/tests/unit/refinement/test_shells.py @@ -1,4 +1,4 @@ -"""Shell helpers shared by sigma_A and sigma_D. +"""Shell helpers used by sigma_A. Pinned: the helpers ``sigma_a`` re-imports are the same objects ``_shells`` defines, so a fit through either module reduces identically; ``equal_count_shells`` reproduces the diff --git a/tests/unit/refinement/test_sigma_d.py b/tests/unit/refinement/test_sigma_d.py deleted file mode 100644 index 39ec8f72..00000000 --- a/tests/unit/refinement/test_sigma_d.py +++ /dev/null @@ -1,358 +0,0 @@ -"""Properties of the sigma_D difference-power estimator. - -Pinned on seeded synthetic differences with a KNOWN power law: the per-shell power is -recovered from a single dataset through ``mean(dF**2) - mean(sigma**2)``, the dark- -amplitude exponent is recovered and can be fixed, the moment identity with a difference -model holds exactly, clamps are counted and an all-noise input is flagged rather than -weighted, degenerate inputs stay finite, per-reflection weights lie in ``[0, 1)`` with the -shell mean of the power preserved, the fit is deterministic and device-independent, and -the cached estimator resets on demand. -""" - -import pytest -import torch - -from torchref.refinement.model_error_estimation._shells import interp_in_dss -from torchref.refinement.model_error_estimation.sigma_d import ( - GAMMA_DEFAULT, - SigmaDConfig, - SigmaDEstimator, - estimate_sigma_d, - sigma_d_per_reflection, -) - -#: Tolerance on the recovered shell power relative to the truth. A shell of 140 -#: reflections estimates ``B`` with a relative sd of ``sqrt(2/140) = 12%``; the line -#: shrinkage pools shells, and the decile means below average ~14 shells, so 10% is -#: ~3 sd of what remains. -POWER_RTOL = 0.10 -#: Tolerance on the fitted exponent. Its standard error at 30 000 reflections is ~0.04; -#: 0.15 is well above that and well below the difference between the pure shell model -#: (0) and the default (1). -GAMMA_ATOL = 0.15 - - -def synth_diff( - n=30000, - gamma=1.0, - sig_frac=1.0, - seed=7, - dtype=torch.float32, - device="cpu", - with_model=False, - alpha_true=0.8, -): - """Signed differences with power ``Sigma_N(d*^2) * (F / )**gamma``. - - ``F`` is Wilson-like (the modulus of a complex normal) so the amplitude classes are - populated realistically; ``Sigma_N`` falls with resolution. The measurement sigma is - ``sig_frac`` times the rms true difference, constant across reflections so the - inverse-variance and sigma_D weights differ only through ``S``. - """ - g = torch.Generator().manual_seed(seed) - dss = torch.linspace(0.02, 0.35, n, dtype=torch.float64) - f = ( - torch.randn(n, generator=g, dtype=torch.float64) ** 2 - + torch.randn(n, generator=g, dtype=torch.float64) ** 2 - ).sqrt() * 10.0 - sigma_n = 0.05 * torch.exp(-3.0 * dss) - s_true = sigma_n * (f / f.mean()) ** gamma - d_true = torch.randn(n, generator=g, dtype=torch.float64) * s_true.sqrt() - sig = torch.full( - (n,), float(sig_frac) * float(s_true.mean().sqrt()), dtype=torch.float64 - ) - d_obs = d_true + torch.randn(n, generator=g, dtype=torch.float64) * sig - out = { - "delta_obs": d_obs, - "sigma_diff": sig, - "d_star_sq": dss, - "f_dark": f, - "s_true": s_true, - "fit_mask": torch.ones(n, dtype=torch.bool), - } - if with_model: - beta_true = 0.3 * s_true - out["delta_calc"] = ( - d_true + torch.randn(n, generator=g, dtype=torch.float64) * beta_true.sqrt() - ) / alpha_true - return { - k: ( - v.to(device=device, dtype=dtype) - if v.dtype.is_floating_point - else v.to(device) - ) - for k, v in out.items() - } - - -def _decile_means(values, dss, n_dec=10): - order = torch.argsort(dss) - chunks = torch.chunk(values[order], n_dec) - return torch.stack([c.mean() for c in chunks]) - - -@pytest.mark.unit -def test_recovers_shell_power_from_one_dataset(any_device): - d = synth_diff(device=any_device) - sh = estimate_sigma_d( - d["delta_obs"], - d["sigma_diff"], - None, - d["d_star_sq"], - d["f_dark"], - d["fit_mask"], - ) - est = sigma_d_per_reflection(sh, d["d_star_sq"], None, d["f_dark"], d["sigma_diff"]) - assert not sh.degenerate and not sh.all_zero - assert (sh.Sigma_N > 0).all() - got = _decile_means(est.S, d["d_star_sq"]) - want = _decile_means(d["s_true"], d["d_star_sq"]) - assert torch.allclose(got, want, rtol=POWER_RTOL) - - -@pytest.mark.unit -@pytest.mark.parametrize("gamma", [1.0, 0.5]) -def test_recovers_the_amplitude_exponent(gamma): - d = synth_diff(gamma=gamma, dtype=torch.float64) - sh = estimate_sigma_d( - d["delta_obs"], - d["sigma_diff"], - None, - d["d_star_sq"], - d["f_dark"], - d["fit_mask"], - ) - assert sh.gamma_fitted and sh.diagnostics["gamma_reason"] == "fitted" - assert abs(sh.gamma - gamma) < GAMMA_ATOL - assert sh.diagnostics["gamma_se"] < GAMMA_ATOL - - -@pytest.mark.unit -def test_fixed_exponent_is_honoured(): - d = synth_diff(n=5000) - sh = estimate_sigma_d( - d["delta_obs"], - d["sigma_diff"], - None, - d["d_star_sq"], - d["f_dark"], - d["fit_mask"], - gamma=0.7, - ) - assert sh.gamma == 0.7 and not sh.gamma_fitted - assert sh.diagnostics["gamma_reason"] == "fixed" - with pytest.raises(ValueError): - SigmaDConfig(gamma=3.0) - - -@pytest.mark.unit -def test_without_dark_amplitude_the_power_is_flat_within_a_shell(): - d = synth_diff(n=5000) - sh = estimate_sigma_d( - d["delta_obs"], d["sigma_diff"], None, d["d_star_sq"], None, d["fit_mask"] - ) - assert sh.gamma == GAMMA_DEFAULT and not sh.gamma_fitted - assert sh.diagnostics["gamma_reason"] == "no_f_dark" - est = sigma_d_per_reflection(sh, d["d_star_sq"], None, None, d["sigma_diff"]) - # Reflections at the same resolution share the power regardless of amplitude. - order = torch.argsort(d["d_star_sq"]) - close = est.S[order][:200] - assert float(close.max() / close.min()) < 1.05 - - -@pytest.mark.unit -@pytest.mark.parametrize("dtype,rtol", [(torch.float32, 1e-4), (torch.float64, 1e-10)]) -def test_moment_identity_with_a_difference_model(dtype, rtol): - d = synth_diff(dtype=dtype, with_model=True) - sh = estimate_sigma_d( - d["delta_obs"], - d["sigma_diff"], - None, - d["d_star_sq"], - d["f_dark"], - d["fit_mask"], - delta_calc=d["delta_calc"], - shrink=False, - ) - assert sh.has_model - assert sh.diagnostics["n_s2_clamped"] == 0 - # Sampling noise can push alpha**2 Sigma_P above Sigma_N in a few shells; the clamp - # there is counted, and the identity is exact everywhere it did not fire. - unclamped = sh.alpha**2 * sh.Sigma_P <= sh.Sigma_N - assert ( - int(unclamped.sum()) - == sh.diagnostics["n_shell"] - sh.diagnostics["n_beta_clamped"] - ) - assert float(unclamped.float().mean()) > 0.8 - lhs = sh.alpha**2 * sh.Sigma_P + sh.beta_model + sh.S2 - assert torch.allclose(lhs[unclamped], sh.B[unclamped], rtol=rtol) - # alpha is the Gaussian coupling S / (S + beta_true) / alpha_true-scaled slope; it must - # be positive and below one for this generator. - assert (sh.alpha > 0).all() and (sh.alpha < 1).all() - - -@pytest.mark.unit -def test_clamps_are_counted_and_all_noise_is_flagged(): - d = synth_diff(n=5000, sig_frac=5.0) - sh = estimate_sigma_d( - d["delta_obs"], - d["sigma_diff"], - None, - d["d_star_sq"], - d["f_dark"], - d["fit_mask"], - ) - assert sh.diagnostics["n_s2_clamped"] > 0 - noise = synth_diff(n=5000, sig_frac=50.0) - sh2 = estimate_sigma_d( - noise["delta_obs"], - noise["sigma_diff"] * 1.2, - None, - noise["d_star_sq"], - noise["f_dark"], - noise["fit_mask"], - shrink=False, - ) - assert sh2.all_zero - est = sigma_d_per_reflection( - sh2, noise["d_star_sq"], None, noise["f_dark"], noise["sigma_diff"] - ) - assert torch.equal(est.w, torch.zeros_like(est.w)) - - -@pytest.mark.unit -def test_pure_noise_with_calibrated_sigma_gets_no_power(): - """Nothing tells the estimator whether a difference exists: on pure noise with - calibrated sigmas the shrinkage must not manufacture power from the positive half - of the noise in ``B - S2``, and the weights collapse onto inverse variance.""" - g = torch.Generator().manual_seed(3) - n = 30000 - dss = torch.linspace(0.02, 0.35, n, dtype=torch.float64) - f = torch.rand(n, generator=g, dtype=torch.float64) * 20.0 + 1.0 - sig = 0.2 + 0.8 * dss - d_obs = torch.randn(n, generator=g, dtype=torch.float64) * sig - mask = torch.ones(n, dtype=torch.bool) - sh = estimate_sigma_d(d_obs, sig, None, dss, f, mask) - assert not sh.degenerate - # Shell power is below a few per cent of the noise power in every shell. - assert (sh.Sigma_N <= 0.05 * sh.S2).all() - est = sigma_d_per_reflection(sh, dss, None, f, sig) - ivw = 1.0 / sig**2 - ivw = ivw / ivw.mean() - w_sd = est.w / est.w.mean().clamp(min=1e-30) - if not sh.all_zero: - assert torch.corrcoef(torch.stack([w_sd, ivw]))[0, 1] > 0.97 - - -@pytest.mark.unit -def test_degenerate_input_stays_finite(): - d = synth_diff(n=100) - mask = torch.zeros(100, dtype=torch.bool) - mask[0] = True - sh = estimate_sigma_d( - d["delta_obs"], d["sigma_diff"], None, d["d_star_sq"], d["f_dark"], mask - ) - assert sh.degenerate and not sh.all_zero - est = sigma_d_per_reflection(sh, d["d_star_sq"], None, d["f_dark"], d["sigma_diff"]) - assert torch.isfinite(est.S).all() and torch.isfinite(est.w).all() - assert est.S.shape == (100,) - - -@pytest.mark.unit -def test_per_reflection_weights_and_shell_mean(any_device): - d = synth_diff(device=any_device) - eps = torch.where( - torch.arange(d["delta_obs"].numel(), device=any_device) % 7 == 0, 2.0, 1.0 - ).to(d["delta_obs"].dtype) - sh = estimate_sigma_d( - d["delta_obs"], d["sigma_diff"], eps, d["d_star_sq"], d["f_dark"], d["fit_mask"] - ) - est = sigma_d_per_reflection(sh, d["d_star_sq"], eps, d["f_dark"], d["sigma_diff"]) - assert (est.w >= 0).all() and (est.w < 1).all() - assert est.S.device == d["delta_obs"].device - # The multiplier has shell mean one, so S / epsilon averages to Sigma_N over a shell. - counts = sh.counts.to(torch.long) # dtype-ok: split sizes; PyTorch requires int64 - order = torch.argsort(d["d_star_sq"]) - per_shell = torch.stack( - [c.mean() for c in torch.split((est.S / eps)[order], counts.tolist())] - ) - assert torch.allclose(per_shell, sh.Sigma_N, rtol=0.15) - # A missing dark amplitude means a multiplier of one. - f_missing = d["f_dark"].clone() - f_missing[:50] = float("nan") - est2 = sigma_d_per_reflection(sh, d["d_star_sq"], eps, f_missing, d["sigma_diff"]) - log_sn = interp_in_dss(d["d_star_sq"][:50], sh.bin_dss, torch.log(sh.Sigma_N)) - assert torch.allclose(est2.S[:50], eps[:50] * torch.exp(log_sn), rtol=1e-4) - - -@pytest.mark.unit -def test_deterministic_and_device_independent(any_device): - d_cpu = synth_diff() - a = estimate_sigma_d( - d_cpu["delta_obs"], - d_cpu["sigma_diff"], - None, - d_cpu["d_star_sq"], - d_cpu["f_dark"], - d_cpu["fit_mask"], - ) - b = estimate_sigma_d( - d_cpu["delta_obs"], - d_cpu["sigma_diff"], - None, - d_cpu["d_star_sq"], - d_cpu["f_dark"], - d_cpu["fit_mask"], - ) - assert torch.equal(a.Sigma_N, b.Sigma_N) and a.gamma == b.gamma - d_dev = synth_diff(device=any_device) - c = estimate_sigma_d( - d_dev["delta_obs"], - d_dev["sigma_diff"], - None, - d_dev["d_star_sq"], - d_dev["f_dark"], - d_dev["fit_mask"], - ) - assert torch.allclose(c.Sigma_N.cpu(), a.Sigma_N, rtol=1e-4) - assert abs(c.gamma - a.gamma) < 1e-3 - - -@pytest.mark.unit -def test_estimator_caches_until_reset_and_remaps(): - d = synth_diff(n=5000) - est = SigmaDEstimator(SigmaDConfig(gamma=1.0)) - first = est.get( - d["delta_obs"], - d["sigma_diff"], - None, - d["d_star_sq"], - d["f_dark"], - d["fit_mask"], - ) - assert ( - est.get( - d["delta_obs"], - d["sigma_diff"], - None, - d["d_star_sq"], - d["f_dark"], - d["fit_mask"], - ) - is first - ) - est.reset() - assert est._cache is None - target = d["d_star_sq"][:1000] - remapped = est.get( - d["delta_obs"], - d["sigma_diff"], - None, - d["d_star_sq"], - d["f_dark"], - d["fit_mask"], - target_dss=target, - out_f_dark=d["f_dark"][:1000], - out_sigma_diff=d["sigma_diff"][:1000], - ) - assert remapped.S.shape == (1000,) and est.shells is not None diff --git a/torchref/cli/_common.py b/torchref/cli/_common.py index 8285df4f..ecb5b610 100644 --- a/torchref/cli/_common.py +++ b/torchref/cli/_common.py @@ -390,7 +390,7 @@ def add_all_columns_arg(parser: argparse.ArgumentParser) -> None: def add_ded_weight_args(parser: argparse.ArgumentParser) -> None: - """Add ``--ded-weight`` and ``--sigma-d-gamma`` for the difference-map writers. + """Add ``--ded-weight`` and ``--difference-gamma`` for the difference-map writers. Every registered scheme's weight is written to the difference MTZ regardless; the choice here decides which one the headline products (validate-ded correlations, @@ -409,13 +409,13 @@ def add_ded_weight_args(parser: argparse.ArgumentParser) -> None: f"{DEFAULT_SCHEME}). All weights are written as columns.", ) parser.add_argument( - "--sigma-d-gamma", + "--difference-gamma", type=float, default=None, metavar="GAMMA", - help="Fix the dark-amplitude exponent of the difference power law in [0, 2] " - "instead of fitting it; used by the q-weight and the sigma_D difference " - "target (default: fitted).", + help="Fix the dark-amplitude exponent of the difference power law in [-1, 3] " + "instead of fitting it; used by the q-weight on amplitude data and by the " + "difference_sd target (default: fitted, or 0 on intensity data).", ) @@ -448,12 +448,14 @@ def intensity_difference(data_dark, data_light, mask=None): return I_light - I_dark, (sig_dark**2 + sig_light**2).sqrt() -def sigma_d_config_from_args(args: argparse.Namespace): - """The :class:`~torchref.refinement.model_error_estimation.sigma_d.SigmaDConfig` - selected by ``--sigma-d-gamma``.""" - from torchref.refinement.model_error_estimation.sigma_d import SigmaDConfig +def difference_config_from_args(args: argparse.Namespace): + """The :class:`~torchref.refinement.model_error_estimation.difference_power. + DifferencePowerConfig` selected by ``--difference-gamma``.""" + from torchref.refinement.model_error_estimation.difference_power import ( + DifferencePowerConfig, + ) - return SigmaDConfig(gamma=getattr(args, "sigma_d_gamma", None)) + return DifferencePowerConfig(gamma=getattr(args, "difference_gamma", None)) def add_output_format_args(parser: argparse.ArgumentParser) -> None: diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index 4f942757..2204c41f 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -48,7 +48,7 @@ parse_weights, register_timing, intensity_difference, - sigma_d_config_from_args, + difference_config_from_args, validate_cif_files, validate_files, ) @@ -239,7 +239,7 @@ def setup_loss_state( similarity_alpha=2.0, two_moment=False, difference_target="difference", - sigma_d_config=None, + difference_config=None, ): """Build LossState with collection-aware targets. @@ -256,8 +256,8 @@ def setup_loss_state( Which difference row the weight schedule drives. Both are registered, as ``xray/difference`` and ``xray/difference_sd``; the other keeps the weight in ``target_weights`` (zero by default). - sigma_d_config : SigmaDConfig, optional - Exponent and shrinkage settings of the ``difference_sd`` row's estimator. + difference_config : DifferencePowerConfig, optional + Its ``gamma`` fixes the dark-amplitude exponent of the ``difference_sd`` fit. """ from torchref.refinement import LossState from torchref.refinement.targets import TotalADPTarget, TotalGeometryTarget @@ -280,7 +280,7 @@ def setup_loss_state( dataset_collection, model_collection, scaler=scaler, - sigma_d_config=sigma_d_config, + difference_config=difference_config, ) selected_diff = {"difference": diff_target, "difference_sd": diff_sd_target}[ difference_target @@ -889,7 +889,7 @@ def write_results_mtz( all_columns=False, verbose=1, ded_weight=DEFAULT_SCHEME, - sigma_d_config=None, + difference_config=None, ): """Write the difference map, and map coefficients when a light model is given. @@ -929,7 +929,7 @@ def write_results_mtz( ded_weight : str, optional Weight scheme for the model-phased and two-moment difference columns; one of :data:`torchref.maps.ded_weights.SCHEMES`. - sigma_d_config : SigmaDConfig, optional + difference_config : DifferencePowerConfig, optional Its ``gamma`` fixes the ``F_dark`` exponent of the ``q`` scheme. Returns @@ -992,7 +992,7 @@ def write_results_mtz( f_dark=Fobs_dark_vals, delta_intensity=delta_I, sigma_delta_intensity=sig_delta_I, - gamma=sigma_d_config.gamma if sigma_d_config is not None else None, + gamma=difference_config.gamma if difference_config is not None else None, ) all_w = all_ded_weights(**snr_inputs) selected = all_w[ded_weight] @@ -1184,7 +1184,7 @@ def main(): default="difference", help="Difference row the weight schedule drives: 'difference' is the Gaussian " "under the measurement variance, 'difference_sd' centres on " - "alpha*dF_calc with the sigma_D unexplained power added to the variance " + "alpha*dF_calc with the model's unexplained power added to the variance " "(default: difference).", ) refine.add_argument( @@ -1441,7 +1441,7 @@ def main(): similarity_alpha=args.similarity_alpha, two_moment=args.two_moment, difference_target=args.difference_target, - sigma_d_config=sigma_d_config_from_args(args), + difference_config=difference_config_from_args(args), ) if args.verbose > 0: @@ -1744,7 +1744,7 @@ def _mtz_to_cif(mtz_path, cif_path): all_columns=args.all_columns, verbose=args.verbose, ded_weight=args.ded_weight, - sigma_d_config=sigma_d_config_from_args(args), + difference_config=difference_config_from_args(args), ) # --- JSON summary --- diff --git a/torchref/cli/difference_map.py b/torchref/cli/difference_map.py index e2c2d2a6..612f2b7d 100644 --- a/torchref/cli/difference_map.py +++ b/torchref/cli/difference_map.py @@ -48,7 +48,7 @@ configure_unbuffered_output, register_timing, parse_device_str, - sigma_d_config_from_args, + difference_config_from_args, validate_cif_files, validate_files, ) @@ -229,7 +229,7 @@ def main(): all_columns=args.all_columns, verbose=args.verbose, ded_weight=args.ded_weight, - sigma_d_config=sigma_d_config_from_args(args), + difference_config=difference_config_from_args(args), ) if args.verbose > 0: diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index 2cf3c73a..b28d6360 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -43,7 +43,7 @@ load_reflection_data, parse_device_str, register_timing, - sigma_d_config_from_args, + difference_config_from_args, validate_cif_files, validate_files, ) @@ -206,7 +206,7 @@ def setup_ded_context( n_bins=20, verbose=0, ded_weight=DEFAULT_SCHEME, - sigma_d_config=None, + difference_config=None, ): """Load reflection data and prepare shared state for DED validation. @@ -304,7 +304,7 @@ def setup_ded_context( cell=data_dark.cell, spacegroup=data_dark.spacegroup, f_dark=F_dark, - gamma=sigma_d_config.gamma if sigma_d_config is not None else None, + gamma=difference_config.gamma if difference_config is not None else None, ) selected = all_w[ded_weight] weights = selected.weights @@ -652,7 +652,7 @@ def run_validation(args): n_bins=args.n_bins, verbose=args.verbose, ded_weight=args.ded_weight, - sigma_d_config=sigma_d_config_from_args(args), + difference_config=difference_config_from_args(args), ) fallback_messages = [ str(w.message) diff --git a/torchref/refinement/model_error_estimation/__init__.py b/torchref/refinement/model_error_estimation/__init__.py index fc01891e..36a0d876 100644 --- a/torchref/refinement/model_error_estimation/__init__.py +++ b/torchref/refinement/model_error_estimation/__init__.py @@ -1,4 +1,4 @@ -"""The two model-error estimators, one module each. +"""The model-error estimators, one module each. A refined model disagrees with the data for two reasons that need separating: the measurement is noisy (``sigma_obs``, which the data carries) and the *model is wrong*. @@ -12,8 +12,13 @@ *structure alone* via the diagonal Fisher information, never seeing ``F_obs`` or ``F_calc``. It agrees with ``beta`` on shape but not magnitude, which is why it takes a caller-supplied scale. +* :mod:`.difference_power` -- **difference-driven**. The expected power of a + light-minus-dark difference and, given a model difference, its coupling and the power + the model leaves unexplained, by one per-reflection likelihood fit with no resolution + shells. What the difference-map weights, the extrapolated amplitudes and the + ``difference_sd`` collection target consume. -**This ``__init__`` deliberately imports nothing.** Both modules are heavy and +**This ``__init__`` deliberately imports nothing.** The modules are heavy and ``sigma_a`` is imported from inside :mod:`torchref.scaling` methods to avoid closing a ``scaling`` <-> ``refinement`` cycle; pulling them in here would defeat that, and re-exporting from :mod:`torchref.refinement` would make the import an attribute lookup on diff --git a/torchref/refinement/model_error_estimation/_shells.py b/torchref/refinement/model_error_estimation/_shells.py index 3d6e9a7d..ecb1b56a 100644 --- a/torchref/refinement/model_error_estimation/_shells.py +++ b/torchref/refinement/model_error_estimation/_shells.py @@ -2,9 +2,9 @@ Equal-count shells over ``d*^2``, atomic-free segment sums, linear interpolation of per-shell values back to reflections, and DerSimonian-Laird shrinkage of noisy per-shell -estimates toward a weighted straight line. :mod:`.sigma_a` and :mod:`.sigma_d` both -build on these; ``estimate_beta`` keeps its own module-level aliases so that its body -resolves the same globals it always did. +estimates toward a weighted straight line. :mod:`.sigma_a` builds on these; +``estimate_beta`` keeps its own module-level aliases so that its body resolves the same +globals it always did. Plain tensors in and out. Every result lives on the device of its inputs, and float work happens in the dtype of the inputs, so callers control both by what they pass. diff --git a/torchref/refinement/model_error_estimation/difference_power.py b/torchref/refinement/model_error_estimation/difference_power.py index bf50f46a..1709dc1e 100644 --- a/torchref/refinement/model_error_estimation/difference_power.py +++ b/torchref/refinement/model_error_estimation/difference_power.py @@ -27,6 +27,8 @@ power the model does not explain -- the coupling and unexplained power a difference likelihood needs, fitted in one pass instead of from per-shell cross moments. +:class:`DifferencePowerEstimator` caches one fit for a target that re-evaluates it every +forward, the counterpart of :class:`~.sigma_a.SigmaAEstimator`. :func:`bounded_wiener_weight` turns a fit into a weight that down-weights noisy reflections but never removes one, the resolution-continuous counterpart of the q-weight's floor. @@ -67,6 +69,25 @@ _LOG_CENTRIC_BOUNDS = (-7.0, 7.0) +@dataclass(frozen=True) +class DifferencePowerConfig: + """The user-facing knob of the difference-power fit, as one value. + + ``gamma=None`` fits the dark-amplitude exponent; a float fixes it, within the + fit's bounds. Frozen, so two consumers sharing a config cannot drift apart. + """ + + gamma: float | None = None + + def __post_init__(self): + if self.gamma is not None: + g = float(self.gamma) + lo, hi = _GAMMA_BOUNDS + if not (lo <= g <= hi): + raise ValueError(f"gamma must lie in {_GAMMA_BOUNDS}, got {g}") + object.__setattr__(self, "gamma", g) + + @dataclass(frozen=True) class DifferencePowerFit: """A fitted difference-power model; evaluate it with :meth:`signal_power`. @@ -414,6 +435,54 @@ def nll(t): ) +class DifferencePowerEstimator: + """Lazy, cached difference-power fit. + + Thin stateful wrapper around :func:`fit_difference_power`: fits on the first + :meth:`get` and returns the cached, detached fit until :meth:`reset`. **The owning + target must call :meth:`reset` whenever the models or data change** (``LossState`` + reaches it through ``maintenance()``); otherwise every later :meth:`get` returns + the stale fit and ignores its arguments. + + Parameters + ---------- + config : DifferencePowerConfig, optional + Its ``gamma`` is passed to every fit that does not name one; module defaults + when omitted. + """ + + def __init__(self, config: DifferencePowerConfig | None = None): + self.config = config if config is not None else DifferencePowerConfig() + self._fit: DifferencePowerFit | None = None + + def reset(self) -> None: + """Invalidate the cache so the next :meth:`get` fits again.""" + self._fit = None + + @property + def fit(self) -> DifferencePowerFit | None: + """The last fit, or ``None`` before the first :meth:`get` and after + :meth:`reset`.""" + return self._fit + + def get( + self, + delta_obs: torch.Tensor, + sigma_diff: torch.Tensor, + d_star_sq: torch.Tensor, + **kwargs, + ) -> DifferencePowerFit: + """Return the cached fit, or fit and cache it. + + Arguments are those of :func:`fit_difference_power` and are used only when no + fit is cached. + """ + if self._fit is None: + kwargs.setdefault("gamma", self.config.gamma) + self._fit = fit_difference_power(delta_obs, sigma_diff, d_star_sq, **kwargs) + return self._fit + + def bounded_wiener_weight( snr: torch.Tensor, snr_floor: float = DEFAULT_SNR_FLOOR ) -> torch.Tensor: @@ -445,6 +514,8 @@ def bounded_wiener_weight( "DEFAULT_ORDER", "DEFAULT_SNR_FLOOR", "SIGMA_SCALE_BOUNDS", + "DifferencePowerConfig", + "DifferencePowerEstimator", "DifferencePowerFit", "bounded_wiener_weight", "fit_difference_power", diff --git a/torchref/refinement/model_error_estimation/sigma_d.py b/torchref/refinement/model_error_estimation/sigma_d.py deleted file mode 100644 index 8ead3617..00000000 --- a/torchref/refinement/model_error_estimation/sigma_d.py +++ /dev/null @@ -1,795 +0,0 @@ -"""Difference-driven error estimation: the expected difference power ``sigma_D``. - -A light-minus-dark difference coefficient ``dF_obs = dF_true + noise`` carries a true -signal whose power ``S = E[dF_true**2]`` varies with resolution and with the dark -amplitude, and a measurement noise ``sigma_diff**2`` the merge reports. The best linear -estimate of ``dF_true`` from ``dF_obs`` is ``w * dF_obs`` with the Wiener weight -``w = S / (S + sigma_diff**2)``, so a difference map needs ``S`` per reflection. Inverse -variance alone weights by precision and treats every reflection as carrying the same -expected difference, which suppresses the strong reflections whose difference power is -ten to seventy times that of weak ones. - -``S`` needs no half datasets. Per resolution shell the second moment of the observed -differences is ``B = S + S2`` with ``S2`` the mean measurement variance, so -``Sigma_N = B - S2`` is the expected true difference power, the same identity -:mod:`.sigma_a` uses for amplitudes. Within a shell the power follows the dark amplitude -as ``(F_dark / )**gamma`` with one fitted exponent, carried by the per-reflection -multiplier ``epsilon`` exactly as the reflection multiplicity is. With a difference model -``dF_calc`` the shell moments also give the Gaussian coupling ``alpha = / -`` and the unexplained power ``beta_model = Sigma_N - alpha**2 Sigma_P``, the -extra variance a difference likelihood adds to ``sigma_diff**2``. - -Differences are signed and small, so the statistics are Gaussian throughout: there is no -Rice branch and no centric distinction. Plain tensors in and out, no ``ReflectionData`` -or ``Scaler`` coupling, so :mod:`torchref.maps` and :mod:`torchref.cli` can import this -module without closing an import cycle. Every result lives on the device of its inputs. -""" - -import math -from dataclasses import dataclass - -import torch - -from torchref.config import get_float_dtype - -from ._shells import equal_count_shells, interp_in_dss, segsum -from .sigma_a import SHRINK_ENABLED - -# --- sigma_D estimator constants ------------------------------------------------- -#: Shell construction, matching the ``estimate_beta`` defaults so a sigma_A and a sigma_D -#: fit on the same reflections use the same shells. -PER_BIN = 140 -MIN_BINS = 5 -MIN_PER_BIN = 40 -#: Exponent of the dark-amplitude power law when it cannot be fitted. The fitted value on -#: the small-molecule, TD1 and bacteriorhodopsin A/B campaigns was 0.8-1.2, so 1.0 (power -#: proportional to the amplitude) is the informed default; 0.0 would be the pure shell model. -GAMMA_DEFAULT = 1.0 -#: Bounds on the fitted exponent. Outside [0, 2] the class means are dominated by one -#: amplitude decile and the regression is on noise; 2 is proportionality to the intensity. -GAMMA_BOUNDS = (0.0, 2.0) -#: Dark-amplitude quantile classes per shell for the exponent regression. Four keeps at -#: least 35 reflections per class at the default shell size; more classes did not move -#: the fitted exponent on the campaign data. -N_F_CLASSES = 4 -#: Minimum number of usable (shell, class) cells before the exponent is fitted at all. -GAMMA_MIN_CLASSES = 8 -#: Minimum reflections in a class for its moment to enter the exponent regression. -MIN_PER_CLASS = 5 -#: Largest standard error at which a fitted exponent is used. Above it the class -#: moments are noise (a null dataset gives se ~ 1 or more), and the default is safer than -#: a random exponent that would redistribute weight between amplitude classes. -GAMMA_SE_MAX = 0.5 -#: Gauss-Newton iterations for the decaying-curve fit; the problem is two-parameter and -#: well conditioned, so this is far more than it needs. -CURVE_ITERS = 60 -#: Floor on ``F_dark`` relative to its shell mean before the power law is evaluated, so a -#: zero or near-zero dark amplitude cannot delete the expected difference power. -F_FLOOR_FRAC = 0.05 -#: Positive floor used where a logarithm of a clamped-to-zero power is needed. -_TINY = 1e-30 - - -@dataclass(frozen=True) -class SigmaDConfig: - """The estimator's knobs, as one value. - - ``gamma=None`` fits the dark-amplitude exponent; a float fixes it. ``shrink=None`` - means the module default shared with sigma_A, normalised here so consumers never - handle ``None``. Frozen, so two consumers sharing a config cannot drift apart. - """ - - gamma: float | None = None - shrink: bool | None = None - - def __post_init__(self): - if self.gamma is not None: - g = float(self.gamma) - lo, hi = GAMMA_BOUNDS - if not (lo <= g <= hi): - raise ValueError(f"gamma must lie in {GAMMA_BOUNDS}, got {g}") - object.__setattr__(self, "gamma", g) - object.__setattr__( - self, "shrink", bool(SHRINK_ENABLED if self.shrink is None else self.shrink) - ) - - -@dataclass(frozen=True) -class SigmaDShells: - """Per-shell output of :func:`estimate_sigma_d`. - - All power quantities are ``epsilon``-reduced and in ``F**2`` units of the input. - - Attributes - ---------- - B, S2 - Raw second moment of the observed differences and mean measurement variance. - Sigma_N_raw, Sigma_N - Expected true difference power ``(B - S2)`` clamped at zero, before and after the - shrinkage of the signed value toward a decaying curve ``exp(a + b d*^2)`` fitted - to every shell. ``Sigma_N`` is what the weights use. - Sigma_P, C, alpha, beta_model - Model power, cross moment, Gaussian coupling ``C / Sigma_P`` and unexplained power - ``(Sigma_N - alpha**2 Sigma_P)`` clamped at zero. Without a model ``Sigma_P`` and - ``C`` are zero, ``alpha`` one and ``beta_model == Sigma_N``. - counts, bin_dss - Reflections per shell and its mean ``d*^2`` in A^-2, the interpolation abscissa. - bin_log_fbar, bin_log_z - Log of the shell-mean dark amplitude and of the shell mean of - ``(F / Fbar)**gamma``, so the per-reflection multiplier - ``(F / Fbar)**gamma / Z`` has shell mean one and ``Sigma_N`` stays the shell mean - of the per-reflection power. Zero when no dark amplitude was supplied. - shrink_w, tau, curve_a, curve_b - Shrinkage weight per shell, the between-shell sd about the fitted curve and its - coefficients ``exp(a + b d*^2)`` (NaN when no curve was fitted; the curve is then - zero everywhere). - gamma, gamma_fitted - The exponent used and whether it was fitted rather than fixed or defaulted. - has_model, degenerate, all_zero - Whether ``dF_calc`` was supplied, whether fewer than two usable reflections - existed, and whether every shell's ``Sigma_N`` is zero (weights would vanish). - diagnostics - Counters: ``n_dropped, n_fit, n_shell, n_s2_clamped, n_beta_clamped, - n_f_floored, n_class_dropped, n_class_used, gamma_se, gamma_at_bound, - gamma_reason``. - """ - - B: torch.Tensor - S2: torch.Tensor - Sigma_N_raw: torch.Tensor - Sigma_N: torch.Tensor - Sigma_P: torch.Tensor - C: torch.Tensor - alpha: torch.Tensor - beta_model: torch.Tensor - counts: torch.Tensor - bin_dss: torch.Tensor - bin_log_fbar: torch.Tensor - bin_log_z: torch.Tensor - shrink_w: torch.Tensor - tau: float - curve_a: float - curve_b: float - gamma: float - gamma_fitted: bool - has_model: bool - degenerate: bool - all_zero: bool - diagnostics: dict - - -@dataclass(frozen=True) -class SigmaDEstimate: - """Everything a consumer needs from one estimate, per reflection and detached. - - Attributes - ---------- - S - Expected true difference power ``epsilon * Sigma_N(d*^2) * g(F_dark)``. - sigma_sq - The measurement variance the weight was formed with (``sigma_diff**2``). - w - Wiener weight ``S / (S + sigma_sq)`` in ``[0, 1)``, not normalised. - alpha, beta_model - Coupling and unexplained power, interpolated per shell; ``beta_model`` carries - the same ``epsilon * g`` multiplier as ``S``. - epsilon - The multiplicity actually applied. - shells - The :class:`SigmaDShells` this was interpolated from. - """ - - S: torch.Tensor - sigma_sq: torch.Tensor - w: torch.Tensor - alpha: torch.Tensor - beta_model: torch.Tensor - epsilon: torch.Tensor - shells: SigmaDShells - - -def _working_dtype(t: torch.Tensor) -> torch.dtype: - dtype = torch.promote_types(get_float_dtype(), t.dtype) - # dtype-ok: MPS capability guard, not an allocation - if dtype == torch.float64 and t.device.type == "mps": - raise RuntimeError( - "MPS has no float64; set the defaults float dtype to float32 or use CPU" - ) - return dtype - - -def _degenerate( - delta_obs: torch.Tensor, gamma: float, has_model: bool, diagnostics: dict, out_dtype -) -> SigmaDShells: - """One conservative shell: the mean squared difference as the power, alpha one.""" - ok = torch.isfinite(delta_obs) - b = (delta_obs[ok] ** 2).mean() if bool(ok.any()) else delta_obs.new_ones(()) - one = torch.ones(1, device=delta_obs.device, dtype=out_dtype) - zero = torch.zeros(1, device=delta_obs.device, dtype=out_dtype) - b1 = (one * b).to(out_dtype) - return SigmaDShells( - B=b1, - S2=zero, - Sigma_N_raw=b1, - Sigma_N=b1, - Sigma_P=zero, - C=zero, - alpha=one, - beta_model=b1, - counts=zero, - bin_dss=zero, - bin_log_fbar=zero, - bin_log_z=zero, - shrink_w=zero, - tau=0.0, - curve_a=float("nan"), - curve_b=float("nan"), - gamma=gamma, - gamma_fitted=False, - has_model=has_model, - degenerate=True, - all_zero=False, - diagnostics=diagnostics, - ) - - -def _fit_gamma( - d2e: torch.Tensor, - s2e: torch.Tensor, - log_ratio: torch.Tensor, - seg: torch.Tensor, - n_bins: int, -) -> tuple[float, float, bool, int, int, str]: - """Fit the dark-amplitude exponent from within-shell amplitude classes. - - Each shell is split into ``N_F_CLASSES`` quantile classes of the dark amplitude. A - class contributes ``log( - )`` against its mean log amplitude - ratio when that difference power is positive. One slope is fitted across all shells - with the shell means removed (fixed effects), weighted by ``n_c / 2``: the log of a - mean of ``n`` squared Gaussians has variance ``2 / n``. - - Returns ``(gamma, gamma_se, at_bound, n_used, n_dropped, reason)``; ``reason`` is - ``"fitted"`` or names why the default was taken. - """ - xs, ys, ws, shell_id = [], [], [], [] - n_dropped = 0 - for k in range(n_bins): - in_shell = torch.nonzero(seg == k, as_tuple=True)[0] - n_k = int(in_shell.numel()) - if n_k < N_F_CLASSES * MIN_PER_CLASS: - n_dropped += N_F_CLASSES - continue - order = torch.argsort(log_ratio[in_shell], stable=True) - idx = in_shell[order] - cls = ( - torch.arange(n_k, device=seg.device) * N_F_CLASSES - ) // n_k # dtype-ok: bincount input; PyTorch requires int64 - lengths = torch.bincount(cls, minlength=N_F_CLASSES).to(d2e.dtype) - m = (segsum(d2e[idx], lengths) - segsum(s2e[idx], lengths)) / lengths - xc = segsum(log_ratio[idx], lengths) / lengths - keep = (m > 0) & (lengths >= MIN_PER_CLASS) - n_dropped += int((~keep).sum()) - if int(keep.sum()) < 2: - continue - xs.append(xc[keep]) - ys.append(torch.log(m[keep])) - ws.append(lengths[keep] / 2.0) - shell_id.append(torch.full_like(xc[keep], float(k))) - if not xs: - return GAMMA_DEFAULT, float("nan"), False, 0, n_dropped, "too_few_classes" - x = torch.cat(xs) - y = torch.cat(ys) - w = torch.cat(ws) - sid = torch.cat(shell_id) - n_used = int(x.numel()) - if n_used < GAMMA_MIN_CLASSES: - return GAMMA_DEFAULT, float("nan"), False, n_used, n_dropped, "too_few_classes" - # Remove each shell's weighted mean from x and y: the slope is then estimated from - # within-shell contrasts only, so shell-to-shell differences in power cannot leak in. - xc = x.clone() - yc = y.clone() - n_shells_used = 0 - for k in torch.unique(sid): - s = sid == k - n_shells_used += 1 - wk = w[s] - xc[s] = x[s] - (wk * x[s]).sum() / wk.sum() - yc[s] = y[s] - (wk * y[s]).sum() / wk.sum() - sxx = (w * xc * xc).sum() - if float(sxx) <= 0.0: - return ( - GAMMA_DEFAULT, - float("nan"), - False, - n_used, - n_dropped, - "no_amplitude_spread", - ) - gamma = float((w * xc * yc).sum() / sxx) - dof = n_used - n_shells_used - 1 - if dof > 0: - resid = yc - gamma * xc - s2 = float((w * resid * resid).sum() / dof) - gamma_se = math.sqrt(max(s2, 0.0) / float(sxx)) - else: - gamma_se = float("nan") - if not math.isfinite(gamma_se) or gamma_se > GAMMA_SE_MAX: - return GAMMA_DEFAULT, gamma_se, False, n_used, n_dropped, "too_uncertain" - lo, hi = GAMMA_BOUNDS - clamped = min(max(gamma, lo), hi) - return clamped, gamma_se, clamped != gamma, n_used, n_dropped, "fitted" - - -def _fit_decay(y: torch.Tensor, var: torch.Tensor, x: torch.Tensor): - """Weighted fit of ``exp(a + b x)``, ``b <= 0``, to signed per-shell power. - - Works on the signed ``B - S2`` of every shell, so shells whose power is zero or - negative by sampling noise pull the curve down instead of being ignored: on a null - dataset the curve goes to zero rather than to the winner's curse of the positive - shells. Gauss-Newton with step halving on the weighted least squares; the fit is - two-parameter and well conditioned. - - Returns ``(curve, a, b)``; ``curve`` is zeros with NaN coefficients when fewer than - four shells are usable or the weighted mean power is not positive. - """ - nan = float("nan") - usable = torch.isfinite(y) & torch.isfinite(var) & (var > 0) - if int(usable.sum()) < 4: - return torch.zeros_like(y), nan, nan - w = torch.where(usable, 1.0 / var.clamp(min=_TINY), torch.zeros_like(var)) - yz = torch.where(usable, y, torch.zeros_like(y)) - mean = float((w * yz).sum() / w.sum()) - if mean <= 0.0: - return torch.zeros_like(y), nan, nan - a = torch.tensor(math.log(mean), dtype=y.dtype, device=y.device) - b = torch.zeros((), dtype=y.dtype, device=y.device) - - def loss(a_, b_): - r = yz - torch.exp(a_ + b_ * x) - return float((w * r * r).sum()) - - current = loss(a, b) - for _ in range(CURVE_ITERS): - f = torch.exp(a + b * x) - r = yz - f - # Jacobian of f with respect to (a, b): f and f*x. - j_a, j_b = f, f * x - g = torch.stack([(w * r * j_a).sum(), (w * r * j_b).sum()]) - h = torch.stack( - [ - torch.stack([(w * j_a * j_a).sum(), (w * j_a * j_b).sum()]), - torch.stack([(w * j_a * j_b).sum(), (w * j_b * j_b).sum()]), - ] - ) - h = ( - h - + 1e-12 * torch.eye(2, dtype=h.dtype, device=h.device) * h.diagonal().max() - ) - step = torch.linalg.solve(h, g) - scale = 1.0 - improved = False - for _ in range(12): - a_new = a + scale * step[0] - b_new = (b + scale * step[1]).clamp(max=0.0) - new = loss(a_new, b_new) - if new < current: - a, b, current, improved = a_new, b_new, new, True - break - scale *= 0.5 - if not improved or float(step.abs().max()) < 1e-9: - break - return torch.exp(a + b * x), float(a), float(b) - - -def estimate_sigma_d( - delta_obs: torch.Tensor, - sigma_diff: torch.Tensor, - epsilon: torch.Tensor | None, - d_star_sq: torch.Tensor, - f_dark: torch.Tensor | None, - fit_mask: torch.Tensor, - *, - delta_calc: torch.Tensor | None = None, - gamma: float | None = None, - shrink: bool | None = None, - per_bin: int = PER_BIN, - min_bins: int = MIN_BINS, - min_per_bin: int = MIN_PER_BIN, -) -> SigmaDShells: - """Per-shell expected difference power, with the dark-amplitude exponent and, - given a difference model, its coupling and unexplained power. - - Runs under ``torch.no_grad()``. The working dtype is the wider of the configured - float dtype and ``delta_obs.dtype``; results are cast back to ``delta_obs.dtype``. - - Parameters - ---------- - delta_obs : torch.Tensor - Signed observed differences ``F_light - F_dark``, shape ``(N,)``, on one common - amplitude scale. - sigma_diff : torch.Tensor - Propagated uncertainty of ``delta_obs``, shape ``(N,)``, same units. - epsilon : torch.Tensor or None - Reflection multiplicity, shape ``(N,)``; ``None`` means ones. - d_star_sq : torch.Tensor - ``1/d**2`` per reflection, shape ``(N,)``, in A^-2. - f_dark : torch.Tensor or None - Dark amplitude for the power law, shape ``(N,)``. ``None`` disables the amplitude - dependence (``gamma`` reported as the default with reason ``"no_f_dark"``). - fit_mask : torch.Tensor - Boolean ``(N,)``: which reflections enter the fit. - delta_calc : torch.Tensor, optional - Model differences ``|F_calc_light| - |F_calc_dark|``, shape ``(N,)``, on the - observed scale. Enables ``alpha`` and ``beta_model``. - gamma : float, optional - Fix the exponent instead of fitting it. - shrink : bool, optional - Shrink the signed shell power toward a decaying curve in ``d*^2``; default the - module setting. - per_bin, min_bins, min_per_bin : int, optional - Shell construction, see :func:`~._shells.equal_count_shells`. - - Returns - ------- - SigmaDShells - One frozen record of per-shell quantities plus counters. - """ - device = delta_obs.device - out_dtype = delta_obs.dtype - dtype = _working_dtype(delta_obs) - shrink = bool(SHRINK_ENABLED if shrink is None else shrink) - if gamma is not None: - lo, hi = GAMMA_BOUNDS - if not (lo <= float(gamma) <= hi): - raise ValueError(f"gamma must lie in {GAMMA_BOUNDS}, got {gamma}") - - with torch.no_grad(): - d_all = delta_obs.reshape(-1).to(dtype) - s_all = sigma_diff.reshape(-1).to(dtype) - x_all = d_star_sq.reshape(-1).to(dtype) - e_all = ( - epsilon.reshape(-1).to(dtype) - if epsilon is not None - else torch.ones_like(d_all) - ) - has_f = f_dark is not None - f_all = f_dark.reshape(-1).to(dtype) if has_f else None - has_model = delta_calc is not None - c_all = delta_calc.reshape(-1).to(dtype) if has_model else None - - finite = ( - torch.isfinite(d_all) - & torch.isfinite(s_all) - & torch.isfinite(x_all) - & torch.isfinite(e_all) - & (s_all >= 0.0) - & (e_all > 0.0) - ) - if has_f: - finite &= torch.isfinite(f_all) - if has_model: - finite &= torch.isfinite(c_all) - fit = fit_mask.reshape(-1).to(torch.bool) - usable = fit & finite - n_dropped = int((fit & ~finite).sum()) - idx = torch.nonzero(usable, as_tuple=True)[0] - n_fit = int(idx.numel()) - - diagnostics = { - "n_dropped": n_dropped, - "n_fit": n_fit, - "n_shell": 0, - "n_s2_clamped": 0, - "n_beta_clamped": 0, - "n_f_floored": 0, - "n_class_dropped": 0, - "n_class_used": 0, - "gamma_se": float("nan"), - "gamma_at_bound": False, - "gamma_reason": "degenerate", - } - if n_fit < 2: - g = float(gamma) if gamma is not None else GAMMA_DEFAULT - return _degenerate(d_all, g, has_model, diagnostics, out_dtype) - - order, seg, seg_lengths, n_bins = equal_count_shells( - x_all[idx], per_bin=per_bin, min_bins=min_bins, min_per_bin=min_per_bin - ) - sel = idx[order] - d, s, e, x = d_all[sel], s_all[sel], e_all[sel], x_all[sel] - counts = seg_lengths.to(dtype) - d2e = d * d / e - s2e = s * s / e - - B = segsum(d2e, seg_lengths) / counts - S2 = segsum(s2e, seg_lengths) / counts - Sigma_N_raw = (B - S2).clamp(min=0.0) - n_s2_clamped = int((S2 >= B).sum()) - bin_dss = segsum(x, seg_lengths) / counts - - # --- dark-amplitude power law ------------------------------------------- - if has_f: - f = f_all[sel] - fbar = (segsum(f, seg_lengths) / counts).clamp(min=_TINY) - fbar_h = fbar[seg] - floor = F_FLOOR_FRAC * fbar_h - n_f_floored = int((f < floor).sum()) - f_fl = torch.maximum(f, floor) - log_ratio = torch.log(f_fl) - torch.log(fbar_h) - bin_log_fbar = torch.log(fbar) - else: - n_f_floored = 0 - log_ratio = torch.zeros_like(d) - bin_log_fbar = torch.zeros_like(bin_dss) - - if gamma is not None: - g_used, g_se, at_bound, n_used, n_cls_dropped, reason = ( - float(gamma), - float("nan"), - False, - 0, - 0, - "fixed", - ) - fitted = False - elif not has_f: - g_used, g_se, at_bound, n_used, n_cls_dropped, reason = ( - GAMMA_DEFAULT, - float("nan"), - False, - 0, - 0, - "no_f_dark", - ) - fitted = False - else: - g_used, g_se, at_bound, n_used, n_cls_dropped, reason = _fit_gamma( - d2e, s2e, log_ratio, seg, n_bins - ) - fitted = reason == "fitted" - - if has_f: - g_raw = torch.exp(g_used * log_ratio) - Z = (segsum(g_raw, seg_lengths) / counts).clamp(min=_TINY) - bin_log_z = torch.log(Z) - else: - bin_log_z = torch.zeros_like(bin_dss) - - # --- difference model ------------------------------------------------------ - if has_model: - c = c_all[sel] - Sigma_P = segsum(c * c / e, seg_lengths) / counts - C = segsum(d * c / e, seg_lengths) / counts - alpha = C / Sigma_P.clamp(min=_TINY) - else: - Sigma_P = torch.zeros_like(B) - C = torch.zeros_like(B) - alpha = torch.ones_like(B) - - # --- stability shrinkage of the signed power toward a decaying curve --------- - # The signed B - S2 keeps every shell as evidence: a shell below zero by noise - # says the power there is small, and it must count. var(B) is 2 B**2 / n for - # Gaussian differences, so that is the sampling variance of each shell's value. - signed = B - S2 - var_s = 2.0 * B * B / counts - if shrink: - curve, curve_a, curve_b = _fit_decay(signed, var_s, bin_dss) - resid = signed - curve - prec = 1.0 / var_s.clamp(min=_TINY) - Q = (prec * resid * resid).sum() - dof = float(max(int(signed.numel()) - 2, 1)) - c = (prec.sum() - (prec * prec).sum() / prec.sum().clamp(min=_TINY)).clamp( - min=_TINY - ) - # Q < dof means the shells scatter no more than their noise: take the curve. - tau_sq = ((Q - dof) / c).clamp(min=0.0) - shrink_w = var_s / (var_s + tau_sq).clamp(min=_TINY) - Sigma_N = ((1.0 - shrink_w) * signed + shrink_w * curve).clamp(min=0.0) - else: - Sigma_N = Sigma_N_raw - shrink_w, tau_sq = torch.zeros_like(B), B.new_zeros(()) - curve_a = curve_b = float("nan") - Sigma_N = torch.where(torch.isfinite(Sigma_N), Sigma_N, torch.zeros_like(B)) - - beta_model = (Sigma_N - alpha * alpha * Sigma_P).clamp(min=0.0) - n_beta_clamped = int(((Sigma_N - alpha * alpha * Sigma_P) < 0.0).sum()) - all_zero = bool((Sigma_N <= 0.0).all()) - - diagnostics.update( - n_shell=int(n_bins), - n_s2_clamped=n_s2_clamped, - n_beta_clamped=n_beta_clamped, - n_f_floored=n_f_floored, - n_class_dropped=n_cls_dropped, - n_class_used=n_used, - gamma_se=g_se, - gamma_at_bound=at_bound, - gamma_reason=reason, - ) - - to = lambda t: t.to(out_dtype) - return SigmaDShells( - B=to(B), - S2=to(S2), - Sigma_N_raw=to(Sigma_N_raw), - Sigma_N=to(Sigma_N), - Sigma_P=to(Sigma_P), - C=to(C), - alpha=to(alpha), - beta_model=to(beta_model), - counts=to(counts), - bin_dss=to(bin_dss), - bin_log_fbar=to(bin_log_fbar), - bin_log_z=to(bin_log_z), - shrink_w=to(shrink_w), - tau=float(tau_sq.clamp(min=0.0).sqrt()), - curve_a=curve_a, - curve_b=curve_b, - gamma=float(g_used), - gamma_fitted=fitted, - has_model=has_model, - degenerate=False, - all_zero=all_zero, - diagnostics=diagnostics, - ) - - -def sigma_d_per_reflection( - shells: SigmaDShells, - d_star_sq: torch.Tensor, - epsilon: torch.Tensor | None, - f_dark: torch.Tensor | None, - sigma_diff: torch.Tensor, -) -> SigmaDEstimate: - """Interpolate a shell estimate onto reflections and form the Wiener weights. - - Parameters - ---------- - shells : SigmaDShells - The shell estimate. - d_star_sq : torch.Tensor - ``1/d**2`` of the output reflections, shape ``(M,)``, in A^-2. - epsilon : torch.Tensor or None - Multiplicity of the output reflections, shape ``(M,)``; ``None`` means ones. - f_dark : torch.Tensor or None - Dark amplitude of the output reflections for the power law; reflections with a - missing or non-finite value get a multiplier of one. - sigma_diff : torch.Tensor - Propagated uncertainty of the output differences, shape ``(M,)``. Non-finite - entries give a weight of zero. - - Returns - ------- - SigmaDEstimate - Per-reflection, detached fields all of length ``M``. - """ - with torch.no_grad(): - dtype = shells.Sigma_N.dtype - grid = d_star_sq.reshape(-1).to(dtype) - eps = ( - epsilon.reshape(-1).to(dtype) - if epsilon is not None - else torch.ones_like(grid) - ) - sig = sigma_diff.reshape(-1).to(dtype) - if shells.degenerate or shells.bin_dss.numel() == 0: - sigma_n = torch.full_like(grid, float(shells.Sigma_N[0])) - alpha = torch.full_like(grid, float(shells.alpha[0])) - beta_model = torch.full_like(grid, float(shells.beta_model[0])) - g = torch.ones_like(grid) - else: - log_sn = interp_in_dss( - grid, shells.bin_dss, torch.log(shells.Sigma_N.clamp(min=_TINY)) - ) - sigma_n = torch.exp(log_sn) - sigma_n = torch.where( - sigma_n > 10.0 * _TINY, sigma_n, torch.zeros_like(sigma_n) - ) - alpha = interp_in_dss(grid, shells.bin_dss, shells.alpha) - log_bm = interp_in_dss( - grid, shells.bin_dss, torch.log(shells.beta_model.clamp(min=_TINY)) - ) - beta_model = torch.exp(log_bm) - beta_model = torch.where( - beta_model > 10.0 * _TINY, beta_model, torch.zeros_like(beta_model) - ) - if f_dark is not None and shells.gamma != 0.0: - f = f_dark.reshape(-1).to(dtype) - log_fbar = interp_in_dss(grid, shells.bin_dss, shells.bin_log_fbar) - log_z = interp_in_dss(grid, shells.bin_dss, shells.bin_log_z) - fbar = torch.exp(log_fbar) - f_fl = torch.maximum(f, F_FLOOR_FRAC * fbar) - g = torch.exp(shells.gamma * (torch.log(f_fl) - log_fbar) - log_z) - g = torch.where(torch.isfinite(f) & (fbar > 0), g, torch.ones_like(g)) - else: - g = torch.ones_like(grid) - S = eps * sigma_n * g - sigma_sq = sig * sig - w = torch.where( - torch.isfinite(sigma_sq), - S / (S + sigma_sq).clamp(min=_TINY), - torch.zeros_like(S), - ) - return SigmaDEstimate( - S=S.detach(), - sigma_sq=sigma_sq.detach(), - w=w.detach(), - alpha=alpha.detach(), - beta_model=(eps * beta_model * g).detach(), - epsilon=eps.detach(), - shells=shells, - ) - - -class SigmaDEstimator: - """Lazy, cached difference-power estimate. - - Thin stateful wrapper around :func:`estimate_sigma_d` and - :func:`sigma_d_per_reflection`: caches the detached estimate and re-estimates only - after :meth:`reset`. **The owning target must call :meth:`reset` from its - ``maintenance()`` hook**, otherwise the estimate is frozen for the whole run. Holds - no tensors of its own beyond the cache, so it has no device to move. - - Parameters - ---------- - config : SigmaDConfig, optional - Exponent and shrinkage settings; the module defaults when omitted. - """ - - def __init__(self, config: SigmaDConfig | None = None): - self.config = config if config is not None else SigmaDConfig() - self._cache: SigmaDEstimate | None = None - self._shells: SigmaDShells | None = None - - def reset(self) -> None: - """Invalidate the cache so the next :meth:`get` re-estimates.""" - self._cache = None - - @property - def shells(self) -> SigmaDShells | None: - """Last shell estimate, for diagnostics; ``None`` until the first call.""" - return self._shells - - def get( - self, - delta_obs: torch.Tensor, - sigma_diff: torch.Tensor, - epsilon: torch.Tensor | None, - d_star_sq: torch.Tensor, - f_dark: torch.Tensor | None, - fit_mask: torch.Tensor, - *, - delta_calc: torch.Tensor | None = None, - target_dss: torch.Tensor | None = None, - out_epsilon: torch.Tensor | None = None, - out_f_dark: torch.Tensor | None = None, - out_sigma_diff: torch.Tensor | None = None, - ) -> SigmaDEstimate: - """Return the cached-or-recomputed :class:`SigmaDEstimate`. - - The fit inputs may be a pooled, flattened set (several datasets end to end); the - ``target_*`` / ``out_*`` arguments map the result onto another reflection list, - defaulting to the fit inputs themselves. - """ - if self._cache is not None: - return self._cache - shells = estimate_sigma_d( - delta_obs, - sigma_diff, - epsilon, - d_star_sq, - f_dark, - fit_mask, - delta_calc=delta_calc, - gamma=self.config.gamma, - shrink=self.config.shrink, - ) - self._shells = shells - self._cache = sigma_d_per_reflection( - shells, - d_star_sq if target_dss is None else target_dss, - epsilon if out_epsilon is None else out_epsilon, - f_dark if out_f_dark is None else out_f_dark, - sigma_diff if out_sigma_diff is None else out_sigma_diff, - ) - return self._cache diff --git a/torchref/refinement/targets/collection/_specs.py b/torchref/refinement/targets/collection/_specs.py index 2d1b9436..5e27fa1b 100644 --- a/torchref/refinement/targets/collection/_specs.py +++ b/torchref/refinement/targets/collection/_specs.py @@ -13,7 +13,7 @@ ``difference`` amplitude ``F_i - F_mean`` against the model's own spread ``difference_i`` intensity the same, in intensities ``difference_sd`` amplitude ``F_i - F_mean`` against ``alpha dF_calc``, variance - ``beta_model + sigma^2`` from a sigma_D fit + ``beta_model + sigma^2`` from a shell-free fit ``two_moment`` intensity ``|F(alpha)|^2 + sigma_alpha^2 |dF|^2`` ``ml`` amplitude each dataset absolutely, at a shared Luzzati beta ==================== =========== ============================================== @@ -141,7 +141,7 @@ def by_name(self, name: str) -> CollectionXrayTargetSpec: name="difference_sd", target_cls=CollectionDifferenceSigmaDTarget, doc="As 'difference', centred on alpha * dF_calc with the unexplained " - "difference power beta_model (sigma_D, fitted on the free set) added to " + "difference power beta_model (fitted on the free set) added to " "the measurement variance.", ), CollectionXrayTargetSpec( diff --git a/torchref/refinement/targets/collection/base.py b/torchref/refinement/targets/collection/base.py index a274ca68..af45714f 100644 --- a/torchref/refinement/targets/collection/base.py +++ b/torchref/refinement/targets/collection/base.py @@ -89,7 +89,7 @@ class CollectionSigmaALossInputs(NamedTuple): class CollectionSigmaDLossInputs(NamedTuple): - """:class:`CollectionLossInputs` plus the sigma_D difference-error estimate. + """:class:`CollectionLossInputs` plus the difference model-error estimate. ``alpha`` and ``beta_model`` live on the **common HKL**, shape ``(n_hkl,)``, and broadcast over the dataset axis: the coupling of the model difference to the true diff --git a/torchref/refinement/targets/collection/xray.py b/torchref/refinement/targets/collection/xray.py index 0e6a40dd..20fa4d71 100644 --- a/torchref/refinement/targets/collection/xray.py +++ b/torchref/refinement/targets/collection/xray.py @@ -9,8 +9,8 @@ The same on intensities -- the whole class is one ``observable`` declaration. :class:`CollectionDifferenceSigmaDTarget` The amplitude difference centred on ``alpha * dF_calc`` with the unexplained - difference power ``beta_model`` added to the measurement variance, both from a - sigma_D fit on the free set. + difference power ``beta_model`` added to the measurement variance, both from one + shell-free fit on the free set. :class:`CollectionMLTarget` Read MLF per dataset at one shared Luzzati ``beta``, pooled over every dataset's free reflections and owned by the target rather than the scaler. The absolute channel. @@ -37,10 +37,9 @@ epsilon_from_hkl, ) from torchref.refinement.model_error_estimation.difference_power import ( - DifferencePowerFit, - fit_difference_power, + DifferencePowerConfig, + DifferencePowerEstimator, ) -from torchref.refinement.model_error_estimation.sigma_d import SigmaDConfig from torchref.utils.stats import VERBOSITY_STANDARD, StatEntry, stat from ._util import common_geom @@ -219,7 +218,7 @@ class CollectionDifferenceSigmaDTarget(CollectionDifferenceTarget): Parameters ---------- - sigma_d_config : SigmaDConfig, optional + difference_config : DifferencePowerConfig, optional Its ``gamma`` fixes the dark-amplitude exponent; fitted when omitted. """ @@ -234,7 +233,7 @@ def __init__( use_work_set: bool = True, use_set: str = None, verbose: int = 0, - sigma_d_config: SigmaDConfig = None, + difference_config: DifferencePowerConfig = None, ): super().__init__( dataset_collection, @@ -245,9 +244,8 @@ def __init__( use_set=use_set, verbose=verbose, ) - self._config = sigma_d_config if sigma_d_config is not None else SigmaDConfig() - # Fitted on first use; maintenance() clears it. - self._fit: DifferencePowerFit = None + # Constructed once; the cached fit lives until maintenance() resets it. + self._estimator = DifferencePowerEstimator(difference_config) self._eps_common: torch.Tensor = None self._dss_common: torch.Tensor = None self._centric_common: torch.Tensor = None @@ -288,7 +286,8 @@ def _loss_inputs(self, recalc: bool = False): eps, dss, centric = self._common_geom() eps, dss = eps.to(dtype), dss.to(dtype) f_dark = ctx.obs[dark] - if self._fit is None: + fit = self._estimator.fit + if fit is None: delta_obs, delta_calc, sigma_diff = self._difference_terms(ctx) rows = [i for i in range(len(ctx.keys)) if i != dark] or [dark] dc = self._dataset_collection @@ -298,7 +297,7 @@ def _loss_inputs(self, recalc: bool = False): fit_mask = torch.cat( [dc[ctx.keys[i]].free.mask.to(ctx.mask.device) for i in rows] ) - self._fit = fit_difference_power( + fit = self._estimator.get( torch.cat([delta_obs[i] for i in rows]), torch.cat([sigma_diff[i] for i in rows]), dss.repeat(n_rows), @@ -307,7 +306,6 @@ def _loss_inputs(self, recalc: bool = False): centric=centric.repeat(n_rows) if centric is not None else None, fit_mask=fit_mask, delta_calc=torch.cat([delta_calc[i].detach() for i in rows]), - gamma=self._config.gamma, fit_sigma_scale=False, ) # A reflection missing from the dark row has no amplitude for the power law; @@ -317,8 +315,8 @@ def _loss_inputs(self, recalc: bool = False): f_eval = torch.where( finite, f_dark, f_dark[finite].median() if bool(finite.any()) else 1.0 ) - alpha = self._fit.alpha_at(dss) - beta = self._fit.signal_power( + alpha = fit.alpha_at(dss) + beta = fit.signal_power( dss, epsilon=eps, f_dark=f_eval, centric=centric ) return CollectionSigmaDLossInputs( @@ -341,14 +339,14 @@ def _per_refl(self, ctx) -> torch.Tensor: return torch.where(torch.isfinite(nll), nll, torch.full_like(nll, 1e6)) def maintenance(self) -> None: - """Clear the fit so it is redone from the updated models on the next forward - (``LossState`` calls this after each optimizer-step block).""" - self._fit = None + """Reset the estimator so it refits from the updated models on the next + forward (``LossState`` calls this after each optimizer-step block).""" + self._estimator.reset() def stats(self) -> Dict[str, StatEntry]: """Base collection X-ray stats plus the difference-power fit summary.""" out = super().stats() - fit = self._fit + fit = self._estimator.fit if fit is not None: lo, hi = fit.stol_range ends = (2.0 * torch.tensor([lo, hi], dtype=fit.coeffs.dtype)) ** 2 From bf4396c66a2c6846a7c2e9c03f8ad89e2646a536 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Wed, 30 Sep 2026 08:39:34 +0200 Subject: [PATCH 209/250] Keep the difference fit defined without residual power and share it in the writer The fit's starting power is floored below the noise, and its amplitude unit falls back to the rms sigma, so all-zero differences or differences exactly explained by the model no longer take the log of zero. The MTZ writer fits the difference SNR once for the q weights and the extrapolation; a failed fit is passed on as its ValueError, so the weights fall back to inverse variance without retrying and the extrapolated amplitudes are written unshrunk. The shrinkage weight is computed as 1 / (1 + 1/snr), exact at an infinite SNR. Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/maps/test_ded_weights.py | 57 +++++++++++++++++-- .../unit/refinement/test_difference_power.py | 34 +++++++++-- torchref/cli/collection_difference_refine.py | 40 ++++++++++--- torchref/maps/ded_weights.py | 17 +++++- .../difference_power.py | 17 ++++-- 5 files changed, 139 insertions(+), 26 deletions(-) diff --git a/tests/unit/maps/test_ded_weights.py b/tests/unit/maps/test_ded_weights.py index 802c0a03..751c6534 100644 --- a/tests/unit/maps/test_ded_weights.py +++ b/tests/unit/maps/test_ded_weights.py @@ -4,9 +4,10 @@ ``inverse_variance`` has mean one and floors a zero sigma; ``q`` gives strong reflections more weight than weak ones where inverse variance cannot, keeps every reflection at or above its floor when noise dominates, and falls back to inverse -variance with a warning that names why when too few reflections exist to fit. The -SNR prefers intensity differences and calibrates their sigmas; the extrapolated -shrinkage keeps every reflection and its weight does not depend on the occupancy. +variance with a warning that names why when too few reflections exist to fit. The SNR +prefers intensity differences, calibrates their sigmas and is reused when supplied; the +extrapolated shrinkage keeps every reflection and its weight does not depend on the +occupancy. """ import pytest @@ -172,8 +173,8 @@ def test_too_few_reflections_fall_back_to_inverse_variance_with_a_warning(): def _intensity_inputs(n=20000, inflation=1.0, seed=3): """Dark/light intensities with known measurement noise, and crude amplitudes.""" g = torch.Generator().manual_seed(seed) - f_dark = (torch.randn(n, generator=g) ** 2 + torch.randn(n, generator=g) ** 2).sqrt() - f_dark = 100.0 * f_dark + re, im = torch.randn(n, generator=g), torch.randn(n, generator=g) + f_dark = 100.0 * (re**2 + im**2).sqrt() f_light = f_dark + torch.randn(n, generator=g) * 5.0 sig_i = 200.0 + 0.05 * f_dark**2 * torch.exp(0.3 * torch.randn(n, generator=g)) i_dark = f_dark**2 + torch.randn(n, generator=g) * sig_i @@ -235,3 +236,49 @@ def test_extrapolated_shrinkage_keeps_every_reflection_and_ignores_occupancy(): assert torch.allclose(var, sig_dark**2 + w * (noise / f) ** 2) assert bool((var > 0).all()) assert torch.equal(out[0.2][2], out[0.5][2]) + # The fallback after a failed fit: an infinite SNR gives the unshrunk amplitude, a + # zero SNR the dark one, both finite. + unshrunk = f_dark + (f_light - f_dark) / 0.2 + for s_val, expect in ((float("inf"), unshrunk), (0.0, f_dark)): + f_ext_b, var, w = compute_bayes_extrapolated_amplitudes( + f_dark, + f_light, + sig_dark, + phi, + phi, + 0.2, + snr=torch.full((n,), s_val), + noise=noise, + ) + assert bool(torch.isfinite(f_ext_b).all()) and bool(torch.isfinite(var).all()) + assert torch.allclose(f_ext_b, expect, atol=1e-4) + + +@pytest.mark.unit +def test_q_reuses_a_supplied_snr_estimate(monkeypatch): + kw = _intensity_inputs(n=5000) + est = difference_snr(**kw) + from torchref.maps import ded_weights as module + + def refuse(**_): + raise AssertionError("fitted again") + + monkeypatch.setattr(module, "difference_snr", refuse) + reused = compute_ded_weights("q", **kw, snr_estimate=est) + monkeypatch.undo() + assert torch.allclose(reused.weights, compute_ded_weights("q", **kw).weights) + + +@pytest.mark.unit +def test_q_takes_a_failed_fit_without_retrying(monkeypatch): + kw = _intensity_inputs(n=5000) + from torchref.maps import ded_weights as module + + def refuse(**_): + raise AssertionError("fitted again") + + monkeypatch.setattr(module, "difference_snr", refuse) + with pytest.warns(DedWeightFallbackWarning) as record: + q = compute_ded_weights("q", **kw, snr_estimate=ValueError("too few")) + assert len(record) == 1 and "too few" in str(record[0].message) + assert q.applied == "inverse_variance" diff --git a/tests/unit/refinement/test_difference_power.py b/tests/unit/refinement/test_difference_power.py index 0b364fb8..87c1b769 100644 --- a/tests/unit/refinement/test_difference_power.py +++ b/tests/unit/refinement/test_difference_power.py @@ -6,9 +6,10 @@ difference's resolution-dependent coupling and the power it leaves unexplained are recovered together; the bounded Wiener weight never falls below its floor, so no reflection or resolution range is removed even when the data hold no signal; the sigma -scale stays within its bounds when the differences hold no noise; the estimator caches -one fit until reset and applies its configured exponent; the fit runs under -``torch.no_grad()`` and on every available device. +scale stays within its bounds when the differences hold no noise; the fit stays defined +when the residual holds no power; the estimator caches one fit until reset and applies +its configured exponent; the fit runs under ``torch.no_grad()`` and on every available +device. """ import pytest @@ -60,9 +61,13 @@ def test_recovers_power_exponent_and_sigma_scale(inflation): @pytest.mark.unit def test_fixed_gamma_is_kept(): s = synth() - fit = fit_difference_power(s["delta"], s["sigma"], s["dss"], f_dark=s["f"], gamma=0.0) + fit = fit_difference_power( + s["delta"], s["sigma"], s["dss"], f_dark=s["f"], gamma=0.0 + ) assert fit.gamma == 0.0 - fit = fit_difference_power(s["delta"], s["sigma"], s["dss"], f_dark=s["f"], gamma=2.0) + fit = fit_difference_power( + s["delta"], s["sigma"], s["dss"], f_dark=s["f"], gamma=2.0 + ) assert fit.gamma == 2.0 @@ -151,3 +156,22 @@ def test_estimator_caches_until_reset_and_applies_its_config(): assert est.get(s["delta"], s["sigma"], s["dss"], f_dark=s["f"]) is not first with pytest.raises(ValueError): DifferencePowerConfig(gamma=10.0) + + +@pytest.mark.unit +def test_fit_is_defined_when_the_residual_holds_no_power(): + # All differences zero, and differences exactly explained by the model: the + # residual power is zero, which the estimator's target path meets on identical data. + s = synth(n=5000) + cases = { + "zero": (torch.zeros_like(s["delta"]), None), + "explained": (s["delta"], s["delta"].clone()), + } + for name, (delta, calc) in cases.items(): + est = DifferencePowerEstimator() + fit = est.get(delta, s["sigma"], s["dss"], f_dark=s["f"], delta_calc=calc) + snr = fit.snr(s["sigma"], d_star_sq=s["dss"], f_dark=s["f"]) + assert bool(torch.isfinite(snr).all()), name + assert float(snr.max()) < 1e-2, name + if calc is not None: + assert float((fit.alpha_at(s["dss"]) - 1.0).abs().max()) < 1e-3 diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index 2204c41f..e73a16db 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -26,6 +26,7 @@ import itertools import json import sys +import warnings from pathlib import Path import torch @@ -55,6 +56,7 @@ from torchref.maps.ded_weights import ( DEFAULT_SCHEME, WEIGHT_COLUMNS, + DedWeightFallbackWarning, all_ded_weights, difference_snr, ) @@ -374,7 +376,8 @@ def compute_bayes_extrapolated_amplitudes( F_dark_phased = Fobs_dark * torch.exp(1j * phi_dark) F_light_phased = Fobs_light * torch.exp(1j * phi_mixed) F_ext = torch.abs(F_dark_phased + (F_light_phased - F_dark_phased) / f) - w = snr / (1.0 + snr) + # snr / (1 + snr), written so an infinite SNR gives exactly 1 rather than inf/inf. + w = 1.0 / (1.0 + 1.0 / snr) var_ext_bayes = sig_dark**2 + w * (noise / f) ** 2 # Shrink the amplitude toward Fo_dark -- scalar, so no phase interference. F_ext_bayes = Fobs_dark + w * (F_ext - Fobs_dark) @@ -689,7 +692,8 @@ def _extrapolation_columns( """Extrapolated light-state amplitudes and the map to view them in. ``snr_est`` is the :class:`~torchref.maps.ded_weights.DifferenceSNR` of the - light-minus-dark differences on the same reflections. + light-minus-dark differences on the same reflections; ``None`` (the fit failed) + writes the unshrunk amplitudes. Three constructions of the same quantity, all needing the light model: @@ -736,6 +740,12 @@ def _fit(amp, sig): sig_light_vals**2 + w_dark**2 * sig_dark_vals**2 ) / w_light + if snr_est is not None: + snr, noise, source = snr_est.snr, snr_est.noise, snr_est.source + else: + snr = torch.full_like(Fobs_dark_vals, float("inf")) + noise = torch.sqrt(sig_dark_vals**2 + sig_light_vals**2) + source = "none" F_ext_bayes_amp, var_ext_bayes, w_shrinkage = compute_bayes_extrapolated_amplitudes( Fobs_dark_vals, Fobs_light_vals, @@ -743,8 +753,8 @@ def _fit(amp, sig): phi_dark, ctx["phi_mixed"], w_light, - snr=snr_est.snr, - noise=snr_est.noise, + snr=snr, + noise=noise, ) sig_ext_bayes = torch.sqrt(var_ext_bayes) @@ -765,7 +775,7 @@ def _np(t): if verbose > 0: print(" Bayes extrapolation rfactors:", rfactor_work_free(data_bayes, amp_calc_bayes)) - print(f" Shrinkage: SNR from {snr_est.source} differences, " + print(f" Shrinkage: SNR from {source} differences, " f"mean w(h) = {w_shrinkage.mean().item():.3f}, " f"min w(h) = {w_shrinkage.min().item():.3g}") @@ -805,7 +815,7 @@ def _np(t): rfactor_work_free(data_scalar, amp_calc_scalar)) diagnostics = { - "shrinkage_source": snr_est.source, + "shrinkage_source": source, "w_shrinkage_mean": float(w_shrinkage.mean().item()), "w_shrinkage_min": float(w_shrinkage.min().item()), } @@ -994,7 +1004,21 @@ def write_results_mtz( sigma_delta_intensity=sig_delta_I, gamma=difference_config.gamma if difference_config is not None else None, ) - all_w = all_ded_weights(**snr_inputs) + # One fit serves the weights and the extrapolation, with one failure policy: when it + # cannot be made, the q weights fall back to inverse variance and the extrapolated + # amplitudes are written unshrunk. + try: + snr_est = difference_snr(**snr_inputs) + except ValueError as err: + snr_est = err + # The q-weight fallback warns for itself; this one covers the extrapolation. + warnings.warn( + f"difference SNR fit failed ({err}); the extrapolated amplitudes are " + "written unshrunk", + DedWeightFallbackWarning, + stacklevel=2, + ) + all_w = all_ded_weights(**snr_inputs, snr_estimate=snr_est) selected = all_w[ded_weight] weights = selected.weights.detach().cpu().numpy() diff_Fobs = diff_t.detach().cpu().numpy() @@ -1066,7 +1090,7 @@ def write_results_mtz( phi_dark=phi_dark, ctx=ctx, rfree_flags_masked=rfree_flags_masked, - snr_est=difference_snr(**snr_inputs), + snr_est=None if isinstance(snr_est, ValueError) else snr_est, all_columns=all_columns, verbose=verbose, ) diff --git a/torchref/maps/ded_weights.py b/torchref/maps/ded_weights.py index 60cdb418..14bf469c 100644 --- a/torchref/maps/ded_weights.py +++ b/torchref/maps/ded_weights.py @@ -193,7 +193,9 @@ def difference_snr( else None ) use_intensity = ( - delta_intensity is not None and sigma_delta_intensity is not None and f is not None + delta_intensity is not None + and sigma_delta_intensity is not None + and f is not None ) if use_intensity: # A floor on the divisor: a near-zero dark amplitude would send both terms to @@ -253,6 +255,7 @@ def compute_ded_weights( fit_mask: torch.Tensor | None = None, gamma: float | None = None, snr_floor: float | None = None, + snr_estimate: DifferenceSNR | ValueError | None = None, ) -> DedWeights: """Per-reflection weights for one scheme. @@ -283,6 +286,10 @@ def compute_ded_weights( Signal-to-noise floor of the ``q`` weight; default :data:`~torchref.refinement.model_error_estimation.difference_power. DEFAULT_SNR_FLOOR`. + snr_estimate : DifferenceSNR or ValueError, optional + A :func:`difference_snr` result on the same inputs, reused instead of fitting + again, so a caller that also needs the SNR fits once; or the ``ValueError`` that + call raised, taken as the failure without retrying. Returns ------- @@ -308,7 +315,9 @@ def compute_ded_weights( floor = DEFAULT_SNR_FLOOR if snr_floor is None else float(snr_floor) try: - est = difference_snr( + if isinstance(snr_estimate, ValueError): + raise snr_estimate + est = snr_estimate or difference_snr( delta_obs=delta_obs, sigma_diff=sigma_diff, hkl=hkl, @@ -339,7 +348,9 @@ def compute_ded_weights( diagnostics = { "source": est.source, "gamma": fit.gamma, - "gamma_fitted": gamma is None and est.source == "amplitude" and f_dark is not None, + "gamma_fitted": ( + gamma is None and est.source == "amplitude" and f_dark is not None + ), "sigma_scale": fit.sigma_scale, "sigma_scale_at_bound": fit.sigma_scale_at_bound, "centric_factor": fit.centric_factor, diff --git a/torchref/refinement/model_error_estimation/difference_power.py b/torchref/refinement/model_error_estimation/difference_power.py index 1709dc1e..f1ec7ed1 100644 --- a/torchref/refinement/model_error_estimation/difference_power.py +++ b/torchref/refinement/model_error_estimation/difference_power.py @@ -299,8 +299,11 @@ def fit_difference_power( d, sig, dss = d[ok], sig[ok], dss[ok] # Work in units of the rms difference so every term of the likelihood is O(1) and - # the Hessian is well conditioned in float32. - amp_scale = float(d.square().mean().sqrt().clamp(min=1e-12)) + # the Hessian is well conditioned in float32; the rms sigma when every difference + # is zero. + amp_scale = float(d.square().mean().sqrt()) + if not amp_scale > 0.0: + amp_scale = float(sig.square().mean().sqrt().clamp(min=1e-12)) d_std = d / amp_scale d2 = d_std.square() log_sig2 = 2.0 * torch.log(sig / amp_scale) @@ -342,9 +345,13 @@ def fit_difference_power( cc = float(c_std.square().mean()) theta[i_a] = float((d_std * c_std).mean()) / cc if cc > 0 else 1.0 d2 = (d_std - theta[i_a] * c_std).square() - excess = float((d2.mean() - log_sig2.exp().mean())) - theta[0] = math.log(max(excess, 0.1 * float(d2.mean()))) - theta[n_c] = float(gamma) if (use_f and gamma is not None) else (1.0 if use_f else 0.0) + noise = float(log_sig2.exp().mean()) + excess = float(d2.mean()) - noise + # Start below the noise when the residual holds no power -- all differences zero, or + # exactly explained by the model -- so the log is always defined. + theta[0] = math.log(max(excess, 0.1 * float(d2.mean()), 1e-3 * noise)) + if use_f: + theta[n_c] = float(gamma) if gamma is not None else 1.0 def nll(t): log_s = basis @ t[:n_c] + t[n_c] * log_f From 23ff661c7a2fbd3cbf1cabf6dca29fc6f477303f Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Wed, 30 Sep 2026 08:50:06 +0200 Subject: [PATCH 210/250] Update the changelog for the shell-free difference weights and target Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 0071b197..44910074 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -6,15 +6,15 @@ Unreleased ---------- - The difference MTZ groups its columns into named datasets -- ``observed``, ``difference``, ``light_model``, ``extrapolated_light``, ``two_moment`` -- with one history line describing each, so ``FWT``/``PHWT`` reads as ``/torchref/extrapolated_light/FWT`` (the extrapolated light-state map ``2*FEXT - Fc``). Labels are unchanged and Coot still auto-opens it - ``SpaceGroup.canonicalize_hkl`` gains ``sort=False``, which keeps the input row order and returns ``None`` for ``sort_indices``; the ASU mapping runs as threaded torch operations and reuses the rotated indices for the Friedel mate, so large reflection lists map faster with unchanged outputs -- ``torchref.difference-map``, ``torchref.difference-refine`` and ``torchref.validate-ded`` gain ``--ded-weight {sigma_d,inverse_variance,none}`` and ``--sigma-d-gamma``. The difference MTZ now carries the unweighted ``DF``/``SIGDF`` on ``PHDELWT`` with one mean-one weight column per scheme, ``W_SD`` and ``W_IVW`` (MTZ type W), and the observed-to-model scale ``KSCALE``; ``DELFWT`` is no longer written, build the map with ``torchref.mtz2map -csf DF -cw W_IVW -cphi PHDELWT``. Registered in ``torchref.maps.ded_weights`` -- Added the ``sigma_D`` estimator (``torchref.refinement.model_error_estimation.sigma_d``): the expected true difference power per resolution shell, ``mean(dF_obs^2) - mean(sigma^2)`` with a fitted ``F_dark^gamma`` amplitude law and DerSimonian-Laird shrinkage of the signed shell power toward a decaying exponential in ``d*^2`` (fitted on all shells, so a dataset without a difference yields no power instead of the positive half of its noise), giving the Wiener weight ``S/(S + sigma^2)`` and, with a difference model, ``alpha``/``beta_model``. Inverse-variance weights suppress the strong reflections whose difference power is 10-70x that of weak ones; on independent half-datasets a Wiener weight with the true power raised map agreement 1.2-1.8x in effective patterns. The single-dataset estimate inherits the calibration of the reported sigmas, and on the campaign TD1 data (sigmas ~1.5x too large at high resolution) it emptied 60-90 % of the shells, so inverse variance stays the default; ``sigma_d`` reports its clamped-shell count and falls back to inverse variance with a warning when every shell is empty -- Added the ``difference_sd`` collection target (``CollectionDifferenceSigmaDTarget``): the difference Gaussian centred on ``alpha * dF_calc`` with variance ``beta_model + sigma_diff^2`` from ``sigma_D`` fitted on the free set. Selected with ``torchref.difference-refine --difference-target difference_sd``; ``difference`` stays the default +- ``torchref.difference-map``, ``torchref.difference-refine`` and ``torchref.validate-ded`` gain ``--ded-weight {q,inverse_variance,none}`` (default ``q``) and ``--difference-gamma``. The difference MTZ now carries the unweighted ``dF``/``SIGdF`` on ``PHDELWT`` with one mean-one weight column per scheme, ``W_Q`` and ``W_InVa`` (MTZ type W), and the observed-to-model scale ``KSCALE``; ``DELFWT`` is no longer written, build the map with ``torchref.mtz2map -csf dF -cw W_Q -cphi PHDELWT``. Registered in ``torchref.maps.ded_weights`` +- Added ``torchref.refinement.model_error_estimation.difference_power``: ``fit_difference_power`` fits the expected power of a light-minus-dark difference by per-reflection maximum likelihood -- ``log S`` a Chebyshev series in resolution plus ``gamma log F_dark``, a fitted scale ``k`` on the reported sigmas and a centric factor -- with no resolution shells; given a model difference it also fits the coupling ``alpha`` and the unexplained power. ``DifferencePowerEstimator`` caches one fit for a target, ``DifferencePowerConfig`` fixes ``gamma``. ``torchref.maps.ded_weights.difference_snr`` turns it into a per-reflection signal-to-noise ratio, fitted on intensity differences when the data carry ``I``/``SIGI``, because French-Wilson amplitude sigmas overstate the noise of a difference. The ``q`` weight is ``(snr + 1/2) / (snr + 3/2)``, which down-weights a noisy reflection to no less than a third and never removes a resolution range +- Added the ``difference_sd`` collection target (``CollectionDifferenceSigmaDTarget``): the difference Gaussian centred on ``alpha * dF_calc`` with variance ``beta_model + sigma_diff^2``, both from one ``fit_difference_power`` on the free set with the reported sigmas taken as calibrated. Selected with ``torchref.difference-refine --difference-target difference_sd``; ``difference`` stays the default. Its stats report ``difference_gamma`` and ``difference_alpha_low_res``/``difference_alpha_high_res`` - ``torchref.mtz2map`` gains ``--column-weight``/``-cw`` (multiply the amplitudes by a weight column before the FFT) and ``--units {sigma,electrons,raw}``: ``electrons`` writes e/A^3 as ``(1/V) sum_h F(h) exp(-2 pi i h.x)`` with the amplitudes divided by the ``--column-scale``/``-ck`` factor (``KSCALE`` by default). ``-n``/``--normalize`` is a deprecated alias. ``Map`` and ``DifferenceMap`` take ``units`` too, and ``DifferenceMap`` an optional per-reflection ``scale`` - ``ScalerBase.multiplicative_scale()`` returns the per-reflection ``K_overall * b_overall * anisotropy`` factor, every multiplicative component of ``forward`` and none of the additive solvent term -- ``torchref.validate-ded`` reports the real- and reciprocal-space correlations for unweighted, inverse-variance and ``sigma_D`` weights side by side (``by_weight`` in the JSON, a table in the summary) and records when ``sigma_D`` fell back to inverse variance -- The extrapolated-map Bayes shrinkage estimates its signal variance per resolution shell through ``sigma_D`` instead of one global ``tau^2``; ``tau_sq`` in the summary is now the count-weighted shell mean +- ``torchref.validate-ded`` reports the real- and reciprocal-space correlations for unweighted, inverse-variance and ``q`` weights side by side (``by_weight`` in the JSON, a table in the summary) and records when ``q`` fell back to inverse variance +- The extrapolated amplitudes ``FEXT`` shrink toward ``Fo_dark`` by ``w = snr / (1 + snr)`` from the same ``difference_snr`` the ``q`` weight uses, positive for every reflection with signal, and ``SIGFEXT`` includes the dark measurement sigma; a failed fit writes them unshrunk - Inverse-variance difference weights floor ``sigma_diff`` at a tenth of its median, so a zero sigma yields a large finite weight rather than an infinite one -- Shell construction, segment sums, interpolation and the shrinkage line fit shared by ``sigma_A`` and ``sigma_D`` live in ``torchref.refinement.model_error_estimation._shells``; ``estimate_beta`` is unchanged +- Shell construction, segment sums, interpolation and the shrinkage line fit used by ``sigma_A`` live in ``torchref.refinement.model_error_estimation._shells``; ``estimate_beta`` is unchanged - Align difference-refinement tests with fixture ownership, integration placement and configured dtype/device conventions; share fresh collection setup and exercise noise statistics on deposited amplitudes. - Read scaled observations directly in subset and collection accessors, avoiding unused sigma/amplitude corrections and removing redundant internal forwarding helpers. - Remove one-off diagnostic scripts and consolidate difference-refinement regression tests while retaining numerical and output-format coverage. @@ -53,8 +53,8 @@ Unreleased - Fixed the empirical-Bayes extrapolation propagating ``sigma_ext^2 = (sigma_L^2 + sigma_D^2)/f^2``. ``F_ext`` is linear in the observations with ``dF_ext/dF_dark = -(1-f)/f``, so the dark term carries a ``(1-f)^2`` weight: at f = 0.22 the term was over-weighted 1.64x, biasing ``tau^2`` low and over-shrinking every reflection. The estimator now takes the caller's already-correct propagation instead of rebuilding it, so the shrinkage weight and the written ``SIGFEXT`` cannot disagree. It also returns the shrunk amplitude it is named for; the caller used to discard the return value and recompute it. Correcting it raises tau^2 and shrinks less, which moves the default extrapolated map - ``FWT``/``PHWT`` carries the phase of its own scale fit. The Bayes coefficients had no phase column and would have been paired with one fitted against the phase-aware amplitudes; the scaler contributes a phase through ``f_sol``, so the three extrapolation fits do not agree - ``write_results_mtz`` is four layers, each declaring its columns' MTZ types beside the values, and it now calls ``infer_mtz_dtypes`` as the canonical writer does. Types used to come from four parallel name lists with no fallback, so a column added to the output dict and missed in the lists was written with whatever dtype numpy produced. A column with no declared type is now an error. Its column table moved onto the function from the module docstring four hundred lines away -- ``DFc_complex`` is ``DFc_phased``: the column holds a real amplitude, and the old name said otherwise. The undocumented ``Fextp``/``Fextc``/``Fextb`` suffixes are ``FEXT_PHASED``/``FEXT_SCALAR``/``FEXT``, and ``SIGFEXT_PHASED`` is written at last -- it was computed all along while the docs claimed it existed -- ``tau_sq`` and mean ``w(h)`` go in the JSON summary. With the Bayes extrapolation as the default map, they are what says whether it is over-shrunk +- ``DF``/``SIGDF`` are ``dF``/``SIGdF`` and ``DFc`` is ``dFc``, a lower-case delta for the amplitude differences. ``DFc_complex`` is ``dFc_phased``: the column holds a real amplitude, and the old name said otherwise. The undocumented ``Fextp``/``Fextc``/``Fextb`` suffixes are ``FEXT_PHASED``/``FEXT_SCALAR``/``FEXT``, and ``SIGFEXT_PHASED`` is written at last -- it was computed all along while the docs claimed it existed +- ``shrinkage_source`` and the mean and least ``w(h)`` of the extrapolation shrinkage go in the JSON summary. With the shrunk extrapolation as the default map, they are what says whether it is over-shrunk - Fixed ``paper/make_ded_maps.py`` pairing ``WDF`` with ``PHIC_diff`` -- the right amplitude on the model difference phase, which is not the map ``validate-ded`` reports - Refinement output no longer inherits the input file's refinement header. It used to copy the whole thing and then append its own ``REMARK 3``, so a refined 3GR5 carried 420 header lines asserting two refinements at once -- ``PROGRAM : REFMAC 5.1.24`` with R-work 0.213 at line 5, ours at line 389 -- and a reader taking the first ``REMARK 3`` got REFMAC. The inherited block was not merely stale but contradicted the data beside it: it claimed a 5.1% / 1072-reflection test set, while the MTZ shipped with it holds 9.85% / 2063 (which torchref reads correctly). The passthrough was also inverted, keeping the statistics refinement invalidates and dropping the chemistry it does not -- SEQRES, SSBOND, DBREF, EXPDTA, COMPND, SOURCE, KEYWDS, SEQADV, HETNAM, FORMUL and SITE were all absent from the output. Now a whitelist carries the crystal, sample and chemistry records through in mandated record order (TITLE used to be emitted after REMARK 900), REMARK 2, 3 and 500 are dropped, and AUTHOR and JRNL are not inherited because they credit the deposition rather than this run. 283 header lines for the same file, 41 structural records preserved, one refinement block. ``add-metadata`` is exempt through ``supersede_refinement=False``: annotating a file is not re-refining it, so nothing there supersedes the existing REMARK 3 or AUTHOR records and both are kept - Prior refinements are tracked through mmCIF's ``_software`` loop, which is the only place either format has room for them: ``_refine`` is singular by design, so a previous program's statistics cannot be kept without contradicting the current ones. ``pdbx_ordinal`` was hardcoded to ``1`` and the incoming loop was never read, truncating the chain to one link on every write; it now reads the input's loop and appends at ``max(ordinal) + 1``, carrying each entry's ``description`` so the chain says what every program did and not just that it ran. Added ``_pdbx_initial_refinement_model`` and ``_refine.pdbx_starting_model``, which name what the refinement started from, and ``_refine.pdbx_R_Free_selection_details``, which names the test set the reported R-free conditions on. ``from_cif_file`` no longer carries the input's ``_refine`` items through -- that was the mmCIF form of the duplicated ``REMARK 3`` From 854fa4c55d984512e77aac878210bcc63d71762a Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 18:44:33 +0200 Subject: [PATCH 211/250] Replace the FrenchWilson module with french_wilson_auto ReflectionData.load kept the FrenchWilson nn.Module it built, but its per-row d-spacing and centric buffers were in the pre-canonicalization row order. difference-refine reused it whenever the length still matched, so on files stored off the CCP4 ASU order (6G9X) the corrected light amplitudes were converted against the wrong reflections. The conversion runs once per dataset, so it is now a plain function: french_wilson_auto gains the NaN handling the module had and returns (F, sigma_F, valid_mask). The dataset no longer holds an nn.Module. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 2 + tests/helpers/device_cases.py | 1 - tests/integration/test_crystfel_hkl.py | 2 +- tests/unit/io/test_data.py | 6 +- tests/unit/io/test_wilson_outlier_masks.py | 47 ++-- torchref/base/__init__.py | 6 +- torchref/base/french_wilson.py | 245 +++---------------- torchref/cli/collection_difference_refine.py | 21 +- torchref/io/datasets/base.py | 6 +- torchref/io/datasets/reflection_data.py | 16 +- 10 files changed, 91 insertions(+), 261 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 0071b197..d5aa5f85 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,8 @@ Changelog Unreleased ---------- +- Removed the ``FrenchWilson`` module and ``ReflectionData._FrenchWilson``; use ``french_wilson_auto(I, sigma_I, hkl, d_spacings, space_group)``, which returns ``(F, sigma_F, valid_mask)`` +- Fixed ``torchref.difference-refine`` reusing the French-Wilson estimator built at load time, whose d-spacings and centric flags were in the pre-canonicalization row order; on files stored off the CCP4 ASU order (e.g. 6G9X) the corrected light amplitudes were computed against the wrong reflections - The difference MTZ groups its columns into named datasets -- ``observed``, ``difference``, ``light_model``, ``extrapolated_light``, ``two_moment`` -- with one history line describing each, so ``FWT``/``PHWT`` reads as ``/torchref/extrapolated_light/FWT`` (the extrapolated light-state map ``2*FEXT - Fc``). Labels are unchanged and Coot still auto-opens it - ``SpaceGroup.canonicalize_hkl`` gains ``sort=False``, which keeps the input row order and returns ``None`` for ``sort_indices``; the ASU mapping runs as threaded torch operations and reuses the rotated indices for the Friedel mate, so large reflection lists map faster with unchanged outputs - ``torchref.difference-map``, ``torchref.difference-refine`` and ``torchref.validate-ded`` gain ``--ded-weight {sigma_d,inverse_variance,none}`` and ``--sigma-d-gamma``. The difference MTZ now carries the unweighted ``DF``/``SIGDF`` on ``PHDELWT`` with one mean-one weight column per scheme, ``W_SD`` and ``W_IVW`` (MTZ type W), and the observed-to-model scale ``KSCALE``; ``DELFWT`` is no longer written, build the map with ``torchref.mtz2map -csf DF -cw W_IVW -cphi PHDELWT``. Registered in ``torchref.maps.ded_weights`` diff --git a/tests/helpers/device_cases.py b/tests/helpers/device_cases.py index d2555971..aa1de517 100644 --- a/tests/helpers/device_cases.py +++ b/tests/helpers/device_cases.py @@ -532,7 +532,6 @@ class TargetDeviceCase: "_SharedMixedModel": "internal view owned by ModelCollection", "Scaler": "needs a loaded model + data; covered in integration", "Restraints": "needs a model + monomer library", - "FrenchWilson": "needs loaded intensities", "DatasetCollection": "needs several loaded datasets", "FcalcDataset": "needs computed structure factors", "Map": "needs data + model", diff --git a/tests/integration/test_crystfel_hkl.py b/tests/integration/test_crystfel_hkl.py index 4dc19c88..65a955bd 100644 --- a/tests/integration/test_crystfel_hkl.py +++ b/tests/integration/test_crystfel_hkl.py @@ -33,7 +33,7 @@ def test_observations_and_metadata(halves, index, test_files_dir): assert data.I.shape == data.I_sigma.shape == data.F.shape == (len(data),) assert torch.isfinite(data.I_sigma).all() and (data.I_sigma >= 0).all() assert (data.I < 0).any() - assert (data.F >= 0).all() and data._FrenchWilson is not None + assert (data.F >= 0).all() and data.FRENCH_WILSON_MASK_KEY in data.masks torch.testing.assert_close(data.cell.data, data.cell.data.new_tensor(CELL)) assert data.spacegroup.number == 1 diff --git a/tests/unit/io/test_data.py b/tests/unit/io/test_data.py index a615bfba..e4e58a6d 100644 --- a/tests/unit/io/test_data.py +++ b/tests/unit/io/test_data.py @@ -224,7 +224,7 @@ def test_french_wilson_off_uses_amplitudes_directly(self): # F should be exactly the sentinel amplitude column, untouched. assert torch.allclose(data.F, torch.full_like(data.F, 7.0)) # French-Wilson must not have run. - assert data._FrenchWilson is None + assert data.FRENCH_WILSON_MASK_KEY not in data.masks @pytest.mark.unit def test_french_wilson_on_derives_from_intensities(self): @@ -236,7 +236,7 @@ def test_french_wilson_on_derives_from_intensities(self): # F is computed from I, so it differs from the sentinel 7.0 column. assert not torch.allclose(data.F, torch.full_like(data.F, 7.0)) - assert data._FrenchWilson is not None + assert data.FRENCH_WILSON_MASK_KEY in data.masks @pytest.mark.unit def test_french_wilson_off_falls_back_when_no_amplitudes(self): @@ -252,4 +252,4 @@ def test_french_wilson_off_falls_back_when_no_amplitudes(self): # No amplitude columns => French-Wilson runs regardless of the flag. assert data.F is not None - assert data._FrenchWilson is not None + assert data.FRENCH_WILSON_MASK_KEY in data.masks diff --git a/tests/unit/io/test_wilson_outlier_masks.py b/tests/unit/io/test_wilson_outlier_masks.py index 23f01b26..6372046f 100644 --- a/tests/unit/io/test_wilson_outlier_masks.py +++ b/tests/unit/io/test_wilson_outlier_masks.py @@ -16,6 +16,7 @@ import pytest import torch +from torchref.base.french_wilson import french_wilson_auto from torchref.io.datasets.reflection_data import ReflectionData CELL = (50.0, 60.0, 70.0, 90.0, 90.0, 90.0) @@ -145,15 +146,34 @@ def test_intensity_path_keeps_french_wilsons_guard_under_its_own_key(mtz_dir): assert data.I is not None, "4BX9 should load via the intensity path" assert ReflectionData.FRENCH_WILSON_MASK_KEY in data.masks - torch.testing.assert_close( - data.masks[ReflectionData.FRENCH_WILSON_MASK_KEY].sum(), - data._FrenchWilson.valid_mask.sum(), + _, _, keep = french_wilson_auto( + data.I, data.I_sigma, data.hkl, data.resolution, data.spacegroup ) + torch.testing.assert_close(data.masks[ReflectionData.FRENCH_WILSON_MASK_KEY], keep) # And the outlier test still ran on top of it, rather than being skipped # because a mask was already present. assert ReflectionData.WILSON_MASK_KEY in data.masks +@pytest.mark.unit +def test_french_wilson_is_row_aligned_after_canonicalization(mtz_dir): + """Converting the loaded intensities again reproduces the loaded amplitudes. + + 6G9X is stored off the CCP4 ASU order, so ``load`` reorders its rows after + French-Wilson has run; F, sigma_F and the guard mask must move with them. + """ + data = ReflectionData(verbose=0).load_mtz(str(mtz_dir / "6G9X.mtz")) + assert data.I is not None, "6G9X should load via the intensity path" + + F, sigma_F, keep = french_wilson_auto( + data.I, data.I_sigma, data.hkl, data.resolution, data.spacegroup + ) + + torch.testing.assert_close(data.masks[ReflectionData.FRENCH_WILSON_MASK_KEY], keep) + torch.testing.assert_close(F[keep], data.F[keep]) + torch.testing.assert_close(sigma_F[keep], data.F_sigma[keep]) + + # ============================================================================= # Deposited data # ============================================================================= @@ -276,21 +296,18 @@ def test_too_few_reflections_are_left_alone(): @pytest.mark.unit -def test_french_wilson_records_its_own_mask_full_size(): - from torchref.base.french_wilson import FrenchWilson - +def test_french_wilson_returns_its_own_mask_full_size(): hkl = torch.tensor([[1, 0, 0], [2, 0, 0], [3, 0, 0], [4, 0, 0]]) - fw = FrenchWilson(hkl, torch.tensor(CELL), "P 1", verbose=0) - assert fw.valid_mask is None - + d = CELL[0] / hkl[:, 0].float() I = torch.tensor([100.0, 50.0, -5.0, float("nan")]) sigma_I = torch.tensor([10.0, 8.0, 7.0, 5.0]) - fw(I, sigma_I) - assert fw.valid_mask is not None - assert fw.valid_mask.shape == I.shape - assert fw.valid_mask.dtype == torch.bool + F, sigma_F, keep = french_wilson_auto(I, sigma_I, hkl, d, "P 1") + + assert keep.shape == I.shape + assert keep.dtype == torch.bool # Well-measured reflections survive; the NaN row never converted, so it is # not kept on the strength of a comparison that was never made. - assert fw.valid_mask[:2].all() - assert not bool(fw.valid_mask[3]) + assert keep[:2].all() + assert not bool(keep[3]) + assert torch.isnan(F[3]) and torch.isnan(sigma_F[3]) diff --git a/torchref/base/__init__.py b/torchref/base/__init__.py index e62d4747..6428f054 100644 --- a/torchref/base/__init__.py +++ b/torchref/base/__init__.py @@ -74,9 +74,9 @@ ) # ============================================================================= -# Main classes +# French-Wilson # ============================================================================= -from .french_wilson import FrenchWilson +from .french_wilson import french_wilson_auto # ============================================================================= # Coordinate transformations (from coordinates submodule) @@ -248,7 +248,7 @@ # ------------------------------------------------------------------------- # Classes # ------------------------------------------------------------------------- - "FrenchWilson", + "french_wilson_auto", "ReciprocalSymmetryExtractor", # ------------------------------------------------------------------------- # Coordinate transforms diff --git a/torchref/base/french_wilson.py b/torchref/base/french_wilson.py index a03c6df2..d2060919 100644 --- a/torchref/base/french_wilson.py +++ b/torchref/base/french_wilson.py @@ -4,42 +4,23 @@ Reference: French, S. & Wilson, K. (1978). Acta Cryst. A34, 517-525 Based on Phenix implementation in cctbx/french_wilson.py -Usage - PyTorch Module (Recommended):: - - import torch - from torchref.base.french_wilson import FrenchWilson - - # Miller indices for your reflections - hkl = torch.tensor([[1, 2, 3], [2, 0, 0], [0, 3, 0], [1, 1, 1]]) - - # Cell: [a, b, c, alpha, beta, gamma] in Å and degrees - cell = [50.0, 60.0, 70.0, 90.0, 90.0, 90.0] - - # Create module (does all preprocessing) - fw_module = FrenchWilson(hkl, cell, space_group='P212121') - - # Apply conversion (can be called repeatedly with different I, sigma_I) - I = torch.tensor([100.0, 50.0, 30.0, 200.0]) - sigma_I = torch.tensor([10.0, 8.0, 7.0, 15.0]) - F, sigma_F = fw_module(I, sigma_I) - print(f"F = {F}") - -Usage - Functional API (for one-off conversions):: +Usage:: from torchref.base.french_wilson import french_wilson_auto F, sigma_F, valid = french_wilson_auto( I, sigma_I, hkl, d_spacings, space_group='P212121' ) + +This is a plain function on purpose: the conversion runs once per dataset, and +a cached estimator holding per-row buffers goes stale the moment the rows are +reordered (as ``ReflectionData`` canonicalization does). """ import torch -import torch.nn as nn -from torchref.base import math_torch from torchref.config import get_float_dtype from torchref.symmetry import SpaceGroup, SpaceGroupLike -from torchref.utils.device_mixin import DeviceMixin # Acentric lookup tables from French-Wilson supplement (1978) AC_ZJ = torch.tensor( @@ -1279,13 +1260,10 @@ def french_wilson_auto( h_min: float = -4.0, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """ - Automatic French-Wilson conversion with binning and centric determination. + Convert intensities to amplitudes, estimating shell means and centricity. - This function automatically: - 1. Bins reflections by resolution - 2. Calculates mean intensity per bin - 3. Determines centric vs acentric from Miller indices - 4. Applies appropriate French-Wilson conversion + Every per-reflection input must be row-aligned: the shell mean and centric + flag of row ``i`` are taken from ``hkl[i]`` and ``d_spacings[i]``. Parameters ---------- @@ -1313,7 +1291,9 @@ def french_wilson_auto( sigma_F : torch.Tensor Standard deviations of F of shape (n_reflections,). valid_mask : torch.Tensor - Boolean mask indicating valid (not rejected) reflections. + Boolean mask, ``True`` = keep. ``False`` both for rows French-Wilson + rejects as too negative and for rows with NaN ``I`` or ``sigma_I``, + whose ``F`` and ``sigma_F`` are NaN. Examples -------- @@ -1325,191 +1305,26 @@ def french_wilson_auto( d_spacings = torch.tensor([2.5, 3.0, 2.8, 2.0]) F, sigma_F, valid = french_wilson_auto(I, sigma_I, hkl, d_spacings, "P212121") """ - # Step 1: Estimate mean intensity by resolution - mean_intensity = estimate_mean_intensity_by_resolution( - I, d_spacings, n_bins=n_bins, min_per_bin=min_per_bin - ) + F = torch.full_like(I, float("nan")) + sigma_F = torch.full_like(sigma_I, float("nan")) + # NaN rows are never converted, so they are not kept either. + valid_mask = torch.zeros_like(I, dtype=torch.bool) - # Step 2: Determine centric reflections from Miller indices - is_centric = SpaceGroup(space_group, device=hkl.device).is_centric(hkl) + finite = ~(torch.isnan(I) | torch.isnan(sigma_I)) + if not finite.any(): + return F, sigma_F, valid_mask - # Step 3: Apply French-Wilson conversion - F, sigma_F, valid_mask = french_wilson( - I, sigma_I, mean_intensity, is_centric=is_centric, h_min=h_min + # NaN rows are dropped before binning, or they would poison the shell means. + mean_intensity = estimate_mean_intensity_by_resolution( + I[finite], d_spacings[finite], n_bins=n_bins, min_per_bin=min_per_bin + ) + is_centric = SpaceGroup(space_group, device=hkl.device).is_centric(hkl[finite]) + + F[finite], sigma_F[finite], valid_mask[finite] = french_wilson( + I[finite], + sigma_I[finite], + mean_intensity, + is_centric=is_centric, + h_min=h_min, ) - return F, sigma_F, valid_mask - - -class FrenchWilson(DeviceMixin, nn.Module): - """ - PyTorch module for French-Wilson conversion from intensities to structure factors. - - Pre-computes all necessary metadata (d-spacings, centric flags, resolution bins) - during initialization, so forward pass only needs I and sigma_I. - - Parameters - ---------- - hkl : torch.Tensor - Miller indices of shape (n_reflections, 3), integer tensor. - cell : torch.Tensor - Unit cell parameters [a, b, c, alpha, beta, gamma] in Å and degrees. - space_group : str, int, or gemmi.SpaceGroup, optional - Space group specification (e.g., 'P21', 4, gemmi.SpaceGroup('P 21')). Default is "P1". - n_bins : int, optional - Number of resolution bins for mean intensity estimation. Default is 60. - min_per_bin : int, optional - Minimum reflections per bin. Default is 40. - h_min : float, optional - Minimum h value for rejection. Default is -4.0. - verbose : int, optional - Verbosity level (0=silent, 1=basic, 2=detailed). Default is 1. - - Attributes - ---------- - hkl : torch.Tensor - Miller indices. - d_spacings : torch.Tensor - Resolution for each reflection in Å. - is_centric : torch.Tensor - Boolean mask for centric reflections. - valid_mask : torch.Tensor or None - French-Wilson's rejection criterion from the most recent :meth:`forward` - (``True`` = keep). ``None`` before the first call. - - Examples - -------- - :: - - hkl = torch.tensor([[1, 2, 3], [2, 0, 0], [0, 3, 0], [1, 1, 1]]) - cell = [50.0, 60.0, 70.0, 90.0, 90.0, 90.0] - fw_module = FrenchWilson(hkl, cell, 'P212121') - I = torch.tensor([100.0, 50.0, 30.0, 200.0]) - sigma_I = torch.tensor([10.0, 8.0, 7.0, 15.0]) - F, sigma_F = fw_module(I, sigma_I) - """ - - def __init__( - self, - hkl: torch.Tensor, - cell: torch.Tensor, - space_group: SpaceGroupLike = "P1", - n_bins: int = 60, - min_per_bin: int = 40, - h_min: float = -4.0, - verbose: int = 1, - ): - super().__init__() - - # Store parameters - self.n_reflections = len(hkl) - self.space_group = space_group - self.n_bins = n_bins - self.min_per_bin = min_per_bin - self.h_min = h_min - self.verbose = verbose - - # Register HKL as buffer (will be moved to device with model) - self.register_buffer("hkl", hkl.long()) - - # Calculate d-spacings from cell and HKL - d_spacings = math_torch.get_d_spacing(hkl, cell) - self.register_buffer("d_spacings", d_spacings) - - # Determine centric reflections - is_centric = SpaceGroup(space_group, device=hkl.device).is_centric(hkl) - self.register_buffer("is_centric", is_centric) - - # Set by forward(); None until the first conversion. Not a buffer -- it - # is per-call output, not model state to serialize or move. - self.valid_mask = None - - # Verbosity level 1: Basic initialization info (most important) - if self.verbose >= 1: - print("FrenchWilson initialized:") - print(f" Reflections: {self.n_reflections}") - print(f" Resolution: {d_spacings.min():.2f} - {d_spacings.max():.2f} Å") - print(f" Space group: {space_group}") - print( - f" Centric: {is_centric.sum()} ({100*is_centric.sum()/self.n_reflections:.1f}%)" - ) - - # Verbosity level 2: Additional detailed info (less important) - if self.verbose >= 2: - print(f" Binning: {n_bins} bins, min {min_per_bin} reflections/bin") - print(f" Rejection threshold: h_min = {h_min}") - print(f" Device: {hkl.device}") - - def forward( - self, I: torch.Tensor, sigma_I: torch.Tensor - ) -> tuple[torch.Tensor, torch.Tensor]: - """ - Apply French-Wilson conversion. - - Parameters - ---------- - I : torch.Tensor - Measured intensities of shape (n_reflections,). - sigma_I : torch.Tensor - Standard deviations of intensities of shape (n_reflections,). - - Returns - ------- - F : torch.Tensor - Structure factor amplitudes of shape (n_reflections,). - sigma_F : torch.Tensor - Standard deviations of F of shape (n_reflections,). - - Notes - ----- - The rejection criterion is recorded on :attr:`valid_mask` (full size, - ``True`` = keep, ``False`` for both NaN input and French-Wilson - rejections) rather than returned, so this stays a two-tuple for the - documented usage above. It is overwritten on each call. - """ - # Check for NaN values in input - nan_mask = torch.isnan(I) | torch.isnan(sigma_I) - - # If all values are NaN, return NaN arrays - if nan_mask.all(): - self.valid_mask = torch.zeros_like(I, dtype=torch.bool) - return torch.full_like(I, float("nan")), torch.full_like( - sigma_I, float("nan") - ) - - # Filter out NaN values and corresponding metadata - finite_mask = ~nan_mask - I_clean = I[finite_mask] - sigma_I_clean = sigma_I[finite_mask] - d_spacings_clean = self.d_spacings[finite_mask] - is_centric_clean = self.is_centric[finite_mask] - - # Estimate mean intensity by resolution (only for valid reflections) - mean_intensity = estimate_mean_intensity_by_resolution( - I_clean, d_spacings_clean, n_bins=self.n_bins, min_per_bin=self.min_per_bin - ) - - # Apply French-Wilson conversion - F_clean, sigma_F_clean, keep_clean = french_wilson( - I_clean, - sigma_I_clean, - mean_intensity, - is_centric=is_centric_clean, - h_min=self.h_min, - ) - - # Create output arrays with NaNs for invalid reflections - F_full = torch.full_like(I, float("nan")) - sigma_F_full = torch.full_like(sigma_I, float("nan")) - - # Insert computed values for valid reflections - F_full[finite_mask] = F_clean - sigma_F_full[finite_mask] = sigma_F_clean - - # Expand the rejection criterion back to full size. NaN rows are - # rejected too -- they were never converted. - keep_full = torch.zeros_like(I, dtype=torch.bool) - keep_full[finite_mask] = keep_clean - self.valid_mask = keep_full - - return F_full, sigma_F_full diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index a6b36909..be365a99 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -30,6 +30,7 @@ import torch +from torchref.base.french_wilson import french_wilson_auto from torchref.cli._common import ( add_all_columns_arg, add_ded_weight_args, @@ -465,24 +466,20 @@ def _two_moment_columns(mc, dc, mask, fcalc_dark_full, fcalc_mixed_full, return empty with torch.no_grad(): - # Full-size, so French-Wilson sees the reflection list it was fitted on. + # Full-size, so French-Wilson estimates its shell means from every reflection. delta_F_full = fcalc_mixed_full - fcalc_dark_full variance_full = mc.sigma_alpha_sq * delta_F_full.abs() ** 2 I_light_full, sig_I_full = data_light.get_corrected_intensities() I_corrected_full = I_light_full - variance_full - # The retained estimator is fitted on the dataset's HKL list *as loaded*; - # joining a collection expands the dataset onto the common grid, so it can be - # the wrong length by then. Rebuild against the current list when that happens. - fw = data_light._FrenchWilson - if fw is None or len(fw.d_spacings) != len(I_corrected_full): - from torchref.base.french_wilson import FrenchWilson - - fw = FrenchWilson( - data_light.hkl, data_light.cell, data_light.spacegroup, verbose=0 - ) - F_corr_full, sig_F_corr_full = fw(I_corrected_full, sig_I_full) + F_corr_full, sig_F_corr_full, _ = french_wilson_auto( + I_corrected_full, + sig_I_full, + data_light.hkl, + data_light.resolution, + data_light.spacegroup or "P1", + ) def _np(t): return t[mask].detach().cpu().numpy() diff --git a/torchref/io/datasets/base.py b/torchref/io/datasets/base.py index 5935fac3..da55081d 100644 --- a/torchref/io/datasets/base.py +++ b/torchref/io/datasets/base.py @@ -121,13 +121,13 @@ def _get_state(self) -> Dict[str, Any]: """Return observation fields and masks with tensors on CPU. Cell/device/space group are flattened to tensors or strings. Loading - provenance and French-Wilson conversion caches are omitted. + provenance is omitted. """ state = {} for f in fields(self): - if f.name in {"source", "reader", "dataset", "_FrenchWilson"}: - # Loading provenance and conversion caches are not observation state. + if f.name in {"source", "reader", "dataset"}: + # Loading provenance is not observation state. state[f.name] = None continue val = getattr(self, f.name) diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index 58722edc..cbb83f06 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -16,7 +16,7 @@ import torch from torchref.base import math_torch -from torchref.base.french_wilson import FrenchWilson +from torchref.base.french_wilson import french_wilson_auto from torchref.config import dtypes, normalize_device from torchref.io import cif, mtz from torchref.io.datasets.base import CrystalDataset @@ -209,7 +209,6 @@ class ReflectionData(CrystalDataset, DebugMixin): # Cached properties (not serialized) _centric: Optional[torch.Tensor] = field(default=None, repr=False) _n_bins: Optional[int] = field(default=None, repr=False) - _FrenchWilson: Optional[FrenchWilson] = field(default=None, repr=False) # Dynamic fields used by various methods source: Optional["ReflectionData"] = field(default=None, repr=False) @@ -796,12 +795,13 @@ def load(self, reader, french_wilson: bool = True): requires_grad=False, ) self.intensity_source = data_dict.get("I_col", "Unknown") - self._FrenchWilson = FrenchWilson( - self.hkl, self.cell.data, self.spacegroup, verbose=self.verbose + self.F, self.F_sigma, fw_keep = french_wilson_auto( + self.I, + self.I_sigma, + self.hkl, + self.resolution, + self.spacegroup or "P1", ) - F, F_sigma = self._FrenchWilson(self.I, self.I_sigma) - self.F = F - self.F_sigma = F_sigma # Record French-Wilson's own input criterion, evaluated on the true # intensities. This is strictly better than anything reconstructible # from the amplitudes afterwards: F is a positive posterior mean, so @@ -812,7 +812,7 @@ def load(self, reader, french_wilson: bool = True): # outlier mask -- this one guards the posterior integral against # unphysical input, which is a different question from whether an # observation is an outlier. - self._set_french_wilson_mask(self._FrenchWilson.valid_mask) + self._set_french_wilson_mask(fw_keep) elif "F" in data_dict: self.F = torch.tensor( data_dict["F"], From 7722a948abf4539b640c3a1b041ab49ef7525cb5 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 18:53:04 +0200 Subject: [PATCH 212/250] Add merge_to_spacegroup and drop ReflectionData.reduce_to_spacegroup merge_to_spacegroup expands every usable observation under the source space group, maps each copy onto the target's CCP4 ASU and merges them, counting each source observation at most once per target reflection so expansion copies never pose as independent measurements. It returns the merged dataset (French-Wilson amplitudes from the merged intensities, Bijvoet mates sharing one R-free flag) and MergeStats with per-shell Rmerge, Rmeas and CC_sym; with a higher target this is a symmetry test. from_tensors takes I, I_sigma and validation_flags so they are reordered with every other row during canonicalization. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 2 + tests/unit/io/test_merging.py | 225 +++++++++++ tests/unit/io/test_reflection_data_reindex.py | 15 +- torchref/io/__init__.py | 4 + torchref/io/datasets/__init__.py | 6 + torchref/io/datasets/merging.py | 373 ++++++++++++++++++ torchref/io/datasets/reflection_data.py | 259 +----------- 7 files changed, 640 insertions(+), 244 deletions(-) create mode 100644 tests/unit/io/test_merging.py create mode 100644 torchref/io/datasets/merging.py diff --git a/docs/changelog.rst b/docs/changelog.rst index d5aa5f85..1ce7c452 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,8 @@ Changelog Unreleased ---------- +- Added ``merge_to_spacegroup(data, spacegroup)``, which merges a dataset into another space group via P1 and returns the merged ``ReflectionData`` with per-shell Rmerge, Rmeas and CC_sym (``MergeStats``); a symmetry test when the target is higher than the source. It replaces ``ReflectionData.reduce_to_spacegroup``, which is removed +- ``ReflectionData.from_tensors`` accepts ``I``, ``I_sigma`` and ``validation_flags``, reordered with the other rows during canonicalization - Removed the ``FrenchWilson`` module and ``ReflectionData._FrenchWilson``; use ``french_wilson_auto(I, sigma_I, hkl, d_spacings, space_group)``, which returns ``(F, sigma_F, valid_mask)`` - Fixed ``torchref.difference-refine`` reusing the French-Wilson estimator built at load time, whose d-spacings and centric flags were in the pre-canonicalization row order; on files stored off the CCP4 ASU order (e.g. 6G9X) the corrected light amplitudes were computed against the wrong reflections - The difference MTZ groups its columns into named datasets -- ``observed``, ``difference``, ``light_model``, ``extrapolated_light``, ``two_moment`` -- with one history line describing each, so ``FWT``/``PHWT`` reads as ``/torchref/extrapolated_light/FWT`` (the extrapolated light-state map ``2*FEXT - Fc``). Labels are unchanged and Coot still auto-opens it diff --git a/tests/unit/io/test_merging.py b/tests/unit/io/test_merging.py new file mode 100644 index 00000000..23474d6b --- /dev/null +++ b/tests/unit/io/test_merging.py @@ -0,0 +1,225 @@ +"""merge_to_spacegroup: symmetry merging via P1 and its merging statistics. + +Pinned behaviour: expansion copies are never counted as independent +observations, row order of the input does not matter, masked or unusable +observations do not contribute, Bijvoet mates stay apart (and share one R-free +flag) only when asked, and the R values are ``None`` rather than zero when no +reflection has two observations. +""" + +import math + +import pytest +import torch + +from torchref.io import ReflectionData, merge_to_spacegroup +from torchref.symmetry import SpaceGroup + +CELL = (50.0, 60.0, 70.0, 90.0, 90.0, 90.0) +SG = "P 21 21 21" + + +def _asu_grid(n=6, sg=SG): + """Unique ASU indices of ``sg`` from a block of Miller indices, minus absences.""" + r = torch.arange(-n, n + 1) + grid = torch.stack(torch.meshgrid(r, r, r, indexing="ij"), dim=-1).reshape(-1, 3) + grid = grid[(grid != 0).any(dim=-1)].to(torch.int32) + group = SpaceGroup(sg) + canon, *_ = group.canonicalize_hkl(grid, include_friedel=True) + uniq = torch.unique(canon, dim=0) + return uniq[~group.is_absent(uniq)] + + +def _invariant(hkl): + """A value shared by all P212121 equivalents (they differ only in signs).""" + h, k, l = (hkl.to(torch.float32).abs().T) + return 10.0 + h + 2.0 * k + 3.0 * l + + +def _data(hkl, I, sigma, sg=SG, **kw): + return ReflectionData.from_tensors( + hkl, + I.clamp(min=0).sqrt(), + sigma / (2.0 * I.clamp(min=1e-3).sqrt()), + CELL, + sg, + verbose=0, + device="cpu", + I=I, + I_sigma=sigma, + **kw, + ) + + +def _synthetic(n=6): + hkl = _asu_grid(n) + I = _invariant(hkl) + return _data(hkl, I, torch.ones_like(I)) + + +def _by_hkl(data, values): + return {tuple(h): float(v) for h, v in zip(data.hkl.tolist(), values)} + + +@pytest.mark.unit +def test_p1_round_trip_recovers_values(): + d = _synthetic() + p1 = d.expand_to_p1() + p1.verbose = 0 + merged, stats = merge_to_spacegroup(p1, SG) + + merged._assert_per_reflection_consistent() + assert len(merged.hkl) == len(d.hkl) + orig = _by_hkl(d, d.I) + for h, v in _by_hkl(merged, merged.I).items(): + assert v == pytest.approx(orig[h], rel=1e-5) + assert stats.overall.r_merge == pytest.approx(0.0, abs=1e-6) + # Every P1 copy counts once: at most n_ops * 2 (Friedel) copies per reflection. + assert 1.0 < stats.overall.multiplicity <= 8.0 + + +@pytest.mark.unit +def test_noise_gives_the_expected_rmeas(): + """For constant I and Gaussian noise, Rmeas -> (sigma / I) * sqrt(2 / pi).""" + hkl = _asu_grid(8) + I0, sigma = 100.0, 10.0 + p1 = _data(hkl, torch.full((len(hkl),), I0), torch.full((len(hkl),), sigma)) + p1 = p1.expand_to_p1() + p1.verbose = 0 + g = torch.Generator().manual_seed(0) + p1.I = p1.I + sigma * torch.randn(len(p1.I), generator=g) + + _, stats = merge_to_spacegroup(p1, SG) + + expected = sigma / I0 * math.sqrt(2.0 / math.pi) + assert stats.overall.r_meas == pytest.approx(expected, rel=0.05) + assert stats.overall.r_merge < stats.overall.r_meas + + +@pytest.mark.unit +def test_same_group_has_no_agreement_to_measure(): + d = _synthetic() + merged, stats = merge_to_spacegroup(d, SG) + assert len(merged.hkl) == len(d.hkl) + assert stats.overall.multiplicity == 1.0 + assert stats.overall.r_merge is None + assert stats.overall.cc_sym is None + + +@pytest.mark.unit +def test_lowering_symmetry_generates_the_missing_reflections(): + d = _synthetic() + merged, stats = merge_to_spacegroup(d, "P 1 21 1") + assert len(merged.hkl) > len(d.hkl) + assert stats.overall.multiplicity == 1.0 + orig = _invariant(merged.hkl) + torch.testing.assert_close(merged.I, orig) + + +@pytest.mark.unit +def test_incompatible_cell_is_refused(): + d = _synthetic() + with pytest.raises(ValueError, match="incompatible"): + merge_to_spacegroup(d, "P 4 21 2") + + +@pytest.mark.unit +def test_row_order_does_not_matter(): + p1 = _synthetic().expand_to_p1() + p1.verbose = 0 + g = torch.Generator().manual_seed(1) + p1.I = p1.I + torch.randn(len(p1.I), generator=g) + shuffled = p1.__select__(torch.randperm(len(p1.hkl), generator=g)) + + a, sa = merge_to_spacegroup(p1, SG) + b, sb = merge_to_spacegroup(shuffled, SG) + + assert torch.equal(a.hkl, b.hkl) + torch.testing.assert_close(a.I, b.I) + assert sa.overall.r_merge == pytest.approx(sb.overall.r_merge) + + +@pytest.mark.unit +def test_same_seed_gives_identical_statistics(): + p1 = _synthetic().expand_to_p1() + p1.verbose = 0 + p1.I = p1.I + torch.randn(len(p1.I), generator=torch.Generator().manual_seed(2)) + assert merge_to_spacegroup(p1, SG, seed=3)[1] == merge_to_spacegroup(p1, SG, seed=3)[1] + + +@pytest.mark.unit +def test_masked_and_unusable_observations_do_not_contribute(): + p1 = _synthetic().expand_to_p1() + p1.verbose = 0 + canon, _, _, order = SpaceGroup(SG).canonicalize_hkl(p1.hkl, include_friedel=True) + per_row = torch.empty_like(canon) + per_row[order] = canon + first = per_row[0] + members = (per_row == first).all(dim=-1) + # Every copy of one reflection masked; a NaN sigma and a huge outlier + # elsewhere, the outlier masked too. + other = torch.nonzero(~members).squeeze(-1) + p1.I_sigma[other[0]] = float("nan") + p1.I[other[1]] = 1e6 + keep = ~members + keep[other[1]] = False + p1.masks["test"] = keep + + merged, stats = merge_to_spacegroup(p1, SG) + + assert tuple(first.tolist()) not in _by_hkl(merged, merged.I) + assert float(merged.I.max()) < 1e3 + assert stats.overall.r_merge == pytest.approx(0.0, abs=1e-6) + + +@pytest.mark.unit +def test_bijvoet_mates_stay_apart_and_share_rfree(): + hkl = _asu_grid() + acentric = ~SpaceGroup(SG).is_centric(hkl) + hkl = hkl[acentric] + both = torch.cat([hkl, -hkl]) + I = torch.cat([_invariant(hkl), _invariant(hkl) * 1.1]) + rfree = torch.ones(len(both), dtype=torch.bool) + rfree[: len(hkl) // 10] = False # free on the "+" member only + d = _data(both, I, torch.ones_like(I), rfree_flags=rfree, friedel_merged=False) + assert d.friedel_merged is False + + merged, _ = merge_to_spacegroup(d, SG) + assert len(merged.hkl) == len(both) + assert int(merged.friedel_flags.sum()) == len(hkl) + gid, n = merged.asu_group_indices() + flags = merged.rfree_flags.to(torch.bool) + split = ReflectionData._group_any(flags, gid, n) & ReflectionData._group_any( + ~flags, gid, n + ) + assert not bool(split.any()) + assert int((~flags).sum()) == 2 * (len(hkl) // 10) + + pooled, stats = merge_to_spacegroup(d, SG, anomalous=False) + assert len(pooled.hkl) == len(hkl) + assert stats.overall.r_merge > 0 + + +@pytest.mark.unit +def test_amplitude_only_data_merge_on_f_squared(mtz_dir): + d = ReflectionData(verbose=0, device="cpu").load_mtz(str(mtz_dir / "3E98.mtz")) + assert d.I is None + p1 = d.expand_to_p1() + p1.verbose = 0 + _, stats = merge_to_spacegroup(p1, d.spacegroup) + assert stats.on == "F^2" + assert stats.overall.r_merge == pytest.approx(0.0, abs=1e-5) + + +@pytest.mark.unit +def test_from_tensors_keeps_intensities_and_validation_row_aligned(): + """Canonicalization reorders rows; I and validation flags must move with them.""" + hkl = _asu_grid() + # Off-ASU equivalents in reverse order force a non-trivial permutation. + raw = torch.flip(hkl * torch.tensor([-1, 1, -1], dtype=hkl.dtype), dims=[0]) + I = _invariant(raw) + validation = (I.round().to(torch.int64) % 2) == 0 + d = _data(raw, I, torch.ones_like(I), validation_flags=validation) + torch.testing.assert_close(d.I, _invariant(d.hkl)) + expected = (_invariant(d.hkl).round().to(torch.int64) % 2) == 0 + assert torch.equal(d.validation_flags, expected) diff --git a/tests/unit/io/test_reflection_data_reindex.py b/tests/unit/io/test_reflection_data_reindex.py index 8b8e5058..b398904e 100644 --- a/tests/unit/io/test_reflection_data_reindex.py +++ b/tests/unit/io/test_reflection_data_reindex.py @@ -1,6 +1,6 @@ """Regression tests for per-reflection field reindexing. -``validate_hkl`` / ``remap`` / ``reduce_to_spacegroup`` must carry EVERY +``validate_hkl`` / ``remap`` / ``merge_to_spacegroup`` must carry EVERY per-reflection field onto the new HKL grid, not a hand-maintained subset. The historical bug left ``hkl_anomalous`` (read by ``_hkl_for_sf``) at the pre-alignment length, which crashed difference refinement whenever the dark and @@ -75,8 +75,8 @@ def test_identical_hkl_preserves_count(self): class TestP1RoundTripReindex: - """The same class of bug lived latently in remap/expand_to_p1 and - reduce_to_spacegroup (silent data loss rather than a crash).""" + """remap/expand_to_p1 and merge_to_spacegroup keep every per-reflection + field at the new length.""" def test_expand_to_p1_carries_validation_flags(self): grid = _base_grid(6, 6, 6) @@ -90,11 +90,16 @@ def test_expand_to_p1_carries_validation_flags(self): assert p1.validation_flags.shape[0] == len(p1.hkl) p1._assert_per_reflection_consistent() - def test_reduce_to_spacegroup_consistent(self): + def test_merge_to_spacegroup_consistent(self): + from torchref.io import merge_to_spacegroup + grid = _base_grid(6, 6, 6) d = _synthetic(grid, seed=4) + d.generate_validation_set(val_fraction_of_free=0.5, seed=0) p1 = d.expand_to_p1() - back = p1.reduce_to_spacegroup("P 21 21 21") + p1.verbose = 0 + back, _ = merge_to_spacegroup(p1, "P 21 21 21") + assert back.validation_flags is not None back._assert_per_reflection_consistent() diff --git a/torchref/io/__init__.py b/torchref/io/__init__.py index 79c69f9c..a82a82b0 100644 --- a/torchref/io/__init__.py +++ b/torchref/io/__init__.py @@ -26,8 +26,10 @@ CrystalDataset, DatasetCollection, FcalcDataset, + MergeStats, ReflectionData, ScaledDataset, + merge_to_spacegroup, ) # IHM ensemble support (mapping always available; reader/writer need python-ihm) @@ -50,6 +52,8 @@ "ScaledDataset", "DatasetCollection", "FcalcDataset", + "merge_to_spacegroup", + "MergeStats", # Top-level readers "read_mtz", "read_cif", diff --git a/torchref/io/datasets/__init__.py b/torchref/io/datasets/__init__.py index 65e4d9da..e2736af9 100644 --- a/torchref/io/datasets/__init__.py +++ b/torchref/io/datasets/__init__.py @@ -6,11 +6,14 @@ - :class:`ScaledDataset` -- live observations corrected by a shared DatasetScaler - :class:`FcalcDataset` -- calculated structure factors on a generated HKL set - :class:`DatasetCollection` -- several ReflectionData on one common HKL grid +- :func:`merge_to_spacegroup` -- merge a dataset into another space group, with + Rmerge / Rmeas / CC_sym as :class:`MergeStats` """ from .base import CrystalDataset from .collection import DatasetCollection from .fcalc_data import FcalcDataset +from .merging import MergeShell, MergeStats, merge_to_spacegroup from .reflection_data import ReflectionData from .scaled_dataset import ScaledDataset @@ -20,4 +23,7 @@ "ScaledDataset", "FcalcDataset", "DatasetCollection", + "merge_to_spacegroup", + "MergeStats", + "MergeShell", ] diff --git a/torchref/io/datasets/merging.py b/torchref/io/datasets/merging.py new file mode 100644 index 00000000..84676874 --- /dev/null +++ b/torchref/io/datasets/merging.py @@ -0,0 +1,373 @@ +""" +Merge a reflection dataset into another space group and report how well it merges. + +:func:`merge_to_spacegroup` takes any :class:`ReflectionData`, generates every +source-symmetry equivalent of every usable observation (the route through P1), +maps each onto the target group's CCP4 asymmetric unit, and merges what lands +together. The accompanying :class:`MergeStats` answers whether the target +symmetry is real: when the target is higher than the source, reflections that +were independent measurements become symmetry mates, and their agreement +(Rmerge, Rmeas, CC_sym) is the evidence. When the target is the same or lower, +every merged reflection has a single observation and the R values are ``None``. +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import List, Optional, Tuple + +import gemmi +import torch + +from torchref.base.french_wilson import french_wilson_auto +from torchref.base.reciprocal.hkl import get_d_spacing +from torchref.io.datasets.reflection_data import ReflectionData +from torchref.symmetry import SpaceGroup, SpaceGroupLike + +__all__ = ["merge_to_spacegroup", "MergeStats", "MergeShell"] + + +@dataclass +class MergeShell: + """Merging statistics for one resolution shell (or overall). + + ``r_merge``, ``r_meas`` and ``cc_sym`` are ``None`` when no reflection in + the shell has two or more observations, since agreement is then undefined. + """ + + d_max: float + d_min: float + n_unique: int + n_obs: int + r_merge: Optional[float] + r_meas: Optional[float] + cc_sym: Optional[float] + + @property + def multiplicity(self) -> float: + """Mean number of observations per merged reflection.""" + return self.n_obs / self.n_unique if self.n_unique else 0.0 + + +@dataclass +class MergeStats: + """Result of :func:`merge_to_spacegroup`; ``str()`` gives a table. + + Attributes + ---------- + source, target : str + Space-group symbols merged from and into. + on : str + ``"I"`` when the intensities were merged, ``"F^2"`` when only + amplitudes were available. R values on ``F^2`` are not comparable with + intensity R values from data processing. + anomalous : bool + Whether Bijvoet mates were kept apart. + overall : MergeShell + Statistics over all reflections. + shells : list of MergeShell + Per resolution shell, low to high resolution. + n_absent_obs : int + Source observations that land on reflections systematically absent in + the target; they are dropped from the merge. + absent_mean_i_over_sigma : float or None + Their mean I/sigma(I). Clearly above zero is evidence against the + target's screw axes or centring. + """ + + source: str + target: str + on: str + anomalous: bool + overall: MergeShell + shells: List[MergeShell] = field(default_factory=list) + n_absent_obs: int = 0 + absent_mean_i_over_sigma: Optional[float] = None + + def __str__(self) -> str: + def fmt(v, spec): + return format(v, spec) if v is not None else "-".rjust(len(format(0.0, spec))) + + head = ( + f"Merge {self.source} -> {self.target} on {self.on}" + f"{' (anomalous)' if self.anomalous else ''}\n" + f"{'d_max':>7} {'d_min':>6} {'n_uniq':>8} {'n_obs':>8} {'mult':>5} " + f"{'Rmerge':>7} {'Rmeas':>7} {'CC_sym':>7}" + ) + rows = [] + for s in self.shells + [self.overall]: + rows.append( + f"{s.d_max:7.2f} {s.d_min:6.2f} {s.n_unique:8d} {s.n_obs:8d} " + f"{s.multiplicity:5.2f} {fmt(s.r_merge, '7.3f')} " + f"{fmt(s.r_meas, '7.3f')} {fmt(s.cc_sym, '7.3f')}" + ) + rows.insert(len(self.shells), "-" * 62) + tail = "" + if self.n_absent_obs: + tail = ( + f"\n{self.n_absent_obs} observations on reflections absent in " + f"{self.target}, mean I/sigma {self.absent_mean_i_over_sigma:.2f}" + ) + return "\n".join([head, *rows]) + tail + + +def merge_to_spacegroup( + data: ReflectionData, + spacegroup: SpaceGroupLike, + *, + anomalous: Optional[bool] = None, + n_bins: int = 10, + seed: int = 0, +) -> Tuple[ReflectionData, MergeStats]: + """Merge ``data`` into ``spacegroup`` and measure the agreement of symmetry mates. + + Every usable source observation is expanded under the source space group, + each copy is mapped onto the target's CCP4 asymmetric unit, and each source + observation contributes at most once to any target reflection -- expansion + copies are never counted as independent measurements. Merged intensities + are inverse-variance weighted means; amplitudes are then derived by + French-Wilson, exactly as on load. + + Parameters + ---------- + data : ReflectionData + Source dataset. Only rows passing ``data.masks()`` with finite + observations and finite, positive sigmas are used. + spacegroup : SpaceGroupLike + Target space group, in the same cell and setting as ``data``. + anomalous : bool, optional + Keep Bijvoet mates apart (acentric reflections only). Defaults to + ``not data.friedel_merged``. + n_bins : int, optional + Number of resolution shells in the statistics, equal in unique + reflections. Default 10. + seed : int, optional + Seed for the random half split behind ``cc_sym``, drawn from a private + generator so the result is deterministic and the global RNG untouched. + + Returns + ------- + merged : ReflectionData + One row per merged reflection (two for a separated Bijvoet pair), with + ``I``/``I_sigma`` and French-Wilson ``F``/``F_sigma``. Reflections + French-Wilson rejects are masked. A merged reflection is free if any of + its observations was free, and Bijvoet mates share one flag. Phases and + figures of merit are not carried. + stats : MergeStats + Merging statistics; printed when ``data.verbose > 0``. + + Raises + ------ + ValueError + If the cell is incompatible with the target's lattice metric -- merging + then produces R values that look like evidence against the symmetry -- + or if no observation is usable. + + Notes + ----- + Without intensities the merge runs on ``F^2`` with ``sigma = 2 F sigma_F``, + and the statistics say so. The computation runs on CPU; the result is + moved to ``data.device``. + """ + # Fresh CPU copies: .to() on the dataset's own objects would move them in place. + target = SpaceGroup(spacegroup, device="cpu") + source = SpaceGroup(data.spacegroup, device="cpu") + cell = data.cell.data.detach().cpu() + if not gemmi.UnitCell(*cell.tolist()).is_compatible_with_spacegroup(target.gemmi): + raise ValueError( + f"Cell {[round(c, 3) for c in cell.tolist()]} is incompatible with " + f"{target.hm}; merging would measure the metric mismatch, not the " + "symmetry." + ) + if anomalous is None: + anomalous = not data.friedel_merged + + if data.I is not None and data.I_sigma is not None: + obs, sig, on = data.I, data.I_sigma, "I" + else: + obs, sig, on = data.F**2, 2.0 * data.F * data.F_sigma, "F^2" + obs, sig = obs.detach().cpu(), sig.detach().cpu() + use = data.masks().cpu() & torch.isfinite(obs) & torch.isfinite(sig) & (sig > 0) + rows = torch.nonzero(use).squeeze(-1) + if rows.numel() == 0: + raise ValueError("No usable observations to merge.") + obs, sig = obs[rows], sig[rows] + + # The signed index keeps a Bijvoet row on its own side of reciprocal space. + hkl_src = data.hkl_anomalous if anomalous and data.hkl_anomalous is not None else data.hkl + hkl_src = hkl_src.detach().cpu()[rows] + n_src = len(rows) + + rotated = source.reciprocal.apply_rotations(hkl_src) + cand = torch.round(rotated).to(hkl_src.dtype).reshape(-1, 3) + cand_src = torch.arange(n_src).repeat(source.n_ops) + + canon, _, friedel, sort_idx = target.canonicalize_hkl(cand, include_friedel=True) + # canonicalize_hkl returns its outputs sorted; the source row must follow. + cand_src = cand_src[sort_idx] + + absent = target.is_absent(canon) + n_absent_obs, absent_isig = 0, None + if bool(absent.any()): + absent_src = torch.unique(cand_src[absent]) + n_absent_obs = int(absent_src.numel()) + absent_isig = float((obs[absent_src] / sig[absent_src]).mean()) + canon, friedel, cand_src = canon[~absent], friedel[~absent], cand_src[~absent] + + # Centric Bijvoet mates are the same reflection, so only acentric ones split. + if anomalous: + side = (friedel & ~target.is_centric(canon)).to(canon.dtype) + key = torch.cat([canon, side.unsqueeze(-1)], dim=-1) + else: + side = torch.zeros(len(canon), dtype=canon.dtype) + key = canon + _, merge_id = torch.unique(key, dim=0, return_inverse=True) + n_merge = int(merge_id.max()) + 1 + + # One contribution per (merged reflection, source observation). + pair = merge_id * n_src + cand_src + order = torch.argsort(pair, stable=True) + first = torch.ones(len(order), dtype=torch.bool) + first[1:] = pair[order][1:] != pair[order][:-1] + sel = order[first] + gid, src = merge_id[sel], cand_src[sel] + + g_hkl = torch.empty((n_merge, 3), dtype=canon.dtype) + g_hkl[gid] = canon[sel] + g_side = torch.empty(n_merge, dtype=canon.dtype) + g_side[gid] = side[sel] + _, g_asu = torch.unique(g_hkl, dim=0, return_inverse=True) + n_asu = int(g_asu.max()) + 1 + + I_o, s_o = obs[src], sig[src] + w = s_o.pow(-2) + sum_w = torch.zeros(n_merge, dtype=w.dtype).index_add_(0, gid, w) + I_m = torch.zeros_like(sum_w).index_add_(0, gid, w * I_o) / sum_w + s_m = sum_w.rsqrt() + n_g = torch.zeros_like(sum_w).index_add_(0, gid, torch.ones_like(w)) + + d_g = get_d_spacing(g_hkl, cell.to(w.dtype)).to(w.dtype) + stats = _merge_stats( + I_o, gid, n_g, d_g, n_bins, torch.Generator().manual_seed(seed) + ) + stats.source, stats.target, stats.on, stats.anomalous = ( + source.hm, + target.hm, + on, + bool(anomalous), + ) + stats.n_absent_obs, stats.absent_mean_i_over_sigma = n_absent_obs, absent_isig + + F_m, sF_m, keep = french_wilson_auto(I_m, s_m, g_hkl, d_g, target) + F_m = torch.where(keep, F_m, torch.full_like(F_m, float("nan"))) + + def _asu_any(flags: torch.Tensor) -> torch.Tensor: + per_obs = flags.detach().cpu().to(torch.bool)[rows][src] + hit = ReflectionData._group_any(per_obs, g_asu[gid], n_asu) + return hit[g_asu] + + rfree = None + if data.rfree_flags is not None: + rfree = ~_asu_any(~data.rfree_flags.detach().cpu().to(torch.bool)) + validation = None + if data.validation_flags is not None: + validation = _asu_any(data.validation_flags) + + # Signed indices, so canonicalization inside from_tensors rebuilds + # friedel_flags / hkl_anomalous for a separated Bijvoet pair. + hkl_out = torch.where(g_side.bool().unsqueeze(-1), -g_hkl, g_hkl) + merged = ReflectionData.from_tensors( + hkl_out.to(data.hkl.dtype), + F_m, + sF_m, + data.cell.clone(), + SpaceGroup(target, device=data.device), + rfree_flags=rfree, + device=data.device, + verbose=data.verbose, + friedel_merged=not anomalous, + I=I_m, + I_sigma=s_m, + validation_flags=validation, + ) + merged.source = data + merged.last_op = f"merge_to_spacegroup({target.hm})" + if data.verbose > 0: + print(stats) + return merged, stats + + +def _merge_stats( + I_o: torch.Tensor, + gid: torch.Tensor, + n_g: torch.Tensor, + d_g: torch.Tensor, + n_bins: int, + gen: torch.Generator, +) -> MergeStats: + """Rmerge / Rmeas / CC_sym overall and per equal-count resolution shell. + + R values use the unweighted mean of each reflection's observations, as the + conventional definitions do; only reflections with two or more + observations contribute. + """ + n_merge = len(n_g) + mean_u = torch.zeros_like(n_g).index_add_(0, gid, I_o) / n_g + absdev = (I_o - mean_u[gid]).abs() + multi_o = n_g[gid] >= 2 + meas_w = torch.where(multi_o, (n_g / (n_g - 1).clamp(min=1)).sqrt()[gid], 0.0) + + # Random half split within each reflection: sort by (reflection, random key) + # and send the first floor(n/2) members of each to half A. + rnd = torch.rand(len(gid), generator=gen, dtype=I_o.dtype) + order = torch.argsort(rnd) + order = order[torch.argsort(gid[order], stable=True)] + g_sorted = gid[order] + start = torch.searchsorted(g_sorted, torch.arange(n_merge)) + pos = torch.arange(len(gid)) - start[g_sorted] + in_a = torch.zeros(len(gid), dtype=torch.bool) + in_a[order] = pos < (n_g[g_sorted] // 2) + n_a = torch.zeros_like(n_g).index_add_(0, gid, in_a.to(I_o.dtype)) + half_a = torch.zeros_like(n_g).index_add_(0, gid, torch.where(in_a, I_o, 0.0)) + half_b = torch.zeros_like(n_g).index_add_(0, gid, torch.where(in_a, 0.0, I_o)) + half_a = half_a / n_a.clamp(min=1) + half_b = half_b / (n_g - n_a).clamp(min=1) + + # Shells equal in unique reflections, low resolution first. + rank = torch.empty(n_merge, dtype=gid.dtype) + rank[torch.argsort(d_g, descending=True)] = torch.arange(n_merge) + shell = (rank * n_bins) // max(n_merge, 1) + + def summarise(sel_g: torch.Tensor) -> MergeShell: + sel_o = sel_g[gid] + multi_g = sel_g & (n_g >= 2) + m_o = sel_o & multi_o + denom = float(I_o[m_o].sum()) + r_merge = float(absdev[m_o].sum()) / denom if bool(m_o.any()) and denom else None + r_meas = ( + float((meas_w * absdev)[m_o].sum()) / denom + if r_merge is not None + else None + ) + cc = None + if int(multi_g.sum()) >= 3: + a = half_a[multi_g] - half_a[multi_g].mean() + b = half_b[multi_g] - half_b[multi_g].mean() + norm = float(a.norm() * b.norm()) + cc = float((a * b).sum()) / norm if norm > 0 else None + d_sel = d_g[sel_g] + return MergeShell( + d_max=float(d_sel.max()) if len(d_sel) else 0.0, + d_min=float(d_sel.min()) if len(d_sel) else 0.0, + n_unique=int(sel_g.sum()), + n_obs=int(sel_o.sum()), + r_merge=r_merge, + r_meas=r_meas, + cc_sym=cc, + ) + + shells = [summarise(shell == b) for b in range(n_bins) if bool((shell == b).any())] + overall = summarise(torch.ones(n_merge, dtype=torch.bool)) + return MergeStats( + source="", target="", on="", anomalous=False, overall=overall, shells=shells + ) diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index cbb83f06..b820a5dd 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -908,10 +908,17 @@ def from_tensors( verbose: int = 1, friedel_merged: Optional[bool] = None, detach: bool = True, + I: Optional[torch.Tensor] = None, + I_sigma: Optional[torch.Tensor] = None, + validation_flags: Optional[torch.Tensor] = None, ) -> "ReflectionData": """ Construct ReflectionData directly from tensors. + Every per-reflection tensor must be row-aligned with ``hkl`` as passed. + Canonicalization then reorders all of them together, so pass them here + rather than assigning them to the returned object. + Parameters ---------- hkl : torch.Tensor @@ -941,6 +948,11 @@ def from_tensors( True (default) stores constant observations, dropping the caller's autograd graph; False keeps it so gradients reach whatever produced ``F``/``F_sigma``. + I, I_sigma : torch.Tensor, optional + Intensities and their uncertainties of shape (N,). Stored as given; + ``F`` is not derived from them. + validation_flags : torch.Tensor, optional + Boolean validation-set flags of shape (N,). Returns ------- @@ -978,6 +990,14 @@ def _prep(t: torch.Tensor) -> torch.Tensor: data.rfree_flags = _prep(rfree_flags).to( device=data.device, dtype=torch.bool ) + if I is not None: + data.I = _prep(I).to(device=data.device) + if I_sigma is not None: + data.I_sigma = _prep(I_sigma).to(device=data.device) + if validation_flags is not None: + data.validation_flags = _prep(validation_flags).to( + device=data.device, dtype=torch.bool + ) # Set before canonicalization, which is what detects real Bijvoet mates # and downgrades this to False -- the same seam load() goes through. @@ -3125,245 +3145,6 @@ def expand_to_p1( op_name=f"expand_to_p1(include_friedel={include_friedel})", ) - def reduce_to_spacegroup( - self, spacegroup, include_friedel: bool = True, aggregation: str = "mean" - ) -> "ReflectionData": - """ - Reduce P1 reflection data to asymmetric unit of a target spacegroup. - - This is the inverse of expand_to_p1(). Takes reflection data in P1 and - merges symmetry-equivalent reflections into single ASU reflections using - the specified aggregation function. - - Parameters - ---------- - spacegroup : str, int, or gemmi.SpaceGroup - Target space group specification. - include_friedel : bool, default True - If True, also merge Friedel mates when reducing. - aggregation : str, default 'mean' - Aggregation function for merging equivalent reflections: - - 'mean': Average values (default, good for amplitudes) - - 'sum': Sum values - - 'first': Take first valid value (no averaging) - - Returns - ------- - ReflectionData - New ReflectionData with merged reflections in the target spacegroup. - - Notes - ----- - Per-field handling: F/I by ``aggregation``; sigmas propagated as - ``sqrt(sum(sigma²))/n`` for ``'mean'`` and ``sqrt(sum(sigma²))`` for - ``'sum'`` (first valid value for ``'first'``); ``phase`` by - amplitude-weighted complex averaging (so it wraps correctly) with - ``fom`` from the resultant - length; ``rfree_flags`` free if any equivalent is free, and - ``validation_flags`` set if any equivalent is set. Any other - per-reflection field takes its first valid equivalent. - """ - from torchref.symmetry.spacegroup import SpaceGroup - - if self.hkl is None: - raise ValueError("ReflectionData has no Miller indices loaded") - - # Get reduction mapping - hkl_asu, reduction_indices, phase_shifts = SpaceGroup( - spacegroup, device=self.device - ).reduce_hkl(self.hkl, include_friedel=include_friedel, device=self.device) - - n_asu = len(hkl_asu) - n_equiv = reduction_indices.shape[1] - valid_mask = reduction_indices >= 0 # (n_asu, n_equiv) - count_valid = valid_mask.sum(dim=1).clamp(min=1).float() # (n_asu,) - - # Helper function for aggregating 1D tensors - def _aggregate_tensor(tensor, agg_func="mean", fill_value=0.0): - if tensor is None: - return None - - # Gather values: (n_asu, n_equiv) - # Use clamp(min=0) to avoid indexing errors, then mask invalid - gathered = tensor[reduction_indices.clamp(min=0)] - gathered = torch.where(valid_mask, gathered, torch.zeros_like(gathered)) - - if agg_func == "mean": - return gathered.sum(dim=1) / count_valid - elif agg_func == "sum": - return gathered.sum(dim=1) - elif agg_func == "first": - # Take first valid value - first_valid_idx = valid_mask.to(dtype=dtypes.int).argmax(dim=1) - return gathered[ - torch.arange(n_asu, device=self.device), first_valid_idx - ] - else: - raise ValueError(f"Unknown aggregation: {agg_func}") - - def _aggregate_sigma(tensor, agg_func="mean"): - """Propagate uncertainty correctly for averaging.""" - if tensor is None: - return None - - # Gather values - gathered = tensor[reduction_indices.clamp(min=0)] - gathered = torch.where(valid_mask, gathered, torch.zeros_like(gathered)) - - if agg_func == "mean": - # For averaging: sigma_mean = sqrt(sum(sigma^2)) / n - variance_sum = (gathered**2).sum(dim=1) - return torch.sqrt(variance_sum) / count_valid - elif agg_func == "sum": - # For summing: sigma_sum = sqrt(sum(sigma^2)) - variance_sum = (gathered**2).sum(dim=1) - return torch.sqrt(variance_sum) - elif agg_func == "first": - first_valid_idx = valid_mask.to(dtype=dtypes.int).argmax(dim=1) - return gathered[ - torch.arange(n_asu, device=self.device), first_valid_idx - ] - else: - raise ValueError(f"Unknown aggregation: {agg_func}") - - # Create new ReflectionData - reduced = ReflectionData(verbose=self.verbose, device=self.device) - - # Set HKL - reduced.hkl = hkl_asu.to(device=self.device) - - # Aggregate amplitude and intensity fields - reduced.F = _aggregate_tensor(self.F, aggregation) - reduced.F_sigma = _aggregate_sigma(self.F_sigma, aggregation) - reduced.I = _aggregate_tensor(self.I, aggregation) - reduced.I_sigma = _aggregate_sigma(self.I_sigma, aggregation) - - # Handle phases via complex averaging - if self.phase is not None: - # Gather phases and apply phase shifts for proper averaging - phases_gathered = self.phase[reduction_indices.clamp(min=0)] - phases_gathered = phases_gathered + phase_shifts - phases_gathered = torch.where( - valid_mask, phases_gathered, torch.zeros_like(phases_gathered) - ) - - # Get weights (amplitudes or FOM) - if self.fom is not None: - weights = self.fom[reduction_indices.clamp(min=0)] - elif self.F is not None: - weights = self.F[reduction_indices.clamp(min=0)] - else: - weights = torch.ones_like(phases_gathered) - weights = torch.where(valid_mask, weights, torch.zeros_like(weights)) - - # Complex averaging: mean of F*exp(i*phi) then extract angle - complex_sf = weights * torch.exp(1j * phases_gathered) - complex_mean = complex_sf.sum(dim=1) / count_valid - reduced.phase = torch.angle(complex_mean).float() - - # FOM as magnitude of normalized mean complex vector - if self.fom is not None: - norm_weights = weights / weights.sum(dim=1, keepdim=True).clamp( - min=1e-10 - ) - unit_vectors = torch.exp(1j * phases_gathered) - mean_vector = (norm_weights * unit_vectors).sum(dim=1) - reduced.fom = torch.abs(mean_vector).float() - else: - reduced.phase = None - reduced.fom = ( - _aggregate_tensor(self.fom, aggregation) - if self.fom is not None - else None - ) - - # Handle rfree_flags: OR operation (free if any equivalent is free) - if self.rfree_flags is not None: - rfree_gathered = self.rfree_flags[reduction_indices.clamp(min=0)].to( - dtypes.int - ) - rfree_gathered = torch.where( - valid_mask, - rfree_gathered, - torch.ones_like(rfree_gathered), # Default to work set - ) - # 0 = free, non-zero = work. Take min to get free if any is free. - reduced.rfree_flags = rfree_gathered.min(dim=1).values != 0 - - # Boolean per-reflection flags: an equivalent's flag propagates to the - # merged reflection if ANY contributor has it set (validation is a - # conservative "exclude if any"). - def _aggregate_any(tensor): - if tensor is None: - return None - gathered = tensor[reduction_indices.clamp(min=0)].to(torch.bool) - gathered = gathered & valid_mask - return gathered.any(dim=1) - - if self.validation_flags is not None: - reduced.validation_flags = _aggregate_any(self.validation_flags) - - # Completeness pass: carry any remaining per-reflection dataclass tensor - # field not handled above so the merge never silently drops data. - # Derived-from-HKL fields are recomputed/invalidated below, not - # aggregated; 'first' is a safe representative for the rest. - from dataclasses import fields as dc_fields - - _already_set = { - "hkl", - "F", - "F_sigma", - "I", - "I_sigma", - "phase", - "fom", - "rfree_flags", - "validation_flags", - } - _recomputed = set(self._REINDEX_DERIVED) | {"hkl_anomalous", "friedel_flags"} - n_src = len(self.hkl) - for f in dc_fields(self): - name = f.name - if name in _already_set or name in _recomputed: - continue - val = getattr(self, name) - if not isinstance(val, torch.Tensor): - continue - if not (val.shape and val.shape[0] == n_src): - continue - setattr(reduced, name, _aggregate_tensor(val, "first")) - - # Clone cell - reduced.cell = self.cell.clone() if self.cell is not None else None - - # Set spacegroup (on the reduced dataset's device, not the global default) - reduced.spacegroup = SpaceGroup(spacegroup, device=reduced.device) - - # Recalculate resolution - if reduced.cell is not None and reduced.hkl is not None: - reduced._calculate_resolution() - - # Invalidate derived-from-HKL fields (recomputed lazily for the new ASU). - # hkl_anomalous / friedel_flags are left unset: the merged ASU is - # Friedel-merged, so _hkl_for_sf() correctly falls back to hkl. - reduced.bin_indices = None - reduced._centric_flags = None - - # Copy metadata sources - reduced.amplitude_source = self.amplitude_source - reduced.intensity_source = self.intensity_source - reduced.phase_source = self.phase_source - reduced.rfree_source = self.rfree_source - - # Track provenance - reduced.source = self - reduced.last_op = ( - f"reduce_to_spacegroup({spacegroup}, aggregation={aggregation})" - ) - - reduced._assert_per_reflection_consistent() - return reduced - def canonicalize(self, include_friedel: bool = True) -> "ReflectionData": """Return new ReflectionData with HKL in standard CCP4 ASU form. From e0a97af1547b7b0cf37a2a5e2af3021590bffe5c Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 18:56:43 +0200 Subject: [PATCH 213/250] Remove unused ReflectionData methods and fields Fourteen methods with no caller in the package, tests, docs or notebooks, and the dataset/reader/_centric/_expansion_phase_shifts fields that were never read. Checkpoints that still carry the removed fields load with the existing stale-key warning. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 1 + torchref/io/datasets/base.py | 2 +- torchref/io/datasets/reflection_data.py | 418 +----------------------- 3 files changed, 8 insertions(+), 413 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 1ce7c452..70b5019c 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Removed unused ``ReflectionData`` methods: ``data_fill_masked``, ``mean_F_per_bin``, ``mean_sigma_per_bin``, ``calc_patterson``, ``fill``, ``possible_hkl``, ``get_structure_factors``, ``get_structure_factors_with_sigma``, ``get_hkl``, ``list_cif_data_blocks``, ``dump``, ``check_all_data_types``, ``unpack_one`` and ``get_min_res``, and the never-populated ``dataset`` and ``reader`` fields - Added ``merge_to_spacegroup(data, spacegroup)``, which merges a dataset into another space group via P1 and returns the merged ``ReflectionData`` with per-shell Rmerge, Rmeas and CC_sym (``MergeStats``); a symmetry test when the target is higher than the source. It replaces ``ReflectionData.reduce_to_spacegroup``, which is removed - ``ReflectionData.from_tensors`` accepts ``I``, ``I_sigma`` and ``validation_flags``, reordered with the other rows during canonicalization - Removed the ``FrenchWilson`` module and ``ReflectionData._FrenchWilson``; use ``french_wilson_auto(I, sigma_I, hkl, d_spacings, space_group)``, which returns ``(F, sigma_F, valid_mask)`` diff --git a/torchref/io/datasets/base.py b/torchref/io/datasets/base.py index da55081d..46854278 100644 --- a/torchref/io/datasets/base.py +++ b/torchref/io/datasets/base.py @@ -126,7 +126,7 @@ def _get_state(self) -> Dict[str, Any]: state = {} for f in fields(self): - if f.name in {"source", "reader", "dataset"}: + if f.name == "source": # Loading provenance is not observation state. state[f.name] = None continue diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index b820a5dd..911e2903 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -9,7 +9,7 @@ import warnings from dataclasses import dataclass, field from pathlib import Path -from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Optional, Tuple, Union import numpy as np import pandas as pd @@ -207,14 +207,11 @@ class ReflectionData(CrystalDataset, DebugMixin): # Note: Most fields are inherited from CrystalDataset dataclass # Cached properties (not serialized) - _centric: Optional[torch.Tensor] = field(default=None, repr=False) _n_bins: Optional[int] = field(default=None, repr=False) - # Dynamic fields used by various methods + # Provenance: the dataset this one was derived from, and the operation. source: Optional["ReflectionData"] = field(default=None, repr=False) - dataset: Optional[pd.DataFrame] = field(default=None, repr=False) last_op: Optional[str] = field(default=None, repr=False) - reader: Optional[Any] = field(default=None, repr=False) def __post_init__(self): """ @@ -526,9 +523,6 @@ def _canonicalize_in_place(self) -> None: dict.__setitem__(self.masks, name, mask_tensor[sort_indices]) self.masks._updated = True - if hasattr(self, "dataset") and self.dataset is not None: - self.dataset = self.dataset.iloc[sort_indices.cpu().numpy()].copy() - # friedel_flags comes back already in sorted (canonical) order, matching # self.hkl. hkl_anomalous carries the SIGNED index used for # structure-factor evaluation -- canonical for the (+) member, negated @@ -1111,31 +1105,10 @@ def load_cif( ReflectionData Self, for method chaining. """ - self.reader = cif.ReflectionCIFReader( + reader = cif.ReflectionCIFReader( str(path), verbose=self.verbose, data_block=data_block, anomalous=anomalous ) - return self.load(self.reader) - - @staticmethod - def list_cif_data_blocks(path: str) -> List[str]: - """ - List all data blocks available in a CIF file without loading data. - - Useful for multi-dataset CIF files to inspect available blocks - before loading a specific one. - - Parameters - ---------- - path : str - Path to CIF file. - - Returns - ------- - list of str - Names of all data blocks in the CIF file, in file order; pass one to - ``load_cif(data_block=...)``. - """ - return cif.list_data_blocks(path) + return self.load(reader) def _generate_rfree_flags( self, @@ -1389,74 +1362,6 @@ def mean_res_per_bin(self) -> torch.Tensor: mean_resolutions = mean_resolutions / count_per_bin.clamp(min=1).float() return mean_resolutions - def mean_F_per_bin(self) -> torch.Tensor: - """ - Calculate mean structure factor amplitude per resolution bin. - - Returns - ------- - torch.Tensor - Mean F per bin of shape (n_bins,). - - Raises - ------ - ValueError - If bins have not been created yet. - """ - if self.bin_indices is None: - self.get_bins() - if self.F is None: - raise ValueError("No amplitude data loaded") - - mean_F = torch.zeros(self._n_bins, dtype=dtypes.float, device=self.device) - count_per_bin = torch.zeros(self._n_bins, dtype=dtypes.int, device=self.device) - mask = self.masks() - mean_F = torch.scatter_add( - mean_F, 0, self.bin_indices[mask].to(torch.int64), self.F[mask] # dtype-ok: bin indices for scatter_add index arg; PyTorch requires int64 - ) - count_per_bin = torch.scatter_add( - count_per_bin, - 0, - self.bin_indices[mask].to(torch.int64), # dtype-ok: bin indices for scatter_add index arg; PyTorch requires int64 - torch.ones_like(self.F[mask], dtype=dtypes.int), - ) - mean_F = mean_F / count_per_bin.clamp(min=1).float() - return mean_F - - def mean_sigma_per_bin(self) -> Optional[torch.Tensor]: - """ - Calculate mean structure factor uncertainty per resolution bin. - - Returns - ------- - torch.Tensor or None - Mean sigma_F per bin of shape (n_bins,), or None if no uncertainties. - - Raises - ------ - ValueError - If bins have not been created yet. - """ - if self.bin_indices is None: - self.get_bins() - if self.F is None: - raise ValueError("No amplitude data loaded") - - mean_sigma = torch.zeros(self._n_bins, dtype=dtypes.float, device=self.device) - count_per_bin = torch.zeros(self._n_bins, dtype=dtypes.int, device=self.device) - mask = self.masks() - mean_sigma = torch.scatter_add( - mean_sigma, 0, self.bin_indices[mask].to(torch.int64), self.F_sigma[mask] # dtype-ok: bin indices for scatter_add index arg; PyTorch requires int64 - ) - count_per_bin = torch.scatter_add( - count_per_bin, - 0, - self.bin_indices[mask].to(torch.int64), # dtype-ok: bin indices for scatter_add index arg; PyTorch requires int64 - torch.ones_like(self.F_sigma[mask], dtype=dtypes.int), - ) - mean_sigma = mean_sigma / count_per_bin.clamp(min=1).float() - return mean_sigma - def regenerate_rfree_flags( self, free_fraction: float = 0.02, @@ -1754,75 +1659,6 @@ def _fit_two_component_wilson( return B_struct.item(), B_sol.item(), k.item() - def get_structure_factors(self, as_complex: bool = False) -> torch.Tensor: - """ - Get structure factors, optionally as complex numbers. - - Parameters - ---------- - as_complex : bool, optional - If True and phases available, return F*exp(i*phi). Default is False. - - Returns - ------- - torch.Tensor - Structure factor amplitudes or complex structure factors. - - Raises - ------ - ValueError - If no amplitude data is loaded. - """ - if self.F is None: - raise ValueError("No amplitude data loaded") - - if as_complex and self.phase is not None: - return self.F * torch.exp(1j * self.phase) - else: - return self.F - - def get_structure_factors_with_sigma( - self, - ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: - """ - Get structure factor amplitudes and their uncertainties. - - Returns - ------- - F : torch.Tensor - Structure factor amplitudes of shape (N,). - F_sigma : torch.Tensor or None - Uncertainties of shape (N,), or None if not available. - - Raises - ------ - ValueError - If no amplitude data is loaded. - """ - if self.F is None: - raise ValueError("No amplitude data loaded") - - return self.F, self.F_sigma - - def get_hkl(self): - """ - Return Miller indices for valid reflections. - - Returns - ------- - torch.Tensor - Miller indices of the valid subset, shape (M, 3) with M <= N - (M = number of valid reflections), dtype int32. - - Raises - ------ - ValueError - If no Miller indices are loaded. - """ - if self.hkl is None: - raise ValueError("No Miller indices loaded") - return self.hkl[self.masks()] - def filter_by_resolution( self, d_min: Optional[float] = None, d_max: Optional[float] = None ) -> "ReflectionData": @@ -1894,13 +1730,6 @@ def get_max_res(self) -> Optional[float]: mask = self.masks() return float(self.resolution[mask].min().item()) - def get_min_res(self) -> Optional[float]: - """Largest d-spacing among valid reflections, in Ångströms.""" - if self.resolution is None: - self._calculate_resolution() - mask = self.masks() - return float(self.resolution[mask].max().item()) - def __len__(self) -> int: """Number of reflections (full array, ignoring masks).""" return len(self.hkl) if self.hkl is not None else 0 @@ -1971,65 +1800,6 @@ def data_indexed( return hkl, F, F_sigma, rfree_flags - def data_fill_masked( - self, mode="mean" - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - """ - Return data tensors with missing or flagged reflections filled in. - - Parameters - ---------- - mode : str, optional - Fill strategy for missing/flagged reflections. Default is 'mean'. - - - 'mean' : fill with the per-bin mean of the present data. - - 'zero' : fill with zero. - - Returns - ------- - hkl : torch.Tensor - Miller indices of shape (N, 3). - F : torch.Tensor - Structure factor amplitudes of shape (N,) with gaps filled. - F_sigma : torch.Tensor - Amplitude uncertainties of shape (N,) with gaps filled. - rfree : torch.Tensor - R-free flags of shape (N,); filled-in reflections are assigned to - the work set (True). - """ - hkl, F, F_sigma = self.hkl, self.F, self.F_sigma - if F is None or F_sigma is None: - raise ValueError("Amplitude observations and uncertainties are required") - mask = self.masks() - if not bool(mask.any()): - raise ValueError("No valid reflections to fill from") - rfree = ( - self.rfree_flags.clone() - if self.rfree_flags is not None - else torch.ones_like(mask) - ) - - if mode == "mean": - mean_F = self.mean_F_per_bin() - mean_F_sigma = self.mean_sigma_per_bin() - F_data = F.clone() - F_sigma_data = F_sigma.clone() - F_data[~mask] = mean_F[self.bin_indices[~mask]] - F_sigma_data[~mask] = mean_F_sigma[self.bin_indices[~mask]] - rfree[~mask] = True # set missing to work set - return hkl, F_data, F_sigma_data, rfree - - elif mode == "zero": - F_data = F.clone() - F_sigma_data = F_sigma.clone() - F_data[~mask] = 0.0 - F_sigma_data[~mask] = 0.0 - rfree[~mask] = True # set missing to work set - return hkl, F_data, F_sigma_data, rfree - - else: - raise ValueError(f"Unknown fill mode: {mode}") - def __getitem__(self, key): """ Index into the reflection dataset. @@ -2100,11 +1870,6 @@ def __select__(self, indices: torch.Tensor, op=None) -> "ReflectionData": new_masks[name] = mask_tensor[indices] selected.masks = new_masks - # Handle DataFrame - if hasattr(self, "dataset") and self.dataset is not None: - idx_np = indices.cpu().numpy() - selected.dataset = self.dataset.iloc[idx_np].copy() - selected.source = self selected.last_op = op return selected @@ -2161,20 +1926,6 @@ def sanitize_F(self): self.F_sigma[mask] = 0.0 return self - def check_all_data_types(self): - """Print dtype/shape (or type/value) of every attribute, for debugging.""" - for key in self.__dict__: - if self.__dict__[key] is not None and isinstance( - self.__dict__[key], torch.Tensor - ): - print( - f"{key}: {self.__dict__[key].dtype}, shape: {self.__dict__[key].shape}" - ) - elif self.__dict__[key] is not None: - print(f"{key}: {type(self.__dict__[key])}, value: {self.__dict__[key]}") - else: - print(f"{key}: None") - def validate_hkl( self, hkl_ref: torch.Tensor, *, identity_hkl: Optional[torch.Tensor] = None ) -> "ReflectionData": @@ -2274,21 +2025,6 @@ def validate_hkl( self._assert_per_reflection_consistent() return self - def unpack_one(self): - """ - Unpack one level of source. - - Does not recurse fully and does not flag. - - Returns - ------- - ReflectionData - Parent source or self if no source. - """ - if self.source is not None: - return self.source - return self - WILSON_MASK_KEY = "wilson_valid" FRENCH_WILSON_MASK_KEY = "french_wilson_valid" @@ -2465,22 +2201,6 @@ def flag_suspicious_sigma(self, z_threshold: float = 5.0) -> None: ) self.masks["flagged_sigma"] = ~flagged - def dump(self): - """ - Dump all reflection data to console for debugging. - - Prints type, shape, and device information for all attributes. - """ - print("ReflectionData dump:") - for key in self.__dict__: - value = self.__dict__[key] - if isinstance(value, torch.Tensor): - print( - f" {key}: dtype={value.dtype}, shape={value.shape}, device={value.device}" - ) - else: - print(f" {key}: type={type(value)}, value={value}") - def _build_anomalous_dataframe( self, fcalc: Optional[torch.Tensor] = None ) -> pd.DataFrame: @@ -2874,80 +2594,6 @@ def centric(self): return self._centric_flags - def calc_patterson( - self, - grid_size: Optional[Tuple[int, int, int]] = None, - grid_sampling: Optional[float] = 1, - ) -> torch.Tensor: - """ - Calculate Patterson map of the dataset. - - The Patterson function P(u,v,w) = Σ|F(hkl)|² exp(-2πi(hu+kv+lw)) - is computed via inverse FFT of F². Data is expanded to P1 symmetry - using only observed reflections (no filling of missing data). - - Parameters - ---------- - grid_size : tuple of int, optional - Grid dimensions (Nx, Ny, Nz). If None, automatically determined - from unit cell and resolution. - grid_sampling : float, optional - Sampling interval for the grid. Default is 1. - This sets the grid so that we sample twice as much as normal for a given resolution - - Returns - ------- - torch.Tensor - Real-valued Patterson map of shape (Nx, Ny, Nz). - Origin is at grid position [0, 0, 0]. - """ - from torchref.base.fourier import find_grid_size - from torchref.base.reciprocal import place_on_grid - - # Expand to P1 symmetry (don't fill missing reflections - use only observed data) - data = self.expand_to_p1() - - max_res = data.resolution.min() * grid_sampling - - if grid_size is None: - grid_size = find_grid_size(data.cell, max_res) - - # Use data_indexed to get only valid (observed) reflections - hkl, F, _, _ = data.data_indexed() - - F_2 = F**2 - - # Place F² on reciprocal grid (don't enforce Hermitian since we have P1 expansion) - grid = place_on_grid(hkl, F_2, grid_size, enforce_hermitian=False) - - patterson = torch.fft.ifftn(grid, dim=(0, 1, 2), norm="forward").real - - return patterson - - def possible_hkl(self) -> torch.Tensor: - """ - Generate all possible HKL indices within the resolution limit. - - Returns - ------- - torch.Tensor - Tensor of shape (M, 3) containing all possible Miller indices - within the resolution limit defined by self.resolution. - """ - from torchref.base.reciprocal import generate_possible_hkl - - if self.cell is None or self.resolution is None: - raise ValueError( - "Cell and resolution must be defined to generate possible HKL" - ) - - max_res = self.resolution.min().item() - possible_hkl = generate_possible_hkl( - self.cell.data, max_res, device=self.device - ) - - return possible_hkl - def remap( self, new_hkl: torch.Tensor, @@ -3017,12 +2663,8 @@ def _remap_mask(tensor, fill_value): self._reindex_per_reflection(index_mapping, new_hkl, target=remapped) # Apply optional phase shifts (e.g. from symmetry translations). - if remapped.phase is not None: - if phase_shifts is not None: - remapped.phase = remapped.phase + phase_shifts.to(device=self.device) - elif phase_shifts is not None: - # No original phases: store the shifts for later phase reconstruction. - remapped._expansion_phase_shifts = phase_shifts.to(device=self.device) + if remapped.phase is not None and phase_shifts is not None: + remapped.phase = remapped.phase + phase_shifts.to(device=self.device) # Carry forward prior combined mask if available. prior_mask = self.masks() @@ -3049,54 +2691,6 @@ def _remap_mask(tensor, fill_value): remapped._assert_per_reflection_consistent() return remapped - def fill(self, d_min: Optional[float] = None) -> "ReflectionData": - """ - Fill missing reflections within resolution limit. - - Generates all possible reflections for the current spacegroup within - the resolution limit, identifies which are missing, and creates a - complete dataset. Missing reflections are filled with default values. - - Parameters - ---------- - d_min : float, optional - High resolution limit in Angstroms. If None, uses the minimum - resolution from the current dataset. - - Returns - ------- - ReflectionData - New ReflectionData with complete set of reflections. - Missing reflections have F/I/phase/fom = 0.0, - F_sigma/I_sigma = 1.0 and ``masks['missing'] = True``. - """ - if self.hkl is None: - raise ValueError("ReflectionData has no Miller indices loaded") - if self.cell is None: - raise ValueError("ReflectionData has no unit cell defined") - - # Use current resolution limit if not specified - if d_min is None: - if self.resolution is None: - raise ValueError("Resolution not available - specify d_min") - d_min = self.resolution.min().item() - - # Get complete HKL set with index mapping - sg = self.spacegroup or SpaceGroup("P1", device=self.device) - filled_hkl, indices, missing = sg.complete_hkl( - self.hkl, self.cell.data, d_min, device=self.device - ) - - # Use remap to create the new dataset - remapped = self.remap( - new_hkl=filled_hkl, - index_mapping=indices, - spacegroup=self.spacegroup, # Keep same spacegroup - op_name=f"fill(d_min={d_min:.2f})", - ) - - return remapped - def expand_to_p1( self, include_friedel: bool = True, remove_absences: bool = True ) -> "ReflectionData": From 037bd40f35aa1c733b7ee3deb023204251df4622 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 18:57:49 +0200 Subject: [PATCH 214/250] Remove ReflectionData.canonicalize and flag_suspicious_sigma canonicalize duplicated the in-place canonicalization load already runs, and flag_suspicious_sigma was a diagnostic superseded by the Wilson outlier mask. Neither had a caller. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 2 +- tests/unit/io/test_wilson_outlier_masks.py | 11 --- torchref/io/datasets/reflection_data.py | 93 ---------------------- 3 files changed, 1 insertion(+), 105 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 70b5019c..9c71d62c 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,7 +4,7 @@ Changelog Unreleased ---------- -- Removed unused ``ReflectionData`` methods: ``data_fill_masked``, ``mean_F_per_bin``, ``mean_sigma_per_bin``, ``calc_patterson``, ``fill``, ``possible_hkl``, ``get_structure_factors``, ``get_structure_factors_with_sigma``, ``get_hkl``, ``list_cif_data_blocks``, ``dump``, ``check_all_data_types``, ``unpack_one`` and ``get_min_res``, and the never-populated ``dataset`` and ``reader`` fields +- Removed unused ``ReflectionData`` methods: ``data_fill_masked``, ``mean_F_per_bin``, ``mean_sigma_per_bin``, ``calc_patterson``, ``fill``, ``possible_hkl``, ``get_structure_factors``, ``get_structure_factors_with_sigma``, ``get_hkl``, ``list_cif_data_blocks``, ``dump``, ``check_all_data_types``, ``unpack_one`` and ``get_min_res``, and the never-populated ``dataset`` and ``reader`` fields. Also removed ``ReflectionData.canonicalize`` (loading already canonicalizes in place; ``SpaceGroup.canonicalize_hkl`` remains) and ``flag_suspicious_sigma`` (superseded by the Wilson outlier mask) - Added ``merge_to_spacegroup(data, spacegroup)``, which merges a dataset into another space group via P1 and returns the merged ``ReflectionData`` with per-shell Rmerge, Rmeas and CC_sym (``MergeStats``); a symmetry test when the target is higher than the source. It replaces ``ReflectionData.reduce_to_spacegroup``, which is removed - ``ReflectionData.from_tensors`` accepts ``I``, ``I_sigma`` and ``validation_flags``, reordered with the other rows during canonicalization - Removed the ``FrenchWilson`` module and ``ReflectionData._FrenchWilson``; use ``french_wilson_auto(I, sigma_I, hkl, d_spacings, space_group)``, which returns ``(F, sigma_F, valid_mask)`` diff --git a/tests/unit/io/test_wilson_outlier_masks.py b/tests/unit/io/test_wilson_outlier_masks.py index 6372046f..0086da15 100644 --- a/tests/unit/io/test_wilson_outlier_masks.py +++ b/tests/unit/io/test_wilson_outlier_masks.py @@ -274,17 +274,6 @@ def test_french_wilson_guard_refuses_an_all_false_mask(): data._set_french_wilson_mask(torch.zeros(len(data.hkl), dtype=torch.bool)) -@pytest.mark.unit -def test_suspicious_sigma_is_no_longer_run_at_load(): - hkl, F, F_sigma = _wilson_grid(half_width=8) - data = _synthetic(F, F_sigma, hkl=hkl) - assert "flagged_sigma" not in data.masks - - # Still available for diagnostics, and still writes its own key. - data.flag_suspicious_sigma() - assert "flagged_sigma" in data.masks - - @pytest.mark.unit def test_too_few_reflections_are_left_alone(): """Wilson statistics cannot be estimated from a handful of reflections, and diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index 911e2903..5629dc87 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -2163,44 +2163,6 @@ def flag_wilson_outliers( device=self.device, dtype=torch.bool ) - def flag_suspicious_sigma(self, z_threshold: float = 5.0) -> None: - """ - Flag sigma values that deviate significantly from expected distribution. - - Sigma values from a detector should follow a log-normal distribution. - Values with z-scores beyond threshold are flagged as suspicious. - - .. note:: - No longer run during loading -- :meth:`flag_wilson_outliers` - supersedes it. The z-score here is taken against a *global* mean and - std of ``log sigma``, but that distribution is a mixture across - resolution shells (sigma tracks the intensity fall-off), so the - global std is inflated by the resolution trend and the test is - correspondingly blunt. It also never looks at ``F`` beside its sigma. - Kept for diagnostics and backwards compatibility. - - Parameters - ---------- - z_threshold : float, optional - Z-score threshold, on ``log(sigma)``, for flagging a sigma as suspicious. - Default is 5.0. - """ - sigmas = self.F_sigma - log_sigmas = torch.log(sigmas) - flagged_initial = torch.isnan(log_sigmas) | torch.isinf(log_sigmas) - mean_log_sigma = torch.mean(log_sigmas[~flagged_initial]) - std_log_sigma = torch.std(log_sigmas[~flagged_initial]) + 1e-5 * mean_log_sigma - z_scores = (log_sigmas - mean_log_sigma) / std_log_sigma - flagged = torch.abs(z_scores) > z_threshold - flagged = flagged | flagged_initial - if self.verbose > 0: - n_flagged = flagged.sum().item() - n_total = len(sigmas) - print( - f"Suspicious sigma detection: {n_flagged}/{n_total} ({100*n_flagged/n_total:.2f}%) reflections flagged" - ) - self.masks["flagged_sigma"] = ~flagged - def _build_anomalous_dataframe( self, fcalc: Optional[torch.Tensor] = None ) -> pd.DataFrame: @@ -2739,61 +2701,6 @@ def expand_to_p1( op_name=f"expand_to_p1(include_friedel={include_friedel})", ) - def canonicalize(self, include_friedel: bool = True) -> "ReflectionData": - """Return new ReflectionData with HKL in standard CCP4 ASU form. - - Remaps all Miller indices to the canonical CCP4 asymmetric unit - representative using ``gemmi.ReciprocalAsu``, adjusts phases - accordingly, and sorts reflections lexicographically by (h, k, l). - - Parameters - ---------- - include_friedel : bool, default True - Whether Friedel mates are considered equivalent. - - Returns - ------- - ReflectionData - New object with canonicalized, sorted Miller indices. - """ - if self.hkl is None: - raise ValueError("ReflectionData has no Miller indices loaded") - - sg = self.spacegroup or SpaceGroup("P1", device=self.device) - canonical_hkl, phase_shifts, friedel_flags, sort_indices = sg.canonicalize_hkl( - self.hkl, include_friedel, device=self.device - ) - - # Reorder all fields using __select__ - result = self.__select__( - sort_indices, op=f"canonicalize(include_friedel={include_friedel})" - ) - - # Overwrite HKL with canonical form (already sorted) - result.hkl = canonical_hkl - - # Fix phases: phi_new = where(friedel, -phi_old, phi_old) + phase_shift - if result.phase is not None: - result.phase = ( - torch.where(friedel_flags, -result.phase, result.phase) + phase_shifts - ) - - # Record anomalous bookkeeping for the canonical result (the stale values - # carried over by __select__ are recomputed here). See _hkl_for_sf. - result.friedel_flags = friedel_flags - result.hkl_anomalous = torch.where( - friedel_flags.unsqueeze(-1), -canonical_hkl, canonical_hkl - ) - - # Recalculate resolution from canonical HKL + cell - if result.cell is not None: - result._calculate_resolution() - - # Invalidate bin_indices - result.bin_indices = None - - return result - # ========== E-VALUE AND ANISOTROPY CORRECTION METHODS ========== def get_scattering_vectors(self) -> torch.Tensor: From 5cc1519a76ef0ef3647e1e40c4fdbc11ac24f9cc Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 19:07:36 +0200 Subject: [PATCH 215/250] Merge ReflectionData's duplicate accessors and R-free generation get_max_res, get_valid_mask and cut_res were aliases of d_min, masks() and filter_by_resolution; callers now use those. _generate_rfree_flags and regenerate_rfree_flags become one public generate_rfree_flags(force=False), and it shares one stratified ASU-group draw with generate_validation_set. Seeded draws are bit-identical to before. get_bins no longer stores bin_indices on the dataset, and mean_res_per_bin takes the bins it averages over: ScalerBase bins once at init, but read the per-bin mean resolution from whatever bins the dataset held last, which R-free or validation-set generation or a least-squares target could have replaced with different settings. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 3 + .../integration/test_rigid_body_isolation.py | 14 +- .../alignment/test_patterson_translation.py | 2 +- tests/unit/io/test_anomalous_reader.py | 4 +- tests/unit/io/test_rfree_generation.py | 13 +- tests/unit/refinement/test_wilson_prior.py | 4 +- torchref/cli/collection_difference_refine.py | 4 +- torchref/cli/simulate_noisy_data.py | 2 +- torchref/cli/validate_ded.py | 4 +- torchref/experimental/alignment/pipeline.py | 4 +- torchref/io/datasets/base.py | 1 - torchref/io/datasets/reflection_data.py | 319 +++++++----------- torchref/refinement/base_refinement.py | 6 +- torchref/refinement/rigid_body_refinement.py | 10 +- torchref/scaling/scaler_base.py | 2 +- 15 files changed, 152 insertions(+), 240 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 9c71d62c..be42b79f 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,9 @@ Changelog Unreleased ---------- +- ``ReflectionData.regenerate_rfree_flags`` is renamed ``generate_rfree_flags`` (same arguments; with existing flags and ``force=False`` it now warns instead of printing), and it prints only when ``verbose > 0``. Seeded draws are unchanged +- ``ReflectionData.get_bins`` no longer stores ``bin_indices`` on the dataset, and ``mean_res_per_bin`` takes the bins it should average over, so a later ``get_bins`` call with other settings (R-free or validation-set generation, a least-squares target) can no longer shift the shells a scaler's per-bin solvent scale was set up on. The ``bin_indices`` field is removed +- Removed the ``ReflectionData`` aliases ``get_max_res`` (use ``d_min``), ``get_valid_mask`` (use ``masks()``) and ``cut_res`` (use ``filter_by_resolution(d_min=, d_max=)``, which now prints only when ``verbose > 0``) - Removed unused ``ReflectionData`` methods: ``data_fill_masked``, ``mean_F_per_bin``, ``mean_sigma_per_bin``, ``calc_patterson``, ``fill``, ``possible_hkl``, ``get_structure_factors``, ``get_structure_factors_with_sigma``, ``get_hkl``, ``list_cif_data_blocks``, ``dump``, ``check_all_data_types``, ``unpack_one`` and ``get_min_res``, and the never-populated ``dataset`` and ``reader`` fields. Also removed ``ReflectionData.canonicalize`` (loading already canonicalizes in place; ``SpaceGroup.canonicalize_hkl`` remains) and ``flag_suspicious_sigma`` (superseded by the Wilson outlier mask) - Added ``merge_to_spacegroup(data, spacegroup)``, which merges a dataset into another space group via P1 and returns the merged ``ReflectionData`` with per-shell Rmerge, Rmeas and CC_sym (``MergeStats``); a symmetry test when the target is higher than the source. It replaces ``ReflectionData.reduce_to_spacegroup``, which is removed - ``ReflectionData.from_tensors`` accepts ``I``, ``I_sigma`` and ``validation_flags``, reordered with the other rows during canonicalization diff --git a/tests/integration/test_rigid_body_isolation.py b/tests/integration/test_rigid_body_isolation.py index 97ca1587..b5ef2b68 100644 --- a/tests/integration/test_rigid_body_isolation.py +++ b/tests/integration/test_rigid_body_isolation.py @@ -74,29 +74,29 @@ def test_targets_and_data_are_not_replaced(refinement): def test_resolution_range_survives_a_coarse_only_cutoff_list(refinement): """The caller's resolution range must be what it was, not the last cutoff's. - `cut_res` masks in place and returns `self`, so a `cutoffs` list ending above - the native d_min can leave the caller truncated. Object identity does not - catch it -- the data object is the same one throughout. + `filter_by_resolution` masks in place and returns `self`, so a `cutoffs` list + ending above the native d_min can leave the caller truncated. Object identity + does not catch it -- the data object is the same one throughout. """ ref = refinement() data = ref.reflection_data ref.get_scales() n_before = int(data.masks().sum()) n_work_before = int(data.work.mask.sum()) - d_min_before = data.get_max_res() + d_min_before = data.d_min # Deliberately coarse-only, and deliberately not ending at the native d_min. ref.refine_rigid_body(iterations_per_step=5, cutoffs=[6.0, 4.0]) assert int(data.masks().sum()) == n_before assert int(data.work.mask.sum()) == n_work_before - assert data.get_max_res() == pytest.approx(d_min_before) + assert data.d_min == pytest.approx(d_min_before) def test_a_caller_supplied_resolution_limit_is_not_widened(refinement): """A refinement built with `max_res` keeps that limit across a rigid-body run.""" ref = refinement() - ref.reflection_data.cut_res(highres=3.5) + ref.reflection_data.filter_by_resolution(d_min=3.5) data = ref.reflection_data ref.get_scales() n_before = int(data.masks().sum()) @@ -104,7 +104,7 @@ def test_a_caller_supplied_resolution_limit_is_not_widened(refinement): ref.refine_rigid_body(iterations_per_step=5) assert int(data.masks().sum()) == n_before - assert data.get_max_res() >= 3.5 + assert data.d_min >= 3.5 def test_refined_coordinates_still_reach_the_caller(refinement): diff --git a/tests/unit/alignment/test_patterson_translation.py b/tests/unit/alignment/test_patterson_translation.py index d5545ef2..2b64eaba 100644 --- a/tests/unit/alignment/test_patterson_translation.py +++ b/tests/unit/alignment/test_patterson_translation.py @@ -45,7 +45,7 @@ def setup(): rec = data.cell.reciprocal_basis_matrix.cpu().to(torch.float64) s = (data.hkl.cpu().to(torch.float64) @ rec).norm(dim=-1) window = ((s >= 1.0 / 15.0) & (s <= 1.0 / 4.0)).to(data.hkl.device) - mask = data.get_valid_mask() & window + mask = data.masks() & window return canonical, data, mask diff --git a/tests/unit/io/test_anomalous_reader.py b/tests/unit/io/test_anomalous_reader.py index d72e0f8e..09cd6edc 100644 --- a/tests/unit/io/test_anomalous_reader.py +++ b/tests/unit/io/test_anomalous_reader.py @@ -126,7 +126,7 @@ def test_generated_rfree_shared_across_mates(anomalous_two_column_mtz): assert d.friedel_merged is False assert bool(d.friedel_flags.any()) - d.regenerate_rfree_flags(force=True, seed=0) + d.generate_rfree_flags(force=True, seed=0) # The seed is part of the provenance: "generated" without it names a draw # nobody can reproduce. assert d.rfree_source == ( @@ -141,7 +141,7 @@ def test_generated_validation_set_shared_across_mates(anomalous_two_column_mtz): path, _ = anomalous_two_column_mtz d = ReflectionData(verbose=0) d.load_mtz(path) - d.regenerate_rfree_flags(force=True, seed=0) + d.generate_rfree_flags(force=True, seed=0) d.generate_validation_set(val_fraction_of_free=0.5, seed=0) assert bool(d.validation_flags.any()) diff --git a/tests/unit/io/test_rfree_generation.py b/tests/unit/io/test_rfree_generation.py index 997feba3..0c2c8a65 100644 --- a/tests/unit/io/test_rfree_generation.py +++ b/tests/unit/io/test_rfree_generation.py @@ -1,8 +1,7 @@ """Unit tests for resolution-stratified R-free flag generation. -Covers the binning contract of ``ReflectionData._generate_rfree_flags`` / -``regenerate_rfree_flags``: each resolution bin holds >= min_per_bin (1000) -reflections and contributes >= min_free_per_bin (50) free reflections, the +Covers the binning contract of ``ReflectionData.generate_rfree_flags``: each +resolution bin holds >= min_per_bin (1000) reflections and contributes >= min_free_per_bin (50) free reflections, the flags are binary, generation is reproducible under a seed, and tiny datasets degrade gracefully to a single clamped bin. """ @@ -18,7 +17,7 @@ def _synthetic_data(h=12, k=12, lmax=30, seed=0, device="cpu"): Friedel folding changes the count). (2h+1)(2k+1)*lmax reflections. Built with placeholder (all-work) flags, then flags are (re)generated via the - public ``regenerate_rfree_flags`` once the validity masks exist. + public ``generate_rfree_flags`` once the validity masks exist. """ hs = torch.arange(-h, h + 1) ks = torch.arange(-k, k + 1) @@ -35,7 +34,7 @@ def _synthetic_data(h=12, k=12, lmax=30, seed=0, device="cpu"): rfree_flags=torch.ones(n, dtype=torch.bool), device=device, verbose=0, friedel_merged=True, ) - data.regenerate_rfree_flags(force=True, seed=seed) + data.generate_rfree_flags(force=True, seed=seed) return data @@ -105,7 +104,7 @@ def test_default_generation_matches_min_free_floor(): def test_seed_reproducible(): data = _synthetic_data(seed=42) flags_a = data.rfree_flags.clone() - data.regenerate_rfree_flags(force=True, seed=42) + data.generate_rfree_flags(force=True, seed=42) assert torch.equal(flags_a, data.rfree_flags) @@ -186,7 +185,7 @@ def test_free_set_excludes_masked_reflections(): keep = torch.ones(len(data.hkl), dtype=torch.bool, device=data.device) keep[::3] = False data.masks["unusable"] = keep - data.regenerate_rfree_flags(force=True, seed=0) + data.generate_rfree_flags(force=True, seed=0) valid = data.masks().to(torch.bool) assert not bool(valid.all()), "expected some masked-out reflections" diff --git a/tests/unit/refinement/test_wilson_prior.py b/tests/unit/refinement/test_wilson_prior.py index 75635189..3a9c1e1f 100644 --- a/tests/unit/refinement/test_wilson_prior.py +++ b/tests/unit/refinement/test_wilson_prior.py @@ -28,11 +28,11 @@ def setup_target(): data._calculate_wilson_b() ens = EnsembleModel.from_single( TEST_PDB, n_members=4, perturb_sigma=0.0, b_const=5.0, - seed=42, verbose=0, max_res=data.get_max_res(), + seed=42, verbose=0, max_res=data.d_min, ) ens.cell = data.cell ens.spacegroup = data.spacegroup - ens.max_res = data.get_max_res() + ens.max_res = data.d_min scaler = Scaler(model=ens, data=data, nbins=10, verbose=0) fcalc0 = ens(data.hkl) diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index be365a99..6346305e 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -142,8 +142,8 @@ def setup_dataset_collection(sf_dark, sf_light, d_min, device, sf_light, device=device, column_names=column_names_light, ) if d_min is not None: - data_dark.cut_res(highres=d_min) - data_light.cut_res(highres=d_min) + data_dark.filter_by_resolution(d_min=d_min) + data_light.filter_by_resolution(d_min=d_min) dc = DatasetCollection(device=device) dc.add_dataset("dark", data_dark) diff --git a/torchref/cli/simulate_noisy_data.py b/torchref/cli/simulate_noisy_data.py index 7e37361d..d25f15d9 100644 --- a/torchref/cli/simulate_noisy_data.py +++ b/torchref/cli/simulate_noisy_data.py @@ -168,7 +168,7 @@ def _run_reference_mode(args, model, device) -> int: args.reference_hkl, cell=model.cell, spacegroup=model.spacegroup, ) # Prune reference tensors by resolution so the simulation output only - # covers the requested range. cut_res() only masks — we want the HKL + # covers the requested range. filter_by_resolution() only masks — we want the HKL # list itself to be filtered so Scaler, model(), and add_noise all see # the same reflection set. if args.d_min is not None or args.d_max is not None: diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index 18b6d9a7..20fa008c 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -251,8 +251,8 @@ def setup_ded_context( str(light_sf), device=device, column_names=col_light, verbose=0 ) if dmin is not None: - data_dark.cut_res(highres=dmin) - data_light.cut_res(highres=dmin) + data_dark.filter_by_resolution(d_min=dmin) + data_light.filter_by_resolution(d_min=dmin) collection = DatasetCollection(device=str(device)) collection.add_dataset("dark", data_dark) diff --git a/torchref/experimental/alignment/pipeline.py b/torchref/experimental/alignment/pipeline.py index 256df019..d23767f2 100644 --- a/torchref/experimental/alignment/pipeline.py +++ b/torchref/experimental/alignment/pipeline.py @@ -593,8 +593,8 @@ def _prepare_translation_arrays(self) -> None: device = self.device hkl_full = data.hkl F_obs_full = data.F - if hasattr(data, "get_valid_mask"): - tmask = data.get_valid_mask() + if getattr(data, "masks", None) is not None: + tmask = data.masks() else: tmask = torch.ones( F_obs_full.shape[0], dtype=torch.bool, device=F_obs_full.device, diff --git a/torchref/io/datasets/base.py b/torchref/io/datasets/base.py index 46854278..966122f8 100644 --- a/torchref/io/datasets/base.py +++ b/torchref/io/datasets/base.py @@ -51,7 +51,6 @@ class CrystalDataset(DeviceMovementMixin): # reflections are carved out of BOTH the work and free sets (disjoint). validation_flags: Optional[torch.Tensor] = None # (N,), bool resolution: Optional[torch.Tensor] = None # Resolution per reflection (N,) - bin_indices: Optional[torch.Tensor] = None # Resolution bin assignments (N,), int32 phase: Optional[torch.Tensor] = None # Phases in radians (N,) fom: Optional[torch.Tensor] = None # Figure of merit (N,) _centric_flags: Optional[torch.Tensor] = None # Centric flags (N,), bool diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index 5629dc87..ffa94802 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -9,7 +9,7 @@ import warnings from dataclasses import dataclass, field from pathlib import Path -from typing import TYPE_CHECKING, Optional, Tuple, Union +from typing import TYPE_CHECKING, Callable, Optional, Tuple, Union import numpy as np import pandas as pd @@ -206,9 +206,6 @@ class ReflectionData(CrystalDataset, DebugMixin): # Additional fields specific to ReflectionData (beyond CrystalDataset) # Note: Most fields are inherited from CrystalDataset dataclass - # Cached properties (not serialized) - _n_bins: Optional[int] = field(default=None, repr=False) - # Provenance: the dataset this one was derived from, and the operation. source: Optional["ReflectionData"] = field(default=None, repr=False) last_op: Optional[str] = field(default=None, repr=False) @@ -382,9 +379,9 @@ def I_sigma_raw(self) -> Optional[torch.Tensor]: # Per-reflection fields that are pure functions of (hkl, cell, spacegroup): # never gathered/aggregated, always recomputed or invalidated after an HKL - # change (``resolution`` recomputed; ``bin_indices`` / ``_centric_flags`` - # lazily rebuilt by ``get_bins`` / the ``centric`` property). - _REINDEX_DERIVED = ("resolution", "bin_indices", "_centric_flags") + # change (``resolution`` recomputed; ``_centric_flags`` lazily rebuilt by + # the ``centric`` property). + _REINDEX_DERIVED = ("resolution", "_centric_flags") def _reindex_per_reflection( self, @@ -458,7 +455,6 @@ def _reindex_per_reflection( # Install the new HKL and recompute / invalidate derived-from-HKL fields. target.hkl = new_hkl - target.bin_indices = None target._centric_flags = None if target.cell is not None: target._calculate_resolution() @@ -863,9 +859,9 @@ def load(self, reader, french_wilson: bool = True): # Generate only after canonicalization: the free set must be drawn on # unique ASU reflections, and the Bijvoet grouping that requires does # not exist until _canonicalize_in_place has run. See - # asu_group_indices / _generate_rfree_flags. + # asu_group_indices / generate_rfree_flags. if self.rfree_flags is None: - self._generate_rfree_flags() + self.generate_rfree_flags() return self @@ -1002,7 +998,7 @@ def _prep(t: torch.Tensor) -> torch.Tensor: # As in load(): generate after canonicalization so the draw can group # Bijvoet mates onto a shared canonical index. if data.rfree_flags is None: - data._generate_rfree_flags() + data.generate_rfree_flags() return data @@ -1110,18 +1106,20 @@ def load_cif( ) return self.load(reader) - def _generate_rfree_flags( + def generate_rfree_flags( self, free_fraction: float = 0.02, n_bins: int = 10, min_per_bin: int = 1000, min_free_per_bin: int = 50, seed: Optional[int] = None, + force: bool = False, ) -> None: """ Generate R-free flags with resolution-stratified sampling. Sets ``rfree_flags`` (int32, 1=work/0=free) and ``rfree_source``. + ``load`` and ``from_tensors`` call this when the input carries no flags. The draw is over *unique ASU reflections*, not rows: the two members of a Bijvoet pair share a canonical index (see :meth:`asu_group_indices`) @@ -1133,8 +1131,6 @@ def _generate_rfree_flags( Only reflections passing the validity masks are drawn from, so the counts below describe usable reflections rather than raw rows. - Must be called *after* canonicalization -- see :meth:`load`. - Parameters ---------- free_fraction : float, optional @@ -1150,7 +1146,11 @@ def _generate_rfree_flags( Minimum free unique reflections per bin, clamped to the number the bin holds. seed : int, optional - Random seed for reproducibility. Default is None. + Seeds the **global** torch and numpy RNGs before the draw, so the + same seed on the same data reproduces the same set. + force : bool, optional + Overwrite existing flags. Default False: existing flags are kept and + the call only warns. Raises ------ @@ -1160,38 +1160,35 @@ def _generate_rfree_flags( If the data have not been canonicalized (via :meth:`asu_group_indices`). """ + if self.rfree_flags is not None and not force: + warnings.warn( + f"R-free flags already exist ({self.rfree_source}); " + "pass force=True to overwrite them." + ) + return if self.resolution is None: raise ValueError("Resolution information required to generate R-free flags") + if self.verbose > 0: + if self.rfree_flags is not None: + print(f"Overwriting existing R-free flags ({self.rfree_source})") + print( + f"Generating R-free flags: {free_fraction*100:.1f}% free, " + f"{n_bins} bins of >= {min_per_bin}, >= {min_free_per_bin} free per bin" + ) - print("Generating R-free flags:") - print(f" Target free fraction: {free_fraction*100:.1f}%") - print(f" Target bins: {n_bins}") - print(f" Minimum per bin: {min_per_bin} reflections") - print(f" Minimum free per bin: {min_free_per_bin} reflections") - - # Set random seed for reproducibility if seed is not None: np.random.seed(seed) torch.manual_seed(seed) - n_refl = len(self.resolution) - - # Create resolution bins bin_indices, actual_n_bins = self.get_bins( n_bins=n_bins, min_per_bin=min_per_bin ) - - print(f" Created {actual_n_bins} resolution bins") - - # Draw on unique ASU reflections rather than rows, so Bijvoet mates - # (which share a canonical index) cannot be split across work/free. group_id, n_groups = self.asu_group_indices() # A group is eligible if any of its rows survives the validity masks; # spending the free quota on masked-out rows would silently shrink the # usable free set below min_free_per_bin. - valid = self.masks().to(torch.bool) - group_valid = self._group_any(valid, group_id, n_groups) + group_valid = self._group_any(self.masks().to(torch.bool), group_id, n_groups) if not bool(group_valid.any()): warnings.warn( "No reflections pass the validity masks; drawing R-free flags " @@ -1199,31 +1196,16 @@ def _generate_rfree_flags( ) group_valid = torch.ones_like(group_valid) - # One bin per group, from a representative row. group_bin = bin_indices[self._group_representative_rows(group_id, n_groups)] - - group_free = torch.zeros(n_groups, dtype=torch.bool, device=self.device) - for bin_idx in range(actual_n_bins): - eligible = torch.where((group_bin == bin_idx) & group_valid)[0] - n_bin_groups = int(eligible.numel()) - if n_bin_groups == 0: - continue - - # At least min_free_per_bin unique reflections, otherwise - # free_fraction of the bin; never more than the bin holds. - n_free_in_bin = min( - n_bin_groups, - max(min_free_per_bin, int(n_bin_groups * free_fraction)), - ) - perm = torch.randperm(n_bin_groups, device=eligible.device)[:n_free_in_bin] - group_free[eligible[perm]] = True - - # Broadcast each group's decision to every row sharing its ASU index. - flags = torch.ones( - n_refl, dtype=dtypes.int, device=self.device, requires_grad=False + group_free = self._stratified_group_draw( + group_valid, + group_bin, + actual_n_bins, + lambda n: min(n, max(min_free_per_bin, int(n * free_fraction))), ) - flags[group_free[group_id]] = 0 + flags = torch.ones(len(self.resolution), dtype=dtypes.int, device=self.device) + flags[group_free[group_id]] = 0 self.rfree_flags = flags # The seed belongs in the provenance string: without it "generated" # names a draw nobody can reproduce. @@ -1233,23 +1215,50 @@ def _generate_rfree_flags( + ")" ) - n_free = (flags == 0).sum().item() - n_work = (flags != 0).sum().item() - free_pct = 100.0 * n_free / n_refl + if self.verbose > 0: + n_free = int((flags == 0).sum()) + print( + f" {n_free} free ({100.0 * n_free / len(flags):.1f}%) in " + f"{actual_n_bins} bins, drawn over {int(group_valid.sum())} unique " + "ASU reflections; Bijvoet mates share a flag" + ) - print( - f" ✓ Generated flags: {n_free} free ({free_pct:.1f}%), {n_work} work ({100-free_pct:.1f}%)" - ) - print( - f" Drawn over {int(group_valid.sum())} unique ASU reflections " - f"({n_groups} groups total); Bijvoet mates share a flag" - ) + @staticmethod + def _stratified_group_draw( + eligible: torch.Tensor, + group_bin: torch.Tensor, + n_bins: int, + n_to_draw: Callable[[int], int], + ) -> torch.Tensor: + """Draw ``n_to_draw(n)`` of the ``n`` eligible groups in each resolution bin. + + Shared by R-free and validation-set generation so both split whole ASU + groups the same way. Uses the global torch RNG, one ``randperm`` per + non-empty bin in bin order, so a seeded caller is reproducible. + + Returns + ------- + torch.Tensor + Boolean mask of shape ``(n_groups,)``, True for drawn groups. + """ + drawn = torch.zeros_like(eligible, dtype=torch.bool) + for b in range(n_bins): + members = torch.where((group_bin == b) & eligible)[0] + n = int(members.numel()) + if n == 0: + continue + perm = torch.randperm(n, device=members.device)[: n_to_draw(n)] + drawn[members[perm]] = True + return drawn def get_bins( self, n_bins: int = 20, min_per_bin: int = 100 ) -> Tuple[torch.Tensor, int]: """ - Create resolution bins with approximately equal reflection counts. + Create resolution bins with approximately equal counts of valid reflections. + + Pure: nothing is stored on the dataset, so callers that need the same + bins later (e.g. :meth:`mean_res_per_bin`) must keep the returned tensor. Parameters ---------- @@ -1321,92 +1330,35 @@ def get_bins( ) if actual_n_bins > 20: print(f" ... ({actual_n_bins - 20} more bins)") - self.bin_indices = bin_indices - self._n_bins = actual_n_bins return bin_indices, actual_n_bins - def mean_res_per_bin(self) -> torch.Tensor: + def mean_res_per_bin(self, bin_indices: torch.Tensor, n_bins: int) -> torch.Tensor: """ - Calculate mean resolution for each bin. + Mean resolution of the valid reflections in each bin. + + Parameters + ---------- + bin_indices : torch.Tensor + Bin of each reflection, shape (N,), as returned by :meth:`get_bins`. + n_bins : int + Number of bins, as returned by :meth:`get_bins`. Returns ------- torch.Tensor - Mean resolution for each bin in Ångströms. - - Raises - ------ - ValueError - If bins have not been created yet. + Mean resolution per bin in Ångströms, shape (n_bins,); 0 for an + empty bin. """ - if self.bin_indices is None or self.resolution is None: - raise ValueError("Bins have not been created yet") - - mean_resolutions = torch.zeros( - self._n_bins, dtype=dtypes.float, device=self.device - ) - count_per_bin = torch.zeros(self._n_bins, dtype=dtypes.int, device=self.device) + if self.resolution is None: + self._calculate_resolution() mask = self.masks() - mean_resolutions = torch.scatter_add( - mean_resolutions, - 0, - self.bin_indices[mask].to(torch.int64), # dtype-ok: bin indices for scatter_add/index; PyTorch requires int64 - self.resolution[mask], - ) - count_per_bin = torch.scatter_add( - count_per_bin, - 0, - self.bin_indices[mask].to(torch.int64), # dtype-ok: bin indices for scatter_add/index; PyTorch requires int64 - torch.ones_like(self.resolution[mask], dtype=dtypes.int), - ) - mean_resolutions = mean_resolutions / count_per_bin.clamp(min=1).float() - return mean_resolutions - - def regenerate_rfree_flags( - self, - free_fraction: float = 0.02, - n_bins: int = 10, - min_per_bin: int = 1000, - min_free_per_bin: int = 50, - seed: Optional[int] = None, - force: bool = False, - ) -> None: - """ - Regenerate R-free flags with resolution-stratified sampling. - - Parameters - ---------- - free_fraction : float, optional - Fraction of reflections to mark as free. Default is 0.02 (2%). - n_bins : int, optional - Target number of resolution bins. Default is 10. - min_per_bin : int, optional - Minimum reflections per resolution bin. Default is 1000. - min_free_per_bin : int, optional - Minimum free reflections per resolution bin. Default is 50. - seed : int, optional - Random seed for reproducibility. Default is None. - force : bool, optional - If True, overwrite existing flags. Default False, in which case an - existing set is kept and the call is a no-op (warning only). - """ - if self.rfree_flags is not None and not force: - print("⚠️ WARNING: R-free flags already exist!") - print(f" Current source: {self.rfree_source}") - print(" Use force=True to overwrite existing flags") - return - - if self.rfree_flags is not None and force: - print("⚠️ WARNING: Overwriting existing R-free flags") - print(f" Old source: {self.rfree_source}") - - self._generate_rfree_flags( - free_fraction=free_fraction, - n_bins=n_bins, - min_per_bin=min_per_bin, - min_free_per_bin=min_free_per_bin, - seed=seed, + idx = bin_indices[mask].to(torch.int64) # dtype-ok: index_add_ requires int64 indices + res = self.resolution[mask] + total = torch.zeros(n_bins, dtype=res.dtype, device=res.device).index_add_( + 0, idx, res ) + count = torch.zeros_like(total).index_add_(0, idx, torch.ones_like(res)) + return total / count.clamp(min=1) def _calculate_resolution(self) -> None: """Set ``self.resolution`` to per-reflection d-spacing in Ångströms. @@ -1691,53 +1643,26 @@ def filter_by_resolution( self.masks["resolution"] = mask - valid = self.masks().sum().item() - print( - f"Filtering: {mask.sum()}/{len(mask)} reflections in range " - f"[{d_max if d_max else 'inf'} - {d_min if d_min else 'inf'}] " - f"\u00c5 ({valid} valid after all masks)" - ) + if self.verbose > 0: + valid = self.masks().sum().item() + print( + f"Filtering: {mask.sum()}/{len(mask)} reflections in range " + f"[{d_max if d_max else 'inf'} - {d_min if d_min else 'inf'}] " + f"\u00c5 ({valid} valid after all masks)" + ) return self - def cut_res( - self, highres: Optional[float] = None, lowres: Optional[float] = None - ) -> "ReflectionData": - """ - Filter reflections by resolution range (alias for filter_by_resolution). - - Masks rather than deletes: reflections outside the range stay in the - arrays but are excluded by ``masks()``. - - Parameters - ---------- - highres : float, optional - High-resolution cutoff (small d, e.g. 1.5 Å); keeps d >= highres. - lowres : float, optional - Low-resolution cutoff (large d, e.g. 50.0 Å); keeps d <= lowres. - - Returns - ------- - ReflectionData - Self, for method chaining. - """ - return self.filter_by_resolution(d_min=highres, d_max=lowres) - - def get_max_res(self) -> Optional[float]: - """Smallest d-spacing among valid reflections, in Ångströms.""" - if self.resolution is None: - self._calculate_resolution() - mask = self.masks() - return float(self.resolution[mask].min().item()) - def __len__(self) -> int: """Number of reflections (full array, ignoring masks).""" return len(self.hkl) if self.hkl is not None else 0 @property def d_min(self) -> Optional[float]: - """High-resolution limit: the smallest d-spacing, in Ångströms.""" - return self.get_max_res() + """High-resolution limit: smallest d-spacing of the valid reflections, in Å.""" + if self.resolution is None: + self._calculate_resolution() + return float(self.resolution[self.masks()].min().item()) def __repr__(self) -> str: """Count, data sources, resolution range and space group.""" @@ -1755,17 +1680,6 @@ def __repr__(self) -> str: return ", ".join(parts) + ")" - def get_valid_mask(self) -> torch.Tensor: - """ - Return the combined validity mask over all active filters. - - Returns - ------- - torch.Tensor - Boolean mask of shape (N,); True = valid/included. - """ - return self.masks() - def data_indexed( self, ) -> Tuple[ @@ -2620,7 +2534,7 @@ def _remap_mask(tensor, fill_value): remapped.spacegroup = self.spacegroup # Reindex ALL per-reflection dataclass fields onto new_hkl: sets hkl / - # resolution, invalidates bin_indices and _centric_flags, and fills + # resolution, invalidates _centric_flags, and fills # missing rows (index -1) per _REINDEX_FILL. self._reindex_per_reflection(index_mapping, new_hkl, target=remapped) @@ -2677,7 +2591,7 @@ def expand_to_p1( New object at ``spacegroup="P1"`` holding every symmetry-equivalent reflection (duplicates removed). Per-reflection fields are indexed from the original, ``phase`` additionally gets the translation phase - shift, ``resolution`` is recomputed and ``bin_indices`` is cleared. + shift, and ``resolution`` is recomputed. ``source``/``last_op`` record the provenance. """ if self.hkl is None: @@ -2764,7 +2678,7 @@ def generate_validation_set( :attr:`rfree_flags` untouched. The work/free/validation subsets are disjoint (validation is carved out of free) -- see :meth:`_subset_indices` and the ``work``/``free``/``validation`` - accessors. Like :meth:`_generate_rfree_flags`, the split is over whole + accessors. Like :meth:`generate_rfree_flags`, the split is over whole ASU groups so Bijvoet mates stay together (see :meth:`asu_group_indices`). @@ -2790,26 +2704,21 @@ def generate_validation_set( rwork = self.rfree_flags.to(torch.bool) free_mask = ~rwork - # Split whole ASU groups, exactly as _generate_rfree_flags does -- a + # Split whole ASU groups, exactly as generate_rfree_flags does -- a # per-row draw here would re-open the Friedel leak at the free/validation # boundary. The free set is already group-consistent, so a group is # wholly free or wholly work. group_id, n_groups = self.asu_group_indices() group_free = self._group_any(free_mask, group_id, n_groups) - # Reuse get_bins for resolution-stratified sampling. bin_indices, n_bins = self.get_bins(n_bins=20, min_per_bin=20) group_bin = bin_indices[self._group_representative_rows(group_id, n_groups)] - - group_val = torch.zeros(n_groups, dtype=torch.bool, device=self.device) - for b in range(n_bins): - bin_free_groups = torch.where((group_bin == b) & group_free)[0] - n_bin_free = int(bin_free_groups.numel()) - if n_bin_free == 0: - continue - n_val = max(1, int(n_bin_free * val_fraction_of_free)) - perm = torch.randperm(n_bin_free, device=bin_free_groups.device)[:n_val] - group_val[bin_free_groups[perm]] = True + group_val = self._stratified_group_draw( + group_free, + group_bin, + n_bins, + lambda n: max(1, int(n * val_fraction_of_free)), + ) # Broadcast to rows, staying within the free set. val_flags = group_val[group_id] & free_mask diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index 96ee5799..4311b6bd 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -322,10 +322,12 @@ def __init__( raise ValueError(f"max_res must be a float > 0, got {max_res!r}") if max_res_val <= 0: raise ValueError(f"max_res must be > 0, got {max_res_val}") - self.reflection_data = self.reflection_data.cut_res(max_res_val) + self.reflection_data = self.reflection_data.filter_by_resolution( + d_min=max_res_val + ) self.max_res = max_res_val else: - self.max_res = self.reflection_data.get_max_res() + self.max_res = self.reflection_data.d_min self.model = ModelFT( verbose=self.verbose, max_res=self.max_res, diff --git a/torchref/refinement/rigid_body_refinement.py b/torchref/refinement/rigid_body_refinement.py index 5718c19b..524ae4b7 100644 --- a/torchref/refinement/rigid_body_refinement.py +++ b/torchref/refinement/rigid_body_refinement.py @@ -144,16 +144,16 @@ def _run(self): ref = self.refinement original_data = ref.reflection_data - native_dmin = float(original_data.get_max_res()) + native_dmin = float(original_data.d_min) cutoffs = ( self.cutoffs if self.cutoffs is not None else self.default_cutoffs(native_dmin) ) - # ``cut_res`` masks in place and returns ``self``, so each cutoff below - # stamps ``masks["resolution"]`` on the caller's own object and rebinding - # restores nothing. Snapshot it (or its absence) to put back. + # ``filter_by_resolution`` masks in place and returns ``self``, so each + # cutoff below stamps ``masks["resolution"]`` on the caller's own object and + # rebinding restores nothing. Snapshot it (or its absence) to put back. had_resolution_mask = "resolution" in original_data.masks saved_resolution_mask = ( original_data.masks["resolution"].clone() if had_resolution_mask else None @@ -173,7 +173,7 @@ def restore_resolution_mask(): for d_min in cutoffs: xray_mode = self._xray_mode_for_cutoff(d_min) self._rebind_for_data( - original_data.cut_res(highres=float(d_min)), + original_data.filter_by_resolution(d_min=float(d_min)), xray_mode=xray_mode, ) step_state = self._run_one_cutoff(d_min) diff --git a/torchref/scaling/scaler_base.py b/torchref/scaling/scaler_base.py index b8350e48..f291baff 100644 --- a/torchref/scaling/scaler_base.py +++ b/torchref/scaling/scaler_base.py @@ -320,7 +320,7 @@ def setup_binwise_solvent_scale(self): Once this exists, :meth:`forward` uses it *instead of* the solvent model's global ``k_sol``/``B_sol``, which then stop affecting the result. """ - mean_res = self._data.mean_res_per_bin() + mean_res = self._data.mean_res_per_bin(self.bins, self.nbins) # Seeded from k_sol * exp(-B s^2) with Phenix-like k=0.35, B=46. s_per_bin = 1.0 / (2.0 * mean_res + 1e-6) # sin(theta)/lambda From d7a07b4b2ab19e8cbbb7633cfc8fcc8ce31e3d55 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 19:13:44 +0200 Subject: [PATCH 216/250] Share one row-gathering path across ReflectionData's reindexing __select__, _canonicalize_in_place, _reindex_per_reflection, validate_hkl and remap each enumerated the per-reflection fields and gathered the masks with their own loop. They now use _per_row_fields, _gather_rows and _gathered_masks. Outputs are identical field for field and mask for mask on 6G9X and 1DAW. Co-Authored-By: Claude Opus 5.5 (1M context) --- torchref/io/datasets/reflection_data.py | 190 ++++++++++-------------- 1 file changed, 76 insertions(+), 114 deletions(-) diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index ffa94802..a5506725 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -7,7 +7,7 @@ """ import warnings -from dataclasses import dataclass, field +from dataclasses import dataclass, field, fields from pathlib import Path from typing import TYPE_CHECKING, Callable, Optional, Tuple, Union @@ -22,6 +22,7 @@ from torchref.io.datasets.base import CrystalDataset from torchref.symmetry import Cell, SpaceGroup from torchref.utils.debug_utils import DebugMixin +from torchref.utils.utils import TensorMasks if TYPE_CHECKING: from torchref.model.model_ft import ModelFT @@ -383,6 +384,47 @@ def I_sigma_raw(self) -> Optional[torch.Tensor]: # the ``centric`` property). _REINDEX_DERIVED = ("resolution", "_centric_flags") + def _per_row_fields(self): + """Yield ``(name, tensor)`` for each per-reflection dataclass field. + + Enumerated generically (``shape[0] == len(hkl)``) so a new per-reflection + field is carried by every reindexing operation without being listed. + """ + n = len(self.hkl) if self.hkl is not None else 0 + for f in fields(self): + val = getattr(self, f.name) + if isinstance(val, torch.Tensor) and val.shape and val.shape[0] == n: + yield f.name, val + + @staticmethod + def _gather_rows(val: torch.Tensor, index: torch.Tensor, fill) -> torch.Tensor: + """``val[index]``, with rows where ``index == -1`` set to ``fill``.""" + present = index >= 0 + if bool(present.all()): + return val[index] + shape = (len(index),) + tuple(val.shape[1:]) + out = torch.full(shape, fill, dtype=val.dtype, device=val.device) + out[present] = val[index[present]] + return out + + def _gathered_masks(self, index: torch.Tensor) -> TensorMasks: + """Every mask gathered by ``index``; rows with ``index == -1`` are masked out. + + Call before ``hkl`` changes length: masks of any other length are dropped. + """ + n = len(self.hkl) if self.hkl is not None else 0 + out = TensorMasks(device=self.device) + for name, mask in self.masks.items(): + if mask is not None and len(mask) == n: + out[name] = self._gather_rows(mask, index, False) + return out + + def _replace_masks(self, new: TensorMasks) -> None: + """Swap in ``new``'s masks, keeping the existing ``TensorMasks`` object.""" + self.masks.clear() + for name, mask in new.items(): + self.masks[name] = mask + def _reindex_per_reflection( self, index_map: torch.Tensor, @@ -413,45 +455,26 @@ def _reindex_per_reflection( Boolean presence mask (``index_map >= 0``), for building the caller's ``hkl_present`` / ``missing`` masks. """ - from dataclasses import fields as dc_fields - if target is None: target = self - n_src = len(self.hkl) if self.hkl is not None else 0 new_hkl = new_hkl.to(dtype=dtypes.int, device=self.device) - n_out = len(new_hkl) index_map = index_map.to(device=self.device, dtype=torch.long) # dtype-ok: index map used for indexing/gather; PyTorch requires int64 present = index_map >= 0 - src_idx = index_map[present] - derived = set(self._REINDEX_DERIVED) - for f in dc_fields(self): - name = f.name - if name == "hkl" or name in derived: - continue - val = getattr(self, name) - if not isinstance(val, torch.Tensor): - continue - if not (val.shape and val.shape[0] == n_src): - # Non-per-reflection tensor: leave target's own value untouched - # (a fresh default when target is a new instance). - continue - if name == "hkl_anomalous": - # Present rows keep their signed (anomalous) index; missing rows - # fall back to the canonical reference HKL (never a 0,0,0 row). - out = new_hkl.clone() - out[present] = val[src_idx] - else: - fill = self._REINDEX_FILL.get(name, 0) - out = torch.full( - (n_out,) + tuple(val.shape[1:]), - fill, - dtype=val.dtype, - device=self.device, - ) - out[present] = val[src_idx] - setattr(target, name, out) + skip = {"hkl", *self._REINDEX_DERIVED} + # Collected first: writing into self (the in-place case) changes the + # row count _per_row_fields keys on. + gathered = { + name: self._gather_rows(val, index_map, self._REINDEX_FILL.get(name, 0)) + for name, val in self._per_row_fields() + if name not in skip + } + if "hkl_anomalous" in gathered: + # Missing rows fall back to the reference HKL, never a 0,0,0 row. + gathered["hkl_anomalous"][~present] = new_hkl[~present] + for name, val in gathered.items(): + setattr(target, name, val) # Install the new HKL and recompute / invalidate derived-from-HKL fields. target.hkl = new_hkl @@ -468,11 +491,9 @@ def _assert_per_reflection_consistent(self) -> None: Post-condition for the reindex routines; raises rather than letting a stale-length field surface as a downstream shape mismatch. """ - from dataclasses import fields as dc_fields - n = len(self.hkl) if self.hkl is not None else 0 bad = [] - for f in dc_fields(self): + for f in fields(self): val = getattr(self, f.name) if isinstance(val, torch.Tensor) and val.ndim >= 1 and val.shape[0] != n: bad.append((f.name, tuple(val.shape))) @@ -483,8 +504,6 @@ def _assert_per_reflection_consistent(self) -> None: def _canonicalize_in_place(self) -> None: """Remap HKL to canonical CCP4 ASU form and reorder all data in-place.""" - from dataclasses import fields as dc_fields - if self.hkl is None or self.spacegroup is None: return @@ -494,13 +513,10 @@ def _canonicalize_in_place(self) -> None: ) ) - n_refl = len(self.hkl) - - for f in dc_fields(self): - val = getattr(self, f.name) - if isinstance(val, torch.Tensor) and val.shape and val.shape[0] == n_refl: - setattr(self, f.name, val[sort_indices]) - + masks = self._gathered_masks(sort_indices) + for name, val in list(self._per_row_fields()): + setattr(self, name, val[sort_indices]) + self._replace_masks(masks) self.hkl = canonical_hkl if self.phase is not None: @@ -511,14 +527,6 @@ def _canonicalize_in_place(self) -> None: if self.cell is not None: self._calculate_resolution() - if hasattr(self, "masks") and self.masks is not None: - for name in list(self.masks.keys()): - mask_tensor = self.masks[name] - if mask_tensor is not None: - # Bypass __setitem__ validation (reordering preserves True count) - dict.__setitem__(self.masks, name, mask_tensor[sort_indices]) - self.masks._updated = True - # friedel_flags comes back already in sorted (canonical) order, matching # self.hkl. hkl_anomalous carries the SIGNED index used for # structure-factor evaluation -- canonical for the (+) member, negated @@ -1751,38 +1759,20 @@ def __select__(self, indices: torch.Tensor, op=None) -> "ReflectionData": ReflectionData New ReflectionData object with selected reflections. """ - from dataclasses import fields as dc_fields - - from torchref.utils.utils import TensorMasks - - n_refl = len(self.hkl) if self.hkl is not None else 0 - - # Create new instance with same device + if indices.dtype == torch.bool: + indices = torch.nonzero(indices).squeeze(-1) selected = ReflectionData(verbose=self.verbose, device=self.device) - for f in dc_fields(self): + per_row = dict(self._per_row_fields()) + for f in fields(self): val = getattr(self, f.name) - if val is None: - continue - if isinstance(val, torch.Tensor): - if val.shape and val.shape[0] == n_refl: - setattr(selected, f.name, val[indices]) - else: - # Preserve scalar tensor metadata. - setattr(selected, f.name, val.clone()) - elif isinstance(val, Cell): + if f.name in per_row: + setattr(selected, f.name, val[indices]) + elif isinstance(val, (torch.Tensor, Cell)): setattr(selected, f.name, val.clone()) - else: - # Scalars, strings, None, gemmi objects, etc. + elif val is not None: setattr(selected, f.name, val) - - # Handle masks (not a dataclass field) - if hasattr(self, "masks") and self.masks is not None and len(self.masks) > 0: - new_masks = TensorMasks(device=self.device) - for name, mask_tensor in self.masks.items(): - if mask_tensor is not None: - new_masks[name] = mask_tensor[indices] - selected.masks = new_masks + selected.masks = self._gathered_masks(indices) selected.source = self selected.last_op = op @@ -1904,26 +1894,13 @@ def validate_hkl( # Reindex EVERY per-reflection field via the shared primitive. Masks are # handled separately below because they are not dataclass fields. + masks = self._gathered_masks(valid_indices) presence_mask = self._reindex_per_reflection(valid_indices, hkl_ref) if identity_hkl is not None: self.hkl_anomalous = identity_hkl.to(self.hkl).clone() self.friedel_flags = (self.hkl_anomalous != self.hkl).any(dim=-1) - - # Transfer existing masks to new indexing - old_masks = dict(self.masks.items()) - # Clear existing masks - self.masks.clear() - self.masks._updated = True - - for name, old_mask in old_masks.items(): - if old_mask is not None and len(old_mask) == n_data: - # Expand mask: missing reflections are masked out (False) - new_mask = torch.zeros(n_ref, dtype=torch.bool, device=self.device) - mask = valid_indices >= 0 - new_mask[mask] = old_mask[valid_indices[mask]] - self.masks[name] = new_mask - - # Add presence mask - this is the key mask that marks real vs placeholder data + self._replace_masks(masks) + # The mask that tells real reflections from placeholder rows. self.masks["hkl_present"] = presence_mask n_present = presence_mask.sum().item() @@ -2509,21 +2486,6 @@ def remap( """ from torchref.symmetry.spacegroup import SpaceGroup - # Mask remapper. Masks are not dataclass fields, so the shared - # per-reflection reindexer below does not touch them. - def _remap_mask(tensor, fill_value): - if tensor is None: - return None - valid_mask = index_mapping >= 0 - result = torch.full( - (len(new_hkl),) + tensor.shape[1:], - fill_value, - dtype=tensor.dtype, - device=self.device, - ) - result[valid_mask] = tensor[index_mapping[valid_mask]] - return result - # Create new ReflectionData; set cell/spacegroup first so the shared # reindexer can recompute resolution on the new grid. remapped = ReflectionData(verbose=self.verbose, device=self.device) @@ -2545,9 +2507,9 @@ def _remap_mask(tensor, fill_value): # Carry forward prior combined mask if available. prior_mask = self.masks() if prior_mask is not None: - remapped.masks["prior_flagged"] = _remap_mask( - prior_mask.to(dtype=dtypes.int), fill_value=0 - ).to(torch.bool) + remapped.masks["prior_flagged"] = self._gather_rows( + prior_mask, index_mapping.to(self.device), False + ) # Copy metadata sources remapped.amplitude_source = self.amplitude_source From 8a6bf4bd85ca3f557341cd0c9e6267c06a544528 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 19:18:09 +0200 Subject: [PATCH 217/250] Move the Wilson B fit out of ReflectionData fit_wilson_b(F, d) in torchref.scaling.wilson holds the same binned two-component fit and returns the structure B, so the dataset no longer carries four fit results as fields. The ensemble Wilson prior, its only consumer, calls it directly. B is unchanged on 1DAW, 6G9X, 3E98 and 5BOV. The fallback for fewer than three usable shells is now an explicit constant, at the 200 A^2 the fit actually returned: it chose between 50 and 200 by testing its label for "struct", which the structure-B call's "high-res" label never contained. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 1 + tests/unit/io/test_data.py | 9 - tests/unit/refinement/test_wilson_prior.py | 1 - .../experimental/ensemble/wilson_prior.py | 13 +- torchref/io/datasets/base.py | 6 - torchref/io/datasets/reflection_data.py | 236 ------------------ torchref/scaling/wilson.py | 208 ++++++++++++++- 7 files changed, 213 insertions(+), 261 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index be42b79f..50ca256e 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- The Wilson B fit moved out of ``ReflectionData`` into ``torchref.scaling.wilson.fit_wilson_b(F, d)``, which returns the structure B (same values as before). Removed the ``wilson_b``, ``wilson_b_structure``, ``wilson_b_solvent`` and ``wilson_k_sol`` dataset fields; nothing but the ensemble Wilson prior read them - ``ReflectionData.regenerate_rfree_flags`` is renamed ``generate_rfree_flags`` (same arguments; with existing flags and ``force=False`` it now warns instead of printing), and it prints only when ``verbose > 0``. Seeded draws are unchanged - ``ReflectionData.get_bins`` no longer stores ``bin_indices`` on the dataset, and ``mean_res_per_bin`` takes the bins it should average over, so a later ``get_bins`` call with other settings (R-free or validation-set generation, a least-squares target) can no longer shift the shells a scaler's per-bin solvent scale was set up on. The ``bin_indices`` field is removed - Removed the ``ReflectionData`` aliases ``get_max_res`` (use ``d_min``), ``get_valid_mask`` (use ``masks()``) and ``cut_res`` (use ``filter_by_resolution(d_min=, d_max=)``, which now prints only when ``verbose > 0``) diff --git a/tests/unit/io/test_data.py b/tests/unit/io/test_data.py index e4e58a6d..74e64aac 100644 --- a/tests/unit/io/test_data.py +++ b/tests/unit/io/test_data.py @@ -124,15 +124,6 @@ def test_has_verbose_attribute(self): class TestReflectionDataProperties: """Tests for ReflectionData computed properties.""" - @pytest.mark.unit - def test_wilson_b_default_none(self): - """Wilson B should be None initially.""" - from torchref.io import ReflectionData - - data = ReflectionData() - - assert data.wilson_b is None - @pytest.mark.unit def test_spacegroup_default_none(self): """Space group should be None initially.""" diff --git a/tests/unit/refinement/test_wilson_prior.py b/tests/unit/refinement/test_wilson_prior.py index 3a9c1e1f..a0f4780d 100644 --- a/tests/unit/refinement/test_wilson_prior.py +++ b/tests/unit/refinement/test_wilson_prior.py @@ -25,7 +25,6 @@ def setup_target(): data = ReflectionData(verbose=0) data.load_mtz(TEST_MTZ) - data._calculate_wilson_b() ens = EnsembleModel.from_single( TEST_PDB, n_members=4, perturb_sigma=0.0, b_const=5.0, seed=42, verbose=0, max_res=data.d_min, diff --git a/torchref/experimental/ensemble/wilson_prior.py b/torchref/experimental/ensemble/wilson_prior.py index 44490f9c..435922eb 100644 --- a/torchref/experimental/ensemble/wilson_prior.py +++ b/torchref/experimental/ensemble/wilson_prior.py @@ -26,8 +26,7 @@ loss = mean_bin( ( log<|F_calc|^2>_bin - log Wilson_expected(s_bin) )^2 ) The reference curve is fit once from the observed data: -``B_W = data.wilson_b`` (already computed by -``ReflectionData._calculate_wilson_b``) and ``K`` from a single +``B_W`` from :func:`torchref.scaling.wilson.fit_wilson_b` and ``K`` from a single least-squares fit at first ``forward()`` call. Used as ``'regularization/wilson'`` in the ensemble refinement LossState. @@ -92,8 +91,7 @@ class WilsonPriorTarget(DataTarget): Parameters ---------- data : ReflectionData - Reflection data. Must have ``wilson_b`` populated (the loader - already does this). + Reflection data; its Wilson B is fitted from ``F`` and ``resolution``. model : ModelFT Atomic model used to compute F_calc. scaler : Scaler @@ -194,10 +192,11 @@ def _fit_K_from_observed(self) -> None: Fit the prefactor ``K`` of the Wilson curve from observed binned intensities so the prior is centered on the observed scale. """ - wilson_b = getattr(self._data, "wilson_b", None) + from torchref.scaling.wilson import fit_wilson_b + + wilson_b = fit_wilson_b(self._data.F, self._data.resolution) if wilson_b is None: - self._data._calculate_wilson_b() - wilson_b = self._data.wilson_b + raise ValueError("Too few reflections to fit a Wilson B for the prior.") device = self._data.device self._B_W = torch.tensor(float(wilson_b), device=device) diff --git a/torchref/io/datasets/base.py b/torchref/io/datasets/base.py index 966122f8..e2a47753 100644 --- a/torchref/io/datasets/base.py +++ b/torchref/io/datasets/base.py @@ -80,12 +80,6 @@ class CrystalDataset(DeviceMovementMixin): intensity_source: Optional[str] = None phase_source: Optional[str] = None - # === Wilson B-factors === - wilson_b: Optional[float] = None - wilson_b_structure: Optional[float] = None - wilson_b_solvent: Optional[float] = None - wilson_k_sol: Optional[float] = None - # === Masks (initialized in __post_init__) === # Note: masks is not a dataclass field to avoid serialization issues # It's initialized in __post_init__ and handled specially diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index a5506725..3306aeaa 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -200,8 +200,6 @@ class ReflectionData(CrystalDataset, DebugMixin): ``torchref.symmetry.SpaceGroup`` object here. resolution : torch.Tensor Resolution per reflection in Ångströms of shape (N,). - wilson_b : float - Overall Wilson B-factor in Ų. """ # Additional fields specific to ReflectionData (beyond CrystalDataset) @@ -1385,240 +1383,6 @@ def _calculate_resolution(self) -> None: resolution = 1.0 / torch.linalg.norm(s, axis=1) self.resolution = resolution - def _calculate_wilson_b(self, n_bins: int = 30) -> None: - """Fit a two-component Wilson plot, `` ∝ A_s·exp(-2B_s·s²) + A_sol·exp(-2B_sol·s²)``. - - Sets ``wilson_b`` (= the structure B), ``wilson_b_structure``, - ``wilson_b_solvent`` and ``wilson_k_sol`` (clamped to 0.01-0.9). Returns - early without setting anything when there are too few valid reflections - or bins, so callers must not assume the attributes exist afterwards. - ``n_bins`` (default 30) is the averaging bin count. - """ - if self.F is None or self.resolution is None: - return - - # Get valid reflections - F = self.F - d = self.resolution - valid = torch.isfinite(F) & (F > 0) & torch.isfinite(d) - - if valid.sum() < 100: - if self.verbose > 0: - print( - f" Wilson B: insufficient data ({valid.sum()} reflections), skipping" - ) - return - - F_valid = F[valid] - d_valid = d[valid] - - # Calculate s² = 1/(4d²) - s_sq = 1.0 / (4.0 * d_valid**2) - F_sq = F_valid**2 - - # Bin the data for noise reduction - s_sq_min, s_sq_max = s_sq.min(), s_sq.max() - bin_edges = torch.linspace(s_sq_min, s_sq_max, n_bins + 1, device=self.device) - bin_centers = (bin_edges[:-1] + bin_edges[1:]) / 2 - bin_idx = torch.bucketize(s_sq, bin_edges[1:-1]) - - # Calculate mean F² per bin - bin_sums = torch.zeros(n_bins, device=self.device, dtype=F_sq.dtype) - bin_counts = torch.zeros(n_bins, device=self.device, dtype=F_sq.dtype) - bin_sums.scatter_add_(0, bin_idx, F_sq) - bin_counts.scatter_add_(0, bin_idx, torch.ones_like(F_sq)) - - valid_bins = bin_counts > 5 - if valid_bins.sum() < 5: - if self.verbose > 0: - print(f" Wilson B: insufficient bins ({valid_bins.sum()}), skipping") - return - - mean_F_sq = bin_sums[valid_bins] / bin_counts[valid_bins] - s_sq_bins = bin_centers[valid_bins] - - # Convert s² back to d-spacing for resolution-based selection - d_bins = 1.0 / (2.0 * torch.sqrt(s_sq_bins)) - - # Stage 1: Fit high-resolution region (d < 3.5 Å) for structure B - high_res_mask = d_bins < 3.5 - B_struct = self._fit_single_wilson( - s_sq_bins, mean_F_sq, high_res_mask, "high-res" - ) - - # Stage 2: Fit low-resolution region (d > 6 Å) for solvent B - low_res_mask = d_bins > 6.0 - B_sol = self._fit_single_wilson(s_sq_bins, mean_F_sq, low_res_mask, "low-res") - - # Stage 3: Two-component fit across all data - B_struct_final, B_sol_final, k_sol = self._fit_two_component_wilson( - s_sq_bins, mean_F_sq, B_struct, B_sol - ) - - # Store results - self.wilson_b_structure = B_struct_final - self.wilson_b_solvent = B_sol_final - self.wilson_k_sol = k_sol - - # Overall Wilson B is the structure B (what people usually mean by "Wilson B") - self.wilson_b = B_struct_final - - if self.verbose > 0: - print(f" Wilson B-factor (structure): {B_struct_final:.1f} Ų") - print(f" Wilson B-factor (solvent): {B_sol_final:.1f} Ų") - print(f" Solvent fraction (k_sol): {k_sol:.3f}") - - def _fit_single_wilson( - self, - s_sq: torch.Tensor, - mean_F_sq: torch.Tensor, - mask: torch.Tensor, - label: str, - ) -> float: - """Fit ``ln(F²) = c - 2B·s²`` over the masked bins, clamped to 0-300 Ų. - - ``label`` is not only cosmetic: with fewer than 3 usable bins the - fallback is 50.0 when it contains "struct", else 200.0. - """ - if mask.sum() < 3: - # Not enough data, return reasonable default - if self.verbose > 1: - print( - f" Wilson {label}: insufficient bins ({mask.sum()}), using default" - ) - return 50.0 if "struct" in label else 200.0 - - x = s_sq[mask] - y = torch.log(mean_F_sq[mask]) - - # Linear regression: ln(F²) = const - 2B*s² - x_mean = x.mean() - y_mean = y.mean() - - numerator = ((x - x_mean) * (y - y_mean)).sum() - denominator = ((x - x_mean) ** 2).sum() - - if denominator < 1e-12: - return 50.0 if "struct" in label else 200.0 - - slope = numerator / denominator - B = -slope.item() / 2.0 - - # Sanity bounds - B = max(0.0, min(B, 300.0)) - - return B - - def _fit_two_component_wilson( - self, - s_sq: torch.Tensor, - mean_F_sq: torch.Tensor, - B_struct_init: float, - B_sol_init: float, - n_iter: int = 50, - ) -> Tuple[float, float, float]: - """Refine ``F² = A·[(1-k)·exp(-2B_s·s²) + k·exp(-2B_sol·s²)]`` by finite-difference descent. - - ``k`` is the relative solvent contribution at s²=0. Returns - ``(B_struct, B_sol, k_sol)``, constrained to B_s in 1-200, B_sol in - 50-500 with ``B_sol >= B_struct + 20``, and k in 0.01-0.9 -- so a - returned value sitting exactly on a bound means the fit hit the clamp. - """ - # Normalize F² for numerical stability - F_sq_max = mean_F_sq.max() - y = mean_F_sq / F_sq_max - x = s_sq - - # Initialize parameters - B_struct = torch.tensor(B_struct_init, device=self.device, dtype=x.dtype) - B_sol = torch.tensor(B_sol_init, device=self.device, dtype=x.dtype) - - # Estimate initial k from ratio of low-res to high-res decay - # At low resolution, solvent contributes more - d_from_s = 1.0 / (2.0 * torch.sqrt(x)) - low_res_val = y[d_from_s > 5.0].mean() if (d_from_s > 5.0).any() else y[0] - high_res_val = y[d_from_s < 3.0].mean() if (d_from_s < 3.0).any() else y[-1] - - # k estimates solvent fraction - if low res is much higher than expected - # from structure alone, there's solvent contribution - struct_decay = torch.exp(-2 * B_struct * x) - expected_low = ( - struct_decay[d_from_s > 5.0].mean() - if (d_from_s > 5.0).any() - else struct_decay[0] - ) - - if expected_low > 1e-6 and low_res_val > expected_low: - k_init = min(0.5, (low_res_val - expected_low).item() / low_res_val.item()) - else: - k_init = 0.1 - - k = torch.tensor(max(0.01, min(0.5, k_init)), device=self.device, dtype=x.dtype) - - # Simple gradient descent refinement - lr = 0.1 - - for _ in range(n_iter): - # Compute model - struct_term = (1 - k) * torch.exp(-2 * B_struct * x) - sol_term = k * torch.exp(-2 * B_sol * x) - model = struct_term + sol_term - - # Compute scale factor analytically - A = (y * model).sum() / (model * model).sum() - model_scaled = A * model - - # Compute gradients (simplified, using finite differences for robustness) - eps = 0.1 - - # B_struct gradient - model_plus = A * ((1 - k) * torch.exp(-2 * (B_struct + eps) * x) + sol_term) - model_minus = A * ( - (1 - k) * torch.exp(-2 * (B_struct - eps) * x) + sol_term - ) - loss_plus = ((y - model_plus) ** 2).sum() - loss_minus = ((y - model_minus) ** 2).sum() - grad_B_struct = (loss_plus - loss_minus) / (2 * eps) - - # B_sol gradient - model_plus = A * (struct_term + k * torch.exp(-2 * (B_sol + eps) * x)) - model_minus = A * (struct_term + k * torch.exp(-2 * (B_sol - eps) * x)) - loss_plus = ((y - model_plus) ** 2).sum() - loss_minus = ((y - model_minus) ** 2).sum() - grad_B_sol = (loss_plus - loss_minus) / (2 * eps) - - # k gradient - eps_k = 0.01 - k_plus = min(0.9, k + eps_k) - k_minus = max(0.01, k - eps_k) - model_plus = A * ( - (1 - k_plus) * torch.exp(-2 * B_struct * x) - + k_plus * torch.exp(-2 * B_sol * x) - ) - model_minus = A * ( - (1 - k_minus) * torch.exp(-2 * B_struct * x) - + k_minus * torch.exp(-2 * B_sol * x) - ) - loss_plus = ((y - model_plus) ** 2).sum() - loss_minus = ((y - model_minus) ** 2).sum() - grad_k = (loss_plus - loss_minus) / (2 * eps_k) - - # Update parameters - B_struct = B_struct - lr * grad_B_struct - B_sol = B_sol - lr * grad_B_sol - k = k - lr * 0.1 * grad_k # Slower learning rate for k - - # Enforce constraints - B_struct = torch.clamp(B_struct, 1.0, 200.0) - B_sol = torch.clamp(B_sol, 50.0, 500.0) - k = torch.clamp(k, 0.01, 0.9) - - # Ensure B_sol > B_struct (solvent is more disordered) - if B_sol < B_struct + 20: - B_sol = B_struct + 20 - - return B_struct.item(), B_sol.item(), k.item() - def filter_by_resolution( self, d_min: Optional[float] = None, d_max: Optional[float] = None ) -> "ReflectionData": diff --git a/torchref/scaling/wilson.py b/torchref/scaling/wilson.py index e16b048d..175eb533 100644 --- a/torchref/scaling/wilson.py +++ b/torchref/scaling/wilson.py @@ -10,7 +10,7 @@ **Why this exists as one shared class.** The repo grew at least five private answers to the same question -- ``base/wilson_outliers.robust_mean_intensity``, ``base/french_wilson.estimate_mean_intensity_by_resolution``, -``ReflectionData._calculate_wilson_b``, the ``Sigma_N`` estimator in +:func:`fit_wilson_b` below, the ``Sigma_N`` estimator in ``refinement/model_error_estimation/sigma_a``, and a per-shell one inside the alignment package -- differing in whether they use means or medians, whether they divide out ``epsilon``, whether they separate centrics, and where they put @@ -24,6 +24,9 @@ made the previous convention object impossible to reason about: it returned a normalisation and a weight together, so sweeping it moved a gauge quantity and a real one at the same time. + +:func:`fit_wilson_b` is the one-number summary: an overall Wilson B for priors +and reports, not a curve to normalise by. """ from __future__ import annotations @@ -35,7 +38,7 @@ from torchref.config import get_float_dtype from torchref.scaling.basis import chebyshev_design -__all__ = ["WilsonNormaliser"] +__all__ = ["WilsonNormaliser", "fit_wilson_b"] #: Chebyshev terms. Enough to follow a Wilson plot's curvature and the #: low-resolution solvent deficit without chasing shell-to-shell noise. @@ -479,3 +482,204 @@ def __repr__(self) -> str: # pragma: no cover - display f"n_coeff={self.n_coeff}, n_fitted={self.n_fitted}, " f"iters={self.n_iter})" ) + + +#: Fallback structure B (Ų) when the high-resolution fit has fewer than three +#: usable shells. This is the value the fit has always returned there. +_FALLBACK_B = 200.0 + + +def fit_wilson_b( + F: torch.Tensor, + d: torch.Tensor, + n_bins: int = 30, + verbose: int = 0, +) -> Optional[float]: + """Fit an overall Wilson B from amplitudes and resolution. + + Bins ``F^2`` in ``s^2 = 1/(4 d^2)``, seeds a structure B from the shells + below 3.5 Å and a solvent B from those above 6 Å, then refines the + two-component curve `` = A [(1-k) exp(-2 B s^2) + k exp(-2 B_sol s^2)]`` + and returns its structure B -- what "the Wilson B" usually means. This is a + single number for priors and reports; for normalising intensities use + :class:`WilsonNormaliser`. + + Parameters + ---------- + F : torch.Tensor + Amplitudes of shape (N,). Non-finite and non-positive values are ignored. + d : torch.Tensor + Resolution of each reflection in Å, shape (N,). + n_bins : int, optional + Number of equal-width ``s^2`` bins. Default 30. + verbose : int, optional + Print the result when > 0. + + Returns + ------- + float or None + Structure B in Ų, within 1-200. ``None`` with fewer than 100 usable + reflections or fewer than 5 bins holding more than 5 each. + """ + device = F.device + valid = torch.isfinite(F) & (F > 0) & torch.isfinite(d) + if valid.sum() < 100: + if verbose > 0: + print(f" Wilson B: too few reflections ({int(valid.sum())}), skipping") + return None + + s_sq = 1.0 / (4.0 * d[valid] ** 2) + F_sq = F[valid] ** 2 + + bin_edges = torch.linspace(s_sq.min(), s_sq.max(), n_bins + 1, device=device) + bin_centers = (bin_edges[:-1] + bin_edges[1:]) / 2 + bin_idx = torch.bucketize(s_sq, bin_edges[1:-1]) + bin_sums = torch.zeros(n_bins, device=device, dtype=F_sq.dtype) + bin_counts = torch.zeros(n_bins, device=device, dtype=F_sq.dtype) + bin_sums.scatter_add_(0, bin_idx, F_sq) + bin_counts.scatter_add_(0, bin_idx, torch.ones_like(F_sq)) + + valid_bins = bin_counts > 5 + if valid_bins.sum() < 5: + if verbose > 0: + print(f" Wilson B: insufficient bins ({valid_bins.sum()}), skipping") + return None + + mean_F_sq = bin_sums[valid_bins] / bin_counts[valid_bins] + s_sq_bins = bin_centers[valid_bins] + d_bins = 1.0 / (2.0 * torch.sqrt(s_sq_bins)) + + B_struct = _fit_single_wilson(s_sq_bins, mean_F_sq, d_bins < 3.5) + B_sol = _fit_single_wilson(s_sq_bins, mean_F_sq, d_bins > 6.0) + B, _, _ = _fit_two_component_wilson(s_sq_bins, mean_F_sq, B_struct, B_sol) + + if verbose > 0: + print(f" Wilson B-factor (structure): {B:.1f} Ų") + return B + + +def _fit_single_wilson( + s_sq: torch.Tensor, mean_F_sq: torch.Tensor, mask: torch.Tensor +) -> float: + """Fit ``ln(F^2) = c - 2 B s^2`` over the masked shells, clamped to 0-300 Ų.""" + if mask.sum() < 3: + return _FALLBACK_B + x = s_sq[mask] + y = torch.log(mean_F_sq[mask]) + x_c, y_c = x - x.mean(), y - y.mean() + denominator = (x_c**2).sum() + if denominator < 1e-12: + return _FALLBACK_B + B = -((x_c * y_c).sum() / denominator).item() / 2.0 + return max(0.0, min(B, 300.0)) + + +def _fit_two_component_wilson( + s_sq: torch.Tensor, + mean_F_sq: torch.Tensor, + B_struct_init: float, + B_sol_init: float, + n_iter: int = 50, +) -> Tuple[float, float, float]: + """Refine ``F^2 = A [(1-k) exp(-2 B_s s^2) + k exp(-2 B_sol s^2)]``. + + Finite-difference gradient descent from the single-component estimates. + + Returns ``(B_struct, B_sol, k_sol)``, constrained to B_s in 1-200, B_sol in + 50-500 with ``B_sol >= B_struct + 20``, and k in 0.01-0.9 -- a value on a + bound means the fit hit the clamp. + """ + device = s_sq.device + F_sq_max = mean_F_sq.max() + y = mean_F_sq / F_sq_max + x = s_sq + + # Initialize parameters + B_struct = torch.tensor(B_struct_init, device=device, dtype=x.dtype) + B_sol = torch.tensor(B_sol_init, device=device, dtype=x.dtype) + + # Estimate initial k from ratio of low-res to high-res decay + # At low resolution, solvent contributes more + d_from_s = 1.0 / (2.0 * torch.sqrt(x)) + low_res_val = y[d_from_s > 5.0].mean() if (d_from_s > 5.0).any() else y[0] + high_res_val = y[d_from_s < 3.0].mean() if (d_from_s < 3.0).any() else y[-1] + + # k estimates solvent fraction - if low res is much higher than expected + # from structure alone, there's solvent contribution + struct_decay = torch.exp(-2 * B_struct * x) + expected_low = ( + struct_decay[d_from_s > 5.0].mean() + if (d_from_s > 5.0).any() + else struct_decay[0] + ) + + if expected_low > 1e-6 and low_res_val > expected_low: + k_init = min(0.5, (low_res_val - expected_low).item() / low_res_val.item()) + else: + k_init = 0.1 + + k = torch.tensor(max(0.01, min(0.5, k_init)), device=device, dtype=x.dtype) + + # Simple gradient descent refinement + lr = 0.1 + + for _ in range(n_iter): + # Compute model + struct_term = (1 - k) * torch.exp(-2 * B_struct * x) + sol_term = k * torch.exp(-2 * B_sol * x) + model = struct_term + sol_term + + # Compute scale factor analytically + A = (y * model).sum() / (model * model).sum() + model_scaled = A * model + + # Compute gradients (simplified, using finite differences for robustness) + eps = 0.1 + + # B_struct gradient + model_plus = A * ((1 - k) * torch.exp(-2 * (B_struct + eps) * x) + sol_term) + model_minus = A * ( + (1 - k) * torch.exp(-2 * (B_struct - eps) * x) + sol_term + ) + loss_plus = ((y - model_plus) ** 2).sum() + loss_minus = ((y - model_minus) ** 2).sum() + grad_B_struct = (loss_plus - loss_minus) / (2 * eps) + + # B_sol gradient + model_plus = A * (struct_term + k * torch.exp(-2 * (B_sol + eps) * x)) + model_minus = A * (struct_term + k * torch.exp(-2 * (B_sol - eps) * x)) + loss_plus = ((y - model_plus) ** 2).sum() + loss_minus = ((y - model_minus) ** 2).sum() + grad_B_sol = (loss_plus - loss_minus) / (2 * eps) + + # k gradient + eps_k = 0.01 + k_plus = min(0.9, k + eps_k) + k_minus = max(0.01, k - eps_k) + model_plus = A * ( + (1 - k_plus) * torch.exp(-2 * B_struct * x) + + k_plus * torch.exp(-2 * B_sol * x) + ) + model_minus = A * ( + (1 - k_minus) * torch.exp(-2 * B_struct * x) + + k_minus * torch.exp(-2 * B_sol * x) + ) + loss_plus = ((y - model_plus) ** 2).sum() + loss_minus = ((y - model_minus) ** 2).sum() + grad_k = (loss_plus - loss_minus) / (2 * eps_k) + + # Update parameters + B_struct = B_struct - lr * grad_B_struct + B_sol = B_sol - lr * grad_B_sol + k = k - lr * 0.1 * grad_k # Slower learning rate for k + + # Enforce constraints + B_struct = torch.clamp(B_struct, 1.0, 200.0) + B_sol = torch.clamp(B_sol, 50.0, 500.0) + k = torch.clamp(k, 0.01, 0.9) + + # Ensure B_sol > B_struct (solvent is more disordered) + if B_sol < B_struct + 20: + B_sol = B_struct + 20 + + return B_struct.item(), B_sol.item(), k.item() From 0597120780388f1d5d8084a923a1beae3bfd1f7e Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 28 Sep 2026 22:01:00 +0200 Subject: [PATCH 218/250] Fix expand_to_p1 on anomalous data and share one symmetry-expansion path _expand_hkl deduplicated P1 indices by keeping the first input row to reach each one. The two members of a Bijvoet pair share a canonical hkl, so expand_to_p1 on anomalous data kept an arbitrary mate per reflection: on 1DAW with unmeasured F(-) about half the reflections disappeared from the P1 set (24,484 usable against 46,690 for the merged file). expand_to_p1 now expands anomalous data from hkl_anomalous, so F(+) and F(-) keep their own P1 indices, and it sets the P1 rows' hkl_anomalous to their own index. _expand_hkl is vectorised on a new equivalent_hkl primitive (also used by merge_to_spacegroup instead of its own copy) and raises when two input rows are symmetry-equivalent. Output for merged data is identical in all 24 cases checked. Map, DifferenceMap and the real-space targets place each amplitude and its conjugate themselves, so they now average Bijvoet mates first (bijvoet_mean / bijvoet_representatives); merged-data maps are bit-identical. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 2 + tests/unit/io/test_merging.py | 29 +++ tests/unit/io/test_reflection_data_reindex.py | 12 +- tests/unit/maps/test_map.py | 66 +++++++ torchref/experimental/targets/realspace.py | 20 +- torchref/io/datasets/merging.py | 4 +- torchref/io/datasets/reflection_data.py | 107 ++++++++++- torchref/maps/difference_map.py | 14 +- torchref/maps/map.py | 14 +- torchref/symmetry/__init__.py | 3 +- torchref/symmetry/reciprocal_symmetry.py | 171 +++++++++++------- torchref/symmetry/spacegroup.py | 46 +++++ 12 files changed, 395 insertions(+), 93 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 50ca256e..8eacaf99 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,8 @@ Changelog Unreleased ---------- +- Fixed ``ReflectionData.expand_to_p1`` on anomalous data keeping only one member of each Bijvoet pair (whichever came first in row order), which dropped about half the reflections when one mate was unmeasured. Anomalous data now expand from their signed indices, so F(+) and F(-) keep their own P1 reflections, and the P1 dataset's ``hkl_anomalous`` is its own index rather than the source row's. ``SpaceGroup.expand_hkl`` raises when two input rows are symmetry-equivalent instead of silently keeping one; ``SpaceGroup.equivalent_hkl`` returns every symmetry copy with its source row +- ``Map``, ``DifferenceMap`` and the experimental real-space targets average each Bijvoet pair (``ReflectionData.bijvoet_mean`` / ``bijvoet_representatives``) before placing amplitudes on the grid, so anomalous input gives the Friedel-averaged map; maps from merged data are unchanged - The Wilson B fit moved out of ``ReflectionData`` into ``torchref.scaling.wilson.fit_wilson_b(F, d)``, which returns the structure B (same values as before). Removed the ``wilson_b``, ``wilson_b_structure``, ``wilson_b_solvent`` and ``wilson_k_sol`` dataset fields; nothing but the ensemble Wilson prior read them - ``ReflectionData.regenerate_rfree_flags`` is renamed ``generate_rfree_flags`` (same arguments; with existing flags and ``force=False`` it now warns instead of printing), and it prints only when ``verbose > 0``. Seeded draws are unchanged - ``ReflectionData.get_bins`` no longer stores ``bin_indices`` on the dataset, and ``mean_res_per_bin`` takes the bins it should average over, so a later ``get_bins`` call with other settings (R-free or validation-set generation, a least-squares target) can no longer shift the shells a scaler's per-bin solvent scale was set up on. The ``bin_indices`` field is removed diff --git a/tests/unit/io/test_merging.py b/tests/unit/io/test_merging.py index 23474d6b..fff64ee9 100644 --- a/tests/unit/io/test_merging.py +++ b/tests/unit/io/test_merging.py @@ -223,3 +223,32 @@ def test_from_tensors_keeps_intensities_and_validation_row_aligned(): torch.testing.assert_close(d.I, _invariant(d.hkl)) expected = (_invariant(d.hkl).round().to(torch.int64) % 2) == 0 assert torch.equal(d.validation_flags, expected) + + +@pytest.mark.unit +def test_expand_to_p1_keeps_both_bijvoet_mates(): + hkl = _asu_grid() + hkl = hkl[~SpaceGroup(SG).is_centric(hkl)] + both = torch.cat([hkl, -hkl]) + I = torch.cat([_invariant(hkl), _invariant(hkl) * 1.1]) + d = _data(both, I, torch.ones_like(I), friedel_merged=False) + + for include_friedel in (True, False): + p1 = d.expand_to_p1(include_friedel=include_friedel) + p1._assert_per_reflection_consistent() + values = _by_hkl(p1, p1.I) + # Every measurement lands on its own P1 index: none is dropped, and + # F(-) is never replaced by a Friedel copy of F(+). + assert len(values) == len(p1.hkl) + assert torch.equal(p1._hkl_for_sf(), p1.hkl) + for h, v in zip(hkl.tolist(), _invariant(hkl).tolist()): + assert values[tuple(h)] == pytest.approx(v) + assert values[tuple(-x for x in h)] == pytest.approx(1.1 * v) + + +@pytest.mark.unit +def test_expansion_refuses_symmetry_equivalent_rows(): + hkl = _asu_grid() + dup = torch.cat([hkl, hkl[:5] * torch.tensor([-1, -1, 1], dtype=hkl.dtype)]) + with pytest.raises(ValueError, match="symmetry-equivalent"): + SpaceGroup(SG).expand_hkl(dup) diff --git a/tests/unit/io/test_reflection_data_reindex.py b/tests/unit/io/test_reflection_data_reindex.py index b398904e..de8fd132 100644 --- a/tests/unit/io/test_reflection_data_reindex.py +++ b/tests/unit/io/test_reflection_data_reindex.py @@ -25,6 +25,14 @@ def _base_grid(h=10, k=10, lmax=10): ) +def _asu_unique(hkl, sg="P 21 21 21"): + """One row per unique reflection of ``sg``: expansion refuses equivalent rows.""" + from torchref.symmetry import SpaceGroup + + canon, *_ = SpaceGroup(sg).canonicalize_hkl(hkl, include_friedel=True) + return torch.unique(canon, dim=0) + + def _synthetic(hkl, seed=0, device="cpu"): n = hkl.shape[0] g = torch.Generator().manual_seed(seed) @@ -79,7 +87,7 @@ class TestP1RoundTripReindex: field at the new length.""" def test_expand_to_p1_carries_validation_flags(self): - grid = _base_grid(6, 6, 6) + grid = _asu_unique(_base_grid(6, 6, 6)) d = _synthetic(grid, seed=3) d.generate_validation_set(val_fraction_of_free=0.5, seed=0) assert d.validation_flags is not None @@ -93,7 +101,7 @@ def test_expand_to_p1_carries_validation_flags(self): def test_merge_to_spacegroup_consistent(self): from torchref.io import merge_to_spacegroup - grid = _base_grid(6, 6, 6) + grid = _asu_unique(_base_grid(6, 6, 6)) d = _synthetic(grid, seed=4) d.generate_validation_set(val_fraction_of_free=0.5, seed=0) p1 = d.expand_to_p1() diff --git a/tests/unit/maps/test_map.py b/tests/unit/maps/test_map.py index c074bc54..426278f7 100644 --- a/tests/unit/maps/test_map.py +++ b/tests/unit/maps/test_map.py @@ -157,3 +157,69 @@ def test_difference_map_write(self, model_ft_and_data): assert os.path.getsize(filepath) > 0 finally: os.unlink(filepath) + + +def _merged_and_anomalous(data): + """The acentric reflections of ``data``, once merged and once as Bijvoet + pairs F*(1 +/- eps) whose mean is the merged F. Both keep every row, so the + outlier masks recomputed on construction cannot make them differ.""" + keep = data.masks() & ~data.centric + hkl, F, sigF = data.hkl[keep], data.F[keep], data.F_sigma[keep] + common = dict(cell=data.cell, spacegroup=data.spacegroup, verbose=0) + merged = ReflectionData.from_tensors(hkl, F, sigF, **common) + anom = ReflectionData.from_tensors( + torch.cat([hkl, -hkl]), + torch.cat([F * 1.2, F * 0.8]), + torch.cat([sigF, sigF]), + friedel_merged=False, + **common, + ) + for d in (merged, anom): + d.masks.clear() + d.masks["all"] = torch.ones(len(d.hkl), dtype=torch.bool) + return merged, anom + + +class TestAnomalousInput: + """Bijvoet pairs enter a map once, at their mean amplitude.""" + + def test_bijvoet_helpers(self, model_ft_and_data): + _, data, _ = model_ft_and_data + merged, anom = _merged_and_anomalous(data) + assert anom.friedel_merged is False + + rows = anom.bijvoet_representatives() + assert len(rows) == len(merged.hkl) + mean = anom.bijvoet_mean(anom.F) + by_hkl = dict(zip(map(tuple, merged.hkl.tolist()), merged.F.tolist())) + for h, f in zip(anom.hkl[rows].tolist(), mean[rows].tolist()): + assert f == pytest.approx(by_hkl[tuple(h)], rel=1e-5) + + # Merged data pass through untouched. + assert torch.equal(merged.bijvoet_mean(merged.F), merged.F) + assert torch.equal( + merged.bijvoet_representatives(), + torch.arange(len(merged.hkl), device=merged.device), + ) + + def test_map_from_anomalous_data_equals_merged(self, model_ft_and_data): + model, data, _ = model_ft_and_data + merged, anom = _merged_and_anomalous(data) + grid = Map(merged, model, map_type="2Fo-Fc")._determine_gridsize() + + expected = Map(merged, model, gridsize=grid, map_type="2Fo-Fc").calculate() + result = Map(anom, model, gridsize=grid, map_type="2Fo-Fc").calculate() + + torch.testing.assert_close(result, expected, rtol=1e-4, atol=1e-5) + + def test_difference_map_from_anomalous_data_equals_merged(self, model_ft_and_data): + model, data, _ = model_ft_and_data + merged, anom = _merged_and_anomalous(data) + merged_pert, anom_pert = _merged_and_anomalous(data) + merged_pert.F = merged_pert.F * 1.1 + anom_pert.F = anom_pert.F * 1.1 + + expected = DifferenceMap(merged_pert, merged, model).calculate() + result = DifferenceMap(anom_pert, anom, model).calculate() + + torch.testing.assert_close(result, expected, rtol=1e-4, atol=1e-5) diff --git a/torchref/experimental/targets/realspace.py b/torchref/experimental/targets/realspace.py index 2024b595..8bcbc6d4 100644 --- a/torchref/experimental/targets/realspace.py +++ b/torchref/experimental/targets/realspace.py @@ -128,14 +128,16 @@ def _ensure_p1_expansion(self): if self._hkl_p1 is not None: return sg = self._data.spacegroup or SpaceGroup("P1") + # One row per reflection; anomalous F_obs is Bijvoet-averaged below. + rows = self._data.bijvoet_representatives() hkl_p1, indices, phase_shifts = sg.expand_hkl( - self._data.hkl, + self._data.hkl[rows], include_friedel=True, remove_absences=True, device=self._data.hkl.device, ) self._hkl_p1 = hkl_p1 - self._p1_indices = indices + self._p1_indices = rows[indices] self._p1_phase_shifts = phase_shifts def _expand_to_p1(self, fcalc: torch.Tensor) -> torch.Tensor: @@ -172,7 +174,7 @@ def _compute_observed_map(self) -> torch.Tensor: # Expand Fobs to P1 using the same index mapping as Fcalc # (amplitudes are invariant under symmetry, no phase shift needed) - fobs_p1 = self._data.F[self._p1_indices] + fobs_p1 = self._data.bijvoet_mean(self._data.F)[self._p1_indices] # Compute and scale Fcalc at ASU level, then expand to P1 fcalc_asu = self.get_fcalc_scaled() @@ -591,6 +593,12 @@ def _setup_data(self): if valid_dark is not None: valid_mask = valid_mask & valid_dark + F_light = self._data_light.bijvoet_mean(F_light, valid_mask) + F_dark = self._data_light.bijvoet_mean(F_dark, valid_mask) + # A Bijvoet pair is measured if either mate is (no-op for merged data). + as_float = valid_mask.to(F_light.dtype) + valid_mask = self._data_light.bijvoet_mean(as_float, valid_mask) > 0 + # Zero invalid values to prevent NaN propagation F_light = torch.where(valid_mask, F_light, torch.zeros_like(F_light)) F_dark = torch.where(valid_mask, F_dark, torch.zeros_like(F_dark)) @@ -606,14 +614,16 @@ def _ensure_p1_expansion(self): return spacegroup = self._data_light.spacegroup + # One row per reflection; F_obs was Bijvoet-averaged in _setup_data. + rows = self._data_light.bijvoet_representatives() hkl_p1, indices, phase_shifts = spacegroup.expand_hkl( - self._hkl, + self._hkl[rows], include_friedel=True, remove_absences=True, device=self._hkl.device, ) self._hkl_p1 = hkl_p1 - self._p1_indices = indices + self._p1_indices = rows[indices] self._p1_phase_shifts = phase_shifts def _compute_observed_map(self) -> torch.Tensor: diff --git a/torchref/io/datasets/merging.py b/torchref/io/datasets/merging.py index 84676874..131492dc 100644 --- a/torchref/io/datasets/merging.py +++ b/torchref/io/datasets/merging.py @@ -198,9 +198,7 @@ def merge_to_spacegroup( hkl_src = hkl_src.detach().cpu()[rows] n_src = len(rows) - rotated = source.reciprocal.apply_rotations(hkl_src) - cand = torch.round(rotated).to(hkl_src.dtype).reshape(-1, 3) - cand_src = torch.arange(n_src).repeat(source.n_ops) + cand, cand_src, _, _ = source.equivalent_hkl(hkl_src, include_friedel=False) canon, _, friedel, sort_idx = target.canonicalize_hkl(cand, include_friedel=True) # canonicalize_hkl returns its outputs sorted; the source row must follow. diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index 3306aeaa..cd93d2e2 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -680,6 +680,72 @@ def asu_group_indices(self) -> Tuple[torch.Tensor, int]: uniq, inverse = torch.unique(self.hkl.cpu(), dim=0, return_inverse=True) return inverse.to(self.device), int(uniq.shape[0]) + def bijvoet_mean( + self, values: torch.Tensor, valid: Optional[torch.Tensor] = None + ) -> torch.Tensor: + """Replace each row's value by the mean over the valid rows of its Bijvoet pair. + + For consumers that want one value per reflection -- a Hermitian map + puts each amplitude at ``h`` and its conjugate at ``-h``, so feeding it + both mates would count every measured pair twice. Merged data are + returned unchanged. + + Parameters + ---------- + values : torch.Tensor + Per-row real values of shape (N,), e.g. amplitudes or differences. + valid : torch.Tensor, optional + Boolean (N,), rows allowed to contribute. Defaults to ``masks()``. + + Returns + ------- + torch.Tensor + Shape (N,). Rows of a pair with no valid member keep their own value. + """ + if self.friedel_merged: + return values + if valid is None: + valid = self.masks() + if valid is None: + valid = torch.ones_like(values, dtype=torch.bool) + group_id, n_groups = self.asu_group_indices() + w = valid.to(values.dtype) + total = torch.zeros(n_groups, dtype=values.dtype, device=values.device) + total = total.index_add(0, group_id, torch.where(valid, values, 0.0)) + count = torch.zeros_like(total).index_add(0, group_id, w) + mean = (total / count.clamp(min=1))[group_id] + return torch.where(count[group_id] > 0, mean, values) + + def bijvoet_representatives( + self, valid: Optional[torch.Tensor] = None + ) -> torch.Tensor: + """One row index per unique reflection, in row order. + + Pairs with :meth:`bijvoet_mean` to build a Friedel-averaged reflection + list from anomalous data. For merged data every row is its own + representative. + + Parameters + ---------- + valid : torch.Tensor, optional + Boolean (N,). If given, only reflections with at least one valid row + are represented. + + Returns + ------- + torch.Tensor + Row indices, int64, ascending. + """ + if self.friedel_merged: + if valid is None: + return torch.arange(len(self.hkl), device=self.device) + return torch.nonzero(valid).squeeze(-1) + group_id, n_groups = self.asu_group_indices() + rows = self._group_representative_rows(group_id, n_groups) + if valid is not None: + rows = rows[self._group_any(valid, group_id, n_groups)] + return torch.sort(rows).values + @staticmethod def _group_any( mask: torch.Tensor, group_id: torch.Tensor, n_groups: int @@ -2303,6 +2369,14 @@ def expand_to_p1( all symmetry-equivalent reflections. Returns a NEW ReflectionData object with expanded reflections; does not modify self. + Anomalous data (``friedel_merged`` False) expand from their signed + indices, so ``F(+)`` and ``F(-)`` each keep their own P1 reflections + (``h`` and ``-h``) and the result holds both halves of reciprocal space + whatever ``include_friedel`` says; with it, a Friedel copy fills in only + where a mate was not measured. Consumers that want one value per + reflection pair (a Hermitian map) must merge the mates first, e.g. with + ``merge_to_spacegroup(data, data.spacegroup, anomalous=False)``. + Parameters ---------- include_friedel : bool, default True @@ -2315,31 +2389,46 @@ def expand_to_p1( ------- ReflectionData New object at ``spacegroup="P1"`` holding every symmetry-equivalent - reflection (duplicates removed). Per-reflection fields are indexed - from the original, ``phase`` additionally gets the translation phase - shift, and ``resolution`` is recomputed. - ``source``/``last_op`` record the provenance. + reflection. Per-reflection fields are indexed from the original, + ``phase`` additionally gets the translation phase shift, and + ``resolution`` is recomputed. ``hkl_anomalous`` equals ``hkl``: each + P1 row is its own index. ``source``/``last_op`` record the provenance. + + Raises + ------ + ValueError + If the data hold symmetry-equivalent rows (unmerged observations), + which expansion would otherwise silently drop. """ if self.hkl is None: raise ValueError("ReflectionData has no Miller indices loaded") - # Get expanded HKL set with index mapping and phase shifts + anomalous = not self.friedel_merged and self.hkl_anomalous is not None sg = self.spacegroup or SpaceGroup("P1", device=self.device) hkl_p1, indices, phase_shifts = sg.expand_hkl( - self.hkl, + self.hkl_anomalous if anomalous else self.hkl, include_friedel=include_friedel, remove_absences=remove_absences, device=self.device, ) - # Use remap to create the new dataset - return self.remap( + p1 = self.remap( new_hkl=hkl_p1, index_mapping=indices, - phase_shifts=phase_shifts, spacegroup="P1", op_name=f"expand_to_p1(include_friedel={include_friedel})", ) + if p1.phase is not None: + phase = p1.phase + if anomalous: + # A conjugated mate stores the phase of its canonical index, the + # negative of the phase at its own signed index. + phase = torch.where(self.friedel_flags[indices], -phase, phase) + p1.phase = phase + phase_shifts + p1.hkl_anomalous = p1.hkl.clone() + p1.friedel_flags = torch.zeros_like(p1.hkl[:, 0], dtype=torch.bool) + p1.friedel_merged = self.friedel_merged + return p1 # ========== E-VALUE AND ANISOTROPY CORRECTION METHODS ========== diff --git a/torchref/maps/difference_map.py b/torchref/maps/difference_map.py index 4dc010ea..e7140b53 100644 --- a/torchref/maps/difference_map.py +++ b/torchref/maps/difference_map.py @@ -120,9 +120,14 @@ def calculate(self) -> torch.Tensor: # Combined mask: only use reflections valid in both datasets mask_combined = self.data_reference.masks() & self.data_perturbed.masks() - hkl_asu = self.data_reference.hkl[mask_combined] - fobs_ref = F_ref_scaled[mask_combined] - fobs_pert = F_pert_scaled[mask_combined] + delta_f = F_pert_scaled - F_ref_scaled + if self.scale is not None: + delta_f = delta_f / self.scale.to(delta_f) + # One difference per reflection (Bijvoet mates averaged): the Hermitian + # placement below adds each conjugate at -h itself. + rows = self.data_reference.bijvoet_representatives(mask_combined) + delta_f = self.data_reference.bijvoet_mean(delta_f, mask_combined)[rows] + hkl_asu = self.data_reference.hkl[rows] # Expand to P1 without Friedel mates (expand_to_p1() would reset # scaling, so expand manually via expand_hkl) @@ -134,9 +139,6 @@ def calculate(self) -> torch.Tensor: ) # Map scaled amplitudes to P1 (amplitudes are invariant under symmetry) - delta_f = fobs_pert - fobs_ref - if self.scale is not None: - delta_f = delta_f / self.scale.to(delta_f)[mask_combined] delta_f_p1 = delta_f[orig_idx] # Compute Fcalc for P1 hkl (for phases) diff --git a/torchref/maps/map.py b/torchref/maps/map.py index 204df60c..97bc3788 100644 --- a/torchref/maps/map.py +++ b/torchref/maps/map.py @@ -154,8 +154,18 @@ def calculate(self) -> torch.Tensor: """ # Expand to P1 without Friedel mates (place_on_grid handles # Hermitian symmetry via enforce_hermitian=True) - data_p1 = self.data.expand_to_p1(include_friedel=False) - hkl_p1, fobs_p1, _, _ = data_p1.data_indexed() + if self.data.friedel_merged: + data_p1 = self.data.expand_to_p1(include_friedel=False) + hkl_p1, fobs_p1, _, _ = data_p1.data_indexed() + else: + # One amplitude per reflection: the Hermitian placement would + # otherwise count every measured Bijvoet pair twice. + valid = self.data.masks() + rows = self.data.bijvoet_representatives(valid) + fobs_rows = self.data.bijvoet_mean(self.data.F, valid)[rows] + sg = self.data.spacegroup + hkl_p1, idx, _ = sg.expand_hkl(self.data.hkl[rows], include_friedel=False) + fobs_p1 = fobs_rows[idx] # Compute Fcalc for P1-expanded hkl fcalc_p1 = self.model.get_structure_factor(hkl_p1) diff --git a/torchref/symmetry/__init__.py b/torchref/symmetry/__init__.py index b8e3b862..9676f7b4 100644 --- a/torchref/symmetry/__init__.py +++ b/torchref/symmetry/__init__.py @@ -8,7 +8,8 @@ :class:`SpaceGroup` specialises it with the crystallographic identity (Hermann-Mauguin naming, number, point group, crystal system) and the CCP4 asymmetric-unit verbs -(``expand_hkl``, ``reduce_hkl``, ``complete_hkl``, ``canonicalize_hkl``). It accepts a +(``equivalent_hkl``, ``expand_hkl``, ``reduce_hkl``, ``complete_hkl``, +``canonicalize_hkl``). It accepts a name, a number 1-230, a ``gemmi.SpaceGroup``, another instance, or None for P1. :class:`Cell` is separate: it wraps the six cell parameters, not a symmetry group. diff --git a/torchref/symmetry/reciprocal_symmetry.py b/torchref/symmetry/reciprocal_symmetry.py index 8e23a57f..1af5da8c 100644 --- a/torchref/symmetry/reciprocal_symmetry.py +++ b/torchref/symmetry/reciprocal_symmetry.py @@ -1,7 +1,8 @@ """Asymmetric-unit conventions for Miller indices. The algorithms behind :class:`~torchref.symmetry.spacegroup.SpaceGroup`'s HKL verbs: -``expand_hkl`` (ASU -> P1), ``reduce_hkl`` (P1 -> ASU), ``complete_hkl`` (reflections +``equivalent_hkl`` (every symmetry copy, with its source row), ``expand_hkl`` +(ASU -> P1, built on it), ``reduce_hkl`` (P1 -> ASU), ``complete_hkl`` (reflections missing from a dataset, same space group) and ``canonicalize_hkl`` (CCP4 ASU representative). All private -- call them through the space group, which is the only public entry point. @@ -29,6 +30,71 @@ +def _equivalent_hkl( + sym, + hkl: torch.Tensor, + include_friedel: bool = True, + device: Optional[torch.device] = None, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Every symmetry copy of every input row, without deduplication. + + Copies are ordered by operation, then row: all rows under operation 0, + all under operation 1, ..., then (with ``include_friedel``) the Friedel + copies in the same order. Callers that keep the first copy per index + therefore prefer a real measurement over a Friedel copy. + + Parameters + ---------- + sym : SpaceGroup + The space group whose operations are applied. + hkl : torch.Tensor, shape (N, 3) + Input Miller indices. + include_friedel : bool, default True + Append the Friedel copy ``-h'`` of every rotated index. + device : torch.device, optional + Output device. Defaults to ``hkl``'s. + + Returns + ------- + copies : torch.Tensor, shape (M, 3), dtype=int32 + ``M = n_ops * N``, doubled with ``include_friedel``. + source : torch.Tensor, shape (M,), dtype=int64 + Input row of each copy. + phase_shifts : torch.Tensor, shape (M,) + Translation phase offset in radians of each copy. + is_friedel : torch.Tensor, shape (M,), dtype=bool + True for the Friedel copies. + """ + if device is None: + device = hkl.device + hkl_float = hkl.to(dtype=get_float_dtype(), device=device) + n = len(hkl_float) + + # h' = h @ R^T, one batched matmul for all operations. + matrices = sym.reciprocal.matrices.to(device=device, dtype=hkl_float.dtype) + rotated = torch.einsum("oij,nj->oni", matrices, hkl_float) + copies = torch.round(rotated).to( + torch.int32 # dtype-ok: transformed Miller indices (hkl); fixed-width int32 representation + ) + # Phase shift from translation: -2π h·t, for h' = hR under the convention + # F(h) = Σ_j f_j exp(+2πi h·x_j). Do NOT "simplify" the sign: the wrong sign + # costs 4π h·t mod 2π, which is exactly zero for 2₁ screws and centring, so + # P21/P212121/C2 cannot see it. tests/unit/symmetry/test_phase_convention.py. + translations = sym.translations.to(device=device, dtype=hkl_float.dtype) + phase = -2.0 * np.pi * (hkl_float @ translations.T).T + + copies = copies.reshape(-1, 3) + phase = phase.reshape(-1) + source = torch.arange(n, device=device).repeat(sym.n_ops) + is_friedel = torch.zeros(len(copies), dtype=torch.bool, device=device) + if include_friedel: + copies = torch.cat([copies, -copies]) + phase = torch.cat([phase, -phase]) + source = torch.cat([source, source]) + is_friedel = torch.cat([is_friedel, ~is_friedel]) + return copies, source, phase, is_friedel + + def _expand_hkl( sym, hkl: torch.Tensor, @@ -39,14 +105,17 @@ def _expand_hkl( """Expand Miller indices under crystallographic symmetry (ASU -> P1). The low-level primitive: returns the expanded indices plus the index map and - phase offsets needed to expand any associated per-reflection data. + phase offsets needed to expand any associated per-reflection data. Each P1 + index takes its first copy in :func:`_equivalent_hkl` order, so a rotated + measurement wins over a Friedel copy -- which is what lets signed Bijvoet + rows (``+h`` and ``-h`` as separate rows) expand without colliding. Parameters ---------- sym : SpaceGroup The space group whose asymmetric unit convention applies. hkl : torch.Tensor, shape (N, 3) - Input Miller indices (asymmetric unit). + Input Miller indices, no two of them symmetry-equivalent. include_friedel : bool, default True Include Friedel mates (-h, -k, -l). remove_absences : bool, default True @@ -57,82 +126,54 @@ def _expand_hkl( Returns ------- expanded_hkl : torch.Tensor, shape (M, 3), dtype=int32 - All unique expanded Miller indices. + All unique expanded Miller indices, in order of first occurrence. orig_indices : torch.Tensor, shape (M,), dtype=int64 Index mapping expanded → original: ``F_expanded = F_orig[orig_indices]``. phase_shifts : torch.Tensor, shape (M,), dtype=float32 Translation phase offsets in radians: ``phase_expanded = phase_orig[orig_indices] + phase_shifts``. + + Raises + ------ + ValueError + If two different input rows produce the same P1 index by the same kind + of copy (rotation, or Friedel), i.e. the input holds symmetry-equivalent + rows. Keeping either would silently discard the other: merge them first, + or expand anomalous data from its signed indices. """ if device is None: device = hkl.device + copies, source, phase, is_friedel = _equivalent_hkl( + sym, hkl, include_friedel=include_friedel, device=device + ) - # Get symmetry operations - n_ops = sym.n_ops - recip_matrices = sym.reciprocal.matrices.to(device=device) - translations = sym.translations.to(device=device) - - # Convert hkl to float for matrix operations - hkl_float = hkl.to(dtype=get_float_dtype(), device=device) - n_orig = len(hkl_float) - - # Apply all symmetry operations - all_hkl = [] - all_phases = [] - - for i in range(n_ops): - # h' = h @ R^T - hkl_transformed = torch.round(torch.matmul(hkl_float, recip_matrices[i].T)).to( - torch.int32 # dtype-ok: transformed Miller indices (hkl); fixed-width int32 representation + # On CPU: torch.unique(dim=0) is not reliably supported across accelerator + # backends, and this runs once per expansion on integer data. + copies_cpu = copies.cpu() + source_cpu, friedel_cpu = source.cpu(), is_friedel.cpu() + uniq, inverse = torch.unique(copies_cpu, dim=0, return_inverse=True) + position = torch.arange(len(copies_cpu)) + first = torch.full((len(uniq),), len(copies_cpu), dtype=position.dtype) + first.scatter_reduce_(0, inverse, position, reduce="amin") + + competing = friedel_cpu == friedel_cpu[first][inverse] + clash = competing & (source_cpu != source_cpu[first][inverse]) + if bool(clash.any()): + i = int(torch.nonzero(clash)[0]) + raise ValueError( + f"Input rows {int(source_cpu[first][inverse][i])} and {int(source_cpu[i])} " + f"both expand onto {copies_cpu[i].tolist()}; {int(clash.sum())} such " + "collisions. The input holds symmetry-equivalent rows -- merge them " + "first, or expand anomalous data from its signed indices." ) - # Phase shift from translation: -2π h·t, for h' = hR under the convention - # F(h) = Σ_j f_j exp(+2πi h·x_j). Do NOT "simplify" the sign: the wrong sign - # costs 4π h·t mod 2π, which is exactly zero for 2₁ screws and centring, so - # P21/P212121/C2 cannot see it. tests/unit/symmetry/test_phase_convention.py. - phase_shift = -2.0 * np.pi * torch.matmul(hkl_float, translations[i]) - all_hkl.append(hkl_transformed) - all_phases.append(phase_shift) - - # Add Friedel mates if requested - if include_friedel: - for i in range(n_ops): - all_hkl.append(-all_hkl[i]) - all_phases.append(-all_phases[i]) - - # Stack all transformed hkl and phases - hkl_expanded = torch.cat(all_hkl, dim=0) - phases_expanded = torch.cat(all_phases, dim=0) - - # Remove duplicates - keep unique (h,k,l) tuples with index mapping - hkl_np = hkl_expanded.cpu().numpy() - phase_np = phases_expanded.cpu().numpy() - - # Build dictionary: key=(h,k,l), value=(first_occurrence_idx, phase) - unique_dict = {} - for idx, (h, phase) in enumerate(zip(hkl_np, phase_np)): - key = tuple(h) - if key not in unique_dict: - unique_dict[key] = (idx, phase) - - # Extract unique data - unique_indices = [v[0] for v in unique_dict.values()] - unique_phases = [v[1] for v in unique_dict.values()] - - # Map back to original reflection index - n_total_ops = n_ops * (2 if include_friedel else 1) - orig_indices = [idx % n_orig for idx in unique_indices] - - # Build output tensors - expanded_hkl = torch.tensor( - [list(k) for k in unique_dict.keys()], dtype=torch.int32, device=device # dtype-ok: unique Miller indices (hkl); fixed-width int32 representation - ) - phase_shifts = torch.tensor(unique_phases, dtype=get_float_dtype(), device=device) - orig_idx_tensor = torch.tensor(orig_indices, dtype=torch.int64, device=device) # dtype-ok: reflection index mapping; int64 index tensor required + keep = torch.sort(first).values.to(device) + expanded_hkl = copies[keep] + orig_idx_tensor = source[keep] + phase_shifts = phase[keep] if remove_absences and sym.number != 1: keep_mask = ~sym.is_absent(expanded_hkl) - expanded_hkl = expanded_hkl[keep_mask] phase_shifts = phase_shifts[keep_mask] orig_idx_tensor = orig_idx_tensor[keep_mask] diff --git a/torchref/symmetry/spacegroup.py b/torchref/symmetry/spacegroup.py index c7a9354e..34c48346 100644 --- a/torchref/symmetry/spacegroup.py +++ b/torchref/symmetry/spacegroup.py @@ -337,6 +337,13 @@ def expand_hkl( phase_shifts : torch.Tensor Translation phase offsets in radians, shape ``(M,)``: ``phase_exp = phase_orig[orig_indices] + phase_shifts``. + + Raises + ------ + ValueError + If two input rows are symmetry-equivalent, since one would be + silently dropped. Merge them first; anomalous data expand from their + signed indices (see ``ReflectionData.expand_to_p1``). """ from torchref.symmetry.reciprocal_symmetry import _expand_hkl @@ -348,6 +355,45 @@ def expand_hkl( device=device, ) + def equivalent_hkl( + self, + hkl: torch.Tensor, + include_friedel: bool = True, + device: Optional[torch.device] = None, + ): + """Every symmetry copy of every input row, not deduplicated. + + The primitive behind :meth:`expand_hkl`, for callers that must know + which input row each copy came from even where copies coincide (e.g. + merging observations into another space group). + + Parameters + ---------- + hkl : torch.Tensor + Input Miller indices, shape ``(N, 3)``. + include_friedel : bool, default True + Append the Friedel copy of every rotated index. + device : torch.device, optional + Output device. Defaults to ``hkl``'s. + + Returns + ------- + copies : torch.Tensor + Shape ``(M, 3)``, int32, ordered by operation then row, Friedel copies + last; ``M = n_ops * N``, doubled with ``include_friedel``. + source : torch.Tensor + Input row of each copy, shape ``(M,)``. + phase_shifts : torch.Tensor + Translation phase offset in radians, shape ``(M,)``. + is_friedel : torch.Tensor + Boolean, shape ``(M,)``, True for the Friedel copies. + """ + from torchref.symmetry.reciprocal_symmetry import _equivalent_hkl + + return _equivalent_hkl( + self, hkl, include_friedel=include_friedel, device=device + ) + def reduce_hkl( self, hkl_p1: torch.Tensor, From 0cd72734c1cdb3a8edf10aa11733c47d94752df2 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Tue, 29 Sep 2026 09:07:08 +0200 Subject: [PATCH 219/250] Give the orbit test unique reflections expand_hkl now refuses symmetry-equivalent input rows, and the test drew 150 random indices, which contain some. It checks only that each emitted index lies in its parent's orbit, so deduplicating the input first leaves what it pins unchanged. Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/alignment/test_symmetry_conventions.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/unit/alignment/test_symmetry_conventions.py b/tests/unit/alignment/test_symmetry_conventions.py index 41570390..0244c5f2 100644 --- a/tests/unit/alignment/test_symmetry_conventions.py +++ b/tests/unit/alignment/test_symmetry_conventions.py @@ -224,6 +224,9 @@ def test_symmetry_unroll_stays_within_the_true_orbit(hm, non_orthogonal): sg = SpaceGroup(hm) g = torch.Generator().manual_seed(19) hkl = torch.randint(-9, 10, (150, 3), generator=g) + # One row per unique reflection: expand_hkl refuses symmetry-equivalent rows. + canon, *_ = sg.canonicalize_hkl(hkl, include_friedel=True) + hkl = torch.unique(canon, dim=0) unrolled, asu_idx, _ = sg.expand_hkl(hkl, include_friedel=False) unrolled = unrolled.detach().cpu().to(torch.long) From 25d6833b3f31694b4d72ba110c2dcd3c6bff84b7 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Tue, 29 Sep 2026 12:00:52 +0200 Subject: [PATCH 220/250] Fit the Wilson B with form factors and the protein correction fit_wilson_b fitted ln without dividing by sum f^2, so the form-factor falloff was read as B (7-10 A^2 too high on the deposited test data), and it returned a fixed 200 A^2 whenever the data stopped short of 3.5 A. It now fits xtriage's model, = K sum_f2 (1 + gamma) exp(-B d*^2/2), through equal-count shells: sum_f2 from ITC92 form factors of an average residue, gamma the Zwart & Lamzin (2004) protein correction vendored from cctbx (BSD, notice kept in the module), which straightens the plot to ~11 A so data cut at 3.5 A give B within 10 A^2 of the full-data value. It returns a WilsonFit (B, its standard error, the curve) or None with a warning; there is no fallback and no clamp. The ensemble Wilson prior uses the fitted curve shape. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 2 +- tests/unit/scaling/test_wilson_b.py | 145 +++++++ .../experimental/ensemble/wilson_prior.py | 45 +- torchref/scaling/_protein_gamma.py | 142 +++++++ torchref/scaling/wilson.py | 391 ++++++++++-------- 5 files changed, 529 insertions(+), 196 deletions(-) create mode 100644 tests/unit/scaling/test_wilson_b.py create mode 100644 torchref/scaling/_protein_gamma.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 8eacaf99..d7c7f69b 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -6,7 +6,7 @@ Unreleased ---------- - Fixed ``ReflectionData.expand_to_p1`` on anomalous data keeping only one member of each Bijvoet pair (whichever came first in row order), which dropped about half the reflections when one mate was unmeasured. Anomalous data now expand from their signed indices, so F(+) and F(-) keep their own P1 reflections, and the P1 dataset's ``hkl_anomalous`` is its own index rather than the source row's. ``SpaceGroup.expand_hkl`` raises when two input rows are symmetry-equivalent instead of silently keeping one; ``SpaceGroup.equivalent_hkl`` returns every symmetry copy with its source row - ``Map``, ``DifferenceMap`` and the experimental real-space targets average each Bijvoet pair (``ReflectionData.bijvoet_mean`` / ``bijvoet_representatives``) before placing amplitudes on the grid, so anomalous input gives the Friedel-averaged map; maps from merged data are unchanged -- The Wilson B fit moved out of ``ReflectionData`` into ``torchref.scaling.wilson.fit_wilson_b(F, d)``, which returns the structure B (same values as before). Removed the ``wilson_b``, ``wilson_b_structure``, ``wilson_b_solvent`` and ``wilson_k_sol`` dataset fields; nothing but the ensemble Wilson prior read them +- Replaced the Wilson B estimate, which fitted ln without form factors (7-10 Ų too high) and returned a fixed 200 Ų when the data stopped short of 3.5 Å. ``torchref.scaling.wilson.fit_wilson_b(I, d)`` now fits xtriage's model, `` = K sum_f2 (1 + gamma) exp(-B d*^2/2)`` with the Zwart & Lamzin (2004) protein correction vendored from cctbx, returns a ``WilsonFit`` (B, its standard error, the curve) or ``None`` with a warning, and stays within a few Ų when the high-resolution shells are cut. Expect values 2-16 Ų below ctruncate-style deposited Wilson B. The ensemble Wilson prior uses the fitted curve. Removed the ``wilson_b``, ``wilson_b_structure``, ``wilson_b_solvent`` and ``wilson_k_sol`` dataset fields - ``ReflectionData.regenerate_rfree_flags`` is renamed ``generate_rfree_flags`` (same arguments; with existing flags and ``force=False`` it now warns instead of printing), and it prints only when ``verbose > 0``. Seeded draws are unchanged - ``ReflectionData.get_bins`` no longer stores ``bin_indices`` on the dataset, and ``mean_res_per_bin`` takes the bins it should average over, so a later ``get_bins`` call with other settings (R-free or validation-set generation, a least-squares target) can no longer shift the shells a scaler's per-bin solvent scale was set up on. The ``bin_indices`` field is removed - Removed the ``ReflectionData`` aliases ``get_max_res`` (use ``d_min``), ``get_valid_mask`` (use ``masks()``) and ``cut_res`` (use ``filter_by_resolution(d_min=, d_max=)``, which now prints only when ``verbose > 0``) diff --git a/tests/unit/scaling/test_wilson_b.py b/tests/unit/scaling/test_wilson_b.py new file mode 100644 index 00000000..6a09ae28 --- /dev/null +++ b/tests/unit/scaling/test_wilson_b.py @@ -0,0 +1,145 @@ +"""fit_wilson_b: the Wilson B of the form-factor and protein-corrected Wilson plot. + +Pinned behaviour: B is recovered from intensities drawn from the model itself, +at full and at truncated resolution; on deposited data it is stable when the +high-resolution shells are cut away; it gives up with a warning rather than +returning a stand-in value; and the vendored protein correction evaluates +cctbx's Chebyshev series exactly as scitbx does. +""" + +import io +import contextlib +import math + +import pytest +import torch + +from torchref.io import ReflectionData +from torchref.scaling._protein_gamma import ( + _COEFFS, + D_STAR_SQ_HIGH, + D_STAR_SQ_LOW, + protein_gamma, +) +from torchref.scaling.wilson import fit_wilson_b, sum_f_squared + + +def _clenshaw(x, coeffs, lo=D_STAR_SQ_LOW, hi=D_STAR_SQ_HIGH): + """scitbx ``chebyshev_base::cheb_base_f``, transcribed line by line.""" + t = (x - (lo + hi) * 0.5) / (0.5 * (hi - lo)) + x2, d, dd = 2.0 * t, 0.0, 0.0 + for c in reversed(coeffs[1:]): + d, dd = x2 * d - dd + c, d + return t * d - dd + 0.5 * coeffs[0] + + +@pytest.mark.unit +def test_protein_gamma_matches_the_scitbx_series(): + x = torch.linspace(D_STAR_SQ_LOW, D_STAR_SQ_HIGH, 50, dtype=torch.float64) + expected = torch.tensor([_clenshaw(float(v), _COEFFS) for v in x], dtype=x.dtype) + torch.testing.assert_close(protein_gamma(x), expected, rtol=0, atol=1e-10) + # Held at the end values outside the fitted range. + ends = protein_gamma(torch.tensor([0.001, D_STAR_SQ_LOW, 0.9, D_STAR_SQ_HIGH])) + assert ends[0] == ends[1] and ends[2] == ends[3] + + +def _synthetic_intensities(B, d_min, seed=0, cell=40.0): + """Acentric Wilson intensities in a cubic P1 cell, drawn from the fitted model.""" + n = int(cell / d_min) + 1 + r = torch.arange(-n, n + 1) + hkl = torch.stack(torch.meshgrid(r, r, r, indexing="ij"), -1).reshape(-1, 3) + hkl = hkl[(hkl[:, 2] > 0) | ((hkl[:, 2] == 0) & (hkl[:, 1] > 0))] + d = cell / hkl.to(torch.float64).norm(dim=-1) + d = d[(d >= d_min) & (d < 20.0)] + d_star_sq = d.pow(-2) + mean = sum_f_squared(d_star_sq) * (1 + protein_gamma(d_star_sq)) + mean = mean * torch.exp(-0.5 * B * d_star_sq) + g = torch.Generator().manual_seed(seed) + return mean * torch.empty_like(mean).exponential_(generator=g), d + + +@pytest.mark.unit +@pytest.mark.parametrize("B", [20.0, 50.0, 80.0]) +@pytest.mark.parametrize("d_min, tol", [(1.8, 1.5), (3.5, 5.0)]) +def test_recovers_the_b_of_its_own_model(B, d_min, tol): + I, d = _synthetic_intensities(B, d_min) + fit = fit_wilson_b(I, d) + assert fit is not None + assert fit.B == pytest.approx(B, abs=tol) + assert fit.sigma_B < tol + + +@pytest.mark.unit +def test_expected_intensity_follows_the_data(): + I, d = _synthetic_intensities(40.0, 2.0) + fit = fit_wilson_b(I, d) + order = torch.argsort(d) + for shell in torch.tensor_split(order, 10): + predicted = fit.expected_intensity(d[shell]).mean() + assert float(I[shell].mean() / predicted) == pytest.approx(1.0, abs=0.1) + + +@pytest.mark.unit +def test_amplitudes_use_f_squared_plus_variance(): + I, d = _synthetic_intensities(30.0, 2.0) + F = I.sqrt() + sigma = torch.full_like(F, 1.0) + plain = fit_wilson_b(F, d, amplitudes=True) + with_var = fit_wilson_b(F, d, sigma=sigma, amplitudes=True) + assert plain.B == pytest.approx(30.0, abs=1.5) + # Adding a constant variance flattens the falloff, so B comes out lower. + assert with_var.B < plain.B + + +@pytest.mark.unit +def test_gives_up_with_a_warning_instead_of_guessing(): + I, d = _synthetic_intensities(30.0, 2.0) + with pytest.warns(UserWarning, match="No Wilson B"): + assert fit_wilson_b(I[:30], d[:30]) is None + thin = (d > 3.0) & (d < 3.05) + with pytest.warns(UserWarning, match="too short"): + assert fit_wilson_b(I[thin], d[thin]) is None + + +#: B measured with this model on the deposited test data. The deposited +#: ``B_iso_Wilson_estimate`` values follow the plain ctruncate-style fit instead +#: and sit 2-16 Ų higher; these pin the γ-corrected estimate. +MEASURED = { + "1DAW": 28.6, + "3E98": 49.6, + "3K7M": 29.3, + "4BX9": 60.5, + "5BOV": 14.0, + "6G9X": 66.6, +} + + +def _fit_deposited(name, mtz_dir, d_cut=None): + with contextlib.redirect_stdout(io.StringIO()): + data = ReflectionData(verbose=0).load_mtz(str(mtz_dir / f"{name}.mtz")) + keep = data.masks() + if d_cut is not None: + keep = keep & (data.resolution >= d_cut) + eps = data.spacegroup.epsilon(data.hkl).to(data.F)[keep] + if data.I is not None: + return fit_wilson_b(data.I[keep], data.resolution[keep], epsilon=eps) + return fit_wilson_b( + data.F[keep], + data.resolution[keep], + sigma=data.F_sigma[keep], + epsilon=eps, + amplitudes=True, + ) + + +@pytest.mark.unit +@pytest.mark.parametrize("name", sorted(MEASURED)) +def test_deposited_data(name, mtz_dir): + full = _fit_deposited(name, mtz_dir) + assert full.B == pytest.approx(MEASURED[name], abs=2.0) + # Cutting the data at 3.5 Å moves B by at most 9 Ų on these six (6G9X, + # the largest), where the uncorrected fit fell back to a fixed 200 Ų. + cut = _fit_deposited(name, mtz_dir, d_cut=3.5) + assert cut is not None + assert cut.B == pytest.approx(full.B, abs=10.0) + assert math.isfinite(cut.sigma_B) diff --git a/torchref/experimental/ensemble/wilson_prior.py b/torchref/experimental/ensemble/wilson_prior.py index 435922eb..d96332aa 100644 --- a/torchref/experimental/ensemble/wilson_prior.py +++ b/torchref/experimental/ensemble/wilson_prior.py @@ -10,9 +10,11 @@ expected per-resolution-bin mean intensity of a randomly-placed atomic ensemble follows:: - <|F|^2>(s) = K * exp(-2 * B_W * s^2) + <|F|^2>(d) = K * sum_f2(d) * (1 + gamma(d)) * exp(-B_W / (2 d^2)) -where ``s = 1/(2*d)`` and ``B_W`` is the overall Wilson B-factor. Real +where ``sum_f2`` is the random-atom mean intensity of an average protein +residue, ``gamma`` the empirical protein correction and ``B_W`` the overall +Wilson B-factor (see :class:`torchref.scaling.wilson.WilsonFit`). Real calculated intensities should track this curve at low-to-mid resolution. A model that drives the work-set R-factor toward zero by absorbing noise into extra structural detail (e.g. a B-factor-free ensemble of many @@ -25,9 +27,9 @@ loss = mean_bin( ( log<|F_calc|^2>_bin - log Wilson_expected(s_bin) )^2 ) -The reference curve is fit once from the observed data: -``B_W`` from :func:`torchref.scaling.wilson.fit_wilson_b` and ``K`` from a single -least-squares fit at first ``forward()`` call. +The reference curve is fit once from the observed data, at the first +``forward()`` call: its shape from :func:`torchref.scaling.wilson.fit_wilson_b`, +and ``K`` from a least-squares fit to the observed bin intensities. Used as ``'regularization/wilson'`` in the ensemble refinement LossState. """ @@ -131,7 +133,7 @@ def __init__( self.nbins = int(nbins) # ``log_K`` is fit lazily on first forward from observed bin intensities. self._log_K: Optional[torch.Tensor] = None - self._B_W: Optional[torch.Tensor] = None + self._wilson_fit = None # WilsonFit, fitted with log_K # Cached resolution-bin assignment for the work-set reflections # (filled on first forward). self._bin_idx: Optional[torch.Tensor] = None @@ -145,10 +147,7 @@ def __init__( def _wilson_curve(self, mean_res: torch.Tensor) -> torch.Tensor: """Expected ``<|F|^2>`` per bin from the Wilson model.""" - # s = 1/(2d) -> s^2 = 1/(4 d^2) - s_sq = 1.0 / (4.0 * mean_res.clamp(min=1e-3) ** 2) - # <|F|^2> = K * exp(-2 * B_W * s^2) - return torch.exp(self._log_K - 2.0 * self._B_W * s_sq) + return torch.exp(self._log_K) * self._wilson_fit.shape(mean_res.clamp(min=1e-3)) def _build_bin_assignment(self) -> None: """ @@ -194,11 +193,23 @@ def _fit_K_from_observed(self) -> None: """ from torchref.scaling.wilson import fit_wilson_b - wilson_b = fit_wilson_b(self._data.F, self._data.resolution) - if wilson_b is None: + data = self._data + valid = data.masks() + epsilon = data.spacegroup.epsilon(data.hkl).to(data.F)[valid] + if data.I is not None and data.I_sigma is not None: + fit = fit_wilson_b(data.I[valid], data.resolution[valid], epsilon=epsilon) + else: + fit = fit_wilson_b( + data.F[valid], + data.resolution[valid], + sigma=data.F_sigma[valid], + epsilon=epsilon, + amplitudes=True, + ) + if fit is None: raise ValueError("Too few reflections to fit a Wilson B for the prior.") - device = self._data.device - self._B_W = torch.tensor(float(wilson_b), device=device) + self._wilson_fit = fit + device = data.device F_obs = self._data.F.index_select(0, self._refl_subset_idx) I_obs = F_obs.float() ** 2 @@ -209,9 +220,9 @@ def _fit_K_from_observed(self) -> None: counts.scatter_add_(0, self._bin_idx, torch.ones_like(I_obs)) mean_obs = mean_obs / counts.clamp(min=1.0) - s_sq = 1.0 / (4.0 * self._mean_res.clamp(min=1e-3) ** 2) # Solve log K from each bin and average for a robust estimate. - log_K_per_bin = torch.log(mean_obs.clamp(min=self.eps)) + 2.0 * self._B_W * s_sq + shape = fit.shape(self._mean_res.clamp(min=1e-3)).to(mean_obs) + log_K_per_bin = torch.log(mean_obs.clamp(min=self.eps)) - torch.log(shape) self._log_K = log_K_per_bin.mean().detach() # ------------------------------------------------------------------ @@ -230,7 +241,7 @@ def forward(self, fcalc: torch.Tensor = None) -> torch.Tensor: if self._bin_idx is None: self._build_bin_assignment() - if self._log_K is None or self._B_W is None: + if self._log_K is None or self._wilson_fit is None: self._fit_K_from_observed() # Apply scaler to F_calc so it sits on the F_obs scale, then keep diff --git a/torchref/scaling/_protein_gamma.py b/torchref/scaling/_protein_gamma.py new file mode 100644 index 00000000..f0efda42 --- /dev/null +++ b/torchref/scaling/_protein_gamma.py @@ -0,0 +1,142 @@ +""" +Empirical protein correction to the Wilson curve, vendored from cctbx. + +Mean intensities of real protein crystals deviate from the ideal random-atom +curve ``sum f^2 exp(-B d*^2 / 2)`` by a resolution-dependent factor +``1 + gamma(d*^2)`` -- the dip near 6 Å and the secondary-structure bump near +4.5 Å. Dividing it out makes the Wilson plot linear from about 11 Å to 1.2 Å, +which is what lets :func:`torchref.scaling.wilson.fit_wilson_b` fit low +resolution data at all. + +``gamma`` was obtained from experimental data by Zwart & Lamzin, Acta Cryst. +(2004) D60, 220-226. The coefficients below are the 45-term Chebyshev fit +``coefs_mean`` of ``gamma_protein`` in cctbx ``mmtbx/scaling/absolute_scaling.py``, +copied verbatim, with its range 0.008 <= d*^2 <= 0.69 Å^-2 from +``mmtbx/scaling/scaling.h``. They are used under the cctbx licence, whose +notice redistribution must retain: + + cctbx Copyright (c) 2006 - 2026, The Regents of the University of + California, through Lawrence Berkeley National Laboratory (subject to + receipt of any required approvals from the U.S. Dept. of Energy). All + rights reserved. + + Redistribution and use in source and binary forms, with or without + modification, are permitted provided that the following conditions are met: + + (1) Redistributions of source code must retain the above copyright + notice, this list of conditions and the following disclaimer. + + (2) Redistributions in binary form must reproduce the above copyright + notice, this list of conditions and the following disclaimer in the + documentation and/or other materials provided with the distribution. + + (3) Neither the name of the University of California, Lawrence Berkeley + National Laboratory, U.S. Dept. of Energy nor the names of its + contributors may be used to endorse or promote products derived from + this software without specific prior written permission. + + THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS + IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED + TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A + PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT OWNER + OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, + EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, + PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR + PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF + LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING + NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS + SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + + You are under no obligation whatsoever to provide any bug fixes, + patches, or upgrades to the features, functionality or performance of + the source code ("Enhancements") to anyone; however, if you choose to + make your Enhancements available either publicly, or directly to + Lawrence Berkeley National Laboratory, without imposing a separate + written license agreement for such Enhancements, then you hereby grant + the following license: a non-exclusive, royalty-free perpetual license + to install, use, modify, prepare derivative works, incorporate into + other computer software, distribute, and sublicense such enhancements or + derivative works thereof, in binary and source code form. +""" + +import torch + +from torchref.scaling.basis import chebyshev_design + +__all__ = ["protein_gamma", "D_STAR_SQ_LOW", "D_STAR_SQ_HIGH"] + +#: Range of the fit, in Å^-2. Outside it the curve is held at its end values. +D_STAR_SQ_LOW = 0.008 +D_STAR_SQ_HIGH = 0.69 + +_COEFFS = ( + -0.24994838652402987, + 0.15287426147680838, + 0.068108692925184011, + 0.15780196907582875, + -0.07811375753346686, + 0.043211175909300889, + -0.043407219965134192, + 0.024613271516995903, + 0.0035146404613345932, + -0.064118486637211411, + 0.10521875419321854, + -0.10153928782775833, + 0.0335706778430487, + -0.0066629477818811282, + -0.0058221659481290031, + 0.0136026246654981, + -0.013385834361135244, + 0.022526368996167032, + -0.019843844247892727, + 0.018128145323325774, + -0.0091740188657759101, + 0.0068283902389141915, + -0.0060880807366142566, + 0.0004002124110802677, + -0.00065686973991185187, + -0.0039358839200389316, + 0.0056185833386634149, + -0.0075257168326962913, + -0.0015215201587884459, + -0.0036383549957990221, + -0.0064289154284325831, + 0.0059080442658917334, + -0.0089851215734611887, + 0.0036488156067441039, + -0.0047375008148055706, + -0.00090999496111171302, + 0.00096986728652170276, + -0.0051006830761911011, + 0.0046838536228956777, + -0.0031683076118337885, + 0.0037866523617167236, + 0.0015810274077361975, + 0.0011030841357086191, + 0.0015715596895281762, + -0.0041354783162507788, +) + + +def protein_gamma(d_star_sq: torch.Tensor) -> torch.Tensor: + """Fractional deviation of mean protein intensity from the Wilson curve. + + Parameters + ---------- + d_star_sq : torch.Tensor + ``1/d^2`` in Å^-2, any shape. Values outside + [:data:`D_STAR_SQ_LOW`, :data:`D_STAR_SQ_HIGH`] are clamped to the range, + as cctbx does. + + Returns + ------- + torch.Tensor + ``gamma``, same shape and dtype as ``d_star_sq``; the mean intensity is + ``(1 + gamma)`` times the random-atom value. + """ + x = d_star_sq.reshape(-1).clamp(D_STAR_SQ_LOW, D_STAR_SQ_HIGH) + design = chebyshev_design(x, len(_COEFFS), D_STAR_SQ_LOW, D_STAR_SQ_HIGH) + coeffs = torch.tensor(_COEFFS, dtype=x.dtype, device=x.device) + # scitbx's Chebyshev series halves the zeroth coefficient. + coeffs[0] = 0.5 * coeffs[0] + return (design @ coeffs).reshape(d_star_sq.shape) diff --git a/torchref/scaling/wilson.py b/torchref/scaling/wilson.py index 175eb533..0c35439a 100644 --- a/torchref/scaling/wilson.py +++ b/torchref/scaling/wilson.py @@ -31,14 +31,24 @@ from __future__ import annotations -from typing import Optional, Tuple +import math +import warnings +from dataclasses import dataclass +from typing import Dict, Optional, Tuple import torch from torchref.config import get_float_dtype +from torchref.scaling._protein_gamma import protein_gamma as _protein_gamma from torchref.scaling.basis import chebyshev_design -__all__ = ["WilsonNormaliser", "fit_wilson_b"] +__all__ = [ + "WilsonNormaliser", + "WilsonFit", + "fit_wilson_b", + "sum_f_squared", + "PROTEIN_RESIDUE", +] #: Chebyshev terms. Enough to follow a Wilson plot's curvature and the #: low-resolution solvent deficit without chasing shell-to-shell noise. @@ -484,202 +494,227 @@ def __repr__(self) -> str: # pragma: no cover - display ) -#: Fallback structure B (Ų) when the high-resolution fit has fewer than three -#: usable shells. This is the value the fit has always returned there. -_FALLBACK_B = 200.0 +#: Average protein residue, as xtriage assumes when no composition is given. +#: Only the relative amounts matter: the absolute scale goes into ``K``. +PROTEIN_RESIDUE = {"H": 8.0, "C": 5.0, "N": 1.5, "O": 1.2} +#: Without the protein correction the plot is only linear at high resolution; +#: fit reflections from this d-spacing (Å) outward only. +_PLAIN_D_MAX = 4.5 -def fit_wilson_b( - F: torch.Tensor, - d: torch.Tensor, - n_bins: int = 30, - verbose: int = 0, -) -> Optional[float]: - """Fit an overall Wilson B from amplitudes and resolution. - Bins ``F^2`` in ``s^2 = 1/(4 d^2)``, seeds a structure B from the shells - below 3.5 Å and a solvent B from those above 6 Å, then refines the - two-component curve `` = A [(1-k) exp(-2 B s^2) + k exp(-2 B_sol s^2)]`` - and returns its structure B -- what "the Wilson B" usually means. This is a - single number for priors and reports; for normalising intensities use - :class:`WilsonNormaliser`. +def sum_f_squared( + d_star_sq: torch.Tensor, composition: Optional[Dict[str, float]] = None +) -> torch.Tensor: + """``sum_j n_j f_j(d*^2)^2`` for a composition, with ITC92 form factors. Parameters ---------- - F : torch.Tensor - Amplitudes of shape (N,). Non-finite and non-positive values are ignored. - d : torch.Tensor - Resolution of each reflection in Å, shape (N,). - n_bins : int, optional - Number of equal-width ``s^2`` bins. Default 30. - verbose : int, optional - Print the result when > 0. + d_star_sq : torch.Tensor + ``1/d^2`` in Å^-2, shape (N,). + composition : dict, optional + Element symbol to count. Defaults to :data:`PROTEIN_RESIDUE`. Returns ------- - float or None - Structure B in Ų, within 1-200. ``None`` with fewer than 100 usable - reflections or fewer than 5 bins holding more than 5 each. + torch.Tensor + Shape (N,), in electrons squared, dtype of ``d_star_sq``. """ - device = F.device - valid = torch.isfinite(F) & (F > 0) & torch.isfinite(d) - if valid.sum() < 100: - if verbose > 0: - print(f" Wilson B: too few reflections ({int(valid.sum())}), skipping") - return None - - s_sq = 1.0 / (4.0 * d[valid] ** 2) - F_sq = F[valid] ** 2 - - bin_edges = torch.linspace(s_sq.min(), s_sq.max(), n_bins + 1, device=device) - bin_centers = (bin_edges[:-1] + bin_edges[1:]) / 2 - bin_idx = torch.bucketize(s_sq, bin_edges[1:-1]) - bin_sums = torch.zeros(n_bins, device=device, dtype=F_sq.dtype) - bin_counts = torch.zeros(n_bins, device=device, dtype=F_sq.dtype) - bin_sums.scatter_add_(0, bin_idx, F_sq) - bin_counts.scatter_add_(0, bin_idx, torch.ones_like(F_sq)) - - valid_bins = bin_counts > 5 - if valid_bins.sum() < 5: - if verbose > 0: - print(f" Wilson B: insufficient bins ({valid_bins.sum()}), skipping") - return None - - mean_F_sq = bin_sums[valid_bins] / bin_counts[valid_bins] - s_sq_bins = bin_centers[valid_bins] - d_bins = 1.0 / (2.0 * torch.sqrt(s_sq_bins)) - - B_struct = _fit_single_wilson(s_sq_bins, mean_F_sq, d_bins < 3.5) - B_sol = _fit_single_wilson(s_sq_bins, mean_F_sq, d_bins > 6.0) - B, _, _ = _fit_two_component_wilson(s_sq_bins, mean_F_sq, B_struct, B_sol) + from torchref.base.scattering.scattering_table import ( + elements_to_z, + get_scattering_params_by_z, + ) - if verbose > 0: - print(f" Wilson B-factor (structure): {B:.1f} Ų") - return B - - -def _fit_single_wilson( - s_sq: torch.Tensor, mean_F_sq: torch.Tensor, mask: torch.Tensor -) -> float: - """Fit ``ln(F^2) = c - 2 B s^2`` over the masked shells, clamped to 0-300 Ų.""" - if mask.sum() < 3: - return _FALLBACK_B - x = s_sq[mask] - y = torch.log(mean_F_sq[mask]) - x_c, y_c = x - x.mean(), y - y.mean() - denominator = (x_c**2).sum() - if denominator < 1e-12: - return _FALLBACK_B - B = -((x_c * y_c).sum() / denominator).item() / 2.0 - return max(0.0, min(B, 300.0)) - - -def _fit_two_component_wilson( - s_sq: torch.Tensor, - mean_F_sq: torch.Tensor, - B_struct_init: float, - B_sol_init: float, - n_iter: int = 50, -) -> Tuple[float, float, float]: - """Refine ``F^2 = A [(1-k) exp(-2 B_s s^2) + k exp(-2 B_sol s^2)]``. - - Finite-difference gradient descent from the single-component estimates. - - Returns ``(B_struct, B_sol, k_sol)``, constrained to B_s in 1-200, B_sol in - 50-500 with ``B_sol >= B_struct + 20``, and k in 0.01-0.9 -- a value on a - bound means the fit hit the clamp. - """ - device = s_sq.device - F_sq_max = mean_F_sq.max() - y = mean_F_sq / F_sq_max - x = s_sq - - # Initialize parameters - B_struct = torch.tensor(B_struct_init, device=device, dtype=x.dtype) - B_sol = torch.tensor(B_sol_init, device=device, dtype=x.dtype) - - # Estimate initial k from ratio of low-res to high-res decay - # At low resolution, solvent contributes more - d_from_s = 1.0 / (2.0 * torch.sqrt(x)) - low_res_val = y[d_from_s > 5.0].mean() if (d_from_s > 5.0).any() else y[0] - high_res_val = y[d_from_s < 3.0].mean() if (d_from_s < 3.0).any() else y[-1] - - # k estimates solvent fraction - if low res is much higher than expected - # from structure alone, there's solvent contribution - struct_decay = torch.exp(-2 * B_struct * x) - expected_low = ( - struct_decay[d_from_s > 5.0].mean() - if (d_from_s > 5.0).any() - else struct_decay[0] + composition = PROTEIN_RESIDUE if composition is None else composition + elements = list(composition) + counts = torch.tensor([composition[e] for e in elements]).to(d_star_sq) + A, B = get_scattering_params_by_z( + elements_to_z(elements).to(d_star_sq.device), dtype=d_star_sq.dtype ) + # f(s) = sum_k A_k exp(-B_k d*^2 / 4), ITC92 in sin^2(theta)/lambda^2. + f = (A[None] * torch.exp(-B[None] * (d_star_sq[:, None, None] / 4.0))).sum(-1) + return (f**2 * counts[None]).sum(-1) - if expected_low > 1e-6 and low_res_val > expected_low: - k_init = min(0.5, (low_res_val - expected_low).item() / low_res_val.item()) - else: - k_init = 0.1 - k = torch.tensor(max(0.01, min(0.5, k_init)), device=device, dtype=x.dtype) +@dataclass +class WilsonFit: + """Result of :func:`fit_wilson_b`. - # Simple gradient descent refinement - lr = 0.1 + The model is `` = K * sum_f2(d*^2) * (1 + gamma(d*^2)) * exp(-B d*^2 / 2)``, + with ``gamma`` the empirical protein correction (zero when + ``protein_gamma`` is False). - for _ in range(n_iter): - # Compute model - struct_term = (1 - k) * torch.exp(-2 * B_struct * x) - sol_term = k * torch.exp(-2 * B_sol * x) - model = struct_term + sol_term + Attributes + ---------- + B : float + Wilson B in Ų. + sigma_B : float + Standard error of ``B`` from the scatter of the shell means about the + line; it reflects how straight the plot is, not the measurement error. + log_scale : float + ``ln K``. + d_max, d_min : float + Resolution range of the fitted reflections, Å. + n_reflections, n_shells : int + Reflections and equal-count shells in the fit. + composition : dict + Composition used for ``sum_f2``. + protein_gamma : bool + Whether the protein correction was applied. + """ - # Compute scale factor analytically - A = (y * model).sum() / (model * model).sum() - model_scaled = A * model + B: float + sigma_B: float + log_scale: float + d_max: float + d_min: float + n_reflections: int + n_shells: int + composition: Dict[str, float] + protein_gamma: bool + + def shape(self, d: torch.Tensor) -> torch.Tensor: + """``sum_f2 * (1 + gamma) * exp(-B d*^2 / 2)`` at resolution ``d`` (Å); no K.""" + d_star_sq = d.pow(-2) + curve = sum_f_squared(d_star_sq, self.composition) + if self.protein_gamma: + curve = curve * (1.0 + _protein_gamma(d_star_sq)) + return curve * torch.exp(-0.5 * self.B * d_star_sq) + + def expected_intensity( + self, d: torch.Tensor, epsilon: Optional[torch.Tensor] = None + ) -> torch.Tensor: + """Expected mean intensity at resolution ``d`` (Å), times ``epsilon``.""" + out = math.exp(self.log_scale) * self.shape(d) + return out if epsilon is None else out * epsilon.to(out) - # Compute gradients (simplified, using finite differences for robustness) - eps = 0.1 - # B_struct gradient - model_plus = A * ((1 - k) * torch.exp(-2 * (B_struct + eps) * x) + sol_term) - model_minus = A * ( - (1 - k) * torch.exp(-2 * (B_struct - eps) * x) + sol_term - ) - loss_plus = ((y - model_plus) ** 2).sum() - loss_minus = ((y - model_minus) ** 2).sum() - grad_B_struct = (loss_plus - loss_minus) / (2 * eps) - - # B_sol gradient - model_plus = A * (struct_term + k * torch.exp(-2 * (B_sol + eps) * x)) - model_minus = A * (struct_term + k * torch.exp(-2 * (B_sol - eps) * x)) - loss_plus = ((y - model_plus) ** 2).sum() - loss_minus = ((y - model_minus) ** 2).sum() - grad_B_sol = (loss_plus - loss_minus) / (2 * eps) - - # k gradient - eps_k = 0.01 - k_plus = min(0.9, k + eps_k) - k_minus = max(0.01, k - eps_k) - model_plus = A * ( - (1 - k_plus) * torch.exp(-2 * B_struct * x) - + k_plus * torch.exp(-2 * B_sol * x) - ) - model_minus = A * ( - (1 - k_minus) * torch.exp(-2 * B_struct * x) - + k_minus * torch.exp(-2 * B_sol * x) - ) - loss_plus = ((y - model_plus) ** 2).sum() - loss_minus = ((y - model_minus) ** 2).sum() - grad_k = (loss_plus - loss_minus) / (2 * eps_k) +def fit_wilson_b( + I: torch.Tensor, + d: torch.Tensor, + *, + sigma: Optional[torch.Tensor] = None, + epsilon: Optional[torch.Tensor] = None, + amplitudes: bool = False, + composition: Optional[Dict[str, float]] = None, + protein_gamma: bool = True, + n_shells: int = 20, + min_per_shell: int = 20, + verbose: int = 0, +) -> Optional[WilsonFit]: + """Fit the overall Wilson B, as xtriage does, by a line through shell means. + + Fits ``ln `` against ``d*^2`` in shells of + equal reflection count: the slope is ``-B/2``. ``sum_f2`` is the random-atom + mean intensity of ``composition``, and ``gamma`` the empirical protein + correction of Zwart & Lamzin (2004), which straightens the plot down to + about 11 Å, so data at 3.5-4 Å still give a usable B. Without either + correction the fit reads the form-factor falloff as B and comes out 7-10 Ų + too high on protein data. The correction also lowers B relative to the plain + ``d <= 4.5`` Å fit that ctruncate reports and most deposited + ``B_iso_Wilson_estimate`` values follow -- by 10-15 Ų on some datasets -- + so compare like with like. A single number for priors and reports; for + normalising intensities use :class:`WilsonNormaliser`. - # Update parameters - B_struct = B_struct - lr * grad_B_struct - B_sol = B_sol - lr * grad_B_sol - k = k - lr * 0.1 * grad_k # Slower learning rate for k + Parameters + ---------- + I : torch.Tensor + Intensities of shape (N,), or amplitudes with ``amplitudes=True``. + Pass only the reflections to use (e.g. ``data.masks()`` applied). + d : torch.Tensor + Resolution of each reflection in Å, shape (N,). + sigma : torch.Tensor, optional + Uncertainties of ``I`` (N,). Used only with ``amplitudes=True``, where + the intensity proxy is ``F^2 + sigma_F^2``: a French-Wilson ``F^2`` alone + underestimates the mean intensity of weak reflections. + epsilon : torch.Tensor, optional + Reflection multiplicity (N,) from ``SpaceGroup.epsilon``; intensities + are divided by it. + amplitudes : bool, optional + ``I`` holds amplitudes. + composition : dict, optional + Element symbol to count; defaults to :data:`PROTEIN_RESIDUE`. The + result barely depends on it for proteins. + protein_gamma : bool, optional + Apply the protein correction (default True). Without it only + reflections with d <= 4.5 Å are fitted; turn it off for nucleic acids. + n_shells : int, optional + Number of equal-count shells, default 20, reduced so each holds at + least ``min_per_shell`` reflections. + min_per_shell : int, optional + Default 20. + verbose : int, optional + Print the result when > 0. - # Enforce constraints - B_struct = torch.clamp(B_struct, 1.0, 200.0) - B_sol = torch.clamp(B_sol, 50.0, 500.0) - k = torch.clamp(k, 0.01, 0.9) + Returns + ------- + WilsonFit or None + ``None``, with a warning saying why, when the data cannot support a + fit: fewer than three shells, a ``d*^2`` range under 0.03 Å^-2, or a + non-positive shell mean. There is no fallback value and ``B`` is not + clamped -- an implausible B means implausible data. + """ + from torchref.scaling._protein_gamma import D_STAR_SQ_HIGH, D_STAR_SQ_LOW + + I = I.detach().reshape(-1) + d = d.detach().to(I).reshape(-1) + if amplitudes: + y = I**2 if sigma is None else I**2 + sigma.detach().to(I) ** 2 + else: + y = I + if epsilon is not None: + y = y / epsilon.detach().to(I) + d_star_sq = d.pow(-2) + keep = torch.isfinite(y) & torch.isfinite(d_star_sq) + if protein_gamma: + keep &= (d_star_sq > D_STAR_SQ_LOW) & (d_star_sq < D_STAR_SQ_HIGH) + else: + keep &= d <= _PLAIN_D_MAX + y, d_star_sq = y[keep], d_star_sq[keep] - # Ensure B_sol > B_struct (solvent is more disordered) - if B_sol < B_struct + 20: - B_sol = B_struct + 20 + def _give_up(reason: str) -> None: + warnings.warn(f"No Wilson B: {reason}.", stacklevel=3) + return None - return B_struct.item(), B_sol.item(), k.item() + n_shells = min(n_shells, len(y) // min_per_shell) + if n_shells < 3: + return _give_up(f"{len(y)} usable reflections, need {3 * min_per_shell}") + span = float(d_star_sq.max() - d_star_sq.min()) + if span < 0.03: + return _give_up(f"d*^2 range {span:.3f} Å^-2 is too short for a slope") + + y = y / sum_f_squared(d_star_sq, composition) + if protein_gamma: + y = y / (1.0 + _protein_gamma(d_star_sq)) + + order = torch.argsort(d_star_sq) + shells = torch.tensor_split(order, n_shells) + x = torch.stack([d_star_sq[s].mean() for s in shells]) + mean_y = torch.stack([y[s].mean() for s in shells]) + if bool((mean_y <= 0).any()): + return _give_up("a resolution shell has non-positive mean intensity") + ln_y = torch.log(mean_y) + + # Centred regression keeps the float32 normal equations well conditioned. + x_c, y_c = x - x.mean(), ln_y - ln_y.mean() + sxx = (x_c**2).sum() + slope = (x_c * y_c).sum() / sxx + resid = y_c - slope * x_c + sigma_slope = torch.sqrt((resid**2).sum() / (n_shells - 2) / sxx) + fit = WilsonFit( + B=float(-2.0 * slope), + sigma_B=float(2.0 * sigma_slope), + log_scale=float(ln_y.mean() - slope * x.mean()), + d_max=float(d_star_sq.min().rsqrt()), + d_min=float(d_star_sq.max().rsqrt()), + n_reflections=len(y), + n_shells=n_shells, + composition=dict(PROTEIN_RESIDUE if composition is None else composition), + protein_gamma=protein_gamma, + ) + if verbose > 0: + print( + f" Wilson B: {fit.B:.1f} ± {fit.sigma_B:.1f} Ų " + f"({fit.d_max:.2f}-{fit.d_min:.2f} Å, {fit.n_reflections} reflections)" + ) + return fit From 2d1a95fbcee4a372c56b7fbf7e52e660b9a85164 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Tue, 29 Sep 2026 13:40:09 +0200 Subject: [PATCH 221/250] Pass the scaler's bins to mean_res_per_bin in get_binwise_mean_intensity mean_res_per_bin takes the bins to average over, but ScalerBase.get_binwise_mean_intensity still called it without them and raised TypeError; so did the integration test of the method. Both now pass the bins they own, and a unit test calls the scaler method directly and checks that re-binning the dataset leaves the scaler's shells alone. Also, from review: the Wilson fit's docstring keeps the model and its scope, with the comparison to ctruncate-style values moved to the scaling user guide, and the new changelog entries are shorter. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 4 +-- docs/user_guide/scaling.rst | 28 ++++++++++++++++++ tests/integration/test_io_reflections.py | 2 +- tests/unit/scaling/test_scaler.py | 36 ++++++++++++++++++++++++ torchref/scaling/scaler_base.py | 3 +- torchref/scaling/wilson.py | 14 ++++----- 6 files changed, 74 insertions(+), 13 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index d7c7f69b..e0c5eabc 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,9 +4,9 @@ Changelog Unreleased ---------- -- Fixed ``ReflectionData.expand_to_p1`` on anomalous data keeping only one member of each Bijvoet pair (whichever came first in row order), which dropped about half the reflections when one mate was unmeasured. Anomalous data now expand from their signed indices, so F(+) and F(-) keep their own P1 reflections, and the P1 dataset's ``hkl_anomalous`` is its own index rather than the source row's. ``SpaceGroup.expand_hkl`` raises when two input rows are symmetry-equivalent instead of silently keeping one; ``SpaceGroup.equivalent_hkl`` returns every symmetry copy with its source row +- Fixed ``ReflectionData.expand_to_p1`` dropping one Bijvoet mate per reflection on anomalous data (about half the reflections when one mate was unmeasured). ``SpaceGroup.expand_hkl`` now raises on symmetry-equivalent input rows instead of silently keeping one; ``SpaceGroup.equivalent_hkl`` returns every symmetry copy with its source row - ``Map``, ``DifferenceMap`` and the experimental real-space targets average each Bijvoet pair (``ReflectionData.bijvoet_mean`` / ``bijvoet_representatives``) before placing amplitudes on the grid, so anomalous input gives the Friedel-averaged map; maps from merged data are unchanged -- Replaced the Wilson B estimate, which fitted ln without form factors (7-10 Ų too high) and returned a fixed 200 Ų when the data stopped short of 3.5 Å. ``torchref.scaling.wilson.fit_wilson_b(I, d)`` now fits xtriage's model, `` = K sum_f2 (1 + gamma) exp(-B d*^2/2)`` with the Zwart & Lamzin (2004) protein correction vendored from cctbx, returns a ``WilsonFit`` (B, its standard error, the curve) or ``None`` with a warning, and stays within a few Ų when the high-resolution shells are cut. Expect values 2-16 Ų below ctruncate-style deposited Wilson B. The ensemble Wilson prior uses the fitted curve. Removed the ``wilson_b``, ``wilson_b_structure``, ``wilson_b_solvent`` and ``wilson_k_sol`` dataset fields +- ``torchref.scaling.wilson.fit_wilson_b`` now fits the form-factor and protein-corrected Wilson model and returns a ``WilsonFit`` or ``None`` with a warning; the previous fit was 7-10 Ų too high and returned a fixed 200 Ų on data short of 3.5 Å. Values run below ctruncate-style Wilson B (see the scaling user guide). The ensemble Wilson prior uses the fitted curve. Removed the ``wilson_b``, ``wilson_b_structure``, ``wilson_b_solvent`` and ``wilson_k_sol`` dataset fields - ``ReflectionData.regenerate_rfree_flags`` is renamed ``generate_rfree_flags`` (same arguments; with existing flags and ``force=False`` it now warns instead of printing), and it prints only when ``verbose > 0``. Seeded draws are unchanged - ``ReflectionData.get_bins`` no longer stores ``bin_indices`` on the dataset, and ``mean_res_per_bin`` takes the bins it should average over, so a later ``get_bins`` call with other settings (R-free or validation-set generation, a least-squares target) can no longer shift the shells a scaler's per-bin solvent scale was set up on. The ``bin_indices`` field is removed - Removed the ``ReflectionData`` aliases ``get_max_res`` (use ``d_min``), ``get_valid_mask`` (use ``masks()``) and ``cut_res`` (use ``filter_by_resolution(d_min=, d_max=)``, which now prints only when ``verbose > 0``) diff --git a/docs/user_guide/scaling.rst b/docs/user_guide/scaling.rst index 32f6e6cf..c498ae1c 100644 --- a/docs/user_guide/scaling.rst +++ b/docs/user_guide/scaling.rst @@ -149,3 +149,31 @@ Use ``WilsonNormaliser`` for normalized E values. Work, free and validation observations are accessed through ``data.work``, ``data.free`` and ``data.validation``; full observations are available directly as ``data.F`` and ``data.F_sigma``. + +Wilson B +-------- + +``torchref.scaling.wilson.fit_wilson_b`` fits xtriage's Wilson model, + +.. math:: + + \langle I/\epsilon \rangle = K \, \Sigma f^2(s) \, (1 + \gamma(s)) \, e^{-B d^{*2}/2}, + +through equal-count resolution shells, and returns a ``WilsonFit`` (B, its +standard error, the fitted curve) or ``None`` with a warning when the data +cannot support a fit. :math:`\Sigma f^2` is the random-atom intensity of an +average protein residue; :math:`\gamma` is the empirical protein correction of +Zwart & Lamzin (2004, Acta Cryst. D60, 220-226), taken from cctbx. + +Both corrections matter. Without :math:`\Sigma f^2` the fit reads the +form-factor falloff as B, 7-10 Ų too high on the deposited test data. Without +:math:`\gamma` the plot is only linear below about 4.5 Å, so data that stop at +3.5-4 Å have nothing straight left to fit; with it, cutting 1DAW, 3E98, 3K7M, +4BX9, 5BOV and 6G9X at 3.5 Å moves B by at most 9 Ų. + +The correction also changes the slope between 4.5 and 2.5 Å, so B comes out +below the plain :math:`d \le 4.5` Å fit that ctruncate reports and that most +deposited ``B_iso_Wilson_estimate`` values follow -- by 2-16 Ų on the test +data (e.g. 4BX9: 60.5 against 76.1 deposited). Compare Wilson B values only +within one convention. + diff --git a/tests/integration/test_io_reflections.py b/tests/integration/test_io_reflections.py index cb9e9a1b..ab427eac 100644 --- a/tests/integration/test_io_reflections.py +++ b/tests/integration/test_io_reflections.py @@ -62,7 +62,7 @@ def test_resolution_bins(loaded_reflection_data) -> None: groups = [(bins == i) & data.masks() for i in range(n_bins)] assert all(group.any() for group in groups) expected = torch.stack([data.resolution[group].mean() for group in groups]) - torch.testing.assert_close(data.mean_res_per_bin(), expected) + torch.testing.assert_close(data.mean_res_per_bin(bins, n_bins), expected) @pytest.mark.integration diff --git a/tests/unit/scaling/test_scaler.py b/tests/unit/scaling/test_scaler.py index 44353856..e9636acf 100644 --- a/tests/unit/scaling/test_scaler.py +++ b/tests/unit/scaling/test_scaler.py @@ -201,3 +201,39 @@ def test_u_matrix_symmetric(self, mock_aniso_u): for i in range(5): mat = U_matrices[i] assert torch.allclose(mat, mat.T, atol=1e-6) + + +class TestBinwiseMeans: + """The scaler's per-bin means use the scaler's own bins.""" + + @pytest.fixture + def scaler_and_data(self, mtz_dir): + from torchref.io import ReflectionData + from torchref.scaling.scaler_base import ScalerBase + + data = ReflectionData(verbose=0, device="cpu").load_mtz( + str(mtz_dir / "1DAW.mtz") + ) + return ScalerBase(data=data, nbins=10, verbose=0), data + + @pytest.mark.unit + def test_mean_resolution_is_per_scaler_bin(self, scaler_and_data): + scaler, data = scaler_and_data + fcalc = data.F.to(torch.complex64) + _, _, mean_res = scaler.get_binwise_mean_intensity(fcalc) + + valid = data.masks() + per_bin = [(scaler.bins == b) & valid for b in range(scaler.nbins)] + expected = torch.stack([data.resolution[sel].mean() for sel in per_bin]) + torch.testing.assert_close(mean_res, expected) + + @pytest.mark.unit + def test_later_binning_of_the_dataset_does_not_move_the_shells( + self, scaler_and_data + ): + scaler, data = scaler_and_data + fcalc = data.F.to(torch.complex64) + before = scaler.get_binwise_mean_intensity(fcalc)[2] + data.get_bins(n_bins=3, min_per_bin=10) + after = scaler.get_binwise_mean_intensity(fcalc)[2] + torch.testing.assert_close(after, before) diff --git a/torchref/scaling/scaler_base.py b/torchref/scaling/scaler_base.py index f291baff..2cfd8842 100644 --- a/torchref/scaling/scaler_base.py +++ b/torchref/scaling/scaler_base.py @@ -432,7 +432,8 @@ def get_binwise_mean_intensity(self, fcalc: torch.Tensor): counts = torch.scatter_add(counts, 0, bins_sel, counts_vals[sel]) mean_obs_intensity = mean_obs_intensity / (counts + 1e-6) mean_calc_intensity = mean_calc_intensity / (counts + 1e-6) - return mean_obs_intensity, mean_calc_intensity, self._data.mean_res_per_bin() + mean_res = self._data.mean_res_per_bin(self.bins, self.nbins) + return mean_obs_intensity, mean_calc_intensity, mean_res def screen_solvent_params( self, diff --git a/torchref/scaling/wilson.py b/torchref/scaling/wilson.py index 0c35439a..9a63558a 100644 --- a/torchref/scaling/wilson.py +++ b/torchref/scaling/wilson.py @@ -605,15 +605,11 @@ def fit_wilson_b( """Fit the overall Wilson B, as xtriage does, by a line through shell means. Fits ``ln `` against ``d*^2`` in shells of - equal reflection count: the slope is ``-B/2``. ``sum_f2`` is the random-atom - mean intensity of ``composition``, and ``gamma`` the empirical protein - correction of Zwart & Lamzin (2004), which straightens the plot down to - about 11 Å, so data at 3.5-4 Å still give a usable B. Without either - correction the fit reads the form-factor falloff as B and comes out 7-10 Ų - too high on protein data. The correction also lowers B relative to the plain - ``d <= 4.5`` Å fit that ctruncate reports and most deposited - ``B_iso_Wilson_estimate`` values follow -- by 10-15 Ų on some datasets -- - so compare like with like. A single number for priors and reports; for + equal reflection count; the slope is ``-B/2``. ``sum_f2`` is the random-atom + mean intensity of ``composition``, ``gamma`` the empirical protein + correction of Zwart & Lamzin (2004), which keeps the plot linear from about + 11 Å to 1.2 Å. Protein-specific: set ``protein_gamma=False`` otherwise. + Values run below ctruncate-style Wilson B; see the scaling user guide. For normalising intensities use :class:`WilsonNormaliser`. Parameters From fe24fde622427a01491e7b3b1d06e36fbc6b9fe1 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Tue, 29 Sep 2026 13:50:46 +0200 Subject: [PATCH 222/250] Drop the trailing blank line in the scaling user guide Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/user_guide/scaling.rst | 1 - 1 file changed, 1 deletion(-) diff --git a/docs/user_guide/scaling.rst b/docs/user_guide/scaling.rst index c498ae1c..b3c17146 100644 --- a/docs/user_guide/scaling.rst +++ b/docs/user_guide/scaling.rst @@ -176,4 +176,3 @@ below the plain :math:`d \le 4.5` Å fit that ctruncate reports and that most deposited ``B_iso_Wilson_estimate`` values follow -- by 2-16 Ų on the test data (e.g. 4BX9: 60.5 against 76.1 deposited). Compare Wilson B values only within one convention. - From d50214ccc0cee46fca905138e9fc282db1968866 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Wed, 30 Sep 2026 09:05:14 +0200 Subject: [PATCH 223/250] Move MTZ writing to io/mtz.py and fix the 2Fo-Fc sign write_mtz and _build_anomalous_dataframe move into torchref.io.mtz as reflection_table / write_reflections; ReflectionData.write_mtz is a thin wrapper. The 2Fo-Fc and Fo-Fc coefficients, formed separately by Map and by both MTZ layouts, now come from one base function, torchref.base.fourier.map_coefficients. The MTZ writer stored |2Fo - |Fc|| with the unflipped model phase, so where 2Fo < |Fc| the 2Fo-Fc coefficient had the wrong sign: 1-7 % of reflections on 1DAW, 3K7M and 5BOV after an initial scale. Masked reflections are now filled with Fc and get zero Fo-Fc in both layouts. With the same fcalc, every other written column is unchanged. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/changelog.rst | 1 + torchref/base/fourier/__init__.py | 6 +- torchref/base/fourier/coefficients.py | 51 ++++ torchref/io/datasets/reflection_data.py | 366 ++---------------------- torchref/io/mtz.py | 237 ++++++++++++++- torchref/maps/map.py | 10 +- 6 files changed, 315 insertions(+), 356 deletions(-) create mode 100644 torchref/base/fourier/coefficients.py diff --git a/docs/changelog.rst b/docs/changelog.rst index e0c5eabc..6f6018b2 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Fixed the 2Fo-Fc coefficients (FWT/PHWT) written to MTZ: where 2Fo < |Fc| the amplitude was made positive without flipping the phase, reversing the sign of 1-7 % of reflections on the test structures. Unmeasured (masked) reflections are now filled with Fc and get zero Fo-Fc in both layouts. ``Map`` and the MTZ writer share ``torchref.base.fourier.map_coefficients``, and the writer moved to ``torchref.io.mtz.write_reflections`` (``ReflectionData.write_mtz`` calls it) - Fixed ``ReflectionData.expand_to_p1`` dropping one Bijvoet mate per reflection on anomalous data (about half the reflections when one mate was unmeasured). ``SpaceGroup.expand_hkl`` now raises on symmetry-equivalent input rows instead of silently keeping one; ``SpaceGroup.equivalent_hkl`` returns every symmetry copy with its source row - ``Map``, ``DifferenceMap`` and the experimental real-space targets average each Bijvoet pair (``ReflectionData.bijvoet_mean`` / ``bijvoet_representatives``) before placing amplitudes on the grid, so anomalous input gives the Friedel-averaged map; maps from merged data are unchanged - ``torchref.scaling.wilson.fit_wilson_b`` now fits the form-factor and protein-corrected Wilson model and returns a ``WilsonFit`` or ``None`` with a warning; the previous fit was 7-10 Ų too high and returned a fixed 200 Ų on data short of 3.5 Å. Values run below ctruncate-style Wilson B (see the scaling user guide). The ensemble Wilson prior uses the fitted curve. Removed the ``wilson_b``, ``wilson_b_structure``, ``wilson_b_solvent`` and ``wilson_k_sol`` dataset fields diff --git a/torchref/base/fourier/__init__.py b/torchref/base/fourier/__init__.py index ca14471c..4131e1ce 100644 --- a/torchref/base/fourier/__init__.py +++ b/torchref/base/fourier/__init__.py @@ -1,5 +1,7 @@ -"""Fourier transforms in the crystallographic convention, plus grid generation.""" +"""Fourier transforms in the crystallographic convention, grid generation, and +the 2Fo-Fc / Fo-Fc map coefficients.""" +from .coefficients import map_coefficients from .fft import fft, ifft from .grid import ( @@ -20,4 +22,6 @@ "get_real_grid_numpy", "get_grids", "put_hkl_on_grid", + # Map coefficients + "map_coefficients", ] diff --git a/torchref/base/fourier/coefficients.py b/torchref/base/fourier/coefficients.py new file mode 100644 index 00000000..129c6d23 --- /dev/null +++ b/torchref/base/fourier/coefficients.py @@ -0,0 +1,51 @@ +"""Map coefficients from observed amplitudes and a model. + +The one place the 2Fo-Fc and Fo-Fc coefficients are formed, for real-space maps +(:class:`torchref.maps.Map`) and for the FWT/DELFWT columns an MTZ carries. +""" + +from typing import Optional, Tuple + +import torch + +__all__ = ["map_coefficients"] + + +def map_coefficients( + fobs: torch.Tensor, + fcalc: torch.Tensor, + observed: Optional[torch.Tensor] = None, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Unweighted ``2Fo-Fc`` and ``Fo-Fc`` coefficients on the model phases. + + ``2Fo-Fc = (2 Fo - |Fc|) exp(i phi_c)`` and ``Fo-Fc = (Fo - |Fc|) exp(i phi_c)``, + both signed: where ``2 Fo < |Fc|`` the 2Fo-Fc coefficient points opposite to + the model phase, which an amplitude/phase pair must carry as a 180° flip. + These are the m = 1, D = 1 forms, not likelihood-weighted 2mFo-DFc maps. + + Parameters + ---------- + fobs : torch.Tensor + Observed amplitudes, shape (N,). + fcalc : torch.Tensor + Complex model structure factors on the same scale, shape (N,). + observed : torch.Tensor, optional + Boolean (N,), reflections with a usable measurement. Unobserved ones get + ``Fc`` in 2Fo-Fc (the model fills the missing term) and zero in Fo-Fc. + Default: all observed. + + Returns + ------- + two_fo_fc, fo_fc : torch.Tensor + Complex coefficients, shape (N,). + """ + fcalc_amp = fcalc.abs() + phase = torch.exp(1j * torch.angle(fcalc)) + fobs = fobs.to(fcalc_amp) + two_fo_fc = (2.0 * fobs - fcalc_amp) * phase + fo_fc = (fobs - fcalc_amp) * phase + if observed is not None: + observed = observed.to(device=fcalc.device, dtype=torch.bool) + two_fo_fc = torch.where(observed, two_fo_fc, fcalc) + fo_fc = torch.where(observed, fo_fc, torch.zeros_like(fo_fc)) + return two_fo_fc, fo_fc diff --git a/torchref/io/datasets/reflection_data.py b/torchref/io/datasets/reflection_data.py index cd93d2e2..f051273f 100644 --- a/torchref/io/datasets/reflection_data.py +++ b/torchref/io/datasets/reflection_data.py @@ -12,7 +12,6 @@ from typing import TYPE_CHECKING, Callable, Optional, Tuple, Union import numpy as np -import pandas as pd import torch from torchref.base import math_torch @@ -1884,212 +1883,6 @@ def flag_wilson_outliers( device=self.device, dtype=torch.bool ) - def _build_anomalous_dataframe( - self, fcalc: Optional[torch.Tensor] = None - ) -> pd.DataFrame: - """Build a phenix-style anomalous MTZ DataFrame on the canonical ASU. - - Bijvoet mates share a canonical ASU index in :attr:`hkl`; here they are - (a) merged by mean amplitude for the display maps / ``F-obs`` / ``F-model`` - and (b) unstacked into ``(+)/(-)`` columns. No negative-ASU Miller - indices are emitted, so the display maps render normally in Coot while - the anomalous columns are available for anomalous difference maps. - - Parameters - ---------- - fcalc : torch.Tensor, optional - Complex per-row structure factors in the canonical-ASU convention, - row-aligned with :attr:`hkl` (see :meth:`structure_factors`). If - None, only the observation columns (``F-obs``, ``F-obs(+/-)``, - ``SIGF-obs(+/-)``, R-free) are written -- the model-derived columns - (``F-model``, ``PHIF-model``, display maps and ``ANOM``/``PANOM``, - which need the model phase) are omitted. - - Returns - ------- - pandas.DataFrame - One row per unique canonical ASU reflection. - """ - if fcalc is not None and not torch.is_complex(fcalc): - raise ValueError("anomalous fcalc, when provided, must be complex") - has_model = fcalc is not None - if self.friedel_flags is None: - raise ValueError( - "anomalous output requires canonicalized data with friedel_flags; " - "load via load_mtz so Friedel bookkeeping is populated." - ) - - hkl = self.hkl.detach().cpu() - N = hkl.shape[0] - flag = self.friedel_flags.detach().cpu() - - # Group rows by unique canonical ASU index; inverse maps row -> group. - inverse, M = self.asu_group_indices() - inverse = inverse.cpu() - # One row per group carries that group's canonical index by definition. - uniq = hkl[self._group_representative_rows(inverse, M)] - - # The (+) member is the unconjugated row, (-) is the Friedel-flagged row. - arange = torch.arange(N) - plus_idx = torch.full((M,), -1, dtype=torch.long) # dtype-ok: Friedel-mate index map (-1 sentinel) for indexing; PyTorch requires int64 - minus_idx = torch.full((M,), -1, dtype=torch.long) # dtype-ok: Friedel-mate index map (-1 sentinel) for indexing; PyTorch requires int64 - # A Bijvoet mate only counts as present if it is a real, positive - # observation. Stacked anomalous input (rs.stack_anomalous) carries a - # row for every *absent* mate with a NaN intensity, which French-Wilson - # maps to F=0; pairing such a phantom with its observed mate would yield - # a spurious ANOM = |F_obs - 0| = |F_obs| -- the whole amplitude, not a - # Bijvoet difference. Gate membership on the same validity convention as - # sanitize_F (finite, positive F and finite sigma) so single-mate - # reflections drop to NaN ANOM/PANOM, matching phenix. - F_cpu = self.F.detach().cpu() - observed = torch.isfinite(F_cpu) & (F_cpu > 0) - if self.F_sigma is not None: - observed = observed & torch.isfinite(self.F_sigma.detach().cpu()) - plus_sel = (~flag) & observed - minus_sel = flag & observed - plus_idx[inverse[plus_sel]] = arange[plus_sel] - minus_idx[inverse[minus_sel]] = arange[minus_sel] - has_plus = (plus_idx >= 0).numpy() - has_minus = (minus_idx >= 0).numpy() - pi = plus_idx.clamp(min=0).numpy() - mi = minus_idx.clamp(min=0).numpy() - - # Centric flags per ASU group (centrics obey Friedel's law: F(+)=F(-)). - cen_full = self.centric - centric = np.zeros(M, dtype=bool) - if cen_full is not None: - cen_full = cen_full.detach().cpu().numpy() - centric[has_plus] = cen_full[pi][has_plus] - centric[has_minus] = cen_full[mi][has_minus] - - if has_model: - fc = fcalc.detach().cpu().numpy() - Fc_amp = np.abs(fc) - # The (+)/(-) phase columns describe each mate at its own index, so - # they read the signed convention. conjugate_friedel is its own - # inverse, so it recovers that from the canonical input. - Fc_ph = np.angle( - self.conjugate_friedel(fcalc).detach().cpu().numpy(), deg=True - ) - F = self.F.detach().cpu().numpy() - Fsig = self.F_sigma.detach().cpu().numpy() if self.F_sigma is not None else None - rfree = ( - self.rfree_flags.detach().cpu().numpy().astype(int) - if self.rfree_flags is not None - else None - ) - - def plus_of(src): - out = np.full(M, np.nan, dtype=np.float64) - out[has_plus] = src[pi][has_plus] - return out - - def minus_of(src): - out = np.full(M, np.nan, dtype=np.float64) - out[has_minus] = src[mi][has_minus] - return out - - def mirror_centric(plus, minus): - # For centrics, the absent mate equals the present one. - p = np.where(centric & ~np.isfinite(plus) & np.isfinite(minus), minus, plus) - m = np.where(centric & ~np.isfinite(minus) & np.isfinite(plus), plus, minus) - return p, m - - # Observed (+/-) amplitudes (always available, no model required). - Fobs_p, Fobs_m = plus_of(F), minus_of(F) - Fobs_p_out, Fobs_m_out = mirror_centric(Fobs_p, Fobs_m) - - # Merged observed amplitude: mean over present mates. Groups with neither - # mate observed average to NaN (expected "empty slice"); nan_to_num'd below. - with warnings.catch_warnings(), np.errstate(invalid="ignore"): - warnings.simplefilter("ignore", category=RuntimeWarning) - Fobs_disp = np.nanmean(np.vstack([Fobs_p, Fobs_m]), axis=0) - - uniq_np = uniq.numpy() - data = { - "H": uniq_np[:, 0], - "K": uniq_np[:, 1], - "L": uniq_np[:, 2], - "F-obs": Fobs_disp, - "F-obs(+)": Fobs_p_out, - "F-obs(-)": Fobs_m_out, - } - # Columns that must be FFT-safe (no NaN). Model/map columns are appended - # to this list only when a model is supplied. - fft_safe = ["F-obs"] - - if has_model: - Fmod_p, Fmod_m = plus_of(Fc_amp), minus_of(Fc_amp) - Phi_p, Phi_m = plus_of(Fc_ph), minus_of(Fc_ph) - Fmod_p_out, Fmod_m_out = mirror_centric(Fmod_p, Fmod_m) - Phi_p_out, Phi_m_out = mirror_centric(Phi_p, Phi_m) - - # ASU representative structure factor: the + member, else the - - # member. Both rows are already on the canonical index. - fc_disp = np.full(M, np.nan, dtype=complex) - fc_disp[has_plus] = fc[pi][has_plus] - only_minus = has_minus & ~has_plus - fc_disp[only_minus] = fc[mi][only_minus] - - Fc_disp_amp = np.abs(fc_disp) - ph_disp = np.angle(fc_disp, deg=True) - - # Map coefficients (same convention as the legacy per-row path). - two_mfo = np.abs(2.0 * Fobs_disp - Fc_disp_amp) - mfo_complex = Fobs_disp * np.exp(1j * np.deg2rad(ph_disp)) - fc_disp - delf = np.abs(mfo_complex) - delph = np.angle(mfo_complex, deg=True) - - # Anomalous-difference Fourier: signed dF = |F(+)| - |F(-)| with - # phase (phi_model - 90deg). Stored in the phenix convention -- - # ANOM = |dF| (always positive) with the sign of dF carried by a - # 180deg flip in PANOM (the (-) member maps to phi-270 = phi+90) -- - # so ANOM*exp(i*PANOM) reproduces the signed dF*exp(i(phi-90)). - anom = Fobs_p_out - Fobs_m_out - panom = np.where(anom < 0.0, ph_disp - 270.0, ph_disp - 90.0) - anom = np.abs(anom) - # Centrics obey Friedel's law even under anomalous scattering, so their - # Bijvoet difference is exactly zero; any measured value is noise that - # inflates the anomalous-map RMS. Phenix omits centrics -- match that. - anom[centric] = np.nan - panom[centric] = np.nan - - data.update( - { - "F-model": Fc_disp_amp, - "PH-model": ph_disp, - "F-model(+)": Fmod_p_out, - "PHIF-model(+)": Phi_p_out, - "F-model(-)": Fmod_m_out, - "PHIF-model(-)": Phi_m_out, - "FWT": two_mfo, - "PHWT": ph_disp, - "DELFWT": delf, - "PHDELWT": delph, - "ANOM": anom, - "PANOM": panom, - } - ) - fft_safe += ["F-model", "PH-model", "FWT", "PHWT", "DELFWT", "PHDELWT"] - - if Fsig is not None: - data["SIGF-obs(+)"], data["SIGF-obs(-)"] = mirror_centric( - plus_of(Fsig), minus_of(Fsig) - ) - if rfree is not None: - rf = np.zeros(M, dtype=int) - rf[has_minus] = rfree[mi][has_minus] - rf[has_plus] = rfree[pi][has_plus] # both mates share a flag - data["R-free-flags"] = rf - - # The display-map / merged columns must be FFT-safe (no NaN); the - # anomalous (+/-) columns may legitimately carry NaN where a mate is - # absent (incomplete anomalous data), matching phenix output. - for key in fft_safe: - data[key] = np.nan_to_num(data[key], nan=0.0) - - return pd.DataFrame(data) - def write_mtz( self, fname: str, @@ -2097,48 +1890,31 @@ def write_mtz( model_ft: Optional["ModelFT"] = None, anomalous: Optional[bool] = None, ) -> None: - """ - Write reflection data to MTZ file with optional map coefficients. + """Write this dataset, and optionally a model's map coefficients, to MTZ. + + A thin wrapper over :func:`torchref.io.mtz.write_reflections`, which + documents the layouts and on-disk labels. Parameters ---------- fname : str Output MTZ filename. fcalc : torch.Tensor, optional - Complex calculated structure factors of shape (N,), in the - canonical-ASU convention and row-aligned with :attr:`hkl` -- as - returned by :meth:`structure_factors`. If provided, computes phases - and map coefficients. + Complex structure factors of shape (N,), row-aligned with + :attr:`hkl` in the canonical-ASU convention (as returned by + :meth:`structure_factors`) and on the scale of ``F``. Adds model + and 2Fo-Fc / Fo-Fc columns. model_ft : ModelFT, optional - ModelFT object to compute fcalc if not provided. + Used to compute ``fcalc`` when it is not given. anomalous : bool, optional - If True, write a phenix-style anomalous MTZ on the canonical ASU: - display maps (FWT/PHWT, DELFWT/PHDELWT) and merged F-obs/F-model - with Friedel mates merged by mean amplitude, plus unstacked - F-obs(+/-), SIGF-obs(+/-), F-model(+/-), PHIF-model(+/-) and - ANOM/PANOM columns. No negative-ASU indices are emitted. If False, - the legacy per-row layout is written. If None (default), this is - chosen automatically from the data: anomalous output when the data - were loaded as Bijvoet pairs (``friedel_merged`` is False), legacy - layout otherwise. - - Notes - ----- - Final on-disk labels (``mtz.write`` remaps the intermediate DataFrame - keys ``F-obs``/``SIGF-obs``/``I-obs``/``SIGI-obs``/``R-free-flags``): - FP, SIGFP, I, SIGI, FreeR_flag, plus FWT/PHWT and DELFWT/PHDELWT when - ``fcalc`` is given. - - The map coefficients use the standard Coot names but are the - *unweighted* forms ``2Fo-Fc`` and ``Fo-Fc`` (m=1, D=1) -- not - likelihood-weighted 2mFo-DFc / mFo-DFc maps. - """ - from torchref.io.mtz import write - - # Auto: write anomalous (+)/(-) columns when the data are Friedel pairs. - if anomalous is None: - anomalous = not self.friedel_merged + Phenix-style anomalous layout; default when the data hold Bijvoet + pairs (``friedel_merged`` False). + Raises + ------ + ValueError + If ``fcalc`` is not row-aligned with :attr:`hkl`. + """ # One fallback for both layouts, so ``fcalc`` means the same thing # whether the caller supplied it or it was derived here. cached=False # keeps a no-grad write from leaving a detached tensor in the model's @@ -2150,113 +1926,9 @@ def write_mtz( f"fcalc has {fcalc.shape[0]} rows but this dataset has " f"{len(self.hkl)}; it must be row-aligned with hkl." ) - - if anomalous: - df = self._build_anomalous_dataframe(fcalc) - write(df, self.cell.data, self.spacegroup, fname) - if self.verbose > 0: - print(f"✓ Wrote phenix-style anomalous MTZ: {fname}") - print(f" ASU reflections: {len(df)}") - print(f" Columns: {', '.join(df.columns)}") - return - - # Convert data to numpy for DataFrame creation - hkl_np = self.hkl.detach().cpu().numpy() - - # Create DataFrame with HKL indices - data_dict = { - "H": hkl_np[:, 0], - "K": hkl_np[:, 1], - "L": hkl_np[:, 2], - } - - # Add observed amplitudes (canonical names: FP, SIGFP) - if self.F is not None: - data_dict["F-obs"] = self.F.detach().cpu().numpy() - if self.F_sigma is not None: - data_dict["SIGF-obs"] = self.F_sigma.detach().cpu().numpy() - - # Add observed intensities (canonical names: I, SIGI) - if self.I is not None: - data_dict["I-obs"] = self.I.detach().cpu().numpy() - if self.I_sigma is not None: - data_dict["SIGI-obs"] = self.I_sigma.detach().cpu().numpy() - - # Add R-free flags (canonical name: FreeR_flag). - # The work/free split lives in the binary ``rfree_flags`` (1=work, - # 0=free); the optional held-out validation set lives in the separate - # boolean ``validation_flags``. They are written as two standard columns - # so external crystallography tools keep working: - # FreeR_flag: 1 = work, 0 = free (classical "1 = refined against") - # Validation_flag: 1 = validation, 0 = otherwise (optional column) - if self.rfree_flags is not None: - flags_np = self.rfree_flags.detach().cpu().numpy() - # rfree_flags is binary work/free (bool or {0,1}); write 1=work. - data_dict["R-free-flags"] = (flags_np != 0).astype(int) - # Emit Validation_flag column only if a validation set exists. - if self.validation_flags is not None and bool( - self.validation_flags.any() - ): - val_np = self.validation_flags.detach().cpu().numpy() - data_dict["Validation_flag"] = (val_np != 0).astype(int) - - mask = self.masks().detach().cpu().numpy() - # Add map coefficients if fcalc is provided - if fcalc is not None: - # Ensure fcalc is complex - if not torch.is_complex(fcalc): - raise ValueError("fcalc must be a complex tensor") - - # Convert to numpy - fcalc_np = fcalc.detach().cpu().numpy() - F_obs = self.F.detach().cpu().numpy() - - # Compute phases in degrees - phases = np.angle(fcalc_np, deg=True) - F_calc_amp = np.abs(fcalc_np) - - # Compute map coefficients - # 2Fo-Fc (unweighted, m=D=1): observed amplitudes with calculated phases - # When 2*Fobs - Fcalc < 0, flip phase by 180° and use absolute amplitude - two_mfo_dfc_raw = 2.0 * F_obs - F_calc_amp - two_mfo_dfc_amp = np.abs(two_mfo_dfc_raw) - two_mfo_dfc_phase = phases.copy() - - # Fo-Fc: Difference map (unweighted, m=D=1) - mfo_dfc_complex = F_obs * np.exp(1j * np.deg2rad(phases)) - fcalc_np - mfo_dfc_complex[~mask] = 0.0 # Zero out reflections outside mask - mfo_dfc_amp = np.abs(mfo_dfc_complex) - mfo_dfc_phase = np.angle(mfo_dfc_complex, deg=True) - - # Add 2Fo-Fc map coefficients (standard Coot names: FWT, PHWT) - data_dict["FWT"] = two_mfo_dfc_amp - data_dict["PHWT"] = two_mfo_dfc_phase - - # Add Fo-Fc map coefficients (standard Coot names: DELFWT, PHDELWT) - data_dict["DELFWT"] = mfo_dfc_amp - data_dict["PHDELWT"] = mfo_dfc_phase - - data_dict["F-model"] = F_calc_amp - data_dict["PH-model"] = phases - - if self.verbose > 0: - print("Added map coefficients:") - print(" 2Fo-Fc: FWT, PHWT") - print(" Fo-Fc: DELFWT, PHDELWT") - print( - f" Resolution range: {self.resolution.min().item():.2f} - {self.resolution.max().item():.2f} Å" - ) - - # Create DataFrame - df = pd.DataFrame(data_dict) - - # Write MTZ file - write(df, self.cell.data, self.spacegroup, fname) - - if self.verbose > 0: - print(f"✓ Wrote MTZ file: {fname}") - print(f" Reflections: {len(df)}") - print(f" Columns: {', '.join(df.columns)}") + mtz.write_reflections( + self, fname, fcalc=fcalc, anomalous=anomalous, verbose=self.verbose + ) @property def centric(self): diff --git a/torchref/io/mtz.py b/torchref/io/mtz.py index ba11fbc7..8f453c11 100644 --- a/torchref/io/mtz.py +++ b/torchref/io/mtz.py @@ -5,12 +5,14 @@ data_dict, cell, spacegroup = mtz.read('data.mtz')() mtz.write(df, cell, spacegroup, 'output.mtz') + mtz.write_reflections(data, 'output.mtz', fcalc=fcalc) # a ReflectionData The space group comes back as an H-M symbol **string** (``"P 21 21 21"``), not a SpaceGroup object; callers wrap it themselves. """ -from typing import Optional, Tuple, Union +import warnings +from typing import TYPE_CHECKING, Optional, Tuple, Union import gemmi import numpy as np @@ -18,6 +20,11 @@ import reciprocalspaceship as rs import torch +from torchref.base.fourier.coefficients import map_coefficients + +if TYPE_CHECKING: + from torchref.io.datasets.reflection_data import ReflectionData + class MTZReader: """ @@ -667,6 +674,234 @@ def write( return 1 +def _np(t: Optional[torch.Tensor]) -> Optional[np.ndarray]: + return None if t is None else t.detach().cpu().numpy() + + +def _amplitude_phase(coeff: torch.Tensor) -> Tuple[np.ndarray, np.ndarray]: + """``|c|`` and ``arg(c)`` in degrees: a negative coefficient becomes a 180° flip.""" + return _np(coeff.abs()), _np(torch.rad2deg(torch.angle(coeff))) + + +def reflection_table( + data: "ReflectionData", + fcalc: Optional[torch.Tensor] = None, + anomalous: bool = False, +) -> pd.DataFrame: + """The DataFrame :func:`write_reflections` writes, before MTZ typing. + + Parameters + ---------- + data : ReflectionData + Canonicalized dataset. + fcalc : torch.Tensor, optional + Complex structure factors row-aligned with ``data.hkl``, in the + canonical-ASU convention (``data.structure_factors``) and on the scale + of ``data.F``. Adds the model and map-coefficient columns. + anomalous : bool, optional + Phenix-style anomalous layout: one row per unique reflection, Bijvoet + mates merged by mean amplitude for the display columns and unstacked + into ``(+)/(-)`` columns, plus ANOM/PANOM. Otherwise one row per + dataset row. + + Returns + ------- + pandas.DataFrame + Columns keyed by the intermediate names :func:`write` maps to MTZ + labels (``F-obs``, ``SIGF-obs``, ``I-obs``, ``R-free-flags``, ...). + """ + if fcalc is not None and not torch.is_complex(fcalc): + raise ValueError("fcalc must be a complex tensor") + if anomalous: + return _anomalous_table(data, fcalc) + return _merged_table(data, fcalc) + + +def _merged_table(data, fcalc): + hkl = _np(data.hkl) + table = {"H": hkl[:, 0], "K": hkl[:, 1], "L": hkl[:, 2]} + if data.F is not None: + table["F-obs"] = _np(data.F) + if data.F_sigma is not None: + table["SIGF-obs"] = _np(data.F_sigma) + if data.I is not None: + table["I-obs"] = _np(data.I) + if data.I_sigma is not None: + table["SIGI-obs"] = _np(data.I_sigma) + # FreeR_flag is 1 = work, 0 = free; the optional held-out validation set is + # a separate Validation_flag column so external tools keep reading FreeR. + if data.rfree_flags is not None: + table["R-free-flags"] = (_np(data.rfree_flags) != 0).astype(int) + if data.validation_flags is not None and bool(data.validation_flags.any()): + table["Validation_flag"] = (_np(data.validation_flags) != 0).astype(int) + if fcalc is not None: + two_fo_fc, fo_fc = map_coefficients(data.F, fcalc, observed=data.masks()) + table["FWT"], table["PHWT"] = _amplitude_phase(two_fo_fc) + table["DELFWT"], table["PHDELWT"] = _amplitude_phase(fo_fc) + table["F-model"], table["PH-model"] = _amplitude_phase(fcalc) + return pd.DataFrame(table) + + +def _anomalous_table(data, fcalc): + if data.friedel_flags is None: + raise ValueError( + "anomalous output requires canonicalized data with friedel_flags; " + "load via load_mtz so Friedel bookkeeping is populated." + ) + hkl = data.hkl.detach().cpu() + n = hkl.shape[0] + flag = data.friedel_flags.detach().cpu() + inverse, m = data.asu_group_indices() + inverse = inverse.cpu() + uniq = hkl[data._group_representative_rows(inverse, m)] + + # A mate counts as present only if it is a real, positive observation: + # stacked input carries a NaN row for every absent mate, which French-Wilson + # maps to F=0, and pairing that phantom with its observed mate would write + # the whole amplitude as the Bijvoet difference. + F_cpu = data.F.detach().cpu() + observed = torch.isfinite(F_cpu) & (F_cpu > 0) + if data.F_sigma is not None: + observed = observed & torch.isfinite(data.F_sigma.detach().cpu()) + arange = torch.arange(n) + plus_idx = torch.full((m,), -1, dtype=torch.long) # dtype-ok: Friedel-mate index map (-1 sentinel) for indexing; PyTorch requires int64 + minus_idx = torch.full((m,), -1, dtype=torch.long) # dtype-ok: Friedel-mate index map (-1 sentinel) for indexing; PyTorch requires int64 + plus_sel, minus_sel = (~flag) & observed, flag & observed + plus_idx[inverse[plus_sel]] = arange[plus_sel] + minus_idx[inverse[minus_sel]] = arange[minus_sel] + has_plus, has_minus = (plus_idx >= 0).numpy(), (minus_idx >= 0).numpy() + pi, mi = plus_idx.clamp(min=0).numpy(), minus_idx.clamp(min=0).numpy() + + # Centrics obey Friedel's law, F(+) = F(-). + centric = np.zeros(m, dtype=bool) + if data.centric is not None: + cen = _np(data.centric) + centric[has_plus] = cen[pi][has_plus] + centric[has_minus] = cen[mi][has_minus] + + def plus_of(src): + out = np.full(m, np.nan, dtype=np.float64) + out[has_plus] = src[pi][has_plus] + return out + + def minus_of(src): + out = np.full(m, np.nan, dtype=np.float64) + out[has_minus] = src[mi][has_minus] + return out + + def mirror_centric(plus, minus): + p = np.where(centric & ~np.isfinite(plus) & np.isfinite(minus), minus, plus) + q = np.where(centric & ~np.isfinite(minus) & np.isfinite(plus), plus, minus) + return p, q + + F = _np(data.F) + Fobs_p, Fobs_m = plus_of(F), minus_of(F) + Fobs_p_out, Fobs_m_out = mirror_centric(Fobs_p, Fobs_m) + # Mean over present mates; NaN where neither was measured. + with warnings.catch_warnings(), np.errstate(invalid="ignore"): + warnings.simplefilter("ignore", category=RuntimeWarning) + Fobs_disp = np.nanmean(np.vstack([Fobs_p, Fobs_m]), axis=0) + measured = np.isfinite(Fobs_disp) + + uniq_np = uniq.numpy() + table = { + "H": uniq_np[:, 0], + "K": uniq_np[:, 1], + "L": uniq_np[:, 2], + "F-obs": np.nan_to_num(Fobs_disp, nan=0.0), + "F-obs(+)": Fobs_p_out, + "F-obs(-)": Fobs_m_out, + } + + if fcalc is not None: + fc = _np(fcalc) + # The (+)/(-) phase columns describe each mate at its own index, the + # signed convention; conjugate_friedel is its own inverse. + Fc_ph = np.angle(_np(data.conjugate_friedel(fcalc)), deg=True) + fc_amp = np.abs(fc) + Fmod_p_out, Fmod_m_out = mirror_centric(plus_of(fc_amp), minus_of(fc_amp)) + Phi_p_out, Phi_m_out = mirror_centric(plus_of(Fc_ph), minus_of(Fc_ph)) + # Representative model value per reflection: the (+) row, else the (-). + fc_disp = np.zeros(m, dtype=complex) + fc_disp[has_plus] = fc[pi][has_plus] + only_minus = has_minus & ~has_plus + fc_disp[only_minus] = fc[mi][only_minus] + fc_disp_t = torch.from_numpy(fc_disp).to(fcalc.dtype) + two_fo_fc, fo_fc = map_coefficients( + torch.from_numpy(np.nan_to_num(Fobs_disp, nan=0.0)), + fc_disp_t, + observed=torch.from_numpy(measured), + ) + ph_disp = np.angle(fc_disp, deg=True) + # Anomalous difference Fourier, phenix convention: ANOM = |F(+) - F(-)| + # with the sign carried by a 180° flip in PANOM, so ANOM exp(i PANOM) + # is (F(+) - F(-)) exp(i (phi - 90°)). Centric differences are exactly + # zero, so any measured value is noise; they are omitted, as in phenix. + anom = Fobs_p_out - Fobs_m_out + panom = np.where(anom < 0.0, ph_disp - 270.0, ph_disp - 90.0) + anom = np.abs(anom) + anom[centric] = np.nan + panom[centric] = np.nan + table["F-model"], table["PH-model"] = _amplitude_phase(fc_disp_t) + table["F-model(+)"], table["PHIF-model(+)"] = Fmod_p_out, Phi_p_out + table["F-model(-)"], table["PHIF-model(-)"] = Fmod_m_out, Phi_m_out + table["FWT"], table["PHWT"] = _amplitude_phase(two_fo_fc) + table["DELFWT"], table["PHDELWT"] = _amplitude_phase(fo_fc) + table["ANOM"], table["PANOM"] = anom, panom + + if data.F_sigma is not None: + sig = _np(data.F_sigma) + table["SIGF-obs(+)"], table["SIGF-obs(-)"] = mirror_centric( + plus_of(sig), minus_of(sig) + ) + if data.rfree_flags is not None: + rfree = _np(data.rfree_flags).astype(int) + rf = np.zeros(m, dtype=int) + rf[has_minus] = rfree[mi][has_minus] + rf[has_plus] = rfree[pi][has_plus] # both mates share a flag + table["R-free-flags"] = rf + return pd.DataFrame(table) + + +def write_reflections( + data: "ReflectionData", + filepath: str, + fcalc: Optional[torch.Tensor] = None, + anomalous: Optional[bool] = None, + verbose: int = 0, +) -> None: + """Write a :class:`ReflectionData` (and optional model) to an MTZ file. + + Labels on disk: FP, SIGFP, I, SIGI, FreeR_flag (1 = work), Validation_flag; + with ``fcalc`` also FWT/PHWT (2Fo-Fc), DELFWT/PHDELWT (Fo-Fc) and + F-model/PH-model -- the unweighted m = 1, D = 1 coefficients of + :func:`~torchref.base.fourier.map_coefficients`, not 2mFo-DFc. + + Parameters + ---------- + data : ReflectionData + Dataset to write. + filepath : str + Output path. + fcalc : torch.Tensor, optional + See :func:`reflection_table`. + anomalous : bool, optional + See :func:`reflection_table`. Default: anomalous exactly when the data + hold Bijvoet pairs (``friedel_merged`` False). + verbose : int, optional + Print a summary when > 0. + """ + if anomalous is None: + anomalous = not data.friedel_merged + df = reflection_table(data, fcalc, anomalous=anomalous) + write(df, data.cell.data, data.spacegroup, filepath) + if verbose > 0: + layout = "anomalous (phenix-style)" if anomalous else "merged" + print(f"✓ Wrote {layout} MTZ: {filepath}") + print(f" Reflections: {len(df)}") + print(f" Columns: {', '.join(df.columns)}") + + # Deprecated alias kept for backwards compatibility; prefer MTZReader. # Slated for removal in a future release. This is a public symbol. MTZ = MTZReader diff --git a/torchref/maps/map.py b/torchref/maps/map.py index 97bc3788..8f0eec98 100644 --- a/torchref/maps/map.py +++ b/torchref/maps/map.py @@ -21,6 +21,7 @@ import torch +from torchref.base.fourier.coefficients import map_coefficients from torchref.base.reciprocal.grid_operations import place_on_grid from torchref.io.cif import write_map from torchref.utils.device_mixin import DeviceMixin @@ -136,13 +137,8 @@ def _compute_map_coefficients( if self.map_type == "Fcalc": return fcalc - # 2Fo-Fc: (2*Fobs - |Fcalc|) * exp(i * phi_calc). Note this is a plain - # 2Fo-Fc map: no figure-of-merit ``m`` weights Fobs and no sigma-A - # coefficient ``D`` scales Fcalc (i.e. m=1, D=1), so it is not a true - # likelihood-weighted 2mFo-DFc map. - fcalc_amp = fcalc.abs() - phi_calc = torch.angle(fcalc) - return (2.0 * fobs - fcalc_amp) * torch.exp(1j * phi_calc) + # Plain 2Fo-Fc (m=1, D=1), not a likelihood-weighted 2mFo-DFc map. + return map_coefficients(fobs, fcalc)[0] def calculate(self) -> torch.Tensor: """Compute the electron density map. From 0cb849bd16398f4273af97042723aa9d51beb733 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 09:34:24 +0000 Subject: [PATCH 224/250] Format the uniform-rfree module, CLI and tests with black Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01BKDn7EaPFPeiG4rLBKxseN --- tests/integration/test_cli_uniform_rfree.py | 26 +++++++++--- tests/unit/io/test_uniform_rfree.py | 44 ++++++++++++++++----- torchref/cli/uniform_rfree.py | 22 +++++++---- torchref/io/rfree.py | 44 +++++++++++++++++---- 4 files changed, 105 insertions(+), 31 deletions(-) diff --git a/tests/integration/test_cli_uniform_rfree.py b/tests/integration/test_cli_uniform_rfree.py index 9a0f628f..ea8313d9 100644 --- a/tests/integration/test_cli_uniform_rfree.py +++ b/tests/integration/test_cli_uniform_rfree.py @@ -36,7 +36,9 @@ def inputs(mtz_dir, cif_sf_dir, tmp_path): L = light.get_hkls()[:, 2].astype(float) f = 1.7 * np.exp(-2.0 * (L / L.max()) ** 2) light["FP"] = rs.DataSeries(light.FP.to_numpy() * f, index=light.index, dtype="F") - light["SIGFP"] = rs.DataSeries(light.SIGFP.to_numpy() * f, index=light.index, dtype="Q") + light["SIGFP"] = rs.DataSeries( + light.SIGFP.to_numpy() * f, index=light.index, dtype="Q" + ) paths = [tmp_path / "dark.mtz", tmp_path / "light.mtz"] dark.write_mtz(str(paths[0])) light.write_mtz(str(paths[1])) @@ -70,7 +72,9 @@ def test_uniform_flags_mtz_and_cif(inputs, tmp_path): def test_scale_onto_reference(inputs, tmp_path): full, paths = inputs out = tmp_path / "out" - res = _run(*paths[:2], "-o", out, "--scale", "--scale-reference", "dark", "--device", "cpu") + res = _run( + *paths[:2], "-o", out, "--scale", "--scale-reference", "dark", "--device", "cpu" + ) assert res.returncode == 0, res.stderr light = rs.read_mtz(str(out / "light_rfree.mtz")) ratio = light.FP.to_numpy() / full.loc[light.index].FP.to_numpy() @@ -92,7 +96,9 @@ def test_torchref_reads_flags(inputs, tmp_path): def test_mismatched_cell_fails(inputs, tmp_path): _, paths = inputs other = rs.read_mtz(str(paths[1])) - other.cell = gemmi.UnitCell(*(np.array(other.cell.parameters) * [1.05, 1, 1, 1, 1, 1])) + other.cell = gemmi.UnitCell( + *(np.array(other.cell.parameters) * [1.05, 1, 1, 1, 1, 1]) + ) bad = tmp_path / "bad.mtz" other.write_mtz(str(bad)) res = _run(paths[0], bad, "-o", tmp_path / "out") @@ -102,7 +108,9 @@ def test_mismatched_cell_fails(inputs, tmp_path): def _strip_flags(src, dst): ds = rs.read_mtz(str(src)) - ds.drop(columns=[c for c in ds.columns if c in rfree.FLAG_COLUMN_NAMES]).write_mtz(str(dst)) + ds.drop(columns=[c for c in ds.columns if c in rfree.FLAG_COLUMN_NAMES]).write_mtz( + str(dst) + ) def test_check_mode(inputs, tmp_path): @@ -175,10 +183,16 @@ def test_excluded_flags_survive(inputs, tmp_path): back = rs.read_mtz(str(out / "excl_rfree.mtz")) assert (back["FreeR_flag"].to_numpy()[:100] == -1).all() # the light file does not inherit dark's exclusions - assert (rs.read_mtz(str(out / "light_rfree.mtz"))["FreeR_flag"].to_numpy() >= 0).all() + assert ( + rs.read_mtz(str(out / "light_rfree.mtz"))["FreeR_flag"].to_numpy() >= 0 + ).all() for f in ("excl_rfree.mtz", "excl_rfree.cif"): data = ReflectionData(device="cpu", verbose=0) - data.load_mtz(str(out / f)) if f.endswith("mtz") else data.load_cif(str(out / f)) + ( + data.load_mtz(str(out / f)) + if f.endswith("mtz") + else data.load_cif(str(out / f)) + ) assert int((~data.masks["flagged_initial"]).sum()) == 100, f diff --git a/tests/unit/io/test_uniform_rfree.py b/tests/unit/io/test_uniform_rfree.py index 60142f96..1052da33 100644 --- a/tests/unit/io/test_uniform_rfree.py +++ b/tests/unit/io/test_uniform_rfree.py @@ -30,7 +30,9 @@ def _table(ds, flags): def test_complete_table_exact_per_shell(full): - keys, flags = rfree.complete_flag_table(full.cell, full.spacegroup, 2.0, 20, shell_size=1000) + keys, flags = rfree.complete_flag_table( + full.cell, full.spacegroup, 2.0, 20, shell_size=1000 + ) assert len(flags) % 1000 == 0 hkl = rfree._unkey(keys) np.testing.assert_array_equal(rfree.hkl_keys(hkl), keys) @@ -113,7 +115,9 @@ def test_reference_free_set_is_inherited(full, convention): def test_check_compatible_flags_mismatch(full): other = full.copy() - other.cell = gemmi.UnitCell(*(np.array(full.cell.parameters) * [1.05, 1, 1, 1, 1, 1])) + other.cell = gemmi.UnitCell( + *(np.array(full.cell.parameters) * [1.05, 1, 1, 1, 1, 1]) + ) assert rfree.check_compatible({"a": full, "b": full.copy()}) == [] assert rfree.check_compatible({"a": full, "b": other}) @@ -132,15 +136,31 @@ def test_apply_flags_replaces_existing_columns(full): def test_scale_columns_amplitude_and_intensity(): ds = rs.DataSet( { - "H": [1, 2], "K": [0, 0], "L": [0, 0], - "FP": [10.0, 20.0], "SIGFP": [1.0, 2.0], - "I": [100.0, 400.0], "SIGI": [10.0, 20.0], + "H": [1, 2], + "K": [0, 0], + "L": [0, 0], + "FP": [10.0, 20.0], + "SIGFP": [1.0, 2.0], + "I": [100.0, 400.0], + "SIGI": [10.0, 20.0], "PHI": [30.0, 40.0], - "FWT": [5.0, 6.0], "PHWT": [0.0, 90.0], + "FWT": [5.0, 6.0], + "PHWT": [0.0, 90.0], }, - cell=[50, 50, 50, 90, 90, 90], spacegroup=1, + cell=[50, 50, 50, 90, 90, 90], + spacegroup=1, ).set_index(["H", "K", "L"]) - ds = ds.astype({"FP": "F", "SIGFP": "Q", "I": "J", "SIGI": "Q", "PHI": "P", "FWT": "F", "PHWT": "P"}) + ds = ds.astype( + { + "FP": "F", + "SIGFP": "Q", + "I": "J", + "SIGI": "Q", + "PHI": "P", + "FWT": "F", + "PHWT": "P", + } + ) out, cols = rfree.scale_columns(ds, np.array([2.0, 0.5])) assert cols == ["FP", "SIGFP", "I", "SIGI"] np.testing.assert_allclose(out.FP, [20, 10]) @@ -211,7 +231,9 @@ def small(mtz_dir): def test_all_excluded_rows_stay_excluded(small): - small["FreeR_flag"] = rs.DataSeries(np.full(len(small), -1), index=small.index, dtype="I") + small["FreeR_flag"] = rs.DataSeries( + np.full(len(small), -1), index=small.index, dtype="I" + ) flags, info = rfree.uniform_rfree({"x": small}, seed=0) assert info["n_excluded"]["x"] == len(small) assert (flags["x"] == -1).all() @@ -235,7 +257,9 @@ def test_conflicting_equivalents_are_inconsistent(small): def test_reference_without_free_reflections_is_rejected(small): - small["FreeR_flag"] = rs.DataSeries(np.ones(len(small)), index=small.index, dtype="I") + small["FreeR_flag"] = rs.DataSeries( + np.ones(len(small)), index=small.index, dtype="I" + ) with pytest.raises(ValueError, match="no reflection as free"): rfree.uniform_rfree({"x": small}, reference=small) report = rfree.compare_free_sets({"x": small}) diff --git a/torchref/cli/uniform_rfree.py b/torchref/cli/uniform_rfree.py index 6ba911f5..c8b23f7a 100644 --- a/torchref/cli/uniform_rfree.py +++ b/torchref/cli/uniform_rfree.py @@ -69,9 +69,7 @@ def _parse_args(argv=None): ) inp = parser.add_argument_group("Input") - inp.add_argument( - "files", nargs="+", help="Structure-factor files (.mtz or .cif)" - ) + inp.add_argument("files", nargs="+", help="Structure-factor files (.mtz or .cif)") inp.add_argument( "--cif-block", default=None, help="Data block to read from CIF inputs" ) @@ -191,7 +189,9 @@ def _parse_args(argv=None): add_device_arg(scl) out = parser.add_argument_group("Output") - add_outdir_arg(out, required=False, help="Output directory (required unless --check)") + add_outdir_arg( + out, required=False, help="Output directory (required unless --check)" + ) out.add_argument( "--format", nargs="+", @@ -297,7 +297,10 @@ def _new_report(datasets, flags, info, ref_label, args): f" inherited from {ref_label} ({r['column']}, {r['convention']}, " f"{r['dmin']:.2f} A): {info['n_inherited']} kept, {info['n_generated']} new" ) - n_beyond, n_gaps = info["n_generated_beyond_reference"], info["n_gaps_in_reference"] + n_beyond, n_gaps = ( + info["n_generated_beyond_reference"], + info["n_gaps_in_reference"], + ) if n_beyond: print( f" Warning: reference ends at {r['dmin']:.2f} A; {n_beyond} reflections " @@ -489,11 +492,16 @@ def main(argv=None): else rfree.read_sf_file(args.reference, cif_block=args.cif_block) ) except Exception as exc: # noqa: BLE001 - print(f"Error: cannot read reference {args.reference}: {exc}", file=sys.stderr) + print( + f"Error: cannot read reference {args.reference}: {exc}", file=sys.stderr + ) return 1 ref_label = ref_label or args.reference if rfree.flag_column(reference, args.reference_column) is None: - print(f"Error: reference {args.reference} has no R-free column", file=sys.stderr) + print( + f"Error: reference {args.reference} has no R-free column", + file=sys.stderr, + ) return 1 problems = rfree.check_compatible( {"inputs": next(iter(datasets.values())), "reference": reference}, diff --git a/torchref/io/rfree.py b/torchref/io/rfree.py index 1d0aac26..521ca7e6 100644 --- a/torchref/io/rfree.py +++ b/torchref/io/rfree.py @@ -207,10 +207,17 @@ def compare_free_sets( try: fs = read_free_set(ds, column) except ValueError as exc: - files[name] = {"column": None, "n": len(ds), "dmin": dmin, "problem": str(exc)} + files[name] = { + "column": None, + "n": len(ds), + "dmin": dmin, + "problem": str(exc), + } continue keep = ~fs["excluded"] - ukeys, ufree, conflict = _group_free(hkl_keys(asu_hkl(ds))[keep], fs["free"][keep]) + ukeys, ufree, conflict = _group_free( + hkl_keys(asu_hkl(ds))[keep], fs["free"][keep] + ) tables[name] = (ukeys[~conflict], ufree[~conflict]) files[name] = { "column": fs["column"], @@ -227,11 +234,15 @@ def compare_free_sets( for b in names[i + 1 :]: ka, fa = tables[a] kb, fb = tables[b] - common, ia, ib = np.intersect1d(ka, kb, assume_unique=True, return_indices=True) + common, ia, ib = np.intersect1d( + ka, kb, assume_unique=True, return_indices=True + ) pairs[(a, b)] = (len(common), int((fa[ia] != fb[ib]).sum())) consistent = ( len(tables) == len(datasets) - and all(files[n]["n_free"] > 0 and files[n]["n_conflicting"] == 0 for n in tables) + and all( + files[n]["n_free"] > 0 and files[n]["n_conflicting"] == 0 for n in tables + ) and all(d == 0 for _, d in pairs.values()) ) return {"files": files, "pairs": pairs, "consistent": consistent} @@ -307,7 +318,9 @@ def hkl_keys(hkl: np.ndarray) -> np.ndarray: def _unkey(keys: np.ndarray) -> np.ndarray: """Inverse of :func:`hkl_keys`.""" span = 2 * _KEY_OFFSET - return np.stack([keys // span**2, keys // span % span, keys % span], 1) - _KEY_OFFSET + return ( + np.stack([keys // span**2, keys // span % span, keys % span], 1) - _KEY_OFFSET + ) def _lookup(table_keys: np.ndarray, table_values: np.ndarray, keys: np.ndarray): @@ -563,9 +576,14 @@ def uniform_rfree( else: n_flags, source = 20, "default" if max_free is not None: - n_complete = len(rs.utils.generate_reciprocal_asu(cell, sg, dmin, anomalous=False)) + n_complete = len( + rs.utils.generate_reciprocal_asu(cell, sg, dmin, anomalous=False) + ) if n_complete / n_flags > max_free: - n_flags, source = int(np.ceil(n_complete / max_free)), f"max_free={max_free}" + n_flags, source = ( + int(np.ceil(n_complete / max_free)), + f"max_free={max_free}", + ) if rinfo is not None and rinfo["convention"] != "ccp4": rkeys, rflags, rinfo = reference_flags( reference, n_flags, column=reference_column, seed=seed @@ -655,7 +673,17 @@ def apply_flags( _INTENSITY_TYPES = {"J", "K", "M"} # I, I(+/-), sigma I(+/-) -_CALC_PREFIXES = ("FC", "FCALC", "FMODEL", "F-MODEL", "FCAL", "FWT", "DELFWT", "2FOFC", "FOFC") +_CALC_PREFIXES = ( + "FC", + "FCALC", + "FMODEL", + "F-MODEL", + "FCAL", + "FWT", + "DELFWT", + "2FOFC", + "FOFC", +) def _is_calculated(name: str) -> bool: From 4b7a0d6cced2072c9e920e045fe64668a8253798 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 09:35:26 +0000 Subject: [PATCH 225/250] Honour --reference-column when the reference is picked automatically With --reference auto, the existing-flag survey ignored the named column, so an input whose free set lives in a non-standard column was reported as having none and a new set was generated in place of inheriting it. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01BKDn7EaPFPeiG4rLBKxseN --- tests/integration/test_cli_uniform_rfree.py | 23 +++++++++++++++++++++ torchref/cli/uniform_rfree.py | 11 ++++++++-- 2 files changed, 32 insertions(+), 2 deletions(-) diff --git a/tests/integration/test_cli_uniform_rfree.py b/tests/integration/test_cli_uniform_rfree.py index ea8313d9..0eee9eff 100644 --- a/tests/integration/test_cli_uniform_rfree.py +++ b/tests/integration/test_cli_uniform_rfree.py @@ -148,6 +148,29 @@ def test_auto_inherits_or_generates(inputs, tmp_path): assert "generating a new free set" in res.stdout +def test_auto_reference_honours_reference_column(inputs, tmp_path): + """A non-standard flag column named by --reference-column is inherited.""" + _, paths = inputs + ds = rs.read_mtz(str(paths[0])) + renamed = tmp_path / "renamed.mtz" + ds.rename(columns={"FreeR_flag": "MYFREE"}).write_mtz(str(renamed)) + bare = tmp_path / "bare.mtz" + _strip_flags(paths[1], bare) + + res = _run(bare, renamed, "--reference-column", "MYFREE", "--check") + assert "MYFREE" in res.stdout + + out = tmp_path / "out" + res = _run(bare, renamed, "--reference-column", "MYFREE", "-o", out) + assert res.returncode == 0, res.stderr + assert "inherited from renamed" in res.stdout + orig = _free_tables([paths[0]])[0] + new = _free_tables([out / "bare_rfree.mtz"])[0] + common = set(new) & set(orig) + assert common + assert all(new[k] == orig[k] for k in common) + + def test_mixed_resolution_and_max_free(inputs, tmp_path): _, paths = inputs out = tmp_path / "out" diff --git a/torchref/cli/uniform_rfree.py b/torchref/cli/uniform_rfree.py index c8b23f7a..bdb23ebe 100644 --- a/torchref/cli/uniform_rfree.py +++ b/torchref/cli/uniform_rfree.py @@ -164,7 +164,10 @@ def _parse_args(argv=None): flg.add_argument( "--reference-column", default=None, - help="R-free column in --reference (default: auto-detect)", + help=( + "R-free column in --reference, or in the inputs with --reference auto " + "and --check (default: auto-detect)" + ), ) scl = parser.add_argument_group("Scaling") @@ -441,7 +444,11 @@ def main(argv=None): return 1 # ---- existing flags ------------------------------------------------- - report = rfree.compare_free_sets(datasets) + # With --reference auto the reference is one of the inputs, so the named + # column is the one to look for in them; an explicit reference file's column + # says nothing about the inputs' own flags. + input_column = args.reference_column if args.reference == "auto" else None + report = rfree.compare_free_sets(datasets, input_column) if args.check or args.verbose > 1 or (args.verbose and not report["consistent"]): _existing_report(report) elif args.verbose: From f0f92303a954b00dcc950494ab71034492ed95f2 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 09:36:20 +0000 Subject: [PATCH 226/250] Reject Miller indices the R-free key encoding cannot hold, and max_free < 1 hkl_keys packs h, k, l into one int64 with a fixed 2048-wide field, so an index of magnitude 1024 or more carried into its neighbour and two reflections silently shared a flag. It now raises. uniform_rfree also rejects max_free < 1, which divided by zero. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01BKDn7EaPFPeiG4rLBKxseN --- tests/unit/io/test_uniform_rfree.py | 15 +++++++++++++++ torchref/io/rfree.py | 30 +++++++++++++++++++++++++++-- 2 files changed, 43 insertions(+), 2 deletions(-) diff --git a/tests/unit/io/test_uniform_rfree.py b/tests/unit/io/test_uniform_rfree.py index 1052da33..0965f2f2 100644 --- a/tests/unit/io/test_uniform_rfree.py +++ b/tests/unit/io/test_uniform_rfree.py @@ -264,3 +264,18 @@ def test_reference_without_free_reflections_is_rejected(small): rfree.uniform_rfree({"x": small}, reference=small) report = rfree.compare_free_sets({"x": small}) assert report["files"]["x"]["n_free"] == 0 and not report["consistent"] + + +def test_hkl_keys_reject_indices_beyond_the_encoding(): + edge = np.array([[1023, -1023, 0], [-1023, 1023, 1023]]) + np.testing.assert_array_equal(rfree._unkey(rfree.hkl_keys(edge)), edge) + # (0, 1024, 0) and (1, -1024, 0) would share a key + with pytest.raises(ValueError, match="supported range"): + rfree.hkl_keys(np.array([[0, 1024, 0]])) + with pytest.raises(ValueError, match="supported range"): + rfree.hkl_keys(np.array([[1, -1024, 0]])) + + +def test_max_free_must_be_positive(small): + with pytest.raises(ValueError, match="max_free"): + rfree.uniform_rfree({"x": small}, max_free=0) diff --git a/torchref/io/rfree.py b/torchref/io/rfree.py index 521ca7e6..dc219f57 100644 --- a/torchref/io/rfree.py +++ b/torchref/io/rfree.py @@ -309,8 +309,32 @@ def asu_hkl(ds: rs.DataSet) -> np.ndarray: def hkl_keys(hkl: np.ndarray) -> np.ndarray: - """Encode integer Miller indices (N, 3) as unique int64 scalars.""" - h = hkl.astype(np.int64) + _KEY_OFFSET + """Encode integer Miller indices as unique int64 scalars. + + Parameters + ---------- + hkl : np.ndarray + Miller indices, shape (N, 3), integer-valued. + + Returns + ------- + np.ndarray + int64 keys, shape (N,), ordered like ``hkl`` and decoded by + :func:`_unkey`. + + Raises + ------ + ValueError + If any index lies outside ``[-1023, 1023]``, where the fixed-width + encoding would silently map two reflections onto one key. + """ + h = hkl.astype(np.int64) + if h.size and np.abs(h).max() >= _KEY_OFFSET: + raise ValueError( + f"Miller index {int(np.abs(h).max())} exceeds the supported range " + f"|h|, |k|, |l| < {_KEY_OFFSET}" + ) + h = h + _KEY_OFFSET span = 2 * _KEY_OFFSET return (h[:, 0] * span + h[:, 1]) * span + h[:, 2] @@ -545,6 +569,8 @@ def uniform_rfree( """ if free_fraction is not None and not 0 < free_fraction < 1: raise ValueError("free_fraction must be between 0 and 1") + if max_free is not None and max_free < 1: + raise ValueError("max_free must be at least 1") # the flag table lives on the reference's cell when there is one, so the # result does not depend on input order From 031c292daf8752866c11085c32c517f026af0a5b Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 09:41:04 +0000 Subject: [PATCH 227/250] Describe converted index tensors by the configured int dtype Docstrings and comments on tensors this branch moved to get_int_dtype() still called them int64 or long: reduce_hkl's reduction_indices, canonicalize_hkl's outputs, the map-symmetry index grid, the disorder-field anchors, EdgeBlock indices, the riding-hydrogen topology fields, the VDW cell-list arrays, the rotation-function beta_starts and the expand_reciprocal note in the rotation search. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01BKDn7EaPFPeiG4rLBKxseN --- .../experimental/alignment/frf/sitelist_ang.py | 2 +- .../experimental/alignment/rotation_search.py | 2 +- .../ensemble/quasi_crystal_amber.py | 2 +- torchref/model/disorder_field.py | 2 +- torchref/symmetry/map_symmetry.py | 2 +- torchref/symmetry/reciprocal_symmetry.py | 8 ++++---- torchref/topology/atom_graph.py | 2 +- torchref/topology/edges.py | 7 ++++--- torchref/topology/nonbonded.py | 18 +++++++++--------- torchref/topology/restraint_sets.py | 2 +- torchref/topology/riding.py | 14 +++++++------- 11 files changed, 31 insertions(+), 30 deletions(-) diff --git a/torchref/experimental/alignment/frf/sitelist_ang.py b/torchref/experimental/alignment/frf/sitelist_ang.py index be9bcb76..1658d5f4 100644 --- a/torchref/experimental/alignment/frf/sitelist_ang.py +++ b/torchref/experimental/alignment/frf/sitelist_ang.py @@ -135,7 +135,7 @@ def build_adaptive_sample_list( alphas : (N_samples,) α value per sample betas : (N_samples,) β value per sample gammas : (N_samples,) γ value per sample - beta_starts: (bmax + 1,) int64 slice [beta_starts[b]:beta_starts[b+1]] + beta_starts: (bmax + 1,) int slice [beta_starts[b]:beta_starts[b+1]] is the samples at β = b · Δ beta_grid : (bmax,) the β values in radians diff --git a/torchref/experimental/alignment/rotation_search.py b/torchref/experimental/alignment/rotation_search.py index 245c049b..85328681 100644 --- a/torchref/experimental/alignment/rotation_search.py +++ b/torchref/experimental/alignment/rotation_search.py @@ -408,7 +408,7 @@ def search_peaks( # were measured with -- a different row order changes the summation order # in the later index_add_/unique and the last bits with it. # - # It rounds to int64 internally, so the products are exact and the cast + # It rounds to integers internally, so the products are exact and the cast # below loses nothing. It also returns on the SPACE GROUP's device rather # than the caller's, so the move is load-bearing whenever they differ. hkl_unrolled = ( diff --git a/torchref/experimental/ensemble/quasi_crystal_amber.py b/torchref/experimental/ensemble/quasi_crystal_amber.py index bf90cd9e..fa7e2e51 100644 --- a/torchref/experimental/ensemble/quasi_crystal_amber.py +++ b/torchref/experimental/ensemble/quasi_crystal_amber.py @@ -562,7 +562,7 @@ def _ensure_torch_buffers(self, device: torch.device, dtype: torch.dtype) -> Non N = self._n_members n_omm = self._n_omm_per_member - # Index pairs (long) for the scatter from model atoms into OMM slots. + # Index pairs for the scatter from model atoms into OMM slots. self._src_model_idx_torch = torch.from_numpy(self._src_model_idx_np).to( device=device, dtype=get_int_dtype(), diff --git a/torchref/model/disorder_field.py b/torchref/model/disorder_field.py index 582997e1..d4417e66 100644 --- a/torchref/model/disorder_field.py +++ b/torchref/model/disorder_field.py @@ -73,7 +73,7 @@ def farthest_point_anchors(xyz: torch.Tensor, n_nodes: int) -> torch.Tensor: Returns ------- torch.Tensor - ``(K,)`` int64 atom indices, sorted ascending. + ``(K,)`` atom indices in the configured int dtype, sorted ascending. """ n_atoms = int(xyz.shape[0]) n_nodes = max(1, min(int(n_nodes), n_atoms)) diff --git a/torchref/symmetry/map_symmetry.py b/torchref/symmetry/map_symmetry.py index b58b65b2..4356b15c 100644 --- a/torchref/symmetry/map_symmetry.py +++ b/torchref/symmetry/map_symmetry.py @@ -125,7 +125,7 @@ def _index_grid(self, op_index: int) -> torch.Tensor: Returns ------- torch.Tensor - Shape ``(nx, ny, nz, 3)``, dtype ``int64``. + Shape ``(nx, ny, nz, 3)``, in the configured int dtype. Notes ----- diff --git a/torchref/symmetry/reciprocal_symmetry.py b/torchref/symmetry/reciprocal_symmetry.py index e8b0c170..34673078 100644 --- a/torchref/symmetry/reciprocal_symmetry.py +++ b/torchref/symmetry/reciprocal_symmetry.py @@ -239,7 +239,7 @@ def _reduce_hkl( ------- hkl_asu : torch.Tensor, shape (M, 3), dtype int32 Unique Miller indices in the asymmetric unit. - reduction_indices : torch.Tensor, shape (M, n_equiv), dtype int64 + reduction_indices : torch.Tensor, shape (M, n_equiv), configured int dtype Indices into ``hkl_p1`` for each ASU reflection's equivalents, **-1 where no P1 reflection exists** -- mask or clamp before gathering, or a -1 will silently read the last row: ``F_asu = aggregate(F_p1[reduction_indices], dim=1)``. @@ -413,7 +413,7 @@ def _canonicalize_hkl( ---------- sym : SpaceGroup The space group whose asymmetric unit convention applies. - hkl : torch.Tensor, shape (N, 3), dtype int32 + hkl : torch.Tensor, shape (N, 3), integer dtype Input Miller indices. include_friedel : bool, default True Whether Friedel mates are considered equivalent. @@ -425,9 +425,9 @@ def _canonicalize_hkl( Returns ------- - canonical_hkl : torch.Tensor, shape (N, 3), dtype int32 + canonical_hkl : torch.Tensor, shape (N, 3), dtype of ``hkl`` Remapped indices, sorted lexicographically by (h, k, l) when ``sort``. - phase_shifts : torch.Tensor, shape (N,), dtype float32 + phase_shifts : torch.Tensor, shape (N,), configured float dtype Additive phase correction in radians, in the same row order. friedel_flags : torch.Tensor, shape (N,), dtype bool True where Friedel conjugation was applied, in the same row order. diff --git a/torchref/topology/atom_graph.py b/torchref/topology/atom_graph.py index 7084b2c4..d86e2ace 100644 --- a/torchref/topology/atom_graph.py +++ b/torchref/topology/atom_graph.py @@ -28,7 +28,7 @@ def _build_csr(bonds: torch.Tensor, n_atoms: int) -> Tuple[torch.Tensor, torch.T Parameters ---------- bonds : torch.Tensor - Bond atom indices, shape ``(E, 2)``, dtype ``int64``. + Bond atom indices, shape ``(E, 2)``, integer dtype. n_atoms : int Number of atoms, so isolated trailing atoms still get an entry. diff --git a/torchref/topology/edges.py b/torchref/topology/edges.py index 50b3877b..a44c3eb5 100644 --- a/torchref/topology/edges.py +++ b/torchref/topology/edges.py @@ -7,8 +7,9 @@ the block: a view that shares storage, costs nothing to take, and reflects an in-place edit to the block immediately. -Nothing here is refinable. Indices are ``int64``, no gradient reaches them, and the -block is a constant for the lifetime of a topology unless the topology is mutated. +Nothing here is refinable. Indices are in the configured int dtype, no gradient reaches +them, and the block is a constant for the lifetime of a topology unless the topology is +mutated. """ from dataclasses import dataclass, field @@ -136,7 +137,7 @@ class EdgeBlock(DeviceMixin): Parameters ---------- indices : torch.Tensor - Atom indices, shape ``(E, k)``, dtype ``int64``, in canonical order. + Atom indices, shape ``(E, k)``, integer dtype, in canonical order. origin_bounds : dict ``{origin: (start, end)}`` half-open row ranges into ``indices``. Ranges are contiguous and cover the block. diff --git a/torchref/topology/nonbonded.py b/torchref/topology/nonbonded.py index e1127643..d27aa2d1 100644 --- a/torchref/topology/nonbonded.py +++ b/torchref/topology/nonbonded.py @@ -51,8 +51,8 @@ def prefilter_symop_offsets( Returns ------- - op_indices : (M,) long – symop indices for each valid combo - cell_offsets : (M, 3) long – integer cell translations + op_indices : (M,) int – symop indices for each valid combo + cell_offsets : (M, 3) int – integer cell translations """ device = xyz_frac.device fdtype = dtypes.float @@ -110,8 +110,8 @@ def assign_to_grid( xyz_frac : (N, 3) cell : Cell sg : SpaceGroup - op_indices : (M,) long - cell_offsets : (M, 3) long + op_indices : (M,) int + cell_offsets : (M, 3) int grid_dims : (3,) long – number of grid cells per axis Returns @@ -178,8 +178,8 @@ def build_cell_list( ------- sort_order : (E,) long unique_cells : (C,) long – occupied cell indices - starts : (C+1,) long – CSR boundaries into sorted arrays - cell_lookup : (n_grid_total,) long – maps flat cell → index in + starts : (C+1,) int – CSR boundaries into sorted arrays + cell_lookup : (n_grid_total,) int – maps flat cell → index in unique_cells, or -1 if empty. """ device = flat_cell.device @@ -312,9 +312,9 @@ def find_pairs_periodic_grid_v2( ASU atom index and (symop, offset) combo index per entry. unique_cells : (C,) long Occupied flat grid-cell indices. - starts : (C+1,) long + starts : (C+1,) int CSR boundaries into the sorted arrays. - cell_lookup : (n_grid_total,) long + cell_lookup : (n_grid_total,) int Maps a flat cell index to its position in ``unique_cells`` (-1 empty). grid_dims : (3,) long Number of grid cells per axis. @@ -392,7 +392,7 @@ def find_pairs_periodic_grid_v2( + nb_ijk[:, 1] * gz + nb_ijk[:, 2] ) - nb_occ_idx = cell_lookup[nb_flat] # (C,) long, -1 empty + nb_occ_idx = cell_lookup[nb_flat] # (C,) int, -1 empty has_nb = nb_occ_idx >= 0 nb_occ_safe = nb_occ_idx.clamp(min=0) diff --git a/torchref/topology/restraint_sets.py b/torchref/topology/restraint_sets.py index a3ae4320..2ab8e1ec 100644 --- a/torchref/topology/restraint_sets.py +++ b/torchref/topology/restraint_sets.py @@ -29,7 +29,7 @@ "torsion": ("intra", "disulfide"), } -#: Integer-valued edge properties, kept as ``int64`` rather than the float dtype. +#: Integer-valued edge properties, kept in the configured int dtype, not the float one. _INTEGER_PROPERTIES = frozenset({"periods", "symop_indices", "cell_offsets"}) #: Boolean edge properties. diff --git a/torchref/topology/riding.py b/torchref/topology/riding.py index 29e1872e..7a3651b9 100644 --- a/torchref/topology/riding.py +++ b/torchref/topology/riding.py @@ -75,24 +75,24 @@ class HydrogenTopology(DeviceMixin): Attributes ---------- h_parent_idx : torch.Tensor - Heavy-atom index of each riding H's parent, ``(N_h,)`` long. + Heavy-atom index of each riding H's parent, ``(N_h,)`` int. h_bond_length : torch.Tensor Ideal H-parent bond length in Angstroms, ``(N_h,)``. h_vdw_radius : torch.Tensor Van der Waals radius per H (1.20 A), ``(N_h,)``. h_placement_type : torch.Tensor - Placement-geometry enum, ``(N_h,)`` long; see the module-level constants. + Placement-geometry enum, ``(N_h,)`` int; see the module-level constants. h_slot_in_parent : torch.Tensor - Ordinal among sibling H atoms on the same parent (0, 1, 2), ``(N_h,)`` long. + Ordinal among sibling H atoms on the same parent (0, 1, 2), ``(N_h,)`` int. parent_neighbor_idx : torch.Tensor - Heavy-atom neighbours of the parent, ``(N_h, MAX_HEAVY_NB)`` long, ``-1`` + Heavy-atom neighbours of the parent, ``(N_h, MAX_HEAVY_NB)`` int, ``-1`` padded. parent_neighbor_count : torch.Tensor - Heavy-atom neighbour count per parent, ``(N_h,)`` long. + Heavy-atom neighbour count per parent, ``(N_h,)`` int. h_chainid_enc : torch.Tensor - Encoded chain ID, ``(N_h,)`` long, for same-residue filtering. + Encoded chain ID, ``(N_h,)`` int, for same-residue filtering. h_resseq : torch.Tensor - Residue sequence number, ``(N_h,)`` long, for same-residue filtering. + Residue sequence number, ``(N_h,)`` int, for same-residue filtering. type_bounds : dict ``{placement_type: (start, end)}`` bounds into the type-sorted arrays. cand_idx_i, cand_idx_j, cand_symop_idx, cand_cell_offset : torch.Tensor From ca2242bcc8bcc969a84a594764bb7f965f904718 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 09:41:04 +0000 Subject: [PATCH 228/250] Merge duplicate config imports and rename the riding row kwargs sh.py and sitelist_ang.py imported get_int_dtype from torchref.config beside an existing relative import of the same module. The row-buffer kwargs in _DerivedRowsMixin were still named ``long`` though they carry the configured int dtype. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01BKDn7EaPFPeiG4rLBKxseN --- torchref/experimental/alignment/frf/sitelist_ang.py | 4 +--- torchref/experimental/alignment/sh.py | 4 +--- torchref/model/riding_xyz.py | 12 ++++++------ 3 files changed, 8 insertions(+), 12 deletions(-) diff --git a/torchref/experimental/alignment/frf/sitelist_ang.py b/torchref/experimental/alignment/frf/sitelist_ang.py index 1658d5f4..eb226827 100644 --- a/torchref/experimental/alignment/frf/sitelist_ang.py +++ b/torchref/experimental/alignment/frf/sitelist_ang.py @@ -40,9 +40,7 @@ import torch -from torchref.config import get_int_dtype - -from ....config import canonical_device +from ....config import canonical_device, get_int_dtype from ....symmetry.symmetry import find_fft_friendly_size from .types import AdaptiveRotationFunction from .wigner_d import wigner_contraction_per_beta diff --git a/torchref/experimental/alignment/sh.py b/torchref/experimental/alignment/sh.py index 6b94c1f5..90ccaa4d 100644 --- a/torchref/experimental/alignment/sh.py +++ b/torchref/experimental/alignment/sh.py @@ -29,9 +29,7 @@ import torch -from torchref.config import get_int_dtype - -from ...config import get_float_dtype +from ...config import get_float_dtype, get_int_dtype def legendre_recurrence_coefficients(L: int, dtype, device): diff --git a/torchref/model/riding_xyz.py b/torchref/model/riding_xyz.py index e809fe19..9446b6fd 100644 --- a/torchref/model/riding_xyz.py +++ b/torchref/model/riding_xyz.py @@ -66,18 +66,18 @@ def _register_rows(self, n_full: int, frames: HydrogenFrames, device) -> None: if len(rows) and is_riding[rows[rows >= 0]].any(): raise ValueError(f"{name} must reference stored rows, not riding ones") - long = dict(dtype=get_int_dtype(), device=device) - self.register_buffer("base_row", torch.as_tensor(base, **long)) - self.register_buffer("h_row", torch.as_tensor(h, **long)) + idx = dict(dtype=get_int_dtype(), device=device) + self.register_buffer("base_row", torch.as_tensor(base, **idx)) + self.register_buffer("h_row", torch.as_tensor(h, **idx)) self.register_buffer( "parent_row", - torch.as_tensor(np.asarray(frames.parent_row, dtype=np.int64), **long), + torch.as_tensor(np.asarray(frames.parent_row, dtype=np.int64), **idx), ) self.register_buffer( - "n1_row", torch.as_tensor(np.asarray(frames.n1_row, dtype=np.int64), **long) + "n1_row", torch.as_tensor(np.asarray(frames.n1_row, dtype=np.int64), **idx) ) self.register_buffer( - "n2_row", torch.as_tensor(np.asarray(frames.n2_row, dtype=np.int64), **long) + "n2_row", torch.as_tensor(np.asarray(frames.n2_row, dtype=np.int64), **idx) ) self.register_buffer( "frame_valid", From 4e105fb26498f55f5700acfcd89da27d89db5288 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 09:41:15 +0000 Subject: [PATCH 229/250] Complete the uniform-rfree docstrings and correct what the docs promise NumPy Parameters/Returns sections, imperative summaries and units for the public helpers in torchref.io.rfree. The MTZ reader documents its int32 R-free flags. The docs now say that only MTZ output keeps every input column (mmCIF keeps those gemmi maps to _refln items), and that a fresh table is tied to the exact cell, so a dataset added later matches only when run with an already flagged file or --reference. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01BKDn7EaPFPeiG4rLBKxseN --- docs/changelog.rst | 2 +- docs/user_guide/cli.rst | 14 ++- tests/unit/io/test_uniform_rfree.py | 2 +- torchref/io/mtz.py | 4 +- torchref/io/rfree.py | 185 +++++++++++++++++++++++++--- 5 files changed, 185 insertions(+), 22 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 951a9786..06e54761 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -7,7 +7,7 @@ Unreleased - Added ``torchref.uniform-rfree``: gives any number of MTZ / SF-mmCIF files of one cell and space group a shared CCP4 ``FreeR_flag``. - An existing free set (CCP4, Phenix or mmCIF convention) is inherited by default and extended, at its own fraction, to reflections it lacks, including those beyond its resolution. That extension is seeded with a hash of the reference's free/work partition, so it is reproducible whatever the file format or convention. New sets are stratified by resolution shell on the complete ASU and depend only on cell, space group and seed. - Options: ``--check`` reports whether the inputs' free sets agree; ``--max-free`` caps the size of a new set; ``--scale`` optionally scales the datasets together with ``DatasetCollection.scale``. - - Excluded reflections (``-1`` / ``x``) are preserved per file. All input columns are kept, and output is MTZ and/or mmCIF. Library helpers are in ``torchref.io.rfree``. + - Excluded reflections (``-1`` / ``x``) are preserved per file. Output is MTZ (all input columns kept) and/or mmCIF (columns with an mmCIF equivalent). Library helpers are in ``torchref.io.rfree``. - The MTZ reader keeps R-free flags as integers, so ``-1`` (excluded) reflections are masked on load instead of becoming work reflections. - ``ModelFT.create_from_state_dict`` restores the restraint dictionary path (``cif_path``) as ``Model`` does, and ``Refinement.create_from_state_dict`` no longer builds a stray ``Restraints`` from the model; the restored model builds its own on first access. - The difference MTZ groups its columns into named datasets -- ``observed``, ``difference``, ``light_model``, ``extrapolated_light``, ``two_moment`` -- with one history line describing each, so ``FWT``/``PHWT`` reads as ``/torchref/extrapolated_light/FWT`` (the extrapolated light-state map ``2*FEXT - Fc``). Labels are unchanged and Coot still auto-opens it diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index 7118cc3d..7a81521d 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -128,13 +128,19 @@ reflection is free in one dataset and work in another. - **Reproducibility.** New flags are drawn on the complete reciprocal ASU (Friedel mates and symmetry equivalents share a flag), with exactly the free fraction in every resolution shell of ``--shell-size`` reflections. Each flag - depends only on cell, space group, fraction and ``--seed``, so a dataset - added later gets the same flags. + depends only on cell, space group, fraction and ``--seed``. The cell is that + of the reference, or of the first input when there is none, and the flags + are sensitive to it at the 1e-5 level. To give a dataset added later the same + flags, run it together with an already flagged file (inherited by default) + or name that file with ``--reference``. A separate ``--fresh`` run on a + dataset with its own cell gives a different free set. - **Excluded reflections** (MTZ flag ``-1``, mmCIF ``status x``) stay excluded in the file that marked them, as ``-1`` / ``x``. They are not copied to the other files. -- **Output.** Every input column is kept; existing flag columns are replaced - by a CCP4 ``FreeR_flag`` (``0..N-1``, ``0`` = free). mmCIF output writes +- **Output.** MTZ output keeps every input column; existing flag columns are + replaced by a CCP4 ``FreeR_flag`` (``0..N-1``, ``0`` = free). mmCIF output + keeps only the columns gemmi's MTZ-to-mmCIF conversion maps to a ``_refln`` + item (e.g. ``DANO`` and custom columns are dropped), and writes ``_refln.status`` ``f``/``o``/``x``, so only the free/work split is kept. .. code-block:: bash diff --git a/tests/unit/io/test_uniform_rfree.py b/tests/unit/io/test_uniform_rfree.py index 0965f2f2..493fc230 100644 --- a/tests/unit/io/test_uniform_rfree.py +++ b/tests/unit/io/test_uniform_rfree.py @@ -223,7 +223,7 @@ def test_extension_is_identical_across_targets(full): @pytest.fixture def small(mtz_dir): - """30 rows of deposited 1DAW (the reviewer's probe set).""" + """The first 30 rows of deposited 1DAW.""" path = mtz_dir / "1DAW.mtz" if not path.exists(): pytest.skip("1DAW.mtz not found") diff --git a/torchref/io/mtz.py b/torchref/io/mtz.py index a06052ca..fe8c7399 100644 --- a/torchref/io/mtz.py +++ b/torchref/io/mtz.py @@ -305,7 +305,9 @@ def __call__(self) -> Tuple[dict, np.ndarray, str]: file, and may include: ``"HKL"`` (int32 Miller indices); ``"F"`` / ``"SIGF"`` and/or ``"I"`` / ``"SIGI"`` (float32 data, with ``"*_col"`` provenance keys recording the source column names); - ``"R-free-flags"`` (a **bool** mask) and ``"R-free-source"``; + ``"R-free-flags"`` (int32: ``0`` = free, positive = work, + negative = excluded; a column whose majority value is ``0`` is + flipped to this convention) and ``"R-free-source"``; ``"Validation-flags"`` (a **bool** mask) and ``"Validation-source"``; and ``"friedel_merged"`` (bool) indicating the Bijvoet state of the returned data (False when anomalous F(+)/F(-) pairs were stacked). diff --git a/torchref/io/rfree.py b/torchref/io/rfree.py index dc219f57..c16b9496 100644 --- a/torchref/io/rfree.py +++ b/torchref/io/rfree.py @@ -12,7 +12,9 @@ ``0`` = free, so a different test set ``k`` can still be selected later. All functions here operate on :class:`reciprocalspaceship.DataSet` objects so -that every original column of the input files is preserved on output. +that every original column of the input files is preserved in MTZ output; +mmCIF output keeps only the columns gemmi's MTZ-to-mmCIF conversion maps to a +``_refln`` item (see :func:`write_sf_file`). """ import hashlib @@ -71,8 +73,21 @@ def read_sf_file(path: str, cif_block: Optional[str] = None) -> rs.DataSet: def write_sf_file(ds: rs.DataSet, path: str) -> None: """Write a DataSet as MTZ or SF-mmCIF depending on the extension. - CIF output uses gemmi's MTZ-to-mmCIF conversion, with ``FreeR_flag == 0`` - written as ``_refln.status 'f'`` and negative (excluded) flags as ``'x'``. + Parameters + ---------- + ds : rs.DataSet + Reflections with cell and space group attached. + path : str + Output file; ``.mtz`` writes every column, ``.cif`` / ``.mmcif`` goes + through gemmi's MTZ-to-mmCIF conversion. + + Notes + ----- + mmCIF output silently drops columns that gemmi's default conversion does + not map to a ``_refln`` item (e.g. ``DANO`` or custom columns); write MTZ + to keep them. ``FreeR_flag == 0`` becomes ``_refln.status 'f'``, negative + (excluded) flags ``'x'`` and all other values ``'o'``, so only the + free/work split survives, not the CCP4 test-set number. """ suffix = Path(path).suffix.lower() if suffix == ".mtz": @@ -113,18 +128,46 @@ def _mark_excluded(text: str, ds: rs.DataSet) -> str: def flag_column(ds: rs.DataSet, column: Optional[str] = None) -> Optional[str]: - """Name of the R-free column in ``ds`` (``column`` if given), or None.""" + """Return the name of the R-free column in ``ds``. + + Parameters + ---------- + ds : rs.DataSet + Reflections to search. + column : str, optional + Column to look for. By default the first of :data:`FLAG_COLUMN_NAMES` + present in ``ds``. + + Returns + ------- + str or None + The column name, or None when ``ds`` has no such column. + """ if column is not None: return column if column in ds.columns else None return next((c for c in FLAG_COLUMN_NAMES if c in ds.columns), None) def excluded_rows(ds: rs.DataSet, column: Optional[str] = None) -> np.ndarray: - """Rows marked excluded (negative or missing flag, CIF ``x``), shape (N,). + """Mark the rows an R-free column excludes. + + A row is excluded when its flag is negative or missing: MTZ ``-1`` or + missing-number flags, and every mmCIF ``_refln.status`` other than ``o`` + and ``f`` (``x``, ``<``, ``-``, ``h``, ``l``), which gemmi reads as missing. + The result does not depend on whether the remaining rows form a valid + partition, so a column that excludes every row still excludes every row. + + Parameters + ---------- + ds : rs.DataSet + Reflections, N rows. + column : str, optional + R-free column; auto-detected by :func:`flag_column` by default. - All-false when ``ds`` has no R-free column. Independent of whether the - remaining rows form a valid partition, so a column that excludes every - row still excludes every row. + Returns + ------- + np.ndarray + Boolean mask, shape (N,); all False when ``ds`` has no R-free column. """ column = flag_column(ds, column) if column is None: @@ -141,11 +184,23 @@ def read_free_set(ds: rs.DataSet, column: Optional[str] = None) -> dict: (0 = free); a binary column takes its majority value as work, which covers both CCP4 ``0 = free`` and Phenix ``1 = free``. + Parameters + ---------- + ds : rs.DataSet + Reflections, N rows. + column : str, optional + R-free column; auto-detected by :func:`flag_column` by default. + Returns ------- dict ``column``, ``convention``, ``raw`` (int, -1 where excluded), - ``free`` and ``excluded`` (bool per row). + ``free`` and ``excluded`` (bool per row), each array of shape (N,). + + Raises + ------ + ValueError + If ``ds`` has no R-free column, or the column has no valid value. """ column = flag_column(ds, column) if column is None: @@ -189,6 +244,14 @@ def compare_free_sets( ) -> dict: """Report whether the existing free sets of several datasets agree. + Parameters + ---------- + datasets : dict of str to rs.DataSet + Datasets to compare, by name. + column : str, optional + R-free column to read in every dataset; auto-detected per dataset by + default. A dataset without it is reported as having no free set. + Returns ------- dict @@ -303,7 +366,18 @@ def check_compatible( def asu_hkl(ds: rs.DataSet) -> np.ndarray: - """Friedel-merged reciprocal-ASU indices for every row, shape (N, 3).""" + """Map every row of ``ds`` to its Friedel-merged reciprocal-ASU index. + + Parameters + ---------- + ds : rs.DataSet + Reflections, N rows, with a space group attached. + + Returns + ------- + np.ndarray + int64 Miller indices, shape (N, 3), in the reciprocalspaceship ASU. + """ hkl = ds.get_hkls() return rs.utils.hkl_to_asu(hkl, ds.spacegroup)[0].astype(np.int64) @@ -361,7 +435,21 @@ def _lookup(table_keys: np.ndarray, table_values: np.ndarray, keys: np.ndarray): def resolution_bins(dstar2: np.ndarray, n_bins: int) -> np.ndarray: - """Equal-count resolution bin index (0..n_bins-1) for each 1/d^2 value.""" + """Assign equal-count resolution bins. + + Parameters + ---------- + dstar2 : np.ndarray + ``1/d^2`` per reflection in Å⁻², shape (N,). + n_bins : int + Requested number of bins, clipped to ``[1, N]``. + + Returns + ------- + np.ndarray + Bin index ``0..n_bins-1`` per reflection, shape (N,), 0 at low + resolution. + """ n_bins = max(1, min(n_bins, len(dstar2))) order = np.argsort(dstar2, kind="stable") bins = np.empty(len(dstar2), dtype=np.int64) @@ -379,11 +467,23 @@ def _hash_flags(hkl: np.ndarray, n_flags: int, seed: int) -> np.ndarray: def partition_seed(keys: np.ndarray, free: np.ndarray) -> int: - """Seed derived from a free/work partition (SHA-256 of keys and free mask). + """Derive a seed from a free/work partition (SHA-256 of keys and free mask). Depends only on which unique reflections are free and which are work, not on file format, row order or flag convention, so every extension of the same deposited free set is identical. + + Parameters + ---------- + keys : np.ndarray + Unique ASU keys from :func:`hkl_keys`, shape (M,), any order. + free : np.ndarray + Boolean free mask aligned with ``keys``, shape (M,). + + Returns + ------- + int + Non-negative 63-bit seed. """ order = np.argsort(keys) digest = hashlib.sha256( @@ -401,7 +501,7 @@ def complete_flag_table( seed: int = 0, shell_size: int = 1000, ) -> Tuple[np.ndarray, np.ndarray]: - """Stratified CCP4 flags on the complete reciprocal ASU out to ``dmin``. + """Deal stratified CCP4 flags on the complete reciprocal ASU out to ``dmin``. The complete ASU is sorted by resolution (ties by index) and cut into consecutive shells of ``shell_size`` reflections. Each shell is shuffled @@ -412,6 +512,27 @@ def complete_flag_table( cell, space group, ``n_flags``, ``shell_size`` and ``seed``: a larger ``dmin`` (a later, better dataset) never changes existing flags. + A flag is tied to a reflection's rank in resolution, so it depends on the + exact cell: a relative change of 1e-5 in one cell edge already moves about + a tenth of the free set, and the ~0.05 % that separates two crystals of + one form leaves the sets nearly independent. Pass the same cell (the same + reference file) to reproduce a table. + + Parameters + ---------- + cell : gemmi.UnitCell + Unit cell; lengths in Å, angles in degrees. + spacegroup : gemmi.SpaceGroup + Space group defining the reciprocal ASU. + dmin : float + High-resolution limit in Å. + n_flags : int + Number of flag values; the free fraction is ``1 / n_flags``. + seed : int + Base seed of the per-shell shuffles. + shell_size : int + Reflections per shell, rounded down to a multiple of ``n_flags``. + Returns ------- keys : np.ndarray @@ -455,6 +576,18 @@ def reference_flags( columns free becomes ``0`` and work reflections are dealt pseudo-random values ``1..n_flags-1``. + Parameters + ---------- + ref : rs.DataSet + Reference reflections carrying an R-free column. + n_flags : int + Number of flag values for the work values of a binary column; unused + for a CCP4 column. + column : str, optional + R-free column in ``ref``; auto-detected by default. + seed : int + Seed of the pseudo-random work values of a binary column. + Returns ------- keys : np.ndarray @@ -541,8 +674,8 @@ def uniform_rfree( flags assigned to reflections the reference lacks (e.g. its missing high-resolution shells) are a deterministic function of the reference. dmin : float, optional - High-resolution limit of the flag table; the best resolution of the - inputs is used if it is finer. + High-resolution limit of the flag table in Å; the best resolution of + the inputs is used if it is finer. reference : rs.DataSet, optional Dataset whose existing flags are inherited; reflections it lacks are newly assigned. @@ -680,6 +813,21 @@ def apply_flags( Existing flag columns are dropped, or renamed ``_orig`` with ``keep_old``. Row order and all other columns are unchanged. + + Parameters + ---------- + ds : rs.DataSet + Reflections, N rows; not modified. + flags : np.ndarray + Integer ``FreeR_flag`` per row, shape (N,), aligned with ``ds``. + keep_old : bool + Keep existing flag columns under ``_orig`` instead of dropping + them. + + Returns + ------- + rs.DataSet + Copy of ``ds`` with a ``FreeR_flag`` column of MTZ type ``I``. """ out = ds.copy() for col in [c for c in out.columns if c in FLAG_COLUMN_NAMES]: @@ -727,6 +875,13 @@ def scale_columns(ds: rs.DataSet, factor: np.ndarray) -> Tuple[rs.DataSet, List[ or a ``FC``/``FCALC``/``FMODEL``-style name), phases, weights, flags and other columns are untouched. + Parameters + ---------- + ds : rs.DataSet + Reflections, N rows; not modified. + factor : np.ndarray + Amplitude scale factor per row, shape (N,), aligned with ``ds``. + Returns ------- rs.DataSet From 49894d6deb7b53f34d1b75e8c773b9ebc37ece73 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 09:43:56 +0000 Subject: [PATCH 230/250] Fit the centric factor only with both classes; give an infinite SNR full weight On all-centric data (every reflection of a centrosymmetric group) the centric factor is collinear with the constant Chebyshev term, so the Newton solve in fit_difference_power raised a singular-matrix LinAlgError that neither the q-weight fallback nor the writer catches. The factor is now fitted only when both centric and acentric reflections are usable, and stays 1 otherwise. bounded_wiener_weight returned inf/inf = NaN for a zero reported sigma, which normalise_mean_one turned into a zero weight; it now returns the limit, 1. Also: import order (ruff I001) and black on lines this branch added, shapes in the fit_difference_power docstring, the stale tau_sq in the write_results_mtz Returns section, and the changelog line. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01BKDn7EaPFPeiG4rLBKxseN --- docs/changelog.rst | 2 +- .../unit/refinement/test_difference_power.py | 40 ++++++++++++++++--- torchref/cli/collection_difference_refine.py | 25 +++++++----- torchref/cli/validate_ded.py | 7 +--- .../difference_power.py | 40 +++++++++++++------ .../refinement/targets/collection/xray.py | 12 +++--- 6 files changed, 85 insertions(+), 41 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 44910074..02f1d302 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -7,7 +7,7 @@ Unreleased - The difference MTZ groups its columns into named datasets -- ``observed``, ``difference``, ``light_model``, ``extrapolated_light``, ``two_moment`` -- with one history line describing each, so ``FWT``/``PHWT`` reads as ``/torchref/extrapolated_light/FWT`` (the extrapolated light-state map ``2*FEXT - Fc``). Labels are unchanged and Coot still auto-opens it - ``SpaceGroup.canonicalize_hkl`` gains ``sort=False``, which keeps the input row order and returns ``None`` for ``sort_indices``; the ASU mapping runs as threaded torch operations and reuses the rotated indices for the Friedel mate, so large reflection lists map faster with unchanged outputs - ``torchref.difference-map``, ``torchref.difference-refine`` and ``torchref.validate-ded`` gain ``--ded-weight {q,inverse_variance,none}`` (default ``q``) and ``--difference-gamma``. The difference MTZ now carries the unweighted ``dF``/``SIGdF`` on ``PHDELWT`` with one mean-one weight column per scheme, ``W_Q`` and ``W_InVa`` (MTZ type W), and the observed-to-model scale ``KSCALE``; ``DELFWT`` is no longer written, build the map with ``torchref.mtz2map -csf dF -cw W_Q -cphi PHDELWT``. Registered in ``torchref.maps.ded_weights`` -- Added ``torchref.refinement.model_error_estimation.difference_power``: ``fit_difference_power`` fits the expected power of a light-minus-dark difference by per-reflection maximum likelihood -- ``log S`` a Chebyshev series in resolution plus ``gamma log F_dark``, a fitted scale ``k`` on the reported sigmas and a centric factor -- with no resolution shells; given a model difference it also fits the coupling ``alpha`` and the unexplained power. ``DifferencePowerEstimator`` caches one fit for a target, ``DifferencePowerConfig`` fixes ``gamma``. ``torchref.maps.ded_weights.difference_snr`` turns it into a per-reflection signal-to-noise ratio, fitted on intensity differences when the data carry ``I``/``SIGI``, because French-Wilson amplitude sigmas overstate the noise of a difference. The ``q`` weight is ``(snr + 1/2) / (snr + 3/2)``, which down-weights a noisy reflection to no less than a third and never removes a resolution range +- Added ``torchref.refinement.model_error_estimation.difference_power``: ``fit_difference_power`` fits the expected power of a light-minus-dark difference by per-reflection maximum likelihood -- ``log S`` a Chebyshev series in resolution plus ``gamma log F_dark``, a fitted scale ``k`` on the reported sigmas and a centric factor (fitted when both centric and acentric reflections are present, so centrosymmetric data fit too) -- with no resolution shells; given a model difference it also fits the coupling ``alpha`` and the unexplained power. ``DifferencePowerEstimator`` caches one fit for a target, ``DifferencePowerConfig`` fixes ``gamma``. ``torchref.maps.ded_weights.difference_snr`` turns it into a per-reflection signal-to-noise ratio, fitted on intensity differences when the data carry ``I``/``SIGI``, because French-Wilson amplitude sigmas overstate the noise of a difference. The ``q`` weight is ``(snr + 1/2) / (snr + 3/2)``, which down-weights a noisy reflection to no less than a third and never removes a resolution range; a zero reported sigma gets full weight - Added the ``difference_sd`` collection target (``CollectionDifferenceSigmaDTarget``): the difference Gaussian centred on ``alpha * dF_calc`` with variance ``beta_model + sigma_diff^2``, both from one ``fit_difference_power`` on the free set with the reported sigmas taken as calibrated. Selected with ``torchref.difference-refine --difference-target difference_sd``; ``difference`` stays the default. Its stats report ``difference_gamma`` and ``difference_alpha_low_res``/``difference_alpha_high_res`` - ``torchref.mtz2map`` gains ``--column-weight``/``-cw`` (multiply the amplitudes by a weight column before the FFT) and ``--units {sigma,electrons,raw}``: ``electrons`` writes e/A^3 as ``(1/V) sum_h F(h) exp(-2 pi i h.x)`` with the amplitudes divided by the ``--column-scale``/``-ck`` factor (``KSCALE`` by default). ``-n``/``--normalize`` is a deprecated alias. ``Map`` and ``DifferenceMap`` take ``units`` too, and ``DifferenceMap`` an optional per-reflection ``scale`` - ``ScalerBase.multiplicative_scale()`` returns the per-reflection ``K_overall * b_overall * anisotropy`` factor, every multiplicative component of ``forward`` and none of the additive solvent term diff --git a/tests/unit/refinement/test_difference_power.py b/tests/unit/refinement/test_difference_power.py index 87c1b769..32dbc450 100644 --- a/tests/unit/refinement/test_difference_power.py +++ b/tests/unit/refinement/test_difference_power.py @@ -5,11 +5,12 @@ when the reported sigmas are uniformly inflated; a fixed exponent stays fixed; a model difference's resolution-dependent coupling and the power it leaves unexplained are recovered together; the bounded Wiener weight never falls below its floor, so no -reflection or resolution range is removed even when the data hold no signal; the sigma -scale stays within its bounds when the differences hold no noise; the fit stays defined -when the residual holds no power; the estimator caches one fit until reset and applies -its configured exponent; the fit runs under ``torch.no_grad()`` and on every available -device. +reflection or resolution range is removed even when the data hold no signal, and an +infinite SNR gets full weight; the centric factor is fitted, and recovered, only when +both centric and acentric reflections are present; the sigma scale stays within its +bounds when the differences hold no noise; the fit stays defined when the residual holds +no power; the estimator caches one fit until reset and applies its configured exponent; +the fit runs under ``torch.no_grad()`` and on every available device. """ import pytest @@ -82,10 +83,39 @@ def test_weight_never_removes_a_reflection_without_signal(): assert float(w.min()) >= 1.0 / 3.0 - 1e-6 assert float(w.max()) < 1.0 assert float(bounded_wiener_weight(torch.zeros(1), 0.0)) == 0.0 + # A zero reported sigma gives an infinite SNR: full weight, not a NaN. + assert float(bounded_wiener_weight(torch.tensor([float("inf")]), 0.5)) == 1.0 with pytest.raises(ValueError): bounded_wiener_weight(snr, -0.1) +@pytest.mark.unit +def test_centric_factor_is_fitted_only_when_both_classes_are_present(): + # Every reflection of a centrosymmetric group is centric: the factor would be + # collinear with the constant term, so it is left at one and the fit matches the + # one without flags. Several seeds, because a singular Hessian is seed-dependent. + for seed in range(4): + s = synth(n=5000, seed=seed) + all_centric = torch.ones_like(s["delta"], dtype=torch.bool) + fit = fit_difference_power( + s["delta"], s["sigma"], s["dss"], f_dark=s["f"], centric=all_centric + ) + plain = fit_difference_power(s["delta"], s["sigma"], s["dss"], f_dark=s["f"]) + assert fit.centric_factor == 1.0 + assert torch.allclose(fit.coeffs, plain.coeffs) + # A doubled centric power is recovered when both classes are present. + s = synth() + centric = torch.rand(len(s["delta"]), generator=torch.Generator().manual_seed(7)) + centric = centric < 0.3 + g = torch.Generator().manual_seed(8) + extra = torch.randn(len(s["delta"]), generator=g) * s["s_true"].sqrt() + delta = torch.where(centric, s["delta"] + extra, s["delta"]) + fit = fit_difference_power( + delta, s["sigma"], s["dss"], f_dark=s["f"], centric=centric + ) + assert fit.centric_factor == pytest.approx(2.0, rel=0.15) + + @pytest.mark.unit def test_runs_on_device(any_device): s = synth(n=5000, device=any_device) diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index e73a16db..622b1388 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -43,13 +43,13 @@ add_weights_arg, build_dual_column_names, configure_unbuffered_output, + difference_config_from_args, + intensity_difference, load_model, load_reflection_data, parse_device_str, parse_weights, register_timing, - intensity_difference, - difference_config_from_args, validate_cif_files, validate_files, ) @@ -773,11 +773,15 @@ def _np(t): types = {"FEXT": "F", "SIGFEXT": "Q", "FWT": "F", "PHWT": "P"} if verbose > 0: - print(" Bayes extrapolation rfactors:", - rfactor_work_free(data_bayes, amp_calc_bayes)) - print(f" Shrinkage: SNR from {source} differences, " - f"mean w(h) = {w_shrinkage.mean().item():.3f}, " - f"min w(h) = {w_shrinkage.min().item():.3g}") + print( + " Bayes extrapolation rfactors:", + rfactor_work_free(data_bayes, amp_calc_bayes), + ) + print( + f" Shrinkage: SNR from {source} differences, " + f"mean w(h) = {w_shrinkage.mean().item():.3f}, " + f"min w(h) = {w_shrinkage.min().item():.3g}" + ) if all_columns: amp_phased = torch.abs(F_light_extra) @@ -945,9 +949,10 @@ def write_results_mtz( Returns ------- dict - Diagnostics worth recording outside the file -- currently the Bayes shrinkage's - ``tau_sq`` and mean ``w(h)``, which say whether the default extrapolated map is - over-shrunk. Empty when no light model was given. + Diagnostics worth recording outside the file: the ``q``-weight fit under + ``ded_weights`` and, with a light model, the extrapolation shrinkage's + ``shrinkage_source`` and mean and least ``w(h)``, which say whether the default + extrapolated map is over-shrunk. """ import reciprocalspaceship as rs diff --git a/torchref/cli/validate_ded.py b/torchref/cli/validate_ded.py index b28d6360..1ce19892 100644 --- a/torchref/cli/validate_ded.py +++ b/torchref/cli/validate_ded.py @@ -38,12 +38,12 @@ add_outdir_arg, build_dual_column_names, configure_unbuffered_output, + difference_config_from_args, intensity_difference, load_model, load_reflection_data, parse_device_str, register_timing, - difference_config_from_args, validate_cif_files, validate_files, ) @@ -357,10 +357,7 @@ def setup_ded_context( "weights_by_scheme": weights_by_scheme, "ded_weight": ded_weight, "ded_weight_applied": selected.applied, - "ded_weight_diagnostics": { - k: v - for k, v in all_w["q"].diagnostics.items() - }, + "ded_weight_diagnostics": dict(all_w["q"].diagnostics), "d_spacing": d_spacing, "cell_t": cell_t, "cell_np": cell_np, diff --git a/torchref/refinement/model_error_estimation/difference_power.py b/torchref/refinement/model_error_estimation/difference_power.py index f1ec7ed1..9d511a05 100644 --- a/torchref/refinement/model_error_estimation/difference_power.py +++ b/torchref/refinement/model_error_estimation/difference_power.py @@ -57,7 +57,8 @@ DEFAULT_ALPHA_ORDER = 2 #: Degrees of freedom of the Student-t likelihood when ``robust=True``. DEFAULT_NU = 4.0 -#: Newton iterations; the problem has at most eight parameters and converges in ~10. +#: Newton iterations; the problem has at most eight parameters (eleven with a model +#: difference) and converges in ~10. MAX_ITER = 60 _GAMMA_BOUNDS = (-1.0, 3.0) #: Bounds on the sigma scale ``k``. No merge misreports its sigmas tenfold; outside @@ -174,8 +175,9 @@ def signal_power( if self.gamma != 0.0: if f_dark is None: raise ValueError("this fit uses F_dark; pass f_dark") - log_s = log_s + self.gamma * (_log_amp(f_dark.to(dss), self.amp_scale) - - self.log_f_ref) + log_s = log_s + self.gamma * ( + _log_amp(f_dark.to(dss), self.amp_scale) - self.log_f_ref + ) if centric is not None and self.centric_factor != 1.0: log_s = log_s + math.log(self.centric_factor) * centric.to(dss) power = torch.exp(log_s) * self.amp_scale**2 @@ -241,13 +243,16 @@ def fit_difference_power( d_star_sq : torch.Tensor ``1/d**2`` in A^-2, shape ``(N,)``. epsilon : torch.Tensor, optional - Reflection multiplicity; ones when omitted. + Reflection multiplicity, shape ``(N,)``; ones when omitted. f_dark : torch.Tensor, optional - Dark amplitudes for the ``F_dark**gamma`` term; without them ``gamma`` is 0. + Dark amplitudes for the ``F_dark**gamma`` term, shape ``(N,)``, on the scale + of ``delta_obs``; without them ``gamma`` is 0. centric : torch.Tensor, optional - Boolean centric flags; given, a centric power factor is fitted. + Boolean centric flags, shape ``(N,)``. A centric power factor is fitted when + the usable reflections hold both centric and acentric ones; otherwise it is 1. fit_mask : torch.Tensor, optional - Reflections entering the fit; default every finite one with positive sigma. + Boolean, shape ``(N,)``: reflections entering the fit; default every finite + one with positive sigma. delta_calc : torch.Tensor, optional Model difference, shape ``(N,)``, on the scale of ``delta_obs``. Given, the mean is ``alpha * delta_calc`` and ``S`` is the unexplained power. @@ -271,7 +276,8 @@ def fit_difference_power( DifferencePowerFit The fitted model, detached from the inputs. Standard errors are NaN for fixed parameters. Runs with gradients enabled internally, so it works under - ``torch.no_grad()``. + ``torch.no_grad()``. Every trial step reads the objective back to the host, + one GPU->CPU sync each on an accelerator. Raises ------ @@ -321,8 +327,12 @@ def fit_difference_power( log_f = log_f - log_f_ref else: log_f, log_f_ref = torch.zeros_like(d), 0.0 - has_centric = centric is not None and bool(centric.reshape(-1)[ok].any()) - cen = centric.reshape(-1).to(dev)[ok].to(dtype) if has_centric else None + cen_ok = centric.reshape(-1).to(dev)[ok] if centric is not None else None + # Fitted only when both classes are present: an all-centric set (every reflection + # of a centrosymmetric group) makes the factor collinear with the constant term + # and the Newton Hessian singular. + has_centric = cen_ok is not None and bool(cen_ok.any()) and not bool(cen_ok.all()) + cen = cen_ok.to(dtype) if has_centric else None n_c = order + 1 n_a = alpha_order + 1 if delta_calc is not None else 0 @@ -496,8 +506,9 @@ def bounded_wiener_weight( """Wiener weight with the signal-to-noise ratio floored smoothly at ``snr_floor``. ``w = (snr + snr_floor) / (snr + 1 + snr_floor)``, which lies in - ``[snr_floor / (1 + snr_floor), 1)``: a reflection is down-weighted by its noise - fraction but never removed. ``snr_floor = 0`` is the plain Wiener weight. + ``[snr_floor / (1 + snr_floor), 1)`` and is 1 at an infinite SNR (a zero sigma): + a reflection is down-weighted by its noise fraction but never removed. + ``snr_floor = 0`` is the plain Wiener weight. Parameters ---------- @@ -513,7 +524,10 @@ def bounded_wiener_weight( """ if snr_floor < 0: raise ValueError("snr_floor must be non-negative") - return (snr + snr_floor) / (snr + 1.0 + snr_floor) + w = (snr + snr_floor) / (snr + 1.0 + snr_floor) + # A zero reported sigma gives an infinite SNR, and inf/inf would be a NaN weight + # that removes the reflection; its limit is one. + return torch.where(torch.isposinf(snr), torch.ones_like(w), w) __all__ = [ diff --git a/torchref/refinement/targets/collection/xray.py b/torchref/refinement/targets/collection/xray.py index 20fa4d71..e68d2136 100644 --- a/torchref/refinement/targets/collection/xray.py +++ b/torchref/refinement/targets/collection/xray.py @@ -32,14 +32,14 @@ gaussian_per_refl, rice_per_refl, ) -from torchref.refinement.model_error_estimation.sigma_a import ( - SigmaAEstimator, - epsilon_from_hkl, -) from torchref.refinement.model_error_estimation.difference_power import ( DifferencePowerConfig, DifferencePowerEstimator, ) +from torchref.refinement.model_error_estimation.sigma_a import ( + SigmaAEstimator, + epsilon_from_hkl, +) from torchref.utils.stats import VERBOSITY_STANDARD, StatEntry, stat from ._util import common_geom @@ -316,9 +316,7 @@ def _loss_inputs(self, recalc: bool = False): finite, f_dark, f_dark[finite].median() if bool(finite.any()) else 1.0 ) alpha = fit.alpha_at(dss) - beta = fit.signal_power( - dss, epsilon=eps, f_dark=f_eval, centric=centric - ) + beta = fit.signal_power(dss, epsilon=eps, f_dark=f_eval, centric=centric) return CollectionSigmaDLossInputs( *ctx, alpha=alpha.to(dtype).detach(), beta_model=beta.to(dtype).detach() ) From 34740ffc088f11735873d51b960dc3f1a1389939 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 09:51:01 +0000 Subject: [PATCH 231/250] Import math for the flat anisotropic ADP field set_adp_mode("field_aniso", init="flat") flattens through B_eq with math.pi, which model.py never imported, so it raised NameError. Also guard the pandas name the annotations use under TYPE_CHECKING. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01BKDn7EaPFPeiG4rLBKxseN --- docs/changelog.rst | 1 + tests/unit/model/test_adp_field_mode.py | 15 +++++++++++++++ torchref/model/model.py | 6 +++++- 3 files changed, 21 insertions(+), 1 deletion(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 8e9365c4..506d884a 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -15,6 +15,7 @@ Unreleased - The per-atom-table queries moved to the context: ``model.ctx.chain_sequences``, ``model.ctx.chain_residues`` (was ``Model.get_chain_residues``), ``model.ctx.occupancy_groups`` and ``model.ctx.register_altlocs`` - ``Restraints`` no longer borrow accessors from the model and live on its context (``model.ctx.restraints``; ``model.restraints`` still builds them on first access). The constructor takes the coordinates to build over (``xyz=``) instead of ``xyz_fn``/``adp_fn``/``vdw_radii_fn``, and every evaluation takes the coordinates or B-factors it scores, e.g. ``bond_deviations(model.xyz())``. ``Model.set_restraints_cif`` is replaced by ``model.ctx.set_cif_path``, and ``Model.bond_deviations``/``angle_deviations``/``torsion_deviations_with_sigmas`` are removed. Without a cell and space group the pair list is searched in an isolated P1 box - ``ModelFT.create_from_state_dict`` restores the restraint dictionary path (``cif_path``) as ``Model`` does, and ``Refinement.create_from_state_dict`` no longer builds a stray ``Restraints`` from the model; the restored model builds its own on first access. +- ``Model.set_adp_mode("field_aniso", init="flat")`` no longer raises ``NameError``; the field starts at one isotropic U from the median equivalent isotropic B. - The difference MTZ groups its columns into named datasets -- ``observed``, ``difference``, ``light_model``, ``extrapolated_light``, ``two_moment`` -- with one history line describing each, so ``FWT``/``PHWT`` reads as ``/torchref/extrapolated_light/FWT`` (the extrapolated light-state map ``2*FEXT - Fc``). Labels are unchanged and Coot still auto-opens it - ``torchref.difference-map``, ``torchref.difference-refine`` and ``torchref.validate-ded`` gain ``--ded-weight {sigma_d,inverse_variance,none}`` and ``--sigma-d-gamma``. The difference MTZ now carries the unweighted ``DF``/``SIGDF`` on ``PHDELWT`` with one mean-one weight column per scheme, ``W_SD`` and ``W_IVW`` (MTZ type W), and the observed-to-model scale ``KSCALE``; ``DELFWT`` is no longer written, build the map with ``torchref.mtz2map -csf DF -cw W_IVW -cphi PHDELWT``. Registered in ``torchref.maps.ded_weights`` - Added the ``sigma_D`` estimator (``torchref.refinement.model_error_estimation.sigma_d``): the expected true difference power per resolution shell, ``mean(dF_obs^2) - mean(sigma^2)`` with a fitted ``F_dark^gamma`` amplitude law and DerSimonian-Laird shrinkage of the signed shell power toward a decaying exponential in ``d*^2`` (fitted on all shells, so a dataset without a difference yields no power instead of the positive half of its noise), giving the Wiener weight ``S/(S + sigma^2)`` and, with a difference model, ``alpha``/``beta_model``. Inverse-variance weights suppress the strong reflections whose difference power is 10-70x that of weak ones; on independent half-datasets a Wiener weight with the true power raised map agreement 1.2-1.8x in effective patterns. The single-dataset estimate inherits the calibration of the reported sigmas, and on the campaign TD1 data (sigmas ~1.5x too large at high resolution) it emptied 60-90 % of the shells, so inverse variance stays the default; ``sigma_d`` reports its clamped-shell count and falls back to inverse variance with a warning when every shell is empty diff --git a/tests/unit/model/test_adp_field_mode.py b/tests/unit/model/test_adp_field_mode.py index 14a8c1d1..a40637a9 100644 --- a/tests/unit/model/test_adp_field_mode.py +++ b/tests/unit/model/test_adp_field_mode.py @@ -120,6 +120,21 @@ def test_field_mode_collapses_anisotropic_atoms_first(pdb_path): assert resid < spread, "field ignored the equivalent isotropic B" +@pytest.mark.unit +def test_flat_anisotropic_field_starts_at_the_median_b_eq(pdb_path): + """``init="flat"`` on the U slot flattens through B_eq to one isotropic U.""" + model = _model(pdb_path) + beq = (8.0 * math.pi**2 / 3.0) * model.adp_u6().detach()[:, :3].sum(dim=1) + + model.set_adp_mode("field_aniso", n_nodes=8, k_neighbors=8, init="flat") + + u6 = model.adp_u6().detach() + assert torch.isfinite(u6).all() + got = (8.0 * math.pi**2 / 3.0) * u6[:, :3].sum(dim=1) + assert torch.allclose(got, beq.median().expand_as(got), rtol=1e-4) + assert torch.allclose(u6[:, 3:], torch.zeros_like(u6[:, 3:]), atol=1e-6) + + @pytest.mark.unit def test_sf_indices_and_flags_stay_consistent(pdb_path): """Everything keyed off the iso/aniso split is refreshed, not left stale.""" diff --git a/torchref/model/model.py b/torchref/model/model.py index 89583382..54df2e08 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -12,8 +12,9 @@ - f_calc/f_obs: Complex structure factors (lowercase = complex) """ +import math import warnings -from typing import Dict, Iterable, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Dict, Iterable, List, Optional, Tuple, Union import gemmi import numpy as np @@ -40,6 +41,9 @@ from torchref.utils.device_mixin import DeviceMovementMixin from torchref.utils.utils import sanitize_pdb_dataframe +if TYPE_CHECKING: + import pandas + class Model(DeviceMovementMixin, DebugMixin, nn.Module): """ From 3e95da2e95f5411b7e8c552f9611686b4c9a5acb Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 09:51:10 +0000 Subject: [PATCH 232/250] Point docs, the notebook and comments at the new restraint and hydrogen APIs The restraints guide and the targets notebook called bond_deviations() and angle_deviations() without coordinates, which now raises, and the guide still said Restraints takes an atom table. Docstrings and test comments still named strip_H and add_hydrogens. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01BKDn7EaPFPeiG4rLBKxseN --- docs/user_guide/restraints.rst | 9 +++++---- example_notebooks/targets_and_weighting.ipynb | 2 +- tests/integration/test_io_cif.py | 7 ++++--- tests/integration/test_model_operations.py | 7 ++++--- tests/unit/io/test_hkl_convention.py | 7 ++++--- tests/unit/monomer/test_link_modifications.py | 4 ++-- tests/unit/topology/test_hydrogens.py | 4 ++-- torchref/topology/restraints.py | 13 +++++++------ torchref/topology/riding.py | 5 +++-- 9 files changed, 32 insertions(+), 26 deletions(-) diff --git a/docs/user_guide/restraints.rst b/docs/user_guide/restraints.rst index de5128e8..9af51b91 100644 --- a/docs/user_guide/restraints.rst +++ b/docs/user_guide/restraints.rst @@ -22,8 +22,9 @@ extra CIF definitions *before* that first access: Restraints hold no reference to the model: every evaluation takes the coordinates (or B-factors) it scores, and the non-bonded pair list is rebuilt from the coordinates the non-bonded target passes in. ``Restraints.__init__`` -takes an atom table and the coordinates to build over (``pdb, cif_path, xyz, -cell, spacegroup, links, verbose, nonbonded``). +takes a node-only topology (``model.ctx.topology``, or +``Topology.from_table(table)``) and the coordinates to build over +(``topology, cif_path, xyz, cell, spacegroup, links, verbose, nonbonded``). Residues for which no restraints could be built are frozen in ``xyz`` rather than refined unrestrained, so a missing ligand definition shows up as an @@ -120,8 +121,8 @@ the geometry *targets*, not off ``Restraints``: .. code-block:: python # Raw per-restraint deviations and their sigmas - deviations, sigmas = restraints.bond_deviations() - deviations, sigmas = restraints.angle_deviations() + deviations, sigmas = restraints.bond_deviations(model.xyz()) + deviations, sigmas = restraints.angle_deviations(model.xyz()) # Summary statistics, keyed by component: bond, angle, torsion, planarity, # chiral, nonbonded, ramachandran diff --git a/example_notebooks/targets_and_weighting.ipynb b/example_notebooks/targets_and_weighting.ipynb index ec0fd5f1..39916d58 100644 --- a/example_notebooks/targets_and_weighting.ipynb +++ b/example_notebooks/targets_and_weighting.ipynb @@ -365,7 +365,7 @@ " name = 'custom/tight_bonds'\n", "\n", " def forward(self):\n", - " d, sigma = self.model.restraints.bond_deviations()\n", + " d, sigma = self.model.restraints.bond_deviations(self.model.xyz())\n", " return ((d / sigma) ** 2).sum()\n", "\n", "ref_custom = LBFGSRefinement(data_file=mtz_file, pdb=pdb_file, verbose=0)\n", diff --git a/tests/integration/test_io_cif.py b/tests/integration/test_io_cif.py index d91df0bc..5ebece42 100644 --- a/tests/integration/test_io_cif.py +++ b/tests/integration/test_io_cif.py @@ -58,9 +58,10 @@ def test_save_and_reload_cif(self, sample_cif_file, tmp_path): assert output_path.exists() - # add_hydrogens=False on reload: what is under test is whether the written - # file round-trips, not whether generation reruns. Regenerating on reload can - # legitimately differ, because ``write_pdb`` does not emit LINK records -- so a + # The default hydrogens="keep" on reload: what is under test is whether the + # written file round-trips, not whether generation reruns. Regenerating on + # reload can legitimately differ, because ``write_pdb`` does not emit LINK + # records -- so a # metal-coordinated nitrogen comes back with a free valence and takes a hydrogen # it did not have before. model2 = Model() diff --git a/tests/integration/test_model_operations.py b/tests/integration/test_model_operations.py index 55980f75..4cb84962 100644 --- a/tests/integration/test_model_operations.py +++ b/tests/integration/test_model_operations.py @@ -289,9 +289,10 @@ def test_model_roundtrip_pdb(self, sample_cif_file, tmp_path): output_path = tmp_path / "output.pdb" model1.write_pdb(str(output_path)) - # add_hydrogens=False on reload: what is under test is whether the written - # file round-trips, not whether generation reruns. Regenerating on reload can - # legitimately differ, because ``write_pdb`` does not emit LINK records -- so a + # The default hydrogens="keep" on reload: what is under test is whether the + # written file round-trips, not whether generation reruns. Regenerating on + # reload can legitimately differ, because ``write_pdb`` does not emit LINK + # records -- so a # metal-coordinated nitrogen comes back with a free valence and takes a hydrogen # it did not have before. model2 = Model() diff --git a/tests/unit/io/test_hkl_convention.py b/tests/unit/io/test_hkl_convention.py index c04db5ff..459e0c66 100644 --- a/tests/unit/io/test_hkl_convention.py +++ b/tests/unit/io/test_hkl_convention.py @@ -57,9 +57,10 @@ def anomalous_data(mtz_dir, tmp_path): def _model(pdb_dir, data): - # strip_H: what is under test is the phase convention, and the absolute check - # compares against a gemmi calculation that calls ``remove_hydrogens``. Letting - # torchref generate hydrogens would have it computing a different structure. + # hydrogens="strip": what is under test is the phase convention, and the absolute + # check compares against a gemmi calculation that calls ``remove_hydrogens``. + # Keeping or generating hydrogens would have torchref computing a different + # structure. m = ModelFT(verbose=0, max_res=2.0, hydrogens="strip") m.load_pdb(str(pdb_dir / f"{CODE}.pdb")) m.cell, m.spacegroup = data.cell, data.spacegroup diff --git a/tests/unit/monomer/test_link_modifications.py b/tests/unit/monomer/test_link_modifications.py index 95ffc459..076b84d3 100644 --- a/tests/unit/monomer/test_link_modifications.py +++ b/tests/unit/monomer/test_link_modifications.py @@ -226,8 +226,8 @@ def test_peptide_modifications_carry_the_linked_backbone_targets(): def _built(pdb_path, strip_H=True): """Build a model's restraints and return ``(model, table accessor)``. - ``add_hydrogens=False``: these tests read the restraint targets of the hydrogens the - file carries. Generating more would add a chain-terminal ``CA-N-H``, which correctly + ``hydrogens="keep"`` or ``"strip"``, never ``"add"``: these tests read the restraint + targets of the hydrogens the file carries. Generating more would add a chain-terminal ``CA-N-H``, which correctly keeps the free-amino-acid target of 109.6 degrees rather than the linked 118.7 and so is outside what they assert. """ diff --git a/tests/unit/topology/test_hydrogens.py b/tests/unit/topology/test_hydrogens.py index b0c1e2a5..10b6385b 100644 --- a/tests/unit/topology/test_hydrogens.py +++ b/tests/unit/topology/test_hydrogens.py @@ -30,8 +30,8 @@ def built(pdb_dir): def _build(code): if code not in cache: - # add_hydrogens=False: these tests exercise generation itself, so the model - # has to arrive without the hydrogens the loader would otherwise add. + # hydrogens="strip": these tests exercise generation itself, so the model + # has to arrive without hydrogens. model = Model(verbose=0, hydrogens="strip") model.load_pdb(str(pdb_dir / f"{code}.pdb")) model.ctx.set_cif_path(None) diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index 95b10e2b..e0b35a26 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -1,8 +1,9 @@ """The restraint layer over a topology, and what it takes to build one. -:class:`Restraints` is the orchestrator. Given an atom table it resolves the monomer -dictionaries, builds the :class:`~torchref.topology.topology.Topology`, layers the ideal -values over its edges, derives the non-bonded pair list, and exposes the whole thing as +:class:`Restraints` is the orchestrator. Given a node-only +:class:`~torchref.topology.topology.Topology` and the coordinates to build over, it +resolves the monomer dictionaries, connects the topology, layers the ideal values over +its edges, derives the non-bonded pair list, and exposes the whole thing as ``restraints[edge_type][origin][property]`` -- three dict lookups into a mapping assembled once, because the geometry targets read it on every iteration. @@ -14,7 +15,7 @@ it is held apart from the rest; * the Ramachandran map, a residue-level product of the same build. -Deliberately decoupled from :class:`~torchref.model.Model`: it takes an atom table and +Deliberately decoupled from :class:`~torchref.model.Model`: it takes a topology and holds no reference back to whatever owns the coordinates. Every evaluation takes the coordinates (or ADPs) it scores as an argument, and the pair list is rebuilt from the coordinates it is handed, so the same object serves any model that shares the atom set. @@ -165,8 +166,8 @@ def _multi_atom_resnames(topology) -> list: def _riding_table(self, xyz: torch.Tensor) -> pd.DataFrame: """The identity-plus-coordinates table :mod:`torchref.topology.riding` reads. - That module still takes an atom table; it goes when the phantom-hydrogen path - is deleted. + That module takes an atom table rather than a topology, so one is assembled + here from :attr:`topology` and ``xyz`` (Cartesian, Å, shape ``(n_atoms, 3)``). """ columns = self.topology.columns() coords = xyz.detach().cpu().numpy() diff --git a/torchref/topology/riding.py b/torchref/topology/riding.py index f36a681c..15662c80 100644 --- a/torchref/topology/riding.py +++ b/torchref/topology/riding.py @@ -1,6 +1,7 @@ """Riding hydrogens: the sterics of hydrogens a model does not carry. -For a model loaded with ``strip_H=True``, whose atoms are heavy only. A static map +For a model whose atoms are heavy only (loaded with ``hydrogens="strip"``, or from a +file without hydrogens). A static map built once at restraint-construction time says how to reconstruct each absent hydrogen from its parent and the parent's bonded neighbours; ``place_riding_hydrogens`` then produces those positions in one vectorized pass at every non-bonded evaluation and @@ -298,7 +299,7 @@ def build_hydrogen_topology( Parameters ---------- pdb : pd.DataFrame - Heavy-atom DataFrame (``strip_H=True``). + Heavy-atom DataFrame (no hydrogen rows). device : torch.device Target device for tensors. verbose : int From 51957b05f83367bd63b141e59782c40d3128f849 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 09:51:10 +0000 Subject: [PATCH 233/250] Build test restraints over model.ctx.topology, not the deprecated model.pdb The new restraint tests rebuilt a topology from Model.pdb, the view AGENTS.md says never to read; the context already holds it. Black the two new test modules. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01BKDn7EaPFPeiG4rLBKxseN --- tests/fixtures/objects.py | 3 +- .../functional/test_restraints_functional.py | 53 ++++--------------- tests/integration/test_refinement_pipeline.py | 5 +- tests/unit/model/test_strip_altlocs.py | 10 +++- tests/unit/topology/test_topology_identity.py | 19 +++++-- 5 files changed, 35 insertions(+), 55 deletions(-) diff --git a/tests/fixtures/objects.py b/tests/fixtures/objects.py index 47f67141..95aa15e9 100644 --- a/tests/fixtures/objects.py +++ b/tests/fixtures/objects.py @@ -111,10 +111,9 @@ def initialized_scaler(model_and_data: dict[str, Any]) -> Scaler: def model_with_restraints(loaded_model: Model) -> dict[str, Any]: """Build restraints around a fresh model.""" from torchref.topology.restraints import Restraints - from torchref.topology.topology import Topology restraints = Restraints( - topology=Topology.from_table(loaded_model.pdb), + topology=loaded_model.ctx.topology, xyz=loaded_model.xyz(), verbose=0, ) diff --git a/tests/functional/test_restraints_functional.py b/tests/functional/test_restraints_functional.py index 44ae43f4..25731cdc 100644 --- a/tests/functional/test_restraints_functional.py +++ b/tests/functional/test_restraints_functional.py @@ -16,14 +16,11 @@ def test_build_restraints_from_cif(self, sample_cif_file): """Test building restraints from a real CIF file.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints - from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) - restraints = Restraints( - topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 - ) + restraints = Restraints(topology=model.ctx.topology, xyz=model.xyz(), verbose=0) # Should have built some restraints assert restraints.restraints is not None @@ -34,14 +31,11 @@ def test_bond_restraints_built(self, sample_cif_file): """Test that bond restraints are built correctly.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints - from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) - restraints = Restraints( - topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 - ) + restraints = Restraints(topology=model.ctx.topology, xyz=model.xyz(), verbose=0) # Check bond restraints exist assert "bond" in restraints.restraints @@ -72,14 +66,11 @@ def test_angle_restraints_built(self, sample_cif_file): """Test that angle restraints are built correctly.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints - from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) - restraints = Restraints( - topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 - ) + restraints = Restraints(topology=model.ctx.topology, xyz=model.xyz(), verbose=0) # Check angle restraints exist assert "angle" in restraints.restraints @@ -103,14 +94,11 @@ def test_torsion_restraints_built(self, sample_cif_file): """Test that torsion restraints are built correctly.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints - from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) - restraints = Restraints( - topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 - ) + restraints = Restraints(topology=model.ctx.topology, xyz=model.xyz(), verbose=0) # Check torsion restraints exist assert "torsion" in restraints.restraints @@ -132,14 +120,11 @@ def test_plane_restraints_built(self, sample_cif_file): """Test that plane restraints are built correctly.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints - from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) - restraints = Restraints( - topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 - ) + restraints = Restraints(topology=model.ctx.topology, xyz=model.xyz(), verbose=0) # Check plane restraints exist assert "plane" in restraints.restraints @@ -164,14 +149,11 @@ def test_bond_deviations(self, sample_cif_file): """Test computing bond length deviations.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints - from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) - restraints = Restraints( - topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 - ) + restraints = Restraints(topology=model.ctx.topology, xyz=model.xyz(), verbose=0) # Compute bond deviations if hasattr(restraints, "bond_deviations"): @@ -188,14 +170,11 @@ def test_angle_deviations(self, sample_cif_file): """Test computing angle deviations.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints - from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) - restraints = Restraints( - topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 - ) + restraints = Restraints(topology=model.ctx.topology, xyz=model.xyz(), verbose=0) # Compute angle deviations if hasattr(restraints, "angle_deviations"): @@ -213,11 +192,10 @@ class TestRestraintsMultipleStructures: def test_restraints_multiple_cif_files(self, compatibility_model): """Each extended crystal supplies bond and angle restraints.""" from torchref.topology.restraints import Restraints - from torchref.topology.topology import Topology model = compatibility_model restraints = Restraints( - topology=Topology.from_table(model.pdb), + topology=model.ctx.topology, xyz=model.xyz(), verbose=0, ) @@ -233,14 +211,11 @@ def test_restraints_device_movement(self, sample_cif_file, cpu_device): """Test moving restraints to different devices.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints - from torchref.topology.topology import Topology model = Model(device=cpu_device) model.load_cif(str(sample_cif_file)) - restraints = Restraints( - topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 - ) + restraints = Restraints(topology=model.ctx.topology, xyz=model.xyz(), verbose=0) # Check that tensors are on the correct device if "bond" in restraints.restraints and "intra" in restraints.restraints["bond"]: @@ -256,14 +231,11 @@ def test_cif_dict_loaded(self, sample_cif_file): """Test that CIF dictionary is loaded correctly.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints - from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) - restraints = Restraints( - topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 - ) + restraints = Restraints(topology=model.ctx.topology, xyz=model.xyz(), verbose=0) # CIF dict should be populated with residue restraints assert restraints.cif_dict is not None @@ -283,14 +255,11 @@ def test_unique_residues_detected(self, sample_cif_file): """Test that unique residues are detected from model.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints - from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) - restraints = Restraints( - topology=Topology.from_table(model.pdb), xyz=model.xyz(), verbose=0 - ) + restraints = Restraints(topology=model.ctx.topology, xyz=model.xyz(), verbose=0) # Should have detected unique residues assert restraints.unique_residues is not None diff --git a/tests/integration/test_refinement_pipeline.py b/tests/integration/test_refinement_pipeline.py index 2947f4a7..08730a45 100644 --- a/tests/integration/test_refinement_pipeline.py +++ b/tests/integration/test_refinement_pipeline.py @@ -45,15 +45,12 @@ def test_restraints_from_model(self, sample_cif_file): """Test building restraints from a loaded model.""" from torchref.model.model import Model from torchref.topology.restraints import Restraints - from torchref.topology.topology import Topology model = Model() model.load_cif(str(sample_cif_file)) # Build restraints - restraints = Restraints( - topology=Topology.from_table(model.pdb), xyz=model.xyz() - ) + restraints = Restraints(topology=model.ctx.topology, xyz=model.xyz()) # Should have some restraints assert restraints.restraints is not None diff --git a/tests/unit/model/test_strip_altlocs.py b/tests/unit/model/test_strip_altlocs.py index 2b508e4f..6a0bc04d 100644 --- a/tests/unit/model/test_strip_altlocs.py +++ b/tests/unit/model/test_strip_altlocs.py @@ -44,7 +44,11 @@ def test_insertion_coded_residues_are_not_conformers_of_each_other(table, baseli first, second = _residue_rows(df, 10), _residue_rows(df, 11) df.loc[first, ["altloc", "occupancy"]] = ["A", 0.6] df.loc[second, ["resseq", "icode", "altloc", "occupancy", "resname"]] = [ - 10, "A", "B", 0.4, df.loc[first[0], "resname"] + 10, + "A", + "B", + 0.4, + df.loc[first[0], "resname"], ] model = _model(df, cell, sg) @@ -81,7 +85,9 @@ def test_the_higher_occupancy_conformer_survives(table, baseline): kept = kept[(kept.chainid == "A") & (kept.resseq == 20)] original = df[(df.chainid == "A") & (df.resseq == 20) & (df.altloc != "B")] np.testing.assert_allclose( - kept[["x", "y", "z"]].to_numpy(), original[["x", "y", "z"]].to_numpy(), atol=1e-4 + kept[["x", "y", "z"]].to_numpy(), + original[["x", "y", "z"]].to_numpy(), + atol=1e-4, ) diff --git a/tests/unit/topology/test_topology_identity.py b/tests/unit/topology/test_topology_identity.py index d49aca64..4664f348 100644 --- a/tests/unit/topology/test_topology_identity.py +++ b/tests/unit/topology/test_topology_identity.py @@ -108,7 +108,8 @@ def test_water_and_polymer_masks(pdb_dir): np.testing.assert_array_equal(nodes.is_water, (df.resname == "HOH").to_numpy()) np.testing.assert_array_equal(nodes.is_polymer, (df.ATOM == "ATOM").to_numpy()) np.testing.assert_array_equal( - nodes.atoms.is_hydrogen.cpu().numpy(), (df.element.str.strip() == "H").to_numpy() + nodes.atoms.is_hydrogen.cpu().numpy(), + (df.element.str.strip() == "H").to_numpy(), ) @@ -118,7 +119,9 @@ def test_hydrogen_insertion_matches_the_table_insertion(pdb_dir, code): """Same row maps and the same identity, row for row, as the table-level insertion.""" model = Model(verbose=0, hydrogens="strip").load_pdb(str(pdb_dir / f"{code}.pdb")) restraints = model.ctx.build_restraints(model.xyz(), nonbonded=False, verbose=0) - plan = plan_hydrogens(restraints.topology, restraints.cif_dict, model.xyz().detach()) + plan = plan_hydrogens( + restraints.topology, restraints.cif_dict, model.xyz().detach() + ) assert plan.n_hydrogens > 0 augmented, old_to_new, plan_to_new = augment_atom_table_with_maps( @@ -135,7 +138,9 @@ def test_hydrogen_insertion_matches_the_table_insertion(pdb_dir, code): expected = Topology.from_table(augmented) for key, value in expected.columns().items(): np.testing.assert_array_equal(nodes.columns()[key], value, err_msg=key) - np.testing.assert_array_equal(nodes.residues.atom_start, expected.residues.atom_start) + np.testing.assert_array_equal( + nodes.residues.atom_start, expected.residues.atom_start + ) @pytest.mark.unit @@ -151,7 +156,9 @@ def test_padded_strings_read_like_clean_ones(pdb_dir): @pytest.mark.unit -@pytest.mark.parametrize("dropped", [["icode"], ["altloc"], ["ATOM"], ["charge"], ["element"]]) +@pytest.mark.parametrize( + "dropped", [["icode"], ["altloc"], ["ATOM"], ["charge"], ["element"]] +) def test_optional_columns_fall_back_to_defaults(pdb_dir, dropped): df = _table(pdb_dir, "1DAW") full = Topology.from_table(df) @@ -161,7 +168,9 @@ def test_optional_columns_fall_back_to_defaults(pdb_dir, dropped): column = {"ATOM": "is_hetatm"}.get(dropped[0], dropped[0]) assert (reduced.columns()[column] == defaults[dropped[0]]).all() if dropped[0] in ("charge", "element"): - np.testing.assert_array_equal(reduced.residues.atom_start, full.residues.atom_start) + np.testing.assert_array_equal( + reduced.residues.atom_start, full.residues.atom_start + ) @pytest.mark.unit From 1a7a58953abaab1fdaf389a70bd1af1d977f9d42 Mon Sep 17 00:00:00 2001 From: Hans Peter Seidel <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 30 Sep 2026 12:17:40 +0200 Subject: [PATCH 234/250] Reject lossy SF-CIF conversions and preserve supported data --- docs/changelog.rst | 1 + docs/user_guide/cli.rst | 9 +- tests/integration/test_cli_uniform_rfree.py | 20 +++ tests/unit/io/test_uniform_rfree.py | 67 +++++++++ torchref/io/rfree.py | 148 +++++++++++++++++--- 5 files changed, 224 insertions(+), 21 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 06e54761..acda6d98 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- SF-CIF export preserves supported numerical columns and free-flag values, and rejects unsupported columns before writing; use MTZ to retain custom columns or saved original flags. - Added ``torchref.uniform-rfree``: gives any number of MTZ / SF-mmCIF files of one cell and space group a shared CCP4 ``FreeR_flag``. - An existing free set (CCP4, Phenix or mmCIF convention) is inherited by default and extended, at its own fraction, to reflections it lacks, including those beyond its resolution. That extension is seeded with a hash of the reference's free/work partition, so it is reproducible whatever the file format or convention. New sets are stratified by resolution shell on the complete ASU and depend only on cell, space group and seed. - Options: ``--check`` reports whether the inputs' free sets agree; ``--max-free`` caps the size of a new set; ``--scale`` optionally scales the datasets together with ``DatasetCollection.scale``. diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index 7a81521d..864e1bb1 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -138,10 +138,11 @@ reflection is free in one dataset and work in another. in the file that marked them, as ``-1`` / ``x``. They are not copied to the other files. - **Output.** MTZ output keeps every input column; existing flag columns are - replaced by a CCP4 ``FreeR_flag`` (``0..N-1``, ``0`` = free). mmCIF output - keeps only the columns gemmi's MTZ-to-mmCIF conversion maps to a ``_refln`` - item (e.g. ``DANO`` and custom columns are dropped), and writes - ``_refln.status`` ``f``/``o``/``x``, so only the free/work split is kept. + replaced unless ``--keep-old-flags`` retains them as ``_orig``. SF-mmCIF + output retains supported mapped measurements and numeric free-flag values, + including CCP4 work-set numbers. Unsupported columns, including saved original + flag columns, cause an error before writing the CIF; use MTZ for these columns. + Standard CIF aliases may rename measurements (for example ``I`` to ``IMEAN``). .. code-block:: bash diff --git a/tests/integration/test_cli_uniform_rfree.py b/tests/integration/test_cli_uniform_rfree.py index 0eee9eff..ad957907 100644 --- a/tests/integration/test_cli_uniform_rfree.py +++ b/tests/integration/test_cli_uniform_rfree.py @@ -236,3 +236,23 @@ def test_unusable_reference_is_skipped_or_rejected(inputs, tmp_path): res = _run(allwork, bare, "--reference", "allwork", "-o", tmp_path / "explicit") assert res.returncode == 1 assert "no reflection as free" in res.stderr and "Traceback" not in res.stderr + + +@pytest.mark.parametrize("suffix", ["mtz", "cif"]) +def test_keep_old_flags_preserved_or_rejected(mtz_dir, tmp_path, suffix): + """The CLI either retains requested original flags or reports unsupported export.""" + source = mtz_dir / "1DAW.mtz" + out = tmp_path / "out" + res = _run(source, "-o", out, "--fresh", "--keep-old-flags", "--format", suffix) + path = out / ("1DAW_rfree." + suffix) + if suffix == "cif": + assert res.returncode == 1 + assert "FreeR_flag_orig" in res.stderr and "use MTZ" in res.stderr + assert not path.exists() + else: + assert res.returncode == 0, res.stderr + original = rs.read_mtz(str(source)) + restored = rfree.read_sf_file(str(path)) + np.testing.assert_array_equal( + restored.FreeR_flag_orig.to_numpy(), original.FreeR_flag.to_numpy() + ) diff --git a/tests/unit/io/test_uniform_rfree.py b/tests/unit/io/test_uniform_rfree.py index 493fc230..e552fcf4 100644 --- a/tests/unit/io/test_uniform_rfree.py +++ b/tests/unit/io/test_uniform_rfree.py @@ -279,3 +279,70 @@ def test_hkl_keys_reject_indices_beyond_the_encoding(): def test_max_free_must_be_positive(small): with pytest.raises(ValueError, match="max_free"): rfree.uniform_rfree({"x": small}, max_free=0) + + +@pytest.mark.parametrize("suffix", [".mtz", ".cif"]) +def test_supported_columns_and_numeric_flags_roundtrip(mtz_dir, tmp_path, suffix): + """Amplitude, intensity, sigma and numeric flag values survive SF export.""" + ds = rs.read_mtz(str(mtz_dir / "1DAW.mtz")) + flags, _ = rfree.uniform_rfree({"data": ds}) + out = rfree.apply_flags(ds, flags["data"]) + path = tmp_path / ("data" + suffix) + rfree.write_sf_file(out, str(path)) + restored = rfree.read_sf_file(str(path)) + np.testing.assert_array_equal(restored.get_hkls(), out.get_hkls()) + for col in out.columns: + alias = ( + {"I": "IMEAN", "SIGI": "SIGIMEAN"}.get(col, col) + if suffix == ".cif" + else col + ) + np.testing.assert_array_equal(restored[alias].to_numpy(), out[col].to_numpy()) + + +@pytest.mark.parametrize("extra", ["EXTRA_F", "FreeR_flag_orig", "I_DUP"]) +def test_cif_rejects_unmapped_columns_without_touching_destination( + mtz_dir, tmp_path, extra +): + """Custom and original-flag columns must not disappear during conversion.""" + ds = rs.read_mtz(str(mtz_dir / "1DAW.mtz")) + flags, _ = rfree.uniform_rfree({"data": ds}) + ds = rfree.apply_flags(ds, flags["data"], keep_old=extra == "FreeR_flag_orig") + if extra != "FreeR_flag_orig": + ds[extra] = rs.DataSeries( + np.arange(len(ds)), index=ds.index, dtype="J" if extra == "I_DUP" else "F" + ) + path = tmp_path / "data.cif" + path.write_text("existing destination") + with pytest.raises(ValueError, match=extra): + rfree.write_sf_file(ds, str(path)) + assert path.read_text() == "existing destination" + mtz = tmp_path / "data.mtz" + rfree.write_sf_file(ds, str(mtz)) + restored = rfree.read_sf_file(str(mtz)) + assert set(restored.columns) == set(ds.columns) + np.testing.assert_array_equal(restored[extra].to_numpy(), ds[extra].to_numpy()) + + +def test_cif_read_rejects_unmapped_measurement_column(mtz_dir, tmp_path): + """An unrecognised CIF measurement cannot silently disappear on input.""" + ds = rs.read_mtz(str(mtz_dir / "1DAW.mtz")) + path = tmp_path / "data.cif" + rfree.write_sf_file(ds, str(path)) + doc = gemmi.cif.read(str(path)) + doc[0].find_loop("_refln.index_h").get_loop().add_columns(["_refln.EXTRA_F"], "1") + doc.write_file(str(path)) + with pytest.raises(ValueError, match="EXTRA_F"): + rfree.read_sf_file(str(path)) + + +def test_cif_read_rejects_colliding_measurement_aliases(mtz_dir, tmp_path): + """Two measurements cannot silently collapse into one conventional MTZ column.""" + ds = rs.read_mtz(str(mtz_dir / "1DAW.mtz")) + path = tmp_path / "data.cif" + rfree.write_sf_file(ds, str(path)) + doc = gemmi.cif.read(str(path)) + doc[0].find_loop("_refln.index_h").get_loop().add_columns(["_refln.F_meas"], "1") + doc.write_file(str(path)) + with pytest.raises(ValueError, match="F_meas"): + rfree.read_sf_file(str(path)) diff --git a/torchref/io/rfree.py b/torchref/io/rfree.py index c16b9496..9d6c07ee 100644 --- a/torchref/io/rfree.py +++ b/torchref/io/rfree.py @@ -13,11 +13,12 @@ All functions here operate on :class:`reciprocalspaceship.DataSet` objects so that every original column of the input files is preserved in MTZ output; -mmCIF output keeps only the columns gemmi's MTZ-to-mmCIF conversion maps to a -``_refln`` item (see :func:`write_sf_file`). +mmCIF output preserves supported mapped measurement columns and rejects +unsupported columns before writing (see :func:`write_sf_file`). """ import hashlib +import re from pathlib import Path from typing import Dict, List, Optional, Tuple @@ -29,6 +30,46 @@ FREE_COLUMN = "FreeR_flag" +# Gemmi merged-reflection mappings, including aliases accepted on input. +_CIF_COLUMNS = { + "pdbx_r_free_flag": ("FreeR_flag", "I"), + "status": ("FreeR_flag", "I"), + "intensity_meas": ("IMEAN", "J"), + "F_squared_meas": ("IMEAN", "J"), + "intensity_sigma": ("SIGIMEAN", "Q"), + "F_squared_sigma": ("SIGIMEAN", "Q"), + "pdbx_I_plus": ("I(+)", "K"), + "pdbx_I_plus_sigma": ("SIGI(+)", "M"), + "pdbx_I_minus": ("I(-)", "K"), + "pdbx_I_minus_sigma": ("SIGI(-)", "M"), + "F_meas": ("FP", "F"), + "F_meas_au": ("FP", "F"), + "F_meas_sigma": ("SIGFP", "Q"), + "F_meas_sigma_au": ("SIGFP", "Q"), + "pdbx_F_plus": ("F(+)", "G"), + "pdbx_F_plus_sigma": ("SIGF(+)", "L"), + "pdbx_F_minus": ("F(-)", "G"), + "pdbx_F_minus_sigma": ("SIGF(-)", "L"), + "pdbx_anom_difference": ("DP", "D"), + "pdbx_anom_difference_sigma": ("SIGDP", "Q"), + "F_calc": ("FC", "F"), + "F_calc_au": ("FC", "F"), + "phase_calc": ("PHIC", "P"), + "pdbx_F_calc_with_solvent": ("F-model", "F"), + "pdbx_phase_calc_with_solvent": ("PHIF-model", "P"), + "fom": ("FOM", "W"), + "weight": ("FOM", "W"), + "pdbx_HL_A_iso": ("HLA", "A"), + "pdbx_HL_B_iso": ("HLB", "A"), + "pdbx_HL_C_iso": ("HLC", "A"), + "pdbx_HL_D_iso": ("HLD", "A"), + "pdbx_FWT": ("FWT", "F"), + "pdbx_PHWT": ("PHWT", "P"), + "pdbx_DELFWT": ("DELFWT", "F"), + "pdbx_DELPHWT": ("PHDELWT", "P"), +} + + # Existing flag columns that are replaced on output. FLAG_COLUMN_NAMES = tuple(dict.fromkeys([*MTZReader.RFREE_FLAG_NAMES, FREE_COLUMN])) @@ -41,7 +82,12 @@ def read_sf_file(path: str, cif_block: Optional[str] = None) -> rs.DataSet: - """Read an MTZ or SF-mmCIF file into a DataSet, keeping every column. + """Read MTZ columns or supported merged SF-mmCIF measurements. + + CIF measurement aliases use conventional MTZ labels. Crystal, wavelength and + scale-group identifiers are metadata rather than MTZ measurements. Unsupported + measurement columns and aliases that collide on one MTZ label raise ValueError. + Parameters ---------- @@ -65,7 +111,44 @@ def read_sf_file(path: str, cif_block: Optional[str] = None) -> rs.DataSet: blocks = [b for b in blocks if b.block.name == cif_block] if not blocks: raise ValueError(f"No reflection block found in {path}") - mtz = gemmi.CifToMtz().convert_block_to_mtz(blocks[0]) + block = blocks[0] + tags = [tag.rsplit(".", 1)[-1] for tag in block.column_labels()] + metadata = { + "index_h", + "index_k", + "index_l", + "crystal_id", + "wavelength_id", + "scale_group_code", + } + unsupported = set(tags) - set(_CIF_COLUMNS) - metadata + if unsupported: + raise ValueError( + "Unsupported SF-CIF columns would be lost: " + + ", ".join(sorted(unsupported)) + ) + by_label = {} + for tag in tags: + if tag in _CIF_COLUMNS: + by_label.setdefault(_CIF_COLUMNS[tag][0], []).append(tag) + collisions = [ + ", ".join(group) + for label, group in by_label.items() + if label != FREE_COLUMN and len(group) > 1 + ] + if collisions: + raise ValueError( + "CIF columns map to the same MTZ label: " + "; ".join(collisions) + ) + converter = gemmi.CifToMtz() + converter.spec_lines = [ + f"{tag} {label} {kind} 1" + (" o=1,f=0,x=-1" if tag == "status" else "") + for tag in tags + if tag in _CIF_COLUMNS + for label, kind in [_CIF_COLUMNS[tag]] + if tag != "status" or "pdbx_r_free_flag" not in tags + ] + mtz = converter.convert_block_to_mtz(block) return rs.io.from_gemmi(mtz) raise ValueError(f"Unsupported structure-factor format: {path}") @@ -73,21 +156,23 @@ def read_sf_file(path: str, cif_block: Optional[str] = None) -> rs.DataSet: def write_sf_file(ds: rs.DataSet, path: str) -> None: """Write a DataSet as MTZ or SF-mmCIF depending on the extension. + CIF output preserves supported numerical columns and numeric free flags; + ``FreeR_flag == 0`` also writes status ``f`` and negative flags status ``x``. + Standard CIF aliases may rename columns (for example I to IMEAN). + Columns without a supported CIF mapping, including saved original flags, + raise ValueError before the destination is written; use MTZ for these. + Parameters ---------- - ds : rs.DataSet - Reflections with cell and space group attached. + ds : reciprocalspaceship.DataSet + Reflections with MTZ column types, cell and space group. path : str - Output file; ``.mtz`` writes every column, ``.cif`` / ``.mmcif`` goes - through gemmi's MTZ-to-mmCIF conversion. - - Notes - ----- - mmCIF output silently drops columns that gemmi's default conversion does - not map to a ``_refln`` item (e.g. ``DANO`` or custom columns); write MTZ - to keep them. ``FreeR_flag == 0`` becomes ``_refln.status 'f'``, negative - (excluded) flags ``'x'`` and all other values ``'o'``, so only the - free/work split survives, not the CCP4 test-set number. + Output MTZ or SF-mmCIF filename. + + Raises + ------ + ValueError + If the output format or a CIF column mapping is unsupported. """ suffix = Path(path).suffix.lower() if suffix == ".mtz": @@ -95,7 +180,36 @@ def write_sf_file(ds: rs.DataSet, path: str) -> None: elif suffix in (".cif", ".mmcif"): converter = gemmi.MtzToCif() converter.free_flag_value = 0 - text = converter.write_cif_to_string(ds.to_gemmi()) + converter.skip_empty = False + converter.skip_negative_sigi = False + mtz = ds.to_gemmi() + text = converter.write_cif_to_string(mtz) + # Gemmi reports the columns it selected; default recipes can choose only + # one of several columns with the same MTZ type or familiar label. + mappings = re.findall(r"^# .* / (\S+) -> (\S+)$", text, re.MULTILINE) + selected = {label for label, _ in mappings} + unsupported = set(ds.columns) - selected + unsupported.update( + label + for label, tag in mappings + if label in ds.columns and tag not in _CIF_COLUMNS + ) + if unsupported: + raise ValueError( + "Unsupported CIF output columns would be lost: " + + ", ".join(sorted(unsupported)) + + "; use MTZ output instead" + ) + types = {c.label: c.type for c in mtz.columns} + converter.spec_lines = [ + f"{label} {types[label]} {tag} {'S' if tag == 'status' else '.9g'}" + for label, tag in mappings + ] + if FREE_COLUMN in ds.columns: + # status encodes only free/work, whereas CCP4 flags carry work-set + # numbers too. Retain those numbers in the standard numeric tag. + converter.spec_lines += [f"{FREE_COLUMN} I pdbx_r_free_flag .9g"] + text = converter.write_cif_to_string(mtz) if FREE_COLUMN in ds.columns: text = _mark_excluded(text, ds) Path(path).write_text(text) From d07432d49ffbcef18895605155ea9d4ab635f440 Mon Sep 17 00:00:00 2001 From: Hans Peter Seidel <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 30 Sep 2026 12:19:24 +0200 Subject: [PATCH 235/250] Propagate correlated measurement uncertainty in extrapolated amplitudes --- docs/changelog.rst | 1 + tests/unit/maps/test_ded_weights.py | 86 +++++++++++++++++++- torchref/cli/collection_difference_refine.py | 58 ++++++++----- 3 files changed, 123 insertions(+), 22 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 02f1d302..d6ffa112 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Extrapolated amplitude sigmas propagate independent dark and light measurement errors with their shared-difference covariance and phase-aware derivatives, for finite shrinkage and unshrunk fallback alike; they describe measurement uncertainty rather than latent-difference posterior variance. - The difference MTZ groups its columns into named datasets -- ``observed``, ``difference``, ``light_model``, ``extrapolated_light``, ``two_moment`` -- with one history line describing each, so ``FWT``/``PHWT`` reads as ``/torchref/extrapolated_light/FWT`` (the extrapolated light-state map ``2*FEXT - Fc``). Labels are unchanged and Coot still auto-opens it - ``SpaceGroup.canonicalize_hkl`` gains ``sort=False``, which keeps the input row order and returns ``None`` for ``sort_indices``; the ASU mapping runs as threaded torch operations and reuses the rotated indices for the Friedel mate, so large reflection lists map faster with unchanged outputs - ``torchref.difference-map``, ``torchref.difference-refine`` and ``torchref.validate-ded`` gain ``--ded-weight {q,inverse_variance,none}`` (default ``q``) and ``--difference-gamma``. The difference MTZ now carries the unweighted ``dF``/``SIGdF`` on ``PHDELWT`` with one mean-one weight column per scheme, ``W_Q`` and ``W_InVa`` (MTZ type W), and the observed-to-model scale ``KSCALE``; ``DELFWT`` is no longer written, build the map with ``torchref.mtz2map -csf dF -cw W_Q -cphi PHDELWT``. Registered in ``torchref.maps.ded_weights`` diff --git a/tests/unit/maps/test_ded_weights.py b/tests/unit/maps/test_ded_weights.py index 751c6534..ef165916 100644 --- a/tests/unit/maps/test_ded_weights.py +++ b/tests/unit/maps/test_ded_weights.py @@ -224,7 +224,14 @@ def test_extrapolated_shrinkage_keeps_every_reflection_and_ignores_occupancy(): sig_dark = torch.full((n,), 0.5) out = { f: compute_bayes_extrapolated_amplitudes( - f_dark, f_light, sig_dark, phi, phi, f, snr=snr, noise=noise + f_dark, + f_light, + sig_dark, + phi, + phi, + f, + snr=snr, + sig_light=torch.sqrt(noise**2 - sig_dark**2), ) for f in (0.2, 0.5) } @@ -233,7 +240,10 @@ def test_extrapolated_shrinkage_keeps_every_reflection_and_ignores_occupancy(): assert bool((w > 0).all()) and bool((w < 1).all()) lo, hi = torch.minimum(f_dark, f_ext), torch.maximum(f_dark, f_ext) assert bool((f_ext_b >= lo - 1e-4).all()) and bool((f_ext_b <= hi + 1e-4).all()) - assert torch.allclose(var, sig_dark**2 + w * (noise / f) ** 2) + assert torch.allclose( + var, + (1 - w / f) ** 2 * sig_dark**2 + (w / f) ** 2 * (noise**2 - sig_dark**2), + ) assert bool((var > 0).all()) assert torch.equal(out[0.2][2], out[0.5][2]) # The fallback after a failed fit: an infinite SNR gives the unshrunk amplitude, a @@ -248,7 +258,7 @@ def test_extrapolated_shrinkage_keeps_every_reflection_and_ignores_occupancy(): phi, 0.2, snr=torch.full((n,), s_val), - noise=noise, + sig_light=torch.sqrt(noise**2 - sig_dark**2), ) assert bool(torch.isfinite(f_ext_b).all()) and bool(torch.isfinite(var).all()) assert torch.allclose(f_ext_b, expect, atol=1e-4) @@ -282,3 +292,73 @@ def refuse(**_): q = compute_ded_weights("q", **kw, snr_estimate=ValueError("too few")) assert len(record) == 1 and "too few" in str(record[0].message) assert q.applied == "inverse_variance" + + +@pytest.mark.parametrize("occupancy", [0.2, 0.5, 1.0]) +@pytest.mark.parametrize("snr_value", [0.0, 0.5, 2.0, float("inf")]) +def test_extrapolation_propagates_shared_dark_noise(occupancy, snr_value): + """Shrinking a shared-noise difference propagates both independent measurements.""" + from torchref.cli.collection_difference_refine import ( + compute_bayes_extrapolated_amplitudes, + ) + from torchref.config import get_float_dtype + + dark, light, sd, sl, phi, snr = [ + torch.tensor([value], dtype=get_float_dtype()) + for value in (10.0, 12.0, 2.0, 3.0, 0.0, snr_value) + ] + amplitude, variance, w = compute_bayes_extrapolated_amplitudes( + dark, light, sd, phi, phi, occupancy, snr=snr, sig_light=sl + ) + coefficient = w / occupancy + torch.testing.assert_close(amplitude, dark + coefficient * (light - dark)) + torch.testing.assert_close( + variance, (1 - coefficient) ** 2 * sd**2 + coefficient**2 * sl**2 + ) + if occupancy == 1.0 and snr_value == float("inf"): + assert amplitude.item() == 12.0 + assert variance.item() == 9.0 + + +def test_extrapolation_phase_derivatives_on_deposited_amplitudes(mtz_dir): + """Reported variance agrees with derivatives of the phase-aware shrunk amplitude.""" + from torchref import ReflectionData + from torchref.cli.collection_difference_refine import ( + compute_bayes_extrapolated_amplitudes, + ) + + data = ReflectionData(device="cpu", verbose=0).load_mtz(str(mtz_dir / "1DAW.mtz")) + dark, sd = data.get_corrected_data() + dark, sd = dark[:64].detach().clone().requires_grad_(), sd[:64].detach() + light = (dark.detach() * 1.2).requires_grad_() + sl = sd * 1.5 + phi_d = torch.linspace(-1.5, 1.5, len(dark), dtype=dark.dtype) + phi_l = phi_d + 0.8 + snr = torch.linspace(0.1, 3.0, len(dark), dtype=dark.dtype) + amp, var, _ = compute_bayes_extrapolated_amplitudes( + dark, light, sd, phi_d, phi_l, 0.35, snr=snr, sig_light=sl + ) + jd, jl = torch.autograd.grad(amp.sum(), (dark, light)) + torch.testing.assert_close(var, jd**2 * sd**2 + jl**2 * sl**2) + + +def test_intensity_snr_controls_weight_with_amplitude_uncertainty(): + """Intensity noise sets shrinkage while amplitude sigmas describe propagated error.""" + from torchref.cli.collection_difference_refine import ( + compute_bayes_extrapolated_amplitudes, + ) + + kw = _intensity_inputs(n=5000) + est = difference_snr(**kw) + assert est.source == "intensity" + dark = kw["f_dark"] + light = dark + kw["delta_obs"] + sd = torch.ones_like(dark) + sl = 2 * sd + phi = torch.zeros_like(dark) + amp, var, w = compute_bayes_extrapolated_amplitudes( + dark, light, sd, phi, phi, 1.0, snr=est.snr, sig_light=sl + ) + torch.testing.assert_close(w, 1 / (1 + 1 / est.snr)) + torch.testing.assert_close(var, (1 - w) ** 2 * sd**2 + w**2 * sl**2) + assert torch.isfinite(amp).all() diff --git a/torchref/cli/collection_difference_refine.py b/torchref/cli/collection_difference_refine.py index 622b1388..48a0f024 100644 --- a/torchref/cli/collection_difference_refine.py +++ b/torchref/cli/collection_difference_refine.py @@ -323,16 +323,16 @@ def setup_loss_state( def compute_bayes_extrapolated_amplitudes( - Fobs_dark, - Fobs_light, - sig_dark, - phi_dark, - phi_mixed, - f, + Fobs_dark: torch.Tensor, + Fobs_light: torch.Tensor, + sig_dark: torch.Tensor, + phi_dark: torch.Tensor, + phi_mixed: torch.Tensor, + f: float | torch.Tensor, *, - snr, - noise, -): + snr: torch.Tensor, + sig_light: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Shrink the extrapolated amplitudes toward ``Fo_dark`` by their signal fraction. The posterior mean of the extrapolated deviation under a Gaussian prior of the @@ -362,23 +362,44 @@ def compute_bayes_extrapolated_amplitudes( snr : Tensor (N,) Per-reflection signal-to-noise ratio of the difference, from :func:`torchref.maps.ded_weights.difference_snr`. - noise : Tensor (N,) - Calibrated noise of the amplitude difference, from the same call. + sig_light : Tensor (N,) + Sigma of the independent light amplitude measurement, in amplitude units. + Intensity-derived difference noise controls the SNR and shrinkage weight, + not these marginal amplitude uncertainties. Returns ------- tuple ``(F_ext_bayes, var_ext_bayes, w_shrinkage)`` -- the shrunk extrapolated - amplitude, its variance ``sig_dark**2 + w (noise / f)**2`` (the dark - measurement plus the posterior variance of the deviation) and the weight per - reflection. + amplitude, its first-order propagated measurement variance and the weight + per reflection. The variance holds model phases and the fitted weight fixed; + it is not the Gaussian posterior variance of a latent difference and does + not include uncertainty in the fit, phases or occupancy. For equal phases + and positive extrapolated amplitude it is + ``(1 - w/f)**2 * sig_dark**2 + (w/f)**2 * sig_light**2``. + The covariance with the dark component is thereby included. At exactly + zero extrapolated complex amplitude, where the norm has no derivative, + the directional upper bound is used. """ F_dark_phased = Fobs_dark * torch.exp(1j * phi_dark) F_light_phased = Fobs_light * torch.exp(1j * phi_mixed) - F_ext = torch.abs(F_dark_phased + (F_light_phased - F_dark_phased) / f) + z = F_dark_phased + (F_light_phased - F_dark_phased) / f + F_ext = torch.abs(z) # snr / (1 + snr), written so an infinite SNR gives exactly 1 rather than inf/inf. w = 1.0 / (1.0 + 1.0 / snr) - var_ext_bayes = sig_dark**2 + w * (noise / f) ** 2 + unit = z / F_ext.clamp_min(torch.finfo(Fobs_dark.dtype).tiny) + a, b = 1.0 - 1.0 / f, 1.0 / f + d_dark = a * (unit.conj() * torch.exp(1j * phi_dark)).real + d_light = b * (unit.conj() * torch.exp(1j * phi_mixed)).real + j_dark = (1.0 - w) + w * d_dark + j_light = w * d_light + var_ext_bayes = ( + j_dark.square() * sig_dark.square() + j_light.square() * sig_light.square() + ) + zero_bound = ((1.0 - w) + w * abs(a)).square() * sig_dark.square() + ( + w * abs(b) + ).square() * sig_light.square() + var_ext_bayes = torch.where(F_ext > 0, var_ext_bayes, zero_bound) # Shrink the amplitude toward Fo_dark -- scalar, so no phase interference. F_ext_bayes = Fobs_dark + w * (F_ext - Fobs_dark) return F_ext_bayes, var_ext_bayes, w @@ -741,10 +762,9 @@ def _fit(amp, sig): ) / w_light if snr_est is not None: - snr, noise, source = snr_est.snr, snr_est.noise, snr_est.source + snr, source = snr_est.snr, snr_est.source else: snr = torch.full_like(Fobs_dark_vals, float("inf")) - noise = torch.sqrt(sig_dark_vals**2 + sig_light_vals**2) source = "none" F_ext_bayes_amp, var_ext_bayes, w_shrinkage = compute_bayes_extrapolated_amplitudes( Fobs_dark_vals, @@ -754,7 +774,7 @@ def _fit(amp, sig): ctx["phi_mixed"], w_light, snr=snr, - noise=noise, + sig_light=sig_light_vals, ) sig_ext_bayes = torch.sqrt(var_ext_bayes) From 9f82c79ad89022747a40cdb917589ac814d32224 Mon Sep 17 00:00:00 2001 From: Hans Peter Seidel <89108105+HatPdotS@users.noreply.github.com> Date: Wed, 30 Sep 2026 12:21:22 +0200 Subject: [PATCH 236/250] Preserve chemical identity of alternate residue conformers --- docs/changelog.rst | 1 + tests/unit/model/test_strip_altlocs.py | 138 +++++++++++++++++++++++++ torchref/model/model.py | 5 + torchref/topology/atom_graph.py | 6 ++ torchref/topology/build.py | 88 ++++++++++++---- torchref/topology/builders.py | 64 ++++++++---- torchref/topology/hydrogens.py | 32 +++--- torchref/topology/restraints.py | 9 +- torchref/topology/topology.py | 22 ++-- 9 files changed, 300 insertions(+), 65 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 506d884a..f1f48cd9 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Preserve alternate residue types at one sequence position through model loading, selections, occupancy grouping, conformer-specific restraint templates and coordinate writers; stripping altlocs retains the winning conformer's residue name. - A model no longer keeps an atom table. Atom identity lives on ``model.ctx.topology`` (a node-only ``Topology``) and every refinable value only on the parameter wrappers; the table is read once at construction (``ModelContext.from_atoms`` splits it into the topology and ``AtomValues``) and written by ``Model.to_dataframe()``, which joins identity and current values afresh on every call. ``Model.update_pdb`` is removed, ``Model.pdb`` is a deprecated read-only view of ``to_dataframe()`` (writing into it changes nothing), ``model.n_atoms`` replaces ``len(model.pdb)``, and checkpoints keep storing the table under ``"pdb"`` so older ones still restore - ``Model.strip_altlocs`` compares conformers within one residue, ``(chain, resseq, icode)``, so residues 100 and 100A never compete and alternates with different residue names are resolved to one; the kept conformer is chosen by current occupancy rather than the occupancies the file was loaded with - Copying a model's context copies its LINK records as a table; it previously turned them into a list of column names, which broke a later restraint rebuild on the copy diff --git a/tests/unit/model/test_strip_altlocs.py b/tests/unit/model/test_strip_altlocs.py index 6a0bc04d..c87836ba 100644 --- a/tests/unit/model/test_strip_altlocs.py +++ b/tests/unit/model/test_strip_altlocs.py @@ -110,3 +110,141 @@ def test_microheterogeneity_keeps_one_residue(table): kept = kept[(kept.chainid == "A") & (kept.resseq == 30)] assert len(kept) == len(rows) assert (kept.resname != "XAA").all() + + +@pytest.fixture +def microheterogeneous_model(table): + """Deposited THR A30 with a higher-occupancy ALA alternate at the same position.""" + df, cell, sg = table + rows = _residue_rows(df, 30) + other = df.loc[rows].copy() + other = other[other.name.str.strip().isin(["N", "CA", "C", "O", "CB"])].copy() + assert len(other) == 5 + other[["altloc", "occupancy", "resname"]] = ["B", 0.65, "ALA"] + df = df.copy() + df.loc[rows, ["altloc", "occupancy"]] = ["A", 0.35] + df = pd.concat([df.loc[: rows[-1]], other, df.loc[rows[-1] + 1 :]]).reset_index( + drop=True + ) + return _model(df, cell, sg) + + +def test_microheterogeneity_identity_and_current_occupancy(microheterogeneous_model): + """Each alternate retains its chemical name and selection through model derivation.""" + model = microheterogeneous_model + query = "chain A and resseq 30 and resname ALA" + assert model.get_selection_mask(query).sum().item() == 5 + cols = model.ctx.topology.columns() + at_position = (cols["chain"] == "A") & (cols["resseq"] == 30) + assert len(set(model.ctx.topology.atoms.residue_of[at_position].tolist())) == 1 + assert (cols["resname"][at_position & (cols["altloc"] == "B")] == "ALA").all() + groups = model.ctx._residue_groups(with_altloc=True) + assert len(groups[("ALA", 30, "A", "B")]) == 5 + assert len(groups[("THR", 30, "A", "A")]) > 5 + np.testing.assert_allclose( + model.occupancy().detach().numpy()[at_position & (cols["altloc"] == "B")], 0.65 + ) + for derived in ( + model.copy(), + Model.create_from_state_dict(model.state_dict(), device="cpu", verbose=0), + ): + assert derived.get_selection_mask(query).sum().item() == 5 + selected = model.select(query) + assert selected.n_atoms == 5 + assert (selected.to_dataframe().resname == "ALA").all() + stripped = model.strip_altlocs().to_dataframe() + kept = stripped[(stripped.chainid == "A") & (stripped.resseq == 30)] + assert len(kept) == 5 + assert (kept.resname == "ALA").all() + assert (kept.altloc.str.strip() == "").all() + + +def test_microheterogeneity_restraint_templates(microheterogeneous_model): + """The ALA alternate has its own methyl template, independently of THR.""" + model = microheterogeneous_model + connected = model.restraints.topology + cols = connected.columns() + for identity, altloc, energy, h_count in ( + ("THR", "A", "CH1", 1), + ("ALA", "B", "CH3", 3), + ): + cb = np.nonzero( + (cols["chain"] == "A") + & (cols["resseq"] == 30) + & (cols["altloc"] == altloc) + & (cols["name"] == "CB") + )[0] + assert len(cb) == 1 + row = cb[0] + assert connected.resname_of_atom(row) == identity + assert connected.atoms.energy_type[row] == energy + assert connected.atoms.template_h_count[row].item() == h_count + b_rows = set( + np.nonzero( + (cols["chain"] == "A") & (cols["resseq"] == 30) & (cols["altloc"] == "B") + )[0] + ) + bonds = connected.atoms.bonds.indices.cpu().numpy() + own_bonds = { + tuple(sorted(cols["name"][[a, b]])) + for a, b in bonds + if a in b_rows and b in b_rows + } + assert own_bonds == { + tuple(sorted(pair)) + for pair in [("N", "CA"), ("CA", "C"), ("C", "O"), ("CA", "CB")] + } + assert (cols["resname"][list(b_rows)] == "ALA").all() + + +@pytest.mark.parametrize("suffix", ["pdb", "cif"]) +def test_microheterogeneity_writers(microheterogeneous_model, tmp_path, suffix): + """Coordinate writers retain both chemical identities and the winning identity.""" + import gemmi + + model = microheterogeneous_model + for current, stripped in ((model, False), (model.strip_altlocs(), True)): + path = tmp_path / (("stripped" if stripped else "alternates") + "." + suffix) + getattr(current, "write_" + suffix)(str(path)) + structure = gemmi.read_structure(str(path)) + rows = [ + (res.name, atom.altloc) + for chain in structure[0] + if chain.name == "A" + for res in chain + if res.seqid.num == 30 + for atom in res + ] + if stripped: + assert len(rows) == 5 and all(name == "ALA" for name, _ in rows) + else: + assert sum(name == "ALA" and alt == "B" for name, alt in rows) == 5 + assert any(name == "THR" and alt == "A" for name, alt in rows) + + +def test_microheterogeneity_shared_atoms_take_winning_name(microheterogeneous_model): + """Shared backbone atoms become part of the retained chemical residue on stripping.""" + model = microheterogeneous_model + df = model.to_dataframe() + at_position = (df.chainid == "A") & (df.resseq == 30) + backbone = df.name.str.strip().isin(["N", "CA", "C", "O"]) + df.loc[at_position & backbone & (df.altloc == "A"), ["altloc", "occupancy"]] = [ + "", + 1.0, + ] + df = df[~(at_position & backbone & (df.altloc == "B"))].reset_index(drop=True) + shared = _model(df, np.asarray(model.cell.data.cpu()), model.spacegroup.hm) + kept = shared.strip_altlocs().to_dataframe() + kept = kept[(kept.chainid == "A") & (kept.resseq == 30)] + assert len(kept) == 5 + assert (kept.resname == "ALA").all() + connected = shared.restraints.topology + columns = connected.columns() + cb_b = np.nonzero( + (columns["chain"] == "A") + & (columns["resseq"] == 30) + & (columns["altloc"] == "B") + & (columns["name"] == "CB") + )[0][0] + neighbours = connected.atoms.neighbors(int(cb_b)).tolist() + assert [columns["name"][row] for row in neighbours] == ["CA"] diff --git a/torchref/model/model.py b/torchref/model/model.py index 54df2e08..8e9cb2ea 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -1776,6 +1776,7 @@ def strip_altlocs(self) -> "Model": altloc = topology.atoms.altloc occupancy = self.occupancy().detach().cpu().numpy() keep = np.ones(self.n_atoms, dtype=bool) + resnames = topology.columns()["resname"] for residue in range(topology.n_residues): rows = np.arange( int(topology.residues.atom_start[residue]), @@ -1786,11 +1787,15 @@ def strip_altlocs(self) -> "Model": continue means = [occupancy[rows[altloc[rows] == label]].mean() for label in labels] best = labels[int(np.argmax(means))] + # Shared atoms belong to the retained chemical conformer in a model + # with no altlocs, even when their deposited name was the other type. + resnames[rows] = resnames[rows[altloc[rows] == best][0]] keep[rows[(altloc[rows] != " ") & (altloc[rows] != best)]] = False rows = np.nonzero(keep)[0] columns = {key: value[rows] for key, value in topology.columns().items()} columns["altloc"] = np.full(len(rows), " ") + columns["resname"] = resnames[rows] return self._derive_from( Topology.from_columns(columns), self._current_values().gather(rows), diff --git a/torchref/topology/atom_graph.py b/torchref/topology/atom_graph.py index 0c557eff..85dc2e4a 100644 --- a/torchref/topology/atom_graph.py +++ b/torchref/topology/atom_graph.py @@ -115,6 +115,9 @@ class AtomGraph(DeviceMixin): Per-atom identifiers, shape ``(N,)``. Strings, so NumPy rather than tensors; residue-level identity is reached through ``residue_of`` rather than duplicated here. ``altloc`` is ``' '`` for atoms in no alternative conformation. + resname : numpy.ndarray, optional + Chemical residue identity per atom, shape ``(N,)``, preserving identities + of alternate conformers at one sequence position. residue_of : torch.Tensor Residue index per atom, shape ``(N,)``, dtype ``int64``. is_hetatm : numpy.ndarray, optional @@ -160,6 +163,7 @@ class AtomGraph(DeviceMixin): energy_type: Optional[np.ndarray] = None template_h_count: Optional[torch.Tensor] = None hb_type: Optional[torch.Tensor] = None + resname: Optional[np.ndarray] = None _adj_indptr: Optional[torch.Tensor] = field(default=None, repr=False) _adj_indices: Optional[torch.Tensor] = field(default=None, repr=False) @@ -243,6 +247,7 @@ def vdw_radii(self) -> np.ndarray: def copy(self) -> "AtomGraph": """An independent copy sharing no storage with this one.""" return AtomGraph( + resname=None if self.resname is None else self.resname.copy(), name=self.name.copy(), element=self.element.copy(), altloc=self.altloc.copy(), @@ -310,6 +315,7 @@ def subset(self, remap: torch.Tensor, residue_remap: torch.Tensor) -> "AtomGraph keep_t = torch.as_tensor(keep, device=self.residue_of.device) return AtomGraph( + resname=None if self.resname is None else self.resname[keep], name=self.name[keep], element=self.element[keep], altloc=self.altloc[keep], diff --git a/torchref/topology/build.py b/torchref/topology/build.py index de5f5932..335e77df 100644 --- a/torchref/topology/build.py +++ b/torchref/topology/build.py @@ -52,6 +52,42 @@ def _atom_columns(topology: Topology) -> Dict[str, np.ndarray]: cols["index"] = np.arange(topology.n_atoms, dtype=np.int64) return cols + +def _chemical_nodes(cols, nodes, peptide_pairs): + """Expand sequence positions into chemical identities for template matching. + + Blank-altloc atoms participate in every chemical conformer at their position. + ``index`` continues to address the original atom order, including when shared + atoms are duplicated in this temporary matching view. + """ + rows, identities, owners, starts, ends = [], [], [], [], [] + variants = {} + for r, (start, end) in enumerate(zip(nodes["atom_start"], nodes["atom_end"])): + source = np.arange(int(start), int(end)) + names = list(dict.fromkeys(cols["resname"][source].tolist())) + variants[r] = [] + for rn in names: + variants[r].append(len(owners)) + chosen = source[ + (cols["resname"][source] == rn) | (cols["altloc"][source] == " ") + ] + starts.append(len(rows)) + rows.extend(chosen.tolist()) + identities.extend([rn] * len(chosen)) + ends.append(len(rows)) + owners.append(r) + rows = np.asarray(rows, dtype=np.int64) + owners = np.asarray(owners, dtype=np.int64) + expanded = {k: v[rows] for k, v in cols.items()} + expanded["resname"] = np.asarray(identities) + chemical = {k: v[owners] for k, v in nodes.items()} + chemical["resname"] = expanded["resname"][starts] + chemical["atom_start"] = np.asarray(starts, dtype=np.int64) + chemical["atom_end"] = np.asarray(ends, dtype=np.int64) + pairs = [(a, b) for i, j in peptide_pairs for a in variants[i] for b in variants[j]] + return expanded, chemical, pairs, rows, owners + + def _conformers( cols: Dict[str, np.ndarray], start: int, end: int ) -> List[Tuple[np.ndarray, np.ndarray]]: @@ -570,17 +606,15 @@ def _lookup_link_atom( key = (str(chainid), int(resseq), str(icode).strip()) candidates = residue_by_key.get(key, []) wanted = str(resname).strip() if resname else "" - if wanted: - tied = [ - r for r in candidates if str(topology.residues.resname[r]).strip() == wanted - ] - candidates = tied or candidates rows = [ row for r in candidates for row in topology.residues.atom_rows(r) if str(topology.atoms.name[row]).strip() == str(name).strip() ] + if wanted: + tied = [row for row in rows if topology.resname_of_atom(row).strip() == wanted] + rows = tied or rows if not rows: return None altlocs = [str(topology.atoms.altloc[row]) for row in rows] @@ -801,33 +835,47 @@ def build_topology_with_values( (int(polymer_map[a]), int(polymer_map[b])) for a, b in peptide_local ] - comp_dict, template_key = resolve_template_keys( - nodes["resname"], peptide_pairs, cif_dict, link_list, verbose=verbose + match_cols, chemical_nodes, chemical_pairs, source_rows, owners = _chemical_nodes( + cols, nodes, peptide_pairs ) + comp_dict, chemical_keys = resolve_template_keys( + chemical_nodes["resname"], chemical_pairs, cif_dict, link_list, verbose=verbose + ) + template_key = np.asarray(nodes["resname"], dtype=object).copy() + _, first_variant = np.unique(owners, return_index=True) + template_key[:] = chemical_keys[first_variant] pp_cif = PreprocessedCIF(comp_dict) - match_cols = dict(cols) - match_cols["name"] = cols["name"].copy() + match_cols["name"] = match_cols["name"].copy() # PDB terminal H1 is the monomer dictionary's H. Resolve the alias only # for matching, preserving the model's atom names and row identities. - for r in range(n_res): - start, end = int(nodes["atom_start"][r]), int(nodes["atom_end"][r]) + for r in range(len(chemical_keys)): + start, end = int(chemical_nodes["atom_start"][r]), int( + chemical_nodes["atom_end"][r] + ) names = match_cols["name"][start:end] if "H1" not in names or "H" in names: continue - component = comp_dict.get(str(template_key[r]), {}) + component = comp_dict.get(str(chemical_keys[r]), {}) atom_table = component.get("atoms") if atom_table is None: continue template_names = set(atom_table["atom_id"].astype(str).str.strip()) if "H" in template_names and "H1" not in template_names: names[names == "H1"] = "H" - energy_type, template_h_count = _atom_types( - match_cols, nodes, template_key, comp_dict + chemical_energy, chemical_h_count = _atom_types( + match_cols, chemical_nodes, chemical_keys, comp_dict + ) + energy_type = np.full(topology.n_atoms, "", dtype=chemical_energy.dtype) + template_h_count = np.full(topology.n_atoms, -1, dtype=np.int8) + own_identity = match_cols["resname"] == cols["resname"][source_rows] + energy_type[source_rows[own_identity]] = chemical_energy[own_identity] + template_h_count[source_rows[own_identity]] = chemical_h_count[own_identity] + + intra, intra_values = _match_intra( + match_cols, chemical_nodes, chemical_keys, pp_cif ) - - intra, intra_values = _match_intra(match_cols, nodes, template_key, pp_cif) intra_planes, intra_plane_values = _match_intra_planes( - match_cols, nodes, template_key, pp_cif + match_cols, chemical_nodes, chemical_keys, pp_cif ) inter, inter_values, extras = _inter_residue_edges( PeptideResidues(topology, peptide_pairs, xyz.detach().cpu().numpy()), @@ -970,6 +1018,7 @@ def build_topology_with_values( ) atoms = AtomGraph( + resname=cols["resname"].copy(), name=cols["name"], element=cols["element"], altloc=cols["altloc"], @@ -990,8 +1039,9 @@ def build_topology_with_values( planes=plane_blocks, energy_type=energy_type, template_h_count=torch.as_tensor( - # dtype-ok: small per-atom count; int8 is AtomGraph's documented storage - template_h_count, dtype=torch.int8, device=device + template_h_count, + dtype=torch.int8, # dtype-ok: small per-atom count; AtomGraph storage + device=device, ), ) diff --git a/torchref/topology/builders.py b/torchref/topology/builders.py index 2e2a95c0..3909f30d 100644 --- a/torchref/topology/builders.py +++ b/torchref/topology/builders.py @@ -91,10 +91,34 @@ class PeptideResidues: def __init__(self, topology, pairs, xyz): self.pairs = [(int(a), int(b)) for a, b in pairs] self.resnames = np.char.strip(np.asarray(topology.residues.resname).astype(str)) + self.atom_altlocs = topology.atoms.altloc + self.atom_resnames = topology.columns()["resname"] self.xyz = np.asarray(xyz, dtype=np.float64) involved = sorted({r for pair in self.pairs for r in pair}) self.conformer_maps = {r: _conformer_maps(topology, r) for r in involved} + def conformer_resname(self, mapping: Dict[str, int]) -> str: + """Return the chemical identity of a conformer atom-name map. + + Parameters + ---------- + mapping : dict + Atom names mapped to topology rows for one conformer. + + Returns + ------- + str + Residue name of the labelled atoms, or the shared atoms if unlabelled. + """ + rows = list(mapping.values()) + names = self.atom_resnames[rows] + # Shared atoms can retain the first conformer's name. A conformer's + # distinct chemical identity belongs to its labelled atoms. + for row in rows: + if self.atom_altlocs[row] != " ": + return str(self.atom_resnames[row]) + return str(names[0]) + class PreprocessedCIF: """ @@ -705,16 +729,19 @@ def build( n_angles = len(angles["atom1"]) for res_i_idx, res_next_idx in pairs: - # Filter by next residue name if requested - if next_resname_filter is not None: - if residues.resnames[res_next_idx] != next_resname_filter: - continue - if exclude_next_resname is not None: - if residues.resnames[res_next_idx] == exclude_next_resname: - continue - for map_i in conf_maps[res_i_idx]: for map_next in conf_maps[res_next_idx]: + next_name = residues.conformer_resname(map_next) + if ( + next_resname_filter is not None + and next_name != next_resname_filter + ): + continue + if ( + exclude_next_resname is not None + and next_name == exclude_next_resname + ): + continue for a in range(n_angles): comp1, comp2, comp3 = ( @@ -953,12 +980,13 @@ def build( from torchref.topology.ramachandran import classify_residue for res_i_idx, res_next_idx in pairs: - resname_i = residues.resnames[res_i_idx] - resname_next = residues.resnames[res_next_idx] - is_proline = resname_next == "PRO" - for map_i in conf_maps[res_i_idx]: for map_next in conf_maps[res_next_idx]: + resname_i = residues.conformer_resname(map_i) + resname_next = residues.conformer_resname(map_next) + is_proline = resname_next == "PRO" + key_i = (res_i_idx, resname_i) + key_next = (res_next_idx, resname_next) # Track which residue each phi/psi belongs to pair_phi = None # phi from this pair belongs to res_next_idx @@ -1013,19 +1041,19 @@ def build( # phi: C(i) - N(j) - CA(j) - C(j) → belongs to residue j # psi: N(i) - CA(i) - C(i) - N(j) → belongs to residue i if pair_phi is not None: - phi_by_residue[res_next_idx] = pair_phi + phi_by_residue[key_next] = pair_phi if pair_psi is not None: - psi_by_residue[res_i_idx] = pair_psi + psi_by_residue[key_i] = pair_psi # Track residue names and next-residue names for classification - resname_by_residue[res_i_idx] = resname_i - resname_by_residue[res_next_idx] = resname_next - next_resname_by_residue[res_i_idx] = resname_next + resname_by_residue[key_i] = resname_i + resname_by_residue[key_next] = resname_next + next_resname_by_residue[key_i] = resname_next # Compute omega for PRO cis/trans detection if omega_data["indices"]: omega_deg = self._torsion_angle_np( coords_np, *omega_data["indices"][-1] ) - omega_by_residue[res_next_idx] = omega_deg + omega_by_residue[key_next] = omega_deg result = {} diff --git a/torchref/topology/hydrogens.py b/torchref/topology/hydrogens.py index 4325836e..d8a9126d 100644 --- a/torchref/topology/hydrogens.py +++ b/torchref/topology/hydrogens.py @@ -536,27 +536,27 @@ def plan_hydrogens(topology, cif_dict: Dict, xyz, verbose: int = 0) -> HydrogenP n_unplaceable = 0 n_no_template = 0 + atom_resnames = topology.columns()["resname"] for residue in range(residues.n_residues): - resname = str(residues.resname[residue]).strip() - template = _template(cif_dict, resname) - if template is None: - # Atoms without a dictionary cannot supply either bond geometry or - # hydrogen identities; leave those residues unchanged. - n_no_template += 1 - continue - rows = np.arange( int(residues.atom_start[residue]), int(residues.atom_end[residue]) ) - present = set(names[rows]) - h1_alias = "H" in template["h_names"] and "H1" not in template["h_names"] - if h1_alias and "H1" in present: - present.add("H") - candidates = [h for h in template["h_names"] if h not in present] - if not candidates: - continue - for altloc, conformer in _conformer_rows(rows, altlocs): + labelled = conformer[altlocs[conformer] != " "] + identity_row = labelled[0] if len(labelled) else conformer[0] + resname = str(atom_resnames[identity_row]).strip() + template = _template(cif_dict, resname) + if template is None: + n_no_template += 1 + continue + present = set(names[conformer]) + h1_alias = "H" in template["h_names"] and "H1" not in template["h_names"] + if h1_alias and "H1" in present: + present.add("H") + candidates = [h for h in template["h_names"] if h not in present] + if not candidates: + continue + name_to_row = {} for row in conformer: name = "H" if h1_alias and names[row] == "H1" else names[row] diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index e0b35a26..c33a1995 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -155,12 +155,9 @@ def _multi_atom_resnames(topology) -> list: Single-atom residues (ions, lone waters) need no dictionary lookup. """ names_by_resname: dict = {} - resnames = np.char.strip(topology.residues.resname.astype(str)) - for r, resname in enumerate(resnames): - rows = topology.residues.atom_rows(r) - names_by_resname.setdefault(str(resname), set()).update( - topology.atoms.name[rows.start : rows.stop].tolist() - ) + columns = topology.columns() + for resname, atom_name in zip(columns["resname"], columns["name"]): + names_by_resname.setdefault(str(resname), set()).add(str(atom_name)) return [name for name, atoms in names_by_resname.items() if len(atoms) > 1] def _riding_table(self, xyz: torch.Tensor) -> pd.DataFrame: diff --git a/torchref/topology/topology.py b/torchref/topology/topology.py index 5561cdd9..dcdf5718 100644 --- a/torchref/topology/topology.py +++ b/torchref/topology/topology.py @@ -2,8 +2,9 @@ :class:`Topology` is where a model's atom identity and connectivity live. The residue level carries the sequence and the inter-residue links; the atom level carries the -atoms, the typed edge blocks and the bond adjacency. Per-atom residue identity is -reached through ``atoms.residue_of`` rather than duplicated per atom. +atoms, the typed edge blocks and the bond adjacency. Sequence position is reached +through ``atoms.residue_of``; chemical residue identity is per atom so alternate +residue types can share a sequence position. Identity comes first. :meth:`Topology.from_table` is the one place an atom table's identity columns become arrays; the result is a node-only topology -- names, elements, @@ -172,6 +173,7 @@ def from_columns(cls, columns: Mapping[str, np.ndarray], device=None) -> "Topolo device=device, ) atoms = AtomGraph( + resname=np.asarray(columns["resname"]).copy(), name=np.asarray(columns["name"]), element=np.asarray(columns["element"]), altloc=np.asarray(columns["altloc"]), @@ -182,7 +184,7 @@ def from_columns(cls, columns: Mapping[str, np.ndarray], device=None) -> "Topolo return cls(residues=residues, atoms=atoms) def columns(self) -> Dict[str, np.ndarray]: - """Per-atom identity arrays, residue fields broadcast to atoms. + """Per-atom identity arrays, sequence-position fields broadcast to atoms. Returns ------- @@ -197,7 +199,11 @@ def columns(self) -> Dict[str, np.ndarray]: "chain": self.residues.chain[of], "resseq": self.residues.resseq[of], "icode": self.residues.icode[of], - "resname": self.residues.resname[of], + "resname": ( + self.residues.resname[of] + if self.atoms.resname is None + else self.atoms.resname.copy() + ), "is_hetatm": self.atoms.is_hetatm.copy(), "charge": self.atoms.charge.copy(), } @@ -243,7 +249,9 @@ def select(self, selection: str) -> torch.Tensor: @property def is_water(self) -> np.ndarray: """True for atoms of water residues, shape ``(N,)``.""" - return self.residues.is_water[self.atoms.residue_of.cpu().numpy()] + from torchref.topology.residue_graph import WATER_RESNAMES + + return np.isin(self.columns()["resname"], list(WATER_RESNAMES)) @property def is_polymer(self) -> np.ndarray: @@ -424,7 +432,9 @@ def residue_of_atom(self, i: int) -> int: return int(self.atoms.residue_of[i]) def resname_of_atom(self, i: int) -> str: - """Residue name of atom ``i``, joined through the residue graph.""" + """Chemical residue name of atom ``i``, including alternate residue types.""" + if self.atoms.resname is not None: + return str(self.atoms.resname[i]) return str(self.residues.resname[self.residue_of_atom(i)]) def edge_block(self, edge_type: str): From 370758de2c15a91a1c5c49e299fa687f9e7ade71 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 20:47:31 +0000 Subject: [PATCH 237/250] Return expand_hkl's index map in the configured int dtype _equivalent_hkl built its source-row map with torch.arange's int64 default, so expand_hkl returned orig_indices in int64 and equivalent_hkl returned its source rows in int64. Every caller only indexes with them, so build the map in the configured int dtype, as the rotated indices already are. merge_to_spacegroup's pair key merge_id * n_src + source stays int64 because merge_id comes from torch.unique. The docstrings of expand_hkl, equivalent_hkl and reduce_hkl name the configured int dtype instead of int32 or int64. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- torchref/symmetry/reciprocal_symmetry.py | 10 +++++----- torchref/symmetry/spacegroup.py | 14 ++++++++------ 2 files changed, 13 insertions(+), 11 deletions(-) diff --git a/torchref/symmetry/reciprocal_symmetry.py b/torchref/symmetry/reciprocal_symmetry.py index e0bcdb7f..ecc312aa 100644 --- a/torchref/symmetry/reciprocal_symmetry.py +++ b/torchref/symmetry/reciprocal_symmetry.py @@ -55,9 +55,9 @@ def _equivalent_hkl( Returns ------- - copies : torch.Tensor, shape (M, 3), dtype=int32 + copies : torch.Tensor, shape (M, 3), configured int dtype ``M = n_ops * N``, doubled with ``include_friedel``. - source : torch.Tensor, shape (M,), dtype=int64 + source : torch.Tensor, shape (M,), configured int dtype Input row of each copy. phase_shifts : torch.Tensor, shape (M,) Translation phase offset in radians of each copy. @@ -82,7 +82,7 @@ def _equivalent_hkl( copies = copies.reshape(-1, 3) phase = phase.reshape(-1) - source = torch.arange(n, device=device).repeat(sym.n_ops) + source = torch.arange(n, dtype=get_int_dtype(), device=device).repeat(sym.n_ops) is_friedel = torch.zeros(len(copies), dtype=torch.bool, device=device) if include_friedel: copies = torch.cat([copies, -copies]) @@ -122,9 +122,9 @@ def _expand_hkl( Returns ------- - expanded_hkl : torch.Tensor, shape (M, 3), dtype=int32 + expanded_hkl : torch.Tensor, shape (M, 3), configured int dtype All unique expanded Miller indices, in order of first occurrence. - orig_indices : torch.Tensor, shape (M,), dtype=int64 + orig_indices : torch.Tensor, shape (M,), configured int dtype Index mapping expanded → original: ``F_expanded = F_orig[orig_indices]``. phase_shifts : torch.Tensor, shape (M,), dtype=float32 Translation phase offsets in radians: diff --git a/torchref/symmetry/spacegroup.py b/torchref/symmetry/spacegroup.py index 34c48346..e24f7a8f 100644 --- a/torchref/symmetry/spacegroup.py +++ b/torchref/symmetry/spacegroup.py @@ -331,9 +331,10 @@ def expand_hkl( Returns ------- expanded_hkl : torch.Tensor - Expanded indices, shape ``(M, 3)``, dtype ``int32``. + Expanded indices, shape ``(M, 3)``, in the configured int dtype. orig_indices : torch.Tensor - Map expanded -> original, shape ``(M,)``: ``F_exp = F_orig[orig_indices]``. + Map expanded -> original, shape ``(M,)``, in the configured int dtype: + ``F_exp = F_orig[orig_indices]``. phase_shifts : torch.Tensor Translation phase offsets in radians, shape ``(M,)``: ``phase_exp = phase_orig[orig_indices] + phase_shifts``. @@ -379,10 +380,11 @@ def equivalent_hkl( Returns ------- copies : torch.Tensor - Shape ``(M, 3)``, int32, ordered by operation then row, Friedel copies - last; ``M = n_ops * N``, doubled with ``include_friedel``. + Shape ``(M, 3)``, in the configured int dtype, ordered by operation + then row, Friedel copies last; ``M = n_ops * N``, doubled with + ``include_friedel``. source : torch.Tensor - Input row of each copy, shape ``(M,)``. + Input row of each copy, shape ``(M,)``, in the configured int dtype. phase_shifts : torch.Tensor Translation phase offset in radians, shape ``(M,)``. is_friedel : torch.Tensor @@ -416,7 +418,7 @@ def reduce_hkl( Returns ------- hkl_asu : torch.Tensor - Unique ASU indices, shape ``(M, 3)``, dtype ``int32``. + Unique ASU indices, shape ``(M, 3)``, in the configured int dtype. reduction_indices : torch.Tensor Indices into ``hkl_p1`` per equivalent, shape ``(M, n_equiv)``, **-1 where no P1 reflection exists** -- mask or clamp before gathering, or a -1 From dfc675ebf54964aed1a3292fe16c6e0553992a2a Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 08:08:47 +0000 Subject: [PATCH 238/250] Refuse --free-fraction and --max-free with an inherited free set An inherited set is extended at its own fraction. Applying another fraction to the reflections the reference lacks would leave a mixed partition that the reported fraction describes for neither part. uniform_rfree now raises ValueError when free_fraction or max_free is given together with reference, and the CLI refuses either option while a set is inherited, pointing to --fresh. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- docs/changelog.rst | 2 +- docs/user_guide/cli.rst | 15 ++++----- tests/integration/test_cli_uniform_rfree.py | 12 +++++++ tests/unit/io/test_uniform_rfree.py | 7 +++++ torchref/cli/uniform_rfree.py | 35 ++++++++++++++++----- torchref/io/rfree.py | 26 ++++++++++++--- 6 files changed, 77 insertions(+), 20 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index a10aeaff..4734e1b5 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -7,7 +7,7 @@ Unreleased - SF-CIF export preserves supported numerical columns and free-flag values, and rejects unsupported columns before writing; use MTZ to retain custom columns or saved original flags. - Added ``torchref.uniform-rfree``: gives any number of MTZ / SF-mmCIF files of one cell and space group a shared CCP4 ``FreeR_flag``. - An existing free set (CCP4, Phenix or mmCIF convention) is inherited by default and extended, at its own fraction, to reflections it lacks, including those beyond its resolution. That extension is seeded with a hash of the reference's free/work partition, so it is reproducible whatever the file format or convention. New sets are stratified by resolution shell on the complete ASU and depend only on cell, space group and seed. - - Options: ``--check`` reports whether the inputs' free sets agree; ``--max-free`` caps the size of a new set; ``--scale`` optionally scales the datasets together with ``DatasetCollection.scale``. + - Options: ``--check`` reports whether the inputs' free sets agree; ``--free-fraction`` and ``--max-free`` size a new set and are refused while a set is inherited; ``--scale`` optionally scales the datasets together with ``DatasetCollection.scale``. - Excluded reflections (``-1`` / ``x``) are preserved per file. Output is MTZ (all input columns kept) and/or mmCIF (columns with an mmCIF equivalent). Library helpers are in ``torchref.io.rfree``. - The MTZ reader keeps R-free flags as integers, so ``-1`` (excluded) reflections are masked on load instead of becoming work reflections. - ``ModelFT.create_from_state_dict`` restores the restraint dictionary path (``cif_path``) as ``Model`` does, and ``Refinement.create_from_state_dict`` no longer builds a stray ``Restraints`` from the model; the restored model builds its own on first access. diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index 6c2fe9ee..4224d842 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -111,11 +111,12 @@ reflection is free in one dataset and work in another. - **Existing flags are kept by default.** If any input already has an R-free column (CCP4 ``0 = free``, Phenix ``1 = free`` or mmCIF ``status``), its - free set is inherited and extended to the reflections it lacks. The source - is the first input with flags, or the file named with ``--reference``. - Without any flags a new set is generated. ``--fresh`` always generates one. - Replacing a set that a model was already refined against makes that model's - R-free meaningless. + free set is inherited and extended, at its own fraction, to the reflections + it lacks. The source is the first input with flags, or the file named with + ``--reference``. Without any flags a new set is generated. ``--fresh`` always + generates one. ``--free-fraction`` and ``--max-free`` size a new set, so they + are refused while a set is inherited. Replacing a set that a model was + already refined against makes that model's R-free meaningless. - **Mixed resolution cutoffs.** If the reference stops short of the data resolution, a warning is printed and the higher-resolution shells, plus any gaps in the reference, are generated at the reference's free fraction, @@ -164,8 +165,8 @@ alone. By default everything goes onto the shared consensus scale; ``--scale-reference`` leaves one input unchanged instead. **Key options:** ``--check``, ``--reference {auto,FILE}``/``--fresh``, -``--reference-column``, ``--free-fraction`` (default: the reference's, else -0.05), ``--max-free`` (because the cap depends on resolution, the run prints +``--reference-column``, ``--free-fraction`` (new sets, default 0.05), +``--max-free`` (new sets; because the cap depends on resolution, the run prints the ``--free-fraction`` that reproduces it), ``--seed`` (default: 0 for a new set, the reference hash when extending), ``--shell-size``, ``--format {mtz,cif}``, ``--suffix``, ``--keep-old-flags``, diff --git a/tests/integration/test_cli_uniform_rfree.py b/tests/integration/test_cli_uniform_rfree.py index ad957907..e6f7fac2 100644 --- a/tests/integration/test_cli_uniform_rfree.py +++ b/tests/integration/test_cli_uniform_rfree.py @@ -190,6 +190,18 @@ def test_mixed_resolution_and_max_free(inputs, tmp_path): assert abs(free[d < 2.3].mean() - free[d >= 2.3].mean()) < 0.02 +@pytest.mark.parametrize("size", [["--free-fraction", "0.1"], ["--max-free", "500"]]) +def test_new_set_size_needs_fresh_while_inheriting(inputs, tmp_path, size): + """A new-set size is refused while a set is inherited, and works with --fresh.""" + _, paths = inputs + res = _run(*paths[:2], "-o", tmp_path / "inh", *size) + assert res.returncode == 1 + assert size[0] in res.stderr and "--fresh" in res.stderr + assert not (tmp_path / "inh").exists() + res = _run(*paths[:2], "-o", tmp_path / "new", *size, "--fresh") + assert res.returncode == 0, res.stderr + + def test_excluded_flags_survive(inputs, tmp_path): from torchref.io.datasets.reflection_data import ReflectionData diff --git a/tests/unit/io/test_uniform_rfree.py b/tests/unit/io/test_uniform_rfree.py index e552fcf4..959073b3 100644 --- a/tests/unit/io/test_uniform_rfree.py +++ b/tests/unit/io/test_uniform_rfree.py @@ -281,6 +281,13 @@ def test_max_free_must_be_positive(small): rfree.uniform_rfree({"x": small}, max_free=0) +@pytest.mark.parametrize("size", [{"free_fraction": 0.1}, {"max_free": 500}]) +def test_new_set_size_is_refused_with_a_reference(small, size): + """An inherited set keeps its own fraction, so no new-set size applies.""" + with pytest.raises(ValueError, match="reference=None"): + rfree.uniform_rfree({"x": small}, reference=small, **size) + + @pytest.mark.parametrize("suffix", [".mtz", ".cif"]) def test_supported_columns_and_numeric_flags_roundtrip(mtz_dir, tmp_path, suffix): """Amplitude, intensity, sigma and numeric flag values survive SF export.""" diff --git a/torchref/cli/uniform_rfree.py b/torchref/cli/uniform_rfree.py index bdb23ebe..fb039340 100644 --- a/torchref/cli/uniform_rfree.py +++ b/torchref/cli/uniform_rfree.py @@ -10,9 +10,10 @@ (:meth:`torchref.io.datasets.collection.DatasetCollection.scale`). If any input already carries an R-free column, its free set is inherited and -extended to the reflections it lacks (``--reference auto``, the default); -otherwise a new set is generated. ``--check`` only reports whether the -inputs' existing free sets agree. +extended, at its own fraction, to the reflections it lacks (``--reference +auto``, the default); otherwise a new set is generated. ``--free-fraction`` and +``--max-free`` size a new set, so they need ``--fresh`` when an input has +flags. ``--check`` only reports whether the inputs' existing free sets agree. Usage:: @@ -105,8 +106,9 @@ def _parse_args(argv=None): type=float, default=None, help=( - "Free-set fraction; FreeR_flag takes round(1/f) values (default: the " - "reference's own fraction when inheriting, else 0.05)" + "Fraction of a new free set; FreeR_flag takes round(1/f) values " + "(default: 0.05). An inherited set keeps its own fraction, so this " + "needs --fresh when an input has flags" ), ) flg.add_argument( @@ -114,8 +116,9 @@ def _parse_args(argv=None): type=int, default=None, help=( - "Cap the free set at this many reflections of the complete set to the " - "best resolution (e.g. 2000, as in Phenix); lowers the fraction" + "Cap a new free set at this many reflections of the complete set to " + "the best resolution (e.g. 2000, as in Phenix); lowers the fraction. " + "Needs --fresh when an input has flags" ), ) flg.add_argument( @@ -520,6 +523,24 @@ def main(argv=None): print("Error: " + p, file=sys.stderr) return 1 + sizing = [ + option + for option, value in [ + ("--free-fraction", args.free_fraction), + ("--max-free", args.max_free), + ] + if value is not None + ] + if reference is not None and sizing: + options = " and ".join(sizing) + print( + f"Error: cannot combine {options} with the free set inherited from " + f"{ref_label!r}, which is extended at its own fraction. Drop {options} " + "to extend it, or add --fresh to generate a new set instead.", + file=sys.stderr, + ) + return 1 + # ---- flags ------------------------------------------------------------ try: flags, info = rfree.uniform_rfree( diff --git a/torchref/io/rfree.py b/torchref/io/rfree.py index 9d6c07ee..27c07c45 100644 --- a/torchref/io/rfree.py +++ b/torchref/io/rfree.py @@ -778,8 +778,9 @@ def uniform_rfree( datasets : dict Name to DataSet (same cell / space group). free_fraction : float, optional - Target free fraction; the number of flag values is ``round(1/f)``. - Defaults to the reference's own fraction when inheriting, else 0.05. + Free fraction of a new set; the number of flag values is ``round(1/f)``. + Defaults to 0.05. Not allowed with ``reference``: an inherited set is + extended at its own fraction. shell_size : int Reflections per stratification shell (see :func:`complete_flag_table`). seed : int, optional @@ -796,9 +797,10 @@ def uniform_rfree( reference_column : str, optional Flag column in ``reference`` (auto-detected by default). max_free : int, optional - Cap on the number of free reflections in the complete set to ``dmin`` - (Phenix-style); lowers the fraction for large datasets. The cap depends - on ``dmin``, so reuse the reported fraction to reproduce a flag set. + Cap on the number of free reflections of a new set, counted in the + complete set to ``dmin`` (Phenix-style); lowers the fraction for large + datasets. The cap depends on ``dmin``, so reuse the reported fraction to + reproduce a flag set. Not allowed with ``reference``. keep_excluded : bool Rows a dataset itself marks as excluded (negative flag, CIF ``x``) stay ``-1`` in that dataset's output. @@ -813,11 +815,25 @@ def uniform_rfree( ``n_unique``, ``n_off_asu``, ``n_inherited``, ``n_generated``, ``n_generated_beyond_reference``, ``n_excluded`` (per file) and the reference ``info``. + + Raises + ------ + ValueError + If ``free_fraction`` is outside (0, 1), ``max_free`` is below 1, either + is given together with ``reference``, or ``reference`` has no usable + R-free column (none recognised, no valid value, or no free reflection). """ if free_fraction is not None and not 0 < free_fraction < 1: raise ValueError("free_fraction must be between 0 and 1") if max_free is not None and max_free < 1: raise ValueError("max_free must be at least 1") + # Extending at another fraction than the inherited set's would leave a mixed + # partition that no single reported fraction describes. + if reference is not None and (free_fraction is not None or max_free is not None): + raise ValueError( + "free_fraction and max_free size a new free set; an inherited set is " + "extended at its own fraction, so pass reference=None to replace it" + ) # the flag table lives on the reference's cell when there is one, so the # result does not depend on input order From 7b028845c39ea20d7790355be3914f3a22f837b1 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 08:20:19 +0000 Subject: [PATCH 239/250] Carry the ensemble state through EnsembleModel.copy Model.copy builds a fresh instance of the same class, then copies the context, buffers and parameter wrappers. For an EnsembleModel that left n_members and n_atoms_per_member at 0, dropped the single-copy table, the dropout and population settings, and the per-member occ_logits and b_raw parameters. It also dropped a low-rank or PCA xyz, which has no copy of its own. xyz_per_member, member_weights and write_pdb then failed on the copy. EnsembleModel.copy now carries all of that state. enable_low_rank and enable_pca read self.verbose, which the model does not have, so both raised AttributeError. They now read ctx.verbose. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- docs/changelog.rst | 1 + tests/unit/model/test_ensemble_model.py | 30 +++++++++++ .../experimental/ensemble/ensemble_model.py | 53 +++++++++++++++++-- 3 files changed, 80 insertions(+), 4 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 7a5bdcab..3c234fc3 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- ``EnsembleModel.copy()`` (experimental) returns a complete ensemble: member layout, single-copy atom table, per-member ``occ_logits`` / ``b_raw``, dropout and population settings, and a low-rank or PCA ``xyz``. ``enable_low_rank`` and ``enable_pca`` no longer fail on a missing ``verbose`` attribute. - Preserve alternate residue types at one sequence position through model loading, selections, occupancy grouping, conformer-specific restraint templates and coordinate writers; stripping altlocs retains the winning conformer's residue name. - A model no longer keeps an atom table. Atom identity lives on ``model.ctx.topology`` (a node-only ``Topology``) and every refinable value only on the parameter wrappers; the table is read once at construction (``ModelContext.from_atoms`` splits it into the topology and ``AtomValues``) and written by ``Model.to_dataframe()``, which joins identity and current values afresh on every call. ``Model.update_pdb`` is removed, ``Model.pdb`` is a deprecated read-only view of ``to_dataframe()`` (writing into it changes nothing), ``model.n_atoms`` replaces ``len(model.pdb)``, and checkpoints keep storing the table under ``"pdb"`` so older ones still restore - ``Model.strip_altlocs`` compares conformers within one residue, ``(chain, resseq, icode)``, so residues 100 and 100A never compete and alternates with different residue names are resolved to one; the kept conformer is chosen by current occupancy rather than the occupancies the file was loaded with diff --git a/tests/unit/model/test_ensemble_model.py b/tests/unit/model/test_ensemble_model.py index a5e7585a..52ef4ad5 100644 --- a/tests/unit/model/test_ensemble_model.py +++ b/tests/unit/model/test_ensemble_model.py @@ -166,3 +166,33 @@ def test_dropout_disable_restores_full_occupancy(small_ensemble): assert torch.allclose( ens._dropout_occ_mult, torch.ones_like(ens._dropout_occ_mult) ) + + +def test_copy_carries_the_ensemble(tmp_path, small_ensemble): + """A copy keeps the member layout, per-member levers and single-copy table.""" + ens = small_ensemble + ens.enable_population_refinement(True) + with torch.no_grad(): + ens.occ_logits[0] = 1.0 + dup = ens.copy() + assert type(dup) is EnsembleModel + assert dup.n_members == ens.n_members + assert dup.n_atoms_per_member == ens.n_atoms_per_member + assert torch.equal(dup.xyz_per_member, ens.xyz_per_member) + assert torch.equal(dup.member_weights(), ens.member_weights()) + assert dup.occ_logits.requires_grad and not dup.b_raw.requires_grad + with torch.no_grad(): + dup.occ_logits[1] = 2.0 + assert ens.occ_logits[1] == 0 + assert len(dup.pdb_single) == ens.n_atoms_per_member + dup.write_pdb(str(tmp_path / "copy.pdb")) + + +def test_copy_keeps_a_low_rank_xyz(small_ensemble): + ens = small_ensemble + ens.enable_low_rank(2) + dup = ens.copy() + assert torch.allclose(dup.xyz(), ens.xyz()) + with torch.no_grad(): + dup.xyz.amplitudes.add_(1.0) + assert not torch.allclose(dup.xyz(), ens.xyz()) diff --git a/torchref/experimental/ensemble/ensemble_model.py b/torchref/experimental/ensemble/ensemble_model.py index 43521193..b59bbf97 100644 --- a/torchref/experimental/ensemble/ensemble_model.py +++ b/torchref/experimental/ensemble/ensemble_model.py @@ -630,9 +630,54 @@ def _finalize_ensemble(self, n_members: int, n_atoms_per_member: int) -> None: try: self.freeze(tgt) except Exception: - if self.verbose > 0: + if self.ctx.verbose > 0: print(f" EnsembleModel: freeze({tgt!r}) failed (ignored)") + def copy(self) -> "EnsembleModel": + """Create a deep copy of the ensemble, of the same class. + + :meth:`Model.copy` carries the context, buffers and parameter wrappers. + This adds the state an ensemble holds outside them: the member layout, + the single-copy atom table, the dropout and population-refinement + settings, the per-member ``occ_logits`` / ``b_raw`` (``requires_grad`` + kept), and a low-rank or PCA ``xyz``, which has no ``copy`` of its own. + + Returns + ------- + EnsembleModel + A new, fully independent ensemble. + """ + import copy as copy_module + + duplicate = super().copy() + for name in ( + "n_members", + "n_atoms_per_member", + "dropout_active", + "dropout_min", + "dropout_max", + "_refine_population", + "_refine_member_b", + ): + if hasattr(self, name): + setattr(duplicate, name, getattr(self, name)) + if self._pdb_single is not None: + duplicate._pdb_single = self._pdb_single.copy(deep=True) + for name, param in self._parameters.items(): + if param is not None: + setattr( + duplicate, + name, + torch.nn.Parameter( + param.detach().clone(), requires_grad=param.requires_grad + ), + ) + if not hasattr(self.xyz, "copy"): + duplicate.xyz = copy_module.deepcopy(self.xyz) + duplicate._repoint_coordinate_accessors() + duplicate.reset_cache() + return duplicate + # ------------------------------------------------------------------ # Per-member occupancy + ADP + birth/death population dynamics # ------------------------------------------------------------------ @@ -760,7 +805,7 @@ def enable_low_rank(self, K: int) -> float: K = int(K) max_rank = max(1, N - 1) if K > max_rank: - if self.verbose > 0: + if self.ctx.verbose > 0: print( f" EnsembleModel.enable_low_rank: K={K} exceeds rank " f"N-1={max_rank}; clamping to {max_rank}." @@ -793,7 +838,7 @@ def enable_low_rank(self, K: int) -> float: self.xyz = lowrank self.reset_cache() - if self.verbose > 0: + if self.ctx.verbose > 0: print( f" EnsembleModel.enable_low_rank: K={K} modes, " f"DOF {N * n_atoms * 3} -> {N * K} " @@ -823,7 +868,7 @@ def enable_pca(self, K: Optional[int] = None) -> float: ) self.xyz = pca.to(self.device) self.reset_cache() - if self.verbose > 0: + if self.ctx.verbose > 0: print( f" EnsembleModel.enable_pca: K={self.xyz.K} modes (refine μ,A,V), " f"explained variance = {self.xyz.explained_variance * 100:.2f}%" From 8307e19146669fa15cda4f020208edc46d676e47 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 08:55:26 +0000 Subject: [PATCH 240/250] Build the low-rank copy test's ensemble on CPU enable_low_rank seeds its basis with a float64 SVD, and MPS has no float64, so the test failed on the MPS runner before it reached the copy. It now builds its ensemble on CPU, as it tests copy, not enable_low_rank's device support. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- tests/unit/model/test_ensemble_model.py | 13 +++++++++++-- 1 file changed, 11 insertions(+), 2 deletions(-) diff --git a/tests/unit/model/test_ensemble_model.py b/tests/unit/model/test_ensemble_model.py index 52ef4ad5..7eab11f7 100644 --- a/tests/unit/model/test_ensemble_model.py +++ b/tests/unit/model/test_ensemble_model.py @@ -188,8 +188,17 @@ def test_copy_carries_the_ensemble(tmp_path, small_ensemble): dup.write_pdb(str(tmp_path / "copy.pdb")) -def test_copy_keeps_a_low_rank_xyz(small_ensemble): - ens = small_ensemble +def test_copy_keeps_a_low_rank_xyz(): + # enable_low_rank seeds its basis with a float64 SVD, which MPS cannot run + ens = EnsembleModel.from_single( + TEST_PDB, + n_members=5, + perturb_sigma=0.2, + b_const=5.0, + seed=42, + verbose=0, + device="cpu", + ) ens.enable_low_rank(2) dup = ens.copy() assert torch.allclose(dup.xyz(), ens.xyz()) From 9d4a15e8fd64d25e26b62ef92ff5dbbc7326501e Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 14:45:48 +0000 Subject: [PATCH 241/250] Write the cell metric once and derive the reciprocal basis from it reciprocal_basis_matrix built its own copy of the cell matrix and squared the cosines in the triple-product term of the volume factor (2 cos^2a cos^2b cos^2g instead of 2 cos a cos b cos g). The term vanishes when any cell angle is 90 degrees, so only triclinic and rhombohedral R-setting cells were affected: on 5BOV d-spacings were off by up to 1.3 % and direct-summation |F| by up to 5 %. The same formula was written out five times (torch and NumPy direct basis, torch and NumPy reciprocal basis, Cell volume). get_fractional_matrix is now the only copy: the reciprocal basis is its inverse (Cell.reciprocal_basis_matrix is Cell.inv_fractional_matrix), the volume its determinant, and the matrix is built with torch.stack so it is differentiable in the cell. The NumPy copies (transforms_numpy, reciprocal_basis_matrix_numpy, get_scattering_vectors_numpy, get_s, get_real_grid_numpy, get_grids) had no callers outside one test and are removed. test_cell_geometry pins the direct basis, reciprocal basis, volume and d-spacings to gemmi on triclinic, rhombohedral, monoclinic, hexagonal and orthorhombic cells; on the previous code its triclinic and rhombohedral cases fail. The structure-factor test scene selected reflections with s <= 1/d_min although ten of its reflections lie exactly on that sphere, so the last bit of a* decided membership and the stride that follows shifted the whole subsample; it now takes reflections strictly inside, and the calibration table measured on that scene is re-measured. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- docs/changelog.rst | 2 + tests/unit/base/test_cell_geometry.py | 78 +++++++++ tests/unit/scaling/test_scaler.py | 37 ++--- tests/unit/structure_factor/helpers.py | 24 +-- torchref/base/__init__.py | 20 +-- torchref/base/coordinates/__init__.py | 17 +- torchref/base/coordinates/transforms_numpy.py | 139 ---------------- torchref/base/coordinates/transforms_torch.py | 35 ++-- torchref/base/fourier/__init__.py | 4 - torchref/base/fourier/grid.py | 82 ---------- torchref/base/reciprocal/__init__.py | 6 - torchref/base/reciprocal/basis.py | 150 ++---------------- torchref/symmetry/cell.py | 35 +--- 13 files changed, 144 insertions(+), 485 deletions(-) create mode 100644 tests/unit/base/test_cell_geometry.py delete mode 100644 torchref/base/coordinates/transforms_numpy.py diff --git a/docs/changelog.rst b/docs/changelog.rst index eac7d3a4..0b495d59 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,8 @@ Changelog Unreleased ---------- +- Fixed the reciprocal basis of cells with three non-90° angles (triclinic, rhombohedral R setting): it squared the cosines in the cell-volume term, so d-spacings, resolution cuts and bins, scaling and direct-summation structure factors were off (d by up to 1.3 % on 5BOV, |F| by up to 5 %); cells with a 90° angle were exact. The cell metric is now written once, in ``get_fractional_matrix``, which is differentiable in the cell: the reciprocal basis is its inverse (``Cell.reciprocal_basis_matrix`` is ``Cell.inv_fractional_matrix``) and the cell volume its determinant +- Removed the unused NumPy copies of the cell geometry: the ``torchref.base.coordinates.transforms_numpy`` module (``cartesian_to_fractional``, ``fractional_to_cartesian``, ``get_fractional_matrix_numpy``, ``get_inv_fractional_matrix``, ``convert_coords_to_fractional``), ``reciprocal_basis_matrix_numpy``, ``get_scattering_vectors_numpy``, ``get_s``, ``get_real_grid_numpy`` and ``get_grids``; use the torch functions - ``EnsembleModel.copy()`` (experimental) returns a complete ensemble: member layout, single-copy atom table, per-member ``occ_logits`` / ``b_raw``, dropout and population settings, and a low-rank or PCA ``xyz``. ``enable_low_rank`` and ``enable_pca`` no longer fail on a missing ``verbose`` attribute. - Preserve alternate residue types at one sequence position through model loading, selections, occupancy grouping, conformer-specific restraint templates and coordinate writers; stripping altlocs retains the winning conformer's residue name. - A model no longer keeps an atom table. Atom identity lives on ``model.ctx.topology`` (a node-only ``Topology``) and every refinable value only on the parameter wrappers; the table is read once at construction (``ModelContext.from_atoms`` splits it into the topology and ``AtomValues``) and written by ``Model.to_dataframe()``, which joins identity and current values afresh on every call. ``Model.update_pdb`` is removed, ``Model.pdb`` is a deprecated read-only view of ``to_dataframe()`` (writing into it changes nothing), ``model.n_atoms`` replaces ``len(model.pdb)``, and checkpoints keep storing the table under ``"pdb"`` so older ones still restore diff --git a/tests/unit/base/test_cell_geometry.py b/tests/unit/base/test_cell_geometry.py new file mode 100644 index 00000000..fd6a8ff8 --- /dev/null +++ b/tests/unit/base/test_cell_geometry.py @@ -0,0 +1,78 @@ +"""Cell metric pinned to gemmi: direct basis, reciprocal basis, volume, d-spacings. + +Every cell quantity derives from ``get_fractional_matrix``. The triple-cosine term of +the metric vanishes whenever one cell angle is 90 degrees, so the triclinic and +rhombohedral (R-setting) cells here are the ones that exercise it. +""" + +import gemmi +import numpy as np +import pytest +import torch + +from torchref.base.coordinates.transforms_torch import get_fractional_matrix +from torchref.base.reciprocal.basis import reciprocal_basis_matrix +from torchref.base.reciprocal.hkl import get_d_spacing +from torchref.config import get_float_dtype +from torchref.symmetry import Cell + +pytestmark = pytest.mark.unit + +CELLS = { + "triclinic_5BOV": (44.199, 81.866, 89.925, 100.997, 106.903, 100.84), + "triclinic_skewed": (30.0, 40.0, 50.0, 70.0, 80.0, 60.0), + "rhombohedral_R": (60.0, 60.0, 60.0, 80.0, 80.0, 80.0), + "rhombohedral_acute": (40.0, 40.0, 40.0, 60.0, 60.0, 60.0), + "monoclinic": (50.0, 60.0, 70.0, 90.0, 105.0, 90.0), + "hexagonal": (80.0, 80.0, 120.0, 90.0, 90.0, 120.0), + "orthorhombic": (40.0, 50.0, 60.0, 90.0, 90.0, 90.0), +} + +_RANGE = np.arange(-6, 7) +HKL = np.array( + [(h, k, l) for h in _RANGE for k in _RANGE for l in _RANGE if (h, k, l) != (0, 0, 0)] +) + + +def _gemmi_matrices(params): + cell = gemmi.UnitCell(*params) + return np.array(cell.orth.mat.tolist()), np.array(cell.frac.mat.tolist()) + + +@pytest.mark.parametrize("params", CELLS.values(), ids=CELLS.keys()) +def test_d_spacing_matches_gemmi(params): + cell = torch.tensor(params, dtype=get_float_dtype()) + d = get_d_spacing(torch.from_numpy(HKL), cell).double().numpy() + uc = gemmi.UnitCell(*params) + d_ref = np.array([uc.calculate_d(list(map(int, hkl))) for hkl in HKL]) + np.testing.assert_allclose(d, d_ref, rtol=2e-6) + + +@pytest.mark.parametrize("params", CELLS.values(), ids=CELLS.keys()) +def test_direct_and_reciprocal_bases_match_gemmi(params): + orth_ref, frac_ref = _gemmi_matrices(params) + cell = torch.tensor(params, dtype=torch.float64) # dtype-ok: reference precision + np.testing.assert_allclose(get_fractional_matrix(cell).numpy(), orth_ref, atol=1e-10) + np.testing.assert_allclose(reciprocal_basis_matrix(cell).numpy(), frac_ref, atol=1e-12) + + +@pytest.mark.parametrize("params", CELLS.values(), ids=CELLS.keys()) +def test_cell_object_is_consistent_with_gemmi(params): + _, frac_ref = _gemmi_matrices(params) + cell = Cell(list(params), dtype=torch.float64) # dtype-ok: reference precision + np.testing.assert_allclose( + cell.reciprocal_basis_matrix.numpy(), frac_ref, atol=1e-12 + ) + np.testing.assert_allclose( + (cell.reciprocal_basis_matrix @ cell.fractional_matrix).numpy(), + np.eye(3), + atol=1e-12, + ) + assert float(cell.volume) == pytest.approx(gemmi.UnitCell(*params).volume, rel=1e-12) + + +def test_fractional_matrix_is_differentiable_in_the_cell(): + cell = torch.tensor( # dtype-ok: gradcheck needs double precision + CELLS["triclinic_skewed"], dtype=torch.float64, requires_grad=True + ) + assert torch.autograd.gradcheck(get_fractional_matrix, (cell,)) diff --git a/tests/unit/scaling/test_scaler.py b/tests/unit/scaling/test_scaler.py index e9636acf..512c5087 100644 --- a/tests/unit/scaling/test_scaler.py +++ b/tests/unit/scaling/test_scaler.py @@ -90,17 +90,14 @@ class TestScalingCalculations: @pytest.mark.unit def test_resolution_binning_logic(self, mock_hkl_indices, mock_cell): """Test resolution binning creates correct number of bins.""" - from torchref.base.reciprocal import get_s + from torchref.base.reciprocal import get_scattering_vectors + + hkl = mock_hkl_indices(n_reflections=1000) + s = get_scattering_vectors(hkl, mock_cell).norm(dim=1) - hkl = mock_hkl_indices(n_reflections=1000).numpy() - cell = mock_cell.numpy() - - # Calculate s values - s = get_s(hkl, cell) - # Create bins nbins = 10 - s_sorted = torch.tensor(sorted(s)) + s_sorted = torch.sort(s).values bin_edges = torch.linspace(s_sorted[0], s_sorted[-1], nbins + 1) assert len(bin_edges) == nbins + 1 @@ -137,11 +134,10 @@ class TestBFactorScaling: @pytest.mark.unit def test_b_factor_debye_waller(self, mock_hkl_indices, mock_cell): """Test Debye-Waller factor calculation.""" - from torchref.base.reciprocal import get_s + from torchref.base.reciprocal import get_scattering_vectors - hkl = mock_hkl_indices(n_reflections=100).numpy() - cell = mock_cell.numpy() - s = torch.tensor(get_s(hkl, cell)) + hkl = mock_hkl_indices(n_reflections=100) + s = get_scattering_vectors(hkl, mock_cell).norm(dim=1) B_factor = 20.0 # Ų @@ -155,20 +151,15 @@ def test_b_factor_debye_waller(self, mock_hkl_indices, mock_cell): @pytest.mark.unit def test_b_factor_high_resolution_attenuation(self, mock_cell): """Higher resolution (larger s) should have more attenuation.""" - from torchref.base.reciprocal import get_s + from torchref.base.reciprocal import get_scattering_vectors - cell = mock_cell.numpy() - # Low and high resolution reflections - hkl_low = torch.tensor([[1, 0, 0]], dtype=torch.float64).numpy() - hkl_high = torch.tensor([[10, 10, 10]], dtype=torch.float64).numpy() - - s_low = get_s(hkl_low, cell)[0] - s_high = get_s(hkl_high, cell)[0] - + hkl = torch.tensor([[1, 0, 0], [10, 10, 10]]) + s_low, s_high = get_scattering_vectors(hkl, mock_cell).norm(dim=1) + B_factor = 20.0 - dw_low = torch.exp(torch.tensor(-B_factor * (s_low ** 2) / 4)) - dw_high = torch.exp(torch.tensor(-B_factor * (s_high ** 2) / 4)) + dw_low = torch.exp(-B_factor * s_low**2 / 4) + dw_high = torch.exp(-B_factor * s_high**2 / 4) # High resolution should be more attenuated assert dw_high < dw_low diff --git a/tests/unit/structure_factor/helpers.py b/tests/unit/structure_factor/helpers.py index fcebf894..2dc16d2a 100644 --- a/tests/unit/structure_factor/helpers.py +++ b/tests/unit/structure_factor/helpers.py @@ -70,26 +70,26 @@ # oracle. Amplitudes are rel L2 on complex F; derivatives are of ``ls_target``: # # fineness spacing gridsize amplitude g_xyz g_xyz cos HVP HVP cos -# 0.667 d_min/2 (30, 36, 27) 1.04e-01 1.25e+00 0.4207 2.77e-01 0.9622 -# 1.000 d_min/3 (45, 50, 40) 4.11e-03 4.30e-02 0.999077 2.06e-02 0.999814 -# 1.300 d_min/3.9 (60, 64, 54) 8.01e-04 5.86e-03 0.999983 1.16e-03 0.999999 -# 1.600 d_min/4.8 (72, 80, 64) 8.03e-04 5.87e-03 0.999983 1.16e-03 0.999999 -# 2.200 d_min/6.6 (100,108, 90) 7.98e-04 6.00e-03 0.999982 1.16e-03 0.999999 +# 0.667 d_min/2 (30, 36, 27) 9.35e-02 1.29e+00 0.5873 3.04e-01 0.9568 +# 1.000 d_min/3 (45, 50, 40) 3.63e-03 3.93e-02 0.999242 2.22e-02 0.999824 +# 1.300 d_min/3.9 (60, 64, 54) 8.94e-04 9.04e-03 0.999960 1.44e-03 0.999999 +# 1.600 d_min/4.8 (72, 80, 64) 9.01e-04 9.12e-03 0.999959 1.34e-03 0.999999 +# 2.200 d_min/6.6 (100,108, 90) 8.97e-04 9.13e-03 0.999959 1.34e-03 0.999999 # # Three things follow. # # 1. **Bare Nyquist is unusable**, which is why ``NYQUIST_OVERSAMPLING`` is 3 and not 2. At -# oversampling 2 the xyz gradient cosine against the analytic answer collapses to 0.42 -# and amplitudes are 10% out. The factor of 3 is buying a great deal. +# oversampling 2 the xyz gradient cosine against the analytic answer collapses to 0.59 +# and amplitudes are 9% out. The factor of 3 is buying a great deal. # 2. **Production sits one step before convergence.** Everything is converged from fineness -# 1.3. At production the residuals are ~5x larger in amplitude and ~7x in the xyz +# 1.3. At production the residuals are ~4x larger in amplitude and in the xyz # gradient, but direction stays excellent (cos 0.999) so the residual is predominantly # magnitude. On a *real* structure the production numbers are better still -- 7L84 gives # amplitude 2.28e-03 and xyz gradient 1.04e-02 -- because derivative aliasing cancels # across atoms as ~1/sqrt(N). See the gate constants in ``__init__.py``; absolute # accuracy gates are calibrated there rather than here. # 3. **The sigma cutoff is not the binding constraint at production.** Sweeping n_sigma at -# fineness 1.0 moves the amplitude residual 5.43e-3 -> 4.11e-3 -> 4.05e-3 and then +# fineness 1.0 moves the amplitude residual 5.22e-3 -> 3.63e-3 -> 3.56e-3 and then # flatlines; grid sampling dominates. Tests that mean to exercise the cutoff therefore # pass an explicit finer ``fineness`` -- see # ``test_forward.py::test_nsigma_reduces_truncation_error``. @@ -200,7 +200,11 @@ def _hkl_within(cell: Cell, d_min: float, dtype: torch.dtype, cap: Optional[int] ] hkl = torch.tensor(cand, dtype=dtype) s = get_scattering_vectors(hkl, cell.data, recB).norm(dim=1) - keep = (s > 0) & (s <= 1.0 / d_min) + # Strictly inside, by a margin far above rounding: when a reflection sits exactly on + # the sphere (a = 24 A with d_min = 1.6 A puts (+-15, 0, 0) there), its membership + # would otherwise hang on the last bit of a*, and the stride below would then shift + # every reflection after it. + keep = (s > 0) & (s < (1.0 / d_min) * (1.0 - 1e-9)) hkl = hkl[keep] if cap is not None and hkl.shape[0] > cap: # Even stride, so the kept set still spans the full resolution range rather diff --git a/torchref/base/__init__.py b/torchref/base/__init__.py index 6428f054..323f0978 100644 --- a/torchref/base/__init__.py +++ b/torchref/base/__init__.py @@ -1,7 +1,7 @@ """ Mathematical functions for crystallographic computations. -This module provides PyTorch and NumPy implementations of: +This module provides PyTorch implementations of: - Coordinate transformations (Cartesian <-> fractional) - Structure factor calculations - R-factor computations @@ -86,10 +86,6 @@ fractional_to_cartesian_torch, get_fractional_matrix, get_inv_fractional_matrix_torch, - cartesian_to_fractional, - fractional_to_cartesian, - get_inv_fractional_matrix, - convert_coords_to_fractional, smallest_diff, smallest_diff_aniso, ) @@ -100,10 +96,7 @@ from .reciprocal import ( # Basis reciprocal_basis_matrix, - reciprocal_basis_matrix_numpy, get_scattering_vectors, - get_scattering_vectors_numpy, - get_s, # HKL get_d_spacing, compute_d_spacing_batch, @@ -159,8 +152,6 @@ ifft, get_real_grid, find_grid_size, - get_real_grid_numpy, - get_grids, put_hkl_on_grid, ) @@ -257,20 +248,13 @@ "fractional_to_cartesian_torch", "get_fractional_matrix", "get_inv_fractional_matrix_torch", - "cartesian_to_fractional", - "fractional_to_cartesian", - "get_inv_fractional_matrix", - "convert_coords_to_fractional", "smallest_diff", "smallest_diff_aniso", # ------------------------------------------------------------------------- # Reciprocal space # ------------------------------------------------------------------------- "reciprocal_basis_matrix", - "reciprocal_basis_matrix_numpy", "get_scattering_vectors", - "get_scattering_vectors_numpy", - "get_s", "get_d_spacing", "compute_d_spacing_batch", "generate_possible_hkl", @@ -311,8 +295,6 @@ "ifft", "get_real_grid", "find_grid_size", - "get_real_grid_numpy", - "get_grids", "put_hkl_on_grid", # ------------------------------------------------------------------------- # alignment diff --git a/torchref/base/coordinates/__init__.py b/torchref/base/coordinates/__init__.py index 5e5dd1cf..a4ab6d41 100644 --- a/torchref/base/coordinates/__init__.py +++ b/torchref/base/coordinates/__init__.py @@ -7,7 +7,8 @@ - Periodic boundary condition handling - Transformation matrix computations -Both PyTorch (GPU-accelerated) and NumPy (CPU) implementations are provided. +All of them are PyTorch functions; the cell metric is written out once, in +:func:`get_fractional_matrix`. """ from .transforms_torch import ( @@ -17,14 +18,6 @@ get_inv_fractional_matrix_torch, ) -from .transforms_numpy import ( - cartesian_to_fractional, - fractional_to_cartesian, - get_fractional_matrix as get_fractional_matrix_numpy, - get_inv_fractional_matrix, - convert_coords_to_fractional, -) - from .periodic_boundary import ( smallest_diff, smallest_diff_aniso, @@ -43,12 +36,6 @@ "fractional_to_cartesian_torch", "get_fractional_matrix", "get_inv_fractional_matrix_torch", - # NumPy implementations - "cartesian_to_fractional", - "fractional_to_cartesian", - "get_fractional_matrix_numpy", - "get_inv_fractional_matrix", - "convert_coords_to_fractional", # Periodic boundary "smallest_diff", "smallest_diff_aniso", diff --git a/torchref/base/coordinates/transforms_numpy.py b/torchref/base/coordinates/transforms_numpy.py deleted file mode 100644 index 641ddaa1..00000000 --- a/torchref/base/coordinates/transforms_numpy.py +++ /dev/null @@ -1,139 +0,0 @@ -""" -NumPy implementations of coordinate transformation functions. - -These functions provide CPU-based coordinate transformations -for use when GPU acceleration is not needed or available. -""" - -import numpy as np - - -def get_fractional_matrix(cell): - """ - Calculate the fractional-to-Cartesian transformation matrix. - - Constructs the matrix B that transforms fractional coordinates to - Cartesian coordinates based on the unit cell parameters. - - Parameters - ---------- - cell : numpy.ndarray or list - Unit cell parameters [a, b, c, alpha, beta, gamma] where lengths are - in Angstroms and angles are in degrees. - - Returns - ------- - numpy.ndarray - 3x3 transformation matrix B such that cart = frac @ B.T. - """ - a, b, c = cell[:3] - alpha, beta, gamma = np.radians(cell[3:]) - cos_alpha, cos_beta, cos_gamma = np.cos(alpha), np.cos(beta), np.cos(gamma) - sin_gamma = np.sin(gamma) - volume = np.sqrt( - 1 - - cos_alpha**2 - - cos_beta**2 - - cos_gamma**2 - + 2 * cos_alpha * cos_beta * cos_gamma - ) - B = np.array( - [ - [a, b * cos_gamma, c * cos_beta], - [0, b * sin_gamma, c * (cos_alpha - cos_beta * cos_gamma) / sin_gamma], - [0, 0, c * volume / sin_gamma], - ] - ) - return B - - -def get_inv_fractional_matrix(cell): - """ - Calculate the Cartesian-to-fractional transformation matrix. - - Computes the inverse of the fractional matrix for converting Cartesian - coordinates to fractional coordinates. - - Parameters - ---------- - cell : numpy.ndarray or list - Unit cell parameters [a, b, c, alpha, beta, gamma] where lengths are - in Angstroms and angles are in degrees. - - Returns - ------- - numpy.ndarray - 3x3 inverse transformation matrix B_inv such that frac = cart @ B_inv.T. - """ - B = get_fractional_matrix(cell) - B_inv = np.linalg.inv(B) - return B_inv - - -def cartesian_to_fractional(xyz, cell): - """ - Convert Cartesian coordinates to fractional coordinates. - - Parameters - ---------- - xyz : numpy.ndarray - Cartesian coordinates with shape (N, 3). - cell : numpy.ndarray or list - Unit cell parameters [a, b, c, alpha, beta, gamma] where lengths are - in Angstroms and angles are in degrees. - - Returns - ------- - numpy.ndarray - Fractional coordinates with shape (N, 3). - """ - B_inv = get_inv_fractional_matrix(cell) - xyz_fractional = np.dot(xyz, B_inv.T) - return xyz_fractional - - -def fractional_to_cartesian(xyz_fractional, cell): - """ - Convert fractional coordinates to Cartesian coordinates. - - Parameters - ---------- - xyz_fractional : numpy.ndarray - Fractional coordinates with shape (N, 3). - cell : numpy.ndarray or list - Unit cell parameters [a, b, c, alpha, beta, gamma] where lengths are - in Angstroms and angles are in degrees. - - Returns - ------- - numpy.ndarray - Cartesian coordinates with shape (N, 3). - """ - B = get_fractional_matrix(cell) - xyz = np.dot(xyz_fractional, B.T) - return xyz - - -def convert_coords_to_fractional(df, cell): - """ - Convert coordinates from a DataFrame to fractional coordinates. - - Extracts x, y, z columns from a DataFrame and converts them from - Cartesian to fractional coordinates. - - Parameters - ---------- - df : pandas.DataFrame - DataFrame containing 'x', 'y', 'z' columns with Cartesian coordinates. - cell : numpy.ndarray or list - Unit cell parameters [a, b, c, alpha, beta, gamma] where lengths are - in Angstroms and angles are in degrees. - - Returns - ------- - numpy.ndarray - Fractional coordinates with shape (N, 3). - """ - xyz = df[["x", "y", "z"]].values - xyz_fractional = cartesian_to_fractional(xyz, cell) - return xyz_fractional diff --git a/torchref/base/coordinates/transforms_torch.py b/torchref/base/coordinates/transforms_torch.py index 4e21bfe0..9c2d50df 100644 --- a/torchref/base/coordinates/transforms_torch.py +++ b/torchref/base/coordinates/transforms_torch.py @@ -4,10 +4,9 @@ These functions are GPU-accelerated and support automatic differentiation for use in optimization and refinement. -.. note:: - ``get_fractional_matrix`` is an exception: it assembles its output via - ``torch.tensor([...])``, which detaches from the autograd graph, so - gradients do not flow back to the input ``cell`` parameters. +:func:`get_fractional_matrix` is the one place the cell metric is written out; the +reciprocal basis (:func:`~torchref.base.reciprocal.basis.reciprocal_basis_matrix`) and +:attr:`torchref.symmetry.Cell.volume` are derived from it. """ import torch @@ -80,33 +79,31 @@ def get_fractional_matrix(cell): Returns ------- torch.Tensor - 3x3 transformation matrix B such that cart = frac @ B.T. - - Notes - ----- - The matrix is assembled with ``torch.tensor([...])``, which detaches the - result from the autograd graph. Gradients therefore do not propagate back - to ``cell``; this function is non-differentiable in the cell parameters. + 3x3 upper-triangular matrix B such that cart = frac @ B.T, in the PDB + orientation (a along x, b in the xy plane), on ``cell``'s dtype and device. + Differentiable in ``cell``. """ - a, b, c = cell[:3] + a, b, c = cell[0], cell[1], cell[2] alpha, beta, gamma = torch.deg2rad(cell[3:]) cos_alpha, cos_beta, cos_gamma = torch.cos(alpha), torch.cos(beta), torch.cos(gamma) sin_gamma = torch.sin(gamma) - volume = torch.sqrt( + volume_factor = torch.sqrt( 1 - cos_alpha**2 - cos_beta**2 - cos_gamma**2 + 2 * cos_alpha * cos_beta * cos_gamma ) - B = torch.tensor( + zero = torch.zeros_like(a) + return torch.stack( [ - [a, b * cos_gamma, c * cos_beta], - [0, b * sin_gamma, c * (cos_alpha - cos_beta * cos_gamma) / sin_gamma], - [0, 0, c * volume / sin_gamma], - ], dtype=cell.dtype, device=cell.device + torch.stack([a, b * cos_gamma, c * cos_beta]), + torch.stack( + [zero, b * sin_gamma, c * (cos_alpha - cos_beta * cos_gamma) / sin_gamma] + ), + torch.stack([zero, zero, c * volume_factor / sin_gamma]), + ] ) - return B def get_inv_fractional_matrix_torch(cell): diff --git a/torchref/base/fourier/__init__.py b/torchref/base/fourier/__init__.py index 4131e1ce..0cd79762 100644 --- a/torchref/base/fourier/__init__.py +++ b/torchref/base/fourier/__init__.py @@ -7,8 +7,6 @@ from .grid import ( get_real_grid, find_grid_size, - get_real_grid_numpy, - get_grids, put_hkl_on_grid, ) @@ -19,8 +17,6 @@ # Grid functions "get_real_grid", "find_grid_size", - "get_real_grid_numpy", - "get_grids", "put_hkl_on_grid", # Map coefficients "map_coefficients", diff --git a/torchref/base/fourier/grid.py b/torchref/base/fourier/grid.py index 4e824757..79956865 100644 --- a/torchref/base/fourier/grid.py +++ b/torchref/base/fourier/grid.py @@ -15,9 +15,6 @@ fractional_to_cartesian_torch, get_fractional_matrix, ) -from torchref.base.coordinates.transforms_numpy import ( - fractional_to_cartesian, -) def get_real_grid(cell=None, fractional_matrix=None, max_res=0.8, gridsize=None, device=None): @@ -102,85 +99,6 @@ def find_grid_size(cell: torch.Tensor, max_res: float): return torch.floor(cell[:3] / max_res * NYQUIST_OVERSAMPLING).to(dtypes.int) -def get_real_grid_numpy(cell, max_res=0.8, gridsize=None): - """ - Generate a real-space grid of Cartesian coordinates (NumPy version). - - Creates a 3D grid in fractional coordinates and converts it to Cartesian - coordinates. Grid points are placed at cell edges following CCTBX convention. - - Parameters - ---------- - cell : numpy.ndarray or list - Unit cell parameters [a, b, c, alpha, beta, gamma] where lengths are - in Angstroms and angles are in degrees. - max_res : float, optional - Maximum resolution in Angstroms for grid spacing. Default is 0.8. - Ignored if gridsize is provided. - gridsize : list or numpy.ndarray, optional - Explicit grid dimensions [nx, ny, nz]. If provided, overrides max_res. - - Returns - ------- - numpy.ndarray - Real-space grid coordinates with shape (nx, ny, nz, 3). - """ - if gridsize is not None: - nsteps = np.array(gridsize, dtype=int) - else: - nsteps = np.astype(np.floor(cell[:3] / max_res * NYQUIST_OVERSAMPLING), int) - x = np.arange(nsteps[0]) / nsteps[0] - y = np.arange(nsteps[1]) / nsteps[1] - z = np.arange(nsteps[2]) / nsteps[2] - x, y, z = np.meshgrid(x, y, z, indexing="ij") - array_shape = x.shape - x = x.reshape((*x.shape, 1)) - y = y.reshape((*y.shape, 1)) - z = z.reshape((*z.shape, 1)) - xyz = np.concatenate((x, y, z), axis=3).reshape(-1, 3) - xyz_real_grid = fractional_to_cartesian(xyz, cell) - xyz_real_grid = xyz_real_grid.reshape((*array_shape, 3)) - return xyz_real_grid - - -def get_grids(cell, max_res=0.8): - """ - Generate real-space and reciprocal-space grids for Fourier transforms. - - Creates a 3D grid in fractional coordinates and converts it to Cartesian - coordinates, along with an empty reciprocal space grid. - - Parameters - ---------- - cell : numpy.ndarray or list - Unit cell parameters [a, b, c, alpha, beta, gamma] where lengths are - in Angstroms and angles are in degrees. - max_res : float, optional - Maximum resolution in Angstroms for grid spacing. Default is 0.8. - - Returns - ------- - recgrid : numpy.ndarray - Empty reciprocal space grid with shape determined by resolution. - xyz_real_grid : numpy.ndarray - Real-space grid coordinates with shape (nx, ny, nz, 3). - """ - nsteps = np.astype(np.floor(cell[:3] / max_res * NYQUIST_OVERSAMPLING), int) - x = np.arange(nsteps[0]) / nsteps[0] - y = np.arange(nsteps[1]) / nsteps[1] - z = np.arange(nsteps[2]) / nsteps[2] - x, y, z = np.meshgrid(x, y, z, indexing="ij") - array_shape = x.shape - x = x.reshape((*x.shape, 1)) - y = y.reshape((*y.shape, 1)) - z = z.reshape((*z.shape, 1)) - xyz = np.concatenate((x, y, z), axis=3).reshape(-1, 3) - xyz_real_grid = fractional_to_cartesian(xyz, cell) - xyz_real_grid = xyz_real_grid.reshape((*array_shape, 3)) - recgrid = np.zeros(array_shape, dtype=float) - return recgrid, xyz_real_grid - - def put_hkl_on_grid(real_space_grid, diff, hkl): """ Place structure factors on a zero-filled reciprocal space grid. diff --git a/torchref/base/reciprocal/__init__.py b/torchref/base/reciprocal/__init__.py index cdaf6c20..cc048cf3 100644 --- a/torchref/base/reciprocal/__init__.py +++ b/torchref/base/reciprocal/__init__.py @@ -7,10 +7,7 @@ from .basis import ( reciprocal_basis_matrix, - reciprocal_basis_matrix_numpy, get_scattering_vectors, - get_scattering_vectors_numpy, - get_s, ) from .hkl import ( @@ -38,10 +35,7 @@ __all__ = [ # Basis functions "reciprocal_basis_matrix", - "reciprocal_basis_matrix_numpy", "get_scattering_vectors", - "get_scattering_vectors_numpy", - "get_s", # HKL functions "generate_possible_hkl", "get_d_spacing", diff --git a/torchref/base/reciprocal/basis.py b/torchref/base/reciprocal/basis.py index 18238edc..b72a7819 100644 --- a/torchref/base/reciprocal/basis.py +++ b/torchref/base/reciprocal/basis.py @@ -5,14 +5,20 @@ from unit cell parameters. """ -import numpy as np import torch +from torchref.base.coordinates.transforms_torch import get_inv_fractional_matrix_torch + def reciprocal_basis_matrix(cell: torch.Tensor): """ Compute the reciprocal space basis matrix from unit cell parameters. + The rows of the Cartesian-to-fractional matrix are a*, b*, c*, so this is + the inverse of + :func:`~torchref.base.coordinates.transforms_torch.get_fractional_matrix` and + shares its cell metric. + Parameters ---------- cell : torch.Tensor @@ -23,94 +29,9 @@ def reciprocal_basis_matrix(cell: torch.Tensor): Returns ------- torch.Tensor - Reciprocal basis matrix of shape (3, 3) with a*, b*, c* as rows. - """ - # Extract cell parameters - angles_rad = torch.deg2rad(cell[3:]) - # Compute real-space basis vectors - angles_cos = torch.cos(angles_rad) - cos_squared = angles_cos**2 - sin_gamma = torch.sin(angles_rad[2]) - volume = torch.sqrt( - 1 - - cos_squared[0] - - cos_squared[1] - - cos_squared[2] - + 2 * cos_squared[0] * cos_squared[1] * cos_squared[2] - ) - a_vec = torch.tensor( - [cell[0], 0, 0], dtype=cell.dtype, device=cell.device - ) - b_vec = torch.tensor( - [cell[1] * angles_cos[2], cell[1] * sin_gamma, 0], - dtype=cell.dtype, - device=cell.device, - ) - c_vec = torch.tensor( - [ - cell[2] * angles_cos[1], - cell[2] * (angles_cos[0] - angles_cos[1] * angles_cos[2]) / sin_gamma, - cell[2] * volume / sin_gamma, - ], - dtype=cell.dtype, - device=cell.device, - ) - # Compute reciprocal basis vectors - volume_real = torch.dot(a_vec, torch.linalg.cross(b_vec, c_vec)) - a_star = torch.linalg.cross(b_vec, c_vec) / volume_real - b_star = torch.linalg.cross(c_vec, a_vec) / volume_real - c_star = torch.linalg.cross(a_vec, b_vec) / volume_real - # Assemble reciprocal basis matrix - return torch.stack([a_star, b_star, c_star]) - - -def reciprocal_basis_matrix_numpy(cell): + Reciprocal basis matrix of shape (3, 3) with a*, b*, c* as rows, in Å⁻¹. """ - Calculate the reciprocal basis matrix from unit cell parameters (NumPy version). - - Computes the reciprocal space basis vectors (a*, b*, c*) that define - the transformation from Miller indices to scattering vectors. - - Parameters - ---------- - cell : numpy.ndarray or list - Unit cell parameters [a, b, c, alpha, beta, gamma] where lengths are - in Angstroms and angles are in degrees. - - Returns - ------- - numpy.ndarray - 3x3 matrix containing reciprocal basis vectors as rows [a*, b*, c*]. - """ - # Extract cell parameters - a, b, c, alpha, beta, gamma = cell - alpha, beta, gamma = np.radians([alpha, beta, gamma]) - # Compute real-space basis vectors - cos_alpha, cos_beta, cos_gamma = np.cos(alpha), np.cos(beta), np.cos(gamma) - sin_gamma = np.sin(gamma) - volume = np.sqrt( - 1 - - cos_alpha**2 - - cos_beta**2 - - cos_gamma**2 - + 2 * cos_alpha * cos_beta * cos_gamma - ) - a_vec = np.array([a, 0, 0]) - b_vec = np.array([b * cos_gamma, b * sin_gamma, 0]) - c_vec = np.array( - [ - c * cos_beta, - c * (cos_alpha - cos_beta * cos_gamma) / sin_gamma, - c * volume / sin_gamma, - ] - ) - # Compute reciprocal basis vectors - volume_real = np.dot(a_vec, np.cross(b_vec, c_vec)) - a_star = np.cross(b_vec, c_vec) / volume_real - b_star = np.cross(c_vec, a_vec) / volume_real - c_star = np.cross(a_vec, b_vec) / volume_real - # Assemble reciprocal basis matrix - return np.array([a_star, b_star, c_star]) + return get_inv_fractional_matrix_torch(cell) def get_scattering_vectors(hkl: torch.Tensor, cell: torch.Tensor, recB=None): @@ -135,56 +56,3 @@ def get_scattering_vectors(hkl: torch.Tensor, cell: torch.Tensor, recB=None): recB = reciprocal_basis_matrix(cell) s = torch.matmul(hkl.to(cell.dtype), recB) return s - - -def get_scattering_vectors_numpy(hkl, cell): - """ - Calculate scattering vectors from Miller indices and unit cell (NumPy version). - - Transforms Miller indices to reciprocal space scattering vectors - using the reciprocal basis matrix. - - Parameters - ---------- - hkl : numpy.ndarray or list - Miller indices with shape (N, 3). - cell : numpy.ndarray or list - Unit cell parameters [a, b, c, alpha, beta, gamma] where lengths are - in Angstroms and angles are in degrees. - - Returns - ------- - numpy.ndarray - Scattering vectors in reciprocal space with shape (N, 3). - """ - recB = reciprocal_basis_matrix_numpy(cell) - hkl = np.array(hkl) # Ensure hkl is a numpy array - s = np.dot(hkl, recB) - return s - - -def get_s(hkl, cell): - """ - Calculate the magnitude of scattering vectors for given Miller indices. - - Computes |s| = 1/d where d is the interplanar spacing for each reflection. - - This is a NumPy-only helper (it delegates to the NumPy scattering-vector - path); it does not accept or return torch tensors. - - Parameters - ---------- - hkl : numpy.ndarray - Miller indices with shape (N, 3). - cell : numpy.ndarray or list - Unit cell parameters [a, b, c, alpha, beta, gamma] where lengths are - in Angstroms and angles are in degrees. - - Returns - ------- - numpy.ndarray - Magnitude of scattering vectors with shape (N,). - """ - s = get_scattering_vectors_numpy(hkl, cell) - s = np.sum(s**2, axis=1) ** 0.5 - return s diff --git a/torchref/symmetry/cell.py b/torchref/symmetry/cell.py index 63a47e94..73616d7c 100644 --- a/torchref/symmetry/cell.py +++ b/torchref/symmetry/cell.py @@ -308,17 +308,16 @@ def reciprocal_basis_matrix(self) -> torch.Tensor: """ Reciprocal basis matrix with [a*, b*, c*] as rows. + The rows of the fractionalization matrix are the reciprocal basis vectors, + so this returns the same cached tensor as :attr:`inv_fractional_matrix`; + do not modify it in place. + Returns ------- torch.Tensor - Shape (3, 3) matrix where rows are the reciprocal basis vectors. + Shape (3, 3) matrix where rows are the reciprocal basis vectors, in Å⁻¹. """ - self._assert_unmodified() - if "reciprocal_basis_matrix" not in self._cache: - self._cache["reciprocal_basis_matrix"] = ( - self._compute_reciprocal_basis_matrix() - ) - return self._cache["reciprocal_basis_matrix"] + return self.inv_fractional_matrix # ========================================================================= # Internal computation methods @@ -331,26 +330,8 @@ def _compute_fractional_matrix(self) -> torch.Tensor: return math_torch.get_fractional_matrix(self._data) def _compute_volume(self) -> torch.Tensor: - """V = abc·sqrt(1 - Σcos²angle + 2·cosα·cosβ·cosγ).""" - a, b, c = self._data[0], self._data[1], self._data[2] - angles_rad = torch.deg2rad(self._data[3:]) - cos_alpha, cos_beta, cos_gamma = torch.cos(angles_rad) - - volume_factor = torch.sqrt( - 1 - - cos_alpha**2 - - cos_beta**2 - - cos_gamma**2 - + 2 * cos_alpha * cos_beta * cos_gamma - ) - - return a * b * c * volume_factor - - def _compute_reciprocal_basis_matrix(self) -> torch.Tensor: - """Reciprocal basis, via ``math_torch.reciprocal_basis_matrix``.""" - from torchref.base import math_torch - - return math_torch.reciprocal_basis_matrix(self._data) + """V = det(B); B is upper triangular, so that is its diagonal product.""" + return torch.diagonal(self.fractional_matrix).prod() # ========================================================================= # Grid computation methods From eb41a10ffc3c3ba6c41b3a3d4017f36297f6e113 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 20:06:26 +0000 Subject: [PATCH 242/250] Compare the held-out scaling fit within float32 jitter, not bit for bit test_free_and_validation_changes_do_not_affect_fit asserted that two fits differing only in held-out amplitudes are bit-identical. The fit is not bit-reproducible run to run: on unchanged dev, repeated fits with the same inputs already differ by one ulp (3e-8) in a parameter, so the assertion held only by the luck of reduction order and flipped with unrelated allocation changes. A leak is orders of magnitude larger (filling 30 work reflections moves the parameters by 6e-4), so a 1e-6 absolute tolerance keeps the check. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- tests/integration/test_dataset_scaler.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/tests/integration/test_dataset_scaler.py b/tests/integration/test_dataset_scaler.py index ec858edc..938ebb36 100644 --- a/tests/integration/test_dataset_scaler.py +++ b/tests/integration/test_dataset_scaler.py @@ -201,7 +201,10 @@ def test_free_and_validation_changes_do_not_affect_fit(loaded_reflection_data): assert torch.equal(dc.hkl, dc.scaler.hkl) assert not dc.scaler.fit_mask[:, mismatched].any() results.append(dc.scaler.raw_parameters.detach().clone()) - assert torch.equal(*results) + # Not bit equality: the fit's float32 reductions differ by an ulp (~3e-8) between + # otherwise identical runs, while filling just 30 work reflections the same way + # moves the parameters by ~6e-4. + torch.testing.assert_close(*results, rtol=0.0, atol=1e-6) def test_permutation_and_partial_overlap_chain(loaded_reflection_data): From da2c489061ffc55a7f6d065c908c2a9a1482b583 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 21:56:06 +0000 Subject: [PATCH 243/250] Group occupancies once, by topology residue, and restore saved groups ModelContext.occupancy_groups seeded every atom with its own index and numbered groups from 0 in the same space, so atoms it left ungrouped (the independent atoms of residues whose occupancies disagree, and the blank-altloc atoms of altloc residues) were merged by the compaction step with whichever unrelated group had the same number. Deposited occupancies changed silently on load (6G9X lost seven partial S occupancies, 5BOV GLN A46 N went from 1.00 to 0.66). Its residue key (resname, resseq, chain) also dropped the insertion code and pooled residues 100 and 100A, which Model.strip_altlocs had already been taught to keep apart in its own loop. OccupancyTensor.from_residue_groups was a third copy of the grouping, used only by its own test, carrying a collision fix the live path never got. There is now one implementation. ModelContext._residue_parts walks the topology residues (chain, resseq, icode); altloc_residues derives each residue's conformers from it, and occupancy_groups, register_altlocs and Model.strip_altlocs all take conformers from there. occupancy_groups gives every atom exactly one group id in sequence, so there is no compaction and no group spans two residues. Alternates with different residue names at one position are now conformers of one residue and sum to 1. from_residue_groups and its test are deleted. A restore re-derived the groups from the checkpoint's occupancy values, so any checkpoint whose values sat on the other side of the 0.01 sharing deadband failed with IndexError or a size mismatch: every 3E98, 5BOV and 6G9X checkpoint (because of the merges above), and any checkpoint after occupancy refinement. Model._build_occupancy now builds the wrapper with OccupancyTensor.from_saved_groups from the saved expansion_mask, linked_occ_ and refinable_mask, so the checkpoint is the source of truth for the grouping; only a fresh load groups atoms from values. OccupancyTensor.copy shares the linked-buffer decoding. Checkpoints written before this change restore exactly as saved. Model.load_state on an empty Model() left xyz, adp, u and occupancy None: the constructor's None placeholders are plain attributes and survived the __dict__ merge, shadowing the restored modules. It now replaces the instance state. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- tests/unit/model/test_from_residue_groups.py | 46 ---- tests/unit/model/test_occupancy_groups.py | 186 +++++++++++++++ tests/unit/model/test_strip_altlocs.py | 10 +- torchref/model/context.py | 229 +++++++++++-------- torchref/model/model.py | 80 ++++--- torchref/model/parameter_wrappers.py | 122 +++++----- 6 files changed, 425 insertions(+), 248 deletions(-) delete mode 100644 tests/unit/model/test_from_residue_groups.py create mode 100644 tests/unit/model/test_occupancy_groups.py diff --git a/tests/unit/model/test_from_residue_groups.py b/tests/unit/model/test_from_residue_groups.py deleted file mode 100644 index 8d419a83..00000000 --- a/tests/unit/model/test_from_residue_groups.py +++ /dev/null @@ -1,46 +0,0 @@ -"""Regression test: from_residue_groups must not merge singletons into groups. - -It labeled singleton atoms with arange ids (0..n-1) and multi-atom groups with -ids 0,1,2,... in the SAME number space, then compacted with torch.unique. A -group id equal to a surviving singleton's arange id silently merged them. The -fix starts group ids at n_atoms. See TORCHREF_AUDIT.md. - -Layout below reproduces the documented collision: residues (by resseq) iterate -as [0,1], [2], [3,4,5], [6], [7,8], [9]; pre-fix the [7,8] group was assigned -id 2, colliding with singleton atom index 2. -""" - -import pandas as pd -import pytest -import torch - -from torchref.model.parameter_wrappers import OccupancyTensor - - -@pytest.mark.unit -def test_from_residue_groups_no_singleton_collision(): - df = pd.DataFrame( - { - "index": list(range(10)), - "resname": ["ALA"] * 10, - "resseq": [1, 1, 2, 3, 3, 3, 4, 5, 5, 6], - "chainid": ["A"] * 10, - "altloc": [""] * 10, - } - ) - init = torch.full((10,), 0.9) - - occ = OccupancyTensor.from_residue_groups(init, df) - g = occ.expansion_mask.cpu().tolist() # per-atom collapsed group id - - # The real sharing groups: - assert g[0] == g[1] - assert g[3] == g[4] == g[5] - assert g[7] == g[8] - # The regression: singleton atom 2 must NOT be merged into the [7,8] group. - assert g[2] != g[7] - # All six residues are distinct groups (one id per residue). - assert len({g[0], g[2], g[3], g[6], g[7], g[9]}) == 6 - # Singletons are independent of every other atom. - for singleton in (2, 6, 9): - assert sum(1 for x in g if x == g[singleton]) == 1 diff --git a/tests/unit/model/test_occupancy_groups.py b/tests/unit/model/test_occupancy_groups.py new file mode 100644 index 00000000..3f2a7cbb --- /dev/null +++ b/tests/unit/model/test_occupancy_groups.py @@ -0,0 +1,186 @@ +"""Occupancy grouping on load, and its round trip through a checkpoint. + +A sharing group never spans two residues, residues are ``(chain, resseq, icode)`` so +100 and 100A are two, and loading keeps every deposited occupancy except where the +atoms of one altloc conformer disagree (a conformer is one group by design). A +checkpoint restores the grouping it was saved with, whatever the occupancies have +become. +""" + +import numpy as np +import pytest +import torch + +from torchref.io.pdb import PDBReader +from torchref.model.model import Model + +pytestmark = pytest.mark.unit + +GLY = ("N", "CA", "C", "O") +MET = ("N", "CA", "C", "O", "CB", "CG", "SD", "CE") +THR = ("N", "CA", "C", "O", "CB", "OG1", "CG2") +ALA = ("N", "CA", "C", "O", "CB") + +#: ``(resname, resseq, icode, [(atom name, altloc, occupancy), ...])``, chain A. +INSERTION_CODES = [ + ("GLY", 99, "", [(n, " ", 1.0) for n in GLY]), + ("GLY", 100, "", [(n, "A", 0.7) for n in GLY] + [(n, "B", 0.3) for n in GLY]), + ("GLY", 100, "A", [(n, "A", 0.4) for n in GLY] + [(n, "B", 0.6) for n in GLY]), + ("GLY", 101, "", [(n, " ", 0.8) for n in GLY]), + ("GLY", 101, "A", [(n, " ", 0.5) for n in GLY]), + ("MET", 102, "", [(n, " ", 0.38 if n == "SD" else 1.0) for n in MET]), +] + +#: THR and ALA alternates at one position, between two glycines. +MICROHETEROGENEITY = [ + ("GLY", 29, "", [(n, " ", 1.0) for n in GLY]), + ("THR", 30, "", [(n, "A", 0.35) for n in THR]), + ("ALA", 30, "", [(n, "B", 0.65) for n in ALA]), + ("GLY", 31, "", [(n, " ", 1.0) for n in GLY]), +] + +STRUCTURES = ["3E98", "5BOV", "6G9X"] + + +def _write_pdb(path, residues): + """Write ``residues`` as a P1 PDB file and return the deposited occupancies.""" + lines = ["CRYST1 40.000 40.000 40.000 90.00 90.00 90.00 P 1 1"] + occupancies = [] + for resname, resseq, icode, atoms in residues: + for name, altloc, occupancy in atoms: + serial = len(occupancies) + 1 + lines.append( + f"ATOM {serial:5d} {name:<3s}{altloc}{resname} A{resseq:4d}{icode:1s}" + f" {0.8 * serial:8.3f}{5.0:8.3f}{5.0:8.3f}{occupancy:6.2f}{20.0:6.2f}" + f" {name[0]:>2s}" + ) + occupancies.append(occupancy) + lines.append("END") + path.write_text("\n".join(lines) + "\n") + return np.asarray(occupancies) + + +def _load(path): + return Model(verbose=0, device="cpu").load_pdb(str(path)) + + +def _residues_per_group(model): + """Number of distinct topology residues in each sharing group.""" + groups = model.occupancy.expansion_mask.cpu().numpy() + residue = model.ctx.topology.atoms.residue_of.cpu().numpy() + pairs = np.unique(np.stack([groups, residue]), axis=1) + return np.bincount(pairs[0]) + + +def _deposited(path): + """The reader's table with the rows a model drops removed, and its occupancies.""" + table, _, _ = PDBReader(verbose=0).read(str(path))() + table = table.dropna(subset=["x", "y", "z", "tempfactor", "occupancy"]) + table = table.reset_index(drop=True) + return table, table["occupancy"].clip(0, 1).to_numpy() + + +def _disagreeing_conformers(table): + """Row lists of altloc conformers whose deposited occupancies are not uniform.""" + altloc = table["altloc"].astype(str).str.strip() + out = [] + for _, residue in table[altloc != ""].groupby(["chainid", "resseq", "icode"]): + if residue["altloc"].nunique() < 2: + continue + for _, conformer in residue.groupby("altloc"): + if conformer["occupancy"].nunique() > 1: + out.append(conformer.index.to_numpy()) + return out + + +def test_insertion_codes_and_ungrouped_atoms_keep_their_occupancies(tmp_path): + """100/100A stay apart, and atoms left ungrouped never join another residue.""" + deposited = _write_pdb(tmp_path / "icode.pdb", INSERTION_CODES) + model = _load(tmp_path / "icode.pdb") + + np.testing.assert_allclose(model.occupancy().detach().numpy(), deposited, atol=1e-5) + assert (_residues_per_group(model) == 1).all() + # 100 and 100A: two conformers each; 99, 101 and 101A: one group each; MET 102's + # atoms disagree, so one group per atom. + assert model.occupancy.collapsed_shape == (2 + 2 + 3 + len(MET),) + pairs = [[rows.tolist() for rows in pair] for pair in model.ctx.altloc_pairs] + assert pairs == [ + [[4, 5, 6, 7], [8, 9, 10, 11]], + [[12, 13, 14, 15], [16, 17, 18, 19]], + ] + + +def test_microheterogeneous_alternates_are_one_residues_conformers(tmp_path): + """THR A and ALA B at one position are linked conformers that sum to 1.""" + deposited = _write_pdb(tmp_path / "micro.pdb", MICROHETEROGENEITY) + model = _load(tmp_path / "micro.pdb") + + ((residue, labels, conformers),) = model.ctx.altloc_residues() + assert model.ctx.topology.residues.key(residue) == ("A", 30, "") + assert labels == ["A", "B"] + assert (len(conformers["A"]), len(conformers["B"])) == (len(THR), len(ALA)) + np.testing.assert_allclose(model.occupancy().detach().numpy(), deposited, atol=1e-5) + + model.occupancy[torch.tensor(conformers["B"])] = 0.9 + occupancy = model.occupancy().detach() + total = occupancy[conformers["A"][0]] + occupancy[conformers["B"][0]] + assert total.item() == pytest.approx(1.0, abs=1e-5) + + +@pytest.mark.parametrize("code", STRUCTURES) +def test_loading_keeps_every_deposited_occupancy(pdb_dir, code): + """Values survive the load; only a disagreeing conformer collapses to one value.""" + table, deposited = _deposited(pdb_dir / f"{code}.pdb") + model = _load(pdb_dir / f"{code}.pdb") + loaded = model.occupancy().detach().numpy() + + assert (_residues_per_group(model) == 1).all() + shared = np.zeros(len(table), dtype=bool) + for rows in _disagreeing_conformers(table): + shared[rows] = True + assert np.ptp(loaded[rows]) < 1e-6 + np.testing.assert_allclose(loaded[~shared], deposited[~shared], atol=1e-5) + + +@pytest.mark.parametrize("code", STRUCTURES) +def test_checkpoint_round_trip_keeps_grouping_and_values(pdb_dir, code, tmp_path): + """``create_from_state_dict``, and ``load_state`` into an empty model, restore + groups and values.""" + model = _load(pdb_dir / f"{code}.pdb") + path = tmp_path / "model.pt" + model.save_state(str(path)) + from_file = Model(verbose=0, device="cpu") + from_file.load_state(str(path)) + from_dict = Model.create_from_state_dict( + model.state_dict(), device="cpu", verbose=0 + ) + + for restored in (from_dict, from_file): + assert torch.equal( + restored.occupancy.expansion_mask, model.occupancy.expansion_mask + ) + assert torch.equal( + restored.occupancy.refinable_mask, model.occupancy.refinable_mask + ) + assert torch.equal(restored.occupancy(), model.occupancy()) + + +def test_checkpoint_restores_the_saved_grouping_after_values_move(pdb_dir): + """Occupancies moved across the sharing deadband restore into the saved groups. + + With every occupancy at 1.0 a fresh load would pool far more atoms than the + saved groups hold. + """ + model = _load(pdb_dir / "6G9X.pdb") + model.occupancy[:] = 1.0 + + restored = Model.create_from_state_dict(model.state_dict(), device="cpu", verbose=0) + + assert torch.equal( + restored.occupancy.expansion_mask, model.occupancy.expansion_mask + ) + assert torch.equal(restored.occupancy(), model.occupancy()) + assert ( + restored.occupancy.get_refinable_count() + == model.occupancy.get_refinable_count() + ) diff --git a/tests/unit/model/test_strip_altlocs.py b/tests/unit/model/test_strip_altlocs.py index c87836ba..dfb00cbd 100644 --- a/tests/unit/model/test_strip_altlocs.py +++ b/tests/unit/model/test_strip_altlocs.py @@ -138,9 +138,13 @@ def test_microheterogeneity_identity_and_current_occupancy(microheterogeneous_mo at_position = (cols["chain"] == "A") & (cols["resseq"] == 30) assert len(set(model.ctx.topology.atoms.residue_of[at_position].tolist())) == 1 assert (cols["resname"][at_position & (cols["altloc"] == "B")] == "ALA").all() - groups = model.ctx._residue_groups(with_altloc=True) - assert len(groups[("ALA", 30, "A", "B")]) == 5 - assert len(groups[("THR", 30, "A", "A")]) > 5 + (conformers,) = [ + rows + for residue, _, rows in model.ctx.altloc_residues() + if model.ctx.topology.residues.key(residue) == ("A", 30, "") + ] + assert len(conformers["B"]) == 5 + assert len(conformers["A"]) > 5 np.testing.assert_allclose( model.occupancy().detach().numpy()[at_position & (cols["altloc"] == "B")], 0.65 ) diff --git a/torchref/model/context.py b/torchref/model/context.py index d69d914b..45a43b63 100644 --- a/torchref/model/context.py +++ b/torchref/model/context.py @@ -499,124 +499,151 @@ def build_restraints( # Identity queries # ------------------------------------------------------------------ - def _residue_groups(self, with_altloc: bool) -> Dict[tuple, List[int]]: - """Atom rows grouped by ``(resname, resseq, chain[, altloc])``, keys sorted. - - The key order is the one pandas' sorted ``groupby`` gives. Occupancy groups are - numbered in it, and checkpoints store occupancies in group space, so it must - not change. + def _residue_parts(self) -> List[Tuple[int, Dict[Tuple[str, str], List[int]]]]: + """Each residue's atom rows, split by ``(resname, altloc)``, in a fixed order. + + A residue is a topology residue, ``(chain, resseq, icode)``, the identity the + restraints use: 100 and 100A are two residues, while alternates with different + residue names at one position (microheterogeneity) are parts of one. Residues + are sorted by ``(resname, resseq, chain, icode)`` and parts by + ``(resname, altloc)``, so occupancy groups are numbered the same whatever + order the file lists its residues in. A blank altloc is ``" "``. """ - columns = self.topology.columns() - altloc = np.where(columns["altloc"] == " ", "", columns["altloc"]) - keys: Dict[tuple, List[int]] = {} - for row in range(self.topology.n_atoms): - key = ( - str(columns["resname"][row]), - int(columns["resseq"][row]), - str(columns["chain"][row]), - ) - if with_altloc: - key = key + (str(altloc[row]),) - keys.setdefault(key, []).append(row) - return {key: keys[key] for key in sorted(keys)} + residues = self.topology.residues + resname = self.topology.columns()["resname"] + altloc = self.topology.atoms.altloc + order = sorted( + range(residues.n_residues), + key=lambda r: ( + str(residues.resname[r]), + int(residues.resseq[r]), + str(residues.chain[r]), + str(residues.icode[r]), + ), + ) + out = [] + for residue in order: + parts: Dict[Tuple[str, str], List[int]] = {} + for row in residues.atom_rows(residue): + parts.setdefault((str(resname[row]), str(altloc[row])), []).append(row) + out.append((residue, {key: parts[key] for key in sorted(parts)})) + return out - def _altloc_residues(self) -> List[Tuple[tuple, List[str], Dict[str, List[int]]]]: - """Residues with several altlocs: ``(key, sorted altlocs, rows per altloc)``. + @staticmethod + def _conformers(parts: Dict[Tuple[str, str], List[int]]) -> Dict[str, List[int]]: + """A residue's atom rows per altloc label, or ``{}`` below two labels.""" + rows: Dict[str, List[int]] = {} + for (_, label), part in parts.items(): + if label != " ": + rows.setdefault(label, []).extend(part) + if len(rows) < 2: + return {} + return {label: sorted(rows[label]) for label in sorted(rows)} - Keys are ``(resname, resseq, chain)``, sorted; blank-altloc atoms are not part - of any conformer. + def altloc_residues(self) -> List[Tuple[int, List[str], Dict[str, List[int]]]]: + """The residues that carry more than one conformer. + + Occupancy grouping, :attr:`altloc_pairs` and ``Model.strip_altlocs`` all take + a residue's conformers from here. + + Returns + ------- + list of tuple + ``(residue, labels, rows)`` per topology residue with at least two altloc + labels, sorted by residue name, then ``resseq``, ``chain`` and ``icode``: + the residue index in :attr:`topology`, its sorted labels, and each label's + atom rows in ascending order. A conformer includes every residue name it + carries; blank-altloc atoms belong to no conformer. """ - altloc = self.topology.atoms.altloc out = [] - for key, rows in self._residue_groups(with_altloc=False).items(): - by_altloc: Dict[str, List[int]] = {} - for row in rows: - if altloc[row] != " ": - by_altloc.setdefault(str(altloc[row]), []).append(row) - if len(by_altloc) > 1: - labels = sorted(by_altloc) - out.append((key, labels, {a: by_altloc[a] for a in labels})) + for residue, parts in self._residue_parts(): + rows = self._conformers(parts) + if rows: + out.append((residue, list(rows), rows)) return out - def occupancy_groups(self, initial_occ): - """``(sharing_groups, altloc_groups, refinable_mask)`` for an + def occupancy_groups( + self, initial_occ: torch.Tensor + ) -> Tuple[torch.Tensor, List[tuple], torch.Tensor]: + """Sharing groups, altloc groups and refinable mask for an :class:`~torchref.model.parameter_wrappers.OccupancyTensor` over these atoms. - Altloc conformations share one collapsed index each; other residues share - one only when their occupancies agree to within 0.01, and an occupancy is - refinable only if it differs from 1.0 by more than that same deadband. + Every conformer of a residue with several altlocs is one group, whatever its + atoms' occupancies, so the sum-to-1 normalization over a residue's conformers + acts on whole conformers. Every other part of a residue -- its blank-altloc + atoms, or, in a residue with at most one altloc label, its atoms split by + residue name and altloc -- is one group when its occupancies agree to within + 0.01 and one group per atom otherwise. No group spans two residues; a starting + occupancy changes only where the atoms of one group disagree (a conformer's + atoms, or a part's within the deadband), which collapse to one shared value. + + Parameters + ---------- + initial_occ : torch.Tensor + Occupancies, shape ``(n_atoms,)``. + + Returns + ------- + sharing_groups : torch.Tensor + Group index per atom, shape ``(n_atoms,)``, contiguous from 0: the + conformers first, then the other parts, both in the residue order of + :meth:`altloc_residues`. + altloc_groups : list of tuple + Per residue with several conformers, the atom rows of each conformer. + refinable_mask : torch.Tensor + Boolean, shape ``(n_atoms,)``: occupancy (a shared group's mean) differs + from 1.0 by more than 0.01. + + Raises + ------ + ValueError + If ``initial_occ`` does not hold one value per atom. """ n_atoms = len(initial_occ) + if n_atoms != self.n_atoms: + raise ValueError( + f"initial_occ has {n_atoms} values for a context of {self.n_atoms} atoms" + ) + sharing_groups = torch.full((n_atoms,), -1, dtype=get_int_dtype()) + refinable_mask = (initial_occ - 1.0).abs() > 0.01 altloc_groups = [] - refinable_mask = torch.zeros(n_atoms, dtype=torch.bool) - - sharing_groups_tensor = torch.arange(n_atoms, dtype=get_int_dtype()) - collapsed_idx = 0 - - # First pass: altlocs. ALL atoms of one conformation must share a collapsed - # index whatever their individual occupancies, or the sum-to-1 - # normalization in OccupancyTensor.forward() acts on the wrong group. - altloc_residues = set() - for key, labels, rows_by_altloc in self._altloc_residues(): - altloc_residues.add(key) - conformation_atom_lists = [] - for label in labels: - indices = rows_by_altloc[label] - sharing_groups_tensor[indices] = collapsed_idx - for idx in indices: - if abs(initial_occ[idx].item() - 1.0) > 0.01: - refinable_mask[idx] = True - conformation_atom_lists.append(indices) - collapsed_idx += 1 - altloc_groups.append(tuple(conformation_atom_lists)) - - # Second pass: non-altloc residues, sharing by occupancy similarity. - for key, indices in self._residue_groups(with_altloc=True).items(): - if key[:3] in altloc_residues: - continue - - residue_occs = initial_occ[indices] - - occ_min = residue_occs.min().item() - occ_max = residue_occs.max().item() - occ_mean = residue_occs.mean().item() - - if (occ_max - occ_min) <= 0.01: - sharing_groups_tensor[indices] = collapsed_idx - collapsed_idx += 1 - - if abs(occ_mean - 1.0) > 0.01: - for idx in indices: - refinable_mask[idx] = True - else: - # Occupancies disagree within the residue: keep atoms independent. - for idx in indices: - if abs(initial_occ[idx].item() - 1.0) > 0.01: - refinable_mask[idx] = True - - # Compact to contiguous indices 0..n_collapsed-1. - unique_indices = torch.unique(sharing_groups_tensor, sorted=True) - index_map = torch.zeros(n_atoms, dtype=get_int_dtype()) - for new_idx, old_idx in enumerate(unique_indices): - mask = sharing_groups_tensor == old_idx - sharing_groups_tensor[mask] = new_idx + others = [] + for _, parts in self._residue_parts(): + conformers = self._conformers(parts) + if conformers: + altloc_groups.append(tuple(conformers.values())) + others.extend( + part + for (_, label), part in parts.items() + if not conformers or label == " " + ) - n_collapsed = len(unique_indices) + n_groups = 0 + for conformers in altloc_groups: + for rows in conformers: + sharing_groups[rows] = n_groups + n_groups += 1 + for rows in others: + occ = initial_occ[rows] + if occ.max().item() - occ.min().item() <= 0.01: + sharing_groups[rows] = n_groups + n_groups += 1 + refinable_mask[rows] = abs(occ.mean().item() - 1.0) > 0.01 + else: + sharing_groups[rows] = torch.arange( + n_groups, n_groups + len(rows), dtype=get_int_dtype() + ) + n_groups += len(rows) if self.verbose > 1: - n_groups = n_collapsed - n_independent = n_atoms - n_collapsed - n_refinable = refinable_mask.sum().item() - n_altloc_groups = len(altloc_groups) - print("\nOccupancy Setup:") print(f" Total atoms: {n_atoms}") - print(f" Collapsed indices: {n_collapsed}") - print(f" Alternative conformation groups: {n_altloc_groups}") - print(f" Refinable atoms: {n_refinable}") - print(f" Compression ratio: {n_atoms / n_collapsed:.2f}x") + print(f" Collapsed indices: {n_groups}") + print(f" Alternative conformation groups: {len(altloc_groups)}") + print(f" Refinable atoms: {refinable_mask.sum().item()}") + print(f" Compression ratio: {n_atoms / max(n_groups, 1):.2f}x") - return sharing_groups_tensor, altloc_groups, refinable_mask + return sharing_groups, altloc_groups, refinable_mask def register_altlocs(self) -> None: """Rebuild :attr:`altloc_pairs` from the topology's altlocs. @@ -631,7 +658,7 @@ def register_altlocs(self) -> None: torch.tensor(rows_by_altloc[label], dtype=get_int_dtype()) for label in labels ) - for _, labels, rows_by_altloc in self._altloc_residues() + for _, labels, rows_by_altloc in self.altloc_residues() ] @property diff --git a/torchref/model/model.py b/torchref/model/model.py index 0ed0da4f..5c552dc3 100644 --- a/torchref/model/model.py +++ b/torchref/model/model.py @@ -549,9 +549,10 @@ def _install_parameters( values : AtomValues Starting values, row-aligned with ``ctx.topology``. state : dict, optional - A state dict about to be loaded. Its saved refinable masks, riding frames - and node-field ADP layout fix the wrappers' shapes, and the default masks - are **not** applied; ``load_state_dict`` supplies the values afterwards. + A state dict about to be loaded. Its saved refinable masks, occupancy + groups, riding frames and node-field ADP layout fix the wrappers' shapes, + and the default masks are **not** applied; ``load_state_dict`` supplies + the values afterwards. xyz : MixedTensor, optional Coordinate wrapper to install as is, instead of building one from ``values``. @@ -573,25 +574,7 @@ def _install_parameters( self.u = self._restore_adp_slot( "u", state, values, dtype, self.xyz, self.device ) - - # Residue-level sharing plus altloc sum-to-1 groups. - initial_occ = torch.tensor(values.occupancy, dtype=dtype) - sharing_groups, altloc_groups, refinable_mask = self.ctx.occupancy_groups( - initial_occ - ) - saved_occ_mask = state.get("occupancy.refinable_mask") - if saved_occ_mask is not None: - # Saved in group space; expanded back over atoms. - refinable_mask = saved_occ_mask.to(sharing_groups.device)[sharing_groups] - self.occupancy = OccupancyTensor( - initial_values=initial_occ, - sharing_groups=sharing_groups, - altloc_groups=altloc_groups, - refinable_mask=refinable_mask, - dtype=dtype, - device=self.device, - name="occupancy", - ) + self.occupancy = self._build_occupancy(values, state) if restoring: # Placeholders: the saved masks arrive with load_state_dict, and applying @@ -617,6 +600,36 @@ def _install_parameters( ) self._repoint_coordinate_accessors() + def _build_occupancy(self, values: AtomValues, state: dict) -> OccupancyTensor: + """The occupancy wrapper over ``values``, grouped as ``state`` saved it. + + A saved grouping is taken as saved, never re-derived from the saved + occupancies: which atoms share a group depends on the values (the 0.01 + deadband of :meth:`ModelContext.occupancy_groups`), and refinement moves them, + so a re-derivation can disagree with the group-space parameters being loaded. + Without one -- a fresh load -- the context groups the atoms. + """ + initial = torch.tensor(values.occupancy, dtype=self.dtype_float) + settings = { + "dtype": self.dtype_float, + "device": self.device, + "name": "occupancy", + } + if state.get("occupancy.expansion_mask") is not None: + return OccupancyTensor.from_saved_groups( + initial, state, prefix="occupancy.", **settings + ) + sharing_groups, altloc_groups, refinable_mask = self.ctx.occupancy_groups( + initial + ) + return OccupancyTensor( + initial_values=initial, + sharing_groups=sharing_groups, + altloc_groups=altloc_groups, + refinable_mask=refinable_mask, + **settings, + ) + def _build_xyz(self, values: AtomValues, state: dict): """The coordinate wrapper over ``values``, riding if ``state`` saved one. @@ -1773,24 +1786,19 @@ def strip_altlocs(self) -> "Model": from torchref.topology import Topology topology = self.ctx.topology - altloc = topology.atoms.altloc occupancy = self.occupancy().detach().cpu().numpy() keep = np.ones(self.n_atoms, dtype=bool) resnames = topology.columns()["resname"] - for residue in range(topology.n_residues): - rows = np.arange( - int(topology.residues.atom_start[residue]), - int(topology.residues.atom_end[residue]), - ) - labels = sorted(set(altloc[rows].tolist()) - {" "}) - if len(labels) < 2: - continue - means = [occupancy[rows[altloc[rows] == label]].mean() for label in labels] + for residue, labels, conformers in self.ctx.altloc_residues(): + means = [occupancy[conformers[label]].mean() for label in labels] best = labels[int(np.argmax(means))] # Shared atoms belong to the retained chemical conformer in a model # with no altlocs, even when their deposited name was the other type. - resnames[rows] = resnames[rows[altloc[rows] == best][0]] - keep[rows[(altloc[rows] != " ") & (altloc[rows] != best)]] = False + rows = list(topology.residues.atom_rows(residue)) + resnames[rows] = resnames[conformers[best][0]] + for label in labels: + if label != best: + keep[conformers[label]] = False rows = np.nonzero(keep)[0] columns = {key: value[rows] for key, value in topology.columns().items()} @@ -1915,7 +1923,9 @@ def load_state(self, path: str, strict: bool = True, device=None): loaded = type(self).create_from_state_dict( state_dict, device=target_device, verbose=self.ctx.verbose ) - # Adopt the fully-built model's state wholesale. + # Replace rather than merge: an empty model's ``None`` wrapper placeholders are + # plain attributes, and left in place they would shadow the restored modules. + self.__dict__.clear() self.__dict__.update(loaded.__dict__) if self.ctx.verbose > 0: print(f"Loaded model state from {path}") diff --git a/torchref/model/parameter_wrappers.py b/torchref/model/parameter_wrappers.py index 288919c1..ca61205d 100644 --- a/torchref/model/parameter_wrappers.py +++ b/torchref/model/parameter_wrappers.py @@ -9,7 +9,7 @@ """ import warnings -from typing import Optional, Union +from typing import Iterable, List, Mapping, Optional, Union import torch from torch import nn @@ -1989,69 +1989,78 @@ def update_refinable_mask( self._build_index_cache() - @staticmethod - def from_residue_groups( + @classmethod + def from_saved_groups( + cls, initial_values: torch.Tensor, - pdb_dataframe, - refinable_mask: Optional[torch.Tensor] = None, + state: Mapping[str, torch.Tensor], + prefix: str = "", **kwargs, ) -> "OccupancyTensor": - """ - Create an OccupancyTensor where all atoms in each residue share occupancy. + """An OccupancyTensor grouped exactly as the one ``state`` was saved from. - Residues are grouped by ``(resname, resseq, chainid, altloc)``. + The sharing groups are the saved ``expansion_mask``, the altloc groups the + saved ``linked_occ_`` buffers and the refinable groups the saved + ``refinable_mask``, so ``load_state_dict`` then finds every buffer and the + refinable parameters at their saved shapes. Nothing is re-derived from the + occupancies, which decide the grouping of a fresh load and which refinement + is free to move. Parameters ---------- initial_values : torch.Tensor - Initial occupancy values for all atoms. - pdb_dataframe : pandas.DataFrame - DataFrame with PDB data (must have 'resname', 'resseq', 'chainid'). - refinable_mask : torch.Tensor, optional - Mask for refinable atoms. + Occupancies for all atoms, shape ``(n_atoms,)``. Placeholders: the saved + values arrive with ``load_state_dict``. + state : mapping + State dict holding ``expansion_mask`` and, where saved, + ``linked_occ_`` and ``refinable_mask``. + prefix : str, default "" + This wrapper's key prefix in ``state``, e.g. ``"occupancy."``. **kwargs - Additional arguments passed to OccupancyTensor constructor. + Passed to the constructor: ``dtype``, ``device``, ``name``, ... Returns ------- OccupancyTensor - OccupancyTensor with residue-based sharing groups. """ - # Group atoms by residue - grouped = pdb_dataframe.groupby(["resname", "resseq", "chainid", "altloc"]) - - n_atoms = len(initial_values) - sharing_groups_tensor = torch.arange(n_atoms, dtype=get_int_dtype()) - # Singletons keep their arange ids (0..n_atoms-1); start multi-atom - # group ids past that range so a group id can never collide with a - # singleton's leftover arange id (the torch.unique compaction below - # would otherwise silently merge them into one sharing group). - collapsed_idx = n_atoms - - for (resname, resseq, chainid, altloc), group in grouped: - indices = group["index"].tolist() - if len(indices) > 1: # Only create group if more than one atom - sharing_groups_tensor[indices] = collapsed_idx - collapsed_idx += 1 - - # Compact the indices - unique_indices = torch.unique(sharing_groups_tensor, sorted=True) - for new_idx, old_idx in enumerate(unique_indices): - mask = sharing_groups_tensor == old_idx - sharing_groups_tensor[mask] = new_idx - - return OccupancyTensor( + expansion_mask = state[prefix + "expansion_mask"] + linked = [ + value + for key, value in state.items() + if key.startswith(prefix + "linked_occ_") + ] + saved_mask = state.get(prefix + "refinable_mask") + refinable_mask = ( + None + if saved_mask is None + else saved_mask.to(expansion_mask.device)[expansion_mask] + ) + return cls( initial_values=initial_values, - sharing_groups=sharing_groups_tensor, + sharing_groups=expansion_mask, + altloc_groups=cls._linked_atoms(expansion_mask, linked), refinable_mask=refinable_mask, - name="occupancy", **kwargs, ) + @staticmethod + def _linked_atoms( + expansion_mask: torch.Tensor, linked: Iterable[torch.Tensor] + ) -> List[tuple]: + """Altloc groups in atom space, from ``linked_occ_`` rows of group indices.""" + return [ + tuple( + (expansion_mask == group).nonzero(as_tuple=True)[0].tolist() + for group in row + ) + for links in linked + for row in links.tolist() + ] + def copy(self) -> "OccupancyTensor": """ - Deep-copy, rebuilding the sharing groups, altloc groups and collapsed - storage from the current occupancies. + Deep-copy, keeping the sharing and altloc groups and rebuilding the + collapsed storage from the current occupancies. Returns ------- @@ -2061,26 +2070,13 @@ def copy(self) -> "OccupancyTensor": current_occ = self.forward().detach() full_refinable_mask = self._expand_values(self.refinable_mask.float()).bool() - - # Rebuild the altloc groups from the linked_occ buffers. - altloc_groups = [] - if hasattr(self, "linked_occ_sizes"): - for n_conf in self.linked_occ_sizes: - linked_indices = getattr( - self, f"linked_occ_{n_conf}" - ) # shape (N_groups, n_conf) - - for group_collapsed_indices in linked_indices: - conf_atom_lists = [] - for collapsed_idx in group_collapsed_indices: - atom_indices = ( - (self.expansion_mask == collapsed_idx) - .nonzero(as_tuple=False) - .squeeze(-1) - ) - conf_atom_lists.append(atom_indices.tolist()) - - altloc_groups.append(tuple(conf_atom_lists)) + altloc_groups = self._linked_atoms( + self.expansion_mask, + [ + getattr(self, f"linked_occ_{n_conf}") + for n_conf in getattr(self, "linked_occ_sizes", []) + ], + ) new_tensor = OccupancyTensor( initial_values=current_occ, From 0a39265f93beb9fc5fdc5a6b4096f36169910816 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 21:56:42 +0000 Subject: [PATCH 244/250] Measure torsions with the IUPAC sign, in one shared dihedral Restraints.torsions, base/targets/_common.torsions_from_xyz and the topology builders' NumPy _torsion_angle_np each carried their own copy of the dihedral formula, and all three returned atan2((n1 x b2_hat).n2, n1.n2): the negative of the IUPAC dihedral that gemmi.calculate_dihedral returns and the CCP4/AceDRG monomer-library references are written in. Every torsion whose reference is not symmetric under negation for its period was restrained toward its mirror image. On 3A5V the NAG/MAN/BMA torsions scored rms z 26-31 at the deposited geometry (1.5-2.6 with the correct sign), 6G9X's BNG 18.2 (8.0), and nucleotide sugar rings flip the same way. Both Ramachandran kernels negated the value to compensate. Amino-acid references are all sign-symmetric, so proteins were unaffected. torsions_from_xyz is now the only eager dihedral and returns the IUPAC sign. Restraints.torsions delegates to it, and the builders measure the omegas that classify cis/trans proline with it, in one call after the pairing loop. The Triton dihedral_and_grad helper uses the same formula with the matching Blondel-Karplus forces, F1 = -(|b2|/|n1|^2) n1 and F4 = (|b2|/|n2|^2) n2, from which F2 and F3 follow. The compensating negation is gone from the eager and Triton Ramachandran kernels; the Ramachandran loss on 1DAW is bitwise unchanged and its gradient agrees to float32 rounding. Tests pin the dihedral against gemmi on random and deposited atoms, the 3A5V sugar torsions against their references, the Ramachandran NLL against a gemmi host reference, and the cis-proline surface choice. A CUDA test compares Triton with eager on 3A5V, where a sign mismatch is visible; the 1DAW sweep cannot see one in the torsion target. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- docs/user_guide/restraints.rst | 8 +- .../test_triton_vs_eager_targets.py | 55 +++++++++++ tests/unit/base/test_target_values.py | 96 ++++++++++++++++++- tests/unit/monomer/test_restraints.py | 42 ++++---- .../unit/topology/test_restraint_torsions.py | 91 ++++++++++++++++++ torchref/base/targets/_common.py | 34 ++++++- torchref/base/targets/ramachandran.py | 8 +- torchref/base/targets/triton/_dihedral.py | 43 ++++----- torchref/base/targets/triton/ramachandran.py | 22 ++--- torchref/topology/builders.py | 36 +++---- torchref/topology/restraints.py | 46 ++------- 11 files changed, 347 insertions(+), 134 deletions(-) create mode 100644 tests/unit/topology/test_restraint_torsions.py diff --git a/docs/user_guide/restraints.rst b/docs/user_guide/restraints.rst index 9af51b91..f73d129a 100644 --- a/docs/user_guide/restraints.rst +++ b/docs/user_guide/restraints.rst @@ -106,9 +106,11 @@ in the deviation from ideal: + \log \sigma_i + \tfrac{1}{2}\log 2\pi \right] with :math:`q` the interatomic distance or the bond angle (angles in radians, -sigmas converted from the CIF's degrees). Torsions use a periodic von Mises NLL, -planarity restrains a group's out-of-plane deviations, chirality preserves -stereochemistry, and the non-bonded term is a steep PROLSQ-style repulsion +sigmas converted from the CIF's degrees). Torsions are measured with the IUPAC sign +the monomer library uses (the same as ``gemmi.calculate_dihedral``) and scored with +a periodic von Mises NLL, planarity restrains a group's out-of-plane deviations, +chirality preserves stereochemistry, and the non-bonded term is a steep PROLSQ-style +repulsion (:math:`E \sim \text{violation}^4`) applied to symmetry mates as well as to the asymmetric unit. diff --git a/tests/integration/test_triton_vs_eager_targets.py b/tests/integration/test_triton_vs_eager_targets.py index 2f5e3a82..72cbc869 100644 --- a/tests/integration/test_triton_vs_eager_targets.py +++ b/tests/integration/test_triton_vs_eager_targets.py @@ -262,6 +262,61 @@ def test_triton_matches_eager_per_target(target_name, gpu_refinement, gpu_state) _assert_close(target_name, eager, triton, atol, rtol) +@pytest.fixture(scope="module") +def gpu_glycoprotein(pdb_dir): + """3A5V on CUDA; its NAG/MAN/BMA torsion references are sign-sensitive.""" + from torchref.model.model import Model + + pdb = pdb_dir / "3A5V.pdb" + if not pdb.exists(): + pytest.skip("3A5V fixture not present") + model = Model(verbose=0, device=torch.device("cuda")) + model.load_pdb(str(pdb)) + return model + + +@pytest.mark.cuda +@pytest.mark.integration +@pytest.mark.parametrize("target_name", ["geometry/torsion", "geometry/ramachandran"]) +def test_triton_matches_eager_where_the_dihedral_sign_matters( + target_name, gpu_glycoprotein +): + """The Triton dihedral carries the eager sign, value and forces alike. + + The 1DAW sweep above cannot see a flipped Triton dihedral in the torsion target: + every amino-acid reference is symmetric under negation for its period, and + ``sin(-d) * (-F) = sin(d) * F`` leaves the gradient unchanged too. 3A5V's sugar + torsions are not symmetric, and the Ramachandran surfaces are not symmetric under + (phi, psi) -> (-phi, -psi), so both targets here differ if the signs disagree. + """ + from torchref.refinement.targets import RamachandranTarget, TorsionTarget + from torchref.utils import use_portable + + model = gpu_glycoprotein + target = { + "geometry/torsion": TorsionTarget, + "geometry/ramachandran": RamachandranTarget, + }[target_name](model) + + def loss_and_grads(): + model.zero_grad(set_to_none=True) + loss = target() + loss.backward() + grads = { + name: p.grad.detach().clone() + for name, p in model.named_parameters() + if p.grad is not None + } + return float(loss.detach().item()), grads + + with use_portable(): + eager = loss_and_grads() + triton = loss_and_grads() + + assert eager[1], "no gradient reached the model parameters" + _assert_close(target_name, eager, triton, *_tol_for(target_name)) + + @pytest.mark.cuda @pytest.mark.integration @pytest.mark.parametrize("target_mode", _triton_xray_modes()) diff --git a/tests/unit/base/test_target_values.py b/tests/unit/base/test_target_values.py index bf27f0a6..22dc5799 100644 --- a/tests/unit/base/test_target_values.py +++ b/tests/unit/base/test_target_values.py @@ -6,12 +6,13 @@ import pytest import torch -from torchref.base.targets._common import EPS +from torchref.base.targets._common import EPS, torsions_from_xyz from torchref.base.targets.adp import adp_simu_math from torchref.base.targets.angle import angle_math from torchref.base.targets.bond import bond_math from torchref.base.targets.chiral import chiral_math from torchref.base.targets.planarity import planarity_math +from torchref.base.targets.ramachandran import ramachandran_math from torchref.base.targets.xray_ls import ls_xray_loss_math from torchref.config import get_default_device, get_float_dtype, get_int_dtype @@ -19,19 +20,45 @@ @pytest.fixture(scope="module") -def deposited_atoms(sample_cif_file): - """Return detached Cartesian coordinates (Å) and isotropic B-factors (Ų).""" +def deposited_model(sample_cif_file): + """The sample structure (1DAW) as deposited.""" from torchref.model import Model model = Model(verbose=0) model.load_cif(str(sample_cif_file)) - return model.xyz().detach().clone(), model.adp().detach().clone() + return model + + +@pytest.fixture(scope="module") +def deposited_atoms(deposited_model): + """Return detached Cartesian coordinates (Å) and isotropic B-factors (Ų).""" + return ( + deposited_model.xyz().detach().clone(), + deposited_model.adp().detach().clone(), + ) def _indices(rows, device): return torch.tensor(rows, dtype=get_int_dtype(), device=device) +def _gemmi_dihedrals(xyz: torch.Tensor, rows) -> np.ndarray: + """``gemmi.calculate_dihedral`` in degrees for each atom quadruple in ``rows``.""" + import gemmi + + host = xyz.detach().cpu().double().numpy() + return np.degrees( + [ + gemmi.calculate_dihedral(*(gemmi.Position(*host[i]) for i in row)) + for row in rows + ] + ) + + +def _wrapped(degrees: np.ndarray) -> np.ndarray: + return (degrees + 180.0) % 360.0 - 180.0 + + def _gaussian_sum(residual, sigma): return np.sum( 0.5 * (residual / sigma) ** 2 + np.log(sigma) + 0.5 * math.log(2 * math.pi) @@ -71,6 +98,67 @@ def test_angle_value(deposited_atoms) -> None: ) +def test_dihedral_sign_is_iupac() -> None: + """Viewed along B→C, a far bond turned clockwise from the near bond is positive.""" + c, s = math.cos(math.radians(60.0)), math.sin(math.radians(60.0)) + xyz = torch.tensor( + [[1.0, 0, 0], [0, 0, 0], [0, 0, 1.0], [c, s, 1.0], [c, -s, 1.0]], + dtype=get_float_dtype(), + device=get_default_device(), + ) + # Looking along +z, +x -> (c, s) is a clockwise turn. Reversing the atom order + # leaves a dihedral unchanged; the mirror image negates it. + idx = _indices([[0, 1, 2, 3], [3, 2, 1, 0], [0, 1, 2, 4]], xyz.device) + torch.testing.assert_close( + torsions_from_xyz(xyz, idx), xyz.new_tensor([60.0, 60.0, -60.0]) + ) + + +def test_dihedral_matches_gemmi(deposited_atoms) -> None: + """The eager dihedral is gemmi's, on deposited atoms and on random quadruples.""" + xyz, _ = deposited_atoms + generator = torch.Generator().manual_seed(0) + scattered = (3.0 * torch.randn(400, 3, generator=generator)).to(xyz) + for points in (xyz[:400], scattered): + rows = [[i, i + 1, i + 2, i + 3] for i in range(len(points) - 3)] + ours = torsions_from_xyz(points, _indices(rows, points.device)) + delta = _wrapped(ours.cpu().double().numpy() - _gemmi_dihedrals(points, rows)) + assert np.abs(delta).max() < 1e-3 + + +def test_ramachandran_reads_surfaces_at_iupac_phi_psi(deposited_model) -> None: + """The Ramachandran NLL is the surface interpolated at gemmi's phi and psi. + + The surfaces are tabulated in the IUPAC convention and are not symmetric under + (phi, psi) -> (-phi, -psi): read at the mirror image, 1DAW scores several times + higher, which this comparison would not survive. + """ + restraints = deposited_model.restraints + xyz = deposited_model.xyz().detach() + phi_idx, psi_idx = restraints._rama_phi_indices, restraints._rama_psi_indices + surfaces, kind = restraints._rama_surfaces, restraints._rama_surface_type + + phi = (_gemmi_dihedrals(xyz, phi_idx.tolist()) + 180.0) % 360.0 + psi = (_gemmi_dihedrals(xyz, psi_idx.tolist()) + 180.0) % 360.0 + grid = surfaces.cpu().double().numpy()[kind.cpu().numpy()] + i0, j0 = np.floor(phi).astype(int) % 360, np.floor(psi).astype(int) % 360 + i1, j1 = (i0 + 1) % 360, (j0 + 1) % 360 + u, v = phi - np.floor(phi), psi - np.floor(psi) + n = np.arange(len(phi)) + expected = np.sum( + (1 - u) * (1 - v) * grid[n, i0, j0] + + (1 - u) * v * grid[n, i0, j1] + + u * (1 - v) * grid[n, i1, j0] + + u * v * grid[n, i1, j1] + ) + torch.testing.assert_close( + ramachandran_math(xyz, phi_idx, psi_idx, surfaces, kind), + xyz.new_tensor(expected), + rtol=1e-5, + atol=1e-3, + ) + + def test_chiral_value(deposited_atoms) -> None: """The signed scalar triple product, without a 1/6 factor, sets chirality.""" xyz, _ = deposited_atoms diff --git a/tests/unit/monomer/test_restraints.py b/tests/unit/monomer/test_restraints.py index 2e558dbb..8b399373 100644 --- a/tests/unit/monomer/test_restraints.py +++ b/tests/unit/monomer/test_restraints.py @@ -128,32 +128,22 @@ class TestTorsionRestraintCalculations: @pytest.mark.unit def test_torsion_calculation(self): - """Test torsion angle calculation between four atoms.""" - # Create atoms with known torsion - coords = torch.tensor([ - [0.0, 0.0, 0.0], - [1.5, 0.0, 0.0], - [2.0, 1.5, 0.0], - [3.5, 1.5, 0.5] - ], dtype=torch.float32) - - # Vectors along bonds - b1 = coords[1] - coords[0] - b2 = coords[2] - coords[1] - b3 = coords[3] - coords[2] - - # Normal vectors to planes - n1 = torch.cross(b1, b2) - n2 = torch.cross(b2, b3) - - # Torsion angle - m1 = torch.cross(n1, b2 / torch.norm(b2)) - x = torch.dot(n1, n2) - y = torch.dot(m1, n2) - torsion = torch.atan2(y, x) * 180 / torch.pi - - assert torch.isfinite(torsion) - assert -180 <= torsion <= 180 + """``Restraints.torsions`` gives gemmi's (IUPAC) dihedral for four atoms.""" + import gemmi + + from torchref.config import get_float_dtype, get_int_dtype + from torchref.topology.restraints import Restraints + + rows = [[0.0, 0.0, 0.0], [1.5, 0.0, 0.0], [2.0, 1.5, 0.0], [3.5, 1.5, 0.5]] + coords = torch.tensor(rows, dtype=get_float_dtype()) + idx = torch.tensor([[0, 1, 2, 3]], dtype=get_int_dtype()) + expected = np.degrees( + gemmi.calculate_dihedral(*(gemmi.Position(*row) for row in rows)) + ) + + torsion = Restraints().torsions(idx, coords) + + torch.testing.assert_close(torsion, coords.new_tensor([expected])) class TestRestraintDeviceHandling: diff --git a/tests/unit/topology/test_restraint_torsions.py b/tests/unit/topology/test_restraint_torsions.py new file mode 100644 index 00000000..72e8b54a --- /dev/null +++ b/tests/unit/topology/test_restraint_torsions.py @@ -0,0 +1,91 @@ +"""Torsion restraints measure the IUPAC dihedral the monomer-library references use. + +Every amino-acid torsion reference is symmetric under negation for its period, so a +protein cannot tell the two signs apart. Sugar rings and their substituents can: 3A5V +carries NAG, MAN and BMA, whose AceDRG references are period-1 or otherwise +sign-sensitive, and under the opposite sign they are restrained toward the mirror image. +""" + +import numpy as np +import pytest +import torch + +from torchref.model.model import Model +from torchref.topology.ramachandran import TYPE_CIS_PROLINE, TYPE_TRANS_PROLINE + +pytestmark = pytest.mark.unit + +_SUGARS = ("NAG", "MAN", "BMA") + + +def _deposited(path): + model = Model(verbose=0) + model.load_pdb(str(path)) + return model.xyz().detach(), model.restraints + + +@pytest.fixture(scope="module") +def glycoprotein(pdb_dir): + """3A5V coordinates (Å) and restraints, as deposited.""" + return _deposited(pdb_dir / "3A5V.pdb") + + +def _gemmi_dihedrals(xyz: torch.Tensor, rows) -> np.ndarray: + import gemmi + + host = xyz.cpu().double().numpy() + return np.degrees( + [ + gemmi.calculate_dihedral(*(gemmi.Position(*host[i]) for i in row)) + for row in rows + ] + ) + + +def _rms(values: torch.Tensor) -> float: + return float(values.square().mean().sqrt()) + + +def test_restraint_torsions_match_gemmi(glycoprotein): + """``Restraints.torsions`` equals ``gemmi.calculate_dihedral`` on every restraint.""" + xyz, restraints = glycoprotein + idx = restraints.restraints["torsion"]["all"]["indices"] + ours = restraints.torsions(idx, xyz).cpu().double().numpy() + delta = (ours - _gemmi_dihedrals(xyz, idx.tolist()) + 180.0) % 360.0 - 180.0 + assert np.abs(delta).max() < 1e-3 + + +def test_sugar_torsions_sit_near_their_references(glycoprotein): + """Deposited sugars score rms z of a few; their mirror images score above 15.""" + xyz, restraints = glycoprotein + group = restraints.restraints["torsion"]["all"] + deviations, sigmas_deg = restraints.torsion_deviations_with_sigmas(xyz) + sigmas = torch.deg2rad(sigmas_deg) + mirror = -restraints.torsions(group["indices"], xyz) + mirrored = restraints._wrap_torsion_periodicity( + torch.deg2rad(mirror - group["references"]), group["periods"] + ) + resname = np.asarray(restraints.topology.columns()["resname"]).astype(str) + owner = resname[group["indices"][:, 1].cpu().numpy()] + + for sugar in _SUGARS: + mask = torch.as_tensor(owner == sugar, device=deviations.device) + assert int(mask.sum()) > 0, sugar + assert _rms(deviations[mask] / sigmas[mask]) < 4.0, sugar + assert _rms(mirrored[mask] / sigmas[mask]) > 15.0, sugar + + +def test_cis_proline_takes_the_cis_surface(pdb_dir): + """A proline reads the cis or trans surface by the |omega| < 90° of its peptide.""" + xyz, restraints = _deposited(pdb_dir / "1DAW.pdb") + omega_idx = restraints.restraints["torsion"]["omega"]["indices"].tolist() + # omega CA-C-N-CA ends on the C(i-1), N, CA that open phi C(i-1)-N-CA-C. + omega_by_tail = {tuple(row[1:]): row for row in omega_idx} + kind = restraints._rama_surface_type.cpu() + proline = (kind == TYPE_CIS_PROLINE) | (kind == TYPE_TRANS_PROLINE) + phi_rows = restraints._rama_phi_indices.cpu()[proline].tolist() + + omega = _gemmi_dihedrals(xyz, [omega_by_tail[tuple(row[:3])] for row in phi_rows]) + is_cis = kind[proline].numpy() == TYPE_CIS_PROLINE + assert is_cis.any() and not is_cis.all() + np.testing.assert_array_equal(is_cis, np.abs(omega) < 90.0) diff --git a/torchref/base/targets/_common.py b/torchref/base/targets/_common.py index 4a343b6b..8594ee86 100644 --- a/torchref/base/targets/_common.py +++ b/torchref/base/targets/_common.py @@ -1,4 +1,9 @@ -"""Shared helpers for target math kernels.""" +"""Shared helpers for target math kernels. + +Also home to :func:`torsions_from_xyz`, the package's only eager dihedral, which the +topology layer uses as well; :mod:`torchref.base.targets.triton._dihedral` is its Triton +counterpart. +""" import numpy as np import torch @@ -21,14 +26,31 @@ def torsions_from_xyz(xyz: torch.Tensor, idx: torch.Tensor) -> torch.Tensor: """Compute dihedral angles in degrees from 4-atom indices. - Matches the sign convention of ``Restraints.torsions``. + The one eager dihedral in the package: ``Restraints.torsions``, the omega that + classifies cis/trans proline for the Ramachandran map and every eager geometry + target call it. + + The sign is IUPAC, the same as ``gemmi.calculate_dihedral`` and the convention the + CCP4/AceDRG monomer-library references are written in: for atoms A-B-C-D viewed + along the B→C bond, the angle is positive when the far bond C-D is rotated + clockwise from the near bond B-A. References in the opposite convention would + restrain every torsion that is not symmetric under negation for its period + (nucleotide and carbohydrate sugar rings among them) toward its mirror image. Parameters ---------- xyz : torch.Tensor - (N_atoms, 3) Cartesian coordinates. + Cartesian coordinates in Å, shape (n_atoms, 3). idx : torch.Tensor - (N, 4) atom indices defining each dihedral. + Atom indices A, B, C, D of each dihedral, shape (n_torsions, 4), integer + dtype. + + Returns + ------- + torch.Tensor + Dihedral angles in degrees in [-180, 180], shape (n_torsions,), in the dtype + of ``xyz``. A fully degenerate quadruple (coincident or collinear atoms) gives + 0 with a zero gradient rather than NaN. """ p1 = xyz[idx[:, 0]] p2 = xyz[idx[:, 1]] @@ -44,7 +66,9 @@ def torsions_from_xyz(xyz: torch.Tensor, idx: torch.Tensor) -> torch.Tensor: # Floor the |b2| divisor so collinear atoms (|b2| -> 0) give a finite # gradient instead of 0/0 = NaN. b2_norm = torch.linalg.norm(b2, dim=-1, keepdim=True).clamp_min(EPS) - m1 = torch.cross(n1, b2 / b2_norm, dim=-1) + # b2_hat x n1, not n1 x b2_hat: the operand order is what makes the sign IUPAC. + # Textbook forms built on n1 x b2_hat carry a compensating minus on the atan2. + m1 = torch.cross(b2 / b2_norm, n1, dim=-1) x = torch.sum(n1 * n2, dim=-1) y = torch.sum(m1 * n2, dim=-1) diff --git a/torchref/base/targets/ramachandran.py b/torchref/base/targets/ramachandran.py index 351996aa..7b6d1b88 100644 --- a/torchref/base/targets/ramachandran.py +++ b/torchref/base/targets/ramachandran.py @@ -13,8 +13,8 @@ def _ramachandran_math_eager( nll_surfaces: torch.Tensor, surface_type: torch.Tensor, ) -> torch.Tensor: - phi_deg = -torsions_from_xyz(xyz, phi_idx) - psi_deg = -torsions_from_xyz(xyz, psi_idx) + phi_deg = torsions_from_xyz(xyz, phi_idx) + psi_deg = torsions_from_xyz(xyz, psi_idx) phi_idx_grid = (phi_deg + 180.0) % 360.0 psi_idx_grid = (psi_deg + 180.0) % 360.0 @@ -62,7 +62,9 @@ def ramachandran_math( xyz : torch.Tensor (N_atoms, 3) Cartesian coordinates. phi_idx, psi_idx : torch.Tensor - (N, 4) atom indices for the two backbone dihedrals. + (N, 4) atom indices for the two backbone dihedrals, measured by + :func:`~torchref.base.targets._common.torsions_from_xyz` with the IUPAC sign + the surfaces are tabulated in. nll_surfaces : torch.Tensor (n_surface_types, 360, 360) precomputed NLL = -log P(φ, ψ | type). surface_type : torch.Tensor diff --git a/torchref/base/targets/triton/_dihedral.py b/torchref/base/targets/triton/_dihedral.py index 2467707c..078309fc 100644 --- a/torchref/base/targets/triton/_dihedral.py +++ b/torchref/base/targets/triton/_dihedral.py @@ -1,22 +1,22 @@ """Shared Triton helpers for dihedral-angle computation and gradients. -The forces use the sign convention of -:func:`torchref.base.targets._common.torsions_from_xyz`, which defines the -angle as ``atan2(m·n2, n1·n2)``. The canonical Bekker / OpenMM formulas are -**overall-negated** relative to this convention; the forms below (and the code) -are the sign-corrected versions, verified by finite differences against the -eager forward and bitwise via the equivalence tests. +:func:`dihedral_and_grad` is the Triton counterpart of +:func:`torchref.base.targets._common.torsions_from_xyz` and computes the angle with the +same formula, ``atan2(m1·n2, n1·n2)`` with ``m1 = (b2/|b2|) x n1``, so it carries the +same IUPAC sign as ``gemmi.calculate_dihedral`` and the monomer-library references. +The forces are the Blondel-Karplus / Bekker forms for that sign. Given four atoms p1, p2, p3, p4 with bonds b1 = p2-p1, b2 = p3-p2, b3 = p4-p3: n1 = b1 x b2, n2 = b2 x b3, b2_len = |b2| - ∂ω/∂p1 = (b2_len / |n1|²) · n1 (call F1) - ∂ω/∂p4 = -(b2_len / |n2|²) · n2 (call F4) + ∂ω/∂p1 = -(b2_len / |n1|²) · n1 (call F1) + ∂ω/∂p4 = (b2_len / |n2|²) · n2 (call F4) ∂ω/∂p2 = -((b1·b2) / |b2|² + 1) · F1 + ((b3·b2) / |b2|²) · F4 (call F2) ∂ω/∂p3 = −F1 − F2 − F4 (returned as F3) -The opposite (Bekker) signs would give the wrong-signed forces here. +The angle and the forces change sign together: F2 and F3 are linear in F1 and F4, +so negating ω means negating F1 and F4 and nothing else. """ from __future__ import annotations @@ -49,7 +49,8 @@ def dihedral_and_grad( ): """Return (omega_rad, F1.., F2.., F3.., F4..) — 1 angle + 12 gradient comps. - All inputs are SIMD lanes of float32 from a Triton block. + ``omega_rad`` is the IUPAC dihedral in radians and ``Fi = ∂omega/∂pi``. All + inputs are SIMD lanes of float32 from a Triton block. """ b1x = p2x - p1x b1y = p2y - p1y @@ -73,12 +74,12 @@ def dihedral_and_grad( c22_safe = c22 + _EPS b2_len = tl.sqrt(c22) - # angle via the same atan2 form as the eager helper - # m1 = n1 x (b2 / |b2|) + # m1 = (b2 / |b2|) x n1, the operand order of torsions_from_xyz: it is what + # makes the sign IUPAC. inv_b2 = 1.0 / tl.sqrt(c22_safe) - m1x = (n1y * b2z - n1z * b2y) * inv_b2 - m1y = (n1z * b2x - n1x * b2z) * inv_b2 - m1z = (n1x * b2y - n1y * b2x) * inv_b2 + m1x = (b2y * n1z - b2z * n1y) * inv_b2 + m1y = (b2z * n1x - b2x * n1z) * inv_b2 + m1z = (b2x * n1y - b2y * n1x) * inv_b2 y = m1x * n2x + m1y * n2y + m1z * n2z x = n1x * n2x + n1y * n2y + n1z * n2z omega = libdevice.atan2(y, x) @@ -86,16 +87,12 @@ def dihedral_and_grad( N1 = n1x * n1x + n1y * n1y + n1z * n1z N2 = n2x * n2x + n2y * n2y + n2z * n2z - # Sign convention matches torsions_from_xyz, which uses - # atan2(m·n2, n1·n2). Finite-difference verification on a known case - # (φ=+90°, p1.z perturbation → ∂φ/∂p1.z = -1) showed the canonical - # Bekker formula is overall-negated relative to this convention, so: # Floor |n1|², |n2|² so collinear b1∥b2 / b2∥b3 give finite forces. - f1c = b2_len / (N1 + _EPS) # F1 = +(b2_len / N1) · n1 + f1c = -b2_len / (N1 + _EPS) F1x = f1c * n1x F1y = f1c * n1y F1z = f1c * n1z - f4c = -b2_len / (N2 + _EPS) # F4 = -(b2_len / N2) · n2 + f4c = b2_len / (N2 + _EPS) F4x = f4c * n2x F4y = f4c * n2y F4z = f4c * n2z @@ -103,10 +100,6 @@ def dihedral_and_grad( c12 = b1x * b2x + b1y * b2y + b1z * b2z c23 = b2x * b3x + b2y * b3y + b2z * b3z - # F2 derived by numerical fit against autograd on the eager forward - # (atan2(m·n2, n1·n2)). The canonical Bekker textbook form doesn't - # match this sign convention; this one does: - # F2 = −(1 + c12/c22) · F1 + (c23/c22) · F4 a = -(c12 / c22_safe + 1.0) b = c23 / c22_safe diff --git a/torchref/base/targets/triton/ramachandran.py b/torchref/base/targets/triton/ramachandran.py index 5ccc2a5f..a991a4ed 100644 --- a/torchref/base/targets/triton/ramachandran.py +++ b/torchref/base/targets/triton/ramachandran.py @@ -7,10 +7,10 @@ d(NLL)/d(phi_frac) = (1-ψf)(v10−v00) + ψf·(v11−v01) d(NLL)/d(psi_frac) = (1-φf)(v01−v00) + φf·(v11−v10) d(phi_frac)/d(phi_deg) = 1 - d(phi_deg)/d(positions) = −(180/π) · F_dihedral + d(phi_deg)/d(positions) = (180/π) · F_dihedral -(the leading minus is because the eager target uses ``-torsions_from_xyz``; -F_dihedral comes from :mod:`_dihedral`.) Same for ψ. +(F_dihedral comes from :mod:`_dihedral`, whose angle carries the IUPAC sign the +surfaces are tabulated in.) Same for ψ. """ from __future__ import annotations @@ -62,7 +62,7 @@ def _rama_nll_fwd_kernel( _F3x, _F3y, _F3z, _F4x, _F4y, _F4z) = dihedral_and_grad( pax, pay, paz, pbx, pby, pbz, pcx, pcy, pcz, pdx, pdy, pdz, ) - phi_deg = -phi_rad * (180.0 / 3.141592653589793) + phi_deg = phi_rad * (180.0 / 3.141592653589793) a = tl.load(psi_idx_ptr + offs * 4 + 0, mask=mask, other=0) b = tl.load(psi_idx_ptr + offs * 4 + 1, mask=mask, other=0) @@ -84,7 +84,7 @@ def _rama_nll_fwd_kernel( _F3x, _F3y, _F3z, _F4x, _F4y, _F4z) = dihedral_and_grad( pax, pay, paz, pbx, pby, pbz, pcx, pcy, pcz, pdx, pdy, pdz, ) - psi_deg = -psi_rad * (180.0 / 3.141592653589793) + psi_deg = psi_rad * (180.0 / 3.141592653589793) phi_g = (phi_deg + 180.0) % 360.0 psi_g = (psi_deg + 180.0) % 360.0 @@ -150,7 +150,7 @@ def _rama_nll_bwd_kernel( pax, pay, paz, pbx, pby, pbz, pcx, pcy, pcz, pdx, pdy, pdz, ) phi_a = a; phi_b = b; phi_c = c; phi_d = d - phi_deg = -phi_rad * RAD2DEG + phi_deg = phi_rad * RAD2DEG # --- psi --- a = tl.load(psi_idx_ptr + offs * 4 + 0, mask=mask, other=0) @@ -174,7 +174,7 @@ def _rama_nll_bwd_kernel( pax, pay, paz, pbx, pby, pbz, pcx, pcy, pcz, pdx, pdy, pdz, ) psi_a = a; psi_b = b; psi_c = c; psi_d = d - psi_deg = -psi_rad * RAD2DEG + psi_deg = psi_rad * RAD2DEG # --- bilinear: gather corner values, compute fractional gradients --- phi_g = (phi_deg + 180.0) % 360.0 @@ -200,10 +200,10 @@ def _rama_nll_bwd_kernel( # phi_g = (phi_deg + 180) % 360 → dphi_g/dphi_deg = 1 a.e. # phi_frac = phi_g - floor(phi_g.detach()) → dphi_frac/dphi_g = 1 - # phi_deg = -phi_rad · RAD2DEG → dphi_deg/dphi_rad = -RAD2DEG - # So dNLL/dphi_rad = -RAD2DEG · dNLL_dphi_frac - coef_phi = grad_out * (-RAD2DEG) * dNLL_dphi_frac - coef_psi = grad_out * (-RAD2DEG) * dNLL_dpsi_frac + # phi_deg = phi_rad · RAD2DEG → dphi_deg/dphi_rad = RAD2DEG + # So dNLL/dphi_rad = RAD2DEG · dNLL_dphi_frac + coef_phi = grad_out * RAD2DEG * dNLL_dphi_frac + coef_psi = grad_out * RAD2DEG * dNLL_dpsi_frac # Scatter phi forces tl.atomic_add(dxyz_ptr + phi_a * 3 + 0, coef_phi * PF1x, mask=mask) diff --git a/torchref/topology/builders.py b/torchref/topology/builders.py index 87650db9..4f115e40 100644 --- a/torchref/topology/builders.py +++ b/torchref/topology/builders.py @@ -18,6 +18,7 @@ import pandas as pd import torch +from torchref.base.targets._common import torsions_from_xyz from torchref.config import get_float_dtype, get_int_dtype @@ -927,19 +928,6 @@ def disulfide_count(self) -> int: """Return total number of disulfide torsion restraints accumulated.""" return self._disulfide_count - @staticmethod - def _torsion_angle_np(coords: np.ndarray, i1, i2, i3, i4) -> float: - """Compute torsion angle (degrees) from coordinates for 4 atom indices.""" - p = coords[[i1, i2, i3, i4]] - b1, b2, b3 = p[1] - p[0], p[2] - p[1], p[3] - p[2] - n1, n2 = np.cross(b1, b2), np.cross(b2, b3) - n1_len, n2_len = np.linalg.norm(n1), np.linalg.norm(n2) - if n1_len < 1e-10 or n2_len < 1e-10: - return 180.0 - n1, n2 = n1 / n1_len, n2 / n2_len - m1 = np.cross(n1, b2 / np.linalg.norm(b2)) - return float(np.degrees(np.arctan2(np.dot(m1, n2), np.dot(n1, n2)))) - def build( self, residues: "PeptideResidues", @@ -963,8 +951,6 @@ def build( if not pairs: return None - coords_np = residues.xyz - # Separate accumulators for phi, psi, omega phi_data = {"indices": [], "periods": []} psi_data = {"indices": [], "periods": []} @@ -980,7 +966,7 @@ def build( # psi from pair (i, j) belongs to residue i (first residue) phi_by_residue = {} # res_idx -> atom indices psi_by_residue = {} # res_idx -> atom indices - omega_by_residue = {} # res_idx -> omega_deg (for cis/trans PRO detection) + omega_idx_by_residue = {} # res_idx -> omega atom indices (cis/trans PRO) resname_by_residue = {} # res_idx -> resname next_resname_by_residue = {} # res_idx -> next resname (for pre-PRO) @@ -1058,12 +1044,9 @@ def build( resname_by_residue[key_i] = resname_i resname_by_residue[key_next] = resname_next next_resname_by_residue[key_i] = resname_next - # Compute omega for PRO cis/trans detection + # The omega that decides PRO cis/trans, measured after the loop if omega_data["indices"]: - omega_deg = self._torsion_angle_np( - coords_np, *omega_data["indices"][-1] - ) - omega_by_residue[key_next] = omega_deg + omega_idx_by_residue[key_next] = omega_data["indices"][-1] result = {} @@ -1125,6 +1108,17 @@ def build( set(phi_by_residue.keys()) & set(psi_by_residue.keys()) ) if rama_residues: + omega_keys = [r for r in rama_residues if r in omega_idx_by_residue] + omega_by_residue = {} + if omega_keys: + omega_values = torsions_from_xyz( + torch.as_tensor(residues.xyz, dtype=get_float_dtype()), + torch.as_tensor( + [omega_idx_by_residue[r] for r in omega_keys], + dtype=get_int_dtype(), + ), + ) + omega_by_residue = dict(zip(omega_keys, omega_values.tolist())) rama_phi = [] rama_psi = [] rama_types = [] diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index e07281b1..c7ce436c 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -26,6 +26,7 @@ import torch from torch.nn import Module +from torchref.base.targets._common import torsions_from_xyz from torchref.config import get_float_dtype from torchref.topology.monomer.cif import ( find_cif_file_in_library, @@ -858,53 +859,26 @@ def cat_dict(self): if self.topology is not None and "all" not in self._entries.get("bond", {}): self._rebuild_entries() - def torsions(self, idx, xyz: torch.Tensor): - """ - Compute current torsion angle values for all torsion restraints. + def torsions(self, idx: torch.Tensor, xyz: torch.Tensor) -> torch.Tensor: + """Compute current torsion angles, IUPAC sign, in degrees. + + Delegates to :func:`torchref.base.targets._common.torsions_from_xyz`, the + package's one eager dihedral, whose sign is the convention the monomer-library + references are written in (the same as ``gemmi.calculate_dihedral``). Parameters ---------- idx : torch.Tensor - Torsion indices tensor of shape (N, 4). + Torsion atom indices of shape (n_torsions, 4), integer dtype. xyz : torch.Tensor Cartesian coordinates in Å, shape (n_atoms, 3). Returns ------- torch.Tensor - Tensor of shape (n_torsions,) with current torsion values in degrees. + Torsion angles in degrees in [-180, 180], shape (n_torsions,). """ - pos1 = xyz[idx[:, 0], :] - pos2 = xyz[idx[:, 1], :] - pos3 = xyz[idx[:, 2], :] - pos4 = xyz[idx[:, 3], :] - - # Compute torsion angles using vector math - b1 = pos2 - pos1 - b2 = pos3 - pos2 - b3 = pos4 - pos3 - - # Normalize b2 for projection - b2_norm = torch.linalg.norm(b2, dim=-1, keepdim=True) - b2_unit = b2 / b2_norm - - # Compute normals to planes - n1 = torch.cross(b1, b2, dim=-1) - n2 = torch.cross(b2, b3, dim=-1) - - # Normalize normals - n1_unit = n1 / torch.linalg.norm(n1, dim=-1, keepdim=True) - n2_unit = n2 / torch.linalg.norm(n2, dim=-1, keepdim=True) - - # Compute angle between normals - m1 = torch.cross(n1_unit, b2_unit, dim=-1) - - x = torch.sum(n1_unit * n2_unit, dim=-1) - y = torch.sum(m1 * n2_unit, dim=-1) - - torsions_rad = torch.atan2(y, x) - torsions_deg = torch.rad2deg(torsions_rad) - return torsions_deg + return torsions_from_xyz(xyz, idx) def _wrap_torsion_periodicity(self, diff_rad, periods): """Smallest angular deviation under n-fold rotational symmetry. From 231080015b223745a61be6aef291bd3e2258674b Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 21:57:16 +0000 Subject: [PATCH 245/250] Write insertion codes through one atom-record formatter The PDB writer formatted the atom-identity columns in three places: the ATOM line in write(), the ATOM line in write_multi_model() and the ANISOU line. Both ATOM copies right-justified the insertion code in a 4-wide field (column 30, not 27) and ANISOU wrote two blanks there, so every insertion code was lost on write and residues such as 52/52A merged. Columns 7-27 are now built once, by _format_atom_identity, which both the ATOM/HETATM and the ANISOU record use, and one row writer, _write_atom_records, serves write() and write_multi_model(). ANISOU values are rounded to U*10^4 instead of truncated; truncation changed 3351 of 7254 deposited 7L84 values on a plain load/write cycle. sanitize_pdb_dataframe, which Model.write_pdb and write_cif run on every output, spelled its duplicate key three times, each without icode, and renumbered per atom serial. One same-name insertion-code pair (GLY 90/90A) therefore tore every GLY in the chain into one-atom residues. The atom key and the residue key are now defined once and include icode. Residues (contiguous runs, split where an atom repeats) are renumbered whole, and only HETATM residues whose identifiers are already taken move. ATOM records are never renumbered. dataframe_to_gemmi_structure built gemmi.SeqId from a string, which lower-cases the insertion code, and its groupby dropped atoms whose chain ID is blank (NaN from the PDB reader). Converting such a PDB to mmCIF lost those atoms. It now builds SeqId(num, icode) and keeps NaN keys. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- tests/unit/io/test_atom_identity_roundtrip.py | 173 +++++++++++ .../unit/utils/test_sanitize_pdb_dataframe.py | 90 ++++++ torchref/io/cif.py | 21 +- torchref/io/pdb.py | 272 +++++++++--------- torchref/utils/utils.py | 159 +++++----- 5 files changed, 496 insertions(+), 219 deletions(-) create mode 100644 tests/unit/io/test_atom_identity_roundtrip.py create mode 100644 tests/unit/utils/test_sanitize_pdb_dataframe.py diff --git a/tests/unit/io/test_atom_identity_roundtrip.py b/tests/unit/io/test_atom_identity_roundtrip.py new file mode 100644 index 00000000..bc1c418e --- /dev/null +++ b/tests/unit/io/test_atom_identity_roundtrip.py @@ -0,0 +1,173 @@ +"""Atom identity survives every coordinate writer. + +Residues 10 and 10A differ only in their insertion code: column 27 of an ATOM, HETATM +or ANISOU record, ``pdbx_PDB_ins_code`` in mmCIF. Each test writes a deposited +structure renumbered to carry insertion codes and reads the file back with gemmi, an +independent reader, so a writer that drops, moves or changes the case of that one +character fails here. 7L84 is the base because it carries ANISOU records, which repeat +the identity columns, as well as altlocs and waters. + +Residues 10/10A share a name, so losing the code would merge them into one residue; +23/23A/23B have different names. +""" + +import gemmi +import pytest + +from torchref.io import cif, pdb +from torchref.model.model import Model + +BASE = "7L84" + +#: Deposited resseq -> the (resseq, icode) written in its place. +RENUMBER = {11: (10, "A"), 24: (23, "A"), 25: (23, "B")} + +RECORDS = ("ATOM", "HETATM", "ANISOU") + + +def _rewrite_with_insertion_codes(source, destination): + """Copy a PDB file, renumbering :data:`RENUMBER`; only columns 23-27 change.""" + out = [] + for line in source.read_text().splitlines(keepends=True): + if line.startswith(RECORDS): + resseq = int(line[22:26]) + if resseq in RENUMBER: + new_resseq, icode = RENUMBER[resseq] + line = f"{line[:22]}{new_resseq:>4d}{icode}{line[27:]}" + out.append(line) + destination.write_text("".join(out)) + + +def _residues(model): + """``(chain, seqnum, icode, resname, n_atoms)`` of every residue in a gemmi model.""" + return [ + (chain.name, res.seqid.num, res.seqid.icode, res.name, len(res)) + for chain in model + for res in chain + ] + + +def _read(path): + """Residues and the number of anisotropic atoms in the first model, per gemmi.""" + model = gemmi.read_structure(str(path))[0] + n_aniso = sum( + atom.aniso.nonzero() for chain in model for res in chain for atom in res + ) + return _residues(model), n_aniso + + +@pytest.fixture(scope="module") +def inserted(pdb_dir, tmp_path_factory): + """Path of the renumbered PDB file.""" + path = tmp_path_factory.mktemp("icode") / f"{BASE}_icode.pdb" + _rewrite_with_insertion_codes(pdb_dir / f"{BASE}.pdb", path) + return path + + +@pytest.fixture(scope="module") +def table(inserted): + """The renumbered file as the PDB reader's atom table.""" + df, _, _ = pdb.read(str(inserted))() + return df + + +@pytest.mark.unit +def test_the_rewrite_produced_insertion_codes(inserted, table): + """Guard the fixture: if the rewrite silently failed the rest proves nothing.""" + residues, n_aniso = _read(inserted) + seqids = {(num, icode) for _, num, icode, _, _ in residues} + assert {(10, " "), (10, "A"), (23, " "), (23, "A"), (23, "B")} <= seqids + assert n_aniso > 0 + assert set(table["icode"]) == {"", "A", "B"} + assert table.loc[table["icode"] == "A", "anisou_flag"].any() + + +@pytest.mark.unit +def test_pdb_write_puts_the_identity_in_columns_7_to_27(inserted, table, tmp_path): + """ATOM, HETATM and ANISOU records carry the input's columns 1-27 byte for byte.""" + out = tmp_path / "out.pdb" + pdb.write(table, str(out)) + + def identities(path): + return [ + line[:27] + for line in path.read_text().splitlines() + if line.startswith(RECORDS) + ] + + assert identities(out) == identities(inserted) + assert all( + len(line) == 80 + for line in out.read_text().splitlines() + if line.startswith(RECORDS) + ) + + +@pytest.mark.unit +def test_pdb_write_round_trips_insertion_codes(inserted, table, tmp_path): + out = tmp_path / "out.pdb" + pdb.write(table, str(out)) + + assert _read(out) == _read(inserted) + back, _, _ = pdb.read(str(out))() + assert back["icode"].tolist() == table["icode"].tolist() + + +@pytest.mark.unit +def test_write_multi_model_round_trips_insertion_codes(inserted, table, tmp_path): + out = tmp_path / "multi.pdb" + pdb.write_multi_model([table, table], str(out)) + + structure = gemmi.read_structure(str(out)) + expected, _ = _read(inserted) + assert len(structure) == 2 + assert all(_residues(model) == expected for model in structure) + + +@pytest.fixture(scope="module") +def model(inserted): + """The renumbered file loaded as a Model, every atom kept.""" + model = Model(verbose=0, hydrogens="keep") + model.load_pdb(str(inserted)) + return model + + +@pytest.mark.unit +def test_model_write_pdb_round_trips_insertion_codes(inserted, model, tmp_path): + """Through Model.write_pdb, which sanitizes the table before writing it.""" + out = tmp_path / "model.pdb" + model.write_pdb(str(out)) + + assert _read(out) == _read(inserted) + + +@pytest.mark.unit +def test_model_write_cif_round_trips_insertion_codes(inserted, model, tmp_path): + """The same through mmCIF, insertion codes in their original (upper) case.""" + out = tmp_path / "model.cif" + model.write_cif(str(out)) + + assert _read(out) == _read(inserted) + back, _, _ = cif.read_model(str(out))() + assert set(back["icode"]) == {"", "A", "B"} + + +@pytest.mark.unit +def test_blank_chain_atoms_survive_pdb_to_cif(pdb_dir, tmp_path): + """Atoms with a blank chain ID, which the PDB reader reads as NaN, are written.""" + source = tmp_path / "blank_chain.pdb" + lines = [] + for line in (pdb_dir / "1DAW.pdb").read_text().splitlines(keepends=True): + if line.startswith(RECORDS) and line[17:20] == "HOH": + line = f"{line[:21]} {line[22:]}" + lines.append(line) + source.write_text("".join(lines)) + table, _, _ = pdb.read(str(source))() + assert table["chainid"].fillna("").eq("").sum() == 285 + + out = tmp_path / "blank_chain.cif" + cif.write_model(table, str(out)) + + residues, _ = _read(out) + assert sum(n_atoms for *_, n_atoms in residues) == len(table) + assert sum(resname == "HOH" for _, _, _, resname, _ in residues) == 285 diff --git a/tests/unit/utils/test_sanitize_pdb_dataframe.py b/tests/unit/utils/test_sanitize_pdb_dataframe.py new file mode 100644 index 00000000..34d12e38 --- /dev/null +++ b/tests/unit/utils/test_sanitize_pdb_dataframe.py @@ -0,0 +1,90 @@ +"""What ``sanitize_pdb_dataframe`` renumbers before a model is written, and what not. + +``Model.write_pdb`` and ``Model.write_cif`` pass every table through it. A HETATM +residue that repeats an atom identifier ``(chainid, resseq, icode, name, altloc)`` -- +unnumbered waters, a copied ligand -- gets a new resseq as a whole. Nothing else +changes: an insertion code keeps residues 90 and 90A apart, and ATOM records are never +renumbered. Tables come from 1DAW (protein, ANP 340, MG 341-342, waters 350-634). +""" + +import pandas as pd +import pytest + +from torchref.io import pdb +from torchref.utils import sanitize_pdb_dataframe + +ATOM_KEY = ["chainid", "resseq", "icode", "name", "altloc"] + + +@pytest.fixture(scope="module") +def table(pdb_dir): + """1DAW as the PDB reader's atom table.""" + df, _, _ = pdb.read(str(pdb_dir / "1DAW.pdb"))() + return df + + +def _resseq(df): + return df["resseq"].to_numpy().tolist() + + +@pytest.mark.unit +def test_a_clean_model_is_untouched(table): + pd.testing.assert_frame_equal(sanitize_pdb_dataframe(table), table) + + +@pytest.mark.unit +def test_an_insertion_code_pair_is_untouched(table): + """GLY 90 and GLY 90A share a name and a number; the insertion code separates them.""" + inserted = table.copy() + inserted.loc[inserted["resseq"] == 91, ["resseq", "icode"]] = [90, "A"] + assert inserted.duplicated(["chainid", "resseq", "name", "altloc"]).any() + + pd.testing.assert_frame_equal(sanitize_pdb_dataframe(inserted), inserted) + + +@pytest.mark.unit +def test_a_copied_ligand_is_renumbered_whole(table): + """A second ANP 340 becomes one residue at the next free number.""" + anp = table[table["resname"] == "ANP"] + df = pd.concat([table, anp], ignore_index=True) + + out = sanitize_pdb_dataframe(df) + + assert _resseq(out.iloc[: len(table)]) == _resseq(table) + assert set(out["resseq"].iloc[len(table) :]) == {table["resseq"].max() + 1} + + +@pytest.mark.unit +def test_unnumbered_waters_become_one_residue_each(table): + df = table.copy() + water = (df["resname"] == "HOH").to_numpy() + df.loc[water, "resseq"] = 0 + + out = sanitize_pdb_dataframe(df) + + assert not out.duplicated(ATOM_KEY).any() + assert out.loc[water, "resseq"].nunique() == water.sum() + assert _resseq(out.loc[~water]) == _resseq(table.loc[~water]) + assert (df.loc[water, "resseq"] == 0).all(), "the input table was modified" + + +@pytest.mark.unit +def test_a_water_colliding_with_a_polymer_residue_moves(table): + """A water numbered 90 collides with GLY 90's O; the water moves, though it is first.""" + water = table[table["resname"] == "HOH"].iloc[[0]].assign(resseq=90) + df = pd.concat([water, table], ignore_index=True) + + out = sanitize_pdb_dataframe(df) + + assert out["resseq"].iloc[0] == table["resseq"].max() + 1 + assert _resseq(out.iloc[1:]) == _resseq(table) + + +@pytest.mark.unit +def test_atom_records_are_never_renumbered(table): + """A duplicated polymer residue is left as it is rather than torn apart.""" + df = pd.concat([table, table[table["resseq"] == 90]], ignore_index=True) + + out = sanitize_pdb_dataframe(df) + + assert _resseq(out) == _resseq(df) diff --git a/torchref/io/cif.py b/torchref/io/cif.py index d616abcb..cd068998 100644 --- a/torchref/io/cif.py +++ b/torchref/io/cif.py @@ -162,7 +162,8 @@ def dataframe_to_gemmi_structure(df, cell, spacegroup): Returns ------- gemmi.Structure - The constructed gemmi Structure object. + The constructed gemmi Structure object. A blank or NaN chain ID is named + ``A``, the same name as a real chain ``A`` if the model has one. """ import gemmi @@ -180,18 +181,22 @@ def dataframe_to_gemmi_structure(df, cell, spacegroup): model = gemmi.Model("1") - # Group by chain, then by (resseq, icode, resname) for residues - for chain_id, chain_group in df.groupby("chainid", sort=False): - chain = gemmi.Chain(str(chain_id) if chain_id and str(chain_id) != "nan" else "A") + # Group by chain, then by (resseq, icode, resname) for residues. NaN keys are + # kept: the PDB reader reads a blank chain ID as NaN, and groupby would + # otherwise drop those atoms. + for chain_id, chain_group in df.groupby("chainid", sort=False, dropna=False): + chain = gemmi.Chain( + str(chain_id) if chain_id and str(chain_id) != "nan" else "A" + ) for (resseq, icode, resname), res_group in chain_group.groupby( - ["resseq", "icode", "resname"], sort=False + ["resseq", "icode", "resname"], sort=False, dropna=False ): residue = gemmi.Residue() residue.name = str(resname).strip() - seq_str = str(int(resseq)) - icode_str = str(icode).strip() if icode and str(icode) not in ("nan", " ") else "" - residue.seqid = gemmi.SeqId(seq_str + icode_str) + icode_str = str(icode).strip() if icode and str(icode) != "nan" else "" + # Not gemmi.SeqId("52A"): parsing a string lower-cases the insertion code. + residue.seqid = gemmi.SeqId(int(resseq), icode_str or " ") # Set het flag based on ATOM/HETATM first_atom_type = res_group.iloc[0]["ATOM"] diff --git a/torchref/io/pdb.py b/torchref/io/pdb.py index c7d749f8..1f4a9b2b 100644 --- a/torchref/io/pdb.py +++ b/torchref/io/pdb.py @@ -494,6 +494,126 @@ def extract_link_records(filepath: str, verbose: int = 0) -> pd.DataFrame: return df +_ATOM_COLUMNS = ( + "ATOM", + "serial", + "name", + "altloc", + "resname", + "chainid", + "resseq", + "icode", + "x", + "y", + "z", + "occupancy", + "tempfactor", + "element", + "charge", +) + +_U_COLUMNS = ("u11", "u22", "u33", "u12", "u13", "u23") + + +def _text(value) -> str: + """``value`` as stripped text, with None, NaN and the string ``'nan'`` blank. + + The reader leaves a blank chain ID as NaN, which ``astype(str)`` downstream + turns into ``'nan'``; both mean the field is empty. + """ + if pd.isna(value): + return "" + text = str(value).strip() + return "" if text == "nan" else text + + +def _format_charge(charge) -> str: + """Formal charge for columns 79-80: blank when neutral, else ``+1`` / ``-2``.""" + charge = 0 if pd.isna(charge) else int(charge) + return f"{charge:+d}" if charge else "" + + +def _format_atom_identity(row) -> str: + """Columns 7-27 of an ATOM, HETATM or ANISOU record: which atom it describes. + + wwPDB v3.3 layout: serial 7-11, atom name 13-16, altLoc 17, resName 18-20, + chainID 22, resSeq 23-26, iCode 27. A two-character chain ID takes columns + 21-22, where gemmi reads and writes it. + + Parameters + ---------- + row : mapping + One atom; reads ``serial``, ``name``, ``element``, ``altloc``, + ``resname``, ``chainid``, ``resseq`` and ``icode``. + + Returns + ------- + str + Exactly 21 characters for in-range values. A value wider than its field + (serial > 99999, a 4-character residue name) is not truncated and shifts + every later column. + """ + name = _format_pdb_atom_name(row["name"], _text(row["element"])) + return ( + f"{int(row['serial']):>5} {name}{_text(row['altloc']):1}" + f"{_text(row['resname']):>3}{_text(row['chainid']):>2}" + f"{int(row['resseq']):>4}{_text(row['icode']):1}" + ) + + +def _format_atom_records(row, anisou: bool) -> str: + """The ATOM or HETATM record of one atom, then its ANISOU record if ``anisou``. + + Both records take columns 7-27 from :func:`_format_atom_identity`, so they + cannot disagree about the atom. ANISOU holds round(U * 10^4) with U in Ų. + + Parameters + ---------- + row : mapping + One atom with the columns :func:`write` requires, plus ``u11`` ... + ``u23`` when ``anisou`` is true. + anisou : bool + Whether to append the ANISOU record. + + Returns + ------- + str + One or two newline-terminated 80-column records. + """ + identity = _format_atom_identity(row) + element_charge = f"{_text(row['element']):>2}{_format_charge(row['charge']):>2}" + records = ( + f"{_text(row['ATOM']):<6}{identity} " + f"{row['x']:8.3f}{row['y']:8.3f}{row['z']:8.3f}" + f"{row['occupancy']:6.2f}{row['tempfactor']:6.2f}" + f"{'':10}{element_charge}\n" + ) + if anisou: + u = "".join(f"{round(float(row[c]) * 1e4):7d}" for c in _U_COLUMNS) + records += f"ANISOU{identity} {u}{'':6}{element_charge}\n" + return records + + +def _write_atom_records(handle, df: pd.DataFrame, anisou: bool) -> None: + """Write one ATOM/HETATM record per row of ``df``, in row order. + + With ``anisou``, rows whose ``anisou_flag`` is set also get an ANISOU + record. A row that cannot be formatted is skipped whole, with a printed + warning, so one bad value costs one atom rather than the file. + """ + anisou = anisou and "anisou_flag" in df.columns and bool(df["anisou_flag"].any()) + columns = list(_ATOM_COLUMNS) + if anisou: + columns += ["anisou_flag", *_U_COLUMNS] + for i, row in enumerate(df[columns].to_dict("records")): + try: + records = _format_atom_records(row, anisou and bool(row["anisou_flag"])) + except (TypeError, ValueError) as error: + print(f"Skipping atom row {i}, which cannot be formatted: {error}") + continue + handle.write(records) + + def write(df: pd.DataFrame, filepath: str, metadata=None) -> None: """ Write a DataFrame to a PDB file. @@ -501,14 +621,20 @@ def write(df: pd.DataFrame, filepath: str, metadata=None) -> None: Parameters ---------- df : pandas.DataFrame - DataFrame containing atom data with columns: ATOM, serial, name, - altloc, resname, chainid, resseq, icode, x, y, z, occupancy, - tempfactor, element, charge. + Atom table with columns ATOM, serial, name, altloc, resname, chainid, + resseq, icode, x, y, z (Cartesian, Å), occupancy, tempfactor (Ų), + element and charge. Rows whose optional ``anisou_flag`` is set also get + an ANISOU record from ``u11`` ... ``u23`` (Ų). filepath : str Output PDB filename. metadata : RefinementMetadata, optional Metadata to render as PDB header (REMARK 3, TITLE, etc.). + Raises + ------ + KeyError + If a required column is missing. + Notes ----- The CRYST1 record is sourced from the DataFrame attributes @@ -517,7 +643,9 @@ def write(df: pd.DataFrame, filepath: str, metadata=None) -> None: the file is written without a CRYST1 record and a warning is printed. Rows that fail to format are skipped with a printed warning; the - remaining rows are still written. + remaining rows are still written. Nothing is renumbered: duplicated atom + identifiers are written as they are (see + :func:`torchref.utils.sanitize_pdb_dataframe`). """ with open(filepath, "w") as n: # Write metadata header if provided (before CRYST1) @@ -548,85 +676,7 @@ def write(df: pd.DataFrame, filepath: str, metadata=None) -> None: except: print("No cell information found, writing without cell and spacegroup") - # Write atom records - for i, row in df.iterrows(): - ( - ATOM, - serial, - name, - altloc, - resname, - chainid, - resseq, - icode, - x, - y, - z_coord, - occupancy, - tempfactor, - element, - charge, - ) = row[ - [ - "ATOM", - "serial", - "name", - "altloc", - "resname", - "chainid", - "resseq", - "icode", - "x", - "y", - "z", - "occupancy", - "tempfactor", - "element", - "charge", - ] - ] - - if charge > 0: - charge = "+" + str(charge) - elif charge == 0: - charge = "" - else: - charge = str(charge) - - # 4-character PDB atom-name field (cols 13-16); preceded by the - # blank col 12 in the format string below. - name_field = _format_pdb_atom_name(name, element) - - if chainid is None or str(chainid) == "nan": - chainid = "" - - try: - s = ( - f"{str(ATOM):<6}{int(serial):>5} {name_field}{str(altloc):>1}" - f"{str(resname):>3}{str(chainid):>2}{int(resseq):>4}{str(icode):>4}" - f"{x:>8.3f}{y:>8.3f}{z_coord:>8.3f}" - f"{occupancy:>6.2f}{tempfactor:>6.2f}" - f"{str(element):>12}{charge:>2}\n" - ) - n.write(s) - except: - print("row", i, "failed") - print(row) - - # Write ANISOU record if present - if row["anisou_flag"]: - u11, u22, u33, u12, u13, u23 = row[ - ["u11", "u22", "u33", "u12", "u13", "u23"] - ] - s = ( - f"ANISOU{int(serial):>5} {name_field}{str(altloc):>1}" - f"{str(resname):>3}{str(chainid):>2}{int(resseq):>4} " - f"{int(u11 * 1e4):>{7}}{int(u22 * 1e4):>{7}}{int(u33 * 1e4):>{7}}" - f"{int(u12 * 1e4):>{7}}{int(u13 * 1e4):>{7}}{int(u23 * 1e4):>{7}}" - f" {str(element):>{2}}{str(charge):>2}\n" - ) - n.write(s) - + _write_atom_records(n, df, anisou=True) n.write("END") @@ -639,17 +689,24 @@ def write_multi_model( Write multiple models to a single PDB file with MODEL/ENDMDL records. Each DataFrame is wrapped in a MODEL/ENDMDL pair, producing a - multi-model PDB file suitable for ensemble or time-resolved data. + multi-model PDB file suitable for ensemble or time-resolved data. Atom + records are formatted as by :func:`write`, but without ANISOU records: + each model's ADPs are its isotropic ``tempfactor``. Parameters ---------- dataframes : list of pandas.DataFrame - List of atom DataFrames (same format as ``write()`` expects). + List of atom DataFrames, with the columns :func:`write` requires. filepath : str Output PDB filename. model_names : list of str, optional Names for each model (written as REMARK before each MODEL record). If None, models are numbered sequentially. + + Raises + ------ + KeyError + If a DataFrame lacks a required column. """ if not dataframes: return @@ -685,50 +742,7 @@ def write_multi_model( if model_names and model_idx < len(model_names): f.write(f"REMARK 3 MODEL {model_num}: {model_names[model_idx]}\n") f.write(f"MODEL {model_num:>4}\n") - - for i, row in df.iterrows(): - ATOM = row.get("ATOM", "ATOM") - serial = row.get("serial", i + 1) - name = str(row.get("name", "CA")) - altloc = str(row.get("altloc", "")) - resname = str(row.get("resname", "UNK")) - chainid = str(row.get("chainid", "")) - resseq = int(row.get("resseq", 1)) - icode = str(row.get("icode", "")) - x = float(row.get("x", 0.0)) - y = float(row.get("y", 0.0)) - z_coord = float(row.get("z", 0.0)) - occupancy = float(row.get("occupancy", 1.0)) - tempfactor = float(row.get("tempfactor", 20.0)) - element = str(row.get("element", "C")) - charge = row.get("charge", 0) - - if charge > 0: - charge_str = "+" + str(charge) - elif charge == 0: - charge_str = "" - else: - charge_str = str(charge) - - # 4-character PDB atom-name field (cols 13-16); preceded by the - # blank col 12 in the format string below. - name_field = _format_pdb_atom_name(name, element) - - if chainid is None or chainid == "nan": - chainid = "" - - try: - s = ( - f"{str(ATOM):<6}{int(serial):>5} {name_field}{altloc:>1}" - f"{resname:>3}{chainid:>2}{resseq:>4}{icode:>4}" - f"{x:>8.3f}{y:>8.3f}{z_coord:>8.3f}" - f"{occupancy:>6.2f}{tempfactor:>6.2f}" - f"{element:>12}{charge_str:>2}\n" - ) - f.write(s) - except Exception: - pass - + _write_atom_records(f, df, anisou=False) f.write("ENDMDL\n") f.write("END\n") diff --git a/torchref/utils/utils.py b/torchref/utils/utils.py index f5a1e669..5b37e412 100644 --- a/torchref/utils/utils.py +++ b/torchref/utils/utils.py @@ -6,8 +6,8 @@ - :class:`TensorDict` -- dict-like tensor container backed by ``nn.Module`` buffers. - :class:`TensorMasks` -- ``dict`` of boolean masks with device movement and a cached combined (logical-AND) mask. -- :func:`sanitize_pdb_dataframe` -- repair duplicate atom identifiers and over-long - residue names in a PDB/CIF DataFrame. +- :func:`sanitize_pdb_dataframe` -- renumber HETATM residues whose atom identifiers + repeat, and truncate over-long residue names, before an atom table is written. - :func:`parse_phenix_selection` / :func:`create_selection_mask` -- Phenix-style atom-selection strings to boolean masks. """ @@ -307,27 +307,59 @@ def __repr__(self): return f"TensorMasks({{{mask_info}}}, device={self.device})" -def sanitize_pdb_dataframe(pdb: pd.DataFrame, verbose: int = 0) -> pd.DataFrame: +#: What identifies one atom in a PDB or mmCIF file. +_ATOM_KEY = ["chainid", "resseq", "icode", "name", "altloc"] + +#: What the rows of one residue share. +_RESIDUE_KEY = ["chainid", "resseq", "icode", "resname"] + + +def _residue_blocks(pdb: pd.DataFrame) -> np.ndarray: + """Residue index of every row, shape ``(n_atoms,)``, non-decreasing down the table. + + A residue is a contiguous run of one ``(chainid, resseq, icode, resname)``, split + wherever an atom ``(name, altloc)`` repeats: unnumbered waters share that whole key, + yet each is a residue of its own. """ - Repair a PDB/CIF DataFrame so ``(chainid, resseq, name, altloc)`` is unique. + keys = pdb.groupby(_RESIDUE_KEY, sort=False, dropna=False).ngroup().to_numpy() + names = ["name", "altloc"] + atoms = pdb.groupby(names, sort=False, dropna=False).ngroup().to_numpy() + blocks = np.empty(len(pdb), dtype=np.int64) + block, current, seen = -1, None, set() + for row, (key, atom) in enumerate(zip(keys.tolist(), atoms.tolist())): + if key != current or atom in seen: + block, current, seen = block + 1, key, set() + seen.add(atom) + blocks[row] = block + return blocks + - Fixes duplicate ``resseq`` on HETATM records (waters are often all 0) by renumbering - within the chain, and truncates residue names to 3 characters. Returns a copy; the - input is not modified. +def sanitize_pdb_dataframe(pdb: pd.DataFrame, verbose: int = 0) -> pd.DataFrame: + """ + Prepare an atom table for writing: unique HETATM residues, 3-character names. + + Truncates residue names to the 3 characters a PDB file holds. Then each HETATM + residue that repeats an atom identifier ``(chainid, resseq, icode, name, altloc)`` + already taken by an ATOM record or an earlier residue -- typically waters all + numbered 0, or a ligand copied without renumbering -- gets a new ``resseq``, counting + up from the chain's highest. A residue is a contiguous run of rows sharing + ``(chainid, resseq, icode, resname)``, split where an atom ``(name, altloc)`` + repeats, so it moves whole and a run of unnumbered waters becomes one residue per + water. ATOM records are never renumbered, and residues kept apart by an insertion + code (52 and 52A) are not duplicates. Returns a copy; the input is not modified. Parameters ---------- pdb : pandas.DataFrame - DataFrame with PDB data (must have columns: ATOM, chainid, resseq, name, altloc, - resname, serial). + Atom table with columns ATOM, chainid, resseq, icode, resname, name and altloc. verbose : int, default 0 Verbosity level (0=silent, 1=info, 2=debug). Returns ------- pandas.DataFrame - Sanitized copy. Renumbering can fail to converge on pathological input, in which - case duplicates remain and a warning is printed at ``verbose > 0``. + Sanitized copy. Duplicated identifiers among ATOM records are left as they are, + with a warning printed at ``verbose > 0``. """ pdb = pdb.copy() @@ -335,7 +367,6 @@ def sanitize_pdb_dataframe(pdb: pd.DataFrame, verbose: int = 0) -> pd.DataFrame: print("Sanitizing PDB DataFrame...") print(f" Initial atoms: {len(pdb)}") - # 1. Standardize residue names to max 3 characters long_resnames = pdb["resname"].str.len() > 3 if long_resnames.any(): n_long = long_resnames.sum() @@ -346,80 +377,44 @@ def sanitize_pdb_dataframe(pdb: pd.DataFrame, verbose: int = 0) -> pd.DataFrame: ) pdb.loc[long_resnames, "resname"] = pdb.loc[long_resnames, "resname"].str[:3] - # 2. Fix duplicate atom identifiers by reassigning resseq - dup_mask = pdb.duplicated( - subset=["chainid", "resseq", "name", "altloc"], keep=False - ) - - if dup_mask.any(): - n_dup = dup_mask.sum() + het = (pdb["ATOM"].astype(str).str.strip() == "HETATM").to_numpy() + residue = _residue_blocks(pdb) + # ATOM records come first, so of a polymer residue and a HETATM residue that + # collide it is always the HETATM one that moves. + order = np.argsort(het, kind="stable") + taken = np.empty(len(pdb), dtype=bool) + taken[order] = pdb.iloc[order].duplicated(subset=_ATOM_KEY).to_numpy() + moved = np.unique(residue[taken & het]) + + if len(moved): + first_row = np.searchsorted(residue, moved) + chain = pd.Series(pdb["chainid"].to_numpy()[first_row]) + by_chain = pdb.groupby("chainid", sort=False, dropna=False)["resseq"] + top = by_chain.transform("max").to_numpy()[first_row] + offset = chain.groupby(chain, sort=False, dropna=False).cumcount().to_numpy() + new = np.where(top > 0, top + 1, 1) + offset + rows = np.isin(residue, moved) + new_resseq = pd.Series(new, index=moved).loc[residue[rows]] + pdb.loc[rows, "resseq"] = new_resseq.to_numpy() if verbose > 0: - print(f" Found {n_dup} atoms with duplicate identifiers") - - # This ensures we only renumber within the same molecule type and chain - for (chainid, resname, atom_type), group in pdb.groupby( - ["chainid", "resname", "ATOM"] - ): - group_indices = group.index - - group_dup_mask = group.duplicated( - subset=["chainid", "resseq", "name", "altloc"], keep=False + print( + f" Renumbered {len(moved)} HETATM residues ({rows.sum()} atoms) " + "whose atom identifiers were already taken" ) - - if group_dup_mask.any(): - chain_data = pdb[pdb["chainid"] == chainid] - max_resseq = chain_data["resseq"].max() - - new_resseq_start = ( - max_resseq + 1 if pd.notna(max_resseq) and max_resseq > 0 else 1 - ) - - # Group by (serial) to keep atoms of the same residue together - unique_serials = group["serial"].unique() - residue_counter = new_resseq_start - - for serial in unique_serials: - serial_mask = pdb["serial"] == serial - pdb.loc[serial_mask, "resseq"] = residue_counter - residue_counter += 1 - - if verbose > 1: - n_fixed = len(unique_serials) - print( - f" Fixed {n_fixed} {resname} residues in chain {chainid} (resseq {new_resseq_start}-{residue_counter-1})" - ) - - final_dup_mask = pdb.duplicated( - subset=["chainid", "resseq", "name", "altloc"], keep=False - ) - if final_dup_mask.any(): - remaining_dups = final_dup_mask.sum() - if verbose > 0: - print( - f" WARNING: Still have {remaining_dups} duplicate identifiers after sanitization" - ) - dups = pdb[final_dup_mask].sort_values(["chainid", "resseq", "name"]) - print( - dups[ - [ - "ATOM", - "serial", - "name", - "resname", - "chainid", - "resseq", - "altloc", - ] - ].head(10) - ) - else: - if verbose > 0: - print(" ✓ All duplicate identifiers resolved") - else: - if verbose > 0: - print(" ✓ No duplicate atom identifiers found") + if verbose > 1: + for chainid, numbers in pd.Series(new).groupby(chain, dropna=False): + print(f" chain {chainid}: resseq {numbers.min()}-{numbers.max()}") + elif verbose > 0: + print(" No HETATM residue needed renumbering") if verbose > 0: + remaining = pdb.duplicated(subset=_ATOM_KEY, keep=False) + if remaining.any(): + print( + f" WARNING: {remaining.sum()} ATOM records share an atom identifier " + "and are left as they are" + ) + print(pdb.loc[remaining, ["ATOM", "resname", *_ATOM_KEY]].head(10)) print(f" Final atoms: {len(pdb)}") return pdb From 605a3ce9d7fbd591ea88664426c776aef27a684d Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 22:08:05 +0000 Subject: [PATCH 246/250] Compute anomalous f' and f'' through the density and FFT path ModelFT added the anomalous correction to F as a second, separate structure-factor sum, occ * (f' + i f'') * exp(2 pi i h.x), over the asymmetric-unit atoms only and without any temperature factor, on top of an F0 that the FFT path had already B-damped and symmetry-expanded. F_calc was therefore wrong in every space group whenever Se, Br or a heavier atom was present, which is the default (wavelength = 1.0 A): symmetry mates no longer shared |F|, the top shell of 3E98 was 12.5% off, the missing Debye-Waller factor made it wrong in P1 too, and Bijvoet differences correlated with the true ones at CC 0.3-0.6. f' and f'' do not depend on the scattering angle, so they are zero-width terms of an atom's form factor: the slot ITC92's constant c already occupies (column CONSTANT_TERM, B = 0). ModelFT now adds f' to that column of the amplitudes it hands to SfFFT, and puts f'' into the same column of a copy of the anomalous atoms, which SfFFT splats into the imaginary part of a complex density, F = FT(rho') + i FT(rho''). The one splat applies each atom's isotropic or anisotropic temperature factor and the one FFT path the symmetry (in reciprocal space, or on both maps), so every term of F_calc has a single implementation. The separate sum, _apply_anomalous_correction, is deleted. No kernel changed: every backend already takes per-atom (n, 5) Gaussian coefficients. The anomalous terms and the forward cache now also key on wavelength and anomalous_threshold, so changing either takes effect without recalc. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- tests/unit/model/test_model_ft_anomalous.py | 372 +++++++++++++++++++ tests/unit/scattering/test_anomalous.py | 42 ++- torchref/base/scattering/scattering_table.py | 8 +- torchref/model/model_ft.py | 257 +++++++------ torchref/model/sf_fft.py | 77 +++- 5 files changed, 616 insertions(+), 140 deletions(-) create mode 100644 tests/unit/model/test_model_ft_anomalous.py diff --git a/tests/unit/model/test_model_ft_anomalous.py b/tests/unit/model/test_model_ft_anomalous.py new file mode 100644 index 00000000..869a275c --- /dev/null +++ b/tests/unit/model/test_model_ft_anomalous.py @@ -0,0 +1,372 @@ +"""Anomalous scattering in ``ModelFT``'s F_calc, pinned against independent references. + +f' and f'' enter F_calc as zero-width terms of each atom's form factor, so the one +density splat and FFT give them the atom's own temperature factor (isotropic or +anisotropic) and the space-group symmetry, as for f0. The references share none of +that path: + +* f' -- gemmi's ``StructureFactorCalculatorX`` with ``addends``, which adds a constant + to an element's form factor inside gemmi's own symmetry and Debye-Waller sums. +* f'' -- gemmi has no imaginary addend, so :func:`_constant_term_sum` sums it + explicitly in float64 over every symmetry mate with each atom's temperature factor. + That helper is itself checked against gemmi's addends, with the real f'. + +Both sides are driven by the model's own atoms (the gemmi structure is built from the +model's tensors), so reader differences cannot pose as structure-factor error. The +overall accuracy is gated against the FFT path's own error without anomalous terms, +measured against the same reference: grid sampling and truncation set the precision +this path can deliver, and the anomalous terms must not add to it. Each structure is +evaluated at its deposited resolution, on the grid production would build for it. + +Structures: 3E98 (P 1 21 1, isotropic, Se), 5BOV (P 1, anisotropic, Se) and 6G9X +(P 21 21 2, anisotropic, Hg, whose f'' of 10 e makes large Bijvoet differences). +""" + +from types import SimpleNamespace + +import gemmi +import numpy as np +import pytest +import torch + +from torchref.base.scattering.anomalous_table import get_significant_elements +from torchref.model import ModelFT + +N_REFL = 400 +WAVELENGTH = 1.0 + +#: Factor by which the anomalous terms may grow the FFT path's own error, overall and +#: in the top resolution quarter. Measured at most 1.011; terms summed over the +#: asymmetric unit without temperature factors grow it 1.9x-12x (top quarter 16x-46x). +ERROR_GROWTH = 1.05 + +#: Relative error allowed on the anomalous contribution itself (``F - F(f0)`` against +#: the reference's) and on the Bijvoet differences. Measured 0.03%-0.09% and +#: 0.23%-0.65%; without symmetry mates and temperature factors both exceed 100%. +TERM_TOL = 0.05 + +STRUCTURES = [ + pytest.param(("3E98", 2.5), id="3E98-P1211-iso"), + pytest.param(("5BOV", 1.6), id="5BOV-P1-aniso"), + pytest.param(("6G9X", 2.3), id="6G9X-P21212-aniso-Hg"), +] + + +def _model(pdb_dir, code, d_min, **kwargs) -> ModelFT: + path = pdb_dir / f"{code}.pdb" + if not path.exists(): + pytest.skip(f"{code}.pdb fixture not present") + kwargs.setdefault("wavelength", WAVELENGTH) + return ModelFT(max_res=d_min, verbose=0, device="cpu", **kwargs).load_pdb(str(path)) + + +def _gemmi_structure(model: ModelFT) -> gemmi.Structure: + """A gemmi structure holding exactly the model's atoms, ADPs and occupancies.""" + st = gemmi.Structure() + st.cell = gemmi.UnitCell(*[float(v) for v in model.cell.data.tolist()]) + st.spacegroup_hm = model.spacegroup.hm + xyz = model.xyz().detach().double().numpy() + adp = model.adp().detach().double().numpy() + occ = model.occupancy().detach().double().numpy() + u = model.u().detach().double().numpy() + aniso = model.aniso_flag.numpy() + chain = gemmi.Chain("A") + for i, element in enumerate(model.ctx.topology.atoms.element.tolist()): + atom = gemmi.Atom() + atom.name = "X" + atom.element = gemmi.Element(element) + atom.pos = gemmi.Position(*xyz[i]) + atom.occ = float(occ[i]) + atom.b_iso = float(adp[i]) + if aniso[i]: + atom.aniso = gemmi.SMat33f(*[float(v) for v in u[i]]) + residue = gemmi.Residue() + residue.name = "UNK" + residue.seqid = gemmi.SeqId(i + 1, " ") + residue.add_atom(atom) + chain.add_residue(residue) + gm = gemmi.Model("1") + gm.add_chain(chain) + st.add_model(gm) + st.setup_cell_images() + return st + + +def _gemmi_sf(structure, hkl: np.ndarray, addends=None) -> torch.Tensor: + calc = gemmi.StructureFactorCalculatorX(structure.cell) + for element, value in (addends or {}).items(): + calc.addends.set(gemmi.Element(element), value) + return torch.tensor( + [ + complex(calc.calculate_sf_from_model(structure[0], [int(v) for v in h])) + for h in hkl + ], + dtype=torch.complex128, + ) + + +def _constant_term_sum(model: ModelFT, hkl: np.ndarray, values: dict) -> torch.Tensor: + """``sum_mates occ_j c_j T_j exp(2 pi i h.x_j)`` in float64, for c_j = ``values``. + + Every symmetry mate is generated explicitly, ``x' = R x + t``, with the mate's + temperature factor ``exp(-B s^2 / 4)`` or ``exp(-2 pi^2 s^T U' s)`` from the rotated + Cartesian ``U' = R_c U R_c^T``. Only atoms whose element is in ``values`` are summed. + """ + orth = np.array(gemmi.UnitCell(*model.cell.data.tolist()).orth.mat.tolist()) + frac = np.linalg.inv(orth) + h = np.asarray(hkl, dtype=np.float64) + s = h @ frac # Cartesian scattering vectors, row-wise + s2 = (s * s).sum(axis=1) + + xyz = model.xyz().detach().double().numpy() + adp = model.adp().detach().double().numpy() + occ = model.occupancy().detach().double().numpy() + u = model.u().detach().double().numpy() + aniso = model.aniso_flag.numpy() + elements = model.ctx.topology.atoms.element.tolist() + ops = gemmi.find_spacegroup_by_name(model.spacegroup.hm).operations() + + total = np.zeros(len(h), dtype=np.complex128) + for op in ops: + rot = np.array(op.rot, dtype=np.float64) / op.DEN + tran = np.array(op.tran, dtype=np.float64) / op.DEN + rot_cart = orth @ rot @ frac + for i, element in enumerate(elements): + if element not in values: + continue + phase = np.exp(2j * np.pi * (h @ (rot @ (frac @ xyz[i]) + tran))) + if aniso[i]: + u11, u22, u33, u12, u13, u23 = u[i] + U = np.array([[u11, u12, u13], [u12, u22, u23], [u13, u23, u33]]) + U = rot_cart @ U @ rot_cart.T + dwf = np.exp(-2 * np.pi**2 * np.einsum("ri,ij,rj->r", s, U, s)) + else: + dwf = np.exp(-adp[i] * s2 / 4.0) + total += occ[i] * values[element] * dwf * phase + return torch.from_numpy(total) + + +def _anomalous_terms(model: ModelFT): + """``({element: f'}, {element: f''})`` for the elements the model treats as anomalous.""" + significant = get_significant_elements( + sorted(set(model.ctx.topology.atoms.element.tolist())), + model.wavelength, + model.anomalous_threshold, + ) + assert significant, "no significant anomalous scatterer: the test would be vacuous" + return ( + {e: fp for e, (fp, _) in significant.items()}, + {e: fdp for e, (_, fdp) in significant.items()}, + ) + + +def _asu_hkl(cell, spacegroup, d_min: float, n: int = N_REFL) -> np.ndarray: + """``n`` ASU reflections spread evenly over the resolution range to ``d_min``.""" + hkl = gemmi.make_miller_array(cell, spacegroup, d_min) + hkl = hkl[np.argsort(cell.calculate_d_array(hkl))] + return hkl[np.linspace(0, len(hkl) - 1, n).round().astype(int)] + + +def _model_hkl(model: ModelFT, d_min: float, n: int = N_REFL) -> np.ndarray: + cell = gemmi.UnitCell(*[float(v) for v in model.cell.data.tolist()]) + return _asu_hkl(cell, gemmi.find_spacegroup_by_name(model.spacegroup.hm), d_min, n) + + +def _F(model: ModelFT, hkl: np.ndarray, **kwargs) -> torch.Tensor: + with torch.no_grad(): + F = model(torch.tensor(hkl, dtype=torch.int32), recalc=True, **kwargs) + return F.to(torch.complex128) + + +def _rel(got: torch.Tensor, ref: torch.Tensor) -> float: + return float((got - ref).norm() / ref.norm()) + + +@pytest.fixture(scope="module", params=STRUCTURES) +def scene(request, pdb_dir): + """One structure's model, reflections ``(asu, -asu)`` and every F the tests compare. + + ``F_f0`` has no anomalous term, ``F_fp`` adds f' and ``F_fdp`` adds f' and f''. + ``G_f0`` / ``G_fp`` are gemmi's without and with the f' addends. + """ + code, d_min = request.param + model = _model(pdb_dir, code, d_min, apply_bijvoet=True) + st = _gemmi_structure(model) + f_prime, f_double_prime = _anomalous_terms(model) + asu = _model_hkl(model, d_min) + hkl = np.concatenate([asu, -asu]) + d = st.cell.calculate_d_array(hkl) + + F_fdp = _F(model, hkl) + F_f0 = _F(model, hkl, apply_anomalous=False) + model.anomalous_bijvoet.fill_(False) + F_fp = _F(model, hkl) + model.anomalous_bijvoet.fill_(True) + return SimpleNamespace( + code=code, + model=model, + hkl=hkl, + n=len(asu), + top=torch.from_numpy(d < np.quantile(d, 0.25)), + f_prime=f_prime, + f_double_prime=f_double_prime, + F_f0=F_f0, + F_fp=F_fp, + F_fdp=F_fdp, + G_f0=_gemmi_sf(st, hkl), + G_fp=_gemmi_sf(st, hkl, f_prime), + ) + + +@pytest.mark.unit +def test_f_prime_matches_gemmi_addends(scene): + """f' (Bijvoet off) matches gemmi with addends as closely as f0 alone matches gemmi. + + The top resolution quarter is gated on its own because a dispersive term missing + the Debye-Waller factor is worst there, and the f' contribution ``F - F(f0)`` is + gated directly because in the totals it is diluted by the f0 grid error. + """ + s, top = scene, scene.top + err_f0, err_fp = _rel(s.F_f0, s.G_f0), _rel(s.F_fp, s.G_fp) + top_f0, top_fp = _rel(s.F_f0[top], s.G_f0[top]), _rel(s.F_fp[top], s.G_fp[top]) + term = _rel(s.F_fp - s.F_f0, s.G_fp - s.G_f0) + print( + f"\n {s.code} f' {s.f_prime}: rel L2 vs gemmi {err_fp:.3e} (f0 alone " + f"{err_f0:.3e}); top quarter {top_fp:.3e} ({top_f0:.3e}); f' term {term:.3e}" + ) + assert err_fp < ERROR_GROWTH * err_f0 + assert top_fp < ERROR_GROWTH * top_f0 + assert term < TERM_TOL + + +@pytest.mark.unit +def test_f_double_prime_matches_explicit_sum(scene): + """With f'' on, F, the f'' term and the Bijvoet differences match the reference. + + The reference is gemmi's F with the f' addends plus ``i`` times the explicit f'' + sum. The Bijvoet differences ``|F(h)| - |F(-h)|`` come from f'' alone. + """ + s, n = scene, scene.n + # The explicit sum must reproduce gemmi's own f' contribution before it is + # trusted with f''. + assert _rel(_constant_term_sum(s.model, s.hkl, s.f_prime), s.G_fp - s.G_f0) < 1e-6 + fdp_ref = _constant_term_sum(s.model, s.hkl, s.f_double_prime) + ref = s.G_fp + 1j * fdp_ref + + err_f0, err_fdp = _rel(s.F_f0, s.G_f0), _rel(s.F_fdp, ref) + term = _rel((s.F_fdp - s.F_fp) / 1j, fdp_ref) + bijvoet = s.F_fdp[:n].abs() - s.F_fdp[n:].abs() + bijvoet_ref = ref[:n].abs() - ref[n:].abs() + rms_ref = float(bijvoet_ref.pow(2).mean().sqrt()) + bijvoet_err = float((bijvoet - bijvoet_ref).pow(2).mean().sqrt()) / rms_ref + cc = float(np.corrcoef(bijvoet.numpy(), bijvoet_ref.numpy())[0, 1]) + print( + f"\n {s.code} f'' {s.f_double_prime}: rel L2 {err_fdp:.3e} (f0 alone " + f"{err_f0:.3e}); f'' term {term:.3e}; Bijvoet differences rms " + f"{rms_ref:.3f} e, error {bijvoet_err:.2%}, CC {cc:.6f}" + ) + assert err_fdp < ERROR_GROWTH * err_f0 + assert term < TERM_TOL + assert bijvoet_err < TERM_TOL + assert cc > 0.999 + + +@pytest.mark.unit +@pytest.mark.parametrize("code, d_min", [("3E98", 2.5), ("6G9X", 2.3)]) +def test_symmetry_equivalents_share_amplitude(pdb_dir, code, d_min): + """Symmetry mates of h share |F|; Friedel mates share it only without f''. + + Both groups are chiral, so every equivalent ``hR`` is in the same Bijvoet class as + ``h`` and every ``-hR`` in the other. f'' separates the two classes and nothing + may separate members of one class. + """ + model = _model(pdb_dir, code, d_min) + ops = gemmi.find_spacegroup_by_name(model.spacegroup.hm).operations() + asu = _model_hkl(model, d_min, n=150) + mates = np.array( + [[op.apply_to_hkl([int(v) for v in h]) for op in ops] for h in asu] + ) + n_refl, n_ops = mates.shape[:2] + flat = mates.reshape(-1, 3) + hkl = np.concatenate([flat, -flat]) + + for bijvoet in (False, True): + model.anomalous_bijvoet.fill_(bijvoet) + amp = _F(model, hkl).abs().reshape(2, n_refl, n_ops) + scale = float(amp.pow(2).mean().sqrt()) + spread = float((amp.max(dim=2).values - amp.min(dim=2).values).max()) / scale + friedel = float((amp[0, :, 0] - amp[1, :, 0]).pow(2).mean().sqrt()) / scale + print( + f"\n {code} bijvoet={bijvoet}: worst spread across {n_ops} symmetry " + f"mates {spread:.2e}, rms Friedel difference {friedel:.2e} (of rms |F|)" + ) + assert spread < 1e-4 + if bijvoet: + assert friedel > 100 * max(spread, 1e-7) + else: + assert friedel < 1e-4 + + +@pytest.mark.unit +def test_imaginary_density_follows_early_symmetry(pdb_dir): + """The map-space symmetry path treats the f'' density as the reciprocal one does. + + Without late symmetry both parts of the complex density are symmetrized on the + grid before the FFT; the result must agree with the reciprocal-space expansion. + """ + model = _model(pdb_dir, "3E98", 2.5, apply_bijvoet=True) + asu = _model_hkl(model, 2.5, n=150) + hkl = np.concatenate([asu, -asu]) + late = _F(model, hkl) + model.fft.use_late_symmetry = False + early = _F(model, hkl) + assert model.ed.is_complex() + assert _rel(early, late) < 1e-5 + + +@pytest.mark.unit +def test_f_double_prime_reaches_the_gradient(pdb_dir): + """A Bijvoet-difference target differentiates through the imaginary density.""" + model = _model(pdb_dir, "3E98", 2.5, apply_bijvoet=True) + f_prime, _ = _anomalous_terms(model) + asu = torch.tensor(_model_hkl(model, 2.5, n=150), dtype=torch.int32) + elements = model.ctx.topology.atoms.element.tolist() + rows = [i for i, e in enumerate(elements) if e in f_prime] + + def bijvoet_grad_norm(apply_bijvoet): + model.anomalous_bijvoet.fill_(apply_bijvoet) + F = model(torch.cat([asu, -asu]), recalc=True) + loss = (F[: len(asu)].abs() - F[len(asu) :].abs()).pow(2).sum() + (grad,) = torch.autograd.grad(loss, model.xyz.refinable_params) + assert torch.isfinite(grad).all() + return float(grad[rows].norm()) + + # Without f'' the Bijvoet differences are float32 noise, and so is their gradient. + assert bijvoet_grad_norm(True) > 100 * bijvoet_grad_norm(False) + + +@pytest.mark.unit +def test_wavelength_none_is_the_f0_path(pdb_dir): + """``wavelength=None`` gives exactly the f0 transform, and so does a wavelength + at which no element is significant.""" + model = _model(pdb_dir, "3E98", 2.5, wavelength=None) + hkl = torch.tensor(_model_hkl(model, 2.5, n=150), dtype=torch.int32) + + with torch.no_grad(): + F_none = model(hkl, recalc=True) + F_f0, _ = model.fft.compute_structure_factors( + hkl, *model.get_iso(), *model.get_aniso(), apply_symmetry=True + ) + F_off = _model(pdb_dir, "3E98", 2.5)(hkl, recalc=True, apply_anomalous=False) + assert not model.ed.is_complex() + assert torch.equal(F_none, F_f0) + assert torch.equal(F_none, F_off) + + # 1DAW: Mg, P and S stay below the 0.5 e threshold at 1 A. + light = _model(pdb_dir, "1DAW", 2.5, apply_bijvoet=True) + assert light._get_anomalous_cache() is None + hkl = torch.tensor([[1, 2, 3], [-1, -2, -3], [4, 0, 2]], dtype=torch.int32) + with torch.no_grad(): + assert torch.equal( + light(hkl, recalc=True), light(hkl, recalc=True, apply_anomalous=False) + ) diff --git a/tests/unit/scattering/test_anomalous.py b/tests/unit/scattering/test_anomalous.py index 6fb0951b..297c498a 100644 --- a/tests/unit/scattering/test_anomalous.py +++ b/tests/unit/scattering/test_anomalous.py @@ -307,17 +307,11 @@ def test_friedel_pair_asymmetry(self, test_pdb_file): # with h, so it makes F(h) != F(-h)*. f' is real and dispersive and leaves the # conjugate relation intact. # - # In this model f'' is gated on ``apply_bijvoet`` (``ModelFT.__init__``, applied at - # ``model_ft.py:951`` via ``include_fdp``), which defaults to False because merged - # data is the usual target and Friedel-preserving F is correct for it. So the - # default path deliberately does *not* break Friedel's law -- both branches are - # asserted here rather than only the one this test originally assumed. - # - # History: this test previously computed ``is_conjugate`` and then ended in - # ``pass``, asserting nothing. A first attempt to fix it asserted breakdown on the - # default path and failed, because that path is Friedel-preserving by design. - mask, _, _, _, _ = model._get_anomalous_cache() - assert mask.any(), ( + # In this model f'' is gated on ``apply_bijvoet`` (``ModelFT.__init__``), which + # defaults to False because merged data is the usual target and + # Friedel-preserving F is correct for it. So the default path deliberately does + # *not* break Friedel's law, and both branches are asserted. + assert model._get_anomalous_cache() is not None, ( "no anomalous scatterers in this structure, so neither branch below is " "meaningful -- pick a structure with an anomalous element" ) @@ -388,21 +382,29 @@ def test_gradient_flow(self, test_pdb_file): assert model.xyz.refinable_params.grad is not None, "Gradients should flow to xyz" def test_cache_invalidation(self, test_pdb_file): - """Test that anomalous cache is invalidated when elements change.""" + """A new wavelength reaches F on the next call, with no ``recalc``. + + ``wavelength`` is a plain attribute, so both the anomalous terms and the + forward cache have to key on it rather than on parameters and buffers alone. + """ from torchref.model import ModelFT model = ModelFT(wavelength=1.0, verbose=0) model.load_pdb(test_pdb_file) + hkl = torch.tensor([[1, 2, 3], [2, 1, 0]], dtype=torch.int32) - # Access cache - _ = model._get_anomalous_cache() - original_hash = model._anomalous_elements_hash + first = model(hkl).detach().clone() + assert model._get_anomalous_cache() is not None + model.wavelength = 1.5418 # Cu K-alpha: Fe f' goes from +0.28 to -1.14 + second = model(hkl).detach() - # The hash should be set - assert original_hash is not None - - # If we modify the element list (hypothetically), the cache should be invalidated - # This is tested implicitly by checking the hash mechanism works + assert not torch.allclose(first, second), ( + "changing the wavelength left F unchanged: a stale anomalous or forward " + "cache served the old f'" + ) + fresh = ModelFT(wavelength=1.5418, verbose=0) + fresh.load_pdb(test_pdb_file) + torch.testing.assert_close(second, fresh(hkl).detach()) class TestAnomalousValuesRealistic: diff --git a/torchref/base/scattering/scattering_table.py b/torchref/base/scattering/scattering_table.py index 8a3d8332..a3dfcc1f 100644 --- a/torchref/base/scattering/scattering_table.py +++ b/torchref/base/scattering/scattering_table.py @@ -14,6 +14,11 @@ from torchref.config import get_float_dtype, get_int_dtype +#: Column of the ITC92 constant ``c`` in the ``(n, 5)`` ``A`` / ``B`` coefficient arrays. +#: ``c`` is stored as a fifth Gaussian of zero width (``B[:, CONSTANT_TERM] == 0``), so +#: this column is also where any other angle-independent term, such as f' or f'', goes. +CONSTANT_TERM = 4 + # Global cache for the loaded table _TABLE_CACHE: Optional[dict] = None @@ -49,7 +54,8 @@ def load_scattering_table( Returns ------- dict - - 'A', 'B': Tensor(max_z + 1, 5), neutral coefficients indexed by Z + - 'A', 'B': Tensor(max_z + 1, 5), neutral coefficients indexed by Z; column + :data:`CONSTANT_TERM` holds ``c`` with ``B = 0`` - 'element_to_z' / 'z_to_element': symbol/number mappings - 'ions': ion key -> (A, B) - 'metadata': source information diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index a58814eb..af63aaa6 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -2,11 +2,13 @@ Adds the electron-density / FFT path (an :class:`~torchref.model.SfFFT` submodule that reads the crystal off the model's context and sizes its grid lazily), the -ITC92 scattering parametrization, and the anomalous f' / f'' correction. +ITC92 scattering parametrization, and the anomalous f' / f'' terms. Those enter the +same density as f0 rather than a separate sum, so every term of F_calc gets the same +temperature factors and symmetry expansion (see :meth:`ModelFT.forward`). """ import math -from typing import Optional, Tuple +from typing import NamedTuple, Optional, Tuple import gemmi import numpy as np @@ -20,6 +22,26 @@ from torchref.utils.caching import CachedForwardMixin +class _AnomalousTerms(NamedTuple): + """f' and f'' laid out for :meth:`ModelFT._add_anomalous_scattering`. + + ``f_prime_iso`` / ``f_prime_aniso`` are ``(n_iso, 5)`` / ``(n_aniso, 5)`` addends to + the ITC92 amplitudes of :meth:`ModelFT.get_iso` / :meth:`ModelFT.get_aniso`: f' in + electrons in the zero-width column (``CONSTANT_TERM``) for atoms above + ``anomalous_threshold``, zero everywhere else. ``rows_iso`` / ``rows_aniso`` index + those atoms within the two subsets, and ``f_double_prime_iso`` / + ``f_double_prime_aniso`` are their f'' amplitudes, ``(len(rows), 5)``, laid out the + same way. + """ + + f_prime_iso: torch.Tensor + f_prime_aniso: torch.Tensor + rows_iso: torch.Tensor + rows_aniso: torch.Tensor + f_double_prime_iso: torch.Tensor + f_double_prime_aniso: torch.Tensor + + class ModelFT(CachedForwardMixin, Model): """ Model subclass for FFT-based electron density and structure factors. @@ -131,10 +153,8 @@ def __init__( torch.tensor(bool(apply_bijvoet), device=self.device), persistent=True, ) - self._anomalous_cache = None # Will hold (mask, f_prime, f_double_prime) - self._anomalous_elements_hash = ( - None # Hash of element list for cache invalidation - ) + # (key, partition, _AnomalousTerms or None); see _get_anomalous_cache. + self._anomalous_cache = None # ========================================================================= # Engine binding and grid inputs @@ -176,12 +196,17 @@ def grid_key(self): return self.fft.grid_key def _fingerprint_state(self): - """Fold the grid key into the forward-cache key. + """Fold the grid key and the anomalous settings into the forward-cache key. Parameters and buffers alone would miss a cell, space-group or resolution - change that leaves the grid buffers untouched until the next forward. + change that leaves the grid buffers untouched until the next forward, and a + new ``wavelength`` or ``anomalous_threshold``, which are plain attributes. """ - return super()._fingerprint_state() + (self.fft.grid_key,) + return super()._fingerprint_state() + ( + self.wavelength, + self.anomalous_threshold, + self.fft.grid_key, + ) # ========================================================================= # Backward-compatible properties for scattering parameters @@ -455,7 +480,6 @@ def reset_cache(self): # Drop the anomalous scattering cache; it is recomputed on next use # and would otherwise hold tensors on the previous device. self._anomalous_cache = None - self._anomalous_elements_hash = None for module in self.children(): if hasattr(module, "reset_forward_cache"): module.reset_forward_cache() @@ -465,109 +489,128 @@ def invalidate_cache(self): self.reset_cache() # ========================================================================= - # Anomalous Scattering Correction Methods + # Anomalous scattering # ========================================================================= - def _get_anomalous_cache( - self, - ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Cached ``(mask, f_prime, f_double_prime, has_anomalous, indices)``. + def _get_anomalous_cache(self) -> Optional[_AnomalousTerms]: + """f' and f'' of the atoms above ``anomalous_threshold``; None if there are none. - ``mask`` is per-atom; ``f_prime`` / ``f_double_prime`` cover only the - significant scatterers. Recomputed when the element list changes. + Rebuilt when the element list, ``wavelength``, ``anomalous_threshold`` or the + iso/aniso partition changes. Building it costs a device sync; using it, none. + + Raises + ------ + RuntimeError + If an anomalous atom's ITC92 column ``CONSTANT_TERM`` has a nonzero width, + which would spread f' and f'' like an f0 Gaussian. """ from torchref.base.scattering.anomalous_table import ( get_anomalous_corrections_by_indices, get_significant_elements, ) - - element_list = self.ctx.topology.atoms.element.tolist() - elements_hash = hash(tuple(element_list)) - - if ( - self._anomalous_cache is None - or self._anomalous_elements_hash != elements_hash - ): - unique_elements = list(set(element_list)) - significant = get_significant_elements( - unique_elements, self.wavelength, self.anomalous_threshold - ) - - if self.ctx.verbose > 1 and significant: + from torchref.base.scattering.scattering_table import CONSTANT_TERM + + elements = self.ctx.topology.atoms.element.tolist() + key = (hash(tuple(elements)), self.wavelength, self.anomalous_threshold) + # Compared by identity: ``_sf_partition`` hands back the same tuple until the + # aniso flags or the hydrogen choice change. + partition = self._sf_partition() + cached = self._anomalous_cache + if cached is not None and cached[0] == key and cached[1] is partition: + return cached[2] + + terms = None + significant = get_significant_elements( + sorted(set(elements)), self.wavelength, self.anomalous_threshold + ) + if significant: + if self.ctx.verbose > 1: print( f"Anomalous scatterers at {self.wavelength:.4f} Å: " - f"{list(significant.keys())}" + f"{sorted(significant)}" ) - mask, f_prime, f_double_prime = get_anomalous_corrections_by_indices( - element_list, significant, self.device, self.dtype_float + elements, significant, self.device, self.dtype_float ) - - # Pre-compute integer indices to avoid boolean indexing GPU sync - has_anomalous = bool(mask.any().item()) - anomalous_indices = ( - mask.nonzero(as_tuple=True)[0] if has_anomalous else None - ) - self._anomalous_cache = ( - mask, - f_prime, - f_double_prime, - has_anomalous, - anomalous_indices, + rows = mask.nonzero(as_tuple=True)[0] + if bool((self.B[rows, CONSTANT_TERM] != 0).any()): + raise RuntimeError( + f"{type(self).__name__}: ITC92 column {CONSTANT_TERM} must be the " + "zero-width constant term to carry f' and f'', but an anomalous " + "atom has a nonzero width there." + ) + addend_fp = self.A.new_zeros(len(elements), self.A.shape[1]) + addend_fdp = torch.zeros_like(addend_fp) + addend_fp[rows, CONSTANT_TERM] = f_prime + addend_fdp[rows, CONSTANT_TERM] = f_double_prime + + iso_idx, aniso_idx = partition[0], partition[1] + rows_iso = mask[iso_idx].nonzero(as_tuple=True)[0].to(dtypes.int) + rows_aniso = mask[aniso_idx].nonzero(as_tuple=True)[0].to(dtypes.int) + terms = _AnomalousTerms( + f_prime_iso=addend_fp[iso_idx], + f_prime_aniso=addend_fp[aniso_idx], + rows_iso=rows_iso, + rows_aniso=rows_aniso, + f_double_prime_iso=addend_fdp[iso_idx][rows_iso], + f_double_prime_aniso=addend_fdp[aniso_idx][rows_aniso], ) - self._anomalous_elements_hash = elements_hash - return self._anomalous_cache + self._anomalous_cache = (key, partition, terms) + return terms - def _apply_anomalous_correction( - self, - sf: torch.Tensor, - hkl: torch.Tensor, - include_fdp: bool = True, - ) -> torch.Tensor: - """Add ``ΔF(h) = Σ (f' + i f'') exp(2πi h·r) occ`` to ``sf``. - - Only the significant scatterers (|f'| or |f''| above - ``anomalous_threshold``) contribute. ``include_fdp=False`` zeroes f'', - keeping Friedel's law intact -- the correct choice for merged data. - """ - mask, f_prime, f_double_prime, has_anomalous, anomalous_indices = ( - self._get_anomalous_cache() - ) - - if not has_anomalous: - return sf # No significant anomalous scatterers + def _add_anomalous_scattering(self, iso, aniso, include_fdp: bool): + """Put f' and f'' into the atoms :meth:`forward` hands to the FFT engine. - # Integer indices, not the boolean mask: boolean indexing forces a GPU sync. - xyz_frac = self.xyz_fractional()[anomalous_indices] # (n_significant, 3) - occ = self.occupancy()[anomalous_indices] # (n_significant,) + Neither term depends on the scattering angle, so each is a zero-width Gaussian + in the atom's form factor -- the slot ITC92's constant ``c`` already occupies. + Placed there, both get exactly what f0 gets: the splat widens the term by the + atom's own isotropic or anisotropic displacement, and the FFT path expands it + over the symmetry operators. f' joins the real amplitudes. f'' becomes the only + amplitude of a copy of the anomalous atoms, which the engine splats into the + imaginary part of the density; each copy keeps its atom's ITC92 widths so both + parts are truncated at the same per-atom radius. - # Phase factors exp(2πi h·r), h·r over fractional coordinates - h_dot_r = torch.matmul( - hkl.to(dtype=self.dtype_float, device=xyz_frac.device), xyz_frac.T - ) # (n_refl, n_significant) - phase = 2 * torch.pi * h_dot_r - - cos_phase = torch.cos(phase) - sin_phase = torch.sin(phase) + Parameters + ---------- + iso, aniso : tuple of torch.Tensor + :meth:`get_iso` and :meth:`get_aniso` -- read from here rather than from the + parameter wrappers, so a subclass that adjusts those (``EnsembleModel``) + applies to the anomalous terms too. + include_fdp : bool + Build the f'' atoms. False keeps ``F(-h) = F(h)*``, which merged data need. - f_prime_occ = f_prime * occ # (n_significant,) - f_double_prime_occ = f_double_prime * occ # (n_significant,) + Returns + ------- + iso, aniso : tuple of torch.Tensor + The inputs with f' added to ``A``; returned as given without significant + scatterers. + imaginary : tuple of torch.Tensor or None + The f'' atoms in the layout of ``(*iso, *aniso)``, or None. + """ + terms = self._get_anomalous_cache() + if terms is None: + return iso, aniso, None + xyz_i, adp_i, occ_i, A_i, B_i = iso + xyz_a, u_a, occ_a, A_a, B_a = aniso + iso = (xyz_i, adp_i, occ_i, A_i + terms.f_prime_iso, B_i) + aniso = (xyz_a, u_a, occ_a, A_a + terms.f_prime_aniso, B_a) if not include_fdp: - # Dispersive f' only, so Friedel's law is preserved (merged data). - f_double_prime_occ = torch.zeros_like(f_double_prime_occ) - - # For each reflection: - # Real part: Σ [f'·cos(φ) - f''·sin(φ)] × occ - # Imag part: Σ [f'·sin(φ) + f''·cos(φ)] × occ - delta_real = torch.sum( - f_prime_occ * cos_phase - f_double_prime_occ * sin_phase, dim=-1 - ) - delta_imag = torch.sum( - f_prime_occ * sin_phase + f_double_prime_occ * cos_phase, dim=-1 + return iso, aniso, None + ri, ra = terms.rows_iso, terms.rows_aniso + imaginary = ( + xyz_i[ri], + adp_i[ri], + occ_i[ri], + terms.f_double_prime_iso, + B_i[ri], + xyz_a[ra], + u_a[ra], + occ_a[ra], + terms.f_double_prime_aniso, + B_a[ra], ) - - return sf + torch.complex(delta_real, delta_imag) + return iso, aniso, imaginary def get_structure_factor( self, hkl: torch.Tensor, recalc=False, apply_anomalous: bool = True @@ -597,8 +640,9 @@ def get_structure_factor( Notes ----- The full scattering factor is ``f(s, λ) = f₀(s) + f'(λ) + i f''(λ)``, - with f₀ from the FFT and the wavelength-dependent f' / f'' applied only - to atoms above ``anomalous_threshold``. + with the wavelength-dependent f' / f'' applied only to atoms above + ``anomalous_threshold``. All three terms go through the same density and + FFT, so each carries the atom's temperature factor and symmetry mates. """ return self(hkl, recalc=recalc, apply_anomalous=apply_anomalous) @@ -648,21 +692,24 @@ def forward(self, hkl, apply_anomalous: bool = True) -> torch.Tensor: ------- torch.Tensor Calculated complex structure factors with shape (n_reflections,). + + Notes + ----- + f' and f'' are not added to F afterwards: they are folded into the atoms' + form factors before the density is built, so the one splat and FFT apply + each atom's isotropic or anisotropic temperature factor and the space-group + symmetry to them exactly as to f0. With f'' the density is complex, and so is + the map left in ``self.ed``. """ self._check_forward_dtype(hkl) - sf, self.ed = self.fft.compute_structure_factors( - hkl, - *self.get_iso(), - *self.get_aniso(), - apply_symmetry=True, - ) - - # Apply anomalous correction as post-processing. f' always applies when a - # wavelength is set; f'' only for unmerged (Bijvoet) data. + iso, aniso, imaginary = self.get_iso(), self.get_aniso(), None if apply_anomalous and self.wavelength is not None: - sf = self._apply_anomalous_correction( - sf, hkl, include_fdp=bool(self.anomalous_bijvoet) + iso, aniso, imaginary = self._add_anomalous_scattering( + iso, aniso, include_fdp=bool(self.anomalous_bijvoet) ) + sf, self.ed = self.fft.compute_structure_factors( + hkl, *iso, *aniso, apply_symmetry=True, imaginary=imaginary + ) if self.ctx.verbose > 2: assert torch.all( diff --git a/torchref/model/sf_fft.py b/torchref/model/sf_fft.py index 81a76863..1c96fef2 100644 --- a/torchref/model/sf_fft.py +++ b/torchref/model/sf_fft.py @@ -387,6 +387,7 @@ def build_density_map( A_aniso: Optional[torch.Tensor] = None, B_aniso: Optional[torch.Tensor] = None, apply_symmetry: bool = True, + imaginary: Optional[Tuple[Optional[torch.Tensor], ...]] = None, ) -> torch.Tensor: """ Build electron density map from atomic parameters. @@ -394,28 +395,73 @@ def build_density_map( Parameters ---------- xyz_iso, adp_iso, occ_iso : torch.Tensor - Isotropic atoms: coordinates ``(n_iso, 3)``, ADPs ``(n_iso,)``, - occupancies ``(n_iso,)``. + Isotropic atoms: Cartesian coordinates ``(n_iso, 3)`` in Å, B-factors + ``(n_iso,)`` in Ų, occupancies ``(n_iso,)``. A_iso, B_iso : torch.Tensor - ITC92 amplitudes / widths for the isotropic atoms, ``(n_iso, 5)``. + ITC92 amplitudes (electrons) / widths (Ų) for the isotropic atoms, + ``(n_iso, 5)``. xyz_aniso, u_aniso, occ_aniso : torch.Tensor, optional - Anisotropic atoms: coordinates ``(n_aniso, 3)``, U components - ``(n_aniso, 6)``, occupancies ``(n_aniso,)``. + Anisotropic atoms: Cartesian coordinates ``(n_aniso, 3)`` in Å, U + components ``(n_aniso, 6)`` in Ų, occupancies ``(n_aniso,)``. A_aniso, B_aniso : torch.Tensor, optional ITC92 amplitudes / widths for the anisotropic atoms, ``(n_aniso, 5)``. apply_symmetry : bool, optional If True, apply crystallographic symmetry to the map. Default is True. + imaginary : tuple of torch.Tensor, optional + Atoms of the imaginary part of the density: ten tensors in the order, + shapes and units of ``xyz_iso`` ... ``B_aniso`` above. They are splatted + like the atoms above, into a second map that becomes the imaginary part. + This is how an absorptive f'' term enters the structure factors. Returns ------- torch.Tensor - Electron density map with shape (nx, ny, nz). + Electron density map with shape (nx, ny, nz); complex when + ``imaginary`` is given, real otherwise. """ self._require_grid() + density_map = self._splat( + xyz_iso, + adp_iso, + occ_iso, + A_iso, + B_iso, + xyz_aniso, + u_aniso, + occ_aniso, + A_aniso, + B_aniso, + ) + density_imag = None if imaginary is None else self._splat(*imaginary) + + if apply_symmetry: + symmetrize = self.ctx.spacegroup.symmetrize_map + density_map = symmetrize(density_map) + if density_imag is not None: + density_imag = symmetrize(density_imag) + + if density_imag is not None: + density_map = torch.complex(density_map, density_imag) + return density_map + + def _splat( + self, + xyz_iso, + adp_iso, + occ_iso, + A_iso, + B_iso, + xyz_aniso=None, + u_aniso=None, + occ_aniso=None, + A_aniso=None, + B_aniso=None, + ) -> torch.Tensor: + """Real P1 density of one set of atoms on this grid, ``(nx, ny, nz)``.""" from torchref.base.electron_density.main import build_electron_density - density_map = build_electron_density( + return build_electron_density( grid_shape=self.grid_shape, device=self.device, xyz_iso=xyz_iso, @@ -433,11 +479,6 @@ def build_density_map( dtype=self.dtype_float, ) - if apply_symmetry: - density_map = self.ctx.spacegroup.symmetrize_map(density_map) - - return density_map - # ========================================================================= # Structure Factor Methods # ========================================================================= @@ -454,7 +495,7 @@ def map_to_structure_factors( Parameters ---------- density_map : torch.Tensor - Electron density map with shape (nx, ny, nz). + Electron density map with shape (nx, ny, nz), real or complex. If apply_symmetry=True, this should be a P1 density map. hkl : torch.Tensor Miller indices with shape (n_reflections, 3). @@ -491,6 +532,7 @@ def compute_structure_factors( A_aniso: Optional[torch.Tensor] = None, B_aniso: Optional[torch.Tensor] = None, apply_symmetry: bool = True, + imaginary: Optional[Tuple[Optional[torch.Tensor], ...]] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: """ Compute structure factors from atomic parameters (end-to-end). @@ -516,13 +558,19 @@ def compute_structure_factors( ITC92 amplitudes / widths for the anisotropic atoms. apply_symmetry : bool, optional If True, apply crystallographic symmetry. Default is True. + imaginary : tuple of torch.Tensor, optional + Atoms of the imaginary part of the density, as in + :meth:`build_density_map`. The complex density is transformed in one + FFT, so ``F = FT(ρ') + i FT(ρ'')`` -- which no longer obeys + ``F(-h) = F(h)*`` -- with symmetry applied to both parts alike. Returns ------- sf : torch.Tensor Complex structure factors with shape (n_reflections,). density_map : torch.Tensor - Electron density map with shape (nx, ny, nz). + Electron density map with shape (nx, ny, nz), complex when + ``imaginary`` is given. Note: When using late symmetry, this is the P1 map (without symmetry). """ # Resolve the grid first: the late-symmetry flag belongs to the grid the @@ -544,6 +592,7 @@ def compute_structure_factors( A_aniso=A_aniso, B_aniso=B_aniso, apply_symmetry=not use_late and apply_symmetry, # Early symmetry + imaginary=imaginary, ) sf = self.map_to_structure_factors( density_map, From 2088c37c278a8f5920ee35c4cdbb55cd272e0848 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 22:29:55 +0000 Subject: [PATCH 247/250] Find every symmetry mate and score lattice images at their image distance The VDW symmetry pre-filter tried cell offsets -1..1 around the raw coordinates, so a model deposited away from the origin cell lost the mates that need a larger translation (6G9X: 0 of gemmi's 2270 symmetry contacts, 3E98: 1800 of 4638), and the count changed under a lattice shift of the model. prefilter_symop_offsets now returns every (symop, offset) whose centroid displacement is within 2r + cutoff, bounding the offsets per axis by that reach over the lattice-plane spacing. Scoring decided per pair list whether to apply the symmetry transform, looking at the operations only, so a list whose operations were all the identity (every P1 list) dropped its cell offsets and scored lattice contacts at the untranslated intra-ASU distance (5BOV: 24798 pairs at 38-105 A instead of 2-6 A, 512 clashes unscored). That gate was written three times: the eager kernel, the Triton wrapper and NonBondedTarget._compute_positions, with a fourth copy of the image formula in NonBondedHTarget. There is now one implementation, symmetry_image_positions in torchref.base.coordinates (B (R B^-1 x + t + n) for every entry, identity included), with is_symmetry_image as the one "image, not the atom itself" test. The pair builder forms the images it searches with it (prefilter and assign_to_grid, on unwrapped Cartesian coordinates), and nonbonded_pair_positions places stored pairs with it for the eager kernel, every NonBondedTarget mode and statistic, and the riding-H term, whose violations and statistics now see images too. The Triton wrapper follows the same rule: has_symmetry whenever symop tensors are passed. The duplicated gates and image formulas are deleted. Riding-H candidates: for an image pair the builder also emitted "H on B against heavy A" with no operation, scoring it at the ASU distance (on 1DAW 3799 candidates up to 62 A away). The reverse pair already gives that contact with the image on the right atom, so it is now emitted for intra-ASU pairs only. The candidate dedup key packed offsets as -1..1 and the operation into a 1000 stride; it is now row-wise. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- tests/unit/base/test_nonbonded_kernel.py | 189 +++++++++++++++ tests/unit/base/test_symmetry_images.py | 119 ++++++++++ tests/unit/topology/test_vdw_pair_search.py | 8 +- .../topology/test_vdw_symmetry_contacts.py | 219 ++++++++++++++++++ torchref/base/coordinates/__init__.py | 9 + torchref/base/coordinates/symmetry_images.py | 96 ++++++++ torchref/base/targets/nonbonded.py | 115 +++++++-- torchref/base/targets/triton/nonbonded.py | 23 +- .../refinement/targets/geometry/non_bonded.py | 78 +++---- .../targets/geometry/non_bonded_h.py | 133 +++++------ torchref/topology/nonbonded.py | 152 ++++++------ torchref/topology/restraints.py | 7 +- torchref/topology/riding.py | 64 ++--- 13 files changed, 960 insertions(+), 252 deletions(-) create mode 100644 tests/unit/base/test_nonbonded_kernel.py create mode 100644 tests/unit/base/test_symmetry_images.py create mode 100644 tests/unit/topology/test_vdw_symmetry_contacts.py create mode 100644 torchref/base/coordinates/symmetry_images.py diff --git a/tests/unit/base/test_nonbonded_kernel.py b/tests/unit/base/test_nonbonded_kernel.py new file mode 100644 index 00000000..dded8431 --- /dev/null +++ b/tests/unit/base/test_nonbonded_kernel.py @@ -0,0 +1,189 @@ +"""Non-bonded scoring of image pairs: pair positions, the eager kernel, and Triton. + +A pair ``(i, j, symop, offset)`` is scored at the distance from atom ``i`` to the +``(symop, offset)`` image of atom ``j``, whatever the operation. In P1 every crystal +contact is a pure lattice translation under symop 0, so 5BOV (P 1) checks that the +offset is applied when no pair has another operation, and 1DAW (C 1 2 1) checks a list +mixing operations. Distances are compared with gemmi's own orthogonalization. The +CUDA-only Triton comparison skips without a CUDA device. +""" + +import math + +import gemmi +import numpy as np +import pytest +import torch + +from torchref.base.coordinates import is_symmetry_image +from torchref.base.targets._dispatch import use_triton +from torchref.base.targets.nonbonded import ( + _nonbonded_heavy_math_eager, + nonbonded_heavy_math, + nonbonded_pair_positions, +) +from torchref.config import get_float_dtype +from torchref.symmetry import SpaceGroup +from torchref.symmetry.cell import Cell +from torchref.topology import nonbonded as nb + +pytestmark = pytest.mark.unit + +CUTOFF = 6.0 +# The production NonBondedTarget defaults: sigma 0.3 Å, r_exp 4, no buffer. +_SIGMA, _R_EXP = 0.3, 4.0 +_C_REP = 1.0 / (_R_EXP * _SIGMA**_R_EXP) + + +def _image_pairs(path): + """Image pairs of one deposited model from builder steps 1-4, with radius sums.""" + st = gemmi.read_structure(str(path)) + atoms = [a for ch in st[0] for r in ch for a in r] + xyz = torch.tensor( + [[a.pos.x, a.pos.y, a.pos.z] for a in atoms], dtype=get_float_dtype() + ) + c = st.cell + cell = Cell([c.a, c.b, c.c, c.alpha, c.beta, c.gamma]) + sg = SpaceGroup(st.spacegroup_hm) + ops, offsets = nb.prefilter_symop_offsets(cell, sg, xyz, CUTOFF) + identity = (~is_symmetry_image(ops, offsets)).nonzero()[0].item() + grid_dims = torch.clamp( + (torch.stack([cell.a, cell.b, cell.c]) / CUTOFF).long(), min=1 + ) + _, atom_idx, combo_idx, cart = nb.assign_to_grid( + xyz, cell, sg, ops, offsets, grid_dims + ) + i, j, combo = nb.find_pairs_kdtree(cart, atom_idx, combo_idx, CUTOFF, identity) + image = combo != identity + i, j, combo = i[image], j[image], combo[image] + radii = torch.as_tensor(nb.vdw_radii_for_elements([a.element.name for a in atoms])) + return { + "st": st, + "xyz": xyz, + "indices": torch.stack([i, j], dim=1), + "symop_indices": ops[combo], + "cell_offsets": offsets[combo], + "min_distances": (radii[i] + radii[j]).to(get_float_dtype()), + "tables": ( + sg.matrices, + sg.translations, + cell.fractional_matrix, + cell.inv_fractional_matrix, + ), + } + + +def _gemmi_distances(pairs): + """Image distance of every pair through gemmi's own operations and cell, float64. + + TorchRef's :class:`SpaceGroup` takes its operations in gemmi's order, so a stored + ``symop`` indexes gemmi's list directly. + """ + st = pairs["st"] + atoms = [a for ch in st[0] for r in ch for a in r] + ops = list(st.find_spacegroup().operations()) + out = [] + for (i, j), s, n in zip( + pairs["indices"].tolist(), + pairs["symop_indices"].tolist(), + pairs["cell_offsets"].tolist(), + ): + f = st.cell.fractionalize(atoms[j].pos) + x, y, z = ops[s].apply_to_xyz([f.x, f.y, f.z]) + image = gemmi.Fractional(x + n[0], y + n[1], z + n[2]) + out.append(atoms[i].pos.dist(st.cell.orthogonalize(image))) + return np.array(out) + + +@pytest.fixture(scope="module", params=["5BOV.pdb", "1DAW.pdb"]) +def pairs(request, pdb_dir): + out = _image_pairs(pdb_dir / request.param) + out["name"] = request.param + return out + + +def _args(pairs): + return ( + pairs["xyz"], + pairs["indices"], + pairs["symop_indices"], + pairs["cell_offsets"], + *pairs["tables"], + ) + + +def test_image_pairs_are_placed_at_gemmi_distances(pairs): + pos1, pos2 = nonbonded_pair_positions(*_args(pairs)) + got = (pos2 - pos1).norm(dim=1).double().numpy() + want = _gemmi_distances(pairs) + assert np.abs(got - want).max() < 1e-3 + assert want.max() < CUTOFF + if pairs["name"] == "5BOV.pdb": + assert bool((pairs["symop_indices"] == 0).all()), "P1: lattice images only" + + +def test_loss_scores_every_image_pair(pairs): + xyz = pairs["xyz"] + one = torch.ones((), dtype=xyz.dtype) + loss = nonbonded_heavy_math( + xyz, + pairs["indices"], + pairs["min_distances"], + pairs["symop_indices"], + pairs["cell_offsets"], + *pairs["tables"], + _C_REP * one, + _R_EXP * one, + 0.0, + _SIGMA * one, + ) + # The kernel's sqrt epsilon matters here: 1DAW has a water on the 2-fold axis + # whose image under that axis is itself, at distance 0. + distance = np.sqrt(_gemmi_distances(pairs) ** 2 + 1e-8) + overlap = np.clip(pairs["min_distances"].double().numpy() - distance, 0.0, None) + assert np.count_nonzero(overlap) > 0, "the fixture must contain clashing mates" + n = len(overlap) + want = _C_REP * (overlap**_R_EXP).sum() + n * ( + math.log(_SIGMA) + 0.5 * math.log(2.0 * math.pi) + ) + assert float(loss) == pytest.approx(want, rel=1e-5, abs=1e-2) + + +@pytest.mark.cuda +@pytest.mark.parametrize("pdb", ["5BOV.pdb", "3E98.pdb"]) +def test_triton_matches_eager_on_image_pairs(pdb, pdb_dir): + """The Triton kernel images every pair as the eager path does, offsets included. + + 5BOV (P 1) holds lattice images only, all under symop 0; 3E98 (P 1 21 1) mixes + them with screw-axis images. Neither has an atom on a special position, whose + self-image at distance ~0 would make its gradient direction float32 noise. + """ + from tests.helpers.grad_asserts import assert_grads_agree + + host = _image_pairs(pdb_dir / pdb) + dev = torch.device("cuda") + # float32 throughout, whatever the configured dtype: that is the Triton contract. + xyz = host["xyz"].to(dev, torch.float32) + min_distances = host["min_distances"].to(dev, torch.float32) + tables = [t.to(dev, torch.float32) for t in host["tables"]] + indices = host["indices"].to(dev) + symop_indices = host["symop_indices"].to(dev) + cell_offsets = host["cell_offsets"].to(dev) + one = torch.ones((), device=dev, dtype=torch.float32) + scalars = (_C_REP * one, _R_EXP * one, 0.0, _SIGMA * one) + assert use_triton(xyz), "the Triton arm would compare eager against eager" + + def run(fn, offsets): + x = xyz.clone().requires_grad_(True) + loss = fn(x, indices, min_distances, symop_indices, offsets, *tables, *scalars) + (grad,) = torch.autograd.grad(loss, x) + return loss.detach(), grad + + loss_t, grad_t = run(nonbonded_heavy_math, cell_offsets) + loss_e, grad_e = run(_nonbonded_heavy_math_eager, cell_offsets) + loss_unshifted, _ = run(_nonbonded_heavy_math_eager, torch.zeros_like(cell_offsets)) + + # Non-vacuity: dropping the offsets must move the loss far past the tolerance. + assert abs(float(loss_unshifted - loss_e)) > 1.0 + torch.testing.assert_close(loss_t, loss_e, rtol=1e-5, atol=1e-2) + assert_grads_agree([grad_t], [grad_e], min_cos=0.9999, ratio_tol=1e-3, ctx="vdw ") diff --git a/tests/unit/base/test_symmetry_images.py b/tests/unit/base/test_symmetry_images.py new file mode 100644 index 00000000..c9288e9e --- /dev/null +++ b/tests/unit/base/test_symmetry_images.py @@ -0,0 +1,119 @@ +"""``symmetry_image_positions`` and ``is_symmetry_image``, the one image construction. + +The non-bonded pair builder forms the images it searches with this function and the +non-bonded scoring places the stored pairs with it. These pin the contract both rely +on: ``B (R B^-1 x + t + n)`` for every entry, the identity included, with a pure +lattice translation kept as a real image. +""" + +import gemmi +import pytest +import torch + +from torchref.base.coordinates import is_symmetry_image, symmetry_image_positions +from torchref.config import get_float_dtype, get_int_dtype +from torchref.symmetry import SpaceGroup +from torchref.symmetry.cell import Cell + +pytestmark = pytest.mark.unit + +# Tolerance on a position in Å. The coordinates reach ~100 Å, so float32 rounding of +# the fractional round trip is ~1e-5 Å. +_ATOL = 1e-4 + + +@pytest.fixture(scope="module") +def crystal(pdb_dir): + """1DAW (C 1 2 1, monoclinic, so B is not diagonal): coordinates, cell, group.""" + st = gemmi.read_structure(str(pdb_dir / "1DAW.pdb")) + xyz = torch.tensor( + [[a.pos.x, a.pos.y, a.pos.z] for ch in st[0] for r in ch for a in r][:200], + dtype=get_float_dtype(), + ) + c = st.cell + cell = Cell([c.a, c.b, c.c, c.alpha, c.beta, c.gamma]) + sg = SpaceGroup(st.spacegroup_hm) + tables = ( + sg.matrices, + sg.translations, + cell.fractional_matrix, + cell.inv_fractional_matrix, + ) + return xyz, cell, sg, tables + + +def _per_point(n_points, op, offset): + ops = torch.full((n_points,), op, dtype=get_int_dtype()) + offsets = torch.tensor(offset, dtype=get_int_dtype()).expand(n_points, 3) + return ops, offsets + + +def test_matches_the_space_group_expansion(crystal): + xyz, cell, sg, tables = crystal + expanded = sg.expand_positions(cell.cartesian_to_fractional(xyz)) + for op in range(sg.n_ops): + for offset in ([0, 0, 0], [1, 0, -2], [-3, 2, 1]): + ops, offsets = _per_point(len(xyz), op, offset) + got = symmetry_image_positions(xyz, ops, offsets, *tables) + shift = torch.tensor(offset, dtype=xyz.dtype) + want = cell.fractional_to_cartesian(expanded[op] + shift) + torch.testing.assert_close(got, want, atol=_ATOL, rtol=0) + + +def test_a_pure_lattice_translation_is_an_image(crystal): + xyz, cell, _, tables = crystal + ops, offsets = _per_point(len(xyz), 0, [2, -1, 1]) + got = symmetry_image_positions(xyz, ops, offsets, *tables) + shift = cell.fractional_to_cartesian(offsets[0].to(xyz.dtype)) + torch.testing.assert_close(got, xyz + shift, atol=_ATOL, rtol=0) + assert bool(is_symmetry_image(ops, offsets).all()) + + +def test_the_identity_returns_the_point(crystal): + xyz, _, _, tables = crystal + ops, offsets = _per_point(len(xyz), 0, [0, 0, 0]) + got = symmetry_image_positions(xyz, ops, offsets, *tables) + torch.testing.assert_close(got, xyz, atol=_ATOL, rtol=0) + assert not bool(is_symmetry_image(ops, offsets).any()) + torch.testing.assert_close( + symmetry_image_positions(xyz, ops, None, *tables), got, atol=0, rtol=0 + ) + + +def test_is_symmetry_image_looks_at_operation_and_offset(): + ops = torch.tensor([0, 0, 1, 1], dtype=get_int_dtype()) + offsets = torch.tensor( + [[0, 0, 0], [0, -1, 0], [0, 0, 0], [1, 0, 0]], dtype=get_int_dtype() + ) + assert is_symmetry_image(ops, offsets).tolist() == [False, True, True, True] + + +def test_broadcasting_matches_one_entry_per_point(crystal): + xyz, _, sg, tables = crystal + ops = torch.tensor([0, 1, 0, sg.n_ops - 1], dtype=get_int_dtype()) + offsets = torch.tensor( + [[0, 0, 0], [0, 0, 1], [-1, 2, 0], [1, 1, 1]], dtype=get_int_dtype() + ) + grid = symmetry_image_positions(xyz[:, None, :], ops, offsets, *tables) + assert grid.shape == (len(xyz), len(ops), 3) + flat = symmetry_image_positions( + xyz.repeat_interleave(len(ops), dim=0), + ops.repeat(len(xyz)), + offsets.repeat(len(xyz), 1), + *tables, + ) + torch.testing.assert_close(grid.reshape(-1, 3), flat, atol=_ATOL, rtol=0) + + +def test_gradient_is_the_cartesian_rotation(crystal): + xyz, cell, sg, tables = crystal + op = 1 + ops, offsets = _per_point(1, op, [1, 0, -1]) + + def image(x): + return symmetry_image_positions(x[None, :], ops, offsets, *tables)[0] + + jacobian = torch.autograd.functional.jacobian(image, xyz[0].clone()) + B, B_inv = cell.fractional_matrix, cell.inv_fractional_matrix + want = B @ sg.matrices[op].to(B.dtype) @ B_inv + torch.testing.assert_close(jacobian, want, atol=1e-5, rtol=0) diff --git a/tests/unit/topology/test_vdw_pair_search.py b/tests/unit/topology/test_vdw_pair_search.py index 35370679..a448a240 100644 --- a/tests/unit/topology/test_vdw_pair_search.py +++ b/tests/unit/topology/test_vdw_pair_search.py @@ -32,13 +32,13 @@ def _image_table(path): model = Model(verbose=0, device=torch.device("cpu")) model.load_pdb(str(path)) cell, sg = model.ctx.cell, model.ctx.spacegroup - xyz_frac = cell.cartesian_to_fractional(model.xyz().detach().to(dtypes.float)) - op_indices, offsets = nb.prefilter_symop_offsets(cell, sg, xyz_frac, CUTOFF) + xyz = model.xyz().detach().to(dtypes.float) + op_indices, offsets = nb.prefilter_symop_offsets(cell, sg, xyz, CUTOFF) identity = ((op_indices == 0) & (offsets == 0).all(dim=1)).nonzero()[0].item() lengths = torch.stack([cell.a, cell.b, cell.c]).to(dtypes.float) grid_dims = torch.clamp((lengths / CUTOFF).long(), min=1) flat_cell, atom_idx, combo_idx, cart_pos = nb.assign_to_grid( - xyz_frac, cell, sg, op_indices, offsets, grid_dims + xyz, cell, sg, op_indices, offsets, grid_dims ) # Perpendicular width of one grid cell along each axis: lattice-plane spacing # V / |face| over the number of cells. @@ -50,7 +50,7 @@ def _image_table(path): for k, (u, v) in enumerate(faces) ] return dict( - n_atoms=xyz_frac.shape[0], n_combos=len(op_indices), identity=identity, + n_atoms=xyz.shape[0], n_combos=len(op_indices), identity=identity, grid_dims=grid_dims, flat_cell=flat_cell, atom_idx=atom_idx, combo_idx=combo_idx, cart_pos=cart_pos, min_width=min(widths), ) diff --git a/tests/unit/topology/test_vdw_symmetry_contacts.py b/tests/unit/topology/test_vdw_symmetry_contacts.py new file mode 100644 index 00000000..b53ab236 --- /dev/null +++ b/tests/unit/topology/test_vdw_symmetry_contacts.py @@ -0,0 +1,219 @@ +"""Symmetry contacts of the VDW pair builder, against gemmi and under lattice shifts. + +Which unit cell the deposited coordinates sit in is arbitrary, so the image pairs the +builder finds must not depend on it, and they must be every contact gemmi finds with +the full space group. 6G9X (P 21 21 2, centroid at fractional x = 1.08) and 3E98 +(P 1 21 1, centroid at z = 1.04) are deposited away from the origin cell; 1DAW (C 1 2 1) +runs the production builder and the riding-hydrogen candidates. +""" + +from collections import defaultdict + +import gemmi +import numpy as np +import pytest +import torch + +from torchref.base.coordinates import is_symmetry_image, symmetry_image_positions +from torchref.base.targets.nonbonded import nonbonded_pair_positions +from torchref.config import get_float_dtype +from torchref.model.model import Model +from torchref.symmetry import SpaceGroup +from torchref.symmetry.cell import Cell +from torchref.topology import nonbonded as nb +from torchref.topology.riding import place_riding_hydrogens + +pytestmark = pytest.mark.unit + +CUTOFF = 6.0 +# float32 positions near 100 Å against gemmi's float64 ones. +_DIST_ATOL = 1e-3 +_SHIFTS = ([-1, 0, 0], [0, 0, -1], [2, -3, 1], [5, 4, -6]) + + +def _read(path): + """Coordinates in gemmi's atom order, the cell and the space group of one model.""" + st = gemmi.read_structure(str(path)) + atoms = [ + (ch.name, r.seqid.num, r.seqid.icode, a.name, a.altloc) + for ch in st[0] + for r in ch + for a in r + ] + xyz = torch.tensor( + [[a.pos.x, a.pos.y, a.pos.z] for ch in st[0] for r in ch for a in r], + dtype=get_float_dtype(), + ) + c = st.cell + cell = Cell([c.a, c.b, c.c, c.alpha, c.beta, c.gamma]) + return ( + st, + {key: k for k, key in enumerate(atoms)}, + xyz, + cell, + SpaceGroup(st.spacegroup_hm), + ) + + +def _image_contacts(xyz, cell, sg): + """Builder steps 1-4 on CPU: ``{(i, j): sorted image distances}`` for image pairs.""" + ops, offsets = nb.prefilter_symop_offsets(cell, sg, xyz, CUTOFF) + identity = (~is_symmetry_image(ops, offsets)).nonzero()[0].item() + lengths = torch.stack([cell.a, cell.b, cell.c]) + grid_dims = torch.clamp((lengths / CUTOFF).long(), min=1) + _, atom_idx, combo_idx, cart = nb.assign_to_grid( + xyz, cell, sg, ops, offsets, grid_dims + ) + i, j, combo = nb.find_pairs_kdtree(cart, atom_idx, combo_idx, CUTOFF, identity) + image = combo != identity + i, j, combo = i[image], j[image], combo[image] + partner = symmetry_image_positions( + xyz[j], + ops[combo], + offsets[combo], + sg.matrices, + sg.translations, + cell.fractional_matrix, + cell.inv_fractional_matrix, + ) + return _group(i, j, (partner - xyz[i]).norm(dim=1)) + + +def _group(i, j, d): + out = defaultdict(list) + for a, b, x in zip(i.tolist(), j.tolist(), d.tolist()): + out[(a, b)].append(x) + return {k: sorted(v) for k, v in out.items()} + + +def _gemmi_contacts(st, index): + """gemmi's symmetry and lattice contacts below ``CUTOFF``, from both ends.""" + ns = gemmi.NeighborSearch(st[0], st.cell, 5).populate() + search = gemmi.ContactSearch(CUTOFF) + search.ignore = gemmi.ContactSearch.Ignore.Nothing + search.twice = True + search.special_pos_cutoff_sq = 0.0 + out = defaultdict(list) + for hit in search.find_contacts(ns): + p1, p2 = hit.partner1, hit.partner2 + image = st.cell.find_nearest_pbc_image(p1.atom.pos, p2.atom.pos, hit.image_idx) + if image.same_asu(): + continue + keys = [ + ( + p.chain.name, + p.residue.seqid.num, + p.residue.seqid.icode, + p.atom.name, + p.atom.altloc, + ) + for p in (p1, p2) + ] + out[(index[keys[0]], index[keys[1]])].append(hit.dist) + return {k: sorted(v) for k, v in out.items()} + + +def _assert_same_contacts(got, want): + assert got.keys() == want.keys(), ( + f"{len(set(want) - set(got))} contacts missing, " + f"{len(set(got) - set(want))} extra" + ) + for key, distances in want.items(): + assert len(got[key]) == len(distances), key + assert np.allclose(got[key], distances, atol=_DIST_ATOL, rtol=0), key + + +@pytest.fixture(scope="module", params=["6G9X.pdb", "3E98.pdb"]) +def deposited(request, pdb_dir): + st, index, xyz, cell, sg = _read(pdb_dir / request.param) + return st, index, xyz, cell, sg, _image_contacts(xyz, cell, sg) + + +def test_image_contacts_match_gemmi(deposited): + st, index, _, _, _, contacts = deposited + _assert_same_contacts(contacts, _gemmi_contacts(st, index)) + + +@pytest.mark.parametrize("shift", _SHIFTS) +def test_image_contacts_do_not_depend_on_the_unit_cell(deposited, shift): + _, _, xyz, cell, sg, contacts = deposited + translation = cell.fractional_to_cartesian(torch.tensor(shift, dtype=xyz.dtype)) + _assert_same_contacts(_image_contacts(xyz + translation, cell, sg), contacts) + + +def test_every_image_contact_is_listed_from_both_ends(deposited): + contacts = deposited[-1] + reverse = {(j, i): d for (i, j), d in contacts.items()} + _assert_same_contacts(reverse, contacts) + + +@pytest.fixture(scope="module") +def model_1daw(pdb_dir): + model = Model(verbose=0, device=torch.device("cpu")) + model.load_pdb(str(pdb_dir / "1DAW.pdb")) + return model + + +def test_production_builder_keeps_its_contacts_under_a_lattice_shift(model_1daw): + restraints = model_1daw.restraints + cell, sg = model_1daw.cell, model_1daw.spacegroup + xyz = model_1daw.xyz().detach() + + def image_distances(coords): + vdw = nb.build_vdw_restraints_gpu( + xyz=coords, + vdw_radii=restraints._vdw_radii, + cell=cell, + sg=sg, + topology=restraints.topology, + exclusion_set=restraints.topology.atoms.exclusions_from_restraint_edges(), + cutoff=CUTOFF, + inter_residue_only=False, + ) + tables = ( + sg.matrices, + sg.translations, + cell.fractional_matrix, + cell.inv_fractional_matrix, + ) + pos1, pos2 = nonbonded_pair_positions( + coords, + vdw["indices"], + vdw["symop_indices"], + vdw["cell_offsets"], + *tables, + ) + image = is_symmetry_image(vdw["symop_indices"], vdw["cell_offsets"]) + idx = vdw["indices"][image] + return _group(idx[:, 0], idx[:, 1], (pos2 - pos1).norm(dim=1)[image]) + + deposited = image_distances(xyz) + assert sum(len(d) for d in deposited.values()) > 5000 + shift = torch.tensor([-1.0, 0.0, 0.0], dtype=xyz.dtype) + _assert_same_contacts( + image_distances(xyz + cell.fractional_to_cartesian(shift)), deposited + ) + + +def test_riding_h_candidates_are_scored_near_their_heavy_contact(model_1daw): + """An H candidate comes from a heavy pair closer than the cutoff, so with the image + on the right atom it lies within the cutoff plus two X-H bonds.""" + h_topo = model_1daw.restraints.h_topo + assert h_topo is not None and h_topo.has_candidates + cell, sg = model_1daw.cell, model_1daw.spacegroup + xyz = model_1daw.xyz().detach() + xyz_all = torch.cat([xyz, place_riding_hydrogens(xyz, h_topo)]) + pos_i, pos_j = nonbonded_pair_positions( + xyz_all, + torch.stack([h_topo.cand_idx_i, h_topo.cand_idx_j], dim=1), + h_topo.cand_symop_idx, + h_topo.cand_cell_offset, + sg.matrices, + sg.translations, + cell.fractional_matrix, + cell.inv_fractional_matrix, + ) + image = is_symmetry_image(h_topo.cand_symop_idx, h_topo.cand_cell_offset) + assert bool(image.any()) + reach = CUTOFF + 2.0 * float(h_topo.h_bond_length.max()) + _DIST_ATOL + assert float((pos_j - pos_i).norm(dim=1).max()) < reach diff --git a/torchref/base/coordinates/__init__.py b/torchref/base/coordinates/__init__.py index a4ab6d41..3c983d1e 100644 --- a/torchref/base/coordinates/__init__.py +++ b/torchref/base/coordinates/__init__.py @@ -6,6 +6,7 @@ - Cartesian <-> fractional coordinate conversions - Periodic boundary condition handling - Transformation matrix computations +- Symmetry images: a position under a space-group operation and lattice translation All of them are PyTorch functions; the cell metric is written out once, in :func:`get_fractional_matrix`. @@ -23,6 +24,11 @@ smallest_diff_aniso, ) +from .symmetry_images import ( + is_symmetry_image, + symmetry_image_positions, +) + from .local_frame import ( frame_is_degenerate, local_frame_axes, @@ -39,6 +45,9 @@ # Periodic boundary "smallest_diff", "smallest_diff_aniso", + # Symmetry images (non-bonded pair building and scoring) + "symmetry_image_positions", + "is_symmetry_image", # Local frames (riding hydrogens) "local_frame_axes", "place_local_frame", diff --git a/torchref/base/coordinates/symmetry_images.py b/torchref/base/coordinates/symmetry_images.py new file mode 100644 index 00000000..cdb5aa05 --- /dev/null +++ b/torchref/base/coordinates/symmetry_images.py @@ -0,0 +1,96 @@ +"""Positions of atoms under a symmetry operation followed by a lattice translation. + +An image is named by a pair ``(symop, cell_offset)``: operation ``symop`` of the space +group, ``x -> R x + t`` in fractional coordinates, then the integer lattice translation +``cell_offset``. :func:`symmetry_image_positions` is the one place such a pair becomes a +Cartesian position, and :func:`is_symmetry_image` the one test for "this is an image, +not the atom itself". The non-bonded pair builder (:mod:`torchref.topology.nonbonded`) +forms the images it searches with the first, and the non-bonded scoring +(:func:`torchref.base.targets.nonbonded.nonbonded_pair_positions`) places the stored +pairs with it, so a pair is scored at the distance it was found at. + +Offsets apply to the coordinates as given; nothing here wraps into the unit cell. A +stored ``cell_offset`` therefore stays valid only together with the unwrapped +coordinates it was chosen for. +""" + +import torch + + +def symmetry_image_positions( + xyz: torch.Tensor, + symop_indices: torch.Tensor, + cell_offsets: torch.Tensor | None, + symop_matrices: torch.Tensor, + symop_translations: torch.Tensor, + fractional_matrix: torch.Tensor, + inv_fractional_matrix: torch.Tensor, +) -> torch.Tensor: + """Place each point's image: ``B (R_s B^-1 x + t_s + n)``. + + Parameters + ---------- + xyz : torch.Tensor + Cartesian source positions in Å, shape ``(..., 3)``. + symop_indices : torch.Tensor + Integer index ``s`` into the operation table per point, any shape that + broadcasts against ``xyz[..., 0]``; 0 is the identity. + cell_offsets : torch.Tensor or None + Integer lattice translation ``n`` per point, in fractional units, any shape + that broadcasts against ``xyz``. None means no translation. + symop_matrices : torch.Tensor + Rotation part of every operation in the fractional basis, ``(n_ops, 3, 3)``. + symop_translations : torch.Tensor + Fractional translation part of every operation, ``(n_ops, 3)``. + fractional_matrix : torch.Tensor + Orthogonalization matrix ``B`` (fractional to Cartesian), ``(3, 3)``, as + :attr:`torchref.symmetry.Cell.fractional_matrix`. + inv_fractional_matrix : torch.Tensor + Its inverse, ``(3, 3)``, as :attr:`torchref.symmetry.Cell.inv_fractional_matrix`. + + Returns + ------- + torch.Tensor + Cartesian image positions in Å, of the broadcast shape ``(..., 3)`` and in + ``xyz``'s dtype. Differentiable in ``xyz`` and in both cell matrices. + + Notes + ----- + Every point is transformed, identity entries included: ``(0, 0)`` returns ``xyz`` + up to the rounding of the ``B^-1``/``B`` round trip. There is deliberately no + shortcut for lists whose operations are all the identity, because such a list can + still carry lattice translations -- in P1 every crystal contact is one. + """ + dtype = xyz.dtype + frac = xyz @ inv_fractional_matrix.to(dtype).T + rot = symop_matrices.to(dtype)[symop_indices] + # einsum rather than a broadcast matmul: with xyz (N, 1, 3) against M operations, + # matmul copies one 3x3 matrix per output point, einsum only writes the output. + frac_image = torch.einsum("...ij,...j->...i", rot, frac) + frac_image = frac_image + symop_translations.to(dtype)[symop_indices] + if cell_offsets is not None: + frac_image = frac_image + cell_offsets.to(dtype) + return frac_image @ fractional_matrix.to(dtype).T + + +def is_symmetry_image( + symop_indices: torch.Tensor, cell_offsets: torch.Tensor +) -> torch.Tensor: + """Whether each ``(symop, cell_offset)`` names an image rather than the atom itself. + + True wherever the operation is not the identity *or* the lattice translation is + nonzero: a pure lattice translation is a real image, and the only kind P1 has. + + Parameters + ---------- + symop_indices : torch.Tensor + Integer operation index per entry, shape ``(...)``; 0 is the identity. + cell_offsets : torch.Tensor + Integer lattice translation per entry, shape ``(..., 3)``. + + Returns + ------- + torch.Tensor + Boolean mask of shape ``symop_indices.shape``. + """ + return (symop_indices != 0) | (cell_offsets != 0).any(dim=-1) diff --git a/torchref/base/targets/nonbonded.py b/torchref/base/targets/nonbonded.py index 76c08cc3..838c0be1 100644 --- a/torchref/base/targets/nonbonded.py +++ b/torchref/base/targets/nonbonded.py @@ -1,13 +1,85 @@ -"""Non-bonded (VDW) heavy-heavy repulsion NLL — prolsq mode with symmetry mates.""" +"""Non-bonded (VDW) heavy-heavy repulsion NLL — prolsq mode with symmetry mates. + +:func:`nonbonded_pair_positions` places both ends of every pair. The eager kernel +here, the inline modes and statistics of ``NonBondedTarget`` and the riding-hydrogen +term of ``NonBondedHTarget`` all read their positions from it; the Triton kernel +(:mod:`torchref.base.targets.triton.nonbonded`) computes the same positions in-kernel. +""" from typing import Optional import torch +from torchref.base.coordinates.symmetry_images import symmetry_image_positions + from ._common import LOG_2PI from ._dispatch import use_triton +def nonbonded_pair_positions( + xyz: torch.Tensor, + indices: torch.Tensor, + symop_indices: torch.Tensor | None, + cell_offsets: torch.Tensor | None, + symop_matrices: torch.Tensor | None, + symop_translations: torch.Tensor | None, + fractional_matrix: torch.Tensor | None, + inv_fractional_matrix: torch.Tensor | None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Cartesian positions of both ends of every non-bonded pair. + + The first atom of a pair is the asymmetric-unit atom itself; the second is the + image ``(symop, cell_offset)`` of an ASU atom, placed by + :func:`~torchref.base.coordinates.symmetry_image_positions` from the current + coordinates, so gradients reach both atoms. + + Parameters + ---------- + xyz : torch.Tensor + Cartesian ASU coordinates in Å, shape ``(N_atoms, 3)``. + indices : torch.Tensor + Atom indices per pair, ``(N, 2)`` integer. + symop_indices : torch.Tensor or None + Operation index per pair, ``(N,)`` integer; 0 is the identity. None means + every partner is the ASU atom itself, and the four symmetry arguments below + are then not read. + cell_offsets : torch.Tensor or None + Integer lattice translation per pair, ``(N, 3)``, in fractional units, + relative to ``xyz`` as given (unwrapped). None means no translation. + symop_matrices, symop_translations : torch.Tensor or None + The operation table: ``(n_ops, 3, 3)`` rotations and ``(n_ops, 3)`` + translations, in the fractional basis. + fractional_matrix, inv_fractional_matrix : torch.Tensor or None + ``Cell.fractional_matrix`` and its inverse, ``(3, 3)``. + + Returns + ------- + pos1, pos2 : torch.Tensor + Each ``(N, 3)`` in Å, Cartesian. + + Notes + ----- + With ``symop_indices`` given, every pair is transformed, intra-ASU ones included + (they come out as the identity). There is no shortcut for lists whose operations + are all 0: those still carry lattice translations, and in P1 they are all of the + crystal contacts. + """ + pos1 = xyz[indices[:, 0]] + partner = xyz[indices[:, 1]] + if symop_indices is None: + return pos1, partner + pos2 = symmetry_image_positions( + partner, + symop_indices, + cell_offsets, + symop_matrices, + symop_translations, + fractional_matrix, + inv_fractional_matrix, + ) + return pos1, pos2 + + def _nonbonded_heavy_math_eager( xyz: torch.Tensor, indices: torch.Tensor, @@ -23,25 +95,16 @@ def _nonbonded_heavy_math_eager( buffer: float, sigma_vdw: torch.Tensor, ) -> torch.Tensor: - pos1 = xyz[indices[:, 0]] - - has_symmetry = ( - symop_indices is not None - and symop_indices.numel() > 0 - and not bool((symop_indices == 0).all()) + pos1, pos2 = nonbonded_pair_positions( + xyz, + indices, + symop_indices, + cell_offsets, + symop_matrices, + symop_translations, + fractional_matrix, + inv_fractional_matrix, ) - - if not has_symmetry: - pos2 = xyz[indices[:, 1]] - else: - mate_source = xyz[indices[:, 1]] - frac = mate_source @ inv_fractional_matrix.T - R = symop_matrices[symop_indices].to(frac.dtype) - t = symop_translations[symop_indices].to(frac.dtype) - offsets = cell_offsets.to(frac.dtype) - frac_transformed = torch.bmm(R, frac.unsqueeze(-1)).squeeze(-1) + t + offsets - pos2 = frac_transformed @ fractional_matrix.T - diff = pos2 - pos1 actual_distances = torch.sqrt((diff ** 2).sum(dim=-1) + 1e-8) violations = torch.clamp(min_distances + buffer - actual_distances, min=0.0) @@ -67,8 +130,8 @@ def nonbonded_heavy_math( ) -> torch.Tensor: """Heavy-heavy VDW prolsq repulsion NLL. - Matches the prolsq branch of ``NonBondedTarget.forward`` plus the symmetry-aware - gather of ``NonBondedTarget._compute_positions``. The H-VDW term ``NonBondedHTarget`` + Matches the prolsq branch of ``NonBondedTarget.forward``, with pair positions + from :func:`nonbonded_pair_positions`. The H-VDW term ``NonBondedHTarget`` adds is **not** included here. Dispatches to :func:`torchref.base.targets.triton.nonbonded_heavy_math_triton` on CUDA float32 (the gain coming mostly from the analytic backward), eager otherwise. @@ -76,15 +139,17 @@ def nonbonded_heavy_math( Parameters ---------- xyz : torch.Tensor - (N_atoms, 3) Cartesian coordinates of the ASU. + (N_atoms, 3) Cartesian coordinates of the ASU in Å. indices : torch.Tensor (N, 2) per-pair atom indices. min_distances : torch.Tensor - (N,) VDW threshold per pair. + (N,) VDW threshold per pair in Å. symop_indices : torch.Tensor, optional - (N,) symmetry-operator index per pair; 0 = identity. + (N,) symmetry-operator index per pair; 0 = identity. None treats every + partner as the ASU atom itself; otherwise every pair is imaged, see + :func:`nonbonded_pair_positions`. cell_offsets : torch.Tensor, optional - (N, 3) fractional cell offsets per pair. + (N, 3) integer lattice translations per pair, in fractional units. symop_matrices, symop_translations : torch.Tensor, optional (n_symops, 3, 3) and (n_symops, 3) — the symmetry operator table. fractional_matrix, inv_fractional_matrix : torch.Tensor diff --git a/torchref/base/targets/triton/nonbonded.py b/torchref/base/targets/triton/nonbonded.py index 2a729e24..c2e35dac 100644 --- a/torchref/base/targets/triton/nonbonded.py +++ b/torchref/base/targets/triton/nonbonded.py @@ -13,7 +13,10 @@ Cartesian symmetry transforms (M·R·M⁻¹ and M·t) are precomputed once on the host, so the kernel only does a 3×3 matvec (forward) and 3×3 -transposed matvec (backward) per pair. +transposed matvec (backward) per pair. As in the eager +:func:`~torchref.base.targets.nonbonded.nonbonded_pair_positions`, every pair +is transformed whenever symmetry tensors are passed: +``M·R·M⁻¹·x + M·t + M·n = M·(R·M⁻¹·x + t + n)``, the eager image term for term. """ from __future__ import annotations @@ -226,18 +229,22 @@ def forward(ctx, xyz, indices, min_distances, N = indices.shape[0] nll = torch.empty(N, dtype=xyz.dtype, device=xyz.device) - has_sym = ( - symop_indices is not None - and symop_indices.numel() > 0 - and not bool((symop_indices == 0).all()) - ) + # The eager rule (``nonbonded_pair_positions``): with symmetry tensors present + # every pair is imaged, identity rows included. Testing the operations alone + # would drop the lattice translation of a pair whose operation is the identity, + # which is every crystal contact in P1. + has_sym = symop_indices is not None if has_sym: cart_mat, cart_off = _build_cartesian_symops( symop_matrices, symop_translations, fractional_matrix, inv_fractional_matrix, ) - cell_off_cart = (cell_offsets.to(xyz.dtype) - @ fractional_matrix.T).contiguous() + if cell_offsets is None: + cell_off_cart = torch.zeros(N, 3, device=xyz.device, dtype=xyz.dtype) + else: + cell_off_cart = ( + cell_offsets.to(xyz.dtype) @ fractional_matrix.T + ).contiguous() symop_i32 = symop_indices.to(torch.int32).contiguous() else: cart_mat = torch.zeros(1, 3, 3, device=xyz.device, dtype=xyz.dtype) diff --git a/torchref/refinement/targets/geometry/non_bonded.py b/torchref/refinement/targets/geometry/non_bonded.py index ee75930e..4e249798 100644 --- a/torchref/refinement/targets/geometry/non_bonded.py +++ b/torchref/refinement/targets/geometry/non_bonded.py @@ -9,6 +9,7 @@ import torch from typing import TYPE_CHECKING, Dict, Tuple +from torchref.base.coordinates.symmetry_images import is_symmetry_image from torchref.config import get_int_dtype from torchref.utils.stats import ( VERBOSITY_DEBUG, @@ -213,52 +214,45 @@ def maintenance(self) -> None: ) r.rebuild_vdw_restraints(self._model.xyz().detach()) + def _symmetry_tables(self) -> tuple[torch.Tensor, ...]: + """The model's operation table and cell matrices, as the pair kernels take them. + + Returns + ------- + tuple of torch.Tensor + ``(symop_matrices, symop_translations, fractional_matrix, + inv_fractional_matrix)``. + """ + sg = self.model.spacegroup + cell = self.model.cell + return ( + sg.matrices, + sg.translations, + cell.fractional_matrix, + cell.inv_fractional_matrix, + ) + def _compute_positions( self, xyz: torch.Tensor ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Per-pair ``(pos1, pos2, min_distances)`` from ASU coordinates (N, 3). - One vectorized pass. Mate positions are recomputed through the symmetry - transform rather than looked up, so gradients reach both atoms; intra-ASU - pairs (symop=0, offset=0) come out as the identity. + Positions come from + :func:`~torchref.base.targets.nonbonded.nonbonded_pair_positions`, as in the + prolsq kernel, so every mode and the statistics see a pair at one distance. + Mate positions are recomputed from ``xyz``, so gradients reach both atoms. """ - vdw_data = self.restraints.restraints["vdw"] - indices = vdw_data["indices"] - min_distances = vdw_data["min_distances"] - symop_indices = vdw_data.get("symop_indices") - cell_offsets = vdw_data.get("cell_offsets") + from torchref.base.targets.nonbonded import nonbonded_pair_positions - pos1 = xyz[indices[:, 0]] - - has_symmetry = ( - symop_indices is not None - and len(symop_indices) > 0 - and not (symop_indices == 0).all() - ) - - if not has_symmetry: - # Fast path: all pairs are intra-ASU. - pos2 = xyz[indices[:, 1]] - return pos1, pos2, min_distances - - cell = self.model.cell - sg = self.model.spacegroup - - mate_source = xyz[indices[:, 1]] # (N_pairs, 3) -- gradients flow - frac = cell.cartesian_to_fractional(mate_source) - - R = sg.matrices[symop_indices].to(frac.dtype) # (N_pairs, 3, 3) - t = sg.translations[symop_indices].to(frac.dtype) # (N_pairs, 3) - offsets = cell_offsets.to(frac.dtype) # (N_pairs, 3) - - # R @ frac + t + offset, batched; identity for intra-ASU pairs. - frac_transformed = ( - torch.bmm(R, frac.unsqueeze(-1)).squeeze(-1) + t + offsets + vdw_data = self.restraints.restraints["vdw"] + pos1, pos2 = nonbonded_pair_positions( + xyz, + vdw_data["indices"], + vdw_data.get("symop_indices"), + vdw_data.get("cell_offsets"), + *self._symmetry_tables(), ) - - pos2 = cell.fractional_to_cartesian(frac_transformed) - - return pos1, pos2, min_distances + return pos1, pos2, vdw_data["min_distances"] def forward(self) -> torch.Tensor: """Summed VDW repulsion loss; 0.0 if the model has no VDW pair list.""" @@ -286,10 +280,7 @@ def forward(self) -> torch.Tensor: vdw_data["min_distances"], vdw_data.get("symop_indices"), vdw_data.get("cell_offsets"), - self.model.spacegroup.matrices, - self.model.spacegroup.translations, - self.model.cell.fractional_matrix, - self.model.cell.inv_fractional_matrix, + *self._symmetry_tables(), self._c_rep, self._r_exp, self._buffer, self._sigma_vdw, ) @@ -434,8 +425,7 @@ def stats(self) -> Dict[str, any]: symop_indices = vdw_data.get("symop_indices") cell_offsets = vdw_data.get("cell_offsets") if symop_indices is not None and len(symop_indices) > 0: - is_sym = (symop_indices != 0) | (cell_offsets != 0).any(dim=-1) - n_sym = is_sym.sum().item() + n_sym = is_symmetry_image(symop_indices, cell_offsets).sum().item() if n_sym > 0: result["n_symmetry"] = stat(n_sym, VERBOSITY_DETAILED) diff --git a/torchref/refinement/targets/geometry/non_bonded_h.py b/torchref/refinement/targets/geometry/non_bonded_h.py index 8a249b24..42d309f5 100644 --- a/torchref/refinement/targets/geometry/non_bonded_h.py +++ b/torchref/refinement/targets/geometry/non_bonded_h.py @@ -9,6 +9,7 @@ import torch from typing import TYPE_CHECKING, Dict +from torchref.base.coordinates.symmetry_images import is_symmetry_image from torchref.config import dtypes from torchref.utils.stats import ( VERBOSITY_DEBUG, @@ -78,6 +79,52 @@ def __init__( # H-VDW loss via precomputed candidates # ------------------------------------------------------------------ + @staticmethod + def _h_candidates( + xyz: torch.Tensor, h_topo: "HydrogenTopology" + ) -> tuple[torch.Tensor, torch.Tensor]: + """``[heavy | riding H]`` coordinates and the candidate pairs indexing them. + + Parameters + ---------- + xyz : torch.Tensor + Heavy-atom Cartesian coordinates in Å, ``(N_heavy, 3)``. + h_topo : HydrogenTopology + Riding topology with candidate pairs built. + + Returns + ------- + xyz_all : torch.Tensor + ``(N_heavy + N_h, 3)`` in Å; the hydrogens are placed from ``xyz`` on + every call, differentiably. + indices : torch.Tensor + ``(P, 2)`` contiguous candidate pairs into ``xyz_all``. + """ + from torchref.topology.riding import place_riding_hydrogens + + xyz_all = torch.cat([xyz, place_riding_hydrogens(xyz, h_topo)], dim=0) + indices = torch.stack([h_topo.cand_idx_i, h_topo.cand_idx_j], dim=1) + return xyz_all, indices.contiguous() + + def _h_pair_positions( + self, xyz: torch.Tensor, h_topo: "HydrogenTopology" + ) -> tuple[torch.Tensor, torch.Tensor]: + """Both ends of every H candidate pair, ``(P, 3)`` each in Å. + + From :func:`~torchref.base.targets.nonbonded.nonbonded_pair_positions`, the + positions the prolsq kernel scores, symmetry and lattice images included. + """ + from torchref.base.targets.nonbonded import nonbonded_pair_positions + + xyz_all, indices = self._h_candidates(xyz, h_topo) + return nonbonded_pair_positions( + xyz_all, + indices, + h_topo.cand_symop_idx, + h_topo.cand_cell_offset, + *self._symmetry_tables(), + ) + def _compute_h_vdw_loss( self, xyz: torch.Tensor, @@ -86,20 +133,12 @@ def _compute_h_vdw_loss( """VDW loss over the precomputed H-heavy candidate pairs. Places riding hydrogens differentiably, so the gradient runs - loss -> H_pos -> ``xyz[parent_idx]`` -> model parameters. The candidate list - must stay sorted ASU-then-sym: the ``prolsq`` fast path hands - ``cand_symop_idx`` straight to - :func:`torchref.base.targets.nonbonded_heavy_math`, which relies on that - ordering to do identity and real symmetry transforms in one pass. Other modes - take the inline eager path below. + loss -> H_pos -> ``xyz[parent_idx]`` -> model parameters. The ``prolsq`` + mode goes through :func:`torchref.base.targets.nonbonded_heavy_math` (Triton on + CUDA float32); the others score the positions of :meth:`_h_pair_positions`. """ - from torchref.topology.riding import place_riding_hydrogens - device = xyz.device - xyz_h = place_riding_hydrogens(xyz, h_topo) - xyz_all = torch.cat([xyz, xyz_h], dim=0) # [heavy | H] - n_cand = h_topo.cand_idx_i.shape[0] if n_cand == 0: return torch.tensor(0.0, device=device) @@ -107,51 +146,20 @@ def _compute_h_vdw_loss( # Fast path: prolsq goes through the dispatcher (Triton on CUDA fp32). if self.mode == "prolsq": from torchref.base.targets.nonbonded import nonbonded_heavy_math - indices = torch.stack( - [h_topo.cand_idx_i, h_topo.cand_idx_j], dim=1 - ).contiguous() + + xyz_all, indices = self._h_candidates(xyz, h_topo) return nonbonded_heavy_math( xyz_all, indices, h_topo.cand_min_dist, h_topo.cand_symop_idx, h_topo.cand_cell_offset, - self.model.spacegroup.matrices, - self.model.spacegroup.translations, - self.model.cell.fractional_matrix, - self.model.cell.inv_fractional_matrix, + *self._symmetry_tables(), self._c_rep, self._r_exp, float(self._buffer), self._sigma_vdw, ) - # Slow path: gaussian / soft modes, inline eager. - pos_i = xyz_all[h_topo.cand_idx_i] - n_asu = h_topo.n_asu_candidates - n_sym = n_cand - n_asu + pos_i, pos_j = self._h_pair_positions(xyz, h_topo) + actual_dist = torch.sqrt(((pos_j - pos_i) ** 2).sum(dim=-1) + 1e-8) min_dist = h_topo.cand_min_dist - if n_asu > 0: - pos_j_asu = xyz_all[h_topo.cand_idx_j[:n_asu]] - diff_asu = pos_j_asu - pos_i[:n_asu] - dist_asu = torch.sqrt((diff_asu ** 2).sum(dim=-1) + 1e-8) - - if n_sym > 0: - cell = self.model.cell - sg = self.model.spacegroup - sym_source = xyz_all[h_topo.cand_idx_j[n_asu:]] - frac = cell.cartesian_to_fractional(sym_source) - R = sg.matrices[h_topo.cand_symop_idx[n_asu:]].to(frac.dtype) - t = sg.translations[h_topo.cand_symop_idx[n_asu:]].to(frac.dtype) - offs = h_topo.cand_cell_offset[n_asu:].to(frac.dtype) - frac_t = torch.bmm(R, frac.unsqueeze(-1)).squeeze(-1) + t + offs - pos_j_sym = cell.fractional_to_cartesian(frac_t) - diff_sym = pos_j_sym - pos_i[n_asu:] - dist_sym = torch.sqrt((diff_sym ** 2).sum(dim=-1) + 1e-8) - - if n_asu > 0 and n_sym > 0: - actual_dist = torch.cat([dist_asu, dist_sym]) - elif n_asu > 0: - actual_dist = dist_asu - else: - actual_dist = dist_sym - violations = torch.clamp(min_dist + self._buffer - actual_dist, min=0.0) if self.mode == "gaussian": @@ -196,12 +204,9 @@ def forward(self) -> torch.Tensor: def get_violations(self, threshold: float = 0.0) -> Dict[str, torch.Tensor]: """Parent VDW violations plus ``h_*`` entries for H-involving contacts. - The H distances here ignore symmetry -- both positions are read straight out of - ``xyz_all`` -- so symmetry-mate H contacts are reported at their intra-ASU - separation, unlike in the loss. + H distances are taken between the positions the loss scores, symmetry and + lattice images included. """ - from torchref.topology.riding import place_riding_hydrogens - result = super().get_violations(threshold) restraints = self.restraints @@ -211,14 +216,7 @@ def get_violations(self, threshold: float = 0.0) -> Dict[str, torch.Tensor]: if h_topo is None or h_topo.n_hydrogens == 0 or not h_topo.has_candidates: return result - xyz = self.model.xyz() - device = xyz.device - xyz_h = place_riding_hydrogens(xyz, h_topo) - xyz_all = torch.cat([xyz, xyz_h], dim=0) - - pos_i = xyz_all[h_topo.cand_idx_i] - pos_j = xyz_all[h_topo.cand_idx_j] - + pos_i, pos_j = self._h_pair_positions(self.model.xyz(), h_topo) actual_dist = torch.norm(pos_j - pos_i, dim=-1) violations = torch.clamp(h_topo.cand_min_dist - actual_dist, min=0.0) @@ -234,8 +232,6 @@ def get_violations(self, threshold: float = 0.0) -> Dict[str, torch.Tensor]: def stats(self) -> Dict[str, any]: """Get statistics including H-VDW contacts.""" - from torchref.topology.riding import place_riding_hydrogens - result = super().stats() restraints = self.restraints @@ -245,14 +241,7 @@ def stats(self) -> Dict[str, any]: if h_topo is None or h_topo.n_hydrogens == 0 or not h_topo.has_candidates: return result - xyz = self.model.xyz() - device = xyz.device - xyz_h = place_riding_hydrogens(xyz, h_topo) - xyz_all = torch.cat([xyz, xyz_h], dim=0) - - pos_i = xyz_all[h_topo.cand_idx_i] - pos_j = xyz_all[h_topo.cand_idx_j] - + pos_i, pos_j = self._h_pair_positions(self.model.xyz(), h_topo) actual_dist = torch.norm(pos_j - pos_i, dim=-1) violations = torch.clamp(h_topo.cand_min_dist - actual_dist, min=0.0) @@ -269,8 +258,8 @@ def stats(self) -> Dict[str, any]: result["h_rms_violation"] = stat(rms, VERBOSITY_DETAILED) result["h_max_violation"] = stat(violations.max().item(), VERBOSITY_DEBUG) - n_sym = ((h_topo.cand_symop_idx != 0) - | (h_topo.cand_cell_offset != 0).any(dim=1)).sum().item() + is_sym = is_symmetry_image(h_topo.cand_symop_idx, h_topo.cand_cell_offset) + n_sym = is_sym.sum().item() if n_sym > 0: result["h_n_symmetry"] = stat(n_sym, VERBOSITY_DETAILED) diff --git a/torchref/topology/nonbonded.py b/torchref/topology/nonbonded.py index 37a081e5..e8a00525 100644 --- a/torchref/topology/nonbonded.py +++ b/torchref/topology/nonbonded.py @@ -17,6 +17,10 @@ import numpy as np import torch +from torchref.base.coordinates.symmetry_images import ( + is_symmetry_image, + symmetry_image_positions, +) from torchref.config import dtypes, get_float_dtype, get_int_dtype if TYPE_CHECKING: @@ -69,13 +73,18 @@ def vdw_radii_for_elements(elements) -> np.ndarray: def prefilter_symop_offsets( cell: "Cell", sg: "SpaceGroup", - xyz_frac: torch.Tensor, + xyz: torch.Tensor, cutoff: float, ) -> Tuple[torch.Tensor, torch.Tensor]: - """Select (symop, cell_offset) combos that could produce contacts. + """Select every (symop, cell_offset) combination whose image can reach the model. - Uses the ASU centroid and molecule radius to eliminate obviously - distant combinations. Always includes identity (op=0, offset=0). + An image can hold an atom within ``cutoff`` of the model only if it moves the + centroid by at most ``2 r + cutoff``, ``r`` being the largest atom-centroid + distance. Every combination meeting that bound is returned, however far the + coordinates sit from the origin cell: the offsets are relative to ``xyz`` as + given, which is how :func:`assign_to_grid` and the VDW kernels form images from + them (:func:`~torchref.base.coordinates.symmetry_image_positions`, on unwrapped + coordinates). Always includes the identity (op=0, offset=0). Parameters ---------- @@ -83,8 +92,8 @@ def prefilter_symop_offsets( Crystallographic unit cell. sg : SpaceGroup Space group providing the symmetry operators. - xyz_frac : torch.Tensor - ``(N, 3)`` fractional ASU coordinates. + xyz : torch.Tensor + ``(N, 3)`` Cartesian ASU coordinates in Å. cutoff : float Cartesian cutoff in Angstrom. @@ -92,42 +101,51 @@ def prefilter_symop_offsets( ------- op_indices : (M,) int – symop indices for each valid combo cell_offsets : (M, 3) int – integer cell translations + Ordered by operation, then offset, lexicographically. """ - device = xyz_frac.device - fdtype = dtypes.float - - centroid_frac = xyz_frac.mean(dim=0) - centroid_cart = cell.fractional_to_cartesian(xyz_frac).mean(dim=0) - xyz_cart = cell.fractional_to_cartesian(xyz_frac) - molecule_radius = (xyz_cart - centroid_cart).norm(dim=1).max().item() - threshold = 2.0 * molecule_radius + cutoff + device = xyz.device + fdtype = get_float_dtype() + int_dtype = get_int_dtype() + xyz = xyz.to(fdtype) B = cell.fractional_matrix.to(device=device, dtype=fdtype) - I_mat = torch.eye(3, dtype=fdtype, device=device) - + B_inv = cell.inv_fractional_matrix.to(device=device, dtype=fdtype) matrices = sg.matrices.to(device=device, dtype=fdtype) translations = sg.translations.to(device=device, dtype=fdtype) + tables = (matrices, translations, B, B_inv) - valid_ops = [] - valid_offsets = [] + centroid = xyz.mean(dim=0) + reach = 2.0 * (xyz - centroid).norm(dim=1).max() + cutoff - for op_idx in range(sg.n_ops): - R = matrices[op_idx] - t = translations[op_idx] - for dx in range(-1, 2): - for dy in range(-1, 2): - for dz in range(-1, 2): - offset = torch.tensor([dx, dy, dz], dtype=fdtype, - device=device) - d_frac = (R - I_mat) @ centroid_frac + t + offset - d_cart = B @ d_frac - if d_cart.norm().item() <= threshold: - valid_ops.append(op_idx) - valid_offsets.append([dx, dy, dz]) + n_ops = matrices.shape[0] + ops = torch.arange(n_ops, dtype=int_dtype, device=device) + no_offset = torch.zeros(n_ops, 3, dtype=int_dtype, device=device) + image0 = symmetry_image_positions( + centroid.expand(n_ops, 3), ops, no_offset, *tables + ) + # Fractional centroid displacement under each operation before any lattice + # translation. An offset n qualifies only if |B (u + n)| <= reach, which along + # axis k bounds |u_k + n_k| by reach * |row k of B^-1| (reach over the spacing of + # the lattice planes normal to that axis), so this box holds every candidate. + u = (image0 - centroid) @ B_inv.T + half_width = reach * B_inv.norm(dim=1) + lo = torch.ceil(-u - half_width) + hi = torch.floor(-u + half_width) + + box = (hi - lo).max(dim=0).values.to(int_dtype) + 1 + steps = torch.cartesian_prod( + *(torch.arange(int(k), dtype=fdtype, device=device) for k in box) + ) + candidates = lo[:, None, :] + steps[None, :, :] # (n_ops, K, 3) + in_box = (candidates <= hi[:, None, :]).all(dim=-1) + candidate_ops = ops[:, None].expand(-1, steps.shape[0]) + offsets = candidates.to(int_dtype) + image = symmetry_image_positions( + centroid.expand_as(candidates), candidate_ops, offsets, *tables + ) + keep = in_box & ((image - centroid).norm(dim=-1) <= reach) - op_indices = torch.tensor(valid_ops, dtype=get_int_dtype(), device=device) - cell_offsets = torch.tensor(valid_offsets, dtype=get_int_dtype(), device=device) - return op_indices, cell_offsets + return candidate_ops[keep], offsets[keep] # ------------------------------------------------------------------ # @@ -135,7 +153,7 @@ def prefilter_symop_offsets( # ------------------------------------------------------------------ # def assign_to_grid( - xyz_frac: torch.Tensor, + xyz: torch.Tensor, cell: "Cell", sg: "SpaceGroup", op_indices: torch.Tensor, @@ -144,9 +162,15 @@ def assign_to_grid( ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: """Compute Cartesian image positions and assign to grid cells. + Images come from :func:`~torchref.base.coordinates.symmetry_image_positions`, + the function the VDW kernels place stored pairs with, so a pair is scored at the + distance it is found at here. + Parameters ---------- - xyz_frac : (N, 3) + xyz : (N, 3) + Cartesian ASU coordinates in Å, unwrapped, as passed to + :func:`prefilter_symop_offsets`. cell : Cell sg : SpaceGroup op_indices : (M,) int @@ -160,29 +184,26 @@ def assign_to_grid( combo_idx : (N*M,) long – index into op_indices / cell_offsets cart_pos : (N*M, 3) float – Cartesian positions (reused in step 4) """ - device = xyz_frac.device - fdtype = dtypes.float - N = xyz_frac.shape[0] + device = xyz.device + fdtype = get_float_dtype() + N = xyz.shape[0] M = op_indices.shape[0] - R_sel = sg.matrices[op_indices].to(dtype=fdtype) # (M, 3, 3) - t_sel = sg.translations[op_indices].to(dtype=fdtype) # (M, 3) - offs = cell_offsets.to(dtype=fdtype) # (M, 3) - - # (N, M, 3) = einsum over symops applied to each atom - frac_images = ( - torch.einsum("mij,nj->nmi", R_sel, xyz_frac.to(fdtype)) - + t_sel[None, :, :] - + offs[None, :, :] - ) - - # Cartesian positions (stored for reuse) - cart_pos = cell.fractional_to_cartesian( - frac_images.reshape(-1, 3) - ) # (N*M, 3) - - # Wrap to [0, 1) for grid assignment - frac_wrapped = frac_images % 1.0 + B = cell.fractional_matrix.to(device=device, dtype=fdtype) + B_inv = cell.inv_fractional_matrix.to(device=device, dtype=fdtype) + images = symmetry_image_positions( + xyz.to(fdtype)[:, None, :], + op_indices.to(device), + cell_offsets.to(device), + sg.matrices.to(device=device, dtype=fdtype), + sg.translations.to(device=device, dtype=fdtype), + B, + B_inv, + ) # (N, M, 3) + cart_pos = images.reshape(-1, 3) + + # Wrap to [0, 1) for grid assignment only; distances use the unwrapped images. + frac_wrapped = (images @ B_inv.T) % 1.0 gd = grid_dims.to(device=device, dtype=fdtype) cell_ijk = (frac_wrapped * gd[None, None, :]).long() cell_ijk = cell_ijk.clamp( @@ -732,20 +753,15 @@ def build_vdw_restraints_gpu( } # Step 1: prefilter symop combos - xyz_frac = cell.cartesian_to_fractional(xyz.detach().to(fdtype)) - op_indices, cell_offsets_valid = prefilter_symop_offsets( - cell, sg, xyz_frac, cutoff - ) + xyz_asu = xyz.detach().to(fdtype) + op_indices, cell_offsets_valid = prefilter_symop_offsets(cell, sg, xyz_asu, cutoff) M = len(op_indices) if verbose > 0: print(f" Symmetry expansion: {M} valid (symop, offset) combos") # Find the identity combo index - is_identity = ( - (op_indices == 0) - & (cell_offsets_valid == 0).all(dim=1) - ) + is_identity = ~is_symmetry_image(op_indices, cell_offsets_valid) identity_indices = is_identity.nonzero(as_tuple=True)[0] if len(identity_indices) == 0: # Identity not in valid combos — should not happen, but add it @@ -773,7 +789,7 @@ def build_vdw_restraints_gpu( ) # (3,) flat_cell, atom_idx, combo_idx, cart_pos = assign_to_grid( - xyz_frac, cell, sg, op_indices, cell_offsets_valid, grid_dims + xyz_asu, cell, sg, op_indices, cell_offsets_valid, grid_dims ) if device.type == "cpu": @@ -885,9 +901,7 @@ def build_vdw_restraints_gpu( } if verbose > 0: - n_sym = ( - (symop_indices != 0) | (pair_cell_offsets != 0).any(dim=1) - ).sum().item() + n_sym = is_symmetry_image(symop_indices, pair_cell_offsets).sum().item() print(f" Built {len(indices)} VDW restraints, {n_sym} symmetry contacts") return result diff --git a/torchref/topology/restraints.py b/torchref/topology/restraints.py index e07281b1..dee0a841 100644 --- a/torchref/topology/restraints.py +++ b/torchref/topology/restraints.py @@ -620,8 +620,11 @@ def get_count(rtype, origin): symop_indices = self.restraints["vdw"].get("symop_indices") cell_offsets = self.restraints["vdw"].get("cell_offsets") if symop_indices is not None and len(symop_indices) > 0: - import torch as _torch - is_sym = (symop_indices != 0) | (cell_offsets != 0).any(dim=-1) + from torchref.base.coordinates.symmetry_images import ( + is_symmetry_image, + ) + + is_sym = is_symmetry_image(symop_indices, cell_offsets) vdw_sym_count = int(is_sym.sum().item()) vdw_asu_count = vdw_count - vdw_sym_count if vdw_sym_count > 0: diff --git a/torchref/topology/riding.py b/torchref/topology/riding.py index ee582e2b..d2447474 100644 --- a/torchref/topology/riding.py +++ b/torchref/topology/riding.py @@ -23,6 +23,7 @@ import numpy as np import torch +from torchref.base.coordinates.symmetry_images import is_symmetry_image from torchref.config import dtypes, get_int_dtype, normalize_device from torchref.utils.device_resolution import resolve_device from torchref.utils.device_mixin import DeviceMixin @@ -715,11 +716,14 @@ def build_h_candidate_pairs( """Precompute candidate H-involving VDW pairs from the heavy-atom pair list. From each heavy-heavy pair (A, B, symop, offset), derives the H-heavy pairs - where an H riding on A could reach B and vice versa, applying the exclusion and - same-residue filters now so the forward pass only computes distances. Mutates - ``h_topo`` in place, registering ``cand_idx_i``/``cand_idx_j`` (combined-array - atom indices), ``cand_symop_idx`` and ``cand_cell_offset`` (for the heavy atom) - and ``cand_min_dist`` (H + heavy radius sum). + where an H riding on A could reach B, and for intra-ASU pairs vice versa, + applying the exclusion and same-residue filters now so the forward pass only + computes distances. An image pair's other direction comes from its reverse + entry, B against A under the inverse operation, so the heavy list must hold both + directions of every image contact, as ``build_vdw_restraints_gpu`` emits them. + Mutates ``h_topo`` in place, registering ``cand_idx_i``/``cand_idx_j`` + (combined-array atom indices), ``cand_symop_idx`` and ``cand_cell_offset`` (the + image of the ``cand_idx_j`` end) and ``cand_min_dist`` (H + heavy radius sum). Parameters ---------- @@ -778,6 +782,7 @@ def build_h_candidate_pairs( idx_B = heavy_indices[:, 1].cpu().numpy() symop_np = heavy_symop.cpu().numpy() offsets_np = heavy_offsets.cpu().numpy() + is_image_np = is_symmetry_image(heavy_symop, heavy_offsets).cpu().numpy() # Per-pair VDW radius sums are not computed here: the cand_min_dist # buffer is allocated as zeros below and is populated by the caller, @@ -798,7 +803,7 @@ def _same_res(chain_a, resseq_a, chain_b, resseq_b): A, B = int(idx_A[p_idx]), int(idx_B[p_idx]) sym = int(symop_np[p_idx]) off = offsets_np[p_idx] - is_intra_asu = (sym == 0) and (off == 0).all() + is_intra_asu = not is_image_np[p_idx] h_on_A = parent_to_h.get(A, []) h_on_B = parent_to_h.get(B, []) @@ -815,15 +820,23 @@ def _same_res(chain_a, resseq_a, chain_b, resseq_b): acc_offset.append(off) # --- H on B ↔ heavy A --- - for hi in h_on_B: - if is_intra_asu and _same_res( - h_chain_np[hi], h_resseq_np[hi], heavy_chain_np[A], heavy_resseq_np[A] - ): - continue - acc_idx_i.append(n_heavy + hi) - acc_idx_j.append(A) - acc_symop.append(0) - acc_offset.append(np.zeros(3, dtype=np.int64)) + # Intra-ASU pairs only. An image pair -- A against B under (symop, offset) + # -- is listed together with B against A under the inverse operation, whose + # "H on A" branch above emits this contact with the image on the right atom. + # From here it could only carry this pair's operation, which images B, not A. + if is_intra_asu: + for hi in h_on_B: + if _same_res( + h_chain_np[hi], + h_resseq_np[hi], + heavy_chain_np[A], + heavy_resseq_np[A], + ): + continue + acc_idx_i.append(n_heavy + hi) + acc_idx_j.append(A) + acc_symop.append(0) + acc_offset.append(np.zeros(3, dtype=np.int64)) # --- H on A ↔ H on B (H-H contacts) --- for hi_a in h_on_A: @@ -859,7 +872,7 @@ def _same_res(chain_a, resseq_a, chain_b, resseq_b): # Apply 1-2 / 1-3 exclusions for intra-ASU candidates if h_excl_hash is not None and len(h_excl_hash) > 0: - is_intra = (cand_sym == 0) & (cand_off == 0).all(dim=1) + is_intra = ~is_symmetry_image(cand_sym, cand_off) if is_intra.any(): max_idx = n_heavy + n_h norm_i = torch.minimum(cand_i, cand_j) @@ -876,18 +889,13 @@ def _same_res(chain_a, resseq_a, chain_b, resseq_b): cand_sym = cand_sym[keep] cand_off = cand_off[keep] - # Deduplicate + # Deduplicate on whole (i, j, symop, offset) rows. No fixed-stride packed key is + # safe: offsets are not confined to -1..1, nor operations to a small count. if len(cand_i) > 0: - n_all = n_heavy + n_h - dedup_key = ( - cand_i.long() * (n_all * 1000) - + cand_j.long() * 1000 - + cand_sym.long() * 27 - + (cand_off[:, 0] + 1) * 9 - + (cand_off[:, 1] + 1) * 3 - + (cand_off[:, 2] + 1) + rows = torch.cat( + [torch.stack([cand_i, cand_j, cand_sym], dim=1), cand_off], dim=1 ) - _, first_idx = torch.unique(dedup_key, return_inverse=True) + _, first_idx = torch.unique(rows, dim=0, return_inverse=True) # MPS does not support int64 scatter_reduce; use configured int dtype. _int_dtype = dtypes.int first_idx_i = first_idx.to(_int_dtype) @@ -905,7 +913,7 @@ def _same_res(chain_a, resseq_a, chain_b, resseq_b): cand_off = cand_off[mask] # Sort: ASU candidates first, symmetry last - is_asu = (cand_sym == 0) & (cand_off == 0).all(dim=1) + is_asu = ~is_symmetry_image(cand_sym, cand_off) sort_order = (~is_asu).long().argsort(stable=True) cand_i = cand_i[sort_order] cand_j = cand_j[sort_order] @@ -923,7 +931,7 @@ def _same_res(chain_a, resseq_a, chain_b, resseq_b): if verbose > 0: n_hh = ((cand_i >= n_heavy) & (cand_j >= n_heavy)).sum().item() - n_sym = ((cand_sym != 0) | (cand_off != 0).any(dim=1)).sum().item() + n_sym = (~is_asu).sum().item() print( f" H candidate pairs: {len(cand_i)} " f"({n_hh} H-H, {len(cand_i)-n_hh} H-heavy, {n_sym} symmetry)" From fcd18b2ceaaadab7aa693e1b46f6bf09688b80c1 Mon Sep 17 00:00:00 2001 From: Claude Date: Fri, 2 Oct 2026 01:11:16 +0000 Subject: [PATCH 248/250] Model anomalous scattering only when the data's wavelength is given The 1.0 A default gave every model f'/f'' and, through the Bijvoet auto-detect, read F(+)/F(-) as separate observations, at a wavelength that is rarely the one the data were collected at: near an absorption edge f' and f'' change by several electrons over a few hundredths of an Angstrom. ModelFT, Refinement, EnsembleModel and --wavelength now default to None, meaning no f'/f'' and a Friedel-merged read, which 0 also selects. A given wavelength enables both, and anomalous=True without one raises instead of reading Bijvoet pairs the model cannot tell apart. A ModelFT state dict without a wavelength restores with none, as the constructor does; every state dict ModelFT writes carries its wavelength. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01KuUn93S8goPxwJLBWjjhht --- docs/user_guide/cli.rst | 5 +- tests/integration/test_cli_refinement.py | 55 ++++++++++++++++--- torchref/cli/_common.py | 12 ++-- torchref/cli/refine.py | 4 +- .../experimental/ensemble/ensemble_model.py | 2 +- torchref/model/model_ft.py | 18 +++--- torchref/refinement/base_refinement.py | 40 ++++++++------ 7 files changed, 91 insertions(+), 45 deletions(-) diff --git a/docs/user_guide/cli.rst b/docs/user_guide/cli.rst index b5686401..b34425c8 100644 --- a/docs/user_guide/cli.rst +++ b/docs/user_guide/cli.rst @@ -63,8 +63,9 @@ and a ``refinement_history.json`` log. Ramachandran with ``--weights '{"geometry/ramachandran": 1.0}'`` * ``--with-rigid-body`` run rigid-body first (``--rigid-body-iter``, ``--rigid-body-cutoffs``) -* ``--wavelength`` Å, for anomalous f'/f''. ``0`` disables anomalous refinement - and forces a Friedel-merged read of the data +* ``--wavelength`` Å, the wavelength the data were collected at. Given, the model + includes anomalous f'/f'' and reads F(+)/F(-) as Bijvoet pairs; without it (or + with ``0``) there is no anomalous scattering and the read is Friedel-merged * ``--dmin`` resolution cutoff * ``--output-format`` ``pdb`` / ``cif`` / ``both`` (default both) * ``--device`` ``auto`` (default) / ``cpu`` / ``cuda``. ``auto`` picks CUDA only diff --git a/tests/integration/test_cli_refinement.py b/tests/integration/test_cli_refinement.py index c3325871..e0d6e4c5 100644 --- a/tests/integration/test_cli_refinement.py +++ b/tests/integration/test_cli_refinement.py @@ -149,14 +149,32 @@ def small_pair(self, test_files_dir): pytest.skip("1DAW test files not found") return {"pdb": str(pdb), "mtz": str(mtz)} + @pytest.fixture + def anomalous_mtz(self, small_pair, tmp_path): + """1DAW written with FP(+)/FP(-) columns, so the file offers Bijvoet pairs.""" + import reciprocalspaceship as rs + + ds = rs.read_mtz(small_pair["mtz"]) + out = tmp_path / "1DAW_anomalous.mtz" + ds[["FP", "SIGFP", "FreeR_flag"]].copy().unstack_anomalous( + columns=["FP", "SIGFP"] + ).write_mtz(str(out)) + return str(out) + @pytest.mark.integration - def test_wavelength_zero_disables_anomalous_and_merges(self, small_pair): - """wavelength=0 -> no anomalous correction + forced Friedel-merged read.""" + @pytest.mark.parametrize( + "kwargs", [{}, {"wavelength": 0}], ids=["default", "wavelength_zero"] + ) + def test_no_wavelength_means_no_anomalous(self, small_pair, anomalous_mtz, kwargs): + """Without a wavelength the model has no f'/f'' and F(+)/F(-) are merged.""" from torchref.refinement.lbfgs_refinement import LBFGSRefinement ref = LBFGSRefinement( - data_file=small_pair["mtz"], pdb=small_pair["pdb"], - device=torch.device("cpu"), verbose=0, wavelength=0, + data_file=anomalous_mtz, + pdb=small_pair["pdb"], + device=torch.device("cpu"), + verbose=0, + **kwargs, ) assert ref.wavelength is None assert ref.anomalous is False @@ -165,12 +183,31 @@ def test_wavelength_zero_disables_anomalous_and_merges(self, small_pair): assert bool(ref.model.anomalous_bijvoet) is False @pytest.mark.integration - def test_wavelength_default_preserved(self, small_pair): + def test_wavelength_reads_bijvoet_pairs(self, small_pair, anomalous_mtz): + """A wavelength gives the model f'/f'' and reads F(+)/F(-) as Bijvoet pairs.""" from torchref.refinement.lbfgs_refinement import LBFGSRefinement ref = LBFGSRefinement( - data_file=small_pair["mtz"], pdb=small_pair["pdb"], - device=torch.device("cpu"), verbose=0, + data_file=anomalous_mtz, + pdb=small_pair["pdb"], + device=torch.device("cpu"), + verbose=0, + wavelength=1.54, ) - assert ref.wavelength == 1.0 - assert ref.model.wavelength == 1.0 + assert bool(ref.reflection_data.friedel_merged) is False + assert ref.model.wavelength == 1.54 + assert bool(ref.model.anomalous_bijvoet) is True + + @pytest.mark.integration + def test_anomalous_without_wavelength_raises(self, small_pair): + """``anomalous=True`` needs the wavelength its f'' term is computed at.""" + from torchref.refinement.lbfgs_refinement import LBFGSRefinement + + with pytest.raises(ValueError, match="wavelength"): + LBFGSRefinement( + data_file=small_pair["mtz"], + pdb=small_pair["pdb"], + device=torch.device("cpu"), + verbose=0, + anomalous=True, + ) diff --git a/torchref/cli/_common.py b/torchref/cli/_common.py index dc749106..b6c115ed 100644 --- a/torchref/cli/_common.py +++ b/torchref/cli/_common.py @@ -194,15 +194,15 @@ def add_adp_mode_arg(parser: argparse.ArgumentParser) -> None: def add_wavelength_arg(parser: argparse.ArgumentParser) -> None: - """Add ``--wavelength`` argument (Angstroms; 0 disables anomalous).""" + """Add ``--wavelength`` argument (Angstroms; unset or 0 means no anomalous).""" parser.add_argument( "--wavelength", type=float, - default=1.0, - help="X-ray wavelength in Angstroms, used for anomalous (f'/f'') " - "scattering. Set to 0 to disable anomalous refinement entirely, which " - "also forces a Friedel-merged read of the data (no F(+)/F(-) Bijvoet " - "pairs). Default 1.0.", + default=None, + help="X-ray wavelength of the data in Angstroms. Given, the model includes " + "anomalous (f'/f'') scattering and F(+)/F(-) columns are read as Bijvoet " + "pairs. Default: no anomalous scattering and a Friedel-merged read, which " + "0 also selects.", ) diff --git a/torchref/cli/refine.py b/torchref/cli/refine.py index d93a3c41..01c52852 100644 --- a/torchref/cli/refine.py +++ b/torchref/cli/refine.py @@ -294,8 +294,8 @@ def main(): f"{args.anisotropic_selection or 'not resname HOH and not element H'})" ) print(adp_line) - if args.wavelength == 0: - print("Anomalous: off (wavelength 0 -> Friedel-merged read)") + if not args.wavelength: + print("Anomalous: off (no --wavelength; Friedel-merged read)") else: print(f"Wavelength: {args.wavelength:.4g} A") if manual_weights: diff --git a/torchref/experimental/ensemble/ensemble_model.py b/torchref/experimental/ensemble/ensemble_model.py index b59bbf97..949affe1 100644 --- a/torchref/experimental/ensemble/ensemble_model.py +++ b/torchref/experimental/ensemble/ensemble_model.py @@ -344,7 +344,7 @@ def __init__( hydrogens: str = "keep", max_res: float = 1.0, gridsize: Optional[Tuple[int, int, int]] = None, - wavelength: float = 1.0, + wavelength: Optional[float] = None, anomalous_threshold: float = 0.5, apply_bijvoet: bool = False, cif_path=None, diff --git a/torchref/model/model_ft.py b/torchref/model/model_ft.py index af63aaa6..40f0df56 100644 --- a/torchref/model/model_ft.py +++ b/torchref/model/model_ft.py @@ -58,9 +58,10 @@ class ModelFT(CachedForwardMixin, Model): gridsize : tuple of int, optional Explicit grid size (nx, ny, nz). If None, computed from cell and max_res. wavelength : float or None, optional - X-ray wavelength in Angstroms for anomalous scattering correction. - Default is 1.0 (standard synchrotron, ~12.4 keV). Set to None to - disable anomalous corrections entirely. + X-ray wavelength of the data in Angstroms, which sets the anomalous f' + and f''. Default None: no anomalous scattering, f0 only. f' and f'' are + strongly wavelength-dependent near an absorption edge, so pass the + wavelength the data were collected at, not a nominal one. anomalous_threshold : float, optional Significance threshold for anomalous scattering in electrons. Atoms with |f'| > threshold or |f''| > threshold will have @@ -72,7 +73,7 @@ class ModelFT(CachedForwardMixin, Model): Attributes ---------- - max_res, wavelength, anomalous_threshold : float + max_res, wavelength, anomalous_threshold : float or None The constructor arguments above, readable back as attributes. gridsize : torch.Tensor or None Grid dimensions ``(nx, ny, nz)``, derived by the ``SfFFT`` submodule from @@ -91,7 +92,7 @@ def __init__( *args, max_res=1.0, gridsize: Optional[Tuple[int, int, int]] = None, - wavelength: Optional[float] = 1.0, + wavelength: Optional[float] = None, anomalous_threshold: float = 0.5, apply_bijvoet: bool = False, **kwargs, @@ -111,9 +112,8 @@ def __init__( gridsize : tuple of int, optional Explicit grid size tuple (nx, ny, nz). If None, computed automatically. wavelength : float or None, optional - X-ray wavelength in Angstroms for anomalous scattering correction. - Default is 1.0 (standard synchrotron, ~12.4 keV). Set to None to - disable anomalous corrections entirely. + X-ray wavelength of the data in Angstroms, which sets the anomalous + f' and f''. Default None: no anomalous scattering, f0 only. anomalous_threshold : float, optional Significance threshold for anomalous scattering in electrons. Atoms with |f'| > threshold or |f''| > threshold will have @@ -775,7 +775,7 @@ def _pop_subclass_state(cls, state_dict: dict) -> dict: return { "max_res": state_dict.pop("max_res", 1.0), "gridsize": state_dict.pop("explicit_gridsize", None), - "wavelength": state_dict.pop("wavelength", 1.0), + "wavelength": state_dict.pop("wavelength", None), "anomalous_threshold": state_dict.pop("anomalous_threshold", 0.5), } diff --git a/torchref/refinement/base_refinement.py b/torchref/refinement/base_refinement.py index b3311fc1..eb9d8308 100644 --- a/torchref/refinement/base_refinement.py +++ b/torchref/refinement/base_refinement.py @@ -124,7 +124,7 @@ def __init__( nbins: int = 10, n_iso_coeff: int = 6, column_names: Optional[Dict[str, str]] = None, - wavelength: Optional[float] = 1.0, + wavelength: Optional[float] = None, anomalous_threshold: float = 0.5, french_wilson: bool = True, anomalous: Optional[bool] = None, @@ -168,18 +168,20 @@ def __init__( column_names : dict, optional Mapping of logical column roles to MTZ column labels. wavelength : float, optional - X-ray wavelength in Angstroms for the anomalous (f'/f'') correction. - ``0`` means "no anomalous refinement": it disables the correction and - forces a Friedel-merged read, **overriding** ``anomalous`` to False. + X-ray wavelength of the data in Angstroms. Given, the model includes the + anomalous f'/f'' and the data may be read as Bijvoet pairs (see + ``anomalous``). Default None, as is ``0``: no anomalous scattering and a + Friedel-merged read. anomalous_threshold : float, optional Threshold controlling anomalous data handling. Default 0.5. french_wilson : bool, optional Derive amplitudes from intensities via French-Wilson. Set False to use existing ``F``/``SIGF`` columns when the MTZ also carries intensities. anomalous : bool, optional - Anomalous (Bijvoet) load preference. None auto-detects ``F(+)/F(-)`` - (or ``I(+)/I(-)``) and loads Friedel pairs when present, enabling the - model's f'' term; True forces it, False forces a merged load. + Anomalous (Bijvoet) load preference. None (default) loads Friedel pairs + when a ``wavelength`` is given and the file has ``F(+)/F(-)`` (or + ``I(+)/I(-)``), enabling the model's f'' term; True forces it; False + forces a merged load. adp_mode : str, optional ADP parametrization: ``"isotropic"`` (default) refines a per-atom B-factor, ``"anisotropic"`` a 6-component U tensor for the atoms @@ -216,6 +218,11 @@ def __init__( hydrogens_in_xray : bool, optional Whether hydrogens contribute to the structure factors. Default True. They take part in the restraints either way. + + Raises + ------ + ValueError + If ``anomalous=True`` is given without a ``wavelength``. """ super().__init__() # Refinement constructs its own submodules from file paths, so @@ -230,14 +237,9 @@ def __init__( self.nbins = nbins self.n_iso_coeff = n_iso_coeff self.lr = 1e-3 - # Wavelength drives f'/f'' anomalous scattering corrections in ModelFT. - # Default 1.0 preserves prior behavior; set to the experimental wavelength - # for anomalous (Bijvoet) refinement, or None to disable entirely. self.wavelength = wavelength self.anomalous_threshold = anomalous_threshold self.french_wilson = french_wilson - # Anomalous (Bijvoet) load preference: None auto-detects and prefers - # anomalous data when present; True forces it; False forces a merged load. self.anomalous = anomalous # ADP parametrization: 'isotropic' (default) refines per-atom B; # 'anisotropic' refines a 6-component U for atoms matched by @@ -258,10 +260,16 @@ def __init__( self.xray_mode = xray_mode self.sigma_a_max = sigma_a_max self.shrink = shrink - # A wavelength of 0 means "no anomalous refinement": disable the f'/f'' - # correction (model wavelength None) and force a Friedel-merged read so - # F(+)/F(-) are not loaded as Bijvoet pairs. - if self.wavelength is not None and float(self.wavelength) == 0.0: + # Without a wavelength there is no f'' to tell Friedel mates apart, so a + # Bijvoet read would only split each acentric reflection into two + # observations of one modelled amplitude. + if self.wavelength is None or float(self.wavelength) == 0.0: + if self.anomalous: + raise ValueError( + "anomalous=True reads F(+)/F(-) as Bijvoet pairs for the f'' " + "term, which needs the wavelength the data were collected at; " + "pass wavelength=..." + ) self.wavelength = None self.anomalous = False From 2cf9a632c28f4442a04a40cf086f692f109dd11f Mon Sep 17 00:00:00 2001 From: Claude Date: Fri, 2 Oct 2026 11:12:38 +0000 Subject: [PATCH 249/250] Put the MPS-failing tests' tensors on the device their objects live on The Accelerator job runs the suite with TORCHREF_DEVICE=mps, where Cell, SpaceGroup and ModelFT default to MPS while a bare torch.tensor() lands on CPU. The new symmetry-image, non-bonded and VDW-contact tests, the float64 Cell reference check and the anomalous cache test mixed the two. - test_nonbonded_kernel / test_vdw_symmetry_contacts: build the cell and space group on CPU, matching their host-side builder and gemmi comparison (the Triton test already moves the host copy to CUDA). - test_symmetry_images: create every input on the cell's device, so the image kernel itself is exercised on the accelerator. - test_cell_geometry: pin the float64 reference Cell to CPU (MPS has no float64). - test_anomalous::test_cache_invalidation: hkl on model.device, as the neighbouring tests already do. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01DMJru8rtHiyy9tD2rEHth9 --- tests/unit/base/test_cell_geometry.py | 3 +- tests/unit/base/test_nonbonded_kernel.py | 7 +++-- tests/unit/base/test_symmetry_images.py | 31 +++++++++++-------- tests/unit/scattering/test_anomalous.py | 4 ++- .../topology/test_vdw_symmetry_contacts.py | 5 +-- 5 files changed, 31 insertions(+), 19 deletions(-) diff --git a/tests/unit/base/test_cell_geometry.py b/tests/unit/base/test_cell_geometry.py index fd6a8ff8..bd8fa9d6 100644 --- a/tests/unit/base/test_cell_geometry.py +++ b/tests/unit/base/test_cell_geometry.py @@ -59,7 +59,8 @@ def test_direct_and_reciprocal_bases_match_gemmi(params): @pytest.mark.parametrize("params", CELLS.values(), ids=CELLS.keys()) def test_cell_object_is_consistent_with_gemmi(params): _, frac_ref = _gemmi_matrices(params) - cell = Cell(list(params), dtype=torch.float64) # dtype-ok: reference precision + # dtype-ok: reference precision, on CPU because MPS has no float64. + cell = Cell(list(params), dtype=torch.float64, device="cpu") np.testing.assert_allclose( cell.reciprocal_basis_matrix.numpy(), frac_ref, atol=1e-12 ) diff --git a/tests/unit/base/test_nonbonded_kernel.py b/tests/unit/base/test_nonbonded_kernel.py index dded8431..e203b545 100644 --- a/tests/unit/base/test_nonbonded_kernel.py +++ b/tests/unit/base/test_nonbonded_kernel.py @@ -39,12 +39,15 @@ def _image_pairs(path): """Image pairs of one deposited model from builder steps 1-4, with radius sums.""" st = gemmi.read_structure(str(path)) atoms = [a for ch in st[0] for r in ch for a in r] + # On CPU whatever the configured device: the distances are compared with gemmi's + # on the host, and the Triton test moves this host copy to CUDA itself. + cpu = torch.device("cpu") xyz = torch.tensor( [[a.pos.x, a.pos.y, a.pos.z] for a in atoms], dtype=get_float_dtype() ) c = st.cell - cell = Cell([c.a, c.b, c.c, c.alpha, c.beta, c.gamma]) - sg = SpaceGroup(st.spacegroup_hm) + cell = Cell([c.a, c.b, c.c, c.alpha, c.beta, c.gamma], device=cpu) + sg = SpaceGroup(st.spacegroup_hm, device=cpu) ops, offsets = nb.prefilter_symop_offsets(cell, sg, xyz, CUTOFF) identity = (~is_symmetry_image(ops, offsets)).nonzero()[0].item() grid_dims = torch.clamp( diff --git a/tests/unit/base/test_symmetry_images.py b/tests/unit/base/test_symmetry_images.py index c9288e9e..cfc9ffe4 100644 --- a/tests/unit/base/test_symmetry_images.py +++ b/tests/unit/base/test_symmetry_images.py @@ -26,12 +26,13 @@ def crystal(pdb_dir): """1DAW (C 1 2 1, monoclinic, so B is not diagonal): coordinates, cell, group.""" st = gemmi.read_structure(str(pdb_dir / "1DAW.pdb")) + c = st.cell + cell = Cell([c.a, c.b, c.c, c.alpha, c.beta, c.gamma]) xyz = torch.tensor( [[a.pos.x, a.pos.y, a.pos.z] for ch in st[0] for r in ch for a in r][:200], dtype=get_float_dtype(), + device=cell.device, ) - c = st.cell - cell = Cell([c.a, c.b, c.c, c.alpha, c.beta, c.gamma]) sg = SpaceGroup(st.spacegroup_hm) tables = ( sg.matrices, @@ -42,10 +43,10 @@ def crystal(pdb_dir): return xyz, cell, sg, tables -def _per_point(n_points, op, offset): - ops = torch.full((n_points,), op, dtype=get_int_dtype()) - offsets = torch.tensor(offset, dtype=get_int_dtype()).expand(n_points, 3) - return ops, offsets +def _per_point(xyz, op, offset): + ops = torch.full((len(xyz),), op, dtype=get_int_dtype(), device=xyz.device) + offsets = torch.tensor(offset, dtype=get_int_dtype(), device=xyz.device) + return ops, offsets.expand(len(xyz), 3) def test_matches_the_space_group_expansion(crystal): @@ -53,16 +54,16 @@ def test_matches_the_space_group_expansion(crystal): expanded = sg.expand_positions(cell.cartesian_to_fractional(xyz)) for op in range(sg.n_ops): for offset in ([0, 0, 0], [1, 0, -2], [-3, 2, 1]): - ops, offsets = _per_point(len(xyz), op, offset) + ops, offsets = _per_point(xyz, op, offset) got = symmetry_image_positions(xyz, ops, offsets, *tables) - shift = torch.tensor(offset, dtype=xyz.dtype) + shift = torch.tensor(offset, dtype=xyz.dtype, device=xyz.device) want = cell.fractional_to_cartesian(expanded[op] + shift) torch.testing.assert_close(got, want, atol=_ATOL, rtol=0) def test_a_pure_lattice_translation_is_an_image(crystal): xyz, cell, _, tables = crystal - ops, offsets = _per_point(len(xyz), 0, [2, -1, 1]) + ops, offsets = _per_point(xyz, 0, [2, -1, 1]) got = symmetry_image_positions(xyz, ops, offsets, *tables) shift = cell.fractional_to_cartesian(offsets[0].to(xyz.dtype)) torch.testing.assert_close(got, xyz + shift, atol=_ATOL, rtol=0) @@ -71,7 +72,7 @@ def test_a_pure_lattice_translation_is_an_image(crystal): def test_the_identity_returns_the_point(crystal): xyz, _, _, tables = crystal - ops, offsets = _per_point(len(xyz), 0, [0, 0, 0]) + ops, offsets = _per_point(xyz, 0, [0, 0, 0]) got = symmetry_image_positions(xyz, ops, offsets, *tables) torch.testing.assert_close(got, xyz, atol=_ATOL, rtol=0) assert not bool(is_symmetry_image(ops, offsets).any()) @@ -90,9 +91,13 @@ def test_is_symmetry_image_looks_at_operation_and_offset(): def test_broadcasting_matches_one_entry_per_point(crystal): xyz, _, sg, tables = crystal - ops = torch.tensor([0, 1, 0, sg.n_ops - 1], dtype=get_int_dtype()) + ops = torch.tensor( + [0, 1, 0, sg.n_ops - 1], dtype=get_int_dtype(), device=xyz.device + ) offsets = torch.tensor( - [[0, 0, 0], [0, 0, 1], [-1, 2, 0], [1, 1, 1]], dtype=get_int_dtype() + [[0, 0, 0], [0, 0, 1], [-1, 2, 0], [1, 1, 1]], + dtype=get_int_dtype(), + device=xyz.device, ) grid = symmetry_image_positions(xyz[:, None, :], ops, offsets, *tables) assert grid.shape == (len(xyz), len(ops), 3) @@ -108,7 +113,7 @@ def test_broadcasting_matches_one_entry_per_point(crystal): def test_gradient_is_the_cartesian_rotation(crystal): xyz, cell, sg, tables = crystal op = 1 - ops, offsets = _per_point(1, op, [1, 0, -1]) + ops, offsets = _per_point(xyz[:1], op, [1, 0, -1]) def image(x): return symmetry_image_positions(x[None, :], ops, offsets, *tables)[0] diff --git a/tests/unit/scattering/test_anomalous.py b/tests/unit/scattering/test_anomalous.py index 297c498a..9e616ab6 100644 --- a/tests/unit/scattering/test_anomalous.py +++ b/tests/unit/scattering/test_anomalous.py @@ -391,7 +391,9 @@ def test_cache_invalidation(self, test_pdb_file): model = ModelFT(wavelength=1.0, verbose=0) model.load_pdb(test_pdb_file) - hkl = torch.tensor([[1, 2, 3], [2, 1, 0]], dtype=torch.int32) + hkl = torch.tensor( + [[1, 2, 3], [2, 1, 0]], dtype=torch.int32, device=model.device + ) first = model(hkl).detach().clone() assert model._get_anomalous_cache() is not None diff --git a/tests/unit/topology/test_vdw_symmetry_contacts.py b/tests/unit/topology/test_vdw_symmetry_contacts.py index b53ab236..8b5ae60c 100644 --- a/tests/unit/topology/test_vdw_symmetry_contacts.py +++ b/tests/unit/topology/test_vdw_symmetry_contacts.py @@ -45,13 +45,14 @@ def _read(path): dtype=get_float_dtype(), ) c = st.cell - cell = Cell([c.a, c.b, c.c, c.alpha, c.beta, c.gamma]) + cpu = torch.device("cpu") + cell = Cell([c.a, c.b, c.c, c.alpha, c.beta, c.gamma], device=cpu) return ( st, {key: k for k, key in enumerate(atoms)}, xyz, cell, - SpaceGroup(st.spacegroup_hm), + SpaceGroup(st.spacegroup_hm, device=cpu), ) From cbedbbdf92755e34c66520f1b4a7979f41897943 Mon Sep 17 00:00:00 2001 From: HatPdotS <89108105+HatPdotS@users.noreply.github.com> Date: Fri, 2 Oct 2026 13:19:54 +0200 Subject: [PATCH 250/250] Bound generate_possible_hkl's index box by the cell lengths h = s . a, so the sphere |s| <= 1/d_min needs |h| <= a/d_min per axis. The box from the reciprocal axis lengths, ceil(1/(|a*| d_min)), is smaller than that in oblique cells and missed reflections near d_min in hexagonal and rhombohedral cells. A new unit test checks the full sphere and Friedel closure in five cell systems against a brute-force box. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01H1SNmK471JmLYUagiN9Rai --- docs/changelog.rst | 1 + tests/unit/base/test_generate_possible_hkl.py | 58 +++++++++++++++++++ torchref/base/reciprocal/hkl.py | 12 ++-- 3 files changed, 63 insertions(+), 8 deletions(-) create mode 100644 tests/unit/base/test_generate_possible_hkl.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 8aa719b8..fc4f3ec4 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- ``generate_possible_hkl``, and with it ``SpaceGroup.complete_hkl``, returns the full resolution sphere in every cell. Its index box was sized by the reciprocal axis lengths, ``ceil(1 / (|a*| d_min))``, which is smaller than ``a / d_min`` when the other two axes are not perpendicular to ``a``, so hexagonal, rhombohedral and some triclinic cells lost reflections near ``d_min``. The box is now bounded by the cell lengths, ``ceil(a / d_min)`` per axis - Anomalous scattering is modelled only when the data's wavelength is given (``--wavelength``, ``wavelength=``). The 1.0 Å default, which gave every model f'/f'' at an arbitrary wavelength and read F(+)/F(-) as Bijvoet pairs, is gone: without a wavelength there is no f'/f'' and the read is Friedel-merged, as ``0`` already selected, and ``anomalous=True`` raises - Anomalous f'/f'' go through the same density and FFT as f0, so each atom's isotropic or anisotropic temperature factor and the space-group symmetry apply to them. They were summed over the asymmetric unit without temperature factors, which skewed F_calc and Bijvoet differences in every space group (on 6G9X the F_calc error was 12.5 times the FFT path's own, and Bijvoet differences correlated with the true ones at CC 0.34) - The non-bonded pair list includes every symmetry mate in reach wherever the model sits relative to the origin cell. Mates that needed a cell offset beyond ±1 were dropped: all 2,270 of 6G9X's symmetry contacts and 2,838 of 3E98's 4,638. Contact counts now match gemmi and do not change when the model is shifted by a lattice vector diff --git a/tests/unit/base/test_generate_possible_hkl.py b/tests/unit/base/test_generate_possible_hkl.py new file mode 100644 index 00000000..6caf0523 --- /dev/null +++ b/tests/unit/base/test_generate_possible_hkl.py @@ -0,0 +1,58 @@ +"""generate_possible_hkl returns the full resolution sphere in every cell. + +The reference is a brute-force box twice as wide as the bound |h| <= a / d_min +(likewise k by b, l by c), filtered by d-spacing. In oblique cells 1 / (|a*| d_min) +is smaller than a / d_min, so a box sized by the reciprocal axis lengths misses +part of the sphere; the hexagonal and rhombohedral cells here exercise that. +""" + +import math + +import pytest +import torch + +from torchref.base.reciprocal.hkl import generate_possible_hkl, get_d_spacing + +pytestmark = pytest.mark.unit + +CELLS = { + "orthorhombic": (40.0, 50.0, 60.0, 90.0, 90.0, 90.0), + "monoclinic": (50.0, 60.0, 70.0, 90.0, 105.0, 90.0), + "hexagonal": (80.0, 80.0, 120.0, 90.0, 90.0, 120.0), + "rhombohedral_acute": (40.0, 40.0, 40.0, 60.0, 60.0, 60.0), + "triclinic_skewed": (30.0, 40.0, 50.0, 70.0, 80.0, 60.0), +} +D_MIN = 5.0 + + +def _keys(hkl): + return set(map(tuple, hkl.tolist())) + + +def _sphere(cell, d_min): + n = 2 * math.ceil(max(cell[:3]) / d_min) + r = torch.arange(-n, n + 1) + box = torch.cartesian_prod(r, r, r) + box = box[(box != 0).any(dim=1)] + inside = box[get_d_spacing(box, cell) >= d_min] + assert inside.abs().max() < n, "reference box too small" + return inside + + +@pytest.mark.parametrize("name", CELLS) +def test_full_sphere(name): + cell = torch.tensor(CELLS[name], dtype=torch.float64) + hkl = generate_possible_hkl(cell, D_MIN) + + assert len(_keys(hkl)) == len(hkl) + assert _keys(hkl) == _keys(_sphere(cell, D_MIN)) + + +@pytest.mark.parametrize("name", CELLS) +def test_friedel_closed_without_origin(name): + cell = torch.tensor(CELLS[name], dtype=torch.float64) + hkl = generate_possible_hkl(cell, D_MIN) + + assert (0, 0, 0) not in _keys(hkl) + assert _keys(hkl) == _keys(-hkl) + assert bool((get_d_spacing(hkl, cell) >= D_MIN).all()) diff --git a/torchref/base/reciprocal/hkl.py b/torchref/base/reciprocal/hkl.py index 2d08e1da..7d03c076 100644 --- a/torchref/base/reciprocal/hkl.py +++ b/torchref/base/reciprocal/hkl.py @@ -1,5 +1,6 @@ """Miller index generation and d-spacing calculation.""" +import math from typing import Optional import torch @@ -67,16 +68,11 @@ def generate_possible_hkl( cell = cell.to(device) recB = reciprocal_basis_matrix(cell) - a_star = torch.linalg.norm(recB[0]) - b_star = torch.linalg.norm(recB[1]) - c_star = torch.linalg.norm(recB[2]) - # Per-axis bound ceil(s_max / a*) over-covers the sphere; the resolution - # filter at the end trims it back. + # h = s . a, so |s| <= s_max bounds |h| by s_max * |a| (likewise k by b, + # l by c) in any cell; the resolution filter at the end trims the box. s_max = 1.0 / d_min - h_max = int(torch.ceil(s_max / a_star).item()) - k_max = int(torch.ceil(s_max / b_star).item()) - l_max = int(torch.ceil(s_max / c_star).item()) + h_max, k_max, l_max = (math.ceil(s_max * float(length)) for length in cell[:3]) h_range = torch.arange(-h_max, h_max + 1, device=device, dtype=dtypes.int) k_range = torch.arange(-k_max, k_max + 1, device=device, dtype=dtypes.int)